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
9pub 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 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}