#![allow(clippy::needless_range_loop)]
use super::calib;
use super::moe;
use super::nn;
use super::rswa::{self, BatchedRingCache, RingCache};
use super::tensor::{Mat, QInt4, QInt8, WeightLayout};
use super::weights::{DType, Weights};
use crate::error::{FocrError, FocrResult};
use crate::quant::calib as quant_calib;
use crate::simd;
use rayon::prelude::*;
#[cfg(not(target_arch = "wasm32"))]
use std::time::Instant;
#[cfg(target_arch = "wasm32")]
use web_time::Instant;
fn checked_shape_mul(context: &str, lhs: usize, rhs: usize, expression: &str) -> FocrResult<usize> {
lhs.checked_mul(rhs).ok_or_else(|| {
FocrError::Other(anyhow::anyhow!(
"{context}: usize overflow computing {expression} ({lhs} * {rhs})"
))
})
}
fn checked_mat_len(context: &str, x: &Mat) -> FocrResult<usize> {
let expected = checked_shape_mul(context, x.rows, x.cols, "rows*cols")?;
if x.data.len() != expected {
return Err(FocrError::Other(anyhow::anyhow!(
"{context}: data len {} != rows*cols {} for shape [{}, {}]",
x.data.len(),
expected,
x.rows,
x.cols
)));
}
Ok(expected)
}
pub mod config {
pub const HIDDEN_SIZE: usize = 1280;
pub const INTERMEDIATE_SIZE: usize = 6848;
pub const NUM_HIDDEN_LAYERS: usize = 12;
pub const NUM_ATTENTION_HEADS: usize = 10;
pub const NUM_KEY_VALUE_HEADS: usize = 10;
pub const HEAD_DIM: usize = 128;
pub const VOCAB_SIZE: usize = 129280;
pub const ROPE_THETA: f32 = 10000.0;
pub const RMS_NORM_EPS: f32 = 1e-6;
pub const FIRST_K_DENSE_REPLACE: usize = 1;
pub const RING_WINDOW: usize = 128;
pub const BOS_TOKEN_ID: u32 = 0;
pub const EOS_TOKEN_ID: u32 = 1;
}
pub fn embed_tokens(table: &[f32], vocab: usize, hidden: usize, ids: &[u32]) -> FocrResult<Mat> {
let expected_table_len = checked_shape_mul("embed_tokens", vocab, hidden, "vocab*hidden")?;
if table.len() != expected_table_len {
return Err(FocrError::Other(anyhow::anyhow!(
"embed_tokens: table len {} != vocab*hidden {}",
table.len(),
expected_table_len
)));
}
let seq = ids.len();
let out_len = checked_shape_mul("embed_tokens", seq, hidden, "seq*hidden")?;
let mut out = vec![0.0f32; out_len];
for (t, &id) in ids.iter().enumerate() {
let row = id as usize;
if row >= vocab {
return Err(FocrError::Other(anyhow::anyhow!(
"embed_tokens: id {row} out of range (vocab {vocab})"
)));
}
let src = &table[row * hidden..(row + 1) * hidden];
out[t * hidden..(t + 1) * hidden].copy_from_slice(src);
}
Ok(Mat::from_vec(seq, hidden, out))
}
#[derive(Debug, Clone)]
pub struct RopeTable {
pub cos: Vec<f32>,
pub sin: Vec<f32>,
pub head_dim: usize,
}
impl RopeTable {
#[must_use]
pub fn build(position_ids: &[usize], head_dim: usize, theta: f32) -> Self {
assert!(
head_dim.is_multiple_of(2),
"RopeTable: head_dim must be even"
);
let half = head_dim / 2;
let seq = position_ids.len();
let table_len = seq.checked_mul(head_dim);
assert!(
table_len.is_some(),
"RopeTable: seq*head_dim overflow ({seq} * {head_dim})"
);
let table_len = table_len.unwrap_or(0);
if seq == 0 {
return Self {
cos: Vec::new(),
sin: Vec::new(),
head_dim,
};
}
let inv_freq: Vec<f32> = (0..half)
.map(|i| {
let exponent = (2 * i) as f64 / head_dim as f64;
(1.0 / (theta as f64).powf(exponent)) as f32
})
.collect();
let mut cos = vec![0.0f32; table_len];
let mut sin = vec![0.0f32; table_len];
for (p_idx, &pos) in position_ids.iter().enumerate() {
let base = p_idx * head_dim;
for i in 0..half {
let angle = pos as f32 * inv_freq[i];
let (s, c) = angle.sin_cos();
cos[base + i] = c;
cos[base + half + i] = c;
sin[base + i] = s;
sin[base + half + i] = s;
}
}
Self { cos, sin, head_dim }
}
}
pub fn apply_rope(x: &mut Mat, rope: &RopeTable) -> FocrResult<()> {
checked_mat_len("apply_rope x", x)?;
let head_dim = rope.head_dim;
if head_dim == 0 || !x.cols.is_multiple_of(head_dim) {
return Err(FocrError::Other(anyhow::anyhow!(
"apply_rope: cols {} not a multiple of head_dim {head_dim}",
x.cols
)));
}
let seq = x.rows;
if rope.cos.len() != seq * head_dim {
return Err(FocrError::Other(anyhow::anyhow!(
"apply_rope: rope built for {} positions, x has {seq} rows",
rope.cos.len() / head_dim
)));
}
let num_heads = x.cols / head_dim;
let half = head_dim / 2;
for t in 0..seq {
let rope_base = t * head_dim;
let row = x.row_mut(t);
for h in 0..num_heads {
let hb = h * head_dim;
for i in 0..half {
let a = row[hb + i]; let b = row[hb + half + i]; let cos_a = rope.cos[rope_base + i];
let sin_a = rope.sin[rope_base + i];
row[hb + i] = a * cos_a - b * sin_a;
row[hb + half + i] = b * cos_a + a * sin_a;
}
}
}
Ok(())
}
pub fn dense_mlp(
x: &Mat,
gate_w: &[f32],
up_w: &[f32],
down_w: &[f32],
hidden: usize,
inter: usize,
) -> FocrResult<Mat> {
if x.cols != hidden {
return Err(FocrError::Other(anyhow::anyhow!(
"dense_mlp: x.cols {} != hidden {hidden}",
x.cols
)));
}
let mut gate = linear_no_bias(x, gate_w, hidden, inter)?;
let up = linear_no_bias(x, up_w, hidden, inter)?;
nn::silu(&mut gate);
for (g, u) in gate.data.iter_mut().zip(up.data.iter()) {
*g *= *u;
}
linear_no_bias(&gate, down_w, inter, hidden)
}
pub(crate) fn linear_no_bias(x: &Mat, w: &[f32], in_: usize, out: usize) -> FocrResult<Mat> {
checked_mat_len("linear_no_bias x", x)?;
if x.cols != in_ {
return Err(FocrError::Other(anyhow::anyhow!(
"linear_no_bias: x.cols {} != in {in_}",
x.cols
)));
}
let expected_weight_len = checked_shape_mul("linear_no_bias", out, in_, "out*in")?;
if w.len() != expected_weight_len {
return Err(FocrError::Other(anyhow::anyhow!(
"linear_no_bias: w len {} != out*in {}",
w.len(),
expected_weight_len
)));
}
let mut wt = vec![0.0f32; expected_weight_len];
for i in 0..in_ {
let dst = &mut wt[i * out..(i + 1) * out];
for (o, slot) in dst.iter_mut().enumerate() {
*slot = w[o * in_ + i];
}
}
let w_mat = Mat::from_vec(in_, out, wt);
nn::matmul(x, &w_mat)
}
pub fn norm_and_lm_head(
hidden: &Mat,
norm_w: &[f32],
head_w: &[f32],
vocab: usize,
eps: f32,
) -> FocrResult<Mat> {
let normed = nn::rms_norm(hidden, Some(norm_w), eps)?;
lm_head_proj(&normed, head_w, vocab)
}
pub(crate) fn norm_and_lm_head_pretransposed(
hidden: &Mat,
norm_w: &[f32],
head_wt: &Mat,
eps: f32,
) -> FocrResult<Mat> {
let normed = nn::rms_norm(hidden, Some(norm_w), eps)?;
nn::matmul(&normed, head_wt)
}
pub fn lm_head_proj(hidden: &Mat, head_w: &[f32], vocab: usize) -> FocrResult<Mat> {
let h = hidden.cols;
linear_no_bias(hidden, head_w, h, vocab)
}
#[derive(Debug, Clone)]
pub struct LayerWeights<'a> {
pub input_ln: &'a [f32],
pub post_attn_ln: &'a [f32],
pub q_proj: &'a [f32],
pub k_proj: &'a [f32],
pub v_proj: &'a [f32],
pub o_proj: &'a [f32],
pub gate_w: &'a [f32],
pub up_w: &'a [f32],
pub down_w: &'a [f32],
}
pub fn qkv_with_rope(
normed: &Mat,
lw: &LayerWeights<'_>,
rope: &RopeTable,
hidden: usize,
qkv_dim: usize,
) -> FocrResult<(Mat, Mat, Mat)> {
let mut q = linear_no_bias(normed, lw.q_proj, hidden, qkv_dim)?;
let mut k = linear_no_bias(normed, lw.k_proj, hidden, qkv_dim)?;
let v = linear_no_bias(normed, lw.v_proj, hidden, qkv_dim)?;
apply_rope(&mut q, rope)?;
apply_rope(&mut k, rope)?;
Ok((q, k, v))
}
pub fn attn_output_proj(
context: &Mat,
o_proj: &[f32],
hidden: usize,
qkv_dim: usize,
) -> FocrResult<Mat> {
linear_no_bias(context, o_proj, qkv_dim, hidden)
}
#[allow(clippy::too_many_arguments)]
pub fn layer_forward<A, M>(
x: &Mat,
lw: &LayerWeights<'_>,
rope: &RopeTable,
hidden: usize,
qkv_dim: usize,
eps: f32,
attn_context: A,
mlp: M,
) -> FocrResult<Mat>
where
A: FnOnce(&Mat, &Mat, &Mat) -> FocrResult<Mat>,
M: FnOnce(&Mat) -> FocrResult<Mat>,
{
let normed = nn::rms_norm(x, Some(lw.input_ln), eps)?;
let (q, k, v) = qkv_with_rope(&normed, lw, rope, hidden, qkv_dim)?;
let context = attn_context(&q, &k, &v)?;
let attn_out = attn_output_proj(&context, lw.o_proj, hidden, qkv_dim)?;
let h = add_residual(x, &attn_out)?;
let normed2 = nn::rms_norm(&h, Some(lw.post_attn_ln), eps)?;
let mlp_out = mlp(&normed2)?;
add_residual(&h, &mlp_out)
}
pub fn add_residual(a: &Mat, b: &Mat) -> FocrResult<Mat> {
checked_mat_len("add_residual lhs", a)?;
checked_mat_len("add_residual rhs", b)?;
if a.shape() != b.shape() {
return Err(FocrError::Other(anyhow::anyhow!(
"add_residual: shape mismatch {:?} vs {:?}",
a.shape(),
b.shape()
)));
}
let data = a
.data
.iter()
.zip(b.data.iter())
.map(|(x, y)| x + y)
.collect::<Vec<_>>();
Ok(Mat::from_vec(a.rows, a.cols, data))
}
pub fn prefill_attention(
q: &Mat,
k: &Mat,
v: &Mat,
num_heads: usize,
head_dim: usize,
) -> FocrResult<Mat> {
checked_mat_len("prefill_attention q", q)?;
checked_mat_len("prefill_attention k", k)?;
checked_mat_len("prefill_attention v", v)?;
let dim = checked_shape_mul(
"prefill_attention",
num_heads,
head_dim,
"num_heads*head_dim",
)?;
let seq = q.rows;
if q.cols != dim || k.cols != dim || v.cols != dim {
return Err(FocrError::Other(anyhow::anyhow!(
"prefill_attention: q/k/v cols ({},{},{}) != num_heads*head_dim {dim}",
q.cols,
k.cols,
v.cols
)));
}
if k.rows != seq || v.rows != seq {
return Err(FocrError::Other(anyhow::anyhow!(
"prefill_attention: q/k/v rows disagree ({}, {}, {})",
seq,
k.rows,
v.rows
)));
}
let span = seq * head_dim;
let mut qh = vec![0.0f32; num_heads * span];
let mut kh = vec![0.0f32; num_heads * span];
let mut vh = vec![0.0f32; num_heads * span];
for s in 0..seq {
let (qr, kr, vr) = (q.row(s), k.row(s), v.row(s));
for h in 0..num_heads {
let src = h * head_dim;
let dst = h * span + s * head_dim;
qh[dst..dst + head_dim].copy_from_slice(&qr[src..src + head_dim]);
kh[dst..dst + head_dim].copy_from_slice(&kr[src..src + head_dim]);
vh[dst..dst + head_dim].copy_from_slice(&vr[src..src + head_dim]);
}
}
let scale = 1.0f32 / (head_dim as f32).sqrt();
let ctx = nn::sdpa(
&qh, &kh, &vh, num_heads, seq, seq, head_dim, head_dim, scale, true,
);
let expected = num_heads * span;
if ctx.len() != expected {
return Err(FocrError::Other(anyhow::anyhow!(
"prefill_attention: sdpa context len {} != expected {expected}",
ctx.len()
)));
}
let mut out = Mat::zeros(seq, dim);
for h in 0..num_heads {
for s in 0..seq {
let src = h * span + s * head_dim;
let dst = s * dim + h * head_dim;
out.data[dst..dst + head_dim].copy_from_slice(&ctx[src..src + head_dim]);
}
}
Ok(out)
}
pub fn lm_head(weights: &Weights, hidden: &Mat) -> FocrResult<Mat> {
let norm_w = weights.vec("model.norm.weight")?;
let head = weights.mat("lm_head.weight")?;
norm_and_lm_head(
hidden,
&norm_w,
&head.data,
config::VOCAB_SIZE,
config::RMS_NORM_EPS,
)
}
pub fn forward(weights: &Weights, inputs_embeds: &Mat) -> FocrResult<Mat> {
checked_mat_len("decoder::forward inputs_embeds", inputs_embeds)?;
let hidden = config::HIDDEN_SIZE;
let qkv_dim = checked_shape_mul(
"decoder::forward",
config::NUM_ATTENTION_HEADS,
config::HEAD_DIM,
"num_heads*head_dim",
)?;
let eps = config::RMS_NORM_EPS;
if inputs_embeds.cols != hidden {
return Err(FocrError::Other(anyhow::anyhow!(
"decoder::forward: inputs_embeds cols {} != hidden {hidden}",
inputs_embeds.cols
)));
}
let seq = inputs_embeds.rows;
let positions: Vec<usize> = (0..seq).collect();
let rope = RopeTable::build(&positions, config::HEAD_DIM, config::ROPE_THETA);
let mut x = inputs_embeds.clone();
for layer in 0..config::NUM_HIDDEN_LAYERS {
let prefix = format!("model.layers.{layer}");
let input_ln = weights.vec(&format!("{prefix}.input_layernorm.weight"))?;
let post_attn_ln = weights.vec(&format!("{prefix}.post_attention_layernorm.weight"))?;
let q_proj = weights.mat(&format!("{prefix}.self_attn.q_proj.weight"))?;
let k_proj = weights.mat(&format!("{prefix}.self_attn.k_proj.weight"))?;
let v_proj = weights.mat(&format!("{prefix}.self_attn.v_proj.weight"))?;
let o_proj = weights.mat(&format!("{prefix}.self_attn.o_proj.weight"))?;
let lw = LayerWeights {
input_ln: &input_ln,
post_attn_ln: &post_attn_ln,
q_proj: &q_proj.data,
k_proj: &k_proj.data,
v_proj: &v_proj.data,
o_proj: &o_proj.data,
gate_w: &[],
up_w: &[],
down_w: &[],
};
x = layer_forward(
&x,
&lw,
&rope,
hidden,
qkv_dim,
eps,
|q, k, v| prefill_attention(q, k, v, config::NUM_ATTENTION_HEADS, config::HEAD_DIM),
|normed| {
if layer < config::FIRST_K_DENSE_REPLACE {
moe::dense_forward(weights, normed)
} else {
moe::forward(weights, normed, layer)
}
},
)?;
}
Ok(x)
}
fn token_major_to_head_major(
k: &Mat,
v: &Mat,
seq: usize,
num_heads: usize,
head_dim: usize,
) -> FocrResult<(Vec<f32>, Vec<f32>)> {
let span = checked_shape_mul("token_major_to_head_major", seq, head_dim, "seq*head_dim")?;
let total = checked_shape_mul(
"token_major_to_head_major",
num_heads,
span,
"num_heads*seq*head_dim",
)?;
let mut kh = vec![0.0f32; total];
let mut vh = vec![0.0f32; total];
for s in 0..seq {
let (kr, vr) = (k.row(s), v.row(s));
for h in 0..num_heads {
let src = h * head_dim;
let dst = h * span + s * head_dim;
kh[dst..dst + head_dim].copy_from_slice(&kr[src..src + head_dim]);
vh[dst..dst + head_dim].copy_from_slice(&vr[src..src + head_dim]);
}
}
Ok((kh, vh))
}
pub struct DecoderWeightCache {
layers: Vec<CachedLayer>,
final_norm: Vec<f32>,
lm_head: Vec<f32>,
}
struct CachedLayer {
input_ln: Vec<f32>,
post_attn_ln: Vec<f32>,
q_proj: Vec<f32>,
k_proj: Vec<f32>,
v_proj: Vec<f32>,
o_proj: Vec<f32>,
mlp: CachedMlpI8,
}
impl DecoderWeightCache {
pub fn build(weights: &Weights) -> FocrResult<Self> {
let mut layers = Vec::with_capacity(config::NUM_HIDDEN_LAYERS);
for layer in 0..config::NUM_HIDDEN_LAYERS {
let prefix = format!("model.layers.{layer}");
let input_ln = weights.vec(&format!("{prefix}.input_layernorm.weight"))?;
let post_attn_ln = weights.vec(&format!("{prefix}.post_attention_layernorm.weight"))?;
let q_proj = weights
.mat(&format!("{prefix}.self_attn.q_proj.weight"))?
.data;
let k_proj = weights
.mat(&format!("{prefix}.self_attn.k_proj.weight"))?
.data;
let v_proj = weights
.mat(&format!("{prefix}.self_attn.v_proj.weight"))?
.data;
let o_proj = weights
.mat(&format!("{prefix}.self_attn.o_proj.weight"))?
.data;
let mlp = if layer < config::FIRST_K_DENSE_REPLACE {
let p = format!("{prefix}.mlp");
let inter = moe::config::DENSE_INTERMEDIATE_SIZE;
CachedMlpI8::Dense {
gate: quant_ffn_loaded(weights, &format!("{p}.gate_proj.weight"), inter)?,
up: quant_ffn_loaded(weights, &format!("{p}.up_proj.weight"), inter)?,
down: quant_ffn_loaded(
weights,
&format!("{p}.down_proj.weight"),
config::HIDDEN_SIZE,
)?,
}
} else {
let p = format!("{prefix}.mlp");
let gate = weights.mat(&format!("{p}.gate.weight"))?.data;
let inter = moe::config::MOE_INTERMEDIATE_SIZE;
let mut experts = Vec::with_capacity(moe::config::N_ROUTED_EXPERTS);
for e in 0..moe::config::N_ROUTED_EXPERTS {
experts.push([
quant_ffn_loaded(
weights,
&format!("{p}.experts.{e}.gate_proj.weight"),
inter,
)?,
quant_ffn_loaded(
weights,
&format!("{p}.experts.{e}.up_proj.weight"),
inter,
)?,
quant_ffn_loaded(
weights,
&format!("{p}.experts.{e}.down_proj.weight"),
config::HIDDEN_SIZE,
)?,
]);
}
let shared_inter = moe::config::SHARED_INTERMEDIATE_SIZE;
let shared = [
quant_ffn_loaded(
weights,
&format!("{p}.shared_experts.gate_proj.weight"),
shared_inter,
)?,
quant_ffn_loaded(
weights,
&format!("{p}.shared_experts.up_proj.weight"),
shared_inter,
)?,
quant_ffn_loaded(
weights,
&format!("{p}.shared_experts.down_proj.weight"),
config::HIDDEN_SIZE,
)?,
];
CachedMlpI8::Moe {
gate,
experts,
shared,
}
};
layers.push(CachedLayer {
input_ln,
post_attn_ln,
q_proj,
k_proj,
v_proj,
o_proj,
mlp,
});
}
let final_norm = weights.vec("model.norm.weight")?;
let lm_head = weights.mat("lm_head.weight")?.data;
Ok(Self {
layers,
final_norm,
lm_head,
})
}
}
fn cached_layer_weights(cl: &CachedLayer) -> LayerWeights<'_> {
LayerWeights {
input_ln: &cl.input_ln,
post_attn_ln: &cl.post_attn_ln,
q_proj: &cl.q_proj,
k_proj: &cl.k_proj,
v_proj: &cl.v_proj,
o_proj: &cl.o_proj,
gate_w: &[],
up_w: &[],
down_w: &[],
}
}
fn cached_mlp(mlp: &CachedMlpI8, normed: &Mat, calib_layer: Option<usize>) -> FocrResult<Mat> {
prefill_mlp_i8(mlp, normed, calib_layer)
}
pub mod prof {
use std::sync::OnceLock;
use std::sync::atomic::{AtomicU64, Ordering};
pub static LMHEAD_NS: AtomicU64 = AtomicU64::new(0);
pub static ATTN_NS: AtomicU64 = AtomicU64::new(0);
pub static EXPERTS_NS: AtomicU64 = AtomicU64::new(0);
pub static ROUTE_NS: AtomicU64 = AtomicU64::new(0);
pub fn enabled() -> bool {
static FLAG: OnceLock<bool> = OnceLock::new();
*FLAG.get_or_init(|| std::env::var_os("FOCR_PROFILE_DECODE").is_some())
}
#[inline]
pub fn add(c: &AtomicU64, ns: u64) {
c.fetch_add(ns, Ordering::Relaxed);
}
pub fn reset() {
for c in [&LMHEAD_NS, &ATTN_NS, &EXPERTS_NS, &ROUTE_NS] {
c.store(0, Ordering::Relaxed);
}
}
pub fn snapshot_ms() -> (f64, f64, f64, f64) {
let ms = |c: &AtomicU64| c.load(Ordering::Relaxed) as f64 / 1e6;
(ms(&LMHEAD_NS), ms(&ATTN_NS), ms(&EXPERTS_NS), ms(&ROUTE_NS))
}
}
#[inline]
fn dot_f32(a: &[f32], b: &[f32]) -> f32 {
debug_assert_eq!(a.len(), b.len());
let n = a.len();
let mut acc = [0.0f32; 8];
let chunks = n / 8;
for c in 0..chunks {
let base = c * 8;
for (l, slot) in acc.iter_mut().enumerate() {
*slot += a[base + l] * b[base + l];
}
}
let mut s = ((acc[0] + acc[1]) + (acc[2] + acc[3])) + ((acc[4] + acc[5]) + (acc[6] + acc[7]));
for i in (chunks * 8)..n {
s += a[i] * b[i];
}
s
}
fn gemv(x: &[f32], w: &[f32], n: usize, k: usize) -> Vec<f32> {
debug_assert_eq!(x.len(), k);
debug_assert_eq!(w.len(), n * k);
let mut y = vec![0.0f32; n];
y.par_chunks_mut(64).enumerate().for_each(|(blk, ys)| {
let base = blk * 64;
for (j, slot) in ys.iter_mut().enumerate() {
let o = base + j;
*slot = dot_f32(x, &w[o * k..o * k + k]);
}
});
y
}
#[inline]
fn silu(x: f32) -> f32 {
x / (1.0 + (-x).exp())
}
fn decode_mlp(mlp: &CachedMlpI8, normed: &Mat, calib_layer: Option<usize>) -> FocrResult<Vec<f32>> {
decode_mlp_i8(mlp, normed, calib_layer)
}
pub fn lm_head_cached(wc: &DecoderWeightCache, hidden: &Mat) -> FocrResult<Mat> {
if hidden.rows != 1 {
return Err(FocrError::Other(anyhow::anyhow!(
"decoder::lm_head_cached: expected a single decode row, got {} rows",
hidden.rows
)));
}
let t = prof::enabled().then(Instant::now);
let normed = nn::rms_norm(hidden, Some(&wc.final_norm), config::RMS_NORM_EPS)?;
let row = normed.row(0);
if calib::enabled() {
calib::record_row(quant_calib::LM_HEAD_IN_KEY, row);
}
let logits = if lmhead_shard_enabled() {
gemv_sharded(
row,
&wc.lm_head,
config::VOCAB_SIZE,
config::HIDDEN_SIZE,
lmhead_shard_tiles(),
)
} else {
gemv(row, &wc.lm_head, config::VOCAB_SIZE, config::HIDDEN_SIZE)
};
if let Some(t) = t {
prof::add(&prof::LMHEAD_NS, t.elapsed().as_nanos() as u64);
}
Ok(Mat::from_vec(1, config::VOCAB_SIZE, logits))
}
const PREFILL_CHUNK_ENV: &str = "FOCR_PREFILL_CHUNK";
const DEFAULT_PREFILL_CHUNK: usize = 256;
#[must_use]
pub fn prefill_chunk_size() -> Option<usize> {
static SIZE: std::sync::OnceLock<Option<usize>> = std::sync::OnceLock::new();
*SIZE.get_or_init(|| {
std::env::var_os(PREFILL_CHUNK_ENV)?;
Some(
std::env::var(PREFILL_CHUNK_ENV)
.ok()
.and_then(|v| v.trim().parse::<usize>().ok())
.filter(|&n| n > 0)
.unwrap_or(DEFAULT_PREFILL_CHUNK),
)
})
}
fn prefill_chunk_bounds(seq: usize, chunk: usize) -> Vec<(usize, usize)> {
let chunk = chunk.max(1);
let mut bounds = Vec::with_capacity(seq.div_ceil(chunk));
let mut c0 = 0usize;
while c0 < seq {
let c1 = (c0 + chunk).min(seq);
bounds.push((c0, c1));
c0 = c1;
}
bounds
}
pub fn chunk_prefill_attention(
q_chunk: &Mat,
k_prefix: &Mat,
v_prefix: &Mat,
num_heads: usize,
head_dim: usize,
c0: usize,
) -> FocrResult<Mat> {
checked_mat_len("chunk_prefill_attention q_chunk", q_chunk)?;
checked_mat_len("chunk_prefill_attention k_prefix", k_prefix)?;
checked_mat_len("chunk_prefill_attention v_prefix", v_prefix)?;
let dim = checked_shape_mul(
"chunk_prefill_attention",
num_heads,
head_dim,
"num_heads*head_dim",
)?;
let c1 = k_prefix.rows;
let cs = q_chunk.rows;
if q_chunk.cols != dim || k_prefix.cols != dim || v_prefix.cols != dim {
return Err(FocrError::Other(anyhow::anyhow!(
"chunk_prefill_attention: q/k/v cols ({},{},{}) != num_heads*head_dim {dim}",
q_chunk.cols,
k_prefix.cols,
v_prefix.cols
)));
}
if v_prefix.rows != c1 {
return Err(FocrError::Other(anyhow::anyhow!(
"chunk_prefill_attention: k/v prefix rows disagree ({}, {})",
c1,
v_prefix.rows
)));
}
if c0 > c1 || c0 + cs != c1 {
return Err(FocrError::Other(anyhow::anyhow!(
"chunk_prefill_attention: chunk [{c0}, {}) of {cs} rows does not fit prefix {c1}",
c0 + cs
)));
}
let mut q_padded = Mat::zeros(c1, dim);
q_padded.data[c0 * dim..c1 * dim].copy_from_slice(&q_chunk.data);
let ctx_full = prefill_attention(&q_padded, k_prefix, v_prefix, num_heads, head_dim)?;
Ok(Mat::from_vec(
cs,
dim,
ctx_full.data[c0 * dim..c1 * dim].to_vec(),
))
}
pub fn prefill_with_cache(
wc: &DecoderWeightCache,
inputs_embeds: &Mat,
) -> FocrResult<(Mat, Vec<RingCache>)> {
prefill_with_cache_chunked(wc, inputs_embeds, prefill_chunk_size())
}
pub fn prefill_with_cache_chunked(
wc: &DecoderWeightCache,
inputs_embeds: &Mat,
chunk: Option<usize>,
) -> FocrResult<(Mat, Vec<RingCache>)> {
checked_mat_len("decoder::prefill_with_cache inputs_embeds", inputs_embeds)?;
let hidden = config::HIDDEN_SIZE;
let num_heads = config::NUM_ATTENTION_HEADS;
let head_dim = config::HEAD_DIM;
let qkv_dim = checked_shape_mul(
"decoder::prefill_with_cache",
num_heads,
head_dim,
"num_heads*head_dim",
)?;
let eps = config::RMS_NORM_EPS;
if inputs_embeds.cols != hidden {
return Err(FocrError::Other(anyhow::anyhow!(
"decoder::prefill_with_cache: inputs_embeds cols {} != hidden {hidden}",
inputs_embeds.cols
)));
}
let seq = inputs_embeds.rows;
let mut caches: Vec<RingCache> = (0..config::NUM_HIDDEN_LAYERS)
.map(|_| RingCache::new(seq.max(1)))
.collect();
if let Some(chunk) = chunk {
let mut out = Mat::zeros(seq, hidden);
let mut k_full: Vec<Mat> = (0..config::NUM_HIDDEN_LAYERS)
.map(|_| Mat::zeros(seq, qkv_dim))
.collect();
let mut v_full: Vec<Mat> = (0..config::NUM_HIDDEN_LAYERS)
.map(|_| Mat::zeros(seq, qkv_dim))
.collect();
for (c0, c1) in prefill_chunk_bounds(seq, chunk) {
let positions: Vec<usize> = (c0..c1).collect();
let rope = RopeTable::build(&positions, head_dim, config::ROPE_THETA);
let mut x = Mat::from_vec(
c1 - c0,
hidden,
inputs_embeds.data[c0 * hidden..c1 * hidden].to_vec(),
);
for layer in 0..config::NUM_HIDDEN_LAYERS {
let cl = &wc.layers[layer];
let lw = cached_layer_weights(cl);
let normed = nn::rms_norm(&x, Some(lw.input_ln), eps)?;
if calib::enabled() {
calib::record_rows(&quant_calib::attn_in_key(layer), &normed.data, normed.cols);
}
let (q, k, v) = qkv_with_rope(&normed, &lw, &rope, hidden, qkv_dim)?;
k_full[layer].data[c0 * qkv_dim..c1 * qkv_dim].copy_from_slice(&k.data);
v_full[layer].data[c0 * qkv_dim..c1 * qkv_dim].copy_from_slice(&v.data);
let kpre = Mat::from_vec(c1, qkv_dim, k_full[layer].data[..c1 * qkv_dim].to_vec());
let vpre = Mat::from_vec(c1, qkv_dim, v_full[layer].data[..c1 * qkv_dim].to_vec());
let context = chunk_prefill_attention(&q, &kpre, &vpre, num_heads, head_dim, c0)?;
if calib::enabled() {
calib::record_rows(&quant_calib::o_in_key(layer), &context.data, context.cols);
}
let attn_out = attn_output_proj(&context, lw.o_proj, hidden, qkv_dim)?;
let h = add_residual(&x, &attn_out)?;
let normed2 = nn::rms_norm(&h, Some(lw.post_attn_ln), eps)?;
let mlp_out = cached_mlp(&cl.mlp, &normed2, calib::enabled().then_some(layer))?;
x = add_residual(&h, &mlp_out)?;
}
out.data[c0 * hidden..c1 * hidden].copy_from_slice(&x.data);
}
for layer in 0..config::NUM_HIDDEN_LAYERS {
let (kh, vh) = token_major_to_head_major(
&k_full[layer],
&v_full[layer],
seq,
num_heads,
head_dim,
)?;
caches[layer].record_prefill(&kh, &vh, seq)?;
}
return Ok((out, caches));
}
let positions: Vec<usize> = (0..seq).collect();
let rope = RopeTable::build(&positions, head_dim, config::ROPE_THETA);
let mut x = inputs_embeds.clone();
for layer in 0..config::NUM_HIDDEN_LAYERS {
let cl = &wc.layers[layer];
let lw = cached_layer_weights(cl);
let normed = nn::rms_norm(&x, Some(lw.input_ln), eps)?;
if calib::enabled() {
calib::record_rows(&quant_calib::attn_in_key(layer), &normed.data, normed.cols);
}
let (q, k, v) = qkv_with_rope(&normed, &lw, &rope, hidden, qkv_dim)?;
let (kh, vh) = token_major_to_head_major(&k, &v, seq, num_heads, head_dim)?;
caches[layer].record_prefill(&kh, &vh, seq)?;
let context = prefill_attention(&q, &k, &v, num_heads, head_dim)?;
if calib::enabled() {
calib::record_rows(&quant_calib::o_in_key(layer), &context.data, context.cols);
}
let attn_out = attn_output_proj(&context, lw.o_proj, hidden, qkv_dim)?;
let h = add_residual(&x, &attn_out)?;
let normed2 = nn::rms_norm(&h, Some(lw.post_attn_ln), eps)?;
let mlp_out = cached_mlp(&cl.mlp, &normed2, calib::enabled().then_some(layer))?;
x = add_residual(&h, &mlp_out)?;
}
Ok((x, caches))
}
pub fn decode_step_with_cache(
wc: &DecoderWeightCache,
caches: &mut [RingCache],
token_embed: &Mat,
position: usize,
) -> FocrResult<Mat> {
let hidden = config::HIDDEN_SIZE;
let qkv_dim = checked_shape_mul(
"decoder::decode_step_with_cache",
config::NUM_ATTENTION_HEADS,
config::HEAD_DIM,
"num_heads*head_dim",
)?;
let eps = config::RMS_NORM_EPS;
if token_embed.rows != 1 || token_embed.cols != hidden {
return Err(FocrError::Other(anyhow::anyhow!(
"decoder::decode_step_with_cache: token_embed shape [{}, {}] != [1, {hidden}]",
token_embed.rows,
token_embed.cols
)));
}
if caches.len() != config::NUM_HIDDEN_LAYERS {
return Err(FocrError::Other(anyhow::anyhow!(
"decoder::decode_step_with_cache: {} caches != {} layers",
caches.len(),
config::NUM_HIDDEN_LAYERS
)));
}
let rope = RopeTable::build(&[position], config::HEAD_DIM, config::ROPE_THETA);
let profiling = prof::enabled();
let mut x = token_embed.clone();
for layer in 0..config::NUM_HIDDEN_LAYERS {
let cl = &wc.layers[layer];
let t_attn = profiling.then(Instant::now);
let normed = nn::rms_norm(&x, Some(&cl.input_ln), eps)?;
let nrow = normed.row(0);
if calib::enabled() {
calib::record_row(&quant_calib::attn_in_key(layer), nrow);
}
let mut q = Mat::from_vec(1, qkv_dim, gemv(nrow, &cl.q_proj, qkv_dim, hidden));
let mut k = Mat::from_vec(1, qkv_dim, gemv(nrow, &cl.k_proj, qkv_dim, hidden));
let v = gemv(nrow, &cl.v_proj, qkv_dim, hidden);
apply_rope(&mut q, &rope)?;
apply_rope(&mut k, &rope)?;
caches[layer].write_decode_step(&k.data, &v)?;
let context = rswa::decode_attention(&caches[layer], &q.data)?;
if calib::enabled() {
calib::record_row(&quant_calib::o_in_key(layer), &context.data);
}
let attn_out = Mat::from_vec(1, hidden, gemv(&context.data, &cl.o_proj, hidden, qkv_dim));
let h = add_residual(&x, &attn_out)?;
if let Some(t) = t_attn {
prof::add(&prof::ATTN_NS, t.elapsed().as_nanos() as u64);
}
let normed2 = nn::rms_norm(&h, Some(&cl.post_attn_ln), eps)?;
let mlp_out = Mat::from_vec(
1,
hidden,
decode_mlp(&cl.mlp, &normed2, calib::enabled().then_some(layer))?,
);
x = add_residual(&h, &mlp_out)?;
}
Ok(x)
}
const I8_GEMV_BLOCK: usize = 64;
#[inline]
fn quantize_row_i8(x: &[f32]) -> (Vec<i8>, f32) {
let amax = x.iter().fold(0.0f32, |m, &v| m.max(v.abs()));
let a_scale = if amax > 0.0 { amax / 127.0 } else { 1.0 };
let xq: Vec<i8> = x
.iter()
.map(|&v| (v / a_scale).round_ties_even().clamp(-127.0, 127.0) as i8)
.collect();
(xq, a_scale)
}
fn gemv_i8(x: &[f32], qw: &QInt8) -> Vec<f32> {
debug_assert_eq!(x.len(), qw.k);
let (xq, a_scale) = quantize_row_i8(x);
gemv_i8_prequant(&xq, a_scale, qw)
}
fn igemm_i8_block(qw: &QInt8, xq: &[i8], m: usize, base: usize, cnt: usize, acc: &mut [i32]) {
let k = qw.k;
match qw.layout {
WeightLayout::RowMajor => {
acc.fill(0);
simd::igemm_s8s8(xq, &qw.w[base * k..(base + cnt) * k], m, k, cnt, acc);
}
WeightLayout::SmmlaPanels => {
debug_assert_eq!(base % 2, 0, "channel blocks must be pair-aligned");
let kb = k.div_ceil(8);
let start = (base / 2) * kb * 16;
let len = cnt.div_ceil(2) * kb * 16;
simd::igemm_s8s8_packed_b(xq, &qw.w[start..start + len], m, k, cnt, acc);
}
}
}
fn gemv_i8_prequant(xq: &[i8], a_scale: f32, qw: &QInt8) -> Vec<f32> {
let k = qw.k;
let n = qw.n;
debug_assert_eq!(xq.len(), k);
debug_assert_eq!(qw.w.len(), qw.expected_w_len());
debug_assert_eq!(qw.scales.len(), n);
let mut y = vec![0.0f32; n];
y.par_chunks_mut(I8_GEMV_BLOCK)
.enumerate()
.for_each(|(blk, ys)| {
let base = blk * I8_GEMV_BLOCK;
let cnt = ys.len();
let mut acc = vec![0i32; cnt];
igemm_i8_block(qw, xq, 1, base, cnt, &mut acc);
for (j, slot) in ys.iter_mut().enumerate() {
*slot = acc[j] as f32 * a_scale * qw.scales[base + j];
}
});
y
}
pub(crate) fn gemm_i8_bias_prequant_batched(
rows: &[(&[i8], f32)],
qw: &QInt8,
bias: Option<&[f32]>,
) -> Vec<Vec<f32>> {
let b = rows.len();
let k = qw.k;
let n = qw.n;
debug_assert_eq!(qw.w.len(), qw.expected_w_len());
debug_assert_eq!(qw.scales.len(), n);
if b == 0 {
return Vec::new();
}
let mut xq = vec![0i8; b * k];
for (r, (row, _)) in rows.iter().enumerate() {
debug_assert_eq!(row.len(), k);
xq[r * k..(r + 1) * k].copy_from_slice(row);
}
let mut ycm = vec![0.0f32; n * b];
ycm.par_chunks_mut(I8_GEMV_BLOCK * b)
.enumerate()
.for_each(|(blk, ys)| {
let base = blk * I8_GEMV_BLOCK;
let cnt = ys.len() / b;
let mut acc = vec![0i32; b * cnt];
igemm_i8_block(qw, &xq, b, base, cnt, &mut acc);
for j in 0..cnt {
let scale_o = qw.scales[base + j];
let bias_o = bias.map_or(0.0, |bb| bb[base + j]);
for (r, &(_, a_scale)) in rows.iter().enumerate() {
ys[j * b + r] = acc[r * cnt + j] as f32 * a_scale * scale_o + bias_o;
}
}
});
let mut out: Vec<Vec<f32>> = (0..b).map(|_| vec![0.0f32; n]).collect();
for o in 0..n {
let col = o * b;
for (r, row) in out.iter_mut().enumerate() {
row[o] = ycm[col + r];
}
}
out
}
#[inline]
pub(crate) fn quantize_row_i8_te(x: &[f32]) -> (Vec<i8>, f32) {
quantize_row_i8(x)
}
pub(crate) fn gemv_i8_bias_prequant(
xq: &[i8],
a_scale: f32,
qw: &QInt8,
bias: Option<&[f32]>,
) -> Vec<f32> {
let mut y = gemv_i8_prequant(xq, a_scale, qw);
if let Some(b) = bias {
debug_assert_eq!(b.len(), y.len());
for (v, &bb) in y.iter_mut().zip(b) {
*v += bb;
}
}
y
}
const FUSE_NORM_QUANT_ENV: &str = "FOCR_FUSE_NORM_QUANT";
fn fuse_norm_quant_enabled() -> bool {
static FLAG: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*FLAG.get_or_init(|| std::env::var_os(FUSE_NORM_QUANT_ENV).is_some())
}
fn rms_norm_quant_i8(x: &Mat, weight: Option<&[f32]>, eps: f32) -> FocrResult<(Vec<i8>, f32)> {
let normed = nn::rms_norm(x, weight, eps)?;
Ok(quantize_row_i8(normed.row(0)))
}
fn expert_gemv_i8(
x: &[f32],
gate: &QLinear,
up: &QLinear,
down: &QLinear,
calib: Option<(usize, quant_calib::MlpUnit)>,
) -> Vec<f32> {
if let Some((layer, unit)) = calib.filter(|_| calib::enabled()) {
calib::record_row(&quant_calib::ffn_in_key(layer, unit), x);
}
let g = gemv_q(x, gate);
let u = gemv_q(x, up);
let inter = gate.n();
let mut act = vec![0.0f32; inter];
for i in 0..inter {
act[i] = silu(g[i]) * u[i];
}
if let Some((layer, unit)) = calib.filter(|_| calib::enabled()) {
calib::record_row(&quant_calib::down_in_key(layer, unit), &act);
}
gemv_q(&act, down)
}
fn gemv_i8_serial(xq: &[i8], a_scale: f32, qw: &QInt8) -> Vec<f32> {
let k = qw.k;
let n = qw.n;
debug_assert_eq!(xq.len(), k);
let mut acc = vec![0i32; n];
igemm_i8_block(qw, xq, 1, 0, n, &mut acc);
let mut y = vec![0.0f32; n];
for (o, slot) in y.iter_mut().enumerate() {
*slot = acc[o] as f32 * a_scale * qw.scales[o];
}
y
}
fn expert_gemv_i8_serial(
x: &[f32],
gate: &QLinear,
up: &QLinear,
down: &QLinear,
calib: Option<(usize, quant_calib::MlpUnit)>,
) -> Vec<f32> {
if let Some((layer, unit)) = calib.filter(|_| calib::enabled()) {
calib::record_row(&quant_calib::ffn_in_key(layer, unit), x);
}
let (xq, a_scale) = quantize_row_i8(x);
let g = gemv_q_serial(&xq, a_scale, gate);
let u = gemv_q_serial(&xq, a_scale, up);
let inter = gate.n();
let mut act = vec![0.0f32; inter];
for i in 0..inter {
act[i] = silu(g[i]) * u[i];
}
if let Some((layer, unit)) = calib.filter(|_| calib::enabled()) {
calib::record_row(&quant_calib::down_in_key(layer, unit), &act);
}
let (aq, a_scale2) = quantize_row_i8(&act);
gemv_q_serial(&aq, a_scale2, down)
}
fn quant_oc(w: &[f32], out: usize, in_: usize, name: &str) -> FocrResult<QInt8> {
let expected = out.checked_mul(in_).ok_or_else(|| {
FocrError::FormatMismatch(format!(
"tensor {name:?}: output/input shape [{out}, {in_}] overflows usize"
))
})?;
if w.len() != expected {
return Err(FocrError::FormatMismatch(format!(
"tensor {name:?}: {} elements != expected output/input shape [{out}, {in_}] ({expected})",
w.len()
)));
}
Ok(nn::quantize_int8(w, out, in_))
}
pub(crate) fn quant_oc_loaded(weights: &Weights, name: &str, out: usize) -> FocrResult<QInt8> {
if matches!(
weights.record(name).map(|rec| rec.dtype),
Some(DType::QInt8PerChan)
) {
let q = weights.qint8(name)?;
if q.n != out {
return Err(FocrError::FormatMismatch(format!(
"QInt8 tensor {name:?} has {} output rows; expected {out}",
q.n
)));
}
return Ok(q);
}
let record = weights.record(name).ok_or_else(|| {
FocrError::FormatMismatch(format!("tensor {name:?} not found in weights directory"))
})?;
let [rows, in_] = record.shape.as_slice() else {
return Err(FocrError::FormatMismatch(format!(
"tensor {name:?} has rank {}; expected 2 ([out, in])",
record.shape.len()
)));
};
if *rows != out {
return Err(FocrError::FormatMismatch(format!(
"tensor {name:?} has {rows} output rows; expected {out}"
)));
}
let mat = weights.mat(name)?;
quant_oc(&mat.data, out, *in_, name)
}
#[derive(Debug, Clone, PartialEq)]
pub(crate) enum QLinear {
I8(QInt8),
I4(QInt4),
}
impl QLinear {
pub(crate) fn n(&self) -> usize {
match self {
QLinear::I8(q) => q.n,
QLinear::I4(q) => q.n,
}
}
#[cfg_attr(not(test), allow(dead_code))]
pub(crate) fn k(&self) -> usize {
match self {
QLinear::I8(q) => q.k,
QLinear::I4(q) => q.k,
}
}
}
pub(crate) fn quant_ffn_loaded(weights: &Weights, name: &str, out: usize) -> FocrResult<QLinear> {
if matches!(
weights.record(name).map(|rec| rec.dtype),
Some(DType::QInt4PerGroup)
) {
let q = weights.qint4(name)?;
if q.n != out {
return Err(FocrError::FormatMismatch(format!(
"QInt4 tensor {name:?} has {} output rows; expected {out}",
q.n
)));
}
return Ok(QLinear::I4(q));
}
Ok(QLinear::I8(quant_oc_loaded(weights, name, out)?))
}
fn gemv_i4(x: &[f32], q: &QInt4) -> Vec<f32> {
debug_assert_eq!(x.len(), q.k);
let (xq, a_scale) = quantize_row_i8(x);
let kbytes = q.k / 2;
let groups = q.k / q.group_size;
let mut y = vec![0.0f32; q.n];
y.par_chunks_mut(I8_GEMV_BLOCK)
.enumerate()
.for_each(|(blk, ys)| {
let base = blk * I8_GEMV_BLOCK;
let cnt = ys.len();
let mut scale_scratch = Vec::new();
let scales =
q.scales
.slice_f32(base * groups, (base + cnt) * groups, &mut scale_scratch);
simd::int4::igemm_s4s8_packed(
&xq,
&q.packed[base * kbytes..(base + cnt) * kbytes],
scales,
q.group_size,
1,
q.k,
cnt,
ys,
);
for slot in ys.iter_mut() {
*slot *= a_scale;
}
});
y
}
fn gemv_i4_serial(xq: &[i8], a_scale: f32, q: &QInt4) -> Vec<f32> {
debug_assert_eq!(xq.len(), q.k);
let mut y = vec![0.0f32; q.n];
let mut scale_scratch = Vec::new();
let scales = q.scales.slice_f32(0, q.scales.len(), &mut scale_scratch);
simd::int4::igemm_s4s8_packed(xq, &q.packed, scales, q.group_size, 1, q.k, q.n, &mut y);
for slot in y.iter_mut() {
*slot *= a_scale;
}
y
}
fn gemv_q(x: &[f32], w: &QLinear) -> Vec<f32> {
match w {
QLinear::I8(q) => gemv_i8(x, q),
QLinear::I4(q) => gemv_i4(x, q),
}
}
fn gemv_q_serial(xq: &[i8], a_scale: f32, w: &QLinear) -> Vec<f32> {
match w {
QLinear::I8(q) => gemv_i8_serial(xq, a_scale, q),
QLinear::I4(q) => gemv_i4_serial(xq, a_scale, q),
}
}
fn linear_int4_dynamic(x: &Mat, q: &QInt4) -> FocrResult<Mat> {
checked_mat_len("linear_int4_dynamic x", x)?;
if x.cols != q.k {
return Err(FocrError::Other(anyhow::anyhow!(
"linear_int4_dynamic: x.cols {} != w.k {}",
x.cols,
q.k
)));
}
let (m, k, n) = (x.rows, x.cols, q.n);
let mut xq = vec![0i8; m * k];
let mut a_scales = vec![0.0f32; m];
for r in 0..m {
let (row_q, scale) = quantize_row_i8(x.row(r));
xq[r * k..(r + 1) * k].copy_from_slice(&row_q);
a_scales[r] = scale;
}
let mut out = vec![0.0f32; m * n];
let mut scale_scratch = Vec::new();
let scales = q.scales.slice_f32(0, q.scales.len(), &mut scale_scratch);
simd::int4::igemm_s4s8_packed(&xq, &q.packed, scales, q.group_size, m, k, n, &mut out);
for (r, &scale) in a_scales.iter().enumerate() {
for value in &mut out[r * n..(r + 1) * n] {
*value *= scale;
}
}
Ok(Mat::from_vec(m, n, out))
}
fn linear_q_dynamic(x: &Mat, w: &QLinear) -> FocrResult<Mat> {
match w {
QLinear::I8(q) => nn::linear_int8_dynamic(x, q, None),
QLinear::I4(q) => linear_int4_dynamic(x, q),
}
}
const LMHEAD_SHARD_ENV: &str = "FOCR_LMHEAD_SHARD";
const LMHEAD_SHARD_TILES_ENV: &str = "FOCR_LMHEAD_SHARD_TILES";
const DEFAULT_LMHEAD_SHARD_TILES: usize = 16;
#[must_use]
pub fn lmhead_shard_enabled() -> bool {
static FLAG: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*FLAG.get_or_init(|| std::env::var_os(LMHEAD_SHARD_ENV).is_some())
}
#[must_use]
pub fn lmhead_shard_tiles() -> usize {
static TILES: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
*TILES.get_or_init(|| {
std::env::var(LMHEAD_SHARD_TILES_ENV)
.ok()
.and_then(|v| v.trim().parse::<usize>().ok())
.filter(|&n| n > 0)
.unwrap_or(DEFAULT_LMHEAD_SHARD_TILES)
})
}
fn vocab_tile_ranges(n: usize, tiles: usize) -> Vec<(usize, usize)> {
let tiles = tiles.clamp(1, n.max(1));
let base = n / tiles;
let rem = n % tiles;
let mut ranges = Vec::with_capacity(tiles);
let mut start = 0usize;
for t in 0..tiles {
let len = base + usize::from(t < rem);
let end = start + len;
ranges.push((start, end));
start = end;
}
debug_assert_eq!(start, n, "vocab_tile_ranges must cover [0, n)");
ranges
}
fn gemv_sharded(x: &[f32], w: &[f32], n: usize, k: usize, tiles: usize) -> Vec<f32> {
debug_assert_eq!(x.len(), k);
debug_assert_eq!(w.len(), n * k);
let mut y = vec![0.0f32; n];
for (start, end) in vocab_tile_ranges(n, tiles) {
y[start..end]
.par_chunks_mut(64)
.enumerate()
.for_each(|(blk, ys)| {
let base = start + blk * 64;
for (j, slot) in ys.iter_mut().enumerate() {
let o = base + j;
*slot = dot_f32(x, &w[o * k..o * k + k]);
}
});
}
y
}
fn gemv_i8_sharded(x: &[f32], qw: &QInt8, tiles: usize) -> Vec<f32> {
let k = qw.k;
let n = qw.n;
debug_assert_eq!(x.len(), k);
debug_assert_eq!(qw.w.len(), qw.expected_w_len());
debug_assert_eq!(qw.scales.len(), n);
let (xq, a_scale) = quantize_row_i8(x);
let mut y = vec![0.0f32; n];
for (start, end) in vocab_tile_ranges(n, tiles) {
y[start..end]
.par_chunks_mut(I8_GEMV_BLOCK)
.enumerate()
.for_each(|(blk, ys)| {
let base = start + blk * I8_GEMV_BLOCK;
let cnt = ys.len();
let mut acc = vec![0i32; cnt];
igemm_i8_block(qw, &xq, 1, base, cnt, &mut acc);
for (j, slot) in ys.iter_mut().enumerate() {
*slot = acc[j] as f32 * a_scale * qw.scales[base + j];
}
});
}
y
}
const QKV_FUSED_ENV: &str = "FOCR_QKV_FUSED";
fn qkv_fused_enabled() -> bool {
static FLAG: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*FLAG.get_or_init(|| {
!matches!(
std::env::var(QKV_FUSED_ENV)
.ok()
.map(|v| v.trim().to_ascii_lowercase())
.as_deref(),
Some("0" | "off" | "false" | "no")
)
})
}
fn fuse_qkv(q: &QInt8, k: &QInt8, v: &QInt8) -> QInt8 {
debug_assert_eq!(q.k, k.k, "fuse_qkv: k contraction dim mismatch");
debug_assert_eq!(q.k, v.k, "fuse_qkv: k contraction dim mismatch");
debug_assert_eq!(q.n, k.n, "fuse_qkv: q/k output dim mismatch");
debug_assert_eq!(q.n, v.n, "fuse_qkv: q/v output dim mismatch");
let (n, kk) = (q.n, q.k);
let mut w = Vec::with_capacity(q.w.len() + k.w.len() + v.w.len());
w.extend_from_slice(&q.w);
w.extend_from_slice(&k.w);
w.extend_from_slice(&v.w);
let mut scales = Vec::with_capacity(n * 3);
scales.extend_from_slice(&q.scales);
scales.extend_from_slice(&k.scales);
scales.extend_from_slice(&v.scales);
if q.layout == WeightLayout::SmmlaPanels
&& k.layout == WeightLayout::SmmlaPanels
&& v.layout == WeightLayout::SmmlaPanels
&& q.n.is_multiple_of(2)
&& k.n.is_multiple_of(2)
{
return QInt8::new_smmla_panels(w, scales, 3 * n, kk);
}
debug_assert!(
q.layout == WeightLayout::RowMajor
&& k.layout == WeightLayout::RowMajor
&& v.layout == WeightLayout::RowMajor,
"fuse_qkv: mixed or odd-row packed layouts are unreachable from the loader"
);
QInt8::new(w, scales, 3 * n, kk)
}
enum CachedMlpI8 {
Dense {
gate: QLinear,
up: QLinear,
down: QLinear,
},
Moe {
gate: Vec<f32>,
experts: Vec<[QLinear; 3]>,
shared: [QLinear; 3],
},
}
struct CachedLayerI8 {
input_ln: Vec<f32>,
post_attn_ln: Vec<f32>,
q_proj: QInt8,
k_proj: QInt8,
v_proj: QInt8,
qkv: Option<QInt8>,
o_proj: QInt8,
mlp: CachedMlpI8,
}
pub struct DecoderWeightCacheI8 {
layers: Vec<CachedLayerI8>,
final_norm: Vec<f32>,
lm_head: QInt8,
}
impl DecoderWeightCacheI8 {
pub fn build(weights: &Weights) -> FocrResult<Self> {
let hidden = config::HIDDEN_SIZE;
let qkv_dim = config::NUM_ATTENTION_HEADS * config::HEAD_DIM;
let mut layers = Vec::with_capacity(config::NUM_HIDDEN_LAYERS);
for layer in 0..config::NUM_HIDDEN_LAYERS {
let prefix = format!("model.layers.{layer}");
let input_ln = weights.vec(&format!("{prefix}.input_layernorm.weight"))?;
let post_attn_ln = weights.vec(&format!("{prefix}.post_attention_layernorm.weight"))?;
let q_proj = quant_oc_loaded(
weights,
&format!("{prefix}.self_attn.q_proj.weight"),
qkv_dim,
)?;
let k_proj = quant_oc_loaded(
weights,
&format!("{prefix}.self_attn.k_proj.weight"),
qkv_dim,
)?;
let v_proj = quant_oc_loaded(
weights,
&format!("{prefix}.self_attn.v_proj.weight"),
qkv_dim,
)?;
let o_proj = quant_oc_loaded(
weights,
&format!("{prefix}.self_attn.o_proj.weight"),
hidden,
)?;
let qkv = qkv_fused_enabled().then(|| fuse_qkv(&q_proj, &k_proj, &v_proj));
let mlp = if layer < config::FIRST_K_DENSE_REPLACE {
let p = format!("{prefix}.mlp");
let inter = moe::config::DENSE_INTERMEDIATE_SIZE;
CachedMlpI8::Dense {
gate: quant_ffn_loaded(weights, &format!("{p}.gate_proj.weight"), inter)?,
up: quant_ffn_loaded(weights, &format!("{p}.up_proj.weight"), inter)?,
down: quant_ffn_loaded(weights, &format!("{p}.down_proj.weight"), hidden)?,
}
} else {
let p = format!("{prefix}.mlp");
let gate = weights.mat(&format!("{p}.gate.weight"))?.data;
let inter = moe::config::MOE_INTERMEDIATE_SIZE;
let mut experts = Vec::with_capacity(moe::config::N_ROUTED_EXPERTS);
for e in 0..moe::config::N_ROUTED_EXPERTS {
experts.push([
quant_ffn_loaded(
weights,
&format!("{p}.experts.{e}.gate_proj.weight"),
inter,
)?,
quant_ffn_loaded(
weights,
&format!("{p}.experts.{e}.up_proj.weight"),
inter,
)?,
quant_ffn_loaded(
weights,
&format!("{p}.experts.{e}.down_proj.weight"),
hidden,
)?,
]);
}
let si = moe::config::SHARED_INTERMEDIATE_SIZE;
let shared = [
quant_ffn_loaded(weights, &format!("{p}.shared_experts.gate_proj.weight"), si)?,
quant_ffn_loaded(weights, &format!("{p}.shared_experts.up_proj.weight"), si)?,
quant_ffn_loaded(
weights,
&format!("{p}.shared_experts.down_proj.weight"),
hidden,
)?,
];
CachedMlpI8::Moe {
gate,
experts,
shared,
}
};
layers.push(CachedLayerI8 {
input_ln,
post_attn_ln,
q_proj,
k_proj,
v_proj,
qkv,
o_proj,
mlp,
});
}
let final_norm = weights.vec("model.norm.weight")?;
let lm_head = quant_oc_loaded(weights, "lm_head.weight", config::VOCAB_SIZE)?;
Ok(Self {
layers,
final_norm,
lm_head,
})
}
}
pub(crate) fn expert_mlp_i8(x: &Mat, gate: &QInt8, up: &QInt8, down: &QInt8) -> FocrResult<Mat> {
let mut g = nn::linear_int8_dynamic(x, gate, None)?;
nn::silu(&mut g);
let u = nn::linear_int8_dynamic(x, up, None)?;
if g.data.len() != u.data.len() {
return Err(FocrError::Other(anyhow::anyhow!(
"decoder::expert_mlp_i8: gate/up shape mismatch ({} vs {})",
g.data.len(),
u.data.len()
)));
}
for (a, &b) in g.data.iter_mut().zip(u.data.iter()) {
*a *= b;
}
nn::linear_int8_dynamic(&g, down, None)
}
pub(crate) fn expert_mlp_q(
x: &Mat,
gate: &QLinear,
up: &QLinear,
down: &QLinear,
calib: Option<(usize, quant_calib::MlpUnit)>,
) -> FocrResult<Mat> {
if let Some((layer, unit)) = calib.filter(|_| calib::enabled()) {
calib::record_rows(&quant_calib::ffn_in_key(layer, unit), &x.data, x.cols);
}
let mut g = linear_q_dynamic(x, gate)?;
nn::silu(&mut g);
let u = linear_q_dynamic(x, up)?;
if g.data.len() != u.data.len() {
return Err(FocrError::Other(anyhow::anyhow!(
"decoder::expert_mlp_q: gate/up shape mismatch ({} vs {})",
g.data.len(),
u.data.len()
)));
}
for (a, &b) in g.data.iter_mut().zip(u.data.iter()) {
*a *= b;
}
if let Some((layer, unit)) = calib.filter(|_| calib::enabled()) {
calib::record_rows(&quant_calib::down_in_key(layer, unit), &g.data, g.cols);
}
linear_q_dynamic(&g, down)
}
#[allow(clippy::needless_range_loop)]
fn moe_block_i8(
hidden: &Mat,
gate: &[f32],
experts: &[[QLinear; 3]],
shared: &[QLinear; 3],
calib_layer: Option<usize>,
) -> FocrResult<Mat> {
let n_tok = hidden.rows;
let h = hidden.cols;
let routing = moe::route_default(hidden, gate)?;
let mut out = Mat::zeros(n_tok, h);
let contribution_len = n_tok
.checked_mul(moe::config::NUM_EXPERTS_PER_TOK)
.and_then(|rows| rows.checked_mul(h))
.ok_or_else(|| {
FocrError::Other(anyhow::anyhow!(
"decoder::moe_block_i8: n_tok*top_k*hidden overflow"
))
})?;
let mut contributions = vec![0.0f32; contribution_len];
let mut per_expert: Vec<Vec<(usize, usize, f32)>> =
vec![Vec::new(); moe::config::N_ROUTED_EXPERTS];
for t in 0..n_tok {
for j in 0..moe::config::NUM_EXPERTS_PER_TOK {
per_expert[routing.indices[t][j]].push((t, j, routing.weights[t][j]));
}
}
for (e, members) in per_expert.iter().enumerate() {
if members.is_empty() {
continue;
}
let m = members.len();
let mut sub = Mat::zeros(m, h);
for (r, &(t, _slot, _w)) in members.iter().enumerate() {
sub.row_mut(r).copy_from_slice(hidden.row(t));
}
let y = expert_mlp_q(
&sub,
&experts[e][0],
&experts[e][1],
&experts[e][2],
calib_layer.map(|l| (l, quant_calib::MlpUnit::Expert(e))),
)?;
for (r, &(t, slot, w)) in members.iter().enumerate() {
let yr = y.row(r);
let base = (t * moe::config::NUM_EXPERTS_PER_TOK + slot) * h;
let dst = &mut contributions[base..base + h];
for (value, &expert_value) in dst.iter_mut().zip(yr.iter()) {
*value = w * expert_value;
}
}
}
for t in 0..n_tok {
let token_base = t * moe::config::NUM_EXPERTS_PER_TOK * h;
let rows: [&[f32]; moe::config::NUM_EXPERTS_PER_TOK] = std::array::from_fn(|slot| {
let base = token_base + slot * h;
&contributions[base..base + h]
});
moe::combine_routed_rows(rows, &routing.indices[t], out.row_mut(t))?;
}
let shared_out = expert_mlp_q(
hidden,
&shared[0],
&shared[1],
&shared[2],
calib_layer.map(|l| (l, quant_calib::MlpUnit::Shared)),
)?;
for (o, &s) in out.data.iter_mut().zip(shared_out.data.iter()) {
*o += s;
}
Ok(out)
}
fn prefill_mlp_i8(mlp: &CachedMlpI8, normed: &Mat, calib_layer: Option<usize>) -> FocrResult<Mat> {
match mlp {
CachedMlpI8::Dense { gate, up, down } => expert_mlp_q(
normed,
gate,
up,
down,
calib_layer.map(|l| (l, quant_calib::MlpUnit::Dense)),
),
CachedMlpI8::Moe {
gate,
experts,
shared,
} => moe_block_i8(normed, gate, experts, shared, calib_layer),
}
}
fn decode_mlp_i8(
mlp: &CachedMlpI8,
normed: &Mat,
calib_layer: Option<usize>,
) -> FocrResult<Vec<f32>> {
let hidden = config::HIDDEN_SIZE;
let row = normed.row(0);
let profiling = prof::enabled();
match mlp {
CachedMlpI8::Dense { gate, up, down } => {
let t = profiling.then(Instant::now);
let y = expert_gemv_i8(
row,
gate,
up,
down,
calib_layer.map(|l| (l, quant_calib::MlpUnit::Dense)),
);
if let Some(t) = t {
prof::add(&prof::EXPERTS_NS, t.elapsed().as_nanos() as u64);
}
Ok(y)
}
CachedMlpI8::Moe {
gate,
experts,
shared,
} => {
let tr = profiling.then(Instant::now);
let routing = moe::route_default(normed, gate)?;
if let Some(t) = tr {
prof::add(&prof::ROUTE_NS, t.elapsed().as_nanos() as u64);
}
let te = profiling.then(Instant::now);
let idx = routing.indices[0];
let wts = routing.weights[0];
let (partials, s) = rayon::join(
|| {
(0..moe::config::NUM_EXPERTS_PER_TOK)
.into_par_iter()
.map(|j| {
let e = idx[j];
let w = wts[j];
let mut y = expert_gemv_i8_serial(
row,
&experts[e][0],
&experts[e][1],
&experts[e][2],
calib_layer.map(|l| (l, quant_calib::MlpUnit::Expert(e))),
);
for v in y.iter_mut() {
*v *= w;
}
y
})
.collect::<Vec<Vec<f32>>>()
},
|| {
expert_gemv_i8(
row,
&shared[0],
&shared[1],
&shared[2],
calib_layer.map(|l| (l, quant_calib::MlpUnit::Shared)),
)
},
);
let mut out = vec![0.0f32; hidden];
let rows: [&[f32]; moe::config::NUM_EXPERTS_PER_TOK] =
std::array::from_fn(|slot| partials[slot].as_slice());
moe::combine_routed_rows(rows, &idx, &mut out)?;
for c in 0..hidden {
out[c] += s[c];
}
if let Some(t) = te {
prof::add(&prof::EXPERTS_NS, t.elapsed().as_nanos() as u64);
}
Ok(out)
}
}
}
pub fn lm_head_cached_i8(wc: &DecoderWeightCacheI8, hidden: &Mat) -> FocrResult<Mat> {
if hidden.rows != 1 {
return Err(FocrError::Other(anyhow::anyhow!(
"decoder::lm_head_cached_i8: expected a single decode row, got {} rows",
hidden.rows
)));
}
let t = prof::enabled().then(Instant::now);
let normed = nn::rms_norm(hidden, Some(&wc.final_norm), config::RMS_NORM_EPS)?;
let row = normed.row(0);
let logits = if lmhead_shard_enabled() {
gemv_i8_sharded(row, &wc.lm_head, lmhead_shard_tiles())
} else {
gemv_i8(row, &wc.lm_head)
};
if let Some(t) = t {
prof::add(&prof::LMHEAD_NS, t.elapsed().as_nanos() as u64);
}
Ok(Mat::from_vec(1, config::VOCAB_SIZE, logits))
}
const FUSE_NGRAM_LMHEAD_ENV: &str = "FOCR_FUSE_NGRAM_LMHEAD";
pub fn fuse_ngram_lmhead_enabled() -> bool {
static FLAG: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*FLAG.get_or_init(|| std::env::var_os(FUSE_NGRAM_LMHEAD_ENV).is_some())
}
fn gemv_i8_ngram_masked(x: &[f32], qw: &QInt8, banned: &[u32]) -> Vec<f32> {
let n = qw.n;
if banned.is_empty() {
return gemv_i8(x, qw);
}
let mut ban_mask = vec![false; n];
for &b in banned {
let bi = b as usize;
if bi < n {
ban_mask[bi] = true;
}
}
let k = qw.k;
debug_assert_eq!(x.len(), k);
debug_assert_eq!(qw.w.len(), qw.expected_w_len());
debug_assert_eq!(qw.scales.len(), n);
let (xq, a_scale) = quantize_row_i8(x);
let mut y = vec![0.0f32; n];
y.par_chunks_mut(I8_GEMV_BLOCK)
.enumerate()
.for_each(|(blk, ys)| {
let base = blk * I8_GEMV_BLOCK;
let cnt = ys.len();
let mut acc = vec![0i32; cnt];
igemm_i8_block(qw, &xq, 1, base, cnt, &mut acc);
for (j, slot) in ys.iter_mut().enumerate() {
*slot = if ban_mask[base + j] {
f32::NEG_INFINITY
} else {
acc[j] as f32 * a_scale * qw.scales[base + j]
};
}
});
y
}
pub fn lm_head_cached_i8_ngram_masked(
wc: &DecoderWeightCacheI8,
hidden: &Mat,
banned: &[u32],
) -> FocrResult<Mat> {
if hidden.rows != 1 {
return Err(FocrError::Other(anyhow::anyhow!(
"decoder::lm_head_cached_i8_ngram_masked: expected a single decode row, got {} rows",
hidden.rows
)));
}
let t = prof::enabled().then(Instant::now);
let normed = nn::rms_norm(hidden, Some(&wc.final_norm), config::RMS_NORM_EPS)?;
let row = normed.row(0);
let logits = gemv_i8_ngram_masked(row, &wc.lm_head, banned);
if let Some(t) = t {
prof::add(&prof::LMHEAD_NS, t.elapsed().as_nanos() as u64);
}
Ok(Mat::from_vec(1, config::VOCAB_SIZE, logits))
}
pub fn prefill_with_cache_i8(
wc: &DecoderWeightCacheI8,
inputs_embeds: &Mat,
) -> FocrResult<(Mat, Vec<RingCache>)> {
prefill_with_cache_i8_chunked(wc, inputs_embeds, prefill_chunk_size())
}
pub fn prefill_with_cache_i8_chunked(
wc: &DecoderWeightCacheI8,
inputs_embeds: &Mat,
chunk: Option<usize>,
) -> FocrResult<(Mat, Vec<RingCache>)> {
checked_mat_len(
"decoder::prefill_with_cache_i8 inputs_embeds",
inputs_embeds,
)?;
let hidden = config::HIDDEN_SIZE;
let num_heads = config::NUM_ATTENTION_HEADS;
let head_dim = config::HEAD_DIM;
let qkv_dim = checked_shape_mul(
"decoder::prefill_with_cache_i8",
num_heads,
head_dim,
"num_heads*head_dim",
)?;
let eps = config::RMS_NORM_EPS;
if inputs_embeds.cols != hidden {
return Err(FocrError::Other(anyhow::anyhow!(
"decoder::prefill_with_cache_i8: inputs_embeds cols {} != hidden {hidden}",
inputs_embeds.cols
)));
}
let seq = inputs_embeds.rows;
let mut caches: Vec<RingCache> = (0..config::NUM_HIDDEN_LAYERS)
.map(|_| RingCache::new(seq.max(1)))
.collect();
if let Some(chunk) = chunk {
let mut out = Mat::zeros(seq, hidden);
let mut k_full: Vec<Mat> = (0..config::NUM_HIDDEN_LAYERS)
.map(|_| Mat::zeros(seq, qkv_dim))
.collect();
let mut v_full: Vec<Mat> = (0..config::NUM_HIDDEN_LAYERS)
.map(|_| Mat::zeros(seq, qkv_dim))
.collect();
for (c0, c1) in prefill_chunk_bounds(seq, chunk) {
let positions: Vec<usize> = (c0..c1).collect();
let rope = RopeTable::build(&positions, head_dim, config::ROPE_THETA);
let mut x = Mat::from_vec(
c1 - c0,
hidden,
inputs_embeds.data[c0 * hidden..c1 * hidden].to_vec(),
);
for layer in 0..config::NUM_HIDDEN_LAYERS {
let cl = &wc.layers[layer];
let normed = nn::rms_norm(&x, Some(&cl.input_ln), eps)?;
let mut q = nn::linear_int8_dynamic(&normed, &cl.q_proj, None)?;
let mut k = nn::linear_int8_dynamic(&normed, &cl.k_proj, None)?;
let v = nn::linear_int8_dynamic(&normed, &cl.v_proj, None)?;
apply_rope(&mut q, &rope)?;
apply_rope(&mut k, &rope)?;
k_full[layer].data[c0 * qkv_dim..c1 * qkv_dim].copy_from_slice(&k.data);
v_full[layer].data[c0 * qkv_dim..c1 * qkv_dim].copy_from_slice(&v.data);
let kpre = Mat::from_vec(c1, qkv_dim, k_full[layer].data[..c1 * qkv_dim].to_vec());
let vpre = Mat::from_vec(c1, qkv_dim, v_full[layer].data[..c1 * qkv_dim].to_vec());
let context = chunk_prefill_attention(&q, &kpre, &vpre, num_heads, head_dim, c0)?;
let attn_out = nn::linear_int8_dynamic(&context, &cl.o_proj, None)?;
let h = add_residual(&x, &attn_out)?;
let normed2 = nn::rms_norm(&h, Some(&cl.post_attn_ln), eps)?;
let mlp_out = prefill_mlp_i8(&cl.mlp, &normed2, calib::enabled().then_some(layer))?;
x = add_residual(&h, &mlp_out)?;
}
out.data[c0 * hidden..c1 * hidden].copy_from_slice(&x.data);
}
for layer in 0..config::NUM_HIDDEN_LAYERS {
let (kh, vh) = token_major_to_head_major(
&k_full[layer],
&v_full[layer],
seq,
num_heads,
head_dim,
)?;
caches[layer].record_prefill(&kh, &vh, seq)?;
}
return Ok((out, caches));
}
let positions: Vec<usize> = (0..seq).collect();
let rope = RopeTable::build(&positions, head_dim, config::ROPE_THETA);
let mut x = inputs_embeds.clone();
for layer in 0..config::NUM_HIDDEN_LAYERS {
let cl = &wc.layers[layer];
let normed = nn::rms_norm(&x, Some(&cl.input_ln), eps)?;
let mut q = nn::linear_int8_dynamic(&normed, &cl.q_proj, None)?;
let mut k = nn::linear_int8_dynamic(&normed, &cl.k_proj, None)?;
let v = nn::linear_int8_dynamic(&normed, &cl.v_proj, None)?;
apply_rope(&mut q, &rope)?;
apply_rope(&mut k, &rope)?;
let (kh, vh) = token_major_to_head_major(&k, &v, seq, num_heads, head_dim)?;
caches[layer].record_prefill(&kh, &vh, seq)?;
let context = prefill_attention(&q, &k, &v, num_heads, head_dim)?;
let attn_out = nn::linear_int8_dynamic(&context, &cl.o_proj, None)?;
let h = add_residual(&x, &attn_out)?;
let normed2 = nn::rms_norm(&h, Some(&cl.post_attn_ln), eps)?;
let mlp_out = prefill_mlp_i8(&cl.mlp, &normed2, calib::enabled().then_some(layer))?;
x = add_residual(&h, &mlp_out)?;
}
Ok((x, caches))
}
pub fn decode_step_with_cache_i8(
wc: &DecoderWeightCacheI8,
caches: &mut [RingCache],
token_embed: &Mat,
position: usize,
) -> FocrResult<Mat> {
let hidden = config::HIDDEN_SIZE;
let qkv_dim = checked_shape_mul(
"decoder::decode_step_with_cache_i8",
config::NUM_ATTENTION_HEADS,
config::HEAD_DIM,
"num_heads*head_dim",
)?;
let eps = config::RMS_NORM_EPS;
if token_embed.rows != 1 || token_embed.cols != hidden {
return Err(FocrError::Other(anyhow::anyhow!(
"decoder::decode_step_with_cache_i8: token_embed shape [{}, {}] != [1, {hidden}]",
token_embed.rows,
token_embed.cols
)));
}
if caches.len() != config::NUM_HIDDEN_LAYERS {
return Err(FocrError::Other(anyhow::anyhow!(
"decoder::decode_step_with_cache_i8: {} caches != {} layers",
caches.len(),
config::NUM_HIDDEN_LAYERS
)));
}
let rope = RopeTable::build(&[position], config::HEAD_DIM, config::ROPE_THETA);
let profiling = prof::enabled();
let mut x = token_embed.clone();
for layer in 0..config::NUM_HIDDEN_LAYERS {
let cl = &wc.layers[layer];
let t_attn = profiling.then(Instant::now);
let (mut q, mut k, v) = if fuse_norm_quant_enabled() {
let (xq, a_scale) = rms_norm_quant_i8(&x, Some(&cl.input_ln), eps)?;
if let Some(qkv) = &cl.qkv {
let out = gemv_i8_prequant(&xq, a_scale, qkv);
let q = Mat::from_vec(1, qkv_dim, out[0..qkv_dim].to_vec());
let k = Mat::from_vec(1, qkv_dim, out[qkv_dim..2 * qkv_dim].to_vec());
let v = out[2 * qkv_dim..3 * qkv_dim].to_vec();
(q, k, v)
} else {
let q = Mat::from_vec(1, qkv_dim, gemv_i8_prequant(&xq, a_scale, &cl.q_proj));
let k = Mat::from_vec(1, qkv_dim, gemv_i8_prequant(&xq, a_scale, &cl.k_proj));
let v = gemv_i8_prequant(&xq, a_scale, &cl.v_proj);
(q, k, v)
}
} else {
let normed = nn::rms_norm(&x, Some(&cl.input_ln), eps)?;
let nrow = normed.row(0);
if let Some(qkv) = &cl.qkv {
let out = gemv_i8(nrow, qkv);
let q = Mat::from_vec(1, qkv_dim, out[0..qkv_dim].to_vec());
let k = Mat::from_vec(1, qkv_dim, out[qkv_dim..2 * qkv_dim].to_vec());
let v = out[2 * qkv_dim..3 * qkv_dim].to_vec();
(q, k, v)
} else {
let q = Mat::from_vec(1, qkv_dim, gemv_i8(nrow, &cl.q_proj));
let k = Mat::from_vec(1, qkv_dim, gemv_i8(nrow, &cl.k_proj));
let v = gemv_i8(nrow, &cl.v_proj);
(q, k, v)
}
};
apply_rope(&mut q, &rope)?;
apply_rope(&mut k, &rope)?;
caches[layer].write_decode_step(&k.data, &v)?;
let context = rswa::decode_attention(&caches[layer], &q.data)?;
let attn_out = Mat::from_vec(1, hidden, gemv_i8(&context.data, &cl.o_proj));
let h = add_residual(&x, &attn_out)?;
if let Some(t) = t_attn {
prof::add(&prof::ATTN_NS, t.elapsed().as_nanos() as u64);
}
let normed2 = nn::rms_norm(&h, Some(&cl.post_attn_ln), eps)?;
let mlp_out = Mat::from_vec(
1,
hidden,
decode_mlp_i8(&cl.mlp, &normed2, calib::enabled().then_some(layer))?,
);
x = add_residual(&h, &mlp_out)?;
}
Ok(x)
}
const BATCH_SPINE_ENV: &str = "FOCR_BATCH_SPINE";
const BATCH_SIZE_ENV: &str = "FOCR_BATCH_SIZE";
const DEFAULT_BATCH_SIZE: usize = 8;
#[must_use]
pub fn batch_spine_enabled() -> bool {
static FLAG: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*FLAG.get_or_init(|| match std::env::var(BATCH_SPINE_ENV) {
Ok(v) => !matches!(
v.trim().to_ascii_lowercase().as_str(),
"" | "0" | "off" | "false" | "no"
),
Err(_) => false,
})
}
#[must_use]
pub fn batch_size_cap() -> usize {
static CAP: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
*CAP.get_or_init(|| {
std::env::var(BATCH_SIZE_ENV)
.ok()
.and_then(|v| v.trim().parse::<usize>().ok())
.filter(|&n| n > 0)
.unwrap_or(DEFAULT_BATCH_SIZE)
})
}
fn gemv_i8_batched(rows: &[&[f32]], qw: &QInt8) -> Vec<Vec<f32>> {
let b = rows.len();
let k = qw.k;
let n = qw.n;
debug_assert_eq!(qw.w.len(), qw.expected_w_len());
debug_assert_eq!(qw.scales.len(), n);
if b == 0 {
return Vec::new();
}
let mut xq = vec![0i8; b * k];
let mut a_scales = vec![0.0f32; b];
for (r, &row) in rows.iter().enumerate() {
debug_assert_eq!(row.len(), k);
let (q, a_scale) = quantize_row_i8(row);
xq[r * k..(r + 1) * k].copy_from_slice(&q);
a_scales[r] = a_scale;
}
let mut ycm = vec![0.0f32; n * b];
ycm.par_chunks_mut(I8_GEMV_BLOCK * b)
.enumerate()
.for_each(|(blk, ys)| {
let base = blk * I8_GEMV_BLOCK;
let cnt = ys.len() / b;
let mut acc = vec![0i32; b * cnt];
igemm_i8_block(qw, &xq, b, base, cnt, &mut acc);
for j in 0..cnt {
let scale_o = qw.scales[base + j];
for r in 0..b {
ys[j * b + r] = acc[r * cnt + j] as f32 * a_scales[r] * scale_o;
}
}
});
let mut out: Vec<Vec<f32>> = (0..b).map(|_| vec![0.0f32; n]).collect();
for o in 0..n {
let col = o * b;
for r in 0..b {
out[r][o] = ycm[col + r];
}
}
out
}
pub fn batched_decode_step_i8(
wc: &DecoderWeightCacheI8,
caches: &mut BatchedRingCache,
token_embeds: &[Mat],
positions: &[usize],
) -> FocrResult<Vec<Mat>> {
let stream_ids: Vec<usize> = (0..caches.num_streams()).collect();
batched_decode_step_i8_streams(wc, caches, &stream_ids, token_embeds, positions)
}
pub fn batched_decode_step_i8_streams(
wc: &DecoderWeightCacheI8,
caches: &mut BatchedRingCache,
stream_ids: &[usize],
token_embeds: &[Mat],
positions: &[usize],
) -> FocrResult<Vec<Mat>> {
let hidden = config::HIDDEN_SIZE;
let qkv_dim = checked_shape_mul(
"decoder::batched_decode_step_i8",
config::NUM_ATTENTION_HEADS,
config::HEAD_DIM,
"num_heads*head_dim",
)?;
let eps = config::RMS_NORM_EPS;
let b = stream_ids.len();
if token_embeds.len() != b || positions.len() != b {
return Err(FocrError::Other(anyhow::anyhow!(
"decoder::batched_decode_step_i8: token_embeds {} / positions {} != active streams {b}",
token_embeds.len(),
positions.len()
)));
}
if caches.num_layers() != config::NUM_HIDDEN_LAYERS {
return Err(FocrError::Other(anyhow::anyhow!(
"decoder::batched_decode_step_i8: {} cache layers != {} layers",
caches.num_layers(),
config::NUM_HIDDEN_LAYERS
)));
}
let n_streams = caches.num_streams();
for (k, &sid) in stream_ids.iter().enumerate() {
if sid >= n_streams {
return Err(FocrError::Other(anyhow::anyhow!(
"decoder::batched_decode_step_i8: active slot {k} stream id {sid} >= cache streams {n_streams}"
)));
}
}
for (s, te) in token_embeds.iter().enumerate() {
if te.rows != 1 || te.cols != hidden {
return Err(FocrError::Other(anyhow::anyhow!(
"decoder::batched_decode_step_i8: token_embeds[{s}] shape [{}, {}] != [1, {hidden}]",
te.rows,
te.cols
)));
}
}
let profiling = prof::enabled();
let mut x: Vec<Mat> = token_embeds.to_vec();
for layer in 0..config::NUM_HIDDEN_LAYERS {
let cl = &wc.layers[layer];
let t_attn = profiling.then(Instant::now);
let mut normed: Vec<Mat> = Vec::with_capacity(b);
for s in 0..b {
normed.push(nn::rms_norm(&x[s], Some(&cl.input_ln), eps)?);
}
let nrows: Vec<&[f32]> = normed.iter().map(|m| m.row(0)).collect();
let (mut q_mats, mut k_mats, v_rows): (Vec<Mat>, Vec<Mat>, Vec<Vec<f32>>) =
if let Some(qkv) = &cl.qkv {
let outs = gemv_i8_batched(&nrows, qkv);
let mut qm = Vec::with_capacity(b);
let mut km = Vec::with_capacity(b);
let mut vr = Vec::with_capacity(b);
for row in outs {
qm.push(Mat::from_vec(1, qkv_dim, row[0..qkv_dim].to_vec()));
km.push(Mat::from_vec(
1,
qkv_dim,
row[qkv_dim..2 * qkv_dim].to_vec(),
));
vr.push(row[2 * qkv_dim..3 * qkv_dim].to_vec());
}
(qm, km, vr)
} else {
let q_out = gemv_i8_batched(&nrows, &cl.q_proj);
let k_out = gemv_i8_batched(&nrows, &cl.k_proj);
let v_out = gemv_i8_batched(&nrows, &cl.v_proj);
let qm: Vec<Mat> = q_out
.into_iter()
.map(|r| Mat::from_vec(1, qkv_dim, r))
.collect();
let km: Vec<Mat> = k_out
.into_iter()
.map(|r| Mat::from_vec(1, qkv_dim, r))
.collect();
(qm, km, v_out)
};
let mut context_rows: Vec<Mat> = Vec::with_capacity(b);
for s in 0..b {
let rope = RopeTable::build(&[positions[s]], config::HEAD_DIM, config::ROPE_THETA);
apply_rope(&mut q_mats[s], &rope)?;
apply_rope(&mut k_mats[s], &rope)?;
caches.write_decode_step(stream_ids[s], layer, &k_mats[s].data, &v_rows[s])?;
context_rows.push(rswa::decode_attention(
caches.layer(stream_ids[s], layer),
&q_mats[s].data,
)?);
}
let ctx_refs: Vec<&[f32]> = context_rows.iter().map(|m| m.row(0)).collect();
let attn_rows = gemv_i8_batched(&ctx_refs, &cl.o_proj);
if let Some(t) = t_attn {
prof::add(&prof::ATTN_NS, t.elapsed().as_nanos() as u64);
}
for (s, attn) in attn_rows.into_iter().enumerate() {
let attn_out = Mat::from_vec(1, hidden, attn);
let h = add_residual(&x[s], &attn_out)?;
let normed2 = nn::rms_norm(&h, Some(&cl.post_attn_ln), eps)?;
let mlp_out = Mat::from_vec(
1,
hidden,
decode_mlp_i8(&cl.mlp, &normed2, calib::enabled().then_some(layer))?,
);
x[s] = add_residual(&h, &mlp_out)?;
}
}
Ok(x)
}
pub fn batched_lm_head_i8(wc: &DecoderWeightCacheI8, hiddens: &[Mat]) -> FocrResult<Vec<Mat>> {
let t = prof::enabled().then(Instant::now);
let mut normed: Vec<Mat> = Vec::with_capacity(hiddens.len());
for (s, h) in hiddens.iter().enumerate() {
if h.rows != 1 {
return Err(FocrError::Other(anyhow::anyhow!(
"decoder::batched_lm_head_i8: hiddens[{s}] has {} rows, expected 1",
h.rows
)));
}
normed.push(nn::rms_norm(h, Some(&wc.final_norm), config::RMS_NORM_EPS)?);
}
let rows: Vec<&[f32]> = normed.iter().map(|m| m.row(0)).collect();
let outs = gemv_i8_batched(&rows, &wc.lm_head);
let logits = outs
.into_iter()
.map(|y| Mat::from_vec(1, config::VOCAB_SIZE, y))
.collect();
if let Some(t) = t {
prof::add(&prof::LMHEAD_NS, t.elapsed().as_nanos() as u64);
}
Ok(logits)
}
const SPEC_VERIFY_ENV: &str = "FOCR_SPEC_VERIFY";
#[must_use]
pub fn spec_verify_enabled() -> bool {
static FLAG: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*FLAG.get_or_init(|| std::env::var_os(SPEC_VERIFY_ENV).is_some())
}
#[allow(dead_code)] fn check_verify_inputs(
context: &str,
caches: &[RingCache],
token_embeds: &[Mat],
) -> FocrResult<()> {
let hidden = config::HIDDEN_SIZE;
if caches.len() != config::NUM_HIDDEN_LAYERS {
return Err(FocrError::Other(anyhow::anyhow!(
"decoder::{context}: {} caches != {} layers",
caches.len(),
config::NUM_HIDDEN_LAYERS
)));
}
for (i, te) in token_embeds.iter().enumerate() {
if te.rows != 1 || te.cols != hidden {
return Err(FocrError::Other(anyhow::anyhow!(
"decoder::{context}: token_embeds[{i}] shape [{}, {}] != [1, {hidden}]",
te.rows,
te.cols
)));
}
}
Ok(())
}
#[allow(dead_code)] pub(crate) fn verify_forward(
wc: &DecoderWeightCache,
caches: &[RingCache],
token_embeds: &[Mat],
base_position: usize,
) -> FocrResult<Vec<Mat>> {
let hidden = config::HIDDEN_SIZE;
let qkv_dim = checked_shape_mul(
"decoder::verify_forward",
config::NUM_ATTENTION_HEADS,
config::HEAD_DIM,
"num_heads*head_dim",
)?;
let eps = config::RMS_NORM_EPS;
check_verify_inputs("verify_forward", caches, token_embeds)?;
let k = token_embeds.len();
let mut x: Vec<Mat> = token_embeds.to_vec();
for layer in 0..config::NUM_HIDDEN_LAYERS {
let cl = &wc.layers[layer];
let mut q_mats: Vec<Mat> = Vec::with_capacity(k);
let mut k_mats: Vec<Mat> = Vec::with_capacity(k);
let mut v_rows: Vec<Vec<f32>> = Vec::with_capacity(k);
for i in 0..k {
let normed = nn::rms_norm(&x[i], Some(&cl.input_ln), eps)?;
let nrow = normed.row(0);
let mut q = Mat::from_vec(1, qkv_dim, gemv(nrow, &cl.q_proj, qkv_dim, hidden));
let mut kk = Mat::from_vec(1, qkv_dim, gemv(nrow, &cl.k_proj, qkv_dim, hidden));
let v = gemv(nrow, &cl.v_proj, qkv_dim, hidden);
let rope = RopeTable::build(&[base_position + i], config::HEAD_DIM, config::ROPE_THETA);
apply_rope(&mut q, &rope)?;
apply_rope(&mut kk, &rope)?;
q_mats.push(q);
k_mats.push(kk);
v_rows.push(v);
}
let q_refs: Vec<&[f32]> = q_mats.iter().map(|m| m.data.as_slice()).collect();
let k_refs: Vec<&[f32]> = k_mats.iter().map(|m| m.data.as_slice()).collect();
let v_refs: Vec<&[f32]> = v_rows.iter().map(|r| r.as_slice()).collect();
let contexts = rswa::verify_attention(&caches[layer], &q_refs, &k_refs, &v_refs)?;
for i in 0..k {
let o = gemv(&contexts[i].data, &cl.o_proj, hidden, qkv_dim);
let attn_out = Mat::from_vec(1, hidden, o);
let h = add_residual(&x[i], &attn_out)?;
let normed2 = nn::rms_norm(&h, Some(&cl.post_attn_ln), eps)?;
let mlp_out = Mat::from_vec(
1,
hidden,
decode_mlp(&cl.mlp, &normed2, calib::enabled().then_some(layer))?,
);
x[i] = add_residual(&h, &mlp_out)?;
}
}
x.iter().map(|h| lm_head_cached(wc, h)).collect()
}
#[allow(dead_code)] pub(crate) fn verify_forward_i8(
wc: &DecoderWeightCacheI8,
caches: &[RingCache],
token_embeds: &[Mat],
base_position: usize,
) -> FocrResult<Vec<Mat>> {
let hidden = config::HIDDEN_SIZE;
let qkv_dim = checked_shape_mul(
"decoder::verify_forward_i8",
config::NUM_ATTENTION_HEADS,
config::HEAD_DIM,
"num_heads*head_dim",
)?;
let eps = config::RMS_NORM_EPS;
check_verify_inputs("verify_forward_i8", caches, token_embeds)?;
let k = token_embeds.len();
let mut x: Vec<Mat> = token_embeds.to_vec();
for layer in 0..config::NUM_HIDDEN_LAYERS {
let cl = &wc.layers[layer];
let mut normed: Vec<Mat> = Vec::with_capacity(k);
for i in 0..k {
normed.push(nn::rms_norm(&x[i], Some(&cl.input_ln), eps)?);
}
let nrows: Vec<&[f32]> = normed.iter().map(|m| m.row(0)).collect();
let (mut q_mats, mut k_mats, v_rows): (Vec<Mat>, Vec<Mat>, Vec<Vec<f32>>) =
if let Some(qkv) = &cl.qkv {
let outs = gemv_i8_batched(&nrows, qkv);
let mut qm = Vec::with_capacity(k);
let mut km = Vec::with_capacity(k);
let mut vr = Vec::with_capacity(k);
for row in outs {
qm.push(Mat::from_vec(1, qkv_dim, row[0..qkv_dim].to_vec()));
km.push(Mat::from_vec(
1,
qkv_dim,
row[qkv_dim..2 * qkv_dim].to_vec(),
));
vr.push(row[2 * qkv_dim..3 * qkv_dim].to_vec());
}
(qm, km, vr)
} else {
let q_out = gemv_i8_batched(&nrows, &cl.q_proj);
let k_out = gemv_i8_batched(&nrows, &cl.k_proj);
let v_out = gemv_i8_batched(&nrows, &cl.v_proj);
let qm: Vec<Mat> = q_out
.into_iter()
.map(|r| Mat::from_vec(1, qkv_dim, r))
.collect();
let km: Vec<Mat> = k_out
.into_iter()
.map(|r| Mat::from_vec(1, qkv_dim, r))
.collect();
(qm, km, v_out)
};
for i in 0..k {
let rope = RopeTable::build(&[base_position + i], config::HEAD_DIM, config::ROPE_THETA);
apply_rope(&mut q_mats[i], &rope)?;
apply_rope(&mut k_mats[i], &rope)?;
}
let q_refs: Vec<&[f32]> = q_mats.iter().map(|m| m.data.as_slice()).collect();
let k_refs: Vec<&[f32]> = k_mats.iter().map(|m| m.data.as_slice()).collect();
let v_refs: Vec<&[f32]> = v_rows.iter().map(|r| r.as_slice()).collect();
let contexts = rswa::verify_attention(&caches[layer], &q_refs, &k_refs, &v_refs)?;
let ctx_refs: Vec<&[f32]> = contexts.iter().map(|m| m.row(0)).collect();
let attn_rows = gemv_i8_batched(&ctx_refs, &cl.o_proj);
for (i, attn) in attn_rows.into_iter().enumerate() {
let attn_out = Mat::from_vec(1, hidden, attn);
let h = add_residual(&x[i], &attn_out)?;
let normed2 = nn::rms_norm(&h, Some(&cl.post_attn_ln), eps)?;
let mlp_out = Mat::from_vec(
1,
hidden,
decode_mlp_i8(&cl.mlp, &normed2, calib::enabled().then_some(layer))?,
);
x[i] = add_residual(&h, &mlp_out)?;
}
}
x.iter().map(|h| lm_head_cached_i8(wc, h)).collect()
}
#[cfg(test)]
mod tests {
use super::*;
const MOE_POLICY_CHILD_ENV: &str = "FOCR_MOE_POLICY_TEST_CASE";
fn moe_policy_child_case() -> Option<(&'static str, bool)> {
match std::env::var(MOE_POLICY_CHILD_ENV).ok().as_deref() {
None => None,
Some("unset") => Some(("unset", false)),
Some("zero") => Some(("zero", false)),
Some("one") => Some(("one", true)),
Some(other) => {
eprintln!("unknown MoE policy subprocess case {other:?}");
None
}
}
}
fn int8_moe_policy_fixture(
rollback: bool,
) -> (Mat, Vec<f32>, Vec<[QLinear; 3]>, [QLinear; 3], moe::Routing) {
const TARGETS: [f32; moe::config::NUM_EXPERTS_PER_TOK] =
[16_777_216.0, 1.0, -16_777_216.0, 1.0, 1.0, 1.0];
let fixture: serde_json::Value =
serde_json::from_str(include_str!("../../tests/fixtures/moe_torch_2_10_cpu.json"))
.expect("valid pinned torch MoE oracle fixture");
let unique = &fixture["cases"][0];
let scores = unique["scores_f32_bits"]
.as_array()
.expect("fixture score bits")
.iter()
.map(|bits| f32::from_bits(bits.as_u64().expect("fixture f32 bits") as u32))
.collect::<Vec<_>>();
let hidden_size = moe::config::HIDDEN_SIZE;
let mut hidden_data = vec![0.0f32; hidden_size];
hidden_data[0] = 1.0;
let hidden = Mat::from_vec(1, hidden_size, hidden_data);
let mut gate = vec![0.0f32; moe::config::N_ROUTED_EXPERTS * hidden_size];
for (expert, score) in scores.into_iter().enumerate() {
gate[expert * hidden_size] = score;
}
let routing = moe::route_default(&hidden, &gate).expect("public route_default succeeds");
let expected_field = if rollback {
"torch_sorted_indices"
} else {
"torch_unsorted_indices"
};
let expected = unique[expected_field]
.as_array()
.expect("fixture route indices")
.iter()
.map(|item| item.as_u64().expect("fixture expert index") as usize)
.collect::<Vec<_>>();
assert_eq!(routing.indices[0].as_slice(), expected.as_slice());
let mut unit_weights = vec![0i8; hidden_size];
unit_weights[0] = 127;
let unit = QInt8::new(unit_weights, vec![1.0 / 127.0], 1, hidden_size);
let zero = QInt8::new(vec![0; hidden_size], vec![1.0], 1, hidden_size);
let silu_one = 1.0f32 / (1.0 + (-1.0f32).exp());
let mut expert_outputs = vec![0.0f32; moe::config::N_ROUTED_EXPERTS];
for (slot, &expert) in routing.indices[0].iter().enumerate() {
expert_outputs[expert] = TARGETS[slot] / routing.weights[0][slot];
}
let experts = expert_outputs
.into_iter()
.map(|expert_output| {
let down = if expert_output == 0.0 {
QInt8::new(vec![0; hidden_size], vec![1.0; hidden_size], hidden_size, 1)
} else {
QInt8::new(
vec![127; hidden_size],
vec![expert_output / (127.0 * silu_one); hidden_size],
hidden_size,
1,
)
};
[
QLinear::I8(unit.clone()),
QLinear::I8(unit.clone()),
QLinear::I8(down),
]
})
.collect::<Vec<_>>();
let shared = [
QLinear::I8(zero.clone()),
QLinear::I8(zero),
QLinear::I8(QInt8::new(
vec![0; hidden_size],
vec![1.0; hidden_size],
hidden_size,
1,
)),
];
(hidden, gate, experts, shared, routing)
}
fn fold_int8_moe_contributions(
contributions: &[Vec<f32>; moe::config::NUM_EXPERTS_PER_TOK],
indices: &[usize; moe::config::NUM_EXPERTS_PER_TOK],
ascending_expert: bool,
) -> Vec<f32> {
let mut slots: [usize; moe::config::NUM_EXPERTS_PER_TOK] = std::array::from_fn(|slot| slot);
if ascending_expert {
slots.sort_unstable_by_key(|&slot| indices[slot]);
}
let mut out = vec![0.0f32; moe::config::HIDDEN_SIZE];
for slot in slots {
for (dst, &value) in out.iter_mut().zip(contributions[slot].iter()) {
*dst += value;
}
}
out
}
#[test]
fn moe_policy_subprocess_probe_int8() -> FocrResult<()> {
let Some((case, rollback)) = moe_policy_child_case() else {
return Ok(());
};
let (hidden, gate, experts, shared, routing) = int8_moe_policy_fixture(rollback);
let indices = routing.indices[0];
let weights = routing.weights[0];
let mut prefill_contributions: [Vec<f32>; moe::config::NUM_EXPERTS_PER_TOK] =
std::array::from_fn(|_| vec![0.0; moe::config::HIDDEN_SIZE]);
let mut decode_contributions: [Vec<f32>; moe::config::NUM_EXPERTS_PER_TOK] =
std::array::from_fn(|_| vec![0.0; moe::config::HIDDEN_SIZE]);
for slot in 0..moe::config::NUM_EXPERTS_PER_TOK {
let expert = indices[slot];
let prefill = expert_mlp_q(
&hidden,
&experts[expert][0],
&experts[expert][1],
&experts[expert][2],
None,
)?;
let decode = expert_gemv_i8_serial(
hidden.row(0),
&experts[expert][0],
&experts[expert][1],
&experts[expert][2],
None,
);
for channel in 0..moe::config::HIDDEN_SIZE {
prefill_contributions[slot][channel] = weights[slot] * prefill.data[channel];
decode_contributions[slot][channel] = weights[slot] * decode[channel];
}
}
let expected_prefill =
fold_int8_moe_contributions(&prefill_contributions, &indices, rollback);
let alternate_prefill =
fold_int8_moe_contributions(&prefill_contributions, &indices, !rollback);
let expected_decode =
fold_int8_moe_contributions(&decode_contributions, &indices, rollback);
let alternate_decode =
fold_int8_moe_contributions(&decode_contributions, &indices, !rollback);
assert_ne!(
expected_prefill[0].to_bits(),
alternate_prefill[0].to_bits()
);
assert_ne!(expected_decode[0].to_bits(), alternate_decode[0].to_bits());
let mut batched_data = hidden.data.clone();
batched_data.extend_from_slice(&hidden.data);
let batched_hidden = Mat::from_vec(2, moe::config::HIDDEN_SIZE, batched_data);
let prefill = moe_block_i8(&batched_hidden, &gate, &experts, &shared, None)?;
assert_eq!(prefill.row(0), expected_prefill.as_slice());
assert_eq!(prefill.row(1), expected_prefill.as_slice());
let mlp = CachedMlpI8::Moe {
gate,
experts,
shared,
};
let decode = decode_mlp_i8(&mlp, &hidden, None)?;
assert_eq!(decode, expected_decode);
eprintln!("FOCR_MOE_POLICY_PROBE_INT8={case}");
Ok(())
}
fn quant_oc_fixture_builder() -> crate::quant::focrq::FocrqBuilder {
let arch = crate::native_engine::model_arch::arch_by_id("got-ocr2")
.expect("got-ocr2 test architecture must be registered");
crate::quant::focrq::FocrqBuilder::new()
.with_model_id(arch.id())
.with_license_notice(arch.license_notice())
}
#[test]
fn quant_oc_loaded_rejects_prequantized_output_row_mismatch() {
let mut builder = quant_oc_fixture_builder();
builder
.add_quantized(
"w",
crate::quant::focrq::WriteDType::QInt8PerChan,
vec![3, 2],
vec![0; 6],
vec![0; 3 * std::mem::size_of::<f32>()],
0,
0,
)
.expect("valid synthetic QInt8 record");
let weights = Weights::from_bytes(builder.build()).expect("synthetic artifact loads");
let q = quant_oc_loaded(&weights, "w", 3).expect("matching output rows load");
assert_eq!((q.n, q.k), (3, 2), "input/K dimension must be preserved");
let err = quant_oc_loaded(&weights, "w", 2).expect_err("row mismatch must fail");
assert!(matches!(err, FocrError::FormatMismatch(_)));
assert!(format!("{err}").contains("3 output rows; expected 2"));
}
#[test]
fn quant_oc_loaded_rejects_high_precision_output_row_mismatch_and_wrong_rank() {
let mut builder = quant_oc_fixture_builder();
builder
.add_tensor(
"matrix",
crate::quant::focrq::WriteDType::Bf16,
vec![3, 2],
vec![0; 3 * 2 * std::mem::size_of::<u16>()],
)
.expect("valid synthetic BF16 matrix");
builder
.add_tensor(
"vector",
crate::quant::focrq::WriteDType::Bf16,
vec![6],
vec![0; 6 * std::mem::size_of::<u16>()],
)
.expect("valid synthetic BF16 vector");
let weights = Weights::from_bytes(builder.build()).expect("synthetic artifact loads");
let q = quant_oc_loaded(&weights, "matrix", 3).expect("matching matrix loads");
assert_eq!((q.n, q.k), (3, 2), "input/K dimension must be preserved");
let err = quant_oc_loaded(&weights, "matrix", 2).expect_err("row mismatch must fail");
assert!(matches!(err, FocrError::FormatMismatch(_)));
assert!(format!("{err}").contains("3 output rows; expected 2"));
let err = quant_oc_loaded(&weights, "vector", 1)
.expect_err("a rank-1 vector is not an [out, in] weight");
assert!(matches!(err, FocrError::FormatMismatch(_)));
assert!(format!("{err}").contains("rank 1; expected 2"));
}
fn synthetic_qint4(n: usize, k: usize, group_size: usize, salt: u64) -> QInt4 {
let packed: Vec<u8> = (0..n * k / 2)
.map(|i| ((i as u64).wrapping_mul(2654435761).wrapping_add(salt * 977) & 0xFF) as u8)
.collect();
let groups = k / group_size;
let scales: Vec<f32> = (0..n * groups)
.map(|i| {
let raw = ((i as u64).wrapping_mul(40503).wrapping_add(salt) % 4096) as f32;
(raw / 4096.0 - 0.5) * 1.7e-2 + 1.0e-3
})
.collect();
QInt4 {
packed: packed.into(),
scales: scales.into(),
n,
k,
group_size,
tier: 1,
}
}
fn reference_gemv_i4(xq: &[i8], a_scale: f32, q: &QInt4) -> Vec<f32> {
let b = simd::int4::unpack_to_i8(&q.packed, q.n, q.k);
let scales = q.scales.to_vec();
let groups = q.k / q.group_size;
let mut y = vec![0.0f32; q.n];
for o in 0..q.n {
let mut acc_f = 0.0f32;
for g in 0..groups {
let mut acc_i: i32 = 0;
for kk in g * q.group_size..(g + 1) * q.group_size {
acc_i += i32::from(xq[kk]) * i32::from(b[o * q.k + kk]);
}
acc_f += scales[o * groups + g] * acc_i as f32;
}
y[o] = acc_f * a_scale;
}
y
}
#[test]
fn gemv_i4_bitexact_vs_unpack_oracle_over_real_expert_dims() {
let cases: [(usize, usize, usize); 5] = [
(896, 1280, 32),
(1280, 896, 16),
(6848, 1280, 32),
(1280, 6848, 16),
(130, 96, 16), ];
for (case, &(n, k, group)) in cases.iter().enumerate() {
let q = synthetic_qint4(n, k, group, case as u64 + 1);
let x: Vec<f32> = (0..k)
.map(|i| (i as f32 * 0.37).sin() * 2.5 - 0.4)
.collect();
let (xq, a_scale) = quantize_row_i8(&x);
let want = reference_gemv_i4(&xq, a_scale, &q);
let parallel = gemv_i4(&x, &q);
let serial = gemv_i4_serial(&xq, a_scale, &q);
let bits = |s: &[f32]| s.iter().map(|f| f.to_bits()).collect::<Vec<u32>>();
assert_eq!(
bits(¶llel),
bits(&want),
"gemv_i4 != oracle for (n={n},k={k},g={group})"
);
assert_eq!(
bits(&serial),
bits(&want),
"gemv_i4_serial != oracle for (n={n},k={k},g={group})"
);
}
}
#[test]
fn linear_int4_dynamic_rows_bitexact_vs_serial_gemv() {
let (n, k, group) = (96usize, 64usize, 16usize);
let q = synthetic_qint4(n, k, group, 7);
let m = 3usize;
let x = Mat::from_vec(
m,
k,
(0..m * k)
.map(|i| (i as f32 * 0.11).cos() * 1.9 + 0.2)
.collect(),
);
let out = linear_int4_dynamic(&x, &q).expect("shapes agree");
assert_eq!(out.shape(), (m, n));
for r in 0..m {
let (xq, a_scale) = quantize_row_i8(x.row(r));
let want = gemv_i4_serial(&xq, a_scale, &q);
assert_eq!(
out.row(r).iter().map(|f| f.to_bits()).collect::<Vec<u32>>(),
want.iter().map(|f| f.to_bits()).collect::<Vec<u32>>(),
"prefill row {r} != m=1 serial GEMV"
);
}
let via_dispatch =
linear_q_dynamic(&x, &QLinear::I4(q.clone())).expect("dispatch shapes agree");
assert_eq!(via_dispatch, out);
let bad = Mat::zeros(1, k + group);
assert!(linear_int4_dynamic(&bad, &q).is_err());
}
#[test]
fn expert_forward_int4_matches_dequantized_f32_reference() {
use crate::quant::int4::{dequantize_int4, pack_int4_f32};
let (hidden, inter) = (1280usize, 896usize);
let wgen = |len: usize, salt: u64| -> Vec<f32> {
(0..len)
.map(|i| {
let raw = ((i as u64)
.wrapping_mul(6364136223846793005)
.wrapping_add(salt)
>> 33) as u32;
(raw as f32 / u32::MAX as f32 - 0.5) * 0.08
})
.collect()
};
let gate_f = wgen(inter * hidden, 1);
let up_f = wgen(inter * hidden, 2);
let down_f = wgen(hidden * inter, 3);
let gate_q = pack_int4_f32(&gate_f, inter, hidden, 16);
let up_q = pack_int4_f32(&up_f, inter, hidden, 32);
let down_q = pack_int4_f32(&down_f, hidden, inter, 16);
let to_qint4 = |q: &crate::quant::int4::QuantizedInt4| QInt4 {
packed: q.packed.clone().into(),
scales: q.scales.clone().into(),
n: q.n,
k: q.k,
group_size: q.group_size,
tier: 1,
};
let gate = QLinear::I4(to_qint4(&gate_q));
let up = QLinear::I4(to_qint4(&up_q));
let down = QLinear::I4(to_qint4(&down_q));
let x: Vec<f32> = (0..hidden)
.map(|i| (i as f32 * 0.019).sin() * 1.3 - 0.1)
.collect();
let decode = expert_gemv_i8_serial(&x, &gate, &up, &down, None);
let parallel = expert_gemv_i8(&x, &gate, &up, &down, None);
assert_eq!(
decode.iter().map(|f| f.to_bits()).collect::<Vec<u32>>(),
parallel.iter().map(|f| f.to_bits()).collect::<Vec<u32>>(),
"serial vs parallel int4 expert paths must be bit-identical"
);
let xm = Mat::from_vec(1, hidden, x.clone());
let prefill = expert_mlp_q(&xm, &gate, &up, &down, None).expect("prefill expert runs");
assert_eq!(
prefill
.row(0)
.iter()
.map(|f| f.to_bits())
.collect::<Vec<u32>>(),
decode.iter().map(|f| f.to_bits()).collect::<Vec<u32>>(),
"prefill m=1 vs decode int4 expert paths must be bit-identical"
);
let gate_d = dequantize_int4(&gate_q);
let up_d = dequantize_int4(&up_q);
let down_d = dequantize_int4(&down_q);
let dot = |w: &[f32], row: usize, x: &[f32]| -> f32 {
let k = x.len();
w[row * k..(row + 1) * k]
.iter()
.zip(x)
.map(|(&a, &b)| a * b)
.sum()
};
let mut act = vec![0.0f32; inter];
for o in 0..inter {
let g = dot(&gate_d, o, &x);
let u = dot(&up_d, o, &x);
act[o] = silu(g) * u;
}
let reference: Vec<f32> = (0..hidden).map(|o| dot(&down_d, o, &act)).collect();
let amax = reference.iter().fold(0.0f32, |m, &v| m.max(v.abs()));
assert!(amax > 0.0, "reference expert output must be non-degenerate");
for (o, (&got, &want)) in decode.iter().zip(reference.iter()).enumerate() {
assert!(
(got - want).abs() <= 0.02 * amax,
"int4 expert channel {o}: {got} vs f32 reference {want} (amax {amax})"
);
}
}
#[test]
fn quant_ffn_loaded_keeps_int4_packed_and_dispatches_int8_unchanged() {
let (n, k, group) = (4usize, 32usize, 16usize);
let q = synthetic_qint4(n, k, group, 11);
let mut builder = quant_oc_fixture_builder();
builder
.add_quantized(
"w4",
crate::quant::focrq::WriteDType::QInt4PerGroup,
vec![n, k],
q.packed.to_vec(),
q.scales
.to_vec()
.iter()
.flat_map(|s| s.to_le_bytes())
.collect(),
group,
1,
)
.expect("valid synthetic QInt4 record");
builder
.add_quantized(
"w8",
crate::quant::focrq::WriteDType::QInt8PerChan,
vec![3, 2],
vec![0; 6],
vec![0; 3 * std::mem::size_of::<f32>()],
0,
0,
)
.expect("valid synthetic QInt8 record");
let weights = Weights::from_bytes(builder.build()).expect("synthetic artifact loads");
let loaded = quant_ffn_loaded(&weights, "w4", n).expect("int4 tensor loads");
assert_eq!((loaded.n(), loaded.k()), (n, k));
let QLinear::I4(loaded_q) = &loaded else {
panic!("QInt4PerGroup record must take the packed int4 arm");
};
assert_eq!(
loaded_q.packed, q.packed,
"nibbles must stay packed verbatim"
);
assert_eq!(loaded_q.scales, q.scales);
assert_eq!(loaded_q.group_size, group);
let x: Vec<f32> = (0..k).map(|i| (i as f32) * 0.21 - 2.0).collect();
assert_eq!(gemv_q(&x, &loaded), gemv_i4(&x, &q));
let err = quant_ffn_loaded(&weights, "w4", n + 1).expect_err("row mismatch must fail");
assert!(matches!(err, FocrError::FormatMismatch(_)));
assert!(format!("{err}").contains("output rows; expected"));
let loaded8 = quant_ffn_loaded(&weights, "w8", 3).expect("int8 tensor loads");
assert!(
matches!(loaded8, QLinear::I8(_)),
"QInt8PerChan must keep the historical int8 arm"
);
}
#[test]
fn decode_mlp_over_artifact_int4_weights_is_bitexact() {
let (hidden, inter, group) = (64usize, 48usize, 16usize);
let gate = synthetic_qint4(inter, hidden, group, 21);
let up = synthetic_qint4(inter, hidden, 32, 22);
let down = synthetic_qint4(hidden, inter, group, 23);
let mut builder = quant_oc_fixture_builder();
for (name, q) in [("gate", &gate), ("up", &up), ("down", &down)] {
builder
.add_quantized(
name,
crate::quant::focrq::WriteDType::QInt4PerGroup,
vec![q.n, q.k],
q.packed.to_vec(),
q.scales
.to_vec()
.iter()
.flat_map(|s| s.to_le_bytes())
.collect(),
q.group_size,
1,
)
.expect("valid synthetic QInt4 record");
}
let weights = Weights::from_bytes(builder.build()).expect("synthetic artifact loads");
let mlp = CachedMlpI8::Dense {
gate: quant_ffn_loaded(&weights, "gate", inter).expect("gate loads"),
up: quant_ffn_loaded(&weights, "up", inter).expect("up loads"),
down: quant_ffn_loaded(&weights, "down", hidden).expect("down loads"),
};
let x = Mat::from_vec(
1,
hidden,
(0..hidden).map(|i| (i as f32 * 0.31).sin() * 1.1).collect(),
);
let decoded = decode_mlp_i8(&mlp, &x, None).expect("decode over artifact int4 runs");
let want = expert_gemv_i8(
x.row(0),
&QLinear::I4(gate),
&QLinear::I4(up),
&QLinear::I4(down),
None,
);
assert_eq!(
decoded.iter().map(|f| f.to_bits()).collect::<Vec<u32>>(),
want.iter().map(|f| f.to_bits()).collect::<Vec<u32>>(),
"artifact-loaded int4 dense MLP must match direct packed compute"
);
let prefill = prefill_mlp_i8(&mlp, &x, None).expect("prefill over artifact int4 runs");
assert_eq!(
prefill
.row(0)
.iter()
.map(|f| f.to_bits())
.collect::<Vec<u32>>(),
decoded.iter().map(|f| f.to_bits()).collect::<Vec<u32>>(),
);
}
#[test]
fn activation_quantization_uses_true_division_and_ties_to_even() {
let ties = [127.0, 0.5, 2.5, -0.5, -2.5];
let (q, scale) = quantize_row_i8(&ties);
assert_eq!(scale, 1.0);
assert_eq!(q, [127, 0, 2, 0, -2]);
let boundary = [f32::from_bits(0x3cdd_67c9), f32::from_bits(0xbcda_ca56)];
let (q, scale) = quantize_row_i8(&boundary);
assert_eq!(scale.to_bits(), 0x395f_2615);
assert_eq!(q, [127, -125]);
let reciprocal = (boundary[1] * (1.0 / scale)).round_ties_even() as i8;
assert_eq!(
reciprocal, -126,
"test vector must distinguish the two formulas"
);
}
#[test]
fn smmla_panel_layout_is_byte_identical_through_the_gemv_paths() {
let (n, k) = (130usize, 96usize);
let mut w = vec![0i8; n * k];
for (i, v) in w.iter_mut().enumerate() {
*v = (((i as i64 * 37 + 5) % 255) - 127) as i8;
}
let scales: Vec<f32> = (0..n).map(|o| 1.0e-3 + o as f32 * 3.0e-5).collect();
let row_major = QInt8::new(w.clone(), scales.clone(), n, k);
let (panels, _, _) = crate::simd::pack::smmla_pack_panels(&w, 0, n, k, k);
let packed = QInt8::new_smmla_panels(panels, scales, n, k);
let x: Vec<f32> = (0..k).map(|i| (i as f32) * 0.11 - 4.0).collect();
let a = gemv_i8(&x, &row_major);
let b = gemv_i8(&x, &packed);
assert_eq!(a, b, "gemv_i8 must be layout-invariant (bit-exact)");
let xq: Vec<i8> = (0..k)
.map(|i| (((i * 91) % 255) as i32 - 127) as i8)
.collect();
let a = gemv_i8_prequant(&xq, 0.017, &row_major);
let b = gemv_i8_prequant(&xq, 0.017, &packed);
assert_eq!(
a, b,
"gemv_i8_prequant must be layout-invariant (bit-exact)"
);
println!(r#"{{"check":"smmla_panel_gemv_parity","n":{n},"k":{k},"result":"pass"}}"#);
}
#[test]
fn fuse_qkv_concatenates_smmla_panels_losslessly() {
let (qkv_dim, hidden) = (64usize, 48usize);
let mk = |salt: i64| -> (Vec<i8>, Vec<f32>) {
let mut w = vec![0i8; qkv_dim * hidden];
for (i, v) in w.iter_mut().enumerate() {
*v = (((i as i64 * 31 + salt * 101) % 255) - 127) as i8;
}
let scales: Vec<f32> = (0..qkv_dim)
.map(|o| 1.0e-3 + (o as f32 + salt as f32 * 0.5) * 1.0e-4)
.collect();
(w, scales)
};
let (qw, qs) = mk(1);
let (kw, ks) = mk(2);
let (vw, vs) = mk(3);
let pack = |w: &[i8]| crate::simd::pack::smmla_pack_panels(w, 0, qkv_dim, hidden, hidden).0;
let fused_rm = fuse_qkv(
&QInt8::new(qw.clone(), qs.clone(), qkv_dim, hidden),
&QInt8::new(kw.clone(), ks.clone(), qkv_dim, hidden),
&QInt8::new(vw.clone(), vs.clone(), qkv_dim, hidden),
);
let fused_pk = fuse_qkv(
&QInt8::new_smmla_panels(pack(&qw), qs, qkv_dim, hidden),
&QInt8::new_smmla_panels(pack(&kw), ks, qkv_dim, hidden),
&QInt8::new_smmla_panels(pack(&vw), vs, qkv_dim, hidden),
);
assert_eq!(fused_pk.layout, WeightLayout::SmmlaPanels);
let repacked = crate::simd::pack::smmla_pack_panels(
&fused_rm.w,
0,
fused_rm.n,
fused_rm.k,
fused_rm.k,
)
.0;
assert_eq!(
&fused_pk.w[..],
&repacked[..],
"panel concat must equal pack of the fused matrix"
);
let x: Vec<f32> = (0..hidden).map(|i| (i as f32) * 0.07 - 1.5).collect();
assert_eq!(
gemv_i8(&x, &fused_rm),
gemv_i8(&x, &fused_pk),
"fused panel GEMV must be bit-exact"
);
println!(r#"{{"check":"smmla_panel_fuse_qkv","result":"pass"}}"#);
}
#[test]
fn fused_qkv_gemv_is_byte_identical_to_three_calls() {
let qkv_dim = 70usize; let hidden = 96usize;
let mk = |salt: i32| -> QInt8 {
let mut w = vec![0i8; qkv_dim * hidden];
for o in 0..qkv_dim {
for i in 0..hidden {
let raw = ((o as i32 * 31 + i as i32 * 7 + salt * 101) % 255) - 127;
w[o * hidden + i] = raw as i8;
}
}
let scales: Vec<f32> = (0..qkv_dim)
.map(|o| 1.0e-3 + (o as f32 + salt as f32 * 0.5) * 1.0e-4)
.collect();
QInt8::new(w, scales, qkv_dim, hidden)
};
let q_proj = mk(0);
let k_proj = mk(1);
let v_proj = mk(2);
let nrow: Vec<f32> = (0..hidden)
.map(|i| (i as f32 * 0.37).sin() * 2.5 - 0.4)
.collect();
let q_sep = gemv_i8(&nrow, &q_proj);
let k_sep = gemv_i8(&nrow, &k_proj);
let v_sep = gemv_i8(&nrow, &v_proj);
let fused = fuse_qkv(&q_proj, &k_proj, &v_proj);
assert_eq!(fused.n, 3 * qkv_dim);
assert_eq!(fused.k, hidden);
let y = gemv_i8(&nrow, &fused);
let bits = |s: &[f32]| s.iter().map(|f| f.to_bits()).collect::<Vec<u32>>();
assert_eq!(bits(&q_sep), bits(&y[0..qkv_dim]), "q mismatch");
assert_eq!(bits(&k_sep), bits(&y[qkv_dim..2 * qkv_dim]), "k mismatch");
assert_eq!(
bits(&v_sep),
bits(&y[2 * qkv_dim..3 * qkv_dim]),
"v mismatch"
);
}
#[test]
fn fused_norm_quant_is_byte_identical_to_norm_then_gemv() {
let hidden = 96usize;
let qkv_dim = 70usize; let mk = |salt: i32| -> QInt8 {
let mut w = vec![0i8; qkv_dim * hidden];
for o in 0..qkv_dim {
for i in 0..hidden {
let raw = ((o as i32 * 29 + i as i32 * 11 + salt * 97) % 255) - 127;
w[o * hidden + i] = raw as i8;
}
}
let scales: Vec<f32> = (0..qkv_dim)
.map(|o| 1.0e-3 + (o as f32 + salt as f32 * 0.5) * 1.0e-4)
.collect();
QInt8::new(w, scales, qkv_dim, hidden)
};
let q_proj = mk(0);
let k_proj = mk(1);
let v_proj = mk(2);
let ln_w: Vec<f32> = (0..hidden).map(|i| 0.5 + (i as f32) * 0.03).collect();
let x = Mat::from_vec(
1,
hidden,
(0..hidden)
.map(|i| (i as f32 * 0.41).cos() * 3.0 - 0.7)
.collect(),
);
let eps = config::RMS_NORM_EPS;
let bits = |s: &[f32]| s.iter().map(|f| f.to_bits()).collect::<Vec<u32>>();
let normed = nn::rms_norm(&x, Some(&ln_w), eps).unwrap();
let nrow = normed.row(0);
let q_def = gemv_i8(nrow, &q_proj);
let k_def = gemv_i8(nrow, &k_proj);
let v_def = gemv_i8(nrow, &v_proj);
let (xq, a_scale) = rms_norm_quant_i8(&x, Some(&ln_w), eps).unwrap();
let q_fused = gemv_i8_prequant(&xq, a_scale, &q_proj);
let k_fused = gemv_i8_prequant(&xq, a_scale, &k_proj);
let v_fused = gemv_i8_prequant(&xq, a_scale, &v_proj);
assert_eq!(bits(&q_def), bits(&q_fused), "q mismatch");
assert_eq!(bits(&k_def), bits(&k_fused), "k mismatch");
assert_eq!(bits(&v_def), bits(&v_fused), "v mismatch");
let (xq_ref, a_ref) = quantize_row_i8(nrow);
assert_eq!(xq, xq_ref, "int8 activation bytes mismatch");
assert_eq!(a_scale.to_bits(), a_ref.to_bits(), "a_scale bits mismatch");
}
#[test]
fn fused_ngram_lmhead_is_byte_identical_to_separate_mask() {
use super::super::sampler;
let vocab = 130usize;
let hidden = 48usize;
let mut w = vec![0i8; vocab * hidden];
for o in 0..vocab {
for i in 0..hidden {
let raw = ((o as i32 * 13 + i as i32 * 7 + 5) % 255) - 127;
w[o * hidden + i] = raw as i8;
}
}
let scales: Vec<f32> = (0..vocab).map(|o| 1.0e-3 + o as f32 * 1.0e-4).collect();
let qw = QInt8::new(w, scales, vocab, hidden);
let x: Vec<f32> = (0..hidden)
.map(|i| (i as f32 * 0.31).sin() * 2.0 - 0.3)
.collect();
let bits = |s: &[f32]| s.iter().map(|f| f.to_bits()).collect::<Vec<u32>>();
let ngram_size = 2usize;
let window = 16usize;
let seq: Vec<u32> = vec![3, 7, 3];
let logits = gemv_i8(&x, &qw);
let expected =
sampler::masked_sliding_window_logits_if_needed(&logits, &seq, ngram_size, window, &[])
.unwrap_or_else(|| logits.clone());
let banned =
sampler::collect_sliding_window_ngram_bans(&seq, ngram_size, window, &[], vocab);
let fused = gemv_i8_ngram_masked(&x, &qw, &banned);
assert!(
!banned.is_empty(),
"test sequence should ban at least one token"
);
assert_eq!(
fused[7].to_bits(),
f32::NEG_INFINITY.to_bits(),
"banned token 7 should be masked to -inf"
);
assert_eq!(bits(&fused), bits(&expected), "masked logits row mismatch");
let none = gemv_i8_ngram_masked(&x, &qw, &[]);
assert_eq!(bits(&none), bits(&logits), "no-ban head must equal gemv_i8");
}
#[test]
fn batched_gemv_i8_is_byte_identical_to_per_row() {
let mk = |n: usize, k: usize, salt: i32| -> QInt8 {
let mut w = vec![0i8; n * k];
for o in 0..n {
for i in 0..k {
let raw = ((o as i32 * 17 + i as i32 * 5 + salt * 101) % 255) - 127;
w[o * k + i] = raw as i8;
}
}
let scales: Vec<f32> = (0..n)
.map(|o| 1.0e-3 + (o as f32 + salt as f32 * 0.5) * 1.0e-4)
.collect();
QInt8::new(w, scales, n, k)
};
let mkrow = |k: usize, seed: f32| -> Vec<f32> {
(0..k)
.map(|i| ((i as f32 + seed) * 0.37).sin() * 2.5 - 0.4)
.collect()
};
let bits = |s: &[f32]| s.iter().map(|f| f.to_bits()).collect::<Vec<u32>>();
for &(n, k) in &[(96usize, 96usize), (3 * 70, 96), (130, 70)] {
let w = mk(n, k, 3);
for b in 1..=3usize {
let rows_owned: Vec<Vec<f32>> = (0..b).map(|s| mkrow(k, s as f32 * 1.7)).collect();
let rows: Vec<&[f32]> = rows_owned.iter().map(|r| r.as_slice()).collect();
let batched = gemv_i8_batched(&rows, &w);
assert_eq!(batched.len(), b);
for r in 0..b {
let single = gemv_i8(&rows_owned[r], &w);
assert_eq!(
bits(&single),
bits(&batched[r]),
"batched row {r} != per-row gemv_i8 (n={n} k={k} b={b})"
);
}
}
}
}
#[test]
fn batch_spine_defaults_off_with_positive_cap() {
if std::env::var_os("FOCR_BATCH_SPINE").is_none() {
assert!(!batch_spine_enabled());
}
assert!(batch_size_cap() >= 1);
}
#[test]
fn prefill_chunk_defaults_off() {
if std::env::var_os("FOCR_PREFILL_CHUNK").is_none() {
assert_eq!(prefill_chunk_size(), None);
}
}
#[test]
fn prefill_chunk_bounds_cover_contiguously_in_order() {
for &seq in &[0usize, 1, 2, 5, 7, 16, 17] {
for &chunk in &[1usize, 2, 3, 5, 7, 16, 64] {
let bounds = prefill_chunk_bounds(seq, chunk);
let mut cursor = 0usize;
for &(c0, c1) in &bounds {
assert_eq!(c0, cursor, "seq={seq} chunk={chunk}: gap/overlap");
assert!(c1 > c0, "seq={seq} chunk={chunk}: empty/reversed chunk");
assert!(c1 - c0 <= chunk, "seq={seq} chunk={chunk}: chunk too wide");
assert!(c1 <= seq, "seq={seq} chunk={chunk}: chunk past seq");
cursor = c1;
}
assert_eq!(cursor, seq, "seq={seq} chunk={chunk}: must cover [0, seq)");
}
}
}
#[test]
fn chunked_attention_is_byte_identical_to_monolithic() {
let (num_heads, head_dim, seq) = (3usize, 4usize, 11usize);
let dim = num_heads * head_dim;
let mk = |salt: f32| -> Mat {
let data: Vec<f32> = (0..seq * dim)
.map(|i| ((i as f32 + salt) * 0.37).sin() * 1.7 - 0.2)
.collect();
Mat::from_vec(seq, dim, data)
};
let (q, k, v) = (mk(0.0), mk(11.0), mk(23.0));
let monolithic = prefill_attention(&q, &k, &v, num_heads, head_dim).unwrap();
let bits = |s: &[f32]| s.iter().map(|f| f.to_bits()).collect::<Vec<u32>>();
for &chunk in &[1usize, 2, 3, seq] {
let mut assembled = Mat::zeros(seq, dim);
for (c0, c1) in prefill_chunk_bounds(seq, chunk) {
let q_chunk = Mat::from_vec(c1 - c0, dim, q.data[c0 * dim..c1 * dim].to_vec());
let kpre = Mat::from_vec(c1, dim, k.data[..c1 * dim].to_vec());
let vpre = Mat::from_vec(c1, dim, v.data[..c1 * dim].to_vec());
let ctx = chunk_prefill_attention(&q_chunk, &kpre, &vpre, num_heads, head_dim, c0)
.unwrap();
assembled.data[c0 * dim..c1 * dim].copy_from_slice(&ctx.data);
}
assert_eq!(
bits(&monolithic.data),
bits(&assembled.data),
"chunk={chunk}: chunked attention != monolithic prefill_attention"
);
}
}
#[test]
fn vocab_tile_ranges_cover_contiguously_in_order() {
for &n in &[0usize, 1, 7, 64, 65, 129_280] {
for &tiles in &[1usize, 2, 3, 7, 16, 1000, 200_000] {
let ranges = vocab_tile_ranges(n, tiles);
let mut cursor = 0usize;
for &(start, end) in &ranges {
assert_eq!(start, cursor, "n={n} tiles={tiles}: tile gap/overlap");
assert!(end >= start, "n={n} tiles={tiles}: reversed tile");
assert!(end <= n, "n={n} tiles={tiles}: tile end out of bounds");
cursor = end;
}
assert_eq!(cursor, n, "n={n} tiles={tiles}: ranges must cover [0, n)");
assert!(!ranges.is_empty(), "n={n} tiles={tiles}: at least one tile");
assert!(
ranges.len() <= n.max(1),
"n={n} tiles={tiles}: more tiles than channels"
);
}
}
}
#[test]
fn lmhead_shard_i8_is_byte_identical_to_monolithic() {
let mk = |n: usize, k: usize, salt: i32| -> QInt8 {
let mut w = vec![0i8; n * k];
for o in 0..n {
for i in 0..k {
let raw = ((o as i32 * 13 + i as i32 * 7 + salt * 101) % 255) - 127;
w[o * k + i] = raw as i8;
}
}
let scales: Vec<f32> = (0..n)
.map(|o| 1.0e-3 + (o as f32 + salt as f32 * 0.5) * 1.0e-4)
.collect();
QInt8::new(w, scales, n, k)
};
let bits = |s: &[f32]| s.iter().map(|f| f.to_bits()).collect::<Vec<u32>>();
let (n, k) = (150usize, 96usize);
let qw = mk(n, k, 3);
let x: Vec<f32> = (0..k)
.map(|i| (i as f32 * 0.41).sin() * 2.5 - 0.4)
.collect();
let monolithic = gemv_i8(&x, &qw);
for &tiles in &[1usize, 2, 7, 13, 64, 150, 200] {
let sharded = gemv_i8_sharded(&x, &qw, tiles);
assert_eq!(
bits(&monolithic),
bits(&sharded),
"int8 lm_head: {tiles}-tile shard != monolithic gemv_i8 (n={n} k={k})"
);
}
}
#[test]
fn lmhead_shard_f32_is_byte_identical_to_monolithic() {
let (n, k) = (150usize, 96usize);
let w: Vec<f32> = (0..n * k)
.map(|idx| ((idx as f32 * 0.013).sin() * 1.7) - 0.3)
.collect();
let x: Vec<f32> = (0..k)
.map(|i| (i as f32 * 0.29).cos() * 2.1 + 0.2)
.collect();
let bits = |s: &[f32]| s.iter().map(|f| f.to_bits()).collect::<Vec<u32>>();
let monolithic = gemv(&x, &w, n, k);
for &tiles in &[1usize, 2, 7, 13, 64, 150, 200] {
let sharded = gemv_sharded(&x, &w, n, k, tiles);
assert_eq!(
bits(&monolithic),
bits(&sharded),
"f32 lm_head: {tiles}-tile shard != monolithic gemv (n={n} k={k})"
);
}
}
#[test]
fn lmhead_shard_defaults_off_with_positive_tiles() {
if std::env::var_os("FOCR_LMHEAD_SHARD").is_none() {
assert!(!lmhead_shard_enabled());
}
assert!(lmhead_shard_tiles() >= 1);
}
fn assert_err_contains<T>(res: FocrResult<T>, needle: &str) {
let message = match res {
Ok(_) => String::from("<ok>"),
Err(err) => err.to_string(),
};
assert!(
message.contains(needle),
"error {message:?} did not contain {needle:?}"
);
}
#[test]
fn token_major_to_head_major_transposes_kv() -> FocrResult<()> {
let k = Mat::from_vec(2, 4, vec![0.0, 1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0]);
let v = Mat::from_vec(2, 4, vec![10.0, 11.0, 12.0, 13.0, 14.0, 15.0, 16.0, 17.0]);
let (kh, vh) = token_major_to_head_major(&k, &v, 2, 2, 2)?;
assert_eq!(kh, vec![0.0, 1.0, 4.0, 5.0, 2.0, 3.0, 6.0, 7.0]);
assert_eq!(vh, vec![10.0, 11.0, 14.0, 15.0, 12.0, 13.0, 16.0, 17.0]);
Ok(())
}
#[test]
fn embed_tokens_gathers_rows() -> FocrResult<()> {
let table = vec![0.0, 1.0, 2.0, 3.0, 4.0, 5.0];
let out = embed_tokens(&table, 3, 2, &[2, 0, 1, 2])?;
assert_eq!(out.shape(), (4, 2));
assert_eq!(out.row(0), &[4.0, 5.0]);
assert_eq!(out.row(1), &[0.0, 1.0]);
assert_eq!(out.row(2), &[2.0, 3.0]);
assert_eq!(out.row(3), &[4.0, 5.0]);
Ok(())
}
#[test]
fn embed_tokens_rejects_oob_id() {
let table = vec![0.0, 1.0, 2.0, 3.0];
assert!(embed_tokens(&table, 2, 2, &[5]).is_err());
}
#[test]
fn embed_tokens_rejects_bad_table_len() {
let table = vec![0.0, 1.0, 2.0];
assert!(embed_tokens(&table, 2, 2, &[0]).is_err());
}
#[test]
fn embed_tokens_rejects_table_shape_product_overflow_without_panic() {
assert_err_contains(embed_tokens(&[], usize::MAX, 2, &[]), "vocab*hidden");
}
#[test]
fn embed_tokens_rejects_output_shape_product_overflow_without_panic() {
assert_err_contains(embed_tokens(&[], 0, usize::MAX, &[0, 1]), "seq*hidden");
}
#[test]
fn rope_position_zero_is_identity() -> FocrResult<()> {
let rope = RopeTable::build(&[0], 4, 10000.0);
let mut x = Mat::from_vec(1, 4, vec![1.0, 2.0, 3.0, 4.0]);
apply_rope(&mut x, &rope)?;
for (got, want) in x.data.iter().zip([1.0f32, 2.0, 3.0, 4.0].iter()) {
assert!((got - want).abs() < 1e-6, "{got} != {want}");
}
Ok(())
}
#[test]
fn rope_matches_hand_computed_rotate_half() -> FocrResult<()> {
let theta = 10000.0f32;
let pos = 1usize;
let rope = RopeTable::build(&[pos], 2, theta);
let (a, b) = (3.0f32, 5.0f32);
let mut x = Mat::from_vec(1, 2, vec![a, b]);
apply_rope(&mut x, &rope)?;
let (s, c) = (pos as f32).sin_cos();
let want0 = a * c - b * s;
let want1 = b * c + a * s;
assert!((x.data[0] - want0).abs() < 1e-5, "{} != {want0}", x.data[0]);
assert!((x.data[1] - want1).abs() < 1e-5, "{} != {want1}", x.data[1]);
Ok(())
}
#[test]
fn rope_preserves_norm_per_head() -> FocrResult<()> {
let rope = RopeTable::build(&[7, 13], 4, 10000.0);
let mut x = Mat::from_vec(2, 8, (0..16).map(|i| (i as f32) * 0.25 - 2.0).collect());
let head_norms = |data: &[f32]| -> Vec<f32> {
let mut out = Vec::with_capacity(4);
for t in 0..2 {
for h in 0..2 {
let r = &data[t * 8 + h * 4..t * 8 + h * 4 + 4];
out.push(r.iter().map(|v| v * v).sum::<f32>());
}
}
out
};
let before = head_norms(&x.data);
apply_rope(&mut x, &rope)?;
let after = head_norms(&x.data);
for (b, a) in before.iter().zip(after.iter()) {
assert!((b - a).abs() < 1e-4, "norm changed: {b} -> {a}");
}
Ok(())
}
#[test]
fn rope_two_heads_share_phase() -> FocrResult<()> {
let rope = RopeTable::build(&[2], 2, 10000.0);
let mut x = Mat::from_vec(1, 4, vec![1.0, 0.0, 0.0, 1.0]);
apply_rope(&mut x, &rope)?;
let (s, c) = (2.0f32).sin_cos();
assert!((x.data[0] - c).abs() < 1e-5);
assert!((x.data[1] - s).abs() < 1e-5);
assert!((x.data[2] - (-s)).abs() < 1e-5);
assert!((x.data[3] - c).abs() < 1e-5);
Ok(())
}
#[test]
fn rope_rejects_bad_shape() {
let rope = RopeTable::build(&[0, 1], 4, 10000.0);
let mut bad = Mat::zeros(2, 3);
assert!(apply_rope(&mut bad, &rope).is_err());
let mut wrong_rows = Mat::zeros(3, 4);
assert!(apply_rope(&mut wrong_rows, &rope).is_err());
}
#[test]
fn rope_rejects_malformed_backing_data_without_panic() {
let rope = RopeTable::build(&[0], 4, 10000.0);
let mut bad = Mat {
rows: 1,
cols: 4,
data: vec![1.0, 2.0, 3.0],
};
assert_err_contains(apply_rope(&mut bad, &rope), "apply_rope x: data len 3");
}
#[test]
#[should_panic(expected = "RopeTable: seq*head_dim overflow")]
fn rope_table_rejects_shape_product_overflow_before_allocating() {
let _ = RopeTable::build(&[0, 1], usize::MAX - 1, 10000.0);
}
#[test]
fn rope_table_empty_sequence_does_not_allocate_by_head_dim() {
let rope = RopeTable::build(&[], usize::MAX - 1, 10000.0);
assert_eq!(rope.head_dim, usize::MAX - 1);
assert!(rope.cos.is_empty());
assert!(rope.sin.is_empty());
}
#[test]
fn dense_mlp_matches_hand_computed() -> FocrResult<()> {
let x = Mat::from_vec(1, 2, vec![1.0, 0.0]);
let gate_w = vec![1.0, 0.0, 0.0, 1.0];
let up_w = vec![1.0, 1.0, 1.0, 1.0];
let down_w = vec![1.0, 0.0, 0.0, 1.0];
let out = dense_mlp(&x, &gate_w, &up_w, &down_w, 2, 2)?;
assert_eq!(out.shape(), (1, 2));
let silu1 = 1.0f32 / (1.0 + (-1.0f32).exp());
assert!((out.data[0] - silu1).abs() < 1e-5, "{}", out.data[0]);
assert!(out.data[1].abs() < 1e-6, "{}", out.data[1]);
Ok(())
}
#[test]
fn dense_mlp_rejects_bad_hidden() {
let x = Mat::from_vec(1, 3, vec![1.0, 0.0, 0.0]);
assert!(dense_mlp(&x, &[1.0, 0.0], &[1.0, 0.0], &[1.0, 0.0], 2, 1).is_err());
}
#[test]
fn linear_no_bias_transposes_pytorch_layout() -> FocrResult<()> {
let x = Mat::from_vec(1, 2, vec![1.0, 2.0]);
let w = vec![1.0, 0.0, 0.0, 1.0, 1.0, 1.0];
let y = linear_no_bias(&x, &w, 2, 3)?;
assert_eq!(y.shape(), (1, 3));
assert_eq!(y.data, vec![1.0, 2.0, 3.0]);
Ok(())
}
#[test]
fn linear_no_bias_transposes_multirow_nonsquare_layout() -> FocrResult<()> {
let x = Mat::from_vec(2, 3, vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0]);
let w = vec![
1.0, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0, 0.0, 1.0, 1.0, 1.0, 1.0, ];
let y = linear_no_bias(&x, &w, 3, 4)?;
assert_eq!(y.shape(), (2, 4));
assert_eq!(y.data, vec![1.0, 2.0, 3.0, 6.0, 4.0, 5.0, 6.0, 15.0]);
Ok(())
}
#[test]
fn linear_no_bias_rejects_weight_shape_product_overflow_without_panic() {
let x = Mat::zeros(1, 2);
assert_err_contains(linear_no_bias(&x, &[], 2, usize::MAX), "out*in");
}
#[test]
fn linear_no_bias_rejects_malformed_input_backing_data_without_panic() {
let x = Mat {
rows: 1,
cols: 2,
data: vec![1.0],
};
assert_err_contains(
linear_no_bias(&x, &[1.0, 0.0], 2, 1),
"linear_no_bias x: data len 1",
);
}
#[test]
fn add_residual_sums_elementwise() -> FocrResult<()> {
let a = Mat::from_vec(2, 2, vec![1.0, 2.0, 3.0, 4.0]);
let b = Mat::from_vec(2, 2, vec![10.0, 20.0, 30.0, 40.0]);
let c = add_residual(&a, &b)?;
assert_eq!(c.data, vec![11.0, 22.0, 33.0, 44.0]);
Ok(())
}
#[test]
fn add_residual_rejects_shape_mismatch() {
let a = Mat::zeros(2, 2);
let b = Mat::zeros(2, 3);
assert!(add_residual(&a, &b).is_err());
}
#[test]
fn add_residual_rejects_malformed_backing_data_without_panic() {
let a = Mat {
rows: 1,
cols: 2,
data: vec![1.0],
};
let b = Mat::zeros(1, 2);
assert_err_contains(add_residual(&a, &b), "add_residual lhs: data len 1");
}
#[test]
fn lm_head_proj_matches_matmul() -> FocrResult<()> {
let h = Mat::from_vec(1, 2, vec![3.0, 4.0]);
let head_w = vec![1.0, 0.0, 0.0, 1.0, 1.0, 1.0];
let logits = lm_head_proj(&h, &head_w, 3)?;
assert_eq!(logits.shape(), (1, 3));
assert_eq!(logits.data, vec![3.0, 4.0, 7.0]);
Ok(())
}
#[test]
fn norm_and_lm_head_composes_rmsnorm_then_head() -> FocrResult<()> {
let h = Mat::from_vec(1, 2, vec![3.0, 4.0]);
let norm_w = vec![1.0, 1.0];
let head_w = vec![1.0, 0.0, 0.0, 1.0];
let logits = norm_and_lm_head(&h, &norm_w, &head_w, 2, 0.0)?;
let rstd = 1.0f32 / 12.5f32.sqrt();
assert!((logits.data[0] - 3.0 * rstd).abs() < 1e-5);
assert!((logits.data[1] - 4.0 * rstd).abs() < 1e-5);
Ok(())
}
#[test]
fn lm_head_last_row_is_full_last_row() -> FocrResult<()> {
let (seq, hidden, vocab) = (5usize, 8usize, 7usize);
let h: Vec<f32> = (0..seq * hidden)
.map(|i| ((i as f32) * 0.37).sin())
.collect();
let full = Mat::from_vec(seq, hidden, h);
let norm_w: Vec<f32> = (0..hidden).map(|i| 0.5 + (i as f32) * 0.1).collect();
let head_w: Vec<f32> = (0..vocab * hidden)
.map(|i| ((i as f32) * 0.13).cos())
.collect();
let eps = 1e-6;
let full_logits = norm_and_lm_head(&full, &norm_w, &head_w, vocab, eps)?;
let last = Mat::from_vec(1, hidden, full.row(seq - 1).to_vec());
let last_logits = norm_and_lm_head(&last, &norm_w, &head_w, vocab, eps)?;
assert_eq!((last_logits.rows, last_logits.cols), (1, vocab));
let full_last = &full_logits.data[(seq - 1) * vocab..seq * vocab];
assert_eq!(
last_logits.data, full_last,
"last-row lm_head must be bit-identical to full[last]"
);
Ok(())
}
#[test]
fn attn_output_proj_projects_context() -> FocrResult<()> {
let ctx = Mat::from_vec(1, 2, vec![1.0, 1.0]);
let o = vec![2.0, 0.0, 0.0, 3.0];
let out = attn_output_proj(&ctx, &o, 2, 2)?;
assert_eq!(out.data, vec![2.0, 3.0]);
Ok(())
}
#[test]
fn qkv_with_rope_shapes_and_pos0_identity() -> FocrResult<()> {
let normed = Mat::from_vec(1, 2, vec![1.0, 2.0]);
let ident = vec![1.0, 0.0, 0.0, 1.0];
let lw = LayerWeights {
input_ln: &[1.0, 1.0],
post_attn_ln: &[1.0, 1.0],
q_proj: &ident,
k_proj: &ident,
v_proj: &ident,
o_proj: &ident,
gate_w: &[],
up_w: &[],
down_w: &[],
};
let rope = RopeTable::build(&[0], 2, 10000.0);
let (q, k, v) = qkv_with_rope(&normed, &lw, &rope, 2, 2)?;
assert_eq!(q.shape(), (1, 2));
assert_eq!(v.data, vec![1.0, 2.0]); assert_eq!(q.data, vec![1.0, 2.0]); assert_eq!(k.data, vec![1.0, 2.0]);
Ok(())
}
#[test]
fn layer_forward_pre_norm_residual_identity_path() -> FocrResult<()> {
let hidden = 2usize;
let qkv_dim = 2usize;
let x = Mat::from_vec(2, hidden, vec![1.0, 2.0, 3.0, 4.0]);
let ident = vec![1.0, 0.0, 0.0, 1.0];
let lw = LayerWeights {
input_ln: &[1.0, 1.0],
post_attn_ln: &[1.0, 1.0],
q_proj: &ident,
k_proj: &ident,
v_proj: &ident,
o_proj: &ident,
gate_w: &[],
up_w: &[],
down_w: &[],
};
let rope = RopeTable::build(&[0, 1], qkv_dim, 10000.0);
let out = layer_forward(
&x,
&lw,
&rope,
hidden,
qkv_dim,
1e-6,
|q, _k, _v| Ok(Mat::zeros(q.rows, q.cols)),
|n| Ok(Mat::zeros(n.rows, n.cols)),
)?;
assert_eq!(out.shape(), (2, hidden));
for (got, want) in out.data.iter().zip(x.data.iter()) {
assert!((got - want).abs() < 1e-6, "{got} != {want}");
}
Ok(())
}
#[test]
fn layer_forward_adds_both_sublayers() -> FocrResult<()> {
let hidden = 2usize;
let qkv_dim = 2usize;
let x = Mat::from_vec(1, hidden, vec![10.0, 20.0]);
let ident = vec![1.0, 0.0, 0.0, 1.0];
let lw = LayerWeights {
input_ln: &[1.0, 1.0],
post_attn_ln: &[1.0, 1.0],
q_proj: &ident,
k_proj: &ident,
v_proj: &ident,
o_proj: &ident, gate_w: &[],
up_w: &[],
down_w: &[],
};
let rope = RopeTable::build(&[0], qkv_dim, 10000.0);
let out = layer_forward(
&x,
&lw,
&rope,
hidden,
qkv_dim,
1e-6,
|_q, _k, _v| Ok(Mat::from_vec(1, qkv_dim, vec![1.0, 1.0])),
|_n| Ok(Mat::from_vec(1, hidden, vec![100.0, 100.0])),
)?;
assert!((out.data[0] - 111.0).abs() < 1e-4, "{}", out.data[0]);
assert!((out.data[1] - 121.0).abs() < 1e-4, "{}", out.data[1]);
Ok(())
}
#[test]
fn top_level_entrypoints_error_cleanly_on_empty_weights() {
let w = Weights::default();
let h = Mat::zeros(1, config::HIDDEN_SIZE);
assert!(matches!(forward(&w, &h), Err(FocrError::FormatMismatch(_))));
assert!(matches!(lm_head(&w, &h), Err(FocrError::FormatMismatch(_))));
}
#[test]
fn prefill_attention_shapes_and_first_token_self_only() -> FocrResult<()> {
let (num_heads, head_dim, seq) = (1usize, 2usize, 2usize);
let q = Mat::from_vec(seq, 2, vec![1.0, 0.0, 0.0, 1.0]);
let k = Mat::from_vec(seq, 2, vec![1.0, 0.0, 0.0, 1.0]);
let v = Mat::from_vec(seq, 2, vec![5.0, 6.0, 7.0, 8.0]);
let out = prefill_attention(&q, &k, &v, num_heads, head_dim)?;
assert_eq!(out.shape(), (seq, num_heads * head_dim));
assert!((out.data[0] - 5.0).abs() < 1e-5, "{}", out.data[0]);
assert!((out.data[1] - 6.0).abs() < 1e-5, "{}", out.data[1]);
assert!(out.data[2] > 5.0 && out.data[2] < 7.0, "{}", out.data[2]);
assert!(out.data[3] > 6.0 && out.data[3] < 8.0, "{}", out.data[3]);
Ok(())
}
#[test]
fn prefill_attention_rejects_bad_shape() {
let q = Mat::zeros(2, 3); let k = Mat::zeros(2, 4);
let v = Mat::zeros(2, 4);
assert!(prefill_attention(&q, &k, &v, 2, 2).is_err());
}
#[test]
fn prefill_attention_rejects_malformed_backing_data_without_panic() {
let q = Mat {
rows: 2,
cols: 4,
data: vec![0.0; 7],
};
let k = Mat::zeros(2, 4);
let v = Mat::zeros(2, 4);
assert_err_contains(
prefill_attention(&q, &k, &v, 2, 2),
"prefill_attention q: data len 7",
);
}
}