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}