use ferrox_core::matmul::{gelu, layer_norm};
use ferrox_core::weight_matrix::WeightMatrix;
use crate::encoder::{EncodeError, TextEncoder};
use crate::pooling::PoolingType;
#[derive(Debug, Clone)]
pub struct BertHparams {
pub arch: String,
pub n_layer: usize,
pub n_embd: usize,
pub n_ff: usize,
pub n_head: usize,
pub n_head_kv: usize,
pub n_ctx_train: usize,
pub n_token_types: usize,
pub layer_norm_eps: f32,
pub pooling: PoolingType,
pub cls_id: u32,
pub sep_id: u32,
}
impl BertHparams {
pub fn head_dim(&self) -> usize {
self.n_embd / self.n_head
}
}
pub struct BertLayer {
pub wq: WeightMatrix,
pub bq: Option<Vec<f32>>,
pub wk: WeightMatrix,
pub bk: Option<Vec<f32>>,
pub wv: WeightMatrix,
pub bv: Option<Vec<f32>>,
pub wo: WeightMatrix,
pub bo: Option<Vec<f32>>,
pub attn_out_norm_w: Vec<f32>,
pub attn_out_norm_b: Vec<f32>,
pub ffn_up: WeightMatrix,
pub ffn_up_b: Option<Vec<f32>>,
pub ffn_down: WeightMatrix,
pub ffn_down_b: Option<Vec<f32>>,
pub layer_out_norm_w: Vec<f32>,
pub layer_out_norm_b: Vec<f32>,
}
pub struct BertEncoder {
pub hp: BertHparams,
pub tok_embd: WeightMatrix,
pub type_embd_row0: Option<Vec<f32>>,
pub pos_embd: WeightMatrix,
pub tok_norm_w: Vec<f32>,
pub tok_norm_b: Vec<f32>,
pub layers: Vec<BertLayer>,
}
fn add_bias_rows(rows: &mut [f32], width: usize, bias: Option<&Vec<f32>>) {
let Some(b) = bias else { return };
debug_assert_eq!(b.len(), width);
for row in rows.chunks_exact_mut(width) {
for (x, bv) in row.iter_mut().zip(b.iter()) {
*x += bv;
}
}
}
fn layer_norm_rows(rows: &mut [f32], width: usize, weight: &[f32], bias: &[f32], eps: f32) {
for row in rows.chunks_exact_mut(width) {
let normed = layer_norm(row, weight, bias, eps);
row.copy_from_slice(&normed);
}
}
fn softmax_row(scores: &mut [f32]) {
let max = scores.iter().copied().fold(f32::NEG_INFINITY, f32::max);
let mut sum = 0.0f32;
for s in scores.iter_mut() {
*s = (*s - max).exp();
sum += *s;
}
let inv = 1.0 / sum;
for s in scores.iter_mut() {
*s *= inv;
}
}
fn bidirectional_attention(
q: &[f32],
k: &[f32],
v: &[f32],
n: usize,
n_head: usize,
n_head_kv: usize,
head_dim: usize,
) -> Vec<f32> {
let q_width = n_head * head_dim;
let kv_width = n_head_kv * head_dim;
let heads_per_kv = n_head / n_head_kv;
let scale = 1.0 / (head_dim as f32).sqrt();
let mut out = vec![0.0f32; n * q_width];
let mut scores = vec![0.0f32; n];
for h in 0..n_head {
let kv_h = h / heads_per_kv;
let q_off = h * head_dim;
let kv_off = kv_h * head_dim;
for i in 0..n {
let qi = &q[i * q_width + q_off..i * q_width + q_off + head_dim];
for (j, s) in scores.iter_mut().enumerate() {
let kj = &k[j * kv_width + kv_off..j * kv_width + kv_off + head_dim];
*s = qi.iter().zip(kj).map(|(a, b)| a * b).sum::<f32>() * scale;
}
softmax_row(&mut scores);
let dst = &mut out[i * q_width + q_off..i * q_width + q_off + head_dim];
for (j, &p) in scores.iter().enumerate() {
let vj = &v[j * kv_width + kv_off..j * kv_width + kv_off + head_dim];
for (o, &vv) in dst.iter_mut().zip(vj) {
*o += p * vv;
}
}
}
}
out
}
impl BertEncoder {
pub fn vocab_size(&self) -> usize {
self.tok_embd.rows()
}
}
impl TextEncoder for BertEncoder {
fn n_embd(&self) -> usize {
self.hp.n_embd
}
fn n_ctx_train(&self) -> usize {
self.hp.n_ctx_train
}
fn pooling_type(&self) -> PoolingType {
self.hp.pooling
}
fn wrap_special(&self, pieces: &[u32]) -> Vec<u32> {
let mut out = Vec::with_capacity(pieces.len() + 2);
out.push(self.hp.cls_id);
out.extend_from_slice(pieces);
out.push(self.hp.sep_id);
out
}
fn wrap_special_pair(&self, a: &[u32], b: &[u32]) -> Option<Vec<u32>> {
let mut out = Vec::with_capacity(a.len() + b.len() + 3);
out.push(self.hp.cls_id);
out.extend_from_slice(a);
out.push(self.hp.sep_id);
out.extend_from_slice(b);
out.push(self.hp.sep_id);
Some(out)
}
fn encode_tokens(&self, tokens: &[u32]) -> Result<Vec<f32>, EncodeError> {
let n = tokens.len();
if n == 0 {
return Err(EncodeError::EmptySequence);
}
if n > self.hp.n_ctx_train {
return Err(EncodeError::TooLong {
got: n,
max: self.hp.n_ctx_train,
arch: self.hp.arch.clone(),
});
}
let d = self.hp.n_embd;
let vocab_size = self.vocab_size();
let mut h = vec![0.0f32; n * d];
for (i, &t) in tokens.iter().enumerate() {
if t as usize >= vocab_size {
return Err(EncodeError::TokenOutOfRange { id: t, vocab_size });
}
let tok = self.tok_embd.dequant_row(t as usize);
let pos = self.pos_embd.dequant_row(i);
let row = &mut h[i * d..(i + 1) * d];
for (j, slot) in row.iter_mut().enumerate() {
*slot = tok[j] + pos[j];
}
if let Some(ty) = &self.type_embd_row0 {
for (slot, tv) in row.iter_mut().zip(ty.iter()) {
*slot += tv;
}
}
}
layer_norm_rows(
&mut h,
d,
&self.tok_norm_w,
&self.tok_norm_b,
self.hp.layer_norm_eps,
);
let head_dim = self.hp.head_dim();
for layer in &self.layers {
let mut q = layer.wq.apply_batch(&h, n);
let mut k = layer.wk.apply_batch(&h, n);
let mut v = layer.wv.apply_batch(&h, n);
add_bias_rows(&mut q, self.hp.n_head * head_dim, layer.bq.as_ref());
add_bias_rows(&mut k, self.hp.n_head_kv * head_dim, layer.bk.as_ref());
add_bias_rows(&mut v, self.hp.n_head_kv * head_dim, layer.bv.as_ref());
let attn =
bidirectional_attention(&q, &k, &v, n, self.hp.n_head, self.hp.n_head_kv, head_dim);
let mut x = layer.wo.apply_batch(&attn, n);
add_bias_rows(&mut x, d, layer.bo.as_ref());
for (xv, hv) in x.iter_mut().zip(h.iter()) {
*xv += hv;
}
layer_norm_rows(
&mut x,
d,
&layer.attn_out_norm_w,
&layer.attn_out_norm_b,
self.hp.layer_norm_eps,
);
let mut up = layer.ffn_up.apply_batch(&x, n);
add_bias_rows(&mut up, self.hp.n_ff, layer.ffn_up_b.as_ref());
for a in up.iter_mut() {
*a = gelu(*a);
}
let mut down = layer.ffn_down.apply_batch(&up, n);
add_bias_rows(&mut down, d, layer.ffn_down_b.as_ref());
for (dv, xv) in down.iter_mut().zip(x.iter()) {
*dv += xv;
}
layer_norm_rows(
&mut down,
d,
&layer.layer_out_norm_w,
&layer.layer_out_norm_b,
self.hp.layer_norm_eps,
);
h = down;
}
Ok(h)
}
}
#[cfg(test)]
mod tests {
use super::*;
use ferrox_core::tensor::Tensor;
struct Lcg(u64);
impl Lcg {
fn next_f32(&mut self) -> f32 {
self.0 = self.0.wrapping_mul(6364136223846793005).wrapping_add(1);
((self.0 >> 33) as f32 / (1u64 << 31) as f32) - 0.5
}
fn vec(&mut self, n: usize) -> Vec<f32> {
(0..n).map(|_| self.next_f32()).collect()
}
fn matrix(&mut self, rows: usize, cols: usize) -> WeightMatrix {
WeightMatrix::F32(Tensor::new(self.vec(rows * cols), vec![rows, cols]))
}
}
const D: usize = 8;
const FF: usize = 16;
const HEADS: usize = 2;
const VOCAB: usize = 20;
const CTX: usize = 12;
const EPS: f32 = 1e-12;
fn fixture(n_layer: usize) -> BertEncoder {
let mut r = Lcg(0x5EED);
let tok_embd = r.matrix(VOCAB, D);
let pos_embd = r.matrix(CTX, D);
let type_embd_row0 = Some(r.vec(D));
let tok_norm_w = r.vec(D);
let tok_norm_b = r.vec(D);
let layers = (0..n_layer)
.map(|_| BertLayer {
wq: r.matrix(D, D),
bq: Some(r.vec(D)),
wk: r.matrix(D, D),
bk: Some(r.vec(D)),
wv: r.matrix(D, D),
bv: Some(r.vec(D)),
wo: r.matrix(D, D),
bo: Some(r.vec(D)),
attn_out_norm_w: r.vec(D),
attn_out_norm_b: r.vec(D),
ffn_up: r.matrix(FF, D),
ffn_up_b: Some(r.vec(FF)),
ffn_down: r.matrix(D, FF),
ffn_down_b: Some(r.vec(D)),
layer_out_norm_w: r.vec(D),
layer_out_norm_b: r.vec(D),
})
.collect();
BertEncoder {
hp: BertHparams {
arch: "bert".into(),
n_layer,
n_embd: D,
n_ff: FF,
n_head: HEADS,
n_head_kv: HEADS,
n_ctx_train: CTX,
n_token_types: 2,
layer_norm_eps: EPS,
pooling: PoolingType::Cls,
cls_id: 1,
sep_id: 2,
},
tok_embd,
type_embd_row0,
pos_embd,
tok_norm_w,
tok_norm_b,
layers,
}
}
fn reference_forward(m: &BertEncoder, tokens: &[u32]) -> Vec<f64> {
let d = m.hp.n_embd;
let n = tokens.len();
let hd = m.hp.head_dim();
let dense = |w: &WeightMatrix| -> Vec<Vec<f64>> {
(0..w.rows())
.map(|r| w.dequant_row(r).iter().map(|&v| v as f64).collect())
.collect()
};
let matvec = |w: &Vec<Vec<f64>>, x: &[f64]| -> Vec<f64> {
w.iter()
.map(|row| row.iter().zip(x).map(|(a, b)| a * b).sum())
.collect()
};
let ln = |x: &[f64], wt: &[f32], b: &[f32]| -> Vec<f64> {
let mean = x.iter().sum::<f64>() / x.len() as f64;
let var = x.iter().map(|v| (v - mean).powi(2)).sum::<f64>() / x.len() as f64;
let inv = 1.0 / (var + m.hp.layer_norm_eps as f64).sqrt();
x.iter()
.zip(wt)
.zip(b)
.map(|((v, w), bb)| (v - mean) * inv * (*w as f64) + (*bb as f64))
.collect()
};
let mut h: Vec<Vec<f64>> = tokens
.iter()
.enumerate()
.map(|(i, &t)| {
let tok = m.tok_embd.dequant_row(t as usize);
let pos = m.pos_embd.dequant_row(i);
let ty = m.type_embd_row0.clone().unwrap_or(vec![0.0; d]);
let row: Vec<f64> = (0..d)
.map(|j| tok[j] as f64 + pos[j] as f64 + ty[j] as f64)
.collect();
ln(&row, &m.tok_norm_w, &m.tok_norm_b)
})
.collect();
for layer in &m.layers {
let (wq, wk, wv, wo) = (
dense(&layer.wq),
dense(&layer.wk),
dense(&layer.wv),
dense(&layer.wo),
);
let (wu, wd) = (dense(&layer.ffn_up), dense(&layer.ffn_down));
let bias = |v: &mut Vec<f64>, b: &Option<Vec<f32>>| {
if let Some(b) = b {
for (x, bb) in v.iter_mut().zip(b) {
*x += *bb as f64;
}
}
};
let mut q = Vec::new();
let mut k = Vec::new();
let mut v = Vec::new();
for row in &h {
let mut a = matvec(&wq, row);
bias(&mut a, &layer.bq);
q.push(a);
let mut a = matvec(&wk, row);
bias(&mut a, &layer.bk);
k.push(a);
let mut a = matvec(&wv, row);
bias(&mut a, &layer.bv);
v.push(a);
}
let mut attn = vec![vec![0.0f64; d]; n];
for head in 0..m.hp.n_head {
let off = head * hd;
for i in 0..n {
let raw: Vec<f64> = (0..n)
.map(|j| {
(0..hd).map(|c| q[i][off + c] * k[j][off + c]).sum::<f64>()
/ (hd as f64).sqrt()
})
.collect();
let mx = raw.iter().cloned().fold(f64::NEG_INFINITY, f64::max);
let ex: Vec<f64> = raw.iter().map(|s| (s - mx).exp()).collect();
let sum: f64 = ex.iter().sum();
for j in 0..n {
let p = ex[j] / sum;
for c in 0..hd {
attn[i][off + c] += p * v[j][off + c];
}
}
}
}
let mut next = Vec::new();
for i in 0..n {
let mut o = matvec(&wo, &attn[i]);
bias(&mut o, &layer.bo);
for (x, hv) in o.iter_mut().zip(&h[i]) {
*x += hv;
}
let x = ln(&o, &layer.attn_out_norm_w, &layer.attn_out_norm_b);
let mut up = matvec(&wu, &x);
bias(&mut up, &layer.ffn_up_b);
let act: Vec<f64> = up
.iter()
.map(|&u| {
const K: f64 = 0.797_884_560_802_865_4;
const C: f64 = 0.044_715;
0.5 * u * (1.0 + (K * (u + C * u * u * u)).tanh())
})
.collect();
let mut down = matvec(&wd, &act);
bias(&mut down, &layer.ffn_down_b);
for (dv, xv) in down.iter_mut().zip(&x) {
*dv += xv;
}
next.push(ln(&down, &layer.layer_out_norm_w, &layer.layer_out_norm_b));
}
h = next;
}
h.into_iter().flatten().collect()
}
#[test]
fn matches_an_independent_f64_transcription_of_the_graph() {
let m = fixture(3);
let tokens = [1u32, 7, 13, 4, 9, 2];
let got = m.encode_tokens(&tokens).unwrap();
let want = reference_forward(&m, &tokens);
assert_eq!(got.len(), want.len());
for (i, (g, w)) in got.iter().zip(&want).enumerate() {
assert!(
(*g as f64 - w).abs() < 2e-4,
"element {i}: {g} vs reference {w}"
);
}
}
#[test]
fn attention_is_bidirectional_not_causal() {
let m = fixture(2);
let a = m.encode_tokens(&[5u32, 6, 7, 8]).unwrap();
let b = m.encode_tokens(&[5u32, 6, 7, 19]).unwrap();
let moved: f32 = a[..D].iter().zip(&b[..D]).map(|(x, y)| (x - y).abs()).sum();
assert!(
moved > 1e-3,
"row 0 barely moved ({moved}) when the last token changed — \
attention is behaving causally"
);
}
#[test]
fn position_embeddings_make_the_same_token_differ_by_index() {
let m = fixture(1);
let out = m.encode_tokens(&[11u32, 11]).unwrap();
let delta: f32 = out[..D]
.iter()
.zip(&out[D..2 * D])
.map(|(x, y)| (x - y).abs())
.sum();
assert!(
delta > 1e-3,
"identical tokens gave identical rows: {delta}"
);
}
#[test]
fn the_last_op_is_a_mean_subtracting_layer_norm() {
let mut m = fixture(2);
let last = m.layers.last_mut().unwrap();
last.layer_out_norm_w = vec![1.0; D];
last.layer_out_norm_b = vec![0.0; D];
let out = m.encode_tokens(&[3u32, 4, 5]).unwrap();
for row in out.as_chunks::<D>().0 {
let mean: f32 = row.iter().sum::<f32>() / D as f32;
let var: f32 = row.iter().map(|v| (v - mean).powi(2)).sum::<f32>() / D as f32;
assert!(mean.abs() < 1e-4, "row mean {mean} is not zero");
assert!((var - 1.0).abs() < 1e-3, "row variance {var} is not one");
}
}
#[test]
fn refuses_an_empty_sequence_and_one_past_the_position_table() {
let m = fixture(1);
assert!(matches!(
m.encode_tokens(&[]),
Err(EncodeError::EmptySequence)
));
let long: Vec<u32> = (0..CTX as u32 + 1).map(|i| i % VOCAB as u32).collect();
let err = m.encode_tokens(&long).unwrap_err();
assert!(
matches!(err, EncodeError::TooLong { got, max, .. } if got == CTX + 1 && max == CTX)
);
assert!(matches!(
m.encode_tokens(&[VOCAB as u32]),
Err(EncodeError::TokenOutOfRange { .. })
));
}
#[test]
fn wrap_special_brackets_the_pieces_with_cls_and_sep() {
let m = fixture(1);
assert_eq!(m.wrap_special(&[7, 8]), vec![1, 7, 8, 2]);
assert_eq!(m.wrap_special(&[]), vec![1, 2]);
}
#[test]
fn the_pair_form_separates_the_two_halves_and_closes_the_second() {
let m = fixture(1);
assert_eq!(
m.wrap_special_pair(&[7, 8], &[9]).unwrap(),
vec![1, 7, 8, 2, 9, 2]
);
assert_eq!(m.wrap_special_pair(&[], &[]).unwrap(), vec![1, 2, 2]);
assert_ne!(
m.wrap_special_pair(&[7, 8], &[9]).unwrap(),
m.wrap_special(&[7, 8, 9])
);
}
#[test]
fn embed_tokens_pools_the_way_the_hparams_say() {
let m = fixture(2);
let tokens = [1u32, 9, 4, 2];
let hidden = m.encode_tokens(&tokens).unwrap();
assert_eq!(m.embed_tokens(&tokens).unwrap(), hidden[..D].to_vec());
}
}