def run(seed: int, steps: int | None, accelerator: str) -> tuple[dict, dict]:
lit.seed_everything(seed, workers=True)
model = rf.Model(
name="request",
d_model=64,
n_layers=2,
n_heads=4,
reduction=rf.Attention(n_layers=2),
batch_size=128,
optimizer=lambda module: torch.optim.Adam(module.parameters(), lr=0.001),
items=rf.Branch(
length=6,
n_layers=2,
reduction=rf.Attention(n_layers=2),
value=rf.Number,
group=rf.Category(size=len(GROUPS), p_unavailable=0.0),
),
answer=rf.Number(mask=True, objective="mse"),
operation=rf.Category(size=len(OPERATIONS), p_unavailable=0.0),
selected_group=rf.Category(size=len(GROUPS), p_unavailable=0.0),
)
data = rf.SyntheticDataModule(
model=model,
train=partial(records, bags=512, items_per_group=2, seed=seed + 147),
validate=partial(records, bags=64, items_per_group=2, seed=seed + 148),
)
trainer = lit.Trainer(
accelerator=accelerator,
max_steps=800 if steps is None else min(steps, 800),
max_epochs=-1,
logger=False,
enable_progress_bar=False,
enable_model_summary=False,
enable_checkpointing=False,
deterministic=True,
num_sanity_val_steps=0,
)
trainer.fit(model, datamodule=data)
train = list(records(bags=512, items_per_group=2, seed=seed + 147))
test = list(records(bags=128, items_per_group=2, seed=seed + 149))
intact_prediction = predict(model, test)
intact = scores(train, test, intact_prediction, ("selected_group",))
metrics = {f"intact/{cell}/{name}": value for cell, score in intact.items() for name, value in score.items()}
corrupted = scores(train, test, predict(model, corrupt_labels(test, seed + 150)), ("selected_group",))
metrics.update(
{f"corrupted/{cell}/{name}": value for cell, score in corrupted.items() for name, value in score.items()}
)
checks = {f"Group {cell} nRMSE <= 0.55": score["nrmse"] <= 0.55 for cell, score in intact.items()}
checks.update(
{
f"Group {cell} corruption increases nRMSE by >= 0.25": corrupted[cell]["nrmse"] >= score["nrmse"] + 0.25
for cell, score in intact.items()
}
)
return metrics, checks