Skip to main content

pictor_runtime/
beam_search.rs

1//! Beam search decoding for Pictor.
2//!
3//! Beam search maintains `beam_width` candidate sequences simultaneously,
4//! expanding each at every step and keeping the top-`beam_width` by
5//! cumulative log-probability (with optional length penalty).
6//!
7//! # Example
8//!
9//! ```rust
10//! use pictor_runtime::beam_search::{BeamSearchConfig, BeamSearchEngine};
11//!
12//! let config = BeamSearchConfig {
13//!     beam_width: 2,
14//!     max_tokens: 10,
15//!     eos_token_id: 2,
16//!     ..Default::default()
17//! };
18//! let engine = BeamSearchEngine::new(config);
19//!
20//! // Mock logits: always prefer token 5
21//! let result = engine.search(vec![1, 2], 10, |_tokens, _step| {
22//!     let mut logits = vec![0.0f32; 10];
23//!     logits[5] = 10.0;
24//!     logits[2] = -10.0; // EOS gets low score
25//!     logits
26//! });
27//!
28//! assert!(!result.best().is_empty());
29//! ```
30
31// ─── Config ────────────────────────────────────────────────────────────────
32
33/// Configuration for beam search decoding.
34#[derive(Debug, Clone)]
35pub struct BeamSearchConfig {
36    /// Number of parallel beams to maintain (typical: 4–8).
37    pub beam_width: usize,
38    /// Maximum tokens to generate per beam.
39    pub max_tokens: usize,
40    /// Length penalty exponent α: `score = log_prob / len^α`.
41    ///
42    /// Values in [0.6, 1.0] are typical. α = 1.0 is neutral; α < 1.0
43    /// rewards longer sequences; α > 1.0 penalises them.
44    pub length_penalty: f32,
45    /// Block any token that would create a repeated n-gram of this size.
46    /// Set to 0 to disable (default).
47    pub no_repeat_ngram_size: usize,
48    /// Stop as soon as the best beam generates an EOS token.
49    pub early_stopping: bool,
50    /// Token ID that marks end of sequence.
51    pub eos_token_id: u32,
52}
53
54impl Default for BeamSearchConfig {
55    fn default() -> Self {
56        Self {
57            beam_width: 4,
58            max_tokens: 256,
59            length_penalty: 0.6,
60            no_repeat_ngram_size: 0,
61            early_stopping: true,
62            eos_token_id: 2,
63        }
64    }
65}
66
67// ─── Beam ──────────────────────────────────────────────────────────────────
68
69/// One candidate sequence in the beam search.
70#[derive(Debug, Clone)]
71pub struct Beam {
72    /// All token IDs in this candidate (prompt + generated so far).
73    pub tokens: Vec<u32>,
74    /// Cumulative log-probability of this sequence.
75    pub log_prob: f64,
76    /// Whether this beam has hit an EOS token and is finished.
77    pub is_done: bool,
78}
79
80impl Beam {
81    /// Create a new beam seeded with the given initial tokens.
82    pub fn new(initial_tokens: Vec<u32>) -> Self {
83        Self {
84            tokens: initial_tokens,
85            log_prob: 0.0,
86            is_done: false,
87        }
88    }
89
90    /// Length-normalised score used for beam ranking.
91    ///
92    /// `score = log_prob / (len ^ length_penalty)`
93    ///
94    /// Avoids division-by-zero by treating a zero-length sequence as length 1.
95    pub fn score(&self, length_penalty: f32) -> f64 {
96        let len = self.tokens.len().max(1) as f64;
97        self.log_prob / len.powf(length_penalty as f64)
98    }
99
100    /// Extend the beam with one more token, returning a new beam.
101    pub fn extend(&self, token: u32, log_prob: f64) -> Self {
102        let mut tokens = self.tokens.clone();
103        tokens.push(token);
104        Self {
105            tokens,
106            log_prob: self.log_prob + log_prob,
107            is_done: false,
108        }
109    }
110
111    /// Total number of tokens in this beam.
112    pub fn len(&self) -> usize {
113        self.tokens.len()
114    }
115
116    /// `true` when the beam contains no tokens.
117    pub fn is_empty(&self) -> bool {
118        self.tokens.is_empty()
119    }
120}
121
122// ─── Result ────────────────────────────────────────────────────────────────
123
124/// Output of a beam search run.
125#[derive(Debug)]
126pub struct BeamSearchResult {
127    /// All completed sequences, ordered best-first.
128    pub sequences: Vec<Vec<u32>>,
129    /// Length-normalised score for each sequence.
130    pub scores: Vec<f64>,
131    /// Number of generation steps taken.
132    pub num_steps: usize,
133}
134
135impl BeamSearchResult {
136    /// The highest-scoring token sequence.
137    pub fn best(&self) -> &[u32] {
138        self.sequences.first().map(|s| s.as_slice()).unwrap_or(&[])
139    }
140
141    /// Score of the highest-scoring sequence.
142    pub fn best_score(&self) -> f64 {
143        self.scores.first().copied().unwrap_or(f64::NEG_INFINITY)
144    }
145}
146
147// ─── Engine ────────────────────────────────────────────────────────────────
148
149/// Beam search engine.
150///
151/// Decoupled from the model via a `get_logits` closure so it can be used
152/// with any inference backend.
153pub struct BeamSearchEngine {
154    /// Search configuration.
155    pub config: BeamSearchConfig,
156}
157
158impl BeamSearchEngine {
159    /// Create a new engine with the given configuration.
160    pub fn new(config: BeamSearchConfig) -> Self {
161        Self { config }
162    }
163
164    /// Run beam search.
165    ///
166    /// `get_logits(beam_tokens, step)` is called for every live beam at every
167    /// step and must return a logit vector of length `vocab_size`.
168    pub fn search<F>(
169        &self,
170        initial_tokens: Vec<u32>,
171        _vocab_size: usize,
172        mut get_logits: F,
173    ) -> BeamSearchResult
174    where
175        F: FnMut(&[u32], usize) -> Vec<f32>,
176    {
177        let cfg = &self.config;
178        let bw = cfg.beam_width.max(1);
179
180        // Initialise with a single beam
181        let mut beams: Vec<Beam> = vec![Beam::new(initial_tokens)];
182        let mut completed: Vec<Beam> = Vec::new();
183        let mut steps = 0;
184
185        for step in 0..cfg.max_tokens {
186            steps = step + 1;
187
188            // Collect live (non-done) beams
189            let live: Vec<Beam> = beams.iter().filter(|b| !b.is_done).cloned().collect();
190
191            if live.is_empty() {
192                steps = step;
193                break;
194            }
195
196            // Expand every live beam
197            let mut candidates: Vec<Beam> = Vec::new();
198
199            for beam in &live {
200                let mut logits = get_logits(&beam.tokens, step);
201
202                // Apply no-repeat-ngram masking if configured
203                if cfg.no_repeat_ngram_size > 0 {
204                    Self::apply_no_repeat_ngram(
205                        &mut logits,
206                        &beam.tokens,
207                        cfg.no_repeat_ngram_size,
208                    );
209                }
210
211                // Get top-k (token, log_prob) candidates from this beam
212                let top = Self::top_k_log_probs(&logits, bw);
213
214                for (token, lp) in top {
215                    let mut new_beam = beam.extend(token, lp);
216
217                    if token == cfg.eos_token_id {
218                        new_beam.is_done = true;
219                        if cfg.early_stopping {
220                            completed.push(new_beam);
221                            continue;
222                        }
223                    }
224                    candidates.push(new_beam);
225                }
226            }
227
228            // Keep any already-done beams from the previous round
229            // Use drain to avoid moving `beams` so we can still use it after break
230            let done_indices: Vec<usize> = beams
231                .iter()
232                .enumerate()
233                .filter(|(_, b)| b.is_done)
234                .map(|(i, _)| i)
235                .collect();
236            // Remove done beams in reverse index order to preserve indices
237            for &idx in done_indices.iter().rev() {
238                completed.push(beams.remove(idx));
239            }
240
241            if candidates.is_empty() {
242                break;
243            }
244
245            // Prune to beam_width
246            beams = Self::prune_beams(candidates, bw, cfg.length_penalty);
247
248            // Early-stop when best completed beam outscores every live beam
249            if cfg.early_stopping && !completed.is_empty() {
250                let best_completed_score = completed
251                    .iter()
252                    .map(|b| b.score(cfg.length_penalty))
253                    .fold(f64::NEG_INFINITY, f64::max);
254
255                let best_live_score = beams
256                    .iter()
257                    .map(|b| b.score(cfg.length_penalty))
258                    .fold(f64::NEG_INFINITY, f64::max);
259
260                if best_completed_score >= best_live_score {
261                    steps = step + 1;
262                    break;
263                }
264            }
265        }
266
267        // Gather all remaining live beams as completed
268        for b in beams {
269            completed.push(b);
270        }
271
272        // Sort by score descending
273        completed.sort_by(|a, b| {
274            b.score(cfg.length_penalty)
275                .partial_cmp(&a.score(cfg.length_penalty))
276                .unwrap_or(std::cmp::Ordering::Equal)
277        });
278
279        // Keep at most beam_width results
280        completed.truncate(bw);
281
282        let scores: Vec<f64> = completed
283            .iter()
284            .map(|b| b.score(cfg.length_penalty))
285            .collect();
286        let sequences: Vec<Vec<u32>> = completed.into_iter().map(|b| b.tokens).collect();
287
288        BeamSearchResult {
289            sequences,
290            scores,
291            num_steps: steps,
292        }
293    }
294
295    /// Zero out (set to −∞) any token that would create a repeated n-gram.
296    ///
297    /// For each position in `tokens` where the last `ngram_size - 1` tokens
298    /// match the trailing `ngram_size - 1` tokens of the current sequence,
299    /// the following token is forbidden.
300    pub fn apply_no_repeat_ngram(logits: &mut [f32], tokens: &[u32], ngram_size: usize) {
301        if ngram_size == 0 || tokens.len() < ngram_size {
302            return;
303        }
304
305        // The suffix we want to avoid repeating is the last (ngram_size - 1) tokens
306        let prefix_len = ngram_size - 1;
307        let suffix = &tokens[tokens.len() - prefix_len..];
308
309        // Scan all valid n-gram starting positions in the existing token sequence
310        for start in 0..tokens.len().saturating_sub(prefix_len) {
311            let window = &tokens[start..start + prefix_len];
312            if window == suffix {
313                // The token that would complete the n-gram is at `start + prefix_len`
314                let banned_token = tokens[start + prefix_len] as usize;
315                if banned_token < logits.len() {
316                    logits[banned_token] = f32::NEG_INFINITY;
317                }
318            }
319        }
320    }
321
322    /// Return the top-`k` `(token_id, log_prob)` pairs from a logit vector.
323    ///
324    /// Logits are converted to log-probabilities via log-softmax.
325    pub fn top_k_log_probs(logits: &[f32], k: usize) -> Vec<(u32, f64)> {
326        if logits.is_empty() {
327            return Vec::new();
328        }
329
330        // Numerical stability: subtract max before exp
331        let max_logit = logits
332            .iter()
333            .copied()
334            .filter(|v| v.is_finite())
335            .fold(f32::NEG_INFINITY, f32::max);
336
337        // Compute log-softmax: log_prob_i = logit_i - max - log(sum(exp(logit_j - max)))
338        let shifted: Vec<f32> = logits
339            .iter()
340            .map(|&v| {
341                if v.is_finite() {
342                    v - max_logit
343                } else {
344                    f32::NEG_INFINITY
345                }
346            })
347            .collect();
348
349        let log_sum_exp = shifted.iter().copied().map(|v| v.exp()).sum::<f32>().ln();
350
351        let mut indexed: Vec<(u32, f64)> = shifted
352            .iter()
353            .enumerate()
354            .map(|(i, &v)| (i as u32, (v - log_sum_exp) as f64))
355            .collect();
356
357        // Sort by log-prob descending
358        indexed.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
359        indexed.truncate(k);
360        indexed
361    }
362
363    /// Keep the top `beam_width` beams by length-normalised score.
364    pub fn prune_beams(mut beams: Vec<Beam>, beam_width: usize, length_penalty: f32) -> Vec<Beam> {
365        beams.sort_by(|a, b| {
366            b.score(length_penalty)
367                .partial_cmp(&a.score(length_penalty))
368                .unwrap_or(std::cmp::Ordering::Equal)
369        });
370        beams.truncate(beam_width);
371        beams
372    }
373}
374
375// ─── Tests ─────────────────────────────────────────────────────────────────
376
377#[cfg(test)]
378mod tests {
379    use super::*;
380
381    // ── Beam unit tests ────────────────────────────────────────────────────
382
383    #[test]
384    fn test_beam_new_initial() {
385        let tokens = vec![1u32, 2, 3];
386        let beam = Beam::new(tokens.clone());
387        assert_eq!(beam.tokens, tokens);
388        assert!((beam.log_prob - 0.0).abs() < f64::EPSILON);
389        assert!(!beam.is_done);
390        assert_eq!(beam.len(), 3);
391        assert!(!beam.is_empty());
392    }
393
394    #[test]
395    fn test_beam_score_length_penalty() {
396        let beam = Beam {
397            tokens: vec![1, 2, 3, 4],
398            log_prob: -4.0,
399            is_done: false,
400        };
401        // score = -4.0 / 4^0.6
402        let expected = -4.0_f64 / (4.0_f64.powf(0.6));
403        let score = beam.score(0.6);
404        assert!(
405            (score - expected).abs() < 1e-6,
406            "score={score}, expected={expected}"
407        );
408    }
409
410    #[test]
411    fn test_beam_score_zero_length() {
412        // An empty beam should not panic — treated as length 1
413        let beam = Beam {
414            tokens: vec![],
415            log_prob: -1.0,
416            is_done: false,
417        };
418        let score = beam.score(0.6);
419        assert!((score - -1.0_f64).abs() < 1e-10);
420    }
421
422    #[test]
423    fn test_beam_extend() {
424        let beam = Beam {
425            tokens: vec![1, 2],
426            log_prob: -1.5,
427            is_done: false,
428        };
429        let extended = beam.extend(3, -0.5);
430        assert_eq!(extended.tokens, vec![1, 2, 3]);
431        assert!((extended.log_prob - -2.0).abs() < 1e-10);
432        assert!(!extended.is_done);
433    }
434
435    // ── top_k_log_probs tests ──────────────────────────────────────────────
436
437    #[test]
438    fn test_top_k_log_probs_returns_k_best() {
439        // logits with clear winner at index 3
440        let logits = vec![0.0f32, 1.0, 2.0, 10.0, 0.5];
441        let result = BeamSearchEngine::top_k_log_probs(&logits, 2);
442        assert_eq!(result.len(), 2);
443        // Best token should be index 3
444        assert_eq!(result[0].0, 3);
445        // Log-probs should be in descending order
446        assert!(result[0].1 >= result[1].1);
447    }
448
449    #[test]
450    fn test_top_k_log_probs_k_larger_than_vocab() {
451        let logits = vec![1.0f32, 2.0, 3.0];
452        let result = BeamSearchEngine::top_k_log_probs(&logits, 10);
453        assert_eq!(result.len(), 3);
454    }
455
456    #[test]
457    fn test_top_k_log_probs_empty() {
458        let result = BeamSearchEngine::top_k_log_probs(&[], 4);
459        assert!(result.is_empty());
460    }
461
462    // ── prune_beams tests ──────────────────────────────────────────────────
463
464    #[test]
465    fn test_prune_beams_keeps_best() {
466        let beams = vec![
467            Beam {
468                tokens: vec![1],
469                log_prob: -10.0,
470                is_done: false,
471            },
472            Beam {
473                tokens: vec![2],
474                log_prob: -1.0,
475                is_done: false,
476            },
477            Beam {
478                tokens: vec![3],
479                log_prob: -5.0,
480                is_done: false,
481            },
482            Beam {
483                tokens: vec![4],
484                log_prob: -2.0,
485                is_done: false,
486            },
487        ];
488        let pruned = BeamSearchEngine::prune_beams(beams, 2, 1.0);
489        assert_eq!(pruned.len(), 2);
490        // Best beam has log_prob = -1.0 → tokens = [2]
491        assert_eq!(pruned[0].tokens, vec![2]);
492        // Second-best has log_prob = -2.0 → tokens = [4]
493        assert_eq!(pruned[1].tokens, vec![4]);
494    }
495
496    #[test]
497    fn test_prune_beams_fewer_than_width() {
498        let beams = vec![Beam {
499            tokens: vec![1],
500            log_prob: -3.0,
501            is_done: false,
502        }];
503        let pruned = BeamSearchEngine::prune_beams(beams, 4, 0.6);
504        assert_eq!(pruned.len(), 1);
505    }
506
507    // ── apply_no_repeat_ngram tests ───────────────────────────────────────
508
509    #[test]
510    fn test_apply_no_repeat_ngram_blocks_repeated() {
511        // tokens = [1, 2, 3]; ngram_size = 2 → last prefix is [3]
512        // If [3] appeared before at position 1 (tokens[1]=2≠3), skip.
513        // If [3] appeared before at position 2 (tokens[2]=3), following token is tokens[3] — but
514        // tokens only has length 3, so that would be out of bounds. Let's use a longer sequence.
515        //
516        // tokens = [1, 2, 1, 2]; ngram_size = 2 → suffix = [2]
517        // position 1: tokens[1]=2 matches; next token = tokens[2]=1 → ban token 1
518        let tokens = vec![1u32, 2, 1, 2];
519        let mut logits = vec![0.0f32; 5];
520        BeamSearchEngine::apply_no_repeat_ngram(&mut logits, &tokens, 2);
521        assert_eq!(logits[1], f32::NEG_INFINITY, "token 1 should be banned");
522        // token 2 not yet banned (the last occurrence of [2] is at the very end,
523        // no following token exists in history)
524        assert!(logits[2].is_finite());
525    }
526
527    #[test]
528    fn test_no_repeat_ngram_no_effect_when_disabled() {
529        let tokens = vec![1u32, 2, 1, 2];
530        let original = vec![1.0f32, 2.0, 3.0, 4.0, 5.0];
531        let mut logits = original.clone();
532        BeamSearchEngine::apply_no_repeat_ngram(&mut logits, &tokens, 0);
533        assert_eq!(
534            logits, original,
535            "ngram_size=0 should leave logits unchanged"
536        );
537    }
538
539    #[test]
540    fn test_no_repeat_ngram_too_short_sequence() {
541        // Sequence shorter than ngram_size → no banning
542        let tokens = vec![1u32];
543        let mut logits = vec![1.0f32; 5];
544        BeamSearchEngine::apply_no_repeat_ngram(&mut logits, &tokens, 3);
545        for &v in &logits {
546            assert!(v.is_finite());
547        }
548    }
549
550    // ── BeamSearchEngine::search integration tests ────────────────────────
551
552    #[test]
553    fn test_beam_search_greedy_equivalent_width1() {
554        // With beam_width=1 and greedy logits, beam search is equivalent to greedy decoding.
555        let config = BeamSearchConfig {
556            beam_width: 1,
557            max_tokens: 5,
558            length_penalty: 1.0,
559            no_repeat_ngram_size: 0,
560            early_stopping: false,
561            eos_token_id: 99, // Never generated
562        };
563        let engine = BeamSearchEngine::new(config);
564
565        // Always return token 7 as the best
566        let result = engine.search(vec![0u32], 10, |_tokens, _step| {
567            let mut logits = vec![0.0f32; 10];
568            logits[7] = 100.0;
569            logits
570        });
571
572        assert_eq!(result.num_steps, 5);
573        let best = result.best();
574        // First token is initial (0), remaining should all be 7
575        assert!(best.iter().skip(1).all(|&t| t == 7));
576    }
577
578    #[test]
579    fn test_beam_search_with_eos() {
580        // Beam search should stop early when EOS is generated (early_stopping=true).
581        let eos = 3u32;
582        let config = BeamSearchConfig {
583            beam_width: 2,
584            max_tokens: 20,
585            length_penalty: 0.6,
586            no_repeat_ngram_size: 0,
587            early_stopping: true,
588            eos_token_id: eos,
589        };
590        let engine = BeamSearchEngine::new(config);
591
592        let step_counter = std::cell::Cell::new(0usize);
593        let result = engine.search(vec![1u32], 5, |_tokens, _step| {
594            step_counter.set(step_counter.get() + 1);
595            // After 2 calls produce EOS as best token
596            let mut logits = vec![0.0f32; 5];
597            if step_counter.get() >= 2 {
598                logits[eos as usize] = 100.0;
599            } else {
600                logits[1] = 5.0;
601            }
602            logits
603        });
604
605        // Should not have run all 20 steps
606        assert!(
607            result.num_steps < 20,
608            "expected early stop, got {} steps",
609            result.num_steps
610        );
611        assert!(!result.sequences.is_empty());
612    }
613
614    #[test]
615    fn test_beam_search_result_best() {
616        let result = BeamSearchResult {
617            sequences: vec![vec![1, 2, 3], vec![4, 5, 6]],
618            scores: vec![-0.5, -1.0],
619            num_steps: 3,
620        };
621        assert_eq!(result.best(), &[1, 2, 3]);
622        assert!((result.best_score() - -0.5).abs() < f64::EPSILON);
623    }
624
625    #[test]
626    fn test_beam_search_result_empty() {
627        let result = BeamSearchResult {
628            sequences: vec![],
629            scores: vec![],
630            num_steps: 0,
631        };
632        assert_eq!(result.best(), &[] as &[u32]);
633        assert_eq!(result.best_score(), f64::NEG_INFINITY);
634    }
635}