code wiki / _hdl_build / nx_ssm_gate.nx

nx_ssm_gate.nx source

↩ module page · 185 lines · 8328 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" 13import "nx_gate_verdict.nx" 14 15const SS_N: i64 = 4 16const SS_D: i64 = 2 17const SS_ND: i64 = 8 18const SS_LOG: *u8 = "knowledge/status/ssm.log" 19 20func 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 } 21func ss_wn(fd: i64, v: i64) -> i64 { 22 let bb: *u8 = sys_mmap(28); var m: i64 = v 23 if m < 0 { m = 0 - m; sys_write(fd, "-" as *u8, 1) } 24 let t: *u8 = sys_mmap(28); var k: i64 = 0 25 if m == 0 { t[0] = 48; k = 1 } 26 while m > 0 { t[k] = (48 + (m % 10)) as u8; m = m / 10; k = k + 1 } 27 var i: i64 = 0 28 while i < k { bb[i] = t[k - 1 - i]; i = i + 1 } 29 sys_write(fd, bb, k); return 0 30} 31 32func ss_build(tape: *i64, vals: *i64, st: *i64, a: *i64, x: *i64, target: *i64, ax: *i64) -> i64 { 33 st[0] = 0; st[1] = 0 34 let an: i64 = ta_leaf(tape, vals, st, SS_D, 1, a, 0) 35 let xn: i64 = ta_leaf(tape, vals, st, SS_N, SS_D, x, 0) 36 let y: i64 = ta_ssm(tape, vals, st, an, xn) 37 let tn: i64 = ta_leaf(tape, vals, st, SS_N, SS_D, target, 0) 38 let loss: i64 = ta_mse(tape, vals, st, y, tn) 39 ax[0] = an; ax[1] = xn 40 return loss 41} 42func ss_loss(tape: *i64, vals: *i64, st: *i64, a: *i64, x: *i64, target: *i64) -> i64 { 43 let ax: *i64 = (sys_mmap(2 * 8)) as *i64 44 let loss: i64 = ss_build(tape, vals, st, a, x, target, ax) 45 return ta_val(tape, vals, loss, 0) 46} 47// forward only: write all SS_ND output cells of ta_ssm(a,x) into yout 48func ss_forward(tape: *i64, vals: *i64, st: *i64, a: *i64, x: *i64, yout: *i64) -> i64 { 49 st[0] = 0; st[1] = 0 50 let an: i64 = ta_leaf(tape, vals, st, SS_D, 1, a, 0) 51 let xn: i64 = ta_leaf(tape, vals, st, SS_N, SS_D, x, 0) 52 let y: i64 = ta_ssm(tape, vals, st, an, xn) 53 var i: i64 = 0 54 while i < SS_ND { yout[i] = ta_val(tape, vals, y, i); i = i + 1 } 55 return 0 56} 57 58func main() -> i64 { 59 var ok: i64 = 1 60 let tape: *i64 = (sys_mmap(256 * 7 * 8)) as *i64 61 let vals: *i64 = (sys_mmap(4096 * 8)) as *i64 62 let grads: *i64 = (sys_mmap(4096 * 8)) as *i64 63 let st: *i64 = (sys_mmap(2 * 8)) as *i64 64 65 let a: *i64 = (sys_mmap(SS_D * 8)) as *i64 66 a[0] = ta_constf(1, 2); a[1] = ta_constf(7, 10) // decays 0.5, 0.7 (stable) 67 let x: *i64 = (sys_mmap(SS_ND * 8)) as *i64 68 let target: *i64 = (sys_mmap(SS_ND * 8)) as *i64 69 var i: i64 = 0 70 while i < SS_ND { 71 x[i] = ta_constf((i % 5) - 2, 2) // -1 .. 1 72 target[i] = ta_constf((i % 7) - 3, 4) // -0.75 .. 0.75 73 i = i + 1 74 } 75 76 // ---------- G1: gradcheck dx and da ---------- 77 let ax: *i64 = (sys_mmap(2 * 8)) as *i64 78 let loss: i64 = ss_build(tape, vals, st, a, x, target, ax) 79 ta_backward(tape, vals, grads, st[0], loss) 80 let ana_dx: *i64 = (sys_mmap(SS_ND * 8)) as *i64 81 let ana_da: *i64 = (sys_mmap(SS_D * 8)) as *i64 82 i = 0 83 while i < SS_ND { ana_dx[i] = ta_grad(tape, grads, ax[1], i); i = i + 1 } 84 i = 0 85 while i < SS_D { ana_da[i] = ta_grad(tape, grads, ax[0], i); i = i + 1 } 86 87 let h: i64 = ta_constf(1, 128) 88 let flo: i64 = ta_constf(1, 64) 89 let tol: i64 = ta_constf(1, 32) 90 var gcPass: i64 = 1 91 var gcWorst: i64 = 0 92 // dx gradcheck 93 let xp: *i64 = (sys_mmap(SS_ND * 8)) as *i64 94 let xm: *i64 = (sys_mmap(SS_ND * 8)) as *i64 95 var c: i64 = 0 96 while c < SS_ND { 97 var j: i64 = 0 98 while j < SS_ND { xp[j] = x[j]; xm[j] = x[j]; j = j + 1 } 99 xp[c] = nx_f32_add(x[c], h); xm[c] = nx_f32_sub(x[c], h) 100 let lp: i64 = ss_loss(tape, vals, st, a, xp, target) 101 let lm: i64 = ss_loss(tape, vals, st, a, xm, target) 102 let fd: i64 = nx_f32_div(nx_f32_sub(lp, lm), nx_f32_add(h, h)) 103 let num: i64 = nx_f32_abs(nx_f32_sub(fd, ana_dx[c])) 104 var den: i64 = nx_f32_abs(ana_dx[c]) 105 if nx_f32_lt(den, flo) == 1 { den = flo } 106 if nx_f32_lt(num, nx_f32_mul(tol, den)) != 1 { gcPass = 0 } 107 let nm: i64 = ta_f32_to_milli(num) 108 if nm > gcWorst { gcWorst = nm } 109 c = c + 1 110 } 111 // da gradcheck 112 let ap: *i64 = (sys_mmap(SS_D * 8)) as *i64 113 let am: *i64 = (sys_mmap(SS_D * 8)) as *i64 114 c = 0 115 while c < SS_D { 116 var j: i64 = 0 117 while j < SS_D { ap[j] = a[j]; am[j] = a[j]; j = j + 1 } 118 ap[c] = nx_f32_add(a[c], h); am[c] = nx_f32_sub(a[c], h) 119 let lp: i64 = ss_loss(tape, vals, st, ap, x, target) 120 let lm: i64 = ss_loss(tape, vals, st, am, x, target) 121 let fd: i64 = nx_f32_div(nx_f32_sub(lp, lm), nx_f32_add(h, h)) 122 let num: i64 = nx_f32_abs(nx_f32_sub(fd, ana_da[c])) 123 var den: i64 = nx_f32_abs(ana_da[c]) 124 if nx_f32_lt(den, flo) == 1 { den = flo } 125 if nx_f32_lt(num, nx_f32_mul(tol, den)) != 1 { gcPass = 0 } 126 let nm: i64 = ta_f32_to_milli(num) 127 if nm > gcWorst { gcWorst = nm } 128 c = c + 1 129 } 130 if gcPass != 1 { ok = 0 } 131 132 // ---------- G2: causality (perturb last time step; earlier outputs must not change) ---------- 133 let y1: *i64 = (sys_mmap(SS_ND * 8)) as *i64 134 let y2: *i64 = (sys_mmap(SS_ND * 8)) as *i64 135 ss_forward(tape, vals, st, a, x, y1) 136 let x2: *i64 = (sys_mmap(SS_ND * 8)) as *i64 137 i = 0 138 while i < SS_ND { x2[i] = x[i]; i = i + 1 } 139 var jj: i64 = 0 140 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 } 141 ss_forward(tape, vals, st, a, x2, y2) 142 var causalPass: i64 = 1 143 var earlier_changed: i64 = 0 144 var last_changed: i64 = 0 145 var t: i64 = 0 146 while t < SS_N { 147 var j: i64 = 0 148 while j < SS_D { 149 let same: i64 = nx_f32_eq(y1[t * SS_D + j], y2[t * SS_D + j]) 150 if t < SS_N - 1 { if same != 1 { earlier_changed = 1 } } 151 else { if same != 1 { last_changed = 1 } } 152 j = j + 1 153 } 154 t = t + 1 155 } 156 if earlier_changed != 0 { causalPass = 0 } // future input must NOT affect past outputs 157 if last_changed != 1 { causalPass = 0 } // but it MUST affect the current output 158 if causalPass != 1 { ok = 0 } 159 160 var fdi: i64 = 1 161 while fdi >= 0 { 162 var out: i64 = 1 163 if fdi == 0 { out = sys_openat_append(SS_LOG, 420) } 164 if out >= 0 { 165 ss_w(out, "SSMGATE authored=organ op=causal-state-space-scan mixer=autoregressive-subquadratic" as *u8) 166 ss_w(out, " | G1_gradcheck_dx+da_pass=" as *u8); ss_wn(out, gcPass) 167 ss_w(out, " worst_|fd-analytic|_milli=" as *u8); ss_wn(out, gcWorst) 168 ss_w(out, " | G2_causality_pass=" as *u8); ss_wn(out, causalPass) 169 ss_w(out, " earlier_changed=" as *u8); ss_wn(out, earlier_changed) 170 ss_w(out, " last_changed=" as *u8); ss_wn(out, last_changed) 171 if ok == 1 { ss_w(out, " verdict=GREEN\n" as *u8) } else { ss_w(out, " verdict=RED\n" as *u8) } 172 if fdi == 0 { sys_close(out) } 173 } 174 fdi = fdi - 1 175 } 176 // MIGRATED onto nx_gate_verdict by nx_gate_dry_apply (D001, minimal form): every check 177 // row above is untouched, so the PASS/FAIL vector cannot change; only the hand-rolled 178 // verdict emission is replaced by the ONE shared base class. Proven by nx_gate_migrate verify. 179 let ctr__dry: *i64 = gv_ctr() 180 ctr__dry[0] = ok 181 ctr__dry[1] = 1 182 let rc__dry: i64 = gv_verdict("SSM-GATE" as *u8, ctr__dry, "teeth unchanged; verdict emission migrated onto the shared base class" as *u8) 183 sys_exit(rc__dry) 184 return rc__dry 185}