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=2048, 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)))
tolerance = 1e-6 * float(np.std([row["y"] for row in train]))
oracle = errors(actual, np.asarray([row["a"] for row in test]), baseline)
rng = np.random.default_rng(seed + 4)
order = rng.permutation(len(test))
shuffled = [{**row, "b": test[index]["b"]} for row, index in zip(test, order, strict=True)]
changed = [{**row, "b": -row["b"]} for row in test]
model = rf.Model(
d_model=32,
n_layers=1,
n_heads=4,
dropout=0.0,
batch_size=128,
a=rf.Number,
b=rf.Number,
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)
source = errors(actual, reference, baseline)
corruption = errors(actual, prediction(model, shuffled), baseline)
zero_filled = errors(actual, prediction(model, [{**row, "b": 0.0} for row in test]), baseline)
schema = deepcopy(model.schema.model_dump())
state = deepcopy(model.state_dict())
selected = rf.where("address") == "record/b"
learned = source["nrmse"] < 0.25
checks = {
"Source learns both-input relationship below 0.25 nRMSE": learned,
"Shuffling b destroys useful signal": corruption["nrmse"] > source["nrmse"] + 0.35
and corruption["nrmse"] >= 0.9 * oracle["nrmse"],
"Source hidden target values cannot affect predictions": bool(
np.allclose(reference, prediction(model, test, corrupt=True), rtol=1e-5, atol=tolerance)
),
}
edits = []
def observe(label: str, *, inactive: bool) -> np.ndarray:
"""Evaluate the appropriate information-loss or restoration invariant."""
values = prediction(model, test)
measured = errors(actual, values, baseline)
same_state = equal(state, model.state_dict())
checks[f"{label}: all learned state survives"] = same_state
if inactive:
omitted = prediction(model, test, include_b=False)
flipped = prediction(model, changed)
checks[f"{label}: input is inactive"] = not model.schema.requests["record/b"].active
checks[f"{label}: b values cannot affect predictions"] = bool(
np.allclose(values, omitted, rtol=1e-5, atol=tolerance)
and np.allclose(values, flipped, rtol=1e-5, atol=tolerance)
)
checks[f"{label}: removing b loses information"] = (
measured["nrmse"] > source["nrmse"] + 0.35 and measured["nrmse"] >= 0.9 * oracle["nrmse"]
)
else:
checks[f"{label}: original schema is restored"] = model.schema.model_dump() == schema
checks[f"{label}: trained predictions are restored"] = bool(
np.allclose(reference, values, rtol=1e-5, atol=tolerance)
)
edits.append(
{
"edit": label,
"inactive": inactive,
"scores": measured,
"max_source_prediction_drift": float(np.max(np.abs(values - reference))),
"learned_state_preserved": same_state,
}
)
return values
with TemporaryDirectory(prefix="relflow-mutation-") as directory:
checkpoint = Path(directory) / "source.ckpt"
model.save(checkpoint)
unchanged = rf.Model.load(checkpoint).to(model.device).eval()
checks["Unchanged checkpoint preserves predictions"] = bool(
np.allclose(reference, prediction(unchanged, test), rtol=1e-5, atol=tolerance)
)
for cycle in range(3):
model.update(selected, active=False)
inactive_values = observe(f"cycle {cycle + 1} deactivate", inactive=True)
if cycle == 0:
path = Path(directory) / "inactive.ckpt"
model.save(path)
loaded = rf.Model.load(path).to(model.device).eval()
checks["Inactive checkpoint preserves schema and state"] = (
loaded.schema.model_dump() == model.schema.model_dump() and equal(state, loaded.state_dict())
)
checks["Inactive checkpoint preserves predictions"] = bool(
np.allclose(inactive_values, prediction(loaded, test), rtol=1e-5, atol=tolerance)
)
loaded.update(selected, active=True)
checks["Loaded inactive input can restore its trained function"] = bool(
loaded.schema.model_dump() == schema
and equal(state, loaded.state_dict())
and np.allclose(reference, prediction(loaded, test), rtol=1e-5, atol=tolerance)
)
model.update(selected, active=True)
observe(f"cycle {cycle + 1} reactivate", inactive=False)
with model.override(selected, active=False):
observe("temporary override", inactive=True)
observe("normal override exit", inactive=False)
marker = RuntimeError("intentional override exit")
try:
with model.override(selected, active=False):
observe("override before exception", inactive=True)
raise marker
except RuntimeError as error:
if error is not marker:
raise
observe("exceptional override exit", inactive=False)
return {
"source": source,
"source_steps": trainer.global_step,
"source_prerequisite_met": learned,
"downstream_interpretable": learned,
"controls": {"shuffled_b": corruption, "zero_filled_b": zero_filled, "only_a_oracle": oracle},
"edits": edits,
"test_rows": len(test),
"optimizer_policy": "Fresh AdamW for source; no fitting inside updates or overrides",
"mutation": "update record/b active=False/True three times; override with normal and exceptional exits",
}, checks