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=512, seed=seed + 1))
test = list(records(bags=256, 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),
value=rf.Number,
weight=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=512, seed=seed + 1),
validate=lambda: records(bags=128, seed=seed + 2),
seed=seed,
)
trainer = lit.Trainer(
accelerator=accelerator,
devices=1,
max_epochs=-1,
max_steps=1100 if steps is None else min(steps, 1100),
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)
labels = score(train=train, test=test, predicted=prediction(model, rotate_group_labels(test)))
pairing = score(train=train, test=test, predicted=prediction(model, swap_weights_within_groups(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,
"labels": labels,
"pairing": pairing,
"permutation_drift": permutation_drift,
"steps": trainer.global_step,
}
checks = {
"All measurements finite": bool(
all(
(
np.isfinite(value)
for value in (intact["nrmse"], labels["nrmse"], pairing["nrmse"], permutation_drift)
)
)
),
"Raw grouped sum nRMSE below 0.50": bool(intact["nrmse"] < 0.5),
"Rotated labels raise nRMSE by at least 0.20": bool(labels["nrmse"] >= intact["nrmse"] + 0.2),
"Swapped weights raise nRMSE by at least 0.20": bool(pairing["nrmse"] >= intact["nrmse"] + 0.2),
"Permutation drift below 0.10 target SD": bool(permutation_drift < 0.1),
}
return metrics, checks