Files
rustytorch/crates/production/rtx-inference/src/beam_search.rs
T
osobhandClaude Sonnet 5 4aaa36a57a style: cargo fmt --workspace (whitespace/wrapping only, no semantic change)
Whole-workspace rustfmt pass picked up while iterating on Mamba GPU
backward work. Verified formatting-only via diff sampling; no logic
changed.

Co-Authored-By: Claude Sonnet 5 <[email protected]>
2026-08-10 07:09:36 -07:00

967 lines
34 KiB
Rust
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
//! Beam search decoding for autoregressive language models.
//!
//! Implements standard beam search (Sutskever et al. 2014) and diverse beam search
//! (Vijayakumar et al. 2018) with Wu et al. 2016 length normalization.
//!
//! # Example
//!
//! ```rust
//! use rtx_inference::beam_search::{BeamSearchConfig, BeamSearchDecoder};
//!
//! let config = BeamSearchConfig {
//! beam_width: 4,
//! max_length: 32,
//! eos_token_id: 2,
//! length_penalty_alpha: 0.6,
//! min_length: 1,
//! no_repeat_ngram_size: 0,
//! };
//! let decoder = BeamSearchDecoder::new(config);
//!
//! let vocab_size = 32_000;
//! let results = decoder.decode(&[1], vocab_size, 2, |tokens| {
//! let mut lp = vec![-10.0_f32; vocab_size];
//! lp[5] = -0.1;
//! lp[2] = -0.5; // EOS
//! lp
//! });
//! assert!(!results.is_empty());
//! ```
use std::collections::HashSet;
// ---------------------------------------------------------------------------
// Beam
// ---------------------------------------------------------------------------
/// A single beam hypothesis tracking generated tokens and cumulative log-probability.
#[derive(Debug, Clone)]
pub struct Beam {
/// Token IDs generated so far (including any initial prompt tokens).
pub tokens: Vec<u32>,
/// Cumulative sum of log-probabilities for all generated tokens.
pub log_prob: f32,
/// Whether this beam has emitted the EOS token and is considered finished.
pub finished: bool,
}
impl Beam {
/// Create a new, empty beam with the given initial log-probability.
///
/// The token list starts empty; use `tokens` field to pre-populate prompt tokens.
#[must_use]
pub fn new(log_prob: f32) -> Self {
Self {
tokens: Vec::new(),
log_prob,
finished: false,
}
}
/// Length-normalized score using the Wu et al. 2016 formula.
///
/// `score = log_prob / length_penalty(len)`
/// `length_penalty(l) = ((5 + l) / 6) ^ alpha`
///
/// With `alpha = 0` the raw log-probability is returned unchanged.
/// With `alpha = 1` full length normalization is applied.
///
/// An empty token sequence returns `log_prob` directly to avoid division issues.
#[must_use]
pub fn normalized_score(&self, alpha: f32) -> f32 {
let len = self.tokens.len();
if len == 0 || alpha == 0.0 {
return self.log_prob;
}
let penalty = ((5.0 + len as f32) / 6.0_f32).powf(alpha);
self.log_prob / penalty
}
}
// ---------------------------------------------------------------------------
// BeamSearchConfig
// ---------------------------------------------------------------------------
/// Configuration parameters for the beam search decoder.
#[derive(Debug, Clone)]
pub struct BeamSearchConfig {
/// Number of beams to maintain at each step.
pub beam_width: usize,
/// Maximum number of tokens to generate per beam.
pub max_length: usize,
/// Token ID that marks end-of-sequence.
pub eos_token_id: u32,
/// Length penalty exponent (Wu et al. 2016). 0 = no penalty, 1 = full normalization.
pub length_penalty_alpha: f32,
/// Suppress EOS emission before this many generated tokens.
pub min_length: usize,
/// If > 0, block any token that would complete an already-seen n-gram of this size.
pub no_repeat_ngram_size: usize,
}
impl Default for BeamSearchConfig {
fn default() -> Self {
Self {
beam_width: 4,
max_length: 128,
eos_token_id: 2,
length_penalty_alpha: 0.6,
min_length: 1,
no_repeat_ngram_size: 0,
}
}
}
// ---------------------------------------------------------------------------
// BeamSearchDecoder
// ---------------------------------------------------------------------------
/// Standard beam search decoder.
pub struct BeamSearchDecoder {
config: BeamSearchConfig,
}
impl BeamSearchDecoder {
/// Create a decoder with a custom configuration.
#[must_use]
pub fn new(config: BeamSearchConfig) -> Self {
Self { config }
}
/// Create a decoder with default configuration.
#[must_use]
pub fn default_decoder() -> Self {
Self::new(BeamSearchConfig::default())
}
/// Run beam search given a scoring function.
///
/// # Arguments
///
/// * `initial_tokens` — Prompt tokens prepended to every beam (not generated; used as
/// context only). Their length offsets `min_length` enforcement.
/// * `vocab_size` — Total vocabulary size; used as the upper bound for candidate tokens.
/// * `num_return` — Number of best beams to return.
/// * `score_fn` — Callable that receives the full current token sequence (prompt +
/// generated) and returns a `Vec<f32>` of length `vocab_size` containing log-probs.
///
/// # Returns
///
/// Up to `num_return` beams sorted by descending normalized score. If no beam has
/// emitted EOS the best still-active beams are returned instead.
pub fn decode<F>(
&self,
initial_tokens: &[u32],
vocab_size: usize,
num_return: usize,
score_fn: F,
) -> Vec<Beam>
where
F: Fn(&[u32]) -> Vec<f32>,
{
let bw = self.config.beam_width;
let alpha = self.config.length_penalty_alpha;
// Initialise with a single beam carrying all prompt tokens.
let mut active: Vec<Beam> = vec![Beam {
tokens: initial_tokens.to_vec(),
log_prob: 0.0,
finished: false,
}];
let mut finished: Vec<Beam> = Vec::new();
for _step in 0..self.config.max_length {
if active.is_empty() {
break;
}
let mut candidates: Vec<Beam> = Vec::new();
for beam in &active {
let log_probs = score_fn(&beam.tokens);
// Determine which token ids are blocked by the n-gram filter.
let blocked = if self.config.no_repeat_ngram_size > 0 {
Self::blocked_tokens_for_ngram(&beam.tokens, self.config.no_repeat_ngram_size)
} else {
HashSet::new()
};
// EOS is suppressed until the beam has generated at least `min_length` tokens
// beyond the prompt.
let generated_so_far = beam.tokens.len().saturating_sub(initial_tokens.len());
let suppress_eos = generated_so_far < self.config.min_length;
// Request extra tokens so that after filtering (EOS suppression, n-gram
// blocking) we still have at least `bw` live candidates. The +2 absorbs
// one EOS suppression slot and one n-gram block slot simultaneously.
let fetch_k = bw * 2 + 2;
let top = Self::top_k_indices(&log_probs, fetch_k);
let mut added = 0usize;
for (token_id, lp) in top {
if added >= bw {
break;
}
let token_id = token_id as u32;
if blocked.contains(&token_id) {
continue;
}
if suppress_eos && token_id == self.config.eos_token_id {
continue;
}
let mut new_tokens = beam.tokens.clone();
new_tokens.push(token_id);
let new_lp = beam.log_prob + lp;
let is_finished = token_id == self.config.eos_token_id;
candidates.push(Beam {
tokens: new_tokens,
log_prob: new_lp,
finished: is_finished,
});
added += 1;
}
}
// Sort all candidates by normalized score descending.
candidates.sort_by(|a, b| {
b.normalized_score(alpha)
.partial_cmp(&a.normalized_score(alpha))
.unwrap_or(std::cmp::Ordering::Equal)
});
// Partition into finished and active, keeping at most `bw` active beams.
active.clear();
for beam in candidates.into_iter().take(bw * 2) {
if beam.finished {
finished.push(beam);
} else if active.len() < bw {
active.push(beam);
}
}
// Early stopping: if we already have enough finished beams and no active beam
// can possibly outscore the worst finished beam, we are done.
if finished.len() >= num_return && !active.is_empty() {
let worst_finished = finished
.iter()
.map(|b| b.normalized_score(alpha))
.fold(f32::INFINITY, f32::min);
let best_active = active
.iter()
.map(|b| b.normalized_score(alpha))
.fold(f32::NEG_INFINITY, f32::max);
if best_active < worst_finished {
break;
}
}
}
// Any remaining active beams are folded into finished so the caller always gets
// results even when no EOS was produced within `max_length`.
finished.extend(active);
finished.sort_by(|a, b| {
b.normalized_score(alpha)
.partial_cmp(&a.normalized_score(alpha))
.unwrap_or(std::cmp::Ordering::Equal)
});
finished.into_iter().take(num_return).collect()
}
/// Return the top-`k` `(index, log_prob)` pairs from `log_probs`, sorted descending.
///
/// If `k` is greater than or equal to the length of `log_probs`, all entries are
/// returned (still sorted descending).
#[must_use]
pub fn top_k_indices(log_probs: &[f32], k: usize) -> Vec<(usize, f32)> {
let k = k.min(log_probs.len());
// Partial sort: O(n log k) via a min-heap maintained at size k.
// For typical vocab sizes (32k256k) and small k (beam width ≤ 16) this is
// substantially faster than a full sort.
use std::cmp::Reverse;
use std::collections::BinaryHeap;
// BinaryHeap is a max-heap; we wrap in Reverse to get a min-heap so we can
// efficiently evict the smallest element as we scan.
// The heap stores (ordered_float_bits, index) where bits are from f32::to_bits
// reinterpreted for ordering. We use a small wrapper instead.
#[derive(PartialEq)]
struct OrdF32(f32);
impl Eq for OrdF32 {}
impl PartialOrd for OrdF32 {
fn partial_cmp(&self, other: &Self) -> Option<std::cmp::Ordering> {
Some(self.cmp(other))
}
}
impl Ord for OrdF32 {
fn cmp(&self, other: &Self) -> std::cmp::Ordering {
self.0
.partial_cmp(&other.0)
.unwrap_or(std::cmp::Ordering::Equal)
}
}
// min-heap of size k: smallest value is at the top and will be evicted first.
let mut heap: BinaryHeap<Reverse<(OrdF32, usize)>> = BinaryHeap::with_capacity(k + 1);
for (idx, &lp) in log_probs.iter().enumerate() {
heap.push(Reverse((OrdF32(lp), idx)));
if heap.len() > k {
heap.pop();
}
}
// Drain heap into a Vec and sort descending.
let mut result: Vec<(usize, f32)> = heap
.into_iter()
.map(|Reverse((OrdF32(lp), idx))| (idx, lp))
.collect();
result.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
result
}
/// Compute the set of token IDs that would complete an already-seen n-gram.
///
/// For a token sequence `tokens` and n-gram size `n`, we look at the last `n-1`
/// tokens as a suffix key and find all tokens `t` such that `suffix ++ [t]` appears
/// earlier in `tokens`. Those tokens are "blocked" because appending them would
/// produce a repeated n-gram.
///
/// Returns an empty set when there are fewer than `n-1` tokens (no suffix to match).
#[must_use]
pub fn blocked_tokens_for_ngram(tokens: &[u32], n: usize) -> HashSet<u32> {
if n == 0 || tokens.len() < n {
return HashSet::new();
}
let suffix_len = n - 1;
let suffix = &tokens[tokens.len() - suffix_len..];
let mut blocked = HashSet::new();
// Slide a window of length `suffix_len` over all earlier positions.
// For each position i where tokens[i..i+suffix_len] == suffix, the token at
// position i+suffix_len would complete a repeated n-gram.
let end = tokens.len() - suffix_len; // exclusive upper bound for window starts
for i in 0..end {
if &tokens[i..i + suffix_len] == suffix {
blocked.insert(tokens[i + suffix_len]);
}
}
blocked
}
}
impl Default for BeamSearchDecoder {
fn default() -> Self {
Self::new(BeamSearchConfig::default())
}
}
// ---------------------------------------------------------------------------
// DiverseBeamSearchDecoder
// ---------------------------------------------------------------------------
/// Diverse beam search decoder (Vijayakumar et al. 2018).
///
/// Divides the total beam budget into `num_groups` groups of size
/// `config.beam_width / num_groups`. Each group runs standard beam search at each
/// step, but tokens already chosen by earlier groups in the same step are penalised
/// by `diversity_penalty` subtracted from their log-probability.
///
/// This encourages the returned hypotheses to differ from one another while still
/// being high-quality.
pub struct DiverseBeamSearchDecoder {
config: BeamSearchConfig,
/// Number of beam groups. Must divide `config.beam_width` evenly.
pub num_groups: usize,
/// Score penalty applied per token already chosen by a preceding group at the same step.
pub diversity_penalty: f32,
}
impl DiverseBeamSearchDecoder {
/// Create a diverse beam search decoder.
///
/// # Panics
///
/// Panics in debug mode if `num_groups` is 0 or exceeds `config.beam_width`.
#[must_use]
pub fn new(config: BeamSearchConfig, num_groups: usize, diversity_penalty: f32) -> Self {
debug_assert!(num_groups > 0, "num_groups must be at least 1");
debug_assert!(
num_groups <= config.beam_width,
"num_groups must not exceed beam_width"
);
Self {
config,
num_groups,
diversity_penalty,
}
}
/// Run diverse beam search.
///
/// Groups run sequentially at each step. Group `g` penalises every token that was
/// already selected as the first new token by groups `0..g` at the current step.
///
/// Returns up to `num_return` beams sorted by descending normalized score.
pub fn decode<F>(
&self,
initial_tokens: &[u32],
vocab_size: usize,
num_return: usize,
score_fn: F,
) -> Vec<Beam>
where
F: Fn(&[u32]) -> Vec<f32>,
{
let alpha = self.config.length_penalty_alpha;
let num_groups = self.num_groups.max(1);
// Each group maintains its own set of active beams.
let beams_per_group = (self.config.beam_width / num_groups).max(1);
let initial_beam = Beam {
tokens: initial_tokens.to_vec(),
log_prob: 0.0,
finished: false,
};
// One Vec<Beam> per group.
let mut group_active: Vec<Vec<Beam>> = (0..num_groups)
.map(|_| vec![initial_beam.clone()])
.collect();
let mut finished: Vec<Beam> = Vec::new();
for _step in 0..self.config.max_length {
// Track which tokens have been committed by each group this step so that
// subsequent groups can be penalised.
let mut chosen_tokens: Vec<HashSet<u32>> = vec![HashSet::new(); num_groups];
for g in 0..num_groups {
let active = &group_active[g];
if active.is_empty() {
continue;
}
// Tokens penalised are those committed by groups 0..g.
let penalised: HashSet<u32> = chosen_tokens[..g]
.iter()
.flat_map(|s| s.iter().copied())
.collect();
let mut candidates: Vec<Beam> = Vec::new();
for beam in active {
let log_probs = score_fn(&beam.tokens);
let blocked = if self.config.no_repeat_ngram_size > 0 {
BeamSearchDecoder::blocked_tokens_for_ngram(
&beam.tokens,
self.config.no_repeat_ngram_size,
)
} else {
HashSet::new()
};
let generated_so_far = beam.tokens.len().saturating_sub(initial_tokens.len());
let suppress_eos = generated_so_far < self.config.min_length;
// Apply diversity penalty before selecting top-k.
let adjusted: Vec<f32> = log_probs
.iter()
.enumerate()
.map(|(i, &lp)| {
if penalised.contains(&(i as u32)) {
lp - self.diversity_penalty
} else {
lp
}
})
.collect();
let fetch_k = beams_per_group * 2 + 2;
let top = BeamSearchDecoder::top_k_indices(&adjusted, fetch_k);
let mut added = 0usize;
for (token_id, lp) in top {
if added >= beams_per_group {
break;
}
let token_id = token_id as u32;
if blocked.contains(&token_id) {
continue;
}
if suppress_eos && token_id == self.config.eos_token_id {
continue;
}
// Record what each group committed.
chosen_tokens[g].insert(token_id);
added += 1;
// Use the un-penalised log-prob to keep beam scores comparable.
let raw_lp = if token_id < vocab_size as u32 {
log_probs[token_id as usize]
} else {
lp
};
let mut new_tokens = beam.tokens.clone();
new_tokens.push(token_id);
let new_lp = beam.log_prob + raw_lp;
let is_finished = token_id == self.config.eos_token_id;
candidates.push(Beam {
tokens: new_tokens,
log_prob: new_lp,
finished: is_finished,
});
}
}
// Sort and prune.
candidates.sort_by(|a, b| {
b.normalized_score(alpha)
.partial_cmp(&a.normalized_score(alpha))
.unwrap_or(std::cmp::Ordering::Equal)
});
let mut new_active: Vec<Beam> = Vec::new();
for beam in candidates.into_iter().take(beams_per_group * 2) {
if beam.finished {
finished.push(beam);
} else if new_active.len() < beams_per_group {
new_active.push(beam);
}
}
group_active[g] = new_active;
}
// Check whether all groups are exhausted.
if group_active.iter().all(|g| g.is_empty()) {
break;
}
}
// Collect remaining active beams as fallback results.
for group in group_active {
finished.extend(group);
}
finished.sort_by(|a, b| {
b.normalized_score(alpha)
.partial_cmp(&a.normalized_score(alpha))
.unwrap_or(std::cmp::Ordering::Equal)
});
finished.into_iter().take(num_return).collect()
}
}
// ---------------------------------------------------------------------------
// Tests
// ---------------------------------------------------------------------------
#[cfg(test)]
mod tests {
use super::*;
// ---- Helper ---------------------------------------------------------------
fn make_score_fn(vocab_size: usize, eos_id: u32) -> impl Fn(&[u32]) -> Vec<f32> {
move |_tokens: &[u32]| {
let mut lp = vec![-10.0_f32; vocab_size];
lp[5] = -0.1;
lp[eos_id as usize] = -0.5;
lp
}
}
// ---- Beam::normalized_score -----------------------------------------------
#[test]
fn test_beam_normalized_score_alpha0() {
// alpha=0 => score == log_prob regardless of length
let mut beam = Beam::new(-3.0);
beam.tokens = vec![1, 2, 3];
let score = beam.normalized_score(0.0);
assert!(
(score - (-3.0)).abs() < 1e-6,
"alpha=0 should return raw log_prob, got {score}"
);
}
#[test]
fn test_beam_normalized_score_alpha1() {
// alpha=1, len=1 => penalty = ((5+1)/6)^1 = 1.0 => score == log_prob
let mut beam = Beam::new(-6.0);
beam.tokens = vec![42];
let score = beam.normalized_score(1.0);
let expected = -6.0_f32 / ((5.0 + 1.0) / 6.0_f32).powf(1.0);
assert!(
(score - expected).abs() < 1e-5,
"Expected {expected}, got {score}"
);
}
#[test]
fn test_beam_normalized_score_longer_beam_wins() {
// Two beams with the same total log_prob; the longer one gets a better
// (less negative) normalised score when alpha > 0.
let mut short_beam = Beam::new(-6.0);
short_beam.tokens = vec![1];
let mut long_beam = Beam::new(-6.0);
long_beam.tokens = vec![1, 2, 3, 4, 5];
let alpha = 0.6;
let short_score = short_beam.normalized_score(alpha);
let long_score = long_beam.normalized_score(alpha);
assert!(
long_score > short_score,
"Longer beam should have higher normalised score: long={long_score}, short={short_score}"
);
}
// ---- top_k_indices --------------------------------------------------------
#[test]
fn test_top_k_indices_returns_k() {
let lp: Vec<f32> = (0..20).map(|i| -(i as f32)).collect();
let result = BeamSearchDecoder::top_k_indices(&lp, 5);
assert_eq!(result.len(), 5, "Should return exactly 5 elements");
}
#[test]
fn test_top_k_indices_descending() {
let lp = vec![-3.0, -1.0, -4.0, -1.5, -2.0];
let result = BeamSearchDecoder::top_k_indices(&lp, 3);
assert_eq!(result.len(), 3);
// Verify descending order.
for window in result.windows(2) {
assert!(
window[0].1 >= window[1].1,
"Not sorted descending: {:?}",
result
);
}
// Top element should be index 1 (log_prob -1.0).
assert_eq!(result[0].0, 1);
}
#[test]
fn test_top_k_larger_than_vocab() {
let lp = vec![-1.0, -2.0, -3.0];
// k=10 > vocab_size=3 => return all 3
let result = BeamSearchDecoder::top_k_indices(&lp, 10);
assert_eq!(result.len(), 3);
}
// ---- blocked_tokens_for_ngram ---------------------------------------------
#[test]
fn test_blocked_ngrams_empty() {
let blocked = BeamSearchDecoder::blocked_tokens_for_ngram(&[], 2);
assert!(blocked.is_empty());
}
#[test]
fn test_blocked_ngrams_bigram() {
// Sequence: [A, B, C, A] — suffix is [A], look for earlier [A] => B is blocked.
let tokens: Vec<u32> = vec![10, 20, 30, 10];
let blocked = BeamSearchDecoder::blocked_tokens_for_ngram(&tokens, 2);
// Earlier [10] at position 0 is followed by 20 => 20 is blocked.
assert!(
blocked.contains(&20),
"Token 20 should be blocked; got {blocked:?}"
);
}
#[test]
fn test_blocked_ngrams_trigram() {
// [A, B, C, A, B] — suffix is [A, B], look for earlier [A, B] at pos 0 => C(=30) blocked.
let tokens: Vec<u32> = vec![10, 20, 30, 10, 20];
let blocked = BeamSearchDecoder::blocked_tokens_for_ngram(&tokens, 3);
assert!(
blocked.contains(&30),
"Token 30 should be blocked; got {blocked:?}"
);
}
#[test]
fn test_blocked_ngrams_no_repeat() {
// No repeated n-gram => empty blocked set.
let tokens: Vec<u32> = vec![1, 2, 3, 4, 5];
let blocked = BeamSearchDecoder::blocked_tokens_for_ngram(&tokens, 2);
assert!(
blocked.is_empty(),
"No repeated bigrams, should be empty; got {blocked:?}"
);
}
// ---- decode ---------------------------------------------------------------
#[test]
fn test_decode_single_step() {
// With beam_width=1 and a deterministic scorer, we should get exactly one beam
// after one step (prompt=[1], max_length=1).
let vocab_size = 10;
let eos_id = 2_u32;
let config = BeamSearchConfig {
beam_width: 1,
max_length: 1,
eos_token_id: eos_id,
length_penalty_alpha: 0.0,
min_length: 0,
no_repeat_ngram_size: 0,
};
let decoder = BeamSearchDecoder::new(config);
let score_fn = |_: &[u32]| {
let mut lp = vec![-10.0_f32; vocab_size];
lp[5] = -0.1;
lp[eos_id as usize] = -0.5;
lp
};
let results = decoder.decode(&[1], vocab_size, 1, score_fn);
assert_eq!(results.len(), 1);
// With min_length=0, EOS is allowed. But token 5 has higher lp, so beam should end with 5.
let last = *results[0].tokens.last().unwrap();
assert_eq!(last, 5, "Greedy first token should be 5");
}
#[test]
fn test_decode_returns_num_return() {
let vocab_size = 32;
let config = BeamSearchConfig {
beam_width: 4,
max_length: 10,
eos_token_id: 2,
length_penalty_alpha: 0.6,
min_length: 1,
no_repeat_ngram_size: 0,
};
let decoder = BeamSearchDecoder::new(config);
let results = decoder.decode(&[1], vocab_size, 2, make_score_fn(vocab_size, 2));
assert_eq!(results.len(), 2, "Should return exactly 2 beams");
}
#[test]
fn test_decode_eos_stops_beam() {
// Force EOS to be the highest-scoring token so the beam finishes immediately.
let vocab_size = 10;
let eos_id = 2_u32;
let config = BeamSearchConfig {
beam_width: 1,
max_length: 50,
eos_token_id: eos_id,
length_penalty_alpha: 0.0,
min_length: 0,
no_repeat_ngram_size: 0,
};
let decoder = BeamSearchDecoder::new(config);
let score_fn = |_: &[u32]| {
let mut lp = vec![-10.0_f32; vocab_size];
lp[eos_id as usize] = -0.1; // EOS is best
lp
};
let results = decoder.decode(&[1], vocab_size, 1, score_fn);
assert_eq!(results.len(), 1);
assert!(
results[0].finished,
"Beam should be marked finished after EOS"
);
assert_eq!(
*results[0].tokens.last().unwrap(),
eos_id,
"Last token should be EOS"
);
}
#[test]
fn test_decode_min_length_suppresses_eos() {
// With min_length=5 the beam must generate at least 5 tokens before EOS is allowed.
let vocab_size = 10;
let eos_id = 2_u32;
let config = BeamSearchConfig {
beam_width: 1,
max_length: 20,
eos_token_id: eos_id,
length_penalty_alpha: 0.0,
min_length: 5,
no_repeat_ngram_size: 0,
};
let decoder = BeamSearchDecoder::new(config);
// EOS is always the best token, but is suppressed until min_length is met.
let score_fn = |_: &[u32]| {
let mut lp = vec![-10.0_f32; vocab_size];
lp[5] = -0.2;
lp[eos_id as usize] = -0.1; // EOS beats others, but suppressed until len>=5
lp
};
let results = decoder.decode(&[1], vocab_size, 1, score_fn);
assert!(!results.is_empty());
// The beam must have at least 1 (prompt) + 5 (generated) = 6 tokens before finishing.
let generated = results[0].tokens.len() - 1; // subtract prompt token
assert!(
generated >= 5,
"Should have generated at least 5 tokens before EOS, got {generated}"
);
}
#[test]
fn test_decode_always_returns_result() {
// Even if no EOS is emitted within max_length, active beams are returned.
let vocab_size = 10;
let config = BeamSearchConfig {
beam_width: 2,
max_length: 5,
eos_token_id: 99, // EOS id that will never be produced by the scorer
length_penalty_alpha: 0.0,
min_length: 0,
no_repeat_ngram_size: 0,
};
let decoder = BeamSearchDecoder::new(config);
let score_fn = |_: &[u32]| {
let mut lp = vec![-10.0_f32; vocab_size];
lp[5] = -0.1;
lp[3] = -0.2;
lp
};
let results = decoder.decode(&[1], vocab_size, 1, score_fn);
assert!(
!results.is_empty(),
"Should always return at least one beam"
);
}
#[test]
fn test_decode_beam_width_1_greedy() {
// beam_width=1 should behave identically to greedy decoding.
let vocab_size = 10;
let eos_id = 2_u32;
let config = BeamSearchConfig {
beam_width: 1,
max_length: 8,
eos_token_id: eos_id,
length_penalty_alpha: 0.0,
min_length: 3,
no_repeat_ngram_size: 0,
};
let decoder = BeamSearchDecoder::new(config);
// Scorer alternates between two tokens to create a deterministic sequence.
let score_fn = |tokens: &[u32]| {
let mut lp = vec![-10.0_f32; vocab_size];
if tokens.len() % 2 == 0 {
lp[7] = -0.1;
} else {
lp[3] = -0.1;
}
lp[eos_id as usize] = -0.5;
lp
};
let results = decoder.decode(&[1], vocab_size, 1, score_fn);
assert_eq!(results.len(), 1);
// Manually simulate greedy: prompt=[1](len=1, odd), so first pick is lp[3].
// Sequence: [1, 3, 7, 3, ...] until min_length=3 allows EOS but EOS < 0.5 vs -0.1 so it
// never fires. We just verify the length constraint held.
let generated = results[0].tokens.len() - 1;
assert!(generated >= 3 || results[0].finished);
}
#[test]
fn test_decode_prefers_higher_logprob() {
// With two beams and clear score separation, the higher-lp beam should rank first.
let vocab_size = 10;
let eos_id = 2_u32;
let config = BeamSearchConfig {
beam_width: 2,
max_length: 2,
eos_token_id: eos_id,
length_penalty_alpha: 0.0,
min_length: 0,
no_repeat_ngram_size: 0,
};
let decoder = BeamSearchDecoder::new(config);
// Always give token 5 a much better score.
let score_fn = |_: &[u32]| {
let mut lp = vec![-20.0_f32; vocab_size];
lp[5] = -0.01; // very good
lp[3] = -5.0; // much worse
lp[eos_id as usize] = -0.5;
lp
};
let results = decoder.decode(&[1], vocab_size, 2, score_fn);
// The top beam should contain token 5.
assert!(
results[0].tokens.contains(&5),
"Top beam should contain token 5"
);
}
#[test]
fn test_default_config() {
let cfg = BeamSearchConfig::default();
assert_eq!(cfg.beam_width, 4);
assert_eq!(cfg.max_length, 128);
assert_eq!(cfg.eos_token_id, 2);
assert!((cfg.length_penalty_alpha - 0.6).abs() < 1e-6);
assert_eq!(cfg.min_length, 1);
assert_eq!(cfg.no_repeat_ngram_size, 0);
}
// ---- DiverseBeamSearchDecoder --------------------------------------------
#[test]
fn test_diverse_decode_returns_results() {
let vocab_size = 20;
let config = BeamSearchConfig {
beam_width: 4,
max_length: 8,
eos_token_id: 2,
length_penalty_alpha: 0.6,
min_length: 1,
no_repeat_ngram_size: 0,
};
let decoder = DiverseBeamSearchDecoder::new(config, 2, 1.0);
let results = decoder.decode(&[1], vocab_size, 2, make_score_fn(vocab_size, 2));
assert!(
!results.is_empty(),
"Diverse beam search should return at least one beam"
);
}
#[test]
fn test_diverse_decode_group1_equals_standard() {
// With num_groups=1 and diversity_penalty=0, diverse beam search should be
// equivalent to standard beam search.
let vocab_size = 20;
let config = BeamSearchConfig {
beam_width: 4,
max_length: 6,
eos_token_id: 2,
length_penalty_alpha: 0.0,
min_length: 1,
no_repeat_ngram_size: 0,
};
let std_decoder = BeamSearchDecoder::new(config.clone());
let div_decoder = DiverseBeamSearchDecoder::new(config, 1, 0.0);
let sf = make_score_fn(vocab_size, 2);
let sf2 = make_score_fn(vocab_size, 2);
let std_results = std_decoder.decode(&[1], vocab_size, 1, sf);
let div_results = div_decoder.decode(&[1], vocab_size, 1, sf2);
assert!(!std_results.is_empty());
assert!(!div_results.is_empty());
// Both should produce the same top beam sequence.
assert_eq!(
std_results[0].tokens, div_results[0].tokens,
"1-group diverse should match standard beam search"
);
}
}