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}