def run(seed: int, steps: int | None, accelerator: str) -> tuple[dict, dict]:
lit.seed_everything(seed, workers=True)
model = rf.Model(
name="collection",
d_model=48,
n_layers=3,
n_heads=4,
reduction=None,
batch_size=128,
optimizer=lambda module: torch.optim.Adam(module.parameters(), lr=1e-3),
items=rf.Branch(
length=ITEMS,
overflow="error",
n_layers=2,
reduction=None,
value=rf.Number,
group=rf.Category(size=len(GROUPS), p_unavailable=0.0),
deviation=rf.Number(mask=True, objective="mse", n_linear=2),
),
)
data = rf.SyntheticDataModule(
model=model,
train=partial(records, rows=2048, seed=seed + 1),
validate=partial(records, rows=512, seed=seed + 2),
seed=seed,
)
trainer = lit.Trainer(
accelerator=accelerator,
max_steps=900 if steps is None else min(steps, 900),
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(rows=2048, seed=seed + 1))
test = list(records(rows=768, seed=seed + 3))
actual = targets(test)
baseline = rmse(actual, float(targets(train).mean()))
predicted = predict(model, test)
intact_error = rmse(actual, predicted)
intact_nrmse = intact_error / baseline
metrics = {"intact_rmse": intact_error, "baseline_rmse": baseline, "intact_nrmse": intact_nrmse}
checks = {"Baseline RMSE > 0.000001": baseline > 1e-6}
corrupted = corrupt_groups(test)
corrupted_error = rmse(actual, predict(model, corrupted))
corrupted_nrmse = corrupted_error / baseline
oracle = rmse(implied_deviation(corrupted), actual) / baseline
metrics.update(corrupted_rmse=corrupted_error, corrupted_nrmse=corrupted_nrmse, oracle_corruption_nrmse=oracle)
checks["Label corruption changes the oracle by > 0.75 nRMSE"] = oracle > 0.75
translated = translate(test, seed + 4)
translated_error = rmse(targets(translated), predict(model, translated))
translated_nrmse = translated_error / baseline
permuted, order = permute_items(test, seed + 5)
drift = rmse(predict(model, permuted), predicted[order]) / baseline
metrics.update(translated_rmse=translated_error, translated_nrmse=translated_nrmse, permutation_drift=drift)
checks.update(
{
"Training deviations are exactly centered": abs(float(targets(train).mean())) < 1e-12,
"Translation preserves target shape": targets(translated).shape == actual.shape,
"Translation preserves all targets": rmse(targets(translated), actual) == 0.0,
"Permutation preserves the target multiset": sorted(targets(permuted).tolist()) == sorted(actual.tolist()),
"Normalized permutation drift < 0.20": drift < 0.20,
}
)
checks.update(
{
"Grouped deviation nRMSE < 0.40": intact_nrmse < 0.40,
"Group-translated nRMSE < 0.50": translated_nrmse < 0.50,
"Label corruption increases nRMSE by >= 0.25": corrupted_nrmse >= intact_nrmse + 0.25,
}
)
checks["Finite calibration and model metrics"] = bool(np.isfinite(list(metrics.values())).all())
return metrics, checks