1use 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 type_vocab_size: usize = 2,
29 rms_norm_eps: f32 = 1e-12,
31});
32
33impl BertConfig {
34 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, ..Default::default()
45 }
46 }
47}
48
49pub struct BertModelV2 {
51 config: BertConfig,
52 device: Device,
53 embeddings: BertEmbeddings,
54 encoder: BertEncoder,
55 pooler: Option<BertPooler>,
56}
57
58pub 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
68pub struct BertEncoder {
70 layers: Vec<BertLayer>,
71 #[allow(dead_code)]
72 config: BertConfig,
73}
74
75pub struct BertLayer {
77 attention: BertAttention,
78 intermediate: BertIntermediate,
79 output: BertOutput,
80}
81
82pub struct BertAttention {
84 self_attention: BertSelfAttention,
85 output: BertSelfOutput,
86}
87
88pub 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
101pub 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
110pub struct BertIntermediate {
112 dense: Tensor,
113 dense_bias: Tensor,
114 hidden_act: String,
115}
116
117pub 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
126pub 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 let embedding_output = self.embeddings.forward(input_ids, token_type_ids)?;
166
167 let encoder_outputs = self.encoder.forward(&embedding_output, attention_mask)?;
169
170 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, 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 let inputs_embeds = ops_fn::embedding(input_ids, &self.word_embeddings)?;
250
251 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 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 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 let embeddings = ops_fn::add(&inputs_embeds, &position_embeds)?;
277 let embeddings = ops_fn::add(&embeddings, &token_type_embeds)?;
278
279 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 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 let attention_output = self.attention.forward(hidden_states, attention_mask)?;
419 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 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 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 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 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 let masked_scores = if let Some(mask) = attention_mask {
538 let mask_candle = mask.to_candle()?;
539 let mask_expanded = if mask_candle.dims().len() == 2 {
541 mask_candle.unsqueeze(1)?.unsqueeze(1)?
542 } else {
543 mask_candle
544 };
545 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 let attention_weights = candle_nn::ops::softmax_last_dim(&masked_scores)?;
555
556 let v_contiguous = v_reshaped.contiguous()?;
558 let attn_output = attention_weights.matmul(&v_contiguous)?;
559
560 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 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 let hidden_states = ops_fn::matmul(hidden_states, &self.dense)?;
633 let hidden_states = ops_fn::add(&hidden_states, &self.dense_bias)?;
634 let hidden_states = ops_fn::add(&hidden_states, input_tensor)?;
636 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 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), }
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 let hidden_states = ops_fn::matmul(hidden_states, &self.dense)?;
761 let hidden_states = ops_fn::add(&hidden_states, &self.dense_bias)?;
762 let hidden_states = ops_fn::add(&hidden_states, input_tensor)?;
764 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 let candle_tensor = hidden_states.to_candle()?;
833 let shape = candle_tensor.dims();
834
835 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 let pooled = ops_fn::matmul(&first_token, &self.dense)?;
846 let pooled = ops_fn::add(&pooled, &self.dense_bias)?;
847
848 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]); assert!(pooled.is_some());
924 let pooled = pooled.unwrap();
925 assert_eq!(pooled.shape(), &[2, 64]); }
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 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 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 let outputs = model.forward(&inputs);
970 assert!(outputs.is_ok());
971 }
972}