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}