def fit(
*,
dateparts: Sequence[str],
classes: int,
train: Callable[[], Iterator[dict]],
validate: Callable[[], Iterator[dict]],
epochs: int,
seed: int,
steps: int | None,
accelerator: str,
) -> rf.Model:
"""Train on the selected calendar coordinates with the remaining schema fixed."""
lit.seed_everything(seed, workers=True)
model = rf.Model(
name="calendar",
d_model=64,
n_layers=3,
n_heads=4,
batch_size=128,
observed_at=rf.DateParts(dateparts=list(dateparts)),
target=rf.Category(mask=True, size=classes, p_unavailable=0.0),
)
model.optimizer = lambda module: torch.optim.AdamW(module.parameters(), lr=3e-3)
data = rf.SyntheticDataModule(model=model, train=train, validate=validate, seed=seed)
trainer = lit.Trainer(
accelerator=accelerator,
max_epochs=epochs,
max_steps=steps if steps is not None else -1,
logger=False,
enable_progress_bar=False,
enable_model_summary=False,
enable_checkpointing=False,
deterministic=True,
num_sanity_val_steps=0,
)
trainer.fit(model=model, datamodule=data)
return model
def accuracy(model: rf.Model, records: Callable[[], Iterator[dict]], accelerator: str) -> float:
"""Evaluate held-out labels without updating the Category vocabulary."""
data = rf.SyntheticDataModule(model=model, test=records)
trainer = lit.Trainer(
accelerator=accelerator,
logger=False,
enable_progress_bar=False,
enable_model_summary=False,
enable_checkpointing=False,
deterministic=True,
)
metrics = trainer.test(model=model, datamodule=data, verbose=False)[0]
return float(metrics["calendar.target/test.accuracy.content"])
def run(seed: int, steps: int | None, accelerator: str) -> tuple[dict, dict]:
train = partial(records, start=date(2017, 1, 2), weeks=52, rows=1024, seed=seed + 1)
validate = partial(records, start=date(2021, 1, 4), weeks=26, rows=512, seed=seed + 2)
test = partial(records, start=date(2025, 1, 6), weeks=52, rows=1024, seed=seed + 3)
day_only = fit(
dateparts=("day_of_week",),
classes=2,
train=train,
validate=validate,
epochs=20,
seed=seed,
steps=steps,
accelerator=accelerator,
)
day_accuracy = accuracy(day_only, test, accelerator)
hour_only = fit(
dateparts=("hour_of_day",),
classes=2,
train=train,
validate=validate,
epochs=20,
seed=seed,
steps=steps,
accelerator=accelerator,
)
hour_accuracy = accuracy(hour_only, test, accelerator)
composed = fit(
dateparts=("day_of_week", "hour_of_day"),
classes=2,
train=train,
validate=validate,
epochs=20,
seed=seed,
steps=steps,
accelerator=accelerator,
)
composed_accuracy = accuracy(composed, test, accelerator)
gap = composed_accuracy - max(day_accuracy, hour_accuracy)
return {
"day_only_accuracy": day_accuracy,
"hour_only_accuracy": hour_accuracy,
"composed_accuracy": composed_accuracy,
"accuracy_gap": gap,
}, {
"Weekday-only accuracy is at most 0.82": day_accuracy <= 0.82,
"Hour-only accuracy is at most 0.82": hour_accuracy <= 0.82,
"Combined coordinates reach 0.95 accuracy": composed_accuracy >= 0.95,
"Composition improves over either coordinate by at least 0.15": gap >= 0.15,
}