Skip to main content

rlx_moshi/
lm.rs

1use crate::config::LmConfig;
2use crate::depformer::DepFormer;
3use crate::nn::{Embedding, linear, rms_norm};
4use crate::transformer::StreamingTransformer;
5use anyhow::{Context, Result};
6use ndarray::{Array1, Array2};
7use std::collections::HashMap;
8
9/// Moshi temporal LM (Helium) + optional DepFormer depth decoder.
10pub struct LmModel {
11    cfg: LmConfig,
12    text_emb: Embedding,
13    audio_embs: Vec<Embedding>,
14    text_linear: Array2<f32>,
15    out_norm_alpha: Array1<f32>,
16    transformer: StreamingTransformer,
17    depformer: Option<DepFormer>,
18}
19
20impl LmModel {
21    pub fn open(cfg: LmConfig, weights: HashMap<String, (Vec<f32>, Vec<usize>)>) -> Result<Self> {
22        let text_emb = Embedding {
23            weight: take_mat(&weights, "text_emb.weight")?,
24        };
25        let mut audio_embs = Vec::with_capacity(cfg.audio_codebooks);
26        for i in 0..cfg.audio_codebooks {
27            audio_embs.push(Embedding {
28                weight: take_mat(&weights, &format!("emb.{i}.weight"))?,
29            });
30        }
31        let text_linear = take_mat(&weights, "text_linear.weight")?;
32        let out_norm_alpha = take_vec1(&weights, "out_norm.alpha")?;
33        let transformer = StreamingTransformer::build(&cfg.transformer, &weights)?;
34        let depformer = match &cfg.depformer {
35            None => None,
36            Some(df) => Some(DepFormer::build(
37                df,
38                cfg.text_in_vocab_size,
39                cfg.audio_vocab_size,
40                cfg.transformer.d_model,
41                &weights,
42            )?),
43        };
44        Ok(Self {
45            cfg,
46            text_emb,
47            audio_embs,
48            text_linear,
49            out_norm_alpha,
50            transformer,
51            depformer,
52        })
53    }
54
55    pub fn config(&self) -> &LmConfig {
56        &self.cfg
57    }
58
59    pub fn reset_state(&mut self) {
60        self.transformer.reset_state();
61    }
62
63    pub fn text_start_token(&self) -> u32 {
64        self.cfg.text_in_vocab_size as u32 - 1
65    }
66
67    pub fn audio_pad_token(&self) -> u32 {
68        self.cfg.audio_vocab_size as u32 - 1
69    }
70
71    /// Single streaming step: sum embeddings → temporal transformer → text logits + hidden.
72    pub fn forward_step(
73        &mut self,
74        text_token: Option<u32>,
75        audio_tokens: &[Option<u32>],
76    ) -> Result<(Array1<f32>, Array1<f32>)> {
77        let d = self.cfg.transformer.d_model;
78        let mut emb = vec![0.0f32; d];
79        if let Some(tt) = text_token {
80            let e = self.text_emb.forward_one(tt);
81            for (i, v) in e.iter().enumerate() {
82                emb[i] += v;
83            }
84        }
85        for (cb, tok) in audio_tokens.iter().zip(self.audio_embs.iter()) {
86            if let Some(t) = cb {
87                let e = tok.forward_one(*t);
88                for (i, v) in e.iter().enumerate() {
89                    emb[i] += v;
90                }
91            }
92        }
93        let x = Array2::from_shape_vec((1, d), emb)?;
94        let h = self.transformer.forward(&x);
95        let normed = rms_norm(h.view(), &self.out_norm_alpha);
96        let logits = linear(normed.view(), &self.text_linear);
97        Ok((logits.row(0).to_owned(), h.row(0).to_owned()))
98    }
99
100    pub fn depformer_sample(
101        &mut self,
102        hidden: &Array1<f32>,
103        text_token: Option<u32>,
104        forced: &[Option<u32>],
105        lp: &mut crate::sampling::LogitsProcessor,
106    ) -> Result<Option<Vec<u32>>> {
107        match self.depformer.as_mut() {
108            None => Ok(None),
109            Some(df) => Ok(Some(df.sample(hidden, text_token, forced, lp)?)),
110        }
111    }
112}
113
114fn take_mat(weights: &HashMap<String, (Vec<f32>, Vec<usize>)>, key: &str) -> Result<Array2<f32>> {
115    let (data, shape) = weights
116        .get(key)
117        .with_context(|| format!("missing weight {key}"))?;
118    Ok(Array2::from_shape_vec((shape[0], shape[1]), data.clone())?)
119}
120
121fn take_vec1(weights: &HashMap<String, (Vec<f32>, Vec<usize>)>, key: &str) -> Result<Array1<f32>> {
122    let (data, shape) = weights
123        .get(key)
124        .with_context(|| format!("missing weight {key}"))?;
125    let _n: usize = shape.iter().product();
126    Ok(Array1::from_vec(data.clone()))
127}