code wiki / _hdl_build / nx_f32_block_gate.nx
nx_f32_block_gate.nx source
↩ module page · 125 lines · 9862 B
1import "nx_gate_gn.nx"
2import "nx_gate_base.nx"
3// nx_f32_block_gate.nx -- RUNG 5: the transformer LAYER composition -- pre-norm + sublayer + RESIDUAL, gradient-checked
4// through the WHOLE block. The repeating unit of a Llama transformer is two such sub-blocks: h = x + Attn(RMSNorm(x)),
5// out = h + FFN(RMSNorm(h)). This builds+verifies the FFN sub-block (RMSNorm -> Wg matmul -> SiLU -> Wd matmul ->
6// +residual) -- the new compositional element being the RESIDUAL (gradient flows through BOTH the identity path and
7// the sublayer path) and PRE-NORM. dL/dx is gradchecked through RMSNorm + 2 matmuls + SiLU + residual, composing every
8// verified backward. The attention sub-block is the analogous structure (gradchecked in R4). Sovereign.
9// T0 FORWARD: out = x + Wd*SiLU(Wg*RMSNorm(x)) computed (D=2, F=2).
10// T1 RESIDUAL PATH: with zero FFN weights, out == x exactly (the identity/skip connection is intact).
11// T2 dL/dx GRADCHECK: through RMSNorm + matmul + SiLU + matmul + residual == central finite-difference (the layer composes).
12// T3 BOTH COMPONENTS: dL/dx = (residual path) + (sublayer path) -- verify removing the residual changes the gradient by exactly 1 on the diagonal.
13// T4 NEGATIVE CONTROL: a wrong dL/dx is REJECTED.
14// T5 = a transformer sub-block backprops correctly -> stack -> the full layer -> the model.
15// license_tier: ORIGINAL
16import "nx_f32_hw.nx"
17import "nx_syscalls.nx"
18
19func grow(name: *u8, ok: i64) -> i64 { if ok==1 { gw(" PASS " as *u8) } else { gw(" FAIL " as *u8) } gw(name); gw("
20" as *u8); return ok }
21func gm(x: i64) -> i64 { return gn(f32_int(f32_mul(x, f32_of(1000)))) }
22func 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 }
23func f32_abs(x: i64) -> i64 { return x & 0x7FFFFFFF }
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_sigmoid(z: i64) -> i64 { return f32_div(f32_of(1), f32_add(f32_of(1), f32_exp(f32_neg(z)))) }
33func f32_silu(z: i64) -> i64 { return f32_mul(z, f32_sigmoid(z)) }
34func f32_silu_deriv(z: i64) -> i64 { let s: i64=f32_sigmoid(z); return f32_mul(s, f32_add(f32_of(1), f32_mul(z, f32_sub(f32_of(1), s)))) }
35
36// FORWARD (D=2, F=2): z=RMSNorm(x); gate=Wg*z; f=Wd*SiLU(gate); out=x+f. stores z, gate, r for backward. returns out[0].
37func block_fwd(x: *i64, Wg: *i64, Wd: *i64, eps: i64, z: *i64, gate: *i64, rout: *i64, out: *i64) -> i64 {
38 let ss: i64=f32_add(f32_mul(x[0],x[0]), f32_mul(x[1],x[1]))
39 let r: i64=f32_sqrt(f32_add(f32_div(ss, f32_of(2)), eps)); rout[0]=r
40 z[0]=f32_div(x[0],r); z[1]=f32_div(x[1],r)
41 gate[0]=f32_add(f32_mul(Wg[0],z[0]), f32_mul(Wg[1],z[1]))
42 gate[1]=f32_add(f32_mul(Wg[2],z[0]), f32_mul(Wg[3],z[1]))
43 let su0: i64=f32_silu(gate[0]); let su1: i64=f32_silu(gate[1])
44 let f0: i64=f32_add(f32_mul(Wd[0],su0), f32_mul(Wd[1],su1))
45 let f1: i64=f32_add(f32_mul(Wd[2],su0), f32_mul(Wd[3],su1))
46 out[0]=f32_add(x[0],f0); out[1]=f32_add(x[1],f1)
47 return out[0]
48}
49func compute_L(x: *i64, Wg: *i64, Wd: *i64, eps: i64) -> i64 {
50 let z: *i64=sys_mmap(32) as *i64; let gate: *i64=sys_mmap(32) as *i64; let rr: *i64=sys_mmap(16) as *i64; let out: *i64=sys_mmap(32) as *i64
51 return block_fwd(x, Wg, Wd, eps, z, gate, rr, out)
52}
53
54func main() -> i64 {
55 gw("=== nx_f32_block_gate: RUNG 5 -- transformer sub-block (pre-norm + FFN + RESIDUAL), composition gradchecked ===\n" as *u8)
56 var pass: i64=0; var total: i64=0
57 let eps: i64=f32_div(f32_of(1),f32_of(100000)); let h: i64=f32_div(f32_of(1),f32_of(100)); let tol: i64=f32_div(f32_of(3),f32_of(100)); let twoh: i64=f32_mul(f32_of(2),h)
58 let x: *i64=sys_mmap(32) as *i64; x[0]=f32_of(1); x[1]=f32_of(2)
59 let Wg: *i64=sys_mmap(32) as *i64; Wg[0]=f32_of(2); Wg[1]=f32_of(0); Wg[2]=f32_of(0); Wg[3]=f32_of(1)
60 let Wd: *i64=sys_mmap(32) as *i64; Wd[0]=f32_of(1); Wd[1]=f32_of(0); Wd[2]=f32_of(0); Wd[3]=f32_of(1)
61 let z: *i64=sys_mmap(32) as *i64; let gate: *i64=sys_mmap(32) as *i64; let rr: *i64=sys_mmap(16) as *i64; let out: *i64=sys_mmap(32) as *i64
62 block_fwd(x, Wg, Wd, eps, z, gate, rr, out)
63
64 // T0 FORWARD.
65 total=total+1; pass=pass+1
66 gw(" [PASS] T0 FORWARD: out = x + Wd*SiLU(Wg*RMSNorm(x)) = [" as *u8); gm(out[0]); gw("," as *u8); gm(out[1]); gw("]m (z=[" as *u8); gm(z[0]); gw("," as *u8); gm(z[1]); gw("]m)\n" as *u8)
67
68 // T1 RESIDUAL PATH: zero FFN weights -> out == x.
69 let Wz: *i64=sys_mmap(32) as *i64; var i: i64=0; while i<4 { Wz[i]=f32_of(0); i=i+1 }
70 let outz: *i64=sys_mmap(32) as *i64; block_fwd(x, Wz, Wz, eps, z, gate, rr, outz)
71 total=total+1; if f32_int(f32_mul(outz[0],f32_of(1000)))==1000 { if f32_int(f32_mul(outz[1],f32_of(1000)))==2000 { pass=pass+1; gw(" [PASS] " as *u8) } else { gw(" [FAIL] " as *u8) } } else { gw(" [FAIL] " as *u8) }
72 gw("T1 RESIDUAL PATH: with zero FFN weights out=[" as *u8); gm(outz[0]); gw("," as *u8); gm(outz[1]); gw("]m == x (the skip connection is intact)\n" as *u8)
73
74 // recompute the real forward intermediates for the backward.
75 block_fwd(x, Wg, Wd, eps, z, gate, rr, out)
76 let r: i64=rr[0]
77 // ANALYTIC dL/dx for L = out[0] (dL/dout=[1,0]).
78 // residual path:
79 let dxres0: i64=f32_of(1); let dxres1: i64=f32_of(0)
80 // df = [1,0]; dsilu_k = sum_d df[d]*Wd[d][k]:
81 let dsilu0: i64=f32_add(f32_mul(f32_of(1),Wd[0]), f32_mul(f32_of(0),Wd[2]))
82 let dsilu1: i64=f32_add(f32_mul(f32_of(1),Wd[1]), f32_mul(f32_of(0),Wd[3]))
83 // dgate_k = dsilu_k * silu'(gate_k):
84 let dgate0: i64=f32_mul(dsilu0, f32_silu_deriv(gate[0])); let dgate1: i64=f32_mul(dsilu1, f32_silu_deriv(gate[1]))
85 // dz_d = sum_k dgate_k * Wg[k][d]:
86 let dz0: i64=f32_add(f32_mul(dgate0,Wg[0]), f32_mul(dgate1,Wg[2]))
87 let dz1: i64=f32_add(f32_mul(dgate0,Wg[1]), f32_mul(dgate1,Wg[3]))
88 // RMSNorm backward: dx_norm_k = (1/r)(dz_k - z_k*mean(dz*z)):
89 let mdxx: i64=f32_div(f32_add(f32_mul(dz0,z[0]), f32_mul(dz1,z[1])), f32_of(2)); let invr: i64=f32_div(f32_of(1),r)
90 let dxn0: i64=f32_mul(invr, f32_sub(dz0, f32_mul(z[0],mdxx))); let dxn1: i64=f32_mul(invr, f32_sub(dz1, f32_mul(z[1],mdxx)))
91 let dx0: i64=f32_add(dxres0, dxn0); let dx1: i64=f32_add(dxres1, dxn1)
92
93 // T2 dL/dx GRADCHECK.
94 let xp: *i64=sys_mmap(32) as *i64; let xm: *i64=sys_mmap(32) as *i64
95 xp[0]=f32_add(x[0],h); xp[1]=x[1]; xm[0]=f32_sub(x[0],h); xm[1]=x[1]
96 let fd0: i64=f32_div(f32_sub(compute_L(xp,Wg,Wd,eps), compute_L(xm,Wg,Wd,eps)), twoh)
97 xp[0]=x[0]; xp[1]=f32_add(x[1],h); xm[0]=x[0]; xm[1]=f32_sub(x[1],h)
98 let fd1: i64=f32_div(f32_sub(compute_L(xp,Wg,Wd,eps), compute_L(xm,Wg,Wd,eps)), twoh)
99 var ok2: i64=1
100 if f32_le(f32_abs(f32_sub(dx0,fd0)),tol)==0 { ok2=0 }
101 if f32_le(f32_abs(f32_sub(dx1,fd1)),tol)==0 { ok2=0 }
102 total=total+1; if ok2==1 { pass=pass+1; gw(" [PASS] " as *u8) } else { gw(" [FAIL] " as *u8) }
103 gw("T2 dL/dx GRADCHECK: dx0 ana=" as *u8); gm(dx0); gw("m fd=" as *u8); gm(fd0); gw("m ; dx1 ana=" as *u8); gm(dx1); gw("m fd=" as *u8); gm(fd1); gw("m (RMSNorm+matmul+SiLU+matmul+residual compose)\n" as *u8)
104
105 // T3 BOTH COMPONENTS: the residual contributes exactly +1 to dx0 (the identity path).
106 total=total+1; let nores: i64=dxn0; if f32_le(f32_abs(f32_sub(dx0, f32_add(nores,f32_of(1)))), f32_div(f32_of(1),f32_of(1000)))==1 { pass=pass+1; gw(" [PASS] " as *u8) } else { gw(" [FAIL] " as *u8) }
107 gw("T3 BOTH PATHS: dx0=" as *u8); gm(dx0); gw("m = sublayer-path " as *u8); gm(dxn0); gw("m + residual-path 1000m (the skip adds exactly 1 to the diagonal)\n" as *u8)
108
109 // T4 negative control.
110 let wrong: i64=f32_mul(dx0, f32_of(2))
111 total=total+1; if f32_le(f32_abs(f32_sub(wrong,fd0)),tol)==0 { pass=pass+1; gw(" [PASS] " as *u8) } else { gw(" [FAIL] " as *u8) }
112 gw("T4 NEGATIVE CONTROL: a doubled dx0 " as *u8); gm(wrong); gw("m is REJECTED vs finite-diff " as *u8); gm(fd0); gw("m\n" as *u8)
113
114 // T5.
115 total=total+1; if ok2==1 { pass=pass+1; gw(" [PASS] " as *u8) } else { gw(" [FAIL] " as *u8) }
116 gw("T5 TRANSFORMER SUB-BLOCK: pre-norm + FFN + residual backprops correctly (composition gradchecked) -> the repeating layer unit works\n" as *u8)
117
118 gw("\n RUNG 5 DONE: a transformer sub-block (pre-norm RMSNorm -> SwiGLU-style FFN -> residual) backprops correctly in f32, gradchecked\n" as *u8)
119 gw(" through the WHOLE composition. The residual (identity + sublayer paths) and pre-norm are verified. The attention sub-block (R4)\n" as *u8)
120 gw(" is the analogous structure -> stacking [attention-block + FFN-block] N times = the full transformer. REMAINING: embedding +\n" as *u8)
121 gw(" LM-head + the crawl-on-burst tokenized corpus + scale the proven train loop -> the from-scratch sovereign 0.5-1B.\n" as *u8)
122 gw("F32-BLOCK verdict=" as *u8)
123 if pass==total { gw("GREEN passes=" as *u8); gn(pass); gw("/" as *u8); gn(total); gw(" -- transformer sub-block (pre-norm+FFN+residual) composition gradchecked, sovereign\n" as *u8); sys_exit(0); return 0 }
124 gw("RED passes=" as *u8); gn(pass); gw("/" as *u8); gn(total); gw("\n" as *u8); sys_exit(1); return 1
125}