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}