nx_f32_kv_cache.nx source
↩ module page · 162 lines · 5969 B
1// nx_f32_kv_cache.nx -- bits-up f32 KV cache substrate.
2//
3// L7 / L8 composition brick. Closes the autoregressive-decode
4// O(n^2) → O(n) gap: instead of recomputing K/V from the prompt
5// every step, cache them per layer and only project the new token.
6//
7// Cache layout (linear i64 arrays of f32 bits):
8// cache_K: [n_layers, max_seq_len, kv_dim] row-major
9// cache_V: [n_layers, max_seq_len, kv_dim]
10// kv_dim = n_kv_heads * head_dim
11// seq_len: scalar count of filled rows (advanced once per fwd pass)
12//
13// Per-layer append + advance pattern:
14// For each layer L:
15// 1. Project x -> K_new, V_new (caller does this)
16// 2. nx_f32_kv_cache_append_layer(cache, L, K_new, V_new, n_new)
17// copies K_new/V_new into cache_K[L, seq_len:seq_len+n_new, :]
18// 3. Read K_all = nx_f32_kv_cache_get_K(cache, L)
19// Read V_all = nx_f32_kv_cache_get_V(cache, L)
20// Both are full pointers (caller uses seq_len + n_new rows)
21// 4. Attention over those (n_new + cached) rows
22// Once ALL layers done:
23// nx_f32_kv_cache_advance(cache, n_new)
24// Increments seq_len by n_new.
25//
26// Substrate-honest: this v1 stores f32 outputs (post-projection).
27// Production KV caches sometimes use int8 or even Q4_K -- queued
28// as v2 quantized-KV variant once a quant primitive lands.
29//
30// genealogy_id: standard_kv_cache_canon + decoder_lm_decode_pattern
31// lineage_id: substrate_f32_kv_cache_v1
32
33import "nx_syscalls.nx"
34import "nx_tier.nx"
35
36const NX_F32_KVC_OK: nx_int = 0
37const NX_F32_KVC_ERR_BAD_DIM: nx_int = 1
38const NX_F32_KVC_ERR_OVERFLOW: nx_int = 2
39const NX_F32_KVC_ERR_BAD_LAYER: nx_int = 3
40const NX_F32_KVC_N_VERDICTS: nx_int = 4
41
42func nx_f32_kvc_verdict_is_valid(v: nx_int) -> nx_int {
43 if v < 0 { return 0 }
44 if v >= NX_F32_KVC_N_VERDICTS { return 0 }
45 return 1
46}
47
48struct NxF32KVCache {
49 n_layers: nx_int,
50 n_kv_heads: nx_int,
51 max_seq_len: nx_int,
52 head_dim: nx_int,
53 seq_len: nx_int,
54 cache_K: *i64, // [n_layers * max_seq_len * kv_dim]
55 cache_V: *i64
56}
57
58const NX_F32_KVC_STRUCT_BYTES: nx_int = 56 // 7 fields * 8
59
60// Allocate + initialize a new KV cache.
61
62func nx_f32_kv_cache_alloc(n_layers: nx_int, n_kv_heads: nx_int,
63 max_seq_len: nx_int, head_dim: nx_int) -> *NxF32KVCache {
64 if n_layers <= 0 { return 0 as *NxF32KVCache }
65 if n_kv_heads <= 0 { return 0 as *NxF32KVCache }
66 if max_seq_len <= 0 { return 0 as *NxF32KVCache }
67 if head_dim <= 0 { return 0 as *NxF32KVCache }
68
69 let kv_dim: nx_int = n_kv_heads * head_dim
70 let layer_bytes: i64 = max_seq_len * kv_dim * 8
71 let total_bytes: i64 = n_layers * layer_bytes
72
73 let c: *NxF32KVCache = sys_mmap(NX_F32_KVC_STRUCT_BYTES) as *NxF32KVCache
74 c.n_layers = n_layers
75 c.n_kv_heads = n_kv_heads
76 c.max_seq_len = max_seq_len
77 c.head_dim = head_dim
78 c.seq_len = 0
79 c.cache_K = sys_mmap(total_bytes) as *i64
80 c.cache_V = sys_mmap(total_bytes) as *i64
81 return c
82}
83
84// Reset cache to empty without freeing memory.
85
86func nx_f32_kv_cache_reset(c: *NxF32KVCache) -> nx_int {
87 if c == (0 as *NxF32KVCache) { return NX_F32_KVC_ERR_BAD_DIM }
88 c.seq_len = 0
89 return NX_F32_KVC_OK
90}
91
92func nx_f32_kv_cache_get_seq_len(c: *NxF32KVCache) -> nx_int {
93 return c.seq_len
94}
95
96// Pointer to the start of layer L's K-cache.
97// First (seq_len + freshly-appended-rows) rows are valid kv_dim entries each.
98
99func nx_f32_kv_cache_get_K_layer(c: *NxF32KVCache, layer: nx_int) -> *i64 {
100 if layer < 0 { return 0 as *i64 }
101 if layer >= c.n_layers { return 0 as *i64 }
102 let kv_dim: nx_int = c.n_kv_heads * c.head_dim
103 let layer_offset: i64 = layer * c.max_seq_len * kv_dim * 8
104 let base: i64 = (c.cache_K as i64) + layer_offset
105 return base as *i64
106}
107
108func nx_f32_kv_cache_get_V_layer(c: *NxF32KVCache, layer: nx_int) -> *i64 {
109 if layer < 0 { return 0 as *i64 }
110 if layer >= c.n_layers { return 0 as *i64 }
111 let kv_dim: nx_int = c.n_kv_heads * c.head_dim
112 let layer_offset: i64 = layer * c.max_seq_len * kv_dim * 8
113 let base: i64 = (c.cache_V as i64) + layer_offset
114 return base as *i64
115}
116
117// Append n_new rows of K and V for the given layer at the CURRENT
118// seq_len position. Does NOT advance seq_len -- caller must
119// call nx_f32_kv_cache_advance ONCE after all layers are appended
120// for this forward pass.
121
122func nx_f32_kv_cache_append_layer(c: *NxF32KVCache, layer: nx_int,
123 K_new: *i64, V_new: *i64,
124 n_new: nx_int) -> nx_int {
125 if c == (0 as *NxF32KVCache) { return NX_F32_KVC_ERR_BAD_DIM }
126 if layer < 0 { return NX_F32_KVC_ERR_BAD_LAYER }
127 if layer >= c.n_layers { return NX_F32_KVC_ERR_BAD_LAYER }
128 if n_new <= 0 { return NX_F32_KVC_ERR_BAD_DIM }
129 if c.seq_len + n_new > c.max_seq_len {
130 return NX_F32_KVC_ERR_OVERFLOW
131 }
132
133 let kv_dim: nx_int = c.n_kv_heads * c.head_dim
134 let K_base: *i64 = nx_f32_kv_cache_get_K_layer(c, layer)
135 let V_base: *i64 = nx_f32_kv_cache_get_V_layer(c, layer)
136 let start_idx: nx_int = c.seq_len * kv_dim
137
138 var t: nx_int = 0
139 while t < n_new {
140 var d: nx_int = 0
141 while d < kv_dim {
142 K_base[start_idx + t * kv_dim + d] = K_new[t * kv_dim + d]
143 V_base[start_idx + t * kv_dim + d] = V_new[t * kv_dim + d]
144 d = d + 1
145 }
146 t = t + 1
147 }
148 return NX_F32_KVC_OK
149}
150
151// Advance seq_len by n_new. Caller invokes ONCE per forward pass
152// after all layers have been appended.
153
154func nx_f32_kv_cache_advance(c: *NxF32KVCache, n_new: nx_int) -> nx_int {
155 if c == (0 as *NxF32KVCache) { return NX_F32_KVC_ERR_BAD_DIM }
156 if n_new <= 0 { return NX_F32_KVC_ERR_BAD_DIM }
157 if c.seq_len + n_new > c.max_seq_len {
158 return NX_F32_KVC_ERR_OVERFLOW
159 }
160 c.seq_len = c.seq_len + n_new
161 return NX_F32_KVC_OK
162}