code wiki / _hdl_build / nx_ssm_gate.nx
nx_ssm_gate.nx source
↩ module page · 177 lines · 7799 B
1// nx_ssm_gate.nx -- GATE for SSM-001: the CAUSAL state-space mixer (the autoregressive sub-quadratic lane,
2// vs FNet's bidirectional encoder lane). Proves, by RUNNING:
3// G1 GRADCHECK through the autograd: x[4,2] + learnable decay a[2] -> ta_ssm -> mse(.,target); analytic
4// gradients for BOTH the input dx AND the decay da (backprop-through-time) vs central finite difference
5// (h=1/128, rel<1/32 floor 1/64). Proves the reverse-scan backward is correct.
6// G2 CAUSALITY: perturbing the input at the LAST time step leaves every EARLIER output unchanged, and
7// changes only the last output -- i.e. the operator's Jacobian is lower-triangular (y_t sees only x_<=t).
8// This is what FNet cannot do and what autoregressive generation requires.
9//
10// Evidence -> knowledge/status/ssm.log (SSMGATE authored=organ ... verdict=GREEN). license_tier: ORIGINAL
11import "nx_autograd_tensor.nx"
12import "nx_syscalls.nx"
13
14const SS_N: i64 = 4
15const SS_D: i64 = 2
16const SS_ND: i64 = 8
17const SS_LOG: *u8 = "knowledge/status/ssm.log"
18
19func ss_w(fd: i64, s: *u8) -> i64 { var n: i64 = 0; while s[n] != (0 as u8) { n = n + 1 } sys_write(fd, s, n); return 0 }
20func ss_wn(fd: i64, v: i64) -> i64 {
21 let bb: *u8 = sys_mmap(28); var m: i64 = v
22 if m < 0 { m = 0 - m; sys_write(fd, "-" as *u8, 1) }
23 let t: *u8 = sys_mmap(28); var k: i64 = 0
24 if m == 0 { t[0] = 48; k = 1 }
25 while m > 0 { t[k] = (48 + (m % 10)) as u8; m = m / 10; k = k + 1 }
26 var i: i64 = 0
27 while i < k { bb[i] = t[k - 1 - i]; i = i + 1 }
28 sys_write(fd, bb, k); return 0
29}
30
31func ss_build(tape: *i64, vals: *i64, st: *i64, a: *i64, x: *i64, target: *i64, ax: *i64) -> i64 {
32 st[0] = 0; st[1] = 0
33 let an: i64 = ta_leaf(tape, vals, st, SS_D, 1, a, 0)
34 let xn: i64 = ta_leaf(tape, vals, st, SS_N, SS_D, x, 0)
35 let y: i64 = ta_ssm(tape, vals, st, an, xn)
36 let tn: i64 = ta_leaf(tape, vals, st, SS_N, SS_D, target, 0)
37 let loss: i64 = ta_mse(tape, vals, st, y, tn)
38 ax[0] = an; ax[1] = xn
39 return loss
40}
41func ss_loss(tape: *i64, vals: *i64, st: *i64, a: *i64, x: *i64, target: *i64) -> i64 {
42 let ax: *i64 = (sys_mmap(2 * 8)) as *i64
43 let loss: i64 = ss_build(tape, vals, st, a, x, target, ax)
44 return ta_val(tape, vals, loss, 0)
45}
46// forward only: write all SS_ND output cells of ta_ssm(a,x) into yout
47func ss_forward(tape: *i64, vals: *i64, st: *i64, a: *i64, x: *i64, yout: *i64) -> i64 {
48 st[0] = 0; st[1] = 0
49 let an: i64 = ta_leaf(tape, vals, st, SS_D, 1, a, 0)
50 let xn: i64 = ta_leaf(tape, vals, st, SS_N, SS_D, x, 0)
51 let y: i64 = ta_ssm(tape, vals, st, an, xn)
52 var i: i64 = 0
53 while i < SS_ND { yout[i] = ta_val(tape, vals, y, i); i = i + 1 }
54 return 0
55}
56
57func main() -> i64 {
58 var ok: i64 = 1
59 let tape: *i64 = (sys_mmap(256 * 7 * 8)) as *i64
60 let vals: *i64 = (sys_mmap(4096 * 8)) as *i64
61 let grads: *i64 = (sys_mmap(4096 * 8)) as *i64
62 let st: *i64 = (sys_mmap(2 * 8)) as *i64
63
64 let a: *i64 = (sys_mmap(SS_D * 8)) as *i64
65 a[0] = ta_constf(1, 2); a[1] = ta_constf(7, 10) // decays 0.5, 0.7 (stable)
66 let x: *i64 = (sys_mmap(SS_ND * 8)) as *i64
67 let target: *i64 = (sys_mmap(SS_ND * 8)) as *i64
68 var i: i64 = 0
69 while i < SS_ND {
70 x[i] = ta_constf((i % 5) - 2, 2) // -1 .. 1
71 target[i] = ta_constf((i % 7) - 3, 4) // -0.75 .. 0.75
72 i = i + 1
73 }
74
75 // ---------- G1: gradcheck dx and da ----------
76 let ax: *i64 = (sys_mmap(2 * 8)) as *i64
77 let loss: i64 = ss_build(tape, vals, st, a, x, target, ax)
78 ta_backward(tape, vals, grads, st[0], loss)
79 let ana_dx: *i64 = (sys_mmap(SS_ND * 8)) as *i64
80 let ana_da: *i64 = (sys_mmap(SS_D * 8)) as *i64
81 i = 0
82 while i < SS_ND { ana_dx[i] = ta_grad(tape, grads, ax[1], i); i = i + 1 }
83 i = 0
84 while i < SS_D { ana_da[i] = ta_grad(tape, grads, ax[0], i); i = i + 1 }
85
86 let h: i64 = ta_constf(1, 128)
87 let flo: i64 = ta_constf(1, 64)
88 let tol: i64 = ta_constf(1, 32)
89 var gcPass: i64 = 1
90 var gcWorst: i64 = 0
91 // dx gradcheck
92 let xp: *i64 = (sys_mmap(SS_ND * 8)) as *i64
93 let xm: *i64 = (sys_mmap(SS_ND * 8)) as *i64
94 var c: i64 = 0
95 while c < SS_ND {
96 var j: i64 = 0
97 while j < SS_ND { xp[j] = x[j]; xm[j] = x[j]; j = j + 1 }
98 xp[c] = nx_f32_add(x[c], h); xm[c] = nx_f32_sub(x[c], h)
99 let lp: i64 = ss_loss(tape, vals, st, a, xp, target)
100 let lm: i64 = ss_loss(tape, vals, st, a, xm, target)
101 let fd: i64 = nx_f32_div(nx_f32_sub(lp, lm), nx_f32_add(h, h))
102 let num: i64 = nx_f32_abs(nx_f32_sub(fd, ana_dx[c]))
103 var den: i64 = nx_f32_abs(ana_dx[c])
104 if nx_f32_lt(den, flo) == 1 { den = flo }
105 if nx_f32_lt(num, nx_f32_mul(tol, den)) != 1 { gcPass = 0 }
106 let nm: i64 = ta_f32_to_milli(num)
107 if nm > gcWorst { gcWorst = nm }
108 c = c + 1
109 }
110 // da gradcheck
111 let ap: *i64 = (sys_mmap(SS_D * 8)) as *i64
112 let am: *i64 = (sys_mmap(SS_D * 8)) as *i64
113 c = 0
114 while c < SS_D {
115 var j: i64 = 0
116 while j < SS_D { ap[j] = a[j]; am[j] = a[j]; j = j + 1 }
117 ap[c] = nx_f32_add(a[c], h); am[c] = nx_f32_sub(a[c], h)
118 let lp: i64 = ss_loss(tape, vals, st, ap, x, target)
119 let lm: i64 = ss_loss(tape, vals, st, am, x, target)
120 let fd: i64 = nx_f32_div(nx_f32_sub(lp, lm), nx_f32_add(h, h))
121 let num: i64 = nx_f32_abs(nx_f32_sub(fd, ana_da[c]))
122 var den: i64 = nx_f32_abs(ana_da[c])
123 if nx_f32_lt(den, flo) == 1 { den = flo }
124 if nx_f32_lt(num, nx_f32_mul(tol, den)) != 1 { gcPass = 0 }
125 let nm: i64 = ta_f32_to_milli(num)
126 if nm > gcWorst { gcWorst = nm }
127 c = c + 1
128 }
129 if gcPass != 1 { ok = 0 }
130
131 // ---------- G2: causality (perturb last time step; earlier outputs must not change) ----------
132 let y1: *i64 = (sys_mmap(SS_ND * 8)) as *i64
133 let y2: *i64 = (sys_mmap(SS_ND * 8)) as *i64
134 ss_forward(tape, vals, st, a, x, y1)
135 let x2: *i64 = (sys_mmap(SS_ND * 8)) as *i64
136 i = 0
137 while i < SS_ND { x2[i] = x[i]; i = i + 1 }
138 var jj: i64 = 0
139 while jj < SS_D { x2[(SS_N - 1) * SS_D + jj] = nx_f32_add(x2[(SS_N - 1) * SS_D + jj], ta_constf(1, 2)); jj = jj + 1 }
140 ss_forward(tape, vals, st, a, x2, y2)
141 var causalPass: i64 = 1
142 var earlier_changed: i64 = 0
143 var last_changed: i64 = 0
144 var t: i64 = 0
145 while t < SS_N {
146 var j: i64 = 0
147 while j < SS_D {
148 let same: i64 = nx_f32_eq(y1[t * SS_D + j], y2[t * SS_D + j])
149 if t < SS_N - 1 { if same != 1 { earlier_changed = 1 } }
150 else { if same != 1 { last_changed = 1 } }
151 j = j + 1
152 }
153 t = t + 1
154 }
155 if earlier_changed != 0 { causalPass = 0 } // future input must NOT affect past outputs
156 if last_changed != 1 { causalPass = 0 } // but it MUST affect the current output
157 if causalPass != 1 { ok = 0 }
158
159 var fdi: i64 = 1
160 while fdi >= 0 {
161 var out: i64 = 1
162 if fdi == 0 { out = sys_openat_append(SS_LOG, 420) }
163 if out >= 0 {
164 ss_w(out, "SSMGATE authored=organ op=causal-state-space-scan mixer=autoregressive-subquadratic" as *u8)
165 ss_w(out, " | G1_gradcheck_dx+da_pass=" as *u8); ss_wn(out, gcPass)
166 ss_w(out, " worst_|fd-analytic|_milli=" as *u8); ss_wn(out, gcWorst)
167 ss_w(out, " | G2_causality_pass=" as *u8); ss_wn(out, causalPass)
168 ss_w(out, " earlier_changed=" as *u8); ss_wn(out, earlier_changed)
169 ss_w(out, " last_changed=" as *u8); ss_wn(out, last_changed)
170 if ok == 1 { ss_w(out, " verdict=GREEN\n" as *u8) } else { ss_w(out, " verdict=RED\n" as *u8) }
171 if fdi == 0 { sys_close(out) }
172 }
173 fdi = fdi - 1
174 }
175 if ok == 1 { return 0 }
176 return 1
177}