1use crate::model_config;
9use super::traits::*;
10use anyhow::Result;
11use serde::{Serialize, Deserialize};
12
13model_config!(WhisperConfig {
15 vocab_size: usize = 51865,
16 d_model: usize = 512,
17 encoder_layers: usize = 6,
18 decoder_layers: usize = 6,
19 encoder_attention_heads: usize = 8,
20 decoder_attention_heads: usize = 8,
21 encoder_ffn_dim: usize = 2048,
22 decoder_ffn_dim: usize = 2048,
23 dropout: f32 = 0.0,
24 attention_dropout: f32 = 0.0,
25 activation_dropout: f32 = 0.0,
26 activation_function: String = "gelu".to_string(),
27 init_std: f32 = 0.02,
28 layer_norm_eps: f32 = 1e-5,
29 scale_embedding: bool = false,
30 use_cache: bool = true,
31 is_encoder_decoder: bool = true,
32 pad_token_id: i64 = 50257,
33 bos_token_id: i64 = 50258,
34 eos_token_id: i64 = 50257,
35 decoder_start_token_id: i64 = 50258,
36 max_source_positions: usize = 1500,
38 max_target_positions: usize = 448,
39 num_mel_bins: usize = 80,
40 num_hidden_layers: usize = 6,
42 hidden_size: usize = 512,
43});
44
45impl WhisperConfig {
46 pub fn from_gguf_config(gguf: &crate::weight_loader_core::GGUFModelConfig) -> Self {
48 Self {
51 vocab_size: gguf.vocab_size,
52 d_model: gguf.hidden_size,
53 hidden_size: gguf.hidden_size,
54 encoder_layers: gguf.num_hidden_layers / 2, decoder_layers: gguf.num_hidden_layers / 2,
56 num_hidden_layers: gguf.num_hidden_layers,
57 encoder_attention_heads: gguf.num_attention_heads,
58 decoder_attention_heads: gguf.num_attention_heads,
59 encoder_ffn_dim: gguf.intermediate_size,
60 decoder_ffn_dim: gguf.intermediate_size,
61 layer_norm_eps: gguf.rms_norm_eps,
62 ..Default::default()
63 }
64 }
65
66 pub fn encoder_head_dim(&self) -> usize {
68 self.d_model / self.encoder_attention_heads
69 }
70
71 pub fn decoder_head_dim(&self) -> usize {
73 self.d_model / self.decoder_attention_heads
74 }
75}
76
77pub struct WhisperModelV2 {
79 config: WhisperConfig,
80 device: Device,
81 encoder: WhisperEncoder,
82 decoder: WhisperDecoder,
83 proj_out: Tensor, }
85
86pub struct WhisperEncoder {
88 conv1_weight: Tensor, conv1_bias: Tensor, conv2_weight: Tensor, conv2_bias: Tensor, embed_positions: Tensor, layers: Vec<WhisperEncoderLayer>,
97 layer_norm: Tensor,
99 layer_norm_bias: Option<Tensor>,
100 config: WhisperConfig,
101}
102
103pub struct WhisperDecoder {
105 embed_tokens: Tensor, embed_positions: Tensor, layers: Vec<WhisperDecoderLayer>,
111 layer_norm: Tensor,
113 layer_norm_bias: Option<Tensor>,
114 config: WhisperConfig,
115}
116
117pub struct WhisperEncoderLayer {
119 self_attn: WhisperAttention,
120 self_attn_layer_norm: Tensor,
121 self_attn_layer_norm_bias: Option<Tensor>,
122 fc1: Tensor,
123 fc1_bias: Tensor,
124 fc2: Tensor,
125 fc2_bias: Tensor,
126 final_layer_norm: Tensor,
127 final_layer_norm_bias: Option<Tensor>,
128 config: WhisperConfig,
129}
130
131pub struct WhisperDecoderLayer {
133 self_attn: WhisperAttention,
134 self_attn_layer_norm: Tensor,
135 self_attn_layer_norm_bias: Option<Tensor>,
136 encoder_attn: WhisperAttention,
137 encoder_attn_layer_norm: Tensor,
138 encoder_attn_layer_norm_bias: Option<Tensor>,
139 fc1: Tensor,
140 fc1_bias: Tensor,
141 fc2: Tensor,
142 fc2_bias: Tensor,
143 final_layer_norm: Tensor,
144 final_layer_norm_bias: Option<Tensor>,
145 config: WhisperConfig,
146}
147
148pub struct WhisperAttention {
150 k_proj: Tensor,
151 k_proj_bias: Option<Tensor>,
152 v_proj: Tensor,
153 v_proj_bias: Option<Tensor>,
154 q_proj: Tensor,
155 q_proj_bias: Option<Tensor>,
156 out_proj: Tensor,
157 out_proj_bias: Option<Tensor>,
158 num_heads: usize,
159 head_dim: usize,
160 scale: f32,
161 is_causal: bool, }
163
164impl Model for WhisperModelV2 {
165 type Config = WhisperConfig;
166
167 fn new(config: WhisperConfig) -> Result<Self> {
168 let device = Device::CPU;
169 let encoder = WhisperEncoder::new(&config, &device)?;
170 let decoder = WhisperDecoder::new(&config, &device)?;
171
172 let proj_out = ops_fn::zeros(&[config.d_model, config.vocab_size], DataType::Float32, &device)?;
174
175 Ok(Self { config, device, encoder, decoder, proj_out })
176 }
177
178 fn from_weights(config: WhisperConfig, weights: ModelWeights) -> Result<Self> {
179 let mut model = Self::new(config)?;
180 model.encoder.load_weights(&weights)?;
181 model.decoder.load_weights(&weights)?;
182
183 if let Some(w) = weights.get("proj_out.weight") {
185 model.proj_out = ops_fn::transpose(w)?;
186 } else if let Some(w) = weights.get("model.decoder.embed_tokens.weight") {
187 model.proj_out = ops_fn::transpose(w)?;
189 }
190
191 Ok(model)
192 }
193
194 fn forward(&self, inputs: &ModelInputs) -> Result<ModelOutputs> {
195 match inputs {
196 ModelInputs::Audio { input_features, attention_mask } => {
197 let encoder_outputs = self.encoder.forward(input_features)?;
199
200 let start_token = self.config.decoder_start_token_id;
203 let batch_size = input_features.shape()[0];
204
205 let decoder_input_ids: Vec<i64> = vec![start_token; batch_size];
207 let decoder_input = Tensor::from_i64_slice(
208 &decoder_input_ids,
209 &[batch_size, 1],
210 &self.device
211 )?;
212
213 let decoder_outputs = self.decoder.forward(&decoder_input, Some(&encoder_outputs))?;
215
216 let logits = ops_fn::matmul(&decoder_outputs, &self.proj_out)?;
218
219 Ok(ModelOutputs::Sequence {
220 logits,
221 encoder_hidden_states: Some(encoder_outputs),
222 decoder_hidden_states: Some(decoder_outputs),
223 })
224 },
225 _ => Err(anyhow::anyhow!("Whisper expects Audio input")),
226 }
227 }
228
229 fn generate(&self, _prompt: &str, config: &GenerationConfig) -> Result<String> {
230 Ok(format!("[Whisper: Use transcribe() method with audio input. Max tokens: {}]",
240 config.max_new_tokens))
241 }
242
243 fn config(&self) -> &Self::Config { &self.config }
244
245 fn memory_requirements(&self) -> MemoryRequirements {
246 let d_model = self.config.d_model;
247 let enc_layers = self.config.encoder_layers;
248 let dec_layers = self.config.decoder_layers;
249 let enc_ffn = self.config.encoder_ffn_dim;
250 let dec_ffn = self.config.decoder_ffn_dim;
251
252 let encoder_params = enc_layers * (4 * d_model * d_model + 2 * d_model * enc_ffn);
254 let decoder_params = dec_layers * (8 * d_model * d_model + 2 * d_model * dec_ffn); let embedding_params = self.config.vocab_size * d_model;
256 let conv_params = self.config.num_mel_bins * d_model * 3 + d_model * d_model * 3;
257
258 let total_params = encoder_params + decoder_params + embedding_params + conv_params;
259 let param_bytes = total_params * 4; let kv_cache_bytes = (self.config.max_source_positions + self.config.max_target_positions)
262 * d_model * 2 * 4;
263
264 MemoryRequirements {
265 gpu_memory: param_bytes,
266 cpu_memory: param_bytes / 4,
267 kv_cache_memory: kv_cache_bytes,
268 peak_memory: param_bytes + kv_cache_bytes,
269 }
270 }
271
272 fn to_device(&mut self, device: &Device) -> Result<()> {
273 self.device = device.clone();
274 self.encoder.to_device(device)?;
275 self.decoder.to_device(device)?;
276 self.proj_out = self.proj_out.to_device(device)?;
277 Ok(())
278 }
279}
280
281impl WhisperModelV2 {
282 pub fn transcribe(&self, mel_spectrogram: &Tensor, config: &GenerationConfig) -> Result<Vec<u32>> {
285 let encoder_outputs = self.encoder.forward(mel_spectrogram)?;
287
288 let batch_size = mel_spectrogram.shape()[0];
289
290 let mut tokens: Vec<u32> = vec![self.config.decoder_start_token_id as u32];
292
293 for _ in 0..config.max_new_tokens {
295 let tokens_i64: Vec<i64> = tokens.iter().map(|&t| t as i64).collect();
297 let decoder_input = Tensor::from_i64_slice(
298 &tokens_i64,
299 &[batch_size, tokens.len()],
300 &self.device
301 )?;
302
303 let decoder_outputs = self.decoder.forward(&decoder_input, Some(&encoder_outputs))?;
305
306 let logits = ops_fn::matmul(&decoder_outputs, &self.proj_out)?;
308
309 let logits_candle = logits.to_candle()?;
311 let shape = logits_candle.dims();
312 let seq_len = shape[1];
313
314 let last_logits = logits_candle
315 .narrow(1, seq_len - 1, 1)?
316 .squeeze(1)?
317 .squeeze(0)?;
318
319 let logits_vec: Vec<f32> = last_logits.to_vec1()?;
320
321 let next_token = {
323 let mut max_idx = 0;
324 let mut max_val = logits_vec[0];
325 for (idx, &val) in logits_vec.iter().enumerate() {
326 if val > max_val {
328 max_val = val;
329 max_idx = idx;
330 }
331 }
332 max_idx as u32
333 };
334
335 if next_token == config.eos_token_id {
337 break;
338 }
339
340 tokens.push(next_token);
341 }
342
343 Ok(tokens)
344 }
345}
346
347impl WhisperEncoder {
348 fn new(config: &WhisperConfig, device: &Device) -> Result<Self> {
349 let mut layers = Vec::new();
350 for _ in 0..config.encoder_layers {
351 layers.push(WhisperEncoderLayer::new(config, device, false)?); }
353
354 let conv1_weight = ops_fn::zeros(&[config.d_model, config.num_mel_bins, 3], DataType::Float32, device)?;
358 let conv1_bias = ops_fn::zeros(&[config.d_model], DataType::Float32, device)?;
359 let conv2_weight = ops_fn::zeros(&[config.d_model, config.d_model, 3], DataType::Float32, device)?;
360 let conv2_bias = ops_fn::zeros(&[config.d_model], DataType::Float32, device)?;
361
362 let embed_positions = create_sinusoidal_embeddings(config.max_source_positions, config.d_model, device)?;
364
365 Ok(Self {
366 conv1_weight,
367 conv1_bias,
368 conv2_weight,
369 conv2_bias,
370 embed_positions,
371 layers,
372 layer_norm: ops_fn::zeros(&[config.d_model], DataType::Float32, device)?,
373 layer_norm_bias: None,
374 config: config.clone(),
375 })
376 }
377
378 fn forward(&self, mel_spectrogram: &Tensor) -> Result<Tensor> {
379 let mut hidden_states = self.apply_conv1d(mel_spectrogram)?;
382
383 let hidden_candle = hidden_states.to_candle()?;
388 let transposed = hidden_candle.transpose(1, 2)?;
389 hidden_states = Tensor::from_candle(transposed);
390
391 let seq_len = hidden_states.shape()[1];
393 let pos_emb = self.get_position_embeddings(seq_len)?;
394 hidden_states = ops_fn::add(&hidden_states, &pos_emb)?;
395
396 for layer in &self.layers {
398 hidden_states = layer.forward(&hidden_states, false)?; }
400
401 let result = ops_fn::layer_norm(&hidden_states, &self.layer_norm, self.layer_norm_bias.as_ref(), self.config.layer_norm_eps)?;
403
404 Ok(result)
405 }
406
407 fn apply_conv1d(&self, mel_spectrogram: &Tensor) -> Result<Tensor> {
409 let input = mel_spectrogram.to_candle()?;
415 let shape = input.dims();
416 let (batch_size, n_mels, n_frames) = (shape[0], shape[1], shape[2]);
417 let d_model = self.config.d_model;
418
419 let conv1_out = self.conv1d_forward(&input, &self.conv1_weight, &self.conv1_bias, 3, 1, 1)?;
425 let conv1_activated = conv1_out.gelu()?;
426
427 let conv2_out = self.conv1d_forward(&conv1_activated, &self.conv2_weight, &self.conv2_bias, 3, 2, 1)?;
430 let conv2_activated = conv2_out.gelu()?;
431
432 Ok(Tensor::from_candle(conv2_activated))
433 }
434
435 fn conv1d_forward(
438 &self,
439 input: &candle_core::Tensor,
440 weight: &Tensor,
441 bias: &Tensor,
442 kernel_size: usize,
443 stride: usize,
444 padding: usize,
445 ) -> Result<candle_core::Tensor> {
446 let weight_candle = weight.to_candle()?;
447 let bias_candle = bias.to_candle()?;
448
449 let shape = input.dims();
450 let (batch_size, in_channels, in_length) = (shape[0], shape[1], shape[2]);
451 let out_channels = weight_candle.dims()[0];
452
453 let out_length = (in_length + 2 * padding - kernel_size) / stride + 1;
455
456 let padded = if padding > 0 {
458 let zeros_shape = &[batch_size, in_channels, padding];
460 let zero_pad = candle_core::Tensor::zeros(zeros_shape, input.dtype(), input.device())?;
461 candle_core::Tensor::cat(&[&zero_pad, input, &zero_pad], 2)?
462 } else {
463 input.clone()
464 };
465
466 let mut output_slices = Vec::new();
469
470 for i in 0..out_length {
471 let start = i * stride;
472 let patch = padded.narrow(2, start, kernel_size)?; let patch_flat = patch.reshape(&[batch_size, in_channels * kernel_size])?;
476
477 let weight_flat = weight_candle.reshape(&[out_channels, in_channels * kernel_size])?;
479
480 let weight_t = weight_flat.t()?;
482 let out_pos = patch_flat.matmul(&weight_t)?;
483
484 output_slices.push(out_pos.unsqueeze(2)?); }
486
487 let refs: Vec<&candle_core::Tensor> = output_slices.iter().collect();
489 let output = candle_core::Tensor::cat(&refs, 2)?; let bias_expanded = bias_candle.unsqueeze(0)?.unsqueeze(2)?;
493 let output_with_bias = output.broadcast_add(&bias_expanded)?;
494
495 Ok(output_with_bias)
496 }
497
498 fn get_position_embeddings(&self, seq_len: usize) -> Result<Tensor> {
500 let emb = self.embed_positions.to_candle()?;
501 let sliced = emb.narrow(0, 0, seq_len)?;
502 let expanded = sliced.unsqueeze(0)?; Ok(Tensor::from_candle(expanded))
504 }
505
506 fn load_weights(&mut self, weights: &ModelWeights) -> Result<()> {
507 if let Some(w) = weights.get("model.encoder.conv1.weight") {
509 self.conv1_weight = w.clone();
510 }
511 if let Some(w) = weights.get("model.encoder.conv1.bias") {
512 self.conv1_bias = w.clone();
513 }
514 if let Some(w) = weights.get("model.encoder.conv2.weight") {
515 self.conv2_weight = w.clone();
516 }
517 if let Some(w) = weights.get("model.encoder.conv2.bias") {
518 self.conv2_bias = w.clone();
519 }
520
521 if let Some(w) = weights.get("model.encoder.embed_positions.weight") {
523 self.embed_positions = w.clone();
524 }
525
526 if let Some(w) = weights.get("model.encoder.layer_norm.weight") {
528 self.layer_norm = w.clone();
529 }
530 if let Some(w) = weights.get("model.encoder.layer_norm.bias") {
531 self.layer_norm_bias = Some(w.clone());
532 }
533
534 for (i, layer) in self.layers.iter_mut().enumerate() {
536 layer.load_weights(weights, i)?;
537 }
538
539 Ok(())
540 }
541
542 fn to_device(&mut self, device: &Device) -> Result<()> {
543 self.conv1_weight = self.conv1_weight.to_device(device)?;
544 self.conv1_bias = self.conv1_bias.to_device(device)?;
545 self.conv2_weight = self.conv2_weight.to_device(device)?;
546 self.conv2_bias = self.conv2_bias.to_device(device)?;
547 self.embed_positions = self.embed_positions.to_device(device)?;
548 self.layer_norm = self.layer_norm.to_device(device)?;
549 if let Some(ref mut b) = self.layer_norm_bias {
550 *b = b.to_device(device)?;
551 }
552 for layer in &mut self.layers {
553 layer.to_device(device)?;
554 }
555 Ok(())
556 }
557}
558
559impl WhisperDecoder {
560 fn new(config: &WhisperConfig, device: &Device) -> Result<Self> {
561 let mut layers = Vec::new();
562 for _ in 0..config.decoder_layers {
563 layers.push(WhisperDecoderLayer::new(config, device)?);
564 }
565
566 let embed_tokens = ops_fn::zeros(&[config.vocab_size, config.d_model], DataType::Float32, device)?;
568
569 let embed_positions = ops_fn::zeros(&[config.max_target_positions, config.d_model], DataType::Float32, device)?;
571
572 Ok(Self {
573 embed_tokens,
574 embed_positions,
575 layers,
576 layer_norm: ops_fn::zeros(&[config.d_model], DataType::Float32, device)?,
577 layer_norm_bias: None,
578 config: config.clone(),
579 })
580 }
581
582 fn forward(&self, input_ids: &Tensor, encoder_hidden_states: Option<&Tensor>) -> Result<Tensor> {
583 let mut hidden_states = ops_fn::embedding(input_ids, &self.embed_tokens)?;
585
586 let seq_len = input_ids.shape()[1];
588 let pos_emb = self.get_position_embeddings(seq_len)?;
589 hidden_states = ops_fn::add(&hidden_states, &pos_emb)?;
590
591 for layer in &self.layers {
593 hidden_states = layer.forward(&hidden_states, encoder_hidden_states)?;
594 }
595
596 let result = ops_fn::layer_norm(&hidden_states, &self.layer_norm, self.layer_norm_bias.as_ref(), self.config.layer_norm_eps)?;
598
599 Ok(result)
600 }
601
602 fn get_position_embeddings(&self, seq_len: usize) -> Result<Tensor> {
604 let emb = self.embed_positions.to_candle()?;
605 let sliced = emb.narrow(0, 0, seq_len)?;
606 let expanded = sliced.unsqueeze(0)?; Ok(Tensor::from_candle(expanded))
608 }
609
610 fn load_weights(&mut self, weights: &ModelWeights) -> Result<()> {
611 if let Some(w) = weights.get("model.decoder.embed_tokens.weight") {
613 self.embed_tokens = w.clone();
614 }
615 if let Some(w) = weights.get("model.decoder.embed_positions.weight") {
616 self.embed_positions = w.clone();
617 }
618
619 if let Some(w) = weights.get("model.decoder.layer_norm.weight") {
621 self.layer_norm = w.clone();
622 }
623 if let Some(w) = weights.get("model.decoder.layer_norm.bias") {
624 self.layer_norm_bias = Some(w.clone());
625 }
626
627 for (i, layer) in self.layers.iter_mut().enumerate() {
629 layer.load_weights(weights, i)?;
630 }
631
632 Ok(())
633 }
634
635 fn to_device(&mut self, device: &Device) -> Result<()> {
636 self.embed_tokens = self.embed_tokens.to_device(device)?;
637 self.embed_positions = self.embed_positions.to_device(device)?;
638 self.layer_norm = self.layer_norm.to_device(device)?;
639 if let Some(ref mut b) = self.layer_norm_bias {
640 *b = b.to_device(device)?;
641 }
642 for layer in &mut self.layers {
643 layer.to_device(device)?;
644 }
645 Ok(())
646 }
647}
648
649impl WhisperEncoderLayer {
650 fn new(config: &WhisperConfig, device: &Device, is_causal: bool) -> Result<Self> {
651 let head_dim = config.encoder_head_dim();
652
653 Ok(Self {
654 self_attn: WhisperAttention::new(
655 config.d_model,
656 config.encoder_attention_heads,
657 head_dim,
658 device,
659 is_causal,
660 )?,
661 self_attn_layer_norm: ops_fn::zeros(&[config.d_model], DataType::Float32, device)?,
662 self_attn_layer_norm_bias: None,
663 fc1: ops_fn::zeros(&[config.d_model, config.encoder_ffn_dim], DataType::Float32, device)?,
664 fc1_bias: ops_fn::zeros(&[config.encoder_ffn_dim], DataType::Float32, device)?,
665 fc2: ops_fn::zeros(&[config.encoder_ffn_dim, config.d_model], DataType::Float32, device)?,
666 fc2_bias: ops_fn::zeros(&[config.d_model], DataType::Float32, device)?,
667 final_layer_norm: ops_fn::zeros(&[config.d_model], DataType::Float32, device)?,
668 final_layer_norm_bias: None,
669 config: config.clone(),
670 })
671 }
672
673 fn forward(&self, hidden_states: &Tensor, is_causal: bool) -> Result<Tensor> {
674 let residual = hidden_states.clone();
677 let hidden_states = ops_fn::layer_norm(
678 hidden_states,
679 &self.self_attn_layer_norm,
680 self.self_attn_layer_norm_bias.as_ref(),
681 self.config.layer_norm_eps
682 )?;
683 let hidden_states = self.self_attn.forward(&hidden_states, None, is_causal)?;
684 let hidden_states = ops_fn::add(&residual, &hidden_states)?;
685
686 let residual = hidden_states.clone();
688 let hidden_states = ops_fn::layer_norm(
689 &hidden_states,
690 &self.final_layer_norm,
691 self.final_layer_norm_bias.as_ref(),
692 self.config.layer_norm_eps
693 )?;
694
695 let hidden_states = ops_fn::matmul(&hidden_states, &self.fc1)?;
697 let hidden_states = self.add_bias(&hidden_states, &self.fc1_bias)?;
698 let hidden_states = ops_fn::gelu(&hidden_states)?;
699 let hidden_states = ops_fn::matmul(&hidden_states, &self.fc2)?;
700 let hidden_states = self.add_bias(&hidden_states, &self.fc2_bias)?;
701
702 ops_fn::add(&residual, &hidden_states)
703 }
704
705 fn add_bias(&self, x: &Tensor, bias: &Tensor) -> Result<Tensor> {
706 ops_fn::add(x, bias)
707 }
708
709 fn load_weights(&mut self, weights: &ModelWeights, layer_idx: usize) -> Result<()> {
710 let prefix = format!("model.encoder.layers.{}", layer_idx);
711
712 if let Some(w) = weights.get(&format!("{}.self_attn_layer_norm.weight", prefix)) {
714 self.self_attn_layer_norm = w.clone();
715 }
716 if let Some(w) = weights.get(&format!("{}.self_attn_layer_norm.bias", prefix)) {
717 self.self_attn_layer_norm_bias = Some(w.clone());
718 }
719 if let Some(w) = weights.get(&format!("{}.final_layer_norm.weight", prefix)) {
720 self.final_layer_norm = w.clone();
721 }
722 if let Some(w) = weights.get(&format!("{}.final_layer_norm.bias", prefix)) {
723 self.final_layer_norm_bias = Some(w.clone());
724 }
725
726 if let Some(w) = weights.get(&format!("{}.fc1.weight", prefix)) {
728 self.fc1 = ops_fn::transpose(w)?;
729 }
730 if let Some(w) = weights.get(&format!("{}.fc1.bias", prefix)) {
731 self.fc1_bias = w.clone();
732 }
733 if let Some(w) = weights.get(&format!("{}.fc2.weight", prefix)) {
734 self.fc2 = ops_fn::transpose(w)?;
735 }
736 if let Some(w) = weights.get(&format!("{}.fc2.bias", prefix)) {
737 self.fc2_bias = w.clone();
738 }
739
740 self.self_attn.load_weights(weights, &format!("{}.self_attn", prefix))?;
742
743 Ok(())
744 }
745
746 fn to_device(&mut self, device: &Device) -> Result<()> {
747 self.self_attn_layer_norm = self.self_attn_layer_norm.to_device(device)?;
748 if let Some(ref mut b) = self.self_attn_layer_norm_bias {
749 *b = b.to_device(device)?;
750 }
751 self.final_layer_norm = self.final_layer_norm.to_device(device)?;
752 if let Some(ref mut b) = self.final_layer_norm_bias {
753 *b = b.to_device(device)?;
754 }
755 self.fc1 = self.fc1.to_device(device)?;
756 self.fc1_bias = self.fc1_bias.to_device(device)?;
757 self.fc2 = self.fc2.to_device(device)?;
758 self.fc2_bias = self.fc2_bias.to_device(device)?;
759 self.self_attn.to_device(device)?;
760 Ok(())
761 }
762}
763
764impl WhisperDecoderLayer {
765 fn new(config: &WhisperConfig, device: &Device) -> Result<Self> {
766 let head_dim = config.decoder_head_dim();
767
768 Ok(Self {
769 self_attn: WhisperAttention::new(
771 config.d_model,
772 config.decoder_attention_heads,
773 head_dim,
774 device,
775 true, )?,
777 self_attn_layer_norm: ops_fn::zeros(&[config.d_model], DataType::Float32, device)?,
778 self_attn_layer_norm_bias: None,
779 encoder_attn: WhisperAttention::new(
781 config.d_model,
782 config.decoder_attention_heads,
783 head_dim,
784 device,
785 false, )?,
787 encoder_attn_layer_norm: ops_fn::zeros(&[config.d_model], DataType::Float32, device)?,
788 encoder_attn_layer_norm_bias: None,
789 fc1: ops_fn::zeros(&[config.d_model, config.decoder_ffn_dim], DataType::Float32, device)?,
790 fc1_bias: ops_fn::zeros(&[config.decoder_ffn_dim], DataType::Float32, device)?,
791 fc2: ops_fn::zeros(&[config.decoder_ffn_dim, config.d_model], DataType::Float32, device)?,
792 fc2_bias: ops_fn::zeros(&[config.d_model], DataType::Float32, device)?,
793 final_layer_norm: ops_fn::zeros(&[config.d_model], DataType::Float32, device)?,
794 final_layer_norm_bias: None,
795 config: config.clone(),
796 })
797 }
798
799 fn forward(&self, hidden_states: &Tensor, encoder_hidden_states: Option<&Tensor>) -> Result<Tensor> {
800 let residual = hidden_states.clone();
802 let hidden_states = ops_fn::layer_norm(
803 hidden_states,
804 &self.self_attn_layer_norm,
805 self.self_attn_layer_norm_bias.as_ref(),
806 self.config.layer_norm_eps
807 )?;
808 let hidden_states = self.self_attn.forward(&hidden_states, None, true)?; let hidden_states = ops_fn::add(&residual, &hidden_states)?;
810
811 let hidden_states = if let Some(encoder_states) = encoder_hidden_states {
813 let residual = hidden_states.clone();
814 let normed = ops_fn::layer_norm(
815 &hidden_states,
816 &self.encoder_attn_layer_norm,
817 self.encoder_attn_layer_norm_bias.as_ref(),
818 self.config.layer_norm_eps
819 )?;
820 let attn_out = self.encoder_attn.forward(&normed, Some(encoder_states), false)?;
821 ops_fn::add(&residual, &attn_out)?
822 } else {
823 hidden_states
824 };
825
826 let residual = hidden_states.clone();
828 let hidden_states = ops_fn::layer_norm(
829 &hidden_states,
830 &self.final_layer_norm,
831 self.final_layer_norm_bias.as_ref(),
832 self.config.layer_norm_eps
833 )?;
834
835 let hidden_states = ops_fn::matmul(&hidden_states, &self.fc1)?;
836 let hidden_states = self.add_bias(&hidden_states, &self.fc1_bias)?;
837 let hidden_states = ops_fn::gelu(&hidden_states)?;
838 let hidden_states = ops_fn::matmul(&hidden_states, &self.fc2)?;
839 let hidden_states = self.add_bias(&hidden_states, &self.fc2_bias)?;
840
841 ops_fn::add(&residual, &hidden_states)
842 }
843
844 fn add_bias(&self, x: &Tensor, bias: &Tensor) -> Result<Tensor> {
845 ops_fn::add(x, bias)
846 }
847
848 fn load_weights(&mut self, weights: &ModelWeights, layer_idx: usize) -> Result<()> {
849 let prefix = format!("model.decoder.layers.{}", layer_idx);
850
851 if let Some(w) = weights.get(&format!("{}.self_attn_layer_norm.weight", prefix)) {
853 self.self_attn_layer_norm = w.clone();
854 }
855 if let Some(w) = weights.get(&format!("{}.self_attn_layer_norm.bias", prefix)) {
856 self.self_attn_layer_norm_bias = Some(w.clone());
857 }
858 if let Some(w) = weights.get(&format!("{}.encoder_attn_layer_norm.weight", prefix)) {
859 self.encoder_attn_layer_norm = w.clone();
860 }
861 if let Some(w) = weights.get(&format!("{}.encoder_attn_layer_norm.bias", prefix)) {
862 self.encoder_attn_layer_norm_bias = Some(w.clone());
863 }
864 if let Some(w) = weights.get(&format!("{}.final_layer_norm.weight", prefix)) {
865 self.final_layer_norm = w.clone();
866 }
867 if let Some(w) = weights.get(&format!("{}.final_layer_norm.bias", prefix)) {
868 self.final_layer_norm_bias = Some(w.clone());
869 }
870
871 if let Some(w) = weights.get(&format!("{}.fc1.weight", prefix)) {
873 self.fc1 = ops_fn::transpose(w)?;
874 }
875 if let Some(w) = weights.get(&format!("{}.fc1.bias", prefix)) {
876 self.fc1_bias = w.clone();
877 }
878 if let Some(w) = weights.get(&format!("{}.fc2.weight", prefix)) {
879 self.fc2 = ops_fn::transpose(w)?;
880 }
881 if let Some(w) = weights.get(&format!("{}.fc2.bias", prefix)) {
882 self.fc2_bias = w.clone();
883 }
884
885 self.self_attn.load_weights(weights, &format!("{}.self_attn", prefix))?;
887 self.encoder_attn.load_weights(weights, &format!("{}.encoder_attn", prefix))?;
888
889 Ok(())
890 }
891
892 fn to_device(&mut self, device: &Device) -> Result<()> {
893 self.self_attn_layer_norm = self.self_attn_layer_norm.to_device(device)?;
894 if let Some(ref mut b) = self.self_attn_layer_norm_bias {
895 *b = b.to_device(device)?;
896 }
897 self.encoder_attn_layer_norm = self.encoder_attn_layer_norm.to_device(device)?;
898 if let Some(ref mut b) = self.encoder_attn_layer_norm_bias {
899 *b = b.to_device(device)?;
900 }
901 self.final_layer_norm = self.final_layer_norm.to_device(device)?;
902 if let Some(ref mut b) = self.final_layer_norm_bias {
903 *b = b.to_device(device)?;
904 }
905 self.fc1 = self.fc1.to_device(device)?;
906 self.fc1_bias = self.fc1_bias.to_device(device)?;
907 self.fc2 = self.fc2.to_device(device)?;
908 self.fc2_bias = self.fc2_bias.to_device(device)?;
909 self.self_attn.to_device(device)?;
910 self.encoder_attn.to_device(device)?;
911 Ok(())
912 }
913}
914
915impl WhisperAttention {
916 fn new(d_model: usize, num_heads: usize, head_dim: usize, device: &Device, is_causal: bool) -> Result<Self> {
917 let scale = 1.0 / (head_dim as f32).sqrt();
918
919 Ok(Self {
920 k_proj: ops_fn::zeros(&[d_model, d_model], DataType::Float32, device)?,
921 k_proj_bias: None,
922 v_proj: ops_fn::zeros(&[d_model, d_model], DataType::Float32, device)?,
923 v_proj_bias: None,
924 q_proj: ops_fn::zeros(&[d_model, d_model], DataType::Float32, device)?,
925 q_proj_bias: None,
926 out_proj: ops_fn::zeros(&[d_model, d_model], DataType::Float32, device)?,
927 out_proj_bias: None,
928 num_heads,
929 head_dim,
930 scale,
931 is_causal,
932 })
933 }
934
935 fn forward(&self, hidden_states: &Tensor, encoder_hidden_states: Option<&Tensor>, is_causal: bool) -> Result<Tensor> {
936 let shape = hidden_states.shape();
937 let (batch_size, seq_len, _) = if shape.len() == 3 {
938 (shape[0], shape[1], shape[2])
939 } else {
940 (1, shape[0], shape[1])
941 };
942
943 let query = ops_fn::matmul(hidden_states, &self.q_proj)?;
945 let query = if let Some(ref bias) = self.q_proj_bias {
946 ops_fn::add(&query, bias)?
947 } else {
948 query
949 };
950
951 let kv_source = encoder_hidden_states.unwrap_or(hidden_states);
953 let kv_seq_len = kv_source.shape()[1];
954
955 let key = ops_fn::matmul(kv_source, &self.k_proj)?;
956 let key = if let Some(ref bias) = self.k_proj_bias {
957 ops_fn::add(&key, bias)?
958 } else {
959 key
960 };
961
962 let value = ops_fn::matmul(kv_source, &self.v_proj)?;
963 let value = if let Some(ref bias) = self.v_proj_bias {
964 ops_fn::add(&value, bias)?
965 } else {
966 value
967 };
968
969 let q_candle = query.to_candle()?;
972 let k_candle = key.to_candle()?;
973 let v_candle = value.to_candle()?;
974
975 let q_reshaped = q_candle
976 .reshape(&[batch_size, seq_len, self.num_heads, self.head_dim])?
977 .transpose(1, 2)?;
978
979 let k_reshaped = k_candle
980 .reshape(&[batch_size, kv_seq_len, self.num_heads, self.head_dim])?
981 .transpose(1, 2)?;
982
983 let v_reshaped = v_candle
984 .reshape(&[batch_size, kv_seq_len, self.num_heads, self.head_dim])?
985 .transpose(1, 2)?;
986
987 let k_t = k_reshaped.transpose(2, 3)?;
990 let q_contiguous = q_reshaped.contiguous()?;
991 let k_contiguous = k_t.contiguous()?;
992
993 let scores = q_contiguous.matmul(&k_contiguous)?;
994 let scaled_scores = (scores * (self.scale as f64))?;
995
996 let masked_scores = if is_causal && encoder_hidden_states.is_none() {
998 let device = scaled_scores.device();
999 let causal_mask = {
1000 let mut mask_data = vec![0.0f32; seq_len * seq_len];
1001 for i in 0..seq_len {
1002 for j in 0..seq_len {
1003 if j > i {
1004 mask_data[i * seq_len + j] = f32::NEG_INFINITY;
1005 }
1006 }
1007 }
1008 candle_core::Tensor::from_vec(mask_data, &[1, 1, seq_len, seq_len], device)?
1009 };
1010 scaled_scores.broadcast_add(&causal_mask)?
1011 } else {
1012 scaled_scores
1013 };
1014
1015 let attention_weights = candle_nn::ops::softmax_last_dim(&masked_scores)?;
1017
1018 let v_contiguous = v_reshaped.contiguous()?;
1020 let attn_output = attention_weights.matmul(&v_contiguous)?;
1021
1022 let attn_output = attn_output
1024 .transpose(1, 2)?
1025 .reshape(&[batch_size, seq_len, self.num_heads * self.head_dim])?;
1026
1027 let attn_output = Tensor::from_candle(attn_output);
1028
1029 let output = ops_fn::matmul(&attn_output, &self.out_proj)?;
1031 let output = if let Some(ref bias) = self.out_proj_bias {
1032 ops_fn::add(&output, bias)?
1033 } else {
1034 output
1035 };
1036
1037 Ok(output)
1038 }
1039
1040 fn load_weights(&mut self, weights: &ModelWeights, prefix: &str) -> Result<()> {
1041 if let Some(w) = weights.get(&format!("{}.k_proj.weight", prefix)) {
1043 self.k_proj = ops_fn::transpose(w)?;
1044 }
1045 if let Some(w) = weights.get(&format!("{}.k_proj.bias", prefix)) {
1046 self.k_proj_bias = Some(w.clone());
1047 }
1048 if let Some(w) = weights.get(&format!("{}.v_proj.weight", prefix)) {
1049 self.v_proj = ops_fn::transpose(w)?;
1050 }
1051 if let Some(w) = weights.get(&format!("{}.v_proj.bias", prefix)) {
1052 self.v_proj_bias = Some(w.clone());
1053 }
1054 if let Some(w) = weights.get(&format!("{}.q_proj.weight", prefix)) {
1055 self.q_proj = ops_fn::transpose(w)?;
1056 }
1057 if let Some(w) = weights.get(&format!("{}.q_proj.bias", prefix)) {
1058 self.q_proj_bias = Some(w.clone());
1059 }
1060 if let Some(w) = weights.get(&format!("{}.out_proj.weight", prefix)) {
1061 self.out_proj = ops_fn::transpose(w)?;
1062 }
1063 if let Some(w) = weights.get(&format!("{}.out_proj.bias", prefix)) {
1064 self.out_proj_bias = Some(w.clone());
1065 }
1066 Ok(())
1067 }
1068
1069 fn to_device(&mut self, device: &Device) -> Result<()> {
1070 self.k_proj = self.k_proj.to_device(device)?;
1071 if let Some(ref mut b) = self.k_proj_bias { *b = b.to_device(device)?; }
1072 self.v_proj = self.v_proj.to_device(device)?;
1073 if let Some(ref mut b) = self.v_proj_bias { *b = b.to_device(device)?; }
1074 self.q_proj = self.q_proj.to_device(device)?;
1075 if let Some(ref mut b) = self.q_proj_bias { *b = b.to_device(device)?; }
1076 self.out_proj = self.out_proj.to_device(device)?;
1077 if let Some(ref mut b) = self.out_proj_bias { *b = b.to_device(device)?; }
1078 Ok(())
1079 }
1080}
1081
1082fn create_sinusoidal_embeddings(max_len: usize, d_model: usize, device: &Device) -> Result<Tensor> {
1086 let mut embeddings = Vec::with_capacity(max_len * d_model);
1087
1088 for pos in 0..max_len {
1089 for i in 0..d_model {
1090 let angle = (pos as f32) / 10000_f32.powf((2 * (i / 2)) as f32 / d_model as f32);
1091 let value = if i % 2 == 0 {
1092 angle.sin()
1093 } else {
1094 angle.cos()
1095 };
1096 embeddings.push(value);
1097 }
1098 }
1099
1100 Tensor::from_f32_slice(&embeddings, &[max_len, d_model], device)
1101}
1102
1103#[cfg(test)]
1104mod tests {
1105 use super::*;
1106
1107 #[test]
1108 fn test_whisper_model_creation() {
1109 let config = WhisperConfig {
1110 vocab_size: 1000,
1111 d_model: 128,
1112 hidden_size: 128,
1113 encoder_layers: 2,
1114 decoder_layers: 2,
1115 num_hidden_layers: 4,
1116 encoder_attention_heads: 4,
1117 decoder_attention_heads: 4,
1118 encoder_ffn_dim: 512,
1119 decoder_ffn_dim: 512,
1120 num_mel_bins: 80,
1121 max_source_positions: 100,
1122 max_target_positions: 50,
1123 ..Default::default()
1124 };
1125
1126 let model = WhisperModelV2::new(config).unwrap();
1127 assert_eq!(model.config().vocab_size(), 1000);
1128 assert_eq!(model.config().hidden_size(), 128);
1129 }
1130
1131 #[test]
1132 fn test_whisper_encoder_forward() {
1133 let config = WhisperConfig {
1134 d_model: 64,
1135 hidden_size: 64,
1136 num_hidden_layers: 1,
1137 encoder_layers: 1,
1138 decoder_layers: 1,
1139 encoder_attention_heads: 2,
1140 decoder_attention_heads: 2,
1141 encoder_ffn_dim: 256,
1142 decoder_ffn_dim: 256,
1143 num_mel_bins: 40,
1144 max_source_positions: 50,
1145 max_target_positions: 25,
1146 ..Default::default()
1147 };
1148
1149 let encoder = WhisperEncoder::new(&config, &Device::CPU).unwrap();
1150
1151 let mel = ops_fn::zeros(&[1, 40, 100], DataType::Float32, &Device::CPU).unwrap();
1153
1154 let output = encoder.forward(&mel).unwrap();
1155 assert_eq!(output.shape()[0], 1); assert_eq!(output.shape()[1], 50); assert_eq!(output.shape()[2], 64); }
1160
1161 #[test]
1162 fn test_whisper_decoder_forward() {
1163 let config = WhisperConfig {
1164 vocab_size: 100,
1165 d_model: 64,
1166 hidden_size: 64,
1167 num_hidden_layers: 1,
1168 encoder_layers: 1,
1169 decoder_layers: 1,
1170 encoder_attention_heads: 2,
1171 decoder_attention_heads: 2,
1172 encoder_ffn_dim: 256,
1173 decoder_ffn_dim: 256,
1174 max_target_positions: 25,
1175 ..Default::default()
1176 };
1177
1178 let decoder = WhisperDecoder::new(&config, &Device::CPU).unwrap();
1179
1180 let input_ids = ops_fn::zeros(&[1, 5], DataType::Int64, &Device::CPU).unwrap();
1182 let encoder_hidden = ops_fn::zeros(&[1, 20, 64], DataType::Float32, &Device::CPU).unwrap();
1183
1184 let output = decoder.forward(&input_ids, Some(&encoder_hidden)).unwrap();
1185 assert_eq!(output.shape(), &[1, 5, 64]); }
1187
1188 #[test]
1189 fn test_whisper_full_forward() {
1190 let config = WhisperConfig {
1191 vocab_size: 100,
1192 d_model: 64,
1193 hidden_size: 64,
1194 num_hidden_layers: 2,
1195 encoder_layers: 1,
1196 decoder_layers: 1,
1197 encoder_attention_heads: 2,
1198 decoder_attention_heads: 2,
1199 encoder_ffn_dim: 256,
1200 decoder_ffn_dim: 256,
1201 num_mel_bins: 40,
1202 max_source_positions: 50,
1203 max_target_positions: 25,
1204 decoder_start_token_id: 1, bos_token_id: 1,
1206 eos_token_id: 2,
1207 pad_token_id: 0,
1208 ..Default::default()
1209 };
1210
1211 let model = WhisperModelV2::new(config).unwrap();
1212
1213 let mel = ops_fn::zeros(&[1, 40, 100], DataType::Float32, &Device::CPU).unwrap();
1215
1216 let inputs = ModelInputs::Audio {
1217 input_features: mel,
1218 attention_mask: None,
1219 };
1220
1221 let outputs = model.forward(&inputs).unwrap();
1222
1223 match outputs {
1224 ModelOutputs::Sequence { logits, encoder_hidden_states, decoder_hidden_states } => {
1225 assert_eq!(logits.shape()[0], 1); assert_eq!(logits.shape()[1], 1); assert_eq!(logits.shape()[2], 100); assert!(encoder_hidden_states.is_some());
1229 assert!(decoder_hidden_states.is_some());
1230 }
1231 _ => panic!("Expected Sequence output"),
1232 }
1233 }
1234
1235 #[test]
1236 fn test_sinusoidal_embeddings() {
1237 let embeddings = create_sinusoidal_embeddings(100, 64, &Device::CPU).unwrap();
1238 assert_eq!(embeddings.shape(), &[100, 64]);
1239 }
1240}