1use 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 let features = self.feature_extractor.forward(input_features)?;
108
109 let hidden = ops_fn::matmul(&features, &self.feature_projection)?;
111
112 let hidden = self.encoder.forward(&hidden)?;
114
115 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; 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 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 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 ops_fn::layer_norm(&conv_out, ln, None, 1e-5)?
203 } else {
204 conv_out
205 };
206
207 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 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 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 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}