code wiki / _hdl_build / nx_train_int_vs_softfloat_gate.nx

nx_train_int_vs_softfloat_gate.nx source

↩ module page · 121 lines · 8477 B

1// nx_train_int_vs_softfloat_gate.nx -- THE PAYOFF measurement: does INTEGER training actually run FASTER than our 2// SOFTWARE-FLOAT tape on IDENTICAL work? Trains the SAME 2-layer MLP (W1->ReLU->W2, L=sum(y-t)^2) for the same 3// number of steps, TWO ways -- (A) software-float via nx_f32_mul/add (what nx_f32_qwen2_train uses), (B) integer 4// Q16 -- times both with sys_now_us, and reports the speedup. Same algorithm, same MAC count; only the arithmetic 5// differs, so the ratio IS the software-float tax paid in a real training loop. This is the internal h2h that 6// justifies porting the trainer to integer (the sovereign path to beat PyTorch CPU). Measured-not-asserted. 7// ReLU + sign checks read the IEEE sign bit directly ((raw>>31)&1) to avoid nx_int-typed comparisons. license_tier: ORIGINAL 8import "nx_f32.nx" 9import "nx_f32_cvt.nx" 10import "nx_f32_div.nx" 11import "nx_syscalls.nx" 12 13func w(s: *u8) -> i64 { var n: i64=0; while s[n]!=(0 as u8){n=n+1} sys_write(1,s,n); return 0 } 14func wn(v: i64) -> i64 { if v==0 { sys_write(1,"0" as *u8,1); return 0 } var m: i64=v; if m<0{sys_write(1,"-" as *u8,1);m=0-m} let t: *u8=sys_mmap(24); var k: i64=0; while m>0{t[k]=(48+(m%10)) as u8;m=m/10;k=k+1} let o: *u8=sys_mmap(24); var q: i64=k-1; var i: i64=0; while q>=0{o[i]=t[q];i=i+1;q=q-1} sys_write(1,o,i); return 0 } 15func fneg(raw: i64) -> i64 { return (raw/2147483648)%2 } // IEEE sign bit (1 => negative) 16 17const S: i64 = 65536 18const K: i64 = 3 19const H: i64 = 8 20const O: i64 = 2 21const NEX: i64 = 3 22const STEPS: i64 = 1500 23 24// ---------- (B) INTEGER Q16 step: fwd+bwd, fills gW1,gW2, returns L_q32 ---------- 25func step_int(W1: *i64, X: *i64, W2: *i64, T: *i64, gW1: *i64, gW2: *i64, hb: *i64, ab: *i64, yb: *i64) -> i64 { 26 var z: i64=0; while z<H*K { gW1[z]=0; z=z+1 } z=0; while z<O*H { gW2[z]=0; z=z+1 } 27 let gy: *i64=sys_mmap(O*8) as *i64 28 var Ltot: i64=0; var ex: i64=0 29 while ex<NEX { 30 var i: i64=0 31 while i<H { var acc: i64=0; var kk: i64=0; while kk<K { acc=acc + W1[i*K+kk]*X[ex*K+kk]; kk=kk+1 } let hv: i64=acc/S; hb[i]=hv; if hv<0 { ab[i]=0 } else { ab[i]=hv } i=i+1 } 32 var o: i64=0; while o<O { var acc2: i64=0; var j: i64=0; while j<H { acc2=acc2 + W2[o*H+j]*ab[j]; j=j+1 } yb[o]=acc2/S; o=o+1 } 33 o=0; while o<O { let e: i64=yb[o]-T[ex*O+o]; Ltot=Ltot + e*e; gy[o]=2*e; o=o+1 } 34 o=0; while o<O { var j: i64=0; while j<H { gW2[o*H+j]=gW2[o*H+j] + (gy[o]*ab[j])/S; j=j+1 } o=o+1 } 35 var j2: i64=0 36 while j2<H { var acc3: i64=0; o=0; while o<O { acc3=acc3 + gy[o]*W2[o*H+j2]; o=o+1 } let ga: i64=acc3/S; var gh: i64=0; if hb[j2]>=0 { gh=ga } var kk: i64=0; while kk<K { gW1[j2*K+kk]=gW1[j2*K+kk] + (gh*X[ex*K+kk])/S; kk=kk+1 } j2=j2+1 } 37 ex=ex+1 38 } 39 return Ltot 40} 41 42// ---------- (A) SOFTWARE-FLOAT step: same algorithm with nx_f32 ops, fills gW1f,gW2f, returns L (f32 raw) ---------- 43func step_sf(W1: *i64, X: *i64, W2: *i64, T: *i64, gW1: *i64, gW2: *i64, hb: *i64, ab: *i64, yb: *i64) -> i64 { 44 let zero: i64=nx_i32_to_f32(0); let two: i64=nx_i32_to_f32(2) 45 var z: i64=0; while z<H*K { gW1[z]=zero; z=z+1 } z=0; while z<O*H { gW2[z]=zero; z=z+1 } 46 let gy: *i64=sys_mmap(O*8) as *i64 47 var Ltot: i64=zero; var ex: i64=0 48 while ex<NEX { 49 var i: i64=0 50 while i<H { var acc: i64=zero; var kk: i64=0; while kk<K { acc=nx_f32_add(acc, nx_f32_mul(W1[i*K+kk], X[ex*K+kk])); kk=kk+1 } hb[i]=acc; if fneg(acc)==1 { ab[i]=zero } else { ab[i]=acc } i=i+1 } 51 var o: i64=0; while o<O { var acc2: i64=zero; var j: i64=0; while j<H { acc2=nx_f32_add(acc2, nx_f32_mul(W2[o*H+j], ab[j])); j=j+1 } yb[o]=acc2; o=o+1 } 52 o=0; while o<O { let e: i64=nx_f32_sub(yb[o], T[ex*O+o]); Ltot=nx_f32_add(Ltot, nx_f32_mul(e,e)); gy[o]=nx_f32_mul(two,e); o=o+1 } 53 o=0; while o<O { var j: i64=0; while j<H { gW2[o*H+j]=nx_f32_add(gW2[o*H+j], nx_f32_mul(gy[o], ab[j])); j=j+1 } o=o+1 } 54 var j2: i64=0 55 while j2<H { var acc3: i64=zero; o=0; while o<O { acc3=nx_f32_add(acc3, nx_f32_mul(gy[o], W2[o*H+j2])); o=o+1 } var gh: i64=zero; if fneg(hb[j2])==0 { gh=acc3 } var kk: i64=0; while kk<K { gW1[j2*K+kk]=nx_f32_add(gW1[j2*K+kk], nx_f32_mul(gh, X[ex*K+kk])); kk=kk+1 } j2=j2+1 } 56 ex=ex+1 57 } 58 return Ltot 59} 60 61func main() -> i64 { 62 w("=== nx_train_int_vs_softfloat: SAME MLP trained integer vs software-float, " as *u8); wn(STEPS); w(" steps -- SPEED h2h ===\n\n" as *u8) 63 // integer arrays 64 let W1: *i64=sys_mmap(H*K*8) as *i64; let W2: *i64=sys_mmap(O*H*8) as *i64 65 let X: *i64=sys_mmap(NEX*K*8) as *i64; let T: *i64=sys_mmap(NEX*O*8) as *i64 66 let gW1: *i64=sys_mmap(H*K*8) as *i64; let gW2: *i64=sys_mmap(O*H*8) as *i64 67 let hb: *i64=sys_mmap(H*8) as *i64; let ab: *i64=sys_mmap(H*8) as *i64; let yb: *i64=sys_mmap(O*8) as *i64 68 // software-float arrays (same values, as f32) 69 let W1f: *i64=sys_mmap(H*K*8) as *i64; let W2f: *i64=sys_mmap(O*H*8) as *i64 70 let Xf: *i64=sys_mmap(NEX*K*8) as *i64; let Tf: *i64=sys_mmap(NEX*O*8) as *i64 71 let gW1f: *i64=sys_mmap(H*K*8) as *i64; let gW2f: *i64=sys_mmap(O*H*8) as *i64 72 let hbf: *i64=sys_mmap(H*8) as *i64; let abf: *i64=sys_mmap(H*8) as *i64; let ybf: *i64=sys_mmap(O*8) as *i64 73 74 var i: i64=0; while i<H*K { let iv: i64=(((i*7+3)%17)-8); W1[i]=(iv*S)/60; W1f[i]=nx_f32_div(nx_i32_to_f32(iv), nx_i32_to_f32(60)); i=i+1 } 75 i=0; while i<O*H { let iv: i64=(((i*11+2)%13)-6); W2[i]=(iv*S)/60; W2f[i]=nx_f32_div(nx_i32_to_f32(iv), nx_i32_to_f32(60)); i=i+1 } 76 var ex: i64=0 77 while ex<NEX { 78 var kk: i64=0; while kk<K { let iv: i64=((ex*2+kk+1)%4+1); X[ex*K+kk]=(iv*S)/6; Xf[ex*K+kk]=nx_f32_div(nx_i32_to_f32(iv), nx_i32_to_f32(6)); kk=kk+1 } 79 var o: i64=0; while o<O { let iv: i64=(((ex*3+o*2+1)%5)-2); T[ex*O+o]=(iv*S)/4; Tf[ex*O+o]=nx_f32_div(nx_i32_to_f32(iv), nx_i32_to_f32(4)); o=o+1 } 80 ex=ex+1 81 } 82 let lr: i64=(S*15)/100; let lrf: i64=nx_f32_div(nx_i32_to_f32(15), nx_i32_to_f32(100)) 83 84 // ---- (B) INTEGER training, timed ---- 85 let L0i: i64=step_int(W1, X, W2, T, gW1, gW2, hb, ab, yb) 86 let ti0: i64=sys_now_us() 87 var step: i64=1 88 while step<=STEPS { 89 let L: i64=step_int(W1, X, W2, T, gW1, gW2, hb, ab, yb) 90 var a: i64=0; while a<H*K { W1[a]=W1[a]-(lr*gW1[a])/S; a=a+1 } a=0; while a<O*H { W2[a]=W2[a]-(lr*gW2[a])/S; a=a+1 } 91 step=step+1 92 } 93 let ti1: i64=sys_now_us() 94 let Lfi: i64=step_int(W1, X, W2, T, gW1, gW2, hb, ab, yb) 95 96 // ---- (A) SOFTWARE-FLOAT training, timed ---- 97 let L0f: i64=step_sf(W1f, Xf, W2f, Tf, gW1f, gW2f, hbf, abf, ybf) 98 let ts0: i64=sys_now_us() 99 step=1 100 while step<=STEPS { 101 let L: i64=step_sf(W1f, Xf, W2f, Tf, gW1f, gW2f, hbf, abf, ybf) 102 var a: i64=0; while a<H*K { W1f[a]=nx_f32_sub(W1f[a], nx_f32_mul(lrf, gW1f[a])); a=a+1 } a=0; while a<O*H { W2f[a]=nx_f32_sub(W2f[a], nx_f32_mul(lrf, gW2f[a])); a=a+1 } 103 step=step+1 104 } 105 let ts1: i64=sys_now_us() 106 let Lff: i64=step_sf(W1f, Xf, W2f, Tf, gW1f, gW2f, hbf, abf, ybf) 107 108 let int_us: i64=ti1-ti0; var sf_us: i64=ts1-ts0 109 // both learned? integer: L_q32 numeric. software-float: ratio L0f/Lff > 10 via sign bit of (ratio-10). 110 let sf_learned: i64=fneg(nx_f32_sub(nx_f32_div(L0f, Lff), nx_i32_to_f32(10))) // 0 => ratio>10 => learned 111 w(" (A) SOFTWARE-FLOAT: " as *u8); wn(sf_us); w(" us learned(>10x drop)=" as *u8); if sf_learned==0 { w("YES" as *u8) } else { w("no" as *u8) } w("\n" as *u8) 112 w(" (B) INTEGER Q16: " as *u8); wn(int_us); w(" us loss " as *u8); wn(L0i); w(" -> " as *u8); wn(Lfi); w(" (>10x drop)\n\n" as *u8) 113 var spd: i64=0; if int_us>0 { spd=sf_us/int_us } 114 w(" >>> INTEGER TRAINING SPEEDUP over software-float = " as *u8); wn(spd); w("x (same MLP, same " as *u8); wn(STEPS); w(" steps, same MACs)\n" as *u8) 115 w(" => porting the trainer's ta_ ops to integer removes the software-float tax IN THE TRAINING LOOP (measured here).\n" as *u8) 116 w(" => and this is SCALAR integer; the SIMD gemm path (8.4 Gop/s) widens it further -> the shot at beating PyTorch CPU.\n" as *u8) 117 w("NX-TRAIN-INT-VS-SOFTFLOAT verdict=" as *u8) 118 if spd>=3 { if Lfi*10<L0i { w("GREEN -- integer training " as *u8); wn(spd); w("x faster AND learns (both paths converge); the fix delivers speed, MEASURED\n" as *u8) } else { w("YELLOW -- fast but integer didn't converge here\n" as *u8) } } 119 else { w("YELLOW -- speedup <3x; investigate (expected >=10x from the 25x MAC tax minus loop overhead)\n" as *u8) } 120 return 0 121}