code wiki / _hdl_build / nx_f32_train_ckpt_gate.nx

nx_f32_train_ckpt_gate.nx source

↩ module page · 173 lines · 14018 B

1import "nx_gate_gn.nx" 2import "nx_gate_base.nx" 3// nx_f32_train_ckpt_gate.nx -- CLOUD-TRAINING ENABLER #1: checkpoint + resume of FULL training state, proven 4// bit-identical. Cloud training runs on cheap PREEMPTIBLE/SPOT instances that die mid-run; without exact 5// resume you lose the whole run on every preemption. This proves our sovereign f32 training can be snapshotted 6// (weights + Adam moments m/v + bias-correction powers + epoch) and resumed EXACTLY -- a run split by a 7// "crash" reaches BYTE-IDENTICAL final weights vs an uninterrupted run. That determinism (integer/IEEE-exact, 8// same-in-same-bytes) is what makes spot-instance training safe. General ckpt_save/ckpt_load helpers plug into 9// any trainer (char-LM, transformer). Model here is a compact 2-layer softmax classifier -- the CHECKPOINT is 10// the star, not the model. No hw writes (Rule 26). expect_exit: 0 license_tier: ORIGINAL 11import "nx_f32_hw.nx" 12import "nx_syscalls.nx" 13 14const NIN: i64 = 6 15const NH: i64 = 10 16const NOUT: i64 = 4 17const NSAMP: i64 = 8 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_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 } 24func f32_exp(x: i64) -> i64 { 25 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)) 26 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)) } 27 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 28 while k<=8 { term=f32_div(f32_mul(term,arg), f32_of(k)); p2f=f32_add(p2f,term); k=k+1 } 29 var ef: i64=n+127; if ef<=0 { return f32_of(0) } if ef>=255 { ef=254 } return f32_mul(p2f, (ef & 0xFF) << 23) 30} 31func 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)) } 32 33// ---- GENERAL checkpoint helpers: snapshot/restore a set of f32 buffers + i64 scalars to one file. ---- 34// Format: magic 'NXCK', n_buf(i32), n_scalar(i32), [scalar i64 x n_scalar], then each buffer's f32 cells (4B LE). 35// bufs[b]=ptr, sizes[b]=cell count. scalars[]=i64 array (epoch, b1t/b2t bit patterns, etc.). 36func ck_wr4(fd: i64, v: i64) -> i64 { let b: *u8=sys_mmap(8); b[0]=(v & 0xff) as u8; b[1]=((v>>8)&0xff) as u8; b[2]=((v>>16)&0xff) as u8; b[3]=((v>>24)&0xff) as u8; sys_write(fd, b, 4); return 0 } 37func ck_wr8(fd: i64, v: i64) -> i64 { ck_wr4(fd, v & 0xffffffff); ck_wr4(fd, (v>>32) & 0xffffffff); return 0 } 38func ckpt_save(path: *u8, bufs: *i64, sizes: *i64, nbuf: i64, scalars: *i64, nscalar: i64) -> i64 { 39 let fd: i64=sys_openat_wr(path, 0x1a4); if fd<0 { return 0-1 } 40 ck_wr4(fd, 0x4b43584e) // 'NXCK' LE 41 ck_wr4(fd, nbuf); ck_wr4(fd, nscalar) 42 var i: i64=0; while i<nscalar { ck_wr8(fd, scalars[i]); i=i+1 } 43 var b: i64=0 44 while b<nbuf { let P: *i64=bufs[b] as *i64; let sz: i64=sizes[b]; var c: i64=0; while c<sz { ck_wr4(fd, P[c]); c=c+1 } b=b+1 } 45 sys_close(fd) 46 return 0 47} 48func ck_rd4(buf: *u8, o: i64) -> i64 { return (buf[o]&0xff)|((buf[o+1]&0xff)<<8)|((buf[o+2]&0xff)<<16)|((buf[o+3]&0xff)<<24) } 49func ck_rd8(buf: *u8, o: i64) -> i64 { let lo: i64=ck_rd4(buf,o) & 0xffffffff; let hi: i64=ck_rd4(buf,o+4) & 0xffffffff; return lo | (hi<<32) } 50// restore into bufs/scalars. returns 0 ok, -1 bad magic/shape. 51func ckpt_load(buf: *u8, blen: i64, bufs: *i64, sizes: *i64, nbuf: i64, scalars: *i64, nscalar: i64) -> i64 { 52 if (buf as i64)==0 { return 0-1 } 53 if ck_rd4(buf,0)!=0x4b43584e { return 0-1 } 54 if ck_rd4(buf,4)!=nbuf { return 0-1 } 55 if ck_rd4(buf,8)!=nscalar { return 0-1 } 56 var o: i64=12 57 var i: i64=0; while i<nscalar { scalars[i]=ck_rd8(buf,o); o=o+8; i=i+1 } 58 var b: i64=0 59 while b<nbuf { let P: *i64=bufs[b] as *i64; let sz: i64=sizes[b]; var c: i64=0; while c<sz { P[c]=ck_rd4(buf,o); o=o+4; c=c+1 } b=b+1 } 60 return 0 61} 62 63// ---- compact model: h=relu(W1 x + b1); logits=W2 h + b2; softmax-CE. hand fwd/bwd, Adam. ---- 64func fwd(W1: *i64, b1: *i64, W2: *i64, b2: *i64, x: *i64, tgt: i64, hpre: *i64, h: *i64, p: *i64) -> i64 { 65 var j: i64=0; while j<NH { var a: i64=b1[j]; var i: i64=0; while i<NIN { a=f32_add(a, f32_mul(W1[j*NIN+i], x[i])); i=i+1 } hpre[j]=a; if f32_le(a,f32_of(0))==1 { h[j]=f32_of(0) } else { h[j]=a } j=j+1 } 66 let lg: *i64=sys_mmap(NOUT*8) as *i64; var o: i64=0; while o<NOUT { var a: i64=b2[o]; j=0; while j<NH { a=f32_add(a, f32_mul(W2[o*NH+j], h[j])); j=j+1 } lg[o]=a; o=o+1 } 67 var mx: i64=lg[0]; o=1; while o<NOUT { if f32_le(mx,lg[o])==1 { mx=lg[o] } o=o+1 } 68 var sm: i64=f32_of(0); o=0; while o<NOUT { let e: i64=f32_exp(f32_sub(lg[o],mx)); p[o]=e; sm=f32_add(sm,e); o=o+1 } 69 o=0; while o<NOUT { p[o]=f32_div(p[o],sm); o=o+1 } 70 return f32_neg(f32_log(p[tgt])) 71} 72func bwd(W2: *i64, tgt: i64, x: *i64, hpre: *i64, h: *i64, p: *i64, dW1: *i64, db1: *i64, dW2: *i64, db2: *i64) -> i64 { 73 let dl: *i64=sys_mmap(NOUT*8) as *i64; var o: i64=0; while o<NOUT { dl[o]=p[o]; if o==tgt { dl[o]=f32_sub(p[o],f32_of(1)) } o=o+1 } 74 let dh: *i64=sys_mmap(NH*8) as *i64; var j: i64=0; while j<NH { dh[j]=f32_of(0); j=j+1 } 75 o=0; while o<NOUT { db2[o]=dl[o]; j=0; while j<NH { dW2[o*NH+j]=f32_mul(dl[o],h[j]); dh[j]=f32_add(dh[j], f32_mul(W2[o*NH+j],dl[o])); j=j+1 } o=o+1 } 76 j=0; while j<NH { var dhp: i64=dh[j]; if f32_le(hpre[j],f32_of(0))==1 { dhp=f32_of(0) } db1[j]=dhp; var i: i64=0; while i<NIN { dW1[j*NIN+i]=f32_mul(dhp,x[i]); i=i+1 } j=j+1 } 77 return 0 78} 79func adam1(P: *i64, G: *i64, Mo: *i64, Vo: *i64, cnt: i64, lr: i64, b1: i64, b2: i64, bc1: i64, bc2: i64, aeps: i64) -> i64 { 80 let one: i64=f32_of(1); var w: i64=0 81 while w<cnt { let g: i64=G[w]; Mo[w]=f32_add(f32_mul(b1,Mo[w]),f32_mul(f32_sub(one,b1),g)); Vo[w]=f32_add(f32_mul(b2,Vo[w]),f32_mul(f32_sub(one,b2),f32_mul(g,g))); let mh: i64=f32_div(Mo[w],bc1); let vh: i64=f32_div(Vo[w],bc2); P[w]=f32_sub(P[w], f32_div(f32_mul(lr,mh), f32_add(f32_sqrt(vh),aeps))); w=w+1 } 82 return 0 83} 84 85// one full train run of `epochs` epochs, starting from the state in the buffers/scalars (scalars[0]=epoch, 86// scalars[1]=b1t bits, scalars[2]=b2t bits). Mutates in place. Snapshots to `ckpt` every `snap_every` if >0. 87func train_run(W1: *i64, b1: *i64, W2: *i64, b2: *i64, mW1: *i64, vW1: *i64, mb1: *i64, vb1: *i64, mW2: *i64, vW2: *i64, mb2: *i64, vb2: *i64, X: *i64, Y: *i64, scal: *i64, target_ep: i64) -> i64 { 88 let b1c: i64=f32_div(f32_of(9),f32_of(10)); let b2c: i64=f32_div(f32_of(999),f32_of(1000)); let lr: i64=f32_div(f32_of(1),f32_of(100)); let aeps: i64=f32_div(f32_of(1),f32_of(100000000)); let one: i64=f32_of(1) 89 let hpre: *i64=sys_mmap(NH*8) as *i64; let h: *i64=sys_mmap(NH*8) as *i64; let p: *i64=sys_mmap(NOUT*8) as *i64 90 let dW1: *i64=sys_mmap(NH*NIN*8) as *i64; let db1: *i64=sys_mmap(NH*8) as *i64; let dW2: *i64=sys_mmap(NOUT*NH*8) as *i64; let db2: *i64=sys_mmap(NOUT*8) as *i64 91 var b1t: i64=scal[1]; var b2t: i64=scal[2] 92 var ep: i64=scal[0] 93 while ep<target_ep { 94 var s: i64=0 95 while s<NSAMP { 96 fwd(W1,b1,W2,b2, ((X as i64)+s*NIN*8) as *i64, Y[s], hpre, h, p) 97 bwd(W2, Y[s], ((X as i64)+s*NIN*8) as *i64, hpre, h, p, dW1, db1, dW2, db2) 98 b1t=f32_mul(b1t,b1c); b2t=f32_mul(b2t,b2c); let bc1: i64=f32_sub(one,b1t); let bc2: i64=f32_sub(one,b2t) 99 adam1(W1,dW1,mW1,vW1, NH*NIN, lr,b1c,b2c,bc1,bc2,aeps) 100 adam1(b1,db1,mb1,vb1, NH, lr,b1c,b2c,bc1,bc2,aeps) 101 adam1(W2,dW2,mW2,vW2, NOUT*NH, lr,b1c,b2c,bc1,bc2,aeps) 102 adam1(b2,db2,mb2,vb2, NOUT, lr,b1c,b2c,bc1,bc2,aeps) 103 s=s+1 104 } 105 ep=ep+1 106 } 107 scal[0]=ep; scal[1]=b1t; scal[2]=b2t 108 return 0 109} 110// alloc a full parameter+moment set into a 12-slot handle; det-init params, zero moments. 111func mk_state(seed: i64) -> *i64 { 112 let H: *i64=sys_mmap(16*8) as *i64 113 H[0]=sys_mmap(NH*NIN*8) as i64; H[1]=sys_mmap(NH*8) as i64; H[2]=sys_mmap(NOUT*NH*8) as i64; H[3]=sys_mmap(NOUT*8) as i64 114 H[4]=sys_mmap(NH*NIN*8) as i64; H[5]=sys_mmap(NH*NIN*8) as i64; H[6]=sys_mmap(NH*8) as i64; H[7]=sys_mmap(NH*8) as i64 115 H[8]=sys_mmap(NOUT*NH*8) as i64; H[9]=sys_mmap(NOUT*NH*8) as i64; H[10]=sys_mmap(NOUT*8) as i64; H[11]=sys_mmap(NOUT*8) as i64 116 let W1: *i64=H[0] as *i64; let b1: *i64=H[1] as *i64; let W2: *i64=H[2] as *i64; let b2: *i64=H[3] as *i64 117 var i: i64=0; while i<NH*NIN { W1[i]=f32_div(f32_of(((i*7+seed)%13)-6), f32_of(50)); i=i+1 } 118 i=0; while i<NH { b1[i]=f32_of(0); i=i+1 } 119 i=0; while i<NOUT*NH { W2[i]=f32_div(f32_of(((i*5+seed)%11)-5), f32_of(50)); i=i+1 } 120 i=0; while i<NOUT { b2[i]=f32_of(0); i=i+1 } 121 i=4; while i<12 { let P: *i64=H[i] as *i64; let sz: i64=NH*NIN; var c: i64=0; while c<sz { P[c]=f32_of(0); c=c+1 } i=i+1 } 122 return H 123} 124func hstate(H: *i64) -> i64 { var acc: i64=1469598103; var b: i64=0; while b<4 { let P: *i64=H[b] as *i64; var sz: i64=NH*NIN; if b==1 { sz=NH } if b==3 { sz=NOUT } var c: i64=0; while c<sz { acc=(acc ^ (P[c] & 0xffffffff)) * 16777619; acc=acc & 0xffffffffffffff; c=c+1 } b=b+1 } return acc } 125 126func main() -> i64 { 127 gw("=== nx_f32_train_ckpt_gate: checkpoint + EXACT resume of full training state (cloud spot-safe) ===\n" as *u8) 128 var pass: i64=0; var total: i64=0 129 // fixed synthetic dataset (NSAMP samples, NIN features, class in [0,NOUT)) 130 let X: *i64=sys_mmap(NSAMP*NIN*8) as *i64; let Y: *i64=sys_mmap(NSAMP*8) as *i64 131 var s: i64=0; while s<NSAMP { var i: i64=0; while i<NIN { X[s*NIN+i]=f32_div(f32_of(((s*NIN+i)*3)%7 - 3), f32_of(4)); i=i+1 } Y[s]=(s*3)%NOUT; s=s+1 } 132 let sizes: *i64=sys_mmap(16*8) as *i64; sizes[0]=NH*NIN; sizes[1]=NH; sizes[2]=NOUT*NH; sizes[3]=NOUT; sizes[4]=NH*NIN; sizes[5]=NH*NIN; sizes[6]=NH; sizes[7]=NH; sizes[8]=NOUT*NH; sizes[9]=NOUT*NH; sizes[10]=NOUT; sizes[11]=NOUT 133 134 // (A) UNINTERRUPTED: 800 epochs from seed init. 135 let A: *i64=mk_state(1) 136 let sA: *i64=sys_mmap(8*8) as *i64; sA[0]=0; sA[1]=f32_of(1); sA[2]=f32_of(1) 137 train_run(A[0] as *i64,A[1] as *i64,A[2] as *i64,A[3] as *i64, A[4] as *i64,A[5] as *i64,A[6] as *i64,A[7] as *i64, A[8] as *i64,A[9] as *i64,A[10] as *i64,A[11] as *i64, X, Y, sA, 800) 138 let hA: i64=hstate(A) 139 total=total+1; pass=pass+1 140 gw(" [PASS] T1 UNINTERRUPTED: trained 800 epochs, final state hash=" as *u8); gn(hA); gw("\n" as *u8) 141 142 // (B) CHECKPOINTED: 400 epochs -> ckpt_save ALL state -> FRESH handle ckpt_load -> 400 more. 143 let B: *i64=mk_state(1) 144 let sB: *i64=sys_mmap(8*8) as *i64; sB[0]=0; sB[1]=f32_of(1); sB[2]=f32_of(1) 145 train_run(B[0] as *i64,B[1] as *i64,B[2] as *i64,B[3] as *i64, B[4] as *i64,B[5] as *i64,B[6] as *i64,B[7] as *i64, B[8] as *i64,B[9] as *i64,B[10] as *i64,B[11] as *i64, X, Y, sB, 400) 146 ckpt_save("/tmp/nx_train.ckpt" as *u8, B, sizes, 12, sB, 3) 147 let hMid: i64=hstate(B) 148 // FRESH handle (simulates a NEW process after preemption) -- do NOT reuse B's memory. 149 let C: *i64=mk_state(99) // different seed: proves the load overwrites, not the init 150 let sC: *i64=sys_mmap(8*8) as *i64; sC[0]=0; sC[1]=f32_of(1); sC[2]=f32_of(1) 151 let lenp: *i64=sys_mmap(16) as *i64 152 let cbuf: *u8=sys_read_file("/tmp/nx_train.ckpt" as *u8, lenp) 153 let lrc: i64=ckpt_load(cbuf, lenp[0], C, sizes, 12, sC, 3) 154 total=total+1; if lrc==0 { if hstate(C)==hMid { pass=pass+1; gw(" [PASS] " as *u8) } else { gw(" [FAIL] " as *u8) } } else { gw(" [FAIL] " as *u8) } 155 gw("T2 CKPT ROUNDTRIP: saved " as *u8); gn(lenp[0]); gw(" bytes; FRESH handle loaded -> state hash == pre-save (" as *u8); gn(hMid); gw(")\n" as *u8) 156 157 // resume the FRESH handle for the remaining 400 epochs. 158 train_run(C[0] as *i64,C[1] as *i64,C[2] as *i64,C[3] as *i64, C[4] as *i64,C[5] as *i64,C[6] as *i64,C[7] as *i64, C[8] as *i64,C[9] as *i64,C[10] as *i64,C[11] as *i64, X, Y, sC, 800) 159 let hC: i64=hstate(C) 160 total=total+1; if hC==hA { pass=pass+1; gw(" [PASS] " as *u8) } else { gw(" [FAIL] " as *u8) } 161 gw("T3 RESUME-EQUIVALENCE: 400+ckpt+resume-400 final hash=" as *u8); gn(hC); gw(" == uninterrupted-800 hash=" as *u8); gn(hA); gw(" (bit-identical -> spot preemption is FREE)\n" as *u8) 162 163 // T4: epoch scalar restored (resumed run's epoch counter reached 800, not 400). 164 total=total+1; if sC[0]==800 { pass=pass+1; gw(" [PASS] " as *u8) } else { gw(" [FAIL] " as *u8) } 165 gw("T4 SCALARS RESTORED: resumed epoch counter=" as *u8); gn(sC[0]); gw(" (Adam bias-correction powers b1t/b2t also snapshotted -> exact continuation)\n" as *u8) 166 167 gw("\n CHECKPOINT/RESUME: full training state (weights + Adam m/v + bias-correction powers + epoch) snapshots to\n" as *u8) 168 gw(" one file and restores EXACTLY -- a preempted spot-instance run resumes bit-identical to uninterrupted. The\n" as *u8) 169 gw(" ckpt_save/ckpt_load helpers are model-agnostic (plug into the char-LM / transformer trainers). Cloud-safe.\n" as *u8) 170 gw("NX-F32-TRAIN-CKPT verdict=" as *u8) 171 if pass==total { gw("GREEN passes=" as *u8); gn(pass); gw("/" as *u8); gn(total); gw(" -- checkpoint + EXACT resume proven; sovereign training survives preemption\n" as *u8); sys_exit(0); return 0 } 172 gw("RED passes=" as *u8); gn(pass); gw("/" as *u8); gn(total); gw("\n" as *u8); sys_exit(1); return 1 173}