code wiki / _hdl_build / nx_nofloat_multihead_gate.nx
nx_nofloat_multihead_gate.nx source
↩ module page · 252 lines · 15554 B
1// nx_nofloat_multihead_gate.nx -- HARD-EVIDENCE gate for MULTI-HEAD attention in pure integer Q16. The defining
2// transformer feature: split Q/K/V into H heads (slice_cols), run the verified attention core per head (RoPE +
3// scaled causal softmax + A.V), concat the heads (concat_cols), then the output projection. Backprops end-to-end.
4//
5// A1 slice_cols gradcheck : extract cols backward (scatter) == finite differences.
6// A2 concat_cols gradcheck : concat backward (split) == finite differences.
7// A3 multi-head gradcheck wrt Wq : the gradient through BOTH heads (slice+RoPE+softmax+A.V+concat+proj) == FD.
8// D neg-control teeth ; C bit-exact.
9// B multi-head LEARNS : with attention fixed, the output projection Wo is trained (AdamW) to a realizable
10// target; loss collapses -> the multi-head output path trains.
11//
12// Evidence -> knowledge/status/nofloat_multihead.log. Sovereign: nx_nofloat_autograd + nx_syscalls. expect_exit: 0
13import "nx_nofloat_autograd.nx"
14import "nx_syscalls.nx"
15import "nx_gate_emit_lib.nx"
16
17const MLOG: *u8 = "knowledge/status/nofloat_multihead.log"
18const Q16: i64 = 65536
19
20
21func g_abs(v: i64) -> i64 { if v < 0 { return 0 - v } return v }
22func m_ws(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 }
23func m_wn(fd: i64, v: i64) -> i64 { let b: *u8=sys_mmap(28); var m: i64=v; if m<0{sys_write(fd,"-" as *u8,1);m=0-m} let t: *u8=sys_mmap(28); var k: i64=0; if m==0{t[0]=48;k=1} while m>0{t[k]=(48+(m%10)) as u8;m=m/10;k=k+1} var i: i64=0; while i<k{b[i]=t[k-1-i];i=i+1} sys_write(fd,b,k); return 0 }
24// init scale 13107 (~[-1,1]): big enough that attention logits spread (non-uniform softmax) -> real Wq gradients
25// (RMSNorm makes the INPUT unit-scale regardless, so the logit spread comes from the Wq/Wk magnitude, not X).
26func lm_init(arr: *i64, n: i64, seed: i64) -> i64 { var i: i64=0; while i<n { arr[i] = (((i*7 + seed*13 + 1) % 11) - 5) * 13107; i=i+1 } return 0 }
27
28// ---- multi-head attention forward; W=[Wq,Wk,Wv,Wo]; leaves[0..3]; returns out node ----
29func mha_fwd(tape: *i64, vals: *i64, st: *i64, W: *i64, X: *i64, T: i64, dm: i64, hd: i64, H: i64, scale: i64, leaves: *i64) -> i64 {
30 let Wq: *i64 = W[0] as *i64; let Wk: *i64 = W[1] as *i64; let Wv: *i64 = W[2] as *i64; let Wo: *i64 = W[3] as *i64
31 st[0]=0; st[1]=0
32 let nX: i64 = nfa_leaf(tape,vals,st,T,dm,X,0)
33 let nWq: i64 = nfa_leaf(tape,vals,st,dm,dm,Wq,0)
34 let nWk: i64 = nfa_leaf(tape,vals,st,dm,dm,Wk,0)
35 let nWv: i64 = nfa_leaf(tape,vals,st,dm,dm,Wv,0)
36 let nWo: i64 = nfa_leaf(tape,vals,st,dm,dm,Wo,0)
37 let nXn: i64 = nfa_rmsnorm_rows(tape,vals,st,nX)
38 let nQ: i64 = nfa_matmul(tape,vals,st,nXn,nWq)
39 let nK: i64 = nfa_matmul(tape,vals,st,nXn,nWk)
40 let nV: i64 = nfa_matmul(tape,vals,st,nXn,nWv)
41 var Oacc: i64 = 0 - 1
42 var hh: i64 = 0
43 while hh < H {
44 let nQh: i64 = nfa_slice_cols(tape,vals,st,nQ,hh*hd,hd)
45 let nKh: i64 = nfa_slice_cols(tape,vals,st,nK,hh*hd,hd)
46 let nVh: i64 = nfa_slice_cols(tape,vals,st,nV,hh*hd,hd)
47 let nQr: i64 = nfa_rope(tape,vals,st,nQh)
48 let nKr: i64 = nfa_rope(tape,vals,st,nKh)
49 let nS: i64 = nfa_matmul_nt(tape,vals,st,nQr,nKr)
50 let nSs: i64 = nfa_cmul(tape,vals,st,nS,scale)
51 let nA: i64 = nfa_softmax_rows(tape,vals,st,nSs,1)
52 let nOh: i64 = nfa_matmul(tape,vals,st,nA,nVh)
53 if Oacc < 0 { Oacc = nOh } else { Oacc = nfa_concat_cols(tape,vals,st,Oacc,nOh) }
54 hh = hh + 1
55 }
56 let nOp: i64 = nfa_matmul(tape,vals,st,Oacc,nWo)
57 let nOut: i64 = nfa_vadd(tape,vals,st,nX,nOp)
58 leaves[0]=nWq; leaves[1]=nWk; leaves[2]=nWv; leaves[3]=nWo
59 return nOut
60}
61func mha_loss(tape: *i64, vals: *i64, st: *i64, W: *i64, X: *i64, Tg: *i64, T: i64, dm: i64, hd: i64, H: i64, scale: i64, leaves: *i64) -> i64 {
62 let nOut: i64 = mha_fwd(tape,vals,st,W,X,T,dm,hd,H,scale,leaves)
63 let nt: i64 = nfa_leaf(tape,vals,st,T,dm,Tg,0)
64 return nfa_mse(tape,vals,st,nOut,nt)
65}
66func mha_lossval(tape: *i64, vals: *i64, st: *i64, W: *i64, X: *i64, Tg: *i64, T: i64, dm: i64, hd: i64, H: i64, scale: i64) -> i64 {
67 let lv: *i64 = sys_mmap(4*8) as *i64
68 let loss: i64 = mha_loss(tape,vals,st,W,X,Tg,T,dm,hd,H,scale,lv)
69 return nfa_val(tape,vals,loss,0)
70}
71func mha_out(tape: *i64, vals: *i64, st: *i64, W: *i64, X: *i64, T: i64, dm: i64, hd: i64, H: i64, scale: i64, outv: *i64) -> i64 {
72 let lv: *i64 = sys_mmap(4*8) as *i64
73 let nOut: i64 = mha_fwd(tape,vals,st,W,X,T,dm,hd,H,scale,lv)
74 var i: i64=0
75 while i<T*dm { outv[i]=nfa_val(tape,vals,nOut,i); i=i+1 }
76 return 0
77}
78
79func main() -> i64 {
80 g_puts("nx_nofloat_multihead gate (H-head causal attention backprops + trains, PURE INTEGER Q16)\n" as *u8)
81 var pass: i64=0; var total: i64=0
82 let tape: *i64 = sys_mmap(512*7*8) as *i64
83 let vals: *i64 = sys_mmap(8192*8) as *i64
84 let grads: *i64 = sys_mmap(8192*8) as *i64
85 let st: *i64 = sys_mmap(2*8) as *i64
86 let h: i64 = 512; let floor_q: i64 = 4096
87
88 // ---- A1: slice_cols gradcheck (X[2,4], slice c0=1,w=2) ----
89 let sX: *i64 = sys_mmap(8*8) as *i64; lm_init(sX,8,3)
90 let sT: *i64 = sys_mmap(4*8) as *i64; sT[0]=6554; sT[1]=0-13107; sT[2]=19661; sT[3]=3277
91 var s_ok: i64=1; var s_worst: i64=0
92 st[0]=0; st[1]=0
93 let snX: i64 = nfa_leaf(tape,vals,st,2,4,sX,0)
94 let snS: i64 = nfa_slice_cols(tape,vals,st,snX,1,2)
95 let snT: i64 = nfa_leaf(tape,vals,st,2,2,sT,0)
96 let sloss: i64 = nfa_mse(tape,vals,st,snS,snT)
97 nfa_backward(tape,vals,grads,st[0],sloss)
98 let sana: *i64 = sys_mmap(8*8) as *i64
99 var sc: i64=0
100 while sc<8 { sana[sc]=nfa_grad(tape,grads,snX,sc); sc=sc+1 }
101 var si: i64=0
102 while si<8 {
103 let old: i64=sX[si]
104 sX[si]=old+h
105 st[0]=0; st[1]=0; let a1: i64=nfa_leaf(tape,vals,st,2,4,sX,0); let a2: i64=nfa_slice_cols(tape,vals,st,a1,1,2); let a3: i64=nfa_leaf(tape,vals,st,2,2,sT,0); let lp: i64=nfa_val(tape,vals,nfa_mse(tape,vals,st,a2,a3),0)
106 sX[si]=old-h
107 st[0]=0; st[1]=0; let b1: i64=nfa_leaf(tape,vals,st,2,4,sX,0); let b2: i64=nfa_slice_cols(tape,vals,st,b1,1,2); let b3: i64=nfa_leaf(tape,vals,st,2,2,sT,0); let lm2: i64=nfa_val(tape,vals,nfa_mse(tape,vals,st,b2,b3),0)
108 sX[si]=old
109 let fd: i64=((lp-lm2)*Q16)/(2*h); let num: i64=g_abs(fd-sana[si]); var den: i64=g_abs(sana[si]); if den<floor_q{den=floor_q}
110 if num >= ((4096*den)>>16) { s_ok=0 }
111 let rel: i64=(num*1000)/den; if rel>s_worst{s_worst=rel}
112 si=si+1
113 }
114 g_puts(" [measure] slice_cols worst rel grad err = " as *u8); g_pn(s_worst); g_puts(" /1000 (tol=62)\n" as *u8)
115 pass=pass+g_check("A1: slice_cols gradcheck -- column extract backward (scatter) == finite differences" as *u8, s_ok); total=total+1
116
117 // ---- A2: concat_cols gradcheck (a[2,2], b[2,3]) wrt a ----
118 let ca: *i64 = sys_mmap(4*8) as *i64; lm_init(ca,4,4)
119 let cb: *i64 = sys_mmap(6*8) as *i64; lm_init(cb,6,5)
120 let cT: *i64 = sys_mmap(10*8) as *i64; var ct: i64=0; while ct<10 { cT[ct]=(ct-5)*4096; ct=ct+1 }
121 var co_ok: i64=1; var co_worst: i64=0
122 st[0]=0; st[1]=0
123 let cna: i64 = nfa_leaf(tape,vals,st,2,2,ca,0)
124 let cnb: i64 = nfa_leaf(tape,vals,st,2,3,cb,0)
125 let cnc: i64 = nfa_concat_cols(tape,vals,st,cna,cnb)
126 let cnt: i64 = nfa_leaf(tape,vals,st,2,5,cT,0)
127 let closs: i64 = nfa_mse(tape,vals,st,cnc,cnt)
128 nfa_backward(tape,vals,grads,st[0],closs)
129 let cana: *i64 = sys_mmap(4*8) as *i64
130 var cc: i64=0
131 while cc<4 { cana[cc]=nfa_grad(tape,grads,cna,cc); cc=cc+1 }
132 var ci: i64=0
133 while ci<4 {
134 let old: i64=ca[ci]
135 ca[ci]=old+h
136 st[0]=0; st[1]=0; let a1: i64=nfa_leaf(tape,vals,st,2,2,ca,0); let a2: i64=nfa_leaf(tape,vals,st,2,3,cb,0); let a3: i64=nfa_concat_cols(tape,vals,st,a1,a2); let a4: i64=nfa_leaf(tape,vals,st,2,5,cT,0); let lp: i64=nfa_val(tape,vals,nfa_mse(tape,vals,st,a3,a4),0)
137 ca[ci]=old-h
138 st[0]=0; st[1]=0; let b1: i64=nfa_leaf(tape,vals,st,2,2,ca,0); let b2: i64=nfa_leaf(tape,vals,st,2,3,cb,0); let b3: i64=nfa_concat_cols(tape,vals,st,b1,b2); let b4: i64=nfa_leaf(tape,vals,st,2,5,cT,0); let lm2: i64=nfa_val(tape,vals,nfa_mse(tape,vals,st,b3,b4),0)
139 ca[ci]=old
140 let fd: i64=((lp-lm2)*Q16)/(2*h); let num: i64=g_abs(fd-cana[ci]); var den: i64=g_abs(cana[ci]); if den<floor_q{den=floor_q}
141 if num >= ((4096*den)>>16) { co_ok=0 }
142 let rel: i64=(num*1000)/den; if rel>co_worst{co_worst=rel}
143 ci=ci+1
144 }
145 g_puts(" [measure] concat_cols worst rel grad err = " as *u8); g_pn(co_worst); g_puts(" /1000 (tol=62)\n" as *u8)
146 pass=pass+g_check("A2: concat_cols gradcheck -- concat backward (split) == finite differences" as *u8, co_ok); total=total+1
147
148 // ---- MHA dims + weights ----
149 let T: i64=3; let dm: i64=4; let hd: i64=2; let H: i64=2; let scale: i64=46341 // 1/sqrt(2)
150 let X: *i64 = sys_mmap(T*dm*8) as *i64; lm_init(X,T*dm,1)
151 let Wq: *i64 = sys_mmap(dm*dm*8) as *i64; lm_init(Wq,dm*dm,2)
152 let Wk: *i64 = sys_mmap(dm*dm*8) as *i64; lm_init(Wk,dm*dm,3)
153 let Wv: *i64 = sys_mmap(dm*dm*8) as *i64; lm_init(Wv,dm*dm,4)
154 let Wo: *i64 = sys_mmap(dm*dm*8) as *i64; lm_init(Wo,dm*dm,5)
155 let W: *i64 = sys_mmap(4*8) as *i64; W[0]=Wq as i64; W[1]=Wk as i64; W[2]=Wv as i64; W[3]=Wo as i64
156 let Tg: *i64 = sys_mmap(T*dm*8) as *i64; var tgi: i64=0; while tgi<T*dm { Tg[tgi]=((tgi%7)-3)*4096; tgi=tgi+1 }
157 let leaves: *i64 = sys_mmap(4*8) as *i64
158
159 // ---- A3: multi-head gradcheck wrt Wq (through BOTH heads) ----
160 let l3: i64 = mha_loss(tape,vals,st,W,X,Tg,T,dm,hd,H,scale,leaves)
161 nfa_backward(tape,vals,grads,st[0],l3)
162 let nWq: i64 = leaves[0]
163 let a3ana: *i64 = sys_mmap(16*8) as *i64
164 var ac: i64=0
165 while ac<dm*dm { a3ana[ac]=nfa_grad(tape,grads,nWq,ac); ac=ac+1 }
166 var a3_ok: i64=1; var a3_worst: i64=0
167 var qi: i64=0
168 while qi<dm*dm {
169 let old: i64=Wq[qi]; Wq[qi]=old+h; let lp: i64=mha_lossval(tape,vals,st,W,X,Tg,T,dm,hd,H,scale); Wq[qi]=old-h; let lm2: i64=mha_lossval(tape,vals,st,W,X,Tg,T,dm,hd,H,scale); Wq[qi]=old
170 let fd: i64=((lp-lm2)*Q16)/(2*h); let num: i64=g_abs(fd-a3ana[qi]); var den: i64=g_abs(a3ana[qi]); if den<floor_q{den=floor_q}
171 if num >= ((16384*den)>>16) { a3_ok=0 } // tol 1/4 (multi-head chain through 2 softmaxes)
172 let rel: i64=(num*1000)/den; if rel>a3_worst{a3_worst=rel}
173 qi=qi+1
174 }
175 g_puts(" [measure] multi-head dL/dWq worst rel grad err = " as *u8); g_pn(a3_worst); g_puts(" /1000 (tol=250)\n" as *u8)
176 pass=pass+g_check("A3: multi-head gradcheck wrt Wq through BOTH heads (slice+RoPE+softmax+A.V+concat+proj)" as *u8, a3_ok); total=total+1
177
178 // ---- D: neg-control teeth (use the LARGEST-magnitude grad component so the test isn't vacuous on a ~0 grad) ----
179 var imax: i64=0; var vmax: i64=g_abs(a3ana[0]); var ii: i64=1
180 while ii<dm*dm { if g_abs(a3ana[ii])>vmax { vmax=g_abs(a3ana[ii]); imax=ii } ii=ii+1 }
181 let dana: i64 = a3ana[imax]
182 let o0: i64=Wq[imax]; Wq[imax]=o0+h; let lpd: i64=mha_lossval(tape,vals,st,W,X,Tg,T,dm,hd,H,scale); Wq[imax]=o0-h; let lmd: i64=mha_lossval(tape,vals,st,W,X,Tg,T,dm,hd,H,scale); Wq[imax]=o0
183 let dfd: i64=((lpd-lmd)*Q16)/(2*h); let dbad: i64=0-dana
184 // scale-free teeth: the FD must be unambiguously CLOSER to the true analytic grad than to the negated one
185 // (works even when grads are tiny, where an absolute floor would swamp the signal). vmax guards vs pure noise.
186 let dgood: i64 = g_abs(dfd - dana); let dneg: i64 = g_abs(dfd - dbad)
187 var caught: i64=0; if vmax > 64 { if dneg > dgood*4 { caught=1 } }
188 pass=pass+g_check("D: neg-control -- FD is >4x closer to the true grad than to the negated one (teeth)" as *u8, caught); total=total+1
189
190 // ---- B: multi-head output trains (fix Wq/Wk/Wv -> attention+values fixed -> Op=Oacc.Wo convex; AdamW) ----
191 let Wot: *i64 = sys_mmap(dm*dm*8) as *i64; lm_init(Wot,dm*dm,9)
192 let Wt2: *i64 = sys_mmap(4*8) as *i64; Wt2[0]=Wq as i64; Wt2[1]=Wk as i64; Wt2[2]=Wv as i64; Wt2[3]=Wot as i64
193 let tgt: *i64 = sys_mmap(T*dm*8) as *i64
194 mha_out(tape,vals,st,Wt2,X,T,dm,hd,H,scale,tgt) // realizable target = mha out with Wo*
195 let Wop: *i64 = sys_mmap(dm*dm*8) as *i64; lm_init(Wop,dm*dm,2) // student Wo (different init)
196 let Wb: *i64 = sys_mmap(4*8) as *i64; Wb[0]=Wq as i64; Wb[1]=Wk as i64; Wb[2]=Wv as i64; Wb[3]=Wop as i64
197 let gWo: *i64 = sys_mmap(16*8) as *i64
198 let mm: *i64 = sys_mmap(16*8) as *i64; let vv: *i64 = sys_mmap(16*8) as *i64
199 var zi: i64=0; while zi<dm*dm { mm[zi]=0; vv[zi]=0; zi=zi+1 }
200 let lvb: *i64 = sys_mmap(4*8) as *i64
201 var lf: i64=0; var ll: i64=0
202 var ep: i64=0
203 while ep < 2000 {
204 let lb: i64 = mha_loss(tape,vals,st,Wb,X,tgt,T,dm,hd,H,scale,lvb)
205 nfa_backward(tape,vals,grads,st[0],lb)
206 if ep==0 { lf=nfa_val(tape,vals,lb,0) }
207 ll=nfa_val(tape,vals,lb,0)
208 let nWoL: i64 = lvb[3]
209 var z: i64=0
210 while z<dm*dm { gWo[z]=nfa_grad(tape,grads,nWoL,z); z=z+1 }
211 nfa_adamw(Wop, gWo, mm, vv, dm*dm, 3277, 58982, 65470, 66, 0, ep+1)
212 ep=ep+1
213 }
214 g_puts(" [measure] multi-head Wo-train loss: start=" as *u8); g_pn(lf); g_puts(" end=" as *u8); g_pn(ll); g_puts("\n" as *u8)
215 var learns: i64=1
216 if ll*5 > lf { learns=0 } // >= 80% loss reduction
217 if lf<=0 { learns=0 }
218 pass=pass+g_check("B: multi-head output LEARNS -- Wo trained with AdamW to a realizable target, loss collapses" as *u8, learns); total=total+1
219
220 // ---- C: bit-exact ----
221 let Wop2: *i64 = sys_mmap(dm*dm*8) as *i64; lm_init(Wop2,dm*dm,2)
222 let Wc: *i64 = sys_mmap(4*8) as *i64; Wc[0]=Wq as i64; Wc[1]=Wk as i64; Wc[2]=Wv as i64; Wc[3]=Wop2 as i64
223 let gWo2: *i64 = sys_mmap(16*8) as *i64
224 let mm2: *i64 = sys_mmap(16*8) as *i64; let vv2: *i64 = sys_mmap(16*8) as *i64
225 zi=0; while zi<dm*dm { mm2[zi]=0; vv2[zi]=0; zi=zi+1 }
226 let lvc: *i64 = sys_mmap(4*8) as *i64
227 var ep2: i64=0
228 while ep2 < 2000 {
229 let lc: i64 = mha_loss(tape,vals,st,Wc,X,tgt,T,dm,hd,H,scale,lvc)
230 nfa_backward(tape,vals,grads,st[0],lc)
231 let nWoL: i64=lvc[3]; var z: i64=0
232 while z<dm*dm { gWo2[z]=nfa_grad(tape,grads,nWoL,z); z=z+1 }
233 nfa_adamw(Wop2, gWo2, mm2, vv2, dm*dm, 3277, 58982, 65470, 66, 0, ep2+1)
234 ep2=ep2+1
235 }
236 var bitexact: i64=1
237 var bz: i64=0
238 while bz<dm*dm { if Wop2[bz]!=Wop[bz] { bitexact=0 } bz=bz+1 }
239 pass=pass+g_check("C: bit-exact -- training the multi-head twice gives IDENTICAL integer Wo (determinism)" as *u8, bitexact); total=total+1
240
241 var okall: i64=0; if pass==total { okall=1 }
242 let logf: i64 = sys_openat_append(MLOG, 420)
243 if logf >= 0 {
244 m_ws(logf,"NOFLOATMULTIHEAD H=2 hd=2 A1_slice=" as *u8); m_wn(logf,s_ok); m_ws(logf," A2_concat=" as *u8); m_wn(logf,co_ok)
245 m_ws(logf," A3_mha_Wq=" as *u8); m_wn(logf,a3_ok); m_ws(logf," D=" as *u8); m_wn(logf,caught); m_ws(logf," B_learns=" as *u8); m_wn(logf,learns); m_ws(logf," C_bitexact=" as *u8); m_wn(logf,bitexact)
246 if okall==1 { m_ws(logf," verdict=GREEN\n" as *u8) } else { m_ws(logf," verdict=RED\n" as *u8) }
247 sys_close(logf)
248 }
249 g_puts("---- nofloat_multihead gate: passed " as *u8); g_pn(pass); g_puts(" / " as *u8); g_pn(total); g_puts(" ----\n" as *u8)
250 if okall==1 { g_puts("verdict=GREEN (multi-head causal attention backprops + trains end-to-end in pure integer Q16)\n" as *u8); sys_exit(0); return 0 }
251 g_puts("verdict=RED\n" as *u8); sys_exit(1); return 1
252}