def run(seed: int, steps: int | None, accelerator: str) -> tuple[dict, dict]:
lit.seed_everything(seed, workers=True)
model = rf.Model(
name="request",
d_model=64,
n_layers=2,
n_heads=4,
reduction=rf.Attention(n_layers=2),
batch_size=128,
optimizer=lambda module: torch.optim.Adam(module.parameters(), lr=0.001),
items=rf.Branch(length=8, n_layers=2, reduction=rf.Attention(n_layers=2), value=rf.Number),
answer=rf.Number(mask=True, objective="mse"),
operation=rf.Category(size=len(OPERATIONS), p_unavailable=0.0),
)
data = rf.SyntheticDataModule(
model=model,
train=partial(records, bags=768, length=8, seed=seed + 135),
validate=partial(records, bags=96, length=8, seed=seed + 136),
)
trainer = lit.Trainer(
accelerator=accelerator,
max_steps=800 if steps is None else min(steps, 800),
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(bags=768, length=8, seed=seed + 135))
test = list(records(bags=192, length=8, seed=seed + 137))
intact_prediction = predict(model, test)
intact = scores(train, test, intact_prediction, ("operation",))
metrics = {f"intact/{cell}/{name}": value for cell, score in intact.items() for name, value in score.items()}
permuted_prediction = predict(model, permute_items(test, seed + 138))
permuted = scores(train, test, permuted_prediction, ("operation",))
hidden_prediction = predict(model, [{**row, "operation": None} for row in test])
hidden = scores(train, test, hidden_prediction, ("operation",))
for label, result in (("permuted", permuted), ("hidden", hidden)):
metrics.update(
{f"{label}/{cell}/{name}": value for cell, score in result.items() for name, value in score.items()}
)
scale = float(np.std([row["answer"] for row in test]))
delta = float(np.sqrt(np.mean(np.square(intact_prediction - permuted_prediction)))) / scale
hidden_spread = float(np.max(np.ptp(hidden_prediction.reshape(-1, len(OPERATIONS)), axis=1)))
metrics.update(permutation_delta=delta, hidden_request_spread=hidden_spread)
checks = {f"{cell} intact nRMSE <= 0.20": score["nrmse"] <= 0.20 for cell, score in intact.items()}
checks.update({f"{cell} permuted nRMSE <= 0.25": score["nrmse"] <= 0.25 for cell, score in permuted.items()})
checks.update(
{
"Permutation changes predictions by <= 0.05 target SD": delta <= 0.05,
"Hidden operations yield identical predictions within each bag": hidden_spread <= 1e-6,
"At least three hidden operations have nRMSE >= 0.75": sum(
score["nrmse"] >= 0.75 for score in hidden.values()
)
>= 3,
}
)
return metrics, checks