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}