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}