brain2qwerty 0.0.1

Brain2Qwerty V1/V2 MEG neural decoding inference in Rust (parity-tested vs Python)
Documentation
//! HuggingFace-compatible beam search with inputs_embeds prefix.

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)
}