Skip to main content

runtime/models_v2/
wav2vec2.rs

1//! Wav2Vec2 Model V2 - Self-Supervised Audio Encoder
2//!
3//! Audio encoder with CNN feature extractor + transformer encoder
4
5use crate::model_config;
6use super::traits::*;
7use anyhow::Result;
8use serde::{Serialize, Deserialize};
9
10model_config!(Wav2Vec2Config {
11    vocab_size: usize = 32,
12    hidden_size: usize = 768,
13    intermediate_size: usize = 3072,
14    num_hidden_layers: usize = 12,
15    num_attention_heads: usize = 12,
16    conv_dim: Vec<usize> = vec![512, 512, 512, 512, 512, 512, 512],
17    conv_kernel: Vec<usize> = vec![10, 3, 3, 3, 3, 2, 2],
18    conv_stride: Vec<usize> = vec![5, 2, 2, 2, 2, 2, 2],
19    feat_extract_norm: String = "group".to_string(),
20    layer_norm_eps: f32 = 1e-5,
21    hidden_dropout: f32 = 0.1,
22    attention_dropout: f32 = 0.1,
23    final_dropout: f32 = 0.1,
24    pad_token_id: i64 = 0,
25    bos_token_id: i64 = 1,
26    eos_token_id: i64 = 2,
27});
28
29impl Wav2Vec2Config {
30    pub fn from_gguf_config(gguf: &crate::weight_loader_core::GGUFModelConfig) -> Self {
31        Self {
32            hidden_size: gguf.hidden_size,
33            num_hidden_layers: gguf.num_hidden_layers,
34            num_attention_heads: gguf.num_attention_heads,
35            ..Default::default()
36        }
37    }
38}
39
40pub struct Wav2Vec2ModelV2 {
41    config: Wav2Vec2Config,
42    device: Device,
43    feature_extractor: Wav2Vec2FeatureExtractor,
44    feature_projection: Tensor,
45    encoder: Wav2Vec2Encoder,
46    lm_head: Tensor,
47}
48
49pub struct Wav2Vec2FeatureExtractor {
50    conv_layers: Vec<Wav2Vec2ConvLayer>,
51}
52
53pub struct Wav2Vec2ConvLayer {
54    conv_weight: Tensor,
55    conv_bias: Option<Tensor>,
56    layer_norm: Option<Tensor>,
57    in_channels: usize,
58    out_channels: usize,
59    kernel_size: usize,
60    stride: usize,
61}
62
63pub struct Wav2Vec2Encoder {
64    layers: Vec<Wav2Vec2EncoderLayer>,
65    pos_conv_embed: Tensor,
66    layer_norm: Tensor,
67}
68
69pub struct Wav2Vec2EncoderLayer {
70    self_attn_q: Tensor,
71    self_attn_k: Tensor,
72    self_attn_v: Tensor,
73    self_attn_o: Tensor,
74    fc1: Tensor,
75    fc2: Tensor,
76    layer_norm: Tensor,
77    final_layer_norm: Tensor,
78    num_heads: usize,
79    head_dim: usize,
80}
81
82impl Model for Wav2Vec2ModelV2 {
83    type Config = Wav2Vec2Config;
84
85    fn new(config: Wav2Vec2Config) -> Result<Self> {
86        let device = Device::CPU;
87
88        let feature_extractor = Wav2Vec2FeatureExtractor::new(&config, &device)?;
89        let last_conv_dim = *config.conv_dim.last().unwrap_or(&512);
90        let feature_projection = ops_fn::zeros(&[last_conv_dim, config.hidden_size], DataType::Float32, &device)?;
91        let encoder = Wav2Vec2Encoder::new(&config, &device)?;
92        let lm_head = ops_fn::zeros(&[config.hidden_size, config.vocab_size], DataType::Float32, &device)?;
93
94        Ok(Self { config, device, feature_extractor, feature_projection, encoder, lm_head })
95    }
96
97    fn from_weights(config: Wav2Vec2Config, weights: ModelWeights) -> Result<Self> {
98        let mut model = Self::new(config)?;
99        if let Some(w) = weights.get("lm_head.weight") { model.lm_head = ops_fn::transpose(w)?; }
100        Ok(model)
101    }
102
103    fn forward(&self, inputs: &ModelInputs) -> Result<ModelOutputs> {
104        match inputs {
105            ModelInputs::Audio { input_features, .. } => {
106                // Feature extraction (CNN)
107                let features = self.feature_extractor.forward(input_features)?;
108
109                // Feature projection
110                let hidden = ops_fn::matmul(&features, &self.feature_projection)?;
111
112                // Transformer encoder
113                let hidden = self.encoder.forward(&hidden)?;
114
115                // LM head
116                let logits = ops_fn::matmul(&hidden, &self.lm_head)?;
117
118                Ok(ModelOutputs::Logits { logits, hidden_states: Some(hidden) })
119            }
120            _ => Err(anyhow::anyhow!("Wav2Vec2 requires audio input")),
121        }
122    }
123
124    fn generate(&self, _prompt: &str, _config: &GenerationConfig) -> Result<String> {
125        Err(anyhow::anyhow!("Wav2Vec2 is an encoder model, use forward() for transcription"))
126    }
127
128    fn config(&self) -> &Self::Config { &self.config }
129
130    fn memory_requirements(&self) -> MemoryRequirements {
131        let p = (self.config.hidden_size * self.config.hidden_size * 4 * self.config.num_hidden_layers) * 4;
132        MemoryRequirements { gpu_memory: p, cpu_memory: p / 4, kv_cache_memory: 0, peak_memory: p * 2 }
133    }
134
135    fn to_device(&mut self, device: &Device) -> Result<()> {
136        self.device = device.clone();
137        self.feature_projection = self.feature_projection.to_device(device)?;
138        self.lm_head = self.lm_head.to_device(device)?;
139        Ok(())
140    }
141}
142
143impl Wav2Vec2FeatureExtractor {
144    fn new(config: &Wav2Vec2Config, device: &Device) -> Result<Self> {
145        let mut conv_layers = Vec::new();
146        let mut in_channels = 1; // Raw audio waveform
147
148        for i in 0..config.conv_dim.len() {
149            let out_channels = config.conv_dim[i];
150            let kernel_size = config.conv_kernel.get(i).copied().unwrap_or(3);
151            let stride = config.conv_stride.get(i).copied().unwrap_or(1);
152
153            conv_layers.push(Wav2Vec2ConvLayer::new(in_channels, out_channels, kernel_size, stride, i == 0, device)?);
154            in_channels = out_channels;
155        }
156
157        Ok(Self { conv_layers })
158    }
159
160    fn forward(&self, input: &Tensor) -> Result<Tensor> {
161        let mut hidden = input.clone();
162
163        for layer in &self.conv_layers {
164            hidden = layer.forward(&hidden)?;
165        }
166
167        // Transpose for transformer: [batch, channels, time] -> [batch, time, channels]
168        let hidden_candle = hidden.to_candle()?;
169        let transposed = hidden_candle.transpose(1, 2)?;
170        Ok(Tensor::from_candle(transposed))
171    }
172}
173
174impl Wav2Vec2ConvLayer {
175    fn new(in_channels: usize, out_channels: usize, kernel_size: usize, stride: usize, has_layer_norm: bool, device: &Device) -> Result<Self> {
176        let conv_weight = ops_fn::zeros(&[out_channels, in_channels, kernel_size], DataType::Float32, device)?;
177        let conv_bias = Some(ops_fn::zeros(&[out_channels], DataType::Float32, device)?);
178        let layer_norm = if has_layer_norm {
179            Some(ops_fn::zeros(&[out_channels], DataType::Float32, device)?)
180        } else {
181            None
182        };
183
184        Ok(Self { conv_weight, conv_bias, layer_norm, in_channels, out_channels, kernel_size, stride })
185    }
186
187    fn forward(&self, input: &Tensor) -> Result<Tensor> {
188        // Simplified conv1d implementation
189        let conv_out = ops_fn::conv1d(input, &self.conv_weight, None, self.stride, 0)?;
190
191        let conv_out = if let Some(ref bias) = self.conv_bias {
192            let bias_candle = bias.to_candle()?;
193            let out_candle = conv_out.to_candle()?;
194            let bias_expanded = bias_candle.reshape(&[1, self.out_channels, 1])?;
195            Tensor::from_candle(out_candle.broadcast_add(&bias_expanded)?)
196        } else {
197            conv_out
198        };
199
200        let conv_out = if let Some(ref ln) = self.layer_norm {
201            // Group norm for first layer
202            ops_fn::layer_norm(&conv_out, ln, None, 1e-5)?
203        } else {
204            conv_out
205        };
206
207        // GELU activation
208        ops_fn::gelu(&conv_out)
209    }
210}
211
212impl Wav2Vec2Encoder {
213    fn new(config: &Wav2Vec2Config, device: &Device) -> Result<Self> {
214        let mut layers = Vec::with_capacity(config.num_hidden_layers);
215        for _ in 0..config.num_hidden_layers {
216            layers.push(Wav2Vec2EncoderLayer::new(config, device)?);
217        }
218
219        let pos_conv_embed = ops_fn::zeros(&[config.hidden_size, config.hidden_size, 128], DataType::Float32, device)?;
220        let layer_norm = ops_fn::zeros(&[config.hidden_size], DataType::Float32, device)?;
221
222        Ok(Self { layers, pos_conv_embed, layer_norm })
223    }
224
225    fn forward(&self, hidden_states: &Tensor) -> Result<Tensor> {
226        let mut hidden = hidden_states.clone();
227
228        // Positional embeddings (simplified - would be conv in full impl)
229        hidden = ops_fn::layer_norm(&hidden, &self.layer_norm, None, 1e-5)?;
230
231        for layer in &self.layers {
232            hidden = layer.forward(&hidden)?;
233        }
234
235        Ok(hidden)
236    }
237}
238
239impl Wav2Vec2EncoderLayer {
240    fn new(config: &Wav2Vec2Config, device: &Device) -> Result<Self> {
241        let head_dim = config.hidden_size / config.num_attention_heads;
242
243        Ok(Self {
244            self_attn_q: ops_fn::zeros(&[config.hidden_size, config.hidden_size], DataType::Float32, device)?,
245            self_attn_k: ops_fn::zeros(&[config.hidden_size, config.hidden_size], DataType::Float32, device)?,
246            self_attn_v: ops_fn::zeros(&[config.hidden_size, config.hidden_size], DataType::Float32, device)?,
247            self_attn_o: ops_fn::zeros(&[config.hidden_size, config.hidden_size], DataType::Float32, device)?,
248            fc1: ops_fn::zeros(&[config.hidden_size, config.intermediate_size], DataType::Float32, device)?,
249            fc2: ops_fn::zeros(&[config.intermediate_size, config.hidden_size], DataType::Float32, device)?,
250            layer_norm: ops_fn::zeros(&[config.hidden_size], DataType::Float32, device)?,
251            final_layer_norm: ops_fn::zeros(&[config.hidden_size], DataType::Float32, device)?,
252            num_heads: config.num_attention_heads,
253            head_dim,
254        })
255    }
256
257    fn forward(&self, hidden_states: &Tensor) -> Result<Tensor> {
258        let shape = hidden_states.shape();
259        let (batch_size, seq_len, _) = (shape[0], shape[1], shape[2]);
260
261        // Self-attention
262        let residual = hidden_states.clone();
263        let hidden = ops_fn::layer_norm(hidden_states, &self.layer_norm, None, 1e-5)?;
264
265        let q = ops_fn::matmul(&hidden, &self.self_attn_q)?.to_candle()?
266            .reshape(&[batch_size, seq_len, self.num_heads, self.head_dim])?.transpose(1, 2)?;
267        let k = ops_fn::matmul(&hidden, &self.self_attn_k)?.to_candle()?
268            .reshape(&[batch_size, seq_len, self.num_heads, self.head_dim])?.transpose(1, 2)?;
269        let v = ops_fn::matmul(&hidden, &self.self_attn_v)?.to_candle()?
270            .reshape(&[batch_size, seq_len, self.num_heads, self.head_dim])?.transpose(1, 2)?;
271
272        let scale = (self.head_dim as f32).powf(-0.5);
273        let scores = (q.contiguous()?.matmul(&k.transpose(2, 3)?.contiguous()?)? * (scale as f64))?;
274        let attn = candle_nn::ops::softmax_last_dim(&scores)?.matmul(&v.contiguous()?)?
275            .transpose(1, 2)?.reshape(&[batch_size, seq_len, self.num_heads * self.head_dim])?;
276
277        let hidden = ops_fn::add(&residual, &ops_fn::matmul(&Tensor::from_candle(attn), &self.self_attn_o)?)?;
278
279        // FFN
280        let residual = hidden.clone();
281        let hidden = ops_fn::layer_norm(&hidden, &self.final_layer_norm, None, 1e-5)?;
282        let hidden = ops_fn::gelu(&ops_fn::matmul(&hidden, &self.fc1)?)?;
283        ops_fn::add(&residual, &ops_fn::matmul(&hidden, &self.fc2)?)
284    }
285}
286
287#[cfg(test)]
288mod tests {
289    use super::*;
290
291    #[test]
292    fn test_wav2vec2_config() {
293        let config = Wav2Vec2Config::default();
294        assert_eq!(config.hidden_size, 768);
295        assert_eq!(config.num_hidden_layers, 12);
296    }
297
298    #[test]
299    fn test_wav2vec2_model_creation() {
300        let config = Wav2Vec2Config {
301            hidden_size: 64,
302            intermediate_size: 256,
303            num_hidden_layers: 2,
304            num_attention_heads: 4,
305            conv_dim: vec![32, 32],
306            conv_kernel: vec![5, 3],
307            conv_stride: vec![2, 2],
308            ..Default::default()
309        };
310
311        let model = Wav2Vec2ModelV2::new(config).unwrap();
312        assert_eq!(model.config().hidden_size(), 64);
313    }
314}