1use crate::attention::{self, QwenAttnCfg};
10use crate::inference;
11
12#[path = "mimo_mtp.rs"]
15pub mod mimo_mtp;
16use crate::kv_cache::KvCache;
17use crate::linear_core::{
18 GdnCfg, GdnWeights, ShortConvCfg, ShortConvWeights, VmfPhaseCfg, VmfPhaseWeights, gdn_forward,
19 gdn_pair, short_conv_forward, short_conv_forward_batch, short_conv_pair, vmf_phase_forward,
20 vmf_phase_pair,
21};
22use crate::pool::Pool;
23use crate::qtensor::QTensor;
24use crate::sampler::{self, SamplerConfig, SamplerScratch, SplitMix64};
25use crate::tokenizer::Tokenizer;
26use cortiq_core::mask::TaskMask;
27use cortiq_core::types::NormStyle;
28
29pub static GLOBAL_USE_GPU: std::sync::atomic::AtomicBool =
30 std::sync::atomic::AtomicBool::new(false);
31
32struct ForwardScratch {
36 n1: Vec<f32>,
37 n2: Vec<f32>,
38 p1: Vec<f32>,
39 p2: Vec<f32>,
40}
41
42impl ForwardScratch {
43 fn new(hidden: usize) -> Self {
44 Self {
45 n1: vec![0.0; hidden],
46 n2: vec![0.0; hidden],
47 p1: vec![0.0; hidden],
48 p2: vec![0.0; hidden],
49 }
50 }
51}
52
53pub struct Pipeline {
55 gpu_plan: Option<std::sync::Arc<Vec<(usize, usize, usize)>>>,
60 pub tokenizer: std::sync::Arc<Tokenizer>,
63 pub kv_cache: KvCache,
64 pub sampler_config: SamplerConfig,
65 pub weights: PipelineWeights,
66 pub hidden_size: usize,
67 pub intermediate_size: usize,
68 pub num_heads: usize,
69 pub num_kv_heads: usize,
70 pub head_dim: usize,
71 pub num_layers: usize,
73 pub physical_layers: usize,
75 pub loop_final_norm: bool,
77 pub vocab_size: usize,
78 pub rms_eps: f64,
79 pub rope_base: f32,
80 pub norm_style: NormStyle,
81 pub rotary_dim: usize,
83 pub attention_heads_per_layer: Option<Vec<usize>>,
85 pub kv_heads_per_layer: Option<Vec<usize>>,
90 pub v_head_dim: Option<usize>,
95 pub layer_dump: Option<std::path::PathBuf>,
108 graph_declines: std::cell::RefCell<Vec<(&'static str, &'static str)>>,
111 pub(crate) mimo_moe: crate::mimo_moe::Slot,
114 pub vmf_cfg: Option<VmfPhaseCfg>,
116 pub gdn_cfg: Option<GdnCfg>,
118 pub logit_multiplier: Option<f32>,
120 pub cancel: std::sync::Arc<std::sync::atomic::AtomicBool>,
125 graph_failed: std::sync::atomic::AtomicBool,
130 pub kv_history: Vec<u32>,
135 pub kv_history_device: bool,
138 pub kda_cfg: Option<crate::linear_core::KdaCfg>,
140 pub g3n: Option<Box<(crate::g3n::G3nGlobals, Vec<crate::g3n::G3nLayer>)>>,
143 pub dsv4: Option<
147 Box<(
148 crate::dsv4::Dsv4Globals,
149 Vec<crate::dsv4::Dsv4Layer>,
150 crate::dsv4::Dsv4Cfg,
151 crate::dsv4::Dsv4State,
152 )>,
153 >,
154 pub dsv41: Option<
158 Box<(
159 crate::dsv41::Dsv41Globals,
160 Vec<crate::dsv41::Dsv41Layer>,
161 crate::dsv41::Dsv41Cfg,
162 crate::dsv41::Dsv41State,
163 )>,
164 >,
165 pub dsv41_vision: Option<crate::dsv41_vision::VisionModel>,
167 dsv41_prefill: Option<(Vec<Option<Vec<f32>>>, Vec<bool>)>,
169 pub qwen4_exp: Option<
172 Box<(
173 crate::qwen4_exp::Globals,
174 Vec<crate::qwen4_exp::Layer>,
175 crate::qwen4_exp::Cfg,
176 crate::qwen4_exp::State,
177 )>,
178 >,
179 pub dsv4_mtp: Vec<crate::dsv4::Dsv4Mtp>,
183 pub dspark: Option<crate::dsv4::DsparkState>,
185 pub dspark_pending: Vec<(usize, Vec<u32>, bool, usize)>,
188 pub dspark_hist: Vec<usize>,
190 pub dspark_real: Vec<u32>,
194 pub dspark_trunk_picks: Vec<Vec<(usize, Vec<usize>)>>,
198 pub dspark_exp: Vec<(usize, usize, usize, usize)>,
200 pub dspark_draft_ns: u128,
204 pub short_conv_cfg: Option<ShortConvCfg>,
207 pub mtp: Option<MtpModule>,
209 pub mimo_mtp: Option<mimo_mtp::MimoMtp>,
213 verify_exact_moe: bool,
216 pub speculative: bool,
218 pub ignore_eos: bool,
224 pub draft_full_streak: u32,
230 pub spec_k_adapt: Option<usize>,
238 pub spec_acc_ewma: f32,
240 rng: SplitMix64,
241 sampler_scratch: SamplerScratch,
242 spec_forced: Option<u32>,
248 spec_q: Vec<Vec<f32>>,
249 spec_p: Vec<f32>,
250 spec_res: Vec<f32>,
251 spec_qs: Vec<sampler::Sparse>,
253 spec_ps: sampler::Sparse,
254 spec_ress: sampler::Sparse,
255 mtp_graph_mode: Option<bool>,
262 #[cfg(target_os = "macos")]
265 metal_verify: Option<MetalVerifyPending>,
266 pub(crate) inv_freq: std::sync::Arc<Vec<f32>>,
270 ws: ForwardScratch,
274 pool: Option<std::sync::Arc<Pool>>,
276 pub(crate) model: Option<std::sync::Arc<cortiq_core::CmfModel>>,
280 pub(crate) dyn_force_f32: bool,
282 pub(crate) dyn_skill_layers: Vec<Option<Vec<usize>>>,
287 pub(crate) dyn_active: Option<usize>,
293 pub(crate) dyn_blend_loaded: bool,
297 pub(crate) dyn_phi_layer: Option<usize>,
300 dyn_phi_ema: Vec<f32>,
302 dyn_phi_seen: usize,
303 pub dyn_router: Option<crate::swarm::DynRouter>,
306 o1_cfg: Option<crate::nystrom::O1Cfg>,
309 o1_epoch: u64,
312 o1_flags: Vec<bool>,
314 trace: bool,
317 calib_temp: f32,
320 #[cfg_attr(not(target_os = "macos"), allow(dead_code))]
322 graph_kv_id: u64,
323 #[cfg_attr(not(target_os = "macos"), allow(dead_code))]
326 graph_want_logits: bool,
327 #[cfg_attr(not(target_os = "macos"), allow(dead_code))]
331 graph_head_required: bool,
332 graph_logits: Option<Vec<f32>>,
335 embryo_graph: Option<std::sync::Arc<crate::gpu::EmbryoGraphModel>>,
338 graph_refused: std::sync::atomic::AtomicBool,
344 pub embed_multiplier: f32,
346 pub attn_scale: f32,
349 pub swa: Option<(usize, usize)>,
352 pub sliding_layers: Option<Vec<bool>>,
355 pub anchor_core: Option<cortiq_core::AnchorCoreConfig>,
361 bounded_rope: Option<std::sync::Arc<crate::bounded::BoundedRope>>,
364 pub kv_prefix: KvPrefix,
368 pub last_prefill_tokens: usize,
371 pub inv_freq_local: Option<std::sync::Arc<Vec<f32>>>,
374 pub rotary_dim_local: Option<usize>,
375 pub rope_scale: f32,
376 pub rope_scale_local: f32,
377 pub global_attn: Option<(usize, usize)>,
380 pub inv_freq_global: Option<std::sync::Arc<Vec<f32>>>,
383 pub attn_v_norm: bool,
385 pub qk_norm_after_rope: bool,
387 pub final_softcap: Option<f32>,
389 pub head_clusters: Option<std::sync::Arc<Vec<f32>>>,
393 pub attn_softcap: f32,
395 confidence_on: bool,
399 #[cfg(test)]
402 nll_test_fail_at: Option<usize>,
403 #[cfg(test)]
406 nll_test_force_serial: bool,
407}
408
409#[cfg(target_os = "macos")]
410impl Drop for Pipeline {
411 fn drop(&mut self) {
412 let _ = crate::gpu_metal::wait_replay();
414 crate::gpu::kv_mirror_drop(self.graph_kv_id);
415 }
416}
417
418#[cfg(not(target_os = "macos"))]
419impl Drop for Pipeline {
420 fn drop(&mut self) {
421 crate::gpu::graph_kv_reset(self.graph_kv_id);
425 }
426}
427
428pub struct PipelineWeights {
433 pub embed_tokens: QTensor,
435 pub layers: Vec<LayerWeights>,
437 pub lm_head: QTensor,
439 pub final_norm: Vec<f32>,
441}
442
443pub struct LayerWeights {
445 pub input_norm: Vec<f32>,
446 pub post_norm: Vec<f32>,
449 pub attn_out_norm: Option<Vec<f32>>,
452 pub layer_scale: Option<f32>,
454 pub ffn_out_norm: Option<Vec<f32>>,
457 pub ffn: FfnKind,
458 pub attn: AttnKind,
459}
460
461#[derive(Clone, Copy, PartialEq, Debug, Default)]
464pub enum Act {
465 #[default]
466 Silu,
467 GeluTanh,
468 Situ {
471 beta: f32,
472 linear_beta: f32,
473 },
474}
475
476impl Act {
477 pub fn from_arch(name: &str) -> Self {
478 if name == "gelu_tanh" {
479 Self::GeluTanh
480 } else {
481 Self::Silu
482 }
483 }
484
485 pub fn from_arch_full(arch: &cortiq_core::ModelArch) -> Self {
487 match arch.hidden_act.as_str() {
488 "situ" => Self::Situ {
489 beta: arch.activation_situ_beta.unwrap_or(1.0) as f32,
490 linear_beta: arch.activation_situ_linear_beta.unwrap_or(0.0) as f32,
491 },
492 other => Self::from_arch(other),
493 }
494 }
495
496 #[inline]
497 pub fn apply(self, x: f32) -> f32 {
498 match self {
499 Self::Silu => inference::silu(x),
500 Self::GeluTanh => inference::gelu_tanh(x),
501 Self::Situ { beta, .. } => beta * (x / beta).tanh() * (1.0 / (1.0 + (-x).exp())),
502 }
503 }
504
505 #[inline]
508 pub fn combine(self, g: f32, u: f32) -> f32 {
509 match self {
510 Self::Situ { linear_beta, .. } if linear_beta > 0.0 => {
511 self.apply(g) * (linear_beta * (u / linear_beta).tanh())
512 }
513 _ => self.apply(g) * u,
514 }
515 }
516}
517
518pub struct DenseFfn {
520 pub gate_proj: QTensor,
521 pub up_proj: QTensor,
522 pub down_proj: QTensor,
523 pub act: Act,
525 pub down_t: Option<QTensor>,
531 pub segs: Vec<FfnSeg>,
538}
539
540pub struct FfnSeg {
545 pub gate: QTensor,
546 pub up: QTensor,
547 pub down: QTensor,
548 pub start: usize,
549 pub width: usize,
550}
551
552pub enum FfnKind {
555 Dense(DenseFfn),
556 Moe(MoeFfn),
560 DenseMoe(Box<DenseMoeFfn>),
567}
568
569pub struct DenseMoeFfn {
571 pub dense: DenseFfn,
572 pub moe: MoeFfn,
573 pub post_norm_1: Vec<f32>,
575 pub pre_norm_2: Vec<f32>,
578 pub post_norm_2: Vec<f32>,
580}
581
582pub struct MoeFfn {
583 pub router: QTensor,
585 pub experts: Vec<DenseFfn>,
586 pub top_k: usize,
587 pub norm_topk_prob: bool,
588 pub router_sigmoid: bool,
591 pub expert_bias: Option<Vec<f32>>,
595 pub routed_scaling: f32,
598 pub route_tau: Option<f32>,
604 pub shared: Option<(DenseFfn, Option<QTensor>)>,
607 pub stats: std::cell::RefCell<Vec<u64>>,
611 pub act_sq: std::cell::RefCell<Vec<f64>>,
618 pub act_rows: std::cell::RefCell<Vec<f32>>,
624 pub mask: Option<Vec<bool>>,
629 pub per_expert_scale: Option<Vec<f32>>,
632 pub router_input_norm: bool,
636 pub resonance: Option<Resonance>,
640 pub grown: Vec<GrownExpert>,
646}
647
648#[derive(Debug, Clone, PartialEq, Eq)]
650pub struct GrownExpert {
651 pub record: String,
653 pub record_index: usize,
655 pub layer: usize,
656 pub expert: usize,
660}
661
662pub struct Resonance {
664 pub mu: Vec<f32>,
666 pub u: Vec<f32>,
668 pub k: usize,
669 pub bias: Vec<f32>,
671 pub shell: Vec<f32>,
677}
678
679static GROWTH_SHELL: std::sync::atomic::AtomicU8 = std::sync::atomic::AtomicU8::new(0);
682
683pub fn growth_shell_enabled() -> bool {
688 use std::sync::atomic::Ordering;
689 match GROWTH_SHELL.load(Ordering::Relaxed) {
690 1 => true,
691 2 => false,
692 _ => {
693 let off = std::env::var("CMF_GROWTH_SHELL")
694 .map(|v| v.eq_ignore_ascii_case("off") || v == "0")
695 .unwrap_or(false);
696 GROWTH_SHELL.store(if off { 2 } else { 1 }, Ordering::Relaxed);
697 !off
698 }
699 }
700}
701
702pub fn set_growth_shell(on: Option<bool>) {
707 GROWTH_SHELL.store(
708 match on {
709 Some(true) => 1,
710 Some(false) => 2,
711 None => 0,
712 },
713 std::sync::atomic::Ordering::Relaxed,
714 );
715}
716
717impl Resonance {
718 pub fn has_shell(&self) -> bool {
720 self.shell.iter().any(|s| s.is_finite())
721 }
722
723 pub fn effective_shell(&self, ne: usize) -> Vec<f32> {
726 let mut out = vec![f32::INFINITY; ne];
727 if growth_shell_enabled() {
728 for (o, s) in out.iter_mut().zip(&self.shell) {
729 *o = *s;
730 }
731 }
732 out
733 }
734
735 pub fn scores(&self, x: &[f32], out: &mut [f32]) {
740 let h = x.len();
741 let ne = out.len();
742 let shell_on = growth_shell_enabled() && !self.shell.is_empty();
743 for e in 0..ne {
744 let mu = &self.mu[e * h..(e + 1) * h];
745 let mut d2 = 0.0f32;
746 for j in 0..h {
747 let d = x[j] - mu[j];
748 d2 += d * d;
749 }
750 let mut proj = 0.0f32;
751 for i in 0..self.k {
752 let u = &self.u[(e * self.k + i) * h..(e * self.k + i + 1) * h];
753 let mut p = 0.0f32;
754 for j in 0..h {
755 p += (x[j] - mu[j]) * u[j];
756 }
757 proj += p * p;
758 }
759 let err = d2 - proj;
760 out[e] = self.bias.get(e).copied().unwrap_or(0.0) - err;
761 if shell_on && err > self.shell.get(e).copied().unwrap_or(f32::INFINITY) {
762 out[e] = f32::NEG_INFINITY;
763 }
764 }
765 }
766}
767
768pub enum AttnKind {
771 Full {
773 wq: QTensor,
774 wk: QTensor,
775 wv: QTensor,
776 wo: QTensor,
777 q_norm: Option<Vec<f32>>,
778 k_norm: Option<Vec<f32>>,
779 output_gate: bool,
780 softplus_gate: Option<(QTensor, bool)>,
784 bias: Option<(Vec<f32>, Vec<f32>, Vec<f32>)>,
786 },
787 Linear(VmfPhaseWeights),
789 LinearGdn(GdnWeights),
791 ShortConv(ShortConvWeights),
794 Mla(Box<MlaWeights>),
802 Kda(Box<crate::linear_core::KdaWeights>),
806 Bounded(Box<crate::bounded::BoundedWeights>),
811}
812
813pub struct MlaWeights {
815 pub q_proj: QTensor,
819 pub q_a: Option<QTensor>,
822 pub q_a_norm: Option<Vec<f32>>,
823 pub kv_a: QTensor,
825 pub kv_a_norm: Vec<f32>,
827 pub kv_b: QTensor,
829 pub o_proj: QTensor,
831 pub nh: usize,
832 pub qk_rope: usize,
833 pub qk_nope: usize,
834 pub v_dim: usize,
835 pub lora: usize,
836 pub scale: f32,
838 pub nope: bool,
840}
841
842pub struct MtpModule {
847 pub enorm: Vec<f32>,
848 pub hnorm: Vec<f32>,
849 pub eh_proj: QTensor,
851 pub layer: LayerWeights,
852 pub final_norm: Vec<f32>,
853 pub kv: crate::kv_cache::LayerKvCache,
854}
855
856#[cfg(target_os = "macos")]
863enum MetalRowsItem<'a> {
864 Gdn {
865 run: Vec<crate::gpu_metal::GdnGpuLayer<'a>>,
866 first: usize,
867 },
868 Attn {
869 l: crate::gpu_metal::AttnGpuLayer<'a>,
870 li: usize,
871 q_norm: Option<&'a [f32]>,
872 k_norm: Option<&'a [f32]>,
873 output_gate: bool,
874 },
875}
876
877#[cfg(target_os = "macos")]
878struct MetalVerifyPending {
879 graph: crate::gpu_metal::VerifyGraph,
880 gdn_layers: Vec<usize>,
881 attn_layers: Vec<(usize, usize)>,
882}
883
884#[cfg(target_os = "macos")]
888struct MetalWarmPending {
889 graph: crate::gpu_metal::VerifyGraph,
890 cpu_stored: usize,
891 b: usize,
892}
893
894#[cfg(target_os = "macos")]
895enum MetalRowsRun {
896 Declined,
898 Failed,
901 Completed(MetalVerifyPending),
902}
903
904#[cfg(target_os = "macos")]
905enum MetalPrefillOutcome {
906 Declined,
907 Failed,
908 Completed(Vec<f32>),
909}
910
911#[cfg(target_os = "macos")]
912enum MetalBatchNllOutcome {
913 Declined,
914 Failed(String),
915 Completed(f64, usize),
916}
917
918#[derive(Clone, Copy)]
922enum SpecTrial {
923 Spec {
924 t0: std::time::Instant,
925 gen0: usize,
926 rounds: usize,
927 },
928 Plain {
929 t0: std::time::Instant,
930 gen0: usize,
931 },
932 Decided {
933 spec: bool,
934 recheck_at: usize,
935 },
936}
937
938pub(crate) fn spec_time_level() -> u8 {
942 static L: std::sync::OnceLock<u8> = std::sync::OnceLock::new();
943 *L.get_or_init(|| match std::env::var("CMF_GRAPH_SPEC_TIME") {
944 Ok(v) => v.trim().parse::<u8>().map(|n| n.max(1)).unwrap_or(1),
945 Err(_) => 0,
946 })
947}
948
949struct SpecStampLog {
955 t_last: std::time::Instant,
956 items: Vec<(&'static str, f32)>,
957}
958
959static SPEC_STAMPS: std::sync::Mutex<Option<SpecStampLog>> = std::sync::Mutex::new(None);
960
961pub(crate) fn spec_stamp(name: &'static str) {
962 if spec_time_level() == 0 {
963 return;
964 }
965 if let Ok(mut g) = SPEC_STAMPS.lock() {
966 if let Some(log) = g.as_mut() {
967 let now = std::time::Instant::now();
968 log.items
969 .push((name, (now - log.t_last).as_secs_f32() * 1e3));
970 log.t_last = now;
971 }
972 }
973}
974
975fn spec_stamps_begin() {
976 if spec_time_level() == 0 {
977 return;
978 }
979 if let Ok(mut g) = SPEC_STAMPS.lock() {
980 *g = Some(SpecStampLog {
981 t_last: std::time::Instant::now(),
982 items: Vec::with_capacity(64),
983 });
984 }
985}
986
987fn spec_stamps_take() -> Vec<(&'static str, f32)> {
988 SPEC_STAMPS
989 .lock()
990 .ok()
991 .and_then(|mut g| g.take())
992 .map(|l| l.items)
993 .unwrap_or_default()
994}
995
996fn spec_stamps_format(items: &[(&'static str, f32)]) -> String {
999 let mut agg: Vec<(&'static str, f32, u32)> = Vec::with_capacity(items.len());
1000 for &(n, ms) in items {
1001 match agg.iter_mut().find(|e| e.0 == n) {
1002 Some(e) => {
1003 e.1 += ms;
1004 e.2 += 1;
1005 }
1006 None => agg.push((n, ms, 1)),
1007 }
1008 }
1009 let mut s = String::with_capacity(agg.len() * 16);
1010 for (n, ms, k) in agg {
1011 if k > 1 {
1012 s.push_str(&format!("{n} {ms:.1}/{k} "));
1013 } else {
1014 s.push_str(&format!("{n} {ms:.1} "));
1015 }
1016 }
1017 s
1018}
1019
1020#[derive(Default, Clone, Copy)]
1042struct SpecMon {
1043 round_ms: f64,
1044 tokens: f64,
1045 plain_ms: f64,
1046 n: u32,
1047 fails: u32,
1048 metal: bool,
1049}
1050
1051const SPEC_PROXY_TOKENS: f64 = 3.5;
1054const SPEC_PLAIN_MIN_MS: f64 = 200.0;
1057
1058impl SpecMon {
1059 fn round(&mut self, dt_ms: f64, produced: usize) {
1060 self.n += 1;
1061 if self.n == 1 {
1062 return; }
1064 let a = if self.n == 2 { 1.0 } else { 0.3 };
1065 self.round_ms += a * (dt_ms - self.round_ms);
1066 self.tokens += a * (produced as f64 - self.tokens);
1067 }
1068 fn pays(&self) -> bool {
1069 if self.plain_ms > 0.0 {
1070 self.tokens * self.plain_ms > self.round_ms * 1.03
1071 } else {
1072 self.metal && self.tokens >= SPEC_PROXY_TOKENS
1073 }
1074 }
1075 fn plain_done(&self, t0: std::time::Instant, gen0: usize, generated: usize) -> bool {
1077 let n = generated.saturating_sub(gen0);
1078 if n >= 8 {
1079 return true;
1080 }
1081 self.metal && n >= 2 && t0.elapsed().as_secs_f64() * 1e3 >= SPEC_PLAIN_MIN_MS
1082 }
1083}
1084
1085pub const KV_PREFIX_TAIL: usize = 128;
1088
1089#[derive(Debug, Clone, Default)]
1095pub struct KvPrefix {
1096 len: usize,
1097 hash: u64,
1098 tail: Vec<u32>,
1099 device: bool,
1104}
1105
1106impl KvPrefix {
1107 #[inline]
1108 fn fold(mut h: u64, ids: &[u32]) -> u64 {
1109 for &id in ids {
1110 h ^= id as u64;
1111 h = h.wrapping_mul(0x100000001b3);
1112 h ^= h >> 29;
1113 }
1114 h
1115 }
1116
1117 pub fn clear(&mut self) {
1118 self.len = 0;
1119 self.hash = 0xcbf29ce484222325;
1120 self.tail.clear();
1121 self.device = false;
1122 }
1123
1124 pub fn on_device(&self) -> bool {
1126 self.device
1127 }
1128
1129 pub fn set_on_device(&mut self, device: bool) {
1131 self.device = device;
1132 }
1133
1134 pub fn len(&self) -> usize {
1136 self.len
1137 }
1138
1139 pub fn is_empty(&self) -> bool {
1140 self.len == 0
1141 }
1142
1143 pub fn tail_len(&self) -> usize {
1145 self.tail.len()
1146 }
1147
1148 pub fn set(&mut self, ids: &[u32]) {
1150 self.clear();
1151 self.extend(ids);
1152 }
1153
1154 pub fn extend(&mut self, more: &[u32]) {
1156 if self.len == 0 && self.hash == 0 {
1157 self.hash = 0xcbf29ce484222325;
1158 }
1159 self.hash = Self::fold(self.hash, more);
1160 self.len += more.len();
1161 if more.len() >= KV_PREFIX_TAIL {
1162 self.tail.clear();
1163 self.tail.extend_from_slice(&more[more.len() - KV_PREFIX_TAIL..]);
1164 } else {
1165 let drop = (self.tail.len() + more.len()).saturating_sub(KV_PREFIX_TAIL);
1166 self.tail.drain(..drop);
1167 self.tail.extend_from_slice(more);
1168 }
1169 }
1170
1171 pub fn extension(&self, ids: &[u32]) -> usize {
1175 if self.len == 0 || ids.len() <= self.len {
1176 return 0;
1177 }
1178 let t = self.tail.len();
1179 if ids[self.len - t..self.len] != self.tail[..] {
1180 return 0;
1181 }
1182 if Self::fold(0xcbf29ce484222325, &ids[..self.len]) != self.hash {
1183 return 0;
1184 }
1185 self.len
1186 }
1187}
1188
1189pub struct GenerateResult {
1191 pub text: String,
1192 pub token_ids: Vec<u32>,
1193 pub prompt_tokens: usize,
1194 pub tokens_generated: usize,
1195 pub finish_reason: String,
1196 pub mtp_drafted: usize,
1198 pub mtp_accepted: usize,
1199 pub token_confidence: Vec<f32>,
1204 pub traces: Vec<TokenTrace>,
1207}
1208
1209#[derive(Clone, Debug)]
1214pub struct TokenTrace {
1215 pub t: usize,
1217 pub token_id: u32,
1219 pub confidence: f32,
1221 pub active_skill: Option<String>,
1223 pub recon: Option<f32>,
1227 pub switched: bool,
1230}
1231
1232#[cfg_attr(not(test), allow(dead_code))]
1237fn top1_prob_t(logits: &[f32], id: u32, temp: f32) -> f32 {
1238 let t = if temp > 1e-3 { temp } else { 1.0 };
1239 let max = logits.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
1240 let sum: f32 = logits.iter().map(|&v| ((v - max) / t).exp()).sum();
1241 if sum > 0.0 {
1242 (((logits[id as usize] - max) / t).exp()) / sum
1243 } else {
1244 0.0
1245 }
1246}
1247
1248fn prefill_batched() -> bool {
1251 std::env::var("CMF_PREFILL")
1252 .map(|v| v != "seq")
1253 .unwrap_or(true)
1254}
1255
1256#[inline]
1260fn nll_graph_policy(
1261 unmasked: bool,
1262 prefer_graph: bool,
1263 native_metal: bool,
1264) -> (bool, bool) {
1265 let graph_quality = unmasked && prefer_graph;
1266 let fused_head_quality = graph_quality && native_metal;
1267 (graph_quality, fused_head_quality)
1268}
1269
1270#[derive(Clone, Copy)]
1274enum PrefillIn<'a> {
1275 Ids(&'a [u32]),
1276 Hidden(&'a [f32]),
1277}
1278
1279impl Pipeline {
1286 fn can_prefill_batched(&self) -> bool {
1287 #[cfg(test)]
1288 let force_serial = self.nll_test_force_serial;
1289 #[cfg(not(test))]
1290 let force_serial = false;
1291 prefill_batched() && !force_serial && !self.weights.layers.is_empty()
1292 }
1293
1294 fn automatic_gpu_prefix(&self) -> Option<usize> {
1297 let (model, _, _, _) = self.weights.embed_tokens.graph_weight()?;
1298 crate::gpu::automatic_layer_prefix(&model, self.num_layers, self.physical_layers)
1299 }
1300
1301 pub fn prefill_chunk(&self) -> usize {
1305 let env = env_prefill_chunk();
1306 if env.is_some() || ChunkHost::here() != ChunkHost::Other {
1307 return prefill_chunk_rule(env, ChunkHost::here(), false);
1308 }
1309 prefill_chunk_rule(None, ChunkHost::Other, self.chunk_stack_facts().dense_on_discrete())
1310 }
1311
1312 fn chunk_stack_facts(&self) -> ChunkStackFacts {
1313 let plain_dense = !self.weights.layers.is_empty()
1314 && self.g3n.is_none()
1315 && self.dsv4.is_none()
1316 && self.dsv41.is_none()
1317 && self.qwen4_exp.is_none()
1318 && self.weights.layers.iter().all(|lw| {
1319 matches!(lw.attn, AttnKind::Full { .. }) && matches!(lw.ffn, FfnKind::Dense(_))
1320 });
1321 let gpu_on = crate::gpu::enabled();
1322 ChunkStackFacts {
1323 plain_dense,
1324 discrete: gpu_on && crate::gpu::discrete(),
1325 gpu_on,
1326 capacity_split: std::env::var_os("CMF_GPU_LAYERS").is_some()
1329 || (plain_dense && gpu_on && self.automatic_gpu_prefix().is_some()),
1330 multi_gpu: self.gpu_plan.is_some(),
1331 o1: self.o1_active(),
1332 }
1333 }
1334}
1335
1336pub fn prefill_chunk() -> usize {
1345 prefill_chunk_rule(env_prefill_chunk(), ChunkHost::here(), false)
1346}
1347
1348fn env_prefill_chunk() -> Option<usize> {
1349 std::env::var("CMF_PREFILL_CHUNK")
1350 .ok()
1351 .and_then(|v| v.parse::<usize>().ok())
1352}
1353
1354#[derive(Clone, Copy, Debug, PartialEq, Eq)]
1356enum ChunkHost {
1357 Macos,
1358 Aarch64,
1360 Other,
1362}
1363
1364impl ChunkHost {
1365 fn here() -> Self {
1366 if cfg!(target_os = "macos") {
1367 ChunkHost::Macos
1368 } else if cfg!(target_arch = "aarch64") {
1369 ChunkHost::Aarch64
1370 } else {
1371 ChunkHost::Other
1372 }
1373 }
1374}
1375
1376const DISCRETE_DENSE_PREFILL_CHUNK: usize = 512;
1383
1384fn prefill_chunk_rule(env: Option<usize>, host: ChunkHost, dense_on_discrete: bool) -> usize {
1390 if let Some(n) = env {
1391 return n.max(1);
1392 }
1393 match host {
1394 ChunkHost::Macos => 512,
1395 ChunkHost::Aarch64 => 256,
1398 ChunkHost::Other if dense_on_discrete => DISCRETE_DENSE_PREFILL_CHUNK,
1399 ChunkHost::Other => 48,
1400 }
1401}
1402
1403#[derive(Clone, Copy, Debug, Default)]
1405struct ChunkStackFacts {
1406 plain_dense: bool,
1409 discrete: bool,
1411 gpu_on: bool,
1413 capacity_split: bool,
1415 multi_gpu: bool,
1417 o1: bool,
1419}
1420
1421impl ChunkStackFacts {
1422 fn dense_on_discrete(self) -> bool {
1423 self.plain_dense
1424 && self.discrete
1425 && self.gpu_on
1426 && !self.capacity_split
1427 && !self.multi_gpu
1428 && !self.o1
1429 }
1430}
1431
1432#[inline]
1438fn mtp_prefill_pair_count(start: usize, end: usize, input_len: usize) -> usize {
1439 if end <= start || start >= input_len {
1440 return 0;
1441 }
1442 let rows = (end.min(input_len) - start).min(input_len - start);
1443 if end < input_len {
1444 rows
1445 } else {
1446 rows.saturating_sub(1)
1447 }
1448}
1449
1450pub type TokenCallback = Box<dyn FnMut(&str) -> bool + Send>;
1452
1453#[derive(Debug, Clone, Copy, PartialEq, Eq)]
1455pub(crate) struct ReuseLayer {
1456 pub full: bool,
1459 pub host_rows: usize,
1461 pub device_rows: Option<usize>,
1463 pub device_state: bool,
1465}
1466
1467#[derive(Debug, Clone, PartialEq, Eq)]
1469pub(crate) enum ReusePlan {
1470 Ready,
1472 Pull(Vec<(usize, usize, usize)>),
1475 Fresh,
1477}
1478
1479pub(crate) fn kv_reuse_plan(reuse_from: usize, layers: &[ReuseLayer]) -> ReusePlan {
1489 let mut pulls = Vec::new();
1490 for (li, l) in layers.iter().enumerate() {
1491 if !l.full {
1492 if l.device_state {
1493 return ReusePlan::Fresh;
1494 }
1495 continue;
1496 }
1497 if l.host_rows == reuse_from {
1498 continue;
1499 }
1500 if l.host_rows < reuse_from && l.device_rows.is_some_and(|d| d >= reuse_from) {
1501 pulls.push((li, l.host_rows, reuse_from));
1502 continue;
1503 }
1504 return ReusePlan::Fresh;
1505 }
1506 if pulls.is_empty() {
1507 ReusePlan::Ready
1508 } else {
1509 ReusePlan::Pull(pulls)
1510 }
1511}
1512
1513impl Pipeline {
1514 fn clear_sequence_state(&mut self) {
1522 #[cfg(target_os = "macos")]
1525 let _ = crate::gpu_metal::wait_replay();
1526 self.kv_cache.clear();
1527 self.clear_history();
1530 self.graph_logits = None;
1531 if let Some(b) = &mut self.dsv41 {
1532 b.3.clear();
1533 }
1534 crate::gpu::graph_kv_reset(self.graph_kv_id);
1535 crate::gpu::graph_kv_reset(self.mtp_kv_id());
1540 }
1541
1542 fn prepare_kv_reuse(&mut self, reuse_from: usize) -> bool {
1549 if self.graph_prefill_preferred() {
1550 return true;
1551 }
1552 let kv_id = self.graph_kv_id;
1553 let layers: Vec<ReuseLayer> = (0..self.num_layers)
1554 .map(|li| {
1555 let full = matches!(
1556 self.weights.layers[self.phys_layer(li)].attn,
1557 AttnKind::Full { .. }
1558 );
1559 ReuseLayer {
1560 full,
1561 host_rows: self.kv_cache.layers[li].seq_len,
1562 device_rows: crate::gpu::graph_kv_stored(kv_id, li),
1563 device_state: crate::gpu::graph_state_resident(kv_id, li),
1564 }
1565 })
1566 .collect();
1567 if layers
1571 .iter()
1572 .all(|l| l.device_rows.is_none() && !l.device_state)
1573 {
1574 return true;
1575 }
1576 let plan = kv_reuse_plan(reuse_from, &layers);
1577 let (what, rows, n) = match &plan {
1578 ReusePlan::Ready => ("host ready", 0, 0),
1579 ReusePlan::Fresh => ("fresh", 0, 0),
1580 ReusePlan::Pull(p) => (
1581 "pull",
1582 p.iter().map(|&(_, a, b)| b - a).max().unwrap_or(0),
1583 p.len(),
1584 ),
1585 };
1586 let t0 = std::time::Instant::now();
1587 let ok = self.apply_kv_reuse_plan(reuse_from, plan, &layers);
1588 if std::env::var("CMF_PREFILL_PROF").is_ok() {
1589 eprintln!(
1590 "kv-reuse: {what}{}: {rows} device row(s) × {n} layer(s) to the host in {:.2} ms",
1591 if ok { "" } else { " (failed → fresh)" },
1592 t0.elapsed().as_secs_f64() * 1e3
1593 );
1594 }
1595 ok
1596 }
1597
1598 fn apply_kv_reuse_plan(
1599 &mut self,
1600 reuse_from: usize,
1601 plan: ReusePlan,
1602 layers: &[ReuseLayer],
1603 ) -> bool {
1604 let kv_id = self.graph_kv_id;
1605 match plan {
1606 ReusePlan::Fresh => return false,
1607 ReusePlan::Ready => {}
1608 ReusePlan::Pull(pulls) => {
1609 let (nkv, hd) = {
1614 let c = &self.kv_cache.layers[pulls[0].0];
1615 (c.num_kv_heads, c.head_dim)
1616 };
1617 let uniform = pulls.iter().all(|&(li, _, _)| {
1618 let c = &self.kv_cache.layers[li];
1619 (c.num_kv_heads, c.head_dim) == (nkv, hd)
1620 });
1621 let batched = if uniform {
1622 crate::gpu::graph_kv_read_rows(kv_id, &pulls, nkv, hd)
1623 } else {
1624 None
1625 };
1626 let rows: Vec<(Vec<f32>, Vec<f32>)> = match batched {
1627 Some(rows) => rows,
1628 None => {
1629 let mut rows = Vec::with_capacity(pulls.len());
1630 for &(li, from, to) in &pulls {
1631 let (lnkv, lhd) = {
1632 let c = &self.kv_cache.layers[li];
1633 (c.num_kv_heads, c.head_dim)
1634 };
1635 let Some((k, v, first_valid)) =
1636 crate::gpu::graph_kv_pull_host(kv_id, li, from, to, lnkv, lhd)
1637 else {
1638 return false;
1639 };
1640 let need_from = match self.layer_window(li) {
1644 Some(w) => from.max((to + 1).saturating_sub(w)),
1645 None => from,
1646 };
1647 if first_valid > need_from {
1648 return false;
1649 }
1650 rows.push((k, v));
1651 }
1652 rows
1653 }
1654 };
1655 for ((li, from, to), (k, v)) in pulls.into_iter().zip(rows) {
1656 let cache = &mut self.kv_cache.layers[li];
1657 let row = cache.num_kv_heads * cache.head_dim;
1658 for p in 0..to - from {
1659 cache.append(&k[p * row..(p + 1) * row], &v[p * row..(p + 1) * row], &[]);
1660 }
1661 if cache.seq_len != to {
1662 return false;
1663 }
1664 }
1665 }
1666 }
1667 for (li, l) in layers.iter().enumerate() {
1671 if l.full
1672 && l.device_rows.is_some_and(|d| d > reuse_from)
1673 && !crate::gpu::graph_kv_set_stored(kv_id, li, reuse_from)
1674 {
1675 return false;
1676 }
1677 }
1678 true
1679 }
1680
1681 fn finish_generation(
1687 &mut self,
1688 mtp: &mut Option<MtpModule>,
1689 router: &mut Option<crate::swarm::DynRouter>,
1690 clear_sequence: bool,
1691 ) {
1692 if router.is_some() {
1696 let _ = self.set_active_skill(None);
1697 }
1698 #[cfg(target_os = "macos")]
1705 let clear_sequence = clear_sequence || !crate::gpu_metal::wait_replay();
1706 if clear_sequence {
1707 self.clear_sequence_state();
1708 if let Some(m) = mtp.as_mut() {
1709 m.kv.clear();
1715 }
1716 if let Some(m) = self.mtp.as_mut() {
1717 m.kv.clear();
1721 }
1722 }
1723 self.graph_want_logits = false;
1724 self.graph_head_required = false;
1725 self.graph_logits = None;
1726 self.graph_failed
1727 .store(false, std::sync::atomic::Ordering::Relaxed);
1728 self.cancel
1729 .store(false, std::sync::atomic::Ordering::Relaxed);
1730 self.dyn_router = router.take().or(self.dyn_router.take());
1731 self.mtp = mtp.take().or(self.mtp.take());
1732 self.mtp_graph_mode = None;
1733 self.spec_forced = None;
1734 }
1735
1736 fn check_forward_graph(&mut self, phase: &str, pos: usize) -> Result<(), String> {
1740 if self
1741 .graph_failed
1742 .swap(false, std::sync::atomic::Ordering::Relaxed)
1743 {
1744 self.cancel
1745 .store(false, std::sync::atomic::Ordering::Relaxed);
1746 self.clear_sequence_state();
1747 self.graph_logits = None;
1748 self.graph_want_logits = false;
1749 self.graph_head_required = false;
1750 return Err(format!("GPU graph failed during {phase} at position {pos}"));
1751 }
1752 Ok(())
1753 }
1754
1755 #[cfg(target_os = "macos")]
1756 fn fail_metal_graph(&mut self, reason: &str) {
1757 crate::pipeline::METAL_GRAPH_ERRORS
1758 .fetch_add(1, std::sync::atomic::Ordering::Relaxed);
1759 self.clear_sequence_state();
1760 self.graph_logits = None;
1761 self.graph_failed
1762 .store(true, std::sync::atomic::Ordering::Relaxed);
1763 self.cancel
1764 .store(true, std::sync::atomic::Ordering::Relaxed);
1765 tracing::error!("native Metal TokenGraph failed closed: {reason}");
1766 }
1767
1768 fn nll_begin(&mut self) -> Result<(), String> {
1773 if self
1774 .graph_failed
1775 .swap(false, std::sync::atomic::Ordering::Relaxed)
1776 {
1777 self.cancel
1778 .store(false, std::sync::atomic::Ordering::Relaxed);
1779 self.clear_sequence_state();
1780 self.graph_logits = None;
1781 self.graph_want_logits = false;
1782 self.graph_head_required = false;
1783 return Err("GPU graph failed before NLL scoring".to_string());
1784 }
1785 self.clear_sequence_state();
1786 self.graph_logits = None;
1787 self.graph_want_logits = false;
1788 self.graph_head_required = false;
1789 Ok(())
1790 }
1791
1792 fn nll_end(&mut self) {
1796 self.clear_sequence_state();
1797 self.graph_logits = None;
1798 self.graph_want_logits = false;
1799 self.graph_head_required = false;
1800 self.graph_failed
1801 .store(false, std::sync::atomic::Ordering::Relaxed);
1802 }
1803
1804 fn nll_check_graph(&mut self, phase: &str, pos: usize) -> Result<(), String> {
1807 #[cfg(test)]
1808 if self.nll_test_fail_at == Some(pos) {
1809 self.nll_test_fail_at = None;
1810 self.graph_failed
1811 .store(true, std::sync::atomic::Ordering::Relaxed);
1812 self.cancel
1813 .store(true, std::sync::atomic::Ordering::Relaxed);
1814 }
1815 if self
1816 .graph_failed
1817 .swap(false, std::sync::atomic::Ordering::Relaxed)
1818 {
1819 self.cancel
1820 .store(false, std::sync::atomic::Ordering::Relaxed);
1821 self.clear_sequence_state();
1822 self.graph_logits = None;
1823 self.graph_want_logits = false;
1824 return Err(format!(
1825 "GPU graph failed during NLL {phase} at position {pos}"
1826 ));
1827 }
1828 Ok(())
1829 }
1830
1831 #[inline]
1835 pub fn phys_layer(&self, virtual_idx: usize) -> usize {
1836 virtual_idx % self.physical_layers
1837 }
1838
1839 #[inline]
1842 pub fn is_loop_end(&self, virtual_idx: usize) -> bool {
1843 self.loop_final_norm && (virtual_idx + 1) % self.physical_layers == 0
1844 }
1845
1846 #[allow(clippy::too_many_arguments)]
1848
1849 #[cfg(target_os = "macos")]
1868 fn graph_prefill_preferred(&self) -> bool {
1869 let graph_force = crate::gpu::q1_force() || crate::gpu::q2tp_gpu_opt_in();
1870 if !crate::gpu::enabled_here()
1871 || !graph_force
1872 || std::env::var("CMF_GPU_BLOCK")
1873 .map(|v| v == "0")
1874 .unwrap_or(false)
1875 || std::env::var("CMF_PREFILL_GRAPH").as_deref() == Ok("0")
1878 {
1879 return false;
1880 }
1881 self.weights
1882 .layers
1883 .iter()
1884 .any(|lw| {
1885 matches!(&lw.attn, AttnKind::LinearGdn(w) if w.in_proj_qkv.metal_graph_parts().is_some())
1886 })
1887 }
1888
1889 #[cfg(not(target_os = "macos"))]
1896 fn batch_prefix_prefill(&self) -> bool {
1897 let forced = match std::env::var("CMF_BATCH_PREFIX").as_deref() {
1898 Ok("0") => return false,
1899 Ok("1") => true,
1900 _ => false,
1901 };
1902 crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Prefill)
1903 && crate::gpu::enabled_here()
1904 && !self.graph_refused()
1905 && (forced || self.graph_attn_decline_reason().is_some())
1906 && self.wgpu_graph_attn_decline().is_none()
1907 && self.attn_softcap == 0.0
1908 && self
1909 .weights
1910 .layers
1911 .iter()
1912 .any(|lw| matches!(&lw.ffn, FfnKind::Moe(_)))
1913 && self.automatic_gpu_prefix().is_some()
1914 }
1915
1916 #[cfg(not(target_os = "macos"))]
1917 fn graph_prefill_preferred(&self) -> bool {
1918 let graph_on = crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Prefill);
1926 if !graph_on || !crate::gpu::enabled_here() {
1927 return false;
1928 }
1929 if self.embryo_resident_eligible() {
1933 return crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Decode);
1938 }
1939 if self.o1_active() {
1953 return false;
1954 }
1955 if self.wgpu_graph_attn_decline().is_some() {
1959 return false;
1960 }
1961 if self
1962 .weights
1963 .layers
1964 .iter()
1965 .any(|lw| matches!(&lw.attn, AttnKind::LinearGdn(_)))
1966 {
1967 return true;
1968 }
1969 self.weights
1980 .layers
1981 .iter()
1982 .any(|lw| matches!(&lw.ffn, FfnKind::Moe(_)))
1983 && self.automatic_gpu_prefix().is_none()
1984 }
1985
1986 #[cfg(target_os = "macos")]
1987 fn q1_graph_gpu(
1988 &mut self,
1989 start: usize,
1990 upto: Option<usize>,
1991 position: usize,
1992 h: &mut [f32],
1993 ) -> usize {
1994 let _mt0 = std::time::Instant::now(); use crate::gpu::{AttnGpuLayer, GdnGpuCfg, GdnGpuLayer, GraphDims, MetalFfn, TokenGraph};
1996 let graph_force = crate::gpu::q1_force() || crate::gpu::q2tp_gpu_opt_in();
1997 if self.attn_softcap > 0.0 || !crate::gpu::enabled_here()
1999 || !graph_force
2000 || std::env::var("CMF_GPU_BLOCK")
2001 .map(|v| v == "0")
2002 .unwrap_or(false)
2003 {
2004 if std::env::var("CMF_GRAPH_DBG").is_ok() {
2005 eprintln!(
2006 "block-graph: front gate (softcap={} enabled_here={} graph_force={})",
2007 self.attn_softcap > 0.0,
2008 crate::gpu::enabled_here(),
2009 graph_force,
2010 );
2011 }
2012 if self.graph_head_required {
2013 self.fail_metal_graph("native graph front gate refused");
2014 }
2015 return start;
2016 }
2017 if self.swa.is_some()
2021 || self.global_attn.is_some()
2022 || self.attention_heads_per_layer.is_some()
2023 || self.attn_v_norm
2024 || self.graph_attn_decline_reason().is_some()
2026 || self.weights.layers.iter().any(|lw| {
2027 lw.attn_out_norm.is_some()
2028 || lw.ffn_out_norm.is_some()
2029 || lw.layer_scale.is_some()
2030 || matches!(&lw.ffn, FfnKind::Dense(d) if d.act != Act::Silu)
2031 })
2032 {
2033 if let Some(reason) = self.graph_attn_decline_reason() {
2036 self.note_graph_decline("metal block graph", reason);
2037 }
2038 if std::env::var("CMF_GRAPH_DBG").is_ok() {
2039 eprintln!(
2040 "block-graph: arch ineligible (swa={} gattn={} hpl={} vnorm={} scale_delta={:.2e})",
2041 self.swa.is_some(),
2042 self.global_attn.is_some(),
2043 self.attention_heads_per_layer.is_some(),
2044 self.attn_v_norm,
2045 (self.attn_scale - 1.0 / (self.head_dim as f32).sqrt()).abs(),
2046 );
2047 }
2048 if self.graph_head_required {
2049 self.fail_metal_graph("native graph architecture gate refused");
2050 }
2051 return start;
2052 }
2053 let limit = upto
2056 .map(|u| u + 1)
2057 .unwrap_or(self.num_layers)
2058 .min(self.num_layers);
2059
2060 enum Item<'a> {
2061 Gdn {
2062 run: Vec<GdnGpuLayer<'a>>,
2063 first: usize,
2064 },
2065 Attn {
2066 l: AttnGpuLayer<'a>,
2067 li: usize,
2068 q_norm: Option<&'a [f32]>,
2069 k_norm: Option<&'a [f32]>,
2070 output_gate: bool,
2071 bias: Option<(&'a [f32], &'a [f32], &'a [f32])>,
2072 full_gpu: bool,
2075 },
2076 }
2077
2078 let attend_mode = std::env::var("CMF_GPU_ATTEND").unwrap_or_else(|_| "auto".into());
2085 let attend_contract = attend_mode != "0"
2086 && attend_mode != "off"
2087 && self.head_dim % 4 == 0
2088 && self.head_dim <= 256
2089 && self.rotary_dim >= 2
2090 && self.rotary_dim <= self.head_dim
2091 && (self.rotary_dim / 2) % 32 == 0
2092 && self.num_kv_heads > 0
2093 && self.num_heads % self.num_kv_heads == 0;
2094
2095 let mut plan: Vec<Item> = Vec::new();
2096 let mut model_ref: Option<std::sync::Arc<cortiq_core::CmfModel>> = None;
2097 let block_diag = std::env::var("CMF_GRAPH_DBG").is_ok();
2099 let mut scan = start;
2100 while scan < limit {
2101 let lw = &self.weights.layers[self.phys_layer(scan)];
2102 let ffn = match &lw.ffn {
2103 FfnKind::Dense(d) if d.segs.is_empty() => {
2104 let (Some(g), Some(u), Some(dn)) = (
2105 d.gate_proj.metal_graph_parts(),
2106 d.up_proj.metal_graph_parts(),
2107 d.down_proj.metal_graph_parts(),
2108 ) else {
2109 if block_diag {
2110 eprintln!(
2111 "block-graph: L{scan} FFN trio not graph-mappable — run ends"
2112 );
2113 }
2114 break;
2115 };
2116 MetalFfn::Dense {
2117 gate: g,
2118 up: u,
2119 down: dn,
2120 }
2121 }
2122 FfnKind::Moe(m) => {
2123 let Some(moe) = metal_moe_graph_parts(m, self.hidden_size) else {
2124 if block_diag {
2125 eprintln!(
2126 "block-graph: L{scan} MoE outside the graph contract — run ends"
2127 );
2128 }
2129 break;
2130 };
2131 if let QTensor::Mapped { model, .. } = &m.experts[0].gate_proj {
2132 model_ref.get_or_insert_with(|| model.clone());
2133 }
2134 MetalFfn::Moe(moe)
2135 }
2136 _ => {
2137 if block_diag {
2138 eprintln!("block-graph: L{scan} non-graph FFN — run ends");
2139 }
2140 break;
2141 }
2142 };
2143 match &lw.attn {
2144 AttnKind::LinearGdn(w) if self.gdn_cfg.is_some() => {
2145 let parts = (
2146 w.in_proj_qkv.metal_graph_parts(),
2147 w.in_proj_z.metal_graph_parts(),
2148 w.in_proj_a.f32_parts(),
2149 w.in_proj_b.f32_parts(),
2150 w.out_proj.metal_graph_parts(),
2151 );
2152 let (Some(qkv), Some(z), Some(a), Some(b), Some(out)) = parts else {
2153 if block_diag {
2154 eprintln!(
2155 "block-graph: L{scan} GDN parts refused (qkv={} z={} a_f32={} b_f32={} out={})",
2156 w.in_proj_qkv.metal_graph_parts().is_some(),
2157 w.in_proj_z.metal_graph_parts().is_some(),
2158 w.in_proj_a.f32_parts().is_some(),
2159 w.in_proj_b.f32_parts().is_some(),
2160 w.out_proj.metal_graph_parts().is_some(),
2161 );
2162 }
2163 break;
2164 };
2165 if let QTensor::Mapped { model, .. } = &w.in_proj_qkv {
2166 model_ref.get_or_insert_with(|| model.clone());
2167 }
2168 let gl = GdnGpuLayer {
2169 attn_norm: &lw.input_norm,
2170 post_norm: &lw.post_norm,
2171 qkv,
2172 z,
2173 a,
2174 b,
2175 out,
2176 ffn,
2177 conv1d: &w.conv1d,
2178 a_log: &w.a_log,
2179 dt_bias: &w.dt_bias,
2180 gnorm: &w.norm,
2181 };
2182 match plan.last_mut() {
2183 Some(Item::Gdn { run, .. }) => run.push(gl),
2184 _ => plan.push(Item::Gdn {
2185 run: vec![gl],
2186 first: scan,
2187 }),
2188 }
2189 }
2190 AttnKind::Full {
2191 wq,
2192 wk,
2193 wv,
2194 wo,
2195 q_norm,
2196 k_norm,
2197 output_gate,
2198 softplus_gate: None,
2199 bias,
2200 } if !self.kv_cache.layers[scan].o1_sealed()
2201 || std::env::var("CMF_O1_METAL").as_deref() == Ok("1") =>
2206 {
2207 let parts = (
2208 wq.metal_graph_parts(),
2209 wk.metal_graph_parts(),
2210 wv.metal_graph_parts(),
2211 wo.metal_graph_parts(),
2212 );
2213 let (Some(pq), Some(pk), Some(pv), Some(po)) = parts else {
2214 break;
2215 };
2216 if let QTensor::Mapped { model, .. } = wq {
2217 model_ref.get_or_insert_with(|| model.clone());
2218 }
2219 let cache = &self.kv_cache.layers[scan];
2220 let o1_metal = cache.o1.is_some()
2224 && std::env::var("CMF_O1_METAL").as_deref() == Ok("1")
2225 && cache.o1_views().is_some();
2226 let full_gpu = attend_contract
2227 && cache.mode == crate::kv_cache::KvMode::F32
2228 && (cache.o1.is_none() || o1_metal)
2229 && bias.is_none()
2230 && pq.1 == self.num_heads * self.head_dim * (1 + *output_gate as usize)
2231 && pk.1 == self.num_kv_heads * self.head_dim
2232 && pv.1 == self.num_kv_heads * self.head_dim
2233 && po.2 == self.num_heads * self.head_dim;
2234 plan.push(Item::Attn {
2235 l: AttnGpuLayer {
2236 attn_norm: &lw.input_norm,
2237 post_norm: &lw.post_norm,
2238 wq: pq,
2239 wk: pk,
2240 wv: pv,
2241 wo: po,
2242 ffn,
2243 },
2244 li: scan,
2245 q_norm: q_norm.as_deref(),
2246 k_norm: k_norm.as_deref(),
2247 output_gate: *output_gate,
2248 bias: bias
2249 .as_ref()
2250 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
2251 full_gpu,
2252 });
2253 }
2254 _ => break,
2255 }
2256 scan += 1;
2257 }
2258 let Some(model) = model_ref else {
2259 if std::env::var("CMF_GRAPH_DBG").is_ok() {
2260 eprintln!("q1-graph: no model ref (start {start}, scanned to {scan})");
2261 }
2262 if self.graph_head_required {
2263 self.fail_metal_graph("native graph has no mapped model reference");
2264 }
2265 return start;
2266 };
2267 if plan.is_empty() {
2268 if std::env::var("CMF_GRAPH_DBG").is_ok() {
2269 eprintln!("q1-graph: empty plan at layer {start}");
2270 }
2271 if self.graph_head_required {
2272 self.fail_metal_graph("native graph plan is empty");
2273 }
2274 return start;
2275 }
2276 let has_moe = plan.iter().any(|it| match it {
2277 Item::Gdn { run, .. } => run.iter().any(|l| matches!(l.ffn, MetalFfn::Moe(_))),
2278 Item::Attn { l, .. } => matches!(l.ffn, MetalFfn::Moe(_)),
2279 });
2280 let has_gdn = plan.iter().any(|it| matches!(it, Item::Gdn { .. }));
2281 let dev_attend = attend_contract
2282 && (self.head_dim <= 128
2283 || has_moe
2284 || (self.head_dim <= 256 && has_gdn)
2290 || attend_mode == "force"
2291 || attend_mode == "256");
2292 if !dev_attend {
2293 for it in &mut plan {
2294 if let Item::Attn { li, full_gpu, .. } = it {
2295 let keep_o1 = self.kv_cache.layers[*li].o1.is_some()
2298 && std::env::var("CMF_O1_METAL").as_deref() == Ok("1");
2299 if !keep_o1 {
2300 *full_gpu = false;
2301 }
2302 }
2303 }
2304 }
2305 if std::env::var("CMF_GRAPH_DBG").is_ok() {
2306 use std::sync::atomic::{AtomicBool, Ordering};
2307 static SAID: AtomicBool = AtomicBool::new(false);
2308 if !SAID.swap(true, Ordering::Relaxed) {
2309 let fg = plan
2310 .iter()
2311 .filter(|it| matches!(it, Item::Attn { full_gpu: true, .. }))
2312 .count();
2313 let att = plan
2314 .iter()
2315 .filter(|it| matches!(it, Item::Attn { .. }))
2316 .count();
2317 eprintln!(
2318 "q1-graph: plan of {} items from layer {start} to {scan} | dev_attend={dev_attend} full_gpu {fg}/{att} | hd={} rd={} nkv={} nh={}",
2319 plan.len(),
2320 self.head_dim,
2321 self.rotary_dim,
2322 self.num_kv_heads,
2323 self.num_heads,
2324 );
2325 }
2326 }
2327 let dims = GraphDims {
2328 hidden: self.hidden_size,
2329 eps: self.rms_eps as f32,
2330 gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
2331 };
2332 let Some(mut graph) = TokenGraph::new(&model, dims, h) else {
2333 if self.graph_head_required {
2334 self.fail_metal_graph("native TokenGraph allocation refused");
2335 }
2336 return start;
2337 };
2338 let gcfg = self.gdn_cfg.map(|cfg| GdnGpuCfg {
2339 nv: cfg.num_v_heads,
2340 nk: cfg.num_k_heads,
2341 dk: cfg.key_head_dim,
2342 dv: cfg.value_head_dim,
2343 kk: cfg.conv_kernel,
2344 hidden: self.hidden_size,
2345 inter: self.intermediate_size,
2346 c_dim: cfg.conv_dim(),
2347 eps: cfg.rms_eps as f32,
2348 gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
2349 });
2350 let mut valid = 0usize;
2354 let mut end = start;
2355 crate::gpu::stageprof(1, _mt0.elapsed()); if std::env::var("CMF_PLAN_DUMP").is_ok() {
2357 static ONCE: std::sync::Once = std::sync::Once::new();
2358 ONCE.call_once(|| {
2359 for it in &plan {
2360 match it {
2361 Item::Gdn { first, run } => {
2362 eprintln!("plan: Gdn first={first} len={}", run.len())
2363 }
2364 Item::Attn { li, full_gpu, .. } => {
2365 eprintln!("plan: Attn li={li} full_gpu={full_gpu}")
2366 }
2367 }
2368 }
2369 });
2370 }
2371 for item in &plan {
2372 let ok = match item {
2373 Item::Gdn { run, .. } => gcfg
2374 .as_ref()
2375 .map(|gc| run.iter().all(|l| graph.gdn_ok(l, gc)))
2376 .unwrap_or(false),
2377 Item::Attn { l, .. } => graph.attn_ok(l),
2378 };
2379 if !ok {
2380 if block_diag {
2381 eprintln!(
2382 "block-graph: plan item {} ({}) failed graph preflight",
2383 valid,
2384 match item {
2385 Item::Gdn { run, first } => format!("GDN run L{first}+{}", run.len()),
2386 Item::Attn { li, .. } => format!("Attn L{li}"),
2387 }
2388 );
2389 }
2390 break;
2391 }
2392 valid += 1;
2393 end += match item {
2394 Item::Gdn { run, .. } => run.len(),
2395 Item::Attn { .. } => 1,
2396 };
2397 }
2398 plan.truncate(valid);
2399 if plan.is_empty() {
2400 if self.graph_head_required {
2401 self.fail_metal_graph("native graph preflight produced no valid items");
2402 }
2403 return start;
2404 }
2405
2406 if self.graph_head_required && (upto.is_some() || end != self.num_layers) {
2407 self.fail_metal_graph("fused-head NLL requires a complete 64-layer graph");
2408 return start;
2409 }
2410
2411 let one_pass = |t: (usize, usize, usize)| {
2420 use cortiq_core::TensorDtype as D;
2421 matches!(
2422 model.tensors[t.0].dtype,
2423 D::Q4TiledP | D::Q4Tiled | D::Q4Block | D::Q8Row | D::Q8_2f | D::Q1
2424 )
2425 };
2426 let dense_fast = plan.iter().all(|it| match it {
2427 Item::Attn {
2428 l, li, full_gpu, ..
2429 } => {
2430 *full_gpu
2431 && self.kv_cache.layers[*li].o1.is_none()
2432 && [l.wq, l.wk, l.wv, l.wo].into_iter().all(one_pass)
2433 && match l.ffn {
2434 MetalFfn::Dense { gate, up, down } => {
2435 one_pass(gate) && one_pass(up) && one_pass(down)
2436 }
2437 _ => false,
2438 }
2439 }
2440 Item::Gdn { .. } => false,
2441 });
2442 let ab = crate::gpu_metal::dense_ab_arm().filter(|_| dense_fast);
2443 let _mv_fast = match ab {
2444 Some((bits, _)) => {
2445 graph.set_dense_concurrent_raw(bits & crate::gpu_metal::DENSE_CONC != 0);
2446 crate::gpu_metal::MvFastGuard::set_raw(bits)
2447 }
2448 None => {
2449 graph.set_dense_concurrent(dense_fast);
2450 crate::gpu_metal::MvFastGuard::set_bits(if dense_fast {
2451 crate::gpu_metal::DENSE_MV | crate::gpu_metal::DENSE_FUSE
2452 } else {
2453 0
2454 })
2455 }
2456 };
2457
2458 let inv_freq = self.inv_freq.clone();
2459 let pool = self.pool.clone();
2460 let (nh, nkv, hd, hs, rd, eps) = (
2461 self.num_heads,
2462 self.num_kv_heads,
2463 self.head_dim,
2464 self.hidden_size,
2465 self.rotary_dim,
2466 self.rms_eps,
2467 );
2468 let norm_style = self.norm_style;
2469 let gemma = norm_style == cortiq_core::NormStyle::Gemma;
2470 let want = self.gdn_cfg.map(|c| c.state_len()).unwrap_or(0);
2471 let kv_id = self.graph_kv_id;
2472 let mut pending: Vec<(usize, usize)> = Vec::new();
2475 let mut dev_attn: Vec<usize> = Vec::new();
2478 for item in &plan {
2479 let _xt0 = std::time::Instant::now();
2480 let _xkind: u32 = match item {
2481 Item::Gdn { .. } => 2,
2482 Item::Attn { .. } => 3,
2483 };
2484 if self.loop_final_norm {
2486 let item_start = match item {
2487 Item::Gdn { first, .. } => *first,
2488 Item::Attn { li, .. } => *li,
2489 };
2490 if item_start > start && self.is_loop_end(item_start - 1) {
2491 graph.encode_loop_norm(&self.weights.final_norm);
2492 }
2493 }
2494 match item {
2495 Item::Gdn { run, first } => {
2496 for l in &mut self.kv_cache.layers[*first..*first + run.len()] {
2497 if l.linear_state.len() != want {
2498 l.linear_state = vec![0f32; want];
2499 }
2500 }
2501 let ro: Vec<&[f32]> = self.kv_cache.layers[*first..*first + run.len()]
2502 .iter()
2503 .map(|l| l.linear_state.as_slice())
2504 .collect();
2505 let _ig = std::time::Instant::now();
2506 if !graph.encode_gdn_run(run, &ro, gcfg.as_ref().unwrap()) {
2507 tracing::error!("q1 graph: GDN run refused after validation");
2509 return start;
2510 }
2511 graph.commit_kind = 2;
2514 graph.commit();
2515 crate::gpu::stageprof(0, _ig.elapsed());
2516 pending.push((*first, run.len()));
2517 }
2518 Item::Attn {
2519 l,
2520 li,
2521 q_norm,
2522 k_norm,
2523 output_gate,
2524 bias,
2525 full_gpu,
2526 } => {
2527 let _ia = std::time::Instant::now();
2528 if *full_gpu {
2530 let cache = &self.kv_cache.layers[*li];
2531 let o1p = if cache.o1.is_some() {
2532 match cache.o1_views() {
2533 Some(views) => Some(crate::gpu::O1AttnParams {
2534 views,
2535 epoch: self.o1_epoch,
2536 }),
2537 None => None,
2539 }
2540 } else {
2541 None
2542 };
2543 let o1_layer = cache.o1.is_some();
2544 if o1_layer && o1p.is_none() {
2545 }
2547 let cpu_k: Vec<&[f32]> = (0..nkv).map(|g| cache.head_keys(g)).collect();
2548 let cpu_v: Vec<&[f32]> = (0..nkv).map(|g| cache.head_values(g)).collect();
2549 let cpu_stored = if o1_layer { 0 } else { cpu_k[0].len() / hd };
2550 let p = crate::gpu::AttnDeviceParams {
2551 kv_id,
2552 layer: *li,
2553 nh,
2554 nkv,
2555 hd,
2556 rd,
2557 position,
2558 scale: self.attn_scale,
2559 eps: eps as f32,
2560 gemma,
2561 late_qk_norm: self.qk_norm_after_rope,
2562 output_gate: *output_gate,
2563 q_norm: *q_norm,
2564 k_norm: *k_norm,
2565 inv_freq: &inv_freq,
2566 cpu_k,
2567 cpu_v,
2568 cpu_stored,
2569 o1: o1p,
2570 };
2571 let o1_bad = o1_layer && p.o1.is_none();
2572 if !o1_bad && graph.attn_device_ok(l, &p) && graph.encode_attn_device(l, &p)
2573 {
2574 if p.o1.is_none() {
2576 dev_attn.push(*li);
2577 }
2578 graph.commit_kind = 3;
2579 graph.commit();
2580 crate::gpu::stageprof(_xkind, _xt0.elapsed());
2584 continue;
2585 }
2586 }
2588 graph.encode_attn_prefix(l);
2589 if let Err(err) = graph.sync_checked() {
2590 self.fail_metal_graph(&err);
2591 return start;
2592 }
2593 if !pending.is_empty() {
2594 let idxs: Vec<usize> =
2595 pending.drain(..).flat_map(|(f, n)| f..f + n).collect();
2596 let mut outs: Vec<&mut [f32]> = self
2597 .kv_cache
2598 .layers
2599 .iter_mut()
2600 .enumerate()
2601 .filter(|(i, _)| idxs.binary_search(i).is_ok())
2602 .map(|(_, s)| s.linear_state.as_mut_slice())
2603 .collect();
2604 graph.read_states(&mut outs);
2605 }
2606 let mut q_raw = attention::take_buf(l.wq.1);
2607 let mut k = attention::take_buf(l.wk.1);
2608 let mut v = attention::take_buf(l.wv.1);
2609 graph.read_qkv(&mut q_raw, &mut k, &mut v);
2610 let cfg = QwenAttnCfg {
2611 num_heads: nh,
2612 num_kv_heads: nkv,
2613 head_dim: hd,
2614 hidden_size: hs,
2615 position,
2616 inv_freq: &inv_freq,
2617 rotary_dim: rd,
2618 scale: self.attn_scale,
2619 softcap: self.attn_softcap,
2620 window: None,
2621 v_norm: false,
2622 qk_norm_after_rope: self.qk_norm_after_rope,
2623 q_norm: *q_norm,
2624 k_norm: *k_norm,
2625 output_gate: *output_gate,
2626 softplus_gate: None,
2627 rope_scale: 1.0,
2628 bias: *bias,
2629 rms_eps: eps,
2630 norm_style,
2631 pool: pool.as_deref(),
2632 v_head_dim: hd,
2633 };
2634 let oracle = std::env::var("CMF_ATTN_ORACLE").as_deref() == Ok("1")
2637 || std::env::var("CMF_ATTN_DUMP").is_ok();
2638 let _ = full_gpu;
2639 let oracle_in = oracle.then(|| (q_raw.clone(), k.clone(), v.clone()));
2640 let mut ao = attention::qwen_attention_core(
2641 q_raw,
2642 k,
2643 v,
2644 &mut self.kv_cache.layers[*li],
2645 &cfg,
2646 );
2647 if let Ok(dir) = std::env::var("CMF_ATTN_DUMP") {
2651 if let Some((qr0, k0, v0)) = oracle_in.clone() {
2652 let (cq, _cg, _ck, _cv) =
2653 attention::finish_projection_debug(qr0, k0, v0, &cfg, position);
2654 let cache = &self.kv_cache.layers[*li];
2655 let n = cache.head_keys(0).len() / hd;
2656 let mut bytes: Vec<u8> = Vec::new();
2657 for v in [nh as u32, nkv as u32, hd as u32, n as u32, position as u32] {
2658 bytes.extend_from_slice(&v.to_le_bytes());
2659 }
2660 for v in &cq {
2661 bytes.extend_from_slice(&v.to_le_bytes());
2662 }
2663 for g in 0..nkv {
2664 for v in cache.head_keys(g) {
2665 bytes.extend_from_slice(&v.to_le_bytes());
2666 }
2667 }
2668 for g in 0..nkv {
2669 for v in cache.head_values(g) {
2670 bytes.extend_from_slice(&v.to_le_bytes());
2671 }
2672 }
2673 let _ =
2674 std::fs::write(format!("{dir}/L{li}_pos{position}.bin"), &bytes);
2675 }
2676 }
2677 if let Some((qr0, k0, v0)) =
2678 oracle_in.filter(|_| std::env::var("CMF_ATTN_ORACLE").as_deref() == Ok("1"))
2679 {
2680 let (cq, _cg, ck, cv) =
2681 attention::finish_projection_debug(qr0, k0, v0, &cfg, position);
2682 let mut h_now = vec![0f32; hs];
2683 graph.read_h(&mut h_now);
2684 let cache = &self.kv_cache.layers[*li];
2685 let n_after = cache.head_keys(0).len() / hd;
2686 let stored = n_after.saturating_sub(1);
2690 let cpu_k: Vec<&[f32]> = (0..nkv)
2691 .map(|g| &cache.head_keys(g)[..stored * hd])
2692 .collect();
2693 let cpu_v: Vec<&[f32]> = (0..nkv)
2694 .map(|g| &cache.head_values(g)[..stored * hd])
2695 .collect();
2696 let p = crate::gpu::AttnDeviceParams {
2697 kv_id,
2698 layer: *li,
2699 nh,
2700 nkv,
2701 hd,
2702 rd,
2703 position,
2704 scale: self.attn_scale,
2705 eps: eps as f32,
2706 gemma,
2707 late_qk_norm: self.qk_norm_after_rope,
2708 output_gate: *output_gate,
2709 q_norm: *q_norm,
2710 k_norm: *k_norm,
2711 inv_freq: &inv_freq,
2712 cpu_k,
2713 cpu_v,
2714 cpu_stored: stored,
2715 o1: None,
2716 };
2717 if let Some((dq, dk, dv, dao)) = graph.debug_attn_device(l, &p, &h_now) {
2718 let md = |a: &[f32], b: &[f32]| {
2719 a.iter().zip(b).fold(0f32, |m, (x, y)| m.max((x - y).abs()))
2720 };
2721 let nn = |a: &[f32]| a.iter().map(|x| x * x).sum::<f32>().sqrt();
2722 eprintln!(
2723 "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}",
2724 nn(&cq),
2725 md(&cq, &dq),
2726 nn(&ck),
2727 md(&ck, &dk),
2728 nn(&cv),
2729 md(&cv, &dv),
2730 nn(&ao),
2731 md(&ao, &dao)
2732 );
2733 } else {
2734 eprintln!("attn-oracle L{li}: device probe declined");
2735 }
2736 }
2737 graph.encode_attn_suffix(l, &ao);
2738 graph.commit();
2741 attention::recycle_buf(&mut ao);
2742 }
2743 }
2744
2745 crate::gpu::stageprof(_xkind, _xt0.elapsed());
2746 }
2747 let mut lm_rows = None;
2752 if self.graph_want_logits
2753 && upto.is_none()
2754 && end == self.num_layers
2755 && std::env::var("CMF_GPU_LMHEAD")
2756 .map(|v| v != "0")
2757 .unwrap_or(true)
2758 {
2759 if let Some(lm) = self.weights.lm_head.metal_graph_parts() {
2760 if graph.lm_head_ok(lm) {
2761 graph.encode_lm_head(&self.weights.final_norm, lm);
2762 lm_rows = Some(lm.1);
2763 }
2764 }
2765 }
2766 if self.graph_head_required && lm_rows.is_none() {
2767 METAL_GRAPH_HEAD_MISS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
2768 self.fail_metal_graph("fused graph head was requested but not encodable");
2769 return start;
2770 }
2771 let _sy0 = std::time::Instant::now();
2772 if let Err(err) = graph.sync_checked() {
2773 self.fail_metal_graph(&err);
2774 return start;
2775 }
2776 let _rs0 = std::time::Instant::now();
2777 if !pending.is_empty() {
2778 let idxs: Vec<usize> = pending.drain(..).flat_map(|(f, n)| f..f + n).collect();
2779 let mut outs: Vec<&mut [f32]> = self
2780 .kv_cache
2781 .layers
2782 .iter_mut()
2783 .enumerate()
2784 .filter(|(i, _)| idxs.binary_search(i).is_ok())
2785 .map(|(_, s)| s.linear_state.as_mut_slice())
2786 .collect();
2787 graph.read_states(&mut outs);
2788 }
2789 if std::env::var("CMF_GRAPH_HOSTPROF").as_deref() == Ok("1") {
2790 use std::sync::atomic::{AtomicU64, Ordering};
2791 static SY: AtomicU64 = AtomicU64::new(0);
2792 static RS: AtomicU64 = AtomicU64::new(0);
2793 static N: AtomicU64 = AtomicU64::new(0);
2794 SY.fetch_add((_rs0 - _sy0).as_nanos() as u64, Ordering::Relaxed);
2795 RS.fetch_add(_rs0.elapsed().as_nanos() as u64, Ordering::Relaxed);
2796 let n = N.fetch_add(1, Ordering::Relaxed) + 1;
2797 if n % 100 == 0 {
2798 eprintln!(
2799 "postprof: sync-wait {:.1} ms/ток | read_states {:.1} ms/ток ({n})",
2800 SY.load(Ordering::Relaxed) as f64 / n as f64 / 1e6,
2801 RS.load(Ordering::Relaxed) as f64 / n as f64 / 1e6
2802 );
2803 }
2804 }
2805 if let Some(rows) = lm_rows {
2806 crate::gpu::hostprof_encode_done(_mt0);
2807 let mut lg = attention::take_buf(rows.min(self.vocab_size));
2808 graph.read_logits(&mut lg);
2809 crate::gpu::hostprof_total(_mt0);
2810 lg.resize(self.vocab_size, 0.0);
2811 if let Some(c) = self.final_softcap {
2812 for l in lg.iter_mut() {
2813 *l = c * (*l / c).tanh();
2814 }
2815 }
2816 self.graph_logits = Some(lg);
2817 }
2818 graph.read_h(h);
2819 if self.graph_head_required && self.graph_logits.is_none() {
2820 METAL_GRAPH_HEAD_MISS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
2821 self.fail_metal_graph("fused graph head completed without logits readback");
2822 return start;
2823 }
2824 METAL_GRAPH_TOK_OK.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
2825 METAL_GRAPH_LAYERS.fetch_add(
2826 end.saturating_sub(start) as u64,
2827 std::sync::atomic::Ordering::Relaxed,
2828 );
2829 if self.graph_head_required {
2830 METAL_GRAPH_HEAD_OK.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
2831 }
2832 for li in dev_attn {
2836 let mut krow = attention::take_buf(nkv * hd);
2837 let mut vrow = attention::take_buf(nkv * hd);
2838 if crate::gpu::kv_mirror_read_last(kv_id, li, nkv, hd, &mut krow, &mut vrow) {
2839 let cache = &mut self.kv_cache.layers[li];
2840 cache.append(&krow, &vrow, &[]);
2841 let n = cache.seq_len;
2842 let mut imp = attention::take_buf(n);
2843 crate::gpu::kv_mirror_take_imp(kv_id, li, &mut imp);
2844 cache.accumulate_imp(&imp);
2845 attention::recycle_buf(&mut imp);
2846 }
2847 attention::recycle_buf(&mut krow);
2848 attention::recycle_buf(&mut vrow);
2849 }
2850 if let Some((_, arm)) = ab {
2851 crate::gpu_metal::dense_ab_record(arm, _mt0.elapsed());
2852 }
2853 end
2854 }
2855
2856 pub fn new(
2857 tokenizer: Tokenizer,
2858 weights: PipelineWeights,
2859 hidden_size: usize,
2860 intermediate_size: usize,
2861 num_heads: usize,
2862 num_kv_heads: usize,
2863 head_dim: usize,
2864 num_layers: usize,
2865 physical_layers: usize,
2866 loop_final_norm: bool,
2867 vocab_size: usize,
2868 rms_eps: f64,
2869 rope_base: f32,
2870 norm_style: NormStyle,
2871 max_seq_len: usize,
2872 sampler_config: SamplerConfig,
2873 ) -> Self {
2874 let rng = match sampler_config.seed {
2875 Some(s) => SplitMix64::new(s),
2876 None => SplitMix64::from_entropy(),
2877 };
2878 let inv_freq = std::sync::Arc::new(attention::rope_inv_freq(head_dim, rope_base));
2879 let pool = Pool::from_env();
2880 if let Some(p) = &pool {
2881 tracing::info!("worker pool: {} threads", p.n_workers());
2882 if let Some(model) = weights
2884 .lm_head
2885 .model_arc()
2886 .or_else(|| weights.embed_tokens.model_arc())
2887 {
2888 let regions: Vec<&[u8]> =
2889 model.tensors.iter().map(|t| model.entry_bytes(t)).collect();
2890 p.bind_numa(®ions);
2891 }
2892 }
2893 Self {
2894 gpu_plan: None,
2895 tokenizer: std::sync::Arc::new(tokenizer),
2896 kv_cache: KvCache::new(num_layers, num_kv_heads, head_dim, max_seq_len),
2897 sampler_config,
2898 weights,
2899 hidden_size,
2900 intermediate_size,
2901 num_heads,
2902 num_kv_heads,
2903 head_dim,
2904 num_layers,
2905 physical_layers,
2906 loop_final_norm,
2907 vocab_size,
2908 rms_eps,
2909 rope_base,
2910 norm_style,
2911 rotary_dim: head_dim,
2912 attention_heads_per_layer: None,
2913 kv_heads_per_layer: None,
2914 v_head_dim: None,
2915 layer_dump: std::env::var_os("CMF_LAYER_DUMP")
2916 .filter(|v| !v.is_empty())
2917 .map(std::path::PathBuf::from),
2918 graph_declines: std::cell::RefCell::new(Vec::new()),
2919 mimo_moe: Default::default(),
2920 vmf_cfg: None,
2921 gdn_cfg: None,
2922 kda_cfg: None,
2923 g3n: None,
2924 dsv4: None,
2925 dsv41: None,
2926 dsv41_vision: None,
2927 dsv41_prefill: None,
2928 qwen4_exp: None,
2929 dsv4_mtp: Vec::new(),
2930 dspark: None,
2931 dspark_pending: Vec::new(),
2932 dspark_hist: Vec::new(),
2933 dspark_real: Vec::new(),
2934 dspark_trunk_picks: Vec::new(),
2935 dspark_exp: Vec::new(),
2936 dspark_draft_ns: 0,
2937 logit_multiplier: None,
2938 cancel: std::sync::Arc::new(std::sync::atomic::AtomicBool::new(false)),
2939 graph_failed: std::sync::atomic::AtomicBool::new(false),
2940 kv_history: Vec::new(),
2941 kv_history_device: false,
2942 short_conv_cfg: None,
2943 mtp: None,
2944 mimo_mtp: None,
2945 verify_exact_moe: false,
2946 speculative: std::env::var("CMF_MTP").map(|v| v != "0").unwrap_or(true),
2947 ignore_eos: false,
2948 draft_full_streak: 0,
2949 spec_k_adapt: None,
2950 spec_acc_ewma: 0.7,
2951 rng,
2952 sampler_scratch: SamplerScratch::default(),
2953 spec_forced: None,
2954 spec_q: Vec::new(),
2955 spec_p: Vec::new(),
2956 spec_res: Vec::new(),
2957 spec_qs: Vec::new(),
2958 spec_ps: Vec::new(),
2959 spec_ress: Vec::new(),
2960 mtp_graph_mode: None,
2961 #[cfg(target_os = "macos")]
2962 metal_verify: None,
2963 inv_freq,
2964 ws: ForwardScratch::new(hidden_size),
2965 pool,
2966 model: None,
2967 dyn_force_f32: false,
2968 dyn_skill_layers: Vec::new(),
2969 dyn_active: None,
2970 dyn_blend_loaded: false,
2971 dyn_phi_layer: None,
2972 dyn_phi_ema: Vec::new(),
2973 dyn_phi_seen: 0,
2974 dyn_router: None,
2975 o1_cfg: None,
2976 o1_epoch: 0,
2977 o1_flags: Vec::new(),
2978 trace: false,
2979 calib_temp: 1.0,
2980 confidence_on: true,
2981 embed_multiplier: 1.0,
2982 attn_scale: 1.0 / (head_dim as f32).sqrt(),
2983 swa: None,
2984 sliding_layers: None,
2985 anchor_core: None,
2986 bounded_rope: None,
2987 kv_prefix: KvPrefix::default(),
2988 last_prefill_tokens: 0,
2989 inv_freq_local: None,
2990 rotary_dim_local: None,
2991 rope_scale: 1.0,
2992 rope_scale_local: 1.0,
2993 global_attn: None,
2994 inv_freq_global: None,
2995 attn_v_norm: false,
2996 qk_norm_after_rope: false,
2997 final_softcap: None,
2998 head_clusters: None,
2999 attn_softcap: 0.0,
3000 graph_want_logits: false,
3001 graph_head_required: false,
3002 graph_logits: None,
3003 embryo_graph: None,
3004 graph_refused: std::sync::atomic::AtomicBool::new(false),
3005 graph_kv_id: {
3006 static NEXT: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(1);
3007 NEXT.fetch_add(1, std::sync::atomic::Ordering::Relaxed)
3008 },
3009 #[cfg(test)]
3010 nll_test_fail_at: None,
3011 #[cfg(test)]
3012 nll_test_force_serial: false,
3013 }
3014 }
3015
3016 pub fn set_o1(&mut self, cfg: Option<crate::nystrom::O1Cfg>) {
3024 if let Err(e) = self.try_set_o1(cfg) {
3025 tracing::error!("{e}");
3026 }
3027 }
3028
3029 pub fn bounded_native(&self) -> bool {
3033 self.anchor_core.is_some()
3034 }
3035
3036 pub fn device_state_bytes(&self) -> Option<(u64, u64)> {
3039 crate::gpu::embryo_device_state_bytes(self.graph_kv_id)
3040 }
3041
3042 pub fn o1_refusal(&self) -> Option<String> {
3044 self.anchor_core.as_ref().map(|ac| {
3045 format!(
3046 "--o1 / CMF_O1 refused: the anchor is native bounded \
3047 (anchor_core kind={} window={} sink={}); the file's operator \
3048 is executed as-is and no post-hoc Nyström overlay applies",
3049 ac.kind, ac.window, ac.sink
3050 )
3051 })
3052 }
3053
3054 pub fn try_set_o1(&mut self, cfg: Option<crate::nystrom::O1Cfg>) -> Result<(), String> {
3056 if let Some(c) = &cfg {
3057 if let Some(why) = self.o1_refusal() {
3058 self.o1_flags = Vec::new();
3059 self.o1_cfg = None;
3060 return Err(why);
3061 }
3062 if crate::nystrom::o1_deferred_boundary(c.w, c.sink).is_none() {
3063 self.o1_flags.clear();
3064 self.o1_cfg = None;
3065 return Err(format!(
3066 "o1 disabled: w + sink + slack + 1 overflows usize (w={}, sink={})",
3067 c.w, c.sink
3068 ));
3069 }
3070 }
3071 self.o1_flags = match &cfg {
3072 Some(c) => {
3073 let mut flags = c.layer_flags(self.num_layers);
3074 for (li, f) in flags.iter_mut().enumerate() {
3075 if *f
3081 && (!matches!(
3082 self.weights.layers[self.phys_layer(li)].attn,
3083 AttnKind::Full { .. }
3084 ) || self.layer_window(li).is_some()
3085 || self.kv_cache.layers[li].sinks.is_some()
3086 || self.layer_v_dim(li) != self.layer_geom(li).1)
3087 {
3088 *f = false;
3089 }
3090 }
3091 flags
3092 }
3093 None => Vec::new(),
3094 };
3095 if let Some(c) = &cfg {
3096 let n = self.o1_flags.iter().filter(|&&f| f).count();
3097 tracing::info!(
3098 "o1 nystrom attention: {n}/{} layer(s), m={} w={} sink={} rect={:?}",
3099 self.num_layers,
3100 c.m,
3101 c.w,
3102 c.sink,
3103 c.rect
3104 );
3105 }
3106 self.o1_cfg = cfg;
3107 Ok(())
3108 }
3109
3110 pub fn install_bounded(
3116 &mut self,
3117 cfg: &cortiq_core::AnchorCoreConfig,
3118 ) -> Result<(), String> {
3119 if !cortiq_core::AnchorCoreConfig::KINDS.contains(&cfg.kind.as_str()) {
3120 return Err(format!(
3121 "anchor_core kind '{}' is not executable by this runtime",
3122 cfg.kind
3123 ));
3124 }
3125 if cfg.window == 0 {
3126 return Err("anchor_core.window must be >= 1".into());
3127 }
3128 let mut n = 0usize;
3129 for li in 0..self.num_layers {
3130 let pl = self.phys_layer(li);
3131 if let AttnKind::Bounded(w) = &self.weights.layers[pl].attn {
3132 if w.window != cfg.window || w.sink != cfg.sink {
3133 return Err(format!(
3134 "layer {li}: bounded weights (window {} sink {}) disagree with \
3135 anchor_core (window {} sink {})",
3136 w.window, w.sink, cfg.window, cfg.sink
3137 ));
3138 }
3139 self.kv_cache.layers[li].install_bounded(cfg.window);
3140 n += 1;
3141 }
3142 }
3143 if n == 0 {
3144 return Err("anchor_core is present but no layer executes it".into());
3145 }
3146 let rope = crate::bounded::BoundedRope::new(cfg.window, &self.inv_freq, self.rope_scale);
3147 self.bounded_rope = Some(std::sync::Arc::new(rope));
3148 self.anchor_core = Some(cfg.clone());
3149 self.embryo_graph = None;
3150 tracing::info!(
3151 "bounded anchor {}: {n} layer(s), window {} sink {} — {} B of ring per layer",
3152 cfg.kind,
3153 cfg.window,
3154 cfg.sink,
3155 self.kv_cache.layers.iter().map(|l| l.bounded_state_bytes()).max().unwrap_or(0)
3156 );
3157 Ok(())
3158 }
3159
3160 pub fn install_wire_identity(&mut self, identity: u64) {
3164 for li in 0..self.kv_cache.layers.len() {
3165 let pl = self.phys_layer(li);
3166 let kind = match self.weights.layers.get(pl).map(|l| &l.attn) {
3167 Some(AttnKind::Bounded(_)) => crate::kv_cache::WireKind::Bounded,
3168 Some(AttnKind::Linear(_))
3169 | Some(AttnKind::LinearGdn(_))
3170 | Some(AttnKind::ShortConv(_))
3171 | Some(AttnKind::Kda(_)) => crate::kv_cache::WireKind::Linear,
3172 _ => crate::kv_cache::WireKind::Full,
3173 };
3174 let l = &mut self.kv_cache.layers[li];
3175 l.wire_kind = kind;
3176 l.wire_identity = identity;
3177 l.wire_layer = li as u32;
3180 }
3181 }
3182
3183 pub fn clear_history(&mut self) {
3185 self.kv_history.clear();
3186 self.kv_history_device = false;
3187 self.kv_prefix.clear();
3188 }
3189
3190 pub fn graph_refused(&self) -> bool {
3194 self.graph_refused
3195 .load(std::sync::atomic::Ordering::Relaxed)
3196 }
3197
3198 pub fn mark_graph_refused(&self) {
3200 if !self
3201 .graph_refused
3202 .swap(true, std::sync::atomic::Ordering::Relaxed)
3203 {
3204 tracing::info!(
3205 "token graph: unsupported for this pipeline (seq {}) — not retrying",
3206 self.graph_kv_id
3207 );
3208 }
3209 }
3210
3211 pub fn device_sequence_position(&self) -> Option<usize> {
3215 crate::gpu::embryo_device_next_position(self.graph_kv_id)
3216 }
3217
3218 fn prefix_owner_matches(&self, n: usize, recorded_on_device: bool) -> bool {
3223 let dev = self.device_sequence_position();
3224 if recorded_on_device {
3225 dev == Some(n) && self.embryo_resident_wanted()
3226 } else {
3227 dev.is_none()
3228 }
3229 }
3230
3231 pub(crate) fn invalidate_for_weight_change(&mut self) {
3239 self.clear_sequence_state();
3240 self.embryo_graph = None;
3241 }
3242
3243 fn cached_prefix_len(&self, input_ids: &[u32]) -> usize {
3249 let (n, on_device) = if self.bounded_native() {
3250 (self.kv_prefix.extension(input_ids), self.kv_prefix.on_device())
3251 } else {
3252 let h = &self.kv_history;
3253 if !h.is_empty() && h.len() < input_ids.len() && input_ids[..h.len()] == h[..] {
3254 (h.len(), self.kv_history_device)
3255 } else {
3256 (0, false)
3257 }
3258 };
3259 if n > 0 && !self.prefix_owner_matches(n, on_device) {
3262 tracing::warn!(
3263 "kv-reuse refused: the cached prefix ({n} positions) was built on the {} path, \
3264 the device now holds {:?} — re-prefilling from zero",
3265 if on_device { "resident device" } else { "host" },
3266 self.device_sequence_position()
3267 );
3268 return 0;
3269 }
3270 n
3271 }
3272
3273 pub fn reusable_prefix_len(&self, input_ids: &[u32]) -> usize {
3276 self.cached_prefix_len(input_ids)
3277 }
3278
3279 fn record_consumed_prefix(&mut self, consumed: &[u32], reused: usize) {
3285 let on_device = self.device_sequence_position().is_some();
3286 if self.bounded_native() {
3287 let keep = reused > 0 && reused == self.kv_prefix.len() && reused <= consumed.len();
3288 let prev_device = self.kv_prefix.on_device();
3289 self.kv_history.clear();
3290 self.kv_history_device = false;
3291 if keep && prev_device == on_device {
3292 self.kv_prefix.extend(&consumed[reused..]);
3293 } else {
3294 self.kv_prefix.set(consumed);
3295 }
3296 self.kv_prefix.set_on_device(on_device);
3297 } else {
3298 self.kv_history = consumed.to_vec();
3299 self.kv_history_device = on_device;
3300 }
3301 }
3302
3303 pub fn o1_active(&self) -> bool {
3305 self.o1_cfg.is_some() && self.o1_flags.iter().any(|&f| f)
3306 }
3307
3308 pub fn generation_batch_k(&self) -> usize {
3320 if let Some(k) = std::env::var("CMF_BATCH_K")
3321 .ok()
3322 .and_then(|v| v.parse::<usize>().ok())
3323 {
3324 return k;
3325 }
3326 #[cfg(not(target_os = "macos"))]
3327 if self.graph_prefill_preferred() && !self.o1_active() {
3328 return 32;
3329 }
3330 0
3331 }
3332
3333 pub fn generation_graph_prefill(&self) -> bool {
3334 let graph = self.graph_prefill_preferred();
3335 #[cfg(not(target_os = "macos"))]
3346 if graph
3347 && self.generation_batch_k() > 0
3348 && std::env::var("CMF_MTP_CHAIN_PROBE").is_err()
3349 {
3350 return false;
3351 }
3352 graph
3353 }
3354
3355 pub fn o1_device_stats(&self) -> (usize, u64) {
3360 crate::gpu::o1_device_stats(self.graph_kv_id)
3361 }
3362
3363 pub fn o1_begin(&mut self) {
3368 self.o1_begin_with_prefix(None);
3369 }
3370
3371 pub fn o1_begin_with_prefix(&mut self, requested_prefix: Option<usize>) {
3375 if let Some(c) = &self.o1_cfg {
3376 let (m, w, sink, rect) = (c.m, c.w, c.sink, c.rect);
3377 let boundary = requested_prefix.map(|p| {
3378 p.max(
3379 crate::nystrom::o1_deferred_boundary(w, sink)
3380 .expect("o1 config boundary validated in set_o1"),
3381 )
3382 });
3383 for (li, &f) in self.o1_flags.iter().enumerate() {
3384 if f {
3385 self.kv_cache.layers[li].o1_begin_with_boundary(m, w, sink, rect, boundary);
3386 }
3387 }
3388 }
3389 }
3390
3391 fn o1_effective_boundary(&self, requested_prefix: usize) -> Option<usize> {
3393 self.o1_cfg.as_ref().and_then(|c| {
3394 crate::nystrom::o1_deferred_boundary(c.w, c.sink)
3395 .map(|floor| requested_prefix.max(floor))
3396 })
3397 }
3398
3399 fn o1_note_transition(&mut self) {
3400 let mut transitioned = false;
3404 for (li, &flagged) in self.o1_flags.iter().enumerate() {
3405 if flagged {
3406 transitioned |= self.kv_cache.layers[li].take_o1_transition();
3407 }
3408 }
3409 if transitioned {
3410 self.o1_epoch = self.o1_epoch.wrapping_add(1);
3411 }
3412 }
3413
3414 fn o1_pending(&self) -> bool {
3415 self.o1_flags.iter().enumerate().any(|(li, &f)| {
3416 f && self.kv_cache.layers[li].seq_len > 0
3417 && self.kv_cache.layers[li].o1_pending_boundary().is_some()
3418 })
3419 }
3420
3421 fn o1_fail(&mut self, err: String) {
3422 tracing::error!("o1 deferred seal failed; terminating sequence: {err}");
3423 self.clear_sequence_state();
3424 self.graph_failed
3425 .store(true, std::sync::atomic::Ordering::Relaxed);
3426 self.cancel
3427 .store(true, std::sync::atomic::Ordering::Relaxed);
3428 }
3429
3430 pub fn o1_seal_checked(&mut self) -> Result<bool, String> {
3435 if self.o1_cfg.is_none() {
3436 return Ok(false);
3437 }
3438 let mut participating = false;
3439 for li in 0..self.num_layers {
3440 if !self.o1_flags.get(li).copied().unwrap_or(false) {
3441 continue;
3442 }
3443 if let Some(err) = self.kv_cache.layers[li].take_o1_error() {
3444 return Err(err);
3445 }
3446 if self.kv_cache.layers[li].seq_len == 0 {
3447 continue;
3448 }
3449 participating = true;
3450 let num_heads = self.layer_num_heads(li);
3451 self.kv_cache.layers[li].o1_seal_checked(num_heads)?;
3452 }
3453 self.o1_note_transition();
3454 for li in 0..self.num_layers {
3455 if self.o1_flags.get(li).copied().unwrap_or(false) {
3456 if let Some(err) = self.kv_cache.layers[li].take_o1_error() {
3457 return Err(err);
3458 }
3459 }
3460 }
3461 Ok(participating
3462 && (0..self.num_layers).all(|li| {
3463 !self.o1_flags.get(li).copied().unwrap_or(false)
3464 || self.kv_cache.layers[li].seq_len == 0
3465 || self.kv_cache.layers[li].o1_sealed()
3466 }))
3467 }
3468
3469 fn o1_progress(&mut self) {
3472 if !self.o1_active() {
3473 return;
3474 }
3475 for li in 0..self.num_layers {
3476 if self.o1_flags.get(li).copied().unwrap_or(false) {
3477 if let Some(err) = self.kv_cache.layers[li].take_o1_error() {
3478 self.o1_fail(err);
3479 return;
3480 }
3481 }
3482 }
3483 self.o1_note_transition();
3487 if !self.o1_pending() {
3488 return;
3489 }
3490 if let Err(err) = self.o1_seal_checked() {
3491 self.o1_fail(err);
3492 }
3493 }
3494
3495 fn check_o1_progress_failure(&mut self, phase: &str) -> Result<(), String> {
3500 if self
3501 .graph_failed
3502 .swap(false, std::sync::atomic::Ordering::Relaxed)
3503 {
3504 self.cancel
3505 .store(false, std::sync::atomic::Ordering::Relaxed);
3506 self.clear_sequence_state();
3507 return Err(format!("{phase}: deferred O(1) transition failed"));
3508 }
3509 Ok(())
3510 }
3511
3512 pub fn o1_seal(&mut self) {
3516 if let Err(err) = self.o1_seal_checked() {
3517 self.o1_fail(err);
3518 }
3519 }
3520
3521 pub fn set_trace(&mut self, on: bool) {
3523 self.trace = on;
3524 }
3525
3526 pub fn set_sampler_config(&mut self, config: SamplerConfig) {
3529 self.rng = match config.seed {
3530 Some(seed) => SplitMix64::new(seed),
3531 None => SplitMix64::from_entropy(),
3532 };
3533 self.sampler_config = config;
3534 }
3535
3536 pub fn set_confidence(&mut self, on: bool) {
3541 self.confidence_on = on;
3542 }
3543
3544 pub fn set_calib_temp(&mut self, t: f32) {
3547 self.calib_temp = if t > 1e-3 { t } else { 1.0 };
3548 }
3549
3550 pub fn calib_temp(&self) -> f32 {
3552 self.calib_temp
3553 }
3554
3555 pub fn set_rotary(&mut self, rotary_dim: usize, base: f32) {
3558 self.rotary_dim = rotary_dim.min(self.head_dim);
3559 self.inv_freq = std::sync::Arc::new(attention::rope_inv_freq(self.rotary_dim, base));
3560 self.embryo_graph = None;
3564 }
3565
3566 fn attn_cfg(&self, position: usize) -> QwenAttnCfg<'_> {
3567 QwenAttnCfg {
3568 num_heads: self.num_heads,
3569 num_kv_heads: self.num_kv_heads,
3570 head_dim: self.head_dim,
3571 hidden_size: self.hidden_size,
3572 position,
3573 inv_freq: &self.inv_freq,
3574 rotary_dim: self.rotary_dim,
3575 scale: self.attn_scale,
3576 softcap: self.attn_softcap,
3577 window: None,
3578 v_norm: false,
3579 qk_norm_after_rope: self.qk_norm_after_rope,
3580 q_norm: None,
3581 k_norm: None,
3582 output_gate: false,
3583 softplus_gate: None,
3584 rope_scale: self.rope_scale,
3585 bias: None,
3586 rms_eps: self.rms_eps,
3587 norm_style: self.norm_style,
3588 pool: self.pool.as_deref(),
3589 v_head_dim: self.v_head_dim.unwrap_or(self.head_dim),
3590 }
3591 }
3592
3593 pub fn generate(
3595 &mut self,
3596 prompt: &str,
3597 max_tokens: usize,
3598 task_mask: Option<&TaskMask>,
3599 on_token: Option<TokenCallback>,
3600 ) -> Result<GenerateResult, String> {
3601 let input_ids = self.tokenizer.with_bos(self.tokenizer.encode(prompt));
3602 self.generate_from_ids(&input_ids, max_tokens, task_mask, on_token)
3603 }
3604
3605 pub fn generate_from_vl(
3608 &mut self,
3609 input: &crate::dsv41_vision::PreparedVlInputs,
3610 max_tokens: usize,
3611 task_mask: Option<&TaskMask>,
3612 on_token: Option<TokenCallback>,
3613 ) -> Result<GenerateResult, String> {
3614 let Some(dsv41) = &self.dsv41 else {
3615 return Err("V4.1 multimodal input requires a DeepSeek-V4.1 pipeline".into());
3616 };
3617 if input.token_ids.is_empty() {
3618 return Err("empty V4.1 multimodal prompt".into());
3619 }
3620 if input.token_types.len() != input.token_ids.len() {
3621 return Err(format!(
3622 "V4.1 token type count {} != token count {}",
3623 input.token_types.len(),
3624 input.token_ids.len()
3625 ));
3626 }
3627 let dim = dsv41.2.dim;
3628 let mut embeddings = vec![None; input.token_ids.len()];
3629 let mut participates = vec![true; input.token_ids.len()];
3630 if !input.images.is_empty() {
3631 let vision = self
3632 .dsv41_vision
3633 .as_ref()
3634 .ok_or_else(|| "V4.1 image prompt has no loaded vision tower".to_string())?;
3635 for image in &input.images {
3636 let end = image.start.saturating_add(image.types.len());
3637 if end > input.token_ids.len() {
3638 return Err(format!(
3639 "V4.1 image span {}..{} exceeds prompt length {}",
3640 image.start,
3641 end,
3642 input.token_ids.len()
3643 ));
3644 }
3645 let mut span = vec![0.0f32; image.types.len() * dim];
3646 vision.fill_image_span(image, &mut span, self.pool.as_deref())?;
3647 for (offset, &kind) in image.types.iter().enumerate() {
3648 let pos = image.start + offset;
3649 if input.token_types[pos] != kind {
3650 return Err(format!(
3651 "V4.1 image type mismatch at position {pos}: {} != {kind}",
3652 input.token_types[pos]
3653 ));
3654 }
3655 embeddings[pos] = Some(span[offset * dim..(offset + 1) * dim].to_vec());
3656 participates[pos] = false;
3657 }
3658 }
3659 }
3660 for (pos, &kind) in input.token_types.iter().enumerate() {
3661 if kind == crate::dsv41_vision::TEXT && embeddings[pos].is_some() {
3662 return Err(format!("V4.1 text position {pos} has an image embedding"));
3663 }
3664 if kind != crate::dsv41_vision::TEXT && embeddings[pos].is_none() {
3665 return Err(format!("V4.1 image position {pos} has no image embedding"));
3666 }
3667 }
3668 self.dsv41_prefill = Some((embeddings, participates));
3669 let result = self.generate_from_ids(&input.token_ids, max_tokens, task_mask, on_token);
3670 self.dsv41_prefill = None;
3671 result
3672 }
3673
3674 fn drop_open_mask<'m>(&self, m: Option<&'m TaskMask>) -> Option<&'m TaskMask> {
3676 m.filter(|m| !m.fully_open(self.intermediate_size, self.num_heads))
3677 }
3678
3679 pub fn generate_from_ids(
3687 &mut self,
3688 input_ids: &[u32],
3689 max_tokens: usize,
3690 task_mask: Option<&TaskMask>,
3691 on_token: Option<TokenCallback>,
3692 ) -> Result<GenerateResult, String> {
3693 self.generate_with_prompt_rows(input_ids, None, max_tokens, task_mask, on_token)
3694 }
3695
3696 pub fn generate_from_embeds(
3702 &mut self,
3703 input_ids: &[u32],
3704 prompt_rows: &[f32],
3705 max_tokens: usize,
3706 task_mask: Option<&TaskMask>,
3707 on_token: Option<TokenCallback>,
3708 ) -> Result<GenerateResult, String> {
3709 if input_ids.is_empty()
3710 || input_ids.len().checked_mul(self.hidden_size) != Some(prompt_rows.len())
3711 {
3712 return Err("embedded prompt dimensions must be [tokens, hidden_size]".into());
3713 }
3714 if prompt_rows.iter().any(|x| !x.is_finite()) {
3715 return Err("embedded prompt contains non-finite values".into());
3716 }
3717 if !self.can_prefill_batched() || self.dyn_router.is_some()
3718 || self.o1_active() || self.mtp.is_some() || self.gpu_plan.is_some()
3719 {
3720 return Err("embedded prompts require the ordinary transformer path without O(1), dynamic routing, GPU splitting or a generic MTP head".into());
3721 }
3722 self.generate_with_prompt_rows(input_ids, Some(prompt_rows), max_tokens, task_mask, on_token)
3723 }
3724
3725 fn generate_with_prompt_rows(
3726 &mut self,
3727 input_ids: &[u32],
3728 prompt_rows: Option<&[f32]>,
3729 max_tokens: usize,
3730 task_mask: Option<&TaskMask>,
3731 mut on_token: Option<TokenCallback>,
3732 ) -> Result<GenerateResult, String> {
3733 #[cfg(target_os = "macos")]
3734 crate::gpu_metal::set_io_namespace(self.graph_kv_id);
3735 if std::env::var("CMF_TRACE_H").is_ok() {
3736 eprintln!("input_ids: {input_ids:?}");
3737 }
3738 if input_ids.is_empty() {
3739 return Err("empty prompt: nothing to generate from".to_string());
3740 }
3741 self.graph_failed
3745 .store(false, std::sync::atomic::Ordering::Relaxed);
3746 let task_mask = self.drop_open_mask(task_mask);
3751
3752 let mut reuse_from = {
3760 let on = !std::env::var("CMF_KV_REUSE").is_ok_and(|v| v == "0");
3761 if on
3762 && prompt_rows.is_none()
3763 && task_mask.is_none()
3764 && self.mtp.is_none()
3765 && !(self.mimo_mtp.is_some() && self.speculative)
3766 && self.o1_cfg.is_none()
3767 && self.dsv41.is_none()
3768 {
3769 self.cached_prefix_len(input_ids)
3770 } else {
3771 0
3772 }
3773 };
3774 if reuse_from > 0 && !self.prepare_kv_reuse(reuse_from) {
3777 reuse_from = 0;
3778 }
3779 self.last_prefill_tokens = input_ids.len() - reuse_from;
3780 let bounded_native = self.bounded_native();
3781 if reuse_from == 0 {
3782 self.clear_sequence_state();
3784 } else if std::env::var("CMF_PREFILL_PROF").is_ok() {
3785 eprintln!(
3786 "kv-reuse: {} of {} prompt positions already cached",
3787 reuse_from,
3788 input_ids.len()
3789 );
3790 }
3791 crate::gpu::graph_race_begin_generation();
3792 let o1_prefill = if self.o1_active() && task_mask.is_none() {
3796 std::env::var("CMF_O1_PREFILL")
3797 .ok()
3798 .and_then(|v| v.parse::<usize>().ok())
3799 .filter(|&p| p > 0)
3800 } else {
3801 None
3802 };
3803 if task_mask.is_none() {
3804 self.o1_begin_with_prefix(o1_prefill);
3805 }
3806
3807 let graph_on = crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Decode);
3813 #[cfg(target_os = "macos")]
3852 let metal_graph = crate::gpu::q1_force()
3853 && crate::gpu::enabled_here()
3854 && std::env::var("CMF_GPU_BLOCK")
3855 .map(|v| v != "0")
3856 .unwrap_or(true);
3857 #[cfg(not(target_os = "macos"))]
3858 let metal_graph = false;
3859 let spec_sample_env = std::env::var("CMF_GRAPH_SPEC_SAMPLE").ok();
3860 let spec_cheap_round = self.sampler_config.temperature < 1e-6
3864 || sampler::sparse_ok(&self.sampler_config);
3865 let spec_sampling_ok = self.sampler_config.temperature < 1e-6
3866 || match spec_sample_env.as_deref() {
3867 Some("1") => true,
3868 Some(_) => false,
3869 None => metal_graph && spec_cheap_round,
3870 };
3871 let (mut dense_n, mut dense_q4tp) = (0usize, 0usize);
3887 for lw in &self.weights.layers {
3888 if let FfnKind::Dense(d) = &lw.ffn {
3889 dense_n += 1;
3890 if matches!(d.gate_proj.graph_weight(), Some((_, _, 6, _)))
3891 && matches!(d.up_proj.graph_weight(), Some((_, _, 6, _)))
3892 && matches!(d.down_proj.graph_weight(), Some((_, _, 6, _)))
3893 {
3894 dense_q4tp += 1;
3895 }
3896 }
3897 }
3898 let spec_default_ok = dense_n == 0 || dense_q4tp * 10 >= dense_n * 9;
3899 let penalized = !metal_graph
3911 && (self.sampler_config.repetition_penalty != 1.0
3912 || self.sampler_config.presence_penalty != 0.0
3913 || !self.sampler_config.suppress_tokens.is_empty());
3914 #[cfg(feature = "gpu")]
3919 let metal_wgpu = graph_on && crate::gpu_wgpu::wgpu_backend_is_metal();
3920 #[cfg(not(feature = "gpu"))]
3921 let metal_wgpu = false;
3922 let spec_env = std::env::var("CMF_GRAPH_SPEC").ok();
3923 let spec_wanted = match spec_env.as_deref() {
3924 Some("0") => false,
3925 Some(_) => {
3926 if metal_wgpu {
3927 tracing::warn!(
3928 "CMF_GRAPH_SPEC forced on wgpu/Metal: the batched verify graph is not \
3929 verified on this backend (garbage measured on Qwen3.5-0.8B)"
3930 );
3931 }
3932 true
3933 }
3934 None => spec_default_ok && !penalized && !metal_wgpu,
3935 };
3936 let graph_spec = self.speculative
3940 && (graph_on || metal_graph)
3941 && self.mtp.is_some()
3942 && task_mask.is_none()
3943 && !self.o1_active()
3944 && spec_sampling_ok
3945 && spec_wanted;
3946 #[cfg(target_os = "macos")]
3950 if metal_graph {
3951 static SAID: std::sync::Once = std::sync::Once::new();
3952 SAID.call_once(|| {
3953 let spec = if graph_spec {
3954 let k = std::env::var("CMF_GRAPH_SPEC_K")
3955 .ok()
3956 .and_then(|v| v.parse::<usize>().ok())
3957 .filter(|&v| (1..=8).contains(&v))
3958 .unwrap_or(7);
3959 let arm = if self.sampler_config.temperature < 1e-6 {
3960 "greedy"
3961 } else {
3962 "sampling"
3963 };
3964 format!(
3965 "spec k={k} {arm} (batched verify, draft shortlist {}, trial: proxy)",
3966 Self::draft_vocab_rows(usize::MAX)
3967 )
3968 } else if !self.speculative {
3969 "spec off (CMF_MTP=0)".to_string()
3970 } else if self.mtp.is_none() {
3971 "spec off (no MTP head)".to_string()
3972 } else if !spec_sampling_ok {
3973 if spec_cheap_round {
3974 "spec off (CMF_GRAPH_SPEC_SAMPLE=0)".to_string()
3975 } else {
3976 "spec off (sampling without a top-k: the dense chain \
3977 costs more than it saves)"
3978 .to_string()
3979 }
3980 } else if !spec_wanted {
3981 "spec off (CMF_GRAPH_SPEC=0 or non-q4tp FFNs)".to_string()
3982 } else if task_mask.is_some() {
3983 "spec off (task mask)".to_string()
3984 } else {
3985 "spec off (O(1) attention)".to_string()
3986 };
3987 let on = |var: &str| {
3988 if std::env::var(var).as_deref() == Ok("0") {
3989 "off"
3990 } else {
3991 "on"
3992 }
3993 };
3994 tracing::info!(
3995 "metal native: {spec}, state4 {}, async replay {}, prefill graph {}, \
3996 MTP graph {}, attend {}, probe {}",
3997 if crate::gpu_metal::state4_on() { "on" } else { "off" },
3998 if crate::gpu_metal::async_replay_on() { "on" } else { "off" },
3999 on("CMF_METAL_PREFILL"),
4000 on("CMF_MTP_GRAPH"),
4001 std::env::var("CMF_GPU_ATTEND").unwrap_or_else(|_| "auto".into()),
4002 if crate::gpu::probe_enabled() { "bypassed (q1 force)" } else { "off" },
4003 );
4004 });
4005 }
4006 let pair_pays = self.gdn_cfg.is_none() || std::env::var("CMF_MTP").as_deref() == Ok("1");
4013 let spec_active = self.speculative
4014 && self.mtp.is_some()
4015 && task_mask.is_none()
4016 && !self.o1_active()
4017 && ((!graph_on && pair_pays && self.sampler_config.temperature < 1e-6) || graph_spec);
4018 let mut mtp = if spec_active { self.mtp.take() } else { None };
4021 if std::env::var("CMF_MTP_CHAIN_PROBE").is_ok() {
4022 eprintln!(
4023 "mtp-probe gate: spec_active={spec_active} mtp={} speculative={} graph_on={graph_on} temp_ok={}",
4024 mtp.is_some(),
4025 self.speculative,
4026 self.sampler_config.temperature < 1e-6,
4027 );
4028 }
4029 if let Some(m) = &mut mtp {
4030 m.kv.clear();
4031 crate::gpu::graph_kv_reset(self.mtp_kv_id());
4033 self.mtp_graph_mode = None;
4034 }
4035 let mimo_spec = self.speculative
4039 && self.mimo_mtp.is_some()
4040 && task_mask.is_none()
4041 && !self.o1_active()
4042 && self.dyn_router.is_none()
4043 && self.sampler_config.temperature < 1e-6
4044 && std::env::var("CMF_MIMO_MTP").as_deref() != Ok("0");
4045 if let Some(st) = self.mimo_mtp.as_mut() {
4046 st.reset();
4047 if mimo_spec && std::env::var_os("CMF_MIMO_MTP_PROBE").is_some() {
4048 Self::mimo_mtp_hist_cap(st, input_ids.len());
4049 }
4050 }
4051 let mut router = if mtp.is_none() {
4055 self.dyn_router.take()
4056 } else {
4057 None
4058 };
4059 let mut reuse_from = reuse_from;
4060 if let Some(r) = &mut router {
4061 r.reset(); self.dyn_phi_seen = 0; if self.dyn_active.is_some() {
4064 let _ = self.set_active_skill(None);
4067 reuse_from = 0;
4068 self.last_prefill_tokens = input_ids.len();
4069 }
4070 }
4071
4072 let mut all_ids = input_ids.to_vec();
4073 let mut generated = 0usize;
4074 let mut finish_reason = "max_tokens".to_string();
4075 let mut drafted = 0usize;
4076 let mut accepted = 0usize;
4077 let mut dsv4_spec_bad = 0usize;
4084 let mut dsv4_spec_retry_at = 0usize;
4085 let mut confidence: Vec<f32> = Vec::new();
4086 let trace_on = self.trace;
4087 let calib_temp = self.calib_temp;
4088 let mut traces: Vec<TokenTrace> = Vec::new();
4089
4090 let mut hidden = vec![0.0f32; self.hidden_size];
4096 let mut pos = reuse_from;
4097 let fuse_lm = mtp.is_none()
4106 && router.is_none()
4107 && std::env::var("CMF_GPU_LMHEAD").as_deref() != Ok("0");
4108 self.graph_logits = None;
4109 self.graph_want_logits = false;
4110 let _tpf = std::time::Instant::now();
4111 let batch_k = self.generation_batch_k();
4112 if let Some(rows) = prompt_rows {
4113 let hs = self.hidden_size;
4114 let chunk = self.prefill_chunk().max(1);
4115 while pos < input_ids.len() && !self.cancel.load(std::sync::atomic::Ordering::Relaxed) {
4116 let end = (pos + chunk).min(input_ids.len());
4117 let hb = match self.prefill_input_rows(
4118 PrefillIn::Hidden(&rows[pos * hs..end * hs]), pos, task_mask,
4119 ) {
4120 Ok(hb) => hb,
4121 Err(err) => {
4122 self.finish_generation(&mut mtp, &mut router, true);
4123 return Err(err);
4124 }
4125 };
4126 if mimo_spec { self.mimo_note_rows(&hb, pos); }
4127 hidden.copy_from_slice(&hb[hb.len() - hs..]);
4128 pos = end;
4129 }
4130 }
4131 while self.qwen4_exp.is_some()
4142 && mtp.is_none()
4143 && pos < input_ids.len()
4144 && !self.cancel.load(std::sync::atomic::Ordering::Relaxed)
4145 {
4146 let end = (pos + crate::qwen4_exp::prefill_chunk()).min(input_ids.len());
4149 let want_logits = end == input_ids.len();
4150 let mut lg = Vec::new();
4151 if let Some(b) = &mut self.qwen4_exp {
4152 crate::qwen4_exp::forward_tokens(
4153 &b.0,
4154 &b.1,
4155 &b.2,
4156 &mut b.3,
4157 &input_ids[pos..end],
4158 pos,
4159 &self.inv_freq,
4160 self.pool.as_deref(),
4161 &mut lg,
4162 want_logits,
4163 );
4164 }
4165 if want_logits {
4166 self.graph_logits = Some(lg);
4167 }
4168 pos = end;
4169 hidden.fill(0.0);
4170 }
4171 while self.dsv4.is_some()
4172 && mtp.is_none()
4173 && pos < input_ids.len()
4174 && !self.cancel.load(std::sync::atomic::Ordering::Relaxed)
4175 {
4176 let end = (pos + prefill_chunk()).min(input_ids.len());
4177 let ids: Vec<u32> = input_ids[pos..end].to_vec();
4178 let mut lg = Vec::new();
4179 if let Some(b) = &mut self.dsv4 {
4180 let (g, layers, cfg, st) = (&b.0, &b.1, b.2, &mut b.3);
4181 crate::dsv4::forward_chunk(
4182 g,
4183 layers,
4184 &cfg,
4185 st,
4186 &ids,
4187 pos,
4188 &self.inv_freq,
4189 self.pool.as_deref(),
4190 &mut lg,
4191 end == input_ids.len(),
4192 );
4193 }
4194 if end == input_ids.len() {
4195 self.graph_logits = Some(lg);
4196 }
4197 pos = end;
4198 hidden = vec![0.0; self.hidden_size];
4199 }
4200 let dsv41_prefill = self.dsv41_prefill.take();
4201 while self.dsv41.is_some()
4202 && mtp.is_none()
4203 && pos < input_ids.len()
4204 && !self.cancel.load(std::sync::atomic::Ordering::Relaxed)
4205 {
4206 let end = (pos + prefill_chunk()).min(input_ids.len());
4207 let ids: Vec<u32> = input_ids[pos..end].to_vec();
4208 let mut lg = Vec::new();
4209 if let Some(b) = &mut self.dsv41 {
4210 let (g, layers, cfg, st) = (&b.0, &b.1, &b.2, &mut b.3);
4211 if let Some((embeddings, participates)) = dsv41_prefill.as_ref() {
4212 crate::dsv41::forward_chunk_masked_with_embeddings(
4213 g,
4214 layers,
4215 cfg,
4216 st,
4217 &ids,
4218 pos,
4219 &embeddings[pos..end],
4220 &participates[pos..end],
4221 self.pool.as_deref(),
4222 &mut lg,
4223 );
4224 } else {
4225 crate::dsv41::forward_chunk(
4226 g,
4227 layers,
4228 cfg,
4229 st,
4230 &ids,
4231 pos,
4232 self.pool.as_deref(),
4233 &mut lg,
4234 );
4235 }
4236 }
4237 if end == input_ids.len() {
4238 self.graph_logits = Some(lg);
4239 }
4240 pos = end;
4241 hidden = vec![0.0; self.hidden_size];
4242 }
4243 let dyn_prefill = router.is_some();
4248 let o1_prefill_limit = o1_prefill
4256 .and_then(|requested| self.o1_effective_boundary(requested))
4257 .map(|boundary| boundary.min(input_ids.len()));
4258 let mut o1_sealed = false;
4259 if let Some(limit) = o1_prefill_limit {
4260 if self.can_prefill_batched() && limit > 2 {
4263 let chunk = self.prefill_chunk();
4264 let hs = self.hidden_size;
4265 while pos < limit && !self.cancel.load(std::sync::atomic::Ordering::Relaxed) {
4266 let end = (pos + chunk).min(limit);
4267 let hb = self.prefill_batch(&input_ids[pos..end], pos);
4268 hidden.copy_from_slice(&hb[(end - pos - 1) * hs..]);
4269 pos = end;
4270 }
4271 } else {
4272 while pos < limit && !self.cancel.load(std::sync::atomic::Ordering::Relaxed) {
4273 hidden = self.forward_layers(&self.embed_single(input_ids[pos]), pos, None);
4274 pos += 1;
4275 }
4276 }
4277 if pos >= limit {
4278 o1_sealed = match self.o1_seal_checked() {
4279 Ok(sealed) => sealed,
4280 Err(err) => {
4281 self.finish_generation(&mut mtp, &mut router, true);
4282 return Err(err);
4283 }
4284 };
4285 tracing::info!(
4286 "o1 bounded prompt prefix: requested={} effective={} processed={} of {} token(s)",
4287 o1_prefill.unwrap_or(0),
4288 self.o1_effective_boundary(o1_prefill.unwrap_or(0))
4289 .unwrap_or(limit),
4290 limit,
4291 input_ids.len()
4292 );
4293 }
4294 }
4295 let graph_prefill = self.graph_prefill_preferred();
4301 #[cfg(target_os = "macos")]
4309 if task_mask.is_none()
4310 && !dyn_prefill
4311 && (crate::gpu::q1_force() || crate::gpu::q2tp_gpu_opt_in())
4312 && crate::gpu::enabled_here()
4313 && self.gdn_cfg.is_some()
4314 && self.g3n.is_none()
4315 && input_ids.len() > 8
4316 && std::env::var("CMF_MTP_CHAIN_PROBE").is_err()
4317 && std::env::var("CMF_METAL_PREFILL").as_deref() != Ok("0")
4318 {
4319 let chunk: usize = std::env::var("CMF_METAL_PREFILL_CHUNK")
4320 .ok()
4321 .and_then(|v| v.parse().ok())
4322 .filter(|&v| (16..=512).contains(&v))
4323 .unwrap_or(256);
4324 let hs = self.hidden_size;
4325 let _tp = std::time::Instant::now();
4326 while pos < input_ids.len() && !self.cancel.load(std::sync::atomic::Ordering::Relaxed) {
4327 let end = (pos + chunk).min(input_ids.len());
4328 let hb = match self.prefill_batch_metal(&input_ids[pos..end], pos) {
4329 MetalPrefillOutcome::Completed(hb) => hb,
4330 MetalPrefillOutcome::Declined => break,
4331 MetalPrefillOutcome::Failed => {
4332 self.finish_generation(&mut mtp, &mut router, true);
4333 return Err("ordinary Metal prefill failed after admission".into());
4334 }
4335 };
4336 if let Some(m) = &mut mtp {
4337 let n_pairs = if end < input_ids.len() {
4338 end - pos
4339 } else {
4340 end - pos - 1
4341 };
4342 if n_pairs > 0 {
4343 let pairs: Vec<(&[f32], u32)> = (0..n_pairs)
4344 .map(|j| (&hb[j * hs..(j + 1) * hs], input_ids[pos + j + 1]))
4345 .collect();
4346 if !self.mtp_warm_batch_metal(m, &pairs, pos) {
4347 for (j, (h, t)) in pairs.iter().enumerate() {
4348 let h = h.to_vec();
4349 let _ = self.mtp_step(m, &h, *t, pos + j);
4350 }
4351 }
4352 }
4353 }
4354 hidden.copy_from_slice(&hb[(end - pos - 1) * hs..]);
4355 pos = end;
4356 }
4357 if std::env::var("CMF_PREFILL_PROF").is_ok() {
4358 eprintln!(
4359 "metal-prefill: {} of {} tokens in {:.1} ms",
4360 pos,
4361 input_ids.len(),
4362 _tp.elapsed().as_secs_f64() * 1e3
4363 );
4364 }
4365 }
4366 self.mimo_moe_prepare();
4367 #[cfg(not(target_os = "macos"))]
4373 if task_mask.is_none()
4374 && !dyn_prefill
4375 && !graph_prefill
4376 && mtp.is_none()
4377 && o1_prefill.is_none()
4378 && !self.o1_active()
4379 && input_ids.len() > 2
4380 && self.batch_prefix_prefill()
4381 {
4382 let chunk = self.prefill_chunk().max(1);
4383 let hs = self.hidden_size;
4384 let t_bp = std::time::Instant::now();
4385 let pos0 = pos;
4386 while pos < input_ids.len() && !self.cancel.load(std::sync::atomic::Ordering::Relaxed) {
4387 let end = (pos + chunk).min(input_ids.len());
4388 let bk = end - pos;
4389 let mut hiddens = vec![0f32; bk * hs];
4390 for (j, &id) in input_ids[pos..end].iter().enumerate() {
4391 hiddens[j * hs..(j + 1) * hs].copy_from_slice(&self.embed_single(id));
4392 }
4393 let positions: Vec<usize> = (pos..end).collect();
4394 let mut run = 0usize;
4395 let outcome = self.try_batch_graph_wgpu_prefix(
4396 &mut hiddens,
4397 &positions,
4398 bk,
4399 None,
4400 Some(&mut run),
4401 );
4402 match outcome {
4403 crate::gpu::BatchGraphOutcome::Completed => {
4404 let hb = if run < self.num_layers {
4405 self.prefill_batch_span(
4406 PrefillIn::Hidden(&hiddens),
4407 pos,
4408 None,
4409 run,
4410 self.num_layers,
4411 )
4412 } else {
4413 hiddens
4414 };
4415 if mimo_spec {
4416 self.mimo_note_rows(&hb, pos);
4417 }
4418 hidden.copy_from_slice(&hb[(bk - 1) * hs..]);
4419 pos = end;
4420 }
4421 crate::gpu::BatchGraphOutcome::Failed => {
4422 self.finish_generation(&mut mtp, &mut router, true);
4423 return Err("batched prefix prefill failed after admission".into());
4424 }
4425 crate::gpu::BatchGraphOutcome::Declined => {
4426 #[cfg(feature = "gpu")]
4429 if pos > pos0 {
4430 self.pull_lagging_host_kv(0, self.num_layers, pos);
4431 }
4432 break;
4433 }
4434 }
4435 }
4436 if std::env::var("CMF_PREFILL_PROF").is_ok() {
4437 eprintln!(
4438 "batch-prefix prefill: {} of {} tokens in {:.1} ms",
4439 pos - pos0,
4440 input_ids.len(),
4441 t_bp.elapsed().as_secs_f64() * 1e3
4442 );
4443 }
4444 }
4445 if task_mask.is_none()
4446 && !dyn_prefill
4447 && !graph_prefill
4448 && self.can_prefill_batched()
4449 && self.g3n.is_none()
4450 && o1_prefill.is_none()
4451 && input_ids.len() > 2
4452 {
4453 let chunk = self.prefill_chunk();
4459 let hs = self.hidden_size;
4460 while pos < input_ids.len() && !self.cancel.load(std::sync::atomic::Ordering::Relaxed) {
4461 let end = (pos + chunk).min(input_ids.len());
4462 let hb = self.prefill_batch(&input_ids[pos..end], pos);
4463 if mimo_spec {
4464 self.mimo_note_rows(&hb, pos);
4465 }
4466 if let Some(m) = &mut mtp {
4467 let probe: usize = std::env::var("CMF_MTP_CHAIN_PROBE")
4468 .ok()
4469 .and_then(|v| v.parse().ok())
4470 .unwrap_or(0);
4471 for p in pos..end {
4472 if p + 1 < input_ids.len() {
4473 if probe >= 1 && p + 2 < input_ids.len() {
4474 let (d1, mut hx) = self.mtp_step_h(
4478 m,
4479 &hb[(p - pos) * hs..(p - pos + 1) * hs],
4480 input_ids[p + 1],
4481 p,
4482 );
4483 let mut ok = d1 == input_ids[p + 2];
4484 Self::chain_probe_note(0, ok);
4485 let mut d_prev = d1;
4486 let mut extra = 0usize;
4487 for j in 1..probe {
4488 if p + 2 + j >= input_ids.len() {
4489 break;
4490 }
4491 let (dj, hj) = self.mtp_step_h(m, &hx, d_prev, p + 1 + j);
4492 extra += 1;
4493 ok = ok && dj == input_ids[p + 2 + j];
4494 Self::chain_probe_note(j, ok);
4495 d_prev = dj;
4496 hx = hj;
4497 }
4498 m.kv.truncate_last(extra);
4499 } else {
4500 let _ = self.mtp_step(
4501 m,
4502 &hb[(p - pos) * hs..(p - pos + 1) * hs],
4503 input_ids[p + 1],
4504 p,
4505 );
4506 }
4507 }
4508 }
4509 }
4510 hidden.copy_from_slice(&hb[(end - pos - 1) * hs..]);
4511 pos = end;
4512 }
4513 }
4514 let pair_off = std::env::var("CMF_PAIR").is_ok_and(|v| v == "0");
4515 if task_mask.is_none()
4516 && !dyn_prefill
4517 && !graph_prefill
4518 && !pair_off
4519 && self.pair_supported()
4520 && o1_prefill.is_none()
4521 {
4522 while pos + 1 < input_ids.len()
4523 && !self.cancel.load(std::sync::atomic::Ordering::Relaxed)
4524 {
4525 let e1 = self.embed_single(input_ids[pos]);
4526 let e2 = self.embed_single(input_ids[pos + 1]);
4527 let (h1, h2) = self.forward_pair(&e1, &e2, pos);
4528 if mimo_spec {
4529 self.mimo_note_rows(&h1, pos);
4530 self.mimo_note_rows(&h2, pos + 1);
4531 }
4532 self.commit_linear_scratch();
4534 if let Some(m) = &mut mtp {
4535 let _ = self.mtp_step(m, &h1, input_ids[pos + 1], pos);
4536 if pos + 2 < input_ids.len() {
4537 let probe: usize = std::env::var("CMF_MTP_CHAIN_PROBE")
4538 .ok()
4539 .and_then(|v| v.parse().ok())
4540 .unwrap_or(0);
4541 if probe >= 1 && pos + 3 < input_ids.len() {
4542 let (d1, mut hx) = self.mtp_step_h(m, &h2, input_ids[pos + 2], pos + 1);
4546 let mut ok = d1 == input_ids[pos + 3];
4547 Self::chain_probe_note(0, ok);
4548 let mut d_prev = d1;
4549 let mut extra = 0usize;
4550 for j in 1..probe {
4551 if pos + 3 + j >= input_ids.len() {
4552 break;
4553 }
4554 let (dj, hj) = self.mtp_step_h(m, &hx, d_prev, pos + 2 + j);
4555 extra += 1;
4556 ok = ok && dj == input_ids[pos + 3 + j];
4557 Self::chain_probe_note(j, ok);
4558 d_prev = dj;
4559 hx = hj;
4560 }
4561 m.kv.truncate_last(extra);
4562 } else {
4563 let _ = self.mtp_step(m, &h2, input_ids[pos + 2], pos + 1);
4564 }
4565 }
4566 }
4567 hidden = h2;
4568 pos += 2;
4569 }
4570 }
4571 let o1_batch_ready = o1_sealed
4584 && o1_prefill.is_some()
4585 && mtp.is_none()
4586 && std::env::var("CMF_O1_GPU").as_deref() == Ok("1")
4587 && (0..self.num_layers).all(|li| {
4588 let cache = &self.kv_cache.layers[self.phys_layer(li)];
4589 cache.o1.is_none() || cache.o1_views().is_some()
4590 });
4591 let mtp_batch_prefill = mtp.is_some()
4596 && graph_prefill
4597 && task_mask.is_none()
4598 && !dyn_prefill
4599 && !self.o1_active()
4600 && std::env::var("CMF_MTP_CHAIN_PROBE").is_err();
4601 if batch_k > 0
4602 && (graph_prefill || o1_batch_ready)
4603 && task_mask.is_none()
4604 && (!self.o1_active() || o1_batch_ready)
4605 && (mtp.is_none() || mtp_batch_prefill)
4606 && !dyn_prefill
4607 && pos + 1 < input_ids.len()
4608 {
4609 let hs = self.hidden_size;
4610 let chunk = batch_k;
4611 while pos < input_ids.len() {
4612 let end = (pos + chunk).min(input_ids.len());
4613 let bk = end - pos;
4614 let mut hiddens = vec![0f32; bk * hs];
4615 for (j, &id) in input_ids[pos..end].iter().enumerate() {
4616 hiddens[j * hs..(j + 1) * hs].copy_from_slice(&self.embed_single(id));
4617 }
4618 let positions: Vec<usize> = (pos..end).collect();
4619 let t_chunk = std::time::Instant::now();
4620 let outcome = self.try_batch_graph_wgpu(&mut hiddens, &positions, bk, None);
4621 let ok_b = outcome == crate::gpu::BatchGraphOutcome::Completed;
4622 if std::env::var("CMF_GRAPH_PROF").is_ok() {
4623 let ms = t_chunk.elapsed().as_secs_f64() * 1000.0;
4624 eprintln!(
4625 "batch-chunk: phase=prompt mode={} k={bk} outcome={outcome:?} {ms:.1} ms ({:.1} tok/s)",
4626 if o1_batch_ready {
4627 "o1"
4628 } else if mtp_batch_prefill {
4629 "ordinary_mtp"
4630 } else {
4631 "ordinary"
4632 },
4633 bk as f64 / (ms / 1000.0)
4634 );
4635 }
4636 {
4637 use std::sync::atomic::{AtomicBool, Ordering};
4638 static SAID: AtomicBool = AtomicBool::new(false);
4639 if !SAID.swap(true, Ordering::Relaxed) {
4640 if ok_b {
4641 tracing::info!(
4642 "batched prefill: ACTIVE mode={} (k={bk})",
4643 if o1_batch_ready {
4644 "o1"
4645 } else if mtp_batch_prefill {
4646 "ordinary_mtp"
4647 } else {
4648 "ordinary"
4649 }
4650 );
4651 } else {
4652 tracing::warn!("batched prefill {:?} — per-position graph", outcome);
4653 }
4654 }
4655 }
4656 if ok_b {
4657 if mimo_spec {
4658 self.mimo_note_rows(&hiddens, pos);
4659 }
4660 if mtp_batch_prefill {
4661 let n_pairs = mtp_prefill_pair_count(pos, end, input_ids.len());
4662 if n_pairs > 0 {
4663 let rows: Vec<Vec<f32>> = (0..n_pairs)
4669 .map(|j| hiddens[j * hs..(j + 1) * hs].to_vec())
4670 .collect();
4671 let pairs: Vec<(&[f32], u32)> = rows
4672 .iter()
4673 .enumerate()
4674 .map(|(j, row)| (row.as_slice(), input_ids[pos + j + 1]))
4675 .collect();
4676 if std::env::var("CMF_GRAPH_PROF").is_ok() {
4677 eprintln!(
4678 "mtp-warm: phase=prompt mode=ordinary_mtp first_pos={} pairs={} last_pos={}",
4679 pos,
4680 n_pairs,
4681 pos + n_pairs - 1,
4682 );
4683 }
4684 let warm_error = if let Some(m) = mtp.as_mut() {
4685 self.mtp_warm_prefill_pairs(m, &pairs, pos).err()
4686 } else {
4687 None
4688 };
4689 if let Some(err) = warm_error {
4690 self.finish_generation(&mut mtp, &mut router, true);
4695 return Err(err.to_string());
4696 }
4697 }
4698 }
4699 hidden.copy_from_slice(&hiddens[(bk - 1) * hs..]);
4700 pos = end;
4701 } else if outcome == crate::gpu::BatchGraphOutcome::Failed {
4702 self.finish_generation(&mut mtp, &mut router, true);
4707 return Err(if o1_batch_ready {
4708 "sealed O(1) batch graph failed after admission".to_string()
4709 } else {
4710 "ordinary recurrent batch graph failed after admission".to_string()
4711 });
4712 } else {
4713 break; }
4715 }
4716 }
4717 if graph_prefill
4721 && task_mask.is_none()
4722 && mtp.is_none()
4723 && !dyn_prefill
4724 && pos == 0
4725 && input_ids.len() > 1
4726 && !self.cancel.load(std::sync::atomic::Ordering::Relaxed)
4727 {
4728 if let Some(lg) = self.embryo_prefill_chunked(input_ids, 0) {
4729 self.graph_logits = Some(lg);
4730 hidden = vec![0.0; self.hidden_size];
4731 pos = input_ids.len();
4732 }
4733 }
4734 while pos < input_ids.len() && !self.cancel.load(std::sync::atomic::Ordering::Relaxed) {
4735 self.graph_want_logits = fuse_lm && pos + 1 == input_ids.len();
4736 hidden = self.forward_layers(&self.embed_single(input_ids[pos]), pos, task_mask);
4737 if mimo_spec {
4738 self.mimo_note_rows(&hidden, pos);
4739 }
4740 if let Some(m) = &mut mtp {
4741 if pos + 1 < input_ids.len() {
4742 let probe: usize = std::env::var("CMF_MTP_CHAIN_PROBE")
4748 .ok()
4749 .and_then(|v| v.parse().ok())
4750 .unwrap_or(0);
4751 if probe >= 1 && pos + 2 < input_ids.len() {
4752 let (d1, mut hx) = self.mtp_step_h(m, &hidden, input_ids[pos + 1], pos);
4753 let mut ok = d1 == input_ids[pos + 2];
4754 Self::chain_probe_note(0, ok);
4755 let mut d_prev = d1;
4756 let mut extra = 0usize;
4757 for j in 1..probe {
4758 if pos + 2 + j >= input_ids.len() {
4759 break;
4760 }
4761 let (dj, hj) = self.mtp_step_h(m, &hx, d_prev, pos + 1 + j);
4762 extra += 1;
4763 ok = ok && dj == input_ids[pos + 2 + j];
4764 Self::chain_probe_note(j, ok);
4765 d_prev = dj;
4766 hx = hj;
4767 }
4768 m.kv.truncate_last(extra);
4771 } else {
4772 let _ = self.mtp_step(m, &hidden, input_ids[pos + 1], pos);
4773 }
4774 }
4775 }
4776 pos += 1;
4777 }
4778 if std::env::var("CMF_PREFILL_PROF").is_ok() {
4779 eprintln!(
4780 "prefill: {} tokens in {:.1} ms (batch_k={batch_k})",
4781 input_ids.len(),
4782 _tpf.elapsed().as_secs_f64() * 1000.0
4783 );
4784 }
4785 if self
4786 .graph_failed
4787 .swap(false, std::sync::atomic::Ordering::Relaxed)
4788 {
4789 self.finish_generation(&mut mtp, &mut router, true);
4794 return Err("GPU token graph failed during prefill".to_string());
4795 }
4796 if self
4799 .cancel
4800 .swap(false, std::sync::atomic::Ordering::Relaxed)
4801 {
4802 self.finish_generation(&mut mtp, &mut router, true);
4806 return Ok(GenerateResult {
4807 text: String::new(),
4808 token_ids: Vec::new(),
4809 prompt_tokens: input_ids.len(),
4810 tokens_generated: 0,
4811 finish_reason: "cancelled".to_string(),
4812 mtp_drafted: 0,
4813 mtp_accepted: 0,
4814 token_confidence: Vec::new(),
4815 traces: Vec::new(),
4816 });
4817 }
4818
4819 if !o1_sealed {
4822 match self.o1_seal_checked() {
4823 Ok(_) => {}
4824 Err(err) => {
4825 self.finish_generation(&mut mtp, &mut router, true);
4826 return Err(err);
4827 }
4828 }
4829 }
4830
4831 macro_rules! commit {
4833 ($id:expr) => {{
4834 all_ids.push($id);
4835 generated += 1;
4836 self.note_draft_id($id);
4837 if self.tokenizer.is_eos($id) && !self.ignore_eos {
4838 finish_reason = "stop".to_string();
4839 false
4840 } else {
4841 let token_text = self.tokenizer.decode_token($id);
4842 let mut go = true;
4843 if let Some(ref mut cb) = on_token {
4844 if !cb(&token_text) {
4845 finish_reason = "cancelled".to_string();
4846 go = false;
4847 }
4848 }
4849 go
4850 }
4851 }};
4852 }
4853
4854 let mut spec_trial = SpecTrial::Spec {
4865 t0: std::time::Instant::now(),
4866 gen0: generated,
4867 rounds: 0,
4868 };
4869 let mut spec_mon = SpecMon {
4875 metal: graph_spec && crate::gpu::q1_force() && spec_cheap_round,
4876 ..SpecMon::default()
4877 };
4878 let mut spec_watchdog_off = false;
4879 let mut spec_walls: Vec<f32> = Vec::new();
4882 let mut spec_round_end: Option<std::time::Instant> = None;
4885 if mimo_spec {
4886 if let Ok(path) = std::env::var("CMF_MIMO_MTP_PROBE") {
4887 if let Some(mut st) = self.mimo_mtp.take() {
4888 self.mimo_mtp_probe(&mut st, input_ids, &path);
4889 self.mimo_mtp = Some(st);
4890 }
4891 }
4892 }
4893 let mut next_pos = input_ids.len();
4895 'decode: while generated < max_tokens {
4896 if self
4897 .graph_failed
4898 .swap(false, std::sync::atomic::Ordering::Relaxed)
4899 {
4900 self.finish_generation(&mut mtp, &mut router, true);
4905 return Err("GPU token graph failed during decode".to_string());
4906 }
4907 if self
4908 .cancel
4909 .swap(false, std::sync::atomic::Ordering::Relaxed)
4910 {
4911 finish_reason = "cancelled".to_string();
4912 break 'decode;
4913 }
4914 if mimo_spec && next_pos > 0 {
4919 self.mimo_note_rows(&hidden, next_pos - 1);
4922 }
4923 let forced = self.spec_forced.take();
4924 let mut logits = match (forced, self.graph_logits.take()) {
4925 (Some(_), _) => Vec::new(),
4926 (None, Some(lg)) => lg,
4927 (None, None) => {
4928 let _prof = crate::cpuprof::time(crate::cpuprof::Slot::Head);
4929 inference::rms_norm_into(
4930 &hidden,
4931 &self.weights.final_norm,
4932 self.rms_eps,
4933 self.norm_style,
4934 &mut self.ws.n1,
4935 );
4936 self.lm_head_forward(&self.ws.n1)
4937 }
4938 };
4939 if generated
4942 == std::env::var("CMF_LOGIT_DUMP_STEP")
4943 .ok()
4944 .and_then(|v| v.parse().ok())
4945 .unwrap_or(0)
4946 {
4947 if let Ok(path) = std::env::var("CMF_LOGIT_DUMP") {
4948 let mut bytes: Vec<u8> = Vec::with_capacity((hidden.len() + logits.len()) * 4);
4949 for v in hidden.iter().chain(logits.iter()) {
4950 bytes.extend_from_slice(&v.to_le_bytes());
4951 }
4952 if let Err(e) = std::fs::write(&path, &bytes) {
4953 eprintln!("logit dump: failed to write {path}: {e}");
4954 self.finish_generation(&mut mtp, &mut router, true);
4955 return Err(format!("logit dump write failed: {e}"));
4956 }
4957 }
4958 }
4959 if let Ok(dir) = std::env::var("CMF_LOGIT_DUMP_ALL") {
4963 if !logits.is_empty() {
4964 let path = std::path::Path::new(&dir).join(format!("step{generated:05}.f32"));
4965 let bytes: Vec<u8> = logits.iter().flat_map(|v| v.to_le_bytes()).collect();
4966 if let Err(e) =
4967 std::fs::create_dir_all(&dir).and_then(|_| std::fs::write(&path, &bytes))
4968 {
4969 eprintln!("logit dump: failed to write {}: {e}", path.display());
4970 }
4971 }
4972 }
4973 let t_next = match forced {
4974 Some(c) => c,
4975 None => {
4976 let _prof = crate::cpuprof::time(crate::cpuprof::Slot::Sampler);
4977 sampler::sample_with_scratch_pool(
4978 &logits,
4979 &self.sampler_config,
4980 self.sampler_config.penalty_past(&all_ids, bounded_native),
4981 &mut self.rng,
4982 &mut self.sampler_scratch,
4983 self.pool.as_deref(),
4984 )
4985 }
4986 };
4987 if self.confidence_on {
4988 confidence.push(if logits.is_empty() {
4989 0.0
4990 } else {
4991 sampler::top1_prob_pool(
4992 self.pool.as_deref(),
4993 &mut self.sampler_scratch,
4994 &logits,
4995 t_next,
4996 calib_temp,
4997 )
4998 });
4999 }
5000 if !logits.is_empty() {
5001 attention::recycle_buf(&mut logits);
5002 }
5003 if trace_on {
5004 let skill = router.as_ref().and_then(|r| r.active_id());
5008 traces.push(TokenTrace {
5009 t: generated,
5010 token_id: t_next,
5011 confidence: confidence.last().copied().unwrap_or(0.0),
5012 active_skill: skill,
5013 recon: None,
5014 switched: false,
5015 });
5016 }
5017 if !commit!(t_next) {
5018 break 'decode;
5019 }
5020 if generated >= max_tokens {
5021 break 'decode;
5022 }
5023
5024 if self.dsv41.is_none() && self.kv_cache.needs_eviction() {
5025 static SAID: std::sync::Once = std::sync::Once::new();
5031 SAID.call_once(|| {
5032 tracing::warn!(
5033 "KV cache full at {} positions — evicting half; quality \
5034 will degrade. Raise CMF_MAX_SEQ.",
5035 self.kv_cache.max_seq_len,
5036 );
5037 });
5038 let keep = (self.kv_cache.max_seq_len / 2).max(1);
5039 self.kv_cache.evict(keep);
5040 }
5041
5042 if graph_spec {
5045 match spec_trial {
5046 SpecTrial::Plain { t0, gen0 } if spec_mon.plain_done(t0, gen0, generated) => {
5047 spec_mon.plain_ms =
5048 t0.elapsed().as_secs_f64() * 1e3 / (generated - gen0) as f64;
5049 let keep = spec_mon.pays();
5050 tracing::info!(
5051 "speculation trial: {:.2} tok/round in {:.1} ms vs plain {:.1} ms/tok — {}",
5052 spec_mon.tokens,
5053 spec_mon.round_ms,
5054 spec_mon.plain_ms,
5055 if keep { "speculating" } else { "plain" }
5056 );
5057 spec_mon.fails = 0;
5058 spec_trial = SpecTrial::Decided {
5059 spec: keep,
5060 recheck_at: if keep { usize::MAX } else { generated + 128 },
5061 };
5062 }
5063 SpecTrial::Decided { recheck_at, .. } if generated >= recheck_at => {
5064 spec_mon.n = 0;
5065 spec_trial = SpecTrial::Spec {
5066 t0: std::time::Instant::now(),
5067 gen0: generated,
5068 rounds: 0,
5069 };
5070 }
5071 _ => {}
5072 }
5073 spec_watchdog_off = matches!(
5074 spec_trial,
5075 SpecTrial::Plain { .. } | SpecTrial::Decided { spec: false, .. }
5076 );
5077 }
5078 if mimo_spec && generated + 1 < max_tokens && next_pos > 0 {
5080 let budget = max_tokens - generated - 1;
5081 if let Some(mut st) = self.mimo_mtp.take() {
5082 let k = st.depth.min(budget);
5083 let r = self.mimo_spec_round(&mut st, next_pos, &all_ids, k);
5084 self.mimo_mtp = Some(st);
5085 let r = match r {
5086 Ok(r) => r,
5087 Err(err) => {
5088 self.finish_generation(&mut mtp, &mut router, true);
5089 return Err(err);
5090 }
5091 };
5092 if let Some(r) = r {
5093 drafted += r.drafted;
5094 accepted += r.accepted.len();
5095 let mut stopped = false;
5096 for &id in &r.accepted {
5097 if self.confidence_on {
5098 confidence.push(0.0);
5099 }
5100 if !commit!(id) {
5101 stopped = true;
5102 break;
5103 }
5104 }
5105 if stopped {
5106 break 'decode;
5107 }
5108 next_pos += r.accepted.len() + 1;
5109 hidden = r.hidden;
5110 self.graph_logits = Some(r.logits);
5113 continue 'decode;
5114 }
5115 }
5116 }
5117 #[cfg(feature = "gpu")]
5120 if self.speculative
5121 && self.qwen4_exp.is_some()
5122 && task_mask.is_none()
5123 && self.sampler_config.temperature < 1e-6
5124 && generated + 1 < max_tokens
5125 && next_pos > 0
5126 && std::env::var("CMF_QWEN_MTP").as_deref() != Ok("0")
5127 {
5128 let r = match &mut self.qwen4_exp {
5129 Some(b) => crate::qwen4_exp::spec_round(
5130 &b.0,
5131 &b.1,
5132 &b.2,
5133 &mut b.3,
5134 next_pos,
5135 &all_ids,
5136 &self.inv_freq,
5137 self.pool.as_deref(),
5138 ),
5139 None => None,
5140 };
5141 if let Some(r) = r {
5142 drafted += r.drafted;
5143 accepted += r.accepted.len();
5144 let mut stopped = false;
5145 for &id in &r.accepted {
5146 if self.confidence_on {
5147 confidence.push(0.0);
5148 }
5149 if !commit!(id) {
5150 stopped = true;
5151 break;
5152 }
5153 }
5154 if stopped {
5155 break 'decode;
5156 }
5157 next_pos += r.accepted.len() + 1;
5158 hidden.fill(0.0);
5159 self.graph_logits = Some(r.logits);
5160 continue 'decode;
5161 }
5162 }
5163 match &mut mtp {
5164 #[cfg(feature = "gpu")]
5166 Some(m)
5167 if graph_spec
5168 && !spec_watchdog_off
5169 && generated + 1 < max_tokens
5170 && next_pos > 0 =>
5171 {
5172 let t_round = std::time::Instant::now();
5173 if spec_time_level() >= 2 {
5174 if let Some(t) = spec_round_end.take() {
5175 eprintln!(
5176 "spec-gap {:.2} ms (host between rounds)",
5177 t.elapsed().as_secs_f64() * 1e3
5178 );
5179 }
5180 }
5181 spec_stamps_begin();
5182 #[cfg(target_os = "macos")]
5187 let allocs0 = crate::gpu_metal::IO_BUF_ALLOCS
5188 .load(std::sync::atomic::Ordering::Relaxed);
5189 #[cfg(not(target_os = "macos"))]
5190 let allocs0 = 0u64;
5191 if let Some((extra, n_pos, new_h)) = self.graph_spec_step(
5192 m,
5193 &hidden,
5194 t_next,
5195 next_pos,
5196 &mut drafted,
5197 &mut accepted,
5198 &mut all_ids,
5199 max_tokens - generated,
5200 ) {
5201 next_pos = n_pos;
5202 hidden = new_h;
5203 let level = spec_time_level();
5204 if level > 0 {
5205 let wall = t_round.elapsed().as_secs_f32() * 1e3;
5206 let stamps = spec_stamps_take();
5207 let median = if spec_walls.len() >= 3 {
5210 let mut s = spec_walls.clone();
5211 s.sort_by(|a, b| a.partial_cmp(b).unwrap());
5212 Some(s[s.len() / 2])
5213 } else {
5214 None
5215 };
5216 let outlier = median.is_some_and(|m| wall > 1.4 * m);
5217 #[cfg(target_os = "macos")]
5218 let allocs = crate::gpu_metal::IO_BUF_ALLOCS
5219 .load(std::sync::atomic::Ordering::Relaxed)
5220 - allocs0;
5221 #[cfg(not(target_os = "macos"))]
5222 let allocs = allocs0;
5223 eprintln!(
5224 "spec-round wall {wall:.1} ms → {} tokens{}{}",
5225 extra.len() + 1,
5226 if allocs > 0 {
5227 format!(" [{allocs} new device buffers]")
5228 } else {
5229 String::new()
5230 },
5231 match (outlier, median) {
5232 (true, Some(m)) => format!(" OUTLIER (median {m:.1})"),
5233 _ => String::new(),
5234 }
5235 );
5236 if level >= 2 || outlier {
5237 let sum: f32 = stamps.iter().map(|s| s.1).sum();
5238 eprintln!(
5239 "spec-stamps: {}| untracked {:.1}",
5240 spec_stamps_format(&stamps),
5241 wall - sum
5242 );
5243 }
5244 if spec_mon.n >= 1 {
5245 spec_walls.push(wall);
5246 }
5247 }
5248 spec_mon.round(t_round.elapsed().as_secs_f64() * 1e3, extra.len() + 1);
5252 spec_trial = Self::spec_trial_round(
5255 spec_trial,
5256 &mut spec_mon,
5257 generated + extra.len() + 1,
5258 );
5259 let mut stopped = false;
5260 for &id in &extra {
5261 if self.confidence_on {
5262 confidence.push(0.0);
5263 }
5264 if !commit!(id) {
5265 stopped = true;
5266 break;
5267 }
5268 }
5269 if stopped {
5270 break 'decode;
5271 }
5272 if spec_time_level() >= 2 {
5273 spec_round_end = Some(std::time::Instant::now());
5274 }
5275 continue 'decode;
5276 }
5277 if self
5278 .graph_failed
5279 .swap(false, std::sync::atomic::Ordering::Relaxed)
5280 {
5281 self.finish_generation(&mut mtp, &mut router, true);
5287 return Err("GPU MTP graph failed during speculative decode".to_string());
5288 }
5289 spec_mon.round(t_round.elapsed().as_secs_f64() * 1e3, 1);
5300 spec_mon.tokens = 0.0;
5301 spec_mon.fails = 3;
5302 spec_trial = Self::spec_trial_round(spec_trial, &mut spec_mon, generated + 1);
5303 hidden = self.forward_layers(&self.embed_single(t_next), next_pos, task_mask);
5304 next_pos += 1;
5305 continue 'decode;
5306 }
5307 Some(m) if !graph_spec && generated + 1 < max_tokens => {
5309 let draft = self.mtp_step(m, &hidden, t_next, next_pos - 1);
5310 drafted += 1;
5311 let emb1 = self.embed_single(t_next);
5312 let emb2 = self.embed_single(draft);
5313 let (h1, h2) = self.forward_pair(&emb1, &emb2, next_pos);
5314
5315 inference::rms_norm_into(
5316 &h1,
5317 &self.weights.final_norm,
5318 self.rms_eps,
5319 self.norm_style,
5320 &mut self.ws.n1,
5321 );
5322 let mut logits1 = self.lm_head_forward(&self.ws.n1);
5323 let t_after = sampler::sample_with_scratch_pool(
5324 &logits1,
5325 &self.sampler_config,
5326 self.sampler_config.penalty_past(&all_ids, bounded_native),
5327 &mut self.rng,
5328 &mut self.sampler_scratch,
5329 self.pool.as_deref(),
5330 );
5331 if self.confidence_on {
5332 confidence.push(sampler::top1_prob_pool(
5333 self.pool.as_deref(),
5334 &mut self.sampler_scratch,
5335 &logits1,
5336 t_after,
5337 calib_temp,
5338 ));
5339 }
5340 attention::recycle_buf(&mut logits1);
5341 if trace_on {
5342 traces.push(TokenTrace {
5345 t: generated,
5346 token_id: t_after,
5347 confidence: confidence.last().copied().unwrap_or(0.0),
5348 active_skill: None,
5349 recon: None,
5350 switched: false,
5351 });
5352 }
5353 let stop = !commit!(t_after);
5354
5355 if t_after == draft {
5356 accepted += 1;
5357 self.commit_linear_scratch();
5358 let _ = self.mtp_step(m, &h1, t_after, next_pos);
5359 hidden = h2;
5360 next_pos += 2;
5361 } else {
5362 for layer in &mut self.kv_cache.layers {
5364 layer.truncate_last(1);
5365 }
5366 if !stop {
5367 let _ = self.mtp_step(m, &h1, t_after, next_pos);
5368 hidden = self.forward_layers(
5369 &self.embed_single(t_after),
5370 next_pos + 1,
5371 None,
5372 );
5373 }
5374 next_pos += 2;
5375 }
5376 if stop {
5377 break 'decode;
5378 }
5379 }
5380 _ => {
5382 #[cfg(feature = "gpu")]
5387 if Self::dsv4_spec_on() && self.dsv4.is_some() {
5388 static SAID: std::sync::Once = std::sync::Once::new();
5389 SAID.call_once(|| {
5390 eprintln!(
5391 "dsv4-spec гейт: mtp={} mask={} router={} trace={} temp={} rep={} ",
5392 !self.dsv4_mtp.is_empty(),
5393 task_mask.is_none(),
5394 router.is_none(),
5395 !trace_on,
5396 self.sampler_config.temperature < 1e-6,
5397 self.sampler_config.repetition_penalty == 1.0,
5398 );
5399 });
5400 }
5401 #[cfg(feature = "gpu")]
5402 if Self::dsv4_spec_on()
5403 && self.dsv4.is_some()
5404 && !self.dsv4_mtp.is_empty()
5405 && task_mask.is_none()
5406 && router.is_none()
5407 && !trace_on
5408 && self.sampler_config.temperature < 1e-6
5409 && self.sampler_config.repetition_penalty == 1.0
5410 && generated + 1 < max_tokens
5411 && all_ids.len() >= 2
5412 && generated >= dsv4_spec_retry_at
5413 {
5414 let tip_token = all_ids[all_ids.len() - 2];
5415 let drafted0 = drafted;
5416 let round = self.dsv4_spec_step(
5417 tip_token,
5418 t_next,
5419 next_pos,
5420 max_tokens.saturating_sub(generated),
5421 &mut drafted,
5422 &mut accepted,
5423 );
5424 if drafted > drafted0 {
5425 let useful = round.as_ref().is_some_and(|(extra, _)| !extra.is_empty());
5426 if useful {
5427 dsv4_spec_bad = 0;
5428 } else {
5429 dsv4_spec_bad += 1;
5430 if dsv4_spec_bad >= 2 {
5431 dsv4_spec_bad = 0;
5432 dsv4_spec_retry_at = generated.saturating_add(32);
5433 tracing::info!(
5434 "dsv4: draft не окупился дважды — точный walk на 32 токена"
5435 );
5436 }
5437 }
5438 }
5439 if let Some((extra, n_pos)) = round {
5440 next_pos = n_pos;
5441 let mut stopped = false;
5442 for &id in &extra {
5443 if self.confidence_on {
5444 confidence.push(0.0);
5445 }
5446 if !commit!(id) {
5447 stopped = true;
5448 break;
5449 }
5450 }
5451 if stopped {
5452 break 'decode;
5453 }
5454 continue 'decode;
5455 }
5456 }
5457 self.graph_want_logits = fuse_lm;
5458 let mut t_fwd = t_next;
5464 let pure_greedy = self.sampler_config.temperature < 1e-6
5465 && self.sampler_config.repetition_penalty == 1.0
5466 && self.sampler_config.suppress_tokens.is_empty();
5467 let burst_k = std::env::var("CMF_MULTISTEP")
5472 .ok()
5473 .and_then(|v| v.parse::<usize>().ok())
5474 .unwrap_or(0);
5475 if pure_greedy
5476 && burst_k >= 1
5477 && fuse_lm
5478 && task_mask.is_none()
5479 && router.is_none()
5480 && !trace_on
5481 && !self.confidence_on
5482 {
5483 let mut stopped = false;
5484 loop {
5485 let room = max_tokens.saturating_sub(generated);
5486 if room <= 2 {
5487 break;
5488 }
5489 let k = burst_k.min(room - 1);
5490 if k < 1 {
5491 break;
5492 }
5493 let Some(ids) = self.try_multi_burst(t_fwd, next_pos, k) else {
5494 if self
5495 .graph_failed
5496 .swap(false, std::sync::atomic::Ordering::Relaxed)
5497 {
5498 self.finish_generation(&mut mtp, &mut router, true);
5499 return Err(
5500 "GPU token graph failed during greedy burst".to_string()
5501 );
5502 }
5503 break;
5504 };
5505 next_pos += k;
5506 for &id in &ids {
5507 if !commit!(id) {
5508 stopped = true;
5509 break;
5510 }
5511 }
5512 if stopped {
5513 break;
5514 }
5515 t_fwd = *ids.last().unwrap();
5516 }
5517 if stopped {
5518 break 'decode;
5519 }
5520 }
5521 #[cfg(target_os = "macos")]
5531 if graph_spec
5532 && spec_watchdog_off
5533 && next_pos > 0
5534 && self.mtp_graph_mode == Some(true)
5535 && crate::gpu::q1_force()
5536 {
5537 if let Some(m) = mtp.as_mut() {
5538 let _ = self.mtp_step_metal(m, &hidden, t_fwd, next_pos - 1, false);
5539 }
5540 }
5541 hidden = self.forward_layers(&self.embed_single(t_fwd), next_pos, task_mask);
5542 next_pos += 1;
5543 if let Some(r) = &mut router {
5546 let phi = self.dyn_phi_ema.clone();
5547 let decision = r.step(&phi, generated);
5548 if let Some(new_active) = decision {
5549 let _ = self.set_active_skill(new_active);
5550 }
5551 if trace_on {
5554 if let Some(last) = traces.last_mut() {
5555 let e = r.last_best_e();
5556 last.recon = e.is_finite().then_some(e);
5557 last.switched = decision.is_some();
5558 }
5559 }
5560 }
5561 }
5562 }
5563 }
5564
5565 let cancelled = finish_reason == "cancelled";
5566 let dyn_switched = router.as_ref().is_some_and(|r| !r.switches.is_empty());
5570 if mimo_spec {
5571 if let Some(st) = self.mimo_mtp.as_ref() {
5572 let line = st.stats.line();
5573 tracing::info!("{line}");
5574 if std::env::var_os("CMF_MIMO_MTP_STATS").is_some() {
5575 eprintln!("{line}");
5576 }
5577 }
5578 }
5579 self.finish_generation(&mut mtp, &mut router, cancelled);
5580
5581 let output_ids = &all_ids[input_ids.len()..];
5582 let forwarded = input_ids.len() + output_ids.len().saturating_sub(1);
5586 let consumed = std::mem::take(&mut all_ids);
5591 if dyn_switched {
5592 self.clear_sequence_state();
5593 } else if cancelled || mimo_spec || prompt_rows.is_some() {
5594 self.clear_history();
5595 } else {
5596 self.record_consumed_prefix(&consumed[..forwarded.min(consumed.len())], reuse_from);
5597 }
5598 all_ids = consumed;
5599 let output_ids = &all_ids[input_ids.len()..];
5600 confidence.truncate(output_ids.len()); traces.truncate(output_ids.len());
5602 Ok(GenerateResult {
5603 text: self.tokenizer.decode(output_ids),
5604 token_ids: output_ids.to_vec(),
5605 prompt_tokens: input_ids.len(),
5606 tokens_generated: generated,
5607 finish_reason,
5608 mtp_drafted: drafted,
5609 mtp_accepted: accepted,
5610 token_confidence: confidence,
5611 traces,
5612 })
5613 }
5614
5615 fn mtp_step(
5619 &mut self,
5620 m: &mut MtpModule,
5621 hidden: &[f32],
5622 next_token: u32,
5623 position: usize,
5624 ) -> u32 {
5625 self.mtp_step_h(m, hidden, next_token, position).0
5626 }
5627
5628 fn chain_probe_note(depth: usize, prefix_ok: bool) {
5632 use std::sync::Mutex;
5633 static T: Mutex<Vec<(u64, u64)>> = Mutex::new(Vec::new());
5634 let mut t = T.lock().unwrap();
5635 if t.len() <= depth {
5636 t.resize(depth + 1, (0, 0));
5637 }
5638 t[depth].0 += 1;
5639 t[depth].1 += prefix_ok as u64;
5640 if depth == 0 && t[0].0 % 128 == 0 {
5641 let line: Vec<String> = t
5642 .iter()
5643 .enumerate()
5644 .map(|(d, (n, k))| {
5645 format!(
5646 "d{}={:.0}%({n})",
5647 d + 1,
5648 100.0 * *k as f64 / (*n).max(1) as f64
5649 )
5650 })
5651 .collect();
5652 eprintln!("mtp-chain: {}", line.join(" "));
5653 }
5654 }
5655
5656 fn mtp_step_hl(
5664 &mut self,
5665 m: &mut MtpModule,
5666 hidden: &[f32],
5667 next_token: u32,
5668 position: usize,
5669 ) -> (Vec<f32>, Vec<f32>) {
5670 #[cfg(target_os = "macos")]
5675 if self.mtp_graph_mode != Some(false) && crate::gpu::q1_force() {
5676 if let Some(r) = self.mtp_step_metal(m, hidden, next_token, position, true) {
5677 self.mtp_graph_mode = Some(true);
5678 return r;
5679 }
5680 if self.mtp_graph_mode == Some(true) {
5681 tracing::error!("mtp Metal graph failed after admission");
5682 self.clear_sequence_state();
5683 self.graph_failed
5684 .store(true, std::sync::atomic::Ordering::Relaxed);
5685 self.cancel
5686 .store(true, std::sync::atomic::Ordering::Relaxed);
5687 return (Vec::new(), Vec::new());
5688 }
5689 self.mtp_graph_mode = Some(false);
5690 }
5691 #[cfg(feature = "gpu")]
5692 if self.mtp_graph_mode != Some(false) {
5693 if !self.mtp_graph_ok(m) {
5694 if self.mtp_graph_mode == Some(true) {
5695 tracing::error!("mtp graph became unavailable after admission");
5700 self.clear_sequence_state();
5701 self.graph_failed
5702 .store(true, std::sync::atomic::Ordering::Relaxed);
5703 self.cancel
5704 .store(true, std::sync::atomic::Ordering::Relaxed);
5705 return (Vec::new(), Vec::new());
5706 }
5707 self.mtp_graph_mode = Some(false);
5708 } else {
5709 if let Some(r) = self.mtp_step_graph(m, hidden, next_token, position) {
5710 self.mtp_graph_mode = Some(true);
5711 return r;
5712 }
5713 if self.graph_failed.load(std::sync::atomic::Ordering::Relaxed) {
5714 return (Vec::new(), Vec::new());
5721 }
5722 tracing::error!("mtp graph failed or declined after admission");
5726 self.clear_sequence_state();
5727 self.graph_failed
5728 .store(true, std::sync::atomic::Ordering::Relaxed);
5729 self.cancel
5730 .store(true, std::sync::atomic::Ordering::Relaxed);
5731 return (Vec::new(), Vec::new());
5732 }
5733 }
5734 let e = self.embed_single(next_token);
5738 let mut cat = vec![0.0f32; 2 * self.hidden_size];
5739 let (cat_e, cat_h) = cat.split_at_mut(self.hidden_size);
5740 inference::rms_norm_into(&e, &m.enorm, self.rms_eps, self.norm_style, cat_e);
5741 inference::rms_norm_into(hidden, &m.hnorm, self.rms_eps, self.norm_style, cat_h);
5742 let mut x = vec![0.0f32; self.hidden_size];
5743 m.eh_proj.matvec(&cat, &mut x, self.pool.as_deref());
5744
5745 let lw = &m.layer;
5747 inference::rms_norm_into(
5748 &x,
5749 &lw.input_norm,
5750 self.rms_eps,
5751 self.norm_style,
5752 &mut self.ws.n1,
5753 );
5754 let attn = match &lw.attn {
5755 AttnKind::Mla(_) => unreachable!("MLA has no MTP/pair path"),
5757 AttnKind::Kda(_) => unreachable!("KDA has no MTP/pair path"),
5758 AttnKind::Full {
5759 wq,
5760 wk,
5761 wv,
5762 wo,
5763 q_norm,
5764 k_norm,
5765 output_gate,
5766 softplus_gate,
5767 bias,
5768 } => {
5769 let mut cfg = self.attn_cfg(position);
5770 cfg.q_norm = q_norm.as_deref();
5771 cfg.k_norm = k_norm.as_deref();
5772 cfg.output_gate = *output_gate;
5773 cfg.softplus_gate = softplus_gate
5774 .as_ref()
5775 .map(|(gate, per_head)| (gate, *per_head));
5776 cfg.bias = bias
5777 .as_ref()
5778 .map(|(q, k, v)| (q.as_slice(), k.as_slice(), v.as_slice()));
5779 attention::qwen_attention(&self.ws.n1, wq, wk, wv, wo, &mut m.kv, &cfg)
5780 }
5781 AttnKind::Linear(_)
5782 | AttnKind::LinearGdn(_)
5783 | AttnKind::ShortConv(_)
5784 | AttnKind::Bounded(_) => {
5785 unreachable!("MTP block is full attention")
5786 }
5787 };
5788 for (i, &a) in attn.iter().enumerate() {
5789 x[i] += a;
5790 }
5791 inference::rms_norm_into(
5792 &x,
5793 &lw.post_norm,
5794 self.rms_eps,
5795 self.norm_style,
5796 &mut self.ws.p1,
5797 );
5798 let ffn = ffn_forward(&lw.ffn, &self.ws.p1, self.pool.as_deref(), None);
5799 for (i, &f) in ffn.iter().enumerate() {
5800 x[i] += f;
5801 }
5802
5803 inference::rms_norm_into(
5804 &x,
5805 &m.final_norm,
5806 self.rms_eps,
5807 self.norm_style,
5808 &mut self.ws.n1,
5809 );
5810 let lg = self.lm_head_forward(&self.ws.n1);
5811 (lg, x)
5812 }
5813
5814 fn mtp_step_h(
5816 &mut self,
5817 m: &mut MtpModule,
5818 hidden: &[f32],
5819 next_token: u32,
5820 position: usize,
5821 ) -> (u32, Vec<f32>) {
5822 let (mut lg, x) = self.mtp_step_hl(m, hidden, next_token, position);
5823 let draft = sampler::argmax(&lg);
5824 attention::recycle_buf(&mut lg);
5825 (draft, x)
5826 }
5827
5828 fn spec_trial_round(trial: SpecTrial, mon: &mut SpecMon, generated: usize) -> SpecTrial {
5834 match trial {
5835 SpecTrial::Spec { t0, gen0, rounds } => {
5836 let rounds = rounds + 1;
5837 if rounds >= 5 {
5838 if mon.plain_ms > 0.0 {
5839 let keep = mon.pays();
5840 mon.fails = 0;
5841 tracing::info!(
5842 "speculation re-check: {:.2} tok/round in {:.1} ms vs plain {:.1} ms/tok — {}",
5843 mon.tokens,
5844 mon.round_ms,
5845 mon.plain_ms,
5846 if keep { "speculating" } else { "plain" }
5847 );
5848 SpecTrial::Decided {
5849 spec: keep,
5850 recheck_at: if keep { usize::MAX } else { generated + 128 },
5851 }
5852 } else if mon.pays() {
5853 mon.fails = 0;
5858 tracing::info!(
5859 "speculation trial: {:.2} tok/round in {:.1} ms — speculating (plain not timed)",
5860 mon.tokens,
5861 mon.round_ms,
5862 );
5863 SpecTrial::Decided {
5864 spec: true,
5865 recheck_at: usize::MAX,
5866 }
5867 } else {
5868 SpecTrial::Plain {
5869 t0: std::time::Instant::now(),
5870 gen0: generated,
5871 }
5872 }
5873 } else {
5874 SpecTrial::Spec { t0, gen0, rounds }
5875 }
5876 }
5877 SpecTrial::Decided { spec: true, .. } => {
5878 if mon.pays() {
5879 mon.fails = 0;
5880 trial
5881 } else {
5882 mon.fails += 1;
5883 if mon.fails >= 4 {
5884 if mon.plain_ms <= 0.0 {
5885 tracing::info!(
5889 "speculation doubtful: {:.2} tok/round in {:.1} ms — timing plain",
5890 mon.tokens,
5891 mon.round_ms,
5892 );
5893 return SpecTrial::Plain {
5894 t0: std::time::Instant::now(),
5895 gen0: generated,
5896 };
5897 }
5898 tracing::info!(
5899 "speculation stopped: {:.2} tok/round in {:.1} ms vs plain {:.1} ms/tok",
5900 mon.tokens,
5901 mon.round_ms,
5902 mon.plain_ms
5903 );
5904 SpecTrial::Decided {
5905 spec: false,
5906 recheck_at: generated + 128,
5907 }
5908 } else {
5909 trial
5910 }
5911 }
5912 }
5913 other => other,
5914 }
5915 }
5916
5917 fn mtp_kv_id(&self) -> u64 {
5920 self.graph_kv_id | (1u64 << 40)
5921 }
5922
5923 const MTP_LAYER_BASE: usize = 0;
5928
5929 #[cfg(feature = "gpu")]
5936 fn rewind_mtp_graph_mirror(&self, stored: usize) -> bool {
5937 self.mtp_graph_mode != Some(true)
5938 || crate::gpu::graph_kv_set_stored(self.mtp_kv_id(), Self::MTP_LAYER_BASE, stored)
5939 }
5940
5941 #[cfg(feature = "gpu")]
5947 fn rewind_trunk_graph_mirrors(&self, stored: usize) -> bool {
5948 let mut ok = true;
5949 let mut expected = false;
5950 for li in 0..self.num_layers {
5951 if matches!(
5952 self.weights.layers[self.phys_layer(li)].attn,
5953 AttnKind::Full { .. }
5954 ) {
5955 expected = true;
5956 ok &= crate::gpu::graph_kv_set_stored(self.graph_kv_id, li, stored);
5957 }
5958 }
5959 !expected || ok
5960 }
5961
5962 fn graph_gdn_layer_count(&self) -> usize {
5966 (0..self.num_layers)
5967 .filter(|&li| {
5968 matches!(
5969 &self.weights.layers[self.phys_layer(li)].attn,
5970 AttnKind::LinearGdn(_)
5971 )
5972 })
5973 .count()
5974 }
5975
5976 fn mtp_block_input(&mut self, m: &MtpModule, hidden: &[f32], next_token: u32) -> Vec<f32> {
5979 let e = self.embed_single(next_token);
5980 let mut cat = vec![0.0f32; 2 * self.hidden_size];
5981 let (cat_e, cat_h) = cat.split_at_mut(self.hidden_size);
5982 inference::rms_norm_into(&e, &m.enorm, self.rms_eps, self.norm_style, cat_e);
5983 inference::rms_norm_into(hidden, &m.hnorm, self.rms_eps, self.norm_style, cat_h);
5984 let mut x = vec![0.0f32; self.hidden_size];
5985 m.eh_proj.matvec(&cat, &mut x, self.pool.as_deref());
5986 x
5987 }
5988
5989 #[cfg(feature = "gpu")]
5992 fn mtp_block_graph_ok(&self, m: &MtpModule) -> bool {
5993 if std::env::var("CMF_MTP_GRAPH").as_deref() == Ok("0") {
5994 return false;
5995 }
5996 if !crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Decode)
5997 || !crate::gpu::enabled_here()
5998 || self.attn_softcap > 0.0
5999 || self.attention_heads_per_layer.is_some()
6000 || self.v_head_dim.is_some()
6003 {
6004 return false;
6005 }
6006 matches!(
6007 &m.layer.attn,
6008 AttnKind::Full {
6009 softplus_gate: None,
6010 ..
6011 }
6012 ) && matches!(&m.layer.ffn, FfnKind::Dense(_))
6013 }
6014
6015 #[cfg(feature = "gpu")]
6019 fn mtp_graph_ok(&self, m: &MtpModule) -> bool {
6020 if !self.mtp_block_graph_ok(m) {
6021 return false;
6022 }
6023 let AttnKind::Full { wq, wk, wv, wo, .. } = &m.layer.attn else {
6024 return false;
6025 };
6026 let FfnKind::Dense(d) = &m.layer.ffn else {
6027 return false;
6028 };
6029 d.segs.is_empty()
6030 && wq.graph_weight().is_some()
6031 && wk.graph_weight().is_some()
6032 && wv.graph_weight().is_some()
6033 && wo.graph_weight().is_some()
6034 && d.gate_proj.graph_weight().is_some()
6035 && d.up_proj.graph_weight().is_some()
6036 && d.down_proj.graph_weight().is_some()
6037 && self.weights.lm_head.graph_weight().is_some()
6038 }
6039
6040 #[cfg(feature = "gpu")]
6046 fn mtp_step_graph(
6047 &mut self,
6048 m: &mut MtpModule,
6049 hidden: &[f32],
6050 next_token: u32,
6051 position: usize,
6052 ) -> Option<(Vec<f32>, Vec<f32>)> {
6053 if !self.mtp_graph_ok(m) {
6054 return None;
6055 }
6056 let lw = &m.layer;
6057 let AttnKind::Full {
6058 wq,
6059 wk,
6060 wv,
6061 wo,
6062 q_norm,
6063 k_norm,
6064 output_gate,
6065 softplus_gate,
6066 bias,
6067 } = &lw.attn
6068 else {
6069 return None;
6070 };
6071 if softplus_gate.is_some() {
6072 return None;
6073 }
6074 let FfnKind::Dense(d) = &lw.ffn else {
6075 return None;
6076 };
6077 if !d.segs.is_empty() {
6078 return None; }
6080 let mut x = self.mtp_block_input(m, hidden, next_token);
6083 fn gw(t: &QTensor) -> Option<crate::gpu::GraphW<'_>> {
6084 let (_, i, kind, rs) = t.graph_weight()?;
6085 Some(crate::gpu::GraphW {
6086 idx: i,
6087 kind,
6088 row_scale: rs,
6089 data: &[],
6090 prism: crate::gpu::GraphPrismOp::None,
6091 affine: false,
6092 })
6093 }
6094 let (model, _, _, _) = wq.graph_weight()?;
6095 let model = model.clone();
6096 let (lm_gw, lm_rows) = {
6097 let (_, i, kind, rs) = self.weights.lm_head.graph_weight()?;
6098 let rows = if kind == 6 {
6102 self.draft_head_rows(self.weights.lm_head.rows())
6103 } else {
6104 self.weights.lm_head.rows()
6105 };
6106 (
6107 crate::gpu::GraphW {
6108 idx: i,
6109 kind,
6110 row_scale: rs,
6111 data: &[],
6112 prism: crate::gpu::GraphPrismOp::None,
6113 affine: false,
6114 },
6115 rows,
6116 )
6117 };
6118 let layer = crate::gpu::GraphLayer {
6119 input_norm: &lw.input_norm,
6120 attn: crate::gpu::GraphAttn::Full {
6121 wq: gw(wq)?,
6122 wk: gw(wk)?,
6123 wv: gw(wv)?,
6124 wo: gw(wo)?,
6125 q_norm: q_norm.as_deref(),
6126 k_norm: k_norm.as_deref(),
6127 late_qk_norm: self.qk_norm_after_rope,
6128 bias: bias
6129 .as_ref()
6130 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
6131 output_gate: *output_gate,
6132 cpu_k: m.kv.k_heads(),
6133 cpu_v: m.kv.v_heads(),
6134 geom: None,
6135 },
6136 post_norm: &lw.post_norm,
6137 ffn: crate::gpu::GraphFfn::Dense {
6138 gate: gw(&d.gate_proj)?,
6139 up: gw(&d.up_proj)?,
6140 down: gw(&d.down_proj)?,
6141 },
6142 };
6143 let nh = self.num_heads;
6144 let (nkv, hd, rd) = self.layer_geom(0);
6145 let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
6146 let mut logits = Vec::new();
6147 let ok = crate::gpu::forward_token_graph(
6148 &model,
6149 self.mtp_kv_id(),
6150 std::slice::from_ref(&layer),
6151 &[None],
6152 self.o1_epoch,
6153 &self.inv_freq,
6154 &mut x,
6155 nh,
6156 nkv,
6157 hd,
6158 self.attn_scale,
6159 rd,
6160 self.hidden_size,
6161 self.intermediate_size,
6162 position,
6163 self.kv_cache.max_seq_len,
6164 gemma,
6165 self.rms_eps as f32,
6166 Some((&lm_gw, lm_rows)),
6167 &m.final_norm,
6168 &mut logits,
6169 &[],
6170 1,
6171 None,
6172 None,
6173 None,
6174 Self::MTP_LAYER_BASE,
6175 true,
6176 );
6177 match ok {
6178 crate::gpu::TokenGraphOutcome::Completed => {}
6179 crate::gpu::TokenGraphOutcome::Declined => return None,
6180 crate::gpu::TokenGraphOutcome::Failed => {
6181 self.clear_sequence_state();
6185 self.graph_failed
6186 .store(true, std::sync::atomic::Ordering::Relaxed);
6187 self.cancel
6188 .store(true, std::sync::atomic::Ordering::Relaxed);
6189 return None;
6190 }
6191 }
6192 logits.resize(self.vocab_size, 0.0);
6193 Some((logits, x))
6194 }
6195
6196 #[cfg(feature = "gpu")]
6204 fn mtp_warm_graph(
6205 &mut self,
6206 m: &mut MtpModule,
6207 pairs: &[(&[f32], u32)],
6208 first_pos: usize,
6209 ) -> crate::gpu::BatchGraphOutcome {
6210 if pairs.is_empty() {
6211 return crate::gpu::BatchGraphOutcome::Completed;
6212 }
6213 if !self.mtp_block_graph_ok(m) {
6214 return crate::gpu::BatchGraphOutcome::Declined;
6215 }
6216 let hs = self.hidden_size;
6217 let mut hiddens = Vec::with_capacity(pairs.len() * hs);
6220 for (h, t) in pairs {
6221 hiddens.extend_from_slice(&self.mtp_block_input(m, h, *t));
6222 }
6223 let lw = &m.layer;
6224 let AttnKind::Full {
6225 wq,
6226 wk,
6227 wv,
6228 wo,
6229 q_norm,
6230 k_norm,
6231 output_gate,
6232 bias,
6233 ..
6234 } = &lw.attn
6235 else {
6236 return crate::gpu::BatchGraphOutcome::Declined;
6237 };
6238 let FfnKind::Dense(d) = &lw.ffn else {
6239 return crate::gpu::BatchGraphOutcome::Declined;
6240 };
6241 if !d.segs.is_empty() {
6242 return crate::gpu::BatchGraphOutcome::Declined; }
6244 fn gw(t: &QTensor) -> Option<crate::gpu::GraphW<'_>> {
6245 let (_, i, kind, rs) = t.graph_weight()?;
6246 Some(crate::gpu::GraphW {
6247 idx: i,
6248 kind,
6249 row_scale: rs,
6250 data: &[],
6251 prism: crate::gpu::GraphPrismOp::None,
6252 affine: false,
6253 })
6254 }
6255 let Some((model, _, _, _)) = wq.graph_weight() else {
6256 return crate::gpu::BatchGraphOutcome::Declined;
6257 };
6258 let model = model.clone();
6259 let (Some(gwq), Some(gwk), Some(gwv), Some(gwo), Some(gg), Some(gu), Some(gd)) = (
6260 gw(wq),
6261 gw(wk),
6262 gw(wv),
6263 gw(wo),
6264 gw(&d.gate_proj),
6265 gw(&d.up_proj),
6266 gw(&d.down_proj),
6267 ) else {
6268 return crate::gpu::BatchGraphOutcome::Declined;
6269 };
6270 let layer = crate::gpu::GraphLayer {
6271 input_norm: &lw.input_norm,
6272 attn: crate::gpu::GraphAttn::Full {
6273 wq: gwq,
6274 wk: gwk,
6275 wv: gwv,
6276 wo: gwo,
6277 q_norm: q_norm.as_deref(),
6278 k_norm: k_norm.as_deref(),
6279 late_qk_norm: self.qk_norm_after_rope,
6280 bias: bias
6281 .as_ref()
6282 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
6283 output_gate: *output_gate,
6284 cpu_k: m.kv.k_heads(),
6285 cpu_v: m.kv.v_heads(),
6286 geom: None,
6287 },
6288 post_norm: &lw.post_norm,
6289 ffn: crate::gpu::GraphFfn::Dense {
6290 gate: gg,
6291 up: gu,
6292 down: gd,
6293 },
6294 };
6295 let positions: Vec<usize> = (first_pos..first_pos + pairs.len()).collect();
6296 let nh = self.num_heads;
6297 let (nkv, hd, rd) = self.layer_geom(0);
6298 let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
6299 crate::gpu::forward_batch_graph(
6300 &model,
6301 self.mtp_kv_id(),
6302 std::slice::from_ref(&layer),
6303 &self.inv_freq,
6304 &mut hiddens,
6305 nh,
6306 nkv,
6307 hd,
6308 rd,
6309 hs,
6310 self.intermediate_size,
6311 &positions,
6312 self.kv_cache.max_seq_len,
6313 gemma,
6314 self.rms_eps as f32,
6315 self.attn_scale,
6316 pairs.len(),
6317 &[],
6318 0,
6319 None,
6320 None,
6321 )
6322 }
6323
6324 #[cfg(feature = "gpu")]
6331 fn mtp_warm_graph_fallback(
6332 &mut self,
6333 m: &mut MtpModule,
6334 pairs: &[(&[f32], u32)],
6335 first_pos: usize,
6336 ) -> bool {
6337 if pairs.is_empty() {
6338 return true;
6339 }
6340 let graphable = self.mtp_block_graph_ok(m);
6341 if !graphable {
6342 if self.mtp_graph_mode == Some(true) {
6346 return false;
6347 }
6348 self.mtp_graph_mode = Some(false);
6349 for (j, (h, t)) in pairs.iter().enumerate() {
6350 self.mtp_warm(m, h, *t, first_pos + j);
6351 }
6352 return true;
6353 }
6354
6355 for (j, (h, t)) in pairs.iter().enumerate() {
6360 if self.mtp_step_graph(m, h, *t, first_pos + j).is_none() {
6361 return false;
6362 }
6363 }
6364 self.mtp_graph_mode = Some(true);
6365 true
6366 }
6367
6368 #[cfg(feature = "gpu")]
6373 fn mtp_warm_prefill_pairs(
6374 &mut self,
6375 m: &mut MtpModule,
6376 pairs: &[(&[f32], u32)],
6377 first_pos: usize,
6378 ) -> Result<(), &'static str> {
6379 if self.mtp_graph_mode == Some(false) || !self.mtp_graph_ok(m) {
6384 if self.mtp_graph_mode == Some(true) {
6385 return Err("MTP token graph became unavailable after admission");
6386 }
6387 self.mtp_graph_mode = Some(false);
6388 for (j, (h, t)) in pairs.iter().enumerate() {
6389 self.mtp_warm(m, h, *t, first_pos + j);
6390 }
6391 return Ok(());
6392 }
6393 match self.mtp_warm_graph(m, pairs, first_pos) {
6394 crate::gpu::BatchGraphOutcome::Completed => {
6395 if !pairs.is_empty() {
6396 self.mtp_graph_mode = Some(true);
6397 }
6398 Ok(())
6399 }
6400 crate::gpu::BatchGraphOutcome::Declined => {
6401 if self.mtp_warm_graph_fallback(m, pairs, first_pos) {
6402 Ok(())
6403 } else {
6404 Err("MTP warm-up fallback failed after device admission")
6405 }
6406 }
6407 crate::gpu::BatchGraphOutcome::Failed => {
6408 Err("MTP warm batch graph failed after admission")
6409 }
6410 }
6411 }
6412
6413 #[cfg(not(feature = "gpu"))]
6414 fn mtp_warm_prefill_pairs(
6415 &mut self,
6416 m: &mut MtpModule,
6417 pairs: &[(&[f32], u32)],
6418 first_pos: usize,
6419 ) -> Result<(), &'static str> {
6420 for (j, (h, t)) in pairs.iter().enumerate() {
6421 self.mtp_warm(m, h, *t, first_pos + j);
6422 }
6423 Ok(())
6424 }
6425
6426 fn mtp_warm(&mut self, m: &mut MtpModule, hidden: &[f32], next_token: u32, position: usize) {
6430 let e = self.embed_single(next_token);
6431 let mut cat = vec![0.0f32; 2 * self.hidden_size];
6432 let (cat_e, cat_h) = cat.split_at_mut(self.hidden_size);
6433 inference::rms_norm_into(&e, &m.enorm, self.rms_eps, self.norm_style, cat_e);
6434 inference::rms_norm_into(hidden, &m.hnorm, self.rms_eps, self.norm_style, cat_h);
6435 let mut x = vec![0.0f32; self.hidden_size];
6436 m.eh_proj.matvec(&cat, &mut x, self.pool.as_deref());
6437 inference::rms_norm_into(
6438 &x,
6439 &m.layer.input_norm,
6440 self.rms_eps,
6441 self.norm_style,
6442 &mut self.ws.n1,
6443 );
6444 let attn = match &m.layer.attn {
6445 AttnKind::Full {
6446 wq,
6447 wk,
6448 wv,
6449 wo,
6450 q_norm,
6451 k_norm,
6452 output_gate,
6453 softplus_gate,
6454 bias,
6455 } => {
6456 let mut cfg = self.attn_cfg(position);
6457 cfg.q_norm = q_norm.as_deref();
6458 cfg.k_norm = k_norm.as_deref();
6459 cfg.output_gate = *output_gate;
6460 cfg.softplus_gate = softplus_gate.as_ref().map(|(g, p)| (g, *p));
6461 cfg.bias = bias
6462 .as_ref()
6463 .map(|(q, k, v)| (q.as_slice(), k.as_slice(), v.as_slice()));
6464 attention::qwen_attention(&self.ws.n1, wq, wk, wv, wo, &mut m.kv, &cfg)
6465 }
6466 _ => return,
6467 };
6468 let _ = attn;
6469 }
6470
6471 #[cfg(feature = "gpu")]
6478 #[allow(clippy::too_many_arguments)]
6479 fn graph_spec_step(
6480 &mut self,
6481 m: &mut MtpModule,
6482 hidden: &[f32],
6483 t_next: u32,
6484 next_pos: usize,
6485 drafted: &mut usize,
6486 accepted: &mut usize,
6487 all_ids: &mut Vec<u32>,
6491 room: usize,
6496 ) -> Option<(Vec<u32>, usize, Vec<f32>)> {
6497 #[cfg(target_os = "macos")]
6508 let metal_native = crate::gpu::q1_force();
6509 #[cfg(not(target_os = "macos"))]
6510 let metal_native = false;
6511 #[cfg(feature = "gpu")]
6512 let k_default = if metal_native {
6513 7
6516 } else if crate::gpu_wgpu::verify_i8_on() {
6517 5
6518 } else {
6519 4
6520 };
6521 #[cfg(not(feature = "gpu"))]
6522 let k_default = 4;
6523 let k_env: Option<usize> = std::env::var("CMF_GRAPH_SPEC_K")
6524 .ok()
6525 .and_then(|v| v.parse().ok())
6526 .filter(|&v| (1..=8).contains(&v));
6527 let (k_start, k_max) = if metal_native { (7, 7) } else { (3, k_default.max(5)) };
6532 let k_full: usize = k_env.unwrap_or_else(|| self.spec_k_adapt.unwrap_or(k_start));
6533 let k_spec = k_full.min(room).max(1);
6534 let k_capped = k_spec < k_full;
6537 if next_pos == 0 {
6538 return None;
6539 }
6540 let t_round = std::time::Instant::now();
6541 let subs = || crate::gpu_wgpu::SUBMITS.load(std::sync::atomic::Ordering::Relaxed);
6557 let sub0 = subs();
6558 let cfg = self.sampler_config.clone();
6563 let penalized = !(cfg.repetition_penalty == 1.0
6564 && cfg.presence_penalty == 0.0
6565 && cfg.suppress_tokens.is_empty());
6566 let greedy_pen = cfg.temperature < 1e-6 && penalized;
6571 let sampling = cfg.temperature >= 1e-6;
6572 let sparse = sampling && sampler::sparse_ok(&cfg);
6578 let base_len = all_ids.len();
6579 if sampling && !sparse && self.spec_q.len() < k_spec {
6580 self.spec_q.resize_with(k_spec, Vec::new);
6581 }
6582 if sparse && self.spec_qs.len() < k_spec {
6583 self.spec_qs.resize_with(k_spec, Vec::new);
6584 }
6585 let mut drafts = Vec::with_capacity(k_spec);
6590 let mut hx = hidden.to_vec();
6591 let spec_dbg = std::env::var("CMF_SPEC_DBG").is_ok();
6594 spec_stamp("pro");
6595 #[cfg(target_os = "macos")]
6601 if metal_native && !sampling && !greedy_pen && self.mtp_graph_mode != Some(false) {
6602 match self.mtp_draft_chain_metal(m, hidden, t_next, next_pos - 1, k_spec) {
6603 Ok(ids) => {
6604 self.mtp_graph_mode = Some(true);
6605 drafts = ids;
6606 }
6607 Err(true) => {
6608 tracing::error!("mtp Metal draft chain failed after commit");
6609 self.clear_sequence_state();
6610 self.graph_failed
6611 .store(true, std::sync::atomic::Ordering::Relaxed);
6612 self.cancel
6613 .store(true, std::sync::atomic::Ordering::Relaxed);
6614 return None;
6615 }
6616 Err(false) => {}
6617 }
6618 }
6619 for j in drafts.len()..k_spec {
6620 let tok_in = if j == 0 { t_next } else { drafts[j - 1] };
6621 let mut dbg_ref: Option<(Vec<f32>, Vec<f32>)> = None;
6622 if spec_dbg {
6623 let saved = self.mtp_graph_mode;
6624 self.mtp_graph_mode = Some(false);
6625 let r = self.mtp_step_hl(m, &hx, tok_in, next_pos - 1 + j);
6626 self.mtp_graph_mode = saved;
6627 if self.graph_failed.load(std::sync::atomic::Ordering::Relaxed) {
6628 return None;
6629 }
6630 m.kv.truncate_last(1);
6631 dbg_ref = Some(r);
6632 }
6633 let (mut lg, hj) = self.mtp_step_hl(m, &hx, tok_in, next_pos - 1 + j);
6634 if self.graph_failed.load(std::sync::atomic::Ordering::Relaxed) {
6635 return None;
6636 }
6637 if let Some((lg_cpu, h_cpu)) = dbg_ref {
6638 let n = |v: &[f32]| v.iter().map(|x| x * x).sum::<f32>().sqrt();
6639 let dl = lg
6640 .iter()
6641 .zip(&lg_cpu)
6642 .fold(0f32, |m, (a, b)| m.max((a - b).abs()));
6643 let dh = hj
6644 .iter()
6645 .zip(&h_cpu)
6646 .fold(0f32, |m, (a, b)| m.max((a - b).abs()));
6647 eprintln!(
6648 "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 {}",
6649 next_pos - 1 + j,
6650 sampler::argmax(&lg_cpu),
6651 sampler::argmax(&lg),
6652 n(&h_cpu),
6653 n(&hj),
6654 m.kv.seq_len
6655 );
6656 }
6657 let dj = if sparse {
6658 let mut q = std::mem::take(&mut self.spec_qs[j]);
6659 let ok = sampler::sparse_distribution_into(
6660 &lg,
6661 &cfg,
6662 all_ids,
6663 &mut self.sampler_scratch,
6664 self.pool.as_deref(),
6665 &mut q,
6666 );
6667 let d = if ok {
6668 sampler::draw_sparse(&q, &mut self.rng)
6669 } else {
6670 let t = sampler::argmax(&lg);
6672 q.clear();
6673 q.push((t, 1.0));
6674 t
6675 };
6676 self.spec_qs[j] = q;
6677 all_ids.push(d);
6678 d
6679 } else if sampling {
6680 let mut q = std::mem::take(&mut self.spec_q[j]);
6681 sampler::distribution_into(
6682 &lg,
6683 &cfg,
6684 all_ids,
6685 &mut self.sampler_scratch,
6686 self.pool.as_deref(),
6687 &mut q,
6688 );
6689 let d = sampler::draw(&q, &mut self.rng);
6690 self.spec_q[j] = q;
6691 all_ids.push(d); d
6693 } else if greedy_pen {
6694 let d = sampler::argmax_penalized(
6695 &lg,
6696 &cfg,
6697 all_ids,
6698 &mut self.sampler_scratch,
6699 self.pool.as_deref(),
6700 );
6701 all_ids.push(d);
6702 d
6703 } else {
6704 sampler::argmax(&lg)
6705 };
6706 attention::recycle_buf(&mut lg);
6707 drafts.push(dj);
6708 hx = hj;
6709 spec_stamp("d.pick");
6710 }
6711 all_ids.truncate(base_len);
6712 *drafted += k_spec;
6713 let t_draft = t_round.elapsed();
6714 let sub_draft = subs();
6715 let b = k_spec + 1;
6718 let mut hiddens = vec![0.0f32; b * self.hidden_size];
6719 for (i, &t) in std::iter::once(&t_next).chain(drafts.iter()).enumerate() {
6720 let e = self.embed_single(t);
6721 hiddens[i * self.hidden_size..(i + 1) * self.hidden_size].copy_from_slice(&e);
6722 }
6723 let positions: Vec<usize> = (next_pos..next_pos + b).collect();
6724 spec_stamp("v.emb");
6725 let (lm_gw, lm_rows) = {
6726 let (_, i, kind, rs) = self.weights.lm_head.graph_weight()?;
6727 (
6728 crate::gpu::GraphW {
6729 idx: i,
6730 kind,
6731 row_scale: rs,
6732 data: &[],
6733 prism: crate::gpu::GraphPrismOp::None,
6734 affine: false,
6735 },
6736 self.weights.lm_head.rows(),
6737 )
6738 };
6739 let mut logits = Vec::new();
6740 let final_norm = self.weights.final_norm.clone();
6741 #[cfg(target_os = "macos")]
6750 let greedy_dev = metal_native
6751 && !sampling
6752 && !greedy_pen
6753 && !self.confidence_on
6754 && self.final_softcap.is_none()
6755 && self.vocab_size == lm_rows
6761 && std::env::var_os("CMF_METAL_VERIFY_CHECK").is_none()
6762 && std::env::var_os("CMF_LOGIT_DUMP").is_none()
6763 && std::env::var("CMF_METAL_DEV_ARGMAX").as_deref() != Ok("0");
6764 #[cfg(not(target_os = "macos"))]
6765 let greedy_dev = false;
6766 let mut dev_ids: Vec<u32> = Vec::new();
6767 #[cfg(target_os = "macos")]
6768 let verify_outcome = if metal_native {
6769 let lm = self.weights.lm_head.q1_parts()?;
6770 let n_score = self.vocab_size.min(lm_rows);
6771 self.try_batch_graph_metal(
6772 &mut hiddens,
6773 &positions,
6774 b,
6775 Some((lm, &final_norm, &mut logits)),
6776 if greedy_dev {
6777 Some((n_score, &mut dev_ids))
6778 } else {
6779 None
6780 },
6781 )
6782 } else {
6783 self.try_batch_graph_wgpu(
6784 &mut hiddens,
6785 &positions,
6786 b,
6787 Some(crate::gpu::SpecTail {
6788 lm: lm_gw,
6789 lm_rows,
6790 final_norm: &final_norm,
6791 logits_out: &mut logits,
6792 }),
6793 )
6794 };
6795 #[cfg(not(target_os = "macos"))]
6796 let verify_outcome = self.try_batch_graph_wgpu(
6797 &mut hiddens,
6798 &positions,
6799 b,
6800 Some(crate::gpu::SpecTail {
6801 lm: lm_gw,
6802 lm_rows,
6803 final_norm: &final_norm,
6804 logits_out: &mut logits,
6805 }),
6806 );
6807 match verify_outcome {
6808 crate::gpu::BatchGraphOutcome::Completed => {}
6809 crate::gpu::BatchGraphOutcome::Declined => {
6810 m.kv.truncate_last(k_spec);
6814 if !metal_native && !self.rewind_mtp_graph_mirror(next_pos) {
6815 self.clear_sequence_state();
6816 self.graph_failed
6817 .store(true, std::sync::atomic::Ordering::Relaxed);
6818 self.cancel
6819 .store(true, std::sync::atomic::Ordering::Relaxed);
6820 tracing::error!("MTP graph mirror rewind failed after verify decline");
6821 }
6822 return None;
6823 }
6824 crate::gpu::BatchGraphOutcome::Failed => {
6825 self.clear_sequence_state();
6829 self.graph_failed
6830 .store(true, std::sync::atomic::Ordering::Relaxed);
6831 self.cancel
6832 .store(true, std::sync::atomic::Ordering::Relaxed);
6833 tracing::error!("MTP verify batch graph failed after admission");
6834 return None;
6835 }
6836 }
6837 #[cfg(target_os = "macos")]
6843 if metal_native && std::env::var("CMF_METAL_VERIFY_CHECK").as_deref() == Ok("1") {
6844 let snap: Vec<Vec<f32>> = self
6845 .kv_cache
6846 .layers
6847 .iter()
6848 .map(|l| l.linear_state.clone())
6849 .collect();
6850 let attn_lens: Vec<usize> = self.kv_cache.layers.iter().map(|l| l.seq_len).collect();
6851 let toks: Vec<u32> = std::iter::once(t_next)
6852 .chain(drafts.iter().copied())
6853 .collect();
6854 let want_save = self.graph_want_logits;
6855 self.graph_want_logits = false;
6856 for (i, &t) in toks.iter().enumerate() {
6857 let hi = self.forward_layers(&self.embed_single(t), next_pos + i, None);
6858 let _ = self.graph_logits.take();
6859 if std::env::var("CMF_SPEC_PLAIN_HIDDEN").as_deref() == Ok("1") {
6863 hiddens[i * self.hidden_size..(i + 1) * self.hidden_size].copy_from_slice(&hi);
6864 }
6865 let ref_lg = self.logits_from_hidden(&hi);
6866 let row = &logits[i * lm_rows..(i + 1) * lm_rows];
6867 let ra = sampler::argmax(&ref_lg);
6868 let va = sampler::argmax(row);
6869 let mut md = 0f32;
6870 let mut rms = 0f64;
6871 for j in 0..lm_rows.min(ref_lg.len()) {
6872 let d = (ref_lg[j] - row[j]).abs();
6873 md = md.max(d);
6874 rms += (d as f64) * (d as f64);
6875 }
6876 let mut hd = 0f32;
6877 for j in 0..self.hidden_size {
6878 hd = hd.max((hi[j] - hiddens[i * self.hidden_size + j]).abs());
6879 }
6880 eprintln!(
6881 "verify-check row {i} tok {t} pos {}: ref argmax {ra} verify argmax {va} {} | max|dlogit| {md:.3} rms {:.4} | max|dhidden| {hd:.4}",
6882 next_pos + i,
6883 if ra == va { "OK" } else { "MISMATCH" },
6884 (rms / lm_rows as f64).sqrt()
6885 );
6886 }
6887 self.graph_want_logits = want_save;
6888 for (l, st) in self.kv_cache.layers.iter_mut().zip(snap) {
6891 if l.linear_state.len() == st.len() {
6892 l.linear_state.copy_from_slice(&st);
6893 } else {
6894 l.linear_state = st;
6895 }
6896 }
6897 for (li, (l, n0)) in self.kv_cache.layers.iter_mut().zip(attn_lens).enumerate() {
6898 let extra = l.seq_len.saturating_sub(n0);
6899 if extra > 0 {
6900 l.truncate_last(extra);
6901 crate::gpu_metal::kv_mirror_set_stored(self.graph_kv_id, li, n0);
6902 }
6903 }
6904 }
6905 let t_verify = t_round.elapsed();
6906 let sub_verify = subs();
6907 let mut a = 0usize;
6912 let mut forced: Option<u32> = None;
6913 let ids: Vec<u32> = if sparse {
6914 let mut p = std::mem::take(&mut self.spec_ps);
6915 let mut res = std::mem::take(&mut self.spec_ress);
6916 while a < k_spec {
6917 let ok = sampler::sparse_distribution_into(
6918 &logits[a * lm_rows..(a + 1) * lm_rows],
6919 &cfg,
6920 all_ids,
6921 &mut self.sampler_scratch,
6922 self.pool.as_deref(),
6923 &mut p,
6924 );
6925 if !ok {
6926 let t = sampler::argmax(&logits[a * lm_rows..(a + 1) * lm_rows]);
6927 p.clear();
6928 p.push((t, 1.0));
6929 }
6930 match sampler::spec_accept_or_correct_sparse(
6931 &p,
6932 &self.spec_qs[a],
6933 drafts[a],
6934 &mut self.rng,
6935 &mut res,
6936 ) {
6937 None => {
6938 all_ids.push(drafts[a]);
6939 a += 1;
6940 }
6941 Some(c) => {
6942 forced = Some(c);
6943 break;
6944 }
6945 }
6946 }
6947 all_ids.truncate(base_len);
6948 self.spec_ps = p;
6949 self.spec_ress = res;
6950 drafts.clone()
6951 } else if sampling {
6952 let mut p = std::mem::take(&mut self.spec_p);
6953 let mut res = std::mem::take(&mut self.spec_res);
6954 while a < k_spec {
6955 sampler::distribution_into(
6956 &logits[a * lm_rows..(a + 1) * lm_rows],
6957 &cfg,
6958 all_ids,
6959 &mut self.sampler_scratch,
6960 self.pool.as_deref(),
6961 &mut p,
6962 );
6963 match sampler::spec_accept_or_correct(
6964 &p,
6965 &self.spec_q[a],
6966 drafts[a],
6967 &mut self.rng,
6968 &mut res,
6969 self.pool.as_deref(),
6970 ) {
6971 None => {
6972 all_ids.push(drafts[a]);
6973 a += 1;
6974 }
6975 Some(c) => {
6976 forced = Some(c);
6977 break;
6978 }
6979 }
6980 }
6981 all_ids.truncate(base_len);
6982 self.spec_p = p;
6983 self.spec_res = res;
6984 drafts.clone()
6986 } else if greedy_pen {
6987 let mut ids: Vec<u32> = Vec::with_capacity(b);
6991 for i in 0..b {
6992 let t = sampler::argmax_penalized(
6993 &logits[i * lm_rows..(i + 1) * lm_rows],
6994 &cfg,
6995 all_ids,
6996 &mut self.sampler_scratch,
6997 self.pool.as_deref(),
6998 );
6999 ids.push(t);
7000 if i < k_spec && t == drafts[i] {
7001 all_ids.push(t);
7002 } else {
7003 break;
7004 }
7005 }
7006 all_ids.truncate(base_len);
7007 while a < k_spec && a < ids.len() && ids[a] == drafts[a] {
7008 a += 1;
7009 }
7010 ids
7013 } else if greedy_dev && dev_ids.len() == b {
7014 let ids = std::mem::take(&mut dev_ids);
7015 while a < k_spec && ids[a] == drafts[a] {
7016 a += 1;
7017 }
7018 ids
7019 } else {
7020 if logits.len() < b * lm_rows {
7021 self.clear_sequence_state();
7024 self.graph_failed
7025 .store(true, std::sync::atomic::Ordering::Relaxed);
7026 self.cancel
7027 .store(true, std::sync::atomic::Ordering::Relaxed);
7028 tracing::error!("Metal verify returned neither logits nor argmax ids");
7029 return None;
7030 }
7031 let ids: Vec<u32> = (0..b)
7032 .map(|i| sampler::argmax(&logits[i * lm_rows..(i + 1) * lm_rows]))
7033 .collect();
7034 while a < k_spec && ids[a] == drafts[a] {
7035 a += 1;
7036 }
7037 ids
7038 };
7039 spec_stamp("acc");
7040 if spec_dbg {
7041 eprintln!(
7042 "spec-dbg round: t_next {t_next} drafts {:?} verified {:?} accepted {a}",
7043 drafts, ids
7044 );
7045 }
7046 #[cfg(target_os = "macos")]
7050 let commit_ref: Option<(Vec<Vec<f32>>, Vec<(usize, Vec<f32>, Vec<f32>)>)> = if metal_native
7051 && std::env::var("CMF_METAL_VERIFY_CHECK").as_deref() == Ok("2")
7052 {
7053 let snap: Vec<Vec<f32>> = self
7054 .kv_cache
7055 .layers
7056 .iter()
7057 .map(|l| l.linear_state.clone())
7058 .collect();
7059 let attn_lens: Vec<usize> = self.kv_cache.layers.iter().map(|l| l.seq_len).collect();
7060 let toks: Vec<u32> = std::iter::once(t_next)
7061 .chain(drafts.iter().copied())
7062 .collect();
7063 let want_save = self.graph_want_logits;
7064 self.graph_want_logits = false;
7065 for (i, &t) in toks.iter().take(a + 1).enumerate() {
7066 let _ = self.forward_layers(&self.embed_single(t), next_pos + i, None);
7067 let _ = self.graph_logits.take();
7068 }
7069 self.graph_want_logits = want_save;
7070 let plain_states: Vec<Vec<f32>> = self
7071 .kv_cache
7072 .layers
7073 .iter()
7074 .map(|l| l.linear_state.clone())
7075 .collect();
7076 let (nkv, hd) = (self.num_kv_heads, self.head_dim);
7077 let mut rows = Vec::new();
7078 for (li, (l, n0)) in self
7079 .kv_cache
7080 .layers
7081 .iter_mut()
7082 .zip(attn_lens.iter())
7083 .enumerate()
7084 {
7085 let extra = l.seq_len.saturating_sub(*n0);
7086 if extra > 0 {
7087 let mut kk = Vec::new();
7088 let mut vv = Vec::new();
7089 for g in 0..nkv {
7090 kk.extend_from_slice(&l.head_keys(g)[n0 * hd..]);
7091 vv.extend_from_slice(&l.head_values(g)[n0 * hd..]);
7092 }
7093 rows.push((li, kk, vv));
7094 l.truncate_last(extra);
7095 crate::gpu_metal::kv_mirror_set_stored(self.graph_kv_id, li, *n0);
7096 }
7097 }
7098 for (l, st) in self.kv_cache.layers.iter_mut().zip(snap) {
7099 if l.linear_state.len() == st.len() {
7100 l.linear_state.copy_from_slice(&st);
7101 } else {
7102 l.linear_state = st;
7103 }
7104 }
7105 Some((plain_states, rows))
7106 } else {
7107 None
7108 };
7109 let warm_off = std::env::var("CMF_SPEC_WARM").is_ok_and(|v| v == "0");
7110 #[cfg(target_os = "macos")]
7119 let mut warm_pending: Option<MetalWarmPending> = None;
7120 #[cfg(target_os = "macos")]
7121 if metal_native {
7122 m.kv.truncate_last(k_spec.saturating_sub(1));
7123 if self.mtp_graph_mode == Some(true) {
7124 crate::gpu_metal::kv_mirror_set_stored(
7127 self.mtp_kv_id(),
7128 Self::MTP_LAYER_BASE,
7129 m.kv.seq_len,
7130 );
7131 if !warm_off && a > 0 {
7132 let pairs: Vec<(&[f32], u32)> = (0..a)
7133 .map(|j| {
7134 (
7135 &hiddens[j * self.hidden_size..(j + 1) * self.hidden_size],
7136 ids[j],
7137 )
7138 })
7139 .collect();
7140 warm_pending = self.mtp_warm_batch_submit(m, &pairs, next_pos);
7141 }
7142 }
7143 spec_stamp("c.wsub");
7144 }
7145 #[cfg(target_os = "macos")]
7147 if metal_native {
7148 if !self.metal_verify_commit(a) {
7151 self.clear_sequence_state();
7152 self.graph_failed
7153 .store(true, std::sync::atomic::Ordering::Relaxed);
7154 self.cancel
7155 .store(true, std::sync::atomic::Ordering::Relaxed);
7156 tracing::error!("Metal verify state/KV handoff failed after admission");
7157 return None;
7158 }
7159 if let Some((plain_states, rows)) = commit_ref {
7160 crate::gpu_metal::queue_fence();
7161 let _ = crate::gpu_metal::wait_replay();
7164 let (nkv, hd) = (self.num_kv_heads, self.head_dim);
7165 let mut worst_s = 0f32;
7166 let mut worst_li = 0usize;
7167 for (li, (l, ps)) in self.kv_cache.layers.iter().zip(&plain_states).enumerate() {
7168 if l.linear_state.len() != ps.len() || ps.is_empty() {
7169 continue;
7170 }
7171 let d = l
7172 .linear_state
7173 .iter()
7174 .zip(ps)
7175 .fold(0f32, |m, (x, y)| m.max((x - y).abs()));
7176 let n = ps.iter().fold(0f32, |m, y| m.max(y.abs()));
7177 let rel = d / n.max(1e-6);
7178 if rel > worst_s {
7179 worst_s = rel;
7180 worst_li = li;
7181 }
7182 }
7183 let mut worst_k = 0f32;
7184 for (li, kk, vv) in &rows {
7185 let l = &self.kv_cache.layers[*li];
7186 let n0 = l.seq_len - (kk.len() / (nkv * hd));
7187 let mut ck = Vec::new();
7188 let mut cv = Vec::new();
7189 for g in 0..nkv {
7190 ck.extend_from_slice(&l.head_keys(g)[n0 * hd..]);
7191 cv.extend_from_slice(&l.head_values(g)[n0 * hd..]);
7192 }
7193 if ck.len() == kk.len() {
7194 let dk = ck
7195 .iter()
7196 .zip(kk)
7197 .fold(0f32, |m, (x, y)| m.max((x - y).abs()));
7198 let dv = cv
7199 .iter()
7200 .zip(vv)
7201 .fold(0f32, |m, (x, y)| m.max((x - y).abs()));
7202 worst_k = worst_k.max(dk).max(dv);
7203 } else {
7204 eprintln!(
7205 "commit-check L{li}: kv row count mismatch {} vs {}",
7206 ck.len(),
7207 kk.len()
7208 );
7209 }
7210 }
7211 eprintln!(
7212 "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}"
7213 );
7214 }
7215 }
7216 if !metal_native && a + 1 < b {
7217 let expected_gdn_layers = self.graph_gdn_layer_count();
7218 if expected_gdn_layers > 0
7219 && !crate::gpu::gdn_spec_restore(self.graph_kv_id, a, next_pos, expected_gdn_layers)
7220 {
7221 self.clear_sequence_state();
7222 self.graph_failed
7223 .store(true, std::sync::atomic::Ordering::Relaxed);
7224 self.cancel
7225 .store(true, std::sync::atomic::Ordering::Relaxed);
7226 tracing::error!("GDN speculative restore failed after verify");
7227 return None;
7228 }
7229 }
7230 if !metal_native && !self.rewind_trunk_graph_mirrors(next_pos + a + 1) {
7231 self.clear_sequence_state();
7236 self.graph_failed
7237 .store(true, std::sync::atomic::Ordering::Relaxed);
7238 self.cancel
7239 .store(true, std::sync::atomic::Ordering::Relaxed);
7240 tracing::error!("trunk graph KV rewind failed after speculative verify");
7241 return None;
7242 }
7243 *accepted += a;
7244 if !metal_native {
7255 m.kv.truncate_last(k_spec.saturating_sub(1));
7257 }
7258 spec_stamp("c.trunc");
7259 if !metal_native
7260 && self.mtp_graph_mode == Some(true)
7261 && !self.rewind_mtp_graph_mirror(next_pos)
7262 {
7263 self.clear_sequence_state();
7267 self.graph_failed
7268 .store(true, std::sync::atomic::Ordering::Relaxed);
7269 self.cancel
7270 .store(true, std::sync::atomic::Ordering::Relaxed);
7271 tracing::error!("MTP graph mirror rewind failed after verify commit");
7272 return None;
7273 }
7274 if !warm_off && a > 0 {
7275 let mut warmed = false;
7278 #[cfg(target_os = "macos")]
7279 if metal_native && self.mtp_graph_mode == Some(true) {
7280 warmed = match warm_pending.take() {
7284 Some(p) => self.mtp_warm_batch_finish(m, p),
7285 None => false,
7286 };
7287 if !warmed {
7288 warmed = true;
7289 for j in 0..a {
7290 let row =
7291 hiddens[j * self.hidden_size..(j + 1) * self.hidden_size].to_vec();
7292 if self
7293 .mtp_step_metal(m, &row, ids[j], next_pos + j, false)
7294 .is_none()
7295 {
7296 warmed = false;
7297 break;
7298 }
7299 }
7300 }
7301 }
7302 if !warmed && self.mtp_graph_mode != Some(false) && !metal_native {
7303 let rows: Vec<Vec<f32>> = (0..a)
7304 .map(|j| hiddens[j * self.hidden_size..(j + 1) * self.hidden_size].to_vec())
7305 .collect();
7306 let pairs: Vec<(&[f32], u32)> = rows
7307 .iter()
7308 .zip(ids.iter())
7309 .map(|(r, &t)| (r.as_slice(), t))
7310 .collect();
7311 match self.mtp_warm_prefill_pairs(m, &pairs, next_pos) {
7312 Ok(()) => warmed = true,
7313 Err(err) => {
7314 tracing::error!("{err}");
7320 self.clear_sequence_state();
7321 self.graph_failed
7322 .store(true, std::sync::atomic::Ordering::Relaxed);
7323 self.cancel
7324 .store(true, std::sync::atomic::Ordering::Relaxed);
7325 return None;
7326 }
7327 }
7328 }
7329 if !warmed {
7330 for j in 0..a {
7331 let row = &hiddens[j * self.hidden_size..(j + 1) * self.hidden_size];
7332 let row = row.to_vec();
7333 self.mtp_warm(m, &row, ids[j], next_pos + j);
7334 }
7335 }
7336 }
7337 spec_stamp("c.warm");
7341 if let Some(c) = forced {
7342 self.spec_forced = Some(c);
7343 self.graph_logits = None;
7344 } else if greedy_dev && logits.is_empty() {
7345 self.spec_forced = Some(ids[a]);
7348 self.graph_logits = None;
7349 } else {
7350 let mut row = logits[a * lm_rows..(a + 1) * lm_rows].to_vec();
7351 row.resize(self.vocab_size, 0.0);
7352 if let Some(c) = self.final_softcap {
7353 for l in row.iter_mut() {
7354 *l = c * (*l / c).tanh();
7355 }
7356 }
7357 self.graph_logits = Some(row);
7358 }
7359 let new_hidden = hiddens[a * self.hidden_size..(a + 1) * self.hidden_size].to_vec();
7360 spec_stamp("c.row");
7361 if std::env::var("CMF_GRAPH_SPEC_TIME").is_ok() {
7367 let end = subs();
7368 eprintln!(
7369 "spec-round: draft {:.1} ms/{} sub | verify {:.1} ms/{} sub | \
7370 commit {:.1} ms/{} sub (accepted {a} of {k_spec}, full-head streak {})",
7371 t_draft.as_secs_f64() * 1e3,
7372 sub_draft - sub0,
7373 (t_verify - t_draft).as_secs_f64() * 1e3,
7374 sub_verify - sub_draft,
7375 (t_round.elapsed() - t_verify).as_secs_f64() * 1e3,
7376 end - sub_verify,
7377 self.draft_full_streak,
7378 );
7379 }
7380 if k_env.is_none() && !metal_native && !k_capped {
7385 let f = a as f32 / k_spec.max(1) as f32;
7389 self.spec_acc_ewma += 0.2 * (f - self.spec_acc_ewma);
7390 let mut k_next = k_spec;
7391 if self.spec_acc_ewma >= 0.75 && k_spec < k_max {
7392 k_next = k_spec + 1;
7393 } else if self.spec_acc_ewma < 0.4 && k_spec > 2 {
7394 k_next = k_spec - 1;
7395 }
7396 if k_next != k_spec {
7397 self.spec_acc_ewma = 0.6;
7398 if std::env::var("CMF_GRAPH_SPEC_TIME").is_ok() {
7399 eprintln!("spec-k: {k_spec} → {k_next}");
7400 }
7401 }
7402 self.spec_k_adapt = Some(k_next);
7403 }
7404 spec_stamp("end");
7405 Some((drafts[..a].to_vec(), next_pos + a + 1, new_hidden))
7406 }
7407
7408 pub fn measure_pair_fusion(&mut self, iters: usize) -> (f64, f64) {
7417 if !self.pair_supported() {
7418 return (0.0, 0.0);
7419 }
7420 let graph_env = std::env::var_os("CMF_GPU_WGPU_GRAPH");
7427 unsafe { std::env::set_var("CMF_GPU_WGPU_GRAPH", "0") };
7428 let emb1 = self.embed_single(1);
7429 let emb2 = self.embed_single(2);
7430 let pos = self.kv_cache.seq_len();
7431
7432 let t0 = std::time::Instant::now();
7433 for _ in 0..iters {
7434 let _ = self.forward_layers(&emb1, pos, None);
7435 let _ = self.forward_layers(&emb2, pos + 1, None);
7436 for l in &mut self.kv_cache.layers {
7437 l.truncate_last(2);
7438 }
7439 }
7440 let singles_ms = t0.elapsed().as_secs_f64() * 1000.0 / iters as f64;
7441
7442 let t1 = std::time::Instant::now();
7443 for _ in 0..iters {
7444 let _ = self.forward_pair(&emb1, &emb2, pos);
7445 for l in &mut self.kv_cache.layers {
7446 l.truncate_last(2);
7447 }
7448 }
7449 let pair_ms = t1.elapsed().as_secs_f64() * 1000.0 / iters as f64;
7450 match graph_env {
7451 Some(value) => unsafe { std::env::set_var("CMF_GPU_WGPU_GRAPH", value) },
7452 None => unsafe { std::env::remove_var("CMF_GPU_WGPU_GRAPH") },
7453 }
7454 (singles_ms, pair_ms)
7455 }
7456
7457 fn pair_supported(&self) -> bool {
7465 !self.weights.layers.is_empty()
7472 && self.g3n.is_none()
7473 && !self
7474 .weights
7475 .layers
7476 .iter()
7477 .any(|lw| matches!(&lw.attn, AttnKind::Mla(_) | AttnKind::Kda(_)))
7478 }
7479
7480 fn forward_pair(
7481 &mut self,
7482 emb1: &[f32],
7483 emb2: &[f32],
7484 position: usize,
7485 ) -> (Vec<f32>, Vec<f32>) {
7486 self.mimo_moe_prepare();
7489 let mut h1 = emb1.to_vec();
7490 let mut h2 = emb2.to_vec();
7491 let (_nkv, _hd, hs, _rd, eps) = (
7492 self.num_kv_heads,
7493 self.head_dim,
7494 self.hidden_size,
7495 self.rotary_dim,
7496 self.rms_eps,
7497 );
7498 let pool = self.pool.clone();
7499
7500 for li in 0..self.num_layers {
7501 let lw = &self.weights.layers[self.phys_layer(li)];
7502 inference::rms_norm_into(
7505 &h1,
7506 &lw.input_norm,
7507 self.rms_eps,
7508 self.norm_style,
7509 &mut self.ws.n1,
7510 );
7511 inference::rms_norm_into(
7512 &h2,
7513 &lw.input_norm,
7514 self.rms_eps,
7515 self.norm_style,
7516 &mut self.ws.n2,
7517 );
7518
7519 let (a1, a2) = match &lw.attn {
7520 AttnKind::Mla(_) => unreachable!("MLA has no MTP/pair path"),
7521 AttnKind::Kda(_) => unreachable!("KDA has no MTP/pair path"),
7522 AttnKind::Bounded(w) => {
7523 let rope = self
7526 .bounded_rope
7527 .clone()
7528 .expect("bounded layer without an installed rotation table");
7529 let cfg = crate::bounded::BoundedAttnCfg {
7530 num_heads: self.num_heads,
7531 num_kv_heads: self.num_kv_heads,
7532 head_dim: self.head_dim,
7533 hidden_size: hs,
7534 scale: self.attn_scale,
7535 rope: &rope,
7536 pool: pool.as_deref(),
7537 };
7538 let a1 = crate::bounded::bounded_attention(
7539 &self.ws.n1,
7540 w,
7541 &mut self.kv_cache.layers[li],
7542 &cfg,
7543 );
7544 let a2 = crate::bounded::bounded_attention(
7545 &self.ws.n2,
7546 w,
7547 &mut self.kv_cache.layers[li],
7548 &cfg,
7549 );
7550 (a1, a2)
7551 }
7552 AttnKind::Linear(w) => {
7553 let cfg = self.vmf_cfg.expect("linear layer without vmf_cfg");
7554 let layer = &mut self.kv_cache.layers[li];
7555 let (state, scratch) = (&mut layer.linear_state, &mut layer.linear_scratch);
7556 vmf_phase_pair(
7557 &self.ws.n1,
7558 &self.ws.n2,
7559 w,
7560 &cfg,
7561 state,
7562 scratch,
7563 self.pool.as_deref(),
7564 )
7565 }
7566 AttnKind::LinearGdn(w) => {
7567 let cfg = self.gdn_cfg.expect("gdn layer without gdn_cfg");
7568 let layer = &mut self.kv_cache.layers[li];
7569 let (state, scratch) = (&mut layer.linear_state, &mut layer.linear_scratch);
7570 gdn_pair(
7571 &self.ws.n1,
7572 &self.ws.n2,
7573 w,
7574 &cfg,
7575 state,
7576 scratch,
7577 self.pool.as_deref(),
7578 )
7579 }
7580 AttnKind::ShortConv(w) => {
7581 let cfg = self
7582 .short_conv_cfg
7583 .expect("short-conv layer without short_conv_cfg");
7584 let layer = &mut self.kv_cache.layers[li];
7585 let (state, scratch) = (&mut layer.linear_state, &mut layer.linear_scratch);
7586 short_conv_pair(
7587 &self.ws.n1,
7588 &self.ws.n2,
7589 w,
7590 &cfg,
7591 state,
7592 scratch,
7593 self.pool.as_deref(),
7594 )
7595 }
7596 AttnKind::Full {
7597 wq,
7598 wk,
7599 wv,
7600 wo,
7601 q_norm,
7602 k_norm,
7603 output_gate,
7604 softplus_gate,
7605 bias,
7606 } => {
7607 let inv_freq_l = self.layer_inv_freq(li);
7608 let (nkv_l, hd_l, rd_l) = self.layer_geom(li);
7609 let cfg = QwenAttnCfg {
7610 num_heads: self.layer_num_heads(li),
7611 num_kv_heads: nkv_l,
7612 head_dim: hd_l,
7613 hidden_size: hs,
7614 position,
7615 inv_freq: &inv_freq_l,
7616 rotary_dim: rd_l,
7617 scale: self.attn_scale,
7618 softcap: self.attn_softcap,
7619 window: self.layer_window(li),
7620 v_norm: self.attn_v_norm,
7621 qk_norm_after_rope: self.qk_norm_after_rope,
7622 q_norm: q_norm.as_deref(),
7623 k_norm: k_norm.as_deref(),
7624 output_gate: *output_gate,
7625 softplus_gate: softplus_gate
7626 .as_ref()
7627 .map(|(gate, per_head)| (gate, *per_head)),
7628 rope_scale: self.layer_rope_scale(li),
7629 bias: bias
7630 .as_ref()
7631 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
7632 rms_eps: eps,
7633 norm_style: self.norm_style,
7634 pool: pool.as_deref(),
7635 v_head_dim: self.layer_v_dim(li),
7636 };
7637 attention::qwen_attention_pair(
7638 &self.ws.n1,
7639 &self.ws.n2,
7640 wq,
7641 wk,
7642 wv,
7643 wo,
7644 &mut self.kv_cache.layers[li],
7645 &cfg,
7646 )
7647 }
7648 };
7649 let (a1, a2) = match &self.weights.layers[self.phys_layer(li)].attn_out_norm {
7650 Some(w) => (
7651 inference::rms_norm(&a1, w, self.rms_eps, self.norm_style),
7652 inference::rms_norm(&a2, w, self.rms_eps, self.norm_style),
7653 ),
7654 None => (a1, a2),
7655 };
7656 for i in 0..self.hidden_size {
7657 h1[i] += a1[i];
7658 h2[i] += a2[i];
7659 }
7660 let (mut a1, mut a2) = (a1, a2);
7661 attention::recycle_buf(&mut a1);
7662 attention::recycle_buf(&mut a2);
7663
7664 let lw = &self.weights.layers[self.phys_layer(li)];
7665 inference::rms_norm_into(
7666 &h1,
7667 &lw.post_norm,
7668 self.rms_eps,
7669 self.norm_style,
7670 &mut self.ws.p1,
7671 );
7672 inference::rms_norm_into(
7673 &h2,
7674 &lw.post_norm,
7675 self.rms_eps,
7676 self.norm_style,
7677 &mut self.ws.p2,
7678 );
7679 let (f1, f2) = match &lw.ffn {
7680 FfnKind::DenseMoe(dm) => (
7683 dense_moe_ffn(
7684 dm,
7685 &self.ws.p1,
7686 &h1,
7687 self.rms_eps,
7688 self.norm_style,
7689 self.pool.as_deref(),
7690 ),
7691 dense_moe_ffn(
7692 dm,
7693 &self.ws.p2,
7694 &h2,
7695 self.rms_eps,
7696 self.norm_style,
7697 self.pool.as_deref(),
7698 ),
7699 ),
7700 FfnKind::Moe(m) if self.mimo_moe.is_dynamic(li, false) => (
7701 moe_ffn_banked(&mut self.mimo_moe, li, m, &self.ws.p1, self.pool.as_deref()),
7702 moe_ffn_banked(&mut self.mimo_moe, li, m, &self.ws.p2, self.pool.as_deref()),
7703 ),
7704 _ => ffn_forward_pair(
7705 &lw.ffn,
7706 &self.ws.p1,
7707 &self.ws.p2,
7708 self.pool.as_deref(),
7709 None,
7710 ),
7711 };
7712 let (f1, f2) = match &self.weights.layers[self.phys_layer(li)].ffn_out_norm {
7713 Some(w) => (
7714 inference::rms_norm(&f1, w, self.rms_eps, self.norm_style),
7715 inference::rms_norm(&f2, w, self.rms_eps, self.norm_style),
7716 ),
7717 None => (f1, f2),
7718 };
7719 for i in 0..self.hidden_size {
7720 h1[i] += f1[i];
7721 h2[i] += f2[i];
7722 }
7723 let (mut f1, mut f2) = (f1, f2);
7724 attention::recycle_buf(&mut f1);
7725 attention::recycle_buf(&mut f2);
7726 if let Some(sc) = self.weights.layers[self.phys_layer(li)].layer_scale {
7727 for i in 0..self.hidden_size {
7728 h1[i] *= sc;
7729 h2[i] *= sc;
7730 }
7731 }
7732 if self.is_loop_end(li) && li + 1 < self.num_layers {
7734 h1 = inference::rms_norm(
7735 &h1,
7736 &self.weights.final_norm,
7737 self.rms_eps,
7738 self.norm_style,
7739 );
7740 h2 = inference::rms_norm(
7741 &h2,
7742 &self.weights.final_norm,
7743 self.rms_eps,
7744 self.norm_style,
7745 );
7746 }
7747 }
7748 if self.o1_active() {
7754 self.commit_linear_scratch();
7755 }
7756 self.o1_progress();
7757 (h1, h2)
7758 }
7759
7760 fn commit_linear_scratch(&mut self) {
7762 for layer in &mut self.kv_cache.layers {
7763 if !layer.linear_scratch.is_empty() {
7764 std::mem::swap(&mut layer.linear_state, &mut layer.linear_scratch);
7765 layer.linear_scratch.clear();
7766 }
7767 }
7768 }
7769
7770 pub fn forward_ids(
7773 &mut self,
7774 ids: &[u32],
7775 task_mask: Option<&TaskMask>,
7776 ) -> Result<Vec<f32>, String> {
7777 #[cfg(target_os = "macos")]
7778 crate::gpu_metal::set_io_namespace(self.graph_kv_id);
7779 if ids.is_empty() {
7780 return Err("empty id sequence".to_string());
7781 }
7782 self.clear_sequence_state();
7783 self.check_forward_graph("forward_ids setup", 0)?;
7784 if task_mask.is_none() {
7785 self.o1_begin();
7786 }
7787 let mut hidden = vec![0.0f32; self.hidden_size];
7788 let mut pos = 0usize;
7789 if let Some(b) = &mut self.dsv41 {
7790 let pool = self.pool.clone();
7791 let mut logits = Vec::new();
7792 crate::dsv41::forward_chunk(
7793 &b.0,
7794 &b.1,
7795 &b.2,
7796 &mut b.3,
7797 ids,
7798 0,
7799 pool.as_deref(),
7800 &mut logits,
7801 );
7802 if let Err(err) = self.o1_seal_checked() {
7803 self.clear_sequence_state();
7804 return Err(err);
7805 }
7806 return Ok(logits);
7807 }
7808 if self.can_prefill_batched() && !self.graph_prefill_preferred() && ids.len() > 2 {
7816 let chunk = self.prefill_chunk();
7820 let hs = self.hidden_size;
7821 while pos < ids.len() {
7822 let end = (pos + chunk).min(ids.len());
7823 let hb = self.prefill_rows(&ids[pos..end], pos, task_mask)?;
7824 self.check_forward_graph("forward_ids batched prefill", end - 1)?;
7825 hidden.copy_from_slice(&hb[(end - pos - 1) * hs..]);
7826 pos = end;
7827 }
7828 }
7829 if task_mask.is_none()
7838 && !self.graph_prefill_preferred()
7839 && !std::env::var("CMF_PAIR").is_ok_and(|v| v == "0")
7840 && self.pair_supported()
7841 {
7842 while pos + 1 < ids.len() {
7843 let e1 = self.embed_single(ids[pos]);
7844 let e2 = self.embed_single(ids[pos + 1]);
7845 let (_, h2) = self.forward_pair(&e1, &e2, pos);
7846 self.check_forward_graph("forward_ids pair", pos + 1)?;
7847 self.commit_linear_scratch();
7848 hidden = h2;
7849 pos += 2;
7850 }
7851 }
7852 if task_mask.is_none() && pos == 0 && ids.len() > 1 {
7855 if let Some(lg) = self.embryo_prefill_chunked(ids, 0) {
7856 self.graph_logits = Some(lg);
7857 hidden = vec![0.0; self.hidden_size];
7858 pos = ids.len();
7859 }
7860 }
7861 while pos < ids.len() {
7862 hidden = self.forward_layers(&self.embed_single(ids[pos]), pos, task_mask);
7863 self.check_forward_graph("forward_ids", pos)?;
7864 pos += 1;
7865 }
7866 if let Some(logits) = self.graph_logits.take() {
7867 if let Err(err) = self.o1_seal_checked() {
7871 self.clear_sequence_state();
7872 return Err(err);
7873 }
7874 return Ok(logits);
7875 }
7876 if let Err(err) = self.o1_seal_checked() {
7880 self.clear_sequence_state();
7881 return Err(err);
7882 }
7883 let normed = inference::rms_norm(
7884 &hidden,
7885 &self.weights.final_norm,
7886 self.rms_eps,
7887 self.norm_style,
7888 );
7889 Ok(self.lm_head_forward(&normed))
7890 }
7891
7892 #[doc(hidden)]
7896 pub fn dsv41_serial_logits(&mut self, ids: &[u32]) -> Result<Vec<Vec<f32>>, String> {
7897 #[cfg(target_os = "macos")]
7898 crate::gpu_metal::set_io_namespace(self.graph_kv_id);
7899 if ids.is_empty() {
7900 return Err("empty id sequence".to_string());
7901 }
7902 self.clear_sequence_state();
7903 self.dsv41
7904 .as_ref()
7905 .ok_or_else(|| "dsv41 serial logits require a DeepSeek-V4.1 model".to_string())?;
7906 self.o1_begin();
7907 let rows = {
7908 let pool = self.pool.clone();
7909 let b = self
7910 .dsv41
7911 .as_mut()
7912 .expect("dsv41 checked above; state cannot change during forward");
7913 let mut rows = Vec::with_capacity(ids.len());
7914 for (position, &id) in ids.iter().enumerate() {
7915 let mut logits = Vec::new();
7916 crate::dsv41::forward_token(
7917 &b.0,
7918 &b.1,
7919 &b.2,
7920 &mut b.3,
7921 id,
7922 position,
7923 pool.as_deref(),
7924 &mut logits,
7925 );
7926 rows.push(logits);
7927 }
7928 rows
7929 };
7930 self.o1_seal();
7931 Ok(rows)
7932 }
7933
7934 pub fn ppl_ids(&mut self, ids: &[u32]) -> Result<f64, String> {
7941 let (nll, cnt) = self.nll_ids_from(ids, 0)?;
7942 Ok((nll / cnt.max(1) as f64).exp())
7943 }
7944
7945 pub fn probe_ffn_mass(&mut self, ids: &[u32]) -> Vec<Vec<f64>> {
7950 self.clear_sequence_state();
7951 FFN_PROBE.with(|p| {
7952 *p.borrow_mut() = Some(vec![vec![0f64; self.intermediate_size]; self.num_layers]);
7953 });
7954 crate::gpu::cpu_scope(|| {
7955 for (pos, &id) in ids.iter().enumerate() {
7956 let emb = self.embed_single(id);
7957 let _ = self.forward_layers(&emb, pos, None);
7958 }
7959 });
7960 self.clear_sequence_state();
7961 FFN_PROBE
7962 .with(|p| p.borrow_mut().take())
7963 .unwrap_or_default()
7964 }
7965
7966 pub fn probe_ffn_mass_batch(&mut self, ids: &[u32]) -> Result<Vec<Vec<f64>>, String> {
7970 if let Err(err) = self.nll_begin() {
7971 let _ = FFN_PROBE.with(|p| p.borrow_mut().take());
7975 self.nll_end();
7976 return Err(err);
7977 }
7978 FFN_PROBE.with(|p| {
7979 *p.borrow_mut() = Some(vec![vec![0f64; self.intermediate_size]; self.num_layers]);
7980 });
7981 let result: Result<(), String> = (|| {
7982 for chunk in ids.chunks(256) {
7983 if chunk.len() < 2 {
7984 continue;
7985 }
7986 self.nll_ids_masked(chunk, 0, None)?;
7987 }
7988 Ok(())
7989 })();
7990 self.nll_end();
7991 let probe = FFN_PROBE
7992 .with(|p| p.borrow_mut().take())
7993 .unwrap_or_default();
7994 match result {
7995 Ok(()) => Ok(probe),
7996 Err(err) => {
7997 drop(probe);
7998 Err(err)
7999 }
8000 }
8001 }
8002
8003 pub fn ppl_ids_masked(&mut self, ids: &[u32], mask: &TaskMask) -> Result<f64, String> {
8007 self.nll_begin()?;
8008 let result: Result<f64, String> = (|| {
8009 let mut nll = 0f64;
8010 let mut cnt = 0usize;
8011 let mut hidden = vec![0f32; self.hidden_size];
8012 for (pos, &id) in ids.iter().enumerate() {
8013 if pos > 0 {
8014 inference::rms_norm_into(
8015 &hidden,
8016 &self.weights.final_norm,
8017 self.rms_eps,
8018 self.norm_style,
8019 &mut self.ws.n1,
8020 );
8021 let mut logits = self.lm_head_forward(&self.ws.n1);
8022 let max = logits.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
8023 let sum: f64 = logits.iter().map(|&v| ((v - max) as f64).exp()).sum();
8024 let p = ((logits[id as usize] - max) as f64).exp() / sum.max(1e-300);
8025 nll -= p.max(1e-300).ln();
8026 cnt += 1;
8027 attention::recycle_buf(&mut logits);
8028 }
8029 let emb = self.embed_single(id);
8030 hidden = self.forward_layers(&emb, pos, Some(mask));
8031 self.nll_check_graph("masked serial forward", pos)?;
8032 let _ = self.graph_logits.take();
8036 }
8037 Ok((nll / cnt.max(1) as f64).exp())
8038 })();
8039 self.nll_end();
8040 result
8041 }
8042
8043 pub fn nll_ids_masked(
8062 &mut self,
8063 ids: &[u32],
8064 start: usize,
8065 task_mask: Option<&TaskMask>,
8066 ) -> Result<(f64, usize), String> {
8067 let task_mask = self.drop_open_mask(task_mask);
8068 self.nll_ids_inner(ids, start, task_mask)
8069 }
8070
8071 pub fn nll_ids_from(&mut self, ids: &[u32], start: usize) -> Result<(f64, usize), String> {
8072 self.nll_ids_inner(ids, start, None)
8073 }
8074
8075 fn nll_ids_inner(
8076 &mut self,
8077 ids: &[u32],
8078 start: usize,
8079 task_mask: Option<&TaskMask>,
8080 ) -> Result<(f64, usize), String> {
8081 self.nll_begin()?;
8082 let result: Result<(f64, usize), String> = (|| {
8083 let mut nll = 0f64;
8084 let mut cnt = 0usize;
8085 let (graph_quality, fused_head_quality) = nll_graph_policy(
8098 task_mask.is_none(),
8099 self.graph_prefill_preferred(),
8100 crate::gpu::q1_force(),
8101 );
8102 self.graph_head_required = fused_head_quality;
8103 self.graph_want_logits = fused_head_quality;
8104 #[cfg(target_os = "macos")]
8105 if graph_quality && std::env::var("CMF_METAL_BATCH_NLL").as_deref() != Ok("0") {
8106 match self.nll_batch_metal(ids, start) {
8107 MetalBatchNllOutcome::Completed(nll, count) => {
8108 return Ok((nll, count));
8109 }
8110 MetalBatchNllOutcome::Declined => {}
8111 MetalBatchNllOutcome::Failed(err) => return Err(err),
8112 }
8113 }
8114 if self.can_prefill_batched() && !graph_quality {
8115 const CHUNK: usize = 128;
8121 const LM_SUB: usize = 32;
8122 let n = ids.len().saturating_sub(1);
8123 let hs = self.hidden_size;
8124 let rows = self.weights.lm_head.rows();
8125 let mut pos = 0usize;
8126 let state_trace = std::env::var("CMF_STATE_TRACE").is_ok();
8127 while pos < n {
8128 let end = (pos + CHUNK).min(n);
8129 let bsz = end - pos;
8130 let hb = self.prefill_rows(&ids[pos..end], pos, task_mask)?;
8131 self.nll_check_graph("batched prefill", pos)?;
8132 if state_trace && end % 256 == 0 {
8133 self.trace_recurrent_state(end);
8134 }
8135 let mut k0 = 0usize;
8136 while k0 < bsz {
8137 let k1 = (k0 + LM_SUB).min(bsz);
8138 let sb = k1 - k0;
8139 if pos + k1 <= start {
8142 k0 = k1;
8143 continue;
8144 }
8145 let mut normed = vec![0.0f32; sb * hs];
8146 for k in 0..sb {
8147 let r = inference::rms_norm(
8148 &hb[(k0 + k) * hs..(k0 + k + 1) * hs],
8149 &self.weights.final_norm,
8150 self.rms_eps,
8151 self.norm_style,
8152 );
8153 normed[k * hs..(k + 1) * hs].copy_from_slice(&r);
8154 }
8155 let mut logits = vec![0.0f32; sb * rows];
8156 self.weights
8157 .lm_head
8158 .matmat(&normed, sb, &mut logits, self.pool.as_deref());
8159 for k in 0..sb {
8160 if pos + k0 + k < start {
8161 continue;
8162 }
8163 self.nll_check_graph("batched score row", pos + k0 + k)?;
8164 let lg = &mut logits[k * rows..k * rows + self.vocab_size.min(rows)];
8165 if let Some(mu) = self.logit_multiplier {
8166 for v in lg.iter_mut() {
8167 *v *= mu;
8168 }
8169 }
8170 if let Some(c) = self.final_softcap {
8174 for v in lg.iter_mut() {
8175 *v = c * (*v / c).tanh();
8176 }
8177 }
8178 if let Some(cm) = self.head_clusters.clone() {
8181 self.hierarchical_head_logprobs(
8182 &normed[k * hs..(k + 1) * hs],
8183 &cm,
8184 lg,
8185 );
8186 }
8187 let lg = &logits[k * rows..k * rows + self.vocab_size.min(rows)];
8188 let target = ids[pos + k0 + k + 1] as usize;
8189 let max = lg.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
8190 let lse: f64 = lg
8191 .iter()
8192 .map(|&v| ((v - max) as f64).exp())
8193 .sum::<f64>()
8194 .ln()
8195 + max as f64;
8196 nll += lse - lg[target] as f64;
8197 cnt += 1;
8198 if std::env::var("CMF_PPL_TRACE").is_ok() {
8199 let top = lg
8200 .iter()
8201 .enumerate()
8202 .max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
8203 .map(|(i, _)| i)
8204 .unwrap_or(0);
8205 eprintln!(
8206 "BTRACE pos {} target {} nll {:.4} top {} lg_t {:.3} lg_top {:.3}",
8207 pos + k0 + k,
8208 target,
8209 lse - lg[target] as f64,
8210 top,
8211 lg[target],
8212 lg[top]
8213 );
8214 }
8215 }
8216 k0 = k1;
8217 }
8218 pos = end;
8219 }
8220 return Ok((nll, cnt));
8221 }
8222 for pos in 0..ids.len().saturating_sub(1) {
8223 let hidden = self.forward_layers(&self.embed_single(ids[pos]), pos, task_mask);
8224 self.nll_check_graph("serial forward", pos)?;
8225 let out_of_band = self.graph_logits.take();
8233 if self.graph_head_required && out_of_band.is_none() {
8234 METAL_GRAPH_HEAD_MISS.fetch_add(
8235 1,
8236 std::sync::atomic::Ordering::Relaxed,
8237 );
8238 return Err(format!(
8239 "fused Metal graph head did not complete at NLL position {pos}"
8240 ));
8241 }
8242 if pos < start {
8243 continue;
8244 }
8245 let logits = match out_of_band {
8246 Some(lg) => lg,
8247 None => {
8248 let normed = inference::rms_norm(
8249 &hidden,
8250 &self.weights.final_norm,
8251 self.rms_eps,
8252 self.norm_style,
8253 );
8254 self.lm_head_forward(&normed)
8258 }
8259 };
8260 let target = ids[pos + 1] as usize;
8261 let max = logits.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
8262 let lse: f64 = logits
8263 .iter()
8264 .map(|&v| ((v - max) as f64).exp())
8265 .sum::<f64>()
8266 .ln()
8267 + max as f64;
8268 let tok_nll = lse - logits[target] as f64;
8269 if std::env::var("CMF_PPL_TRACE").is_ok() && pos < 48 {
8270 let top = logits
8271 .iter()
8272 .enumerate()
8273 .max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
8274 .map(|(i, _)| i)
8275 .unwrap_or(0);
8276 eprintln!(
8277 "pos {pos:3} tgt {target:6} nll {tok_nll:7.3} | top1 {top:6} lg[t]={:.2} lg[top]={:.2}",
8278 logits[target], logits[top]
8279 );
8280 }
8281 nll += tok_nll;
8282 cnt += 1;
8283 }
8284 Ok((nll, cnt))
8285 })();
8286 self.nll_end();
8287 result
8288 }
8289
8290 fn nll_from_hidden(&mut self, hidden: &[f32], target: u32, pos: usize) -> f64 {
8295 let normed = inference::rms_norm(
8296 hidden,
8297 &self.weights.final_norm,
8298 self.rms_eps,
8299 self.norm_style,
8300 );
8301 let mut logits = self.lm_head_forward(&normed);
8304 let target = target as usize;
8305 let max = logits.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
8306 let lse: f64 = logits
8307 .iter()
8308 .map(|&v| ((v - max) as f64).exp())
8309 .sum::<f64>()
8310 .ln()
8311 + max as f64;
8312 let tok_nll = lse - logits[target] as f64;
8313 if std::env::var("CMF_PPL_TRACE").is_ok() && pos < 48 {
8314 let top = logits
8315 .iter()
8316 .enumerate()
8317 .max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
8318 .map(|(i, _)| i)
8319 .unwrap_or(0);
8320 eprintln!(
8321 "pos {pos:3} tgt {target:6} nll {tok_nll:7.3} | top1 {top:6} lg[t]={:.2} lg[top]={:.2}",
8322 logits[target], logits[top]
8323 );
8324 }
8325 attention::recycle_buf(&mut logits);
8326 tok_nll
8327 }
8328
8329 fn trace_recurrent_state(&self, pos: usize) {
8337 let stats = |v: &[f32]| -> (f64, f64) {
8338 if v.is_empty() {
8339 return (0.0, 0.0);
8340 }
8341 let (mut ss, mut mx) = (0f64, 0f64);
8342 for &x in v {
8343 ss += (x as f64) * (x as f64);
8344 mx = mx.max((x as f64).abs());
8345 }
8346 ((ss / v.len() as f64).sqrt(), mx)
8347 };
8348 for (li, l) in self.kv_cache.layers.iter().enumerate() {
8349 let lw = &self.weights.layers[self.phys_layer(li)];
8350 let (kind, s_len) = match &lw.attn {
8351 AttnKind::Linear(_) => (
8352 "vmf",
8353 self.vmf_cfg.map(|c| c.state_len()).unwrap_or(0),
8354 ),
8355 AttnKind::LinearGdn(_) => (
8356 "gdn",
8357 self.gdn_cfg.map(|c| c.state_len()).unwrap_or(0),
8358 ),
8359 AttnKind::Bounded(_) => ("bounded", 0),
8360 AttnKind::Full { .. } => ("full", 0),
8361 _ => ("other", 0),
8362 };
8363 let (rms, max) = stats(&l.linear_state);
8364 let s_part = if kind == "vmf" {
8365 &l.linear_state[..s_len.min(l.linear_state.len())]
8366 } else if kind == "gdn" {
8367 let ring = l.linear_state.len().saturating_sub(s_len).min(l.linear_state.len());
8368 &l.linear_state[ring..]
8369 } else {
8370 &l.linear_state[..0]
8371 };
8372 let (s_rms, s_max) = stats(s_part);
8373 let (ring_rms, ring_len) = match &l.bounded {
8374 Some(b) => (stats(&b.ring_k).0, b.ring_k.len()),
8375 None => (0.0, 0),
8376 };
8377 eprintln!(
8378 "STATE pos={pos} layer={li} kind={kind} state_len={} rms={rms:.5} max={max:.4} \
8379 S_rms={s_rms:.5} S_max={s_max:.4} ring_k_rms={ring_rms:.5} ring_len={ring_len} kv_rows={}",
8380 l.linear_state.len(),
8381 l.seq_len
8382 );
8383 }
8384 }
8385
8386 pub fn nll_ids_o1(&mut self, ids: &[u32], prefill: usize) -> Result<(f64, usize), String> {
8404 self.nll_begin()?;
8409 let requested_prefix = (prefill > 0).then_some(prefill);
8410 self.o1_begin_with_prefix(requested_prefix);
8411 let n = ids.len().saturating_sub(1);
8412 let requested_start = prefill.min(n);
8413 let exact_end = if self.o1_active() {
8418 match requested_prefix {
8419 Some(requested) => self.o1_effective_boundary(requested),
8420 None => self
8421 .o1_cfg
8422 .as_ref()
8423 .and_then(|c| crate::nystrom::o1_deferred_boundary(c.w, c.sink)),
8424 }
8425 .unwrap_or(requested_start)
8426 .min(n)
8427 } else {
8428 requested_start
8429 };
8430 let mut nll = 0f64;
8431 let mut cnt = 0usize;
8432
8433 let mut pos = 0usize;
8437 if self.can_prefill_batched() {
8438 const CHUNK: usize = 128;
8439 while pos < exact_end {
8440 let end = (pos + CHUNK).min(exact_end);
8441 let hiddens = self.prefill_batch(&ids[pos..end], pos);
8442 if self
8443 .graph_failed
8444 .swap(false, std::sync::atomic::Ordering::Relaxed)
8445 {
8446 self.cancel
8447 .store(false, std::sync::atomic::Ordering::Relaxed);
8448 self.nll_end();
8449 return Err("GPU graph failed during O(1) NLL prefix".into());
8450 }
8451 for row in 0..end - pos {
8452 let score_pos = pos + row;
8453 if score_pos >= requested_start && score_pos < n {
8454 nll += self.nll_from_hidden(
8455 &hiddens[row * self.hidden_size..(row + 1) * self.hidden_size],
8456 ids[score_pos + 1],
8457 score_pos,
8458 );
8459 cnt += 1;
8460 }
8461 }
8462 pos = end;
8463 }
8464 } else {
8465 while pos < exact_end {
8466 let hidden = self.forward_layers(&self.embed_single(ids[pos]), pos, None);
8467 if self
8468 .graph_failed
8469 .swap(false, std::sync::atomic::Ordering::Relaxed)
8470 {
8471 self.cancel
8472 .store(false, std::sync::atomic::Ordering::Relaxed);
8473 self.nll_end();
8474 return Err("GPU graph failed during O(1) NLL prefix".into());
8475 }
8476 if pos >= requested_start && pos < n {
8477 nll += self.nll_from_hidden(&hidden, ids[pos + 1], pos);
8478 cnt += 1;
8479 }
8480 pos += 1;
8481 }
8482 }
8483 self.o1_seal_checked().map_err(|err| {
8484 self.nll_end();
8485 err
8486 })?;
8487
8488 let batch_k = std::env::var("CMF_BATCH_K")
8497 .ok()
8498 .and_then(|v| v.parse::<usize>().ok())
8499 .unwrap_or(0);
8500 let batch_admitted = batch_k > 0
8501 && self.can_prefill_batched()
8502 && self.o1_active()
8503 && std::env::var("CMF_O1_GPU").as_deref() == Ok("1")
8504 && (0..self.num_layers).all(|li| {
8505 let cache = &self.kv_cache.layers[self.phys_layer(li)];
8506 cache.o1.is_none() || cache.o1_views().is_some()
8507 });
8508 if std::env::var("CMF_GRAPH_PROF").is_ok() {
8509 eprintln!(
8510 "nll-batch: phase=post-seal admission={} requested_k={} scored_rows={}",
8511 batch_admitted,
8512 batch_k,
8513 n.saturating_sub(exact_end),
8514 );
8515 }
8516 let mut batch_completed = false;
8517 if batch_admitted && exact_end < n {
8518 let hs = self.hidden_size;
8519 let mut batch_pos = exact_end;
8520 while batch_pos < n {
8521 let end = (batch_pos + batch_k).min(n);
8522 let bk = end - batch_pos;
8523 let mut hiddens = vec![0.0f32; bk * hs];
8524 for (row, &id) in ids[batch_pos..end].iter().enumerate() {
8525 hiddens[row * hs..(row + 1) * hs].copy_from_slice(&self.embed_single(id));
8526 }
8527 let positions: Vec<usize> = (batch_pos..end).collect();
8528 let t_batch = std::time::Instant::now();
8529 let outcome = self.try_batch_graph_wgpu(&mut hiddens, &positions, bk, None);
8530 if std::env::var("CMF_GRAPH_PROF").is_ok() {
8531 let ms = t_batch.elapsed().as_secs_f64() * 1000.0;
8532 eprintln!(
8533 "nll-batch: phase=post-seal mode=o1 k={bk} pos={}..{} outcome={outcome:?} {ms:.1} ms ({:.1} tok/s)",
8534 batch_pos,
8535 end.saturating_sub(1),
8536 bk as f64 / (ms / 1000.0),
8537 );
8538 }
8539 if let Err(err) = self.nll_check_graph("batch graph", batch_pos) {
8540 self.nll_end();
8541 return Err(err);
8542 }
8543 match outcome {
8544 crate::gpu::BatchGraphOutcome::Completed => {
8545 batch_completed = true;
8546 for row in 0..bk {
8547 nll += self.nll_from_hidden(
8548 &hiddens[row * hs..(row + 1) * hs],
8549 ids[batch_pos + row + 1],
8550 batch_pos + row,
8551 );
8552 cnt += 1;
8553 }
8554 batch_pos = end;
8555 }
8556 crate::gpu::BatchGraphOutcome::Declined => {
8557 if batch_completed {
8558 self.nll_end();
8559 return Err(format!(
8560 "O(1) NLL batch declined after completed chunk at position {batch_pos}"
8561 ));
8562 }
8563 break;
8564 }
8565 crate::gpu::BatchGraphOutcome::Failed => {
8566 self.nll_end();
8567 return Err(format!(
8568 "O(1) NLL batch graph failed after admission at position {batch_pos}"
8569 ));
8570 }
8571 }
8572 }
8573 if batch_completed && cnt == n.saturating_sub(requested_start) {
8574 self.nll_end();
8575 return Ok((nll, cnt));
8576 }
8577 }
8578
8579 for pos in exact_end..n {
8584 let hidden = self.forward_layers(&self.embed_single(ids[pos]), pos, None);
8585 if self
8586 .graph_failed
8587 .swap(false, std::sync::atomic::Ordering::Relaxed)
8588 {
8589 self.cancel
8590 .store(false, std::sync::atomic::Ordering::Relaxed);
8591 self.nll_end();
8592 return Err(format!(
8593 "GPU graph failed during O(1) NLL serial scoring at position {pos}"
8594 ));
8595 }
8596 nll += self.nll_from_hidden(&hidden, ids[pos + 1], pos);
8597 cnt += 1;
8598 }
8599 self.nll_end();
8600 Ok((nll, cnt))
8601 }
8602
8603 pub fn calib_ids(&mut self, ids: &[u32], temps: &[f32]) -> (Vec<bool>, Vec<Vec<f32>>) {
8611 self.clear_sequence_state();
8612 let n = ids.len().saturating_sub(1);
8613 let mut correct = Vec::with_capacity(n);
8614 let mut pmax = Vec::with_capacity(n);
8615 for pos in 0..n {
8616 let emb = self.embed_single(ids[pos]);
8617 let hidden = self.forward_layers(&emb, pos, None);
8618 let logits = if let Some(logits) = self.graph_logits.take() {
8619 logits
8620 } else {
8621 let normed = inference::rms_norm(
8622 &hidden,
8623 &self.weights.final_norm,
8624 self.rms_eps,
8625 self.norm_style,
8626 );
8627 self.lm_head_forward(&normed)
8631 };
8632 let target = ids[pos + 1] as usize;
8633 let (mut amax, mut mval) = (0usize, f32::NEG_INFINITY);
8634 for (i, &v) in logits.iter().enumerate() {
8635 if v > mval {
8636 mval = v;
8637 amax = i;
8638 }
8639 }
8640 correct.push(amax == target);
8641 let row: Vec<f32> = temps
8642 .iter()
8643 .map(|&t| {
8644 let tt = t.max(1e-3);
8645 let s: f32 = logits.iter().map(|&v| ((v - mval) / tt).exp()).sum();
8646 1.0 / s.max(1e-12) })
8648 .collect();
8649 pmax.push(row);
8650 }
8651 self.clear_sequence_state();
8652 (correct, pmax)
8653 }
8654
8655 pub fn ppl_ids_dynamic(&mut self, ids: &[u32]) -> Result<(f64, usize), String> {
8662 if self.dyn_router.is_none() {
8663 return Ok((self.ppl_ids(ids)?, 0));
8664 }
8665 self.nll_begin()?;
8666 let saved_active = self.dyn_active;
8667 let mut router = self
8668 .dyn_router
8669 .take()
8670 .ok_or_else(|| "dynamic router disappeared before PPL scoring".to_string())?;
8671 router.reset();
8672 self.dyn_phi_seen = 0;
8673 let _ = self.set_active_skill(None);
8674
8675 let result: Result<(f64, usize), String> = (|| {
8676 let mut nll = 0f64;
8677 let mut cnt = 0usize;
8678 for pos in 0..ids.len().saturating_sub(1) {
8679 let hidden = self.forward_layers(&self.embed_single(ids[pos]), pos, None);
8680 self.nll_check_graph("dynamic serial forward", pos)?;
8681 let out_of_band = self.graph_logits.take();
8682 let mut logits = match out_of_band {
8683 Some(lg) => lg,
8684 None => {
8685 let normed = inference::rms_norm(
8686 &hidden,
8687 &self.weights.final_norm,
8688 self.rms_eps,
8689 self.norm_style,
8690 );
8691 self.lm_head_forward(&normed)
8695 }
8696 };
8697 let target = ids[pos + 1] as usize;
8698 let max = logits.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
8699 let lse: f64 = logits
8700 .iter()
8701 .map(|&v| ((v - max) as f64).exp())
8702 .sum::<f64>()
8703 .ln()
8704 + max as f64;
8705 let tok_nll = lse - logits[target] as f64;
8706 if std::env::var("CMF_PPL_TRACE").is_ok() && pos < 48 {
8707 let top = logits
8708 .iter()
8709 .enumerate()
8710 .max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
8711 .map(|(i, _)| i)
8712 .unwrap_or(0);
8713 eprintln!(
8714 "pos {pos:3} tgt {target:6} nll {tok_nll:7.3} | top1 {top:6} lg[t]={:.2} lg[top]={:.2}",
8715 logits[target], logits[top]
8716 );
8717 }
8718 nll += tok_nll;
8719 cnt += 1;
8720 attention::recycle_buf(&mut logits);
8721 let phi = self.dyn_phi_ema.clone();
8723 if let Some(new_active) = router.step(&phi, pos) {
8724 let _ = self.set_active_skill(new_active);
8725 }
8726 }
8727 Ok(((nll / cnt.max(1) as f64).exp(), router.switches.len()))
8728 })();
8729
8730 let _ = self.set_active_skill(saved_active);
8733 self.dyn_router = Some(router);
8734 self.nll_end();
8735 result
8736 }
8737
8738 pub fn probe_phi(&mut self, ids: &[u32], layer: usize) -> Vec<f32> {
8740 self.clear_sequence_state();
8741 let mut acc = vec![0f32; self.hidden_size];
8742 for (pos, &id) in ids.iter().enumerate() {
8743 let h = self.forward_layers_upto(&self.embed_single(id), pos, None, Some(layer));
8744 for (a, v) in acc.iter_mut().zip(&h) {
8745 *a += v;
8746 }
8747 }
8748 let n = ids.len().max(1) as f32;
8749 for a in acc.iter_mut() {
8750 *a /= n;
8751 }
8752 self.clear_sequence_state();
8753 acc
8754 }
8755
8756 pub fn probe_phi_span(
8770 &mut self,
8771 ids: &[u32],
8772 layer: usize,
8773 span: std::ops::Range<usize>,
8774 ) -> Vec<f32> {
8775 #[cfg(target_os = "macos")]
8776 crate::gpu_metal::set_io_namespace(self.graph_kv_id);
8777 let end = span.end.min(ids.len());
8778 let start = span.start.min(end);
8779 let reset = |p: &mut Self| p.clear_sequence_state();
8780 reset(self);
8781 let mut acc = vec![0f32; self.hidden_size];
8782 for (pos, &id) in ids[..end].iter().enumerate() {
8783 let h = self.forward_layers_upto(&self.embed_single(id), pos, None, Some(layer));
8784 if pos >= start {
8785 for (a, v) in acc.iter_mut().zip(&h) {
8786 *a += v;
8787 }
8788 }
8789 }
8790 let n = end - start;
8791 if n > 0 {
8792 let n = n as f32;
8793 for a in acc.iter_mut() {
8794 *a /= n;
8795 }
8796 }
8797 reset(self);
8798 acc
8799 }
8800
8801 pub fn decode_step_logits(&mut self, token: u32, position: usize) -> Vec<f32> {
8808 #[cfg(target_os = "macos")]
8809 crate::gpu_metal::set_io_namespace(self.graph_kv_id);
8810 self.graph_logits = None;
8811 let hidden = self.forward_layers(&self.embed_single(token), position, None);
8812 if let Some(logits) = self.graph_logits.take() {
8813 return logits;
8814 }
8815 inference::rms_norm_into(
8816 &hidden,
8817 &self.weights.final_norm,
8818 self.rms_eps,
8819 self.norm_style,
8820 &mut self.ws.n1,
8821 );
8822 self.lm_head_forward(&self.ws.n1)
8823 }
8824
8825 fn prefill_batch(&mut self, ids: &[u32], start_pos: usize) -> Vec<f32> {
8831 self.prefill_batch_masked(ids, start_pos, None)
8832 }
8833
8834 fn prefill_batch_masked(
8840 &mut self,
8841 ids: &[u32],
8842 start_pos: usize,
8843 task_mask: Option<&TaskMask>,
8844 ) -> Vec<f32> {
8845 self.prefill_batch_span(PrefillIn::Ids(ids), start_pos, task_mask, 0, usize::MAX)
8846 }
8847
8848 fn prefill_rows(
8856 &mut self,
8857 ids: &[u32],
8858 pos: usize,
8859 task_mask: Option<&TaskMask>,
8860 ) -> Result<Vec<f32>, String> {
8861 self.prefill_input_rows(PrefillIn::Ids(ids), pos, task_mask)
8862 }
8863
8864 fn prefill_input_rows(
8865 &mut self,
8866 input: PrefillIn<'_>,
8867 pos: usize,
8868 task_mask: Option<&TaskMask>,
8869 ) -> Result<Vec<f32>, String> {
8870 self.mimo_moe_prepare();
8871 let hs = self.hidden_size;
8872 let bk = match input {
8873 PrefillIn::Ids(ids) => ids.len(),
8874 PrefillIn::Hidden(rows) => rows.len() / hs,
8875 };
8876 #[cfg(not(target_os = "macos"))]
8877 if task_mask.is_none()
8878 && !self.o1_active()
8879 && bk > 1
8880 && (self.batch_prefix_prefill()
8881 || (self.verify_exact_moe
8882 && crate::gpu::enabled_here()
8883 && crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Decode)))
8884 {
8885 let mut hiddens = match input {
8886 PrefillIn::Hidden(rows) => rows.to_vec(),
8887 PrefillIn::Ids(ids) => ids.iter().flat_map(|&id| self.embed_single(id)).collect(),
8888 };
8889 let positions: Vec<usize> = (pos..pos + bk).collect();
8890 let mut run = 0usize;
8891 match self.try_batch_graph_wgpu_prefix(
8892 &mut hiddens,
8893 &positions,
8894 bk,
8895 None,
8896 Some(&mut run),
8897 ) {
8898 crate::gpu::BatchGraphOutcome::Completed => {
8899 let out = if run < self.num_layers {
8900 self.prefill_batch_span(
8901 PrefillIn::Hidden(&hiddens),
8902 pos,
8903 None,
8904 run,
8905 self.num_layers,
8906 )
8907 } else {
8908 hiddens
8909 };
8910 return if self.graph_failed.load(std::sync::atomic::Ordering::Relaxed) {
8911 Err("MiMo attention graph failed after admission".into())
8912 } else { Ok(out) };
8913 }
8914 crate::gpu::BatchGraphOutcome::Failed => {
8915 return Err("batched prefix prefill failed after admission".into());
8916 }
8917 crate::gpu::BatchGraphOutcome::Declined => {
8918 #[cfg(feature = "gpu")]
8920 self.pull_lagging_host_kv(0, self.num_layers, pos);
8921 }
8922 }
8923 }
8924 let out = self.prefill_batch_span(input, pos, task_mask, 0, usize::MAX);
8925 if self.graph_failed.load(std::sync::atomic::Ordering::Relaxed) {
8926 Err("batch tail graph failed after admission".into())
8927 } else { Ok(out) }
8928 }
8929
8930 fn prefill_batch_span(
8936 &mut self,
8937 input: PrefillIn<'_>,
8938 start_pos: usize,
8939 task_mask: Option<&TaskMask>,
8940 from: usize,
8941 upto_excl: usize,
8942 ) -> Vec<f32> {
8943 let hs = self.hidden_size;
8944 let b = match input {
8945 PrefillIn::Ids(ids) => ids.len(),
8946 PrefillIn::Hidden(hb) => hb.len() / hs,
8947 };
8948 let upto_excl = upto_excl.min(self.num_layers);
8949 let mut h: Vec<f32>;
8953 let mut h_ready;
8954 match input {
8955 PrefillIn::Ids(_) => {
8956 h = vec![0.0; b * hs];
8957 h_ready = false;
8958 }
8959 PrefillIn::Hidden(hb) => {
8960 h = hb.to_vec();
8961 h_ready = true;
8962 }
8963 }
8964 let fill_h = |h: &mut Vec<f32>, me: &Self| {
8965 if let PrefillIn::Ids(ids) = input {
8966 for (bi, &id) in ids.iter().enumerate() {
8967 let e = me.embed_single(id);
8968 h[bi * hs..(bi + 1) * hs].copy_from_slice(&e);
8969 }
8970 if let Ok(tp) = std::env::var("CMF_TRACE_POS") {
8971 if let Ok(t) = tp.parse::<usize>() {
8972 if t >= start_pos && t < start_pos + ids.len() {
8973 let bi = t - start_pos;
8974 let row = &h[bi * hs..(bi + 1) * hs];
8975 let n: f32 = row.iter().map(|x| x * x).sum::<f32>().sqrt();
8976 eprintln!(
8977 "BATCH pos {t} embed: id {} |h| = {n:.6} h0 {:.6} h1 {:.6} | b={} start={start_pos} ids[..8]={:?}",
8978 ids[bi],
8979 row[0],
8980 row[1],
8981 ids.len(),
8982 &ids[..ids.len().min(8)]
8983 );
8984 }
8985 }
8986 }
8987 }
8988 };
8989 let (_nkv, _hd, _rd, eps) = (
8990 self.num_kv_heads,
8991 self.head_dim,
8992 self.rotary_dim,
8993 self.rms_eps,
8994 );
8995 let pool = self.pool.clone();
8996 let norm_style = self.norm_style;
8997 self.mimo_moe_prepare();
8998 let automatic_gpu_prefix = self.automatic_gpu_prefix();
8999
9000 #[cfg(target_os = "macos")]
9001 let mut chunk_skip_until = 0usize;
9002 for li in from..upto_excl {
9003 let _capacity_tail = automatic_gpu_prefix
9004 .filter(|&prefix| {
9005 li >= prefix && !(self.verify_exact_moe && self.mimo_moe.is_dynamic(li, false))
9006 })
9007 .map(|_| crate::gpu::enter_cpu_scope());
9008 crate::gpu::set_layer(li as i64); #[cfg(target_os = "macos")]
9015 if task_mask.is_none() {
9016 if li < chunk_skip_until {
9017 continue;
9018 }
9019 if !h_ready && li == 0 && self.weights.embed_tokens.q8_row_parts().is_none() {
9025 fill_h(&mut h, self);
9026 h_ready = true;
9027 }
9028 let ids_for_embed = match input {
9029 PrefillIn::Ids(ids) => (!h_ready && li == 0).then_some(ids),
9030 PrefillIn::Hidden(_) => None,
9031 };
9032 let end = self.chunk_run_gpu(li, &mut h, b, start_pos, ids_for_embed, upto_excl);
9033 if end > li {
9034 h_ready = true;
9035 chunk_skip_until = end;
9036 if self.is_loop_end(end - 1) && end < self.num_layers {
9039 for bi in 0..b {
9040 let normed = inference::rms_norm(
9041 &h[bi * hs..(bi + 1) * hs],
9042 &self.weights.final_norm,
9043 eps,
9044 norm_style,
9045 );
9046 h[bi * hs..(bi + 1) * hs].copy_from_slice(&normed);
9047 }
9048 }
9049 continue;
9050 }
9051 }
9052 if !h_ready {
9053 fill_h(&mut h, self);
9054 h_ready = true;
9055 }
9056 if task_mask.is_none() && self.verify_exact_moe {
9057 let positions: Vec<_> = (start_pos..start_pos + b).collect();
9058 match self.mimo_graph_layer_rows(li, &mut h, &positions) {
9059 crate::gpu::BatchGraphOutcome::Completed => continue,
9060 crate::gpu::BatchGraphOutcome::Failed => return h,
9061 crate::gpu::BatchGraphOutcome::Declined => {},
9062 }
9063 }
9064 #[cfg(feature = "gpu")]
9065 self.pull_lagging_host_kv(li, li + 1, start_pos);
9066 let lw = &self.weights.layers[self.phys_layer(li)];
9067 match &lw.attn {
9069 AttnKind::Kda(w) => {
9070 let cfg = self.kda_cfg.expect("kda layer without kda_cfg");
9072 let mut normed = vec![0.0f32; b * hs];
9073 for bi in 0..b {
9074 inference::rms_norm_into(
9075 &h[bi * hs..(bi + 1) * hs],
9076 &lw.input_norm,
9077 eps,
9078 norm_style,
9079 &mut normed[bi * hs..(bi + 1) * hs],
9080 );
9081 }
9082 let attn = crate::linear_core::kda_forward_batch(
9083 &normed,
9084 b,
9085 w,
9086 &cfg,
9087 &mut self.kv_cache.layers[li].linear_state,
9088 pool.as_deref(),
9089 );
9090 for (dst, &a) in h.iter_mut().zip(&attn) {
9091 *dst += a;
9092 }
9093 }
9094 AttnKind::LinearGdn(w) => {
9095 let cfg = self.gdn_cfg.expect("gdn layer without gdn_cfg");
9097 let mut normed = vec![0.0f32; b * hs];
9098 for bi in 0..b {
9099 let r = inference::rms_norm(
9100 &h[bi * hs..(bi + 1) * hs],
9101 &lw.input_norm,
9102 eps,
9103 norm_style,
9104 );
9105 normed[bi * hs..(bi + 1) * hs].copy_from_slice(&r);
9106 }
9107 let attn = crate::linear_core::gdn_forward_batch(
9108 &normed,
9109 b,
9110 w,
9111 &cfg,
9112 &mut self.kv_cache.layers[li].linear_state,
9113 pool.as_deref(),
9114 );
9115 for (dst, &a) in h.iter_mut().zip(&attn) {
9116 *dst += a;
9117 }
9118 }
9119 AttnKind::ShortConv(w) => {
9120 let cfg = self
9123 .short_conv_cfg
9124 .expect("short-conv layer without short_conv_cfg");
9125 let mut normed = vec![0.0f32; b * hs];
9126 for bi in 0..b {
9127 inference::rms_norm_into(
9128 &h[bi * hs..(bi + 1) * hs],
9129 &lw.input_norm,
9130 eps,
9131 norm_style,
9132 &mut normed[bi * hs..(bi + 1) * hs],
9133 );
9134 }
9135 let attn = short_conv_forward_batch(
9136 &normed,
9137 b,
9138 w,
9139 &cfg,
9140 &mut self.kv_cache.layers[li].linear_state,
9141 pool.as_deref(),
9142 );
9143 for (dst, &a) in h.iter_mut().zip(&attn) {
9144 *dst += a;
9145 }
9146 }
9147 AttnKind::Mla(w) => {
9148 let inv_freq_l = self.layer_inv_freq(li);
9151 let rs = self.layer_rope_scale(li);
9152 let mut normed = vec![0.0f32; hs];
9153 for bi in 0..b {
9154 inference::rms_norm_into(
9155 &h[bi * hs..(bi + 1) * hs],
9156 &lw.input_norm,
9157 eps,
9158 norm_style,
9159 &mut normed,
9160 );
9161 let ao = mla_attention(
9162 w,
9163 &normed,
9164 &mut self.kv_cache.layers[li],
9165 start_pos + bi,
9166 &inv_freq_l,
9167 rs,
9168 eps,
9169 pool.as_deref(),
9170 );
9171 for (dst, &a) in h[bi * hs..(bi + 1) * hs].iter_mut().zip(&ao) {
9172 *dst += a;
9173 }
9174 }
9175 }
9176 AttnKind::Full {
9177 wq,
9178 wk,
9179 wv,
9180 wo,
9181 q_norm,
9182 k_norm,
9183 output_gate,
9184 softplus_gate,
9185 bias,
9186 } => {
9187 let mut normed = vec![0.0f32; b * hs];
9191 for bi in 0..b {
9192 inference::rms_norm_into(
9193 &h[bi * hs..(bi + 1) * hs],
9194 &lw.input_norm,
9195 eps,
9196 norm_style,
9197 &mut normed[bi * hs..(bi + 1) * hs],
9198 );
9199 }
9200 let inv_freq_l = self.layer_inv_freq(li);
9201 let (nkv_l, hd_l, rd_l) = self.layer_geom(li);
9202 let cfg = QwenAttnCfg {
9203 num_heads: self.layer_num_heads(li),
9204 num_kv_heads: nkv_l,
9205 head_dim: hd_l,
9206 hidden_size: hs,
9207 position: start_pos,
9208 inv_freq: &inv_freq_l,
9209 rotary_dim: rd_l,
9210 scale: self.attn_scale,
9211 softcap: self.attn_softcap,
9212 window: self.layer_window(li),
9213 v_norm: self.attn_v_norm,
9214 qk_norm_after_rope: self.qk_norm_after_rope,
9215 q_norm: q_norm.as_deref(),
9216 k_norm: k_norm.as_deref(),
9217 output_gate: *output_gate,
9218 softplus_gate: softplus_gate
9219 .as_ref()
9220 .map(|(gate, per_head)| (gate, *per_head)),
9221 rope_scale: self.layer_rope_scale(li),
9222 bias: bias
9223 .as_ref()
9224 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
9225 rms_eps: eps,
9226 norm_style,
9227 pool: pool.as_deref(),
9228 v_head_dim: self.layer_v_dim(li),
9229 };
9230 let mut attn = attention::qwen_attention_batch(
9231 &normed,
9232 b,
9233 wq,
9234 wk,
9235 wv,
9236 wo,
9237 &mut self.kv_cache.layers[li],
9238 &cfg,
9239 );
9240 if let Some(w) = &lw.attn_out_norm {
9241 for bi in 0..b {
9242 inference::rms_norm_into(
9243 &attn[bi * hs..(bi + 1) * hs],
9244 w,
9245 eps,
9246 norm_style,
9247 &mut normed[bi * hs..(bi + 1) * hs],
9248 );
9249 }
9250 attn.copy_from_slice(&normed);
9251 }
9252 for (dst, &a) in h.iter_mut().zip(&attn) {
9253 *dst += a;
9254 }
9255 }
9256 AttnKind::Bounded(w) => {
9257 let mut normed = vec![0.0f32; b * hs];
9260 for bi in 0..b {
9261 inference::rms_norm_into(
9262 &h[bi * hs..(bi + 1) * hs],
9263 &lw.input_norm,
9264 eps,
9265 norm_style,
9266 &mut normed[bi * hs..(bi + 1) * hs],
9267 );
9268 }
9269 let rope = self
9270 .bounded_rope
9271 .clone()
9272 .expect("bounded layer without an installed rotation table");
9273 let cfg = crate::bounded::BoundedAttnCfg {
9274 num_heads: self.num_heads,
9275 num_kv_heads: self.num_kv_heads,
9276 head_dim: self.head_dim,
9277 hidden_size: hs,
9278 scale: self.attn_scale,
9279 rope: &rope,
9280 pool: pool.as_deref(),
9281 };
9282 let mut attn = crate::bounded::bounded_attention_batch(
9283 &normed,
9284 b,
9285 w,
9286 &mut self.kv_cache.layers[li],
9287 &cfg,
9288 );
9289 if let Some(wn) = &lw.attn_out_norm {
9290 for bi in 0..b {
9291 inference::rms_norm_into(
9292 &attn[bi * hs..(bi + 1) * hs],
9293 wn,
9294 eps,
9295 norm_style,
9296 &mut normed[bi * hs..(bi + 1) * hs],
9297 );
9298 }
9299 attn.copy_from_slice(&normed);
9300 }
9301 for (dst, &a) in h.iter_mut().zip(&attn) {
9302 *dst += a;
9303 }
9304 attention::recycle_buf(&mut attn);
9305 }
9306 AttnKind::Linear(w) => {
9307 for bi in 0..b {
9308 let normed = inference::rms_norm(
9309 &h[bi * hs..(bi + 1) * hs],
9310 &lw.input_norm,
9311 eps,
9312 norm_style,
9313 );
9314 vmf_phase_forward(
9315 &normed,
9316 w,
9317 &self.vmf_cfg.expect("linear layer without vmf_cfg"),
9318 &mut self.kv_cache.layers[li].linear_state,
9319 pool.as_deref(),
9320 )
9321 .iter()
9322 .enumerate()
9323 .for_each(|(i, &a)| h[bi * hs + i] += a);
9324 }
9325 }
9326 }
9327
9328 let lw = &self.weights.layers[self.phys_layer(li)];
9330 let mut post = vec![0.0f32; b * hs];
9331 for bi in 0..b {
9332 let r =
9333 inference::rms_norm(&h[bi * hs..(bi + 1) * hs], &lw.post_norm, eps, norm_style);
9334 post[bi * hs..(bi + 1) * hs].copy_from_slice(&r);
9335 }
9336 let mask_row = task_mask
9339 .filter(|m| m.ffn_active_count(li) < self.intermediate_size)
9340 .and_then(|m| m.ffn_masks.get(li))
9341 .map(|v| v.as_slice());
9342 let mut ffn = match &lw.ffn {
9343 FfnKind::Dense(d) if !d.segs.is_empty() => {
9344 tube_ffn(d, &post, b, pool.as_deref(), mask_row)
9345 }
9346 FfnKind::Dense(d) => dense_ffn_batch(d, &post, b, pool.as_deref(), mask_row),
9347 FfnKind::Moe(m) if self.verify_exact_moe && self.mimo_moe.is_dynamic(li, false) => {
9348 moe_ffn_banked_rows(&mut self.mimo_moe, li, m, &post, b, hs, pool.as_deref())
9349 }
9350 FfnKind::Moe(m) if self.verify_exact_moe => {
9351 moe_ffn_rows_exact(m, &post, b, hs, pool.as_deref())
9352 }
9353 FfnKind::Moe(m) if self.mimo_moe.is_dynamic(li, false) => {
9356 let before = m.stats.borrow().clone();
9357 let out = crate::gpu::cpu_scope(|| {
9358 moe_ffn_batch(m, &post, b, hs, pool.as_deref(), None)
9359 });
9360 self.mimo_moe.prime(li, m, &before);
9361 out
9362 }
9363 FfnKind::Moe(m) => moe_ffn_batch(m, &post, b, hs, pool.as_deref(), None),
9364 FfnKind::DenseMoe(dm) => {
9367 let mut out = vec![0.0f32; b * hs];
9368 for bi in 0..b {
9369 let r = dense_moe_ffn(
9370 dm,
9371 &post[bi * hs..(bi + 1) * hs],
9372 &h[bi * hs..(bi + 1) * hs],
9373 eps,
9374 norm_style,
9375 pool.as_deref(),
9376 );
9377 out[bi * hs..(bi + 1) * hs].copy_from_slice(&r);
9378 }
9379 out
9380 }
9381 };
9382 if let Some(w) = &lw.ffn_out_norm {
9383 for bi in 0..b {
9384 inference::rms_norm_into(
9385 &ffn[bi * hs..(bi + 1) * hs],
9386 w,
9387 eps,
9388 norm_style,
9389 &mut post[bi * hs..(bi + 1) * hs],
9390 );
9391 }
9392 ffn.copy_from_slice(&post);
9393 }
9394 for (dst, &f) in h.iter_mut().zip(&ffn) {
9395 *dst += f;
9396 }
9397 if let Some(sc) = lw.layer_scale {
9398 for v in h.iter_mut() {
9399 *v *= sc;
9400 }
9401 }
9402 if self.layer_dump.is_some() {
9404 for bi in 0..b {
9405 self.dump_layer_row(start_pos + bi, li, &h[bi * hs..(bi + 1) * hs]);
9406 }
9407 }
9408 if let Ok(tp) = std::env::var("CMF_TRACE_POS") {
9409 if let Ok(t) = tp.parse::<usize>() {
9410 if t >= start_pos && t < start_pos + b {
9411 let bi = t - start_pos;
9412 let row = &h[bi * hs..(bi + 1) * hs];
9413 let n: f32 = row.iter().map(|x| x * x).sum::<f32>().sqrt();
9414 eprintln!(
9415 "BATCH pos {t} after layer {li}: |h| = {n:.6} h0 {:.6} h1 {:.6}",
9416 row[0], row[1]
9417 );
9418 }
9419 }
9420 }
9421 if std::env::var("CMF_DEBUG_LAYERS").is_ok() {
9425 let row = &h[(b - 1) * hs..b * hs];
9426 let rms =
9427 (row.iter().map(|&v| (v as f64) * (v as f64)).sum::<f64>() / hs as f64).sqrt();
9428 let mx = row.iter().fold(0f32, |m, &v| m.max(v.abs()));
9429 eprintln!(
9430 "layer {li:>3} {:>10} ffn={:<5} rms={rms:>12.4} max={mx:>12.4}",
9431 match &self.weights.layers[self.phys_layer(li)].attn {
9432 AttnKind::LinearGdn(_) => "gdn",
9433 AttnKind::Linear(_) => "vmf",
9434 AttnKind::ShortConv(_) => "conv",
9435 _ => "attn",
9436 },
9437 match &lw.ffn {
9438 FfnKind::Moe(_) => "moe",
9439 FfnKind::Dense(_) => "dense",
9440 FfnKind::DenseMoe(_) => "dense+moe",
9441 },
9442 );
9443 }
9444 if self.is_loop_end(li) && li + 1 < self.num_layers {
9446 for bi in 0..b {
9447 let normed = inference::rms_norm(
9448 &h[bi * hs..(bi + 1) * hs],
9449 &self.weights.final_norm,
9450 eps,
9451 norm_style,
9452 );
9453 h[bi * hs..(bi + 1) * hs].copy_from_slice(&normed);
9454 }
9455 }
9456 if std::env::var("CMF_TRACE_H").is_ok() {
9457 let n = h[..hs].iter().map(|v| v.abs()).sum::<f32>() / hs as f32;
9458 let mx = h[..hs].iter().fold(0.0f32, |a, &v| a.max(v.abs()));
9459 eprintln!(
9460 "layer {li}: mean|h|={n:.4} max|h|={mx:.2} scale={:?}",
9461 lw.layer_scale
9462 );
9463 }
9464 }
9465 crate::gpu::set_layer(-1); self.o1_progress();
9471 h
9472 }
9473
9474 fn embed_single(&self, id: u32) -> Vec<f32> {
9476 let mut out = vec![0.0f32; self.hidden_size];
9477 if (id as usize) < self.weights.embed_tokens.rows() {
9478 self.weights.embed_tokens.row_f32(id as usize, &mut out);
9479 }
9480 if self.embed_multiplier != 1.0 {
9481 for v in out.iter_mut() {
9482 *v *= self.embed_multiplier;
9483 }
9484 }
9485 if self.dsv4.is_some()
9489 || self.dsv41.is_some()
9490 || self.qwen4_exp.is_some()
9491 {
9492 let mut v = vec![0.0f32; self.hidden_size.max(1)];
9493 v[0] = id as f32;
9494 return v;
9495 }
9496 if let Some(b) = &self.g3n {
9499 return b.0.extend_embedding(id, &out, self.pool.as_deref());
9500 }
9501 out
9502 }
9503
9504 #[cfg(target_os = "macos")]
9510 fn chunk_run_gpu(
9511 &mut self,
9512 li0: usize,
9513 h: &mut [f32],
9514 b: usize,
9515 pos0: usize,
9516 embed_ids: Option<&[u32]>,
9517 cap: usize,
9518 ) -> usize {
9519 if !crate::gpu::enabled_here()
9523 || std::env::var("CMF_GPU_CHUNK")
9524 .map(|v| v == "0")
9525 .unwrap_or(false)
9526 || b < 32
9527 || self.swa.is_some()
9528 || self.global_attn.is_some()
9529 || self.graph_attn_decline_reason().is_some()
9531 || self.o1_active()
9534 || self.attn_v_norm
9535 || (self.attn_scale - 1.0 / (self.head_dim as f32).sqrt()).abs() > 1e-9
9536 {
9537 return li0;
9538 }
9539 let Some(model) = self.model.clone() else {
9540 return li0;
9541 };
9542 let inv_freq = self.inv_freq.clone();
9543 let (nh, nkv, hd, hs) = (
9544 self.num_heads,
9545 self.num_kv_heads,
9546 self.head_dim,
9547 self.hidden_size,
9548 );
9549 let loop_end = if self.loop_final_norm {
9553 ((li0 / self.physical_layers) + 1) * self.physical_layers
9554 } else {
9555 self.num_layers
9556 };
9557 let mut layers: Vec<crate::gpu_metal::ChunkLayer> = Vec::new();
9558 let mut stored_at: Vec<usize> = Vec::new();
9559 for li in li0..self.num_layers.min(loop_end).min(cap) {
9560 let lw = &self.weights.layers[self.phys_layer(li)];
9561 if lw.attn_out_norm.is_some() || lw.ffn_out_norm.is_some() || lw.layer_scale.is_some() {
9562 break;
9563 }
9564 let AttnKind::Full {
9565 wq,
9566 wk,
9567 wv,
9568 wo,
9569 q_norm,
9570 k_norm,
9571 output_gate: false,
9572 softplus_gate: None,
9573 bias,
9574 } = &lw.attn
9575 else {
9576 break;
9577 };
9578 let FfnKind::Dense(d) = &lw.ffn else { break };
9579 if d.act != Act::Silu || !d.segs.is_empty() {
9580 break;
9581 }
9582 fn cw(t: &QTensor) -> Option<(usize, usize, usize, &[f32])> {
9587 t.q8_row_parts()
9588 .or_else(|| t.q4t_parts().map(|(i, r, c)| (i, r, c, &[][..])))
9589 .or_else(|| t.q4tp_parts().map(|(i, r, c)| (i, r, c, &[][..])))
9590 }
9591 let parts = (
9592 cw(wq),
9593 cw(wk),
9594 cw(wv),
9595 cw(wo),
9596 cw(&d.gate_proj),
9597 cw(&d.up_proj),
9598 cw(&d.down_proj),
9599 );
9600 let (Some(pq), Some(pk), Some(pv), Some(po), Some(pg), Some(pu), Some(pd)) = parts
9601 else {
9602 break;
9603 };
9604 let layer = &self.kv_cache.layers[li];
9605 if layer.mode != crate::kv_cache::KvMode::F32 || layer.o1.is_some() {
9606 break;
9607 }
9608 stored_at.push(layer.head_len(0));
9609 layers.push(crate::gpu_metal::ChunkLayer {
9610 model: &model,
9611 kv_id: self.graph_kv_id,
9612 layer: li,
9613 wq: pq,
9614 wk: pk,
9615 wv: pv,
9616 wo: po,
9617 gate: pg,
9618 up: pu,
9619 down: pd,
9620 input_norm: &lw.input_norm,
9621 post_norm: &lw.post_norm,
9622 bias: bias
9623 .as_ref()
9624 .map(|(a, bb, cc)| (a.as_slice(), bb.as_slice(), cc.as_slice())),
9625 q_norm: q_norm.as_deref(),
9626 k_norm: k_norm.as_deref(),
9627 inv_freq: &inv_freq,
9628 rd: self.rotary_dim,
9629 nh,
9630 nkv,
9631 hd,
9632 hs,
9633 inter: d.gate_proj.rows(),
9634 gemma: matches!(self.norm_style, cortiq_core::NormStyle::Gemma),
9635 late_qk_norm: self.qk_norm_after_rope,
9636 eps: self.rms_eps as f32,
9637 });
9638 }
9639 if layers.is_empty() {
9640 return li0;
9641 }
9642 let row = nkv * hd;
9643 let mut store: Vec<(Vec<f32>, Vec<f32>, Vec<f32>)> = stored_at
9644 .iter()
9645 .map(|&st| (vec![0f32; b * row], vec![0f32; b * row], vec![0f32; st + b]))
9646 .collect();
9647 let mut io: Vec<crate::gpu_metal::ChunkIo> = Vec::with_capacity(layers.len());
9648 for (i, (ok, ov, oi)) in store.iter_mut().enumerate() {
9649 let li = layers[i].layer;
9650 let layer = &self.kv_cache.layers[li];
9651 io.push(crate::gpu_metal::ChunkIo {
9652 cpu_stored: stored_at[i],
9653 cpu_k: (0..nkv).map(|g| layer.head_keys(g)).collect(),
9654 cpu_v: (0..nkv).map(|g| layer.head_values(g)).collect(),
9655 out_k: ok,
9656 out_v: ov,
9657 imp: oi,
9658 });
9659 }
9660 let n_run = layers.len();
9661 let last = layers.last().map(|l| l.layer + 1).unwrap_or(li0);
9662 let ep = embed_ids.and_then(|ids| {
9665 self.weights
9666 .embed_tokens
9667 .q8_row_parts()
9668 .map(|(idx, rows, _c, rs)| crate::gpu_metal::ChunkEmbed {
9669 idx,
9670 rows,
9671 row_scale: rs,
9672 ids,
9673 mult: self.embed_multiplier,
9674 })
9675 });
9676 if embed_ids.is_some() && ep.is_none() {
9677 return li0;
9678 }
9679 if !crate::gpu_metal::chunk_run_gpu(&layers, &mut io, h, b, pos0, ep.as_ref()) {
9680 return li0;
9681 }
9682 drop(io);
9683 drop(layers);
9684 for (i, (ok, ov, oi)) in store.iter().enumerate().take(n_run) {
9687 let li = li0 + i;
9688 let layer = &mut self.kv_cache.layers[li];
9689 for bi in 0..b {
9690 layer.append(
9691 &ok[bi * row..(bi + 1) * row],
9692 &ov[bi * row..(bi + 1) * row],
9693 &[],
9694 );
9695 }
9696 layer.accumulate_imp(oi);
9697 }
9698 last
9699 }
9700
9701 fn layer_is_local(&self, li: usize) -> bool {
9704 if let Some(layers) = &self.sliding_layers {
9705 return layers.get(li).copied().unwrap_or(false);
9706 }
9707 match self.swa {
9708 Some((_, pattern)) => (li + 1) % pattern.max(1) != 0,
9709 None => false,
9710 }
9711 }
9712
9713 fn layer_inv_freq(&self, li: usize) -> std::sync::Arc<Vec<f32>> {
9716 if self.layer_is_local(li) {
9717 if let Some(f) = &self.inv_freq_local {
9718 return f.clone();
9719 }
9720 } else if let Some(f) = &self.inv_freq_global {
9721 return f.clone();
9722 }
9723 self.inv_freq.clone()
9724 }
9725
9726 fn layer_window(&self, li: usize) -> Option<usize> {
9728 self.swa
9729 .and_then(|(w, _)| self.layer_is_local(li).then_some(w))
9730 }
9731
9732 fn layer_num_heads(&self, li: usize) -> usize {
9733 self.attention_heads_per_layer
9734 .as_ref()
9735 .and_then(|v| v.get(li).copied())
9736 .unwrap_or(self.num_heads)
9737 }
9738
9739 fn layer_rope_scale(&self, li: usize) -> f32 {
9740 if self.layer_is_local(li) {
9741 self.rope_scale_local
9742 } else {
9743 self.rope_scale
9744 }
9745 }
9746
9747 fn layer_geom(&self, li: usize) -> (usize, usize, usize) {
9750 if !self.layer_is_local(li) {
9751 if let Some((ghd, gkv)) = self.global_attn {
9752 return (gkv, ghd, ghd);
9753 }
9754 }
9755 (
9756 self.layer_num_kv_heads(li),
9757 self.head_dim,
9758 if self.layer_is_local(li) {
9759 self.rotary_dim_local.unwrap_or(self.rotary_dim)
9760 } else {
9761 self.rotary_dim
9762 },
9763 )
9764 }
9765
9766 fn layer_num_kv_heads(&self, li: usize) -> usize {
9769 self.kv_heads_per_layer
9770 .as_ref()
9771 .and_then(|v| v.get(self.phys_layer(li)).copied())
9772 .unwrap_or(self.num_kv_heads)
9773 }
9774
9775 fn layer_v_dim(&self, li: usize) -> usize {
9777 let (_, hd, _) = self.layer_geom(li);
9778 self.v_head_dim.unwrap_or(hd).min(hd)
9779 }
9780
9781 pub fn set_attn_geometry(
9790 &mut self,
9791 kv_heads_per_layer: Option<Vec<usize>>,
9792 v_head_dim: Option<usize>,
9793 ) -> Result<(), String> {
9794 if kv_heads_per_layer.is_some() || v_head_dim.is_some() {
9795 if self.global_attn.is_some() {
9796 return Err(
9797 "per-layer KV heads / v_head_dim cannot combine with Gemma-4 global \
9798 attention geometry"
9799 .into(),
9800 );
9801 }
9802 if self
9803 .weights
9804 .layers
9805 .iter()
9806 .any(|lw| matches!(lw.attn, AttnKind::Mla(_)))
9807 {
9808 return Err("per-layer KV heads / v_head_dim cannot combine with MLA".into());
9809 }
9810 }
9811 if let Some(vd) = v_head_dim {
9812 if vd == 0 || vd > self.head_dim {
9813 return Err(format!(
9814 "v_head_dim {vd} must be in 1..={} (head_dim)",
9815 self.head_dim
9816 ));
9817 }
9818 }
9819 if let Some(v) = &kv_heads_per_layer {
9820 if v.len() != self.physical_layers {
9821 return Err(format!(
9822 "kv_heads_per_layer has {} entries, expected {} layers",
9823 v.len(),
9824 self.physical_layers
9825 ));
9826 }
9827 for (li, &nkv) in v.iter().enumerate() {
9828 let is_attn = matches!(
9829 self.weights.layers.get(li).map(|lw| &lw.attn),
9830 Some(AttnKind::Full { .. }) | None
9831 );
9832 if !is_attn {
9833 continue;
9834 }
9835 let nh = self
9836 .attention_heads_per_layer
9837 .as_ref()
9838 .and_then(|h| h.get(li).copied())
9839 .unwrap_or(self.num_heads);
9840 if nkv == 0 || nh % nkv != 0 {
9841 return Err(format!(
9842 "layer {li}: {nkv} KV heads must be nonzero and divide {nh} Q heads"
9843 ));
9844 }
9845 }
9846 }
9847 self.kv_heads_per_layer = kv_heads_per_layer;
9848 self.v_head_dim = v_head_dim.filter(|&vd| vd != self.head_dim);
9849 if self.kv_heads_per_layer.is_some() {
9850 for li in 0..self.kv_cache.layers.len() {
9851 let full = matches!(
9852 self.weights
9853 .layers
9854 .get(self.phys_layer(li))
9855 .map(|lw| &lw.attn),
9856 Some(AttnKind::Full { .. })
9857 );
9858 let nkv = self.layer_num_kv_heads(li);
9859 let cache = &self.kv_cache.layers[li];
9860 if full && (cache.num_kv_heads != nkv || cache.head_dim != self.head_dim) {
9861 let sinks = cache.sinks.clone();
9862 self.kv_cache.layers[li] =
9863 crate::kv_cache::LayerKvCache::new(nkv, self.head_dim);
9864 self.kv_cache.layers[li].sinks = sinks;
9865 }
9866 }
9867 }
9868 Ok(())
9869 }
9870
9871 pub fn set_layer_sinks(&mut self, phys: usize, sinks: Vec<f32>) -> Result<(), String> {
9875 let Some(lw) = self.weights.layers.get(phys) else {
9876 return Err(format!("sinks for layer {phys}: no such layer"));
9877 };
9878 if !matches!(lw.attn, AttnKind::Full { .. }) {
9879 return Err(format!(
9880 "sinks for layer {phys}: only softmax (Full) attention layers take sinks"
9881 ));
9882 }
9883 let nh = self
9884 .attention_heads_per_layer
9885 .as_ref()
9886 .and_then(|h| h.get(phys).copied())
9887 .unwrap_or(self.num_heads);
9888 if sinks.len() != nh {
9889 return Err(format!(
9890 "sinks for layer {phys}: {} values, expected one per Q head ({nh})",
9891 sinks.len()
9892 ));
9893 }
9894 if let Some(bad) = sinks.iter().find(|v| !v.is_finite()) {
9895 return Err(format!("sinks for layer {phys}: non-finite value {bad}"));
9896 }
9897 for li in 0..self.kv_cache.layers.len() {
9898 if self.phys_layer(li) == phys {
9899 self.kv_cache.layers[li].sinks = Some(sinks.clone());
9900 }
9901 }
9902 Ok(())
9903 }
9904
9905 pub fn graph_attn_decline_reason(&self) -> Option<&'static str> {
9915 if self.kv_heads_per_layer.is_some() {
9916 return Some("per-layer KV head counts");
9917 }
9918 if self.v_head_dim.is_some_and(|vd| vd != self.head_dim) {
9919 return Some("V heads narrower than Q/K heads");
9920 }
9921 if self.kv_cache.layers.iter().any(|l| l.sinks.is_some()) {
9922 return Some("learned attention sinks");
9923 }
9924 if self.swa.is_some() || self.sliding_layers.is_some() {
9925 return Some("sliding-window layers");
9926 }
9927 None
9928 }
9929
9930 pub fn wgpu_graph_attn_decline(&self) -> Option<&'static str> {
9937 self.graph_attn_decline_reason()?;
9938 if self.global_attn.is_some() {
9939 return Some("per-layer head width (Gemma-4 global layers) with per-layer geometry");
9940 }
9941 if self.attention_heads_per_layer.is_some() {
9942 return Some("per-layer Q head counts with per-layer geometry");
9943 }
9944 if self.attn_v_norm {
9945 return Some("V norm with per-layer geometry");
9946 }
9947 if (0..self.num_layers).any(|li| self.layer_rope_scale(li) != 1.0) {
9948 return Some("scaled RoPE positions with per-layer geometry");
9949 }
9950 if self.weights.layers.iter().any(|lw| {
9951 lw.attn_out_norm.is_some() || lw.ffn_out_norm.is_some() || lw.layer_scale.is_some()
9952 }) {
9953 return Some("sandwich norms / layer scale with per-layer geometry");
9954 }
9955 if self.weights.layers.iter().any(|lw| {
9956 matches!(
9957 &lw.attn,
9958 AttnKind::Full {
9959 output_gate: true,
9960 ..
9961 }
9962 )
9963 }) && self.v_head_dim.is_some()
9964 {
9965 return Some("gated attention with V narrower than K");
9966 }
9967 if (0..self.num_layers).any(|li| {
9968 self.layer_is_local(li)
9969 && self.inv_freq_local.is_none()
9970 && self.rotary_dim_local.is_some_and(|r| r != self.rotary_dim)
9971 }) {
9972 return Some("local rotary width without a local RoPE table");
9973 }
9974 None
9975 }
9976
9977 fn graph_attn_geom(&self, li: usize) -> Option<crate::gpu::GraphAttnGeom<'_>> {
9982 self.graph_attn_decline_reason()?;
9983 let (nkv, _hd, rd) = self.layer_geom(li);
9984 let invf: &[f32] = if self.layer_is_local(li) {
9985 match &self.inv_freq_local {
9986 Some(f) => f.as_slice(),
9987 None => self.inv_freq.as_slice(),
9988 }
9989 } else {
9990 match &self.inv_freq_global {
9991 Some(f) => f.as_slice(),
9992 None => self.inv_freq.as_slice(),
9993 }
9994 };
9995 Some(crate::gpu::GraphAttnGeom {
9996 nkv,
9997 dv: self.layer_v_dim(li),
9998 rd,
9999 invf,
10000 window: self.layer_window(li),
10001 sink: self.kv_cache.layers[li].sinks.as_deref(),
10002 })
10003 }
10004
10005 #[cfg(feature = "gpu")]
10014 fn pull_lagging_host_kv(&mut self, from: usize, upto: usize, position: usize) {
10015 let kv_id = self.graph_kv_id;
10016 for li in from..upto.min(self.num_layers) {
10017 if !matches!(
10018 self.weights.layers[self.phys_layer(li)].attn,
10019 AttnKind::Full { .. }
10020 ) {
10021 continue;
10022 }
10023 let host = self.kv_cache.layers[li].seq_len;
10024 if host >= position {
10025 continue;
10026 }
10027 let Some(dev) = crate::gpu::graph_kv_stored(kv_id, li) else {
10028 continue;
10029 };
10030 let to = dev.min(position);
10031 if to <= host {
10032 continue;
10033 }
10034 let (nkv, hd) = {
10035 let c = &self.kv_cache.layers[li];
10036 (c.num_kv_heads, c.head_dim)
10037 };
10038 let Some((k, v, first_valid)) =
10039 crate::gpu::graph_kv_pull_host(kv_id, li, host, to, nkv, hd)
10040 else {
10041 continue;
10042 };
10043 let need_from = match self.layer_window(li) {
10046 Some(w) => host.max((position + 1).saturating_sub(w)),
10047 None => host,
10048 };
10049 if first_valid > need_from {
10050 tracing::warn!(
10051 "layer {li}: device KV rows {host}..{to} no longer resident \
10052 (from {first_valid}); host attention will miss them"
10053 );
10054 }
10055 let row = nkv * hd;
10056 let cache = &mut self.kv_cache.layers[li];
10057 for p in 0..to - host {
10058 cache.append(&k[p * row..(p + 1) * row], &v[p * row..(p + 1) * row], &[]);
10059 }
10060 }
10061 }
10062
10063 fn note_graph_decline(&self, site: &'static str, reason: &'static str) {
10066 let mut seen = self.graph_declines.borrow_mut();
10067 if !seen.iter().any(|&(s, r)| s == site && r == reason) {
10068 tracing::warn!("{site} declined: {reason} (CPU attention path)");
10069 seen.push((site, reason));
10070 }
10071 }
10072
10073 pub fn graph_declines(&self) -> Vec<String> {
10076 self.graph_declines
10077 .borrow()
10078 .iter()
10079 .map(|(site, reason)| format!("{site} declined: {reason} (CPU attention path)"))
10080 .collect()
10081 }
10082
10083 fn dump_layer_row(&self, pos: usize, li: usize, row: &[f32]) {
10092 let Some(dir) = &self.layer_dump else {
10093 return;
10094 };
10095 let mut bytes = Vec::with_capacity(row.len() * 4);
10096 for v in row {
10097 bytes.extend_from_slice(&v.to_le_bytes());
10098 }
10099 let path = dir.join(format!("p{pos:06}_l{li:02}.f32"));
10100 if let Err(e) = std::fs::create_dir_all(dir).and_then(|_| std::fs::write(&path, &bytes)) {
10101 use std::sync::atomic::{AtomicBool, Ordering};
10102 static SAID: AtomicBool = AtomicBool::new(false);
10103 if !SAID.swap(true, Ordering::Relaxed) {
10104 tracing::error!("CMF_LAYER_DUMP: cannot write {}: {e}", path.display());
10105 }
10106 }
10107 }
10108
10109 fn mimo_moe_prepare(&mut self) {
10112 if !self.mimo_moe.is_undecided() {
10113 return;
10114 }
10115 let slot = {
10116 let layers: Vec<(usize, &MoeFfn)> = (0..self.num_layers)
10117 .filter_map(
10118 |li| match &self.weights.layers.get(self.phys_layer(li))?.ffn {
10119 FfnKind::Moe(m) => Some((li, m)),
10120 _ => None,
10121 },
10122 )
10123 .collect();
10124 if layers.is_empty()
10127 || self.physical_layers != self.num_layers
10128 || self.gpu_plan.is_some()
10129 {
10130 crate::mimo_moe::Slot::Off
10131 } else {
10132 let graph_prefix =
10136 self.wgpu_graph_attn_decline().is_none() && crate::gpu::wgpu_graph_default();
10137 crate::mimo_moe::Slot::decide(&layers, self.num_layers, graph_prefix)
10138 }
10139 };
10140 self.mimo_moe = slot;
10141 }
10142
10143 #[cfg(test)]
10144 pub(crate) fn test_graph_kv_id(&self) -> u64 {
10145 self.graph_kv_id
10146 }
10147
10148 pub(crate) fn mimo_graph_layer_rows(
10152 &mut self,
10153 li: usize,
10154 h: &mut [f32],
10155 positions: &[usize],
10156 ) -> crate::gpu::BatchGraphOutcome {
10157 use crate::gpu::BatchGraphOutcome as Out;
10158 let b = positions.len();
10159 if !(1..=4).contains(&b)
10160 || h.len() != b * self.hidden_size
10161 || !self.mimo_moe.is_dynamic(li, true)
10162 || !crate::gpu::enabled_here()
10163 || !crate::gpu::wgpu_active()
10164 || self.o1_active()
10165 || self.physical_layers != self.num_layers
10166 || std::env::var("CMF_GPU_WGPU_GRAPH").as_deref() == Ok("0")
10170 || std::env::var("CMF_MIMO_ATTN_GRAPH").as_deref() == Ok("0")
10171 || self.wgpu_graph_attn_decline().is_some()
10172 {
10173 return Out::Declined;
10174 }
10175 let attn_started = std::time::Instant::now();
10176 let outcome = {
10177 let lw = &self.weights.layers[li];
10178 if lw.attn_out_norm.is_some() || lw.ffn_out_norm.is_some() || lw.layer_scale.is_some() {
10179 return Out::Declined;
10180 }
10181 let FfnKind::Moe(m) = &lw.ffn else {
10182 return Out::Declined;
10183 };
10184 let AttnKind::Full {
10185 wq,
10186 wk,
10187 wv,
10188 wo,
10189 q_norm,
10190 k_norm,
10191 output_gate,
10192 softplus_gate,
10193 bias,
10194 } = &lw.attn
10195 else {
10196 return Out::Declined;
10197 };
10198 if *output_gate || softplus_gate.is_some() {
10199 return Out::Declined;
10200 }
10201 let Some(model) = wq.graph_weight().map(|(m, ..)| m.clone()).or_else(|| {
10202 m.experts
10203 .first()?
10204 .gate_proj
10205 .mapped_q4tp()
10206 .map(|(m, _)| m.clone())
10207 }) else {
10208 return Out::Declined;
10209 };
10210 fn gw<'a>(
10211 t: &'a QTensor,
10212 owner: &std::sync::Arc<cortiq_core::CmfModel>,
10213 ) -> Option<crate::gpu::GraphW<'a>> {
10214 if let Some((m, idx, kind, rs)) = t.graph_weight() {
10215 if m.uid() != owner.uid() || t.has_prism_contract() {
10216 return None;
10217 }
10218 return Some(crate::gpu::GraphW {
10219 idx,
10220 kind,
10221 row_scale: rs,
10222 data: &[],
10223 prism: crate::gpu::GraphPrismOp::None,
10224 affine: false,
10225 });
10226 }
10227 t.as_f32().map(|data| crate::gpu::GraphW {
10228 idx: 0,
10229 kind: 4,
10230 row_scale: &[],
10231 data,
10232 prism: crate::gpu::GraphPrismOp::None,
10233 affine: false,
10234 })
10235 }
10236 let (Some(q), Some(k), Some(v), Some(o)) = (
10237 gw(wq, &model),
10238 gw(wk, &model),
10239 gw(wv, &model),
10240 gw(wo, &model),
10241 ) else {
10242 return Out::Declined;
10243 };
10244 let layer = crate::gpu::GraphLayer {
10245 input_norm: &lw.input_norm,
10246 post_norm: &lw.post_norm,
10247 ffn: crate::gpu::GraphFfn::AttentionOnly,
10248 attn: crate::gpu::GraphAttn::Full {
10249 wq: q,
10250 wk: k,
10251 wv: v,
10252 wo: o,
10253 q_norm: q_norm.as_deref(),
10254 k_norm: k_norm.as_deref(),
10255 late_qk_norm: self.qk_norm_after_rope,
10256 bias: bias
10257 .as_ref()
10258 .map(|(q, k, v)| (q.as_slice(), k.as_slice(), v.as_slice())),
10259 output_gate: false,
10260 cpu_k: self.kv_cache.layers[li].k_heads(),
10261 cpu_v: self.kv_cache.layers[li].v_heads(),
10262 geom: self.graph_attn_geom(li),
10263 },
10264 };
10265 let (nkv, hd, rd) = self.layer_geom(li);
10266 crate::gpu::forward_batch_graph_at(
10267 &model,
10268 self.graph_kv_id,
10269 li,
10270 &[layer],
10271 &self.inv_freq,
10272 h,
10273 self.layer_num_heads(li),
10274 nkv,
10275 hd,
10276 rd,
10277 self.hidden_size,
10278 1,
10279 positions,
10280 self.kv_cache.max_seq_len,
10281 self.norm_style == cortiq_core::NormStyle::Gemma,
10282 self.rms_eps as f32,
10283 self.attn_scale,
10284 b,
10285 &[],
10286 self.o1_epoch,
10287 None,
10288 None,
10289 )
10290 };
10291 match outcome {
10292 Out::Completed => {}
10293 Out::Declined => return Out::Declined,
10294 Out::Failed => {
10295 self.graph_failed
10296 .store(true, std::sync::atomic::Ordering::Relaxed);
10297 return Out::Failed;
10298 }
10299 }
10300 let attn_ns = attn_started.elapsed().as_nanos() as u64;
10301 let hs = self.hidden_size;
10302 let lw = &self.weights.layers[li];
10303 let FfnKind::Moe(m) = &lw.ffn else {
10304 unreachable!()
10305 };
10306 let mut post = vec![0.0; h.len()];
10307 for (x, y) in h.chunks_exact(hs).zip(post.chunks_exact_mut(hs)) {
10308 inference::rms_norm_into(x, &lw.post_norm, self.rms_eps, self.norm_style, y);
10309 }
10310 let mut ffn = if b == 1 {
10311 moe_ffn_banked(&mut self.mimo_moe, li, m, &post, self.pool.as_deref())
10312 } else {
10313 moe_ffn_banked_rows(
10314 &mut self.mimo_moe,
10315 li,
10316 m,
10317 &post,
10318 b,
10319 hs,
10320 self.pool.as_deref(),
10321 )
10322 };
10323 for (x, &f) in h.iter_mut().zip(&ffn) {
10324 *x += f;
10325 }
10326 attention::recycle_buf(&mut ffn);
10327 if self.layer_dump.is_some() {
10328 for (&pos, row) in positions.iter().zip(h.chunks_exact(hs)) {
10329 self.dump_layer_row(pos, li, row);
10330 }
10331 }
10332 crate::mimo_moe::note_attention_graph(b, attn_ns);
10333 Out::Completed
10334 }
10335
10336 fn layer_attn_plain(&self, li: usize) -> bool {
10337 self.kv_heads_per_layer.is_none()
10338 && self.v_head_dim.is_none()
10339 && self.global_attn.is_none()
10340 && self.layer_window(li).is_none()
10341 && self.kv_cache.layers[li].sinks.is_none()
10342 }
10343
10344 fn forward_layers(
10346 &mut self,
10347 hidden: &[f32],
10348 position: usize,
10349 task_mask: Option<&TaskMask>,
10350 ) -> Vec<f32> {
10351 let out = self.forward_layers_upto(hidden, position, task_mask, None);
10352 self.o1_progress();
10353 out
10354 }
10355
10356 pub fn embed_id(&self, id: u32) -> Vec<f32> {
10364 self.embed_single(id)
10365 }
10366
10367 pub fn split_supported(&self) -> Result<(), String> {
10371 if self.dsv4.is_some() {
10372 return Err(
10373 "network split: DeepSeek-V4 runs its own fused stack (not splittable yet)".into(),
10374 );
10375 }
10376 if self.dsv41.is_some() {
10377 return Err(
10378 "network split: DeepSeek-V4.1 owns the shared CED/CSA2 state (not splittable)"
10379 .into(),
10380 );
10381 }
10382 if self.qwen4_exp.is_some() {
10383 return Err(
10384 "network split: Qwen3.8-Flash-Next hyper/QSA stack is not splittable yet".into(),
10385 );
10386 }
10387 if self.g3n.is_some() {
10388 return Err(
10389 "network split: Gemma-3n runs its own AltUp stack (not splittable yet)".into(),
10390 );
10391 }
10392 Ok(())
10393 }
10394
10395 pub fn forward_span(
10400 &mut self,
10401 hidden: &[f32],
10402 position: usize,
10403 from: usize,
10404 upto: usize,
10405 task_mask: Option<&TaskMask>,
10406 ) -> Result<Vec<f32>, String> {
10407 self.split_supported()?;
10408 if from > upto || upto >= self.num_layers {
10409 return Err(format!(
10410 "forward_span: layer range {from}..={upto} outside 0..{}",
10411 self.num_layers
10412 ));
10413 }
10414 if hidden.len() != self.hidden_size {
10415 return Err(format!(
10416 "forward_span: hidden len {} ≠ hidden_size {}",
10417 hidden.len(),
10418 self.hidden_size
10419 ));
10420 }
10421 let out = self.forward_layers_span(hidden, position, task_mask, from, Some(upto));
10422 self.o1_progress();
10423 if self
10424 .graph_failed
10425 .swap(false, std::sync::atomic::Ordering::Relaxed)
10426 {
10427 self.cancel
10428 .store(false, std::sync::atomic::Ordering::Relaxed);
10429 self.clear_sequence_state();
10430 return Err("forward_span: deferred O(1) transition failed".into());
10431 }
10432 Ok(out)
10433 }
10434
10435 pub fn logits_from_hidden(&mut self, hidden: &[f32]) -> Vec<f32> {
10438 let normed = inference::rms_norm(
10439 hidden,
10440 &self.weights.final_norm,
10441 self.rms_eps,
10442 self.norm_style,
10443 );
10444 self.lm_head_forward(&normed)
10445 }
10446
10447 pub fn sample_next(&mut self, logits: &[f32], past_tokens: &[u32]) -> u32 {
10449 sampler::sample_with_scratch(
10450 logits,
10451 &self.sampler_config,
10452 past_tokens,
10453 &mut self.rng,
10454 &mut self.sampler_scratch,
10455 )
10456 }
10457
10458 pub fn reset_session(&mut self) {
10460 self.clear_sequence_state();
10461 }
10462
10463 pub fn prefill_span_ids(
10469 &mut self,
10470 ids: &[u32],
10471 start_pos: usize,
10472 upto: usize,
10473 task_mask: Option<&TaskMask>,
10474 ) -> Result<Vec<f32>, String> {
10475 self.split_supported()?;
10476 if upto >= self.num_layers {
10477 return Err(format!(
10478 "prefill_span_ids: upto {upto} outside 0..{}",
10479 self.num_layers
10480 ));
10481 }
10482 if self.can_prefill_batched() && !self.graph_prefill_preferred() {
10486 let out =
10487 self.prefill_batch_span(PrefillIn::Ids(ids), start_pos, task_mask, 0, upto + 1);
10488 self.check_o1_progress_failure("prefill_span_ids")?;
10489 Ok(out)
10490 } else {
10491 let hs = self.hidden_size;
10492 let mut out = Vec::with_capacity(ids.len() * hs);
10493 for (i, &id) in ids.iter().enumerate() {
10494 let emb = self.embed_id(id);
10495 out.extend_from_slice(&self.forward_span(
10496 &emb,
10497 start_pos + i,
10498 0,
10499 upto,
10500 task_mask,
10501 )?);
10502 }
10503 Ok(out)
10504 }
10505 }
10506
10507 pub fn prefill_span_hidden(
10510 &mut self,
10511 hidden: &[f32],
10512 start_pos: usize,
10513 from: usize,
10514 upto: usize,
10515 task_mask: Option<&TaskMask>,
10516 ) -> Result<Vec<f32>, String> {
10517 self.split_supported()?;
10518 let hs = self.hidden_size;
10519 if hidden.is_empty() || hidden.len() % hs != 0 {
10520 return Err(format!(
10521 "prefill_span_hidden: {} floats is not a multiple of hidden {hs}",
10522 hidden.len()
10523 ));
10524 }
10525 if from > upto || upto >= self.num_layers {
10526 return Err(format!(
10527 "prefill_span_hidden: layer range {from}..={upto} outside 0..{}",
10528 self.num_layers
10529 ));
10530 }
10531 if self.can_prefill_batched() && !self.graph_prefill_preferred() {
10532 let out = self.prefill_batch_span(
10533 PrefillIn::Hidden(hidden),
10534 start_pos,
10535 task_mask,
10536 from,
10537 upto + 1,
10538 );
10539 self.check_o1_progress_failure("prefill_span_hidden")?;
10540 Ok(out)
10541 } else {
10542 let b = hidden.len() / hs;
10543 let mut out = Vec::with_capacity(hidden.len());
10544 for i in 0..b {
10545 let h = self.forward_span(
10546 &hidden[i * hs..(i + 1) * hs],
10547 start_pos + i,
10548 from,
10549 upto,
10550 task_mask,
10551 )?;
10552 out.extend_from_slice(&h);
10553 }
10554 Ok(out)
10555 }
10556 }
10557
10558 fn try_token_graph_wgpu(
10562 &self,
10563 hidden: &[f32],
10564 position: usize,
10565 logits_out: &mut Vec<f32>,
10566 layers_run: &mut usize,
10567 ) -> Option<Result<Vec<f32>, ()>> {
10568 self.try_token_graph_wgpu_steps(
10569 hidden,
10570 position,
10571 logits_out,
10572 1,
10573 None,
10574 Some(layers_run),
10575 0,
10576 self.num_layers,
10577 )
10578 }
10579
10580 fn try_token_graph_wgpu_span(
10584 &self,
10585 hidden: &[f32],
10586 position: usize,
10587 logits_out: &mut Vec<f32>,
10588 from: usize,
10589 upto_excl: usize,
10590 layers_run: &mut usize,
10591 ) -> Option<Result<Vec<f32>, ()>> {
10592 self.try_token_graph_wgpu_steps(
10593 hidden,
10594 position,
10595 logits_out,
10596 1,
10597 None,
10598 Some(layers_run),
10599 from,
10600 upto_excl,
10601 )
10602 }
10603
10604 fn try_multi_burst(&self, t_next: u32, position: usize, k: usize) -> Option<Vec<u32>> {
10608 if self.o1_active() || self.attn_softcap > 0.0 {
10609 return None;
10610 }
10611 if let Some(reason) = self.wgpu_graph_attn_decline() {
10614 self.note_graph_decline("wgpu multi-burst", reason);
10615 return None;
10616 }
10617 let graph_on = crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Decode);
10618 if !graph_on || self.graph_refused() {
10619 return None;
10626 }
10627 let emb = self.embed_single(t_next);
10628 let mut lg = Vec::new();
10629 let mut ids = Vec::new();
10630 match self.try_token_graph_wgpu_steps(
10631 &emb,
10632 position,
10633 &mut lg,
10634 k,
10635 Some(&mut ids),
10636 None,
10637 0,
10638 self.num_layers,
10639 ) {
10640 Some(Ok(_)) => {}
10641 Some(Err(())) => {
10642 self.graph_failed
10647 .store(true, std::sync::atomic::Ordering::Relaxed);
10648 return None;
10649 }
10650 None => return None,
10651 }
10652 (ids.len() == k).then_some(ids)
10653 }
10654
10655 fn try_token_graph_wgpu_steps(
10659 &self,
10660 hidden: &[f32],
10661 position: usize,
10662 logits_out: &mut Vec<f32>,
10663 steps: usize,
10664 ids_out: Option<&mut Vec<u32>>,
10665 layers_run: Option<&mut usize>,
10666 from: usize,
10667 upto_excl: usize,
10668 ) -> Option<Result<Vec<f32>, ()>> {
10669 let upto_excl = match self.mimo_moe.graph_prefix_end() {
10672 Some(end) if end < upto_excl => {
10673 if steps != 1 || layers_run.is_none() || from >= end {
10674 return None;
10675 }
10676 end
10677 }
10678 _ => upto_excl,
10679 };
10680 let o1_gpu = std::env::var("CMF_O1_GPU").as_deref() == Ok("1");
10683 if (self.o1_active() && !o1_gpu) || self.attn_softcap > 0.0 {
10684 return None;
10688 }
10689 if let Some(reason) = self.wgpu_graph_attn_decline() {
10696 self.note_graph_decline("wgpu token graph", reason);
10697 return None;
10698 }
10699 let o1_views: Vec<Option<Vec<crate::nystrom::O1DeviceView<'_>>>> = (from..upto_excl)
10704 .map(|li| {
10705 if !o1_gpu {
10706 return None;
10707 }
10708 self.kv_cache.layers[self.phys_layer(li)].o1_views()
10709 })
10710 .collect();
10711 if self.o1_active() && o1_gpu {
10712 let want: usize = (from..upto_excl)
10715 .filter(|li| self.kv_cache.layers[self.phys_layer(*li)].o1.is_some())
10716 .count();
10717 let have = o1_views.iter().filter(|v| v.is_some()).count();
10718 if want == 0 || have != want {
10719 use std::sync::atomic::{AtomicUsize, Ordering};
10729 static LAST: AtomicUsize = AtomicUsize::new(usize::MAX);
10730 let code = have * 1000 + want;
10731 if LAST.swap(code, Ordering::Relaxed) != code {
10732 tracing::warn!(
10733 "o1 graph: {have} of {want} layers sealed — per-op until all seal"
10734 );
10735 }
10736 return None;
10737 }
10738 }
10739 let nh = self.num_heads;
10740 let (nkv, hd, rd) = self.layer_geom(0);
10741 let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
10742 let mut layers = Vec::with_capacity(upto_excl - from);
10743 let mut model = None;
10744 let dbg = std::env::var("CMF_GRAPH_DEBUG").is_ok();
10745 fn gw(t: &QTensor) -> Option<crate::gpu::GraphW<'_>> {
10746 if let Some((m, i, kind, rs)) = t
10747 .graph_weight()
10748 .or_else(|| t.graph_weight_descriptor())
10749 {
10750 let name = &m.tensors[i].name;
10751 let prism = if crate::prism::is_inverse_embedding(m, name) {
10752 crate::gpu::GraphPrismOp::InverseEmbedding
10753 } else if crate::prism::is_forward_weight(m, name) {
10754 crate::gpu::GraphPrismOp::Forward
10755 } else {
10756 crate::gpu::GraphPrismOp::None
10757 };
10758 return Some(crate::gpu::GraphW {
10759 idx: i,
10760 kind,
10761 row_scale: rs,
10762 data: &[],
10763 prism,
10764 affine: crate::prism::is_affine_target(m, name),
10765 });
10766 }
10767 match t.as_f32() {
10769 Some(d) => Some(crate::gpu::GraphW {
10770 idx: 0,
10771 kind: 4,
10772 row_scale: &[],
10773 data: d,
10774 prism: crate::gpu::GraphPrismOp::None,
10775 affine: false,
10776 }),
10777 None => {
10778 if std::env::var_os("CMF_BATCH_DEBUG").is_some() {
10779 eprintln!("batch graph: weight has no graph/f32 representation");
10780 }
10781 None
10782 }
10783 }
10784 }
10785 for li in from..upto_excl {
10786 let lw = &self.weights.layers[self.phys_layer(li)];
10787 if dbg {
10788 let ak = match &lw.attn {
10789 AttnKind::Mla(_) => "Mla".into(),
10790 AttnKind::Full {
10791 output_gate, bias, ..
10792 } => format!("Full gate={output_gate} bias={}", bias.is_some()),
10793 AttnKind::LinearGdn(_) => "LinearGdn".into(),
10794 AttnKind::Kda(_) => "Kda".into(),
10795 AttnKind::Linear(_) => "Linear".into(),
10796 AttnKind::ShortConv(_) => "ShortConv".into(),
10797 AttnKind::Bounded(_) => "Bounded".into(),
10798 };
10799 let fk = match &lw.ffn {
10800 FfnKind::Dense(_) => "Dense",
10801 FfnKind::Moe(_) => "Moe",
10802 FfnKind::DenseMoe(_) => "DenseMoe",
10803 };
10804 eprintln!("graph L{li}: attn={ak} ffn={fk}");
10805 }
10806 let gffn = match &lw.ffn {
10807 FfnKind::DenseMoe(_) => return None, FfnKind::Dense(d) if !d.segs.is_empty() => return None,
10811 FfnKind::Dense(d) => crate::gpu::GraphFfn::Dense {
10812 gate: gw(&d.gate_proj)?,
10813 up: gw(&d.up_proj)?,
10814 down: gw(&d.down_proj)?,
10815 },
10816 FfnKind::Moe(m) => {
10817 if m.route_tau.is_some() || m.mask.is_some() {
10825 return None;
10826 }
10827 let shared = m.shared.as_ref();
10828 let has_shared = shared.is_some();
10829 let shared_gated = matches!(shared, Some((_, Some(_))));
10830 let sgate = match shared {
10831 Some((_, Some(sg))) => gw(sg)?,
10832 _ => gw(&m.router)?,
10836 };
10837 let router = gw(&m.router)?;
10838 if router.prism != crate::gpu::GraphPrismOp::None
10844 || sgate.prism != crate::gpu::GraphPrismOp::None
10845 || router.affine
10846 || sgate.affine
10847 {
10848 tracing::warn!(
10849 "resident MoE declined: Prism/affine router or shared gate transform is not implemented"
10850 );
10851 return None;
10852 }
10853 let inter = m.experts.first()?.gate_proj.rows();
10854 let mut experts = Vec::with_capacity(m.experts.len() + 1);
10855 let mut q4tp: Option<bool> = None;
10858 let mut gu_q2: Option<bool> = None;
10861 for e in m.experts.iter().chain(shared.map(|(se, _)| se)) {
10862 if !matches!(e.act, Act::Silu)
10863 || e.gate_proj.rows() != inter
10864 || e.up_proj.rows() != inter
10865 {
10866 return None;
10867 }
10868 for expert_weight in [&e.gate_proj, &e.up_proj, &e.down_proj] {
10873 let Some((em, ei, _, _)) = expert_weight
10874 .graph_weight()
10875 .or_else(|| expert_weight.graph_weight_descriptor())
10876 else {
10877 return None;
10878 };
10879 let name = &em.tensors[ei].name;
10880 if crate::prism::is_forward_weight(em, name)
10881 || crate::prism::is_inverse_embedding(em, name)
10882 || crate::prism::is_affine_target(em, name)
10883 {
10884 tracing::warn!(
10885 "resident MoE declined: expert Prism/affine transform is not implemented"
10886 );
10887 return None;
10888 }
10889 }
10890 let (mm, gi, ui, di, is_p, is_q2) = match e.gate_proj.mapped_q4t() {
10891 Some((mm, gi)) => (
10892 mm,
10893 gi,
10894 e.up_proj.mapped_q4t()?.1,
10895 e.down_proj.mapped_q4t()?.1,
10896 false,
10897 false,
10898 ),
10899 None => match e.gate_proj.mapped_q2tp() {
10900 Some((mm, gi)) => (
10901 mm,
10902 gi,
10903 e.up_proj.mapped_q2tp()?.1,
10904 e.down_proj.mapped_q4tp()?.1,
10905 true,
10906 true,
10907 ),
10908 None => {
10909 let (mm, gi) = e.gate_proj.mapped_q4tp()?;
10910 (
10911 mm,
10912 gi,
10913 e.up_proj.mapped_q4tp()?.1,
10914 e.down_proj.mapped_q4tp()?.1,
10915 true,
10916 false,
10917 )
10918 }
10919 },
10920 };
10921 if *q4tp.get_or_insert(is_p) != is_p || *gu_q2.get_or_insert(is_q2) != is_q2
10922 {
10923 tracing::warn!(
10929 "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."
10930 );
10931 return None;
10932 }
10933 model.get_or_insert_with(|| mm.clone());
10934 experts.push((gi, ui, di));
10935 }
10936 crate::gpu::GraphFfn::Moe {
10937 router,
10938 shared_gate: sgate,
10939 experts,
10940 n_exp: m.experts.len(),
10941 top_k: std::env::var("CMF_TOPK_PROBE")
10947 .ok()
10948 .and_then(|v| v.parse::<usize>().ok())
10949 .filter(|k| *k > 0 && *k <= m.top_k)
10950 .unwrap_or(m.top_k),
10951 inter,
10952 norm_topk: m.norm_topk_prob,
10953 q4tp: q4tp?,
10954 gu_q2: gu_q2.unwrap_or(false),
10955 sigmoid: m.router_sigmoid,
10956 bias: m.expert_bias.as_deref(),
10957 has_shared,
10958 shared_gated,
10959 route_scale: m.routed_scaling,
10960 }
10961 }
10962 };
10963 let attn = match &lw.attn {
10964 AttnKind::Full {
10965 wq,
10966 wk,
10967 wv,
10968 wo,
10969 q_norm,
10970 k_norm,
10971 output_gate,
10972 softplus_gate,
10973 bias,
10974 } => {
10975 if softplus_gate.is_some() || self.attention_heads_per_layer.is_some() {
10976 return None;
10977 }
10978 let (m, _, _, _) = wq
10979 .graph_weight()
10980 .or_else(|| wq.graph_weight_descriptor())?;
10981 model = Some(m.clone());
10982 crate::gpu::GraphAttn::Full {
10983 wq: gw(wq)?,
10984 wk: gw(wk)?,
10985 wv: gw(wv)?,
10986 wo: gw(wo)?,
10987 q_norm: q_norm.as_deref(),
10988 k_norm: k_norm.as_deref(),
10989 late_qk_norm: self.qk_norm_after_rope,
10990 bias: bias
10991 .as_ref()
10992 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
10993 output_gate: *output_gate,
10994 cpu_k: self.kv_cache.layers[li].k_heads(),
10995 cpu_v: self.kv_cache.layers[li].v_heads(),
10996 geom: self.graph_attn_geom(li),
10997 }
10998 }
10999 AttnKind::LinearGdn(w) => {
11000 let cfg = self.gdn_cfg?;
11001 let (m, _, _, _) = w
11002 .in_proj_qkv
11003 .graph_weight()
11004 .or_else(|| w.in_proj_qkv.graph_weight_descriptor())?;
11005 model = Some(m.clone());
11006 crate::gpu::GraphAttn::Gdn {
11007 qkv: gw(&w.in_proj_qkv)?,
11008 z: gw(&w.in_proj_z)?,
11009 a: gw(&w.in_proj_a)?,
11010 b: gw(&w.in_proj_b)?,
11011 out: gw(&w.out_proj)?,
11012 conv1d: &w.conv1d,
11013 a_log: &w.a_log,
11014 dt_bias: &w.dt_bias,
11015 norm: &w.norm,
11016 nv: cfg.num_v_heads,
11017 nk: cfg.num_k_heads,
11018 dk: cfg.key_head_dim,
11019 dv: cfg.value_head_dim,
11020 kk: cfg.conv_kernel,
11021 cpu_state: &self.kv_cache.layers[self.phys_layer(li)].linear_state,
11022 }
11023 }
11024 AttnKind::ShortConv(w) => {
11025 let cfg = self.short_conv_cfg?;
11026 let (m, _, _, _) = w
11027 .in_proj
11028 .graph_weight()
11029 .or_else(|| w.in_proj.graph_weight_descriptor())?;
11030 model = Some(m.clone());
11031 crate::gpu::GraphAttn::ShortConv {
11032 inp: gw(&w.in_proj)?,
11033 out: gw(&w.out_proj)?,
11034 taps: &w.conv,
11035 kernel: cfg.kernel,
11036 cpu_state: &self.kv_cache.layers[self.phys_layer(li)].linear_state,
11037 }
11038 }
11039 _ => return None,
11040 };
11041 layers.push(crate::gpu::GraphLayer {
11042 input_norm: &lw.input_norm,
11043 attn,
11044 post_norm: &lw.post_norm,
11045 ffn: gffn,
11046 });
11047 }
11048 let model = model?;
11049 let lm_gw = if upto_excl == self.num_layers
11055 && self.graph_want_logits
11056 && std::env::var("CMF_GPU_LMHEAD")
11057 .map(|v| v != "0")
11058 .unwrap_or(true)
11059 {
11060 self.weights
11061 .lm_head
11062 .graph_weight()
11063 .or_else(|| self.weights.lm_head.graph_weight_descriptor())
11064 .map(|(m, i, kind, rs)| {
11065 let name = &m.tensors[i].name;
11066 let prism = if crate::prism::is_inverse_embedding(m, name) {
11067 crate::gpu::GraphPrismOp::InverseEmbedding
11068 } else if crate::prism::is_forward_weight(m, name) {
11069 crate::gpu::GraphPrismOp::Forward
11070 } else {
11071 crate::gpu::GraphPrismOp::None
11072 };
11073 (
11074 crate::gpu::GraphW {
11075 idx: i,
11076 kind,
11077 row_scale: rs,
11078 data: &[],
11079 prism,
11080 affine: crate::prism::is_affine_target(m, name),
11081 },
11082 self.weights.lm_head.rows(),
11083 )
11084 })
11085 } else {
11086 None
11087 };
11088 let lm = lm_gw.as_ref().map(|(gw, rows)| (gw, *rows));
11089 let emb_gw = if steps > 1 {
11091 self.weights
11092 .embed_tokens
11093 .graph_weight()
11094 .or_else(|| self.weights.embed_tokens.graph_weight_descriptor())
11095 .map(|(m, i, kind, rs)| {
11096 let name = &m.tensors[i].name;
11097 let prism = if crate::prism::is_inverse_embedding(m, name) {
11098 crate::gpu::GraphPrismOp::InverseEmbedding
11099 } else if crate::prism::is_forward_weight(m, name) {
11100 crate::gpu::GraphPrismOp::Forward
11101 } else {
11102 crate::gpu::GraphPrismOp::None
11103 };
11104 (
11105 crate::gpu::GraphW {
11106 idx: i,
11107 kind,
11108 row_scale: rs,
11109 data: &[],
11110 prism,
11111 affine: crate::prism::is_affine_target(m, name),
11112 },
11113 self.weights.embed_tokens.rows(),
11114 self.embed_multiplier,
11115 )
11116 })
11117 } else {
11118 None
11119 };
11120
11121 let loop_norm_at: Vec<usize> = if self.loop_final_norm {
11127 (from..upto_excl.min(self.num_layers - 1))
11128 .filter(|&li| (li + 1) % self.physical_layers == 0)
11129 .map(|li| li - from)
11130 .collect()
11131 } else {
11132 Vec::new()
11133 };
11134 let mut h = hidden.to_vec();
11135 let dump_hidden = std::env::var_os("CMF_LOGIT_DUMP").is_some();
11141 let outcome = crate::gpu::forward_token_graph(
11142 &model,
11143 self.graph_kv_id,
11144 &layers,
11145 &o1_views,
11146 self.o1_epoch,
11147 &self.inv_freq,
11148 &mut h,
11149 nh,
11150 nkv,
11151 hd,
11152 self.attn_scale,
11153 rd,
11154 self.hidden_size,
11155 self.intermediate_size,
11156 position,
11157 self.kv_cache.max_seq_len,
11158 gemma,
11159 self.rms_eps as f32,
11160 lm,
11161 &self.weights.final_norm,
11162 logits_out,
11163 &loop_norm_at,
11164 steps,
11165 emb_gw.as_ref().map(|(gw, rows, m)| (gw, *rows, *m)),
11166 ids_out,
11167 layers_run,
11168 from,
11169 dump_hidden,
11170 );
11171 match outcome {
11172 crate::gpu::TokenGraphOutcome::Completed => Some(Ok(h)),
11173 crate::gpu::TokenGraphOutcome::Failed => Some(Err(())),
11174 crate::gpu::TokenGraphOutcome::Declined => None,
11175 }
11176 }
11177
11178 #[cfg(target_os = "macos")]
11187 #[allow(clippy::type_complexity)]
11188 fn metal_rows_plan(
11189 &self,
11190 ) -> Option<(
11191 Vec<MetalRowsItem<'_>>,
11192 std::sync::Arc<cortiq_core::CmfModel>,
11193 Option<crate::gpu_metal::GdnGpuCfg>,
11194 )> {
11195 use crate::gpu_metal::{AttnGpuLayer, GdnGpuCfg, GdnGpuLayer, MetalFfn};
11196 let graph_force = crate::gpu::q1_force() || crate::gpu::q2tp_gpu_opt_in();
11197 if !graph_force
11198 || !crate::gpu::enabled_here()
11199 || std::env::var("CMF_GPU_BLOCK")
11200 .map(|v| v == "0")
11201 .unwrap_or(false)
11202 || self.attn_softcap > 0.0
11203 || self.o1_active()
11204 || self.swa.is_some()
11205 || self.global_attn.is_some()
11206 || self.attention_heads_per_layer.is_some()
11207 || self.graph_attn_decline_reason().is_some()
11209 || self.attn_v_norm
11210 || self.loop_final_norm
11211 {
11212 return None;
11213 }
11214 let attend_contract = self.head_dim % 4 == 0
11215 && self.head_dim <= 256
11216 && self.rotary_dim >= 2
11217 && self.rotary_dim <= self.head_dim
11218 && (self.rotary_dim / 2) % 32 == 0
11219 && self.num_kv_heads > 0
11220 && self.num_heads % self.num_kv_heads == 0;
11221 if !attend_contract {
11222 return None;
11223 }
11224 let mut plan: Vec<MetalRowsItem> = Vec::new();
11225 let mut model_ref: Option<std::sync::Arc<cortiq_core::CmfModel>> = None;
11226 for li in 0..self.num_layers {
11227 let lw = &self.weights.layers[self.phys_layer(li)];
11228 if lw.attn_out_norm.is_some() || lw.ffn_out_norm.is_some() || lw.layer_scale.is_some() {
11229 return None;
11230 }
11231 let ffn = match &lw.ffn {
11232 FfnKind::Dense(d) if d.act == Act::Silu && d.segs.is_empty() => {
11233 let (Some(g), Some(u), Some(dn)) = (
11234 d.gate_proj.metal_graph_parts(),
11235 d.up_proj.metal_graph_parts(),
11236 d.down_proj.metal_graph_parts(),
11237 ) else {
11238 return None;
11239 };
11240 MetalFfn::Dense {
11241 gate: g,
11242 up: u,
11243 down: dn,
11244 }
11245 }
11246 _ => return None,
11247 };
11248 match &lw.attn {
11249 AttnKind::LinearGdn(w) if self.gdn_cfg.is_some() => {
11250 let (Some(qkv), Some(z), Some(a), Some(bb), Some(out)) = (
11251 w.in_proj_qkv.metal_graph_parts(),
11252 w.in_proj_z.metal_graph_parts(),
11253 w.in_proj_a.f32_parts(),
11254 w.in_proj_b.f32_parts(),
11255 w.out_proj.metal_graph_parts(),
11256 ) else {
11257 return None;
11258 };
11259 if let QTensor::Mapped { model, .. } = &w.in_proj_qkv {
11260 model_ref.get_or_insert_with(|| model.clone());
11261 }
11262 let gl = GdnGpuLayer {
11263 attn_norm: &lw.input_norm,
11264 post_norm: &lw.post_norm,
11265 qkv,
11266 z,
11267 a,
11268 b: bb,
11269 out,
11270 ffn,
11271 conv1d: &w.conv1d,
11272 a_log: &w.a_log,
11273 dt_bias: &w.dt_bias,
11274 gnorm: &w.norm,
11275 };
11276 match plan.last_mut() {
11277 Some(MetalRowsItem::Gdn { run, .. }) => run.push(gl),
11278 _ => plan.push(MetalRowsItem::Gdn {
11279 run: vec![gl],
11280 first: li,
11281 }),
11282 }
11283 }
11284 AttnKind::Full {
11285 wq,
11286 wk,
11287 wv,
11288 wo,
11289 q_norm,
11290 k_norm,
11291 output_gate,
11292 softplus_gate: None,
11293 bias: None,
11294 } => {
11295 let (Some(pq), Some(pk), Some(pv), Some(po)) =
11296 (
11297 wq.metal_graph_parts(),
11298 wk.metal_graph_parts(),
11299 wv.metal_graph_parts(),
11300 wo.metal_graph_parts(),
11301 )
11302 else {
11303 return None;
11304 };
11305 if let QTensor::Mapped { model, .. } = wq {
11306 model_ref.get_or_insert_with(|| model.clone());
11307 }
11308 let cache = &self.kv_cache.layers[li];
11309 if cache.mode != crate::kv_cache::KvMode::F32 || cache.o1.is_some() {
11310 return None;
11311 }
11312 plan.push(MetalRowsItem::Attn {
11313 l: AttnGpuLayer {
11314 attn_norm: &lw.input_norm,
11315 post_norm: &lw.post_norm,
11316 wq: pq,
11317 wk: pk,
11318 wv: pv,
11319 wo: po,
11320 ffn,
11321 },
11322 li,
11323 q_norm: q_norm.as_deref(),
11324 k_norm: k_norm.as_deref(),
11325 output_gate: *output_gate,
11326 });
11327 }
11328 _ => return None,
11329 }
11330 }
11331 let model = model_ref?;
11332 let gcfg = self.gdn_cfg.map(|cfg| GdnGpuCfg {
11333 nv: cfg.num_v_heads,
11334 nk: cfg.num_k_heads,
11335 dk: cfg.key_head_dim,
11336 dv: cfg.value_head_dim,
11337 kk: cfg.conv_kernel,
11338 hidden: self.hidden_size,
11339 inter: self.intermediate_size,
11340 c_dim: cfg.conv_dim(),
11341 eps: cfg.rms_eps as f32,
11342 gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
11343 });
11344 Some((plan, model, gcfg))
11345 }
11346
11347 #[cfg(target_os = "macos")]
11349 #[allow(clippy::too_many_arguments)]
11350 fn metal_attn_params<'a>(
11351 li: usize,
11352 cache: &'a crate::kv_cache::LayerKvCache,
11353 q_norm: Option<&'a [f32]>,
11354 k_norm: Option<&'a [f32]>,
11355 output_gate: bool,
11356 inv_freq: &'a [f32],
11357 geom: (usize, usize, usize, usize),
11358 pos0: usize,
11359 kv_id: u64,
11360 scale: f32,
11361 eps: f32,
11362 gemma: bool,
11363 late_qk_norm: bool,
11364 ) -> (crate::gpu_metal::AttnDeviceParams<'a>, usize) {
11365 let (nh, nkv, hd, rd) = geom;
11366 let cpu_k: Vec<&[f32]> = (0..nkv).map(|g| cache.head_keys(g)).collect();
11367 let cpu_v: Vec<&[f32]> = (0..nkv).map(|g| cache.head_values(g)).collect();
11368 let cpu_stored = cpu_k[0].len() / hd;
11369 (
11370 crate::gpu_metal::AttnDeviceParams {
11371 kv_id,
11372 layer: li,
11373 nh,
11374 nkv,
11375 hd,
11376 rd,
11377 position: pos0,
11378 scale,
11379 eps,
11380 gemma,
11381 late_qk_norm,
11382 output_gate,
11383 q_norm,
11384 k_norm,
11385 inv_freq,
11386 cpu_k,
11387 cpu_v,
11388 cpu_stored,
11389 o1: None,
11390 },
11391 cpu_stored,
11392 )
11393 }
11394
11395 #[cfg(target_os = "macos")]
11400 #[allow(clippy::type_complexity)]
11401 fn metal_rows_run(
11402 &mut self,
11403 hiddens: &mut [f32],
11404 pos0: usize,
11405 b: usize,
11406 prefill: bool,
11407 spec: Option<((usize, usize, usize), &[f32], &mut Vec<f32>)>,
11408 mut argmax_out: Option<(usize, &mut Vec<u32>)>,
11412 ) -> MetalRowsRun {
11413 use crate::gpu_metal::{GraphDims, VerifyGraph};
11414 if !crate::gpu_metal::wait_replay() {
11420 tracing::error!("Metal rows graph: the pending async replay failed");
11421 return MetalRowsRun::Failed;
11422 }
11423 spec_stamp("v.wait");
11424 let want = self.gdn_cfg.map(|c| c.state_len()).unwrap_or(0);
11430 if want > 0 {
11431 let phys = self.physical_layers.max(1);
11432 for (li, l) in self.kv_cache.layers.iter_mut().enumerate() {
11433 let is_gdn = self
11434 .weights
11435 .layers
11436 .get(li % phys)
11437 .is_some_and(|lw| matches!(lw.attn, AttnKind::LinearGdn(_)));
11438 if is_gdn && l.linear_state.len() != want {
11439 l.linear_state = vec![0f32; want];
11440 }
11441 }
11442 }
11443 let Some((plan, model, gcfg)) = self.metal_rows_plan() else {
11444 return MetalRowsRun::Declined;
11445 };
11446 spec_stamp("v.plan");
11447 let dims = GraphDims {
11448 hidden: self.hidden_size,
11449 eps: self.rms_eps as f32,
11450 gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
11451 };
11452 let Some(mut graph) = (if prefill {
11453 VerifyGraph::new_prefill(&model, dims, hiddens, b)
11454 } else {
11455 VerifyGraph::new(&model, dims, hiddens, b)
11456 }) else {
11457 return MetalRowsRun::Declined;
11458 };
11459 let geom = (
11460 self.num_heads,
11461 self.num_kv_heads,
11462 self.head_dim,
11463 self.rotary_dim,
11464 );
11465 let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
11466 let eps = self.rms_eps as f32;
11467 let kv_id = self.graph_kv_id;
11468 let inv_freq = self.inv_freq.clone();
11469 for item in &plan {
11470 let ok = match item {
11471 MetalRowsItem::Gdn { run, .. } => gcfg
11472 .as_ref()
11473 .map(|gc| run.iter().all(|l| graph.gdn_ok(l, gc)))
11474 .unwrap_or(false),
11475 MetalRowsItem::Attn {
11476 l,
11477 li,
11478 q_norm,
11479 k_norm,
11480 output_gate,
11481 } => {
11482 let (p, _) = Self::metal_attn_params(
11483 *li,
11484 &self.kv_cache.layers[*li],
11485 *q_norm,
11486 *k_norm,
11487 *output_gate,
11488 &inv_freq,
11489 geom,
11490 pos0,
11491 kv_id,
11492 self.attn_scale,
11493 eps,
11494 gemma,
11495 self.qk_norm_after_rope,
11496 );
11497 graph.attn_ok(l, &p)
11498 }
11499 };
11500 if !ok {
11501 use std::sync::atomic::{AtomicBool, Ordering};
11502 static SAID: AtomicBool = AtomicBool::new(false);
11503 if !SAID.swap(true, Ordering::Relaxed) {
11504 tracing::warn!("metal rows graph: a layer failed preflight — declining");
11505 }
11506 return MetalRowsRun::Declined;
11507 }
11508 }
11509 let lm = match &spec {
11510 Some((lm, _, _)) => {
11511 if !graph.lm_head_ok(*lm) {
11512 return MetalRowsRun::Declined;
11513 }
11514 Some(*lm)
11515 }
11516 None => None,
11517 };
11518 let mut gdn_layers = Vec::new();
11519 let mut attn_layers = Vec::new();
11520 for item in &plan {
11521 match item {
11522 MetalRowsItem::Gdn { run, first } => {
11523 let ro: Vec<&[f32]> = self.kv_cache.layers[*first..*first + run.len()]
11524 .iter()
11525 .map(|l| l.linear_state.as_slice())
11526 .collect();
11527 if !graph.encode_gdn_run_b(run, &ro, gcfg.as_ref().unwrap()) {
11528 return MetalRowsRun::Declined;
11529 }
11530 gdn_layers.extend(*first..*first + run.len());
11531 }
11532 MetalRowsItem::Attn {
11533 l,
11534 li,
11535 q_norm,
11536 k_norm,
11537 output_gate,
11538 } => {
11539 let (p, cpu_stored) = Self::metal_attn_params(
11540 *li,
11541 &self.kv_cache.layers[*li],
11542 *q_norm,
11543 *k_norm,
11544 *output_gate,
11545 &inv_freq,
11546 geom,
11547 pos0,
11548 kv_id,
11549 self.attn_scale,
11550 eps,
11551 gemma,
11552 self.qk_norm_after_rope,
11553 );
11554 if !graph.encode_attn_b(l, &p) {
11555 return MetalRowsRun::Declined;
11556 }
11557 attn_layers.push((*li, cpu_stored));
11558 }
11559 }
11560 }
11561 if let (Some(lm), Some((_, final_norm, _))) = (lm, spec.as_ref()) {
11562 if !graph.encode_lm_head_b(final_norm, lm) {
11563 return MetalRowsRun::Declined;
11564 }
11565 if let Some((n, _)) = argmax_out.as_ref() {
11570 if !graph.encode_argmax_b(*n) {
11571 argmax_out = None;
11572 }
11573 }
11574 }
11575 spec_stamp("v.enc");
11576 if !graph.sync() {
11577 return MetalRowsRun::Failed;
11578 }
11579 spec_stamp("v.gpu");
11580 match (spec, argmax_out) {
11581 (Some(_), Some((_, ids))) => {
11582 ids.resize(b, 0);
11583 if !graph.read_argmax(ids) {
11584 return MetalRowsRun::Failed;
11585 }
11586 spec_stamp("v.am");
11587 }
11588 (Some((lm, _, logits)), None) => {
11589 logits.resize(b * lm.1, 0.0);
11590 if !graph.read_logits(logits) {
11591 return MetalRowsRun::Failed;
11592 }
11593 spec_stamp("v.lg");
11594 }
11595 (None, _) => {}
11596 }
11597 if !graph.read_hidden(hiddens) {
11598 return MetalRowsRun::Failed;
11599 }
11600 spec_stamp("v.hid");
11601 MetalRowsRun::Completed(MetalVerifyPending {
11602 graph,
11603 gdn_layers,
11604 attn_layers,
11605 })
11606 }
11607
11608 #[cfg(target_os = "macos")]
11614 fn try_batch_graph_metal(
11615 &mut self,
11616 hiddens: &mut [f32],
11617 positions: &[usize],
11618 b: usize,
11619 spec: Option<((usize, usize, usize), &[f32], &mut Vec<f32>)>,
11620 argmax_out: Option<(usize, &mut Vec<u32>)>,
11621 ) -> crate::gpu::BatchGraphOutcome {
11622 let _t0 = std::time::Instant::now();
11623 if positions.len() != b
11624 || positions.windows(2).any(|w| w[1] != w[0] + 1)
11625 || hiddens.len() != b * self.hidden_size
11626 {
11627 return crate::gpu::BatchGraphOutcome::Declined;
11628 }
11629 let pending = match self.metal_rows_run(hiddens, positions[0], b, false, spec, argmax_out) {
11630 MetalRowsRun::Declined => return crate::gpu::BatchGraphOutcome::Declined,
11631 MetalRowsRun::Failed => return crate::gpu::BatchGraphOutcome::Failed,
11632 MetalRowsRun::Completed(pending) => pending,
11633 };
11634 if std::env::var("CMF_GRAPH_SPEC_TIME").is_ok() {
11635 eprintln!(
11636 "metal-verify: {:.1} ms | b={b}",
11637 _t0.elapsed().as_secs_f64() * 1e3
11638 );
11639 }
11640 self.metal_verify = Some(pending);
11641 crate::gpu::BatchGraphOutcome::Completed
11642 }
11643
11644 #[cfg(target_os = "macos")]
11649 fn prefill_rows_metal(
11650 &mut self,
11651 ids: &[u32],
11652 start_pos: usize,
11653 spec: Option<((usize, usize, usize), &[f32], &mut Vec<f32>)>,
11654 ) -> MetalPrefillOutcome {
11655 let b = ids.len();
11656 if b == 0 || b > 512 {
11657 return MetalPrefillOutcome::Declined;
11658 }
11659 METAL_PREFILL_CHUNKS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
11660 let with_head = spec.is_some();
11661 let hs = self.hidden_size;
11662 let mut hiddens = vec![0f32; b * hs];
11663 for (j, &id) in ids.iter().enumerate() {
11664 let e = self.embed_single(id);
11665 hiddens[j * hs..(j + 1) * hs].copy_from_slice(&e);
11666 }
11667 let mut pending = match self.metal_rows_run(&mut hiddens, start_pos, b, true, spec, None) {
11668 MetalRowsRun::Declined => return MetalPrefillOutcome::Declined,
11669 MetalRowsRun::Failed => {
11670 METAL_PREFILL_ERRORS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
11671 return MetalPrefillOutcome::Failed;
11672 }
11673 MetalRowsRun::Completed(pending) => pending,
11674 };
11675 let idxs = pending.gdn_layers.clone();
11677 let mut outs: Vec<&mut [f32]> = self
11678 .kv_cache
11679 .layers
11680 .iter_mut()
11681 .enumerate()
11682 .filter(|(i, _)| idxs.binary_search(i).is_ok())
11683 .map(|(_, l)| l.linear_state.as_mut_slice())
11684 .collect();
11685 if !pending.graph.finish_states(&mut outs) {
11686 METAL_PREFILL_ERRORS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
11687 return MetalPrefillOutcome::Failed;
11688 }
11689 let (nkv, hd) = (self.num_kv_heads, self.head_dim);
11690 let mut rows = Vec::with_capacity(pending.attn_layers.len());
11694 for (li, cpu_stored) in &pending.attn_layers {
11695 let mut kbuf = vec![0f32; b * nkv * hd];
11696 let mut vbuf = vec![0f32; b * nkv * hd];
11697 if !crate::gpu_metal::kv_mirror_read_rows(
11698 self.graph_kv_id,
11699 *li,
11700 nkv,
11701 hd,
11702 *cpu_stored,
11703 b,
11704 &mut kbuf,
11705 &mut vbuf,
11706 ) {
11707 METAL_PREFILL_ERRORS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
11708 return MetalPrefillOutcome::Failed;
11709 }
11710 rows.push((*li, *cpu_stored, kbuf, vbuf));
11711 }
11712 for (li, cpu_stored, kbuf, vbuf) in rows {
11713 let cache = &mut self.kv_cache.layers[li];
11714 for r in 0..b {
11715 cache.append(
11716 &kbuf[r * nkv * hd..(r + 1) * nkv * hd],
11717 &vbuf[r * nkv * hd..(r + 1) * nkv * hd],
11718 &[],
11719 );
11720 }
11721 crate::gpu_metal::kv_mirror_set_stored(self.graph_kv_id, li, cpu_stored + b);
11722 }
11723 METAL_PREFILL_ROWS.fetch_add(b as u64, std::sync::atomic::Ordering::Relaxed);
11724 if with_head {
11725 METAL_PREFILL_HEAD_ROWS.fetch_add(b as u64, std::sync::atomic::Ordering::Relaxed);
11726 }
11727 MetalPrefillOutcome::Completed(hiddens)
11728 }
11729
11730 #[cfg(target_os = "macos")]
11731 fn prefill_batch_metal(&mut self, ids: &[u32], start_pos: usize) -> MetalPrefillOutcome {
11732 self.prefill_rows_metal(ids, start_pos, None)
11733 }
11734
11735 #[cfg(target_os = "macos")]
11740 fn nll_batch_metal(&mut self, ids: &[u32], start: usize) -> MetalBatchNllOutcome {
11741 if ids.len() < 2 || self.o1_active() || self.head_clusters.is_some() {
11742 return MetalBatchNllOutcome::Declined;
11743 }
11744 let Some(lm) = self.weights.lm_head.metal_graph_parts() else {
11745 return MetalBatchNllOutcome::Declined;
11746 };
11747 let chunk = std::env::var("CMF_METAL_PREFILL_CHUNK")
11748 .ok()
11749 .and_then(|v| v.parse::<usize>().ok())
11750 .filter(|&v| (1..=512).contains(&v))
11751 .unwrap_or(32);
11752 let final_norm = self.weights.final_norm.clone();
11753 let mut nll = 0.0f64;
11754 let mut count = 0usize;
11755 let mut pos = 0usize;
11756 let mut completed = 0usize;
11757 while pos < ids.len() {
11758 let end = (pos + chunk).min(ids.len());
11759 let mut logits = Vec::new();
11760 let outcome = self.prefill_rows_metal(
11761 &ids[pos..end],
11762 pos,
11763 Some((lm, &final_norm, &mut logits)),
11764 );
11765 match outcome {
11766 MetalPrefillOutcome::Declined => {
11767 return if completed == 0 {
11768 MetalBatchNllOutcome::Declined
11769 } else {
11770 MetalBatchNllOutcome::Failed(format!(
11771 "ordinary Metal NLL batch declined after {completed} chunks"
11772 ))
11773 };
11774 }
11775 MetalPrefillOutcome::Failed => {
11776 return MetalBatchNllOutcome::Failed(
11777 "ordinary Metal NLL batch failed after admission".to_string(),
11778 );
11779 }
11780 MetalPrefillOutcome::Completed(_) => {}
11781 }
11782 completed += 1;
11783 let vocab = self.vocab_size.min(lm.1);
11784 if logits.len() != (end - pos) * lm.1 || vocab == 0 {
11785 return MetalBatchNllOutcome::Failed(
11786 "ordinary Metal NLL head returned an invalid shape".to_string(),
11787 );
11788 }
11789 for row in 0..(end - pos) {
11790 let absolute = pos + row;
11791 if absolute < start || absolute + 1 >= ids.len() {
11792 continue;
11793 }
11794 let lg = &mut logits[row * lm.1..row * lm.1 + vocab];
11795 if let Some(mu) = self.logit_multiplier {
11796 for v in lg.iter_mut() {
11797 *v *= mu;
11798 }
11799 }
11800 if let Some(c) = self.final_softcap {
11801 for v in lg.iter_mut() {
11802 *v = c * (*v / c).tanh();
11803 }
11804 }
11805 let target = ids[absolute + 1] as usize;
11806 if target >= vocab {
11807 return MetalBatchNllOutcome::Failed(format!(
11808 "target token {target} exceeds Metal head rows {vocab}"
11809 ));
11810 }
11811 let max = lg.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
11812 let lse: f64 = lg
11813 .iter()
11814 .map(|&v| ((v - max) as f64).exp())
11815 .sum::<f64>()
11816 .ln()
11817 + max as f64;
11818 nll += lse - lg[target] as f64;
11819 count += 1;
11820 }
11821 pos = end;
11822 }
11823 MetalBatchNllOutcome::Completed(nll, count)
11824 }
11825
11826 #[cfg(target_os = "macos")]
11830 fn metal_verify_commit(&mut self, a: usize) -> bool {
11831 let Some(mut pending) = self.metal_verify.take() else {
11832 return false;
11833 };
11834 let n = a + 1;
11835 let idxs = pending.gdn_layers.clone();
11837 let mut outs: Vec<&mut [f32]> = self
11838 .kv_cache
11839 .layers
11840 .iter_mut()
11841 .enumerate()
11842 .filter(|(i, _)| idxs.binary_search(i).is_ok())
11843 .map(|(_, l)| l.linear_state.as_mut_slice())
11844 .collect();
11845 if !pending.graph.commit(n, &mut outs) {
11846 return false;
11847 }
11848 spec_stamp("c.replay");
11849 let (nkv, hd) = (self.num_kv_heads, self.head_dim);
11850 let mut rows = Vec::with_capacity(pending.attn_layers.len());
11854 for (li, cpu_stored) in &pending.attn_layers {
11855 let mut kbuf = vec![0f32; n * nkv * hd];
11856 let mut vbuf = vec![0f32; n * nkv * hd];
11857 if !crate::gpu_metal::kv_mirror_read_rows(
11858 self.graph_kv_id,
11859 *li,
11860 nkv,
11861 hd,
11862 *cpu_stored,
11863 n,
11864 &mut kbuf,
11865 &mut vbuf,
11866 ) {
11867 return false;
11868 }
11869 rows.push((*li, *cpu_stored, kbuf, vbuf));
11870 }
11871 for (li, cpu_stored, kbuf, vbuf) in rows {
11872 let cache = &mut self.kv_cache.layers[li];
11873 for r in 0..n {
11874 cache.append(
11875 &kbuf[r * nkv * hd..(r + 1) * nkv * hd],
11876 &vbuf[r * nkv * hd..(r + 1) * nkv * hd],
11877 &[],
11878 );
11879 }
11880 crate::gpu_metal::kv_mirror_set_stored(self.graph_kv_id, li, cpu_stored + n);
11881 }
11882 spec_stamp("c.kv");
11883 true
11884 }
11885
11886 #[cfg(target_os = "macos")]
11893 fn mtp_warm_batch_submit(
11894 &mut self,
11895 m: &mut MtpModule,
11896 pairs: &[(&[f32], u32)],
11897 first_pos: usize,
11898 ) -> Option<MetalWarmPending> {
11899 use crate::gpu_metal::{AttnDeviceParams, AttnGpuLayer, GraphDims, MetalFfn, VerifyGraph};
11900 let b = pairs.len();
11901 if b == 0 || b > 512 || m.kv.mode != crate::kv_cache::KvMode::F32 || m.kv.o1.is_some() {
11902 return None;
11903 }
11904 let AttnKind::Full {
11905 wq,
11906 wk,
11907 wv,
11908 wo,
11909 q_norm,
11910 k_norm,
11911 output_gate,
11912 softplus_gate: None,
11913 bias: None,
11914 } = &m.layer.attn
11915 else {
11916 return None;
11917 };
11918 let FfnKind::Dense(d) = &m.layer.ffn else {
11919 return None;
11920 };
11921 if !d.segs.is_empty() {
11922 return None;
11923 }
11924 let (Some(pq), Some(pk), Some(pv), Some(po)) =
11925 (wq.q1_parts(), wk.q1_parts(), wv.q1_parts(), wo.q1_parts())
11926 else {
11927 return None;
11928 };
11929 let (Some(g), Some(u), Some(dn)) = (
11930 d.gate_proj.q1_parts(),
11931 d.up_proj.q1_parts(),
11932 d.down_proj.q1_parts(),
11933 ) else {
11934 return None;
11935 };
11936 let Some(eh) = m.eh_proj.q1_parts() else {
11937 return None;
11938 };
11939 let QTensor::Mapped { model, .. } = wq else {
11940 return None;
11941 };
11942 let model = model.clone();
11943 let hs = self.hidden_size;
11944 let mut cat = vec![0f32; b * 2 * hs];
11946 for (j, (h, tok)) in pairs.iter().enumerate() {
11947 let e = self.embed_single(*tok);
11948 let (ce, ch) = cat[j * 2 * hs..(j + 1) * 2 * hs].split_at_mut(hs);
11949 inference::rms_norm_into(&e, &m.enorm, self.rms_eps, self.norm_style, ce);
11950 inference::rms_norm_into(h, &m.hnorm, self.rms_eps, self.norm_style, ch);
11951 }
11952 let dims = GraphDims {
11953 hidden: hs,
11954 eps: self.rms_eps as f32,
11955 gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
11956 };
11957 spec_stamp("w.cat");
11958 let Some(mut graph) = VerifyGraph::new_via_proj(&model, dims, eh, &cat, b) else {
11959 return None;
11960 };
11961 spec_stamp("w.new");
11962 let l = AttnGpuLayer {
11963 attn_norm: &m.layer.input_norm,
11964 post_norm: &m.layer.post_norm,
11965 wq: pq,
11966 wk: pk,
11967 wv: pv,
11968 wo: po,
11969 ffn: MetalFfn::Dense {
11970 gate: g,
11971 up: u,
11972 down: dn,
11973 },
11974 };
11975 let (nh, nkv, hd, rd) = (
11976 self.num_heads,
11977 self.num_kv_heads,
11978 self.head_dim,
11979 self.rotary_dim,
11980 );
11981 let inv_freq = self.inv_freq.clone();
11982 let cpu_stored;
11983 {
11984 let cache = &m.kv;
11985 let cpu_k: Vec<&[f32]> = (0..nkv).map(|g| cache.head_keys(g)).collect();
11986 let cpu_v: Vec<&[f32]> = (0..nkv).map(|g| cache.head_values(g)).collect();
11987 cpu_stored = cpu_k[0].len() / hd;
11988 if cpu_stored > first_pos {
11993 spec_stamp("w.decl");
11994 return None;
11995 }
11996 let p = AttnDeviceParams {
11997 kv_id: self.mtp_kv_id(),
11998 layer: Self::MTP_LAYER_BASE,
11999 nh,
12000 nkv,
12001 hd,
12002 rd,
12003 position: first_pos,
12004 scale: self.attn_scale,
12005 eps: self.rms_eps as f32,
12006 gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
12007 late_qk_norm: self.qk_norm_after_rope,
12008 output_gate: *output_gate,
12009 q_norm: q_norm.as_deref(),
12010 k_norm: k_norm.as_deref(),
12011 inv_freq: &inv_freq,
12012 cpu_k,
12013 cpu_v,
12014 cpu_stored,
12015 o1: None,
12016 };
12017 if !graph.attn_ok(&l, &p) || !graph.encode_attn_b(&l, &p) {
12018 return None;
12019 }
12020 }
12021 spec_stamp("w.enc");
12022 if !graph.submit() {
12023 return None;
12024 }
12025 spec_stamp("w.sub");
12026 Some(MetalWarmPending {
12027 graph,
12028 cpu_stored,
12029 b,
12030 })
12031 }
12032
12033 #[cfg(target_os = "macos")]
12036 fn mtp_warm_batch_metal(
12037 &mut self,
12038 m: &mut MtpModule,
12039 pairs: &[(&[f32], u32)],
12040 first_pos: usize,
12041 ) -> bool {
12042 match self.mtp_warm_batch_submit(m, pairs, first_pos) {
12043 Some(p) => self.mtp_warm_batch_finish(m, p),
12044 None => false,
12045 }
12046 }
12047
12048 #[cfg(target_os = "macos")]
12053 fn mtp_warm_batch_finish(&mut self, m: &mut MtpModule, pending: MetalWarmPending) -> bool {
12054 let MetalWarmPending {
12055 mut graph,
12056 cpu_stored,
12057 b,
12058 } = pending;
12059 let (nkv, hd) = (self.num_kv_heads, self.head_dim);
12060 if !graph.sync() {
12061 return false;
12062 }
12063 spec_stamp("w.gpu");
12064 let mut kbuf = vec![0f32; b * nkv * hd];
12065 let mut vbuf = vec![0f32; b * nkv * hd];
12066 if !crate::gpu_metal::kv_mirror_read_rows(
12067 self.mtp_kv_id(),
12068 Self::MTP_LAYER_BASE,
12069 nkv,
12070 hd,
12071 cpu_stored,
12072 b,
12073 &mut kbuf,
12074 &mut vbuf,
12075 ) {
12076 return false;
12077 }
12078 for r in 0..b {
12079 m.kv.append(
12080 &kbuf[r * nkv * hd..(r + 1) * nkv * hd],
12081 &vbuf[r * nkv * hd..(r + 1) * nkv * hd],
12082 &[],
12083 );
12084 }
12085 crate::gpu_metal::kv_mirror_set_stored(
12086 self.mtp_kv_id(),
12087 Self::MTP_LAYER_BASE,
12088 cpu_stored + b,
12089 );
12090 spec_stamp("w.kv");
12091 true
12092 }
12093
12094 pub(crate) fn note_draft_id(&mut self, id: u32) {
12101 let cut = Self::draft_vocab_rows(usize::MAX).max(131_072);
12102 if (id as usize) >= cut {
12103 self.draft_full_streak = 16;
12104 } else {
12105 self.draft_full_streak = self.draft_full_streak.saturating_sub(1);
12106 }
12107 }
12108
12109 fn draft_head_rows(&self, head_rows: usize) -> usize {
12112 if self.draft_full_streak > 0 {
12113 head_rows
12114 } else {
12115 Self::draft_vocab_rows(head_rows)
12116 }
12117 }
12118
12119 fn draft_vocab_rows(head_rows: usize) -> usize {
12122 static N: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
12123 let n = *N.get_or_init(|| {
12124 std::env::var("CMF_DRAFT_VOCAB")
12125 .ok()
12126 .and_then(|v| v.parse().ok())
12127 .unwrap_or(65536)
12128 });
12129 if n == 0 { head_rows } else { n.min(head_rows) }
12130 }
12131
12132 #[cfg(target_os = "macos")]
12137 fn mtp_step_metal(
12138 &mut self,
12139 m: &mut MtpModule,
12140 hidden: &[f32],
12141 next_token: u32,
12142 position: usize,
12143 want_logits: bool,
12144 ) -> Option<(Vec<f32>, Vec<f32>)> {
12145 use crate::gpu_metal::{AttnDeviceParams, AttnGpuLayer, GraphDims, MetalFfn, TokenGraph};
12146 if std::env::var("CMF_MTP_GRAPH").as_deref() == Ok("0")
12147 || !crate::gpu::q1_force()
12148 || !crate::gpu::enabled_here()
12149 || self.attn_softcap > 0.0
12150 || self.attention_heads_per_layer.is_some()
12151 || m.kv.mode != crate::kv_cache::KvMode::F32
12152 || m.kv.o1.is_some()
12153 {
12154 return None;
12155 }
12156 let AttnKind::Full {
12157 wq,
12158 wk,
12159 wv,
12160 wo,
12161 q_norm,
12162 k_norm,
12163 output_gate,
12164 softplus_gate: None,
12165 bias: None,
12166 } = &m.layer.attn
12167 else {
12168 return None;
12169 };
12170 let FfnKind::Dense(d) = &m.layer.ffn else {
12171 return None;
12172 };
12173 if d.act != Act::Silu || !d.segs.is_empty() {
12174 return None;
12175 }
12176 let (pq, pk, pv, po) = (
12177 wq.q1_parts()?,
12178 wk.q1_parts()?,
12179 wv.q1_parts()?,
12180 wo.q1_parts()?,
12181 );
12182 let (g, u, dn) = (
12183 d.gate_proj.q1_parts()?,
12184 d.up_proj.q1_parts()?,
12185 d.down_proj.q1_parts()?,
12186 );
12187 let QTensor::Mapped { model, .. } = wq else {
12188 return None;
12189 };
12190 let model = model.clone();
12191 let lm = if want_logits {
12192 Some(self.weights.lm_head.q1_parts()?)
12193 } else {
12194 None
12195 };
12196 let dims = GraphDims {
12197 hidden: self.hidden_size,
12198 eps: self.rms_eps as f32,
12199 gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
12200 };
12201 let hs = self.hidden_size;
12204 let mut x = vec![0f32; hs];
12205 let mut graph = TokenGraph::new(&model, dims, &x)?;
12206 let mut folded = false;
12207 if let Some(eh) = m.eh_proj.q1_parts() {
12208 let e = self.embed_single(next_token);
12209 let mut cat = vec![0.0f32; 2 * hs];
12210 let (cat_e, cat_h) = cat.split_at_mut(hs);
12211 inference::rms_norm_into(&e, &m.enorm, self.rms_eps, self.norm_style, cat_e);
12212 inference::rms_norm_into(hidden, &m.hnorm, self.rms_eps, self.norm_style, cat_h);
12213 folded = graph.encode_input_proj(eh, &cat);
12214 }
12215 if !folded {
12216 x = self.mtp_block_input(m, hidden, next_token);
12217 graph = TokenGraph::new(&model, dims, &x)?;
12218 }
12219 spec_stamp("d.in");
12220 let l = AttnGpuLayer {
12221 attn_norm: &m.layer.input_norm,
12222 post_norm: &m.layer.post_norm,
12223 wq: pq,
12224 wk: pk,
12225 wv: pv,
12226 wo: po,
12227 ffn: MetalFfn::Dense {
12228 gate: g,
12229 up: u,
12230 down: dn,
12231 },
12232 };
12233 let (nh, nkv, hd, rd) = (
12234 self.num_heads,
12235 self.num_kv_heads,
12236 self.head_dim,
12237 self.rotary_dim,
12238 );
12239 let inv_freq = self.inv_freq.clone();
12240 {
12241 let cache = &m.kv;
12242 let cpu_k: Vec<&[f32]> = (0..nkv).map(|g| cache.head_keys(g)).collect();
12243 let cpu_v: Vec<&[f32]> = (0..nkv).map(|g| cache.head_values(g)).collect();
12244 let cpu_stored = cpu_k[0].len() / hd;
12245 let p = AttnDeviceParams {
12246 kv_id: self.mtp_kv_id(),
12247 layer: Self::MTP_LAYER_BASE,
12248 nh,
12249 nkv,
12250 hd,
12251 rd,
12252 position,
12253 scale: self.attn_scale,
12254 eps: self.rms_eps as f32,
12255 gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
12256 late_qk_norm: self.qk_norm_after_rope,
12257 output_gate: *output_gate,
12258 q_norm: q_norm.as_deref(),
12259 k_norm: k_norm.as_deref(),
12260 inv_freq: &inv_freq,
12261 cpu_k,
12262 cpu_v,
12263 cpu_stored,
12264 o1: None,
12265 };
12266 if !graph.attn_device_ok(&l, &p) || !graph.encode_attn_device(&l, &p) {
12267 return None;
12268 }
12269 }
12270 let draft_rows = if let Some(lm) = lm {
12276 self.draft_head_rows(lm.1)
12277 } else {
12278 0
12279 };
12280 if let Some(lm) = lm {
12281 if !graph.lm_head_ok(lm) {
12282 return None;
12283 }
12284 if draft_rows < lm.1 {
12285 if !graph.encode_lm_head_part(&m.final_norm, lm, draft_rows) {
12286 return None;
12287 }
12288 } else {
12289 graph.encode_lm_head(&m.final_norm, lm);
12290 }
12291 }
12292 spec_stamp("d.enc");
12293 if graph.sync_checked().is_err() {
12294 return None;
12295 }
12296 spec_stamp("d.gpu");
12297 let mut logits = Vec::new();
12298 if let Some(lm) = lm {
12299 let n_read = draft_rows.min(lm.1).min(self.vocab_size);
12300 logits = attention::take_buf(n_read);
12301 graph.read_logits(&mut logits);
12302 logits.resize(self.vocab_size, f32::NEG_INFINITY);
12304 }
12305 graph.finish(&mut x);
12306 let mut krow = attention::take_buf(nkv * hd);
12307 let mut vrow = attention::take_buf(nkv * hd);
12308 if crate::gpu_metal::kv_mirror_read_last(
12309 self.mtp_kv_id(),
12310 Self::MTP_LAYER_BASE,
12311 nkv,
12312 hd,
12313 &mut krow,
12314 &mut vrow,
12315 ) {
12316 m.kv.append(&krow, &vrow, &[]);
12317 }
12318 attention::recycle_buf(&mut krow);
12319 attention::recycle_buf(&mut vrow);
12320 spec_stamp("d.rd");
12321 Some((logits, x))
12322 }
12323
12324 fn mtp_chain_on() -> bool {
12337 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
12338 *ON.get_or_init(|| std::env::var("CMF_MTP_CHAIN").as_deref() != Ok("0"))
12339 }
12340
12341 #[cfg(target_os = "macos")]
12353 fn mtp_draft_chain_metal(
12354 &mut self,
12355 m: &mut MtpModule,
12356 hidden: &[f32],
12357 t_next: u32,
12358 position: usize,
12359 k: usize,
12360 ) -> Result<Vec<u32>, bool> {
12361 use crate::gpu_metal::{AttnDeviceParams, AttnGpuLayer, GraphDims, MetalFfn, TokenGraph};
12362 if k == 0
12363 || k > 64
12364 || !Self::mtp_chain_on()
12365 || std::env::var("CMF_MTP_GRAPH").as_deref() == Ok("0")
12366 || !crate::gpu::q1_force()
12367 || !crate::gpu::enabled_here()
12368 || self.attn_softcap > 0.0
12369 || self.attention_heads_per_layer.is_some()
12370 || m.kv.mode != crate::kv_cache::KvMode::F32
12371 || m.kv.o1.is_some()
12372 || self.dsv4.is_some()
12374 || self.dsv41.is_some()
12375 || self.qwen4_exp.is_some()
12376 || self.g3n.is_some()
12377 {
12378 return Err(false);
12379 }
12380 let AttnKind::Full {
12381 wq,
12382 wk,
12383 wv,
12384 wo,
12385 q_norm,
12386 k_norm,
12387 output_gate,
12388 softplus_gate: None,
12389 bias: None,
12390 } = &m.layer.attn
12391 else {
12392 return Err(false);
12393 };
12394 let FfnKind::Dense(d) = &m.layer.ffn else {
12395 return Err(false);
12396 };
12397 if d.act != Act::Silu || !d.segs.is_empty() {
12398 return Err(false);
12399 }
12400 let (Some(pq), Some(pk), Some(pv), Some(po)) =
12401 (wq.q1_parts(), wk.q1_parts(), wv.q1_parts(), wo.q1_parts())
12402 else {
12403 return Err(false);
12404 };
12405 let (Some(g), Some(u), Some(dn)) = (
12406 d.gate_proj.q1_parts(),
12407 d.up_proj.q1_parts(),
12408 d.down_proj.q1_parts(),
12409 ) else {
12410 return Err(false);
12411 };
12412 let (Some(eh), Some(lm)) = (m.eh_proj.q1_parts(), self.weights.lm_head.q1_parts()) else {
12413 return Err(false);
12414 };
12415 let QTensor::Mapped { model, .. } = wq else {
12416 return Err(false);
12417 };
12418 let model = model.clone();
12419 let QTensor::Mapped {
12422 model: em,
12423 idx: eidx,
12424 dtype: cortiq_core::TensorDtype::Q4TiledP,
12425 ..
12426 } = &self.weights.embed_tokens
12427 else {
12428 return Err(false);
12429 };
12430 if !std::sync::Arc::ptr_eq(em, &model)
12431 || crate::prism::is_inverse_embedding(&model, &model.tensors[*eidx].name)
12432 {
12433 return Err(false);
12434 }
12435 let embed = (
12436 *eidx,
12437 self.weights.embed_tokens.rows(),
12438 self.weights.embed_tokens.cols(),
12439 );
12440 if embed.2 != self.hidden_size || hidden.len() != self.hidden_size {
12441 return Err(false);
12442 }
12443 let dims = GraphDims {
12444 hidden: self.hidden_size,
12445 eps: self.rms_eps as f32,
12446 gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
12447 };
12448 let Some(mut graph) = TokenGraph::new(&model, dims, hidden) else {
12449 return Err(false);
12450 };
12451 if !graph.chain_embed_ok(embed) || !graph.lm_head_ok(lm) {
12452 return Err(false);
12453 }
12454 let l = AttnGpuLayer {
12455 attn_norm: &m.layer.input_norm,
12456 post_norm: &m.layer.post_norm,
12457 wq: pq,
12458 wk: pk,
12459 wv: pv,
12460 wo: po,
12461 ffn: MetalFfn::Dense {
12462 gate: g,
12463 up: u,
12464 down: dn,
12465 },
12466 };
12467 let (nh, nkv, hd, rd) = (
12468 self.num_heads,
12469 self.num_kv_heads,
12470 self.head_dim,
12471 self.rotary_dim,
12472 );
12473 let inv_freq = self.inv_freq.clone();
12474 let draft_rows = self.draft_head_rows(lm.1);
12475 let n_arg = draft_rows.min(lm.1).min(self.vocab_size);
12476 if n_arg == 0 {
12477 return Err(false);
12478 }
12479 let split = std::env::var("CMF_MTP_CHAIN_SPLIT").as_deref() == Ok("1");
12486 let t_chain = std::time::Instant::now();
12487 graph.chain_ids_init(t_next, k);
12488 let cpu_stored;
12489 {
12490 let cache = &m.kv;
12491 let cpu_k: Vec<&[f32]> = (0..nkv).map(|g| cache.head_keys(g)).collect();
12492 let cpu_v: Vec<&[f32]> = (0..nkv).map(|g| cache.head_values(g)).collect();
12493 cpu_stored = cpu_k[0].len() / hd;
12494 for j in 0..k {
12495 if !graph.encode_chain_input(
12496 embed,
12497 j as u32,
12498 &m.enorm,
12499 &m.hnorm,
12500 self.embed_multiplier,
12501 eh,
12502 ) {
12503 return Err(false);
12504 }
12505 let p = AttnDeviceParams {
12509 kv_id: self.mtp_kv_id(),
12510 layer: Self::MTP_LAYER_BASE,
12511 nh,
12512 nkv,
12513 hd,
12514 rd,
12515 position: position + j,
12516 scale: self.attn_scale,
12517 eps: self.rms_eps as f32,
12518 gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
12519 late_qk_norm: self.qk_norm_after_rope,
12520 output_gate: *output_gate,
12521 q_norm: q_norm.as_deref(),
12522 k_norm: k_norm.as_deref(),
12523 inv_freq: &inv_freq,
12524 cpu_k: cpu_k.clone(),
12525 cpu_v: cpu_v.clone(),
12526 cpu_stored: cpu_stored + j,
12527 o1: None,
12528 };
12529 if !graph.attn_device_ok(&l, &p) || !graph.encode_attn_device(&l, &p) {
12530 return Err(false);
12531 }
12532 if draft_rows < lm.1 {
12533 if !graph.encode_lm_head_part(&m.final_norm, lm, draft_rows) {
12534 return Err(false);
12535 }
12536 } else {
12537 graph.encode_lm_head(&m.final_norm, lm);
12538 }
12539 if !graph.encode_argmax(n_arg, j as u32 + 1) {
12540 return Err(false);
12541 }
12542 if split {
12543 graph.commit();
12546 }
12547 }
12548 }
12549 let t_enc = t_chain.elapsed();
12550 if graph.sync_checked().is_err() {
12551 return Err(true);
12552 }
12553 if std::env::var_os("CMF_GRAPH_SPEC_TIME").is_some() {
12554 eprintln!(
12555 "mtp-chain: encode {:.1} ms | wait {:.1} ms (k={k}, head rows {draft_rows}{})",
12556 t_enc.as_secs_f64() * 1e3,
12557 (t_chain.elapsed() - t_enc).as_secs_f64() * 1e3,
12558 if split { ", split" } else { "" }
12559 );
12560 }
12561 let mut ids = vec![0u32; k];
12562 if !graph.chain_ids_read(&mut ids) {
12563 return Err(true);
12564 }
12565 let mut kbuf = vec![0f32; k * nkv * hd];
12566 let mut vbuf = vec![0f32; k * nkv * hd];
12567 if !crate::gpu_metal::kv_mirror_read_rows(
12568 self.mtp_kv_id(),
12569 Self::MTP_LAYER_BASE,
12570 nkv,
12571 hd,
12572 cpu_stored,
12573 k,
12574 &mut kbuf,
12575 &mut vbuf,
12576 ) {
12577 return Err(true);
12578 }
12579 for r in 0..k {
12580 m.kv.append(
12581 &kbuf[r * nkv * hd..(r + 1) * nkv * hd],
12582 &vbuf[r * nkv * hd..(r + 1) * nkv * hd],
12583 &[],
12584 );
12585 }
12586 Ok(ids)
12587 }
12588
12589 fn try_batch_graph_wgpu(
12590 &self,
12591 hiddens: &mut [f32],
12592 positions: &[usize],
12593 k: usize,
12594 spec: Option<crate::gpu::SpecTail<'_>>,
12595 ) -> crate::gpu::BatchGraphOutcome {
12596 self.try_batch_graph_wgpu_prefix(hiddens, positions, k, spec, None)
12597 }
12598
12599 fn try_batch_graph_wgpu_prefix(
12604 &self,
12605 hiddens: &mut [f32],
12606 positions: &[usize],
12607 k: usize,
12608 spec: Option<crate::gpu::SpecTail<'_>>,
12609 layers_run: Option<&mut usize>,
12610 ) -> crate::gpu::BatchGraphOutcome {
12611 let graph_end = match self.mimo_moe.graph_prefix_end() {
12612 Some(end) if end < self.num_layers => {
12613 if layers_run.is_none() || spec.is_some() || end == 0 {
12614 return crate::gpu::BatchGraphOutcome::Declined;
12615 }
12616 end
12617 }
12618 _ => self.num_layers,
12619 };
12620 let _tb = std::time::Instant::now();
12621 let batch_debug = std::env::var_os("CMF_BATCH_DEBUG").is_some();
12622 if self.attn_softcap > 0.0 {
12623 return crate::gpu::BatchGraphOutcome::Declined; }
12625 if let Some(reason) = self.wgpu_graph_attn_decline() {
12628 self.note_graph_decline("wgpu batch graph", reason);
12629 return crate::gpu::BatchGraphOutcome::Declined;
12630 }
12631 let nh = self.num_heads;
12632 let (nkv, hd, rd) = self.layer_geom(0);
12633 let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
12634 fn gw(t: &QTensor) -> Option<crate::gpu::GraphW<'_>> {
12635 if let Some((m, i, kind, rs)) = t
12636 .graph_weight()
12637 .or_else(|| t.graph_weight_descriptor())
12638 {
12639 let name = &m.tensors[i].name;
12640 let prism = if crate::prism::is_inverse_embedding(m, name) {
12641 crate::gpu::GraphPrismOp::InverseEmbedding
12642 } else if crate::prism::is_forward_weight(m, name) {
12643 crate::gpu::GraphPrismOp::Forward
12644 } else {
12645 crate::gpu::GraphPrismOp::None
12646 };
12647 return Some(crate::gpu::GraphW {
12648 idx: i,
12649 kind,
12650 row_scale: rs,
12651 data: &[],
12652 prism,
12653 affine: crate::prism::is_affine_target(m, name),
12654 });
12655 }
12656 if std::env::var_os("CMF_BATCH_DEBUG").is_some() {
12657 eprintln!(
12658 "batch graph: tensor has no graph descriptor/f32 fallback rows={} cols={}",
12659 t.rows(),
12660 t.cols()
12661 );
12662 }
12663 t.as_f32().map(|d| crate::gpu::GraphW {
12664 idx: 0,
12665 kind: 4,
12666 row_scale: &[],
12667 data: d,
12668 prism: crate::gpu::GraphPrismOp::None,
12669 affine: false,
12670 })
12671 }
12672 let built: Option<(
12673 Vec<crate::gpu::GraphLayer<'_>>,
12674 std::sync::Arc<cortiq_core::CmfModel>,
12675 )> = (|| {
12676 let mut layers = Vec::with_capacity(graph_end);
12677 let mut model = None;
12678 for li in 0..graph_end {
12679 let lw = &self.weights.layers[self.phys_layer(li)];
12680 let gffn = match &lw.ffn {
12687 FfnKind::Dense(d) if !d.segs.is_empty() => {
12688 if batch_debug {
12689 eprintln!("batch graph: dense segmented FFN at layer {li}");
12690 }
12691 return None;
12692 }
12693 FfnKind::Dense(d) => crate::gpu::GraphFfn::Dense {
12694 gate: gw(&d.gate_proj)?,
12695 up: gw(&d.up_proj)?,
12696 down: gw(&d.down_proj)?,
12697 },
12698 FfnKind::Moe(m) => {
12699 if m.route_tau.is_some() || m.mask.is_some() {
12706 return None;
12707 }
12708 let shared = m.shared.as_ref();
12712 let has_shared = shared.is_some();
12713 let shared_gated = matches!(shared, Some((_, Some(_))));
12714 let sgate = match shared {
12715 Some((_, Some(sg))) => gw(sg)?,
12716 _ => gw(&m.router)?,
12720 };
12721 let router = gw(&m.router)?;
12722 if router.prism != crate::gpu::GraphPrismOp::None
12728 || router.affine
12729 || sgate.prism != crate::gpu::GraphPrismOp::None
12730 || sgate.affine
12731 {
12732 return None;
12733 }
12734 let inter = m.experts.first()?.gate_proj.rows();
12735 let mut experts = Vec::with_capacity(m.experts.len() + 1);
12736 let mut q4tp: Option<bool> = None;
12737 let mut gu_q2: Option<bool> = None;
12738 for e in m.experts.iter().chain(shared.map(|(se, _)| se)) {
12739 if !matches!(e.act, Act::Silu)
12740 || e.gate_proj.rows() != inter
12741 || e.up_proj.rows() != inter
12742 {
12743 return None;
12744 }
12745 let (mm, gi, ui, di, is_p, is_q2) = match e.gate_proj.mapped_q4t() {
12749 Some((mm, gi)) => (
12750 mm,
12751 gi,
12752 e.up_proj.mapped_q4t()?.1,
12753 e.down_proj.mapped_q4t()?.1,
12754 false,
12755 false,
12756 ),
12757 None => match e.gate_proj.mapped_q2tp() {
12758 Some((mm, gi)) => (
12759 mm,
12760 gi,
12761 e.up_proj.mapped_q2tp()?.1,
12762 e.down_proj.mapped_q4tp()?.1,
12763 true,
12764 true,
12765 ),
12766 None => {
12767 let (mm, gi) = e.gate_proj.mapped_q4tp()?;
12768 (
12769 mm,
12770 gi,
12771 e.up_proj.mapped_q4tp()?.1,
12772 e.down_proj.mapped_q4tp()?.1,
12773 true,
12774 false,
12775 )
12776 }
12777 },
12778 };
12779 if *q4tp.get_or_insert(is_p) != is_p
12780 || *gu_q2.get_or_insert(is_q2) != is_q2
12781 {
12782 return None;
12783 }
12784 if [gi, ui, di].into_iter().any(|idx| {
12785 mm.tensors
12786 .get(idx)
12787 .is_some_and(|t| {
12788 crate::prism::is_forward_weight(mm, &t.name)
12789 || crate::prism::is_affine_target(mm, &t.name)
12790 })
12791 }) {
12792 return None;
12793 }
12794 model.get_or_insert_with(|| mm.clone());
12795 experts.push((gi, ui, di));
12796 }
12797 crate::gpu::GraphFfn::Moe {
12798 router,
12799 shared_gate: sgate,
12800 experts,
12801 n_exp: m.experts.len(),
12802 top_k: m.top_k,
12803 inter,
12804 norm_topk: m.norm_topk_prob,
12805 q4tp: q4tp?,
12806 gu_q2: gu_q2.unwrap_or(false),
12807 sigmoid: m.router_sigmoid,
12808 bias: m.expert_bias.as_deref(),
12809 has_shared,
12810 shared_gated,
12811 route_scale: m.routed_scaling,
12812 }
12813 }
12814 _ => return None,
12815 };
12816 let attn = match &lw.attn {
12817 AttnKind::Full {
12818 wq,
12819 wk,
12820 wv,
12821 wo,
12822 q_norm,
12823 k_norm,
12824 output_gate,
12825 softplus_gate,
12826 bias,
12827 } => {
12828 if softplus_gate.is_some() || self.attention_heads_per_layer.is_some() {
12829 if batch_debug {
12830 eprintln!(
12831 "batch graph: unsupported Full attention gate at layer {li} softplus={} heads={}",
12832 softplus_gate.is_some(),
12833 self.attention_heads_per_layer.is_some()
12834 );
12835 }
12836 return None;
12837 }
12838 let (m, _, _, _) = wq
12839 .graph_weight()
12840 .or_else(|| wq.graph_weight_descriptor())?;
12841 model = Some(m.clone());
12842 crate::gpu::GraphAttn::Full {
12843 wq: gw(wq)?,
12844 wk: gw(wk)?,
12845 wv: gw(wv)?,
12846 wo: gw(wo)?,
12847 q_norm: q_norm.as_deref(),
12848 k_norm: k_norm.as_deref(),
12849 late_qk_norm: self.qk_norm_after_rope,
12850 bias: bias
12851 .as_ref()
12852 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
12853 output_gate: *output_gate,
12854 cpu_k: self.kv_cache.layers[li].k_heads(),
12855 cpu_v: self.kv_cache.layers[li].v_heads(),
12856 geom: self.graph_attn_geom(li),
12857 }
12858 }
12859 AttnKind::LinearGdn(w) => {
12860 let Some(cfg) = self.gdn_cfg else {
12861 if batch_debug {
12862 eprintln!("batch graph: no GDN config at layer {li}");
12863 }
12864 return None;
12865 };
12866 let (m, _, _, _) = w
12867 .in_proj_qkv
12868 .graph_weight()
12869 .or_else(|| w.in_proj_qkv.graph_weight_descriptor())?;
12870 model = Some(m.clone());
12871 crate::gpu::GraphAttn::Gdn {
12872 qkv: gw(&w.in_proj_qkv)?,
12873 z: gw(&w.in_proj_z)?,
12874 a: gw(&w.in_proj_a)?,
12875 b: gw(&w.in_proj_b)?,
12876 out: gw(&w.out_proj)?,
12877 conv1d: &w.conv1d,
12878 a_log: &w.a_log,
12879 dt_bias: &w.dt_bias,
12880 norm: &w.norm,
12881 nv: cfg.num_v_heads,
12882 nk: cfg.num_k_heads,
12883 dk: cfg.key_head_dim,
12884 dv: cfg.value_head_dim,
12885 kk: cfg.conv_kernel,
12886 cpu_state: &self.kv_cache.layers[self.phys_layer(li)].linear_state,
12887 }
12888 }
12889 _ => return None,
12890 };
12891 layers.push(crate::gpu::GraphLayer {
12892 input_norm: &lw.input_norm,
12893 attn,
12894 post_norm: &lw.post_norm,
12895 ffn: gffn,
12896 });
12897 }
12898 Some((layers, model?))
12899 })();
12900 let Some((layers, model)) = built else {
12901 {
12902 use std::sync::atomic::{AtomicBool, Ordering};
12903 static SAID: AtomicBool = AtomicBool::new(false);
12904 if !SAID.swap(true, Ordering::Relaxed) {
12905 tracing::warn!("batch graph: BUILDER refused (layer weights/kinds)");
12906 }
12907 }
12908 return crate::gpu::BatchGraphOutcome::Declined;
12909 };
12910 if std::env::var("CMF_GRAPH_SPEC_TIME").is_ok() {
12911 eprintln!("batch-build: {:.1} ms", _tb.elapsed().as_secs_f64() * 1e3);
12912 }
12913 crate::gpu::forward_batch_graph(
12914 &model,
12915 self.graph_kv_id,
12916 &layers,
12917 &self.inv_freq,
12918 hiddens,
12919 nh,
12920 nkv,
12921 hd,
12922 rd,
12923 self.hidden_size,
12924 self.intermediate_size,
12925 positions,
12926 self.kv_cache.max_seq_len,
12927 gemma,
12928 self.rms_eps as f32,
12929 self.attn_scale,
12930 k,
12931 &(0..graph_end)
12932 .map(|li| self.kv_cache.layers[self.phys_layer(li)].o1_views())
12933 .collect::<Vec<_>>(),
12934 self.o1_epoch,
12935 spec,
12936 layers_run,
12937 )
12938 }
12939
12940 fn draft_probe() -> bool {
12944 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
12945 *ON.get_or_init(|| std::env::var("CMF_DSV4_DRAFT_PROBE").is_ok_and(|v| v != "0"))
12946 }
12947
12948 #[cfg(feature = "gpu")]
12960 fn dsv4_spec_on() -> bool {
12961 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
12962 *ON.get_or_init(|| {
12963 if let Ok(v) = std::env::var("CMF_DSV4_SPEC_RUN") {
12967 return v != "0";
12968 }
12969 std::env::var("CMF_DSV4_SPEC")
12976 .map(|v| v != "0")
12977 .unwrap_or_else(|_| {
12978 crate::gpu_wgpu::DRAFT_RESERVE.load(std::sync::atomic::Ordering::Relaxed) > 0
12979 })
12980 })
12981 }
12982
12983 #[cfg(feature = "gpu")]
12990 fn dsv4_spec_step(
12991 &mut self,
12992 tip_token: u32,
12993 t_next: u32,
12994 next_pos: usize,
12995 max_extra: usize,
12996 drafted: &mut usize,
12997 accepted_ctr: &mut usize,
12998 ) -> Option<(Vec<u32>, usize)> {
12999 let t_all = std::time::Instant::now();
13000 if std::env::var("CMF_DSV4_SPEC_TIME").is_ok() {
13001 thread_local! {
13002 static LAST: std::cell::Cell<Option<std::time::Instant>> =
13003 const { std::cell::Cell::new(None) };
13004 }
13005 LAST.with(|l| {
13006 if let Some(prev) = l.get() {
13007 eprintln!(
13008 "между раундами {:.1} мс",
13009 prev.elapsed().as_secs_f64() * 1e3
13010 );
13011 }
13012 l.set(Some(std::time::Instant::now()));
13013 });
13014 }
13015 if std::env::var("CMF_DSV4_SPEC_DEBUG").is_ok() {
13016 eprintln!("spec_step: вход pos={next_pos}");
13017 }
13018 let n_layers = self.dsv4.as_ref().map(|b| b.1.len())?;
13019 let cfg = self.dsv4.as_ref().map(|b| b.2)?;
13020 if self.dspark.is_none() {
13022 let t = crate::dsv4::dspark_targets(&self.dsv4_mtp, &cfg, n_layers);
13023 if t.is_empty() {
13024 return None;
13025 }
13026 crate::dsv4::dspark_arm(&t, cfg.dim);
13027 self.dspark = Some(crate::dsv4::DsparkState::new(
13028 self.dsv4_mtp.len(),
13029 &cfg,
13030 t.len(),
13031 ));
13032 }
13033 let targets = crate::dsv4::dspark_targets(&self.dsv4_mtp, &cfg, n_layers);
13034 let pack = crate::dsv4::dspark_pack_get(&self.dsv4_mtp, &cfg);
13035 if pack.is_none() && std::env::var("CMF_DSV4_SPEC_DEBUG").is_ok() {
13036 eprintln!("spec_step: пак не построился (targets {targets:?})");
13037 }
13038 let pack = pack?;
13039 let block = crate::dsv4::dspark_block();
13040 let b_box = self.dsv4.as_mut()?;
13041 let (g, layers, st) = (&b_box.0, &b_box.1, &mut b_box.3);
13042 let ds = self.dspark.as_mut()?;
13043 let dbg = std::env::var("CMF_DSV4_SPEC_DEBUG").is_ok();
13046 if !crate::dsv4::dspark_take(&mut ds.main_hidden) && !ds.have_hidden {
13047 if dbg {
13048 eprintln!("spec_step: нет захвата");
13049 }
13050 return None;
13051 }
13052 ds.have_hidden = true;
13053 let tip_pos = next_pos.checked_sub(1)?;
13054 let draft_started = std::time::Instant::now();
13055 let mut conf = Vec::new();
13056 let props = crate::dsv4::dspark_draft_gpu(
13057 g,
13058 &self.dsv4_mtp,
13059 &cfg,
13060 ds,
13061 pack,
13062 st.kv_id,
13063 tip_token,
13064 tip_pos,
13065 self.pool.as_deref(),
13066 &mut conf,
13067 );
13068 self.dspark_draft_ns += draft_started.elapsed().as_nanos();
13069 *drafted += block;
13070 if props.is_empty() || props[0] != t_next {
13071 if dbg {
13072 eprintln!(
13073 "spec_step: черновик {} (props0={:?} t_next={t_next})",
13074 if props.is_empty() {
13075 "пуст"
13076 } else {
13077 "мимо"
13078 },
13079 props.first()
13080 );
13081 }
13082 return None;
13083 }
13084 let mut k_verify = crate::dsv4::dspark_verify_k()
13091 .min(props.len())
13092 .min(max_extra.saturating_add(1));
13093 let conf_min = {
13099 static M: std::sync::OnceLock<f32> = std::sync::OnceLock::new();
13100 *M.get_or_init(|| {
13101 std::env::var("CMF_DSPARK_CONF_MIN")
13102 .ok()
13103 .and_then(|v| v.parse().ok())
13104 .unwrap_or(0.0)
13105 })
13106 };
13107 if conf_min > 0.0 && conf.len() >= props.len() {
13108 let mut keep = 1usize;
13109 while keep < k_verify && conf.get(keep).copied().unwrap_or(0.0) >= conf_min {
13110 keep += 1;
13111 }
13112 k_verify = k_verify.min(keep.max(2));
13113 }
13114 if k_verify < 2 {
13115 return None;
13116 }
13117 let mut fed = Vec::with_capacity(k_verify);
13118 fed.push(t_next);
13119 fed.extend_from_slice(&props[1..k_verify]);
13120 let mut argmax = Vec::new();
13121 let mut logits_all = Vec::new();
13122 let mut walked = Vec::new();
13123 let txn = crate::dsv4::dsv4_verify_chunk(
13124 g,
13125 layers,
13126 &cfg,
13127 st,
13128 &fed,
13129 next_pos,
13130 &self.inv_freq,
13131 self.pool.as_deref(),
13132 &targets,
13133 &mut argmax,
13134 &mut logits_all,
13135 &mut walked,
13136 );
13137 if txn.is_none() && dbg {
13138 eprintln!("spec_step: verify отказал");
13139 }
13140 let txn = txn?;
13141 let spec_gpu_end = txn.gpu_end;
13142 let b = fed.len();
13143 let mut accepted = 1usize;
13144 while accepted < b && fed[accepted] == argmax[accepted - 1] {
13145 accepted += 1;
13146 }
13147 if std::env::var("CMF_DSV4_SPEC_FORCE_REJECT").is_ok_and(|v| v != "0") {
13152 accepted = 1;
13153 }
13154 if std::env::var("CMF_DSV4_SPEC_TRACE").is_ok() {
13155 eprintln!("spec@{next_pos}: fed={fed:?} argmax={argmax:?} accepted={accepted}");
13156 }
13157 let t_fin = std::time::Instant::now();
13158 if !crate::dsv4::dsv4_spec_finish(
13159 g,
13160 layers,
13161 &cfg,
13162 st,
13163 txn,
13164 accepted,
13165 &fed,
13166 &self.inv_freq,
13167 self.pool.as_deref(),
13168 ) {
13169 tracing::warn!("dsv4: спекулятивный откат не удался — состояние подозрительно");
13170 return None;
13171 }
13172 if std::env::var("CMF_DSV4_SPEC_TIME").is_ok() {
13173 eprintln!(
13174 "finish(k={accepted}): {:.1} мс",
13175 t_fin.elapsed().as_secs_f64() * 1e3
13176 );
13177 }
13178 *accepted_ctr += accepted - 1;
13179 let (hc, dim) = (cfg.hc_mult, cfg.dim);
13184 let dev_caps: Vec<usize> = targets
13189 .iter()
13190 .copied()
13191 .filter(|&t| t < spec_gpu_end)
13192 .collect();
13193 let mut caps_all = vec![0.0f32; dev_caps.len() * b * hc * dim];
13194 if !crate::gpu_wgpu::dsv4_spec_cap_read_all(b, dev_caps.len(), hc * dim, &mut caps_all) {
13195 return None;
13196 }
13197 for t in 0..accepted {
13198 let tip = t + 1 == accepted;
13199 for (slot, &tl) in targets.iter().enumerate() {
13200 if let Some(di) = dev_caps.iter().position(|&d| d == tl) {
13201 let lo = (di * b + t) * hc * dim;
13202 crate::dsv4::dspark_capture(
13203 &caps_all[lo..lo + hc * dim],
13204 &cfg,
13205 slot,
13206 &mut ds.main_hidden,
13207 );
13208 } else if tip
13209 && crate::dsv4::dspark_peek_slot(slot, dim, {
13210 let lo = slot * dim;
13211 &mut ds.main_hidden[lo..lo + dim]
13212 })
13213 {
13214 } else {
13219 crate::dsv4::dspark_capture(
13223 &walked[t * hc * dim..(t + 1) * hc * dim],
13224 &cfg,
13225 slot,
13226 &mut ds.main_hidden,
13227 );
13228 }
13229 }
13230 crate::dsv4::dspark_ring_append(
13231 g,
13232 &self.dsv4_mtp,
13233 &cfg,
13234 ds,
13235 next_pos + t,
13236 self.pool.as_deref(),
13237 );
13238 }
13239 let row = logits_all[(accepted - 1) * cfg.vocab..accepted * cfg.vocab].to_vec();
13240 self.graph_logits = Some(row);
13241 if std::env::var("CMF_DSV4_TRUNK_PICK_DUMP").is_ok() {
13246 crate::dsv4::trunk_freq_note(&crate::dsv4::pick_tally_take());
13247 crate::dsv4::pick_tally_arm();
13248 }
13249 if std::env::var("CMF_DSV4_SPEC_TIME").is_ok() {
13250 eprintln!(
13251 "spec_step total {:.1} мс (k={accepted})",
13252 t_all.elapsed().as_secs_f64() * 1e3
13253 );
13254 }
13255 Some((fed[1..accepted].to_vec(), next_pos + accepted))
13256 }
13257
13258 fn dspark_probe(&mut self, position: usize, token_id: u32) {
13259 if self.dsv4_mtp.is_empty() || !Self::draft_probe() {
13260 return;
13261 }
13262 let trunk_now = crate::dsv4::pick_tally_take();
13264 crate::dsv4::trunk_freq_note(&trunk_now);
13265 if !trunk_now.is_empty() {
13266 self.dspark_trunk_picks.push(trunk_now);
13267 let keep = crate::dsv4::dspark_block();
13268 if self.dspark_trunk_picks.len() > keep {
13269 self.dspark_trunk_picks.remove(0);
13270 }
13271 }
13272 for p in std::mem::take(&mut self.dspark_pending) {
13275 let Some(i) = position.checked_sub(p.0 + 1) else {
13276 continue;
13277 };
13278 let mut p = p;
13279 if i < p.1.len() {
13280 if p.2 && p.1[i] == token_id {
13281 p.3 = i + 1;
13282 } else {
13283 p.2 = false;
13284 }
13285 if i + 1 < p.1.len() {
13286 self.dspark_pending.push(p);
13287 continue;
13288 }
13289 }
13290 self.dspark_hist.push(p.3);
13291 self.dspark_real.push(token_id);
13292 }
13293 let Some(b) = &mut self.dsv4 else { return };
13294 let (g, layers, cfg) = (&b.0, &b.1, b.2);
13295 let n_layers = layers.len();
13296 if self.dspark.is_none() {
13297 let t = crate::dsv4::dspark_targets(&self.dsv4_mtp, &cfg, n_layers);
13298 if t.is_empty() {
13299 return;
13300 }
13301 eprintln!(
13302 "DSpark: захват со слоёв {t:?}, блок {}",
13303 crate::dsv4::dspark_block()
13304 );
13305 crate::dsv4::dspark_arm(&t, cfg.dim);
13306 self.dspark = Some(crate::dsv4::DsparkState::new(
13307 self.dsv4_mtp.len(),
13308 &cfg,
13309 t.len(),
13310 ));
13311 }
13312 let ds = self.dspark.as_mut().unwrap();
13313 if !crate::dsv4::dspark_take(&mut ds.main_hidden) {
13314 return; }
13316 let mut conf = Vec::new();
13317 crate::dsv4::pick_tally_arm();
13318 let draft_started = std::time::Instant::now();
13323 #[cfg(feature = "gpu")]
13324 let gpu_draft = crate::dsv4::dspark_gpu_on();
13325 #[cfg(not(feature = "gpu"))]
13326 let gpu_draft = false;
13327 let props = if gpu_draft {
13328 #[cfg(feature = "gpu")]
13329 {
13330 let kv_id = b.3.kv_id;
13331 match crate::dsv4::dspark_pack_get(&self.dsv4_mtp, &cfg) {
13332 Some(pk) => crate::dsv4::dspark_draft_gpu(
13333 g,
13334 &self.dsv4_mtp,
13335 &cfg,
13336 ds,
13337 pk,
13338 kv_id,
13339 token_id,
13340 position,
13341 self.pool.as_deref(),
13342 &mut conf,
13343 ),
13344 None => Vec::new(),
13345 }
13346 }
13347 #[cfg(not(feature = "gpu"))]
13348 Vec::new()
13349 } else {
13350 crate::gpu::cpu_scope(|| {
13351 crate::dsv4::dspark_draft(
13352 g,
13353 &self.dsv4_mtp,
13354 &cfg,
13355 ds,
13356 token_id,
13357 position,
13358 self.pool.as_deref(),
13359 &mut conf,
13360 )
13361 })
13362 };
13363 self.dspark_draft_ns += draft_started.elapsed().as_nanos();
13364 let draft_picks = crate::dsv4::pick_tally_take();
13365 crate::dsv4::dspark_freq_note(&draft_picks);
13366 crate::dsv4::pick_tally_arm();
13369 if !props.is_empty() {
13370 let (tu, tt) = {
13374 let flat: Vec<(usize, Vec<usize>)> = self
13375 .dspark_trunk_picks
13376 .iter()
13377 .flat_map(|v| v.iter().cloned())
13378 .collect();
13379 let mut per: std::collections::HashMap<usize, Vec<usize>> =
13381 std::collections::HashMap::new();
13382 for (li, picks) in flat {
13383 per.entry(li).or_default().extend(picks);
13384 }
13385 let n = per.len().max(1);
13386 let mut u = 0usize;
13387 let mut t = 0usize;
13388 for (_, v) in per {
13389 t += v.len();
13390 u += v.iter().collect::<std::collections::HashSet<_>>().len();
13391 }
13392 (u / n, t / n)
13393 };
13394 let (du, dt) = crate::dsv4::tally_unique(&draft_picks);
13395 self.dspark_exp.push((tu, tt, du, dt));
13396 self.dspark_pending.push((position, props, true, 0));
13397 }
13398 if self.dspark_hist.len() >= 8 && self.dspark_hist.len() % 8 == 0 {
13399 let n = self.dspark_hist.len() as f32;
13400 let mean: f32 = self.dspark_hist.iter().sum::<usize>() as f32 / n;
13401 let block = crate::dsv4::dspark_block();
13402 let mut at = vec![0usize; block + 1];
13403 for &k in &self.dspark_hist {
13404 at[k] += 1;
13405 }
13406 let mut surv = Vec::with_capacity(block);
13408 for i in 1..=block {
13409 let k = at[i..].iter().sum::<usize>() as f32 / n;
13410 surv.push(format!("{k:.2}"));
13411 }
13412 let distinct = self
13413 .dspark_real
13414 .iter()
13415 .collect::<std::collections::HashSet<_>>()
13416 .len();
13417 let (tu, tt, du, dt) = self.dspark_exp.iter().fold((0, 0, 0, 0), |a, b| {
13418 (a.0 + b.0, a.1 + b.1, a.2 + b.2, a.3 + b.3)
13419 });
13420 let m = self.dspark_exp.len().max(1);
13421 eprintln!(
13422 "DSpark: черновиков {}, принято в среднем {mean:.2} из {block} \
13423 (токенов за проход {:.2}), распределение {at:?}, выживание [{}]",
13424 self.dspark_hist.len(),
13425 mean + 1.0,
13426 surv.join(" ")
13427 );
13428 eprintln!(
13429 "DSpark: разных токенов {distinct} из {} (вырожденность), \
13430 эксперты ствол {}/{} на слой за {block} токенов, \
13431 черновик {}/{} за блок, draft {:.2} мс/блок",
13432 self.dspark_real.len(),
13433 tu / m,
13434 tt / m,
13435 du / m,
13436 dt / m,
13437 self.dspark_draft_ns as f64 / self.dspark_exp.len().max(1) as f64 / 1e6
13438 );
13439 }
13440 }
13441
13442 fn forward_layers_upto(
13443 &mut self,
13444 hidden: &[f32],
13445 position: usize,
13446 task_mask: Option<&TaskMask>,
13447 upto: Option<usize>,
13448 ) -> Vec<f32> {
13449 if let Some(plan) = self.gpu_plan.clone() {
13455 if upto.is_none() && plan.len() > 1 {
13456 let mut h = hidden.to_vec();
13457 for &(dev, from, upto_incl) in plan.iter() {
13458 h = crate::gpu::with_device(dev, || {
13459 self.forward_layers_span(&h, position, task_mask, from, Some(upto_incl))
13460 });
13461 }
13462 return h;
13463 }
13464 }
13465 self.forward_layers_span(hidden, position, task_mask, 0, upto)
13466 }
13467
13468 pub fn set_gpu_plan(&mut self, devices: Option<&[usize]>) -> Result<(), String> {
13473 self.set_gpu_plan_at(devices, None)
13474 }
13475
13476 pub fn set_gpu_plan_at(
13480 &mut self,
13481 devices: Option<&[usize]>,
13482 at: Option<usize>,
13483 ) -> Result<(), String> {
13484 let Some(devs) = devices.filter(|d| d.len() > 1) else {
13485 self.gpu_plan = None;
13486 return Ok(());
13487 };
13488 self.split_supported()?;
13489 let n = self.num_layers;
13490 if devs.len() > n {
13491 return Err(format!("{} devices for {n} layers", devs.len()));
13492 }
13493 if let Some(k) = at {
13494 if k == 0 || k >= n {
13495 return Err(format!("split at {k}: the model has {n} layers"));
13496 }
13497 if devs.len() == 2 {
13498 self.gpu_plan = Some(std::sync::Arc::new(vec![
13499 (devs[0], 0, k - 1),
13500 (devs[1], k, n - 1),
13501 ]));
13502 return Ok(());
13503 }
13504 return Err(format!(
13505 "an explicit split point takes exactly 2 devices, got {}",
13506 devs.len()
13507 ));
13508 }
13509 let per = n.div_ceil(devs.len());
13510 let mut plan = Vec::with_capacity(devs.len());
13511 let mut from = 0usize;
13512 for &d in devs {
13513 if from >= n {
13514 break;
13515 }
13516 let upto = (from + per - 1).min(n - 1);
13517 plan.push((d, from, upto));
13518 from = upto + 1;
13519 }
13520 self.gpu_plan = Some(std::sync::Arc::new(plan));
13521 Ok(())
13522 }
13523
13524 pub fn gpu_plan(&self) -> Option<Vec<(usize, usize, usize)>> {
13526 self.gpu_plan.as_ref().map(|p| p.as_ref().clone())
13527 }
13528
13529 fn embryo_qtensor_f32(t: &QTensor) -> Vec<f32> {
13535 if let Some(x) = t.as_f32() {
13536 return x.to_vec();
13537 }
13538 let mut out = vec![0.0; t.rows() * t.cols()];
13539 for r in 0..t.rows() {
13540 t.row_f32(r, &mut out[r * t.cols()..(r + 1) * t.cols()]);
13541 }
13542 out
13543 }
13544
13545 fn embryo_resident_eligible(&self) -> bool {
13546 if (self.vmf_cfg.is_none() && self.gdn_cfg.is_none())
13549 || self.num_layers != self.physical_layers
13550 || self.loop_final_norm
13551 || self.weights.layers.len() != self.num_layers
13552 || self.head_clusters.is_none()
13553 || self.final_softcap.is_some()
13554 || self.logit_multiplier.is_some()
13555 || self.attn_softcap != 0.0
13556 || self.mtp.is_some()
13557 || self.g3n.is_some()
13558 || self.dsv4.is_some()
13559 || self.dsv41.is_some()
13560 || self.qwen4_exp.is_some()
13561 || self.dyn_router.is_some()
13568 || self.dyn_phi_layer.is_some()
13569 || self.dyn_blend_loaded
13570 || self.o1_cfg.is_some()
13571 || self.swa.is_some()
13572 || self.sliding_layers.is_some()
13573 || self.global_attn.is_some()
13574 || self.attention_heads_per_layer.is_some()
13575 || self.attn_v_norm
13576 || self
13577 .kv_cache
13578 .layers
13579 .iter()
13580 .any(|l| l.mode != crate::kv_cache::KvMode::F32)
13581 || self.rope_scale != 1.0
13582 || self.rope_scale_local != 1.0
13583 || self.attn_scale != 1.0 / (self.head_dim as f32).sqrt()
13584 || self.hidden_size == 0
13585 || self.hidden_size > 1024
13586 || self.intermediate_size > 1024
13587 || self.num_heads == 0
13588 || self.num_kv_heads == 0
13589 || self.num_heads % self.num_kv_heads != 0
13590 || self.num_heads.saturating_mul(self.head_dim) > 1024
13591 || self.num_kv_heads.saturating_mul(self.head_dim) > 1024
13592 || self.vocab_size == 0
13593 || self.kv_cache.max_seq_len == 0
13594 || self.rotary_dim == 0
13595 || self.rotary_dim > self.head_dim
13596 || self.rotary_dim % 2 != 0
13597 || self.inv_freq.len() < self.rotary_dim / 2
13598 {
13599 return false;
13600 }
13601 if self.weights.lm_head.as_f32().is_none()
13612 || self.weights.embed_tokens.as_f32().is_none()
13613 || self.weights.lm_head.rows() < self.vocab_size
13614 || self.weights.lm_head.cols() != self.hidden_size
13615 || self.weights.embed_tokens.rows() < self.vocab_size
13616 || self.weights.embed_tokens.cols() != self.hidden_size
13617 || self.weights.final_norm.len() != self.hidden_size
13618 {
13619 return false;
13620 }
13621 if let Some(cfg) = self.vmf_cfg {
13622 if cfg.num_heads.saturating_mul(cfg.nphase) > 1024
13623 || cfg.num_heads.saturating_mul(cfg.nphase.saturating_add(1)) > 1024
13624 || cfg.num_heads.saturating_mul(cfg.value_head_dim) > 1024
13625 || cfg.state_len() == 0
13626 {
13627 return false;
13628 }
13629 }
13630 if let Some(g) = self.gdn_cfg {
13631 if g.num_v_heads == 0
13636 || g.num_k_heads == 0
13637 || g.num_v_heads % g.num_k_heads != 0
13638 || g.key_head_dim == 0
13639 || g.key_head_dim > 128
13640 || g.value_head_dim == 0
13641 || g.value_head_dim > 256
13642 || g.value_head_dim % 4 != 0
13643 || g.conv_kernel == 0
13644 || g.num_v_heads > 512
13645 || g.num_v_heads.saturating_mul(g.value_head_dim) > 1024
13646 || g.conv_dim() > 2048
13647 || g.conv_dim() % 4 != 0
13648 || g.hidden_size != self.hidden_size
13649 || g.output_gate_sigmoid
13650 || g.rms_eps != self.rms_eps
13651 || g.state_len() == 0
13652 {
13653 return false;
13654 }
13655 }
13656 let mut full_seen = false;
13657 for lw in &self.weights.layers {
13658 if lw.attn_out_norm.is_some() || lw.ffn_out_norm.is_some() || lw.layer_scale.is_some() {
13659 return false;
13660 }
13661 match &lw.attn {
13662 AttnKind::LinearGdn(w) => {
13663 let Some(g) = self.gdn_cfg else {
13664 return false;
13665 };
13666 let (nv, dv, kk) = (g.num_v_heads, g.value_head_dim, g.conv_kernel);
13667 if w.in_proj_qkv.rows() != g.conv_dim()
13668 || w.in_proj_qkv.cols() != self.hidden_size
13669 || w.in_proj_qkv.as_f32().is_none()
13670 || w.in_proj_z.rows() != nv * dv
13671 || w.in_proj_z.cols() != self.hidden_size
13672 || w.in_proj_z.as_f32().is_none()
13673 || w.in_proj_a.rows() != nv
13674 || w.in_proj_a.cols() != self.hidden_size
13675 || w.in_proj_a.as_f32().is_none()
13676 || w.in_proj_b.rows() != nv
13677 || w.in_proj_b.cols() != self.hidden_size
13678 || w.in_proj_b.as_f32().is_none()
13679 || w.conv1d.len() != g.conv_dim() * kk
13680 || w.a_log.len() != nv
13681 || w.dt_bias.len() != nv
13682 || w.norm.len() != dv
13683 || w.out_proj.rows() != self.hidden_size
13684 || w.out_proj.cols() != nv * dv
13685 || w.out_proj.as_f32().is_none()
13686 {
13687 return false;
13688 }
13689 }
13690 AttnKind::Linear(w) => {
13691 let Some(cfg) = self.vmf_cfg else {
13692 return false;
13693 };
13694 if w.thq.rows() != cfg.num_heads * cfg.nphase
13695 || w.thq.cols() != self.hidden_size
13696 || w.thq.as_f32().is_none()
13697 || w.thk.rows() != cfg.num_heads * cfg.nphase
13698 || w.thk.cols() != self.hidden_size
13699 || w.thk.as_f32().is_none()
13700 || w.v_proj.rows() != cfg.num_heads * cfg.value_head_dim
13701 || w.v_proj.cols() != self.hidden_size
13702 || w.v_proj.as_f32().is_none()
13703 || w.out_proj.rows() != self.hidden_size
13704 || w.out_proj.cols() != cfg.num_heads * cfg.value_head_dim
13705 || w.out_proj.as_f32().is_none()
13706 || w.decay.len() != cfg.num_heads * 2 * cfg.nphase
13707 {
13708 return false;
13709 }
13710 if let Some((kg, kb)) = &w.k_gate {
13711 if kg.rows() != cfg.num_heads
13712 || kg.cols() != self.hidden_size
13713 || kg.as_f32().is_none()
13714 || kb.len() != cfg.num_heads
13715 {
13716 return false;
13717 }
13718 }
13719 if let Some(conv) = &w.conv {
13720 if conv.len() % self.hidden_size != 0 || conv.len() / self.hidden_size < 2 {
13721 return false;
13722 }
13723 }
13724 }
13725 AttnKind::Full {
13726 wq,
13727 wk,
13728 wv,
13729 wo,
13730 q_norm,
13731 k_norm,
13732 output_gate,
13733 softplus_gate,
13734 bias,
13735 } => {
13736 if full_seen
13737 || q_norm.is_some()
13738 || k_norm.is_some()
13739 || *output_gate
13740 || softplus_gate.is_some()
13741 || bias.is_some()
13742 || wq.as_f32().is_none()
13743 || wk.as_f32().is_none()
13744 || wv.as_f32().is_none()
13745 || wo.as_f32().is_none()
13746 || wq.rows() != self.num_heads * self.head_dim
13747 || wk.rows() != self.num_kv_heads * self.head_dim
13748 || wv.rows() != self.num_kv_heads * self.head_dim
13749 || wq.cols() != self.hidden_size
13750 || wk.cols() != self.hidden_size
13751 || wv.cols() != self.hidden_size
13752 || wo.rows() != self.hidden_size
13753 || wo.cols() != self.num_heads * self.head_dim
13754 {
13755 return false;
13756 }
13757 full_seen = true;
13758 }
13759 AttnKind::Bounded(w) => {
13760 let Some(ac) = self.anchor_core.as_ref() else {
13763 return false;
13764 };
13765 if self.bounded_rope.is_none()
13766 || w.window != ac.window
13767 || w.sink != ac.sink
13768 || w.window == 0
13769 || w.window + w.sink > 256
13770 || w.sink_k.len() != self.num_kv_heads * w.sink * self.head_dim
13771 || w.sink_v.len() != self.num_kv_heads * w.sink * self.head_dim
13772 || w.wq.as_f32().is_none()
13773 || w.wk.as_f32().is_none()
13774 || w.wv.as_f32().is_none()
13775 || w.wo.as_f32().is_none()
13776 || w.wq.rows() != self.num_heads * self.head_dim
13777 || w.wk.rows() != self.num_kv_heads * self.head_dim
13778 || w.wv.rows() != self.num_kv_heads * self.head_dim
13779 || w.wq.cols() != self.hidden_size
13780 || w.wk.cols() != self.hidden_size
13781 || w.wv.cols() != self.hidden_size
13782 || w.wo.rows() != self.hidden_size
13783 || w.wo.cols() != self.num_heads * self.head_dim
13784 {
13785 return false;
13786 }
13787 }
13788 _ => return false,
13789 }
13790 match &lw.ffn {
13791 FfnKind::Dense(d) => {
13792 if d.act != Act::Silu
13793 || !d.segs.is_empty()
13794 || d.gate_proj.as_f32().is_none()
13795 || d.up_proj.as_f32().is_none()
13796 || d.down_proj.as_f32().is_none()
13797 || d.gate_proj.rows() != self.intermediate_size
13798 || d.gate_proj.cols() != self.hidden_size
13799 || d.up_proj.rows() != self.intermediate_size
13800 || d.up_proj.cols() != self.hidden_size
13801 || d.down_proj.rows() != self.hidden_size
13802 || d.down_proj.cols() != self.intermediate_size
13803 {
13804 return false;
13805 }
13806 }
13807 FfnKind::Moe(m) => {
13808 if m.resonance.is_none()
13809 || m.top_k != 1
13810 || m.router_sigmoid
13811 || !m.norm_topk_prob
13812 || m.expert_bias.is_some()
13813 || m.routed_scaling != 1.0
13814 || m.route_tau.is_some()
13815 || m.shared.is_none()
13816 || m.mask.is_some()
13817 || m.per_expert_scale.is_some()
13818 || m.router_input_norm
13819 || m.experts.is_empty()
13820 || m.experts.len() > 8
13821 {
13822 return false;
13823 }
13824 let r = m.resonance.as_ref().unwrap();
13825 if r.mu.len() != m.experts.len() * self.hidden_size
13826 || r.bias.len() != m.experts.len()
13827 || r.u.len() != m.experts.len() * r.k * self.hidden_size
13828 || r.k > 128
13829 {
13830 return false;
13831 }
13832 let Some((shared, gate)) = &m.shared else {
13833 return false;
13834 };
13835 if gate.is_some() || shared.act != Act::Silu || !shared.segs.is_empty() {
13836 return false;
13837 }
13838 if shared.gate_proj.as_f32().is_none()
13839 || shared.up_proj.as_f32().is_none()
13840 || shared.down_proj.as_f32().is_none()
13841 || shared.gate_proj.rows() != self.intermediate_size
13842 || shared.gate_proj.cols() != self.hidden_size
13843 || shared.up_proj.rows() != self.intermediate_size
13844 || shared.up_proj.cols() != self.hidden_size
13845 || shared.down_proj.rows() != self.hidden_size
13846 || shared.down_proj.cols() != self.intermediate_size
13847 {
13848 return false;
13849 }
13850 for e in &m.experts {
13851 if e.act != Act::Silu
13852 || !e.segs.is_empty()
13853 || e.gate_proj.as_f32().is_none()
13854 || e.up_proj.as_f32().is_none()
13855 || e.down_proj.as_f32().is_none()
13856 || e.gate_proj.rows() != self.intermediate_size
13857 || e.gate_proj.cols() != self.hidden_size
13858 || e.up_proj.rows() != self.intermediate_size
13859 || e.up_proj.cols() != self.hidden_size
13860 || e.down_proj.rows() != self.hidden_size
13861 || e.down_proj.cols() != self.intermediate_size
13862 {
13863 return false;
13864 }
13865 }
13866 }
13867 FfnKind::DenseMoe(_) => return false,
13868 }
13869 if lw.input_norm.len() != self.hidden_size || lw.post_norm.len() != self.hidden_size {
13870 return false;
13871 }
13872 }
13873 if full_seen && self.anchor_core.is_some() {
13874 return false;
13875 }
13876 full_seen || self.num_layers > 0
13877 }
13878
13879 fn embryo_resident_wanted(&self) -> bool {
13884 crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Decode)
13885 && crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Prefill)
13886 && matches!(
13887 std::env::var("CMF_EMBRYO_RESIDENT").as_deref(),
13888 Ok("1") | Ok("parallel")
13889 )
13890 && crate::gpu::enabled_here()
13891 && !self.graph_refused()
13892 && self.embryo_resident_eligible()
13893 }
13894
13895 fn embryo_prefill_chunked(&mut self, ids: &[u32], start: usize) -> Option<Vec<f32>> {
13903 if ids.len() < 2
13904 || std::env::var("CMF_EMBRYO_CHUNK").as_deref() == Ok("0")
13905 || !self.embryo_resident_wanted()
13906 {
13907 return None;
13908 }
13909 let model = self.ensure_embryo_graph()?;
13910 let cmax = std::env::var("CMF_EMBRYO_CHUNK")
13911 .ok()
13912 .and_then(|v| v.parse::<usize>().ok())
13913 .filter(|&v| v >= 1)
13914 .unwrap_or(crate::gpu::EMBRYO_CHUNK_MAX)
13915 .min(crate::gpu::EMBRYO_CHUNK_MAX);
13916 let hs = self.hidden_size;
13917 let n = ids.len();
13918 let mut pos = start;
13919 let mut last = None;
13920 let mut rows = Vec::with_capacity(cmax * hs);
13921 while pos < n {
13922 let end = (pos + cmax).min(n);
13923 rows.clear();
13924 for &id in &ids[pos..end] {
13925 rows.extend_from_slice(&self.embed_single(id));
13926 }
13927 let mut lg = Vec::new();
13928 if !crate::gpu::forward_embryo_graph_chunk(
13929 &model,
13930 self.graph_kv_id,
13931 &rows,
13932 pos,
13933 end - pos,
13934 &mut lg,
13935 ) {
13936 if pos == start {
13937 if std::env::var("CMF_EMBRYO_DBG").as_deref() == Ok("1") {
13938 eprintln!("embryo-dbg: chunk prefill refused at position {pos}");
13939 }
13940 return None;
13941 }
13942 self.kv_cache.clear();
13946 self.clear_history();
13947 crate::gpu::graph_kv_reset(self.graph_kv_id);
13948 panic!(
13949 "resident Embryo chunk prefill refused at position {pos}; refusing an in-flight host fallback"
13950 );
13951 }
13952 last = Some(lg);
13953 pos = end;
13954 }
13955 if std::env::var("CMF_EMBRYO_DBG").as_deref() == Ok("1") {
13956 eprintln!(
13957 "embryo-dbg: chunk prefill {} ids from {start} in {} submits (chunk {cmax})",
13958 n - start,
13959 (n - start).div_ceil(cmax)
13960 );
13961 }
13962 last
13963 }
13964
13965 fn ensure_embryo_graph(&mut self) -> Option<std::sync::Arc<crate::gpu::EmbryoGraphModel>> {
13966 if self.embryo_graph.is_none() && self.embryo_resident_eligible() {
13967 const UMAX: u32 = u32::MAX;
13968 const HEADER: usize = crate::gpu::EMBRYO_META_HEADER;
13969 const REC: usize = 64;
13970 struct Pack {
13971 data: Vec<f32>,
13972 }
13973 impl Pack {
13974 fn put(&mut self, x: &[f32]) -> u32 {
13975 if x.is_empty() {
13976 return u32::MAX;
13977 }
13978 let off = self.data.len();
13979 self.data.extend_from_slice(x);
13980 off as u32
13981 }
13982 }
13983 let vmf = self.vmf_cfg;
13987 let gdn = self.gdn_cfg;
13988 let mut pack = Pack { data: Vec::new() };
13989 let mut meta = vec![0u32; HEADER];
13990 meta[0] = self.hidden_size as u32;
13991 meta[1] = self.intermediate_size as u32;
13992 meta[2] = self.vocab_size as u32;
13993 meta[3] = self.num_layers as u32;
13994 meta[4] = vmf.map(|c| c.num_heads).unwrap_or(0) as u32;
13995 meta[5] = vmf.map(|c| c.nphase).unwrap_or(0) as u32;
13996 meta[6] = vmf.map(|c| c.value_head_dim).unwrap_or(0) as u32;
13997 if let Some(g) = gdn {
13998 meta[24] = g.num_v_heads as u32;
13999 meta[25] = g.num_k_heads as u32;
14000 meta[26] = g.key_head_dim as u32;
14001 meta[27] = g.value_head_dim as u32;
14002 meta[28] = g.conv_kernel as u32;
14003 meta[29] = g.conv_dim() as u32;
14004 }
14005 meta[7] = self.num_heads as u32;
14006 meta[8] = self.num_kv_heads as u32;
14007 meta[9] = self.head_dim as u32;
14008 meta[10] = self.kv_cache.max_seq_len as u32;
14009 let clusters = self.head_clusters.as_ref().unwrap();
14010 let cluster_count = clusters.len() / self.hidden_size;
14011 if clusters.len() % self.hidden_size != 0
14012 || cluster_count == 0
14013 || cluster_count > 1024
14014 || self.vocab_size % cluster_count != 0
14015 || self.weights.lm_head.rows() < self.vocab_size
14016 || self.weights.final_norm.len() != self.hidden_size
14017 {
14018 return None;
14019 }
14020 meta[11] = cluster_count as u32;
14021 meta[12] = (self.vocab_size / cluster_count) as u32;
14022 meta[13] = vmf.map(|c| c.state_len()).unwrap_or(0) as u32;
14023 meta[16] = self.rotary_dim as u32;
14024 meta[17] = matches!(self.norm_style, NormStyle::Gemma) as u32;
14025 meta[19] = (self.rms_eps as f32).to_bits();
14026 let max_conv = self
14027 .weights
14028 .layers
14029 .iter()
14030 .filter_map(|lw| match &lw.attn {
14031 AttnKind::Linear(w) => w.conv.as_ref().map(|v| v.len() / self.hidden_size),
14032 _ => None,
14033 })
14034 .max()
14035 .unwrap_or(1);
14036 let phase_stride = vmf
14042 .map(|c| c.state_len() + max_conv.saturating_sub(1) * self.hidden_size)
14043 .unwrap_or(0);
14044 let gdn_stride = gdn.map(|g| g.state_len()).unwrap_or(0);
14045 let state_stride = phase_stride.max(gdn_stride);
14046 let bounded = self.anchor_core.clone();
14050 let (anchor_window, anchor_sink) = bounded
14051 .as_ref()
14052 .map(|ac| (ac.window, ac.sink))
14053 .unwrap_or((0, 0));
14054 let kv_stride = if bounded.is_some() {
14055 2usize
14056 .saturating_mul(self.num_kv_heads)
14057 .saturating_mul(anchor_window)
14058 .saturating_mul(self.head_dim)
14059 } else {
14060 2usize
14061 .saturating_mul(self.num_kv_heads)
14062 .saturating_mul(self.kv_cache.max_seq_len)
14063 .saturating_mul(self.head_dim)
14064 };
14065 meta[14] = state_stride as u32;
14066 meta[15] = kv_stride as u32;
14067 meta[18] = anchor_window as u32;
14068 meta[20] = anchor_sink as u32;
14069 meta[21] = match &self.bounded_rope {
14070 Some(rope) => {
14071 let off = pack.put(&rope.cos);
14073 let _ = pack.put(&rope.sin);
14074 off
14075 }
14076 None => UMAX,
14077 };
14078 let mut full_seen = false;
14079 let mut bounded_seen = 0usize;
14080 let mut phase_seen = 0usize;
14084 let mut gdn_seen = 0usize;
14085 for (li, lw) in self.weights.layers.iter().enumerate() {
14086 let base = meta.len();
14087 meta.resize(base + REC, UMAX);
14088 meta[base] = match &lw.attn {
14089 AttnKind::Linear(w) if w.phase_delta => 1,
14090 AttnKind::Linear(_) => 0,
14091 AttnKind::Full { .. } => 2,
14092 AttnKind::Bounded(_) => 3,
14093 AttnKind::LinearGdn(_) => 4,
14094 _ => UMAX,
14095 };
14096 meta[base + 1] = pack.put(&lw.input_norm);
14097 meta[base + 2] = pack.put(&lw.post_norm);
14098 meta[base + 25] = match &lw.attn {
14099 AttnKind::Linear(_) => {
14100 let off = ((phase_seen + gdn_seen) * state_stride) as u32;
14101 phase_seen += 1;
14102 off
14103 }
14104 AttnKind::LinearGdn(_) => {
14105 let off = ((phase_seen + gdn_seen) * state_stride) as u32;
14106 gdn_seen += 1;
14107 off
14108 }
14109 _ => UMAX,
14110 };
14111 match &lw.attn {
14112 AttnKind::LinearGdn(w) => {
14113 meta[base + 56] = pack.put(&Self::embryo_qtensor_f32(&w.in_proj_qkv));
14116 meta[base + 57] = pack.put(&Self::embryo_qtensor_f32(&w.in_proj_z));
14117 meta[base + 58] = pack.put(&Self::embryo_qtensor_f32(&w.in_proj_a));
14118 meta[base + 59] = pack.put(&Self::embryo_qtensor_f32(&w.in_proj_b));
14119 meta[base + 60] = pack.put(&w.conv1d);
14120 meta[base + 61] = pack.put(&w.a_log);
14121 meta[base + 62] = pack.put(&w.dt_bias);
14122 meta[base + 63] = pack.put(&w.norm);
14123 meta[base + 29] = pack.put(&Self::embryo_qtensor_f32(&w.out_proj));
14124 meta[base + 24] = 0;
14125 }
14126 AttnKind::Linear(w) => {
14127 meta[base + 3] = pack.put(&Self::embryo_qtensor_f32(&w.thq));
14128 meta[base + 4] = pack.put(&Self::embryo_qtensor_f32(&w.thk));
14129 meta[base + 5] = pack.put(&Self::embryo_qtensor_f32(&w.v_proj));
14130 meta[base + 6] = pack.put(&Self::embryo_qtensor_f32(&w.out_proj));
14131 let decay: Vec<f32> = w.decay.iter().map(|&x| x as f32).collect();
14132 meta[base + 7] = pack.put(&decay);
14133 if let Some((kg, kb)) = &w.k_gate {
14134 meta[base + 8] = pack.put(&Self::embryo_qtensor_f32(kg));
14135 meta[base + 9] = pack.put(kb);
14136 }
14137 if let Some(conv) = &w.conv {
14138 meta[base + 10] = pack.put(conv);
14139 meta[base + 24] = (conv.len() / self.hidden_size) as u32;
14140 } else {
14141 meta[base + 24] = 0;
14142 }
14143 }
14144 AttnKind::Full { wq, wk, wv, wo, .. } => {
14145 full_seen = true;
14146 meta[base + 11] = pack.put(&Self::embryo_qtensor_f32(wq));
14147 meta[base + 12] = pack.put(&Self::embryo_qtensor_f32(wk));
14148 meta[base + 13] = pack.put(&Self::embryo_qtensor_f32(wv));
14149 meta[base + 14] = pack.put(&Self::embryo_qtensor_f32(wo));
14150 meta[base + 26] = (li * kv_stride) as u32;
14151 }
14152 AttnKind::Bounded(w) => {
14153 meta[base + 11] = pack.put(&Self::embryo_qtensor_f32(&w.wq));
14154 meta[base + 12] = pack.put(&Self::embryo_qtensor_f32(&w.wk));
14155 meta[base + 13] = pack.put(&Self::embryo_qtensor_f32(&w.wv));
14156 meta[base + 14] = pack.put(&Self::embryo_qtensor_f32(&w.wo));
14157 meta[base + 26] = (bounded_seen * kv_stride) as u32;
14159 meta[base + 27] = pack.put(&w.sink_k);
14160 meta[base + 28] = pack.put(&w.sink_v);
14161 bounded_seen += 1;
14162 }
14163 _ => return None,
14164 }
14165 match &lw.ffn {
14166 FfnKind::Dense(d) => {
14167 meta[base + 15] = 0;
14168 meta[base + 16] = 0;
14169 meta[base + 21] = pack.put(&Self::embryo_qtensor_f32(&d.gate_proj));
14170 meta[base + 22] = pack.put(&Self::embryo_qtensor_f32(&d.up_proj));
14171 meta[base + 23] = pack.put(&Self::embryo_qtensor_f32(&d.down_proj));
14172 }
14173 FfnKind::Moe(m) => {
14174 let r = m.resonance.as_ref().unwrap();
14175 let (shared, _) = m.shared.as_ref().unwrap();
14176 meta[base + 15] = 1;
14177 meta[base + 16] = m.experts.len() as u32;
14178 meta[base + 17] = pack.put(&r.mu);
14179 meta[base + 18] = pack.put(&r.u);
14180 meta[base + 19] = pack.put(&r.bias);
14181 meta[base + 20] = r.k as u32;
14182 let mut shell = r.effective_shell(m.experts.len());
14191 shell.push(f32::NEG_INFINITY);
14192 meta[base + 30] = pack.put(&shell);
14193 meta[base + 21] = pack.put(&Self::embryo_qtensor_f32(&shared.gate_proj));
14194 meta[base + 22] = pack.put(&Self::embryo_qtensor_f32(&shared.up_proj));
14195 meta[base + 23] = pack.put(&Self::embryo_qtensor_f32(&shared.down_proj));
14196 for (e, ex) in m.experts.iter().enumerate() {
14197 meta[base + 32 + e * 3] =
14198 pack.put(&Self::embryo_qtensor_f32(&ex.gate_proj));
14199 meta[base + 33 + e * 3] =
14200 pack.put(&Self::embryo_qtensor_f32(&ex.up_proj));
14201 meta[base + 34 + e * 3] =
14202 pack.put(&Self::embryo_qtensor_f32(&ex.down_proj));
14203 }
14204 }
14205 FfnKind::DenseMoe(_) => return None,
14206 }
14207 }
14208 if !full_seen && self.num_layers == 0 {
14209 return None;
14210 }
14211 let id = {
14212 static NEXT: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(1);
14213 NEXT.fetch_add(1, std::sync::atomic::Ordering::Relaxed)
14214 };
14215 let model = crate::gpu::EmbryoGraphModel {
14216 id,
14217 hidden: self.hidden_size,
14218 intermediate: self.intermediate_size,
14219 vocab: self.vocab_size,
14220 layers: self.num_layers,
14221 phase_heads: vmf.map(|c| c.num_heads).unwrap_or(0),
14222 nphase: vmf.map(|c| c.nphase).unwrap_or(0),
14223 phase_dv: vmf.map(|c| c.value_head_dim).unwrap_or(0),
14224 anchor_q_heads: self.num_heads,
14225 anchor_kv_heads: self.num_kv_heads,
14226 anchor_head_dim: self.head_dim,
14227 rotary_dim: self.rotary_dim,
14228 max_seq: self.kv_cache.max_seq_len,
14229 cluster_count,
14230 cluster_size: self.vocab_size / cluster_count,
14231 phase_state_len: vmf.map(|c| c.state_len()).unwrap_or(0),
14232 state_stride,
14233 kv_stride,
14234 norm_gemma: matches!(self.norm_style, NormStyle::Gemma),
14235 phase_mass: vmf.map(|c| c.phase_mass).unwrap_or(0.0),
14236 weights: pack.data,
14237 meta,
14238 lm_head: Self::embryo_qtensor_f32(&self.weights.lm_head),
14239 clusters: clusters.as_ref().clone(),
14240 final_norm: self.weights.final_norm.clone(),
14241 inv_freq: self.inv_freq.as_ref().clone(),
14242 bounded: bounded.is_some(),
14243 kv_layers: if bounded.is_some() {
14244 bounded_seen
14245 } else {
14246 self.num_layers
14247 },
14248 state_layers: phase_seen + gdn_seen,
14249 anchor_window,
14250 anchor_sink,
14251 phase_layers: phase_seen,
14252 gdn_layers: gdn_seen,
14253 gdn_heads: gdn.map(|g| g.num_v_heads).unwrap_or(0),
14254 gdn_k_heads: gdn.map(|g| g.num_k_heads).unwrap_or(0),
14255 gdn_dk: gdn.map(|g| g.key_head_dim).unwrap_or(0),
14256 gdn_dv: gdn.map(|g| g.value_head_dim).unwrap_or(0),
14257 gdn_kk: gdn.map(|g| g.conv_kernel).unwrap_or(0),
14258 };
14259 self.embryo_graph = Some(std::sync::Arc::new(model));
14260 }
14261 self.embryo_graph.clone()
14262 }
14263
14264 fn forward_layers_span(
14265 &mut self,
14266 hidden: &[f32],
14267 position: usize,
14268 task_mask: Option<&TaskMask>,
14269 from: usize,
14270 upto: Option<usize>,
14271 ) -> Vec<f32> {
14272 debug_assert!(
14273 from == 0
14274 || (self.dsv4.is_none()
14275 && self.dsv41.is_none()
14276 && self.qwen4_exp.is_none()
14277 && self.g3n.is_none())
14278 );
14279 #[cfg(target_os = "macos")]
14285 if !crate::gpu_metal::wait_replay() {
14286 self.fail_metal_graph("the pending async replay failed before a plain forward");
14287 return vec![0.0; self.hidden_size];
14288 }
14289 if let Some(b) = &mut self.qwen4_exp {
14290 let _ = (task_mask, upto);
14291 let token_id = hidden.first().copied().unwrap_or(0.0) as u32;
14292 let mut logits = Vec::new();
14293 crate::qwen4_exp::forward_token(
14294 &b.0,
14295 &b.1,
14296 &b.2,
14297 &mut b.3,
14298 token_id,
14299 position,
14300 &self.inv_freq,
14301 self.pool.as_deref(),
14302 &mut logits,
14303 true,
14304 );
14305 self.graph_logits = Some(logits);
14306 return vec![0.0; self.hidden_size];
14307 }
14308 if let Some(b) = &mut self.dsv4 {
14314 let _ = (task_mask, upto);
14315 let token_id = hidden.first().copied().unwrap_or(0.0) as u32;
14316 let (g, layers, cfg, st) = (&b.0, &b.1, b.2, &mut b.3);
14317 st.pos = position;
14318 let mut logits = Vec::new();
14319 crate::dsv4::forward_token(
14320 g,
14321 layers,
14322 &cfg,
14323 st,
14324 token_id,
14325 &self.inv_freq,
14326 self.pool.as_deref(),
14327 &mut logits,
14328 );
14329 self.graph_logits = Some(logits);
14330 self.dspark_probe(position, token_id);
14331 return vec![0.0; self.hidden_size];
14334 }
14335 if let Some(b) = &mut self.dsv41 {
14337 let _ = (task_mask, upto);
14338 let token_id = hidden.first().copied().unwrap_or(0.0) as u32;
14339 let mut logits = Vec::new();
14340 crate::dsv41::forward_token(
14341 &b.0,
14342 &b.1,
14343 &b.2,
14344 &mut b.3,
14345 token_id,
14346 position,
14347 self.pool.as_deref(),
14348 &mut logits,
14349 );
14350 self.graph_logits = Some(logits);
14351 return vec![0.0; self.hidden_size];
14352 }
14353 if let Some(b) = &self.g3n {
14356 let _ = (task_mask, upto);
14357 return crate::g3n::g3n_forward(
14358 &b.0,
14359 &b.1,
14360 hidden,
14361 position,
14362 &mut self.kv_cache.layers,
14363 self.num_heads,
14364 self.num_kv_heads,
14365 self.head_dim,
14366 self.pool.as_deref(),
14367 );
14368 }
14369 if from == 0
14376 && upto.is_none()
14377 && task_mask.is_none()
14378 && self.anchor_core.is_some()
14379 && std::env::var("CMF_EMBRYO_DBG").as_deref() == Ok("1")
14380 {
14381 static ONCE: std::sync::Once = std::sync::Once::new();
14382 ONCE.call_once(|| {
14383 eprintln!(
14384 "embryo-dbg: graph_on decode={} prefill={} resident_env={:?} enabled_here={} \
14385 unsupported={} eligible={}",
14386 crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Decode),
14387 crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Prefill),
14388 std::env::var("CMF_EMBRYO_RESIDENT").ok(),
14389 crate::gpu::enabled_here(),
14390 self.graph_refused(),
14391 self.embryo_resident_eligible(),
14392 );
14393 });
14394 }
14395 if from == 0
14396 && upto.is_none()
14397 && task_mask.is_none()
14398 && crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Decode)
14402 && crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Prefill)
14403 && matches!(
14407 std::env::var("CMF_EMBRYO_RESIDENT").as_deref(),
14408 Ok("1") | Ok("parallel")
14409 )
14410 && crate::gpu::enabled_here()
14411 && !self.graph_refused()
14412 && (position == 0 || self.device_sequence_position().is_some())
14418 && self.embryo_resident_eligible()
14419 && let Some(model) = self.ensure_embryo_graph()
14420 {
14421 let mut lg = Vec::new();
14422 if crate::gpu::forward_embryo_graph(&model, self.graph_kv_id, hidden, position, &mut lg)
14423 {
14424 self.graph_logits = Some(lg);
14425 return vec![0.0; self.hidden_size];
14426 }
14427 if std::env::var("CMF_EMBRYO_DBG").as_deref() == Ok("1") {
14428 eprintln!("embryo-dbg: forward_embryo_graph refused at position {position}");
14429 }
14430 self.mark_graph_refused();
14436 if position != 0 {
14437 self.kv_cache.clear();
14441 self.clear_history();
14442 crate::gpu::graph_kv_reset(self.graph_kv_id);
14443 panic!(
14444 "resident Embryo graph refused at position {position}; refusing an in-flight host fallback"
14445 );
14446 }
14447 }
14448 let mut h = hidden.to_vec();
14449 self.mimo_moe_prepare();
14452 let _mimo_q8 = self.mimo_moe.is_on()
14453 .then(crate::qtensor::enter_full_gpu_q8_scope);
14454 let (nh, _nkv, _hd, hs, _rd, eps) = (
14457 self.num_heads,
14458 self.num_kv_heads,
14459 self.head_dim,
14460 self.hidden_size,
14461 self.rotary_dim,
14462 self.rms_eps,
14463 );
14464 let pool = self.pool.clone();
14465 let graph_env = std::env::var("CMF_GPU_WGPU_GRAPH").ok();
14477 let graph_on = match graph_env.as_deref() {
14478 Some("0") => false,
14479 Some("prefill") => false, Some(_) => true,
14481 None => crate::gpu::wgpu_graph_default(),
14487 };
14488 let graph_trusted =
14489 graph_env.is_some() || crate::gpu::wgpu_graph_default() || self.gdn_cfg.is_some();
14490 let race_eligible = graph_on
14491 && upto.is_none()
14492 && task_mask.is_none()
14493 && from == 0
14494 && !self.graph_refused();
14495 let mut tail_start = 0usize;
14496 if race_eligible && crate::gpu::graph_race_use_graph(graph_trusted) {
14497 let t_graph = std::time::Instant::now();
14498 let mut lg = Vec::new();
14499 let mut gl = 0usize;
14500 let built = self.try_token_graph_wgpu(hidden, position, &mut lg, &mut gl);
14501 let declined = built.is_none();
14502 let built = match built {
14503 Some(Ok(hh)) => Some(hh),
14504 Some(Err(())) => {
14505 self.clear_sequence_state();
14509 self.graph_failed
14510 .store(true, std::sync::atomic::Ordering::Relaxed);
14511 self.cancel
14512 .store(true, std::sync::atomic::Ordering::Relaxed);
14513 tracing::error!("token graph failed after admission; sequence state cleared");
14514 return vec![0.0; self.hidden_size];
14515 }
14516 None => None,
14517 };
14518 if declined && !self.o1_active() && self.attn_softcap == 0.0 {
14523 self.mark_graph_refused();
14524 }
14525 graph_note(built.is_some(), gl, self.num_layers);
14526 if let Some(hh) = built {
14527 let dur = t_graph.elapsed();
14528 if std::env::var("CMF_GRAPH_PROF").is_ok() {
14529 eprintln!("graph-call: {:.2} ms total", dur.as_secs_f64() * 1000.0);
14530 }
14531 if gl > 0 && gl < self.num_layers {
14532 h = hh;
14538 tail_start = gl;
14539 } else if graph_trusted || !crate::gpu::graph_race_first_token_hopeless(dur) {
14540 if !graph_trusted {
14541 crate::gpu::graph_race_record(true, dur);
14542 }
14543 if !lg.is_empty() {
14544 lg.resize(self.vocab_size, 0.0);
14547 if let Some(c) = self.final_softcap {
14548 for l in lg.iter_mut() {
14549 *l = c * (*l / c).tanh();
14550 }
14551 }
14552 self.graph_logits = Some(lg);
14553 }
14554 return hh;
14555 }
14556 }
14562 }
14563 let span = from > 0 || upto.is_some();
14587 if span && graph_on && task_mask.is_none() && graph_trusted {
14588 let upto_excl = upto.map_or(self.num_layers, |u| u + 1);
14589 let mut lg = Vec::new();
14590 let mut gl = 0usize;
14591 let span_res =
14592 self.try_token_graph_wgpu_span(hidden, position, &mut lg, from, upto_excl, &mut gl);
14593 let span_res = match span_res {
14594 Some(Ok(hh)) => Some(hh),
14595 Some(Err(())) => {
14596 self.clear_sequence_state();
14597 self.graph_failed
14598 .store(true, std::sync::atomic::Ordering::Relaxed);
14599 self.cancel
14600 .store(true, std::sync::atomic::Ordering::Relaxed);
14601 tracing::error!(
14602 "span token graph failed after admission; sequence state cleared"
14603 );
14604 return vec![0.0; self.hidden_size];
14605 }
14606 None => None,
14607 };
14608 graph_note(span_res.is_some(), gl, upto_excl - from);
14609 if std::env::var("CMF_GPU_DEBUG").is_ok() {
14610 static SEEN: std::sync::atomic::AtomicU32 = std::sync::atomic::AtomicU32::new(0);
14614 if SEEN.fetch_add(1, std::sync::atomic::Ordering::Relaxed) < 4 {
14615 eprintln!(
14616 "span graph: covered {gl} of {} layers [{from}..{upto_excl}) res={}",
14617 upto_excl - from,
14618 span_res.is_some()
14619 );
14620 }
14621 }
14622 if let Some(hh) = span_res {
14623 if gl == upto_excl - from {
14624 if !lg.is_empty() {
14625 lg.resize(self.vocab_size, 0.0);
14626 if let Some(c) = self.final_softcap {
14627 for l in lg.iter_mut() {
14628 *l = c * (*l / c).tanh();
14629 }
14630 }
14631 self.graph_logits = Some(lg);
14632 }
14633 crate::gpu::set_layer(-1);
14634 return hh;
14635 }
14636 h = hh;
14638 tail_start = from + gl;
14639 }
14640 }
14641 let t_race_cpu = (race_eligible && !graph_trusted).then(std::time::Instant::now);
14646
14647 let host_tail = tail_start > from;
14657 let _host_tail = (host_tail && !self.mimo_moe.is_on()).then(crate::gpu::enter_cpu_scope);
14658 let automatic_gpu_prefix = self.automatic_gpu_prefix();
14659
14660 let _prof_layers = crate::cpuprof::time(crate::cpuprof::Slot::Layers);
14661 #[cfg(target_os = "macos")]
14662 let mut gpu_skip_until = 0usize;
14663 for li in tail_start.max(from)..self.num_layers {
14664 let _capacity_tail = automatic_gpu_prefix
14665 .filter(|&prefix| li >= prefix && !self.mimo_moe.is_dynamic(li, host_tail))
14666 .map(|_| crate::gpu::enter_cpu_scope());
14667 crate::gpu::set_layer(li as i64); if let Some(u) = upto {
14669 if li > u {
14670 break;
14671 }
14672 }
14673 if let Some(mask) = task_mask {
14674 if !mask.layer_alive(li) {
14675 continue; }
14677 }
14678 #[cfg(target_os = "macos")]
14682 {
14683 if li < gpu_skip_until {
14684 continue;
14685 }
14686 if task_mask.is_none() {
14687 let end = self.q1_graph_gpu(li, upto, position, &mut h);
14688 if self
14689 .graph_failed
14690 .load(std::sync::atomic::Ordering::Relaxed)
14691 {
14692 return vec![0.0; self.hidden_size];
14696 }
14697 if end > li {
14698 gpu_skip_until = end;
14699 if self.is_loop_end(end - 1) && end < self.num_layers {
14702 h = inference::rms_norm(
14703 &h,
14704 &self.weights.final_norm,
14705 self.rms_eps,
14706 self.norm_style,
14707 );
14708 }
14709 continue;
14710 }
14711 }
14712 }
14713
14714 if task_mask.is_none() {
14715 match self.mimo_graph_layer_rows(li, &mut h, &[position]) {
14716 crate::gpu::BatchGraphOutcome::Completed => continue,
14717 crate::gpu::BatchGraphOutcome::Failed => return vec![0.0; self.hidden_size],
14718 crate::gpu::BatchGraphOutcome::Declined => {},
14719 }
14720 }
14721 #[cfg(feature = "gpu")]
14722 self.pull_lagging_host_kv(li, li + 1, position);
14723 let lw = &self.weights.layers[self.phys_layer(li)];
14724 if let Ok(tp) = std::env::var("CMF_TRACE_POS") {
14725 if tp.parse::<usize>().ok() == Some(position) {
14726 let n: f32 = h.iter().map(|x| x * x).sum::<f32>().sqrt();
14727 eprintln!(
14728 "TRACE pos {position} layer {li}: |h| = {n:.6} h0 {:.6} h1 {:.6}",
14729 h[0], h[1]
14730 );
14731 }
14732 }
14733 let prof = crate::cpuprof::time(crate::cpuprof::Slot::Norms);
14736 inference::rms_norm_into(
14737 &h,
14738 &lw.input_norm,
14739 self.rms_eps,
14740 self.norm_style,
14741 &mut self.ws.n1,
14742 );
14743 drop(prof);
14744
14745 let attn_out = match &lw.attn {
14746 AttnKind::Mla(w) => {
14747 let inv_freq_l = self.layer_inv_freq(li);
14748 let rs = self.layer_rope_scale(li);
14749 let eps = self.rms_eps;
14750 let pool = self.pool.clone();
14751 mla_attention(
14752 w,
14753 &self.ws.n1,
14754 &mut self.kv_cache.layers[li],
14755 position,
14756 &inv_freq_l,
14757 rs,
14758 eps,
14759 pool.as_deref(),
14760 )
14761 }
14762 AttnKind::Linear(w) => {
14763 let cfg = self.vmf_cfg.expect("linear layer without vmf_cfg");
14764 vmf_phase_forward(
14765 &self.ws.n1,
14766 w,
14767 &cfg,
14768 &mut self.kv_cache.layers[li].linear_state,
14769 self.pool.as_deref(),
14770 )
14771 }
14772 AttnKind::Kda(w) => {
14773 let cfg = self.kda_cfg.expect("kda layer without kda_cfg");
14774 crate::linear_core::kda_forward(
14775 &self.ws.n1,
14776 w,
14777 &cfg,
14778 &mut self.kv_cache.layers[li].linear_state,
14779 self.pool.as_deref(),
14780 )
14781 }
14782 AttnKind::LinearGdn(w) => {
14783 let cfg = self.gdn_cfg.expect("gdn layer without gdn_cfg");
14784 gdn_forward(
14785 &self.ws.n1,
14786 w,
14787 &cfg,
14788 &mut self.kv_cache.layers[li].linear_state,
14789 self.pool.as_deref(),
14790 )
14791 }
14792 AttnKind::ShortConv(w) => {
14793 let cfg = self
14794 .short_conv_cfg
14795 .expect("short-conv layer without short_conv_cfg");
14796 short_conv_forward(
14797 &self.ws.n1,
14798 w,
14799 &cfg,
14800 &mut self.kv_cache.layers[li].linear_state,
14801 self.pool.as_deref(),
14802 )
14803 }
14804 AttnKind::Bounded(w) => {
14805 let rope = self
14808 .bounded_rope
14809 .clone()
14810 .expect("bounded layer without an installed rotation table");
14811 let cfg = crate::bounded::BoundedAttnCfg {
14812 num_heads: self.num_heads,
14813 num_kv_heads: self.num_kv_heads,
14814 head_dim: self.head_dim,
14815 hidden_size: hs,
14816 scale: self.attn_scale,
14817 rope: &rope,
14818 pool: pool.as_deref(),
14819 };
14820 crate::bounded::bounded_attention(
14821 &self.ws.n1,
14822 w,
14823 &mut self.kv_cache.layers[li],
14824 &cfg,
14825 )
14826 }
14827 AttnKind::Full {
14828 wq,
14829 wk,
14830 wv,
14831 wo,
14832 q_norm,
14833 k_norm,
14834 output_gate,
14835 softplus_gate,
14836 bias,
14837 } if self.kv_cache.layers[li].o1_sealed() => {
14838 let inv_freq_l = self.layer_inv_freq(li);
14841 let (nkv_l, hd_l, rd_l) = self.layer_geom(li);
14842 let cfg = QwenAttnCfg {
14843 num_heads: self.layer_num_heads(li),
14844 num_kv_heads: nkv_l,
14845 head_dim: hd_l,
14846 hidden_size: hs,
14847 position,
14848 inv_freq: &inv_freq_l,
14849 rotary_dim: rd_l,
14850 scale: self.attn_scale,
14851 softcap: self.attn_softcap,
14852 window: None,
14853 v_norm: self.attn_v_norm,
14854 qk_norm_after_rope: self.qk_norm_after_rope,
14855 q_norm: q_norm.as_deref(),
14856 k_norm: k_norm.as_deref(),
14857 output_gate: *output_gate,
14858 softplus_gate: softplus_gate
14859 .as_ref()
14860 .map(|(gate, per_head)| (gate, *per_head)),
14861 rope_scale: self.layer_rope_scale(li),
14862 bias: bias
14863 .as_ref()
14864 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
14865 rms_eps: eps,
14866 norm_style: self.norm_style,
14867 pool: pool.as_deref(),
14868 v_head_dim: self.layer_v_dim(li),
14869 };
14870 attention::qwen_attention_nystrom(
14871 &self.ws.n1,
14872 wq,
14873 wk,
14874 wv,
14875 wo,
14876 &mut self.kv_cache.layers[li],
14877 &cfg,
14878 )
14879 }
14880 AttnKind::Full {
14881 wq,
14882 wk,
14883 wv,
14884 wo,
14885 q_norm,
14886 k_norm,
14887 output_gate,
14888 softplus_gate,
14889 bias,
14890 } => 'attn: {
14891 let dropin_reason =
14896 graph_on.then(|| self.graph_attn_decline_reason()).flatten();
14897 if let Some(reason) = dropin_reason {
14898 self.note_graph_decline("wgpu attn dropin", reason);
14899 }
14900 if graph_on
14901 && dropin_reason.is_none()
14902 && !*output_gate
14903 && softplus_gate.is_none()
14904 && self.attention_heads_per_layer.is_none()
14905 && bias.is_none()
14906 && task_mask.is_none()
14907 {
14908 let inv_freq_l = self.layer_inv_freq(li);
14909 let (nkv_l, hd_l, rd_l) = self.layer_geom(li);
14910 let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
14911 if let (Some((gm, qi)), Some((_, ki)), Some((_, vi)), Some((_, oi))) = (
14912 wq.mapped_q1(),
14913 wk.mapped_q1(),
14914 wv.mapped_q1(),
14915 wo.mapped_q1(),
14916 ) {
14917 let gm = gm.clone();
14918 let mut out = vec![0f32; hs];
14919 let cache = &self.kv_cache.layers[li];
14920 if crate::gpu::attn_dropin(
14921 &gm,
14922 self.graph_kv_id,
14923 li,
14924 &self.ws.n1,
14925 qi,
14926 ki,
14927 vi,
14928 oi,
14929 q_norm.as_deref(),
14930 k_norm.as_deref(),
14931 self.qk_norm_after_rope,
14932 &inv_freq_l,
14933 nh,
14934 nkv_l,
14935 hd_l,
14936 rd_l,
14937 hs,
14938 position,
14939 self.kv_cache.max_seq_len,
14940 gemma,
14941 eps as f32,
14942 cache.k_heads(),
14943 cache.v_heads(),
14944 &mut out,
14945 ) {
14946 break 'attn out;
14947 }
14948 }
14949 }
14950 let masked = task_mask
14951 .map(|m| m.head_flags(li, self.num_heads).iter().any(|&a| !a))
14952 .unwrap_or(false);
14953 let f32_view = (wq.as_f32(), wk.as_f32(), wv.as_f32(), wo.as_f32());
14954 let plain = self.layer_attn_plain(li);
14957 match (masked, f32_view) {
14958 (true, (Some(q), Some(k), Some(v), Some(o))) if plain => {
14961 let active_heads = task_mask.unwrap().head_flags(li, self.num_heads);
14962 attention::multi_head_attention(
14963 &self.ws.n1,
14964 q,
14965 k,
14966 v,
14967 o,
14968 &mut self.kv_cache.layers[li],
14969 self.num_heads,
14970 self.num_kv_heads,
14971 self.head_dim,
14972 self.hidden_size,
14973 position,
14974 &active_heads,
14975 &self.inv_freq,
14976 )
14977 }
14978 (masked, _) => {
14979 if masked {
14980 tracing::warn!(
14981 "layer {li}: head mask on quantized weights or on a \
14982 window/sink/per-layer-geometry layer not supported \
14983 yet — executing dense"
14984 );
14985 }
14986 let inv_freq_l = self.layer_inv_freq(li);
14987 let (nkv_l, hd_l, rd_l) = self.layer_geom(li);
14988 let cfg = QwenAttnCfg {
14989 num_heads: self.layer_num_heads(li),
14990 num_kv_heads: nkv_l,
14991 head_dim: hd_l,
14992 hidden_size: hs,
14993 position,
14994 inv_freq: &inv_freq_l,
14995 rotary_dim: rd_l,
14996 scale: self.attn_scale,
14997 softcap: self.attn_softcap,
14998 window: self.layer_window(li),
14999 v_norm: self.attn_v_norm,
15000 qk_norm_after_rope: self.qk_norm_after_rope,
15001 q_norm: q_norm.as_deref(),
15002 k_norm: k_norm.as_deref(),
15003 output_gate: *output_gate,
15004 softplus_gate: softplus_gate
15005 .as_ref()
15006 .map(|(gate, per_head)| (gate, *per_head)),
15007 rope_scale: self.layer_rope_scale(li),
15008 bias: bias
15009 .as_ref()
15010 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
15011 rms_eps: eps,
15012 norm_style: self.norm_style,
15013 pool: pool.as_deref(),
15014 v_head_dim: self.layer_v_dim(li),
15015 };
15016 attention::qwen_attention(
15017 &self.ws.n1,
15018 wq,
15019 wk,
15020 wv,
15021 wo,
15022 &mut self.kv_cache.layers[li],
15023 &cfg,
15024 )
15025 }
15026 }
15027 }
15028 };
15029 let attn_out = match &self.weights.layers[self.phys_layer(li)].attn_out_norm {
15032 Some(w) => inference::rms_norm(&attn_out, w, self.rms_eps, self.norm_style),
15033 None => attn_out,
15034 };
15035 let lw = &self.weights.layers[self.phys_layer(li)];
15036 let prof = crate::cpuprof::time(crate::cpuprof::Slot::Norms);
15037 inference::add_rmsnorm_fused_into(
15038 &mut h,
15039 &attn_out,
15040 &lw.post_norm,
15041 self.rms_eps,
15042 self.norm_style,
15043 &mut self.ws.p1,
15044 );
15045 drop(prof);
15046 let mut attn_out = attn_out;
15047 attention::recycle_buf(&mut attn_out);
15048 let post_normed = &self.ws.p1;
15049
15050 let ffn_masked = task_mask
15051 .map(|m| m.ffn_active_count(li) < self.intermediate_size)
15052 .unwrap_or(false);
15053 let ffn_out = match (ffn_masked, &lw.ffn) {
15065 (_, FfnKind::Dense(d)) if !d.segs.is_empty() => {
15069 let row = task_mask
15070 .and_then(|tm| tm.ffn_masks.get(li))
15071 .map(|v| v.as_slice());
15072 tube_ffn(d, post_normed, 1, self.pool.as_deref(), row)
15073 }
15074 (true, FfnKind::Dense(d)) => {
15075 let tm = task_mask.unwrap();
15076 let alive = tm.ffn_active_count(li);
15077 let deep = alive * 2 <= self.intermediate_size;
15078 if deep && d.down_proj.sparse_col_ok() && !d.gate_proj.has_prism_contract() {
15079 let active = tm.ffn_active_indices(li);
15080 sparse_ffn_quant(
15081 d,
15082 post_normed,
15083 &active,
15084 self.hidden_size,
15085 self.pool.as_deref(),
15086 )
15087 } else if deep
15088 && let (Some(g), Some(u), Some(dn)) = (
15089 d.gate_proj.as_f32(),
15090 d.up_proj.as_f32(),
15091 d.down_proj.as_f32(),
15092 )
15093 {
15094 let active = tm.ffn_active_indices(li);
15095 inference::sparse_ffn_forward(
15096 post_normed,
15097 g,
15098 u,
15099 dn,
15100 self.hidden_size,
15101 self.intermediate_size,
15102 &active,
15103 self.pool.as_deref(),
15104 )
15105 } else {
15106 let row = tm.ffn_masks.get(li).map(|v| v.as_slice());
15107 dense_ffn_batch(d, post_normed, 1, self.pool.as_deref(), row)
15108 }
15109 }
15110 (true, FfnKind::Moe(m)) => {
15111 let allowed = task_mask.and_then(|tm| tm.expert_flags(li, m.experts.len()));
15115 ffn_forward(
15116 &lw.ffn,
15117 post_normed,
15118 self.pool.as_deref(),
15119 allowed.as_deref(),
15120 )
15121 }
15122 (true, FfnKind::DenseMoe(dm)) => dense_moe_ffn(
15123 dm,
15124 post_normed,
15125 &h,
15126 self.rms_eps,
15127 self.norm_style,
15128 self.pool.as_deref(),
15129 ),
15130 (false, _) => match &lw.ffn {
15131 FfnKind::DenseMoe(dm) => dense_moe_ffn(
15132 dm,
15133 post_normed,
15134 &h,
15135 self.rms_eps,
15136 self.norm_style,
15137 self.pool.as_deref(),
15138 ),
15139 FfnKind::Moe(m)
15140 if task_mask.is_none() && self.mimo_moe.is_dynamic(li, host_tail) =>
15141 {
15142 moe_ffn_banked(&mut self.mimo_moe, li, m, post_normed, self.pool.as_deref())
15143 }
15144 _ => {
15145 let allowed = match (&lw.ffn, task_mask) {
15146 (FfnKind::Moe(m), Some(tm)) => tm.expert_flags(li, m.experts.len()),
15147 _ => None,
15148 };
15149 ffn_forward(
15150 &lw.ffn,
15151 post_normed,
15152 self.pool.as_deref(),
15153 allowed.as_deref(),
15154 )
15155 }
15156 },
15157 };
15158 let ffn_out = match &self.weights.layers[self.phys_layer(li)].ffn_out_norm {
15159 Some(w) => inference::rms_norm(&ffn_out, w, self.rms_eps, self.norm_style),
15160 None => ffn_out,
15161 };
15162 for (i, &f) in ffn_out.iter().enumerate() {
15163 h[i] += f;
15164 }
15165 let mut ffn_out = ffn_out;
15166 attention::recycle_buf(&mut ffn_out);
15167
15168 if let Some(sc) = self.weights.layers[self.phys_layer(li)].layer_scale {
15170 for v in h.iter_mut() {
15171 *v *= sc;
15172 }
15173 }
15174 if self.layer_dump.is_some() {
15176 self.dump_layer_row(position, li, &h);
15177 }
15178
15179 if self.is_loop_end(li) && li + 1 < self.num_layers {
15182 h = inference::rms_norm(
15183 &h,
15184 &self.weights.final_norm,
15185 self.rms_eps,
15186 self.norm_style,
15187 );
15188 }
15189
15190 if self.dyn_phi_layer == Some(li) {
15194 self.update_dyn_phi(&h);
15195 }
15196 }
15197 crate::gpu::set_layer(-1); if let Some(t) = t_race_cpu {
15199 crate::gpu::graph_race_record(false, t.elapsed());
15200 }
15201
15202 h
15203 }
15204
15205 fn update_dyn_phi(&mut self, h: &[f32]) {
15208 const A: f32 = 0.2;
15209 if self.dyn_phi_ema.len() != h.len() {
15210 self.dyn_phi_ema = vec![0.0; h.len()];
15211 self.dyn_phi_seen = 0;
15212 }
15213 if self.dyn_phi_seen == 0 {
15214 self.dyn_phi_ema.copy_from_slice(h);
15215 } else {
15216 for (e, &v) in self.dyn_phi_ema.iter_mut().zip(h) {
15217 *e = (1.0 - A) * *e + A * v;
15218 }
15219 }
15220 self.dyn_phi_seen += 1;
15221 }
15222
15223 pub fn dyn_phi(&self) -> &[f32] {
15225 &self.dyn_phi_ema
15226 }
15227
15228 pub fn set_dyn_phi_layer(&mut self, layer: Option<usize>) {
15230 self.dyn_phi_layer = layer;
15231 self.dyn_phi_ema.clear();
15232 self.dyn_phi_seen = 0;
15233 }
15234
15235 pub fn dynamic_skills(&self) -> Vec<(usize, String, usize)> {
15237 let Some(model) = &self.model else {
15238 return Vec::new();
15239 };
15240 model
15241 .header
15242 .skills
15243 .iter()
15244 .enumerate()
15245 .filter_map(|(i, sk)| {
15246 if sk.is_v2() {
15250 return None;
15251 }
15252 let ok = matches!(self.dyn_skill_layers.get(i), Some(Some(_)));
15253 let sel = sk.selection.as_ref()?;
15254 (ok).then(|| (i, sk.id.clone(), sel.phi_layer))
15255 })
15256 .collect()
15257 }
15258
15259 pub fn active_skill(&self) -> Option<usize> {
15261 self.dyn_active
15262 }
15263
15264 pub fn enable_dynamic_routing(&mut self) -> usize {
15269 use crate::swarm::{DynRouter, RoutableSkill};
15270 let Some(model) = self.model.clone() else {
15271 return 0;
15272 };
15273 if let Some(r) = &model.header.router {
15280 tracing::warn!(
15281 "dynamic routing disabled: this file declares router policy '{}' with \
15282 granularity \"{}\" — the request-level decision applies instead",
15283 r.policy,
15284 r.granularity
15285 );
15286 return 0;
15287 }
15288 if model.required_features & cortiq_core::format::features::SKILLS_V2 != 0
15294 || model.header.skills.iter().any(|s| s.is_v2())
15295 {
15296 tracing::warn!(
15297 "dynamic routing disabled: this file carries format-v2 skill records \
15298 (SKILLS_V2) — they route per request through a router policy only"
15299 );
15300 return 0;
15301 }
15302 if self.dyn_blend_loaded {
15305 tracing::warn!("dynamic routing unavailable on a blend-loaded pipeline");
15306 return 0;
15307 }
15308 if let Some(a) = self.dyn_active {
15312 if !matches!(self.dyn_skill_layers.get(a), Some(Some(_))) {
15313 tracing::warn!("loaded skill is not FFN-eligible — dynamic routing unavailable");
15314 return 0;
15315 }
15316 }
15317 let hidden = self.hidden_size;
15318 let mut skills = Vec::new();
15319 for (idx, id, _phi) in self.dynamic_skills() {
15320 if let Some(sel) = model.header.skills[idx].selection.as_ref() {
15321 if let Some(rs) = RoutableSkill::from_descriptor(idx, id, sel, hidden) {
15322 skills.push(rs);
15323 }
15324 }
15325 }
15326 if skills.is_empty() {
15327 return 0;
15328 }
15329 let phi = skills[0].phi_layer;
15331 if skills.iter().any(|s| s.phi_layer != phi) {
15332 tracing::warn!("routable skills disagree on phi_layer; using {phi}");
15333 }
15334 let n = skills.len();
15335 self.set_dyn_phi_layer(Some(phi));
15336 self.dyn_router = Some(DynRouter::new(skills));
15337 n
15338 }
15339
15340 pub fn route_switches(&self) -> Vec<(usize, Option<String>, Option<String>)> {
15342 self.dyn_router
15343 .as_ref()
15344 .map(|r| r.switches.clone())
15345 .unwrap_or_default()
15346 }
15347
15348 fn lm_head_forward(&self, hidden: &[f32]) -> Vec<f32> {
15351 let _mimo_q8 = self.mimo_moe.is_on()
15352 .then(crate::qtensor::enter_full_gpu_q8_scope);
15353 let rows = self.weights.lm_head.rows();
15354 let mut logits = attention::take_buf(rows.min(self.vocab_size));
15355 let served = self.mimo_moe.is_on() && crate::gpu::mimo_q8_short_enabled()
15359 && rows == self.vocab_size && !self.weights.lm_head.has_prism_contract()
15360 && self.weights.lm_head.graph_weight().is_some_and(|(model, idx, kind, _)| {
15361 kind == 7 && crate::gpu::q82_short_rows(model, idx, hidden, 1,
15362 rows, self.hidden_size, &mut logits)
15363 });
15364 if !served {
15365 self.weights.lm_head.matvec(hidden, &mut logits, self.pool.as_deref());
15366 }
15367 logits.resize(self.vocab_size, 0.0);
15368 if let Some(m) = self.logit_multiplier {
15369 for l in logits.iter_mut() {
15370 *l *= m;
15371 }
15372 }
15373 if let Some(c) = self.final_softcap {
15374 for l in logits.iter_mut() {
15375 *l = c * (*l / c).tanh();
15376 }
15377 }
15378 if let Some(cm) = self.head_clusters.as_ref() {
15379 self.hierarchical_head_logprobs(hidden, cm, &mut logits);
15380 }
15381 logits
15382 }
15383
15384 fn hierarchical_head_logprobs(&self, hidden: &[f32], cm: &[f32], logits: &mut [f32]) {
15387 let h = hidden.len();
15388 let ncl = cm.len() / h.max(1);
15389 if ncl == 0 || logits.len() % ncl != 0 {
15390 return;
15391 }
15392 let cs = logits.len() / ncl;
15393 let mut lc = vec![0.0f32; ncl];
15395 for c in 0..ncl {
15396 let row = &cm[c * h..(c + 1) * h];
15397 let mut s = 0.0f32;
15398 for j in 0..h {
15399 s += row[j] * hidden[j];
15400 }
15401 lc[c] = s;
15402 }
15403 let mx = lc.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
15404 let lse: f32 = mx + lc.iter().map(|v| (v - mx).exp()).sum::<f32>().ln();
15405 for c in 0..ncl {
15406 let blk = &mut logits[c * cs..(c + 1) * cs];
15407 let bm = blk.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
15408 let bl: f32 = bm + blk.iter().map(|v| (v - bm).exp()).sum::<f32>().ln();
15409 let add = lc[c] - lse - bl;
15410 for v in blk.iter_mut() {
15411 *v += add;
15412 }
15413 }
15414 }
15415
15416 pub fn prefill_next_logits(&mut self, ids: &[u32], task_mask: Option<&TaskMask>) -> Vec<f32> {
15421 #[cfg(target_os = "macos")]
15422 crate::gpu_metal::set_io_namespace(self.graph_kv_id);
15423 self.clear_sequence_state();
15424 crate::gpu::graph_race_begin_generation();
15428 if task_mask.is_none() {
15429 self.o1_begin();
15430 }
15431 let mut hidden = vec![0.0f32; self.hidden_size];
15432 for (pos, &id) in ids.iter().enumerate() {
15433 let emb = self.embed_single(id);
15434 hidden = self.forward_layers(&emb, pos, task_mask);
15435 }
15436 if let Err(err) = self.o1_seal_checked() {
15437 self.o1_fail(err);
15438 }
15439 if let Some(logits) = self.graph_logits.take() {
15442 return logits;
15443 }
15444 inference::rms_norm_into(
15445 &hidden,
15446 &self.weights.final_norm,
15447 self.rms_eps,
15448 self.norm_style,
15449 &mut self.ws.n1,
15450 );
15451 self.lm_head_forward(&self.ws.n1)
15452 }
15453}
15454
15455pub fn create_test_pipeline(
15457 hidden_size: usize,
15458 intermediate_size: usize,
15459 num_heads: usize,
15460 num_kv_heads: usize,
15461 head_dim: usize,
15462 num_layers: usize,
15463 vocab_size: usize,
15464) -> Pipeline {
15465 let synth = |n: usize, salt: usize| -> Vec<f32> {
15468 (0..n)
15469 .map(|i| (((i * 31 + salt * 17 + 7) % 97) as f32 / 97.0 - 0.5) * 0.2)
15470 .collect()
15471 };
15472 let qt = |rows: usize, cols: usize, salt: usize| -> QTensor {
15473 QTensor::from_f32(synth(rows * cols, salt), rows, cols)
15474 };
15475 let layer_weights: Vec<LayerWeights> = (0..num_layers)
15476 .map(|li| LayerWeights {
15477 input_norm: vec![1.0; hidden_size],
15478 post_norm: vec![1.0; hidden_size],
15479 attn_out_norm: None,
15480 ffn_out_norm: None,
15481 layer_scale: None,
15482 ffn: FfnKind::Dense(DenseFfn {
15483 gate_proj: qt(intermediate_size, hidden_size, li * 10 + 5),
15484 up_proj: qt(intermediate_size, hidden_size, li * 10 + 6),
15485 down_proj: qt(hidden_size, intermediate_size, li * 10 + 7),
15486 act: Act::Silu,
15487 down_t: None,
15488 segs: Vec::new(),
15489 }),
15490 attn: AttnKind::Full {
15491 bias: None,
15492 wq: qt(num_heads * head_dim, hidden_size, li * 10 + 1),
15493 wk: qt(num_kv_heads * head_dim, hidden_size, li * 10 + 2),
15494 wv: qt(num_kv_heads * head_dim, hidden_size, li * 10 + 3),
15495 wo: qt(hidden_size, num_heads * head_dim, li * 10 + 4),
15496 q_norm: None,
15497 k_norm: None,
15498 output_gate: false,
15499 softplus_gate: None,
15500 },
15501 })
15502 .collect();
15503
15504 Pipeline::new(
15505 Tokenizer::byte_level(),
15506 PipelineWeights {
15507 embed_tokens: qt(vocab_size, hidden_size, 100),
15508 layers: layer_weights,
15509 lm_head: qt(vocab_size, hidden_size, 200),
15510 final_norm: vec![1.0; hidden_size],
15511 },
15512 hidden_size,
15513 intermediate_size,
15514 num_heads,
15515 num_kv_heads,
15516 head_dim,
15517 num_layers,
15518 num_layers, false, vocab_size,
15521 1e-6,
15522 10_000.0,
15523 NormStyle::Qwen,
15524 4096,
15525 SamplerConfig {
15526 seed: Some(42),
15527 ..Default::default()
15528 },
15529 )
15530}
15531
15532#[inline]
15537fn mask_bit(row: &[u8], j: usize) -> bool {
15538 (row.get(j >> 3).copied().unwrap_or(0) >> (j & 7)) & 1 != 0
15539}
15540
15541fn mask_gain() -> f32 {
15552 static G: std::sync::OnceLock<f32> = std::sync::OnceLock::new();
15553 *G.get_or_init(|| {
15554 std::env::var("CMF_FFN_MASK_GAIN")
15555 .ok()
15556 .and_then(|v| v.parse().ok())
15557 .unwrap_or(1.0)
15558 })
15559}
15560
15561fn zero_masked_cols(g: &mut [f32], rows: usize, inter: usize, row: &[u8]) {
15562 let fill = meanfill().and_then(|(i, v)| {
15565 let li = crate::gpu::cur_layer();
15566 (*i == inter && li >= 0).then(|| &v[li as usize * inter..(li as usize + 1) * inter])
15567 });
15568 for r in 0..rows {
15569 let base = r * inter;
15570 for (bi, &byte) in row.iter().enumerate() {
15571 if byte == 0xFF {
15572 continue;
15573 }
15574 let j0 = bi * 8;
15575 for bit in 0..8 {
15576 let j = j0 + bit;
15577 if j < inter && byte & (1 << bit) == 0 {
15578 g[base + j] = fill.map_or(0.0, |f| f[j]);
15579 }
15580 }
15581 }
15582 }
15583 let gain = mask_gain();
15584 if gain != 1.0 {
15585 for v in g[..rows * inter].iter_mut() {
15586 *v *= gain;
15587 }
15588 }
15589}
15590
15591#[inline]
15593fn tube_bit(row: Option<&[u8]>, i: usize) -> bool {
15594 row.is_none_or(|r| mask_bit(r, i))
15595}
15596
15597fn all_bits_on(row: &[u8], n: usize) -> bool {
15600 (0..n).all(|i| mask_bit(row, i))
15601}
15602
15603fn tube_topk() -> usize {
15611 static K: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
15612 *K.get_or_init(|| {
15613 std::env::var("CMF_TUBE_TOPK")
15614 .ok()
15615 .and_then(|v| v.parse().ok())
15616 .unwrap_or(0)
15617 })
15618}
15619
15620fn tube_score_oracle() -> bool {
15621 static O: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
15622 *O.get_or_init(|| std::env::var("CMF_TUBE_SCORE").is_ok_and(|v| v == "oracle"))
15623}
15624
15625fn tube_ffn_routed(
15632 d: &DenseFfn,
15633 xs: &[f32],
15634 b: usize,
15635 pool: Option<&Pool>,
15636 mask_row: Option<&[u8]>,
15637 k: usize,
15638) -> Vec<f32> {
15639 let hidden = d.down_proj.rows();
15640 let core = d.gate_proj.rows();
15641 let core_full = mask_row.is_none_or(|r| all_bits_on(r, core));
15642 let mut out = match (b, core_full, mask_row) {
15643 (1, true, _) => dense_ffn(d, xs, pool),
15644 (1, false, Some(row)) => dense_ffn_masked(d, xs, pool, row),
15645 (_, true, _) => dense_ffn_batch(d, xs, b, pool, None),
15646 (_, false, row) => dense_ffn_batch(d, xs, b, pool, row),
15647 };
15648 let cand: Vec<usize> = (0..d.segs.len())
15649 .filter(|&i| tube_bit(mask_row, d.segs[i].start))
15650 .collect();
15651 if cand.is_empty() {
15652 return out;
15653 }
15654 let oracle = tube_score_oracle();
15658 let mut acts: Vec<Vec<f32>> = Vec::with_capacity(cand.len());
15659 let mut scores = vec![0f32; b * cand.len()];
15660 for (ci, &i) in cand.iter().enumerate() {
15661 let seg = &d.segs[i];
15662 let w = seg.width;
15663 let mut g = vec![0.0f32; b * w];
15664 if b == 1 {
15665 seg.gate.matvec(xs, &mut g, pool);
15666 } else {
15667 seg.gate.matmat(xs, b, &mut g, pool);
15668 }
15669 for v in g.iter_mut() {
15670 *v = Act::Silu.combine(*v, 1.0);
15671 }
15672 if !oracle {
15673 for t in 0..b {
15674 scores[t * cand.len() + ci] =
15675 g[t * w..(t + 1) * w].iter().map(|v| v * v).sum::<f32>();
15676 }
15677 }
15678 if oracle || b > 1 {
15679 let mut u = vec![0.0f32; b * w];
15680 if b == 1 {
15681 seg.up.matvec(xs, &mut u, pool);
15682 } else {
15683 seg.up.matmat(xs, b, &mut u, pool);
15684 }
15685 for (a, &v) in g.iter_mut().zip(u.iter()) {
15686 *a *= v;
15687 }
15688 if oracle {
15689 for t in 0..b {
15690 scores[t * cand.len() + ci] =
15691 g[t * w..(t + 1) * w].iter().map(|v| v * v).sum::<f32>();
15692 }
15693 }
15694 }
15695 acts.push(g);
15696 }
15697 let keep = k.min(cand.len());
15699 let mut scratch: Vec<f32> = Vec::new();
15700 for t in 0..b {
15701 let mut sc: Vec<(f32, usize)> = (0..cand.len())
15702 .map(|ci| (scores[t * cand.len() + ci], ci))
15703 .collect();
15704 sc.sort_unstable_by(|x, y| y.0.total_cmp(&x.0));
15705 let mut alive = vec![false; cand.len()];
15706 for &(_, ci) in sc.iter().take(keep) {
15707 alive[ci] = true;
15708 }
15709 if b > 1 {
15710 for (ci, a) in acts.iter_mut().enumerate() {
15711 if !alive[ci] {
15712 let w = d.segs[cand[ci]].width;
15713 a[t * w..(t + 1) * w].fill(0.0);
15714 }
15715 }
15716 } else {
15717 for (ci, &i) in cand.iter().enumerate() {
15721 if !alive[ci] {
15722 continue;
15723 }
15724 let seg = &d.segs[i];
15725 let w = seg.width;
15726 let g = &mut acts[ci];
15727 if !tube_score_oracle() {
15728 scratch.clear();
15729 scratch.resize(w, 0.0);
15730 seg.up.matvec(xs, &mut scratch, pool);
15731 for (a, &v) in g.iter_mut().zip(scratch.iter()) {
15732 *a *= v;
15733 }
15734 }
15735 let mut acc = vec![0.0f32; hidden];
15736 seg.down.matvec(g, &mut acc, pool);
15737 for (o, a) in out.iter_mut().zip(&acc) {
15738 *o += *a;
15739 }
15740 }
15741 }
15742 }
15743 if b > 1 {
15744 for (ci, &i) in cand.iter().enumerate() {
15745 let seg = &d.segs[i];
15746 let mut acc = vec![0.0f32; b * hidden];
15747 seg.down.matmat(&acts[ci], b, &mut acc, pool);
15748 for (o, a) in out.iter_mut().zip(&acc) {
15749 *o += *a;
15750 }
15751 }
15752 }
15753 out
15754}
15755
15756fn tube_ffn(
15762 d: &DenseFfn,
15763 xs: &[f32],
15764 b: usize,
15765 pool: Option<&Pool>,
15766 mask_row: Option<&[u8]>,
15767) -> Vec<f32> {
15768 if tube_topk() > 0 {
15769 return tube_ffn_routed(d, xs, b, pool, mask_row, tube_topk());
15770 }
15771 let hidden = d.down_proj.rows();
15772 let core = d.gate_proj.rows();
15773 let core_full = mask_row.is_none_or(|r| all_bits_on(r, core));
15774 let mut out = match (b, core_full, mask_row) {
15775 (1, true, _) => dense_ffn(d, xs, pool),
15776 (1, false, Some(row)) => dense_ffn_masked(d, xs, pool, row),
15777 (_, true, _) => dense_ffn_batch(d, xs, b, pool, None),
15778 (_, false, row) => dense_ffn_batch(d, xs, b, pool, row),
15779 };
15780 TUBE_SCRATCH.with(|sc| {
15781 let mut sc = sc.borrow_mut();
15782 let [g, u, acc] = &mut *sc;
15783 for seg in &d.segs {
15784 if !tube_bit(mask_row, seg.start) {
15785 continue;
15786 }
15787 let w = seg.width;
15788 g.resize(b * w, 0.0);
15789 if b == 1
15790 && d.act == Act::Silu
15791 && QTensor::matvec_silu_mul(&seg.gate, &seg.up, xs, g, pool)
15792 {
15793 } else {
15795 u.resize(b * w, 0.0);
15796 if b == 1 {
15797 QTensor::matvec_many([&seg.gate, &seg.up], xs, [g, u], pool);
15798 } else {
15799 seg.gate.matmat(xs, b, g, pool);
15800 seg.up.matmat(xs, b, u, pool);
15801 }
15802 for i in 0..b * w {
15803 g[i] = d.act.combine(g[i], u[i]);
15804 }
15805 }
15806 acc.resize(b * hidden, 0.0);
15807 acc.fill(0.0);
15808 if b == 1 {
15809 seg.down.matvec(g, acc, pool);
15810 } else {
15811 seg.down.matmat(g, b, acc, pool);
15812 }
15813 for (o, a) in out.iter_mut().zip(acc.iter()) {
15814 *o += *a;
15815 }
15816 }
15817 out
15818 })
15819}
15820
15821thread_local! {
15822 static TUBE_SCRATCH: std::cell::RefCell<[Vec<f32>; 3]> =
15826 const { std::cell::RefCell::new([Vec::new(), Vec::new(), Vec::new()]) };
15827}
15828
15829fn dense_ffn_batch(
15830 d: &DenseFfn,
15831 xs: &[f32],
15832 b: usize,
15833 pool: Option<&Pool>,
15834 mask_row: Option<&[u8]>,
15835) -> Vec<f32> {
15836 let inter = d.gate_proj.rows();
15837 let hidden = d.down_proj.rows();
15838 if mask_row.is_none()
15846 && d.act == Act::Silu
15847 && b >= 32
15848 && crate::gpu::enabled_here()
15849 && !crate::gpu::mm_killed()
15850 && refit_dir().is_none()
15855 && !ffn_probe_active()
15860 {
15861 if let (Some((model, w1)), Some((_, w3)), Some((_, w2))) = (
15862 d.gate_proj.mapped_q4t(),
15863 d.up_proj.mapped_q4t(),
15864 d.down_proj.mapped_q4t(),
15865 ) {
15866 let mut out = vec![0.0f32; b * hidden];
15867 if crate::gpu::q4t_ffn(model, w1, w3, w2, xs, b, hidden, inter, &mut out) {
15868 return out;
15869 }
15870 }
15871 if let (Some((model, w1)), Some((_, w3)), Some((_, w2))) = (
15876 d.gate_proj.mapped_q4tp(),
15877 d.up_proj.mapped_q4tp(),
15878 d.down_proj.mapped_q4tp(),
15879 ) {
15880 let mut out = vec![0.0f32; b * hidden];
15881 if crate::gpu::q4tp_ffn(model, w1, w3, w2, xs, b, hidden, inter, &mut out) {
15882 return out;
15883 }
15884 }
15885 }
15886 let mut g = vec![0.0f32; b * inter];
15887 d.gate_proj.matmat(xs, b, &mut g, pool);
15888 let mut u = vec![0.0f32; b * inter];
15889 d.up_proj.matmat(xs, b, &mut u, pool);
15890 if gate_topk() > 0 && d.act == Act::Silu {
15891 for t in 0..b {
15892 let row = &mut g[t * inter..(t + 1) * inter];
15893 for v in row.iter_mut() {
15894 *v = Act::Silu.combine(*v, 1.0);
15895 }
15896 keep_top_k(row, gate_topk());
15897 }
15898 for i in 0..b * inter {
15899 g[i] *= u[i];
15900 }
15901 } else {
15902 for i in 0..b * inter {
15903 g[i] = d.act.combine(g[i], u[i]);
15904 }
15905 }
15906 if let Some(row) = mask_row {
15907 zero_masked_cols(&mut g, b, inter, row);
15908 }
15909 if oracle_topk() > 0 {
15910 for t in 0..b {
15911 keep_top_k(&mut g[t * inter..(t + 1) * inter], oracle_topk());
15912 }
15913 }
15914 let mut out = vec![0.0f32; b * hidden];
15915 d.down_proj.matmat(&g, b, &mut out, pool);
15916 if refit_dir().is_some() {
15917 let li = crate::gpu::cur_layer();
15918 if li >= 0 {
15919 refit_accumulate(li as usize, &g, b, inter, &out, hidden, pool);
15920 }
15921 }
15922 FFN_PROBE.with(|pr| {
15926 if let Some(acc) = pr.borrow_mut().as_mut() {
15927 let li = crate::gpu::cur_layer();
15928 if li < 0 {
15929 return;
15930 }
15931 let Some(row) = acc.get_mut(li as usize) else {
15932 return;
15933 };
15934 let sq = probe_sq();
15935 for t in 0..b {
15936 for (a, &v) in row.iter_mut().zip(&g[t * inter..(t + 1) * inter]) {
15937 *a += if sq {
15938 (v as f64) * (v as f64)
15939 } else {
15940 (v as f64).abs()
15941 };
15942 }
15943 }
15944 }
15945 });
15946 out
15947}
15948
15949fn accumulate_act(m: &MoeFfn, xs: &[f32], b: usize) {
15954 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
15955 static DUMP: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
15956 let on = *ON.get_or_init(|| std::env::var("CMF_RMS_TRACE").is_ok());
15957 let dump = *DUMP.get_or_init(|| std::env::var("CMF_ACT_DUMP").is_ok());
15958 if (!on && !dump) || b == 0 {
15959 return;
15960 }
15961 let hidden = xs.len() / b;
15962 if on {
15963 let mut acc = m.act_sq.borrow_mut();
15964 if acc.len() < hidden {
15965 acc.resize(hidden, 0.0);
15966 }
15967 for t in 0..b {
15968 let row = &xs[t * hidden..(t + 1) * hidden];
15969 for (a, &v) in acc.iter_mut().zip(row) {
15970 *a += (v as f64) * (v as f64);
15971 }
15972 }
15973 }
15974 if dump {
15975 let cap: usize = std::env::var("CMF_ACT_DUMP_ROWS")
15978 .ok()
15979 .and_then(|v| v.parse().ok())
15980 .unwrap_or(4096);
15981 let mut rows = m.act_rows.borrow_mut();
15982 if rows.len() < cap * hidden {
15983 let take = b.min((cap * hidden - rows.len()) / hidden.max(1));
15984 rows.extend_from_slice(&xs[..take * hidden]);
15985 }
15986 }
15987}
15988
15989#[derive(Clone, Copy)]
15992struct SendVecs(*mut Vec<f32>);
15993unsafe impl Send for SendVecs {}
15994unsafe impl Sync for SendVecs {}
15995impl SendVecs {
15996 #[inline]
15997 fn at(self, i: usize) -> *mut Vec<f32> {
15998 unsafe { self.0.add(i) }
15999 }
16000}
16001
16002fn moe_ffn_batch(
16003 m: &MoeFfn,
16004 xs: &[f32],
16005 b: usize,
16006 hidden: usize,
16007 pool: Option<&Pool>,
16008 allowed: Option<&[bool]>,
16009) -> Vec<f32> {
16010 accumulate_act(m, xs, b);
16011 let ne = m.experts.len();
16012 let mut logits = vec![0.0f32; b * ne];
16013 match &m.resonance {
16014 Some(r) => {
16015 let hdim = xs.len() / b.max(1);
16016 for bi in 0..b {
16017 r.scores(
16018 &xs[bi * hdim..(bi + 1) * hdim],
16019 &mut logits[bi * ne..(bi + 1) * ne],
16020 );
16021 }
16022 }
16023 None => m.router.matmat(xs, b, &mut logits, pool),
16024 }
16025
16026 let mut assign: Vec<Vec<(usize, f32)>> = vec![Vec::new(); ne];
16029 {
16030 let mut st = m.stats.borrow_mut();
16031 if st.len() < ne {
16032 st.resize(ne, 0);
16033 }
16034 for bi in 0..b {
16035 let (idx, p, wsum) = moe_route(&logits[bi * ne..(bi + 1) * ne], m, allowed);
16036 for &e in &idx {
16037 st[e] += 1;
16038 assign[e].push((bi, p[e] / wsum));
16039 }
16040 }
16041 }
16042
16043 let mut out = vec![0.0f32; b * hidden];
16044 let cols = m.experts[0].gate_proj.cols();
16045 let run_expert = |d: &DenseFfn, list: &[(usize, f32)], out: &mut [f32]| {
16046 let sb = list.len();
16047 let mut sub = vec![0.0f32; sb * cols];
16048 for (k, &(bi, _)) in list.iter().enumerate() {
16049 sub[k * cols..(k + 1) * cols].copy_from_slice(&xs[bi * cols..(bi + 1) * cols]);
16050 }
16051 let eo = dense_ffn_batch(d, &sub, sb, pool, None);
16052 for (k, &(bi, w)) in list.iter().enumerate() {
16053 for i in 0..hidden {
16054 out[bi * hidden + i] += w * eo[k * hidden + i];
16055 }
16056 }
16057 };
16058 let active: Vec<usize> = (0..ne).filter(|&e| !assign[e].is_empty()).collect();
16064 if pool.is_some() && active.len() >= 8 {
16065 let mut panels: Vec<Vec<f32>> = vec![Vec::new(); active.len()];
16066 {
16067 let panel_ptr = SendVecs(panels.as_mut_ptr());
16068 let experts = &m.experts;
16071 let (active_r, assign_r) = (&active, &assign);
16072 let inherit_cpu = crate::gpu::inherit_cpu_scope();
16073 let run = |start: usize, end: usize| {
16074 let _cpu_scope = inherit_cpu();
16075 for ai in start..end {
16076 let e = active_r[ai];
16077 let list = &assign_r[e];
16078 let sb = list.len();
16079 let mut sub = vec![0.0f32; sb * cols];
16080 for (k, &(bi, _)) in list.iter().enumerate() {
16081 sub[k * cols..(k + 1) * cols]
16082 .copy_from_slice(&xs[bi * cols..(bi + 1) * cols]);
16083 }
16084 unsafe {
16086 *panel_ptr.at(ai) = dense_ffn_batch(&experts[e], &sub, sb, None, None);
16087 }
16088 }
16089 };
16090 match pool {
16091 Some(p) => p.run_rows(active.len(), &run),
16092 None => run(0, active.len()),
16093 }
16094 }
16095 for (ai, &e) in active.iter().enumerate() {
16096 for (k, &(bi, w)) in assign[e].iter().enumerate() {
16097 let eo = &panels[ai][k * hidden..(k + 1) * hidden];
16098 for i in 0..hidden {
16099 out[bi * hidden + i] += w * eo[i];
16100 }
16101 }
16102 }
16103 } else {
16104 for &e in &active {
16105 run_expert(&m.experts[e], &assign[e], &mut out);
16106 }
16107 }
16108 if let Some((se, gate)) = &m.shared {
16109 let all: Vec<(usize, f32)> = if let Some(gate) = gate {
16110 let mut gl = vec![0.0f32; b];
16111 gate.matmat(xs, b, &mut gl, pool);
16112 (0..b)
16113 .map(|bi| (bi, 1.0 / (1.0 + (-gl[bi]).exp())))
16114 .collect()
16115 } else {
16116 (0..b).map(|bi| (bi, 1.0)).collect()
16117 };
16118 run_expert(se, &all, &mut out);
16119 }
16120 out
16121}
16122
16123fn moe_ffn_rows_exact(
16135 m: &MoeFfn,
16136 xs: &[f32],
16137 b: usize,
16138 hidden: usize,
16139 pool: Option<&Pool>,
16140) -> Vec<f32> {
16141 let mut out = vec![0.0f32; b * hidden];
16142 let per_row = |out: &mut [f32]| {
16143 for r in 0..b {
16144 let o = moe_ffn(m, &xs[r * hidden..(r + 1) * hidden], pool, None);
16145 out[r * hidden..(r + 1) * hidden].copy_from_slice(&o);
16146 }
16147 };
16148 let covered = !crate::gpu::enabled_here()
16149 && moe_batch_enabled()
16150 && m.shared.is_none()
16151 && m.resonance.is_none()
16152 && FFN_PROBE.with(|pr| pr.borrow().is_none())
16153 && m.experts.iter().all(|d| d.act == Act::Silu);
16154 if !covered {
16155 per_row(&mut out);
16156 return out;
16157 }
16158 let ne = m.experts.len();
16159 let mut routes: Vec<(Vec<usize>, Vec<f32>)> = Vec::with_capacity(b);
16161 for r in 0..b {
16162 let x = &xs[r * hidden..(r + 1) * hidden];
16163 accumulate_act(m, x, 1);
16164 let mut logits = vec![0.0f32; ne];
16165 m.router.matvec(x, &mut logits, pool);
16166 let (idx, p, wsum) = moe_route(&logits, m, None);
16167 {
16168 let mut st = m.stats.borrow_mut();
16169 if st.len() < ne {
16170 st.resize(ne, 0);
16171 }
16172 for &e in &idx {
16173 st[e] += 1;
16174 }
16175 }
16176 let w: Vec<f32> = idx
16177 .iter()
16178 .map(|&e| p[e] / wsum * m.per_expert_scale.as_ref().map_or(1.0, |v| v[e]))
16179 .collect();
16180 routes.push((idx, w));
16181 }
16182 if routes.iter().any(|(idx, _)| idx.is_empty()) {
16183 per_row(&mut out);
16184 return out;
16185 }
16186 let mut experts: Vec<usize> = Vec::new();
16188 let mut groups: Vec<Vec<usize>> = Vec::new();
16189 for (r, (idx, _)) in routes.iter().enumerate() {
16190 for &e in idx {
16191 match experts.iter().position(|&x| x == e) {
16192 Some(g) => groups[g].push(r),
16193 None => {
16194 experts.push(e);
16195 groups.push(vec![r]);
16196 }
16197 }
16198 }
16199 }
16200 let n_pairs: usize = groups.iter().map(|g| g.len()).sum();
16201 let inter = m.experts[experts[0]].gate_proj.rows();
16202 let pairs: Vec<(&QTensor, &QTensor)> = experts
16203 .iter()
16204 .map(|&e| (&m.experts[e].gate_proj, &m.experts[e].up_proj))
16205 .collect();
16206 let mut gs: Vec<Vec<f32>> = (0..n_pairs).map(|_| vec![0f32; inter]).collect();
16207 if !QTensor::moe_gate_up_rows(&pairs, &groups, xs, &mut gs, pool) {
16208 per_row(&mut out);
16209 return out;
16210 }
16211 let downs: Vec<&QTensor> = experts.iter().map(|&e| &m.experts[e].down_proj).collect();
16212 let lens: Vec<usize> = groups.iter().map(|g| g.len()).collect();
16213 let mut ds: Vec<Vec<f32>> = (0..n_pairs).map(|_| vec![0f32; hidden]).collect();
16214 if !QTensor::moe_down_rows(&downs, &lens, &gs, &mut ds, pool) {
16215 per_row(&mut out);
16216 return out;
16217 }
16218 let mut slot = std::collections::HashMap::with_capacity(n_pairs);
16220 let mut p = 0usize;
16221 for (g, &e) in experts.iter().enumerate() {
16222 for &r in &groups[g] {
16223 slot.insert((r, e), p);
16224 p += 1;
16225 }
16226 }
16227 for (r, (idx, w)) in routes.iter().enumerate() {
16228 let terms: Vec<(&[f32], f32)> = idx
16229 .iter()
16230 .zip(w)
16231 .map(|(&e, &we)| (ds[slot[&(r, e)]].as_slice(), we))
16232 .collect();
16233 let row = &mut out[r * hidden..(r + 1) * hidden];
16234 for (i, dst) in row.iter_mut().enumerate() {
16235 let mut acc = 0f32;
16237 for (d, we) in &terms {
16238 acc += we * d[i];
16239 }
16240 *dst = acc;
16241 }
16242 }
16243 out
16244}
16245
16246thread_local! {
16247 static FFN_SCRATCH: std::cell::RefCell<[Vec<f32>; 4]> =
16251 const { std::cell::RefCell::new([Vec::new(), Vec::new(), Vec::new(), Vec::new()]) };
16252}
16253
16254fn dense_ffn(d: &DenseFfn, x: &[f32], pool: Option<&Pool>) -> Vec<f32> {
16256 if gate_topk() > 0
16259 && let Some(out) = dense_ffn_dynamic(d, x, pool, gate_topk())
16260 {
16261 return out;
16262 }
16263 let prism_body = d.gate_proj.has_prism_contract()
16279 || d.up_proj.has_prism_contract()
16280 || d.down_proj.has_prism_contract();
16281 if !prism_body
16282 && crate::gpu::enabled_here()
16283 && (d.gate_proj.rows() >= crate::gpu::min_rows() || d.gate_proj.is_q1())
16284 {
16285 let arm = if d.gate_proj.is_q1() && crate::gpu::q1_force() {
16286 crate::gpu::ProbeArm::Gpu
16287 } else {
16288 crate::gpu::probe_arm(crate::gpu::OpClass::Ffn)
16289 };
16290 match arm {
16291 crate::gpu::ProbeArm::Gpu => {
16292 let t0 = std::time::Instant::now();
16293 if let Some(out) = dense_ffn_gpu(d, x, pool) {
16294 crate::gpu::probe_record(crate::gpu::OpClass::Ffn, true, t0.elapsed());
16295 return out;
16296 }
16297 crate::gpu::probe_note_decline(crate::gpu::OpClass::Ffn);
16301 }
16302 crate::gpu::ProbeArm::CpuTimed => {
16303 let t0 = std::time::Instant::now();
16304 let out = crate::gpu::cpu_scope(|| dense_ffn_cpu(d, x, pool));
16305 crate::gpu::probe_record(crate::gpu::OpClass::Ffn, false, t0.elapsed());
16306 return out;
16307 }
16308 crate::gpu::ProbeArm::Cpu => {
16309 return crate::gpu::cpu_scope(|| dense_ffn_cpu(d, x, pool));
16310 }
16311 }
16312 }
16313 dense_ffn_cpu(d, x, pool)
16314}
16315
16316fn dense_ffn_cpu(d: &DenseFfn, x: &[f32], pool: Option<&Pool>) -> Vec<f32> {
16318 let inter = d.gate_proj.rows();
16319 FFN_SCRATCH.with(|s| {
16320 let mut s = s.borrow_mut();
16321 let [g, u, ..] = &mut *s;
16322 g.resize(inter, 0.0);
16323 if gate_topk() > 0 {
16326 u.resize(inter, 0.0);
16330 QTensor::matvec_many([&d.gate_proj, &d.up_proj], x, [g, u], pool);
16331 for i in 0..inter {
16332 g[i] = Act::Silu.combine(g[i], 1.0);
16333 }
16334 keep_top_k(g, gate_topk());
16335 for i in 0..inter {
16336 g[i] *= u[i];
16337 }
16338 } else if d.act == Act::Silu && {
16339 let _prof = crate::cpuprof::time(crate::cpuprof::Slot::FfnGateUp);
16340 QTensor::matvec_silu_mul(&d.gate_proj, &d.up_proj, x, g, pool)
16341 } {
16342 } else {
16344 u.resize(inter, 0.0);
16345 let _prof = crate::cpuprof::time(crate::cpuprof::Slot::FfnGateUp);
16347 QTensor::matvec_many([&d.gate_proj, &d.up_proj], x, [g, u], pool);
16348 for i in 0..inter {
16349 g[i] = d.act.combine(g[i], u[i]);
16350 }
16351 }
16352 FFN_PROBE.with(|pr| {
16360 if let Some(acc) = pr.borrow_mut().as_mut() {
16361 let li = crate::gpu::cur_layer();
16362 if li >= 0 {
16363 if let Some(row) = acc.get_mut(li as usize) {
16364 match probe_topk() {
16365 0 if probe_sq() => {
16366 for (a, &v) in row.iter_mut().zip(g.iter()) {
16367 *a += (v as f64) * (v as f64);
16368 }
16369 }
16370 0 if probe_signed() => {
16371 for (a, &v) in row.iter_mut().zip(g.iter()) {
16372 *a += v as f64;
16373 }
16374 }
16375 0 => {
16376 for (a, &v) in row.iter_mut().zip(g.iter()) {
16377 *a += (v as f64).abs();
16378 }
16379 }
16380 k => {
16381 let n = g.len();
16382 let k = k.min(n);
16383 let mut mag: Vec<f32> = g.iter().map(|v| v.abs()).collect();
16384 let (_, kth, _) = mag.select_nth_unstable_by(k - 1, |a, b| {
16385 b.partial_cmp(a).unwrap_or(std::cmp::Ordering::Equal)
16386 });
16387 let thr = *kth;
16388 for (a, &v) in row.iter_mut().zip(g.iter()) {
16389 if v.abs() >= thr {
16390 *a += 1.0;
16391 }
16392 }
16393 }
16394 }
16395 }
16396 }
16397 }
16398 });
16399 if oracle_topk() > 0 {
16400 keep_top_k(g, oracle_topk());
16401 }
16402 {
16403 let li = crate::gpu::cur_layer();
16404 if li >= 0 {
16405 adump_row(li as usize, g);
16406 }
16407 }
16408 let mut out = attention::take_buf(d.down_proj.rows());
16409 let _prof = crate::cpuprof::time(crate::cpuprof::Slot::FfnDown);
16410 d.down_proj.matvec(g, &mut out, pool);
16411 out
16412 })
16413}
16414
16415pub struct RefitAcc {
16428 pub support: Vec<u32>,
16429 pub gss: Vec<f32>,
16430 pub ya: Vec<f32>,
16431 pub hidden: usize,
16432 pub tokens: u64,
16433 pub buf_g: Vec<f32>,
16439 pub buf_o: Vec<f32>,
16440 pub buf_t: usize,
16441}
16442
16443type RefitState = (std::collections::HashMap<usize, RefitAcc>, Vec<f32>);
16447
16448static REFIT: std::sync::OnceLock<Option<(String, std::sync::Mutex<RefitState>)>> =
16449 std::sync::OnceLock::new();
16450
16451fn ffn_probe_active() -> bool {
16454 FFN_PROBE.with(|p| p.borrow().is_some())
16455}
16456
16457fn refit_dir() -> Option<&'static (String, std::sync::Mutex<RefitState>)> {
16458 REFIT
16459 .get_or_init(|| {
16460 std::env::var("CMF_FFN_REFIT").ok().map(|d| {
16461 (
16462 d,
16463 std::sync::Mutex::new((std::collections::HashMap::new(), Vec::new())),
16464 )
16465 })
16466 })
16467 .as_ref()
16468}
16469
16470fn refit_accumulate(
16472 li: usize,
16473 g: &[f32],
16474 b: usize,
16475 inter: usize,
16476 out: &[f32],
16477 hidden: usize,
16478 pool: Option<&Pool>,
16479) {
16480 let Some((dir, map)) = refit_dir() else {
16481 return;
16482 };
16483 static SPAN: std::sync::OnceLock<(usize, usize)> = std::sync::OnceLock::new();
16484 let (from, to) = *SPAN.get_or_init(|| {
16485 let g = |k: &str, d: usize| {
16486 std::env::var(k)
16487 .ok()
16488 .and_then(|v| v.parse().ok())
16489 .unwrap_or(d)
16490 };
16491 (
16492 g("CMF_FFN_REFIT_FROM", 0),
16493 g("CMF_FFN_REFIT_TO", usize::MAX),
16494 )
16495 });
16496 if li < from || li > to {
16497 return;
16498 }
16499 let mut guard = map.lock().unwrap();
16500 let (map, shared) = &mut *guard;
16501 let acc = match map.entry(li) {
16502 std::collections::hash_map::Entry::Occupied(e) => e.into_mut(),
16503 std::collections::hash_map::Entry::Vacant(e) => {
16504 let path = format!("{dir}/support.{li}.u32");
16505 let Ok(bytes) = std::fs::read(&path) else {
16506 eprintln!("refit: no {path} — layer {li} skipped");
16507 return;
16508 };
16509 let n = u32::from_le_bytes(bytes[0..4].try_into().unwrap()) as usize;
16510 let support: Vec<u32> = bytes[4..4 + n * 4]
16511 .chunks_exact(4)
16512 .map(|c| u32::from_le_bytes([c[0], c[1], c[2], c[3]]))
16513 .collect();
16514 eprintln!(
16515 "refit: layer {li} support {n} ({:.0} MB of accumulator)",
16516 (n * n + hidden * n) as f64 * 4.0 / 1e6
16517 );
16518 e.insert(RefitAcc {
16519 gss: vec![0.0; n * n],
16520 ya: vec![0.0; hidden * n],
16521 buf_g: Vec::new(),
16522 buf_o: Vec::new(),
16523 buf_t: 0,
16524 support,
16525 hidden,
16526 tokens: 0,
16527 })
16528 }
16529 };
16530 let ns = acc.support.len();
16531 let cap = refit_batch();
16533 if acc.buf_g.is_empty() {
16534 acc.buf_g = vec![0.0; ns * cap];
16535 acc.buf_o = vec![0.0; hidden * cap];
16536 }
16537 let take = b.min(cap - acc.buf_t);
16538 for t in 0..take {
16539 let col = acc.buf_t + t;
16540 for (j, &n) in acc.support.iter().enumerate() {
16541 acc.buf_g[j * cap + col] = g[t * inter + n as usize];
16542 }
16543 for h in 0..hidden {
16544 acc.buf_o[h * cap + col] = out[t * hidden + h];
16545 }
16546 }
16547 acc.buf_t += take;
16548 acc.tokens += take as u64;
16549 if acc.buf_t < cap {
16550 return;
16551 }
16552 let bt = acc.buf_t;
16553 acc.buf_t = 0;
16554 let RefitAcc {
16564 gss,
16565 ya,
16566 buf_g,
16567 buf_o,
16568 ..
16569 } = acc;
16570 let need = (ns * ns).max(hidden * ns);
16571 if shared.len() < need {
16572 shared.resize(need, 0.0);
16573 }
16574 let scratch = &mut shared[..];
16575 let _ = bt;
16576 if crate::gpu::gemm_nt_f32_transient(buf_g, buf_g, &mut scratch[..ns * ns], ns, cap, ns) {
16577 add_into(gss, &scratch[..ns * ns], pool);
16578 if crate::gpu::gemm_nt_f32_transient(
16579 buf_o,
16580 buf_g,
16581 &mut scratch[..hidden * ns],
16582 hidden,
16583 cap,
16584 ns,
16585 ) {
16586 add_into(ya, &scratch[..hidden * ns], pool);
16587 } else {
16588 accum_outer_t(ya, hidden, ns, cap, buf_o, buf_g, pool);
16589 }
16590 } else {
16591 accum_outer_t(gss, ns, ns, cap, buf_g, buf_g, pool);
16592 accum_outer_t(ya, hidden, ns, cap, buf_o, buf_g, pool);
16593 }
16594 }
16598
16599fn refit_batch() -> usize {
16601 static B: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
16602 *B.get_or_init(|| {
16603 std::env::var("CMF_FFN_REFIT_BATCH")
16604 .ok()
16605 .and_then(|v| v.parse().ok())
16606 .unwrap_or(4096)
16607 })
16608}
16609
16610fn accum_outer_t(
16613 c: &mut [f32],
16614 m: usize,
16615 n: usize,
16616 b: usize,
16617 left: &[f32],
16618 right: &[f32],
16619 pool: Option<&Pool>,
16620) {
16621 let ptr = SendMut(c.as_mut_ptr());
16622 let body = |i: usize| {
16623 let ptr = &ptr;
16624 let row = unsafe { std::slice::from_raw_parts_mut(ptr.0.add(i * n), n) };
16625 for t in 0..b {
16626 let a = left[i * b + t];
16627 if a == 0.0 {
16628 continue;
16629 }
16630 for (j, o) in row.iter_mut().enumerate() {
16631 *o += a * right[j * b + t];
16632 }
16633 }
16634 };
16635 match pool {
16636 Some(p) if m > 1 => p.run_rows(m, &|s, e| {
16637 for i in s..e {
16638 body(i);
16639 }
16640 }),
16641 _ => {
16642 for i in 0..m {
16643 body(i);
16644 }
16645 }
16646 }
16647}
16648
16649fn add_into(dst: &mut [f32], src: &[f32], pool: Option<&Pool>) {
16652 let n = dst.len().min(src.len());
16653 match pool {
16654 Some(p) if n >= 1 << 16 => {
16655 let ptr = SendMut(dst.as_mut_ptr());
16656 let f = |s: usize, e: usize| {
16657 let ptr = &ptr;
16658 for blk in s..e {
16659 let (a, b) = (blk * 4096, ((blk + 1) * 4096).min(n));
16660 for i in a..b {
16661 unsafe { *ptr.0.add(i) += src[i] };
16662 }
16663 }
16664 };
16665 p.run_rows(n.div_ceil(4096), &f);
16666 }
16667 _ => {
16668 for (d, v) in dst.iter_mut().zip(&src[..n]) {
16669 *d += *v;
16670 }
16671 }
16672 }
16673}
16674
16675fn accum_outer(
16680 c: &mut [f32],
16681 m: usize,
16682 n: usize,
16683 b: usize,
16684 left: &[f32],
16685 right: &[f32],
16686 pool: Option<&Pool>,
16687) {
16688 const TILE: usize = 32;
16689 let tiles = m.div_ceil(TILE);
16690 let cp = SendMut(c.as_mut_ptr());
16691 let body = |ti: usize| {
16692 let cp = &cp;
16693 let i0 = ti * TILE;
16694 let i1 = (i0 + TILE).min(m);
16695 for t in 0..b {
16696 let r = &right[t * n..t * n + n];
16697 for i in i0..i1 {
16698 let a = left[i * b + t];
16699 if a == 0.0 {
16700 continue;
16701 }
16702 let row = unsafe { std::slice::from_raw_parts_mut(cp.0.add(i * n), n) };
16704 for (o, v) in row.iter_mut().zip(r) {
16705 *o += a * *v;
16706 }
16707 }
16708 }
16709 };
16710 match pool {
16711 Some(p) if tiles > 1 => p.run_rows(tiles, &|s, e| {
16712 for ti in s..e {
16713 body(ti);
16714 }
16715 }),
16716 _ => {
16717 for ti in 0..tiles {
16718 body(ti);
16719 }
16720 }
16721 }
16722}
16723
16724pub fn refit_flush() -> usize {
16726 let Some((dir, map)) = refit_dir() else {
16727 return 0;
16728 };
16729 let guard = map.lock().unwrap();
16730 let mut n = 0;
16731 for (li, acc) in guard.0.iter() {
16732 let w = |name: &str, v: &[f32]| {
16735 let path = format!("{dir}/{name}.{li}.f32");
16736 let bytes: Vec<u8> = v.iter().flat_map(|x| x.to_le_bytes()).collect();
16737 match std::fs::write(&path, &bytes) {
16738 Ok(()) => {}
16739 Err(e) => eprintln!(
16740 "refit: FAILED to write {path} ({} MB): {e}",
16741 bytes.len() / 1_000_000
16742 ),
16743 }
16744 };
16745 w("gss", &acc.gss);
16746 w("ya", &acc.ya);
16747 println!(
16748 "refit L{li}: {} support, {} tokens, hidden {}",
16749 acc.support.len(),
16750 acc.tokens,
16751 acc.hidden
16752 );
16753 n += 1;
16754 }
16755 n
16756}
16757
16758fn adump_row(li: usize, g: &[f32]) {
16763 use std::io::Write as _;
16764 static FILES: std::sync::OnceLock<
16765 Option<(
16766 String,
16767 std::sync::Mutex<std::collections::HashMap<usize, std::fs::File>>,
16768 )>,
16769 > = std::sync::OnceLock::new();
16770 let Some((prefix, map)) = FILES
16771 .get_or_init(|| {
16772 std::env::var("CMF_FFN_ADUMP")
16773 .ok()
16774 .map(|p| (p, std::sync::Mutex::new(std::collections::HashMap::new())))
16775 })
16776 .as_ref()
16777 else {
16778 return;
16779 };
16780 static SPAN: std::sync::OnceLock<(usize, usize)> = std::sync::OnceLock::new();
16783 let (from, to) = *SPAN.get_or_init(|| {
16784 let g = |k: &str, d: usize| {
16785 std::env::var(k)
16786 .ok()
16787 .and_then(|v| v.parse().ok())
16788 .unwrap_or(d)
16789 };
16790 (
16791 g("CMF_FFN_ADUMP_FROM", 0),
16792 g("CMF_FFN_ADUMP_TO", usize::MAX),
16793 )
16794 });
16795 if li < from || li > to {
16796 return;
16797 }
16798 let mut map = map.lock().unwrap();
16799 let f = map.entry(li).or_insert_with(|| {
16800 std::fs::File::create(format!("{prefix}.{li}.f16")).expect("adump file")
16801 });
16802 let mut bytes = Vec::with_capacity(g.len() * 2);
16803 for v in g {
16804 bytes.extend_from_slice(&cortiq_core::quant::f32_to_f16(*v).to_le_bytes());
16805 }
16806 let _ = f.write_all(&bytes);
16807}
16808
16809fn oracle_topk() -> usize {
16815 static K: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
16816 *K.get_or_init(|| {
16817 std::env::var("CMF_FFN_ORACLE_TOPK")
16818 .ok()
16819 .and_then(|v| v.parse().ok())
16820 .unwrap_or(0)
16821 })
16822}
16823
16824fn gate_topk() -> usize {
16830 static K: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
16831 *K.get_or_init(|| {
16832 std::env::var("CMF_FFN_GATE_TOPK")
16833 .ok()
16834 .and_then(|v| v.parse().ok())
16835 .unwrap_or(0)
16836 })
16837}
16838
16839fn gate_block() -> usize {
16846 static B: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
16847 *B.get_or_init(|| {
16848 std::env::var("CMF_FFN_GATE_BLOCK")
16849 .ok()
16850 .and_then(|v| v.parse().ok())
16851 .unwrap_or(1)
16852 })
16853}
16854
16855fn keep_top_blocks(g: &mut [f32], keep_n: usize, block: usize) {
16857 let n = g.len();
16858 let nb = n.div_ceil(block);
16859 let kb = (keep_n.div_ceil(block)).clamp(1, nb);
16860 if kb >= nb {
16861 return;
16862 }
16863 let mut score: Vec<f32> = (0..nb)
16864 .map(|b| {
16865 g[b * block..((b + 1) * block).min(n)]
16866 .iter()
16867 .map(|v| v * v)
16868 .sum::<f32>()
16869 })
16870 .collect();
16871 let mut ord = score.clone();
16872 let (_, kth, _) = ord.select_nth_unstable_by(kb - 1, |a, b| {
16873 b.partial_cmp(a).unwrap_or(std::cmp::Ordering::Equal)
16874 });
16875 let thr = *kth;
16876 for b in 0..nb {
16877 if score[b] < thr {
16878 g[b * block..((b + 1) * block).min(n)].fill(0.0);
16879 }
16880 }
16881 score.clear();
16882}
16883
16884fn keep_top_k(g: &mut [f32], k: usize) {
16886 if gate_block() > 1 {
16887 return keep_top_blocks(g, k, gate_block());
16888 }
16889 let n = g.len();
16890 if k == 0 || k >= n {
16891 return;
16892 }
16893 let mut mag: Vec<f32> = g.iter().map(|v| v.abs()).collect();
16894 let (_, kth, _) = mag.select_nth_unstable_by(k - 1, |a, b| {
16895 b.partial_cmp(a).unwrap_or(std::cmp::Ordering::Equal)
16896 });
16897 let thr = *kth;
16898 for v in g.iter_mut() {
16899 if v.abs() < thr {
16900 *v = 0.0;
16901 }
16902 }
16903}
16904
16905fn probe_sq() -> bool {
16909 static S: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
16910 *S.get_or_init(|| std::env::var("CMF_FFN_PROBE_SQ").is_ok())
16911}
16912
16913fn probe_signed() -> bool {
16917 static S: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
16918 *S.get_or_init(|| std::env::var("CMF_FFN_PROBE_SIGNED").is_ok())
16919}
16920
16921fn meanfill() -> Option<&'static (usize, Vec<f32>)> {
16929 static M: std::sync::OnceLock<Option<(usize, Vec<f32>)>> = std::sync::OnceLock::new();
16930 M.get_or_init(|| {
16931 let p = std::env::var("CMF_FFN_MEANFILL").ok()?;
16932 let b = std::fs::read(&p).ok()?;
16933 let inter = u32::from_le_bytes(b[4..8].try_into().ok()?) as usize;
16934 let vals: Vec<f32> = b[8..]
16935 .chunks_exact(4)
16936 .map(|c| f32::from_le_bytes([c[0], c[1], c[2], c[3]]))
16937 .collect();
16938 eprintln!("meanfill: {} value(s), inter {inter}", vals.len());
16939 Some((inter, vals))
16940 })
16941 .as_ref()
16942}
16943
16944fn probe_topk() -> usize {
16947 static K: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
16948 *K.get_or_init(|| {
16949 std::env::var("CMF_FFN_PROBE_TOPK")
16950 .ok()
16951 .and_then(|v| v.parse().ok())
16952 .unwrap_or(0)
16953 })
16954}
16955
16956thread_local! {
16957 static FFN_PROBE: std::cell::RefCell<Option<Vec<Vec<f64>>>> =
16960 const { std::cell::RefCell::new(None) };
16961}
16962
16963fn dense_ffn_dynamic(d: &DenseFfn, x: &[f32], pool: Option<&Pool>, k: usize) -> Option<Vec<f32>> {
16976 if d.gate_proj.has_prism_contract()
16981 || d.up_proj.has_prism_contract()
16982 || d.down_proj.has_prism_contract()
16983 {
16984 return None;
16985 }
16986 let dt = d.down_t.as_ref()?;
16987 let inter = d.gate_proj.rows();
16988 let hidden = dt.cols();
16989 if k == 0 || k >= inter || d.act != Act::Silu {
16990 return None;
16991 }
16992 DYN_SCRATCH.with(|sc| {
16993 let mut sc = sc.borrow_mut();
16994 let DynScratch {
16995 g,
16996 mag,
16997 live,
16998 parts,
16999 } = &mut *sc;
17000 g.resize(inter, 0.0);
17001 d.gate_proj.matvec(x, g, pool);
17002 for v in g.iter_mut() {
17003 *v = inference::silu(*v);
17004 }
17005 mag.clear();
17008 mag.extend(g.iter().map(|v| v.abs()));
17009 let (_, kth, _) = mag.select_nth_unstable_by(k - 1, |a, b| {
17010 b.partial_cmp(a).unwrap_or(std::cmp::Ordering::Equal)
17011 });
17012 let thr = *kth;
17013 live.clear();
17014 live.extend((0..inter as u32).filter(|&n| g[n as usize].abs() >= thr));
17015 let mut out = vec![0.0f32; hidden];
17016 match pool {
17017 Some(p) if live.len() >= 64 => {
17018 let nw = p.n_workers() + 1;
17019 parts.clear();
17020 parts.resize(nw * hidden, 0.0);
17021 let ptr = SendMut(parts.as_mut_ptr());
17022 let n = live.len();
17023 let live_ref: &[u32] = live;
17024 let g_ref: &[f32] = g;
17025 p.run(&|w, workers| {
17026 let chunk = n.div_ceil(workers);
17027 let (s, e) = (w * chunk, ((w + 1) * chunk).min(n));
17028 if s >= e {
17029 return;
17030 }
17031 WORKER_SCRATCH.with(|ws| {
17032 let mut ws = ws.borrow_mut();
17033 let [scratch, acc] = &mut *ws;
17034 scratch.resize(hidden.max(x.len()), 0.0);
17035 acc.clear();
17036 acc.resize(hidden, 0.0);
17037 for (o, &nrm) in live_ref[s..e].iter().enumerate() {
17038 if let Some(&nx) = live_ref[s..e].get(o + 1) {
17041 d.up_proj.prefetch_row(nx as usize);
17042 dt.prefetch_row(nx as usize);
17043 }
17044 let idx = nrm as usize;
17045 let up = d.up_proj.row_dot(idx, x, scratch);
17046 let a = g_ref[idx] * up;
17047 if a != 0.0 {
17048 dt.add_row_scaled(idx, a, acc, scratch);
17049 }
17050 }
17051 for (j, v) in acc.iter().enumerate() {
17052 unsafe { *ptr.at(w * hidden + j) = *v };
17053 }
17054 });
17055 });
17056 for w in 0..nw {
17057 for (j, o) in out.iter_mut().enumerate() {
17058 *o += parts[w * hidden + j];
17059 }
17060 }
17061 }
17062 _ => {
17063 WORKER_SCRATCH.with(|ws| {
17064 let mut ws = ws.borrow_mut();
17065 let [scratch, _acc] = &mut *ws;
17066 scratch.resize(hidden.max(x.len()), 0.0);
17067 for &nrm in live.iter() {
17068 let idx = nrm as usize;
17069 let up = d.up_proj.row_dot(idx, x, scratch);
17070 let a = g[idx] * up;
17071 if a != 0.0 {
17072 dt.add_row_scaled(idx, a, &mut out, scratch);
17073 }
17074 }
17075 });
17076 }
17077 }
17078 Some(out)
17079 })
17080}
17081
17082struct DynScratch {
17085 g: Vec<f32>,
17086 mag: Vec<f32>,
17087 live: Vec<u32>,
17088 parts: Vec<f32>,
17089}
17090
17091thread_local! {
17092 static DYN_SCRATCH: std::cell::RefCell<DynScratch> = const {
17093 std::cell::RefCell::new(DynScratch {
17094 g: Vec::new(),
17095 mag: Vec::new(),
17096 live: Vec::new(),
17097 parts: Vec::new(),
17098 })
17099 };
17100 static WORKER_SCRATCH: std::cell::RefCell<[Vec<f32>; 2]> =
17102 const { std::cell::RefCell::new([Vec::new(), Vec::new()]) };
17103}
17104
17105fn dense_ffn_masked(d: &DenseFfn, x: &[f32], pool: Option<&Pool>, mask_row: &[u8]) -> Vec<f32> {
17110 let inter = d.gate_proj.rows();
17111 FFN_SCRATCH.with(|s| {
17112 let mut s = s.borrow_mut();
17113 let [g, u, ..] = &mut *s;
17114 g.resize(inter, 0.0);
17115 if d.act == Act::Silu && QTensor::matvec_silu_mul(&d.gate_proj, &d.up_proj, x, g, pool) {
17116 } else {
17118 u.resize(inter, 0.0);
17119 QTensor::matvec_many([&d.gate_proj, &d.up_proj], x, [g, u], pool);
17120 for i in 0..inter {
17121 g[i] = d.act.combine(g[i], u[i]);
17122 }
17123 }
17124 zero_masked_cols(g, 1, inter, mask_row);
17125 let mut out = attention::take_buf(d.down_proj.rows());
17126 d.down_proj.matvec(g, &mut out, pool);
17127 out
17128 })
17129}
17130
17131fn dense_ffn_gpu(d: &DenseFfn, x: &[f32], _pool: Option<&Pool>) -> Option<Vec<f32>> {
17137 if d.gate_proj.has_prism_contract()
17138 || d.up_proj.has_prism_contract()
17139 || d.down_proj.has_prism_contract()
17140 {
17141 return None;
17142 }
17143 if d.act != Act::Silu {
17145 return None;
17146 }
17147 if d.gate_proj.rows() < crate::gpu::min_rows() && !d.gate_proj.is_q1() {
17150 return None;
17151 }
17152 let mut jobs: Vec<crate::gpu::MoeJob> = Vec::with_capacity(1);
17153 let mut model_ref = None;
17154 moe_push_job(d, x, 1.0, &mut jobs, &mut model_ref)?;
17155 let model = model_ref?;
17156 let hidden = jobs[0].down.1;
17157 let mut out = attention::take_buf(hidden);
17158 if crate::gpu::moe_block(&model, &jobs, &mut out) {
17159 Some(out)
17160 } else {
17161 let mut out = out;
17162 attention::recycle_buf(&mut out);
17163 None
17164 }
17165}
17166
17167#[allow(clippy::type_complexity)]
17172#[allow(clippy::type_complexity)]
17173pub(crate) fn moe_parts(
17174 t: &QTensor,
17175) -> Option<(
17176 &std::sync::Arc<cortiq_core::CmfModel>,
17177 usize,
17178 usize,
17179 usize,
17180 &[f32],
17181 &[f32],
17182 bool,
17183 bool,
17184 bool,
17185)> {
17186 match t {
17187 QTensor::Mapped {
17188 model,
17189 idx,
17190 dtype: dt @ (cortiq_core::TensorDtype::Q8_2f | cortiq_core::TensorDtype::Q8Row),
17191 rows,
17192 cols,
17193 row_scale,
17194 col_field,
17195 ..
17196 } if (*dt == cortiq_core::TensorDtype::Q8Row) || !col_field.is_empty() => Some((
17197 model, *idx, *rows, *cols, row_scale, col_field, false, false, false,
17198 )),
17199 QTensor::Mapped {
17201 model,
17202 idx,
17203 dtype: cortiq_core::TensorDtype::Q1,
17204 rows,
17205 cols,
17206 ..
17207 } => Some((
17208 model,
17209 *idx,
17210 *rows,
17211 *cols,
17212 &[][..],
17213 &[][..],
17214 true,
17215 false,
17216 false,
17217 )),
17218 QTensor::Mapped {
17220 model,
17221 idx,
17222 dtype: cortiq_core::TensorDtype::Q4Tiled,
17223 rows,
17224 cols,
17225 ..
17226 } => Some((
17227 model,
17228 *idx,
17229 *rows,
17230 *cols,
17231 &[][..],
17232 &[][..],
17233 false,
17234 true,
17235 false,
17236 )),
17237 QTensor::Mapped {
17239 model,
17240 idx,
17241 dtype: cortiq_core::TensorDtype::Q4TiledP,
17242 rows,
17243 cols,
17244 ..
17245 } => Some((
17246 model,
17247 *idx,
17248 *rows,
17249 *cols,
17250 &[][..],
17251 &[][..],
17252 false,
17253 true,
17254 false,
17255 )),
17256 QTensor::Mapped {
17260 model,
17261 idx,
17262 dtype: cortiq_core::TensorDtype::Q2TiledP,
17263 rows,
17264 cols,
17265 ..
17266 } => Some((
17267 model,
17268 *idx,
17269 *rows,
17270 *cols,
17271 &[][..],
17272 &[][..],
17273 false,
17274 true,
17275 true,
17276 )),
17277 _ => None,
17278 }
17279}
17280
17281#[cfg(target_os = "macos")]
17289fn metal_moe_graph_parts(m: &MoeFfn, hidden: usize) -> Option<crate::gpu::GpuMoe<'_>> {
17290 if m.router_input_norm
17291 || m.route_tau.is_some()
17292 || m.mask.is_some()
17293 || m.per_expert_scale.is_some()
17294 || m.experts.is_empty()
17295 || m.top_k == 0
17296 || m.resonance.is_some()
17297 {
17298 return None;
17299 }
17300 let (sh, sg) = match &m.shared {
17303 Some((sh, sg)) => (sh, sg.as_ref()),
17304 None => return None,
17305 };
17306 let (rf, rr, rc) = m.router.f32_parts()?;
17307 if rr != m.experts.len() || rc != hidden {
17308 return None;
17309 }
17310 let shared_gated = sg.is_some();
17311 let sf = match sg {
17312 Some(sg) => {
17313 let (sf, sr, sc) = sg.f32_parts()?;
17314 if sr * sc != hidden {
17315 return None;
17316 }
17317 sf
17318 }
17319 None => &rf[..hidden],
17322 };
17323 if let Some(b) = &m.expert_bias {
17324 if b.len() != m.experts.len() {
17325 return None;
17326 }
17327 }
17328 let inter = m.experts[0].gate_proj.rows();
17329 let gu_q2 = m.experts[0].gate_proj.mapped_q2tp().is_some();
17332 let trio = |e: &DenseFfn| -> Option<(usize, usize, usize)> {
17333 if e.act != Act::Silu
17334 || e.gate_proj.rows() != inter
17335 || e.gate_proj.cols() != hidden
17336 || e.up_proj.rows() != inter
17337 || e.up_proj.cols() != hidden
17338 || e.down_proj.rows() != hidden
17339 || e.down_proj.cols() != inter
17340 {
17341 return None;
17342 }
17343 let pick = |t: &QTensor| -> Option<usize> {
17344 if gu_q2 {
17345 t.mapped_q2tp().map(|(_, i)| i)
17346 } else {
17347 t.mapped_q4tp().map(|(_, i)| i)
17348 }
17349 };
17350 Some((
17351 pick(&e.gate_proj)?,
17352 pick(&e.up_proj)?,
17353 e.down_proj.mapped_q4tp().map(|(_, i)| i)?,
17354 ))
17355 };
17356 let experts = m.experts.iter().map(trio).collect::<Option<Vec<_>>>()?;
17357 let shared = trio(sh)?;
17358 Some(crate::gpu::GpuMoe {
17359 router: rf,
17360 sgate: sf,
17361 experts,
17362 shared,
17363 n_exp: m.experts.len(),
17364 top_k: m.top_k,
17365 inter,
17366 norm_topk: m.norm_topk_prob,
17367 route_scale: m.routed_scaling,
17368 gu_q2,
17369 sigmoid: m.router_sigmoid,
17370 bias: m.expert_bias.as_deref(),
17371 shared_gated,
17372 })
17373}
17374
17375pub(crate) fn moe_push_job_parts<'a>(
17379 gate: &'a QTensor,
17380 up: &'a QTensor,
17381 down: &'a QTensor,
17382 x: &[f32],
17383 w: f32,
17384 swiglu_limit: f32,
17385 jobs: &mut Vec<crate::gpu::MoeJob<'a>>,
17386 model_ref: &mut Option<std::sync::Arc<cortiq_core::CmfModel>>,
17387) -> Option<()> {
17388 use crate::qtensor::prescale;
17389 let (gm, gi, gr, gc, grs, gcf, gq1, gq4, gq2) = moe_parts(gate)?;
17390 let (_, ui, ur, uc, urs, ucf, uq1, uq4, uq2) = moe_parts(up)?;
17391 let (_, di, dr, dc, drs, dcf, dq1, dq4, dq2) = moe_parts(down)?;
17392 if gq1 != uq1 || uq1 != dq1 || gq4 != uq4 || uq4 != dq4 || gq2 != uq2 {
17393 return None; }
17395 if gq2 && (dq2 || !dq4 || down.mapped_q4tp().is_none()) {
17398 return None;
17399 }
17400 if !gq2 && dq2 {
17401 return None;
17402 }
17403 model_ref.get_or_insert_with(|| gm.clone());
17404 let dt = |cf: &[f32]| {
17405 if cf.is_empty() {
17406 cortiq_core::TensorDtype::Q8Row
17407 } else {
17408 cortiq_core::TensorDtype::Q8_2f
17409 }
17410 };
17411 jobs.push(crate::gpu::MoeJob {
17412 gate: (gi, gr, gc, grs),
17413 up: (ui, ur, uc, urs),
17414 down: (di, dr, dc, drs),
17415 xs_gate: prescale(x, gcf, dt(gcf)).into_owned(),
17416 xs_up: prescale(x, ucf, dt(ucf)).into_owned(),
17417 down_col: dcf,
17418 w,
17419 q1: gq1,
17420 q4t: gq4 && !gq2 && gate.mapped_q4tp().is_none(),
17421 q4tp: gq4 && (gq2 || gate.mapped_q4tp().is_some()),
17422 gu_q2: gq2,
17423 swiglu_limit,
17424 });
17425 Some(())
17426}
17427
17428fn moe_push_job<'a>(
17430 d: &'a DenseFfn,
17431 x: &[f32],
17432 w: f32,
17433 jobs: &mut Vec<crate::gpu::MoeJob<'a>>,
17434 model_ref: &mut Option<std::sync::Arc<cortiq_core::CmfModel>>,
17435) -> Option<()> {
17436 use crate::qtensor::prescale;
17437 if d.act != Act::Silu {
17438 return None; }
17440 let (gm, gi, gr, gc, grs, gcf, gq1, gq4, gq2) = moe_parts(&d.gate_proj)?;
17441 let (_, ui, ur, uc, urs, ucf, uq1, uq4, uq2) = moe_parts(&d.up_proj)?;
17442 let (_, di, dr, dc, drs, dcf, dq1, dq4, dq2) = moe_parts(&d.down_proj)?;
17443 if gq1 != uq1 || uq1 != dq1 || gq4 != uq4 || uq4 != dq4 || gq2 != uq2 {
17444 return None; }
17446 if gq2 && (dq2 || !dq4 || d.down_proj.mapped_q4tp().is_none()) {
17447 return None;
17448 }
17449 if !gq2 && dq2 {
17450 return None;
17451 }
17452 model_ref.get_or_insert_with(|| gm.clone());
17453 let gdt = if gcf.is_empty() {
17454 cortiq_core::TensorDtype::Q8Row
17455 } else {
17456 cortiq_core::TensorDtype::Q8_2f
17457 };
17458 let udt = if ucf.is_empty() {
17459 cortiq_core::TensorDtype::Q8Row
17460 } else {
17461 cortiq_core::TensorDtype::Q8_2f
17462 };
17463 jobs.push(crate::gpu::MoeJob {
17464 gate: (gi, gr, gc, grs),
17465 up: (ui, ur, uc, urs),
17466 down: (di, dr, dc, drs),
17467 xs_gate: prescale(x, gcf, gdt).into_owned(),
17468 xs_up: prescale(x, ucf, udt).into_owned(),
17469 down_col: dcf,
17470 w,
17471 q1: gq1,
17472 q4t: gq4 && !gq2 && d.gate_proj.mapped_q4tp().is_none(),
17473 q4tp: gq4 && (gq2 || d.gate_proj.mapped_q4tp().is_some()),
17474 gu_q2: gq2,
17475 swiglu_limit: 0.0,
17476 });
17477 Some(())
17478}
17479
17480fn sparse_ffn_quant(
17487 d: &DenseFfn,
17488 x: &[f32],
17489 active: &[u16],
17490 hidden: usize,
17491 pool: Option<&Pool>,
17492) -> Vec<f32> {
17493 let n = active.len();
17494 let inter = d.gate_proj.rows();
17495 let mut act = vec![0.0f32; n];
17496 let need_scratch = !(d.gate_proj.sparse_col_ok() && d.up_proj.sparse_col_ok());
17499 let compute = |ai: usize| -> f32 {
17500 let idx = active[ai] as usize;
17501 if idx >= inter {
17502 return 0.0; }
17504 let mut s = if need_scratch {
17505 vec![0.0f32; hidden]
17506 } else {
17507 Vec::new()
17508 };
17509 let gate = d.gate_proj.row_dot(idx, x, &mut s);
17510 let up = d.up_proj.row_dot(idx, x, &mut s);
17511 d.act.combine(gate, up)
17512 };
17513 match pool {
17514 Some(p) if n >= 256 => {
17515 let ptr = SendMut(act.as_mut_ptr());
17516 p.run(&|widx, nw| {
17517 let chunk = n.div_ceil(nw);
17518 let (s, e) = (widx * chunk, ((widx + 1) * chunk).min(n));
17519 for ai in s..e {
17520 unsafe { *ptr.at(ai) = compute(ai) };
17521 }
17522 });
17523 }
17524 _ => {
17525 for (ai, a) in act.iter_mut().enumerate() {
17526 *a = compute(ai);
17527 }
17528 }
17529 }
17530 let mut out = vec![0.0f32; hidden];
17532 for (ai, &idx) in active.iter().enumerate() {
17533 let w = act[ai];
17534 if w.abs() >= 1e-12 && (idx as usize) < inter {
17535 d.down_proj.add_col_scaled(idx as usize, w, &mut out);
17536 }
17537 }
17538 out
17539}
17540
17541#[doc(hidden)]
17543pub fn sparse_ffn_quant_for_test(
17544 d: &DenseFfn,
17545 x: &[f32],
17546 active: &[u16],
17547 hidden: usize,
17548) -> Vec<f32> {
17549 sparse_ffn_quant(d, x, active, hidden, None)
17550}
17551
17552fn dequant_dense_f32(d: &DenseFfn) -> (Vec<f32>, Vec<f32>, Vec<f32>) {
17556 let deq = |t: &QTensor| -> Vec<f32> {
17557 let (rows, cols) = (t.rows(), t.cols());
17558 let mut out = vec![0.0f32; rows * cols];
17559 for r in 0..rows {
17560 t.row_f32(r, &mut out[r * cols..(r + 1) * cols]);
17561 }
17562 out
17563 };
17564 (deq(&d.gate_proj), deq(&d.up_proj), deq(&d.down_proj))
17565}
17566
17567struct SendMut(*mut f32);
17569unsafe impl Send for SendMut {}
17570unsafe impl Sync for SendMut {}
17571impl SendMut {
17572 #[inline]
17573 #[allow(clippy::mut_from_ref)]
17576 unsafe fn at(&self, i: usize) -> &mut f32 {
17577 unsafe { &mut *self.0.add(i) }
17578 }
17579}
17580
17581pub(crate) fn moe_route(
17594 logits: &[f32],
17595 m: &MoeFfn,
17596 allowed: Option<&[bool]>,
17597) -> (Vec<usize>, Vec<f32>, f32) {
17598 moe_route_with_eps(logits, m, allowed, 1e-6)
17599}
17600
17601pub(crate) fn moe_route_with_eps(
17610 logits: &[f32],
17611 m: &MoeFfn,
17612 allowed: Option<&[bool]>,
17613 sigmoid_denom_eps: f32,
17614) -> (Vec<usize>, Vec<f32>, f32) {
17615 let ne = logits.len();
17616 let admit = |e: usize| {
17622 m.mask.as_ref().is_none_or(|mk| mk[e])
17623 && allowed.is_none_or(|a| a.get(e).copied().unwrap_or(false))
17624 };
17625 if m.resonance.is_some() && m.top_k == 1 {
17638 let mut best: Option<usize> = None;
17639 for e in (0..ne).filter(|&e| admit(e)) {
17640 let l = logits[e];
17641 if l == f32::NEG_INFINITY || l.is_nan() {
17642 continue;
17643 }
17644 if best.is_none_or(|b| l > logits[b]) {
17645 best = Some(e);
17646 }
17647 }
17648 if let Some(b) = best {
17649 let mut p = vec![0.0f32; ne];
17650 p[b] = 1.0;
17651 return (vec![b], p, 1.0 / m.routed_scaling);
17652 }
17653 }
17654 let p: Vec<f32> = if m.router_sigmoid {
17660 logits.iter().map(|&l| 1.0 / (1.0 + (-l).exp())).collect()
17661 } else {
17662 let mx = logits.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
17663 if mx == f32::NEG_INFINITY {
17664 vec![1.0 / ne.max(1) as f32; ne]
17665 } else {
17666 let mut e: Vec<f32> = logits.iter().map(|&l| (l - mx).exp()).collect();
17667 let s: f32 = e.iter().sum();
17668 for v in &mut e {
17669 *v /= s;
17670 }
17671 e
17672 }
17673 };
17674 let mut idx: Vec<usize> = (0..ne).filter(|&e| admit(e)).collect();
17675 match &m.expert_bias {
17677 Some(b) => idx.sort_unstable_by(|&x, &y| {
17678 (p[y] + b[y])
17679 .partial_cmp(&(p[x] + b[x]))
17680 .unwrap()
17681 .then(x.cmp(&y))
17682 }),
17683 None => idx.sort_unstable_by(|&x, &y| p[y].partial_cmp(&p[x]).unwrap().then(x.cmp(&y))),
17684 }
17685 idx.truncate(m.top_k);
17686 if let Some(tau) = m.route_tau {
17690 let total: f32 = idx.iter().map(|&e| p[e]).sum();
17691 if total > 0.0 {
17692 let mut acc = 0.0f32;
17693 let mut keep = idx.len();
17694 for (i, &e) in idx.iter().enumerate() {
17695 acc += p[e];
17696 if acc >= tau * total {
17697 keep = i + 1;
17698 break;
17699 }
17700 }
17701 idx.truncate(keep);
17702 }
17703 }
17704 let wsum: f32 = if m.norm_topk_prob {
17705 let s: f32 = idx.iter().map(|&e| p[e]).sum();
17706 (if m.router_sigmoid {
17710 s + sigmoid_denom_eps
17711 } else {
17712 s
17713 }) / m.routed_scaling
17714 } else {
17715 1.0 / m.routed_scaling
17716 };
17717 (idx, p, wsum)
17718}
17719
17720fn moe_trace(idx: &[usize]) {
17722 moe_trace_at(crate::gpu::cur_layer() as i32, idx)
17723}
17724
17725pub(crate) fn moe_trace_at(li: i32, idx: &[usize]) {
17728 use std::io::Write;
17729 static F: std::sync::OnceLock<Option<std::sync::Mutex<std::fs::File>>> =
17730 std::sync::OnceLock::new();
17731 let Some(f) = F.get_or_init(|| {
17732 let p = std::env::var("CMF_MOE_TRACE").ok()?;
17733 Some(std::sync::Mutex::new(
17734 std::fs::OpenOptions::new()
17735 .create(true)
17736 .append(true)
17737 .open(p)
17738 .ok()?,
17739 ))
17740 }) else {
17741 return;
17742 };
17743 let ids: Vec<String> = idx.iter().map(|e| e.to_string()).collect();
17744 let _ = writeln!(f.lock().unwrap(), "{li}:{}", ids.join(","));
17745}
17746
17747pub(crate) fn moe_ffn(
17750 m: &MoeFfn,
17751 x: &[f32],
17752 pool: Option<&Pool>,
17753 allowed: Option<&[bool]>,
17754) -> Vec<f32> {
17755 let r = moe_ffn_route(m, x, pool, allowed);
17756 moe_ffn_experts(m, x, &r, pool)
17757}
17758
17759pub(crate) struct MoeRoute {
17763 pub idx: Vec<usize>,
17764 pub p: Vec<f32>,
17765 pub wsum: f32,
17766 pub logits: Vec<f32>,
17767}
17768
17769pub(crate) fn moe_ffn_route(
17775 m: &MoeFfn,
17776 x: &[f32],
17777 pool: Option<&Pool>,
17778 allowed: Option<&[bool]>,
17779) -> MoeRoute {
17780 accumulate_act(m, x, 1);
17781 let ne = m.experts.len();
17782 let mut logits = vec![0.0f32; ne];
17783 match &m.resonance {
17784 Some(r) => r.scores(x, &mut logits),
17785 None => m.router.matvec(x, &mut logits, pool),
17786 }
17787 let (idx, p, wsum) = moe_route(&logits, m, allowed);
17788 {
17789 let mut st = m.stats.borrow_mut();
17790 if st.len() < ne {
17791 st.resize(ne, 0);
17792 }
17793 for &e in &idx {
17794 st[e] += 1;
17795 }
17796 }
17797 moe_trace(&idx);
17803 MoeRoute {
17804 idx,
17805 p,
17806 wsum,
17807 logits,
17808 }
17809}
17810
17811pub(crate) fn moe_ffn_experts(
17814 m: &MoeFfn,
17815 x: &[f32],
17816 r: &MoeRoute,
17817 pool: Option<&Pool>,
17818) -> Vec<f32> {
17819 let (idx, p, wsum) = (&r.idx, &r.p, r.wsum);
17820 if crate::gpu::enabled_here() {
17825 match crate::gpu::probe_arm(crate::gpu::OpClass::Ffn) {
17826 crate::gpu::ProbeArm::Gpu => {
17827 let t0 = std::time::Instant::now();
17828 if let Some(out) = moe_ffn_gpu(m, x, &idx, &p, wsum, pool) {
17829 crate::gpu::probe_record(crate::gpu::OpClass::Ffn, true, t0.elapsed());
17830 return out;
17831 }
17832 }
17833 crate::gpu::ProbeArm::CpuTimed => {
17834 let t0 = std::time::Instant::now();
17835 let out = crate::gpu::cpu_scope(|| moe_ffn_cpu(m, x, &idx, &p, wsum, pool));
17836 crate::gpu::probe_record(crate::gpu::OpClass::Ffn, false, t0.elapsed());
17837 return out;
17838 }
17839 crate::gpu::ProbeArm::Cpu => {
17840 return crate::gpu::cpu_scope(|| moe_ffn_cpu(m, x, &idx, &p, wsum, pool));
17841 }
17842 }
17843 }
17844 moe_ffn_cpu(m, x, &idx, &p, wsum, pool)
17845}
17846
17847fn moe_ffn_banked(
17851 slot: &mut crate::mimo_moe::Slot,
17852 li: usize,
17853 m: &MoeFfn,
17854 x: &[f32],
17855 pool: Option<&Pool>,
17856) -> Vec<f32> {
17857 let t0 = std::time::Instant::now();
17858 let r = moe_ffn_route(m, x, pool, None);
17859 slot.note_route(t0.elapsed().as_nanos() as u64);
17860 match slot.forward(li, m, x, &r, pool) {
17861 Some(out) => out,
17862 None => crate::qtensor::float_activations_scope(|| {
17863 crate::gpu::cpu_scope(|| moe_ffn_experts(m, x, &r, pool))
17864 }),
17865 }
17866}
17867
17868fn moe_ffn_banked_rows(
17870 slot: &mut crate::mimo_moe::Slot,
17871 li: usize,
17872 m: &MoeFfn,
17873 xs: &[f32],
17874 b: usize,
17875 hidden: usize,
17876 pool: Option<&Pool>,
17877) -> Vec<f32> {
17878 let t0 = std::time::Instant::now();
17879 let routes: Vec<_> = xs
17880 .chunks_exact(hidden)
17881 .map(|x| moe_ffn_route(m, x, pool, None))
17882 .collect();
17883 slot.note_route(t0.elapsed().as_nanos() as u64);
17884 if let Some(out) = slot.forward_rows(li, m, xs, &routes, pool) {
17885 return out;
17886 }
17887 let mut out = Vec::with_capacity(b * hidden);
17888 for (x, r) in xs.chunks_exact(hidden).zip(&routes) {
17889 let row = slot.forward(li, m, x, r, pool).unwrap_or_else(|| {
17890 crate::qtensor::float_activations_scope(|| {
17892 crate::gpu::cpu_scope(|| moe_ffn_experts(m, x, r, pool))
17893 })
17894 });
17895 out.extend(row);
17896 }
17897 out
17898}
17899
17900fn graph_note(built: bool, layers_run: usize, total_layers: usize) {
17905 use std::sync::atomic::{AtomicBool, Ordering};
17906 if built {
17907 GRAPH_TOK_OK.fetch_add(1, Ordering::Relaxed);
17908 if total_layers > 0 && layers_run < total_layers {
17909 GRAPH_TOK_PREFIX.fetch_add(1, Ordering::Relaxed);
17910 } else {
17911 GRAPH_TOK_FULL.fetch_add(1, Ordering::Relaxed);
17912 }
17913 } else {
17914 GRAPH_TOK_MISS.fetch_add(1, Ordering::Relaxed);
17915 }
17916 static SAID: AtomicBool = AtomicBool::new(false);
17917 if !SAID.swap(true, Ordering::Relaxed) {
17918 if built {
17919 tracing::info!("wgpu whole-token graph: ACTIVE");
17920 } else {
17921 tracing::warn!("wgpu whole-token graph refused — per-op path");
17922 }
17923 }
17924}
17925
17926pub static GRAPH_TOK_OK: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
17930pub static GRAPH_TOK_MISS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
17931pub static GRAPH_TOK_PREFIX: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
17935pub static GRAPH_TOK_FULL: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
17937
17938pub static METAL_GRAPH_TOK_OK: std::sync::atomic::AtomicU64 =
17942 std::sync::atomic::AtomicU64::new(0);
17943pub static METAL_GRAPH_HEAD_OK: std::sync::atomic::AtomicU64 =
17944 std::sync::atomic::AtomicU64::new(0);
17945pub static METAL_GRAPH_HEAD_MISS: std::sync::atomic::AtomicU64 =
17946 std::sync::atomic::AtomicU64::new(0);
17947pub static METAL_GRAPH_LAYERS: std::sync::atomic::AtomicU64 =
17948 std::sync::atomic::AtomicU64::new(0);
17949pub static METAL_GRAPH_ERRORS: std::sync::atomic::AtomicU64 =
17950 std::sync::atomic::AtomicU64::new(0);
17951pub static METAL_PREFILL_CHUNKS: std::sync::atomic::AtomicU64 =
17955 std::sync::atomic::AtomicU64::new(0);
17956pub static METAL_PREFILL_ROWS: std::sync::atomic::AtomicU64 =
17957 std::sync::atomic::AtomicU64::new(0);
17958pub static METAL_PREFILL_HEAD_ROWS: std::sync::atomic::AtomicU64 =
17959 std::sync::atomic::AtomicU64::new(0);
17960pub static METAL_PREFILL_ERRORS: std::sync::atomic::AtomicU64 =
17961 std::sync::atomic::AtomicU64::new(0);
17962
17963fn moe_batch_enabled() -> bool {
17966 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
17967 *ON.get_or_init(|| std::env::var("CMF_MOE_BATCH").as_deref() != Ok("0"))
17968}
17969
17970fn moe_ffn_cpu_batched(
17976 m: &MoeFfn,
17977 x: &[f32],
17978 idx: &[usize],
17979 p: &[f32],
17980 wsum: f32,
17981 pool: Option<&Pool>,
17982) -> Option<Vec<f32>> {
17983 if idx.is_empty() || !moe_batch_enabled() {
17984 return None;
17985 }
17986 if FFN_PROBE.with(|pr| pr.borrow().is_some()) {
17990 return None;
17991 }
17992 let n = idx.len() + usize::from(m.shared.is_some());
17993 let mut pairs = Vec::with_capacity(n);
17994 let mut downs = Vec::with_capacity(n);
17995 let mut ws = Vec::with_capacity(n);
17996 for &e in idx {
17997 let d = &m.experts[e];
17998 if d.act != Act::Silu {
17999 return None;
18000 }
18001 pairs.push((&d.gate_proj, &d.up_proj));
18002 downs.push(&d.down_proj);
18003 ws.push(p[e] / wsum * m.per_expert_scale.as_ref().map_or(1.0, |v| v[e]));
18004 }
18005 if let Some((se, gate)) = &m.shared {
18008 if se.act != Act::Silu {
18009 return None;
18010 }
18011 let g = gate.as_ref().map_or(1.0, |gate| {
18012 let mut gl = [0.0f32; 1];
18013 gate.matvec(x, &mut gl, pool);
18014 1.0 / (1.0 + (-gl[0]).exp())
18015 });
18016 pairs.push((&se.gate_proj, &se.up_proj));
18017 downs.push(&se.down_proj);
18018 ws.push(g);
18019 }
18020 let inter = pairs[0].0.rows();
18021 let mut gs: Vec<Vec<f32>> = (0..pairs.len()).map(|_| vec![0f32; inter]).collect();
18022 if !QTensor::moe_gate_up_many(&pairs, x, &mut gs, pool) {
18023 return None;
18024 }
18025 let mut out = attention::take_buf(x.len());
18026 if !QTensor::moe_down_many(&downs, &gs, &ws, &mut out, pool) {
18027 attention::recycle_buf(&mut out);
18028 return None;
18029 }
18030 Some(out)
18031}
18032
18033pub(crate) fn moe_cold_experts_cpu(
18039 experts: &[(&DenseFfn, f32)],
18040 x: &[f32],
18041 pool: Option<&Pool>,
18042) -> Vec<f32> {
18043 let mut out = attention::take_buf(x.len());
18044 if experts.is_empty() {
18045 return out;
18046 }
18047 let pairs: Vec<_> = experts
18048 .iter()
18049 .map(|(e, _)| (&e.gate_proj, &e.up_proj))
18050 .collect();
18051 let downs: Vec<_> = experts.iter().map(|(e, _)| &e.down_proj).collect();
18052 let weights: Vec<_> = experts.iter().map(|(_, w)| *w).collect();
18053 let inter = experts[0].0.gate_proj.rows();
18054 let mut activations: Vec<Vec<f32>> = (0..experts.len()).map(|_| vec![0.0; inter]).collect();
18055 if QTensor::moe_gate_up_many(&pairs, x, &mut activations, pool)
18056 && QTensor::moe_down_many(&downs, &activations, &weights, &mut out, pool)
18057 {
18058 return out;
18059 }
18060 out.fill(0.0);
18061 for &(expert, weight) in experts {
18062 let mut one = dense_ffn(expert, x, pool);
18063 for (o, v) in out.iter_mut().zip(&one) {
18064 *o += weight * v;
18065 }
18066 attention::recycle_buf(&mut one);
18067 }
18068 out
18069}
18070
18071pub(crate) fn moe_cold_experts_rows_cpu(
18075 jobs: &[Vec<(&DenseFfn, f32)>],
18076 xs: &[f32],
18077 hidden: usize,
18078 pool: Option<&Pool>,
18079) -> Vec<f32> {
18080 let mut out = vec![0.0; xs.len()];
18081 let mut experts: Vec<&DenseFfn> = Vec::new();
18082 let mut groups: Vec<Vec<usize>> = Vec::new();
18083 let mut terms = vec![Vec::new(); jobs.len()];
18084 for (r, row) in jobs.iter().enumerate() {
18085 for &(e, w) in row {
18086 let g = match experts.iter().position(|&d| std::ptr::eq(d, e)) {
18087 Some(g) => g,
18088 None => {
18089 experts.push(e);
18090 groups.push(Vec::new());
18091 groups.len() - 1
18092 }
18093 };
18094 terms[r].push((g, groups[g].len(), w));
18095 groups[g].push(r);
18096 }
18097 }
18098 if experts.is_empty() {
18099 return out;
18100 }
18101 let pairs: Vec<_> = experts.iter().map(|e| (&e.gate_proj, &e.up_proj)).collect();
18102 let downs: Vec<_> = experts.iter().map(|e| &e.down_proj).collect();
18103 let lens: Vec<_> = groups.iter().map(Vec::len).collect();
18104 let count: usize = lens.iter().sum();
18105 let mut acts = vec![vec![0.0; experts[0].gate_proj.rows()]; count];
18106 let mut ds = vec![vec![0.0; hidden]; count];
18107 if QTensor::moe_gate_up_rows(&pairs, &groups, xs, &mut acts, pool)
18108 && QTensor::moe_down_rows(&downs, &lens, &acts, &mut ds, pool)
18109 {
18110 let mut offset = 0;
18111 let offsets: Vec<_> = lens
18112 .iter()
18113 .map(|&n| {
18114 let start = offset;
18115 offset += n;
18116 start
18117 })
18118 .collect();
18119 for (r, terms) in terms.iter().enumerate() {
18120 for &(g, slot, w) in terms {
18121 for (o, &v) in out[r * hidden..(r + 1) * hidden]
18122 .iter_mut()
18123 .zip(&ds[offsets[g] + slot])
18124 {
18125 *o += w * v;
18126 }
18127 }
18128 }
18129 } else {
18130 for (r, jobs) in jobs.iter().enumerate() {
18131 if !jobs.is_empty() {
18132 let mut row = moe_cold_experts_cpu(jobs, &xs[r * hidden..(r + 1) * hidden], pool);
18133 out[r * hidden..(r + 1) * hidden].copy_from_slice(&row);
18134 attention::recycle_buf(&mut row);
18135 }
18136 }
18137 }
18138 out
18139}
18140
18141fn moe_ffn_cpu(
18143 m: &MoeFfn,
18144 x: &[f32],
18145 idx: &[usize],
18146 p: &[f32],
18147 wsum: f32,
18148 pool: Option<&Pool>,
18149) -> Vec<f32> {
18150 if let Some(out) = moe_ffn_cpu_batched(m, x, idx, p, wsum, pool) {
18151 return out;
18152 }
18153 let mut out = attention::take_buf(x.len());
18154 for &e in idx {
18155 let mut eo = dense_ffn(&m.experts[e], x, pool);
18156 let w = p[e] / wsum * m.per_expert_scale.as_ref().map_or(1.0, |v| v[e]);
18157 for i in 0..out.len() {
18158 out[i] += w * eo[i];
18159 }
18160 attention::recycle_buf(&mut eo);
18161 }
18162 if let Some((se, gate)) = &m.shared {
18163 let mut so = dense_ffn(se, x, pool);
18164 let g = gate.as_ref().map_or(1.0, |gate| {
18165 let mut gl = [0.0f32; 1];
18166 gate.matvec(x, &mut gl, pool);
18167 1.0 / (1.0 + (-gl[0]).exp())
18168 });
18169 for i in 0..out.len() {
18170 out[i] += g * so[i];
18171 }
18172 attention::recycle_buf(&mut so);
18173 }
18174 out
18175}
18176
18177#[allow(clippy::too_many_arguments)]
18185pub(crate) fn mla_attention(
18186 w: &MlaWeights,
18187 normed: &[f32],
18188 cache: &mut crate::kv_cache::LayerKvCache,
18189 position: usize,
18190 inv_freq: &[f32],
18191 rope_scale: f32,
18192 eps: f64,
18193 pool: Option<&Pool>,
18194) -> Vec<f32> {
18195 let (nh, dr, dn, dv, lora) = (w.nh, w.qk_rope, w.qk_nope, w.v_dim, w.lora);
18196 let hd = dr + dn;
18197 let mut q = vec![0.0f32; nh * hd];
18198 match (&w.q_a, &w.q_a_norm) {
18199 (Some(qa), Some(qn)) => {
18200 let mut t = vec![0.0f32; qa.rows()];
18201 qa.matvec(normed, &mut t, pool);
18202 let tn = inference::rms_norm(&t, qn, eps, NormStyle::Qwen);
18203 w.q_proj.matvec(&tn, &mut q, pool);
18204 }
18205 _ => w.q_proj.matvec(normed, &mut q, pool),
18206 }
18207 let mut ca = vec![0.0f32; lora + dr];
18208 w.kv_a.matvec(normed, &mut ca, pool);
18209 let (c_lat, k_rope) = ca.split_at_mut(lora);
18210 let latn = inference::rms_norm(c_lat, &w.kv_a_norm, eps, NormStyle::Qwen);
18211 let mut kvb = vec![0.0f32; nh * (dn + dv)];
18212 w.kv_b.matvec(&latn, &mut kvb, pool);
18213 if !w.nope {
18214 attention::rope_rotate_scaled(k_rope, position, inv_freq, rope_scale);
18215 }
18216 for h in 0..nh {
18217 if !w.nope {
18218 attention::rope_rotate_scaled(
18219 &mut q[h * hd..h * hd + dr],
18220 position,
18221 inv_freq,
18222 rope_scale,
18223 );
18224 }
18225 }
18226 let mut k = vec![0.0f32; nh * hd];
18227 let mut v = vec![0.0f32; nh * hd];
18228 for h in 0..nh {
18229 k[h * hd..h * hd + dr].copy_from_slice(k_rope);
18230 k[h * hd + dr..(h + 1) * hd].copy_from_slice(&kvb[h * (dn + dv)..h * (dn + dv) + dn]);
18231 v[h * hd..h * hd + dv].copy_from_slice(&kvb[h * (dn + dv) + dn..(h + 1) * (dn + dv)]);
18232 }
18233 cache.append(&k, &v, &vec![true; nh]);
18234 let (ao, mut imp) = attention::attend_all_heads(&q, cache, nh, 1, hd, w.scale, None, 0.0);
18235 attention::recycle_buf(&mut imp);
18236 let mut ov = vec![0.0f32; nh * dv];
18237 for h in 0..nh {
18238 ov[h * dv..(h + 1) * dv].copy_from_slice(&ao[h * hd..h * hd + dv]);
18239 }
18240 let mut out = vec![0.0f32; w.o_proj.rows()];
18241 w.o_proj.matvec(&ov, &mut out, pool);
18242 out
18243}
18244
18245fn dense_moe_ffn(
18252 dm: &DenseMoeFfn,
18253 x_normed: &[f32],
18254 h_raw: &[f32],
18255 eps: f64,
18256 norm_style: NormStyle,
18257 pool: Option<&Pool>,
18258) -> Vec<f32> {
18259 let mut d = dense_ffn(&dm.dense, x_normed, pool);
18260 d = inference::rms_norm(&d, &dm.post_norm_1, eps, norm_style);
18261 let m = &dm.moe;
18262 let ne = m.experts.len();
18263 let mut logits = vec![0.0f32; ne];
18264 if m.router_input_norm {
18265 let ss: f32 = h_raw.iter().map(|v| v * v).sum::<f32>() / h_raw.len() as f32;
18266 let inv = 1.0 / (ss + eps as f32).sqrt();
18267 let xr: Vec<f32> = h_raw.iter().map(|v| v * inv).collect();
18268 m.router.matvec(&xr, &mut logits, pool);
18269 } else {
18270 m.router.matvec(h_raw, &mut logits, pool);
18271 }
18272 let (idx, p, wsum) = moe_route(&logits, m, None);
18273 {
18274 let mut st = m.stats.borrow_mut();
18275 if st.len() < ne {
18276 st.resize(ne, 0);
18277 }
18278 for &e in &idx {
18279 st[e] += 1;
18280 }
18281 }
18282 let x2 = inference::rms_norm(h_raw, &dm.pre_norm_2, eps, norm_style);
18283 let mo = moe_ffn_cpu(m, &x2, &idx, &p, wsum, pool);
18284 let mo = inference::rms_norm(&mo, &dm.post_norm_2, eps, norm_style);
18285 for (di, mi) in d.iter_mut().zip(&mo) {
18286 *di += mi;
18287 }
18288 d
18289}
18290
18291fn moe_gpu_refused(why: &'static str) {
18298 use std::sync::atomic::{AtomicBool, Ordering};
18299 static SAID: AtomicBool = AtomicBool::new(false);
18300 if !SAID.swap(true, Ordering::Relaxed) {
18301 tracing::warn!("MoE GPU block refused ({why}) — experts run on the CPU");
18302 }
18303}
18304
18305fn moe_ffn_gpu(
18306 m: &MoeFfn,
18307 x: &[f32],
18308 idx: &[usize],
18309 p: &[f32],
18310 wsum: f32,
18311 pool: Option<&Pool>,
18312) -> Option<Vec<f32>> {
18313 use crate::gpu::MoeJob;
18314
18315 let mut jobs: Vec<MoeJob> = Vec::with_capacity(idx.len() + 1);
18316 let mut model_ref = None;
18317 for &e in idx {
18318 if moe_push_job(&m.experts[e], x, p[e] / wsum, &mut jobs, &mut model_ref).is_none() {
18319 moe_gpu_refused("push_job(expert)");
18320 return None;
18321 }
18322 }
18323 if let Some((se, gate)) = &m.shared {
18324 let g = gate.as_ref().map_or(1.0, |gate| {
18325 let mut gl = [0.0f32; 1];
18326 gate.matvec(x, &mut gl, pool);
18327 1.0 / (1.0 + (-gl[0]).exp())
18328 });
18329 if moe_push_job(se, x, g, &mut jobs, &mut model_ref).is_none() {
18330 moe_gpu_refused("push_job(shared)");
18331 return None;
18332 }
18333 }
18334 let Some(model) = model_ref else {
18335 moe_gpu_refused("no model_ref");
18336 return None;
18337 };
18338 let hidden = jobs[0].down.1;
18339 let mut out = vec![0.0f32; hidden];
18340 if crate::gpu::moe_block(&model, &jobs, &mut out) {
18341 Some(out)
18342 } else {
18343 moe_gpu_refused("gpu::moe_block");
18344 None
18345 }
18346}
18347
18348fn ffn_forward(
18350 ffn: &FfnKind,
18351 x: &[f32],
18352 pool: Option<&Pool>,
18353 experts_allowed: Option<&[bool]>,
18354) -> Vec<f32> {
18355 match ffn {
18356 FfnKind::Dense(d) if !d.segs.is_empty() => tube_ffn(d, x, 1, pool, None),
18357 FfnKind::Dense(d) => dense_ffn(d, x, pool),
18358 FfnKind::Moe(m) => moe_ffn(m, x, pool, experts_allowed),
18359 FfnKind::DenseMoe(_) => unreachable!("DenseMoe dispatches via dense_moe_ffn"),
18363 }
18364}
18365
18366fn ffn_forward_pair(
18370 ffn: &FfnKind,
18371 x1: &[f32],
18372 x2: &[f32],
18373 pool: Option<&Pool>,
18374 experts_allowed: Option<&[bool]>,
18375) -> (Vec<f32>, Vec<f32>) {
18376 let d = match ffn {
18377 FfnKind::Dense(d) if !d.segs.is_empty() => {
18380 return (
18381 tube_ffn(d, x1, 1, pool, None),
18382 tube_ffn(d, x2, 1, pool, None),
18383 );
18384 }
18385 FfnKind::Dense(d) => d,
18386 FfnKind::Moe(m) => {
18387 return (
18388 moe_ffn(m, x1, pool, experts_allowed),
18389 moe_ffn(m, x2, pool, experts_allowed),
18390 );
18391 }
18392 FfnKind::DenseMoe(_) => unreachable!("DenseMoe dispatches via dense_moe_ffn"),
18393 };
18394 let inter = d.gate_proj.rows();
18395 FFN_SCRATCH.with(|s| {
18396 let mut s = s.borrow_mut();
18397 let [g1, g2, u1, u2] = &mut *s;
18398 g1.resize(inter, 0.0);
18399 g2.resize(inter, 0.0);
18400 u1.resize(inter, 0.0);
18401 u2.resize(inter, 0.0);
18402 QTensor::matvec2_many(
18405 [&d.gate_proj, &d.up_proj],
18406 x1,
18407 x2,
18408 [g1.as_mut_slice(), u1.as_mut_slice()],
18409 [g2.as_mut_slice(), u2.as_mut_slice()],
18410 pool,
18411 );
18412 for i in 0..inter {
18413 g1[i] = d.act.combine(g1[i], u1[i]);
18414 g2[i] = d.act.combine(g2[i], u2[i]);
18415 }
18416 let mut o1 = attention::take_buf(d.down_proj.rows());
18417 let mut o2 = attention::take_buf(d.down_proj.rows());
18418 d.down_proj.matvec2(g1, g2, &mut o1, &mut o2, pool);
18419 (o1, o2)
18420 })
18421}
18422
18423#[cfg(test)]
18424mod tests {
18425
18426 #[test]
18431 fn prefill_chunk_rule_widens_only_dense_on_discrete() {
18432 use super::{
18433 prefill_chunk_rule, ChunkHost, ChunkStackFacts, DISCRETE_DENSE_PREFILL_CHUNK,
18434 };
18435 let dense_card = ChunkStackFacts {
18436 plain_dense: true,
18437 discrete: true,
18438 gpu_on: true,
18439 ..Default::default()
18440 };
18441 assert!(dense_card.dense_on_discrete());
18442 assert_eq!(
18444 prefill_chunk_rule(None, ChunkHost::Other, dense_card.dense_on_discrete()),
18445 DISCRETE_DENSE_PREFILL_CHUNK
18446 );
18447 assert!(DISCRETE_DENSE_PREFILL_CHUNK > 48);
18448 for (label, facts) in [
18449 ("GDN hybrid / MoE / DeepSeek stack", ChunkStackFacts { plain_dense: false, ..dense_card }),
18450 ("integrated GPU", ChunkStackFacts { discrete: false, ..dense_card }),
18451 ("CPU only", ChunkStackFacts { gpu_on: false, discrete: false, ..dense_card }),
18452 ("capacity split", ChunkStackFacts { capacity_split: true, ..dense_card }),
18453 ("multi-GPU plan", ChunkStackFacts { multi_gpu: true, ..dense_card }),
18454 ("O(1) layers", ChunkStackFacts { o1: true, ..dense_card }),
18455 ] {
18456 assert!(!facts.dense_on_discrete(), "{label}");
18457 assert_eq!(
18458 prefill_chunk_rule(None, ChunkHost::Other, facts.dense_on_discrete()),
18459 48,
18460 "{label} keeps the historical x86 chunk"
18461 );
18462 }
18463 for dense in [false, true] {
18465 assert_eq!(prefill_chunk_rule(None, ChunkHost::Macos, dense), 512);
18466 assert_eq!(prefill_chunk_rule(None, ChunkHost::Aarch64, dense), 256);
18467 }
18468 for host in [ChunkHost::Macos, ChunkHost::Aarch64, ChunkHost::Other] {
18470 for dense in [false, true] {
18471 assert_eq!(prefill_chunk_rule(Some(48), host, dense), 48);
18472 assert_eq!(prefill_chunk_rule(Some(0), host, dense), 1);
18473 }
18474 }
18475 }
18476
18477 #[test]
18478 fn kv_reuse_plan_pulls_rows_decode_wrote_only_on_the_device() {
18479 use super::{ReuseLayer, ReusePlan, kv_reuse_plan};
18480 let full = |host_rows, device_rows| ReuseLayer {
18481 full: true,
18482 host_rows,
18483 device_rows,
18484 device_state: false,
18485 };
18486 assert_eq!(
18489 kv_reuse_plan(339, &[full(300, Some(339)), full(300, Some(339))]),
18490 ReusePlan::Pull(vec![(0, 300, 339), (1, 300, 339)])
18491 );
18492 assert_eq!(kv_reuse_plan(339, &[full(339, None)]), ReusePlan::Ready);
18494 assert_eq!(kv_reuse_plan(339, &[full(339, Some(345))]), ReusePlan::Ready);
18496 assert_eq!(
18498 kv_reuse_plan(339, &[full(300, Some(339)), full(339, None)]),
18499 ReusePlan::Pull(vec![(0, 300, 339)])
18500 );
18501 assert_eq!(kv_reuse_plan(339, &[full(300, Some(320))]), ReusePlan::Fresh);
18503 assert_eq!(kv_reuse_plan(339, &[full(300, None)]), ReusePlan::Fresh);
18504 assert_eq!(kv_reuse_plan(339, &[full(350, None)]), ReusePlan::Fresh);
18505 let conv = |device_state| ReuseLayer {
18508 full: false,
18509 host_rows: 0,
18510 device_rows: None,
18511 device_state,
18512 };
18513 assert_eq!(
18514 kv_reuse_plan(339, &[conv(true), full(300, Some(339))]),
18515 ReusePlan::Fresh
18516 );
18517 assert_eq!(kv_reuse_plan(339, &[conv(false), full(339, None)]), ReusePlan::Ready);
18518 }
18519
18520 #[test]
18521 fn nll_graph_policy_scopes_only_the_fused_head() {
18522 for (label, unmasked, prefer_graph, native_metal, want_graph, want_head) in [
18523 ("vulkan graph", true, true, false, true, false),
18525 ("native Metal graph", true, true, true, true, true),
18527 ("masked", false, true, false, false, false),
18529 ("graph disabled", true, false, true, false, false),
18530 ] {
18531 let (graph_quality, graph_head_required) =
18532 super::nll_graph_policy(unmasked, prefer_graph, native_metal);
18533 assert_eq!(graph_quality, want_graph, "{label}: graph quality");
18534 assert_eq!(graph_head_required, want_head, "{label}: fused head");
18535 }
18536 }
18537
18538 #[test]
18539 fn mtp_prefill_pair_boundaries_skip_only_final_prompt_row() {
18540 assert_eq!(mtp_prefill_pair_count(0, 128, 256), 128);
18541 assert_eq!(mtp_prefill_pair_count(128, 256, 256), 127);
18542 assert_eq!(mtp_prefill_pair_count(0, 256, 256), 255);
18543 assert_eq!(mtp_prefill_pair_count(256, 256, 256), 0);
18544 assert_eq!(mtp_prefill_pair_count(300, 320, 256), 0);
18545 }
18546
18547 #[test]
18548 fn cancel_flag_stops_generation() {
18549 let mut p = create_test_pipeline(16, 32, 2, 2, 8, 2, 32);
18550 p.cancel.store(true, std::sync::atomic::Ordering::Relaxed);
18553 let r = p.generate_from_ids(&[1, 2, 3], 8, None, None).unwrap();
18554 assert_eq!(r.finish_reason, "cancelled");
18555 assert!(
18556 r.token_ids.is_empty(),
18557 "no tokens after cancel: {:?}",
18558 r.token_ids
18559 );
18560 assert_eq!(p.kv_cache.seq_len(), 0);
18561 assert!(p.kv_history.is_empty());
18562 assert!(!p.graph_want_logits);
18563 assert!(p.graph_logits.is_none());
18564 let r2 = p.generate_from_ids(&[1, 2, 3], 4, None, None).unwrap();
18566 assert_ne!(r2.finish_reason, "cancelled");
18567 }
18568 use super::*;
18569
18570 #[test]
18578 fn dynamic_ffn_equals_the_zeroing_arm() {
18579 let (hidden, inter) = (8usize, 32usize);
18580 let synth = |n: usize, salt: usize| -> Vec<f32> {
18581 (0..n)
18582 .map(|i| (((i * 29 + salt * 13 + 7) % 89) as f32 / 89.0 - 0.5) * 0.6)
18583 .collect()
18584 };
18585 let down = synth(hidden * inter, 3);
18586 let mut down_t = vec![0.0f32; inter * hidden];
18587 for r in 0..hidden {
18588 for c in 0..inter {
18589 down_t[c * hidden + r] = down[r * inter + c];
18590 }
18591 }
18592 let d = DenseFfn {
18593 gate_proj: QTensor::from_f32(synth(inter * hidden, 1), inter, hidden),
18594 up_proj: QTensor::from_f32(synth(inter * hidden, 2), inter, hidden),
18595 down_proj: QTensor::from_f32(down.clone(), hidden, inter),
18596 act: Act::Silu,
18597 down_t: Some(QTensor::from_f32(down_t, inter, hidden)),
18598 segs: Vec::new(),
18599 };
18600 let x = synth(hidden, 11);
18601 let k = 12usize;
18602 let got = dense_ffn_dynamic(&d, &x, None, k).expect("down_t present");
18603 let mut g = vec![0.0f32; inter];
18605 d.gate_proj.matvec(&x, &mut g, None);
18606 let mut u = vec![0.0f32; inter];
18607 d.up_proj.matvec(&x, &mut u, None);
18608 for v in g.iter_mut() {
18609 *v = inference::silu(*v);
18610 }
18611 keep_top_k(&mut g, k);
18612 for i in 0..inter {
18613 g[i] *= u[i];
18614 }
18615 let mut want = vec![0.0f32; hidden];
18616 d.down_proj.matvec(&g, &mut want, None);
18617 for (a, b) in want.iter().zip(&got) {
18618 assert!((a - b).abs() < 1e-5, "dynamic {b} vs reference {a}");
18619 }
18620 }
18621
18622 #[test]
18628 fn tube_ffn_open_equals_dense_and_closed_equals_masked() {
18629 let (hidden, core, tube) = (8usize, 12usize, 8usize);
18630 let inter = core + tube;
18631 let synth = |n: usize, salt: usize| -> Vec<f32> {
18632 (0..n)
18633 .map(|i| (((i * 41 + salt * 17 + 5) % 97) as f32 / 97.0 - 0.5) * 0.5)
18634 .collect()
18635 };
18636 let (g_all, u_all) = (synth(inter * hidden, 1), synth(inter * hidden, 2));
18637 let d_all = synth(hidden * inter, 3);
18638 let dense = DenseFfn {
18640 gate_proj: QTensor::from_f32(g_all.clone(), inter, hidden),
18641 up_proj: QTensor::from_f32(u_all.clone(), inter, hidden),
18642 down_proj: QTensor::from_f32(d_all.clone(), hidden, inter),
18643 act: Act::Silu,
18644 down_t: None,
18645 segs: Vec::new(),
18646 };
18647 let rows =
18648 |v: &[f32], a: usize, b: usize| -> Vec<f32> { v[a * hidden..b * hidden].to_vec() };
18649 let cols = |v: &[f32], a: usize, b: usize| -> Vec<f32> {
18650 let mut o = Vec::with_capacity(hidden * (b - a));
18651 for r in 0..hidden {
18652 o.extend_from_slice(&v[r * inter + a..r * inter + b]);
18653 }
18654 o
18655 };
18656 let tubed = DenseFfn {
18657 down_t: None,
18658 gate_proj: QTensor::from_f32(rows(&g_all, 0, core), core, hidden),
18659 up_proj: QTensor::from_f32(rows(&u_all, 0, core), core, hidden),
18660 down_proj: QTensor::from_f32(cols(&d_all, 0, core), hidden, core),
18661 act: Act::Silu,
18662 segs: vec![FfnSeg {
18663 gate: QTensor::from_f32(rows(&g_all, core, inter), tube, hidden),
18664 up: QTensor::from_f32(rows(&u_all, core, inter), tube, hidden),
18665 down: QTensor::from_f32(cols(&d_all, core, inter), hidden, tube),
18666 start: core,
18667 width: tube,
18668 }],
18669 };
18670 let x = synth(hidden, 7);
18671 let want = dense_ffn(&dense, &x, None);
18672 let got = tube_ffn(&tubed, &x, 1, None, None);
18673 for (a, b) in want.iter().zip(&got) {
18674 assert!((a - b).abs() < 1e-5, "open tube: {a} vs {b}");
18675 }
18676 let mut bits = vec![0u8; inter.div_ceil(8)];
18678 for n in 0..core {
18679 bits[n / 8] |= 1 << (n % 8);
18680 }
18681 let closed = tube_ffn(&tubed, &x, 1, None, Some(&bits));
18682 let masked = dense_ffn_masked(&dense, &x, None, &bits);
18683 for (a, b) in masked.iter().zip(&closed) {
18684 assert!((a - b).abs() < 1e-5, "closed tube: {a} vs {b}");
18685 }
18686 let batch = tube_ffn(&tubed, &x, 1, None, Some(&bits));
18688 for (a, b) in closed.iter().zip(&batch) {
18689 assert_eq!(a, b, "batch arm disagrees with decode arm");
18690 }
18691 }
18692
18693 #[test]
18695 fn sparse_ffn_quant_equals_dense_with_inactive_zeroed() {
18696 let (hidden, inter) = (16usize, 40usize);
18697 let synth = |n: usize, salt: usize| -> Vec<f32> {
18698 (0..n)
18699 .map(|i| (((i * 37 + salt * 11 + 3) % 101) as f32 / 101.0 - 0.5) * 0.4)
18700 .collect()
18701 };
18702 let d = DenseFfn {
18703 gate_proj: QTensor::from_f32(synth(inter * hidden, 1), inter, hidden),
18704 up_proj: QTensor::from_f32(synth(inter * hidden, 2), inter, hidden),
18705 down_proj: QTensor::from_f32(synth(hidden * inter, 3), hidden, inter),
18706 act: Act::Silu,
18707 down_t: None,
18708 segs: Vec::new(),
18709 };
18710 let x = synth(hidden, 9);
18711 let active: Vec<u16> = (0..inter as u16).filter(|i| i % 3 == 0).collect();
18713
18714 let sparse = sparse_ffn_quant(&d, &x, &active, hidden, None);
18715
18716 let mut g = vec![0.0f32; inter];
18718 d.gate_proj.matvec(&x, &mut g, None);
18719 let mut u = vec![0.0f32; inter];
18720 d.up_proj.matvec(&x, &mut u, None);
18721 let act_set: std::collections::HashSet<u16> = active.iter().copied().collect();
18722 for i in 0..inter {
18723 g[i] = if act_set.contains(&(i as u16)) {
18724 inference::silu(g[i]) * u[i]
18725 } else {
18726 0.0
18727 };
18728 }
18729 let mut reference = vec![0.0f32; hidden];
18730 d.down_proj.matvec(&g, &mut reference, None);
18731
18732 let max_d = sparse
18733 .iter()
18734 .zip(&reference)
18735 .map(|(a, b)| (a - b).abs())
18736 .fold(0.0f32, f32::max);
18737 assert!(max_d < 1e-5, "sparse != dense-zeroed: max|Δ| = {max_d}");
18738 }
18739
18740 fn attach_test_mtp(p: &mut Pipeline) {
18742 let (h, inter, heads, kv, hd) = (
18743 p.hidden_size,
18744 p.intermediate_size,
18745 p.num_heads,
18746 p.num_kv_heads,
18747 p.head_dim,
18748 );
18749 let synth = |n: usize, salt: usize| -> Vec<f32> {
18750 (0..n)
18751 .map(|i| (((i * 29 + salt * 23 + 5) % 101) as f32 / 101.0 - 0.5) * 0.2)
18752 .collect()
18753 };
18754 let qt = |rows: usize, cols: usize, salt: usize| -> QTensor {
18755 QTensor::from_f32(synth(rows * cols, salt), rows, cols)
18756 };
18757 p.mtp = Some(MtpModule {
18758 enorm: vec![1.0; h],
18759 hnorm: vec![1.0; h],
18760 eh_proj: qt(h, 2 * h, 301),
18761 layer: LayerWeights {
18762 input_norm: vec![1.0; h],
18763 post_norm: vec![1.0; h],
18764 attn_out_norm: None,
18765 ffn_out_norm: None,
18766 layer_scale: None,
18767 ffn: FfnKind::Dense(DenseFfn {
18768 gate_proj: qt(inter, h, 315),
18769 up_proj: qt(inter, h, 316),
18770 down_proj: qt(h, inter, 317),
18771 act: Act::Silu,
18772 down_t: None,
18773 segs: Vec::new(),
18774 }),
18775 attn: AttnKind::Full {
18776 bias: None,
18777 wq: qt(heads * hd, h, 311),
18778 wk: qt(kv * hd, h, 312),
18779 wv: qt(kv * hd, h, 313),
18780 wo: qt(h, heads * hd, 314),
18781 q_norm: None,
18782 k_norm: None,
18783 output_gate: false,
18784 softplus_gate: None,
18785 },
18786 },
18787 final_norm: vec![1.0; h],
18788 kv: crate::kv_cache::LayerKvCache::new(kv, hd),
18789 });
18790 }
18791
18792 #[test]
18793 fn speculative_equals_vanilla_greedy() {
18794 unsafe { std::env::set_var("CMF_GPU_WGPU_GRAPH", "0") };
18798 let run = |spec: bool| {
18799 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
18800 p.sampler_config.temperature = 0.0;
18801 attach_test_mtp(&mut p);
18802 p.speculative = spec;
18803 let r = p.generate("abcdef", 12, None, None).unwrap();
18804 (r.token_ids, r.mtp_drafted, r.mtp_accepted)
18805 };
18806 let (vanilla, d0, _) = run(false);
18807 let (spec, d1, a1) = run(true);
18808 assert_eq!(d0, 0, "vanilla path must not draft");
18809 assert!(d1 > 0, "speculative path must draft");
18810 assert_eq!(
18811 vanilla, spec,
18812 "speculative must reproduce the exact greedy sequence (accepted {a1}/{d1})"
18813 );
18814 }
18815
18816 #[test]
18817 fn speculative_accepts_constant_oracle() {
18818 unsafe { std::env::set_var("CMF_GPU_WGPU_GRAPH", "0") };
18820 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
18821 p.sampler_config.temperature = 0.0;
18822 p.sampler_config.repetition_penalty = 1.0;
18823 p.weights.lm_head = QTensor::from_f32(vec![0.01; 64 * 8], 64, 8);
18826 attach_test_mtp(&mut p);
18827 p.speculative = true;
18828 let r = p.generate("abcd", 10, None, None).unwrap();
18829 assert!(r.mtp_drafted > 0);
18830 assert_eq!(
18831 r.mtp_accepted, r.mtp_drafted,
18832 "constant logits → every draft accepted"
18833 );
18834 assert!(r.token_ids.windows(2).all(|w| w[0] == w[1]));
18837 }
18838
18839 #[test]
18840 fn empty_prompt_is_an_error_not_a_panic() {
18841 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 260);
18842 let r = p.generate("", 4, None, None);
18843 assert!(r.is_err(), "empty prompt must be a clean error");
18844 }
18845
18846 #[test]
18847 fn every_token_enters_kv_exactly_once() {
18848 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
18849 p.sampler_config.temperature = 0.0;
18851 let r = p.generate("abc", 2, None, None).unwrap();
18852 assert_eq!(r.prompt_tokens, 3);
18853 assert_eq!(
18857 p.kv_cache.seq_len(),
18858 3 + r.tokens_generated - 1,
18859 "each token must be cached exactly once (v1 cached the last prompt token twice)"
18860 );
18861 }
18862
18863 #[test]
18864 fn generation_is_reproducible_with_seed() {
18865 let run = || {
18866 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
18867 p.generate("hello", 8, None, None).unwrap().token_ids
18868 };
18869 assert_eq!(run(), run());
18870 }
18871
18872 #[test]
18873 fn resetting_sampler_restarts_the_seeded_stream() {
18874 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
18875 let config = SamplerConfig {
18876 seed: Some(1234),
18877 ..SamplerConfig::default()
18878 };
18879 p.set_sampler_config(config.clone());
18880 let first = p.generate("hello", 8, None, None).unwrap().token_ids;
18881 p.set_sampler_config(config);
18882 let second = p.generate("hello", 8, None, None).unwrap().token_ids;
18883 assert_eq!(first, second);
18884 }
18885
18886 #[test]
18887 fn eviction_bounds_the_cache() {
18888 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 260);
18889 p.kv_cache.max_seq_len = 6;
18890 p.sampler_config.temperature = 0.0;
18891 let _ = p.generate("abcd", 12, None, None).unwrap();
18892 assert!(
18893 p.kv_cache.seq_len() <= 6 + 1,
18894 "cache must stay bounded by max_seq_len (got {})",
18895 p.kv_cache.seq_len()
18896 );
18897 }
18898
18899 #[test]
18900 fn confidence_matches_tokens_and_is_a_probability() {
18901 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
18902 p.sampler_config.temperature = 0.0;
18903 p.sampler_config.repetition_penalty = 1.0;
18904 let r = p.generate("abcd", 10, None, None).unwrap();
18905 assert_eq!(
18906 r.token_confidence.len(),
18907 r.token_ids.len(),
18908 "one confidence per emitted token"
18909 );
18910 for &c in &r.token_confidence {
18911 assert!((0.0..=1.0).contains(&c), "confidence out of [0,1]: {c}");
18912 }
18913 let logits = [1.0f32, 3.0, 0.5, 3.0];
18915 let p0 = top1_prob_t(&logits, 1, 1.0);
18916 let p1 = top1_prob_t(&logits, 3, 1.0);
18917 assert!((p0 - p1).abs() < 1e-6, "equal logits → equal prob");
18918 assert!(p0 > 0.0 && p0 < 1.0);
18919 let sharp = top1_prob_t(&logits, 1, 1.0);
18921 let soft = top1_prob_t(&logits, 1, 2.0);
18922 assert!(soft < sharp, "higher temperature lowers peak confidence");
18923 }
18924
18925 #[test]
18926 fn trace_is_opt_in_and_parallels_the_output() {
18927 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
18929 p.sampler_config.temperature = 0.0;
18930 p.sampler_config.repetition_penalty = 1.0;
18931 let r = p.generate("abcd", 10, None, None).unwrap();
18932 assert!(r.traces.is_empty(), "trace must be empty unless enabled");
18933
18934 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
18936 p.sampler_config.temperature = 0.0;
18937 p.sampler_config.repetition_penalty = 1.0;
18938 p.set_trace(true);
18939 let r = p.generate("abcd", 10, None, None).unwrap();
18940 assert_eq!(r.traces.len(), r.token_ids.len(), "one trace row per token");
18941 for (i, tr) in r.traces.iter().enumerate() {
18942 assert_eq!(tr.t, i, "trace index is sequential");
18943 assert_eq!(tr.token_id, r.token_ids[i], "trace token_id matches output");
18944 assert_eq!(
18945 tr.confidence, r.token_confidence[i],
18946 "trace confidence matches the confidence channel"
18947 );
18948 assert!(tr.active_skill.is_none() && tr.recon.is_none() && !tr.switched);
18950 }
18951 }
18952
18953 #[test]
18954 fn explain_prefill_logits_match_greedy_first_token() {
18955 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
18959 p.sampler_config.temperature = 0.0;
18960 p.sampler_config.repetition_penalty = 1.0;
18961 let ids = p.tokenizer.encode("abcd");
18962 let logits = p.prefill_next_logits(&ids, None);
18963 let argmax = logits
18964 .iter()
18965 .enumerate()
18966 .max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
18967 .unwrap()
18968 .0 as u32;
18969 let r = p.generate("abcd", 1, None, None).unwrap();
18970 assert_eq!(
18971 argmax, r.token_ids[0],
18972 "explain preview must match greedy emit"
18973 );
18974 }
18975
18976 #[test]
18977 fn laguna_shared_expert_is_unconditionally_added() {
18978 let matrix = |values: Vec<f32>| QTensor::from_f32(values, 2, 2);
18979 let identity = || matrix(vec![1.0, 0.0, 0.0, 1.0]);
18980 let zero_dense = || DenseFfn {
18981 gate_proj: matrix(vec![0.0; 4]),
18982 up_proj: matrix(vec![0.0; 4]),
18983 down_proj: matrix(vec![0.0; 4]),
18984 act: Act::Silu,
18985 down_t: None,
18986 segs: Vec::new(),
18987 };
18988 let shared = DenseFfn {
18989 gate_proj: identity(),
18990 up_proj: identity(),
18991 down_proj: identity(),
18992 act: Act::Silu,
18993 down_t: None,
18994 segs: Vec::new(),
18995 };
18996 let x = [1.0, 2.0];
18997 let expected = dense_ffn(&shared, &x, None);
18998 let moe = MoeFfn {
18999 router: QTensor::from_f32(vec![0.0, 0.0], 1, 2),
19000 experts: vec![zero_dense()],
19001 top_k: 1,
19002 norm_topk_prob: true,
19003 router_sigmoid: true,
19004 expert_bias: None,
19005 routed_scaling: 1.0,
19006 route_tau: None,
19007 shared: Some((shared, None)),
19008 stats: std::cell::RefCell::new(Vec::new()),
19009 act_sq: std::cell::RefCell::new(Vec::new()),
19010 act_rows: std::cell::RefCell::new(Vec::new()),
19011 mask: None,
19012 per_expert_scale: None,
19013 router_input_norm: false,
19014 resonance: None,
19015 grown: Vec::new(),
19016 };
19017 let actual = moe_ffn_cpu(&moe, &x, &[0], &[0.0], 1.0, None);
19018 for (actual, expected) in actual.iter().zip(expected) {
19019 assert!((actual - expected).abs() < 1e-6);
19020 }
19021 }
19022
19023 fn mimo_test_pipeline() -> Pipeline {
19032 let (hs, inter, nh, hd, vd, vocab) = (16usize, 24usize, 4usize, 8usize, 4usize, 64usize);
19033 let kvh = [1usize, 2, 2, 1];
19034 let synth = |n: usize, salt: usize| -> Vec<f32> {
19035 (0..n)
19036 .map(|i| (((i * 31 + salt * 17 + 7) % 97) as f32 / 97.0 - 0.5) * 0.4)
19037 .collect()
19038 };
19039 let qt = |rows: usize, cols: usize, salt: usize| {
19040 QTensor::from_f32(synth(rows * cols, salt), rows, cols)
19041 };
19042 let dense = |inter: usize, salt: usize| DenseFfn {
19043 gate_proj: qt(inter, hs, salt),
19044 up_proj: qt(inter, hs, salt + 1),
19045 down_proj: qt(hs, inter, salt + 2),
19046 act: Act::Silu,
19047 down_t: None,
19048 segs: Vec::new(),
19049 };
19050 let layers: Vec<LayerWeights> = (0..4)
19051 .map(|li| LayerWeights {
19052 input_norm: vec![1.0; hs],
19053 post_norm: vec![1.0; hs],
19054 attn_out_norm: None,
19055 ffn_out_norm: None,
19056 layer_scale: None,
19057 ffn: if li == 0 {
19058 FfnKind::Dense(dense(inter, 50))
19059 } else {
19060 FfnKind::Moe(MoeFfn {
19061 router: qt(4, hs, 60 + li),
19062 experts: (0..4).map(|e| dense(8, 70 + li * 10 + e * 3)).collect(),
19063 top_k: 2,
19064 norm_topk_prob: true,
19065 router_sigmoid: true,
19066 expert_bias: Some(vec![0.02, -0.03, 0.01, 0.0]),
19067 routed_scaling: 1.0,
19068 route_tau: None,
19069 shared: None,
19070 stats: std::cell::RefCell::new(Vec::new()),
19071 act_sq: std::cell::RefCell::new(Vec::new()),
19072 act_rows: std::cell::RefCell::new(Vec::new()),
19073 mask: None,
19074 per_expert_scale: None,
19075 router_input_norm: false,
19076 resonance: None,
19077 grown: Vec::new(),
19078 })
19079 },
19080 attn: AttnKind::Full {
19081 wq: qt(nh * hd, hs, li * 10 + 1),
19082 wk: qt(kvh[li] * hd, hs, li * 10 + 2),
19083 wv: qt(kvh[li] * vd, hs, li * 10 + 3),
19084 wo: qt(hs, nh * vd, li * 10 + 4),
19085 q_norm: None,
19086 k_norm: None,
19087 output_gate: false,
19088 softplus_gate: None,
19089 bias: None,
19090 },
19091 })
19092 .collect();
19093 let mut p = Pipeline::new(
19094 Tokenizer::byte_level(),
19095 PipelineWeights {
19096 embed_tokens: qt(vocab, hs, 100),
19097 layers,
19098 lm_head: qt(vocab, hs, 200),
19099 final_norm: vec![1.0; hs],
19100 },
19101 hs,
19102 inter,
19103 nh,
19104 1, hd,
19106 4,
19107 4,
19108 false,
19109 vocab,
19110 1e-6,
19111 1e7,
19112 NormStyle::Qwen,
19113 4096,
19114 SamplerConfig {
19115 seed: Some(7),
19116 ..Default::default()
19117 },
19118 );
19119 p.layer_dump = None;
19121 p.set_rotary(4, 1e7);
19122 p.sliding_layers = Some(vec![false, true, true, false]);
19123 p.swa = Some((3, usize::MAX));
19124 p.rotary_dim_local = Some(4);
19125 p.inv_freq_local = Some(std::sync::Arc::new(attention::rope_inv_freq(4, 1e4)));
19126 p.set_attn_geometry(Some(kvh.to_vec()), Some(vd)).unwrap();
19127 p.set_layer_sinks(1, vec![0.5, -1.0, 1.5, 0.0]).unwrap();
19128 p.set_layer_sinks(2, vec![-0.25, 2.0, 0.75, -1.5]).unwrap();
19129 p
19130 }
19131
19132 #[test]
19133 fn mimo_embedded_prompt_uses_rows_and_never_reuses_token_only_kv() {
19134 let mut p = mimo_test_pipeline();
19135 p.speculative = false;
19136 p.ignore_eos = true;
19137 p.sampler_config.temperature = 0.0;
19138 p.sampler_config.repetition_penalty = 1.0;
19139 let a = vec![3, 5, 7, 9, 11, 13];
19140 let b = vec![4, 8, 12, 16, 20, 24];
19141 let rows: Vec<_> = b.iter().flat_map(|&id| p.embed_id(id)).collect();
19142 let expected = p.generate_from_ids(&b, 8, None, None).unwrap().token_ids;
19143 let actual = p.generate_from_embeds(&a, &rows, 8, None, None).unwrap().token_ids;
19147 assert_eq!(actual, expected);
19148 assert!(p.kv_history.is_empty());
19149 let mut extended = a.clone();
19150 extended.push(17);
19151 let after_media = p.generate_from_ids(&extended, 8, None, None).unwrap().token_ids;
19152 p.reset_session();
19153 let fresh = p.generate_from_ids(&extended, 8, None, None).unwrap().token_ids;
19154 assert_eq!(after_media, fresh);
19155 assert!(p.generate_from_embeds(&a, &rows[..rows.len()-1], 1, None, None).is_err());
19156 p.reset_session();
19160 p.generate_from_ids(&a, 1, None, None).unwrap();
19161 let mut media_ids = p.kv_history.clone();
19162 assert!(!media_ids.is_empty());
19163 media_ids.extend_from_slice(&[19, 21, 23]);
19164 let source_ids: Vec<_> = (0..media_ids.len()).map(|i| b[i % b.len()]).collect();
19165 let source_rows: Vec<_> = source_ids.iter().flat_map(|&id| p.embed_id(id)).collect();
19166 let mut oracle = mimo_test_pipeline();
19167 oracle.speculative = false;
19168 oracle.ignore_eos = true;
19169 oracle.sampler_config.temperature = 0.0;
19170 oracle.sampler_config.repetition_penalty = 1.0;
19171 let expected = oracle.generate_from_ids(&source_ids, 8, None, None).unwrap().token_ids;
19172 assert_eq!(p.generate_from_embeds(&media_ids, &source_rows, 8, None, None).unwrap().token_ids, expected);
19173 assert!(p.kv_history.is_empty());
19174 let mut bad = rows;
19175 bad[0] = f32::NAN;
19176 assert!(p.generate_from_embeds(&a, &bad, 1, None, None).is_err());
19177 }
19178
19179 fn f32_bits(v: &[f32]) -> Vec<u32> {
19180 v.iter().map(|x| x.to_bits()).collect()
19181 }
19182
19183 #[test]
19190 fn mimo_shaped_decode_matches_prefill_batch_bitwise() {
19191 let mut p = mimo_test_pipeline();
19192 let kv: Vec<usize> = p.kv_cache.layers.iter().map(|l| l.num_kv_heads).collect();
19193 assert_eq!(kv, vec![1, 2, 2, 1]);
19194 assert!(p.kv_cache.layers[0].sinks.is_none() && p.kv_cache.layers[3].sinks.is_none());
19195 assert!(p.kv_cache.layers[1].sinks.is_some() && p.kv_cache.layers[2].sinks.is_some());
19196 let ids: Vec<u32> = (0..12u32).map(|i| (i * 7 + 3) % 64).collect();
19197 let hs = p.hidden_size;
19198 let mut decode = Vec::new();
19199 for (pos, &id) in ids.iter().enumerate() {
19200 let e = p.embed_single(id);
19201 let h = p.forward_layers(&e, pos, None);
19202 decode.push(p.logits_from_hidden(&h));
19203 }
19204 for l in &p.kv_cache.layers {
19205 assert_eq!(l.seq_len, 12);
19206 assert_eq!(l.head_values(0).len(), 12 * 8);
19208 }
19209 assert!(decode.iter().flatten().all(|v| v.is_finite()));
19210
19211 p.clear_sequence_state();
19212 let hb = p.prefill_batch_span(PrefillIn::Ids(&ids), 0, None, 0, p.num_layers);
19213 for pos in 0..ids.len() {
19214 let lg = p.logits_from_hidden(&hb[pos * hs..(pos + 1) * hs]);
19215 assert_eq!(
19216 f32_bits(&decode[pos]),
19217 f32_bits(&lg),
19218 "whole prompt, pos {pos}"
19219 );
19220 }
19221
19222 p.clear_sequence_state();
19223 let a = p.prefill_batch_span(PrefillIn::Ids(&ids[..5]), 0, None, 0, p.num_layers);
19224 let b = p.prefill_batch_span(PrefillIn::Ids(&ids[5..]), 5, None, 0, p.num_layers);
19225 for pos in 0..ids.len() {
19226 let row = if pos < 5 {
19227 &a[pos * hs..(pos + 1) * hs]
19228 } else {
19229 &b[(pos - 5) * hs..(pos - 4) * hs]
19230 };
19231 let lg = p.logits_from_hidden(row);
19232 assert_eq!(
19233 f32_bits(&decode[pos]),
19234 f32_bits(&lg),
19235 "two chunks, pos {pos}"
19236 );
19237 }
19238
19239 let last = |p: &mut Pipeline| {
19242 p.clear_sequence_state();
19243 let hb = p.prefill_batch_span(PrefillIn::Ids(&ids), 0, None, 0, p.num_layers);
19244 p.logits_from_hidden(&hb[11 * hs..12 * hs])
19245 };
19246 let base = last(&mut p);
19247 let mut no_sinks = mimo_test_pipeline();
19248 for l in &mut no_sinks.kv_cache.layers {
19249 l.sinks = None;
19250 }
19251 assert_ne!(
19252 f32_bits(&last(&mut no_sinks)),
19253 f32_bits(&base),
19254 "sinks are live"
19255 );
19256 let mut wide = mimo_test_pipeline();
19257 wide.swa = Some((64, usize::MAX));
19258 assert_ne!(
19259 f32_bits(&last(&mut wide)),
19260 f32_bits(&base),
19261 "window is live"
19262 );
19263
19264 p.clear_sequence_state();
19266 p.ignore_eos = true;
19267 let r = p.generate_from_ids(&ids, 4, None, None).unwrap();
19268 assert_eq!(r.token_ids.len(), 4);
19269 }
19270
19271 fn mimo_test_mtp(n: usize, gain: f32) -> mimo_mtp::MimoMtp {
19274 let (hs, inter, nh, hd, vd, nkv) = (16usize, 24usize, 4usize, 8usize, 4usize, 2usize);
19275 let synth = |len: usize, salt: usize| -> Vec<f32> {
19276 (0..len)
19277 .map(|i| (((i * 37 + salt * 13 + 3) % 89) as f32 / 89.0 - 0.5) * 0.6 * gain)
19278 .collect()
19279 };
19280 let qt = |rows: usize, cols: usize, salt: usize| {
19281 QTensor::from_f32(synth(rows * cols, salt), rows, cols)
19282 };
19283 let layers = (0..n)
19284 .map(|k| {
19285 let s = 500 + k * 40;
19286 let mut kv = crate::kv_cache::LayerKvCache::new(nkv, hd);
19287 kv.sinks = Some(vec![0.3, -0.7, 1.1, 0.0]);
19288 MtpModule {
19289 enorm: vec![1.0; hs],
19290 hnorm: vec![1.0; hs],
19291 eh_proj: qt(hs, 2 * hs, s),
19292 layer: LayerWeights {
19293 input_norm: vec![1.0; hs],
19294 post_norm: vec![1.0; hs],
19295 attn_out_norm: None,
19296 ffn_out_norm: None,
19297 layer_scale: None,
19298 attn: AttnKind::Full {
19299 wq: qt(nh * hd, hs, s + 1),
19300 wk: qt(nkv * hd, hs, s + 2),
19301 wv: qt(nkv * vd, hs, s + 3),
19302 wo: qt(hs, nh * vd, s + 4),
19303 q_norm: None,
19304 k_norm: None,
19305 output_gate: false,
19306 softplus_gate: None,
19307 bias: None,
19308 },
19309 ffn: FfnKind::Dense(DenseFfn {
19310 gate_proj: qt(inter, hs, s + 5),
19311 up_proj: qt(inter, hs, s + 6),
19312 down_proj: qt(hs, inter, s + 7),
19313 act: Act::Silu,
19314 down_t: None,
19315 segs: Vec::new(),
19316 }),
19317 },
19318 final_norm: vec![1.0; hs],
19319 kv,
19320 }
19321 })
19322 .collect();
19323 mimo_mtp::MimoMtp::from_layers(layers)
19324 }
19325
19326 fn mimo_greedy(p: &mut Pipeline, ids: &[u32], n: usize, spec: bool) -> GenerateResult {
19327 p.clear_sequence_state();
19328 p.speculative = spec;
19329 p.ignore_eos = true;
19330 p.sampler_config.temperature = 0.0;
19331 p.generate_from_ids(ids, n, None, None).unwrap()
19332 }
19333
19334 #[test]
19340 fn mimo_mtp_incremental_rounds_equal_the_teacher_forced_table() {
19341 for post in [false, true] {
19344 let mut p = mimo_test_pipeline();
19345 p.weights.final_norm = (0..p.hidden_size).map(|i| 0.5 + 0.1 * i as f32).collect();
19347 let mut st0 = mimo_test_mtp(3, 1.0);
19348 st0.post_norm_hidden = post;
19349 p.mimo_mtp = Some(st0);
19350 let ids: Vec<u32> = (0..14u32).map(|i| (i * 11 + 5) % 64).collect();
19351 let hs = p.hidden_size;
19352 let hb = p.prefill_batch_span(PrefillIn::Ids(&ids), 0, None, 0, p.num_layers);
19353 p.mimo_note_rows(&hb, 0);
19354 let mut st = p.mimo_mtp.take().unwrap();
19355 let k = 3;
19358 let mut inc = Vec::new();
19359 for t in 0..ids.len() - k - 1 {
19360 inc.push(p.mimo_mtp_draft(&mut st, t, &ids, k));
19361 }
19362 let s = ids.len();
19365 let mut reference = vec![vec![0u32; k]; s - k - 1];
19366 let mut fresh = mimo_test_mtp(3, 1.0);
19367 for (layer, m) in fresh.layers.iter_mut().enumerate() {
19368 let n = s - layer - 1;
19369 let mut cats = vec![0.0f32; n * 2 * hs];
19370 for j in 0..n {
19371 let e = p.embed_single(ids[j + layer + 1]);
19372 let raw = &hb[j * hs..(j + 1) * hs];
19373 let g = if post {
19374 inference::rms_norm(raw, &p.weights.final_norm, p.rms_eps, p.norm_style)
19375 } else {
19376 raw.to_vec()
19377 };
19378 let (ce, ch) = cats[j * 2 * hs..(j + 1) * 2 * hs].split_at_mut(hs);
19379 inference::rms_norm_into(&e, &m.enorm, p.rms_eps, p.norm_style, ce);
19380 inference::rms_norm_into(&g, &m.hnorm, p.rms_eps, p.norm_style, ch);
19381 }
19382 let mut x = vec![0.0f32; n * hs];
19383 m.eh_proj.matmat(&cats, n, &mut x, None);
19384 p.mimo_mtp_block(m, &mut x, n, 0);
19385 for (t, row) in reference.iter_mut().enumerate() {
19386 let y = inference::rms_norm(
19387 &x[t * hs..(t + 1) * hs],
19388 &m.final_norm,
19389 p.rms_eps,
19390 p.norm_style,
19391 );
19392 row[layer] = sampler::argmax(&p.lm_head_forward(&y));
19393 }
19394 }
19395 assert_eq!(inc, reference, "post_norm_hidden = {post}");
19396 let distinct: std::collections::HashSet<u32> =
19398 inc.iter().flatten().copied().collect();
19399 assert!(distinct.len() > 3, "{inc:?}");
19400 let last_t = ids.len() - k - 2;
19402 for m in &st.layers {
19403 assert_eq!(m.kv.seq_len, last_t + 1);
19404 }
19405 }
19406 }
19407
19408 #[test]
19415 fn mimo_speculative_greedy_equals_plain_greedy() {
19416 unsafe { std::env::set_var("CMF_GPU_WGPU_GRAPH", "0") };
19417 let ids: Vec<u32> = (0..9u32).map(|i| (i * 7 + 3) % 64).collect();
19418 let n = 24;
19419 let mut p = mimo_test_pipeline();
19420 let plain = mimo_greedy(&mut p, &ids, n, false);
19421 assert_eq!(plain.mtp_drafted, 0);
19422 assert_eq!(plain.token_ids.len(), n);
19423 let plain_kv = p.kv_cache.layers[0].seq_len;
19424
19425 p.mimo_mtp = Some(mimo_test_mtp(3, 1.0));
19427 let spec = mimo_greedy(&mut p, &ids, n, true);
19428 assert!(spec.mtp_drafted > 0, "the round must draft");
19429 assert_eq!(spec.token_ids, plain.token_ids);
19430 assert_eq!(p.kv_cache.layers[0].seq_len, plain_kv);
19431
19432 let mut truth: Vec<u32> = ids.clone();
19435 truth.extend(&plain.token_ids);
19436 let mut noisy = truth.clone();
19437 for (i, t) in noisy.iter_mut().enumerate() {
19438 if i % 5 == 0 {
19439 *t = (*t + 1) % 64;
19440 }
19441 }
19442 let mut st = mimo_test_mtp(3, 1.0);
19443 st.draft_override = Some(noisy);
19444 p.mimo_mtp = Some(st);
19445 let spec = mimo_greedy(&mut p, &ids, n, true);
19446 assert_eq!(spec.token_ids, plain.token_ids);
19447 assert_eq!(p.kv_cache.layers[0].seq_len, plain_kv);
19448 let stats = p.mimo_mtp.as_ref().unwrap().stats.clone();
19449 assert_eq!(stats.accepted as usize, spec.mtp_accepted);
19450 assert!(spec.mtp_accepted > 0 && spec.mtp_accepted < spec.mtp_drafted);
19451 assert!(
19452 stats.accept_hist.iter().filter(|&&c| c > 0).count() >= 3,
19453 "{:?}",
19454 stats.accept_hist
19455 );
19456 assert!(stats.tokens_per_round() > 1.5, "{}", stats.line());
19457
19458 let mut st = mimo_test_mtp(3, 1.0);
19461 st.draft_override = Some(truth);
19462 p.mimo_mtp = Some(st);
19463 let spec = mimo_greedy(&mut p, &ids, n, true);
19464 assert_eq!(spec.token_ids, plain.token_ids);
19465 assert_eq!(spec.mtp_accepted, spec.mtp_drafted);
19466 assert_eq!(p.kv_cache.layers[0].seq_len, plain_kv);
19467
19468 let off = mimo_greedy(&mut p, &ids, n, false);
19470 assert_eq!(off.token_ids, plain.token_ids);
19471 assert_eq!(off.mtp_drafted, 0);
19472 }
19473
19474 #[test]
19483 fn mimo_shaped_model_rides_the_wgpu_graph_geometry() {
19484 let p = mimo_test_pipeline();
19485 assert_eq!(
19486 p.graph_attn_decline_reason(),
19487 Some("per-layer KV head counts")
19488 );
19489 assert_eq!(p.wgpu_graph_attn_decline(), None);
19490 let g0 = p.graph_attn_geom(0).expect("full layer geometry");
19491 assert_eq!(
19492 (g0.nkv, g0.dv, g0.rd, g0.window, g0.sink.is_some()),
19493 (1, 4, 4, None, false)
19494 );
19495 assert_eq!(g0.invf, p.inv_freq.as_slice());
19496 let g1 = p.graph_attn_geom(1).expect("sliding layer geometry");
19497 assert_eq!((g1.nkv, g1.dv, g1.rd, g1.window), (2, 4, 4, Some(3)));
19498 assert_eq!(g1.sink, Some(&[0.5f32, -1.0, 1.5, 0.0][..]));
19499 assert_eq!(g1.invf, p.inv_freq_local.as_ref().unwrap().as_slice());
19500 assert_ne!(g0.invf, g1.invf, "two RoPE tables");
19501 let g3 = p.graph_attn_geom(3).expect("full layer geometry");
19502 assert_eq!((g3.nkv, g3.window, g3.sink.is_some()), (1, None, false));
19503
19504 let emb = p.embed_single(3);
19507 let mut lg = Vec::new();
19508 assert!(
19509 p.try_token_graph_wgpu_steps(&emb, 0, &mut lg, 1, None, None, 0, p.num_layers)
19510 .is_none()
19511 );
19512 let mut hid = emb.clone();
19513 assert_eq!(
19514 p.try_batch_graph_wgpu(&mut hid, &[0], 1, None),
19515 crate::gpu::BatchGraphOutcome::Declined
19516 );
19517 assert_eq!(hid, emb, "a declined batch graph leaves the rows untouched");
19518 assert!(p.try_multi_burst(3, 0, 4).is_none());
19519 assert!(
19520 p.graph_declines().is_empty(),
19521 "no attention decline logged: {:?}",
19522 p.graph_declines()
19523 );
19524 let plain = || create_test_pipeline(8, 16, 2, 1, 4, 2, 32);
19529 assert_eq!(plain().graph_attn_decline_reason(), None);
19530 assert_eq!(plain().wgpu_graph_attn_decline(), None);
19531 assert!(
19532 plain().graph_attn_geom(0).is_none(),
19533 "uniform models keep the historical arms"
19534 );
19535 let mut q = plain();
19536 q.set_layer_sinks(1, vec![0.25, -0.25]).unwrap();
19537 assert_eq!(
19538 q.graph_attn_decline_reason(),
19539 Some("learned attention sinks")
19540 );
19541 assert_eq!(
19542 q.graph_attn_geom(1).unwrap().sink,
19543 Some(&[0.25f32, -0.25][..])
19544 );
19545 let mut q = plain();
19546 q.set_attn_geometry(None, Some(2)).unwrap();
19547 assert_eq!(
19548 q.graph_attn_decline_reason(),
19549 Some("V heads narrower than Q/K heads")
19550 );
19551 assert_eq!(q.graph_attn_geom(0).unwrap().dv, 2);
19552 let mut q = plain();
19553 q.sliding_layers = Some(vec![true, false]);
19554 q.swa = Some((4, usize::MAX));
19555 assert_eq!(q.graph_attn_decline_reason(), Some("sliding-window layers"));
19556 assert_eq!(q.graph_attn_geom(0).unwrap().window, Some(4));
19557 assert_eq!(q.graph_attn_geom(1).unwrap().window, None);
19558
19559 let mut q = mimo_test_pipeline();
19562 q.rope_scale = 2.0;
19563 assert_eq!(
19564 q.wgpu_graph_attn_decline(),
19565 Some("scaled RoPE positions with per-layer geometry")
19566 );
19567 let emb = q.embed_single(3);
19568 assert!(
19569 q.try_token_graph_wgpu_steps(&emb, 0, &mut lg, 1, None, None, 0, q.num_layers)
19570 .is_none()
19571 );
19572 let _ = q.try_token_graph_wgpu_steps(&emb, 1, &mut lg, 1, None, None, 0, q.num_layers);
19573 let lines = q.graph_declines();
19574 assert_eq!(
19575 lines
19576 .iter()
19577 .filter(|l| l.starts_with("wgpu token graph") && l.contains("scaled RoPE"))
19578 .count(),
19579 1,
19580 "{lines:?}"
19581 );
19582 }
19583
19584 #[test]
19585 fn mimo_verify_rewind_preserves_lagging_host_caches() {
19586 let mut p = mimo_test_pipeline();
19587 for (li, layer) in p.kv_cache.layers.iter_mut().enumerate() {
19588 let row = vec![0.0; layer.num_kv_heads * layer.head_dim];
19589 for _ in 0..if li == 0 { 2 } else { 12 } {
19590 layer.append(&row, &row, &[]);
19591 }
19592 }
19593 p.mimo_verify_rewind(9).unwrap();
19594 assert_eq!(p.kv_cache.layers[0].seq_len, 2);
19595 for layer in &p.kv_cache.layers[1..] {
19596 assert_eq!(layer.seq_len, 9);
19597 }
19598 }
19599
19600 #[test]
19604 fn layer_dump_covers_every_position_and_layer_on_both_walks() {
19605 let dir = std::env::temp_dir().join(format!("cmf-layer-dump-{}", std::process::id()));
19606 let _ = std::fs::remove_dir_all(&dir);
19607 let mut p = mimo_test_pipeline();
19608 let hs = p.hidden_size;
19609 let ids = [5u32, 9, 11, 2, 40];
19610 p.layer_dump = Some(dir.join("decode"));
19611 for (pos, &id) in ids.iter().enumerate() {
19612 let e = p.embed_single(id);
19613 let _ = p.forward_layers(&e, pos, None);
19614 }
19615 p.clear_sequence_state();
19616 p.layer_dump = Some(dir.join("prefill"));
19617 let hb = p.prefill_batch_span(PrefillIn::Ids(&ids), 0, None, 0, p.num_layers);
19618 for pos in 0..ids.len() {
19619 for li in 0..p.num_layers {
19620 let name = format!("p{pos:06}_l{li:02}.f32");
19621 let a = std::fs::read(dir.join("decode").join(&name)).unwrap();
19622 let b = std::fs::read(dir.join("prefill").join(&name)).unwrap();
19623 assert_eq!(a.len(), hs * 4, "{name}");
19624 assert_eq!(a, b, "{name}");
19625 }
19626 }
19627 let last = std::fs::read(dir.join("prefill").join("p000004_l03.f32")).unwrap();
19628 let vals: Vec<f32> = last
19629 .chunks(4)
19630 .map(|c| f32::from_le_bytes(c.try_into().unwrap()))
19631 .collect();
19632 assert_eq!(f32_bits(&vals), f32_bits(&hb[4 * hs..5 * hs]));
19633 let _ = std::fs::remove_dir_all(&dir);
19634 }
19635
19636 #[test]
19637 fn attn_geometry_and_sinks_are_validated() {
19638 let mut p = create_test_pipeline(8, 16, 4, 2, 4, 2, 32);
19639 assert!(
19640 p.set_attn_geometry(Some(vec![2]), None).is_err(),
19641 "one entry per layer"
19642 );
19643 assert!(
19644 p.set_attn_geometry(Some(vec![2, 3]), None).is_err(),
19645 "3 does not divide 4"
19646 );
19647 assert!(p.set_attn_geometry(Some(vec![2, 0]), None).is_err());
19648 assert!(p.set_attn_geometry(None, Some(0)).is_err());
19649 assert!(
19650 p.set_attn_geometry(None, Some(5)).is_err(),
19651 "V wider than the head"
19652 );
19653 p.set_attn_geometry(None, Some(4)).unwrap();
19654 assert_eq!(
19655 p.v_head_dim, None,
19656 "v_head_dim == head_dim is the uniform case"
19657 );
19658 p.set_layer_sinks(1, vec![0.1; 4]).unwrap();
19659 p.set_attn_geometry(Some(vec![1, 4]), None).unwrap();
19660 assert_eq!(p.kv_cache.layers[0].num_kv_heads, 1);
19661 assert_eq!(p.kv_cache.layers[1].num_kv_heads, 4);
19662 assert!(
19663 p.kv_cache.layers[1].sinks.is_some(),
19664 "a reshape keeps the layer's sinks"
19665 );
19666 assert_eq!(p.layer_geom(1).0, 4);
19667 assert!(
19668 p.set_layer_sinks(0, vec![0.0; 3]).is_err(),
19669 "one sink per Q head"
19670 );
19671 assert!(p.set_layer_sinks(7, vec![0.0; 4]).is_err());
19672 assert!(p.set_layer_sinks(0, vec![f32::NAN, 0.0, 0.0, 0.0]).is_err());
19673 }
19674
19675 #[test]
19678 fn o1_is_never_armed_on_sink_window_or_narrow_v_layers() {
19679 let cfg = || {
19680 Some(crate::nystrom::O1Cfg {
19681 layers: crate::nystrom::O1Layers::All,
19682 m: 4,
19683 w: 8,
19684 sink: 2,
19685 rect: crate::nystrom::O1Rect::Aggregate,
19686 })
19687 };
19688 let mut p = mimo_test_pipeline();
19689 p.set_o1(cfg());
19690 assert!(!p.o1_active(), "every MiMo-shaped layer is ineligible");
19691 let mut q = create_test_pipeline(8, 16, 2, 1, 4, 3, 64);
19692 q.set_layer_sinks(1, vec![0.0, 0.0]).unwrap();
19693 q.sliding_layers = Some(vec![false, false, true]);
19694 q.swa = Some((4, usize::MAX));
19695 q.set_o1(cfg());
19696 assert_eq!(q.o1_flags, vec![true, false, false]);
19697 }
19698
19699 #[test]
19700 fn o1_batch_transition_publishes_one_epoch_before_serial_handoff() {
19701 const B: usize = 19;
19702 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
19703 p.set_o1(Some(crate::nystrom::O1Cfg {
19704 layers: crate::nystrom::O1Layers::All,
19705 m: 4,
19706 w: 8,
19707 sink: 2,
19708 rect: crate::nystrom::O1Rect::Aggregate,
19709 }));
19710 p.o1_begin_with_prefix(Some(B));
19711 let ids: Vec<u32> = (0..B as u32).collect();
19712 let _ = p.prefill_batch_span(PrefillIn::Ids(&ids), 0, None, 0, p.num_layers);
19713
19714 assert_eq!(p.o1_epoch, 1, "all layers publish one completed transition");
19715 assert!(p.kv_cache.layers.iter().all(|l| l.o1_sealed()));
19716 let next = p.embed_single(B as u32);
19717 let _ = p.forward_layers(&next, B, None);
19718 assert_eq!(p.o1_epoch, 1, "sealed handoff must not republish the epoch");
19719 }
19720
19721 #[test]
19722 fn o1_pair_transition_commits_scratch_before_epoch_publication() {
19723 const B: usize = 19;
19724 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
19725 let gdn_cfg = crate::linear_core::GdnCfg {
19729 num_v_heads: 2,
19730 num_k_heads: 1,
19731 key_head_dim: 2,
19732 value_head_dim: 4,
19733 conv_kernel: 3,
19734 hidden_size: 8,
19735 rms_eps: 1e-6,
19736 output_gate_sigmoid: false,
19737 };
19738 let synth = |n: usize, salt: usize| -> Vec<f32> {
19739 (0..n)
19740 .map(|i| (((i * 13 + salt * 7) % 97) as f32 / 97.0 - 0.5) * 0.4)
19741 .collect()
19742 };
19743 let qt = |rows: usize, cols: usize, salt: usize| {
19744 crate::qtensor::QTensor::from_f32(synth(rows * cols, salt), rows, cols)
19745 };
19746 let c_dim = gdn_cfg.conv_dim();
19747 let vd = gdn_cfg.num_v_heads * gdn_cfg.value_head_dim;
19748 p.weights.layers[0].attn = AttnKind::LinearGdn(crate::linear_core::GdnWeights {
19749 in_proj_qkv: qt(c_dim, 8, 1),
19750 in_proj_z: qt(vd, 8, 2),
19751 in_proj_a: qt(gdn_cfg.num_v_heads, 8, 3),
19752 in_proj_b: qt(gdn_cfg.num_v_heads, 8, 4),
19753 conv1d: synth(c_dim * gdn_cfg.conv_kernel, 5),
19754 a_log: vec![0.2, 0.5],
19755 dt_bias: synth(gdn_cfg.num_v_heads, 6),
19756 norm: vec![1.0; gdn_cfg.value_head_dim],
19757 out_proj: qt(8, vd, 7),
19758 });
19759 p.gdn_cfg = Some(gdn_cfg);
19760 p.set_o1(Some(crate::nystrom::O1Cfg {
19761 layers: crate::nystrom::O1Layers::All,
19762 m: 4,
19763 w: 8,
19764 sink: 2,
19765 rect: crate::nystrom::O1Rect::Aggregate,
19766 }));
19767 p.o1_begin_with_prefix(Some(B));
19768 for pos in 0..B - 2 {
19769 let emb = p.embed_single(pos as u32);
19770 let _ = p.forward_layers(&emb, pos, None);
19771 }
19772 let lane1_state = p.kv_cache.layers[0].linear_state.clone();
19773
19774 let e1 = p.embed_single((B - 2) as u32);
19775 let e2 = p.embed_single((B - 1) as u32);
19776 let _ = p.forward_pair(&e1, &e2, B - 2);
19777
19778 assert_eq!(p.o1_epoch, 1, "pair crossing B publishes one epoch");
19779 assert!(
19780 p.kv_cache
19781 .layers
19782 .iter()
19783 .enumerate()
19784 .all(|(li, l)| !p.o1_flags[li] || l.o1_sealed())
19785 );
19786 assert!(!p.kv_cache.layers[0].linear_state.is_empty());
19787 assert_ne!(
19788 p.kv_cache.layers[0].linear_state, lane1_state,
19789 "real pair must commit GDN lane 2 before returning"
19790 );
19791 assert!(p.kv_cache.layers[0].linear_scratch.is_empty());
19792 let next = p.embed_single(B as u32);
19793 let _ = p.forward_layers(&next, B, None);
19794 assert_eq!(p.o1_epoch, 1, "serial continuation must reuse the epoch");
19795 }
19796
19797 #[test]
19798 fn o1_error_observation_stays_terminal_until_reset() {
19799 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
19800 p.set_o1(Some(crate::nystrom::O1Cfg {
19801 layers: crate::nystrom::O1Layers::All,
19802 m: 4,
19803 w: 8,
19804 sink: 2,
19805 rect: crate::nystrom::O1Rect::Aggregate,
19806 }));
19807 p.o1_begin();
19808 p.kv_cache.layers[0].o1_abort("synthetic transition failure".into());
19809
19810 assert!(p.o1_seal_checked().is_err());
19811 assert!(
19812 p.o1_seal_checked().is_err(),
19813 "retry must see the sticky error"
19814 );
19815 let k = vec![0.2f32; 4];
19816 let v = vec![0.3f32; 4];
19817 p.kv_cache.layers[0].append(&k, &v, &[]);
19818 assert_eq!(p.kv_cache.layers[0].seq_len, 0);
19819
19820 p.reset_session();
19821 p.o1_begin();
19822 p.kv_cache.layers[0].append(&k, &v, &[]);
19823 assert_eq!(p.kv_cache.layers[0].seq_len, 1);
19824 }
19825
19826 #[test]
19827 fn nll_graph_failure_is_terminal_and_request_is_reusable() {
19828 let ids = vec![1u32, 2, 3, 4, 5, 6];
19829 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
19830 p.graph_logits = Some(vec![123.0]);
19831 p.graph_want_logits = true;
19832 p.graph_failed
19833 .store(true, std::sync::atomic::Ordering::Relaxed);
19834 p.cancel.store(true, std::sync::atomic::Ordering::Relaxed);
19835 let err = p.nll_ids_from(&ids, 0).expect_err("prior graph failure");
19836 assert!(err.contains("before NLL"));
19837 assert!(p.graph_logits.is_none());
19838 assert!(!p.graph_want_logits);
19839 assert!(!p.graph_failed.load(std::sync::atomic::Ordering::Relaxed));
19840 assert!(!p.cancel.load(std::sync::atomic::Ordering::Relaxed));
19841
19842 let mut fresh = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
19843 let expected = fresh.nll_ids_from(&ids, 0).expect("fresh NLL");
19844 let actual = p.nll_ids_from(&ids, 0).expect("reused NLL");
19845 assert_eq!(actual.1, expected.1);
19846 assert!((actual.0 - expected.0).abs() < 1e-9);
19847 }
19848
19849 #[test]
19850 fn nll_forward_failure_discards_partial_score_and_clears_sidechannels() {
19851 let ids = vec![1u32, 2, 3, 4, 5, 6];
19852 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
19853 p.nll_test_fail_at = Some(1);
19854 let err = p
19855 .nll_ids_from(&ids, 0)
19856 .expect_err("one-shot forward failure");
19857 assert!(err.contains("forward") || err.contains("score row"));
19858 assert!(!p.graph_failed.load(std::sync::atomic::Ordering::Relaxed));
19859 assert!(!p.graph_want_logits);
19860 assert!(p.graph_logits.is_none());
19861 assert!(p.kv_history.is_empty());
19862
19863 let mut fresh = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
19864 let expected = fresh.nll_ids_from(&ids, 0).expect("fresh NLL");
19865 let actual = p.nll_ids_from(&ids, 0).expect("reused NLL");
19866 assert_eq!(actual.1, expected.1);
19867 assert!((actual.0 - expected.0).abs() < 1e-9);
19868 }
19869
19870 #[test]
19871 fn nll_serial_failure_before_first_row_is_reported() {
19872 let ids = vec![1u32, 2, 3, 4];
19873 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
19874 p.nll_test_force_serial = true;
19875 p.nll_test_fail_at = Some(0);
19876 let err = p.nll_ids_from(&ids, 0).expect_err("serial forward failure");
19877 assert!(err.contains("serial forward"));
19878 assert!(p.kv_history.is_empty());
19879 assert!(!p.graph_failed.load(std::sync::atomic::Ordering::Relaxed));
19880 assert!(!p.cancel.load(std::sync::atomic::Ordering::Relaxed));
19881 }
19882
19883 #[test]
19884 fn ffn_probe_failure_discards_recorder_and_state() {
19885 let ids = vec![1u32, 2, 3, 4];
19886 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
19887 p.nll_test_fail_at = Some(0);
19888 let err = p
19889 .probe_ffn_mass_batch(&ids)
19890 .expect_err("probe forward failure");
19891 assert!(err.contains("NLL"));
19892 assert!(FFN_PROBE.with(|probe| probe.borrow().is_none()));
19893 assert!(p.kv_history.is_empty());
19894 assert!(!p.graph_failed.load(std::sync::atomic::Ordering::Relaxed));
19895 }
19896
19897 #[test]
19898 fn nll_test_controls_are_pipeline_scoped() {
19899 let ids = vec![1u32, 2, 3, 4];
19900 let mut failing = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
19901 let mut unaffected = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
19902 failing.nll_test_force_serial = true;
19903 failing.nll_test_fail_at = Some(0);
19904
19905 assert!(!failing.can_prefill_batched());
19906 assert!(unaffected.can_prefill_batched());
19907 let expected = unaffected
19908 .nll_ids_from(&ids, 0)
19909 .expect("unaffected pipeline remains usable");
19910 let err = failing
19911 .nll_ids_from(&ids, 0)
19912 .expect_err("failure injection belongs to failing pipeline");
19913 assert!(err.contains("serial forward"));
19914 assert!(failing.nll_test_fail_at.is_none());
19915 assert!(unaffected.can_prefill_batched());
19916 let actual = unaffected
19917 .nll_ids_from(&ids, 0)
19918 .expect("unaffected pipeline remains reusable");
19919 assert_eq!(actual.1, expected.1);
19920 assert!((actual.0 - expected.0).abs() < 1e-9);
19921 }
19922
19923 #[test]
19924 fn forward_ids_failure_channel_is_terminal_and_reusable() {
19925 let ids = vec![1u32, 2, 3, 4, 5, 6];
19926 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
19927 p.graph_logits = Some(vec![123.0]);
19928 p.graph_want_logits = true;
19929 p.graph_failed
19930 .store(true, std::sync::atomic::Ordering::Relaxed);
19931 p.cancel.store(true, std::sync::atomic::Ordering::Relaxed);
19932
19933 let err = p
19934 .forward_ids(&ids, None)
19935 .expect_err("a failed forward must not become a valid head result");
19936 assert!(err.contains("forward_ids setup"));
19937 assert!(p.graph_logits.is_none());
19938 assert!(!p.graph_want_logits);
19939 assert!(!p.graph_failed.load(std::sync::atomic::Ordering::Relaxed));
19940 assert!(!p.cancel.load(std::sync::atomic::Ordering::Relaxed));
19941 assert_eq!(p.kv_cache.seq_len(), 0);
19942
19943 let expected = create_test_pipeline(8, 16, 2, 1, 4, 1, 64)
19944 .forward_ids(&ids, None)
19945 .expect("fresh forward_ids");
19946 let actual = p
19947 .forward_ids(&ids, None)
19948 .expect("pipeline remains reusable after a failed forward");
19949 assert_eq!(actual.len(), expected.len());
19950 assert!(
19951 actual
19952 .iter()
19953 .zip(expected)
19954 .all(|(a, b)| (a - b).abs() < 1e-9)
19955 );
19956 assert_eq!(p.kv_cache.seq_len(), ids.len());
19957 }
19958
19959 #[test]
19960 fn sigmoid_router_floor_is_explicit_per_architecture() {
19961 let zero = || DenseFfn {
19967 gate_proj: QTensor::from_f32(vec![0.0; 4], 2, 2),
19968 up_proj: QTensor::from_f32(vec![0.0; 4], 2, 2),
19969 down_proj: QTensor::from_f32(vec![0.0; 4], 2, 2),
19970 act: Act::Silu,
19971 down_t: None,
19972 segs: Vec::new(),
19973 };
19974 let m = MoeFfn {
19975 router: QTensor::from_f32(vec![0.0; 4], 2, 2),
19976 experts: vec![zero(), zero()],
19977 top_k: 1,
19978 norm_topk_prob: true,
19979 router_sigmoid: true,
19980 expert_bias: None,
19981 routed_scaling: 2.5,
19982 route_tau: None,
19983 shared: None,
19984 stats: std::cell::RefCell::new(Vec::new()),
19985 act_sq: std::cell::RefCell::new(Vec::new()),
19986 act_rows: std::cell::RefCell::new(Vec::new()),
19987 mask: None,
19988 per_expert_scale: None,
19989 router_input_norm: false,
19990 resonance: None,
19991 grown: Vec::new(),
19992 };
19993 let logits = [-20.0f32, -20.0];
19994 let (_, p, glm_wsum) = moe_route_with_eps(&logits, &m, None, 1e-20);
19995 let (_, _, generic_wsum) = moe_route(&logits, &m, None);
19996 let expected = (p[0] + 1e-20) / m.routed_scaling;
19997 assert!((glm_wsum - expected).abs() < 1e-15);
19998 assert!(generic_wsum > glm_wsum * 100.0);
19999 }
20000
20001 #[test]
20002 fn resonance_scores_match_formula_and_stable_tie() {
20003 let r = Resonance {
20004 mu: vec![0.0, 0.0, 1.0, 0.0, 0.0, 1.0],
20006 u: vec![1.0, 0.0, 0.0, 1.0, 1.0, 1.0],
20007 k: 1,
20008 bias: vec![1.5, 0.5, 0.0],
20009 shell: Vec::new(),
20010 };
20011 let x = [1.0f32, 1.0];
20012 let mut got = vec![0.0; 3];
20013 r.scores(&x, &mut got);
20014 assert!((got[0] - 0.5).abs() < 1e-6);
20018 assert!((got[1] - 0.5).abs() < 1e-6);
20019 assert!(got[2].abs() < 1e-6);
20020 let best = got
20021 .iter()
20022 .enumerate()
20023 .max_by(|(ia, a), (ib, b)| a.partial_cmp(b).unwrap().then(ib.cmp(ia)))
20024 .map(|(i, _)| i);
20025 assert_eq!(best, Some(0));
20026 assert!(got.iter().all(|v| v.is_finite()));
20027 }
20028
20029 #[test]
20034 fn resonance_shell_masks_outside_keeps_inside_and_trunk_bits() {
20035 let plain = Resonance {
20041 mu: vec![0.0, 0.0, 1.0, 0.0, 0.0, 1.0, 0.0, 1.0],
20042 u: vec![1.0, 0.0, 0.0, 1.0, 1.0, 1.0, 1.0, 1.0],
20043 k: 1,
20044 bias: vec![1.5, 0.5, 0.0, 0.0],
20045 shell: Vec::new(),
20046 };
20047 let shelled = Resonance {
20048 mu: plain.mu.clone(),
20049 u: plain.u.clone(),
20050 k: 1,
20051 bias: plain.bias.clone(),
20052 shell: vec![f32::INFINITY, f32::INFINITY, 6.0, 0.25],
20053 };
20054 assert!(!plain.has_shell());
20055 assert!(shelled.has_shell());
20056 let bits = |v: &[f32]| v.iter().map(|x| x.to_bits()).collect::<Vec<_>>();
20057 let (mut a, mut b) = (vec![0.0; 4], vec![0.0; 4]);
20058 set_growth_shell(Some(true));
20059 assert!(growth_shell_enabled());
20060 let x = [1.0f32, 1.0];
20063 plain.scores(&x, &mut a);
20064 shelled.scores(&x, &mut b);
20065 assert_eq!(bits(&a), bits(&b), "inside every shell: unchanged");
20066 assert!(a[2] == 0.0 && a[3] == 0.0);
20067 let xo = [3.0f32, 0.0];
20070 plain.scores(&xo, &mut a);
20071 shelled.scores(&xo, &mut b);
20072 assert_eq!(a[2], -6.0);
20073 assert_eq!(a[3], -6.0);
20074 assert_eq!(bits(&a[..3]), bits(&b[..3]), "trunk rows + the boundary row");
20075 assert_eq!(b[3], f32::NEG_INFINITY, "outside the shell: −∞");
20076 assert_eq!(shelled.effective_shell(4), shelled.shell);
20077 set_growth_shell(Some(false));
20080 assert!(!growth_shell_enabled());
20081 assert_eq!(shelled.effective_shell(4), vec![f32::INFINITY; 4]);
20082 shelled.scores(&xo, &mut b);
20083 assert_eq!(bits(&a), bits(&b));
20084 set_growth_shell(None);
20085 let short = Resonance {
20088 shell: vec![f32::INFINITY, f32::INFINITY],
20089 ..shelled
20090 };
20091 set_growth_shell(Some(true));
20092 short.scores(&xo, &mut b);
20093 assert_eq!(bits(&a), bits(&b));
20094 set_growth_shell(None);
20095 }
20096
20097 #[test]
20101 fn moe_route_neg_inf_logits_pick_the_best_finite_expert_with_weight_one() {
20102 let zero = || DenseFfn {
20103 gate_proj: QTensor::from_f32(vec![0.0; 4], 2, 2),
20104 up_proj: QTensor::from_f32(vec![0.0; 4], 2, 2),
20105 down_proj: QTensor::from_f32(vec![0.0; 4], 2, 2),
20106 act: Act::Silu,
20107 down_t: None,
20108 segs: Vec::new(),
20109 };
20110 let moe = |sigmoid: bool| MoeFfn {
20111 router: QTensor::from_f32(vec![0.0; 8], 4, 2),
20112 experts: vec![zero(), zero(), zero(), zero()],
20113 top_k: 1,
20114 norm_topk_prob: true,
20115 router_sigmoid: sigmoid,
20116 expert_bias: None,
20117 routed_scaling: 1.0,
20118 route_tau: None,
20119 shared: None,
20120 stats: std::cell::RefCell::new(Vec::new()),
20121 act_sq: std::cell::RefCell::new(Vec::new()),
20122 act_rows: std::cell::RefCell::new(Vec::new()),
20123 mask: None,
20124 per_expert_scale: None,
20125 router_input_norm: false,
20126 resonance: None,
20127 grown: Vec::new(),
20128 };
20129 let logits = [-1.0f32, f32::NEG_INFINITY, -0.5, f32::NEG_INFINITY];
20130 for sigmoid in [false, true] {
20131 let m = moe(sigmoid);
20132 let (idx, p, wsum) = moe_route(&logits, &m, None);
20133 assert_eq!(idx, vec![2], "sigmoid {sigmoid}");
20134 assert_eq!(p[1], 0.0);
20135 assert_eq!(p[3], 0.0);
20136 assert!(p[2] > p[0] && p[0] > 0.0);
20137 assert!(p.iter().all(|v| v.is_finite()));
20138 let w = p[2] / wsum;
20139 if sigmoid {
20140 assert!((w - 1.0).abs() < 1e-5, "sigmoid: weight {w}");
20142 } else {
20143 assert_eq!(w, 1.0, "softmax: the renormalized top-1 weight is exactly 1");
20144 }
20145 let (idx, _, _) = moe_route(&logits, &m, Some(&[true, true, true, true]));
20148 assert_eq!(idx, vec![2]);
20149 }
20150 let m = moe(false);
20152 let (idx, _, _) = moe_route(&[f32::NEG_INFINITY, f32::NEG_INFINITY, f32::NEG_INFINITY, -9.0], &m, None);
20153 assert_eq!(idx, vec![3]);
20154 let (idx, p, wsum) = moe_route(&[f32::NEG_INFINITY; 4], &m, None);
20157 assert_eq!(idx, vec![0]);
20158 assert!(p.iter().all(|&v| v == 0.25));
20159 assert!(wsum.is_finite() && wsum > 0.0);
20160 }
20161
20162 #[test]
20167 fn resonance_route_picks_the_raw_argmax_not_the_softmax_tie() {
20168 let zero = || DenseFfn {
20169 gate_proj: QTensor::from_f32(vec![0.0; 4], 2, 2),
20170 up_proj: QTensor::from_f32(vec![0.0; 4], 2, 2),
20171 down_proj: QTensor::from_f32(vec![0.0; 4], 2, 2),
20172 act: Act::Silu,
20173 down_t: None,
20174 segs: Vec::new(),
20175 };
20176 let moe = |resonant: bool, norm_topk: bool| MoeFfn {
20177 router: QTensor::from_f32(vec![0.0; 6], 3, 2),
20178 experts: vec![zero(), zero(), zero()],
20179 top_k: 1,
20180 norm_topk_prob: norm_topk,
20181 router_sigmoid: false,
20182 expert_bias: None,
20183 routed_scaling: 1.0,
20184 route_tau: None,
20185 shared: None,
20186 stats: std::cell::RefCell::new(Vec::new()),
20187 act_sq: std::cell::RefCell::new(Vec::new()),
20188 act_rows: std::cell::RefCell::new(Vec::new()),
20189 mask: None,
20190 per_expert_scale: None,
20191 router_input_norm: false,
20192 resonance: resonant.then(|| Resonance {
20193 mu: vec![0.0; 6],
20194 u: Vec::new(),
20195 k: 0,
20196 bias: vec![0.0; 3],
20197 shell: Vec::new(),
20198 }),
20199 grown: Vec::new(),
20200 };
20201 let lo = -0.1f32;
20204 let hi = f32::from_bits(lo.to_bits() - 1);
20205 assert!(hi > lo && hi - lo < 2f32.powi(-25));
20206 assert_eq!((lo - hi).exp(), 1.0, "the tie this test is about");
20207 let (idx, _, _) = moe_route(&[lo, hi], &moe(false, true), None);
20209 assert_eq!(idx, vec![0], "softmax tie → lower index (the gated contract)");
20210 for norm in [true, false] {
20213 let m = moe(true, norm);
20214 let (idx, p, wsum) = moe_route(&[lo, hi, f32::NEG_INFINITY], &m, None);
20215 assert_eq!(idx, vec![1], "norm_topk {norm}");
20216 assert_eq!(p, vec![0.0, 1.0, 0.0]);
20217 assert_eq!(p[1] / wsum, 1.0);
20218 let (idx, _, _) = moe_route(&[hi, hi, lo], &m, None);
20221 assert_eq!(idx, vec![0]);
20222 let (idx, _, _) = moe_route(&[f32::NEG_INFINITY, lo, hi], &m, None);
20224 assert_eq!(idx, vec![2]);
20225 let (idx, p, wsum) = moe_route(&[lo, hi, hi], &m, Some(&[true, false, false]));
20226 assert_eq!(idx, vec![0]);
20227 assert_eq!(p[0] / wsum, 1.0);
20228 let (idx, p, wsum) = moe_route(&[f32::NEG_INFINITY; 3], &m, None);
20231 assert_eq!(idx, vec![0]);
20232 assert!(p.iter().all(|v| v.is_finite()) && wsum.is_finite() && wsum > 0.0);
20233 }
20234 }
20235}