nx_token_sample.nx source
↩ module page · 318 lines · 10951 B
1// nx_token_sample.nx -- logit sampling primitives for inference.
2//
3// Closes the OUTPUT half of the transformer inference loop.
4// Composes against shipped primitives -- no new math:
5//
6// nx_prng (canonical RNG)
7// nx_attention._attn_exp_q10 (numerically-stable softmax)
8// Standard top-K mask + categorical sample math
9//
10// The typical decoder-only LLM inference loop:
11//
12// for each generation step:
13// logits = forward_pass(...)[-1, :] // last token, vocab-wide
14// nx_logit_apply_temperature(logits, vocab, temp_q10)
15// nx_logit_top_k_mask(logits, vocab, K)
16// softmax(logits) // (caller via nx_attn_softmax_row_q10
17// // reshaping to [1, vocab])
18// next_token = nx_sample_categorical(probs, vocab, prng_state)
19//
20// Top-P (nucleus) sampling queued -- needs a sort primitive (or
21// partial-sort) which isn't shipped yet. Per the no-skipping
22// cardinal, top-P lands after nx_sort.
23//
24// ===== Math =======================================================
25//
26// Temperature: logit_i' = logit_i * Q10 / temp_q10
27// * temp = Q10 (1.0) -> unchanged
28// * temp = Q10/2 (0.5) -> sharper distribution (sampled = more deterministic)
29// * temp = 2*Q10 (2.0) -> flatter distribution (more random)
30//
31// Top-K masking: keep K highest logits, set rest to NEG_INF so they
32// get 0 weight in subsequent softmax. Implemented via repeated
33// max-find -- O(n*k); for K=40 on 50k vocab that's 2M ops, fine.
34//
35// Categorical sample: given probs in Q10 that nominally sum to Q10,
36// draw u in [0, sum_q10) and walk CDF. Returns index where the
37// CDF crosses u.
38//
39// Per the bounded-loop + bits-up cardinals.
40//
41// genealogy_id: holtzman_2020_nucleus_sampling + fan_2018_top_k +
42// ackley_hinton_sejnowski_1985_simulated_annealing_temperature
43// lineage_id: substrate_token_sample_v1
44
45// nx_safety_envelope:
46// intended_use: AUTO_APPLIED -- primitive-specific tuning queued
47// sil_target: SIL1
48// evidence: [bulk_applied_2026-05-16, see-file-comment-for-detail]
49// verdict: NOT_YET_EVALUATED
50
51import "nx_syscalls.nx"
52import "nx_tier.nx"
53import "nx_loop.nx"
54import "nx_prng.nx"
55
56const NX_TS_Q10: nx_int = 1024
57const NX_TS_NEG_INF: nx_int = -1000000000
58
59// ===== Sealed-enum: SampleVerdict =================================
60
61const NX_TS_OK: nx_int = 0
62const NX_TS_ERR_BAD_VOCAB: nx_int = 1
63const NX_TS_ERR_BAD_TEMP: nx_int = 2
64const NX_TS_ERR_BAD_K: nx_int = 3
65const NX_TS_ERR_DEGENERATE: nx_int = 4 // all probs zero
66const NX_TS_N_VERDICTS: nx_int = 5
67
68func nx_ts_verdict_is_valid(v: nx_int) -> nx_int {
69 if v < 0 { return 0 }
70 if v >= NX_TS_N_VERDICTS { return 0 }
71 return 1
72}
73
74// ===== Temperature scaling =======================================
75//
76// In-place: logits[i] = logits[i] * Q10 / temp_q10.
77//
78// At temp = Q10 (1.0), result == input. Smaller temp -> sharper.
79
80func nx_logit_apply_temperature(logits: *i64, vocab: nx_int, temp_q10: nx_int) -> nx_int {
81 if vocab <= 0 { return NX_TS_ERR_BAD_VOCAB }
82 if temp_q10 <= 0 { return NX_TS_ERR_BAD_TEMP }
83 var i: nx_int = 0
84 var iter: nx_int = 0
85 var verdict: nx_int = NX_LOOP_RUNNING
86 let BUDGET: nx_int = vocab
87 while verdict == NX_LOOP_RUNNING && iter < BUDGET {
88 logits[i] = (logits[i] * NX_TS_Q10) / temp_q10
89 i = i + 1
90 iter = iter + 1
91 }
92 return NX_TS_OK
93}
94
95// ===== Top-K masking =============================================
96//
97// Find the K highest values, set everything else to NEG_INF.
98// O(n*k) -- iterate K times, each time finding + marking the
99// current max, then re-scanning to "skip" already-marked entries.
100//
101// Implementation: use a parallel `kept` array (boolean) and walk K
102// passes. Caller-friendly: we allocate kept[] internally via
103// sys_mmap.
104
105func nx_logit_top_k_mask(logits: *i64, vocab: nx_int, k: nx_int) -> nx_int {
106 if vocab <= 0 { return NX_TS_ERR_BAD_VOCAB }
107 if k <= 0 { return NX_TS_ERR_BAD_K }
108 if k >= vocab { return NX_TS_OK } // nothing to mask
109
110 let kept: *i64 = sys_mmap(vocab * 8) as *i64
111 var z: nx_int = 0
112 while z < vocab { kept[z] = 0; z = z + 1 }
113
114 // Pick the top-K via K passes of max-find.
115 var pick: nx_int = 0
116 var pick_iter: nx_int = 0
117 var pick_verdict: nx_int = NX_LOOP_RUNNING
118 let PICK_BUDGET: nx_int = k
119 while pick_verdict == NX_LOOP_RUNNING && pick_iter < PICK_BUDGET {
120 var max_v: i64 = NX_TS_NEG_INF
121 var max_i: nx_int = -1
122 var j: nx_int = 0
123 var j_iter: nx_int = 0
124 var j_verdict: nx_int = NX_LOOP_RUNNING
125 let J_BUDGET: nx_int = vocab
126 while j_verdict == NX_LOOP_RUNNING && j_iter < J_BUDGET {
127 if kept[j] == 0 {
128 if logits[j] > max_v {
129 max_v = logits[j]
130 max_i = j
131 }
132 }
133 j = j + 1
134 j_iter = j_iter + 1
135 }
136 if max_i < 0 { pick_verdict = NX_LOOP_DONE_EXIT }
137 if pick_verdict == NX_LOOP_RUNNING {
138 kept[max_i] = 1
139 }
140 pick = pick + 1
141 pick_iter = pick_iter + 1
142 }
143
144 // Mask everything not kept.
145 var m: nx_int = 0
146 var m_iter: nx_int = 0
147 var m_verdict: nx_int = NX_LOOP_RUNNING
148 let M_BUDGET: nx_int = vocab
149 while m_verdict == NX_LOOP_RUNNING && m_iter < M_BUDGET {
150 if kept[m] == 0 {
151 logits[m] = NX_TS_NEG_INF
152 }
153 m = m + 1
154 m_iter = m_iter + 1
155 }
156 return NX_TS_OK
157}
158
159// ===== Categorical sample ========================================
160//
161// Given probs[0..vocab) in Q10 that nominally sum to Q10, draw a
162// random Q10 value u in [0, sum) and return the index where the
163// cumulative sum first exceeds u.
164//
165// Robustness: we compute the actual sum first (probs may not be
166// exactly Q10 due to rounding). Argmax-fallback if sum == 0.
167
168func nx_sample_categorical(probs: *i64, vocab: nx_int, prng_state: *i64) -> nx_int {
169 if vocab <= 0 { return 0 }
170
171 // Total mass.
172 var sum: i64 = 0
173 var i: nx_int = 0
174 var iter: nx_int = 0
175 var verdict: nx_int = NX_LOOP_RUNNING
176 let BUDGET: nx_int = vocab
177 while verdict == NX_LOOP_RUNNING && iter < BUDGET {
178 if probs[i] > 0 { sum = sum + probs[i] }
179 i = i + 1
180 iter = iter + 1
181 }
182 if sum <= 0 {
183 // Degenerate: return argmax as the safe fallback.
184 var max_v: i64 = NX_TS_NEG_INF
185 var max_i: nx_int = 0
186 var k: nx_int = 0
187 while k < vocab {
188 if probs[k] > max_v { max_v = probs[k]; max_i = k }
189 k = k + 1
190 }
191 return max_i
192 }
193
194 // Draw u in [0, sum).
195 let u: i64 = nx_prng_range(prng_state, sum)
196 var acc: i64 = 0
197 var j: nx_int = 0
198 var j_iter: nx_int = 0
199 var j_verdict: nx_int = NX_LOOP_RUNNING
200 while j_verdict == NX_LOOP_RUNNING && j_iter < BUDGET {
201 if probs[j] > 0 { acc = acc + probs[j] }
202 if acc > u { j_verdict = NX_LOOP_DONE_EXIT }
203 if j_verdict == NX_LOOP_RUNNING {
204 j = j + 1
205 j_iter = j_iter + 1
206 }
207 }
208 // j is the chosen index (or vocab-1 if we somehow rounded past).
209 if j >= vocab { return vocab - 1 }
210 return j
211}
212
213// ===== Self-test ==================================================
214//
215// (a) Temperature = Q10 leaves logits unchanged (within bits).
216// (b) Temperature = Q10 / 2 doubles each logit.
217// (c) Top-K keeps exactly K non-NEG_INF entries.
218// (d) Top-K preserves the K largest values.
219// (e) Categorical sample on a one-hot distribution always returns
220// that index.
221// (f) Categorical sample on uniform distribution covers all indices
222// across many draws (statistical, weak check).
223
224func main() -> i64 {
225 let vocab: nx_int = 16
226 let logits: *i64 = sys_mmap(vocab * 8) as *i64
227 let backup: *i64 = sys_mmap(vocab * 8) as *i64
228
229 // Seed logits with i (0..15).
230 var i: nx_int = 0
231 while i < vocab { logits[i] = i; backup[i] = i; i = i + 1 }
232
233 // --- (a) temp = Q10 -> identity ---
234 let v_t1: nx_int = nx_logit_apply_temperature(logits, vocab, NX_TS_Q10)
235 if v_t1 != NX_TS_OK { return 10 + v_t1 }
236 var c: nx_int = 0
237 while c < vocab {
238 if logits[c] != backup[c] { return 20 }
239 c = c + 1
240 }
241
242 // --- (b) temp = Q10/2 -> doubles ---
243 nx_logit_apply_temperature(logits, vocab, NX_TS_Q10 / 2)
244 var c2: nx_int = 0
245 while c2 < vocab {
246 if logits[c2] != 2 * backup[c2] { return 30 }
247 c2 = c2 + 1
248 }
249
250 // Restore baseline before top-K test.
251 var r: nx_int = 0
252 while r < vocab { logits[r] = r; r = r + 1 }
253
254 // --- (c, d) Top-K keeps 4 -- highest values are 12,13,14,15 ---
255 let v_tk: nx_int = nx_logit_top_k_mask(logits, vocab, 4)
256 if v_tk != NX_TS_OK { return 40 + v_tk }
257 var kept_count: nx_int = 0
258 var ck: nx_int = 0
259 while ck < vocab {
260 if logits[ck] != NX_TS_NEG_INF { kept_count = kept_count + 1 }
261 ck = ck + 1
262 }
263 if kept_count != 4 { return 50 }
264 // The four kept must be exactly indices 12, 13, 14, 15.
265 if logits[12] == NX_TS_NEG_INF { return 51 }
266 if logits[13] == NX_TS_NEG_INF { return 52 }
267 if logits[14] == NX_TS_NEG_INF { return 53 }
268 if logits[15] == NX_TS_NEG_INF { return 54 }
269 if logits[11] != NX_TS_NEG_INF { return 55 }
270 if logits[0] != NX_TS_NEG_INF { return 56 }
271
272 // --- (e) One-hot categorical sample ---
273 let probs: *i64 = sys_mmap(vocab * 8) as *i64
274 var p: nx_int = 0
275 while p < vocab { probs[p] = 0; p = p + 1 }
276 probs[7] = NX_TS_Q10
277 let prng: *i64 = sys_mmap(8) as *i64
278 nx_prng_init(prng, 0xdeadbeef)
279 var ti: nx_int = 0
280 while ti < 16 {
281 let idx: nx_int = nx_sample_categorical(probs, vocab, prng)
282 if idx != 7 { return 60 + ti }
283 ti = ti + 1
284 }
285
286 // --- (f) Uniform distribution covers (at least) several indices ---
287 var u: nx_int = 0
288 while u < vocab { probs[u] = NX_TS_Q10 / 4; u = u + 1 }
289 let coverage: *i64 = sys_mmap(vocab * 8) as *i64
290 var cov: nx_int = 0
291 while cov < vocab { coverage[cov] = 0; cov = cov + 1 }
292 var n: nx_int = 0
293 while n < 200 {
294 let idx: nx_int = nx_sample_categorical(probs, vocab, prng)
295 if idx >= 0 {
296 if idx < vocab { coverage[idx] = coverage[idx] + 1 }
297 }
298 n = n + 1
299 }
300 // Expect at least half the vocab indices to be hit at least once
301 // across 200 uniform draws.
302 var hits: nx_int = 0
303 var h: nx_int = 0
304 while h < vocab {
305 if coverage[h] > 0 { hits = hits + 1 }
306 h = h + 1
307 }
308 if hits < 8 { return 80 }
309
310 // --- (g) Verdict gate ---
311 var vi: nx_int = 0
312 while vi < NX_TS_N_VERDICTS {
313 if nx_ts_verdict_is_valid(vi) != 1 { return 90 + vi }
314 vi = vi + 1
315 }
316
317 return 0
318}