code wiki / (root) / nx_f32_kv_cache.nx

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}