code wiki / _hdl_build / nx_f32_softmax_gate.nx
nx_f32_softmax_gate.nx source
↩ module page · 123 lines · 8359 B
1import "nx_gate_gn.nx"
2import "nx_gate_base.nx"
3// nx_f32_softmax_gate.nx -- RUNG 2 of the from-scratch 0.5B training loop: the transcendentals (f32_exp via Taylor,
4// f32_sqrt via Newton) + SOFTMAX forward & backward, gradient-checked. Softmax is the attention core; its backward
5// (the Jacobian-vector product) is what backprop through attention needs. Builds on the R1 f32 autograd + nx_f32_hw.
6// f32_sqrt is tested here too (the layernorm rung, R2b, uses it). Same proof discipline as R1: analytic gradient ==
7// central finite difference, with a negative control.
8// T0 f32_exp: exp(0)=1, exp(1)=e=2.718, exp(2)=7.389 (Taylor, accurate on the logit range).
9// T1 f32_sqrt: sqrt(4)=2, sqrt(2)=1.414, sqrt(9)=3 (Newton).
10// T2 SOFTMAX FORWARD: softmax([1,2,3]) sums to 1.0 (max-subtracted for stability).
11// T3 SOFTMAX BACKWARD GRADCHECK: dL/dx (L=softmax_0) autograd == central finite-difference, all components.
12// T4 NEGATIVE CONTROL: a doubled (wrong) gradient is REJECTED.
13// T5 = softmax backward correct -> attention backprop unblocked; next R2b: f32_sqrt -> layernorm backward.
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)))) } // f32 -> milliunits (honest display)
21
22func f32_abs(x: i64) -> i64 { return x & 0x7FFFFFFF }
23func 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 }
24func f32_max2(a: i64, b: i64) -> i64 { if f32_le(a,b)==1 { return b } return a }
25
26// f32_exp via Taylor: sum_{k=0..N} x^k/k! (N=16; accurate for |x| <~ 6, and softmax feeds it x-max <= 0).
27func f32_exp(x: i64) -> i64 {
28 var sum: i64=f32_of(1); var term: i64=f32_of(1); var k: i64=1
29 while k<=16 { term=f32_div(f32_mul(term,x), f32_of(k)); sum=f32_add(sum,term); k=k+1 }
30 return sum
31}
32// f32_sqrt via Newton: y <- (y + x/y)/2 from y0=x (>0); 14 iters covers the demo range.
33func f32_sqrt(x: i64) -> i64 {
34 if (x & 0x7FFFFFFF)==0 { return f32_of(0) }
35 var y: i64=x; var i: i64=0
36 while i<14 { y=f32_div(f32_add(y, f32_div(x,y)), f32_of(2)); i=i+1 }
37 return y
38}
39
40// softmax(x[0..n]) -> out, max-subtracted for stability.
41func softmax(x: *i64, n: i64, out: *i64) -> i64 {
42 var mx: i64=x[0]; var i: i64=1; while i<n { mx=f32_max2(mx, x[i]); i=i+1 }
43 var sum: i64=f32_of(0); i=0; while i<n { out[i]=f32_exp(f32_sub(x[i], mx)); sum=f32_add(sum, out[i]); i=i+1 }
44 i=0; while i<n { out[i]=f32_div(out[i], sum); i=i+1 }
45 return 0
46}
47// softmax backward (JVP): given upstream g, dL/dx_j = s_j*(g_j - sum_i g_i s_i).
48func softmax_backward(s: *i64, g: *i64, n: i64, dx: *i64) -> i64 {
49 var dot: i64=f32_of(0); var i: i64=0; while i<n { dot=f32_add(dot, f32_mul(g[i], s[i])); i=i+1 }
50 i=0; while i<n { dx[i]=f32_mul(s[i], f32_sub(g[i], dot)); i=i+1 }
51 return 0
52}
53// L(x) = sum_i w_i * softmax(x)_i (forward, for the finite-difference probe).
54func compute_L(x: *i64, w: *i64, n: i64) -> i64 {
55 let s: *i64=sys_mmap(64) as *i64; softmax(x, n, s)
56 var L: i64=f32_of(0); var i: i64=0; while i<n { L=f32_add(L, f32_mul(w[i], s[i])); i=i+1 }
57 return L
58}
59
60func main() -> i64 {
61 gw("=== nx_f32_softmax_gate: RUNG 2 -- f32_exp + f32_sqrt + softmax forward/backward (gradchecked) ===\n" as *u8)
62 var pass: i64=0; var total: i64=0
63 let eps: i64=f32_div(f32_of(1), f32_of(100)) // 0.01
64 let tol: i64=f32_div(f32_of(2), f32_of(100)) // 0.02
65
66 // T0 f32_exp.
67 let e0: i64=f32_int(f32_mul(f32_exp(f32_of(0)),f32_of(1000)))
68 let e1: i64=f32_int(f32_mul(f32_exp(f32_of(1)),f32_of(1000)))
69 let e2: i64=f32_int(f32_mul(f32_exp(f32_of(2)),f32_of(1000)))
70 total=total+1; if e0==1000 { if e1>=2710 { if e1<=2725 { if e2>=7380 { if e2<=7398 { pass=pass+1; gw(" [PASS] " as *u8) } else { gw(" [FAIL] " as *u8) } } else { gw(" [FAIL] " as *u8) } } else { gw(" [FAIL] " as *u8) } } else { gw(" [FAIL] " as *u8) } } else { gw(" [FAIL] " as *u8) }
71 gw("T0 f32_exp (milliunits): exp(0)=" as *u8); gn(e0); gw(" exp(1)=" as *u8); gn(e1); gw(" exp(2)=" as *u8); gn(e2); gw(" (expect 1000/2718/7389)\n" as *u8)
72
73 // T1 f32_sqrt.
74 let q4: i64=f32_int(f32_mul(f32_sqrt(f32_of(4)),f32_of(1000)))
75 let q2: i64=f32_int(f32_mul(f32_sqrt(f32_of(2)),f32_of(1000)))
76 let q9: i64=f32_int(f32_mul(f32_sqrt(f32_of(9)),f32_of(1000)))
77 total=total+1; if q4>=1998 { if q4<=2002 { if q2>=1412 { if q2<=1416 { if q9>=2998 { if q9<=3002 { pass=pass+1; gw(" [PASS] " as *u8) } else { gw(" [FAIL] " as *u8) } } else { gw(" [FAIL] " as *u8) } } else { gw(" [FAIL] " as *u8) } } else { gw(" [FAIL] " as *u8) } } else { gw(" [FAIL] " as *u8) } } else { gw(" [FAIL] " as *u8) }
78 gw("T1 f32_sqrt (milliunits): sqrt(4)=" as *u8); gn(q4); gw(" sqrt(2)=" as *u8); gn(q2); gw(" sqrt(9)=" as *u8); gn(q9); gw(" (expect 2000/1414/3000)\n" as *u8)
79
80 // T2 SOFTMAX FORWARD sums to 1.
81 let x: *i64=sys_mmap(64) as *i64; x[0]=f32_of(1); x[1]=f32_of(2); x[2]=f32_of(3)
82 let s: *i64=sys_mmap(64) as *i64; softmax(x, 3, s)
83 let ssum: i64=f32_int(f32_mul(f32_add(f32_add(s[0],s[1]),s[2]), f32_of(1000)))
84 total=total+1; if ssum>=999 { if ssum<=1001 { pass=pass+1; gw(" [PASS] " as *u8) } else { gw(" [FAIL] " as *u8) } } else { gw(" [FAIL] " as *u8) }
85 gw("T2 SOFTMAX FORWARD: softmax([1,2,3]) = [" as *u8); gm(s[0]); gw("," as *u8); gm(s[1]); gw("," as *u8); gm(s[2]); gw("]m sums to " as *u8); gn(ssum); gw("m (expect 1000)\n" as *u8)
86
87 // T3 SOFTMAX BACKWARD GRADCHECK -- L = softmax_0 (w=[1,0,0]); dL/dx vs central finite-diff.
88 let w: *i64=sys_mmap(64) as *i64; w[0]=f32_of(1); w[1]=f32_of(0); w[2]=f32_of(0)
89 let dx: *i64=sys_mmap(64) as *i64; softmax_backward(s, w, 3, dx)
90 let twoeps: i64=f32_mul(f32_of(2), eps)
91 var gok: i64=1; var j: i64=0
92 while j<3 {
93 let xp: *i64=sys_mmap(64) as *i64; let xm: *i64=sys_mmap(64) as *i64
94 var c: i64=0; while c<3 { xp[c]=x[c]; xm[c]=x[c]; c=c+1 }
95 xp[j]=f32_add(x[j],eps); xm[j]=f32_sub(x[j],eps)
96 let fd: i64=f32_div(f32_sub(compute_L(xp,w,3), compute_L(xm,w,3)), twoeps)
97 if f32_le(f32_abs(f32_sub(dx[j], fd)), tol)==0 { gok=0 }
98 gw(" d/dx" as *u8); gn(j); gw(": autograd=" as *u8); gm(dx[j]); gw("m finite-diff=" as *u8); gm(fd); gw("m\n" as *u8)
99 j=j+1
100 }
101 total=total+1; if gok==1 { pass=pass+1; gw(" [PASS] " as *u8) } else { gw(" [FAIL] " as *u8) }
102 gw("T3 SOFTMAX BACKWARD GRADCHECK: all dL/dx autograd == finite-diff (within tol) -> attention backprop is correct\n" as *u8)
103
104 // T4 NEGATIVE CONTROL.
105 let xp0: *i64=sys_mmap(64) as *i64; let xm0: *i64=sys_mmap(64) as *i64
106 var c2: i64=0; while c2<3 { xp0[c2]=x[c2]; xm0[c2]=x[c2]; c2=c2+1 }
107 xp0[0]=f32_add(x[0],eps); xm0[0]=f32_sub(x[0],eps)
108 let fd0: i64=f32_div(f32_sub(compute_L(xp0,w,3), compute_L(xm0,w,3)), twoeps)
109 let wrong: i64=f32_mul(dx[0], f32_of(2))
110 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) }
111 gw("T4 NEGATIVE CONTROL: a doubled gradient " as *u8); gm(wrong); gw("m is REJECTED vs finite-diff " as *u8); gm(fd0); gw("m\n" as *u8)
112
113 // T5.
114 total=total+1; if gok==1 { pass=pass+1; gw(" [PASS] " as *u8) } else { gw(" [FAIL] " as *u8) }
115 gw("T5 RUNG 2 CORE: f32_exp + f32_sqrt + softmax forward/backward gradchecked -> attention backprop unblocked\n" as *u8)
116
117 gw("\n RUNG 2 (core) DONE: the transcendentals exist (f32_exp Taylor, f32_sqrt Newton) and softmax backward is gradchecked --\n" as *u8)
118 gw(" the attention nonlinearity backprops correctly in f32. NEXT (R2b): layernorm forward/backward (uses f32_sqrt), gradchecked;\n" as *u8)
119 gw(" then R3 Adam, R4 data pipeline (crawl-on-burst), R5 train loop. Each rung finite-difference-verified, same as this.\n" as *u8)
120 gw("F32-SOFTMAX verdict=" as *u8)
121 if pass==total { gw("GREEN passes=" as *u8); gn(pass); gw("/" as *u8); gn(total); gw(" -- transcendentals + softmax backward correct (gradchecked), rung 2 core\n" as *u8); sys_exit(0); return 0 }
122 gw("RED passes=" as *u8); gn(pass); gw("/" as *u8); gn(total); gw("\n" as *u8); sys_exit(1); return 1
123}