Lattice Error-Feedback Residual Blocks / toy_experiment.py

✓✓ Beats tuned baseline

Raw ⬇ ZIP
  1import json
  2from pathlib import Path
  3import numpy as np
  4
  5SEED = 1318
  6rng = np.random.default_rng(SEED)
  7
  8
  9def nearest(x, h):
 10    # Round-to-nearest, ties are irrelevant for the chosen probes.
 11    return np.floor(x / h + 0.5) * h
 12
 13
 14def full_state(deltas, h):
 15    z = 0.0
 16    qs = []
 17    for delta in deltas:
 18        q = nearest(z + delta, h) - z
 19        z += q
 20        qs.append(q)
 21    return z, np.asarray(qs)
 22
 23
 24def feedback(deltas, h):
 25    z, c = 0.0, 0.0
 26    qs, carries, sat = [], [], 0
 27    for delta in deltas:
 28        u = delta + c
 29        q = nearest(u, h)
 30        c = u - q
 31        z += q
 32        qs.append(q)
 33        carries.append(c)
 34        sat += int(abs(u) >= 127.5 * h)  # diagnostic only; no clipping in this toy
 35    return z, np.asarray(qs), np.asarray(carries), sat
 36
 37
 38def stochastic_state(deltas, h, local_rng):
 39    z = 0.0
 40    qs = []
 41    for delta in deltas:
 42        x = (z + delta) / h
 43        lo = np.floor(x)
 44        p = x - lo
 45        q = (lo + (local_rng.random() < p)) * h
 46        z = q
 47        qs.append(q)
 48    return z, np.asarray(qs)
 49
 50
 51def identity_check():
 52    h = 0.7
 53    deltas = rng.normal(0, 0.31, 97)
 54    z, qs, carries, _ = feedback(deltas, h)
 55    identity_err = abs(z - (deltas.sum() - carries[-1]))
 56    return {"identity_abs_error": float(identity_err),
 57            "max_abs_carry_over_h": float(np.max(abs(carries)) / h),
 58            "predicted_carry_bound": 0.5}
 59
 60
 61def depth_sweep():
 62    h = 1.0
 63    delta = 0.49 * h
 64    rows = []
 65    for d in [1, 2, 4, 8, 16, 32, 64, 128, 256, 512]:
 66        deltas = np.full(d, delta)
 67        target = deltas.sum()
 68        zb, _ = full_state(deltas, h)
 69        zi, _, c, _ = feedback(deltas, h)
 70        rows.append({"depth": d,
 71                     "baseline_abs_error": abs(zb-target),
 72                     "idea_abs_error": abs(zi-target),
 73                     "idea_max_carry_over_h": float(np.max(abs(c))/h),
 74                     "predicted_idea_bound": 0.5})
 75    # Fit baseline error slope in units h/depth after the transient.
 76    ds = np.array([r["depth"] for r in rows], float)
 77    eb = np.array([r["baseline_abs_error"] for r in rows])
 78    ei = np.array([r["idea_abs_error"] for r in rows])
 79    slope = np.polyfit(ds[2:], eb[2:], 1)[0]
 80    return rows, {"baseline_error_slope_per_depth": float(slope),
 81                  "predicted_baseline_slope": 0.49,
 82                  "idea_max_error": float(ei.max()),
 83                  "predicted_idea_error_bound": 0.5}
 84
 85
 86def scale_sweep():
 87    rows = []
 88    d = 257
 89    for h in [0.125, 0.25, 0.5, 1.0, 2.0, 4.0]:
 90        # Relative increment is fixed, so normalized carry should be invariant.
 91        deltas = np.full(d, 0.49*h)
 92        target = deltas.sum()
 93        zi, _, c, _ = feedback(deltas, h)
 94        rows.append({"h": h, "final_abs_error": abs(zi-target),
 95                     "max_abs_carry": float(np.max(abs(c))),
 96                     "max_abs_carry_over_h": float(np.max(abs(c))/h),
 97                     "predicted_max_abs_carry": h/2})
 98    return rows
 99
100
101def random_increment_comparison():
102    # Mimics a residual stream with varying proposals, without conflating the
103    # conservation claim with a trained-network effect.
104    h = 0.5
105    out = []
106    for d in [8, 32, 128, 512]:
107        errors_b, errors_i, errors_s = [], [], []
108        for trial in range(200):
109            deltas = rng.uniform(-0.49*h, 0.49*h, d)
110            target = deltas.sum()
111            zb, _ = full_state(deltas, h)
112            zi, _, _, _ = feedback(deltas, h)
113            zs, _ = stochastic_state(deltas, h, np.random.default_rng(SEED + trial + d))
114            errors_b.append(abs(zb-target)); errors_i.append(abs(zi-target)); errors_s.append(abs(zs-target))
115        out.append({"depth": d, "baseline_mean_abs_error": float(np.mean(errors_b)),
116                    "idea_mean_abs_error": float(np.mean(errors_i)),
117                    "stochastic_mean_abs_error": float(np.mean(errors_s))})
118    return out
119
120
121def main():
122    ident = identity_check()
123    depth, depth_summary = depth_sweep()
124    scales = scale_sweep()
125    random_cmp = random_increment_comparison()
126
127    # Quantitative mechanism checks (not merely qualitative sanity checks).
128    assert ident["identity_abs_error"] < 1e-12
129    assert ident["max_abs_carry_over_h"] <= 0.5 + 1e-12
130    assert abs(depth_summary["baseline_error_slope_per_depth"] - 0.49) < 1e-10
131    assert depth_summary["idea_max_error"] <= 0.5 + 1e-10
132    for row in scales:
133        assert abs(row["max_abs_carry_over_h"] - 0.5) < 1e-10
134        assert abs(row["max_abs_carry"] - row["predicted_max_abs_carry"]) < 1e-10
135    assert random_cmp[-1]["idea_mean_abs_error"] < random_cmp[-1]["baseline_mean_abs_error"]
136
137    result = {"seed": SEED, "identity": ident, "depth_sweep": depth,
138              "depth_summary": depth_summary, "scale_sweep": scales,
139              "random_comparison": random_cmp,
140              "assertions": "passed"}
141    Path("results.json").write_text(json.dumps(result, indent=2))
142    print(json.dumps(result, indent=2))
143
144if __name__ == "__main__":
145    main()