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