1use ferrox_core::matmul::{gelu, layer_norm};
64use ferrox_core::weight_matrix::WeightMatrix;
65
66use crate::encoder::{EncodeError, PairSequence, TextEncoder};
67use crate::pooling::PoolingType;
68
69#[derive(Debug, Clone)]
71pub struct BertHparams {
72 pub arch: String,
73 pub n_layer: usize,
74 pub n_embd: usize,
75 pub n_ff: usize,
76 pub n_head: usize,
77 pub n_head_kv: usize,
78 pub n_ctx_train: usize,
80 pub n_token_types: usize,
81 pub layer_norm_eps: f32,
82 pub pooling: PoolingType,
83 pub cls_id: u32,
87 pub sep_id: u32,
88}
89
90impl BertHparams {
91 pub fn head_dim(&self) -> usize {
92 self.n_embd / self.n_head
93 }
94}
95
96pub struct BertLayer {
100 pub wq: WeightMatrix,
101 pub bq: Option<Vec<f32>>,
102 pub wk: WeightMatrix,
103 pub bk: Option<Vec<f32>>,
104 pub wv: WeightMatrix,
105 pub bv: Option<Vec<f32>>,
106 pub wo: WeightMatrix,
107 pub bo: Option<Vec<f32>>,
108 pub attn_out_norm_w: Vec<f32>,
110 pub attn_out_norm_b: Vec<f32>,
111 pub ffn_up: WeightMatrix,
112 pub ffn_up_b: Option<Vec<f32>>,
113 pub ffn_down: WeightMatrix,
114 pub ffn_down_b: Option<Vec<f32>>,
115 pub layer_out_norm_w: Vec<f32>,
117 pub layer_out_norm_b: Vec<f32>,
118}
119
120pub struct BertEncoder {
121 pub hp: BertHparams,
122 pub tok_embd: WeightMatrix,
123 pub type_embd: Option<Vec<Vec<f32>>>,
130 pub pos_embd: WeightMatrix,
131 pub tok_norm_w: Vec<f32>,
132 pub tok_norm_b: Vec<f32>,
133 pub layers: Vec<BertLayer>,
134}
135
136fn add_bias_rows(rows: &mut [f32], width: usize, bias: Option<&Vec<f32>>) {
138 let Some(b) = bias else { return };
139 debug_assert_eq!(b.len(), width);
140 for row in rows.chunks_exact_mut(width) {
141 for (x, bv) in row.iter_mut().zip(b.iter()) {
142 *x += bv;
143 }
144 }
145}
146
147fn layer_norm_rows(rows: &mut [f32], width: usize, weight: &[f32], bias: &[f32], eps: f32) {
149 for row in rows.chunks_exact_mut(width) {
150 let normed = layer_norm(row, weight, bias, eps);
151 row.copy_from_slice(&normed);
152 }
153}
154
155fn softmax_row(scores: &mut [f32]) {
157 let max = scores.iter().copied().fold(f32::NEG_INFINITY, f32::max);
158 let mut sum = 0.0f32;
159 for s in scores.iter_mut() {
160 *s = (*s - max).exp();
161 sum += *s;
162 }
163 let inv = 1.0 / sum;
164 for s in scores.iter_mut() {
165 *s *= inv;
166 }
167}
168
169fn bidirectional_attention(
176 q: &[f32],
177 k: &[f32],
178 v: &[f32],
179 n: usize,
180 n_head: usize,
181 n_head_kv: usize,
182 head_dim: usize,
183) -> Vec<f32> {
184 let q_width = n_head * head_dim;
185 let kv_width = n_head_kv * head_dim;
186 let heads_per_kv = n_head / n_head_kv;
187 let scale = 1.0 / (head_dim as f32).sqrt();
188 let mut out = vec![0.0f32; n * q_width];
189 let mut scores = vec![0.0f32; n];
190 for h in 0..n_head {
191 let kv_h = h / heads_per_kv;
192 let q_off = h * head_dim;
193 let kv_off = kv_h * head_dim;
194 for i in 0..n {
195 let qi = &q[i * q_width + q_off..i * q_width + q_off + head_dim];
196 for (j, s) in scores.iter_mut().enumerate() {
197 let kj = &k[j * kv_width + kv_off..j * kv_width + kv_off + head_dim];
198 *s = qi.iter().zip(kj).map(|(a, b)| a * b).sum::<f32>() * scale;
199 }
200 softmax_row(&mut scores);
201 let dst = &mut out[i * q_width + q_off..i * q_width + q_off + head_dim];
202 for (j, &p) in scores.iter().enumerate() {
203 let vj = &v[j * kv_width + kv_off..j * kv_width + kv_off + head_dim];
204 for (o, &vv) in dst.iter_mut().zip(vj) {
205 *o += p * vv;
206 }
207 }
208 }
209 }
210 out
211}
212
213impl BertEncoder {
214 pub fn vocab_size(&self) -> usize {
215 self.tok_embd.rows()
216 }
217}
218
219impl TextEncoder for BertEncoder {
220 fn n_embd(&self) -> usize {
221 self.hp.n_embd
222 }
223
224 fn n_ctx_train(&self) -> usize {
225 self.hp.n_ctx_train
226 }
227
228 fn pooling_type(&self) -> PoolingType {
229 self.hp.pooling
230 }
231
232 fn wrap_special(&self, pieces: &[u32]) -> Vec<u32> {
238 let mut out = Vec::with_capacity(pieces.len() + 2);
239 out.push(self.hp.cls_id);
240 out.extend_from_slice(pieces);
241 out.push(self.hp.sep_id);
242 out
243 }
244
245 fn n_segments(&self) -> usize {
249 self.type_embd.as_ref().map(Vec::len).unwrap_or(1)
250 }
251
252 fn wrap_special_pair(&self, a: &[u32], b: &[u32]) -> Option<PairSequence> {
263 let mut tokens = Vec::with_capacity(a.len() + b.len() + 3);
264 tokens.push(self.hp.cls_id);
265 tokens.extend_from_slice(a);
266 tokens.push(self.hp.sep_id);
267 let first_half = tokens.len();
268 tokens.extend_from_slice(b);
269 tokens.push(self.hp.sep_id);
270 let mut segments = vec![0u32; tokens.len()];
271 for s in segments[first_half..].iter_mut() {
272 *s = 1;
273 }
274 Some(PairSequence { tokens, segments })
275 }
276
277 fn encode_on_worker(
278 &self,
279 tokens: &[u32],
280 segments: Option<&[u32]>,
281 ) -> Result<Vec<f32>, EncodeError> {
282 let n = tokens.len();
283 if n == 0 {
284 return Err(EncodeError::EmptySequence);
285 }
286 if let Some(seg) = segments {
287 if seg.len() != n {
288 return Err(EncodeError::RaggedSegments {
289 tokens: n,
290 segments: seg.len(),
291 });
292 }
293 }
294 if n > self.hp.n_ctx_train {
295 return Err(EncodeError::TooLong {
296 got: n,
297 max: self.hp.n_ctx_train,
298 arch: self.hp.arch.clone(),
299 });
300 }
301 let d = self.hp.n_embd;
302 let vocab_size = self.vocab_size();
303
304 let mut h = vec![0.0f32; n * d];
306 for (i, &t) in tokens.iter().enumerate() {
307 if t as usize >= vocab_size {
308 return Err(EncodeError::TokenOutOfRange { id: t, vocab_size });
309 }
310 let tok = self.tok_embd.dequant_row(t as usize);
311 let pos = self.pos_embd.dequant_row(i);
312 let row = &mut h[i * d..(i + 1) * d];
313 for (j, slot) in row.iter_mut().enumerate() {
314 *slot = tok[j] + pos[j];
315 }
316 if let Some(table) = &self.type_embd {
317 let seg = segments.map(|s| s[i]).unwrap_or(0);
318 let ty = table
319 .get(seg as usize)
320 .ok_or(EncodeError::SegmentOutOfRange {
321 id: seg,
322 pos: i,
323 n_segments: table.len(),
324 })?;
325 for (slot, tv) in row.iter_mut().zip(ty.iter()) {
326 *slot += tv;
327 }
328 }
329 }
330 layer_norm_rows(
331 &mut h,
332 d,
333 &self.tok_norm_w,
334 &self.tok_norm_b,
335 self.hp.layer_norm_eps,
336 );
337
338 let head_dim = self.hp.head_dim();
339 for layer in &self.layers {
340 let mut q = layer.wq.apply_batch(&h, n);
341 let mut k = layer.wk.apply_batch(&h, n);
342 let mut v = layer.wv.apply_batch(&h, n);
343 add_bias_rows(&mut q, self.hp.n_head * head_dim, layer.bq.as_ref());
344 add_bias_rows(&mut k, self.hp.n_head_kv * head_dim, layer.bk.as_ref());
345 add_bias_rows(&mut v, self.hp.n_head_kv * head_dim, layer.bv.as_ref());
346
347 let attn =
348 bidirectional_attention(&q, &k, &v, n, self.hp.n_head, self.hp.n_head_kv, head_dim);
349
350 let mut x = layer.wo.apply_batch(&attn, n);
351 add_bias_rows(&mut x, d, layer.bo.as_ref());
352 for (xv, hv) in x.iter_mut().zip(h.iter()) {
354 *xv += hv;
355 }
356 layer_norm_rows(
357 &mut x,
358 d,
359 &layer.attn_out_norm_w,
360 &layer.attn_out_norm_b,
361 self.hp.layer_norm_eps,
362 );
363
364 let mut up = layer.ffn_up.apply_batch(&x, n);
367 add_bias_rows(&mut up, self.hp.n_ff, layer.ffn_up_b.as_ref());
368 for a in up.iter_mut() {
369 *a = gelu(*a);
370 }
371 let mut down = layer.ffn_down.apply_batch(&up, n);
372 add_bias_rows(&mut down, d, layer.ffn_down_b.as_ref());
373 for (dv, xv) in down.iter_mut().zip(x.iter()) {
374 *dv += xv;
375 }
376 layer_norm_rows(
377 &mut down,
378 d,
379 &layer.layer_out_norm_w,
380 &layer.layer_out_norm_b,
381 self.hp.layer_norm_eps,
382 );
383 h = down;
384 }
385 Ok(h)
386 }
387}
388
389#[cfg(test)]
390mod tests {
391 use super::*;
392 use ferrox_core::tensor::Tensor;
393
394 struct Lcg(u64);
397 impl Lcg {
398 fn next_f32(&mut self) -> f32 {
399 self.0 = self.0.wrapping_mul(6364136223846793005).wrapping_add(1);
400 ((self.0 >> 33) as f32 / (1u64 << 31) as f32) - 0.5
401 }
402 fn vec(&mut self, n: usize) -> Vec<f32> {
403 (0..n).map(|_| self.next_f32()).collect()
404 }
405 fn matrix(&mut self, rows: usize, cols: usize) -> WeightMatrix {
406 WeightMatrix::F32(Tensor::new(self.vec(rows * cols), vec![rows, cols]))
407 }
408 }
409
410 const D: usize = 8;
411 const FF: usize = 16;
412 const HEADS: usize = 2;
413 const VOCAB: usize = 20;
414 const CTX: usize = 12;
415 const EPS: f32 = 1e-12;
416
417 fn fixture(n_layer: usize) -> BertEncoder {
418 let mut r = Lcg(0x5EED);
419 let tok_embd = r.matrix(VOCAB, D);
420 let pos_embd = r.matrix(CTX, D);
421 let type_embd = Some(vec![r.vec(D), r.vec(D)]);
423 let tok_norm_w = r.vec(D);
424 let tok_norm_b = r.vec(D);
425 let layers = (0..n_layer)
426 .map(|_| BertLayer {
427 wq: r.matrix(D, D),
428 bq: Some(r.vec(D)),
429 wk: r.matrix(D, D),
430 bk: Some(r.vec(D)),
431 wv: r.matrix(D, D),
432 bv: Some(r.vec(D)),
433 wo: r.matrix(D, D),
434 bo: Some(r.vec(D)),
435 attn_out_norm_w: r.vec(D),
436 attn_out_norm_b: r.vec(D),
437 ffn_up: r.matrix(FF, D),
438 ffn_up_b: Some(r.vec(FF)),
439 ffn_down: r.matrix(D, FF),
440 ffn_down_b: Some(r.vec(D)),
441 layer_out_norm_w: r.vec(D),
442 layer_out_norm_b: r.vec(D),
443 })
444 .collect();
445 BertEncoder {
446 hp: BertHparams {
447 arch: "bert".into(),
448 n_layer,
449 n_embd: D,
450 n_ff: FF,
451 n_head: HEADS,
452 n_head_kv: HEADS,
453 n_ctx_train: CTX,
454 n_token_types: 2,
455 layer_norm_eps: EPS,
456 pooling: PoolingType::Cls,
457 cls_id: 1,
458 sep_id: 2,
459 },
460 tok_embd,
461 type_embd,
462 pos_embd,
463 tok_norm_w,
464 tok_norm_b,
465 layers,
466 }
467 }
468
469 fn reference_forward(m: &BertEncoder, tokens: &[u32], segments: &[u32]) -> Vec<f64> {
476 let d = m.hp.n_embd;
477 let n = tokens.len();
478 let hd = m.hp.head_dim();
479
480 let dense = |w: &WeightMatrix| -> Vec<Vec<f64>> {
481 (0..w.rows())
482 .map(|r| w.dequant_row(r).iter().map(|&v| v as f64).collect())
483 .collect()
484 };
485 let matvec = |w: &Vec<Vec<f64>>, x: &[f64]| -> Vec<f64> {
486 w.iter()
487 .map(|row| row.iter().zip(x).map(|(a, b)| a * b).sum())
488 .collect()
489 };
490 let ln = |x: &[f64], wt: &[f32], b: &[f32]| -> Vec<f64> {
491 let mean = x.iter().sum::<f64>() / x.len() as f64;
492 let var = x.iter().map(|v| (v - mean).powi(2)).sum::<f64>() / x.len() as f64;
493 let inv = 1.0 / (var + m.hp.layer_norm_eps as f64).sqrt();
494 x.iter()
495 .zip(wt)
496 .zip(b)
497 .map(|((v, w), bb)| (v - mean) * inv * (*w as f64) + (*bb as f64))
498 .collect()
499 };
500
501 let mut h: Vec<Vec<f64>> = tokens
502 .iter()
503 .enumerate()
504 .map(|(i, &t)| {
505 let tok = m.tok_embd.dequant_row(t as usize);
506 let pos = m.pos_embd.dequant_row(i);
507 let ty = m
508 .type_embd
509 .as_ref()
510 .map(|t| t[segments[i] as usize].clone())
511 .unwrap_or_else(|| vec![0.0; d]);
512 let row: Vec<f64> = (0..d)
513 .map(|j| tok[j] as f64 + pos[j] as f64 + ty[j] as f64)
514 .collect();
515 ln(&row, &m.tok_norm_w, &m.tok_norm_b)
516 })
517 .collect();
518
519 for layer in &m.layers {
520 let (wq, wk, wv, wo) = (
521 dense(&layer.wq),
522 dense(&layer.wk),
523 dense(&layer.wv),
524 dense(&layer.wo),
525 );
526 let (wu, wd) = (dense(&layer.ffn_up), dense(&layer.ffn_down));
527 let bias = |v: &mut Vec<f64>, b: &Option<Vec<f32>>| {
528 if let Some(b) = b {
529 for (x, bb) in v.iter_mut().zip(b) {
530 *x += *bb as f64;
531 }
532 }
533 };
534 let mut q = Vec::new();
535 let mut k = Vec::new();
536 let mut v = Vec::new();
537 for row in &h {
538 let mut a = matvec(&wq, row);
539 bias(&mut a, &layer.bq);
540 q.push(a);
541 let mut a = matvec(&wk, row);
542 bias(&mut a, &layer.bk);
543 k.push(a);
544 let mut a = matvec(&wv, row);
545 bias(&mut a, &layer.bv);
546 v.push(a);
547 }
548 let mut attn = vec![vec![0.0f64; d]; n];
549 for head in 0..m.hp.n_head {
550 let off = head * hd;
551 for i in 0..n {
552 let raw: Vec<f64> = (0..n)
553 .map(|j| {
554 (0..hd).map(|c| q[i][off + c] * k[j][off + c]).sum::<f64>()
555 / (hd as f64).sqrt()
556 })
557 .collect();
558 let mx = raw.iter().cloned().fold(f64::NEG_INFINITY, f64::max);
559 let ex: Vec<f64> = raw.iter().map(|s| (s - mx).exp()).collect();
560 let sum: f64 = ex.iter().sum();
561 for j in 0..n {
562 let p = ex[j] / sum;
563 for c in 0..hd {
564 attn[i][off + c] += p * v[j][off + c];
565 }
566 }
567 }
568 }
569 let mut next = Vec::new();
570 for i in 0..n {
571 let mut o = matvec(&wo, &attn[i]);
572 bias(&mut o, &layer.bo);
573 for (x, hv) in o.iter_mut().zip(&h[i]) {
574 *x += hv;
575 }
576 let x = ln(&o, &layer.attn_out_norm_w, &layer.attn_out_norm_b);
577 let mut up = matvec(&wu, &x);
578 bias(&mut up, &layer.ffn_up_b);
579 let act: Vec<f64> = up
580 .iter()
581 .map(|&u| {
582 const K: f64 = 0.797_884_560_802_865_4;
583 const C: f64 = 0.044_715;
584 0.5 * u * (1.0 + (K * (u + C * u * u * u)).tanh())
585 })
586 .collect();
587 let mut down = matvec(&wd, &act);
588 bias(&mut down, &layer.ffn_down_b);
589 for (dv, xv) in down.iter_mut().zip(&x) {
590 *dv += xv;
591 }
592 next.push(ln(&down, &layer.layer_out_norm_w, &layer.layer_out_norm_b));
593 }
594 h = next;
595 }
596 h.into_iter().flatten().collect()
597 }
598
599 #[test]
600 fn matches_an_independent_f64_transcription_of_the_graph() {
601 let m = fixture(3);
602 let tokens = [1u32, 7, 13, 4, 9, 2];
603 let got = m.encode_tokens(&tokens).unwrap();
604 let want = reference_forward(&m, &tokens, &[0; 6]);
605 assert_eq!(got.len(), want.len());
606 for (i, (g, w)) in got.iter().zip(&want).enumerate() {
607 assert!(
608 (*g as f64 - w).abs() < 2e-4,
609 "element {i}: {g} vs reference {w}"
610 );
611 }
612 }
613
614 #[test]
620 fn the_segment_id_selects_the_token_type_row_at_every_position() {
621 let m = fixture(3);
622 let tokens = [1u32, 7, 13, 4, 9, 2];
623 let segments = [0u32, 0, 0, 1, 1, 1];
624 let got = m.encode(&tokens, Some(&segments)).unwrap();
625 let want = reference_forward(&m, &tokens, &segments);
626 for (i, (g, w)) in got.iter().zip(&want).enumerate() {
627 assert!(
628 (*g as f64 - w).abs() < 2e-4,
629 "element {i}: {g} vs reference {w}"
630 );
631 }
632 let all_zero = m.encode_tokens(&tokens).unwrap();
636 let moved: f32 = all_zero
637 .iter()
638 .zip(&got)
639 .map(|(x, y)| (x - y).abs())
640 .sum::<f32>();
641 assert!(moved > 1e-3, "segment 1 changed nothing ({moved})");
642 }
643
644 #[test]
647 fn a_segment_id_off_the_table_and_a_ragged_segment_list_are_refused() {
648 let m = fixture(1);
649 assert!(matches!(
650 m.encode(&[1, 7, 2], Some(&[0, 2, 0])),
651 Err(EncodeError::SegmentOutOfRange { id: 2, pos: 1, .. })
652 ));
653 assert!(matches!(
654 m.encode(&[1, 7, 2], Some(&[0, 0])),
655 Err(EncodeError::RaggedSegments {
656 tokens: 3,
657 segments: 2
658 })
659 ));
660 }
661
662 #[test]
666 fn attention_is_bidirectional_not_causal() {
667 let m = fixture(2);
668 let a = m.encode_tokens(&[5u32, 6, 7, 8]).unwrap();
669 let b = m.encode_tokens(&[5u32, 6, 7, 19]).unwrap();
670 let moved: f32 = a[..D].iter().zip(&b[..D]).map(|(x, y)| (x - y).abs()).sum();
671 assert!(
672 moved > 1e-3,
673 "row 0 barely moved ({moved}) when the last token changed — \
674 attention is behaving causally"
675 );
676 }
677
678 #[test]
681 fn position_embeddings_make_the_same_token_differ_by_index() {
682 let m = fixture(1);
683 let out = m.encode_tokens(&[11u32, 11]).unwrap();
684 let delta: f32 = out[..D]
685 .iter()
686 .zip(&out[D..2 * D])
687 .map(|(x, y)| (x - y).abs())
688 .sum();
689 assert!(
690 delta > 1e-3,
691 "identical tokens gave identical rows: {delta}"
692 );
693 }
694
695 #[test]
699 fn the_last_op_is_a_mean_subtracting_layer_norm() {
700 let mut m = fixture(2);
701 let last = m.layers.last_mut().unwrap();
702 last.layer_out_norm_w = vec![1.0; D];
703 last.layer_out_norm_b = vec![0.0; D];
704 let out = m.encode_tokens(&[3u32, 4, 5]).unwrap();
705 for row in out.as_chunks::<D>().0 {
706 let mean: f32 = row.iter().sum::<f32>() / D as f32;
707 let var: f32 = row.iter().map(|v| (v - mean).powi(2)).sum::<f32>() / D as f32;
708 assert!(mean.abs() < 1e-4, "row mean {mean} is not zero");
709 assert!((var - 1.0).abs() < 1e-3, "row variance {var} is not one");
710 }
711 }
712
713 #[test]
714 fn refuses_an_empty_sequence_and_one_past_the_position_table() {
715 let m = fixture(1);
716 assert!(matches!(
717 m.encode_tokens(&[]),
718 Err(EncodeError::EmptySequence)
719 ));
720 let long: Vec<u32> = (0..CTX as u32 + 1).map(|i| i % VOCAB as u32).collect();
721 let err = m.encode_tokens(&long).unwrap_err();
722 assert!(
723 matches!(err, EncodeError::TooLong { got, max, .. } if got == CTX + 1 && max == CTX)
724 );
725 assert!(matches!(
726 m.encode_tokens(&[VOCAB as u32]),
727 Err(EncodeError::TokenOutOfRange { .. })
728 ));
729 }
730
731 #[test]
732 fn wrap_special_brackets_the_pieces_with_cls_and_sep() {
733 let m = fixture(1);
734 assert_eq!(m.wrap_special(&[7, 8]), vec![1, 7, 8, 2]);
735 assert_eq!(m.wrap_special(&[]), vec![1, 2]);
736 }
737
738 #[test]
750 fn the_pair_form_separates_the_two_halves_and_labels_each_one() {
751 let m = fixture(1);
752 let pair = m.wrap_special_pair(&[7, 8], &[9]).unwrap();
753 assert_eq!(pair.tokens, vec![1, 7, 8, 2, 9, 2]);
754 assert_eq!(pair.segments, vec![0, 0, 0, 0, 1, 1]);
755 let empty = m.wrap_special_pair(&[], &[]).unwrap();
757 assert_eq!(empty.tokens, vec![1, 2, 2]);
758 assert_eq!(empty.segments, vec![0, 0, 1]);
759 assert_ne!(pair.tokens, m.wrap_special(&[7, 8, 9]));
762 }
763
764 #[test]
768 fn n_segments_is_the_height_of_the_token_type_table() {
769 let mut m = fixture(1);
770 assert_eq!(m.n_segments(), 2);
771 m.type_embd = Some(vec![vec![0.0; D]]);
772 assert_eq!(m.n_segments(), 1);
773 m.type_embd = None;
774 assert_eq!(m.n_segments(), 1);
775 }
776
777 #[test]
780 fn embed_tokens_pools_the_way_the_hparams_say() {
781 let m = fixture(2);
782 let tokens = [1u32, 9, 4, 2];
783 let hidden = m.encode_tokens(&tokens).unwrap();
784 assert_eq!(m.embed_tokens(&tokens).unwrap(), hidden[..D].to_vec());
785 }
786}