def run(seed: int, steps: int | None, accelerator: str) -> tuple[dict, dict]:
lit.seed_everything(seed, workers=True)
model = rf.Model(
name="overlap",
d_model=48,
n_layers=3,
n_heads=4,
reduction=None,
batch_size=128,
optimizer=lambda module: torch.optim.AdamW(module.parameters(), lr=0.003),
left=rf.Branch(length=MAX_LENGTH, n_layers=1, reduction=None, entity_id=rf.Hash(n_hashes=4, n_bands=8)),
right=rf.Branch(length=MAX_LENGTH, n_layers=1, reduction=None, entity_id=rf.Hash(n_hashes=4, n_bands=8)),
has_overlap=rf.Boolean(mask=True),
)
# Restart the training RNG independently of parameter initialization.
lit.seed_everything(seed, workers=True)
data = rf.SyntheticDataModule(
model=model,
train=partial(records, rows=3072, seed=seed + 1, namespace="train"),
validate=partial(records, rows=768, seed=seed + 2, namespace="validate"),
seed=seed,
)
trainer = lit.Trainer(
accelerator=accelerator,
max_steps=700 if steps is None else min(steps, 700),
max_epochs=-1,
logger=False,
enable_progress_bar=False,
enable_model_summary=False,
enable_checkpointing=False,
deterministic=True,
check_val_every_n_epoch=5,
)
trainer.fit(model, datamodule=data)
test = list(records(rows=1024, seed=seed + 3, namespace="test"))
target = np.asarray([row["has_overlap"] for row in test], dtype=np.bool_)
intact = probabilities(model, test)
invariant = list(rename_and_permute(test, seed=seed + 4))
broken = list(break_overlaps(test))
invariant_probability = probabilities(model, invariant)
broken_probability = probabilities(model, broken)
metrics = {
"intact_auc": auc(target, intact),
"invariant_auc": auc(target, invariant_probability),
"broken_auc": auc(target, broken_probability),
"invariant_drift": float(np.mean(np.abs(intact - invariant_probability))),
"positive_break_drop": float(np.mean(intact[target] - broken_probability[target])),
}
checks = {
"Labels are exactly balanced": float(target.mean()) == 0.5,
"Labels match actual overlap": bool(np.array_equal(target, overlap_truth(test))),
"Renaming and permutation preserve overlap": bool(np.array_equal(target, overlap_truth(invariant))),
"Intervention removes all overlaps": not bool(overlap_truth(broken).any()),
"Finite metrics": bool(np.isfinite(list(metrics.values())).all()),
"Intact AUC remains in the calibrated chance band": 0.35 <= metrics["intact_auc"] <= 0.65,
"Renamed and permuted AUC remains in the chance band": 0.35 <= metrics["invariant_auc"] <= 0.65,
"Broken AUC remains in the chance band": 0.35 <= metrics["broken_auc"] <= 0.65,
"Invariant probability drift <= 0.15": metrics["invariant_drift"] <= 0.15,
"Absolute positive break response <= 0.10": abs(metrics["positive_break_drop"]) <= 0.10,
}
# These are generator invariants, separate from the model's expected limitation.
for split, rows in (
("train", list(records(rows=3072, seed=seed + 1, namespace="train"))),
("validate", list(records(rows=768, seed=seed + 2, namespace="validate"))),
("test", test),
("invariant", invariant),
):
labels = np.asarray([row["has_overlap"] for row in rows])
checks[f"{split} generator has balanced truthful labels"] = float(labels.mean()) == 0.5 and bool(
np.array_equal(labels, overlap_truth(rows))
)
checks[f"{split} sides have two to five unique members"] = all(
MIN_LENGTH <= len(row[side]) <= MAX_LENGTH
and len(row[side]) == len({item["entity_id"] for item in row[side]})
for row in rows
for side in ("left", "right")
)
checks[f"{split} positive rows have exactly one overlap"] = all(
len({item["entity_id"] for item in row["left"]} & {item["entity_id"] for item in row["right"]})
== int(row["has_overlap"])
for row in rows
)
return metrics, checks