code wiki / (root) / nx_token_sample.nx

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}