use half::bf16;
use oxibonsai_kernels::softmax_simd;
use crate::gemm::{dot, gemm_abt};
use crate::math::silu;
use crate::te::config::{TeConfig, STACK_LAYERS};
use crate::te::error::{TeError, TeResult};
use crate::te::rope::Qwen3Rope;
use crate::te::weights::TeWeights;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Precision {
F32,
Bf16Storage,
}
fn round_bf16(x: &mut [f32]) {
par_flat(x, 1 << 16, |chunk| {
for v in chunk.iter_mut() {
*v = bf16::from_f32(*v).to_f32();
}
});
}
fn par_flat<F>(data: &mut [f32], min_chunk: usize, body: F)
where
F: Fn(&mut [f32]) + Sync,
{
let n = data.len();
let threads = std::thread::available_parallelism()
.map(|t| t.get())
.unwrap_or(1);
if threads <= 1 || n < 2 * min_chunk {
body(data);
return;
}
let nthreads = threads.min(n / min_chunk).max(1);
let per = n.div_ceil(nthreads);
let body_ref = &body;
std::thread::scope(|scope| {
for chunk in data.chunks_mut(per) {
scope.spawn(move || body_ref(chunk));
}
});
}
pub struct TeOutput {
pub hidden_states: Vec<Vec<f32>>,
pub seq: usize,
pub hidden: usize,
}
impl TeOutput {
pub fn cond_7680(&self) -> TeResult<Vec<f32>> {
let seq = self.seq;
let hidden = self.hidden;
for &l in &STACK_LAYERS {
if l >= self.hidden_states.len() {
return Err(TeError::Shape(format!(
"stack layer {l} out of range (have {})",
self.hidden_states.len()
)));
}
}
let groups = STACK_LAYERS.len();
let mut out = vec![0.0f32; seq * groups * hidden];
for t in 0..seq {
for (g, &l) in STACK_LAYERS.iter().enumerate() {
let src = &self.hidden_states[l][t * hidden..(t + 1) * hidden];
let dst_base = t * groups * hidden + g * hidden;
out[dst_base..dst_base + hidden].copy_from_slice(src);
}
}
Ok(out)
}
}
pub struct TextEncoder<'w> {
weights: &'w TeWeights,
cfg: TeConfig,
precision: Precision,
}
impl<'w> TextEncoder<'w> {
pub fn new(weights: &'w TeWeights) -> Self {
Self::with_precision(weights, Precision::F32)
}
pub fn with_precision(weights: &'w TeWeights, precision: Precision) -> Self {
let cfg = weights.config().clone();
Self {
weights,
cfg,
precision,
}
}
pub fn config(&self) -> &TeConfig {
&self.cfg
}
pub fn precision(&self) -> Precision {
self.precision
}
#[inline]
fn quantize(&self, x: &mut [f32]) {
if self.precision == Precision::Bf16Storage {
round_bf16(x);
}
}
pub fn forward(&self, input_ids: &[u32], attention_mask: &[i32]) -> TeResult<TeOutput> {
let seq = input_ids.len();
if attention_mask.len() != seq {
return Err(TeError::Shape(format!(
"attention_mask len {} != input_ids len {seq}",
attention_mask.len()
)));
}
let hidden = self.cfg.hidden_size;
let head_dim = self.cfg.head_dim;
let n_q = self.cfg.num_attention_heads;
let n_kv = self.cfg.num_key_value_heads;
let eps = self.cfg.rms_norm_eps;
let mut h = self
.weights
.embed_gather("embed_tokens", input_ids, hidden)?;
self.quantize(&mut h);
let mut hidden_states: Vec<Vec<f32>> = Vec::with_capacity(self.cfg.num_layers + 1);
hidden_states.push(h.clone());
let rope = Qwen3Rope::new(seq, head_dim, self.cfg.rope_theta);
let mask = build_mask(attention_mask, seq);
let timing = std::env::var("OXI_IMAGE_TIMING").is_ok();
TE_MATMUL_NS.store(0, std::sync::atomic::Ordering::Relaxed);
TE_ATTN_NS.store(0, std::sync::atomic::Ordering::Relaxed);
for layer in 0..self.cfg.num_layers {
self.decoder_layer(
&mut h, layer, seq, &rope, &mask, eps, hidden, head_dim, n_q, n_kv,
)?;
hidden_states.push(h.clone());
}
if timing {
let mm = TE_MATMUL_NS.load(std::sync::atomic::Ordering::Relaxed) as f64 / 1e9;
let at = TE_ATTN_NS.load(std::sync::atomic::Ordering::Relaxed) as f64 / 1e9;
eprintln!(
"[timing] TE matmul(GPU)={mm:.2}s attention(CPU)={at:.2}s \
(remainder = norms/rope/quantize/reshape/embed)"
);
}
Ok(TeOutput {
hidden_states,
seq,
hidden,
})
}
#[allow(clippy::too_many_arguments)]
fn decoder_layer(
&self,
h: &mut [f32],
layer: usize,
seq: usize,
rope: &Qwen3Rope,
mask: &[f32],
eps: f32,
hidden: usize,
head_dim: usize,
n_q: usize,
n_kv: usize,
) -> TeResult<()> {
let pfx = format!("layers.{layer}");
let q_dim = self.cfg.q_dim();
let kv_dim = self.cfg.kv_dim();
let in_ln = self
.weights
.vec1(&format!("{pfx}.input_layernorm"), hidden)?;
let mut x = rms_norm(h, seq, hidden, &in_ln.data, eps);
self.quantize(&mut x);
let wq = self
.weights
.linear(&format!("{pfx}.self_attn.q_proj"), q_dim, hidden)?;
let wk = self
.weights
.linear(&format!("{pfx}.self_attn.k_proj"), kv_dim, hidden)?;
let wv = self
.weights
.linear(&format!("{pfx}.self_attn.v_proj"), kv_dim, hidden)?;
let mut q = matmul(&x, &wq.data, seq, q_dim, hidden)?;
let mut k = matmul(&x, &wk.data, seq, kv_dim, hidden)?;
let mut v = matmul(&x, &wv.data, seq, kv_dim, hidden)?;
self.quantize(&mut q);
self.quantize(&mut k);
self.quantize(&mut v);
let q_norm = self
.weights
.vec1(&format!("{pfx}.self_attn.q_norm"), head_dim)?;
let k_norm = self
.weights
.vec1(&format!("{pfx}.self_attn.k_norm"), head_dim)?;
rms_norm_heads(&mut q, seq * n_q, head_dim, &q_norm.data, eps);
rms_norm_heads(&mut k, seq * n_kv, head_dim, &k_norm.data, eps);
self.quantize(&mut q);
self.quantize(&mut k);
let mut qh = token_to_head_major(&q, seq, n_q, head_dim);
let mut kh = token_to_head_major(&k, seq, n_kv, head_dim);
let vh = token_to_head_major(&v, seq, n_kv, head_dim);
rope.apply(&mut qh, n_q, seq);
rope.apply(&mut kh, n_kv, seq);
self.quantize(&mut qh);
self.quantize(&mut kh);
let t_attn = std::time::Instant::now();
let mut attn = self.gqa_attention(&qh, &kh, &vh, mask, seq, head_dim, n_q, n_kv)?;
TE_ATTN_NS.fetch_add(
t_attn.elapsed().as_nanos() as u64,
std::sync::atomic::Ordering::Relaxed,
);
self.quantize(&mut attn);
let wo = self
.weights
.linear(&format!("{pfx}.self_attn.o_proj"), hidden, q_dim)?;
let mut attn_out = matmul(&attn, &wo.data, seq, hidden, q_dim)?;
self.quantize(&mut attn_out);
for (hv, av) in h.iter_mut().zip(attn_out.iter()) {
*hv += *av;
}
self.quantize(h);
let post_ln = self
.weights
.vec1(&format!("{pfx}.post_attention_layernorm"), hidden)?;
x = rms_norm(h, seq, hidden, &post_ln.data, eps);
self.quantize(&mut x);
let inter = self.cfg.intermediate_size;
let wgate = self
.weights
.linear(&format!("{pfx}.mlp.gate_proj"), inter, hidden)?;
let wup = self
.weights
.linear(&format!("{pfx}.mlp.up_proj"), inter, hidden)?;
let mut gate = matmul(&x, &wgate.data, seq, inter, hidden)?;
let mut up = matmul(&x, &wup.data, seq, inter, hidden)?;
self.quantize(&mut gate);
self.quantize(&mut up);
let mut act = vec![0.0f32; seq * inter];
par_heads(&mut act, seq, inter, |r, dst| {
let base = r * inter;
for (j, d) in dst.iter_mut().enumerate() {
*d = silu(gate[base + j]) * up[base + j];
}
});
self.quantize(&mut act);
let wdown = self
.weights
.linear(&format!("{pfx}.mlp.down_proj"), hidden, inter)?;
let mut down = matmul(&act, &wdown.data, seq, hidden, inter)?;
self.quantize(&mut down);
for (hv, dv) in h.iter_mut().zip(down.iter()) {
*hv += *dv;
}
self.quantize(h);
Ok(())
}
#[allow(clippy::too_many_arguments)]
fn gqa_attention(
&self,
q: &[f32],
k: &[f32],
v: &[f32],
mask: &[f32],
seq: usize,
head_dim: usize,
n_q: usize,
n_kv: usize,
) -> TeResult<Vec<f32>> {
let kv_group = n_q / n_kv;
let scale = 1.0f32 / (head_dim as f32).sqrt();
let inner = n_q * head_dim;
let mut out = vec![0.0f32; seq * inner];
let attend = |hq: usize, dst: &mut [f32]| {
let kv = hq / kv_group;
let q_off = hq * seq * head_dim;
let kv_off = kv * seq * head_dim;
let mut scores = vec![0.0f32; seq];
for qi in 0..seq {
let q_row = &q[q_off + qi * head_dim..q_off + (qi + 1) * head_dim];
let mrow = &mask[qi * seq..(qi + 1) * seq];
for (ki, score) in scores.iter_mut().enumerate() {
let k_row = &k[kv_off + ki * head_dim..kv_off + (ki + 1) * head_dim];
*score = dot(q_row, k_row, head_dim) * scale + mrow[ki];
}
softmax_simd(&mut scores);
let o = &mut dst[qi * head_dim..(qi + 1) * head_dim];
for d in o.iter_mut() {
*d = 0.0;
}
for (ki, &w) in scores.iter().enumerate() {
if w == 0.0 {
continue;
}
let v_row = &v[kv_off + ki * head_dim..kv_off + (ki + 1) * head_dim];
crate::gemm::axpy(o, w, v_row, head_dim);
}
}
};
let mut head_out = vec![0.0f32; n_q * seq * head_dim];
par_heads(&mut head_out, n_q, seq * head_dim, attend);
for hq in 0..n_q {
for qi in 0..seq {
let src = &head_out[(hq * seq + qi) * head_dim..(hq * seq + qi + 1) * head_dim];
let dst = &mut out[qi * inner + hq * head_dim..qi * inner + (hq + 1) * head_dim];
dst.copy_from_slice(src);
}
}
Ok(out)
}
}
fn par_heads<F>(out: &mut [f32], heads: usize, width: usize, body: F)
where
F: Fn(usize, &mut [f32]) + Sync,
{
debug_assert_eq!(out.len(), heads * width);
let threads = std::thread::available_parallelism()
.map(|n| n.get())
.unwrap_or(1)
.min(heads.max(1));
if threads <= 1 || heads < 4 {
for (hidx, chunk) in out.chunks_mut(width).enumerate() {
body(hidx, chunk);
}
return;
}
let per = heads.div_ceil(threads);
let body_ref = &body;
std::thread::scope(|scope| {
let mut base = 0usize;
for chunk in out.chunks_mut(per * width) {
let start = base;
let chunk_heads = chunk.len() / width;
base += chunk_heads;
scope.spawn(move || {
for r in 0..chunk_heads {
let slab = &mut chunk[r * width..(r + 1) * width];
body_ref(start + r, slab);
}
});
}
});
}
fn build_mask(attention_mask: &[i32], seq: usize) -> Vec<f32> {
let neg_inf = f32::NEG_INFINITY;
let mut mask = vec![0.0f32; seq * seq];
for i in 0..seq {
for j in 0..seq {
let blocked = j > i || attention_mask[j] == 0;
mask[i * seq + j] = if blocked { neg_inf } else { 0.0 };
}
}
mask
}
static TE_MATMUL_NS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
static TE_ATTN_NS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
fn matmul(input: &[f32], weight: &[f32], m: usize, n: usize, k: usize) -> TeResult<Vec<f32>> {
let t = std::time::Instant::now();
let r = matmul_inner(input, weight, m, n, k);
TE_MATMUL_NS.fetch_add(
t.elapsed().as_nanos() as u64,
std::sync::atomic::Ordering::Relaxed,
);
r
}
fn matmul_inner(input: &[f32], weight: &[f32], m: usize, n: usize, k: usize) -> TeResult<Vec<f32>> {
if input.len() != m * k {
return Err(TeError::Shape(format!(
"matmul input len {} != m*k {}",
input.len(),
m * k
)));
}
if weight.len() != n * k {
return Err(TeError::Shape(format!(
"matmul weight len {} != n*k {}",
weight.len(),
n * k
)));
}
let mut out = vec![0.0f32; m * n];
#[cfg(all(feature = "metal", target_os = "macos"))]
{
if crate::te::gpu::te_gpu_enabled() {
match crate::te::gpu::te_matmul_gpu(weight, input, &mut out, m, n, k) {
Ok(()) => return Ok(out),
Err(_e) => {
}
}
}
}
#[cfg(all(
feature = "native-cuda",
any(target_os = "linux", target_os = "windows")
))]
{
if crate::te::cuda_gpu::te_gpu_enabled() {
match crate::te::cuda_gpu::te_matmul_gpu(weight, input, &mut out, m, n, k) {
Ok(()) => return Ok(out),
Err(_e) => {
}
}
}
}
gemm_abt(input, weight, &mut out, m, n, k);
Ok(out)
}
fn rms_norm(x: &[f32], rows: usize, dim: usize, weight: &[f32], eps: f32) -> Vec<f32> {
debug_assert_eq!(x.len(), rows * dim);
debug_assert_eq!(weight.len(), dim);
let inv_dim = 1.0f32 / dim as f32;
let mut out = vec![0.0f32; rows * dim];
par_heads(&mut out, rows, dim, |r, dst| {
let src = &x[r * dim..(r + 1) * dim];
let mut ms = 0.0f32;
for &v in src {
ms += v * v;
}
ms *= inv_dim;
let inv_rms = 1.0f32 / (ms + eps).sqrt();
for i in 0..dim {
dst[i] = weight[i] * src[i] * inv_rms;
}
});
out
}
fn rms_norm_heads(x: &mut [f32], rows: usize, head_dim: usize, weight: &[f32], eps: f32) {
debug_assert_eq!(weight.len(), head_dim);
let inv_dim = 1.0f32 / head_dim as f32;
par_heads(x, rows, head_dim, |_, row| {
let mut ms = 0.0f32;
for &v in row.iter() {
ms += v * v;
}
ms *= inv_dim;
let inv_rms = 1.0f32 / (ms + eps).sqrt();
for (i, v) in row.iter_mut().enumerate() {
*v = weight[i] * *v * inv_rms;
}
});
}
fn token_to_head_major(x: &[f32], seq: usize, num_heads: usize, head_dim: usize) -> Vec<f32> {
let inner = num_heads * head_dim;
debug_assert_eq!(x.len(), seq * inner);
let mut out = vec![0.0f32; seq * inner];
par_heads(&mut out, num_heads, seq * head_dim, |hh, head_block| {
for t in 0..seq {
let src = &x[t * inner + hh * head_dim..t * inner + (hh + 1) * head_dim];
head_block[t * head_dim..(t + 1) * head_dim].copy_from_slice(src);
}
});
out
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn rms_norm_unit_weight_normalises() {
let x = vec![3.0f32, 4.0, 0.0, 0.0];
let w = vec![1.0f32; 4];
let y = rms_norm(&x, 1, 4, &w, 0.0);
assert!((y[0] - 1.2).abs() < 1e-6);
assert!((y[1] - 1.6).abs() < 1e-6);
}
#[test]
fn mask_is_causal_and_padding() {
let am = [1, 1, 1, 0];
let m = build_mask(&am, 4);
assert_eq!(m[0], 0.0);
assert!(m[1].is_infinite());
assert_eq!(m[2 * 4], 0.0);
assert_eq!(m[2 * 4 + 2], 0.0);
assert!(m[3 * 4 + 3].is_infinite());
}
#[test]
fn token_head_roundtrip() {
let x = vec![0.0, 1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0];
let hm = token_to_head_major(&x, 2, 2, 2);
assert_eq!(hm, vec![0.0, 1.0, 4.0, 5.0, 2.0, 3.0, 6.0, 7.0]);
}
}