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}