1use 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
13fn 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
23struct 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
50struct 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 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
85struct MultiHeadAttention {
88 query: Linear,
89 key: Linear,
90 value: Linear,
91 out: Linear,
92 n_head: usize,
93 cross_kv_cache: Option<(Tensor, Tensor)>,
95 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 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 (
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 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 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
197struct 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
269fn 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
286struct 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
352struct TextDecoder {
355 token_embedding: Tensor, positional_embedding: Tensor, blocks: Vec<ResidualAttentionBlock>,
358 ln: LayerNorm,
359 mask: Tensor,
360 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 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 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
426pub 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 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 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 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 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 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 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}