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}