Skip to main content

runtime/models_v2/
whisper.rs

1//! Whisper Model V2 - Clean implementation using solid abstractions
2//!
3//! This implements the Whisper speech recognition encoder-decoder architecture:
4//! - Audio encoder: Conv1D for feature extraction + transformer with bidirectional self-attention
5//! - Text decoder: transformer with causal self-attention and cross-attention
6//! - Supports Whisper-tiny, Whisper-base, Whisper-small, Whisper-medium, Whisper-large
7
8use crate::model_config;
9use super::traits::*;
10use anyhow::Result;
11use serde::{Serialize, Deserialize};
12
13/// Whisper model configuration using the model_config macro
14model_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    // Whisper specific
37    max_source_positions: usize = 1500,
38    max_target_positions: usize = 448,
39    num_mel_bins: usize = 80,
40    // Required by model_config macro but not used directly
41    num_hidden_layers: usize = 6,
42    hidden_size: usize = 512,
43});
44
45impl WhisperConfig {
46    /// Create WhisperConfig from GGUF model configuration
47    pub fn from_gguf_config(gguf: &crate::weight_loader_core::GGUFModelConfig) -> Self {
48        // Map GGUF config to Whisper config
49        // Whisper in GGUF has different field names
50        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, // Approximate split
55            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    /// Get head dimension for encoder
67    pub fn encoder_head_dim(&self) -> usize {
68        self.d_model / self.encoder_attention_heads
69    }
70
71    /// Get head dimension for decoder
72    pub fn decoder_head_dim(&self) -> usize {
73        self.d_model / self.decoder_attention_heads
74    }
75}
76
77/// Main Whisper model implementation
78pub struct WhisperModelV2 {
79    config: WhisperConfig,
80    device: Device,
81    encoder: WhisperEncoder,
82    decoder: WhisperDecoder,
83    proj_out: Tensor, // Output projection (tied with embed_tokens)
84}
85
86/// Whisper audio encoder with Conv1D preprocessing and transformer layers
87pub struct WhisperEncoder {
88    // Conv1D layers for mel spectrogram feature extraction
89    conv1_weight: Tensor, // [d_model, n_mels, 3] - kernel size 3
90    conv1_bias: Tensor,   // [d_model]
91    conv2_weight: Tensor, // [d_model, d_model, 3] - kernel size 3, stride 2
92    conv2_bias: Tensor,   // [d_model]
93    // Sinusoidal position embeddings
94    embed_positions: Tensor, // [max_source_positions, d_model]
95    // Transformer layers
96    layers: Vec<WhisperEncoderLayer>,
97    // Final layer norm
98    layer_norm: Tensor,
99    layer_norm_bias: Option<Tensor>,
100    config: WhisperConfig,
101}
102
103/// Whisper text decoder with causal self-attention and cross-attention
104pub struct WhisperDecoder {
105    // Token embedding
106    embed_tokens: Tensor, // [vocab_size, d_model]
107    // Learned position embeddings
108    embed_positions: Tensor, // [max_target_positions, d_model]
109    // Transformer layers
110    layers: Vec<WhisperDecoderLayer>,
111    // Final layer norm
112    layer_norm: Tensor,
113    layer_norm_bias: Option<Tensor>,
114    config: WhisperConfig,
115}
116
117/// Whisper encoder transformer layer with bidirectional self-attention
118pub 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
131/// Whisper decoder transformer layer with causal self-attention and cross-attention
132pub 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
148/// Whisper multi-head attention
149pub 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, // True for decoder self-attention
162}
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        // Output projection (tied with decoder embed_tokens)
173        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        // Load output projection (may be tied with embed_tokens)
184        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            // Tied weights - transpose for projection
188            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                // Encoder forward: process mel spectrogram
198                let encoder_outputs = self.encoder.forward(input_features)?;
199
200                // For inference, we need decoder_input_ids
201                // If not provided, use start token
202                let start_token = self.config.decoder_start_token_id;
203                let batch_size = input_features.shape()[0];
204
205                // Create initial decoder input [batch, 1] with start token
206                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                // Decoder forward
214                let decoder_outputs = self.decoder.forward(&decoder_input, Some(&encoder_outputs))?;
215
216                // Project to vocabulary
217                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        // For Whisper, generate() would typically be called with audio input
231        // This is a text-based fallback that returns a placeholder
232        // Real usage should call transcribe() with audio data
233
234        // In practice, Whisper generation works as follows:
235        // 1. Encode mel spectrogram with encoder
236        // 2. Autoregressively decode with decoder using encoder outputs
237        // 3. Sample tokens until EOS or max_length
238
239        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        // Approximate parameter count
253        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); // 8 for self + cross attn
255        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; // float32
260
261        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    /// Transcribe audio mel spectrogram to text
283    /// mel_spectrogram shape: [batch, n_mels, n_frames]
284    pub fn transcribe(&self, mel_spectrogram: &Tensor, config: &GenerationConfig) -> Result<Vec<u32>> {
285        // 1. Encode audio
286        let encoder_outputs = self.encoder.forward(mel_spectrogram)?;
287
288        let batch_size = mel_spectrogram.shape()[0];
289
290        // 2. Initialize decoder with start token
291        let mut tokens: Vec<u32> = vec![self.config.decoder_start_token_id as u32];
292
293        // 3. Autoregressive decoding loop
294        for _ in 0..config.max_new_tokens {
295            // Create decoder input from current tokens
296            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            // Decoder forward pass
304            let decoder_outputs = self.decoder.forward(&decoder_input, Some(&encoder_outputs))?;
305
306            // Get logits for last position
307            let logits = ops_fn::matmul(&decoder_outputs, &self.proj_out)?;
308
309            // Extract last position logits
310            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            // Greedy decoding (simplified)
322            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                    // Skip suppressed tokens
327                    if val > max_val {
328                        max_val = val;
329                        max_idx = idx;
330                    }
331                }
332                max_idx as u32
333            };
334
335            // Check for EOS
336            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)?); // bidirectional
352        }
353
354        // Conv1D: [out_channels, in_channels, kernel_size]
355        // conv1: 80 mel bins -> d_model, kernel=3, stride=1, padding=1
356        // conv2: d_model -> d_model, kernel=3, stride=2, padding=1
357        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        // Sinusoidal position embeddings for encoder
363        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        // Input: [batch, n_mels, n_frames]
380        // 1. Apply Conv1D layers
381        let mut hidden_states = self.apply_conv1d(mel_spectrogram)?;
382
383        // After conv2 with stride 2, sequence length is halved
384        // hidden_states shape: [batch, d_model, n_frames/2]
385
386        // 2. Transpose to [batch, seq, d_model] for transformer
387        let hidden_candle = hidden_states.to_candle()?;
388        let transposed = hidden_candle.transpose(1, 2)?;
389        hidden_states = Tensor::from_candle(transposed);
390
391        // 3. Add sinusoidal positional embeddings
392        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        // 4. Apply transformer encoder layers (bidirectional)
397        for layer in &self.layers {
398            hidden_states = layer.forward(&hidden_states, false)?; // is_causal=false
399        }
400
401        // 5. Final layer norm
402        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    /// Apply Conv1D feature extraction
408    fn apply_conv1d(&self, mel_spectrogram: &Tensor) -> Result<Tensor> {
409        // Input: [batch, n_mels, n_frames]
410        // Whisper uses two 1D convolutions:
411        // conv1: kernel=3, stride=1, padding=1 (preserves length)
412        // conv2: kernel=3, stride=2, padding=1 (halves length)
413
414        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        // For simplicity, implement conv1d as a series of operations
420        // In practice, this should use optimized conv1d kernel
421
422        // Conv1: [batch, n_mels, n_frames] -> [batch, d_model, n_frames]
423        // Unfold with kernel_size=3, stride=1, padding=1
424        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        // Conv2: [batch, d_model, n_frames] -> [batch, d_model, n_frames/2]
428        // Unfold with kernel_size=3, stride=2, padding=1
429        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    /// Simple Conv1D implementation
436    /// weight shape: [out_channels, in_channels, kernel_size]
437    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        // Calculate output length
454        let out_length = (in_length + 2 * padding - kernel_size) / stride + 1;
455
456        // Pad input if needed
457        let padded = if padding > 0 {
458            // Pad along the last dimension
459            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        // Simple implementation: unfold + matmul
467        // For each output position, extract kernel_size elements and multiply
468        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)?; // [batch, in_ch, kernel]
473
474            // Flatten patch: [batch, in_ch * kernel]
475            let patch_flat = patch.reshape(&[batch_size, in_channels * kernel_size])?;
476
477            // Reshape weight: [out_ch, in_ch * kernel]
478            let weight_flat = weight_candle.reshape(&[out_channels, in_channels * kernel_size])?;
479
480            // Matmul: [batch, in_ch * kernel] @ [in_ch * kernel, out_ch] -> [batch, out_ch]
481            let weight_t = weight_flat.t()?;
482            let out_pos = patch_flat.matmul(&weight_t)?;
483
484            output_slices.push(out_pos.unsqueeze(2)?); // [batch, out_ch, 1]
485        }
486
487        // Concatenate along sequence dimension
488        let refs: Vec<&candle_core::Tensor> = output_slices.iter().collect();
489        let output = candle_core::Tensor::cat(&refs, 2)?; // [batch, out_ch, out_len]
490
491        // Add bias: [out_ch] broadcast to [batch, out_ch, out_len]
492        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    /// Get sinusoidal position embeddings for the given sequence length
499    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)?; // [1, seq, d_model] for broadcasting
503        Ok(Tensor::from_candle(expanded))
504    }
505
506    fn load_weights(&mut self, weights: &ModelWeights) -> Result<()> {
507        // Load conv weights (transpose for our conv1d impl)
508        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        // Load position embeddings (usually not loaded - computed)
522        if let Some(w) = weights.get("model.encoder.embed_positions.weight") {
523            self.embed_positions = w.clone();
524        }
525
526        // Load final layer norm
527        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        // Load transformer layer weights
535        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        // Token embeddings
567        let embed_tokens = ops_fn::zeros(&[config.vocab_size, config.d_model], DataType::Float32, device)?;
568
569        // Learned position embeddings for decoder
570        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        // 1. Token embedding lookup
584        let mut hidden_states = ops_fn::embedding(input_ids, &self.embed_tokens)?;
585
586        // 2. Add learned positional embeddings
587        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        // 3. Apply transformer decoder layers
592        for layer in &self.layers {
593            hidden_states = layer.forward(&hidden_states, encoder_hidden_states)?;
594        }
595
596        // 4. Final layer norm
597        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    /// Get learned position embeddings for the given sequence length
603    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)?; // [1, seq, d_model] for broadcasting
607        Ok(Tensor::from_candle(expanded))
608    }
609
610    fn load_weights(&mut self, weights: &ModelWeights) -> Result<()> {
611        // Load embeddings (no transpose - used for lookup)
612        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        // Load final layer norm
620        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        // Load transformer layer weights
628        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        // Pre-norm architecture
675        // 1. Self-attention with residual
676        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        // 2. Feed-forward with residual
687        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        // FFN: fc1 -> activation -> fc2
696        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        // Load layer norms
713        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        // Load FFN weights (transpose for matmul)
727        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        // Load attention weights
741        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            // Causal self-attention
770            self_attn: WhisperAttention::new(
771                config.d_model,
772                config.decoder_attention_heads,
773                head_dim,
774                device,
775                true, // causal
776            )?,
777            self_attn_layer_norm: ops_fn::zeros(&[config.d_model], DataType::Float32, device)?,
778            self_attn_layer_norm_bias: None,
779            // Cross-attention (non-causal, uses encoder outputs)
780            encoder_attn: WhisperAttention::new(
781                config.d_model,
782                config.decoder_attention_heads,
783                head_dim,
784                device,
785                false, // not causal for cross-attention
786            )?,
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        // 1. Causal self-attention with residual
801        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)?; // causal=true
809        let hidden_states = ops_fn::add(&residual, &hidden_states)?;
810
811        // 2. Cross-attention with encoder outputs (if provided)
812        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        // 3. Feed-forward with residual
827        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        // Load layer norms
852        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        // Load FFN weights (transpose for matmul)
872        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        // Load attention weights
886        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        // Project query from decoder hidden states
944        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        // For cross-attention, K/V come from encoder; for self-attention, from same input
952        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        // Reshape for multi-head attention
970        // [batch, seq, d_model] -> [batch, seq, heads, head_dim] -> [batch, heads, seq, head_dim]
971        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        // Scaled dot-product attention
988        // scores = Q @ K^T / sqrt(head_dim)
989        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        // Apply causal mask if needed (decoder self-attention)
997        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        // Softmax
1016        let attention_weights = candle_nn::ops::softmax_last_dim(&masked_scores)?;
1017
1018        // Apply attention to values
1019        let v_contiguous = v_reshaped.contiguous()?;
1020        let attn_output = attention_weights.matmul(&v_contiguous)?;
1021
1022        // Reshape back: [batch, heads, seq, head_dim] -> [batch, seq, d_model]
1023        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        // Output projection
1030        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        // Load projection weights (transpose for matmul: [out, in] -> [in, out])
1042        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
1082/// Create sinusoidal positional embeddings
1083/// PE(pos, 2i) = sin(pos / 10000^(2i/d_model))
1084/// PE(pos, 2i+1) = cos(pos / 10000^(2i/d_model))
1085fn 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        // Create dummy mel spectrogram [batch=1, n_mels=40, n_frames=100]
1152        let mel = ops_fn::zeros(&[1, 40, 100], DataType::Float32, &Device::CPU).unwrap();
1153
1154        let output = encoder.forward(&mel).unwrap();
1155        // After conv2 with stride 2: n_frames/2 = 50
1156        assert_eq!(output.shape()[0], 1); // batch
1157        assert_eq!(output.shape()[1], 50); // seq (n_frames/2)
1158        assert_eq!(output.shape()[2], 64); // d_model
1159    }
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        // Create dummy input
1181        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]); // [batch, seq, d_model]
1186    }
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, // Use a valid token ID for the test vocab size
1205            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        // Create dummy mel spectrogram
1214        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); // batch
1226                assert_eq!(logits.shape()[1], 1); // seq (just start token)
1227                assert_eq!(logits.shape()[2], 100); // vocab
1228                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}