use crate::attention::{self, QwenAttnCfg};
use crate::inference;
use crate::kv_cache::KvCache;
use crate::linear_core::{
GdnCfg, GdnWeights, ShortConvCfg, ShortConvWeights, VmfPhaseCfg, VmfPhaseWeights, gdn_forward,
gdn_pair, short_conv_forward, short_conv_forward_batch, short_conv_pair, vmf_phase_forward,
vmf_phase_pair,
};
use crate::pool::Pool;
use crate::qtensor::QTensor;
use crate::sampler::{self, SamplerConfig, SamplerScratch, SplitMix64};
use crate::tokenizer::Tokenizer;
use cortiq_core::mask::TaskMask;
use cortiq_core::types::NormStyle;
pub static GLOBAL_USE_GPU: std::sync::atomic::AtomicBool =
std::sync::atomic::AtomicBool::new(false);
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 {
gpu_plan: Option<std::sync::Arc<Vec<(usize, usize, usize)>>>,
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 physical_layers: usize,
pub loop_final_norm: bool,
pub vocab_size: usize,
pub rms_eps: f64,
pub rope_base: f32,
pub norm_style: NormStyle,
pub rotary_dim: usize,
pub attention_heads_per_layer: Option<Vec<usize>>,
pub vmf_cfg: Option<VmfPhaseCfg>,
pub gdn_cfg: Option<GdnCfg>,
pub logit_multiplier: Option<f32>,
pub cancel: std::sync::Arc<std::sync::atomic::AtomicBool>,
graph_failed: std::sync::atomic::AtomicBool,
pub kv_history: Vec<u32>,
pub kda_cfg: Option<crate::linear_core::KdaCfg>,
pub g3n: Option<Box<(crate::g3n::G3nGlobals, Vec<crate::g3n::G3nLayer>)>>,
pub dsv4: Option<
Box<(
crate::dsv4::Dsv4Globals,
Vec<crate::dsv4::Dsv4Layer>,
crate::dsv4::Dsv4Cfg,
crate::dsv4::Dsv4State,
)>,
>,
pub dsv41: Option<
Box<(
crate::dsv41::Dsv41Globals,
Vec<crate::dsv41::Dsv41Layer>,
crate::dsv41::Dsv41Cfg,
crate::dsv41::Dsv41State,
)>,
>,
pub dsv41_vision: Option<crate::dsv41_vision::VisionModel>,
dsv41_prefill: Option<(Vec<Option<Vec<f32>>>, Vec<bool>)>,
pub qwen4_exp: Option<
Box<(
crate::qwen4_exp::Globals,
Vec<crate::qwen4_exp::Layer>,
crate::qwen4_exp::Cfg,
crate::qwen4_exp::State,
)>,
>,
pub dsv4_mtp: Vec<crate::dsv4::Dsv4Mtp>,
pub dspark: Option<crate::dsv4::DsparkState>,
pub dspark_pending: Vec<(usize, Vec<u32>, bool, usize)>,
pub dspark_hist: Vec<usize>,
pub dspark_real: Vec<u32>,
pub dspark_trunk_picks: Vec<Vec<(usize, Vec<usize>)>>,
pub dspark_exp: Vec<(usize, usize, usize, usize)>,
pub dspark_draft_ns: u128,
pub short_conv_cfg: Option<ShortConvCfg>,
pub mtp: Option<MtpModule>,
pub speculative: bool,
pub ignore_eos: bool,
pub draft_full_streak: u32,
pub spec_k_adapt: Option<usize>,
pub spec_acc_ewma: f32,
rng: SplitMix64,
sampler_scratch: SamplerScratch,
spec_forced: Option<u32>,
spec_q: Vec<Vec<f32>>,
spec_p: Vec<f32>,
spec_res: Vec<f32>,
spec_qs: Vec<sampler::Sparse>,
spec_ps: sampler::Sparse,
spec_ress: sampler::Sparse,
mtp_graph_mode: Option<bool>,
#[cfg(target_os = "macos")]
metal_verify: Option<MetalVerifyPending>,
pub(crate) 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_epoch: u64,
o1_flags: Vec<bool>,
trace: bool,
calib_temp: f32,
#[cfg_attr(not(target_os = "macos"), allow(dead_code))]
graph_kv_id: u64,
#[cfg_attr(not(target_os = "macos"), allow(dead_code))]
graph_want_logits: bool,
#[cfg_attr(not(target_os = "macos"), allow(dead_code))]
graph_head_required: bool,
graph_logits: Option<Vec<f32>>,
pub embed_multiplier: f32,
pub attn_scale: f32,
pub swa: Option<(usize, usize)>,
pub sliding_layers: Option<Vec<bool>>,
pub inv_freq_local: Option<std::sync::Arc<Vec<f32>>>,
pub rotary_dim_local: Option<usize>,
pub rope_scale: f32,
pub rope_scale_local: f32,
pub global_attn: Option<(usize, usize)>,
pub inv_freq_global: Option<std::sync::Arc<Vec<f32>>>,
pub attn_v_norm: bool,
pub qk_norm_after_rope: bool,
pub final_softcap: Option<f32>,
pub head_clusters: Option<std::sync::Arc<Vec<f32>>>,
pub attn_softcap: f32,
confidence_on: bool,
#[cfg(test)]
nll_test_fail_at: Option<usize>,
#[cfg(test)]
nll_test_force_serial: bool,
}
#[cfg(target_os = "macos")]
impl Drop for Pipeline {
fn drop(&mut self) {
let _ = crate::gpu_metal::wait_replay();
crate::gpu::kv_mirror_drop(self.graph_kv_id);
}
}
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 attn_out_norm: Option<Vec<f32>>,
pub layer_scale: Option<f32>,
pub ffn_out_norm: Option<Vec<f32>>,
pub ffn: FfnKind,
pub attn: AttnKind,
}
#[derive(Clone, Copy, PartialEq, Debug, Default)]
pub enum Act {
#[default]
Silu,
GeluTanh,
Situ {
beta: f32,
linear_beta: f32,
},
}
impl Act {
pub fn from_arch(name: &str) -> Self {
if name == "gelu_tanh" {
Self::GeluTanh
} else {
Self::Silu
}
}
pub fn from_arch_full(arch: &cortiq_core::ModelArch) -> Self {
match arch.hidden_act.as_str() {
"situ" => Self::Situ {
beta: arch.activation_situ_beta.unwrap_or(1.0) as f32,
linear_beta: arch.activation_situ_linear_beta.unwrap_or(0.0) as f32,
},
other => Self::from_arch(other),
}
}
#[inline]
pub fn apply(self, x: f32) -> f32 {
match self {
Self::Silu => inference::silu(x),
Self::GeluTanh => inference::gelu_tanh(x),
Self::Situ { beta, .. } => beta * (x / beta).tanh() * (1.0 / (1.0 + (-x).exp())),
}
}
#[inline]
pub fn combine(self, g: f32, u: f32) -> f32 {
match self {
Self::Situ { linear_beta, .. } if linear_beta > 0.0 => {
self.apply(g) * (linear_beta * (u / linear_beta).tanh())
}
_ => self.apply(g) * u,
}
}
}
pub struct DenseFfn {
pub gate_proj: QTensor,
pub up_proj: QTensor,
pub down_proj: QTensor,
pub act: Act,
pub down_t: Option<QTensor>,
pub segs: Vec<FfnSeg>,
}
pub struct FfnSeg {
pub gate: QTensor,
pub up: QTensor,
pub down: QTensor,
pub start: usize,
pub width: usize,
}
pub enum FfnKind {
Dense(DenseFfn),
Moe(MoeFfn),
DenseMoe(Box<DenseMoeFfn>),
}
pub struct DenseMoeFfn {
pub dense: DenseFfn,
pub moe: MoeFfn,
pub post_norm_1: Vec<f32>,
pub pre_norm_2: Vec<f32>,
pub post_norm_2: Vec<f32>,
}
pub struct MoeFfn {
pub router: QTensor,
pub experts: Vec<DenseFfn>,
pub top_k: usize,
pub norm_topk_prob: bool,
pub router_sigmoid: bool,
pub expert_bias: Option<Vec<f32>>,
pub routed_scaling: f32,
pub route_tau: Option<f32>,
pub shared: Option<(DenseFfn, Option<QTensor>)>,
pub stats: std::cell::RefCell<Vec<u64>>,
pub act_sq: std::cell::RefCell<Vec<f64>>,
pub act_rows: std::cell::RefCell<Vec<f32>>,
pub mask: Option<Vec<bool>>,
pub per_expert_scale: Option<Vec<f32>>,
pub router_input_norm: bool,
pub resonance: Option<Resonance>,
}
pub struct Resonance {
pub mu: Vec<f32>,
pub u: Vec<f32>,
pub k: usize,
pub bias: Vec<f32>,
}
impl Resonance {
pub fn scores(&self, x: &[f32], out: &mut [f32]) {
let h = x.len();
let ne = out.len();
for e in 0..ne {
let mu = &self.mu[e * h..(e + 1) * h];
let mut d2 = 0.0f32;
for j in 0..h {
let d = x[j] - mu[j];
d2 += d * d;
}
let mut proj = 0.0f32;
for i in 0..self.k {
let u = &self.u[(e * self.k + i) * h..(e * self.k + i + 1) * h];
let mut p = 0.0f32;
for j in 0..h {
p += (x[j] - mu[j]) * u[j];
}
proj += p * p;
}
out[e] = self.bias.get(e).copied().unwrap_or(0.0) - (d2 - proj);
}
}
}
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,
softplus_gate: Option<(QTensor, bool)>,
bias: Option<(Vec<f32>, Vec<f32>, Vec<f32>)>,
},
Linear(VmfPhaseWeights),
LinearGdn(GdnWeights),
ShortConv(ShortConvWeights),
Mla(Box<MlaWeights>),
Kda(Box<crate::linear_core::KdaWeights>),
}
pub struct MlaWeights {
pub q_proj: QTensor,
pub q_a: Option<QTensor>,
pub q_a_norm: Option<Vec<f32>>,
pub kv_a: QTensor,
pub kv_a_norm: Vec<f32>,
pub kv_b: QTensor,
pub o_proj: QTensor,
pub nh: usize,
pub qk_rope: usize,
pub qk_nope: usize,
pub v_dim: usize,
pub lora: usize,
pub scale: f32,
pub nope: bool,
}
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,
}
#[cfg(target_os = "macos")]
enum MetalRowsItem<'a> {
Gdn {
run: Vec<crate::gpu_metal::GdnGpuLayer<'a>>,
first: usize,
},
Attn {
l: crate::gpu_metal::AttnGpuLayer<'a>,
li: usize,
q_norm: Option<&'a [f32]>,
k_norm: Option<&'a [f32]>,
output_gate: bool,
},
}
#[cfg(target_os = "macos")]
struct MetalVerifyPending {
graph: crate::gpu_metal::VerifyGraph,
gdn_layers: Vec<usize>,
attn_layers: Vec<(usize, usize)>,
}
#[cfg(target_os = "macos")]
struct MetalWarmPending {
graph: crate::gpu_metal::VerifyGraph,
cpu_stored: usize,
b: usize,
}
#[cfg(target_os = "macos")]
enum MetalRowsRun {
Declined,
Failed,
Completed(MetalVerifyPending),
}
#[cfg(target_os = "macos")]
enum MetalPrefillOutcome {
Declined,
Failed,
Completed(Vec<f32>),
}
#[cfg(target_os = "macos")]
enum MetalBatchNllOutcome {
Declined,
Failed(String),
Completed(f64, usize),
}
#[derive(Clone, Copy)]
enum SpecTrial {
Spec {
t0: std::time::Instant,
gen0: usize,
rounds: usize,
},
Plain {
t0: std::time::Instant,
gen0: usize,
},
Decided {
spec: bool,
recheck_at: usize,
},
}
pub(crate) fn spec_time_level() -> u8 {
static L: std::sync::OnceLock<u8> = std::sync::OnceLock::new();
*L.get_or_init(|| match std::env::var("CMF_GRAPH_SPEC_TIME") {
Ok(v) => v.trim().parse::<u8>().map(|n| n.max(1)).unwrap_or(1),
Err(_) => 0,
})
}
struct SpecStampLog {
t_last: std::time::Instant,
items: Vec<(&'static str, f32)>,
}
static SPEC_STAMPS: std::sync::Mutex<Option<SpecStampLog>> = std::sync::Mutex::new(None);
pub(crate) fn spec_stamp(name: &'static str) {
if spec_time_level() == 0 {
return;
}
if let Ok(mut g) = SPEC_STAMPS.lock() {
if let Some(log) = g.as_mut() {
let now = std::time::Instant::now();
log.items
.push((name, (now - log.t_last).as_secs_f32() * 1e3));
log.t_last = now;
}
}
}
fn spec_stamps_begin() {
if spec_time_level() == 0 {
return;
}
if let Ok(mut g) = SPEC_STAMPS.lock() {
*g = Some(SpecStampLog {
t_last: std::time::Instant::now(),
items: Vec::with_capacity(64),
});
}
}
fn spec_stamps_take() -> Vec<(&'static str, f32)> {
SPEC_STAMPS
.lock()
.ok()
.and_then(|mut g| g.take())
.map(|l| l.items)
.unwrap_or_default()
}
fn spec_stamps_format(items: &[(&'static str, f32)]) -> String {
let mut agg: Vec<(&'static str, f32, u32)> = Vec::with_capacity(items.len());
for &(n, ms) in items {
match agg.iter_mut().find(|e| e.0 == n) {
Some(e) => {
e.1 += ms;
e.2 += 1;
}
None => agg.push((n, ms, 1)),
}
}
let mut s = String::with_capacity(agg.len() * 16);
for (n, ms, k) in agg {
if k > 1 {
s.push_str(&format!("{n} {ms:.1}/{k} "));
} else {
s.push_str(&format!("{n} {ms:.1} "));
}
}
s
}
#[derive(Default, Clone, Copy)]
struct SpecMon {
round_ms: f64,
tokens: f64,
plain_ms: f64,
n: u32,
fails: u32,
metal: bool,
}
const SPEC_PROXY_TOKENS: f64 = 3.5;
const SPEC_PLAIN_MIN_MS: f64 = 200.0;
impl SpecMon {
fn round(&mut self, dt_ms: f64, produced: usize) {
self.n += 1;
if self.n == 1 {
return; }
let a = if self.n == 2 { 1.0 } else { 0.3 };
self.round_ms += a * (dt_ms - self.round_ms);
self.tokens += a * (produced as f64 - self.tokens);
}
fn pays(&self) -> bool {
if self.plain_ms > 0.0 {
self.tokens * self.plain_ms > self.round_ms * 1.03
} else {
self.metal && self.tokens >= SPEC_PROXY_TOKENS
}
}
fn plain_done(&self, t0: std::time::Instant, gen0: usize, generated: usize) -> bool {
let n = generated.saturating_sub(gen0);
if n >= 8 {
return true;
}
self.metal && n >= 2 && t0.elapsed().as_secs_f64() * 1e3 >= SPEC_PLAIN_MIN_MS
}
}
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,
}
#[cfg_attr(not(test), allow(dead_code))]
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)
}
#[inline]
fn nll_graph_policy(
unmasked: bool,
prefer_graph: bool,
native_metal: bool,
) -> (bool, bool) {
let graph_quality = unmasked && prefer_graph;
let fused_head_quality = graph_quality && native_metal;
(graph_quality, fused_head_quality)
}
#[derive(Clone, Copy)]
enum PrefillIn<'a> {
Ids(&'a [u32]),
Hidden(&'a [f32]),
}
impl Pipeline {
fn can_prefill_batched(&self) -> bool {
#[cfg(test)]
let force_serial = self.nll_test_force_serial;
#[cfg(not(test))]
let force_serial = false;
prefill_batched() && !force_serial && !self.weights.layers.is_empty()
}
fn automatic_gpu_prefix(&self) -> Option<usize> {
let (model, _, _, _) = self.weights.embed_tokens.graph_weight()?;
crate::gpu::automatic_layer_prefix(&model, self.num_layers, self.physical_layers)
}
pub fn prefill_chunk(&self) -> usize {
let env = env_prefill_chunk();
if env.is_some() || ChunkHost::here() != ChunkHost::Other {
return prefill_chunk_rule(env, ChunkHost::here(), false);
}
prefill_chunk_rule(None, ChunkHost::Other, self.chunk_stack_facts().dense_on_discrete())
}
fn chunk_stack_facts(&self) -> ChunkStackFacts {
let plain_dense = !self.weights.layers.is_empty()
&& self.g3n.is_none()
&& self.dsv4.is_none()
&& self.dsv41.is_none()
&& self.qwen4_exp.is_none()
&& self.weights.layers.iter().all(|lw| {
matches!(lw.attn, AttnKind::Full { .. }) && matches!(lw.ffn, FfnKind::Dense(_))
});
let gpu_on = crate::gpu::enabled();
ChunkStackFacts {
plain_dense,
discrete: gpu_on && crate::gpu::discrete(),
gpu_on,
capacity_split: std::env::var_os("CMF_GPU_LAYERS").is_some()
|| (plain_dense && gpu_on && self.automatic_gpu_prefix().is_some()),
multi_gpu: self.gpu_plan.is_some(),
o1: self.o1_active(),
}
}
}
pub fn prefill_chunk() -> usize {
prefill_chunk_rule(env_prefill_chunk(), ChunkHost::here(), false)
}
fn env_prefill_chunk() -> Option<usize> {
std::env::var("CMF_PREFILL_CHUNK")
.ok()
.and_then(|v| v.parse::<usize>().ok())
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
enum ChunkHost {
Macos,
Aarch64,
Other,
}
impl ChunkHost {
fn here() -> Self {
if cfg!(target_os = "macos") {
ChunkHost::Macos
} else if cfg!(target_arch = "aarch64") {
ChunkHost::Aarch64
} else {
ChunkHost::Other
}
}
}
const DISCRETE_DENSE_PREFILL_CHUNK: usize = 512;
fn prefill_chunk_rule(env: Option<usize>, host: ChunkHost, dense_on_discrete: bool) -> usize {
if let Some(n) = env {
return n.max(1);
}
match host {
ChunkHost::Macos => 512,
ChunkHost::Aarch64 => 256,
ChunkHost::Other if dense_on_discrete => DISCRETE_DENSE_PREFILL_CHUNK,
ChunkHost::Other => 48,
}
}
#[derive(Clone, Copy, Debug, Default)]
struct ChunkStackFacts {
plain_dense: bool,
discrete: bool,
gpu_on: bool,
capacity_split: bool,
multi_gpu: bool,
o1: bool,
}
impl ChunkStackFacts {
fn dense_on_discrete(self) -> bool {
self.plain_dense
&& self.discrete
&& self.gpu_on
&& !self.capacity_split
&& !self.multi_gpu
&& !self.o1
}
}
#[inline]
fn mtp_prefill_pair_count(start: usize, end: usize, input_len: usize) -> usize {
if end <= start || start >= input_len {
return 0;
}
let rows = (end.min(input_len) - start).min(input_len - start);
if end < input_len {
rows
} else {
rows.saturating_sub(1)
}
}
pub type TokenCallback = Box<dyn FnMut(&str) -> bool + Send>;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) struct ReuseLayer {
pub full: bool,
pub host_rows: usize,
pub device_rows: Option<usize>,
pub device_state: bool,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub(crate) enum ReusePlan {
Ready,
Pull(Vec<(usize, usize, usize)>),
Fresh,
}
pub(crate) fn kv_reuse_plan(reuse_from: usize, layers: &[ReuseLayer]) -> ReusePlan {
let mut pulls = Vec::new();
for (li, l) in layers.iter().enumerate() {
if !l.full {
if l.device_state {
return ReusePlan::Fresh;
}
continue;
}
if l.host_rows == reuse_from {
continue;
}
if l.host_rows < reuse_from && l.device_rows.is_some_and(|d| d >= reuse_from) {
pulls.push((li, l.host_rows, reuse_from));
continue;
}
return ReusePlan::Fresh;
}
if pulls.is_empty() {
ReusePlan::Ready
} else {
ReusePlan::Pull(pulls)
}
}
impl Pipeline {
fn clear_sequence_state(&mut self) {
#[cfg(target_os = "macos")]
let _ = crate::gpu_metal::wait_replay();
self.kv_cache.clear();
self.kv_history.clear();
if let Some(b) = &mut self.dsv41 {
b.3.clear();
}
crate::gpu::graph_kv_reset(self.graph_kv_id);
crate::gpu::graph_kv_reset(self.mtp_kv_id());
}
fn prepare_kv_reuse(&mut self, reuse_from: usize) -> bool {
if self.graph_prefill_preferred() {
return true;
}
let kv_id = self.graph_kv_id;
let layers: Vec<ReuseLayer> = (0..self.num_layers)
.map(|li| {
let full = matches!(
self.weights.layers[self.phys_layer(li)].attn,
AttnKind::Full { .. }
);
ReuseLayer {
full,
host_rows: self.kv_cache.layers[li].seq_len,
device_rows: crate::gpu::graph_kv_stored(kv_id, li),
device_state: crate::gpu::graph_state_resident(kv_id, li),
}
})
.collect();
if layers
.iter()
.all(|l| l.device_rows.is_none() && !l.device_state)
{
return true;
}
let plan = kv_reuse_plan(reuse_from, &layers);
let (what, rows, n) = match &plan {
ReusePlan::Ready => ("host ready", 0, 0),
ReusePlan::Fresh => ("fresh", 0, 0),
ReusePlan::Pull(p) => (
"pull",
p.iter().map(|&(_, a, b)| b - a).max().unwrap_or(0),
p.len(),
),
};
let t0 = std::time::Instant::now();
let ok = self.apply_kv_reuse_plan(reuse_from, plan, &layers);
if std::env::var("CMF_PREFILL_PROF").is_ok() {
eprintln!(
"kv-reuse: {what}{}: {rows} device row(s) × {n} layer(s) to the host in {:.2} ms",
if ok { "" } else { " (failed → fresh)" },
t0.elapsed().as_secs_f64() * 1e3
);
}
ok
}
fn apply_kv_reuse_plan(
&mut self,
reuse_from: usize,
plan: ReusePlan,
layers: &[ReuseLayer],
) -> bool {
let kv_id = self.graph_kv_id;
match plan {
ReusePlan::Fresh => return false,
ReusePlan::Ready => {}
ReusePlan::Pull(pulls) => {
let (nkv, hd) = {
let c = &self.kv_cache.layers[pulls[0].0];
(c.num_kv_heads, c.head_dim)
};
if pulls.iter().any(|&(li, _, _)| {
let c = &self.kv_cache.layers[li];
(c.num_kv_heads, c.head_dim) != (nkv, hd)
}) {
return false;
}
let Some(rows) = crate::gpu::graph_kv_read_rows(kv_id, &pulls, nkv, hd) else {
return false;
};
for ((li, from, to), (k, v)) in pulls.into_iter().zip(rows) {
let cache = &mut self.kv_cache.layers[li];
let row = nkv * hd;
for p in 0..to - from {
cache.append(&k[p * row..(p + 1) * row], &v[p * row..(p + 1) * row], &[]);
}
if cache.seq_len != to {
return false;
}
}
}
}
for (li, l) in layers.iter().enumerate() {
if l.full
&& l.device_rows.is_some_and(|d| d > reuse_from)
&& !crate::gpu::graph_kv_set_stored(kv_id, li, reuse_from)
{
return false;
}
}
true
}
fn finish_generation(
&mut self,
mtp: &mut Option<MtpModule>,
router: &mut Option<crate::swarm::DynRouter>,
clear_sequence: bool,
) {
if router.is_some() {
let _ = self.set_active_skill(None);
}
#[cfg(target_os = "macos")]
let clear_sequence = clear_sequence || !crate::gpu_metal::wait_replay();
if clear_sequence {
self.clear_sequence_state();
if let Some(m) = mtp.as_mut() {
m.kv.clear();
}
if let Some(m) = self.mtp.as_mut() {
m.kv.clear();
}
}
self.graph_want_logits = false;
self.graph_head_required = false;
self.graph_logits = None;
self.graph_failed
.store(false, std::sync::atomic::Ordering::Relaxed);
self.cancel
.store(false, std::sync::atomic::Ordering::Relaxed);
self.dyn_router = router.take().or(self.dyn_router.take());
self.mtp = mtp.take().or(self.mtp.take());
self.mtp_graph_mode = None;
self.spec_forced = None;
}
fn check_forward_graph(&mut self, phase: &str, pos: usize) -> Result<(), String> {
if self
.graph_failed
.swap(false, std::sync::atomic::Ordering::Relaxed)
{
self.cancel
.store(false, std::sync::atomic::Ordering::Relaxed);
self.clear_sequence_state();
self.graph_logits = None;
self.graph_want_logits = false;
self.graph_head_required = false;
return Err(format!("GPU graph failed during {phase} at position {pos}"));
}
Ok(())
}
#[cfg(target_os = "macos")]
fn fail_metal_graph(&mut self, reason: &str) {
crate::pipeline::METAL_GRAPH_ERRORS
.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
self.clear_sequence_state();
self.graph_logits = None;
self.graph_failed
.store(true, std::sync::atomic::Ordering::Relaxed);
self.cancel
.store(true, std::sync::atomic::Ordering::Relaxed);
tracing::error!("native Metal TokenGraph failed closed: {reason}");
}
fn nll_begin(&mut self) -> Result<(), String> {
if self
.graph_failed
.swap(false, std::sync::atomic::Ordering::Relaxed)
{
self.cancel
.store(false, std::sync::atomic::Ordering::Relaxed);
self.clear_sequence_state();
self.graph_logits = None;
self.graph_want_logits = false;
self.graph_head_required = false;
return Err("GPU graph failed before NLL scoring".to_string());
}
self.clear_sequence_state();
self.graph_logits = None;
self.graph_want_logits = false;
self.graph_head_required = false;
Ok(())
}
fn nll_end(&mut self) {
self.clear_sequence_state();
self.graph_logits = None;
self.graph_want_logits = false;
self.graph_head_required = false;
self.graph_failed
.store(false, std::sync::atomic::Ordering::Relaxed);
}
fn nll_check_graph(&mut self, phase: &str, pos: usize) -> Result<(), String> {
#[cfg(test)]
if self.nll_test_fail_at == Some(pos) {
self.nll_test_fail_at = None;
self.graph_failed
.store(true, std::sync::atomic::Ordering::Relaxed);
self.cancel
.store(true, std::sync::atomic::Ordering::Relaxed);
}
if self
.graph_failed
.swap(false, std::sync::atomic::Ordering::Relaxed)
{
self.cancel
.store(false, std::sync::atomic::Ordering::Relaxed);
self.clear_sequence_state();
self.graph_logits = None;
self.graph_want_logits = false;
return Err(format!(
"GPU graph failed during NLL {phase} at position {pos}"
));
}
Ok(())
}
#[inline]
pub fn phys_layer(&self, virtual_idx: usize) -> usize {
virtual_idx % self.physical_layers
}
#[inline]
pub fn is_loop_end(&self, virtual_idx: usize) -> bool {
self.loop_final_norm && (virtual_idx + 1) % self.physical_layers == 0
}
#[allow(clippy::too_many_arguments)]
#[cfg(target_os = "macos")]
fn graph_prefill_preferred(&self) -> bool {
let graph_force = crate::gpu::q1_force() || crate::gpu::q2tp_gpu_opt_in();
if !crate::gpu::enabled_here()
|| !graph_force
|| std::env::var("CMF_GPU_BLOCK")
.map(|v| v == "0")
.unwrap_or(false)
|| std::env::var("CMF_PREFILL_GRAPH").as_deref() == Ok("0")
{
return false;
}
self.weights
.layers
.iter()
.any(|lw| {
matches!(&lw.attn, AttnKind::LinearGdn(w) if w.in_proj_qkv.metal_graph_parts().is_some())
})
}
#[cfg(not(target_os = "macos"))]
fn graph_prefill_preferred(&self) -> bool {
let graph_on = crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Prefill);
if !graph_on || !crate::gpu::enabled_here() {
return false;
}
if self.o1_active() {
return false;
}
if self
.weights
.layers
.iter()
.any(|lw| matches!(&lw.attn, AttnKind::LinearGdn(_)))
{
return true;
}
self.weights
.layers
.iter()
.any(|lw| matches!(&lw.ffn, FfnKind::Moe(_)))
&& self.automatic_gpu_prefix().is_none()
}
#[cfg(target_os = "macos")]
fn q1_graph_gpu(
&mut self,
start: usize,
upto: Option<usize>,
position: usize,
h: &mut [f32],
) -> usize {
let _mt0 = std::time::Instant::now(); use crate::gpu::{AttnGpuLayer, GdnGpuCfg, GdnGpuLayer, GraphDims, MetalFfn, TokenGraph};
let graph_force = crate::gpu::q1_force() || crate::gpu::q2tp_gpu_opt_in();
if self.attn_softcap > 0.0 || !crate::gpu::enabled_here()
|| !graph_force
|| std::env::var("CMF_GPU_BLOCK")
.map(|v| v == "0")
.unwrap_or(false)
{
if std::env::var("CMF_GRAPH_DBG").is_ok() {
eprintln!(
"block-graph: front gate (softcap={} enabled_here={} graph_force={})",
self.attn_softcap > 0.0,
crate::gpu::enabled_here(),
graph_force,
);
}
if self.graph_head_required {
self.fail_metal_graph("native graph front gate refused");
}
return start;
}
if self.swa.is_some()
|| self.global_attn.is_some()
|| self.attention_heads_per_layer.is_some()
|| self.attn_v_norm
|| self.weights.layers.iter().any(|lw| {
lw.attn_out_norm.is_some()
|| lw.ffn_out_norm.is_some()
|| lw.layer_scale.is_some()
|| matches!(&lw.ffn, FfnKind::Dense(d) if d.act != Act::Silu)
})
{
if std::env::var("CMF_GRAPH_DBG").is_ok() {
eprintln!(
"block-graph: arch ineligible (swa={} gattn={} hpl={} vnorm={} scale_delta={:.2e})",
self.swa.is_some(),
self.global_attn.is_some(),
self.attention_heads_per_layer.is_some(),
self.attn_v_norm,
(self.attn_scale - 1.0 / (self.head_dim as f32).sqrt()).abs(),
);
}
if self.graph_head_required {
self.fail_metal_graph("native graph architecture gate refused");
}
return start;
}
let limit = upto
.map(|u| u + 1)
.unwrap_or(self.num_layers)
.min(self.num_layers);
enum Item<'a> {
Gdn {
run: Vec<GdnGpuLayer<'a>>,
first: usize,
},
Attn {
l: AttnGpuLayer<'a>,
li: usize,
q_norm: Option<&'a [f32]>,
k_norm: Option<&'a [f32]>,
output_gate: bool,
bias: Option<(&'a [f32], &'a [f32], &'a [f32])>,
full_gpu: bool,
},
}
let attend_mode = std::env::var("CMF_GPU_ATTEND").unwrap_or_else(|_| "auto".into());
let attend_contract = attend_mode != "0"
&& attend_mode != "off"
&& self.head_dim % 4 == 0
&& self.head_dim <= 256
&& self.rotary_dim >= 2
&& self.rotary_dim <= self.head_dim
&& (self.rotary_dim / 2) % 32 == 0
&& self.num_kv_heads > 0
&& self.num_heads % self.num_kv_heads == 0;
let mut plan: Vec<Item> = Vec::new();
let mut model_ref: Option<std::sync::Arc<cortiq_core::CmfModel>> = None;
let block_diag = std::env::var("CMF_GRAPH_DBG").is_ok();
let mut scan = start;
while scan < limit {
let lw = &self.weights.layers[self.phys_layer(scan)];
let ffn = match &lw.ffn {
FfnKind::Dense(d) if d.segs.is_empty() => {
let (Some(g), Some(u), Some(dn)) = (
d.gate_proj.metal_graph_parts(),
d.up_proj.metal_graph_parts(),
d.down_proj.metal_graph_parts(),
) else {
if block_diag {
eprintln!(
"block-graph: L{scan} FFN trio not graph-mappable — run ends"
);
}
break;
};
MetalFfn::Dense {
gate: g,
up: u,
down: dn,
}
}
FfnKind::Moe(m) => {
let Some(moe) = metal_moe_graph_parts(m, self.hidden_size) else {
if block_diag {
eprintln!(
"block-graph: L{scan} MoE outside the graph contract — run ends"
);
}
break;
};
if let QTensor::Mapped { model, .. } = &m.experts[0].gate_proj {
model_ref.get_or_insert_with(|| model.clone());
}
MetalFfn::Moe(moe)
}
_ => {
if block_diag {
eprintln!("block-graph: L{scan} non-graph FFN — run ends");
}
break;
}
};
match &lw.attn {
AttnKind::LinearGdn(w) if self.gdn_cfg.is_some() => {
let parts = (
w.in_proj_qkv.metal_graph_parts(),
w.in_proj_z.metal_graph_parts(),
w.in_proj_a.f32_parts(),
w.in_proj_b.f32_parts(),
w.out_proj.metal_graph_parts(),
);
let (Some(qkv), Some(z), Some(a), Some(b), Some(out)) = parts else {
if block_diag {
eprintln!(
"block-graph: L{scan} GDN parts refused (qkv={} z={} a_f32={} b_f32={} out={})",
w.in_proj_qkv.metal_graph_parts().is_some(),
w.in_proj_z.metal_graph_parts().is_some(),
w.in_proj_a.f32_parts().is_some(),
w.in_proj_b.f32_parts().is_some(),
w.out_proj.metal_graph_parts().is_some(),
);
}
break;
};
if let QTensor::Mapped { model, .. } = &w.in_proj_qkv {
model_ref.get_or_insert_with(|| model.clone());
}
let gl = GdnGpuLayer {
attn_norm: &lw.input_norm,
post_norm: &lw.post_norm,
qkv,
z,
a,
b,
out,
ffn,
conv1d: &w.conv1d,
a_log: &w.a_log,
dt_bias: &w.dt_bias,
gnorm: &w.norm,
};
match plan.last_mut() {
Some(Item::Gdn { run, .. }) => run.push(gl),
_ => plan.push(Item::Gdn {
run: vec![gl],
first: scan,
}),
}
}
AttnKind::Full {
wq,
wk,
wv,
wo,
q_norm,
k_norm,
output_gate,
softplus_gate: None,
bias,
} if !self.kv_cache.layers[scan].o1_sealed()
|| std::env::var("CMF_O1_METAL").as_deref() == Ok("1") =>
{
let parts = (
wq.metal_graph_parts(),
wk.metal_graph_parts(),
wv.metal_graph_parts(),
wo.metal_graph_parts(),
);
let (Some(pq), Some(pk), Some(pv), Some(po)) = parts else {
break;
};
if let QTensor::Mapped { model, .. } = wq {
model_ref.get_or_insert_with(|| model.clone());
}
let cache = &self.kv_cache.layers[scan];
let o1_metal = cache.o1.is_some()
&& std::env::var("CMF_O1_METAL").as_deref() == Ok("1")
&& cache.o1_views().is_some();
let full_gpu = attend_contract
&& cache.mode == crate::kv_cache::KvMode::F32
&& (cache.o1.is_none() || o1_metal)
&& bias.is_none()
&& pq.1 == self.num_heads * self.head_dim * (1 + *output_gate as usize)
&& pk.1 == self.num_kv_heads * self.head_dim
&& pv.1 == self.num_kv_heads * self.head_dim
&& po.2 == self.num_heads * self.head_dim;
plan.push(Item::Attn {
l: AttnGpuLayer {
attn_norm: &lw.input_norm,
post_norm: &lw.post_norm,
wq: pq,
wk: pk,
wv: pv,
wo: po,
ffn,
},
li: scan,
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())),
full_gpu,
});
}
_ => break,
}
scan += 1;
}
let Some(model) = model_ref else {
if std::env::var("CMF_GRAPH_DBG").is_ok() {
eprintln!("q1-graph: no model ref (start {start}, scanned to {scan})");
}
if self.graph_head_required {
self.fail_metal_graph("native graph has no mapped model reference");
}
return start;
};
if plan.is_empty() {
if std::env::var("CMF_GRAPH_DBG").is_ok() {
eprintln!("q1-graph: empty plan at layer {start}");
}
if self.graph_head_required {
self.fail_metal_graph("native graph plan is empty");
}
return start;
}
let has_moe = plan.iter().any(|it| match it {
Item::Gdn { run, .. } => run.iter().any(|l| matches!(l.ffn, MetalFfn::Moe(_))),
Item::Attn { l, .. } => matches!(l.ffn, MetalFfn::Moe(_)),
});
let has_gdn = plan.iter().any(|it| matches!(it, Item::Gdn { .. }));
let dev_attend = attend_contract
&& (self.head_dim <= 128
|| has_moe
|| (self.head_dim <= 256 && has_gdn)
|| attend_mode == "force"
|| attend_mode == "256");
if !dev_attend {
for it in &mut plan {
if let Item::Attn { li, full_gpu, .. } = it {
let keep_o1 = self.kv_cache.layers[*li].o1.is_some()
&& std::env::var("CMF_O1_METAL").as_deref() == Ok("1");
if !keep_o1 {
*full_gpu = false;
}
}
}
}
if std::env::var("CMF_GRAPH_DBG").is_ok() {
use std::sync::atomic::{AtomicBool, Ordering};
static SAID: AtomicBool = AtomicBool::new(false);
if !SAID.swap(true, Ordering::Relaxed) {
let fg = plan
.iter()
.filter(|it| matches!(it, Item::Attn { full_gpu: true, .. }))
.count();
let att = plan
.iter()
.filter(|it| matches!(it, Item::Attn { .. }))
.count();
eprintln!(
"q1-graph: plan of {} items from layer {start} to {scan} | dev_attend={dev_attend} full_gpu {fg}/{att} | hd={} rd={} nkv={} nh={}",
plan.len(),
self.head_dim,
self.rotary_dim,
self.num_kv_heads,
self.num_heads,
);
}
}
let dims = GraphDims {
hidden: self.hidden_size,
eps: self.rms_eps as f32,
gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
};
let Some(mut graph) = TokenGraph::new(&model, dims, h) else {
if self.graph_head_required {
self.fail_metal_graph("native TokenGraph allocation refused");
}
return start;
};
let gcfg = self.gdn_cfg.map(|cfg| GdnGpuCfg {
nv: cfg.num_v_heads,
nk: cfg.num_k_heads,
dk: cfg.key_head_dim,
dv: cfg.value_head_dim,
kk: cfg.conv_kernel,
hidden: self.hidden_size,
inter: self.intermediate_size,
c_dim: cfg.conv_dim(),
eps: cfg.rms_eps as f32,
gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
});
let mut valid = 0usize;
let mut end = start;
crate::gpu::stageprof(1, _mt0.elapsed()); if std::env::var("CMF_PLAN_DUMP").is_ok() {
static ONCE: std::sync::Once = std::sync::Once::new();
ONCE.call_once(|| {
for it in &plan {
match it {
Item::Gdn { first, run } => {
eprintln!("plan: Gdn first={first} len={}", run.len())
}
Item::Attn { li, full_gpu, .. } => {
eprintln!("plan: Attn li={li} full_gpu={full_gpu}")
}
}
}
});
}
for item in &plan {
let ok = match item {
Item::Gdn { run, .. } => gcfg
.as_ref()
.map(|gc| run.iter().all(|l| graph.gdn_ok(l, gc)))
.unwrap_or(false),
Item::Attn { l, .. } => graph.attn_ok(l),
};
if !ok {
if block_diag {
eprintln!(
"block-graph: plan item {} ({}) failed graph preflight",
valid,
match item {
Item::Gdn { run, first } => format!("GDN run L{first}+{}", run.len()),
Item::Attn { li, .. } => format!("Attn L{li}"),
}
);
}
break;
}
valid += 1;
end += match item {
Item::Gdn { run, .. } => run.len(),
Item::Attn { .. } => 1,
};
}
plan.truncate(valid);
if plan.is_empty() {
if self.graph_head_required {
self.fail_metal_graph("native graph preflight produced no valid items");
}
return start;
}
if self.graph_head_required && (upto.is_some() || end != self.num_layers) {
self.fail_metal_graph("fused-head NLL requires a complete 64-layer graph");
return start;
}
let one_pass = |t: (usize, usize, usize)| {
use cortiq_core::TensorDtype as D;
matches!(
model.tensors[t.0].dtype,
D::Q4TiledP | D::Q4Tiled | D::Q4Block | D::Q8Row | D::Q8_2f | D::Q1
)
};
let dense_fast = plan.iter().all(|it| match it {
Item::Attn {
l, li, full_gpu, ..
} => {
*full_gpu
&& self.kv_cache.layers[*li].o1.is_none()
&& [l.wq, l.wk, l.wv, l.wo].into_iter().all(one_pass)
&& match l.ffn {
MetalFfn::Dense { gate, up, down } => {
one_pass(gate) && one_pass(up) && one_pass(down)
}
_ => false,
}
}
Item::Gdn { .. } => false,
});
let ab = crate::gpu_metal::dense_ab_arm().filter(|_| dense_fast);
let _mv_fast = match ab {
Some((bits, _)) => {
graph.set_dense_concurrent_raw(bits & crate::gpu_metal::DENSE_CONC != 0);
crate::gpu_metal::MvFastGuard::set_raw(bits)
}
None => {
graph.set_dense_concurrent(dense_fast);
crate::gpu_metal::MvFastGuard::set_bits(if dense_fast {
crate::gpu_metal::DENSE_MV | crate::gpu_metal::DENSE_FUSE
} else {
0
})
}
};
let inv_freq = self.inv_freq.clone();
let pool = self.pool.clone();
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 norm_style = self.norm_style;
let gemma = norm_style == cortiq_core::NormStyle::Gemma;
let want = self.gdn_cfg.map(|c| c.state_len()).unwrap_or(0);
let kv_id = self.graph_kv_id;
let mut pending: Vec<(usize, usize)> = Vec::new();
let mut dev_attn: Vec<usize> = Vec::new();
for item in &plan {
let _xt0 = std::time::Instant::now();
let _xkind: u32 = match item {
Item::Gdn { .. } => 2,
Item::Attn { .. } => 3,
};
if self.loop_final_norm {
let item_start = match item {
Item::Gdn { first, .. } => *first,
Item::Attn { li, .. } => *li,
};
if item_start > start && self.is_loop_end(item_start - 1) {
graph.encode_loop_norm(&self.weights.final_norm);
}
}
match item {
Item::Gdn { run, first } => {
for l in &mut self.kv_cache.layers[*first..*first + run.len()] {
if l.linear_state.len() != want {
l.linear_state = vec![0f32; want];
}
}
let ro: Vec<&[f32]> = self.kv_cache.layers[*first..*first + run.len()]
.iter()
.map(|l| l.linear_state.as_slice())
.collect();
let _ig = std::time::Instant::now();
if !graph.encode_gdn_run(run, &ro, gcfg.as_ref().unwrap()) {
tracing::error!("q1 graph: GDN run refused after validation");
return start;
}
graph.commit_kind = 2;
graph.commit();
crate::gpu::stageprof(0, _ig.elapsed());
pending.push((*first, run.len()));
}
Item::Attn {
l,
li,
q_norm,
k_norm,
output_gate,
bias,
full_gpu,
} => {
let _ia = std::time::Instant::now();
if *full_gpu {
let cache = &self.kv_cache.layers[*li];
let o1p = if cache.o1.is_some() {
match cache.o1_views() {
Some(views) => Some(crate::gpu::O1AttnParams {
views,
epoch: self.o1_epoch,
}),
None => None,
}
} else {
None
};
let o1_layer = cache.o1.is_some();
if o1_layer && o1p.is_none() {
}
let cpu_k: Vec<&[f32]> = (0..nkv).map(|g| cache.head_keys(g)).collect();
let cpu_v: Vec<&[f32]> = (0..nkv).map(|g| cache.head_values(g)).collect();
let cpu_stored = if o1_layer { 0 } else { cpu_k[0].len() / hd };
let p = crate::gpu::AttnDeviceParams {
kv_id,
layer: *li,
nh,
nkv,
hd,
rd,
position,
scale: self.attn_scale,
eps: eps as f32,
gemma,
late_qk_norm: self.qk_norm_after_rope,
output_gate: *output_gate,
q_norm: *q_norm,
k_norm: *k_norm,
inv_freq: &inv_freq,
cpu_k,
cpu_v,
cpu_stored,
o1: o1p,
};
let o1_bad = o1_layer && p.o1.is_none();
if !o1_bad && graph.attn_device_ok(l, &p) && graph.encode_attn_device(l, &p)
{
if p.o1.is_none() {
dev_attn.push(*li);
}
graph.commit_kind = 3;
graph.commit();
crate::gpu::stageprof(_xkind, _xt0.elapsed());
continue;
}
}
graph.encode_attn_prefix(l);
if let Err(err) = graph.sync_checked() {
self.fail_metal_graph(&err);
return start;
}
if !pending.is_empty() {
let idxs: Vec<usize> =
pending.drain(..).flat_map(|(f, n)| f..f + n).collect();
let mut outs: Vec<&mut [f32]> = self
.kv_cache
.layers
.iter_mut()
.enumerate()
.filter(|(i, _)| idxs.binary_search(i).is_ok())
.map(|(_, s)| s.linear_state.as_mut_slice())
.collect();
graph.read_states(&mut outs);
}
let mut q_raw = attention::take_buf(l.wq.1);
let mut k = attention::take_buf(l.wk.1);
let mut v = attention::take_buf(l.wv.1);
graph.read_qkv(&mut q_raw, &mut k, &mut v);
let cfg = QwenAttnCfg {
num_heads: nh,
num_kv_heads: nkv,
head_dim: hd,
hidden_size: hs,
position,
inv_freq: &inv_freq,
rotary_dim: rd,
scale: self.attn_scale,
softcap: self.attn_softcap,
window: None,
v_norm: false,
qk_norm_after_rope: self.qk_norm_after_rope,
q_norm: *q_norm,
k_norm: *k_norm,
output_gate: *output_gate,
softplus_gate: None,
rope_scale: 1.0,
bias: *bias,
rms_eps: eps,
norm_style,
pool: pool.as_deref(),
};
let oracle = std::env::var("CMF_ATTN_ORACLE").as_deref() == Ok("1")
|| std::env::var("CMF_ATTN_DUMP").is_ok();
let _ = full_gpu;
let oracle_in = oracle.then(|| (q_raw.clone(), k.clone(), v.clone()));
let mut ao = attention::qwen_attention_core(
q_raw,
k,
v,
&mut self.kv_cache.layers[*li],
&cfg,
);
if let Ok(dir) = std::env::var("CMF_ATTN_DUMP") {
if let Some((qr0, k0, v0)) = oracle_in.clone() {
let (cq, _cg, _ck, _cv) =
attention::finish_projection_debug(qr0, k0, v0, &cfg, position);
let cache = &self.kv_cache.layers[*li];
let n = cache.head_keys(0).len() / hd;
let mut bytes: Vec<u8> = Vec::new();
for v in [nh as u32, nkv as u32, hd as u32, n as u32, position as u32] {
bytes.extend_from_slice(&v.to_le_bytes());
}
for v in &cq {
bytes.extend_from_slice(&v.to_le_bytes());
}
for g in 0..nkv {
for v in cache.head_keys(g) {
bytes.extend_from_slice(&v.to_le_bytes());
}
}
for g in 0..nkv {
for v in cache.head_values(g) {
bytes.extend_from_slice(&v.to_le_bytes());
}
}
let _ =
std::fs::write(format!("{dir}/L{li}_pos{position}.bin"), &bytes);
}
}
if let Some((qr0, k0, v0)) =
oracle_in.filter(|_| std::env::var("CMF_ATTN_ORACLE").as_deref() == Ok("1"))
{
let (cq, _cg, ck, cv) =
attention::finish_projection_debug(qr0, k0, v0, &cfg, position);
let mut h_now = vec![0f32; hs];
graph.read_h(&mut h_now);
let cache = &self.kv_cache.layers[*li];
let n_after = cache.head_keys(0).len() / hd;
let stored = n_after.saturating_sub(1);
let cpu_k: Vec<&[f32]> = (0..nkv)
.map(|g| &cache.head_keys(g)[..stored * hd])
.collect();
let cpu_v: Vec<&[f32]> = (0..nkv)
.map(|g| &cache.head_values(g)[..stored * hd])
.collect();
let p = crate::gpu::AttnDeviceParams {
kv_id,
layer: *li,
nh,
nkv,
hd,
rd,
position,
scale: self.attn_scale,
eps: eps as f32,
gemma,
late_qk_norm: self.qk_norm_after_rope,
output_gate: *output_gate,
q_norm: *q_norm,
k_norm: *k_norm,
inv_freq: &inv_freq,
cpu_k,
cpu_v,
cpu_stored: stored,
o1: None,
};
if let Some((dq, dk, dv, dao)) = graph.debug_attn_device(l, &p, &h_now) {
let md = |a: &[f32], b: &[f32]| {
a.iter().zip(b).fold(0f32, |m, (x, y)| m.max((x - y).abs()))
};
let nn = |a: &[f32]| a.iter().map(|x| x * x).sum::<f32>().sqrt();
eprintln!(
"attn-oracle L{li} pos {position}: |q| {:.2} max|dq| {:.4} | |k| {:.2} max|dk| {:.4} | |v| {:.2} max|dv| {:.4} | |ao| {:.2} max|dao| {:.4}",
nn(&cq),
md(&cq, &dq),
nn(&ck),
md(&ck, &dk),
nn(&cv),
md(&cv, &dv),
nn(&ao),
md(&ao, &dao)
);
} else {
eprintln!("attn-oracle L{li}: device probe declined");
}
}
graph.encode_attn_suffix(l, &ao);
graph.commit();
attention::recycle_buf(&mut ao);
}
}
crate::gpu::stageprof(_xkind, _xt0.elapsed());
}
let mut lm_rows = None;
if self.graph_want_logits
&& upto.is_none()
&& end == self.num_layers
&& std::env::var("CMF_GPU_LMHEAD")
.map(|v| v != "0")
.unwrap_or(true)
{
if let Some(lm) = self.weights.lm_head.metal_graph_parts() {
if graph.lm_head_ok(lm) {
graph.encode_lm_head(&self.weights.final_norm, lm);
lm_rows = Some(lm.1);
}
}
}
if self.graph_head_required && lm_rows.is_none() {
METAL_GRAPH_HEAD_MISS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
self.fail_metal_graph("fused graph head was requested but not encodable");
return start;
}
let _sy0 = std::time::Instant::now();
if let Err(err) = graph.sync_checked() {
self.fail_metal_graph(&err);
return start;
}
let _rs0 = std::time::Instant::now();
if !pending.is_empty() {
let idxs: Vec<usize> = pending.drain(..).flat_map(|(f, n)| f..f + n).collect();
let mut outs: Vec<&mut [f32]> = self
.kv_cache
.layers
.iter_mut()
.enumerate()
.filter(|(i, _)| idxs.binary_search(i).is_ok())
.map(|(_, s)| s.linear_state.as_mut_slice())
.collect();
graph.read_states(&mut outs);
}
if std::env::var("CMF_GRAPH_HOSTPROF").as_deref() == Ok("1") {
use std::sync::atomic::{AtomicU64, Ordering};
static SY: AtomicU64 = AtomicU64::new(0);
static RS: AtomicU64 = AtomicU64::new(0);
static N: AtomicU64 = AtomicU64::new(0);
SY.fetch_add((_rs0 - _sy0).as_nanos() as u64, Ordering::Relaxed);
RS.fetch_add(_rs0.elapsed().as_nanos() as u64, Ordering::Relaxed);
let n = N.fetch_add(1, Ordering::Relaxed) + 1;
if n % 100 == 0 {
eprintln!(
"postprof: sync-wait {:.1} ms/ток | read_states {:.1} ms/ток ({n})",
SY.load(Ordering::Relaxed) as f64 / n as f64 / 1e6,
RS.load(Ordering::Relaxed) as f64 / n as f64 / 1e6
);
}
}
if let Some(rows) = lm_rows {
crate::gpu::hostprof_encode_done(_mt0);
let mut lg = attention::take_buf(rows.min(self.vocab_size));
graph.read_logits(&mut lg);
crate::gpu::hostprof_total(_mt0);
lg.resize(self.vocab_size, 0.0);
if let Some(c) = self.final_softcap {
for l in lg.iter_mut() {
*l = c * (*l / c).tanh();
}
}
self.graph_logits = Some(lg);
}
graph.read_h(h);
if self.graph_head_required && self.graph_logits.is_none() {
METAL_GRAPH_HEAD_MISS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
self.fail_metal_graph("fused graph head completed without logits readback");
return start;
}
METAL_GRAPH_TOK_OK.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
METAL_GRAPH_LAYERS.fetch_add(
end.saturating_sub(start) as u64,
std::sync::atomic::Ordering::Relaxed,
);
if self.graph_head_required {
METAL_GRAPH_HEAD_OK.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
}
for li in dev_attn {
let mut krow = attention::take_buf(nkv * hd);
let mut vrow = attention::take_buf(nkv * hd);
if crate::gpu::kv_mirror_read_last(kv_id, li, nkv, hd, &mut krow, &mut vrow) {
let cache = &mut self.kv_cache.layers[li];
cache.append(&krow, &vrow, &[]);
let n = cache.seq_len;
let mut imp = attention::take_buf(n);
crate::gpu::kv_mirror_take_imp(kv_id, li, &mut imp);
cache.accumulate_imp(&imp);
attention::recycle_buf(&mut imp);
}
attention::recycle_buf(&mut krow);
attention::recycle_buf(&mut vrow);
}
if let Some((_, arm)) = ab {
crate::gpu_metal::dense_ab_record(arm, _mt0.elapsed());
}
end
}
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,
physical_layers: usize,
loop_final_norm: bool,
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());
if let Some(model) = weights
.lm_head
.model_arc()
.or_else(|| weights.embed_tokens.model_arc())
{
let regions: Vec<&[u8]> =
model.tensors.iter().map(|t| model.entry_bytes(t)).collect();
p.bind_numa(®ions);
}
}
Self {
gpu_plan: None,
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,
physical_layers,
loop_final_norm,
vocab_size,
rms_eps,
rope_base,
norm_style,
rotary_dim: head_dim,
attention_heads_per_layer: None,
vmf_cfg: None,
gdn_cfg: None,
kda_cfg: None,
g3n: None,
dsv4: None,
dsv41: None,
dsv41_vision: None,
dsv41_prefill: None,
qwen4_exp: None,
dsv4_mtp: Vec::new(),
dspark: None,
dspark_pending: Vec::new(),
dspark_hist: Vec::new(),
dspark_real: Vec::new(),
dspark_trunk_picks: Vec::new(),
dspark_exp: Vec::new(),
dspark_draft_ns: 0,
logit_multiplier: None,
cancel: std::sync::Arc::new(std::sync::atomic::AtomicBool::new(false)),
graph_failed: std::sync::atomic::AtomicBool::new(false),
kv_history: Vec::new(),
short_conv_cfg: None,
mtp: None,
speculative: std::env::var("CMF_MTP").map(|v| v != "0").unwrap_or(true),
ignore_eos: false,
draft_full_streak: 0,
spec_k_adapt: None,
spec_acc_ewma: 0.7,
rng,
sampler_scratch: SamplerScratch::default(),
spec_forced: None,
spec_q: Vec::new(),
spec_p: Vec::new(),
spec_res: Vec::new(),
spec_qs: Vec::new(),
spec_ps: Vec::new(),
spec_ress: Vec::new(),
mtp_graph_mode: None,
#[cfg(target_os = "macos")]
metal_verify: None,
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_epoch: 0,
o1_flags: Vec::new(),
trace: false,
calib_temp: 1.0,
confidence_on: true,
embed_multiplier: 1.0,
attn_scale: 1.0 / (head_dim as f32).sqrt(),
swa: None,
sliding_layers: None,
inv_freq_local: None,
rotary_dim_local: None,
rope_scale: 1.0,
rope_scale_local: 1.0,
global_attn: None,
inv_freq_global: None,
attn_v_norm: false,
qk_norm_after_rope: false,
final_softcap: None,
head_clusters: None,
attn_softcap: 0.0,
graph_want_logits: false,
graph_head_required: false,
graph_logits: None,
graph_kv_id: {
static NEXT: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(1);
NEXT.fetch_add(1, std::sync::atomic::Ordering::Relaxed)
},
#[cfg(test)]
nll_test_fail_at: None,
#[cfg(test)]
nll_test_force_serial: false,
}
}
pub fn set_o1(&mut self, cfg: Option<crate::nystrom::O1Cfg>) {
if let Some(c) = &cfg {
if crate::nystrom::o1_deferred_boundary(c.w, c.sink).is_none() {
tracing::error!(
"o1 disabled: w + sink + slack + 1 overflows usize (w={}, sink={})",
c.w,
c.sink
);
self.o1_flags.clear();
self.o1_cfg = None;
return;
}
}
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[self.phys_layer(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)
}
pub fn generation_batch_k(&self) -> usize {
if let Some(k) = std::env::var("CMF_BATCH_K")
.ok()
.and_then(|v| v.parse::<usize>().ok())
{
return k;
}
#[cfg(not(target_os = "macos"))]
if self.graph_prefill_preferred() && !self.o1_active() {
return 32;
}
0
}
pub fn generation_graph_prefill(&self) -> bool {
let graph = self.graph_prefill_preferred();
#[cfg(not(target_os = "macos"))]
if graph
&& self.generation_batch_k() > 0
&& std::env::var("CMF_MTP_CHAIN_PROBE").is_err()
{
return false;
}
graph
}
pub fn o1_device_stats(&self) -> (usize, u64) {
crate::gpu::o1_device_stats(self.graph_kv_id)
}
pub fn o1_begin(&mut self) {
self.o1_begin_with_prefix(None);
}
pub fn o1_begin_with_prefix(&mut self, requested_prefix: Option<usize>) {
if let Some(c) = &self.o1_cfg {
let (m, w, sink, rect) = (c.m, c.w, c.sink, c.rect);
let boundary = requested_prefix.map(|p| {
p.max(
crate::nystrom::o1_deferred_boundary(w, sink)
.expect("o1 config boundary validated in set_o1"),
)
});
for (li, &f) in self.o1_flags.iter().enumerate() {
if f {
self.kv_cache.layers[li].o1_begin_with_boundary(m, w, sink, rect, boundary);
}
}
}
}
fn o1_effective_boundary(&self, requested_prefix: usize) -> Option<usize> {
self.o1_cfg.as_ref().and_then(|c| {
crate::nystrom::o1_deferred_boundary(c.w, c.sink)
.map(|floor| requested_prefix.max(floor))
})
}
fn o1_note_transition(&mut self) {
let mut transitioned = false;
for (li, &flagged) in self.o1_flags.iter().enumerate() {
if flagged {
transitioned |= self.kv_cache.layers[li].take_o1_transition();
}
}
if transitioned {
self.o1_epoch = self.o1_epoch.wrapping_add(1);
}
}
fn o1_pending(&self) -> bool {
self.o1_flags.iter().enumerate().any(|(li, &f)| {
f && self.kv_cache.layers[li].seq_len > 0
&& self.kv_cache.layers[li].o1_pending_boundary().is_some()
})
}
fn o1_fail(&mut self, err: String) {
tracing::error!("o1 deferred seal failed; terminating sequence: {err}");
self.clear_sequence_state();
self.graph_failed
.store(true, std::sync::atomic::Ordering::Relaxed);
self.cancel
.store(true, std::sync::atomic::Ordering::Relaxed);
}
pub fn o1_seal_checked(&mut self) -> Result<bool, String> {
if self.o1_cfg.is_none() {
return Ok(false);
}
let mut participating = false;
for li in 0..self.num_layers {
if !self.o1_flags.get(li).copied().unwrap_or(false) {
continue;
}
if let Some(err) = self.kv_cache.layers[li].take_o1_error() {
return Err(err);
}
if self.kv_cache.layers[li].seq_len == 0 {
continue;
}
participating = true;
let num_heads = self.layer_num_heads(li);
self.kv_cache.layers[li].o1_seal_checked(num_heads)?;
}
self.o1_note_transition();
for li in 0..self.num_layers {
if self.o1_flags.get(li).copied().unwrap_or(false) {
if let Some(err) = self.kv_cache.layers[li].take_o1_error() {
return Err(err);
}
}
}
Ok(participating
&& (0..self.num_layers).all(|li| {
!self.o1_flags.get(li).copied().unwrap_or(false)
|| self.kv_cache.layers[li].seq_len == 0
|| self.kv_cache.layers[li].o1_sealed()
}))
}
fn o1_progress(&mut self) {
if !self.o1_active() {
return;
}
for li in 0..self.num_layers {
if self.o1_flags.get(li).copied().unwrap_or(false) {
if let Some(err) = self.kv_cache.layers[li].take_o1_error() {
self.o1_fail(err);
return;
}
}
}
self.o1_note_transition();
if !self.o1_pending() {
return;
}
if let Err(err) = self.o1_seal_checked() {
self.o1_fail(err);
}
}
fn check_o1_progress_failure(&mut self, phase: &str) -> Result<(), String> {
if self
.graph_failed
.swap(false, std::sync::atomic::Ordering::Relaxed)
{
self.cancel
.store(false, std::sync::atomic::Ordering::Relaxed);
self.clear_sequence_state();
return Err(format!("{phase}: deferred O(1) transition failed"));
}
Ok(())
}
pub fn o1_seal(&mut self) {
if let Err(err) = self.o1_seal_checked() {
self.o1_fail(err);
}
}
pub fn set_trace(&mut self, on: bool) {
self.trace = on;
}
pub fn set_sampler_config(&mut self, config: SamplerConfig) {
self.rng = match config.seed {
Some(seed) => SplitMix64::new(seed),
None => SplitMix64::from_entropy(),
};
self.sampler_config = config;
}
pub fn set_confidence(&mut self, on: bool) {
self.confidence_on = 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,
scale: self.attn_scale,
softcap: self.attn_softcap,
window: None,
v_norm: false,
qk_norm_after_rope: self.qk_norm_after_rope,
q_norm: None,
k_norm: None,
output_gate: false,
softplus_gate: None,
rope_scale: self.rope_scale,
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.with_bos(self.tokenizer.encode(prompt));
self.generate_from_ids(&input_ids, max_tokens, task_mask, on_token)
}
pub fn generate_from_vl(
&mut self,
input: &crate::dsv41_vision::PreparedVlInputs,
max_tokens: usize,
task_mask: Option<&TaskMask>,
on_token: Option<TokenCallback>,
) -> Result<GenerateResult, String> {
let Some(dsv41) = &self.dsv41 else {
return Err("V4.1 multimodal input requires a DeepSeek-V4.1 pipeline".into());
};
if input.token_ids.is_empty() {
return Err("empty V4.1 multimodal prompt".into());
}
if input.token_types.len() != input.token_ids.len() {
return Err(format!(
"V4.1 token type count {} != token count {}",
input.token_types.len(),
input.token_ids.len()
));
}
let dim = dsv41.2.dim;
let mut embeddings = vec![None; input.token_ids.len()];
let mut participates = vec![true; input.token_ids.len()];
if !input.images.is_empty() {
let vision = self
.dsv41_vision
.as_ref()
.ok_or_else(|| "V4.1 image prompt has no loaded vision tower".to_string())?;
for image in &input.images {
let end = image.start.saturating_add(image.types.len());
if end > input.token_ids.len() {
return Err(format!(
"V4.1 image span {}..{} exceeds prompt length {}",
image.start,
end,
input.token_ids.len()
));
}
let mut span = vec![0.0f32; image.types.len() * dim];
vision.fill_image_span(image, &mut span, self.pool.as_deref())?;
for (offset, &kind) in image.types.iter().enumerate() {
let pos = image.start + offset;
if input.token_types[pos] != kind {
return Err(format!(
"V4.1 image type mismatch at position {pos}: {} != {kind}",
input.token_types[pos]
));
}
embeddings[pos] = Some(span[offset * dim..(offset + 1) * dim].to_vec());
participates[pos] = false;
}
}
}
for (pos, &kind) in input.token_types.iter().enumerate() {
if kind == crate::dsv41_vision::TEXT && embeddings[pos].is_some() {
return Err(format!("V4.1 text position {pos} has an image embedding"));
}
if kind != crate::dsv41_vision::TEXT && embeddings[pos].is_none() {
return Err(format!("V4.1 image position {pos} has no image embedding"));
}
}
self.dsv41_prefill = Some((embeddings, participates));
let result = self.generate_from_ids(&input.token_ids, max_tokens, task_mask, on_token);
self.dsv41_prefill = None;
result
}
fn drop_open_mask<'m>(&self, m: Option<&'m TaskMask>) -> Option<&'m TaskMask> {
m.filter(|m| !m.fully_open(self.intermediate_size, self.num_heads))
}
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 std::env::var("CMF_TRACE_H").is_ok() {
eprintln!("input_ids: {input_ids:?}");
}
if input_ids.is_empty() {
return Err("empty prompt: nothing to generate from".to_string());
}
self.graph_failed
.store(false, std::sync::atomic::Ordering::Relaxed);
let task_mask = self.drop_open_mask(task_mask);
let mut reuse_from = {
let on = !std::env::var("CMF_KV_REUSE").is_ok_and(|v| v == "0");
let h = &self.kv_history;
if on
&& task_mask.is_none()
&& self.mtp.is_none()
&& self.o1_cfg.is_none()
&& self.dsv41.is_none()
&& !h.is_empty()
&& h.len() < input_ids.len()
&& input_ids[..h.len()] == h[..]
{
h.len()
} else {
0
}
};
if reuse_from > 0 && !self.prepare_kv_reuse(reuse_from) {
reuse_from = 0;
}
if reuse_from == 0 {
self.clear_sequence_state();
} else if std::env::var("CMF_PREFILL_PROF").is_ok() {
eprintln!(
"kv-reuse: {} of {} prompt positions already cached",
reuse_from,
input_ids.len()
);
}
crate::gpu::graph_race_begin_generation();
let o1_prefill = if self.o1_active() && task_mask.is_none() {
std::env::var("CMF_O1_PREFILL")
.ok()
.and_then(|v| v.parse::<usize>().ok())
.filter(|&p| p > 0)
} else {
None
};
if task_mask.is_none() {
self.o1_begin_with_prefix(o1_prefill);
}
let graph_on = crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Decode);
#[cfg(target_os = "macos")]
let metal_graph = crate::gpu::q1_force()
&& crate::gpu::enabled_here()
&& std::env::var("CMF_GPU_BLOCK")
.map(|v| v != "0")
.unwrap_or(true);
#[cfg(not(target_os = "macos"))]
let metal_graph = false;
let spec_sample_env = std::env::var("CMF_GRAPH_SPEC_SAMPLE").ok();
let spec_cheap_round = self.sampler_config.temperature < 1e-6
|| sampler::sparse_ok(&self.sampler_config);
let spec_sampling_ok = self.sampler_config.temperature < 1e-6
|| match spec_sample_env.as_deref() {
Some("1") => true,
Some(_) => false,
None => metal_graph && spec_cheap_round,
};
let (mut dense_n, mut dense_q4tp) = (0usize, 0usize);
for lw in &self.weights.layers {
if let FfnKind::Dense(d) = &lw.ffn {
dense_n += 1;
if matches!(d.gate_proj.graph_weight(), Some((_, _, 6, _)))
&& matches!(d.up_proj.graph_weight(), Some((_, _, 6, _)))
&& matches!(d.down_proj.graph_weight(), Some((_, _, 6, _)))
{
dense_q4tp += 1;
}
}
}
let spec_default_ok = dense_n == 0 || dense_q4tp * 10 >= dense_n * 9;
let penalized = !metal_graph
&& (self.sampler_config.repetition_penalty != 1.0
|| self.sampler_config.presence_penalty != 0.0
|| !self.sampler_config.suppress_tokens.is_empty());
#[cfg(feature = "gpu")]
let metal_wgpu = graph_on && crate::gpu_wgpu::wgpu_backend_is_metal();
#[cfg(not(feature = "gpu"))]
let metal_wgpu = false;
let spec_env = std::env::var("CMF_GRAPH_SPEC").ok();
let spec_wanted = match spec_env.as_deref() {
Some("0") => false,
Some(_) => {
if metal_wgpu {
tracing::warn!(
"CMF_GRAPH_SPEC forced on wgpu/Metal: the batched verify graph is not \
verified on this backend (garbage measured on Qwen3.5-0.8B)"
);
}
true
}
None => spec_default_ok && !penalized && !metal_wgpu,
};
let graph_spec = self.speculative
&& (graph_on || metal_graph)
&& self.mtp.is_some()
&& task_mask.is_none()
&& !self.o1_active()
&& spec_sampling_ok
&& spec_wanted;
#[cfg(target_os = "macos")]
if metal_graph {
static SAID: std::sync::Once = std::sync::Once::new();
SAID.call_once(|| {
let spec = if graph_spec {
let k = std::env::var("CMF_GRAPH_SPEC_K")
.ok()
.and_then(|v| v.parse::<usize>().ok())
.filter(|&v| (1..=8).contains(&v))
.unwrap_or(7);
let arm = if self.sampler_config.temperature < 1e-6 {
"greedy"
} else {
"sampling"
};
format!(
"spec k={k} {arm} (batched verify, draft shortlist {}, trial: proxy)",
Self::draft_vocab_rows(usize::MAX)
)
} else if !self.speculative {
"spec off (CMF_MTP=0)".to_string()
} else if self.mtp.is_none() {
"spec off (no MTP head)".to_string()
} else if !spec_sampling_ok {
if spec_cheap_round {
"spec off (CMF_GRAPH_SPEC_SAMPLE=0)".to_string()
} else {
"spec off (sampling without a top-k: the dense chain \
costs more than it saves)"
.to_string()
}
} else if !spec_wanted {
"spec off (CMF_GRAPH_SPEC=0 or non-q4tp FFNs)".to_string()
} else if task_mask.is_some() {
"spec off (task mask)".to_string()
} else {
"spec off (O(1) attention)".to_string()
};
let on = |var: &str| {
if std::env::var(var).as_deref() == Ok("0") {
"off"
} else {
"on"
}
};
tracing::info!(
"metal native: {spec}, state4 {}, async replay {}, prefill graph {}, \
MTP graph {}, attend {}, probe {}",
if crate::gpu_metal::state4_on() { "on" } else { "off" },
if crate::gpu_metal::async_replay_on() { "on" } else { "off" },
on("CMF_METAL_PREFILL"),
on("CMF_MTP_GRAPH"),
std::env::var("CMF_GPU_ATTEND").unwrap_or_else(|_| "auto".into()),
if crate::gpu::probe_enabled() { "bypassed (q1 force)" } else { "off" },
);
});
}
let pair_pays = self.gdn_cfg.is_none() || std::env::var("CMF_MTP").as_deref() == Ok("1");
let spec_active = self.speculative
&& self.mtp.is_some()
&& task_mask.is_none()
&& !self.o1_active()
&& ((!graph_on && pair_pays && self.sampler_config.temperature < 1e-6) || graph_spec);
let mut mtp = if spec_active { self.mtp.take() } else { None };
if std::env::var("CMF_MTP_CHAIN_PROBE").is_ok() {
eprintln!(
"mtp-probe gate: spec_active={spec_active} mtp={} speculative={} graph_on={graph_on} temp_ok={}",
mtp.is_some(),
self.speculative,
self.sampler_config.temperature < 1e-6,
);
}
if let Some(m) = &mut mtp {
m.kv.clear();
crate::gpu::graph_kv_reset(self.mtp_kv_id());
self.mtp_graph_mode = None;
}
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 dsv4_spec_bad = 0usize;
let mut dsv4_spec_retry_at = 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 = reuse_from;
let fuse_lm = mtp.is_none()
&& router.is_none()
&& std::env::var("CMF_GPU_LMHEAD").as_deref() != Ok("0");
self.graph_logits = None;
self.graph_want_logits = false;
let _tpf = std::time::Instant::now();
let batch_k = self.generation_batch_k();
while self.qwen4_exp.is_some()
&& mtp.is_none()
&& pos < input_ids.len()
&& !self.cancel.load(std::sync::atomic::Ordering::Relaxed)
{
let token_id = input_ids[pos];
let want_logits = pos + 1 == input_ids.len();
let mut lg = Vec::new();
if let Some(b) = &mut self.qwen4_exp {
crate::qwen4_exp::forward_token(
&b.0,
&b.1,
&b.2,
&mut b.3,
token_id,
pos,
&self.inv_freq,
self.pool.as_deref(),
&mut lg,
want_logits,
);
}
if want_logits {
self.graph_logits = Some(lg);
}
pos += 1;
hidden.fill(0.0);
}
while self.dsv4.is_some()
&& mtp.is_none()
&& pos < input_ids.len()
&& !self.cancel.load(std::sync::atomic::Ordering::Relaxed)
{
let end = (pos + prefill_chunk()).min(input_ids.len());
let ids: Vec<u32> = input_ids[pos..end].to_vec();
let mut lg = Vec::new();
if let Some(b) = &mut self.dsv4 {
let (g, layers, cfg, st) = (&b.0, &b.1, b.2, &mut b.3);
crate::dsv4::forward_chunk(
g,
layers,
&cfg,
st,
&ids,
pos,
&self.inv_freq,
self.pool.as_deref(),
&mut lg,
end == input_ids.len(),
);
}
if end == input_ids.len() {
self.graph_logits = Some(lg);
}
pos = end;
hidden = vec![0.0; self.hidden_size];
}
let dsv41_prefill = self.dsv41_prefill.take();
while self.dsv41.is_some()
&& mtp.is_none()
&& pos < input_ids.len()
&& !self.cancel.load(std::sync::atomic::Ordering::Relaxed)
{
let end = (pos + prefill_chunk()).min(input_ids.len());
let ids: Vec<u32> = input_ids[pos..end].to_vec();
let mut lg = Vec::new();
if let Some(b) = &mut self.dsv41 {
let (g, layers, cfg, st) = (&b.0, &b.1, &b.2, &mut b.3);
if let Some((embeddings, participates)) = dsv41_prefill.as_ref() {
crate::dsv41::forward_chunk_masked_with_embeddings(
g,
layers,
cfg,
st,
&ids,
pos,
&embeddings[pos..end],
&participates[pos..end],
self.pool.as_deref(),
&mut lg,
);
} else {
crate::dsv41::forward_chunk(
g,
layers,
cfg,
st,
&ids,
pos,
self.pool.as_deref(),
&mut lg,
);
}
}
if end == input_ids.len() {
self.graph_logits = Some(lg);
}
pos = end;
hidden = vec![0.0; self.hidden_size];
}
let dyn_prefill = router.is_some();
let o1_prefill_limit = o1_prefill
.and_then(|requested| self.o1_effective_boundary(requested))
.map(|boundary| boundary.min(input_ids.len()));
let mut o1_sealed = false;
if let Some(limit) = o1_prefill_limit {
if self.can_prefill_batched() && limit > 2 {
let chunk = self.prefill_chunk();
let hs = self.hidden_size;
while pos < limit && !self.cancel.load(std::sync::atomic::Ordering::Relaxed) {
let end = (pos + chunk).min(limit);
let hb = self.prefill_batch(&input_ids[pos..end], pos);
hidden.copy_from_slice(&hb[(end - pos - 1) * hs..]);
pos = end;
}
} else {
while pos < limit && !self.cancel.load(std::sync::atomic::Ordering::Relaxed) {
hidden = self.forward_layers(&self.embed_single(input_ids[pos]), pos, None);
pos += 1;
}
}
if pos >= limit {
o1_sealed = match self.o1_seal_checked() {
Ok(sealed) => sealed,
Err(err) => {
self.finish_generation(&mut mtp, &mut router, true);
return Err(err);
}
};
tracing::info!(
"o1 bounded prompt prefix: requested={} effective={} processed={} of {} token(s)",
o1_prefill.unwrap_or(0),
self.o1_effective_boundary(o1_prefill.unwrap_or(0))
.unwrap_or(limit),
limit,
input_ids.len()
);
}
}
let graph_prefill = self.graph_prefill_preferred();
#[cfg(target_os = "macos")]
if task_mask.is_none()
&& !dyn_prefill
&& (crate::gpu::q1_force() || crate::gpu::q2tp_gpu_opt_in())
&& crate::gpu::enabled_here()
&& self.gdn_cfg.is_some()
&& self.g3n.is_none()
&& input_ids.len() > 8
&& std::env::var("CMF_MTP_CHAIN_PROBE").is_err()
&& std::env::var("CMF_METAL_PREFILL").as_deref() != Ok("0")
{
let chunk: usize = std::env::var("CMF_METAL_PREFILL_CHUNK")
.ok()
.and_then(|v| v.parse().ok())
.filter(|&v| (16..=512).contains(&v))
.unwrap_or(256);
let hs = self.hidden_size;
let _tp = std::time::Instant::now();
while pos < input_ids.len() && !self.cancel.load(std::sync::atomic::Ordering::Relaxed) {
let end = (pos + chunk).min(input_ids.len());
let hb = match self.prefill_batch_metal(&input_ids[pos..end], pos) {
MetalPrefillOutcome::Completed(hb) => hb,
MetalPrefillOutcome::Declined => break,
MetalPrefillOutcome::Failed => {
self.finish_generation(&mut mtp, &mut router, true);
return Err("ordinary Metal prefill failed after admission".into());
}
};
if let Some(m) = &mut mtp {
let n_pairs = if end < input_ids.len() {
end - pos
} else {
end - pos - 1
};
if n_pairs > 0 {
let pairs: Vec<(&[f32], u32)> = (0..n_pairs)
.map(|j| (&hb[j * hs..(j + 1) * hs], input_ids[pos + j + 1]))
.collect();
if !self.mtp_warm_batch_metal(m, &pairs, pos) {
for (j, (h, t)) in pairs.iter().enumerate() {
let h = h.to_vec();
let _ = self.mtp_step(m, &h, *t, pos + j);
}
}
}
}
hidden.copy_from_slice(&hb[(end - pos - 1) * hs..]);
pos = end;
}
if std::env::var("CMF_PREFILL_PROF").is_ok() {
eprintln!(
"metal-prefill: {} of {} tokens in {:.1} ms",
pos,
input_ids.len(),
_tp.elapsed().as_secs_f64() * 1e3
);
}
}
if task_mask.is_none()
&& !dyn_prefill
&& !graph_prefill
&& self.can_prefill_batched()
&& self.g3n.is_none()
&& o1_prefill.is_none()
&& input_ids.len() > 2
{
let chunk = self.prefill_chunk();
let hs = self.hidden_size;
while pos < input_ids.len() && !self.cancel.load(std::sync::atomic::Ordering::Relaxed) {
let end = (pos + chunk).min(input_ids.len());
let hb = self.prefill_batch(&input_ids[pos..end], pos);
if let Some(m) = &mut mtp {
let probe: usize = std::env::var("CMF_MTP_CHAIN_PROBE")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(0);
for p in pos..end {
if p + 1 < input_ids.len() {
if probe >= 1 && p + 2 < input_ids.len() {
let (d1, mut hx) = self.mtp_step_h(
m,
&hb[(p - pos) * hs..(p - pos + 1) * hs],
input_ids[p + 1],
p,
);
let mut ok = d1 == input_ids[p + 2];
Self::chain_probe_note(0, ok);
let mut d_prev = d1;
let mut extra = 0usize;
for j in 1..probe {
if p + 2 + j >= input_ids.len() {
break;
}
let (dj, hj) = self.mtp_step_h(m, &hx, d_prev, p + 1 + j);
extra += 1;
ok = ok && dj == input_ids[p + 2 + j];
Self::chain_probe_note(j, ok);
d_prev = dj;
hx = hj;
}
m.kv.truncate_last(extra);
} else {
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;
}
}
let pair_off = std::env::var("CMF_PAIR").is_ok_and(|v| v == "0");
if task_mask.is_none()
&& !dyn_prefill
&& !graph_prefill
&& !pair_off
&& self.pair_supported()
&& o1_prefill.is_none()
{
while pos + 1 < input_ids.len()
&& !self.cancel.load(std::sync::atomic::Ordering::Relaxed)
{
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 probe: usize = std::env::var("CMF_MTP_CHAIN_PROBE")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(0);
if probe >= 1 && pos + 3 < input_ids.len() {
let (d1, mut hx) = self.mtp_step_h(m, &h2, input_ids[pos + 2], pos + 1);
let mut ok = d1 == input_ids[pos + 3];
Self::chain_probe_note(0, ok);
let mut d_prev = d1;
let mut extra = 0usize;
for j in 1..probe {
if pos + 3 + j >= input_ids.len() {
break;
}
let (dj, hj) = self.mtp_step_h(m, &hx, d_prev, pos + 2 + j);
extra += 1;
ok = ok && dj == input_ids[pos + 3 + j];
Self::chain_probe_note(j, ok);
d_prev = dj;
hx = hj;
}
m.kv.truncate_last(extra);
} else {
let _ = self.mtp_step(m, &h2, input_ids[pos + 2], pos + 1);
}
}
}
hidden = h2;
pos += 2;
}
}
let o1_batch_ready = o1_sealed
&& o1_prefill.is_some()
&& mtp.is_none()
&& std::env::var("CMF_O1_GPU").as_deref() == Ok("1")
&& (0..self.num_layers).all(|li| {
let cache = &self.kv_cache.layers[self.phys_layer(li)];
cache.o1.is_none() || cache.o1_views().is_some()
});
let mtp_batch_prefill = mtp.is_some()
&& graph_prefill
&& task_mask.is_none()
&& !dyn_prefill
&& !self.o1_active()
&& std::env::var("CMF_MTP_CHAIN_PROBE").is_err();
if batch_k > 0
&& (graph_prefill || o1_batch_ready)
&& task_mask.is_none()
&& (!self.o1_active() || o1_batch_ready)
&& (mtp.is_none() || mtp_batch_prefill)
&& !dyn_prefill
&& pos + 1 < input_ids.len()
{
let hs = self.hidden_size;
let chunk = batch_k;
while pos < input_ids.len() {
let end = (pos + chunk).min(input_ids.len());
let bk = end - pos;
let mut hiddens = vec![0f32; bk * hs];
for (j, &id) in input_ids[pos..end].iter().enumerate() {
hiddens[j * hs..(j + 1) * hs].copy_from_slice(&self.embed_single(id));
}
let positions: Vec<usize> = (pos..end).collect();
let t_chunk = std::time::Instant::now();
let outcome = self.try_batch_graph_wgpu(&mut hiddens, &positions, bk, None);
let ok_b = outcome == crate::gpu::BatchGraphOutcome::Completed;
if std::env::var("CMF_GRAPH_PROF").is_ok() {
let ms = t_chunk.elapsed().as_secs_f64() * 1000.0;
eprintln!(
"batch-chunk: phase=prompt mode={} k={bk} outcome={outcome:?} {ms:.1} ms ({:.1} tok/s)",
if o1_batch_ready {
"o1"
} else if mtp_batch_prefill {
"ordinary_mtp"
} else {
"ordinary"
},
bk as f64 / (ms / 1000.0)
);
}
{
use std::sync::atomic::{AtomicBool, Ordering};
static SAID: AtomicBool = AtomicBool::new(false);
if !SAID.swap(true, Ordering::Relaxed) {
if ok_b {
tracing::info!(
"batched prefill: ACTIVE mode={} (k={bk})",
if o1_batch_ready {
"o1"
} else if mtp_batch_prefill {
"ordinary_mtp"
} else {
"ordinary"
}
);
} else {
tracing::warn!("batched prefill {:?} — per-position graph", outcome);
}
}
}
if ok_b {
if mtp_batch_prefill {
let n_pairs = mtp_prefill_pair_count(pos, end, input_ids.len());
if n_pairs > 0 {
let rows: Vec<Vec<f32>> = (0..n_pairs)
.map(|j| hiddens[j * hs..(j + 1) * hs].to_vec())
.collect();
let pairs: Vec<(&[f32], u32)> = rows
.iter()
.enumerate()
.map(|(j, row)| (row.as_slice(), input_ids[pos + j + 1]))
.collect();
if std::env::var("CMF_GRAPH_PROF").is_ok() {
eprintln!(
"mtp-warm: phase=prompt mode=ordinary_mtp first_pos={} pairs={} last_pos={}",
pos,
n_pairs,
pos + n_pairs - 1,
);
}
let warm_error = if let Some(m) = mtp.as_mut() {
self.mtp_warm_prefill_pairs(m, &pairs, pos).err()
} else {
None
};
if let Some(err) = warm_error {
self.finish_generation(&mut mtp, &mut router, true);
return Err(err.to_string());
}
}
}
hidden.copy_from_slice(&hiddens[(bk - 1) * hs..]);
pos = end;
} else if outcome == crate::gpu::BatchGraphOutcome::Failed {
self.finish_generation(&mut mtp, &mut router, true);
return Err(if o1_batch_ready {
"sealed O(1) batch graph failed after admission".to_string()
} else {
"ordinary recurrent batch graph failed after admission".to_string()
});
} else {
break; }
}
}
while pos < input_ids.len() && !self.cancel.load(std::sync::atomic::Ordering::Relaxed) {
self.graph_want_logits = fuse_lm && pos + 1 == 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 probe: usize = std::env::var("CMF_MTP_CHAIN_PROBE")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(0);
if probe >= 1 && pos + 2 < input_ids.len() {
let (d1, mut hx) = self.mtp_step_h(m, &hidden, input_ids[pos + 1], pos);
let mut ok = d1 == input_ids[pos + 2];
Self::chain_probe_note(0, ok);
let mut d_prev = d1;
let mut extra = 0usize;
for j in 1..probe {
if pos + 2 + j >= input_ids.len() {
break;
}
let (dj, hj) = self.mtp_step_h(m, &hx, d_prev, pos + 1 + j);
extra += 1;
ok = ok && dj == input_ids[pos + 2 + j];
Self::chain_probe_note(j, ok);
d_prev = dj;
hx = hj;
}
m.kv.truncate_last(extra);
} else {
let _ = self.mtp_step(m, &hidden, input_ids[pos + 1], pos);
}
}
}
pos += 1;
}
if std::env::var("CMF_PREFILL_PROF").is_ok() {
eprintln!(
"prefill: {} tokens in {:.1} ms (batch_k={batch_k})",
input_ids.len(),
_tpf.elapsed().as_secs_f64() * 1000.0
);
}
if self
.graph_failed
.swap(false, std::sync::atomic::Ordering::Relaxed)
{
self.finish_generation(&mut mtp, &mut router, true);
return Err("GPU token graph failed during prefill".to_string());
}
if self
.cancel
.swap(false, std::sync::atomic::Ordering::Relaxed)
{
self.finish_generation(&mut mtp, &mut router, true);
return Ok(GenerateResult {
text: String::new(),
token_ids: Vec::new(),
prompt_tokens: input_ids.len(),
tokens_generated: 0,
finish_reason: "cancelled".to_string(),
mtp_drafted: 0,
mtp_accepted: 0,
token_confidence: Vec::new(),
traces: Vec::new(),
});
}
if !o1_sealed {
match self.o1_seal_checked() {
Ok(_) => {}
Err(err) => {
self.finish_generation(&mut mtp, &mut router, true);
return Err(err);
}
}
}
macro_rules! commit {
($id:expr) => {{
all_ids.push($id);
generated += 1;
self.note_draft_id($id);
if self.tokenizer.is_eos($id) && !self.ignore_eos {
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 spec_trial = SpecTrial::Spec {
t0: std::time::Instant::now(),
gen0: generated,
rounds: 0,
};
let mut spec_mon = SpecMon {
metal: graph_spec && crate::gpu::q1_force() && spec_cheap_round,
..SpecMon::default()
};
let mut spec_watchdog_off = false;
let mut spec_walls: Vec<f32> = Vec::new();
let mut spec_round_end: Option<std::time::Instant> = None;
let mut next_pos = input_ids.len();
'decode: while generated < max_tokens {
if self
.graph_failed
.swap(false, std::sync::atomic::Ordering::Relaxed)
{
self.finish_generation(&mut mtp, &mut router, true);
return Err("GPU token graph failed during decode".to_string());
}
if self
.cancel
.swap(false, std::sync::atomic::Ordering::Relaxed)
{
finish_reason = "cancelled".to_string();
break 'decode;
}
let forced = self.spec_forced.take();
let mut logits = match (forced, self.graph_logits.take()) {
(Some(_), _) => Vec::new(),
(None, Some(lg)) => lg,
(None, None) => {
let _prof = crate::cpuprof::time(crate::cpuprof::Slot::Head);
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)
}
};
if generated
== std::env::var("CMF_LOGIT_DUMP_STEP")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(0)
{
if let Ok(path) = std::env::var("CMF_LOGIT_DUMP") {
let mut bytes: Vec<u8> = Vec::with_capacity((hidden.len() + logits.len()) * 4);
for v in hidden.iter().chain(logits.iter()) {
bytes.extend_from_slice(&v.to_le_bytes());
}
if let Err(e) = std::fs::write(&path, &bytes) {
eprintln!("logit dump: failed to write {path}: {e}");
self.finish_generation(&mut mtp, &mut router, true);
return Err(format!("logit dump write failed: {e}"));
}
}
}
let t_next = match forced {
Some(c) => c,
None => {
let _prof = crate::cpuprof::time(crate::cpuprof::Slot::Sampler);
sampler::sample_with_scratch_pool(
&logits,
&self.sampler_config,
&all_ids,
&mut self.rng,
&mut self.sampler_scratch,
self.pool.as_deref(),
)
}
};
if self.confidence_on {
confidence.push(if logits.is_empty() {
0.0
} else {
sampler::top1_prob_pool(
self.pool.as_deref(),
&mut self.sampler_scratch,
&logits,
t_next,
calib_temp,
)
});
}
if !logits.is_empty() {
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().copied().unwrap_or(0.0),
active_skill: skill,
recon: None,
switched: false,
});
}
if !commit!(t_next) {
break 'decode;
}
if generated >= max_tokens {
break 'decode;
}
if self.dsv41.is_none() && self.kv_cache.needs_eviction() {
static SAID: std::sync::Once = std::sync::Once::new();
SAID.call_once(|| {
tracing::warn!(
"KV cache full at {} positions — evicting half; quality \
will degrade. Raise CMF_MAX_SEQ.",
self.kv_cache.max_seq_len,
);
});
let keep = (self.kv_cache.max_seq_len / 2).max(1);
self.kv_cache.evict(keep);
}
if graph_spec {
match spec_trial {
SpecTrial::Plain { t0, gen0 } if spec_mon.plain_done(t0, gen0, generated) => {
spec_mon.plain_ms =
t0.elapsed().as_secs_f64() * 1e3 / (generated - gen0) as f64;
let keep = spec_mon.pays();
tracing::info!(
"speculation trial: {:.2} tok/round in {:.1} ms vs plain {:.1} ms/tok — {}",
spec_mon.tokens,
spec_mon.round_ms,
spec_mon.plain_ms,
if keep { "speculating" } else { "plain" }
);
spec_mon.fails = 0;
spec_trial = SpecTrial::Decided {
spec: keep,
recheck_at: if keep { usize::MAX } else { generated + 128 },
};
}
SpecTrial::Decided { recheck_at, .. } if generated >= recheck_at => {
spec_mon.n = 0;
spec_trial = SpecTrial::Spec {
t0: std::time::Instant::now(),
gen0: generated,
rounds: 0,
};
}
_ => {}
}
spec_watchdog_off = matches!(
spec_trial,
SpecTrial::Plain { .. } | SpecTrial::Decided { spec: false, .. }
);
}
match &mut mtp {
#[cfg(feature = "gpu")]
Some(m)
if graph_spec
&& !spec_watchdog_off
&& generated + 1 < max_tokens
&& next_pos > 0 =>
{
let t_round = std::time::Instant::now();
if spec_time_level() >= 2 {
if let Some(t) = spec_round_end.take() {
eprintln!(
"spec-gap {:.2} ms (host between rounds)",
t.elapsed().as_secs_f64() * 1e3
);
}
}
spec_stamps_begin();
#[cfg(target_os = "macos")]
let allocs0 = crate::gpu_metal::IO_BUF_ALLOCS
.load(std::sync::atomic::Ordering::Relaxed);
#[cfg(not(target_os = "macos"))]
let allocs0 = 0u64;
if let Some((extra, n_pos, new_h)) = self.graph_spec_step(
m,
&hidden,
t_next,
next_pos,
&mut drafted,
&mut accepted,
&mut all_ids,
max_tokens - generated,
) {
next_pos = n_pos;
hidden = new_h;
let level = spec_time_level();
if level > 0 {
let wall = t_round.elapsed().as_secs_f32() * 1e3;
let stamps = spec_stamps_take();
let median = if spec_walls.len() >= 3 {
let mut s = spec_walls.clone();
s.sort_by(|a, b| a.partial_cmp(b).unwrap());
Some(s[s.len() / 2])
} else {
None
};
let outlier = median.is_some_and(|m| wall > 1.4 * m);
#[cfg(target_os = "macos")]
let allocs = crate::gpu_metal::IO_BUF_ALLOCS
.load(std::sync::atomic::Ordering::Relaxed)
- allocs0;
#[cfg(not(target_os = "macos"))]
let allocs = allocs0;
eprintln!(
"spec-round wall {wall:.1} ms → {} tokens{}{}",
extra.len() + 1,
if allocs > 0 {
format!(" [{allocs} new device buffers]")
} else {
String::new()
},
match (outlier, median) {
(true, Some(m)) => format!(" OUTLIER (median {m:.1})"),
_ => String::new(),
}
);
if level >= 2 || outlier {
let sum: f32 = stamps.iter().map(|s| s.1).sum();
eprintln!(
"spec-stamps: {}| untracked {:.1}",
spec_stamps_format(&stamps),
wall - sum
);
}
if spec_mon.n >= 1 {
spec_walls.push(wall);
}
}
spec_mon.round(t_round.elapsed().as_secs_f64() * 1e3, extra.len() + 1);
spec_trial = Self::spec_trial_round(
spec_trial,
&mut spec_mon,
generated + extra.len() + 1,
);
let mut stopped = false;
for &id in &extra {
if self.confidence_on {
confidence.push(0.0);
}
if !commit!(id) {
stopped = true;
break;
}
}
if stopped {
break 'decode;
}
if spec_time_level() >= 2 {
spec_round_end = Some(std::time::Instant::now());
}
continue 'decode;
}
if self
.graph_failed
.swap(false, std::sync::atomic::Ordering::Relaxed)
{
self.finish_generation(&mut mtp, &mut router, true);
return Err("GPU MTP graph failed during speculative decode".to_string());
}
spec_mon.round(t_round.elapsed().as_secs_f64() * 1e3, 1);
spec_mon.tokens = 0.0;
spec_mon.fails = 3;
spec_trial = Self::spec_trial_round(spec_trial, &mut spec_mon, generated + 1);
hidden = self.forward_layers(&self.embed_single(t_next), next_pos, task_mask);
next_pos += 1;
continue 'decode;
}
Some(m) if !graph_spec && 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_with_scratch_pool(
&logits1,
&self.sampler_config,
&all_ids,
&mut self.rng,
&mut self.sampler_scratch,
self.pool.as_deref(),
);
if self.confidence_on {
confidence.push(sampler::top1_prob_pool(
self.pool.as_deref(),
&mut self.sampler_scratch,
&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().copied().unwrap_or(0.0),
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;
}
}
_ => {
#[cfg(feature = "gpu")]
if Self::dsv4_spec_on() && self.dsv4.is_some() {
static SAID: std::sync::Once = std::sync::Once::new();
SAID.call_once(|| {
eprintln!(
"dsv4-spec гейт: mtp={} mask={} router={} trace={} temp={} rep={} ",
!self.dsv4_mtp.is_empty(),
task_mask.is_none(),
router.is_none(),
!trace_on,
self.sampler_config.temperature < 1e-6,
self.sampler_config.repetition_penalty == 1.0,
);
});
}
#[cfg(feature = "gpu")]
if Self::dsv4_spec_on()
&& self.dsv4.is_some()
&& !self.dsv4_mtp.is_empty()
&& task_mask.is_none()
&& router.is_none()
&& !trace_on
&& self.sampler_config.temperature < 1e-6
&& self.sampler_config.repetition_penalty == 1.0
&& generated + 1 < max_tokens
&& all_ids.len() >= 2
&& generated >= dsv4_spec_retry_at
{
let tip_token = all_ids[all_ids.len() - 2];
let drafted0 = drafted;
let round = self.dsv4_spec_step(
tip_token,
t_next,
next_pos,
max_tokens.saturating_sub(generated),
&mut drafted,
&mut accepted,
);
if drafted > drafted0 {
let useful = round.as_ref().is_some_and(|(extra, _)| !extra.is_empty());
if useful {
dsv4_spec_bad = 0;
} else {
dsv4_spec_bad += 1;
if dsv4_spec_bad >= 2 {
dsv4_spec_bad = 0;
dsv4_spec_retry_at = generated.saturating_add(32);
tracing::info!(
"dsv4: draft не окупился дважды — точный walk на 32 токена"
);
}
}
}
if let Some((extra, n_pos)) = round {
next_pos = n_pos;
let mut stopped = false;
for &id in &extra {
if self.confidence_on {
confidence.push(0.0);
}
if !commit!(id) {
stopped = true;
break;
}
}
if stopped {
break 'decode;
}
continue 'decode;
}
}
self.graph_want_logits = fuse_lm;
let mut t_fwd = t_next;
let pure_greedy = self.sampler_config.temperature < 1e-6
&& self.sampler_config.repetition_penalty == 1.0
&& self.sampler_config.suppress_tokens.is_empty();
let burst_k = std::env::var("CMF_MULTISTEP")
.ok()
.and_then(|v| v.parse::<usize>().ok())
.unwrap_or(0);
if pure_greedy
&& burst_k >= 1
&& fuse_lm
&& task_mask.is_none()
&& router.is_none()
&& !trace_on
&& !self.confidence_on
{
let mut stopped = false;
loop {
let room = max_tokens.saturating_sub(generated);
if room <= 2 {
break;
}
let k = burst_k.min(room - 1);
if k < 1 {
break;
}
let Some(ids) = self.try_multi_burst(t_fwd, next_pos, k) else {
if self
.graph_failed
.swap(false, std::sync::atomic::Ordering::Relaxed)
{
self.finish_generation(&mut mtp, &mut router, true);
return Err(
"GPU token graph failed during greedy burst".to_string()
);
}
break;
};
next_pos += k;
for &id in &ids {
if !commit!(id) {
stopped = true;
break;
}
}
if stopped {
break;
}
t_fwd = *ids.last().unwrap();
}
if stopped {
break 'decode;
}
}
#[cfg(target_os = "macos")]
if graph_spec
&& spec_watchdog_off
&& next_pos > 0
&& self.mtp_graph_mode == Some(true)
&& crate::gpu::q1_force()
{
if let Some(m) = mtp.as_mut() {
let _ = self.mtp_step_metal(m, &hidden, t_fwd, next_pos - 1, false);
}
}
hidden = self.forward_layers(&self.embed_single(t_fwd), 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();
}
}
}
}
}
}
let cancelled = finish_reason == "cancelled";
self.finish_generation(&mut mtp, &mut router, cancelled);
let output_ids = &all_ids[input_ids.len()..];
let forwarded = input_ids.len() + output_ids.len().saturating_sub(1);
if cancelled {
self.kv_history.clear();
} else {
self.kv_history = all_ids[..forwarded.min(all_ids.len())].to_vec();
}
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 {
self.mtp_step_h(m, hidden, next_token, position).0
}
fn chain_probe_note(depth: usize, prefix_ok: bool) {
use std::sync::Mutex;
static T: Mutex<Vec<(u64, u64)>> = Mutex::new(Vec::new());
let mut t = T.lock().unwrap();
if t.len() <= depth {
t.resize(depth + 1, (0, 0));
}
t[depth].0 += 1;
t[depth].1 += prefix_ok as u64;
if depth == 0 && t[0].0 % 128 == 0 {
let line: Vec<String> = t
.iter()
.enumerate()
.map(|(d, (n, k))| {
format!(
"d{}={:.0}%({n})",
d + 1,
100.0 * *k as f64 / (*n).max(1) as f64
)
})
.collect();
eprintln!("mtp-chain: {}", line.join(" "));
}
}
fn mtp_step_hl(
&mut self,
m: &mut MtpModule,
hidden: &[f32],
next_token: u32,
position: usize,
) -> (Vec<f32>, Vec<f32>) {
#[cfg(target_os = "macos")]
if self.mtp_graph_mode != Some(false) && crate::gpu::q1_force() {
if let Some(r) = self.mtp_step_metal(m, hidden, next_token, position, true) {
self.mtp_graph_mode = Some(true);
return r;
}
if self.mtp_graph_mode == Some(true) {
tracing::error!("mtp Metal graph failed after admission");
self.clear_sequence_state();
self.graph_failed
.store(true, std::sync::atomic::Ordering::Relaxed);
self.cancel
.store(true, std::sync::atomic::Ordering::Relaxed);
return (Vec::new(), Vec::new());
}
self.mtp_graph_mode = Some(false);
}
#[cfg(feature = "gpu")]
if self.mtp_graph_mode != Some(false) {
if !self.mtp_graph_ok(m) {
if self.mtp_graph_mode == Some(true) {
tracing::error!("mtp graph became unavailable after admission");
self.clear_sequence_state();
self.graph_failed
.store(true, std::sync::atomic::Ordering::Relaxed);
self.cancel
.store(true, std::sync::atomic::Ordering::Relaxed);
return (Vec::new(), Vec::new());
}
self.mtp_graph_mode = Some(false);
} else {
if let Some(r) = self.mtp_step_graph(m, hidden, next_token, position) {
self.mtp_graph_mode = Some(true);
return r;
}
if self.graph_failed.load(std::sync::atomic::Ordering::Relaxed) {
return (Vec::new(), Vec::new());
}
tracing::error!("mtp graph failed or declined after admission");
self.clear_sequence_state();
self.graph_failed
.store(true, std::sync::atomic::Ordering::Relaxed);
self.cancel
.store(true, std::sync::atomic::Ordering::Relaxed);
return (Vec::new(), Vec::new());
}
}
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::Mla(_) => unreachable!("MLA has no MTP/pair path"),
AttnKind::Kda(_) => unreachable!("KDA has no MTP/pair path"),
AttnKind::Full {
wq,
wk,
wv,
wo,
q_norm,
k_norm,
output_gate,
softplus_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.softplus_gate = softplus_gate
.as_ref()
.map(|(gate, per_head)| (gate, *per_head));
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(_) | AttnKind::ShortConv(_) => {
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(), None);
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 lg = self.lm_head_forward(&self.ws.n1);
(lg, x)
}
fn mtp_step_h(
&mut self,
m: &mut MtpModule,
hidden: &[f32],
next_token: u32,
position: usize,
) -> (u32, Vec<f32>) {
let (mut lg, x) = self.mtp_step_hl(m, hidden, next_token, position);
let draft = sampler::argmax(&lg);
attention::recycle_buf(&mut lg);
(draft, x)
}
fn spec_trial_round(trial: SpecTrial, mon: &mut SpecMon, generated: usize) -> SpecTrial {
match trial {
SpecTrial::Spec { t0, gen0, rounds } => {
let rounds = rounds + 1;
if rounds >= 5 {
if mon.plain_ms > 0.0 {
let keep = mon.pays();
mon.fails = 0;
tracing::info!(
"speculation re-check: {:.2} tok/round in {:.1} ms vs plain {:.1} ms/tok — {}",
mon.tokens,
mon.round_ms,
mon.plain_ms,
if keep { "speculating" } else { "plain" }
);
SpecTrial::Decided {
spec: keep,
recheck_at: if keep { usize::MAX } else { generated + 128 },
}
} else if mon.pays() {
mon.fails = 0;
tracing::info!(
"speculation trial: {:.2} tok/round in {:.1} ms — speculating (plain not timed)",
mon.tokens,
mon.round_ms,
);
SpecTrial::Decided {
spec: true,
recheck_at: usize::MAX,
}
} else {
SpecTrial::Plain {
t0: std::time::Instant::now(),
gen0: generated,
}
}
} else {
SpecTrial::Spec { t0, gen0, rounds }
}
}
SpecTrial::Decided { spec: true, .. } => {
if mon.pays() {
mon.fails = 0;
trial
} else {
mon.fails += 1;
if mon.fails >= 4 {
if mon.plain_ms <= 0.0 {
tracing::info!(
"speculation doubtful: {:.2} tok/round in {:.1} ms — timing plain",
mon.tokens,
mon.round_ms,
);
return SpecTrial::Plain {
t0: std::time::Instant::now(),
gen0: generated,
};
}
tracing::info!(
"speculation stopped: {:.2} tok/round in {:.1} ms vs plain {:.1} ms/tok",
mon.tokens,
mon.round_ms,
mon.plain_ms
);
SpecTrial::Decided {
spec: false,
recheck_at: generated + 128,
}
} else {
trial
}
}
}
other => other,
}
}
fn mtp_kv_id(&self) -> u64 {
self.graph_kv_id | (1u64 << 40)
}
const MTP_LAYER_BASE: usize = 0;
#[cfg(feature = "gpu")]
fn rewind_mtp_graph_mirror(&self, stored: usize) -> bool {
self.mtp_graph_mode != Some(true)
|| crate::gpu::graph_kv_set_stored(self.mtp_kv_id(), Self::MTP_LAYER_BASE, stored)
}
#[cfg(feature = "gpu")]
fn rewind_trunk_graph_mirrors(&self, stored: usize) -> bool {
let mut ok = true;
let mut expected = false;
for li in 0..self.num_layers {
if matches!(
self.weights.layers[self.phys_layer(li)].attn,
AttnKind::Full { .. }
) {
expected = true;
ok &= crate::gpu::graph_kv_set_stored(self.graph_kv_id, li, stored);
}
}
!expected || ok
}
fn graph_gdn_layer_count(&self) -> usize {
(0..self.num_layers)
.filter(|&li| {
matches!(
&self.weights.layers[self.phys_layer(li)].attn,
AttnKind::LinearGdn(_)
)
})
.count()
}
fn mtp_block_input(&mut self, m: &MtpModule, hidden: &[f32], next_token: u32) -> Vec<f32> {
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());
x
}
#[cfg(feature = "gpu")]
fn mtp_block_graph_ok(&self, m: &MtpModule) -> bool {
if std::env::var("CMF_MTP_GRAPH").as_deref() == Ok("0") {
return false;
}
if !crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Decode)
|| !crate::gpu::enabled_here()
|| self.attn_softcap > 0.0
|| self.attention_heads_per_layer.is_some()
{
return false;
}
matches!(
&m.layer.attn,
AttnKind::Full {
softplus_gate: None,
..
}
) && matches!(&m.layer.ffn, FfnKind::Dense(_))
}
#[cfg(feature = "gpu")]
fn mtp_graph_ok(&self, m: &MtpModule) -> bool {
if !self.mtp_block_graph_ok(m) {
return false;
}
let AttnKind::Full { wq, wk, wv, wo, .. } = &m.layer.attn else {
return false;
};
let FfnKind::Dense(d) = &m.layer.ffn else {
return false;
};
d.segs.is_empty()
&& wq.graph_weight().is_some()
&& wk.graph_weight().is_some()
&& wv.graph_weight().is_some()
&& wo.graph_weight().is_some()
&& d.gate_proj.graph_weight().is_some()
&& d.up_proj.graph_weight().is_some()
&& d.down_proj.graph_weight().is_some()
&& self.weights.lm_head.graph_weight().is_some()
}
#[cfg(feature = "gpu")]
fn mtp_step_graph(
&mut self,
m: &mut MtpModule,
hidden: &[f32],
next_token: u32,
position: usize,
) -> Option<(Vec<f32>, Vec<f32>)> {
if !self.mtp_graph_ok(m) {
return None;
}
let lw = &m.layer;
let AttnKind::Full {
wq,
wk,
wv,
wo,
q_norm,
k_norm,
output_gate,
softplus_gate,
bias,
} = &lw.attn
else {
return None;
};
if softplus_gate.is_some() {
return None;
}
let FfnKind::Dense(d) = &lw.ffn else {
return None;
};
if !d.segs.is_empty() {
return None; }
let mut x = self.mtp_block_input(m, hidden, next_token);
fn gw(t: &QTensor) -> Option<crate::gpu::GraphW<'_>> {
let (_, i, kind, rs) = t.graph_weight()?;
Some(crate::gpu::GraphW {
idx: i,
kind,
row_scale: rs,
data: &[],
prism: crate::gpu::GraphPrismOp::None,
affine: false,
})
}
let (model, _, _, _) = wq.graph_weight()?;
let model = model.clone();
let (lm_gw, lm_rows) = {
let (_, i, kind, rs) = self.weights.lm_head.graph_weight()?;
let rows = if kind == 6 {
self.draft_head_rows(self.weights.lm_head.rows())
} else {
self.weights.lm_head.rows()
};
(
crate::gpu::GraphW {
idx: i,
kind,
row_scale: rs,
data: &[],
prism: crate::gpu::GraphPrismOp::None,
affine: false,
},
rows,
)
};
let layer = crate::gpu::GraphLayer {
input_norm: &lw.input_norm,
attn: crate::gpu::GraphAttn::Full {
wq: gw(wq)?,
wk: gw(wk)?,
wv: gw(wv)?,
wo: gw(wo)?,
q_norm: q_norm.as_deref(),
k_norm: k_norm.as_deref(),
late_qk_norm: self.qk_norm_after_rope,
bias: bias
.as_ref()
.map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
output_gate: *output_gate,
cpu_k: m.kv.k_heads(),
cpu_v: m.kv.v_heads(),
},
post_norm: &lw.post_norm,
ffn: crate::gpu::GraphFfn::Dense {
gate: gw(&d.gate_proj)?,
up: gw(&d.up_proj)?,
down: gw(&d.down_proj)?,
},
};
let nh = self.num_heads;
let (nkv, hd, rd) = self.layer_geom(0);
let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
let mut logits = Vec::new();
let ok = crate::gpu::forward_token_graph(
&model,
self.mtp_kv_id(),
std::slice::from_ref(&layer),
&[None],
self.o1_epoch,
&self.inv_freq,
&mut x,
nh,
nkv,
hd,
self.attn_scale,
rd,
self.hidden_size,
self.intermediate_size,
position,
self.kv_cache.max_seq_len,
gemma,
self.rms_eps as f32,
Some((&lm_gw, lm_rows)),
&m.final_norm,
&mut logits,
&[],
1,
None,
None,
None,
Self::MTP_LAYER_BASE,
true,
);
match ok {
crate::gpu::TokenGraphOutcome::Completed => {}
crate::gpu::TokenGraphOutcome::Declined => return None,
crate::gpu::TokenGraphOutcome::Failed => {
self.clear_sequence_state();
self.graph_failed
.store(true, std::sync::atomic::Ordering::Relaxed);
self.cancel
.store(true, std::sync::atomic::Ordering::Relaxed);
return None;
}
}
logits.resize(self.vocab_size, 0.0);
Some((logits, x))
}
#[cfg(feature = "gpu")]
fn mtp_warm_graph(
&mut self,
m: &mut MtpModule,
pairs: &[(&[f32], u32)],
first_pos: usize,
) -> crate::gpu::BatchGraphOutcome {
if pairs.is_empty() {
return crate::gpu::BatchGraphOutcome::Completed;
}
if !self.mtp_block_graph_ok(m) {
return crate::gpu::BatchGraphOutcome::Declined;
}
let hs = self.hidden_size;
let mut hiddens = Vec::with_capacity(pairs.len() * hs);
for (h, t) in pairs {
hiddens.extend_from_slice(&self.mtp_block_input(m, h, *t));
}
let lw = &m.layer;
let AttnKind::Full {
wq,
wk,
wv,
wo,
q_norm,
k_norm,
output_gate,
bias,
..
} = &lw.attn
else {
return crate::gpu::BatchGraphOutcome::Declined;
};
let FfnKind::Dense(d) = &lw.ffn else {
return crate::gpu::BatchGraphOutcome::Declined;
};
if !d.segs.is_empty() {
return crate::gpu::BatchGraphOutcome::Declined; }
fn gw(t: &QTensor) -> Option<crate::gpu::GraphW<'_>> {
let (_, i, kind, rs) = t.graph_weight()?;
Some(crate::gpu::GraphW {
idx: i,
kind,
row_scale: rs,
data: &[],
prism: crate::gpu::GraphPrismOp::None,
affine: false,
})
}
let Some((model, _, _, _)) = wq.graph_weight() else {
return crate::gpu::BatchGraphOutcome::Declined;
};
let model = model.clone();
let (Some(gwq), Some(gwk), Some(gwv), Some(gwo), Some(gg), Some(gu), Some(gd)) = (
gw(wq),
gw(wk),
gw(wv),
gw(wo),
gw(&d.gate_proj),
gw(&d.up_proj),
gw(&d.down_proj),
) else {
return crate::gpu::BatchGraphOutcome::Declined;
};
let layer = crate::gpu::GraphLayer {
input_norm: &lw.input_norm,
attn: crate::gpu::GraphAttn::Full {
wq: gwq,
wk: gwk,
wv: gwv,
wo: gwo,
q_norm: q_norm.as_deref(),
k_norm: k_norm.as_deref(),
late_qk_norm: self.qk_norm_after_rope,
bias: bias
.as_ref()
.map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
output_gate: *output_gate,
cpu_k: m.kv.k_heads(),
cpu_v: m.kv.v_heads(),
},
post_norm: &lw.post_norm,
ffn: crate::gpu::GraphFfn::Dense {
gate: gg,
up: gu,
down: gd,
},
};
let positions: Vec<usize> = (first_pos..first_pos + pairs.len()).collect();
let nh = self.num_heads;
let (nkv, hd, rd) = self.layer_geom(0);
let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
crate::gpu::forward_batch_graph(
&model,
self.mtp_kv_id(),
std::slice::from_ref(&layer),
&self.inv_freq,
&mut hiddens,
nh,
nkv,
hd,
rd,
hs,
self.intermediate_size,
&positions,
self.kv_cache.max_seq_len,
gemma,
self.rms_eps as f32,
self.attn_scale,
pairs.len(),
&[],
0,
None,
)
}
#[cfg(feature = "gpu")]
fn mtp_warm_graph_fallback(
&mut self,
m: &mut MtpModule,
pairs: &[(&[f32], u32)],
first_pos: usize,
) -> bool {
if pairs.is_empty() {
return true;
}
let graphable = self.mtp_block_graph_ok(m);
if !graphable {
if self.mtp_graph_mode == Some(true) {
return false;
}
self.mtp_graph_mode = Some(false);
for (j, (h, t)) in pairs.iter().enumerate() {
self.mtp_warm(m, h, *t, first_pos + j);
}
return true;
}
for (j, (h, t)) in pairs.iter().enumerate() {
if self.mtp_step_graph(m, h, *t, first_pos + j).is_none() {
return false;
}
}
self.mtp_graph_mode = Some(true);
true
}
#[cfg(feature = "gpu")]
fn mtp_warm_prefill_pairs(
&mut self,
m: &mut MtpModule,
pairs: &[(&[f32], u32)],
first_pos: usize,
) -> Result<(), &'static str> {
if self.mtp_graph_mode == Some(false) || !self.mtp_graph_ok(m) {
if self.mtp_graph_mode == Some(true) {
return Err("MTP token graph became unavailable after admission");
}
self.mtp_graph_mode = Some(false);
for (j, (h, t)) in pairs.iter().enumerate() {
self.mtp_warm(m, h, *t, first_pos + j);
}
return Ok(());
}
match self.mtp_warm_graph(m, pairs, first_pos) {
crate::gpu::BatchGraphOutcome::Completed => {
if !pairs.is_empty() {
self.mtp_graph_mode = Some(true);
}
Ok(())
}
crate::gpu::BatchGraphOutcome::Declined => {
if self.mtp_warm_graph_fallback(m, pairs, first_pos) {
Ok(())
} else {
Err("MTP warm-up fallback failed after device admission")
}
}
crate::gpu::BatchGraphOutcome::Failed => {
Err("MTP warm batch graph failed after admission")
}
}
}
#[cfg(not(feature = "gpu"))]
fn mtp_warm_prefill_pairs(
&mut self,
m: &mut MtpModule,
pairs: &[(&[f32], u32)],
first_pos: usize,
) -> Result<(), &'static str> {
for (j, (h, t)) in pairs.iter().enumerate() {
self.mtp_warm(m, h, *t, first_pos + j);
}
Ok(())
}
fn mtp_warm(&mut self, m: &mut MtpModule, hidden: &[f32], next_token: u32, position: usize) {
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());
inference::rms_norm_into(
&x,
&m.layer.input_norm,
self.rms_eps,
self.norm_style,
&mut self.ws.n1,
);
let attn = match &m.layer.attn {
AttnKind::Full {
wq,
wk,
wv,
wo,
q_norm,
k_norm,
output_gate,
softplus_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.softplus_gate = softplus_gate.as_ref().map(|(g, p)| (g, *p));
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)
}
_ => return,
};
let _ = attn;
}
#[cfg(feature = "gpu")]
#[allow(clippy::too_many_arguments)]
fn graph_spec_step(
&mut self,
m: &mut MtpModule,
hidden: &[f32],
t_next: u32,
next_pos: usize,
drafted: &mut usize,
accepted: &mut usize,
all_ids: &mut Vec<u32>,
room: usize,
) -> Option<(Vec<u32>, usize, Vec<f32>)> {
#[cfg(target_os = "macos")]
let metal_native = crate::gpu::q1_force();
#[cfg(not(target_os = "macos"))]
let metal_native = false;
#[cfg(feature = "gpu")]
let k_default = if metal_native {
7
} else if crate::gpu_wgpu::verify_i8_on() {
5
} else {
4
};
#[cfg(not(feature = "gpu"))]
let k_default = 4;
let k_env: Option<usize> = std::env::var("CMF_GRAPH_SPEC_K")
.ok()
.and_then(|v| v.parse().ok())
.filter(|&v| (1..=8).contains(&v));
let (k_start, k_max) = if metal_native { (7, 7) } else { (3, k_default.max(5)) };
let k_full: usize = k_env.unwrap_or_else(|| self.spec_k_adapt.unwrap_or(k_start));
let k_spec = k_full.min(room).max(1);
let k_capped = k_spec < k_full;
if next_pos == 0 {
return None;
}
let t_round = std::time::Instant::now();
let subs = || crate::gpu_wgpu::SUBMITS.load(std::sync::atomic::Ordering::Relaxed);
let sub0 = subs();
let cfg = self.sampler_config.clone();
let penalized = !(cfg.repetition_penalty == 1.0
&& cfg.presence_penalty == 0.0
&& cfg.suppress_tokens.is_empty());
let greedy_pen = cfg.temperature < 1e-6 && penalized;
let sampling = cfg.temperature >= 1e-6;
let sparse = sampling && sampler::sparse_ok(&cfg);
let base_len = all_ids.len();
if sampling && !sparse && self.spec_q.len() < k_spec {
self.spec_q.resize_with(k_spec, Vec::new);
}
if sparse && self.spec_qs.len() < k_spec {
self.spec_qs.resize_with(k_spec, Vec::new);
}
let mut drafts = Vec::with_capacity(k_spec);
let mut hx = hidden.to_vec();
let spec_dbg = std::env::var("CMF_SPEC_DBG").is_ok();
spec_stamp("pro");
#[cfg(target_os = "macos")]
if metal_native && !sampling && !greedy_pen && self.mtp_graph_mode != Some(false) {
match self.mtp_draft_chain_metal(m, hidden, t_next, next_pos - 1, k_spec) {
Ok(ids) => {
self.mtp_graph_mode = Some(true);
drafts = ids;
}
Err(true) => {
tracing::error!("mtp Metal draft chain failed after commit");
self.clear_sequence_state();
self.graph_failed
.store(true, std::sync::atomic::Ordering::Relaxed);
self.cancel
.store(true, std::sync::atomic::Ordering::Relaxed);
return None;
}
Err(false) => {}
}
}
for j in drafts.len()..k_spec {
let tok_in = if j == 0 { t_next } else { drafts[j - 1] };
let mut dbg_ref: Option<(Vec<f32>, Vec<f32>)> = None;
if spec_dbg {
let saved = self.mtp_graph_mode;
self.mtp_graph_mode = Some(false);
let r = self.mtp_step_hl(m, &hx, tok_in, next_pos - 1 + j);
self.mtp_graph_mode = saved;
if self.graph_failed.load(std::sync::atomic::Ordering::Relaxed) {
return None;
}
m.kv.truncate_last(1);
dbg_ref = Some(r);
}
let (mut lg, hj) = self.mtp_step_hl(m, &hx, tok_in, next_pos - 1 + j);
if self.graph_failed.load(std::sync::atomic::Ordering::Relaxed) {
return None;
}
if let Some((lg_cpu, h_cpu)) = dbg_ref {
let n = |v: &[f32]| v.iter().map(|x| x * x).sum::<f32>().sqrt();
let dl = lg
.iter()
.zip(&lg_cpu)
.fold(0f32, |m, (a, b)| m.max((a - b).abs()));
let dh = hj
.iter()
.zip(&h_cpu)
.fold(0f32, |m, (a, b)| m.max((a - b).abs()));
eprintln!(
"spec-dbg j={j} pos {} tok_in {tok_in}: per-op draft {} graph draft {} | max|dlogit| {dl:.3} | |h_cpu| {:.2} |h_graph| {:.2} max|dh| {dh:.3} | kv rows {}",
next_pos - 1 + j,
sampler::argmax(&lg_cpu),
sampler::argmax(&lg),
n(&h_cpu),
n(&hj),
m.kv.seq_len
);
}
let dj = if sparse {
let mut q = std::mem::take(&mut self.spec_qs[j]);
let ok = sampler::sparse_distribution_into(
&lg,
&cfg,
all_ids,
&mut self.sampler_scratch,
self.pool.as_deref(),
&mut q,
);
let d = if ok {
sampler::draw_sparse(&q, &mut self.rng)
} else {
let t = sampler::argmax(&lg);
q.clear();
q.push((t, 1.0));
t
};
self.spec_qs[j] = q;
all_ids.push(d);
d
} else if sampling {
let mut q = std::mem::take(&mut self.spec_q[j]);
sampler::distribution_into(
&lg,
&cfg,
all_ids,
&mut self.sampler_scratch,
self.pool.as_deref(),
&mut q,
);
let d = sampler::draw(&q, &mut self.rng);
self.spec_q[j] = q;
all_ids.push(d); d
} else if greedy_pen {
let d = sampler::argmax_penalized(
&lg,
&cfg,
all_ids,
&mut self.sampler_scratch,
self.pool.as_deref(),
);
all_ids.push(d);
d
} else {
sampler::argmax(&lg)
};
attention::recycle_buf(&mut lg);
drafts.push(dj);
hx = hj;
spec_stamp("d.pick");
}
all_ids.truncate(base_len);
*drafted += k_spec;
let t_draft = t_round.elapsed();
let sub_draft = subs();
let b = k_spec + 1;
let mut hiddens = vec![0.0f32; b * self.hidden_size];
for (i, &t) in std::iter::once(&t_next).chain(drafts.iter()).enumerate() {
let e = self.embed_single(t);
hiddens[i * self.hidden_size..(i + 1) * self.hidden_size].copy_from_slice(&e);
}
let positions: Vec<usize> = (next_pos..next_pos + b).collect();
spec_stamp("v.emb");
let (lm_gw, lm_rows) = {
let (_, i, kind, rs) = self.weights.lm_head.graph_weight()?;
(
crate::gpu::GraphW {
idx: i,
kind,
row_scale: rs,
data: &[],
prism: crate::gpu::GraphPrismOp::None,
affine: false,
},
self.weights.lm_head.rows(),
)
};
let mut logits = Vec::new();
let final_norm = self.weights.final_norm.clone();
#[cfg(target_os = "macos")]
let greedy_dev = metal_native
&& !sampling
&& !greedy_pen
&& !self.confidence_on
&& self.final_softcap.is_none()
&& self.vocab_size == lm_rows
&& std::env::var_os("CMF_METAL_VERIFY_CHECK").is_none()
&& std::env::var_os("CMF_LOGIT_DUMP").is_none()
&& std::env::var("CMF_METAL_DEV_ARGMAX").as_deref() != Ok("0");
#[cfg(not(target_os = "macos"))]
let greedy_dev = false;
let mut dev_ids: Vec<u32> = Vec::new();
#[cfg(target_os = "macos")]
let verify_outcome = if metal_native {
let lm = self.weights.lm_head.q1_parts()?;
let n_score = self.vocab_size.min(lm_rows);
self.try_batch_graph_metal(
&mut hiddens,
&positions,
b,
Some((lm, &final_norm, &mut logits)),
if greedy_dev {
Some((n_score, &mut dev_ids))
} else {
None
},
)
} else {
self.try_batch_graph_wgpu(
&mut hiddens,
&positions,
b,
Some(crate::gpu::SpecTail {
lm: lm_gw,
lm_rows,
final_norm: &final_norm,
logits_out: &mut logits,
}),
)
};
#[cfg(not(target_os = "macos"))]
let verify_outcome = self.try_batch_graph_wgpu(
&mut hiddens,
&positions,
b,
Some(crate::gpu::SpecTail {
lm: lm_gw,
lm_rows,
final_norm: &final_norm,
logits_out: &mut logits,
}),
);
match verify_outcome {
crate::gpu::BatchGraphOutcome::Completed => {}
crate::gpu::BatchGraphOutcome::Declined => {
m.kv.truncate_last(k_spec);
if !metal_native && !self.rewind_mtp_graph_mirror(next_pos) {
self.clear_sequence_state();
self.graph_failed
.store(true, std::sync::atomic::Ordering::Relaxed);
self.cancel
.store(true, std::sync::atomic::Ordering::Relaxed);
tracing::error!("MTP graph mirror rewind failed after verify decline");
}
return None;
}
crate::gpu::BatchGraphOutcome::Failed => {
self.clear_sequence_state();
self.graph_failed
.store(true, std::sync::atomic::Ordering::Relaxed);
self.cancel
.store(true, std::sync::atomic::Ordering::Relaxed);
tracing::error!("MTP verify batch graph failed after admission");
return None;
}
}
#[cfg(target_os = "macos")]
if metal_native && std::env::var("CMF_METAL_VERIFY_CHECK").as_deref() == Ok("1") {
let snap: Vec<Vec<f32>> = self
.kv_cache
.layers
.iter()
.map(|l| l.linear_state.clone())
.collect();
let attn_lens: Vec<usize> = self.kv_cache.layers.iter().map(|l| l.seq_len).collect();
let toks: Vec<u32> = std::iter::once(t_next)
.chain(drafts.iter().copied())
.collect();
let want_save = self.graph_want_logits;
self.graph_want_logits = false;
for (i, &t) in toks.iter().enumerate() {
let hi = self.forward_layers(&self.embed_single(t), next_pos + i, None);
let _ = self.graph_logits.take();
if std::env::var("CMF_SPEC_PLAIN_HIDDEN").as_deref() == Ok("1") {
hiddens[i * self.hidden_size..(i + 1) * self.hidden_size].copy_from_slice(&hi);
}
let ref_lg = self.logits_from_hidden(&hi);
let row = &logits[i * lm_rows..(i + 1) * lm_rows];
let ra = sampler::argmax(&ref_lg);
let va = sampler::argmax(row);
let mut md = 0f32;
let mut rms = 0f64;
for j in 0..lm_rows.min(ref_lg.len()) {
let d = (ref_lg[j] - row[j]).abs();
md = md.max(d);
rms += (d as f64) * (d as f64);
}
let mut hd = 0f32;
for j in 0..self.hidden_size {
hd = hd.max((hi[j] - hiddens[i * self.hidden_size + j]).abs());
}
eprintln!(
"verify-check row {i} tok {t} pos {}: ref argmax {ra} verify argmax {va} {} | max|dlogit| {md:.3} rms {:.4} | max|dhidden| {hd:.4}",
next_pos + i,
if ra == va { "OK" } else { "MISMATCH" },
(rms / lm_rows as f64).sqrt()
);
}
self.graph_want_logits = want_save;
for (l, st) in self.kv_cache.layers.iter_mut().zip(snap) {
if l.linear_state.len() == st.len() {
l.linear_state.copy_from_slice(&st);
} else {
l.linear_state = st;
}
}
for (li, (l, n0)) in self.kv_cache.layers.iter_mut().zip(attn_lens).enumerate() {
let extra = l.seq_len.saturating_sub(n0);
if extra > 0 {
l.truncate_last(extra);
crate::gpu_metal::kv_mirror_set_stored(self.graph_kv_id, li, n0);
}
}
}
let t_verify = t_round.elapsed();
let sub_verify = subs();
let mut a = 0usize;
let mut forced: Option<u32> = None;
let ids: Vec<u32> = if sparse {
let mut p = std::mem::take(&mut self.spec_ps);
let mut res = std::mem::take(&mut self.spec_ress);
while a < k_spec {
let ok = sampler::sparse_distribution_into(
&logits[a * lm_rows..(a + 1) * lm_rows],
&cfg,
all_ids,
&mut self.sampler_scratch,
self.pool.as_deref(),
&mut p,
);
if !ok {
let t = sampler::argmax(&logits[a * lm_rows..(a + 1) * lm_rows]);
p.clear();
p.push((t, 1.0));
}
match sampler::spec_accept_or_correct_sparse(
&p,
&self.spec_qs[a],
drafts[a],
&mut self.rng,
&mut res,
) {
None => {
all_ids.push(drafts[a]);
a += 1;
}
Some(c) => {
forced = Some(c);
break;
}
}
}
all_ids.truncate(base_len);
self.spec_ps = p;
self.spec_ress = res;
drafts.clone()
} else if sampling {
let mut p = std::mem::take(&mut self.spec_p);
let mut res = std::mem::take(&mut self.spec_res);
while a < k_spec {
sampler::distribution_into(
&logits[a * lm_rows..(a + 1) * lm_rows],
&cfg,
all_ids,
&mut self.sampler_scratch,
self.pool.as_deref(),
&mut p,
);
match sampler::spec_accept_or_correct(
&p,
&self.spec_q[a],
drafts[a],
&mut self.rng,
&mut res,
self.pool.as_deref(),
) {
None => {
all_ids.push(drafts[a]);
a += 1;
}
Some(c) => {
forced = Some(c);
break;
}
}
}
all_ids.truncate(base_len);
self.spec_p = p;
self.spec_res = res;
drafts.clone()
} else if greedy_pen {
let mut ids: Vec<u32> = Vec::with_capacity(b);
for i in 0..b {
let t = sampler::argmax_penalized(
&logits[i * lm_rows..(i + 1) * lm_rows],
&cfg,
all_ids,
&mut self.sampler_scratch,
self.pool.as_deref(),
);
ids.push(t);
if i < k_spec && t == drafts[i] {
all_ids.push(t);
} else {
break;
}
}
all_ids.truncate(base_len);
while a < k_spec && a < ids.len() && ids[a] == drafts[a] {
a += 1;
}
ids
} else if greedy_dev && dev_ids.len() == b {
let ids = std::mem::take(&mut dev_ids);
while a < k_spec && ids[a] == drafts[a] {
a += 1;
}
ids
} else {
if logits.len() < b * lm_rows {
self.clear_sequence_state();
self.graph_failed
.store(true, std::sync::atomic::Ordering::Relaxed);
self.cancel
.store(true, std::sync::atomic::Ordering::Relaxed);
tracing::error!("Metal verify returned neither logits nor argmax ids");
return None;
}
let ids: Vec<u32> = (0..b)
.map(|i| sampler::argmax(&logits[i * lm_rows..(i + 1) * lm_rows]))
.collect();
while a < k_spec && ids[a] == drafts[a] {
a += 1;
}
ids
};
spec_stamp("acc");
if spec_dbg {
eprintln!(
"spec-dbg round: t_next {t_next} drafts {:?} verified {:?} accepted {a}",
drafts, ids
);
}
#[cfg(target_os = "macos")]
let commit_ref: Option<(Vec<Vec<f32>>, Vec<(usize, Vec<f32>, Vec<f32>)>)> = if metal_native
&& std::env::var("CMF_METAL_VERIFY_CHECK").as_deref() == Ok("2")
{
let snap: Vec<Vec<f32>> = self
.kv_cache
.layers
.iter()
.map(|l| l.linear_state.clone())
.collect();
let attn_lens: Vec<usize> = self.kv_cache.layers.iter().map(|l| l.seq_len).collect();
let toks: Vec<u32> = std::iter::once(t_next)
.chain(drafts.iter().copied())
.collect();
let want_save = self.graph_want_logits;
self.graph_want_logits = false;
for (i, &t) in toks.iter().take(a + 1).enumerate() {
let _ = self.forward_layers(&self.embed_single(t), next_pos + i, None);
let _ = self.graph_logits.take();
}
self.graph_want_logits = want_save;
let plain_states: Vec<Vec<f32>> = self
.kv_cache
.layers
.iter()
.map(|l| l.linear_state.clone())
.collect();
let (nkv, hd) = (self.num_kv_heads, self.head_dim);
let mut rows = Vec::new();
for (li, (l, n0)) in self
.kv_cache
.layers
.iter_mut()
.zip(attn_lens.iter())
.enumerate()
{
let extra = l.seq_len.saturating_sub(*n0);
if extra > 0 {
let mut kk = Vec::new();
let mut vv = Vec::new();
for g in 0..nkv {
kk.extend_from_slice(&l.head_keys(g)[n0 * hd..]);
vv.extend_from_slice(&l.head_values(g)[n0 * hd..]);
}
rows.push((li, kk, vv));
l.truncate_last(extra);
crate::gpu_metal::kv_mirror_set_stored(self.graph_kv_id, li, *n0);
}
}
for (l, st) in self.kv_cache.layers.iter_mut().zip(snap) {
if l.linear_state.len() == st.len() {
l.linear_state.copy_from_slice(&st);
} else {
l.linear_state = st;
}
}
Some((plain_states, rows))
} else {
None
};
let warm_off = std::env::var("CMF_SPEC_WARM").is_ok_and(|v| v == "0");
#[cfg(target_os = "macos")]
let mut warm_pending: Option<MetalWarmPending> = None;
#[cfg(target_os = "macos")]
if metal_native {
m.kv.truncate_last(k_spec.saturating_sub(1));
if self.mtp_graph_mode == Some(true) {
crate::gpu_metal::kv_mirror_set_stored(
self.mtp_kv_id(),
Self::MTP_LAYER_BASE,
m.kv.seq_len,
);
if !warm_off && a > 0 {
let pairs: Vec<(&[f32], u32)> = (0..a)
.map(|j| {
(
&hiddens[j * self.hidden_size..(j + 1) * self.hidden_size],
ids[j],
)
})
.collect();
warm_pending = self.mtp_warm_batch_submit(m, &pairs, next_pos);
}
}
spec_stamp("c.wsub");
}
#[cfg(target_os = "macos")]
if metal_native {
if !self.metal_verify_commit(a) {
self.clear_sequence_state();
self.graph_failed
.store(true, std::sync::atomic::Ordering::Relaxed);
self.cancel
.store(true, std::sync::atomic::Ordering::Relaxed);
tracing::error!("Metal verify state/KV handoff failed after admission");
return None;
}
if let Some((plain_states, rows)) = commit_ref {
crate::gpu_metal::queue_fence();
let _ = crate::gpu_metal::wait_replay();
let (nkv, hd) = (self.num_kv_heads, self.head_dim);
let mut worst_s = 0f32;
let mut worst_li = 0usize;
for (li, (l, ps)) in self.kv_cache.layers.iter().zip(&plain_states).enumerate() {
if l.linear_state.len() != ps.len() || ps.is_empty() {
continue;
}
let d = l
.linear_state
.iter()
.zip(ps)
.fold(0f32, |m, (x, y)| m.max((x - y).abs()));
let n = ps.iter().fold(0f32, |m, y| m.max(y.abs()));
let rel = d / n.max(1e-6);
if rel > worst_s {
worst_s = rel;
worst_li = li;
}
}
let mut worst_k = 0f32;
for (li, kk, vv) in &rows {
let l = &self.kv_cache.layers[*li];
let n0 = l.seq_len - (kk.len() / (nkv * hd));
let mut ck = Vec::new();
let mut cv = Vec::new();
for g in 0..nkv {
ck.extend_from_slice(&l.head_keys(g)[n0 * hd..]);
cv.extend_from_slice(&l.head_values(g)[n0 * hd..]);
}
if ck.len() == kk.len() {
let dk = ck
.iter()
.zip(kk)
.fold(0f32, |m, (x, y)| m.max((x - y).abs()));
let dv = cv
.iter()
.zip(vv)
.fold(0f32, |m, (x, y)| m.max((x - y).abs()));
worst_k = worst_k.max(dk).max(dv);
} else {
eprintln!(
"commit-check L{li}: kv row count mismatch {} vs {}",
ck.len(),
kk.len()
);
}
}
eprintln!(
"commit-check a={a}: worst GDN state rel-max diff {worst_s:.2e} (L{worst_li}) | worst K/V row abs diff {worst_k:.4}"
);
}
}
if !metal_native && a + 1 < b {
let expected_gdn_layers = self.graph_gdn_layer_count();
if expected_gdn_layers > 0
&& !crate::gpu::gdn_spec_restore(self.graph_kv_id, a, next_pos, expected_gdn_layers)
{
self.clear_sequence_state();
self.graph_failed
.store(true, std::sync::atomic::Ordering::Relaxed);
self.cancel
.store(true, std::sync::atomic::Ordering::Relaxed);
tracing::error!("GDN speculative restore failed after verify");
return None;
}
}
if !metal_native && !self.rewind_trunk_graph_mirrors(next_pos + a + 1) {
self.clear_sequence_state();
self.graph_failed
.store(true, std::sync::atomic::Ordering::Relaxed);
self.cancel
.store(true, std::sync::atomic::Ordering::Relaxed);
tracing::error!("trunk graph KV rewind failed after speculative verify");
return None;
}
*accepted += a;
if !metal_native {
m.kv.truncate_last(k_spec.saturating_sub(1));
}
spec_stamp("c.trunc");
if !metal_native
&& self.mtp_graph_mode == Some(true)
&& !self.rewind_mtp_graph_mirror(next_pos)
{
self.clear_sequence_state();
self.graph_failed
.store(true, std::sync::atomic::Ordering::Relaxed);
self.cancel
.store(true, std::sync::atomic::Ordering::Relaxed);
tracing::error!("MTP graph mirror rewind failed after verify commit");
return None;
}
if !warm_off && a > 0 {
let mut warmed = false;
#[cfg(target_os = "macos")]
if metal_native && self.mtp_graph_mode == Some(true) {
warmed = match warm_pending.take() {
Some(p) => self.mtp_warm_batch_finish(m, p),
None => false,
};
if !warmed {
warmed = true;
for j in 0..a {
let row =
hiddens[j * self.hidden_size..(j + 1) * self.hidden_size].to_vec();
if self
.mtp_step_metal(m, &row, ids[j], next_pos + j, false)
.is_none()
{
warmed = false;
break;
}
}
}
}
if !warmed && self.mtp_graph_mode != Some(false) && !metal_native {
let rows: Vec<Vec<f32>> = (0..a)
.map(|j| hiddens[j * self.hidden_size..(j + 1) * self.hidden_size].to_vec())
.collect();
let pairs: Vec<(&[f32], u32)> = rows
.iter()
.zip(ids.iter())
.map(|(r, &t)| (r.as_slice(), t))
.collect();
match self.mtp_warm_prefill_pairs(m, &pairs, next_pos) {
Ok(()) => warmed = true,
Err(err) => {
tracing::error!("{err}");
self.clear_sequence_state();
self.graph_failed
.store(true, std::sync::atomic::Ordering::Relaxed);
self.cancel
.store(true, std::sync::atomic::Ordering::Relaxed);
return None;
}
}
}
if !warmed {
for j in 0..a {
let row = &hiddens[j * self.hidden_size..(j + 1) * self.hidden_size];
let row = row.to_vec();
self.mtp_warm(m, &row, ids[j], next_pos + j);
}
}
}
spec_stamp("c.warm");
if let Some(c) = forced {
self.spec_forced = Some(c);
self.graph_logits = None;
} else if greedy_dev && logits.is_empty() {
self.spec_forced = Some(ids[a]);
self.graph_logits = None;
} else {
let mut row = logits[a * lm_rows..(a + 1) * lm_rows].to_vec();
row.resize(self.vocab_size, 0.0);
if let Some(c) = self.final_softcap {
for l in row.iter_mut() {
*l = c * (*l / c).tanh();
}
}
self.graph_logits = Some(row);
}
let new_hidden = hiddens[a * self.hidden_size..(a + 1) * self.hidden_size].to_vec();
spec_stamp("c.row");
if std::env::var("CMF_GRAPH_SPEC_TIME").is_ok() {
let end = subs();
eprintln!(
"spec-round: draft {:.1} ms/{} sub | verify {:.1} ms/{} sub | \
commit {:.1} ms/{} sub (accepted {a} of {k_spec}, full-head streak {})",
t_draft.as_secs_f64() * 1e3,
sub_draft - sub0,
(t_verify - t_draft).as_secs_f64() * 1e3,
sub_verify - sub_draft,
(t_round.elapsed() - t_verify).as_secs_f64() * 1e3,
end - sub_verify,
self.draft_full_streak,
);
}
if k_env.is_none() && !metal_native && !k_capped {
let f = a as f32 / k_spec.max(1) as f32;
self.spec_acc_ewma += 0.2 * (f - self.spec_acc_ewma);
let mut k_next = k_spec;
if self.spec_acc_ewma >= 0.75 && k_spec < k_max {
k_next = k_spec + 1;
} else if self.spec_acc_ewma < 0.4 && k_spec > 2 {
k_next = k_spec - 1;
}
if k_next != k_spec {
self.spec_acc_ewma = 0.6;
if std::env::var("CMF_GRAPH_SPEC_TIME").is_ok() {
eprintln!("spec-k: {k_spec} → {k_next}");
}
}
self.spec_k_adapt = Some(k_next);
}
spec_stamp("end");
Some((drafts[..a].to_vec(), next_pos + a + 1, new_hidden))
}
pub fn measure_pair_fusion(&mut self, iters: usize) -> (f64, f64) {
if !self.pair_supported() {
return (0.0, 0.0);
}
let graph_env = std::env::var_os("CMF_GPU_WGPU_GRAPH");
unsafe { std::env::set_var("CMF_GPU_WGPU_GRAPH", "0") };
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;
match graph_env {
Some(value) => unsafe { std::env::set_var("CMF_GPU_WGPU_GRAPH", value) },
None => unsafe { std::env::remove_var("CMF_GPU_WGPU_GRAPH") },
}
(singles_ms, pair_ms)
}
fn pair_supported(&self) -> bool {
!self.weights.layers.is_empty()
&& self.g3n.is_none()
&& !self
.weights
.layers
.iter()
.any(|lw| matches!(&lw.attn, AttnKind::Mla(_) | AttnKind::Kda(_)))
}
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 (_nkv, _hd, hs, _rd, eps) = (
self.num_kv_heads,
self.head_dim,
self.hidden_size,
self.rotary_dim,
self.rms_eps,
);
let pool = self.pool.clone();
for li in 0..self.num_layers {
let lw = &self.weights.layers[self.phys_layer(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::Mla(_) => unreachable!("MLA has no MTP/pair path"),
AttnKind::Kda(_) => unreachable!("KDA has no MTP/pair path"),
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::ShortConv(w) => {
let cfg = self
.short_conv_cfg
.expect("short-conv layer without short_conv_cfg");
let layer = &mut self.kv_cache.layers[li];
let (state, scratch) = (&mut layer.linear_state, &mut layer.linear_scratch);
short_conv_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,
softplus_gate,
bias,
} => {
let inv_freq_l = self.layer_inv_freq(li);
let (nkv_l, hd_l, rd_l) = self.layer_geom(li);
let cfg = QwenAttnCfg {
num_heads: self.layer_num_heads(li),
num_kv_heads: nkv_l,
head_dim: hd_l,
hidden_size: hs,
position,
inv_freq: &inv_freq_l,
rotary_dim: rd_l,
scale: self.attn_scale,
softcap: self.attn_softcap,
window: self.layer_window(li),
v_norm: self.attn_v_norm,
qk_norm_after_rope: self.qk_norm_after_rope,
q_norm: q_norm.as_deref(),
k_norm: k_norm.as_deref(),
output_gate: *output_gate,
softplus_gate: softplus_gate
.as_ref()
.map(|(gate, per_head)| (gate, *per_head)),
rope_scale: self.layer_rope_scale(li),
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,
)
}
};
let (a1, a2) = match &self.weights.layers[self.phys_layer(li)].attn_out_norm {
Some(w) => (
inference::rms_norm(&a1, w, self.rms_eps, self.norm_style),
inference::rms_norm(&a2, w, self.rms_eps, self.norm_style),
),
None => (a1, a2),
};
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[self.phys_layer(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) = match &lw.ffn {
FfnKind::DenseMoe(dm) => (
dense_moe_ffn(
dm,
&self.ws.p1,
&h1,
self.rms_eps,
self.norm_style,
self.pool.as_deref(),
),
dense_moe_ffn(
dm,
&self.ws.p2,
&h2,
self.rms_eps,
self.norm_style,
self.pool.as_deref(),
),
),
_ => ffn_forward_pair(
&lw.ffn,
&self.ws.p1,
&self.ws.p2,
self.pool.as_deref(),
None,
),
};
let (f1, f2) = match &self.weights.layers[self.phys_layer(li)].ffn_out_norm {
Some(w) => (
inference::rms_norm(&f1, w, self.rms_eps, self.norm_style),
inference::rms_norm(&f2, w, self.rms_eps, self.norm_style),
),
None => (f1, f2),
};
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);
if let Some(sc) = self.weights.layers[self.phys_layer(li)].layer_scale {
for i in 0..self.hidden_size {
h1[i] *= sc;
h2[i] *= sc;
}
}
if self.is_loop_end(li) && li + 1 < self.num_layers {
h1 = inference::rms_norm(
&h1,
&self.weights.final_norm,
self.rms_eps,
self.norm_style,
);
h2 = inference::rms_norm(
&h2,
&self.weights.final_norm,
self.rms_eps,
self.norm_style,
);
}
}
if self.o1_active() {
self.commit_linear_scratch();
}
self.o1_progress();
(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.clear_sequence_state();
self.check_forward_graph("forward_ids setup", 0)?;
if task_mask.is_none() {
self.o1_begin();
}
let mut hidden = vec![0.0f32; self.hidden_size];
let mut pos = 0usize;
if let Some(b) = &mut self.dsv41 {
let pool = self.pool.clone();
let mut logits = Vec::new();
crate::dsv41::forward_chunk(
&b.0,
&b.1,
&b.2,
&mut b.3,
ids,
0,
pool.as_deref(),
&mut logits,
);
if let Err(err) = self.o1_seal_checked() {
self.clear_sequence_state();
return Err(err);
}
return Ok(logits);
}
if self.can_prefill_batched() && !self.graph_prefill_preferred() && ids.len() > 2 {
let chunk = self.prefill_chunk();
let hs = self.hidden_size;
while pos < ids.len() {
let end = (pos + chunk).min(ids.len());
let hb = self.prefill_batch_masked(&ids[pos..end], pos, task_mask);
self.check_forward_graph("forward_ids batched prefill", end - 1)?;
hidden.copy_from_slice(&hb[(end - pos - 1) * hs..]);
pos = end;
}
}
if task_mask.is_none()
&& !self.graph_prefill_preferred()
&& !std::env::var("CMF_PAIR").is_ok_and(|v| v == "0")
&& self.pair_supported()
{
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.check_forward_graph("forward_ids pair", pos + 1)?;
self.commit_linear_scratch();
hidden = h2;
pos += 2;
}
}
while pos < ids.len() {
hidden = self.forward_layers(&self.embed_single(ids[pos]), pos, task_mask);
self.check_forward_graph("forward_ids", pos)?;
pos += 1;
}
if let Err(err) = self.o1_seal_checked() {
self.clear_sequence_state();
return Err(err);
}
let normed = inference::rms_norm(
&hidden,
&self.weights.final_norm,
self.rms_eps,
self.norm_style,
);
Ok(self.lm_head_forward(&normed))
}
#[doc(hidden)]
pub fn dsv41_serial_logits(&mut self, ids: &[u32]) -> Result<Vec<Vec<f32>>, String> {
#[cfg(target_os = "macos")]
crate::gpu_metal::set_io_namespace(self.graph_kv_id);
if ids.is_empty() {
return Err("empty id sequence".to_string());
}
self.clear_sequence_state();
self.dsv41
.as_ref()
.ok_or_else(|| "dsv41 serial logits require a DeepSeek-V4.1 model".to_string())?;
self.o1_begin();
let rows = {
let pool = self.pool.clone();
let b = self
.dsv41
.as_mut()
.expect("dsv41 checked above; state cannot change during forward");
let mut rows = Vec::with_capacity(ids.len());
for (position, &id) in ids.iter().enumerate() {
let mut logits = Vec::new();
crate::dsv41::forward_token(
&b.0,
&b.1,
&b.2,
&mut b.3,
id,
position,
pool.as_deref(),
&mut logits,
);
rows.push(logits);
}
rows
};
self.o1_seal();
Ok(rows)
}
pub fn ppl_ids(&mut self, ids: &[u32]) -> Result<f64, String> {
let (nll, cnt) = self.nll_ids_from(ids, 0)?;
Ok((nll / cnt.max(1) as f64).exp())
}
pub fn probe_ffn_mass(&mut self, ids: &[u32]) -> Vec<Vec<f64>> {
self.clear_sequence_state();
FFN_PROBE.with(|p| {
*p.borrow_mut() = Some(vec![vec![0f64; self.intermediate_size]; self.num_layers]);
});
crate::gpu::cpu_scope(|| {
for (pos, &id) in ids.iter().enumerate() {
let emb = self.embed_single(id);
let _ = self.forward_layers(&emb, pos, None);
}
});
self.clear_sequence_state();
FFN_PROBE
.with(|p| p.borrow_mut().take())
.unwrap_or_default()
}
pub fn probe_ffn_mass_batch(&mut self, ids: &[u32]) -> Result<Vec<Vec<f64>>, String> {
if let Err(err) = self.nll_begin() {
let _ = FFN_PROBE.with(|p| p.borrow_mut().take());
self.nll_end();
return Err(err);
}
FFN_PROBE.with(|p| {
*p.borrow_mut() = Some(vec![vec![0f64; self.intermediate_size]; self.num_layers]);
});
let result: Result<(), String> = (|| {
for chunk in ids.chunks(256) {
if chunk.len() < 2 {
continue;
}
self.nll_ids_masked(chunk, 0, None)?;
}
Ok(())
})();
self.nll_end();
let probe = FFN_PROBE
.with(|p| p.borrow_mut().take())
.unwrap_or_default();
match result {
Ok(()) => Ok(probe),
Err(err) => {
drop(probe);
Err(err)
}
}
}
pub fn ppl_ids_masked(&mut self, ids: &[u32], mask: &TaskMask) -> Result<f64, String> {
self.nll_begin()?;
let result: Result<f64, String> = (|| {
let mut nll = 0f64;
let mut cnt = 0usize;
let mut hidden = vec![0f32; self.hidden_size];
for (pos, &id) in ids.iter().enumerate() {
if pos > 0 {
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 max = logits.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
let sum: f64 = logits.iter().map(|&v| ((v - max) as f64).exp()).sum();
let p = ((logits[id as usize] - max) as f64).exp() / sum.max(1e-300);
nll -= p.max(1e-300).ln();
cnt += 1;
attention::recycle_buf(&mut logits);
}
let emb = self.embed_single(id);
hidden = self.forward_layers(&emb, pos, Some(mask));
self.nll_check_graph("masked serial forward", pos)?;
let _ = self.graph_logits.take();
}
Ok((nll / cnt.max(1) as f64).exp())
})();
self.nll_end();
result
}
pub fn nll_ids_masked(
&mut self,
ids: &[u32],
start: usize,
task_mask: Option<&TaskMask>,
) -> Result<(f64, usize), String> {
let task_mask = self.drop_open_mask(task_mask);
self.nll_ids_inner(ids, start, task_mask)
}
pub fn nll_ids_from(&mut self, ids: &[u32], start: usize) -> Result<(f64, usize), String> {
self.nll_ids_inner(ids, start, None)
}
fn nll_ids_inner(
&mut self,
ids: &[u32],
start: usize,
task_mask: Option<&TaskMask>,
) -> Result<(f64, usize), String> {
self.nll_begin()?;
let result: Result<(f64, usize), String> = (|| {
let mut nll = 0f64;
let mut cnt = 0usize;
let (graph_quality, fused_head_quality) = nll_graph_policy(
task_mask.is_none(),
self.graph_prefill_preferred(),
crate::gpu::q1_force(),
);
self.graph_head_required = fused_head_quality;
self.graph_want_logits = fused_head_quality;
#[cfg(target_os = "macos")]
if graph_quality && std::env::var("CMF_METAL_BATCH_NLL").as_deref() != Ok("0") {
match self.nll_batch_metal(ids, start) {
MetalBatchNllOutcome::Completed(nll, count) => {
return Ok((nll, count));
}
MetalBatchNllOutcome::Declined => {}
MetalBatchNllOutcome::Failed(err) => return Err(err),
}
}
if self.can_prefill_batched() && !graph_quality {
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_masked(&ids[pos..end], pos, task_mask);
self.nll_check_graph("batched prefill", 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;
}
self.nll_check_graph("batched score row", pos + k0 + k)?;
let lg = &mut logits[k * rows..k * rows + self.vocab_size.min(rows)];
if let Some(mu) = self.logit_multiplier {
for v in lg.iter_mut() {
*v *= mu;
}
}
if let Some(c) = self.final_softcap {
for v in lg.iter_mut() {
*v = c * (*v / c).tanh();
}
}
if let Some(cm) = self.head_clusters.clone() {
self.hierarchical_head_logprobs(
&normed[k * hs..(k + 1) * hs],
&cm,
lg,
);
}
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;
if std::env::var("CMF_PPL_TRACE").is_ok() {
let top = lg
.iter()
.enumerate()
.max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
.map(|(i, _)| i)
.unwrap_or(0);
eprintln!(
"BTRACE pos {} target {} nll {:.4} top {} lg_t {:.3} lg_top {:.3}",
pos + k0 + k,
target,
lse - lg[target] as f64,
top,
lg[target],
lg[top]
);
}
}
k0 = k1;
}
pos = end;
}
return Ok((nll, cnt));
}
for pos in 0..ids.len().saturating_sub(1) {
let hidden = self.forward_layers(&self.embed_single(ids[pos]), pos, task_mask);
self.nll_check_graph("serial forward", pos)?;
let out_of_band = self.graph_logits.take();
if self.graph_head_required && out_of_band.is_none() {
METAL_GRAPH_HEAD_MISS.fetch_add(
1,
std::sync::atomic::Ordering::Relaxed,
);
return Err(format!(
"fused Metal graph head did not complete at NLL position {pos}"
));
}
if pos < start {
continue;
}
let logits = match out_of_band {
Some(lg) => lg,
None => {
let normed = inference::rms_norm(
&hidden,
&self.weights.final_norm,
self.rms_eps,
self.norm_style,
);
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;
let tok_nll = lse - logits[target] as f64;
if std::env::var("CMF_PPL_TRACE").is_ok() && pos < 48 {
let top = logits
.iter()
.enumerate()
.max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
.map(|(i, _)| i)
.unwrap_or(0);
eprintln!(
"pos {pos:3} tgt {target:6} nll {tok_nll:7.3} | top1 {top:6} lg[t]={:.2} lg[top]={:.2}",
logits[target], logits[top]
);
}
nll += tok_nll;
cnt += 1;
}
Ok((nll, cnt))
})();
self.nll_end();
result
}
fn nll_from_hidden(&mut self, hidden: &[f32], target: u32, pos: usize) -> f64 {
let normed = inference::rms_norm(
hidden,
&self.weights.final_norm,
self.rms_eps,
self.norm_style,
);
let mut logits = self.lm_head_forward(&normed);
let target = target 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;
let tok_nll = lse - logits[target] as f64;
if std::env::var("CMF_PPL_TRACE").is_ok() && pos < 48 {
let top = logits
.iter()
.enumerate()
.max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
.map(|(i, _)| i)
.unwrap_or(0);
eprintln!(
"pos {pos:3} tgt {target:6} nll {tok_nll:7.3} | top1 {top:6} lg[t]={:.2} lg[top]={:.2}",
logits[target], logits[top]
);
}
attention::recycle_buf(&mut logits);
tok_nll
}
pub fn nll_ids_o1(&mut self, ids: &[u32], prefill: usize) -> Result<(f64, usize), String> {
self.nll_begin()?;
let requested_prefix = (prefill > 0).then_some(prefill);
self.o1_begin_with_prefix(requested_prefix);
let n = ids.len().saturating_sub(1);
let requested_start = prefill.min(n);
let exact_end = if self.o1_active() {
match requested_prefix {
Some(requested) => self.o1_effective_boundary(requested),
None => self
.o1_cfg
.as_ref()
.and_then(|c| crate::nystrom::o1_deferred_boundary(c.w, c.sink)),
}
.unwrap_or(requested_start)
.min(n)
} else {
requested_start
};
let mut nll = 0f64;
let mut cnt = 0usize;
let mut pos = 0usize;
if self.can_prefill_batched() {
const CHUNK: usize = 128;
while pos < exact_end {
let end = (pos + CHUNK).min(exact_end);
let hiddens = self.prefill_batch(&ids[pos..end], pos);
if self
.graph_failed
.swap(false, std::sync::atomic::Ordering::Relaxed)
{
self.cancel
.store(false, std::sync::atomic::Ordering::Relaxed);
self.nll_end();
return Err("GPU graph failed during O(1) NLL prefix".into());
}
for row in 0..end - pos {
let score_pos = pos + row;
if score_pos >= requested_start && score_pos < n {
nll += self.nll_from_hidden(
&hiddens[row * self.hidden_size..(row + 1) * self.hidden_size],
ids[score_pos + 1],
score_pos,
);
cnt += 1;
}
}
pos = end;
}
} else {
while pos < exact_end {
let hidden = self.forward_layers(&self.embed_single(ids[pos]), pos, None);
if self
.graph_failed
.swap(false, std::sync::atomic::Ordering::Relaxed)
{
self.cancel
.store(false, std::sync::atomic::Ordering::Relaxed);
self.nll_end();
return Err("GPU graph failed during O(1) NLL prefix".into());
}
if pos >= requested_start && pos < n {
nll += self.nll_from_hidden(&hidden, ids[pos + 1], pos);
cnt += 1;
}
pos += 1;
}
}
self.o1_seal_checked().map_err(|err| {
self.nll_end();
err
})?;
let batch_k = std::env::var("CMF_BATCH_K")
.ok()
.and_then(|v| v.parse::<usize>().ok())
.unwrap_or(0);
let batch_admitted = batch_k > 0
&& self.can_prefill_batched()
&& self.o1_active()
&& std::env::var("CMF_O1_GPU").as_deref() == Ok("1")
&& (0..self.num_layers).all(|li| {
let cache = &self.kv_cache.layers[self.phys_layer(li)];
cache.o1.is_none() || cache.o1_views().is_some()
});
if std::env::var("CMF_GRAPH_PROF").is_ok() {
eprintln!(
"nll-batch: phase=post-seal admission={} requested_k={} scored_rows={}",
batch_admitted,
batch_k,
n.saturating_sub(exact_end),
);
}
let mut batch_completed = false;
if batch_admitted && exact_end < n {
let hs = self.hidden_size;
let mut batch_pos = exact_end;
while batch_pos < n {
let end = (batch_pos + batch_k).min(n);
let bk = end - batch_pos;
let mut hiddens = vec![0.0f32; bk * hs];
for (row, &id) in ids[batch_pos..end].iter().enumerate() {
hiddens[row * hs..(row + 1) * hs].copy_from_slice(&self.embed_single(id));
}
let positions: Vec<usize> = (batch_pos..end).collect();
let t_batch = std::time::Instant::now();
let outcome = self.try_batch_graph_wgpu(&mut hiddens, &positions, bk, None);
if std::env::var("CMF_GRAPH_PROF").is_ok() {
let ms = t_batch.elapsed().as_secs_f64() * 1000.0;
eprintln!(
"nll-batch: phase=post-seal mode=o1 k={bk} pos={}..{} outcome={outcome:?} {ms:.1} ms ({:.1} tok/s)",
batch_pos,
end.saturating_sub(1),
bk as f64 / (ms / 1000.0),
);
}
if let Err(err) = self.nll_check_graph("batch graph", batch_pos) {
self.nll_end();
return Err(err);
}
match outcome {
crate::gpu::BatchGraphOutcome::Completed => {
batch_completed = true;
for row in 0..bk {
nll += self.nll_from_hidden(
&hiddens[row * hs..(row + 1) * hs],
ids[batch_pos + row + 1],
batch_pos + row,
);
cnt += 1;
}
batch_pos = end;
}
crate::gpu::BatchGraphOutcome::Declined => {
if batch_completed {
self.nll_end();
return Err(format!(
"O(1) NLL batch declined after completed chunk at position {batch_pos}"
));
}
break;
}
crate::gpu::BatchGraphOutcome::Failed => {
self.nll_end();
return Err(format!(
"O(1) NLL batch graph failed after admission at position {batch_pos}"
));
}
}
}
if batch_completed && cnt == n.saturating_sub(requested_start) {
self.nll_end();
return Ok((nll, cnt));
}
}
for pos in exact_end..n {
let hidden = self.forward_layers(&self.embed_single(ids[pos]), pos, None);
if self
.graph_failed
.swap(false, std::sync::atomic::Ordering::Relaxed)
{
self.cancel
.store(false, std::sync::atomic::Ordering::Relaxed);
self.nll_end();
return Err(format!(
"GPU graph failed during O(1) NLL serial scoring at position {pos}"
));
}
nll += self.nll_from_hidden(&hidden, ids[pos + 1], pos);
cnt += 1;
}
self.nll_end();
Ok((nll, cnt))
}
pub fn calib_ids(&mut self, ids: &[u32], temps: &[f32]) -> (Vec<bool>, Vec<Vec<f32>>) {
self.clear_sequence_state();
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.clear_sequence_state();
(correct, pmax)
}
pub fn ppl_ids_dynamic(&mut self, ids: &[u32]) -> Result<(f64, usize), String> {
if self.dyn_router.is_none() {
return Ok((self.ppl_ids(ids)?, 0));
}
self.nll_begin()?;
let saved_active = self.dyn_active;
let mut router = self
.dyn_router
.take()
.ok_or_else(|| "dynamic router disappeared before PPL scoring".to_string())?;
router.reset();
self.dyn_phi_seen = 0;
let _ = self.set_active_skill(None);
let result: Result<(f64, usize), String> = (|| {
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);
self.nll_check_graph("dynamic serial forward", pos)?;
let out_of_band = self.graph_logits.take();
let mut logits = match out_of_band {
Some(lg) => lg,
None => {
let normed = inference::rms_norm(
&hidden,
&self.weights.final_norm,
self.rms_eps,
self.norm_style,
);
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;
let tok_nll = lse - logits[target] as f64;
if std::env::var("CMF_PPL_TRACE").is_ok() && pos < 48 {
let top = logits
.iter()
.enumerate()
.max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
.map(|(i, _)| i)
.unwrap_or(0);
eprintln!(
"pos {pos:3} tgt {target:6} nll {tok_nll:7.3} | top1 {top:6} lg[t]={:.2} lg[top]={:.2}",
logits[target], logits[top]
);
}
nll += tok_nll;
cnt += 1;
attention::recycle_buf(&mut logits);
let phi = self.dyn_phi_ema.clone();
if let Some(new_active) = router.step(&phi, pos) {
let _ = self.set_active_skill(new_active);
}
}
Ok(((nll / cnt.max(1) as f64).exp(), router.switches.len()))
})();
let _ = self.set_active_skill(saved_active);
self.dyn_router = Some(router);
self.nll_end();
result
}
pub fn probe_phi(&mut self, ids: &[u32], layer: usize) -> Vec<f32> {
self.clear_sequence_state();
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.clear_sequence_state();
acc
}
fn prefill_batch(&mut self, ids: &[u32], start_pos: usize) -> Vec<f32> {
self.prefill_batch_masked(ids, start_pos, None)
}
fn prefill_batch_masked(
&mut self,
ids: &[u32],
start_pos: usize,
task_mask: Option<&TaskMask>,
) -> Vec<f32> {
self.prefill_batch_span(PrefillIn::Ids(ids), start_pos, task_mask, 0, usize::MAX)
}
fn prefill_batch_span(
&mut self,
input: PrefillIn<'_>,
start_pos: usize,
task_mask: Option<&TaskMask>,
from: usize,
upto_excl: usize,
) -> Vec<f32> {
let hs = self.hidden_size;
let b = match input {
PrefillIn::Ids(ids) => ids.len(),
PrefillIn::Hidden(hb) => hb.len() / hs,
};
let upto_excl = upto_excl.min(self.num_layers);
let mut h: Vec<f32>;
let mut h_ready;
match input {
PrefillIn::Ids(_) => {
h = vec![0.0; b * hs];
h_ready = false;
}
PrefillIn::Hidden(hb) => {
h = hb.to_vec();
h_ready = true;
}
}
let fill_h = |h: &mut Vec<f32>, me: &Self| {
if let PrefillIn::Ids(ids) = input {
for (bi, &id) in ids.iter().enumerate() {
let e = me.embed_single(id);
h[bi * hs..(bi + 1) * hs].copy_from_slice(&e);
}
if let Ok(tp) = std::env::var("CMF_TRACE_POS") {
if let Ok(t) = tp.parse::<usize>() {
if t >= start_pos && t < start_pos + ids.len() {
let bi = t - start_pos;
let row = &h[bi * hs..(bi + 1) * hs];
let n: f32 = row.iter().map(|x| x * x).sum::<f32>().sqrt();
eprintln!(
"BATCH pos {t} embed: id {} |h| = {n:.6} h0 {:.6} h1 {:.6} | b={} start={start_pos} ids[..8]={:?}",
ids[bi],
row[0],
row[1],
ids.len(),
&ids[..ids.len().min(8)]
);
}
}
}
}
};
let (_nkv, _hd, _rd, eps) = (
self.num_kv_heads,
self.head_dim,
self.rotary_dim,
self.rms_eps,
);
let pool = self.pool.clone();
let norm_style = self.norm_style;
let automatic_gpu_prefix = self.automatic_gpu_prefix();
#[cfg(target_os = "macos")]
let mut chunk_skip_until = 0usize;
for li in from..upto_excl {
let _capacity_tail = automatic_gpu_prefix
.filter(|&prefix| li >= prefix)
.map(|_| crate::gpu::enter_cpu_scope());
crate::gpu::set_layer(li as i64); #[cfg(target_os = "macos")]
if task_mask.is_none() {
if li < chunk_skip_until {
continue;
}
if !h_ready && li == 0 && self.weights.embed_tokens.q8_row_parts().is_none() {
fill_h(&mut h, self);
h_ready = true;
}
let ids_for_embed = match input {
PrefillIn::Ids(ids) => (!h_ready && li == 0).then_some(ids),
PrefillIn::Hidden(_) => None,
};
let end = self.chunk_run_gpu(li, &mut h, b, start_pos, ids_for_embed, upto_excl);
if end > li {
h_ready = true;
chunk_skip_until = end;
if self.is_loop_end(end - 1) && end < self.num_layers {
for bi in 0..b {
let normed = inference::rms_norm(
&h[bi * hs..(bi + 1) * hs],
&self.weights.final_norm,
eps,
norm_style,
);
h[bi * hs..(bi + 1) * hs].copy_from_slice(&normed);
}
}
continue;
}
}
if !h_ready {
fill_h(&mut h, self);
h_ready = true;
}
let lw = &self.weights.layers[self.phys_layer(li)];
match &lw.attn {
AttnKind::Kda(w) => {
let cfg = self.kda_cfg.expect("kda layer without kda_cfg");
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 attn = crate::linear_core::kda_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::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::ShortConv(w) => {
let cfg = self
.short_conv_cfg
.expect("short-conv layer without short_conv_cfg");
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 attn = short_conv_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::Mla(w) => {
let inv_freq_l = self.layer_inv_freq(li);
let rs = self.layer_rope_scale(li);
let mut normed = vec![0.0f32; hs];
for bi in 0..b {
inference::rms_norm_into(
&h[bi * hs..(bi + 1) * hs],
&lw.input_norm,
eps,
norm_style,
&mut normed,
);
let ao = mla_attention(
w,
&normed,
&mut self.kv_cache.layers[li],
start_pos + bi,
&inv_freq_l,
rs,
eps,
pool.as_deref(),
);
for (dst, &a) in h[bi * hs..(bi + 1) * hs].iter_mut().zip(&ao) {
*dst += a;
}
}
}
AttnKind::Full {
wq,
wk,
wv,
wo,
q_norm,
k_norm,
output_gate,
softplus_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 inv_freq_l = self.layer_inv_freq(li);
let (nkv_l, hd_l, rd_l) = self.layer_geom(li);
let cfg = QwenAttnCfg {
num_heads: self.layer_num_heads(li),
num_kv_heads: nkv_l,
head_dim: hd_l,
hidden_size: hs,
position: start_pos,
inv_freq: &inv_freq_l,
rotary_dim: rd_l,
scale: self.attn_scale,
softcap: self.attn_softcap,
window: self.layer_window(li),
v_norm: self.attn_v_norm,
qk_norm_after_rope: self.qk_norm_after_rope,
q_norm: q_norm.as_deref(),
k_norm: k_norm.as_deref(),
output_gate: *output_gate,
softplus_gate: softplus_gate
.as_ref()
.map(|(gate, per_head)| (gate, *per_head)),
rope_scale: self.layer_rope_scale(li),
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 mut attn = attention::qwen_attention_batch(
&normed,
b,
wq,
wk,
wv,
wo,
&mut self.kv_cache.layers[li],
&cfg,
);
if let Some(w) = &lw.attn_out_norm {
for bi in 0..b {
inference::rms_norm_into(
&attn[bi * hs..(bi + 1) * hs],
w,
eps,
norm_style,
&mut normed[bi * hs..(bi + 1) * hs],
);
}
attn.copy_from_slice(&normed);
}
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[self.phys_layer(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 mask_row = task_mask
.filter(|m| m.ffn_active_count(li) < self.intermediate_size)
.and_then(|m| m.ffn_masks.get(li))
.map(|v| v.as_slice());
let mut ffn = match &lw.ffn {
FfnKind::Dense(d) if !d.segs.is_empty() => {
tube_ffn(d, &post, b, pool.as_deref(), mask_row)
}
FfnKind::Dense(d) => dense_ffn_batch(d, &post, b, pool.as_deref(), mask_row),
FfnKind::Moe(m) => moe_ffn_batch(m, &post, b, hs, pool.as_deref(), None),
FfnKind::DenseMoe(dm) => {
let mut out = vec![0.0f32; b * hs];
for bi in 0..b {
let r = dense_moe_ffn(
dm,
&post[bi * hs..(bi + 1) * hs],
&h[bi * hs..(bi + 1) * hs],
eps,
norm_style,
pool.as_deref(),
);
out[bi * hs..(bi + 1) * hs].copy_from_slice(&r);
}
out
}
};
if let Some(w) = &lw.ffn_out_norm {
for bi in 0..b {
inference::rms_norm_into(
&ffn[bi * hs..(bi + 1) * hs],
w,
eps,
norm_style,
&mut post[bi * hs..(bi + 1) * hs],
);
}
ffn.copy_from_slice(&post);
}
for (dst, &f) in h.iter_mut().zip(&ffn) {
*dst += f;
}
if let Some(sc) = lw.layer_scale {
for v in h.iter_mut() {
*v *= sc;
}
}
if let Ok(tp) = std::env::var("CMF_TRACE_POS") {
if let Ok(t) = tp.parse::<usize>() {
if t >= start_pos && t < start_pos + b {
let bi = t - start_pos;
let row = &h[bi * hs..(bi + 1) * hs];
let n: f32 = row.iter().map(|x| x * x).sum::<f32>().sqrt();
eprintln!(
"BATCH pos {t} after layer {li}: |h| = {n:.6} h0 {:.6} h1 {:.6}",
row[0], row[1]
);
}
}
}
if std::env::var("CMF_DEBUG_LAYERS").is_ok() {
let row = &h[(b - 1) * hs..b * hs];
let rms =
(row.iter().map(|&v| (v as f64) * (v as f64)).sum::<f64>() / hs as f64).sqrt();
let mx = row.iter().fold(0f32, |m, &v| m.max(v.abs()));
eprintln!(
"layer {li:>3} {:>10} ffn={:<5} rms={rms:>12.4} max={mx:>12.4}",
match &self.weights.layers[self.phys_layer(li)].attn {
AttnKind::LinearGdn(_) => "gdn",
AttnKind::Linear(_) => "vmf",
AttnKind::ShortConv(_) => "conv",
_ => "attn",
},
match &lw.ffn {
FfnKind::Moe(_) => "moe",
FfnKind::Dense(_) => "dense",
FfnKind::DenseMoe(_) => "dense+moe",
},
);
}
if self.is_loop_end(li) && li + 1 < self.num_layers {
for bi in 0..b {
let normed = inference::rms_norm(
&h[bi * hs..(bi + 1) * hs],
&self.weights.final_norm,
eps,
norm_style,
);
h[bi * hs..(bi + 1) * hs].copy_from_slice(&normed);
}
}
if std::env::var("CMF_TRACE_H").is_ok() {
let n = h[..hs].iter().map(|v| v.abs()).sum::<f32>() / hs as f32;
let mx = h[..hs].iter().fold(0.0f32, |a, &v| a.max(v.abs()));
eprintln!(
"layer {li}: mean|h|={n:.4} max|h|={mx:.2} scale={:?}",
lw.layer_scale
);
}
}
crate::gpu::set_layer(-1); self.o1_progress();
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);
}
if self.embed_multiplier != 1.0 {
for v in out.iter_mut() {
*v *= self.embed_multiplier;
}
}
if self.dsv4.is_some() || self.dsv41.is_some() || self.qwen4_exp.is_some() {
let mut v = vec![0.0f32; self.hidden_size.max(1)];
v[0] = id as f32;
return v;
}
if let Some(b) = &self.g3n {
return b.0.extend_embedding(id, &out, self.pool.as_deref());
}
out
}
#[cfg(target_os = "macos")]
fn chunk_run_gpu(
&mut self,
li0: usize,
h: &mut [f32],
b: usize,
pos0: usize,
embed_ids: Option<&[u32]>,
cap: usize,
) -> usize {
if !crate::gpu::enabled_here()
|| std::env::var("CMF_GPU_CHUNK")
.map(|v| v == "0")
.unwrap_or(false)
|| b < 32
|| self.swa.is_some()
|| self.global_attn.is_some()
|| self.o1_active()
|| self.attn_v_norm
|| (self.attn_scale - 1.0 / (self.head_dim as f32).sqrt()).abs() > 1e-9
{
return li0;
}
let Some(model) = self.model.clone() else {
return li0;
};
let inv_freq = self.inv_freq.clone();
let (nh, nkv, hd, hs) = (
self.num_heads,
self.num_kv_heads,
self.head_dim,
self.hidden_size,
);
let loop_end = if self.loop_final_norm {
((li0 / self.physical_layers) + 1) * self.physical_layers
} else {
self.num_layers
};
let mut layers: Vec<crate::gpu_metal::ChunkLayer> = Vec::new();
let mut stored_at: Vec<usize> = Vec::new();
for li in li0..self.num_layers.min(loop_end).min(cap) {
let lw = &self.weights.layers[self.phys_layer(li)];
if lw.attn_out_norm.is_some() || lw.ffn_out_norm.is_some() || lw.layer_scale.is_some() {
break;
}
let AttnKind::Full {
wq,
wk,
wv,
wo,
q_norm,
k_norm,
output_gate: false,
softplus_gate: None,
bias,
} = &lw.attn
else {
break;
};
let FfnKind::Dense(d) = &lw.ffn else { break };
if d.act != Act::Silu || !d.segs.is_empty() {
break;
}
fn cw(t: &QTensor) -> Option<(usize, usize, usize, &[f32])> {
t.q8_row_parts()
.or_else(|| t.q4t_parts().map(|(i, r, c)| (i, r, c, &[][..])))
.or_else(|| t.q4tp_parts().map(|(i, r, c)| (i, r, c, &[][..])))
}
let parts = (
cw(wq),
cw(wk),
cw(wv),
cw(wo),
cw(&d.gate_proj),
cw(&d.up_proj),
cw(&d.down_proj),
);
let (Some(pq), Some(pk), Some(pv), Some(po), Some(pg), Some(pu), Some(pd)) = parts
else {
break;
};
let layer = &self.kv_cache.layers[li];
if layer.mode != crate::kv_cache::KvMode::F32 || layer.o1.is_some() {
break;
}
stored_at.push(layer.head_len(0));
layers.push(crate::gpu_metal::ChunkLayer {
model: &model,
kv_id: self.graph_kv_id,
layer: li,
wq: pq,
wk: pk,
wv: pv,
wo: po,
gate: pg,
up: pu,
down: pd,
input_norm: &lw.input_norm,
post_norm: &lw.post_norm,
bias: bias
.as_ref()
.map(|(a, bb, cc)| (a.as_slice(), bb.as_slice(), cc.as_slice())),
q_norm: q_norm.as_deref(),
k_norm: k_norm.as_deref(),
inv_freq: &inv_freq,
rd: self.rotary_dim,
nh,
nkv,
hd,
hs,
inter: d.gate_proj.rows(),
gemma: matches!(self.norm_style, cortiq_core::NormStyle::Gemma),
late_qk_norm: self.qk_norm_after_rope,
eps: self.rms_eps as f32,
});
}
if layers.is_empty() {
return li0;
}
let row = nkv * hd;
let mut store: Vec<(Vec<f32>, Vec<f32>, Vec<f32>)> = stored_at
.iter()
.map(|&st| (vec![0f32; b * row], vec![0f32; b * row], vec![0f32; st + b]))
.collect();
let mut io: Vec<crate::gpu_metal::ChunkIo> = Vec::with_capacity(layers.len());
for (i, (ok, ov, oi)) in store.iter_mut().enumerate() {
let li = layers[i].layer;
let layer = &self.kv_cache.layers[li];
io.push(crate::gpu_metal::ChunkIo {
cpu_stored: stored_at[i],
cpu_k: (0..nkv).map(|g| layer.head_keys(g)).collect(),
cpu_v: (0..nkv).map(|g| layer.head_values(g)).collect(),
out_k: ok,
out_v: ov,
imp: oi,
});
}
let n_run = layers.len();
let last = layers.last().map(|l| l.layer + 1).unwrap_or(li0);
let ep = embed_ids.and_then(|ids| {
self.weights
.embed_tokens
.q8_row_parts()
.map(|(idx, rows, _c, rs)| crate::gpu_metal::ChunkEmbed {
idx,
rows,
row_scale: rs,
ids,
mult: self.embed_multiplier,
})
});
if embed_ids.is_some() && ep.is_none() {
return li0;
}
if !crate::gpu_metal::chunk_run_gpu(&layers, &mut io, h, b, pos0, ep.as_ref()) {
return li0;
}
drop(io);
drop(layers);
for (i, (ok, ov, oi)) in store.iter().enumerate().take(n_run) {
let li = li0 + i;
let layer = &mut self.kv_cache.layers[li];
for bi in 0..b {
layer.append(
&ok[bi * row..(bi + 1) * row],
&ov[bi * row..(bi + 1) * row],
&[],
);
}
layer.accumulate_imp(oi);
}
last
}
fn layer_is_local(&self, li: usize) -> bool {
if let Some(layers) = &self.sliding_layers {
return layers.get(li).copied().unwrap_or(false);
}
match self.swa {
Some((_, pattern)) => (li + 1) % pattern.max(1) != 0,
None => false,
}
}
fn layer_inv_freq(&self, li: usize) -> std::sync::Arc<Vec<f32>> {
if self.layer_is_local(li) {
if let Some(f) = &self.inv_freq_local {
return f.clone();
}
} else if let Some(f) = &self.inv_freq_global {
return f.clone();
}
self.inv_freq.clone()
}
fn layer_window(&self, li: usize) -> Option<usize> {
self.swa
.and_then(|(w, _)| self.layer_is_local(li).then_some(w))
}
fn layer_num_heads(&self, li: usize) -> usize {
self.attention_heads_per_layer
.as_ref()
.and_then(|v| v.get(li).copied())
.unwrap_or(self.num_heads)
}
fn layer_rope_scale(&self, li: usize) -> f32 {
if self.layer_is_local(li) {
self.rope_scale_local
} else {
self.rope_scale
}
}
fn layer_geom(&self, li: usize) -> (usize, usize, usize) {
if !self.layer_is_local(li) {
if let Some((ghd, gkv)) = self.global_attn {
return (gkv, ghd, ghd);
}
}
(
self.num_kv_heads,
self.head_dim,
if self.layer_is_local(li) {
self.rotary_dim_local.unwrap_or(self.rotary_dim)
} else {
self.rotary_dim
},
)
}
fn forward_layers(
&mut self,
hidden: &[f32],
position: usize,
task_mask: Option<&TaskMask>,
) -> Vec<f32> {
let out = self.forward_layers_upto(hidden, position, task_mask, None);
self.o1_progress();
out
}
pub fn embed_id(&self, id: u32) -> Vec<f32> {
self.embed_single(id)
}
pub fn split_supported(&self) -> Result<(), String> {
if self.dsv4.is_some() {
return Err(
"network split: DeepSeek-V4 runs its own fused stack (not splittable yet)".into(),
);
}
if self.dsv41.is_some() {
return Err(
"network split: DeepSeek-V4.1 owns the shared CED/CSA2 state (not splittable)"
.into(),
);
}
if self.qwen4_exp.is_some() {
return Err(
"network split: Qwen3.8-Flash-Next hyper/QSA stack is not splittable yet".into(),
);
}
if self.g3n.is_some() {
return Err(
"network split: Gemma-3n runs its own AltUp stack (not splittable yet)".into(),
);
}
Ok(())
}
pub fn forward_span(
&mut self,
hidden: &[f32],
position: usize,
from: usize,
upto: usize,
task_mask: Option<&TaskMask>,
) -> Result<Vec<f32>, String> {
self.split_supported()?;
if from > upto || upto >= self.num_layers {
return Err(format!(
"forward_span: layer range {from}..={upto} outside 0..{}",
self.num_layers
));
}
if hidden.len() != self.hidden_size {
return Err(format!(
"forward_span: hidden len {} ≠ hidden_size {}",
hidden.len(),
self.hidden_size
));
}
let out = self.forward_layers_span(hidden, position, task_mask, from, Some(upto));
self.o1_progress();
if self
.graph_failed
.swap(false, std::sync::atomic::Ordering::Relaxed)
{
self.cancel
.store(false, std::sync::atomic::Ordering::Relaxed);
self.clear_sequence_state();
return Err("forward_span: deferred O(1) transition failed".into());
}
Ok(out)
}
pub fn logits_from_hidden(&mut self, hidden: &[f32]) -> Vec<f32> {
let normed = inference::rms_norm(
hidden,
&self.weights.final_norm,
self.rms_eps,
self.norm_style,
);
self.lm_head_forward(&normed)
}
pub fn sample_next(&mut self, logits: &[f32], past_tokens: &[u32]) -> u32 {
sampler::sample_with_scratch(
logits,
&self.sampler_config,
past_tokens,
&mut self.rng,
&mut self.sampler_scratch,
)
}
pub fn reset_session(&mut self) {
self.clear_sequence_state();
}
pub fn prefill_span_ids(
&mut self,
ids: &[u32],
start_pos: usize,
upto: usize,
task_mask: Option<&TaskMask>,
) -> Result<Vec<f32>, String> {
self.split_supported()?;
if upto >= self.num_layers {
return Err(format!(
"prefill_span_ids: upto {upto} outside 0..{}",
self.num_layers
));
}
if self.can_prefill_batched() && !self.graph_prefill_preferred() {
let out =
self.prefill_batch_span(PrefillIn::Ids(ids), start_pos, task_mask, 0, upto + 1);
self.check_o1_progress_failure("prefill_span_ids")?;
Ok(out)
} else {
let hs = self.hidden_size;
let mut out = Vec::with_capacity(ids.len() * hs);
for (i, &id) in ids.iter().enumerate() {
let emb = self.embed_id(id);
out.extend_from_slice(&self.forward_span(
&emb,
start_pos + i,
0,
upto,
task_mask,
)?);
}
Ok(out)
}
}
pub fn prefill_span_hidden(
&mut self,
hidden: &[f32],
start_pos: usize,
from: usize,
upto: usize,
task_mask: Option<&TaskMask>,
) -> Result<Vec<f32>, String> {
self.split_supported()?;
let hs = self.hidden_size;
if hidden.is_empty() || hidden.len() % hs != 0 {
return Err(format!(
"prefill_span_hidden: {} floats is not a multiple of hidden {hs}",
hidden.len()
));
}
if from > upto || upto >= self.num_layers {
return Err(format!(
"prefill_span_hidden: layer range {from}..={upto} outside 0..{}",
self.num_layers
));
}
if self.can_prefill_batched() && !self.graph_prefill_preferred() {
let out = self.prefill_batch_span(
PrefillIn::Hidden(hidden),
start_pos,
task_mask,
from,
upto + 1,
);
self.check_o1_progress_failure("prefill_span_hidden")?;
Ok(out)
} else {
let b = hidden.len() / hs;
let mut out = Vec::with_capacity(hidden.len());
for i in 0..b {
let h = self.forward_span(
&hidden[i * hs..(i + 1) * hs],
start_pos + i,
from,
upto,
task_mask,
)?;
out.extend_from_slice(&h);
}
Ok(out)
}
}
fn try_token_graph_wgpu(
&self,
hidden: &[f32],
position: usize,
logits_out: &mut Vec<f32>,
layers_run: &mut usize,
) -> Option<Result<Vec<f32>, ()>> {
self.try_token_graph_wgpu_steps(
hidden,
position,
logits_out,
1,
None,
Some(layers_run),
0,
self.num_layers,
)
}
fn try_token_graph_wgpu_span(
&self,
hidden: &[f32],
position: usize,
logits_out: &mut Vec<f32>,
from: usize,
upto_excl: usize,
layers_run: &mut usize,
) -> Option<Result<Vec<f32>, ()>> {
self.try_token_graph_wgpu_steps(
hidden,
position,
logits_out,
1,
None,
Some(layers_run),
from,
upto_excl,
)
}
fn try_multi_burst(&self, t_next: u32, position: usize, k: usize) -> Option<Vec<u32>> {
if self.o1_active() || self.attn_softcap > 0.0 {
return None;
}
let graph_on = crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Decode);
if !graph_on || crate::gpu::graph_unsupported() {
return None;
}
let emb = self.embed_single(t_next);
let mut lg = Vec::new();
let mut ids = Vec::new();
match self.try_token_graph_wgpu_steps(
&emb,
position,
&mut lg,
k,
Some(&mut ids),
None,
0,
self.num_layers,
) {
Some(Ok(_)) => {}
Some(Err(())) => {
self.graph_failed
.store(true, std::sync::atomic::Ordering::Relaxed);
return None;
}
None => return None,
}
(ids.len() == k).then_some(ids)
}
fn try_token_graph_wgpu_steps(
&self,
hidden: &[f32],
position: usize,
logits_out: &mut Vec<f32>,
steps: usize,
ids_out: Option<&mut Vec<u32>>,
layers_run: Option<&mut usize>,
from: usize,
upto_excl: usize,
) -> Option<Result<Vec<f32>, ()>> {
let o1_gpu = std::env::var("CMF_O1_GPU").as_deref() == Ok("1");
if (self.o1_active() && !o1_gpu) || self.attn_softcap > 0.0 {
return None;
}
let o1_views: Vec<Option<Vec<crate::nystrom::O1DeviceView<'_>>>> = (from..upto_excl)
.map(|li| {
if !o1_gpu {
return None;
}
self.kv_cache.layers[self.phys_layer(li)].o1_views()
})
.collect();
if self.o1_active() && o1_gpu {
let want: usize = (from..upto_excl)
.filter(|li| self.kv_cache.layers[self.phys_layer(*li)].o1.is_some())
.count();
let have = o1_views.iter().filter(|v| v.is_some()).count();
if want == 0 || have != want {
use std::sync::atomic::{AtomicUsize, Ordering};
static LAST: AtomicUsize = AtomicUsize::new(usize::MAX);
let code = have * 1000 + want;
if LAST.swap(code, Ordering::Relaxed) != code {
tracing::warn!(
"o1 graph: {have} of {want} layers sealed — per-op until all seal"
);
}
return None;
}
}
let nh = self.num_heads;
let (nkv, hd, rd) = self.layer_geom(0);
let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
let mut layers = Vec::with_capacity(upto_excl - from);
let mut model = None;
let dbg = std::env::var("CMF_GRAPH_DEBUG").is_ok();
fn gw(t: &QTensor) -> Option<crate::gpu::GraphW<'_>> {
if let Some((m, i, kind, rs)) = t
.graph_weight()
.or_else(|| t.graph_weight_descriptor())
{
let name = &m.tensors[i].name;
let prism = if crate::prism::is_inverse_embedding(m, name) {
crate::gpu::GraphPrismOp::InverseEmbedding
} else if crate::prism::is_forward_weight(m, name) {
crate::gpu::GraphPrismOp::Forward
} else {
crate::gpu::GraphPrismOp::None
};
return Some(crate::gpu::GraphW {
idx: i,
kind,
row_scale: rs,
data: &[],
prism,
affine: crate::prism::is_affine_target(m, name),
});
}
match t.as_f32() {
Some(d) => Some(crate::gpu::GraphW {
idx: 0,
kind: 4,
row_scale: &[],
data: d,
prism: crate::gpu::GraphPrismOp::None,
affine: false,
}),
None => {
if std::env::var_os("CMF_BATCH_DEBUG").is_some() {
eprintln!("batch graph: weight has no graph/f32 representation");
}
None
}
}
}
for li in from..upto_excl {
let lw = &self.weights.layers[self.phys_layer(li)];
if dbg {
let ak = match &lw.attn {
AttnKind::Mla(_) => "Mla".into(),
AttnKind::Full {
output_gate, bias, ..
} => format!("Full gate={output_gate} bias={}", bias.is_some()),
AttnKind::LinearGdn(_) => "LinearGdn".into(),
AttnKind::Kda(_) => "Kda".into(),
AttnKind::Linear(_) => "Linear".into(),
AttnKind::ShortConv(_) => "ShortConv".into(),
};
let fk = match &lw.ffn {
FfnKind::Dense(_) => "Dense",
FfnKind::Moe(_) => "Moe",
FfnKind::DenseMoe(_) => "DenseMoe",
};
eprintln!("graph L{li}: attn={ak} ffn={fk}");
}
let gffn = match &lw.ffn {
FfnKind::DenseMoe(_) => return None, FfnKind::Dense(d) if !d.segs.is_empty() => return None,
FfnKind::Dense(d) => crate::gpu::GraphFfn::Dense {
gate: gw(&d.gate_proj)?,
up: gw(&d.up_proj)?,
down: gw(&d.down_proj)?,
},
FfnKind::Moe(m) => {
if m.route_tau.is_some() || m.mask.is_some() {
return None;
}
let shared = m.shared.as_ref();
let has_shared = shared.is_some();
let shared_gated = matches!(shared, Some((_, Some(_))));
let sgate = match shared {
Some((_, Some(sg))) => gw(sg)?,
_ => gw(&m.router)?,
};
let router = gw(&m.router)?;
if router.prism != crate::gpu::GraphPrismOp::None
|| sgate.prism != crate::gpu::GraphPrismOp::None
|| router.affine
|| sgate.affine
{
tracing::warn!(
"resident MoE declined: Prism/affine router or shared gate transform is not implemented"
);
return None;
}
let inter = m.experts.first()?.gate_proj.rows();
let mut experts = Vec::with_capacity(m.experts.len() + 1);
let mut q4tp: Option<bool> = None;
let mut gu_q2: Option<bool> = None;
for e in m.experts.iter().chain(shared.map(|(se, _)| se)) {
if !matches!(e.act, Act::Silu)
|| e.gate_proj.rows() != inter
|| e.up_proj.rows() != inter
{
return None;
}
for expert_weight in [&e.gate_proj, &e.up_proj, &e.down_proj] {
let Some((em, ei, _, _)) = expert_weight
.graph_weight()
.or_else(|| expert_weight.graph_weight_descriptor())
else {
return None;
};
let name = &em.tensors[ei].name;
if crate::prism::is_forward_weight(em, name)
|| crate::prism::is_inverse_embedding(em, name)
|| crate::prism::is_affine_target(em, name)
{
tracing::warn!(
"resident MoE declined: expert Prism/affine transform is not implemented"
);
return None;
}
}
let (mm, gi, ui, di, is_p, is_q2) = match e.gate_proj.mapped_q4t() {
Some((mm, gi)) => (
mm,
gi,
e.up_proj.mapped_q4t()?.1,
e.down_proj.mapped_q4t()?.1,
false,
false,
),
None => match e.gate_proj.mapped_q2tp() {
Some((mm, gi)) => (
mm,
gi,
e.up_proj.mapped_q2tp()?.1,
e.down_proj.mapped_q4tp()?.1,
true,
true,
),
None => {
let (mm, gi) = e.gate_proj.mapped_q4tp()?;
(
mm,
gi,
e.up_proj.mapped_q4tp()?.1,
e.down_proj.mapped_q4tp()?.1,
true,
false,
)
}
},
};
if *q4tp.get_or_insert(is_p) != is_p || *gu_q2.get_or_insert(is_q2) != is_q2
{
tracing::warn!(
"MoE layer mixes expert layouts (q4tp={is_p}, q2tp gate/up={is_q2}) — every expert of a layer, INCLUDING the shared one, must share a layout. The whole-token graph declines this layer."
);
return None;
}
model.get_or_insert_with(|| mm.clone());
experts.push((gi, ui, di));
}
crate::gpu::GraphFfn::Moe {
router,
shared_gate: sgate,
experts,
n_exp: m.experts.len(),
top_k: std::env::var("CMF_TOPK_PROBE")
.ok()
.and_then(|v| v.parse::<usize>().ok())
.filter(|k| *k > 0 && *k <= m.top_k)
.unwrap_or(m.top_k),
inter,
norm_topk: m.norm_topk_prob,
q4tp: q4tp?,
gu_q2: gu_q2.unwrap_or(false),
sigmoid: m.router_sigmoid,
bias: m.expert_bias.as_deref(),
has_shared,
shared_gated,
route_scale: m.routed_scaling,
}
}
};
let attn = match &lw.attn {
AttnKind::Full {
wq,
wk,
wv,
wo,
q_norm,
k_norm,
output_gate,
softplus_gate,
bias,
} => {
if softplus_gate.is_some() || self.attention_heads_per_layer.is_some() {
return None;
}
let (m, _, _, _) = wq
.graph_weight()
.or_else(|| wq.graph_weight_descriptor())?;
model = Some(m.clone());
crate::gpu::GraphAttn::Full {
wq: gw(wq)?,
wk: gw(wk)?,
wv: gw(wv)?,
wo: gw(wo)?,
q_norm: q_norm.as_deref(),
k_norm: k_norm.as_deref(),
late_qk_norm: self.qk_norm_after_rope,
bias: bias
.as_ref()
.map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
output_gate: *output_gate,
cpu_k: self.kv_cache.layers[li].k_heads(),
cpu_v: self.kv_cache.layers[li].v_heads(),
}
}
AttnKind::LinearGdn(w) => {
let cfg = self.gdn_cfg?;
let (m, _, _, _) = w
.in_proj_qkv
.graph_weight()
.or_else(|| w.in_proj_qkv.graph_weight_descriptor())?;
model = Some(m.clone());
crate::gpu::GraphAttn::Gdn {
qkv: gw(&w.in_proj_qkv)?,
z: gw(&w.in_proj_z)?,
a: gw(&w.in_proj_a)?,
b: gw(&w.in_proj_b)?,
out: gw(&w.out_proj)?,
conv1d: &w.conv1d,
a_log: &w.a_log,
dt_bias: &w.dt_bias,
norm: &w.norm,
nv: cfg.num_v_heads,
nk: cfg.num_k_heads,
dk: cfg.key_head_dim,
dv: cfg.value_head_dim,
kk: cfg.conv_kernel,
cpu_state: &self.kv_cache.layers[self.phys_layer(li)].linear_state,
}
}
AttnKind::ShortConv(w) => {
let cfg = self.short_conv_cfg?;
let (m, _, _, _) = w
.in_proj
.graph_weight()
.or_else(|| w.in_proj.graph_weight_descriptor())?;
model = Some(m.clone());
crate::gpu::GraphAttn::ShortConv {
inp: gw(&w.in_proj)?,
out: gw(&w.out_proj)?,
taps: &w.conv,
kernel: cfg.kernel,
cpu_state: &self.kv_cache.layers[self.phys_layer(li)].linear_state,
}
}
_ => return None,
};
layers.push(crate::gpu::GraphLayer {
input_norm: &lw.input_norm,
attn,
post_norm: &lw.post_norm,
ffn: gffn,
});
}
let model = model?;
let lm_gw = if upto_excl == self.num_layers
&& self.graph_want_logits
&& std::env::var("CMF_GPU_LMHEAD")
.map(|v| v != "0")
.unwrap_or(true)
{
self.weights
.lm_head
.graph_weight()
.or_else(|| self.weights.lm_head.graph_weight_descriptor())
.map(|(m, i, kind, rs)| {
let name = &m.tensors[i].name;
let prism = if crate::prism::is_inverse_embedding(m, name) {
crate::gpu::GraphPrismOp::InverseEmbedding
} else if crate::prism::is_forward_weight(m, name) {
crate::gpu::GraphPrismOp::Forward
} else {
crate::gpu::GraphPrismOp::None
};
(
crate::gpu::GraphW {
idx: i,
kind,
row_scale: rs,
data: &[],
prism,
affine: crate::prism::is_affine_target(m, name),
},
self.weights.lm_head.rows(),
)
})
} else {
None
};
let lm = lm_gw.as_ref().map(|(gw, rows)| (gw, *rows));
let emb_gw = if steps > 1 {
self.weights
.embed_tokens
.graph_weight()
.or_else(|| self.weights.embed_tokens.graph_weight_descriptor())
.map(|(m, i, kind, rs)| {
let name = &m.tensors[i].name;
let prism = if crate::prism::is_inverse_embedding(m, name) {
crate::gpu::GraphPrismOp::InverseEmbedding
} else if crate::prism::is_forward_weight(m, name) {
crate::gpu::GraphPrismOp::Forward
} else {
crate::gpu::GraphPrismOp::None
};
(
crate::gpu::GraphW {
idx: i,
kind,
row_scale: rs,
data: &[],
prism,
affine: crate::prism::is_affine_target(m, name),
},
self.weights.embed_tokens.rows(),
self.embed_multiplier,
)
})
} else {
None
};
let loop_norm_at: Vec<usize> = if self.loop_final_norm {
(from..upto_excl.min(self.num_layers - 1))
.filter(|&li| (li + 1) % self.physical_layers == 0)
.map(|li| li - from)
.collect()
} else {
Vec::new()
};
let mut h = hidden.to_vec();
let dump_hidden = std::env::var_os("CMF_LOGIT_DUMP").is_some();
let outcome = crate::gpu::forward_token_graph(
&model,
self.graph_kv_id,
&layers,
&o1_views,
self.o1_epoch,
&self.inv_freq,
&mut h,
nh,
nkv,
hd,
self.attn_scale,
rd,
self.hidden_size,
self.intermediate_size,
position,
self.kv_cache.max_seq_len,
gemma,
self.rms_eps as f32,
lm,
&self.weights.final_norm,
logits_out,
&loop_norm_at,
steps,
emb_gw.as_ref().map(|(gw, rows, m)| (gw, *rows, *m)),
ids_out,
layers_run,
from,
dump_hidden,
);
match outcome {
crate::gpu::TokenGraphOutcome::Completed => Some(Ok(h)),
crate::gpu::TokenGraphOutcome::Failed => Some(Err(())),
crate::gpu::TokenGraphOutcome::Declined => None,
}
}
#[cfg(target_os = "macos")]
#[allow(clippy::type_complexity)]
fn metal_rows_plan(
&self,
) -> Option<(
Vec<MetalRowsItem<'_>>,
std::sync::Arc<cortiq_core::CmfModel>,
Option<crate::gpu_metal::GdnGpuCfg>,
)> {
use crate::gpu_metal::{AttnGpuLayer, GdnGpuCfg, GdnGpuLayer, MetalFfn};
let graph_force = crate::gpu::q1_force() || crate::gpu::q2tp_gpu_opt_in();
if !graph_force
|| !crate::gpu::enabled_here()
|| std::env::var("CMF_GPU_BLOCK")
.map(|v| v == "0")
.unwrap_or(false)
|| self.attn_softcap > 0.0
|| self.o1_active()
|| self.swa.is_some()
|| self.global_attn.is_some()
|| self.attention_heads_per_layer.is_some()
|| self.attn_v_norm
|| self.loop_final_norm
{
return None;
}
let attend_contract = self.head_dim % 4 == 0
&& self.head_dim <= 256
&& self.rotary_dim >= 2
&& self.rotary_dim <= self.head_dim
&& (self.rotary_dim / 2) % 32 == 0
&& self.num_kv_heads > 0
&& self.num_heads % self.num_kv_heads == 0;
if !attend_contract {
return None;
}
let mut plan: Vec<MetalRowsItem> = Vec::new();
let mut model_ref: Option<std::sync::Arc<cortiq_core::CmfModel>> = None;
for li in 0..self.num_layers {
let lw = &self.weights.layers[self.phys_layer(li)];
if lw.attn_out_norm.is_some() || lw.ffn_out_norm.is_some() || lw.layer_scale.is_some() {
return None;
}
let ffn = match &lw.ffn {
FfnKind::Dense(d) if d.act == Act::Silu && d.segs.is_empty() => {
let (Some(g), Some(u), Some(dn)) = (
d.gate_proj.metal_graph_parts(),
d.up_proj.metal_graph_parts(),
d.down_proj.metal_graph_parts(),
) else {
return None;
};
MetalFfn::Dense {
gate: g,
up: u,
down: dn,
}
}
_ => return None,
};
match &lw.attn {
AttnKind::LinearGdn(w) if self.gdn_cfg.is_some() => {
let (Some(qkv), Some(z), Some(a), Some(bb), Some(out)) = (
w.in_proj_qkv.metal_graph_parts(),
w.in_proj_z.metal_graph_parts(),
w.in_proj_a.f32_parts(),
w.in_proj_b.f32_parts(),
w.out_proj.metal_graph_parts(),
) else {
return None;
};
if let QTensor::Mapped { model, .. } = &w.in_proj_qkv {
model_ref.get_or_insert_with(|| model.clone());
}
let gl = GdnGpuLayer {
attn_norm: &lw.input_norm,
post_norm: &lw.post_norm,
qkv,
z,
a,
b: bb,
out,
ffn,
conv1d: &w.conv1d,
a_log: &w.a_log,
dt_bias: &w.dt_bias,
gnorm: &w.norm,
};
match plan.last_mut() {
Some(MetalRowsItem::Gdn { run, .. }) => run.push(gl),
_ => plan.push(MetalRowsItem::Gdn {
run: vec![gl],
first: li,
}),
}
}
AttnKind::Full {
wq,
wk,
wv,
wo,
q_norm,
k_norm,
output_gate,
softplus_gate: None,
bias: None,
} => {
let (Some(pq), Some(pk), Some(pv), Some(po)) =
(
wq.metal_graph_parts(),
wk.metal_graph_parts(),
wv.metal_graph_parts(),
wo.metal_graph_parts(),
)
else {
return None;
};
if let QTensor::Mapped { model, .. } = wq {
model_ref.get_or_insert_with(|| model.clone());
}
let cache = &self.kv_cache.layers[li];
if cache.mode != crate::kv_cache::KvMode::F32 || cache.o1.is_some() {
return None;
}
plan.push(MetalRowsItem::Attn {
l: AttnGpuLayer {
attn_norm: &lw.input_norm,
post_norm: &lw.post_norm,
wq: pq,
wk: pk,
wv: pv,
wo: po,
ffn,
},
li,
q_norm: q_norm.as_deref(),
k_norm: k_norm.as_deref(),
output_gate: *output_gate,
});
}
_ => return None,
}
}
let model = model_ref?;
let gcfg = self.gdn_cfg.map(|cfg| GdnGpuCfg {
nv: cfg.num_v_heads,
nk: cfg.num_k_heads,
dk: cfg.key_head_dim,
dv: cfg.value_head_dim,
kk: cfg.conv_kernel,
hidden: self.hidden_size,
inter: self.intermediate_size,
c_dim: cfg.conv_dim(),
eps: cfg.rms_eps as f32,
gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
});
Some((plan, model, gcfg))
}
#[cfg(target_os = "macos")]
#[allow(clippy::too_many_arguments)]
fn metal_attn_params<'a>(
li: usize,
cache: &'a crate::kv_cache::LayerKvCache,
q_norm: Option<&'a [f32]>,
k_norm: Option<&'a [f32]>,
output_gate: bool,
inv_freq: &'a [f32],
geom: (usize, usize, usize, usize),
pos0: usize,
kv_id: u64,
scale: f32,
eps: f32,
gemma: bool,
late_qk_norm: bool,
) -> (crate::gpu_metal::AttnDeviceParams<'a>, usize) {
let (nh, nkv, hd, rd) = geom;
let cpu_k: Vec<&[f32]> = (0..nkv).map(|g| cache.head_keys(g)).collect();
let cpu_v: Vec<&[f32]> = (0..nkv).map(|g| cache.head_values(g)).collect();
let cpu_stored = cpu_k[0].len() / hd;
(
crate::gpu_metal::AttnDeviceParams {
kv_id,
layer: li,
nh,
nkv,
hd,
rd,
position: pos0,
scale,
eps,
gemma,
late_qk_norm,
output_gate,
q_norm,
k_norm,
inv_freq,
cpu_k,
cpu_v,
cpu_stored,
o1: None,
},
cpu_stored,
)
}
#[cfg(target_os = "macos")]
#[allow(clippy::type_complexity)]
fn metal_rows_run(
&mut self,
hiddens: &mut [f32],
pos0: usize,
b: usize,
prefill: bool,
spec: Option<((usize, usize, usize), &[f32], &mut Vec<f32>)>,
mut argmax_out: Option<(usize, &mut Vec<u32>)>,
) -> MetalRowsRun {
use crate::gpu_metal::{GraphDims, VerifyGraph};
if !crate::gpu_metal::wait_replay() {
tracing::error!("Metal rows graph: the pending async replay failed");
return MetalRowsRun::Failed;
}
spec_stamp("v.wait");
let want = self.gdn_cfg.map(|c| c.state_len()).unwrap_or(0);
for l in &mut self.kv_cache.layers {
if l.linear_state.len() != want && want > 0 {
l.linear_state = vec![0f32; want];
}
}
let Some((plan, model, gcfg)) = self.metal_rows_plan() else {
return MetalRowsRun::Declined;
};
spec_stamp("v.plan");
let dims = GraphDims {
hidden: self.hidden_size,
eps: self.rms_eps as f32,
gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
};
let Some(mut graph) = (if prefill {
VerifyGraph::new_prefill(&model, dims, hiddens, b)
} else {
VerifyGraph::new(&model, dims, hiddens, b)
}) else {
return MetalRowsRun::Declined;
};
let geom = (
self.num_heads,
self.num_kv_heads,
self.head_dim,
self.rotary_dim,
);
let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
let eps = self.rms_eps as f32;
let kv_id = self.graph_kv_id;
let inv_freq = self.inv_freq.clone();
for item in &plan {
let ok = match item {
MetalRowsItem::Gdn { run, .. } => gcfg
.as_ref()
.map(|gc| run.iter().all(|l| graph.gdn_ok(l, gc)))
.unwrap_or(false),
MetalRowsItem::Attn {
l,
li,
q_norm,
k_norm,
output_gate,
} => {
let (p, _) = Self::metal_attn_params(
*li,
&self.kv_cache.layers[*li],
*q_norm,
*k_norm,
*output_gate,
&inv_freq,
geom,
pos0,
kv_id,
self.attn_scale,
eps,
gemma,
self.qk_norm_after_rope,
);
graph.attn_ok(l, &p)
}
};
if !ok {
use std::sync::atomic::{AtomicBool, Ordering};
static SAID: AtomicBool = AtomicBool::new(false);
if !SAID.swap(true, Ordering::Relaxed) {
tracing::warn!("metal rows graph: a layer failed preflight — declining");
}
return MetalRowsRun::Declined;
}
}
let lm = match &spec {
Some((lm, _, _)) => {
if !graph.lm_head_ok(*lm) {
return MetalRowsRun::Declined;
}
Some(*lm)
}
None => None,
};
let mut gdn_layers = Vec::new();
let mut attn_layers = Vec::new();
for item in &plan {
match item {
MetalRowsItem::Gdn { run, first } => {
let ro: Vec<&[f32]> = self.kv_cache.layers[*first..*first + run.len()]
.iter()
.map(|l| l.linear_state.as_slice())
.collect();
if !graph.encode_gdn_run_b(run, &ro, gcfg.as_ref().unwrap()) {
return MetalRowsRun::Declined;
}
gdn_layers.extend(*first..*first + run.len());
}
MetalRowsItem::Attn {
l,
li,
q_norm,
k_norm,
output_gate,
} => {
let (p, cpu_stored) = Self::metal_attn_params(
*li,
&self.kv_cache.layers[*li],
*q_norm,
*k_norm,
*output_gate,
&inv_freq,
geom,
pos0,
kv_id,
self.attn_scale,
eps,
gemma,
self.qk_norm_after_rope,
);
if !graph.encode_attn_b(l, &p) {
return MetalRowsRun::Declined;
}
attn_layers.push((*li, cpu_stored));
}
}
}
if let (Some(lm), Some((_, final_norm, _))) = (lm, spec.as_ref()) {
if !graph.encode_lm_head_b(final_norm, lm) {
return MetalRowsRun::Declined;
}
if let Some((n, _)) = argmax_out.as_ref() {
if !graph.encode_argmax_b(*n) {
argmax_out = None;
}
}
}
spec_stamp("v.enc");
if !graph.sync() {
return MetalRowsRun::Failed;
}
spec_stamp("v.gpu");
match (spec, argmax_out) {
(Some(_), Some((_, ids))) => {
ids.resize(b, 0);
if !graph.read_argmax(ids) {
return MetalRowsRun::Failed;
}
spec_stamp("v.am");
}
(Some((lm, _, logits)), None) => {
logits.resize(b * lm.1, 0.0);
if !graph.read_logits(logits) {
return MetalRowsRun::Failed;
}
spec_stamp("v.lg");
}
(None, _) => {}
}
if !graph.read_hidden(hiddens) {
return MetalRowsRun::Failed;
}
spec_stamp("v.hid");
MetalRowsRun::Completed(MetalVerifyPending {
graph,
gdn_layers,
attn_layers,
})
}
#[cfg(target_os = "macos")]
fn try_batch_graph_metal(
&mut self,
hiddens: &mut [f32],
positions: &[usize],
b: usize,
spec: Option<((usize, usize, usize), &[f32], &mut Vec<f32>)>,
argmax_out: Option<(usize, &mut Vec<u32>)>,
) -> crate::gpu::BatchGraphOutcome {
let _t0 = std::time::Instant::now();
if positions.len() != b
|| positions.windows(2).any(|w| w[1] != w[0] + 1)
|| hiddens.len() != b * self.hidden_size
{
return crate::gpu::BatchGraphOutcome::Declined;
}
let pending = match self.metal_rows_run(hiddens, positions[0], b, false, spec, argmax_out) {
MetalRowsRun::Declined => return crate::gpu::BatchGraphOutcome::Declined,
MetalRowsRun::Failed => return crate::gpu::BatchGraphOutcome::Failed,
MetalRowsRun::Completed(pending) => pending,
};
if std::env::var("CMF_GRAPH_SPEC_TIME").is_ok() {
eprintln!(
"metal-verify: {:.1} ms | b={b}",
_t0.elapsed().as_secs_f64() * 1e3
);
}
self.metal_verify = Some(pending);
crate::gpu::BatchGraphOutcome::Completed
}
#[cfg(target_os = "macos")]
fn prefill_rows_metal(
&mut self,
ids: &[u32],
start_pos: usize,
spec: Option<((usize, usize, usize), &[f32], &mut Vec<f32>)>,
) -> MetalPrefillOutcome {
let b = ids.len();
if b == 0 || b > 512 {
return MetalPrefillOutcome::Declined;
}
METAL_PREFILL_CHUNKS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
let with_head = spec.is_some();
let hs = self.hidden_size;
let mut hiddens = vec![0f32; b * hs];
for (j, &id) in ids.iter().enumerate() {
let e = self.embed_single(id);
hiddens[j * hs..(j + 1) * hs].copy_from_slice(&e);
}
let mut pending = match self.metal_rows_run(&mut hiddens, start_pos, b, true, spec, None) {
MetalRowsRun::Declined => return MetalPrefillOutcome::Declined,
MetalRowsRun::Failed => {
METAL_PREFILL_ERRORS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
return MetalPrefillOutcome::Failed;
}
MetalRowsRun::Completed(pending) => pending,
};
let idxs = pending.gdn_layers.clone();
let mut outs: Vec<&mut [f32]> = self
.kv_cache
.layers
.iter_mut()
.enumerate()
.filter(|(i, _)| idxs.binary_search(i).is_ok())
.map(|(_, l)| l.linear_state.as_mut_slice())
.collect();
if !pending.graph.finish_states(&mut outs) {
METAL_PREFILL_ERRORS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
return MetalPrefillOutcome::Failed;
}
let (nkv, hd) = (self.num_kv_heads, self.head_dim);
let mut rows = Vec::with_capacity(pending.attn_layers.len());
for (li, cpu_stored) in &pending.attn_layers {
let mut kbuf = vec![0f32; b * nkv * hd];
let mut vbuf = vec![0f32; b * nkv * hd];
if !crate::gpu_metal::kv_mirror_read_rows(
self.graph_kv_id,
*li,
nkv,
hd,
*cpu_stored,
b,
&mut kbuf,
&mut vbuf,
) {
METAL_PREFILL_ERRORS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
return MetalPrefillOutcome::Failed;
}
rows.push((*li, *cpu_stored, kbuf, vbuf));
}
for (li, cpu_stored, kbuf, vbuf) in rows {
let cache = &mut self.kv_cache.layers[li];
for r in 0..b {
cache.append(
&kbuf[r * nkv * hd..(r + 1) * nkv * hd],
&vbuf[r * nkv * hd..(r + 1) * nkv * hd],
&[],
);
}
crate::gpu_metal::kv_mirror_set_stored(self.graph_kv_id, li, cpu_stored + b);
}
METAL_PREFILL_ROWS.fetch_add(b as u64, std::sync::atomic::Ordering::Relaxed);
if with_head {
METAL_PREFILL_HEAD_ROWS.fetch_add(b as u64, std::sync::atomic::Ordering::Relaxed);
}
MetalPrefillOutcome::Completed(hiddens)
}
#[cfg(target_os = "macos")]
fn prefill_batch_metal(&mut self, ids: &[u32], start_pos: usize) -> MetalPrefillOutcome {
self.prefill_rows_metal(ids, start_pos, None)
}
#[cfg(target_os = "macos")]
fn nll_batch_metal(&mut self, ids: &[u32], start: usize) -> MetalBatchNllOutcome {
if ids.len() < 2 || self.o1_active() || self.head_clusters.is_some() {
return MetalBatchNllOutcome::Declined;
}
let Some(lm) = self.weights.lm_head.metal_graph_parts() else {
return MetalBatchNllOutcome::Declined;
};
let chunk = std::env::var("CMF_METAL_PREFILL_CHUNK")
.ok()
.and_then(|v| v.parse::<usize>().ok())
.filter(|&v| (1..=512).contains(&v))
.unwrap_or(32);
let final_norm = self.weights.final_norm.clone();
let mut nll = 0.0f64;
let mut count = 0usize;
let mut pos = 0usize;
let mut completed = 0usize;
while pos < ids.len() {
let end = (pos + chunk).min(ids.len());
let mut logits = Vec::new();
let outcome = self.prefill_rows_metal(
&ids[pos..end],
pos,
Some((lm, &final_norm, &mut logits)),
);
match outcome {
MetalPrefillOutcome::Declined => {
return if completed == 0 {
MetalBatchNllOutcome::Declined
} else {
MetalBatchNllOutcome::Failed(format!(
"ordinary Metal NLL batch declined after {completed} chunks"
))
};
}
MetalPrefillOutcome::Failed => {
return MetalBatchNllOutcome::Failed(
"ordinary Metal NLL batch failed after admission".to_string(),
);
}
MetalPrefillOutcome::Completed(_) => {}
}
completed += 1;
let vocab = self.vocab_size.min(lm.1);
if logits.len() != (end - pos) * lm.1 || vocab == 0 {
return MetalBatchNllOutcome::Failed(
"ordinary Metal NLL head returned an invalid shape".to_string(),
);
}
for row in 0..(end - pos) {
let absolute = pos + row;
if absolute < start || absolute + 1 >= ids.len() {
continue;
}
let lg = &mut logits[row * lm.1..row * lm.1 + vocab];
if let Some(mu) = self.logit_multiplier {
for v in lg.iter_mut() {
*v *= mu;
}
}
if let Some(c) = self.final_softcap {
for v in lg.iter_mut() {
*v = c * (*v / c).tanh();
}
}
let target = ids[absolute + 1] as usize;
if target >= vocab {
return MetalBatchNllOutcome::Failed(format!(
"target token {target} exceeds Metal head rows {vocab}"
));
}
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;
count += 1;
}
pos = end;
}
MetalBatchNllOutcome::Completed(nll, count)
}
#[cfg(target_os = "macos")]
fn metal_verify_commit(&mut self, a: usize) -> bool {
let Some(mut pending) = self.metal_verify.take() else {
return false;
};
let n = a + 1;
let idxs = pending.gdn_layers.clone();
let mut outs: Vec<&mut [f32]> = self
.kv_cache
.layers
.iter_mut()
.enumerate()
.filter(|(i, _)| idxs.binary_search(i).is_ok())
.map(|(_, l)| l.linear_state.as_mut_slice())
.collect();
if !pending.graph.commit(n, &mut outs) {
return false;
}
spec_stamp("c.replay");
let (nkv, hd) = (self.num_kv_heads, self.head_dim);
let mut rows = Vec::with_capacity(pending.attn_layers.len());
for (li, cpu_stored) in &pending.attn_layers {
let mut kbuf = vec![0f32; n * nkv * hd];
let mut vbuf = vec![0f32; n * nkv * hd];
if !crate::gpu_metal::kv_mirror_read_rows(
self.graph_kv_id,
*li,
nkv,
hd,
*cpu_stored,
n,
&mut kbuf,
&mut vbuf,
) {
return false;
}
rows.push((*li, *cpu_stored, kbuf, vbuf));
}
for (li, cpu_stored, kbuf, vbuf) in rows {
let cache = &mut self.kv_cache.layers[li];
for r in 0..n {
cache.append(
&kbuf[r * nkv * hd..(r + 1) * nkv * hd],
&vbuf[r * nkv * hd..(r + 1) * nkv * hd],
&[],
);
}
crate::gpu_metal::kv_mirror_set_stored(self.graph_kv_id, li, cpu_stored + n);
}
spec_stamp("c.kv");
true
}
#[cfg(target_os = "macos")]
fn mtp_warm_batch_submit(
&mut self,
m: &mut MtpModule,
pairs: &[(&[f32], u32)],
first_pos: usize,
) -> Option<MetalWarmPending> {
use crate::gpu_metal::{AttnDeviceParams, AttnGpuLayer, GraphDims, MetalFfn, VerifyGraph};
let b = pairs.len();
if b == 0 || b > 512 || m.kv.mode != crate::kv_cache::KvMode::F32 || m.kv.o1.is_some() {
return None;
}
let AttnKind::Full {
wq,
wk,
wv,
wo,
q_norm,
k_norm,
output_gate,
softplus_gate: None,
bias: None,
} = &m.layer.attn
else {
return None;
};
let FfnKind::Dense(d) = &m.layer.ffn else {
return None;
};
if !d.segs.is_empty() {
return None;
}
let (Some(pq), Some(pk), Some(pv), Some(po)) =
(wq.q1_parts(), wk.q1_parts(), wv.q1_parts(), wo.q1_parts())
else {
return None;
};
let (Some(g), Some(u), Some(dn)) = (
d.gate_proj.q1_parts(),
d.up_proj.q1_parts(),
d.down_proj.q1_parts(),
) else {
return None;
};
let Some(eh) = m.eh_proj.q1_parts() else {
return None;
};
let QTensor::Mapped { model, .. } = wq else {
return None;
};
let model = model.clone();
let hs = self.hidden_size;
let mut cat = vec![0f32; b * 2 * hs];
for (j, (h, tok)) in pairs.iter().enumerate() {
let e = self.embed_single(*tok);
let (ce, ch) = cat[j * 2 * hs..(j + 1) * 2 * hs].split_at_mut(hs);
inference::rms_norm_into(&e, &m.enorm, self.rms_eps, self.norm_style, ce);
inference::rms_norm_into(h, &m.hnorm, self.rms_eps, self.norm_style, ch);
}
let dims = GraphDims {
hidden: hs,
eps: self.rms_eps as f32,
gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
};
spec_stamp("w.cat");
let Some(mut graph) = VerifyGraph::new_via_proj(&model, dims, eh, &cat, b) else {
return None;
};
spec_stamp("w.new");
let l = AttnGpuLayer {
attn_norm: &m.layer.input_norm,
post_norm: &m.layer.post_norm,
wq: pq,
wk: pk,
wv: pv,
wo: po,
ffn: MetalFfn::Dense {
gate: g,
up: u,
down: dn,
},
};
let (nh, nkv, hd, rd) = (
self.num_heads,
self.num_kv_heads,
self.head_dim,
self.rotary_dim,
);
let inv_freq = self.inv_freq.clone();
let cpu_stored;
{
let cache = &m.kv;
let cpu_k: Vec<&[f32]> = (0..nkv).map(|g| cache.head_keys(g)).collect();
let cpu_v: Vec<&[f32]> = (0..nkv).map(|g| cache.head_values(g)).collect();
cpu_stored = cpu_k[0].len() / hd;
if cpu_stored > first_pos {
spec_stamp("w.decl");
return None;
}
let p = AttnDeviceParams {
kv_id: self.mtp_kv_id(),
layer: Self::MTP_LAYER_BASE,
nh,
nkv,
hd,
rd,
position: first_pos,
scale: self.attn_scale,
eps: self.rms_eps as f32,
gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
late_qk_norm: self.qk_norm_after_rope,
output_gate: *output_gate,
q_norm: q_norm.as_deref(),
k_norm: k_norm.as_deref(),
inv_freq: &inv_freq,
cpu_k,
cpu_v,
cpu_stored,
o1: None,
};
if !graph.attn_ok(&l, &p) || !graph.encode_attn_b(&l, &p) {
return None;
}
}
spec_stamp("w.enc");
if !graph.submit() {
return None;
}
spec_stamp("w.sub");
Some(MetalWarmPending {
graph,
cpu_stored,
b,
})
}
#[cfg(target_os = "macos")]
fn mtp_warm_batch_metal(
&mut self,
m: &mut MtpModule,
pairs: &[(&[f32], u32)],
first_pos: usize,
) -> bool {
match self.mtp_warm_batch_submit(m, pairs, first_pos) {
Some(p) => self.mtp_warm_batch_finish(m, p),
None => false,
}
}
#[cfg(target_os = "macos")]
fn mtp_warm_batch_finish(&mut self, m: &mut MtpModule, pending: MetalWarmPending) -> bool {
let MetalWarmPending {
mut graph,
cpu_stored,
b,
} = pending;
let (nkv, hd) = (self.num_kv_heads, self.head_dim);
if !graph.sync() {
return false;
}
spec_stamp("w.gpu");
let mut kbuf = vec![0f32; b * nkv * hd];
let mut vbuf = vec![0f32; b * nkv * hd];
if !crate::gpu_metal::kv_mirror_read_rows(
self.mtp_kv_id(),
Self::MTP_LAYER_BASE,
nkv,
hd,
cpu_stored,
b,
&mut kbuf,
&mut vbuf,
) {
return false;
}
for r in 0..b {
m.kv.append(
&kbuf[r * nkv * hd..(r + 1) * nkv * hd],
&vbuf[r * nkv * hd..(r + 1) * nkv * hd],
&[],
);
}
crate::gpu_metal::kv_mirror_set_stored(
self.mtp_kv_id(),
Self::MTP_LAYER_BASE,
cpu_stored + b,
);
spec_stamp("w.kv");
true
}
pub(crate) fn note_draft_id(&mut self, id: u32) {
let cut = Self::draft_vocab_rows(usize::MAX).max(131_072);
if (id as usize) >= cut {
self.draft_full_streak = 16;
} else {
self.draft_full_streak = self.draft_full_streak.saturating_sub(1);
}
}
fn draft_head_rows(&self, head_rows: usize) -> usize {
if self.draft_full_streak > 0 {
head_rows
} else {
Self::draft_vocab_rows(head_rows)
}
}
fn draft_vocab_rows(head_rows: usize) -> usize {
static N: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
let n = *N.get_or_init(|| {
std::env::var("CMF_DRAFT_VOCAB")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(65536)
});
if n == 0 { head_rows } else { n.min(head_rows) }
}
#[cfg(target_os = "macos")]
fn mtp_step_metal(
&mut self,
m: &mut MtpModule,
hidden: &[f32],
next_token: u32,
position: usize,
want_logits: bool,
) -> Option<(Vec<f32>, Vec<f32>)> {
use crate::gpu_metal::{AttnDeviceParams, AttnGpuLayer, GraphDims, MetalFfn, TokenGraph};
if std::env::var("CMF_MTP_GRAPH").as_deref() == Ok("0")
|| !crate::gpu::q1_force()
|| !crate::gpu::enabled_here()
|| self.attn_softcap > 0.0
|| self.attention_heads_per_layer.is_some()
|| m.kv.mode != crate::kv_cache::KvMode::F32
|| m.kv.o1.is_some()
{
return None;
}
let AttnKind::Full {
wq,
wk,
wv,
wo,
q_norm,
k_norm,
output_gate,
softplus_gate: None,
bias: None,
} = &m.layer.attn
else {
return None;
};
let FfnKind::Dense(d) = &m.layer.ffn else {
return None;
};
if d.act != Act::Silu || !d.segs.is_empty() {
return None;
}
let (pq, pk, pv, po) = (
wq.q1_parts()?,
wk.q1_parts()?,
wv.q1_parts()?,
wo.q1_parts()?,
);
let (g, u, dn) = (
d.gate_proj.q1_parts()?,
d.up_proj.q1_parts()?,
d.down_proj.q1_parts()?,
);
let QTensor::Mapped { model, .. } = wq else {
return None;
};
let model = model.clone();
let lm = if want_logits {
Some(self.weights.lm_head.q1_parts()?)
} else {
None
};
let dims = GraphDims {
hidden: self.hidden_size,
eps: self.rms_eps as f32,
gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
};
let hs = self.hidden_size;
let mut x = vec![0f32; hs];
let mut graph = TokenGraph::new(&model, dims, &x)?;
let mut folded = false;
if let Some(eh) = m.eh_proj.q1_parts() {
let e = self.embed_single(next_token);
let mut cat = vec![0.0f32; 2 * hs];
let (cat_e, cat_h) = cat.split_at_mut(hs);
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);
folded = graph.encode_input_proj(eh, &cat);
}
if !folded {
x = self.mtp_block_input(m, hidden, next_token);
graph = TokenGraph::new(&model, dims, &x)?;
}
spec_stamp("d.in");
let l = AttnGpuLayer {
attn_norm: &m.layer.input_norm,
post_norm: &m.layer.post_norm,
wq: pq,
wk: pk,
wv: pv,
wo: po,
ffn: MetalFfn::Dense {
gate: g,
up: u,
down: dn,
},
};
let (nh, nkv, hd, rd) = (
self.num_heads,
self.num_kv_heads,
self.head_dim,
self.rotary_dim,
);
let inv_freq = self.inv_freq.clone();
{
let cache = &m.kv;
let cpu_k: Vec<&[f32]> = (0..nkv).map(|g| cache.head_keys(g)).collect();
let cpu_v: Vec<&[f32]> = (0..nkv).map(|g| cache.head_values(g)).collect();
let cpu_stored = cpu_k[0].len() / hd;
let p = AttnDeviceParams {
kv_id: self.mtp_kv_id(),
layer: Self::MTP_LAYER_BASE,
nh,
nkv,
hd,
rd,
position,
scale: self.attn_scale,
eps: self.rms_eps as f32,
gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
late_qk_norm: self.qk_norm_after_rope,
output_gate: *output_gate,
q_norm: q_norm.as_deref(),
k_norm: k_norm.as_deref(),
inv_freq: &inv_freq,
cpu_k,
cpu_v,
cpu_stored,
o1: None,
};
if !graph.attn_device_ok(&l, &p) || !graph.encode_attn_device(&l, &p) {
return None;
}
}
let draft_rows = if let Some(lm) = lm {
self.draft_head_rows(lm.1)
} else {
0
};
if let Some(lm) = lm {
if !graph.lm_head_ok(lm) {
return None;
}
if draft_rows < lm.1 {
if !graph.encode_lm_head_part(&m.final_norm, lm, draft_rows) {
return None;
}
} else {
graph.encode_lm_head(&m.final_norm, lm);
}
}
spec_stamp("d.enc");
if graph.sync_checked().is_err() {
return None;
}
spec_stamp("d.gpu");
let mut logits = Vec::new();
if let Some(lm) = lm {
let n_read = draft_rows.min(lm.1).min(self.vocab_size);
logits = attention::take_buf(n_read);
graph.read_logits(&mut logits);
logits.resize(self.vocab_size, f32::NEG_INFINITY);
}
graph.finish(&mut x);
let mut krow = attention::take_buf(nkv * hd);
let mut vrow = attention::take_buf(nkv * hd);
if crate::gpu_metal::kv_mirror_read_last(
self.mtp_kv_id(),
Self::MTP_LAYER_BASE,
nkv,
hd,
&mut krow,
&mut vrow,
) {
m.kv.append(&krow, &vrow, &[]);
}
attention::recycle_buf(&mut krow);
attention::recycle_buf(&mut vrow);
spec_stamp("d.rd");
Some((logits, x))
}
fn mtp_chain_on() -> bool {
static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*ON.get_or_init(|| std::env::var("CMF_MTP_CHAIN").as_deref() != Ok("0"))
}
#[cfg(target_os = "macos")]
fn mtp_draft_chain_metal(
&mut self,
m: &mut MtpModule,
hidden: &[f32],
t_next: u32,
position: usize,
k: usize,
) -> Result<Vec<u32>, bool> {
use crate::gpu_metal::{AttnDeviceParams, AttnGpuLayer, GraphDims, MetalFfn, TokenGraph};
if k == 0
|| k > 64
|| !Self::mtp_chain_on()
|| std::env::var("CMF_MTP_GRAPH").as_deref() == Ok("0")
|| !crate::gpu::q1_force()
|| !crate::gpu::enabled_here()
|| self.attn_softcap > 0.0
|| self.attention_heads_per_layer.is_some()
|| m.kv.mode != crate::kv_cache::KvMode::F32
|| m.kv.o1.is_some()
|| self.dsv4.is_some()
|| self.dsv41.is_some()
|| self.qwen4_exp.is_some()
|| self.g3n.is_some()
{
return Err(false);
}
let AttnKind::Full {
wq,
wk,
wv,
wo,
q_norm,
k_norm,
output_gate,
softplus_gate: None,
bias: None,
} = &m.layer.attn
else {
return Err(false);
};
let FfnKind::Dense(d) = &m.layer.ffn else {
return Err(false);
};
if d.act != Act::Silu || !d.segs.is_empty() {
return Err(false);
}
let (Some(pq), Some(pk), Some(pv), Some(po)) =
(wq.q1_parts(), wk.q1_parts(), wv.q1_parts(), wo.q1_parts())
else {
return Err(false);
};
let (Some(g), Some(u), Some(dn)) = (
d.gate_proj.q1_parts(),
d.up_proj.q1_parts(),
d.down_proj.q1_parts(),
) else {
return Err(false);
};
let (Some(eh), Some(lm)) = (m.eh_proj.q1_parts(), self.weights.lm_head.q1_parts()) else {
return Err(false);
};
let QTensor::Mapped { model, .. } = wq else {
return Err(false);
};
let model = model.clone();
let QTensor::Mapped {
model: em,
idx: eidx,
dtype: cortiq_core::TensorDtype::Q4TiledP,
..
} = &self.weights.embed_tokens
else {
return Err(false);
};
if !std::sync::Arc::ptr_eq(em, &model)
|| crate::prism::is_inverse_embedding(&model, &model.tensors[*eidx].name)
{
return Err(false);
}
let embed = (
*eidx,
self.weights.embed_tokens.rows(),
self.weights.embed_tokens.cols(),
);
if embed.2 != self.hidden_size || hidden.len() != self.hidden_size {
return Err(false);
}
let dims = GraphDims {
hidden: self.hidden_size,
eps: self.rms_eps as f32,
gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
};
let Some(mut graph) = TokenGraph::new(&model, dims, hidden) else {
return Err(false);
};
if !graph.chain_embed_ok(embed) || !graph.lm_head_ok(lm) {
return Err(false);
}
let l = AttnGpuLayer {
attn_norm: &m.layer.input_norm,
post_norm: &m.layer.post_norm,
wq: pq,
wk: pk,
wv: pv,
wo: po,
ffn: MetalFfn::Dense {
gate: g,
up: u,
down: dn,
},
};
let (nh, nkv, hd, rd) = (
self.num_heads,
self.num_kv_heads,
self.head_dim,
self.rotary_dim,
);
let inv_freq = self.inv_freq.clone();
let draft_rows = self.draft_head_rows(lm.1);
let n_arg = draft_rows.min(lm.1).min(self.vocab_size);
if n_arg == 0 {
return Err(false);
}
let split = std::env::var("CMF_MTP_CHAIN_SPLIT").as_deref() == Ok("1");
let t_chain = std::time::Instant::now();
graph.chain_ids_init(t_next, k);
let cpu_stored;
{
let cache = &m.kv;
let cpu_k: Vec<&[f32]> = (0..nkv).map(|g| cache.head_keys(g)).collect();
let cpu_v: Vec<&[f32]> = (0..nkv).map(|g| cache.head_values(g)).collect();
cpu_stored = cpu_k[0].len() / hd;
for j in 0..k {
if !graph.encode_chain_input(
embed,
j as u32,
&m.enorm,
&m.hnorm,
self.embed_multiplier,
eh,
) {
return Err(false);
}
let p = AttnDeviceParams {
kv_id: self.mtp_kv_id(),
layer: Self::MTP_LAYER_BASE,
nh,
nkv,
hd,
rd,
position: position + j,
scale: self.attn_scale,
eps: self.rms_eps as f32,
gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
late_qk_norm: self.qk_norm_after_rope,
output_gate: *output_gate,
q_norm: q_norm.as_deref(),
k_norm: k_norm.as_deref(),
inv_freq: &inv_freq,
cpu_k: cpu_k.clone(),
cpu_v: cpu_v.clone(),
cpu_stored: cpu_stored + j,
o1: None,
};
if !graph.attn_device_ok(&l, &p) || !graph.encode_attn_device(&l, &p) {
return Err(false);
}
if draft_rows < lm.1 {
if !graph.encode_lm_head_part(&m.final_norm, lm, draft_rows) {
return Err(false);
}
} else {
graph.encode_lm_head(&m.final_norm, lm);
}
if !graph.encode_argmax(n_arg, j as u32 + 1) {
return Err(false);
}
if split {
graph.commit();
}
}
}
let t_enc = t_chain.elapsed();
if graph.sync_checked().is_err() {
return Err(true);
}
if std::env::var_os("CMF_GRAPH_SPEC_TIME").is_some() {
eprintln!(
"mtp-chain: encode {:.1} ms | wait {:.1} ms (k={k}, head rows {draft_rows}{})",
t_enc.as_secs_f64() * 1e3,
(t_chain.elapsed() - t_enc).as_secs_f64() * 1e3,
if split { ", split" } else { "" }
);
}
let mut ids = vec![0u32; k];
if !graph.chain_ids_read(&mut ids) {
return Err(true);
}
let mut kbuf = vec![0f32; k * nkv * hd];
let mut vbuf = vec![0f32; k * nkv * hd];
if !crate::gpu_metal::kv_mirror_read_rows(
self.mtp_kv_id(),
Self::MTP_LAYER_BASE,
nkv,
hd,
cpu_stored,
k,
&mut kbuf,
&mut vbuf,
) {
return Err(true);
}
for r in 0..k {
m.kv.append(
&kbuf[r * nkv * hd..(r + 1) * nkv * hd],
&vbuf[r * nkv * hd..(r + 1) * nkv * hd],
&[],
);
}
Ok(ids)
}
fn try_batch_graph_wgpu(
&self,
hiddens: &mut [f32],
positions: &[usize],
k: usize,
spec: Option<crate::gpu::SpecTail<'_>>,
) -> crate::gpu::BatchGraphOutcome {
let _tb = std::time::Instant::now();
let batch_debug = std::env::var_os("CMF_BATCH_DEBUG").is_some();
if self.attn_softcap > 0.0 {
return crate::gpu::BatchGraphOutcome::Declined; }
let nh = self.num_heads;
let (nkv, hd, rd) = self.layer_geom(0);
let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
fn gw(t: &QTensor) -> Option<crate::gpu::GraphW<'_>> {
if let Some((m, i, kind, rs)) = t
.graph_weight()
.or_else(|| t.graph_weight_descriptor())
{
let name = &m.tensors[i].name;
let prism = if crate::prism::is_inverse_embedding(m, name) {
crate::gpu::GraphPrismOp::InverseEmbedding
} else if crate::prism::is_forward_weight(m, name) {
crate::gpu::GraphPrismOp::Forward
} else {
crate::gpu::GraphPrismOp::None
};
return Some(crate::gpu::GraphW {
idx: i,
kind,
row_scale: rs,
data: &[],
prism,
affine: crate::prism::is_affine_target(m, name),
});
}
if std::env::var_os("CMF_BATCH_DEBUG").is_some() {
eprintln!(
"batch graph: tensor has no graph descriptor/f32 fallback rows={} cols={}",
t.rows(),
t.cols()
);
}
t.as_f32().map(|d| crate::gpu::GraphW {
idx: 0,
kind: 4,
row_scale: &[],
data: d,
prism: crate::gpu::GraphPrismOp::None,
affine: false,
})
}
let built: Option<(
Vec<crate::gpu::GraphLayer<'_>>,
std::sync::Arc<cortiq_core::CmfModel>,
)> = (|| {
let mut layers = Vec::with_capacity(self.num_layers);
let mut model = None;
for li in 0..self.num_layers {
let lw = &self.weights.layers[self.phys_layer(li)];
let gffn = match &lw.ffn {
FfnKind::Dense(d) if !d.segs.is_empty() => {
if batch_debug {
eprintln!("batch graph: dense segmented FFN at layer {li}");
}
return None;
}
FfnKind::Dense(d) => crate::gpu::GraphFfn::Dense {
gate: gw(&d.gate_proj)?,
up: gw(&d.up_proj)?,
down: gw(&d.down_proj)?,
},
FfnKind::Moe(m) => {
if m.route_tau.is_some() || m.mask.is_some() {
return None;
}
let (se, sg) = m.shared.as_ref()?;
let shared_gated = sg.is_some();
let sgate = match sg {
Some(sg) => gw(sg)?,
None => gw(&m.router)?,
};
let router = gw(&m.router)?;
if router.prism != crate::gpu::GraphPrismOp::None
|| router.affine
|| sgate.prism != crate::gpu::GraphPrismOp::None
|| sgate.affine
{
return None;
}
let inter = m.experts.first()?.gate_proj.rows();
let mut experts = Vec::with_capacity(m.experts.len() + 1);
let mut q4tp: Option<bool> = None;
let mut gu_q2: Option<bool> = None;
for e in m.experts.iter().chain(std::iter::once(se)) {
if !matches!(e.act, Act::Silu)
|| e.gate_proj.rows() != inter
|| e.up_proj.rows() != inter
{
return None;
}
let (mm, gi, ui, di, is_p, is_q2) = match e.gate_proj.mapped_q4t() {
Some((mm, gi)) => (
mm,
gi,
e.up_proj.mapped_q4t()?.1,
e.down_proj.mapped_q4t()?.1,
false,
false,
),
None => match e.gate_proj.mapped_q2tp() {
Some((mm, gi)) => (
mm,
gi,
e.up_proj.mapped_q2tp()?.1,
e.down_proj.mapped_q4tp()?.1,
true,
true,
),
None => {
let (mm, gi) = e.gate_proj.mapped_q4tp()?;
(
mm,
gi,
e.up_proj.mapped_q4tp()?.1,
e.down_proj.mapped_q4tp()?.1,
true,
false,
)
}
},
};
if *q4tp.get_or_insert(is_p) != is_p
|| *gu_q2.get_or_insert(is_q2) != is_q2
{
return None;
}
if [gi, ui, di].into_iter().any(|idx| {
mm.tensors
.get(idx)
.is_some_and(|t| {
crate::prism::is_forward_weight(mm, &t.name)
|| crate::prism::is_affine_target(mm, &t.name)
})
}) {
return None;
}
model.get_or_insert_with(|| mm.clone());
experts.push((gi, ui, di));
}
crate::gpu::GraphFfn::Moe {
router,
shared_gate: sgate,
experts,
n_exp: m.experts.len(),
top_k: m.top_k,
inter,
norm_topk: m.norm_topk_prob,
q4tp: q4tp?,
gu_q2: gu_q2.unwrap_or(false),
sigmoid: m.router_sigmoid,
bias: m.expert_bias.as_deref(),
has_shared: true,
shared_gated,
route_scale: m.routed_scaling,
}
}
_ => return None,
};
let attn = match &lw.attn {
AttnKind::Full {
wq,
wk,
wv,
wo,
q_norm,
k_norm,
output_gate,
softplus_gate,
bias,
} => {
if softplus_gate.is_some() || self.attention_heads_per_layer.is_some() {
if batch_debug {
eprintln!(
"batch graph: unsupported Full attention gate at layer {li} softplus={} heads={}",
softplus_gate.is_some(),
self.attention_heads_per_layer.is_some()
);
}
return None;
}
let (m, _, _, _) = wq
.graph_weight()
.or_else(|| wq.graph_weight_descriptor())?;
model = Some(m.clone());
crate::gpu::GraphAttn::Full {
wq: gw(wq)?,
wk: gw(wk)?,
wv: gw(wv)?,
wo: gw(wo)?,
q_norm: q_norm.as_deref(),
k_norm: k_norm.as_deref(),
late_qk_norm: self.qk_norm_after_rope,
bias: bias
.as_ref()
.map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
output_gate: *output_gate,
cpu_k: self.kv_cache.layers[li].k_heads(),
cpu_v: self.kv_cache.layers[li].v_heads(),
}
}
AttnKind::LinearGdn(w) => {
let Some(cfg) = self.gdn_cfg else {
if batch_debug {
eprintln!("batch graph: no GDN config at layer {li}");
}
return None;
};
let (m, _, _, _) = w
.in_proj_qkv
.graph_weight()
.or_else(|| w.in_proj_qkv.graph_weight_descriptor())?;
model = Some(m.clone());
crate::gpu::GraphAttn::Gdn {
qkv: gw(&w.in_proj_qkv)?,
z: gw(&w.in_proj_z)?,
a: gw(&w.in_proj_a)?,
b: gw(&w.in_proj_b)?,
out: gw(&w.out_proj)?,
conv1d: &w.conv1d,
a_log: &w.a_log,
dt_bias: &w.dt_bias,
norm: &w.norm,
nv: cfg.num_v_heads,
nk: cfg.num_k_heads,
dk: cfg.key_head_dim,
dv: cfg.value_head_dim,
kk: cfg.conv_kernel,
cpu_state: &self.kv_cache.layers[self.phys_layer(li)].linear_state,
}
}
_ => return None,
};
layers.push(crate::gpu::GraphLayer {
input_norm: &lw.input_norm,
attn,
post_norm: &lw.post_norm,
ffn: gffn,
});
}
Some((layers, model?))
})();
let Some((layers, model)) = built else {
{
use std::sync::atomic::{AtomicBool, Ordering};
static SAID: AtomicBool = AtomicBool::new(false);
if !SAID.swap(true, Ordering::Relaxed) {
tracing::warn!("batch graph: BUILDER refused (layer weights/kinds)");
}
}
return crate::gpu::BatchGraphOutcome::Declined;
};
if std::env::var("CMF_GRAPH_SPEC_TIME").is_ok() {
eprintln!("batch-build: {:.1} ms", _tb.elapsed().as_secs_f64() * 1e3);
}
crate::gpu::forward_batch_graph(
&model,
self.graph_kv_id,
&layers,
&self.inv_freq,
hiddens,
nh,
nkv,
hd,
rd,
self.hidden_size,
self.intermediate_size,
positions,
self.kv_cache.max_seq_len,
gemma,
self.rms_eps as f32,
self.attn_scale,
k,
&(0..self.num_layers)
.map(|li| self.kv_cache.layers[self.phys_layer(li)].o1_views())
.collect::<Vec<_>>(),
self.o1_epoch,
spec,
)
}
fn draft_probe() -> bool {
static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*ON.get_or_init(|| std::env::var("CMF_DSV4_DRAFT_PROBE").is_ok_and(|v| v != "0"))
}
#[cfg(feature = "gpu")]
fn dsv4_spec_on() -> bool {
static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*ON.get_or_init(|| {
if let Ok(v) = std::env::var("CMF_DSV4_SPEC_RUN") {
return v != "0";
}
std::env::var("CMF_DSV4_SPEC")
.map(|v| v != "0")
.unwrap_or_else(|_| {
crate::gpu_wgpu::DRAFT_RESERVE.load(std::sync::atomic::Ordering::Relaxed) > 0
})
})
}
#[cfg(feature = "gpu")]
fn dsv4_spec_step(
&mut self,
tip_token: u32,
t_next: u32,
next_pos: usize,
max_extra: usize,
drafted: &mut usize,
accepted_ctr: &mut usize,
) -> Option<(Vec<u32>, usize)> {
let t_all = std::time::Instant::now();
if std::env::var("CMF_DSV4_SPEC_TIME").is_ok() {
thread_local! {
static LAST: std::cell::Cell<Option<std::time::Instant>> =
const { std::cell::Cell::new(None) };
}
LAST.with(|l| {
if let Some(prev) = l.get() {
eprintln!(
"между раундами {:.1} мс",
prev.elapsed().as_secs_f64() * 1e3
);
}
l.set(Some(std::time::Instant::now()));
});
}
if std::env::var("CMF_DSV4_SPEC_DEBUG").is_ok() {
eprintln!("spec_step: вход pos={next_pos}");
}
let n_layers = self.dsv4.as_ref().map(|b| b.1.len())?;
let cfg = self.dsv4.as_ref().map(|b| b.2)?;
if self.dspark.is_none() {
let t = crate::dsv4::dspark_targets(&self.dsv4_mtp, &cfg, n_layers);
if t.is_empty() {
return None;
}
crate::dsv4::dspark_arm(&t, cfg.dim);
self.dspark = Some(crate::dsv4::DsparkState::new(
self.dsv4_mtp.len(),
&cfg,
t.len(),
));
}
let targets = crate::dsv4::dspark_targets(&self.dsv4_mtp, &cfg, n_layers);
let pack = crate::dsv4::dspark_pack_get(&self.dsv4_mtp, &cfg);
if pack.is_none() && std::env::var("CMF_DSV4_SPEC_DEBUG").is_ok() {
eprintln!("spec_step: пак не построился (targets {targets:?})");
}
let pack = pack?;
let block = crate::dsv4::dspark_block();
let b_box = self.dsv4.as_mut()?;
let (g, layers, st) = (&b_box.0, &b_box.1, &mut b_box.3);
let ds = self.dspark.as_mut()?;
let dbg = std::env::var("CMF_DSV4_SPEC_DEBUG").is_ok();
if !crate::dsv4::dspark_take(&mut ds.main_hidden) && !ds.have_hidden {
if dbg {
eprintln!("spec_step: нет захвата");
}
return None;
}
ds.have_hidden = true;
let tip_pos = next_pos.checked_sub(1)?;
let draft_started = std::time::Instant::now();
let mut conf = Vec::new();
let props = crate::dsv4::dspark_draft_gpu(
g,
&self.dsv4_mtp,
&cfg,
ds,
pack,
st.kv_id,
tip_token,
tip_pos,
self.pool.as_deref(),
&mut conf,
);
self.dspark_draft_ns += draft_started.elapsed().as_nanos();
*drafted += block;
if props.is_empty() || props[0] != t_next {
if dbg {
eprintln!(
"spec_step: черновик {} (props0={:?} t_next={t_next})",
if props.is_empty() {
"пуст"
} else {
"мимо"
},
props.first()
);
}
return None;
}
let mut k_verify = crate::dsv4::dspark_verify_k()
.min(props.len())
.min(max_extra.saturating_add(1));
let conf_min = {
static M: std::sync::OnceLock<f32> = std::sync::OnceLock::new();
*M.get_or_init(|| {
std::env::var("CMF_DSPARK_CONF_MIN")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(0.0)
})
};
if conf_min > 0.0 && conf.len() >= props.len() {
let mut keep = 1usize;
while keep < k_verify && conf.get(keep).copied().unwrap_or(0.0) >= conf_min {
keep += 1;
}
k_verify = k_verify.min(keep.max(2));
}
if k_verify < 2 {
return None;
}
let mut fed = Vec::with_capacity(k_verify);
fed.push(t_next);
fed.extend_from_slice(&props[1..k_verify]);
let mut argmax = Vec::new();
let mut logits_all = Vec::new();
let mut walked = Vec::new();
let txn = crate::dsv4::dsv4_verify_chunk(
g,
layers,
&cfg,
st,
&fed,
next_pos,
&self.inv_freq,
self.pool.as_deref(),
&targets,
&mut argmax,
&mut logits_all,
&mut walked,
);
if txn.is_none() && dbg {
eprintln!("spec_step: verify отказал");
}
let txn = txn?;
let spec_gpu_end = txn.gpu_end;
let b = fed.len();
let mut accepted = 1usize;
while accepted < b && fed[accepted] == argmax[accepted - 1] {
accepted += 1;
}
if std::env::var("CMF_DSV4_SPEC_FORCE_REJECT").is_ok_and(|v| v != "0") {
accepted = 1;
}
if std::env::var("CMF_DSV4_SPEC_TRACE").is_ok() {
eprintln!("spec@{next_pos}: fed={fed:?} argmax={argmax:?} accepted={accepted}");
}
let t_fin = std::time::Instant::now();
if !crate::dsv4::dsv4_spec_finish(
g,
layers,
&cfg,
st,
txn,
accepted,
&fed,
&self.inv_freq,
self.pool.as_deref(),
) {
tracing::warn!("dsv4: спекулятивный откат не удался — состояние подозрительно");
return None;
}
if std::env::var("CMF_DSV4_SPEC_TIME").is_ok() {
eprintln!(
"finish(k={accepted}): {:.1} мс",
t_fin.elapsed().as_secs_f64() * 1e3
);
}
*accepted_ctr += accepted - 1;
let (hc, dim) = (cfg.hc_mult, cfg.dim);
let dev_caps: Vec<usize> = targets
.iter()
.copied()
.filter(|&t| t < spec_gpu_end)
.collect();
let mut caps_all = vec![0.0f32; dev_caps.len() * b * hc * dim];
if !crate::gpu_wgpu::dsv4_spec_cap_read_all(b, dev_caps.len(), hc * dim, &mut caps_all) {
return None;
}
for t in 0..accepted {
let tip = t + 1 == accepted;
for (slot, &tl) in targets.iter().enumerate() {
if let Some(di) = dev_caps.iter().position(|&d| d == tl) {
let lo = (di * b + t) * hc * dim;
crate::dsv4::dspark_capture(
&caps_all[lo..lo + hc * dim],
&cfg,
slot,
&mut ds.main_hidden,
);
} else if tip
&& crate::dsv4::dspark_peek_slot(slot, dim, {
let lo = slot * dim;
&mut ds.main_hidden[lo..lo + dim]
})
{
} else {
crate::dsv4::dspark_capture(
&walked[t * hc * dim..(t + 1) * hc * dim],
&cfg,
slot,
&mut ds.main_hidden,
);
}
}
crate::dsv4::dspark_ring_append(
g,
&self.dsv4_mtp,
&cfg,
ds,
next_pos + t,
self.pool.as_deref(),
);
}
let row = logits_all[(accepted - 1) * cfg.vocab..accepted * cfg.vocab].to_vec();
self.graph_logits = Some(row);
if std::env::var("CMF_DSV4_TRUNK_PICK_DUMP").is_ok() {
crate::dsv4::trunk_freq_note(&crate::dsv4::pick_tally_take());
crate::dsv4::pick_tally_arm();
}
if std::env::var("CMF_DSV4_SPEC_TIME").is_ok() {
eprintln!(
"spec_step total {:.1} мс (k={accepted})",
t_all.elapsed().as_secs_f64() * 1e3
);
}
Some((fed[1..accepted].to_vec(), next_pos + accepted))
}
fn dspark_probe(&mut self, position: usize, token_id: u32) {
if self.dsv4_mtp.is_empty() || !Self::draft_probe() {
return;
}
let trunk_now = crate::dsv4::pick_tally_take();
crate::dsv4::trunk_freq_note(&trunk_now);
if !trunk_now.is_empty() {
self.dspark_trunk_picks.push(trunk_now);
let keep = crate::dsv4::dspark_block();
if self.dspark_trunk_picks.len() > keep {
self.dspark_trunk_picks.remove(0);
}
}
for p in std::mem::take(&mut self.dspark_pending) {
let Some(i) = position.checked_sub(p.0 + 1) else {
continue;
};
let mut p = p;
if i < p.1.len() {
if p.2 && p.1[i] == token_id {
p.3 = i + 1;
} else {
p.2 = false;
}
if i + 1 < p.1.len() {
self.dspark_pending.push(p);
continue;
}
}
self.dspark_hist.push(p.3);
self.dspark_real.push(token_id);
}
let Some(b) = &mut self.dsv4 else { return };
let (g, layers, cfg) = (&b.0, &b.1, b.2);
let n_layers = layers.len();
if self.dspark.is_none() {
let t = crate::dsv4::dspark_targets(&self.dsv4_mtp, &cfg, n_layers);
if t.is_empty() {
return;
}
eprintln!(
"DSpark: захват со слоёв {t:?}, блок {}",
crate::dsv4::dspark_block()
);
crate::dsv4::dspark_arm(&t, cfg.dim);
self.dspark = Some(crate::dsv4::DsparkState::new(
self.dsv4_mtp.len(),
&cfg,
t.len(),
));
}
let ds = self.dspark.as_mut().unwrap();
if !crate::dsv4::dspark_take(&mut ds.main_hidden) {
return; }
let mut conf = Vec::new();
crate::dsv4::pick_tally_arm();
let draft_started = std::time::Instant::now();
#[cfg(feature = "gpu")]
let gpu_draft = crate::dsv4::dspark_gpu_on();
#[cfg(not(feature = "gpu"))]
let gpu_draft = false;
let props = if gpu_draft {
#[cfg(feature = "gpu")]
{
let kv_id = b.3.kv_id;
match crate::dsv4::dspark_pack_get(&self.dsv4_mtp, &cfg) {
Some(pk) => crate::dsv4::dspark_draft_gpu(
g,
&self.dsv4_mtp,
&cfg,
ds,
pk,
kv_id,
token_id,
position,
self.pool.as_deref(),
&mut conf,
),
None => Vec::new(),
}
}
#[cfg(not(feature = "gpu"))]
Vec::new()
} else {
crate::gpu::cpu_scope(|| {
crate::dsv4::dspark_draft(
g,
&self.dsv4_mtp,
&cfg,
ds,
token_id,
position,
self.pool.as_deref(),
&mut conf,
)
})
};
self.dspark_draft_ns += draft_started.elapsed().as_nanos();
let draft_picks = crate::dsv4::pick_tally_take();
crate::dsv4::dspark_freq_note(&draft_picks);
crate::dsv4::pick_tally_arm();
if !props.is_empty() {
let (tu, tt) = {
let flat: Vec<(usize, Vec<usize>)> = self
.dspark_trunk_picks
.iter()
.flat_map(|v| v.iter().cloned())
.collect();
let mut per: std::collections::HashMap<usize, Vec<usize>> =
std::collections::HashMap::new();
for (li, picks) in flat {
per.entry(li).or_default().extend(picks);
}
let n = per.len().max(1);
let mut u = 0usize;
let mut t = 0usize;
for (_, v) in per {
t += v.len();
u += v.iter().collect::<std::collections::HashSet<_>>().len();
}
(u / n, t / n)
};
let (du, dt) = crate::dsv4::tally_unique(&draft_picks);
self.dspark_exp.push((tu, tt, du, dt));
self.dspark_pending.push((position, props, true, 0));
}
if self.dspark_hist.len() >= 8 && self.dspark_hist.len() % 8 == 0 {
let n = self.dspark_hist.len() as f32;
let mean: f32 = self.dspark_hist.iter().sum::<usize>() as f32 / n;
let block = crate::dsv4::dspark_block();
let mut at = vec![0usize; block + 1];
for &k in &self.dspark_hist {
at[k] += 1;
}
let mut surv = Vec::with_capacity(block);
for i in 1..=block {
let k = at[i..].iter().sum::<usize>() as f32 / n;
surv.push(format!("{k:.2}"));
}
let distinct = self
.dspark_real
.iter()
.collect::<std::collections::HashSet<_>>()
.len();
let (tu, tt, du, dt) = self.dspark_exp.iter().fold((0, 0, 0, 0), |a, b| {
(a.0 + b.0, a.1 + b.1, a.2 + b.2, a.3 + b.3)
});
let m = self.dspark_exp.len().max(1);
eprintln!(
"DSpark: черновиков {}, принято в среднем {mean:.2} из {block} \
(токенов за проход {:.2}), распределение {at:?}, выживание [{}]",
self.dspark_hist.len(),
mean + 1.0,
surv.join(" ")
);
eprintln!(
"DSpark: разных токенов {distinct} из {} (вырожденность), \
эксперты ствол {}/{} на слой за {block} токенов, \
черновик {}/{} за блок, draft {:.2} мс/блок",
self.dspark_real.len(),
tu / m,
tt / m,
du / m,
dt / m,
self.dspark_draft_ns as f64 / self.dspark_exp.len().max(1) as f64 / 1e6
);
}
}
fn forward_layers_upto(
&mut self,
hidden: &[f32],
position: usize,
task_mask: Option<&TaskMask>,
upto: Option<usize>,
) -> Vec<f32> {
if let Some(plan) = self.gpu_plan.clone() {
if upto.is_none() && plan.len() > 1 {
let mut h = hidden.to_vec();
for &(dev, from, upto_incl) in plan.iter() {
h = crate::gpu::with_device(dev, || {
self.forward_layers_span(&h, position, task_mask, from, Some(upto_incl))
});
}
return h;
}
}
self.forward_layers_span(hidden, position, task_mask, 0, upto)
}
pub fn set_gpu_plan(&mut self, devices: Option<&[usize]>) -> Result<(), String> {
self.set_gpu_plan_at(devices, None)
}
pub fn set_gpu_plan_at(
&mut self,
devices: Option<&[usize]>,
at: Option<usize>,
) -> Result<(), String> {
let Some(devs) = devices.filter(|d| d.len() > 1) else {
self.gpu_plan = None;
return Ok(());
};
self.split_supported()?;
let n = self.num_layers;
if devs.len() > n {
return Err(format!("{} devices for {n} layers", devs.len()));
}
if let Some(k) = at {
if k == 0 || k >= n {
return Err(format!("split at {k}: the model has {n} layers"));
}
if devs.len() == 2 {
self.gpu_plan = Some(std::sync::Arc::new(vec![
(devs[0], 0, k - 1),
(devs[1], k, n - 1),
]));
return Ok(());
}
return Err(format!(
"an explicit split point takes exactly 2 devices, got {}",
devs.len()
));
}
let per = n.div_ceil(devs.len());
let mut plan = Vec::with_capacity(devs.len());
let mut from = 0usize;
for &d in devs {
if from >= n {
break;
}
let upto = (from + per - 1).min(n - 1);
plan.push((d, from, upto));
from = upto + 1;
}
self.gpu_plan = Some(std::sync::Arc::new(plan));
Ok(())
}
pub fn gpu_plan(&self) -> Option<Vec<(usize, usize, usize)>> {
self.gpu_plan.as_ref().map(|p| p.as_ref().clone())
}
fn forward_layers_span(
&mut self,
hidden: &[f32],
position: usize,
task_mask: Option<&TaskMask>,
from: usize,
upto: Option<usize>,
) -> Vec<f32> {
debug_assert!(
from == 0
|| (self.dsv4.is_none()
&& self.dsv41.is_none()
&& self.qwen4_exp.is_none()
&& self.g3n.is_none())
);
#[cfg(target_os = "macos")]
if !crate::gpu_metal::wait_replay() {
self.fail_metal_graph("the pending async replay failed before a plain forward");
return vec![0.0; self.hidden_size];
}
if let Some(b) = &mut self.qwen4_exp {
let _ = (task_mask, upto);
let token_id = hidden.first().copied().unwrap_or(0.0) as u32;
let mut logits = Vec::new();
crate::qwen4_exp::forward_token(
&b.0,
&b.1,
&b.2,
&mut b.3,
token_id,
position,
&self.inv_freq,
self.pool.as_deref(),
&mut logits,
true,
);
self.graph_logits = Some(logits);
return vec![0.0; self.hidden_size];
}
if let Some(b) = &mut self.dsv4 {
let _ = (task_mask, upto);
let token_id = hidden.first().copied().unwrap_or(0.0) as u32;
let (g, layers, cfg, st) = (&b.0, &b.1, b.2, &mut b.3);
st.pos = position;
let mut logits = Vec::new();
crate::dsv4::forward_token(
g,
layers,
&cfg,
st,
token_id,
&self.inv_freq,
self.pool.as_deref(),
&mut logits,
);
self.graph_logits = Some(logits);
self.dspark_probe(position, token_id);
return vec![0.0; self.hidden_size];
}
if let Some(b) = &mut self.dsv41 {
let _ = (task_mask, upto);
let token_id = hidden.first().copied().unwrap_or(0.0) as u32;
let mut logits = Vec::new();
crate::dsv41::forward_token(
&b.0,
&b.1,
&b.2,
&mut b.3,
token_id,
position,
self.pool.as_deref(),
&mut logits,
);
self.graph_logits = Some(logits);
return vec![0.0; self.hidden_size];
}
if let Some(b) = &self.g3n {
let _ = (task_mask, upto);
return crate::g3n::g3n_forward(
&b.0,
&b.1,
hidden,
position,
&mut self.kv_cache.layers,
self.num_heads,
self.num_kv_heads,
self.head_dim,
self.pool.as_deref(),
);
}
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 pool = self.pool.clone();
let graph_env = std::env::var("CMF_GPU_WGPU_GRAPH").ok();
let graph_on = match graph_env.as_deref() {
Some("0") => false,
Some("prefill") => false, Some(_) => true,
None => crate::gpu::wgpu_graph_default(),
};
let graph_trusted =
graph_env.is_some() || crate::gpu::wgpu_graph_default() || self.gdn_cfg.is_some();
let race_eligible = graph_on
&& upto.is_none()
&& task_mask.is_none()
&& from == 0
&& !crate::gpu::graph_unsupported();
let mut tail_start = 0usize;
if race_eligible && crate::gpu::graph_race_use_graph(graph_trusted) {
let t_graph = std::time::Instant::now();
let mut lg = Vec::new();
let mut gl = 0usize;
let built = self.try_token_graph_wgpu(hidden, position, &mut lg, &mut gl);
let declined = built.is_none();
let built = match built {
Some(Ok(hh)) => Some(hh),
Some(Err(())) => {
self.clear_sequence_state();
self.graph_failed
.store(true, std::sync::atomic::Ordering::Relaxed);
self.cancel
.store(true, std::sync::atomic::Ordering::Relaxed);
tracing::error!("token graph failed after admission; sequence state cleared");
return vec![0.0; self.hidden_size];
}
None => None,
};
if declined && !self.o1_active() && self.attn_softcap == 0.0 {
crate::gpu::graph_mark_unsupported();
}
graph_note(built.is_some(), gl, self.num_layers);
if let Some(hh) = built {
let dur = t_graph.elapsed();
if std::env::var("CMF_GRAPH_PROF").is_ok() {
eprintln!("graph-call: {:.2} ms total", dur.as_secs_f64() * 1000.0);
}
if gl > 0 && gl < self.num_layers {
h = hh;
tail_start = gl;
} else if graph_trusted || !crate::gpu::graph_race_first_token_hopeless(dur) {
if !graph_trusted {
crate::gpu::graph_race_record(true, dur);
}
if !lg.is_empty() {
lg.resize(self.vocab_size, 0.0);
if let Some(c) = self.final_softcap {
for l in lg.iter_mut() {
*l = c * (*l / c).tanh();
}
}
self.graph_logits = Some(lg);
}
return hh;
}
}
}
let span = from > 0 || upto.is_some();
if span && graph_on && task_mask.is_none() && graph_trusted {
let upto_excl = upto.map_or(self.num_layers, |u| u + 1);
let mut lg = Vec::new();
let mut gl = 0usize;
let span_res =
self.try_token_graph_wgpu_span(hidden, position, &mut lg, from, upto_excl, &mut gl);
let span_res = match span_res {
Some(Ok(hh)) => Some(hh),
Some(Err(())) => {
self.clear_sequence_state();
self.graph_failed
.store(true, std::sync::atomic::Ordering::Relaxed);
self.cancel
.store(true, std::sync::atomic::Ordering::Relaxed);
tracing::error!(
"span token graph failed after admission; sequence state cleared"
);
return vec![0.0; self.hidden_size];
}
None => None,
};
graph_note(span_res.is_some(), gl, upto_excl - from);
if std::env::var("CMF_GPU_DEBUG").is_ok() {
static SEEN: std::sync::atomic::AtomicU32 = std::sync::atomic::AtomicU32::new(0);
if SEEN.fetch_add(1, std::sync::atomic::Ordering::Relaxed) < 4 {
eprintln!(
"span graph: covered {gl} of {} layers [{from}..{upto_excl}) res={}",
upto_excl - from,
span_res.is_some()
);
}
}
if let Some(hh) = span_res {
if gl == upto_excl - from {
if !lg.is_empty() {
lg.resize(self.vocab_size, 0.0);
if let Some(c) = self.final_softcap {
for l in lg.iter_mut() {
*l = c * (*l / c).tanh();
}
}
self.graph_logits = Some(lg);
}
crate::gpu::set_layer(-1);
return hh;
}
h = hh;
tail_start = from + gl;
}
}
let t_race_cpu = (race_eligible && !graph_trusted).then(std::time::Instant::now);
let _host_tail = (tail_start > from).then(crate::gpu::enter_cpu_scope);
let automatic_gpu_prefix = self.automatic_gpu_prefix();
let _prof_layers = crate::cpuprof::time(crate::cpuprof::Slot::Layers);
#[cfg(target_os = "macos")]
let mut gpu_skip_until = 0usize;
for li in tail_start.max(from)..self.num_layers {
let _capacity_tail = automatic_gpu_prefix
.filter(|&prefix| li >= prefix)
.map(|_| crate::gpu::enter_cpu_scope());
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; }
}
#[cfg(target_os = "macos")]
{
if li < gpu_skip_until {
continue;
}
if task_mask.is_none() {
let end = self.q1_graph_gpu(li, upto, position, &mut h);
if self
.graph_failed
.load(std::sync::atomic::Ordering::Relaxed)
{
return vec![0.0; self.hidden_size];
}
if end > li {
gpu_skip_until = end;
if self.is_loop_end(end - 1) && end < self.num_layers {
h = inference::rms_norm(
&h,
&self.weights.final_norm,
self.rms_eps,
self.norm_style,
);
}
continue;
}
}
}
let lw = &self.weights.layers[self.phys_layer(li)];
if let Ok(tp) = std::env::var("CMF_TRACE_POS") {
if tp.parse::<usize>().ok() == Some(position) {
let n: f32 = h.iter().map(|x| x * x).sum::<f32>().sqrt();
eprintln!(
"TRACE pos {position} layer {li}: |h| = {n:.6} h0 {:.6} h1 {:.6}",
h[0], h[1]
);
}
}
let prof = crate::cpuprof::time(crate::cpuprof::Slot::Norms);
inference::rms_norm_into(
&h,
&lw.input_norm,
self.rms_eps,
self.norm_style,
&mut self.ws.n1,
);
drop(prof);
let attn_out = match &lw.attn {
AttnKind::Mla(w) => {
let inv_freq_l = self.layer_inv_freq(li);
let rs = self.layer_rope_scale(li);
let eps = self.rms_eps;
let pool = self.pool.clone();
mla_attention(
w,
&self.ws.n1,
&mut self.kv_cache.layers[li],
position,
&inv_freq_l,
rs,
eps,
pool.as_deref(),
)
}
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::Kda(w) => {
let cfg = self.kda_cfg.expect("kda layer without kda_cfg");
crate::linear_core::kda_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::ShortConv(w) => {
let cfg = self
.short_conv_cfg
.expect("short-conv layer without short_conv_cfg");
short_conv_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,
softplus_gate,
bias,
} if self.kv_cache.layers[li].o1_sealed() => {
let inv_freq_l = self.layer_inv_freq(li);
let (nkv_l, hd_l, rd_l) = self.layer_geom(li);
let cfg = QwenAttnCfg {
num_heads: self.layer_num_heads(li),
num_kv_heads: nkv_l,
head_dim: hd_l,
hidden_size: hs,
position,
inv_freq: &inv_freq_l,
rotary_dim: rd_l,
scale: self.attn_scale,
softcap: self.attn_softcap,
window: None,
v_norm: self.attn_v_norm,
qk_norm_after_rope: self.qk_norm_after_rope,
q_norm: q_norm.as_deref(),
k_norm: k_norm.as_deref(),
output_gate: *output_gate,
softplus_gate: softplus_gate
.as_ref()
.map(|(gate, per_head)| (gate, *per_head)),
rope_scale: self.layer_rope_scale(li),
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,
softplus_gate,
bias,
} => 'attn: {
if graph_on
&& !*output_gate
&& softplus_gate.is_none()
&& self.attention_heads_per_layer.is_none()
&& bias.is_none()
&& task_mask.is_none()
{
let inv_freq_l = self.layer_inv_freq(li);
let (nkv_l, hd_l, rd_l) = self.layer_geom(li);
let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
if let (Some((gm, qi)), Some((_, ki)), Some((_, vi)), Some((_, oi))) = (
wq.mapped_q1(),
wk.mapped_q1(),
wv.mapped_q1(),
wo.mapped_q1(),
) {
let gm = gm.clone();
let mut out = vec![0f32; hs];
let cache = &self.kv_cache.layers[li];
if crate::gpu::attn_dropin(
&gm,
self.graph_kv_id,
li,
&self.ws.n1,
qi,
ki,
vi,
oi,
q_norm.as_deref(),
k_norm.as_deref(),
self.qk_norm_after_rope,
&inv_freq_l,
nh,
nkv_l,
hd_l,
rd_l,
hs,
position,
self.kv_cache.max_seq_len,
gemma,
eps as f32,
cache.k_heads(),
cache.v_heads(),
&mut out,
) {
break 'attn out;
}
}
}
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 inv_freq_l = self.layer_inv_freq(li);
let (nkv_l, hd_l, rd_l) = self.layer_geom(li);
let cfg = QwenAttnCfg {
num_heads: self.layer_num_heads(li),
num_kv_heads: nkv_l,
head_dim: hd_l,
hidden_size: hs,
position,
inv_freq: &inv_freq_l,
rotary_dim: rd_l,
scale: self.attn_scale,
softcap: self.attn_softcap,
window: self.layer_window(li),
v_norm: self.attn_v_norm,
qk_norm_after_rope: self.qk_norm_after_rope,
q_norm: q_norm.as_deref(),
k_norm: k_norm.as_deref(),
output_gate: *output_gate,
softplus_gate: softplus_gate
.as_ref()
.map(|(gate, per_head)| (gate, *per_head)),
rope_scale: self.layer_rope_scale(li),
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,
)
}
}
}
};
let attn_out = match &self.weights.layers[self.phys_layer(li)].attn_out_norm {
Some(w) => inference::rms_norm(&attn_out, w, self.rms_eps, self.norm_style),
None => attn_out,
};
let lw = &self.weights.layers[self.phys_layer(li)];
let prof = crate::cpuprof::time(crate::cpuprof::Slot::Norms);
inference::add_rmsnorm_fused_into(
&mut h,
&attn_out,
&lw.post_norm,
self.rms_eps,
self.norm_style,
&mut self.ws.p1,
);
drop(prof);
let mut attn_out = attn_out;
attention::recycle_buf(&mut attn_out);
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 ffn_out = match (ffn_masked, &lw.ffn) {
(_, FfnKind::Dense(d)) if !d.segs.is_empty() => {
let row = task_mask
.and_then(|tm| tm.ffn_masks.get(li))
.map(|v| v.as_slice());
tube_ffn(d, post_normed, 1, self.pool.as_deref(), row)
}
(true, FfnKind::Dense(d)) => {
let tm = task_mask.unwrap();
let alive = tm.ffn_active_count(li);
let deep = alive * 2 <= self.intermediate_size;
if deep && d.down_proj.sparse_col_ok() && !d.gate_proj.has_prism_contract() {
let active = tm.ffn_active_indices(li);
sparse_ffn_quant(
d,
post_normed,
&active,
self.hidden_size,
self.pool.as_deref(),
)
} else if deep
&& let (Some(g), Some(u), Some(dn)) = (
d.gate_proj.as_f32(),
d.up_proj.as_f32(),
d.down_proj.as_f32(),
)
{
let active = tm.ffn_active_indices(li);
inference::sparse_ffn_forward(
post_normed,
g,
u,
dn,
self.hidden_size,
self.intermediate_size,
&active,
self.pool.as_deref(),
)
} else {
let row = tm.ffn_masks.get(li).map(|v| v.as_slice());
dense_ffn_batch(d, post_normed, 1, self.pool.as_deref(), row)
}
}
(true, FfnKind::Moe(m)) => {
let allowed = task_mask.and_then(|tm| tm.expert_flags(li, m.experts.len()));
ffn_forward(
&lw.ffn,
post_normed,
self.pool.as_deref(),
allowed.as_deref(),
)
}
(true, FfnKind::DenseMoe(dm)) => dense_moe_ffn(
dm,
post_normed,
&h,
self.rms_eps,
self.norm_style,
self.pool.as_deref(),
),
(false, _) => match &lw.ffn {
FfnKind::DenseMoe(dm) => dense_moe_ffn(
dm,
post_normed,
&h,
self.rms_eps,
self.norm_style,
self.pool.as_deref(),
),
_ => {
let allowed = match (&lw.ffn, task_mask) {
(FfnKind::Moe(m), Some(tm)) => tm.expert_flags(li, m.experts.len()),
_ => None,
};
ffn_forward(
&lw.ffn,
post_normed,
self.pool.as_deref(),
allowed.as_deref(),
)
}
},
};
let ffn_out = match &self.weights.layers[self.phys_layer(li)].ffn_out_norm {
Some(w) => inference::rms_norm(&ffn_out, w, self.rms_eps, self.norm_style),
None => ffn_out,
};
for (i, &f) in ffn_out.iter().enumerate() {
h[i] += f;
}
let mut ffn_out = ffn_out;
attention::recycle_buf(&mut ffn_out);
if let Some(sc) = self.weights.layers[self.phys_layer(li)].layer_scale {
for v in h.iter_mut() {
*v *= sc;
}
}
if self.is_loop_end(li) && li + 1 < self.num_layers {
h = inference::rms_norm(
&h,
&self.weights.final_norm,
self.rms_eps,
self.norm_style,
);
}
if self.dyn_phi_layer == Some(li) {
self.update_dyn_phi(&h);
}
}
crate::gpu::set_layer(-1); if let Some(t) = t_race_cpu {
crate::gpu::graph_race_record(false, t.elapsed());
}
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);
if let Some(m) = self.logit_multiplier {
for l in logits.iter_mut() {
*l *= m;
}
}
if let Some(c) = self.final_softcap {
for l in logits.iter_mut() {
*l = c * (*l / c).tanh();
}
}
if let Some(cm) = self.head_clusters.as_ref() {
self.hierarchical_head_logprobs(hidden, cm, &mut logits);
}
logits
}
fn hierarchical_head_logprobs(&self, hidden: &[f32], cm: &[f32], logits: &mut [f32]) {
let h = hidden.len();
let ncl = cm.len() / h.max(1);
if ncl == 0 || logits.len() % ncl != 0 {
return;
}
let cs = logits.len() / ncl;
let mut lc = vec![0.0f32; ncl];
for c in 0..ncl {
let row = &cm[c * h..(c + 1) * h];
let mut s = 0.0f32;
for j in 0..h {
s += row[j] * hidden[j];
}
lc[c] = s;
}
let mx = lc.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
let lse: f32 = mx + lc.iter().map(|v| (v - mx).exp()).sum::<f32>().ln();
for c in 0..ncl {
let blk = &mut logits[c * cs..(c + 1) * cs];
let bm = blk.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
let bl: f32 = bm + blk.iter().map(|v| (v - bm).exp()).sum::<f32>().ln();
let add = lc[c] - lse - bl;
for v in blk.iter_mut() {
*v += add;
}
}
}
pub fn prefill_next_logits(&mut self, ids: &[u32], task_mask: Option<&TaskMask>) -> Vec<f32> {
self.clear_sequence_state();
crate::gpu::graph_race_begin_generation();
if task_mask.is_none() {
self.o1_begin();
}
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);
}
if let Err(err) = self.o1_seal_checked() {
self.o1_fail(err);
}
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],
attn_out_norm: None,
ffn_out_norm: None,
layer_scale: None,
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),
act: Act::Silu,
down_t: None,
segs: Vec::new(),
}),
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,
softplus_gate: None,
},
})
.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,
num_layers, false, vocab_size,
1e-6,
10_000.0,
NormStyle::Qwen,
4096,
SamplerConfig {
seed: Some(42),
..Default::default()
},
)
}
#[inline]
fn mask_bit(row: &[u8], j: usize) -> bool {
(row.get(j >> 3).copied().unwrap_or(0) >> (j & 7)) & 1 != 0
}
fn mask_gain() -> f32 {
static G: std::sync::OnceLock<f32> = std::sync::OnceLock::new();
*G.get_or_init(|| {
std::env::var("CMF_FFN_MASK_GAIN")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(1.0)
})
}
fn zero_masked_cols(g: &mut [f32], rows: usize, inter: usize, row: &[u8]) {
let fill = meanfill().and_then(|(i, v)| {
let li = crate::gpu::cur_layer();
(*i == inter && li >= 0).then(|| &v[li as usize * inter..(li as usize + 1) * inter])
});
for r in 0..rows {
let base = r * inter;
for (bi, &byte) in row.iter().enumerate() {
if byte == 0xFF {
continue;
}
let j0 = bi * 8;
for bit in 0..8 {
let j = j0 + bit;
if j < inter && byte & (1 << bit) == 0 {
g[base + j] = fill.map_or(0.0, |f| f[j]);
}
}
}
}
let gain = mask_gain();
if gain != 1.0 {
for v in g[..rows * inter].iter_mut() {
*v *= gain;
}
}
}
#[inline]
fn tube_bit(row: Option<&[u8]>, i: usize) -> bool {
row.is_none_or(|r| mask_bit(r, i))
}
fn all_bits_on(row: &[u8], n: usize) -> bool {
(0..n).all(|i| mask_bit(row, i))
}
fn tube_topk() -> usize {
static K: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
*K.get_or_init(|| {
std::env::var("CMF_TUBE_TOPK")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(0)
})
}
fn tube_score_oracle() -> bool {
static O: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*O.get_or_init(|| std::env::var("CMF_TUBE_SCORE").is_ok_and(|v| v == "oracle"))
}
fn tube_ffn_routed(
d: &DenseFfn,
xs: &[f32],
b: usize,
pool: Option<&Pool>,
mask_row: Option<&[u8]>,
k: usize,
) -> Vec<f32> {
let hidden = d.down_proj.rows();
let core = d.gate_proj.rows();
let core_full = mask_row.is_none_or(|r| all_bits_on(r, core));
let mut out = match (b, core_full, mask_row) {
(1, true, _) => dense_ffn(d, xs, pool),
(1, false, Some(row)) => dense_ffn_masked(d, xs, pool, row),
(_, true, _) => dense_ffn_batch(d, xs, b, pool, None),
(_, false, row) => dense_ffn_batch(d, xs, b, pool, row),
};
let cand: Vec<usize> = (0..d.segs.len())
.filter(|&i| tube_bit(mask_row, d.segs[i].start))
.collect();
if cand.is_empty() {
return out;
}
let oracle = tube_score_oracle();
let mut acts: Vec<Vec<f32>> = Vec::with_capacity(cand.len());
let mut scores = vec![0f32; b * cand.len()];
for (ci, &i) in cand.iter().enumerate() {
let seg = &d.segs[i];
let w = seg.width;
let mut g = vec![0.0f32; b * w];
if b == 1 {
seg.gate.matvec(xs, &mut g, pool);
} else {
seg.gate.matmat(xs, b, &mut g, pool);
}
for v in g.iter_mut() {
*v = Act::Silu.combine(*v, 1.0);
}
if !oracle {
for t in 0..b {
scores[t * cand.len() + ci] =
g[t * w..(t + 1) * w].iter().map(|v| v * v).sum::<f32>();
}
}
if oracle || b > 1 {
let mut u = vec![0.0f32; b * w];
if b == 1 {
seg.up.matvec(xs, &mut u, pool);
} else {
seg.up.matmat(xs, b, &mut u, pool);
}
for (a, &v) in g.iter_mut().zip(u.iter()) {
*a *= v;
}
if oracle {
for t in 0..b {
scores[t * cand.len() + ci] =
g[t * w..(t + 1) * w].iter().map(|v| v * v).sum::<f32>();
}
}
}
acts.push(g);
}
let keep = k.min(cand.len());
let mut scratch: Vec<f32> = Vec::new();
for t in 0..b {
let mut sc: Vec<(f32, usize)> = (0..cand.len())
.map(|ci| (scores[t * cand.len() + ci], ci))
.collect();
sc.sort_unstable_by(|x, y| y.0.total_cmp(&x.0));
let mut alive = vec![false; cand.len()];
for &(_, ci) in sc.iter().take(keep) {
alive[ci] = true;
}
if b > 1 {
for (ci, a) in acts.iter_mut().enumerate() {
if !alive[ci] {
let w = d.segs[cand[ci]].width;
a[t * w..(t + 1) * w].fill(0.0);
}
}
} else {
for (ci, &i) in cand.iter().enumerate() {
if !alive[ci] {
continue;
}
let seg = &d.segs[i];
let w = seg.width;
let g = &mut acts[ci];
if !tube_score_oracle() {
scratch.clear();
scratch.resize(w, 0.0);
seg.up.matvec(xs, &mut scratch, pool);
for (a, &v) in g.iter_mut().zip(scratch.iter()) {
*a *= v;
}
}
let mut acc = vec![0.0f32; hidden];
seg.down.matvec(g, &mut acc, pool);
for (o, a) in out.iter_mut().zip(&acc) {
*o += *a;
}
}
}
}
if b > 1 {
for (ci, &i) in cand.iter().enumerate() {
let seg = &d.segs[i];
let mut acc = vec![0.0f32; b * hidden];
seg.down.matmat(&acts[ci], b, &mut acc, pool);
for (o, a) in out.iter_mut().zip(&acc) {
*o += *a;
}
}
}
out
}
fn tube_ffn(
d: &DenseFfn,
xs: &[f32],
b: usize,
pool: Option<&Pool>,
mask_row: Option<&[u8]>,
) -> Vec<f32> {
if tube_topk() > 0 {
return tube_ffn_routed(d, xs, b, pool, mask_row, tube_topk());
}
let hidden = d.down_proj.rows();
let core = d.gate_proj.rows();
let core_full = mask_row.is_none_or(|r| all_bits_on(r, core));
let mut out = match (b, core_full, mask_row) {
(1, true, _) => dense_ffn(d, xs, pool),
(1, false, Some(row)) => dense_ffn_masked(d, xs, pool, row),
(_, true, _) => dense_ffn_batch(d, xs, b, pool, None),
(_, false, row) => dense_ffn_batch(d, xs, b, pool, row),
};
TUBE_SCRATCH.with(|sc| {
let mut sc = sc.borrow_mut();
let [g, u, acc] = &mut *sc;
for seg in &d.segs {
if !tube_bit(mask_row, seg.start) {
continue;
}
let w = seg.width;
g.resize(b * w, 0.0);
if b == 1
&& d.act == Act::Silu
&& QTensor::matvec_silu_mul(&seg.gate, &seg.up, xs, g, pool)
{
} else {
u.resize(b * w, 0.0);
if b == 1 {
QTensor::matvec_many([&seg.gate, &seg.up], xs, [g, u], pool);
} else {
seg.gate.matmat(xs, b, g, pool);
seg.up.matmat(xs, b, u, pool);
}
for i in 0..b * w {
g[i] = d.act.combine(g[i], u[i]);
}
}
acc.resize(b * hidden, 0.0);
acc.fill(0.0);
if b == 1 {
seg.down.matvec(g, acc, pool);
} else {
seg.down.matmat(g, b, acc, pool);
}
for (o, a) in out.iter_mut().zip(acc.iter()) {
*o += *a;
}
}
out
})
}
thread_local! {
static TUBE_SCRATCH: std::cell::RefCell<[Vec<f32>; 3]> =
const { std::cell::RefCell::new([Vec::new(), Vec::new(), Vec::new()]) };
}
fn dense_ffn_batch(
d: &DenseFfn,
xs: &[f32],
b: usize,
pool: Option<&Pool>,
mask_row: Option<&[u8]>,
) -> Vec<f32> {
let inter = d.gate_proj.rows();
let hidden = d.down_proj.rows();
if mask_row.is_none()
&& d.act == Act::Silu
&& b >= 32
&& crate::gpu::enabled_here()
&& !crate::gpu::mm_killed()
&& refit_dir().is_none()
&& !ffn_probe_active()
{
if let (Some((model, w1)), Some((_, w3)), Some((_, w2))) = (
d.gate_proj.mapped_q4t(),
d.up_proj.mapped_q4t(),
d.down_proj.mapped_q4t(),
) {
let mut out = vec![0.0f32; b * hidden];
if crate::gpu::q4t_ffn(model, w1, w3, w2, xs, b, hidden, inter, &mut out) {
return out;
}
}
if let (Some((model, w1)), Some((_, w3)), Some((_, w2))) = (
d.gate_proj.mapped_q4tp(),
d.up_proj.mapped_q4tp(),
d.down_proj.mapped_q4tp(),
) {
let mut out = vec![0.0f32; b * hidden];
if crate::gpu::q4tp_ffn(model, w1, w3, w2, xs, b, hidden, inter, &mut out) {
return out;
}
}
}
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);
if gate_topk() > 0 && d.act == Act::Silu {
for t in 0..b {
let row = &mut g[t * inter..(t + 1) * inter];
for v in row.iter_mut() {
*v = Act::Silu.combine(*v, 1.0);
}
keep_top_k(row, gate_topk());
}
for i in 0..b * inter {
g[i] *= u[i];
}
} else {
for i in 0..b * inter {
g[i] = d.act.combine(g[i], u[i]);
}
}
if let Some(row) = mask_row {
zero_masked_cols(&mut g, b, inter, row);
}
if oracle_topk() > 0 {
for t in 0..b {
keep_top_k(&mut g[t * inter..(t + 1) * inter], oracle_topk());
}
}
let mut out = vec![0.0f32; b * hidden];
d.down_proj.matmat(&g, b, &mut out, pool);
if refit_dir().is_some() {
let li = crate::gpu::cur_layer();
if li >= 0 {
refit_accumulate(li as usize, &g, b, inter, &out, hidden, pool);
}
}
FFN_PROBE.with(|pr| {
if let Some(acc) = pr.borrow_mut().as_mut() {
let li = crate::gpu::cur_layer();
if li < 0 {
return;
}
let Some(row) = acc.get_mut(li as usize) else {
return;
};
let sq = probe_sq();
for t in 0..b {
for (a, &v) in row.iter_mut().zip(&g[t * inter..(t + 1) * inter]) {
*a += if sq {
(v as f64) * (v as f64)
} else {
(v as f64).abs()
};
}
}
}
});
out
}
fn accumulate_act(m: &MoeFfn, xs: &[f32], b: usize) {
static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
static DUMP: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
let on = *ON.get_or_init(|| std::env::var("CMF_RMS_TRACE").is_ok());
let dump = *DUMP.get_or_init(|| std::env::var("CMF_ACT_DUMP").is_ok());
if (!on && !dump) || b == 0 {
return;
}
let hidden = xs.len() / b;
if on {
let mut acc = m.act_sq.borrow_mut();
if acc.len() < hidden {
acc.resize(hidden, 0.0);
}
for t in 0..b {
let row = &xs[t * hidden..(t + 1) * hidden];
for (a, &v) in acc.iter_mut().zip(row) {
*a += (v as f64) * (v as f64);
}
}
}
if dump {
let cap: usize = std::env::var("CMF_ACT_DUMP_ROWS")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(4096);
let mut rows = m.act_rows.borrow_mut();
if rows.len() < cap * hidden {
let take = b.min((cap * hidden - rows.len()) / hidden.max(1));
rows.extend_from_slice(&xs[..take * hidden]);
}
}
}
#[derive(Clone, Copy)]
struct SendVecs(*mut Vec<f32>);
unsafe impl Send for SendVecs {}
unsafe impl Sync for SendVecs {}
impl SendVecs {
#[inline]
fn at(self, i: usize) -> *mut Vec<f32> {
unsafe { self.0.add(i) }
}
}
fn moe_ffn_batch(
m: &MoeFfn,
xs: &[f32],
b: usize,
hidden: usize,
pool: Option<&Pool>,
allowed: Option<&[bool]>,
) -> Vec<f32> {
accumulate_act(m, xs, b);
let ne = m.experts.len();
let mut logits = vec![0.0f32; b * ne];
match &m.resonance {
Some(r) => {
let hdim = xs.len() / b.max(1);
for bi in 0..b {
r.scores(
&xs[bi * hdim..(bi + 1) * hdim],
&mut logits[bi * ne..(bi + 1) * ne],
);
}
}
None => 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 (idx, p, wsum) = moe_route(&logits[bi * ne..(bi + 1) * ne], m, allowed);
for &e in &idx {
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 run_expert = |d: &DenseFfn, list: &[(usize, f32)], out: &mut [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, None);
for (k, &(bi, w)) in list.iter().enumerate() {
for i in 0..hidden {
out[bi * hidden + i] += w * eo[k * hidden + i];
}
}
};
let active: Vec<usize> = (0..ne).filter(|&e| !assign[e].is_empty()).collect();
if pool.is_some() && active.len() >= 8 {
let mut panels: Vec<Vec<f32>> = vec![Vec::new(); active.len()];
{
let panel_ptr = SendVecs(panels.as_mut_ptr());
let experts = &m.experts;
let (active_r, assign_r) = (&active, &assign);
let run = |start: usize, end: usize| {
for ai in start..end {
let e = active_r[ai];
let list = &assign_r[e];
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]);
}
unsafe {
*panel_ptr.at(ai) = dense_ffn_batch(&experts[e], &sub, sb, None, None);
}
}
};
match pool {
Some(p) => p.run_rows(active.len(), &run),
None => run(0, active.len()),
}
}
for (ai, &e) in active.iter().enumerate() {
for (k, &(bi, w)) in assign[e].iter().enumerate() {
let eo = &panels[ai][k * hidden..(k + 1) * hidden];
for i in 0..hidden {
out[bi * hidden + i] += w * eo[i];
}
}
}
} else {
for &e in &active {
run_expert(&m.experts[e], &assign[e], &mut out);
}
}
if let Some((se, gate)) = &m.shared {
let all: Vec<(usize, f32)> = if let Some(gate) = gate {
let mut gl = vec![0.0f32; b];
gate.matmat(xs, b, &mut gl, pool);
(0..b)
.map(|bi| (bi, 1.0 / (1.0 + (-gl[bi]).exp())))
.collect()
} else {
(0..b).map(|bi| (bi, 1.0)).collect()
};
run_expert(se, &all, &mut out);
}
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 gate_topk() > 0
&& let Some(out) = dense_ffn_dynamic(d, x, pool, gate_topk())
{
return out;
}
let prism_body = d.gate_proj.has_prism_contract()
|| d.up_proj.has_prism_contract()
|| d.down_proj.has_prism_contract();
if !prism_body
&& crate::gpu::enabled_here()
&& (d.gate_proj.rows() >= crate::gpu::min_rows() || d.gate_proj.is_q1())
{
let arm = if d.gate_proj.is_q1() && crate::gpu::q1_force() {
crate::gpu::ProbeArm::Gpu
} else {
crate::gpu::probe_arm(crate::gpu::OpClass::Ffn)
};
match arm {
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::probe_note_decline(crate::gpu::OpClass::Ffn);
}
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);
if gate_topk() > 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] = Act::Silu.combine(g[i], 1.0);
}
keep_top_k(g, gate_topk());
for i in 0..inter {
g[i] *= u[i];
}
} else if d.act == Act::Silu && {
let _prof = crate::cpuprof::time(crate::cpuprof::Slot::FfnGateUp);
QTensor::matvec_silu_mul(&d.gate_proj, &d.up_proj, x, g, pool)
} {
} else {
u.resize(inter, 0.0);
let _prof = crate::cpuprof::time(crate::cpuprof::Slot::FfnGateUp);
QTensor::matvec_many([&d.gate_proj, &d.up_proj], x, [g, u], pool);
for i in 0..inter {
g[i] = d.act.combine(g[i], u[i]);
}
}
FFN_PROBE.with(|pr| {
if let Some(acc) = pr.borrow_mut().as_mut() {
let li = crate::gpu::cur_layer();
if li >= 0 {
if let Some(row) = acc.get_mut(li as usize) {
match probe_topk() {
0 if probe_sq() => {
for (a, &v) in row.iter_mut().zip(g.iter()) {
*a += (v as f64) * (v as f64);
}
}
0 if probe_signed() => {
for (a, &v) in row.iter_mut().zip(g.iter()) {
*a += v as f64;
}
}
0 => {
for (a, &v) in row.iter_mut().zip(g.iter()) {
*a += (v as f64).abs();
}
}
k => {
let n = g.len();
let k = k.min(n);
let mut mag: Vec<f32> = g.iter().map(|v| v.abs()).collect();
let (_, kth, _) = mag.select_nth_unstable_by(k - 1, |a, b| {
b.partial_cmp(a).unwrap_or(std::cmp::Ordering::Equal)
});
let thr = *kth;
for (a, &v) in row.iter_mut().zip(g.iter()) {
if v.abs() >= thr {
*a += 1.0;
}
}
}
}
}
}
}
});
if oracle_topk() > 0 {
keep_top_k(g, oracle_topk());
}
{
let li = crate::gpu::cur_layer();
if li >= 0 {
adump_row(li as usize, g);
}
}
let mut out = attention::take_buf(d.down_proj.rows());
let _prof = crate::cpuprof::time(crate::cpuprof::Slot::FfnDown);
d.down_proj.matvec(g, &mut out, pool);
out
})
}
pub struct RefitAcc {
pub support: Vec<u32>,
pub gss: Vec<f32>,
pub ya: Vec<f32>,
pub hidden: usize,
pub tokens: u64,
pub buf_g: Vec<f32>,
pub buf_o: Vec<f32>,
pub buf_t: usize,
}
type RefitState = (std::collections::HashMap<usize, RefitAcc>, Vec<f32>);
static REFIT: std::sync::OnceLock<Option<(String, std::sync::Mutex<RefitState>)>> =
std::sync::OnceLock::new();
fn ffn_probe_active() -> bool {
FFN_PROBE.with(|p| p.borrow().is_some())
}
fn refit_dir() -> Option<&'static (String, std::sync::Mutex<RefitState>)> {
REFIT
.get_or_init(|| {
std::env::var("CMF_FFN_REFIT").ok().map(|d| {
(
d,
std::sync::Mutex::new((std::collections::HashMap::new(), Vec::new())),
)
})
})
.as_ref()
}
fn refit_accumulate(
li: usize,
g: &[f32],
b: usize,
inter: usize,
out: &[f32],
hidden: usize,
pool: Option<&Pool>,
) {
let Some((dir, map)) = refit_dir() else {
return;
};
static SPAN: std::sync::OnceLock<(usize, usize)> = std::sync::OnceLock::new();
let (from, to) = *SPAN.get_or_init(|| {
let g = |k: &str, d: usize| {
std::env::var(k)
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(d)
};
(
g("CMF_FFN_REFIT_FROM", 0),
g("CMF_FFN_REFIT_TO", usize::MAX),
)
});
if li < from || li > to {
return;
}
let mut guard = map.lock().unwrap();
let (map, shared) = &mut *guard;
let acc = match map.entry(li) {
std::collections::hash_map::Entry::Occupied(e) => e.into_mut(),
std::collections::hash_map::Entry::Vacant(e) => {
let path = format!("{dir}/support.{li}.u32");
let Ok(bytes) = std::fs::read(&path) else {
eprintln!("refit: no {path} — layer {li} skipped");
return;
};
let n = u32::from_le_bytes(bytes[0..4].try_into().unwrap()) as usize;
let support: Vec<u32> = bytes[4..4 + n * 4]
.chunks_exact(4)
.map(|c| u32::from_le_bytes([c[0], c[1], c[2], c[3]]))
.collect();
eprintln!(
"refit: layer {li} support {n} ({:.0} MB of accumulator)",
(n * n + hidden * n) as f64 * 4.0 / 1e6
);
e.insert(RefitAcc {
gss: vec![0.0; n * n],
ya: vec![0.0; hidden * n],
buf_g: Vec::new(),
buf_o: Vec::new(),
buf_t: 0,
support,
hidden,
tokens: 0,
})
}
};
let ns = acc.support.len();
let cap = refit_batch();
if acc.buf_g.is_empty() {
acc.buf_g = vec![0.0; ns * cap];
acc.buf_o = vec![0.0; hidden * cap];
}
let take = b.min(cap - acc.buf_t);
for t in 0..take {
let col = acc.buf_t + t;
for (j, &n) in acc.support.iter().enumerate() {
acc.buf_g[j * cap + col] = g[t * inter + n as usize];
}
for h in 0..hidden {
acc.buf_o[h * cap + col] = out[t * hidden + h];
}
}
acc.buf_t += take;
acc.tokens += take as u64;
if acc.buf_t < cap {
return;
}
let bt = acc.buf_t;
acc.buf_t = 0;
let RefitAcc {
gss,
ya,
buf_g,
buf_o,
..
} = acc;
let need = (ns * ns).max(hidden * ns);
if shared.len() < need {
shared.resize(need, 0.0);
}
let scratch = &mut shared[..];
let _ = bt;
if crate::gpu::gemm_nt_f32_transient(buf_g, buf_g, &mut scratch[..ns * ns], ns, cap, ns) {
add_into(gss, &scratch[..ns * ns], pool);
if crate::gpu::gemm_nt_f32_transient(
buf_o,
buf_g,
&mut scratch[..hidden * ns],
hidden,
cap,
ns,
) {
add_into(ya, &scratch[..hidden * ns], pool);
} else {
accum_outer_t(ya, hidden, ns, cap, buf_o, buf_g, pool);
}
} else {
accum_outer_t(gss, ns, ns, cap, buf_g, buf_g, pool);
accum_outer_t(ya, hidden, ns, cap, buf_o, buf_g, pool);
}
}
fn refit_batch() -> usize {
static B: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
*B.get_or_init(|| {
std::env::var("CMF_FFN_REFIT_BATCH")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(4096)
})
}
fn accum_outer_t(
c: &mut [f32],
m: usize,
n: usize,
b: usize,
left: &[f32],
right: &[f32],
pool: Option<&Pool>,
) {
let ptr = SendMut(c.as_mut_ptr());
let body = |i: usize| {
let ptr = &ptr;
let row = unsafe { std::slice::from_raw_parts_mut(ptr.0.add(i * n), n) };
for t in 0..b {
let a = left[i * b + t];
if a == 0.0 {
continue;
}
for (j, o) in row.iter_mut().enumerate() {
*o += a * right[j * b + t];
}
}
};
match pool {
Some(p) if m > 1 => p.run_rows(m, &|s, e| {
for i in s..e {
body(i);
}
}),
_ => {
for i in 0..m {
body(i);
}
}
}
}
fn add_into(dst: &mut [f32], src: &[f32], pool: Option<&Pool>) {
let n = dst.len().min(src.len());
match pool {
Some(p) if n >= 1 << 16 => {
let ptr = SendMut(dst.as_mut_ptr());
let f = |s: usize, e: usize| {
let ptr = &ptr;
for blk in s..e {
let (a, b) = (blk * 4096, ((blk + 1) * 4096).min(n));
for i in a..b {
unsafe { *ptr.0.add(i) += src[i] };
}
}
};
p.run_rows(n.div_ceil(4096), &f);
}
_ => {
for (d, v) in dst.iter_mut().zip(&src[..n]) {
*d += *v;
}
}
}
}
fn accum_outer(
c: &mut [f32],
m: usize,
n: usize,
b: usize,
left: &[f32],
right: &[f32],
pool: Option<&Pool>,
) {
const TILE: usize = 32;
let tiles = m.div_ceil(TILE);
let cp = SendMut(c.as_mut_ptr());
let body = |ti: usize| {
let cp = &cp;
let i0 = ti * TILE;
let i1 = (i0 + TILE).min(m);
for t in 0..b {
let r = &right[t * n..t * n + n];
for i in i0..i1 {
let a = left[i * b + t];
if a == 0.0 {
continue;
}
let row = unsafe { std::slice::from_raw_parts_mut(cp.0.add(i * n), n) };
for (o, v) in row.iter_mut().zip(r) {
*o += a * *v;
}
}
}
};
match pool {
Some(p) if tiles > 1 => p.run_rows(tiles, &|s, e| {
for ti in s..e {
body(ti);
}
}),
_ => {
for ti in 0..tiles {
body(ti);
}
}
}
}
pub fn refit_flush() -> usize {
let Some((dir, map)) = refit_dir() else {
return 0;
};
let guard = map.lock().unwrap();
let mut n = 0;
for (li, acc) in guard.0.iter() {
let w = |name: &str, v: &[f32]| {
let path = format!("{dir}/{name}.{li}.f32");
let bytes: Vec<u8> = v.iter().flat_map(|x| x.to_le_bytes()).collect();
match std::fs::write(&path, &bytes) {
Ok(()) => {}
Err(e) => eprintln!(
"refit: FAILED to write {path} ({} MB): {e}",
bytes.len() / 1_000_000
),
}
};
w("gss", &acc.gss);
w("ya", &acc.ya);
println!(
"refit L{li}: {} support, {} tokens, hidden {}",
acc.support.len(),
acc.tokens,
acc.hidden
);
n += 1;
}
n
}
fn adump_row(li: usize, g: &[f32]) {
use std::io::Write as _;
static FILES: std::sync::OnceLock<
Option<(
String,
std::sync::Mutex<std::collections::HashMap<usize, std::fs::File>>,
)>,
> = std::sync::OnceLock::new();
let Some((prefix, map)) = FILES
.get_or_init(|| {
std::env::var("CMF_FFN_ADUMP")
.ok()
.map(|p| (p, std::sync::Mutex::new(std::collections::HashMap::new())))
})
.as_ref()
else {
return;
};
static SPAN: std::sync::OnceLock<(usize, usize)> = std::sync::OnceLock::new();
let (from, to) = *SPAN.get_or_init(|| {
let g = |k: &str, d: usize| {
std::env::var(k)
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(d)
};
(
g("CMF_FFN_ADUMP_FROM", 0),
g("CMF_FFN_ADUMP_TO", usize::MAX),
)
});
if li < from || li > to {
return;
}
let mut map = map.lock().unwrap();
let f = map.entry(li).or_insert_with(|| {
std::fs::File::create(format!("{prefix}.{li}.f16")).expect("adump file")
});
let mut bytes = Vec::with_capacity(g.len() * 2);
for v in g {
bytes.extend_from_slice(&cortiq_core::quant::f32_to_f16(*v).to_le_bytes());
}
let _ = f.write_all(&bytes);
}
fn oracle_topk() -> usize {
static K: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
*K.get_or_init(|| {
std::env::var("CMF_FFN_ORACLE_TOPK")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(0)
})
}
fn gate_topk() -> usize {
static K: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
*K.get_or_init(|| {
std::env::var("CMF_FFN_GATE_TOPK")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(0)
})
}
fn gate_block() -> usize {
static B: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
*B.get_or_init(|| {
std::env::var("CMF_FFN_GATE_BLOCK")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(1)
})
}
fn keep_top_blocks(g: &mut [f32], keep_n: usize, block: usize) {
let n = g.len();
let nb = n.div_ceil(block);
let kb = (keep_n.div_ceil(block)).clamp(1, nb);
if kb >= nb {
return;
}
let mut score: Vec<f32> = (0..nb)
.map(|b| {
g[b * block..((b + 1) * block).min(n)]
.iter()
.map(|v| v * v)
.sum::<f32>()
})
.collect();
let mut ord = score.clone();
let (_, kth, _) = ord.select_nth_unstable_by(kb - 1, |a, b| {
b.partial_cmp(a).unwrap_or(std::cmp::Ordering::Equal)
});
let thr = *kth;
for b in 0..nb {
if score[b] < thr {
g[b * block..((b + 1) * block).min(n)].fill(0.0);
}
}
score.clear();
}
fn keep_top_k(g: &mut [f32], k: usize) {
if gate_block() > 1 {
return keep_top_blocks(g, k, gate_block());
}
let n = g.len();
if k == 0 || k >= n {
return;
}
let mut mag: Vec<f32> = g.iter().map(|v| v.abs()).collect();
let (_, kth, _) = mag.select_nth_unstable_by(k - 1, |a, b| {
b.partial_cmp(a).unwrap_or(std::cmp::Ordering::Equal)
});
let thr = *kth;
for v in g.iter_mut() {
if v.abs() < thr {
*v = 0.0;
}
}
}
fn probe_sq() -> bool {
static S: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*S.get_or_init(|| std::env::var("CMF_FFN_PROBE_SQ").is_ok())
}
fn probe_signed() -> bool {
static S: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*S.get_or_init(|| std::env::var("CMF_FFN_PROBE_SIGNED").is_ok())
}
fn meanfill() -> Option<&'static (usize, Vec<f32>)> {
static M: std::sync::OnceLock<Option<(usize, Vec<f32>)>> = std::sync::OnceLock::new();
M.get_or_init(|| {
let p = std::env::var("CMF_FFN_MEANFILL").ok()?;
let b = std::fs::read(&p).ok()?;
let inter = u32::from_le_bytes(b[4..8].try_into().ok()?) as usize;
let vals: Vec<f32> = b[8..]
.chunks_exact(4)
.map(|c| f32::from_le_bytes([c[0], c[1], c[2], c[3]]))
.collect();
eprintln!("meanfill: {} value(s), inter {inter}", vals.len());
Some((inter, vals))
})
.as_ref()
}
fn probe_topk() -> usize {
static K: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
*K.get_or_init(|| {
std::env::var("CMF_FFN_PROBE_TOPK")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(0)
})
}
thread_local! {
static FFN_PROBE: std::cell::RefCell<Option<Vec<Vec<f64>>>> =
const { std::cell::RefCell::new(None) };
}
fn dense_ffn_dynamic(d: &DenseFfn, x: &[f32], pool: Option<&Pool>, k: usize) -> Option<Vec<f32>> {
if d.gate_proj.has_prism_contract()
|| d.up_proj.has_prism_contract()
|| d.down_proj.has_prism_contract()
{
return None;
}
let dt = d.down_t.as_ref()?;
let inter = d.gate_proj.rows();
let hidden = dt.cols();
if k == 0 || k >= inter || d.act != Act::Silu {
return None;
}
DYN_SCRATCH.with(|sc| {
let mut sc = sc.borrow_mut();
let DynScratch {
g,
mag,
live,
parts,
} = &mut *sc;
g.resize(inter, 0.0);
d.gate_proj.matvec(x, g, pool);
for v in g.iter_mut() {
*v = inference::silu(*v);
}
mag.clear();
mag.extend(g.iter().map(|v| v.abs()));
let (_, kth, _) = mag.select_nth_unstable_by(k - 1, |a, b| {
b.partial_cmp(a).unwrap_or(std::cmp::Ordering::Equal)
});
let thr = *kth;
live.clear();
live.extend((0..inter as u32).filter(|&n| g[n as usize].abs() >= thr));
let mut out = vec![0.0f32; hidden];
match pool {
Some(p) if live.len() >= 64 => {
let nw = p.n_workers() + 1;
parts.clear();
parts.resize(nw * hidden, 0.0);
let ptr = SendMut(parts.as_mut_ptr());
let n = live.len();
let live_ref: &[u32] = live;
let g_ref: &[f32] = g;
p.run(&|w, workers| {
let chunk = n.div_ceil(workers);
let (s, e) = (w * chunk, ((w + 1) * chunk).min(n));
if s >= e {
return;
}
WORKER_SCRATCH.with(|ws| {
let mut ws = ws.borrow_mut();
let [scratch, acc] = &mut *ws;
scratch.resize(hidden.max(x.len()), 0.0);
acc.clear();
acc.resize(hidden, 0.0);
for (o, &nrm) in live_ref[s..e].iter().enumerate() {
if let Some(&nx) = live_ref[s..e].get(o + 1) {
d.up_proj.prefetch_row(nx as usize);
dt.prefetch_row(nx as usize);
}
let idx = nrm as usize;
let up = d.up_proj.row_dot(idx, x, scratch);
let a = g_ref[idx] * up;
if a != 0.0 {
dt.add_row_scaled(idx, a, acc, scratch);
}
}
for (j, v) in acc.iter().enumerate() {
unsafe { *ptr.at(w * hidden + j) = *v };
}
});
});
for w in 0..nw {
for (j, o) in out.iter_mut().enumerate() {
*o += parts[w * hidden + j];
}
}
}
_ => {
WORKER_SCRATCH.with(|ws| {
let mut ws = ws.borrow_mut();
let [scratch, _acc] = &mut *ws;
scratch.resize(hidden.max(x.len()), 0.0);
for &nrm in live.iter() {
let idx = nrm as usize;
let up = d.up_proj.row_dot(idx, x, scratch);
let a = g[idx] * up;
if a != 0.0 {
dt.add_row_scaled(idx, a, &mut out, scratch);
}
}
});
}
}
Some(out)
})
}
struct DynScratch {
g: Vec<f32>,
mag: Vec<f32>,
live: Vec<u32>,
parts: Vec<f32>,
}
thread_local! {
static DYN_SCRATCH: std::cell::RefCell<DynScratch> = const {
std::cell::RefCell::new(DynScratch {
g: Vec::new(),
mag: Vec::new(),
live: Vec::new(),
parts: Vec::new(),
})
};
static WORKER_SCRATCH: std::cell::RefCell<[Vec<f32>; 2]> =
const { std::cell::RefCell::new([Vec::new(), Vec::new()]) };
}
fn dense_ffn_masked(d: &DenseFfn, x: &[f32], pool: Option<&Pool>, mask_row: &[u8]) -> 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);
if d.act == Act::Silu && QTensor::matvec_silu_mul(&d.gate_proj, &d.up_proj, x, g, pool) {
} else {
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] = d.act.combine(g[i], u[i]);
}
}
zero_masked_cols(g, 1, inter, mask_row);
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.has_prism_contract()
|| d.up_proj.has_prism_contract()
|| d.down_proj.has_prism_contract()
{
return None;
}
if d.act != Act::Silu {
return None;
}
if d.gate_proj.rows() < crate::gpu::min_rows() && !d.gate_proj.is_q1() {
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)]
#[allow(clippy::type_complexity)]
pub(crate) fn moe_parts(
t: &QTensor,
) -> Option<(
&std::sync::Arc<cortiq_core::CmfModel>,
usize,
usize,
usize,
&[f32],
&[f32],
bool,
bool,
bool,
)> {
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, false, false, false,
)),
QTensor::Mapped {
model,
idx,
dtype: cortiq_core::TensorDtype::Q1,
rows,
cols,
..
} => Some((
model,
*idx,
*rows,
*cols,
&[][..],
&[][..],
true,
false,
false,
)),
QTensor::Mapped {
model,
idx,
dtype: cortiq_core::TensorDtype::Q4Tiled,
rows,
cols,
..
} => Some((
model,
*idx,
*rows,
*cols,
&[][..],
&[][..],
false,
true,
false,
)),
QTensor::Mapped {
model,
idx,
dtype: cortiq_core::TensorDtype::Q4TiledP,
rows,
cols,
..
} => Some((
model,
*idx,
*rows,
*cols,
&[][..],
&[][..],
false,
true,
false,
)),
QTensor::Mapped {
model,
idx,
dtype: cortiq_core::TensorDtype::Q2TiledP,
rows,
cols,
..
} => Some((
model,
*idx,
*rows,
*cols,
&[][..],
&[][..],
false,
true,
true,
)),
_ => None,
}
}
#[cfg(target_os = "macos")]
fn metal_moe_graph_parts(m: &MoeFfn, hidden: usize) -> Option<crate::gpu::GpuMoe<'_>> {
if m.router_input_norm
|| m.route_tau.is_some()
|| m.mask.is_some()
|| m.per_expert_scale.is_some()
|| m.experts.is_empty()
|| m.top_k == 0
|| m.resonance.is_some()
{
return None;
}
let (sh, sg) = match &m.shared {
Some((sh, sg)) => (sh, sg.as_ref()),
None => return None,
};
let (rf, rr, rc) = m.router.f32_parts()?;
if rr != m.experts.len() || rc != hidden {
return None;
}
let shared_gated = sg.is_some();
let sf = match sg {
Some(sg) => {
let (sf, sr, sc) = sg.f32_parts()?;
if sr * sc != hidden {
return None;
}
sf
}
None => &rf[..hidden],
};
if let Some(b) = &m.expert_bias {
if b.len() != m.experts.len() {
return None;
}
}
let inter = m.experts[0].gate_proj.rows();
let gu_q2 = m.experts[0].gate_proj.mapped_q2tp().is_some();
let trio = |e: &DenseFfn| -> Option<(usize, usize, usize)> {
if e.act != Act::Silu
|| e.gate_proj.rows() != inter
|| e.gate_proj.cols() != hidden
|| e.up_proj.rows() != inter
|| e.up_proj.cols() != hidden
|| e.down_proj.rows() != hidden
|| e.down_proj.cols() != inter
{
return None;
}
let pick = |t: &QTensor| -> Option<usize> {
if gu_q2 {
t.mapped_q2tp().map(|(_, i)| i)
} else {
t.mapped_q4tp().map(|(_, i)| i)
}
};
Some((
pick(&e.gate_proj)?,
pick(&e.up_proj)?,
e.down_proj.mapped_q4tp().map(|(_, i)| i)?,
))
};
let experts = m.experts.iter().map(trio).collect::<Option<Vec<_>>>()?;
let shared = trio(sh)?;
Some(crate::gpu::GpuMoe {
router: rf,
sgate: sf,
experts,
shared,
n_exp: m.experts.len(),
top_k: m.top_k,
inter,
norm_topk: m.norm_topk_prob,
route_scale: m.routed_scaling,
gu_q2,
sigmoid: m.router_sigmoid,
bias: m.expert_bias.as_deref(),
shared_gated,
})
}
pub(crate) fn moe_push_job_parts<'a>(
gate: &'a QTensor,
up: &'a QTensor,
down: &'a QTensor,
x: &[f32],
w: f32,
swiglu_limit: 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, gq1, gq4, gq2) = moe_parts(gate)?;
let (_, ui, ur, uc, urs, ucf, uq1, uq4, uq2) = moe_parts(up)?;
let (_, di, dr, dc, drs, dcf, dq1, dq4, dq2) = moe_parts(down)?;
if gq1 != uq1 || uq1 != dq1 || gq4 != uq4 || uq4 != dq4 || gq2 != uq2 {
return None; }
if gq2 && (dq2 || !dq4 || down.mapped_q4tp().is_none()) {
return None;
}
if !gq2 && dq2 {
return None;
}
model_ref.get_or_insert_with(|| gm.clone());
let dt = |cf: &[f32]| {
if cf.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, dt(gcf)).into_owned(),
xs_up: prescale(x, ucf, dt(ucf)).into_owned(),
down_col: dcf,
w,
q1: gq1,
q4t: gq4 && !gq2 && gate.mapped_q4tp().is_none(),
q4tp: gq4 && (gq2 || gate.mapped_q4tp().is_some()),
gu_q2: gq2,
swiglu_limit,
});
Some(())
}
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;
if d.act != Act::Silu {
return None; }
let (gm, gi, gr, gc, grs, gcf, gq1, gq4, gq2) = moe_parts(&d.gate_proj)?;
let (_, ui, ur, uc, urs, ucf, uq1, uq4, uq2) = moe_parts(&d.up_proj)?;
let (_, di, dr, dc, drs, dcf, dq1, dq4, dq2) = moe_parts(&d.down_proj)?;
if gq1 != uq1 || uq1 != dq1 || gq4 != uq4 || uq4 != dq4 || gq2 != uq2 {
return None; }
if gq2 && (dq2 || !dq4 || d.down_proj.mapped_q4tp().is_none()) {
return None;
}
if !gq2 && dq2 {
return None;
}
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,
q1: gq1,
q4t: gq4 && !gq2 && d.gate_proj.mapped_q4tp().is_none(),
q4tp: gq4 && (gq2 || d.gate_proj.mapped_q4tp().is_some()),
gu_q2: gq2,
swiglu_limit: 0.0,
});
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);
d.act.combine(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) }
}
}
pub(crate) fn moe_route(
logits: &[f32],
m: &MoeFfn,
allowed: Option<&[bool]>,
) -> (Vec<usize>, Vec<f32>, f32) {
let ne = logits.len();
let p: Vec<f32> = if m.router_sigmoid {
logits.iter().map(|&l| 1.0 / (1.0 + (-l).exp())).collect()
} else {
let mx = logits.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
let mut e: Vec<f32> = logits.iter().map(|&l| (l - mx).exp()).collect();
let s: f32 = e.iter().sum();
for v in &mut e {
*v /= s;
}
e
};
let admit = |e: usize| {
m.mask.as_ref().is_none_or(|mk| mk[e])
&& allowed.is_none_or(|a| a.get(e).copied().unwrap_or(false))
};
let mut idx: Vec<usize> = (0..ne).filter(|&e| admit(e)).collect();
match &m.expert_bias {
Some(b) => idx.sort_unstable_by(|&x, &y| {
(p[y] + b[y])
.partial_cmp(&(p[x] + b[x]))
.unwrap()
.then(x.cmp(&y))
}),
None => idx.sort_unstable_by(|&x, &y| p[y].partial_cmp(&p[x]).unwrap().then(x.cmp(&y))),
}
idx.truncate(m.top_k);
if let Some(tau) = m.route_tau {
let total: f32 = idx.iter().map(|&e| p[e]).sum();
if total > 0.0 {
let mut acc = 0.0f32;
let mut keep = idx.len();
for (i, &e) in idx.iter().enumerate() {
acc += p[e];
if acc >= tau * total {
keep = i + 1;
break;
}
}
idx.truncate(keep);
}
}
let wsum: f32 = if m.norm_topk_prob {
let s: f32 = idx.iter().map(|&e| p[e]).sum();
(if m.router_sigmoid { s + 1e-6 } else { s }) / m.routed_scaling
} else {
1.0 / m.routed_scaling
};
(idx, p, wsum)
}
fn moe_trace(idx: &[usize]) {
moe_trace_at(crate::gpu::cur_layer() as i32, idx)
}
pub(crate) fn moe_trace_at(li: i32, idx: &[usize]) {
use std::io::Write;
static F: std::sync::OnceLock<Option<std::sync::Mutex<std::fs::File>>> =
std::sync::OnceLock::new();
let Some(f) = F.get_or_init(|| {
let p = std::env::var("CMF_MOE_TRACE").ok()?;
Some(std::sync::Mutex::new(
std::fs::OpenOptions::new()
.create(true)
.append(true)
.open(p)
.ok()?,
))
}) else {
return;
};
let ids: Vec<String> = idx.iter().map(|e| e.to_string()).collect();
let _ = writeln!(f.lock().unwrap(), "{li}:{}", ids.join(","));
}
pub(crate) fn moe_ffn(
m: &MoeFfn,
x: &[f32],
pool: Option<&Pool>,
allowed: Option<&[bool]>,
) -> Vec<f32> {
accumulate_act(m, x, 1);
let ne = m.experts.len();
let mut logits = vec![0.0f32; ne];
match &m.resonance {
Some(r) => r.scores(x, &mut logits),
None => m.router.matvec(x, &mut logits, pool),
}
let (idx, p, wsum) = moe_route(&logits, m, allowed);
{
let mut st = m.stats.borrow_mut();
if st.len() < ne {
st.resize(ne, 0);
}
for &e in &idx {
st[e] += 1;
}
}
moe_trace(&idx);
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 graph_note(built: bool, layers_run: usize, total_layers: usize) {
use std::sync::atomic::{AtomicBool, Ordering};
if built {
GRAPH_TOK_OK.fetch_add(1, Ordering::Relaxed);
if total_layers > 0 && layers_run < total_layers {
GRAPH_TOK_PREFIX.fetch_add(1, Ordering::Relaxed);
} else {
GRAPH_TOK_FULL.fetch_add(1, Ordering::Relaxed);
}
} else {
GRAPH_TOK_MISS.fetch_add(1, Ordering::Relaxed);
}
static SAID: AtomicBool = AtomicBool::new(false);
if !SAID.swap(true, Ordering::Relaxed) {
if built {
tracing::info!("wgpu whole-token graph: ACTIVE");
} else {
tracing::warn!("wgpu whole-token graph refused — per-op path");
}
}
}
pub static GRAPH_TOK_OK: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
pub static GRAPH_TOK_MISS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
pub static GRAPH_TOK_PREFIX: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
pub static GRAPH_TOK_FULL: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
pub static METAL_GRAPH_TOK_OK: std::sync::atomic::AtomicU64 =
std::sync::atomic::AtomicU64::new(0);
pub static METAL_GRAPH_HEAD_OK: std::sync::atomic::AtomicU64 =
std::sync::atomic::AtomicU64::new(0);
pub static METAL_GRAPH_HEAD_MISS: std::sync::atomic::AtomicU64 =
std::sync::atomic::AtomicU64::new(0);
pub static METAL_GRAPH_LAYERS: std::sync::atomic::AtomicU64 =
std::sync::atomic::AtomicU64::new(0);
pub static METAL_GRAPH_ERRORS: std::sync::atomic::AtomicU64 =
std::sync::atomic::AtomicU64::new(0);
pub static METAL_PREFILL_CHUNKS: std::sync::atomic::AtomicU64 =
std::sync::atomic::AtomicU64::new(0);
pub static METAL_PREFILL_ROWS: std::sync::atomic::AtomicU64 =
std::sync::atomic::AtomicU64::new(0);
pub static METAL_PREFILL_HEAD_ROWS: std::sync::atomic::AtomicU64 =
std::sync::atomic::AtomicU64::new(0);
pub static METAL_PREFILL_ERRORS: std::sync::atomic::AtomicU64 =
std::sync::atomic::AtomicU64::new(0);
fn moe_batch_enabled() -> bool {
static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*ON.get_or_init(|| std::env::var("CMF_MOE_BATCH").as_deref() != Ok("0"))
}
fn moe_ffn_cpu_batched(
m: &MoeFfn,
x: &[f32],
idx: &[usize],
p: &[f32],
wsum: f32,
pool: Option<&Pool>,
) -> Option<Vec<f32>> {
if idx.is_empty() || !moe_batch_enabled() {
return None;
}
if FFN_PROBE.with(|pr| pr.borrow().is_some()) {
return None;
}
let n = idx.len() + usize::from(m.shared.is_some());
let mut pairs = Vec::with_capacity(n);
let mut downs = Vec::with_capacity(n);
let mut ws = Vec::with_capacity(n);
for &e in idx {
let d = &m.experts[e];
if d.act != Act::Silu {
return None;
}
pairs.push((&d.gate_proj, &d.up_proj));
downs.push(&d.down_proj);
ws.push(p[e] / wsum * m.per_expert_scale.as_ref().map_or(1.0, |v| v[e]));
}
if let Some((se, gate)) = &m.shared {
if se.act != Act::Silu {
return None;
}
let g = gate.as_ref().map_or(1.0, |gate| {
let mut gl = [0.0f32; 1];
gate.matvec(x, &mut gl, pool);
1.0 / (1.0 + (-gl[0]).exp())
});
pairs.push((&se.gate_proj, &se.up_proj));
downs.push(&se.down_proj);
ws.push(g);
}
let inter = pairs[0].0.rows();
let mut gs: Vec<Vec<f32>> = (0..pairs.len()).map(|_| vec![0f32; inter]).collect();
if !QTensor::moe_gate_up_many(&pairs, x, &mut gs, pool) {
return None;
}
let mut out = attention::take_buf(x.len());
if !QTensor::moe_down_many(&downs, &gs, &ws, &mut out, pool) {
attention::recycle_buf(&mut out);
return None;
}
Some(out)
}
pub(crate) fn moe_cold_experts_cpu(
experts: &[(&DenseFfn, f32)],
x: &[f32],
pool: Option<&Pool>,
) -> Vec<f32> {
let mut out = attention::take_buf(x.len());
if experts.is_empty() {
return out;
}
let pairs: Vec<_> = experts
.iter()
.map(|(e, _)| (&e.gate_proj, &e.up_proj))
.collect();
let downs: Vec<_> = experts.iter().map(|(e, _)| &e.down_proj).collect();
let weights: Vec<_> = experts.iter().map(|(_, w)| *w).collect();
let inter = experts[0].0.gate_proj.rows();
let mut activations: Vec<Vec<f32>> = (0..experts.len()).map(|_| vec![0.0; inter]).collect();
if QTensor::moe_gate_up_many(&pairs, x, &mut activations, pool)
&& QTensor::moe_down_many(&downs, &activations, &weights, &mut out, pool)
{
return out;
}
out.fill(0.0);
for &(expert, weight) in experts {
let mut one = dense_ffn(expert, x, pool);
for (o, v) in out.iter_mut().zip(&one) {
*o += weight * v;
}
attention::recycle_buf(&mut one);
}
out
}
fn moe_ffn_cpu(
m: &MoeFfn,
x: &[f32],
idx: &[usize],
p: &[f32],
wsum: f32,
pool: Option<&Pool>,
) -> Vec<f32> {
if let Some(out) = moe_ffn_cpu_batched(m, x, idx, p, wsum, pool) {
return out;
}
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 * m.per_expert_scale.as_ref().map_or(1.0, |v| v[e]);
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 g = gate.as_ref().map_or(1.0, |gate| {
let mut gl = [0.0f32; 1];
gate.matvec(x, &mut gl, pool);
1.0 / (1.0 + (-gl[0]).exp())
});
for i in 0..out.len() {
out[i] += g * so[i];
}
attention::recycle_buf(&mut so);
}
out
}
#[allow(clippy::too_many_arguments)]
fn mla_attention(
w: &MlaWeights,
normed: &[f32],
cache: &mut crate::kv_cache::LayerKvCache,
position: usize,
inv_freq: &[f32],
rope_scale: f32,
eps: f64,
pool: Option<&Pool>,
) -> Vec<f32> {
let (nh, dr, dn, dv, lora) = (w.nh, w.qk_rope, w.qk_nope, w.v_dim, w.lora);
let hd = dr + dn;
let mut q = vec![0.0f32; nh * hd];
match (&w.q_a, &w.q_a_norm) {
(Some(qa), Some(qn)) => {
let mut t = vec![0.0f32; qa.rows()];
qa.matvec(normed, &mut t, pool);
let tn = inference::rms_norm(&t, qn, eps, NormStyle::Qwen);
w.q_proj.matvec(&tn, &mut q, pool);
}
_ => w.q_proj.matvec(normed, &mut q, pool),
}
let mut ca = vec![0.0f32; lora + dr];
w.kv_a.matvec(normed, &mut ca, pool);
let (c_lat, k_rope) = ca.split_at_mut(lora);
let latn = inference::rms_norm(c_lat, &w.kv_a_norm, eps, NormStyle::Qwen);
let mut kvb = vec![0.0f32; nh * (dn + dv)];
w.kv_b.matvec(&latn, &mut kvb, pool);
if !w.nope {
attention::rope_rotate_scaled(k_rope, position, inv_freq, rope_scale);
}
for h in 0..nh {
if !w.nope {
attention::rope_rotate_scaled(
&mut q[h * hd..h * hd + dr],
position,
inv_freq,
rope_scale,
);
}
}
let mut k = vec![0.0f32; nh * hd];
let mut v = vec![0.0f32; nh * hd];
for h in 0..nh {
k[h * hd..h * hd + dr].copy_from_slice(k_rope);
k[h * hd + dr..(h + 1) * hd].copy_from_slice(&kvb[h * (dn + dv)..h * (dn + dv) + dn]);
v[h * hd..h * hd + dv].copy_from_slice(&kvb[h * (dn + dv) + dn..(h + 1) * (dn + dv)]);
}
cache.append(&k, &v, &vec![true; nh]);
let (ao, mut imp) = attention::attend_all_heads(&q, cache, nh, 1, hd, w.scale, None, 0.0);
attention::recycle_buf(&mut imp);
let mut ov = vec![0.0f32; nh * dv];
for h in 0..nh {
ov[h * dv..(h + 1) * dv].copy_from_slice(&ao[h * hd..h * hd + dv]);
}
let mut out = vec![0.0f32; w.o_proj.rows()];
w.o_proj.matvec(&ov, &mut out, pool);
out
}
fn dense_moe_ffn(
dm: &DenseMoeFfn,
x_normed: &[f32],
h_raw: &[f32],
eps: f64,
norm_style: NormStyle,
pool: Option<&Pool>,
) -> Vec<f32> {
let mut d = dense_ffn(&dm.dense, x_normed, pool);
d = inference::rms_norm(&d, &dm.post_norm_1, eps, norm_style);
let m = &dm.moe;
let ne = m.experts.len();
let mut logits = vec![0.0f32; ne];
if m.router_input_norm {
let ss: f32 = h_raw.iter().map(|v| v * v).sum::<f32>() / h_raw.len() as f32;
let inv = 1.0 / (ss + eps as f32).sqrt();
let xr: Vec<f32> = h_raw.iter().map(|v| v * inv).collect();
m.router.matvec(&xr, &mut logits, pool);
} else {
m.router.matvec(h_raw, &mut logits, pool);
}
let (idx, p, wsum) = moe_route(&logits, m, None);
{
let mut st = m.stats.borrow_mut();
if st.len() < ne {
st.resize(ne, 0);
}
for &e in &idx {
st[e] += 1;
}
}
let x2 = inference::rms_norm(h_raw, &dm.pre_norm_2, eps, norm_style);
let mo = moe_ffn_cpu(m, &x2, &idx, &p, wsum, pool);
let mo = inference::rms_norm(&mo, &dm.post_norm_2, eps, norm_style);
for (di, mi) in d.iter_mut().zip(&mo) {
*di += mi;
}
d
}
fn moe_gpu_refused(why: &'static str) {
use std::sync::atomic::{AtomicBool, Ordering};
static SAID: AtomicBool = AtomicBool::new(false);
if !SAID.swap(true, Ordering::Relaxed) {
tracing::warn!("MoE GPU block refused ({why}) — experts run on the CPU");
}
}
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 {
if moe_push_job(&m.experts[e], x, p[e] / wsum, &mut jobs, &mut model_ref).is_none() {
moe_gpu_refused("push_job(expert)");
return None;
}
}
if let Some((se, gate)) = &m.shared {
let g = gate.as_ref().map_or(1.0, |gate| {
let mut gl = [0.0f32; 1];
gate.matvec(x, &mut gl, pool);
1.0 / (1.0 + (-gl[0]).exp())
});
if moe_push_job(se, x, g, &mut jobs, &mut model_ref).is_none() {
moe_gpu_refused("push_job(shared)");
return None;
}
}
let Some(model) = model_ref else {
moe_gpu_refused("no model_ref");
return None;
};
let hidden = jobs[0].down.1;
let mut out = vec![0.0f32; hidden];
if crate::gpu::moe_block(&model, &jobs, &mut out) {
Some(out)
} else {
moe_gpu_refused("gpu::moe_block");
None
}
}
fn ffn_forward(
ffn: &FfnKind,
x: &[f32],
pool: Option<&Pool>,
experts_allowed: Option<&[bool]>,
) -> Vec<f32> {
match ffn {
FfnKind::Dense(d) if !d.segs.is_empty() => tube_ffn(d, x, 1, pool, None),
FfnKind::Dense(d) => dense_ffn(d, x, pool),
FfnKind::Moe(m) => moe_ffn(m, x, pool, experts_allowed),
FfnKind::DenseMoe(_) => unreachable!("DenseMoe dispatches via dense_moe_ffn"),
}
}
fn ffn_forward_pair(
ffn: &FfnKind,
x1: &[f32],
x2: &[f32],
pool: Option<&Pool>,
experts_allowed: Option<&[bool]>,
) -> (Vec<f32>, Vec<f32>) {
let d = match ffn {
FfnKind::Dense(d) if !d.segs.is_empty() => {
return (
tube_ffn(d, x1, 1, pool, None),
tube_ffn(d, x2, 1, pool, None),
);
}
FfnKind::Dense(d) => d,
FfnKind::Moe(m) => {
return (
moe_ffn(m, x1, pool, experts_allowed),
moe_ffn(m, x2, pool, experts_allowed),
);
}
FfnKind::DenseMoe(_) => unreachable!("DenseMoe dispatches via dense_moe_ffn"),
};
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] = d.act.combine(g1[i], u1[i]);
g2[i] = d.act.combine(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 {
#[test]
fn prefill_chunk_rule_widens_only_dense_on_discrete() {
use super::{
prefill_chunk_rule, ChunkHost, ChunkStackFacts, DISCRETE_DENSE_PREFILL_CHUNK,
};
let dense_card = ChunkStackFacts {
plain_dense: true,
discrete: true,
gpu_on: true,
..Default::default()
};
assert!(dense_card.dense_on_discrete());
assert_eq!(
prefill_chunk_rule(None, ChunkHost::Other, dense_card.dense_on_discrete()),
DISCRETE_DENSE_PREFILL_CHUNK
);
assert!(DISCRETE_DENSE_PREFILL_CHUNK > 48);
for (label, facts) in [
("GDN hybrid / MoE / DeepSeek stack", ChunkStackFacts { plain_dense: false, ..dense_card }),
("integrated GPU", ChunkStackFacts { discrete: false, ..dense_card }),
("CPU only", ChunkStackFacts { gpu_on: false, discrete: false, ..dense_card }),
("capacity split", ChunkStackFacts { capacity_split: true, ..dense_card }),
("multi-GPU plan", ChunkStackFacts { multi_gpu: true, ..dense_card }),
("O(1) layers", ChunkStackFacts { o1: true, ..dense_card }),
] {
assert!(!facts.dense_on_discrete(), "{label}");
assert_eq!(
prefill_chunk_rule(None, ChunkHost::Other, facts.dense_on_discrete()),
48,
"{label} keeps the historical x86 chunk"
);
}
for dense in [false, true] {
assert_eq!(prefill_chunk_rule(None, ChunkHost::Macos, dense), 512);
assert_eq!(prefill_chunk_rule(None, ChunkHost::Aarch64, dense), 256);
}
for host in [ChunkHost::Macos, ChunkHost::Aarch64, ChunkHost::Other] {
for dense in [false, true] {
assert_eq!(prefill_chunk_rule(Some(48), host, dense), 48);
assert_eq!(prefill_chunk_rule(Some(0), host, dense), 1);
}
}
}
#[test]
fn kv_reuse_plan_pulls_rows_decode_wrote_only_on_the_device() {
use super::{ReuseLayer, ReusePlan, kv_reuse_plan};
let full = |host_rows, device_rows| ReuseLayer {
full: true,
host_rows,
device_rows,
device_state: false,
};
assert_eq!(
kv_reuse_plan(339, &[full(300, Some(339)), full(300, Some(339))]),
ReusePlan::Pull(vec![(0, 300, 339), (1, 300, 339)])
);
assert_eq!(kv_reuse_plan(339, &[full(339, None)]), ReusePlan::Ready);
assert_eq!(kv_reuse_plan(339, &[full(339, Some(345))]), ReusePlan::Ready);
assert_eq!(
kv_reuse_plan(339, &[full(300, Some(339)), full(339, None)]),
ReusePlan::Pull(vec![(0, 300, 339)])
);
assert_eq!(kv_reuse_plan(339, &[full(300, Some(320))]), ReusePlan::Fresh);
assert_eq!(kv_reuse_plan(339, &[full(300, None)]), ReusePlan::Fresh);
assert_eq!(kv_reuse_plan(339, &[full(350, None)]), ReusePlan::Fresh);
let conv = |device_state| ReuseLayer {
full: false,
host_rows: 0,
device_rows: None,
device_state,
};
assert_eq!(
kv_reuse_plan(339, &[conv(true), full(300, Some(339))]),
ReusePlan::Fresh
);
assert_eq!(kv_reuse_plan(339, &[conv(false), full(339, None)]), ReusePlan::Ready);
}
#[test]
fn nll_graph_policy_scopes_only_the_fused_head() {
for (label, unmasked, prefer_graph, native_metal, want_graph, want_head) in [
("vulkan graph", true, true, false, true, false),
("native Metal graph", true, true, true, true, true),
("masked", false, true, false, false, false),
("graph disabled", true, false, true, false, false),
] {
let (graph_quality, graph_head_required) =
super::nll_graph_policy(unmasked, prefer_graph, native_metal);
assert_eq!(graph_quality, want_graph, "{label}: graph quality");
assert_eq!(graph_head_required, want_head, "{label}: fused head");
}
}
#[test]
fn mtp_prefill_pair_boundaries_skip_only_final_prompt_row() {
assert_eq!(mtp_prefill_pair_count(0, 128, 256), 128);
assert_eq!(mtp_prefill_pair_count(128, 256, 256), 127);
assert_eq!(mtp_prefill_pair_count(0, 256, 256), 255);
assert_eq!(mtp_prefill_pair_count(256, 256, 256), 0);
assert_eq!(mtp_prefill_pair_count(300, 320, 256), 0);
}
#[test]
fn cancel_flag_stops_generation() {
let mut p = create_test_pipeline(16, 32, 2, 2, 8, 2, 32);
p.cancel.store(true, std::sync::atomic::Ordering::Relaxed);
let r = p.generate_from_ids(&[1, 2, 3], 8, None, None).unwrap();
assert_eq!(r.finish_reason, "cancelled");
assert!(
r.token_ids.is_empty(),
"no tokens after cancel: {:?}",
r.token_ids
);
assert_eq!(p.kv_cache.seq_len(), 0);
assert!(p.kv_history.is_empty());
assert!(!p.graph_want_logits);
assert!(p.graph_logits.is_none());
let r2 = p.generate_from_ids(&[1, 2, 3], 4, None, None).unwrap();
assert_ne!(r2.finish_reason, "cancelled");
}
use super::*;
#[test]
fn dynamic_ffn_equals_the_zeroing_arm() {
let (hidden, inter) = (8usize, 32usize);
let synth = |n: usize, salt: usize| -> Vec<f32> {
(0..n)
.map(|i| (((i * 29 + salt * 13 + 7) % 89) as f32 / 89.0 - 0.5) * 0.6)
.collect()
};
let down = synth(hidden * inter, 3);
let mut down_t = vec![0.0f32; inter * hidden];
for r in 0..hidden {
for c in 0..inter {
down_t[c * hidden + r] = down[r * inter + c];
}
}
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(down.clone(), hidden, inter),
act: Act::Silu,
down_t: Some(QTensor::from_f32(down_t, inter, hidden)),
segs: Vec::new(),
};
let x = synth(hidden, 11);
let k = 12usize;
let got = dense_ffn_dynamic(&d, &x, None, k).expect("down_t present");
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);
for v in g.iter_mut() {
*v = inference::silu(*v);
}
keep_top_k(&mut g, k);
for i in 0..inter {
g[i] *= u[i];
}
let mut want = vec![0.0f32; hidden];
d.down_proj.matvec(&g, &mut want, None);
for (a, b) in want.iter().zip(&got) {
assert!((a - b).abs() < 1e-5, "dynamic {b} vs reference {a}");
}
}
#[test]
fn tube_ffn_open_equals_dense_and_closed_equals_masked() {
let (hidden, core, tube) = (8usize, 12usize, 8usize);
let inter = core + tube;
let synth = |n: usize, salt: usize| -> Vec<f32> {
(0..n)
.map(|i| (((i * 41 + salt * 17 + 5) % 97) as f32 / 97.0 - 0.5) * 0.5)
.collect()
};
let (g_all, u_all) = (synth(inter * hidden, 1), synth(inter * hidden, 2));
let d_all = synth(hidden * inter, 3);
let dense = DenseFfn {
gate_proj: QTensor::from_f32(g_all.clone(), inter, hidden),
up_proj: QTensor::from_f32(u_all.clone(), inter, hidden),
down_proj: QTensor::from_f32(d_all.clone(), hidden, inter),
act: Act::Silu,
down_t: None,
segs: Vec::new(),
};
let rows =
|v: &[f32], a: usize, b: usize| -> Vec<f32> { v[a * hidden..b * hidden].to_vec() };
let cols = |v: &[f32], a: usize, b: usize| -> Vec<f32> {
let mut o = Vec::with_capacity(hidden * (b - a));
for r in 0..hidden {
o.extend_from_slice(&v[r * inter + a..r * inter + b]);
}
o
};
let tubed = DenseFfn {
down_t: None,
gate_proj: QTensor::from_f32(rows(&g_all, 0, core), core, hidden),
up_proj: QTensor::from_f32(rows(&u_all, 0, core), core, hidden),
down_proj: QTensor::from_f32(cols(&d_all, 0, core), hidden, core),
act: Act::Silu,
segs: vec![FfnSeg {
gate: QTensor::from_f32(rows(&g_all, core, inter), tube, hidden),
up: QTensor::from_f32(rows(&u_all, core, inter), tube, hidden),
down: QTensor::from_f32(cols(&d_all, core, inter), hidden, tube),
start: core,
width: tube,
}],
};
let x = synth(hidden, 7);
let want = dense_ffn(&dense, &x, None);
let got = tube_ffn(&tubed, &x, 1, None, None);
for (a, b) in want.iter().zip(&got) {
assert!((a - b).abs() < 1e-5, "open tube: {a} vs {b}");
}
let mut bits = vec![0u8; inter.div_ceil(8)];
for n in 0..core {
bits[n / 8] |= 1 << (n % 8);
}
let closed = tube_ffn(&tubed, &x, 1, None, Some(&bits));
let masked = dense_ffn_masked(&dense, &x, None, &bits);
for (a, b) in masked.iter().zip(&closed) {
assert!((a - b).abs() < 1e-5, "closed tube: {a} vs {b}");
}
let batch = tube_ffn(&tubed, &x, 1, None, Some(&bits));
for (a, b) in closed.iter().zip(&batch) {
assert_eq!(a, b, "batch arm disagrees with decode arm");
}
}
#[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),
act: Act::Silu,
down_t: None,
segs: Vec::new(),
};
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],
attn_out_norm: None,
ffn_out_norm: None,
layer_scale: None,
ffn: FfnKind::Dense(DenseFfn {
gate_proj: qt(inter, h, 315),
up_proj: qt(inter, h, 316),
down_proj: qt(h, inter, 317),
act: Act::Silu,
down_t: None,
segs: Vec::new(),
}),
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,
softplus_gate: None,
},
},
final_norm: vec![1.0; h],
kv: crate::kv_cache::LayerKvCache::new(kv, hd),
});
}
#[test]
fn speculative_equals_vanilla_greedy() {
unsafe { std::env::set_var("CMF_GPU_WGPU_GRAPH", "0") };
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() {
unsafe { std::env::set_var("CMF_GPU_WGPU_GRAPH", "0") };
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 resetting_sampler_restarts_the_seeded_stream() {
let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
let config = SamplerConfig {
seed: Some(1234),
..SamplerConfig::default()
};
p.set_sampler_config(config.clone());
let first = p.generate("hello", 8, None, None).unwrap().token_ids;
p.set_sampler_config(config);
let second = p.generate("hello", 8, None, None).unwrap().token_ids;
assert_eq!(first, second);
}
#[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"
);
}
#[test]
fn laguna_shared_expert_is_unconditionally_added() {
let matrix = |values: Vec<f32>| QTensor::from_f32(values, 2, 2);
let identity = || matrix(vec![1.0, 0.0, 0.0, 1.0]);
let zero_dense = || DenseFfn {
gate_proj: matrix(vec![0.0; 4]),
up_proj: matrix(vec![0.0; 4]),
down_proj: matrix(vec![0.0; 4]),
act: Act::Silu,
down_t: None,
segs: Vec::new(),
};
let shared = DenseFfn {
gate_proj: identity(),
up_proj: identity(),
down_proj: identity(),
act: Act::Silu,
down_t: None,
segs: Vec::new(),
};
let x = [1.0, 2.0];
let expected = dense_ffn(&shared, &x, None);
let moe = MoeFfn {
router: QTensor::from_f32(vec![0.0, 0.0], 1, 2),
experts: vec![zero_dense()],
top_k: 1,
norm_topk_prob: true,
router_sigmoid: true,
expert_bias: None,
routed_scaling: 1.0,
route_tau: None,
shared: Some((shared, None)),
stats: std::cell::RefCell::new(Vec::new()),
act_sq: std::cell::RefCell::new(Vec::new()),
act_rows: std::cell::RefCell::new(Vec::new()),
mask: None,
per_expert_scale: None,
router_input_norm: false,
resonance: None,
};
let actual = moe_ffn_cpu(&moe, &x, &[0], &[0.0], 1.0, None);
for (actual, expected) in actual.iter().zip(expected) {
assert!((actual - expected).abs() < 1e-6);
}
}
#[test]
fn o1_batch_transition_publishes_one_epoch_before_serial_handoff() {
const B: usize = 19;
let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
p.set_o1(Some(crate::nystrom::O1Cfg {
layers: crate::nystrom::O1Layers::All,
m: 4,
w: 8,
sink: 2,
rect: crate::nystrom::O1Rect::Aggregate,
}));
p.o1_begin_with_prefix(Some(B));
let ids: Vec<u32> = (0..B as u32).collect();
let _ = p.prefill_batch_span(PrefillIn::Ids(&ids), 0, None, 0, p.num_layers);
assert_eq!(p.o1_epoch, 1, "all layers publish one completed transition");
assert!(p.kv_cache.layers.iter().all(|l| l.o1_sealed()));
let next = p.embed_single(B as u32);
let _ = p.forward_layers(&next, B, None);
assert_eq!(p.o1_epoch, 1, "sealed handoff must not republish the epoch");
}
#[test]
fn o1_pair_transition_commits_scratch_before_epoch_publication() {
const B: usize = 19;
let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
let gdn_cfg = crate::linear_core::GdnCfg {
num_v_heads: 2,
num_k_heads: 1,
key_head_dim: 2,
value_head_dim: 4,
conv_kernel: 3,
hidden_size: 8,
rms_eps: 1e-6,
output_gate_sigmoid: false,
};
let synth = |n: usize, salt: usize| -> Vec<f32> {
(0..n)
.map(|i| (((i * 13 + salt * 7) % 97) as f32 / 97.0 - 0.5) * 0.4)
.collect()
};
let qt = |rows: usize, cols: usize, salt: usize| {
crate::qtensor::QTensor::from_f32(synth(rows * cols, salt), rows, cols)
};
let c_dim = gdn_cfg.conv_dim();
let vd = gdn_cfg.num_v_heads * gdn_cfg.value_head_dim;
p.weights.layers[0].attn = AttnKind::LinearGdn(crate::linear_core::GdnWeights {
in_proj_qkv: qt(c_dim, 8, 1),
in_proj_z: qt(vd, 8, 2),
in_proj_a: qt(gdn_cfg.num_v_heads, 8, 3),
in_proj_b: qt(gdn_cfg.num_v_heads, 8, 4),
conv1d: synth(c_dim * gdn_cfg.conv_kernel, 5),
a_log: vec![0.2, 0.5],
dt_bias: synth(gdn_cfg.num_v_heads, 6),
norm: vec![1.0; gdn_cfg.value_head_dim],
out_proj: qt(8, vd, 7),
});
p.gdn_cfg = Some(gdn_cfg);
p.set_o1(Some(crate::nystrom::O1Cfg {
layers: crate::nystrom::O1Layers::All,
m: 4,
w: 8,
sink: 2,
rect: crate::nystrom::O1Rect::Aggregate,
}));
p.o1_begin_with_prefix(Some(B));
for pos in 0..B - 2 {
let emb = p.embed_single(pos as u32);
let _ = p.forward_layers(&emb, pos, None);
}
let lane1_state = p.kv_cache.layers[0].linear_state.clone();
let e1 = p.embed_single((B - 2) as u32);
let e2 = p.embed_single((B - 1) as u32);
let _ = p.forward_pair(&e1, &e2, B - 2);
assert_eq!(p.o1_epoch, 1, "pair crossing B publishes one epoch");
assert!(
p.kv_cache
.layers
.iter()
.enumerate()
.all(|(li, l)| !p.o1_flags[li] || l.o1_sealed())
);
assert!(!p.kv_cache.layers[0].linear_state.is_empty());
assert_ne!(
p.kv_cache.layers[0].linear_state, lane1_state,
"real pair must commit GDN lane 2 before returning"
);
assert!(p.kv_cache.layers[0].linear_scratch.is_empty());
let next = p.embed_single(B as u32);
let _ = p.forward_layers(&next, B, None);
assert_eq!(p.o1_epoch, 1, "serial continuation must reuse the epoch");
}
#[test]
fn o1_error_observation_stays_terminal_until_reset() {
let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
p.set_o1(Some(crate::nystrom::O1Cfg {
layers: crate::nystrom::O1Layers::All,
m: 4,
w: 8,
sink: 2,
rect: crate::nystrom::O1Rect::Aggregate,
}));
p.o1_begin();
p.kv_cache.layers[0].o1_abort("synthetic transition failure".into());
assert!(p.o1_seal_checked().is_err());
assert!(
p.o1_seal_checked().is_err(),
"retry must see the sticky error"
);
let k = vec![0.2f32; 4];
let v = vec![0.3f32; 4];
p.kv_cache.layers[0].append(&k, &v, &[]);
assert_eq!(p.kv_cache.layers[0].seq_len, 0);
p.reset_session();
p.o1_begin();
p.kv_cache.layers[0].append(&k, &v, &[]);
assert_eq!(p.kv_cache.layers[0].seq_len, 1);
}
#[test]
fn nll_graph_failure_is_terminal_and_request_is_reusable() {
let ids = vec![1u32, 2, 3, 4, 5, 6];
let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
p.graph_logits = Some(vec![123.0]);
p.graph_want_logits = true;
p.graph_failed
.store(true, std::sync::atomic::Ordering::Relaxed);
p.cancel.store(true, std::sync::atomic::Ordering::Relaxed);
let err = p.nll_ids_from(&ids, 0).expect_err("prior graph failure");
assert!(err.contains("before NLL"));
assert!(p.graph_logits.is_none());
assert!(!p.graph_want_logits);
assert!(!p.graph_failed.load(std::sync::atomic::Ordering::Relaxed));
assert!(!p.cancel.load(std::sync::atomic::Ordering::Relaxed));
let mut fresh = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
let expected = fresh.nll_ids_from(&ids, 0).expect("fresh NLL");
let actual = p.nll_ids_from(&ids, 0).expect("reused NLL");
assert_eq!(actual.1, expected.1);
assert!((actual.0 - expected.0).abs() < 1e-9);
}
#[test]
fn nll_forward_failure_discards_partial_score_and_clears_sidechannels() {
let ids = vec![1u32, 2, 3, 4, 5, 6];
let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
p.nll_test_fail_at = Some(1);
let err = p
.nll_ids_from(&ids, 0)
.expect_err("one-shot forward failure");
assert!(err.contains("forward") || err.contains("score row"));
assert!(!p.graph_failed.load(std::sync::atomic::Ordering::Relaxed));
assert!(!p.graph_want_logits);
assert!(p.graph_logits.is_none());
assert!(p.kv_history.is_empty());
let mut fresh = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
let expected = fresh.nll_ids_from(&ids, 0).expect("fresh NLL");
let actual = p.nll_ids_from(&ids, 0).expect("reused NLL");
assert_eq!(actual.1, expected.1);
assert!((actual.0 - expected.0).abs() < 1e-9);
}
#[test]
fn nll_serial_failure_before_first_row_is_reported() {
let ids = vec![1u32, 2, 3, 4];
let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
p.nll_test_force_serial = true;
p.nll_test_fail_at = Some(0);
let err = p.nll_ids_from(&ids, 0).expect_err("serial forward failure");
assert!(err.contains("serial forward"));
assert!(p.kv_history.is_empty());
assert!(!p.graph_failed.load(std::sync::atomic::Ordering::Relaxed));
assert!(!p.cancel.load(std::sync::atomic::Ordering::Relaxed));
}
#[test]
fn ffn_probe_failure_discards_recorder_and_state() {
let ids = vec![1u32, 2, 3, 4];
let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
p.nll_test_fail_at = Some(0);
let err = p
.probe_ffn_mass_batch(&ids)
.expect_err("probe forward failure");
assert!(err.contains("NLL"));
assert!(FFN_PROBE.with(|probe| probe.borrow().is_none()));
assert!(p.kv_history.is_empty());
assert!(!p.graph_failed.load(std::sync::atomic::Ordering::Relaxed));
}
#[test]
fn nll_test_controls_are_pipeline_scoped() {
let ids = vec![1u32, 2, 3, 4];
let mut failing = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
let mut unaffected = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
failing.nll_test_force_serial = true;
failing.nll_test_fail_at = Some(0);
assert!(!failing.can_prefill_batched());
assert!(unaffected.can_prefill_batched());
let expected = unaffected
.nll_ids_from(&ids, 0)
.expect("unaffected pipeline remains usable");
let err = failing
.nll_ids_from(&ids, 0)
.expect_err("failure injection belongs to failing pipeline");
assert!(err.contains("serial forward"));
assert!(failing.nll_test_fail_at.is_none());
assert!(unaffected.can_prefill_batched());
let actual = unaffected
.nll_ids_from(&ids, 0)
.expect("unaffected pipeline remains reusable");
assert_eq!(actual.1, expected.1);
assert!((actual.0 - expected.0).abs() < 1e-9);
}
#[test]
fn forward_ids_failure_channel_is_terminal_and_reusable() {
let ids = vec![1u32, 2, 3, 4, 5, 6];
let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
p.graph_logits = Some(vec![123.0]);
p.graph_want_logits = true;
p.graph_failed
.store(true, std::sync::atomic::Ordering::Relaxed);
p.cancel.store(true, std::sync::atomic::Ordering::Relaxed);
let err = p
.forward_ids(&ids, None)
.expect_err("a failed forward must not become a valid head result");
assert!(err.contains("forward_ids setup"));
assert!(p.graph_logits.is_none());
assert!(!p.graph_want_logits);
assert!(!p.graph_failed.load(std::sync::atomic::Ordering::Relaxed));
assert!(!p.cancel.load(std::sync::atomic::Ordering::Relaxed));
assert_eq!(p.kv_cache.seq_len(), 0);
let expected = create_test_pipeline(8, 16, 2, 1, 4, 1, 64)
.forward_ids(&ids, None)
.expect("fresh forward_ids");
let actual = p
.forward_ids(&ids, None)
.expect("pipeline remains reusable after a failed forward");
assert_eq!(actual.len(), expected.len());
assert!(
actual
.iter()
.zip(expected)
.all(|(a, b)| (a - b).abs() < 1e-9)
);
assert_eq!(p.kv_cache.seq_len(), ids.len());
}
}