code wiki / _hdl_build / nx_nofloat_kvcache_gate.nx

nx_nofloat_kvcache_gate.nx source

↩ module page · 145 lines · 8897 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" 12const Q16: i64 = 65536 13 14 15func 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 } 16 17// row-vector(x, len=rows) times matrix(W, rows x cols) -> out(cols); Q16 accumulate-then-shift. opc += rows*cols. 18func matvec(x: *i64, W: *i64, rows: i64, cols: i64, out: *i64, opc: *i64) -> i64 { 19 var c: i64=0 20 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 } 21 opc[0]=opc[0]+rows*cols 22 return 0 23} 24func rmsnorm(x: *i64, n: i64, out: *i64) -> i64 { 25 var ss: i64=0; var i: i64=0; while i<n { ss=ss + x[i]*x[i]; i=i+1 } // Q32 sum of squares 26 var ms: i64=ss/n 27 var rms: i64=nfa_isqrt(ms) // sqrt(Q32) = Q16 rms 28 if rms<1 { rms=1 } 29 i=0; while i<n { out[i]=(x[i]<<16)/rms; i=i+1 } 30 return 0 31} 32func 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 } 33func copyv(src: *i64, dst: *i64, n: i64) -> i64 { var i: i64=0; while i<n { dst[i]=src[i]; i=i+1 } return 0 } 34// simple consistent RoPE: rotate pairs by angle = pos * (Q16>>(i+1)) (form is irrelevant; consistency is what matters) 35func rope(v: *i64, pos: i64, dm: i64) -> i64 { 36 var i: i64=0 37 while i+1<dm { 38 let ang: i64 = pos * (32768 >> (i/2 + 1)) 39 let cs: i64 = nfa_cosf(ang); let sn: i64 = nfa_sinf(ang) 40 let a: i64=v[i]; let b: i64=v[i+1] 41 v[i] = nfa_qmul(a,cs) - nfa_qmul(b,sn) 42 v[i+1] = nfa_qmul(a,sn) + nfa_qmul(b,cs) 43 i=i+2 44 } 45 return 0 46} 47// attention of query q over cache K[0..t],V[0..t] (each dm) -> out(dm). softmax with max-subtract + Q16 exp. 48func attend(q: *i64, Kc: *i64, Vc: *i64, t: i64, dm: i64, scale: i64, out: *i64) -> i64 { 49 let sc: *i64=sys_mmap((t+1)*8) as *i64 50 var s: i64=0; var mx: i64=0-2000000000 51 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 } 52 var sum: i64=0; s=0 53 while s<=t { let e: i64=nfa_fxexp(sc[s]-mx); sc[s]=e; sum=sum+e; s=s+1 } 54 if sum<1 { sum=1 } 55 var i: i64=0; while i<dm { out[i]=0; i=i+1 } 56 s=0 57 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 } 58 return 0 59} 60// 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. 61func project(E: *i64, Wq: *i64, Wk: *i64, Wv: *i64, tok: i64, pos: i64, dm: i64, qq: *i64, kk: *i64, vv: *i64, opc: *i64) -> i64 { 62 let x: *i64=sys_mmap(dm*8) as *i64; copyv(E + tok*dm as i64, x, dm) 63 let xn: *i64=sys_mmap(dm*8) as *i64; rmsnorm(x, dm, xn) 64 matvec(xn, Wq, dm, dm, qq, opc); matvec(xn, Wk, dm, dm, kk, opc); matvec(xn, Wv, dm, dm, vv, opc) 65 rope(qq, pos, dm); rope(kk, pos, dm) 66 return 0 67} 68// finish a position: given x_embed, attn output o -> logits(V). (residual + rmsnorm + Wlm) 69func head(E: *i64, Wo: *i64, Wlm: *i64, tok: i64, o: *i64, dm: i64, V: i64, logits: *i64, opc: *i64) -> i64 { 70 let op: *i64=sys_mmap(dm*8) as *i64; matvec(o, Wo, dm, dm, op, opc) 71 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 } 72 let hn: *i64=sys_mmap(dm*8) as *i64; rmsnorm(h, dm, hn) 73 matvec(hn, Wlm, dm, V, logits, opc) 74 return 0 75} 76 77// CACHED decode: each K/V computed once. logits_out is [T*V]. returns 0; opc accumulates projection MACs. 78func 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 { 79 let Kc: *i64=sys_mmap(T*dm*8) as *i64; let Vc: *i64=sys_mmap(T*dm*8) as *i64 80 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 81 var t: i64=0 82 while t<T { 83 project(E,Wq,Wk,Wv, seq[t], t, dm, qq, kk, vv, opc) 84 copyv(kk, Kc + t*dm as i64, dm); copyv(vv, Vc + t*dm as i64, dm) 85 if corrupt==1 { if t==1 { Kc[1*dm+0]=0 } } // teeth: stale/zeroed cache entry 86 attend(qq, Kc, Vc, t, dm, scale, o) 87 head(E,Wo,Wlm, seq[t], o, dm, V, logits_out + t*V as i64, opc) 88 t=t+1 89 } 90 return 0 91} 92// FULL-RECOMPUTE decode: at each step re-derive ALL past K/V from scratch (no cache). 93func 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 { 94 let Kc: *i64=sys_mmap(T*dm*8) as *i64; let Vc: *i64=sys_mmap(T*dm*8) as *i64 95 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 96 var t: i64=0 97 while t<T { 98 var s: i64=0 99 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 } 100 attend(qq, Kc, Vc, t, dm, scale, o) // qq currently holds position t's query (last projected) 101 head(E,Wo,Wlm, seq[t], o, dm, V, logits_out + t*V as i64, opc) 102 t=t+1 103 } 104 return 0 105} 106 107func main() -> i64 { 108 g_puts("nx_nofloat_kvcache gate (O(T) incremental decode via KV-cache; correct + cheaper than full recompute)\n" as *u8) 109 let dm: i64=8; let V: i64=6; let T: i64=10; let scale: i64=23170 // 1/sqrt(8)~0.3536 Q16 110 let E: *i64=sys_mmap(V*dm*8) as *i64; dini(E,V*dm,1) 111 let Wq: *i64=sys_mmap(dm*dm*8) as *i64; dini(Wq,dm*dm,2) 112 let Wk: *i64=sys_mmap(dm*dm*8) as *i64; dini(Wk,dm*dm,3) 113 let Wv: *i64=sys_mmap(dm*dm*8) as *i64; dini(Wv,dm*dm,4) 114 let Wo: *i64=sys_mmap(dm*dm*8) as *i64; dini(Wo,dm*dm,5) 115 let Wlm: *i64=sys_mmap(dm*V*8) as *i64; dini(Wlm,dm*V,6) 116 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 } 117 118 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 119 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 120 op_full[0]=0; op_cache[0]=0; op_bug[0]=0 121 decode_full(E,Wq,Wk,Wv,Wo,Wlm, seq, T, dm, V, scale, lg_full, op_full) 122 decode_cached(E,Wq,Wk,Wv,Wo,Wlm, seq, T, dm, V, scale, lg_cache, op_cache, 0) 123 decode_cached(E,Wq,Wk,Wv,Wo,Wlm, seq, T, dm, V, scale, lg_bug, op_bug, 1) 124 125 var same: i64=1; var diff_bug: i64=0; i=0 126 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 } 127 128 g_puts(" [measure] logits match (cache vs full): "); if same==1 { g_puts("IDENTICAL" as *u8) } else { g_puts("DIFFER" as *u8) } 129 g_puts(" proj-MACs full="); g_pn(op_full[0]); g_puts(" cached="); g_pn(op_cache[0]); g_puts("\n") 130 131 var pass: i64=0; var total: i64=0 132 pass=pass+g_check("T1: CORRECT -- cached decode logits == full-recompute logits, bit-exact" as *u8, same); total=total+1 133 var t2: i64=0; if op_cache[0]*2 < op_full[0] { t2=1 } 134 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 135 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 136 137 var okall: i64=0; if pass==total { okall=1 } 138 if okall==1 { 139 let logf: i64=sys_openat_append("knowledge/status/nofloat_kvcache.log" as *u8, 420) 140 if logf>=0 { let x0: i64=sys_write(logf,"NOFLOATKVCACHE incremental decode correct + cheaper measured\n" as *u8,60); sys_close(logf) } 141 } 142 g_puts("---- kvcache gate: passed "); g_pn(pass); g_puts(" / "); g_pn(total); g_puts(" ----\n") 143 if okall==1 { g_puts("verdict=GREEN (KV-cache: O(T) incremental decode, bit-exact vs full recompute, fewer ops -- fast sovereign inference)\n" as *u8); sys_exit(0); return 0 } 144 g_puts("verdict=RED\n" as *u8); sys_exit(1); return 1 145}