Reset One Learned Target and Teach It Again

Mutation reset
Reset one hidden regression head, verify immediate loss is localized, and compare relearning against unchanged and complete-reset controls.

Reset should erase selected learned state. This experiment asks whether the loss is visible in predictions, whether an unselected hidden target survives, and whether the selected target can learn again.

Seed 7102 on gpu: 17 of 17 behavioral checks met. See the measurements below.

Insights

Reset erases one hidden head’s learned prediction while its sibling retains it. The full runs show localized changes to state and held-out behavior, followed by successful relearning of the selected target.

Both targets remain supervised during adaptation. An unchanged continuation and a complete reset use the same adaptation examples and update budget. Their scores provide context without assuming selective reset must learn faster. Branch resets can affect shared context and need a separate proof.

Setup

Code
"""P048: localize loss of learned behavior after reset and demonstrate relearning."""

from collections.abc import Iterator
from copy import deepcopy
from functools import partial
from pathlib import Path
from tempfile import TemporaryDirectory

import lightning.pytorch as lit
import numpy as np
import torch
from reporting import report

import relflow as rf

PROOF_ID = "P048"
TARGETS = ("u", "v")

Examples

Draw independent a and b uniformly from [−1, 1]. Both targets are hidden: u = a + 2b, and v = 2a − b.

a: 0.5
b: 0.25
u: 1.0
v: 0.75

Holding a fixed while changing b changes both answers:

a: 0.5
b: -0.25
u: 0.0
v: 1.25

Holding b fixed does not determine either answer:

a: -0.5
b: 0.25
u: 0.0
v: -1.25

Data and comparisons

Code
def records(*, rows: int, seed: int) -> Iterator[dict]:
    rng = np.random.default_rng(seed)
    for _ in range(rows):
        a, b = map(float, rng.uniform(-1, 1, size=2))
        yield {"a": a, "b": b, "u": a + 2 * b, "v": 2 * a - b}


def prediction(model: rf.Model, rows: list[dict], *, corrupt: bool = False) -> dict:
    inputs = [{"a": row["a"], "b": row["b"]} for row in rows]
    if corrupt:
        for row in inputs:
            row.update({name: 1000.0 for name in TARGETS})
    output = model.predict(inputs)["predictions"].to_pylist()
    return {
        name: np.asarray([row[f"record/{name}"]["content"] for row in output], dtype=np.float64) for name in TARGETS
    }


def scores(rows: list[dict], predicted: dict, means: dict) -> dict:
    measured = {}
    for name, values in predicted.items():
        actual = np.asarray([row[name] for row in rows])
        baseline = float(np.sqrt(np.mean((actual - means[name]) ** 2)))
        rmse = float(np.sqrt(np.mean((actual - values) ** 2)))
        measured[name] = {"rmse": rmse, "baseline_rmse": baseline, "nrmse": rmse / baseline}
    return measured


def equal(first, second) -> bool:
    """Compare tensor and extension-owned state, including normalization."""
    if isinstance(first, torch.Tensor):
        return isinstance(second, torch.Tensor) and torch.equal(first, second)
    if isinstance(first, dict):
        return (
            isinstance(second, dict)
            and first.keys() == second.keys()
            and all(equal(value, second[name]) for name, value in first.items())
        )
    if isinstance(first, (tuple, list)):
        return (
            type(first) is type(second)
            and len(first) == len(second)
            and all(equal(a, b) for a, b in zip(first, second, strict=True))
        )
    return type(first) is type(second) and first == second


class Curve(lit.Callback):
    """Observe fixed validation checkpoints without restarting the optimizer."""

    def __init__(self, rows: list[dict], means: dict, budget: int):
        self.rows, self.means = rows, means
        self.checkpoints = {32, 128, budget}
        self.measurements = []

    def on_train_start(self, trainer, pl_module):
        self.measurements.append(
            {
                "step": 0,
                "scores": scores(self.rows, prediction(pl_module, self.rows), self.means),
            }
        )

    def on_train_batch_end(self, trainer, pl_module, outputs, batch, batch_idx):
        if trainer.global_step in self.checkpoints:
            self.measurements.append(
                {
                    "step": trainer.global_step,
                    "scores": scores(self.rows, prediction(pl_module, self.rows), self.means),
                }
            )


def fit(model: rf.Model, *, seed: int, split: int, budget: int, accelerator: str) -> dict:
    lit.seed_everything(seed, workers=True)
    training = list(records(rows=2048, seed=split))
    means = {name: float(np.mean([row[name] for row in training])) for name in TARGETS}
    curve = Curve(list(records(rows=512, seed=split + 1)), means, budget)
    model.optimizer = rf.adamw(learning_rate=3e-3, fused=False)
    data = rf.SyntheticDataModule(
        model=model,
        train=partial(records, rows=2048, seed=split),
        validate=partial(records, rows=512, seed=split + 1),
        seed=seed,
    )
    trainer = lit.Trainer(
        accelerator=accelerator,
        devices=1,
        max_epochs=-1,
        max_steps=budget,
        callbacks=[curve],
        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()
    optimized = {id(p) for group in trainer.optimizers[0].param_groups for p in group["params"]}
    current = {id(p) for p in model.parameters() if p.requires_grad}
    return {
        "steps": trainer.global_step,
        "optimizer_covers_current_parameters": current == optimized,
        "validation_curve": curve.measurements,
    }

Scope of the reset

Record reads a and b to predict hidden u and v. Only the runtime node for u is reset. The root, inputs, and v retain their trained state.

Record reads a and b to predict hidden u and v. Only the runtime node for u is reset. The root, inputs, and v retain their trained state.

Figure 1: Reset replaces the selected target’s runtime state; its schema and task stay fixed.

Source training uses 2,048 records and 512 updates. Adaptation uses 2,048 independently drawn records and 256 updates. The phases each have 512 validation records. A separate, fixed panel of 1,024 records measures immediate loss, retention, and final skill for every arm.

Validation curves record steps 0, 32, 128, and the final update. All arms restart AdamW and rehearse both targets. No test result selects a checkpoint. Test nRMSE divides RMSE by the error of predicting the source-training target mean. Validation curves use each phase’s training mean. The provisional learning gate is 0.25; loss of skill requires nRMSE above 0.8 and three times source error. Strict retention uses rtol=1e-5 and atol=1e-6 times training SD.

Training, reset, and relearning

Code
def run(seed: int, steps: int | None, accelerator: str) -> tuple[dict, dict]:
    source_budget = 512 if steps is None else min(steps, 512)
    adapt_budget = 256 if steps is None else min(steps, 256)
    lit.seed_everything(seed, workers=True)
    training = list(records(rows=2048, seed=seed + 1))
    test = list(records(rows=1024, seed=seed + 3))
    means = {name: float(np.mean([row[name] for row in training])) for name in TARGETS}
    scales = {name: float(np.std([row[name] for row in training])) for name in TARGETS}
    source = rf.Model(
        d_model=32,
        n_layers=1,
        n_heads=4,
        dropout=0.0,
        batch_size=128,
        a=rf.Number,
        b=rf.Number,
        u=rf.Number(mask=True, objective="mse"),
        v=rf.Number(mask=True, objective="mse"),
    )
    source_fit = fit(source, seed=seed, split=seed + 1, budget=source_budget, accelerator=accelerator)
    reference = prediction(source, test)
    initial = scores(test, reference, means)
    learned = all(value["nrmse"] < 0.25 for value in initial.values())
    checks = {"Source learns both targets below 0.25 nRMSE": learned}
    state = deepcopy(source.state_dict())
    arms = {}

    with TemporaryDirectory(prefix="relflow-mutation-") as directory:
        checkpoint = Path(directory) / "source.ckpt"
        source.save(checkpoint)
        for index, arm in enumerate(("selective_reset", "continuation", "complete_reset")):
            lit.seed_everything(seed + 100 + index, workers=True)
            model = rf.Model.load(checkpoint).to(source.device)
            if arm == "selective_reset":
                model.reset(rf.where("address") == "record/u")
            elif arm == "complete_reset":
                model.reset(rf.where("address") == "record", descendants=True)
            model.eval()
            current = model.state_dict()
            changes = [name for name, value in state.items() if name not in current or not equal(value, current[name])]
            before = prediction(model, test)
            immediate = scores(test, before, means)
            if arm == "selective_reset":
                checks["Reset changes selected learned state"] = any(
                    name.startswith("nodes.record/u.") for name in changes
                )
                checks["Reset preserves every unselected state entry"] = all(
                    name.startswith("nodes.record/u.") for name in changes
                )
                checks["Reset preserves the schema"] = model.schema.model_dump() == source.schema.model_dump()
                checks["Selected target immediately loses its skill"] = (
                    immediate["u"]["nrmse"] > 0.8 and immediate["u"]["nrmse"] > 3 * initial["u"]["nrmse"]
                )
                checks["Unselected target immediately preserves predictions"] = bool(
                    np.allclose(reference["v"], before["v"], rtol=1e-5, atol=1e-6 * scales["v"])
                )
                head = {name: p.detach().clone() for name, p in model.nodes["record/u"].named_parameters()}
            if arm == "complete_reset":
                checks["Complete reset loses both learned tasks"] = all(
                    immediate[name]["nrmse"] > 0.8 and immediate[name]["nrmse"] > 3 * initial[name]["nrmse"]
                    for name in TARGETS
                )
            fitting = fit(model, seed=seed + 200, split=seed + 11, budget=adapt_budget, accelerator=accelerator)
            after = prediction(model, test)
            final = scores(test, after, means)
            arms[arm] = {
                "before": immediate,
                "after": final,
                "changed_state_entries": changes,
                "initial_prediction_drift": {
                    name: float(np.max(np.abs(reference[name] - before[name]))) for name in TARGETS
                },
                **fitting,
            }
            checks[f"{arm}: fresh optimizer covers current parameters"] = fitting["optimizer_covers_current_parameters"]
            checks[f"{arm}: both final targets below 0.25 nRMSE"] = all(
                value["nrmse"] < 0.25 for value in final.values()
            )
            if arm == "selective_reset":
                changed = [
                    name for name, p in model.nodes["record/u"].named_parameters() if not torch.equal(head[name], p)
                ]
                arms[arm]["reset_head_updated_parameters"] = changed
                checks["Reset head parameters actually relearn"] = bool(changed)
                corrupted = prediction(model, test, corrupt=True)
                checks["Hidden target placeholders cannot affect predictions"] = all(
                    np.allclose(after[name], corrupted[name], rtol=1e-5, atol=1e-6 * scales[name]) for name in TARGETS
                )
                path = Path(directory) / "relearned.ckpt"
                model.save(path)
                loaded = rf.Model.load(path).to(model.device).eval()
                restored = prediction(loaded, test)
                checks["Relearned checkpoint preserves schema and state"] = (
                    loaded.schema.model_dump() == model.schema.model_dump()
                    and equal(model.state_dict(), loaded.state_dict())
                )
                checks["Relearned checkpoint preserves predictions"] = all(
                    np.allclose(after[name], restored[name], rtol=1e-5, atol=1e-6 * scales[name]) for name in TARGETS
                )
    return {
        "source": initial,
        "source_fit": source_fit,
        "source_prerequisite_met": learned,
        "downstream_interpretable": learned,
        "arms": arms,
        "test_rows": len(test),
        "baseline": "Source-training mean for test comparisons; phase-training mean for validation curves",
        "optimizer_policy": "New AdamW factory and Trainer for every fit; both hidden labels rehearsed",
        "mutation": "reset record/u; complete-reset control resets record with descendants=True",
    }, checks

Evidence

Latest full run

Seed 7102, gpu, recorded 2026-09-15T18:03:47.205739+00:00. Outcome: met.

Source fingerprint: 2252cee7efe14048db1252663050b74fce7ce20d11f143c1b24b93df622d6f7c. Python 3.12.6; Torch 2.12.0.

Measurement Value
source u: rmse: 0.065655; baseline_rmse: 1.29673; nrmse: 0.0506312; v: rmse: 0.0413812; baseline_rmse: 1.26931; nrmse: 0.0326013
source_fit steps: 512; optimizer_covers_current_parameters: True; validation_curve: step: 0; scores: u: rmse: 1.2393; baseline_rmse: 1.27237; nrmse: 0.974008; v: rmse: 1.70215; baseline_rmse: 1.34685; nrmse: 1.2638, step: 32; scores: u: rmse: 0.103551; baseline_rmse: 1.27237; nrmse: 0.0813844; v: rmse: 0.089618; baseline_rmse: 1.34685; nrmse: 0.0665388, step: 128; scores: u: rmse: 0.0397873; baseline_rmse: 1.27237; nrmse: 0.0312701; v: rmse: 0.0490955; baseline_rmse: 1.34685; nrmse: 0.036452, step: 512; scores: u: rmse: 0.0649841; baseline_rmse: 1.27237; nrmse: 0.0510731; v: rmse: 0.0399428; baseline_rmse: 1.34685; nrmse: 0.0296564
source_prerequisite_met True
downstream_interpretable True
arms selective_reset: before: u: rmse: 1.24894; baseline_rmse: 1.29673; nrmse: 0.963143; v: rmse: 0.0413812; baseline_rmse: 1.26931; nrmse: 0.0326013; after: u: rmse: 0.021251; baseline_rmse: 1.29673; nrmse: 0.0163882; v: rmse: 0.0129122; baseline_rmse: 1.26931; nrmse: 0.0101725; changed_state_entries: 37 values; final 8: nodes.record/u.decoder.pool.blocks.0.ffn.3.weight, nodes.record/u.decoder.pool.blocks.0.ffn.3.bias, nodes.record/u.decoder.pool.norm.weight, nodes.record/u.decoder.pool.norm.bias, nodes.record/u.decoder.classification.weight, nodes.record/u.decoder.classification.bias, nodes.record/u.decoder.regression.weight, nodes.record/u.decoder.regression.bias; initial_prediction_drift: u: 2.99012; v: 0; steps: 256; optimizer_covers_current_parameters: True; validation_curve: step: 0; scores: u: rmse: 1.22038; baseline_rmse: 1.27617; nrmse: 0.956291; v: rmse: 0.0413389; baseline_rmse: 1.3344; nrmse: 0.0309794, step: 32; scores: u: rmse: 0.0421101; baseline_rmse: 1.27617; nrmse: 0.0329974; v: rmse: 0.061978; baseline_rmse: 1.3344; nrmse: 0.0464464, step: 128; scores: u: rmse: 0.0182268; baseline_rmse: 1.27617; nrmse: 0.0142825; v: rmse: 0.0132202; baseline_rmse: 1.3344; nrmse: 0.00990719, step: 256; scores: u: rmse: 0.0207765; baseline_rmse: 1.27617; nrmse: 0.0162804; v: rmse: 0.0127284; baseline_rmse: 1.3344; nrmse: 0.0095387; reset_head_updated_parameters: 32 values; final 8: decoder.pool.blocks.0.ffn.3.weight, decoder.pool.blocks.0.ffn.3.bias, decoder.pool.norm.weight, decoder.pool.norm.bias, decoder.classification.weight, decoder.classification.bias, decoder.regression.weight, decoder.regression.bias; continuation: before: u: rmse: 0.065655; baseline_rmse: 1.29673; nrmse: 0.0506312; v: rmse: 0.0413812; baseline_rmse: 1.26931; nrmse: 0.0326013; after: u: rmse: 0.0199771; baseline_rmse: 1.29673; nrmse: 0.0154058; v: rmse: 0.00927444; baseline_rmse: 1.26931; nrmse: 0.00730665; changed_state_entries: ; initial_prediction_drift: u: 0; v: 0; steps: 256; optimizer_covers_current_parameters: True; validation_curve: step: 0; scores: u: rmse: 0.0649856; baseline_rmse: 1.27617; nrmse: 0.0509226; v: rmse: 0.0413389; baseline_rmse: 1.3344; nrmse: 0.0309794, step: 32; scores: u: rmse: 0.0529592; baseline_rmse: 1.27617; nrmse: 0.0414987; v: rmse: 0.104475; baseline_rmse: 1.3344; nrmse: 0.0782937, step: 128; scores: u: rmse: 0.00932613; baseline_rmse: 1.27617; nrmse: 0.00730793; v: rmse: 0.0175418; baseline_rmse: 1.3344; nrmse: 0.0131459, step: 256; scores: u: rmse: 0.0193039; baseline_rmse: 1.27617; nrmse: 0.0151265; v: rmse: 0.0104376; baseline_rmse: 1.3344; nrmse: 0.00782196; complete_reset: before: u: rmse: 1.34865; baseline_rmse: 1.29673; nrmse: 1.04004; v: rmse: 1.20858; baseline_rmse: 1.26931; nrmse: 0.952152; after: u: rmse: 0.0217558; baseline_rmse: 1.29673; nrmse: 0.0167774; v: rmse: 0.012717; baseline_rmse: 1.26931; nrmse: 0.0100188; changed_state_entries: 143 values; final 8: nodes.record.encoder.pool.blocks.0.attention.out_proj.bias, nodes.record.encoder.pool.blocks.0.ffn.0.weight, nodes.record.encoder.pool.blocks.0.ffn.0.bias, nodes.record.encoder.pool.blocks.0.ffn.3.weight, nodes.record.encoder.pool.blocks.0.ffn.3.bias, nodes.record.encoder.pool.norm.weight, nodes.record.encoder.pool.norm.bias, nodes.record.encoder.pool.mass_projection.weight; initial_prediction_drift: u: 3.20103; v: 2.76781; steps: 256; optimizer_covers_current_parameters: True; validation_curve: step: 0; scores: u: rmse: 1.31619; baseline_rmse: 1.27617; nrmse: 1.03136; v: rmse: 1.27491; baseline_rmse: 1.3344; nrmse: 0.955422, step: 32; scores: u: rmse: 0.0763504; baseline_rmse: 1.27617; nrmse: 0.059828; v: rmse: 0.0873164; baseline_rmse: 1.3344; nrmse: 0.065435, step: 128; scores: u: rmse: 0.0492419; baseline_rmse: 1.27617; nrmse: 0.0385859; v: rmse: 0.036472; baseline_rmse: 1.3344; nrmse: 0.0273322, step: 256; scores: u: rmse: 0.0223337; baseline_rmse: 1.27617; nrmse: 0.0175006; v: rmse: 0.0124282; baseline_rmse: 1.3344; nrmse: 0.00931371
test_rows 1024
baseline Source-training mean for test comparisons; phase-training mean for validation curves
optimizer_policy New AdamW factory and Trainer for every fit; both hidden labels rehearsed
mutation reset record/u; complete-reset control resets record with descendants=True
Behavioral checks
Behavioral check Outcome
Source learns both targets below 0.25 nRMSE Met
Reset changes selected learned state Met
Reset preserves every unselected state entry Met
Reset preserves the schema Met
Selected target immediately loses its skill Met
Unselected target immediately preserves predictions Met
selective_reset: fresh optimizer covers current parameters Met
selective_reset: both final targets below 0.25 nRMSE Met
Reset head parameters actually relearn Met
Hidden target placeholders cannot affect predictions Met
Relearned checkpoint preserves schema and state Met
Relearned checkpoint preserves predictions Met
continuation: fresh optimizer covers current parameters Met
continuation: both final targets below 0.25 nRMSE Met
Complete reset loses both learned tasks Met
complete_reset: fresh optimizer covers current parameters Met
complete_reset: both final targets below 0.25 nRMSE Met

Recorded results.

Remaining work

Three GPU seeds support this case; calibrate gates on ten separate seeds. Resetting a context-producing branch, resetting descendants, and relearning without sibling labels require their own experiments. Localized loss here depends on the selected field being hidden from the shared context.

Reproduce

Run by stable ID from the repository root:

uv run python proofs/run.py P048

Or run the self-contained script directly:

PYTHONPATH=proofs uv run python proofs/mutations/reset_prediction_target.py

Add --accelerator gpu for CUDA or --seed 42 for another seeded experiment. --steps 2 checks execution with a short training budget; it is recorded as a smoke run.

Download the complete proof.

Code
if __name__ == "__main__":
    report(PROOF_ID, run, seed=4801)