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(weighted_records(rows=1024, length=ITEMS, seed=seed + 1))
test = list(weighted_records(rows=512, length=ITEMS, seed=seed + 3))
model = rf.Model(
d_model=32,
n_layers=2,
n_heads=4,
reduction=rf.Attention(n_layers=2),
batch_size=64,
optimizer=lambda module: torch.optim.Adam(module.parameters(), lr=0.001),
items=rf.Branch(
length=ITEMS, n_layers=2, reduction=rf.Attention(n_layers=2), value=rf.Number, weight=rf.Number
),
weighted_sum=rf.Number(mask=True, objective="mse"),
)
datamodule = rf.SyntheticDataModule(
model=model,
train=lambda: weighted_records(rows=1024, length=ITEMS, seed=seed + 1),
validate=lambda: weighted_records(rows=256, length=ITEMS, seed=seed + 2),
seed=seed,
)
trainer = lit.Trainer(
accelerator=accelerator,
devices=1,
max_epochs=-1,
max_steps=900 if steps is None else min(steps, 900),
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_table = permute_weights(test, seed=seed + 4)
corrupted = score(train=train, test=test, predicted=prediction(model, corrupted_table))
permuted_prediction = prediction(model, permute_items(test, seed=seed + 5))
factor = 0.75
scaled_table = scale_weights(test, factor=factor)
scaled_prediction = prediction(model, scaled_table)
scaled = score(train=scale_weights(train, factor=factor), test=scaled_table, predicted=scaled_prediction)
target_scale = float(np.std(np.asarray([row["weighted_sum"] for row in test], dtype=np.float64)))
permutation_drift = rmse(intact_prediction, permuted_prediction) / target_scale
scaling_error = rmse(factor * intact_prediction, scaled_prediction) / (factor * target_scale)
metrics = {
"intact": intact,
"corrupted": corrupted,
"scaled": scaled,
"permutation_drift": permutation_drift,
"scaling_error": scaling_error,
"steps": trainer.global_step,
}
checks = {
"Weighted sum nRMSE below 0.30": bool(intact["nrmse"] < 0.3),
"Shuffled weights raise nRMSE by at least 0.35": bool(corrupted["nrmse"] >= intact["nrmse"] + 0.35),
"Scaled-weight nRMSE below 0.35": bool(scaled["nrmse"] < 0.35),
"Permutation drift below 0.08 target SD": bool(permutation_drift < 0.08),
"Weight scaling error below 0.15 target SD": bool(scaling_error < 0.15),
}
return metrics, checks