Skip to main content

runtime/models_v2/
bert.rs

1//! BERT Model V2 - Clean implementation using solid abstractions
2//!
3//! BERT is an encoder-only model with key differences from decoder-only models:
4//! - Returns embeddings (ModelOutputs::Embeddings) instead of logits
5//! - Bidirectional attention (NO causal masking)
6//! - Uses position embeddings (learned) instead of RoPE
7//! - Uses token type embeddings for segment distinction
8//! - Uses LayerNorm (not RMSNorm)
9//! - Has a pooler for [CLS] token representation
10
11use crate::model_config;
12use super::traits::*;
13use anyhow::Result;
14use serde::{Serialize, Deserialize};
15
16model_config!(BertConfig {
17    vocab_size: usize = 30522,
18    hidden_size: usize = 768,
19    intermediate_size: usize = 3072,
20    num_hidden_layers: usize = 12,
21    num_attention_heads: usize = 12,
22    hidden_act: String = "gelu".to_string(),
23    max_position_embeddings: usize = 512,
24    initializer_range: f32 = 0.02,
25    layer_norm_eps: f32 = 1e-12,
26    pad_token_id: i64 = 0,
27    // BERT specific
28    type_vocab_size: usize = 2,
29    // For ModelConfig trait compatibility
30    rms_norm_eps: f32 = 1e-12,
31});
32
33impl BertConfig {
34    /// Create BertConfig from GGUF model configuration
35    pub fn from_gguf_config(gguf: &crate::weight_loader_core::GGUFModelConfig) -> Self {
36        Self {
37            vocab_size: gguf.vocab_size,
38            hidden_size: gguf.hidden_size,
39            intermediate_size: gguf.intermediate_size,
40            num_hidden_layers: gguf.num_hidden_layers,
41            num_attention_heads: gguf.num_attention_heads,
42            max_position_embeddings: gguf.max_position_embeddings,
43            layer_norm_eps: gguf.rms_norm_eps,  // GGUF uses rms_norm_eps field
44            ..Default::default()
45        }
46    }
47}
48
49/// Main BERT model implementation
50pub struct BertModelV2 {
51    config: BertConfig,
52    device: Device,
53    embeddings: BertEmbeddings,
54    encoder: BertEncoder,
55    pooler: Option<BertPooler>,
56}
57
58/// BERT embeddings: word + position + token_type embeddings, then LayerNorm
59pub struct BertEmbeddings {
60    word_embeddings: Tensor,
61    position_embeddings: Tensor,
62    token_type_embeddings: Tensor,
63    layer_norm_weight: Tensor,
64    layer_norm_bias: Tensor,
65    config: BertConfig,
66}
67
68/// BERT encoder: stack of BertLayer
69pub struct BertEncoder {
70    layers: Vec<BertLayer>,
71    #[allow(dead_code)]
72    config: BertConfig,
73}
74
75/// BERT transformer layer
76pub struct BertLayer {
77    attention: BertAttention,
78    intermediate: BertIntermediate,
79    output: BertOutput,
80}
81
82/// BERT attention: self-attention + output projection with residual
83pub struct BertAttention {
84    self_attention: BertSelfAttention,
85    output: BertSelfOutput,
86}
87
88/// BERT self-attention (bidirectional, no causal mask)
89pub struct BertSelfAttention {
90    query: Tensor,
91    query_bias: Tensor,
92    key: Tensor,
93    key_bias: Tensor,
94    value: Tensor,
95    value_bias: Tensor,
96    num_attention_heads: usize,
97    head_dim: usize,
98    scale: f32,
99}
100
101/// BERT self-attention output projection
102pub struct BertSelfOutput {
103    dense: Tensor,
104    dense_bias: Tensor,
105    layer_norm_weight: Tensor,
106    layer_norm_bias: Tensor,
107    layer_norm_eps: f32,
108}
109
110/// BERT intermediate (first FFN layer with activation)
111pub struct BertIntermediate {
112    dense: Tensor,
113    dense_bias: Tensor,
114    hidden_act: String,
115}
116
117/// BERT output (second FFN layer with residual and LayerNorm)
118pub struct BertOutput {
119    dense: Tensor,
120    dense_bias: Tensor,
121    layer_norm_weight: Tensor,
122    layer_norm_bias: Tensor,
123    layer_norm_eps: f32,
124}
125
126/// BERT pooler: takes [CLS] token and projects through dense + tanh
127pub struct BertPooler {
128    dense: Tensor,
129    dense_bias: Tensor,
130}
131
132impl Model for BertModelV2 {
133    type Config = BertConfig;
134
135    fn new(config: BertConfig) -> Result<Self> {
136        let device = Device::CPU;
137        Ok(Self {
138            embeddings: BertEmbeddings::new(&config, &device)?,
139            encoder: BertEncoder::new(&config, &device)?,
140            pooler: Some(BertPooler::new(&config, &device)?),
141            config,
142            device,
143        })
144    }
145
146    fn from_weights(config: BertConfig, weights: ModelWeights) -> Result<Self> {
147        let mut model = Self::new(config)?;
148        model.embeddings.load_weights(&weights)?;
149        model.encoder.load_weights(&weights)?;
150        if let Some(ref mut pooler) = model.pooler {
151            pooler.load_weights(&weights)?;
152        }
153        Ok(model)
154    }
155
156    fn forward(&self, inputs: &ModelInputs) -> Result<ModelOutputs> {
157        let (input_ids, attention_mask, token_type_ids) = match inputs {
158            ModelInputs::Text { input_ids, attention_mask, position_ids: _ } => {
159                (input_ids, attention_mask.as_ref(), None)
160            }
161            _ => return Err(anyhow::anyhow!("BERT expects text input")),
162        };
163
164        // 1. Embeddings: word + position + token_type, then LayerNorm
165        let embedding_output = self.embeddings.forward(input_ids, token_type_ids)?;
166
167        // 2. Encoder: stack of transformer layers with bidirectional attention
168        let encoder_outputs = self.encoder.forward(&embedding_output, attention_mask)?;
169
170        // 3. Pooler: extract [CLS] token representation
171        let pooled_output = if let Some(ref pooler) = self.pooler {
172            Some(pooler.forward(&encoder_outputs)?)
173        } else {
174            None
175        };
176
177        Ok(ModelOutputs::Embeddings {
178            embeddings: encoder_outputs,
179            pooled: pooled_output,
180        })
181    }
182
183    fn generate(&self, _prompt: &str, _config: &GenerationConfig) -> Result<String> {
184        Err(anyhow::anyhow!(
185            "BERT is an encoder-only model and cannot generate text. \
186             Use forward() to get embeddings instead."
187        ))
188    }
189
190    fn config(&self) -> &Self::Config {
191        &self.config
192    }
193
194    fn memory_requirements(&self) -> MemoryRequirements {
195        let param_size = (self.config.vocab_size * self.config.hidden_size
196            + self.config.max_position_embeddings * self.config.hidden_size
197            + self.config.type_vocab_size * self.config.hidden_size
198            + self.config.num_hidden_layers * self.config.hidden_size * self.config.hidden_size * 4
199            + self.config.num_hidden_layers * self.config.hidden_size * self.config.intermediate_size * 2)
200            * 4;
201        MemoryRequirements {
202            gpu_memory: param_size,
203            cpu_memory: param_size / 4,
204            kv_cache_memory: 0,  // BERT doesn't use KV cache (no autoregressive generation)
205            peak_memory: param_size + param_size / 2,
206        }
207    }
208
209    fn to_device(&mut self, device: &Device) -> Result<()> {
210        self.device = device.clone();
211        self.embeddings.to_device(device)?;
212        self.encoder.to_device(device)?;
213        if let Some(ref mut pooler) = self.pooler {
214            pooler.to_device(device)?;
215        }
216        Ok(())
217    }
218}
219
220impl BertEmbeddings {
221    fn new(config: &BertConfig, device: &Device) -> Result<Self> {
222        Ok(Self {
223            word_embeddings: ops_fn::zeros(
224                &[config.vocab_size, config.hidden_size],
225                DataType::Float32,
226                device,
227            )?,
228            position_embeddings: ops_fn::zeros(
229                &[config.max_position_embeddings, config.hidden_size],
230                DataType::Float32,
231                device,
232            )?,
233            token_type_embeddings: ops_fn::zeros(
234                &[config.type_vocab_size, config.hidden_size],
235                DataType::Float32,
236                device,
237            )?,
238            layer_norm_weight: ops_fn::zeros(&[config.hidden_size], DataType::Float32, device)?,
239            layer_norm_bias: ops_fn::zeros(&[config.hidden_size], DataType::Float32, device)?,
240            config: config.clone(),
241        })
242    }
243
244    fn forward(&self, input_ids: &Tensor, token_type_ids: Option<&Tensor>) -> Result<Tensor> {
245        let shape = input_ids.shape();
246        let seq_len = if shape.len() == 2 { shape[1] } else { shape[0] };
247
248        // 1. Word embeddings
249        let inputs_embeds = ops_fn::embedding(input_ids, &self.word_embeddings)?;
250
251        // 2. Position embeddings (learned, absolute)
252        // Create position_ids: [0, 1, 2, ..., seq_len-1]
253        let position_ids: Vec<i64> = (0..seq_len as i64).collect();
254        let batch_size = if shape.len() == 2 { shape[0] } else { 1 };
255        let position_ids_expanded: Vec<i64> = (0..batch_size)
256            .flat_map(|_| position_ids.iter().cloned())
257            .collect();
258        let position_ids_tensor = Tensor::from_i64_slice(
259            &position_ids_expanded,
260            &[batch_size, seq_len],
261            input_ids.device(),
262        )?;
263        let position_embeds = ops_fn::embedding(&position_ids_tensor, &self.position_embeddings)?;
264
265        // 3. Token type embeddings (segment embeddings)
266        let token_type_embeds = if let Some(tt_ids) = token_type_ids {
267            ops_fn::embedding(tt_ids, &self.token_type_embeddings)?
268        } else {
269            // Default to all zeros (segment A)
270            let zeros: Vec<i64> = vec![0; batch_size * seq_len];
271            let tt_tensor = Tensor::from_i64_slice(&zeros, &[batch_size, seq_len], input_ids.device())?;
272            ops_fn::embedding(&tt_tensor, &self.token_type_embeddings)?
273        };
274
275        // 4. Sum all embeddings
276        let embeddings = ops_fn::add(&inputs_embeds, &position_embeds)?;
277        let embeddings = ops_fn::add(&embeddings, &token_type_embeds)?;
278
279        // 5. Layer normalization with bias
280        self.layer_norm(&embeddings)
281    }
282
283    fn layer_norm(&self, input: &Tensor) -> Result<Tensor> {
284        let x = input.to_candle()?;
285        let w = self.layer_norm_weight.to_candle()?;
286        let b = self.layer_norm_bias.to_candle()?;
287
288        let last_dim = x.dims().len() - 1;
289        let mean = x.mean_keepdim(last_dim)?;
290        let x_centered = x.broadcast_sub(&mean)?;
291        let variance = x_centered.sqr()?.mean_keepdim(last_dim)?;
292        let std = (variance + self.config.layer_norm_eps as f64)?.sqrt()?;
293        let normalized = x_centered.broadcast_div(&std)?;
294        let scaled = normalized.broadcast_mul(&w)?;
295        let result = scaled.broadcast_add(&b)?;
296
297        Ok(Tensor::from_candle(result))
298    }
299
300    fn load_weights(&mut self, weights: &ModelWeights) -> Result<()> {
301        // Try different weight naming conventions
302        let word_keys = [
303            "bert.embeddings.word_embeddings.weight",
304            "embeddings.word_embeddings.weight",
305        ];
306        for key in word_keys {
307            if let Some(w) = weights.get(key) {
308                self.word_embeddings = w.clone();
309                break;
310            }
311        }
312
313        let pos_keys = [
314            "bert.embeddings.position_embeddings.weight",
315            "embeddings.position_embeddings.weight",
316        ];
317        for key in pos_keys {
318            if let Some(w) = weights.get(key) {
319                self.position_embeddings = w.clone();
320                break;
321            }
322        }
323
324        let tt_keys = [
325            "bert.embeddings.token_type_embeddings.weight",
326            "embeddings.token_type_embeddings.weight",
327        ];
328        for key in tt_keys {
329            if let Some(w) = weights.get(key) {
330                self.token_type_embeddings = w.clone();
331                break;
332            }
333        }
334
335        let ln_weight_keys = [
336            "bert.embeddings.LayerNorm.weight",
337            "embeddings.LayerNorm.weight",
338            "bert.embeddings.LayerNorm.gamma",
339        ];
340        for key in ln_weight_keys {
341            if let Some(w) = weights.get(key) {
342                self.layer_norm_weight = w.clone();
343                break;
344            }
345        }
346
347        let ln_bias_keys = [
348            "bert.embeddings.LayerNorm.bias",
349            "embeddings.LayerNorm.bias",
350            "bert.embeddings.LayerNorm.beta",
351        ];
352        for key in ln_bias_keys {
353            if let Some(w) = weights.get(key) {
354                self.layer_norm_bias = w.clone();
355                break;
356            }
357        }
358
359        Ok(())
360    }
361
362    fn to_device(&mut self, device: &Device) -> Result<()> {
363        self.word_embeddings = self.word_embeddings.to_device(device)?;
364        self.position_embeddings = self.position_embeddings.to_device(device)?;
365        self.token_type_embeddings = self.token_type_embeddings.to_device(device)?;
366        self.layer_norm_weight = self.layer_norm_weight.to_device(device)?;
367        self.layer_norm_bias = self.layer_norm_bias.to_device(device)?;
368        Ok(())
369    }
370}
371
372impl BertEncoder {
373    fn new(config: &BertConfig, device: &Device) -> Result<Self> {
374        let mut layers = Vec::new();
375        for _ in 0..config.num_hidden_layers {
376            layers.push(BertLayer::new(config, device)?);
377        }
378        Ok(Self {
379            layers,
380            config: config.clone(),
381        })
382    }
383
384    fn forward(&self, hidden_states: &Tensor, attention_mask: Option<&Tensor>) -> Result<Tensor> {
385        let mut hidden_states = hidden_states.clone();
386        for layer in &self.layers {
387            hidden_states = layer.forward(&hidden_states, attention_mask)?;
388        }
389        Ok(hidden_states)
390    }
391
392    fn load_weights(&mut self, weights: &ModelWeights) -> Result<()> {
393        for (i, layer) in self.layers.iter_mut().enumerate() {
394            layer.load_weights(weights, i)?;
395        }
396        Ok(())
397    }
398
399    fn to_device(&mut self, device: &Device) -> Result<()> {
400        for layer in &mut self.layers {
401            layer.to_device(device)?;
402        }
403        Ok(())
404    }
405}
406
407impl BertLayer {
408    fn new(config: &BertConfig, device: &Device) -> Result<Self> {
409        Ok(Self {
410            attention: BertAttention::new(config, device)?,
411            intermediate: BertIntermediate::new(config, device)?,
412            output: BertOutput::new(config, device)?,
413        })
414    }
415
416    fn forward(&self, hidden_states: &Tensor, attention_mask: Option<&Tensor>) -> Result<Tensor> {
417        // Self-attention with residual
418        let attention_output = self.attention.forward(hidden_states, attention_mask)?;
419        // FFN with residual
420        let intermediate_output = self.intermediate.forward(&attention_output)?;
421        self.output.forward(&intermediate_output, &attention_output)
422    }
423
424    fn load_weights(&mut self, weights: &ModelWeights, layer_idx: usize) -> Result<()> {
425        self.attention.load_weights(weights, layer_idx)?;
426        self.intermediate.load_weights(weights, layer_idx)?;
427        self.output.load_weights(weights, layer_idx)?;
428        Ok(())
429    }
430
431    fn to_device(&mut self, device: &Device) -> Result<()> {
432        self.attention.to_device(device)?;
433        self.intermediate.to_device(device)?;
434        self.output.to_device(device)?;
435        Ok(())
436    }
437}
438
439impl BertAttention {
440    fn new(config: &BertConfig, device: &Device) -> Result<Self> {
441        Ok(Self {
442            self_attention: BertSelfAttention::new(config, device)?,
443            output: BertSelfOutput::new(config, device)?,
444        })
445    }
446
447    fn forward(&self, hidden_states: &Tensor, attention_mask: Option<&Tensor>) -> Result<Tensor> {
448        let self_output = self.self_attention.forward(hidden_states, attention_mask)?;
449        self.output.forward(&self_output, hidden_states)
450    }
451
452    fn load_weights(&mut self, weights: &ModelWeights, layer_idx: usize) -> Result<()> {
453        self.self_attention.load_weights(weights, layer_idx)?;
454        self.output.load_weights(weights, layer_idx)?;
455        Ok(())
456    }
457
458    fn to_device(&mut self, device: &Device) -> Result<()> {
459        self.self_attention.to_device(device)?;
460        self.output.to_device(device)?;
461        Ok(())
462    }
463}
464
465impl BertSelfAttention {
466    fn new(config: &BertConfig, device: &Device) -> Result<Self> {
467        let head_dim = config.hidden_size / config.num_attention_heads;
468        let scale = 1.0 / (head_dim as f32).sqrt();
469
470        Ok(Self {
471            query: ops_fn::zeros(
472                &[config.hidden_size, config.hidden_size],
473                DataType::Float32,
474                device,
475            )?,
476            query_bias: ops_fn::zeros(&[config.hidden_size], DataType::Float32, device)?,
477            key: ops_fn::zeros(
478                &[config.hidden_size, config.hidden_size],
479                DataType::Float32,
480                device,
481            )?,
482            key_bias: ops_fn::zeros(&[config.hidden_size], DataType::Float32, device)?,
483            value: ops_fn::zeros(
484                &[config.hidden_size, config.hidden_size],
485                DataType::Float32,
486                device,
487            )?,
488            value_bias: ops_fn::zeros(&[config.hidden_size], DataType::Float32, device)?,
489            num_attention_heads: config.num_attention_heads,
490            head_dim,
491            scale,
492        })
493    }
494
495    fn forward(&self, hidden_states: &Tensor, attention_mask: Option<&Tensor>) -> Result<Tensor> {
496        let shape = hidden_states.shape();
497        let (batch_size, seq_len, _hidden_size) = if shape.len() == 3 {
498            (shape[0], shape[1], shape[2])
499        } else if shape.len() == 2 {
500            (1, shape[0], shape[1])
501        } else {
502            return Err(anyhow::anyhow!("Invalid hidden_states shape: {:?}", shape));
503        };
504
505        // 1. Project to Q, K, V with bias
506        let query_states = self.linear_with_bias(hidden_states, &self.query, &self.query_bias)?;
507        let key_states = self.linear_with_bias(hidden_states, &self.key, &self.key_bias)?;
508        let value_states = self.linear_with_bias(hidden_states, &self.value, &self.value_bias)?;
509
510        // 2. Reshape for multi-head attention
511        let q_candle = query_states.to_candle()?;
512        let k_candle = key_states.to_candle()?;
513        let v_candle = value_states.to_candle()?;
514
515        // [batch, seq, hidden] -> [batch, seq, heads, head_dim] -> [batch, heads, seq, head_dim]
516        let q_reshaped = q_candle
517            .reshape(&[batch_size, seq_len, self.num_attention_heads, self.head_dim])?
518            .transpose(1, 2)?;
519
520        let k_reshaped = k_candle
521            .reshape(&[batch_size, seq_len, self.num_attention_heads, self.head_dim])?
522            .transpose(1, 2)?;
523
524        let v_reshaped = v_candle
525            .reshape(&[batch_size, seq_len, self.num_attention_heads, self.head_dim])?
526            .transpose(1, 2)?;
527
528        // 3. Scaled dot-product attention (BIDIRECTIONAL - no causal mask!)
529        let k_t = k_reshaped.transpose(2, 3)?;
530        let q_contiguous = q_reshaped.contiguous()?;
531        let k_contiguous = k_t.contiguous()?;
532
533        let scores = q_contiguous.matmul(&k_contiguous)?;
534        let scaled_scores = (scores * (self.scale as f64))?;
535
536        // Apply attention mask if provided (for padding tokens)
537        let masked_scores = if let Some(mask) = attention_mask {
538            let mask_candle = mask.to_candle()?;
539            // Expand mask: [batch, seq] -> [batch, 1, 1, seq] for broadcasting
540            let mask_expanded = if mask_candle.dims().len() == 2 {
541                mask_candle.unsqueeze(1)?.unsqueeze(1)?
542            } else {
543                mask_candle
544            };
545            // Convert mask: 1 -> 0, 0 -> -inf
546            let mask_f32 = mask_expanded.to_dtype(candle_core::DType::F32)?;
547            let inverted_mask = ((1.0 - &mask_f32)? * f32::NEG_INFINITY as f64)?;
548            scaled_scores.broadcast_add(&inverted_mask)?
549        } else {
550            scaled_scores
551        };
552
553        // Softmax
554        let attention_weights = candle_nn::ops::softmax_last_dim(&masked_scores)?;
555
556        // Apply attention to values
557        let v_contiguous = v_reshaped.contiguous()?;
558        let attn_output = attention_weights.matmul(&v_contiguous)?;
559
560        // 4. Reshape back: [batch, heads, seq, head_dim] -> [batch, seq, hidden]
561        let attn_output = attn_output
562            .transpose(1, 2)?
563            .reshape(&[batch_size, seq_len, self.num_attention_heads * self.head_dim])?;
564
565        Ok(Tensor::from_candle(attn_output))
566    }
567
568    fn linear_with_bias(&self, input: &Tensor, weight: &Tensor, bias: &Tensor) -> Result<Tensor> {
569        let output = ops_fn::matmul(input, weight)?;
570        ops_fn::add(&output, bias)
571    }
572
573    fn load_weights(&mut self, weights: &ModelWeights, layer_idx: usize) -> Result<()> {
574        let prefixes = [
575            format!("bert.encoder.layer.{}.attention.self", layer_idx),
576            format!("encoder.layer.{}.attention.self", layer_idx),
577        ];
578
579        for prefix in &prefixes {
580            // Transpose weights: [out, in] -> [in, out] for matmul
581            if let Some(w) = weights.get(&format!("{}.query.weight", prefix)) {
582                self.query = ops_fn::transpose(w)?;
583            }
584            if let Some(b) = weights.get(&format!("{}.query.bias", prefix)) {
585                self.query_bias = b.clone();
586            }
587            if let Some(w) = weights.get(&format!("{}.key.weight", prefix)) {
588                self.key = ops_fn::transpose(w)?;
589            }
590            if let Some(b) = weights.get(&format!("{}.key.bias", prefix)) {
591                self.key_bias = b.clone();
592            }
593            if let Some(w) = weights.get(&format!("{}.value.weight", prefix)) {
594                self.value = ops_fn::transpose(w)?;
595            }
596            if let Some(b) = weights.get(&format!("{}.value.bias", prefix)) {
597                self.value_bias = b.clone();
598            }
599        }
600
601        Ok(())
602    }
603
604    fn to_device(&mut self, device: &Device) -> Result<()> {
605        self.query = self.query.to_device(device)?;
606        self.query_bias = self.query_bias.to_device(device)?;
607        self.key = self.key.to_device(device)?;
608        self.key_bias = self.key_bias.to_device(device)?;
609        self.value = self.value.to_device(device)?;
610        self.value_bias = self.value_bias.to_device(device)?;
611        Ok(())
612    }
613}
614
615impl BertSelfOutput {
616    fn new(config: &BertConfig, device: &Device) -> Result<Self> {
617        Ok(Self {
618            dense: ops_fn::zeros(
619                &[config.hidden_size, config.hidden_size],
620                DataType::Float32,
621                device,
622            )?,
623            dense_bias: ops_fn::zeros(&[config.hidden_size], DataType::Float32, device)?,
624            layer_norm_weight: ops_fn::zeros(&[config.hidden_size], DataType::Float32, device)?,
625            layer_norm_bias: ops_fn::zeros(&[config.hidden_size], DataType::Float32, device)?,
626            layer_norm_eps: config.layer_norm_eps,
627        })
628    }
629
630    fn forward(&self, hidden_states: &Tensor, input_tensor: &Tensor) -> Result<Tensor> {
631        // Dense projection with bias
632        let hidden_states = ops_fn::matmul(hidden_states, &self.dense)?;
633        let hidden_states = ops_fn::add(&hidden_states, &self.dense_bias)?;
634        // Residual connection
635        let hidden_states = ops_fn::add(&hidden_states, input_tensor)?;
636        // Layer normalization with bias
637        self.layer_norm(&hidden_states)
638    }
639
640    fn layer_norm(&self, input: &Tensor) -> Result<Tensor> {
641        let x = input.to_candle()?;
642        let w = self.layer_norm_weight.to_candle()?;
643        let b = self.layer_norm_bias.to_candle()?;
644
645        let last_dim = x.dims().len() - 1;
646        let mean = x.mean_keepdim(last_dim)?;
647        let x_centered = x.broadcast_sub(&mean)?;
648        let variance = x_centered.sqr()?.mean_keepdim(last_dim)?;
649        let std = (variance + self.layer_norm_eps as f64)?.sqrt()?;
650        let normalized = x_centered.broadcast_div(&std)?;
651        let scaled = normalized.broadcast_mul(&w)?;
652        let result = scaled.broadcast_add(&b)?;
653
654        Ok(Tensor::from_candle(result))
655    }
656
657    fn load_weights(&mut self, weights: &ModelWeights, layer_idx: usize) -> Result<()> {
658        let prefixes = [
659            format!("bert.encoder.layer.{}.attention.output", layer_idx),
660            format!("encoder.layer.{}.attention.output", layer_idx),
661        ];
662
663        for prefix in &prefixes {
664            if let Some(w) = weights.get(&format!("{}.dense.weight", prefix)) {
665                self.dense = ops_fn::transpose(w)?;
666            }
667            if let Some(b) = weights.get(&format!("{}.dense.bias", prefix)) {
668                self.dense_bias = b.clone();
669            }
670            if let Some(w) = weights.get(&format!("{}.LayerNorm.weight", prefix)) {
671                self.layer_norm_weight = w.clone();
672            }
673            if let Some(b) = weights.get(&format!("{}.LayerNorm.bias", prefix)) {
674                self.layer_norm_bias = b.clone();
675            }
676        }
677
678        Ok(())
679    }
680
681    fn to_device(&mut self, device: &Device) -> Result<()> {
682        self.dense = self.dense.to_device(device)?;
683        self.dense_bias = self.dense_bias.to_device(device)?;
684        self.layer_norm_weight = self.layer_norm_weight.to_device(device)?;
685        self.layer_norm_bias = self.layer_norm_bias.to_device(device)?;
686        Ok(())
687    }
688}
689
690impl BertIntermediate {
691    fn new(config: &BertConfig, device: &Device) -> Result<Self> {
692        Ok(Self {
693            dense: ops_fn::zeros(
694                &[config.hidden_size, config.intermediate_size],
695                DataType::Float32,
696                device,
697            )?,
698            dense_bias: ops_fn::zeros(&[config.intermediate_size], DataType::Float32, device)?,
699            hidden_act: config.hidden_act.clone(),
700        })
701    }
702
703    fn forward(&self, hidden_states: &Tensor) -> Result<Tensor> {
704        let hidden_states = ops_fn::matmul(hidden_states, &self.dense)?;
705        let hidden_states = ops_fn::add(&hidden_states, &self.dense_bias)?;
706        // BERT uses GELU activation
707        match self.hidden_act.as_str() {
708            "gelu" | "gelu_new" => ops_fn::gelu(&hidden_states),
709            "relu" => {
710                let x = hidden_states.to_candle()?;
711                let result = x.relu()?;
712                Ok(Tensor::from_candle(result))
713            }
714            _ => ops_fn::gelu(&hidden_states),  // Default to GELU
715        }
716    }
717
718    fn load_weights(&mut self, weights: &ModelWeights, layer_idx: usize) -> Result<()> {
719        let prefixes = [
720            format!("bert.encoder.layer.{}.intermediate", layer_idx),
721            format!("encoder.layer.{}.intermediate", layer_idx),
722        ];
723
724        for prefix in &prefixes {
725            if let Some(w) = weights.get(&format!("{}.dense.weight", prefix)) {
726                self.dense = ops_fn::transpose(w)?;
727            }
728            if let Some(b) = weights.get(&format!("{}.dense.bias", prefix)) {
729                self.dense_bias = b.clone();
730            }
731        }
732
733        Ok(())
734    }
735
736    fn to_device(&mut self, device: &Device) -> Result<()> {
737        self.dense = self.dense.to_device(device)?;
738        self.dense_bias = self.dense_bias.to_device(device)?;
739        Ok(())
740    }
741}
742
743impl BertOutput {
744    fn new(config: &BertConfig, device: &Device) -> Result<Self> {
745        Ok(Self {
746            dense: ops_fn::zeros(
747                &[config.intermediate_size, config.hidden_size],
748                DataType::Float32,
749                device,
750            )?,
751            dense_bias: ops_fn::zeros(&[config.hidden_size], DataType::Float32, device)?,
752            layer_norm_weight: ops_fn::zeros(&[config.hidden_size], DataType::Float32, device)?,
753            layer_norm_bias: ops_fn::zeros(&[config.hidden_size], DataType::Float32, device)?,
754            layer_norm_eps: config.layer_norm_eps,
755        })
756    }
757
758    fn forward(&self, hidden_states: &Tensor, input_tensor: &Tensor) -> Result<Tensor> {
759        // Dense projection with bias
760        let hidden_states = ops_fn::matmul(hidden_states, &self.dense)?;
761        let hidden_states = ops_fn::add(&hidden_states, &self.dense_bias)?;
762        // Residual connection
763        let hidden_states = ops_fn::add(&hidden_states, input_tensor)?;
764        // Layer normalization with bias
765        self.layer_norm(&hidden_states)
766    }
767
768    fn layer_norm(&self, input: &Tensor) -> Result<Tensor> {
769        let x = input.to_candle()?;
770        let w = self.layer_norm_weight.to_candle()?;
771        let b = self.layer_norm_bias.to_candle()?;
772
773        let last_dim = x.dims().len() - 1;
774        let mean = x.mean_keepdim(last_dim)?;
775        let x_centered = x.broadcast_sub(&mean)?;
776        let variance = x_centered.sqr()?.mean_keepdim(last_dim)?;
777        let std = (variance + self.layer_norm_eps as f64)?.sqrt()?;
778        let normalized = x_centered.broadcast_div(&std)?;
779        let scaled = normalized.broadcast_mul(&w)?;
780        let result = scaled.broadcast_add(&b)?;
781
782        Ok(Tensor::from_candle(result))
783    }
784
785    fn load_weights(&mut self, weights: &ModelWeights, layer_idx: usize) -> Result<()> {
786        let prefixes = [
787            format!("bert.encoder.layer.{}.output", layer_idx),
788            format!("encoder.layer.{}.output", layer_idx),
789        ];
790
791        for prefix in &prefixes {
792            if let Some(w) = weights.get(&format!("{}.dense.weight", prefix)) {
793                self.dense = ops_fn::transpose(w)?;
794            }
795            if let Some(b) = weights.get(&format!("{}.dense.bias", prefix)) {
796                self.dense_bias = b.clone();
797            }
798            if let Some(w) = weights.get(&format!("{}.LayerNorm.weight", prefix)) {
799                self.layer_norm_weight = w.clone();
800            }
801            if let Some(b) = weights.get(&format!("{}.LayerNorm.bias", prefix)) {
802                self.layer_norm_bias = b.clone();
803            }
804        }
805
806        Ok(())
807    }
808
809    fn to_device(&mut self, device: &Device) -> Result<()> {
810        self.dense = self.dense.to_device(device)?;
811        self.dense_bias = self.dense_bias.to_device(device)?;
812        self.layer_norm_weight = self.layer_norm_weight.to_device(device)?;
813        self.layer_norm_bias = self.layer_norm_bias.to_device(device)?;
814        Ok(())
815    }
816}
817
818impl BertPooler {
819    fn new(config: &BertConfig, device: &Device) -> Result<Self> {
820        Ok(Self {
821            dense: ops_fn::zeros(
822                &[config.hidden_size, config.hidden_size],
823                DataType::Float32,
824                device,
825            )?,
826            dense_bias: ops_fn::zeros(&[config.hidden_size], DataType::Float32, device)?,
827        })
828    }
829
830    fn forward(&self, hidden_states: &Tensor) -> Result<Tensor> {
831        // Take the first token ([CLS]) representation
832        let candle_tensor = hidden_states.to_candle()?;
833        let shape = candle_tensor.dims();
834
835        // Extract [CLS] token: [batch, seq, hidden] -> [batch, hidden]
836        let first_token = if shape.len() == 3 {
837            candle_tensor.narrow(1, 0, 1)?.squeeze(1)?
838        } else {
839            candle_tensor.narrow(0, 0, 1)?.squeeze(0)?
840        };
841
842        let first_token = Tensor::from_candle(first_token);
843
844        // Dense projection with bias
845        let pooled = ops_fn::matmul(&first_token, &self.dense)?;
846        let pooled = ops_fn::add(&pooled, &self.dense_bias)?;
847
848        // Tanh activation
849        let pooled_candle = pooled.to_candle()?;
850        let result = pooled_candle.tanh()?;
851
852        Ok(Tensor::from_candle(result))
853    }
854
855    fn load_weights(&mut self, weights: &ModelWeights) -> Result<()> {
856        let weight_keys = ["bert.pooler.dense.weight", "pooler.dense.weight"];
857        let bias_keys = ["bert.pooler.dense.bias", "pooler.dense.bias"];
858
859        for key in weight_keys {
860            if let Some(w) = weights.get(key) {
861                self.dense = ops_fn::transpose(w)?;
862                break;
863            }
864        }
865
866        for key in bias_keys {
867            if let Some(b) = weights.get(key) {
868                self.dense_bias = b.clone();
869                break;
870            }
871        }
872
873        Ok(())
874    }
875
876    fn to_device(&mut self, device: &Device) -> Result<()> {
877        self.dense = self.dense.to_device(device)?;
878        self.dense_bias = self.dense_bias.to_device(device)?;
879        Ok(())
880    }
881}
882
883#[cfg(test)]
884mod tests {
885    use super::*;
886
887    #[test]
888    fn test_bert_model_creation() {
889        let config = BertConfig {
890            vocab_size: 1000,
891            hidden_size: 128,
892            intermediate_size: 512,
893            num_hidden_layers: 2,
894            num_attention_heads: 4,
895            ..Default::default()
896        };
897
898        let model = BertModelV2::new(config).unwrap();
899        assert_eq!(model.config().vocab_size(), 1000);
900        assert_eq!(model.config().hidden_size(), 128);
901        assert_eq!(model.config().num_layers(), 2);
902    }
903
904    #[test]
905    fn test_bert_forward_pass() {
906        let config = BertConfig {
907            vocab_size: 100,
908            hidden_size: 64,
909            intermediate_size: 256,
910            num_hidden_layers: 1,
911            num_attention_heads: 4,
912            ..Default::default()
913        };
914
915        let model = BertModelV2::new(config).unwrap();
916        let input_ids = ops_fn::zeros(&[2, 8], DataType::Int64, &Device::CPU).unwrap();
917        let inputs = ModelInputs::text(input_ids);
918
919        let outputs = model.forward(&inputs).unwrap();
920        match outputs {
921            ModelOutputs::Embeddings { embeddings, pooled } => {
922                assert_eq!(embeddings.shape(), &[2, 8, 64]); // batch, seq, hidden
923                assert!(pooled.is_some());
924                let pooled = pooled.unwrap();
925                assert_eq!(pooled.shape(), &[2, 64]); // batch, hidden
926            }
927            _ => panic!("Expected embeddings output"),
928        }
929    }
930
931    #[test]
932    fn test_bert_generate_returns_error() {
933        let config = BertConfig {
934            vocab_size: 100,
935            hidden_size: 64,
936            intermediate_size: 256,
937            num_hidden_layers: 1,
938            num_attention_heads: 4,
939            ..Default::default()
940        };
941        let model = BertModelV2::new(config).unwrap();
942        let gen_config = GenerationConfig::default();
943
944        let result = model.generate("Hello", &gen_config);
945        assert!(result.is_err());
946        assert!(result.unwrap_err().to_string().contains("encoder-only"));
947    }
948
949    #[test]
950    fn test_bert_bidirectional_attention() {
951        // BERT should allow attending to all positions (no causal mask)
952        let config = BertConfig {
953            vocab_size: 100,
954            hidden_size: 64,
955            intermediate_size: 256,
956            num_hidden_layers: 1,
957            num_attention_heads: 4,
958            ..Default::default()
959        };
960
961        let model = BertModelV2::new(config).unwrap();
962
963        // Create input with some tokens
964        let input_data: Vec<i64> = vec![1, 2, 3, 4, 5, 6, 7, 8];
965        let input_ids = Tensor::from_i64_slice(&input_data, &[1, 8], &Device::CPU).unwrap();
966        let inputs = ModelInputs::text(input_ids);
967
968        // Forward pass should work with bidirectional attention
969        let outputs = model.forward(&inputs);
970        assert!(outputs.is_ok());
971    }
972}