Skip to main content

ferrum_models/multimodal/
whisper.rs

1//! Whisper ASR model — custom forward pass.
2//!
3//! candle loads weights only; forward pass is ours so Metal/CUDA work
4//! without depending on candle-nn's LayerNorm CustomOp.
5
6use candle_core::{DType, Device as CandleDevice, IndexOp, Module, Tensor, D};
7use candle_nn::{Conv1d, Conv1dConfig, VarBuilder};
8use candle_transformers::models::whisper::{self, Config};
9use ferrum_types::{FerrumError, Result};
10use parking_lot::Mutex;
11use tracing::info;
12
13// ── Manual softmax (works on CPU/Metal/CUDA) ────────────────────────────
14
15fn softmax_last_dim(x: &Tensor) -> candle_core::Result<Tensor> {
16    let max = x.max_keepdim(D::Minus1)?;
17    let shifted = x.broadcast_sub(&max)?;
18    let exp = shifted.exp()?;
19    let sum = exp.sum_keepdim(D::Minus1)?;
20    exp.broadcast_div(&sum)
21}
22
23// ── Manual LayerNorm (pure tensor ops — works on CPU/Metal/CUDA) ─────────
24
25struct LayerNorm {
26    weight: Tensor,
27    bias: Tensor,
28    eps: f64,
29}
30
31impl LayerNorm {
32    fn load(size: usize, eps: f64, vb: VarBuilder) -> candle_core::Result<Self> {
33        let weight = vb.get(size, "weight")?;
34        let bias = vb.get(size, "bias")?;
35        Ok(Self { weight, bias, eps })
36    }
37
38    fn forward(&self, x: &Tensor) -> candle_core::Result<Tensor> {
39        let x_dtype = x.dtype();
40        let x = x.to_dtype(DType::F32)?;
41        let mean = x.mean_keepdim(D::Minus1)?;
42        let diff = x.broadcast_sub(&mean)?;
43        let var = diff.sqr()?.mean_keepdim(D::Minus1)?;
44        let norm = diff.broadcast_div(&(var + self.eps)?.sqrt()?)?;
45        let norm = norm.to_dtype(x_dtype)?;
46        norm.broadcast_mul(&self.weight)?.broadcast_add(&self.bias)
47    }
48}
49
50// ── Linear (weight + optional bias) ─────────────────────────────────────
51
52struct Linear {
53    weight: Tensor,
54    bias: Option<Tensor>,
55}
56
57impl Linear {
58    fn load(in_: usize, out: usize, vb: VarBuilder) -> candle_core::Result<Self> {
59        let weight = vb.get((out, in_), "weight")?;
60        let bias = vb.get(out, "bias").ok();
61        Ok(Self { weight, bias })
62    }
63
64    fn load_no_bias(in_: usize, out: usize, vb: VarBuilder) -> candle_core::Result<Self> {
65        let weight = vb.get((out, in_), "weight")?;
66        Ok(Self { weight, bias: None })
67    }
68
69    fn forward(&self, x: &Tensor) -> candle_core::Result<Tensor> {
70        let wt = self.weight.t()?;
71        // Broadcast weight for batched matmul: (out, in) → (batch, in, out)
72        let y = if x.dims().len() == 3 {
73            let b = x.dim(0)?;
74            x.matmul(&wt.broadcast_left(b)?)?
75        } else {
76            x.matmul(&wt)?
77        };
78        match &self.bias {
79            Some(b) => y.broadcast_add(b),
80            None => Ok(y),
81        }
82    }
83}
84
85// ── Multi-Head Attention ────────────────────────────────────────────────
86
87struct MultiHeadAttention {
88    query: Linear,
89    key: Linear,
90    value: Linear,
91    out: Linear,
92    n_head: usize,
93    /// Cross-attention cache (encoder output, computed once)
94    cross_kv_cache: Option<(Tensor, Tensor)>,
95    /// Self-attention cache (accumulated K/V across decode steps)
96    self_kv_cache: Option<(Tensor, Tensor)>,
97}
98
99impl MultiHeadAttention {
100    fn load(n_state: usize, n_head: usize, vb: VarBuilder) -> candle_core::Result<Self> {
101        let query = Linear::load(n_state, n_state, vb.pp("q_proj"))?;
102        let value = Linear::load(n_state, n_state, vb.pp("v_proj"))?;
103        let key = Linear::load_no_bias(n_state, n_state, vb.pp("k_proj"))?;
104        let out = Linear::load(n_state, n_state, vb.pp("out_proj"))?;
105        Ok(Self {
106            query,
107            key,
108            value,
109            out,
110            n_head,
111            cross_kv_cache: None,
112            self_kv_cache: None,
113        })
114    }
115
116    fn forward(
117        &mut self,
118        x: &Tensor,
119        xa: Option<&Tensor>,
120        mask: Option<&Tensor>,
121        flush_cache: bool,
122    ) -> candle_core::Result<Tensor> {
123        let q = self.query.forward(x)?;
124        let (k, v) = match xa {
125            // Self-attention: cache and accumulate K/V
126            None => {
127                let new_k = self.key.forward(x)?;
128                let new_v = self.value.forward(x)?;
129                let (k, v) = if let Some((prev_k, prev_v)) = &self.self_kv_cache {
130                    // Cat along sequence dimension (dim 1)
131                    (
132                        Tensor::cat(&[prev_k, &new_k], 1)?,
133                        Tensor::cat(&[prev_v, &new_v], 1)?,
134                    )
135                } else {
136                    (new_k, new_v)
137                };
138                self.self_kv_cache = Some((k.clone(), v.clone()));
139                (k, v)
140            }
141            // Cross-attention: compute once, cache forever
142            Some(xa_t) => {
143                if flush_cache {
144                    self.cross_kv_cache = None;
145                }
146                if let Some((k, v)) = &self.cross_kv_cache {
147                    (k.clone(), v.clone())
148                } else {
149                    let k = self.key.forward(xa_t)?;
150                    let v = self.value.forward(xa_t)?;
151                    self.cross_kv_cache = Some((k.clone(), v.clone()));
152                    (k, v)
153                }
154            }
155        };
156        let wv = self.qkv_attention(&q, &k, &v, mask)?;
157        self.out.forward(&wv)
158    }
159
160    fn reshape_head(&self, x: &Tensor) -> candle_core::Result<Tensor> {
161        let (b, t, c) = x.dims3()?;
162        x.reshape((b, t, self.n_head, c / self.n_head))?
163            .transpose(1, 2)
164    }
165
166    fn qkv_attention(
167        &self,
168        q: &Tensor,
169        k: &Tensor,
170        v: &Tensor,
171        mask: Option<&Tensor>,
172    ) -> candle_core::Result<Tensor> {
173        let (_, q_len, n_state) = q.dims3()?;
174        let kv_len = k.dim(1)?;
175        let scale = ((n_state / self.n_head) as f64).powf(-0.25);
176        let q = (self.reshape_head(q)? * scale)?;
177        let k = (self.reshape_head(k)?.transpose(2, 3)? * scale)?;
178        let v = self.reshape_head(v)?.contiguous()?;
179        let mut qk = q.matmul(&k)?;
180        if let Some(mask) = mask {
181            // With KV cache: q_len=1, kv_len=accumulated. Take the row
182            // corresponding to the current position in the causal mask.
183            let q_start = kv_len - q_len;
184            let mask = mask.i((q_start..kv_len, 0..kv_len))?;
185            qk = qk.broadcast_add(&mask)?;
186        }
187        let w = softmax_last_dim(&qk)?;
188        w.matmul(&v)?.transpose(1, 2)?.flatten_from(2)
189    }
190
191    fn reset_kv_cache(&mut self) {
192        self.cross_kv_cache = None;
193        self.self_kv_cache = None;
194    }
195}
196
197// ── Residual Attention Block ────────────────────────────────────────────
198
199struct ResidualAttentionBlock {
200    attn: MultiHeadAttention,
201    attn_ln: LayerNorm,
202    cross_attn: Option<(MultiHeadAttention, LayerNorm)>,
203    mlp_linear1: Linear,
204    mlp_linear2: Linear,
205    mlp_ln: LayerNorm,
206}
207
208impl ResidualAttentionBlock {
209    fn load(
210        n_state: usize,
211        n_head: usize,
212        cross_attn: bool,
213        vb: VarBuilder,
214    ) -> candle_core::Result<Self> {
215        let attn = MultiHeadAttention::load(n_state, n_head, vb.pp("self_attn"))?;
216        let attn_ln = LayerNorm::load(n_state, 1e-5, vb.pp("self_attn_layer_norm"))?;
217        let ca = if cross_attn {
218            let ca_attn = MultiHeadAttention::load(n_state, n_head, vb.pp("encoder_attn"))?;
219            let ca_ln = LayerNorm::load(n_state, 1e-5, vb.pp("encoder_attn_layer_norm"))?;
220            Some((ca_attn, ca_ln))
221        } else {
222            None
223        };
224        let n_mlp = n_state * 4;
225        let mlp_linear1 = Linear::load(n_state, n_mlp, vb.pp("fc1"))?;
226        let mlp_linear2 = Linear::load(n_mlp, n_state, vb.pp("fc2"))?;
227        let mlp_ln = LayerNorm::load(n_state, 1e-5, vb.pp("final_layer_norm"))?;
228        Ok(Self {
229            attn,
230            attn_ln,
231            cross_attn: ca,
232            mlp_linear1,
233            mlp_linear2,
234            mlp_ln,
235        })
236    }
237
238    fn forward(
239        &mut self,
240        x: &Tensor,
241        xa: Option<&Tensor>,
242        mask: Option<&Tensor>,
243        flush_kv: bool,
244    ) -> candle_core::Result<Tensor> {
245        let a = self
246            .attn
247            .forward(&self.attn_ln.forward(x)?, None, mask, flush_kv)?;
248        let mut x = (x + a)?;
249        if let Some((ref mut ca, ref ln)) = self.cross_attn {
250            x = (&x + ca.forward(&ln.forward(&x)?, xa, None, flush_kv)?)?;
251        }
252        let mlp = self.mlp_linear2.forward(
253            &self
254                .mlp_linear1
255                .forward(&self.mlp_ln.forward(&x)?)?
256                .gelu()?,
257        )?;
258        x + mlp
259    }
260
261    fn reset_kv_cache(&mut self) {
262        self.attn.reset_kv_cache();
263        if let Some((ref mut ca, _)) = self.cross_attn {
264            ca.reset_kv_cache();
265        }
266    }
267}
268
269// ── Sinusoidal positional encoding ──────────────────────────────────────
270
271fn sinusoids(length: usize, channels: usize, device: &CandleDevice) -> candle_core::Result<Tensor> {
272    let max_timescale = 10000f32;
273    let log_inc = max_timescale.ln() / (channels / 2 - 1) as f32;
274    let inv: Vec<f32> = (0..channels / 2)
275        .map(|i| (i as f32 * (-log_inc)).exp())
276        .collect();
277    let inv = Tensor::new(inv.as_slice(), device)?.unsqueeze(0)?;
278    let arange = Tensor::arange(0, length as u32, device)?
279        .to_dtype(DType::F32)?
280        .unsqueeze(1)?;
281    let sh = (length, channels / 2);
282    let scaled = (arange.broadcast_as(sh)? * inv.broadcast_as(sh)?)?;
283    Tensor::cat(&[scaled.sin()?, scaled.cos()?], 1)
284}
285
286// ── Audio Encoder ───────────────────────────────────────────────────────
287
288struct AudioEncoder {
289    conv1: Conv1d,
290    conv2: Conv1d,
291    positional_embedding: Tensor,
292    blocks: Vec<ResidualAttentionBlock>,
293    ln_post: LayerNorm,
294}
295
296impl AudioEncoder {
297    fn load(vb: VarBuilder, cfg: &Config) -> candle_core::Result<Self> {
298        let n = cfg.d_model;
299        let h = cfg.encoder_attention_heads;
300        let cfg1 = Conv1dConfig {
301            padding: 1,
302            stride: 1,
303            groups: 1,
304            dilation: 1,
305            cudnn_fwd_algo: None,
306        };
307        let cfg2 = Conv1dConfig {
308            padding: 1,
309            stride: 2,
310            groups: 1,
311            dilation: 1,
312            cudnn_fwd_algo: None,
313        };
314        let conv1 = {
315            let w = vb.pp("conv1").get((n, cfg.num_mel_bins, 3), "weight")?;
316            let b = vb.pp("conv1").get(n, "bias")?;
317            Conv1d::new(w, Some(b), cfg1)
318        };
319        let conv2 = {
320            let w = vb.pp("conv2").get((n, n, 3), "weight")?;
321            let b = vb.pp("conv2").get(n, "bias")?;
322            Conv1d::new(w, Some(b), cfg2)
323        };
324        let pe = sinusoids(cfg.max_source_positions, n, vb.device())?;
325        let blocks = (0..cfg.encoder_layers)
326            .map(|i| ResidualAttentionBlock::load(n, h, false, vb.pp(format!("layers.{i}"))))
327            .collect::<candle_core::Result<Vec<_>>>()?;
328        let ln_post = LayerNorm::load(n, 1e-5, vb.pp("layer_norm"))?;
329        Ok(Self {
330            conv1,
331            conv2,
332            positional_embedding: pe,
333            blocks,
334            ln_post,
335        })
336    }
337
338    fn forward(&mut self, x: &Tensor, flush: bool) -> candle_core::Result<Tensor> {
339        let x = self.conv1.forward(x)?.gelu()?;
340        let x = self.conv2.forward(&x)?.gelu()?;
341        let x = x.transpose(1, 2)?;
342        let (_, seq_len, _) = x.dims3()?;
343        let pe = self.positional_embedding.narrow(0, 0, seq_len)?;
344        let mut x = x.broadcast_add(&pe)?;
345        for block in &mut self.blocks {
346            x = block.forward(&x, None, None, flush)?;
347        }
348        self.ln_post.forward(&x)
349    }
350}
351
352// ── Text Decoder ────────────────────────────────────────────────────────
353
354struct TextDecoder {
355    token_embedding: Tensor,      // (vocab, d_model)
356    positional_embedding: Tensor, // (max_target_positions, d_model)
357    blocks: Vec<ResidualAttentionBlock>,
358    ln: LayerNorm,
359    mask: Tensor,
360    /// Tracks number of tokens processed so far (for positional embedding offset)
361    tokens_seen: usize,
362}
363
364impl TextDecoder {
365    fn load(vb: VarBuilder, cfg: &Config) -> candle_core::Result<Self> {
366        let n = cfg.d_model;
367        let h = cfg.decoder_attention_heads;
368        let ctx = cfg.max_target_positions;
369        let token_embedding = vb.get((cfg.vocab_size, n), "embed_tokens.weight")?;
370        let positional_embedding = vb.get((ctx, n), "embed_positions.weight")?;
371        let blocks = (0..cfg.decoder_layers)
372            .map(|i| ResidualAttentionBlock::load(n, h, true, vb.pp(format!("layers.{i}"))))
373            .collect::<candle_core::Result<Vec<_>>>()?;
374        let ln = LayerNorm::load(n, 1e-5, vb.pp("layer_norm"))?;
375        let mask_data: Vec<f32> = (0..ctx)
376            .flat_map(|i| (0..ctx).map(move |j| if j > i { f32::NEG_INFINITY } else { 0.0 }))
377            .collect();
378        let mask = Tensor::from_vec(mask_data, (ctx, ctx), vb.device())?;
379        Ok(Self {
380            token_embedding,
381            positional_embedding,
382            blocks,
383            ln,
384            mask,
385            tokens_seen: 0,
386        })
387    }
388
389    fn forward(
390        &mut self,
391        tokens: &Tensor,
392        xa: &Tensor,
393        flush: bool,
394    ) -> candle_core::Result<Tensor> {
395        let seq_len = tokens.dim(D::Minus1)?;
396        // Embedding lookup: tokens [batch, seq] → embeddings [batch, seq, d_model]
397        let flat_tokens = tokens.flatten_all()?;
398        let te = self.token_embedding.index_select(&flat_tokens, 0)?;
399        let te = te.reshape((tokens.dim(0)?, seq_len, self.token_embedding.dim(1)?))?;
400        // Use positional embedding at the correct offset (not always 0)
401        let pe = self
402            .positional_embedding
403            .narrow(0, self.tokens_seen, seq_len)?;
404        self.tokens_seen += seq_len;
405        let mut x = te.broadcast_add(&pe)?;
406        for block in &mut self.blocks {
407            x = block.forward(&x, Some(xa), Some(&self.mask), flush)?;
408        }
409        self.ln.forward(&x)
410    }
411
412    fn final_linear(&self, x: &Tensor) -> candle_core::Result<Tensor> {
413        let b = x.dim(0)?;
414        let w = self.token_embedding.broadcast_left(b)?;
415        x.matmul(&w.t()?)
416    }
417
418    fn reset_kv_cache(&mut self) {
419        self.tokens_seen = 0;
420        for block in &mut self.blocks {
421            block.reset_kv_cache();
422        }
423    }
424}
425
426// ── Top-level Whisper wrapper ───────────────────────────────────────────
427
428pub struct WhisperModelWrapper {
429    encoder: Mutex<AudioEncoder>,
430    decoder: Mutex<TextDecoder>,
431    config: Config,
432    mel_filters: Vec<f32>,
433    device: CandleDevice,
434    #[allow(dead_code)]
435    dtype: DType,
436}
437
438impl WhisperModelWrapper {
439    /// Load from VarBuilder + config.
440    pub fn new(
441        vb: VarBuilder,
442        config: Config,
443        mel_filters: Vec<f32>,
444        device: CandleDevice,
445        dtype: DType,
446    ) -> Result<Self> {
447        info!(
448            "Loading Whisper (d_model={}, encoder_layers={}, decoder_layers={})",
449            config.d_model, config.encoder_layers, config.decoder_layers
450        );
451        let enc_vb = vb.pp("model.encoder");
452        let dec_vb = vb.pp("model.decoder");
453        let encoder = AudioEncoder::load(enc_vb, &config)
454            .map_err(|e| FerrumError::model(format!("encoder load: {e}")))?;
455        let decoder = TextDecoder::load(dec_vb, &config)
456            .map_err(|e| FerrumError::model(format!("decoder load: {e}")))?;
457        Ok(Self {
458            encoder: Mutex::new(encoder),
459            decoder: Mutex::new(decoder),
460            config,
461            mel_filters,
462            device,
463            dtype,
464        })
465    }
466
467    /// Load from model directory.
468    pub fn from_model_dir(
469        model_dir: &std::path::Path,
470        device: CandleDevice,
471        dtype: DType,
472    ) -> Result<Self> {
473        let config_path = model_dir.join("config.json");
474        let config: Config = serde_json::from_str(
475            &std::fs::read_to_string(&config_path)
476                .map_err(|e| FerrumError::model(format!("read config: {e}")))?,
477        )
478        .map_err(|e| FerrumError::model(format!("parse config: {e}")))?;
479
480        let mel_bytes = match config.num_mel_bins {
481            128 => include_bytes!("mel_filters128.bin").as_slice(),
482            _ => include_bytes!("mel_filters80.bin").as_slice(),
483        };
484        let mut mel_filters = vec![0f32; mel_bytes.len() / 4];
485        for (i, chunk) in mel_bytes.chunks_exact(4).enumerate() {
486            mel_filters[i] = f32::from_le_bytes([chunk[0], chunk[1], chunk[2], chunk[3]]);
487        }
488
489        let safetensors: Vec<_> = std::fs::read_dir(model_dir)
490            .map_err(|e| FerrumError::model(format!("read dir: {e}")))?
491            .filter_map(|e| e.ok())
492            .map(|e| e.path())
493            .filter(|p| p.extension().map_or(false, |ext| ext == "safetensors"))
494            .collect();
495
496        if safetensors.is_empty() {
497            return Err(FerrumError::model("No safetensors files found"));
498        }
499
500        let vb = unsafe {
501            VarBuilder::from_mmaped_safetensors(&safetensors, dtype, &device)
502                .map_err(|e| FerrumError::model(format!("load weights: {e}")))?
503        };
504
505        Self::new(vb, config, mel_filters, device, dtype)
506    }
507
508    /// PCM → mel spectrogram tensor (matching Python whisper exactly).
509    /// Pads or truncates PCM to exactly 30 seconds (N_SAMPLES = 480000).
510    pub fn pcm_to_mel_tensor(&self, pcm: &[f32]) -> Result<Tensor> {
511        let n_samples = whisper::N_SAMPLES;
512        let samples = if pcm.len() >= n_samples {
513            pcm[..n_samples].to_vec()
514        } else {
515            let mut buf = pcm.to_vec();
516            buf.resize(n_samples, 0.0);
517            buf
518        };
519
520        let n_mels = self.config.num_mel_bins;
521        let mel = crate::mel::log_mel_spectrogram(&samples, n_mels, &self.mel_filters);
522        let n_frames = mel.len() / n_mels;
523
524        Tensor::from_vec(mel, (1, n_mels, n_frames), &self.device)
525            .map_err(|e| FerrumError::model(format!("mel tensor: {e}")))
526    }
527
528    /// Encode a mel segment → encoder hidden states.
529    pub fn encode(&self, mel: &Tensor) -> Result<Tensor> {
530        let mut enc = self.encoder.lock();
531        enc.blocks.iter_mut().for_each(|b| b.reset_kv_cache());
532        enc.forward(mel, true)
533            .map_err(|e| FerrumError::model(format!("encode: {e}")))
534    }
535
536    /// Run one decode pass: tokens → logits (includes KV cache, final_linear).
537    /// On first call pass full initial_tokens; on subsequent calls pass only the new token.
538    pub fn decode_step(&self, tokens: &[u32], encoder_out: &Tensor) -> Result<Vec<f32>> {
539        let mut dec = self.decoder.lock();
540        let t = Tensor::new(tokens, &self.device)
541            .and_then(|t| t.unsqueeze(0))
542            .map_err(|e| FerrumError::model(format!("token tensor: {e}")))?;
543        let hidden = dec
544            .forward(&t, encoder_out, false)
545            .map_err(|e| FerrumError::model(format!("decode: {e}")))?;
546        let last_pos = hidden
547            .dim(1)
548            .map_err(|e| FerrumError::model(format!("dim: {e}")))?
549            - 1;
550        let last_hidden = hidden
551            .i((.., last_pos..last_pos + 1))
552            .map_err(|e| FerrumError::model(format!("slice: {e}")))?;
553        let logits = dec
554            .final_linear(&last_hidden)
555            .map_err(|e| FerrumError::model(format!("final_linear: {e}")))?;
556        logits
557            .squeeze(0)
558            .and_then(|t| t.squeeze(0))
559            .and_then(|t| t.to_dtype(DType::F32))
560            .and_then(|t| t.to_vec1::<f32>())
561            .map_err(|e| FerrumError::model(format!("logits to vec: {e}")))
562    }
563
564    /// Reset decoder KV cache (call between segments).
565    pub fn reset_decoder(&self) {
566        self.decoder.lock().reset_kv_cache();
567    }
568
569    pub fn config(&self) -> &Config {
570        &self.config
571    }
572
573    pub fn device(&self) -> &CandleDevice {
574        &self.device
575    }
576}