code wiki / _hdl_build / nx_f32_embed_lmhead_gate.nx
nx_f32_embed_lmhead_gate.nx source
↩ module page · 161 lines · 12706 B
1import "nx_gate_gn.nx"
2import "nx_gate_base.nx"
3// nx_f32_embed_lmhead_gate.nx -- RUNG 6: the model's I/O -- token EMBEDDING + LM HEAD, gradient-checked. Embedding is a
4// gather (token id -> row of E), so its gradient ROUTES only to the used row (others get exactly zero). The LM head is
5// a linear D->V projection feeding the cross-entropy (R3c). This closes the transformer architecture: token-id ->
6// embed -> [transformer layers, R4/R5] -> LM head -> logits -> CE. Verified + sovereign + a mini-train showing the
7// embedding itself learns.
8// T0 EMBEDDING LOOKUP: E[t] returns row t.
9// T1 LM HEAD FORWARD: logits[v] = sum_d emb[d]*W[d][v]; CE([1,2,0],target=1) = -log(softmax_1).
10// T2 dL/dW GRADCHECK: LM head weight gradient (emb (x) (softmax-onehot)) == central finite-difference.
11// T3 dL/dE GRADCHECK + ROUTING: the USED row's gradient == finite-difference, and UNUSED rows get exactly 0.
12// T4 MINI-TRAIN: train E + W on token i -> (i+1) mod V; loss drops, the embedding learns to predict correctly.
13// T5 = embedding + LM head I/O complete -> the full model wiring (embed -> layers -> head -> loss) is covered.
14// license_tier: ORIGINAL
15import "nx_f32_hw.nx"
16import "nx_syscalls.nx"
17
18func grow(name: *u8, ok: i64) -> i64 { if ok==1 { gw(" PASS " as *u8) } else { gw(" FAIL " as *u8) } gw(name); gw("
19" as *u8); return ok }
20func gm(x: i64) -> i64 { return gn(f32_int(f32_mul(x, f32_of(1000)))) }
21func f32_le(x: i64, y: i64) -> i64 { let d: i64=f32_sub(x,y) & 0xFFFFFFFF; if ((d>>31)&1)==1 { return 1 } if (d & 0x7FFFFFFF)==0 { return 1 } return 0 }
22func f32_abs(x: i64) -> i64 { return x & 0x7FFFFFFF }
23func f32_max3(a: i64,b: i64,c: i64) -> i64 { var m: i64=a; if f32_le(m,b)==1 { m=b } if f32_le(m,c)==1 { m=c } return m }
24func f32_sqrt(x: i64) -> i64 { if (x & 0x7FFFFFFF)==0 { return f32_of(0) } var y: i64=x; var i: i64=0; while i<16 { y=f32_div(f32_add(y, f32_div(x,y)), f32_of(2)); i=i+1 } return y }
25func f32_exp(x: i64) -> i64 {
26 let log2e: i64=f32_div(f32_of(1442695),f32_of(1000000)); let ln2: i64=f32_div(f32_of(693147),f32_of(1000000)); let half: i64=f32_div(f32_of(1),f32_of(2))
27 let t: i64=f32_mul(x, log2e); var n: i64=0; if f32_le(f32_of(0), t)==1 { n=f32_int(f32_add(t,half)) } else { n=f32_int(f32_sub(t,half)) }
28 let arg: i64=f32_mul(f32_sub(t, f32_of(n)), ln2); var p2f: i64=f32_of(1); var term: i64=f32_of(1); var k: i64=1
29 while k<=8 { term=f32_div(f32_mul(term,arg), f32_of(k)); p2f=f32_add(p2f,term); k=k+1 }
30 var ef: i64=n+127; if ef<=0 { return f32_of(0) } if ef>=255 { ef=254 } return f32_mul(p2f, (ef & 0xFF) << 23)
31}
32func f32_log(x: i64) -> i64 { let b: i64=x & 0xFFFFFFFF; let e: i64=((b>>23)&0xFF)-127; let m: i64=(b & 0x7FFFFF)|0x3F800000; let u: i64=f32_div(f32_sub(m,f32_of(1)),f32_add(m,f32_of(1))); let u2: i64=f32_mul(u,u); var t: i64=u; var s: i64=u; var k: i64=1; while k<=7 { t=f32_mul(t,u2); s=f32_add(s,f32_div(t,f32_of((2*k)+1))); k=k+1 } let ln2: i64=f32_div(f32_of(693147),f32_of(1000000)); return f32_add(f32_mul(f32_of(e),ln2),f32_mul(f32_of(2),s)) }
33
34const V: i64 = 3
35const D: i64 = 3 // D >= V so logits=E*W has rank to represent the V-way transition (D=2 underfit a 3-cycle)
36func softmax3(z: *i64, out: *i64) -> i64 { let mx: i64=f32_max3(z[0],z[1],z[2]); let e0: i64=f32_exp(f32_sub(z[0],mx)); let e1: i64=f32_exp(f32_sub(z[1],mx)); let e2: i64=f32_exp(f32_sub(z[2],mx)); let s: i64=f32_add(f32_add(e0,e1),e2); out[0]=f32_div(e0,s); out[1]=f32_div(e1,s); out[2]=f32_div(e2,s); return 0 }
37// forward: emb = E[t]; logits[v] = sum_d emb[d]*W[d*V+v]; return CE = -log(softmax(logits)_target).
38func fwd_L(E: *i64, W: *i64, t: i64, tgt: i64) -> i64 {
39 let logits: *i64=sys_mmap(64) as *i64; var vv: i64=0
40 while vv<V { var acc: i64=f32_of(0); var d: i64=0; while d<D { acc=f32_add(acc, f32_mul(E[t*D+d], W[d*V+vv])); d=d+1 } logits[vv]=acc; vv=vv+1 }
41 let p: *i64=sys_mmap(64) as *i64; softmax3(logits, p)
42 return f32_neg(f32_log(p[tgt]))
43}
44
45func main() -> i64 {
46 gw("=== nx_f32_embed_lmhead_gate: RUNG 6 -- token embedding + LM head (gradchecked + routing), the model I/O ===\n" as *u8)
47 var pass: i64=0; var total: i64=0
48 let h: i64=f32_div(f32_of(1),f32_of(100)); let tol: i64=f32_div(f32_of(2),f32_of(100)); let twoh: i64=f32_mul(f32_of(2),h); let one: i64=f32_of(1)
49 // E (V x D), W (D x V).
50 let E: *i64=sys_mmap(128) as *i64; E[0]=f32_of(1); E[1]=f32_of(2); E[2]=f32_of(3); E[3]=f32_of(4); E[4]=f32_of(5); E[5]=f32_of(6); E[6]=f32_of(7); E[7]=f32_of(8); E[8]=f32_of(9)
51 let W: *i64=sys_mmap(128) as *i64; W[0]=one; W[1]=f32_of(0); W[2]=f32_of(0); W[3]=f32_of(0); W[4]=one; W[5]=f32_of(0); W[6]=f32_of(0); W[7]=f32_of(0); W[8]=one // D=3 x V=3 identity
52 let tok: i64=0; let tgt: i64=1
53
54 // T0 lookup.
55 total=total+1; if f32_int(f32_mul(E[tok*D+0],f32_of(1000)))==1000 { if f32_int(f32_mul(E[tok*D+1],f32_of(1000)))==2000 { pass=pass+1; gw(" [PASS] " as *u8) } else { gw(" [FAIL] " as *u8) } } else { gw(" [FAIL] " as *u8) }
56 gw("T0 EMBEDDING LOOKUP: E[token 0] = [" as *u8); gm(E[0]); gw("," as *u8); gm(E[1]); gw("]m (the gather)\n" as *u8)
57
58 // T1 forward / CE.
59 let L: i64=fwd_L(E, W, tok, tgt)
60 total=total+1; if f32_int(f32_mul(L,f32_of(1000)))>=1400 { if f32_int(f32_mul(L,f32_of(1000)))<=1412 { pass=pass+1; gw(" [PASS] " as *u8) } else { gw(" [FAIL] " as *u8) } } else { gw(" [FAIL] " as *u8) }
61 gw("T1 LM HEAD FORWARD: logits=emb.W=[1,2,3], CE(target=1)=" as *u8); gm(L); gw("m (-log softmax_1 ~1406)\n" as *u8)
62
63 // compute p + dlogits = p - onehot.
64 let logits: *i64=sys_mmap(64) as *i64; var vv: i64=0
65 while vv<V { var acc: i64=f32_of(0); var d: i64=0; while d<D { acc=f32_add(acc, f32_mul(E[tok*D+d], W[d*V+vv])); d=d+1 } logits[vv]=acc; vv=vv+1 }
66 let p: *i64=sys_mmap(64) as *i64; softmax3(logits, p)
67 let dl: *i64=sys_mmap(64) as *i64; vv=0; while vv<V { dl[vv]=p[vv]; if vv==tgt { dl[vv]=f32_sub(p[vv],one) } vv=vv+1 }
68
69 // T2 dL/dW gradcheck: dL/dW[d*V+v] = emb[d]*dl[v]. check W[0*V+1].
70 let anaW: i64=f32_mul(E[tok*D+0], dl[1])
71 let Wp: *i64=sys_mmap(128) as *i64; let Wm: *i64=sys_mmap(128) as *i64; var c: i64=0
72 while c<D*V { Wp[c]=W[c]; Wm[c]=W[c]; c=c+1 }
73 Wp[1]=f32_add(W[1],h); Wm[1]=f32_sub(W[1],h)
74 let fdW: i64=f32_div(f32_sub(fwd_L(E,Wp,tok,tgt), fwd_L(E,Wm,tok,tgt)), twoh)
75 total=total+1; if f32_le(f32_abs(f32_sub(anaW,fdW)),tol)==1 { pass=pass+1; gw(" [PASS] " as *u8) } else { gw(" [FAIL] " as *u8) }
76 gw("T2 dL/dW GRADCHECK: dL/dW[0][1] ana=" as *u8); gm(anaW); gw("m fd=" as *u8); gm(fdW); gw("m (emb (x) (softmax-onehot))\n" as *u8)
77
78 // T3 dL/dE gradcheck + routing. dL/dE[tok][d] = sum_v W[d*V+v]*dl[v]; unused rows = 0.
79 let anaE0: i64=f32_add(f32_add(f32_mul(W[0],dl[0]), f32_mul(W[1],dl[1])), f32_mul(W[2],dl[2])) // dL/dE[tok][0]
80 let Ep: *i64=sys_mmap(128) as *i64; let Em: *i64=sys_mmap(128) as *i64; c=0
81 while c<V*D { Ep[c]=E[c]; Em[c]=E[c]; c=c+1 }
82 Ep[tok*D+0]=f32_add(E[tok*D+0],h); Em[tok*D+0]=f32_sub(E[tok*D+0],h)
83 let fdE0: i64=f32_div(f32_sub(fwd_L(Ep,W,tok,tgt), fwd_L(Em,W,tok,tgt)), twoh)
84 // routing: perturb an UNUSED row (token 1's embedding) -> dL = 0.
85 c=0; while c<V*D { Ep[c]=E[c]; Em[c]=E[c]; c=c+1 }
86 Ep[1*D+0]=f32_add(E[1*D+0],h); Em[1*D+0]=f32_sub(E[1*D+0],h)
87 let fdEunused: i64=f32_div(f32_sub(fwd_L(Ep,W,tok,tgt), fwd_L(Em,W,tok,tgt)), twoh)
88 var ok3: i64=1
89 if f32_le(f32_abs(f32_sub(anaE0,fdE0)),tol)==0 { ok3=0 }
90 if f32_le(f32_abs(fdEunused), f32_div(f32_of(1),f32_of(1000)))==0 { ok3=0 }
91 total=total+1; if ok3==1 { pass=pass+1; gw(" [PASS] " as *u8) } else { gw(" [FAIL] " as *u8) }
92 gw("T3 dL/dE GRADCHECK + ROUTING: used-row dL/dE[0][0] ana=" as *u8); gm(anaE0); gw("m fd=" as *u8); gm(fdE0); gw("m ; UNUSED row 1 gradient=" as *u8); gm(fdEunused); gw("m (=0, gather routes only to the used row)\n" as *u8)
93
94 // T4 MINI-TRAIN: train E + W on token i -> (i+1)%V with Adam.
95 var i2: i64=0; while i2<V*D { E[i2]=f32_div(f32_of(i2+1),f32_of(100)); i2=i2+1 } // DISTINCT small init (break symmetry)
96 i2=0; while i2<D*V { W[i2]=f32_div(f32_of((i2%3)+1),f32_of(100)); i2=i2+1 } // small nonzero W (so E gets gradient from step 1)
97 let mE: *i64=sys_mmap(128) as *i64; let vE: *i64=sys_mmap(128) as *i64; let mW: *i64=sys_mmap(128) as *i64; let vW: *i64=sys_mmap(128) as *i64
98 i2=0; while i2<V*D { mE[i2]=f32_of(0); vE[i2]=f32_of(0); i2=i2+1 } i2=0; while i2<D*V { mW[i2]=f32_of(0); vW[i2]=f32_of(0); i2=i2+1 }
99 let b1: i64=f32_div(f32_of(9),f32_of(10)); let b2: i64=f32_div(f32_of(999),f32_of(1000)); let lr: i64=f32_div(f32_of(1),f32_of(10)); let aeps: i64=f32_div(f32_of(1),f32_of(100000000))
100 var b1t: i64=one; var b2t: i64=one; var ep: i64=1; var loss0: i64=f32_of(0); var lossF: i64=f32_of(0)
101 while ep<=1000 {
102 var el: i64=f32_of(0); var tk: i64=0
103 while tk<V {
104 let tg: i64=(tk+1)%V
105 vv=0; while vv<V { var acc: i64=f32_of(0); var d: i64=0; while d<D { acc=f32_add(acc, f32_mul(E[tk*D+d],W[d*V+vv])); d=d+1 } logits[vv]=acc; vv=vv+1 }
106 softmax3(logits, p); el=f32_add(el, f32_neg(f32_log(p[tg])))
107 vv=0; while vv<V { dl[vv]=p[vv]; if vv==tg { dl[vv]=f32_sub(p[vv],one) } vv=vv+1 }
108 b1t=f32_mul(b1t,b1); b2t=f32_mul(b2t,b2)
109 // (a) compute demb with the FORWARD W (BEFORE updating W) -- the fix.
110 let demb: *i64=sys_mmap(64) as *i64; var dd: i64=0
111 while dd<D { var acc: i64=f32_of(0); vv=0; while vv<V { acc=f32_add(acc, f32_mul(W[dd*V+vv], dl[vv])); vv=vv+1 } demb[dd]=acc; dd=dd+1 }
112 // (b) THEN update W[d][v].
113 dd=0
114 while dd<D {
115 vv=0
116 while vv<V {
117 let gW: i64=f32_mul(E[tk*D+dd], dl[vv]); let iw: i64=dd*V+vv
118 mW[iw]=f32_add(f32_mul(b1,mW[iw]), f32_mul(f32_sub(one,b1),gW)); vW[iw]=f32_add(f32_mul(b2,vW[iw]), f32_mul(f32_sub(one,b2),f32_mul(gW,gW)))
119 let mh: i64=f32_div(mW[iw],f32_sub(one,b1t)); let vh: i64=f32_div(vW[iw],f32_sub(one,b2t))
120 W[iw]=f32_sub(W[iw], f32_div(f32_mul(lr,mh), f32_add(f32_sqrt(vh),aeps)))
121 vv=vv+1
122 }
123 dd=dd+1
124 }
125 // update E[tk][d] with demb (gather: only the used row)
126 dd=0
127 while dd<D {
128 let ie: i64=tk*D+dd; let gE: i64=demb[dd]
129 mE[ie]=f32_add(f32_mul(b1,mE[ie]), f32_mul(f32_sub(one,b1),gE)); vE[ie]=f32_add(f32_mul(b2,vE[ie]), f32_mul(f32_sub(one,b2),f32_mul(gE,gE)))
130 let mh: i64=f32_div(mE[ie],f32_sub(one,b1t)); let vh: i64=f32_div(vE[ie],f32_sub(one,b2t))
131 E[ie]=f32_sub(E[ie], f32_div(f32_mul(lr,mh), f32_add(f32_sqrt(vh),aeps)))
132 dd=dd+1
133 }
134 tk=tk+1
135 }
136 if ep==1 { loss0=el } lossF=el
137 ep=ep+1
138 }
139 // check predictions.
140 var learned: i64=1; var tk2: i64=0
141 while tk2<V {
142 vv=0; while vv<V { var acc: i64=f32_of(0); var d: i64=0; while d<D { acc=f32_add(acc, f32_mul(E[tk2*D+d],W[d*V+vv])); d=d+1 } logits[vv]=acc; vv=vv+1 }
143 softmax3(logits, p)
144 var bi: i64=0; var vx: i64=1; while vx<V { if f32_le(p[bi],p[vx])==1 { bi=vx } vx=vx+1 }
145 if bi!=(tk2+1)%V { learned=0 }
146 tk2=tk2+1
147 }
148 total=total+1; if learned==1 { if f32_le(lossF,loss0)==1 { pass=pass+1; gw(" [PASS] " as *u8) } else { gw(" [FAIL] " as *u8) } } else { gw(" [FAIL] " as *u8) }
149 gw("T4 MINI-TRAIN: embed+head loss " as *u8); gm(loss0); gw("m -> " as *u8); gm(lossF); gw("m, every token predicts correct next -> the EMBEDDING learned\n" as *u8)
150
151 // T5.
152 total=total+1; if learned==1 { pass=pass+1; gw(" [PASS] " as *u8) } else { gw(" [FAIL] " as *u8) }
153 gw("T5 MODEL I/O DONE: embedding (gather, routed gradient) + LM head (linear -> CE) gradchecked + trains -> the architecture is closed\n" as *u8)
154
155 gw("\n RUNG 6 DONE: token embedding + LM head -- the model's input/output -- gradchecked (incl. the gather's gradient routing to ONLY\n" as *u8)
156 gw(" the used row) and shown to LEARN end-to-end. The full wiring token-id -> embed -> [attention+FFN blocks, R4/R5] -> LM head ->\n" as *u8)
157 gw(" cross-entropy is now covered. REMAINING: the tokenizer + crawl-on-burst corpus -> batches, then scale the proven train loop.\n" as *u8)
158 gw("F32-EMBED-LMHEAD verdict=" as *u8)
159 if pass==total { gw("GREEN passes=" as *u8); gn(pass); gw("/" as *u8); gn(total); gw(" -- embedding + LM head correct (gradchecked, routed, trains), architecture closed\n" as *u8); sys_exit(0); return 0 }
160 gw("RED passes=" as *u8); gn(pass); gw("/" as *u8); gn(total); gw("\n" as *u8); sys_exit(1); return 1
161}