code wiki / (root) / nx_f32_sampler.nx

nx_f32_sampler.nx source

↩ module page · 344 lines · 12038 B

1// nx_f32_sampler.nx -- temperature + top-k sampler for LLM logits. 2// 3// Production samplers beyond greedy. Composes: 4// nx_f32_lt IEEE 754 compare 5// nx_f32_sub logits shift for numerical stability 6// nx_f32_mul temperature scaling 7// nx_f32_exp softmax exp 8// nx_f32_add cumulative sum 9// nx_q14_to_f32 PRNG uniform conversion 10// nx_prng_uniform_q14 Q14 uniform [0, 16384) 11// 12// Algorithms (canonical): 13// Temperature scaling: Hinton 2015 (Knowledge Distillation) 14// Categorical sampling: Wong 1980 (inverse CDF / Walker's alias) 15// Top-k sampling: Fan 2018 (Hierarchical Neural Story Gen.) 16// 17// genealogy_id: hinton_2015_temperature + fan_2018_top_k + 18// bridle_1990_softmax + wong_1980_inverse_cdf 19// lineage_id: substrate_f32_sampler_v1 20 21import "nx_syscalls.nx" 22import "nx_tier.nx" 23import "nx_prng.nx" 24import "nx_f32.nx" 25import "nx_f32_div.nx" 26import "nx_f32_cvt.nx" 27import "nx_f32_exp.nx" 28const NX_MAGIC_16384: i64 = 16384 29 30const NX_FSP_OK: nx_int = 0 31const NX_FSP_ERR_BAD_DIM: nx_int = 1 32const NX_FSP_ERR_NULL: nx_int = 2 33const NX_FSP_N_VERDICTS: nx_int = 3 34 35func nx_fsp_verdict_is_valid(v: nx_int) -> nx_int { 36 if v < 0 { return 0 } 37 if v >= NX_FSP_N_VERDICTS { return 0 } 38 return 1 39} 40 41// Find argmax over vocab_size logits. 42 43func nx_f32_sampler_argmax(logits: *i64, vocab_size: nx_int) -> nx_int { 44 if vocab_size <= 0 { return 0 } 45 var best_id: nx_int = 0 46 var best_val: i64 = logits[0] 47 var i: nx_int = 1 48 while i < vocab_size { 49 if nx_f32_lt(best_val, logits[i]) != 0 { 50 best_val = logits[i] 51 best_id = i 52 } 53 i = i + 1 54 } 55 return best_id 56} 57 58// Temperature sampler. inv_temp_f32 = 1/temperature in f32 bits. 59// For temp=1.0, pass 0x3F800000. For temp=0.5, pass 0x40000000. 60// As temp -> 0, sampling collapses to argmax. 61 62func nx_f32_sampler_sample_temp(logits: *i64, vocab_size: nx_int, 63 inv_temp_f32: i64, 64 prng_state: *i64) -> nx_int { 65 if vocab_size <= 0 { return 0 } 66 67 // Find max for stability. 68 var max_l: i64 = logits[0] 69 var i: nx_int = 1 70 while i < vocab_size { 71 if nx_f32_lt(max_l, logits[i]) != 0 { max_l = logits[i] } 72 i = i + 1 73 } 74 75 // Compute exp((logits - max) * inv_temp) and accumulate sum. 76 let scaled: *i64 = sys_mmap(vocab_size * 8) as *i64 77 var sum: i64 = 0 78 var j: nx_int = 0 79 while j < vocab_size { 80 let shifted: i64 = nx_f32_sub(logits[j], max_l) 81 let with_t: i64 = nx_f32_mul(shifted, inv_temp_f32) 82 let e: i64 = nx_f32_exp(with_t) 83 scaled[j] = e 84 sum = nx_f32_add(sum, e) 85 j = j + 1 86 } 87 88 // Sample u = r * sum where r in [0, 1). 89 let r_q14: i64 = nx_prng_uniform_q14(prng_state) 90 let r_f32: i64 = nx_q14_to_f32(r_q14) 91 let u: i64 = nx_f32_mul(r_f32, sum) 92 93 // Walk cumulative to find bucket. 94 var acc: i64 = 0 95 var k: nx_int = 0 96 while k < vocab_size { 97 acc = nx_f32_add(acc, scaled[k]) 98 if nx_f32_lt(u, acc) != 0 { return k } 99 k = k + 1 100 } 101 return vocab_size - 1 102} 103 104// Top-k sampler: keeps only the top-k logits, treats the rest as -inf 105// (zero probability after softmax), then temperature-samples. 106// 107// k must be in [1, vocab_size]; k=1 is greedy. 108 109func nx_f32_sampler_sample_top_k(logits: *i64, vocab_size: nx_int, 110 k: nx_int, 111 inv_temp_f32: i64, 112 prng_state: *i64) -> nx_int { 113 if vocab_size <= 0 { return 0 } 114 if k <= 0 { return nx_f32_sampler_argmax(logits, vocab_size) } 115 if k >= vocab_size { 116 return nx_f32_sampler_sample_temp(logits, vocab_size, inv_temp_f32, prng_state) 117 } 118 119 // Find the k-th highest logit value via partial selection. 120 let copy: *i64 = sys_mmap(vocab_size * 8) as *i64 121 var c: nx_int = 0 122 while c < vocab_size { copy[c] = logits[c]; c = c + 1 } 123 124 // Partial selection sort: bubble the top k values to the front. 125 var sel: nx_int = 0 126 while sel < k { 127 var best_idx: nx_int = sel 128 var best_val: i64 = copy[sel] 129 var s: nx_int = sel + 1 130 while s < vocab_size { 131 if nx_f32_lt(best_val, copy[s]) != 0 { 132 best_val = copy[s] 133 best_idx = s 134 } 135 s = s + 1 136 } 137 let tmp: i64 = copy[sel] 138 copy[sel] = copy[best_idx] 139 copy[best_idx] = tmp 140 sel = sel + 1 141 } 142 let threshold: i64 = copy[k - 1] // k-th highest value 143 144 // Build masked logits: keep >= threshold, set rest to a sentinel 145 // that triggers nx_f32_exp's underflow-to-zero threshold at 146 // -87.336. Use -88 (= 0xC2B00000) so that for any max_logit 147 // close to typical LLM ranges and inv_temp >= 1, exp produces 148 // exactly 0 for masked positions, preventing tiny-but-nonzero 149 // probabilities from leaking into the sampler's cumulative walk. 150 let masked: *i64 = sys_mmap(vocab_size * 8) as *i64 151 let neg_safe: i64 = 0xC2B00000 as i64 // -88.0 152 var m: nx_int = 0 153 while m < vocab_size { 154 if nx_f32_lt(logits[m], threshold) != 0 { 155 masked[m] = neg_safe 156 } else { 157 masked[m] = logits[m] 158 } 159 m = m + 1 160 } 161 162 return nx_f32_sampler_sample_temp(masked, vocab_size, inv_temp_f32, prng_state) 163} 164 165// Top-p (nucleus) sampler: keep the smallest set of highest-prob tokens 166// whose cumulative probability exceeds p, then temperature-sample from 167// that set. Holtzman 2019 (curious case of neural text degeneration). 168// 169// p_q14: target cumulative probability in Q14 (= prob * 16384). Common 170// production values: p=0.9 -> p_q14=14746, p=0.95 -> p_q14=15565. 171 172func _fsp_find_top_p_k(logits: *i64, vocab_size: nx_int, 173 p_q14: i64, inv_temp_f32: i64) -> nx_int { 174 // 1. Compute scaled exp values for softmax stability + temperature. 175 var max_l: i64 = logits[0] 176 var i: nx_int = 1 177 while i < vocab_size { 178 if nx_f32_lt(max_l, logits[i]) != 0 { max_l = logits[i] } 179 i = i + 1 180 } 181 let scaled: *i64 = sys_mmap(vocab_size * 8) as *i64 182 var sum: i64 = 0 183 var j: nx_int = 0 184 while j < vocab_size { 185 let shifted: i64 = nx_f32_sub(logits[j], max_l) 186 let with_t: i64 = nx_f32_mul(shifted, inv_temp_f32) 187 let e: i64 = nx_f32_exp(with_t) 188 scaled[j] = e 189 sum = nx_f32_add(sum, e) 190 j = j + 1 191 } 192 // 2. Target = p * sum. 193 let p_f32: i64 = nx_q14_to_f32(p_q14) 194 let target: i64 = nx_f32_mul(p_f32, sum) 195 196 // 3. Iterative selection of largest probs until cumsum >= target. 197 let used: *u8 = sys_mmap(vocab_size) 198 var u: nx_int = 0 199 while u < vocab_size { used[u] = 0 as u8; u = u + 1 } 200 201 var cumsum: i64 = 0 202 var k: nx_int = 0 203 var done: nx_int = 0 204 while done == 0 { 205 if k >= vocab_size { done = 1 } 206 if done == 0 { 207 if nx_f32_lt(cumsum, target) == 0 { done = 1 } 208 } 209 if done == 0 { 210 var best: nx_int = -1 211 var best_val: i64 = 0 212 var s: nx_int = 0 213 while s < vocab_size { 214 if used[s] == (0 as u8) { 215 if best < 0 { 216 best = s 217 best_val = scaled[s] 218 } else { 219 if nx_f32_lt(best_val, scaled[s]) != 0 { 220 best = s 221 best_val = scaled[s] 222 } 223 } 224 } 225 s = s + 1 226 } 227 if best < 0 { 228 done = 1 229 } else { 230 used[best] = 1 as u8 231 cumsum = nx_f32_add(cumsum, best_val) 232 k = k + 1 233 } 234 } 235 } 236 return k 237} 238 239// Repetition penalty (Keskar 2019, CTRL). Mutates logits in place: 240// For each token id `t` in recent_tokens that is in [0, vocab_size): 241// if logits[t] > 0: logits[t] = logits[t] / penalty 242// else: logits[t] = logits[t] * penalty 243// Increasing magnitude in the wrong direction makes the token less 244// likely while preserving its sign, unlike a strict mask that could 245// banish a legitimate next-word continuation. 246// 247// penalty_f32 in [1.0, ~2.0] is the production range; 1.0 disables. 248// Caller passes raw f32 bits. 249 250func nx_f32_sampler_apply_repetition_penalty( 251 logits: *i64, vocab_size: nx_int, 252 recent_tokens: *i64, n_recent: nx_int, 253 penalty_f32: i64) -> nx_int { 254 if vocab_size <= 0 { return 0 } 255 if n_recent <= 0 { return 0 } 256 let zero: i64 = 0 257 var i: nx_int = 0 258 while i < n_recent { 259 let tok: nx_int = recent_tokens[i] as nx_int 260 if tok >= 0 { 261 if tok < vocab_size { 262 let l: i64 = logits[tok] 263 if nx_f32_lt(zero, l) != 0 { 264 logits[tok] = nx_f32_div(l, penalty_f32) 265 } else { 266 logits[tok] = nx_f32_mul(l, penalty_f32) 267 } 268 } 269 } 270 i = i + 1 271 } 272 return 0 273} 274 275// Min-p sampler (Nguyen 2023). Keep tokens whose probability is 276// >= min_p * max_prob; truncate the rest. Threshold scales with 277// the dominant prediction's confidence -- complements top-p well. 278// 279// min_p_q14 in Q14: e.g. 0.05 -> 819, 0.1 -> 1638, 0.2 -> 3277. 280 281func _fsp_find_min_p_k(logits: *i64, vocab_size: nx_int, 282 min_p_q14: i64, inv_temp_f32: i64) -> nx_int { 283 var max_l: i64 = logits[0] 284 var i: nx_int = 1 285 while i < vocab_size { 286 if nx_f32_lt(max_l, logits[i]) != 0 { max_l = logits[i] } 287 i = i + 1 288 } 289 let scaled: *i64 = sys_mmap(vocab_size * 8) as *i64 290 var max_scaled: i64 = 0 291 var sum: i64 = 0 292 var j: nx_int = 0 293 while j < vocab_size { 294 let shifted: i64 = nx_f32_sub(logits[j], max_l) 295 let with_t: i64 = nx_f32_mul(shifted, inv_temp_f32) 296 let e: i64 = nx_f32_exp(with_t) 297 scaled[j] = e 298 if j == 0 { max_scaled = e } else { 299 if nx_f32_lt(max_scaled, e) != 0 { max_scaled = e } 300 } 301 sum = nx_f32_add(sum, e) 302 j = j + 1 303 } 304 let min_p_f32: i64 = nx_q14_to_f32(min_p_q14) 305 let threshold: i64 = nx_f32_mul(min_p_f32, max_scaled) 306 307 var k: nx_int = 0 308 var s: nx_int = 0 309 while s < vocab_size { 310 if nx_f32_lt(scaled[s], threshold) == 0 { k = k + 1 } 311 s = s + 1 312 } 313 return k 314} 315 316func nx_f32_sampler_sample_min_p(logits: *i64, vocab_size: nx_int, 317 min_p_q14: i64, 318 inv_temp_f32: i64, 319 prng_state: *i64) -> nx_int { 320 if vocab_size <= 0 { return 0 } 321 if min_p_q14 <= 0 { 322 return nx_f32_sampler_sample_temp(logits, vocab_size, inv_temp_f32, prng_state) 323 } 324 let k: nx_int = _fsp_find_min_p_k(logits, vocab_size, min_p_q14, inv_temp_f32) 325 if k <= 0 { return nx_f32_sampler_argmax(logits, vocab_size) } 326 return nx_f32_sampler_sample_top_k(logits, vocab_size, k, 327 inv_temp_f32, prng_state) 328} 329 330func nx_f32_sampler_sample_top_p(logits: *i64, vocab_size: nx_int, 331 p_q14: i64, 332 inv_temp_f32: i64, 333 prng_state: *i64) -> nx_int { 334 if vocab_size <= 0 { return 0 } 335 if p_q14 <= 0 { return nx_f32_sampler_argmax(logits, vocab_size) } 336 // p_q14 >= 16384 (= prob >= 1.0) -> no truncation, plain temp sample. 337 if p_q14 >= NX_MAGIC_16384 { 338 return nx_f32_sampler_sample_temp(logits, vocab_size, inv_temp_f32, prng_state) 339 } 340 let k: nx_int = _fsp_find_top_p_k(logits, vocab_size, p_q14, inv_temp_f32) 341 if k <= 0 { return nx_f32_sampler_argmax(logits, vocab_size) } 342 return nx_f32_sampler_sample_top_k(logits, vocab_size, k, 343 inv_temp_f32, prng_state) 344}