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(variable_weighted_mean_records(rows=1536, seed=seed + 1))
test = list(variable_weighted_mean_records(rows=768, 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_mean=rf.Number(mask=True, objective="mse"),
)
datamodule = rf.SyntheticDataModule(
model=model,
train=lambda: variable_weighted_mean_records(rows=1536, seed=seed + 1),
validate=lambda: variable_weighted_mean_records(rows=384, 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)
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))
scaled_prediction = prediction(model, scale_weights(test, factor=0.75))
original_short = [row for row in deepcopy(test) if len(row["items"]) <= ITEMS // 2]
duplicated_prediction = prediction(model, duplicate_items(test, maximum_original=ITEMS // 2))
original_short_prediction = prediction(model, original_short)
target_scale = float(np.std(np.asarray([row["weighted_mean"] for row in test], dtype=np.float64)))
permutation_drift = rmse(intact_prediction, permuted_prediction) / target_scale
scaling_drift = rmse(intact_prediction, scaled_prediction) / target_scale
duplication_drift = rmse(original_short_prediction, duplicated_prediction) / target_scale
metrics = {
"intact": intact,
"corrupted": corrupted,
"permutation_drift": permutation_drift,
"scaling_drift": scaling_drift,
"duplication_drift": duplication_drift,
"steps": trainer.global_step,
}
checks = {
"Weighted mean nRMSE below 0.40": bool(intact["nrmse"] < 0.4),
"Shuffled weights raise nRMSE by at least 0.25": bool(corrupted["nrmse"] >= intact["nrmse"] + 0.25),
"Permutation drift below 0.10 target SD": bool(permutation_drift < 0.1),
"Weight scaling drift below 0.10 target SD": bool(scaling_drift < 0.1),
"Duplication drift below 0.10 target SD": bool(duplication_drift < 0.1),
}
return metrics, checks