1use ferrox_core::matmul::{gelu, layer_norm};
56use ferrox_core::weight_matrix::WeightMatrix;
57
58use crate::encoder::{EncodeError, TextEncoder};
59use crate::pooling::PoolingType;
60
61#[derive(Debug, Clone)]
63pub struct BertHparams {
64 pub arch: String,
65 pub n_layer: usize,
66 pub n_embd: usize,
67 pub n_ff: usize,
68 pub n_head: usize,
69 pub n_head_kv: usize,
70 pub n_ctx_train: usize,
72 pub n_token_types: usize,
73 pub layer_norm_eps: f32,
74 pub pooling: PoolingType,
75 pub cls_id: u32,
79 pub sep_id: u32,
80}
81
82impl BertHparams {
83 pub fn head_dim(&self) -> usize {
84 self.n_embd / self.n_head
85 }
86}
87
88pub struct BertLayer {
92 pub wq: WeightMatrix,
93 pub bq: Option<Vec<f32>>,
94 pub wk: WeightMatrix,
95 pub bk: Option<Vec<f32>>,
96 pub wv: WeightMatrix,
97 pub bv: Option<Vec<f32>>,
98 pub wo: WeightMatrix,
99 pub bo: Option<Vec<f32>>,
100 pub attn_out_norm_w: Vec<f32>,
102 pub attn_out_norm_b: Vec<f32>,
103 pub ffn_up: WeightMatrix,
104 pub ffn_up_b: Option<Vec<f32>>,
105 pub ffn_down: WeightMatrix,
106 pub ffn_down_b: Option<Vec<f32>>,
107 pub layer_out_norm_w: Vec<f32>,
109 pub layer_out_norm_b: Vec<f32>,
110}
111
112pub struct BertEncoder {
113 pub hp: BertHparams,
114 pub tok_embd: WeightMatrix,
115 pub type_embd_row0: Option<Vec<f32>>,
117 pub pos_embd: WeightMatrix,
118 pub tok_norm_w: Vec<f32>,
119 pub tok_norm_b: Vec<f32>,
120 pub layers: Vec<BertLayer>,
121}
122
123fn add_bias_rows(rows: &mut [f32], width: usize, bias: Option<&Vec<f32>>) {
125 let Some(b) = bias else { return };
126 debug_assert_eq!(b.len(), width);
127 for row in rows.chunks_exact_mut(width) {
128 for (x, bv) in row.iter_mut().zip(b.iter()) {
129 *x += bv;
130 }
131 }
132}
133
134fn layer_norm_rows(rows: &mut [f32], width: usize, weight: &[f32], bias: &[f32], eps: f32) {
136 for row in rows.chunks_exact_mut(width) {
137 let normed = layer_norm(row, weight, bias, eps);
138 row.copy_from_slice(&normed);
139 }
140}
141
142fn softmax_row(scores: &mut [f32]) {
144 let max = scores.iter().copied().fold(f32::NEG_INFINITY, f32::max);
145 let mut sum = 0.0f32;
146 for s in scores.iter_mut() {
147 *s = (*s - max).exp();
148 sum += *s;
149 }
150 let inv = 1.0 / sum;
151 for s in scores.iter_mut() {
152 *s *= inv;
153 }
154}
155
156fn bidirectional_attention(
163 q: &[f32],
164 k: &[f32],
165 v: &[f32],
166 n: usize,
167 n_head: usize,
168 n_head_kv: usize,
169 head_dim: usize,
170) -> Vec<f32> {
171 let q_width = n_head * head_dim;
172 let kv_width = n_head_kv * head_dim;
173 let heads_per_kv = n_head / n_head_kv;
174 let scale = 1.0 / (head_dim as f32).sqrt();
175 let mut out = vec![0.0f32; n * q_width];
176 let mut scores = vec![0.0f32; n];
177 for h in 0..n_head {
178 let kv_h = h / heads_per_kv;
179 let q_off = h * head_dim;
180 let kv_off = kv_h * head_dim;
181 for i in 0..n {
182 let qi = &q[i * q_width + q_off..i * q_width + q_off + head_dim];
183 for (j, s) in scores.iter_mut().enumerate() {
184 let kj = &k[j * kv_width + kv_off..j * kv_width + kv_off + head_dim];
185 *s = qi.iter().zip(kj).map(|(a, b)| a * b).sum::<f32>() * scale;
186 }
187 softmax_row(&mut scores);
188 let dst = &mut out[i * q_width + q_off..i * q_width + q_off + head_dim];
189 for (j, &p) in scores.iter().enumerate() {
190 let vj = &v[j * kv_width + kv_off..j * kv_width + kv_off + head_dim];
191 for (o, &vv) in dst.iter_mut().zip(vj) {
192 *o += p * vv;
193 }
194 }
195 }
196 }
197 out
198}
199
200impl BertEncoder {
201 pub fn vocab_size(&self) -> usize {
202 self.tok_embd.rows()
203 }
204}
205
206impl TextEncoder for BertEncoder {
207 fn n_embd(&self) -> usize {
208 self.hp.n_embd
209 }
210
211 fn n_ctx_train(&self) -> usize {
212 self.hp.n_ctx_train
213 }
214
215 fn pooling_type(&self) -> PoolingType {
216 self.hp.pooling
217 }
218
219 fn wrap_special(&self, pieces: &[u32]) -> Vec<u32> {
225 let mut out = Vec::with_capacity(pieces.len() + 2);
226 out.push(self.hp.cls_id);
227 out.extend_from_slice(pieces);
228 out.push(self.hp.sep_id);
229 out
230 }
231
232 fn wrap_special_pair(&self, a: &[u32], b: &[u32]) -> Option<Vec<u32>> {
246 let mut out = Vec::with_capacity(a.len() + b.len() + 3);
247 out.push(self.hp.cls_id);
248 out.extend_from_slice(a);
249 out.push(self.hp.sep_id);
250 out.extend_from_slice(b);
251 out.push(self.hp.sep_id);
252 Some(out)
253 }
254
255 fn encode_tokens(&self, tokens: &[u32]) -> Result<Vec<f32>, EncodeError> {
256 let n = tokens.len();
257 if n == 0 {
258 return Err(EncodeError::EmptySequence);
259 }
260 if n > self.hp.n_ctx_train {
261 return Err(EncodeError::TooLong {
262 got: n,
263 max: self.hp.n_ctx_train,
264 arch: self.hp.arch.clone(),
265 });
266 }
267 let d = self.hp.n_embd;
268 let vocab_size = self.vocab_size();
269
270 let mut h = vec![0.0f32; n * d];
272 for (i, &t) in tokens.iter().enumerate() {
273 if t as usize >= vocab_size {
274 return Err(EncodeError::TokenOutOfRange { id: t, vocab_size });
275 }
276 let tok = self.tok_embd.dequant_row(t as usize);
277 let pos = self.pos_embd.dequant_row(i);
278 let row = &mut h[i * d..(i + 1) * d];
279 for (j, slot) in row.iter_mut().enumerate() {
280 *slot = tok[j] + pos[j];
281 }
282 if let Some(ty) = &self.type_embd_row0 {
283 for (slot, tv) in row.iter_mut().zip(ty.iter()) {
284 *slot += tv;
285 }
286 }
287 }
288 layer_norm_rows(
289 &mut h,
290 d,
291 &self.tok_norm_w,
292 &self.tok_norm_b,
293 self.hp.layer_norm_eps,
294 );
295
296 let head_dim = self.hp.head_dim();
297 for layer in &self.layers {
298 let mut q = layer.wq.apply_batch(&h, n);
299 let mut k = layer.wk.apply_batch(&h, n);
300 let mut v = layer.wv.apply_batch(&h, n);
301 add_bias_rows(&mut q, self.hp.n_head * head_dim, layer.bq.as_ref());
302 add_bias_rows(&mut k, self.hp.n_head_kv * head_dim, layer.bk.as_ref());
303 add_bias_rows(&mut v, self.hp.n_head_kv * head_dim, layer.bv.as_ref());
304
305 let attn =
306 bidirectional_attention(&q, &k, &v, n, self.hp.n_head, self.hp.n_head_kv, head_dim);
307
308 let mut x = layer.wo.apply_batch(&attn, n);
309 add_bias_rows(&mut x, d, layer.bo.as_ref());
310 for (xv, hv) in x.iter_mut().zip(h.iter()) {
312 *xv += hv;
313 }
314 layer_norm_rows(
315 &mut x,
316 d,
317 &layer.attn_out_norm_w,
318 &layer.attn_out_norm_b,
319 self.hp.layer_norm_eps,
320 );
321
322 let mut up = layer.ffn_up.apply_batch(&x, n);
325 add_bias_rows(&mut up, self.hp.n_ff, layer.ffn_up_b.as_ref());
326 for a in up.iter_mut() {
327 *a = gelu(*a);
328 }
329 let mut down = layer.ffn_down.apply_batch(&up, n);
330 add_bias_rows(&mut down, d, layer.ffn_down_b.as_ref());
331 for (dv, xv) in down.iter_mut().zip(x.iter()) {
332 *dv += xv;
333 }
334 layer_norm_rows(
335 &mut down,
336 d,
337 &layer.layer_out_norm_w,
338 &layer.layer_out_norm_b,
339 self.hp.layer_norm_eps,
340 );
341 h = down;
342 }
343 Ok(h)
344 }
345}
346
347#[cfg(test)]
348mod tests {
349 use super::*;
350 use ferrox_core::tensor::Tensor;
351
352 struct Lcg(u64);
355 impl Lcg {
356 fn next_f32(&mut self) -> f32 {
357 self.0 = self.0.wrapping_mul(6364136223846793005).wrapping_add(1);
358 ((self.0 >> 33) as f32 / (1u64 << 31) as f32) - 0.5
359 }
360 fn vec(&mut self, n: usize) -> Vec<f32> {
361 (0..n).map(|_| self.next_f32()).collect()
362 }
363 fn matrix(&mut self, rows: usize, cols: usize) -> WeightMatrix {
364 WeightMatrix::F32(Tensor::new(self.vec(rows * cols), vec![rows, cols]))
365 }
366 }
367
368 const D: usize = 8;
369 const FF: usize = 16;
370 const HEADS: usize = 2;
371 const VOCAB: usize = 20;
372 const CTX: usize = 12;
373 const EPS: f32 = 1e-12;
374
375 fn fixture(n_layer: usize) -> BertEncoder {
376 let mut r = Lcg(0x5EED);
377 let tok_embd = r.matrix(VOCAB, D);
378 let pos_embd = r.matrix(CTX, D);
379 let type_embd_row0 = Some(r.vec(D));
380 let tok_norm_w = r.vec(D);
381 let tok_norm_b = r.vec(D);
382 let layers = (0..n_layer)
383 .map(|_| BertLayer {
384 wq: r.matrix(D, D),
385 bq: Some(r.vec(D)),
386 wk: r.matrix(D, D),
387 bk: Some(r.vec(D)),
388 wv: r.matrix(D, D),
389 bv: Some(r.vec(D)),
390 wo: r.matrix(D, D),
391 bo: Some(r.vec(D)),
392 attn_out_norm_w: r.vec(D),
393 attn_out_norm_b: r.vec(D),
394 ffn_up: r.matrix(FF, D),
395 ffn_up_b: Some(r.vec(FF)),
396 ffn_down: r.matrix(D, FF),
397 ffn_down_b: Some(r.vec(D)),
398 layer_out_norm_w: r.vec(D),
399 layer_out_norm_b: r.vec(D),
400 })
401 .collect();
402 BertEncoder {
403 hp: BertHparams {
404 arch: "bert".into(),
405 n_layer,
406 n_embd: D,
407 n_ff: FF,
408 n_head: HEADS,
409 n_head_kv: HEADS,
410 n_ctx_train: CTX,
411 n_token_types: 2,
412 layer_norm_eps: EPS,
413 pooling: PoolingType::Cls,
414 cls_id: 1,
415 sep_id: 2,
416 },
417 tok_embd,
418 type_embd_row0,
419 pos_embd,
420 tok_norm_w,
421 tok_norm_b,
422 layers,
423 }
424 }
425
426 fn reference_forward(m: &BertEncoder, tokens: &[u32]) -> Vec<f64> {
433 let d = m.hp.n_embd;
434 let n = tokens.len();
435 let hd = m.hp.head_dim();
436
437 let dense = |w: &WeightMatrix| -> Vec<Vec<f64>> {
438 (0..w.rows())
439 .map(|r| w.dequant_row(r).iter().map(|&v| v as f64).collect())
440 .collect()
441 };
442 let matvec = |w: &Vec<Vec<f64>>, x: &[f64]| -> Vec<f64> {
443 w.iter()
444 .map(|row| row.iter().zip(x).map(|(a, b)| a * b).sum())
445 .collect()
446 };
447 let ln = |x: &[f64], wt: &[f32], b: &[f32]| -> Vec<f64> {
448 let mean = x.iter().sum::<f64>() / x.len() as f64;
449 let var = x.iter().map(|v| (v - mean).powi(2)).sum::<f64>() / x.len() as f64;
450 let inv = 1.0 / (var + m.hp.layer_norm_eps as f64).sqrt();
451 x.iter()
452 .zip(wt)
453 .zip(b)
454 .map(|((v, w), bb)| (v - mean) * inv * (*w as f64) + (*bb as f64))
455 .collect()
456 };
457
458 let mut h: Vec<Vec<f64>> = tokens
459 .iter()
460 .enumerate()
461 .map(|(i, &t)| {
462 let tok = m.tok_embd.dequant_row(t as usize);
463 let pos = m.pos_embd.dequant_row(i);
464 let ty = m.type_embd_row0.clone().unwrap_or(vec![0.0; d]);
465 let row: Vec<f64> = (0..d)
466 .map(|j| tok[j] as f64 + pos[j] as f64 + ty[j] as f64)
467 .collect();
468 ln(&row, &m.tok_norm_w, &m.tok_norm_b)
469 })
470 .collect();
471
472 for layer in &m.layers {
473 let (wq, wk, wv, wo) = (
474 dense(&layer.wq),
475 dense(&layer.wk),
476 dense(&layer.wv),
477 dense(&layer.wo),
478 );
479 let (wu, wd) = (dense(&layer.ffn_up), dense(&layer.ffn_down));
480 let bias = |v: &mut Vec<f64>, b: &Option<Vec<f32>>| {
481 if let Some(b) = b {
482 for (x, bb) in v.iter_mut().zip(b) {
483 *x += *bb as f64;
484 }
485 }
486 };
487 let mut q = Vec::new();
488 let mut k = Vec::new();
489 let mut v = Vec::new();
490 for row in &h {
491 let mut a = matvec(&wq, row);
492 bias(&mut a, &layer.bq);
493 q.push(a);
494 let mut a = matvec(&wk, row);
495 bias(&mut a, &layer.bk);
496 k.push(a);
497 let mut a = matvec(&wv, row);
498 bias(&mut a, &layer.bv);
499 v.push(a);
500 }
501 let mut attn = vec![vec![0.0f64; d]; n];
502 for head in 0..m.hp.n_head {
503 let off = head * hd;
504 for i in 0..n {
505 let raw: Vec<f64> = (0..n)
506 .map(|j| {
507 (0..hd).map(|c| q[i][off + c] * k[j][off + c]).sum::<f64>()
508 / (hd as f64).sqrt()
509 })
510 .collect();
511 let mx = raw.iter().cloned().fold(f64::NEG_INFINITY, f64::max);
512 let ex: Vec<f64> = raw.iter().map(|s| (s - mx).exp()).collect();
513 let sum: f64 = ex.iter().sum();
514 for j in 0..n {
515 let p = ex[j] / sum;
516 for c in 0..hd {
517 attn[i][off + c] += p * v[j][off + c];
518 }
519 }
520 }
521 }
522 let mut next = Vec::new();
523 for i in 0..n {
524 let mut o = matvec(&wo, &attn[i]);
525 bias(&mut o, &layer.bo);
526 for (x, hv) in o.iter_mut().zip(&h[i]) {
527 *x += hv;
528 }
529 let x = ln(&o, &layer.attn_out_norm_w, &layer.attn_out_norm_b);
530 let mut up = matvec(&wu, &x);
531 bias(&mut up, &layer.ffn_up_b);
532 let act: Vec<f64> = up
533 .iter()
534 .map(|&u| {
535 const K: f64 = 0.797_884_560_802_865_4;
536 const C: f64 = 0.044_715;
537 0.5 * u * (1.0 + (K * (u + C * u * u * u)).tanh())
538 })
539 .collect();
540 let mut down = matvec(&wd, &act);
541 bias(&mut down, &layer.ffn_down_b);
542 for (dv, xv) in down.iter_mut().zip(&x) {
543 *dv += xv;
544 }
545 next.push(ln(&down, &layer.layer_out_norm_w, &layer.layer_out_norm_b));
546 }
547 h = next;
548 }
549 h.into_iter().flatten().collect()
550 }
551
552 #[test]
553 fn matches_an_independent_f64_transcription_of_the_graph() {
554 let m = fixture(3);
555 let tokens = [1u32, 7, 13, 4, 9, 2];
556 let got = m.encode_tokens(&tokens).unwrap();
557 let want = reference_forward(&m, &tokens);
558 assert_eq!(got.len(), want.len());
559 for (i, (g, w)) in got.iter().zip(&want).enumerate() {
560 assert!(
561 (*g as f64 - w).abs() < 2e-4,
562 "element {i}: {g} vs reference {w}"
563 );
564 }
565 }
566
567 #[test]
571 fn attention_is_bidirectional_not_causal() {
572 let m = fixture(2);
573 let a = m.encode_tokens(&[5u32, 6, 7, 8]).unwrap();
574 let b = m.encode_tokens(&[5u32, 6, 7, 19]).unwrap();
575 let moved: f32 = a[..D].iter().zip(&b[..D]).map(|(x, y)| (x - y).abs()).sum();
576 assert!(
577 moved > 1e-3,
578 "row 0 barely moved ({moved}) when the last token changed — \
579 attention is behaving causally"
580 );
581 }
582
583 #[test]
586 fn position_embeddings_make_the_same_token_differ_by_index() {
587 let m = fixture(1);
588 let out = m.encode_tokens(&[11u32, 11]).unwrap();
589 let delta: f32 = out[..D]
590 .iter()
591 .zip(&out[D..2 * D])
592 .map(|(x, y)| (x - y).abs())
593 .sum();
594 assert!(
595 delta > 1e-3,
596 "identical tokens gave identical rows: {delta}"
597 );
598 }
599
600 #[test]
604 fn the_last_op_is_a_mean_subtracting_layer_norm() {
605 let mut m = fixture(2);
606 let last = m.layers.last_mut().unwrap();
607 last.layer_out_norm_w = vec![1.0; D];
608 last.layer_out_norm_b = vec![0.0; D];
609 let out = m.encode_tokens(&[3u32, 4, 5]).unwrap();
610 for row in out.as_chunks::<D>().0 {
611 let mean: f32 = row.iter().sum::<f32>() / D as f32;
612 let var: f32 = row.iter().map(|v| (v - mean).powi(2)).sum::<f32>() / D as f32;
613 assert!(mean.abs() < 1e-4, "row mean {mean} is not zero");
614 assert!((var - 1.0).abs() < 1e-3, "row variance {var} is not one");
615 }
616 }
617
618 #[test]
619 fn refuses_an_empty_sequence_and_one_past_the_position_table() {
620 let m = fixture(1);
621 assert!(matches!(
622 m.encode_tokens(&[]),
623 Err(EncodeError::EmptySequence)
624 ));
625 let long: Vec<u32> = (0..CTX as u32 + 1).map(|i| i % VOCAB as u32).collect();
626 let err = m.encode_tokens(&long).unwrap_err();
627 assert!(
628 matches!(err, EncodeError::TooLong { got, max, .. } if got == CTX + 1 && max == CTX)
629 );
630 assert!(matches!(
631 m.encode_tokens(&[VOCAB as u32]),
632 Err(EncodeError::TokenOutOfRange { .. })
633 ));
634 }
635
636 #[test]
637 fn wrap_special_brackets_the_pieces_with_cls_and_sep() {
638 let m = fixture(1);
639 assert_eq!(m.wrap_special(&[7, 8]), vec![1, 7, 8, 2]);
640 assert_eq!(m.wrap_special(&[]), vec![1, 2]);
641 }
642
643 #[test]
650 fn the_pair_form_separates_the_two_halves_and_closes_the_second() {
651 let m = fixture(1);
652 assert_eq!(
653 m.wrap_special_pair(&[7, 8], &[9]).unwrap(),
654 vec![1, 7, 8, 2, 9, 2]
655 );
656 assert_eq!(m.wrap_special_pair(&[], &[]).unwrap(), vec![1, 2, 2]);
658 assert_ne!(
661 m.wrap_special_pair(&[7, 8], &[9]).unwrap(),
662 m.wrap_special(&[7, 8, 9])
663 );
664 }
665
666 #[test]
669 fn embed_tokens_pools_the_way_the_hparams_say() {
670 let m = fixture(2);
671 let tokens = [1u32, 9, 4, 2];
672 let hidden = m.encode_tokens(&tokens).unwrap();
673 assert_eq!(m.embed_tokens(&tokens).unwrap(), hidden[..D].to_vec());
674 }
675}