def run(seed: int, steps: int | None, accelerator: str) -> tuple[dict, dict]:
lit.seed_everything(seed, workers=True)
budget = 512 if steps is None else min(steps, 512)
train = list(records(rows=2048, seed=seed + 1))
test = list(records(rows=1024, seed=seed + 3))
actual = np.asarray([row["y"] for row in test])
baseline = float(np.sqrt(np.mean((actual - np.mean([row["y"] for row in train])) ** 2)))
scale = float(np.std([row["y"] for row in train]))
model = rf.Model(
d_model=32,
n_layers=1,
n_heads=4,
dropout=0.0,
batch_size=128,
x=rf.Number,
code=rf.Category(size=8, p_unavailable=0.0),
y=rf.Number(mask=True, objective="mse"),
)
model.optimizer = rf.adamw(learning_rate=3e-3, fused=False)
data = rf.SyntheticDataModule(
model=model,
train=partial(records, rows=2048, seed=seed + 1),
validate=partial(records, rows=512, seed=seed + 2),
seed=seed,
)
trainer = lit.Trainer(
accelerator=accelerator,
devices=1,
max_epochs=-1,
max_steps=budget,
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)
model.eval()
reference = prediction(model, test)
state = deepcopy(model.state_dict())
vocabulary = model.nodes["record/code"].embedder.vocab.snapshot()
root = rf.where("address") == "record"
source = rf.where("address") == "record/x"
measured = errors(actual, reference, baseline)
edits = []
checks = {"Source learns the relationship below 0.25 nRMSE": measured["nrmse"] < 0.25}
def observe(label: str, *, key: str = "x") -> None:
result = prediction(model, test, key=key)
current = model.state_dict()
missing = [name for name, value in state.items() if name not in current or not equal(value, current[name])]
preserved = bool(np.allclose(reference, result, rtol=1e-5, atol=1e-6 * scale))
same_vocabulary = model.nodes["record/code"].embedder.vocab.snapshot() == vocabulary
edits.append(
{
"edit": label,
"max_prediction_drift": float(np.max(np.abs(reference - result))),
"nrmse": errors(actual, result, baseline)["nrmse"],
"changed_or_missing_state_entries": missing,
"vocabulary_preserved": same_vocabulary,
}
)
checks[f"{label}: predictions survive"] = preserved
checks[f"{label}: all original state and vocabulary survive"] = not missing and same_vocabulary
for cycle in range(3):
model.update(source, description=f"Equivalent input, cycle {cycle + 1}")
observe(f"cycle {cycle + 1} metadata")
model.extend(root, unused=rf.Number(active=False))
observe(f"cycle {cycle + 1} inactive extension")
model.delete(rf.where("address") == "record/unused")
observe(f"cycle {cycle + 1} inactive deletion")
model.update(source, query="renamed_x")
observe("equivalent source rebind", key="renamed_x")
model.update(source, query=None)
observe("source rebind restoration")
rejected = False
try:
model.extend(root, x=rf.Number)
except ValueError:
rejected = True
checks["Duplicate-name edit is rejected"] = rejected
observe("rejected duplicate extension")
with TemporaryDirectory(prefix="relflow-mutation-") as directory:
path = Path(directory) / "edited.ckpt"
model.save(path)
loaded = rf.Model.load(path).to(model.device).eval()
restored = prediction(loaded, test)
checks["Edited checkpoint preserves schema"] = loaded.schema.model_dump() == model.schema.model_dump()
checks["Edited checkpoint preserves learned state"] = equal(model.state_dict(), loaded.state_dict())
checks["Edited checkpoint preserves predictions"] = bool(
np.allclose(reference, restored, rtol=1e-5, atol=1e-6 * scale)
)
lit.seed_everything(seed + 100, workers=True)
loaded.reset(root, descendants=True)
reset = errors(actual, prediction(loaded.eval(), test), baseline)
checks["Complete reset loses the learned relationship"] = (
reset["nrmse"] > 0.8 and reset["nrmse"] > 3 * measured["nrmse"]
)
return {
"source": measured,
"source_prerequisite_met": measured["nrmse"] < 0.25,
"downstream_interpretable": measured["nrmse"] < 0.25,
"source_steps": trainer.global_step,
"edits": edits,
"controls": {"complete_reset": reset},
"source_vocabulary": vocabulary,
"test_rows": len(test),
"optimizer_policy": "Fresh AdamW for source training; no post-edit fitting",
}, checks