Skip to main content

runtime/models_v2/
hubert.rs

1//! HuBERT Model V2 - Hidden-Unit BERT for Self-Supervised Speech
2//!
3//! Similar to Wav2Vec2 but uses masked prediction with discrete targets
4
5use 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                // Feature projection (assuming features already extracted)
88                let hidden = ops_fn::matmul(input_features, &self.feature_projection)?;
89
90                // Encoder layers
91                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}