code wiki / _hdl_build / nx_nofloat_kvcache_gate.nx

nx_nofloat_kvcache_gate.nx source

↩ module page · 153 lines · 9280 B

1// nx_nofloat_kvcache_gate.nx -- CAP-NF-KVCACHE: a KV-CACHE for O(T) autoregressive decode (vs O(T^2) 2// full-recompute), pure integer Q16. Single-block causal attention model (random weights -- this is about 3// inference CORRECTNESS + EFFICIENCY, not model quality, so no training needed). Two decoders share the SAME 4// primitives: full-recompute re-derives every past K/V each step; cached computes each K/V once and reuses it. 5// T1 CORRECT: cached logits == full-recompute logits at every position (bit-exact, by construction). 6// T2 CHEAPER: cached does far fewer K/V projections (full = T(T+1)/2, cached = T) -- measured counts. 7// T3 TEETH: a CORRUPTED cache (one stale entry) yields DIFFERENT logits -> the cache content is really used. 8// Sovereign: nx_nofloat_autograd (qmul/isqrt/sin/cos/exp) + nx_syscalls. expect_exit: 0 9import "nx_nofloat_autograd.nx" 10import "nx_syscalls.nx" 11import "nx_gate_emit_lib.nx" 12import "nx_gate_verdict.nx" 13const Q16: i64 = 65536 14 15 16func dini(a: *i64, n: i64, sd: i64) -> i64 { var i: i64=0; while i<n { a[i]=(((i*7+sd*13+1)%11)-5)*9362; i=i+1 } return 0 } 17 18// row-vector(x, len=rows) times matrix(W, rows x cols) -> out(cols); Q16 accumulate-then-shift. opc += rows*cols. 19func matvec(x: *i64, W: *i64, rows: i64, cols: i64, out: *i64, opc: *i64) -> i64 { 20 var c: i64=0 21 while c<cols { var acc: i64=0; var r: i64=0; while r<rows { acc=acc + x[r]*W[r*cols+c]; r=r+1 } out[c]=acc>>16; c=c+1 } 22 opc[0]=opc[0]+rows*cols 23 return 0 24} 25func rmsnorm(x: *i64, n: i64, out: *i64) -> i64 { 26 var ss: i64=0; var i: i64=0; while i<n { ss=ss + x[i]*x[i]; i=i+1 } // Q32 sum of squares 27 var ms: i64=ss/n 28 var rms: i64=nfa_isqrt(ms) // sqrt(Q32) = Q16 rms 29 if rms<1 { rms=1 } 30 i=0; while i<n { out[i]=(x[i]<<16)/rms; i=i+1 } 31 return 0 32} 33func dotq(a: *i64, b: *i64, n: i64) -> i64 { var acc: i64=0; var i: i64=0; while i<n { acc=acc + a[i]*b[i]; i=i+1 } return acc>>16 } 34func copyv(src: *i64, dst: *i64, n: i64) -> i64 { var i: i64=0; while i<n { dst[i]=src[i]; i=i+1 } return 0 } 35// simple consistent RoPE: rotate pairs by angle = pos * (Q16>>(i+1)) (form is irrelevant; consistency is what matters) 36func rope(v: *i64, pos: i64, dm: i64) -> i64 { 37 var i: i64=0 38 while i+1<dm { 39 let ang: i64 = pos * (32768 >> (i/2 + 1)) 40 let cs: i64 = nfa_cosf(ang); let sn: i64 = nfa_sinf(ang) 41 let a: i64=v[i]; let b: i64=v[i+1] 42 v[i] = nfa_qmul(a,cs) - nfa_qmul(b,sn) 43 v[i+1] = nfa_qmul(a,sn) + nfa_qmul(b,cs) 44 i=i+2 45 } 46 return 0 47} 48// attention of query q over cache K[0..t],V[0..t] (each dm) -> out(dm). softmax with max-subtract + Q16 exp. 49func attend(q: *i64, Kc: *i64, Vc: *i64, t: i64, dm: i64, scale: i64, out: *i64) -> i64 { 50 let sc: *i64=sys_mmap((t+1)*8) as *i64 51 var s: i64=0; var mx: i64=0-2000000000 52 while s<=t { let d: i64=nfa_qmul(dotq(q, Kc + s*dm as i64, dm), scale); sc[s]=d; if d>mx { mx=d } s=s+1 } 53 var sum: i64=0; s=0 54 while s<=t { let e: i64=nfa_fxexp(sc[s]-mx); sc[s]=e; sum=sum+e; s=s+1 } 55 if sum<1 { sum=1 } 56 var i: i64=0; while i<dm { out[i]=0; i=i+1 } 57 s=0 58 while s<=t { let a: i64=(sc[s]<<16)/sum; var j: i64=0; while j<dm { out[j]=out[j]+nfa_qmul(a, Vc[s*dm+j]); j=j+1 } s=s+1 } 59 return 0 60} 61// compute K[t],V[t] (and Q[t]) for a given position into kk/vv/qq (dm each). opc counts the 3 projections via matvec. 62func project(E: *i64, Wq: *i64, Wk: *i64, Wv: *i64, tok: i64, pos: i64, dm: i64, qq: *i64, kk: *i64, vv: *i64, opc: *i64) -> i64 { 63 let x: *i64=sys_mmap(dm*8) as *i64; copyv(E + tok*dm as i64, x, dm) 64 let xn: *i64=sys_mmap(dm*8) as *i64; rmsnorm(x, dm, xn) 65 matvec(xn, Wq, dm, dm, qq, opc); matvec(xn, Wk, dm, dm, kk, opc); matvec(xn, Wv, dm, dm, vv, opc) 66 rope(qq, pos, dm); rope(kk, pos, dm) 67 return 0 68} 69// finish a position: given x_embed, attn output o -> logits(V). (residual + rmsnorm + Wlm) 70func head(E: *i64, Wo: *i64, Wlm: *i64, tok: i64, o: *i64, dm: i64, V: i64, logits: *i64, opc: *i64) -> i64 { 71 let op: *i64=sys_mmap(dm*8) as *i64; matvec(o, Wo, dm, dm, op, opc) 72 let h: *i64=sys_mmap(dm*8) as *i64; var i: i64=0; while i<dm { h[i]=E[tok*dm+i]+op[i]; i=i+1 } 73 let hn: *i64=sys_mmap(dm*8) as *i64; rmsnorm(h, dm, hn) 74 matvec(hn, Wlm, dm, V, logits, opc) 75 return 0 76} 77 78// CACHED decode: each K/V computed once. logits_out is [T*V]. returns 0; opc accumulates projection MACs. 79func decode_cached(E: *i64, Wq: *i64, Wk: *i64, Wv: *i64, Wo: *i64, Wlm: *i64, seq: *i64, T: i64, dm: i64, V: i64, scale: i64, logits_out: *i64, opc: *i64, corrupt: i64) -> i64 { 80 let Kc: *i64=sys_mmap(T*dm*8) as *i64; let Vc: *i64=sys_mmap(T*dm*8) as *i64 81 let qq: *i64=sys_mmap(dm*8) as *i64; let kk: *i64=sys_mmap(dm*8) as *i64; let vv: *i64=sys_mmap(dm*8) as *i64; let o: *i64=sys_mmap(dm*8) as *i64 82 var t: i64=0 83 while t<T { 84 project(E,Wq,Wk,Wv, seq[t], t, dm, qq, kk, vv, opc) 85 copyv(kk, Kc + t*dm as i64, dm); copyv(vv, Vc + t*dm as i64, dm) 86 if corrupt==1 { if t==1 { Kc[1*dm+0]=0 } } // teeth: stale/zeroed cache entry 87 attend(qq, Kc, Vc, t, dm, scale, o) 88 head(E,Wo,Wlm, seq[t], o, dm, V, logits_out + t*V as i64, opc) 89 t=t+1 90 } 91 return 0 92} 93// FULL-RECOMPUTE decode: at each step re-derive ALL past K/V from scratch (no cache). 94func decode_full(E: *i64, Wq: *i64, Wk: *i64, Wv: *i64, Wo: *i64, Wlm: *i64, seq: *i64, T: i64, dm: i64, V: i64, scale: i64, logits_out: *i64, opc: *i64) -> i64 { 95 let Kc: *i64=sys_mmap(T*dm*8) as *i64; let Vc: *i64=sys_mmap(T*dm*8) as *i64 96 let qq: *i64=sys_mmap(dm*8) as *i64; let kk: *i64=sys_mmap(dm*8) as *i64; let vv: *i64=sys_mmap(dm*8) as *i64; let o: *i64=sys_mmap(dm*8) as *i64 97 var t: i64=0 98 while t<T { 99 var s: i64=0 100 while s<=t { project(E,Wq,Wk,Wv, seq[s], s, dm, qq, kk, vv, opc); copyv(kk, Kc + s*dm as i64, dm); copyv(vv, Vc + s*dm as i64, dm); s=s+1 } 101 attend(qq, Kc, Vc, t, dm, scale, o) // qq currently holds position t's query (last projected) 102 head(E,Wo,Wlm, seq[t], o, dm, V, logits_out + t*V as i64, opc) 103 t=t+1 104 } 105 return 0 106} 107 108func main() -> i64 { 109 g_puts("nx_nofloat_kvcache gate (O(T) incremental decode via KV-cache; correct + cheaper than full recompute)\n" as *u8) 110 let dm: i64=8; let V: i64=6; let T: i64=10; let scale: i64=23170 // 1/sqrt(8)~0.3536 Q16 111 let E: *i64=sys_mmap(V*dm*8) as *i64; dini(E,V*dm,1) 112 let Wq: *i64=sys_mmap(dm*dm*8) as *i64; dini(Wq,dm*dm,2) 113 let Wk: *i64=sys_mmap(dm*dm*8) as *i64; dini(Wk,dm*dm,3) 114 let Wv: *i64=sys_mmap(dm*dm*8) as *i64; dini(Wv,dm*dm,4) 115 let Wo: *i64=sys_mmap(dm*dm*8) as *i64; dini(Wo,dm*dm,5) 116 let Wlm: *i64=sys_mmap(dm*V*8) as *i64; dini(Wlm,dm*V,6) 117 let seq: *i64=sys_mmap(T*8) as *i64; var i: i64=0; while i<T { seq[i]=(i*3+1)%V; i=i+1 } 118 119 let lg_full: *i64=sys_mmap(T*V*8) as *i64; let lg_cache: *i64=sys_mmap(T*V*8) as *i64; let lg_bug: *i64=sys_mmap(T*V*8) as *i64 120 let op_full: *i64=sys_mmap(8) as *i64; let op_cache: *i64=sys_mmap(8) as *i64; let op_bug: *i64=sys_mmap(8) as *i64 121 op_full[0]=0; op_cache[0]=0; op_bug[0]=0 122 decode_full(E,Wq,Wk,Wv,Wo,Wlm, seq, T, dm, V, scale, lg_full, op_full) 123 decode_cached(E,Wq,Wk,Wv,Wo,Wlm, seq, T, dm, V, scale, lg_cache, op_cache, 0) 124 decode_cached(E,Wq,Wk,Wv,Wo,Wlm, seq, T, dm, V, scale, lg_bug, op_bug, 1) 125 126 var same: i64=1; var diff_bug: i64=0; i=0 127 while i<T*V { if lg_full[i]!=lg_cache[i] { same=0 } if lg_full[i]!=lg_bug[i] { diff_bug=1 } i=i+1 } 128 129 g_puts(" [measure] logits match (cache vs full): "); if same==1 { g_puts("IDENTICAL" as *u8) } else { g_puts("DIFFER" as *u8) } 130 g_puts(" proj-MACs full="); g_pn(op_full[0]); g_puts(" cached="); g_pn(op_cache[0]); g_puts("\n") 131 132 var pass: i64=0; var total: i64=0 133 pass=pass+g_check("T1: CORRECT -- cached decode logits == full-recompute logits, bit-exact" as *u8, same); total=total+1 134 var t2: i64=0; if op_cache[0]*2 < op_full[0] { t2=1 } 135 pass=pass+g_check("T2: CHEAPER -- cached does < half the MACs of full recompute (O(T) vs O(T^2) K/V projection)" as *u8, t2); total=total+1 136 pass=pass+g_check("T3: TEETH -- a CORRUPTED cache entry changes the output (the cache is really used)" as *u8, diff_bug); total=total+1 137 138 var okall: i64=0; if pass==total { okall=1 } 139 if okall==1 { 140 let logf: i64=sys_openat_append("knowledge/status/nofloat_kvcache.log" as *u8, 420) 141 if logf>=0 { let x0: i64=sys_write(logf,"NOFLOATKVCACHE incremental decode correct + cheaper measured\n" as *u8,60); sys_close(logf) } 142 } 143 g_puts("---- kvcache gate: passed "); g_pn(pass); g_puts(" / "); g_pn(total); g_puts(" ----\n") 144 // MIGRATED onto nx_gate_verdict by nx_gate_dry_apply (D001, minimal form): every check 145 // row above is untouched, so the PASS/FAIL vector cannot change; only the hand-rolled 146 // verdict emission is replaced by the ONE shared base class. Proven by nx_gate_migrate verify. 147 let ctr__dry: *i64 = gv_ctr() 148 ctr__dry[0] = pass 149 ctr__dry[1] = total 150 let rc__dry: i64 = gv_verdict("NOFLOAT-KVCACHE-GATE" as *u8, ctr__dry, "KV-cache: O(T) incremental decode, bit-exact vs full recompute, fewer ops -- fast sovereign inference)" as *u8) 151 sys_exit(rc__dry) 152 return rc__dry 153}