def run(seed: int, steps: int | None, accelerator: str) -> tuple[dict, dict]:
test = list(records(rows=1024, seed=seed + 3, namespace="test"))
broken = list(records(rows=1024, seed=seed + 3, namespace="test", broken_identity=True))
metrics = {}
# Fit both routes from the same seed and data, changing only reduction.
for route, reduction in (("compressed", rf.Attention()), ("preserved", None)):
lit.seed_everything(seed, workers=True)
model = rf.Model(
name="association",
d_model=64,
n_layers=2,
n_heads=4,
reduction=reduction,
batch_size=128,
optimizer=lambda module: torch.optim.AdamW(module.parameters(), lr=3e-3),
source=rf.Branch(
length=PAIR_COUNT,
n_layers=2,
reduction=reduction,
entity_id=rf.Category(size=PAIR_COUNT, p_unavailable=0.0),
value=rf.Number,
),
target=rf.Branch(
length=PAIR_COUNT,
n_layers=2,
reduction=reduction,
entity_id=rf.Category(size=PAIR_COUNT, p_unavailable=0.0),
value=rf.Number(mask=True, objective="mse"),
),
)
data = rf.SyntheticDataModule(
model=model,
train=partial(records, rows=2048, seed=seed + 1, namespace="train"),
validate=partial(records, rows=512, seed=seed + 2, namespace="validate"),
seed=seed,
)
trainer = lit.Trainer(
accelerator=accelerator,
max_steps=600 if steps is None else min(steps, 600),
max_epochs=-1,
logger=False,
enable_progress_bar=False,
enable_model_summary=False,
enable_checkpointing=False,
deterministic=True,
check_val_every_n_epoch=5,
)
trainer.fit(model, datamodule=data)
metrics[f"{route}_nrmse"] = normalized_rmse(model, test)
metrics[f"{route}_broken_nrmse"] = normalized_rmse(model, broken)
compressed = metrics["compressed_nrmse"]
intact = metrics["preserved_nrmse"]
broken_error = metrics["preserved_broken_nrmse"]
checks = {
"Finite errors": bool(np.isfinite(list(metrics.values())).all()),
"Preserved transfer calibration nRMSE <= 1.10": intact <= 1.10,
"Preserved transfer nRMSE <= 0.25": intact <= 0.25,
"Broken identities nRMSE >= 0.80": broken_error >= 0.80,
"Identity corruption increases nRMSE by >= 0.50": broken_error - intact >= 0.50,
"Compressed Category transfer nRMSE <= 0.25": compressed <= 0.25,
}
return metrics, checks