1#[derive(Debug, Clone)]
35pub struct BeamSearchConfig {
36 pub beam_width: usize,
38 pub max_tokens: usize,
40 pub length_penalty: f32,
45 pub no_repeat_ngram_size: usize,
48 pub early_stopping: bool,
50 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#[derive(Debug, Clone)]
71pub struct Beam {
72 pub tokens: Vec<u32>,
74 pub log_prob: f64,
76 pub is_done: bool,
78}
79
80impl Beam {
81 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 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 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 pub fn len(&self) -> usize {
113 self.tokens.len()
114 }
115
116 pub fn is_empty(&self) -> bool {
118 self.tokens.is_empty()
119 }
120}
121
122#[derive(Debug)]
126pub struct BeamSearchResult {
127 pub sequences: Vec<Vec<u32>>,
129 pub scores: Vec<f64>,
131 pub num_steps: usize,
133}
134
135impl BeamSearchResult {
136 pub fn best(&self) -> &[u32] {
138 self.sequences.first().map(|s| s.as_slice()).unwrap_or(&[])
139 }
140
141 pub fn best_score(&self) -> f64 {
143 self.scores.first().copied().unwrap_or(f64::NEG_INFINITY)
144 }
145}
146
147pub struct BeamSearchEngine {
154 pub config: BeamSearchConfig,
156}
157
158impl BeamSearchEngine {
159 pub fn new(config: BeamSearchConfig) -> Self {
161 Self { config }
162 }
163
164 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 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 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 let mut candidates: Vec<Beam> = Vec::new();
198
199 for beam in &live {
200 let mut logits = get_logits(&beam.tokens, step);
201
202 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 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 let done_indices: Vec<usize> = beams
231 .iter()
232 .enumerate()
233 .filter(|(_, b)| b.is_done)
234 .map(|(i, _)| i)
235 .collect();
236 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 beams = Self::prune_beams(candidates, bw, cfg.length_penalty);
247
248 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 for b in beams {
269 completed.push(b);
270 }
271
272 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 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 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 let prefix_len = ngram_size - 1;
307 let suffix = &tokens[tokens.len() - prefix_len..];
308
309 for start in 0..tokens.len().saturating_sub(prefix_len) {
311 let window = &tokens[start..start + prefix_len];
312 if window == suffix {
313 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 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 let max_logit = logits
332 .iter()
333 .copied()
334 .filter(|v| v.is_finite())
335 .fold(f32::NEG_INFINITY, f32::max);
336
337 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 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 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#[cfg(test)]
378mod tests {
379 use super::*;
380
381 #[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 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 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 #[test]
438 fn test_top_k_log_probs_returns_k_best() {
439 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 assert_eq!(result[0].0, 3);
445 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 #[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 assert_eq!(pruned[0].tokens, vec![2]);
492 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 #[test]
510 fn test_apply_no_repeat_ngram_blocks_repeated() {
511 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 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 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 #[test]
553 fn test_beam_search_greedy_equivalent_width1() {
554 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, };
563 let engine = BeamSearchEngine::new(config);
564
565 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 assert!(best.iter().skip(1).all(|&t| t == 7));
576 }
577
578 #[test]
579 fn test_beam_search_with_eos() {
580 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 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 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}