def run(seed: int, steps: int | None, accelerator: str) -> tuple[dict, dict]:
lit.seed_everything(seed, workers=True)
# Split seeds are independent; rerunning a generator reproduces the same records.
train = list(records(bags=384, seed=seed + 1))
test = list(records(bags=192, seed=seed + 3))
model = rf.Model(
name="request",
d_model=48,
n_layers=2,
n_heads=4,
reduction=rf.Attention(n_outputs=ITEMS, n_layers=2),
batch_size=96,
optimizer=lambda module: torch.optim.Adam(module.parameters(), lr=0.001),
items=rf.Branch(
length=ITEMS,
n_layers=2,
reduction=None,
group=rf.Category(size=len(GROUPS), p_unavailable=0.0),
contribution=rf.Number,
),
selected_group=rf.Category(size=len(GROUPS), p_unavailable=0.0),
answer=rf.Number(mask=True, objective="mse"),
)
datamodule = rf.SyntheticDataModule(
model=model,
train=lambda: records(bags=384, seed=seed + 1),
validate=lambda: records(bags=96, seed=seed + 2),
seed=seed,
)
trainer = lit.Trainer(
accelerator=accelerator,
devices=1,
max_epochs=-1,
max_steps=800 if steps is None else min(steps, 800),
logger=False,
enable_progress_bar=False,
enable_model_summary=False,
enable_checkpointing=False,
deterministic=True,
num_sanity_val_steps=0,
)
trainer.fit(model, datamodule=datamodule)
# Evaluate held-out answers and retain their original labels in corruption controls.
intact_prediction = prediction(model, test)
intact = score(train=train, test=test, predicted=intact_prediction)
corrupted = score(train=train, test=test, predicted=prediction(model, rotate_group_labels(test)))
permuted_prediction = prediction(model, permute_items(test, seed=seed + 4))
target_scale = float(np.std(np.asarray([row["answer"] for row in test], dtype=np.float64)))
permutation_drift = rmse(intact_prediction, permuted_prediction) / target_scale
metrics = {
"intact": intact,
"corrupted": corrupted,
"permutation_drift": permutation_drift,
"steps": trainer.global_step,
}
checks = {
"All measurements finite": bool(
all((np.isfinite(value) for value in (intact["nrmse"], corrupted["nrmse"], permutation_drift)))
),
"Intact calibration nRMSE below 1.60": bool(intact["nrmse"] < 1.6),
"Corrupt calibration nRMSE below 1.80": bool(corrupted["nrmse"] < 1.8),
"Grouped contribution nRMSE below 0.50": bool(intact["nrmse"] < 0.5),
"Rotated labels raise nRMSE by at least 0.20": bool(corrupted["nrmse"] >= intact["nrmse"] + 0.2),
"Permutation drift below 0.10 target SD": bool(permutation_drift < 0.1),
}
return metrics, checks