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, (2017, 2018, 2019), parity=1)
validate = partial(records, (2021,), parity=1)
test = list(records((2023, 2025), parity=0))
model = fit(
dateparts=("day_of_year",),
classes=12,
train=train,
validate=validate,
epochs=70,
seed=seed,
steps=steps,
accelerator=accelerator,
)
periodic_accuracy = accuracy(model, lambda: iter(test), accelerator)
order = np.random.default_rng(seed + 1).permutation(len(test))
permuted = [{**row, "observed_at": test[index]["observed_at"]} for row, index in zip(test, order, strict=True)]
permuted_accuracy = accuracy(model, lambda: iter(permuted), accelerator)
gap = periodic_accuracy - permuted_accuracy
return {
"periodic_accuracy": periodic_accuracy,
"permuted_accuracy": permuted_accuracy,
"accuracy_gap": gap,
}, {
"Unseen-date month accuracy reaches 0.90": periodic_accuracy >= 0.90,
"Timestamp permutation accuracy is at most 0.20": permuted_accuracy <= 0.20,
"Original dates exceed permutation by at least 0.65 accuracy": gap >= 0.65,
}