brain2qwerty 0.0.1

Brain2Qwerty V1/V2 MEG neural decoding inference in Rust (parity-tested vs Python)
Documentation
//! End-to-end Brain2Qwerty **V2** inference: encoder → CTC → segmenter → LLM beam search.

use crate::config::Brain2QwertyConfig;
use crate::decode::{build_intra_word_pooler, ctc_greedy_decode, CTCSpaceSegmenter};
use crate::llm::{BeamConfig, BeamSearch, LlamaConfig, TinyLlama};
use crate::model::conv_conformer::{ConvConformer, ConvConformerOutput};
use crate::tensor::Tensor;
use crate::weights::load_safetensors;

#[cfg(feature = "rlx-encoder")]
use crate::model_rlx::ConvConformerRlx;
#[cfg(feature = "rlx-encoder")]
use crate::rlx_device::{device_label, resolve_rlx_device};

pub struct InferenceInput {
    /// MEG tensor `(batch, time, channels)`.
    pub neuros: Tensor,
    /// Subject index per batch row (for per-subject layers / merger).
    pub subject_ids: Vec<usize>,
    /// Optional 2-D channel positions `(batch, channels, 2)`.
    pub chan_pos: Option<Tensor>,
}

pub struct PipelineOutput {
    /// Greedy CTC decode string (first batch item).
    pub ctc_text: String,
    /// Beam-search LLM output, or CTC text when no LLM is loaded.
    pub pred_text: String,
    pub z_final: Tensor,
    pub c_out: Tensor,
    pub word_embeds: Vec<Tensor>,
}

/// Encoder backend selection (`rlx` is default when `rlx-encoder` is enabled).
#[derive(Clone, Debug)]
pub struct PipelineOptions {
    /// `auto`, `cpu`, `metal`, `mlx`, `cuda`, …
    pub rlx_device: String,
    /// When true, use the pure-Rust reference encoder instead of RLX.
    pub use_rust_encoder: bool,
}

impl Default for PipelineOptions {
    fn default() -> Self {
        Self {
            rlx_device: "auto".into(),
            use_rust_encoder: false,
        }
    }
}

enum Encoder {
    Rust(ConvConformer),
    #[cfg(feature = "rlx-encoder")]
    Rlx(ConvConformerRlx),
}

impl Encoder {
    fn forward(
        &mut self,
        neuros: &Tensor,
        subject_ids: &[usize],
        chan_pos: Option<&Tensor>,
    ) -> ConvConformerOutput {
        match self {
            Encoder::Rust(e) => e.forward(neuros, subject_ids, chan_pos),
            #[cfg(feature = "rlx-encoder")]
            Encoder::Rlx(e) => e.forward(neuros, subject_ids, chan_pos),
        }
    }
}

pub struct Pipeline {
    pub config: Brain2QwertyConfig,
    encoder: Encoder,
    /// Resolved backend label, e.g. `rlx-cpu` or `rust`.
    pub encoder_backend: String,
    pub segmenter: CTCSpaceSegmenter,
    pub word_adapter_w: Option<Tensor>,
    pub word_adapter_b: Option<Tensor>,
    pub llm: Option<TinyLlama>,
    pub tokenizer: Option<tokenizers::Tokenizer>,
}

impl Pipeline {
    /// Load tiny CI config + weights (matches `generate_parity_refs.py`).
    pub fn from_tiny_paths(encoder_weights: &str, llm_dir: Option<&str>) -> anyhow::Result<Self> {
        Self::from_tiny_paths_with_options(encoder_weights, llm_dir, PipelineOptions::default())
    }

    pub fn from_tiny_paths_with_options(
        encoder_weights: &str,
        llm_dir: Option<&str>,
        opts: PipelineOptions,
    ) -> anyhow::Result<Self> {
        let config = Brain2QwertyConfig::tiny();
        let mut store = load_safetensors(encoder_weights)?;
        Self::build(config, &mut store, llm_dir, encoder_weights, opts)
    }

    pub fn from_paths(
        config_path: &str,
        encoder_weights: &str,
        llm_dir: Option<&str>,
    ) -> anyhow::Result<Self> {
        Self::from_paths_with_options(
            config_path,
            encoder_weights,
            llm_dir,
            PipelineOptions::default(),
        )
    }

    /// Production config from YAML with explicit encoder backend options.
    pub fn from_paths_with_options(
        config_path: &str,
        encoder_weights: &str,
        llm_dir: Option<&str>,
        opts: PipelineOptions,
    ) -> anyhow::Result<Self> {
        let config = Brain2QwertyConfig::from_yaml(config_path)?;
        let mut store = load_safetensors(encoder_weights)?;
        Self::build(config, &mut store, llm_dir, encoder_weights, opts)
    }

    fn build(
        config: Brain2QwertyConfig,
        store: &mut crate::weights::WeightStore,
        llm_dir: Option<&str>,
        encoder_weights: &str,
        opts: PipelineOptions,
    ) -> anyhow::Result<Self> {
        let (encoder, encoder_backend) = build_encoder(&config, store, encoder_weights, &opts)?;
        let pooler = build_intra_word_pooler(
            config.brain_model_config.dim,
            config.inference.word_pool_n_layers,
            Some(store),
        );
        let segmenter = CTCSpaceSegmenter {
            include_blanks: config.inference.seg_include_blanks,
            min_word_frames: 1,
            pooler,
        };
        let (word_adapter_w, word_adapter_b) = (
            store
                .get("word_proj_adapter.weight")
                .map(|p| crate::weights::param_to_tensor(p)),
            store
                .get("word_proj_adapter.bias")
                .map(|p| crate::weights::param_to_tensor(p)),
        );
        let (llm, tokenizer) = if let Some(dir) = llm_dir {
            let llm_cfg = LlamaConfig::from_json(&format!("{dir}/config.json"))?;
            let mut llm_store = load_safetensors(&format!("{dir}/model.safetensors"))?;
            crate::llm::merge_lora_into_base(
                &mut llm_store,
                config.inference.lora_rank,
                config.inference.lora_alpha,
                &config.inference.lora_target_modules,
            )?;
            let llm = TinyLlama::load(&llm_store, llm_cfg, "")?;
            let tok = tokenizers::Tokenizer::from_file(format!("{dir}/tokenizer.json")).ok();
            (Some(llm), tok)
        } else {
            (None, None)
        };
        Ok(Self {
            config,
            encoder,
            encoder_backend,
            segmenter,
            word_adapter_w,
            word_adapter_b,
            llm,
            tokenizer,
        })
    }

    /// Run full V2 inference on one batch.
    pub fn run(&mut self, input: &InferenceInput) -> anyhow::Result<PipelineOutput> {
        let enc = self
            .encoder
            .forward(&input.neuros, &input.subject_ids, input.chan_pos.as_ref());
        let ctc_texts = ctc_greedy_decode(&enc.c_out);
        let ctc_text = ctc_texts.first().cloned().unwrap_or_default();
        let mut word_embeds = self.segmenter.forward(&enc.z_final, &enc.c_out);
        for w in &mut word_embeds {
            if let Some(ref aw) = self.word_adapter_w {
                let adapted = adapt_words(w, aw, self.word_adapter_b.as_ref());
                *w = adapted;
            }
        }
        let pred_text = if let (Some(llm), Some(tok)) = (&self.llm, &self.tokenizer) {
            let (prefix, mask) = build_prefix_embeds(
                llm,
                tok,
                &self.config.inference.sys_prompt,
                &self.config.inference.mid_prompt,
                &self.config.inference.resp_prompt,
                &ctc_text,
                word_embeds.first(),
            )?;
            let beam = BeamSearch {
                model: llm.clone(),
                cfg: BeamConfig {
                    num_beams: self.config.inference.num_beams,
                    max_new_tokens: self.config.inference.max_new_tokens,
                    length_penalty: self.config.inference.length_penalty,
                    eos_token_id: tok.token_to_id("</s>").unwrap_or(1) as usize,
                    pad_token_id: tok.token_to_id("<pad>").unwrap_or(0) as usize,
                },
            };
            let ids = beam.generate(&prefix, &mask);
            tok.decode(&ids.iter().map(|&i| i as u32).collect::<Vec<_>>(), true)
                .unwrap_or_default()
        } else {
            ctc_text.clone()
        };
        Ok(PipelineOutput {
            ctc_text,
            pred_text,
            z_final: enc.z_final,
            c_out: enc.c_out,
            word_embeds,
        })
    }
}

fn build_encoder(
    config: &Brain2QwertyConfig,
    store: &mut crate::weights::WeightStore,
    weights_path: &str,
    opts: &PipelineOptions,
) -> anyhow::Result<(Encoder, String)> {
    #[cfg(feature = "rlx-encoder")]
    if !opts.use_rust_encoder {
        let device = resolve_rlx_device(&opts.rlx_device)?;
        let encoder =
            ConvConformerRlx::from_config_and_weights(config, weights_path)?.with_device(device);
        let label = format!("rlx-{}", device_label(device));
        return Ok((Encoder::Rlx(encoder), label));
    }
    let encoder = ConvConformer::from_config_and_weights(&config.brain_model_config, store, "")?;
    Ok((Encoder::Rust(encoder), "rust".into()))
}

fn adapt_words(w: &Tensor, weight: &Tensor, bias: Option<&Tensor>) -> Tensor {
    let (n, d_in) = (w.shape[0], w.shape[1]);
    let d_out = weight.shape[0];
    let mut out = vec![0.0f32; n * d_out];
    for wi in 0..n {
        for o in 0..d_out {
            let mut sum = 0.0f32;
            for i in 0..d_in {
                sum += w.data[wi * d_in + i] * weight.data[o * d_in + i];
            }
            if let Some(b) = bias {
                sum += b.data[o];
            }
            out[wi * d_out + o] = sum;
        }
    }
    Tensor::from_vec(out, vec![n, d_out])
}

fn build_prefix_embeds(
    llm: &TinyLlama,
    tok: &tokenizers::Tokenizer,
    sys: &str,
    mid: &str,
    resp: &str,
    ctc_text: &str,
    words: Option<&Tensor>,
) -> anyhow::Result<(Tensor, Vec<f32>)> {
    let mut parts: Vec<Tensor> = Vec::new();
    for text in [sys] {
        let enc = tok
            .encode(text, false)
            .map_err(|e| anyhow::anyhow!("{e}"))?;
        let ids: Vec<usize> = enc.get_ids().iter().map(|&id| id as usize).collect();
        if !ids.is_empty() {
            parts.push(llm.embed(&ids));
        }
    }
    let enc = tok
        .encode(ctc_text, false)
        .map_err(|e| anyhow::anyhow!("{e}"))?;
    let ids: Vec<usize> = enc.get_ids().iter().map(|&id| id as usize).collect();
    if !ids.is_empty() {
        parts.push(llm.embed(&ids));
    }
    if let Some(w) = words {
        if w.shape[0] > 0 {
            for text in [mid] {
                let enc = tok
                    .encode(text, false)
                    .map_err(|e| anyhow::anyhow!("{e}"))?;
                let ids: Vec<usize> = enc.get_ids().iter().map(|&id| id as usize).collect();
                if !ids.is_empty() {
                    parts.push(llm.embed(&ids));
                }
            }
            let (n, d) = (w.shape[0], w.shape[1]);
            parts.push(Tensor::from_vec(w.data.clone(), vec![1, n, d]));
        }
    }
    for text in [resp] {
        let enc = tok
            .encode(text, false)
            .map_err(|e| anyhow::anyhow!("{e}"))?;
        let ids: Vec<usize> = enc.get_ids().iter().map(|&id| id as usize).collect();
        if !ids.is_empty() {
            parts.push(llm.embed(&ids));
        }
    }
    if parts.is_empty() {
        return Ok((Tensor::zeros(&[1, 1, llm.config.hidden_size]), vec![1.0]));
    }
    let d = parts[0].shape[2];
    let total_len: usize = parts.iter().map(|p| p.shape[1]).sum();
    let mut data = vec![0.0f32; total_len * d];
    let mut offset = 0usize;
    for p in parts {
        let len = p.shape[1];
        data[offset * d..(offset + len) * d].copy_from_slice(&p.data);
        offset += len;
    }
    Ok((
        Tensor::from_vec(data, vec![1, total_len, d]),
        vec![1.0; total_len],
    ))
}