use burn::module::Module;
use burn::nn::{
Embedding, EmbeddingConfig, Linear, LinearConfig, RmsNorm, RmsNormConfig, RotaryEncoding,
RotaryEncodingConfig,
};
use burn::tensor::{activation, backend::Backend, Int, Tensor, TensorData};
pub(crate) const ROPE_CACHE_LEN: usize = 4096;
#[derive(Clone, Debug)]
pub struct GemmaConfig {
pub vocab_size: usize, pub hidden_size: usize, pub intermediate_size: usize, pub num_layers: usize, pub num_heads: usize, pub num_kv_heads: usize, pub head_dim: usize, pub rope_theta: f32, pub rope_local_base_freq: f32, pub sliding_window: usize, pub sliding_window_pattern: usize, pub query_pre_attn_scalar: f64, pub rms_norm_eps: f64, }
impl GemmaConfig {
pub fn gemma_3_270m() -> Self {
Self {
vocab_size: 262_144,
hidden_size: 640,
intermediate_size: 2048,
num_layers: 18,
num_heads: 4,
num_kv_heads: 1,
head_dim: 256,
rope_theta: 1_000_000.0,
rope_local_base_freq: 10_000.0,
sliding_window: 512,
sliding_window_pattern: 6,
query_pre_attn_scalar: 256.0,
rms_norm_eps: 1e-6,
}
}
pub fn is_full_attention(&self, layer: usize) -> bool {
(layer + 1) % self.sliding_window_pattern == 0
}
pub fn layer_window(&self, layer: usize) -> Option<usize> {
if self.is_full_attention(layer) {
None
} else {
Some(self.sliding_window)
}
}
fn attn_scale(&self) -> f64 {
1.0 / self.query_pre_attn_scalar.sqrt()
}
}
pub(crate) fn attn_blocked(q_abs: usize, k_abs: usize, window: Option<usize>) -> bool {
k_abs > q_abs || window.is_some_and(|w| k_abs + w <= q_abs)
}
pub(crate) fn kept_cache_len(processed: usize, window: Option<usize>) -> usize {
match window {
Some(w) => processed.min(w - 1),
None => processed,
}
}
pub(crate) fn needs_mask(seq: usize, kv_len: usize, window: Option<usize>) -> bool {
seq > 1 || window.is_some_and(|w| kv_len > w)
}
fn trim_to_last<B: Backend>(t: Tensor<B, 4>, keep: usize) -> Tensor<B, 4> {
let [b, h, s, d] = t.dims();
t.slice([0..b, 0..h, s - keep..s, 0..d])
}
#[derive(Module, Debug)]
pub struct Mlp<B: Backend> {
gate_proj: Linear<B>,
up_proj: Linear<B>,
down_proj: Linear<B>,
}
impl<B: Backend> Mlp<B> {
fn init(cfg: &GemmaConfig, device: &B::Device) -> Self {
let lin = |i, o| LinearConfig::new(i, o).with_bias(false).init(device);
Self {
gate_proj: lin(cfg.hidden_size, cfg.intermediate_size),
up_proj: lin(cfg.hidden_size, cfg.intermediate_size),
down_proj: lin(cfg.intermediate_size, cfg.hidden_size),
}
}
fn forward(&self, x: Tensor<B, 3>) -> Tensor<B, 3> {
let gate = activation::gelu(self.gate_proj.forward(x.clone()));
let up = self.up_proj.forward(x);
self.down_proj.forward(gate * up)
}
}
#[derive(Module, Debug)]
pub struct Attention<B: Backend> {
q_proj: Linear<B>,
k_proj: Linear<B>,
v_proj: Linear<B>,
o_proj: Linear<B>,
q_norm: RmsNorm<B>,
k_norm: RmsNorm<B>,
}
impl<B: Backend> Attention<B> {
fn init(cfg: &GemmaConfig, device: &B::Device) -> Self {
let q_out = cfg.num_heads * cfg.head_dim;
let kv_out = cfg.num_kv_heads * cfg.head_dim;
let lin = |i, o| LinearConfig::new(i, o).with_bias(false).init(device);
let norm = || RmsNormConfig::new(cfg.head_dim).with_epsilon(cfg.rms_norm_eps).init(device);
Self {
q_proj: lin(cfg.hidden_size, q_out),
k_proj: lin(cfg.hidden_size, kv_out),
v_proj: lin(cfg.hidden_size, kv_out),
o_proj: lin(q_out, cfg.hidden_size),
q_norm: norm(),
k_norm: norm(),
}
}
#[allow(clippy::too_many_arguments)]
fn forward(
&self,
x: Tensor<B, 3>,
cfg: &GemmaConfig,
rope: &RotaryEncoding<B>,
mask: Option<Tensor<B, 4>>,
cache: &mut Option<(Tensor<B, 4>, Tensor<B, 4>)>,
offset: usize,
window: Option<usize>,
) -> Tensor<B, 3> {
let [batch, seq, _] = x.dims();
let (h, kv, hd) = (cfg.num_heads, cfg.num_kv_heads, cfg.head_dim);
let q = self.q_proj.forward(x.clone()).reshape([batch, seq, h, hd]);
let k = self.k_proj.forward(x.clone()).reshape([batch, seq, kv, hd]);
let v = self.v_proj.forward(x).reshape([batch, seq, kv, hd]);
let q = self.q_norm.forward(q).swap_dims(1, 2); let k = self.k_norm.forward(k).swap_dims(1, 2); let v = v.swap_dims(1, 2);
let q = rope.apply(q, offset);
let k = rope.apply(k, offset);
let (k_all, v_all) = match cache.take() {
Some((ck, cv)) => (Tensor::cat(vec![ck, k], 2), Tensor::cat(vec![cv, v], 2)),
None => (k, v),
};
let kv_len = k_all.dims()[2];
let keep = kept_cache_len(offset + seq, window);
*cache = Some(if keep < kv_len {
(
trim_to_last(k_all.clone(), keep),
trim_to_last(v_all.clone(), keep),
)
} else {
(k_all.clone(), v_all.clone())
});
let k_all = k_all.repeat_dim(1, h / kv);
let v_all = v_all.repeat_dim(1, h / kv);
let mut scores = q
.matmul(k_all.swap_dims(2, 3))
.mul_scalar(cfg.attn_scale());
if let Some(mask) = mask {
scores = scores + mask;
}
let probs = activation::softmax(scores, 3);
let ctx = probs.matmul(v_all);
let ctx = ctx.swap_dims(1, 2).reshape([batch, seq, h * hd]);
self.o_proj.forward(ctx)
}
}
#[derive(Module, Debug)]
pub struct DecoderLayer<B: Backend> {
input_layernorm: RmsNorm<B>,
self_attn: Attention<B>,
post_attention_layernorm: RmsNorm<B>,
pre_feedforward_layernorm: RmsNorm<B>,
mlp: Mlp<B>,
post_feedforward_layernorm: RmsNorm<B>,
}
impl<B: Backend> DecoderLayer<B> {
fn init(cfg: &GemmaConfig, device: &B::Device) -> Self {
let norm = || {
RmsNormConfig::new(cfg.hidden_size)
.with_epsilon(cfg.rms_norm_eps)
.init(device)
};
Self {
input_layernorm: norm(),
self_attn: Attention::init(cfg, device),
post_attention_layernorm: norm(),
pre_feedforward_layernorm: norm(),
mlp: Mlp::init(cfg, device),
post_feedforward_layernorm: norm(),
}
}
#[allow(clippy::too_many_arguments)]
fn forward(
&self,
x: Tensor<B, 3>,
cfg: &GemmaConfig,
rope: &RotaryEncoding<B>,
mask: Option<Tensor<B, 4>>,
cache: &mut Option<(Tensor<B, 4>, Tensor<B, 4>)>,
offset: usize,
window: Option<usize>,
) -> Tensor<B, 3> {
let normed = self.input_layernorm.forward(x.clone());
let attn = self
.self_attn
.forward(normed, cfg, rope, mask, cache, offset, window);
let h = x + self.post_attention_layernorm.forward(attn);
let normed = self.pre_feedforward_layernorm.forward(h.clone());
let ff = self.mlp.forward(normed);
h + self.post_feedforward_layernorm.forward(ff)
}
}
#[derive(Module, Debug)]
pub struct GemmaModel<B: Backend> {
embed: Embedding<B>,
layers: Vec<DecoderLayer<B>>,
norm: RmsNorm<B>,
rope_global: RotaryEncoding<B>,
rope_local: RotaryEncoding<B>,
#[module(skip)]
config: GemmaConfig,
}
impl<B: Backend> GemmaModel<B> {
pub fn init(cfg: GemmaConfig, device: &B::Device) -> Self {
let layers = (0..cfg.num_layers)
.map(|_| DecoderLayer::init(&cfg, device))
.collect();
let rope = |theta| {
RotaryEncodingConfig::new(ROPE_CACHE_LEN, cfg.head_dim)
.with_theta(theta)
.init(device)
};
Self {
embed: EmbeddingConfig::new(cfg.vocab_size, cfg.hidden_size).init(device),
layers,
norm: RmsNormConfig::new(cfg.hidden_size)
.with_epsilon(cfg.rms_norm_eps)
.init(device),
rope_global: rope(cfg.rope_theta),
rope_local: rope(cfg.rope_local_base_freq),
config: cfg,
}
}
pub fn forward(&self, tokens: Tensor<B, 2, Int>) -> Tensor<B, 3> {
let mut cache = self.new_cache();
self.forward_cached(tokens, &mut cache)
}
pub fn new_cache(&self) -> KvCache<B> {
KvCache {
layers: vec![None; self.config.num_layers],
len: 0,
}
}
pub fn forward_cached(&self, tokens: Tensor<B, 2, Int>, cache: &mut KvCache<B>) -> Tensor<B, 3> {
let cfg = &self.config;
let [batch, seq] = tokens.dims();
let device = tokens.device();
let offset = cache.len;
let scale = (cfg.hidden_size as f64).sqrt();
let mut x = self.embed.forward(tokens).mul_scalar(scale);
let build_mask = |window: Option<usize>| -> Option<Tensor<B, 4>> {
let kv_len = kept_cache_len(offset, window) + seq;
needs_mask(seq, kv_len, window).then(|| {
let kv_start = offset + seq - kv_len; let mut rows = vec![0f32; seq * kv_len];
for i in 0..seq {
for j in 0..kv_len {
if attn_blocked(offset + i, kv_start + j, window) {
rows[i * kv_len + j] = f32::NEG_INFINITY;
}
}
}
Tensor::<B, 1>::from_data(TensorData::from(rows.as_slice()), &device)
.reshape([1, 1, seq, kv_len])
})
};
let mask_global = build_mask(None);
let mask_sliding = build_mask(Some(cfg.sliding_window));
let _ = batch;
for (i, layer) in self.layers.iter().enumerate() {
let (rope, mask, window) = if cfg.is_full_attention(i) {
(&self.rope_global, mask_global.clone(), None)
} else {
(
&self.rope_local,
mask_sliding.clone(),
Some(cfg.sliding_window),
)
};
x = layer.forward(x, cfg, rope, mask, &mut cache.layers[i], offset, window);
}
cache.len += seq;
let x = self.norm.forward(x);
let embed_t = self.embed.weight.val().transpose(); x.matmul(embed_t.unsqueeze::<3>())
}
}
pub struct KvCache<B: Backend> {
layers: Vec<Option<(Tensor<B, 4>, Tensor<B, 4>)>>,
len: usize,
}
impl<B: Backend> KvCache<B> {
pub fn len(&self) -> usize {
self.len
}
pub fn is_empty(&self) -> bool {
self.len == 0
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn sliding_mask_matches_hf_truth_table() {
#[rustfmt::skip]
let allowed = [
[true, false, false, false, false],
[true, true, false, false, false],
[true, true, true, false, false],
[false, true, true, true, false],
[false, false, true, true, true ],
];
for (q, row) in allowed.iter().enumerate() {
for (k, &want) in row.iter().enumerate() {
assert_eq!(!attn_blocked(q, k, Some(3)), want, "q={q} k={k}");
}
}
}
#[test]
fn global_mask_is_plain_causal() {
for q in 0..8 {
for k in 0..8 {
assert_eq!(attn_blocked(q, k, None), k > q, "q={q} k={k}");
}
}
}
#[test]
fn windowing_is_noop_below_the_window() {
let w = GemmaConfig::gemma_3_270m().sliding_window;
for q in 0..w {
for k in 0..w {
assert_eq!(attn_blocked(q, k, Some(w)), attn_blocked(q, k, None));
}
}
assert!(attn_blocked(w, 0, Some(w)));
assert!(!attn_blocked(w, 0, None));
let q = 3 * w + 7;
let attended = (0..=q).filter(|&k| !attn_blocked(q, k, Some(w))).count();
assert_eq!(attended, w);
assert!(!attn_blocked(q, q, Some(w)), "self is always attended");
assert!(!attn_blocked(q, q + 1 - w, Some(w)), "oldest in-window key");
assert!(attn_blocked(q, q - w, Some(w)), "first out-of-window key");
}
#[test]
fn layer_pattern_matches_config_layer_types() {
let cfg = GemmaConfig::gemma_3_270m();
let full: Vec<usize> = (0..cfg.num_layers)
.filter(|&l| cfg.is_full_attention(l))
.collect();
assert_eq!(full, [5, 11, 17]);
for l in 0..cfg.num_layers {
let want = if full.contains(&l) { None } else { Some(512) };
assert_eq!(cfg.layer_window(l), want, "layer {l}");
}
}
#[test]
fn cache_trim_caps_and_keeps_decode_mask_free() {
let w = 512;
assert_eq!(kept_cache_len(0, Some(w)), 0);
assert_eq!(kept_cache_len(w - 1, Some(w)), w - 1);
assert_eq!(kept_cache_len(w, Some(w)), w - 1); assert_eq!(kept_cache_len(10_000, Some(w)), w - 1);
assert_eq!(kept_cache_len(10_000, None), 10_000); for offset in [0, 1, w - 1, w, w + 1, 4 * w] {
assert!(!needs_mask(1, kept_cache_len(offset, Some(w)) + 1, Some(w)));
assert!(!needs_mask(1, offset + 1, None));
}
assert!(needs_mask(1, w + 1, Some(w))); assert!(needs_mask(2, 2, None)); assert!(needs_mask(2, 2, Some(w)));
}
#[test]
fn trim_and_mask_builder_agree_on_cache_length() {
let w = 512;
for window in [Some(w), None] {
let mut cached = 0usize; let mut offset = 0usize; for seq in [600usize, 1, 1, 1, 300, 1] {
assert_eq!(cached, kept_cache_len(offset, window), "pre @{offset}");
let kv_len = cached + seq;
let keep = kept_cache_len(offset + seq, window);
assert!(keep <= kv_len, "trim would underflow @{offset}+{seq}");
cached = keep; offset += seq;
}
}
}
}