use crate::llm::tinyllama::TinyLlama;
use crate::tensor::Tensor;
pub struct BeamConfig {
pub num_beams: usize,
pub max_new_tokens: usize,
pub length_penalty: f32,
pub eos_token_id: usize,
pub pad_token_id: usize,
}
pub struct BeamSearch {
pub model: TinyLlama,
pub cfg: BeamConfig,
}
impl BeamSearch {
pub fn generate(&self, prefix_embeds: &Tensor, attention_mask: &[f32]) -> Vec<usize> {
if self.cfg.num_beams <= 1 {
return self.generate_greedy(prefix_embeds, attention_mask);
}
let d = prefix_embeds.shape[2];
let mut beams: Vec<(Vec<usize>, f32, Tensor, Vec<f32>)> = vec![(
Vec::new(),
0.0,
prefix_embeds.clone(),
attention_mask.to_vec(),
)];
for _step in 0..self.cfg.max_new_tokens {
let mut candidates = Vec::new();
for (tokens, score, embeds, mask) in &beams {
let hidden = self.model.forward_embeds(embeds, Some(mask));
let last = hidden.narrow(1, hidden.shape[1] - 1, 1);
let logits = self.model.logits(&last);
let vocab = logits.shape[2];
let log_probs = log_softmax(&logits.data[..vocab]);
let mut ranked: Vec<(usize, f32)> = log_probs
.into_iter()
.enumerate()
.map(|(i, lp)| (i, lp))
.collect();
ranked.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
for (tid, lp) in ranked.into_iter().take(self.cfg.num_beams) {
let mut new_tokens = tokens.clone();
new_tokens.push(tid);
let tok_emb = self.model.embed(&[tid]);
let mut new_embeds_data = embeds.data.clone();
new_embeds_data.extend(tok_emb.data);
let new_len = embeds.shape[1] + 1;
let new_embeds = Tensor::from_vec(new_embeds_data, vec![1, new_len, d]);
let mut new_mask = mask.clone();
new_mask.push(1.0);
let len = new_tokens.len() as f32;
let norm = len.powf(self.cfg.length_penalty);
let new_score = score + lp / norm.max(1.0);
candidates.push((new_tokens, new_score, new_embeds, new_mask));
}
}
candidates.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
beams = candidates.into_iter().take(self.cfg.num_beams).collect();
if beams
.iter()
.any(|(t, _, _, _)| t.last() == Some(&self.cfg.eos_token_id))
{
break;
}
}
beams
.into_iter()
.max_by(|a, b| a.1.partial_cmp(&b.1).unwrap_or(std::cmp::Ordering::Equal))
.map(|(t, _, _, _)| t)
.unwrap_or_default()
}
fn generate_greedy(&self, prefix_embeds: &Tensor, attention_mask: &[f32]) -> Vec<usize> {
let mut tokens = Vec::new();
let mut embeds = prefix_embeds.clone();
let mut mask = attention_mask.to_vec();
for _ in 0..self.cfg.max_new_tokens {
let hidden = self.model.forward_embeds(&embeds, Some(&mask));
let last = hidden.narrow(1, hidden.shape[1] - 1, 1);
let logits = self.model.logits(&last);
let vocab = logits.shape[2];
let tid = argmax(&logits.data[..vocab]);
tokens.push(tid);
if tid == self.cfg.eos_token_id {
break;
}
let tok_emb = self.model.embed(&[tid]);
embeds.data.extend(tok_emb.data);
embeds.shape[1] += 1;
mask.push(1.0);
}
tokens
}
}
fn log_softmax(logits: &[f32]) -> Vec<f32> {
let max_l = logits.iter().copied().fold(f32::NEG_INFINITY, f32::max);
let exps: Vec<f32> = logits.iter().map(|&x| (x - max_l).exp()).collect();
let sum: f32 = exps.iter().sum();
exps.iter().map(|&e| (e / sum).ln()).collect()
}
fn argmax(x: &[f32]) -> usize {
x.iter()
.enumerate()
.max_by(|(_, a), (_, b)| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal))
.map(|(i, _)| i)
.unwrap_or(0)
}