code wiki / (root) / nx_embedding.nx

nx_embedding.nx source

↩ module page · 291 lines · 11359 B

1// nx_embedding.nx -- token-ID -> embedding-vector lookup. 2// 3// Closes the INPUT half of the inference pipeline. With this brick 4// shipped, the substrate can do: 5// 6// token_ids -> nx_embedding_lookup -> [n_tok, hidden_dim] tensor 7// (the input to layer 0) 8// 9// Composes only existing primitives: 10// NxTensor (L1 container) 11// NxQuantBlock / Q8 / Q4_K (L2 quantized weight containers) 12// LoopVerdict (control) 13// 14// ===== Embedding table layout ==================================== 15// 16// An embedding table is a [vocab_size, hidden_dim] matrix. Token ID 17// i selects row i. Two storage paths: 18// 19// Dense: *NxTensor with shape [vocab_size, hidden_dim], i64 20// Q10 values. Simple memcpy of row i to out[t, :]. 21// 22// Quantized: *NxQuantBlock (q4_0) or *NxQuantBlockQ8 (q8_0) or 23// *NxQuantQ4K (q4_K). Row-by-row layout: vocab_size 24// blocks of (hidden_dim / BLOCK_SIZE) sub-blocks each. 25// Lookup dequantizes the requested row to Q10. 26// 27// VRAM impact: storing a 50k x 4096 embedding table dense in Q10 is 28// 50000 * 4096 * 8 = 1.6 GB. Quantized at q4_0 is ~150 MB (10.67x 29// compression). Per nx_quant_policy, embeddings stay at f16 by 30// default (small + sensitive), but for memory-constrained deployment 31// a quantized embedding table is a real ~1.5 GB save. 32// 33// ===== Out-of-vocab handling ===================================== 34// 35// Returns NX_EMB_ERR_OOV if any token ID >= vocab_size. Caller 36// MUST validate token IDs upstream (or accept the error verdict 37// and fall back). No silent UNK substitution -- explicit verdict. 38// 39// genealogy_id: mikolov_2013_word2vec + bengio_2003_neural_lm + 40// rumelhart_hinton_williams_1986_distributed_repr 41// lineage_id: substrate_embedding_v1 42 43// nx_safety_envelope: 44// intended_use: AUTO_APPLIED -- primitive-specific tuning queued 45// sil_target: SIL1 46// evidence: [bulk_applied_2026-05-16, see-file-comment-for-detail] 47// verdict: NOT_YET_EVALUATED 48 49import "nx_syscalls.nx" 50import "nx_tier.nx" 51import "nx_loop.nx" 52import "nx_tensor.nx" 53import "nx_quant_block.nx" 54import "nx_quant_block_q8.nx" 55const NX_MAGIC_3000: i64 = 3000 56const NX_MAGIC_3031: i64 = 3031 57const NX_MAGIC_7000: i64 = 7000 58const NX_MAGIC_7015: i64 = 7015 59 60// ===== Sealed-enum: EmbVerdict ==================================== 61 62const NX_EMB_OK: nx_int = 0 63const NX_EMB_ERR_BAD_DIMS: nx_int = 1 64const NX_EMB_ERR_OOV: nx_int = 2 // token >= vocab_size 65const NX_EMB_ERR_BAD_TID: nx_int = 3 // token < 0 66const NX_EMB_ERR_BAD_TABLE: nx_int = 4 67const NX_EMB_N_VERDICTS: nx_int = 5 68 69func nx_emb_verdict_is_valid(v: nx_int) -> nx_int { 70 if v < 0 { return 0 } 71 if v >= NX_EMB_N_VERDICTS { return 0 } 72 return 1 73} 74 75// ===== Dense lookup (i64 Q10 table) ============================== 76// 77// table: *NxTensor [vocab_size, hidden_dim] Q10 78// token_ids:[n_tokens] raw indices 79// out: *NxTensor [n_tokens, hidden_dim] Q10 80// 81// out[t, :] = table[token_ids[t], :] (simple row gather) 82 83func nx_embedding_lookup(table: *NxTensor, token_ids: *i64, n_tokens: nx_int, 84 out: *NxTensor) -> nx_int { 85 if table.dtype != NX_DT_I64 { return NX_EMB_ERR_BAD_TABLE } 86 if out.dtype != NX_DT_I64 { return NX_EMB_ERR_BAD_TABLE } 87 if table.ndim != 2 { return NX_EMB_ERR_BAD_DIMS } 88 if out.ndim != 2 { return NX_EMB_ERR_BAD_DIMS } 89 let vocab: nx_int = table.shape[0] 90 let hidden: nx_int = table.shape[1] 91 if out.shape[0] != n_tokens { return NX_EMB_ERR_BAD_DIMS } 92 if out.shape[1] != hidden { return NX_EMB_ERR_BAD_DIMS } 93 if nx_t_is_contiguous(table) == 0 { return NX_EMB_ERR_BAD_TABLE } 94 if nx_t_is_contiguous(out) == 0 { return NX_EMB_ERR_BAD_DIMS } 95 96 let pt: *i64 = table.storage as *i64 97 let po: *i64 = out.storage as *i64 98 99 var t: nx_int = 0 100 var iter: nx_int = 0 101 var verdict: nx_int = NX_LOOP_RUNNING 102 let BUDGET: nx_int = n_tokens 103 while verdict == NX_LOOP_RUNNING && iter < BUDGET { 104 let tid: nx_int = token_ids[t] 105 if tid < 0 { verdict = NX_LOOP_ABORTED } 106 if tid >= vocab { verdict = NX_LOOP_ABORTED } 107 if verdict == NX_LOOP_RUNNING { 108 let src_base: nx_int = tid * hidden 109 let dst_base: nx_int = t * hidden 110 var c: nx_int = 0 111 var c_iter: nx_int = 0 112 var c_verdict: nx_int = NX_LOOP_RUNNING 113 let C_BUDGET: nx_int = hidden 114 while c_verdict == NX_LOOP_RUNNING && c_iter < C_BUDGET { 115 po[dst_base + c] = pt[src_base + c] 116 c = c + 1 117 c_iter = c_iter + 1 118 } 119 } 120 t = t + 1 121 iter = iter + 1 122 } 123 if verdict == NX_LOOP_ABORTED { return NX_EMB_ERR_OOV } 124 return NX_EMB_OK 125} 126 127// ===== Q8_0 dequantizing lookup ================================== 128// 129// table: *NxQuantBlockQ8 row-major flat blocks 130// vocab: nx_int number of embedding rows 131// hidden: nx_int dimension per row 132// token_ids: [n_tokens] 133// out: *i64 flat [n_tokens * hidden] output 134// 135// For each requested token, walk the appropriate range of blocks 136// (hidden / BLOCK_SIZE blocks per embedding row) and dequantize 137// to Q10. 138 139func nx_embedding_lookup_q8(table: *NxQuantBlockQ8, vocab: nx_int, hidden: nx_int, 140 token_ids: *i64, n_tokens: nx_int, out: *i64) -> nx_int { 141 let blocks_per_row: nx_int = hidden / NX_QB8_BLOCK_SIZE 142 if blocks_per_row * NX_QB8_BLOCK_SIZE != hidden { return NX_EMB_ERR_BAD_DIMS } 143 if table.n_values < vocab * hidden { return NX_EMB_ERR_BAD_TABLE } 144 145 var t: nx_int = 0 146 var iter: nx_int = 0 147 var verdict: nx_int = NX_LOOP_RUNNING 148 let BUDGET: nx_int = n_tokens 149 while verdict == NX_LOOP_RUNNING && iter < BUDGET { 150 let tid: nx_int = token_ids[t] 151 if tid < 0 { verdict = NX_LOOP_ABORTED } 152 if tid >= vocab { verdict = NX_LOOP_ABORTED } 153 if verdict == NX_LOOP_RUNNING { 154 // Dequantize the tid-th row's blocks_per_row blocks into out[t*hidden..]. 155 let dst_base: nx_int = t * hidden 156 let row_block_start: nx_int = tid * blocks_per_row 157 var b: nx_int = 0 158 var b_iter: nx_int = 0 159 var b_verdict: nx_int = NX_LOOP_RUNNING 160 let B_BUDGET: nx_int = blocks_per_row 161 while b_verdict == NX_LOOP_RUNNING && b_iter < B_BUDGET { 162 let block_id: nx_int = row_block_start + b 163 let scale: i64 = table.scales[block_id] 164 let block_start_val: nx_int = block_id * NX_QB8_BLOCK_SIZE 165 var v: nx_int = 0 166 var v_iter: nx_int = 0 167 var v_verdict: nx_int = NX_LOOP_RUNNING 168 let V_BUDGET: nx_int = NX_QB8_BLOCK_SIZE 169 while v_verdict == NX_LOOP_RUNNING && v_iter < V_BUDGET { 170 let u_byte: nx_int = table.packed[block_start_val + v] as nx_int 171 let signed_byte: nx_int = u_byte - NX_QB8_BYTE_OFFSET 172 out[dst_base + b * NX_QB8_BLOCK_SIZE + v] = signed_byte * scale 173 v = v + 1 174 v_iter = v_iter + 1 175 } 176 b = b + 1 177 b_iter = b_iter + 1 178 } 179 } 180 t = t + 1 181 iter = iter + 1 182 } 183 if verdict == NX_LOOP_ABORTED { return NX_EMB_ERR_OOV } 184 return NX_EMB_OK 185} 186 187// ===== Self-test ================================================== 188// 189// Closed-form invariants: 190// 191// (a) Dense lookup: a known table + ID sequence yields the exact 192// row values bit-equal. 193// (b) OOV detection: token_id == vocab triggers NX_EMB_ERR_OOV. 194// (c) Negative ID detection: token_id == -1 triggers NX_EMB_ERR_OOV. 195// (d) Q8 dequant round-trip: encode a small vocab x hidden table, 196// lookup token 2, verify dequantized values equal scale * stored. 197// (e) Verdict gate. 198 199func main() -> i64 { 200 let vocab: nx_int = 8 201 let hidden: nx_int = 32 // = 1 block @ Q8 block_size 32 202 203 // --- Dense table setup --- 204 let sh: *nx_int = sys_mmap(2 * 8) as *nx_int 205 sh[0] = vocab; sh[1] = hidden 206 let err: *nx_int = sys_mmap(8) as *nx_int 207 err[0] = 0 208 let table: *NxTensor = nx_t_alloc(NX_DT_I64, sh, 2, err) 209 if err[0] != 0 { return 5 } 210 let pt: *i64 = table.storage as *i64 211 // Fill table[i, j] = i * 1000 + j. 212 var i: nx_int = 0 213 while i < vocab { 214 var j: nx_int = 0 215 while j < hidden { 216 pt[i * hidden + j] = i * 1000 + j 217 j = j + 1 218 } 219 i = i + 1 220 } 221 222 // --- (a) Dense lookup --- 223 let n_tok: nx_int = 4 224 let token_ids: *i64 = sys_mmap(n_tok * 8) as *i64 225 token_ids[0] = 0; token_ids[1] = 3; token_ids[2] = 7; token_ids[3] = 1 226 227 let out_sh: *nx_int = sys_mmap(2 * 8) as *nx_int 228 out_sh[0] = n_tok; out_sh[1] = hidden 229 let out: *NxTensor = nx_t_alloc(NX_DT_I64, out_sh, 2, err) 230 if err[0] != 0 { return 6 } 231 232 let v_lookup: nx_int = nx_embedding_lookup(table, token_ids, n_tok, out) 233 if v_lookup != NX_EMB_OK { return 10 + v_lookup } 234 235 let po: *i64 = out.storage as *i64 236 // Verify token 0 -> row 0 -> values 0..31 237 if po[0] != 0 { return 20 } 238 if po[31] != 31 { return 21 } 239 // Token 1 -> row 3 -> values 3000..3031 240 if po[hidden + 0] != NX_MAGIC_3000 { return 22 } 241 if po[hidden + 31] != NX_MAGIC_3031 { return 23 } 242 // Token 2 -> row 7 243 if po[2 * hidden + 0] != NX_MAGIC_7000 { return 24 } 244 if po[2 * hidden + 15] != NX_MAGIC_7015 { return 25 } 245 // Token 3 -> row 1 246 if po[3 * hidden + 0] != 1000 { return 26 } 247 248 // --- (b) OOV: token == vocab --- 249 let bad: *i64 = sys_mmap(8) as *i64 250 bad[0] = vocab // OOV 251 let v_oov: nx_int = nx_embedding_lookup(table, bad, 1, out) 252 if v_oov != NX_EMB_ERR_OOV { return 30 } 253 254 // --- (c) Negative ID --- 255 bad[0] = -1 256 let v_neg: nx_int = nx_embedding_lookup(table, bad, 1, out) 257 if v_neg != NX_EMB_ERR_OOV { return 40 } 258 259 // --- (d) Q8 dequant lookup --- 260 // Build a flat i64 buffer matching the table, then encode block- 261 // by-block via nx_qb8_quantize. 262 let flat: *i64 = sys_mmap(vocab * hidden * 8) as *i64 263 var fi: nx_int = 0 264 while fi < vocab { 265 var fj: nx_int = 0 266 while fj < hidden { 267 flat[fi * hidden + fj] = (fi + 1) * 100 // each row constant 268 fj = fj + 1 269 } 270 fi = fi + 1 271 } 272 let qb: *NxQuantBlockQ8 = nx_qb8_alloc(vocab * hidden) 273 nx_qb8_quantize(flat, vocab * hidden, qb) 274 275 let q_out: *i64 = sys_mmap(n_tok * hidden * 8) as *i64 276 let v_q: nx_int = nx_embedding_lookup_q8(qb, vocab, hidden, token_ids, n_tok, q_out) 277 if v_q != NX_EMB_OK { return 50 + v_q } 278 // Token 0 -> row 0 -> all 100s. Allow scale rounding. 279 if q_out[0] != 100 { return 60 } 280 // Token 1 -> row 3 -> all 400s. 281 if q_out[hidden] != 400 { return 70 } 282 283 // --- (e) Verdict gate --- 284 var vi: nx_int = 0 285 while vi < NX_EMB_N_VERDICTS { 286 if nx_emb_verdict_is_valid(vi) != 1 { return 80 + vi } 287 vi = vi + 1 288 } 289 290 return 0 291}