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 {
pub neuros: Tensor,
pub subject_ids: Vec<usize>,
pub chan_pos: Option<Tensor>,
}
pub struct PipelineOutput {
pub ctc_text: String,
pub pred_text: String,
pub z_final: Tensor,
pub c_out: Tensor,
pub word_embeds: Vec<Tensor>,
}
#[derive(Clone, Debug)]
pub struct PipelineOptions {
pub rlx_device: String,
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,
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 {
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(),
)
}
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,
})
}
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],
))
}