use crate::attention::{self, QwenAttnCfg};
use crate::inference;
use crate::kv_cache::KvCache;
use crate::linear_core::{
gdn_forward, gdn_pair, vmf_phase_forward, vmf_phase_pair, GdnCfg, GdnWeights, VmfPhaseCfg,
VmfPhaseWeights,
};
use crate::pool::Pool;
use crate::qtensor::QTensor;
use crate::sampler::{self, SamplerConfig, SplitMix64};
use crate::tokenizer::Tokenizer;
use cortiq_core::mask::TaskMask;
use cortiq_core::types::NormStyle;
struct ForwardScratch {
n1: Vec<f32>,
n2: Vec<f32>,
p1: Vec<f32>,
p2: Vec<f32>,
}
impl ForwardScratch {
fn new(hidden: usize) -> Self {
Self {
n1: vec![0.0; hidden],
n2: vec![0.0; hidden],
p1: vec![0.0; hidden],
p2: vec![0.0; hidden],
}
}
}
pub struct Pipeline {
pub tokenizer: std::sync::Arc<Tokenizer>,
pub kv_cache: KvCache,
pub sampler_config: SamplerConfig,
pub weights: PipelineWeights,
pub hidden_size: usize,
pub intermediate_size: usize,
pub num_heads: usize,
pub num_kv_heads: usize,
pub head_dim: usize,
pub num_layers: usize,
pub vocab_size: usize,
pub rms_eps: f64,
pub rope_base: f32,
pub norm_style: NormStyle,
pub rotary_dim: usize,
pub vmf_cfg: Option<VmfPhaseCfg>,
pub gdn_cfg: Option<GdnCfg>,
pub mtp: Option<MtpModule>,
pub speculative: bool,
rng: SplitMix64,
inv_freq: std::sync::Arc<Vec<f32>>,
ws: ForwardScratch,
pool: Option<std::sync::Arc<Pool>>,
pub(crate) model: Option<std::sync::Arc<cortiq_core::CmfModel>>,
pub(crate) dyn_force_f32: bool,
pub(crate) dyn_skill_layers: Vec<Option<Vec<usize>>>,
pub(crate) dyn_active: Option<usize>,
pub(crate) dyn_blend_loaded: bool,
pub(crate) dyn_phi_layer: Option<usize>,
dyn_phi_ema: Vec<f32>,
dyn_phi_seen: usize,
pub dyn_router: Option<crate::swarm::DynRouter>,
o1_cfg: Option<crate::nystrom::O1Cfg>,
o1_flags: Vec<bool>,
trace: bool,
calib_temp: f32,
}
pub struct PipelineWeights {
pub embed_tokens: QTensor,
pub layers: Vec<LayerWeights>,
pub lm_head: QTensor,
pub final_norm: Vec<f32>,
}
pub struct LayerWeights {
pub input_norm: Vec<f32>,
pub post_norm: Vec<f32>,
pub ffn: FfnKind,
pub attn: AttnKind,
}
pub struct DenseFfn {
pub gate_proj: QTensor,
pub up_proj: QTensor,
pub down_proj: QTensor,
}
pub enum FfnKind {
Dense(DenseFfn),
Moe(MoeFfn),
}
pub struct MoeFfn {
pub router: QTensor,
pub experts: Vec<DenseFfn>,
pub top_k: usize,
pub norm_topk_prob: bool,
pub shared: Option<(DenseFfn, QTensor)>,
pub stats: std::cell::RefCell<Vec<u64>>,
}
pub enum AttnKind {
Full {
wq: QTensor,
wk: QTensor,
wv: QTensor,
wo: QTensor,
q_norm: Option<Vec<f32>>,
k_norm: Option<Vec<f32>>,
output_gate: bool,
bias: Option<(Vec<f32>, Vec<f32>, Vec<f32>)>,
},
Linear(VmfPhaseWeights),
LinearGdn(GdnWeights),
}
pub struct MtpModule {
pub enorm: Vec<f32>,
pub hnorm: Vec<f32>,
pub eh_proj: QTensor,
pub layer: LayerWeights,
pub final_norm: Vec<f32>,
pub kv: crate::kv_cache::LayerKvCache,
}
pub struct GenerateResult {
pub text: String,
pub token_ids: Vec<u32>,
pub prompt_tokens: usize,
pub tokens_generated: usize,
pub finish_reason: String,
pub mtp_drafted: usize,
pub mtp_accepted: usize,
pub token_confidence: Vec<f32>,
pub traces: Vec<TokenTrace>,
}
#[derive(Clone, Debug)]
pub struct TokenTrace {
pub t: usize,
pub token_id: u32,
pub confidence: f32,
pub active_skill: Option<String>,
pub recon: Option<f32>,
pub switched: bool,
}
fn top1_prob_t(logits: &[f32], id: u32, temp: f32) -> f32 {
let t = if temp > 1e-3 { temp } else { 1.0 };
let max = logits.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
let sum: f32 = logits.iter().map(|&v| ((v - max) / t).exp()).sum();
if sum > 0.0 {
(((logits[id as usize] - max) / t).exp()) / sum
} else {
0.0
}
}
fn prefill_batched() -> bool {
std::env::var("CMF_PREFILL").map(|v| v != "seq").unwrap_or(true)
}
pub type TokenCallback = Box<dyn FnMut(&str) -> bool + Send>;
impl Pipeline {
#[allow(clippy::too_many_arguments)]
pub fn new(
tokenizer: Tokenizer,
weights: PipelineWeights,
hidden_size: usize,
intermediate_size: usize,
num_heads: usize,
num_kv_heads: usize,
head_dim: usize,
num_layers: usize,
vocab_size: usize,
rms_eps: f64,
rope_base: f32,
norm_style: NormStyle,
max_seq_len: usize,
sampler_config: SamplerConfig,
) -> Self {
let rng = match sampler_config.seed {
Some(s) => SplitMix64::new(s),
None => SplitMix64::from_entropy(),
};
let inv_freq = std::sync::Arc::new(attention::rope_inv_freq(head_dim, rope_base));
let pool = Pool::from_env();
if let Some(p) = &pool {
tracing::info!("worker pool: {} threads", p.n_workers());
}
Self {
tokenizer: std::sync::Arc::new(tokenizer),
kv_cache: KvCache::new(num_layers, num_kv_heads, head_dim, max_seq_len),
sampler_config,
weights,
hidden_size,
intermediate_size,
num_heads,
num_kv_heads,
head_dim,
num_layers,
vocab_size,
rms_eps,
rope_base,
norm_style,
rotary_dim: head_dim,
vmf_cfg: None,
gdn_cfg: None,
mtp: None,
speculative: std::env::var("CMF_MTP").map(|v| v != "0").unwrap_or(true),
rng,
inv_freq,
ws: ForwardScratch::new(hidden_size),
pool,
model: None,
dyn_force_f32: false,
dyn_skill_layers: Vec::new(),
dyn_active: None,
dyn_blend_loaded: false,
dyn_phi_layer: None,
dyn_phi_ema: Vec::new(),
dyn_phi_seen: 0,
dyn_router: None,
o1_cfg: None,
o1_flags: Vec::new(),
trace: false,
calib_temp: 1.0,
}
}
pub fn set_o1(&mut self, cfg: Option<crate::nystrom::O1Cfg>) {
self.o1_flags = match &cfg {
Some(c) => {
let mut flags = c.layer_flags(self.num_layers);
for (li, f) in flags.iter_mut().enumerate() {
if *f && !matches!(self.weights.layers[li].attn, AttnKind::Full { .. }) {
*f = false;
}
}
flags
}
None => Vec::new(),
};
if let Some(c) = &cfg {
let n = self.o1_flags.iter().filter(|&&f| f).count();
tracing::info!(
"o1 nystrom attention: {n}/{} layer(s), m={} w={} sink={} rect={:?}",
self.num_layers, c.m, c.w, c.sink, c.rect
);
}
self.o1_cfg = cfg;
}
pub fn o1_active(&self) -> bool {
self.o1_cfg.is_some() && self.o1_flags.iter().any(|&f| f)
}
fn o1_begin(&mut self) {
if let Some(c) = &self.o1_cfg {
let (m, w, sink, rect) = (c.m, c.w, c.sink, c.rect);
for (li, &f) in self.o1_flags.iter().enumerate() {
if f {
self.kv_cache.layers[li].o1_begin(m, w, sink, rect);
}
}
}
}
fn o1_seal(&mut self) {
if self.o1_cfg.is_none() {
return;
}
for li in 0..self.num_layers {
if self.o1_flags.get(li).copied().unwrap_or(false) {
self.kv_cache.layers[li].o1_seal(self.num_heads);
}
}
}
pub fn set_trace(&mut self, on: bool) {
self.trace = on;
}
pub fn set_calib_temp(&mut self, t: f32) {
self.calib_temp = if t > 1e-3 { t } else { 1.0 };
}
pub fn calib_temp(&self) -> f32 {
self.calib_temp
}
pub fn set_rotary(&mut self, rotary_dim: usize, base: f32) {
self.rotary_dim = rotary_dim.min(self.head_dim);
self.inv_freq = std::sync::Arc::new(attention::rope_inv_freq(self.rotary_dim, base));
}
fn attn_cfg(&self, position: usize) -> QwenAttnCfg<'_> {
QwenAttnCfg {
num_heads: self.num_heads,
num_kv_heads: self.num_kv_heads,
head_dim: self.head_dim,
hidden_size: self.hidden_size,
position,
inv_freq: &self.inv_freq,
rotary_dim: self.rotary_dim,
q_norm: None,
k_norm: None,
output_gate: false,
bias: None,
rms_eps: self.rms_eps,
norm_style: self.norm_style,
pool: self.pool.as_deref(),
}
}
pub fn generate(
&mut self,
prompt: &str,
max_tokens: usize,
task_mask: Option<&TaskMask>,
on_token: Option<TokenCallback>,
) -> Result<GenerateResult, String> {
let input_ids = self.tokenizer.encode(prompt);
self.generate_from_ids(&input_ids, max_tokens, task_mask, on_token)
}
pub fn generate_from_ids(
&mut self,
input_ids: &[u32],
max_tokens: usize,
task_mask: Option<&TaskMask>,
mut on_token: Option<TokenCallback>,
) -> Result<GenerateResult, String> {
if input_ids.is_empty() {
return Err("empty prompt: nothing to generate from".to_string());
}
self.kv_cache.clear();
self.o1_begin();
let spec_active = self.speculative
&& self.mtp.is_some()
&& task_mask.is_none()
&& !self.o1_active()
&& self.sampler_config.temperature < 1e-6;
let mut mtp = if spec_active { self.mtp.take() } else { None };
if let Some(m) = &mut mtp {
m.kv.clear();
}
let mut router = if mtp.is_none() { self.dyn_router.take() } else { None };
if let Some(r) = &mut router {
r.reset(); self.dyn_phi_seen = 0; let _ = self.set_active_skill(None);
}
let mut all_ids = input_ids.to_vec();
let mut generated = 0usize;
let mut finish_reason = "max_tokens".to_string();
let mut drafted = 0usize;
let mut accepted = 0usize;
let mut confidence: Vec<f32> = Vec::new();
let trace_on = self.trace;
let calib_temp = self.calib_temp;
let mut traces: Vec<TokenTrace> = Vec::new();
let mut hidden = vec![0.0f32; self.hidden_size];
let mut pos = 0usize;
let dyn_prefill = router.is_some();
if task_mask.is_none() && !dyn_prefill && prefill_batched() && input_ids.len() > 2 {
const CHUNK: usize = 48;
let hs = self.hidden_size;
while pos < input_ids.len() {
let end = (pos + CHUNK).min(input_ids.len());
let hb = self.prefill_batch(&input_ids[pos..end], pos);
if let Some(m) = &mut mtp {
for p in pos..end {
if p + 1 < input_ids.len() {
let _ = self.mtp_step(
m,
&hb[(p - pos) * hs..(p - pos + 1) * hs],
input_ids[p + 1],
p,
);
}
}
}
hidden.copy_from_slice(&hb[(end - pos - 1) * hs..]);
pos = end;
}
}
if task_mask.is_none() && !dyn_prefill {
while pos + 1 < input_ids.len() {
let e1 = self.embed_single(input_ids[pos]);
let e2 = self.embed_single(input_ids[pos + 1]);
let (h1, h2) = self.forward_pair(&e1, &e2, pos);
self.commit_linear_scratch();
if let Some(m) = &mut mtp {
let _ = self.mtp_step(m, &h1, input_ids[pos + 1], pos);
if pos + 2 < input_ids.len() {
let _ = self.mtp_step(m, &h2, input_ids[pos + 2], pos + 1);
}
}
hidden = h2;
pos += 2;
}
}
while pos < input_ids.len() {
hidden = self.forward_layers(&self.embed_single(input_ids[pos]), pos, task_mask);
if let Some(m) = &mut mtp {
if pos + 1 < input_ids.len() {
let _ = self.mtp_step(m, &hidden, input_ids[pos + 1], pos);
}
}
pos += 1;
}
self.o1_seal();
macro_rules! commit {
($id:expr) => {{
all_ids.push($id);
generated += 1;
if self.tokenizer.is_eos($id) {
finish_reason = "stop".to_string();
false
} else {
let token_text = self.tokenizer.decode_token($id);
let mut go = true;
if let Some(ref mut cb) = on_token {
if !cb(&token_text) {
finish_reason = "cancelled".to_string();
go = false;
}
}
go
}
}};
}
let mut next_pos = input_ids.len();
'decode: while generated < max_tokens {
inference::rms_norm_into(
&hidden,
&self.weights.final_norm,
self.rms_eps,
self.norm_style,
&mut self.ws.n1,
);
let mut logits = self.lm_head_forward(&self.ws.n1);
let t_next = sampler::sample(&logits, &self.sampler_config, &all_ids, &mut self.rng);
confidence.push(top1_prob_t(&logits, t_next, calib_temp));
attention::recycle_buf(&mut logits);
if trace_on {
let skill = router.as_ref().and_then(|r| r.active_id());
traces.push(TokenTrace {
t: generated,
token_id: t_next,
confidence: *confidence.last().unwrap(),
active_skill: skill,
recon: None,
switched: false,
});
}
if !commit!(t_next) {
break 'decode;
}
if generated >= max_tokens {
break 'decode;
}
if self.kv_cache.needs_eviction() {
let keep = (self.kv_cache.max_seq_len / 2).max(1);
self.kv_cache.evict(keep);
}
match &mut mtp {
Some(m) if generated + 1 < max_tokens => {
let draft = self.mtp_step(m, &hidden, t_next, next_pos - 1);
drafted += 1;
let emb1 = self.embed_single(t_next);
let emb2 = self.embed_single(draft);
let (h1, h2) = self.forward_pair(&emb1, &emb2, next_pos);
inference::rms_norm_into(
&h1,
&self.weights.final_norm,
self.rms_eps,
self.norm_style,
&mut self.ws.n1,
);
let mut logits1 = self.lm_head_forward(&self.ws.n1);
let t_after =
sampler::sample(&logits1, &self.sampler_config, &all_ids, &mut self.rng);
confidence.push(top1_prob_t(&logits1, t_after, calib_temp));
attention::recycle_buf(&mut logits1);
if trace_on {
traces.push(TokenTrace {
t: generated,
token_id: t_after,
confidence: *confidence.last().unwrap(),
active_skill: None,
recon: None,
switched: false,
});
}
let stop = !commit!(t_after);
if t_after == draft {
accepted += 1;
self.commit_linear_scratch();
let _ = self.mtp_step(m, &h1, t_after, next_pos);
hidden = h2;
next_pos += 2;
} else {
for layer in &mut self.kv_cache.layers {
layer.truncate_last(1);
}
if !stop {
let _ = self.mtp_step(m, &h1, t_after, next_pos);
hidden = self
.forward_layers(&self.embed_single(t_after), next_pos + 1, None);
}
next_pos += 2;
}
if stop {
break 'decode;
}
}
_ => {
hidden = self.forward_layers(&self.embed_single(t_next), next_pos, task_mask);
next_pos += 1;
if let Some(r) = &mut router {
let phi = self.dyn_phi_ema.clone();
let decision = r.step(&phi, generated);
if let Some(new_active) = decision {
let _ = self.set_active_skill(new_active);
}
if trace_on {
if let Some(last) = traces.last_mut() {
let e = r.last_best_e();
last.recon = e.is_finite().then_some(e);
last.switched = decision.is_some();
}
}
}
}
}
}
if router.is_some() {
let _ = self.set_active_skill(None);
}
self.dyn_router = router.or(self.dyn_router.take());
self.mtp = mtp.or(self.mtp.take());
let output_ids = &all_ids[input_ids.len()..];
confidence.truncate(output_ids.len()); traces.truncate(output_ids.len());
Ok(GenerateResult {
text: self.tokenizer.decode(output_ids),
token_ids: output_ids.to_vec(),
prompt_tokens: input_ids.len(),
tokens_generated: generated,
finish_reason,
mtp_drafted: drafted,
mtp_accepted: accepted,
token_confidence: confidence,
traces,
})
}
fn mtp_step(&mut self, m: &mut MtpModule, hidden: &[f32], next_token: u32, position: usize) -> u32 {
let e = self.embed_single(next_token);
let mut cat = vec![0.0f32; 2 * self.hidden_size];
let (cat_e, cat_h) = cat.split_at_mut(self.hidden_size);
inference::rms_norm_into(&e, &m.enorm, self.rms_eps, self.norm_style, cat_e);
inference::rms_norm_into(hidden, &m.hnorm, self.rms_eps, self.norm_style, cat_h);
let mut x = vec![0.0f32; self.hidden_size];
m.eh_proj.matvec(&cat, &mut x, self.pool.as_deref());
let lw = &m.layer;
inference::rms_norm_into(&x, &lw.input_norm, self.rms_eps, self.norm_style, &mut self.ws.n1);
let attn = match &lw.attn {
AttnKind::Full {
wq,
wk,
wv,
wo,
q_norm,
k_norm,
output_gate,
bias,
} => {
let mut cfg = self.attn_cfg(position);
cfg.q_norm = q_norm.as_deref();
cfg.k_norm = k_norm.as_deref();
cfg.output_gate = *output_gate;
cfg.bias = bias
.as_ref()
.map(|(q, k, v)| (q.as_slice(), k.as_slice(), v.as_slice()));
attention::qwen_attention(&self.ws.n1, wq, wk, wv, wo, &mut m.kv, &cfg)
}
AttnKind::Linear(_) | AttnKind::LinearGdn(_) => {
unreachable!("MTP block is full attention")
}
};
for (i, &a) in attn.iter().enumerate() {
x[i] += a;
}
inference::rms_norm_into(&x, &lw.post_norm, self.rms_eps, self.norm_style, &mut self.ws.p1);
let ffn = ffn_forward(&lw.ffn, &self.ws.p1, self.pool.as_deref());
for (i, &f) in ffn.iter().enumerate() {
x[i] += f;
}
inference::rms_norm_into(&x, &m.final_norm, self.rms_eps, self.norm_style, &mut self.ws.n1);
let mut lg = self.lm_head_forward(&self.ws.n1);
let draft = sampler::argmax(&lg);
attention::recycle_buf(&mut lg);
draft
}
pub fn measure_pair_fusion(&mut self, iters: usize) -> (f64, f64) {
let emb1 = self.embed_single(1);
let emb2 = self.embed_single(2);
let pos = self.kv_cache.seq_len();
let t0 = std::time::Instant::now();
for _ in 0..iters {
let _ = self.forward_layers(&emb1, pos, None);
let _ = self.forward_layers(&emb2, pos + 1, None);
for l in &mut self.kv_cache.layers {
l.truncate_last(2);
}
}
let singles_ms = t0.elapsed().as_secs_f64() * 1000.0 / iters as f64;
let t1 = std::time::Instant::now();
for _ in 0..iters {
let _ = self.forward_pair(&emb1, &emb2, pos);
for l in &mut self.kv_cache.layers {
l.truncate_last(2);
}
}
let pair_ms = t1.elapsed().as_secs_f64() * 1000.0 / iters as f64;
(singles_ms, pair_ms)
}
fn forward_pair(&mut self, emb1: &[f32], emb2: &[f32], position: usize) -> (Vec<f32>, Vec<f32>) {
let mut h1 = emb1.to_vec();
let mut h2 = emb2.to_vec();
let (nh, nkv, hd, hs, rd, eps) = (
self.num_heads,
self.num_kv_heads,
self.head_dim,
self.hidden_size,
self.rotary_dim,
self.rms_eps,
);
let inv_freq = self.inv_freq.clone();
let pool = self.pool.clone();
for li in 0..self.num_layers {
let lw = &self.weights.layers[li];
inference::rms_norm_into(&h1, &lw.input_norm, self.rms_eps, self.norm_style, &mut self.ws.n1);
inference::rms_norm_into(&h2, &lw.input_norm, self.rms_eps, self.norm_style, &mut self.ws.n2);
let (a1, a2) = match &lw.attn {
AttnKind::Linear(w) => {
let cfg = self.vmf_cfg.expect("linear layer without vmf_cfg");
let layer = &mut self.kv_cache.layers[li];
let (state, scratch) = (&mut layer.linear_state, &mut layer.linear_scratch);
vmf_phase_pair(&self.ws.n1, &self.ws.n2, w, &cfg, state, scratch, self.pool.as_deref())
}
AttnKind::LinearGdn(w) => {
let cfg = self.gdn_cfg.expect("gdn layer without gdn_cfg");
let layer = &mut self.kv_cache.layers[li];
let (state, scratch) = (&mut layer.linear_state, &mut layer.linear_scratch);
gdn_pair(&self.ws.n1, &self.ws.n2, w, &cfg, state, scratch, self.pool.as_deref())
}
AttnKind::Full {
wq,
wk,
wv,
wo,
q_norm,
k_norm,
output_gate,
bias,
} => {
let cfg = QwenAttnCfg {
num_heads: nh,
num_kv_heads: nkv,
head_dim: hd,
hidden_size: hs,
position,
inv_freq: &inv_freq,
rotary_dim: rd,
q_norm: q_norm.as_deref(),
k_norm: k_norm.as_deref(),
output_gate: *output_gate,
bias: bias
.as_ref()
.map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
rms_eps: eps,
norm_style: self.norm_style,
pool: pool.as_deref(),
};
attention::qwen_attention_pair(
&self.ws.n1,
&self.ws.n2,
wq,
wk,
wv,
wo,
&mut self.kv_cache.layers[li],
&cfg,
)
}
};
for i in 0..self.hidden_size {
h1[i] += a1[i];
h2[i] += a2[i];
}
let (mut a1, mut a2) = (a1, a2);
attention::recycle_buf(&mut a1);
attention::recycle_buf(&mut a2);
let lw = &self.weights.layers[li];
inference::rms_norm_into(&h1, &lw.post_norm, self.rms_eps, self.norm_style, &mut self.ws.p1);
inference::rms_norm_into(&h2, &lw.post_norm, self.rms_eps, self.norm_style, &mut self.ws.p2);
let (f1, f2) =
ffn_forward_pair(&lw.ffn, &self.ws.p1, &self.ws.p2, self.pool.as_deref());
for i in 0..self.hidden_size {
h1[i] += f1[i];
h2[i] += f2[i];
}
let (mut f1, mut f2) = (f1, f2);
attention::recycle_buf(&mut f1);
attention::recycle_buf(&mut f2);
}
(h1, h2)
}
fn commit_linear_scratch(&mut self) {
for layer in &mut self.kv_cache.layers {
if !layer.linear_scratch.is_empty() {
std::mem::swap(&mut layer.linear_state, &mut layer.linear_scratch);
layer.linear_scratch.clear();
}
}
}
pub fn forward_ids(
&mut self,
ids: &[u32],
task_mask: Option<&TaskMask>,
) -> Result<Vec<f32>, String> {
if ids.is_empty() {
return Err("empty id sequence".to_string());
}
self.kv_cache.clear();
self.o1_begin();
let mut hidden = vec![0.0f32; self.hidden_size];
let mut pos = 0usize;
if task_mask.is_none() && prefill_batched() && ids.len() > 2 {
const CHUNK: usize = 48;
let hs = self.hidden_size;
while pos < ids.len() {
let end = (pos + CHUNK).min(ids.len());
let hb = self.prefill_batch(&ids[pos..end], pos);
hidden.copy_from_slice(&hb[(end - pos - 1) * hs..]);
pos = end;
}
}
if task_mask.is_none() {
while pos + 1 < ids.len() {
let e1 = self.embed_single(ids[pos]);
let e2 = self.embed_single(ids[pos + 1]);
let (_, h2) = self.forward_pair(&e1, &e2, pos);
self.commit_linear_scratch();
hidden = h2;
pos += 2;
}
}
while pos < ids.len() {
hidden = self.forward_layers(&self.embed_single(ids[pos]), pos, task_mask);
pos += 1;
}
self.o1_seal();
let normed = inference::rms_norm(
&hidden,
&self.weights.final_norm,
self.rms_eps,
self.norm_style,
);
Ok(self.lm_head_forward(&normed))
}
pub fn ppl_ids(&mut self, ids: &[u32]) -> f64 {
let (nll, cnt) = self.nll_ids_from(ids, 0);
(nll / cnt.max(1) as f64).exp()
}
pub fn nll_ids_from(&mut self, ids: &[u32], start: usize) -> (f64, usize) {
self.kv_cache.clear();
let mut nll = 0f64;
let mut cnt = 0usize;
if prefill_batched() {
const CHUNK: usize = 128;
const LM_SUB: usize = 32;
let n = ids.len().saturating_sub(1);
let hs = self.hidden_size;
let rows = self.weights.lm_head.rows();
let mut pos = 0usize;
while pos < n {
let end = (pos + CHUNK).min(n);
let bsz = end - pos;
let hb = self.prefill_batch(&ids[pos..end], pos);
let mut k0 = 0usize;
while k0 < bsz {
let k1 = (k0 + LM_SUB).min(bsz);
let sb = k1 - k0;
if pos + k1 <= start {
k0 = k1;
continue;
}
let mut normed = vec![0.0f32; sb * hs];
for k in 0..sb {
let r = inference::rms_norm(
&hb[(k0 + k) * hs..(k0 + k + 1) * hs],
&self.weights.final_norm,
self.rms_eps,
self.norm_style,
);
normed[k * hs..(k + 1) * hs].copy_from_slice(&r);
}
let mut logits = vec![0.0f32; sb * rows];
self.weights
.lm_head
.matmat(&normed, sb, &mut logits, self.pool.as_deref());
for k in 0..sb {
if pos + k0 + k < start {
continue;
}
let lg = &logits[k * rows..k * rows + self.vocab_size.min(rows)];
let target = ids[pos + k0 + k + 1] as usize;
let max = lg.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
let lse: f64 = lg
.iter()
.map(|&v| ((v - max) as f64).exp())
.sum::<f64>()
.ln()
+ max as f64;
nll += lse - lg[target] as f64;
cnt += 1;
}
k0 = k1;
}
pos = end;
}
self.kv_cache.clear();
return (nll, cnt);
}
for pos in 0..ids.len().saturating_sub(1) {
let hidden = self.forward_layers(&self.embed_single(ids[pos]), pos, None);
if pos < start {
continue;
}
let normed = inference::rms_norm(
&hidden,
&self.weights.final_norm,
self.rms_eps,
self.norm_style,
);
let logits = self.lm_head_forward(&normed);
let target = ids[pos + 1] as usize;
let max = logits.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
let lse: f64 = logits.iter().map(|&v| ((v - max) as f64).exp()).sum::<f64>().ln()
+ max as f64;
nll += lse - logits[target] as f64;
cnt += 1;
}
self.kv_cache.clear();
(nll, cnt)
}
pub fn nll_ids_o1(&mut self, ids: &[u32], prefill: usize) -> (f64, usize) {
self.kv_cache.clear();
self.o1_begin();
let n = ids.len().saturating_sub(1);
let p = prefill.min(n);
let mut pos = 0usize;
if prefill_batched() {
const CHUNK: usize = 128;
while pos < p {
let end = (pos + CHUNK).min(p);
let _ = self.prefill_batch(&ids[pos..end], pos);
pos = end;
}
} else {
while pos < p {
let _ = self.forward_layers(&self.embed_single(ids[pos]), pos, None);
pos += 1;
}
}
self.o1_seal();
let mut nll = 0f64;
let mut cnt = 0usize;
for pos in p..n {
let hidden = self.forward_layers(&self.embed_single(ids[pos]), pos, None);
let normed = inference::rms_norm(
&hidden,
&self.weights.final_norm,
self.rms_eps,
self.norm_style,
);
let logits = self.lm_head_forward(&normed);
let target = ids[pos + 1] as usize;
let max = logits.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
let lse: f64 = logits.iter().map(|&v| ((v - max) as f64).exp()).sum::<f64>().ln()
+ max as f64;
nll += lse - logits[target] as f64;
cnt += 1;
}
self.kv_cache.clear();
(nll, cnt)
}
pub fn calib_ids(&mut self, ids: &[u32], temps: &[f32]) -> (Vec<bool>, Vec<Vec<f32>>) {
self.kv_cache.clear();
let n = ids.len().saturating_sub(1);
let mut correct = Vec::with_capacity(n);
let mut pmax = Vec::with_capacity(n);
for pos in 0..n {
let emb = self.embed_single(ids[pos]);
let hidden = self.forward_layers(&emb, pos, None);
let normed = inference::rms_norm(
&hidden,
&self.weights.final_norm,
self.rms_eps,
self.norm_style,
);
let logits = self.lm_head_forward(&normed);
let target = ids[pos + 1] as usize;
let (mut amax, mut mval) = (0usize, f32::NEG_INFINITY);
for (i, &v) in logits.iter().enumerate() {
if v > mval {
mval = v;
amax = i;
}
}
correct.push(amax == target);
let row: Vec<f32> = temps
.iter()
.map(|&t| {
let tt = t.max(1e-3);
let s: f32 = logits.iter().map(|&v| ((v - mval) / tt).exp()).sum();
1.0 / s.max(1e-12) })
.collect();
pmax.push(row);
}
self.kv_cache.clear();
(correct, pmax)
}
pub fn ppl_ids_dynamic(&mut self, ids: &[u32]) -> (f64, usize) {
let mut router = match self.dyn_router.take() {
Some(r) => r,
None => return (self.ppl_ids(ids), 0),
};
router.reset();
self.dyn_phi_seen = 0;
let _ = self.set_active_skill(None);
self.kv_cache.clear();
let mut nll = 0f64;
let mut cnt = 0usize;
for pos in 0..ids.len().saturating_sub(1) {
let hidden = self.forward_layers(&self.embed_single(ids[pos]), pos, None);
let normed = inference::rms_norm(
&hidden,
&self.weights.final_norm,
self.rms_eps,
self.norm_style,
);
let logits = self.lm_head_forward(&normed);
let target = ids[pos + 1] as usize;
let max = logits.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
let lse: f64 = logits.iter().map(|&v| ((v - max) as f64).exp()).sum::<f64>().ln()
+ max as f64;
nll += lse - logits[target] as f64;
cnt += 1;
let phi = self.dyn_phi_ema.clone();
if let Some(new_active) = router.step(&phi, pos) {
let _ = self.set_active_skill(new_active);
}
}
let switches = router.switches.len();
let _ = self.set_active_skill(None);
self.dyn_router = Some(router);
self.kv_cache.clear();
((nll / cnt.max(1) as f64).exp(), switches)
}
pub fn probe_phi(&mut self, ids: &[u32], layer: usize) -> Vec<f32> {
self.kv_cache.clear();
let mut acc = vec![0f32; self.hidden_size];
for (pos, &id) in ids.iter().enumerate() {
let h = self.forward_layers_upto(&self.embed_single(id), pos, None, Some(layer));
for (a, v) in acc.iter_mut().zip(&h) {
*a += v;
}
}
let n = ids.len().max(1) as f32;
for a in acc.iter_mut() {
*a /= n;
}
self.kv_cache.clear();
acc
}
fn prefill_batch(&mut self, ids: &[u32], start_pos: usize) -> Vec<f32> {
let b = ids.len();
let hs = self.hidden_size;
let mut h: Vec<f32> = Vec::with_capacity(b * hs);
for &id in ids {
h.extend_from_slice(&self.embed_single(id));
}
let (nh, nkv, hd, rd, eps) = (
self.num_heads,
self.num_kv_heads,
self.head_dim,
self.rotary_dim,
self.rms_eps,
);
let inv_freq = self.inv_freq.clone();
let pool = self.pool.clone();
let norm_style = self.norm_style;
for li in 0..self.num_layers {
crate::gpu::set_layer(li as i64); let lw = &self.weights.layers[li];
match &lw.attn {
AttnKind::LinearGdn(w) => {
let cfg = self.gdn_cfg.expect("gdn layer without gdn_cfg");
let mut normed = vec![0.0f32; b * hs];
for bi in 0..b {
let r = inference::rms_norm(
&h[bi * hs..(bi + 1) * hs], &lw.input_norm, eps, norm_style);
normed[bi * hs..(bi + 1) * hs].copy_from_slice(&r);
}
let attn = crate::linear_core::gdn_forward_batch(
&normed, b, w, &cfg,
&mut self.kv_cache.layers[li].linear_state,
pool.as_deref(),
);
for (dst, &a) in h.iter_mut().zip(&attn) {
*dst += a;
}
}
AttnKind::Full {
wq, wk, wv, wo, q_norm, k_norm, output_gate, bias,
} => {
let mut normed = vec![0.0f32; b * hs];
for bi in 0..b {
inference::rms_norm_into(
&h[bi * hs..(bi + 1) * hs],
&lw.input_norm,
eps,
norm_style,
&mut normed[bi * hs..(bi + 1) * hs],
);
}
let cfg = QwenAttnCfg {
num_heads: nh,
num_kv_heads: nkv,
head_dim: hd,
hidden_size: hs,
position: start_pos,
inv_freq: &inv_freq,
rotary_dim: rd,
q_norm: q_norm.as_deref(),
k_norm: k_norm.as_deref(),
output_gate: *output_gate,
bias: bias.as_ref().map(|(a, b, c)| {
(a.as_slice(), b.as_slice(), c.as_slice())
}),
rms_eps: eps,
norm_style,
pool: pool.as_deref(),
};
let attn = attention::qwen_attention_batch(
&normed, b, wq, wk, wv, wo,
&mut self.kv_cache.layers[li], &cfg);
for (dst, &a) in h.iter_mut().zip(&attn) {
*dst += a;
}
}
AttnKind::Linear(w) => {
for bi in 0..b {
let normed = inference::rms_norm(
&h[bi * hs..(bi + 1) * hs], &lw.input_norm, eps, norm_style);
vmf_phase_forward(
&normed, w,
&self.vmf_cfg.expect("linear layer without vmf_cfg"),
&mut self.kv_cache.layers[li].linear_state,
pool.as_deref(),
)
.iter()
.enumerate()
.for_each(|(i, &a)| h[bi * hs + i] += a);
}
}
}
let lw = &self.weights.layers[li];
let mut post = vec![0.0f32; b * hs];
for bi in 0..b {
let r = inference::rms_norm(
&h[bi * hs..(bi + 1) * hs], &lw.post_norm, eps, norm_style);
post[bi * hs..(bi + 1) * hs].copy_from_slice(&r);
}
let ffn = match &lw.ffn {
FfnKind::Dense(d) => dense_ffn_batch(d, &post, b, pool.as_deref()),
FfnKind::Moe(m) => moe_ffn_batch(m, &post, b, hs, pool.as_deref()),
};
for (dst, &f) in h.iter_mut().zip(&ffn) {
*dst += f;
}
}
crate::gpu::set_layer(-1); h
}
fn embed_single(&self, id: u32) -> Vec<f32> {
let mut out = vec![0.0f32; self.hidden_size];
if (id as usize) < self.weights.embed_tokens.rows() {
self.weights.embed_tokens.row_f32(id as usize, &mut out);
}
out
}
fn forward_layers(
&mut self,
hidden: &[f32],
position: usize,
task_mask: Option<&TaskMask>,
) -> Vec<f32> {
self.forward_layers_upto(hidden, position, task_mask, None)
}
fn forward_layers_upto(
&mut self,
hidden: &[f32],
position: usize,
task_mask: Option<&TaskMask>,
upto: Option<usize>,
) -> Vec<f32> {
let mut h = hidden.to_vec();
let (nh, nkv, hd, hs, rd, eps) = (
self.num_heads,
self.num_kv_heads,
self.head_dim,
self.hidden_size,
self.rotary_dim,
self.rms_eps,
);
let inv_freq = self.inv_freq.clone();
let pool = self.pool.clone();
for li in 0..self.num_layers {
crate::gpu::set_layer(li as i64); if let Some(u) = upto {
if li > u {
break;
}
}
if let Some(mask) = task_mask {
if !mask.layer_alive(li) {
continue; }
}
let lw = &self.weights.layers[li];
inference::rms_norm_into(&h, &lw.input_norm, self.rms_eps, self.norm_style, &mut self.ws.n1);
let attn_out = match &lw.attn {
AttnKind::Linear(w) => {
let cfg = self.vmf_cfg.expect("linear layer without vmf_cfg");
vmf_phase_forward(
&self.ws.n1,
w,
&cfg,
&mut self.kv_cache.layers[li].linear_state,
self.pool.as_deref(),
)
}
AttnKind::LinearGdn(w) => {
let cfg = self.gdn_cfg.expect("gdn layer without gdn_cfg");
gdn_forward(
&self.ws.n1,
w,
&cfg,
&mut self.kv_cache.layers[li].linear_state,
self.pool.as_deref(),
)
}
AttnKind::Full {
wq,
wk,
wv,
wo,
q_norm,
k_norm,
output_gate,
bias,
} if self.kv_cache.layers[li].o1_sealed() => {
let cfg = QwenAttnCfg {
num_heads: nh,
num_kv_heads: nkv,
head_dim: hd,
hidden_size: hs,
position,
inv_freq: &inv_freq,
rotary_dim: rd,
q_norm: q_norm.as_deref(),
k_norm: k_norm.as_deref(),
output_gate: *output_gate,
bias: bias
.as_ref()
.map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
rms_eps: eps,
norm_style: self.norm_style,
pool: pool.as_deref(),
};
attention::qwen_attention_nystrom(
&self.ws.n1,
wq,
wk,
wv,
wo,
&mut self.kv_cache.layers[li],
&cfg,
)
}
AttnKind::Full {
wq,
wk,
wv,
wo,
q_norm,
k_norm,
output_gate,
bias,
} => {
let masked = task_mask
.map(|m| m.head_flags(li, self.num_heads).iter().any(|&a| !a))
.unwrap_or(false);
let f32_view = (wq.as_f32(), wk.as_f32(), wv.as_f32(), wo.as_f32());
match (masked, f32_view) {
(true, (Some(q), Some(k), Some(v), Some(o))) => {
let active_heads = task_mask.unwrap().head_flags(li, self.num_heads);
attention::multi_head_attention(
&self.ws.n1,
q,
k,
v,
o,
&mut self.kv_cache.layers[li],
self.num_heads,
self.num_kv_heads,
self.head_dim,
self.hidden_size,
position,
&active_heads,
&self.inv_freq,
)
}
(masked, _) => {
if masked {
tracing::warn!(
"layer {li}: head mask on quantized weights not \
supported yet — executing dense"
);
}
let cfg = QwenAttnCfg {
num_heads: nh,
num_kv_heads: nkv,
head_dim: hd,
hidden_size: hs,
position,
inv_freq: &inv_freq,
rotary_dim: rd,
q_norm: q_norm.as_deref(),
k_norm: k_norm.as_deref(),
output_gate: *output_gate,
bias: bias
.as_ref()
.map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
rms_eps: eps,
norm_style: self.norm_style,
pool: pool.as_deref(),
};
attention::qwen_attention(
&self.ws.n1,
wq,
wk,
wv,
wo,
&mut self.kv_cache.layers[li],
&cfg,
)
}
}
}
};
for (i, &a) in attn_out.iter().enumerate() {
h[i] += a;
}
let mut attn_out = attn_out;
attention::recycle_buf(&mut attn_out);
let lw = &self.weights.layers[li];
inference::rms_norm_into(&h, &lw.post_norm, self.rms_eps, self.norm_style, &mut self.ws.p1);
let post_normed = &self.ws.p1;
let ffn_masked = task_mask
.map(|m| m.ffn_active_count(li) < self.intermediate_size)
.unwrap_or(false);
let f32_ffn = match &lw.ffn {
FfnKind::Dense(d) => {
(d.gate_proj.as_f32(), d.up_proj.as_f32(), d.down_proj.as_f32())
}
FfnKind::Moe(_) => (None, None, None),
};
let ffn_out = match (ffn_masked, f32_ffn) {
(true, (Some(g), Some(u), Some(d))) => {
let active = task_mask.unwrap().ffn_active_indices(li);
inference::sparse_ffn_forward(
&post_normed,
g,
u,
d,
self.hidden_size,
self.intermediate_size,
&active,
self.pool.as_deref(),
)
}
(true, _) => match &lw.ffn {
FfnKind::Dense(d) if d.down_proj.sparse_col_ok() => {
let active = task_mask.unwrap().ffn_active_indices(li);
sparse_ffn_quant(
d,
&post_normed,
&active,
self.hidden_size,
self.pool.as_deref(),
)
}
FfnKind::Dense(d) => {
let active = task_mask.unwrap().ffn_active_indices(li);
let (gf, uf, df) = dequant_dense_f32(d);
inference::sparse_ffn_forward(
&post_normed,
&gf,
&uf,
&df,
self.hidden_size,
self.intermediate_size,
&active,
self.pool.as_deref(),
)
}
FfnKind::Moe(_) => {
ffn_forward(&lw.ffn, &post_normed, self.pool.as_deref())
}
},
(false, _) => ffn_forward(&lw.ffn, &post_normed, self.pool.as_deref()),
};
for (i, &f) in ffn_out.iter().enumerate() {
h[i] += f;
}
let mut ffn_out = ffn_out;
attention::recycle_buf(&mut ffn_out);
if self.dyn_phi_layer == Some(li) {
self.update_dyn_phi(&h);
}
}
crate::gpu::set_layer(-1);
h
}
fn update_dyn_phi(&mut self, h: &[f32]) {
const A: f32 = 0.2;
if self.dyn_phi_ema.len() != h.len() {
self.dyn_phi_ema = vec![0.0; h.len()];
self.dyn_phi_seen = 0;
}
if self.dyn_phi_seen == 0 {
self.dyn_phi_ema.copy_from_slice(h);
} else {
for (e, &v) in self.dyn_phi_ema.iter_mut().zip(h) {
*e = (1.0 - A) * *e + A * v;
}
}
self.dyn_phi_seen += 1;
}
pub fn dyn_phi(&self) -> &[f32] {
&self.dyn_phi_ema
}
pub fn set_dyn_phi_layer(&mut self, layer: Option<usize>) {
self.dyn_phi_layer = layer;
self.dyn_phi_ema.clear();
self.dyn_phi_seen = 0;
}
pub fn dynamic_skills(&self) -> Vec<(usize, String, usize)> {
let Some(model) = &self.model else { return Vec::new() };
model
.header
.skills
.iter()
.enumerate()
.filter_map(|(i, sk)| {
let ok = matches!(self.dyn_skill_layers.get(i), Some(Some(_)));
let sel = sk.selection.as_ref()?;
(ok).then(|| (i, sk.id.clone(), sel.phi_layer))
})
.collect()
}
pub fn active_skill(&self) -> Option<usize> {
self.dyn_active
}
pub fn enable_dynamic_routing(&mut self) -> usize {
use crate::swarm::{DynRouter, RoutableSkill};
let Some(model) = self.model.clone() else { return 0 };
if self.dyn_blend_loaded {
tracing::warn!("dynamic routing unavailable on a blend-loaded pipeline");
return 0;
}
if let Some(a) = self.dyn_active {
if !matches!(self.dyn_skill_layers.get(a), Some(Some(_))) {
tracing::warn!(
"loaded skill is not FFN-eligible — dynamic routing unavailable"
);
return 0;
}
}
let hidden = self.hidden_size;
let mut skills = Vec::new();
for (idx, id, _phi) in self.dynamic_skills() {
if let Some(sel) = model.header.skills[idx].selection.as_ref() {
if let Some(rs) = RoutableSkill::from_descriptor(idx, id, sel, hidden) {
skills.push(rs);
}
}
}
if skills.is_empty() {
return 0;
}
let phi = skills[0].phi_layer;
if skills.iter().any(|s| s.phi_layer != phi) {
tracing::warn!("routable skills disagree on phi_layer; using {phi}");
}
let n = skills.len();
self.set_dyn_phi_layer(Some(phi));
self.dyn_router = Some(DynRouter::new(skills));
n
}
pub fn route_switches(&self) -> Vec<(usize, Option<String>, Option<String>)> {
self.dyn_router
.as_ref()
.map(|r| r.switches.clone())
.unwrap_or_default()
}
fn lm_head_forward(&self, hidden: &[f32]) -> Vec<f32> {
let rows = self.weights.lm_head.rows();
let mut logits = attention::take_buf(rows.min(self.vocab_size));
self.weights
.lm_head
.matvec(hidden, &mut logits, self.pool.as_deref());
logits.resize(self.vocab_size, 0.0);
logits
}
pub fn prefill_next_logits(&mut self, ids: &[u32], task_mask: Option<&TaskMask>) -> Vec<f32> {
self.kv_cache.clear();
let mut hidden = vec![0.0f32; self.hidden_size];
for (pos, &id) in ids.iter().enumerate() {
let emb = self.embed_single(id);
hidden = self.forward_layers(&emb, pos, task_mask);
}
inference::rms_norm_into(
&hidden,
&self.weights.final_norm,
self.rms_eps,
self.norm_style,
&mut self.ws.n1,
);
self.lm_head_forward(&self.ws.n1)
}
}
pub fn create_test_pipeline(
hidden_size: usize,
intermediate_size: usize,
num_heads: usize,
num_kv_heads: usize,
head_dim: usize,
num_layers: usize,
vocab_size: usize,
) -> Pipeline {
let synth = |n: usize, salt: usize| -> Vec<f32> {
(0..n)
.map(|i| (((i * 31 + salt * 17 + 7) % 97) as f32 / 97.0 - 0.5) * 0.2)
.collect()
};
let qt = |rows: usize, cols: usize, salt: usize| -> QTensor {
QTensor::from_f32(synth(rows * cols, salt), rows, cols)
};
let layer_weights: Vec<LayerWeights> = (0..num_layers)
.map(|li| LayerWeights {
input_norm: vec![1.0; hidden_size],
post_norm: vec![1.0; hidden_size],
ffn: FfnKind::Dense(DenseFfn {
gate_proj: qt(intermediate_size, hidden_size, li * 10 + 5),
up_proj: qt(intermediate_size, hidden_size, li * 10 + 6),
down_proj: qt(hidden_size, intermediate_size, li * 10 + 7),
}),
attn: AttnKind::Full {
bias: None,
wq: qt(num_heads * head_dim, hidden_size, li * 10 + 1),
wk: qt(num_kv_heads * head_dim, hidden_size, li * 10 + 2),
wv: qt(num_kv_heads * head_dim, hidden_size, li * 10 + 3),
wo: qt(hidden_size, num_heads * head_dim, li * 10 + 4),
q_norm: None,
k_norm: None,
output_gate: false,
},
})
.collect();
Pipeline::new(
Tokenizer::byte_level(),
PipelineWeights {
embed_tokens: qt(vocab_size, hidden_size, 100),
layers: layer_weights,
lm_head: qt(vocab_size, hidden_size, 200),
final_norm: vec![1.0; hidden_size],
},
hidden_size,
intermediate_size,
num_heads,
num_kv_heads,
head_dim,
num_layers,
vocab_size,
1e-6,
10_000.0,
NormStyle::Qwen,
4096,
SamplerConfig {
seed: Some(42),
..Default::default()
},
)
}
fn dense_ffn_batch(d: &DenseFfn, xs: &[f32], b: usize, pool: Option<&Pool>) -> Vec<f32> {
let inter = d.gate_proj.rows();
let hidden = d.down_proj.rows();
let mut g = vec![0.0f32; b * inter];
d.gate_proj.matmat(xs, b, &mut g, pool);
let mut u = vec![0.0f32; b * inter];
d.up_proj.matmat(xs, b, &mut u, pool);
for i in 0..b * inter {
g[i] = inference::silu(g[i]) * u[i];
}
let mut out = vec![0.0f32; b * hidden];
d.down_proj.matmat(&g, b, &mut out, pool);
out
}
fn moe_ffn_batch(m: &MoeFfn, xs: &[f32], b: usize, hidden: usize, pool: Option<&Pool>) -> Vec<f32> {
let ne = m.experts.len();
let mut logits = vec![0.0f32; b * ne];
m.router.matmat(xs, b, &mut logits, pool);
let mut assign: Vec<Vec<(usize, f32)>> = vec![Vec::new(); ne];
{
let mut st = m.stats.borrow_mut();
if st.len() < ne {
st.resize(ne, 0);
}
for bi in 0..b {
let lg = &logits[bi * ne..(bi + 1) * ne];
let mx = lg.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
let mut p: Vec<f32> = lg.iter().map(|&l| (l - mx).exp()).collect();
let sum: f32 = p.iter().sum();
for v in &mut p {
*v /= sum;
}
let mut order: Vec<usize> = (0..ne).collect();
order.sort_unstable_by(|&x, &y| p[y].partial_cmp(&p[x]).unwrap().then(x.cmp(&y)));
order.truncate(m.top_k);
let wsum: f32 = if m.norm_topk_prob {
order.iter().map(|&e| p[e]).sum()
} else {
1.0
};
for &e in &order {
st[e] += 1;
assign[e].push((bi, p[e] / wsum));
}
}
}
let mut out = vec![0.0f32; b * hidden];
let cols = m.experts[0].gate_proj.cols();
let mut run_expert = |d: &DenseFfn, list: &[(usize, f32)]| {
let sb = list.len();
let mut sub = vec![0.0f32; sb * cols];
for (k, &(bi, _)) in list.iter().enumerate() {
sub[k * cols..(k + 1) * cols].copy_from_slice(&xs[bi * cols..(bi + 1) * cols]);
}
let eo = dense_ffn_batch(d, &sub, sb, pool);
for (k, &(bi, w)) in list.iter().enumerate() {
for i in 0..hidden {
out[bi * hidden + i] += w * eo[k * hidden + i];
}
}
};
for e in 0..ne {
if !assign[e].is_empty() {
run_expert(&m.experts[e], &assign[e]);
}
}
if let Some((se, gate)) = &m.shared {
let mut gl = vec![0.0f32; b];
gate.matmat(xs, b, &mut gl, pool);
let all: Vec<(usize, f32)> = (0..b)
.map(|bi| (bi, 1.0 / (1.0 + (-gl[bi]).exp())))
.collect();
run_expert(se, &all);
}
out
}
thread_local! {
static FFN_SCRATCH: std::cell::RefCell<[Vec<f32>; 4]> =
const { std::cell::RefCell::new([Vec::new(), Vec::new(), Vec::new(), Vec::new()]) };
}
fn dense_ffn(d: &DenseFfn, x: &[f32], pool: Option<&Pool>) -> Vec<f32> {
if crate::gpu::enabled_here() && d.gate_proj.rows() >= crate::gpu::min_rows() {
match crate::gpu::probe_arm(crate::gpu::OpClass::Ffn) {
crate::gpu::ProbeArm::Gpu => {
let t0 = std::time::Instant::now();
if let Some(out) = dense_ffn_gpu(d, x, pool) {
crate::gpu::probe_record(crate::gpu::OpClass::Ffn, true, t0.elapsed());
return out;
}
}
crate::gpu::ProbeArm::CpuTimed => {
let t0 = std::time::Instant::now();
let out = crate::gpu::cpu_scope(|| dense_ffn_cpu(d, x, pool));
crate::gpu::probe_record(crate::gpu::OpClass::Ffn, false, t0.elapsed());
return out;
}
crate::gpu::ProbeArm::Cpu => {
return crate::gpu::cpu_scope(|| dense_ffn_cpu(d, x, pool));
}
}
}
dense_ffn_cpu(d, x, pool)
}
fn dense_ffn_cpu(d: &DenseFfn, x: &[f32], pool: Option<&Pool>) -> Vec<f32> {
let inter = d.gate_proj.rows();
FFN_SCRATCH.with(|s| {
let mut s = s.borrow_mut();
let [g, u, ..] = &mut *s;
g.resize(inter, 0.0);
u.resize(inter, 0.0);
QTensor::matvec_many([&d.gate_proj, &d.up_proj], x, [g, u], pool);
for i in 0..inter {
g[i] = inference::silu(g[i]) * u[i];
}
let mut out = attention::take_buf(d.down_proj.rows());
d.down_proj.matvec(g, &mut out, pool);
out
})
}
fn dense_ffn_gpu(d: &DenseFfn, x: &[f32], _pool: Option<&Pool>) -> Option<Vec<f32>> {
if d.gate_proj.rows() < crate::gpu::min_rows() {
return None;
}
let mut jobs: Vec<crate::gpu::MoeJob> = Vec::with_capacity(1);
let mut model_ref = None;
moe_push_job(d, x, 1.0, &mut jobs, &mut model_ref)?;
let model = model_ref?;
let hidden = jobs[0].down.1;
let mut out = attention::take_buf(hidden);
if crate::gpu::moe_block(&model, &jobs, &mut out) {
Some(out)
} else {
let mut out = out;
attention::recycle_buf(&mut out);
None
}
}
#[allow(clippy::type_complexity)]
fn moe_parts(
t: &QTensor,
) -> Option<(&std::sync::Arc<cortiq_core::CmfModel>, usize, usize, usize, &[f32], &[f32])> {
match t {
QTensor::Mapped {
model,
idx,
dtype: dt @ (cortiq_core::TensorDtype::Q8_2f | cortiq_core::TensorDtype::Q8Row),
rows,
cols,
row_scale,
col_field,
..
} if (*dt == cortiq_core::TensorDtype::Q8Row) || !col_field.is_empty() => {
Some((model, *idx, *rows, *cols, row_scale, col_field))
}
_ => None,
}
}
fn moe_push_job<'a>(
d: &'a DenseFfn,
x: &[f32],
w: f32,
jobs: &mut Vec<crate::gpu::MoeJob<'a>>,
model_ref: &mut Option<std::sync::Arc<cortiq_core::CmfModel>>,
) -> Option<()> {
use crate::qtensor::prescale;
let (gm, gi, gr, gc, grs, gcf) = moe_parts(&d.gate_proj)?;
let (_, ui, ur, uc, urs, ucf) = moe_parts(&d.up_proj)?;
let (_, di, dr, dc, drs, dcf) = moe_parts(&d.down_proj)?;
model_ref.get_or_insert_with(|| gm.clone());
let gdt = if gcf.is_empty() { cortiq_core::TensorDtype::Q8Row } else { cortiq_core::TensorDtype::Q8_2f };
let udt = if ucf.is_empty() { cortiq_core::TensorDtype::Q8Row } else { cortiq_core::TensorDtype::Q8_2f };
jobs.push(crate::gpu::MoeJob {
gate: (gi, gr, gc, grs),
up: (ui, ur, uc, urs),
down: (di, dr, dc, drs),
xs_gate: prescale(x, gcf, gdt).into_owned(),
xs_up: prescale(x, ucf, udt).into_owned(),
down_col: dcf,
w,
});
Some(())
}
fn sparse_ffn_quant(
d: &DenseFfn,
x: &[f32],
active: &[u16],
hidden: usize,
pool: Option<&Pool>,
) -> Vec<f32> {
let n = active.len();
let inter = d.gate_proj.rows();
let mut act = vec![0.0f32; n];
let need_scratch = !(d.gate_proj.sparse_col_ok() && d.up_proj.sparse_col_ok());
let compute = |ai: usize| -> f32 {
let idx = active[ai] as usize;
if idx >= inter {
return 0.0; }
let mut s = if need_scratch { vec![0.0f32; hidden] } else { Vec::new() };
let gate = d.gate_proj.row_dot(idx, x, &mut s);
let up = d.up_proj.row_dot(idx, x, &mut s);
inference::silu(gate) * up
};
match pool {
Some(p) if n >= 256 => {
let ptr = SendMut(act.as_mut_ptr());
p.run(&|widx, nw| {
let chunk = n.div_ceil(nw);
let (s, e) = (widx * chunk, ((widx + 1) * chunk).min(n));
for ai in s..e {
unsafe { *ptr.at(ai) = compute(ai) };
}
});
}
_ => {
for (ai, a) in act.iter_mut().enumerate() {
*a = compute(ai);
}
}
}
let mut out = vec![0.0f32; hidden];
for (ai, &idx) in active.iter().enumerate() {
let w = act[ai];
if w.abs() >= 1e-12 && (idx as usize) < inter {
d.down_proj.add_col_scaled(idx as usize, w, &mut out);
}
}
out
}
#[doc(hidden)]
pub fn sparse_ffn_quant_for_test(
d: &DenseFfn,
x: &[f32],
active: &[u16],
hidden: usize,
) -> Vec<f32> {
sparse_ffn_quant(d, x, active, hidden, None)
}
fn dequant_dense_f32(d: &DenseFfn) -> (Vec<f32>, Vec<f32>, Vec<f32>) {
let deq = |t: &QTensor| -> Vec<f32> {
let (rows, cols) = (t.rows(), t.cols());
let mut out = vec![0.0f32; rows * cols];
for r in 0..rows {
t.row_f32(r, &mut out[r * cols..(r + 1) * cols]);
}
out
};
(deq(&d.gate_proj), deq(&d.up_proj), deq(&d.down_proj))
}
struct SendMut(*mut f32);
unsafe impl Send for SendMut {}
unsafe impl Sync for SendMut {}
impl SendMut {
#[inline]
#[allow(clippy::mut_from_ref)]
unsafe fn at(&self, i: usize) -> &mut f32 {
unsafe { &mut *self.0.add(i) }
}
}
fn moe_ffn(m: &MoeFfn, x: &[f32], pool: Option<&Pool>) -> Vec<f32> {
let ne = m.experts.len();
let mut logits = vec![0.0f32; ne];
m.router.matvec(x, &mut logits, pool);
let mx = logits.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
let mut p: Vec<f32> = logits.iter().map(|&l| (l - mx).exp()).collect();
let s: f32 = p.iter().sum();
for v in &mut p {
*v /= s;
}
let mut idx: Vec<usize> = (0..ne).collect();
idx.sort_unstable_by(|&a, &b| p[b].partial_cmp(&p[a]).unwrap().then(a.cmp(&b)));
idx.truncate(m.top_k);
let wsum: f32 = if m.norm_topk_prob {
idx.iter().map(|&e| p[e]).sum()
} else {
1.0
};
{
let mut st = m.stats.borrow_mut();
if st.len() < ne {
st.resize(ne, 0);
}
for &e in &idx {
st[e] += 1;
}
}
if crate::gpu::enabled_here() {
match crate::gpu::probe_arm(crate::gpu::OpClass::Ffn) {
crate::gpu::ProbeArm::Gpu => {
let t0 = std::time::Instant::now();
if let Some(out) = moe_ffn_gpu(m, x, &idx, &p, wsum, pool) {
crate::gpu::probe_record(crate::gpu::OpClass::Ffn, true, t0.elapsed());
return out;
}
}
crate::gpu::ProbeArm::CpuTimed => {
let t0 = std::time::Instant::now();
let out = crate::gpu::cpu_scope(|| moe_ffn_cpu(m, x, &idx, &p, wsum, pool));
crate::gpu::probe_record(crate::gpu::OpClass::Ffn, false, t0.elapsed());
return out;
}
crate::gpu::ProbeArm::Cpu => {
return crate::gpu::cpu_scope(|| moe_ffn_cpu(m, x, &idx, &p, wsum, pool));
}
}
}
moe_ffn_cpu(m, x, &idx, &p, wsum, pool)
}
fn moe_ffn_cpu(
m: &MoeFfn,
x: &[f32],
idx: &[usize],
p: &[f32],
wsum: f32,
pool: Option<&Pool>,
) -> Vec<f32> {
let mut out = attention::take_buf(x.len());
for &e in idx {
let mut eo = dense_ffn(&m.experts[e], x, pool);
let w = p[e] / wsum;
for i in 0..out.len() {
out[i] += w * eo[i];
}
attention::recycle_buf(&mut eo);
}
if let Some((se, gate)) = &m.shared {
let mut so = dense_ffn(se, x, pool);
let mut gl = vec![0.0f32; 1];
gate.matvec(x, &mut gl, pool);
let g = 1.0 / (1.0 + (-gl[0]).exp());
for i in 0..out.len() {
out[i] += g * so[i];
}
attention::recycle_buf(&mut so);
}
out
}
fn moe_ffn_gpu(
m: &MoeFfn,
x: &[f32],
idx: &[usize],
p: &[f32],
wsum: f32,
pool: Option<&Pool>,
) -> Option<Vec<f32>> {
use crate::gpu::MoeJob;
let mut jobs: Vec<MoeJob> = Vec::with_capacity(idx.len() + 1);
let mut model_ref = None;
for &e in idx {
moe_push_job(&m.experts[e], x, p[e] / wsum, &mut jobs, &mut model_ref)?;
}
if let Some((se, gate)) = &m.shared {
let mut gl = vec![0.0f32; 1];
gate.matvec(x, &mut gl, pool);
let g = 1.0 / (1.0 + (-gl[0]).exp());
moe_push_job(se, x, g, &mut jobs, &mut model_ref)?;
}
let model = model_ref?;
let hidden = jobs[0].down.1;
let mut out = vec![0.0f32; hidden];
crate::gpu::moe_block(&model, &jobs, &mut out).then_some(out)
}
fn ffn_forward(ffn: &FfnKind, x: &[f32], pool: Option<&Pool>) -> Vec<f32> {
match ffn {
FfnKind::Dense(d) => dense_ffn(d, x, pool),
FfnKind::Moe(m) => moe_ffn(m, x, pool),
}
}
fn ffn_forward_pair(
ffn: &FfnKind,
x1: &[f32],
x2: &[f32],
pool: Option<&Pool>,
) -> (Vec<f32>, Vec<f32>) {
let d = match ffn {
FfnKind::Dense(d) => d,
FfnKind::Moe(m) => return (moe_ffn(m, x1, pool), moe_ffn(m, x2, pool)),
};
let inter = d.gate_proj.rows();
FFN_SCRATCH.with(|s| {
let mut s = s.borrow_mut();
let [g1, g2, u1, u2] = &mut *s;
g1.resize(inter, 0.0);
g2.resize(inter, 0.0);
u1.resize(inter, 0.0);
u2.resize(inter, 0.0);
QTensor::matvec2_many(
[&d.gate_proj, &d.up_proj],
x1,
x2,
[g1.as_mut_slice(), u1.as_mut_slice()],
[g2.as_mut_slice(), u2.as_mut_slice()],
pool,
);
for i in 0..inter {
g1[i] = inference::silu(g1[i]) * u1[i];
g2[i] = inference::silu(g2[i]) * u2[i];
}
let mut o1 = attention::take_buf(d.down_proj.rows());
let mut o2 = attention::take_buf(d.down_proj.rows());
d.down_proj.matvec2(g1, g2, &mut o1, &mut o2, pool);
(o1, o2)
})
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn sparse_ffn_quant_equals_dense_with_inactive_zeroed() {
let (hidden, inter) = (16usize, 40usize);
let synth = |n: usize, salt: usize| -> Vec<f32> {
(0..n)
.map(|i| (((i * 37 + salt * 11 + 3) % 101) as f32 / 101.0 - 0.5) * 0.4)
.collect()
};
let d = DenseFfn {
gate_proj: QTensor::from_f32(synth(inter * hidden, 1), inter, hidden),
up_proj: QTensor::from_f32(synth(inter * hidden, 2), inter, hidden),
down_proj: QTensor::from_f32(synth(hidden * inter, 3), hidden, inter),
};
let x = synth(hidden, 9);
let active: Vec<u16> = (0..inter as u16).filter(|i| i % 3 == 0).collect();
let sparse = sparse_ffn_quant(&d, &x, &active, hidden, None);
let mut g = vec![0.0f32; inter];
d.gate_proj.matvec(&x, &mut g, None);
let mut u = vec![0.0f32; inter];
d.up_proj.matvec(&x, &mut u, None);
let act_set: std::collections::HashSet<u16> = active.iter().copied().collect();
for i in 0..inter {
g[i] = if act_set.contains(&(i as u16)) {
inference::silu(g[i]) * u[i]
} else {
0.0
};
}
let mut reference = vec![0.0f32; hidden];
d.down_proj.matvec(&g, &mut reference, None);
let max_d = sparse
.iter()
.zip(&reference)
.map(|(a, b)| (a - b).abs())
.fold(0.0f32, f32::max);
assert!(max_d < 1e-5, "sparse != dense-zeroed: max|Δ| = {max_d}");
}
fn attach_test_mtp(p: &mut Pipeline) {
let (h, inter, heads, kv, hd) = (
p.hidden_size,
p.intermediate_size,
p.num_heads,
p.num_kv_heads,
p.head_dim,
);
let synth = |n: usize, salt: usize| -> Vec<f32> {
(0..n)
.map(|i| (((i * 29 + salt * 23 + 5) % 101) as f32 / 101.0 - 0.5) * 0.2)
.collect()
};
let qt = |rows: usize, cols: usize, salt: usize| -> QTensor {
QTensor::from_f32(synth(rows * cols, salt), rows, cols)
};
p.mtp = Some(MtpModule {
enorm: vec![1.0; h],
hnorm: vec![1.0; h],
eh_proj: qt(h, 2 * h, 301),
layer: LayerWeights {
input_norm: vec![1.0; h],
post_norm: vec![1.0; h],
ffn: FfnKind::Dense(DenseFfn {
gate_proj: qt(inter, h, 315),
up_proj: qt(inter, h, 316),
down_proj: qt(h, inter, 317),
}),
attn: AttnKind::Full {
bias: None,
wq: qt(heads * hd, h, 311),
wk: qt(kv * hd, h, 312),
wv: qt(kv * hd, h, 313),
wo: qt(h, heads * hd, 314),
q_norm: None,
k_norm: None,
output_gate: false,
},
},
final_norm: vec![1.0; h],
kv: crate::kv_cache::LayerKvCache::new(kv, hd),
});
}
#[test]
fn speculative_equals_vanilla_greedy() {
let run = |spec: bool| {
let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
p.sampler_config.temperature = 0.0;
attach_test_mtp(&mut p);
p.speculative = spec;
let r = p.generate("abcdef", 12, None, None).unwrap();
(r.token_ids, r.mtp_drafted, r.mtp_accepted)
};
let (vanilla, d0, _) = run(false);
let (spec, d1, a1) = run(true);
assert_eq!(d0, 0, "vanilla path must not draft");
assert!(d1 > 0, "speculative path must draft");
assert_eq!(
vanilla, spec,
"speculative must reproduce the exact greedy sequence (accepted {a1}/{d1})"
);
}
#[test]
fn speculative_accepts_constant_oracle() {
let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
p.sampler_config.temperature = 0.0;
p.sampler_config.repetition_penalty = 1.0;
p.weights.lm_head = QTensor::from_f32(vec![0.01; 64 * 8], 64, 8);
attach_test_mtp(&mut p);
p.speculative = true;
let r = p.generate("abcd", 10, None, None).unwrap();
assert!(r.mtp_drafted > 0);
assert_eq!(
r.mtp_accepted, r.mtp_drafted,
"constant logits → every draft accepted"
);
assert!(r.token_ids.windows(2).all(|w| w[0] == w[1]));
}
#[test]
fn empty_prompt_is_an_error_not_a_panic() {
let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 260);
let r = p.generate("", 4, None, None);
assert!(r.is_err(), "empty prompt must be a clean error");
}
#[test]
fn every_token_enters_kv_exactly_once() {
let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
p.sampler_config.temperature = 0.0;
let r = p.generate("abc", 2, None, None).unwrap();
assert_eq!(r.prompt_tokens, 3);
assert_eq!(
p.kv_cache.seq_len(),
3 + r.tokens_generated - 1,
"each token must be cached exactly once (v1 cached the last prompt token twice)"
);
}
#[test]
fn generation_is_reproducible_with_seed() {
let run = || {
let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
p.generate("hello", 8, None, None).unwrap().token_ids
};
assert_eq!(run(), run());
}
#[test]
fn eviction_bounds_the_cache() {
let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 260);
p.kv_cache.max_seq_len = 6;
p.sampler_config.temperature = 0.0;
let _ = p.generate("abcd", 12, None, None).unwrap();
assert!(
p.kv_cache.seq_len() <= 6 + 1,
"cache must stay bounded by max_seq_len (got {})",
p.kv_cache.seq_len()
);
}
#[test]
fn confidence_matches_tokens_and_is_a_probability() {
let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
p.sampler_config.temperature = 0.0;
p.sampler_config.repetition_penalty = 1.0;
let r = p.generate("abcd", 10, None, None).unwrap();
assert_eq!(
r.token_confidence.len(),
r.token_ids.len(),
"one confidence per emitted token"
);
for &c in &r.token_confidence {
assert!((0.0..=1.0).contains(&c), "confidence out of [0,1]: {c}");
}
let logits = [1.0f32, 3.0, 0.5, 3.0];
let p0 = top1_prob_t(&logits, 1, 1.0);
let p1 = top1_prob_t(&logits, 3, 1.0);
assert!((p0 - p1).abs() < 1e-6, "equal logits → equal prob");
assert!(p0 > 0.0 && p0 < 1.0);
let sharp = top1_prob_t(&logits, 1, 1.0);
let soft = top1_prob_t(&logits, 1, 2.0);
assert!(soft < sharp, "higher temperature lowers peak confidence");
}
#[test]
fn trace_is_opt_in_and_parallels_the_output() {
let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
p.sampler_config.temperature = 0.0;
p.sampler_config.repetition_penalty = 1.0;
let r = p.generate("abcd", 10, None, None).unwrap();
assert!(r.traces.is_empty(), "trace must be empty unless enabled");
let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
p.sampler_config.temperature = 0.0;
p.sampler_config.repetition_penalty = 1.0;
p.set_trace(true);
let r = p.generate("abcd", 10, None, None).unwrap();
assert_eq!(r.traces.len(), r.token_ids.len(), "one trace row per token");
for (i, tr) in r.traces.iter().enumerate() {
assert_eq!(tr.t, i, "trace index is sequential");
assert_eq!(tr.token_id, r.token_ids[i], "trace token_id matches output");
assert_eq!(
tr.confidence, r.token_confidence[i],
"trace confidence matches the confidence channel"
);
assert!(tr.active_skill.is_none() && tr.recon.is_none() && !tr.switched);
}
}
#[test]
fn explain_prefill_logits_match_greedy_first_token() {
let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
p.sampler_config.temperature = 0.0;
p.sampler_config.repetition_penalty = 1.0;
let ids = p.tokenizer.encode("abcd");
let logits = p.prefill_next_logits(&ids, None);
let argmax = logits
.iter()
.enumerate()
.max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
.unwrap()
.0 as u32;
let r = p.generate("abcd", 1, None, None).unwrap();
assert_eq!(argmax, r.token_ids[0], "explain preview must match greedy emit");
}
}