1use crate::model_config;
6use super::traits::*;
7use anyhow::Result;
8use serde::{Serialize, Deserialize};
9
10model_config!(HuBERTConfig {
11 vocab_size: usize = 100,
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 layer_norm_eps: f32 = 1e-5,
20 hidden_dropout: f32 = 0.1,
21 pad_token_id: i64 = 0,
22 bos_token_id: i64 = 1,
23 eos_token_id: i64 = 2,
24});
25
26impl HuBERTConfig {
27 pub fn from_gguf_config(gguf: &crate::weight_loader_core::GGUFModelConfig) -> Self {
28 Self {
29 hidden_size: gguf.hidden_size,
30 num_hidden_layers: gguf.num_hidden_layers,
31 num_attention_heads: gguf.num_attention_heads,
32 ..Default::default()
33 }
34 }
35}
36
37pub struct HuBERTModelV2 {
38 config: HuBERTConfig,
39 device: Device,
40 feature_projection: Tensor,
41 encoder_layers: Vec<HuBERTEncoderLayer>,
42 encoder_norm: Tensor,
43 lm_head: Tensor,
44}
45
46pub struct HuBERTEncoderLayer {
47 self_attn_q: Tensor,
48 self_attn_k: Tensor,
49 self_attn_v: Tensor,
50 self_attn_o: Tensor,
51 fc1: Tensor,
52 fc2: Tensor,
53 layer_norm: Tensor,
54 final_layer_norm: Tensor,
55 num_heads: usize,
56 head_dim: usize,
57}
58
59impl Model for HuBERTModelV2 {
60 type Config = HuBERTConfig;
61
62 fn new(config: HuBERTConfig) -> Result<Self> {
63 let device = Device::CPU;
64
65 let last_conv_dim = *config.conv_dim.last().unwrap_or(&512);
66 let feature_projection = ops_fn::zeros(&[last_conv_dim, config.hidden_size], DataType::Float32, &device)?;
67 let encoder_norm = ops_fn::zeros(&[config.hidden_size], DataType::Float32, &device)?;
68 let lm_head = ops_fn::zeros(&[config.hidden_size, config.vocab_size], DataType::Float32, &device)?;
69
70 let mut encoder_layers = Vec::with_capacity(config.num_hidden_layers);
71 for _ in 0..config.num_hidden_layers {
72 encoder_layers.push(HuBERTEncoderLayer::new(&config, &device)?);
73 }
74
75 Ok(Self { config, device, feature_projection, encoder_layers, encoder_norm, lm_head })
76 }
77
78 fn from_weights(config: HuBERTConfig, weights: ModelWeights) -> Result<Self> {
79 let mut model = Self::new(config)?;
80 if let Some(w) = weights.get("lm_head.weight") { model.lm_head = ops_fn::transpose(w)?; }
81 Ok(model)
82 }
83
84 fn forward(&self, inputs: &ModelInputs) -> Result<ModelOutputs> {
85 match inputs {
86 ModelInputs::Audio { input_features, .. } => {
87 let hidden = ops_fn::matmul(input_features, &self.feature_projection)?;
89
90 let mut hidden = hidden;
92 for layer in &self.encoder_layers {
93 hidden = layer.forward(&hidden)?;
94 }
95
96 hidden = ops_fn::layer_norm(&hidden, &self.encoder_norm, None, self.config.layer_norm_eps)?;
97 let logits = ops_fn::matmul(&hidden, &self.lm_head)?;
98
99 Ok(ModelOutputs::Logits { logits, hidden_states: Some(hidden) })
100 }
101 _ => Err(anyhow::anyhow!("HuBERT requires audio input")),
102 }
103 }
104
105 fn generate(&self, _prompt: &str, _config: &GenerationConfig) -> Result<String> {
106 Err(anyhow::anyhow!("HuBERT is an encoder model"))
107 }
108
109 fn config(&self) -> &Self::Config { &self.config }
110
111 fn memory_requirements(&self) -> MemoryRequirements {
112 let p = (self.config.hidden_size * self.config.hidden_size * 4 * self.config.num_hidden_layers) * 4;
113 MemoryRequirements { gpu_memory: p, cpu_memory: p / 4, kv_cache_memory: 0, peak_memory: p * 2 }
114 }
115
116 fn to_device(&mut self, device: &Device) -> Result<()> {
117 self.device = device.clone();
118 self.feature_projection = self.feature_projection.to_device(device)?;
119 self.encoder_norm = self.encoder_norm.to_device(device)?;
120 self.lm_head = self.lm_head.to_device(device)?;
121 Ok(())
122 }
123}
124
125impl HuBERTEncoderLayer {
126 fn new(config: &HuBERTConfig, device: &Device) -> Result<Self> {
127 let head_dim = config.hidden_size / config.num_attention_heads;
128
129 Ok(Self {
130 self_attn_q: ops_fn::zeros(&[config.hidden_size, config.hidden_size], DataType::Float32, device)?,
131 self_attn_k: ops_fn::zeros(&[config.hidden_size, config.hidden_size], DataType::Float32, device)?,
132 self_attn_v: ops_fn::zeros(&[config.hidden_size, config.hidden_size], DataType::Float32, device)?,
133 self_attn_o: ops_fn::zeros(&[config.hidden_size, config.hidden_size], DataType::Float32, device)?,
134 fc1: ops_fn::zeros(&[config.hidden_size, config.intermediate_size], DataType::Float32, device)?,
135 fc2: ops_fn::zeros(&[config.intermediate_size, config.hidden_size], DataType::Float32, device)?,
136 layer_norm: ops_fn::zeros(&[config.hidden_size], DataType::Float32, device)?,
137 final_layer_norm: ops_fn::zeros(&[config.hidden_size], DataType::Float32, device)?,
138 num_heads: config.num_attention_heads,
139 head_dim,
140 })
141 }
142
143 fn forward(&self, hidden_states: &Tensor) -> Result<Tensor> {
144 let shape = hidden_states.shape();
145 let (batch_size, seq_len, _) = (shape[0], shape[1], shape[2]);
146
147 let residual = hidden_states.clone();
148 let hidden = ops_fn::layer_norm(hidden_states, &self.layer_norm, None, 1e-5)?;
149
150 let q = ops_fn::matmul(&hidden, &self.self_attn_q)?.to_candle()?
151 .reshape(&[batch_size, seq_len, self.num_heads, self.head_dim])?.transpose(1, 2)?;
152 let k = ops_fn::matmul(&hidden, &self.self_attn_k)?.to_candle()?
153 .reshape(&[batch_size, seq_len, self.num_heads, self.head_dim])?.transpose(1, 2)?;
154 let v = ops_fn::matmul(&hidden, &self.self_attn_v)?.to_candle()?
155 .reshape(&[batch_size, seq_len, self.num_heads, self.head_dim])?.transpose(1, 2)?;
156
157 let scale = (self.head_dim as f32).powf(-0.5);
158 let scores = (q.contiguous()?.matmul(&k.transpose(2, 3)?.contiguous()?)? * (scale as f64))?;
159 let attn = candle_nn::ops::softmax_last_dim(&scores)?.matmul(&v.contiguous()?)?
160 .transpose(1, 2)?.reshape(&[batch_size, seq_len, self.num_heads * self.head_dim])?;
161
162 let hidden = ops_fn::add(&residual, &ops_fn::matmul(&Tensor::from_candle(attn), &self.self_attn_o)?)?;
163
164 let residual = hidden.clone();
165 let hidden = ops_fn::layer_norm(&hidden, &self.final_layer_norm, None, 1e-5)?;
166 let hidden = ops_fn::gelu(&ops_fn::matmul(&hidden, &self.fc1)?)?;
167 ops_fn::add(&residual, &ops_fn::matmul(&hidden, &self.fc2)?)
168 }
169}
170
171#[cfg(test)]
172mod tests {
173 use super::*;
174
175 #[test]
176 fn test_hubert_config() {
177 let config = HuBERTConfig::default();
178 assert_eq!(config.hidden_size, 768);
179 }
180
181 #[test]
182 fn test_hubert_model_creation() {
183 let config = HuBERTConfig {
184 hidden_size: 64,
185 intermediate_size: 256,
186 num_hidden_layers: 2,
187 num_attention_heads: 4,
188 ..Default::default()
189 };
190
191 let model = HuBERTModelV2::new(config).unwrap();
192 assert_eq!(model.config().hidden_size(), 64);
193 }
194}