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 token_id = input_ids[pos];
4147 let want_logits = pos + 1 == input_ids.len();
4148 let mut lg = Vec::new();
4149 if let Some(b) = &mut self.qwen4_exp {
4150 crate::qwen4_exp::forward_token(
4151 &b.0,
4152 &b.1,
4153 &b.2,
4154 &mut b.3,
4155 token_id,
4156 pos,
4157 &self.inv_freq,
4158 self.pool.as_deref(),
4159 &mut lg,
4160 want_logits,
4161 );
4162 }
4163 if want_logits {
4164 self.graph_logits = Some(lg);
4165 }
4166 pos += 1;
4167 hidden.fill(0.0);
4168 }
4169 while self.dsv4.is_some()
4170 && mtp.is_none()
4171 && pos < input_ids.len()
4172 && !self.cancel.load(std::sync::atomic::Ordering::Relaxed)
4173 {
4174 let end = (pos + prefill_chunk()).min(input_ids.len());
4175 let ids: Vec<u32> = input_ids[pos..end].to_vec();
4176 let mut lg = Vec::new();
4177 if let Some(b) = &mut self.dsv4 {
4178 let (g, layers, cfg, st) = (&b.0, &b.1, b.2, &mut b.3);
4179 crate::dsv4::forward_chunk(
4180 g,
4181 layers,
4182 &cfg,
4183 st,
4184 &ids,
4185 pos,
4186 &self.inv_freq,
4187 self.pool.as_deref(),
4188 &mut lg,
4189 end == input_ids.len(),
4190 );
4191 }
4192 if end == input_ids.len() {
4193 self.graph_logits = Some(lg);
4194 }
4195 pos = end;
4196 hidden = vec![0.0; self.hidden_size];
4197 }
4198 let dsv41_prefill = self.dsv41_prefill.take();
4199 while self.dsv41.is_some()
4200 && mtp.is_none()
4201 && pos < input_ids.len()
4202 && !self.cancel.load(std::sync::atomic::Ordering::Relaxed)
4203 {
4204 let end = (pos + prefill_chunk()).min(input_ids.len());
4205 let ids: Vec<u32> = input_ids[pos..end].to_vec();
4206 let mut lg = Vec::new();
4207 if let Some(b) = &mut self.dsv41 {
4208 let (g, layers, cfg, st) = (&b.0, &b.1, &b.2, &mut b.3);
4209 if let Some((embeddings, participates)) = dsv41_prefill.as_ref() {
4210 crate::dsv41::forward_chunk_masked_with_embeddings(
4211 g,
4212 layers,
4213 cfg,
4214 st,
4215 &ids,
4216 pos,
4217 &embeddings[pos..end],
4218 &participates[pos..end],
4219 self.pool.as_deref(),
4220 &mut lg,
4221 );
4222 } else {
4223 crate::dsv41::forward_chunk(
4224 g,
4225 layers,
4226 cfg,
4227 st,
4228 &ids,
4229 pos,
4230 self.pool.as_deref(),
4231 &mut lg,
4232 );
4233 }
4234 }
4235 if end == input_ids.len() {
4236 self.graph_logits = Some(lg);
4237 }
4238 pos = end;
4239 hidden = vec![0.0; self.hidden_size];
4240 }
4241 let dyn_prefill = router.is_some();
4246 let o1_prefill_limit = o1_prefill
4254 .and_then(|requested| self.o1_effective_boundary(requested))
4255 .map(|boundary| boundary.min(input_ids.len()));
4256 let mut o1_sealed = false;
4257 if let Some(limit) = o1_prefill_limit {
4258 if self.can_prefill_batched() && limit > 2 {
4261 let chunk = self.prefill_chunk();
4262 let hs = self.hidden_size;
4263 while pos < limit && !self.cancel.load(std::sync::atomic::Ordering::Relaxed) {
4264 let end = (pos + chunk).min(limit);
4265 let hb = self.prefill_batch(&input_ids[pos..end], pos);
4266 hidden.copy_from_slice(&hb[(end - pos - 1) * hs..]);
4267 pos = end;
4268 }
4269 } else {
4270 while pos < limit && !self.cancel.load(std::sync::atomic::Ordering::Relaxed) {
4271 hidden = self.forward_layers(&self.embed_single(input_ids[pos]), pos, None);
4272 pos += 1;
4273 }
4274 }
4275 if pos >= limit {
4276 o1_sealed = match self.o1_seal_checked() {
4277 Ok(sealed) => sealed,
4278 Err(err) => {
4279 self.finish_generation(&mut mtp, &mut router, true);
4280 return Err(err);
4281 }
4282 };
4283 tracing::info!(
4284 "o1 bounded prompt prefix: requested={} effective={} processed={} of {} token(s)",
4285 o1_prefill.unwrap_or(0),
4286 self.o1_effective_boundary(o1_prefill.unwrap_or(0))
4287 .unwrap_or(limit),
4288 limit,
4289 input_ids.len()
4290 );
4291 }
4292 }
4293 let graph_prefill = self.graph_prefill_preferred();
4299 #[cfg(target_os = "macos")]
4307 if task_mask.is_none()
4308 && !dyn_prefill
4309 && (crate::gpu::q1_force() || crate::gpu::q2tp_gpu_opt_in())
4310 && crate::gpu::enabled_here()
4311 && self.gdn_cfg.is_some()
4312 && self.g3n.is_none()
4313 && input_ids.len() > 8
4314 && std::env::var("CMF_MTP_CHAIN_PROBE").is_err()
4315 && std::env::var("CMF_METAL_PREFILL").as_deref() != Ok("0")
4316 {
4317 let chunk: usize = std::env::var("CMF_METAL_PREFILL_CHUNK")
4318 .ok()
4319 .and_then(|v| v.parse().ok())
4320 .filter(|&v| (16..=512).contains(&v))
4321 .unwrap_or(256);
4322 let hs = self.hidden_size;
4323 let _tp = std::time::Instant::now();
4324 while pos < input_ids.len() && !self.cancel.load(std::sync::atomic::Ordering::Relaxed) {
4325 let end = (pos + chunk).min(input_ids.len());
4326 let hb = match self.prefill_batch_metal(&input_ids[pos..end], pos) {
4327 MetalPrefillOutcome::Completed(hb) => hb,
4328 MetalPrefillOutcome::Declined => break,
4329 MetalPrefillOutcome::Failed => {
4330 self.finish_generation(&mut mtp, &mut router, true);
4331 return Err("ordinary Metal prefill failed after admission".into());
4332 }
4333 };
4334 if let Some(m) = &mut mtp {
4335 let n_pairs = if end < input_ids.len() {
4336 end - pos
4337 } else {
4338 end - pos - 1
4339 };
4340 if n_pairs > 0 {
4341 let pairs: Vec<(&[f32], u32)> = (0..n_pairs)
4342 .map(|j| (&hb[j * hs..(j + 1) * hs], input_ids[pos + j + 1]))
4343 .collect();
4344 if !self.mtp_warm_batch_metal(m, &pairs, pos) {
4345 for (j, (h, t)) in pairs.iter().enumerate() {
4346 let h = h.to_vec();
4347 let _ = self.mtp_step(m, &h, *t, pos + j);
4348 }
4349 }
4350 }
4351 }
4352 hidden.copy_from_slice(&hb[(end - pos - 1) * hs..]);
4353 pos = end;
4354 }
4355 if std::env::var("CMF_PREFILL_PROF").is_ok() {
4356 eprintln!(
4357 "metal-prefill: {} of {} tokens in {:.1} ms",
4358 pos,
4359 input_ids.len(),
4360 _tp.elapsed().as_secs_f64() * 1e3
4361 );
4362 }
4363 }
4364 self.mimo_moe_prepare();
4365 #[cfg(not(target_os = "macos"))]
4371 if task_mask.is_none()
4372 && !dyn_prefill
4373 && !graph_prefill
4374 && mtp.is_none()
4375 && o1_prefill.is_none()
4376 && !self.o1_active()
4377 && input_ids.len() > 2
4378 && self.batch_prefix_prefill()
4379 {
4380 let chunk = self.prefill_chunk().max(1);
4381 let hs = self.hidden_size;
4382 let t_bp = std::time::Instant::now();
4383 let pos0 = pos;
4384 while pos < input_ids.len() && !self.cancel.load(std::sync::atomic::Ordering::Relaxed) {
4385 let end = (pos + chunk).min(input_ids.len());
4386 let bk = end - pos;
4387 let mut hiddens = vec![0f32; bk * hs];
4388 for (j, &id) in input_ids[pos..end].iter().enumerate() {
4389 hiddens[j * hs..(j + 1) * hs].copy_from_slice(&self.embed_single(id));
4390 }
4391 let positions: Vec<usize> = (pos..end).collect();
4392 let mut run = 0usize;
4393 let outcome = self.try_batch_graph_wgpu_prefix(
4394 &mut hiddens,
4395 &positions,
4396 bk,
4397 None,
4398 Some(&mut run),
4399 );
4400 match outcome {
4401 crate::gpu::BatchGraphOutcome::Completed => {
4402 let hb = if run < self.num_layers {
4403 self.prefill_batch_span(
4404 PrefillIn::Hidden(&hiddens),
4405 pos,
4406 None,
4407 run,
4408 self.num_layers,
4409 )
4410 } else {
4411 hiddens
4412 };
4413 if mimo_spec {
4414 self.mimo_note_rows(&hb, pos);
4415 }
4416 hidden.copy_from_slice(&hb[(bk - 1) * hs..]);
4417 pos = end;
4418 }
4419 crate::gpu::BatchGraphOutcome::Failed => {
4420 self.finish_generation(&mut mtp, &mut router, true);
4421 return Err("batched prefix prefill failed after admission".into());
4422 }
4423 crate::gpu::BatchGraphOutcome::Declined => {
4424 #[cfg(feature = "gpu")]
4427 if pos > pos0 {
4428 self.pull_lagging_host_kv(0, self.num_layers, pos);
4429 }
4430 break;
4431 }
4432 }
4433 }
4434 if std::env::var("CMF_PREFILL_PROF").is_ok() {
4435 eprintln!(
4436 "batch-prefix prefill: {} of {} tokens in {:.1} ms",
4437 pos - pos0,
4438 input_ids.len(),
4439 t_bp.elapsed().as_secs_f64() * 1e3
4440 );
4441 }
4442 }
4443 if task_mask.is_none()
4444 && !dyn_prefill
4445 && !graph_prefill
4446 && self.can_prefill_batched()
4447 && self.g3n.is_none()
4448 && o1_prefill.is_none()
4449 && input_ids.len() > 2
4450 {
4451 let chunk = self.prefill_chunk();
4457 let hs = self.hidden_size;
4458 while pos < input_ids.len() && !self.cancel.load(std::sync::atomic::Ordering::Relaxed) {
4459 let end = (pos + chunk).min(input_ids.len());
4460 let hb = self.prefill_batch(&input_ids[pos..end], pos);
4461 if mimo_spec {
4462 self.mimo_note_rows(&hb, pos);
4463 }
4464 if let Some(m) = &mut mtp {
4465 let probe: usize = std::env::var("CMF_MTP_CHAIN_PROBE")
4466 .ok()
4467 .and_then(|v| v.parse().ok())
4468 .unwrap_or(0);
4469 for p in pos..end {
4470 if p + 1 < input_ids.len() {
4471 if probe >= 1 && p + 2 < input_ids.len() {
4472 let (d1, mut hx) = self.mtp_step_h(
4476 m,
4477 &hb[(p - pos) * hs..(p - pos + 1) * hs],
4478 input_ids[p + 1],
4479 p,
4480 );
4481 let mut ok = d1 == input_ids[p + 2];
4482 Self::chain_probe_note(0, ok);
4483 let mut d_prev = d1;
4484 let mut extra = 0usize;
4485 for j in 1..probe {
4486 if p + 2 + j >= input_ids.len() {
4487 break;
4488 }
4489 let (dj, hj) = self.mtp_step_h(m, &hx, d_prev, p + 1 + j);
4490 extra += 1;
4491 ok = ok && dj == input_ids[p + 2 + j];
4492 Self::chain_probe_note(j, ok);
4493 d_prev = dj;
4494 hx = hj;
4495 }
4496 m.kv.truncate_last(extra);
4497 } else {
4498 let _ = self.mtp_step(
4499 m,
4500 &hb[(p - pos) * hs..(p - pos + 1) * hs],
4501 input_ids[p + 1],
4502 p,
4503 );
4504 }
4505 }
4506 }
4507 }
4508 hidden.copy_from_slice(&hb[(end - pos - 1) * hs..]);
4509 pos = end;
4510 }
4511 }
4512 let pair_off = std::env::var("CMF_PAIR").is_ok_and(|v| v == "0");
4513 if task_mask.is_none()
4514 && !dyn_prefill
4515 && !graph_prefill
4516 && !pair_off
4517 && self.pair_supported()
4518 && o1_prefill.is_none()
4519 {
4520 while pos + 1 < input_ids.len()
4521 && !self.cancel.load(std::sync::atomic::Ordering::Relaxed)
4522 {
4523 let e1 = self.embed_single(input_ids[pos]);
4524 let e2 = self.embed_single(input_ids[pos + 1]);
4525 let (h1, h2) = self.forward_pair(&e1, &e2, pos);
4526 if mimo_spec {
4527 self.mimo_note_rows(&h1, pos);
4528 self.mimo_note_rows(&h2, pos + 1);
4529 }
4530 self.commit_linear_scratch();
4532 if let Some(m) = &mut mtp {
4533 let _ = self.mtp_step(m, &h1, input_ids[pos + 1], pos);
4534 if pos + 2 < input_ids.len() {
4535 let probe: usize = std::env::var("CMF_MTP_CHAIN_PROBE")
4536 .ok()
4537 .and_then(|v| v.parse().ok())
4538 .unwrap_or(0);
4539 if probe >= 1 && pos + 3 < input_ids.len() {
4540 let (d1, mut hx) = self.mtp_step_h(m, &h2, input_ids[pos + 2], pos + 1);
4544 let mut ok = d1 == input_ids[pos + 3];
4545 Self::chain_probe_note(0, ok);
4546 let mut d_prev = d1;
4547 let mut extra = 0usize;
4548 for j in 1..probe {
4549 if pos + 3 + j >= input_ids.len() {
4550 break;
4551 }
4552 let (dj, hj) = self.mtp_step_h(m, &hx, d_prev, pos + 2 + j);
4553 extra += 1;
4554 ok = ok && dj == input_ids[pos + 3 + j];
4555 Self::chain_probe_note(j, ok);
4556 d_prev = dj;
4557 hx = hj;
4558 }
4559 m.kv.truncate_last(extra);
4560 } else {
4561 let _ = self.mtp_step(m, &h2, input_ids[pos + 2], pos + 1);
4562 }
4563 }
4564 }
4565 hidden = h2;
4566 pos += 2;
4567 }
4568 }
4569 let o1_batch_ready = o1_sealed
4582 && o1_prefill.is_some()
4583 && mtp.is_none()
4584 && std::env::var("CMF_O1_GPU").as_deref() == Ok("1")
4585 && (0..self.num_layers).all(|li| {
4586 let cache = &self.kv_cache.layers[self.phys_layer(li)];
4587 cache.o1.is_none() || cache.o1_views().is_some()
4588 });
4589 let mtp_batch_prefill = mtp.is_some()
4594 && graph_prefill
4595 && task_mask.is_none()
4596 && !dyn_prefill
4597 && !self.o1_active()
4598 && std::env::var("CMF_MTP_CHAIN_PROBE").is_err();
4599 if batch_k > 0
4600 && (graph_prefill || o1_batch_ready)
4601 && task_mask.is_none()
4602 && (!self.o1_active() || o1_batch_ready)
4603 && (mtp.is_none() || mtp_batch_prefill)
4604 && !dyn_prefill
4605 && pos + 1 < input_ids.len()
4606 {
4607 let hs = self.hidden_size;
4608 let chunk = batch_k;
4609 while pos < input_ids.len() {
4610 let end = (pos + chunk).min(input_ids.len());
4611 let bk = end - pos;
4612 let mut hiddens = vec![0f32; bk * hs];
4613 for (j, &id) in input_ids[pos..end].iter().enumerate() {
4614 hiddens[j * hs..(j + 1) * hs].copy_from_slice(&self.embed_single(id));
4615 }
4616 let positions: Vec<usize> = (pos..end).collect();
4617 let t_chunk = std::time::Instant::now();
4618 let outcome = self.try_batch_graph_wgpu(&mut hiddens, &positions, bk, None);
4619 let ok_b = outcome == crate::gpu::BatchGraphOutcome::Completed;
4620 if std::env::var("CMF_GRAPH_PROF").is_ok() {
4621 let ms = t_chunk.elapsed().as_secs_f64() * 1000.0;
4622 eprintln!(
4623 "batch-chunk: phase=prompt mode={} k={bk} outcome={outcome:?} {ms:.1} ms ({:.1} tok/s)",
4624 if o1_batch_ready {
4625 "o1"
4626 } else if mtp_batch_prefill {
4627 "ordinary_mtp"
4628 } else {
4629 "ordinary"
4630 },
4631 bk as f64 / (ms / 1000.0)
4632 );
4633 }
4634 {
4635 use std::sync::atomic::{AtomicBool, Ordering};
4636 static SAID: AtomicBool = AtomicBool::new(false);
4637 if !SAID.swap(true, Ordering::Relaxed) {
4638 if ok_b {
4639 tracing::info!(
4640 "batched prefill: ACTIVE mode={} (k={bk})",
4641 if o1_batch_ready {
4642 "o1"
4643 } else if mtp_batch_prefill {
4644 "ordinary_mtp"
4645 } else {
4646 "ordinary"
4647 }
4648 );
4649 } else {
4650 tracing::warn!("batched prefill {:?} — per-position graph", outcome);
4651 }
4652 }
4653 }
4654 if ok_b {
4655 if mimo_spec {
4656 self.mimo_note_rows(&hiddens, pos);
4657 }
4658 if mtp_batch_prefill {
4659 let n_pairs = mtp_prefill_pair_count(pos, end, input_ids.len());
4660 if n_pairs > 0 {
4661 let rows: Vec<Vec<f32>> = (0..n_pairs)
4667 .map(|j| hiddens[j * hs..(j + 1) * hs].to_vec())
4668 .collect();
4669 let pairs: Vec<(&[f32], u32)> = rows
4670 .iter()
4671 .enumerate()
4672 .map(|(j, row)| (row.as_slice(), input_ids[pos + j + 1]))
4673 .collect();
4674 if std::env::var("CMF_GRAPH_PROF").is_ok() {
4675 eprintln!(
4676 "mtp-warm: phase=prompt mode=ordinary_mtp first_pos={} pairs={} last_pos={}",
4677 pos,
4678 n_pairs,
4679 pos + n_pairs - 1,
4680 );
4681 }
4682 let warm_error = if let Some(m) = mtp.as_mut() {
4683 self.mtp_warm_prefill_pairs(m, &pairs, pos).err()
4684 } else {
4685 None
4686 };
4687 if let Some(err) = warm_error {
4688 self.finish_generation(&mut mtp, &mut router, true);
4693 return Err(err.to_string());
4694 }
4695 }
4696 }
4697 hidden.copy_from_slice(&hiddens[(bk - 1) * hs..]);
4698 pos = end;
4699 } else if outcome == crate::gpu::BatchGraphOutcome::Failed {
4700 self.finish_generation(&mut mtp, &mut router, true);
4705 return Err(if o1_batch_ready {
4706 "sealed O(1) batch graph failed after admission".to_string()
4707 } else {
4708 "ordinary recurrent batch graph failed after admission".to_string()
4709 });
4710 } else {
4711 break; }
4713 }
4714 }
4715 if graph_prefill
4719 && task_mask.is_none()
4720 && mtp.is_none()
4721 && !dyn_prefill
4722 && pos == 0
4723 && input_ids.len() > 1
4724 && !self.cancel.load(std::sync::atomic::Ordering::Relaxed)
4725 {
4726 if let Some(lg) = self.embryo_prefill_chunked(input_ids, 0) {
4727 self.graph_logits = Some(lg);
4728 hidden = vec![0.0; self.hidden_size];
4729 pos = input_ids.len();
4730 }
4731 }
4732 while pos < input_ids.len() && !self.cancel.load(std::sync::atomic::Ordering::Relaxed) {
4733 self.graph_want_logits = fuse_lm && pos + 1 == input_ids.len();
4734 hidden = self.forward_layers(&self.embed_single(input_ids[pos]), pos, task_mask);
4735 if mimo_spec {
4736 self.mimo_note_rows(&hidden, pos);
4737 }
4738 if let Some(m) = &mut mtp {
4739 if pos + 1 < input_ids.len() {
4740 let probe: usize = std::env::var("CMF_MTP_CHAIN_PROBE")
4746 .ok()
4747 .and_then(|v| v.parse().ok())
4748 .unwrap_or(0);
4749 if probe >= 1 && pos + 2 < input_ids.len() {
4750 let (d1, mut hx) = self.mtp_step_h(m, &hidden, input_ids[pos + 1], pos);
4751 let mut ok = d1 == input_ids[pos + 2];
4752 Self::chain_probe_note(0, ok);
4753 let mut d_prev = d1;
4754 let mut extra = 0usize;
4755 for j in 1..probe {
4756 if pos + 2 + j >= input_ids.len() {
4757 break;
4758 }
4759 let (dj, hj) = self.mtp_step_h(m, &hx, d_prev, pos + 1 + j);
4760 extra += 1;
4761 ok = ok && dj == input_ids[pos + 2 + j];
4762 Self::chain_probe_note(j, ok);
4763 d_prev = dj;
4764 hx = hj;
4765 }
4766 m.kv.truncate_last(extra);
4769 } else {
4770 let _ = self.mtp_step(m, &hidden, input_ids[pos + 1], pos);
4771 }
4772 }
4773 }
4774 pos += 1;
4775 }
4776 if std::env::var("CMF_PREFILL_PROF").is_ok() {
4777 eprintln!(
4778 "prefill: {} tokens in {:.1} ms (batch_k={batch_k})",
4779 input_ids.len(),
4780 _tpf.elapsed().as_secs_f64() * 1000.0
4781 );
4782 }
4783 if self
4784 .graph_failed
4785 .swap(false, std::sync::atomic::Ordering::Relaxed)
4786 {
4787 self.finish_generation(&mut mtp, &mut router, true);
4792 return Err("GPU token graph failed during prefill".to_string());
4793 }
4794 if self
4797 .cancel
4798 .swap(false, std::sync::atomic::Ordering::Relaxed)
4799 {
4800 self.finish_generation(&mut mtp, &mut router, true);
4804 return Ok(GenerateResult {
4805 text: String::new(),
4806 token_ids: Vec::new(),
4807 prompt_tokens: input_ids.len(),
4808 tokens_generated: 0,
4809 finish_reason: "cancelled".to_string(),
4810 mtp_drafted: 0,
4811 mtp_accepted: 0,
4812 token_confidence: Vec::new(),
4813 traces: Vec::new(),
4814 });
4815 }
4816
4817 if !o1_sealed {
4820 match self.o1_seal_checked() {
4821 Ok(_) => {}
4822 Err(err) => {
4823 self.finish_generation(&mut mtp, &mut router, true);
4824 return Err(err);
4825 }
4826 }
4827 }
4828
4829 macro_rules! commit {
4831 ($id:expr) => {{
4832 all_ids.push($id);
4833 generated += 1;
4834 self.note_draft_id($id);
4835 if self.tokenizer.is_eos($id) && !self.ignore_eos {
4836 finish_reason = "stop".to_string();
4837 false
4838 } else {
4839 let token_text = self.tokenizer.decode_token($id);
4840 let mut go = true;
4841 if let Some(ref mut cb) = on_token {
4842 if !cb(&token_text) {
4843 finish_reason = "cancelled".to_string();
4844 go = false;
4845 }
4846 }
4847 go
4848 }
4849 }};
4850 }
4851
4852 let mut spec_trial = SpecTrial::Spec {
4863 t0: std::time::Instant::now(),
4864 gen0: generated,
4865 rounds: 0,
4866 };
4867 let mut spec_mon = SpecMon {
4873 metal: graph_spec && crate::gpu::q1_force() && spec_cheap_round,
4874 ..SpecMon::default()
4875 };
4876 let mut spec_watchdog_off = false;
4877 let mut spec_walls: Vec<f32> = Vec::new();
4880 let mut spec_round_end: Option<std::time::Instant> = None;
4883 if mimo_spec {
4884 if let Ok(path) = std::env::var("CMF_MIMO_MTP_PROBE") {
4885 if let Some(mut st) = self.mimo_mtp.take() {
4886 self.mimo_mtp_probe(&mut st, input_ids, &path);
4887 self.mimo_mtp = Some(st);
4888 }
4889 }
4890 }
4891 let mut next_pos = input_ids.len();
4893 'decode: while generated < max_tokens {
4894 if self
4895 .graph_failed
4896 .swap(false, std::sync::atomic::Ordering::Relaxed)
4897 {
4898 self.finish_generation(&mut mtp, &mut router, true);
4903 return Err("GPU token graph failed during decode".to_string());
4904 }
4905 if self
4906 .cancel
4907 .swap(false, std::sync::atomic::Ordering::Relaxed)
4908 {
4909 finish_reason = "cancelled".to_string();
4910 break 'decode;
4911 }
4912 if mimo_spec && next_pos > 0 {
4917 self.mimo_note_rows(&hidden, next_pos - 1);
4920 }
4921 let forced = self.spec_forced.take();
4922 let mut logits = match (forced, self.graph_logits.take()) {
4923 (Some(_), _) => Vec::new(),
4924 (None, Some(lg)) => lg,
4925 (None, None) => {
4926 let _prof = crate::cpuprof::time(crate::cpuprof::Slot::Head);
4927 inference::rms_norm_into(
4928 &hidden,
4929 &self.weights.final_norm,
4930 self.rms_eps,
4931 self.norm_style,
4932 &mut self.ws.n1,
4933 );
4934 self.lm_head_forward(&self.ws.n1)
4935 }
4936 };
4937 if generated
4940 == std::env::var("CMF_LOGIT_DUMP_STEP")
4941 .ok()
4942 .and_then(|v| v.parse().ok())
4943 .unwrap_or(0)
4944 {
4945 if let Ok(path) = std::env::var("CMF_LOGIT_DUMP") {
4946 let mut bytes: Vec<u8> = Vec::with_capacity((hidden.len() + logits.len()) * 4);
4947 for v in hidden.iter().chain(logits.iter()) {
4948 bytes.extend_from_slice(&v.to_le_bytes());
4949 }
4950 if let Err(e) = std::fs::write(&path, &bytes) {
4951 eprintln!("logit dump: failed to write {path}: {e}");
4952 self.finish_generation(&mut mtp, &mut router, true);
4953 return Err(format!("logit dump write failed: {e}"));
4954 }
4955 }
4956 }
4957 if let Ok(dir) = std::env::var("CMF_LOGIT_DUMP_ALL") {
4961 if !logits.is_empty() {
4962 let path = std::path::Path::new(&dir).join(format!("step{generated:05}.f32"));
4963 let bytes: Vec<u8> = logits.iter().flat_map(|v| v.to_le_bytes()).collect();
4964 if let Err(e) =
4965 std::fs::create_dir_all(&dir).and_then(|_| std::fs::write(&path, &bytes))
4966 {
4967 eprintln!("logit dump: failed to write {}: {e}", path.display());
4968 }
4969 }
4970 }
4971 let t_next = match forced {
4972 Some(c) => c,
4973 None => {
4974 let _prof = crate::cpuprof::time(crate::cpuprof::Slot::Sampler);
4975 sampler::sample_with_scratch_pool(
4976 &logits,
4977 &self.sampler_config,
4978 self.sampler_config.penalty_past(&all_ids, bounded_native),
4979 &mut self.rng,
4980 &mut self.sampler_scratch,
4981 self.pool.as_deref(),
4982 )
4983 }
4984 };
4985 if self.confidence_on {
4986 confidence.push(if logits.is_empty() {
4987 0.0
4988 } else {
4989 sampler::top1_prob_pool(
4990 self.pool.as_deref(),
4991 &mut self.sampler_scratch,
4992 &logits,
4993 t_next,
4994 calib_temp,
4995 )
4996 });
4997 }
4998 if !logits.is_empty() {
4999 attention::recycle_buf(&mut logits);
5000 }
5001 if trace_on {
5002 let skill = router.as_ref().and_then(|r| r.active_id());
5006 traces.push(TokenTrace {
5007 t: generated,
5008 token_id: t_next,
5009 confidence: confidence.last().copied().unwrap_or(0.0),
5010 active_skill: skill,
5011 recon: None,
5012 switched: false,
5013 });
5014 }
5015 if !commit!(t_next) {
5016 break 'decode;
5017 }
5018 if generated >= max_tokens {
5019 break 'decode;
5020 }
5021
5022 if self.dsv41.is_none() && self.kv_cache.needs_eviction() {
5023 static SAID: std::sync::Once = std::sync::Once::new();
5029 SAID.call_once(|| {
5030 tracing::warn!(
5031 "KV cache full at {} positions — evicting half; quality \
5032 will degrade. Raise CMF_MAX_SEQ.",
5033 self.kv_cache.max_seq_len,
5034 );
5035 });
5036 let keep = (self.kv_cache.max_seq_len / 2).max(1);
5037 self.kv_cache.evict(keep);
5038 }
5039
5040 if graph_spec {
5043 match spec_trial {
5044 SpecTrial::Plain { t0, gen0 } if spec_mon.plain_done(t0, gen0, generated) => {
5045 spec_mon.plain_ms =
5046 t0.elapsed().as_secs_f64() * 1e3 / (generated - gen0) as f64;
5047 let keep = spec_mon.pays();
5048 tracing::info!(
5049 "speculation trial: {:.2} tok/round in {:.1} ms vs plain {:.1} ms/tok — {}",
5050 spec_mon.tokens,
5051 spec_mon.round_ms,
5052 spec_mon.plain_ms,
5053 if keep { "speculating" } else { "plain" }
5054 );
5055 spec_mon.fails = 0;
5056 spec_trial = SpecTrial::Decided {
5057 spec: keep,
5058 recheck_at: if keep { usize::MAX } else { generated + 128 },
5059 };
5060 }
5061 SpecTrial::Decided { recheck_at, .. } if generated >= recheck_at => {
5062 spec_mon.n = 0;
5063 spec_trial = SpecTrial::Spec {
5064 t0: std::time::Instant::now(),
5065 gen0: generated,
5066 rounds: 0,
5067 };
5068 }
5069 _ => {}
5070 }
5071 spec_watchdog_off = matches!(
5072 spec_trial,
5073 SpecTrial::Plain { .. } | SpecTrial::Decided { spec: false, .. }
5074 );
5075 }
5076 if mimo_spec && generated + 1 < max_tokens && next_pos > 0 {
5078 let budget = max_tokens - generated - 1;
5079 if let Some(mut st) = self.mimo_mtp.take() {
5080 let k = st.depth.min(budget);
5081 let r = self.mimo_spec_round(&mut st, next_pos, &all_ids, k);
5082 self.mimo_mtp = Some(st);
5083 let r = match r {
5084 Ok(r) => r,
5085 Err(err) => {
5086 self.finish_generation(&mut mtp, &mut router, true);
5087 return Err(err);
5088 }
5089 };
5090 if let Some(r) = r {
5091 drafted += r.drafted;
5092 accepted += r.accepted.len();
5093 let mut stopped = false;
5094 for &id in &r.accepted {
5095 if self.confidence_on {
5096 confidence.push(0.0);
5097 }
5098 if !commit!(id) {
5099 stopped = true;
5100 break;
5101 }
5102 }
5103 if stopped {
5104 break 'decode;
5105 }
5106 next_pos += r.accepted.len() + 1;
5107 hidden = r.hidden;
5108 self.graph_logits = Some(r.logits);
5111 continue 'decode;
5112 }
5113 }
5114 }
5115 match &mut mtp {
5116 #[cfg(feature = "gpu")]
5118 Some(m)
5119 if graph_spec
5120 && !spec_watchdog_off
5121 && generated + 1 < max_tokens
5122 && next_pos > 0 =>
5123 {
5124 let t_round = std::time::Instant::now();
5125 if spec_time_level() >= 2 {
5126 if let Some(t) = spec_round_end.take() {
5127 eprintln!(
5128 "spec-gap {:.2} ms (host between rounds)",
5129 t.elapsed().as_secs_f64() * 1e3
5130 );
5131 }
5132 }
5133 spec_stamps_begin();
5134 #[cfg(target_os = "macos")]
5139 let allocs0 = crate::gpu_metal::IO_BUF_ALLOCS
5140 .load(std::sync::atomic::Ordering::Relaxed);
5141 #[cfg(not(target_os = "macos"))]
5142 let allocs0 = 0u64;
5143 if let Some((extra, n_pos, new_h)) = self.graph_spec_step(
5144 m,
5145 &hidden,
5146 t_next,
5147 next_pos,
5148 &mut drafted,
5149 &mut accepted,
5150 &mut all_ids,
5151 max_tokens - generated,
5152 ) {
5153 next_pos = n_pos;
5154 hidden = new_h;
5155 let level = spec_time_level();
5156 if level > 0 {
5157 let wall = t_round.elapsed().as_secs_f32() * 1e3;
5158 let stamps = spec_stamps_take();
5159 let median = if spec_walls.len() >= 3 {
5162 let mut s = spec_walls.clone();
5163 s.sort_by(|a, b| a.partial_cmp(b).unwrap());
5164 Some(s[s.len() / 2])
5165 } else {
5166 None
5167 };
5168 let outlier = median.is_some_and(|m| wall > 1.4 * m);
5169 #[cfg(target_os = "macos")]
5170 let allocs = crate::gpu_metal::IO_BUF_ALLOCS
5171 .load(std::sync::atomic::Ordering::Relaxed)
5172 - allocs0;
5173 #[cfg(not(target_os = "macos"))]
5174 let allocs = allocs0;
5175 eprintln!(
5176 "spec-round wall {wall:.1} ms → {} tokens{}{}",
5177 extra.len() + 1,
5178 if allocs > 0 {
5179 format!(" [{allocs} new device buffers]")
5180 } else {
5181 String::new()
5182 },
5183 match (outlier, median) {
5184 (true, Some(m)) => format!(" OUTLIER (median {m:.1})"),
5185 _ => String::new(),
5186 }
5187 );
5188 if level >= 2 || outlier {
5189 let sum: f32 = stamps.iter().map(|s| s.1).sum();
5190 eprintln!(
5191 "spec-stamps: {}| untracked {:.1}",
5192 spec_stamps_format(&stamps),
5193 wall - sum
5194 );
5195 }
5196 if spec_mon.n >= 1 {
5197 spec_walls.push(wall);
5198 }
5199 }
5200 spec_mon.round(t_round.elapsed().as_secs_f64() * 1e3, extra.len() + 1);
5204 spec_trial = Self::spec_trial_round(
5207 spec_trial,
5208 &mut spec_mon,
5209 generated + extra.len() + 1,
5210 );
5211 let mut stopped = false;
5212 for &id in &extra {
5213 if self.confidence_on {
5214 confidence.push(0.0);
5215 }
5216 if !commit!(id) {
5217 stopped = true;
5218 break;
5219 }
5220 }
5221 if stopped {
5222 break 'decode;
5223 }
5224 if spec_time_level() >= 2 {
5225 spec_round_end = Some(std::time::Instant::now());
5226 }
5227 continue 'decode;
5228 }
5229 if self
5230 .graph_failed
5231 .swap(false, std::sync::atomic::Ordering::Relaxed)
5232 {
5233 self.finish_generation(&mut mtp, &mut router, true);
5239 return Err("GPU MTP graph failed during speculative decode".to_string());
5240 }
5241 spec_mon.round(t_round.elapsed().as_secs_f64() * 1e3, 1);
5252 spec_mon.tokens = 0.0;
5253 spec_mon.fails = 3;
5254 spec_trial = Self::spec_trial_round(spec_trial, &mut spec_mon, generated + 1);
5255 hidden = self.forward_layers(&self.embed_single(t_next), next_pos, task_mask);
5256 next_pos += 1;
5257 continue 'decode;
5258 }
5259 Some(m) if !graph_spec && generated + 1 < max_tokens => {
5261 let draft = self.mtp_step(m, &hidden, t_next, next_pos - 1);
5262 drafted += 1;
5263 let emb1 = self.embed_single(t_next);
5264 let emb2 = self.embed_single(draft);
5265 let (h1, h2) = self.forward_pair(&emb1, &emb2, next_pos);
5266
5267 inference::rms_norm_into(
5268 &h1,
5269 &self.weights.final_norm,
5270 self.rms_eps,
5271 self.norm_style,
5272 &mut self.ws.n1,
5273 );
5274 let mut logits1 = self.lm_head_forward(&self.ws.n1);
5275 let t_after = sampler::sample_with_scratch_pool(
5276 &logits1,
5277 &self.sampler_config,
5278 self.sampler_config.penalty_past(&all_ids, bounded_native),
5279 &mut self.rng,
5280 &mut self.sampler_scratch,
5281 self.pool.as_deref(),
5282 );
5283 if self.confidence_on {
5284 confidence.push(sampler::top1_prob_pool(
5285 self.pool.as_deref(),
5286 &mut self.sampler_scratch,
5287 &logits1,
5288 t_after,
5289 calib_temp,
5290 ));
5291 }
5292 attention::recycle_buf(&mut logits1);
5293 if trace_on {
5294 traces.push(TokenTrace {
5297 t: generated,
5298 token_id: t_after,
5299 confidence: confidence.last().copied().unwrap_or(0.0),
5300 active_skill: None,
5301 recon: None,
5302 switched: false,
5303 });
5304 }
5305 let stop = !commit!(t_after);
5306
5307 if t_after == draft {
5308 accepted += 1;
5309 self.commit_linear_scratch();
5310 let _ = self.mtp_step(m, &h1, t_after, next_pos);
5311 hidden = h2;
5312 next_pos += 2;
5313 } else {
5314 for layer in &mut self.kv_cache.layers {
5316 layer.truncate_last(1);
5317 }
5318 if !stop {
5319 let _ = self.mtp_step(m, &h1, t_after, next_pos);
5320 hidden = self.forward_layers(
5321 &self.embed_single(t_after),
5322 next_pos + 1,
5323 None,
5324 );
5325 }
5326 next_pos += 2;
5327 }
5328 if stop {
5329 break 'decode;
5330 }
5331 }
5332 _ => {
5334 #[cfg(feature = "gpu")]
5339 if Self::dsv4_spec_on() && self.dsv4.is_some() {
5340 static SAID: std::sync::Once = std::sync::Once::new();
5341 SAID.call_once(|| {
5342 eprintln!(
5343 "dsv4-spec гейт: mtp={} mask={} router={} trace={} temp={} rep={} ",
5344 !self.dsv4_mtp.is_empty(),
5345 task_mask.is_none(),
5346 router.is_none(),
5347 !trace_on,
5348 self.sampler_config.temperature < 1e-6,
5349 self.sampler_config.repetition_penalty == 1.0,
5350 );
5351 });
5352 }
5353 #[cfg(feature = "gpu")]
5354 if Self::dsv4_spec_on()
5355 && self.dsv4.is_some()
5356 && !self.dsv4_mtp.is_empty()
5357 && task_mask.is_none()
5358 && router.is_none()
5359 && !trace_on
5360 && self.sampler_config.temperature < 1e-6
5361 && self.sampler_config.repetition_penalty == 1.0
5362 && generated + 1 < max_tokens
5363 && all_ids.len() >= 2
5364 && generated >= dsv4_spec_retry_at
5365 {
5366 let tip_token = all_ids[all_ids.len() - 2];
5367 let drafted0 = drafted;
5368 let round = self.dsv4_spec_step(
5369 tip_token,
5370 t_next,
5371 next_pos,
5372 max_tokens.saturating_sub(generated),
5373 &mut drafted,
5374 &mut accepted,
5375 );
5376 if drafted > drafted0 {
5377 let useful = round.as_ref().is_some_and(|(extra, _)| !extra.is_empty());
5378 if useful {
5379 dsv4_spec_bad = 0;
5380 } else {
5381 dsv4_spec_bad += 1;
5382 if dsv4_spec_bad >= 2 {
5383 dsv4_spec_bad = 0;
5384 dsv4_spec_retry_at = generated.saturating_add(32);
5385 tracing::info!(
5386 "dsv4: draft не окупился дважды — точный walk на 32 токена"
5387 );
5388 }
5389 }
5390 }
5391 if let Some((extra, n_pos)) = round {
5392 next_pos = n_pos;
5393 let mut stopped = false;
5394 for &id in &extra {
5395 if self.confidence_on {
5396 confidence.push(0.0);
5397 }
5398 if !commit!(id) {
5399 stopped = true;
5400 break;
5401 }
5402 }
5403 if stopped {
5404 break 'decode;
5405 }
5406 continue 'decode;
5407 }
5408 }
5409 self.graph_want_logits = fuse_lm;
5410 let mut t_fwd = t_next;
5416 let pure_greedy = self.sampler_config.temperature < 1e-6
5417 && self.sampler_config.repetition_penalty == 1.0
5418 && self.sampler_config.suppress_tokens.is_empty();
5419 let burst_k = std::env::var("CMF_MULTISTEP")
5424 .ok()
5425 .and_then(|v| v.parse::<usize>().ok())
5426 .unwrap_or(0);
5427 if pure_greedy
5428 && burst_k >= 1
5429 && fuse_lm
5430 && task_mask.is_none()
5431 && router.is_none()
5432 && !trace_on
5433 && !self.confidence_on
5434 {
5435 let mut stopped = false;
5436 loop {
5437 let room = max_tokens.saturating_sub(generated);
5438 if room <= 2 {
5439 break;
5440 }
5441 let k = burst_k.min(room - 1);
5442 if k < 1 {
5443 break;
5444 }
5445 let Some(ids) = self.try_multi_burst(t_fwd, next_pos, k) else {
5446 if self
5447 .graph_failed
5448 .swap(false, std::sync::atomic::Ordering::Relaxed)
5449 {
5450 self.finish_generation(&mut mtp, &mut router, true);
5451 return Err(
5452 "GPU token graph failed during greedy burst".to_string()
5453 );
5454 }
5455 break;
5456 };
5457 next_pos += k;
5458 for &id in &ids {
5459 if !commit!(id) {
5460 stopped = true;
5461 break;
5462 }
5463 }
5464 if stopped {
5465 break;
5466 }
5467 t_fwd = *ids.last().unwrap();
5468 }
5469 if stopped {
5470 break 'decode;
5471 }
5472 }
5473 #[cfg(target_os = "macos")]
5483 if graph_spec
5484 && spec_watchdog_off
5485 && next_pos > 0
5486 && self.mtp_graph_mode == Some(true)
5487 && crate::gpu::q1_force()
5488 {
5489 if let Some(m) = mtp.as_mut() {
5490 let _ = self.mtp_step_metal(m, &hidden, t_fwd, next_pos - 1, false);
5491 }
5492 }
5493 hidden = self.forward_layers(&self.embed_single(t_fwd), next_pos, task_mask);
5494 next_pos += 1;
5495 if let Some(r) = &mut router {
5498 let phi = self.dyn_phi_ema.clone();
5499 let decision = r.step(&phi, generated);
5500 if let Some(new_active) = decision {
5501 let _ = self.set_active_skill(new_active);
5502 }
5503 if trace_on {
5506 if let Some(last) = traces.last_mut() {
5507 let e = r.last_best_e();
5508 last.recon = e.is_finite().then_some(e);
5509 last.switched = decision.is_some();
5510 }
5511 }
5512 }
5513 }
5514 }
5515 }
5516
5517 let cancelled = finish_reason == "cancelled";
5518 let dyn_switched = router.as_ref().is_some_and(|r| !r.switches.is_empty());
5522 if mimo_spec {
5523 if let Some(st) = self.mimo_mtp.as_ref() {
5524 let line = st.stats.line();
5525 tracing::info!("{line}");
5526 if std::env::var_os("CMF_MIMO_MTP_STATS").is_some() {
5527 eprintln!("{line}");
5528 }
5529 }
5530 }
5531 self.finish_generation(&mut mtp, &mut router, cancelled);
5532
5533 let output_ids = &all_ids[input_ids.len()..];
5534 let forwarded = input_ids.len() + output_ids.len().saturating_sub(1);
5538 let consumed = std::mem::take(&mut all_ids);
5543 if dyn_switched {
5544 self.clear_sequence_state();
5545 } else if cancelled || mimo_spec || prompt_rows.is_some() {
5546 self.clear_history();
5547 } else {
5548 self.record_consumed_prefix(&consumed[..forwarded.min(consumed.len())], reuse_from);
5549 }
5550 all_ids = consumed;
5551 let output_ids = &all_ids[input_ids.len()..];
5552 confidence.truncate(output_ids.len()); traces.truncate(output_ids.len());
5554 Ok(GenerateResult {
5555 text: self.tokenizer.decode(output_ids),
5556 token_ids: output_ids.to_vec(),
5557 prompt_tokens: input_ids.len(),
5558 tokens_generated: generated,
5559 finish_reason,
5560 mtp_drafted: drafted,
5561 mtp_accepted: accepted,
5562 token_confidence: confidence,
5563 traces,
5564 })
5565 }
5566
5567 fn mtp_step(
5571 &mut self,
5572 m: &mut MtpModule,
5573 hidden: &[f32],
5574 next_token: u32,
5575 position: usize,
5576 ) -> u32 {
5577 self.mtp_step_h(m, hidden, next_token, position).0
5578 }
5579
5580 fn chain_probe_note(depth: usize, prefix_ok: bool) {
5584 use std::sync::Mutex;
5585 static T: Mutex<Vec<(u64, u64)>> = Mutex::new(Vec::new());
5586 let mut t = T.lock().unwrap();
5587 if t.len() <= depth {
5588 t.resize(depth + 1, (0, 0));
5589 }
5590 t[depth].0 += 1;
5591 t[depth].1 += prefix_ok as u64;
5592 if depth == 0 && t[0].0 % 128 == 0 {
5593 let line: Vec<String> = t
5594 .iter()
5595 .enumerate()
5596 .map(|(d, (n, k))| {
5597 format!(
5598 "d{}={:.0}%({n})",
5599 d + 1,
5600 100.0 * *k as f64 / (*n).max(1) as f64
5601 )
5602 })
5603 .collect();
5604 eprintln!("mtp-chain: {}", line.join(" "));
5605 }
5606 }
5607
5608 fn mtp_step_hl(
5616 &mut self,
5617 m: &mut MtpModule,
5618 hidden: &[f32],
5619 next_token: u32,
5620 position: usize,
5621 ) -> (Vec<f32>, Vec<f32>) {
5622 #[cfg(target_os = "macos")]
5627 if self.mtp_graph_mode != Some(false) && crate::gpu::q1_force() {
5628 if let Some(r) = self.mtp_step_metal(m, hidden, next_token, position, true) {
5629 self.mtp_graph_mode = Some(true);
5630 return r;
5631 }
5632 if self.mtp_graph_mode == Some(true) {
5633 tracing::error!("mtp Metal graph failed after admission");
5634 self.clear_sequence_state();
5635 self.graph_failed
5636 .store(true, std::sync::atomic::Ordering::Relaxed);
5637 self.cancel
5638 .store(true, std::sync::atomic::Ordering::Relaxed);
5639 return (Vec::new(), Vec::new());
5640 }
5641 self.mtp_graph_mode = Some(false);
5642 }
5643 #[cfg(feature = "gpu")]
5644 if self.mtp_graph_mode != Some(false) {
5645 if !self.mtp_graph_ok(m) {
5646 if self.mtp_graph_mode == Some(true) {
5647 tracing::error!("mtp graph became unavailable after admission");
5652 self.clear_sequence_state();
5653 self.graph_failed
5654 .store(true, std::sync::atomic::Ordering::Relaxed);
5655 self.cancel
5656 .store(true, std::sync::atomic::Ordering::Relaxed);
5657 return (Vec::new(), Vec::new());
5658 }
5659 self.mtp_graph_mode = Some(false);
5660 } else {
5661 if let Some(r) = self.mtp_step_graph(m, hidden, next_token, position) {
5662 self.mtp_graph_mode = Some(true);
5663 return r;
5664 }
5665 if self.graph_failed.load(std::sync::atomic::Ordering::Relaxed) {
5666 return (Vec::new(), Vec::new());
5673 }
5674 tracing::error!("mtp graph failed or declined after admission");
5678 self.clear_sequence_state();
5679 self.graph_failed
5680 .store(true, std::sync::atomic::Ordering::Relaxed);
5681 self.cancel
5682 .store(true, std::sync::atomic::Ordering::Relaxed);
5683 return (Vec::new(), Vec::new());
5684 }
5685 }
5686 let e = self.embed_single(next_token);
5690 let mut cat = vec![0.0f32; 2 * self.hidden_size];
5691 let (cat_e, cat_h) = cat.split_at_mut(self.hidden_size);
5692 inference::rms_norm_into(&e, &m.enorm, self.rms_eps, self.norm_style, cat_e);
5693 inference::rms_norm_into(hidden, &m.hnorm, self.rms_eps, self.norm_style, cat_h);
5694 let mut x = vec![0.0f32; self.hidden_size];
5695 m.eh_proj.matvec(&cat, &mut x, self.pool.as_deref());
5696
5697 let lw = &m.layer;
5699 inference::rms_norm_into(
5700 &x,
5701 &lw.input_norm,
5702 self.rms_eps,
5703 self.norm_style,
5704 &mut self.ws.n1,
5705 );
5706 let attn = match &lw.attn {
5707 AttnKind::Mla(_) => unreachable!("MLA has no MTP/pair path"),
5709 AttnKind::Kda(_) => unreachable!("KDA has no MTP/pair path"),
5710 AttnKind::Full {
5711 wq,
5712 wk,
5713 wv,
5714 wo,
5715 q_norm,
5716 k_norm,
5717 output_gate,
5718 softplus_gate,
5719 bias,
5720 } => {
5721 let mut cfg = self.attn_cfg(position);
5722 cfg.q_norm = q_norm.as_deref();
5723 cfg.k_norm = k_norm.as_deref();
5724 cfg.output_gate = *output_gate;
5725 cfg.softplus_gate = softplus_gate
5726 .as_ref()
5727 .map(|(gate, per_head)| (gate, *per_head));
5728 cfg.bias = bias
5729 .as_ref()
5730 .map(|(q, k, v)| (q.as_slice(), k.as_slice(), v.as_slice()));
5731 attention::qwen_attention(&self.ws.n1, wq, wk, wv, wo, &mut m.kv, &cfg)
5732 }
5733 AttnKind::Linear(_)
5734 | AttnKind::LinearGdn(_)
5735 | AttnKind::ShortConv(_)
5736 | AttnKind::Bounded(_) => {
5737 unreachable!("MTP block is full attention")
5738 }
5739 };
5740 for (i, &a) in attn.iter().enumerate() {
5741 x[i] += a;
5742 }
5743 inference::rms_norm_into(
5744 &x,
5745 &lw.post_norm,
5746 self.rms_eps,
5747 self.norm_style,
5748 &mut self.ws.p1,
5749 );
5750 let ffn = ffn_forward(&lw.ffn, &self.ws.p1, self.pool.as_deref(), None);
5751 for (i, &f) in ffn.iter().enumerate() {
5752 x[i] += f;
5753 }
5754
5755 inference::rms_norm_into(
5756 &x,
5757 &m.final_norm,
5758 self.rms_eps,
5759 self.norm_style,
5760 &mut self.ws.n1,
5761 );
5762 let lg = self.lm_head_forward(&self.ws.n1);
5763 (lg, x)
5764 }
5765
5766 fn mtp_step_h(
5768 &mut self,
5769 m: &mut MtpModule,
5770 hidden: &[f32],
5771 next_token: u32,
5772 position: usize,
5773 ) -> (u32, Vec<f32>) {
5774 let (mut lg, x) = self.mtp_step_hl(m, hidden, next_token, position);
5775 let draft = sampler::argmax(&lg);
5776 attention::recycle_buf(&mut lg);
5777 (draft, x)
5778 }
5779
5780 fn spec_trial_round(trial: SpecTrial, mon: &mut SpecMon, generated: usize) -> SpecTrial {
5786 match trial {
5787 SpecTrial::Spec { t0, gen0, rounds } => {
5788 let rounds = rounds + 1;
5789 if rounds >= 5 {
5790 if mon.plain_ms > 0.0 {
5791 let keep = mon.pays();
5792 mon.fails = 0;
5793 tracing::info!(
5794 "speculation re-check: {:.2} tok/round in {:.1} ms vs plain {:.1} ms/tok — {}",
5795 mon.tokens,
5796 mon.round_ms,
5797 mon.plain_ms,
5798 if keep { "speculating" } else { "plain" }
5799 );
5800 SpecTrial::Decided {
5801 spec: keep,
5802 recheck_at: if keep { usize::MAX } else { generated + 128 },
5803 }
5804 } else if mon.pays() {
5805 mon.fails = 0;
5810 tracing::info!(
5811 "speculation trial: {:.2} tok/round in {:.1} ms — speculating (plain not timed)",
5812 mon.tokens,
5813 mon.round_ms,
5814 );
5815 SpecTrial::Decided {
5816 spec: true,
5817 recheck_at: usize::MAX,
5818 }
5819 } else {
5820 SpecTrial::Plain {
5821 t0: std::time::Instant::now(),
5822 gen0: generated,
5823 }
5824 }
5825 } else {
5826 SpecTrial::Spec { t0, gen0, rounds }
5827 }
5828 }
5829 SpecTrial::Decided { spec: true, .. } => {
5830 if mon.pays() {
5831 mon.fails = 0;
5832 trial
5833 } else {
5834 mon.fails += 1;
5835 if mon.fails >= 4 {
5836 if mon.plain_ms <= 0.0 {
5837 tracing::info!(
5841 "speculation doubtful: {:.2} tok/round in {:.1} ms — timing plain",
5842 mon.tokens,
5843 mon.round_ms,
5844 );
5845 return SpecTrial::Plain {
5846 t0: std::time::Instant::now(),
5847 gen0: generated,
5848 };
5849 }
5850 tracing::info!(
5851 "speculation stopped: {:.2} tok/round in {:.1} ms vs plain {:.1} ms/tok",
5852 mon.tokens,
5853 mon.round_ms,
5854 mon.plain_ms
5855 );
5856 SpecTrial::Decided {
5857 spec: false,
5858 recheck_at: generated + 128,
5859 }
5860 } else {
5861 trial
5862 }
5863 }
5864 }
5865 other => other,
5866 }
5867 }
5868
5869 fn mtp_kv_id(&self) -> u64 {
5872 self.graph_kv_id | (1u64 << 40)
5873 }
5874
5875 const MTP_LAYER_BASE: usize = 0;
5880
5881 #[cfg(feature = "gpu")]
5888 fn rewind_mtp_graph_mirror(&self, stored: usize) -> bool {
5889 self.mtp_graph_mode != Some(true)
5890 || crate::gpu::graph_kv_set_stored(self.mtp_kv_id(), Self::MTP_LAYER_BASE, stored)
5891 }
5892
5893 #[cfg(feature = "gpu")]
5899 fn rewind_trunk_graph_mirrors(&self, stored: usize) -> bool {
5900 let mut ok = true;
5901 let mut expected = false;
5902 for li in 0..self.num_layers {
5903 if matches!(
5904 self.weights.layers[self.phys_layer(li)].attn,
5905 AttnKind::Full { .. }
5906 ) {
5907 expected = true;
5908 ok &= crate::gpu::graph_kv_set_stored(self.graph_kv_id, li, stored);
5909 }
5910 }
5911 !expected || ok
5912 }
5913
5914 fn graph_gdn_layer_count(&self) -> usize {
5918 (0..self.num_layers)
5919 .filter(|&li| {
5920 matches!(
5921 &self.weights.layers[self.phys_layer(li)].attn,
5922 AttnKind::LinearGdn(_)
5923 )
5924 })
5925 .count()
5926 }
5927
5928 fn mtp_block_input(&mut self, m: &MtpModule, hidden: &[f32], next_token: u32) -> Vec<f32> {
5931 let e = self.embed_single(next_token);
5932 let mut cat = vec![0.0f32; 2 * self.hidden_size];
5933 let (cat_e, cat_h) = cat.split_at_mut(self.hidden_size);
5934 inference::rms_norm_into(&e, &m.enorm, self.rms_eps, self.norm_style, cat_e);
5935 inference::rms_norm_into(hidden, &m.hnorm, self.rms_eps, self.norm_style, cat_h);
5936 let mut x = vec![0.0f32; self.hidden_size];
5937 m.eh_proj.matvec(&cat, &mut x, self.pool.as_deref());
5938 x
5939 }
5940
5941 #[cfg(feature = "gpu")]
5944 fn mtp_block_graph_ok(&self, m: &MtpModule) -> bool {
5945 if std::env::var("CMF_MTP_GRAPH").as_deref() == Ok("0") {
5946 return false;
5947 }
5948 if !crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Decode)
5949 || !crate::gpu::enabled_here()
5950 || self.attn_softcap > 0.0
5951 || self.attention_heads_per_layer.is_some()
5952 || self.v_head_dim.is_some()
5955 {
5956 return false;
5957 }
5958 matches!(
5959 &m.layer.attn,
5960 AttnKind::Full {
5961 softplus_gate: None,
5962 ..
5963 }
5964 ) && matches!(&m.layer.ffn, FfnKind::Dense(_))
5965 }
5966
5967 #[cfg(feature = "gpu")]
5971 fn mtp_graph_ok(&self, m: &MtpModule) -> bool {
5972 if !self.mtp_block_graph_ok(m) {
5973 return false;
5974 }
5975 let AttnKind::Full { wq, wk, wv, wo, .. } = &m.layer.attn else {
5976 return false;
5977 };
5978 let FfnKind::Dense(d) = &m.layer.ffn else {
5979 return false;
5980 };
5981 d.segs.is_empty()
5982 && wq.graph_weight().is_some()
5983 && wk.graph_weight().is_some()
5984 && wv.graph_weight().is_some()
5985 && wo.graph_weight().is_some()
5986 && d.gate_proj.graph_weight().is_some()
5987 && d.up_proj.graph_weight().is_some()
5988 && d.down_proj.graph_weight().is_some()
5989 && self.weights.lm_head.graph_weight().is_some()
5990 }
5991
5992 #[cfg(feature = "gpu")]
5998 fn mtp_step_graph(
5999 &mut self,
6000 m: &mut MtpModule,
6001 hidden: &[f32],
6002 next_token: u32,
6003 position: usize,
6004 ) -> Option<(Vec<f32>, Vec<f32>)> {
6005 if !self.mtp_graph_ok(m) {
6006 return None;
6007 }
6008 let lw = &m.layer;
6009 let AttnKind::Full {
6010 wq,
6011 wk,
6012 wv,
6013 wo,
6014 q_norm,
6015 k_norm,
6016 output_gate,
6017 softplus_gate,
6018 bias,
6019 } = &lw.attn
6020 else {
6021 return None;
6022 };
6023 if softplus_gate.is_some() {
6024 return None;
6025 }
6026 let FfnKind::Dense(d) = &lw.ffn else {
6027 return None;
6028 };
6029 if !d.segs.is_empty() {
6030 return None; }
6032 let mut x = self.mtp_block_input(m, hidden, next_token);
6035 fn gw(t: &QTensor) -> Option<crate::gpu::GraphW<'_>> {
6036 let (_, i, kind, rs) = t.graph_weight()?;
6037 Some(crate::gpu::GraphW {
6038 idx: i,
6039 kind,
6040 row_scale: rs,
6041 data: &[],
6042 prism: crate::gpu::GraphPrismOp::None,
6043 affine: false,
6044 })
6045 }
6046 let (model, _, _, _) = wq.graph_weight()?;
6047 let model = model.clone();
6048 let (lm_gw, lm_rows) = {
6049 let (_, i, kind, rs) = self.weights.lm_head.graph_weight()?;
6050 let rows = if kind == 6 {
6054 self.draft_head_rows(self.weights.lm_head.rows())
6055 } else {
6056 self.weights.lm_head.rows()
6057 };
6058 (
6059 crate::gpu::GraphW {
6060 idx: i,
6061 kind,
6062 row_scale: rs,
6063 data: &[],
6064 prism: crate::gpu::GraphPrismOp::None,
6065 affine: false,
6066 },
6067 rows,
6068 )
6069 };
6070 let layer = crate::gpu::GraphLayer {
6071 input_norm: &lw.input_norm,
6072 attn: crate::gpu::GraphAttn::Full {
6073 wq: gw(wq)?,
6074 wk: gw(wk)?,
6075 wv: gw(wv)?,
6076 wo: gw(wo)?,
6077 q_norm: q_norm.as_deref(),
6078 k_norm: k_norm.as_deref(),
6079 late_qk_norm: self.qk_norm_after_rope,
6080 bias: bias
6081 .as_ref()
6082 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
6083 output_gate: *output_gate,
6084 cpu_k: m.kv.k_heads(),
6085 cpu_v: m.kv.v_heads(),
6086 geom: None,
6087 },
6088 post_norm: &lw.post_norm,
6089 ffn: crate::gpu::GraphFfn::Dense {
6090 gate: gw(&d.gate_proj)?,
6091 up: gw(&d.up_proj)?,
6092 down: gw(&d.down_proj)?,
6093 },
6094 };
6095 let nh = self.num_heads;
6096 let (nkv, hd, rd) = self.layer_geom(0);
6097 let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
6098 let mut logits = Vec::new();
6099 let ok = crate::gpu::forward_token_graph(
6100 &model,
6101 self.mtp_kv_id(),
6102 std::slice::from_ref(&layer),
6103 &[None],
6104 self.o1_epoch,
6105 &self.inv_freq,
6106 &mut x,
6107 nh,
6108 nkv,
6109 hd,
6110 self.attn_scale,
6111 rd,
6112 self.hidden_size,
6113 self.intermediate_size,
6114 position,
6115 self.kv_cache.max_seq_len,
6116 gemma,
6117 self.rms_eps as f32,
6118 Some((&lm_gw, lm_rows)),
6119 &m.final_norm,
6120 &mut logits,
6121 &[],
6122 1,
6123 None,
6124 None,
6125 None,
6126 Self::MTP_LAYER_BASE,
6127 true,
6128 );
6129 match ok {
6130 crate::gpu::TokenGraphOutcome::Completed => {}
6131 crate::gpu::TokenGraphOutcome::Declined => return None,
6132 crate::gpu::TokenGraphOutcome::Failed => {
6133 self.clear_sequence_state();
6137 self.graph_failed
6138 .store(true, std::sync::atomic::Ordering::Relaxed);
6139 self.cancel
6140 .store(true, std::sync::atomic::Ordering::Relaxed);
6141 return None;
6142 }
6143 }
6144 logits.resize(self.vocab_size, 0.0);
6145 Some((logits, x))
6146 }
6147
6148 #[cfg(feature = "gpu")]
6156 fn mtp_warm_graph(
6157 &mut self,
6158 m: &mut MtpModule,
6159 pairs: &[(&[f32], u32)],
6160 first_pos: usize,
6161 ) -> crate::gpu::BatchGraphOutcome {
6162 if pairs.is_empty() {
6163 return crate::gpu::BatchGraphOutcome::Completed;
6164 }
6165 if !self.mtp_block_graph_ok(m) {
6166 return crate::gpu::BatchGraphOutcome::Declined;
6167 }
6168 let hs = self.hidden_size;
6169 let mut hiddens = Vec::with_capacity(pairs.len() * hs);
6172 for (h, t) in pairs {
6173 hiddens.extend_from_slice(&self.mtp_block_input(m, h, *t));
6174 }
6175 let lw = &m.layer;
6176 let AttnKind::Full {
6177 wq,
6178 wk,
6179 wv,
6180 wo,
6181 q_norm,
6182 k_norm,
6183 output_gate,
6184 bias,
6185 ..
6186 } = &lw.attn
6187 else {
6188 return crate::gpu::BatchGraphOutcome::Declined;
6189 };
6190 let FfnKind::Dense(d) = &lw.ffn else {
6191 return crate::gpu::BatchGraphOutcome::Declined;
6192 };
6193 if !d.segs.is_empty() {
6194 return crate::gpu::BatchGraphOutcome::Declined; }
6196 fn gw(t: &QTensor) -> Option<crate::gpu::GraphW<'_>> {
6197 let (_, i, kind, rs) = t.graph_weight()?;
6198 Some(crate::gpu::GraphW {
6199 idx: i,
6200 kind,
6201 row_scale: rs,
6202 data: &[],
6203 prism: crate::gpu::GraphPrismOp::None,
6204 affine: false,
6205 })
6206 }
6207 let Some((model, _, _, _)) = wq.graph_weight() else {
6208 return crate::gpu::BatchGraphOutcome::Declined;
6209 };
6210 let model = model.clone();
6211 let (Some(gwq), Some(gwk), Some(gwv), Some(gwo), Some(gg), Some(gu), Some(gd)) = (
6212 gw(wq),
6213 gw(wk),
6214 gw(wv),
6215 gw(wo),
6216 gw(&d.gate_proj),
6217 gw(&d.up_proj),
6218 gw(&d.down_proj),
6219 ) else {
6220 return crate::gpu::BatchGraphOutcome::Declined;
6221 };
6222 let layer = crate::gpu::GraphLayer {
6223 input_norm: &lw.input_norm,
6224 attn: crate::gpu::GraphAttn::Full {
6225 wq: gwq,
6226 wk: gwk,
6227 wv: gwv,
6228 wo: gwo,
6229 q_norm: q_norm.as_deref(),
6230 k_norm: k_norm.as_deref(),
6231 late_qk_norm: self.qk_norm_after_rope,
6232 bias: bias
6233 .as_ref()
6234 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
6235 output_gate: *output_gate,
6236 cpu_k: m.kv.k_heads(),
6237 cpu_v: m.kv.v_heads(),
6238 geom: None,
6239 },
6240 post_norm: &lw.post_norm,
6241 ffn: crate::gpu::GraphFfn::Dense {
6242 gate: gg,
6243 up: gu,
6244 down: gd,
6245 },
6246 };
6247 let positions: Vec<usize> = (first_pos..first_pos + pairs.len()).collect();
6248 let nh = self.num_heads;
6249 let (nkv, hd, rd) = self.layer_geom(0);
6250 let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
6251 crate::gpu::forward_batch_graph(
6252 &model,
6253 self.mtp_kv_id(),
6254 std::slice::from_ref(&layer),
6255 &self.inv_freq,
6256 &mut hiddens,
6257 nh,
6258 nkv,
6259 hd,
6260 rd,
6261 hs,
6262 self.intermediate_size,
6263 &positions,
6264 self.kv_cache.max_seq_len,
6265 gemma,
6266 self.rms_eps as f32,
6267 self.attn_scale,
6268 pairs.len(),
6269 &[],
6270 0,
6271 None,
6272 None,
6273 )
6274 }
6275
6276 #[cfg(feature = "gpu")]
6283 fn mtp_warm_graph_fallback(
6284 &mut self,
6285 m: &mut MtpModule,
6286 pairs: &[(&[f32], u32)],
6287 first_pos: usize,
6288 ) -> bool {
6289 if pairs.is_empty() {
6290 return true;
6291 }
6292 let graphable = self.mtp_block_graph_ok(m);
6293 if !graphable {
6294 if self.mtp_graph_mode == Some(true) {
6298 return false;
6299 }
6300 self.mtp_graph_mode = Some(false);
6301 for (j, (h, t)) in pairs.iter().enumerate() {
6302 self.mtp_warm(m, h, *t, first_pos + j);
6303 }
6304 return true;
6305 }
6306
6307 for (j, (h, t)) in pairs.iter().enumerate() {
6312 if self.mtp_step_graph(m, h, *t, first_pos + j).is_none() {
6313 return false;
6314 }
6315 }
6316 self.mtp_graph_mode = Some(true);
6317 true
6318 }
6319
6320 #[cfg(feature = "gpu")]
6325 fn mtp_warm_prefill_pairs(
6326 &mut self,
6327 m: &mut MtpModule,
6328 pairs: &[(&[f32], u32)],
6329 first_pos: usize,
6330 ) -> Result<(), &'static str> {
6331 if self.mtp_graph_mode == Some(false) || !self.mtp_graph_ok(m) {
6336 if self.mtp_graph_mode == Some(true) {
6337 return Err("MTP token graph became unavailable after admission");
6338 }
6339 self.mtp_graph_mode = Some(false);
6340 for (j, (h, t)) in pairs.iter().enumerate() {
6341 self.mtp_warm(m, h, *t, first_pos + j);
6342 }
6343 return Ok(());
6344 }
6345 match self.mtp_warm_graph(m, pairs, first_pos) {
6346 crate::gpu::BatchGraphOutcome::Completed => {
6347 if !pairs.is_empty() {
6348 self.mtp_graph_mode = Some(true);
6349 }
6350 Ok(())
6351 }
6352 crate::gpu::BatchGraphOutcome::Declined => {
6353 if self.mtp_warm_graph_fallback(m, pairs, first_pos) {
6354 Ok(())
6355 } else {
6356 Err("MTP warm-up fallback failed after device admission")
6357 }
6358 }
6359 crate::gpu::BatchGraphOutcome::Failed => {
6360 Err("MTP warm batch graph failed after admission")
6361 }
6362 }
6363 }
6364
6365 #[cfg(not(feature = "gpu"))]
6366 fn mtp_warm_prefill_pairs(
6367 &mut self,
6368 m: &mut MtpModule,
6369 pairs: &[(&[f32], u32)],
6370 first_pos: usize,
6371 ) -> Result<(), &'static str> {
6372 for (j, (h, t)) in pairs.iter().enumerate() {
6373 self.mtp_warm(m, h, *t, first_pos + j);
6374 }
6375 Ok(())
6376 }
6377
6378 fn mtp_warm(&mut self, m: &mut MtpModule, hidden: &[f32], next_token: u32, position: usize) {
6382 let e = self.embed_single(next_token);
6383 let mut cat = vec![0.0f32; 2 * self.hidden_size];
6384 let (cat_e, cat_h) = cat.split_at_mut(self.hidden_size);
6385 inference::rms_norm_into(&e, &m.enorm, self.rms_eps, self.norm_style, cat_e);
6386 inference::rms_norm_into(hidden, &m.hnorm, self.rms_eps, self.norm_style, cat_h);
6387 let mut x = vec![0.0f32; self.hidden_size];
6388 m.eh_proj.matvec(&cat, &mut x, self.pool.as_deref());
6389 inference::rms_norm_into(
6390 &x,
6391 &m.layer.input_norm,
6392 self.rms_eps,
6393 self.norm_style,
6394 &mut self.ws.n1,
6395 );
6396 let attn = match &m.layer.attn {
6397 AttnKind::Full {
6398 wq,
6399 wk,
6400 wv,
6401 wo,
6402 q_norm,
6403 k_norm,
6404 output_gate,
6405 softplus_gate,
6406 bias,
6407 } => {
6408 let mut cfg = self.attn_cfg(position);
6409 cfg.q_norm = q_norm.as_deref();
6410 cfg.k_norm = k_norm.as_deref();
6411 cfg.output_gate = *output_gate;
6412 cfg.softplus_gate = softplus_gate.as_ref().map(|(g, p)| (g, *p));
6413 cfg.bias = bias
6414 .as_ref()
6415 .map(|(q, k, v)| (q.as_slice(), k.as_slice(), v.as_slice()));
6416 attention::qwen_attention(&self.ws.n1, wq, wk, wv, wo, &mut m.kv, &cfg)
6417 }
6418 _ => return,
6419 };
6420 let _ = attn;
6421 }
6422
6423 #[cfg(feature = "gpu")]
6430 #[allow(clippy::too_many_arguments)]
6431 fn graph_spec_step(
6432 &mut self,
6433 m: &mut MtpModule,
6434 hidden: &[f32],
6435 t_next: u32,
6436 next_pos: usize,
6437 drafted: &mut usize,
6438 accepted: &mut usize,
6439 all_ids: &mut Vec<u32>,
6443 room: usize,
6448 ) -> Option<(Vec<u32>, usize, Vec<f32>)> {
6449 #[cfg(target_os = "macos")]
6460 let metal_native = crate::gpu::q1_force();
6461 #[cfg(not(target_os = "macos"))]
6462 let metal_native = false;
6463 #[cfg(feature = "gpu")]
6464 let k_default = if metal_native {
6465 7
6468 } else if crate::gpu_wgpu::verify_i8_on() {
6469 5
6470 } else {
6471 4
6472 };
6473 #[cfg(not(feature = "gpu"))]
6474 let k_default = 4;
6475 let k_env: Option<usize> = std::env::var("CMF_GRAPH_SPEC_K")
6476 .ok()
6477 .and_then(|v| v.parse().ok())
6478 .filter(|&v| (1..=8).contains(&v));
6479 let (k_start, k_max) = if metal_native { (7, 7) } else { (3, k_default.max(5)) };
6484 let k_full: usize = k_env.unwrap_or_else(|| self.spec_k_adapt.unwrap_or(k_start));
6485 let k_spec = k_full.min(room).max(1);
6486 let k_capped = k_spec < k_full;
6489 if next_pos == 0 {
6490 return None;
6491 }
6492 let t_round = std::time::Instant::now();
6493 let subs = || crate::gpu_wgpu::SUBMITS.load(std::sync::atomic::Ordering::Relaxed);
6509 let sub0 = subs();
6510 let cfg = self.sampler_config.clone();
6515 let penalized = !(cfg.repetition_penalty == 1.0
6516 && cfg.presence_penalty == 0.0
6517 && cfg.suppress_tokens.is_empty());
6518 let greedy_pen = cfg.temperature < 1e-6 && penalized;
6523 let sampling = cfg.temperature >= 1e-6;
6524 let sparse = sampling && sampler::sparse_ok(&cfg);
6530 let base_len = all_ids.len();
6531 if sampling && !sparse && self.spec_q.len() < k_spec {
6532 self.spec_q.resize_with(k_spec, Vec::new);
6533 }
6534 if sparse && self.spec_qs.len() < k_spec {
6535 self.spec_qs.resize_with(k_spec, Vec::new);
6536 }
6537 let mut drafts = Vec::with_capacity(k_spec);
6542 let mut hx = hidden.to_vec();
6543 let spec_dbg = std::env::var("CMF_SPEC_DBG").is_ok();
6546 spec_stamp("pro");
6547 #[cfg(target_os = "macos")]
6553 if metal_native && !sampling && !greedy_pen && self.mtp_graph_mode != Some(false) {
6554 match self.mtp_draft_chain_metal(m, hidden, t_next, next_pos - 1, k_spec) {
6555 Ok(ids) => {
6556 self.mtp_graph_mode = Some(true);
6557 drafts = ids;
6558 }
6559 Err(true) => {
6560 tracing::error!("mtp Metal draft chain failed after commit");
6561 self.clear_sequence_state();
6562 self.graph_failed
6563 .store(true, std::sync::atomic::Ordering::Relaxed);
6564 self.cancel
6565 .store(true, std::sync::atomic::Ordering::Relaxed);
6566 return None;
6567 }
6568 Err(false) => {}
6569 }
6570 }
6571 for j in drafts.len()..k_spec {
6572 let tok_in = if j == 0 { t_next } else { drafts[j - 1] };
6573 let mut dbg_ref: Option<(Vec<f32>, Vec<f32>)> = None;
6574 if spec_dbg {
6575 let saved = self.mtp_graph_mode;
6576 self.mtp_graph_mode = Some(false);
6577 let r = self.mtp_step_hl(m, &hx, tok_in, next_pos - 1 + j);
6578 self.mtp_graph_mode = saved;
6579 if self.graph_failed.load(std::sync::atomic::Ordering::Relaxed) {
6580 return None;
6581 }
6582 m.kv.truncate_last(1);
6583 dbg_ref = Some(r);
6584 }
6585 let (mut lg, hj) = self.mtp_step_hl(m, &hx, tok_in, next_pos - 1 + j);
6586 if self.graph_failed.load(std::sync::atomic::Ordering::Relaxed) {
6587 return None;
6588 }
6589 if let Some((lg_cpu, h_cpu)) = dbg_ref {
6590 let n = |v: &[f32]| v.iter().map(|x| x * x).sum::<f32>().sqrt();
6591 let dl = lg
6592 .iter()
6593 .zip(&lg_cpu)
6594 .fold(0f32, |m, (a, b)| m.max((a - b).abs()));
6595 let dh = hj
6596 .iter()
6597 .zip(&h_cpu)
6598 .fold(0f32, |m, (a, b)| m.max((a - b).abs()));
6599 eprintln!(
6600 "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 {}",
6601 next_pos - 1 + j,
6602 sampler::argmax(&lg_cpu),
6603 sampler::argmax(&lg),
6604 n(&h_cpu),
6605 n(&hj),
6606 m.kv.seq_len
6607 );
6608 }
6609 let dj = if sparse {
6610 let mut q = std::mem::take(&mut self.spec_qs[j]);
6611 let ok = sampler::sparse_distribution_into(
6612 &lg,
6613 &cfg,
6614 all_ids,
6615 &mut self.sampler_scratch,
6616 self.pool.as_deref(),
6617 &mut q,
6618 );
6619 let d = if ok {
6620 sampler::draw_sparse(&q, &mut self.rng)
6621 } else {
6622 let t = sampler::argmax(&lg);
6624 q.clear();
6625 q.push((t, 1.0));
6626 t
6627 };
6628 self.spec_qs[j] = q;
6629 all_ids.push(d);
6630 d
6631 } else if sampling {
6632 let mut q = std::mem::take(&mut self.spec_q[j]);
6633 sampler::distribution_into(
6634 &lg,
6635 &cfg,
6636 all_ids,
6637 &mut self.sampler_scratch,
6638 self.pool.as_deref(),
6639 &mut q,
6640 );
6641 let d = sampler::draw(&q, &mut self.rng);
6642 self.spec_q[j] = q;
6643 all_ids.push(d); d
6645 } else if greedy_pen {
6646 let d = sampler::argmax_penalized(
6647 &lg,
6648 &cfg,
6649 all_ids,
6650 &mut self.sampler_scratch,
6651 self.pool.as_deref(),
6652 );
6653 all_ids.push(d);
6654 d
6655 } else {
6656 sampler::argmax(&lg)
6657 };
6658 attention::recycle_buf(&mut lg);
6659 drafts.push(dj);
6660 hx = hj;
6661 spec_stamp("d.pick");
6662 }
6663 all_ids.truncate(base_len);
6664 *drafted += k_spec;
6665 let t_draft = t_round.elapsed();
6666 let sub_draft = subs();
6667 let b = k_spec + 1;
6670 let mut hiddens = vec![0.0f32; b * self.hidden_size];
6671 for (i, &t) in std::iter::once(&t_next).chain(drafts.iter()).enumerate() {
6672 let e = self.embed_single(t);
6673 hiddens[i * self.hidden_size..(i + 1) * self.hidden_size].copy_from_slice(&e);
6674 }
6675 let positions: Vec<usize> = (next_pos..next_pos + b).collect();
6676 spec_stamp("v.emb");
6677 let (lm_gw, lm_rows) = {
6678 let (_, i, kind, rs) = self.weights.lm_head.graph_weight()?;
6679 (
6680 crate::gpu::GraphW {
6681 idx: i,
6682 kind,
6683 row_scale: rs,
6684 data: &[],
6685 prism: crate::gpu::GraphPrismOp::None,
6686 affine: false,
6687 },
6688 self.weights.lm_head.rows(),
6689 )
6690 };
6691 let mut logits = Vec::new();
6692 let final_norm = self.weights.final_norm.clone();
6693 #[cfg(target_os = "macos")]
6702 let greedy_dev = metal_native
6703 && !sampling
6704 && !greedy_pen
6705 && !self.confidence_on
6706 && self.final_softcap.is_none()
6707 && self.vocab_size == lm_rows
6713 && std::env::var_os("CMF_METAL_VERIFY_CHECK").is_none()
6714 && std::env::var_os("CMF_LOGIT_DUMP").is_none()
6715 && std::env::var("CMF_METAL_DEV_ARGMAX").as_deref() != Ok("0");
6716 #[cfg(not(target_os = "macos"))]
6717 let greedy_dev = false;
6718 let mut dev_ids: Vec<u32> = Vec::new();
6719 #[cfg(target_os = "macos")]
6720 let verify_outcome = if metal_native {
6721 let lm = self.weights.lm_head.q1_parts()?;
6722 let n_score = self.vocab_size.min(lm_rows);
6723 self.try_batch_graph_metal(
6724 &mut hiddens,
6725 &positions,
6726 b,
6727 Some((lm, &final_norm, &mut logits)),
6728 if greedy_dev {
6729 Some((n_score, &mut dev_ids))
6730 } else {
6731 None
6732 },
6733 )
6734 } else {
6735 self.try_batch_graph_wgpu(
6736 &mut hiddens,
6737 &positions,
6738 b,
6739 Some(crate::gpu::SpecTail {
6740 lm: lm_gw,
6741 lm_rows,
6742 final_norm: &final_norm,
6743 logits_out: &mut logits,
6744 }),
6745 )
6746 };
6747 #[cfg(not(target_os = "macos"))]
6748 let verify_outcome = self.try_batch_graph_wgpu(
6749 &mut hiddens,
6750 &positions,
6751 b,
6752 Some(crate::gpu::SpecTail {
6753 lm: lm_gw,
6754 lm_rows,
6755 final_norm: &final_norm,
6756 logits_out: &mut logits,
6757 }),
6758 );
6759 match verify_outcome {
6760 crate::gpu::BatchGraphOutcome::Completed => {}
6761 crate::gpu::BatchGraphOutcome::Declined => {
6762 m.kv.truncate_last(k_spec);
6766 if !metal_native && !self.rewind_mtp_graph_mirror(next_pos) {
6767 self.clear_sequence_state();
6768 self.graph_failed
6769 .store(true, std::sync::atomic::Ordering::Relaxed);
6770 self.cancel
6771 .store(true, std::sync::atomic::Ordering::Relaxed);
6772 tracing::error!("MTP graph mirror rewind failed after verify decline");
6773 }
6774 return None;
6775 }
6776 crate::gpu::BatchGraphOutcome::Failed => {
6777 self.clear_sequence_state();
6781 self.graph_failed
6782 .store(true, std::sync::atomic::Ordering::Relaxed);
6783 self.cancel
6784 .store(true, std::sync::atomic::Ordering::Relaxed);
6785 tracing::error!("MTP verify batch graph failed after admission");
6786 return None;
6787 }
6788 }
6789 #[cfg(target_os = "macos")]
6795 if metal_native && std::env::var("CMF_METAL_VERIFY_CHECK").as_deref() == Ok("1") {
6796 let snap: Vec<Vec<f32>> = self
6797 .kv_cache
6798 .layers
6799 .iter()
6800 .map(|l| l.linear_state.clone())
6801 .collect();
6802 let attn_lens: Vec<usize> = self.kv_cache.layers.iter().map(|l| l.seq_len).collect();
6803 let toks: Vec<u32> = std::iter::once(t_next)
6804 .chain(drafts.iter().copied())
6805 .collect();
6806 let want_save = self.graph_want_logits;
6807 self.graph_want_logits = false;
6808 for (i, &t) in toks.iter().enumerate() {
6809 let hi = self.forward_layers(&self.embed_single(t), next_pos + i, None);
6810 let _ = self.graph_logits.take();
6811 if std::env::var("CMF_SPEC_PLAIN_HIDDEN").as_deref() == Ok("1") {
6815 hiddens[i * self.hidden_size..(i + 1) * self.hidden_size].copy_from_slice(&hi);
6816 }
6817 let ref_lg = self.logits_from_hidden(&hi);
6818 let row = &logits[i * lm_rows..(i + 1) * lm_rows];
6819 let ra = sampler::argmax(&ref_lg);
6820 let va = sampler::argmax(row);
6821 let mut md = 0f32;
6822 let mut rms = 0f64;
6823 for j in 0..lm_rows.min(ref_lg.len()) {
6824 let d = (ref_lg[j] - row[j]).abs();
6825 md = md.max(d);
6826 rms += (d as f64) * (d as f64);
6827 }
6828 let mut hd = 0f32;
6829 for j in 0..self.hidden_size {
6830 hd = hd.max((hi[j] - hiddens[i * self.hidden_size + j]).abs());
6831 }
6832 eprintln!(
6833 "verify-check row {i} tok {t} pos {}: ref argmax {ra} verify argmax {va} {} | max|dlogit| {md:.3} rms {:.4} | max|dhidden| {hd:.4}",
6834 next_pos + i,
6835 if ra == va { "OK" } else { "MISMATCH" },
6836 (rms / lm_rows as f64).sqrt()
6837 );
6838 }
6839 self.graph_want_logits = want_save;
6840 for (l, st) in self.kv_cache.layers.iter_mut().zip(snap) {
6843 if l.linear_state.len() == st.len() {
6844 l.linear_state.copy_from_slice(&st);
6845 } else {
6846 l.linear_state = st;
6847 }
6848 }
6849 for (li, (l, n0)) in self.kv_cache.layers.iter_mut().zip(attn_lens).enumerate() {
6850 let extra = l.seq_len.saturating_sub(n0);
6851 if extra > 0 {
6852 l.truncate_last(extra);
6853 crate::gpu_metal::kv_mirror_set_stored(self.graph_kv_id, li, n0);
6854 }
6855 }
6856 }
6857 let t_verify = t_round.elapsed();
6858 let sub_verify = subs();
6859 let mut a = 0usize;
6864 let mut forced: Option<u32> = None;
6865 let ids: Vec<u32> = if sparse {
6866 let mut p = std::mem::take(&mut self.spec_ps);
6867 let mut res = std::mem::take(&mut self.spec_ress);
6868 while a < k_spec {
6869 let ok = sampler::sparse_distribution_into(
6870 &logits[a * lm_rows..(a + 1) * lm_rows],
6871 &cfg,
6872 all_ids,
6873 &mut self.sampler_scratch,
6874 self.pool.as_deref(),
6875 &mut p,
6876 );
6877 if !ok {
6878 let t = sampler::argmax(&logits[a * lm_rows..(a + 1) * lm_rows]);
6879 p.clear();
6880 p.push((t, 1.0));
6881 }
6882 match sampler::spec_accept_or_correct_sparse(
6883 &p,
6884 &self.spec_qs[a],
6885 drafts[a],
6886 &mut self.rng,
6887 &mut res,
6888 ) {
6889 None => {
6890 all_ids.push(drafts[a]);
6891 a += 1;
6892 }
6893 Some(c) => {
6894 forced = Some(c);
6895 break;
6896 }
6897 }
6898 }
6899 all_ids.truncate(base_len);
6900 self.spec_ps = p;
6901 self.spec_ress = res;
6902 drafts.clone()
6903 } else if sampling {
6904 let mut p = std::mem::take(&mut self.spec_p);
6905 let mut res = std::mem::take(&mut self.spec_res);
6906 while a < k_spec {
6907 sampler::distribution_into(
6908 &logits[a * lm_rows..(a + 1) * lm_rows],
6909 &cfg,
6910 all_ids,
6911 &mut self.sampler_scratch,
6912 self.pool.as_deref(),
6913 &mut p,
6914 );
6915 match sampler::spec_accept_or_correct(
6916 &p,
6917 &self.spec_q[a],
6918 drafts[a],
6919 &mut self.rng,
6920 &mut res,
6921 self.pool.as_deref(),
6922 ) {
6923 None => {
6924 all_ids.push(drafts[a]);
6925 a += 1;
6926 }
6927 Some(c) => {
6928 forced = Some(c);
6929 break;
6930 }
6931 }
6932 }
6933 all_ids.truncate(base_len);
6934 self.spec_p = p;
6935 self.spec_res = res;
6936 drafts.clone()
6938 } else if greedy_pen {
6939 let mut ids: Vec<u32> = Vec::with_capacity(b);
6943 for i in 0..b {
6944 let t = sampler::argmax_penalized(
6945 &logits[i * lm_rows..(i + 1) * lm_rows],
6946 &cfg,
6947 all_ids,
6948 &mut self.sampler_scratch,
6949 self.pool.as_deref(),
6950 );
6951 ids.push(t);
6952 if i < k_spec && t == drafts[i] {
6953 all_ids.push(t);
6954 } else {
6955 break;
6956 }
6957 }
6958 all_ids.truncate(base_len);
6959 while a < k_spec && a < ids.len() && ids[a] == drafts[a] {
6960 a += 1;
6961 }
6962 ids
6965 } else if greedy_dev && dev_ids.len() == b {
6966 let ids = std::mem::take(&mut dev_ids);
6967 while a < k_spec && ids[a] == drafts[a] {
6968 a += 1;
6969 }
6970 ids
6971 } else {
6972 if logits.len() < b * lm_rows {
6973 self.clear_sequence_state();
6976 self.graph_failed
6977 .store(true, std::sync::atomic::Ordering::Relaxed);
6978 self.cancel
6979 .store(true, std::sync::atomic::Ordering::Relaxed);
6980 tracing::error!("Metal verify returned neither logits nor argmax ids");
6981 return None;
6982 }
6983 let ids: Vec<u32> = (0..b)
6984 .map(|i| sampler::argmax(&logits[i * lm_rows..(i + 1) * lm_rows]))
6985 .collect();
6986 while a < k_spec && ids[a] == drafts[a] {
6987 a += 1;
6988 }
6989 ids
6990 };
6991 spec_stamp("acc");
6992 if spec_dbg {
6993 eprintln!(
6994 "spec-dbg round: t_next {t_next} drafts {:?} verified {:?} accepted {a}",
6995 drafts, ids
6996 );
6997 }
6998 #[cfg(target_os = "macos")]
7002 let commit_ref: Option<(Vec<Vec<f32>>, Vec<(usize, Vec<f32>, Vec<f32>)>)> = if metal_native
7003 && std::env::var("CMF_METAL_VERIFY_CHECK").as_deref() == Ok("2")
7004 {
7005 let snap: Vec<Vec<f32>> = self
7006 .kv_cache
7007 .layers
7008 .iter()
7009 .map(|l| l.linear_state.clone())
7010 .collect();
7011 let attn_lens: Vec<usize> = self.kv_cache.layers.iter().map(|l| l.seq_len).collect();
7012 let toks: Vec<u32> = std::iter::once(t_next)
7013 .chain(drafts.iter().copied())
7014 .collect();
7015 let want_save = self.graph_want_logits;
7016 self.graph_want_logits = false;
7017 for (i, &t) in toks.iter().take(a + 1).enumerate() {
7018 let _ = self.forward_layers(&self.embed_single(t), next_pos + i, None);
7019 let _ = self.graph_logits.take();
7020 }
7021 self.graph_want_logits = want_save;
7022 let plain_states: Vec<Vec<f32>> = self
7023 .kv_cache
7024 .layers
7025 .iter()
7026 .map(|l| l.linear_state.clone())
7027 .collect();
7028 let (nkv, hd) = (self.num_kv_heads, self.head_dim);
7029 let mut rows = Vec::new();
7030 for (li, (l, n0)) in self
7031 .kv_cache
7032 .layers
7033 .iter_mut()
7034 .zip(attn_lens.iter())
7035 .enumerate()
7036 {
7037 let extra = l.seq_len.saturating_sub(*n0);
7038 if extra > 0 {
7039 let mut kk = Vec::new();
7040 let mut vv = Vec::new();
7041 for g in 0..nkv {
7042 kk.extend_from_slice(&l.head_keys(g)[n0 * hd..]);
7043 vv.extend_from_slice(&l.head_values(g)[n0 * hd..]);
7044 }
7045 rows.push((li, kk, vv));
7046 l.truncate_last(extra);
7047 crate::gpu_metal::kv_mirror_set_stored(self.graph_kv_id, li, *n0);
7048 }
7049 }
7050 for (l, st) in self.kv_cache.layers.iter_mut().zip(snap) {
7051 if l.linear_state.len() == st.len() {
7052 l.linear_state.copy_from_slice(&st);
7053 } else {
7054 l.linear_state = st;
7055 }
7056 }
7057 Some((plain_states, rows))
7058 } else {
7059 None
7060 };
7061 let warm_off = std::env::var("CMF_SPEC_WARM").is_ok_and(|v| v == "0");
7062 #[cfg(target_os = "macos")]
7071 let mut warm_pending: Option<MetalWarmPending> = None;
7072 #[cfg(target_os = "macos")]
7073 if metal_native {
7074 m.kv.truncate_last(k_spec.saturating_sub(1));
7075 if self.mtp_graph_mode == Some(true) {
7076 crate::gpu_metal::kv_mirror_set_stored(
7079 self.mtp_kv_id(),
7080 Self::MTP_LAYER_BASE,
7081 m.kv.seq_len,
7082 );
7083 if !warm_off && a > 0 {
7084 let pairs: Vec<(&[f32], u32)> = (0..a)
7085 .map(|j| {
7086 (
7087 &hiddens[j * self.hidden_size..(j + 1) * self.hidden_size],
7088 ids[j],
7089 )
7090 })
7091 .collect();
7092 warm_pending = self.mtp_warm_batch_submit(m, &pairs, next_pos);
7093 }
7094 }
7095 spec_stamp("c.wsub");
7096 }
7097 #[cfg(target_os = "macos")]
7099 if metal_native {
7100 if !self.metal_verify_commit(a) {
7103 self.clear_sequence_state();
7104 self.graph_failed
7105 .store(true, std::sync::atomic::Ordering::Relaxed);
7106 self.cancel
7107 .store(true, std::sync::atomic::Ordering::Relaxed);
7108 tracing::error!("Metal verify state/KV handoff failed after admission");
7109 return None;
7110 }
7111 if let Some((plain_states, rows)) = commit_ref {
7112 crate::gpu_metal::queue_fence();
7113 let _ = crate::gpu_metal::wait_replay();
7116 let (nkv, hd) = (self.num_kv_heads, self.head_dim);
7117 let mut worst_s = 0f32;
7118 let mut worst_li = 0usize;
7119 for (li, (l, ps)) in self.kv_cache.layers.iter().zip(&plain_states).enumerate() {
7120 if l.linear_state.len() != ps.len() || ps.is_empty() {
7121 continue;
7122 }
7123 let d = l
7124 .linear_state
7125 .iter()
7126 .zip(ps)
7127 .fold(0f32, |m, (x, y)| m.max((x - y).abs()));
7128 let n = ps.iter().fold(0f32, |m, y| m.max(y.abs()));
7129 let rel = d / n.max(1e-6);
7130 if rel > worst_s {
7131 worst_s = rel;
7132 worst_li = li;
7133 }
7134 }
7135 let mut worst_k = 0f32;
7136 for (li, kk, vv) in &rows {
7137 let l = &self.kv_cache.layers[*li];
7138 let n0 = l.seq_len - (kk.len() / (nkv * hd));
7139 let mut ck = Vec::new();
7140 let mut cv = Vec::new();
7141 for g in 0..nkv {
7142 ck.extend_from_slice(&l.head_keys(g)[n0 * hd..]);
7143 cv.extend_from_slice(&l.head_values(g)[n0 * hd..]);
7144 }
7145 if ck.len() == kk.len() {
7146 let dk = ck
7147 .iter()
7148 .zip(kk)
7149 .fold(0f32, |m, (x, y)| m.max((x - y).abs()));
7150 let dv = cv
7151 .iter()
7152 .zip(vv)
7153 .fold(0f32, |m, (x, y)| m.max((x - y).abs()));
7154 worst_k = worst_k.max(dk).max(dv);
7155 } else {
7156 eprintln!(
7157 "commit-check L{li}: kv row count mismatch {} vs {}",
7158 ck.len(),
7159 kk.len()
7160 );
7161 }
7162 }
7163 eprintln!(
7164 "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}"
7165 );
7166 }
7167 }
7168 if !metal_native && a + 1 < b {
7169 let expected_gdn_layers = self.graph_gdn_layer_count();
7170 if expected_gdn_layers > 0
7171 && !crate::gpu::gdn_spec_restore(self.graph_kv_id, a, next_pos, expected_gdn_layers)
7172 {
7173 self.clear_sequence_state();
7174 self.graph_failed
7175 .store(true, std::sync::atomic::Ordering::Relaxed);
7176 self.cancel
7177 .store(true, std::sync::atomic::Ordering::Relaxed);
7178 tracing::error!("GDN speculative restore failed after verify");
7179 return None;
7180 }
7181 }
7182 if !metal_native && !self.rewind_trunk_graph_mirrors(next_pos + a + 1) {
7183 self.clear_sequence_state();
7188 self.graph_failed
7189 .store(true, std::sync::atomic::Ordering::Relaxed);
7190 self.cancel
7191 .store(true, std::sync::atomic::Ordering::Relaxed);
7192 tracing::error!("trunk graph KV rewind failed after speculative verify");
7193 return None;
7194 }
7195 *accepted += a;
7196 if !metal_native {
7207 m.kv.truncate_last(k_spec.saturating_sub(1));
7209 }
7210 spec_stamp("c.trunc");
7211 if !metal_native
7212 && self.mtp_graph_mode == Some(true)
7213 && !self.rewind_mtp_graph_mirror(next_pos)
7214 {
7215 self.clear_sequence_state();
7219 self.graph_failed
7220 .store(true, std::sync::atomic::Ordering::Relaxed);
7221 self.cancel
7222 .store(true, std::sync::atomic::Ordering::Relaxed);
7223 tracing::error!("MTP graph mirror rewind failed after verify commit");
7224 return None;
7225 }
7226 if !warm_off && a > 0 {
7227 let mut warmed = false;
7230 #[cfg(target_os = "macos")]
7231 if metal_native && self.mtp_graph_mode == Some(true) {
7232 warmed = match warm_pending.take() {
7236 Some(p) => self.mtp_warm_batch_finish(m, p),
7237 None => false,
7238 };
7239 if !warmed {
7240 warmed = true;
7241 for j in 0..a {
7242 let row =
7243 hiddens[j * self.hidden_size..(j + 1) * self.hidden_size].to_vec();
7244 if self
7245 .mtp_step_metal(m, &row, ids[j], next_pos + j, false)
7246 .is_none()
7247 {
7248 warmed = false;
7249 break;
7250 }
7251 }
7252 }
7253 }
7254 if !warmed && self.mtp_graph_mode != Some(false) && !metal_native {
7255 let rows: Vec<Vec<f32>> = (0..a)
7256 .map(|j| hiddens[j * self.hidden_size..(j + 1) * self.hidden_size].to_vec())
7257 .collect();
7258 let pairs: Vec<(&[f32], u32)> = rows
7259 .iter()
7260 .zip(ids.iter())
7261 .map(|(r, &t)| (r.as_slice(), t))
7262 .collect();
7263 match self.mtp_warm_prefill_pairs(m, &pairs, next_pos) {
7264 Ok(()) => warmed = true,
7265 Err(err) => {
7266 tracing::error!("{err}");
7272 self.clear_sequence_state();
7273 self.graph_failed
7274 .store(true, std::sync::atomic::Ordering::Relaxed);
7275 self.cancel
7276 .store(true, std::sync::atomic::Ordering::Relaxed);
7277 return None;
7278 }
7279 }
7280 }
7281 if !warmed {
7282 for j in 0..a {
7283 let row = &hiddens[j * self.hidden_size..(j + 1) * self.hidden_size];
7284 let row = row.to_vec();
7285 self.mtp_warm(m, &row, ids[j], next_pos + j);
7286 }
7287 }
7288 }
7289 spec_stamp("c.warm");
7293 if let Some(c) = forced {
7294 self.spec_forced = Some(c);
7295 self.graph_logits = None;
7296 } else if greedy_dev && logits.is_empty() {
7297 self.spec_forced = Some(ids[a]);
7300 self.graph_logits = None;
7301 } else {
7302 let mut row = logits[a * lm_rows..(a + 1) * lm_rows].to_vec();
7303 row.resize(self.vocab_size, 0.0);
7304 if let Some(c) = self.final_softcap {
7305 for l in row.iter_mut() {
7306 *l = c * (*l / c).tanh();
7307 }
7308 }
7309 self.graph_logits = Some(row);
7310 }
7311 let new_hidden = hiddens[a * self.hidden_size..(a + 1) * self.hidden_size].to_vec();
7312 spec_stamp("c.row");
7313 if std::env::var("CMF_GRAPH_SPEC_TIME").is_ok() {
7319 let end = subs();
7320 eprintln!(
7321 "spec-round: draft {:.1} ms/{} sub | verify {:.1} ms/{} sub | \
7322 commit {:.1} ms/{} sub (accepted {a} of {k_spec}, full-head streak {})",
7323 t_draft.as_secs_f64() * 1e3,
7324 sub_draft - sub0,
7325 (t_verify - t_draft).as_secs_f64() * 1e3,
7326 sub_verify - sub_draft,
7327 (t_round.elapsed() - t_verify).as_secs_f64() * 1e3,
7328 end - sub_verify,
7329 self.draft_full_streak,
7330 );
7331 }
7332 if k_env.is_none() && !metal_native && !k_capped {
7337 let f = a as f32 / k_spec.max(1) as f32;
7341 self.spec_acc_ewma += 0.2 * (f - self.spec_acc_ewma);
7342 let mut k_next = k_spec;
7343 if self.spec_acc_ewma >= 0.75 && k_spec < k_max {
7344 k_next = k_spec + 1;
7345 } else if self.spec_acc_ewma < 0.4 && k_spec > 2 {
7346 k_next = k_spec - 1;
7347 }
7348 if k_next != k_spec {
7349 self.spec_acc_ewma = 0.6;
7350 if std::env::var("CMF_GRAPH_SPEC_TIME").is_ok() {
7351 eprintln!("spec-k: {k_spec} → {k_next}");
7352 }
7353 }
7354 self.spec_k_adapt = Some(k_next);
7355 }
7356 spec_stamp("end");
7357 Some((drafts[..a].to_vec(), next_pos + a + 1, new_hidden))
7358 }
7359
7360 pub fn measure_pair_fusion(&mut self, iters: usize) -> (f64, f64) {
7369 if !self.pair_supported() {
7370 return (0.0, 0.0);
7371 }
7372 let graph_env = std::env::var_os("CMF_GPU_WGPU_GRAPH");
7379 unsafe { std::env::set_var("CMF_GPU_WGPU_GRAPH", "0") };
7380 let emb1 = self.embed_single(1);
7381 let emb2 = self.embed_single(2);
7382 let pos = self.kv_cache.seq_len();
7383
7384 let t0 = std::time::Instant::now();
7385 for _ in 0..iters {
7386 let _ = self.forward_layers(&emb1, pos, None);
7387 let _ = self.forward_layers(&emb2, pos + 1, None);
7388 for l in &mut self.kv_cache.layers {
7389 l.truncate_last(2);
7390 }
7391 }
7392 let singles_ms = t0.elapsed().as_secs_f64() * 1000.0 / iters as f64;
7393
7394 let t1 = std::time::Instant::now();
7395 for _ in 0..iters {
7396 let _ = self.forward_pair(&emb1, &emb2, pos);
7397 for l in &mut self.kv_cache.layers {
7398 l.truncate_last(2);
7399 }
7400 }
7401 let pair_ms = t1.elapsed().as_secs_f64() * 1000.0 / iters as f64;
7402 match graph_env {
7403 Some(value) => unsafe { std::env::set_var("CMF_GPU_WGPU_GRAPH", value) },
7404 None => unsafe { std::env::remove_var("CMF_GPU_WGPU_GRAPH") },
7405 }
7406 (singles_ms, pair_ms)
7407 }
7408
7409 fn pair_supported(&self) -> bool {
7417 !self.weights.layers.is_empty()
7424 && self.g3n.is_none()
7425 && !self
7426 .weights
7427 .layers
7428 .iter()
7429 .any(|lw| matches!(&lw.attn, AttnKind::Mla(_) | AttnKind::Kda(_)))
7430 }
7431
7432 fn forward_pair(
7433 &mut self,
7434 emb1: &[f32],
7435 emb2: &[f32],
7436 position: usize,
7437 ) -> (Vec<f32>, Vec<f32>) {
7438 self.mimo_moe_prepare();
7441 let mut h1 = emb1.to_vec();
7442 let mut h2 = emb2.to_vec();
7443 let (_nkv, _hd, hs, _rd, eps) = (
7444 self.num_kv_heads,
7445 self.head_dim,
7446 self.hidden_size,
7447 self.rotary_dim,
7448 self.rms_eps,
7449 );
7450 let pool = self.pool.clone();
7451
7452 for li in 0..self.num_layers {
7453 let lw = &self.weights.layers[self.phys_layer(li)];
7454 inference::rms_norm_into(
7457 &h1,
7458 &lw.input_norm,
7459 self.rms_eps,
7460 self.norm_style,
7461 &mut self.ws.n1,
7462 );
7463 inference::rms_norm_into(
7464 &h2,
7465 &lw.input_norm,
7466 self.rms_eps,
7467 self.norm_style,
7468 &mut self.ws.n2,
7469 );
7470
7471 let (a1, a2) = match &lw.attn {
7472 AttnKind::Mla(_) => unreachable!("MLA has no MTP/pair path"),
7473 AttnKind::Kda(_) => unreachable!("KDA has no MTP/pair path"),
7474 AttnKind::Bounded(w) => {
7475 let rope = self
7478 .bounded_rope
7479 .clone()
7480 .expect("bounded layer without an installed rotation table");
7481 let cfg = crate::bounded::BoundedAttnCfg {
7482 num_heads: self.num_heads,
7483 num_kv_heads: self.num_kv_heads,
7484 head_dim: self.head_dim,
7485 hidden_size: hs,
7486 scale: self.attn_scale,
7487 rope: &rope,
7488 pool: pool.as_deref(),
7489 };
7490 let a1 = crate::bounded::bounded_attention(
7491 &self.ws.n1,
7492 w,
7493 &mut self.kv_cache.layers[li],
7494 &cfg,
7495 );
7496 let a2 = crate::bounded::bounded_attention(
7497 &self.ws.n2,
7498 w,
7499 &mut self.kv_cache.layers[li],
7500 &cfg,
7501 );
7502 (a1, a2)
7503 }
7504 AttnKind::Linear(w) => {
7505 let cfg = self.vmf_cfg.expect("linear layer without vmf_cfg");
7506 let layer = &mut self.kv_cache.layers[li];
7507 let (state, scratch) = (&mut layer.linear_state, &mut layer.linear_scratch);
7508 vmf_phase_pair(
7509 &self.ws.n1,
7510 &self.ws.n2,
7511 w,
7512 &cfg,
7513 state,
7514 scratch,
7515 self.pool.as_deref(),
7516 )
7517 }
7518 AttnKind::LinearGdn(w) => {
7519 let cfg = self.gdn_cfg.expect("gdn layer without gdn_cfg");
7520 let layer = &mut self.kv_cache.layers[li];
7521 let (state, scratch) = (&mut layer.linear_state, &mut layer.linear_scratch);
7522 gdn_pair(
7523 &self.ws.n1,
7524 &self.ws.n2,
7525 w,
7526 &cfg,
7527 state,
7528 scratch,
7529 self.pool.as_deref(),
7530 )
7531 }
7532 AttnKind::ShortConv(w) => {
7533 let cfg = self
7534 .short_conv_cfg
7535 .expect("short-conv layer without short_conv_cfg");
7536 let layer = &mut self.kv_cache.layers[li];
7537 let (state, scratch) = (&mut layer.linear_state, &mut layer.linear_scratch);
7538 short_conv_pair(
7539 &self.ws.n1,
7540 &self.ws.n2,
7541 w,
7542 &cfg,
7543 state,
7544 scratch,
7545 self.pool.as_deref(),
7546 )
7547 }
7548 AttnKind::Full {
7549 wq,
7550 wk,
7551 wv,
7552 wo,
7553 q_norm,
7554 k_norm,
7555 output_gate,
7556 softplus_gate,
7557 bias,
7558 } => {
7559 let inv_freq_l = self.layer_inv_freq(li);
7560 let (nkv_l, hd_l, rd_l) = self.layer_geom(li);
7561 let cfg = QwenAttnCfg {
7562 num_heads: self.layer_num_heads(li),
7563 num_kv_heads: nkv_l,
7564 head_dim: hd_l,
7565 hidden_size: hs,
7566 position,
7567 inv_freq: &inv_freq_l,
7568 rotary_dim: rd_l,
7569 scale: self.attn_scale,
7570 softcap: self.attn_softcap,
7571 window: self.layer_window(li),
7572 v_norm: self.attn_v_norm,
7573 qk_norm_after_rope: self.qk_norm_after_rope,
7574 q_norm: q_norm.as_deref(),
7575 k_norm: k_norm.as_deref(),
7576 output_gate: *output_gate,
7577 softplus_gate: softplus_gate
7578 .as_ref()
7579 .map(|(gate, per_head)| (gate, *per_head)),
7580 rope_scale: self.layer_rope_scale(li),
7581 bias: bias
7582 .as_ref()
7583 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
7584 rms_eps: eps,
7585 norm_style: self.norm_style,
7586 pool: pool.as_deref(),
7587 v_head_dim: self.layer_v_dim(li),
7588 };
7589 attention::qwen_attention_pair(
7590 &self.ws.n1,
7591 &self.ws.n2,
7592 wq,
7593 wk,
7594 wv,
7595 wo,
7596 &mut self.kv_cache.layers[li],
7597 &cfg,
7598 )
7599 }
7600 };
7601 let (a1, a2) = match &self.weights.layers[self.phys_layer(li)].attn_out_norm {
7602 Some(w) => (
7603 inference::rms_norm(&a1, w, self.rms_eps, self.norm_style),
7604 inference::rms_norm(&a2, w, self.rms_eps, self.norm_style),
7605 ),
7606 None => (a1, a2),
7607 };
7608 for i in 0..self.hidden_size {
7609 h1[i] += a1[i];
7610 h2[i] += a2[i];
7611 }
7612 let (mut a1, mut a2) = (a1, a2);
7613 attention::recycle_buf(&mut a1);
7614 attention::recycle_buf(&mut a2);
7615
7616 let lw = &self.weights.layers[self.phys_layer(li)];
7617 inference::rms_norm_into(
7618 &h1,
7619 &lw.post_norm,
7620 self.rms_eps,
7621 self.norm_style,
7622 &mut self.ws.p1,
7623 );
7624 inference::rms_norm_into(
7625 &h2,
7626 &lw.post_norm,
7627 self.rms_eps,
7628 self.norm_style,
7629 &mut self.ws.p2,
7630 );
7631 let (f1, f2) = match &lw.ffn {
7632 FfnKind::DenseMoe(dm) => (
7635 dense_moe_ffn(
7636 dm,
7637 &self.ws.p1,
7638 &h1,
7639 self.rms_eps,
7640 self.norm_style,
7641 self.pool.as_deref(),
7642 ),
7643 dense_moe_ffn(
7644 dm,
7645 &self.ws.p2,
7646 &h2,
7647 self.rms_eps,
7648 self.norm_style,
7649 self.pool.as_deref(),
7650 ),
7651 ),
7652 FfnKind::Moe(m) if self.mimo_moe.is_dynamic(li, false) => (
7653 moe_ffn_banked(&mut self.mimo_moe, li, m, &self.ws.p1, self.pool.as_deref()),
7654 moe_ffn_banked(&mut self.mimo_moe, li, m, &self.ws.p2, self.pool.as_deref()),
7655 ),
7656 _ => ffn_forward_pair(
7657 &lw.ffn,
7658 &self.ws.p1,
7659 &self.ws.p2,
7660 self.pool.as_deref(),
7661 None,
7662 ),
7663 };
7664 let (f1, f2) = match &self.weights.layers[self.phys_layer(li)].ffn_out_norm {
7665 Some(w) => (
7666 inference::rms_norm(&f1, w, self.rms_eps, self.norm_style),
7667 inference::rms_norm(&f2, w, self.rms_eps, self.norm_style),
7668 ),
7669 None => (f1, f2),
7670 };
7671 for i in 0..self.hidden_size {
7672 h1[i] += f1[i];
7673 h2[i] += f2[i];
7674 }
7675 let (mut f1, mut f2) = (f1, f2);
7676 attention::recycle_buf(&mut f1);
7677 attention::recycle_buf(&mut f2);
7678 if let Some(sc) = self.weights.layers[self.phys_layer(li)].layer_scale {
7679 for i in 0..self.hidden_size {
7680 h1[i] *= sc;
7681 h2[i] *= sc;
7682 }
7683 }
7684 if self.is_loop_end(li) && li + 1 < self.num_layers {
7686 h1 = inference::rms_norm(
7687 &h1,
7688 &self.weights.final_norm,
7689 self.rms_eps,
7690 self.norm_style,
7691 );
7692 h2 = inference::rms_norm(
7693 &h2,
7694 &self.weights.final_norm,
7695 self.rms_eps,
7696 self.norm_style,
7697 );
7698 }
7699 }
7700 if self.o1_active() {
7706 self.commit_linear_scratch();
7707 }
7708 self.o1_progress();
7709 (h1, h2)
7710 }
7711
7712 fn commit_linear_scratch(&mut self) {
7714 for layer in &mut self.kv_cache.layers {
7715 if !layer.linear_scratch.is_empty() {
7716 std::mem::swap(&mut layer.linear_state, &mut layer.linear_scratch);
7717 layer.linear_scratch.clear();
7718 }
7719 }
7720 }
7721
7722 pub fn forward_ids(
7725 &mut self,
7726 ids: &[u32],
7727 task_mask: Option<&TaskMask>,
7728 ) -> Result<Vec<f32>, String> {
7729 #[cfg(target_os = "macos")]
7730 crate::gpu_metal::set_io_namespace(self.graph_kv_id);
7731 if ids.is_empty() {
7732 return Err("empty id sequence".to_string());
7733 }
7734 self.clear_sequence_state();
7735 self.check_forward_graph("forward_ids setup", 0)?;
7736 if task_mask.is_none() {
7737 self.o1_begin();
7738 }
7739 let mut hidden = vec![0.0f32; self.hidden_size];
7740 let mut pos = 0usize;
7741 if let Some(b) = &mut self.dsv41 {
7742 let pool = self.pool.clone();
7743 let mut logits = Vec::new();
7744 crate::dsv41::forward_chunk(
7745 &b.0,
7746 &b.1,
7747 &b.2,
7748 &mut b.3,
7749 ids,
7750 0,
7751 pool.as_deref(),
7752 &mut logits,
7753 );
7754 if let Err(err) = self.o1_seal_checked() {
7755 self.clear_sequence_state();
7756 return Err(err);
7757 }
7758 return Ok(logits);
7759 }
7760 if self.can_prefill_batched() && !self.graph_prefill_preferred() && ids.len() > 2 {
7768 let chunk = self.prefill_chunk();
7772 let hs = self.hidden_size;
7773 while pos < ids.len() {
7774 let end = (pos + chunk).min(ids.len());
7775 let hb = self.prefill_rows(&ids[pos..end], pos, task_mask)?;
7776 self.check_forward_graph("forward_ids batched prefill", end - 1)?;
7777 hidden.copy_from_slice(&hb[(end - pos - 1) * hs..]);
7778 pos = end;
7779 }
7780 }
7781 if task_mask.is_none()
7790 && !self.graph_prefill_preferred()
7791 && !std::env::var("CMF_PAIR").is_ok_and(|v| v == "0")
7792 && self.pair_supported()
7793 {
7794 while pos + 1 < ids.len() {
7795 let e1 = self.embed_single(ids[pos]);
7796 let e2 = self.embed_single(ids[pos + 1]);
7797 let (_, h2) = self.forward_pair(&e1, &e2, pos);
7798 self.check_forward_graph("forward_ids pair", pos + 1)?;
7799 self.commit_linear_scratch();
7800 hidden = h2;
7801 pos += 2;
7802 }
7803 }
7804 if task_mask.is_none() && pos == 0 && ids.len() > 1 {
7807 if let Some(lg) = self.embryo_prefill_chunked(ids, 0) {
7808 self.graph_logits = Some(lg);
7809 hidden = vec![0.0; self.hidden_size];
7810 pos = ids.len();
7811 }
7812 }
7813 while pos < ids.len() {
7814 hidden = self.forward_layers(&self.embed_single(ids[pos]), pos, task_mask);
7815 self.check_forward_graph("forward_ids", pos)?;
7816 pos += 1;
7817 }
7818 if let Some(logits) = self.graph_logits.take() {
7819 if let Err(err) = self.o1_seal_checked() {
7823 self.clear_sequence_state();
7824 return Err(err);
7825 }
7826 return Ok(logits);
7827 }
7828 if let Err(err) = self.o1_seal_checked() {
7832 self.clear_sequence_state();
7833 return Err(err);
7834 }
7835 let normed = inference::rms_norm(
7836 &hidden,
7837 &self.weights.final_norm,
7838 self.rms_eps,
7839 self.norm_style,
7840 );
7841 Ok(self.lm_head_forward(&normed))
7842 }
7843
7844 #[doc(hidden)]
7848 pub fn dsv41_serial_logits(&mut self, ids: &[u32]) -> Result<Vec<Vec<f32>>, String> {
7849 #[cfg(target_os = "macos")]
7850 crate::gpu_metal::set_io_namespace(self.graph_kv_id);
7851 if ids.is_empty() {
7852 return Err("empty id sequence".to_string());
7853 }
7854 self.clear_sequence_state();
7855 self.dsv41
7856 .as_ref()
7857 .ok_or_else(|| "dsv41 serial logits require a DeepSeek-V4.1 model".to_string())?;
7858 self.o1_begin();
7859 let rows = {
7860 let pool = self.pool.clone();
7861 let b = self
7862 .dsv41
7863 .as_mut()
7864 .expect("dsv41 checked above; state cannot change during forward");
7865 let mut rows = Vec::with_capacity(ids.len());
7866 for (position, &id) in ids.iter().enumerate() {
7867 let mut logits = Vec::new();
7868 crate::dsv41::forward_token(
7869 &b.0,
7870 &b.1,
7871 &b.2,
7872 &mut b.3,
7873 id,
7874 position,
7875 pool.as_deref(),
7876 &mut logits,
7877 );
7878 rows.push(logits);
7879 }
7880 rows
7881 };
7882 self.o1_seal();
7883 Ok(rows)
7884 }
7885
7886 pub fn ppl_ids(&mut self, ids: &[u32]) -> Result<f64, String> {
7893 let (nll, cnt) = self.nll_ids_from(ids, 0)?;
7894 Ok((nll / cnt.max(1) as f64).exp())
7895 }
7896
7897 pub fn probe_ffn_mass(&mut self, ids: &[u32]) -> Vec<Vec<f64>> {
7902 self.clear_sequence_state();
7903 FFN_PROBE.with(|p| {
7904 *p.borrow_mut() = Some(vec![vec![0f64; self.intermediate_size]; self.num_layers]);
7905 });
7906 crate::gpu::cpu_scope(|| {
7907 for (pos, &id) in ids.iter().enumerate() {
7908 let emb = self.embed_single(id);
7909 let _ = self.forward_layers(&emb, pos, None);
7910 }
7911 });
7912 self.clear_sequence_state();
7913 FFN_PROBE
7914 .with(|p| p.borrow_mut().take())
7915 .unwrap_or_default()
7916 }
7917
7918 pub fn probe_ffn_mass_batch(&mut self, ids: &[u32]) -> Result<Vec<Vec<f64>>, String> {
7922 if let Err(err) = self.nll_begin() {
7923 let _ = FFN_PROBE.with(|p| p.borrow_mut().take());
7927 self.nll_end();
7928 return Err(err);
7929 }
7930 FFN_PROBE.with(|p| {
7931 *p.borrow_mut() = Some(vec![vec![0f64; self.intermediate_size]; self.num_layers]);
7932 });
7933 let result: Result<(), String> = (|| {
7934 for chunk in ids.chunks(256) {
7935 if chunk.len() < 2 {
7936 continue;
7937 }
7938 self.nll_ids_masked(chunk, 0, None)?;
7939 }
7940 Ok(())
7941 })();
7942 self.nll_end();
7943 let probe = FFN_PROBE
7944 .with(|p| p.borrow_mut().take())
7945 .unwrap_or_default();
7946 match result {
7947 Ok(()) => Ok(probe),
7948 Err(err) => {
7949 drop(probe);
7950 Err(err)
7951 }
7952 }
7953 }
7954
7955 pub fn ppl_ids_masked(&mut self, ids: &[u32], mask: &TaskMask) -> Result<f64, String> {
7959 self.nll_begin()?;
7960 let result: Result<f64, String> = (|| {
7961 let mut nll = 0f64;
7962 let mut cnt = 0usize;
7963 let mut hidden = vec![0f32; self.hidden_size];
7964 for (pos, &id) in ids.iter().enumerate() {
7965 if pos > 0 {
7966 inference::rms_norm_into(
7967 &hidden,
7968 &self.weights.final_norm,
7969 self.rms_eps,
7970 self.norm_style,
7971 &mut self.ws.n1,
7972 );
7973 let mut logits = self.lm_head_forward(&self.ws.n1);
7974 let max = logits.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
7975 let sum: f64 = logits.iter().map(|&v| ((v - max) as f64).exp()).sum();
7976 let p = ((logits[id as usize] - max) as f64).exp() / sum.max(1e-300);
7977 nll -= p.max(1e-300).ln();
7978 cnt += 1;
7979 attention::recycle_buf(&mut logits);
7980 }
7981 let emb = self.embed_single(id);
7982 hidden = self.forward_layers(&emb, pos, Some(mask));
7983 self.nll_check_graph("masked serial forward", pos)?;
7984 let _ = self.graph_logits.take();
7988 }
7989 Ok((nll / cnt.max(1) as f64).exp())
7990 })();
7991 self.nll_end();
7992 result
7993 }
7994
7995 pub fn nll_ids_masked(
8014 &mut self,
8015 ids: &[u32],
8016 start: usize,
8017 task_mask: Option<&TaskMask>,
8018 ) -> Result<(f64, usize), String> {
8019 let task_mask = self.drop_open_mask(task_mask);
8020 self.nll_ids_inner(ids, start, task_mask)
8021 }
8022
8023 pub fn nll_ids_from(&mut self, ids: &[u32], start: usize) -> Result<(f64, usize), String> {
8024 self.nll_ids_inner(ids, start, None)
8025 }
8026
8027 fn nll_ids_inner(
8028 &mut self,
8029 ids: &[u32],
8030 start: usize,
8031 task_mask: Option<&TaskMask>,
8032 ) -> Result<(f64, usize), String> {
8033 self.nll_begin()?;
8034 let result: Result<(f64, usize), String> = (|| {
8035 let mut nll = 0f64;
8036 let mut cnt = 0usize;
8037 let (graph_quality, fused_head_quality) = nll_graph_policy(
8050 task_mask.is_none(),
8051 self.graph_prefill_preferred(),
8052 crate::gpu::q1_force(),
8053 );
8054 self.graph_head_required = fused_head_quality;
8055 self.graph_want_logits = fused_head_quality;
8056 #[cfg(target_os = "macos")]
8057 if graph_quality && std::env::var("CMF_METAL_BATCH_NLL").as_deref() != Ok("0") {
8058 match self.nll_batch_metal(ids, start) {
8059 MetalBatchNllOutcome::Completed(nll, count) => {
8060 return Ok((nll, count));
8061 }
8062 MetalBatchNllOutcome::Declined => {}
8063 MetalBatchNllOutcome::Failed(err) => return Err(err),
8064 }
8065 }
8066 if self.can_prefill_batched() && !graph_quality {
8067 const CHUNK: usize = 128;
8073 const LM_SUB: usize = 32;
8074 let n = ids.len().saturating_sub(1);
8075 let hs = self.hidden_size;
8076 let rows = self.weights.lm_head.rows();
8077 let mut pos = 0usize;
8078 let state_trace = std::env::var("CMF_STATE_TRACE").is_ok();
8079 while pos < n {
8080 let end = (pos + CHUNK).min(n);
8081 let bsz = end - pos;
8082 let hb = self.prefill_rows(&ids[pos..end], pos, task_mask)?;
8083 self.nll_check_graph("batched prefill", pos)?;
8084 if state_trace && end % 256 == 0 {
8085 self.trace_recurrent_state(end);
8086 }
8087 let mut k0 = 0usize;
8088 while k0 < bsz {
8089 let k1 = (k0 + LM_SUB).min(bsz);
8090 let sb = k1 - k0;
8091 if pos + k1 <= start {
8094 k0 = k1;
8095 continue;
8096 }
8097 let mut normed = vec![0.0f32; sb * hs];
8098 for k in 0..sb {
8099 let r = inference::rms_norm(
8100 &hb[(k0 + k) * hs..(k0 + k + 1) * hs],
8101 &self.weights.final_norm,
8102 self.rms_eps,
8103 self.norm_style,
8104 );
8105 normed[k * hs..(k + 1) * hs].copy_from_slice(&r);
8106 }
8107 let mut logits = vec![0.0f32; sb * rows];
8108 self.weights
8109 .lm_head
8110 .matmat(&normed, sb, &mut logits, self.pool.as_deref());
8111 for k in 0..sb {
8112 if pos + k0 + k < start {
8113 continue;
8114 }
8115 self.nll_check_graph("batched score row", pos + k0 + k)?;
8116 let lg = &mut logits[k * rows..k * rows + self.vocab_size.min(rows)];
8117 if let Some(mu) = self.logit_multiplier {
8118 for v in lg.iter_mut() {
8119 *v *= mu;
8120 }
8121 }
8122 if let Some(c) = self.final_softcap {
8126 for v in lg.iter_mut() {
8127 *v = c * (*v / c).tanh();
8128 }
8129 }
8130 if let Some(cm) = self.head_clusters.clone() {
8133 self.hierarchical_head_logprobs(
8134 &normed[k * hs..(k + 1) * hs],
8135 &cm,
8136 lg,
8137 );
8138 }
8139 let lg = &logits[k * rows..k * rows + self.vocab_size.min(rows)];
8140 let target = ids[pos + k0 + k + 1] as usize;
8141 let max = lg.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
8142 let lse: f64 = lg
8143 .iter()
8144 .map(|&v| ((v - max) as f64).exp())
8145 .sum::<f64>()
8146 .ln()
8147 + max as f64;
8148 nll += lse - lg[target] as f64;
8149 cnt += 1;
8150 if std::env::var("CMF_PPL_TRACE").is_ok() {
8151 let top = lg
8152 .iter()
8153 .enumerate()
8154 .max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
8155 .map(|(i, _)| i)
8156 .unwrap_or(0);
8157 eprintln!(
8158 "BTRACE pos {} target {} nll {:.4} top {} lg_t {:.3} lg_top {:.3}",
8159 pos + k0 + k,
8160 target,
8161 lse - lg[target] as f64,
8162 top,
8163 lg[target],
8164 lg[top]
8165 );
8166 }
8167 }
8168 k0 = k1;
8169 }
8170 pos = end;
8171 }
8172 return Ok((nll, cnt));
8173 }
8174 for pos in 0..ids.len().saturating_sub(1) {
8175 let hidden = self.forward_layers(&self.embed_single(ids[pos]), pos, task_mask);
8176 self.nll_check_graph("serial forward", pos)?;
8177 let out_of_band = self.graph_logits.take();
8185 if self.graph_head_required && out_of_band.is_none() {
8186 METAL_GRAPH_HEAD_MISS.fetch_add(
8187 1,
8188 std::sync::atomic::Ordering::Relaxed,
8189 );
8190 return Err(format!(
8191 "fused Metal graph head did not complete at NLL position {pos}"
8192 ));
8193 }
8194 if pos < start {
8195 continue;
8196 }
8197 let logits = match out_of_band {
8198 Some(lg) => lg,
8199 None => {
8200 let normed = inference::rms_norm(
8201 &hidden,
8202 &self.weights.final_norm,
8203 self.rms_eps,
8204 self.norm_style,
8205 );
8206 self.lm_head_forward(&normed)
8210 }
8211 };
8212 let target = ids[pos + 1] as usize;
8213 let max = logits.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
8214 let lse: f64 = logits
8215 .iter()
8216 .map(|&v| ((v - max) as f64).exp())
8217 .sum::<f64>()
8218 .ln()
8219 + max as f64;
8220 let tok_nll = lse - logits[target] as f64;
8221 if std::env::var("CMF_PPL_TRACE").is_ok() && pos < 48 {
8222 let top = logits
8223 .iter()
8224 .enumerate()
8225 .max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
8226 .map(|(i, _)| i)
8227 .unwrap_or(0);
8228 eprintln!(
8229 "pos {pos:3} tgt {target:6} nll {tok_nll:7.3} | top1 {top:6} lg[t]={:.2} lg[top]={:.2}",
8230 logits[target], logits[top]
8231 );
8232 }
8233 nll += tok_nll;
8234 cnt += 1;
8235 }
8236 Ok((nll, cnt))
8237 })();
8238 self.nll_end();
8239 result
8240 }
8241
8242 fn nll_from_hidden(&mut self, hidden: &[f32], target: u32, pos: usize) -> f64 {
8247 let normed = inference::rms_norm(
8248 hidden,
8249 &self.weights.final_norm,
8250 self.rms_eps,
8251 self.norm_style,
8252 );
8253 let mut logits = self.lm_head_forward(&normed);
8256 let target = target as usize;
8257 let max = logits.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
8258 let lse: f64 = logits
8259 .iter()
8260 .map(|&v| ((v - max) as f64).exp())
8261 .sum::<f64>()
8262 .ln()
8263 + max as f64;
8264 let tok_nll = lse - logits[target] as f64;
8265 if std::env::var("CMF_PPL_TRACE").is_ok() && pos < 48 {
8266 let top = logits
8267 .iter()
8268 .enumerate()
8269 .max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
8270 .map(|(i, _)| i)
8271 .unwrap_or(0);
8272 eprintln!(
8273 "pos {pos:3} tgt {target:6} nll {tok_nll:7.3} | top1 {top:6} lg[t]={:.2} lg[top]={:.2}",
8274 logits[target], logits[top]
8275 );
8276 }
8277 attention::recycle_buf(&mut logits);
8278 tok_nll
8279 }
8280
8281 fn trace_recurrent_state(&self, pos: usize) {
8289 let stats = |v: &[f32]| -> (f64, f64) {
8290 if v.is_empty() {
8291 return (0.0, 0.0);
8292 }
8293 let (mut ss, mut mx) = (0f64, 0f64);
8294 for &x in v {
8295 ss += (x as f64) * (x as f64);
8296 mx = mx.max((x as f64).abs());
8297 }
8298 ((ss / v.len() as f64).sqrt(), mx)
8299 };
8300 for (li, l) in self.kv_cache.layers.iter().enumerate() {
8301 let lw = &self.weights.layers[self.phys_layer(li)];
8302 let (kind, s_len) = match &lw.attn {
8303 AttnKind::Linear(_) => (
8304 "vmf",
8305 self.vmf_cfg.map(|c| c.state_len()).unwrap_or(0),
8306 ),
8307 AttnKind::LinearGdn(_) => (
8308 "gdn",
8309 self.gdn_cfg.map(|c| c.state_len()).unwrap_or(0),
8310 ),
8311 AttnKind::Bounded(_) => ("bounded", 0),
8312 AttnKind::Full { .. } => ("full", 0),
8313 _ => ("other", 0),
8314 };
8315 let (rms, max) = stats(&l.linear_state);
8316 let s_part = if kind == "vmf" {
8317 &l.linear_state[..s_len.min(l.linear_state.len())]
8318 } else if kind == "gdn" {
8319 let ring = l.linear_state.len().saturating_sub(s_len).min(l.linear_state.len());
8320 &l.linear_state[ring..]
8321 } else {
8322 &l.linear_state[..0]
8323 };
8324 let (s_rms, s_max) = stats(s_part);
8325 let (ring_rms, ring_len) = match &l.bounded {
8326 Some(b) => (stats(&b.ring_k).0, b.ring_k.len()),
8327 None => (0.0, 0),
8328 };
8329 eprintln!(
8330 "STATE pos={pos} layer={li} kind={kind} state_len={} rms={rms:.5} max={max:.4} \
8331 S_rms={s_rms:.5} S_max={s_max:.4} ring_k_rms={ring_rms:.5} ring_len={ring_len} kv_rows={}",
8332 l.linear_state.len(),
8333 l.seq_len
8334 );
8335 }
8336 }
8337
8338 pub fn nll_ids_o1(&mut self, ids: &[u32], prefill: usize) -> Result<(f64, usize), String> {
8356 self.nll_begin()?;
8361 let requested_prefix = (prefill > 0).then_some(prefill);
8362 self.o1_begin_with_prefix(requested_prefix);
8363 let n = ids.len().saturating_sub(1);
8364 let requested_start = prefill.min(n);
8365 let exact_end = if self.o1_active() {
8370 match requested_prefix {
8371 Some(requested) => self.o1_effective_boundary(requested),
8372 None => self
8373 .o1_cfg
8374 .as_ref()
8375 .and_then(|c| crate::nystrom::o1_deferred_boundary(c.w, c.sink)),
8376 }
8377 .unwrap_or(requested_start)
8378 .min(n)
8379 } else {
8380 requested_start
8381 };
8382 let mut nll = 0f64;
8383 let mut cnt = 0usize;
8384
8385 let mut pos = 0usize;
8389 if self.can_prefill_batched() {
8390 const CHUNK: usize = 128;
8391 while pos < exact_end {
8392 let end = (pos + CHUNK).min(exact_end);
8393 let hiddens = self.prefill_batch(&ids[pos..end], pos);
8394 if self
8395 .graph_failed
8396 .swap(false, std::sync::atomic::Ordering::Relaxed)
8397 {
8398 self.cancel
8399 .store(false, std::sync::atomic::Ordering::Relaxed);
8400 self.nll_end();
8401 return Err("GPU graph failed during O(1) NLL prefix".into());
8402 }
8403 for row in 0..end - pos {
8404 let score_pos = pos + row;
8405 if score_pos >= requested_start && score_pos < n {
8406 nll += self.nll_from_hidden(
8407 &hiddens[row * self.hidden_size..(row + 1) * self.hidden_size],
8408 ids[score_pos + 1],
8409 score_pos,
8410 );
8411 cnt += 1;
8412 }
8413 }
8414 pos = end;
8415 }
8416 } else {
8417 while pos < exact_end {
8418 let hidden = self.forward_layers(&self.embed_single(ids[pos]), pos, None);
8419 if self
8420 .graph_failed
8421 .swap(false, std::sync::atomic::Ordering::Relaxed)
8422 {
8423 self.cancel
8424 .store(false, std::sync::atomic::Ordering::Relaxed);
8425 self.nll_end();
8426 return Err("GPU graph failed during O(1) NLL prefix".into());
8427 }
8428 if pos >= requested_start && pos < n {
8429 nll += self.nll_from_hidden(&hidden, ids[pos + 1], pos);
8430 cnt += 1;
8431 }
8432 pos += 1;
8433 }
8434 }
8435 self.o1_seal_checked().map_err(|err| {
8436 self.nll_end();
8437 err
8438 })?;
8439
8440 let batch_k = std::env::var("CMF_BATCH_K")
8449 .ok()
8450 .and_then(|v| v.parse::<usize>().ok())
8451 .unwrap_or(0);
8452 let batch_admitted = batch_k > 0
8453 && self.can_prefill_batched()
8454 && self.o1_active()
8455 && std::env::var("CMF_O1_GPU").as_deref() == Ok("1")
8456 && (0..self.num_layers).all(|li| {
8457 let cache = &self.kv_cache.layers[self.phys_layer(li)];
8458 cache.o1.is_none() || cache.o1_views().is_some()
8459 });
8460 if std::env::var("CMF_GRAPH_PROF").is_ok() {
8461 eprintln!(
8462 "nll-batch: phase=post-seal admission={} requested_k={} scored_rows={}",
8463 batch_admitted,
8464 batch_k,
8465 n.saturating_sub(exact_end),
8466 );
8467 }
8468 let mut batch_completed = false;
8469 if batch_admitted && exact_end < n {
8470 let hs = self.hidden_size;
8471 let mut batch_pos = exact_end;
8472 while batch_pos < n {
8473 let end = (batch_pos + batch_k).min(n);
8474 let bk = end - batch_pos;
8475 let mut hiddens = vec![0.0f32; bk * hs];
8476 for (row, &id) in ids[batch_pos..end].iter().enumerate() {
8477 hiddens[row * hs..(row + 1) * hs].copy_from_slice(&self.embed_single(id));
8478 }
8479 let positions: Vec<usize> = (batch_pos..end).collect();
8480 let t_batch = std::time::Instant::now();
8481 let outcome = self.try_batch_graph_wgpu(&mut hiddens, &positions, bk, None);
8482 if std::env::var("CMF_GRAPH_PROF").is_ok() {
8483 let ms = t_batch.elapsed().as_secs_f64() * 1000.0;
8484 eprintln!(
8485 "nll-batch: phase=post-seal mode=o1 k={bk} pos={}..{} outcome={outcome:?} {ms:.1} ms ({:.1} tok/s)",
8486 batch_pos,
8487 end.saturating_sub(1),
8488 bk as f64 / (ms / 1000.0),
8489 );
8490 }
8491 if let Err(err) = self.nll_check_graph("batch graph", batch_pos) {
8492 self.nll_end();
8493 return Err(err);
8494 }
8495 match outcome {
8496 crate::gpu::BatchGraphOutcome::Completed => {
8497 batch_completed = true;
8498 for row in 0..bk {
8499 nll += self.nll_from_hidden(
8500 &hiddens[row * hs..(row + 1) * hs],
8501 ids[batch_pos + row + 1],
8502 batch_pos + row,
8503 );
8504 cnt += 1;
8505 }
8506 batch_pos = end;
8507 }
8508 crate::gpu::BatchGraphOutcome::Declined => {
8509 if batch_completed {
8510 self.nll_end();
8511 return Err(format!(
8512 "O(1) NLL batch declined after completed chunk at position {batch_pos}"
8513 ));
8514 }
8515 break;
8516 }
8517 crate::gpu::BatchGraphOutcome::Failed => {
8518 self.nll_end();
8519 return Err(format!(
8520 "O(1) NLL batch graph failed after admission at position {batch_pos}"
8521 ));
8522 }
8523 }
8524 }
8525 if batch_completed && cnt == n.saturating_sub(requested_start) {
8526 self.nll_end();
8527 return Ok((nll, cnt));
8528 }
8529 }
8530
8531 for pos in exact_end..n {
8536 let hidden = self.forward_layers(&self.embed_single(ids[pos]), pos, None);
8537 if self
8538 .graph_failed
8539 .swap(false, std::sync::atomic::Ordering::Relaxed)
8540 {
8541 self.cancel
8542 .store(false, std::sync::atomic::Ordering::Relaxed);
8543 self.nll_end();
8544 return Err(format!(
8545 "GPU graph failed during O(1) NLL serial scoring at position {pos}"
8546 ));
8547 }
8548 nll += self.nll_from_hidden(&hidden, ids[pos + 1], pos);
8549 cnt += 1;
8550 }
8551 self.nll_end();
8552 Ok((nll, cnt))
8553 }
8554
8555 pub fn calib_ids(&mut self, ids: &[u32], temps: &[f32]) -> (Vec<bool>, Vec<Vec<f32>>) {
8563 self.clear_sequence_state();
8564 let n = ids.len().saturating_sub(1);
8565 let mut correct = Vec::with_capacity(n);
8566 let mut pmax = Vec::with_capacity(n);
8567 for pos in 0..n {
8568 let emb = self.embed_single(ids[pos]);
8569 let hidden = self.forward_layers(&emb, pos, None);
8570 let logits = if let Some(logits) = self.graph_logits.take() {
8571 logits
8572 } else {
8573 let normed = inference::rms_norm(
8574 &hidden,
8575 &self.weights.final_norm,
8576 self.rms_eps,
8577 self.norm_style,
8578 );
8579 self.lm_head_forward(&normed)
8583 };
8584 let target = ids[pos + 1] as usize;
8585 let (mut amax, mut mval) = (0usize, f32::NEG_INFINITY);
8586 for (i, &v) in logits.iter().enumerate() {
8587 if v > mval {
8588 mval = v;
8589 amax = i;
8590 }
8591 }
8592 correct.push(amax == target);
8593 let row: Vec<f32> = temps
8594 .iter()
8595 .map(|&t| {
8596 let tt = t.max(1e-3);
8597 let s: f32 = logits.iter().map(|&v| ((v - mval) / tt).exp()).sum();
8598 1.0 / s.max(1e-12) })
8600 .collect();
8601 pmax.push(row);
8602 }
8603 self.clear_sequence_state();
8604 (correct, pmax)
8605 }
8606
8607 pub fn ppl_ids_dynamic(&mut self, ids: &[u32]) -> Result<(f64, usize), String> {
8614 if self.dyn_router.is_none() {
8615 return Ok((self.ppl_ids(ids)?, 0));
8616 }
8617 self.nll_begin()?;
8618 let saved_active = self.dyn_active;
8619 let mut router = self
8620 .dyn_router
8621 .take()
8622 .ok_or_else(|| "dynamic router disappeared before PPL scoring".to_string())?;
8623 router.reset();
8624 self.dyn_phi_seen = 0;
8625 let _ = self.set_active_skill(None);
8626
8627 let result: Result<(f64, usize), String> = (|| {
8628 let mut nll = 0f64;
8629 let mut cnt = 0usize;
8630 for pos in 0..ids.len().saturating_sub(1) {
8631 let hidden = self.forward_layers(&self.embed_single(ids[pos]), pos, None);
8632 self.nll_check_graph("dynamic serial forward", pos)?;
8633 let out_of_band = self.graph_logits.take();
8634 let mut logits = match out_of_band {
8635 Some(lg) => lg,
8636 None => {
8637 let normed = inference::rms_norm(
8638 &hidden,
8639 &self.weights.final_norm,
8640 self.rms_eps,
8641 self.norm_style,
8642 );
8643 self.lm_head_forward(&normed)
8647 }
8648 };
8649 let target = ids[pos + 1] as usize;
8650 let max = logits.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
8651 let lse: f64 = logits
8652 .iter()
8653 .map(|&v| ((v - max) as f64).exp())
8654 .sum::<f64>()
8655 .ln()
8656 + max as f64;
8657 let tok_nll = lse - logits[target] as f64;
8658 if std::env::var("CMF_PPL_TRACE").is_ok() && pos < 48 {
8659 let top = logits
8660 .iter()
8661 .enumerate()
8662 .max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
8663 .map(|(i, _)| i)
8664 .unwrap_or(0);
8665 eprintln!(
8666 "pos {pos:3} tgt {target:6} nll {tok_nll:7.3} | top1 {top:6} lg[t]={:.2} lg[top]={:.2}",
8667 logits[target], logits[top]
8668 );
8669 }
8670 nll += tok_nll;
8671 cnt += 1;
8672 attention::recycle_buf(&mut logits);
8673 let phi = self.dyn_phi_ema.clone();
8675 if let Some(new_active) = router.step(&phi, pos) {
8676 let _ = self.set_active_skill(new_active);
8677 }
8678 }
8679 Ok(((nll / cnt.max(1) as f64).exp(), router.switches.len()))
8680 })();
8681
8682 let _ = self.set_active_skill(saved_active);
8685 self.dyn_router = Some(router);
8686 self.nll_end();
8687 result
8688 }
8689
8690 pub fn probe_phi(&mut self, ids: &[u32], layer: usize) -> Vec<f32> {
8692 self.clear_sequence_state();
8693 let mut acc = vec![0f32; self.hidden_size];
8694 for (pos, &id) in ids.iter().enumerate() {
8695 let h = self.forward_layers_upto(&self.embed_single(id), pos, None, Some(layer));
8696 for (a, v) in acc.iter_mut().zip(&h) {
8697 *a += v;
8698 }
8699 }
8700 let n = ids.len().max(1) as f32;
8701 for a in acc.iter_mut() {
8702 *a /= n;
8703 }
8704 self.clear_sequence_state();
8705 acc
8706 }
8707
8708 pub fn probe_phi_span(
8722 &mut self,
8723 ids: &[u32],
8724 layer: usize,
8725 span: std::ops::Range<usize>,
8726 ) -> Vec<f32> {
8727 #[cfg(target_os = "macos")]
8728 crate::gpu_metal::set_io_namespace(self.graph_kv_id);
8729 let end = span.end.min(ids.len());
8730 let start = span.start.min(end);
8731 let reset = |p: &mut Self| p.clear_sequence_state();
8732 reset(self);
8733 let mut acc = vec![0f32; self.hidden_size];
8734 for (pos, &id) in ids[..end].iter().enumerate() {
8735 let h = self.forward_layers_upto(&self.embed_single(id), pos, None, Some(layer));
8736 if pos >= start {
8737 for (a, v) in acc.iter_mut().zip(&h) {
8738 *a += v;
8739 }
8740 }
8741 }
8742 let n = end - start;
8743 if n > 0 {
8744 let n = n as f32;
8745 for a in acc.iter_mut() {
8746 *a /= n;
8747 }
8748 }
8749 reset(self);
8750 acc
8751 }
8752
8753 pub fn decode_step_logits(&mut self, token: u32, position: usize) -> Vec<f32> {
8760 #[cfg(target_os = "macos")]
8761 crate::gpu_metal::set_io_namespace(self.graph_kv_id);
8762 self.graph_logits = None;
8763 let hidden = self.forward_layers(&self.embed_single(token), position, None);
8764 if let Some(logits) = self.graph_logits.take() {
8765 return logits;
8766 }
8767 inference::rms_norm_into(
8768 &hidden,
8769 &self.weights.final_norm,
8770 self.rms_eps,
8771 self.norm_style,
8772 &mut self.ws.n1,
8773 );
8774 self.lm_head_forward(&self.ws.n1)
8775 }
8776
8777 fn prefill_batch(&mut self, ids: &[u32], start_pos: usize) -> Vec<f32> {
8783 self.prefill_batch_masked(ids, start_pos, None)
8784 }
8785
8786 fn prefill_batch_masked(
8792 &mut self,
8793 ids: &[u32],
8794 start_pos: usize,
8795 task_mask: Option<&TaskMask>,
8796 ) -> Vec<f32> {
8797 self.prefill_batch_span(PrefillIn::Ids(ids), start_pos, task_mask, 0, usize::MAX)
8798 }
8799
8800 fn prefill_rows(
8808 &mut self,
8809 ids: &[u32],
8810 pos: usize,
8811 task_mask: Option<&TaskMask>,
8812 ) -> Result<Vec<f32>, String> {
8813 self.prefill_input_rows(PrefillIn::Ids(ids), pos, task_mask)
8814 }
8815
8816 fn prefill_input_rows(
8817 &mut self,
8818 input: PrefillIn<'_>,
8819 pos: usize,
8820 task_mask: Option<&TaskMask>,
8821 ) -> Result<Vec<f32>, String> {
8822 self.mimo_moe_prepare();
8823 let hs = self.hidden_size;
8824 let bk = match input {
8825 PrefillIn::Ids(ids) => ids.len(),
8826 PrefillIn::Hidden(rows) => rows.len() / hs,
8827 };
8828 #[cfg(not(target_os = "macos"))]
8829 if task_mask.is_none()
8830 && !self.o1_active()
8831 && bk > 1
8832 && (self.batch_prefix_prefill()
8833 || (self.verify_exact_moe
8834 && crate::gpu::enabled_here()
8835 && crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Decode)))
8836 {
8837 let mut hiddens = match input {
8838 PrefillIn::Hidden(rows) => rows.to_vec(),
8839 PrefillIn::Ids(ids) => ids.iter().flat_map(|&id| self.embed_single(id)).collect(),
8840 };
8841 let positions: Vec<usize> = (pos..pos + bk).collect();
8842 let mut run = 0usize;
8843 match self.try_batch_graph_wgpu_prefix(
8844 &mut hiddens,
8845 &positions,
8846 bk,
8847 None,
8848 Some(&mut run),
8849 ) {
8850 crate::gpu::BatchGraphOutcome::Completed => {
8851 let out = if run < self.num_layers {
8852 self.prefill_batch_span(
8853 PrefillIn::Hidden(&hiddens),
8854 pos,
8855 None,
8856 run,
8857 self.num_layers,
8858 )
8859 } else {
8860 hiddens
8861 };
8862 return if self.graph_failed.load(std::sync::atomic::Ordering::Relaxed) {
8863 Err("MiMo attention graph failed after admission".into())
8864 } else { Ok(out) };
8865 }
8866 crate::gpu::BatchGraphOutcome::Failed => {
8867 return Err("batched prefix prefill failed after admission".into());
8868 }
8869 crate::gpu::BatchGraphOutcome::Declined => {
8870 #[cfg(feature = "gpu")]
8872 self.pull_lagging_host_kv(0, self.num_layers, pos);
8873 }
8874 }
8875 }
8876 let out = self.prefill_batch_span(input, pos, task_mask, 0, usize::MAX);
8877 if self.graph_failed.load(std::sync::atomic::Ordering::Relaxed) {
8878 Err("batch tail graph failed after admission".into())
8879 } else { Ok(out) }
8880 }
8881
8882 fn prefill_batch_span(
8888 &mut self,
8889 input: PrefillIn<'_>,
8890 start_pos: usize,
8891 task_mask: Option<&TaskMask>,
8892 from: usize,
8893 upto_excl: usize,
8894 ) -> Vec<f32> {
8895 let hs = self.hidden_size;
8896 let b = match input {
8897 PrefillIn::Ids(ids) => ids.len(),
8898 PrefillIn::Hidden(hb) => hb.len() / hs,
8899 };
8900 let upto_excl = upto_excl.min(self.num_layers);
8901 let mut h: Vec<f32>;
8905 let mut h_ready;
8906 match input {
8907 PrefillIn::Ids(_) => {
8908 h = vec![0.0; b * hs];
8909 h_ready = false;
8910 }
8911 PrefillIn::Hidden(hb) => {
8912 h = hb.to_vec();
8913 h_ready = true;
8914 }
8915 }
8916 let fill_h = |h: &mut Vec<f32>, me: &Self| {
8917 if let PrefillIn::Ids(ids) = input {
8918 for (bi, &id) in ids.iter().enumerate() {
8919 let e = me.embed_single(id);
8920 h[bi * hs..(bi + 1) * hs].copy_from_slice(&e);
8921 }
8922 if let Ok(tp) = std::env::var("CMF_TRACE_POS") {
8923 if let Ok(t) = tp.parse::<usize>() {
8924 if t >= start_pos && t < start_pos + ids.len() {
8925 let bi = t - start_pos;
8926 let row = &h[bi * hs..(bi + 1) * hs];
8927 let n: f32 = row.iter().map(|x| x * x).sum::<f32>().sqrt();
8928 eprintln!(
8929 "BATCH pos {t} embed: id {} |h| = {n:.6} h0 {:.6} h1 {:.6} | b={} start={start_pos} ids[..8]={:?}",
8930 ids[bi],
8931 row[0],
8932 row[1],
8933 ids.len(),
8934 &ids[..ids.len().min(8)]
8935 );
8936 }
8937 }
8938 }
8939 }
8940 };
8941 let (_nkv, _hd, _rd, eps) = (
8942 self.num_kv_heads,
8943 self.head_dim,
8944 self.rotary_dim,
8945 self.rms_eps,
8946 );
8947 let pool = self.pool.clone();
8948 let norm_style = self.norm_style;
8949 self.mimo_moe_prepare();
8950 let automatic_gpu_prefix = self.automatic_gpu_prefix();
8951
8952 #[cfg(target_os = "macos")]
8953 let mut chunk_skip_until = 0usize;
8954 for li in from..upto_excl {
8955 let _capacity_tail = automatic_gpu_prefix
8956 .filter(|&prefix| {
8957 li >= prefix && !(self.verify_exact_moe && self.mimo_moe.is_dynamic(li, false))
8958 })
8959 .map(|_| crate::gpu::enter_cpu_scope());
8960 crate::gpu::set_layer(li as i64); #[cfg(target_os = "macos")]
8967 if task_mask.is_none() {
8968 if li < chunk_skip_until {
8969 continue;
8970 }
8971 if !h_ready && li == 0 && self.weights.embed_tokens.q8_row_parts().is_none() {
8977 fill_h(&mut h, self);
8978 h_ready = true;
8979 }
8980 let ids_for_embed = match input {
8981 PrefillIn::Ids(ids) => (!h_ready && li == 0).then_some(ids),
8982 PrefillIn::Hidden(_) => None,
8983 };
8984 let end = self.chunk_run_gpu(li, &mut h, b, start_pos, ids_for_embed, upto_excl);
8985 if end > li {
8986 h_ready = true;
8987 chunk_skip_until = end;
8988 if self.is_loop_end(end - 1) && end < self.num_layers {
8991 for bi in 0..b {
8992 let normed = inference::rms_norm(
8993 &h[bi * hs..(bi + 1) * hs],
8994 &self.weights.final_norm,
8995 eps,
8996 norm_style,
8997 );
8998 h[bi * hs..(bi + 1) * hs].copy_from_slice(&normed);
8999 }
9000 }
9001 continue;
9002 }
9003 }
9004 if !h_ready {
9005 fill_h(&mut h, self);
9006 h_ready = true;
9007 }
9008 if task_mask.is_none() && self.verify_exact_moe {
9009 let positions: Vec<_> = (start_pos..start_pos + b).collect();
9010 match self.mimo_graph_layer_rows(li, &mut h, &positions) {
9011 crate::gpu::BatchGraphOutcome::Completed => continue,
9012 crate::gpu::BatchGraphOutcome::Failed => return h,
9013 crate::gpu::BatchGraphOutcome::Declined => {},
9014 }
9015 }
9016 #[cfg(feature = "gpu")]
9017 self.pull_lagging_host_kv(li, li + 1, start_pos);
9018 let lw = &self.weights.layers[self.phys_layer(li)];
9019 match &lw.attn {
9021 AttnKind::Kda(w) => {
9022 let cfg = self.kda_cfg.expect("kda layer without kda_cfg");
9024 let mut normed = vec![0.0f32; b * hs];
9025 for bi in 0..b {
9026 inference::rms_norm_into(
9027 &h[bi * hs..(bi + 1) * hs],
9028 &lw.input_norm,
9029 eps,
9030 norm_style,
9031 &mut normed[bi * hs..(bi + 1) * hs],
9032 );
9033 }
9034 let attn = crate::linear_core::kda_forward_batch(
9035 &normed,
9036 b,
9037 w,
9038 &cfg,
9039 &mut self.kv_cache.layers[li].linear_state,
9040 pool.as_deref(),
9041 );
9042 for (dst, &a) in h.iter_mut().zip(&attn) {
9043 *dst += a;
9044 }
9045 }
9046 AttnKind::LinearGdn(w) => {
9047 let cfg = self.gdn_cfg.expect("gdn layer without gdn_cfg");
9049 let mut normed = vec![0.0f32; b * hs];
9050 for bi in 0..b {
9051 let r = inference::rms_norm(
9052 &h[bi * hs..(bi + 1) * hs],
9053 &lw.input_norm,
9054 eps,
9055 norm_style,
9056 );
9057 normed[bi * hs..(bi + 1) * hs].copy_from_slice(&r);
9058 }
9059 let attn = crate::linear_core::gdn_forward_batch(
9060 &normed,
9061 b,
9062 w,
9063 &cfg,
9064 &mut self.kv_cache.layers[li].linear_state,
9065 pool.as_deref(),
9066 );
9067 for (dst, &a) in h.iter_mut().zip(&attn) {
9068 *dst += a;
9069 }
9070 }
9071 AttnKind::ShortConv(w) => {
9072 let cfg = self
9075 .short_conv_cfg
9076 .expect("short-conv layer without short_conv_cfg");
9077 let mut normed = vec![0.0f32; b * hs];
9078 for bi in 0..b {
9079 inference::rms_norm_into(
9080 &h[bi * hs..(bi + 1) * hs],
9081 &lw.input_norm,
9082 eps,
9083 norm_style,
9084 &mut normed[bi * hs..(bi + 1) * hs],
9085 );
9086 }
9087 let attn = short_conv_forward_batch(
9088 &normed,
9089 b,
9090 w,
9091 &cfg,
9092 &mut self.kv_cache.layers[li].linear_state,
9093 pool.as_deref(),
9094 );
9095 for (dst, &a) in h.iter_mut().zip(&attn) {
9096 *dst += a;
9097 }
9098 }
9099 AttnKind::Mla(w) => {
9100 let inv_freq_l = self.layer_inv_freq(li);
9103 let rs = self.layer_rope_scale(li);
9104 let mut normed = vec![0.0f32; hs];
9105 for bi in 0..b {
9106 inference::rms_norm_into(
9107 &h[bi * hs..(bi + 1) * hs],
9108 &lw.input_norm,
9109 eps,
9110 norm_style,
9111 &mut normed,
9112 );
9113 let ao = mla_attention(
9114 w,
9115 &normed,
9116 &mut self.kv_cache.layers[li],
9117 start_pos + bi,
9118 &inv_freq_l,
9119 rs,
9120 eps,
9121 pool.as_deref(),
9122 );
9123 for (dst, &a) in h[bi * hs..(bi + 1) * hs].iter_mut().zip(&ao) {
9124 *dst += a;
9125 }
9126 }
9127 }
9128 AttnKind::Full {
9129 wq,
9130 wk,
9131 wv,
9132 wo,
9133 q_norm,
9134 k_norm,
9135 output_gate,
9136 softplus_gate,
9137 bias,
9138 } => {
9139 let mut normed = vec![0.0f32; b * hs];
9143 for bi in 0..b {
9144 inference::rms_norm_into(
9145 &h[bi * hs..(bi + 1) * hs],
9146 &lw.input_norm,
9147 eps,
9148 norm_style,
9149 &mut normed[bi * hs..(bi + 1) * hs],
9150 );
9151 }
9152 let inv_freq_l = self.layer_inv_freq(li);
9153 let (nkv_l, hd_l, rd_l) = self.layer_geom(li);
9154 let cfg = QwenAttnCfg {
9155 num_heads: self.layer_num_heads(li),
9156 num_kv_heads: nkv_l,
9157 head_dim: hd_l,
9158 hidden_size: hs,
9159 position: start_pos,
9160 inv_freq: &inv_freq_l,
9161 rotary_dim: rd_l,
9162 scale: self.attn_scale,
9163 softcap: self.attn_softcap,
9164 window: self.layer_window(li),
9165 v_norm: self.attn_v_norm,
9166 qk_norm_after_rope: self.qk_norm_after_rope,
9167 q_norm: q_norm.as_deref(),
9168 k_norm: k_norm.as_deref(),
9169 output_gate: *output_gate,
9170 softplus_gate: softplus_gate
9171 .as_ref()
9172 .map(|(gate, per_head)| (gate, *per_head)),
9173 rope_scale: self.layer_rope_scale(li),
9174 bias: bias
9175 .as_ref()
9176 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
9177 rms_eps: eps,
9178 norm_style,
9179 pool: pool.as_deref(),
9180 v_head_dim: self.layer_v_dim(li),
9181 };
9182 let mut attn = attention::qwen_attention_batch(
9183 &normed,
9184 b,
9185 wq,
9186 wk,
9187 wv,
9188 wo,
9189 &mut self.kv_cache.layers[li],
9190 &cfg,
9191 );
9192 if let Some(w) = &lw.attn_out_norm {
9193 for bi in 0..b {
9194 inference::rms_norm_into(
9195 &attn[bi * hs..(bi + 1) * hs],
9196 w,
9197 eps,
9198 norm_style,
9199 &mut normed[bi * hs..(bi + 1) * hs],
9200 );
9201 }
9202 attn.copy_from_slice(&normed);
9203 }
9204 for (dst, &a) in h.iter_mut().zip(&attn) {
9205 *dst += a;
9206 }
9207 }
9208 AttnKind::Bounded(w) => {
9209 let mut normed = vec![0.0f32; b * hs];
9212 for bi in 0..b {
9213 inference::rms_norm_into(
9214 &h[bi * hs..(bi + 1) * hs],
9215 &lw.input_norm,
9216 eps,
9217 norm_style,
9218 &mut normed[bi * hs..(bi + 1) * hs],
9219 );
9220 }
9221 let rope = self
9222 .bounded_rope
9223 .clone()
9224 .expect("bounded layer without an installed rotation table");
9225 let cfg = crate::bounded::BoundedAttnCfg {
9226 num_heads: self.num_heads,
9227 num_kv_heads: self.num_kv_heads,
9228 head_dim: self.head_dim,
9229 hidden_size: hs,
9230 scale: self.attn_scale,
9231 rope: &rope,
9232 pool: pool.as_deref(),
9233 };
9234 let mut attn = crate::bounded::bounded_attention_batch(
9235 &normed,
9236 b,
9237 w,
9238 &mut self.kv_cache.layers[li],
9239 &cfg,
9240 );
9241 if let Some(wn) = &lw.attn_out_norm {
9242 for bi in 0..b {
9243 inference::rms_norm_into(
9244 &attn[bi * hs..(bi + 1) * hs],
9245 wn,
9246 eps,
9247 norm_style,
9248 &mut normed[bi * hs..(bi + 1) * hs],
9249 );
9250 }
9251 attn.copy_from_slice(&normed);
9252 }
9253 for (dst, &a) in h.iter_mut().zip(&attn) {
9254 *dst += a;
9255 }
9256 attention::recycle_buf(&mut attn);
9257 }
9258 AttnKind::Linear(w) => {
9259 for bi in 0..b {
9260 let normed = inference::rms_norm(
9261 &h[bi * hs..(bi + 1) * hs],
9262 &lw.input_norm,
9263 eps,
9264 norm_style,
9265 );
9266 vmf_phase_forward(
9267 &normed,
9268 w,
9269 &self.vmf_cfg.expect("linear layer without vmf_cfg"),
9270 &mut self.kv_cache.layers[li].linear_state,
9271 pool.as_deref(),
9272 )
9273 .iter()
9274 .enumerate()
9275 .for_each(|(i, &a)| h[bi * hs + i] += a);
9276 }
9277 }
9278 }
9279
9280 let lw = &self.weights.layers[self.phys_layer(li)];
9282 let mut post = vec![0.0f32; b * hs];
9283 for bi in 0..b {
9284 let r =
9285 inference::rms_norm(&h[bi * hs..(bi + 1) * hs], &lw.post_norm, eps, norm_style);
9286 post[bi * hs..(bi + 1) * hs].copy_from_slice(&r);
9287 }
9288 let mask_row = task_mask
9291 .filter(|m| m.ffn_active_count(li) < self.intermediate_size)
9292 .and_then(|m| m.ffn_masks.get(li))
9293 .map(|v| v.as_slice());
9294 let mut ffn = match &lw.ffn {
9295 FfnKind::Dense(d) if !d.segs.is_empty() => {
9296 tube_ffn(d, &post, b, pool.as_deref(), mask_row)
9297 }
9298 FfnKind::Dense(d) => dense_ffn_batch(d, &post, b, pool.as_deref(), mask_row),
9299 FfnKind::Moe(m) if self.verify_exact_moe && self.mimo_moe.is_dynamic(li, false) => {
9300 moe_ffn_banked_rows(&mut self.mimo_moe, li, m, &post, b, hs, pool.as_deref())
9301 }
9302 FfnKind::Moe(m) if self.verify_exact_moe => {
9303 moe_ffn_rows_exact(m, &post, b, hs, pool.as_deref())
9304 }
9305 FfnKind::Moe(m) if self.mimo_moe.is_dynamic(li, false) => {
9308 let before = m.stats.borrow().clone();
9309 let out = crate::gpu::cpu_scope(|| {
9310 moe_ffn_batch(m, &post, b, hs, pool.as_deref(), None)
9311 });
9312 self.mimo_moe.prime(li, m, &before);
9313 out
9314 }
9315 FfnKind::Moe(m) => moe_ffn_batch(m, &post, b, hs, pool.as_deref(), None),
9316 FfnKind::DenseMoe(dm) => {
9319 let mut out = vec![0.0f32; b * hs];
9320 for bi in 0..b {
9321 let r = dense_moe_ffn(
9322 dm,
9323 &post[bi * hs..(bi + 1) * hs],
9324 &h[bi * hs..(bi + 1) * hs],
9325 eps,
9326 norm_style,
9327 pool.as_deref(),
9328 );
9329 out[bi * hs..(bi + 1) * hs].copy_from_slice(&r);
9330 }
9331 out
9332 }
9333 };
9334 if let Some(w) = &lw.ffn_out_norm {
9335 for bi in 0..b {
9336 inference::rms_norm_into(
9337 &ffn[bi * hs..(bi + 1) * hs],
9338 w,
9339 eps,
9340 norm_style,
9341 &mut post[bi * hs..(bi + 1) * hs],
9342 );
9343 }
9344 ffn.copy_from_slice(&post);
9345 }
9346 for (dst, &f) in h.iter_mut().zip(&ffn) {
9347 *dst += f;
9348 }
9349 if let Some(sc) = lw.layer_scale {
9350 for v in h.iter_mut() {
9351 *v *= sc;
9352 }
9353 }
9354 if self.layer_dump.is_some() {
9356 for bi in 0..b {
9357 self.dump_layer_row(start_pos + bi, li, &h[bi * hs..(bi + 1) * hs]);
9358 }
9359 }
9360 if let Ok(tp) = std::env::var("CMF_TRACE_POS") {
9361 if let Ok(t) = tp.parse::<usize>() {
9362 if t >= start_pos && t < start_pos + b {
9363 let bi = t - start_pos;
9364 let row = &h[bi * hs..(bi + 1) * hs];
9365 let n: f32 = row.iter().map(|x| x * x).sum::<f32>().sqrt();
9366 eprintln!(
9367 "BATCH pos {t} after layer {li}: |h| = {n:.6} h0 {:.6} h1 {:.6}",
9368 row[0], row[1]
9369 );
9370 }
9371 }
9372 }
9373 if std::env::var("CMF_DEBUG_LAYERS").is_ok() {
9377 let row = &h[(b - 1) * hs..b * hs];
9378 let rms =
9379 (row.iter().map(|&v| (v as f64) * (v as f64)).sum::<f64>() / hs as f64).sqrt();
9380 let mx = row.iter().fold(0f32, |m, &v| m.max(v.abs()));
9381 eprintln!(
9382 "layer {li:>3} {:>10} ffn={:<5} rms={rms:>12.4} max={mx:>12.4}",
9383 match &self.weights.layers[self.phys_layer(li)].attn {
9384 AttnKind::LinearGdn(_) => "gdn",
9385 AttnKind::Linear(_) => "vmf",
9386 AttnKind::ShortConv(_) => "conv",
9387 _ => "attn",
9388 },
9389 match &lw.ffn {
9390 FfnKind::Moe(_) => "moe",
9391 FfnKind::Dense(_) => "dense",
9392 FfnKind::DenseMoe(_) => "dense+moe",
9393 },
9394 );
9395 }
9396 if self.is_loop_end(li) && li + 1 < self.num_layers {
9398 for bi in 0..b {
9399 let normed = inference::rms_norm(
9400 &h[bi * hs..(bi + 1) * hs],
9401 &self.weights.final_norm,
9402 eps,
9403 norm_style,
9404 );
9405 h[bi * hs..(bi + 1) * hs].copy_from_slice(&normed);
9406 }
9407 }
9408 if std::env::var("CMF_TRACE_H").is_ok() {
9409 let n = h[..hs].iter().map(|v| v.abs()).sum::<f32>() / hs as f32;
9410 let mx = h[..hs].iter().fold(0.0f32, |a, &v| a.max(v.abs()));
9411 eprintln!(
9412 "layer {li}: mean|h|={n:.4} max|h|={mx:.2} scale={:?}",
9413 lw.layer_scale
9414 );
9415 }
9416 }
9417 crate::gpu::set_layer(-1); self.o1_progress();
9423 h
9424 }
9425
9426 fn embed_single(&self, id: u32) -> Vec<f32> {
9428 let mut out = vec![0.0f32; self.hidden_size];
9429 if (id as usize) < self.weights.embed_tokens.rows() {
9430 self.weights.embed_tokens.row_f32(id as usize, &mut out);
9431 }
9432 if self.embed_multiplier != 1.0 {
9433 for v in out.iter_mut() {
9434 *v *= self.embed_multiplier;
9435 }
9436 }
9437 if self.dsv4.is_some()
9441 || self.dsv41.is_some()
9442 || self.qwen4_exp.is_some()
9443 {
9444 let mut v = vec![0.0f32; self.hidden_size.max(1)];
9445 v[0] = id as f32;
9446 return v;
9447 }
9448 if let Some(b) = &self.g3n {
9451 return b.0.extend_embedding(id, &out, self.pool.as_deref());
9452 }
9453 out
9454 }
9455
9456 #[cfg(target_os = "macos")]
9462 fn chunk_run_gpu(
9463 &mut self,
9464 li0: usize,
9465 h: &mut [f32],
9466 b: usize,
9467 pos0: usize,
9468 embed_ids: Option<&[u32]>,
9469 cap: usize,
9470 ) -> usize {
9471 if !crate::gpu::enabled_here()
9475 || std::env::var("CMF_GPU_CHUNK")
9476 .map(|v| v == "0")
9477 .unwrap_or(false)
9478 || b < 32
9479 || self.swa.is_some()
9480 || self.global_attn.is_some()
9481 || self.graph_attn_decline_reason().is_some()
9483 || self.o1_active()
9486 || self.attn_v_norm
9487 || (self.attn_scale - 1.0 / (self.head_dim as f32).sqrt()).abs() > 1e-9
9488 {
9489 return li0;
9490 }
9491 let Some(model) = self.model.clone() else {
9492 return li0;
9493 };
9494 let inv_freq = self.inv_freq.clone();
9495 let (nh, nkv, hd, hs) = (
9496 self.num_heads,
9497 self.num_kv_heads,
9498 self.head_dim,
9499 self.hidden_size,
9500 );
9501 let loop_end = if self.loop_final_norm {
9505 ((li0 / self.physical_layers) + 1) * self.physical_layers
9506 } else {
9507 self.num_layers
9508 };
9509 let mut layers: Vec<crate::gpu_metal::ChunkLayer> = Vec::new();
9510 let mut stored_at: Vec<usize> = Vec::new();
9511 for li in li0..self.num_layers.min(loop_end).min(cap) {
9512 let lw = &self.weights.layers[self.phys_layer(li)];
9513 if lw.attn_out_norm.is_some() || lw.ffn_out_norm.is_some() || lw.layer_scale.is_some() {
9514 break;
9515 }
9516 let AttnKind::Full {
9517 wq,
9518 wk,
9519 wv,
9520 wo,
9521 q_norm,
9522 k_norm,
9523 output_gate: false,
9524 softplus_gate: None,
9525 bias,
9526 } = &lw.attn
9527 else {
9528 break;
9529 };
9530 let FfnKind::Dense(d) = &lw.ffn else { break };
9531 if d.act != Act::Silu || !d.segs.is_empty() {
9532 break;
9533 }
9534 fn cw(t: &QTensor) -> Option<(usize, usize, usize, &[f32])> {
9539 t.q8_row_parts()
9540 .or_else(|| t.q4t_parts().map(|(i, r, c)| (i, r, c, &[][..])))
9541 .or_else(|| t.q4tp_parts().map(|(i, r, c)| (i, r, c, &[][..])))
9542 }
9543 let parts = (
9544 cw(wq),
9545 cw(wk),
9546 cw(wv),
9547 cw(wo),
9548 cw(&d.gate_proj),
9549 cw(&d.up_proj),
9550 cw(&d.down_proj),
9551 );
9552 let (Some(pq), Some(pk), Some(pv), Some(po), Some(pg), Some(pu), Some(pd)) = parts
9553 else {
9554 break;
9555 };
9556 let layer = &self.kv_cache.layers[li];
9557 if layer.mode != crate::kv_cache::KvMode::F32 || layer.o1.is_some() {
9558 break;
9559 }
9560 stored_at.push(layer.head_len(0));
9561 layers.push(crate::gpu_metal::ChunkLayer {
9562 model: &model,
9563 kv_id: self.graph_kv_id,
9564 layer: li,
9565 wq: pq,
9566 wk: pk,
9567 wv: pv,
9568 wo: po,
9569 gate: pg,
9570 up: pu,
9571 down: pd,
9572 input_norm: &lw.input_norm,
9573 post_norm: &lw.post_norm,
9574 bias: bias
9575 .as_ref()
9576 .map(|(a, bb, cc)| (a.as_slice(), bb.as_slice(), cc.as_slice())),
9577 q_norm: q_norm.as_deref(),
9578 k_norm: k_norm.as_deref(),
9579 inv_freq: &inv_freq,
9580 rd: self.rotary_dim,
9581 nh,
9582 nkv,
9583 hd,
9584 hs,
9585 inter: d.gate_proj.rows(),
9586 gemma: matches!(self.norm_style, cortiq_core::NormStyle::Gemma),
9587 late_qk_norm: self.qk_norm_after_rope,
9588 eps: self.rms_eps as f32,
9589 });
9590 }
9591 if layers.is_empty() {
9592 return li0;
9593 }
9594 let row = nkv * hd;
9595 let mut store: Vec<(Vec<f32>, Vec<f32>, Vec<f32>)> = stored_at
9596 .iter()
9597 .map(|&st| (vec![0f32; b * row], vec![0f32; b * row], vec![0f32; st + b]))
9598 .collect();
9599 let mut io: Vec<crate::gpu_metal::ChunkIo> = Vec::with_capacity(layers.len());
9600 for (i, (ok, ov, oi)) in store.iter_mut().enumerate() {
9601 let li = layers[i].layer;
9602 let layer = &self.kv_cache.layers[li];
9603 io.push(crate::gpu_metal::ChunkIo {
9604 cpu_stored: stored_at[i],
9605 cpu_k: (0..nkv).map(|g| layer.head_keys(g)).collect(),
9606 cpu_v: (0..nkv).map(|g| layer.head_values(g)).collect(),
9607 out_k: ok,
9608 out_v: ov,
9609 imp: oi,
9610 });
9611 }
9612 let n_run = layers.len();
9613 let last = layers.last().map(|l| l.layer + 1).unwrap_or(li0);
9614 let ep = embed_ids.and_then(|ids| {
9617 self.weights
9618 .embed_tokens
9619 .q8_row_parts()
9620 .map(|(idx, rows, _c, rs)| crate::gpu_metal::ChunkEmbed {
9621 idx,
9622 rows,
9623 row_scale: rs,
9624 ids,
9625 mult: self.embed_multiplier,
9626 })
9627 });
9628 if embed_ids.is_some() && ep.is_none() {
9629 return li0;
9630 }
9631 if !crate::gpu_metal::chunk_run_gpu(&layers, &mut io, h, b, pos0, ep.as_ref()) {
9632 return li0;
9633 }
9634 drop(io);
9635 drop(layers);
9636 for (i, (ok, ov, oi)) in store.iter().enumerate().take(n_run) {
9639 let li = li0 + i;
9640 let layer = &mut self.kv_cache.layers[li];
9641 for bi in 0..b {
9642 layer.append(
9643 &ok[bi * row..(bi + 1) * row],
9644 &ov[bi * row..(bi + 1) * row],
9645 &[],
9646 );
9647 }
9648 layer.accumulate_imp(oi);
9649 }
9650 last
9651 }
9652
9653 fn layer_is_local(&self, li: usize) -> bool {
9656 if let Some(layers) = &self.sliding_layers {
9657 return layers.get(li).copied().unwrap_or(false);
9658 }
9659 match self.swa {
9660 Some((_, pattern)) => (li + 1) % pattern.max(1) != 0,
9661 None => false,
9662 }
9663 }
9664
9665 fn layer_inv_freq(&self, li: usize) -> std::sync::Arc<Vec<f32>> {
9668 if self.layer_is_local(li) {
9669 if let Some(f) = &self.inv_freq_local {
9670 return f.clone();
9671 }
9672 } else if let Some(f) = &self.inv_freq_global {
9673 return f.clone();
9674 }
9675 self.inv_freq.clone()
9676 }
9677
9678 fn layer_window(&self, li: usize) -> Option<usize> {
9680 self.swa
9681 .and_then(|(w, _)| self.layer_is_local(li).then_some(w))
9682 }
9683
9684 fn layer_num_heads(&self, li: usize) -> usize {
9685 self.attention_heads_per_layer
9686 .as_ref()
9687 .and_then(|v| v.get(li).copied())
9688 .unwrap_or(self.num_heads)
9689 }
9690
9691 fn layer_rope_scale(&self, li: usize) -> f32 {
9692 if self.layer_is_local(li) {
9693 self.rope_scale_local
9694 } else {
9695 self.rope_scale
9696 }
9697 }
9698
9699 fn layer_geom(&self, li: usize) -> (usize, usize, usize) {
9702 if !self.layer_is_local(li) {
9703 if let Some((ghd, gkv)) = self.global_attn {
9704 return (gkv, ghd, ghd);
9705 }
9706 }
9707 (
9708 self.layer_num_kv_heads(li),
9709 self.head_dim,
9710 if self.layer_is_local(li) {
9711 self.rotary_dim_local.unwrap_or(self.rotary_dim)
9712 } else {
9713 self.rotary_dim
9714 },
9715 )
9716 }
9717
9718 fn layer_num_kv_heads(&self, li: usize) -> usize {
9721 self.kv_heads_per_layer
9722 .as_ref()
9723 .and_then(|v| v.get(self.phys_layer(li)).copied())
9724 .unwrap_or(self.num_kv_heads)
9725 }
9726
9727 fn layer_v_dim(&self, li: usize) -> usize {
9729 let (_, hd, _) = self.layer_geom(li);
9730 self.v_head_dim.unwrap_or(hd).min(hd)
9731 }
9732
9733 pub fn set_attn_geometry(
9742 &mut self,
9743 kv_heads_per_layer: Option<Vec<usize>>,
9744 v_head_dim: Option<usize>,
9745 ) -> Result<(), String> {
9746 if kv_heads_per_layer.is_some() || v_head_dim.is_some() {
9747 if self.global_attn.is_some() {
9748 return Err(
9749 "per-layer KV heads / v_head_dim cannot combine with Gemma-4 global \
9750 attention geometry"
9751 .into(),
9752 );
9753 }
9754 if self
9755 .weights
9756 .layers
9757 .iter()
9758 .any(|lw| matches!(lw.attn, AttnKind::Mla(_)))
9759 {
9760 return Err("per-layer KV heads / v_head_dim cannot combine with MLA".into());
9761 }
9762 }
9763 if let Some(vd) = v_head_dim {
9764 if vd == 0 || vd > self.head_dim {
9765 return Err(format!(
9766 "v_head_dim {vd} must be in 1..={} (head_dim)",
9767 self.head_dim
9768 ));
9769 }
9770 }
9771 if let Some(v) = &kv_heads_per_layer {
9772 if v.len() != self.physical_layers {
9773 return Err(format!(
9774 "kv_heads_per_layer has {} entries, expected {} layers",
9775 v.len(),
9776 self.physical_layers
9777 ));
9778 }
9779 for (li, &nkv) in v.iter().enumerate() {
9780 let is_attn = matches!(
9781 self.weights.layers.get(li).map(|lw| &lw.attn),
9782 Some(AttnKind::Full { .. }) | None
9783 );
9784 if !is_attn {
9785 continue;
9786 }
9787 let nh = self
9788 .attention_heads_per_layer
9789 .as_ref()
9790 .and_then(|h| h.get(li).copied())
9791 .unwrap_or(self.num_heads);
9792 if nkv == 0 || nh % nkv != 0 {
9793 return Err(format!(
9794 "layer {li}: {nkv} KV heads must be nonzero and divide {nh} Q heads"
9795 ));
9796 }
9797 }
9798 }
9799 self.kv_heads_per_layer = kv_heads_per_layer;
9800 self.v_head_dim = v_head_dim.filter(|&vd| vd != self.head_dim);
9801 if self.kv_heads_per_layer.is_some() {
9802 for li in 0..self.kv_cache.layers.len() {
9803 let full = matches!(
9804 self.weights
9805 .layers
9806 .get(self.phys_layer(li))
9807 .map(|lw| &lw.attn),
9808 Some(AttnKind::Full { .. })
9809 );
9810 let nkv = self.layer_num_kv_heads(li);
9811 let cache = &self.kv_cache.layers[li];
9812 if full && (cache.num_kv_heads != nkv || cache.head_dim != self.head_dim) {
9813 let sinks = cache.sinks.clone();
9814 self.kv_cache.layers[li] =
9815 crate::kv_cache::LayerKvCache::new(nkv, self.head_dim);
9816 self.kv_cache.layers[li].sinks = sinks;
9817 }
9818 }
9819 }
9820 Ok(())
9821 }
9822
9823 pub fn set_layer_sinks(&mut self, phys: usize, sinks: Vec<f32>) -> Result<(), String> {
9827 let Some(lw) = self.weights.layers.get(phys) else {
9828 return Err(format!("sinks for layer {phys}: no such layer"));
9829 };
9830 if !matches!(lw.attn, AttnKind::Full { .. }) {
9831 return Err(format!(
9832 "sinks for layer {phys}: only softmax (Full) attention layers take sinks"
9833 ));
9834 }
9835 let nh = self
9836 .attention_heads_per_layer
9837 .as_ref()
9838 .and_then(|h| h.get(phys).copied())
9839 .unwrap_or(self.num_heads);
9840 if sinks.len() != nh {
9841 return Err(format!(
9842 "sinks for layer {phys}: {} values, expected one per Q head ({nh})",
9843 sinks.len()
9844 ));
9845 }
9846 if let Some(bad) = sinks.iter().find(|v| !v.is_finite()) {
9847 return Err(format!("sinks for layer {phys}: non-finite value {bad}"));
9848 }
9849 for li in 0..self.kv_cache.layers.len() {
9850 if self.phys_layer(li) == phys {
9851 self.kv_cache.layers[li].sinks = Some(sinks.clone());
9852 }
9853 }
9854 Ok(())
9855 }
9856
9857 pub fn graph_attn_decline_reason(&self) -> Option<&'static str> {
9867 if self.kv_heads_per_layer.is_some() {
9868 return Some("per-layer KV head counts");
9869 }
9870 if self.v_head_dim.is_some_and(|vd| vd != self.head_dim) {
9871 return Some("V heads narrower than Q/K heads");
9872 }
9873 if self.kv_cache.layers.iter().any(|l| l.sinks.is_some()) {
9874 return Some("learned attention sinks");
9875 }
9876 if self.swa.is_some() || self.sliding_layers.is_some() {
9877 return Some("sliding-window layers");
9878 }
9879 None
9880 }
9881
9882 pub fn wgpu_graph_attn_decline(&self) -> Option<&'static str> {
9889 self.graph_attn_decline_reason()?;
9890 if self.global_attn.is_some() {
9891 return Some("per-layer head width (Gemma-4 global layers) with per-layer geometry");
9892 }
9893 if self.attention_heads_per_layer.is_some() {
9894 return Some("per-layer Q head counts with per-layer geometry");
9895 }
9896 if self.attn_v_norm {
9897 return Some("V norm with per-layer geometry");
9898 }
9899 if (0..self.num_layers).any(|li| self.layer_rope_scale(li) != 1.0) {
9900 return Some("scaled RoPE positions with per-layer geometry");
9901 }
9902 if self.weights.layers.iter().any(|lw| {
9903 lw.attn_out_norm.is_some() || lw.ffn_out_norm.is_some() || lw.layer_scale.is_some()
9904 }) {
9905 return Some("sandwich norms / layer scale with per-layer geometry");
9906 }
9907 if self.weights.layers.iter().any(|lw| {
9908 matches!(
9909 &lw.attn,
9910 AttnKind::Full {
9911 output_gate: true,
9912 ..
9913 }
9914 )
9915 }) && self.v_head_dim.is_some()
9916 {
9917 return Some("gated attention with V narrower than K");
9918 }
9919 if (0..self.num_layers).any(|li| {
9920 self.layer_is_local(li)
9921 && self.inv_freq_local.is_none()
9922 && self.rotary_dim_local.is_some_and(|r| r != self.rotary_dim)
9923 }) {
9924 return Some("local rotary width without a local RoPE table");
9925 }
9926 None
9927 }
9928
9929 fn graph_attn_geom(&self, li: usize) -> Option<crate::gpu::GraphAttnGeom<'_>> {
9934 self.graph_attn_decline_reason()?;
9935 let (nkv, _hd, rd) = self.layer_geom(li);
9936 let invf: &[f32] = if self.layer_is_local(li) {
9937 match &self.inv_freq_local {
9938 Some(f) => f.as_slice(),
9939 None => self.inv_freq.as_slice(),
9940 }
9941 } else {
9942 match &self.inv_freq_global {
9943 Some(f) => f.as_slice(),
9944 None => self.inv_freq.as_slice(),
9945 }
9946 };
9947 Some(crate::gpu::GraphAttnGeom {
9948 nkv,
9949 dv: self.layer_v_dim(li),
9950 rd,
9951 invf,
9952 window: self.layer_window(li),
9953 sink: self.kv_cache.layers[li].sinks.as_deref(),
9954 })
9955 }
9956
9957 #[cfg(feature = "gpu")]
9966 fn pull_lagging_host_kv(&mut self, from: usize, upto: usize, position: usize) {
9967 let kv_id = self.graph_kv_id;
9968 for li in from..upto.min(self.num_layers) {
9969 if !matches!(
9970 self.weights.layers[self.phys_layer(li)].attn,
9971 AttnKind::Full { .. }
9972 ) {
9973 continue;
9974 }
9975 let host = self.kv_cache.layers[li].seq_len;
9976 if host >= position {
9977 continue;
9978 }
9979 let Some(dev) = crate::gpu::graph_kv_stored(kv_id, li) else {
9980 continue;
9981 };
9982 let to = dev.min(position);
9983 if to <= host {
9984 continue;
9985 }
9986 let (nkv, hd) = {
9987 let c = &self.kv_cache.layers[li];
9988 (c.num_kv_heads, c.head_dim)
9989 };
9990 let Some((k, v, first_valid)) =
9991 crate::gpu::graph_kv_pull_host(kv_id, li, host, to, nkv, hd)
9992 else {
9993 continue;
9994 };
9995 let need_from = match self.layer_window(li) {
9998 Some(w) => host.max((position + 1).saturating_sub(w)),
9999 None => host,
10000 };
10001 if first_valid > need_from {
10002 tracing::warn!(
10003 "layer {li}: device KV rows {host}..{to} no longer resident \
10004 (from {first_valid}); host attention will miss them"
10005 );
10006 }
10007 let row = nkv * hd;
10008 let cache = &mut self.kv_cache.layers[li];
10009 for p in 0..to - host {
10010 cache.append(&k[p * row..(p + 1) * row], &v[p * row..(p + 1) * row], &[]);
10011 }
10012 }
10013 }
10014
10015 fn note_graph_decline(&self, site: &'static str, reason: &'static str) {
10018 let mut seen = self.graph_declines.borrow_mut();
10019 if !seen.iter().any(|&(s, r)| s == site && r == reason) {
10020 tracing::warn!("{site} declined: {reason} (CPU attention path)");
10021 seen.push((site, reason));
10022 }
10023 }
10024
10025 pub fn graph_declines(&self) -> Vec<String> {
10028 self.graph_declines
10029 .borrow()
10030 .iter()
10031 .map(|(site, reason)| format!("{site} declined: {reason} (CPU attention path)"))
10032 .collect()
10033 }
10034
10035 fn dump_layer_row(&self, pos: usize, li: usize, row: &[f32]) {
10044 let Some(dir) = &self.layer_dump else {
10045 return;
10046 };
10047 let mut bytes = Vec::with_capacity(row.len() * 4);
10048 for v in row {
10049 bytes.extend_from_slice(&v.to_le_bytes());
10050 }
10051 let path = dir.join(format!("p{pos:06}_l{li:02}.f32"));
10052 if let Err(e) = std::fs::create_dir_all(dir).and_then(|_| std::fs::write(&path, &bytes)) {
10053 use std::sync::atomic::{AtomicBool, Ordering};
10054 static SAID: AtomicBool = AtomicBool::new(false);
10055 if !SAID.swap(true, Ordering::Relaxed) {
10056 tracing::error!("CMF_LAYER_DUMP: cannot write {}: {e}", path.display());
10057 }
10058 }
10059 }
10060
10061 fn mimo_moe_prepare(&mut self) {
10064 if !self.mimo_moe.is_undecided() {
10065 return;
10066 }
10067 let slot = {
10068 let layers: Vec<(usize, &MoeFfn)> = (0..self.num_layers)
10069 .filter_map(
10070 |li| match &self.weights.layers.get(self.phys_layer(li))?.ffn {
10071 FfnKind::Moe(m) => Some((li, m)),
10072 _ => None,
10073 },
10074 )
10075 .collect();
10076 if layers.is_empty()
10079 || self.physical_layers != self.num_layers
10080 || self.gpu_plan.is_some()
10081 {
10082 crate::mimo_moe::Slot::Off
10083 } else {
10084 let graph_prefix =
10088 self.wgpu_graph_attn_decline().is_none() && crate::gpu::wgpu_graph_default();
10089 crate::mimo_moe::Slot::decide(&layers, self.num_layers, graph_prefix)
10090 }
10091 };
10092 self.mimo_moe = slot;
10093 }
10094
10095 #[cfg(test)]
10096 pub(crate) fn test_graph_kv_id(&self) -> u64 {
10097 self.graph_kv_id
10098 }
10099
10100 pub(crate) fn mimo_graph_layer_rows(
10104 &mut self,
10105 li: usize,
10106 h: &mut [f32],
10107 positions: &[usize],
10108 ) -> crate::gpu::BatchGraphOutcome {
10109 use crate::gpu::BatchGraphOutcome as Out;
10110 let b = positions.len();
10111 if !(1..=4).contains(&b)
10112 || h.len() != b * self.hidden_size
10113 || !self.mimo_moe.is_dynamic(li, true)
10114 || !crate::gpu::enabled_here()
10115 || !crate::gpu::wgpu_active()
10116 || self.o1_active()
10117 || self.physical_layers != self.num_layers
10118 || std::env::var("CMF_GPU_WGPU_GRAPH").as_deref() == Ok("0")
10122 || std::env::var("CMF_MIMO_ATTN_GRAPH").as_deref() == Ok("0")
10123 || self.wgpu_graph_attn_decline().is_some()
10124 {
10125 return Out::Declined;
10126 }
10127 let attn_started = std::time::Instant::now();
10128 let outcome = {
10129 let lw = &self.weights.layers[li];
10130 if lw.attn_out_norm.is_some() || lw.ffn_out_norm.is_some() || lw.layer_scale.is_some() {
10131 return Out::Declined;
10132 }
10133 let FfnKind::Moe(m) = &lw.ffn else {
10134 return Out::Declined;
10135 };
10136 let AttnKind::Full {
10137 wq,
10138 wk,
10139 wv,
10140 wo,
10141 q_norm,
10142 k_norm,
10143 output_gate,
10144 softplus_gate,
10145 bias,
10146 } = &lw.attn
10147 else {
10148 return Out::Declined;
10149 };
10150 if *output_gate || softplus_gate.is_some() {
10151 return Out::Declined;
10152 }
10153 let Some(model) = wq.graph_weight().map(|(m, ..)| m.clone()).or_else(|| {
10154 m.experts
10155 .first()?
10156 .gate_proj
10157 .mapped_q4tp()
10158 .map(|(m, _)| m.clone())
10159 }) else {
10160 return Out::Declined;
10161 };
10162 fn gw<'a>(
10163 t: &'a QTensor,
10164 owner: &std::sync::Arc<cortiq_core::CmfModel>,
10165 ) -> Option<crate::gpu::GraphW<'a>> {
10166 if let Some((m, idx, kind, rs)) = t.graph_weight() {
10167 if m.uid() != owner.uid() || t.has_prism_contract() {
10168 return None;
10169 }
10170 return Some(crate::gpu::GraphW {
10171 idx,
10172 kind,
10173 row_scale: rs,
10174 data: &[],
10175 prism: crate::gpu::GraphPrismOp::None,
10176 affine: false,
10177 });
10178 }
10179 t.as_f32().map(|data| crate::gpu::GraphW {
10180 idx: 0,
10181 kind: 4,
10182 row_scale: &[],
10183 data,
10184 prism: crate::gpu::GraphPrismOp::None,
10185 affine: false,
10186 })
10187 }
10188 let (Some(q), Some(k), Some(v), Some(o)) = (
10189 gw(wq, &model),
10190 gw(wk, &model),
10191 gw(wv, &model),
10192 gw(wo, &model),
10193 ) else {
10194 return Out::Declined;
10195 };
10196 let layer = crate::gpu::GraphLayer {
10197 input_norm: &lw.input_norm,
10198 post_norm: &lw.post_norm,
10199 ffn: crate::gpu::GraphFfn::AttentionOnly,
10200 attn: crate::gpu::GraphAttn::Full {
10201 wq: q,
10202 wk: k,
10203 wv: v,
10204 wo: o,
10205 q_norm: q_norm.as_deref(),
10206 k_norm: k_norm.as_deref(),
10207 late_qk_norm: self.qk_norm_after_rope,
10208 bias: bias
10209 .as_ref()
10210 .map(|(q, k, v)| (q.as_slice(), k.as_slice(), v.as_slice())),
10211 output_gate: false,
10212 cpu_k: self.kv_cache.layers[li].k_heads(),
10213 cpu_v: self.kv_cache.layers[li].v_heads(),
10214 geom: self.graph_attn_geom(li),
10215 },
10216 };
10217 let (nkv, hd, rd) = self.layer_geom(li);
10218 crate::gpu::forward_batch_graph_at(
10219 &model,
10220 self.graph_kv_id,
10221 li,
10222 &[layer],
10223 &self.inv_freq,
10224 h,
10225 self.layer_num_heads(li),
10226 nkv,
10227 hd,
10228 rd,
10229 self.hidden_size,
10230 1,
10231 positions,
10232 self.kv_cache.max_seq_len,
10233 self.norm_style == cortiq_core::NormStyle::Gemma,
10234 self.rms_eps as f32,
10235 self.attn_scale,
10236 b,
10237 &[],
10238 self.o1_epoch,
10239 None,
10240 None,
10241 )
10242 };
10243 match outcome {
10244 Out::Completed => {}
10245 Out::Declined => return Out::Declined,
10246 Out::Failed => {
10247 self.graph_failed
10248 .store(true, std::sync::atomic::Ordering::Relaxed);
10249 return Out::Failed;
10250 }
10251 }
10252 let attn_ns = attn_started.elapsed().as_nanos() as u64;
10253 let hs = self.hidden_size;
10254 let lw = &self.weights.layers[li];
10255 let FfnKind::Moe(m) = &lw.ffn else {
10256 unreachable!()
10257 };
10258 let mut post = vec![0.0; h.len()];
10259 for (x, y) in h.chunks_exact(hs).zip(post.chunks_exact_mut(hs)) {
10260 inference::rms_norm_into(x, &lw.post_norm, self.rms_eps, self.norm_style, y);
10261 }
10262 let mut ffn = if b == 1 {
10263 moe_ffn_banked(&mut self.mimo_moe, li, m, &post, self.pool.as_deref())
10264 } else {
10265 moe_ffn_banked_rows(
10266 &mut self.mimo_moe,
10267 li,
10268 m,
10269 &post,
10270 b,
10271 hs,
10272 self.pool.as_deref(),
10273 )
10274 };
10275 for (x, &f) in h.iter_mut().zip(&ffn) {
10276 *x += f;
10277 }
10278 attention::recycle_buf(&mut ffn);
10279 if self.layer_dump.is_some() {
10280 for (&pos, row) in positions.iter().zip(h.chunks_exact(hs)) {
10281 self.dump_layer_row(pos, li, row);
10282 }
10283 }
10284 crate::mimo_moe::note_attention_graph(b, attn_ns);
10285 Out::Completed
10286 }
10287
10288 fn layer_attn_plain(&self, li: usize) -> bool {
10289 self.kv_heads_per_layer.is_none()
10290 && self.v_head_dim.is_none()
10291 && self.global_attn.is_none()
10292 && self.layer_window(li).is_none()
10293 && self.kv_cache.layers[li].sinks.is_none()
10294 }
10295
10296 fn forward_layers(
10298 &mut self,
10299 hidden: &[f32],
10300 position: usize,
10301 task_mask: Option<&TaskMask>,
10302 ) -> Vec<f32> {
10303 let out = self.forward_layers_upto(hidden, position, task_mask, None);
10304 self.o1_progress();
10305 out
10306 }
10307
10308 pub fn embed_id(&self, id: u32) -> Vec<f32> {
10316 self.embed_single(id)
10317 }
10318
10319 pub fn split_supported(&self) -> Result<(), String> {
10323 if self.dsv4.is_some() {
10324 return Err(
10325 "network split: DeepSeek-V4 runs its own fused stack (not splittable yet)".into(),
10326 );
10327 }
10328 if self.dsv41.is_some() {
10329 return Err(
10330 "network split: DeepSeek-V4.1 owns the shared CED/CSA2 state (not splittable)"
10331 .into(),
10332 );
10333 }
10334 if self.qwen4_exp.is_some() {
10335 return Err(
10336 "network split: Qwen3.8-Flash-Next hyper/QSA stack is not splittable yet".into(),
10337 );
10338 }
10339 if self.g3n.is_some() {
10340 return Err(
10341 "network split: Gemma-3n runs its own AltUp stack (not splittable yet)".into(),
10342 );
10343 }
10344 Ok(())
10345 }
10346
10347 pub fn forward_span(
10352 &mut self,
10353 hidden: &[f32],
10354 position: usize,
10355 from: usize,
10356 upto: usize,
10357 task_mask: Option<&TaskMask>,
10358 ) -> Result<Vec<f32>, String> {
10359 self.split_supported()?;
10360 if from > upto || upto >= self.num_layers {
10361 return Err(format!(
10362 "forward_span: layer range {from}..={upto} outside 0..{}",
10363 self.num_layers
10364 ));
10365 }
10366 if hidden.len() != self.hidden_size {
10367 return Err(format!(
10368 "forward_span: hidden len {} ≠ hidden_size {}",
10369 hidden.len(),
10370 self.hidden_size
10371 ));
10372 }
10373 let out = self.forward_layers_span(hidden, position, task_mask, from, Some(upto));
10374 self.o1_progress();
10375 if self
10376 .graph_failed
10377 .swap(false, std::sync::atomic::Ordering::Relaxed)
10378 {
10379 self.cancel
10380 .store(false, std::sync::atomic::Ordering::Relaxed);
10381 self.clear_sequence_state();
10382 return Err("forward_span: deferred O(1) transition failed".into());
10383 }
10384 Ok(out)
10385 }
10386
10387 pub fn logits_from_hidden(&mut self, hidden: &[f32]) -> Vec<f32> {
10390 let normed = inference::rms_norm(
10391 hidden,
10392 &self.weights.final_norm,
10393 self.rms_eps,
10394 self.norm_style,
10395 );
10396 self.lm_head_forward(&normed)
10397 }
10398
10399 pub fn sample_next(&mut self, logits: &[f32], past_tokens: &[u32]) -> u32 {
10401 sampler::sample_with_scratch(
10402 logits,
10403 &self.sampler_config,
10404 past_tokens,
10405 &mut self.rng,
10406 &mut self.sampler_scratch,
10407 )
10408 }
10409
10410 pub fn reset_session(&mut self) {
10412 self.clear_sequence_state();
10413 }
10414
10415 pub fn prefill_span_ids(
10421 &mut self,
10422 ids: &[u32],
10423 start_pos: usize,
10424 upto: usize,
10425 task_mask: Option<&TaskMask>,
10426 ) -> Result<Vec<f32>, String> {
10427 self.split_supported()?;
10428 if upto >= self.num_layers {
10429 return Err(format!(
10430 "prefill_span_ids: upto {upto} outside 0..{}",
10431 self.num_layers
10432 ));
10433 }
10434 if self.can_prefill_batched() && !self.graph_prefill_preferred() {
10438 let out =
10439 self.prefill_batch_span(PrefillIn::Ids(ids), start_pos, task_mask, 0, upto + 1);
10440 self.check_o1_progress_failure("prefill_span_ids")?;
10441 Ok(out)
10442 } else {
10443 let hs = self.hidden_size;
10444 let mut out = Vec::with_capacity(ids.len() * hs);
10445 for (i, &id) in ids.iter().enumerate() {
10446 let emb = self.embed_id(id);
10447 out.extend_from_slice(&self.forward_span(
10448 &emb,
10449 start_pos + i,
10450 0,
10451 upto,
10452 task_mask,
10453 )?);
10454 }
10455 Ok(out)
10456 }
10457 }
10458
10459 pub fn prefill_span_hidden(
10462 &mut self,
10463 hidden: &[f32],
10464 start_pos: usize,
10465 from: usize,
10466 upto: usize,
10467 task_mask: Option<&TaskMask>,
10468 ) -> Result<Vec<f32>, String> {
10469 self.split_supported()?;
10470 let hs = self.hidden_size;
10471 if hidden.is_empty() || hidden.len() % hs != 0 {
10472 return Err(format!(
10473 "prefill_span_hidden: {} floats is not a multiple of hidden {hs}",
10474 hidden.len()
10475 ));
10476 }
10477 if from > upto || upto >= self.num_layers {
10478 return Err(format!(
10479 "prefill_span_hidden: layer range {from}..={upto} outside 0..{}",
10480 self.num_layers
10481 ));
10482 }
10483 if self.can_prefill_batched() && !self.graph_prefill_preferred() {
10484 let out = self.prefill_batch_span(
10485 PrefillIn::Hidden(hidden),
10486 start_pos,
10487 task_mask,
10488 from,
10489 upto + 1,
10490 );
10491 self.check_o1_progress_failure("prefill_span_hidden")?;
10492 Ok(out)
10493 } else {
10494 let b = hidden.len() / hs;
10495 let mut out = Vec::with_capacity(hidden.len());
10496 for i in 0..b {
10497 let h = self.forward_span(
10498 &hidden[i * hs..(i + 1) * hs],
10499 start_pos + i,
10500 from,
10501 upto,
10502 task_mask,
10503 )?;
10504 out.extend_from_slice(&h);
10505 }
10506 Ok(out)
10507 }
10508 }
10509
10510 fn try_token_graph_wgpu(
10514 &self,
10515 hidden: &[f32],
10516 position: usize,
10517 logits_out: &mut Vec<f32>,
10518 layers_run: &mut usize,
10519 ) -> Option<Result<Vec<f32>, ()>> {
10520 self.try_token_graph_wgpu_steps(
10521 hidden,
10522 position,
10523 logits_out,
10524 1,
10525 None,
10526 Some(layers_run),
10527 0,
10528 self.num_layers,
10529 )
10530 }
10531
10532 fn try_token_graph_wgpu_span(
10536 &self,
10537 hidden: &[f32],
10538 position: usize,
10539 logits_out: &mut Vec<f32>,
10540 from: usize,
10541 upto_excl: usize,
10542 layers_run: &mut usize,
10543 ) -> Option<Result<Vec<f32>, ()>> {
10544 self.try_token_graph_wgpu_steps(
10545 hidden,
10546 position,
10547 logits_out,
10548 1,
10549 None,
10550 Some(layers_run),
10551 from,
10552 upto_excl,
10553 )
10554 }
10555
10556 fn try_multi_burst(&self, t_next: u32, position: usize, k: usize) -> Option<Vec<u32>> {
10560 if self.o1_active() || self.attn_softcap > 0.0 {
10561 return None;
10562 }
10563 if let Some(reason) = self.wgpu_graph_attn_decline() {
10566 self.note_graph_decline("wgpu multi-burst", reason);
10567 return None;
10568 }
10569 let graph_on = crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Decode);
10570 if !graph_on || self.graph_refused() {
10571 return None;
10578 }
10579 let emb = self.embed_single(t_next);
10580 let mut lg = Vec::new();
10581 let mut ids = Vec::new();
10582 match self.try_token_graph_wgpu_steps(
10583 &emb,
10584 position,
10585 &mut lg,
10586 k,
10587 Some(&mut ids),
10588 None,
10589 0,
10590 self.num_layers,
10591 ) {
10592 Some(Ok(_)) => {}
10593 Some(Err(())) => {
10594 self.graph_failed
10599 .store(true, std::sync::atomic::Ordering::Relaxed);
10600 return None;
10601 }
10602 None => return None,
10603 }
10604 (ids.len() == k).then_some(ids)
10605 }
10606
10607 fn try_token_graph_wgpu_steps(
10611 &self,
10612 hidden: &[f32],
10613 position: usize,
10614 logits_out: &mut Vec<f32>,
10615 steps: usize,
10616 ids_out: Option<&mut Vec<u32>>,
10617 layers_run: Option<&mut usize>,
10618 from: usize,
10619 upto_excl: usize,
10620 ) -> Option<Result<Vec<f32>, ()>> {
10621 let upto_excl = match self.mimo_moe.graph_prefix_end() {
10624 Some(end) if end < upto_excl => {
10625 if steps != 1 || layers_run.is_none() || from >= end {
10626 return None;
10627 }
10628 end
10629 }
10630 _ => upto_excl,
10631 };
10632 let o1_gpu = std::env::var("CMF_O1_GPU").as_deref() == Ok("1");
10635 if (self.o1_active() && !o1_gpu) || self.attn_softcap > 0.0 {
10636 return None;
10640 }
10641 if let Some(reason) = self.wgpu_graph_attn_decline() {
10648 self.note_graph_decline("wgpu token graph", reason);
10649 return None;
10650 }
10651 let o1_views: Vec<Option<Vec<crate::nystrom::O1DeviceView<'_>>>> = (from..upto_excl)
10656 .map(|li| {
10657 if !o1_gpu {
10658 return None;
10659 }
10660 self.kv_cache.layers[self.phys_layer(li)].o1_views()
10661 })
10662 .collect();
10663 if self.o1_active() && o1_gpu {
10664 let want: usize = (from..upto_excl)
10667 .filter(|li| self.kv_cache.layers[self.phys_layer(*li)].o1.is_some())
10668 .count();
10669 let have = o1_views.iter().filter(|v| v.is_some()).count();
10670 if want == 0 || have != want {
10671 use std::sync::atomic::{AtomicUsize, Ordering};
10681 static LAST: AtomicUsize = AtomicUsize::new(usize::MAX);
10682 let code = have * 1000 + want;
10683 if LAST.swap(code, Ordering::Relaxed) != code {
10684 tracing::warn!(
10685 "o1 graph: {have} of {want} layers sealed — per-op until all seal"
10686 );
10687 }
10688 return None;
10689 }
10690 }
10691 let nh = self.num_heads;
10692 let (nkv, hd, rd) = self.layer_geom(0);
10693 let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
10694 let mut layers = Vec::with_capacity(upto_excl - from);
10695 let mut model = None;
10696 let dbg = std::env::var("CMF_GRAPH_DEBUG").is_ok();
10697 fn gw(t: &QTensor) -> Option<crate::gpu::GraphW<'_>> {
10698 if let Some((m, i, kind, rs)) = t
10699 .graph_weight()
10700 .or_else(|| t.graph_weight_descriptor())
10701 {
10702 let name = &m.tensors[i].name;
10703 let prism = if crate::prism::is_inverse_embedding(m, name) {
10704 crate::gpu::GraphPrismOp::InverseEmbedding
10705 } else if crate::prism::is_forward_weight(m, name) {
10706 crate::gpu::GraphPrismOp::Forward
10707 } else {
10708 crate::gpu::GraphPrismOp::None
10709 };
10710 return Some(crate::gpu::GraphW {
10711 idx: i,
10712 kind,
10713 row_scale: rs,
10714 data: &[],
10715 prism,
10716 affine: crate::prism::is_affine_target(m, name),
10717 });
10718 }
10719 match t.as_f32() {
10721 Some(d) => Some(crate::gpu::GraphW {
10722 idx: 0,
10723 kind: 4,
10724 row_scale: &[],
10725 data: d,
10726 prism: crate::gpu::GraphPrismOp::None,
10727 affine: false,
10728 }),
10729 None => {
10730 if std::env::var_os("CMF_BATCH_DEBUG").is_some() {
10731 eprintln!("batch graph: weight has no graph/f32 representation");
10732 }
10733 None
10734 }
10735 }
10736 }
10737 for li in from..upto_excl {
10738 let lw = &self.weights.layers[self.phys_layer(li)];
10739 if dbg {
10740 let ak = match &lw.attn {
10741 AttnKind::Mla(_) => "Mla".into(),
10742 AttnKind::Full {
10743 output_gate, bias, ..
10744 } => format!("Full gate={output_gate} bias={}", bias.is_some()),
10745 AttnKind::LinearGdn(_) => "LinearGdn".into(),
10746 AttnKind::Kda(_) => "Kda".into(),
10747 AttnKind::Linear(_) => "Linear".into(),
10748 AttnKind::ShortConv(_) => "ShortConv".into(),
10749 AttnKind::Bounded(_) => "Bounded".into(),
10750 };
10751 let fk = match &lw.ffn {
10752 FfnKind::Dense(_) => "Dense",
10753 FfnKind::Moe(_) => "Moe",
10754 FfnKind::DenseMoe(_) => "DenseMoe",
10755 };
10756 eprintln!("graph L{li}: attn={ak} ffn={fk}");
10757 }
10758 let gffn = match &lw.ffn {
10759 FfnKind::DenseMoe(_) => return None, FfnKind::Dense(d) if !d.segs.is_empty() => return None,
10763 FfnKind::Dense(d) => crate::gpu::GraphFfn::Dense {
10764 gate: gw(&d.gate_proj)?,
10765 up: gw(&d.up_proj)?,
10766 down: gw(&d.down_proj)?,
10767 },
10768 FfnKind::Moe(m) => {
10769 if m.route_tau.is_some() || m.mask.is_some() {
10777 return None;
10778 }
10779 let shared = m.shared.as_ref();
10780 let has_shared = shared.is_some();
10781 let shared_gated = matches!(shared, Some((_, Some(_))));
10782 let sgate = match shared {
10783 Some((_, Some(sg))) => gw(sg)?,
10784 _ => gw(&m.router)?,
10788 };
10789 let router = gw(&m.router)?;
10790 if router.prism != crate::gpu::GraphPrismOp::None
10796 || sgate.prism != crate::gpu::GraphPrismOp::None
10797 || router.affine
10798 || sgate.affine
10799 {
10800 tracing::warn!(
10801 "resident MoE declined: Prism/affine router or shared gate transform is not implemented"
10802 );
10803 return None;
10804 }
10805 let inter = m.experts.first()?.gate_proj.rows();
10806 let mut experts = Vec::with_capacity(m.experts.len() + 1);
10807 let mut q4tp: Option<bool> = None;
10810 let mut gu_q2: Option<bool> = None;
10813 for e in m.experts.iter().chain(shared.map(|(se, _)| se)) {
10814 if !matches!(e.act, Act::Silu)
10815 || e.gate_proj.rows() != inter
10816 || e.up_proj.rows() != inter
10817 {
10818 return None;
10819 }
10820 for expert_weight in [&e.gate_proj, &e.up_proj, &e.down_proj] {
10825 let Some((em, ei, _, _)) = expert_weight
10826 .graph_weight()
10827 .or_else(|| expert_weight.graph_weight_descriptor())
10828 else {
10829 return None;
10830 };
10831 let name = &em.tensors[ei].name;
10832 if crate::prism::is_forward_weight(em, name)
10833 || crate::prism::is_inverse_embedding(em, name)
10834 || crate::prism::is_affine_target(em, name)
10835 {
10836 tracing::warn!(
10837 "resident MoE declined: expert Prism/affine transform is not implemented"
10838 );
10839 return None;
10840 }
10841 }
10842 let (mm, gi, ui, di, is_p, is_q2) = match e.gate_proj.mapped_q4t() {
10843 Some((mm, gi)) => (
10844 mm,
10845 gi,
10846 e.up_proj.mapped_q4t()?.1,
10847 e.down_proj.mapped_q4t()?.1,
10848 false,
10849 false,
10850 ),
10851 None => match e.gate_proj.mapped_q2tp() {
10852 Some((mm, gi)) => (
10853 mm,
10854 gi,
10855 e.up_proj.mapped_q2tp()?.1,
10856 e.down_proj.mapped_q4tp()?.1,
10857 true,
10858 true,
10859 ),
10860 None => {
10861 let (mm, gi) = e.gate_proj.mapped_q4tp()?;
10862 (
10863 mm,
10864 gi,
10865 e.up_proj.mapped_q4tp()?.1,
10866 e.down_proj.mapped_q4tp()?.1,
10867 true,
10868 false,
10869 )
10870 }
10871 },
10872 };
10873 if *q4tp.get_or_insert(is_p) != is_p || *gu_q2.get_or_insert(is_q2) != is_q2
10874 {
10875 tracing::warn!(
10881 "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."
10882 );
10883 return None;
10884 }
10885 model.get_or_insert_with(|| mm.clone());
10886 experts.push((gi, ui, di));
10887 }
10888 crate::gpu::GraphFfn::Moe {
10889 router,
10890 shared_gate: sgate,
10891 experts,
10892 n_exp: m.experts.len(),
10893 top_k: std::env::var("CMF_TOPK_PROBE")
10899 .ok()
10900 .and_then(|v| v.parse::<usize>().ok())
10901 .filter(|k| *k > 0 && *k <= m.top_k)
10902 .unwrap_or(m.top_k),
10903 inter,
10904 norm_topk: m.norm_topk_prob,
10905 q4tp: q4tp?,
10906 gu_q2: gu_q2.unwrap_or(false),
10907 sigmoid: m.router_sigmoid,
10908 bias: m.expert_bias.as_deref(),
10909 has_shared,
10910 shared_gated,
10911 route_scale: m.routed_scaling,
10912 }
10913 }
10914 };
10915 let attn = match &lw.attn {
10916 AttnKind::Full {
10917 wq,
10918 wk,
10919 wv,
10920 wo,
10921 q_norm,
10922 k_norm,
10923 output_gate,
10924 softplus_gate,
10925 bias,
10926 } => {
10927 if softplus_gate.is_some() || self.attention_heads_per_layer.is_some() {
10928 return None;
10929 }
10930 let (m, _, _, _) = wq
10931 .graph_weight()
10932 .or_else(|| wq.graph_weight_descriptor())?;
10933 model = Some(m.clone());
10934 crate::gpu::GraphAttn::Full {
10935 wq: gw(wq)?,
10936 wk: gw(wk)?,
10937 wv: gw(wv)?,
10938 wo: gw(wo)?,
10939 q_norm: q_norm.as_deref(),
10940 k_norm: k_norm.as_deref(),
10941 late_qk_norm: self.qk_norm_after_rope,
10942 bias: bias
10943 .as_ref()
10944 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
10945 output_gate: *output_gate,
10946 cpu_k: self.kv_cache.layers[li].k_heads(),
10947 cpu_v: self.kv_cache.layers[li].v_heads(),
10948 geom: self.graph_attn_geom(li),
10949 }
10950 }
10951 AttnKind::LinearGdn(w) => {
10952 let cfg = self.gdn_cfg?;
10953 let (m, _, _, _) = w
10954 .in_proj_qkv
10955 .graph_weight()
10956 .or_else(|| w.in_proj_qkv.graph_weight_descriptor())?;
10957 model = Some(m.clone());
10958 crate::gpu::GraphAttn::Gdn {
10959 qkv: gw(&w.in_proj_qkv)?,
10960 z: gw(&w.in_proj_z)?,
10961 a: gw(&w.in_proj_a)?,
10962 b: gw(&w.in_proj_b)?,
10963 out: gw(&w.out_proj)?,
10964 conv1d: &w.conv1d,
10965 a_log: &w.a_log,
10966 dt_bias: &w.dt_bias,
10967 norm: &w.norm,
10968 nv: cfg.num_v_heads,
10969 nk: cfg.num_k_heads,
10970 dk: cfg.key_head_dim,
10971 dv: cfg.value_head_dim,
10972 kk: cfg.conv_kernel,
10973 cpu_state: &self.kv_cache.layers[self.phys_layer(li)].linear_state,
10974 }
10975 }
10976 AttnKind::ShortConv(w) => {
10977 let cfg = self.short_conv_cfg?;
10978 let (m, _, _, _) = w
10979 .in_proj
10980 .graph_weight()
10981 .or_else(|| w.in_proj.graph_weight_descriptor())?;
10982 model = Some(m.clone());
10983 crate::gpu::GraphAttn::ShortConv {
10984 inp: gw(&w.in_proj)?,
10985 out: gw(&w.out_proj)?,
10986 taps: &w.conv,
10987 kernel: cfg.kernel,
10988 cpu_state: &self.kv_cache.layers[self.phys_layer(li)].linear_state,
10989 }
10990 }
10991 _ => return None,
10992 };
10993 layers.push(crate::gpu::GraphLayer {
10994 input_norm: &lw.input_norm,
10995 attn,
10996 post_norm: &lw.post_norm,
10997 ffn: gffn,
10998 });
10999 }
11000 let model = model?;
11001 let lm_gw = if upto_excl == self.num_layers
11007 && self.graph_want_logits
11008 && std::env::var("CMF_GPU_LMHEAD")
11009 .map(|v| v != "0")
11010 .unwrap_or(true)
11011 {
11012 self.weights
11013 .lm_head
11014 .graph_weight()
11015 .or_else(|| self.weights.lm_head.graph_weight_descriptor())
11016 .map(|(m, i, kind, rs)| {
11017 let name = &m.tensors[i].name;
11018 let prism = if crate::prism::is_inverse_embedding(m, name) {
11019 crate::gpu::GraphPrismOp::InverseEmbedding
11020 } else if crate::prism::is_forward_weight(m, name) {
11021 crate::gpu::GraphPrismOp::Forward
11022 } else {
11023 crate::gpu::GraphPrismOp::None
11024 };
11025 (
11026 crate::gpu::GraphW {
11027 idx: i,
11028 kind,
11029 row_scale: rs,
11030 data: &[],
11031 prism,
11032 affine: crate::prism::is_affine_target(m, name),
11033 },
11034 self.weights.lm_head.rows(),
11035 )
11036 })
11037 } else {
11038 None
11039 };
11040 let lm = lm_gw.as_ref().map(|(gw, rows)| (gw, *rows));
11041 let emb_gw = if steps > 1 {
11043 self.weights
11044 .embed_tokens
11045 .graph_weight()
11046 .or_else(|| self.weights.embed_tokens.graph_weight_descriptor())
11047 .map(|(m, i, kind, rs)| {
11048 let name = &m.tensors[i].name;
11049 let prism = if crate::prism::is_inverse_embedding(m, name) {
11050 crate::gpu::GraphPrismOp::InverseEmbedding
11051 } else if crate::prism::is_forward_weight(m, name) {
11052 crate::gpu::GraphPrismOp::Forward
11053 } else {
11054 crate::gpu::GraphPrismOp::None
11055 };
11056 (
11057 crate::gpu::GraphW {
11058 idx: i,
11059 kind,
11060 row_scale: rs,
11061 data: &[],
11062 prism,
11063 affine: crate::prism::is_affine_target(m, name),
11064 },
11065 self.weights.embed_tokens.rows(),
11066 self.embed_multiplier,
11067 )
11068 })
11069 } else {
11070 None
11071 };
11072
11073 let loop_norm_at: Vec<usize> = if self.loop_final_norm {
11079 (from..upto_excl.min(self.num_layers - 1))
11080 .filter(|&li| (li + 1) % self.physical_layers == 0)
11081 .map(|li| li - from)
11082 .collect()
11083 } else {
11084 Vec::new()
11085 };
11086 let mut h = hidden.to_vec();
11087 let dump_hidden = std::env::var_os("CMF_LOGIT_DUMP").is_some();
11093 let outcome = crate::gpu::forward_token_graph(
11094 &model,
11095 self.graph_kv_id,
11096 &layers,
11097 &o1_views,
11098 self.o1_epoch,
11099 &self.inv_freq,
11100 &mut h,
11101 nh,
11102 nkv,
11103 hd,
11104 self.attn_scale,
11105 rd,
11106 self.hidden_size,
11107 self.intermediate_size,
11108 position,
11109 self.kv_cache.max_seq_len,
11110 gemma,
11111 self.rms_eps as f32,
11112 lm,
11113 &self.weights.final_norm,
11114 logits_out,
11115 &loop_norm_at,
11116 steps,
11117 emb_gw.as_ref().map(|(gw, rows, m)| (gw, *rows, *m)),
11118 ids_out,
11119 layers_run,
11120 from,
11121 dump_hidden,
11122 );
11123 match outcome {
11124 crate::gpu::TokenGraphOutcome::Completed => Some(Ok(h)),
11125 crate::gpu::TokenGraphOutcome::Failed => Some(Err(())),
11126 crate::gpu::TokenGraphOutcome::Declined => None,
11127 }
11128 }
11129
11130 #[cfg(target_os = "macos")]
11139 #[allow(clippy::type_complexity)]
11140 fn metal_rows_plan(
11141 &self,
11142 ) -> Option<(
11143 Vec<MetalRowsItem<'_>>,
11144 std::sync::Arc<cortiq_core::CmfModel>,
11145 Option<crate::gpu_metal::GdnGpuCfg>,
11146 )> {
11147 use crate::gpu_metal::{AttnGpuLayer, GdnGpuCfg, GdnGpuLayer, MetalFfn};
11148 let graph_force = crate::gpu::q1_force() || crate::gpu::q2tp_gpu_opt_in();
11149 if !graph_force
11150 || !crate::gpu::enabled_here()
11151 || std::env::var("CMF_GPU_BLOCK")
11152 .map(|v| v == "0")
11153 .unwrap_or(false)
11154 || self.attn_softcap > 0.0
11155 || self.o1_active()
11156 || self.swa.is_some()
11157 || self.global_attn.is_some()
11158 || self.attention_heads_per_layer.is_some()
11159 || self.graph_attn_decline_reason().is_some()
11161 || self.attn_v_norm
11162 || self.loop_final_norm
11163 {
11164 return None;
11165 }
11166 let attend_contract = self.head_dim % 4 == 0
11167 && self.head_dim <= 256
11168 && self.rotary_dim >= 2
11169 && self.rotary_dim <= self.head_dim
11170 && (self.rotary_dim / 2) % 32 == 0
11171 && self.num_kv_heads > 0
11172 && self.num_heads % self.num_kv_heads == 0;
11173 if !attend_contract {
11174 return None;
11175 }
11176 let mut plan: Vec<MetalRowsItem> = Vec::new();
11177 let mut model_ref: Option<std::sync::Arc<cortiq_core::CmfModel>> = None;
11178 for li in 0..self.num_layers {
11179 let lw = &self.weights.layers[self.phys_layer(li)];
11180 if lw.attn_out_norm.is_some() || lw.ffn_out_norm.is_some() || lw.layer_scale.is_some() {
11181 return None;
11182 }
11183 let ffn = match &lw.ffn {
11184 FfnKind::Dense(d) if d.act == Act::Silu && d.segs.is_empty() => {
11185 let (Some(g), Some(u), Some(dn)) = (
11186 d.gate_proj.metal_graph_parts(),
11187 d.up_proj.metal_graph_parts(),
11188 d.down_proj.metal_graph_parts(),
11189 ) else {
11190 return None;
11191 };
11192 MetalFfn::Dense {
11193 gate: g,
11194 up: u,
11195 down: dn,
11196 }
11197 }
11198 _ => return None,
11199 };
11200 match &lw.attn {
11201 AttnKind::LinearGdn(w) if self.gdn_cfg.is_some() => {
11202 let (Some(qkv), Some(z), Some(a), Some(bb), Some(out)) = (
11203 w.in_proj_qkv.metal_graph_parts(),
11204 w.in_proj_z.metal_graph_parts(),
11205 w.in_proj_a.f32_parts(),
11206 w.in_proj_b.f32_parts(),
11207 w.out_proj.metal_graph_parts(),
11208 ) else {
11209 return None;
11210 };
11211 if let QTensor::Mapped { model, .. } = &w.in_proj_qkv {
11212 model_ref.get_or_insert_with(|| model.clone());
11213 }
11214 let gl = GdnGpuLayer {
11215 attn_norm: &lw.input_norm,
11216 post_norm: &lw.post_norm,
11217 qkv,
11218 z,
11219 a,
11220 b: bb,
11221 out,
11222 ffn,
11223 conv1d: &w.conv1d,
11224 a_log: &w.a_log,
11225 dt_bias: &w.dt_bias,
11226 gnorm: &w.norm,
11227 };
11228 match plan.last_mut() {
11229 Some(MetalRowsItem::Gdn { run, .. }) => run.push(gl),
11230 _ => plan.push(MetalRowsItem::Gdn {
11231 run: vec![gl],
11232 first: li,
11233 }),
11234 }
11235 }
11236 AttnKind::Full {
11237 wq,
11238 wk,
11239 wv,
11240 wo,
11241 q_norm,
11242 k_norm,
11243 output_gate,
11244 softplus_gate: None,
11245 bias: None,
11246 } => {
11247 let (Some(pq), Some(pk), Some(pv), Some(po)) =
11248 (
11249 wq.metal_graph_parts(),
11250 wk.metal_graph_parts(),
11251 wv.metal_graph_parts(),
11252 wo.metal_graph_parts(),
11253 )
11254 else {
11255 return None;
11256 };
11257 if let QTensor::Mapped { model, .. } = wq {
11258 model_ref.get_or_insert_with(|| model.clone());
11259 }
11260 let cache = &self.kv_cache.layers[li];
11261 if cache.mode != crate::kv_cache::KvMode::F32 || cache.o1.is_some() {
11262 return None;
11263 }
11264 plan.push(MetalRowsItem::Attn {
11265 l: AttnGpuLayer {
11266 attn_norm: &lw.input_norm,
11267 post_norm: &lw.post_norm,
11268 wq: pq,
11269 wk: pk,
11270 wv: pv,
11271 wo: po,
11272 ffn,
11273 },
11274 li,
11275 q_norm: q_norm.as_deref(),
11276 k_norm: k_norm.as_deref(),
11277 output_gate: *output_gate,
11278 });
11279 }
11280 _ => return None,
11281 }
11282 }
11283 let model = model_ref?;
11284 let gcfg = self.gdn_cfg.map(|cfg| GdnGpuCfg {
11285 nv: cfg.num_v_heads,
11286 nk: cfg.num_k_heads,
11287 dk: cfg.key_head_dim,
11288 dv: cfg.value_head_dim,
11289 kk: cfg.conv_kernel,
11290 hidden: self.hidden_size,
11291 inter: self.intermediate_size,
11292 c_dim: cfg.conv_dim(),
11293 eps: cfg.rms_eps as f32,
11294 gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
11295 });
11296 Some((plan, model, gcfg))
11297 }
11298
11299 #[cfg(target_os = "macos")]
11301 #[allow(clippy::too_many_arguments)]
11302 fn metal_attn_params<'a>(
11303 li: usize,
11304 cache: &'a crate::kv_cache::LayerKvCache,
11305 q_norm: Option<&'a [f32]>,
11306 k_norm: Option<&'a [f32]>,
11307 output_gate: bool,
11308 inv_freq: &'a [f32],
11309 geom: (usize, usize, usize, usize),
11310 pos0: usize,
11311 kv_id: u64,
11312 scale: f32,
11313 eps: f32,
11314 gemma: bool,
11315 late_qk_norm: bool,
11316 ) -> (crate::gpu_metal::AttnDeviceParams<'a>, usize) {
11317 let (nh, nkv, hd, rd) = geom;
11318 let cpu_k: Vec<&[f32]> = (0..nkv).map(|g| cache.head_keys(g)).collect();
11319 let cpu_v: Vec<&[f32]> = (0..nkv).map(|g| cache.head_values(g)).collect();
11320 let cpu_stored = cpu_k[0].len() / hd;
11321 (
11322 crate::gpu_metal::AttnDeviceParams {
11323 kv_id,
11324 layer: li,
11325 nh,
11326 nkv,
11327 hd,
11328 rd,
11329 position: pos0,
11330 scale,
11331 eps,
11332 gemma,
11333 late_qk_norm,
11334 output_gate,
11335 q_norm,
11336 k_norm,
11337 inv_freq,
11338 cpu_k,
11339 cpu_v,
11340 cpu_stored,
11341 o1: None,
11342 },
11343 cpu_stored,
11344 )
11345 }
11346
11347 #[cfg(target_os = "macos")]
11352 #[allow(clippy::type_complexity)]
11353 fn metal_rows_run(
11354 &mut self,
11355 hiddens: &mut [f32],
11356 pos0: usize,
11357 b: usize,
11358 prefill: bool,
11359 spec: Option<((usize, usize, usize), &[f32], &mut Vec<f32>)>,
11360 mut argmax_out: Option<(usize, &mut Vec<u32>)>,
11364 ) -> MetalRowsRun {
11365 use crate::gpu_metal::{GraphDims, VerifyGraph};
11366 if !crate::gpu_metal::wait_replay() {
11372 tracing::error!("Metal rows graph: the pending async replay failed");
11373 return MetalRowsRun::Failed;
11374 }
11375 spec_stamp("v.wait");
11376 let want = self.gdn_cfg.map(|c| c.state_len()).unwrap_or(0);
11382 if want > 0 {
11383 let phys = self.physical_layers.max(1);
11384 for (li, l) in self.kv_cache.layers.iter_mut().enumerate() {
11385 let is_gdn = self
11386 .weights
11387 .layers
11388 .get(li % phys)
11389 .is_some_and(|lw| matches!(lw.attn, AttnKind::LinearGdn(_)));
11390 if is_gdn && l.linear_state.len() != want {
11391 l.linear_state = vec![0f32; want];
11392 }
11393 }
11394 }
11395 let Some((plan, model, gcfg)) = self.metal_rows_plan() else {
11396 return MetalRowsRun::Declined;
11397 };
11398 spec_stamp("v.plan");
11399 let dims = GraphDims {
11400 hidden: self.hidden_size,
11401 eps: self.rms_eps as f32,
11402 gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
11403 };
11404 let Some(mut graph) = (if prefill {
11405 VerifyGraph::new_prefill(&model, dims, hiddens, b)
11406 } else {
11407 VerifyGraph::new(&model, dims, hiddens, b)
11408 }) else {
11409 return MetalRowsRun::Declined;
11410 };
11411 let geom = (
11412 self.num_heads,
11413 self.num_kv_heads,
11414 self.head_dim,
11415 self.rotary_dim,
11416 );
11417 let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
11418 let eps = self.rms_eps as f32;
11419 let kv_id = self.graph_kv_id;
11420 let inv_freq = self.inv_freq.clone();
11421 for item in &plan {
11422 let ok = match item {
11423 MetalRowsItem::Gdn { run, .. } => gcfg
11424 .as_ref()
11425 .map(|gc| run.iter().all(|l| graph.gdn_ok(l, gc)))
11426 .unwrap_or(false),
11427 MetalRowsItem::Attn {
11428 l,
11429 li,
11430 q_norm,
11431 k_norm,
11432 output_gate,
11433 } => {
11434 let (p, _) = Self::metal_attn_params(
11435 *li,
11436 &self.kv_cache.layers[*li],
11437 *q_norm,
11438 *k_norm,
11439 *output_gate,
11440 &inv_freq,
11441 geom,
11442 pos0,
11443 kv_id,
11444 self.attn_scale,
11445 eps,
11446 gemma,
11447 self.qk_norm_after_rope,
11448 );
11449 graph.attn_ok(l, &p)
11450 }
11451 };
11452 if !ok {
11453 use std::sync::atomic::{AtomicBool, Ordering};
11454 static SAID: AtomicBool = AtomicBool::new(false);
11455 if !SAID.swap(true, Ordering::Relaxed) {
11456 tracing::warn!("metal rows graph: a layer failed preflight — declining");
11457 }
11458 return MetalRowsRun::Declined;
11459 }
11460 }
11461 let lm = match &spec {
11462 Some((lm, _, _)) => {
11463 if !graph.lm_head_ok(*lm) {
11464 return MetalRowsRun::Declined;
11465 }
11466 Some(*lm)
11467 }
11468 None => None,
11469 };
11470 let mut gdn_layers = Vec::new();
11471 let mut attn_layers = Vec::new();
11472 for item in &plan {
11473 match item {
11474 MetalRowsItem::Gdn { run, first } => {
11475 let ro: Vec<&[f32]> = self.kv_cache.layers[*first..*first + run.len()]
11476 .iter()
11477 .map(|l| l.linear_state.as_slice())
11478 .collect();
11479 if !graph.encode_gdn_run_b(run, &ro, gcfg.as_ref().unwrap()) {
11480 return MetalRowsRun::Declined;
11481 }
11482 gdn_layers.extend(*first..*first + run.len());
11483 }
11484 MetalRowsItem::Attn {
11485 l,
11486 li,
11487 q_norm,
11488 k_norm,
11489 output_gate,
11490 } => {
11491 let (p, cpu_stored) = Self::metal_attn_params(
11492 *li,
11493 &self.kv_cache.layers[*li],
11494 *q_norm,
11495 *k_norm,
11496 *output_gate,
11497 &inv_freq,
11498 geom,
11499 pos0,
11500 kv_id,
11501 self.attn_scale,
11502 eps,
11503 gemma,
11504 self.qk_norm_after_rope,
11505 );
11506 if !graph.encode_attn_b(l, &p) {
11507 return MetalRowsRun::Declined;
11508 }
11509 attn_layers.push((*li, cpu_stored));
11510 }
11511 }
11512 }
11513 if let (Some(lm), Some((_, final_norm, _))) = (lm, spec.as_ref()) {
11514 if !graph.encode_lm_head_b(final_norm, lm) {
11515 return MetalRowsRun::Declined;
11516 }
11517 if let Some((n, _)) = argmax_out.as_ref() {
11522 if !graph.encode_argmax_b(*n) {
11523 argmax_out = None;
11524 }
11525 }
11526 }
11527 spec_stamp("v.enc");
11528 if !graph.sync() {
11529 return MetalRowsRun::Failed;
11530 }
11531 spec_stamp("v.gpu");
11532 match (spec, argmax_out) {
11533 (Some(_), Some((_, ids))) => {
11534 ids.resize(b, 0);
11535 if !graph.read_argmax(ids) {
11536 return MetalRowsRun::Failed;
11537 }
11538 spec_stamp("v.am");
11539 }
11540 (Some((lm, _, logits)), None) => {
11541 logits.resize(b * lm.1, 0.0);
11542 if !graph.read_logits(logits) {
11543 return MetalRowsRun::Failed;
11544 }
11545 spec_stamp("v.lg");
11546 }
11547 (None, _) => {}
11548 }
11549 if !graph.read_hidden(hiddens) {
11550 return MetalRowsRun::Failed;
11551 }
11552 spec_stamp("v.hid");
11553 MetalRowsRun::Completed(MetalVerifyPending {
11554 graph,
11555 gdn_layers,
11556 attn_layers,
11557 })
11558 }
11559
11560 #[cfg(target_os = "macos")]
11566 fn try_batch_graph_metal(
11567 &mut self,
11568 hiddens: &mut [f32],
11569 positions: &[usize],
11570 b: usize,
11571 spec: Option<((usize, usize, usize), &[f32], &mut Vec<f32>)>,
11572 argmax_out: Option<(usize, &mut Vec<u32>)>,
11573 ) -> crate::gpu::BatchGraphOutcome {
11574 let _t0 = std::time::Instant::now();
11575 if positions.len() != b
11576 || positions.windows(2).any(|w| w[1] != w[0] + 1)
11577 || hiddens.len() != b * self.hidden_size
11578 {
11579 return crate::gpu::BatchGraphOutcome::Declined;
11580 }
11581 let pending = match self.metal_rows_run(hiddens, positions[0], b, false, spec, argmax_out) {
11582 MetalRowsRun::Declined => return crate::gpu::BatchGraphOutcome::Declined,
11583 MetalRowsRun::Failed => return crate::gpu::BatchGraphOutcome::Failed,
11584 MetalRowsRun::Completed(pending) => pending,
11585 };
11586 if std::env::var("CMF_GRAPH_SPEC_TIME").is_ok() {
11587 eprintln!(
11588 "metal-verify: {:.1} ms | b={b}",
11589 _t0.elapsed().as_secs_f64() * 1e3
11590 );
11591 }
11592 self.metal_verify = Some(pending);
11593 crate::gpu::BatchGraphOutcome::Completed
11594 }
11595
11596 #[cfg(target_os = "macos")]
11601 fn prefill_rows_metal(
11602 &mut self,
11603 ids: &[u32],
11604 start_pos: usize,
11605 spec: Option<((usize, usize, usize), &[f32], &mut Vec<f32>)>,
11606 ) -> MetalPrefillOutcome {
11607 let b = ids.len();
11608 if b == 0 || b > 512 {
11609 return MetalPrefillOutcome::Declined;
11610 }
11611 METAL_PREFILL_CHUNKS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
11612 let with_head = spec.is_some();
11613 let hs = self.hidden_size;
11614 let mut hiddens = vec![0f32; b * hs];
11615 for (j, &id) in ids.iter().enumerate() {
11616 let e = self.embed_single(id);
11617 hiddens[j * hs..(j + 1) * hs].copy_from_slice(&e);
11618 }
11619 let mut pending = match self.metal_rows_run(&mut hiddens, start_pos, b, true, spec, None) {
11620 MetalRowsRun::Declined => return MetalPrefillOutcome::Declined,
11621 MetalRowsRun::Failed => {
11622 METAL_PREFILL_ERRORS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
11623 return MetalPrefillOutcome::Failed;
11624 }
11625 MetalRowsRun::Completed(pending) => pending,
11626 };
11627 let idxs = pending.gdn_layers.clone();
11629 let mut outs: Vec<&mut [f32]> = self
11630 .kv_cache
11631 .layers
11632 .iter_mut()
11633 .enumerate()
11634 .filter(|(i, _)| idxs.binary_search(i).is_ok())
11635 .map(|(_, l)| l.linear_state.as_mut_slice())
11636 .collect();
11637 if !pending.graph.finish_states(&mut outs) {
11638 METAL_PREFILL_ERRORS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
11639 return MetalPrefillOutcome::Failed;
11640 }
11641 let (nkv, hd) = (self.num_kv_heads, self.head_dim);
11642 let mut rows = Vec::with_capacity(pending.attn_layers.len());
11646 for (li, cpu_stored) in &pending.attn_layers {
11647 let mut kbuf = vec![0f32; b * nkv * hd];
11648 let mut vbuf = vec![0f32; b * nkv * hd];
11649 if !crate::gpu_metal::kv_mirror_read_rows(
11650 self.graph_kv_id,
11651 *li,
11652 nkv,
11653 hd,
11654 *cpu_stored,
11655 b,
11656 &mut kbuf,
11657 &mut vbuf,
11658 ) {
11659 METAL_PREFILL_ERRORS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
11660 return MetalPrefillOutcome::Failed;
11661 }
11662 rows.push((*li, *cpu_stored, kbuf, vbuf));
11663 }
11664 for (li, cpu_stored, kbuf, vbuf) in rows {
11665 let cache = &mut self.kv_cache.layers[li];
11666 for r in 0..b {
11667 cache.append(
11668 &kbuf[r * nkv * hd..(r + 1) * nkv * hd],
11669 &vbuf[r * nkv * hd..(r + 1) * nkv * hd],
11670 &[],
11671 );
11672 }
11673 crate::gpu_metal::kv_mirror_set_stored(self.graph_kv_id, li, cpu_stored + b);
11674 }
11675 METAL_PREFILL_ROWS.fetch_add(b as u64, std::sync::atomic::Ordering::Relaxed);
11676 if with_head {
11677 METAL_PREFILL_HEAD_ROWS.fetch_add(b as u64, std::sync::atomic::Ordering::Relaxed);
11678 }
11679 MetalPrefillOutcome::Completed(hiddens)
11680 }
11681
11682 #[cfg(target_os = "macos")]
11683 fn prefill_batch_metal(&mut self, ids: &[u32], start_pos: usize) -> MetalPrefillOutcome {
11684 self.prefill_rows_metal(ids, start_pos, None)
11685 }
11686
11687 #[cfg(target_os = "macos")]
11692 fn nll_batch_metal(&mut self, ids: &[u32], start: usize) -> MetalBatchNllOutcome {
11693 if ids.len() < 2 || self.o1_active() || self.head_clusters.is_some() {
11694 return MetalBatchNllOutcome::Declined;
11695 }
11696 let Some(lm) = self.weights.lm_head.metal_graph_parts() else {
11697 return MetalBatchNllOutcome::Declined;
11698 };
11699 let chunk = std::env::var("CMF_METAL_PREFILL_CHUNK")
11700 .ok()
11701 .and_then(|v| v.parse::<usize>().ok())
11702 .filter(|&v| (1..=512).contains(&v))
11703 .unwrap_or(32);
11704 let final_norm = self.weights.final_norm.clone();
11705 let mut nll = 0.0f64;
11706 let mut count = 0usize;
11707 let mut pos = 0usize;
11708 let mut completed = 0usize;
11709 while pos < ids.len() {
11710 let end = (pos + chunk).min(ids.len());
11711 let mut logits = Vec::new();
11712 let outcome = self.prefill_rows_metal(
11713 &ids[pos..end],
11714 pos,
11715 Some((lm, &final_norm, &mut logits)),
11716 );
11717 match outcome {
11718 MetalPrefillOutcome::Declined => {
11719 return if completed == 0 {
11720 MetalBatchNllOutcome::Declined
11721 } else {
11722 MetalBatchNllOutcome::Failed(format!(
11723 "ordinary Metal NLL batch declined after {completed} chunks"
11724 ))
11725 };
11726 }
11727 MetalPrefillOutcome::Failed => {
11728 return MetalBatchNllOutcome::Failed(
11729 "ordinary Metal NLL batch failed after admission".to_string(),
11730 );
11731 }
11732 MetalPrefillOutcome::Completed(_) => {}
11733 }
11734 completed += 1;
11735 let vocab = self.vocab_size.min(lm.1);
11736 if logits.len() != (end - pos) * lm.1 || vocab == 0 {
11737 return MetalBatchNllOutcome::Failed(
11738 "ordinary Metal NLL head returned an invalid shape".to_string(),
11739 );
11740 }
11741 for row in 0..(end - pos) {
11742 let absolute = pos + row;
11743 if absolute < start || absolute + 1 >= ids.len() {
11744 continue;
11745 }
11746 let lg = &mut logits[row * lm.1..row * lm.1 + vocab];
11747 if let Some(mu) = self.logit_multiplier {
11748 for v in lg.iter_mut() {
11749 *v *= mu;
11750 }
11751 }
11752 if let Some(c) = self.final_softcap {
11753 for v in lg.iter_mut() {
11754 *v = c * (*v / c).tanh();
11755 }
11756 }
11757 let target = ids[absolute + 1] as usize;
11758 if target >= vocab {
11759 return MetalBatchNllOutcome::Failed(format!(
11760 "target token {target} exceeds Metal head rows {vocab}"
11761 ));
11762 }
11763 let max = lg.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
11764 let lse: f64 = lg
11765 .iter()
11766 .map(|&v| ((v - max) as f64).exp())
11767 .sum::<f64>()
11768 .ln()
11769 + max as f64;
11770 nll += lse - lg[target] as f64;
11771 count += 1;
11772 }
11773 pos = end;
11774 }
11775 MetalBatchNllOutcome::Completed(nll, count)
11776 }
11777
11778 #[cfg(target_os = "macos")]
11782 fn metal_verify_commit(&mut self, a: usize) -> bool {
11783 let Some(mut pending) = self.metal_verify.take() else {
11784 return false;
11785 };
11786 let n = a + 1;
11787 let idxs = pending.gdn_layers.clone();
11789 let mut outs: Vec<&mut [f32]> = self
11790 .kv_cache
11791 .layers
11792 .iter_mut()
11793 .enumerate()
11794 .filter(|(i, _)| idxs.binary_search(i).is_ok())
11795 .map(|(_, l)| l.linear_state.as_mut_slice())
11796 .collect();
11797 if !pending.graph.commit(n, &mut outs) {
11798 return false;
11799 }
11800 spec_stamp("c.replay");
11801 let (nkv, hd) = (self.num_kv_heads, self.head_dim);
11802 let mut rows = Vec::with_capacity(pending.attn_layers.len());
11806 for (li, cpu_stored) in &pending.attn_layers {
11807 let mut kbuf = vec![0f32; n * nkv * hd];
11808 let mut vbuf = vec![0f32; n * nkv * hd];
11809 if !crate::gpu_metal::kv_mirror_read_rows(
11810 self.graph_kv_id,
11811 *li,
11812 nkv,
11813 hd,
11814 *cpu_stored,
11815 n,
11816 &mut kbuf,
11817 &mut vbuf,
11818 ) {
11819 return false;
11820 }
11821 rows.push((*li, *cpu_stored, kbuf, vbuf));
11822 }
11823 for (li, cpu_stored, kbuf, vbuf) in rows {
11824 let cache = &mut self.kv_cache.layers[li];
11825 for r in 0..n {
11826 cache.append(
11827 &kbuf[r * nkv * hd..(r + 1) * nkv * hd],
11828 &vbuf[r * nkv * hd..(r + 1) * nkv * hd],
11829 &[],
11830 );
11831 }
11832 crate::gpu_metal::kv_mirror_set_stored(self.graph_kv_id, li, cpu_stored + n);
11833 }
11834 spec_stamp("c.kv");
11835 true
11836 }
11837
11838 #[cfg(target_os = "macos")]
11845 fn mtp_warm_batch_submit(
11846 &mut self,
11847 m: &mut MtpModule,
11848 pairs: &[(&[f32], u32)],
11849 first_pos: usize,
11850 ) -> Option<MetalWarmPending> {
11851 use crate::gpu_metal::{AttnDeviceParams, AttnGpuLayer, GraphDims, MetalFfn, VerifyGraph};
11852 let b = pairs.len();
11853 if b == 0 || b > 512 || m.kv.mode != crate::kv_cache::KvMode::F32 || m.kv.o1.is_some() {
11854 return None;
11855 }
11856 let AttnKind::Full {
11857 wq,
11858 wk,
11859 wv,
11860 wo,
11861 q_norm,
11862 k_norm,
11863 output_gate,
11864 softplus_gate: None,
11865 bias: None,
11866 } = &m.layer.attn
11867 else {
11868 return None;
11869 };
11870 let FfnKind::Dense(d) = &m.layer.ffn else {
11871 return None;
11872 };
11873 if !d.segs.is_empty() {
11874 return None;
11875 }
11876 let (Some(pq), Some(pk), Some(pv), Some(po)) =
11877 (wq.q1_parts(), wk.q1_parts(), wv.q1_parts(), wo.q1_parts())
11878 else {
11879 return None;
11880 };
11881 let (Some(g), Some(u), Some(dn)) = (
11882 d.gate_proj.q1_parts(),
11883 d.up_proj.q1_parts(),
11884 d.down_proj.q1_parts(),
11885 ) else {
11886 return None;
11887 };
11888 let Some(eh) = m.eh_proj.q1_parts() else {
11889 return None;
11890 };
11891 let QTensor::Mapped { model, .. } = wq else {
11892 return None;
11893 };
11894 let model = model.clone();
11895 let hs = self.hidden_size;
11896 let mut cat = vec![0f32; b * 2 * hs];
11898 for (j, (h, tok)) in pairs.iter().enumerate() {
11899 let e = self.embed_single(*tok);
11900 let (ce, ch) = cat[j * 2 * hs..(j + 1) * 2 * hs].split_at_mut(hs);
11901 inference::rms_norm_into(&e, &m.enorm, self.rms_eps, self.norm_style, ce);
11902 inference::rms_norm_into(h, &m.hnorm, self.rms_eps, self.norm_style, ch);
11903 }
11904 let dims = GraphDims {
11905 hidden: hs,
11906 eps: self.rms_eps as f32,
11907 gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
11908 };
11909 spec_stamp("w.cat");
11910 let Some(mut graph) = VerifyGraph::new_via_proj(&model, dims, eh, &cat, b) else {
11911 return None;
11912 };
11913 spec_stamp("w.new");
11914 let l = AttnGpuLayer {
11915 attn_norm: &m.layer.input_norm,
11916 post_norm: &m.layer.post_norm,
11917 wq: pq,
11918 wk: pk,
11919 wv: pv,
11920 wo: po,
11921 ffn: MetalFfn::Dense {
11922 gate: g,
11923 up: u,
11924 down: dn,
11925 },
11926 };
11927 let (nh, nkv, hd, rd) = (
11928 self.num_heads,
11929 self.num_kv_heads,
11930 self.head_dim,
11931 self.rotary_dim,
11932 );
11933 let inv_freq = self.inv_freq.clone();
11934 let cpu_stored;
11935 {
11936 let cache = &m.kv;
11937 let cpu_k: Vec<&[f32]> = (0..nkv).map(|g| cache.head_keys(g)).collect();
11938 let cpu_v: Vec<&[f32]> = (0..nkv).map(|g| cache.head_values(g)).collect();
11939 cpu_stored = cpu_k[0].len() / hd;
11940 if cpu_stored > first_pos {
11945 spec_stamp("w.decl");
11946 return None;
11947 }
11948 let p = AttnDeviceParams {
11949 kv_id: self.mtp_kv_id(),
11950 layer: Self::MTP_LAYER_BASE,
11951 nh,
11952 nkv,
11953 hd,
11954 rd,
11955 position: first_pos,
11956 scale: self.attn_scale,
11957 eps: self.rms_eps as f32,
11958 gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
11959 late_qk_norm: self.qk_norm_after_rope,
11960 output_gate: *output_gate,
11961 q_norm: q_norm.as_deref(),
11962 k_norm: k_norm.as_deref(),
11963 inv_freq: &inv_freq,
11964 cpu_k,
11965 cpu_v,
11966 cpu_stored,
11967 o1: None,
11968 };
11969 if !graph.attn_ok(&l, &p) || !graph.encode_attn_b(&l, &p) {
11970 return None;
11971 }
11972 }
11973 spec_stamp("w.enc");
11974 if !graph.submit() {
11975 return None;
11976 }
11977 spec_stamp("w.sub");
11978 Some(MetalWarmPending {
11979 graph,
11980 cpu_stored,
11981 b,
11982 })
11983 }
11984
11985 #[cfg(target_os = "macos")]
11988 fn mtp_warm_batch_metal(
11989 &mut self,
11990 m: &mut MtpModule,
11991 pairs: &[(&[f32], u32)],
11992 first_pos: usize,
11993 ) -> bool {
11994 match self.mtp_warm_batch_submit(m, pairs, first_pos) {
11995 Some(p) => self.mtp_warm_batch_finish(m, p),
11996 None => false,
11997 }
11998 }
11999
12000 #[cfg(target_os = "macos")]
12005 fn mtp_warm_batch_finish(&mut self, m: &mut MtpModule, pending: MetalWarmPending) -> bool {
12006 let MetalWarmPending {
12007 mut graph,
12008 cpu_stored,
12009 b,
12010 } = pending;
12011 let (nkv, hd) = (self.num_kv_heads, self.head_dim);
12012 if !graph.sync() {
12013 return false;
12014 }
12015 spec_stamp("w.gpu");
12016 let mut kbuf = vec![0f32; b * nkv * hd];
12017 let mut vbuf = vec![0f32; b * nkv * hd];
12018 if !crate::gpu_metal::kv_mirror_read_rows(
12019 self.mtp_kv_id(),
12020 Self::MTP_LAYER_BASE,
12021 nkv,
12022 hd,
12023 cpu_stored,
12024 b,
12025 &mut kbuf,
12026 &mut vbuf,
12027 ) {
12028 return false;
12029 }
12030 for r in 0..b {
12031 m.kv.append(
12032 &kbuf[r * nkv * hd..(r + 1) * nkv * hd],
12033 &vbuf[r * nkv * hd..(r + 1) * nkv * hd],
12034 &[],
12035 );
12036 }
12037 crate::gpu_metal::kv_mirror_set_stored(
12038 self.mtp_kv_id(),
12039 Self::MTP_LAYER_BASE,
12040 cpu_stored + b,
12041 );
12042 spec_stamp("w.kv");
12043 true
12044 }
12045
12046 pub(crate) fn note_draft_id(&mut self, id: u32) {
12053 let cut = Self::draft_vocab_rows(usize::MAX).max(131_072);
12054 if (id as usize) >= cut {
12055 self.draft_full_streak = 16;
12056 } else {
12057 self.draft_full_streak = self.draft_full_streak.saturating_sub(1);
12058 }
12059 }
12060
12061 fn draft_head_rows(&self, head_rows: usize) -> usize {
12064 if self.draft_full_streak > 0 {
12065 head_rows
12066 } else {
12067 Self::draft_vocab_rows(head_rows)
12068 }
12069 }
12070
12071 fn draft_vocab_rows(head_rows: usize) -> usize {
12074 static N: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
12075 let n = *N.get_or_init(|| {
12076 std::env::var("CMF_DRAFT_VOCAB")
12077 .ok()
12078 .and_then(|v| v.parse().ok())
12079 .unwrap_or(65536)
12080 });
12081 if n == 0 { head_rows } else { n.min(head_rows) }
12082 }
12083
12084 #[cfg(target_os = "macos")]
12089 fn mtp_step_metal(
12090 &mut self,
12091 m: &mut MtpModule,
12092 hidden: &[f32],
12093 next_token: u32,
12094 position: usize,
12095 want_logits: bool,
12096 ) -> Option<(Vec<f32>, Vec<f32>)> {
12097 use crate::gpu_metal::{AttnDeviceParams, AttnGpuLayer, GraphDims, MetalFfn, TokenGraph};
12098 if std::env::var("CMF_MTP_GRAPH").as_deref() == Ok("0")
12099 || !crate::gpu::q1_force()
12100 || !crate::gpu::enabled_here()
12101 || self.attn_softcap > 0.0
12102 || self.attention_heads_per_layer.is_some()
12103 || m.kv.mode != crate::kv_cache::KvMode::F32
12104 || m.kv.o1.is_some()
12105 {
12106 return None;
12107 }
12108 let AttnKind::Full {
12109 wq,
12110 wk,
12111 wv,
12112 wo,
12113 q_norm,
12114 k_norm,
12115 output_gate,
12116 softplus_gate: None,
12117 bias: None,
12118 } = &m.layer.attn
12119 else {
12120 return None;
12121 };
12122 let FfnKind::Dense(d) = &m.layer.ffn else {
12123 return None;
12124 };
12125 if d.act != Act::Silu || !d.segs.is_empty() {
12126 return None;
12127 }
12128 let (pq, pk, pv, po) = (
12129 wq.q1_parts()?,
12130 wk.q1_parts()?,
12131 wv.q1_parts()?,
12132 wo.q1_parts()?,
12133 );
12134 let (g, u, dn) = (
12135 d.gate_proj.q1_parts()?,
12136 d.up_proj.q1_parts()?,
12137 d.down_proj.q1_parts()?,
12138 );
12139 let QTensor::Mapped { model, .. } = wq else {
12140 return None;
12141 };
12142 let model = model.clone();
12143 let lm = if want_logits {
12144 Some(self.weights.lm_head.q1_parts()?)
12145 } else {
12146 None
12147 };
12148 let dims = GraphDims {
12149 hidden: self.hidden_size,
12150 eps: self.rms_eps as f32,
12151 gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
12152 };
12153 let hs = self.hidden_size;
12156 let mut x = vec![0f32; hs];
12157 let mut graph = TokenGraph::new(&model, dims, &x)?;
12158 let mut folded = false;
12159 if let Some(eh) = m.eh_proj.q1_parts() {
12160 let e = self.embed_single(next_token);
12161 let mut cat = vec![0.0f32; 2 * hs];
12162 let (cat_e, cat_h) = cat.split_at_mut(hs);
12163 inference::rms_norm_into(&e, &m.enorm, self.rms_eps, self.norm_style, cat_e);
12164 inference::rms_norm_into(hidden, &m.hnorm, self.rms_eps, self.norm_style, cat_h);
12165 folded = graph.encode_input_proj(eh, &cat);
12166 }
12167 if !folded {
12168 x = self.mtp_block_input(m, hidden, next_token);
12169 graph = TokenGraph::new(&model, dims, &x)?;
12170 }
12171 spec_stamp("d.in");
12172 let l = AttnGpuLayer {
12173 attn_norm: &m.layer.input_norm,
12174 post_norm: &m.layer.post_norm,
12175 wq: pq,
12176 wk: pk,
12177 wv: pv,
12178 wo: po,
12179 ffn: MetalFfn::Dense {
12180 gate: g,
12181 up: u,
12182 down: dn,
12183 },
12184 };
12185 let (nh, nkv, hd, rd) = (
12186 self.num_heads,
12187 self.num_kv_heads,
12188 self.head_dim,
12189 self.rotary_dim,
12190 );
12191 let inv_freq = self.inv_freq.clone();
12192 {
12193 let cache = &m.kv;
12194 let cpu_k: Vec<&[f32]> = (0..nkv).map(|g| cache.head_keys(g)).collect();
12195 let cpu_v: Vec<&[f32]> = (0..nkv).map(|g| cache.head_values(g)).collect();
12196 let cpu_stored = cpu_k[0].len() / hd;
12197 let p = AttnDeviceParams {
12198 kv_id: self.mtp_kv_id(),
12199 layer: Self::MTP_LAYER_BASE,
12200 nh,
12201 nkv,
12202 hd,
12203 rd,
12204 position,
12205 scale: self.attn_scale,
12206 eps: self.rms_eps as f32,
12207 gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
12208 late_qk_norm: self.qk_norm_after_rope,
12209 output_gate: *output_gate,
12210 q_norm: q_norm.as_deref(),
12211 k_norm: k_norm.as_deref(),
12212 inv_freq: &inv_freq,
12213 cpu_k,
12214 cpu_v,
12215 cpu_stored,
12216 o1: None,
12217 };
12218 if !graph.attn_device_ok(&l, &p) || !graph.encode_attn_device(&l, &p) {
12219 return None;
12220 }
12221 }
12222 let draft_rows = if let Some(lm) = lm {
12228 self.draft_head_rows(lm.1)
12229 } else {
12230 0
12231 };
12232 if let Some(lm) = lm {
12233 if !graph.lm_head_ok(lm) {
12234 return None;
12235 }
12236 if draft_rows < lm.1 {
12237 if !graph.encode_lm_head_part(&m.final_norm, lm, draft_rows) {
12238 return None;
12239 }
12240 } else {
12241 graph.encode_lm_head(&m.final_norm, lm);
12242 }
12243 }
12244 spec_stamp("d.enc");
12245 if graph.sync_checked().is_err() {
12246 return None;
12247 }
12248 spec_stamp("d.gpu");
12249 let mut logits = Vec::new();
12250 if let Some(lm) = lm {
12251 let n_read = draft_rows.min(lm.1).min(self.vocab_size);
12252 logits = attention::take_buf(n_read);
12253 graph.read_logits(&mut logits);
12254 logits.resize(self.vocab_size, f32::NEG_INFINITY);
12256 }
12257 graph.finish(&mut x);
12258 let mut krow = attention::take_buf(nkv * hd);
12259 let mut vrow = attention::take_buf(nkv * hd);
12260 if crate::gpu_metal::kv_mirror_read_last(
12261 self.mtp_kv_id(),
12262 Self::MTP_LAYER_BASE,
12263 nkv,
12264 hd,
12265 &mut krow,
12266 &mut vrow,
12267 ) {
12268 m.kv.append(&krow, &vrow, &[]);
12269 }
12270 attention::recycle_buf(&mut krow);
12271 attention::recycle_buf(&mut vrow);
12272 spec_stamp("d.rd");
12273 Some((logits, x))
12274 }
12275
12276 fn mtp_chain_on() -> bool {
12289 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
12290 *ON.get_or_init(|| std::env::var("CMF_MTP_CHAIN").as_deref() != Ok("0"))
12291 }
12292
12293 #[cfg(target_os = "macos")]
12305 fn mtp_draft_chain_metal(
12306 &mut self,
12307 m: &mut MtpModule,
12308 hidden: &[f32],
12309 t_next: u32,
12310 position: usize,
12311 k: usize,
12312 ) -> Result<Vec<u32>, bool> {
12313 use crate::gpu_metal::{AttnDeviceParams, AttnGpuLayer, GraphDims, MetalFfn, TokenGraph};
12314 if k == 0
12315 || k > 64
12316 || !Self::mtp_chain_on()
12317 || std::env::var("CMF_MTP_GRAPH").as_deref() == Ok("0")
12318 || !crate::gpu::q1_force()
12319 || !crate::gpu::enabled_here()
12320 || self.attn_softcap > 0.0
12321 || self.attention_heads_per_layer.is_some()
12322 || m.kv.mode != crate::kv_cache::KvMode::F32
12323 || m.kv.o1.is_some()
12324 || self.dsv4.is_some()
12326 || self.dsv41.is_some()
12327 || self.qwen4_exp.is_some()
12328 || self.g3n.is_some()
12329 {
12330 return Err(false);
12331 }
12332 let AttnKind::Full {
12333 wq,
12334 wk,
12335 wv,
12336 wo,
12337 q_norm,
12338 k_norm,
12339 output_gate,
12340 softplus_gate: None,
12341 bias: None,
12342 } = &m.layer.attn
12343 else {
12344 return Err(false);
12345 };
12346 let FfnKind::Dense(d) = &m.layer.ffn else {
12347 return Err(false);
12348 };
12349 if d.act != Act::Silu || !d.segs.is_empty() {
12350 return Err(false);
12351 }
12352 let (Some(pq), Some(pk), Some(pv), Some(po)) =
12353 (wq.q1_parts(), wk.q1_parts(), wv.q1_parts(), wo.q1_parts())
12354 else {
12355 return Err(false);
12356 };
12357 let (Some(g), Some(u), Some(dn)) = (
12358 d.gate_proj.q1_parts(),
12359 d.up_proj.q1_parts(),
12360 d.down_proj.q1_parts(),
12361 ) else {
12362 return Err(false);
12363 };
12364 let (Some(eh), Some(lm)) = (m.eh_proj.q1_parts(), self.weights.lm_head.q1_parts()) else {
12365 return Err(false);
12366 };
12367 let QTensor::Mapped { model, .. } = wq else {
12368 return Err(false);
12369 };
12370 let model = model.clone();
12371 let QTensor::Mapped {
12374 model: em,
12375 idx: eidx,
12376 dtype: cortiq_core::TensorDtype::Q4TiledP,
12377 ..
12378 } = &self.weights.embed_tokens
12379 else {
12380 return Err(false);
12381 };
12382 if !std::sync::Arc::ptr_eq(em, &model)
12383 || crate::prism::is_inverse_embedding(&model, &model.tensors[*eidx].name)
12384 {
12385 return Err(false);
12386 }
12387 let embed = (
12388 *eidx,
12389 self.weights.embed_tokens.rows(),
12390 self.weights.embed_tokens.cols(),
12391 );
12392 if embed.2 != self.hidden_size || hidden.len() != self.hidden_size {
12393 return Err(false);
12394 }
12395 let dims = GraphDims {
12396 hidden: self.hidden_size,
12397 eps: self.rms_eps as f32,
12398 gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
12399 };
12400 let Some(mut graph) = TokenGraph::new(&model, dims, hidden) else {
12401 return Err(false);
12402 };
12403 if !graph.chain_embed_ok(embed) || !graph.lm_head_ok(lm) {
12404 return Err(false);
12405 }
12406 let l = AttnGpuLayer {
12407 attn_norm: &m.layer.input_norm,
12408 post_norm: &m.layer.post_norm,
12409 wq: pq,
12410 wk: pk,
12411 wv: pv,
12412 wo: po,
12413 ffn: MetalFfn::Dense {
12414 gate: g,
12415 up: u,
12416 down: dn,
12417 },
12418 };
12419 let (nh, nkv, hd, rd) = (
12420 self.num_heads,
12421 self.num_kv_heads,
12422 self.head_dim,
12423 self.rotary_dim,
12424 );
12425 let inv_freq = self.inv_freq.clone();
12426 let draft_rows = self.draft_head_rows(lm.1);
12427 let n_arg = draft_rows.min(lm.1).min(self.vocab_size);
12428 if n_arg == 0 {
12429 return Err(false);
12430 }
12431 let split = std::env::var("CMF_MTP_CHAIN_SPLIT").as_deref() == Ok("1");
12438 let t_chain = std::time::Instant::now();
12439 graph.chain_ids_init(t_next, k);
12440 let cpu_stored;
12441 {
12442 let cache = &m.kv;
12443 let cpu_k: Vec<&[f32]> = (0..nkv).map(|g| cache.head_keys(g)).collect();
12444 let cpu_v: Vec<&[f32]> = (0..nkv).map(|g| cache.head_values(g)).collect();
12445 cpu_stored = cpu_k[0].len() / hd;
12446 for j in 0..k {
12447 if !graph.encode_chain_input(
12448 embed,
12449 j as u32,
12450 &m.enorm,
12451 &m.hnorm,
12452 self.embed_multiplier,
12453 eh,
12454 ) {
12455 return Err(false);
12456 }
12457 let p = AttnDeviceParams {
12461 kv_id: self.mtp_kv_id(),
12462 layer: Self::MTP_LAYER_BASE,
12463 nh,
12464 nkv,
12465 hd,
12466 rd,
12467 position: position + j,
12468 scale: self.attn_scale,
12469 eps: self.rms_eps as f32,
12470 gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
12471 late_qk_norm: self.qk_norm_after_rope,
12472 output_gate: *output_gate,
12473 q_norm: q_norm.as_deref(),
12474 k_norm: k_norm.as_deref(),
12475 inv_freq: &inv_freq,
12476 cpu_k: cpu_k.clone(),
12477 cpu_v: cpu_v.clone(),
12478 cpu_stored: cpu_stored + j,
12479 o1: None,
12480 };
12481 if !graph.attn_device_ok(&l, &p) || !graph.encode_attn_device(&l, &p) {
12482 return Err(false);
12483 }
12484 if draft_rows < lm.1 {
12485 if !graph.encode_lm_head_part(&m.final_norm, lm, draft_rows) {
12486 return Err(false);
12487 }
12488 } else {
12489 graph.encode_lm_head(&m.final_norm, lm);
12490 }
12491 if !graph.encode_argmax(n_arg, j as u32 + 1) {
12492 return Err(false);
12493 }
12494 if split {
12495 graph.commit();
12498 }
12499 }
12500 }
12501 let t_enc = t_chain.elapsed();
12502 if graph.sync_checked().is_err() {
12503 return Err(true);
12504 }
12505 if std::env::var_os("CMF_GRAPH_SPEC_TIME").is_some() {
12506 eprintln!(
12507 "mtp-chain: encode {:.1} ms | wait {:.1} ms (k={k}, head rows {draft_rows}{})",
12508 t_enc.as_secs_f64() * 1e3,
12509 (t_chain.elapsed() - t_enc).as_secs_f64() * 1e3,
12510 if split { ", split" } else { "" }
12511 );
12512 }
12513 let mut ids = vec![0u32; k];
12514 if !graph.chain_ids_read(&mut ids) {
12515 return Err(true);
12516 }
12517 let mut kbuf = vec![0f32; k * nkv * hd];
12518 let mut vbuf = vec![0f32; k * nkv * hd];
12519 if !crate::gpu_metal::kv_mirror_read_rows(
12520 self.mtp_kv_id(),
12521 Self::MTP_LAYER_BASE,
12522 nkv,
12523 hd,
12524 cpu_stored,
12525 k,
12526 &mut kbuf,
12527 &mut vbuf,
12528 ) {
12529 return Err(true);
12530 }
12531 for r in 0..k {
12532 m.kv.append(
12533 &kbuf[r * nkv * hd..(r + 1) * nkv * hd],
12534 &vbuf[r * nkv * hd..(r + 1) * nkv * hd],
12535 &[],
12536 );
12537 }
12538 Ok(ids)
12539 }
12540
12541 fn try_batch_graph_wgpu(
12542 &self,
12543 hiddens: &mut [f32],
12544 positions: &[usize],
12545 k: usize,
12546 spec: Option<crate::gpu::SpecTail<'_>>,
12547 ) -> crate::gpu::BatchGraphOutcome {
12548 self.try_batch_graph_wgpu_prefix(hiddens, positions, k, spec, None)
12549 }
12550
12551 fn try_batch_graph_wgpu_prefix(
12556 &self,
12557 hiddens: &mut [f32],
12558 positions: &[usize],
12559 k: usize,
12560 spec: Option<crate::gpu::SpecTail<'_>>,
12561 layers_run: Option<&mut usize>,
12562 ) -> crate::gpu::BatchGraphOutcome {
12563 let graph_end = match self.mimo_moe.graph_prefix_end() {
12564 Some(end) if end < self.num_layers => {
12565 if layers_run.is_none() || spec.is_some() || end == 0 {
12566 return crate::gpu::BatchGraphOutcome::Declined;
12567 }
12568 end
12569 }
12570 _ => self.num_layers,
12571 };
12572 let _tb = std::time::Instant::now();
12573 let batch_debug = std::env::var_os("CMF_BATCH_DEBUG").is_some();
12574 if self.attn_softcap > 0.0 {
12575 return crate::gpu::BatchGraphOutcome::Declined; }
12577 if let Some(reason) = self.wgpu_graph_attn_decline() {
12580 self.note_graph_decline("wgpu batch graph", reason);
12581 return crate::gpu::BatchGraphOutcome::Declined;
12582 }
12583 let nh = self.num_heads;
12584 let (nkv, hd, rd) = self.layer_geom(0);
12585 let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
12586 fn gw(t: &QTensor) -> Option<crate::gpu::GraphW<'_>> {
12587 if let Some((m, i, kind, rs)) = t
12588 .graph_weight()
12589 .or_else(|| t.graph_weight_descriptor())
12590 {
12591 let name = &m.tensors[i].name;
12592 let prism = if crate::prism::is_inverse_embedding(m, name) {
12593 crate::gpu::GraphPrismOp::InverseEmbedding
12594 } else if crate::prism::is_forward_weight(m, name) {
12595 crate::gpu::GraphPrismOp::Forward
12596 } else {
12597 crate::gpu::GraphPrismOp::None
12598 };
12599 return Some(crate::gpu::GraphW {
12600 idx: i,
12601 kind,
12602 row_scale: rs,
12603 data: &[],
12604 prism,
12605 affine: crate::prism::is_affine_target(m, name),
12606 });
12607 }
12608 if std::env::var_os("CMF_BATCH_DEBUG").is_some() {
12609 eprintln!(
12610 "batch graph: tensor has no graph descriptor/f32 fallback rows={} cols={}",
12611 t.rows(),
12612 t.cols()
12613 );
12614 }
12615 t.as_f32().map(|d| crate::gpu::GraphW {
12616 idx: 0,
12617 kind: 4,
12618 row_scale: &[],
12619 data: d,
12620 prism: crate::gpu::GraphPrismOp::None,
12621 affine: false,
12622 })
12623 }
12624 let built: Option<(
12625 Vec<crate::gpu::GraphLayer<'_>>,
12626 std::sync::Arc<cortiq_core::CmfModel>,
12627 )> = (|| {
12628 let mut layers = Vec::with_capacity(graph_end);
12629 let mut model = None;
12630 for li in 0..graph_end {
12631 let lw = &self.weights.layers[self.phys_layer(li)];
12632 let gffn = match &lw.ffn {
12639 FfnKind::Dense(d) if !d.segs.is_empty() => {
12640 if batch_debug {
12641 eprintln!("batch graph: dense segmented FFN at layer {li}");
12642 }
12643 return None;
12644 }
12645 FfnKind::Dense(d) => crate::gpu::GraphFfn::Dense {
12646 gate: gw(&d.gate_proj)?,
12647 up: gw(&d.up_proj)?,
12648 down: gw(&d.down_proj)?,
12649 },
12650 FfnKind::Moe(m) => {
12651 if m.route_tau.is_some() || m.mask.is_some() {
12658 return None;
12659 }
12660 let shared = m.shared.as_ref();
12664 let has_shared = shared.is_some();
12665 let shared_gated = matches!(shared, Some((_, Some(_))));
12666 let sgate = match shared {
12667 Some((_, Some(sg))) => gw(sg)?,
12668 _ => gw(&m.router)?,
12672 };
12673 let router = gw(&m.router)?;
12674 if router.prism != crate::gpu::GraphPrismOp::None
12680 || router.affine
12681 || sgate.prism != crate::gpu::GraphPrismOp::None
12682 || sgate.affine
12683 {
12684 return None;
12685 }
12686 let inter = m.experts.first()?.gate_proj.rows();
12687 let mut experts = Vec::with_capacity(m.experts.len() + 1);
12688 let mut q4tp: Option<bool> = None;
12689 let mut gu_q2: Option<bool> = None;
12690 for e in m.experts.iter().chain(shared.map(|(se, _)| se)) {
12691 if !matches!(e.act, Act::Silu)
12692 || e.gate_proj.rows() != inter
12693 || e.up_proj.rows() != inter
12694 {
12695 return None;
12696 }
12697 let (mm, gi, ui, di, is_p, is_q2) = match e.gate_proj.mapped_q4t() {
12701 Some((mm, gi)) => (
12702 mm,
12703 gi,
12704 e.up_proj.mapped_q4t()?.1,
12705 e.down_proj.mapped_q4t()?.1,
12706 false,
12707 false,
12708 ),
12709 None => match e.gate_proj.mapped_q2tp() {
12710 Some((mm, gi)) => (
12711 mm,
12712 gi,
12713 e.up_proj.mapped_q2tp()?.1,
12714 e.down_proj.mapped_q4tp()?.1,
12715 true,
12716 true,
12717 ),
12718 None => {
12719 let (mm, gi) = e.gate_proj.mapped_q4tp()?;
12720 (
12721 mm,
12722 gi,
12723 e.up_proj.mapped_q4tp()?.1,
12724 e.down_proj.mapped_q4tp()?.1,
12725 true,
12726 false,
12727 )
12728 }
12729 },
12730 };
12731 if *q4tp.get_or_insert(is_p) != is_p
12732 || *gu_q2.get_or_insert(is_q2) != is_q2
12733 {
12734 return None;
12735 }
12736 if [gi, ui, di].into_iter().any(|idx| {
12737 mm.tensors
12738 .get(idx)
12739 .is_some_and(|t| {
12740 crate::prism::is_forward_weight(mm, &t.name)
12741 || crate::prism::is_affine_target(mm, &t.name)
12742 })
12743 }) {
12744 return None;
12745 }
12746 model.get_or_insert_with(|| mm.clone());
12747 experts.push((gi, ui, di));
12748 }
12749 crate::gpu::GraphFfn::Moe {
12750 router,
12751 shared_gate: sgate,
12752 experts,
12753 n_exp: m.experts.len(),
12754 top_k: m.top_k,
12755 inter,
12756 norm_topk: m.norm_topk_prob,
12757 q4tp: q4tp?,
12758 gu_q2: gu_q2.unwrap_or(false),
12759 sigmoid: m.router_sigmoid,
12760 bias: m.expert_bias.as_deref(),
12761 has_shared,
12762 shared_gated,
12763 route_scale: m.routed_scaling,
12764 }
12765 }
12766 _ => return None,
12767 };
12768 let attn = match &lw.attn {
12769 AttnKind::Full {
12770 wq,
12771 wk,
12772 wv,
12773 wo,
12774 q_norm,
12775 k_norm,
12776 output_gate,
12777 softplus_gate,
12778 bias,
12779 } => {
12780 if softplus_gate.is_some() || self.attention_heads_per_layer.is_some() {
12781 if batch_debug {
12782 eprintln!(
12783 "batch graph: unsupported Full attention gate at layer {li} softplus={} heads={}",
12784 softplus_gate.is_some(),
12785 self.attention_heads_per_layer.is_some()
12786 );
12787 }
12788 return None;
12789 }
12790 let (m, _, _, _) = wq
12791 .graph_weight()
12792 .or_else(|| wq.graph_weight_descriptor())?;
12793 model = Some(m.clone());
12794 crate::gpu::GraphAttn::Full {
12795 wq: gw(wq)?,
12796 wk: gw(wk)?,
12797 wv: gw(wv)?,
12798 wo: gw(wo)?,
12799 q_norm: q_norm.as_deref(),
12800 k_norm: k_norm.as_deref(),
12801 late_qk_norm: self.qk_norm_after_rope,
12802 bias: bias
12803 .as_ref()
12804 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
12805 output_gate: *output_gate,
12806 cpu_k: self.kv_cache.layers[li].k_heads(),
12807 cpu_v: self.kv_cache.layers[li].v_heads(),
12808 geom: self.graph_attn_geom(li),
12809 }
12810 }
12811 AttnKind::LinearGdn(w) => {
12812 let Some(cfg) = self.gdn_cfg else {
12813 if batch_debug {
12814 eprintln!("batch graph: no GDN config at layer {li}");
12815 }
12816 return None;
12817 };
12818 let (m, _, _, _) = w
12819 .in_proj_qkv
12820 .graph_weight()
12821 .or_else(|| w.in_proj_qkv.graph_weight_descriptor())?;
12822 model = Some(m.clone());
12823 crate::gpu::GraphAttn::Gdn {
12824 qkv: gw(&w.in_proj_qkv)?,
12825 z: gw(&w.in_proj_z)?,
12826 a: gw(&w.in_proj_a)?,
12827 b: gw(&w.in_proj_b)?,
12828 out: gw(&w.out_proj)?,
12829 conv1d: &w.conv1d,
12830 a_log: &w.a_log,
12831 dt_bias: &w.dt_bias,
12832 norm: &w.norm,
12833 nv: cfg.num_v_heads,
12834 nk: cfg.num_k_heads,
12835 dk: cfg.key_head_dim,
12836 dv: cfg.value_head_dim,
12837 kk: cfg.conv_kernel,
12838 cpu_state: &self.kv_cache.layers[self.phys_layer(li)].linear_state,
12839 }
12840 }
12841 _ => return None,
12842 };
12843 layers.push(crate::gpu::GraphLayer {
12844 input_norm: &lw.input_norm,
12845 attn,
12846 post_norm: &lw.post_norm,
12847 ffn: gffn,
12848 });
12849 }
12850 Some((layers, model?))
12851 })();
12852 let Some((layers, model)) = built else {
12853 {
12854 use std::sync::atomic::{AtomicBool, Ordering};
12855 static SAID: AtomicBool = AtomicBool::new(false);
12856 if !SAID.swap(true, Ordering::Relaxed) {
12857 tracing::warn!("batch graph: BUILDER refused (layer weights/kinds)");
12858 }
12859 }
12860 return crate::gpu::BatchGraphOutcome::Declined;
12861 };
12862 if std::env::var("CMF_GRAPH_SPEC_TIME").is_ok() {
12863 eprintln!("batch-build: {:.1} ms", _tb.elapsed().as_secs_f64() * 1e3);
12864 }
12865 crate::gpu::forward_batch_graph(
12866 &model,
12867 self.graph_kv_id,
12868 &layers,
12869 &self.inv_freq,
12870 hiddens,
12871 nh,
12872 nkv,
12873 hd,
12874 rd,
12875 self.hidden_size,
12876 self.intermediate_size,
12877 positions,
12878 self.kv_cache.max_seq_len,
12879 gemma,
12880 self.rms_eps as f32,
12881 self.attn_scale,
12882 k,
12883 &(0..graph_end)
12884 .map(|li| self.kv_cache.layers[self.phys_layer(li)].o1_views())
12885 .collect::<Vec<_>>(),
12886 self.o1_epoch,
12887 spec,
12888 layers_run,
12889 )
12890 }
12891
12892 fn draft_probe() -> bool {
12896 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
12897 *ON.get_or_init(|| std::env::var("CMF_DSV4_DRAFT_PROBE").is_ok_and(|v| v != "0"))
12898 }
12899
12900 #[cfg(feature = "gpu")]
12912 fn dsv4_spec_on() -> bool {
12913 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
12914 *ON.get_or_init(|| {
12915 if let Ok(v) = std::env::var("CMF_DSV4_SPEC_RUN") {
12919 return v != "0";
12920 }
12921 std::env::var("CMF_DSV4_SPEC")
12928 .map(|v| v != "0")
12929 .unwrap_or_else(|_| {
12930 crate::gpu_wgpu::DRAFT_RESERVE.load(std::sync::atomic::Ordering::Relaxed) > 0
12931 })
12932 })
12933 }
12934
12935 #[cfg(feature = "gpu")]
12942 fn dsv4_spec_step(
12943 &mut self,
12944 tip_token: u32,
12945 t_next: u32,
12946 next_pos: usize,
12947 max_extra: usize,
12948 drafted: &mut usize,
12949 accepted_ctr: &mut usize,
12950 ) -> Option<(Vec<u32>, usize)> {
12951 let t_all = std::time::Instant::now();
12952 if std::env::var("CMF_DSV4_SPEC_TIME").is_ok() {
12953 thread_local! {
12954 static LAST: std::cell::Cell<Option<std::time::Instant>> =
12955 const { std::cell::Cell::new(None) };
12956 }
12957 LAST.with(|l| {
12958 if let Some(prev) = l.get() {
12959 eprintln!(
12960 "между раундами {:.1} мс",
12961 prev.elapsed().as_secs_f64() * 1e3
12962 );
12963 }
12964 l.set(Some(std::time::Instant::now()));
12965 });
12966 }
12967 if std::env::var("CMF_DSV4_SPEC_DEBUG").is_ok() {
12968 eprintln!("spec_step: вход pos={next_pos}");
12969 }
12970 let n_layers = self.dsv4.as_ref().map(|b| b.1.len())?;
12971 let cfg = self.dsv4.as_ref().map(|b| b.2)?;
12972 if self.dspark.is_none() {
12974 let t = crate::dsv4::dspark_targets(&self.dsv4_mtp, &cfg, n_layers);
12975 if t.is_empty() {
12976 return None;
12977 }
12978 crate::dsv4::dspark_arm(&t, cfg.dim);
12979 self.dspark = Some(crate::dsv4::DsparkState::new(
12980 self.dsv4_mtp.len(),
12981 &cfg,
12982 t.len(),
12983 ));
12984 }
12985 let targets = crate::dsv4::dspark_targets(&self.dsv4_mtp, &cfg, n_layers);
12986 let pack = crate::dsv4::dspark_pack_get(&self.dsv4_mtp, &cfg);
12987 if pack.is_none() && std::env::var("CMF_DSV4_SPEC_DEBUG").is_ok() {
12988 eprintln!("spec_step: пак не построился (targets {targets:?})");
12989 }
12990 let pack = pack?;
12991 let block = crate::dsv4::dspark_block();
12992 let b_box = self.dsv4.as_mut()?;
12993 let (g, layers, st) = (&b_box.0, &b_box.1, &mut b_box.3);
12994 let ds = self.dspark.as_mut()?;
12995 let dbg = std::env::var("CMF_DSV4_SPEC_DEBUG").is_ok();
12998 if !crate::dsv4::dspark_take(&mut ds.main_hidden) && !ds.have_hidden {
12999 if dbg {
13000 eprintln!("spec_step: нет захвата");
13001 }
13002 return None;
13003 }
13004 ds.have_hidden = true;
13005 let tip_pos = next_pos.checked_sub(1)?;
13006 let draft_started = std::time::Instant::now();
13007 let mut conf = Vec::new();
13008 let props = crate::dsv4::dspark_draft_gpu(
13009 g,
13010 &self.dsv4_mtp,
13011 &cfg,
13012 ds,
13013 pack,
13014 st.kv_id,
13015 tip_token,
13016 tip_pos,
13017 self.pool.as_deref(),
13018 &mut conf,
13019 );
13020 self.dspark_draft_ns += draft_started.elapsed().as_nanos();
13021 *drafted += block;
13022 if props.is_empty() || props[0] != t_next {
13023 if dbg {
13024 eprintln!(
13025 "spec_step: черновик {} (props0={:?} t_next={t_next})",
13026 if props.is_empty() {
13027 "пуст"
13028 } else {
13029 "мимо"
13030 },
13031 props.first()
13032 );
13033 }
13034 return None;
13035 }
13036 let mut k_verify = crate::dsv4::dspark_verify_k()
13043 .min(props.len())
13044 .min(max_extra.saturating_add(1));
13045 let conf_min = {
13051 static M: std::sync::OnceLock<f32> = std::sync::OnceLock::new();
13052 *M.get_or_init(|| {
13053 std::env::var("CMF_DSPARK_CONF_MIN")
13054 .ok()
13055 .and_then(|v| v.parse().ok())
13056 .unwrap_or(0.0)
13057 })
13058 };
13059 if conf_min > 0.0 && conf.len() >= props.len() {
13060 let mut keep = 1usize;
13061 while keep < k_verify && conf.get(keep).copied().unwrap_or(0.0) >= conf_min {
13062 keep += 1;
13063 }
13064 k_verify = k_verify.min(keep.max(2));
13065 }
13066 if k_verify < 2 {
13067 return None;
13068 }
13069 let mut fed = Vec::with_capacity(k_verify);
13070 fed.push(t_next);
13071 fed.extend_from_slice(&props[1..k_verify]);
13072 let mut argmax = Vec::new();
13073 let mut logits_all = Vec::new();
13074 let mut walked = Vec::new();
13075 let txn = crate::dsv4::dsv4_verify_chunk(
13076 g,
13077 layers,
13078 &cfg,
13079 st,
13080 &fed,
13081 next_pos,
13082 &self.inv_freq,
13083 self.pool.as_deref(),
13084 &targets,
13085 &mut argmax,
13086 &mut logits_all,
13087 &mut walked,
13088 );
13089 if txn.is_none() && dbg {
13090 eprintln!("spec_step: verify отказал");
13091 }
13092 let txn = txn?;
13093 let spec_gpu_end = txn.gpu_end;
13094 let b = fed.len();
13095 let mut accepted = 1usize;
13096 while accepted < b && fed[accepted] == argmax[accepted - 1] {
13097 accepted += 1;
13098 }
13099 if std::env::var("CMF_DSV4_SPEC_FORCE_REJECT").is_ok_and(|v| v != "0") {
13104 accepted = 1;
13105 }
13106 if std::env::var("CMF_DSV4_SPEC_TRACE").is_ok() {
13107 eprintln!("spec@{next_pos}: fed={fed:?} argmax={argmax:?} accepted={accepted}");
13108 }
13109 let t_fin = std::time::Instant::now();
13110 if !crate::dsv4::dsv4_spec_finish(
13111 g,
13112 layers,
13113 &cfg,
13114 st,
13115 txn,
13116 accepted,
13117 &fed,
13118 &self.inv_freq,
13119 self.pool.as_deref(),
13120 ) {
13121 tracing::warn!("dsv4: спекулятивный откат не удался — состояние подозрительно");
13122 return None;
13123 }
13124 if std::env::var("CMF_DSV4_SPEC_TIME").is_ok() {
13125 eprintln!(
13126 "finish(k={accepted}): {:.1} мс",
13127 t_fin.elapsed().as_secs_f64() * 1e3
13128 );
13129 }
13130 *accepted_ctr += accepted - 1;
13131 let (hc, dim) = (cfg.hc_mult, cfg.dim);
13136 let dev_caps: Vec<usize> = targets
13141 .iter()
13142 .copied()
13143 .filter(|&t| t < spec_gpu_end)
13144 .collect();
13145 let mut caps_all = vec![0.0f32; dev_caps.len() * b * hc * dim];
13146 if !crate::gpu_wgpu::dsv4_spec_cap_read_all(b, dev_caps.len(), hc * dim, &mut caps_all) {
13147 return None;
13148 }
13149 for t in 0..accepted {
13150 let tip = t + 1 == accepted;
13151 for (slot, &tl) in targets.iter().enumerate() {
13152 if let Some(di) = dev_caps.iter().position(|&d| d == tl) {
13153 let lo = (di * b + t) * hc * dim;
13154 crate::dsv4::dspark_capture(
13155 &caps_all[lo..lo + hc * dim],
13156 &cfg,
13157 slot,
13158 &mut ds.main_hidden,
13159 );
13160 } else if tip
13161 && crate::dsv4::dspark_peek_slot(slot, dim, {
13162 let lo = slot * dim;
13163 &mut ds.main_hidden[lo..lo + dim]
13164 })
13165 {
13166 } else {
13171 crate::dsv4::dspark_capture(
13175 &walked[t * hc * dim..(t + 1) * hc * dim],
13176 &cfg,
13177 slot,
13178 &mut ds.main_hidden,
13179 );
13180 }
13181 }
13182 crate::dsv4::dspark_ring_append(
13183 g,
13184 &self.dsv4_mtp,
13185 &cfg,
13186 ds,
13187 next_pos + t,
13188 self.pool.as_deref(),
13189 );
13190 }
13191 let row = logits_all[(accepted - 1) * cfg.vocab..accepted * cfg.vocab].to_vec();
13192 self.graph_logits = Some(row);
13193 if std::env::var("CMF_DSV4_TRUNK_PICK_DUMP").is_ok() {
13198 crate::dsv4::trunk_freq_note(&crate::dsv4::pick_tally_take());
13199 crate::dsv4::pick_tally_arm();
13200 }
13201 if std::env::var("CMF_DSV4_SPEC_TIME").is_ok() {
13202 eprintln!(
13203 "spec_step total {:.1} мс (k={accepted})",
13204 t_all.elapsed().as_secs_f64() * 1e3
13205 );
13206 }
13207 Some((fed[1..accepted].to_vec(), next_pos + accepted))
13208 }
13209
13210 fn dspark_probe(&mut self, position: usize, token_id: u32) {
13211 if self.dsv4_mtp.is_empty() || !Self::draft_probe() {
13212 return;
13213 }
13214 let trunk_now = crate::dsv4::pick_tally_take();
13216 crate::dsv4::trunk_freq_note(&trunk_now);
13217 if !trunk_now.is_empty() {
13218 self.dspark_trunk_picks.push(trunk_now);
13219 let keep = crate::dsv4::dspark_block();
13220 if self.dspark_trunk_picks.len() > keep {
13221 self.dspark_trunk_picks.remove(0);
13222 }
13223 }
13224 for p in std::mem::take(&mut self.dspark_pending) {
13227 let Some(i) = position.checked_sub(p.0 + 1) else {
13228 continue;
13229 };
13230 let mut p = p;
13231 if i < p.1.len() {
13232 if p.2 && p.1[i] == token_id {
13233 p.3 = i + 1;
13234 } else {
13235 p.2 = false;
13236 }
13237 if i + 1 < p.1.len() {
13238 self.dspark_pending.push(p);
13239 continue;
13240 }
13241 }
13242 self.dspark_hist.push(p.3);
13243 self.dspark_real.push(token_id);
13244 }
13245 let Some(b) = &mut self.dsv4 else { return };
13246 let (g, layers, cfg) = (&b.0, &b.1, b.2);
13247 let n_layers = layers.len();
13248 if self.dspark.is_none() {
13249 let t = crate::dsv4::dspark_targets(&self.dsv4_mtp, &cfg, n_layers);
13250 if t.is_empty() {
13251 return;
13252 }
13253 eprintln!(
13254 "DSpark: захват со слоёв {t:?}, блок {}",
13255 crate::dsv4::dspark_block()
13256 );
13257 crate::dsv4::dspark_arm(&t, cfg.dim);
13258 self.dspark = Some(crate::dsv4::DsparkState::new(
13259 self.dsv4_mtp.len(),
13260 &cfg,
13261 t.len(),
13262 ));
13263 }
13264 let ds = self.dspark.as_mut().unwrap();
13265 if !crate::dsv4::dspark_take(&mut ds.main_hidden) {
13266 return; }
13268 let mut conf = Vec::new();
13269 crate::dsv4::pick_tally_arm();
13270 let draft_started = std::time::Instant::now();
13275 #[cfg(feature = "gpu")]
13276 let gpu_draft = crate::dsv4::dspark_gpu_on();
13277 #[cfg(not(feature = "gpu"))]
13278 let gpu_draft = false;
13279 let props = if gpu_draft {
13280 #[cfg(feature = "gpu")]
13281 {
13282 let kv_id = b.3.kv_id;
13283 match crate::dsv4::dspark_pack_get(&self.dsv4_mtp, &cfg) {
13284 Some(pk) => crate::dsv4::dspark_draft_gpu(
13285 g,
13286 &self.dsv4_mtp,
13287 &cfg,
13288 ds,
13289 pk,
13290 kv_id,
13291 token_id,
13292 position,
13293 self.pool.as_deref(),
13294 &mut conf,
13295 ),
13296 None => Vec::new(),
13297 }
13298 }
13299 #[cfg(not(feature = "gpu"))]
13300 Vec::new()
13301 } else {
13302 crate::gpu::cpu_scope(|| {
13303 crate::dsv4::dspark_draft(
13304 g,
13305 &self.dsv4_mtp,
13306 &cfg,
13307 ds,
13308 token_id,
13309 position,
13310 self.pool.as_deref(),
13311 &mut conf,
13312 )
13313 })
13314 };
13315 self.dspark_draft_ns += draft_started.elapsed().as_nanos();
13316 let draft_picks = crate::dsv4::pick_tally_take();
13317 crate::dsv4::dspark_freq_note(&draft_picks);
13318 crate::dsv4::pick_tally_arm();
13321 if !props.is_empty() {
13322 let (tu, tt) = {
13326 let flat: Vec<(usize, Vec<usize>)> = self
13327 .dspark_trunk_picks
13328 .iter()
13329 .flat_map(|v| v.iter().cloned())
13330 .collect();
13331 let mut per: std::collections::HashMap<usize, Vec<usize>> =
13333 std::collections::HashMap::new();
13334 for (li, picks) in flat {
13335 per.entry(li).or_default().extend(picks);
13336 }
13337 let n = per.len().max(1);
13338 let mut u = 0usize;
13339 let mut t = 0usize;
13340 for (_, v) in per {
13341 t += v.len();
13342 u += v.iter().collect::<std::collections::HashSet<_>>().len();
13343 }
13344 (u / n, t / n)
13345 };
13346 let (du, dt) = crate::dsv4::tally_unique(&draft_picks);
13347 self.dspark_exp.push((tu, tt, du, dt));
13348 self.dspark_pending.push((position, props, true, 0));
13349 }
13350 if self.dspark_hist.len() >= 8 && self.dspark_hist.len() % 8 == 0 {
13351 let n = self.dspark_hist.len() as f32;
13352 let mean: f32 = self.dspark_hist.iter().sum::<usize>() as f32 / n;
13353 let block = crate::dsv4::dspark_block();
13354 let mut at = vec![0usize; block + 1];
13355 for &k in &self.dspark_hist {
13356 at[k] += 1;
13357 }
13358 let mut surv = Vec::with_capacity(block);
13360 for i in 1..=block {
13361 let k = at[i..].iter().sum::<usize>() as f32 / n;
13362 surv.push(format!("{k:.2}"));
13363 }
13364 let distinct = self
13365 .dspark_real
13366 .iter()
13367 .collect::<std::collections::HashSet<_>>()
13368 .len();
13369 let (tu, tt, du, dt) = self.dspark_exp.iter().fold((0, 0, 0, 0), |a, b| {
13370 (a.0 + b.0, a.1 + b.1, a.2 + b.2, a.3 + b.3)
13371 });
13372 let m = self.dspark_exp.len().max(1);
13373 eprintln!(
13374 "DSpark: черновиков {}, принято в среднем {mean:.2} из {block} \
13375 (токенов за проход {:.2}), распределение {at:?}, выживание [{}]",
13376 self.dspark_hist.len(),
13377 mean + 1.0,
13378 surv.join(" ")
13379 );
13380 eprintln!(
13381 "DSpark: разных токенов {distinct} из {} (вырожденность), \
13382 эксперты ствол {}/{} на слой за {block} токенов, \
13383 черновик {}/{} за блок, draft {:.2} мс/блок",
13384 self.dspark_real.len(),
13385 tu / m,
13386 tt / m,
13387 du / m,
13388 dt / m,
13389 self.dspark_draft_ns as f64 / self.dspark_exp.len().max(1) as f64 / 1e6
13390 );
13391 }
13392 }
13393
13394 fn forward_layers_upto(
13395 &mut self,
13396 hidden: &[f32],
13397 position: usize,
13398 task_mask: Option<&TaskMask>,
13399 upto: Option<usize>,
13400 ) -> Vec<f32> {
13401 if let Some(plan) = self.gpu_plan.clone() {
13407 if upto.is_none() && plan.len() > 1 {
13408 let mut h = hidden.to_vec();
13409 for &(dev, from, upto_incl) in plan.iter() {
13410 h = crate::gpu::with_device(dev, || {
13411 self.forward_layers_span(&h, position, task_mask, from, Some(upto_incl))
13412 });
13413 }
13414 return h;
13415 }
13416 }
13417 self.forward_layers_span(hidden, position, task_mask, 0, upto)
13418 }
13419
13420 pub fn set_gpu_plan(&mut self, devices: Option<&[usize]>) -> Result<(), String> {
13425 self.set_gpu_plan_at(devices, None)
13426 }
13427
13428 pub fn set_gpu_plan_at(
13432 &mut self,
13433 devices: Option<&[usize]>,
13434 at: Option<usize>,
13435 ) -> Result<(), String> {
13436 let Some(devs) = devices.filter(|d| d.len() > 1) else {
13437 self.gpu_plan = None;
13438 return Ok(());
13439 };
13440 self.split_supported()?;
13441 let n = self.num_layers;
13442 if devs.len() > n {
13443 return Err(format!("{} devices for {n} layers", devs.len()));
13444 }
13445 if let Some(k) = at {
13446 if k == 0 || k >= n {
13447 return Err(format!("split at {k}: the model has {n} layers"));
13448 }
13449 if devs.len() == 2 {
13450 self.gpu_plan = Some(std::sync::Arc::new(vec![
13451 (devs[0], 0, k - 1),
13452 (devs[1], k, n - 1),
13453 ]));
13454 return Ok(());
13455 }
13456 return Err(format!(
13457 "an explicit split point takes exactly 2 devices, got {}",
13458 devs.len()
13459 ));
13460 }
13461 let per = n.div_ceil(devs.len());
13462 let mut plan = Vec::with_capacity(devs.len());
13463 let mut from = 0usize;
13464 for &d in devs {
13465 if from >= n {
13466 break;
13467 }
13468 let upto = (from + per - 1).min(n - 1);
13469 plan.push((d, from, upto));
13470 from = upto + 1;
13471 }
13472 self.gpu_plan = Some(std::sync::Arc::new(plan));
13473 Ok(())
13474 }
13475
13476 pub fn gpu_plan(&self) -> Option<Vec<(usize, usize, usize)>> {
13478 self.gpu_plan.as_ref().map(|p| p.as_ref().clone())
13479 }
13480
13481 fn embryo_qtensor_f32(t: &QTensor) -> Vec<f32> {
13487 if let Some(x) = t.as_f32() {
13488 return x.to_vec();
13489 }
13490 let mut out = vec![0.0; t.rows() * t.cols()];
13491 for r in 0..t.rows() {
13492 t.row_f32(r, &mut out[r * t.cols()..(r + 1) * t.cols()]);
13493 }
13494 out
13495 }
13496
13497 fn embryo_resident_eligible(&self) -> bool {
13498 if (self.vmf_cfg.is_none() && self.gdn_cfg.is_none())
13501 || self.num_layers != self.physical_layers
13502 || self.loop_final_norm
13503 || self.weights.layers.len() != self.num_layers
13504 || self.head_clusters.is_none()
13505 || self.final_softcap.is_some()
13506 || self.logit_multiplier.is_some()
13507 || self.attn_softcap != 0.0
13508 || self.mtp.is_some()
13509 || self.g3n.is_some()
13510 || self.dsv4.is_some()
13511 || self.dsv41.is_some()
13512 || self.qwen4_exp.is_some()
13513 || self.dyn_router.is_some()
13520 || self.dyn_phi_layer.is_some()
13521 || self.dyn_blend_loaded
13522 || self.o1_cfg.is_some()
13523 || self.swa.is_some()
13524 || self.sliding_layers.is_some()
13525 || self.global_attn.is_some()
13526 || self.attention_heads_per_layer.is_some()
13527 || self.attn_v_norm
13528 || self
13529 .kv_cache
13530 .layers
13531 .iter()
13532 .any(|l| l.mode != crate::kv_cache::KvMode::F32)
13533 || self.rope_scale != 1.0
13534 || self.rope_scale_local != 1.0
13535 || self.attn_scale != 1.0 / (self.head_dim as f32).sqrt()
13536 || self.hidden_size == 0
13537 || self.hidden_size > 1024
13538 || self.intermediate_size > 1024
13539 || self.num_heads == 0
13540 || self.num_kv_heads == 0
13541 || self.num_heads % self.num_kv_heads != 0
13542 || self.num_heads.saturating_mul(self.head_dim) > 1024
13543 || self.num_kv_heads.saturating_mul(self.head_dim) > 1024
13544 || self.vocab_size == 0
13545 || self.kv_cache.max_seq_len == 0
13546 || self.rotary_dim == 0
13547 || self.rotary_dim > self.head_dim
13548 || self.rotary_dim % 2 != 0
13549 || self.inv_freq.len() < self.rotary_dim / 2
13550 {
13551 return false;
13552 }
13553 if self.weights.lm_head.as_f32().is_none()
13564 || self.weights.embed_tokens.as_f32().is_none()
13565 || self.weights.lm_head.rows() < self.vocab_size
13566 || self.weights.lm_head.cols() != self.hidden_size
13567 || self.weights.embed_tokens.rows() < self.vocab_size
13568 || self.weights.embed_tokens.cols() != self.hidden_size
13569 || self.weights.final_norm.len() != self.hidden_size
13570 {
13571 return false;
13572 }
13573 if let Some(cfg) = self.vmf_cfg {
13574 if cfg.num_heads.saturating_mul(cfg.nphase) > 1024
13575 || cfg.num_heads.saturating_mul(cfg.nphase.saturating_add(1)) > 1024
13576 || cfg.num_heads.saturating_mul(cfg.value_head_dim) > 1024
13577 || cfg.state_len() == 0
13578 {
13579 return false;
13580 }
13581 }
13582 if let Some(g) = self.gdn_cfg {
13583 if g.num_v_heads == 0
13588 || g.num_k_heads == 0
13589 || g.num_v_heads % g.num_k_heads != 0
13590 || g.key_head_dim == 0
13591 || g.key_head_dim > 128
13592 || g.value_head_dim == 0
13593 || g.value_head_dim > 256
13594 || g.value_head_dim % 4 != 0
13595 || g.conv_kernel == 0
13596 || g.num_v_heads > 512
13597 || g.num_v_heads.saturating_mul(g.value_head_dim) > 1024
13598 || g.conv_dim() > 2048
13599 || g.conv_dim() % 4 != 0
13600 || g.hidden_size != self.hidden_size
13601 || g.output_gate_sigmoid
13602 || g.rms_eps != self.rms_eps
13603 || g.state_len() == 0
13604 {
13605 return false;
13606 }
13607 }
13608 let mut full_seen = false;
13609 for lw in &self.weights.layers {
13610 if lw.attn_out_norm.is_some() || lw.ffn_out_norm.is_some() || lw.layer_scale.is_some() {
13611 return false;
13612 }
13613 match &lw.attn {
13614 AttnKind::LinearGdn(w) => {
13615 let Some(g) = self.gdn_cfg else {
13616 return false;
13617 };
13618 let (nv, dv, kk) = (g.num_v_heads, g.value_head_dim, g.conv_kernel);
13619 if w.in_proj_qkv.rows() != g.conv_dim()
13620 || w.in_proj_qkv.cols() != self.hidden_size
13621 || w.in_proj_qkv.as_f32().is_none()
13622 || w.in_proj_z.rows() != nv * dv
13623 || w.in_proj_z.cols() != self.hidden_size
13624 || w.in_proj_z.as_f32().is_none()
13625 || w.in_proj_a.rows() != nv
13626 || w.in_proj_a.cols() != self.hidden_size
13627 || w.in_proj_a.as_f32().is_none()
13628 || w.in_proj_b.rows() != nv
13629 || w.in_proj_b.cols() != self.hidden_size
13630 || w.in_proj_b.as_f32().is_none()
13631 || w.conv1d.len() != g.conv_dim() * kk
13632 || w.a_log.len() != nv
13633 || w.dt_bias.len() != nv
13634 || w.norm.len() != dv
13635 || w.out_proj.rows() != self.hidden_size
13636 || w.out_proj.cols() != nv * dv
13637 || w.out_proj.as_f32().is_none()
13638 {
13639 return false;
13640 }
13641 }
13642 AttnKind::Linear(w) => {
13643 let Some(cfg) = self.vmf_cfg else {
13644 return false;
13645 };
13646 if w.thq.rows() != cfg.num_heads * cfg.nphase
13647 || w.thq.cols() != self.hidden_size
13648 || w.thq.as_f32().is_none()
13649 || w.thk.rows() != cfg.num_heads * cfg.nphase
13650 || w.thk.cols() != self.hidden_size
13651 || w.thk.as_f32().is_none()
13652 || w.v_proj.rows() != cfg.num_heads * cfg.value_head_dim
13653 || w.v_proj.cols() != self.hidden_size
13654 || w.v_proj.as_f32().is_none()
13655 || w.out_proj.rows() != self.hidden_size
13656 || w.out_proj.cols() != cfg.num_heads * cfg.value_head_dim
13657 || w.out_proj.as_f32().is_none()
13658 || w.decay.len() != cfg.num_heads * 2 * cfg.nphase
13659 {
13660 return false;
13661 }
13662 if let Some((kg, kb)) = &w.k_gate {
13663 if kg.rows() != cfg.num_heads
13664 || kg.cols() != self.hidden_size
13665 || kg.as_f32().is_none()
13666 || kb.len() != cfg.num_heads
13667 {
13668 return false;
13669 }
13670 }
13671 if let Some(conv) = &w.conv {
13672 if conv.len() % self.hidden_size != 0 || conv.len() / self.hidden_size < 2 {
13673 return false;
13674 }
13675 }
13676 }
13677 AttnKind::Full {
13678 wq,
13679 wk,
13680 wv,
13681 wo,
13682 q_norm,
13683 k_norm,
13684 output_gate,
13685 softplus_gate,
13686 bias,
13687 } => {
13688 if full_seen
13689 || q_norm.is_some()
13690 || k_norm.is_some()
13691 || *output_gate
13692 || softplus_gate.is_some()
13693 || bias.is_some()
13694 || wq.as_f32().is_none()
13695 || wk.as_f32().is_none()
13696 || wv.as_f32().is_none()
13697 || wo.as_f32().is_none()
13698 || wq.rows() != self.num_heads * self.head_dim
13699 || wk.rows() != self.num_kv_heads * self.head_dim
13700 || wv.rows() != self.num_kv_heads * self.head_dim
13701 || wq.cols() != self.hidden_size
13702 || wk.cols() != self.hidden_size
13703 || wv.cols() != self.hidden_size
13704 || wo.rows() != self.hidden_size
13705 || wo.cols() != self.num_heads * self.head_dim
13706 {
13707 return false;
13708 }
13709 full_seen = true;
13710 }
13711 AttnKind::Bounded(w) => {
13712 let Some(ac) = self.anchor_core.as_ref() else {
13715 return false;
13716 };
13717 if self.bounded_rope.is_none()
13718 || w.window != ac.window
13719 || w.sink != ac.sink
13720 || w.window == 0
13721 || w.window + w.sink > 256
13722 || w.sink_k.len() != self.num_kv_heads * w.sink * self.head_dim
13723 || w.sink_v.len() != self.num_kv_heads * w.sink * self.head_dim
13724 || w.wq.as_f32().is_none()
13725 || w.wk.as_f32().is_none()
13726 || w.wv.as_f32().is_none()
13727 || w.wo.as_f32().is_none()
13728 || w.wq.rows() != self.num_heads * self.head_dim
13729 || w.wk.rows() != self.num_kv_heads * self.head_dim
13730 || w.wv.rows() != self.num_kv_heads * self.head_dim
13731 || w.wq.cols() != self.hidden_size
13732 || w.wk.cols() != self.hidden_size
13733 || w.wv.cols() != self.hidden_size
13734 || w.wo.rows() != self.hidden_size
13735 || w.wo.cols() != self.num_heads * self.head_dim
13736 {
13737 return false;
13738 }
13739 }
13740 _ => return false,
13741 }
13742 match &lw.ffn {
13743 FfnKind::Dense(d) => {
13744 if d.act != Act::Silu
13745 || !d.segs.is_empty()
13746 || d.gate_proj.as_f32().is_none()
13747 || d.up_proj.as_f32().is_none()
13748 || d.down_proj.as_f32().is_none()
13749 || d.gate_proj.rows() != self.intermediate_size
13750 || d.gate_proj.cols() != self.hidden_size
13751 || d.up_proj.rows() != self.intermediate_size
13752 || d.up_proj.cols() != self.hidden_size
13753 || d.down_proj.rows() != self.hidden_size
13754 || d.down_proj.cols() != self.intermediate_size
13755 {
13756 return false;
13757 }
13758 }
13759 FfnKind::Moe(m) => {
13760 if m.resonance.is_none()
13761 || m.top_k != 1
13762 || m.router_sigmoid
13763 || !m.norm_topk_prob
13764 || m.expert_bias.is_some()
13765 || m.routed_scaling != 1.0
13766 || m.route_tau.is_some()
13767 || m.shared.is_none()
13768 || m.mask.is_some()
13769 || m.per_expert_scale.is_some()
13770 || m.router_input_norm
13771 || m.experts.is_empty()
13772 || m.experts.len() > 8
13773 {
13774 return false;
13775 }
13776 let r = m.resonance.as_ref().unwrap();
13777 if r.mu.len() != m.experts.len() * self.hidden_size
13778 || r.bias.len() != m.experts.len()
13779 || r.u.len() != m.experts.len() * r.k * self.hidden_size
13780 || r.k > 128
13781 {
13782 return false;
13783 }
13784 let Some((shared, gate)) = &m.shared else {
13785 return false;
13786 };
13787 if gate.is_some() || shared.act != Act::Silu || !shared.segs.is_empty() {
13788 return false;
13789 }
13790 if shared.gate_proj.as_f32().is_none()
13791 || shared.up_proj.as_f32().is_none()
13792 || shared.down_proj.as_f32().is_none()
13793 || shared.gate_proj.rows() != self.intermediate_size
13794 || shared.gate_proj.cols() != self.hidden_size
13795 || shared.up_proj.rows() != self.intermediate_size
13796 || shared.up_proj.cols() != self.hidden_size
13797 || shared.down_proj.rows() != self.hidden_size
13798 || shared.down_proj.cols() != self.intermediate_size
13799 {
13800 return false;
13801 }
13802 for e in &m.experts {
13803 if e.act != Act::Silu
13804 || !e.segs.is_empty()
13805 || e.gate_proj.as_f32().is_none()
13806 || e.up_proj.as_f32().is_none()
13807 || e.down_proj.as_f32().is_none()
13808 || e.gate_proj.rows() != self.intermediate_size
13809 || e.gate_proj.cols() != self.hidden_size
13810 || e.up_proj.rows() != self.intermediate_size
13811 || e.up_proj.cols() != self.hidden_size
13812 || e.down_proj.rows() != self.hidden_size
13813 || e.down_proj.cols() != self.intermediate_size
13814 {
13815 return false;
13816 }
13817 }
13818 }
13819 FfnKind::DenseMoe(_) => return false,
13820 }
13821 if lw.input_norm.len() != self.hidden_size || lw.post_norm.len() != self.hidden_size {
13822 return false;
13823 }
13824 }
13825 if full_seen && self.anchor_core.is_some() {
13826 return false;
13827 }
13828 full_seen || self.num_layers > 0
13829 }
13830
13831 fn embryo_resident_wanted(&self) -> bool {
13836 crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Decode)
13837 && crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Prefill)
13838 && matches!(
13839 std::env::var("CMF_EMBRYO_RESIDENT").as_deref(),
13840 Ok("1") | Ok("parallel")
13841 )
13842 && crate::gpu::enabled_here()
13843 && !self.graph_refused()
13844 && self.embryo_resident_eligible()
13845 }
13846
13847 fn embryo_prefill_chunked(&mut self, ids: &[u32], start: usize) -> Option<Vec<f32>> {
13855 if ids.len() < 2
13856 || std::env::var("CMF_EMBRYO_CHUNK").as_deref() == Ok("0")
13857 || !self.embryo_resident_wanted()
13858 {
13859 return None;
13860 }
13861 let model = self.ensure_embryo_graph()?;
13862 let cmax = std::env::var("CMF_EMBRYO_CHUNK")
13863 .ok()
13864 .and_then(|v| v.parse::<usize>().ok())
13865 .filter(|&v| v >= 1)
13866 .unwrap_or(crate::gpu::EMBRYO_CHUNK_MAX)
13867 .min(crate::gpu::EMBRYO_CHUNK_MAX);
13868 let hs = self.hidden_size;
13869 let n = ids.len();
13870 let mut pos = start;
13871 let mut last = None;
13872 let mut rows = Vec::with_capacity(cmax * hs);
13873 while pos < n {
13874 let end = (pos + cmax).min(n);
13875 rows.clear();
13876 for &id in &ids[pos..end] {
13877 rows.extend_from_slice(&self.embed_single(id));
13878 }
13879 let mut lg = Vec::new();
13880 if !crate::gpu::forward_embryo_graph_chunk(
13881 &model,
13882 self.graph_kv_id,
13883 &rows,
13884 pos,
13885 end - pos,
13886 &mut lg,
13887 ) {
13888 if pos == start {
13889 if std::env::var("CMF_EMBRYO_DBG").as_deref() == Ok("1") {
13890 eprintln!("embryo-dbg: chunk prefill refused at position {pos}");
13891 }
13892 return None;
13893 }
13894 self.kv_cache.clear();
13898 self.clear_history();
13899 crate::gpu::graph_kv_reset(self.graph_kv_id);
13900 panic!(
13901 "resident Embryo chunk prefill refused at position {pos}; refusing an in-flight host fallback"
13902 );
13903 }
13904 last = Some(lg);
13905 pos = end;
13906 }
13907 if std::env::var("CMF_EMBRYO_DBG").as_deref() == Ok("1") {
13908 eprintln!(
13909 "embryo-dbg: chunk prefill {} ids from {start} in {} submits (chunk {cmax})",
13910 n - start,
13911 (n - start).div_ceil(cmax)
13912 );
13913 }
13914 last
13915 }
13916
13917 fn ensure_embryo_graph(&mut self) -> Option<std::sync::Arc<crate::gpu::EmbryoGraphModel>> {
13918 if self.embryo_graph.is_none() && self.embryo_resident_eligible() {
13919 const UMAX: u32 = u32::MAX;
13920 const HEADER: usize = crate::gpu::EMBRYO_META_HEADER;
13921 const REC: usize = 64;
13922 struct Pack {
13923 data: Vec<f32>,
13924 }
13925 impl Pack {
13926 fn put(&mut self, x: &[f32]) -> u32 {
13927 if x.is_empty() {
13928 return u32::MAX;
13929 }
13930 let off = self.data.len();
13931 self.data.extend_from_slice(x);
13932 off as u32
13933 }
13934 }
13935 let vmf = self.vmf_cfg;
13939 let gdn = self.gdn_cfg;
13940 let mut pack = Pack { data: Vec::new() };
13941 let mut meta = vec![0u32; HEADER];
13942 meta[0] = self.hidden_size as u32;
13943 meta[1] = self.intermediate_size as u32;
13944 meta[2] = self.vocab_size as u32;
13945 meta[3] = self.num_layers as u32;
13946 meta[4] = vmf.map(|c| c.num_heads).unwrap_or(0) as u32;
13947 meta[5] = vmf.map(|c| c.nphase).unwrap_or(0) as u32;
13948 meta[6] = vmf.map(|c| c.value_head_dim).unwrap_or(0) as u32;
13949 if let Some(g) = gdn {
13950 meta[24] = g.num_v_heads as u32;
13951 meta[25] = g.num_k_heads as u32;
13952 meta[26] = g.key_head_dim as u32;
13953 meta[27] = g.value_head_dim as u32;
13954 meta[28] = g.conv_kernel as u32;
13955 meta[29] = g.conv_dim() as u32;
13956 }
13957 meta[7] = self.num_heads as u32;
13958 meta[8] = self.num_kv_heads as u32;
13959 meta[9] = self.head_dim as u32;
13960 meta[10] = self.kv_cache.max_seq_len as u32;
13961 let clusters = self.head_clusters.as_ref().unwrap();
13962 let cluster_count = clusters.len() / self.hidden_size;
13963 if clusters.len() % self.hidden_size != 0
13964 || cluster_count == 0
13965 || cluster_count > 1024
13966 || self.vocab_size % cluster_count != 0
13967 || self.weights.lm_head.rows() < self.vocab_size
13968 || self.weights.final_norm.len() != self.hidden_size
13969 {
13970 return None;
13971 }
13972 meta[11] = cluster_count as u32;
13973 meta[12] = (self.vocab_size / cluster_count) as u32;
13974 meta[13] = vmf.map(|c| c.state_len()).unwrap_or(0) as u32;
13975 meta[16] = self.rotary_dim as u32;
13976 meta[17] = matches!(self.norm_style, NormStyle::Gemma) as u32;
13977 meta[19] = (self.rms_eps as f32).to_bits();
13978 let max_conv = self
13979 .weights
13980 .layers
13981 .iter()
13982 .filter_map(|lw| match &lw.attn {
13983 AttnKind::Linear(w) => w.conv.as_ref().map(|v| v.len() / self.hidden_size),
13984 _ => None,
13985 })
13986 .max()
13987 .unwrap_or(1);
13988 let phase_stride = vmf
13994 .map(|c| c.state_len() + max_conv.saturating_sub(1) * self.hidden_size)
13995 .unwrap_or(0);
13996 let gdn_stride = gdn.map(|g| g.state_len()).unwrap_or(0);
13997 let state_stride = phase_stride.max(gdn_stride);
13998 let bounded = self.anchor_core.clone();
14002 let (anchor_window, anchor_sink) = bounded
14003 .as_ref()
14004 .map(|ac| (ac.window, ac.sink))
14005 .unwrap_or((0, 0));
14006 let kv_stride = if bounded.is_some() {
14007 2usize
14008 .saturating_mul(self.num_kv_heads)
14009 .saturating_mul(anchor_window)
14010 .saturating_mul(self.head_dim)
14011 } else {
14012 2usize
14013 .saturating_mul(self.num_kv_heads)
14014 .saturating_mul(self.kv_cache.max_seq_len)
14015 .saturating_mul(self.head_dim)
14016 };
14017 meta[14] = state_stride as u32;
14018 meta[15] = kv_stride as u32;
14019 meta[18] = anchor_window as u32;
14020 meta[20] = anchor_sink as u32;
14021 meta[21] = match &self.bounded_rope {
14022 Some(rope) => {
14023 let off = pack.put(&rope.cos);
14025 let _ = pack.put(&rope.sin);
14026 off
14027 }
14028 None => UMAX,
14029 };
14030 let mut full_seen = false;
14031 let mut bounded_seen = 0usize;
14032 let mut phase_seen = 0usize;
14036 let mut gdn_seen = 0usize;
14037 for (li, lw) in self.weights.layers.iter().enumerate() {
14038 let base = meta.len();
14039 meta.resize(base + REC, UMAX);
14040 meta[base] = match &lw.attn {
14041 AttnKind::Linear(w) if w.phase_delta => 1,
14042 AttnKind::Linear(_) => 0,
14043 AttnKind::Full { .. } => 2,
14044 AttnKind::Bounded(_) => 3,
14045 AttnKind::LinearGdn(_) => 4,
14046 _ => UMAX,
14047 };
14048 meta[base + 1] = pack.put(&lw.input_norm);
14049 meta[base + 2] = pack.put(&lw.post_norm);
14050 meta[base + 25] = match &lw.attn {
14051 AttnKind::Linear(_) => {
14052 let off = ((phase_seen + gdn_seen) * state_stride) as u32;
14053 phase_seen += 1;
14054 off
14055 }
14056 AttnKind::LinearGdn(_) => {
14057 let off = ((phase_seen + gdn_seen) * state_stride) as u32;
14058 gdn_seen += 1;
14059 off
14060 }
14061 _ => UMAX,
14062 };
14063 match &lw.attn {
14064 AttnKind::LinearGdn(w) => {
14065 meta[base + 56] = pack.put(&Self::embryo_qtensor_f32(&w.in_proj_qkv));
14068 meta[base + 57] = pack.put(&Self::embryo_qtensor_f32(&w.in_proj_z));
14069 meta[base + 58] = pack.put(&Self::embryo_qtensor_f32(&w.in_proj_a));
14070 meta[base + 59] = pack.put(&Self::embryo_qtensor_f32(&w.in_proj_b));
14071 meta[base + 60] = pack.put(&w.conv1d);
14072 meta[base + 61] = pack.put(&w.a_log);
14073 meta[base + 62] = pack.put(&w.dt_bias);
14074 meta[base + 63] = pack.put(&w.norm);
14075 meta[base + 29] = pack.put(&Self::embryo_qtensor_f32(&w.out_proj));
14076 meta[base + 24] = 0;
14077 }
14078 AttnKind::Linear(w) => {
14079 meta[base + 3] = pack.put(&Self::embryo_qtensor_f32(&w.thq));
14080 meta[base + 4] = pack.put(&Self::embryo_qtensor_f32(&w.thk));
14081 meta[base + 5] = pack.put(&Self::embryo_qtensor_f32(&w.v_proj));
14082 meta[base + 6] = pack.put(&Self::embryo_qtensor_f32(&w.out_proj));
14083 let decay: Vec<f32> = w.decay.iter().map(|&x| x as f32).collect();
14084 meta[base + 7] = pack.put(&decay);
14085 if let Some((kg, kb)) = &w.k_gate {
14086 meta[base + 8] = pack.put(&Self::embryo_qtensor_f32(kg));
14087 meta[base + 9] = pack.put(kb);
14088 }
14089 if let Some(conv) = &w.conv {
14090 meta[base + 10] = pack.put(conv);
14091 meta[base + 24] = (conv.len() / self.hidden_size) as u32;
14092 } else {
14093 meta[base + 24] = 0;
14094 }
14095 }
14096 AttnKind::Full { wq, wk, wv, wo, .. } => {
14097 full_seen = true;
14098 meta[base + 11] = pack.put(&Self::embryo_qtensor_f32(wq));
14099 meta[base + 12] = pack.put(&Self::embryo_qtensor_f32(wk));
14100 meta[base + 13] = pack.put(&Self::embryo_qtensor_f32(wv));
14101 meta[base + 14] = pack.put(&Self::embryo_qtensor_f32(wo));
14102 meta[base + 26] = (li * kv_stride) as u32;
14103 }
14104 AttnKind::Bounded(w) => {
14105 meta[base + 11] = pack.put(&Self::embryo_qtensor_f32(&w.wq));
14106 meta[base + 12] = pack.put(&Self::embryo_qtensor_f32(&w.wk));
14107 meta[base + 13] = pack.put(&Self::embryo_qtensor_f32(&w.wv));
14108 meta[base + 14] = pack.put(&Self::embryo_qtensor_f32(&w.wo));
14109 meta[base + 26] = (bounded_seen * kv_stride) as u32;
14111 meta[base + 27] = pack.put(&w.sink_k);
14112 meta[base + 28] = pack.put(&w.sink_v);
14113 bounded_seen += 1;
14114 }
14115 _ => return None,
14116 }
14117 match &lw.ffn {
14118 FfnKind::Dense(d) => {
14119 meta[base + 15] = 0;
14120 meta[base + 16] = 0;
14121 meta[base + 21] = pack.put(&Self::embryo_qtensor_f32(&d.gate_proj));
14122 meta[base + 22] = pack.put(&Self::embryo_qtensor_f32(&d.up_proj));
14123 meta[base + 23] = pack.put(&Self::embryo_qtensor_f32(&d.down_proj));
14124 }
14125 FfnKind::Moe(m) => {
14126 let r = m.resonance.as_ref().unwrap();
14127 let (shared, _) = m.shared.as_ref().unwrap();
14128 meta[base + 15] = 1;
14129 meta[base + 16] = m.experts.len() as u32;
14130 meta[base + 17] = pack.put(&r.mu);
14131 meta[base + 18] = pack.put(&r.u);
14132 meta[base + 19] = pack.put(&r.bias);
14133 meta[base + 20] = r.k as u32;
14134 let mut shell = r.effective_shell(m.experts.len());
14143 shell.push(f32::NEG_INFINITY);
14144 meta[base + 30] = pack.put(&shell);
14145 meta[base + 21] = pack.put(&Self::embryo_qtensor_f32(&shared.gate_proj));
14146 meta[base + 22] = pack.put(&Self::embryo_qtensor_f32(&shared.up_proj));
14147 meta[base + 23] = pack.put(&Self::embryo_qtensor_f32(&shared.down_proj));
14148 for (e, ex) in m.experts.iter().enumerate() {
14149 meta[base + 32 + e * 3] =
14150 pack.put(&Self::embryo_qtensor_f32(&ex.gate_proj));
14151 meta[base + 33 + e * 3] =
14152 pack.put(&Self::embryo_qtensor_f32(&ex.up_proj));
14153 meta[base + 34 + e * 3] =
14154 pack.put(&Self::embryo_qtensor_f32(&ex.down_proj));
14155 }
14156 }
14157 FfnKind::DenseMoe(_) => return None,
14158 }
14159 }
14160 if !full_seen && self.num_layers == 0 {
14161 return None;
14162 }
14163 let id = {
14164 static NEXT: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(1);
14165 NEXT.fetch_add(1, std::sync::atomic::Ordering::Relaxed)
14166 };
14167 let model = crate::gpu::EmbryoGraphModel {
14168 id,
14169 hidden: self.hidden_size,
14170 intermediate: self.intermediate_size,
14171 vocab: self.vocab_size,
14172 layers: self.num_layers,
14173 phase_heads: vmf.map(|c| c.num_heads).unwrap_or(0),
14174 nphase: vmf.map(|c| c.nphase).unwrap_or(0),
14175 phase_dv: vmf.map(|c| c.value_head_dim).unwrap_or(0),
14176 anchor_q_heads: self.num_heads,
14177 anchor_kv_heads: self.num_kv_heads,
14178 anchor_head_dim: self.head_dim,
14179 rotary_dim: self.rotary_dim,
14180 max_seq: self.kv_cache.max_seq_len,
14181 cluster_count,
14182 cluster_size: self.vocab_size / cluster_count,
14183 phase_state_len: vmf.map(|c| c.state_len()).unwrap_or(0),
14184 state_stride,
14185 kv_stride,
14186 norm_gemma: matches!(self.norm_style, NormStyle::Gemma),
14187 phase_mass: vmf.map(|c| c.phase_mass).unwrap_or(0.0),
14188 weights: pack.data,
14189 meta,
14190 lm_head: Self::embryo_qtensor_f32(&self.weights.lm_head),
14191 clusters: clusters.as_ref().clone(),
14192 final_norm: self.weights.final_norm.clone(),
14193 inv_freq: self.inv_freq.as_ref().clone(),
14194 bounded: bounded.is_some(),
14195 kv_layers: if bounded.is_some() {
14196 bounded_seen
14197 } else {
14198 self.num_layers
14199 },
14200 state_layers: phase_seen + gdn_seen,
14201 anchor_window,
14202 anchor_sink,
14203 phase_layers: phase_seen,
14204 gdn_layers: gdn_seen,
14205 gdn_heads: gdn.map(|g| g.num_v_heads).unwrap_or(0),
14206 gdn_k_heads: gdn.map(|g| g.num_k_heads).unwrap_or(0),
14207 gdn_dk: gdn.map(|g| g.key_head_dim).unwrap_or(0),
14208 gdn_dv: gdn.map(|g| g.value_head_dim).unwrap_or(0),
14209 gdn_kk: gdn.map(|g| g.conv_kernel).unwrap_or(0),
14210 };
14211 self.embryo_graph = Some(std::sync::Arc::new(model));
14212 }
14213 self.embryo_graph.clone()
14214 }
14215
14216 fn forward_layers_span(
14217 &mut self,
14218 hidden: &[f32],
14219 position: usize,
14220 task_mask: Option<&TaskMask>,
14221 from: usize,
14222 upto: Option<usize>,
14223 ) -> Vec<f32> {
14224 debug_assert!(
14225 from == 0
14226 || (self.dsv4.is_none()
14227 && self.dsv41.is_none()
14228 && self.qwen4_exp.is_none()
14229 && self.g3n.is_none())
14230 );
14231 #[cfg(target_os = "macos")]
14237 if !crate::gpu_metal::wait_replay() {
14238 self.fail_metal_graph("the pending async replay failed before a plain forward");
14239 return vec![0.0; self.hidden_size];
14240 }
14241 if let Some(b) = &mut self.qwen4_exp {
14242 let _ = (task_mask, upto);
14243 let token_id = hidden.first().copied().unwrap_or(0.0) as u32;
14244 let mut logits = Vec::new();
14245 crate::qwen4_exp::forward_token(
14246 &b.0,
14247 &b.1,
14248 &b.2,
14249 &mut b.3,
14250 token_id,
14251 position,
14252 &self.inv_freq,
14253 self.pool.as_deref(),
14254 &mut logits,
14255 true,
14256 );
14257 self.graph_logits = Some(logits);
14258 return vec![0.0; self.hidden_size];
14259 }
14260 if let Some(b) = &mut self.dsv4 {
14266 let _ = (task_mask, upto);
14267 let token_id = hidden.first().copied().unwrap_or(0.0) as u32;
14268 let (g, layers, cfg, st) = (&b.0, &b.1, b.2, &mut b.3);
14269 st.pos = position;
14270 let mut logits = Vec::new();
14271 crate::dsv4::forward_token(
14272 g,
14273 layers,
14274 &cfg,
14275 st,
14276 token_id,
14277 &self.inv_freq,
14278 self.pool.as_deref(),
14279 &mut logits,
14280 );
14281 self.graph_logits = Some(logits);
14282 self.dspark_probe(position, token_id);
14283 return vec![0.0; self.hidden_size];
14286 }
14287 if let Some(b) = &mut self.dsv41 {
14289 let _ = (task_mask, upto);
14290 let token_id = hidden.first().copied().unwrap_or(0.0) as u32;
14291 let mut logits = Vec::new();
14292 crate::dsv41::forward_token(
14293 &b.0,
14294 &b.1,
14295 &b.2,
14296 &mut b.3,
14297 token_id,
14298 position,
14299 self.pool.as_deref(),
14300 &mut logits,
14301 );
14302 self.graph_logits = Some(logits);
14303 return vec![0.0; self.hidden_size];
14304 }
14305 if let Some(b) = &self.g3n {
14308 let _ = (task_mask, upto);
14309 return crate::g3n::g3n_forward(
14310 &b.0,
14311 &b.1,
14312 hidden,
14313 position,
14314 &mut self.kv_cache.layers,
14315 self.num_heads,
14316 self.num_kv_heads,
14317 self.head_dim,
14318 self.pool.as_deref(),
14319 );
14320 }
14321 if from == 0
14328 && upto.is_none()
14329 && task_mask.is_none()
14330 && self.anchor_core.is_some()
14331 && std::env::var("CMF_EMBRYO_DBG").as_deref() == Ok("1")
14332 {
14333 static ONCE: std::sync::Once = std::sync::Once::new();
14334 ONCE.call_once(|| {
14335 eprintln!(
14336 "embryo-dbg: graph_on decode={} prefill={} resident_env={:?} enabled_here={} \
14337 unsupported={} eligible={}",
14338 crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Decode),
14339 crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Prefill),
14340 std::env::var("CMF_EMBRYO_RESIDENT").ok(),
14341 crate::gpu::enabled_here(),
14342 self.graph_refused(),
14343 self.embryo_resident_eligible(),
14344 );
14345 });
14346 }
14347 if from == 0
14348 && upto.is_none()
14349 && task_mask.is_none()
14350 && crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Decode)
14354 && crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Prefill)
14355 && matches!(
14359 std::env::var("CMF_EMBRYO_RESIDENT").as_deref(),
14360 Ok("1") | Ok("parallel")
14361 )
14362 && crate::gpu::enabled_here()
14363 && !self.graph_refused()
14364 && (position == 0 || self.device_sequence_position().is_some())
14370 && self.embryo_resident_eligible()
14371 && let Some(model) = self.ensure_embryo_graph()
14372 {
14373 let mut lg = Vec::new();
14374 if crate::gpu::forward_embryo_graph(&model, self.graph_kv_id, hidden, position, &mut lg)
14375 {
14376 self.graph_logits = Some(lg);
14377 return vec![0.0; self.hidden_size];
14378 }
14379 if std::env::var("CMF_EMBRYO_DBG").as_deref() == Ok("1") {
14380 eprintln!("embryo-dbg: forward_embryo_graph refused at position {position}");
14381 }
14382 self.mark_graph_refused();
14388 if position != 0 {
14389 self.kv_cache.clear();
14393 self.clear_history();
14394 crate::gpu::graph_kv_reset(self.graph_kv_id);
14395 panic!(
14396 "resident Embryo graph refused at position {position}; refusing an in-flight host fallback"
14397 );
14398 }
14399 }
14400 let mut h = hidden.to_vec();
14401 self.mimo_moe_prepare();
14404 let _mimo_q8 = self.mimo_moe.is_on()
14405 .then(crate::qtensor::enter_full_gpu_q8_scope);
14406 let (nh, _nkv, _hd, hs, _rd, eps) = (
14409 self.num_heads,
14410 self.num_kv_heads,
14411 self.head_dim,
14412 self.hidden_size,
14413 self.rotary_dim,
14414 self.rms_eps,
14415 );
14416 let pool = self.pool.clone();
14417 let graph_env = std::env::var("CMF_GPU_WGPU_GRAPH").ok();
14429 let graph_on = match graph_env.as_deref() {
14430 Some("0") => false,
14431 Some("prefill") => false, Some(_) => true,
14433 None => crate::gpu::wgpu_graph_default(),
14439 };
14440 let graph_trusted =
14441 graph_env.is_some() || crate::gpu::wgpu_graph_default() || self.gdn_cfg.is_some();
14442 let race_eligible = graph_on
14443 && upto.is_none()
14444 && task_mask.is_none()
14445 && from == 0
14446 && !self.graph_refused();
14447 let mut tail_start = 0usize;
14448 if race_eligible && crate::gpu::graph_race_use_graph(graph_trusted) {
14449 let t_graph = std::time::Instant::now();
14450 let mut lg = Vec::new();
14451 let mut gl = 0usize;
14452 let built = self.try_token_graph_wgpu(hidden, position, &mut lg, &mut gl);
14453 let declined = built.is_none();
14454 let built = match built {
14455 Some(Ok(hh)) => Some(hh),
14456 Some(Err(())) => {
14457 self.clear_sequence_state();
14461 self.graph_failed
14462 .store(true, std::sync::atomic::Ordering::Relaxed);
14463 self.cancel
14464 .store(true, std::sync::atomic::Ordering::Relaxed);
14465 tracing::error!("token graph failed after admission; sequence state cleared");
14466 return vec![0.0; self.hidden_size];
14467 }
14468 None => None,
14469 };
14470 if declined && !self.o1_active() && self.attn_softcap == 0.0 {
14475 self.mark_graph_refused();
14476 }
14477 graph_note(built.is_some(), gl, self.num_layers);
14478 if let Some(hh) = built {
14479 let dur = t_graph.elapsed();
14480 if std::env::var("CMF_GRAPH_PROF").is_ok() {
14481 eprintln!("graph-call: {:.2} ms total", dur.as_secs_f64() * 1000.0);
14482 }
14483 if gl > 0 && gl < self.num_layers {
14484 h = hh;
14490 tail_start = gl;
14491 } else if graph_trusted || !crate::gpu::graph_race_first_token_hopeless(dur) {
14492 if !graph_trusted {
14493 crate::gpu::graph_race_record(true, dur);
14494 }
14495 if !lg.is_empty() {
14496 lg.resize(self.vocab_size, 0.0);
14499 if let Some(c) = self.final_softcap {
14500 for l in lg.iter_mut() {
14501 *l = c * (*l / c).tanh();
14502 }
14503 }
14504 self.graph_logits = Some(lg);
14505 }
14506 return hh;
14507 }
14508 }
14514 }
14515 let span = from > 0 || upto.is_some();
14539 if span && graph_on && task_mask.is_none() && graph_trusted {
14540 let upto_excl = upto.map_or(self.num_layers, |u| u + 1);
14541 let mut lg = Vec::new();
14542 let mut gl = 0usize;
14543 let span_res =
14544 self.try_token_graph_wgpu_span(hidden, position, &mut lg, from, upto_excl, &mut gl);
14545 let span_res = match span_res {
14546 Some(Ok(hh)) => Some(hh),
14547 Some(Err(())) => {
14548 self.clear_sequence_state();
14549 self.graph_failed
14550 .store(true, std::sync::atomic::Ordering::Relaxed);
14551 self.cancel
14552 .store(true, std::sync::atomic::Ordering::Relaxed);
14553 tracing::error!(
14554 "span token graph failed after admission; sequence state cleared"
14555 );
14556 return vec![0.0; self.hidden_size];
14557 }
14558 None => None,
14559 };
14560 graph_note(span_res.is_some(), gl, upto_excl - from);
14561 if std::env::var("CMF_GPU_DEBUG").is_ok() {
14562 static SEEN: std::sync::atomic::AtomicU32 = std::sync::atomic::AtomicU32::new(0);
14566 if SEEN.fetch_add(1, std::sync::atomic::Ordering::Relaxed) < 4 {
14567 eprintln!(
14568 "span graph: covered {gl} of {} layers [{from}..{upto_excl}) res={}",
14569 upto_excl - from,
14570 span_res.is_some()
14571 );
14572 }
14573 }
14574 if let Some(hh) = span_res {
14575 if gl == upto_excl - from {
14576 if !lg.is_empty() {
14577 lg.resize(self.vocab_size, 0.0);
14578 if let Some(c) = self.final_softcap {
14579 for l in lg.iter_mut() {
14580 *l = c * (*l / c).tanh();
14581 }
14582 }
14583 self.graph_logits = Some(lg);
14584 }
14585 crate::gpu::set_layer(-1);
14586 return hh;
14587 }
14588 h = hh;
14590 tail_start = from + gl;
14591 }
14592 }
14593 let t_race_cpu = (race_eligible && !graph_trusted).then(std::time::Instant::now);
14598
14599 let host_tail = tail_start > from;
14609 let _host_tail = (host_tail && !self.mimo_moe.is_on()).then(crate::gpu::enter_cpu_scope);
14610 let automatic_gpu_prefix = self.automatic_gpu_prefix();
14611
14612 let _prof_layers = crate::cpuprof::time(crate::cpuprof::Slot::Layers);
14613 #[cfg(target_os = "macos")]
14614 let mut gpu_skip_until = 0usize;
14615 for li in tail_start.max(from)..self.num_layers {
14616 let _capacity_tail = automatic_gpu_prefix
14617 .filter(|&prefix| li >= prefix && !self.mimo_moe.is_dynamic(li, host_tail))
14618 .map(|_| crate::gpu::enter_cpu_scope());
14619 crate::gpu::set_layer(li as i64); if let Some(u) = upto {
14621 if li > u {
14622 break;
14623 }
14624 }
14625 if let Some(mask) = task_mask {
14626 if !mask.layer_alive(li) {
14627 continue; }
14629 }
14630 #[cfg(target_os = "macos")]
14634 {
14635 if li < gpu_skip_until {
14636 continue;
14637 }
14638 if task_mask.is_none() {
14639 let end = self.q1_graph_gpu(li, upto, position, &mut h);
14640 if self
14641 .graph_failed
14642 .load(std::sync::atomic::Ordering::Relaxed)
14643 {
14644 return vec![0.0; self.hidden_size];
14648 }
14649 if end > li {
14650 gpu_skip_until = end;
14651 if self.is_loop_end(end - 1) && end < self.num_layers {
14654 h = inference::rms_norm(
14655 &h,
14656 &self.weights.final_norm,
14657 self.rms_eps,
14658 self.norm_style,
14659 );
14660 }
14661 continue;
14662 }
14663 }
14664 }
14665
14666 if task_mask.is_none() {
14667 match self.mimo_graph_layer_rows(li, &mut h, &[position]) {
14668 crate::gpu::BatchGraphOutcome::Completed => continue,
14669 crate::gpu::BatchGraphOutcome::Failed => return vec![0.0; self.hidden_size],
14670 crate::gpu::BatchGraphOutcome::Declined => {},
14671 }
14672 }
14673 #[cfg(feature = "gpu")]
14674 self.pull_lagging_host_kv(li, li + 1, position);
14675 let lw = &self.weights.layers[self.phys_layer(li)];
14676 if let Ok(tp) = std::env::var("CMF_TRACE_POS") {
14677 if tp.parse::<usize>().ok() == Some(position) {
14678 let n: f32 = h.iter().map(|x| x * x).sum::<f32>().sqrt();
14679 eprintln!(
14680 "TRACE pos {position} layer {li}: |h| = {n:.6} h0 {:.6} h1 {:.6}",
14681 h[0], h[1]
14682 );
14683 }
14684 }
14685 let prof = crate::cpuprof::time(crate::cpuprof::Slot::Norms);
14688 inference::rms_norm_into(
14689 &h,
14690 &lw.input_norm,
14691 self.rms_eps,
14692 self.norm_style,
14693 &mut self.ws.n1,
14694 );
14695 drop(prof);
14696
14697 let attn_out = match &lw.attn {
14698 AttnKind::Mla(w) => {
14699 let inv_freq_l = self.layer_inv_freq(li);
14700 let rs = self.layer_rope_scale(li);
14701 let eps = self.rms_eps;
14702 let pool = self.pool.clone();
14703 mla_attention(
14704 w,
14705 &self.ws.n1,
14706 &mut self.kv_cache.layers[li],
14707 position,
14708 &inv_freq_l,
14709 rs,
14710 eps,
14711 pool.as_deref(),
14712 )
14713 }
14714 AttnKind::Linear(w) => {
14715 let cfg = self.vmf_cfg.expect("linear layer without vmf_cfg");
14716 vmf_phase_forward(
14717 &self.ws.n1,
14718 w,
14719 &cfg,
14720 &mut self.kv_cache.layers[li].linear_state,
14721 self.pool.as_deref(),
14722 )
14723 }
14724 AttnKind::Kda(w) => {
14725 let cfg = self.kda_cfg.expect("kda layer without kda_cfg");
14726 crate::linear_core::kda_forward(
14727 &self.ws.n1,
14728 w,
14729 &cfg,
14730 &mut self.kv_cache.layers[li].linear_state,
14731 self.pool.as_deref(),
14732 )
14733 }
14734 AttnKind::LinearGdn(w) => {
14735 let cfg = self.gdn_cfg.expect("gdn layer without gdn_cfg");
14736 gdn_forward(
14737 &self.ws.n1,
14738 w,
14739 &cfg,
14740 &mut self.kv_cache.layers[li].linear_state,
14741 self.pool.as_deref(),
14742 )
14743 }
14744 AttnKind::ShortConv(w) => {
14745 let cfg = self
14746 .short_conv_cfg
14747 .expect("short-conv layer without short_conv_cfg");
14748 short_conv_forward(
14749 &self.ws.n1,
14750 w,
14751 &cfg,
14752 &mut self.kv_cache.layers[li].linear_state,
14753 self.pool.as_deref(),
14754 )
14755 }
14756 AttnKind::Bounded(w) => {
14757 let rope = self
14760 .bounded_rope
14761 .clone()
14762 .expect("bounded layer without an installed rotation table");
14763 let cfg = crate::bounded::BoundedAttnCfg {
14764 num_heads: self.num_heads,
14765 num_kv_heads: self.num_kv_heads,
14766 head_dim: self.head_dim,
14767 hidden_size: hs,
14768 scale: self.attn_scale,
14769 rope: &rope,
14770 pool: pool.as_deref(),
14771 };
14772 crate::bounded::bounded_attention(
14773 &self.ws.n1,
14774 w,
14775 &mut self.kv_cache.layers[li],
14776 &cfg,
14777 )
14778 }
14779 AttnKind::Full {
14780 wq,
14781 wk,
14782 wv,
14783 wo,
14784 q_norm,
14785 k_norm,
14786 output_gate,
14787 softplus_gate,
14788 bias,
14789 } if self.kv_cache.layers[li].o1_sealed() => {
14790 let inv_freq_l = self.layer_inv_freq(li);
14793 let (nkv_l, hd_l, rd_l) = self.layer_geom(li);
14794 let cfg = QwenAttnCfg {
14795 num_heads: self.layer_num_heads(li),
14796 num_kv_heads: nkv_l,
14797 head_dim: hd_l,
14798 hidden_size: hs,
14799 position,
14800 inv_freq: &inv_freq_l,
14801 rotary_dim: rd_l,
14802 scale: self.attn_scale,
14803 softcap: self.attn_softcap,
14804 window: None,
14805 v_norm: self.attn_v_norm,
14806 qk_norm_after_rope: self.qk_norm_after_rope,
14807 q_norm: q_norm.as_deref(),
14808 k_norm: k_norm.as_deref(),
14809 output_gate: *output_gate,
14810 softplus_gate: softplus_gate
14811 .as_ref()
14812 .map(|(gate, per_head)| (gate, *per_head)),
14813 rope_scale: self.layer_rope_scale(li),
14814 bias: bias
14815 .as_ref()
14816 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
14817 rms_eps: eps,
14818 norm_style: self.norm_style,
14819 pool: pool.as_deref(),
14820 v_head_dim: self.layer_v_dim(li),
14821 };
14822 attention::qwen_attention_nystrom(
14823 &self.ws.n1,
14824 wq,
14825 wk,
14826 wv,
14827 wo,
14828 &mut self.kv_cache.layers[li],
14829 &cfg,
14830 )
14831 }
14832 AttnKind::Full {
14833 wq,
14834 wk,
14835 wv,
14836 wo,
14837 q_norm,
14838 k_norm,
14839 output_gate,
14840 softplus_gate,
14841 bias,
14842 } => 'attn: {
14843 let dropin_reason =
14848 graph_on.then(|| self.graph_attn_decline_reason()).flatten();
14849 if let Some(reason) = dropin_reason {
14850 self.note_graph_decline("wgpu attn dropin", reason);
14851 }
14852 if graph_on
14853 && dropin_reason.is_none()
14854 && !*output_gate
14855 && softplus_gate.is_none()
14856 && self.attention_heads_per_layer.is_none()
14857 && bias.is_none()
14858 && task_mask.is_none()
14859 {
14860 let inv_freq_l = self.layer_inv_freq(li);
14861 let (nkv_l, hd_l, rd_l) = self.layer_geom(li);
14862 let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
14863 if let (Some((gm, qi)), Some((_, ki)), Some((_, vi)), Some((_, oi))) = (
14864 wq.mapped_q1(),
14865 wk.mapped_q1(),
14866 wv.mapped_q1(),
14867 wo.mapped_q1(),
14868 ) {
14869 let gm = gm.clone();
14870 let mut out = vec![0f32; hs];
14871 let cache = &self.kv_cache.layers[li];
14872 if crate::gpu::attn_dropin(
14873 &gm,
14874 self.graph_kv_id,
14875 li,
14876 &self.ws.n1,
14877 qi,
14878 ki,
14879 vi,
14880 oi,
14881 q_norm.as_deref(),
14882 k_norm.as_deref(),
14883 self.qk_norm_after_rope,
14884 &inv_freq_l,
14885 nh,
14886 nkv_l,
14887 hd_l,
14888 rd_l,
14889 hs,
14890 position,
14891 self.kv_cache.max_seq_len,
14892 gemma,
14893 eps as f32,
14894 cache.k_heads(),
14895 cache.v_heads(),
14896 &mut out,
14897 ) {
14898 break 'attn out;
14899 }
14900 }
14901 }
14902 let masked = task_mask
14903 .map(|m| m.head_flags(li, self.num_heads).iter().any(|&a| !a))
14904 .unwrap_or(false);
14905 let f32_view = (wq.as_f32(), wk.as_f32(), wv.as_f32(), wo.as_f32());
14906 let plain = self.layer_attn_plain(li);
14909 match (masked, f32_view) {
14910 (true, (Some(q), Some(k), Some(v), Some(o))) if plain => {
14913 let active_heads = task_mask.unwrap().head_flags(li, self.num_heads);
14914 attention::multi_head_attention(
14915 &self.ws.n1,
14916 q,
14917 k,
14918 v,
14919 o,
14920 &mut self.kv_cache.layers[li],
14921 self.num_heads,
14922 self.num_kv_heads,
14923 self.head_dim,
14924 self.hidden_size,
14925 position,
14926 &active_heads,
14927 &self.inv_freq,
14928 )
14929 }
14930 (masked, _) => {
14931 if masked {
14932 tracing::warn!(
14933 "layer {li}: head mask on quantized weights or on a \
14934 window/sink/per-layer-geometry layer not supported \
14935 yet — executing dense"
14936 );
14937 }
14938 let inv_freq_l = self.layer_inv_freq(li);
14939 let (nkv_l, hd_l, rd_l) = self.layer_geom(li);
14940 let cfg = QwenAttnCfg {
14941 num_heads: self.layer_num_heads(li),
14942 num_kv_heads: nkv_l,
14943 head_dim: hd_l,
14944 hidden_size: hs,
14945 position,
14946 inv_freq: &inv_freq_l,
14947 rotary_dim: rd_l,
14948 scale: self.attn_scale,
14949 softcap: self.attn_softcap,
14950 window: self.layer_window(li),
14951 v_norm: self.attn_v_norm,
14952 qk_norm_after_rope: self.qk_norm_after_rope,
14953 q_norm: q_norm.as_deref(),
14954 k_norm: k_norm.as_deref(),
14955 output_gate: *output_gate,
14956 softplus_gate: softplus_gate
14957 .as_ref()
14958 .map(|(gate, per_head)| (gate, *per_head)),
14959 rope_scale: self.layer_rope_scale(li),
14960 bias: bias
14961 .as_ref()
14962 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
14963 rms_eps: eps,
14964 norm_style: self.norm_style,
14965 pool: pool.as_deref(),
14966 v_head_dim: self.layer_v_dim(li),
14967 };
14968 attention::qwen_attention(
14969 &self.ws.n1,
14970 wq,
14971 wk,
14972 wv,
14973 wo,
14974 &mut self.kv_cache.layers[li],
14975 &cfg,
14976 )
14977 }
14978 }
14979 }
14980 };
14981 let attn_out = match &self.weights.layers[self.phys_layer(li)].attn_out_norm {
14984 Some(w) => inference::rms_norm(&attn_out, w, self.rms_eps, self.norm_style),
14985 None => attn_out,
14986 };
14987 let lw = &self.weights.layers[self.phys_layer(li)];
14988 let prof = crate::cpuprof::time(crate::cpuprof::Slot::Norms);
14989 inference::add_rmsnorm_fused_into(
14990 &mut h,
14991 &attn_out,
14992 &lw.post_norm,
14993 self.rms_eps,
14994 self.norm_style,
14995 &mut self.ws.p1,
14996 );
14997 drop(prof);
14998 let mut attn_out = attn_out;
14999 attention::recycle_buf(&mut attn_out);
15000 let post_normed = &self.ws.p1;
15001
15002 let ffn_masked = task_mask
15003 .map(|m| m.ffn_active_count(li) < self.intermediate_size)
15004 .unwrap_or(false);
15005 let ffn_out = match (ffn_masked, &lw.ffn) {
15017 (_, FfnKind::Dense(d)) if !d.segs.is_empty() => {
15021 let row = task_mask
15022 .and_then(|tm| tm.ffn_masks.get(li))
15023 .map(|v| v.as_slice());
15024 tube_ffn(d, post_normed, 1, self.pool.as_deref(), row)
15025 }
15026 (true, FfnKind::Dense(d)) => {
15027 let tm = task_mask.unwrap();
15028 let alive = tm.ffn_active_count(li);
15029 let deep = alive * 2 <= self.intermediate_size;
15030 if deep && d.down_proj.sparse_col_ok() && !d.gate_proj.has_prism_contract() {
15031 let active = tm.ffn_active_indices(li);
15032 sparse_ffn_quant(
15033 d,
15034 post_normed,
15035 &active,
15036 self.hidden_size,
15037 self.pool.as_deref(),
15038 )
15039 } else if deep
15040 && let (Some(g), Some(u), Some(dn)) = (
15041 d.gate_proj.as_f32(),
15042 d.up_proj.as_f32(),
15043 d.down_proj.as_f32(),
15044 )
15045 {
15046 let active = tm.ffn_active_indices(li);
15047 inference::sparse_ffn_forward(
15048 post_normed,
15049 g,
15050 u,
15051 dn,
15052 self.hidden_size,
15053 self.intermediate_size,
15054 &active,
15055 self.pool.as_deref(),
15056 )
15057 } else {
15058 let row = tm.ffn_masks.get(li).map(|v| v.as_slice());
15059 dense_ffn_batch(d, post_normed, 1, self.pool.as_deref(), row)
15060 }
15061 }
15062 (true, FfnKind::Moe(m)) => {
15063 let allowed = task_mask.and_then(|tm| tm.expert_flags(li, m.experts.len()));
15067 ffn_forward(
15068 &lw.ffn,
15069 post_normed,
15070 self.pool.as_deref(),
15071 allowed.as_deref(),
15072 )
15073 }
15074 (true, FfnKind::DenseMoe(dm)) => dense_moe_ffn(
15075 dm,
15076 post_normed,
15077 &h,
15078 self.rms_eps,
15079 self.norm_style,
15080 self.pool.as_deref(),
15081 ),
15082 (false, _) => match &lw.ffn {
15083 FfnKind::DenseMoe(dm) => dense_moe_ffn(
15084 dm,
15085 post_normed,
15086 &h,
15087 self.rms_eps,
15088 self.norm_style,
15089 self.pool.as_deref(),
15090 ),
15091 FfnKind::Moe(m)
15092 if task_mask.is_none() && self.mimo_moe.is_dynamic(li, host_tail) =>
15093 {
15094 moe_ffn_banked(&mut self.mimo_moe, li, m, post_normed, self.pool.as_deref())
15095 }
15096 _ => {
15097 let allowed = match (&lw.ffn, task_mask) {
15098 (FfnKind::Moe(m), Some(tm)) => tm.expert_flags(li, m.experts.len()),
15099 _ => None,
15100 };
15101 ffn_forward(
15102 &lw.ffn,
15103 post_normed,
15104 self.pool.as_deref(),
15105 allowed.as_deref(),
15106 )
15107 }
15108 },
15109 };
15110 let ffn_out = match &self.weights.layers[self.phys_layer(li)].ffn_out_norm {
15111 Some(w) => inference::rms_norm(&ffn_out, w, self.rms_eps, self.norm_style),
15112 None => ffn_out,
15113 };
15114 for (i, &f) in ffn_out.iter().enumerate() {
15115 h[i] += f;
15116 }
15117 let mut ffn_out = ffn_out;
15118 attention::recycle_buf(&mut ffn_out);
15119
15120 if let Some(sc) = self.weights.layers[self.phys_layer(li)].layer_scale {
15122 for v in h.iter_mut() {
15123 *v *= sc;
15124 }
15125 }
15126 if self.layer_dump.is_some() {
15128 self.dump_layer_row(position, li, &h);
15129 }
15130
15131 if self.is_loop_end(li) && li + 1 < self.num_layers {
15134 h = inference::rms_norm(
15135 &h,
15136 &self.weights.final_norm,
15137 self.rms_eps,
15138 self.norm_style,
15139 );
15140 }
15141
15142 if self.dyn_phi_layer == Some(li) {
15146 self.update_dyn_phi(&h);
15147 }
15148 }
15149 crate::gpu::set_layer(-1); if let Some(t) = t_race_cpu {
15151 crate::gpu::graph_race_record(false, t.elapsed());
15152 }
15153
15154 h
15155 }
15156
15157 fn update_dyn_phi(&mut self, h: &[f32]) {
15160 const A: f32 = 0.2;
15161 if self.dyn_phi_ema.len() != h.len() {
15162 self.dyn_phi_ema = vec![0.0; h.len()];
15163 self.dyn_phi_seen = 0;
15164 }
15165 if self.dyn_phi_seen == 0 {
15166 self.dyn_phi_ema.copy_from_slice(h);
15167 } else {
15168 for (e, &v) in self.dyn_phi_ema.iter_mut().zip(h) {
15169 *e = (1.0 - A) * *e + A * v;
15170 }
15171 }
15172 self.dyn_phi_seen += 1;
15173 }
15174
15175 pub fn dyn_phi(&self) -> &[f32] {
15177 &self.dyn_phi_ema
15178 }
15179
15180 pub fn set_dyn_phi_layer(&mut self, layer: Option<usize>) {
15182 self.dyn_phi_layer = layer;
15183 self.dyn_phi_ema.clear();
15184 self.dyn_phi_seen = 0;
15185 }
15186
15187 pub fn dynamic_skills(&self) -> Vec<(usize, String, usize)> {
15189 let Some(model) = &self.model else {
15190 return Vec::new();
15191 };
15192 model
15193 .header
15194 .skills
15195 .iter()
15196 .enumerate()
15197 .filter_map(|(i, sk)| {
15198 if sk.is_v2() {
15202 return None;
15203 }
15204 let ok = matches!(self.dyn_skill_layers.get(i), Some(Some(_)));
15205 let sel = sk.selection.as_ref()?;
15206 (ok).then(|| (i, sk.id.clone(), sel.phi_layer))
15207 })
15208 .collect()
15209 }
15210
15211 pub fn active_skill(&self) -> Option<usize> {
15213 self.dyn_active
15214 }
15215
15216 pub fn enable_dynamic_routing(&mut self) -> usize {
15221 use crate::swarm::{DynRouter, RoutableSkill};
15222 let Some(model) = self.model.clone() else {
15223 return 0;
15224 };
15225 if let Some(r) = &model.header.router {
15232 tracing::warn!(
15233 "dynamic routing disabled: this file declares router policy '{}' with \
15234 granularity \"{}\" — the request-level decision applies instead",
15235 r.policy,
15236 r.granularity
15237 );
15238 return 0;
15239 }
15240 if model.required_features & cortiq_core::format::features::SKILLS_V2 != 0
15246 || model.header.skills.iter().any(|s| s.is_v2())
15247 {
15248 tracing::warn!(
15249 "dynamic routing disabled: this file carries format-v2 skill records \
15250 (SKILLS_V2) — they route per request through a router policy only"
15251 );
15252 return 0;
15253 }
15254 if self.dyn_blend_loaded {
15257 tracing::warn!("dynamic routing unavailable on a blend-loaded pipeline");
15258 return 0;
15259 }
15260 if let Some(a) = self.dyn_active {
15264 if !matches!(self.dyn_skill_layers.get(a), Some(Some(_))) {
15265 tracing::warn!("loaded skill is not FFN-eligible — dynamic routing unavailable");
15266 return 0;
15267 }
15268 }
15269 let hidden = self.hidden_size;
15270 let mut skills = Vec::new();
15271 for (idx, id, _phi) in self.dynamic_skills() {
15272 if let Some(sel) = model.header.skills[idx].selection.as_ref() {
15273 if let Some(rs) = RoutableSkill::from_descriptor(idx, id, sel, hidden) {
15274 skills.push(rs);
15275 }
15276 }
15277 }
15278 if skills.is_empty() {
15279 return 0;
15280 }
15281 let phi = skills[0].phi_layer;
15283 if skills.iter().any(|s| s.phi_layer != phi) {
15284 tracing::warn!("routable skills disagree on phi_layer; using {phi}");
15285 }
15286 let n = skills.len();
15287 self.set_dyn_phi_layer(Some(phi));
15288 self.dyn_router = Some(DynRouter::new(skills));
15289 n
15290 }
15291
15292 pub fn route_switches(&self) -> Vec<(usize, Option<String>, Option<String>)> {
15294 self.dyn_router
15295 .as_ref()
15296 .map(|r| r.switches.clone())
15297 .unwrap_or_default()
15298 }
15299
15300 fn lm_head_forward(&self, hidden: &[f32]) -> Vec<f32> {
15303 let _mimo_q8 = self.mimo_moe.is_on()
15304 .then(crate::qtensor::enter_full_gpu_q8_scope);
15305 let rows = self.weights.lm_head.rows();
15306 let mut logits = attention::take_buf(rows.min(self.vocab_size));
15307 let served = self.mimo_moe.is_on() && crate::gpu::mimo_q8_short_enabled()
15311 && rows == self.vocab_size && !self.weights.lm_head.has_prism_contract()
15312 && self.weights.lm_head.graph_weight().is_some_and(|(model, idx, kind, _)| {
15313 kind == 7 && crate::gpu::q82_short_rows(model, idx, hidden, 1,
15314 rows, self.hidden_size, &mut logits)
15315 });
15316 if !served {
15317 self.weights.lm_head.matvec(hidden, &mut logits, self.pool.as_deref());
15318 }
15319 logits.resize(self.vocab_size, 0.0);
15320 if let Some(m) = self.logit_multiplier {
15321 for l in logits.iter_mut() {
15322 *l *= m;
15323 }
15324 }
15325 if let Some(c) = self.final_softcap {
15326 for l in logits.iter_mut() {
15327 *l = c * (*l / c).tanh();
15328 }
15329 }
15330 if let Some(cm) = self.head_clusters.as_ref() {
15331 self.hierarchical_head_logprobs(hidden, cm, &mut logits);
15332 }
15333 logits
15334 }
15335
15336 fn hierarchical_head_logprobs(&self, hidden: &[f32], cm: &[f32], logits: &mut [f32]) {
15339 let h = hidden.len();
15340 let ncl = cm.len() / h.max(1);
15341 if ncl == 0 || logits.len() % ncl != 0 {
15342 return;
15343 }
15344 let cs = logits.len() / ncl;
15345 let mut lc = vec![0.0f32; ncl];
15347 for c in 0..ncl {
15348 let row = &cm[c * h..(c + 1) * h];
15349 let mut s = 0.0f32;
15350 for j in 0..h {
15351 s += row[j] * hidden[j];
15352 }
15353 lc[c] = s;
15354 }
15355 let mx = lc.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
15356 let lse: f32 = mx + lc.iter().map(|v| (v - mx).exp()).sum::<f32>().ln();
15357 for c in 0..ncl {
15358 let blk = &mut logits[c * cs..(c + 1) * cs];
15359 let bm = blk.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
15360 let bl: f32 = bm + blk.iter().map(|v| (v - bm).exp()).sum::<f32>().ln();
15361 let add = lc[c] - lse - bl;
15362 for v in blk.iter_mut() {
15363 *v += add;
15364 }
15365 }
15366 }
15367
15368 pub fn prefill_next_logits(&mut self, ids: &[u32], task_mask: Option<&TaskMask>) -> Vec<f32> {
15373 #[cfg(target_os = "macos")]
15374 crate::gpu_metal::set_io_namespace(self.graph_kv_id);
15375 self.clear_sequence_state();
15376 crate::gpu::graph_race_begin_generation();
15380 if task_mask.is_none() {
15381 self.o1_begin();
15382 }
15383 let mut hidden = vec![0.0f32; self.hidden_size];
15384 for (pos, &id) in ids.iter().enumerate() {
15385 let emb = self.embed_single(id);
15386 hidden = self.forward_layers(&emb, pos, task_mask);
15387 }
15388 if let Err(err) = self.o1_seal_checked() {
15389 self.o1_fail(err);
15390 }
15391 if let Some(logits) = self.graph_logits.take() {
15394 return logits;
15395 }
15396 inference::rms_norm_into(
15397 &hidden,
15398 &self.weights.final_norm,
15399 self.rms_eps,
15400 self.norm_style,
15401 &mut self.ws.n1,
15402 );
15403 self.lm_head_forward(&self.ws.n1)
15404 }
15405}
15406
15407pub fn create_test_pipeline(
15409 hidden_size: usize,
15410 intermediate_size: usize,
15411 num_heads: usize,
15412 num_kv_heads: usize,
15413 head_dim: usize,
15414 num_layers: usize,
15415 vocab_size: usize,
15416) -> Pipeline {
15417 let synth = |n: usize, salt: usize| -> Vec<f32> {
15420 (0..n)
15421 .map(|i| (((i * 31 + salt * 17 + 7) % 97) as f32 / 97.0 - 0.5) * 0.2)
15422 .collect()
15423 };
15424 let qt = |rows: usize, cols: usize, salt: usize| -> QTensor {
15425 QTensor::from_f32(synth(rows * cols, salt), rows, cols)
15426 };
15427 let layer_weights: Vec<LayerWeights> = (0..num_layers)
15428 .map(|li| LayerWeights {
15429 input_norm: vec![1.0; hidden_size],
15430 post_norm: vec![1.0; hidden_size],
15431 attn_out_norm: None,
15432 ffn_out_norm: None,
15433 layer_scale: None,
15434 ffn: FfnKind::Dense(DenseFfn {
15435 gate_proj: qt(intermediate_size, hidden_size, li * 10 + 5),
15436 up_proj: qt(intermediate_size, hidden_size, li * 10 + 6),
15437 down_proj: qt(hidden_size, intermediate_size, li * 10 + 7),
15438 act: Act::Silu,
15439 down_t: None,
15440 segs: Vec::new(),
15441 }),
15442 attn: AttnKind::Full {
15443 bias: None,
15444 wq: qt(num_heads * head_dim, hidden_size, li * 10 + 1),
15445 wk: qt(num_kv_heads * head_dim, hidden_size, li * 10 + 2),
15446 wv: qt(num_kv_heads * head_dim, hidden_size, li * 10 + 3),
15447 wo: qt(hidden_size, num_heads * head_dim, li * 10 + 4),
15448 q_norm: None,
15449 k_norm: None,
15450 output_gate: false,
15451 softplus_gate: None,
15452 },
15453 })
15454 .collect();
15455
15456 Pipeline::new(
15457 Tokenizer::byte_level(),
15458 PipelineWeights {
15459 embed_tokens: qt(vocab_size, hidden_size, 100),
15460 layers: layer_weights,
15461 lm_head: qt(vocab_size, hidden_size, 200),
15462 final_norm: vec![1.0; hidden_size],
15463 },
15464 hidden_size,
15465 intermediate_size,
15466 num_heads,
15467 num_kv_heads,
15468 head_dim,
15469 num_layers,
15470 num_layers, false, vocab_size,
15473 1e-6,
15474 10_000.0,
15475 NormStyle::Qwen,
15476 4096,
15477 SamplerConfig {
15478 seed: Some(42),
15479 ..Default::default()
15480 },
15481 )
15482}
15483
15484#[inline]
15489fn mask_bit(row: &[u8], j: usize) -> bool {
15490 (row.get(j >> 3).copied().unwrap_or(0) >> (j & 7)) & 1 != 0
15491}
15492
15493fn mask_gain() -> f32 {
15504 static G: std::sync::OnceLock<f32> = std::sync::OnceLock::new();
15505 *G.get_or_init(|| {
15506 std::env::var("CMF_FFN_MASK_GAIN")
15507 .ok()
15508 .and_then(|v| v.parse().ok())
15509 .unwrap_or(1.0)
15510 })
15511}
15512
15513fn zero_masked_cols(g: &mut [f32], rows: usize, inter: usize, row: &[u8]) {
15514 let fill = meanfill().and_then(|(i, v)| {
15517 let li = crate::gpu::cur_layer();
15518 (*i == inter && li >= 0).then(|| &v[li as usize * inter..(li as usize + 1) * inter])
15519 });
15520 for r in 0..rows {
15521 let base = r * inter;
15522 for (bi, &byte) in row.iter().enumerate() {
15523 if byte == 0xFF {
15524 continue;
15525 }
15526 let j0 = bi * 8;
15527 for bit in 0..8 {
15528 let j = j0 + bit;
15529 if j < inter && byte & (1 << bit) == 0 {
15530 g[base + j] = fill.map_or(0.0, |f| f[j]);
15531 }
15532 }
15533 }
15534 }
15535 let gain = mask_gain();
15536 if gain != 1.0 {
15537 for v in g[..rows * inter].iter_mut() {
15538 *v *= gain;
15539 }
15540 }
15541}
15542
15543#[inline]
15545fn tube_bit(row: Option<&[u8]>, i: usize) -> bool {
15546 row.is_none_or(|r| mask_bit(r, i))
15547}
15548
15549fn all_bits_on(row: &[u8], n: usize) -> bool {
15552 (0..n).all(|i| mask_bit(row, i))
15553}
15554
15555fn tube_topk() -> usize {
15563 static K: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
15564 *K.get_or_init(|| {
15565 std::env::var("CMF_TUBE_TOPK")
15566 .ok()
15567 .and_then(|v| v.parse().ok())
15568 .unwrap_or(0)
15569 })
15570}
15571
15572fn tube_score_oracle() -> bool {
15573 static O: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
15574 *O.get_or_init(|| std::env::var("CMF_TUBE_SCORE").is_ok_and(|v| v == "oracle"))
15575}
15576
15577fn tube_ffn_routed(
15584 d: &DenseFfn,
15585 xs: &[f32],
15586 b: usize,
15587 pool: Option<&Pool>,
15588 mask_row: Option<&[u8]>,
15589 k: usize,
15590) -> Vec<f32> {
15591 let hidden = d.down_proj.rows();
15592 let core = d.gate_proj.rows();
15593 let core_full = mask_row.is_none_or(|r| all_bits_on(r, core));
15594 let mut out = match (b, core_full, mask_row) {
15595 (1, true, _) => dense_ffn(d, xs, pool),
15596 (1, false, Some(row)) => dense_ffn_masked(d, xs, pool, row),
15597 (_, true, _) => dense_ffn_batch(d, xs, b, pool, None),
15598 (_, false, row) => dense_ffn_batch(d, xs, b, pool, row),
15599 };
15600 let cand: Vec<usize> = (0..d.segs.len())
15601 .filter(|&i| tube_bit(mask_row, d.segs[i].start))
15602 .collect();
15603 if cand.is_empty() {
15604 return out;
15605 }
15606 let oracle = tube_score_oracle();
15610 let mut acts: Vec<Vec<f32>> = Vec::with_capacity(cand.len());
15611 let mut scores = vec![0f32; b * cand.len()];
15612 for (ci, &i) in cand.iter().enumerate() {
15613 let seg = &d.segs[i];
15614 let w = seg.width;
15615 let mut g = vec![0.0f32; b * w];
15616 if b == 1 {
15617 seg.gate.matvec(xs, &mut g, pool);
15618 } else {
15619 seg.gate.matmat(xs, b, &mut g, pool);
15620 }
15621 for v in g.iter_mut() {
15622 *v = Act::Silu.combine(*v, 1.0);
15623 }
15624 if !oracle {
15625 for t in 0..b {
15626 scores[t * cand.len() + ci] =
15627 g[t * w..(t + 1) * w].iter().map(|v| v * v).sum::<f32>();
15628 }
15629 }
15630 if oracle || b > 1 {
15631 let mut u = vec![0.0f32; b * w];
15632 if b == 1 {
15633 seg.up.matvec(xs, &mut u, pool);
15634 } else {
15635 seg.up.matmat(xs, b, &mut u, pool);
15636 }
15637 for (a, &v) in g.iter_mut().zip(u.iter()) {
15638 *a *= v;
15639 }
15640 if oracle {
15641 for t in 0..b {
15642 scores[t * cand.len() + ci] =
15643 g[t * w..(t + 1) * w].iter().map(|v| v * v).sum::<f32>();
15644 }
15645 }
15646 }
15647 acts.push(g);
15648 }
15649 let keep = k.min(cand.len());
15651 let mut scratch: Vec<f32> = Vec::new();
15652 for t in 0..b {
15653 let mut sc: Vec<(f32, usize)> = (0..cand.len())
15654 .map(|ci| (scores[t * cand.len() + ci], ci))
15655 .collect();
15656 sc.sort_unstable_by(|x, y| y.0.total_cmp(&x.0));
15657 let mut alive = vec![false; cand.len()];
15658 for &(_, ci) in sc.iter().take(keep) {
15659 alive[ci] = true;
15660 }
15661 if b > 1 {
15662 for (ci, a) in acts.iter_mut().enumerate() {
15663 if !alive[ci] {
15664 let w = d.segs[cand[ci]].width;
15665 a[t * w..(t + 1) * w].fill(0.0);
15666 }
15667 }
15668 } else {
15669 for (ci, &i) in cand.iter().enumerate() {
15673 if !alive[ci] {
15674 continue;
15675 }
15676 let seg = &d.segs[i];
15677 let w = seg.width;
15678 let g = &mut acts[ci];
15679 if !tube_score_oracle() {
15680 scratch.clear();
15681 scratch.resize(w, 0.0);
15682 seg.up.matvec(xs, &mut scratch, pool);
15683 for (a, &v) in g.iter_mut().zip(scratch.iter()) {
15684 *a *= v;
15685 }
15686 }
15687 let mut acc = vec![0.0f32; hidden];
15688 seg.down.matvec(g, &mut acc, pool);
15689 for (o, a) in out.iter_mut().zip(&acc) {
15690 *o += *a;
15691 }
15692 }
15693 }
15694 }
15695 if b > 1 {
15696 for (ci, &i) in cand.iter().enumerate() {
15697 let seg = &d.segs[i];
15698 let mut acc = vec![0.0f32; b * hidden];
15699 seg.down.matmat(&acts[ci], b, &mut acc, pool);
15700 for (o, a) in out.iter_mut().zip(&acc) {
15701 *o += *a;
15702 }
15703 }
15704 }
15705 out
15706}
15707
15708fn tube_ffn(
15714 d: &DenseFfn,
15715 xs: &[f32],
15716 b: usize,
15717 pool: Option<&Pool>,
15718 mask_row: Option<&[u8]>,
15719) -> Vec<f32> {
15720 if tube_topk() > 0 {
15721 return tube_ffn_routed(d, xs, b, pool, mask_row, tube_topk());
15722 }
15723 let hidden = d.down_proj.rows();
15724 let core = d.gate_proj.rows();
15725 let core_full = mask_row.is_none_or(|r| all_bits_on(r, core));
15726 let mut out = match (b, core_full, mask_row) {
15727 (1, true, _) => dense_ffn(d, xs, pool),
15728 (1, false, Some(row)) => dense_ffn_masked(d, xs, pool, row),
15729 (_, true, _) => dense_ffn_batch(d, xs, b, pool, None),
15730 (_, false, row) => dense_ffn_batch(d, xs, b, pool, row),
15731 };
15732 TUBE_SCRATCH.with(|sc| {
15733 let mut sc = sc.borrow_mut();
15734 let [g, u, acc] = &mut *sc;
15735 for seg in &d.segs {
15736 if !tube_bit(mask_row, seg.start) {
15737 continue;
15738 }
15739 let w = seg.width;
15740 g.resize(b * w, 0.0);
15741 if b == 1
15742 && d.act == Act::Silu
15743 && QTensor::matvec_silu_mul(&seg.gate, &seg.up, xs, g, pool)
15744 {
15745 } else {
15747 u.resize(b * w, 0.0);
15748 if b == 1 {
15749 QTensor::matvec_many([&seg.gate, &seg.up], xs, [g, u], pool);
15750 } else {
15751 seg.gate.matmat(xs, b, g, pool);
15752 seg.up.matmat(xs, b, u, pool);
15753 }
15754 for i in 0..b * w {
15755 g[i] = d.act.combine(g[i], u[i]);
15756 }
15757 }
15758 acc.resize(b * hidden, 0.0);
15759 acc.fill(0.0);
15760 if b == 1 {
15761 seg.down.matvec(g, acc, pool);
15762 } else {
15763 seg.down.matmat(g, b, acc, pool);
15764 }
15765 for (o, a) in out.iter_mut().zip(acc.iter()) {
15766 *o += *a;
15767 }
15768 }
15769 out
15770 })
15771}
15772
15773thread_local! {
15774 static TUBE_SCRATCH: std::cell::RefCell<[Vec<f32>; 3]> =
15778 const { std::cell::RefCell::new([Vec::new(), Vec::new(), Vec::new()]) };
15779}
15780
15781fn dense_ffn_batch(
15782 d: &DenseFfn,
15783 xs: &[f32],
15784 b: usize,
15785 pool: Option<&Pool>,
15786 mask_row: Option<&[u8]>,
15787) -> Vec<f32> {
15788 let inter = d.gate_proj.rows();
15789 let hidden = d.down_proj.rows();
15790 if mask_row.is_none()
15798 && d.act == Act::Silu
15799 && b >= 32
15800 && crate::gpu::enabled_here()
15801 && !crate::gpu::mm_killed()
15802 && refit_dir().is_none()
15807 && !ffn_probe_active()
15812 {
15813 if let (Some((model, w1)), Some((_, w3)), Some((_, w2))) = (
15814 d.gate_proj.mapped_q4t(),
15815 d.up_proj.mapped_q4t(),
15816 d.down_proj.mapped_q4t(),
15817 ) {
15818 let mut out = vec![0.0f32; b * hidden];
15819 if crate::gpu::q4t_ffn(model, w1, w3, w2, xs, b, hidden, inter, &mut out) {
15820 return out;
15821 }
15822 }
15823 if let (Some((model, w1)), Some((_, w3)), Some((_, w2))) = (
15828 d.gate_proj.mapped_q4tp(),
15829 d.up_proj.mapped_q4tp(),
15830 d.down_proj.mapped_q4tp(),
15831 ) {
15832 let mut out = vec![0.0f32; b * hidden];
15833 if crate::gpu::q4tp_ffn(model, w1, w3, w2, xs, b, hidden, inter, &mut out) {
15834 return out;
15835 }
15836 }
15837 }
15838 let mut g = vec![0.0f32; b * inter];
15839 d.gate_proj.matmat(xs, b, &mut g, pool);
15840 let mut u = vec![0.0f32; b * inter];
15841 d.up_proj.matmat(xs, b, &mut u, pool);
15842 if gate_topk() > 0 && d.act == Act::Silu {
15843 for t in 0..b {
15844 let row = &mut g[t * inter..(t + 1) * inter];
15845 for v in row.iter_mut() {
15846 *v = Act::Silu.combine(*v, 1.0);
15847 }
15848 keep_top_k(row, gate_topk());
15849 }
15850 for i in 0..b * inter {
15851 g[i] *= u[i];
15852 }
15853 } else {
15854 for i in 0..b * inter {
15855 g[i] = d.act.combine(g[i], u[i]);
15856 }
15857 }
15858 if let Some(row) = mask_row {
15859 zero_masked_cols(&mut g, b, inter, row);
15860 }
15861 if oracle_topk() > 0 {
15862 for t in 0..b {
15863 keep_top_k(&mut g[t * inter..(t + 1) * inter], oracle_topk());
15864 }
15865 }
15866 let mut out = vec![0.0f32; b * hidden];
15867 d.down_proj.matmat(&g, b, &mut out, pool);
15868 if refit_dir().is_some() {
15869 let li = crate::gpu::cur_layer();
15870 if li >= 0 {
15871 refit_accumulate(li as usize, &g, b, inter, &out, hidden, pool);
15872 }
15873 }
15874 FFN_PROBE.with(|pr| {
15878 if let Some(acc) = pr.borrow_mut().as_mut() {
15879 let li = crate::gpu::cur_layer();
15880 if li < 0 {
15881 return;
15882 }
15883 let Some(row) = acc.get_mut(li as usize) else {
15884 return;
15885 };
15886 let sq = probe_sq();
15887 for t in 0..b {
15888 for (a, &v) in row.iter_mut().zip(&g[t * inter..(t + 1) * inter]) {
15889 *a += if sq {
15890 (v as f64) * (v as f64)
15891 } else {
15892 (v as f64).abs()
15893 };
15894 }
15895 }
15896 }
15897 });
15898 out
15899}
15900
15901fn accumulate_act(m: &MoeFfn, xs: &[f32], b: usize) {
15906 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
15907 static DUMP: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
15908 let on = *ON.get_or_init(|| std::env::var("CMF_RMS_TRACE").is_ok());
15909 let dump = *DUMP.get_or_init(|| std::env::var("CMF_ACT_DUMP").is_ok());
15910 if (!on && !dump) || b == 0 {
15911 return;
15912 }
15913 let hidden = xs.len() / b;
15914 if on {
15915 let mut acc = m.act_sq.borrow_mut();
15916 if acc.len() < hidden {
15917 acc.resize(hidden, 0.0);
15918 }
15919 for t in 0..b {
15920 let row = &xs[t * hidden..(t + 1) * hidden];
15921 for (a, &v) in acc.iter_mut().zip(row) {
15922 *a += (v as f64) * (v as f64);
15923 }
15924 }
15925 }
15926 if dump {
15927 let cap: usize = std::env::var("CMF_ACT_DUMP_ROWS")
15930 .ok()
15931 .and_then(|v| v.parse().ok())
15932 .unwrap_or(4096);
15933 let mut rows = m.act_rows.borrow_mut();
15934 if rows.len() < cap * hidden {
15935 let take = b.min((cap * hidden - rows.len()) / hidden.max(1));
15936 rows.extend_from_slice(&xs[..take * hidden]);
15937 }
15938 }
15939}
15940
15941#[derive(Clone, Copy)]
15944struct SendVecs(*mut Vec<f32>);
15945unsafe impl Send for SendVecs {}
15946unsafe impl Sync for SendVecs {}
15947impl SendVecs {
15948 #[inline]
15949 fn at(self, i: usize) -> *mut Vec<f32> {
15950 unsafe { self.0.add(i) }
15951 }
15952}
15953
15954fn moe_ffn_batch(
15955 m: &MoeFfn,
15956 xs: &[f32],
15957 b: usize,
15958 hidden: usize,
15959 pool: Option<&Pool>,
15960 allowed: Option<&[bool]>,
15961) -> Vec<f32> {
15962 accumulate_act(m, xs, b);
15963 let ne = m.experts.len();
15964 let mut logits = vec![0.0f32; b * ne];
15965 match &m.resonance {
15966 Some(r) => {
15967 let hdim = xs.len() / b.max(1);
15968 for bi in 0..b {
15969 r.scores(
15970 &xs[bi * hdim..(bi + 1) * hdim],
15971 &mut logits[bi * ne..(bi + 1) * ne],
15972 );
15973 }
15974 }
15975 None => m.router.matmat(xs, b, &mut logits, pool),
15976 }
15977
15978 let mut assign: Vec<Vec<(usize, f32)>> = vec![Vec::new(); ne];
15981 {
15982 let mut st = m.stats.borrow_mut();
15983 if st.len() < ne {
15984 st.resize(ne, 0);
15985 }
15986 for bi in 0..b {
15987 let (idx, p, wsum) = moe_route(&logits[bi * ne..(bi + 1) * ne], m, allowed);
15988 for &e in &idx {
15989 st[e] += 1;
15990 assign[e].push((bi, p[e] / wsum));
15991 }
15992 }
15993 }
15994
15995 let mut out = vec![0.0f32; b * hidden];
15996 let cols = m.experts[0].gate_proj.cols();
15997 let run_expert = |d: &DenseFfn, list: &[(usize, f32)], out: &mut [f32]| {
15998 let sb = list.len();
15999 let mut sub = vec![0.0f32; sb * cols];
16000 for (k, &(bi, _)) in list.iter().enumerate() {
16001 sub[k * cols..(k + 1) * cols].copy_from_slice(&xs[bi * cols..(bi + 1) * cols]);
16002 }
16003 let eo = dense_ffn_batch(d, &sub, sb, pool, None);
16004 for (k, &(bi, w)) in list.iter().enumerate() {
16005 for i in 0..hidden {
16006 out[bi * hidden + i] += w * eo[k * hidden + i];
16007 }
16008 }
16009 };
16010 let active: Vec<usize> = (0..ne).filter(|&e| !assign[e].is_empty()).collect();
16016 if pool.is_some() && active.len() >= 8 {
16017 let mut panels: Vec<Vec<f32>> = vec![Vec::new(); active.len()];
16018 {
16019 let panel_ptr = SendVecs(panels.as_mut_ptr());
16020 let experts = &m.experts;
16023 let (active_r, assign_r) = (&active, &assign);
16024 let inherit_cpu = crate::gpu::inherit_cpu_scope();
16025 let run = |start: usize, end: usize| {
16026 let _cpu_scope = inherit_cpu();
16027 for ai in start..end {
16028 let e = active_r[ai];
16029 let list = &assign_r[e];
16030 let sb = list.len();
16031 let mut sub = vec![0.0f32; sb * cols];
16032 for (k, &(bi, _)) in list.iter().enumerate() {
16033 sub[k * cols..(k + 1) * cols]
16034 .copy_from_slice(&xs[bi * cols..(bi + 1) * cols]);
16035 }
16036 unsafe {
16038 *panel_ptr.at(ai) = dense_ffn_batch(&experts[e], &sub, sb, None, None);
16039 }
16040 }
16041 };
16042 match pool {
16043 Some(p) => p.run_rows(active.len(), &run),
16044 None => run(0, active.len()),
16045 }
16046 }
16047 for (ai, &e) in active.iter().enumerate() {
16048 for (k, &(bi, w)) in assign[e].iter().enumerate() {
16049 let eo = &panels[ai][k * hidden..(k + 1) * hidden];
16050 for i in 0..hidden {
16051 out[bi * hidden + i] += w * eo[i];
16052 }
16053 }
16054 }
16055 } else {
16056 for &e in &active {
16057 run_expert(&m.experts[e], &assign[e], &mut out);
16058 }
16059 }
16060 if let Some((se, gate)) = &m.shared {
16061 let all: Vec<(usize, f32)> = if let Some(gate) = gate {
16062 let mut gl = vec![0.0f32; b];
16063 gate.matmat(xs, b, &mut gl, pool);
16064 (0..b)
16065 .map(|bi| (bi, 1.0 / (1.0 + (-gl[bi]).exp())))
16066 .collect()
16067 } else {
16068 (0..b).map(|bi| (bi, 1.0)).collect()
16069 };
16070 run_expert(se, &all, &mut out);
16071 }
16072 out
16073}
16074
16075fn moe_ffn_rows_exact(
16087 m: &MoeFfn,
16088 xs: &[f32],
16089 b: usize,
16090 hidden: usize,
16091 pool: Option<&Pool>,
16092) -> Vec<f32> {
16093 let mut out = vec![0.0f32; b * hidden];
16094 let per_row = |out: &mut [f32]| {
16095 for r in 0..b {
16096 let o = moe_ffn(m, &xs[r * hidden..(r + 1) * hidden], pool, None);
16097 out[r * hidden..(r + 1) * hidden].copy_from_slice(&o);
16098 }
16099 };
16100 let covered = !crate::gpu::enabled_here()
16101 && moe_batch_enabled()
16102 && m.shared.is_none()
16103 && m.resonance.is_none()
16104 && FFN_PROBE.with(|pr| pr.borrow().is_none())
16105 && m.experts.iter().all(|d| d.act == Act::Silu);
16106 if !covered {
16107 per_row(&mut out);
16108 return out;
16109 }
16110 let ne = m.experts.len();
16111 let mut routes: Vec<(Vec<usize>, Vec<f32>)> = Vec::with_capacity(b);
16113 for r in 0..b {
16114 let x = &xs[r * hidden..(r + 1) * hidden];
16115 accumulate_act(m, x, 1);
16116 let mut logits = vec![0.0f32; ne];
16117 m.router.matvec(x, &mut logits, pool);
16118 let (idx, p, wsum) = moe_route(&logits, m, None);
16119 {
16120 let mut st = m.stats.borrow_mut();
16121 if st.len() < ne {
16122 st.resize(ne, 0);
16123 }
16124 for &e in &idx {
16125 st[e] += 1;
16126 }
16127 }
16128 let w: Vec<f32> = idx
16129 .iter()
16130 .map(|&e| p[e] / wsum * m.per_expert_scale.as_ref().map_or(1.0, |v| v[e]))
16131 .collect();
16132 routes.push((idx, w));
16133 }
16134 if routes.iter().any(|(idx, _)| idx.is_empty()) {
16135 per_row(&mut out);
16136 return out;
16137 }
16138 let mut experts: Vec<usize> = Vec::new();
16140 let mut groups: Vec<Vec<usize>> = Vec::new();
16141 for (r, (idx, _)) in routes.iter().enumerate() {
16142 for &e in idx {
16143 match experts.iter().position(|&x| x == e) {
16144 Some(g) => groups[g].push(r),
16145 None => {
16146 experts.push(e);
16147 groups.push(vec![r]);
16148 }
16149 }
16150 }
16151 }
16152 let n_pairs: usize = groups.iter().map(|g| g.len()).sum();
16153 let inter = m.experts[experts[0]].gate_proj.rows();
16154 let pairs: Vec<(&QTensor, &QTensor)> = experts
16155 .iter()
16156 .map(|&e| (&m.experts[e].gate_proj, &m.experts[e].up_proj))
16157 .collect();
16158 let mut gs: Vec<Vec<f32>> = (0..n_pairs).map(|_| vec![0f32; inter]).collect();
16159 if !QTensor::moe_gate_up_rows(&pairs, &groups, xs, &mut gs, pool) {
16160 per_row(&mut out);
16161 return out;
16162 }
16163 let downs: Vec<&QTensor> = experts.iter().map(|&e| &m.experts[e].down_proj).collect();
16164 let lens: Vec<usize> = groups.iter().map(|g| g.len()).collect();
16165 let mut ds: Vec<Vec<f32>> = (0..n_pairs).map(|_| vec![0f32; hidden]).collect();
16166 if !QTensor::moe_down_rows(&downs, &lens, &gs, &mut ds, pool) {
16167 per_row(&mut out);
16168 return out;
16169 }
16170 let mut slot = std::collections::HashMap::with_capacity(n_pairs);
16172 let mut p = 0usize;
16173 for (g, &e) in experts.iter().enumerate() {
16174 for &r in &groups[g] {
16175 slot.insert((r, e), p);
16176 p += 1;
16177 }
16178 }
16179 for (r, (idx, w)) in routes.iter().enumerate() {
16180 let terms: Vec<(&[f32], f32)> = idx
16181 .iter()
16182 .zip(w)
16183 .map(|(&e, &we)| (ds[slot[&(r, e)]].as_slice(), we))
16184 .collect();
16185 let row = &mut out[r * hidden..(r + 1) * hidden];
16186 for (i, dst) in row.iter_mut().enumerate() {
16187 let mut acc = 0f32;
16189 for (d, we) in &terms {
16190 acc += we * d[i];
16191 }
16192 *dst = acc;
16193 }
16194 }
16195 out
16196}
16197
16198thread_local! {
16199 static FFN_SCRATCH: std::cell::RefCell<[Vec<f32>; 4]> =
16203 const { std::cell::RefCell::new([Vec::new(), Vec::new(), Vec::new(), Vec::new()]) };
16204}
16205
16206fn dense_ffn(d: &DenseFfn, x: &[f32], pool: Option<&Pool>) -> Vec<f32> {
16208 if gate_topk() > 0
16211 && let Some(out) = dense_ffn_dynamic(d, x, pool, gate_topk())
16212 {
16213 return out;
16214 }
16215 let prism_body = d.gate_proj.has_prism_contract()
16231 || d.up_proj.has_prism_contract()
16232 || d.down_proj.has_prism_contract();
16233 if !prism_body
16234 && crate::gpu::enabled_here()
16235 && (d.gate_proj.rows() >= crate::gpu::min_rows() || d.gate_proj.is_q1())
16236 {
16237 let arm = if d.gate_proj.is_q1() && crate::gpu::q1_force() {
16238 crate::gpu::ProbeArm::Gpu
16239 } else {
16240 crate::gpu::probe_arm(crate::gpu::OpClass::Ffn)
16241 };
16242 match arm {
16243 crate::gpu::ProbeArm::Gpu => {
16244 let t0 = std::time::Instant::now();
16245 if let Some(out) = dense_ffn_gpu(d, x, pool) {
16246 crate::gpu::probe_record(crate::gpu::OpClass::Ffn, true, t0.elapsed());
16247 return out;
16248 }
16249 crate::gpu::probe_note_decline(crate::gpu::OpClass::Ffn);
16253 }
16254 crate::gpu::ProbeArm::CpuTimed => {
16255 let t0 = std::time::Instant::now();
16256 let out = crate::gpu::cpu_scope(|| dense_ffn_cpu(d, x, pool));
16257 crate::gpu::probe_record(crate::gpu::OpClass::Ffn, false, t0.elapsed());
16258 return out;
16259 }
16260 crate::gpu::ProbeArm::Cpu => {
16261 return crate::gpu::cpu_scope(|| dense_ffn_cpu(d, x, pool));
16262 }
16263 }
16264 }
16265 dense_ffn_cpu(d, x, pool)
16266}
16267
16268fn dense_ffn_cpu(d: &DenseFfn, x: &[f32], pool: Option<&Pool>) -> Vec<f32> {
16270 let inter = d.gate_proj.rows();
16271 FFN_SCRATCH.with(|s| {
16272 let mut s = s.borrow_mut();
16273 let [g, u, ..] = &mut *s;
16274 g.resize(inter, 0.0);
16275 if gate_topk() > 0 {
16278 u.resize(inter, 0.0);
16282 QTensor::matvec_many([&d.gate_proj, &d.up_proj], x, [g, u], pool);
16283 for i in 0..inter {
16284 g[i] = Act::Silu.combine(g[i], 1.0);
16285 }
16286 keep_top_k(g, gate_topk());
16287 for i in 0..inter {
16288 g[i] *= u[i];
16289 }
16290 } else if d.act == Act::Silu && {
16291 let _prof = crate::cpuprof::time(crate::cpuprof::Slot::FfnGateUp);
16292 QTensor::matvec_silu_mul(&d.gate_proj, &d.up_proj, x, g, pool)
16293 } {
16294 } else {
16296 u.resize(inter, 0.0);
16297 let _prof = crate::cpuprof::time(crate::cpuprof::Slot::FfnGateUp);
16299 QTensor::matvec_many([&d.gate_proj, &d.up_proj], x, [g, u], pool);
16300 for i in 0..inter {
16301 g[i] = d.act.combine(g[i], u[i]);
16302 }
16303 }
16304 FFN_PROBE.with(|pr| {
16312 if let Some(acc) = pr.borrow_mut().as_mut() {
16313 let li = crate::gpu::cur_layer();
16314 if li >= 0 {
16315 if let Some(row) = acc.get_mut(li as usize) {
16316 match probe_topk() {
16317 0 if probe_sq() => {
16318 for (a, &v) in row.iter_mut().zip(g.iter()) {
16319 *a += (v as f64) * (v as f64);
16320 }
16321 }
16322 0 if probe_signed() => {
16323 for (a, &v) in row.iter_mut().zip(g.iter()) {
16324 *a += v as f64;
16325 }
16326 }
16327 0 => {
16328 for (a, &v) in row.iter_mut().zip(g.iter()) {
16329 *a += (v as f64).abs();
16330 }
16331 }
16332 k => {
16333 let n = g.len();
16334 let k = k.min(n);
16335 let mut mag: Vec<f32> = g.iter().map(|v| v.abs()).collect();
16336 let (_, kth, _) = mag.select_nth_unstable_by(k - 1, |a, b| {
16337 b.partial_cmp(a).unwrap_or(std::cmp::Ordering::Equal)
16338 });
16339 let thr = *kth;
16340 for (a, &v) in row.iter_mut().zip(g.iter()) {
16341 if v.abs() >= thr {
16342 *a += 1.0;
16343 }
16344 }
16345 }
16346 }
16347 }
16348 }
16349 }
16350 });
16351 if oracle_topk() > 0 {
16352 keep_top_k(g, oracle_topk());
16353 }
16354 {
16355 let li = crate::gpu::cur_layer();
16356 if li >= 0 {
16357 adump_row(li as usize, g);
16358 }
16359 }
16360 let mut out = attention::take_buf(d.down_proj.rows());
16361 let _prof = crate::cpuprof::time(crate::cpuprof::Slot::FfnDown);
16362 d.down_proj.matvec(g, &mut out, pool);
16363 out
16364 })
16365}
16366
16367pub struct RefitAcc {
16380 pub support: Vec<u32>,
16381 pub gss: Vec<f32>,
16382 pub ya: Vec<f32>,
16383 pub hidden: usize,
16384 pub tokens: u64,
16385 pub buf_g: Vec<f32>,
16391 pub buf_o: Vec<f32>,
16392 pub buf_t: usize,
16393}
16394
16395type RefitState = (std::collections::HashMap<usize, RefitAcc>, Vec<f32>);
16399
16400static REFIT: std::sync::OnceLock<Option<(String, std::sync::Mutex<RefitState>)>> =
16401 std::sync::OnceLock::new();
16402
16403fn ffn_probe_active() -> bool {
16406 FFN_PROBE.with(|p| p.borrow().is_some())
16407}
16408
16409fn refit_dir() -> Option<&'static (String, std::sync::Mutex<RefitState>)> {
16410 REFIT
16411 .get_or_init(|| {
16412 std::env::var("CMF_FFN_REFIT").ok().map(|d| {
16413 (
16414 d,
16415 std::sync::Mutex::new((std::collections::HashMap::new(), Vec::new())),
16416 )
16417 })
16418 })
16419 .as_ref()
16420}
16421
16422fn refit_accumulate(
16424 li: usize,
16425 g: &[f32],
16426 b: usize,
16427 inter: usize,
16428 out: &[f32],
16429 hidden: usize,
16430 pool: Option<&Pool>,
16431) {
16432 let Some((dir, map)) = refit_dir() else {
16433 return;
16434 };
16435 static SPAN: std::sync::OnceLock<(usize, usize)> = std::sync::OnceLock::new();
16436 let (from, to) = *SPAN.get_or_init(|| {
16437 let g = |k: &str, d: usize| {
16438 std::env::var(k)
16439 .ok()
16440 .and_then(|v| v.parse().ok())
16441 .unwrap_or(d)
16442 };
16443 (
16444 g("CMF_FFN_REFIT_FROM", 0),
16445 g("CMF_FFN_REFIT_TO", usize::MAX),
16446 )
16447 });
16448 if li < from || li > to {
16449 return;
16450 }
16451 let mut guard = map.lock().unwrap();
16452 let (map, shared) = &mut *guard;
16453 let acc = match map.entry(li) {
16454 std::collections::hash_map::Entry::Occupied(e) => e.into_mut(),
16455 std::collections::hash_map::Entry::Vacant(e) => {
16456 let path = format!("{dir}/support.{li}.u32");
16457 let Ok(bytes) = std::fs::read(&path) else {
16458 eprintln!("refit: no {path} — layer {li} skipped");
16459 return;
16460 };
16461 let n = u32::from_le_bytes(bytes[0..4].try_into().unwrap()) as usize;
16462 let support: Vec<u32> = bytes[4..4 + n * 4]
16463 .chunks_exact(4)
16464 .map(|c| u32::from_le_bytes([c[0], c[1], c[2], c[3]]))
16465 .collect();
16466 eprintln!(
16467 "refit: layer {li} support {n} ({:.0} MB of accumulator)",
16468 (n * n + hidden * n) as f64 * 4.0 / 1e6
16469 );
16470 e.insert(RefitAcc {
16471 gss: vec![0.0; n * n],
16472 ya: vec![0.0; hidden * n],
16473 buf_g: Vec::new(),
16474 buf_o: Vec::new(),
16475 buf_t: 0,
16476 support,
16477 hidden,
16478 tokens: 0,
16479 })
16480 }
16481 };
16482 let ns = acc.support.len();
16483 let cap = refit_batch();
16485 if acc.buf_g.is_empty() {
16486 acc.buf_g = vec![0.0; ns * cap];
16487 acc.buf_o = vec![0.0; hidden * cap];
16488 }
16489 let take = b.min(cap - acc.buf_t);
16490 for t in 0..take {
16491 let col = acc.buf_t + t;
16492 for (j, &n) in acc.support.iter().enumerate() {
16493 acc.buf_g[j * cap + col] = g[t * inter + n as usize];
16494 }
16495 for h in 0..hidden {
16496 acc.buf_o[h * cap + col] = out[t * hidden + h];
16497 }
16498 }
16499 acc.buf_t += take;
16500 acc.tokens += take as u64;
16501 if acc.buf_t < cap {
16502 return;
16503 }
16504 let bt = acc.buf_t;
16505 acc.buf_t = 0;
16506 let RefitAcc {
16516 gss,
16517 ya,
16518 buf_g,
16519 buf_o,
16520 ..
16521 } = acc;
16522 let need = (ns * ns).max(hidden * ns);
16523 if shared.len() < need {
16524 shared.resize(need, 0.0);
16525 }
16526 let scratch = &mut shared[..];
16527 let _ = bt;
16528 if crate::gpu::gemm_nt_f32_transient(buf_g, buf_g, &mut scratch[..ns * ns], ns, cap, ns) {
16529 add_into(gss, &scratch[..ns * ns], pool);
16530 if crate::gpu::gemm_nt_f32_transient(
16531 buf_o,
16532 buf_g,
16533 &mut scratch[..hidden * ns],
16534 hidden,
16535 cap,
16536 ns,
16537 ) {
16538 add_into(ya, &scratch[..hidden * ns], pool);
16539 } else {
16540 accum_outer_t(ya, hidden, ns, cap, buf_o, buf_g, pool);
16541 }
16542 } else {
16543 accum_outer_t(gss, ns, ns, cap, buf_g, buf_g, pool);
16544 accum_outer_t(ya, hidden, ns, cap, buf_o, buf_g, pool);
16545 }
16546 }
16550
16551fn refit_batch() -> usize {
16553 static B: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
16554 *B.get_or_init(|| {
16555 std::env::var("CMF_FFN_REFIT_BATCH")
16556 .ok()
16557 .and_then(|v| v.parse().ok())
16558 .unwrap_or(4096)
16559 })
16560}
16561
16562fn accum_outer_t(
16565 c: &mut [f32],
16566 m: usize,
16567 n: usize,
16568 b: usize,
16569 left: &[f32],
16570 right: &[f32],
16571 pool: Option<&Pool>,
16572) {
16573 let ptr = SendMut(c.as_mut_ptr());
16574 let body = |i: usize| {
16575 let ptr = &ptr;
16576 let row = unsafe { std::slice::from_raw_parts_mut(ptr.0.add(i * n), n) };
16577 for t in 0..b {
16578 let a = left[i * b + t];
16579 if a == 0.0 {
16580 continue;
16581 }
16582 for (j, o) in row.iter_mut().enumerate() {
16583 *o += a * right[j * b + t];
16584 }
16585 }
16586 };
16587 match pool {
16588 Some(p) if m > 1 => p.run_rows(m, &|s, e| {
16589 for i in s..e {
16590 body(i);
16591 }
16592 }),
16593 _ => {
16594 for i in 0..m {
16595 body(i);
16596 }
16597 }
16598 }
16599}
16600
16601fn add_into(dst: &mut [f32], src: &[f32], pool: Option<&Pool>) {
16604 let n = dst.len().min(src.len());
16605 match pool {
16606 Some(p) if n >= 1 << 16 => {
16607 let ptr = SendMut(dst.as_mut_ptr());
16608 let f = |s: usize, e: usize| {
16609 let ptr = &ptr;
16610 for blk in s..e {
16611 let (a, b) = (blk * 4096, ((blk + 1) * 4096).min(n));
16612 for i in a..b {
16613 unsafe { *ptr.0.add(i) += src[i] };
16614 }
16615 }
16616 };
16617 p.run_rows(n.div_ceil(4096), &f);
16618 }
16619 _ => {
16620 for (d, v) in dst.iter_mut().zip(&src[..n]) {
16621 *d += *v;
16622 }
16623 }
16624 }
16625}
16626
16627fn accum_outer(
16632 c: &mut [f32],
16633 m: usize,
16634 n: usize,
16635 b: usize,
16636 left: &[f32],
16637 right: &[f32],
16638 pool: Option<&Pool>,
16639) {
16640 const TILE: usize = 32;
16641 let tiles = m.div_ceil(TILE);
16642 let cp = SendMut(c.as_mut_ptr());
16643 let body = |ti: usize| {
16644 let cp = &cp;
16645 let i0 = ti * TILE;
16646 let i1 = (i0 + TILE).min(m);
16647 for t in 0..b {
16648 let r = &right[t * n..t * n + n];
16649 for i in i0..i1 {
16650 let a = left[i * b + t];
16651 if a == 0.0 {
16652 continue;
16653 }
16654 let row = unsafe { std::slice::from_raw_parts_mut(cp.0.add(i * n), n) };
16656 for (o, v) in row.iter_mut().zip(r) {
16657 *o += a * *v;
16658 }
16659 }
16660 }
16661 };
16662 match pool {
16663 Some(p) if tiles > 1 => p.run_rows(tiles, &|s, e| {
16664 for ti in s..e {
16665 body(ti);
16666 }
16667 }),
16668 _ => {
16669 for ti in 0..tiles {
16670 body(ti);
16671 }
16672 }
16673 }
16674}
16675
16676pub fn refit_flush() -> usize {
16678 let Some((dir, map)) = refit_dir() else {
16679 return 0;
16680 };
16681 let guard = map.lock().unwrap();
16682 let mut n = 0;
16683 for (li, acc) in guard.0.iter() {
16684 let w = |name: &str, v: &[f32]| {
16687 let path = format!("{dir}/{name}.{li}.f32");
16688 let bytes: Vec<u8> = v.iter().flat_map(|x| x.to_le_bytes()).collect();
16689 match std::fs::write(&path, &bytes) {
16690 Ok(()) => {}
16691 Err(e) => eprintln!(
16692 "refit: FAILED to write {path} ({} MB): {e}",
16693 bytes.len() / 1_000_000
16694 ),
16695 }
16696 };
16697 w("gss", &acc.gss);
16698 w("ya", &acc.ya);
16699 println!(
16700 "refit L{li}: {} support, {} tokens, hidden {}",
16701 acc.support.len(),
16702 acc.tokens,
16703 acc.hidden
16704 );
16705 n += 1;
16706 }
16707 n
16708}
16709
16710fn adump_row(li: usize, g: &[f32]) {
16715 use std::io::Write as _;
16716 static FILES: std::sync::OnceLock<
16717 Option<(
16718 String,
16719 std::sync::Mutex<std::collections::HashMap<usize, std::fs::File>>,
16720 )>,
16721 > = std::sync::OnceLock::new();
16722 let Some((prefix, map)) = FILES
16723 .get_or_init(|| {
16724 std::env::var("CMF_FFN_ADUMP")
16725 .ok()
16726 .map(|p| (p, std::sync::Mutex::new(std::collections::HashMap::new())))
16727 })
16728 .as_ref()
16729 else {
16730 return;
16731 };
16732 static SPAN: std::sync::OnceLock<(usize, usize)> = std::sync::OnceLock::new();
16735 let (from, to) = *SPAN.get_or_init(|| {
16736 let g = |k: &str, d: usize| {
16737 std::env::var(k)
16738 .ok()
16739 .and_then(|v| v.parse().ok())
16740 .unwrap_or(d)
16741 };
16742 (
16743 g("CMF_FFN_ADUMP_FROM", 0),
16744 g("CMF_FFN_ADUMP_TO", usize::MAX),
16745 )
16746 });
16747 if li < from || li > to {
16748 return;
16749 }
16750 let mut map = map.lock().unwrap();
16751 let f = map.entry(li).or_insert_with(|| {
16752 std::fs::File::create(format!("{prefix}.{li}.f16")).expect("adump file")
16753 });
16754 let mut bytes = Vec::with_capacity(g.len() * 2);
16755 for v in g {
16756 bytes.extend_from_slice(&cortiq_core::quant::f32_to_f16(*v).to_le_bytes());
16757 }
16758 let _ = f.write_all(&bytes);
16759}
16760
16761fn oracle_topk() -> usize {
16767 static K: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
16768 *K.get_or_init(|| {
16769 std::env::var("CMF_FFN_ORACLE_TOPK")
16770 .ok()
16771 .and_then(|v| v.parse().ok())
16772 .unwrap_or(0)
16773 })
16774}
16775
16776fn gate_topk() -> usize {
16782 static K: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
16783 *K.get_or_init(|| {
16784 std::env::var("CMF_FFN_GATE_TOPK")
16785 .ok()
16786 .and_then(|v| v.parse().ok())
16787 .unwrap_or(0)
16788 })
16789}
16790
16791fn gate_block() -> usize {
16798 static B: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
16799 *B.get_or_init(|| {
16800 std::env::var("CMF_FFN_GATE_BLOCK")
16801 .ok()
16802 .and_then(|v| v.parse().ok())
16803 .unwrap_or(1)
16804 })
16805}
16806
16807fn keep_top_blocks(g: &mut [f32], keep_n: usize, block: usize) {
16809 let n = g.len();
16810 let nb = n.div_ceil(block);
16811 let kb = (keep_n.div_ceil(block)).clamp(1, nb);
16812 if kb >= nb {
16813 return;
16814 }
16815 let mut score: Vec<f32> = (0..nb)
16816 .map(|b| {
16817 g[b * block..((b + 1) * block).min(n)]
16818 .iter()
16819 .map(|v| v * v)
16820 .sum::<f32>()
16821 })
16822 .collect();
16823 let mut ord = score.clone();
16824 let (_, kth, _) = ord.select_nth_unstable_by(kb - 1, |a, b| {
16825 b.partial_cmp(a).unwrap_or(std::cmp::Ordering::Equal)
16826 });
16827 let thr = *kth;
16828 for b in 0..nb {
16829 if score[b] < thr {
16830 g[b * block..((b + 1) * block).min(n)].fill(0.0);
16831 }
16832 }
16833 score.clear();
16834}
16835
16836fn keep_top_k(g: &mut [f32], k: usize) {
16838 if gate_block() > 1 {
16839 return keep_top_blocks(g, k, gate_block());
16840 }
16841 let n = g.len();
16842 if k == 0 || k >= n {
16843 return;
16844 }
16845 let mut mag: Vec<f32> = g.iter().map(|v| v.abs()).collect();
16846 let (_, kth, _) = mag.select_nth_unstable_by(k - 1, |a, b| {
16847 b.partial_cmp(a).unwrap_or(std::cmp::Ordering::Equal)
16848 });
16849 let thr = *kth;
16850 for v in g.iter_mut() {
16851 if v.abs() < thr {
16852 *v = 0.0;
16853 }
16854 }
16855}
16856
16857fn probe_sq() -> bool {
16861 static S: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
16862 *S.get_or_init(|| std::env::var("CMF_FFN_PROBE_SQ").is_ok())
16863}
16864
16865fn probe_signed() -> bool {
16869 static S: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
16870 *S.get_or_init(|| std::env::var("CMF_FFN_PROBE_SIGNED").is_ok())
16871}
16872
16873fn meanfill() -> Option<&'static (usize, Vec<f32>)> {
16881 static M: std::sync::OnceLock<Option<(usize, Vec<f32>)>> = std::sync::OnceLock::new();
16882 M.get_or_init(|| {
16883 let p = std::env::var("CMF_FFN_MEANFILL").ok()?;
16884 let b = std::fs::read(&p).ok()?;
16885 let inter = u32::from_le_bytes(b[4..8].try_into().ok()?) as usize;
16886 let vals: Vec<f32> = b[8..]
16887 .chunks_exact(4)
16888 .map(|c| f32::from_le_bytes([c[0], c[1], c[2], c[3]]))
16889 .collect();
16890 eprintln!("meanfill: {} value(s), inter {inter}", vals.len());
16891 Some((inter, vals))
16892 })
16893 .as_ref()
16894}
16895
16896fn probe_topk() -> usize {
16899 static K: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
16900 *K.get_or_init(|| {
16901 std::env::var("CMF_FFN_PROBE_TOPK")
16902 .ok()
16903 .and_then(|v| v.parse().ok())
16904 .unwrap_or(0)
16905 })
16906}
16907
16908thread_local! {
16909 static FFN_PROBE: std::cell::RefCell<Option<Vec<Vec<f64>>>> =
16912 const { std::cell::RefCell::new(None) };
16913}
16914
16915fn dense_ffn_dynamic(d: &DenseFfn, x: &[f32], pool: Option<&Pool>, k: usize) -> Option<Vec<f32>> {
16928 if d.gate_proj.has_prism_contract()
16933 || d.up_proj.has_prism_contract()
16934 || d.down_proj.has_prism_contract()
16935 {
16936 return None;
16937 }
16938 let dt = d.down_t.as_ref()?;
16939 let inter = d.gate_proj.rows();
16940 let hidden = dt.cols();
16941 if k == 0 || k >= inter || d.act != Act::Silu {
16942 return None;
16943 }
16944 DYN_SCRATCH.with(|sc| {
16945 let mut sc = sc.borrow_mut();
16946 let DynScratch {
16947 g,
16948 mag,
16949 live,
16950 parts,
16951 } = &mut *sc;
16952 g.resize(inter, 0.0);
16953 d.gate_proj.matvec(x, g, pool);
16954 for v in g.iter_mut() {
16955 *v = inference::silu(*v);
16956 }
16957 mag.clear();
16960 mag.extend(g.iter().map(|v| v.abs()));
16961 let (_, kth, _) = mag.select_nth_unstable_by(k - 1, |a, b| {
16962 b.partial_cmp(a).unwrap_or(std::cmp::Ordering::Equal)
16963 });
16964 let thr = *kth;
16965 live.clear();
16966 live.extend((0..inter as u32).filter(|&n| g[n as usize].abs() >= thr));
16967 let mut out = vec![0.0f32; hidden];
16968 match pool {
16969 Some(p) if live.len() >= 64 => {
16970 let nw = p.n_workers() + 1;
16971 parts.clear();
16972 parts.resize(nw * hidden, 0.0);
16973 let ptr = SendMut(parts.as_mut_ptr());
16974 let n = live.len();
16975 let live_ref: &[u32] = live;
16976 let g_ref: &[f32] = g;
16977 p.run(&|w, workers| {
16978 let chunk = n.div_ceil(workers);
16979 let (s, e) = (w * chunk, ((w + 1) * chunk).min(n));
16980 if s >= e {
16981 return;
16982 }
16983 WORKER_SCRATCH.with(|ws| {
16984 let mut ws = ws.borrow_mut();
16985 let [scratch, acc] = &mut *ws;
16986 scratch.resize(hidden.max(x.len()), 0.0);
16987 acc.clear();
16988 acc.resize(hidden, 0.0);
16989 for (o, &nrm) in live_ref[s..e].iter().enumerate() {
16990 if let Some(&nx) = live_ref[s..e].get(o + 1) {
16993 d.up_proj.prefetch_row(nx as usize);
16994 dt.prefetch_row(nx as usize);
16995 }
16996 let idx = nrm as usize;
16997 let up = d.up_proj.row_dot(idx, x, scratch);
16998 let a = g_ref[idx] * up;
16999 if a != 0.0 {
17000 dt.add_row_scaled(idx, a, acc, scratch);
17001 }
17002 }
17003 for (j, v) in acc.iter().enumerate() {
17004 unsafe { *ptr.at(w * hidden + j) = *v };
17005 }
17006 });
17007 });
17008 for w in 0..nw {
17009 for (j, o) in out.iter_mut().enumerate() {
17010 *o += parts[w * hidden + j];
17011 }
17012 }
17013 }
17014 _ => {
17015 WORKER_SCRATCH.with(|ws| {
17016 let mut ws = ws.borrow_mut();
17017 let [scratch, _acc] = &mut *ws;
17018 scratch.resize(hidden.max(x.len()), 0.0);
17019 for &nrm in live.iter() {
17020 let idx = nrm as usize;
17021 let up = d.up_proj.row_dot(idx, x, scratch);
17022 let a = g[idx] * up;
17023 if a != 0.0 {
17024 dt.add_row_scaled(idx, a, &mut out, scratch);
17025 }
17026 }
17027 });
17028 }
17029 }
17030 Some(out)
17031 })
17032}
17033
17034struct DynScratch {
17037 g: Vec<f32>,
17038 mag: Vec<f32>,
17039 live: Vec<u32>,
17040 parts: Vec<f32>,
17041}
17042
17043thread_local! {
17044 static DYN_SCRATCH: std::cell::RefCell<DynScratch> = const {
17045 std::cell::RefCell::new(DynScratch {
17046 g: Vec::new(),
17047 mag: Vec::new(),
17048 live: Vec::new(),
17049 parts: Vec::new(),
17050 })
17051 };
17052 static WORKER_SCRATCH: std::cell::RefCell<[Vec<f32>; 2]> =
17054 const { std::cell::RefCell::new([Vec::new(), Vec::new()]) };
17055}
17056
17057fn dense_ffn_masked(d: &DenseFfn, x: &[f32], pool: Option<&Pool>, mask_row: &[u8]) -> Vec<f32> {
17062 let inter = d.gate_proj.rows();
17063 FFN_SCRATCH.with(|s| {
17064 let mut s = s.borrow_mut();
17065 let [g, u, ..] = &mut *s;
17066 g.resize(inter, 0.0);
17067 if d.act == Act::Silu && QTensor::matvec_silu_mul(&d.gate_proj, &d.up_proj, x, g, pool) {
17068 } else {
17070 u.resize(inter, 0.0);
17071 QTensor::matvec_many([&d.gate_proj, &d.up_proj], x, [g, u], pool);
17072 for i in 0..inter {
17073 g[i] = d.act.combine(g[i], u[i]);
17074 }
17075 }
17076 zero_masked_cols(g, 1, inter, mask_row);
17077 let mut out = attention::take_buf(d.down_proj.rows());
17078 d.down_proj.matvec(g, &mut out, pool);
17079 out
17080 })
17081}
17082
17083fn dense_ffn_gpu(d: &DenseFfn, x: &[f32], _pool: Option<&Pool>) -> Option<Vec<f32>> {
17089 if d.gate_proj.has_prism_contract()
17090 || d.up_proj.has_prism_contract()
17091 || d.down_proj.has_prism_contract()
17092 {
17093 return None;
17094 }
17095 if d.act != Act::Silu {
17097 return None;
17098 }
17099 if d.gate_proj.rows() < crate::gpu::min_rows() && !d.gate_proj.is_q1() {
17102 return None;
17103 }
17104 let mut jobs: Vec<crate::gpu::MoeJob> = Vec::with_capacity(1);
17105 let mut model_ref = None;
17106 moe_push_job(d, x, 1.0, &mut jobs, &mut model_ref)?;
17107 let model = model_ref?;
17108 let hidden = jobs[0].down.1;
17109 let mut out = attention::take_buf(hidden);
17110 if crate::gpu::moe_block(&model, &jobs, &mut out) {
17111 Some(out)
17112 } else {
17113 let mut out = out;
17114 attention::recycle_buf(&mut out);
17115 None
17116 }
17117}
17118
17119#[allow(clippy::type_complexity)]
17124#[allow(clippy::type_complexity)]
17125pub(crate) fn moe_parts(
17126 t: &QTensor,
17127) -> Option<(
17128 &std::sync::Arc<cortiq_core::CmfModel>,
17129 usize,
17130 usize,
17131 usize,
17132 &[f32],
17133 &[f32],
17134 bool,
17135 bool,
17136 bool,
17137)> {
17138 match t {
17139 QTensor::Mapped {
17140 model,
17141 idx,
17142 dtype: dt @ (cortiq_core::TensorDtype::Q8_2f | cortiq_core::TensorDtype::Q8Row),
17143 rows,
17144 cols,
17145 row_scale,
17146 col_field,
17147 ..
17148 } if (*dt == cortiq_core::TensorDtype::Q8Row) || !col_field.is_empty() => Some((
17149 model, *idx, *rows, *cols, row_scale, col_field, false, false, false,
17150 )),
17151 QTensor::Mapped {
17153 model,
17154 idx,
17155 dtype: cortiq_core::TensorDtype::Q1,
17156 rows,
17157 cols,
17158 ..
17159 } => Some((
17160 model,
17161 *idx,
17162 *rows,
17163 *cols,
17164 &[][..],
17165 &[][..],
17166 true,
17167 false,
17168 false,
17169 )),
17170 QTensor::Mapped {
17172 model,
17173 idx,
17174 dtype: cortiq_core::TensorDtype::Q4Tiled,
17175 rows,
17176 cols,
17177 ..
17178 } => Some((
17179 model,
17180 *idx,
17181 *rows,
17182 *cols,
17183 &[][..],
17184 &[][..],
17185 false,
17186 true,
17187 false,
17188 )),
17189 QTensor::Mapped {
17191 model,
17192 idx,
17193 dtype: cortiq_core::TensorDtype::Q4TiledP,
17194 rows,
17195 cols,
17196 ..
17197 } => Some((
17198 model,
17199 *idx,
17200 *rows,
17201 *cols,
17202 &[][..],
17203 &[][..],
17204 false,
17205 true,
17206 false,
17207 )),
17208 QTensor::Mapped {
17212 model,
17213 idx,
17214 dtype: cortiq_core::TensorDtype::Q2TiledP,
17215 rows,
17216 cols,
17217 ..
17218 } => Some((
17219 model,
17220 *idx,
17221 *rows,
17222 *cols,
17223 &[][..],
17224 &[][..],
17225 false,
17226 true,
17227 true,
17228 )),
17229 _ => None,
17230 }
17231}
17232
17233#[cfg(target_os = "macos")]
17241fn metal_moe_graph_parts(m: &MoeFfn, hidden: usize) -> Option<crate::gpu::GpuMoe<'_>> {
17242 if m.router_input_norm
17243 || m.route_tau.is_some()
17244 || m.mask.is_some()
17245 || m.per_expert_scale.is_some()
17246 || m.experts.is_empty()
17247 || m.top_k == 0
17248 || m.resonance.is_some()
17249 {
17250 return None;
17251 }
17252 let (sh, sg) = match &m.shared {
17255 Some((sh, sg)) => (sh, sg.as_ref()),
17256 None => return None,
17257 };
17258 let (rf, rr, rc) = m.router.f32_parts()?;
17259 if rr != m.experts.len() || rc != hidden {
17260 return None;
17261 }
17262 let shared_gated = sg.is_some();
17263 let sf = match sg {
17264 Some(sg) => {
17265 let (sf, sr, sc) = sg.f32_parts()?;
17266 if sr * sc != hidden {
17267 return None;
17268 }
17269 sf
17270 }
17271 None => &rf[..hidden],
17274 };
17275 if let Some(b) = &m.expert_bias {
17276 if b.len() != m.experts.len() {
17277 return None;
17278 }
17279 }
17280 let inter = m.experts[0].gate_proj.rows();
17281 let gu_q2 = m.experts[0].gate_proj.mapped_q2tp().is_some();
17284 let trio = |e: &DenseFfn| -> Option<(usize, usize, usize)> {
17285 if e.act != Act::Silu
17286 || e.gate_proj.rows() != inter
17287 || e.gate_proj.cols() != hidden
17288 || e.up_proj.rows() != inter
17289 || e.up_proj.cols() != hidden
17290 || e.down_proj.rows() != hidden
17291 || e.down_proj.cols() != inter
17292 {
17293 return None;
17294 }
17295 let pick = |t: &QTensor| -> Option<usize> {
17296 if gu_q2 {
17297 t.mapped_q2tp().map(|(_, i)| i)
17298 } else {
17299 t.mapped_q4tp().map(|(_, i)| i)
17300 }
17301 };
17302 Some((
17303 pick(&e.gate_proj)?,
17304 pick(&e.up_proj)?,
17305 e.down_proj.mapped_q4tp().map(|(_, i)| i)?,
17306 ))
17307 };
17308 let experts = m.experts.iter().map(trio).collect::<Option<Vec<_>>>()?;
17309 let shared = trio(sh)?;
17310 Some(crate::gpu::GpuMoe {
17311 router: rf,
17312 sgate: sf,
17313 experts,
17314 shared,
17315 n_exp: m.experts.len(),
17316 top_k: m.top_k,
17317 inter,
17318 norm_topk: m.norm_topk_prob,
17319 route_scale: m.routed_scaling,
17320 gu_q2,
17321 sigmoid: m.router_sigmoid,
17322 bias: m.expert_bias.as_deref(),
17323 shared_gated,
17324 })
17325}
17326
17327pub(crate) fn moe_push_job_parts<'a>(
17331 gate: &'a QTensor,
17332 up: &'a QTensor,
17333 down: &'a QTensor,
17334 x: &[f32],
17335 w: f32,
17336 swiglu_limit: f32,
17337 jobs: &mut Vec<crate::gpu::MoeJob<'a>>,
17338 model_ref: &mut Option<std::sync::Arc<cortiq_core::CmfModel>>,
17339) -> Option<()> {
17340 use crate::qtensor::prescale;
17341 let (gm, gi, gr, gc, grs, gcf, gq1, gq4, gq2) = moe_parts(gate)?;
17342 let (_, ui, ur, uc, urs, ucf, uq1, uq4, uq2) = moe_parts(up)?;
17343 let (_, di, dr, dc, drs, dcf, dq1, dq4, dq2) = moe_parts(down)?;
17344 if gq1 != uq1 || uq1 != dq1 || gq4 != uq4 || uq4 != dq4 || gq2 != uq2 {
17345 return None; }
17347 if gq2 && (dq2 || !dq4 || down.mapped_q4tp().is_none()) {
17350 return None;
17351 }
17352 if !gq2 && dq2 {
17353 return None;
17354 }
17355 model_ref.get_or_insert_with(|| gm.clone());
17356 let dt = |cf: &[f32]| {
17357 if cf.is_empty() {
17358 cortiq_core::TensorDtype::Q8Row
17359 } else {
17360 cortiq_core::TensorDtype::Q8_2f
17361 }
17362 };
17363 jobs.push(crate::gpu::MoeJob {
17364 gate: (gi, gr, gc, grs),
17365 up: (ui, ur, uc, urs),
17366 down: (di, dr, dc, drs),
17367 xs_gate: prescale(x, gcf, dt(gcf)).into_owned(),
17368 xs_up: prescale(x, ucf, dt(ucf)).into_owned(),
17369 down_col: dcf,
17370 w,
17371 q1: gq1,
17372 q4t: gq4 && !gq2 && gate.mapped_q4tp().is_none(),
17373 q4tp: gq4 && (gq2 || gate.mapped_q4tp().is_some()),
17374 gu_q2: gq2,
17375 swiglu_limit,
17376 });
17377 Some(())
17378}
17379
17380fn moe_push_job<'a>(
17382 d: &'a DenseFfn,
17383 x: &[f32],
17384 w: 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 if d.act != Act::Silu {
17390 return None; }
17392 let (gm, gi, gr, gc, grs, gcf, gq1, gq4, gq2) = moe_parts(&d.gate_proj)?;
17393 let (_, ui, ur, uc, urs, ucf, uq1, uq4, uq2) = moe_parts(&d.up_proj)?;
17394 let (_, di, dr, dc, drs, dcf, dq1, dq4, dq2) = moe_parts(&d.down_proj)?;
17395 if gq1 != uq1 || uq1 != dq1 || gq4 != uq4 || uq4 != dq4 || gq2 != uq2 {
17396 return None; }
17398 if gq2 && (dq2 || !dq4 || d.down_proj.mapped_q4tp().is_none()) {
17399 return None;
17400 }
17401 if !gq2 && dq2 {
17402 return None;
17403 }
17404 model_ref.get_or_insert_with(|| gm.clone());
17405 let gdt = if gcf.is_empty() {
17406 cortiq_core::TensorDtype::Q8Row
17407 } else {
17408 cortiq_core::TensorDtype::Q8_2f
17409 };
17410 let udt = if ucf.is_empty() {
17411 cortiq_core::TensorDtype::Q8Row
17412 } else {
17413 cortiq_core::TensorDtype::Q8_2f
17414 };
17415 jobs.push(crate::gpu::MoeJob {
17416 gate: (gi, gr, gc, grs),
17417 up: (ui, ur, uc, urs),
17418 down: (di, dr, dc, drs),
17419 xs_gate: prescale(x, gcf, gdt).into_owned(),
17420 xs_up: prescale(x, ucf, udt).into_owned(),
17421 down_col: dcf,
17422 w,
17423 q1: gq1,
17424 q4t: gq4 && !gq2 && d.gate_proj.mapped_q4tp().is_none(),
17425 q4tp: gq4 && (gq2 || d.gate_proj.mapped_q4tp().is_some()),
17426 gu_q2: gq2,
17427 swiglu_limit: 0.0,
17428 });
17429 Some(())
17430}
17431
17432fn sparse_ffn_quant(
17439 d: &DenseFfn,
17440 x: &[f32],
17441 active: &[u16],
17442 hidden: usize,
17443 pool: Option<&Pool>,
17444) -> Vec<f32> {
17445 let n = active.len();
17446 let inter = d.gate_proj.rows();
17447 let mut act = vec![0.0f32; n];
17448 let need_scratch = !(d.gate_proj.sparse_col_ok() && d.up_proj.sparse_col_ok());
17451 let compute = |ai: usize| -> f32 {
17452 let idx = active[ai] as usize;
17453 if idx >= inter {
17454 return 0.0; }
17456 let mut s = if need_scratch {
17457 vec![0.0f32; hidden]
17458 } else {
17459 Vec::new()
17460 };
17461 let gate = d.gate_proj.row_dot(idx, x, &mut s);
17462 let up = d.up_proj.row_dot(idx, x, &mut s);
17463 d.act.combine(gate, up)
17464 };
17465 match pool {
17466 Some(p) if n >= 256 => {
17467 let ptr = SendMut(act.as_mut_ptr());
17468 p.run(&|widx, nw| {
17469 let chunk = n.div_ceil(nw);
17470 let (s, e) = (widx * chunk, ((widx + 1) * chunk).min(n));
17471 for ai in s..e {
17472 unsafe { *ptr.at(ai) = compute(ai) };
17473 }
17474 });
17475 }
17476 _ => {
17477 for (ai, a) in act.iter_mut().enumerate() {
17478 *a = compute(ai);
17479 }
17480 }
17481 }
17482 let mut out = vec![0.0f32; hidden];
17484 for (ai, &idx) in active.iter().enumerate() {
17485 let w = act[ai];
17486 if w.abs() >= 1e-12 && (idx as usize) < inter {
17487 d.down_proj.add_col_scaled(idx as usize, w, &mut out);
17488 }
17489 }
17490 out
17491}
17492
17493#[doc(hidden)]
17495pub fn sparse_ffn_quant_for_test(
17496 d: &DenseFfn,
17497 x: &[f32],
17498 active: &[u16],
17499 hidden: usize,
17500) -> Vec<f32> {
17501 sparse_ffn_quant(d, x, active, hidden, None)
17502}
17503
17504fn dequant_dense_f32(d: &DenseFfn) -> (Vec<f32>, Vec<f32>, Vec<f32>) {
17508 let deq = |t: &QTensor| -> Vec<f32> {
17509 let (rows, cols) = (t.rows(), t.cols());
17510 let mut out = vec![0.0f32; rows * cols];
17511 for r in 0..rows {
17512 t.row_f32(r, &mut out[r * cols..(r + 1) * cols]);
17513 }
17514 out
17515 };
17516 (deq(&d.gate_proj), deq(&d.up_proj), deq(&d.down_proj))
17517}
17518
17519struct SendMut(*mut f32);
17521unsafe impl Send for SendMut {}
17522unsafe impl Sync for SendMut {}
17523impl SendMut {
17524 #[inline]
17525 #[allow(clippy::mut_from_ref)]
17528 unsafe fn at(&self, i: usize) -> &mut f32 {
17529 unsafe { &mut *self.0.add(i) }
17530 }
17531}
17532
17533pub(crate) fn moe_route(
17546 logits: &[f32],
17547 m: &MoeFfn,
17548 allowed: Option<&[bool]>,
17549) -> (Vec<usize>, Vec<f32>, f32) {
17550 moe_route_with_eps(logits, m, allowed, 1e-6)
17551}
17552
17553pub(crate) fn moe_route_with_eps(
17562 logits: &[f32],
17563 m: &MoeFfn,
17564 allowed: Option<&[bool]>,
17565 sigmoid_denom_eps: f32,
17566) -> (Vec<usize>, Vec<f32>, f32) {
17567 let ne = logits.len();
17568 let admit = |e: usize| {
17574 m.mask.as_ref().is_none_or(|mk| mk[e])
17575 && allowed.is_none_or(|a| a.get(e).copied().unwrap_or(false))
17576 };
17577 if m.resonance.is_some() && m.top_k == 1 {
17590 let mut best: Option<usize> = None;
17591 for e in (0..ne).filter(|&e| admit(e)) {
17592 let l = logits[e];
17593 if l == f32::NEG_INFINITY || l.is_nan() {
17594 continue;
17595 }
17596 if best.is_none_or(|b| l > logits[b]) {
17597 best = Some(e);
17598 }
17599 }
17600 if let Some(b) = best {
17601 let mut p = vec![0.0f32; ne];
17602 p[b] = 1.0;
17603 return (vec![b], p, 1.0 / m.routed_scaling);
17604 }
17605 }
17606 let p: Vec<f32> = if m.router_sigmoid {
17612 logits.iter().map(|&l| 1.0 / (1.0 + (-l).exp())).collect()
17613 } else {
17614 let mx = logits.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
17615 if mx == f32::NEG_INFINITY {
17616 vec![1.0 / ne.max(1) as f32; ne]
17617 } else {
17618 let mut e: Vec<f32> = logits.iter().map(|&l| (l - mx).exp()).collect();
17619 let s: f32 = e.iter().sum();
17620 for v in &mut e {
17621 *v /= s;
17622 }
17623 e
17624 }
17625 };
17626 let mut idx: Vec<usize> = (0..ne).filter(|&e| admit(e)).collect();
17627 match &m.expert_bias {
17629 Some(b) => idx.sort_unstable_by(|&x, &y| {
17630 (p[y] + b[y])
17631 .partial_cmp(&(p[x] + b[x]))
17632 .unwrap()
17633 .then(x.cmp(&y))
17634 }),
17635 None => idx.sort_unstable_by(|&x, &y| p[y].partial_cmp(&p[x]).unwrap().then(x.cmp(&y))),
17636 }
17637 idx.truncate(m.top_k);
17638 if let Some(tau) = m.route_tau {
17642 let total: f32 = idx.iter().map(|&e| p[e]).sum();
17643 if total > 0.0 {
17644 let mut acc = 0.0f32;
17645 let mut keep = idx.len();
17646 for (i, &e) in idx.iter().enumerate() {
17647 acc += p[e];
17648 if acc >= tau * total {
17649 keep = i + 1;
17650 break;
17651 }
17652 }
17653 idx.truncate(keep);
17654 }
17655 }
17656 let wsum: f32 = if m.norm_topk_prob {
17657 let s: f32 = idx.iter().map(|&e| p[e]).sum();
17658 (if m.router_sigmoid {
17662 s + sigmoid_denom_eps
17663 } else {
17664 s
17665 }) / m.routed_scaling
17666 } else {
17667 1.0 / m.routed_scaling
17668 };
17669 (idx, p, wsum)
17670}
17671
17672fn moe_trace(idx: &[usize]) {
17674 moe_trace_at(crate::gpu::cur_layer() as i32, idx)
17675}
17676
17677pub(crate) fn moe_trace_at(li: i32, idx: &[usize]) {
17680 use std::io::Write;
17681 static F: std::sync::OnceLock<Option<std::sync::Mutex<std::fs::File>>> =
17682 std::sync::OnceLock::new();
17683 let Some(f) = F.get_or_init(|| {
17684 let p = std::env::var("CMF_MOE_TRACE").ok()?;
17685 Some(std::sync::Mutex::new(
17686 std::fs::OpenOptions::new()
17687 .create(true)
17688 .append(true)
17689 .open(p)
17690 .ok()?,
17691 ))
17692 }) else {
17693 return;
17694 };
17695 let ids: Vec<String> = idx.iter().map(|e| e.to_string()).collect();
17696 let _ = writeln!(f.lock().unwrap(), "{li}:{}", ids.join(","));
17697}
17698
17699pub(crate) fn moe_ffn(
17702 m: &MoeFfn,
17703 x: &[f32],
17704 pool: Option<&Pool>,
17705 allowed: Option<&[bool]>,
17706) -> Vec<f32> {
17707 let r = moe_ffn_route(m, x, pool, allowed);
17708 moe_ffn_experts(m, x, &r, pool)
17709}
17710
17711pub(crate) struct MoeRoute {
17715 pub idx: Vec<usize>,
17716 pub p: Vec<f32>,
17717 pub wsum: f32,
17718 pub logits: Vec<f32>,
17719}
17720
17721pub(crate) fn moe_ffn_route(
17727 m: &MoeFfn,
17728 x: &[f32],
17729 pool: Option<&Pool>,
17730 allowed: Option<&[bool]>,
17731) -> MoeRoute {
17732 accumulate_act(m, x, 1);
17733 let ne = m.experts.len();
17734 let mut logits = vec![0.0f32; ne];
17735 match &m.resonance {
17736 Some(r) => r.scores(x, &mut logits),
17737 None => m.router.matvec(x, &mut logits, pool),
17738 }
17739 let (idx, p, wsum) = moe_route(&logits, m, allowed);
17740 {
17741 let mut st = m.stats.borrow_mut();
17742 if st.len() < ne {
17743 st.resize(ne, 0);
17744 }
17745 for &e in &idx {
17746 st[e] += 1;
17747 }
17748 }
17749 moe_trace(&idx);
17755 MoeRoute {
17756 idx,
17757 p,
17758 wsum,
17759 logits,
17760 }
17761}
17762
17763pub(crate) fn moe_ffn_experts(
17766 m: &MoeFfn,
17767 x: &[f32],
17768 r: &MoeRoute,
17769 pool: Option<&Pool>,
17770) -> Vec<f32> {
17771 let (idx, p, wsum) = (&r.idx, &r.p, r.wsum);
17772 if crate::gpu::enabled_here() {
17777 match crate::gpu::probe_arm(crate::gpu::OpClass::Ffn) {
17778 crate::gpu::ProbeArm::Gpu => {
17779 let t0 = std::time::Instant::now();
17780 if let Some(out) = moe_ffn_gpu(m, x, &idx, &p, wsum, pool) {
17781 crate::gpu::probe_record(crate::gpu::OpClass::Ffn, true, t0.elapsed());
17782 return out;
17783 }
17784 }
17785 crate::gpu::ProbeArm::CpuTimed => {
17786 let t0 = std::time::Instant::now();
17787 let out = crate::gpu::cpu_scope(|| moe_ffn_cpu(m, x, &idx, &p, wsum, pool));
17788 crate::gpu::probe_record(crate::gpu::OpClass::Ffn, false, t0.elapsed());
17789 return out;
17790 }
17791 crate::gpu::ProbeArm::Cpu => {
17792 return crate::gpu::cpu_scope(|| moe_ffn_cpu(m, x, &idx, &p, wsum, pool));
17793 }
17794 }
17795 }
17796 moe_ffn_cpu(m, x, &idx, &p, wsum, pool)
17797}
17798
17799fn moe_ffn_banked(
17803 slot: &mut crate::mimo_moe::Slot,
17804 li: usize,
17805 m: &MoeFfn,
17806 x: &[f32],
17807 pool: Option<&Pool>,
17808) -> Vec<f32> {
17809 let t0 = std::time::Instant::now();
17810 let r = moe_ffn_route(m, x, pool, None);
17811 slot.note_route(t0.elapsed().as_nanos() as u64);
17812 match slot.forward(li, m, x, &r, pool) {
17813 Some(out) => out,
17814 None => crate::qtensor::float_activations_scope(|| {
17815 crate::gpu::cpu_scope(|| moe_ffn_experts(m, x, &r, pool))
17816 }),
17817 }
17818}
17819
17820fn moe_ffn_banked_rows(
17822 slot: &mut crate::mimo_moe::Slot,
17823 li: usize,
17824 m: &MoeFfn,
17825 xs: &[f32],
17826 b: usize,
17827 hidden: usize,
17828 pool: Option<&Pool>,
17829) -> Vec<f32> {
17830 let t0 = std::time::Instant::now();
17831 let routes: Vec<_> = xs
17832 .chunks_exact(hidden)
17833 .map(|x| moe_ffn_route(m, x, pool, None))
17834 .collect();
17835 slot.note_route(t0.elapsed().as_nanos() as u64);
17836 if let Some(out) = slot.forward_rows(li, m, xs, &routes, pool) {
17837 return out;
17838 }
17839 let mut out = Vec::with_capacity(b * hidden);
17840 for (x, r) in xs.chunks_exact(hidden).zip(&routes) {
17841 let row = slot.forward(li, m, x, r, pool).unwrap_or_else(|| {
17842 crate::qtensor::float_activations_scope(|| {
17844 crate::gpu::cpu_scope(|| moe_ffn_experts(m, x, r, pool))
17845 })
17846 });
17847 out.extend(row);
17848 }
17849 out
17850}
17851
17852fn graph_note(built: bool, layers_run: usize, total_layers: usize) {
17857 use std::sync::atomic::{AtomicBool, Ordering};
17858 if built {
17859 GRAPH_TOK_OK.fetch_add(1, Ordering::Relaxed);
17860 if total_layers > 0 && layers_run < total_layers {
17861 GRAPH_TOK_PREFIX.fetch_add(1, Ordering::Relaxed);
17862 } else {
17863 GRAPH_TOK_FULL.fetch_add(1, Ordering::Relaxed);
17864 }
17865 } else {
17866 GRAPH_TOK_MISS.fetch_add(1, Ordering::Relaxed);
17867 }
17868 static SAID: AtomicBool = AtomicBool::new(false);
17869 if !SAID.swap(true, Ordering::Relaxed) {
17870 if built {
17871 tracing::info!("wgpu whole-token graph: ACTIVE");
17872 } else {
17873 tracing::warn!("wgpu whole-token graph refused — per-op path");
17874 }
17875 }
17876}
17877
17878pub static GRAPH_TOK_OK: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
17882pub static GRAPH_TOK_MISS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
17883pub static GRAPH_TOK_PREFIX: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
17887pub static GRAPH_TOK_FULL: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
17889
17890pub static METAL_GRAPH_TOK_OK: std::sync::atomic::AtomicU64 =
17894 std::sync::atomic::AtomicU64::new(0);
17895pub static METAL_GRAPH_HEAD_OK: std::sync::atomic::AtomicU64 =
17896 std::sync::atomic::AtomicU64::new(0);
17897pub static METAL_GRAPH_HEAD_MISS: std::sync::atomic::AtomicU64 =
17898 std::sync::atomic::AtomicU64::new(0);
17899pub static METAL_GRAPH_LAYERS: std::sync::atomic::AtomicU64 =
17900 std::sync::atomic::AtomicU64::new(0);
17901pub static METAL_GRAPH_ERRORS: std::sync::atomic::AtomicU64 =
17902 std::sync::atomic::AtomicU64::new(0);
17903pub static METAL_PREFILL_CHUNKS: std::sync::atomic::AtomicU64 =
17907 std::sync::atomic::AtomicU64::new(0);
17908pub static METAL_PREFILL_ROWS: std::sync::atomic::AtomicU64 =
17909 std::sync::atomic::AtomicU64::new(0);
17910pub static METAL_PREFILL_HEAD_ROWS: std::sync::atomic::AtomicU64 =
17911 std::sync::atomic::AtomicU64::new(0);
17912pub static METAL_PREFILL_ERRORS: std::sync::atomic::AtomicU64 =
17913 std::sync::atomic::AtomicU64::new(0);
17914
17915fn moe_batch_enabled() -> bool {
17918 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
17919 *ON.get_or_init(|| std::env::var("CMF_MOE_BATCH").as_deref() != Ok("0"))
17920}
17921
17922fn moe_ffn_cpu_batched(
17928 m: &MoeFfn,
17929 x: &[f32],
17930 idx: &[usize],
17931 p: &[f32],
17932 wsum: f32,
17933 pool: Option<&Pool>,
17934) -> Option<Vec<f32>> {
17935 if idx.is_empty() || !moe_batch_enabled() {
17936 return None;
17937 }
17938 if FFN_PROBE.with(|pr| pr.borrow().is_some()) {
17942 return None;
17943 }
17944 let n = idx.len() + usize::from(m.shared.is_some());
17945 let mut pairs = Vec::with_capacity(n);
17946 let mut downs = Vec::with_capacity(n);
17947 let mut ws = Vec::with_capacity(n);
17948 for &e in idx {
17949 let d = &m.experts[e];
17950 if d.act != Act::Silu {
17951 return None;
17952 }
17953 pairs.push((&d.gate_proj, &d.up_proj));
17954 downs.push(&d.down_proj);
17955 ws.push(p[e] / wsum * m.per_expert_scale.as_ref().map_or(1.0, |v| v[e]));
17956 }
17957 if let Some((se, gate)) = &m.shared {
17960 if se.act != Act::Silu {
17961 return None;
17962 }
17963 let g = gate.as_ref().map_or(1.0, |gate| {
17964 let mut gl = [0.0f32; 1];
17965 gate.matvec(x, &mut gl, pool);
17966 1.0 / (1.0 + (-gl[0]).exp())
17967 });
17968 pairs.push((&se.gate_proj, &se.up_proj));
17969 downs.push(&se.down_proj);
17970 ws.push(g);
17971 }
17972 let inter = pairs[0].0.rows();
17973 let mut gs: Vec<Vec<f32>> = (0..pairs.len()).map(|_| vec![0f32; inter]).collect();
17974 if !QTensor::moe_gate_up_many(&pairs, x, &mut gs, pool) {
17975 return None;
17976 }
17977 let mut out = attention::take_buf(x.len());
17978 if !QTensor::moe_down_many(&downs, &gs, &ws, &mut out, pool) {
17979 attention::recycle_buf(&mut out);
17980 return None;
17981 }
17982 Some(out)
17983}
17984
17985pub(crate) fn moe_cold_experts_cpu(
17991 experts: &[(&DenseFfn, f32)],
17992 x: &[f32],
17993 pool: Option<&Pool>,
17994) -> Vec<f32> {
17995 let mut out = attention::take_buf(x.len());
17996 if experts.is_empty() {
17997 return out;
17998 }
17999 let pairs: Vec<_> = experts
18000 .iter()
18001 .map(|(e, _)| (&e.gate_proj, &e.up_proj))
18002 .collect();
18003 let downs: Vec<_> = experts.iter().map(|(e, _)| &e.down_proj).collect();
18004 let weights: Vec<_> = experts.iter().map(|(_, w)| *w).collect();
18005 let inter = experts[0].0.gate_proj.rows();
18006 let mut activations: Vec<Vec<f32>> = (0..experts.len()).map(|_| vec![0.0; inter]).collect();
18007 if QTensor::moe_gate_up_many(&pairs, x, &mut activations, pool)
18008 && QTensor::moe_down_many(&downs, &activations, &weights, &mut out, pool)
18009 {
18010 return out;
18011 }
18012 out.fill(0.0);
18013 for &(expert, weight) in experts {
18014 let mut one = dense_ffn(expert, x, pool);
18015 for (o, v) in out.iter_mut().zip(&one) {
18016 *o += weight * v;
18017 }
18018 attention::recycle_buf(&mut one);
18019 }
18020 out
18021}
18022
18023pub(crate) fn moe_cold_experts_rows_cpu(
18027 jobs: &[Vec<(&DenseFfn, f32)>],
18028 xs: &[f32],
18029 hidden: usize,
18030 pool: Option<&Pool>,
18031) -> Vec<f32> {
18032 let mut out = vec![0.0; xs.len()];
18033 let mut experts: Vec<&DenseFfn> = Vec::new();
18034 let mut groups: Vec<Vec<usize>> = Vec::new();
18035 let mut terms = vec![Vec::new(); jobs.len()];
18036 for (r, row) in jobs.iter().enumerate() {
18037 for &(e, w) in row {
18038 let g = match experts.iter().position(|&d| std::ptr::eq(d, e)) {
18039 Some(g) => g,
18040 None => {
18041 experts.push(e);
18042 groups.push(Vec::new());
18043 groups.len() - 1
18044 }
18045 };
18046 terms[r].push((g, groups[g].len(), w));
18047 groups[g].push(r);
18048 }
18049 }
18050 if experts.is_empty() {
18051 return out;
18052 }
18053 let pairs: Vec<_> = experts.iter().map(|e| (&e.gate_proj, &e.up_proj)).collect();
18054 let downs: Vec<_> = experts.iter().map(|e| &e.down_proj).collect();
18055 let lens: Vec<_> = groups.iter().map(Vec::len).collect();
18056 let count: usize = lens.iter().sum();
18057 let mut acts = vec![vec![0.0; experts[0].gate_proj.rows()]; count];
18058 let mut ds = vec![vec![0.0; hidden]; count];
18059 if QTensor::moe_gate_up_rows(&pairs, &groups, xs, &mut acts, pool)
18060 && QTensor::moe_down_rows(&downs, &lens, &acts, &mut ds, pool)
18061 {
18062 let mut offset = 0;
18063 let offsets: Vec<_> = lens
18064 .iter()
18065 .map(|&n| {
18066 let start = offset;
18067 offset += n;
18068 start
18069 })
18070 .collect();
18071 for (r, terms) in terms.iter().enumerate() {
18072 for &(g, slot, w) in terms {
18073 for (o, &v) in out[r * hidden..(r + 1) * hidden]
18074 .iter_mut()
18075 .zip(&ds[offsets[g] + slot])
18076 {
18077 *o += w * v;
18078 }
18079 }
18080 }
18081 } else {
18082 for (r, jobs) in jobs.iter().enumerate() {
18083 if !jobs.is_empty() {
18084 let mut row = moe_cold_experts_cpu(jobs, &xs[r * hidden..(r + 1) * hidden], pool);
18085 out[r * hidden..(r + 1) * hidden].copy_from_slice(&row);
18086 attention::recycle_buf(&mut row);
18087 }
18088 }
18089 }
18090 out
18091}
18092
18093fn moe_ffn_cpu(
18095 m: &MoeFfn,
18096 x: &[f32],
18097 idx: &[usize],
18098 p: &[f32],
18099 wsum: f32,
18100 pool: Option<&Pool>,
18101) -> Vec<f32> {
18102 if let Some(out) = moe_ffn_cpu_batched(m, x, idx, p, wsum, pool) {
18103 return out;
18104 }
18105 let mut out = attention::take_buf(x.len());
18106 for &e in idx {
18107 let mut eo = dense_ffn(&m.experts[e], x, pool);
18108 let w = p[e] / wsum * m.per_expert_scale.as_ref().map_or(1.0, |v| v[e]);
18109 for i in 0..out.len() {
18110 out[i] += w * eo[i];
18111 }
18112 attention::recycle_buf(&mut eo);
18113 }
18114 if let Some((se, gate)) = &m.shared {
18115 let mut so = dense_ffn(se, x, pool);
18116 let g = gate.as_ref().map_or(1.0, |gate| {
18117 let mut gl = [0.0f32; 1];
18118 gate.matvec(x, &mut gl, pool);
18119 1.0 / (1.0 + (-gl[0]).exp())
18120 });
18121 for i in 0..out.len() {
18122 out[i] += g * so[i];
18123 }
18124 attention::recycle_buf(&mut so);
18125 }
18126 out
18127}
18128
18129#[allow(clippy::too_many_arguments)]
18137pub(crate) fn mla_attention(
18138 w: &MlaWeights,
18139 normed: &[f32],
18140 cache: &mut crate::kv_cache::LayerKvCache,
18141 position: usize,
18142 inv_freq: &[f32],
18143 rope_scale: f32,
18144 eps: f64,
18145 pool: Option<&Pool>,
18146) -> Vec<f32> {
18147 let (nh, dr, dn, dv, lora) = (w.nh, w.qk_rope, w.qk_nope, w.v_dim, w.lora);
18148 let hd = dr + dn;
18149 let mut q = vec![0.0f32; nh * hd];
18150 match (&w.q_a, &w.q_a_norm) {
18151 (Some(qa), Some(qn)) => {
18152 let mut t = vec![0.0f32; qa.rows()];
18153 qa.matvec(normed, &mut t, pool);
18154 let tn = inference::rms_norm(&t, qn, eps, NormStyle::Qwen);
18155 w.q_proj.matvec(&tn, &mut q, pool);
18156 }
18157 _ => w.q_proj.matvec(normed, &mut q, pool),
18158 }
18159 let mut ca = vec![0.0f32; lora + dr];
18160 w.kv_a.matvec(normed, &mut ca, pool);
18161 let (c_lat, k_rope) = ca.split_at_mut(lora);
18162 let latn = inference::rms_norm(c_lat, &w.kv_a_norm, eps, NormStyle::Qwen);
18163 let mut kvb = vec![0.0f32; nh * (dn + dv)];
18164 w.kv_b.matvec(&latn, &mut kvb, pool);
18165 if !w.nope {
18166 attention::rope_rotate_scaled(k_rope, position, inv_freq, rope_scale);
18167 }
18168 for h in 0..nh {
18169 if !w.nope {
18170 attention::rope_rotate_scaled(
18171 &mut q[h * hd..h * hd + dr],
18172 position,
18173 inv_freq,
18174 rope_scale,
18175 );
18176 }
18177 }
18178 let mut k = vec![0.0f32; nh * hd];
18179 let mut v = vec![0.0f32; nh * hd];
18180 for h in 0..nh {
18181 k[h * hd..h * hd + dr].copy_from_slice(k_rope);
18182 k[h * hd + dr..(h + 1) * hd].copy_from_slice(&kvb[h * (dn + dv)..h * (dn + dv) + dn]);
18183 v[h * hd..h * hd + dv].copy_from_slice(&kvb[h * (dn + dv) + dn..(h + 1) * (dn + dv)]);
18184 }
18185 cache.append(&k, &v, &vec![true; nh]);
18186 let (ao, mut imp) = attention::attend_all_heads(&q, cache, nh, 1, hd, w.scale, None, 0.0);
18187 attention::recycle_buf(&mut imp);
18188 let mut ov = vec![0.0f32; nh * dv];
18189 for h in 0..nh {
18190 ov[h * dv..(h + 1) * dv].copy_from_slice(&ao[h * hd..h * hd + dv]);
18191 }
18192 let mut out = vec![0.0f32; w.o_proj.rows()];
18193 w.o_proj.matvec(&ov, &mut out, pool);
18194 out
18195}
18196
18197fn dense_moe_ffn(
18204 dm: &DenseMoeFfn,
18205 x_normed: &[f32],
18206 h_raw: &[f32],
18207 eps: f64,
18208 norm_style: NormStyle,
18209 pool: Option<&Pool>,
18210) -> Vec<f32> {
18211 let mut d = dense_ffn(&dm.dense, x_normed, pool);
18212 d = inference::rms_norm(&d, &dm.post_norm_1, eps, norm_style);
18213 let m = &dm.moe;
18214 let ne = m.experts.len();
18215 let mut logits = vec![0.0f32; ne];
18216 if m.router_input_norm {
18217 let ss: f32 = h_raw.iter().map(|v| v * v).sum::<f32>() / h_raw.len() as f32;
18218 let inv = 1.0 / (ss + eps as f32).sqrt();
18219 let xr: Vec<f32> = h_raw.iter().map(|v| v * inv).collect();
18220 m.router.matvec(&xr, &mut logits, pool);
18221 } else {
18222 m.router.matvec(h_raw, &mut logits, pool);
18223 }
18224 let (idx, p, wsum) = moe_route(&logits, m, None);
18225 {
18226 let mut st = m.stats.borrow_mut();
18227 if st.len() < ne {
18228 st.resize(ne, 0);
18229 }
18230 for &e in &idx {
18231 st[e] += 1;
18232 }
18233 }
18234 let x2 = inference::rms_norm(h_raw, &dm.pre_norm_2, eps, norm_style);
18235 let mo = moe_ffn_cpu(m, &x2, &idx, &p, wsum, pool);
18236 let mo = inference::rms_norm(&mo, &dm.post_norm_2, eps, norm_style);
18237 for (di, mi) in d.iter_mut().zip(&mo) {
18238 *di += mi;
18239 }
18240 d
18241}
18242
18243fn moe_gpu_refused(why: &'static str) {
18250 use std::sync::atomic::{AtomicBool, Ordering};
18251 static SAID: AtomicBool = AtomicBool::new(false);
18252 if !SAID.swap(true, Ordering::Relaxed) {
18253 tracing::warn!("MoE GPU block refused ({why}) — experts run on the CPU");
18254 }
18255}
18256
18257fn moe_ffn_gpu(
18258 m: &MoeFfn,
18259 x: &[f32],
18260 idx: &[usize],
18261 p: &[f32],
18262 wsum: f32,
18263 pool: Option<&Pool>,
18264) -> Option<Vec<f32>> {
18265 use crate::gpu::MoeJob;
18266
18267 let mut jobs: Vec<MoeJob> = Vec::with_capacity(idx.len() + 1);
18268 let mut model_ref = None;
18269 for &e in idx {
18270 if moe_push_job(&m.experts[e], x, p[e] / wsum, &mut jobs, &mut model_ref).is_none() {
18271 moe_gpu_refused("push_job(expert)");
18272 return None;
18273 }
18274 }
18275 if let Some((se, gate)) = &m.shared {
18276 let g = gate.as_ref().map_or(1.0, |gate| {
18277 let mut gl = [0.0f32; 1];
18278 gate.matvec(x, &mut gl, pool);
18279 1.0 / (1.0 + (-gl[0]).exp())
18280 });
18281 if moe_push_job(se, x, g, &mut jobs, &mut model_ref).is_none() {
18282 moe_gpu_refused("push_job(shared)");
18283 return None;
18284 }
18285 }
18286 let Some(model) = model_ref else {
18287 moe_gpu_refused("no model_ref");
18288 return None;
18289 };
18290 let hidden = jobs[0].down.1;
18291 let mut out = vec![0.0f32; hidden];
18292 if crate::gpu::moe_block(&model, &jobs, &mut out) {
18293 Some(out)
18294 } else {
18295 moe_gpu_refused("gpu::moe_block");
18296 None
18297 }
18298}
18299
18300fn ffn_forward(
18302 ffn: &FfnKind,
18303 x: &[f32],
18304 pool: Option<&Pool>,
18305 experts_allowed: Option<&[bool]>,
18306) -> Vec<f32> {
18307 match ffn {
18308 FfnKind::Dense(d) if !d.segs.is_empty() => tube_ffn(d, x, 1, pool, None),
18309 FfnKind::Dense(d) => dense_ffn(d, x, pool),
18310 FfnKind::Moe(m) => moe_ffn(m, x, pool, experts_allowed),
18311 FfnKind::DenseMoe(_) => unreachable!("DenseMoe dispatches via dense_moe_ffn"),
18315 }
18316}
18317
18318fn ffn_forward_pair(
18322 ffn: &FfnKind,
18323 x1: &[f32],
18324 x2: &[f32],
18325 pool: Option<&Pool>,
18326 experts_allowed: Option<&[bool]>,
18327) -> (Vec<f32>, Vec<f32>) {
18328 let d = match ffn {
18329 FfnKind::Dense(d) if !d.segs.is_empty() => {
18332 return (
18333 tube_ffn(d, x1, 1, pool, None),
18334 tube_ffn(d, x2, 1, pool, None),
18335 );
18336 }
18337 FfnKind::Dense(d) => d,
18338 FfnKind::Moe(m) => {
18339 return (
18340 moe_ffn(m, x1, pool, experts_allowed),
18341 moe_ffn(m, x2, pool, experts_allowed),
18342 );
18343 }
18344 FfnKind::DenseMoe(_) => unreachable!("DenseMoe dispatches via dense_moe_ffn"),
18345 };
18346 let inter = d.gate_proj.rows();
18347 FFN_SCRATCH.with(|s| {
18348 let mut s = s.borrow_mut();
18349 let [g1, g2, u1, u2] = &mut *s;
18350 g1.resize(inter, 0.0);
18351 g2.resize(inter, 0.0);
18352 u1.resize(inter, 0.0);
18353 u2.resize(inter, 0.0);
18354 QTensor::matvec2_many(
18357 [&d.gate_proj, &d.up_proj],
18358 x1,
18359 x2,
18360 [g1.as_mut_slice(), u1.as_mut_slice()],
18361 [g2.as_mut_slice(), u2.as_mut_slice()],
18362 pool,
18363 );
18364 for i in 0..inter {
18365 g1[i] = d.act.combine(g1[i], u1[i]);
18366 g2[i] = d.act.combine(g2[i], u2[i]);
18367 }
18368 let mut o1 = attention::take_buf(d.down_proj.rows());
18369 let mut o2 = attention::take_buf(d.down_proj.rows());
18370 d.down_proj.matvec2(g1, g2, &mut o1, &mut o2, pool);
18371 (o1, o2)
18372 })
18373}
18374
18375#[cfg(test)]
18376mod tests {
18377
18378 #[test]
18383 fn prefill_chunk_rule_widens_only_dense_on_discrete() {
18384 use super::{
18385 prefill_chunk_rule, ChunkHost, ChunkStackFacts, DISCRETE_DENSE_PREFILL_CHUNK,
18386 };
18387 let dense_card = ChunkStackFacts {
18388 plain_dense: true,
18389 discrete: true,
18390 gpu_on: true,
18391 ..Default::default()
18392 };
18393 assert!(dense_card.dense_on_discrete());
18394 assert_eq!(
18396 prefill_chunk_rule(None, ChunkHost::Other, dense_card.dense_on_discrete()),
18397 DISCRETE_DENSE_PREFILL_CHUNK
18398 );
18399 assert!(DISCRETE_DENSE_PREFILL_CHUNK > 48);
18400 for (label, facts) in [
18401 ("GDN hybrid / MoE / DeepSeek stack", ChunkStackFacts { plain_dense: false, ..dense_card }),
18402 ("integrated GPU", ChunkStackFacts { discrete: false, ..dense_card }),
18403 ("CPU only", ChunkStackFacts { gpu_on: false, discrete: false, ..dense_card }),
18404 ("capacity split", ChunkStackFacts { capacity_split: true, ..dense_card }),
18405 ("multi-GPU plan", ChunkStackFacts { multi_gpu: true, ..dense_card }),
18406 ("O(1) layers", ChunkStackFacts { o1: true, ..dense_card }),
18407 ] {
18408 assert!(!facts.dense_on_discrete(), "{label}");
18409 assert_eq!(
18410 prefill_chunk_rule(None, ChunkHost::Other, facts.dense_on_discrete()),
18411 48,
18412 "{label} keeps the historical x86 chunk"
18413 );
18414 }
18415 for dense in [false, true] {
18417 assert_eq!(prefill_chunk_rule(None, ChunkHost::Macos, dense), 512);
18418 assert_eq!(prefill_chunk_rule(None, ChunkHost::Aarch64, dense), 256);
18419 }
18420 for host in [ChunkHost::Macos, ChunkHost::Aarch64, ChunkHost::Other] {
18422 for dense in [false, true] {
18423 assert_eq!(prefill_chunk_rule(Some(48), host, dense), 48);
18424 assert_eq!(prefill_chunk_rule(Some(0), host, dense), 1);
18425 }
18426 }
18427 }
18428
18429 #[test]
18430 fn kv_reuse_plan_pulls_rows_decode_wrote_only_on_the_device() {
18431 use super::{ReuseLayer, ReusePlan, kv_reuse_plan};
18432 let full = |host_rows, device_rows| ReuseLayer {
18433 full: true,
18434 host_rows,
18435 device_rows,
18436 device_state: false,
18437 };
18438 assert_eq!(
18441 kv_reuse_plan(339, &[full(300, Some(339)), full(300, Some(339))]),
18442 ReusePlan::Pull(vec![(0, 300, 339), (1, 300, 339)])
18443 );
18444 assert_eq!(kv_reuse_plan(339, &[full(339, None)]), ReusePlan::Ready);
18446 assert_eq!(kv_reuse_plan(339, &[full(339, Some(345))]), ReusePlan::Ready);
18448 assert_eq!(
18450 kv_reuse_plan(339, &[full(300, Some(339)), full(339, None)]),
18451 ReusePlan::Pull(vec![(0, 300, 339)])
18452 );
18453 assert_eq!(kv_reuse_plan(339, &[full(300, Some(320))]), ReusePlan::Fresh);
18455 assert_eq!(kv_reuse_plan(339, &[full(300, None)]), ReusePlan::Fresh);
18456 assert_eq!(kv_reuse_plan(339, &[full(350, None)]), ReusePlan::Fresh);
18457 let conv = |device_state| ReuseLayer {
18460 full: false,
18461 host_rows: 0,
18462 device_rows: None,
18463 device_state,
18464 };
18465 assert_eq!(
18466 kv_reuse_plan(339, &[conv(true), full(300, Some(339))]),
18467 ReusePlan::Fresh
18468 );
18469 assert_eq!(kv_reuse_plan(339, &[conv(false), full(339, None)]), ReusePlan::Ready);
18470 }
18471
18472 #[test]
18473 fn nll_graph_policy_scopes_only_the_fused_head() {
18474 for (label, unmasked, prefer_graph, native_metal, want_graph, want_head) in [
18475 ("vulkan graph", true, true, false, true, false),
18477 ("native Metal graph", true, true, true, true, true),
18479 ("masked", false, true, false, false, false),
18481 ("graph disabled", true, false, true, false, false),
18482 ] {
18483 let (graph_quality, graph_head_required) =
18484 super::nll_graph_policy(unmasked, prefer_graph, native_metal);
18485 assert_eq!(graph_quality, want_graph, "{label}: graph quality");
18486 assert_eq!(graph_head_required, want_head, "{label}: fused head");
18487 }
18488 }
18489
18490 #[test]
18491 fn mtp_prefill_pair_boundaries_skip_only_final_prompt_row() {
18492 assert_eq!(mtp_prefill_pair_count(0, 128, 256), 128);
18493 assert_eq!(mtp_prefill_pair_count(128, 256, 256), 127);
18494 assert_eq!(mtp_prefill_pair_count(0, 256, 256), 255);
18495 assert_eq!(mtp_prefill_pair_count(256, 256, 256), 0);
18496 assert_eq!(mtp_prefill_pair_count(300, 320, 256), 0);
18497 }
18498
18499 #[test]
18500 fn cancel_flag_stops_generation() {
18501 let mut p = create_test_pipeline(16, 32, 2, 2, 8, 2, 32);
18502 p.cancel.store(true, std::sync::atomic::Ordering::Relaxed);
18505 let r = p.generate_from_ids(&[1, 2, 3], 8, None, None).unwrap();
18506 assert_eq!(r.finish_reason, "cancelled");
18507 assert!(
18508 r.token_ids.is_empty(),
18509 "no tokens after cancel: {:?}",
18510 r.token_ids
18511 );
18512 assert_eq!(p.kv_cache.seq_len(), 0);
18513 assert!(p.kv_history.is_empty());
18514 assert!(!p.graph_want_logits);
18515 assert!(p.graph_logits.is_none());
18516 let r2 = p.generate_from_ids(&[1, 2, 3], 4, None, None).unwrap();
18518 assert_ne!(r2.finish_reason, "cancelled");
18519 }
18520 use super::*;
18521
18522 #[test]
18530 fn dynamic_ffn_equals_the_zeroing_arm() {
18531 let (hidden, inter) = (8usize, 32usize);
18532 let synth = |n: usize, salt: usize| -> Vec<f32> {
18533 (0..n)
18534 .map(|i| (((i * 29 + salt * 13 + 7) % 89) as f32 / 89.0 - 0.5) * 0.6)
18535 .collect()
18536 };
18537 let down = synth(hidden * inter, 3);
18538 let mut down_t = vec![0.0f32; inter * hidden];
18539 for r in 0..hidden {
18540 for c in 0..inter {
18541 down_t[c * hidden + r] = down[r * inter + c];
18542 }
18543 }
18544 let d = DenseFfn {
18545 gate_proj: QTensor::from_f32(synth(inter * hidden, 1), inter, hidden),
18546 up_proj: QTensor::from_f32(synth(inter * hidden, 2), inter, hidden),
18547 down_proj: QTensor::from_f32(down.clone(), hidden, inter),
18548 act: Act::Silu,
18549 down_t: Some(QTensor::from_f32(down_t, inter, hidden)),
18550 segs: Vec::new(),
18551 };
18552 let x = synth(hidden, 11);
18553 let k = 12usize;
18554 let got = dense_ffn_dynamic(&d, &x, None, k).expect("down_t present");
18555 let mut g = vec![0.0f32; inter];
18557 d.gate_proj.matvec(&x, &mut g, None);
18558 let mut u = vec![0.0f32; inter];
18559 d.up_proj.matvec(&x, &mut u, None);
18560 for v in g.iter_mut() {
18561 *v = inference::silu(*v);
18562 }
18563 keep_top_k(&mut g, k);
18564 for i in 0..inter {
18565 g[i] *= u[i];
18566 }
18567 let mut want = vec![0.0f32; hidden];
18568 d.down_proj.matvec(&g, &mut want, None);
18569 for (a, b) in want.iter().zip(&got) {
18570 assert!((a - b).abs() < 1e-5, "dynamic {b} vs reference {a}");
18571 }
18572 }
18573
18574 #[test]
18580 fn tube_ffn_open_equals_dense_and_closed_equals_masked() {
18581 let (hidden, core, tube) = (8usize, 12usize, 8usize);
18582 let inter = core + tube;
18583 let synth = |n: usize, salt: usize| -> Vec<f32> {
18584 (0..n)
18585 .map(|i| (((i * 41 + salt * 17 + 5) % 97) as f32 / 97.0 - 0.5) * 0.5)
18586 .collect()
18587 };
18588 let (g_all, u_all) = (synth(inter * hidden, 1), synth(inter * hidden, 2));
18589 let d_all = synth(hidden * inter, 3);
18590 let dense = DenseFfn {
18592 gate_proj: QTensor::from_f32(g_all.clone(), inter, hidden),
18593 up_proj: QTensor::from_f32(u_all.clone(), inter, hidden),
18594 down_proj: QTensor::from_f32(d_all.clone(), hidden, inter),
18595 act: Act::Silu,
18596 down_t: None,
18597 segs: Vec::new(),
18598 };
18599 let rows =
18600 |v: &[f32], a: usize, b: usize| -> Vec<f32> { v[a * hidden..b * hidden].to_vec() };
18601 let cols = |v: &[f32], a: usize, b: usize| -> Vec<f32> {
18602 let mut o = Vec::with_capacity(hidden * (b - a));
18603 for r in 0..hidden {
18604 o.extend_from_slice(&v[r * inter + a..r * inter + b]);
18605 }
18606 o
18607 };
18608 let tubed = DenseFfn {
18609 down_t: None,
18610 gate_proj: QTensor::from_f32(rows(&g_all, 0, core), core, hidden),
18611 up_proj: QTensor::from_f32(rows(&u_all, 0, core), core, hidden),
18612 down_proj: QTensor::from_f32(cols(&d_all, 0, core), hidden, core),
18613 act: Act::Silu,
18614 segs: vec![FfnSeg {
18615 gate: QTensor::from_f32(rows(&g_all, core, inter), tube, hidden),
18616 up: QTensor::from_f32(rows(&u_all, core, inter), tube, hidden),
18617 down: QTensor::from_f32(cols(&d_all, core, inter), hidden, tube),
18618 start: core,
18619 width: tube,
18620 }],
18621 };
18622 let x = synth(hidden, 7);
18623 let want = dense_ffn(&dense, &x, None);
18624 let got = tube_ffn(&tubed, &x, 1, None, None);
18625 for (a, b) in want.iter().zip(&got) {
18626 assert!((a - b).abs() < 1e-5, "open tube: {a} vs {b}");
18627 }
18628 let mut bits = vec![0u8; inter.div_ceil(8)];
18630 for n in 0..core {
18631 bits[n / 8] |= 1 << (n % 8);
18632 }
18633 let closed = tube_ffn(&tubed, &x, 1, None, Some(&bits));
18634 let masked = dense_ffn_masked(&dense, &x, None, &bits);
18635 for (a, b) in masked.iter().zip(&closed) {
18636 assert!((a - b).abs() < 1e-5, "closed tube: {a} vs {b}");
18637 }
18638 let batch = tube_ffn(&tubed, &x, 1, None, Some(&bits));
18640 for (a, b) in closed.iter().zip(&batch) {
18641 assert_eq!(a, b, "batch arm disagrees with decode arm");
18642 }
18643 }
18644
18645 #[test]
18647 fn sparse_ffn_quant_equals_dense_with_inactive_zeroed() {
18648 let (hidden, inter) = (16usize, 40usize);
18649 let synth = |n: usize, salt: usize| -> Vec<f32> {
18650 (0..n)
18651 .map(|i| (((i * 37 + salt * 11 + 3) % 101) as f32 / 101.0 - 0.5) * 0.4)
18652 .collect()
18653 };
18654 let d = DenseFfn {
18655 gate_proj: QTensor::from_f32(synth(inter * hidden, 1), inter, hidden),
18656 up_proj: QTensor::from_f32(synth(inter * hidden, 2), inter, hidden),
18657 down_proj: QTensor::from_f32(synth(hidden * inter, 3), hidden, inter),
18658 act: Act::Silu,
18659 down_t: None,
18660 segs: Vec::new(),
18661 };
18662 let x = synth(hidden, 9);
18663 let active: Vec<u16> = (0..inter as u16).filter(|i| i % 3 == 0).collect();
18665
18666 let sparse = sparse_ffn_quant(&d, &x, &active, hidden, None);
18667
18668 let mut g = vec![0.0f32; inter];
18670 d.gate_proj.matvec(&x, &mut g, None);
18671 let mut u = vec![0.0f32; inter];
18672 d.up_proj.matvec(&x, &mut u, None);
18673 let act_set: std::collections::HashSet<u16> = active.iter().copied().collect();
18674 for i in 0..inter {
18675 g[i] = if act_set.contains(&(i as u16)) {
18676 inference::silu(g[i]) * u[i]
18677 } else {
18678 0.0
18679 };
18680 }
18681 let mut reference = vec![0.0f32; hidden];
18682 d.down_proj.matvec(&g, &mut reference, None);
18683
18684 let max_d = sparse
18685 .iter()
18686 .zip(&reference)
18687 .map(|(a, b)| (a - b).abs())
18688 .fold(0.0f32, f32::max);
18689 assert!(max_d < 1e-5, "sparse != dense-zeroed: max|Δ| = {max_d}");
18690 }
18691
18692 fn attach_test_mtp(p: &mut Pipeline) {
18694 let (h, inter, heads, kv, hd) = (
18695 p.hidden_size,
18696 p.intermediate_size,
18697 p.num_heads,
18698 p.num_kv_heads,
18699 p.head_dim,
18700 );
18701 let synth = |n: usize, salt: usize| -> Vec<f32> {
18702 (0..n)
18703 .map(|i| (((i * 29 + salt * 23 + 5) % 101) as f32 / 101.0 - 0.5) * 0.2)
18704 .collect()
18705 };
18706 let qt = |rows: usize, cols: usize, salt: usize| -> QTensor {
18707 QTensor::from_f32(synth(rows * cols, salt), rows, cols)
18708 };
18709 p.mtp = Some(MtpModule {
18710 enorm: vec![1.0; h],
18711 hnorm: vec![1.0; h],
18712 eh_proj: qt(h, 2 * h, 301),
18713 layer: LayerWeights {
18714 input_norm: vec![1.0; h],
18715 post_norm: vec![1.0; h],
18716 attn_out_norm: None,
18717 ffn_out_norm: None,
18718 layer_scale: None,
18719 ffn: FfnKind::Dense(DenseFfn {
18720 gate_proj: qt(inter, h, 315),
18721 up_proj: qt(inter, h, 316),
18722 down_proj: qt(h, inter, 317),
18723 act: Act::Silu,
18724 down_t: None,
18725 segs: Vec::new(),
18726 }),
18727 attn: AttnKind::Full {
18728 bias: None,
18729 wq: qt(heads * hd, h, 311),
18730 wk: qt(kv * hd, h, 312),
18731 wv: qt(kv * hd, h, 313),
18732 wo: qt(h, heads * hd, 314),
18733 q_norm: None,
18734 k_norm: None,
18735 output_gate: false,
18736 softplus_gate: None,
18737 },
18738 },
18739 final_norm: vec![1.0; h],
18740 kv: crate::kv_cache::LayerKvCache::new(kv, hd),
18741 });
18742 }
18743
18744 #[test]
18745 fn speculative_equals_vanilla_greedy() {
18746 unsafe { std::env::set_var("CMF_GPU_WGPU_GRAPH", "0") };
18750 let run = |spec: bool| {
18751 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
18752 p.sampler_config.temperature = 0.0;
18753 attach_test_mtp(&mut p);
18754 p.speculative = spec;
18755 let r = p.generate("abcdef", 12, None, None).unwrap();
18756 (r.token_ids, r.mtp_drafted, r.mtp_accepted)
18757 };
18758 let (vanilla, d0, _) = run(false);
18759 let (spec, d1, a1) = run(true);
18760 assert_eq!(d0, 0, "vanilla path must not draft");
18761 assert!(d1 > 0, "speculative path must draft");
18762 assert_eq!(
18763 vanilla, spec,
18764 "speculative must reproduce the exact greedy sequence (accepted {a1}/{d1})"
18765 );
18766 }
18767
18768 #[test]
18769 fn speculative_accepts_constant_oracle() {
18770 unsafe { std::env::set_var("CMF_GPU_WGPU_GRAPH", "0") };
18772 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
18773 p.sampler_config.temperature = 0.0;
18774 p.sampler_config.repetition_penalty = 1.0;
18775 p.weights.lm_head = QTensor::from_f32(vec![0.01; 64 * 8], 64, 8);
18778 attach_test_mtp(&mut p);
18779 p.speculative = true;
18780 let r = p.generate("abcd", 10, None, None).unwrap();
18781 assert!(r.mtp_drafted > 0);
18782 assert_eq!(
18783 r.mtp_accepted, r.mtp_drafted,
18784 "constant logits → every draft accepted"
18785 );
18786 assert!(r.token_ids.windows(2).all(|w| w[0] == w[1]));
18789 }
18790
18791 #[test]
18792 fn empty_prompt_is_an_error_not_a_panic() {
18793 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 260);
18794 let r = p.generate("", 4, None, None);
18795 assert!(r.is_err(), "empty prompt must be a clean error");
18796 }
18797
18798 #[test]
18799 fn every_token_enters_kv_exactly_once() {
18800 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
18801 p.sampler_config.temperature = 0.0;
18803 let r = p.generate("abc", 2, None, None).unwrap();
18804 assert_eq!(r.prompt_tokens, 3);
18805 assert_eq!(
18809 p.kv_cache.seq_len(),
18810 3 + r.tokens_generated - 1,
18811 "each token must be cached exactly once (v1 cached the last prompt token twice)"
18812 );
18813 }
18814
18815 #[test]
18816 fn generation_is_reproducible_with_seed() {
18817 let run = || {
18818 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
18819 p.generate("hello", 8, None, None).unwrap().token_ids
18820 };
18821 assert_eq!(run(), run());
18822 }
18823
18824 #[test]
18825 fn resetting_sampler_restarts_the_seeded_stream() {
18826 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
18827 let config = SamplerConfig {
18828 seed: Some(1234),
18829 ..SamplerConfig::default()
18830 };
18831 p.set_sampler_config(config.clone());
18832 let first = p.generate("hello", 8, None, None).unwrap().token_ids;
18833 p.set_sampler_config(config);
18834 let second = p.generate("hello", 8, None, None).unwrap().token_ids;
18835 assert_eq!(first, second);
18836 }
18837
18838 #[test]
18839 fn eviction_bounds_the_cache() {
18840 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 260);
18841 p.kv_cache.max_seq_len = 6;
18842 p.sampler_config.temperature = 0.0;
18843 let _ = p.generate("abcd", 12, None, None).unwrap();
18844 assert!(
18845 p.kv_cache.seq_len() <= 6 + 1,
18846 "cache must stay bounded by max_seq_len (got {})",
18847 p.kv_cache.seq_len()
18848 );
18849 }
18850
18851 #[test]
18852 fn confidence_matches_tokens_and_is_a_probability() {
18853 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
18854 p.sampler_config.temperature = 0.0;
18855 p.sampler_config.repetition_penalty = 1.0;
18856 let r = p.generate("abcd", 10, None, None).unwrap();
18857 assert_eq!(
18858 r.token_confidence.len(),
18859 r.token_ids.len(),
18860 "one confidence per emitted token"
18861 );
18862 for &c in &r.token_confidence {
18863 assert!((0.0..=1.0).contains(&c), "confidence out of [0,1]: {c}");
18864 }
18865 let logits = [1.0f32, 3.0, 0.5, 3.0];
18867 let p0 = top1_prob_t(&logits, 1, 1.0);
18868 let p1 = top1_prob_t(&logits, 3, 1.0);
18869 assert!((p0 - p1).abs() < 1e-6, "equal logits → equal prob");
18870 assert!(p0 > 0.0 && p0 < 1.0);
18871 let sharp = top1_prob_t(&logits, 1, 1.0);
18873 let soft = top1_prob_t(&logits, 1, 2.0);
18874 assert!(soft < sharp, "higher temperature lowers peak confidence");
18875 }
18876
18877 #[test]
18878 fn trace_is_opt_in_and_parallels_the_output() {
18879 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
18881 p.sampler_config.temperature = 0.0;
18882 p.sampler_config.repetition_penalty = 1.0;
18883 let r = p.generate("abcd", 10, None, None).unwrap();
18884 assert!(r.traces.is_empty(), "trace must be empty unless enabled");
18885
18886 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
18888 p.sampler_config.temperature = 0.0;
18889 p.sampler_config.repetition_penalty = 1.0;
18890 p.set_trace(true);
18891 let r = p.generate("abcd", 10, None, None).unwrap();
18892 assert_eq!(r.traces.len(), r.token_ids.len(), "one trace row per token");
18893 for (i, tr) in r.traces.iter().enumerate() {
18894 assert_eq!(tr.t, i, "trace index is sequential");
18895 assert_eq!(tr.token_id, r.token_ids[i], "trace token_id matches output");
18896 assert_eq!(
18897 tr.confidence, r.token_confidence[i],
18898 "trace confidence matches the confidence channel"
18899 );
18900 assert!(tr.active_skill.is_none() && tr.recon.is_none() && !tr.switched);
18902 }
18903 }
18904
18905 #[test]
18906 fn explain_prefill_logits_match_greedy_first_token() {
18907 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
18911 p.sampler_config.temperature = 0.0;
18912 p.sampler_config.repetition_penalty = 1.0;
18913 let ids = p.tokenizer.encode("abcd");
18914 let logits = p.prefill_next_logits(&ids, None);
18915 let argmax = logits
18916 .iter()
18917 .enumerate()
18918 .max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
18919 .unwrap()
18920 .0 as u32;
18921 let r = p.generate("abcd", 1, None, None).unwrap();
18922 assert_eq!(
18923 argmax, r.token_ids[0],
18924 "explain preview must match greedy emit"
18925 );
18926 }
18927
18928 #[test]
18929 fn laguna_shared_expert_is_unconditionally_added() {
18930 let matrix = |values: Vec<f32>| QTensor::from_f32(values, 2, 2);
18931 let identity = || matrix(vec![1.0, 0.0, 0.0, 1.0]);
18932 let zero_dense = || DenseFfn {
18933 gate_proj: matrix(vec![0.0; 4]),
18934 up_proj: matrix(vec![0.0; 4]),
18935 down_proj: matrix(vec![0.0; 4]),
18936 act: Act::Silu,
18937 down_t: None,
18938 segs: Vec::new(),
18939 };
18940 let shared = DenseFfn {
18941 gate_proj: identity(),
18942 up_proj: identity(),
18943 down_proj: identity(),
18944 act: Act::Silu,
18945 down_t: None,
18946 segs: Vec::new(),
18947 };
18948 let x = [1.0, 2.0];
18949 let expected = dense_ffn(&shared, &x, None);
18950 let moe = MoeFfn {
18951 router: QTensor::from_f32(vec![0.0, 0.0], 1, 2),
18952 experts: vec![zero_dense()],
18953 top_k: 1,
18954 norm_topk_prob: true,
18955 router_sigmoid: true,
18956 expert_bias: None,
18957 routed_scaling: 1.0,
18958 route_tau: None,
18959 shared: Some((shared, None)),
18960 stats: std::cell::RefCell::new(Vec::new()),
18961 act_sq: std::cell::RefCell::new(Vec::new()),
18962 act_rows: std::cell::RefCell::new(Vec::new()),
18963 mask: None,
18964 per_expert_scale: None,
18965 router_input_norm: false,
18966 resonance: None,
18967 grown: Vec::new(),
18968 };
18969 let actual = moe_ffn_cpu(&moe, &x, &[0], &[0.0], 1.0, None);
18970 for (actual, expected) in actual.iter().zip(expected) {
18971 assert!((actual - expected).abs() < 1e-6);
18972 }
18973 }
18974
18975 fn mimo_test_pipeline() -> Pipeline {
18984 let (hs, inter, nh, hd, vd, vocab) = (16usize, 24usize, 4usize, 8usize, 4usize, 64usize);
18985 let kvh = [1usize, 2, 2, 1];
18986 let synth = |n: usize, salt: usize| -> Vec<f32> {
18987 (0..n)
18988 .map(|i| (((i * 31 + salt * 17 + 7) % 97) as f32 / 97.0 - 0.5) * 0.4)
18989 .collect()
18990 };
18991 let qt = |rows: usize, cols: usize, salt: usize| {
18992 QTensor::from_f32(synth(rows * cols, salt), rows, cols)
18993 };
18994 let dense = |inter: usize, salt: usize| DenseFfn {
18995 gate_proj: qt(inter, hs, salt),
18996 up_proj: qt(inter, hs, salt + 1),
18997 down_proj: qt(hs, inter, salt + 2),
18998 act: Act::Silu,
18999 down_t: None,
19000 segs: Vec::new(),
19001 };
19002 let layers: Vec<LayerWeights> = (0..4)
19003 .map(|li| LayerWeights {
19004 input_norm: vec![1.0; hs],
19005 post_norm: vec![1.0; hs],
19006 attn_out_norm: None,
19007 ffn_out_norm: None,
19008 layer_scale: None,
19009 ffn: if li == 0 {
19010 FfnKind::Dense(dense(inter, 50))
19011 } else {
19012 FfnKind::Moe(MoeFfn {
19013 router: qt(4, hs, 60 + li),
19014 experts: (0..4).map(|e| dense(8, 70 + li * 10 + e * 3)).collect(),
19015 top_k: 2,
19016 norm_topk_prob: true,
19017 router_sigmoid: true,
19018 expert_bias: Some(vec![0.02, -0.03, 0.01, 0.0]),
19019 routed_scaling: 1.0,
19020 route_tau: None,
19021 shared: None,
19022 stats: std::cell::RefCell::new(Vec::new()),
19023 act_sq: std::cell::RefCell::new(Vec::new()),
19024 act_rows: std::cell::RefCell::new(Vec::new()),
19025 mask: None,
19026 per_expert_scale: None,
19027 router_input_norm: false,
19028 resonance: None,
19029 grown: Vec::new(),
19030 })
19031 },
19032 attn: AttnKind::Full {
19033 wq: qt(nh * hd, hs, li * 10 + 1),
19034 wk: qt(kvh[li] * hd, hs, li * 10 + 2),
19035 wv: qt(kvh[li] * vd, hs, li * 10 + 3),
19036 wo: qt(hs, nh * vd, li * 10 + 4),
19037 q_norm: None,
19038 k_norm: None,
19039 output_gate: false,
19040 softplus_gate: None,
19041 bias: None,
19042 },
19043 })
19044 .collect();
19045 let mut p = Pipeline::new(
19046 Tokenizer::byte_level(),
19047 PipelineWeights {
19048 embed_tokens: qt(vocab, hs, 100),
19049 layers,
19050 lm_head: qt(vocab, hs, 200),
19051 final_norm: vec![1.0; hs],
19052 },
19053 hs,
19054 inter,
19055 nh,
19056 1, hd,
19058 4,
19059 4,
19060 false,
19061 vocab,
19062 1e-6,
19063 1e7,
19064 NormStyle::Qwen,
19065 4096,
19066 SamplerConfig {
19067 seed: Some(7),
19068 ..Default::default()
19069 },
19070 );
19071 p.layer_dump = None;
19073 p.set_rotary(4, 1e7);
19074 p.sliding_layers = Some(vec![false, true, true, false]);
19075 p.swa = Some((3, usize::MAX));
19076 p.rotary_dim_local = Some(4);
19077 p.inv_freq_local = Some(std::sync::Arc::new(attention::rope_inv_freq(4, 1e4)));
19078 p.set_attn_geometry(Some(kvh.to_vec()), Some(vd)).unwrap();
19079 p.set_layer_sinks(1, vec![0.5, -1.0, 1.5, 0.0]).unwrap();
19080 p.set_layer_sinks(2, vec![-0.25, 2.0, 0.75, -1.5]).unwrap();
19081 p
19082 }
19083
19084 #[test]
19085 fn mimo_embedded_prompt_uses_rows_and_never_reuses_token_only_kv() {
19086 let mut p = mimo_test_pipeline();
19087 p.speculative = false;
19088 p.ignore_eos = true;
19089 p.sampler_config.temperature = 0.0;
19090 p.sampler_config.repetition_penalty = 1.0;
19091 let a = vec![3, 5, 7, 9, 11, 13];
19092 let b = vec![4, 8, 12, 16, 20, 24];
19093 let rows: Vec<_> = b.iter().flat_map(|&id| p.embed_id(id)).collect();
19094 let expected = p.generate_from_ids(&b, 8, None, None).unwrap().token_ids;
19095 let actual = p.generate_from_embeds(&a, &rows, 8, None, None).unwrap().token_ids;
19099 assert_eq!(actual, expected);
19100 assert!(p.kv_history.is_empty());
19101 let mut extended = a.clone();
19102 extended.push(17);
19103 let after_media = p.generate_from_ids(&extended, 8, None, None).unwrap().token_ids;
19104 p.reset_session();
19105 let fresh = p.generate_from_ids(&extended, 8, None, None).unwrap().token_ids;
19106 assert_eq!(after_media, fresh);
19107 assert!(p.generate_from_embeds(&a, &rows[..rows.len()-1], 1, None, None).is_err());
19108 p.reset_session();
19112 p.generate_from_ids(&a, 1, None, None).unwrap();
19113 let mut media_ids = p.kv_history.clone();
19114 assert!(!media_ids.is_empty());
19115 media_ids.extend_from_slice(&[19, 21, 23]);
19116 let source_ids: Vec<_> = (0..media_ids.len()).map(|i| b[i % b.len()]).collect();
19117 let source_rows: Vec<_> = source_ids.iter().flat_map(|&id| p.embed_id(id)).collect();
19118 let mut oracle = mimo_test_pipeline();
19119 oracle.speculative = false;
19120 oracle.ignore_eos = true;
19121 oracle.sampler_config.temperature = 0.0;
19122 oracle.sampler_config.repetition_penalty = 1.0;
19123 let expected = oracle.generate_from_ids(&source_ids, 8, None, None).unwrap().token_ids;
19124 assert_eq!(p.generate_from_embeds(&media_ids, &source_rows, 8, None, None).unwrap().token_ids, expected);
19125 assert!(p.kv_history.is_empty());
19126 let mut bad = rows;
19127 bad[0] = f32::NAN;
19128 assert!(p.generate_from_embeds(&a, &bad, 1, None, None).is_err());
19129 }
19130
19131 fn f32_bits(v: &[f32]) -> Vec<u32> {
19132 v.iter().map(|x| x.to_bits()).collect()
19133 }
19134
19135 #[test]
19142 fn mimo_shaped_decode_matches_prefill_batch_bitwise() {
19143 let mut p = mimo_test_pipeline();
19144 let kv: Vec<usize> = p.kv_cache.layers.iter().map(|l| l.num_kv_heads).collect();
19145 assert_eq!(kv, vec![1, 2, 2, 1]);
19146 assert!(p.kv_cache.layers[0].sinks.is_none() && p.kv_cache.layers[3].sinks.is_none());
19147 assert!(p.kv_cache.layers[1].sinks.is_some() && p.kv_cache.layers[2].sinks.is_some());
19148 let ids: Vec<u32> = (0..12u32).map(|i| (i * 7 + 3) % 64).collect();
19149 let hs = p.hidden_size;
19150 let mut decode = Vec::new();
19151 for (pos, &id) in ids.iter().enumerate() {
19152 let e = p.embed_single(id);
19153 let h = p.forward_layers(&e, pos, None);
19154 decode.push(p.logits_from_hidden(&h));
19155 }
19156 for l in &p.kv_cache.layers {
19157 assert_eq!(l.seq_len, 12);
19158 assert_eq!(l.head_values(0).len(), 12 * 8);
19160 }
19161 assert!(decode.iter().flatten().all(|v| v.is_finite()));
19162
19163 p.clear_sequence_state();
19164 let hb = p.prefill_batch_span(PrefillIn::Ids(&ids), 0, None, 0, p.num_layers);
19165 for pos in 0..ids.len() {
19166 let lg = p.logits_from_hidden(&hb[pos * hs..(pos + 1) * hs]);
19167 assert_eq!(
19168 f32_bits(&decode[pos]),
19169 f32_bits(&lg),
19170 "whole prompt, pos {pos}"
19171 );
19172 }
19173
19174 p.clear_sequence_state();
19175 let a = p.prefill_batch_span(PrefillIn::Ids(&ids[..5]), 0, None, 0, p.num_layers);
19176 let b = p.prefill_batch_span(PrefillIn::Ids(&ids[5..]), 5, None, 0, p.num_layers);
19177 for pos in 0..ids.len() {
19178 let row = if pos < 5 {
19179 &a[pos * hs..(pos + 1) * hs]
19180 } else {
19181 &b[(pos - 5) * hs..(pos - 4) * hs]
19182 };
19183 let lg = p.logits_from_hidden(row);
19184 assert_eq!(
19185 f32_bits(&decode[pos]),
19186 f32_bits(&lg),
19187 "two chunks, pos {pos}"
19188 );
19189 }
19190
19191 let last = |p: &mut Pipeline| {
19194 p.clear_sequence_state();
19195 let hb = p.prefill_batch_span(PrefillIn::Ids(&ids), 0, None, 0, p.num_layers);
19196 p.logits_from_hidden(&hb[11 * hs..12 * hs])
19197 };
19198 let base = last(&mut p);
19199 let mut no_sinks = mimo_test_pipeline();
19200 for l in &mut no_sinks.kv_cache.layers {
19201 l.sinks = None;
19202 }
19203 assert_ne!(
19204 f32_bits(&last(&mut no_sinks)),
19205 f32_bits(&base),
19206 "sinks are live"
19207 );
19208 let mut wide = mimo_test_pipeline();
19209 wide.swa = Some((64, usize::MAX));
19210 assert_ne!(
19211 f32_bits(&last(&mut wide)),
19212 f32_bits(&base),
19213 "window is live"
19214 );
19215
19216 p.clear_sequence_state();
19218 p.ignore_eos = true;
19219 let r = p.generate_from_ids(&ids, 4, None, None).unwrap();
19220 assert_eq!(r.token_ids.len(), 4);
19221 }
19222
19223 fn mimo_test_mtp(n: usize, gain: f32) -> mimo_mtp::MimoMtp {
19226 let (hs, inter, nh, hd, vd, nkv) = (16usize, 24usize, 4usize, 8usize, 4usize, 2usize);
19227 let synth = |len: usize, salt: usize| -> Vec<f32> {
19228 (0..len)
19229 .map(|i| (((i * 37 + salt * 13 + 3) % 89) as f32 / 89.0 - 0.5) * 0.6 * gain)
19230 .collect()
19231 };
19232 let qt = |rows: usize, cols: usize, salt: usize| {
19233 QTensor::from_f32(synth(rows * cols, salt), rows, cols)
19234 };
19235 let layers = (0..n)
19236 .map(|k| {
19237 let s = 500 + k * 40;
19238 let mut kv = crate::kv_cache::LayerKvCache::new(nkv, hd);
19239 kv.sinks = Some(vec![0.3, -0.7, 1.1, 0.0]);
19240 MtpModule {
19241 enorm: vec![1.0; hs],
19242 hnorm: vec![1.0; hs],
19243 eh_proj: qt(hs, 2 * hs, s),
19244 layer: LayerWeights {
19245 input_norm: vec![1.0; hs],
19246 post_norm: vec![1.0; hs],
19247 attn_out_norm: None,
19248 ffn_out_norm: None,
19249 layer_scale: None,
19250 attn: AttnKind::Full {
19251 wq: qt(nh * hd, hs, s + 1),
19252 wk: qt(nkv * hd, hs, s + 2),
19253 wv: qt(nkv * vd, hs, s + 3),
19254 wo: qt(hs, nh * vd, s + 4),
19255 q_norm: None,
19256 k_norm: None,
19257 output_gate: false,
19258 softplus_gate: None,
19259 bias: None,
19260 },
19261 ffn: FfnKind::Dense(DenseFfn {
19262 gate_proj: qt(inter, hs, s + 5),
19263 up_proj: qt(inter, hs, s + 6),
19264 down_proj: qt(hs, inter, s + 7),
19265 act: Act::Silu,
19266 down_t: None,
19267 segs: Vec::new(),
19268 }),
19269 },
19270 final_norm: vec![1.0; hs],
19271 kv,
19272 }
19273 })
19274 .collect();
19275 mimo_mtp::MimoMtp::from_layers(layers)
19276 }
19277
19278 fn mimo_greedy(p: &mut Pipeline, ids: &[u32], n: usize, spec: bool) -> GenerateResult {
19279 p.clear_sequence_state();
19280 p.speculative = spec;
19281 p.ignore_eos = true;
19282 p.sampler_config.temperature = 0.0;
19283 p.generate_from_ids(ids, n, None, None).unwrap()
19284 }
19285
19286 #[test]
19292 fn mimo_mtp_incremental_rounds_equal_the_teacher_forced_table() {
19293 for post in [false, true] {
19296 let mut p = mimo_test_pipeline();
19297 p.weights.final_norm = (0..p.hidden_size).map(|i| 0.5 + 0.1 * i as f32).collect();
19299 let mut st0 = mimo_test_mtp(3, 1.0);
19300 st0.post_norm_hidden = post;
19301 p.mimo_mtp = Some(st0);
19302 let ids: Vec<u32> = (0..14u32).map(|i| (i * 11 + 5) % 64).collect();
19303 let hs = p.hidden_size;
19304 let hb = p.prefill_batch_span(PrefillIn::Ids(&ids), 0, None, 0, p.num_layers);
19305 p.mimo_note_rows(&hb, 0);
19306 let mut st = p.mimo_mtp.take().unwrap();
19307 let k = 3;
19310 let mut inc = Vec::new();
19311 for t in 0..ids.len() - k - 1 {
19312 inc.push(p.mimo_mtp_draft(&mut st, t, &ids, k));
19313 }
19314 let s = ids.len();
19317 let mut reference = vec![vec![0u32; k]; s - k - 1];
19318 let mut fresh = mimo_test_mtp(3, 1.0);
19319 for (layer, m) in fresh.layers.iter_mut().enumerate() {
19320 let n = s - layer - 1;
19321 let mut cats = vec![0.0f32; n * 2 * hs];
19322 for j in 0..n {
19323 let e = p.embed_single(ids[j + layer + 1]);
19324 let raw = &hb[j * hs..(j + 1) * hs];
19325 let g = if post {
19326 inference::rms_norm(raw, &p.weights.final_norm, p.rms_eps, p.norm_style)
19327 } else {
19328 raw.to_vec()
19329 };
19330 let (ce, ch) = cats[j * 2 * hs..(j + 1) * 2 * hs].split_at_mut(hs);
19331 inference::rms_norm_into(&e, &m.enorm, p.rms_eps, p.norm_style, ce);
19332 inference::rms_norm_into(&g, &m.hnorm, p.rms_eps, p.norm_style, ch);
19333 }
19334 let mut x = vec![0.0f32; n * hs];
19335 m.eh_proj.matmat(&cats, n, &mut x, None);
19336 p.mimo_mtp_block(m, &mut x, n, 0);
19337 for (t, row) in reference.iter_mut().enumerate() {
19338 let y = inference::rms_norm(
19339 &x[t * hs..(t + 1) * hs],
19340 &m.final_norm,
19341 p.rms_eps,
19342 p.norm_style,
19343 );
19344 row[layer] = sampler::argmax(&p.lm_head_forward(&y));
19345 }
19346 }
19347 assert_eq!(inc, reference, "post_norm_hidden = {post}");
19348 let distinct: std::collections::HashSet<u32> =
19350 inc.iter().flatten().copied().collect();
19351 assert!(distinct.len() > 3, "{inc:?}");
19352 let last_t = ids.len() - k - 2;
19354 for m in &st.layers {
19355 assert_eq!(m.kv.seq_len, last_t + 1);
19356 }
19357 }
19358 }
19359
19360 #[test]
19367 fn mimo_speculative_greedy_equals_plain_greedy() {
19368 unsafe { std::env::set_var("CMF_GPU_WGPU_GRAPH", "0") };
19369 let ids: Vec<u32> = (0..9u32).map(|i| (i * 7 + 3) % 64).collect();
19370 let n = 24;
19371 let mut p = mimo_test_pipeline();
19372 let plain = mimo_greedy(&mut p, &ids, n, false);
19373 assert_eq!(plain.mtp_drafted, 0);
19374 assert_eq!(plain.token_ids.len(), n);
19375 let plain_kv = p.kv_cache.layers[0].seq_len;
19376
19377 p.mimo_mtp = Some(mimo_test_mtp(3, 1.0));
19379 let spec = mimo_greedy(&mut p, &ids, n, true);
19380 assert!(spec.mtp_drafted > 0, "the round must draft");
19381 assert_eq!(spec.token_ids, plain.token_ids);
19382 assert_eq!(p.kv_cache.layers[0].seq_len, plain_kv);
19383
19384 let mut truth: Vec<u32> = ids.clone();
19387 truth.extend(&plain.token_ids);
19388 let mut noisy = truth.clone();
19389 for (i, t) in noisy.iter_mut().enumerate() {
19390 if i % 5 == 0 {
19391 *t = (*t + 1) % 64;
19392 }
19393 }
19394 let mut st = mimo_test_mtp(3, 1.0);
19395 st.draft_override = Some(noisy);
19396 p.mimo_mtp = Some(st);
19397 let spec = mimo_greedy(&mut p, &ids, n, true);
19398 assert_eq!(spec.token_ids, plain.token_ids);
19399 assert_eq!(p.kv_cache.layers[0].seq_len, plain_kv);
19400 let stats = p.mimo_mtp.as_ref().unwrap().stats.clone();
19401 assert_eq!(stats.accepted as usize, spec.mtp_accepted);
19402 assert!(spec.mtp_accepted > 0 && spec.mtp_accepted < spec.mtp_drafted);
19403 assert!(
19404 stats.accept_hist.iter().filter(|&&c| c > 0).count() >= 3,
19405 "{:?}",
19406 stats.accept_hist
19407 );
19408 assert!(stats.tokens_per_round() > 1.5, "{}", stats.line());
19409
19410 let mut st = mimo_test_mtp(3, 1.0);
19413 st.draft_override = Some(truth);
19414 p.mimo_mtp = Some(st);
19415 let spec = mimo_greedy(&mut p, &ids, n, true);
19416 assert_eq!(spec.token_ids, plain.token_ids);
19417 assert_eq!(spec.mtp_accepted, spec.mtp_drafted);
19418 assert_eq!(p.kv_cache.layers[0].seq_len, plain_kv);
19419
19420 let off = mimo_greedy(&mut p, &ids, n, false);
19422 assert_eq!(off.token_ids, plain.token_ids);
19423 assert_eq!(off.mtp_drafted, 0);
19424 }
19425
19426 #[test]
19435 fn mimo_shaped_model_rides_the_wgpu_graph_geometry() {
19436 let p = mimo_test_pipeline();
19437 assert_eq!(
19438 p.graph_attn_decline_reason(),
19439 Some("per-layer KV head counts")
19440 );
19441 assert_eq!(p.wgpu_graph_attn_decline(), None);
19442 let g0 = p.graph_attn_geom(0).expect("full layer geometry");
19443 assert_eq!(
19444 (g0.nkv, g0.dv, g0.rd, g0.window, g0.sink.is_some()),
19445 (1, 4, 4, None, false)
19446 );
19447 assert_eq!(g0.invf, p.inv_freq.as_slice());
19448 let g1 = p.graph_attn_geom(1).expect("sliding layer geometry");
19449 assert_eq!((g1.nkv, g1.dv, g1.rd, g1.window), (2, 4, 4, Some(3)));
19450 assert_eq!(g1.sink, Some(&[0.5f32, -1.0, 1.5, 0.0][..]));
19451 assert_eq!(g1.invf, p.inv_freq_local.as_ref().unwrap().as_slice());
19452 assert_ne!(g0.invf, g1.invf, "two RoPE tables");
19453 let g3 = p.graph_attn_geom(3).expect("full layer geometry");
19454 assert_eq!((g3.nkv, g3.window, g3.sink.is_some()), (1, None, false));
19455
19456 let emb = p.embed_single(3);
19459 let mut lg = Vec::new();
19460 assert!(
19461 p.try_token_graph_wgpu_steps(&emb, 0, &mut lg, 1, None, None, 0, p.num_layers)
19462 .is_none()
19463 );
19464 let mut hid = emb.clone();
19465 assert_eq!(
19466 p.try_batch_graph_wgpu(&mut hid, &[0], 1, None),
19467 crate::gpu::BatchGraphOutcome::Declined
19468 );
19469 assert_eq!(hid, emb, "a declined batch graph leaves the rows untouched");
19470 assert!(p.try_multi_burst(3, 0, 4).is_none());
19471 assert!(
19472 p.graph_declines().is_empty(),
19473 "no attention decline logged: {:?}",
19474 p.graph_declines()
19475 );
19476 let plain = || create_test_pipeline(8, 16, 2, 1, 4, 2, 32);
19481 assert_eq!(plain().graph_attn_decline_reason(), None);
19482 assert_eq!(plain().wgpu_graph_attn_decline(), None);
19483 assert!(
19484 plain().graph_attn_geom(0).is_none(),
19485 "uniform models keep the historical arms"
19486 );
19487 let mut q = plain();
19488 q.set_layer_sinks(1, vec![0.25, -0.25]).unwrap();
19489 assert_eq!(
19490 q.graph_attn_decline_reason(),
19491 Some("learned attention sinks")
19492 );
19493 assert_eq!(
19494 q.graph_attn_geom(1).unwrap().sink,
19495 Some(&[0.25f32, -0.25][..])
19496 );
19497 let mut q = plain();
19498 q.set_attn_geometry(None, Some(2)).unwrap();
19499 assert_eq!(
19500 q.graph_attn_decline_reason(),
19501 Some("V heads narrower than Q/K heads")
19502 );
19503 assert_eq!(q.graph_attn_geom(0).unwrap().dv, 2);
19504 let mut q = plain();
19505 q.sliding_layers = Some(vec![true, false]);
19506 q.swa = Some((4, usize::MAX));
19507 assert_eq!(q.graph_attn_decline_reason(), Some("sliding-window layers"));
19508 assert_eq!(q.graph_attn_geom(0).unwrap().window, Some(4));
19509 assert_eq!(q.graph_attn_geom(1).unwrap().window, None);
19510
19511 let mut q = mimo_test_pipeline();
19514 q.rope_scale = 2.0;
19515 assert_eq!(
19516 q.wgpu_graph_attn_decline(),
19517 Some("scaled RoPE positions with per-layer geometry")
19518 );
19519 let emb = q.embed_single(3);
19520 assert!(
19521 q.try_token_graph_wgpu_steps(&emb, 0, &mut lg, 1, None, None, 0, q.num_layers)
19522 .is_none()
19523 );
19524 let _ = q.try_token_graph_wgpu_steps(&emb, 1, &mut lg, 1, None, None, 0, q.num_layers);
19525 let lines = q.graph_declines();
19526 assert_eq!(
19527 lines
19528 .iter()
19529 .filter(|l| l.starts_with("wgpu token graph") && l.contains("scaled RoPE"))
19530 .count(),
19531 1,
19532 "{lines:?}"
19533 );
19534 }
19535
19536 #[test]
19537 fn mimo_verify_rewind_preserves_lagging_host_caches() {
19538 let mut p = mimo_test_pipeline();
19539 for (li, layer) in p.kv_cache.layers.iter_mut().enumerate() {
19540 let row = vec![0.0; layer.num_kv_heads * layer.head_dim];
19541 for _ in 0..if li == 0 { 2 } else { 12 } {
19542 layer.append(&row, &row, &[]);
19543 }
19544 }
19545 p.mimo_verify_rewind(9).unwrap();
19546 assert_eq!(p.kv_cache.layers[0].seq_len, 2);
19547 for layer in &p.kv_cache.layers[1..] {
19548 assert_eq!(layer.seq_len, 9);
19549 }
19550 }
19551
19552 #[test]
19556 fn layer_dump_covers_every_position_and_layer_on_both_walks() {
19557 let dir = std::env::temp_dir().join(format!("cmf-layer-dump-{}", std::process::id()));
19558 let _ = std::fs::remove_dir_all(&dir);
19559 let mut p = mimo_test_pipeline();
19560 let hs = p.hidden_size;
19561 let ids = [5u32, 9, 11, 2, 40];
19562 p.layer_dump = Some(dir.join("decode"));
19563 for (pos, &id) in ids.iter().enumerate() {
19564 let e = p.embed_single(id);
19565 let _ = p.forward_layers(&e, pos, None);
19566 }
19567 p.clear_sequence_state();
19568 p.layer_dump = Some(dir.join("prefill"));
19569 let hb = p.prefill_batch_span(PrefillIn::Ids(&ids), 0, None, 0, p.num_layers);
19570 for pos in 0..ids.len() {
19571 for li in 0..p.num_layers {
19572 let name = format!("p{pos:06}_l{li:02}.f32");
19573 let a = std::fs::read(dir.join("decode").join(&name)).unwrap();
19574 let b = std::fs::read(dir.join("prefill").join(&name)).unwrap();
19575 assert_eq!(a.len(), hs * 4, "{name}");
19576 assert_eq!(a, b, "{name}");
19577 }
19578 }
19579 let last = std::fs::read(dir.join("prefill").join("p000004_l03.f32")).unwrap();
19580 let vals: Vec<f32> = last
19581 .chunks(4)
19582 .map(|c| f32::from_le_bytes(c.try_into().unwrap()))
19583 .collect();
19584 assert_eq!(f32_bits(&vals), f32_bits(&hb[4 * hs..5 * hs]));
19585 let _ = std::fs::remove_dir_all(&dir);
19586 }
19587
19588 #[test]
19589 fn attn_geometry_and_sinks_are_validated() {
19590 let mut p = create_test_pipeline(8, 16, 4, 2, 4, 2, 32);
19591 assert!(
19592 p.set_attn_geometry(Some(vec![2]), None).is_err(),
19593 "one entry per layer"
19594 );
19595 assert!(
19596 p.set_attn_geometry(Some(vec![2, 3]), None).is_err(),
19597 "3 does not divide 4"
19598 );
19599 assert!(p.set_attn_geometry(Some(vec![2, 0]), None).is_err());
19600 assert!(p.set_attn_geometry(None, Some(0)).is_err());
19601 assert!(
19602 p.set_attn_geometry(None, Some(5)).is_err(),
19603 "V wider than the head"
19604 );
19605 p.set_attn_geometry(None, Some(4)).unwrap();
19606 assert_eq!(
19607 p.v_head_dim, None,
19608 "v_head_dim == head_dim is the uniform case"
19609 );
19610 p.set_layer_sinks(1, vec![0.1; 4]).unwrap();
19611 p.set_attn_geometry(Some(vec![1, 4]), None).unwrap();
19612 assert_eq!(p.kv_cache.layers[0].num_kv_heads, 1);
19613 assert_eq!(p.kv_cache.layers[1].num_kv_heads, 4);
19614 assert!(
19615 p.kv_cache.layers[1].sinks.is_some(),
19616 "a reshape keeps the layer's sinks"
19617 );
19618 assert_eq!(p.layer_geom(1).0, 4);
19619 assert!(
19620 p.set_layer_sinks(0, vec![0.0; 3]).is_err(),
19621 "one sink per Q head"
19622 );
19623 assert!(p.set_layer_sinks(7, vec![0.0; 4]).is_err());
19624 assert!(p.set_layer_sinks(0, vec![f32::NAN, 0.0, 0.0, 0.0]).is_err());
19625 }
19626
19627 #[test]
19630 fn o1_is_never_armed_on_sink_window_or_narrow_v_layers() {
19631 let cfg = || {
19632 Some(crate::nystrom::O1Cfg {
19633 layers: crate::nystrom::O1Layers::All,
19634 m: 4,
19635 w: 8,
19636 sink: 2,
19637 rect: crate::nystrom::O1Rect::Aggregate,
19638 })
19639 };
19640 let mut p = mimo_test_pipeline();
19641 p.set_o1(cfg());
19642 assert!(!p.o1_active(), "every MiMo-shaped layer is ineligible");
19643 let mut q = create_test_pipeline(8, 16, 2, 1, 4, 3, 64);
19644 q.set_layer_sinks(1, vec![0.0, 0.0]).unwrap();
19645 q.sliding_layers = Some(vec![false, false, true]);
19646 q.swa = Some((4, usize::MAX));
19647 q.set_o1(cfg());
19648 assert_eq!(q.o1_flags, vec![true, false, false]);
19649 }
19650
19651 #[test]
19652 fn o1_batch_transition_publishes_one_epoch_before_serial_handoff() {
19653 const B: usize = 19;
19654 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
19655 p.set_o1(Some(crate::nystrom::O1Cfg {
19656 layers: crate::nystrom::O1Layers::All,
19657 m: 4,
19658 w: 8,
19659 sink: 2,
19660 rect: crate::nystrom::O1Rect::Aggregate,
19661 }));
19662 p.o1_begin_with_prefix(Some(B));
19663 let ids: Vec<u32> = (0..B as u32).collect();
19664 let _ = p.prefill_batch_span(PrefillIn::Ids(&ids), 0, None, 0, p.num_layers);
19665
19666 assert_eq!(p.o1_epoch, 1, "all layers publish one completed transition");
19667 assert!(p.kv_cache.layers.iter().all(|l| l.o1_sealed()));
19668 let next = p.embed_single(B as u32);
19669 let _ = p.forward_layers(&next, B, None);
19670 assert_eq!(p.o1_epoch, 1, "sealed handoff must not republish the epoch");
19671 }
19672
19673 #[test]
19674 fn o1_pair_transition_commits_scratch_before_epoch_publication() {
19675 const B: usize = 19;
19676 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
19677 let gdn_cfg = crate::linear_core::GdnCfg {
19681 num_v_heads: 2,
19682 num_k_heads: 1,
19683 key_head_dim: 2,
19684 value_head_dim: 4,
19685 conv_kernel: 3,
19686 hidden_size: 8,
19687 rms_eps: 1e-6,
19688 output_gate_sigmoid: false,
19689 };
19690 let synth = |n: usize, salt: usize| -> Vec<f32> {
19691 (0..n)
19692 .map(|i| (((i * 13 + salt * 7) % 97) as f32 / 97.0 - 0.5) * 0.4)
19693 .collect()
19694 };
19695 let qt = |rows: usize, cols: usize, salt: usize| {
19696 crate::qtensor::QTensor::from_f32(synth(rows * cols, salt), rows, cols)
19697 };
19698 let c_dim = gdn_cfg.conv_dim();
19699 let vd = gdn_cfg.num_v_heads * gdn_cfg.value_head_dim;
19700 p.weights.layers[0].attn = AttnKind::LinearGdn(crate::linear_core::GdnWeights {
19701 in_proj_qkv: qt(c_dim, 8, 1),
19702 in_proj_z: qt(vd, 8, 2),
19703 in_proj_a: qt(gdn_cfg.num_v_heads, 8, 3),
19704 in_proj_b: qt(gdn_cfg.num_v_heads, 8, 4),
19705 conv1d: synth(c_dim * gdn_cfg.conv_kernel, 5),
19706 a_log: vec![0.2, 0.5],
19707 dt_bias: synth(gdn_cfg.num_v_heads, 6),
19708 norm: vec![1.0; gdn_cfg.value_head_dim],
19709 out_proj: qt(8, vd, 7),
19710 });
19711 p.gdn_cfg = Some(gdn_cfg);
19712 p.set_o1(Some(crate::nystrom::O1Cfg {
19713 layers: crate::nystrom::O1Layers::All,
19714 m: 4,
19715 w: 8,
19716 sink: 2,
19717 rect: crate::nystrom::O1Rect::Aggregate,
19718 }));
19719 p.o1_begin_with_prefix(Some(B));
19720 for pos in 0..B - 2 {
19721 let emb = p.embed_single(pos as u32);
19722 let _ = p.forward_layers(&emb, pos, None);
19723 }
19724 let lane1_state = p.kv_cache.layers[0].linear_state.clone();
19725
19726 let e1 = p.embed_single((B - 2) as u32);
19727 let e2 = p.embed_single((B - 1) as u32);
19728 let _ = p.forward_pair(&e1, &e2, B - 2);
19729
19730 assert_eq!(p.o1_epoch, 1, "pair crossing B publishes one epoch");
19731 assert!(
19732 p.kv_cache
19733 .layers
19734 .iter()
19735 .enumerate()
19736 .all(|(li, l)| !p.o1_flags[li] || l.o1_sealed())
19737 );
19738 assert!(!p.kv_cache.layers[0].linear_state.is_empty());
19739 assert_ne!(
19740 p.kv_cache.layers[0].linear_state, lane1_state,
19741 "real pair must commit GDN lane 2 before returning"
19742 );
19743 assert!(p.kv_cache.layers[0].linear_scratch.is_empty());
19744 let next = p.embed_single(B as u32);
19745 let _ = p.forward_layers(&next, B, None);
19746 assert_eq!(p.o1_epoch, 1, "serial continuation must reuse the epoch");
19747 }
19748
19749 #[test]
19750 fn o1_error_observation_stays_terminal_until_reset() {
19751 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
19752 p.set_o1(Some(crate::nystrom::O1Cfg {
19753 layers: crate::nystrom::O1Layers::All,
19754 m: 4,
19755 w: 8,
19756 sink: 2,
19757 rect: crate::nystrom::O1Rect::Aggregate,
19758 }));
19759 p.o1_begin();
19760 p.kv_cache.layers[0].o1_abort("synthetic transition failure".into());
19761
19762 assert!(p.o1_seal_checked().is_err());
19763 assert!(
19764 p.o1_seal_checked().is_err(),
19765 "retry must see the sticky error"
19766 );
19767 let k = vec![0.2f32; 4];
19768 let v = vec![0.3f32; 4];
19769 p.kv_cache.layers[0].append(&k, &v, &[]);
19770 assert_eq!(p.kv_cache.layers[0].seq_len, 0);
19771
19772 p.reset_session();
19773 p.o1_begin();
19774 p.kv_cache.layers[0].append(&k, &v, &[]);
19775 assert_eq!(p.kv_cache.layers[0].seq_len, 1);
19776 }
19777
19778 #[test]
19779 fn nll_graph_failure_is_terminal_and_request_is_reusable() {
19780 let ids = vec![1u32, 2, 3, 4, 5, 6];
19781 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
19782 p.graph_logits = Some(vec![123.0]);
19783 p.graph_want_logits = true;
19784 p.graph_failed
19785 .store(true, std::sync::atomic::Ordering::Relaxed);
19786 p.cancel.store(true, std::sync::atomic::Ordering::Relaxed);
19787 let err = p.nll_ids_from(&ids, 0).expect_err("prior graph failure");
19788 assert!(err.contains("before NLL"));
19789 assert!(p.graph_logits.is_none());
19790 assert!(!p.graph_want_logits);
19791 assert!(!p.graph_failed.load(std::sync::atomic::Ordering::Relaxed));
19792 assert!(!p.cancel.load(std::sync::atomic::Ordering::Relaxed));
19793
19794 let mut fresh = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
19795 let expected = fresh.nll_ids_from(&ids, 0).expect("fresh NLL");
19796 let actual = p.nll_ids_from(&ids, 0).expect("reused NLL");
19797 assert_eq!(actual.1, expected.1);
19798 assert!((actual.0 - expected.0).abs() < 1e-9);
19799 }
19800
19801 #[test]
19802 fn nll_forward_failure_discards_partial_score_and_clears_sidechannels() {
19803 let ids = vec![1u32, 2, 3, 4, 5, 6];
19804 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
19805 p.nll_test_fail_at = Some(1);
19806 let err = p
19807 .nll_ids_from(&ids, 0)
19808 .expect_err("one-shot forward failure");
19809 assert!(err.contains("forward") || err.contains("score row"));
19810 assert!(!p.graph_failed.load(std::sync::atomic::Ordering::Relaxed));
19811 assert!(!p.graph_want_logits);
19812 assert!(p.graph_logits.is_none());
19813 assert!(p.kv_history.is_empty());
19814
19815 let mut fresh = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
19816 let expected = fresh.nll_ids_from(&ids, 0).expect("fresh NLL");
19817 let actual = p.nll_ids_from(&ids, 0).expect("reused NLL");
19818 assert_eq!(actual.1, expected.1);
19819 assert!((actual.0 - expected.0).abs() < 1e-9);
19820 }
19821
19822 #[test]
19823 fn nll_serial_failure_before_first_row_is_reported() {
19824 let ids = vec![1u32, 2, 3, 4];
19825 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
19826 p.nll_test_force_serial = true;
19827 p.nll_test_fail_at = Some(0);
19828 let err = p.nll_ids_from(&ids, 0).expect_err("serial forward failure");
19829 assert!(err.contains("serial forward"));
19830 assert!(p.kv_history.is_empty());
19831 assert!(!p.graph_failed.load(std::sync::atomic::Ordering::Relaxed));
19832 assert!(!p.cancel.load(std::sync::atomic::Ordering::Relaxed));
19833 }
19834
19835 #[test]
19836 fn ffn_probe_failure_discards_recorder_and_state() {
19837 let ids = vec![1u32, 2, 3, 4];
19838 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
19839 p.nll_test_fail_at = Some(0);
19840 let err = p
19841 .probe_ffn_mass_batch(&ids)
19842 .expect_err("probe forward failure");
19843 assert!(err.contains("NLL"));
19844 assert!(FFN_PROBE.with(|probe| probe.borrow().is_none()));
19845 assert!(p.kv_history.is_empty());
19846 assert!(!p.graph_failed.load(std::sync::atomic::Ordering::Relaxed));
19847 }
19848
19849 #[test]
19850 fn nll_test_controls_are_pipeline_scoped() {
19851 let ids = vec![1u32, 2, 3, 4];
19852 let mut failing = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
19853 let mut unaffected = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
19854 failing.nll_test_force_serial = true;
19855 failing.nll_test_fail_at = Some(0);
19856
19857 assert!(!failing.can_prefill_batched());
19858 assert!(unaffected.can_prefill_batched());
19859 let expected = unaffected
19860 .nll_ids_from(&ids, 0)
19861 .expect("unaffected pipeline remains usable");
19862 let err = failing
19863 .nll_ids_from(&ids, 0)
19864 .expect_err("failure injection belongs to failing pipeline");
19865 assert!(err.contains("serial forward"));
19866 assert!(failing.nll_test_fail_at.is_none());
19867 assert!(unaffected.can_prefill_batched());
19868 let actual = unaffected
19869 .nll_ids_from(&ids, 0)
19870 .expect("unaffected pipeline remains reusable");
19871 assert_eq!(actual.1, expected.1);
19872 assert!((actual.0 - expected.0).abs() < 1e-9);
19873 }
19874
19875 #[test]
19876 fn forward_ids_failure_channel_is_terminal_and_reusable() {
19877 let ids = vec![1u32, 2, 3, 4, 5, 6];
19878 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
19879 p.graph_logits = Some(vec![123.0]);
19880 p.graph_want_logits = true;
19881 p.graph_failed
19882 .store(true, std::sync::atomic::Ordering::Relaxed);
19883 p.cancel.store(true, std::sync::atomic::Ordering::Relaxed);
19884
19885 let err = p
19886 .forward_ids(&ids, None)
19887 .expect_err("a failed forward must not become a valid head result");
19888 assert!(err.contains("forward_ids setup"));
19889 assert!(p.graph_logits.is_none());
19890 assert!(!p.graph_want_logits);
19891 assert!(!p.graph_failed.load(std::sync::atomic::Ordering::Relaxed));
19892 assert!(!p.cancel.load(std::sync::atomic::Ordering::Relaxed));
19893 assert_eq!(p.kv_cache.seq_len(), 0);
19894
19895 let expected = create_test_pipeline(8, 16, 2, 1, 4, 1, 64)
19896 .forward_ids(&ids, None)
19897 .expect("fresh forward_ids");
19898 let actual = p
19899 .forward_ids(&ids, None)
19900 .expect("pipeline remains reusable after a failed forward");
19901 assert_eq!(actual.len(), expected.len());
19902 assert!(
19903 actual
19904 .iter()
19905 .zip(expected)
19906 .all(|(a, b)| (a - b).abs() < 1e-9)
19907 );
19908 assert_eq!(p.kv_cache.seq_len(), ids.len());
19909 }
19910
19911 #[test]
19912 fn sigmoid_router_floor_is_explicit_per_architecture() {
19913 let zero = || DenseFfn {
19919 gate_proj: QTensor::from_f32(vec![0.0; 4], 2, 2),
19920 up_proj: QTensor::from_f32(vec![0.0; 4], 2, 2),
19921 down_proj: QTensor::from_f32(vec![0.0; 4], 2, 2),
19922 act: Act::Silu,
19923 down_t: None,
19924 segs: Vec::new(),
19925 };
19926 let m = MoeFfn {
19927 router: QTensor::from_f32(vec![0.0; 4], 2, 2),
19928 experts: vec![zero(), zero()],
19929 top_k: 1,
19930 norm_topk_prob: true,
19931 router_sigmoid: true,
19932 expert_bias: None,
19933 routed_scaling: 2.5,
19934 route_tau: None,
19935 shared: None,
19936 stats: std::cell::RefCell::new(Vec::new()),
19937 act_sq: std::cell::RefCell::new(Vec::new()),
19938 act_rows: std::cell::RefCell::new(Vec::new()),
19939 mask: None,
19940 per_expert_scale: None,
19941 router_input_norm: false,
19942 resonance: None,
19943 grown: Vec::new(),
19944 };
19945 let logits = [-20.0f32, -20.0];
19946 let (_, p, glm_wsum) = moe_route_with_eps(&logits, &m, None, 1e-20);
19947 let (_, _, generic_wsum) = moe_route(&logits, &m, None);
19948 let expected = (p[0] + 1e-20) / m.routed_scaling;
19949 assert!((glm_wsum - expected).abs() < 1e-15);
19950 assert!(generic_wsum > glm_wsum * 100.0);
19951 }
19952
19953 #[test]
19954 fn resonance_scores_match_formula_and_stable_tie() {
19955 let r = Resonance {
19956 mu: vec![0.0, 0.0, 1.0, 0.0, 0.0, 1.0],
19958 u: vec![1.0, 0.0, 0.0, 1.0, 1.0, 1.0],
19959 k: 1,
19960 bias: vec![1.5, 0.5, 0.0],
19961 shell: Vec::new(),
19962 };
19963 let x = [1.0f32, 1.0];
19964 let mut got = vec![0.0; 3];
19965 r.scores(&x, &mut got);
19966 assert!((got[0] - 0.5).abs() < 1e-6);
19970 assert!((got[1] - 0.5).abs() < 1e-6);
19971 assert!(got[2].abs() < 1e-6);
19972 let best = got
19973 .iter()
19974 .enumerate()
19975 .max_by(|(ia, a), (ib, b)| a.partial_cmp(b).unwrap().then(ib.cmp(ia)))
19976 .map(|(i, _)| i);
19977 assert_eq!(best, Some(0));
19978 assert!(got.iter().all(|v| v.is_finite()));
19979 }
19980
19981 #[test]
19986 fn resonance_shell_masks_outside_keeps_inside_and_trunk_bits() {
19987 let plain = Resonance {
19993 mu: vec![0.0, 0.0, 1.0, 0.0, 0.0, 1.0, 0.0, 1.0],
19994 u: vec![1.0, 0.0, 0.0, 1.0, 1.0, 1.0, 1.0, 1.0],
19995 k: 1,
19996 bias: vec![1.5, 0.5, 0.0, 0.0],
19997 shell: Vec::new(),
19998 };
19999 let shelled = Resonance {
20000 mu: plain.mu.clone(),
20001 u: plain.u.clone(),
20002 k: 1,
20003 bias: plain.bias.clone(),
20004 shell: vec![f32::INFINITY, f32::INFINITY, 6.0, 0.25],
20005 };
20006 assert!(!plain.has_shell());
20007 assert!(shelled.has_shell());
20008 let bits = |v: &[f32]| v.iter().map(|x| x.to_bits()).collect::<Vec<_>>();
20009 let (mut a, mut b) = (vec![0.0; 4], vec![0.0; 4]);
20010 set_growth_shell(Some(true));
20011 assert!(growth_shell_enabled());
20012 let x = [1.0f32, 1.0];
20015 plain.scores(&x, &mut a);
20016 shelled.scores(&x, &mut b);
20017 assert_eq!(bits(&a), bits(&b), "inside every shell: unchanged");
20018 assert!(a[2] == 0.0 && a[3] == 0.0);
20019 let xo = [3.0f32, 0.0];
20022 plain.scores(&xo, &mut a);
20023 shelled.scores(&xo, &mut b);
20024 assert_eq!(a[2], -6.0);
20025 assert_eq!(a[3], -6.0);
20026 assert_eq!(bits(&a[..3]), bits(&b[..3]), "trunk rows + the boundary row");
20027 assert_eq!(b[3], f32::NEG_INFINITY, "outside the shell: −∞");
20028 assert_eq!(shelled.effective_shell(4), shelled.shell);
20029 set_growth_shell(Some(false));
20032 assert!(!growth_shell_enabled());
20033 assert_eq!(shelled.effective_shell(4), vec![f32::INFINITY; 4]);
20034 shelled.scores(&xo, &mut b);
20035 assert_eq!(bits(&a), bits(&b));
20036 set_growth_shell(None);
20037 let short = Resonance {
20040 shell: vec![f32::INFINITY, f32::INFINITY],
20041 ..shelled
20042 };
20043 set_growth_shell(Some(true));
20044 short.scores(&xo, &mut b);
20045 assert_eq!(bits(&a), bits(&b));
20046 set_growth_shell(None);
20047 }
20048
20049 #[test]
20053 fn moe_route_neg_inf_logits_pick_the_best_finite_expert_with_weight_one() {
20054 let zero = || DenseFfn {
20055 gate_proj: QTensor::from_f32(vec![0.0; 4], 2, 2),
20056 up_proj: QTensor::from_f32(vec![0.0; 4], 2, 2),
20057 down_proj: QTensor::from_f32(vec![0.0; 4], 2, 2),
20058 act: Act::Silu,
20059 down_t: None,
20060 segs: Vec::new(),
20061 };
20062 let moe = |sigmoid: bool| MoeFfn {
20063 router: QTensor::from_f32(vec![0.0; 8], 4, 2),
20064 experts: vec![zero(), zero(), zero(), zero()],
20065 top_k: 1,
20066 norm_topk_prob: true,
20067 router_sigmoid: sigmoid,
20068 expert_bias: None,
20069 routed_scaling: 1.0,
20070 route_tau: None,
20071 shared: None,
20072 stats: std::cell::RefCell::new(Vec::new()),
20073 act_sq: std::cell::RefCell::new(Vec::new()),
20074 act_rows: std::cell::RefCell::new(Vec::new()),
20075 mask: None,
20076 per_expert_scale: None,
20077 router_input_norm: false,
20078 resonance: None,
20079 grown: Vec::new(),
20080 };
20081 let logits = [-1.0f32, f32::NEG_INFINITY, -0.5, f32::NEG_INFINITY];
20082 for sigmoid in [false, true] {
20083 let m = moe(sigmoid);
20084 let (idx, p, wsum) = moe_route(&logits, &m, None);
20085 assert_eq!(idx, vec![2], "sigmoid {sigmoid}");
20086 assert_eq!(p[1], 0.0);
20087 assert_eq!(p[3], 0.0);
20088 assert!(p[2] > p[0] && p[0] > 0.0);
20089 assert!(p.iter().all(|v| v.is_finite()));
20090 let w = p[2] / wsum;
20091 if sigmoid {
20092 assert!((w - 1.0).abs() < 1e-5, "sigmoid: weight {w}");
20094 } else {
20095 assert_eq!(w, 1.0, "softmax: the renormalized top-1 weight is exactly 1");
20096 }
20097 let (idx, _, _) = moe_route(&logits, &m, Some(&[true, true, true, true]));
20100 assert_eq!(idx, vec![2]);
20101 }
20102 let m = moe(false);
20104 let (idx, _, _) = moe_route(&[f32::NEG_INFINITY, f32::NEG_INFINITY, f32::NEG_INFINITY, -9.0], &m, None);
20105 assert_eq!(idx, vec![3]);
20106 let (idx, p, wsum) = moe_route(&[f32::NEG_INFINITY; 4], &m, None);
20109 assert_eq!(idx, vec![0]);
20110 assert!(p.iter().all(|&v| v == 0.25));
20111 assert!(wsum.is_finite() && wsum > 0.0);
20112 }
20113
20114 #[test]
20119 fn resonance_route_picks_the_raw_argmax_not_the_softmax_tie() {
20120 let zero = || DenseFfn {
20121 gate_proj: QTensor::from_f32(vec![0.0; 4], 2, 2),
20122 up_proj: QTensor::from_f32(vec![0.0; 4], 2, 2),
20123 down_proj: QTensor::from_f32(vec![0.0; 4], 2, 2),
20124 act: Act::Silu,
20125 down_t: None,
20126 segs: Vec::new(),
20127 };
20128 let moe = |resonant: bool, norm_topk: bool| MoeFfn {
20129 router: QTensor::from_f32(vec![0.0; 6], 3, 2),
20130 experts: vec![zero(), zero(), zero()],
20131 top_k: 1,
20132 norm_topk_prob: norm_topk,
20133 router_sigmoid: false,
20134 expert_bias: None,
20135 routed_scaling: 1.0,
20136 route_tau: None,
20137 shared: None,
20138 stats: std::cell::RefCell::new(Vec::new()),
20139 act_sq: std::cell::RefCell::new(Vec::new()),
20140 act_rows: std::cell::RefCell::new(Vec::new()),
20141 mask: None,
20142 per_expert_scale: None,
20143 router_input_norm: false,
20144 resonance: resonant.then(|| Resonance {
20145 mu: vec![0.0; 6],
20146 u: Vec::new(),
20147 k: 0,
20148 bias: vec![0.0; 3],
20149 shell: Vec::new(),
20150 }),
20151 grown: Vec::new(),
20152 };
20153 let lo = -0.1f32;
20156 let hi = f32::from_bits(lo.to_bits() - 1);
20157 assert!(hi > lo && hi - lo < 2f32.powi(-25));
20158 assert_eq!((lo - hi).exp(), 1.0, "the tie this test is about");
20159 let (idx, _, _) = moe_route(&[lo, hi], &moe(false, true), None);
20161 assert_eq!(idx, vec![0], "softmax tie → lower index (the gated contract)");
20162 for norm in [true, false] {
20165 let m = moe(true, norm);
20166 let (idx, p, wsum) = moe_route(&[lo, hi, f32::NEG_INFINITY], &m, None);
20167 assert_eq!(idx, vec![1], "norm_topk {norm}");
20168 assert_eq!(p, vec![0.0, 1.0, 0.0]);
20169 assert_eq!(p[1] / wsum, 1.0);
20170 let (idx, _, _) = moe_route(&[hi, hi, lo], &m, None);
20173 assert_eq!(idx, vec![0]);
20174 let (idx, _, _) = moe_route(&[f32::NEG_INFINITY, lo, hi], &m, None);
20176 assert_eq!(idx, vec![2]);
20177 let (idx, p, wsum) = moe_route(&[lo, hi, hi], &m, Some(&[true, false, false]));
20178 assert_eq!(idx, vec![0]);
20179 assert_eq!(p[0] / wsum, 1.0);
20180 let (idx, p, wsum) = moe_route(&[f32::NEG_INFINITY; 3], &m, None);
20183 assert_eq!(idx, vec![0]);
20184 assert!(p.iter().all(|v| v.is_finite()) && wsum.is_finite() && wsum > 0.0);
20185 }
20186 }
20187}