def run(seed: int, steps: int | None, accelerator: str) -> tuple[dict, dict]:
lit.seed_everything(seed, workers=True)
train = list(records(bags=1024, seed=seed + 1))
validate = list(records(bags=192, seed=seed + 2))
test = list(records(bags=384, seed=seed + 3))
model = rf.Model(
name="request",
d_model=64,
n_layers=3,
n_heads=4,
reduction=None,
batch_size=128,
optimizer=lambda module: torch.optim.AdamW(module.parameters(), lr=0.003),
items=rf.Branch(length=LENGTH, overflow="error", n_layers=3, n_heads=4, reduction=None, value=rf.Number),
rank=rf.Category(size=len(RANKS), p_unavailable=0.0),
answer=rf.Number(mask=True, objective="mse"),
)
datamodule = rf.SyntheticDataModule(
model=model,
train=lambda: records(bags=1024, seed=seed + 1),
validate=lambda: records(bags=192, seed=seed + 2),
seed=seed,
)
trainer = lit.Trainer(
accelerator=accelerator,
devices=1,
max_epochs=-1,
max_steps=800 if steps is None else min(steps, 800),
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)
intact_prediction = prediction(model, test)
by_rank = scores(train=train, test=test, predicted=intact_prediction, key="rank")
by_shape = scores(train=train, test=test, predicted=intact_prediction, key="shape")
intact = overall_score(train=train, test=test, predicted=intact_prediction)
permuted_prediction = prediction(model, permute_items(test, seed=seed + 4))
target_scale = float(np.std(np.asarray([row["answer"] for row in test]).astype(np.float64)))
permutation_drift = rmse(intact_prediction, permuted_prediction) / target_scale
cycled_prediction = prediction(model, cycle_rank(test))
cycled = overall_score(train=train, test=test, predicted=cycled_prediction)
hidden_prediction = prediction(model, [{**row, "rank": None} for row in test])
hidden = scores(train=train, test=test, predicted=hidden_prediction, key="rank")
hidden_family_spread = float(np.max(np.ptp(hidden_prediction.reshape(-1, len(RANKS)), axis=1)))
train_bags = validate_request_families(train)
validate_bags = validate_request_families(validate)
test_bags = validate_request_families(test)
calibration = np.asarray(
[
target_scale,
intact["rmse"],
intact["baseline_rmse"],
intact["nrmse"],
cycled["rmse"],
cycled["baseline_rmse"],
cycled["nrmse"],
permutation_drift,
hidden_family_spread,
*(score["baseline_rmse"] for score in by_rank.values()),
*(score["baseline_rmse"] for score in by_shape.values()),
]
)
endpoint_passes = max((by_rank[rank]["nrmse"] for rank in ("minimum", "maximum"))) < 0.3
interior_passes = max((by_rank[rank]["nrmse"] for rank in ("q25", "median", "q75"))) < 0.35
shape_passes = all((by_shape[shape]["nrmse"] < 0.45 for shape in SHAPES))
controls_pass = (
permutation_drift < 0.1
and cycled["nrmse"] >= 0.8
and (cycled["nrmse"] >= intact["nrmse"] + 0.35)
and all((hidden[rank]["nrmse"] >= 0.65 for rank in ("minimum", "maximum")))
and (sum((score["nrmse"] >= 0.65 for score in hidden.values())) >= 3)
)
metrics = {
"intact": intact,
"by_rank": by_rank,
"by_shape": by_shape,
"cycled": cycled,
"hidden": hidden,
"permutation_drift": permutation_drift,
"hidden_family_spread": hidden_family_spread,
"target_scale": target_scale,
"steps": trainer.global_step,
}
checks = {
"Expected number of rank requests": bool(len(train) == 1024 * len(RANKS) and len(test) == 384 * len(RANKS)),
"All ranks represented": bool(set(np.asarray([row["rank"] for row in test])) == set(RANKS)),
"All bag shapes represented": bool(set(np.asarray([row["shape"] for row in test])) == set(SHAPES)),
"Training bags separate from validation and test": bool(train_bags.isdisjoint(validate_bags | test_bags)),
"Validation bags separate from test": bool(validate_bags.isdisjoint(test_bags)),
"All measurements finite": bool(np.isfinite(calibration).all()),
"Target SD above 0.80": bool(target_scale > 0.8),
"Per-rank baseline RMSE above 0.60": bool(min((score["baseline_rmse"] for score in by_rank.values())) > 0.6),
"Per-shape baseline RMSE above 0.80": bool(min((score["baseline_rmse"] for score in by_shape.values())) > 0.8),
"Hidden requests yield identical predictions": bool(hidden_family_spread <= 1e-06),
"Endpoint, interior-rank, shape, and causal-control gates": bool(
endpoint_passes and interior_passes and shape_passes and controls_pass
),
}
return metrics, checks