Lattice Error-Feedback Residual Blocks / toy_experiment.py
Beats tuned baseline
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()