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, years=tuple(range(1901, 1951)), rows=512)
validate = partial(records, years=tuple(range(1951, 2001)), rows=128)
test = partial(records, years=tuple(range(2001, 2051)), rows=256)
day_only = fit(
dateparts=("day_of_year",),
classes=2,
train=train,
validate=validate,
epochs=10,
seed=seed,
steps=steps,
accelerator=accelerator,
)
ambiguous_accuracy = accuracy(day_only, test, accelerator)
identified = fit(
dateparts=("day_of_year", "week_of_month"),
classes=2,
train=train,
validate=validate,
epochs=20,
seed=seed,
steps=steps,
accelerator=accelerator,
)
identified_accuracy = accuracy(identified, test, accelerator)
gap = identified_accuracy - ambiguous_accuracy
return {
"day_only_accuracy": ambiguous_accuracy,
"with_week_of_month_accuracy": identified_accuracy,
"accuracy_gap": gap,
}, {
"Day-of-year accuracy remains between 0.49 and 0.51": 0.49 <= ambiguous_accuracy <= 0.51,
"Visible week-of-month accuracy reaches 0.95": identified_accuracy >= 0.95,
"Visible week-of-month improves accuracy by at least 0.40": gap >= 0.40,
}