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 swa_trim: Option<(usize, usize)>,
359 pub sliding_layers: Option<Vec<bool>>,
362 pub anchor_core: Option<cortiq_core::AnchorCoreConfig>,
368 bounded_rope: Option<std::sync::Arc<crate::bounded::BoundedRope>>,
371 pub kv_prefix: KvPrefix,
375 pub last_prefill_tokens: usize,
378 pub inv_freq_local: Option<std::sync::Arc<Vec<f32>>>,
381 pub rotary_dim_local: Option<usize>,
382 pub rope_scale: f32,
383 pub rope_scale_local: f32,
384 pub global_attn: Option<(usize, usize)>,
387 pub inv_freq_global: Option<std::sync::Arc<Vec<f32>>>,
390 pub attn_v_norm: bool,
392 pub qk_norm_after_rope: bool,
394 pub proj_gate_sigmoid: bool,
397 pub final_softcap: Option<f32>,
399 pub head_clusters: Option<std::sync::Arc<Vec<f32>>>,
403 pub attn_softcap: f32,
405 confidence_on: bool,
409 #[cfg(test)]
412 nll_test_fail_at: Option<usize>,
413 #[cfg(test)]
416 nll_test_force_serial: bool,
417}
418
419#[cfg(target_os = "macos")]
420impl Drop for Pipeline {
421 fn drop(&mut self) {
422 let _ = crate::gpu_metal::wait_replay();
424 crate::gpu::kv_mirror_drop(self.graph_kv_id);
425 }
426}
427
428#[cfg(not(target_os = "macos"))]
429impl Drop for Pipeline {
430 fn drop(&mut self) {
431 crate::gpu::graph_kv_reset(self.graph_kv_id);
435 }
436}
437
438pub struct PipelineWeights {
443 pub embed_tokens: QTensor,
445 pub layers: Vec<LayerWeights>,
447 pub lm_head: QTensor,
449 pub final_norm: Vec<f32>,
451}
452
453pub struct LayerWeights {
455 pub input_norm: Vec<f32>,
456 pub post_norm: Vec<f32>,
459 pub attn_out_norm: Option<Vec<f32>>,
462 pub layer_scale: Option<f32>,
464 pub ffn_out_norm: Option<Vec<f32>>,
467 pub ffn: FfnKind,
468 pub attn: AttnKind,
469}
470
471#[derive(Clone, Copy, PartialEq, Debug, Default)]
474pub enum Act {
475 #[default]
476 Silu,
477 GeluTanh,
478 Gelu,
480 Situ {
483 beta: f32,
484 linear_beta: f32,
485 },
486}
487
488impl Act {
489 pub fn from_arch(name: &str) -> Self {
490 if name == "gelu_tanh" {
491 Self::GeluTanh
492 } else if name == "gelu" {
493 Self::Gelu
494 } else {
495 Self::Silu
496 }
497 }
498
499 pub fn from_arch_full(arch: &cortiq_core::ModelArch) -> Self {
501 match arch.hidden_act.as_str() {
502 "situ" => Self::Situ {
503 beta: arch.activation_situ_beta.unwrap_or(1.0) as f32,
504 linear_beta: arch.activation_situ_linear_beta.unwrap_or(0.0) as f32,
505 },
506 other => Self::from_arch(other),
507 }
508 }
509
510 #[inline]
511 pub fn apply(self, x: f32) -> f32 {
512 match self {
513 Self::Silu => inference::silu(x),
514 Self::GeluTanh => inference::gelu_tanh(x),
515 Self::Gelu => inference::gelu_erf(x),
516 Self::Situ { beta, .. } => beta * (x / beta).tanh() * (1.0 / (1.0 + (-x).exp())),
517 }
518 }
519
520 #[inline]
523 pub fn combine(self, g: f32, u: f32) -> f32 {
524 match self {
525 Self::Situ { linear_beta, .. } if linear_beta > 0.0 => {
526 self.apply(g) * (linear_beta * (u / linear_beta).tanh())
527 }
528 _ => self.apply(g) * u,
529 }
530 }
531
532 pub fn graph_act(self) -> Option<crate::gpu::GraphAct> {
536 match self {
537 Self::Silu => Some(crate::gpu::GraphAct::Silu),
538 Self::Gelu => Some(crate::gpu::GraphAct::GeluErf),
539 Self::GeluTanh | Self::Situ { .. } => None,
540 }
541 }
542}
543
544pub struct DenseFfn {
546 pub gate_proj: QTensor,
547 pub up_proj: QTensor,
548 pub down_proj: QTensor,
549 pub act: Act,
551 pub down_t: Option<QTensor>,
557 pub segs: Vec<FfnSeg>,
564}
565
566pub struct FfnSeg {
571 pub gate: QTensor,
572 pub up: QTensor,
573 pub down: QTensor,
574 pub start: usize,
575 pub width: usize,
576}
577
578pub enum FfnKind {
581 Dense(DenseFfn),
582 Moe(MoeFfn),
586 DenseMoe(Box<DenseMoeFfn>),
593}
594
595pub struct DenseMoeFfn {
597 pub dense: DenseFfn,
598 pub moe: MoeFfn,
599 pub post_norm_1: Vec<f32>,
601 pub pre_norm_2: Vec<f32>,
604 pub post_norm_2: Vec<f32>,
606}
607
608pub struct MoeFfn {
609 pub router: QTensor,
611 pub experts: Vec<DenseFfn>,
612 pub top_k: usize,
613 pub norm_topk_prob: bool,
614 pub router_sigmoid: bool,
617 pub expert_bias: Option<Vec<f32>>,
621 pub routed_scaling: f32,
624 pub route_tau: Option<f32>,
630 pub shared: Option<(DenseFfn, Option<QTensor>)>,
633 pub stats: std::cell::RefCell<Vec<u64>>,
637 pub act_sq: std::cell::RefCell<Vec<f64>>,
644 pub act_rows: std::cell::RefCell<Vec<f32>>,
650 pub mask: Option<Vec<bool>>,
655 pub per_expert_scale: Option<Vec<f32>>,
658 pub router_input_norm: bool,
662 pub resonance: Option<Resonance>,
666 pub grown: Vec<GrownExpert>,
672}
673
674#[derive(Debug, Clone, PartialEq, Eq)]
676pub struct GrownExpert {
677 pub record: String,
679 pub record_index: usize,
681 pub layer: usize,
682 pub expert: usize,
686}
687
688pub struct Resonance {
690 pub mu: Vec<f32>,
692 pub u: Vec<f32>,
694 pub k: usize,
695 pub bias: Vec<f32>,
697 pub shell: Vec<f32>,
703}
704
705static GROWTH_SHELL: std::sync::atomic::AtomicU8 = std::sync::atomic::AtomicU8::new(0);
708
709pub fn growth_shell_enabled() -> bool {
714 use std::sync::atomic::Ordering;
715 match GROWTH_SHELL.load(Ordering::Relaxed) {
716 1 => true,
717 2 => false,
718 _ => {
719 let off = std::env::var("CMF_GROWTH_SHELL")
720 .map(|v| v.eq_ignore_ascii_case("off") || v == "0")
721 .unwrap_or(false);
722 GROWTH_SHELL.store(if off { 2 } else { 1 }, Ordering::Relaxed);
723 !off
724 }
725 }
726}
727
728pub fn set_growth_shell(on: Option<bool>) {
733 GROWTH_SHELL.store(
734 match on {
735 Some(true) => 1,
736 Some(false) => 2,
737 None => 0,
738 },
739 std::sync::atomic::Ordering::Relaxed,
740 );
741}
742
743impl Resonance {
744 pub fn has_shell(&self) -> bool {
746 self.shell.iter().any(|s| s.is_finite())
747 }
748
749 pub fn effective_shell(&self, ne: usize) -> Vec<f32> {
752 let mut out = vec![f32::INFINITY; ne];
753 if growth_shell_enabled() {
754 for (o, s) in out.iter_mut().zip(&self.shell) {
755 *o = *s;
756 }
757 }
758 out
759 }
760
761 pub fn scores(&self, x: &[f32], out: &mut [f32]) {
766 let h = x.len();
767 let ne = out.len();
768 let shell_on = growth_shell_enabled() && !self.shell.is_empty();
769 for e in 0..ne {
770 let mu = &self.mu[e * h..(e + 1) * h];
771 let mut d2 = 0.0f32;
772 for j in 0..h {
773 let d = x[j] - mu[j];
774 d2 += d * d;
775 }
776 let mut proj = 0.0f32;
777 for i in 0..self.k {
778 let u = &self.u[(e * self.k + i) * h..(e * self.k + i + 1) * h];
779 let mut p = 0.0f32;
780 for j in 0..h {
781 p += (x[j] - mu[j]) * u[j];
782 }
783 proj += p * p;
784 }
785 let err = d2 - proj;
786 out[e] = self.bias.get(e).copied().unwrap_or(0.0) - err;
787 if shell_on && err > self.shell.get(e).copied().unwrap_or(f32::INFINITY) {
788 out[e] = f32::NEG_INFINITY;
789 }
790 }
791 }
792}
793
794pub enum AttnKind {
797 Full {
799 wq: QTensor,
800 wk: QTensor,
801 wv: QTensor,
802 wo: QTensor,
803 q_norm: Option<Vec<f32>>,
804 k_norm: Option<Vec<f32>>,
805 output_gate: bool,
806 softplus_gate: Option<(QTensor, bool)>,
810 bias: Option<(Vec<f32>, Vec<f32>, Vec<f32>)>,
812 },
813 Linear(VmfPhaseWeights),
815 LinearGdn(GdnWeights),
817 ShortConv(ShortConvWeights),
820 Mla(Box<MlaWeights>),
828 Kda(Box<crate::linear_core::KdaWeights>),
832 Bounded(Box<crate::bounded::BoundedWeights>),
837}
838
839pub struct MlaWeights {
841 pub q_proj: QTensor,
845 pub q_a: Option<QTensor>,
848 pub q_a_norm: Option<Vec<f32>>,
849 pub kv_a: QTensor,
851 pub kv_a_norm: Vec<f32>,
853 pub kv_b: QTensor,
855 pub o_proj: QTensor,
857 pub nh: usize,
858 pub qk_rope: usize,
859 pub qk_nope: usize,
860 pub v_dim: usize,
861 pub lora: usize,
862 pub scale: f32,
864 pub nope: bool,
866}
867
868pub struct MtpModule {
873 pub enorm: Vec<f32>,
874 pub hnorm: Vec<f32>,
875 pub eh_proj: QTensor,
877 pub layer: LayerWeights,
878 pub final_norm: Vec<f32>,
879 pub kv: crate::kv_cache::LayerKvCache,
880}
881
882#[cfg(target_os = "macos")]
889enum MetalRowsItem<'a> {
890 Gdn {
891 run: Vec<crate::gpu_metal::GdnGpuLayer<'a>>,
892 first: usize,
893 },
894 Attn {
895 l: crate::gpu_metal::AttnGpuLayer<'a>,
896 li: usize,
897 q_norm: Option<&'a [f32]>,
898 k_norm: Option<&'a [f32]>,
899 output_gate: bool,
900 },
901}
902
903#[cfg(target_os = "macos")]
904struct MetalVerifyPending {
905 graph: crate::gpu_metal::VerifyGraph,
906 gdn_layers: Vec<usize>,
907 attn_layers: Vec<(usize, usize)>,
908}
909
910#[cfg(target_os = "macos")]
914struct MetalWarmPending {
915 graph: crate::gpu_metal::VerifyGraph,
916 cpu_stored: usize,
917 b: usize,
918}
919
920#[cfg(target_os = "macos")]
921enum MetalRowsRun {
922 Declined,
924 Failed,
927 Completed(MetalVerifyPending),
928}
929
930#[cfg(target_os = "macos")]
931enum MetalPrefillOutcome {
932 Declined,
933 Failed,
934 Completed(Vec<f32>),
935}
936
937#[cfg(target_os = "macos")]
938enum MetalBatchNllOutcome {
939 Declined,
940 Failed(String),
941 Completed(f64, usize),
942}
943
944#[derive(Clone, Copy)]
948enum SpecTrial {
949 Spec {
950 t0: std::time::Instant,
951 gen0: usize,
952 rounds: usize,
953 },
954 Plain {
955 t0: std::time::Instant,
956 gen0: usize,
957 },
958 Decided {
959 spec: bool,
960 recheck_at: usize,
961 },
962}
963
964pub(crate) fn spec_time_level() -> u8 {
968 static L: std::sync::OnceLock<u8> = std::sync::OnceLock::new();
969 *L.get_or_init(|| match std::env::var("CMF_GRAPH_SPEC_TIME") {
970 Ok(v) => v.trim().parse::<u8>().map(|n| n.max(1)).unwrap_or(1),
971 Err(_) => 0,
972 })
973}
974
975struct SpecStampLog {
981 t_last: std::time::Instant,
982 items: Vec<(&'static str, f32)>,
983}
984
985static SPEC_STAMPS: std::sync::Mutex<Option<SpecStampLog>> = std::sync::Mutex::new(None);
986
987pub(crate) fn spec_stamp(name: &'static str) {
988 if spec_time_level() == 0 {
989 return;
990 }
991 if let Ok(mut g) = SPEC_STAMPS.lock() {
992 if let Some(log) = g.as_mut() {
993 let now = std::time::Instant::now();
994 log.items
995 .push((name, (now - log.t_last).as_secs_f32() * 1e3));
996 log.t_last = now;
997 }
998 }
999}
1000
1001fn spec_stamps_begin() {
1002 if spec_time_level() == 0 {
1003 return;
1004 }
1005 if let Ok(mut g) = SPEC_STAMPS.lock() {
1006 *g = Some(SpecStampLog {
1007 t_last: std::time::Instant::now(),
1008 items: Vec::with_capacity(64),
1009 });
1010 }
1011}
1012
1013fn spec_stamps_take() -> Vec<(&'static str, f32)> {
1014 SPEC_STAMPS
1015 .lock()
1016 .ok()
1017 .and_then(|mut g| g.take())
1018 .map(|l| l.items)
1019 .unwrap_or_default()
1020}
1021
1022fn spec_stamps_format(items: &[(&'static str, f32)]) -> String {
1025 let mut agg: Vec<(&'static str, f32, u32)> = Vec::with_capacity(items.len());
1026 for &(n, ms) in items {
1027 match agg.iter_mut().find(|e| e.0 == n) {
1028 Some(e) => {
1029 e.1 += ms;
1030 e.2 += 1;
1031 }
1032 None => agg.push((n, ms, 1)),
1033 }
1034 }
1035 let mut s = String::with_capacity(agg.len() * 16);
1036 for (n, ms, k) in agg {
1037 if k > 1 {
1038 s.push_str(&format!("{n} {ms:.1}/{k} "));
1039 } else {
1040 s.push_str(&format!("{n} {ms:.1} "));
1041 }
1042 }
1043 s
1044}
1045
1046#[derive(Default, Clone, Copy)]
1068struct SpecMon {
1069 round_ms: f64,
1070 tokens: f64,
1071 plain_ms: f64,
1072 n: u32,
1073 fails: u32,
1074 metal: bool,
1075}
1076
1077const SPEC_PROXY_TOKENS: f64 = 3.5;
1080const SPEC_PLAIN_MIN_MS: f64 = 200.0;
1083
1084impl SpecMon {
1085 fn round(&mut self, dt_ms: f64, produced: usize) {
1086 self.n += 1;
1087 if self.n == 1 {
1088 return; }
1090 let a = if self.n == 2 { 1.0 } else { 0.3 };
1091 self.round_ms += a * (dt_ms - self.round_ms);
1092 self.tokens += a * (produced as f64 - self.tokens);
1093 }
1094 fn pays(&self) -> bool {
1095 if self.plain_ms > 0.0 {
1096 self.tokens * self.plain_ms > self.round_ms * 1.03
1097 } else {
1098 self.metal && self.tokens >= SPEC_PROXY_TOKENS
1099 }
1100 }
1101 fn plain_done(&self, t0: std::time::Instant, gen0: usize, generated: usize) -> bool {
1103 let n = generated.saturating_sub(gen0);
1104 if n >= 8 {
1105 return true;
1106 }
1107 self.metal && n >= 2 && t0.elapsed().as_secs_f64() * 1e3 >= SPEC_PLAIN_MIN_MS
1108 }
1109}
1110
1111pub const KV_PREFIX_TAIL: usize = 128;
1114
1115#[derive(Debug, Clone, Default)]
1121pub struct KvPrefix {
1122 len: usize,
1123 hash: u64,
1124 tail: Vec<u32>,
1125 device: bool,
1130}
1131
1132impl KvPrefix {
1133 #[inline]
1134 fn fold(mut h: u64, ids: &[u32]) -> u64 {
1135 for &id in ids {
1136 h ^= id as u64;
1137 h = h.wrapping_mul(0x100000001b3);
1138 h ^= h >> 29;
1139 }
1140 h
1141 }
1142
1143 pub fn clear(&mut self) {
1144 self.len = 0;
1145 self.hash = 0xcbf29ce484222325;
1146 self.tail.clear();
1147 self.device = false;
1148 }
1149
1150 pub fn on_device(&self) -> bool {
1152 self.device
1153 }
1154
1155 pub fn set_on_device(&mut self, device: bool) {
1157 self.device = device;
1158 }
1159
1160 pub fn len(&self) -> usize {
1162 self.len
1163 }
1164
1165 pub fn is_empty(&self) -> bool {
1166 self.len == 0
1167 }
1168
1169 pub fn tail_len(&self) -> usize {
1171 self.tail.len()
1172 }
1173
1174 pub fn set(&mut self, ids: &[u32]) {
1176 self.clear();
1177 self.extend(ids);
1178 }
1179
1180 pub fn extend(&mut self, more: &[u32]) {
1182 if self.len == 0 && self.hash == 0 {
1183 self.hash = 0xcbf29ce484222325;
1184 }
1185 self.hash = Self::fold(self.hash, more);
1186 self.len += more.len();
1187 if more.len() >= KV_PREFIX_TAIL {
1188 self.tail.clear();
1189 self.tail.extend_from_slice(&more[more.len() - KV_PREFIX_TAIL..]);
1190 } else {
1191 let drop = (self.tail.len() + more.len()).saturating_sub(KV_PREFIX_TAIL);
1192 self.tail.drain(..drop);
1193 self.tail.extend_from_slice(more);
1194 }
1195 }
1196
1197 pub fn extension(&self, ids: &[u32]) -> usize {
1201 if self.len == 0 || ids.len() <= self.len {
1202 return 0;
1203 }
1204 let t = self.tail.len();
1205 if ids[self.len - t..self.len] != self.tail[..] {
1206 return 0;
1207 }
1208 if Self::fold(0xcbf29ce484222325, &ids[..self.len]) != self.hash {
1209 return 0;
1210 }
1211 self.len
1212 }
1213}
1214
1215pub struct GenerateResult {
1217 pub text: String,
1218 pub token_ids: Vec<u32>,
1219 pub prompt_tokens: usize,
1220 pub tokens_generated: usize,
1221 pub finish_reason: String,
1222 pub mtp_drafted: usize,
1224 pub mtp_accepted: usize,
1225 pub token_confidence: Vec<f32>,
1230 pub traces: Vec<TokenTrace>,
1233}
1234
1235#[derive(Clone, Debug)]
1240pub struct TokenTrace {
1241 pub t: usize,
1243 pub token_id: u32,
1245 pub confidence: f32,
1247 pub active_skill: Option<String>,
1249 pub recon: Option<f32>,
1253 pub switched: bool,
1256}
1257
1258#[cfg_attr(not(test), allow(dead_code))]
1263fn top1_prob_t(logits: &[f32], id: u32, temp: f32) -> f32 {
1264 let t = if temp > 1e-3 { temp } else { 1.0 };
1265 let max = logits.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
1266 let sum: f32 = logits.iter().map(|&v| ((v - max) / t).exp()).sum();
1267 if sum > 0.0 {
1268 (((logits[id as usize] - max) / t).exp()) / sum
1269 } else {
1270 0.0
1271 }
1272}
1273
1274fn prefill_batched() -> bool {
1277 std::env::var("CMF_PREFILL")
1278 .map(|v| v != "seq")
1279 .unwrap_or(true)
1280}
1281
1282#[inline]
1286fn nll_graph_policy(
1287 unmasked: bool,
1288 prefer_graph: bool,
1289 native_metal: bool,
1290) -> (bool, bool) {
1291 let graph_quality = unmasked && prefer_graph;
1292 let fused_head_quality = graph_quality && native_metal;
1293 (graph_quality, fused_head_quality)
1294}
1295
1296#[derive(Clone, Copy)]
1300enum PrefillIn<'a> {
1301 Ids(&'a [u32]),
1302 Hidden(&'a [f32]),
1303}
1304
1305impl Pipeline {
1312 fn can_prefill_batched(&self) -> bool {
1313 #[cfg(test)]
1314 let force_serial = self.nll_test_force_serial;
1315 #[cfg(not(test))]
1316 let force_serial = false;
1317 prefill_batched() && !force_serial && !self.weights.layers.is_empty()
1318 }
1319
1320 fn automatic_gpu_prefix(&self) -> Option<usize> {
1323 let (model, _, _, _) = self.weights.embed_tokens.graph_weight()?;
1324 crate::gpu::automatic_layer_prefix(&model, self.num_layers, self.physical_layers)
1325 }
1326
1327 pub fn prefill_chunk(&self) -> usize {
1331 let env = env_prefill_chunk();
1332 if env.is_some() || ChunkHost::here() != ChunkHost::Other {
1333 return prefill_chunk_rule(env, ChunkHost::here(), false);
1334 }
1335 prefill_chunk_rule(None, ChunkHost::Other, self.chunk_stack_facts().dense_on_discrete())
1336 }
1337
1338 fn chunk_stack_facts(&self) -> ChunkStackFacts {
1339 let plain_dense = !self.weights.layers.is_empty()
1340 && self.g3n.is_none()
1341 && self.dsv4.is_none()
1342 && self.dsv41.is_none()
1343 && self.qwen4_exp.is_none()
1344 && self.weights.layers.iter().all(|lw| {
1345 matches!(lw.attn, AttnKind::Full { .. }) && matches!(lw.ffn, FfnKind::Dense(_))
1346 });
1347 let gpu_on = crate::gpu::enabled();
1348 ChunkStackFacts {
1349 plain_dense,
1350 discrete: gpu_on && crate::gpu::discrete(),
1351 gpu_on,
1352 capacity_split: std::env::var_os("CMF_GPU_LAYERS").is_some()
1355 || (plain_dense && gpu_on && self.automatic_gpu_prefix().is_some()),
1356 multi_gpu: self.gpu_plan.is_some(),
1357 o1: self.o1_active(),
1358 }
1359 }
1360}
1361
1362pub fn prefill_chunk() -> usize {
1371 prefill_chunk_rule(env_prefill_chunk(), ChunkHost::here(), false)
1372}
1373
1374fn env_prefill_chunk() -> Option<usize> {
1375 std::env::var("CMF_PREFILL_CHUNK")
1376 .ok()
1377 .and_then(|v| v.parse::<usize>().ok())
1378}
1379
1380#[derive(Clone, Copy, Debug, PartialEq, Eq)]
1382enum ChunkHost {
1383 Macos,
1384 Aarch64,
1386 Other,
1388}
1389
1390impl ChunkHost {
1391 fn here() -> Self {
1392 if cfg!(target_os = "macos") {
1393 ChunkHost::Macos
1394 } else if cfg!(target_arch = "aarch64") {
1395 ChunkHost::Aarch64
1396 } else {
1397 ChunkHost::Other
1398 }
1399 }
1400}
1401
1402const DISCRETE_DENSE_PREFILL_CHUNK: usize = 512;
1409
1410fn prefill_chunk_rule(env: Option<usize>, host: ChunkHost, dense_on_discrete: bool) -> usize {
1416 if let Some(n) = env {
1417 return n.max(1);
1418 }
1419 match host {
1420 ChunkHost::Macos => 512,
1421 ChunkHost::Aarch64 => 256,
1424 ChunkHost::Other if dense_on_discrete => DISCRETE_DENSE_PREFILL_CHUNK,
1425 ChunkHost::Other => 48,
1426 }
1427}
1428
1429#[derive(Clone, Copy, Debug, Default)]
1431struct ChunkStackFacts {
1432 plain_dense: bool,
1435 discrete: bool,
1437 gpu_on: bool,
1439 capacity_split: bool,
1441 multi_gpu: bool,
1443 o1: bool,
1445}
1446
1447impl ChunkStackFacts {
1448 fn dense_on_discrete(self) -> bool {
1449 self.plain_dense
1450 && self.discrete
1451 && self.gpu_on
1452 && !self.capacity_split
1453 && !self.multi_gpu
1454 && !self.o1
1455 }
1456}
1457
1458#[inline]
1464fn mtp_prefill_pair_count(start: usize, end: usize, input_len: usize) -> usize {
1465 if end <= start || start >= input_len {
1466 return 0;
1467 }
1468 let rows = (end.min(input_len) - start).min(input_len - start);
1469 if end < input_len {
1470 rows
1471 } else {
1472 rows.saturating_sub(1)
1473 }
1474}
1475
1476pub type TokenCallback = Box<dyn FnMut(&str) -> bool + Send>;
1478
1479#[derive(Debug, Clone, Copy, PartialEq, Eq)]
1481pub(crate) struct ReuseLayer {
1482 pub full: bool,
1485 pub host_rows: usize,
1488 pub device_rows: Option<usize>,
1490 pub device_state: bool,
1492}
1493
1494#[derive(Debug, Clone, PartialEq, Eq)]
1496pub(crate) enum ReusePlan {
1497 Ready,
1499 Pull(Vec<(usize, usize, usize)>),
1502 Fresh,
1504}
1505
1506pub(crate) fn kv_reuse_plan(reuse_from: usize, layers: &[ReuseLayer]) -> ReusePlan {
1516 let mut pulls = Vec::new();
1517 for (li, l) in layers.iter().enumerate() {
1518 if !l.full {
1519 if l.device_state {
1520 return ReusePlan::Fresh;
1521 }
1522 continue;
1523 }
1524 if l.host_rows == reuse_from {
1525 continue;
1526 }
1527 if l.host_rows < reuse_from && l.device_rows.is_some_and(|d| d >= reuse_from) {
1528 pulls.push((li, l.host_rows, reuse_from));
1529 continue;
1530 }
1531 return ReusePlan::Fresh;
1532 }
1533 if pulls.is_empty() {
1534 ReusePlan::Ready
1535 } else {
1536 ReusePlan::Pull(pulls)
1537 }
1538}
1539
1540impl Pipeline {
1541 fn clear_sequence_state(&mut self) {
1549 #[cfg(target_os = "macos")]
1552 let _ = crate::gpu_metal::wait_replay();
1553 self.kv_cache.clear();
1554 self.clear_history();
1557 self.graph_logits = None;
1558 if let Some(b) = &mut self.dsv41 {
1559 b.3.clear();
1560 }
1561 crate::gpu::graph_kv_reset(self.graph_kv_id);
1562 crate::gpu::graph_kv_reset(self.mtp_kv_id());
1567 }
1568
1569 fn prepare_kv_reuse(&mut self, reuse_from: usize) -> bool {
1576 if self.graph_prefill_preferred() {
1577 return true;
1578 }
1579 let kv_id = self.graph_kv_id;
1580 let layers: Vec<ReuseLayer> = (0..self.num_layers)
1581 .map(|li| {
1582 let full = matches!(
1583 self.weights.layers[self.phys_layer(li)].attn,
1584 AttnKind::Full { .. }
1585 );
1586 ReuseLayer {
1587 full,
1588 host_rows: self.kv_cache.layers[li].pos_len(),
1591 device_rows: crate::gpu::graph_kv_stored(kv_id, li),
1592 device_state: crate::gpu::graph_state_resident(kv_id, li),
1593 }
1594 })
1595 .collect();
1596 if layers
1600 .iter()
1601 .all(|l| l.device_rows.is_none() && !l.device_state)
1602 {
1603 return true;
1604 }
1605 let plan = kv_reuse_plan(reuse_from, &layers);
1606 let (what, rows, n) = match &plan {
1607 ReusePlan::Ready => ("host ready", 0, 0),
1608 ReusePlan::Fresh => ("fresh", 0, 0),
1609 ReusePlan::Pull(p) => (
1610 "pull",
1611 p.iter().map(|&(_, a, b)| b - a).max().unwrap_or(0),
1612 p.len(),
1613 ),
1614 };
1615 let t0 = std::time::Instant::now();
1616 let ok = self.apply_kv_reuse_plan(reuse_from, plan, &layers);
1617 if std::env::var("CMF_PREFILL_PROF").is_ok() {
1618 eprintln!(
1619 "kv-reuse: {what}{}: {rows} device row(s) × {n} layer(s) to the host in {:.2} ms",
1620 if ok { "" } else { " (failed → fresh)" },
1621 t0.elapsed().as_secs_f64() * 1e3
1622 );
1623 }
1624 ok
1625 }
1626
1627 fn apply_kv_reuse_plan(
1628 &mut self,
1629 reuse_from: usize,
1630 plan: ReusePlan,
1631 layers: &[ReuseLayer],
1632 ) -> bool {
1633 let kv_id = self.graph_kv_id;
1634 match plan {
1635 ReusePlan::Fresh => return false,
1636 ReusePlan::Ready => {}
1637 ReusePlan::Pull(pulls) => {
1638 let (nkv, hd) = {
1643 let c = &self.kv_cache.layers[pulls[0].0];
1644 (c.num_kv_heads, c.head_dim)
1645 };
1646 let uniform = pulls.iter().all(|&(li, _, _)| {
1647 let c = &self.kv_cache.layers[li];
1648 (c.num_kv_heads, c.head_dim) == (nkv, hd)
1649 });
1650 let batched = if uniform {
1651 crate::gpu::graph_kv_read_rows(kv_id, &pulls, nkv, hd)
1652 } else {
1653 None
1654 };
1655 let rows: Vec<(Vec<f32>, Vec<f32>)> = match batched {
1656 Some(rows) => rows,
1657 None => {
1658 let mut rows = Vec::with_capacity(pulls.len());
1659 for &(li, from, to) in &pulls {
1660 let (lnkv, lhd) = {
1661 let c = &self.kv_cache.layers[li];
1662 (c.num_kv_heads, c.head_dim)
1663 };
1664 let Some((k, v, first_valid)) =
1665 crate::gpu::graph_kv_pull_host(kv_id, li, from, to, lnkv, lhd)
1666 else {
1667 return false;
1668 };
1669 let need_from = match self.layer_window(li) {
1673 Some(w) => from.max((to + 1).saturating_sub(w)),
1674 None => from,
1675 };
1676 if first_valid > need_from {
1677 return false;
1678 }
1679 rows.push((k, v));
1680 }
1681 rows
1682 }
1683 };
1684 for ((li, from, to), (k, v)) in pulls.into_iter().zip(rows) {
1685 let cache = &mut self.kv_cache.layers[li];
1686 let row = cache.num_kv_heads * cache.head_dim;
1687 for p in 0..to - from {
1688 cache.append(&k[p * row..(p + 1) * row], &v[p * row..(p + 1) * row], &[]);
1689 }
1690 if cache.pos_len() != to {
1691 return false;
1692 }
1693 }
1694 }
1695 }
1696 for (li, l) in layers.iter().enumerate() {
1700 if l.full
1701 && l.device_rows.is_some_and(|d| d > reuse_from)
1702 && !crate::gpu::graph_kv_set_stored(kv_id, li, reuse_from)
1703 {
1704 return false;
1705 }
1706 }
1707 true
1708 }
1709
1710 fn finish_generation(
1716 &mut self,
1717 mtp: &mut Option<MtpModule>,
1718 router: &mut Option<crate::swarm::DynRouter>,
1719 clear_sequence: bool,
1720 ) {
1721 if router.is_some() {
1725 let _ = self.set_active_skill(None);
1726 }
1727 #[cfg(target_os = "macos")]
1734 let clear_sequence = clear_sequence || !crate::gpu_metal::wait_replay();
1735 if clear_sequence {
1736 self.clear_sequence_state();
1737 if let Some(m) = mtp.as_mut() {
1738 m.kv.clear();
1744 }
1745 if let Some(m) = self.mtp.as_mut() {
1746 m.kv.clear();
1750 }
1751 }
1752 self.graph_want_logits = false;
1753 self.graph_head_required = false;
1754 self.graph_logits = None;
1755 self.graph_failed
1756 .store(false, std::sync::atomic::Ordering::Relaxed);
1757 self.cancel
1758 .store(false, std::sync::atomic::Ordering::Relaxed);
1759 self.dyn_router = router.take().or(self.dyn_router.take());
1760 self.mtp = mtp.take().or(self.mtp.take());
1761 self.mtp_graph_mode = None;
1762 self.spec_forced = None;
1763 }
1764
1765 fn check_forward_graph(&mut self, phase: &str, pos: usize) -> Result<(), String> {
1769 if self
1770 .graph_failed
1771 .swap(false, std::sync::atomic::Ordering::Relaxed)
1772 {
1773 self.cancel
1774 .store(false, std::sync::atomic::Ordering::Relaxed);
1775 self.clear_sequence_state();
1776 self.graph_logits = None;
1777 self.graph_want_logits = false;
1778 self.graph_head_required = false;
1779 return Err(format!("GPU graph failed during {phase} at position {pos}"));
1780 }
1781 Ok(())
1782 }
1783
1784 #[cfg(target_os = "macos")]
1785 fn fail_metal_graph(&mut self, reason: &str) {
1786 crate::pipeline::METAL_GRAPH_ERRORS
1787 .fetch_add(1, std::sync::atomic::Ordering::Relaxed);
1788 self.clear_sequence_state();
1789 self.graph_logits = None;
1790 self.graph_failed
1791 .store(true, std::sync::atomic::Ordering::Relaxed);
1792 self.cancel
1793 .store(true, std::sync::atomic::Ordering::Relaxed);
1794 tracing::error!("native Metal TokenGraph failed closed: {reason}");
1795 }
1796
1797 fn nll_begin(&mut self) -> Result<(), String> {
1802 if self
1803 .graph_failed
1804 .swap(false, std::sync::atomic::Ordering::Relaxed)
1805 {
1806 self.cancel
1807 .store(false, std::sync::atomic::Ordering::Relaxed);
1808 self.clear_sequence_state();
1809 self.graph_logits = None;
1810 self.graph_want_logits = false;
1811 self.graph_head_required = false;
1812 return Err("GPU graph failed before NLL scoring".to_string());
1813 }
1814 self.clear_sequence_state();
1815 self.graph_logits = None;
1816 self.graph_want_logits = false;
1817 self.graph_head_required = false;
1818 Ok(())
1819 }
1820
1821 fn nll_end(&mut self) {
1825 self.clear_sequence_state();
1826 self.graph_logits = None;
1827 self.graph_want_logits = false;
1828 self.graph_head_required = false;
1829 self.graph_failed
1830 .store(false, std::sync::atomic::Ordering::Relaxed);
1831 }
1832
1833 fn nll_check_graph(&mut self, phase: &str, pos: usize) -> Result<(), String> {
1836 #[cfg(test)]
1837 if self.nll_test_fail_at == Some(pos) {
1838 self.nll_test_fail_at = None;
1839 self.graph_failed
1840 .store(true, std::sync::atomic::Ordering::Relaxed);
1841 self.cancel
1842 .store(true, std::sync::atomic::Ordering::Relaxed);
1843 }
1844 if self
1845 .graph_failed
1846 .swap(false, std::sync::atomic::Ordering::Relaxed)
1847 {
1848 self.cancel
1849 .store(false, std::sync::atomic::Ordering::Relaxed);
1850 self.clear_sequence_state();
1851 self.graph_logits = None;
1852 self.graph_want_logits = false;
1853 return Err(format!(
1854 "GPU graph failed during NLL {phase} at position {pos}"
1855 ));
1856 }
1857 Ok(())
1858 }
1859
1860 #[inline]
1864 pub fn phys_layer(&self, virtual_idx: usize) -> usize {
1865 virtual_idx % self.physical_layers
1866 }
1867
1868 #[inline]
1871 pub fn is_loop_end(&self, virtual_idx: usize) -> bool {
1872 self.loop_final_norm && (virtual_idx + 1) % self.physical_layers == 0
1873 }
1874
1875 #[allow(clippy::too_many_arguments)]
1877
1878 #[cfg(target_os = "macos")]
1897 fn graph_prefill_preferred(&self) -> bool {
1898 let graph_force = crate::gpu::q1_force() || crate::gpu::q2tp_gpu_opt_in();
1899 if !crate::gpu::enabled_here()
1900 || !graph_force
1901 || std::env::var("CMF_GPU_BLOCK")
1902 .map(|v| v == "0")
1903 .unwrap_or(false)
1904 || std::env::var("CMF_PREFILL_GRAPH").as_deref() == Ok("0")
1907 {
1908 return false;
1909 }
1910 self.weights
1911 .layers
1912 .iter()
1913 .any(|lw| {
1914 matches!(&lw.attn, AttnKind::LinearGdn(w) if w.in_proj_qkv.metal_graph_parts().is_some())
1915 })
1916 }
1917
1918 #[cfg(not(target_os = "macos"))]
1925 fn batch_prefix_prefill(&self) -> bool {
1926 let forced = match std::env::var("CMF_BATCH_PREFIX").as_deref() {
1927 Ok("0") => return false,
1928 Ok("1") => true,
1929 _ => false,
1930 };
1931 crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Prefill)
1932 && crate::gpu::enabled_here()
1933 && !self.graph_refused()
1934 && (forced || self.graph_attn_decline_reason().is_some())
1935 && self.wgpu_graph_attn_decline().is_none()
1936 && self.attn_softcap == 0.0
1937 && self
1938 .weights
1939 .layers
1940 .iter()
1941 .any(|lw| matches!(&lw.ffn, FfnKind::Moe(_)))
1942 && self.automatic_gpu_prefix().is_some()
1943 }
1944
1945 #[cfg(not(target_os = "macos"))]
1946 fn graph_prefill_preferred(&self) -> bool {
1947 let graph_on = crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Prefill);
1955 if !graph_on || !crate::gpu::enabled_here() {
1956 return false;
1957 }
1958 if self.embryo_resident_eligible() {
1962 return crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Decode);
1967 }
1968 if self.o1_active() {
1982 return false;
1983 }
1984 if self.wgpu_graph_attn_decline().is_some() {
1988 return false;
1989 }
1990 if self
1991 .weights
1992 .layers
1993 .iter()
1994 .any(|lw| matches!(&lw.attn, AttnKind::LinearGdn(_)))
1995 {
1996 return true;
1997 }
1998 self.weights
2009 .layers
2010 .iter()
2011 .any(|lw| matches!(&lw.ffn, FfnKind::Moe(_)))
2012 && self.automatic_gpu_prefix().is_none()
2013 }
2014
2015 #[cfg(target_os = "macos")]
2016 fn q1_graph_gpu(
2017 &mut self,
2018 start: usize,
2019 upto: Option<usize>,
2020 position: usize,
2021 h: &mut [f32],
2022 ) -> usize {
2023 let _mt0 = std::time::Instant::now(); use crate::gpu::{AttnGpuLayer, GdnGpuCfg, GdnGpuLayer, GraphDims, MetalFfn, TokenGraph};
2025 let graph_force = crate::gpu::q1_force() || crate::gpu::q2tp_gpu_opt_in();
2026 if self.attn_softcap > 0.0 || !crate::gpu::enabled_here()
2028 || !graph_force
2029 || std::env::var("CMF_GPU_BLOCK")
2030 .map(|v| v == "0")
2031 .unwrap_or(false)
2032 {
2033 if std::env::var("CMF_GRAPH_DBG").is_ok() {
2034 eprintln!(
2035 "block-graph: front gate (softcap={} enabled_here={} graph_force={})",
2036 self.attn_softcap > 0.0,
2037 crate::gpu::enabled_here(),
2038 graph_force,
2039 );
2040 }
2041 if self.graph_head_required {
2042 self.fail_metal_graph("native graph front gate refused");
2043 }
2044 return start;
2045 }
2046 let swa_graph = self.metal_graph_swa();
2053 if (self.swa.is_some() && !swa_graph)
2054 || self.global_attn.is_some()
2055 || self.attention_heads_per_layer.is_some()
2056 || self.attn_v_norm
2057 || (self.graph_attn_decline_reason().is_some() && !swa_graph)
2059 || self.weights.layers.iter().any(|lw| {
2060 lw.attn_out_norm.is_some()
2061 || lw.ffn_out_norm.is_some()
2062 || lw.layer_scale.is_some()
2063 || matches!(&lw.ffn, FfnKind::Dense(d) if !matches!(d.act, Act::Silu | Act::Gelu))
2064 })
2065 {
2066 if let Some(reason) = self.graph_attn_decline_reason().filter(|_| !swa_graph) {
2069 self.note_graph_decline("metal block graph", reason);
2070 }
2071 if std::env::var("CMF_GRAPH_DBG").is_ok() {
2072 eprintln!(
2073 "block-graph: arch ineligible (swa={} gattn={} hpl={} vnorm={} scale_delta={:.2e})",
2074 self.swa.is_some(),
2075 self.global_attn.is_some(),
2076 self.attention_heads_per_layer.is_some(),
2077 self.attn_v_norm,
2078 (self.attn_scale - 1.0 / (self.head_dim as f32).sqrt()).abs(),
2079 );
2080 }
2081 if self.graph_head_required {
2082 self.fail_metal_graph("native graph architecture gate refused");
2083 }
2084 return start;
2085 }
2086 let limit = upto
2089 .map(|u| u + 1)
2090 .unwrap_or(self.num_layers)
2091 .min(self.num_layers);
2092
2093 enum Item<'a> {
2094 Gdn {
2095 run: Vec<GdnGpuLayer<'a>>,
2096 first: usize,
2097 },
2098 Attn {
2099 l: AttnGpuLayer<'a>,
2100 li: usize,
2101 q_norm: Option<&'a [f32]>,
2102 k_norm: Option<&'a [f32]>,
2103 output_gate: bool,
2104 proj_gate: Option<(&'a QTensor, bool)>,
2108 head_gate_w: Option<&'a [f32]>,
2111 bias: Option<(&'a [f32], &'a [f32], &'a [f32])>,
2112 full_gpu: bool,
2115 },
2116 }
2117
2118 let attend_mode = std::env::var("CMF_GPU_ATTEND").unwrap_or_else(|_| "auto".into());
2125 let attend_contract = attend_mode != "0"
2126 && attend_mode != "off"
2127 && self.head_dim % 4 == 0
2128 && self.head_dim <= 256
2129 && self.rotary_dim >= 2
2130 && self.rotary_dim <= self.head_dim
2131 && (self.rotary_dim / 2) % 32 == 0
2132 && self.num_kv_heads > 0
2133 && self.num_heads % self.num_kv_heads == 0;
2134
2135 let mut plan: Vec<Item> = Vec::new();
2136 let mut model_ref: Option<std::sync::Arc<cortiq_core::CmfModel>> = None;
2137 let block_diag = std::env::var("CMF_GRAPH_DBG").is_ok();
2139 let mut scan = start;
2140 while scan < limit {
2141 let lw = &self.weights.layers[self.phys_layer(scan)];
2142 let ffn = match &lw.ffn {
2143 FfnKind::Dense(d) if d.segs.is_empty() => {
2144 let (Some(g), Some(u), Some(dn)) = (
2145 d.gate_proj.metal_graph_parts(),
2146 d.up_proj.metal_graph_parts(),
2147 d.down_proj.metal_graph_parts(),
2148 ) else {
2149 if block_diag {
2150 eprintln!(
2151 "block-graph: L{scan} FFN trio not graph-mappable — run ends"
2152 );
2153 }
2154 break;
2155 };
2156 MetalFfn::Dense {
2157 gate: g,
2158 up: u,
2159 down: dn,
2160 gelu: d.act == Act::Gelu,
2162 }
2163 }
2164 FfnKind::Moe(m) => {
2165 let Some(moe) = metal_moe_graph_parts(m, self.hidden_size) else {
2166 if block_diag {
2167 eprintln!(
2168 "block-graph: L{scan} MoE outside the graph contract — run ends"
2169 );
2170 }
2171 break;
2172 };
2173 if let QTensor::Mapped { model, .. } = &m.experts[0].gate_proj {
2174 model_ref.get_or_insert_with(|| model.clone());
2175 }
2176 MetalFfn::Moe(moe)
2177 }
2178 _ => {
2179 if block_diag {
2180 eprintln!("block-graph: L{scan} non-graph FFN — run ends");
2181 }
2182 break;
2183 }
2184 };
2185 match &lw.attn {
2186 AttnKind::LinearGdn(w) if self.gdn_cfg.is_some() => {
2187 let parts = (
2188 w.in_proj_qkv.metal_graph_parts(),
2189 w.in_proj_z.metal_graph_parts(),
2190 w.in_proj_a.f32_parts(),
2191 w.in_proj_b.f32_parts(),
2192 w.out_proj.metal_graph_parts(),
2193 );
2194 let (Some(qkv), Some(z), Some(a), Some(b), Some(out)) = parts else {
2195 if block_diag {
2196 eprintln!(
2197 "block-graph: L{scan} GDN parts refused (qkv={} z={} a_f32={} b_f32={} out={})",
2198 w.in_proj_qkv.metal_graph_parts().is_some(),
2199 w.in_proj_z.metal_graph_parts().is_some(),
2200 w.in_proj_a.f32_parts().is_some(),
2201 w.in_proj_b.f32_parts().is_some(),
2202 w.out_proj.metal_graph_parts().is_some(),
2203 );
2204 }
2205 break;
2206 };
2207 if let QTensor::Mapped { model, .. } = &w.in_proj_qkv {
2208 model_ref.get_or_insert_with(|| model.clone());
2209 }
2210 let gl = GdnGpuLayer {
2211 attn_norm: &lw.input_norm,
2212 post_norm: &lw.post_norm,
2213 qkv,
2214 z,
2215 a,
2216 b,
2217 out,
2218 ffn,
2219 conv1d: &w.conv1d,
2220 a_log: &w.a_log,
2221 dt_bias: &w.dt_bias,
2222 gnorm: &w.norm,
2223 };
2224 match plan.last_mut() {
2225 Some(Item::Gdn { run, .. }) => run.push(gl),
2226 _ => plan.push(Item::Gdn {
2227 run: vec![gl],
2228 first: scan,
2229 }),
2230 }
2231 }
2232 AttnKind::Full {
2233 wq,
2234 wk,
2235 wv,
2236 wo,
2237 q_norm,
2238 k_norm,
2239 output_gate,
2240 softplus_gate,
2241 bias,
2242 } if (!self.kv_cache.layers[scan].o1_sealed()
2243 || std::env::var("CMF_O1_METAL").as_deref() == Ok("1"))
2248 && softplus_gate
2252 .as_ref()
2253 .is_none_or(|(_, per_head)| *per_head && self.proj_gate_sigmoid) =>
2254 {
2255 let parts = (
2256 wq.metal_graph_parts(),
2257 wk.metal_graph_parts(),
2258 wv.metal_graph_parts(),
2259 wo.metal_graph_parts(),
2260 );
2261 let (Some(pq), Some(pk), Some(pv), Some(po)) = parts else {
2262 break;
2263 };
2264 if let QTensor::Mapped { model, .. } = wq {
2265 model_ref.get_or_insert_with(|| model.clone());
2266 }
2267 let cache = &self.kv_cache.layers[scan];
2268 let o1_metal = cache.o1.is_some()
2272 && std::env::var("CMF_O1_METAL").as_deref() == Ok("1")
2273 && cache.o1_views().is_some();
2274 let head_gate_w = softplus_gate
2278 .as_ref()
2279 .and_then(|(g, _)| g.f32_parts())
2280 .filter(|&(_, r, c)| r == self.num_heads && c == self.hidden_size)
2281 .map(|(d, _, _)| d);
2282 let full_gpu = attend_contract
2283 && softplus_gate.is_none() == head_gate_w.is_none()
2284 && cache.mode == crate::kv_cache::KvMode::F32
2285 && (cache.o1.is_none() || o1_metal)
2286 && bias.is_none()
2287 && pq.1 == self.num_heads * self.head_dim * (1 + *output_gate as usize)
2288 && pk.1 == self.num_kv_heads * self.head_dim
2289 && pv.1 == self.num_kv_heads * self.head_dim
2290 && po.2 == self.num_heads * self.head_dim;
2291 plan.push(Item::Attn {
2292 l: AttnGpuLayer {
2293 attn_norm: &lw.input_norm,
2294 post_norm: &lw.post_norm,
2295 wq: pq,
2296 wk: pk,
2297 wv: pv,
2298 wo: po,
2299 ffn,
2300 },
2301 li: scan,
2302 q_norm: q_norm.as_deref(),
2303 k_norm: k_norm.as_deref(),
2304 output_gate: *output_gate,
2305 proj_gate: softplus_gate.as_ref().map(|(g, per_head)| (g, *per_head)),
2306 head_gate_w,
2307 bias: bias
2308 .as_ref()
2309 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
2310 full_gpu,
2311 });
2312 }
2313 _ => break,
2314 }
2315 scan += 1;
2316 }
2317 let Some(model) = model_ref else {
2318 if std::env::var("CMF_GRAPH_DBG").is_ok() {
2319 eprintln!("q1-graph: no model ref (start {start}, scanned to {scan})");
2320 }
2321 if self.graph_head_required {
2322 self.fail_metal_graph("native graph has no mapped model reference");
2323 }
2324 return start;
2325 };
2326 if plan.is_empty() {
2327 if std::env::var("CMF_GRAPH_DBG").is_ok() {
2328 eprintln!("q1-graph: empty plan at layer {start}");
2329 }
2330 if self.graph_head_required {
2331 self.fail_metal_graph("native graph plan is empty");
2332 }
2333 return start;
2334 }
2335 let has_moe = plan.iter().any(|it| match it {
2336 Item::Gdn { run, .. } => run.iter().any(|l| matches!(l.ffn, MetalFfn::Moe(_))),
2337 Item::Attn { l, .. } => matches!(l.ffn, MetalFfn::Moe(_)),
2338 });
2339 let has_gdn = plan.iter().any(|it| matches!(it, Item::Gdn { .. }));
2340 let dev_attend = attend_contract
2341 && (self.head_dim <= 128
2342 || has_moe
2343 || (self.head_dim <= 256 && has_gdn)
2349 || (self.head_dim <= 256 && swa_graph)
2353 || attend_mode == "force"
2354 || attend_mode == "256");
2355 if !dev_attend {
2356 for it in &mut plan {
2357 if let Item::Attn { li, full_gpu, .. } = it {
2358 let keep_o1 = self.kv_cache.layers[*li].o1.is_some()
2361 && std::env::var("CMF_O1_METAL").as_deref() == Ok("1");
2362 if !keep_o1 {
2363 *full_gpu = false;
2364 }
2365 }
2366 }
2367 }
2368 if std::env::var("CMF_GRAPH_DBG").is_ok() {
2369 use std::sync::atomic::{AtomicBool, Ordering};
2370 static SAID: AtomicBool = AtomicBool::new(false);
2371 if !SAID.swap(true, Ordering::Relaxed) {
2372 let fg = plan
2373 .iter()
2374 .filter(|it| matches!(it, Item::Attn { full_gpu: true, .. }))
2375 .count();
2376 let att = plan
2377 .iter()
2378 .filter(|it| matches!(it, Item::Attn { .. }))
2379 .count();
2380 eprintln!(
2381 "q1-graph: plan of {} items from layer {start} to {scan} | dev_attend={dev_attend} full_gpu {fg}/{att} | hd={} rd={} nkv={} nh={}",
2382 plan.len(),
2383 self.head_dim,
2384 self.rotary_dim,
2385 self.num_kv_heads,
2386 self.num_heads,
2387 );
2388 }
2389 }
2390 let dims = GraphDims {
2391 hidden: self.hidden_size,
2392 eps: self.rms_eps as f32,
2393 gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
2394 };
2395 let Some(mut graph) = TokenGraph::new(&model, dims, h) else {
2396 if self.graph_head_required {
2397 self.fail_metal_graph("native TokenGraph allocation refused");
2398 }
2399 return start;
2400 };
2401 if swa_graph && self.head_dim > 128 {
2402 graph.set_attend_blk_from(64);
2408 }
2409 let gcfg = self.gdn_cfg.map(|cfg| GdnGpuCfg {
2410 nv: cfg.num_v_heads,
2411 nk: cfg.num_k_heads,
2412 dk: cfg.key_head_dim,
2413 dv: cfg.value_head_dim,
2414 kk: cfg.conv_kernel,
2415 hidden: self.hidden_size,
2416 inter: self.intermediate_size,
2417 c_dim: cfg.conv_dim(),
2418 eps: cfg.rms_eps as f32,
2419 gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
2420 });
2421 let mut valid = 0usize;
2425 let mut end = start;
2426 crate::gpu::stageprof(1, _mt0.elapsed()); if std::env::var("CMF_PLAN_DUMP").is_ok() {
2428 static ONCE: std::sync::Once = std::sync::Once::new();
2429 ONCE.call_once(|| {
2430 for it in &plan {
2431 match it {
2432 Item::Gdn { first, run } => {
2433 eprintln!("plan: Gdn first={first} len={}", run.len())
2434 }
2435 Item::Attn { li, full_gpu, .. } => {
2436 eprintln!("plan: Attn li={li} full_gpu={full_gpu}")
2437 }
2438 }
2439 }
2440 });
2441 }
2442 for item in &plan {
2443 let ok = match item {
2444 Item::Gdn { run, .. } => gcfg
2445 .as_ref()
2446 .map(|gc| run.iter().all(|l| graph.gdn_ok(l, gc)))
2447 .unwrap_or(false),
2448 Item::Attn { l, .. } => graph.attn_ok(l),
2449 };
2450 if !ok {
2451 if block_diag {
2452 eprintln!(
2453 "block-graph: plan item {} ({}) failed graph preflight",
2454 valid,
2455 match item {
2456 Item::Gdn { run, first } => format!("GDN run L{first}+{}", run.len()),
2457 Item::Attn { li, .. } => format!("Attn L{li}"),
2458 }
2459 );
2460 }
2461 break;
2462 }
2463 valid += 1;
2464 end += match item {
2465 Item::Gdn { run, .. } => run.len(),
2466 Item::Attn { .. } => 1,
2467 };
2468 }
2469 plan.truncate(valid);
2470 if plan.is_empty() {
2471 if self.graph_head_required {
2472 self.fail_metal_graph("native graph preflight produced no valid items");
2473 }
2474 return start;
2475 }
2476
2477 if self.graph_head_required && (upto.is_some() || end != self.num_layers) {
2478 self.fail_metal_graph("fused-head NLL requires a complete 64-layer graph");
2479 return start;
2480 }
2481
2482 let one_pass = |t: (usize, usize, usize)| {
2491 use cortiq_core::TensorDtype as D;
2492 matches!(
2493 model.tensors[t.0].dtype,
2494 D::Q4TiledP | D::Q4Tiled | D::Q4Block | D::Q8Row | D::Q8_2f | D::Q1
2495 )
2496 };
2497 let dense_fast = plan.iter().all(|it| match it {
2498 Item::Attn {
2499 l, li, full_gpu, ..
2500 } => {
2501 *full_gpu
2502 && self.kv_cache.layers[*li].o1.is_none()
2503 && [l.wq, l.wk, l.wv, l.wo].into_iter().all(one_pass)
2504 && match l.ffn {
2505 MetalFfn::Dense { gate, up, down, .. } => {
2506 one_pass(gate) && one_pass(up) && one_pass(down)
2507 }
2508 _ => false,
2509 }
2510 }
2511 Item::Gdn { .. } => false,
2512 });
2513 let ab = crate::gpu_metal::dense_ab_arm().filter(|_| dense_fast);
2514 let _mv_fast = match ab {
2515 Some((bits, _)) => {
2516 graph.set_dense_concurrent_raw(bits & crate::gpu_metal::DENSE_CONC != 0);
2517 crate::gpu_metal::MvFastGuard::set_raw(bits)
2518 }
2519 None => {
2520 graph.set_dense_concurrent(dense_fast);
2521 let q8r4 = if swa_graph {
2524 crate::gpu_metal::DENSE_Q8R4
2525 } else {
2526 0
2527 };
2528 crate::gpu_metal::MvFastGuard::set_bits(if dense_fast {
2529 crate::gpu_metal::DENSE_MV | crate::gpu_metal::DENSE_FUSE | q8r4
2530 } else {
2531 0
2532 })
2533 }
2534 };
2535
2536 let pool = self.pool.clone();
2539 let (nh, nkv, hd, hs, eps) = (
2540 self.num_heads,
2541 self.num_kv_heads,
2542 self.head_dim,
2543 self.hidden_size,
2544 self.rms_eps,
2545 );
2546 let norm_style = self.norm_style;
2547 let gemma = norm_style == cortiq_core::NormStyle::Gemma;
2548 let want = self.gdn_cfg.map(|c| c.state_len()).unwrap_or(0);
2549 let kv_id = self.graph_kv_id;
2550 let mut pending: Vec<(usize, usize)> = Vec::new();
2553 let mut dev_attn: Vec<usize> = Vec::new();
2556 for item in &plan {
2557 let _xt0 = std::time::Instant::now();
2558 let _xkind: u32 = match item {
2559 Item::Gdn { .. } => 2,
2560 Item::Attn { .. } => 3,
2561 };
2562 if self.loop_final_norm {
2564 let item_start = match item {
2565 Item::Gdn { first, .. } => *first,
2566 Item::Attn { li, .. } => *li,
2567 };
2568 if item_start > start && self.is_loop_end(item_start - 1) {
2569 graph.encode_loop_norm(&self.weights.final_norm);
2570 }
2571 }
2572 match item {
2573 Item::Gdn { run, first } => {
2574 for l in &mut self.kv_cache.layers[*first..*first + run.len()] {
2575 if l.linear_state.len() != want {
2576 l.linear_state = vec![0f32; want];
2577 }
2578 }
2579 let ro: Vec<&[f32]> = self.kv_cache.layers[*first..*first + run.len()]
2580 .iter()
2581 .map(|l| l.linear_state.as_slice())
2582 .collect();
2583 let _ig = std::time::Instant::now();
2584 if !graph.encode_gdn_run(run, &ro, gcfg.as_ref().unwrap()) {
2585 tracing::error!("q1 graph: GDN run refused after validation");
2587 return start;
2588 }
2589 graph.commit_kind = 2;
2592 graph.commit();
2593 crate::gpu::stageprof(0, _ig.elapsed());
2594 pending.push((*first, run.len()));
2595 }
2596 Item::Attn {
2597 l,
2598 li,
2599 q_norm,
2600 k_norm,
2601 output_gate,
2602 proj_gate,
2603 head_gate_w,
2604 bias,
2605 full_gpu,
2606 } => {
2607 let _ia = std::time::Instant::now();
2608 let inv_freq_l = self.layer_inv_freq(*li);
2614 let rd_l = self.layer_geom(*li).2;
2615 let window_l = self.layer_window(*li);
2616 let head_gate_w = *head_gate_w;
2617 if *full_gpu {
2619 let cache = &self.kv_cache.layers[*li];
2620 let o1p = if cache.o1.is_some() {
2621 match cache.o1_views() {
2622 Some(views) => Some(crate::gpu::O1AttnParams {
2623 views,
2624 epoch: self.o1_epoch,
2625 }),
2626 None => None,
2628 }
2629 } else {
2630 None
2631 };
2632 let o1_layer = cache.o1.is_some();
2633 if o1_layer && o1p.is_none() {
2634 }
2636 let cpu_k: Vec<&[f32]> = (0..nkv).map(|g| cache.head_keys(g)).collect();
2637 let cpu_v: Vec<&[f32]> = (0..nkv).map(|g| cache.head_values(g)).collect();
2638 let cpu_stored = if o1_layer { 0 } else { cpu_k[0].len() / hd };
2639 debug_assert!(
2642 cache.base() == 0 || window_l.is_some_and(|w| cpu_stored + 1 >= w)
2643 );
2644 let p = crate::gpu::AttnDeviceParams {
2645 kv_id,
2646 layer: *li,
2647 nh,
2648 nkv,
2649 hd,
2650 rd: rd_l,
2651 position,
2652 scale: self.attn_scale,
2653 eps: eps as f32,
2654 gemma,
2655 late_qk_norm: self.qk_norm_after_rope,
2656 output_gate: *output_gate,
2657 q_norm: *q_norm,
2658 k_norm: *k_norm,
2659 inv_freq: &inv_freq_l,
2660 cpu_k,
2661 cpu_v,
2662 cpu_stored,
2663 cpu_gen: cache.generation(),
2664 o1: o1p,
2665 window: window_l,
2666 head_gate: head_gate_w,
2667 };
2668 let o1_bad = o1_layer && p.o1.is_none();
2669 if !o1_bad && graph.attn_device_ok(l, &p) && graph.encode_attn_device(l, &p)
2670 {
2671 if p.o1.is_none() {
2673 dev_attn.push(*li);
2674 }
2675 graph.commit_kind = 3;
2676 graph.commit();
2677 crate::gpu::stageprof(_xkind, _xt0.elapsed());
2681 continue;
2682 }
2683 }
2685 graph.encode_attn_prefix(l);
2686 if let Err(err) = graph.sync_checked() {
2687 self.fail_metal_graph(&err);
2688 return start;
2689 }
2690 if !pending.is_empty() {
2691 let idxs: Vec<usize> =
2692 pending.drain(..).flat_map(|(f, n)| f..f + n).collect();
2693 let mut outs: Vec<&mut [f32]> = self
2694 .kv_cache
2695 .layers
2696 .iter_mut()
2697 .enumerate()
2698 .filter(|(i, _)| idxs.binary_search(i).is_ok())
2699 .map(|(_, s)| s.linear_state.as_mut_slice())
2700 .collect();
2701 graph.read_states(&mut outs);
2702 }
2703 let mut q_raw = attention::take_buf(l.wq.1);
2704 let mut k = attention::take_buf(l.wk.1);
2705 let mut v = attention::take_buf(l.wv.1);
2706 graph.read_qkv(&mut q_raw, &mut k, &mut v);
2707 let mut gate_raw = proj_gate.map(|(gp, _)| {
2710 let mut normed = attention::take_buf(hs);
2711 graph.read_normed(&mut normed);
2712 let mut raw = attention::take_buf(gp.rows());
2713 gp.matvec(&normed, &mut raw, pool.as_deref());
2714 attention::recycle_buf(&mut normed);
2715 raw
2716 });
2717 let cfg = QwenAttnCfg {
2718 num_heads: nh,
2719 num_kv_heads: nkv,
2720 head_dim: hd,
2721 hidden_size: hs,
2722 position,
2723 inv_freq: &inv_freq_l,
2724 rotary_dim: rd_l,
2725 scale: self.attn_scale,
2726 softcap: self.attn_softcap,
2727 window: window_l,
2728 v_norm: false,
2729 qk_norm_after_rope: self.qk_norm_after_rope,
2730 gate_sigmoid: self.proj_gate_sigmoid,
2731 q_norm: *q_norm,
2732 k_norm: *k_norm,
2733 output_gate: *output_gate,
2734 softplus_gate: None,
2735 rope_scale: 1.0,
2736 bias: *bias,
2737 rms_eps: eps,
2738 norm_style,
2739 pool: pool.as_deref(),
2740 v_head_dim: hd,
2741 };
2742 let oracle = std::env::var("CMF_ATTN_ORACLE").as_deref() == Ok("1")
2745 || std::env::var("CMF_ATTN_DUMP").is_ok();
2746 let _ = full_gpu;
2747 let oracle_in = oracle.then(|| (q_raw.clone(), k.clone(), v.clone()));
2748 let mut ao = attention::qwen_attention_core(
2749 q_raw,
2750 k,
2751 v,
2752 &mut self.kv_cache.layers[*li],
2753 &cfg,
2754 );
2755 if let Ok(dir) = std::env::var("CMF_ATTN_DUMP") {
2759 if let Some((qr0, k0, v0)) = oracle_in.clone() {
2760 let (cq, _cg, _ck, _cv) =
2761 attention::finish_projection_debug(qr0, k0, v0, &cfg, position);
2762 let cache = &self.kv_cache.layers[*li];
2763 let n = cache.head_keys(0).len() / hd;
2764 let mut bytes: Vec<u8> = Vec::new();
2765 for v in [nh as u32, nkv as u32, hd as u32, n as u32, position as u32] {
2766 bytes.extend_from_slice(&v.to_le_bytes());
2767 }
2768 for v in &cq {
2769 bytes.extend_from_slice(&v.to_le_bytes());
2770 }
2771 for g in 0..nkv {
2772 for v in cache.head_keys(g) {
2773 bytes.extend_from_slice(&v.to_le_bytes());
2774 }
2775 }
2776 for g in 0..nkv {
2777 for v in cache.head_values(g) {
2778 bytes.extend_from_slice(&v.to_le_bytes());
2779 }
2780 }
2781 let _ =
2782 std::fs::write(format!("{dir}/L{li}_pos{position}.bin"), &bytes);
2783 }
2784 }
2785 if let Some((qr0, k0, v0)) =
2786 oracle_in.filter(|_| std::env::var("CMF_ATTN_ORACLE").as_deref() == Ok("1"))
2787 {
2788 let (cq, _cg, ck, cv) =
2789 attention::finish_projection_debug(qr0, k0, v0, &cfg, position);
2790 let mut h_now = vec![0f32; hs];
2791 graph.read_h(&mut h_now);
2792 let cache = &self.kv_cache.layers[*li];
2793 let n_after = cache.head_keys(0).len() / hd;
2794 let stored = n_after.saturating_sub(1);
2798 let cpu_k: Vec<&[f32]> = (0..nkv)
2799 .map(|g| &cache.head_keys(g)[..stored * hd])
2800 .collect();
2801 let cpu_v: Vec<&[f32]> = (0..nkv)
2802 .map(|g| &cache.head_values(g)[..stored * hd])
2803 .collect();
2804 let p = crate::gpu::AttnDeviceParams {
2805 kv_id,
2806 layer: *li,
2807 nh,
2808 nkv,
2809 hd,
2810 rd: rd_l,
2814 position,
2815 scale: self.attn_scale,
2816 eps: eps as f32,
2817 gemma,
2818 late_qk_norm: self.qk_norm_after_rope,
2819 output_gate: *output_gate,
2820 q_norm: *q_norm,
2821 k_norm: *k_norm,
2822 inv_freq: &inv_freq_l,
2823 cpu_k,
2824 cpu_v,
2825 cpu_stored: stored,
2826 cpu_gen: cache.generation(),
2827 o1: None,
2828 window: window_l,
2829 head_gate: None,
2830 };
2831 if let Some((dq, dk, dv, dao)) = graph.debug_attn_device(l, &p, &h_now) {
2832 let md = |a: &[f32], b: &[f32]| {
2833 a.iter().zip(b).fold(0f32, |m, (x, y)| m.max((x - y).abs()))
2834 };
2835 let nn = |a: &[f32]| a.iter().map(|x| x * x).sum::<f32>().sqrt();
2836 eprintln!(
2837 "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}",
2838 nn(&cq),
2839 md(&cq, &dq),
2840 nn(&ck),
2841 md(&ck, &dk),
2842 nn(&cv),
2843 md(&cv, &dv),
2844 nn(&ao),
2845 md(&ao, &dao)
2846 );
2847 } else {
2848 eprintln!("attn-oracle L{li}: device probe declined");
2849 }
2850 }
2851 if let (Some(raw), Some((_, per_head))) = (gate_raw.as_deref(), *proj_gate) {
2852 attention::apply_projected_gate(
2855 &mut ao,
2856 raw,
2857 per_head,
2858 hd,
2859 self.proj_gate_sigmoid,
2860 );
2861 }
2862 if let Some(mut raw) = gate_raw.take() {
2863 attention::recycle_buf(&mut raw);
2864 }
2865 graph.encode_attn_suffix(l, &ao);
2866 graph.commit();
2869 attention::recycle_buf(&mut ao);
2870 }
2871 }
2872
2873 crate::gpu::stageprof(_xkind, _xt0.elapsed());
2874 }
2875 let mut lm_rows = None;
2880 if self.graph_want_logits
2881 && upto.is_none()
2882 && end == self.num_layers
2883 && std::env::var("CMF_GPU_LMHEAD")
2884 .map(|v| v != "0")
2885 .unwrap_or(true)
2886 {
2887 if let Some(lm) = self.weights.lm_head.metal_graph_parts() {
2888 if graph.lm_head_ok(lm) {
2889 graph.encode_lm_head(&self.weights.final_norm, lm);
2890 lm_rows = Some(lm.1);
2891 }
2892 }
2893 }
2894 if self.graph_head_required && lm_rows.is_none() {
2895 METAL_GRAPH_HEAD_MISS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
2896 self.fail_metal_graph("fused graph head was requested but not encodable");
2897 return start;
2898 }
2899 let _sy0 = std::time::Instant::now();
2900 if let Err(err) = graph.sync_checked() {
2901 self.fail_metal_graph(&err);
2902 return start;
2903 }
2904 let _rs0 = std::time::Instant::now();
2905 if !pending.is_empty() {
2906 let idxs: Vec<usize> = pending.drain(..).flat_map(|(f, n)| f..f + n).collect();
2907 let mut outs: Vec<&mut [f32]> = self
2908 .kv_cache
2909 .layers
2910 .iter_mut()
2911 .enumerate()
2912 .filter(|(i, _)| idxs.binary_search(i).is_ok())
2913 .map(|(_, s)| s.linear_state.as_mut_slice())
2914 .collect();
2915 graph.read_states(&mut outs);
2916 }
2917 if std::env::var("CMF_GRAPH_HOSTPROF").as_deref() == Ok("1") {
2918 use std::sync::atomic::{AtomicU64, Ordering};
2919 static SY: AtomicU64 = AtomicU64::new(0);
2920 static RS: AtomicU64 = AtomicU64::new(0);
2921 static N: AtomicU64 = AtomicU64::new(0);
2922 SY.fetch_add((_rs0 - _sy0).as_nanos() as u64, Ordering::Relaxed);
2923 RS.fetch_add(_rs0.elapsed().as_nanos() as u64, Ordering::Relaxed);
2924 let n = N.fetch_add(1, Ordering::Relaxed) + 1;
2925 if n % 100 == 0 {
2926 eprintln!(
2927 "postprof: sync-wait {:.1} ms/ток | read_states {:.1} ms/ток ({n})",
2928 SY.load(Ordering::Relaxed) as f64 / n as f64 / 1e6,
2929 RS.load(Ordering::Relaxed) as f64 / n as f64 / 1e6
2930 );
2931 }
2932 }
2933 if let Some(rows) = lm_rows {
2934 crate::gpu::hostprof_encode_done(_mt0);
2935 let mut lg = attention::take_buf(rows.min(self.vocab_size));
2936 graph.read_logits(&mut lg);
2937 crate::gpu::hostprof_total(_mt0);
2938 lg.resize(self.vocab_size, 0.0);
2939 if let Some(c) = self.final_softcap {
2940 for l in lg.iter_mut() {
2941 *l = c * (*l / c).tanh();
2942 }
2943 }
2944 self.graph_logits = Some(lg);
2945 }
2946 graph.read_h(h);
2947 if self.graph_head_required && self.graph_logits.is_none() {
2948 METAL_GRAPH_HEAD_MISS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
2949 self.fail_metal_graph("fused graph head completed without logits readback");
2950 return start;
2951 }
2952 METAL_GRAPH_TOK_OK.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
2953 METAL_GRAPH_LAYERS.fetch_add(
2954 end.saturating_sub(start) as u64,
2955 std::sync::atomic::Ordering::Relaxed,
2956 );
2957 if self.graph_head_required {
2958 METAL_GRAPH_HEAD_OK.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
2959 }
2960 for li in dev_attn {
2964 let mut krow = attention::take_buf(nkv * hd);
2965 let mut vrow = attention::take_buf(nkv * hd);
2966 if crate::gpu::kv_mirror_read_last(kv_id, li, nkv, hd, &mut krow, &mut vrow) {
2967 let cache = &mut self.kv_cache.layers[li];
2968 cache.append(&krow, &vrow, &[]);
2969 let n = cache.seq_len;
2970 let mut imp = attention::take_buf(n);
2971 crate::gpu::kv_mirror_take_imp(kv_id, li, &mut imp);
2972 cache.accumulate_imp(&imp);
2973 attention::recycle_buf(&mut imp);
2974 }
2975 attention::recycle_buf(&mut krow);
2976 attention::recycle_buf(&mut vrow);
2977 }
2978 if let Some((_, arm)) = ab {
2979 crate::gpu_metal::dense_ab_record(arm, _mt0.elapsed());
2980 }
2981 end
2982 }
2983
2984 pub fn new(
2985 tokenizer: Tokenizer,
2986 weights: PipelineWeights,
2987 hidden_size: usize,
2988 intermediate_size: usize,
2989 num_heads: usize,
2990 num_kv_heads: usize,
2991 head_dim: usize,
2992 num_layers: usize,
2993 physical_layers: usize,
2994 loop_final_norm: bool,
2995 vocab_size: usize,
2996 rms_eps: f64,
2997 rope_base: f32,
2998 norm_style: NormStyle,
2999 max_seq_len: usize,
3000 sampler_config: SamplerConfig,
3001 ) -> Self {
3002 let rng = match sampler_config.seed {
3003 Some(s) => SplitMix64::new(s),
3004 None => SplitMix64::from_entropy(),
3005 };
3006 let inv_freq = std::sync::Arc::new(attention::rope_inv_freq(head_dim, rope_base));
3007 let pool = Pool::from_env();
3008 if let Some(p) = &pool {
3009 tracing::info!("worker pool: {} threads", p.n_workers());
3010 if let Some(model) = weights
3012 .lm_head
3013 .model_arc()
3014 .or_else(|| weights.embed_tokens.model_arc())
3015 {
3016 let regions: Vec<&[u8]> =
3017 model.tensors.iter().map(|t| model.entry_bytes(t)).collect();
3018 p.bind_numa(®ions);
3019 }
3020 }
3021 Self {
3022 gpu_plan: None,
3023 tokenizer: std::sync::Arc::new(tokenizer),
3024 kv_cache: KvCache::new(num_layers, num_kv_heads, head_dim, max_seq_len),
3025 sampler_config,
3026 weights,
3027 hidden_size,
3028 intermediate_size,
3029 num_heads,
3030 num_kv_heads,
3031 head_dim,
3032 num_layers,
3033 physical_layers,
3034 loop_final_norm,
3035 vocab_size,
3036 rms_eps,
3037 rope_base,
3038 norm_style,
3039 rotary_dim: head_dim,
3040 attention_heads_per_layer: None,
3041 kv_heads_per_layer: None,
3042 v_head_dim: None,
3043 layer_dump: std::env::var_os("CMF_LAYER_DUMP")
3044 .filter(|v| !v.is_empty())
3045 .map(std::path::PathBuf::from),
3046 graph_declines: std::cell::RefCell::new(Vec::new()),
3047 mimo_moe: Default::default(),
3048 vmf_cfg: None,
3049 gdn_cfg: None,
3050 kda_cfg: None,
3051 g3n: None,
3052 dsv4: None,
3053 dsv41: None,
3054 dsv41_vision: None,
3055 dsv41_prefill: None,
3056 qwen4_exp: None,
3057 dsv4_mtp: Vec::new(),
3058 dspark: None,
3059 dspark_pending: Vec::new(),
3060 dspark_hist: Vec::new(),
3061 dspark_real: Vec::new(),
3062 dspark_trunk_picks: Vec::new(),
3063 dspark_exp: Vec::new(),
3064 dspark_draft_ns: 0,
3065 logit_multiplier: None,
3066 cancel: std::sync::Arc::new(std::sync::atomic::AtomicBool::new(false)),
3067 graph_failed: std::sync::atomic::AtomicBool::new(false),
3068 kv_history: Vec::new(),
3069 kv_history_device: false,
3070 short_conv_cfg: None,
3071 mtp: None,
3072 mimo_mtp: None,
3073 verify_exact_moe: false,
3074 speculative: std::env::var("CMF_MTP").map(|v| v != "0").unwrap_or(true),
3075 ignore_eos: false,
3076 draft_full_streak: 0,
3077 spec_k_adapt: None,
3078 spec_acc_ewma: 0.7,
3079 rng,
3080 sampler_scratch: SamplerScratch::default(),
3081 spec_forced: None,
3082 spec_q: Vec::new(),
3083 spec_p: Vec::new(),
3084 spec_res: Vec::new(),
3085 spec_qs: Vec::new(),
3086 spec_ps: Vec::new(),
3087 spec_ress: Vec::new(),
3088 mtp_graph_mode: None,
3089 #[cfg(target_os = "macos")]
3090 metal_verify: None,
3091 inv_freq,
3092 ws: ForwardScratch::new(hidden_size),
3093 pool,
3094 model: None,
3095 dyn_force_f32: false,
3096 dyn_skill_layers: Vec::new(),
3097 dyn_active: None,
3098 dyn_blend_loaded: false,
3099 dyn_phi_layer: None,
3100 dyn_phi_ema: Vec::new(),
3101 dyn_phi_seen: 0,
3102 dyn_router: None,
3103 o1_cfg: None,
3104 o1_epoch: 0,
3105 o1_flags: Vec::new(),
3106 trace: false,
3107 calib_temp: 1.0,
3108 confidence_on: true,
3109 embed_multiplier: 1.0,
3110 attn_scale: 1.0 / (head_dim as f32).sqrt(),
3111 swa: None,
3112 swa_trim: None,
3113 sliding_layers: None,
3114 anchor_core: None,
3115 bounded_rope: None,
3116 kv_prefix: KvPrefix::default(),
3117 last_prefill_tokens: 0,
3118 inv_freq_local: None,
3119 rotary_dim_local: None,
3120 rope_scale: 1.0,
3121 rope_scale_local: 1.0,
3122 global_attn: None,
3123 inv_freq_global: None,
3124 attn_v_norm: false,
3125 qk_norm_after_rope: false,
3126 proj_gate_sigmoid: false,
3127 final_softcap: None,
3128 head_clusters: None,
3129 attn_softcap: 0.0,
3130 graph_want_logits: false,
3131 graph_head_required: false,
3132 graph_logits: None,
3133 embryo_graph: None,
3134 graph_refused: std::sync::atomic::AtomicBool::new(false),
3135 graph_kv_id: {
3136 static NEXT: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(1);
3137 NEXT.fetch_add(1, std::sync::atomic::Ordering::Relaxed)
3138 },
3139 #[cfg(test)]
3140 nll_test_fail_at: None,
3141 #[cfg(test)]
3142 nll_test_force_serial: false,
3143 }
3144 }
3145
3146 pub fn set_o1(&mut self, cfg: Option<crate::nystrom::O1Cfg>) {
3154 if let Err(e) = self.try_set_o1(cfg) {
3155 tracing::error!("{e}");
3156 }
3157 }
3158
3159 pub fn bounded_native(&self) -> bool {
3163 self.anchor_core.is_some()
3164 }
3165
3166 pub fn device_state_bytes(&self) -> Option<(u64, u64)> {
3169 crate::gpu::embryo_device_state_bytes(self.graph_kv_id)
3170 }
3171
3172 pub fn o1_refusal(&self) -> Option<String> {
3174 self.anchor_core.as_ref().map(|ac| {
3175 format!(
3176 "--o1 / CMF_O1 refused: the anchor is native bounded \
3177 (anchor_core kind={} window={} sink={}); the file's operator \
3178 is executed as-is and no post-hoc Nyström overlay applies",
3179 ac.kind, ac.window, ac.sink
3180 )
3181 })
3182 }
3183
3184 pub fn try_set_o1(&mut self, cfg: Option<crate::nystrom::O1Cfg>) -> Result<(), String> {
3186 if let Some(c) = &cfg {
3187 if let Some(why) = self.o1_refusal() {
3188 self.o1_flags = Vec::new();
3189 self.o1_cfg = None;
3190 return Err(why);
3191 }
3192 if crate::nystrom::o1_deferred_boundary(c.w, c.sink).is_none() {
3193 self.o1_flags.clear();
3194 self.o1_cfg = None;
3195 return Err(format!(
3196 "o1 disabled: w + sink + slack + 1 overflows usize (w={}, sink={})",
3197 c.w, c.sink
3198 ));
3199 }
3200 }
3201 self.o1_flags = match &cfg {
3202 Some(c) => {
3203 let mut flags = c.layer_flags(self.num_layers);
3204 for (li, f) in flags.iter_mut().enumerate() {
3205 if *f
3211 && (!matches!(
3212 self.weights.layers[self.phys_layer(li)].attn,
3213 AttnKind::Full { .. }
3214 ) || self.layer_window(li).is_some()
3215 || self.kv_cache.layers[li].sinks.is_some()
3216 || self.layer_v_dim(li) != self.layer_geom(li).1)
3217 {
3218 *f = false;
3219 }
3220 }
3221 flags
3222 }
3223 None => Vec::new(),
3224 };
3225 if let Some(c) = &cfg {
3226 let n = self.o1_flags.iter().filter(|&&f| f).count();
3227 tracing::info!(
3228 "o1 nystrom attention: {n}/{} layer(s), m={} w={} sink={} rect={:?}",
3229 self.num_layers,
3230 c.m,
3231 c.w,
3232 c.sink,
3233 c.rect
3234 );
3235 }
3236 self.o1_cfg = cfg;
3237 Ok(())
3238 }
3239
3240 pub fn install_bounded(
3246 &mut self,
3247 cfg: &cortiq_core::AnchorCoreConfig,
3248 ) -> Result<(), String> {
3249 if !cortiq_core::AnchorCoreConfig::KINDS.contains(&cfg.kind.as_str()) {
3250 return Err(format!(
3251 "anchor_core kind '{}' is not executable by this runtime",
3252 cfg.kind
3253 ));
3254 }
3255 if cfg.window == 0 {
3256 return Err("anchor_core.window must be >= 1".into());
3257 }
3258 let mut n = 0usize;
3259 for li in 0..self.num_layers {
3260 let pl = self.phys_layer(li);
3261 if let AttnKind::Bounded(w) = &self.weights.layers[pl].attn {
3262 if w.window != cfg.window || w.sink != cfg.sink {
3263 return Err(format!(
3264 "layer {li}: bounded weights (window {} sink {}) disagree with \
3265 anchor_core (window {} sink {})",
3266 w.window, w.sink, cfg.window, cfg.sink
3267 ));
3268 }
3269 self.kv_cache.layers[li].install_bounded(cfg.window);
3270 n += 1;
3271 }
3272 }
3273 if n == 0 {
3274 return Err("anchor_core is present but no layer executes it".into());
3275 }
3276 let rope = crate::bounded::BoundedRope::new(cfg.window, &self.inv_freq, self.rope_scale);
3277 self.bounded_rope = Some(std::sync::Arc::new(rope));
3278 self.anchor_core = Some(cfg.clone());
3279 self.embryo_graph = None;
3280 tracing::info!(
3281 "bounded anchor {}: {n} layer(s), window {} sink {} — {} B of ring per layer",
3282 cfg.kind,
3283 cfg.window,
3284 cfg.sink,
3285 self.kv_cache.layers.iter().map(|l| l.bounded_state_bytes()).max().unwrap_or(0)
3286 );
3287 Ok(())
3288 }
3289
3290 pub fn install_wire_identity(&mut self, identity: u64) {
3294 for li in 0..self.kv_cache.layers.len() {
3295 let pl = self.phys_layer(li);
3296 let kind = match self.weights.layers.get(pl).map(|l| &l.attn) {
3297 Some(AttnKind::Bounded(_)) => crate::kv_cache::WireKind::Bounded,
3298 Some(AttnKind::Linear(_))
3299 | Some(AttnKind::LinearGdn(_))
3300 | Some(AttnKind::ShortConv(_))
3301 | Some(AttnKind::Kda(_)) => crate::kv_cache::WireKind::Linear,
3302 _ => crate::kv_cache::WireKind::Full,
3303 };
3304 let l = &mut self.kv_cache.layers[li];
3305 l.wire_kind = kind;
3306 l.wire_identity = identity;
3307 l.wire_layer = li as u32;
3310 }
3311 }
3312
3313 pub fn clear_history(&mut self) {
3315 self.kv_history.clear();
3316 self.kv_history_device = false;
3317 self.kv_prefix.clear();
3318 }
3319
3320 pub fn graph_refused(&self) -> bool {
3324 self.graph_refused
3325 .load(std::sync::atomic::Ordering::Relaxed)
3326 }
3327
3328 pub fn mark_graph_refused(&self) {
3330 if !self
3331 .graph_refused
3332 .swap(true, std::sync::atomic::Ordering::Relaxed)
3333 {
3334 tracing::info!(
3335 "token graph: unsupported for this pipeline (seq {}) — not retrying",
3336 self.graph_kv_id
3337 );
3338 }
3339 }
3340
3341 pub fn device_sequence_position(&self) -> Option<usize> {
3345 crate::gpu::embryo_device_next_position(self.graph_kv_id)
3346 }
3347
3348 fn prefix_owner_matches(&self, n: usize, recorded_on_device: bool) -> bool {
3353 let dev = self.device_sequence_position();
3354 if recorded_on_device {
3355 dev == Some(n) && self.embryo_resident_wanted()
3356 } else {
3357 dev.is_none()
3358 }
3359 }
3360
3361 pub(crate) fn invalidate_for_weight_change(&mut self) {
3369 self.clear_sequence_state();
3370 self.embryo_graph = None;
3371 }
3372
3373 fn cached_prefix_len(&self, input_ids: &[u32]) -> usize {
3379 let (n, on_device) = if self.bounded_native() {
3380 (self.kv_prefix.extension(input_ids), self.kv_prefix.on_device())
3381 } else {
3382 let h = &self.kv_history;
3383 if !h.is_empty() && h.len() < input_ids.len() && input_ids[..h.len()] == h[..] {
3384 (h.len(), self.kv_history_device)
3385 } else {
3386 (0, false)
3387 }
3388 };
3389 if n > 0 && !self.prefix_owner_matches(n, on_device) {
3392 tracing::warn!(
3393 "kv-reuse refused: the cached prefix ({n} positions) was built on the {} path, \
3394 the device now holds {:?} — re-prefilling from zero",
3395 if on_device { "resident device" } else { "host" },
3396 self.device_sequence_position()
3397 );
3398 return 0;
3399 }
3400 n
3401 }
3402
3403 pub fn reusable_prefix_len(&self, input_ids: &[u32]) -> usize {
3406 self.cached_prefix_len(input_ids)
3407 }
3408
3409 fn record_consumed_prefix(&mut self, consumed: &[u32], reused: usize) {
3415 let on_device = self.device_sequence_position().is_some();
3416 if self.bounded_native() {
3417 let keep = reused > 0 && reused == self.kv_prefix.len() && reused <= consumed.len();
3418 let prev_device = self.kv_prefix.on_device();
3419 self.kv_history.clear();
3420 self.kv_history_device = false;
3421 if keep && prev_device == on_device {
3422 self.kv_prefix.extend(&consumed[reused..]);
3423 } else {
3424 self.kv_prefix.set(consumed);
3425 }
3426 self.kv_prefix.set_on_device(on_device);
3427 } else {
3428 self.kv_history = consumed.to_vec();
3429 self.kv_history_device = on_device;
3430 }
3431 }
3432
3433 pub fn o1_active(&self) -> bool {
3435 self.o1_cfg.is_some() && self.o1_flags.iter().any(|&f| f)
3436 }
3437
3438 pub fn generation_batch_k(&self) -> usize {
3450 if let Some(k) = std::env::var("CMF_BATCH_K")
3451 .ok()
3452 .and_then(|v| v.parse::<usize>().ok())
3453 {
3454 return k;
3455 }
3456 #[cfg(not(target_os = "macos"))]
3457 if self.graph_prefill_preferred() && !self.o1_active() {
3458 return 32;
3459 }
3460 0
3461 }
3462
3463 pub fn generation_graph_prefill(&self) -> bool {
3464 let graph = self.graph_prefill_preferred();
3465 #[cfg(not(target_os = "macos"))]
3476 if graph
3477 && self.generation_batch_k() > 0
3478 && std::env::var("CMF_MTP_CHAIN_PROBE").is_err()
3479 {
3480 return false;
3481 }
3482 graph
3483 }
3484
3485 pub fn o1_device_stats(&self) -> (usize, u64) {
3490 crate::gpu::o1_device_stats(self.graph_kv_id)
3491 }
3492
3493 pub fn o1_begin(&mut self) {
3498 self.o1_begin_with_prefix(None);
3499 }
3500
3501 pub fn o1_begin_with_prefix(&mut self, requested_prefix: Option<usize>) {
3505 if let Some(c) = &self.o1_cfg {
3506 let (m, w, sink, rect) = (c.m, c.w, c.sink, c.rect);
3507 let boundary = requested_prefix.map(|p| {
3508 p.max(
3509 crate::nystrom::o1_deferred_boundary(w, sink)
3510 .expect("o1 config boundary validated in set_o1"),
3511 )
3512 });
3513 for (li, &f) in self.o1_flags.iter().enumerate() {
3514 if f {
3515 self.kv_cache.layers[li].o1_begin_with_boundary(m, w, sink, rect, boundary);
3516 }
3517 }
3518 }
3519 }
3520
3521 fn o1_effective_boundary(&self, requested_prefix: usize) -> Option<usize> {
3523 self.o1_cfg.as_ref().and_then(|c| {
3524 crate::nystrom::o1_deferred_boundary(c.w, c.sink)
3525 .map(|floor| requested_prefix.max(floor))
3526 })
3527 }
3528
3529 fn o1_note_transition(&mut self) {
3530 let mut transitioned = false;
3534 for (li, &flagged) in self.o1_flags.iter().enumerate() {
3535 if flagged {
3536 transitioned |= self.kv_cache.layers[li].take_o1_transition();
3537 }
3538 }
3539 if transitioned {
3540 self.o1_epoch = self.o1_epoch.wrapping_add(1);
3541 }
3542 }
3543
3544 fn o1_pending(&self) -> bool {
3545 self.o1_flags.iter().enumerate().any(|(li, &f)| {
3546 f && self.kv_cache.layers[li].seq_len > 0
3547 && self.kv_cache.layers[li].o1_pending_boundary().is_some()
3548 })
3549 }
3550
3551 fn o1_fail(&mut self, err: String) {
3552 tracing::error!("o1 deferred seal failed; terminating sequence: {err}");
3553 self.clear_sequence_state();
3554 self.graph_failed
3555 .store(true, std::sync::atomic::Ordering::Relaxed);
3556 self.cancel
3557 .store(true, std::sync::atomic::Ordering::Relaxed);
3558 }
3559
3560 pub fn o1_seal_checked(&mut self) -> Result<bool, String> {
3565 if self.o1_cfg.is_none() {
3566 return Ok(false);
3567 }
3568 let mut participating = false;
3569 for li in 0..self.num_layers {
3570 if !self.o1_flags.get(li).copied().unwrap_or(false) {
3571 continue;
3572 }
3573 if let Some(err) = self.kv_cache.layers[li].take_o1_error() {
3574 return Err(err);
3575 }
3576 if self.kv_cache.layers[li].seq_len == 0 {
3577 continue;
3578 }
3579 participating = true;
3580 let num_heads = self.layer_num_heads(li);
3581 self.kv_cache.layers[li].o1_seal_checked(num_heads)?;
3582 }
3583 self.o1_note_transition();
3584 for li in 0..self.num_layers {
3585 if self.o1_flags.get(li).copied().unwrap_or(false) {
3586 if let Some(err) = self.kv_cache.layers[li].take_o1_error() {
3587 return Err(err);
3588 }
3589 }
3590 }
3591 Ok(participating
3592 && (0..self.num_layers).all(|li| {
3593 !self.o1_flags.get(li).copied().unwrap_or(false)
3594 || self.kv_cache.layers[li].seq_len == 0
3595 || self.kv_cache.layers[li].o1_sealed()
3596 }))
3597 }
3598
3599 fn o1_progress(&mut self) {
3602 if !self.o1_active() {
3603 return;
3604 }
3605 for li in 0..self.num_layers {
3606 if self.o1_flags.get(li).copied().unwrap_or(false) {
3607 if let Some(err) = self.kv_cache.layers[li].take_o1_error() {
3608 self.o1_fail(err);
3609 return;
3610 }
3611 }
3612 }
3613 self.o1_note_transition();
3617 if !self.o1_pending() {
3618 return;
3619 }
3620 if let Err(err) = self.o1_seal_checked() {
3621 self.o1_fail(err);
3622 }
3623 }
3624
3625 fn check_o1_progress_failure(&mut self, phase: &str) -> Result<(), String> {
3630 if self
3631 .graph_failed
3632 .swap(false, std::sync::atomic::Ordering::Relaxed)
3633 {
3634 self.cancel
3635 .store(false, std::sync::atomic::Ordering::Relaxed);
3636 self.clear_sequence_state();
3637 return Err(format!("{phase}: deferred O(1) transition failed"));
3638 }
3639 Ok(())
3640 }
3641
3642 pub fn o1_seal(&mut self) {
3646 if let Err(err) = self.o1_seal_checked() {
3647 self.o1_fail(err);
3648 }
3649 }
3650
3651 pub fn set_trace(&mut self, on: bool) {
3653 self.trace = on;
3654 }
3655
3656 pub fn set_sampler_config(&mut self, config: SamplerConfig) {
3659 self.rng = match config.seed {
3660 Some(seed) => SplitMix64::new(seed),
3661 None => SplitMix64::from_entropy(),
3662 };
3663 self.sampler_config = config;
3664 }
3665
3666 pub fn set_confidence(&mut self, on: bool) {
3671 self.confidence_on = on;
3672 }
3673
3674 pub fn set_calib_temp(&mut self, t: f32) {
3677 self.calib_temp = if t > 1e-3 { t } else { 1.0 };
3678 }
3679
3680 pub fn calib_temp(&self) -> f32 {
3682 self.calib_temp
3683 }
3684
3685 pub fn set_rotary(&mut self, rotary_dim: usize, base: f32) {
3688 self.rotary_dim = rotary_dim.min(self.head_dim);
3689 self.inv_freq = std::sync::Arc::new(attention::rope_inv_freq(self.rotary_dim, base));
3690 self.embryo_graph = None;
3694 }
3695
3696 fn attn_cfg(&self, position: usize) -> QwenAttnCfg<'_> {
3697 QwenAttnCfg {
3698 num_heads: self.num_heads,
3699 num_kv_heads: self.num_kv_heads,
3700 head_dim: self.head_dim,
3701 hidden_size: self.hidden_size,
3702 position,
3703 inv_freq: &self.inv_freq,
3704 rotary_dim: self.rotary_dim,
3705 scale: self.attn_scale,
3706 softcap: self.attn_softcap,
3707 window: None,
3708 v_norm: false,
3709 qk_norm_after_rope: self.qk_norm_after_rope,
3710 gate_sigmoid: self.proj_gate_sigmoid,
3711 q_norm: None,
3712 k_norm: None,
3713 output_gate: false,
3714 softplus_gate: None,
3715 rope_scale: self.rope_scale,
3716 bias: None,
3717 rms_eps: self.rms_eps,
3718 norm_style: self.norm_style,
3719 pool: self.pool.as_deref(),
3720 v_head_dim: self.v_head_dim.unwrap_or(self.head_dim),
3721 }
3722 }
3723
3724 pub fn generate(
3726 &mut self,
3727 prompt: &str,
3728 max_tokens: usize,
3729 task_mask: Option<&TaskMask>,
3730 on_token: Option<TokenCallback>,
3731 ) -> Result<GenerateResult, String> {
3732 let input_ids = self.tokenizer.with_bos(self.tokenizer.encode(prompt));
3733 self.generate_from_ids(&input_ids, max_tokens, task_mask, on_token)
3734 }
3735
3736 pub fn generate_from_vl(
3739 &mut self,
3740 input: &crate::dsv41_vision::PreparedVlInputs,
3741 max_tokens: usize,
3742 task_mask: Option<&TaskMask>,
3743 on_token: Option<TokenCallback>,
3744 ) -> Result<GenerateResult, String> {
3745 let Some(dsv41) = &self.dsv41 else {
3746 return Err("V4.1 multimodal input requires a DeepSeek-V4.1 pipeline".into());
3747 };
3748 if input.token_ids.is_empty() {
3749 return Err("empty V4.1 multimodal prompt".into());
3750 }
3751 if input.token_types.len() != input.token_ids.len() {
3752 return Err(format!(
3753 "V4.1 token type count {} != token count {}",
3754 input.token_types.len(),
3755 input.token_ids.len()
3756 ));
3757 }
3758 let dim = dsv41.2.dim;
3759 let mut embeddings = vec![None; input.token_ids.len()];
3760 let mut participates = vec![true; input.token_ids.len()];
3761 if !input.images.is_empty() {
3762 let vision = self
3763 .dsv41_vision
3764 .as_ref()
3765 .ok_or_else(|| "V4.1 image prompt has no loaded vision tower".to_string())?;
3766 for image in &input.images {
3767 let end = image.start.saturating_add(image.types.len());
3768 if end > input.token_ids.len() {
3769 return Err(format!(
3770 "V4.1 image span {}..{} exceeds prompt length {}",
3771 image.start,
3772 end,
3773 input.token_ids.len()
3774 ));
3775 }
3776 let mut span = vec![0.0f32; image.types.len() * dim];
3777 vision.fill_image_span(image, &mut span, self.pool.as_deref())?;
3778 for (offset, &kind) in image.types.iter().enumerate() {
3779 let pos = image.start + offset;
3780 if input.token_types[pos] != kind {
3781 return Err(format!(
3782 "V4.1 image type mismatch at position {pos}: {} != {kind}",
3783 input.token_types[pos]
3784 ));
3785 }
3786 embeddings[pos] = Some(span[offset * dim..(offset + 1) * dim].to_vec());
3787 participates[pos] = false;
3788 }
3789 }
3790 }
3791 for (pos, &kind) in input.token_types.iter().enumerate() {
3792 if kind == crate::dsv41_vision::TEXT && embeddings[pos].is_some() {
3793 return Err(format!("V4.1 text position {pos} has an image embedding"));
3794 }
3795 if kind != crate::dsv41_vision::TEXT && embeddings[pos].is_none() {
3796 return Err(format!("V4.1 image position {pos} has no image embedding"));
3797 }
3798 }
3799 self.dsv41_prefill = Some((embeddings, participates));
3800 let result = self.generate_from_ids(&input.token_ids, max_tokens, task_mask, on_token);
3801 self.dsv41_prefill = None;
3802 result
3803 }
3804
3805 fn drop_open_mask<'m>(&self, m: Option<&'m TaskMask>) -> Option<&'m TaskMask> {
3807 m.filter(|m| !m.fully_open(self.intermediate_size, self.num_heads))
3808 }
3809
3810 pub fn generate_from_ids(
3818 &mut self,
3819 input_ids: &[u32],
3820 max_tokens: usize,
3821 task_mask: Option<&TaskMask>,
3822 on_token: Option<TokenCallback>,
3823 ) -> Result<GenerateResult, String> {
3824 self.generate_with_prompt_rows(input_ids, None, max_tokens, task_mask, on_token)
3825 }
3826
3827 pub fn generate_from_embeds(
3833 &mut self,
3834 input_ids: &[u32],
3835 prompt_rows: &[f32],
3836 max_tokens: usize,
3837 task_mask: Option<&TaskMask>,
3838 on_token: Option<TokenCallback>,
3839 ) -> Result<GenerateResult, String> {
3840 if input_ids.is_empty()
3841 || input_ids.len().checked_mul(self.hidden_size) != Some(prompt_rows.len())
3842 {
3843 return Err("embedded prompt dimensions must be [tokens, hidden_size]".into());
3844 }
3845 if prompt_rows.iter().any(|x| !x.is_finite()) {
3846 return Err("embedded prompt contains non-finite values".into());
3847 }
3848 if !self.can_prefill_batched() || self.dyn_router.is_some()
3849 || self.o1_active() || self.mtp.is_some() || self.gpu_plan.is_some()
3850 {
3851 return Err("embedded prompts require the ordinary transformer path without O(1), dynamic routing, GPU splitting or a generic MTP head".into());
3852 }
3853 self.generate_with_prompt_rows(input_ids, Some(prompt_rows), max_tokens, task_mask, on_token)
3854 }
3855
3856 fn generate_with_prompt_rows(
3857 &mut self,
3858 input_ids: &[u32],
3859 prompt_rows: Option<&[f32]>,
3860 max_tokens: usize,
3861 task_mask: Option<&TaskMask>,
3862 mut on_token: Option<TokenCallback>,
3863 ) -> Result<GenerateResult, String> {
3864 #[cfg(target_os = "macos")]
3865 crate::gpu_metal::set_io_namespace(self.graph_kv_id);
3866 if std::env::var("CMF_TRACE_H").is_ok() {
3867 eprintln!("input_ids: {input_ids:?}");
3868 }
3869 if input_ids.is_empty() {
3870 return Err("empty prompt: nothing to generate from".to_string());
3871 }
3872 self.graph_failed
3876 .store(false, std::sync::atomic::Ordering::Relaxed);
3877 let task_mask = self.drop_open_mask(task_mask);
3882
3883 let mut reuse_from = {
3891 let on = !std::env::var("CMF_KV_REUSE").is_ok_and(|v| v == "0");
3892 if on
3893 && prompt_rows.is_none()
3894 && task_mask.is_none()
3895 && self.mtp.is_none()
3896 && !(self.mimo_mtp.is_some() && self.speculative)
3897 && self.o1_cfg.is_none()
3898 && self.dsv41.is_none()
3899 {
3900 self.cached_prefix_len(input_ids)
3901 } else {
3902 0
3903 }
3904 };
3905 if reuse_from > 0 && !self.prepare_kv_reuse(reuse_from) {
3908 reuse_from = 0;
3909 }
3910 self.last_prefill_tokens = input_ids.len() - reuse_from;
3911 let bounded_native = self.bounded_native();
3912 if reuse_from == 0 {
3913 self.clear_sequence_state();
3915 } else if std::env::var("CMF_PREFILL_PROF").is_ok() {
3916 eprintln!(
3917 "kv-reuse: {} of {} prompt positions already cached",
3918 reuse_from,
3919 input_ids.len()
3920 );
3921 }
3922 crate::gpu::graph_race_begin_generation();
3923 let o1_prefill = if self.o1_active() && task_mask.is_none() {
3927 std::env::var("CMF_O1_PREFILL")
3928 .ok()
3929 .and_then(|v| v.parse::<usize>().ok())
3930 .filter(|&p| p > 0)
3931 } else {
3932 None
3933 };
3934 if task_mask.is_none() {
3935 self.o1_begin_with_prefix(o1_prefill);
3936 }
3937
3938 let graph_on = crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Decode);
3944 #[cfg(target_os = "macos")]
3983 let metal_graph = crate::gpu::q1_force()
3984 && crate::gpu::enabled_here()
3985 && std::env::var("CMF_GPU_BLOCK")
3986 .map(|v| v != "0")
3987 .unwrap_or(true);
3988 #[cfg(not(target_os = "macos"))]
3989 let metal_graph = false;
3990 let spec_sample_env = std::env::var("CMF_GRAPH_SPEC_SAMPLE").ok();
3991 let spec_cheap_round = self.sampler_config.temperature < 1e-6
3995 || sampler::sparse_ok(&self.sampler_config);
3996 let spec_sampling_ok = self.sampler_config.temperature < 1e-6
3997 || match spec_sample_env.as_deref() {
3998 Some("1") => true,
3999 Some(_) => false,
4000 None => metal_graph && spec_cheap_round,
4001 };
4002 let (mut dense_n, mut dense_q4tp) = (0usize, 0usize);
4018 for lw in &self.weights.layers {
4019 if let FfnKind::Dense(d) = &lw.ffn {
4020 dense_n += 1;
4021 if matches!(d.gate_proj.graph_weight(), Some((_, _, 6, _)))
4022 && matches!(d.up_proj.graph_weight(), Some((_, _, 6, _)))
4023 && matches!(d.down_proj.graph_weight(), Some((_, _, 6, _)))
4024 {
4025 dense_q4tp += 1;
4026 }
4027 }
4028 }
4029 let spec_default_ok = dense_n == 0 || dense_q4tp * 10 >= dense_n * 9;
4030 let penalized = !metal_graph
4042 && (self.sampler_config.repetition_penalty != 1.0
4043 || self.sampler_config.presence_penalty != 0.0
4044 || !self.sampler_config.suppress_tokens.is_empty());
4045 #[cfg(feature = "gpu")]
4050 let metal_wgpu = graph_on && crate::gpu_wgpu::wgpu_backend_is_metal();
4051 #[cfg(not(feature = "gpu"))]
4052 let metal_wgpu = false;
4053 let spec_env = std::env::var("CMF_GRAPH_SPEC").ok();
4054 let spec_wanted = match spec_env.as_deref() {
4055 Some("0") => false,
4056 Some(_) => {
4057 if metal_wgpu {
4058 tracing::warn!(
4059 "CMF_GRAPH_SPEC forced on wgpu/Metal: the batched verify graph is not \
4060 verified on this backend (garbage measured on Qwen3.5-0.8B)"
4061 );
4062 }
4063 true
4064 }
4065 None => spec_default_ok && !penalized && !metal_wgpu,
4066 };
4067 let graph_spec = self.speculative
4071 && (graph_on || metal_graph)
4072 && self.mtp.is_some()
4073 && task_mask.is_none()
4074 && !self.o1_active()
4075 && spec_sampling_ok
4076 && spec_wanted;
4077 #[cfg(target_os = "macos")]
4081 if metal_graph {
4082 static SAID: std::sync::Once = std::sync::Once::new();
4083 SAID.call_once(|| {
4084 let spec = if graph_spec {
4085 let k = std::env::var("CMF_GRAPH_SPEC_K")
4086 .ok()
4087 .and_then(|v| v.parse::<usize>().ok())
4088 .filter(|&v| (1..=8).contains(&v))
4089 .unwrap_or(7);
4090 let arm = if self.sampler_config.temperature < 1e-6 {
4091 "greedy"
4092 } else {
4093 "sampling"
4094 };
4095 format!(
4096 "spec k={k} {arm} (batched verify, draft shortlist {}, trial: proxy)",
4097 Self::draft_vocab_rows(usize::MAX)
4098 )
4099 } else if !self.speculative {
4100 "spec off (CMF_MTP=0)".to_string()
4101 } else if self.mtp.is_none() {
4102 "spec off (no MTP head)".to_string()
4103 } else if !spec_sampling_ok {
4104 if spec_cheap_round {
4105 "spec off (CMF_GRAPH_SPEC_SAMPLE=0)".to_string()
4106 } else {
4107 "spec off (sampling without a top-k: the dense chain \
4108 costs more than it saves)"
4109 .to_string()
4110 }
4111 } else if !spec_wanted {
4112 "spec off (CMF_GRAPH_SPEC=0 or non-q4tp FFNs)".to_string()
4113 } else if task_mask.is_some() {
4114 "spec off (task mask)".to_string()
4115 } else {
4116 "spec off (O(1) attention)".to_string()
4117 };
4118 let on = |var: &str| {
4119 if std::env::var(var).as_deref() == Ok("0") {
4120 "off"
4121 } else {
4122 "on"
4123 }
4124 };
4125 tracing::info!(
4126 "metal native: {spec}, state4 {}, async replay {}, prefill graph {}, \
4127 MTP graph {}, attend {}, probe {}",
4128 if crate::gpu_metal::state4_on() { "on" } else { "off" },
4129 if crate::gpu_metal::async_replay_on() { "on" } else { "off" },
4130 on("CMF_METAL_PREFILL"),
4131 on("CMF_MTP_GRAPH"),
4132 std::env::var("CMF_GPU_ATTEND").unwrap_or_else(|_| "auto".into()),
4133 if crate::gpu::probe_enabled() { "bypassed (q1 force)" } else { "off" },
4134 );
4135 });
4136 }
4137 let pair_pays = self.gdn_cfg.is_none() || std::env::var("CMF_MTP").as_deref() == Ok("1");
4144 let spec_active = self.speculative
4145 && self.mtp.is_some()
4146 && task_mask.is_none()
4147 && !self.o1_active()
4148 && ((!graph_on && pair_pays && self.sampler_config.temperature < 1e-6) || graph_spec);
4149 let mut mtp = if spec_active { self.mtp.take() } else { None };
4152 if std::env::var("CMF_MTP_CHAIN_PROBE").is_ok() {
4153 eprintln!(
4154 "mtp-probe gate: spec_active={spec_active} mtp={} speculative={} graph_on={graph_on} temp_ok={}",
4155 mtp.is_some(),
4156 self.speculative,
4157 self.sampler_config.temperature < 1e-6,
4158 );
4159 }
4160 if let Some(m) = &mut mtp {
4161 m.kv.clear();
4162 crate::gpu::graph_kv_reset(self.mtp_kv_id());
4164 self.mtp_graph_mode = None;
4165 }
4166 let mimo_spec = self.speculative
4170 && self.mimo_mtp.is_some()
4171 && task_mask.is_none()
4172 && !self.o1_active()
4173 && self.dyn_router.is_none()
4174 && self.sampler_config.temperature < 1e-6
4175 && std::env::var("CMF_MIMO_MTP").as_deref() != Ok("0");
4176 if let Some(st) = self.mimo_mtp.as_mut() {
4177 st.reset();
4178 if mimo_spec && std::env::var_os("CMF_MIMO_MTP_PROBE").is_some() {
4179 Self::mimo_mtp_hist_cap(st, input_ids.len());
4180 }
4181 }
4182 let mut router = if mtp.is_none() {
4186 self.dyn_router.take()
4187 } else {
4188 None
4189 };
4190 let mut reuse_from = reuse_from;
4191 if let Some(r) = &mut router {
4192 r.reset(); self.dyn_phi_seen = 0; if self.dyn_active.is_some() {
4195 let _ = self.set_active_skill(None);
4198 reuse_from = 0;
4199 self.last_prefill_tokens = input_ids.len();
4200 }
4201 }
4202
4203 let mut all_ids = input_ids.to_vec();
4204 let mut generated = 0usize;
4205 let mut finish_reason = "max_tokens".to_string();
4206 let mut drafted = 0usize;
4207 let mut accepted = 0usize;
4208 let mut dsv4_spec_bad = 0usize;
4215 let mut dsv4_spec_retry_at = 0usize;
4216 let mut confidence: Vec<f32> = Vec::new();
4217 let trace_on = self.trace;
4218 let calib_temp = self.calib_temp;
4219 let mut traces: Vec<TokenTrace> = Vec::new();
4220
4221 let mut hidden = vec![0.0f32; self.hidden_size];
4227 let mut pos = reuse_from;
4228 let fuse_lm = mtp.is_none()
4237 && router.is_none()
4238 && std::env::var("CMF_GPU_LMHEAD").as_deref() != Ok("0");
4239 self.graph_logits = None;
4240 self.graph_want_logits = false;
4241 let _tpf = std::time::Instant::now();
4242 let batch_k = self.generation_batch_k();
4243 if let Some(rows) = prompt_rows {
4244 let hs = self.hidden_size;
4245 let chunk = self.prefill_chunk().max(1);
4246 while pos < input_ids.len() && !self.cancel.load(std::sync::atomic::Ordering::Relaxed) {
4247 let end = (pos + chunk).min(input_ids.len());
4248 let hb = match self.prefill_input_rows(
4249 PrefillIn::Hidden(&rows[pos * hs..end * hs]), pos, task_mask,
4250 ) {
4251 Ok(hb) => hb,
4252 Err(err) => {
4253 self.finish_generation(&mut mtp, &mut router, true);
4254 return Err(err);
4255 }
4256 };
4257 if mimo_spec { self.mimo_note_rows(&hb, pos); }
4258 hidden.copy_from_slice(&hb[hb.len() - hs..]);
4259 pos = end;
4260 }
4261 }
4262 while self.qwen4_exp.is_some()
4273 && mtp.is_none()
4274 && pos < input_ids.len()
4275 && !self.cancel.load(std::sync::atomic::Ordering::Relaxed)
4276 {
4277 let end = (pos + crate::qwen4_exp::prefill_chunk()).min(input_ids.len());
4280 let want_logits = end == input_ids.len();
4281 let mut lg = Vec::new();
4282 if let Some(b) = &mut self.qwen4_exp {
4283 crate::qwen4_exp::forward_tokens(
4284 &b.0,
4285 &b.1,
4286 &b.2,
4287 &mut b.3,
4288 &input_ids[pos..end],
4289 pos,
4290 &self.inv_freq,
4291 self.pool.as_deref(),
4292 &mut lg,
4293 want_logits,
4294 );
4295 }
4296 if want_logits {
4297 self.graph_logits = Some(lg);
4298 }
4299 pos = end;
4300 hidden.fill(0.0);
4301 }
4302 while self.dsv4.is_some()
4303 && mtp.is_none()
4304 && pos < input_ids.len()
4305 && !self.cancel.load(std::sync::atomic::Ordering::Relaxed)
4306 {
4307 let end = (pos + prefill_chunk()).min(input_ids.len());
4308 let ids: Vec<u32> = input_ids[pos..end].to_vec();
4309 let mut lg = Vec::new();
4310 if let Some(b) = &mut self.dsv4 {
4311 let (g, layers, cfg, st) = (&b.0, &b.1, b.2, &mut b.3);
4312 crate::dsv4::forward_chunk(
4313 g,
4314 layers,
4315 &cfg,
4316 st,
4317 &ids,
4318 pos,
4319 &self.inv_freq,
4320 self.pool.as_deref(),
4321 &mut lg,
4322 end == input_ids.len(),
4323 );
4324 }
4325 if end == input_ids.len() {
4326 self.graph_logits = Some(lg);
4327 }
4328 pos = end;
4329 hidden = vec![0.0; self.hidden_size];
4330 }
4331 let dsv41_prefill = self.dsv41_prefill.take();
4332 while self.dsv41.is_some()
4333 && mtp.is_none()
4334 && pos < input_ids.len()
4335 && !self.cancel.load(std::sync::atomic::Ordering::Relaxed)
4336 {
4337 let end = (pos + prefill_chunk()).min(input_ids.len());
4338 let ids: Vec<u32> = input_ids[pos..end].to_vec();
4339 let mut lg = Vec::new();
4340 if let Some(b) = &mut self.dsv41 {
4341 let (g, layers, cfg, st) = (&b.0, &b.1, &b.2, &mut b.3);
4342 if let Some((embeddings, participates)) = dsv41_prefill.as_ref() {
4343 crate::dsv41::forward_chunk_masked_with_embeddings(
4344 g,
4345 layers,
4346 cfg,
4347 st,
4348 &ids,
4349 pos,
4350 &embeddings[pos..end],
4351 &participates[pos..end],
4352 self.pool.as_deref(),
4353 &mut lg,
4354 );
4355 } else {
4356 crate::dsv41::forward_chunk(
4357 g,
4358 layers,
4359 cfg,
4360 st,
4361 &ids,
4362 pos,
4363 self.pool.as_deref(),
4364 &mut lg,
4365 );
4366 }
4367 }
4368 if end == input_ids.len() {
4369 self.graph_logits = Some(lg);
4370 }
4371 pos = end;
4372 hidden = vec![0.0; self.hidden_size];
4373 }
4374 let dyn_prefill = router.is_some();
4379 let o1_prefill_limit = o1_prefill
4387 .and_then(|requested| self.o1_effective_boundary(requested))
4388 .map(|boundary| boundary.min(input_ids.len()));
4389 let mut o1_sealed = false;
4390 if let Some(limit) = o1_prefill_limit {
4391 if self.can_prefill_batched() && limit > 2 {
4394 let chunk = self.prefill_chunk();
4395 let hs = self.hidden_size;
4396 while pos < limit && !self.cancel.load(std::sync::atomic::Ordering::Relaxed) {
4397 let end = (pos + chunk).min(limit);
4398 let hb = self.prefill_batch(&input_ids[pos..end], pos);
4399 hidden.copy_from_slice(&hb[(end - pos - 1) * hs..]);
4400 pos = end;
4401 }
4402 } else {
4403 while pos < limit && !self.cancel.load(std::sync::atomic::Ordering::Relaxed) {
4404 hidden = self.forward_layers(&self.embed_single(input_ids[pos]), pos, None);
4405 pos += 1;
4406 }
4407 }
4408 if pos >= limit {
4409 o1_sealed = match self.o1_seal_checked() {
4410 Ok(sealed) => sealed,
4411 Err(err) => {
4412 self.finish_generation(&mut mtp, &mut router, true);
4413 return Err(err);
4414 }
4415 };
4416 tracing::info!(
4417 "o1 bounded prompt prefix: requested={} effective={} processed={} of {} token(s)",
4418 o1_prefill.unwrap_or(0),
4419 self.o1_effective_boundary(o1_prefill.unwrap_or(0))
4420 .unwrap_or(limit),
4421 limit,
4422 input_ids.len()
4423 );
4424 }
4425 }
4426 let graph_prefill = self.graph_prefill_preferred();
4432 #[cfg(target_os = "macos")]
4440 if task_mask.is_none()
4441 && !dyn_prefill
4442 && (crate::gpu::q1_force() || crate::gpu::q2tp_gpu_opt_in())
4443 && crate::gpu::enabled_here()
4444 && self.gdn_cfg.is_some()
4445 && self.g3n.is_none()
4446 && input_ids.len() > 8
4447 && std::env::var("CMF_MTP_CHAIN_PROBE").is_err()
4448 && std::env::var("CMF_METAL_PREFILL").as_deref() != Ok("0")
4449 {
4450 let chunk: usize = std::env::var("CMF_METAL_PREFILL_CHUNK")
4451 .ok()
4452 .and_then(|v| v.parse().ok())
4453 .filter(|&v| (16..=512).contains(&v))
4454 .unwrap_or(256);
4455 let hs = self.hidden_size;
4456 let _tp = std::time::Instant::now();
4457 while pos < input_ids.len() && !self.cancel.load(std::sync::atomic::Ordering::Relaxed) {
4458 let end = (pos + chunk).min(input_ids.len());
4459 let hb = match self.prefill_batch_metal(&input_ids[pos..end], pos) {
4460 MetalPrefillOutcome::Completed(hb) => hb,
4461 MetalPrefillOutcome::Declined => break,
4462 MetalPrefillOutcome::Failed => {
4463 self.finish_generation(&mut mtp, &mut router, true);
4464 return Err("ordinary Metal prefill failed after admission".into());
4465 }
4466 };
4467 if let Some(m) = &mut mtp {
4468 let n_pairs = if end < input_ids.len() {
4469 end - pos
4470 } else {
4471 end - pos - 1
4472 };
4473 if n_pairs > 0 {
4474 let pairs: Vec<(&[f32], u32)> = (0..n_pairs)
4475 .map(|j| (&hb[j * hs..(j + 1) * hs], input_ids[pos + j + 1]))
4476 .collect();
4477 if !self.mtp_warm_batch_metal(m, &pairs, pos) {
4478 for (j, (h, t)) in pairs.iter().enumerate() {
4479 let h = h.to_vec();
4480 let _ = self.mtp_step(m, &h, *t, pos + j);
4481 }
4482 }
4483 }
4484 }
4485 hidden.copy_from_slice(&hb[(end - pos - 1) * hs..]);
4486 pos = end;
4487 }
4488 if std::env::var("CMF_PREFILL_PROF").is_ok() {
4489 eprintln!(
4490 "metal-prefill: {} of {} tokens in {:.1} ms",
4491 pos,
4492 input_ids.len(),
4493 _tp.elapsed().as_secs_f64() * 1e3
4494 );
4495 }
4496 }
4497 self.mimo_moe_prepare();
4498 #[cfg(not(target_os = "macos"))]
4504 if task_mask.is_none()
4505 && !dyn_prefill
4506 && !graph_prefill
4507 && mtp.is_none()
4508 && o1_prefill.is_none()
4509 && !self.o1_active()
4510 && input_ids.len() > 2
4511 && self.batch_prefix_prefill()
4512 {
4513 let chunk = self.prefill_chunk().max(1);
4514 let hs = self.hidden_size;
4515 let t_bp = std::time::Instant::now();
4516 let pos0 = pos;
4517 while pos < input_ids.len() && !self.cancel.load(std::sync::atomic::Ordering::Relaxed) {
4518 let end = (pos + chunk).min(input_ids.len());
4519 let bk = end - pos;
4520 let mut hiddens = vec![0f32; bk * hs];
4521 for (j, &id) in input_ids[pos..end].iter().enumerate() {
4522 hiddens[j * hs..(j + 1) * hs].copy_from_slice(&self.embed_single(id));
4523 }
4524 let positions: Vec<usize> = (pos..end).collect();
4525 let mut run = 0usize;
4526 let outcome = self.try_batch_graph_wgpu_prefix(
4527 &mut hiddens,
4528 &positions,
4529 bk,
4530 None,
4531 Some(&mut run),
4532 );
4533 match outcome {
4534 crate::gpu::BatchGraphOutcome::Completed => {
4535 let hb = if run < self.num_layers {
4536 self.prefill_batch_span(
4537 PrefillIn::Hidden(&hiddens),
4538 pos,
4539 None,
4540 run,
4541 self.num_layers,
4542 )
4543 } else {
4544 hiddens
4545 };
4546 if mimo_spec {
4547 self.mimo_note_rows(&hb, pos);
4548 }
4549 hidden.copy_from_slice(&hb[(bk - 1) * hs..]);
4550 pos = end;
4551 }
4552 crate::gpu::BatchGraphOutcome::Failed => {
4553 self.finish_generation(&mut mtp, &mut router, true);
4554 return Err("batched prefix prefill failed after admission".into());
4555 }
4556 crate::gpu::BatchGraphOutcome::Declined => {
4557 #[cfg(feature = "gpu")]
4560 if pos > pos0 {
4561 self.pull_lagging_host_kv(0, self.num_layers, pos);
4562 }
4563 break;
4564 }
4565 }
4566 }
4567 if std::env::var("CMF_PREFILL_PROF").is_ok() {
4568 eprintln!(
4569 "batch-prefix prefill: {} of {} tokens in {:.1} ms",
4570 pos - pos0,
4571 input_ids.len(),
4572 t_bp.elapsed().as_secs_f64() * 1e3
4573 );
4574 }
4575 }
4576 if task_mask.is_none()
4577 && !dyn_prefill
4578 && !graph_prefill
4579 && self.can_prefill_batched()
4580 && self.g3n.is_none()
4581 && o1_prefill.is_none()
4582 && input_ids.len() > 2
4583 {
4584 let chunk = self.prefill_chunk();
4590 let hs = self.hidden_size;
4591 while pos < input_ids.len() && !self.cancel.load(std::sync::atomic::Ordering::Relaxed) {
4592 let end = (pos + chunk).min(input_ids.len());
4593 let hb = self.prefill_batch(&input_ids[pos..end], pos);
4594 if mimo_spec {
4595 self.mimo_note_rows(&hb, pos);
4596 }
4597 if let Some(m) = &mut mtp {
4598 let probe: usize = std::env::var("CMF_MTP_CHAIN_PROBE")
4599 .ok()
4600 .and_then(|v| v.parse().ok())
4601 .unwrap_or(0);
4602 for p in pos..end {
4603 if p + 1 < input_ids.len() {
4604 if probe >= 1 && p + 2 < input_ids.len() {
4605 let (d1, mut hx) = self.mtp_step_h(
4609 m,
4610 &hb[(p - pos) * hs..(p - pos + 1) * hs],
4611 input_ids[p + 1],
4612 p,
4613 );
4614 let mut ok = d1 == input_ids[p + 2];
4615 Self::chain_probe_note(0, ok);
4616 let mut d_prev = d1;
4617 let mut extra = 0usize;
4618 for j in 1..probe {
4619 if p + 2 + j >= input_ids.len() {
4620 break;
4621 }
4622 let (dj, hj) = self.mtp_step_h(m, &hx, d_prev, p + 1 + j);
4623 extra += 1;
4624 ok = ok && dj == input_ids[p + 2 + j];
4625 Self::chain_probe_note(j, ok);
4626 d_prev = dj;
4627 hx = hj;
4628 }
4629 m.kv.truncate_last(extra);
4630 } else {
4631 let _ = self.mtp_step(
4632 m,
4633 &hb[(p - pos) * hs..(p - pos + 1) * hs],
4634 input_ids[p + 1],
4635 p,
4636 );
4637 }
4638 }
4639 }
4640 }
4641 hidden.copy_from_slice(&hb[(end - pos - 1) * hs..]);
4642 pos = end;
4643 }
4644 }
4645 let pair_off = std::env::var("CMF_PAIR").is_ok_and(|v| v == "0");
4646 if task_mask.is_none()
4647 && !dyn_prefill
4648 && !graph_prefill
4649 && !pair_off
4650 && self.pair_supported()
4651 && o1_prefill.is_none()
4652 {
4653 while pos + 1 < input_ids.len()
4654 && !self.cancel.load(std::sync::atomic::Ordering::Relaxed)
4655 {
4656 let e1 = self.embed_single(input_ids[pos]);
4657 let e2 = self.embed_single(input_ids[pos + 1]);
4658 let (h1, h2) = self.forward_pair(&e1, &e2, pos);
4659 if mimo_spec {
4660 self.mimo_note_rows(&h1, pos);
4661 self.mimo_note_rows(&h2, pos + 1);
4662 }
4663 self.commit_linear_scratch();
4665 if let Some(m) = &mut mtp {
4666 let _ = self.mtp_step(m, &h1, input_ids[pos + 1], pos);
4667 if pos + 2 < input_ids.len() {
4668 let probe: usize = std::env::var("CMF_MTP_CHAIN_PROBE")
4669 .ok()
4670 .and_then(|v| v.parse().ok())
4671 .unwrap_or(0);
4672 if probe >= 1 && pos + 3 < input_ids.len() {
4673 let (d1, mut hx) = self.mtp_step_h(m, &h2, input_ids[pos + 2], pos + 1);
4677 let mut ok = d1 == input_ids[pos + 3];
4678 Self::chain_probe_note(0, ok);
4679 let mut d_prev = d1;
4680 let mut extra = 0usize;
4681 for j in 1..probe {
4682 if pos + 3 + j >= input_ids.len() {
4683 break;
4684 }
4685 let (dj, hj) = self.mtp_step_h(m, &hx, d_prev, pos + 2 + j);
4686 extra += 1;
4687 ok = ok && dj == input_ids[pos + 3 + j];
4688 Self::chain_probe_note(j, ok);
4689 d_prev = dj;
4690 hx = hj;
4691 }
4692 m.kv.truncate_last(extra);
4693 } else {
4694 let _ = self.mtp_step(m, &h2, input_ids[pos + 2], pos + 1);
4695 }
4696 }
4697 }
4698 hidden = h2;
4699 pos += 2;
4700 }
4701 }
4702 let o1_batch_ready = o1_sealed
4715 && o1_prefill.is_some()
4716 && mtp.is_none()
4717 && std::env::var("CMF_O1_GPU").as_deref() == Ok("1")
4718 && (0..self.num_layers).all(|li| {
4719 let cache = &self.kv_cache.layers[self.phys_layer(li)];
4720 cache.o1.is_none() || cache.o1_views().is_some()
4721 });
4722 let mtp_batch_prefill = mtp.is_some()
4727 && graph_prefill
4728 && task_mask.is_none()
4729 && !dyn_prefill
4730 && !self.o1_active()
4731 && std::env::var("CMF_MTP_CHAIN_PROBE").is_err();
4732 if batch_k > 0
4733 && (graph_prefill || o1_batch_ready)
4734 && task_mask.is_none()
4735 && (!self.o1_active() || o1_batch_ready)
4736 && (mtp.is_none() || mtp_batch_prefill)
4737 && !dyn_prefill
4738 && pos + 1 < input_ids.len()
4739 {
4740 let hs = self.hidden_size;
4741 let chunk = batch_k;
4742 while pos < input_ids.len() {
4743 let end = (pos + chunk).min(input_ids.len());
4744 let bk = end - pos;
4745 let mut hiddens = vec![0f32; bk * hs];
4746 for (j, &id) in input_ids[pos..end].iter().enumerate() {
4747 hiddens[j * hs..(j + 1) * hs].copy_from_slice(&self.embed_single(id));
4748 }
4749 let positions: Vec<usize> = (pos..end).collect();
4750 let t_chunk = std::time::Instant::now();
4751 let outcome = self.try_batch_graph_wgpu(&mut hiddens, &positions, bk, None);
4752 let ok_b = outcome == crate::gpu::BatchGraphOutcome::Completed;
4753 if std::env::var("CMF_GRAPH_PROF").is_ok() {
4754 let ms = t_chunk.elapsed().as_secs_f64() * 1000.0;
4755 eprintln!(
4756 "batch-chunk: phase=prompt mode={} k={bk} outcome={outcome:?} {ms:.1} ms ({:.1} tok/s)",
4757 if o1_batch_ready {
4758 "o1"
4759 } else if mtp_batch_prefill {
4760 "ordinary_mtp"
4761 } else {
4762 "ordinary"
4763 },
4764 bk as f64 / (ms / 1000.0)
4765 );
4766 }
4767 {
4768 use std::sync::atomic::{AtomicBool, Ordering};
4769 static SAID: AtomicBool = AtomicBool::new(false);
4770 if !SAID.swap(true, Ordering::Relaxed) {
4771 if ok_b {
4772 tracing::info!(
4773 "batched prefill: ACTIVE mode={} (k={bk})",
4774 if o1_batch_ready {
4775 "o1"
4776 } else if mtp_batch_prefill {
4777 "ordinary_mtp"
4778 } else {
4779 "ordinary"
4780 }
4781 );
4782 } else {
4783 tracing::warn!("batched prefill {:?} — per-position graph", outcome);
4784 }
4785 }
4786 }
4787 if ok_b {
4788 if mimo_spec {
4789 self.mimo_note_rows(&hiddens, pos);
4790 }
4791 if mtp_batch_prefill {
4792 let n_pairs = mtp_prefill_pair_count(pos, end, input_ids.len());
4793 if n_pairs > 0 {
4794 let rows: Vec<Vec<f32>> = (0..n_pairs)
4800 .map(|j| hiddens[j * hs..(j + 1) * hs].to_vec())
4801 .collect();
4802 let pairs: Vec<(&[f32], u32)> = rows
4803 .iter()
4804 .enumerate()
4805 .map(|(j, row)| (row.as_slice(), input_ids[pos + j + 1]))
4806 .collect();
4807 if std::env::var("CMF_GRAPH_PROF").is_ok() {
4808 eprintln!(
4809 "mtp-warm: phase=prompt mode=ordinary_mtp first_pos={} pairs={} last_pos={}",
4810 pos,
4811 n_pairs,
4812 pos + n_pairs - 1,
4813 );
4814 }
4815 let warm_error = if let Some(m) = mtp.as_mut() {
4816 self.mtp_warm_prefill_pairs(m, &pairs, pos).err()
4817 } else {
4818 None
4819 };
4820 if let Some(err) = warm_error {
4821 self.finish_generation(&mut mtp, &mut router, true);
4826 return Err(err.to_string());
4827 }
4828 }
4829 }
4830 hidden.copy_from_slice(&hiddens[(bk - 1) * hs..]);
4831 pos = end;
4832 } else if outcome == crate::gpu::BatchGraphOutcome::Failed {
4833 self.finish_generation(&mut mtp, &mut router, true);
4838 return Err(if o1_batch_ready {
4839 "sealed O(1) batch graph failed after admission".to_string()
4840 } else {
4841 "ordinary recurrent batch graph failed after admission".to_string()
4842 });
4843 } else {
4844 break; }
4846 }
4847 }
4848 if graph_prefill
4852 && task_mask.is_none()
4853 && mtp.is_none()
4854 && !dyn_prefill
4855 && pos == 0
4856 && input_ids.len() > 1
4857 && !self.cancel.load(std::sync::atomic::Ordering::Relaxed)
4858 {
4859 if let Some(lg) = self.embryo_prefill_chunked(input_ids, 0) {
4860 self.graph_logits = Some(lg);
4861 hidden = vec![0.0; self.hidden_size];
4862 pos = input_ids.len();
4863 }
4864 }
4865 while pos < input_ids.len() && !self.cancel.load(std::sync::atomic::Ordering::Relaxed) {
4866 self.graph_want_logits = fuse_lm && pos + 1 == input_ids.len();
4867 hidden = self.forward_layers(&self.embed_single(input_ids[pos]), pos, task_mask);
4868 if mimo_spec {
4869 self.mimo_note_rows(&hidden, pos);
4870 }
4871 if let Some(m) = &mut mtp {
4872 if pos + 1 < input_ids.len() {
4873 let probe: usize = std::env::var("CMF_MTP_CHAIN_PROBE")
4879 .ok()
4880 .and_then(|v| v.parse().ok())
4881 .unwrap_or(0);
4882 if probe >= 1 && pos + 2 < input_ids.len() {
4883 let (d1, mut hx) = self.mtp_step_h(m, &hidden, input_ids[pos + 1], pos);
4884 let mut ok = d1 == input_ids[pos + 2];
4885 Self::chain_probe_note(0, ok);
4886 let mut d_prev = d1;
4887 let mut extra = 0usize;
4888 for j in 1..probe {
4889 if pos + 2 + j >= input_ids.len() {
4890 break;
4891 }
4892 let (dj, hj) = self.mtp_step_h(m, &hx, d_prev, pos + 1 + j);
4893 extra += 1;
4894 ok = ok && dj == input_ids[pos + 2 + j];
4895 Self::chain_probe_note(j, ok);
4896 d_prev = dj;
4897 hx = hj;
4898 }
4899 m.kv.truncate_last(extra);
4902 } else {
4903 let _ = self.mtp_step(m, &hidden, input_ids[pos + 1], pos);
4904 }
4905 }
4906 }
4907 pos += 1;
4908 }
4909 if std::env::var("CMF_PREFILL_PROF").is_ok() {
4910 eprintln!(
4911 "prefill: {} tokens in {:.1} ms (batch_k={batch_k})",
4912 input_ids.len(),
4913 _tpf.elapsed().as_secs_f64() * 1000.0
4914 );
4915 }
4916 if self
4917 .graph_failed
4918 .swap(false, std::sync::atomic::Ordering::Relaxed)
4919 {
4920 self.finish_generation(&mut mtp, &mut router, true);
4925 return Err("GPU token graph failed during prefill".to_string());
4926 }
4927 if self
4930 .cancel
4931 .swap(false, std::sync::atomic::Ordering::Relaxed)
4932 {
4933 self.finish_generation(&mut mtp, &mut router, true);
4937 return Ok(GenerateResult {
4938 text: String::new(),
4939 token_ids: Vec::new(),
4940 prompt_tokens: input_ids.len(),
4941 tokens_generated: 0,
4942 finish_reason: "cancelled".to_string(),
4943 mtp_drafted: 0,
4944 mtp_accepted: 0,
4945 token_confidence: Vec::new(),
4946 traces: Vec::new(),
4947 });
4948 }
4949
4950 if !o1_sealed {
4953 match self.o1_seal_checked() {
4954 Ok(_) => {}
4955 Err(err) => {
4956 self.finish_generation(&mut mtp, &mut router, true);
4957 return Err(err);
4958 }
4959 }
4960 }
4961
4962 macro_rules! commit {
4964 ($id:expr) => {{
4965 all_ids.push($id);
4966 generated += 1;
4967 self.note_draft_id($id);
4968 if self.tokenizer.is_eos($id) && !self.ignore_eos {
4969 finish_reason = "stop".to_string();
4970 false
4971 } else {
4972 let token_text = self.tokenizer.decode_token($id);
4973 let mut go = true;
4974 if let Some(ref mut cb) = on_token {
4975 if !cb(&token_text) {
4976 finish_reason = "cancelled".to_string();
4977 go = false;
4978 }
4979 }
4980 go
4981 }
4982 }};
4983 }
4984
4985 let mut spec_trial = SpecTrial::Spec {
4996 t0: std::time::Instant::now(),
4997 gen0: generated,
4998 rounds: 0,
4999 };
5000 let mut spec_mon = SpecMon {
5006 metal: graph_spec && crate::gpu::q1_force() && spec_cheap_round,
5007 ..SpecMon::default()
5008 };
5009 let mut spec_watchdog_off = false;
5010 let mut spec_walls: Vec<f32> = Vec::new();
5013 let mut spec_round_end: Option<std::time::Instant> = None;
5016 if mimo_spec {
5017 if let Ok(path) = std::env::var("CMF_MIMO_MTP_PROBE") {
5018 if let Some(mut st) = self.mimo_mtp.take() {
5019 self.mimo_mtp_probe(&mut st, input_ids, &path);
5020 self.mimo_mtp = Some(st);
5021 }
5022 }
5023 }
5024 let mut next_pos = input_ids.len();
5026 'decode: while generated < max_tokens {
5027 if self
5028 .graph_failed
5029 .swap(false, std::sync::atomic::Ordering::Relaxed)
5030 {
5031 self.finish_generation(&mut mtp, &mut router, true);
5036 return Err("GPU token graph failed during decode".to_string());
5037 }
5038 if self
5039 .cancel
5040 .swap(false, std::sync::atomic::Ordering::Relaxed)
5041 {
5042 finish_reason = "cancelled".to_string();
5043 break 'decode;
5044 }
5045 if mimo_spec && next_pos > 0 {
5050 self.mimo_note_rows(&hidden, next_pos - 1);
5053 }
5054 let forced = self.spec_forced.take();
5055 let mut logits = match (forced, self.graph_logits.take()) {
5056 (Some(_), _) => Vec::new(),
5057 (None, Some(lg)) => lg,
5058 (None, None) => {
5059 let _prof = crate::cpuprof::time(crate::cpuprof::Slot::Head);
5060 inference::rms_norm_into(
5061 &hidden,
5062 &self.weights.final_norm,
5063 self.rms_eps,
5064 self.norm_style,
5065 &mut self.ws.n1,
5066 );
5067 self.lm_head_forward(&self.ws.n1)
5068 }
5069 };
5070 if generated
5073 == std::env::var("CMF_LOGIT_DUMP_STEP")
5074 .ok()
5075 .and_then(|v| v.parse().ok())
5076 .unwrap_or(0)
5077 {
5078 if let Ok(path) = std::env::var("CMF_LOGIT_DUMP") {
5079 let mut bytes: Vec<u8> = Vec::with_capacity((hidden.len() + logits.len()) * 4);
5080 for v in hidden.iter().chain(logits.iter()) {
5081 bytes.extend_from_slice(&v.to_le_bytes());
5082 }
5083 if let Err(e) = std::fs::write(&path, &bytes) {
5084 eprintln!("logit dump: failed to write {path}: {e}");
5085 self.finish_generation(&mut mtp, &mut router, true);
5086 return Err(format!("logit dump write failed: {e}"));
5087 }
5088 }
5089 }
5090 if let Ok(dir) = std::env::var("CMF_LOGIT_DUMP_ALL") {
5094 if !logits.is_empty() {
5095 let path = std::path::Path::new(&dir).join(format!("step{generated:05}.f32"));
5096 let bytes: Vec<u8> = logits.iter().flat_map(|v| v.to_le_bytes()).collect();
5097 if let Err(e) =
5098 std::fs::create_dir_all(&dir).and_then(|_| std::fs::write(&path, &bytes))
5099 {
5100 eprintln!("logit dump: failed to write {}: {e}", path.display());
5101 }
5102 }
5103 }
5104 let t_next = match forced {
5105 Some(c) => c,
5106 None => {
5107 let _prof = crate::cpuprof::time(crate::cpuprof::Slot::Sampler);
5108 sampler::sample_with_scratch_pool(
5109 &logits,
5110 &self.sampler_config,
5111 self.sampler_config.penalty_past(&all_ids, bounded_native),
5112 &mut self.rng,
5113 &mut self.sampler_scratch,
5114 self.pool.as_deref(),
5115 )
5116 }
5117 };
5118 if self.confidence_on {
5119 confidence.push(if logits.is_empty() {
5120 0.0
5121 } else {
5122 sampler::top1_prob_pool(
5123 self.pool.as_deref(),
5124 &mut self.sampler_scratch,
5125 &logits,
5126 t_next,
5127 calib_temp,
5128 )
5129 });
5130 }
5131 if !logits.is_empty() {
5132 attention::recycle_buf(&mut logits);
5133 }
5134 if trace_on {
5135 let skill = router.as_ref().and_then(|r| r.active_id());
5139 traces.push(TokenTrace {
5140 t: generated,
5141 token_id: t_next,
5142 confidence: confidence.last().copied().unwrap_or(0.0),
5143 active_skill: skill,
5144 recon: None,
5145 switched: false,
5146 });
5147 }
5148 if !commit!(t_next) {
5149 break 'decode;
5150 }
5151 if generated >= max_tokens {
5152 break 'decode;
5153 }
5154
5155 self.swa_trim_tails();
5158 if self.dsv41.is_none() && self.kv_cache.needs_eviction() {
5159 static SAID: std::sync::Once = std::sync::Once::new();
5165 SAID.call_once(|| {
5166 tracing::warn!(
5167 "KV cache full at {} positions — evicting half; quality \
5168 will degrade. Raise CMF_MAX_SEQ.",
5169 self.kv_cache.max_seq_len,
5170 );
5171 });
5172 let keep = (self.kv_cache.max_seq_len / 2).max(1);
5173 self.kv_cache.evict(keep);
5174 }
5175
5176 if graph_spec {
5179 match spec_trial {
5180 SpecTrial::Plain { t0, gen0 } if spec_mon.plain_done(t0, gen0, generated) => {
5181 spec_mon.plain_ms =
5182 t0.elapsed().as_secs_f64() * 1e3 / (generated - gen0) as f64;
5183 let keep = spec_mon.pays();
5184 tracing::info!(
5185 "speculation trial: {:.2} tok/round in {:.1} ms vs plain {:.1} ms/tok — {}",
5186 spec_mon.tokens,
5187 spec_mon.round_ms,
5188 spec_mon.plain_ms,
5189 if keep { "speculating" } else { "plain" }
5190 );
5191 spec_mon.fails = 0;
5192 spec_trial = SpecTrial::Decided {
5193 spec: keep,
5194 recheck_at: if keep { usize::MAX } else { generated + 128 },
5195 };
5196 }
5197 SpecTrial::Decided { recheck_at, .. } if generated >= recheck_at => {
5198 spec_mon.n = 0;
5199 spec_trial = SpecTrial::Spec {
5200 t0: std::time::Instant::now(),
5201 gen0: generated,
5202 rounds: 0,
5203 };
5204 }
5205 _ => {}
5206 }
5207 spec_watchdog_off = matches!(
5208 spec_trial,
5209 SpecTrial::Plain { .. } | SpecTrial::Decided { spec: false, .. }
5210 );
5211 }
5212 if mimo_spec && generated + 1 < max_tokens && next_pos > 0 {
5214 let budget = max_tokens - generated - 1;
5215 if let Some(mut st) = self.mimo_mtp.take() {
5216 let k = st.depth.min(budget);
5217 let r = self.mimo_spec_round(&mut st, next_pos, &all_ids, k);
5218 self.mimo_mtp = Some(st);
5219 let r = match r {
5220 Ok(r) => r,
5221 Err(err) => {
5222 self.finish_generation(&mut mtp, &mut router, true);
5223 return Err(err);
5224 }
5225 };
5226 if let Some(r) = r {
5227 drafted += r.drafted;
5228 accepted += r.accepted.len();
5229 let mut stopped = false;
5230 for &id in &r.accepted {
5231 if self.confidence_on {
5232 confidence.push(0.0);
5233 }
5234 if !commit!(id) {
5235 stopped = true;
5236 break;
5237 }
5238 }
5239 if stopped {
5240 break 'decode;
5241 }
5242 next_pos += r.accepted.len() + 1;
5243 hidden = r.hidden;
5244 self.graph_logits = Some(r.logits);
5247 continue 'decode;
5248 }
5249 }
5250 }
5251 #[cfg(feature = "gpu")]
5254 if self.speculative
5255 && self.qwen4_exp.is_some()
5256 && task_mask.is_none()
5257 && self.sampler_config.temperature < 1e-6
5258 && generated + 1 < max_tokens
5259 && next_pos > 0
5260 && std::env::var("CMF_QWEN_MTP").as_deref() != Ok("0")
5261 {
5262 let r = match &mut self.qwen4_exp {
5263 Some(b) => crate::qwen4_exp::spec_round(
5264 &b.0,
5265 &b.1,
5266 &b.2,
5267 &mut b.3,
5268 next_pos,
5269 &all_ids,
5270 &self.inv_freq,
5271 self.pool.as_deref(),
5272 ),
5273 None => None,
5274 };
5275 if let Some(r) = r {
5276 drafted += r.drafted;
5277 accepted += r.accepted.len();
5278 let mut stopped = false;
5279 for &id in &r.accepted {
5280 if self.confidence_on {
5281 confidence.push(0.0);
5282 }
5283 if !commit!(id) {
5284 stopped = true;
5285 break;
5286 }
5287 }
5288 if stopped {
5289 break 'decode;
5290 }
5291 next_pos += r.accepted.len() + 1;
5292 hidden.fill(0.0);
5293 self.graph_logits = Some(r.logits);
5294 continue 'decode;
5295 }
5296 }
5297 match &mut mtp {
5298 #[cfg(feature = "gpu")]
5300 Some(m)
5301 if graph_spec
5302 && !spec_watchdog_off
5303 && generated + 1 < max_tokens
5304 && next_pos > 0 =>
5305 {
5306 let t_round = std::time::Instant::now();
5307 if spec_time_level() >= 2 {
5308 if let Some(t) = spec_round_end.take() {
5309 eprintln!(
5310 "spec-gap {:.2} ms (host between rounds)",
5311 t.elapsed().as_secs_f64() * 1e3
5312 );
5313 }
5314 }
5315 spec_stamps_begin();
5316 #[cfg(target_os = "macos")]
5321 let allocs0 = crate::gpu_metal::IO_BUF_ALLOCS
5322 .load(std::sync::atomic::Ordering::Relaxed);
5323 #[cfg(not(target_os = "macos"))]
5324 let allocs0 = 0u64;
5325 if let Some((extra, n_pos, new_h)) = self.graph_spec_step(
5326 m,
5327 &hidden,
5328 t_next,
5329 next_pos,
5330 &mut drafted,
5331 &mut accepted,
5332 &mut all_ids,
5333 max_tokens - generated,
5334 ) {
5335 next_pos = n_pos;
5336 hidden = new_h;
5337 let level = spec_time_level();
5338 if level > 0 {
5339 let wall = t_round.elapsed().as_secs_f32() * 1e3;
5340 let stamps = spec_stamps_take();
5341 let median = if spec_walls.len() >= 3 {
5344 let mut s = spec_walls.clone();
5345 s.sort_by(|a, b| a.partial_cmp(b).unwrap());
5346 Some(s[s.len() / 2])
5347 } else {
5348 None
5349 };
5350 let outlier = median.is_some_and(|m| wall > 1.4 * m);
5351 #[cfg(target_os = "macos")]
5352 let allocs = crate::gpu_metal::IO_BUF_ALLOCS
5353 .load(std::sync::atomic::Ordering::Relaxed)
5354 - allocs0;
5355 #[cfg(not(target_os = "macos"))]
5356 let allocs = allocs0;
5357 eprintln!(
5358 "spec-round wall {wall:.1} ms → {} tokens{}{}",
5359 extra.len() + 1,
5360 if allocs > 0 {
5361 format!(" [{allocs} new device buffers]")
5362 } else {
5363 String::new()
5364 },
5365 match (outlier, median) {
5366 (true, Some(m)) => format!(" OUTLIER (median {m:.1})"),
5367 _ => String::new(),
5368 }
5369 );
5370 if level >= 2 || outlier {
5371 let sum: f32 = stamps.iter().map(|s| s.1).sum();
5372 eprintln!(
5373 "spec-stamps: {}| untracked {:.1}",
5374 spec_stamps_format(&stamps),
5375 wall - sum
5376 );
5377 }
5378 if spec_mon.n >= 1 {
5379 spec_walls.push(wall);
5380 }
5381 }
5382 spec_mon.round(t_round.elapsed().as_secs_f64() * 1e3, extra.len() + 1);
5386 spec_trial = Self::spec_trial_round(
5389 spec_trial,
5390 &mut spec_mon,
5391 generated + extra.len() + 1,
5392 );
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 if spec_time_level() >= 2 {
5407 spec_round_end = Some(std::time::Instant::now());
5408 }
5409 continue 'decode;
5410 }
5411 if self
5412 .graph_failed
5413 .swap(false, std::sync::atomic::Ordering::Relaxed)
5414 {
5415 self.finish_generation(&mut mtp, &mut router, true);
5421 return Err("GPU MTP graph failed during speculative decode".to_string());
5422 }
5423 spec_mon.round(t_round.elapsed().as_secs_f64() * 1e3, 1);
5434 spec_mon.tokens = 0.0;
5435 spec_mon.fails = 3;
5436 spec_trial = Self::spec_trial_round(spec_trial, &mut spec_mon, generated + 1);
5437 hidden = self.forward_layers(&self.embed_single(t_next), next_pos, task_mask);
5438 next_pos += 1;
5439 continue 'decode;
5440 }
5441 Some(m) if !graph_spec && generated + 1 < max_tokens => {
5443 let draft = self.mtp_step(m, &hidden, t_next, next_pos - 1);
5444 drafted += 1;
5445 let emb1 = self.embed_single(t_next);
5446 let emb2 = self.embed_single(draft);
5447 let (h1, h2) = self.forward_pair(&emb1, &emb2, next_pos);
5448
5449 inference::rms_norm_into(
5450 &h1,
5451 &self.weights.final_norm,
5452 self.rms_eps,
5453 self.norm_style,
5454 &mut self.ws.n1,
5455 );
5456 let mut logits1 = self.lm_head_forward(&self.ws.n1);
5457 let t_after = sampler::sample_with_scratch_pool(
5458 &logits1,
5459 &self.sampler_config,
5460 self.sampler_config.penalty_past(&all_ids, bounded_native),
5461 &mut self.rng,
5462 &mut self.sampler_scratch,
5463 self.pool.as_deref(),
5464 );
5465 if self.confidence_on {
5466 confidence.push(sampler::top1_prob_pool(
5467 self.pool.as_deref(),
5468 &mut self.sampler_scratch,
5469 &logits1,
5470 t_after,
5471 calib_temp,
5472 ));
5473 }
5474 attention::recycle_buf(&mut logits1);
5475 if trace_on {
5476 traces.push(TokenTrace {
5479 t: generated,
5480 token_id: t_after,
5481 confidence: confidence.last().copied().unwrap_or(0.0),
5482 active_skill: None,
5483 recon: None,
5484 switched: false,
5485 });
5486 }
5487 let stop = !commit!(t_after);
5488
5489 if t_after == draft {
5490 accepted += 1;
5491 self.commit_linear_scratch();
5492 let _ = self.mtp_step(m, &h1, t_after, next_pos);
5493 hidden = h2;
5494 next_pos += 2;
5495 } else {
5496 for layer in &mut self.kv_cache.layers {
5498 layer.truncate_last(1);
5499 }
5500 if !stop {
5501 let _ = self.mtp_step(m, &h1, t_after, next_pos);
5502 hidden = self.forward_layers(
5503 &self.embed_single(t_after),
5504 next_pos + 1,
5505 None,
5506 );
5507 }
5508 next_pos += 2;
5509 }
5510 if stop {
5511 break 'decode;
5512 }
5513 }
5514 _ => {
5516 #[cfg(feature = "gpu")]
5521 if Self::dsv4_spec_on() && self.dsv4.is_some() {
5522 static SAID: std::sync::Once = std::sync::Once::new();
5523 SAID.call_once(|| {
5524 eprintln!(
5525 "dsv4-spec гейт: mtp={} mask={} router={} trace={} temp={} rep={} ",
5526 !self.dsv4_mtp.is_empty(),
5527 task_mask.is_none(),
5528 router.is_none(),
5529 !trace_on,
5530 self.sampler_config.temperature < 1e-6,
5531 self.sampler_config.repetition_penalty == 1.0,
5532 );
5533 });
5534 }
5535 #[cfg(feature = "gpu")]
5536 if Self::dsv4_spec_on()
5537 && self.dsv4.is_some()
5538 && !self.dsv4_mtp.is_empty()
5539 && task_mask.is_none()
5540 && router.is_none()
5541 && !trace_on
5542 && self.sampler_config.temperature < 1e-6
5543 && self.sampler_config.repetition_penalty == 1.0
5544 && generated + 1 < max_tokens
5545 && all_ids.len() >= 2
5546 && generated >= dsv4_spec_retry_at
5547 {
5548 let tip_token = all_ids[all_ids.len() - 2];
5549 let drafted0 = drafted;
5550 let round = self.dsv4_spec_step(
5551 tip_token,
5552 t_next,
5553 next_pos,
5554 max_tokens.saturating_sub(generated),
5555 &mut drafted,
5556 &mut accepted,
5557 );
5558 if drafted > drafted0 {
5559 let useful = round.as_ref().is_some_and(|(extra, _)| !extra.is_empty());
5560 if useful {
5561 dsv4_spec_bad = 0;
5562 } else {
5563 dsv4_spec_bad += 1;
5564 if dsv4_spec_bad >= 2 {
5565 dsv4_spec_bad = 0;
5566 dsv4_spec_retry_at = generated.saturating_add(32);
5567 tracing::info!(
5568 "dsv4: draft не окупился дважды — точный walk на 32 токена"
5569 );
5570 }
5571 }
5572 }
5573 if let Some((extra, n_pos)) = round {
5574 next_pos = n_pos;
5575 let mut stopped = false;
5576 for &id in &extra {
5577 if self.confidence_on {
5578 confidence.push(0.0);
5579 }
5580 if !commit!(id) {
5581 stopped = true;
5582 break;
5583 }
5584 }
5585 if stopped {
5586 break 'decode;
5587 }
5588 continue 'decode;
5589 }
5590 }
5591 self.graph_want_logits = fuse_lm;
5592 let mut t_fwd = t_next;
5598 let pure_greedy = self.sampler_config.temperature < 1e-6
5599 && self.sampler_config.repetition_penalty == 1.0
5600 && self.sampler_config.suppress_tokens.is_empty();
5601 let burst_k = std::env::var("CMF_MULTISTEP")
5606 .ok()
5607 .and_then(|v| v.parse::<usize>().ok())
5608 .unwrap_or(0);
5609 if pure_greedy
5610 && burst_k >= 1
5611 && fuse_lm
5612 && task_mask.is_none()
5613 && router.is_none()
5614 && !trace_on
5615 && !self.confidence_on
5616 {
5617 let mut stopped = false;
5618 loop {
5619 let room = max_tokens.saturating_sub(generated);
5620 if room <= 2 {
5621 break;
5622 }
5623 let k = burst_k.min(room - 1);
5624 if k < 1 {
5625 break;
5626 }
5627 let Some(ids) = self.try_multi_burst(t_fwd, next_pos, k) else {
5628 if self
5629 .graph_failed
5630 .swap(false, std::sync::atomic::Ordering::Relaxed)
5631 {
5632 self.finish_generation(&mut mtp, &mut router, true);
5633 return Err(
5634 "GPU token graph failed during greedy burst".to_string()
5635 );
5636 }
5637 break;
5638 };
5639 next_pos += k;
5640 for &id in &ids {
5641 if !commit!(id) {
5642 stopped = true;
5643 break;
5644 }
5645 }
5646 if stopped {
5647 break;
5648 }
5649 t_fwd = *ids.last().unwrap();
5650 }
5651 if stopped {
5652 break 'decode;
5653 }
5654 }
5655 #[cfg(target_os = "macos")]
5665 if graph_spec
5666 && spec_watchdog_off
5667 && next_pos > 0
5668 && self.mtp_graph_mode == Some(true)
5669 && crate::gpu::q1_force()
5670 {
5671 if let Some(m) = mtp.as_mut() {
5672 let _ = self.mtp_step_metal(m, &hidden, t_fwd, next_pos - 1, false);
5673 }
5674 }
5675 hidden = self.forward_layers(&self.embed_single(t_fwd), next_pos, task_mask);
5676 next_pos += 1;
5677 if let Some(r) = &mut router {
5680 let phi = self.dyn_phi_ema.clone();
5681 let decision = r.step(&phi, generated);
5682 if let Some(new_active) = decision {
5683 let _ = self.set_active_skill(new_active);
5684 }
5685 if trace_on {
5688 if let Some(last) = traces.last_mut() {
5689 let e = r.last_best_e();
5690 last.recon = e.is_finite().then_some(e);
5691 last.switched = decision.is_some();
5692 }
5693 }
5694 }
5695 }
5696 }
5697 }
5698
5699 let cancelled = finish_reason == "cancelled";
5700 let dyn_switched = router.as_ref().is_some_and(|r| !r.switches.is_empty());
5704 if mimo_spec {
5705 if let Some(st) = self.mimo_mtp.as_ref() {
5706 let line = st.stats.line();
5707 tracing::info!("{line}");
5708 if std::env::var_os("CMF_MIMO_MTP_STATS").is_some() {
5709 eprintln!("{line}");
5710 }
5711 }
5712 }
5713 self.finish_generation(&mut mtp, &mut router, cancelled);
5714
5715 let output_ids = &all_ids[input_ids.len()..];
5716 let forwarded = input_ids.len() + output_ids.len().saturating_sub(1);
5720 let consumed = std::mem::take(&mut all_ids);
5725 if dyn_switched {
5726 self.clear_sequence_state();
5727 } else if cancelled || mimo_spec || prompt_rows.is_some() {
5728 self.clear_history();
5729 } else {
5730 self.record_consumed_prefix(&consumed[..forwarded.min(consumed.len())], reuse_from);
5731 }
5732 all_ids = consumed;
5733 let output_ids = &all_ids[input_ids.len()..];
5734 confidence.truncate(output_ids.len()); traces.truncate(output_ids.len());
5736 Ok(GenerateResult {
5737 text: self.tokenizer.decode(output_ids),
5738 token_ids: output_ids.to_vec(),
5739 prompt_tokens: input_ids.len(),
5740 tokens_generated: generated,
5741 finish_reason,
5742 mtp_drafted: drafted,
5743 mtp_accepted: accepted,
5744 token_confidence: confidence,
5745 traces,
5746 })
5747 }
5748
5749 fn mtp_step(
5753 &mut self,
5754 m: &mut MtpModule,
5755 hidden: &[f32],
5756 next_token: u32,
5757 position: usize,
5758 ) -> u32 {
5759 self.mtp_step_h(m, hidden, next_token, position).0
5760 }
5761
5762 fn chain_probe_note(depth: usize, prefix_ok: bool) {
5766 use std::sync::Mutex;
5767 static T: Mutex<Vec<(u64, u64)>> = Mutex::new(Vec::new());
5768 let mut t = T.lock().unwrap();
5769 if t.len() <= depth {
5770 t.resize(depth + 1, (0, 0));
5771 }
5772 t[depth].0 += 1;
5773 t[depth].1 += prefix_ok as u64;
5774 if depth == 0 && t[0].0 % 128 == 0 {
5775 let line: Vec<String> = t
5776 .iter()
5777 .enumerate()
5778 .map(|(d, (n, k))| {
5779 format!(
5780 "d{}={:.0}%({n})",
5781 d + 1,
5782 100.0 * *k as f64 / (*n).max(1) as f64
5783 )
5784 })
5785 .collect();
5786 eprintln!("mtp-chain: {}", line.join(" "));
5787 }
5788 }
5789
5790 fn mtp_step_hl(
5798 &mut self,
5799 m: &mut MtpModule,
5800 hidden: &[f32],
5801 next_token: u32,
5802 position: usize,
5803 ) -> (Vec<f32>, Vec<f32>) {
5804 #[cfg(target_os = "macos")]
5809 if self.mtp_graph_mode != Some(false) && crate::gpu::q1_force() {
5810 if let Some(r) = self.mtp_step_metal(m, hidden, next_token, position, true) {
5811 self.mtp_graph_mode = Some(true);
5812 return r;
5813 }
5814 if self.mtp_graph_mode == Some(true) {
5815 tracing::error!("mtp Metal graph failed after admission");
5816 self.clear_sequence_state();
5817 self.graph_failed
5818 .store(true, std::sync::atomic::Ordering::Relaxed);
5819 self.cancel
5820 .store(true, std::sync::atomic::Ordering::Relaxed);
5821 return (Vec::new(), Vec::new());
5822 }
5823 self.mtp_graph_mode = Some(false);
5824 }
5825 #[cfg(feature = "gpu")]
5826 if self.mtp_graph_mode != Some(false) {
5827 if !self.mtp_graph_ok(m) {
5828 if self.mtp_graph_mode == Some(true) {
5829 tracing::error!("mtp graph became unavailable after admission");
5834 self.clear_sequence_state();
5835 self.graph_failed
5836 .store(true, std::sync::atomic::Ordering::Relaxed);
5837 self.cancel
5838 .store(true, std::sync::atomic::Ordering::Relaxed);
5839 return (Vec::new(), Vec::new());
5840 }
5841 self.mtp_graph_mode = Some(false);
5842 } else {
5843 if let Some(r) = self.mtp_step_graph(m, hidden, next_token, position) {
5844 self.mtp_graph_mode = Some(true);
5845 return r;
5846 }
5847 if self.graph_failed.load(std::sync::atomic::Ordering::Relaxed) {
5848 return (Vec::new(), Vec::new());
5855 }
5856 tracing::error!("mtp graph failed or declined after admission");
5860 self.clear_sequence_state();
5861 self.graph_failed
5862 .store(true, std::sync::atomic::Ordering::Relaxed);
5863 self.cancel
5864 .store(true, std::sync::atomic::Ordering::Relaxed);
5865 return (Vec::new(), Vec::new());
5866 }
5867 }
5868 let e = self.embed_single(next_token);
5872 let mut cat = vec![0.0f32; 2 * self.hidden_size];
5873 let (cat_e, cat_h) = cat.split_at_mut(self.hidden_size);
5874 inference::rms_norm_into(&e, &m.enorm, self.rms_eps, self.norm_style, cat_e);
5875 inference::rms_norm_into(hidden, &m.hnorm, self.rms_eps, self.norm_style, cat_h);
5876 let mut x = vec![0.0f32; self.hidden_size];
5877 m.eh_proj.matvec(&cat, &mut x, self.pool.as_deref());
5878
5879 let lw = &m.layer;
5881 inference::rms_norm_into(
5882 &x,
5883 &lw.input_norm,
5884 self.rms_eps,
5885 self.norm_style,
5886 &mut self.ws.n1,
5887 );
5888 let attn = match &lw.attn {
5889 AttnKind::Mla(_) => unreachable!("MLA has no MTP/pair path"),
5891 AttnKind::Kda(_) => unreachable!("KDA has no MTP/pair path"),
5892 AttnKind::Full {
5893 wq,
5894 wk,
5895 wv,
5896 wo,
5897 q_norm,
5898 k_norm,
5899 output_gate,
5900 softplus_gate,
5901 bias,
5902 } => {
5903 let mut cfg = self.attn_cfg(position);
5904 cfg.q_norm = q_norm.as_deref();
5905 cfg.k_norm = k_norm.as_deref();
5906 cfg.output_gate = *output_gate;
5907 cfg.softplus_gate = softplus_gate
5908 .as_ref()
5909 .map(|(gate, per_head)| (gate, *per_head));
5910 cfg.bias = bias
5911 .as_ref()
5912 .map(|(q, k, v)| (q.as_slice(), k.as_slice(), v.as_slice()));
5913 attention::qwen_attention(&self.ws.n1, wq, wk, wv, wo, &mut m.kv, &cfg)
5914 }
5915 AttnKind::Linear(_)
5916 | AttnKind::LinearGdn(_)
5917 | AttnKind::ShortConv(_)
5918 | AttnKind::Bounded(_) => {
5919 unreachable!("MTP block is full attention")
5920 }
5921 };
5922 for (i, &a) in attn.iter().enumerate() {
5923 x[i] += a;
5924 }
5925 inference::rms_norm_into(
5926 &x,
5927 &lw.post_norm,
5928 self.rms_eps,
5929 self.norm_style,
5930 &mut self.ws.p1,
5931 );
5932 let ffn = ffn_forward(&lw.ffn, &self.ws.p1, self.pool.as_deref(), None);
5933 for (i, &f) in ffn.iter().enumerate() {
5934 x[i] += f;
5935 }
5936
5937 inference::rms_norm_into(
5938 &x,
5939 &m.final_norm,
5940 self.rms_eps,
5941 self.norm_style,
5942 &mut self.ws.n1,
5943 );
5944 let lg = self.lm_head_forward(&self.ws.n1);
5945 (lg, x)
5946 }
5947
5948 fn mtp_step_h(
5950 &mut self,
5951 m: &mut MtpModule,
5952 hidden: &[f32],
5953 next_token: u32,
5954 position: usize,
5955 ) -> (u32, Vec<f32>) {
5956 let (mut lg, x) = self.mtp_step_hl(m, hidden, next_token, position);
5957 let draft = sampler::argmax(&lg);
5958 attention::recycle_buf(&mut lg);
5959 (draft, x)
5960 }
5961
5962 fn spec_trial_round(trial: SpecTrial, mon: &mut SpecMon, generated: usize) -> SpecTrial {
5968 match trial {
5969 SpecTrial::Spec { t0, gen0, rounds } => {
5970 let rounds = rounds + 1;
5971 if rounds >= 5 {
5972 if mon.plain_ms > 0.0 {
5973 let keep = mon.pays();
5974 mon.fails = 0;
5975 tracing::info!(
5976 "speculation re-check: {:.2} tok/round in {:.1} ms vs plain {:.1} ms/tok — {}",
5977 mon.tokens,
5978 mon.round_ms,
5979 mon.plain_ms,
5980 if keep { "speculating" } else { "plain" }
5981 );
5982 SpecTrial::Decided {
5983 spec: keep,
5984 recheck_at: if keep { usize::MAX } else { generated + 128 },
5985 }
5986 } else if mon.pays() {
5987 mon.fails = 0;
5992 tracing::info!(
5993 "speculation trial: {:.2} tok/round in {:.1} ms — speculating (plain not timed)",
5994 mon.tokens,
5995 mon.round_ms,
5996 );
5997 SpecTrial::Decided {
5998 spec: true,
5999 recheck_at: usize::MAX,
6000 }
6001 } else {
6002 SpecTrial::Plain {
6003 t0: std::time::Instant::now(),
6004 gen0: generated,
6005 }
6006 }
6007 } else {
6008 SpecTrial::Spec { t0, gen0, rounds }
6009 }
6010 }
6011 SpecTrial::Decided { spec: true, .. } => {
6012 if mon.pays() {
6013 mon.fails = 0;
6014 trial
6015 } else {
6016 mon.fails += 1;
6017 if mon.fails >= 4 {
6018 if mon.plain_ms <= 0.0 {
6019 tracing::info!(
6023 "speculation doubtful: {:.2} tok/round in {:.1} ms — timing plain",
6024 mon.tokens,
6025 mon.round_ms,
6026 );
6027 return SpecTrial::Plain {
6028 t0: std::time::Instant::now(),
6029 gen0: generated,
6030 };
6031 }
6032 tracing::info!(
6033 "speculation stopped: {:.2} tok/round in {:.1} ms vs plain {:.1} ms/tok",
6034 mon.tokens,
6035 mon.round_ms,
6036 mon.plain_ms
6037 );
6038 SpecTrial::Decided {
6039 spec: false,
6040 recheck_at: generated + 128,
6041 }
6042 } else {
6043 trial
6044 }
6045 }
6046 }
6047 other => other,
6048 }
6049 }
6050
6051 fn mtp_kv_id(&self) -> u64 {
6054 self.graph_kv_id | (1u64 << 40)
6055 }
6056
6057 const MTP_LAYER_BASE: usize = 0;
6062
6063 #[cfg(feature = "gpu")]
6070 fn rewind_mtp_graph_mirror(&self, stored: usize) -> bool {
6071 self.mtp_graph_mode != Some(true)
6072 || crate::gpu::graph_kv_set_stored(self.mtp_kv_id(), Self::MTP_LAYER_BASE, stored)
6073 }
6074
6075 #[cfg(feature = "gpu")]
6081 fn rewind_trunk_graph_mirrors(&self, stored: usize) -> bool {
6082 let mut ok = true;
6083 let mut expected = false;
6084 for li in 0..self.num_layers {
6085 if matches!(
6086 self.weights.layers[self.phys_layer(li)].attn,
6087 AttnKind::Full { .. }
6088 ) {
6089 expected = true;
6090 ok &= crate::gpu::graph_kv_set_stored(self.graph_kv_id, li, stored);
6091 }
6092 }
6093 !expected || ok
6094 }
6095
6096 fn graph_gdn_layer_count(&self) -> usize {
6100 (0..self.num_layers)
6101 .filter(|&li| {
6102 matches!(
6103 &self.weights.layers[self.phys_layer(li)].attn,
6104 AttnKind::LinearGdn(_)
6105 )
6106 })
6107 .count()
6108 }
6109
6110 fn mtp_block_input(&mut self, m: &MtpModule, hidden: &[f32], next_token: u32) -> Vec<f32> {
6113 let e = self.embed_single(next_token);
6114 let mut cat = vec![0.0f32; 2 * self.hidden_size];
6115 let (cat_e, cat_h) = cat.split_at_mut(self.hidden_size);
6116 inference::rms_norm_into(&e, &m.enorm, self.rms_eps, self.norm_style, cat_e);
6117 inference::rms_norm_into(hidden, &m.hnorm, self.rms_eps, self.norm_style, cat_h);
6118 let mut x = vec![0.0f32; self.hidden_size];
6119 m.eh_proj.matvec(&cat, &mut x, self.pool.as_deref());
6120 x
6121 }
6122
6123 #[cfg(feature = "gpu")]
6126 fn mtp_block_graph_ok(&self, m: &MtpModule) -> bool {
6127 if std::env::var("CMF_MTP_GRAPH").as_deref() == Ok("0") {
6128 return false;
6129 }
6130 if !crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Decode)
6131 || !crate::gpu::enabled_here()
6132 || self.attn_softcap > 0.0
6133 || self.attention_heads_per_layer.is_some()
6134 || self.v_head_dim.is_some()
6137 {
6138 return false;
6139 }
6140 matches!(
6141 &m.layer.attn,
6142 AttnKind::Full {
6143 softplus_gate: None,
6144 ..
6145 }
6146 ) && matches!(&m.layer.ffn, FfnKind::Dense(_))
6147 }
6148
6149 #[cfg(feature = "gpu")]
6153 fn mtp_graph_ok(&self, m: &MtpModule) -> bool {
6154 if !self.mtp_block_graph_ok(m) {
6155 return false;
6156 }
6157 let AttnKind::Full { wq, wk, wv, wo, .. } = &m.layer.attn else {
6158 return false;
6159 };
6160 let FfnKind::Dense(d) = &m.layer.ffn else {
6161 return false;
6162 };
6163 d.segs.is_empty()
6164 && wq.graph_weight().is_some()
6165 && wk.graph_weight().is_some()
6166 && wv.graph_weight().is_some()
6167 && wo.graph_weight().is_some()
6168 && d.gate_proj.graph_weight().is_some()
6169 && d.up_proj.graph_weight().is_some()
6170 && d.down_proj.graph_weight().is_some()
6171 && self.weights.lm_head.graph_weight().is_some()
6172 }
6173
6174 #[cfg(feature = "gpu")]
6180 fn mtp_step_graph(
6181 &mut self,
6182 m: &mut MtpModule,
6183 hidden: &[f32],
6184 next_token: u32,
6185 position: usize,
6186 ) -> Option<(Vec<f32>, Vec<f32>)> {
6187 if !self.mtp_graph_ok(m) {
6188 return None;
6189 }
6190 let lw = &m.layer;
6191 let AttnKind::Full {
6192 wq,
6193 wk,
6194 wv,
6195 wo,
6196 q_norm,
6197 k_norm,
6198 output_gate,
6199 softplus_gate,
6200 bias,
6201 } = &lw.attn
6202 else {
6203 return None;
6204 };
6205 if softplus_gate.is_some() {
6206 return None;
6207 }
6208 let FfnKind::Dense(d) = &lw.ffn else {
6209 return None;
6210 };
6211 if !d.segs.is_empty() {
6212 return None; }
6214 let mut x = self.mtp_block_input(m, hidden, next_token);
6217 fn gw(t: &QTensor) -> Option<crate::gpu::GraphW<'_>> {
6218 let (_, i, kind, rs) = t.graph_weight()?;
6219 Some(crate::gpu::GraphW {
6220 idx: i,
6221 kind,
6222 row_scale: rs,
6223 data: &[],
6224 prism: crate::gpu::GraphPrismOp::None,
6225 affine: false,
6226 })
6227 }
6228 let (model, _, _, _) = wq.graph_weight()?;
6229 let model = model.clone();
6230 let (lm_gw, lm_rows) = {
6231 let (_, i, kind, rs) = self.weights.lm_head.graph_weight()?;
6232 let rows = if kind == 6 {
6236 self.draft_head_rows(self.weights.lm_head.rows())
6237 } else {
6238 self.weights.lm_head.rows()
6239 };
6240 (
6241 crate::gpu::GraphW {
6242 idx: i,
6243 kind,
6244 row_scale: rs,
6245 data: &[],
6246 prism: crate::gpu::GraphPrismOp::None,
6247 affine: false,
6248 },
6249 rows,
6250 )
6251 };
6252 let layer = crate::gpu::GraphLayer {
6253 input_norm: &lw.input_norm,
6254 attn: crate::gpu::GraphAttn::Full {
6255 wq: gw(wq)?,
6256 wk: gw(wk)?,
6257 wv: gw(wv)?,
6258 wo: gw(wo)?,
6259 q_norm: q_norm.as_deref(),
6260 k_norm: k_norm.as_deref(),
6261 late_qk_norm: self.qk_norm_after_rope,
6262 bias: bias
6263 .as_ref()
6264 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
6265 output_gate: *output_gate,
6266 cpu_k: m.kv.k_heads(),
6267 cpu_v: m.kv.v_heads(),
6268 cpu_base: m.kv.base(),
6269 geom: None,
6270 head_gate: None,
6271 },
6272 post_norm: &lw.post_norm,
6273 ffn: crate::gpu::GraphFfn::Dense {
6274 gate: gw(&d.gate_proj)?,
6275 up: gw(&d.up_proj)?,
6276 down: gw(&d.down_proj)?,
6277 act: d.act.graph_act()?,
6278 },
6279 };
6280 let nh = self.num_heads;
6281 let (nkv, hd, rd) = self.layer_geom(0);
6282 let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
6283 let mut logits = Vec::new();
6284 let ok = crate::gpu::forward_token_graph(
6285 &model,
6286 self.mtp_kv_id(),
6287 std::slice::from_ref(&layer),
6288 &[None],
6289 self.o1_epoch,
6290 &self.inv_freq,
6291 &mut x,
6292 nh,
6293 nkv,
6294 hd,
6295 self.attn_scale,
6296 rd,
6297 self.hidden_size,
6298 self.intermediate_size,
6299 position,
6300 self.kv_cache.max_seq_len,
6301 gemma,
6302 self.rms_eps as f32,
6303 Some((&lm_gw, lm_rows)),
6304 &m.final_norm,
6305 &mut logits,
6306 &[],
6307 1,
6308 None,
6309 None,
6310 None,
6311 Self::MTP_LAYER_BASE,
6312 true,
6313 );
6314 match ok {
6315 crate::gpu::TokenGraphOutcome::Completed => {}
6316 crate::gpu::TokenGraphOutcome::Declined => return None,
6317 crate::gpu::TokenGraphOutcome::Failed => {
6318 self.clear_sequence_state();
6322 self.graph_failed
6323 .store(true, std::sync::atomic::Ordering::Relaxed);
6324 self.cancel
6325 .store(true, std::sync::atomic::Ordering::Relaxed);
6326 return None;
6327 }
6328 }
6329 logits.resize(self.vocab_size, 0.0);
6330 Some((logits, x))
6331 }
6332
6333 #[cfg(feature = "gpu")]
6341 fn mtp_warm_graph(
6342 &mut self,
6343 m: &mut MtpModule,
6344 pairs: &[(&[f32], u32)],
6345 first_pos: usize,
6346 ) -> crate::gpu::BatchGraphOutcome {
6347 if pairs.is_empty() {
6348 return crate::gpu::BatchGraphOutcome::Completed;
6349 }
6350 if !self.mtp_block_graph_ok(m) {
6351 return crate::gpu::BatchGraphOutcome::Declined;
6352 }
6353 let hs = self.hidden_size;
6354 let mut hiddens = Vec::with_capacity(pairs.len() * hs);
6357 for (h, t) in pairs {
6358 hiddens.extend_from_slice(&self.mtp_block_input(m, h, *t));
6359 }
6360 let lw = &m.layer;
6361 let AttnKind::Full {
6362 wq,
6363 wk,
6364 wv,
6365 wo,
6366 q_norm,
6367 k_norm,
6368 output_gate,
6369 bias,
6370 ..
6371 } = &lw.attn
6372 else {
6373 return crate::gpu::BatchGraphOutcome::Declined;
6374 };
6375 let FfnKind::Dense(d) = &lw.ffn else {
6376 return crate::gpu::BatchGraphOutcome::Declined;
6377 };
6378 if !d.segs.is_empty() {
6379 return crate::gpu::BatchGraphOutcome::Declined; }
6381 let (AttnKind::Full { softplus_gate: None, .. }, Some(gact)) =
6384 (&lw.attn, d.act.graph_act())
6385 else {
6386 return crate::gpu::BatchGraphOutcome::Declined;
6387 };
6388 fn gw(t: &QTensor) -> Option<crate::gpu::GraphW<'_>> {
6389 let (_, i, kind, rs) = t.graph_weight()?;
6390 Some(crate::gpu::GraphW {
6391 idx: i,
6392 kind,
6393 row_scale: rs,
6394 data: &[],
6395 prism: crate::gpu::GraphPrismOp::None,
6396 affine: false,
6397 })
6398 }
6399 let Some((model, _, _, _)) = wq.graph_weight() else {
6400 return crate::gpu::BatchGraphOutcome::Declined;
6401 };
6402 let model = model.clone();
6403 let (Some(gwq), Some(gwk), Some(gwv), Some(gwo), Some(gg), Some(gu), Some(gd)) = (
6404 gw(wq),
6405 gw(wk),
6406 gw(wv),
6407 gw(wo),
6408 gw(&d.gate_proj),
6409 gw(&d.up_proj),
6410 gw(&d.down_proj),
6411 ) else {
6412 return crate::gpu::BatchGraphOutcome::Declined;
6413 };
6414 let layer = crate::gpu::GraphLayer {
6415 input_norm: &lw.input_norm,
6416 attn: crate::gpu::GraphAttn::Full {
6417 wq: gwq,
6418 wk: gwk,
6419 wv: gwv,
6420 wo: gwo,
6421 q_norm: q_norm.as_deref(),
6422 k_norm: k_norm.as_deref(),
6423 late_qk_norm: self.qk_norm_after_rope,
6424 bias: bias
6425 .as_ref()
6426 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
6427 output_gate: *output_gate,
6428 cpu_k: m.kv.k_heads(),
6429 cpu_v: m.kv.v_heads(),
6430 cpu_base: m.kv.base(),
6431 geom: None,
6432 head_gate: None,
6433 },
6434 post_norm: &lw.post_norm,
6435 ffn: crate::gpu::GraphFfn::Dense {
6436 gate: gg,
6437 up: gu,
6438 down: gd,
6439 act: gact,
6440 },
6441 };
6442 let positions: Vec<usize> = (first_pos..first_pos + pairs.len()).collect();
6443 let nh = self.num_heads;
6444 let (nkv, hd, rd) = self.layer_geom(0);
6445 let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
6446 crate::gpu::forward_batch_graph(
6447 &model,
6448 self.mtp_kv_id(),
6449 std::slice::from_ref(&layer),
6450 &self.inv_freq,
6451 &mut hiddens,
6452 nh,
6453 nkv,
6454 hd,
6455 rd,
6456 hs,
6457 self.intermediate_size,
6458 &positions,
6459 self.kv_cache.max_seq_len,
6460 gemma,
6461 self.rms_eps as f32,
6462 self.attn_scale,
6463 pairs.len(),
6464 &[],
6465 0,
6466 None,
6467 None,
6468 )
6469 }
6470
6471 #[cfg(feature = "gpu")]
6478 fn mtp_warm_graph_fallback(
6479 &mut self,
6480 m: &mut MtpModule,
6481 pairs: &[(&[f32], u32)],
6482 first_pos: usize,
6483 ) -> bool {
6484 if pairs.is_empty() {
6485 return true;
6486 }
6487 let graphable = self.mtp_block_graph_ok(m);
6488 if !graphable {
6489 if self.mtp_graph_mode == Some(true) {
6493 return false;
6494 }
6495 self.mtp_graph_mode = Some(false);
6496 for (j, (h, t)) in pairs.iter().enumerate() {
6497 self.mtp_warm(m, h, *t, first_pos + j);
6498 }
6499 return true;
6500 }
6501
6502 for (j, (h, t)) in pairs.iter().enumerate() {
6507 if self.mtp_step_graph(m, h, *t, first_pos + j).is_none() {
6508 return false;
6509 }
6510 }
6511 self.mtp_graph_mode = Some(true);
6512 true
6513 }
6514
6515 #[cfg(feature = "gpu")]
6520 fn mtp_warm_prefill_pairs(
6521 &mut self,
6522 m: &mut MtpModule,
6523 pairs: &[(&[f32], u32)],
6524 first_pos: usize,
6525 ) -> Result<(), &'static str> {
6526 if self.mtp_graph_mode == Some(false) || !self.mtp_graph_ok(m) {
6531 if self.mtp_graph_mode == Some(true) {
6532 return Err("MTP token graph became unavailable after admission");
6533 }
6534 self.mtp_graph_mode = Some(false);
6535 for (j, (h, t)) in pairs.iter().enumerate() {
6536 self.mtp_warm(m, h, *t, first_pos + j);
6537 }
6538 return Ok(());
6539 }
6540 match self.mtp_warm_graph(m, pairs, first_pos) {
6541 crate::gpu::BatchGraphOutcome::Completed => {
6542 if !pairs.is_empty() {
6543 self.mtp_graph_mode = Some(true);
6544 }
6545 Ok(())
6546 }
6547 crate::gpu::BatchGraphOutcome::Declined => {
6548 if self.mtp_warm_graph_fallback(m, pairs, first_pos) {
6549 Ok(())
6550 } else {
6551 Err("MTP warm-up fallback failed after device admission")
6552 }
6553 }
6554 crate::gpu::BatchGraphOutcome::Failed => {
6555 Err("MTP warm batch graph failed after admission")
6556 }
6557 }
6558 }
6559
6560 #[cfg(not(feature = "gpu"))]
6561 fn mtp_warm_prefill_pairs(
6562 &mut self,
6563 m: &mut MtpModule,
6564 pairs: &[(&[f32], u32)],
6565 first_pos: usize,
6566 ) -> Result<(), &'static str> {
6567 for (j, (h, t)) in pairs.iter().enumerate() {
6568 self.mtp_warm(m, h, *t, first_pos + j);
6569 }
6570 Ok(())
6571 }
6572
6573 fn mtp_warm(&mut self, m: &mut MtpModule, hidden: &[f32], next_token: u32, position: usize) {
6577 let e = self.embed_single(next_token);
6578 let mut cat = vec![0.0f32; 2 * self.hidden_size];
6579 let (cat_e, cat_h) = cat.split_at_mut(self.hidden_size);
6580 inference::rms_norm_into(&e, &m.enorm, self.rms_eps, self.norm_style, cat_e);
6581 inference::rms_norm_into(hidden, &m.hnorm, self.rms_eps, self.norm_style, cat_h);
6582 let mut x = vec![0.0f32; self.hidden_size];
6583 m.eh_proj.matvec(&cat, &mut x, self.pool.as_deref());
6584 inference::rms_norm_into(
6585 &x,
6586 &m.layer.input_norm,
6587 self.rms_eps,
6588 self.norm_style,
6589 &mut self.ws.n1,
6590 );
6591 let attn = match &m.layer.attn {
6592 AttnKind::Full {
6593 wq,
6594 wk,
6595 wv,
6596 wo,
6597 q_norm,
6598 k_norm,
6599 output_gate,
6600 softplus_gate,
6601 bias,
6602 } => {
6603 let mut cfg = self.attn_cfg(position);
6604 cfg.q_norm = q_norm.as_deref();
6605 cfg.k_norm = k_norm.as_deref();
6606 cfg.output_gate = *output_gate;
6607 cfg.softplus_gate = softplus_gate.as_ref().map(|(g, p)| (g, *p));
6608 cfg.bias = bias
6609 .as_ref()
6610 .map(|(q, k, v)| (q.as_slice(), k.as_slice(), v.as_slice()));
6611 attention::qwen_attention(&self.ws.n1, wq, wk, wv, wo, &mut m.kv, &cfg)
6612 }
6613 _ => return,
6614 };
6615 let _ = attn;
6616 }
6617
6618 #[cfg(feature = "gpu")]
6625 #[allow(clippy::too_many_arguments)]
6626 fn graph_spec_step(
6627 &mut self,
6628 m: &mut MtpModule,
6629 hidden: &[f32],
6630 t_next: u32,
6631 next_pos: usize,
6632 drafted: &mut usize,
6633 accepted: &mut usize,
6634 all_ids: &mut Vec<u32>,
6638 room: usize,
6643 ) -> Option<(Vec<u32>, usize, Vec<f32>)> {
6644 #[cfg(target_os = "macos")]
6655 let metal_native = crate::gpu::q1_force();
6656 #[cfg(not(target_os = "macos"))]
6657 let metal_native = false;
6658 #[cfg(feature = "gpu")]
6659 let k_default = if metal_native {
6660 7
6663 } else if crate::gpu_wgpu::verify_i8_on() {
6664 5
6665 } else {
6666 4
6667 };
6668 #[cfg(not(feature = "gpu"))]
6669 let k_default = 4;
6670 let k_env: Option<usize> = std::env::var("CMF_GRAPH_SPEC_K")
6671 .ok()
6672 .and_then(|v| v.parse().ok())
6673 .filter(|&v| (1..=8).contains(&v));
6674 let (k_start, k_max) = if metal_native { (7, 7) } else { (3, k_default.max(5)) };
6679 let k_full: usize = k_env.unwrap_or_else(|| self.spec_k_adapt.unwrap_or(k_start));
6680 let k_spec = k_full.min(room).max(1);
6681 let k_capped = k_spec < k_full;
6684 if next_pos == 0 {
6685 return None;
6686 }
6687 let t_round = std::time::Instant::now();
6688 let subs = || crate::gpu_wgpu::SUBMITS.load(std::sync::atomic::Ordering::Relaxed);
6704 let sub0 = subs();
6705 let cfg = self.sampler_config.clone();
6710 let penalized = !(cfg.repetition_penalty == 1.0
6711 && cfg.presence_penalty == 0.0
6712 && cfg.suppress_tokens.is_empty());
6713 let greedy_pen = cfg.temperature < 1e-6 && penalized;
6718 let sampling = cfg.temperature >= 1e-6;
6719 let sparse = sampling && sampler::sparse_ok(&cfg);
6725 let base_len = all_ids.len();
6726 if sampling && !sparse && self.spec_q.len() < k_spec {
6727 self.spec_q.resize_with(k_spec, Vec::new);
6728 }
6729 if sparse && self.spec_qs.len() < k_spec {
6730 self.spec_qs.resize_with(k_spec, Vec::new);
6731 }
6732 let mut drafts = Vec::with_capacity(k_spec);
6737 let mut hx = hidden.to_vec();
6738 let spec_dbg = std::env::var("CMF_SPEC_DBG").is_ok();
6741 spec_stamp("pro");
6742 #[cfg(target_os = "macos")]
6748 if metal_native && !sampling && !greedy_pen && self.mtp_graph_mode != Some(false) {
6749 match self.mtp_draft_chain_metal(m, hidden, t_next, next_pos - 1, k_spec) {
6750 Ok(ids) => {
6751 self.mtp_graph_mode = Some(true);
6752 drafts = ids;
6753 }
6754 Err(true) => {
6755 tracing::error!("mtp Metal draft chain failed after commit");
6756 self.clear_sequence_state();
6757 self.graph_failed
6758 .store(true, std::sync::atomic::Ordering::Relaxed);
6759 self.cancel
6760 .store(true, std::sync::atomic::Ordering::Relaxed);
6761 return None;
6762 }
6763 Err(false) => {}
6764 }
6765 }
6766 for j in drafts.len()..k_spec {
6767 let tok_in = if j == 0 { t_next } else { drafts[j - 1] };
6768 let mut dbg_ref: Option<(Vec<f32>, Vec<f32>)> = None;
6769 if spec_dbg {
6770 let saved = self.mtp_graph_mode;
6771 self.mtp_graph_mode = Some(false);
6772 let r = self.mtp_step_hl(m, &hx, tok_in, next_pos - 1 + j);
6773 self.mtp_graph_mode = saved;
6774 if self.graph_failed.load(std::sync::atomic::Ordering::Relaxed) {
6775 return None;
6776 }
6777 m.kv.truncate_last(1);
6778 dbg_ref = Some(r);
6779 }
6780 let (mut lg, hj) = self.mtp_step_hl(m, &hx, tok_in, next_pos - 1 + j);
6781 if self.graph_failed.load(std::sync::atomic::Ordering::Relaxed) {
6782 return None;
6783 }
6784 if let Some((lg_cpu, h_cpu)) = dbg_ref {
6785 let n = |v: &[f32]| v.iter().map(|x| x * x).sum::<f32>().sqrt();
6786 let dl = lg
6787 .iter()
6788 .zip(&lg_cpu)
6789 .fold(0f32, |m, (a, b)| m.max((a - b).abs()));
6790 let dh = hj
6791 .iter()
6792 .zip(&h_cpu)
6793 .fold(0f32, |m, (a, b)| m.max((a - b).abs()));
6794 eprintln!(
6795 "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 {}",
6796 next_pos - 1 + j,
6797 sampler::argmax(&lg_cpu),
6798 sampler::argmax(&lg),
6799 n(&h_cpu),
6800 n(&hj),
6801 m.kv.seq_len
6802 );
6803 }
6804 let dj = if sparse {
6805 let mut q = std::mem::take(&mut self.spec_qs[j]);
6806 let ok = sampler::sparse_distribution_into(
6807 &lg,
6808 &cfg,
6809 all_ids,
6810 &mut self.sampler_scratch,
6811 self.pool.as_deref(),
6812 &mut q,
6813 );
6814 let d = if ok {
6815 sampler::draw_sparse(&q, &mut self.rng)
6816 } else {
6817 let t = sampler::argmax(&lg);
6819 q.clear();
6820 q.push((t, 1.0));
6821 t
6822 };
6823 self.spec_qs[j] = q;
6824 all_ids.push(d);
6825 d
6826 } else if sampling {
6827 let mut q = std::mem::take(&mut self.spec_q[j]);
6828 sampler::distribution_into(
6829 &lg,
6830 &cfg,
6831 all_ids,
6832 &mut self.sampler_scratch,
6833 self.pool.as_deref(),
6834 &mut q,
6835 );
6836 let d = sampler::draw(&q, &mut self.rng);
6837 self.spec_q[j] = q;
6838 all_ids.push(d); d
6840 } else if greedy_pen {
6841 let d = sampler::argmax_penalized(
6842 &lg,
6843 &cfg,
6844 all_ids,
6845 &mut self.sampler_scratch,
6846 self.pool.as_deref(),
6847 );
6848 all_ids.push(d);
6849 d
6850 } else {
6851 sampler::argmax(&lg)
6852 };
6853 attention::recycle_buf(&mut lg);
6854 drafts.push(dj);
6855 hx = hj;
6856 spec_stamp("d.pick");
6857 }
6858 all_ids.truncate(base_len);
6859 *drafted += k_spec;
6860 let t_draft = t_round.elapsed();
6861 let sub_draft = subs();
6862 let b = k_spec + 1;
6865 let mut hiddens = vec![0.0f32; b * self.hidden_size];
6866 for (i, &t) in std::iter::once(&t_next).chain(drafts.iter()).enumerate() {
6867 let e = self.embed_single(t);
6868 hiddens[i * self.hidden_size..(i + 1) * self.hidden_size].copy_from_slice(&e);
6869 }
6870 let positions: Vec<usize> = (next_pos..next_pos + b).collect();
6871 spec_stamp("v.emb");
6872 let (lm_gw, lm_rows) = {
6873 let (_, i, kind, rs) = self.weights.lm_head.graph_weight()?;
6874 (
6875 crate::gpu::GraphW {
6876 idx: i,
6877 kind,
6878 row_scale: rs,
6879 data: &[],
6880 prism: crate::gpu::GraphPrismOp::None,
6881 affine: false,
6882 },
6883 self.weights.lm_head.rows(),
6884 )
6885 };
6886 let mut logits = Vec::new();
6887 let final_norm = self.weights.final_norm.clone();
6888 #[cfg(target_os = "macos")]
6897 let greedy_dev = metal_native
6898 && !sampling
6899 && !greedy_pen
6900 && !self.confidence_on
6901 && self.final_softcap.is_none()
6902 && self.vocab_size == lm_rows
6908 && std::env::var_os("CMF_METAL_VERIFY_CHECK").is_none()
6909 && std::env::var_os("CMF_LOGIT_DUMP").is_none()
6910 && std::env::var("CMF_METAL_DEV_ARGMAX").as_deref() != Ok("0");
6911 #[cfg(not(target_os = "macos"))]
6912 let greedy_dev = false;
6913 let mut dev_ids: Vec<u32> = Vec::new();
6914 #[cfg(target_os = "macos")]
6915 let verify_outcome = if metal_native {
6916 let lm = self.weights.lm_head.q1_parts()?;
6917 let n_score = self.vocab_size.min(lm_rows);
6918 self.try_batch_graph_metal(
6919 &mut hiddens,
6920 &positions,
6921 b,
6922 Some((lm, &final_norm, &mut logits)),
6923 if greedy_dev {
6924 Some((n_score, &mut dev_ids))
6925 } else {
6926 None
6927 },
6928 )
6929 } else {
6930 self.try_batch_graph_wgpu(
6931 &mut hiddens,
6932 &positions,
6933 b,
6934 Some(crate::gpu::SpecTail {
6935 lm: lm_gw,
6936 lm_rows,
6937 final_norm: &final_norm,
6938 logits_out: &mut logits,
6939 }),
6940 )
6941 };
6942 #[cfg(not(target_os = "macos"))]
6943 let verify_outcome = self.try_batch_graph_wgpu(
6944 &mut hiddens,
6945 &positions,
6946 b,
6947 Some(crate::gpu::SpecTail {
6948 lm: lm_gw,
6949 lm_rows,
6950 final_norm: &final_norm,
6951 logits_out: &mut logits,
6952 }),
6953 );
6954 match verify_outcome {
6955 crate::gpu::BatchGraphOutcome::Completed => {}
6956 crate::gpu::BatchGraphOutcome::Declined => {
6957 m.kv.truncate_last(k_spec);
6961 if !metal_native && !self.rewind_mtp_graph_mirror(next_pos) {
6962 self.clear_sequence_state();
6963 self.graph_failed
6964 .store(true, std::sync::atomic::Ordering::Relaxed);
6965 self.cancel
6966 .store(true, std::sync::atomic::Ordering::Relaxed);
6967 tracing::error!("MTP graph mirror rewind failed after verify decline");
6968 }
6969 return None;
6970 }
6971 crate::gpu::BatchGraphOutcome::Failed => {
6972 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!("MTP verify batch graph failed after admission");
6981 return None;
6982 }
6983 }
6984 #[cfg(target_os = "macos")]
6990 if metal_native && std::env::var("CMF_METAL_VERIFY_CHECK").as_deref() == Ok("1") {
6991 let snap: Vec<Vec<f32>> = self
6992 .kv_cache
6993 .layers
6994 .iter()
6995 .map(|l| l.linear_state.clone())
6996 .collect();
6997 let attn_lens: Vec<usize> = self.kv_cache.layers.iter().map(|l| l.seq_len).collect();
6998 let toks: Vec<u32> = std::iter::once(t_next)
6999 .chain(drafts.iter().copied())
7000 .collect();
7001 let want_save = self.graph_want_logits;
7002 self.graph_want_logits = false;
7003 for (i, &t) in toks.iter().enumerate() {
7004 let hi = self.forward_layers(&self.embed_single(t), next_pos + i, None);
7005 let _ = self.graph_logits.take();
7006 if std::env::var("CMF_SPEC_PLAIN_HIDDEN").as_deref() == Ok("1") {
7010 hiddens[i * self.hidden_size..(i + 1) * self.hidden_size].copy_from_slice(&hi);
7011 }
7012 let ref_lg = self.logits_from_hidden(&hi);
7013 let row = &logits[i * lm_rows..(i + 1) * lm_rows];
7014 let ra = sampler::argmax(&ref_lg);
7015 let va = sampler::argmax(row);
7016 let mut md = 0f32;
7017 let mut rms = 0f64;
7018 for j in 0..lm_rows.min(ref_lg.len()) {
7019 let d = (ref_lg[j] - row[j]).abs();
7020 md = md.max(d);
7021 rms += (d as f64) * (d as f64);
7022 }
7023 let mut hd = 0f32;
7024 for j in 0..self.hidden_size {
7025 hd = hd.max((hi[j] - hiddens[i * self.hidden_size + j]).abs());
7026 }
7027 eprintln!(
7028 "verify-check row {i} tok {t} pos {}: ref argmax {ra} verify argmax {va} {} | max|dlogit| {md:.3} rms {:.4} | max|dhidden| {hd:.4}",
7029 next_pos + i,
7030 if ra == va { "OK" } else { "MISMATCH" },
7031 (rms / lm_rows as f64).sqrt()
7032 );
7033 }
7034 self.graph_want_logits = want_save;
7035 for (l, st) in self.kv_cache.layers.iter_mut().zip(snap) {
7038 if l.linear_state.len() == st.len() {
7039 l.linear_state.copy_from_slice(&st);
7040 } else {
7041 l.linear_state = st;
7042 }
7043 }
7044 for (li, (l, n0)) in self.kv_cache.layers.iter_mut().zip(attn_lens).enumerate() {
7045 let extra = l.seq_len.saturating_sub(n0);
7046 if extra > 0 {
7047 l.truncate_last(extra);
7048 crate::gpu_metal::kv_mirror_set_stored(self.graph_kv_id, li, n0);
7049 }
7050 }
7051 }
7052 let t_verify = t_round.elapsed();
7053 let sub_verify = subs();
7054 let mut a = 0usize;
7059 let mut forced: Option<u32> = None;
7060 let ids: Vec<u32> = if sparse {
7061 let mut p = std::mem::take(&mut self.spec_ps);
7062 let mut res = std::mem::take(&mut self.spec_ress);
7063 while a < k_spec {
7064 let ok = sampler::sparse_distribution_into(
7065 &logits[a * lm_rows..(a + 1) * lm_rows],
7066 &cfg,
7067 all_ids,
7068 &mut self.sampler_scratch,
7069 self.pool.as_deref(),
7070 &mut p,
7071 );
7072 if !ok {
7073 let t = sampler::argmax(&logits[a * lm_rows..(a + 1) * lm_rows]);
7074 p.clear();
7075 p.push((t, 1.0));
7076 }
7077 match sampler::spec_accept_or_correct_sparse(
7078 &p,
7079 &self.spec_qs[a],
7080 drafts[a],
7081 &mut self.rng,
7082 &mut res,
7083 ) {
7084 None => {
7085 all_ids.push(drafts[a]);
7086 a += 1;
7087 }
7088 Some(c) => {
7089 forced = Some(c);
7090 break;
7091 }
7092 }
7093 }
7094 all_ids.truncate(base_len);
7095 self.spec_ps = p;
7096 self.spec_ress = res;
7097 drafts.clone()
7098 } else if sampling {
7099 let mut p = std::mem::take(&mut self.spec_p);
7100 let mut res = std::mem::take(&mut self.spec_res);
7101 while a < k_spec {
7102 sampler::distribution_into(
7103 &logits[a * lm_rows..(a + 1) * lm_rows],
7104 &cfg,
7105 all_ids,
7106 &mut self.sampler_scratch,
7107 self.pool.as_deref(),
7108 &mut p,
7109 );
7110 match sampler::spec_accept_or_correct(
7111 &p,
7112 &self.spec_q[a],
7113 drafts[a],
7114 &mut self.rng,
7115 &mut res,
7116 self.pool.as_deref(),
7117 ) {
7118 None => {
7119 all_ids.push(drafts[a]);
7120 a += 1;
7121 }
7122 Some(c) => {
7123 forced = Some(c);
7124 break;
7125 }
7126 }
7127 }
7128 all_ids.truncate(base_len);
7129 self.spec_p = p;
7130 self.spec_res = res;
7131 drafts.clone()
7133 } else if greedy_pen {
7134 let mut ids: Vec<u32> = Vec::with_capacity(b);
7138 for i in 0..b {
7139 let t = sampler::argmax_penalized(
7140 &logits[i * lm_rows..(i + 1) * lm_rows],
7141 &cfg,
7142 all_ids,
7143 &mut self.sampler_scratch,
7144 self.pool.as_deref(),
7145 );
7146 ids.push(t);
7147 if i < k_spec && t == drafts[i] {
7148 all_ids.push(t);
7149 } else {
7150 break;
7151 }
7152 }
7153 all_ids.truncate(base_len);
7154 while a < k_spec && a < ids.len() && ids[a] == drafts[a] {
7155 a += 1;
7156 }
7157 ids
7160 } else if greedy_dev && dev_ids.len() == b {
7161 let ids = std::mem::take(&mut dev_ids);
7162 while a < k_spec && ids[a] == drafts[a] {
7163 a += 1;
7164 }
7165 ids
7166 } else {
7167 if logits.len() < b * lm_rows {
7168 self.clear_sequence_state();
7171 self.graph_failed
7172 .store(true, std::sync::atomic::Ordering::Relaxed);
7173 self.cancel
7174 .store(true, std::sync::atomic::Ordering::Relaxed);
7175 tracing::error!("Metal verify returned neither logits nor argmax ids");
7176 return None;
7177 }
7178 let ids: Vec<u32> = (0..b)
7179 .map(|i| sampler::argmax(&logits[i * lm_rows..(i + 1) * lm_rows]))
7180 .collect();
7181 while a < k_spec && ids[a] == drafts[a] {
7182 a += 1;
7183 }
7184 ids
7185 };
7186 spec_stamp("acc");
7187 if spec_dbg {
7188 eprintln!(
7189 "spec-dbg round: t_next {t_next} drafts {:?} verified {:?} accepted {a}",
7190 drafts, ids
7191 );
7192 }
7193 #[cfg(target_os = "macos")]
7197 let commit_ref: Option<(Vec<Vec<f32>>, Vec<(usize, Vec<f32>, Vec<f32>)>)> = if metal_native
7198 && std::env::var("CMF_METAL_VERIFY_CHECK").as_deref() == Ok("2")
7199 {
7200 let snap: Vec<Vec<f32>> = self
7201 .kv_cache
7202 .layers
7203 .iter()
7204 .map(|l| l.linear_state.clone())
7205 .collect();
7206 let attn_lens: Vec<usize> = self.kv_cache.layers.iter().map(|l| l.seq_len).collect();
7207 let toks: Vec<u32> = std::iter::once(t_next)
7208 .chain(drafts.iter().copied())
7209 .collect();
7210 let want_save = self.graph_want_logits;
7211 self.graph_want_logits = false;
7212 for (i, &t) in toks.iter().take(a + 1).enumerate() {
7213 let _ = self.forward_layers(&self.embed_single(t), next_pos + i, None);
7214 let _ = self.graph_logits.take();
7215 }
7216 self.graph_want_logits = want_save;
7217 let plain_states: Vec<Vec<f32>> = self
7218 .kv_cache
7219 .layers
7220 .iter()
7221 .map(|l| l.linear_state.clone())
7222 .collect();
7223 let (nkv, hd) = (self.num_kv_heads, self.head_dim);
7224 let mut rows = Vec::new();
7225 for (li, (l, n0)) in self
7226 .kv_cache
7227 .layers
7228 .iter_mut()
7229 .zip(attn_lens.iter())
7230 .enumerate()
7231 {
7232 let extra = l.seq_len.saturating_sub(*n0);
7233 if extra > 0 {
7234 let mut kk = Vec::new();
7235 let mut vv = Vec::new();
7236 for g in 0..nkv {
7237 kk.extend_from_slice(&l.head_keys(g)[n0 * hd..]);
7238 vv.extend_from_slice(&l.head_values(g)[n0 * hd..]);
7239 }
7240 rows.push((li, kk, vv));
7241 l.truncate_last(extra);
7242 crate::gpu_metal::kv_mirror_set_stored(self.graph_kv_id, li, *n0);
7243 }
7244 }
7245 for (l, st) in self.kv_cache.layers.iter_mut().zip(snap) {
7246 if l.linear_state.len() == st.len() {
7247 l.linear_state.copy_from_slice(&st);
7248 } else {
7249 l.linear_state = st;
7250 }
7251 }
7252 Some((plain_states, rows))
7253 } else {
7254 None
7255 };
7256 let warm_off = std::env::var("CMF_SPEC_WARM").is_ok_and(|v| v == "0");
7257 #[cfg(target_os = "macos")]
7266 let mut warm_pending: Option<MetalWarmPending> = None;
7267 #[cfg(target_os = "macos")]
7268 if metal_native {
7269 m.kv.truncate_last(k_spec.saturating_sub(1));
7270 if self.mtp_graph_mode == Some(true) {
7271 crate::gpu_metal::kv_mirror_set_stored(
7274 self.mtp_kv_id(),
7275 Self::MTP_LAYER_BASE,
7276 m.kv.seq_len,
7277 );
7278 if !warm_off && a > 0 {
7279 let pairs: Vec<(&[f32], u32)> = (0..a)
7280 .map(|j| {
7281 (
7282 &hiddens[j * self.hidden_size..(j + 1) * self.hidden_size],
7283 ids[j],
7284 )
7285 })
7286 .collect();
7287 warm_pending = self.mtp_warm_batch_submit(m, &pairs, next_pos);
7288 }
7289 }
7290 spec_stamp("c.wsub");
7291 }
7292 #[cfg(target_os = "macos")]
7294 if metal_native {
7295 if !self.metal_verify_commit(a) {
7298 self.clear_sequence_state();
7299 self.graph_failed
7300 .store(true, std::sync::atomic::Ordering::Relaxed);
7301 self.cancel
7302 .store(true, std::sync::atomic::Ordering::Relaxed);
7303 tracing::error!("Metal verify state/KV handoff failed after admission");
7304 return None;
7305 }
7306 if let Some((plain_states, rows)) = commit_ref {
7307 crate::gpu_metal::queue_fence();
7308 let _ = crate::gpu_metal::wait_replay();
7311 let (nkv, hd) = (self.num_kv_heads, self.head_dim);
7312 let mut worst_s = 0f32;
7313 let mut worst_li = 0usize;
7314 for (li, (l, ps)) in self.kv_cache.layers.iter().zip(&plain_states).enumerate() {
7315 if l.linear_state.len() != ps.len() || ps.is_empty() {
7316 continue;
7317 }
7318 let d = l
7319 .linear_state
7320 .iter()
7321 .zip(ps)
7322 .fold(0f32, |m, (x, y)| m.max((x - y).abs()));
7323 let n = ps.iter().fold(0f32, |m, y| m.max(y.abs()));
7324 let rel = d / n.max(1e-6);
7325 if rel > worst_s {
7326 worst_s = rel;
7327 worst_li = li;
7328 }
7329 }
7330 let mut worst_k = 0f32;
7331 for (li, kk, vv) in &rows {
7332 let l = &self.kv_cache.layers[*li];
7333 let n0 = l.seq_len - (kk.len() / (nkv * hd));
7334 let mut ck = Vec::new();
7335 let mut cv = Vec::new();
7336 for g in 0..nkv {
7337 ck.extend_from_slice(&l.head_keys(g)[n0 * hd..]);
7338 cv.extend_from_slice(&l.head_values(g)[n0 * hd..]);
7339 }
7340 if ck.len() == kk.len() {
7341 let dk = ck
7342 .iter()
7343 .zip(kk)
7344 .fold(0f32, |m, (x, y)| m.max((x - y).abs()));
7345 let dv = cv
7346 .iter()
7347 .zip(vv)
7348 .fold(0f32, |m, (x, y)| m.max((x - y).abs()));
7349 worst_k = worst_k.max(dk).max(dv);
7350 } else {
7351 eprintln!(
7352 "commit-check L{li}: kv row count mismatch {} vs {}",
7353 ck.len(),
7354 kk.len()
7355 );
7356 }
7357 }
7358 eprintln!(
7359 "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}"
7360 );
7361 }
7362 }
7363 if !metal_native && a + 1 < b {
7364 let expected_gdn_layers = self.graph_gdn_layer_count();
7365 if expected_gdn_layers > 0
7366 && !crate::gpu::gdn_spec_restore(self.graph_kv_id, a, next_pos, expected_gdn_layers)
7367 {
7368 self.clear_sequence_state();
7369 self.graph_failed
7370 .store(true, std::sync::atomic::Ordering::Relaxed);
7371 self.cancel
7372 .store(true, std::sync::atomic::Ordering::Relaxed);
7373 tracing::error!("GDN speculative restore failed after verify");
7374 return None;
7375 }
7376 }
7377 if !metal_native && !self.rewind_trunk_graph_mirrors(next_pos + a + 1) {
7378 self.clear_sequence_state();
7383 self.graph_failed
7384 .store(true, std::sync::atomic::Ordering::Relaxed);
7385 self.cancel
7386 .store(true, std::sync::atomic::Ordering::Relaxed);
7387 tracing::error!("trunk graph KV rewind failed after speculative verify");
7388 return None;
7389 }
7390 *accepted += a;
7391 if !metal_native {
7402 m.kv.truncate_last(k_spec.saturating_sub(1));
7404 }
7405 spec_stamp("c.trunc");
7406 if !metal_native
7407 && self.mtp_graph_mode == Some(true)
7408 && !self.rewind_mtp_graph_mirror(next_pos)
7409 {
7410 self.clear_sequence_state();
7414 self.graph_failed
7415 .store(true, std::sync::atomic::Ordering::Relaxed);
7416 self.cancel
7417 .store(true, std::sync::atomic::Ordering::Relaxed);
7418 tracing::error!("MTP graph mirror rewind failed after verify commit");
7419 return None;
7420 }
7421 if !warm_off && a > 0 {
7422 let mut warmed = false;
7425 #[cfg(target_os = "macos")]
7426 if metal_native && self.mtp_graph_mode == Some(true) {
7427 warmed = match warm_pending.take() {
7431 Some(p) => self.mtp_warm_batch_finish(m, p),
7432 None => false,
7433 };
7434 if !warmed {
7435 warmed = true;
7436 for j in 0..a {
7437 let row =
7438 hiddens[j * self.hidden_size..(j + 1) * self.hidden_size].to_vec();
7439 if self
7440 .mtp_step_metal(m, &row, ids[j], next_pos + j, false)
7441 .is_none()
7442 {
7443 warmed = false;
7444 break;
7445 }
7446 }
7447 }
7448 }
7449 if !warmed && self.mtp_graph_mode != Some(false) && !metal_native {
7450 let rows: Vec<Vec<f32>> = (0..a)
7451 .map(|j| hiddens[j * self.hidden_size..(j + 1) * self.hidden_size].to_vec())
7452 .collect();
7453 let pairs: Vec<(&[f32], u32)> = rows
7454 .iter()
7455 .zip(ids.iter())
7456 .map(|(r, &t)| (r.as_slice(), t))
7457 .collect();
7458 match self.mtp_warm_prefill_pairs(m, &pairs, next_pos) {
7459 Ok(()) => warmed = true,
7460 Err(err) => {
7461 tracing::error!("{err}");
7467 self.clear_sequence_state();
7468 self.graph_failed
7469 .store(true, std::sync::atomic::Ordering::Relaxed);
7470 self.cancel
7471 .store(true, std::sync::atomic::Ordering::Relaxed);
7472 return None;
7473 }
7474 }
7475 }
7476 if !warmed {
7477 for j in 0..a {
7478 let row = &hiddens[j * self.hidden_size..(j + 1) * self.hidden_size];
7479 let row = row.to_vec();
7480 self.mtp_warm(m, &row, ids[j], next_pos + j);
7481 }
7482 }
7483 }
7484 spec_stamp("c.warm");
7488 if let Some(c) = forced {
7489 self.spec_forced = Some(c);
7490 self.graph_logits = None;
7491 } else if greedy_dev && logits.is_empty() {
7492 self.spec_forced = Some(ids[a]);
7495 self.graph_logits = None;
7496 } else {
7497 let mut row = logits[a * lm_rows..(a + 1) * lm_rows].to_vec();
7498 row.resize(self.vocab_size, 0.0);
7499 if let Some(c) = self.final_softcap {
7500 for l in row.iter_mut() {
7501 *l = c * (*l / c).tanh();
7502 }
7503 }
7504 self.graph_logits = Some(row);
7505 }
7506 let new_hidden = hiddens[a * self.hidden_size..(a + 1) * self.hidden_size].to_vec();
7507 spec_stamp("c.row");
7508 if std::env::var("CMF_GRAPH_SPEC_TIME").is_ok() {
7514 let end = subs();
7515 eprintln!(
7516 "spec-round: draft {:.1} ms/{} sub | verify {:.1} ms/{} sub | \
7517 commit {:.1} ms/{} sub (accepted {a} of {k_spec}, full-head streak {})",
7518 t_draft.as_secs_f64() * 1e3,
7519 sub_draft - sub0,
7520 (t_verify - t_draft).as_secs_f64() * 1e3,
7521 sub_verify - sub_draft,
7522 (t_round.elapsed() - t_verify).as_secs_f64() * 1e3,
7523 end - sub_verify,
7524 self.draft_full_streak,
7525 );
7526 }
7527 if k_env.is_none() && !metal_native && !k_capped {
7532 let f = a as f32 / k_spec.max(1) as f32;
7536 self.spec_acc_ewma += 0.2 * (f - self.spec_acc_ewma);
7537 let mut k_next = k_spec;
7538 if self.spec_acc_ewma >= 0.75 && k_spec < k_max {
7539 k_next = k_spec + 1;
7540 } else if self.spec_acc_ewma < 0.4 && k_spec > 2 {
7541 k_next = k_spec - 1;
7542 }
7543 if k_next != k_spec {
7544 self.spec_acc_ewma = 0.6;
7545 if std::env::var("CMF_GRAPH_SPEC_TIME").is_ok() {
7546 eprintln!("spec-k: {k_spec} → {k_next}");
7547 }
7548 }
7549 self.spec_k_adapt = Some(k_next);
7550 }
7551 spec_stamp("end");
7552 Some((drafts[..a].to_vec(), next_pos + a + 1, new_hidden))
7553 }
7554
7555 pub fn measure_pair_fusion(&mut self, iters: usize) -> (f64, f64) {
7564 if !self.pair_supported() {
7565 return (0.0, 0.0);
7566 }
7567 let graph_env = std::env::var_os("CMF_GPU_WGPU_GRAPH");
7574 unsafe { std::env::set_var("CMF_GPU_WGPU_GRAPH", "0") };
7575 let emb1 = self.embed_single(1);
7576 let emb2 = self.embed_single(2);
7577 let pos = self.kv_cache.seq_len();
7578
7579 let t0 = std::time::Instant::now();
7580 for _ in 0..iters {
7581 let _ = self.forward_layers(&emb1, pos, None);
7582 let _ = self.forward_layers(&emb2, pos + 1, None);
7583 for l in &mut self.kv_cache.layers {
7584 l.truncate_last(2);
7585 }
7586 }
7587 let singles_ms = t0.elapsed().as_secs_f64() * 1000.0 / iters as f64;
7588
7589 let t1 = std::time::Instant::now();
7590 for _ in 0..iters {
7591 let _ = self.forward_pair(&emb1, &emb2, pos);
7592 for l in &mut self.kv_cache.layers {
7593 l.truncate_last(2);
7594 }
7595 }
7596 let pair_ms = t1.elapsed().as_secs_f64() * 1000.0 / iters as f64;
7597 match graph_env {
7598 Some(value) => unsafe { std::env::set_var("CMF_GPU_WGPU_GRAPH", value) },
7599 None => unsafe { std::env::remove_var("CMF_GPU_WGPU_GRAPH") },
7600 }
7601 (singles_ms, pair_ms)
7602 }
7603
7604 fn pair_supported(&self) -> bool {
7612 !self.weights.layers.is_empty()
7619 && self.g3n.is_none()
7620 && !self
7621 .weights
7622 .layers
7623 .iter()
7624 .any(|lw| matches!(&lw.attn, AttnKind::Mla(_) | AttnKind::Kda(_)))
7625 }
7626
7627 fn forward_pair(
7628 &mut self,
7629 emb1: &[f32],
7630 emb2: &[f32],
7631 position: usize,
7632 ) -> (Vec<f32>, Vec<f32>) {
7633 self.mimo_moe_prepare();
7636 let mut h1 = emb1.to_vec();
7637 let mut h2 = emb2.to_vec();
7638 let (_nkv, _hd, hs, _rd, eps) = (
7639 self.num_kv_heads,
7640 self.head_dim,
7641 self.hidden_size,
7642 self.rotary_dim,
7643 self.rms_eps,
7644 );
7645 let pool = self.pool.clone();
7646
7647 for li in 0..self.num_layers {
7648 let lw = &self.weights.layers[self.phys_layer(li)];
7649 inference::rms_norm_into(
7652 &h1,
7653 &lw.input_norm,
7654 self.rms_eps,
7655 self.norm_style,
7656 &mut self.ws.n1,
7657 );
7658 inference::rms_norm_into(
7659 &h2,
7660 &lw.input_norm,
7661 self.rms_eps,
7662 self.norm_style,
7663 &mut self.ws.n2,
7664 );
7665
7666 let (a1, a2) = match &lw.attn {
7667 AttnKind::Mla(_) => unreachable!("MLA has no MTP/pair path"),
7668 AttnKind::Kda(_) => unreachable!("KDA has no MTP/pair path"),
7669 AttnKind::Bounded(w) => {
7670 let rope = self
7673 .bounded_rope
7674 .clone()
7675 .expect("bounded layer without an installed rotation table");
7676 let cfg = crate::bounded::BoundedAttnCfg {
7677 num_heads: self.num_heads,
7678 num_kv_heads: self.num_kv_heads,
7679 head_dim: self.head_dim,
7680 hidden_size: hs,
7681 scale: self.attn_scale,
7682 rope: &rope,
7683 pool: pool.as_deref(),
7684 };
7685 let a1 = crate::bounded::bounded_attention(
7686 &self.ws.n1,
7687 w,
7688 &mut self.kv_cache.layers[li],
7689 &cfg,
7690 );
7691 let a2 = crate::bounded::bounded_attention(
7692 &self.ws.n2,
7693 w,
7694 &mut self.kv_cache.layers[li],
7695 &cfg,
7696 );
7697 (a1, a2)
7698 }
7699 AttnKind::Linear(w) => {
7700 let cfg = self.vmf_cfg.expect("linear layer without vmf_cfg");
7701 let layer = &mut self.kv_cache.layers[li];
7702 let (state, scratch) = (&mut layer.linear_state, &mut layer.linear_scratch);
7703 vmf_phase_pair(
7704 &self.ws.n1,
7705 &self.ws.n2,
7706 w,
7707 &cfg,
7708 state,
7709 scratch,
7710 self.pool.as_deref(),
7711 )
7712 }
7713 AttnKind::LinearGdn(w) => {
7714 let cfg = self.gdn_cfg.expect("gdn layer without gdn_cfg");
7715 let layer = &mut self.kv_cache.layers[li];
7716 let (state, scratch) = (&mut layer.linear_state, &mut layer.linear_scratch);
7717 gdn_pair(
7718 &self.ws.n1,
7719 &self.ws.n2,
7720 w,
7721 &cfg,
7722 state,
7723 scratch,
7724 self.pool.as_deref(),
7725 )
7726 }
7727 AttnKind::ShortConv(w) => {
7728 let cfg = self
7729 .short_conv_cfg
7730 .expect("short-conv layer without short_conv_cfg");
7731 let layer = &mut self.kv_cache.layers[li];
7732 let (state, scratch) = (&mut layer.linear_state, &mut layer.linear_scratch);
7733 short_conv_pair(
7734 &self.ws.n1,
7735 &self.ws.n2,
7736 w,
7737 &cfg,
7738 state,
7739 scratch,
7740 self.pool.as_deref(),
7741 )
7742 }
7743 AttnKind::Full {
7744 wq,
7745 wk,
7746 wv,
7747 wo,
7748 q_norm,
7749 k_norm,
7750 output_gate,
7751 softplus_gate,
7752 bias,
7753 } => {
7754 let inv_freq_l = self.layer_inv_freq(li);
7755 let (nkv_l, hd_l, rd_l) = self.layer_geom(li);
7756 let cfg = QwenAttnCfg {
7757 num_heads: self.layer_num_heads(li),
7758 num_kv_heads: nkv_l,
7759 head_dim: hd_l,
7760 hidden_size: hs,
7761 position,
7762 inv_freq: &inv_freq_l,
7763 rotary_dim: rd_l,
7764 scale: self.attn_scale,
7765 softcap: self.attn_softcap,
7766 window: self.layer_window(li),
7767 v_norm: self.attn_v_norm,
7768 qk_norm_after_rope: self.qk_norm_after_rope,
7769 gate_sigmoid: self.proj_gate_sigmoid,
7770 q_norm: q_norm.as_deref(),
7771 k_norm: k_norm.as_deref(),
7772 output_gate: *output_gate,
7773 softplus_gate: softplus_gate
7774 .as_ref()
7775 .map(|(gate, per_head)| (gate, *per_head)),
7776 rope_scale: self.layer_rope_scale(li),
7777 bias: bias
7778 .as_ref()
7779 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
7780 rms_eps: eps,
7781 norm_style: self.norm_style,
7782 pool: pool.as_deref(),
7783 v_head_dim: self.layer_v_dim(li),
7784 };
7785 attention::qwen_attention_pair(
7786 &self.ws.n1,
7787 &self.ws.n2,
7788 wq,
7789 wk,
7790 wv,
7791 wo,
7792 &mut self.kv_cache.layers[li],
7793 &cfg,
7794 )
7795 }
7796 };
7797 let (a1, a2) = match &self.weights.layers[self.phys_layer(li)].attn_out_norm {
7798 Some(w) => (
7799 inference::rms_norm(&a1, w, self.rms_eps, self.norm_style),
7800 inference::rms_norm(&a2, w, self.rms_eps, self.norm_style),
7801 ),
7802 None => (a1, a2),
7803 };
7804 for i in 0..self.hidden_size {
7805 h1[i] += a1[i];
7806 h2[i] += a2[i];
7807 }
7808 let (mut a1, mut a2) = (a1, a2);
7809 attention::recycle_buf(&mut a1);
7810 attention::recycle_buf(&mut a2);
7811
7812 let lw = &self.weights.layers[self.phys_layer(li)];
7813 inference::rms_norm_into(
7814 &h1,
7815 &lw.post_norm,
7816 self.rms_eps,
7817 self.norm_style,
7818 &mut self.ws.p1,
7819 );
7820 inference::rms_norm_into(
7821 &h2,
7822 &lw.post_norm,
7823 self.rms_eps,
7824 self.norm_style,
7825 &mut self.ws.p2,
7826 );
7827 let (f1, f2) = match &lw.ffn {
7828 FfnKind::DenseMoe(dm) => (
7831 dense_moe_ffn(
7832 dm,
7833 &self.ws.p1,
7834 &h1,
7835 self.rms_eps,
7836 self.norm_style,
7837 self.pool.as_deref(),
7838 ),
7839 dense_moe_ffn(
7840 dm,
7841 &self.ws.p2,
7842 &h2,
7843 self.rms_eps,
7844 self.norm_style,
7845 self.pool.as_deref(),
7846 ),
7847 ),
7848 FfnKind::Moe(m) if self.mimo_moe.is_dynamic(li, false) => (
7849 moe_ffn_banked(&mut self.mimo_moe, li, m, &self.ws.p1, self.pool.as_deref()),
7850 moe_ffn_banked(&mut self.mimo_moe, li, m, &self.ws.p2, self.pool.as_deref()),
7851 ),
7852 _ => ffn_forward_pair(
7853 &lw.ffn,
7854 &self.ws.p1,
7855 &self.ws.p2,
7856 self.pool.as_deref(),
7857 None,
7858 ),
7859 };
7860 let (f1, f2) = match &self.weights.layers[self.phys_layer(li)].ffn_out_norm {
7861 Some(w) => (
7862 inference::rms_norm(&f1, w, self.rms_eps, self.norm_style),
7863 inference::rms_norm(&f2, w, self.rms_eps, self.norm_style),
7864 ),
7865 None => (f1, f2),
7866 };
7867 for i in 0..self.hidden_size {
7868 h1[i] += f1[i];
7869 h2[i] += f2[i];
7870 }
7871 let (mut f1, mut f2) = (f1, f2);
7872 attention::recycle_buf(&mut f1);
7873 attention::recycle_buf(&mut f2);
7874 if let Some(sc) = self.weights.layers[self.phys_layer(li)].layer_scale {
7875 for i in 0..self.hidden_size {
7876 h1[i] *= sc;
7877 h2[i] *= sc;
7878 }
7879 }
7880 if self.is_loop_end(li) && li + 1 < self.num_layers {
7882 h1 = inference::rms_norm(
7883 &h1,
7884 &self.weights.final_norm,
7885 self.rms_eps,
7886 self.norm_style,
7887 );
7888 h2 = inference::rms_norm(
7889 &h2,
7890 &self.weights.final_norm,
7891 self.rms_eps,
7892 self.norm_style,
7893 );
7894 }
7895 }
7896 if self.o1_active() {
7902 self.commit_linear_scratch();
7903 }
7904 self.o1_progress();
7905 self.swa_trim_tails();
7906 (h1, h2)
7907 }
7908
7909 fn commit_linear_scratch(&mut self) {
7911 for layer in &mut self.kv_cache.layers {
7912 if !layer.linear_scratch.is_empty() {
7913 std::mem::swap(&mut layer.linear_state, &mut layer.linear_scratch);
7914 layer.linear_scratch.clear();
7915 }
7916 }
7917 }
7918
7919 pub fn forward_ids(
7922 &mut self,
7923 ids: &[u32],
7924 task_mask: Option<&TaskMask>,
7925 ) -> Result<Vec<f32>, String> {
7926 #[cfg(target_os = "macos")]
7927 crate::gpu_metal::set_io_namespace(self.graph_kv_id);
7928 if ids.is_empty() {
7929 return Err("empty id sequence".to_string());
7930 }
7931 self.clear_sequence_state();
7932 self.check_forward_graph("forward_ids setup", 0)?;
7933 if task_mask.is_none() {
7934 self.o1_begin();
7935 }
7936 let mut hidden = vec![0.0f32; self.hidden_size];
7937 let mut pos = 0usize;
7938 if let Some(b) = &mut self.dsv41 {
7939 let pool = self.pool.clone();
7940 let mut logits = Vec::new();
7941 crate::dsv41::forward_chunk(
7942 &b.0,
7943 &b.1,
7944 &b.2,
7945 &mut b.3,
7946 ids,
7947 0,
7948 pool.as_deref(),
7949 &mut logits,
7950 );
7951 if let Err(err) = self.o1_seal_checked() {
7952 self.clear_sequence_state();
7953 return Err(err);
7954 }
7955 return Ok(logits);
7956 }
7957 if self.can_prefill_batched() && !self.graph_prefill_preferred() && ids.len() > 2 {
7965 let chunk = self.prefill_chunk();
7969 let hs = self.hidden_size;
7970 while pos < ids.len() {
7971 let end = (pos + chunk).min(ids.len());
7972 let hb = self.prefill_rows(&ids[pos..end], pos, task_mask)?;
7973 self.check_forward_graph("forward_ids batched prefill", end - 1)?;
7974 hidden.copy_from_slice(&hb[(end - pos - 1) * hs..]);
7975 pos = end;
7976 }
7977 }
7978 if task_mask.is_none()
7987 && !self.graph_prefill_preferred()
7988 && !std::env::var("CMF_PAIR").is_ok_and(|v| v == "0")
7989 && self.pair_supported()
7990 {
7991 while pos + 1 < ids.len() {
7992 let e1 = self.embed_single(ids[pos]);
7993 let e2 = self.embed_single(ids[pos + 1]);
7994 let (_, h2) = self.forward_pair(&e1, &e2, pos);
7995 self.check_forward_graph("forward_ids pair", pos + 1)?;
7996 self.commit_linear_scratch();
7997 hidden = h2;
7998 pos += 2;
7999 }
8000 }
8001 if task_mask.is_none() && pos == 0 && ids.len() > 1 {
8004 if let Some(lg) = self.embryo_prefill_chunked(ids, 0) {
8005 self.graph_logits = Some(lg);
8006 hidden = vec![0.0; self.hidden_size];
8007 pos = ids.len();
8008 }
8009 }
8010 while pos < ids.len() {
8011 hidden = self.forward_layers(&self.embed_single(ids[pos]), pos, task_mask);
8012 self.check_forward_graph("forward_ids", pos)?;
8013 pos += 1;
8014 }
8015 if let Some(logits) = self.graph_logits.take() {
8016 if let Err(err) = self.o1_seal_checked() {
8020 self.clear_sequence_state();
8021 return Err(err);
8022 }
8023 return Ok(logits);
8024 }
8025 if let Err(err) = self.o1_seal_checked() {
8029 self.clear_sequence_state();
8030 return Err(err);
8031 }
8032 let normed = inference::rms_norm(
8033 &hidden,
8034 &self.weights.final_norm,
8035 self.rms_eps,
8036 self.norm_style,
8037 );
8038 Ok(self.lm_head_forward(&normed))
8039 }
8040
8041 #[doc(hidden)]
8045 pub fn dsv41_serial_logits(&mut self, ids: &[u32]) -> Result<Vec<Vec<f32>>, String> {
8046 #[cfg(target_os = "macos")]
8047 crate::gpu_metal::set_io_namespace(self.graph_kv_id);
8048 if ids.is_empty() {
8049 return Err("empty id sequence".to_string());
8050 }
8051 self.clear_sequence_state();
8052 self.dsv41
8053 .as_ref()
8054 .ok_or_else(|| "dsv41 serial logits require a DeepSeek-V4.1 model".to_string())?;
8055 self.o1_begin();
8056 let rows = {
8057 let pool = self.pool.clone();
8058 let b = self
8059 .dsv41
8060 .as_mut()
8061 .expect("dsv41 checked above; state cannot change during forward");
8062 let mut rows = Vec::with_capacity(ids.len());
8063 for (position, &id) in ids.iter().enumerate() {
8064 let mut logits = Vec::new();
8065 crate::dsv41::forward_token(
8066 &b.0,
8067 &b.1,
8068 &b.2,
8069 &mut b.3,
8070 id,
8071 position,
8072 pool.as_deref(),
8073 &mut logits,
8074 );
8075 rows.push(logits);
8076 }
8077 rows
8078 };
8079 self.o1_seal();
8080 Ok(rows)
8081 }
8082
8083 pub fn ppl_ids(&mut self, ids: &[u32]) -> Result<f64, String> {
8090 let (nll, cnt) = self.nll_ids_from(ids, 0)?;
8091 Ok((nll / cnt.max(1) as f64).exp())
8092 }
8093
8094 pub fn probe_ffn_mass(&mut self, ids: &[u32]) -> Vec<Vec<f64>> {
8099 self.clear_sequence_state();
8100 FFN_PROBE.with(|p| {
8101 *p.borrow_mut() = Some(vec![vec![0f64; self.intermediate_size]; self.num_layers]);
8102 });
8103 crate::gpu::cpu_scope(|| {
8104 for (pos, &id) in ids.iter().enumerate() {
8105 let emb = self.embed_single(id);
8106 let _ = self.forward_layers(&emb, pos, None);
8107 }
8108 });
8109 self.clear_sequence_state();
8110 FFN_PROBE
8111 .with(|p| p.borrow_mut().take())
8112 .unwrap_or_default()
8113 }
8114
8115 pub fn probe_ffn_mass_batch(&mut self, ids: &[u32]) -> Result<Vec<Vec<f64>>, String> {
8119 if let Err(err) = self.nll_begin() {
8120 let _ = FFN_PROBE.with(|p| p.borrow_mut().take());
8124 self.nll_end();
8125 return Err(err);
8126 }
8127 FFN_PROBE.with(|p| {
8128 *p.borrow_mut() = Some(vec![vec![0f64; self.intermediate_size]; self.num_layers]);
8129 });
8130 let result: Result<(), String> = (|| {
8131 for chunk in ids.chunks(256) {
8132 if chunk.len() < 2 {
8133 continue;
8134 }
8135 self.nll_ids_masked(chunk, 0, None)?;
8136 }
8137 Ok(())
8138 })();
8139 self.nll_end();
8140 let probe = FFN_PROBE
8141 .with(|p| p.borrow_mut().take())
8142 .unwrap_or_default();
8143 match result {
8144 Ok(()) => Ok(probe),
8145 Err(err) => {
8146 drop(probe);
8147 Err(err)
8148 }
8149 }
8150 }
8151
8152 pub fn ppl_ids_masked(&mut self, ids: &[u32], mask: &TaskMask) -> Result<f64, String> {
8156 self.nll_begin()?;
8157 let result: Result<f64, String> = (|| {
8158 let mut nll = 0f64;
8159 let mut cnt = 0usize;
8160 let mut hidden = vec![0f32; self.hidden_size];
8161 for (pos, &id) in ids.iter().enumerate() {
8162 if pos > 0 {
8163 inference::rms_norm_into(
8164 &hidden,
8165 &self.weights.final_norm,
8166 self.rms_eps,
8167 self.norm_style,
8168 &mut self.ws.n1,
8169 );
8170 let mut logits = self.lm_head_forward(&self.ws.n1);
8171 let max = logits.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
8172 let sum: f64 = logits.iter().map(|&v| ((v - max) as f64).exp()).sum();
8173 let p = ((logits[id as usize] - max) as f64).exp() / sum.max(1e-300);
8174 nll -= p.max(1e-300).ln();
8175 cnt += 1;
8176 attention::recycle_buf(&mut logits);
8177 }
8178 let emb = self.embed_single(id);
8179 hidden = self.forward_layers(&emb, pos, Some(mask));
8180 self.nll_check_graph("masked serial forward", pos)?;
8181 let _ = self.graph_logits.take();
8185 }
8186 Ok((nll / cnt.max(1) as f64).exp())
8187 })();
8188 self.nll_end();
8189 result
8190 }
8191
8192 pub fn nll_ids_masked(
8211 &mut self,
8212 ids: &[u32],
8213 start: usize,
8214 task_mask: Option<&TaskMask>,
8215 ) -> Result<(f64, usize), String> {
8216 let task_mask = self.drop_open_mask(task_mask);
8217 self.nll_ids_inner(ids, start, task_mask)
8218 }
8219
8220 pub fn nll_ids_from(&mut self, ids: &[u32], start: usize) -> Result<(f64, usize), String> {
8221 self.nll_ids_inner(ids, start, None)
8222 }
8223
8224 fn nll_ids_inner(
8225 &mut self,
8226 ids: &[u32],
8227 start: usize,
8228 task_mask: Option<&TaskMask>,
8229 ) -> Result<(f64, usize), String> {
8230 self.nll_begin()?;
8231 let result: Result<(f64, usize), String> = (|| {
8232 let mut nll = 0f64;
8233 let mut cnt = 0usize;
8234 let (graph_quality, fused_head_quality) = nll_graph_policy(
8247 task_mask.is_none(),
8248 self.graph_prefill_preferred(),
8249 crate::gpu::q1_force(),
8250 );
8251 self.graph_head_required = fused_head_quality;
8252 self.graph_want_logits = fused_head_quality;
8253 #[cfg(target_os = "macos")]
8254 if graph_quality && std::env::var("CMF_METAL_BATCH_NLL").as_deref() != Ok("0") {
8255 match self.nll_batch_metal(ids, start) {
8256 MetalBatchNllOutcome::Completed(nll, count) => {
8257 return Ok((nll, count));
8258 }
8259 MetalBatchNllOutcome::Declined => {}
8260 MetalBatchNllOutcome::Failed(err) => return Err(err),
8261 }
8262 }
8263 let force_serial = std::env::var("CMF_NLL_SERIAL").as_deref() == Ok("1");
8268 if self.can_prefill_batched() && !graph_quality && !force_serial {
8269 const CHUNK: usize = 128;
8275 const LM_SUB: usize = 32;
8276 let n = ids.len().saturating_sub(1);
8277 let hs = self.hidden_size;
8278 let rows = self.weights.lm_head.rows();
8279 let mut pos = 0usize;
8280 let state_trace = std::env::var("CMF_STATE_TRACE").is_ok();
8281 while pos < n {
8282 let end = (pos + CHUNK).min(n);
8283 let bsz = end - pos;
8284 let hb = self.prefill_rows(&ids[pos..end], pos, task_mask)?;
8285 self.nll_check_graph("batched prefill", pos)?;
8286 if state_trace && end % 256 == 0 {
8287 self.trace_recurrent_state(end);
8288 }
8289 let mut k0 = 0usize;
8290 while k0 < bsz {
8291 let k1 = (k0 + LM_SUB).min(bsz);
8292 let sb = k1 - k0;
8293 if pos + k1 <= start {
8296 k0 = k1;
8297 continue;
8298 }
8299 let mut normed = vec![0.0f32; sb * hs];
8300 for k in 0..sb {
8301 let r = inference::rms_norm(
8302 &hb[(k0 + k) * hs..(k0 + k + 1) * hs],
8303 &self.weights.final_norm,
8304 self.rms_eps,
8305 self.norm_style,
8306 );
8307 normed[k * hs..(k + 1) * hs].copy_from_slice(&r);
8308 }
8309 let mut logits = vec![0.0f32; sb * rows];
8310 self.weights
8311 .lm_head
8312 .matmat(&normed, sb, &mut logits, self.pool.as_deref());
8313 for k in 0..sb {
8314 if pos + k0 + k < start {
8315 continue;
8316 }
8317 self.nll_check_graph("batched score row", pos + k0 + k)?;
8318 let lg = &mut logits[k * rows..k * rows + self.vocab_size.min(rows)];
8319 if let Some(mu) = self.logit_multiplier {
8320 for v in lg.iter_mut() {
8321 *v *= mu;
8322 }
8323 }
8324 if let Some(c) = self.final_softcap {
8328 for v in lg.iter_mut() {
8329 *v = c * (*v / c).tanh();
8330 }
8331 }
8332 if let Some(cm) = self.head_clusters.clone() {
8335 self.hierarchical_head_logprobs(
8336 &normed[k * hs..(k + 1) * hs],
8337 &cm,
8338 lg,
8339 );
8340 }
8341 let lg = &logits[k * rows..k * rows + self.vocab_size.min(rows)];
8342 let target = ids[pos + k0 + k + 1] as usize;
8343 let max = lg.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
8344 let lse: f64 = lg
8345 .iter()
8346 .map(|&v| ((v - max) as f64).exp())
8347 .sum::<f64>()
8348 .ln()
8349 + max as f64;
8350 nll += lse - lg[target] as f64;
8351 cnt += 1;
8352 if std::env::var("CMF_PPL_TRACE").is_ok() {
8353 let top = lg
8354 .iter()
8355 .enumerate()
8356 .max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
8357 .map(|(i, _)| i)
8358 .unwrap_or(0);
8359 eprintln!(
8360 "BTRACE pos {} target {} nll {:.4} top {} lg_t {:.3} lg_top {:.3}",
8361 pos + k0 + k,
8362 target,
8363 lse - lg[target] as f64,
8364 top,
8365 lg[target],
8366 lg[top]
8367 );
8368 }
8369 }
8370 k0 = k1;
8371 }
8372 pos = end;
8373 }
8374 return Ok((nll, cnt));
8375 }
8376 for pos in 0..ids.len().saturating_sub(1) {
8377 let hidden = self.forward_layers(&self.embed_single(ids[pos]), pos, task_mask);
8378 self.nll_check_graph("serial forward", pos)?;
8379 let out_of_band = self.graph_logits.take();
8387 if self.graph_head_required && out_of_band.is_none() {
8388 METAL_GRAPH_HEAD_MISS.fetch_add(
8389 1,
8390 std::sync::atomic::Ordering::Relaxed,
8391 );
8392 return Err(format!(
8393 "fused Metal graph head did not complete at NLL position {pos}"
8394 ));
8395 }
8396 if pos < start {
8397 continue;
8398 }
8399 let logits = match out_of_band {
8400 Some(lg) => lg,
8401 None => {
8402 let normed = inference::rms_norm(
8403 &hidden,
8404 &self.weights.final_norm,
8405 self.rms_eps,
8406 self.norm_style,
8407 );
8408 self.lm_head_forward(&normed)
8412 }
8413 };
8414 let target = ids[pos + 1] as usize;
8415 let max = logits.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
8416 let lse: f64 = logits
8417 .iter()
8418 .map(|&v| ((v - max) as f64).exp())
8419 .sum::<f64>()
8420 .ln()
8421 + max as f64;
8422 let tok_nll = lse - logits[target] as f64;
8423 if std::env::var("CMF_PPL_TRACE").is_ok() && pos < 48 {
8424 let top = logits
8425 .iter()
8426 .enumerate()
8427 .max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
8428 .map(|(i, _)| i)
8429 .unwrap_or(0);
8430 eprintln!(
8431 "pos {pos:3} tgt {target:6} nll {tok_nll:7.3} | top1 {top:6} lg[t]={:.2} lg[top]={:.2}",
8432 logits[target], logits[top]
8433 );
8434 }
8435 nll += tok_nll;
8436 cnt += 1;
8437 }
8438 Ok((nll, cnt))
8439 })();
8440 self.nll_end();
8441 result
8442 }
8443
8444 fn nll_from_hidden(&mut self, hidden: &[f32], target: u32, pos: usize) -> f64 {
8449 let normed = inference::rms_norm(
8450 hidden,
8451 &self.weights.final_norm,
8452 self.rms_eps,
8453 self.norm_style,
8454 );
8455 let mut logits = self.lm_head_forward(&normed);
8458 let target = target as usize;
8459 let max = logits.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
8460 let lse: f64 = logits
8461 .iter()
8462 .map(|&v| ((v - max) as f64).exp())
8463 .sum::<f64>()
8464 .ln()
8465 + max as f64;
8466 let tok_nll = lse - logits[target] as f64;
8467 if std::env::var("CMF_PPL_TRACE").is_ok() && pos < 48 {
8468 let top = logits
8469 .iter()
8470 .enumerate()
8471 .max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
8472 .map(|(i, _)| i)
8473 .unwrap_or(0);
8474 eprintln!(
8475 "pos {pos:3} tgt {target:6} nll {tok_nll:7.3} | top1 {top:6} lg[t]={:.2} lg[top]={:.2}",
8476 logits[target], logits[top]
8477 );
8478 }
8479 attention::recycle_buf(&mut logits);
8480 tok_nll
8481 }
8482
8483 fn trace_recurrent_state(&self, pos: usize) {
8491 let stats = |v: &[f32]| -> (f64, f64) {
8492 if v.is_empty() {
8493 return (0.0, 0.0);
8494 }
8495 let (mut ss, mut mx) = (0f64, 0f64);
8496 for &x in v {
8497 ss += (x as f64) * (x as f64);
8498 mx = mx.max((x as f64).abs());
8499 }
8500 ((ss / v.len() as f64).sqrt(), mx)
8501 };
8502 for (li, l) in self.kv_cache.layers.iter().enumerate() {
8503 let lw = &self.weights.layers[self.phys_layer(li)];
8504 let (kind, s_len) = match &lw.attn {
8505 AttnKind::Linear(_) => (
8506 "vmf",
8507 self.vmf_cfg.map(|c| c.state_len()).unwrap_or(0),
8508 ),
8509 AttnKind::LinearGdn(_) => (
8510 "gdn",
8511 self.gdn_cfg.map(|c| c.state_len()).unwrap_or(0),
8512 ),
8513 AttnKind::Bounded(_) => ("bounded", 0),
8514 AttnKind::Full { .. } => ("full", 0),
8515 _ => ("other", 0),
8516 };
8517 let (rms, max) = stats(&l.linear_state);
8518 let s_part = if kind == "vmf" {
8519 &l.linear_state[..s_len.min(l.linear_state.len())]
8520 } else if kind == "gdn" {
8521 let ring = l.linear_state.len().saturating_sub(s_len).min(l.linear_state.len());
8522 &l.linear_state[ring..]
8523 } else {
8524 &l.linear_state[..0]
8525 };
8526 let (s_rms, s_max) = stats(s_part);
8527 let (ring_rms, ring_len) = match &l.bounded {
8528 Some(b) => (stats(&b.ring_k).0, b.ring_k.len()),
8529 None => (0.0, 0),
8530 };
8531 eprintln!(
8532 "STATE pos={pos} layer={li} kind={kind} state_len={} rms={rms:.5} max={max:.4} \
8533 S_rms={s_rms:.5} S_max={s_max:.4} ring_k_rms={ring_rms:.5} ring_len={ring_len} kv_rows={}",
8534 l.linear_state.len(),
8535 l.seq_len
8536 );
8537 }
8538 }
8539
8540 pub fn nll_ids_o1(&mut self, ids: &[u32], prefill: usize) -> Result<(f64, usize), String> {
8558 self.nll_begin()?;
8563 let requested_prefix = (prefill > 0).then_some(prefill);
8564 self.o1_begin_with_prefix(requested_prefix);
8565 let n = ids.len().saturating_sub(1);
8566 let requested_start = prefill.min(n);
8567 let exact_end = if self.o1_active() {
8572 match requested_prefix {
8573 Some(requested) => self.o1_effective_boundary(requested),
8574 None => self
8575 .o1_cfg
8576 .as_ref()
8577 .and_then(|c| crate::nystrom::o1_deferred_boundary(c.w, c.sink)),
8578 }
8579 .unwrap_or(requested_start)
8580 .min(n)
8581 } else {
8582 requested_start
8583 };
8584 let mut nll = 0f64;
8585 let mut cnt = 0usize;
8586
8587 let mut pos = 0usize;
8591 if self.can_prefill_batched() {
8592 const CHUNK: usize = 128;
8593 while pos < exact_end {
8594 let end = (pos + CHUNK).min(exact_end);
8595 let hiddens = self.prefill_batch(&ids[pos..end], pos);
8596 if self
8597 .graph_failed
8598 .swap(false, std::sync::atomic::Ordering::Relaxed)
8599 {
8600 self.cancel
8601 .store(false, std::sync::atomic::Ordering::Relaxed);
8602 self.nll_end();
8603 return Err("GPU graph failed during O(1) NLL prefix".into());
8604 }
8605 for row in 0..end - pos {
8606 let score_pos = pos + row;
8607 if score_pos >= requested_start && score_pos < n {
8608 nll += self.nll_from_hidden(
8609 &hiddens[row * self.hidden_size..(row + 1) * self.hidden_size],
8610 ids[score_pos + 1],
8611 score_pos,
8612 );
8613 cnt += 1;
8614 }
8615 }
8616 pos = end;
8617 }
8618 } else {
8619 while pos < exact_end {
8620 let hidden = self.forward_layers(&self.embed_single(ids[pos]), pos, None);
8621 if self
8622 .graph_failed
8623 .swap(false, std::sync::atomic::Ordering::Relaxed)
8624 {
8625 self.cancel
8626 .store(false, std::sync::atomic::Ordering::Relaxed);
8627 self.nll_end();
8628 return Err("GPU graph failed during O(1) NLL prefix".into());
8629 }
8630 if pos >= requested_start && pos < n {
8631 nll += self.nll_from_hidden(&hidden, ids[pos + 1], pos);
8632 cnt += 1;
8633 }
8634 pos += 1;
8635 }
8636 }
8637 self.o1_seal_checked().map_err(|err| {
8638 self.nll_end();
8639 err
8640 })?;
8641
8642 let batch_k = std::env::var("CMF_BATCH_K")
8651 .ok()
8652 .and_then(|v| v.parse::<usize>().ok())
8653 .unwrap_or(0);
8654 let batch_admitted = batch_k > 0
8655 && self.can_prefill_batched()
8656 && self.o1_active()
8657 && std::env::var("CMF_O1_GPU").as_deref() == Ok("1")
8658 && (0..self.num_layers).all(|li| {
8659 let cache = &self.kv_cache.layers[self.phys_layer(li)];
8660 cache.o1.is_none() || cache.o1_views().is_some()
8661 });
8662 if std::env::var("CMF_GRAPH_PROF").is_ok() {
8663 eprintln!(
8664 "nll-batch: phase=post-seal admission={} requested_k={} scored_rows={}",
8665 batch_admitted,
8666 batch_k,
8667 n.saturating_sub(exact_end),
8668 );
8669 }
8670 let mut batch_completed = false;
8671 if batch_admitted && exact_end < n {
8672 let hs = self.hidden_size;
8673 let mut batch_pos = exact_end;
8674 while batch_pos < n {
8675 let end = (batch_pos + batch_k).min(n);
8676 let bk = end - batch_pos;
8677 let mut hiddens = vec![0.0f32; bk * hs];
8678 for (row, &id) in ids[batch_pos..end].iter().enumerate() {
8679 hiddens[row * hs..(row + 1) * hs].copy_from_slice(&self.embed_single(id));
8680 }
8681 let positions: Vec<usize> = (batch_pos..end).collect();
8682 let t_batch = std::time::Instant::now();
8683 let outcome = self.try_batch_graph_wgpu(&mut hiddens, &positions, bk, None);
8684 if std::env::var("CMF_GRAPH_PROF").is_ok() {
8685 let ms = t_batch.elapsed().as_secs_f64() * 1000.0;
8686 eprintln!(
8687 "nll-batch: phase=post-seal mode=o1 k={bk} pos={}..{} outcome={outcome:?} {ms:.1} ms ({:.1} tok/s)",
8688 batch_pos,
8689 end.saturating_sub(1),
8690 bk as f64 / (ms / 1000.0),
8691 );
8692 }
8693 if let Err(err) = self.nll_check_graph("batch graph", batch_pos) {
8694 self.nll_end();
8695 return Err(err);
8696 }
8697 match outcome {
8698 crate::gpu::BatchGraphOutcome::Completed => {
8699 batch_completed = true;
8700 for row in 0..bk {
8701 nll += self.nll_from_hidden(
8702 &hiddens[row * hs..(row + 1) * hs],
8703 ids[batch_pos + row + 1],
8704 batch_pos + row,
8705 );
8706 cnt += 1;
8707 }
8708 batch_pos = end;
8709 }
8710 crate::gpu::BatchGraphOutcome::Declined => {
8711 if batch_completed {
8712 self.nll_end();
8713 return Err(format!(
8714 "O(1) NLL batch declined after completed chunk at position {batch_pos}"
8715 ));
8716 }
8717 break;
8718 }
8719 crate::gpu::BatchGraphOutcome::Failed => {
8720 self.nll_end();
8721 return Err(format!(
8722 "O(1) NLL batch graph failed after admission at position {batch_pos}"
8723 ));
8724 }
8725 }
8726 }
8727 if batch_completed && cnt == n.saturating_sub(requested_start) {
8728 self.nll_end();
8729 return Ok((nll, cnt));
8730 }
8731 }
8732
8733 for pos in exact_end..n {
8738 let hidden = self.forward_layers(&self.embed_single(ids[pos]), pos, None);
8739 if self
8740 .graph_failed
8741 .swap(false, std::sync::atomic::Ordering::Relaxed)
8742 {
8743 self.cancel
8744 .store(false, std::sync::atomic::Ordering::Relaxed);
8745 self.nll_end();
8746 return Err(format!(
8747 "GPU graph failed during O(1) NLL serial scoring at position {pos}"
8748 ));
8749 }
8750 nll += self.nll_from_hidden(&hidden, ids[pos + 1], pos);
8751 cnt += 1;
8752 }
8753 self.nll_end();
8754 Ok((nll, cnt))
8755 }
8756
8757 pub fn calib_ids(&mut self, ids: &[u32], temps: &[f32]) -> (Vec<bool>, Vec<Vec<f32>>) {
8765 self.clear_sequence_state();
8766 let n = ids.len().saturating_sub(1);
8767 let mut correct = Vec::with_capacity(n);
8768 let mut pmax = Vec::with_capacity(n);
8769 for pos in 0..n {
8770 let emb = self.embed_single(ids[pos]);
8771 let hidden = self.forward_layers(&emb, pos, None);
8772 let logits = if let Some(logits) = self.graph_logits.take() {
8773 logits
8774 } else {
8775 let normed = inference::rms_norm(
8776 &hidden,
8777 &self.weights.final_norm,
8778 self.rms_eps,
8779 self.norm_style,
8780 );
8781 self.lm_head_forward(&normed)
8785 };
8786 let target = ids[pos + 1] as usize;
8787 let (mut amax, mut mval) = (0usize, f32::NEG_INFINITY);
8788 for (i, &v) in logits.iter().enumerate() {
8789 if v > mval {
8790 mval = v;
8791 amax = i;
8792 }
8793 }
8794 correct.push(amax == target);
8795 let row: Vec<f32> = temps
8796 .iter()
8797 .map(|&t| {
8798 let tt = t.max(1e-3);
8799 let s: f32 = logits.iter().map(|&v| ((v - mval) / tt).exp()).sum();
8800 1.0 / s.max(1e-12) })
8802 .collect();
8803 pmax.push(row);
8804 }
8805 self.clear_sequence_state();
8806 (correct, pmax)
8807 }
8808
8809 pub fn ppl_ids_dynamic(&mut self, ids: &[u32]) -> Result<(f64, usize), String> {
8816 if self.dyn_router.is_none() {
8817 return Ok((self.ppl_ids(ids)?, 0));
8818 }
8819 self.nll_begin()?;
8820 let saved_active = self.dyn_active;
8821 let mut router = self
8822 .dyn_router
8823 .take()
8824 .ok_or_else(|| "dynamic router disappeared before PPL scoring".to_string())?;
8825 router.reset();
8826 self.dyn_phi_seen = 0;
8827 let _ = self.set_active_skill(None);
8828
8829 let result: Result<(f64, usize), String> = (|| {
8830 let mut nll = 0f64;
8831 let mut cnt = 0usize;
8832 for pos in 0..ids.len().saturating_sub(1) {
8833 let hidden = self.forward_layers(&self.embed_single(ids[pos]), pos, None);
8834 self.nll_check_graph("dynamic serial forward", pos)?;
8835 let out_of_band = self.graph_logits.take();
8836 let mut logits = match out_of_band {
8837 Some(lg) => lg,
8838 None => {
8839 let normed = inference::rms_norm(
8840 &hidden,
8841 &self.weights.final_norm,
8842 self.rms_eps,
8843 self.norm_style,
8844 );
8845 self.lm_head_forward(&normed)
8849 }
8850 };
8851 let target = ids[pos + 1] as usize;
8852 let max = logits.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
8853 let lse: f64 = logits
8854 .iter()
8855 .map(|&v| ((v - max) as f64).exp())
8856 .sum::<f64>()
8857 .ln()
8858 + max as f64;
8859 let tok_nll = lse - logits[target] as f64;
8860 if std::env::var("CMF_PPL_TRACE").is_ok() && pos < 48 {
8861 let top = logits
8862 .iter()
8863 .enumerate()
8864 .max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
8865 .map(|(i, _)| i)
8866 .unwrap_or(0);
8867 eprintln!(
8868 "pos {pos:3} tgt {target:6} nll {tok_nll:7.3} | top1 {top:6} lg[t]={:.2} lg[top]={:.2}",
8869 logits[target], logits[top]
8870 );
8871 }
8872 nll += tok_nll;
8873 cnt += 1;
8874 attention::recycle_buf(&mut logits);
8875 let phi = self.dyn_phi_ema.clone();
8877 if let Some(new_active) = router.step(&phi, pos) {
8878 let _ = self.set_active_skill(new_active);
8879 }
8880 }
8881 Ok(((nll / cnt.max(1) as f64).exp(), router.switches.len()))
8882 })();
8883
8884 let _ = self.set_active_skill(saved_active);
8887 self.dyn_router = Some(router);
8888 self.nll_end();
8889 result
8890 }
8891
8892 pub fn probe_phi(&mut self, ids: &[u32], layer: usize) -> Vec<f32> {
8894 self.clear_sequence_state();
8895 let mut acc = vec![0f32; self.hidden_size];
8896 for (pos, &id) in ids.iter().enumerate() {
8897 let h = self.forward_layers_upto(&self.embed_single(id), pos, None, Some(layer));
8898 for (a, v) in acc.iter_mut().zip(&h) {
8899 *a += v;
8900 }
8901 }
8902 let n = ids.len().max(1) as f32;
8903 for a in acc.iter_mut() {
8904 *a /= n;
8905 }
8906 self.clear_sequence_state();
8907 acc
8908 }
8909
8910 pub fn probe_phi_span(
8924 &mut self,
8925 ids: &[u32],
8926 layer: usize,
8927 span: std::ops::Range<usize>,
8928 ) -> Vec<f32> {
8929 #[cfg(target_os = "macos")]
8930 crate::gpu_metal::set_io_namespace(self.graph_kv_id);
8931 let end = span.end.min(ids.len());
8932 let start = span.start.min(end);
8933 let reset = |p: &mut Self| p.clear_sequence_state();
8934 reset(self);
8935 let mut acc = vec![0f32; self.hidden_size];
8936 for (pos, &id) in ids[..end].iter().enumerate() {
8937 let h = self.forward_layers_upto(&self.embed_single(id), pos, None, Some(layer));
8938 if pos >= start {
8939 for (a, v) in acc.iter_mut().zip(&h) {
8940 *a += v;
8941 }
8942 }
8943 }
8944 let n = end - start;
8945 if n > 0 {
8946 let n = n as f32;
8947 for a in acc.iter_mut() {
8948 *a /= n;
8949 }
8950 }
8951 reset(self);
8952 acc
8953 }
8954
8955 pub fn decode_step_logits(&mut self, token: u32, position: usize) -> Vec<f32> {
8962 #[cfg(target_os = "macos")]
8963 crate::gpu_metal::set_io_namespace(self.graph_kv_id);
8964 self.graph_logits = None;
8965 let hidden = self.forward_layers(&self.embed_single(token), position, None);
8966 if let Some(logits) = self.graph_logits.take() {
8967 return logits;
8968 }
8969 inference::rms_norm_into(
8970 &hidden,
8971 &self.weights.final_norm,
8972 self.rms_eps,
8973 self.norm_style,
8974 &mut self.ws.n1,
8975 );
8976 self.lm_head_forward(&self.ws.n1)
8977 }
8978
8979 fn prefill_batch(&mut self, ids: &[u32], start_pos: usize) -> Vec<f32> {
8985 self.prefill_batch_masked(ids, start_pos, None)
8986 }
8987
8988 fn prefill_batch_masked(
8994 &mut self,
8995 ids: &[u32],
8996 start_pos: usize,
8997 task_mask: Option<&TaskMask>,
8998 ) -> Vec<f32> {
8999 self.prefill_batch_span(PrefillIn::Ids(ids), start_pos, task_mask, 0, usize::MAX)
9000 }
9001
9002 fn prefill_rows(
9010 &mut self,
9011 ids: &[u32],
9012 pos: usize,
9013 task_mask: Option<&TaskMask>,
9014 ) -> Result<Vec<f32>, String> {
9015 self.prefill_input_rows(PrefillIn::Ids(ids), pos, task_mask)
9016 }
9017
9018 fn prefill_input_rows(
9019 &mut self,
9020 input: PrefillIn<'_>,
9021 pos: usize,
9022 task_mask: Option<&TaskMask>,
9023 ) -> Result<Vec<f32>, String> {
9024 self.mimo_moe_prepare();
9025 let hs = self.hidden_size;
9026 let bk = match input {
9027 PrefillIn::Ids(ids) => ids.len(),
9028 PrefillIn::Hidden(rows) => rows.len() / hs,
9029 };
9030 #[cfg(not(target_os = "macos"))]
9031 if task_mask.is_none()
9032 && !self.o1_active()
9033 && bk > 1
9034 && (self.batch_prefix_prefill()
9035 || (self.verify_exact_moe
9036 && crate::gpu::enabled_here()
9037 && crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Decode)))
9038 {
9039 let mut hiddens = match input {
9040 PrefillIn::Hidden(rows) => rows.to_vec(),
9041 PrefillIn::Ids(ids) => ids.iter().flat_map(|&id| self.embed_single(id)).collect(),
9042 };
9043 let positions: Vec<usize> = (pos..pos + bk).collect();
9044 let mut run = 0usize;
9045 match self.try_batch_graph_wgpu_prefix(
9046 &mut hiddens,
9047 &positions,
9048 bk,
9049 None,
9050 Some(&mut run),
9051 ) {
9052 crate::gpu::BatchGraphOutcome::Completed => {
9053 let out = if run < self.num_layers {
9054 self.prefill_batch_span(
9055 PrefillIn::Hidden(&hiddens),
9056 pos,
9057 None,
9058 run,
9059 self.num_layers,
9060 )
9061 } else {
9062 hiddens
9063 };
9064 return if self.graph_failed.load(std::sync::atomic::Ordering::Relaxed) {
9065 Err("MiMo attention graph failed after admission".into())
9066 } else { Ok(out) };
9067 }
9068 crate::gpu::BatchGraphOutcome::Failed => {
9069 return Err("batched prefix prefill failed after admission".into());
9070 }
9071 crate::gpu::BatchGraphOutcome::Declined => {
9072 #[cfg(feature = "gpu")]
9074 self.pull_lagging_host_kv(0, self.num_layers, pos);
9075 }
9076 }
9077 }
9078 let out = self.prefill_batch_span(input, pos, task_mask, 0, usize::MAX);
9079 if self.graph_failed.load(std::sync::atomic::Ordering::Relaxed) {
9080 Err("batch tail graph failed after admission".into())
9081 } else { Ok(out) }
9082 }
9083
9084 fn prefill_batch_span(
9090 &mut self,
9091 input: PrefillIn<'_>,
9092 start_pos: usize,
9093 task_mask: Option<&TaskMask>,
9094 from: usize,
9095 upto_excl: usize,
9096 ) -> Vec<f32> {
9097 let hs = self.hidden_size;
9098 let b = match input {
9099 PrefillIn::Ids(ids) => ids.len(),
9100 PrefillIn::Hidden(hb) => hb.len() / hs,
9101 };
9102 let upto_excl = upto_excl.min(self.num_layers);
9103 let mut h: Vec<f32>;
9107 let mut h_ready;
9108 match input {
9109 PrefillIn::Ids(_) => {
9110 h = vec![0.0; b * hs];
9111 h_ready = false;
9112 }
9113 PrefillIn::Hidden(hb) => {
9114 h = hb.to_vec();
9115 h_ready = true;
9116 }
9117 }
9118 let fill_h = |h: &mut Vec<f32>, me: &Self| {
9119 if let PrefillIn::Ids(ids) = input {
9120 for (bi, &id) in ids.iter().enumerate() {
9121 let e = me.embed_single(id);
9122 h[bi * hs..(bi + 1) * hs].copy_from_slice(&e);
9123 }
9124 if let Ok(tp) = std::env::var("CMF_TRACE_POS") {
9125 if let Ok(t) = tp.parse::<usize>() {
9126 if t >= start_pos && t < start_pos + ids.len() {
9127 let bi = t - start_pos;
9128 let row = &h[bi * hs..(bi + 1) * hs];
9129 let n: f32 = row.iter().map(|x| x * x).sum::<f32>().sqrt();
9130 eprintln!(
9131 "BATCH pos {t} embed: id {} |h| = {n:.6} h0 {:.6} h1 {:.6} | b={} start={start_pos} ids[..8]={:?}",
9132 ids[bi],
9133 row[0],
9134 row[1],
9135 ids.len(),
9136 &ids[..ids.len().min(8)]
9137 );
9138 }
9139 }
9140 }
9141 }
9142 };
9143 let (_nkv, _hd, _rd, eps) = (
9144 self.num_kv_heads,
9145 self.head_dim,
9146 self.rotary_dim,
9147 self.rms_eps,
9148 );
9149 let pool = self.pool.clone();
9150 let norm_style = self.norm_style;
9151 self.mimo_moe_prepare();
9152 let automatic_gpu_prefix = self.automatic_gpu_prefix();
9153
9154 #[cfg(target_os = "macos")]
9155 let mut chunk_skip_until = 0usize;
9156 for li in from..upto_excl {
9157 let _capacity_tail = automatic_gpu_prefix
9158 .filter(|&prefix| {
9159 li >= prefix && !(self.verify_exact_moe && self.mimo_moe.is_dynamic(li, false))
9160 })
9161 .map(|_| crate::gpu::enter_cpu_scope());
9162 crate::gpu::set_layer(li as i64); #[cfg(feature = "gpu")]
9167 let _fast_gemm = self
9168 .proj_gate_sigmoid
9169 .then(crate::gpu::enter_prefill_fast_gemm);
9170 #[cfg(target_os = "macos")]
9176 if task_mask.is_none() {
9177 if li < chunk_skip_until {
9178 continue;
9179 }
9180 if !h_ready && li == 0 && self.weights.embed_tokens.q8_row_parts().is_none() {
9186 fill_h(&mut h, self);
9187 h_ready = true;
9188 }
9189 let ids_for_embed = match input {
9190 PrefillIn::Ids(ids) => (!h_ready && li == 0).then_some(ids),
9191 PrefillIn::Hidden(_) => None,
9192 };
9193 let end = self.chunk_run_gpu(li, &mut h, b, start_pos, ids_for_embed, upto_excl);
9194 if end > li {
9195 h_ready = true;
9196 chunk_skip_until = end;
9197 if self.is_loop_end(end - 1) && end < self.num_layers {
9200 for bi in 0..b {
9201 let normed = inference::rms_norm(
9202 &h[bi * hs..(bi + 1) * hs],
9203 &self.weights.final_norm,
9204 eps,
9205 norm_style,
9206 );
9207 h[bi * hs..(bi + 1) * hs].copy_from_slice(&normed);
9208 }
9209 }
9210 continue;
9211 }
9212 }
9213 if !h_ready {
9214 fill_h(&mut h, self);
9215 h_ready = true;
9216 }
9217 if task_mask.is_none() && self.verify_exact_moe {
9218 let positions: Vec<_> = (start_pos..start_pos + b).collect();
9219 match self.mimo_graph_layer_rows(li, &mut h, &positions) {
9220 crate::gpu::BatchGraphOutcome::Completed => continue,
9221 crate::gpu::BatchGraphOutcome::Failed => return h,
9222 crate::gpu::BatchGraphOutcome::Declined => {},
9223 }
9224 }
9225 #[cfg(feature = "gpu")]
9226 self.pull_lagging_host_kv(li, li + 1, start_pos);
9227 let lw = &self.weights.layers[self.phys_layer(li)];
9228 let t_attn = std::time::Instant::now();
9229 match &lw.attn {
9231 AttnKind::Kda(w) => {
9232 let cfg = self.kda_cfg.expect("kda layer without kda_cfg");
9234 let mut normed = vec![0.0f32; b * hs];
9235 for bi in 0..b {
9236 inference::rms_norm_into(
9237 &h[bi * hs..(bi + 1) * hs],
9238 &lw.input_norm,
9239 eps,
9240 norm_style,
9241 &mut normed[bi * hs..(bi + 1) * hs],
9242 );
9243 }
9244 let attn = crate::linear_core::kda_forward_batch(
9245 &normed,
9246 b,
9247 w,
9248 &cfg,
9249 &mut self.kv_cache.layers[li].linear_state,
9250 pool.as_deref(),
9251 );
9252 for (dst, &a) in h.iter_mut().zip(&attn) {
9253 *dst += a;
9254 }
9255 }
9256 AttnKind::LinearGdn(w) => {
9257 let cfg = self.gdn_cfg.expect("gdn layer without gdn_cfg");
9259 let mut normed = vec![0.0f32; b * hs];
9260 for bi in 0..b {
9261 let r = inference::rms_norm(
9262 &h[bi * hs..(bi + 1) * hs],
9263 &lw.input_norm,
9264 eps,
9265 norm_style,
9266 );
9267 normed[bi * hs..(bi + 1) * hs].copy_from_slice(&r);
9268 }
9269 let attn = crate::linear_core::gdn_forward_batch(
9270 &normed,
9271 b,
9272 w,
9273 &cfg,
9274 &mut self.kv_cache.layers[li].linear_state,
9275 pool.as_deref(),
9276 );
9277 for (dst, &a) in h.iter_mut().zip(&attn) {
9278 *dst += a;
9279 }
9280 }
9281 AttnKind::ShortConv(w) => {
9282 let cfg = self
9285 .short_conv_cfg
9286 .expect("short-conv layer without short_conv_cfg");
9287 let mut normed = vec![0.0f32; b * hs];
9288 for bi in 0..b {
9289 inference::rms_norm_into(
9290 &h[bi * hs..(bi + 1) * hs],
9291 &lw.input_norm,
9292 eps,
9293 norm_style,
9294 &mut normed[bi * hs..(bi + 1) * hs],
9295 );
9296 }
9297 let attn = short_conv_forward_batch(
9298 &normed,
9299 b,
9300 w,
9301 &cfg,
9302 &mut self.kv_cache.layers[li].linear_state,
9303 pool.as_deref(),
9304 );
9305 for (dst, &a) in h.iter_mut().zip(&attn) {
9306 *dst += a;
9307 }
9308 }
9309 AttnKind::Mla(w) => {
9310 let inv_freq_l = self.layer_inv_freq(li);
9313 let rs = self.layer_rope_scale(li);
9314 let mut normed = vec![0.0f32; hs];
9315 for bi in 0..b {
9316 inference::rms_norm_into(
9317 &h[bi * hs..(bi + 1) * hs],
9318 &lw.input_norm,
9319 eps,
9320 norm_style,
9321 &mut normed,
9322 );
9323 let ao = mla_attention(
9324 w,
9325 &normed,
9326 &mut self.kv_cache.layers[li],
9327 start_pos + bi,
9328 &inv_freq_l,
9329 rs,
9330 eps,
9331 pool.as_deref(),
9332 );
9333 for (dst, &a) in h[bi * hs..(bi + 1) * hs].iter_mut().zip(&ao) {
9334 *dst += a;
9335 }
9336 }
9337 }
9338 AttnKind::Full {
9339 wq,
9340 wk,
9341 wv,
9342 wo,
9343 q_norm,
9344 k_norm,
9345 output_gate,
9346 softplus_gate,
9347 bias,
9348 } => {
9349 let mut normed = vec![0.0f32; b * hs];
9353 for bi in 0..b {
9354 inference::rms_norm_into(
9355 &h[bi * hs..(bi + 1) * hs],
9356 &lw.input_norm,
9357 eps,
9358 norm_style,
9359 &mut normed[bi * hs..(bi + 1) * hs],
9360 );
9361 }
9362 let inv_freq_l = self.layer_inv_freq(li);
9363 let (nkv_l, hd_l, rd_l) = self.layer_geom(li);
9364 let cfg = QwenAttnCfg {
9365 num_heads: self.layer_num_heads(li),
9366 num_kv_heads: nkv_l,
9367 head_dim: hd_l,
9368 hidden_size: hs,
9369 position: start_pos,
9370 inv_freq: &inv_freq_l,
9371 rotary_dim: rd_l,
9372 scale: self.attn_scale,
9373 softcap: self.attn_softcap,
9374 window: self.layer_window(li),
9375 v_norm: self.attn_v_norm,
9376 qk_norm_after_rope: self.qk_norm_after_rope,
9377 gate_sigmoid: self.proj_gate_sigmoid,
9378 q_norm: q_norm.as_deref(),
9379 k_norm: k_norm.as_deref(),
9380 output_gate: *output_gate,
9381 softplus_gate: softplus_gate
9382 .as_ref()
9383 .map(|(gate, per_head)| (gate, *per_head)),
9384 rope_scale: self.layer_rope_scale(li),
9385 bias: bias
9386 .as_ref()
9387 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
9388 rms_eps: eps,
9389 norm_style,
9390 pool: pool.as_deref(),
9391 v_head_dim: self.layer_v_dim(li),
9392 };
9393 #[cfg(feature = "gpu")]
9394 let _mirror = self
9395 .prefill_mirror_target(li)
9396 .map(crate::gpu::enter_prefill_mirror);
9397 let mut attn = attention::qwen_attention_batch(
9398 &normed,
9399 b,
9400 wq,
9401 wk,
9402 wv,
9403 wo,
9404 &mut self.kv_cache.layers[li],
9405 &cfg,
9406 );
9407 if let Some(w) = &lw.attn_out_norm {
9408 for bi in 0..b {
9409 inference::rms_norm_into(
9410 &attn[bi * hs..(bi + 1) * hs],
9411 w,
9412 eps,
9413 norm_style,
9414 &mut normed[bi * hs..(bi + 1) * hs],
9415 );
9416 }
9417 attn.copy_from_slice(&normed);
9418 }
9419 for (dst, &a) in h.iter_mut().zip(&attn) {
9420 *dst += a;
9421 }
9422 }
9423 AttnKind::Bounded(w) => {
9424 let mut normed = vec![0.0f32; b * hs];
9427 for bi in 0..b {
9428 inference::rms_norm_into(
9429 &h[bi * hs..(bi + 1) * hs],
9430 &lw.input_norm,
9431 eps,
9432 norm_style,
9433 &mut normed[bi * hs..(bi + 1) * hs],
9434 );
9435 }
9436 let rope = self
9437 .bounded_rope
9438 .clone()
9439 .expect("bounded layer without an installed rotation table");
9440 let cfg = crate::bounded::BoundedAttnCfg {
9441 num_heads: self.num_heads,
9442 num_kv_heads: self.num_kv_heads,
9443 head_dim: self.head_dim,
9444 hidden_size: hs,
9445 scale: self.attn_scale,
9446 rope: &rope,
9447 pool: pool.as_deref(),
9448 };
9449 let mut attn = crate::bounded::bounded_attention_batch(
9450 &normed,
9451 b,
9452 w,
9453 &mut self.kv_cache.layers[li],
9454 &cfg,
9455 );
9456 if let Some(wn) = &lw.attn_out_norm {
9457 for bi in 0..b {
9458 inference::rms_norm_into(
9459 &attn[bi * hs..(bi + 1) * hs],
9460 wn,
9461 eps,
9462 norm_style,
9463 &mut normed[bi * hs..(bi + 1) * hs],
9464 );
9465 }
9466 attn.copy_from_slice(&normed);
9467 }
9468 for (dst, &a) in h.iter_mut().zip(&attn) {
9469 *dst += a;
9470 }
9471 attention::recycle_buf(&mut attn);
9472 }
9473 AttnKind::Linear(w) => {
9474 for bi in 0..b {
9475 let normed = inference::rms_norm(
9476 &h[bi * hs..(bi + 1) * hs],
9477 &lw.input_norm,
9478 eps,
9479 norm_style,
9480 );
9481 vmf_phase_forward(
9482 &normed,
9483 w,
9484 &self.vmf_cfg.expect("linear layer without vmf_cfg"),
9485 &mut self.kv_cache.layers[li].linear_state,
9486 pool.as_deref(),
9487 )
9488 .iter()
9489 .enumerate()
9490 .for_each(|(i, &a)| h[bi * hs + i] += a);
9491 }
9492 }
9493 }
9494
9495 let lw = &self.weights.layers[self.phys_layer(li)];
9497 let mut post = vec![0.0f32; b * hs];
9498 for bi in 0..b {
9499 let r =
9500 inference::rms_norm(&h[bi * hs..(bi + 1) * hs], &lw.post_norm, eps, norm_style);
9501 post[bi * hs..(bi + 1) * hs].copy_from_slice(&r);
9502 }
9503 let mask_row = task_mask
9506 .filter(|m| m.ffn_active_count(li) < self.intermediate_size)
9507 .and_then(|m| m.ffn_masks.get(li))
9508 .map(|v| v.as_slice());
9509 let attn_ns = t_attn.elapsed().as_nanos() as u64;
9510 let t_ffn = std::time::Instant::now();
9511 let mut ffn = match &lw.ffn {
9512 FfnKind::Dense(d) if !d.segs.is_empty() => {
9513 tube_ffn(d, &post, b, pool.as_deref(), mask_row)
9514 }
9515 FfnKind::Dense(d) => dense_ffn_batch(d, &post, b, pool.as_deref(), mask_row),
9516 FfnKind::Moe(m) if self.verify_exact_moe && self.mimo_moe.is_dynamic(li, false) => {
9517 moe_ffn_banked_rows(&mut self.mimo_moe, li, m, &post, b, hs, pool.as_deref())
9518 }
9519 FfnKind::Moe(m) if self.verify_exact_moe => {
9520 moe_ffn_rows_exact(m, &post, b, hs, pool.as_deref())
9521 }
9522 FfnKind::Moe(m) if self.mimo_moe.is_dynamic(li, false) => {
9525 let before = m.stats.borrow().clone();
9526 let out = crate::gpu::cpu_scope(|| {
9527 moe_ffn_batch(m, &post, b, hs, pool.as_deref(), None)
9528 });
9529 self.mimo_moe.prime(li, m, &before);
9530 out
9531 }
9532 FfnKind::Moe(m) => moe_ffn_batch(m, &post, b, hs, pool.as_deref(), None),
9533 FfnKind::DenseMoe(dm) => {
9536 let mut out = vec![0.0f32; b * hs];
9537 for bi in 0..b {
9538 let r = dense_moe_ffn(
9539 dm,
9540 &post[bi * hs..(bi + 1) * hs],
9541 &h[bi * hs..(bi + 1) * hs],
9542 eps,
9543 norm_style,
9544 pool.as_deref(),
9545 );
9546 out[bi * hs..(bi + 1) * hs].copy_from_slice(&r);
9547 }
9548 out
9549 }
9550 };
9551 if prefill_prof_on() {
9552 PREFILL_SPLIT[0].fetch_add(attn_ns, std::sync::atomic::Ordering::Relaxed);
9553 PREFILL_SPLIT[1].fetch_add(t_ffn.elapsed().as_nanos() as u64, std::sync::atomic::Ordering::Relaxed);
9554 }
9555 if let Some(w) = &lw.ffn_out_norm {
9556 for bi in 0..b {
9557 inference::rms_norm_into(
9558 &ffn[bi * hs..(bi + 1) * hs],
9559 w,
9560 eps,
9561 norm_style,
9562 &mut post[bi * hs..(bi + 1) * hs],
9563 );
9564 }
9565 ffn.copy_from_slice(&post);
9566 }
9567 for (dst, &f) in h.iter_mut().zip(&ffn) {
9568 *dst += f;
9569 }
9570 if let Some(sc) = lw.layer_scale {
9571 for v in h.iter_mut() {
9572 *v *= sc;
9573 }
9574 }
9575 if self.layer_dump.is_some() {
9577 for bi in 0..b {
9578 self.dump_layer_row(start_pos + bi, li, &h[bi * hs..(bi + 1) * hs]);
9579 }
9580 }
9581 if let Ok(tp) = std::env::var("CMF_TRACE_POS") {
9582 if let Ok(t) = tp.parse::<usize>() {
9583 if t >= start_pos && t < start_pos + b {
9584 let bi = t - start_pos;
9585 let row = &h[bi * hs..(bi + 1) * hs];
9586 let n: f32 = row.iter().map(|x| x * x).sum::<f32>().sqrt();
9587 eprintln!(
9588 "BATCH pos {t} after layer {li}: |h| = {n:.6} h0 {:.6} h1 {:.6}",
9589 row[0], row[1]
9590 );
9591 }
9592 }
9593 }
9594 if std::env::var("CMF_DEBUG_LAYERS").is_ok() {
9598 let row = &h[(b - 1) * hs..b * hs];
9599 let rms =
9600 (row.iter().map(|&v| (v as f64) * (v as f64)).sum::<f64>() / hs as f64).sqrt();
9601 let mx = row.iter().fold(0f32, |m, &v| m.max(v.abs()));
9602 eprintln!(
9603 "layer {li:>3} {:>10} ffn={:<5} rms={rms:>12.4} max={mx:>12.4}",
9604 match &self.weights.layers[self.phys_layer(li)].attn {
9605 AttnKind::LinearGdn(_) => "gdn",
9606 AttnKind::Linear(_) => "vmf",
9607 AttnKind::ShortConv(_) => "conv",
9608 _ => "attn",
9609 },
9610 match &lw.ffn {
9611 FfnKind::Moe(_) => "moe",
9612 FfnKind::Dense(_) => "dense",
9613 FfnKind::DenseMoe(_) => "dense+moe",
9614 },
9615 );
9616 }
9617 if self.is_loop_end(li) && li + 1 < self.num_layers {
9619 for bi in 0..b {
9620 let normed = inference::rms_norm(
9621 &h[bi * hs..(bi + 1) * hs],
9622 &self.weights.final_norm,
9623 eps,
9624 norm_style,
9625 );
9626 h[bi * hs..(bi + 1) * hs].copy_from_slice(&normed);
9627 }
9628 }
9629 if std::env::var("CMF_TRACE_H").is_ok() {
9630 let n = h[..hs].iter().map(|v| v.abs()).sum::<f32>() / hs as f32;
9631 let mx = h[..hs].iter().fold(0.0f32, |a, &v| a.max(v.abs()));
9632 eprintln!(
9633 "layer {li}: mean|h|={n:.4} max|h|={mx:.2} scale={:?}",
9634 lw.layer_scale
9635 );
9636 }
9637 }
9638 crate::gpu::set_layer(-1); if prefill_prof_on() {
9640 let ms = |a: &std::sync::atomic::AtomicU64| {
9641 a.load(std::sync::atomic::Ordering::Relaxed) as f64 / 1e6
9642 };
9643 let sp = &attention::ATTN_SPLIT;
9644 let (qkv_enq, qkv_wait, ca_wait, zi) = wgpu_prefill_counters();
9645 eprintln!(
9646 "prefill-split: attention {:.1} ms (proj {:.1} [qkv enqueue {:.1} wait {:.1}, \
9647 host gate {:.1}], host loop {:.1}, attend {:.1} \
9648 [q pack {:.1}, device {:.1} of which readback {:.1}], o-proj {:.1}), \
9649 ffn {:.1} ms (cumulative; zi_mm GEMMs {zi})",
9650 ms(&PREFILL_SPLIT[0]),
9651 ms(&sp[0]),
9652 qkv_enq,
9653 qkv_wait,
9654 ms(&sp[6]),
9655 ms(&sp[1]),
9656 ms(&sp[2]),
9657 ms(&sp[4]),
9658 ms(&sp[5]),
9659 ca_wait,
9660 ms(&sp[3]),
9661 ms(&PREFILL_SPLIT[1]),
9662 );
9663 let [calls, attended, wo, ring, reseed, f16] = wgpu_mirror_counters();
9664 eprintln!(
9665 "prefill-mirror: {calls} calls, {attended} attended ({wo} with O on the card), \
9666 {ring} ring refusals, {reseed} reseeds, {f16} f16 fallbacks, {} upload \
9667 fallbacks (cumulative)",
9668 attention::MIRROR_UPLOADS.load(std::sync::atomic::Ordering::Relaxed),
9669 );
9670 }
9671 self.o1_progress();
9676 self.swa_trim_tails();
9678 h
9679 }
9680
9681 fn embed_single(&self, id: u32) -> Vec<f32> {
9683 let mut out = vec![0.0f32; self.hidden_size];
9684 if (id as usize) < self.weights.embed_tokens.rows() {
9685 self.weights.embed_tokens.row_f32(id as usize, &mut out);
9686 }
9687 if self.embed_multiplier != 1.0 {
9688 for v in out.iter_mut() {
9689 *v *= self.embed_multiplier;
9690 }
9691 }
9692 if self.dsv4.is_some()
9696 || self.dsv41.is_some()
9697 || self.qwen4_exp.is_some()
9698 {
9699 let mut v = vec![0.0f32; self.hidden_size.max(1)];
9700 v[0] = id as f32;
9701 return v;
9702 }
9703 if let Some(b) = &self.g3n {
9706 return b.0.extend_embedding(id, &out, self.pool.as_deref());
9707 }
9708 out
9709 }
9710
9711 #[cfg(target_os = "macos")]
9717 fn chunk_run_gpu(
9718 &mut self,
9719 li0: usize,
9720 h: &mut [f32],
9721 b: usize,
9722 pos0: usize,
9723 embed_ids: Option<&[u32]>,
9724 cap: usize,
9725 ) -> usize {
9726 if !crate::gpu::enabled_here()
9730 || std::env::var("CMF_GPU_CHUNK")
9731 .map(|v| v == "0")
9732 .unwrap_or(false)
9733 || b < 32
9734 || (self.swa.is_some() && !self.metal_graph_swa())
9737 || self.global_attn.is_some()
9738 || (self.graph_attn_decline_reason().is_some() && !self.metal_graph_swa())
9740 || self.o1_active()
9743 || self.attn_v_norm
9744 || (self.attn_scale - 1.0 / (self.head_dim as f32).sqrt()).abs() > 1e-9
9745 {
9746 return li0;
9747 }
9748 let Some(model) = self.model.clone() else {
9749 return li0;
9750 };
9751 let (nh, nkv, hd, hs) = (
9752 self.num_heads,
9753 self.num_kv_heads,
9754 self.head_dim,
9755 self.hidden_size,
9756 );
9757 let loop_end = if self.loop_final_norm {
9761 ((li0 / self.physical_layers) + 1) * self.physical_layers
9762 } else {
9763 self.num_layers
9764 };
9765 let mut layers: Vec<crate::gpu_metal::ChunkLayer> = Vec::new();
9766 let mut stored_at: Vec<usize> = Vec::new();
9767 let run_end = self.num_layers.min(loop_end).min(cap);
9768 let tables: Vec<std::sync::Arc<Vec<f32>>> =
9771 (li0..run_end).map(|li| self.layer_inv_freq(li)).collect();
9772 for li in li0..run_end {
9773 let lw = &self.weights.layers[self.phys_layer(li)];
9774 if lw.attn_out_norm.is_some() || lw.ffn_out_norm.is_some() || lw.layer_scale.is_some() {
9775 break;
9776 }
9777 let AttnKind::Full {
9778 wq,
9779 wk,
9780 wv,
9781 wo,
9782 q_norm,
9783 k_norm,
9784 output_gate: false,
9785 softplus_gate,
9786 bias,
9787 } = &lw.attn
9788 else {
9789 break;
9790 };
9791 let head_gate = match softplus_gate {
9793 None => None,
9794 Some((g, true)) if self.proj_gate_sigmoid => match g.f32_parts() {
9795 Some((d, r, c)) if r == nh && c == hs => Some(d),
9796 _ => break,
9797 },
9798 Some(_) => break,
9799 };
9800 let FfnKind::Dense(d) = &lw.ffn else { break };
9801 if !matches!(d.act, Act::Silu | Act::Gelu) || !d.segs.is_empty() {
9802 break;
9803 }
9804 fn cw(t: &QTensor) -> Option<(usize, usize, usize, &[f32])> {
9809 t.q8_row_parts()
9810 .or_else(|| t.q4t_parts().map(|(i, r, c)| (i, r, c, &[][..])))
9811 .or_else(|| t.q4tp_parts().map(|(i, r, c)| (i, r, c, &[][..])))
9812 }
9813 let parts = (
9814 cw(wq),
9815 cw(wk),
9816 cw(wv),
9817 cw(wo),
9818 cw(&d.gate_proj),
9819 cw(&d.up_proj),
9820 cw(&d.down_proj),
9821 );
9822 let (Some(pq), Some(pk), Some(pv), Some(po), Some(pg), Some(pu), Some(pd)) = parts
9823 else {
9824 break;
9825 };
9826 let layer = &self.kv_cache.layers[li];
9827 if layer.mode != crate::kv_cache::KvMode::F32 || layer.o1.is_some() {
9828 break;
9829 }
9830 debug_assert!(
9832 layer.base() == 0
9833 || self
9834 .layer_window(li)
9835 .is_some_and(|w| layer.head_len(0) + 1 >= w)
9836 );
9837 stored_at.push(layer.head_len(0));
9838 layers.push(crate::gpu_metal::ChunkLayer {
9839 model: &model,
9840 kv_id: self.graph_kv_id,
9841 layer: li,
9842 wq: pq,
9843 wk: pk,
9844 wv: pv,
9845 wo: po,
9846 gate: pg,
9847 up: pu,
9848 down: pd,
9849 input_norm: &lw.input_norm,
9850 post_norm: &lw.post_norm,
9851 bias: bias
9852 .as_ref()
9853 .map(|(a, bb, cc)| (a.as_slice(), bb.as_slice(), cc.as_slice())),
9854 q_norm: q_norm.as_deref(),
9855 k_norm: k_norm.as_deref(),
9856 inv_freq: &tables[li - li0],
9857 rd: self.layer_geom(li).2,
9858 nh,
9859 nkv,
9860 hd,
9861 hs,
9862 inter: d.gate_proj.rows(),
9863 gemma: matches!(self.norm_style, cortiq_core::NormStyle::Gemma),
9864 late_qk_norm: self.qk_norm_after_rope,
9865 eps: self.rms_eps as f32,
9866 window: self.layer_window(li),
9867 head_gate,
9868 gelu: d.act == Act::Gelu,
9869 });
9870 }
9871 if layers.is_empty() {
9872 return li0;
9873 }
9874 let row = nkv * hd;
9875 let mut store: Vec<(Vec<f32>, Vec<f32>, Vec<f32>)> = stored_at
9876 .iter()
9877 .map(|&st| (vec![0f32; b * row], vec![0f32; b * row], vec![0f32; st + b]))
9878 .collect();
9879 let mut io: Vec<crate::gpu_metal::ChunkIo> = Vec::with_capacity(layers.len());
9880 for (i, (ok, ov, oi)) in store.iter_mut().enumerate() {
9881 let li = layers[i].layer;
9882 let layer = &self.kv_cache.layers[li];
9883 io.push(crate::gpu_metal::ChunkIo {
9884 cpu_stored: stored_at[i],
9885 cpu_gen: layer.generation(),
9886 cpu_k: (0..nkv).map(|g| layer.head_keys(g)).collect(),
9887 cpu_v: (0..nkv).map(|g| layer.head_values(g)).collect(),
9888 out_k: ok,
9889 out_v: ov,
9890 imp: oi,
9891 });
9892 }
9893 let n_run = layers.len();
9894 let last = layers.last().map(|l| l.layer + 1).unwrap_or(li0);
9895 let ep = embed_ids.and_then(|ids| {
9898 self.weights
9899 .embed_tokens
9900 .q8_row_parts()
9901 .map(|(idx, rows, _c, rs)| crate::gpu_metal::ChunkEmbed {
9902 idx,
9903 rows,
9904 row_scale: rs,
9905 ids,
9906 mult: self.embed_multiplier,
9907 })
9908 });
9909 if embed_ids.is_some() && ep.is_none() {
9910 return li0;
9911 }
9912 if !crate::gpu_metal::chunk_run_gpu(&layers, &mut io, h, b, pos0, ep.as_ref()) {
9913 return li0;
9914 }
9915 drop(io);
9916 drop(layers);
9917 for (i, (ok, ov, oi)) in store.iter().enumerate().take(n_run) {
9920 let li = li0 + i;
9921 let layer = &mut self.kv_cache.layers[li];
9922 for bi in 0..b {
9923 layer.append(
9924 &ok[bi * row..(bi + 1) * row],
9925 &ov[bi * row..(bi + 1) * row],
9926 &[],
9927 );
9928 }
9929 layer.accumulate_imp(oi);
9930 }
9931 last
9932 }
9933
9934 fn layer_is_local(&self, li: usize) -> bool {
9937 if let Some(layers) = &self.sliding_layers {
9938 return layers.get(li).copied().unwrap_or(false);
9939 }
9940 match self.swa {
9941 Some((_, pattern)) => (li + 1) % pattern.max(1) != 0,
9942 None => false,
9943 }
9944 }
9945
9946 fn layer_inv_freq(&self, li: usize) -> std::sync::Arc<Vec<f32>> {
9949 if self.layer_is_local(li) {
9950 if let Some(f) = &self.inv_freq_local {
9951 return f.clone();
9952 }
9953 } else if let Some(f) = &self.inv_freq_global {
9954 return f.clone();
9955 }
9956 self.inv_freq.clone()
9957 }
9958
9959 fn layer_window(&self, li: usize) -> Option<usize> {
9961 self.swa
9962 .and_then(|(w, _)| self.layer_is_local(li).then_some(w))
9963 }
9964
9965 fn swa_trim_tails(&mut self) {
9974 let Some((slack, align)) = self.swa_trim else {
9975 return;
9976 };
9977 if self.mtp.is_some() || self.mimo_mtp.is_some() || !self.dsv4_mtp.is_empty() {
9978 return;
9979 }
9980 for li in 0..self.num_layers.min(self.kv_cache.layers.len()) {
9981 if let Some(w) = self.layer_window(li) {
9982 self.kv_cache.layers[li].trim_window(w, slack, align);
9983 }
9984 }
9985 }
9986
9987 fn layer_num_heads(&self, li: usize) -> usize {
9988 self.attention_heads_per_layer
9989 .as_ref()
9990 .and_then(|v| v.get(li).copied())
9991 .unwrap_or(self.num_heads)
9992 }
9993
9994 fn layer_rope_scale(&self, li: usize) -> f32 {
9995 if self.layer_is_local(li) {
9996 self.rope_scale_local
9997 } else {
9998 self.rope_scale
9999 }
10000 }
10001
10002 fn layer_geom(&self, li: usize) -> (usize, usize, usize) {
10005 if !self.layer_is_local(li) {
10006 if let Some((ghd, gkv)) = self.global_attn {
10007 return (gkv, ghd, ghd);
10008 }
10009 }
10010 (
10011 self.layer_num_kv_heads(li),
10012 self.head_dim,
10013 if self.layer_is_local(li) {
10014 self.rotary_dim_local.unwrap_or(self.rotary_dim)
10015 } else {
10016 self.rotary_dim
10017 },
10018 )
10019 }
10020
10021 fn layer_num_kv_heads(&self, li: usize) -> usize {
10024 self.kv_heads_per_layer
10025 .as_ref()
10026 .and_then(|v| v.get(self.phys_layer(li)).copied())
10027 .unwrap_or(self.num_kv_heads)
10028 }
10029
10030 fn layer_v_dim(&self, li: usize) -> usize {
10032 let (_, hd, _) = self.layer_geom(li);
10033 self.v_head_dim.unwrap_or(hd).min(hd)
10034 }
10035
10036 pub fn set_attn_geometry(
10045 &mut self,
10046 kv_heads_per_layer: Option<Vec<usize>>,
10047 v_head_dim: Option<usize>,
10048 ) -> Result<(), String> {
10049 if kv_heads_per_layer.is_some() || v_head_dim.is_some() {
10050 if self.global_attn.is_some() {
10051 return Err(
10052 "per-layer KV heads / v_head_dim cannot combine with Gemma-4 global \
10053 attention geometry"
10054 .into(),
10055 );
10056 }
10057 if self
10058 .weights
10059 .layers
10060 .iter()
10061 .any(|lw| matches!(lw.attn, AttnKind::Mla(_)))
10062 {
10063 return Err("per-layer KV heads / v_head_dim cannot combine with MLA".into());
10064 }
10065 }
10066 if let Some(vd) = v_head_dim {
10067 if vd == 0 || vd > self.head_dim {
10068 return Err(format!(
10069 "v_head_dim {vd} must be in 1..={} (head_dim)",
10070 self.head_dim
10071 ));
10072 }
10073 }
10074 if let Some(v) = &kv_heads_per_layer {
10075 if v.len() != self.physical_layers {
10076 return Err(format!(
10077 "kv_heads_per_layer has {} entries, expected {} layers",
10078 v.len(),
10079 self.physical_layers
10080 ));
10081 }
10082 for (li, &nkv) in v.iter().enumerate() {
10083 let is_attn = matches!(
10084 self.weights.layers.get(li).map(|lw| &lw.attn),
10085 Some(AttnKind::Full { .. }) | None
10086 );
10087 if !is_attn {
10088 continue;
10089 }
10090 let nh = self
10091 .attention_heads_per_layer
10092 .as_ref()
10093 .and_then(|h| h.get(li).copied())
10094 .unwrap_or(self.num_heads);
10095 if nkv == 0 || nh % nkv != 0 {
10096 return Err(format!(
10097 "layer {li}: {nkv} KV heads must be nonzero and divide {nh} Q heads"
10098 ));
10099 }
10100 }
10101 }
10102 self.kv_heads_per_layer = kv_heads_per_layer;
10103 self.v_head_dim = v_head_dim.filter(|&vd| vd != self.head_dim);
10104 if self.kv_heads_per_layer.is_some() {
10105 for li in 0..self.kv_cache.layers.len() {
10106 let full = matches!(
10107 self.weights
10108 .layers
10109 .get(self.phys_layer(li))
10110 .map(|lw| &lw.attn),
10111 Some(AttnKind::Full { .. })
10112 );
10113 let nkv = self.layer_num_kv_heads(li);
10114 let cache = &self.kv_cache.layers[li];
10115 if full && (cache.num_kv_heads != nkv || cache.head_dim != self.head_dim) {
10116 let sinks = cache.sinks.clone();
10117 self.kv_cache.layers[li] =
10118 crate::kv_cache::LayerKvCache::new(nkv, self.head_dim);
10119 self.kv_cache.layers[li].sinks = sinks;
10120 }
10121 }
10122 }
10123 Ok(())
10124 }
10125
10126 pub fn set_layer_sinks(&mut self, phys: usize, sinks: Vec<f32>) -> Result<(), String> {
10130 let Some(lw) = self.weights.layers.get(phys) else {
10131 return Err(format!("sinks for layer {phys}: no such layer"));
10132 };
10133 if !matches!(lw.attn, AttnKind::Full { .. }) {
10134 return Err(format!(
10135 "sinks for layer {phys}: only softmax (Full) attention layers take sinks"
10136 ));
10137 }
10138 let nh = self
10139 .attention_heads_per_layer
10140 .as_ref()
10141 .and_then(|h| h.get(phys).copied())
10142 .unwrap_or(self.num_heads);
10143 if sinks.len() != nh {
10144 return Err(format!(
10145 "sinks for layer {phys}: {} values, expected one per Q head ({nh})",
10146 sinks.len()
10147 ));
10148 }
10149 if let Some(bad) = sinks.iter().find(|v| !v.is_finite()) {
10150 return Err(format!("sinks for layer {phys}: non-finite value {bad}"));
10151 }
10152 for li in 0..self.kv_cache.layers.len() {
10153 if self.phys_layer(li) == phys {
10154 self.kv_cache.layers[li].sinks = Some(sinks.clone());
10155 }
10156 }
10157 Ok(())
10158 }
10159
10160 pub fn graph_attn_decline_reason(&self) -> Option<&'static str> {
10170 if self.kv_heads_per_layer.is_some() {
10171 return Some("per-layer KV head counts");
10172 }
10173 if self.v_head_dim.is_some_and(|vd| vd != self.head_dim) {
10174 return Some("V heads narrower than Q/K heads");
10175 }
10176 if self.kv_cache.layers.iter().any(|l| l.sinks.is_some()) {
10177 return Some("learned attention sinks");
10178 }
10179 if self.swa.is_some() || self.sliding_layers.is_some() {
10180 return Some("sliding-window layers");
10181 }
10182 None
10183 }
10184
10185 #[cfg(target_os = "macos")]
10199 fn metal_graph_swa(&self) -> bool {
10200 (self.swa.is_some() || self.sliding_layers.is_some())
10201 && self.attention_heads_per_layer.is_none()
10202 && !self.attn_v_norm
10203 && self.attn_softcap == 0.0
10204 && self.kv_heads_per_layer.is_none()
10205 && self.v_head_dim.map_or(true, |vd| vd == self.head_dim)
10206 && !self.kv_cache.layers.iter().any(|l| l.sinks.is_some())
10207 && self.global_attn.is_none()
10208 && self.inv_freq_global.is_none()
10209 && (0..self.num_layers).all(|li| self.layer_rope_scale(li) == 1.0)
10210 && !(0..self.num_layers).any(|li| {
10211 self.layer_is_local(li)
10212 && self.inv_freq_local.is_none()
10213 && self.rotary_dim_local.is_some_and(|r| r != self.rotary_dim)
10214 })
10215 }
10216
10217 pub fn wgpu_graph_attn_decline(&self) -> Option<&'static str> {
10224 self.graph_attn_decline_reason()?;
10225 if self.global_attn.is_some() {
10226 return Some("per-layer head width (Gemma-4 global layers) with per-layer geometry");
10227 }
10228 if self.attention_heads_per_layer.is_some() {
10229 return Some("per-layer Q head counts with per-layer geometry");
10230 }
10231 if self.attn_v_norm {
10232 return Some("V norm with per-layer geometry");
10233 }
10234 if (0..self.num_layers).any(|li| self.layer_rope_scale(li) != 1.0) {
10235 return Some("scaled RoPE positions with per-layer geometry");
10236 }
10237 if self.weights.layers.iter().any(|lw| {
10238 lw.attn_out_norm.is_some() || lw.ffn_out_norm.is_some() || lw.layer_scale.is_some()
10239 }) {
10240 return Some("sandwich norms / layer scale with per-layer geometry");
10241 }
10242 if self.weights.layers.iter().any(|lw| {
10243 matches!(
10244 &lw.attn,
10245 AttnKind::Full {
10246 output_gate: true,
10247 ..
10248 }
10249 )
10250 }) && self.v_head_dim.is_some()
10251 {
10252 return Some("gated attention with V narrower than K");
10253 }
10254 if (0..self.num_layers).any(|li| {
10255 self.layer_is_local(li)
10256 && self.inv_freq_local.is_none()
10257 && self.rotary_dim_local.is_some_and(|r| r != self.rotary_dim)
10258 }) {
10259 return Some("local rotary width without a local RoPE table");
10260 }
10261 None
10262 }
10263
10264 fn graph_attn_geom(&self, li: usize) -> Option<crate::gpu::GraphAttnGeom<'_>> {
10269 self.graph_attn_decline_reason()?;
10270 let (nkv, _hd, rd) = self.layer_geom(li);
10271 let invf: &[f32] = if self.layer_is_local(li) {
10272 match &self.inv_freq_local {
10273 Some(f) => f.as_slice(),
10274 None => self.inv_freq.as_slice(),
10275 }
10276 } else {
10277 match &self.inv_freq_global {
10278 Some(f) => f.as_slice(),
10279 None => self.inv_freq.as_slice(),
10280 }
10281 };
10282 Some(crate::gpu::GraphAttnGeom {
10283 nkv,
10284 dv: self.layer_v_dim(li),
10285 rd,
10286 invf,
10287 window: self.layer_window(li),
10288 sink: self.kv_cache.layers[li].sinks.as_deref(),
10289 })
10290 }
10291
10292 #[cfg(feature = "gpu")]
10304 fn prefill_mirror_target(&self, li: usize) -> Option<crate::gpu::PrefillMirror> {
10305 if !self.proj_gate_sigmoid
10306 || std::env::var("CMF_PREFILL_MIRROR").as_deref() == Ok("0")
10307 || !crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Decode)
10308 || self.graph_refused()
10309 || self.o1_active()
10310 || self.attn_softcap > 0.0
10311 || self.wgpu_graph_attn_decline().is_some()
10312 {
10313 return None;
10314 }
10315 let g = self.graph_attn_geom(li)?;
10316 let (nkv, hd, _) = self.layer_geom(li);
10317 if g.nkv != nkv || g.dv != hd || g.sink.is_some() || hd != self.layer_geom(0).1 {
10319 return None;
10320 }
10321 Some(crate::gpu::PrefillMirror {
10322 kv_id: self.graph_kv_id,
10323 layer: li,
10324 limit: self.kv_cache.max_seq_len,
10325 })
10326 }
10327
10328 #[cfg(feature = "gpu")]
10337 fn pull_lagging_host_kv(&mut self, from: usize, upto: usize, position: usize) {
10338 let kv_id = self.graph_kv_id;
10339 for li in from..upto.min(self.num_layers) {
10340 if !matches!(
10341 self.weights.layers[self.phys_layer(li)].attn,
10342 AttnKind::Full { .. }
10343 ) {
10344 continue;
10345 }
10346 let host = self.kv_cache.layers[li].pos_len();
10348 if host >= position {
10349 continue;
10350 }
10351 let Some(dev) = crate::gpu::graph_kv_stored(kv_id, li) else {
10352 continue;
10353 };
10354 let to = dev.min(position);
10355 if to <= host {
10356 continue;
10357 }
10358 let (nkv, hd) = {
10359 let c = &self.kv_cache.layers[li];
10360 (c.num_kv_heads, c.head_dim)
10361 };
10362 let Some((k, v, first_valid)) =
10363 crate::gpu::graph_kv_pull_host(kv_id, li, host, to, nkv, hd)
10364 else {
10365 continue;
10366 };
10367 let need_from = match self.layer_window(li) {
10370 Some(w) => host.max((position + 1).saturating_sub(w)),
10371 None => host,
10372 };
10373 if first_valid > need_from {
10374 tracing::warn!(
10375 "layer {li}: device KV rows {host}..{to} no longer resident \
10376 (from {first_valid}); host attention will miss them"
10377 );
10378 }
10379 let row = nkv * hd;
10380 let cache = &mut self.kv_cache.layers[li];
10381 for p in 0..to - host {
10382 cache.append(&k[p * row..(p + 1) * row], &v[p * row..(p + 1) * row], &[]);
10383 }
10384 }
10385 }
10386
10387 fn note_graph_decline(&self, site: &'static str, reason: &'static str) {
10390 let mut seen = self.graph_declines.borrow_mut();
10391 if !seen.iter().any(|&(s, r)| s == site && r == reason) {
10392 tracing::warn!("{site} declined: {reason} (CPU attention path)");
10393 seen.push((site, reason));
10394 }
10395 }
10396
10397 pub fn graph_declines(&self) -> Vec<String> {
10400 self.graph_declines
10401 .borrow()
10402 .iter()
10403 .map(|(site, reason)| format!("{site} declined: {reason} (CPU attention path)"))
10404 .collect()
10405 }
10406
10407 fn dump_layer_row(&self, pos: usize, li: usize, row: &[f32]) {
10416 let Some(dir) = &self.layer_dump else {
10417 return;
10418 };
10419 let mut bytes = Vec::with_capacity(row.len() * 4);
10420 for v in row {
10421 bytes.extend_from_slice(&v.to_le_bytes());
10422 }
10423 let path = dir.join(format!("p{pos:06}_l{li:02}.f32"));
10424 if let Err(e) = std::fs::create_dir_all(dir).and_then(|_| std::fs::write(&path, &bytes)) {
10425 use std::sync::atomic::{AtomicBool, Ordering};
10426 static SAID: AtomicBool = AtomicBool::new(false);
10427 if !SAID.swap(true, Ordering::Relaxed) {
10428 tracing::error!("CMF_LAYER_DUMP: cannot write {}: {e}", path.display());
10429 }
10430 }
10431 }
10432
10433 fn mimo_moe_prepare(&mut self) {
10436 if !self.mimo_moe.is_undecided() {
10437 return;
10438 }
10439 let slot = {
10440 let layers: Vec<(usize, &MoeFfn)> = (0..self.num_layers)
10441 .filter_map(
10442 |li| match &self.weights.layers.get(self.phys_layer(li))?.ffn {
10443 FfnKind::Moe(m) => Some((li, m)),
10444 _ => None,
10445 },
10446 )
10447 .collect();
10448 if layers.is_empty()
10451 || self.physical_layers != self.num_layers
10452 || self.gpu_plan.is_some()
10453 {
10454 crate::mimo_moe::Slot::Off
10455 } else {
10456 let graph_prefix =
10460 self.wgpu_graph_attn_decline().is_none() && crate::gpu::wgpu_graph_default();
10461 crate::mimo_moe::Slot::decide(&layers, self.num_layers, graph_prefix)
10462 }
10463 };
10464 self.mimo_moe = slot;
10465 }
10466
10467 #[cfg(test)]
10468 pub(crate) fn test_graph_kv_id(&self) -> u64 {
10469 self.graph_kv_id
10470 }
10471
10472 pub(crate) fn mimo_graph_layer_rows(
10476 &mut self,
10477 li: usize,
10478 h: &mut [f32],
10479 positions: &[usize],
10480 ) -> crate::gpu::BatchGraphOutcome {
10481 use crate::gpu::BatchGraphOutcome as Out;
10482 let b = positions.len();
10483 if !(1..=4).contains(&b)
10484 || h.len() != b * self.hidden_size
10485 || !self.mimo_moe.is_dynamic(li, true)
10486 || !crate::gpu::enabled_here()
10487 || !crate::gpu::wgpu_active()
10488 || self.o1_active()
10489 || self.physical_layers != self.num_layers
10490 || std::env::var("CMF_GPU_WGPU_GRAPH").as_deref() == Ok("0")
10494 || std::env::var("CMF_MIMO_ATTN_GRAPH").as_deref() == Ok("0")
10495 || self.wgpu_graph_attn_decline().is_some()
10496 {
10497 return Out::Declined;
10498 }
10499 let attn_started = std::time::Instant::now();
10500 let outcome = {
10501 let lw = &self.weights.layers[li];
10502 if lw.attn_out_norm.is_some() || lw.ffn_out_norm.is_some() || lw.layer_scale.is_some() {
10503 return Out::Declined;
10504 }
10505 let FfnKind::Moe(m) = &lw.ffn else {
10506 return Out::Declined;
10507 };
10508 let AttnKind::Full {
10509 wq,
10510 wk,
10511 wv,
10512 wo,
10513 q_norm,
10514 k_norm,
10515 output_gate,
10516 softplus_gate,
10517 bias,
10518 } = &lw.attn
10519 else {
10520 return Out::Declined;
10521 };
10522 if *output_gate || softplus_gate.is_some() {
10523 return Out::Declined;
10524 }
10525 let Some(model) = wq.graph_weight().map(|(m, ..)| m.clone()).or_else(|| {
10526 m.experts
10527 .first()?
10528 .gate_proj
10529 .mapped_q4tp()
10530 .map(|(m, _)| m.clone())
10531 }) else {
10532 return Out::Declined;
10533 };
10534 fn gw<'a>(
10535 t: &'a QTensor,
10536 owner: &std::sync::Arc<cortiq_core::CmfModel>,
10537 ) -> Option<crate::gpu::GraphW<'a>> {
10538 if let Some((m, idx, kind, rs)) = t.graph_weight() {
10539 if m.uid() != owner.uid() || t.has_prism_contract() {
10540 return None;
10541 }
10542 return Some(crate::gpu::GraphW {
10543 idx,
10544 kind,
10545 row_scale: rs,
10546 data: &[],
10547 prism: crate::gpu::GraphPrismOp::None,
10548 affine: false,
10549 });
10550 }
10551 t.as_f32().map(|data| crate::gpu::GraphW {
10552 idx: 0,
10553 kind: 4,
10554 row_scale: &[],
10555 data,
10556 prism: crate::gpu::GraphPrismOp::None,
10557 affine: false,
10558 })
10559 }
10560 let (Some(q), Some(k), Some(v), Some(o)) = (
10561 gw(wq, &model),
10562 gw(wk, &model),
10563 gw(wv, &model),
10564 gw(wo, &model),
10565 ) else {
10566 return Out::Declined;
10567 };
10568 let layer = crate::gpu::GraphLayer {
10569 input_norm: &lw.input_norm,
10570 post_norm: &lw.post_norm,
10571 ffn: crate::gpu::GraphFfn::AttentionOnly,
10572 attn: crate::gpu::GraphAttn::Full {
10573 wq: q,
10574 wk: k,
10575 wv: v,
10576 wo: o,
10577 q_norm: q_norm.as_deref(),
10578 k_norm: k_norm.as_deref(),
10579 late_qk_norm: self.qk_norm_after_rope,
10580 bias: bias
10581 .as_ref()
10582 .map(|(q, k, v)| (q.as_slice(), k.as_slice(), v.as_slice())),
10583 output_gate: false,
10584 cpu_k: self.kv_cache.layers[li].k_heads(),
10585 cpu_v: self.kv_cache.layers[li].v_heads(),
10586 cpu_base: self.kv_cache.layers[li].base(),
10587 geom: self.graph_attn_geom(li),
10588 head_gate: None,
10589 },
10590 };
10591 let (nkv, hd, rd) = self.layer_geom(li);
10592 crate::gpu::forward_batch_graph_at(
10593 &model,
10594 self.graph_kv_id,
10595 li,
10596 &[layer],
10597 &self.inv_freq,
10598 h,
10599 self.layer_num_heads(li),
10600 nkv,
10601 hd,
10602 rd,
10603 self.hidden_size,
10604 1,
10605 positions,
10606 self.kv_cache.max_seq_len,
10607 self.norm_style == cortiq_core::NormStyle::Gemma,
10608 self.rms_eps as f32,
10609 self.attn_scale,
10610 b,
10611 &[],
10612 self.o1_epoch,
10613 None,
10614 None,
10615 )
10616 };
10617 match outcome {
10618 Out::Completed => {}
10619 Out::Declined => return Out::Declined,
10620 Out::Failed => {
10621 self.graph_failed
10622 .store(true, std::sync::atomic::Ordering::Relaxed);
10623 return Out::Failed;
10624 }
10625 }
10626 let attn_ns = attn_started.elapsed().as_nanos() as u64;
10627 let hs = self.hidden_size;
10628 let lw = &self.weights.layers[li];
10629 let FfnKind::Moe(m) = &lw.ffn else {
10630 unreachable!()
10631 };
10632 let mut post = vec![0.0; h.len()];
10633 for (x, y) in h.chunks_exact(hs).zip(post.chunks_exact_mut(hs)) {
10634 inference::rms_norm_into(x, &lw.post_norm, self.rms_eps, self.norm_style, y);
10635 }
10636 let mut ffn = if b == 1 {
10637 moe_ffn_banked(&mut self.mimo_moe, li, m, &post, self.pool.as_deref())
10638 } else {
10639 moe_ffn_banked_rows(
10640 &mut self.mimo_moe,
10641 li,
10642 m,
10643 &post,
10644 b,
10645 hs,
10646 self.pool.as_deref(),
10647 )
10648 };
10649 for (x, &f) in h.iter_mut().zip(&ffn) {
10650 *x += f;
10651 }
10652 attention::recycle_buf(&mut ffn);
10653 if self.layer_dump.is_some() {
10654 for (&pos, row) in positions.iter().zip(h.chunks_exact(hs)) {
10655 self.dump_layer_row(pos, li, row);
10656 }
10657 }
10658 crate::mimo_moe::note_attention_graph(b, attn_ns);
10659 Out::Completed
10660 }
10661
10662 fn layer_attn_plain(&self, li: usize) -> bool {
10663 self.kv_heads_per_layer.is_none()
10664 && self.v_head_dim.is_none()
10665 && self.global_attn.is_none()
10666 && self.layer_window(li).is_none()
10667 && self.kv_cache.layers[li].sinks.is_none()
10668 }
10669
10670 fn forward_layers(
10672 &mut self,
10673 hidden: &[f32],
10674 position: usize,
10675 task_mask: Option<&TaskMask>,
10676 ) -> Vec<f32> {
10677 let out = self.forward_layers_upto(hidden, position, task_mask, None);
10678 self.o1_progress();
10679 self.swa_trim_tails();
10680 out
10681 }
10682
10683 pub fn embed_id(&self, id: u32) -> Vec<f32> {
10691 self.embed_single(id)
10692 }
10693
10694 pub fn split_supported(&self) -> Result<(), String> {
10698 if self.dsv4.is_some() {
10699 return Err(
10700 "network split: DeepSeek-V4 runs its own fused stack (not splittable yet)".into(),
10701 );
10702 }
10703 if self.dsv41.is_some() {
10704 return Err(
10705 "network split: DeepSeek-V4.1 owns the shared CED/CSA2 state (not splittable)"
10706 .into(),
10707 );
10708 }
10709 if self.qwen4_exp.is_some() {
10710 return Err(
10711 "network split: Qwen3.8-Flash-Next hyper/QSA stack is not splittable yet".into(),
10712 );
10713 }
10714 if self.g3n.is_some() {
10715 return Err(
10716 "network split: Gemma-3n runs its own AltUp stack (not splittable yet)".into(),
10717 );
10718 }
10719 Ok(())
10720 }
10721
10722 pub fn forward_span(
10727 &mut self,
10728 hidden: &[f32],
10729 position: usize,
10730 from: usize,
10731 upto: usize,
10732 task_mask: Option<&TaskMask>,
10733 ) -> Result<Vec<f32>, String> {
10734 self.split_supported()?;
10735 if from > upto || upto >= self.num_layers {
10736 return Err(format!(
10737 "forward_span: layer range {from}..={upto} outside 0..{}",
10738 self.num_layers
10739 ));
10740 }
10741 if hidden.len() != self.hidden_size {
10742 return Err(format!(
10743 "forward_span: hidden len {} ≠ hidden_size {}",
10744 hidden.len(),
10745 self.hidden_size
10746 ));
10747 }
10748 let out = self.forward_layers_span(hidden, position, task_mask, from, Some(upto));
10749 self.o1_progress();
10750 self.swa_trim_tails();
10751 if self
10752 .graph_failed
10753 .swap(false, std::sync::atomic::Ordering::Relaxed)
10754 {
10755 self.cancel
10756 .store(false, std::sync::atomic::Ordering::Relaxed);
10757 self.clear_sequence_state();
10758 return Err("forward_span: deferred O(1) transition failed".into());
10759 }
10760 Ok(out)
10761 }
10762
10763 pub fn logits_from_hidden(&mut self, hidden: &[f32]) -> Vec<f32> {
10766 let normed = inference::rms_norm(
10767 hidden,
10768 &self.weights.final_norm,
10769 self.rms_eps,
10770 self.norm_style,
10771 );
10772 self.lm_head_forward(&normed)
10773 }
10774
10775 pub fn sample_next(&mut self, logits: &[f32], past_tokens: &[u32]) -> u32 {
10777 sampler::sample_with_scratch(
10778 logits,
10779 &self.sampler_config,
10780 past_tokens,
10781 &mut self.rng,
10782 &mut self.sampler_scratch,
10783 )
10784 }
10785
10786 pub fn reset_session(&mut self) {
10788 self.clear_sequence_state();
10789 }
10790
10791 pub fn prefill_span_ids(
10797 &mut self,
10798 ids: &[u32],
10799 start_pos: usize,
10800 upto: usize,
10801 task_mask: Option<&TaskMask>,
10802 ) -> Result<Vec<f32>, String> {
10803 self.split_supported()?;
10804 if upto >= self.num_layers {
10805 return Err(format!(
10806 "prefill_span_ids: upto {upto} outside 0..{}",
10807 self.num_layers
10808 ));
10809 }
10810 if self.can_prefill_batched() && !self.graph_prefill_preferred() {
10814 let out =
10815 self.prefill_batch_span(PrefillIn::Ids(ids), start_pos, task_mask, 0, upto + 1);
10816 self.check_o1_progress_failure("prefill_span_ids")?;
10817 Ok(out)
10818 } else {
10819 let hs = self.hidden_size;
10820 let mut out = Vec::with_capacity(ids.len() * hs);
10821 for (i, &id) in ids.iter().enumerate() {
10822 let emb = self.embed_id(id);
10823 out.extend_from_slice(&self.forward_span(
10824 &emb,
10825 start_pos + i,
10826 0,
10827 upto,
10828 task_mask,
10829 )?);
10830 }
10831 Ok(out)
10832 }
10833 }
10834
10835 pub fn prefill_span_hidden(
10838 &mut self,
10839 hidden: &[f32],
10840 start_pos: usize,
10841 from: usize,
10842 upto: usize,
10843 task_mask: Option<&TaskMask>,
10844 ) -> Result<Vec<f32>, String> {
10845 self.split_supported()?;
10846 let hs = self.hidden_size;
10847 if hidden.is_empty() || hidden.len() % hs != 0 {
10848 return Err(format!(
10849 "prefill_span_hidden: {} floats is not a multiple of hidden {hs}",
10850 hidden.len()
10851 ));
10852 }
10853 if from > upto || upto >= self.num_layers {
10854 return Err(format!(
10855 "prefill_span_hidden: layer range {from}..={upto} outside 0..{}",
10856 self.num_layers
10857 ));
10858 }
10859 if self.can_prefill_batched() && !self.graph_prefill_preferred() {
10860 let out = self.prefill_batch_span(
10861 PrefillIn::Hidden(hidden),
10862 start_pos,
10863 task_mask,
10864 from,
10865 upto + 1,
10866 );
10867 self.check_o1_progress_failure("prefill_span_hidden")?;
10868 Ok(out)
10869 } else {
10870 let b = hidden.len() / hs;
10871 let mut out = Vec::with_capacity(hidden.len());
10872 for i in 0..b {
10873 let h = self.forward_span(
10874 &hidden[i * hs..(i + 1) * hs],
10875 start_pos + i,
10876 from,
10877 upto,
10878 task_mask,
10879 )?;
10880 out.extend_from_slice(&h);
10881 }
10882 Ok(out)
10883 }
10884 }
10885
10886 fn try_token_graph_wgpu(
10890 &self,
10891 hidden: &[f32],
10892 position: usize,
10893 logits_out: &mut Vec<f32>,
10894 layers_run: &mut usize,
10895 ) -> Option<Result<Vec<f32>, ()>> {
10896 self.try_token_graph_wgpu_steps(
10897 hidden,
10898 position,
10899 logits_out,
10900 1,
10901 None,
10902 Some(layers_run),
10903 0,
10904 self.num_layers,
10905 )
10906 }
10907
10908 fn try_token_graph_wgpu_span(
10912 &self,
10913 hidden: &[f32],
10914 position: usize,
10915 logits_out: &mut Vec<f32>,
10916 from: usize,
10917 upto_excl: usize,
10918 layers_run: &mut usize,
10919 ) -> Option<Result<Vec<f32>, ()>> {
10920 self.try_token_graph_wgpu_steps(
10921 hidden,
10922 position,
10923 logits_out,
10924 1,
10925 None,
10926 Some(layers_run),
10927 from,
10928 upto_excl,
10929 )
10930 }
10931
10932 fn try_multi_burst(&self, t_next: u32, position: usize, k: usize) -> Option<Vec<u32>> {
10936 if self.o1_active() || self.attn_softcap > 0.0 {
10937 return None;
10938 }
10939 if let Some(reason) = self.wgpu_graph_attn_decline() {
10942 self.note_graph_decline("wgpu multi-burst", reason);
10943 return None;
10944 }
10945 let graph_on = crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Decode);
10946 if !graph_on || self.graph_refused() {
10947 return None;
10954 }
10955 let emb = self.embed_single(t_next);
10956 let mut lg = Vec::new();
10957 let mut ids = Vec::new();
10958 match self.try_token_graph_wgpu_steps(
10959 &emb,
10960 position,
10961 &mut lg,
10962 k,
10963 Some(&mut ids),
10964 None,
10965 0,
10966 self.num_layers,
10967 ) {
10968 Some(Ok(_)) => {}
10969 Some(Err(())) => {
10970 self.graph_failed
10975 .store(true, std::sync::atomic::Ordering::Relaxed);
10976 return None;
10977 }
10978 None => return None,
10979 }
10980 (ids.len() == k).then_some(ids)
10981 }
10982
10983 fn try_token_graph_wgpu_steps(
10987 &self,
10988 hidden: &[f32],
10989 position: usize,
10990 logits_out: &mut Vec<f32>,
10991 steps: usize,
10992 ids_out: Option<&mut Vec<u32>>,
10993 layers_run: Option<&mut usize>,
10994 from: usize,
10995 upto_excl: usize,
10996 ) -> Option<Result<Vec<f32>, ()>> {
10997 let upto_excl = match self.mimo_moe.graph_prefix_end() {
11000 Some(end) if end < upto_excl => {
11001 if steps != 1 || layers_run.is_none() || from >= end {
11002 return None;
11003 }
11004 end
11005 }
11006 _ => upto_excl,
11007 };
11008 let o1_gpu = std::env::var("CMF_O1_GPU").as_deref() == Ok("1");
11011 if (self.o1_active() && !o1_gpu) || self.attn_softcap > 0.0 {
11012 return None;
11016 }
11017 if let Some(reason) = self.wgpu_graph_attn_decline() {
11024 self.note_graph_decline("wgpu token graph", reason);
11025 return None;
11026 }
11027 let o1_views: Vec<Option<Vec<crate::nystrom::O1DeviceView<'_>>>> = (from..upto_excl)
11032 .map(|li| {
11033 if !o1_gpu {
11034 return None;
11035 }
11036 self.kv_cache.layers[self.phys_layer(li)].o1_views()
11037 })
11038 .collect();
11039 if self.o1_active() && o1_gpu {
11040 let want: usize = (from..upto_excl)
11043 .filter(|li| self.kv_cache.layers[self.phys_layer(*li)].o1.is_some())
11044 .count();
11045 let have = o1_views.iter().filter(|v| v.is_some()).count();
11046 if want == 0 || have != want {
11047 use std::sync::atomic::{AtomicUsize, Ordering};
11057 static LAST: AtomicUsize = AtomicUsize::new(usize::MAX);
11058 let code = have * 1000 + want;
11059 if LAST.swap(code, Ordering::Relaxed) != code {
11060 tracing::warn!(
11061 "o1 graph: {have} of {want} layers sealed — per-op until all seal"
11062 );
11063 }
11064 return None;
11065 }
11066 }
11067 let nh = self.num_heads;
11068 let (nkv, hd, rd) = self.layer_geom(0);
11069 let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
11070 let mut layers = Vec::with_capacity(upto_excl - from);
11071 let mut model = None;
11072 let dbg = std::env::var("CMF_GRAPH_DEBUG").is_ok();
11073 fn gw(t: &QTensor) -> Option<crate::gpu::GraphW<'_>> {
11074 if let Some((m, i, kind, rs)) = t
11075 .graph_weight()
11076 .or_else(|| t.graph_weight_descriptor())
11077 {
11078 let name = &m.tensors[i].name;
11079 let prism = if crate::prism::is_inverse_embedding(m, name) {
11080 crate::gpu::GraphPrismOp::InverseEmbedding
11081 } else if crate::prism::is_forward_weight(m, name) {
11082 crate::gpu::GraphPrismOp::Forward
11083 } else {
11084 crate::gpu::GraphPrismOp::None
11085 };
11086 return Some(crate::gpu::GraphW {
11087 idx: i,
11088 kind,
11089 row_scale: rs,
11090 data: &[],
11091 prism,
11092 affine: crate::prism::is_affine_target(m, name),
11093 });
11094 }
11095 match t.as_f32() {
11097 Some(d) => Some(crate::gpu::GraphW {
11098 idx: 0,
11099 kind: 4,
11100 row_scale: &[],
11101 data: d,
11102 prism: crate::gpu::GraphPrismOp::None,
11103 affine: false,
11104 }),
11105 None => {
11106 if std::env::var_os("CMF_BATCH_DEBUG").is_some() {
11107 eprintln!("batch graph: weight has no graph/f32 representation");
11108 }
11109 None
11110 }
11111 }
11112 }
11113 for li in from..upto_excl {
11114 let lw = &self.weights.layers[self.phys_layer(li)];
11115 if dbg {
11116 let ak = match &lw.attn {
11117 AttnKind::Mla(_) => "Mla".into(),
11118 AttnKind::Full {
11119 output_gate, bias, ..
11120 } => format!("Full gate={output_gate} bias={}", bias.is_some()),
11121 AttnKind::LinearGdn(_) => "LinearGdn".into(),
11122 AttnKind::Kda(_) => "Kda".into(),
11123 AttnKind::Linear(_) => "Linear".into(),
11124 AttnKind::ShortConv(_) => "ShortConv".into(),
11125 AttnKind::Bounded(_) => "Bounded".into(),
11126 };
11127 let fk = match &lw.ffn {
11128 FfnKind::Dense(_) => "Dense",
11129 FfnKind::Moe(_) => "Moe",
11130 FfnKind::DenseMoe(_) => "DenseMoe",
11131 };
11132 eprintln!("graph L{li}: attn={ak} ffn={fk}");
11133 }
11134 let gffn = match &lw.ffn {
11135 FfnKind::DenseMoe(_) => return None, FfnKind::Dense(d) if !d.segs.is_empty() => return None,
11139 FfnKind::Dense(d) => {
11140 let Some(act) = d.act.graph_act() else {
11144 self.note_graph_decline(
11145 "wgpu token graph",
11146 "dense FFN activation without a graph kernel",
11147 );
11148 return None;
11149 };
11150 crate::gpu::GraphFfn::Dense {
11151 gate: gw(&d.gate_proj)?,
11152 up: gw(&d.up_proj)?,
11153 down: gw(&d.down_proj)?,
11154 act,
11155 }
11156 }
11157 FfnKind::Moe(m) => {
11158 if m.route_tau.is_some() || m.mask.is_some() {
11166 return None;
11167 }
11168 let shared = m.shared.as_ref();
11169 let has_shared = shared.is_some();
11170 let shared_gated = matches!(shared, Some((_, Some(_))));
11171 let sgate = match shared {
11172 Some((_, Some(sg))) => gw(sg)?,
11173 _ => gw(&m.router)?,
11177 };
11178 let router = gw(&m.router)?;
11179 if router.prism != crate::gpu::GraphPrismOp::None
11185 || sgate.prism != crate::gpu::GraphPrismOp::None
11186 || router.affine
11187 || sgate.affine
11188 {
11189 tracing::warn!(
11190 "resident MoE declined: Prism/affine router or shared gate transform is not implemented"
11191 );
11192 return None;
11193 }
11194 let inter = m.experts.first()?.gate_proj.rows();
11195 let mut experts = Vec::with_capacity(m.experts.len() + 1);
11196 let mut q4tp: Option<bool> = None;
11199 let mut gu_q2: Option<bool> = None;
11202 for e in m.experts.iter().chain(shared.map(|(se, _)| se)) {
11203 if !matches!(e.act, Act::Silu)
11204 || e.gate_proj.rows() != inter
11205 || e.up_proj.rows() != inter
11206 {
11207 return None;
11208 }
11209 for expert_weight in [&e.gate_proj, &e.up_proj, &e.down_proj] {
11214 let Some((em, ei, _, _)) = expert_weight
11215 .graph_weight()
11216 .or_else(|| expert_weight.graph_weight_descriptor())
11217 else {
11218 return None;
11219 };
11220 let name = &em.tensors[ei].name;
11221 if crate::prism::is_forward_weight(em, name)
11222 || crate::prism::is_inverse_embedding(em, name)
11223 || crate::prism::is_affine_target(em, name)
11224 {
11225 tracing::warn!(
11226 "resident MoE declined: expert Prism/affine transform is not implemented"
11227 );
11228 return None;
11229 }
11230 }
11231 let (mm, gi, ui, di, is_p, is_q2) = match e.gate_proj.mapped_q4t() {
11232 Some((mm, gi)) => (
11233 mm,
11234 gi,
11235 e.up_proj.mapped_q4t()?.1,
11236 e.down_proj.mapped_q4t()?.1,
11237 false,
11238 false,
11239 ),
11240 None => match e.gate_proj.mapped_q2tp() {
11241 Some((mm, gi)) => (
11242 mm,
11243 gi,
11244 e.up_proj.mapped_q2tp()?.1,
11245 e.down_proj.mapped_q4tp()?.1,
11246 true,
11247 true,
11248 ),
11249 None => {
11250 let (mm, gi) = e.gate_proj.mapped_q4tp()?;
11251 (
11252 mm,
11253 gi,
11254 e.up_proj.mapped_q4tp()?.1,
11255 e.down_proj.mapped_q4tp()?.1,
11256 true,
11257 false,
11258 )
11259 }
11260 },
11261 };
11262 if *q4tp.get_or_insert(is_p) != is_p || *gu_q2.get_or_insert(is_q2) != is_q2
11263 {
11264 tracing::warn!(
11270 "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."
11271 );
11272 return None;
11273 }
11274 model.get_or_insert_with(|| mm.clone());
11275 experts.push((gi, ui, di));
11276 }
11277 crate::gpu::GraphFfn::Moe {
11278 router,
11279 shared_gate: sgate,
11280 experts,
11281 n_exp: m.experts.len(),
11282 top_k: std::env::var("CMF_TOPK_PROBE")
11288 .ok()
11289 .and_then(|v| v.parse::<usize>().ok())
11290 .filter(|k| *k > 0 && *k <= m.top_k)
11291 .unwrap_or(m.top_k),
11292 inter,
11293 norm_topk: m.norm_topk_prob,
11294 q4tp: q4tp?,
11295 gu_q2: gu_q2.unwrap_or(false),
11296 sigmoid: m.router_sigmoid,
11297 bias: m.expert_bias.as_deref(),
11298 has_shared,
11299 shared_gated,
11300 route_scale: m.routed_scaling,
11301 }
11302 }
11303 };
11304 let attn = match &lw.attn {
11305 AttnKind::Full {
11306 wq,
11307 wk,
11308 wv,
11309 wo,
11310 q_norm,
11311 k_norm,
11312 output_gate,
11313 softplus_gate,
11314 bias,
11315 } => {
11316 if self.attention_heads_per_layer.is_some() {
11317 return None;
11318 }
11319 let head_gate = match softplus_gate {
11323 None => None,
11324 Some((g, true)) if self.proj_gate_sigmoid && !*output_gate => {
11325 Some(gw(g)?)
11326 }
11327 Some(_) => {
11328 self.note_graph_decline(
11329 "wgpu token graph",
11330 "projected softplus / per-element output gate",
11331 );
11332 return None;
11333 }
11334 };
11335 let (m, _, _, _) = wq
11336 .graph_weight()
11337 .or_else(|| wq.graph_weight_descriptor())?;
11338 model = Some(m.clone());
11339 crate::gpu::GraphAttn::Full {
11340 wq: gw(wq)?,
11341 wk: gw(wk)?,
11342 wv: gw(wv)?,
11343 wo: gw(wo)?,
11344 q_norm: q_norm.as_deref(),
11345 k_norm: k_norm.as_deref(),
11346 late_qk_norm: self.qk_norm_after_rope,
11347 bias: bias
11348 .as_ref()
11349 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
11350 output_gate: *output_gate,
11351 cpu_k: self.kv_cache.layers[li].k_heads(),
11352 cpu_v: self.kv_cache.layers[li].v_heads(),
11353 cpu_base: self.kv_cache.layers[li].base(),
11354 geom: self.graph_attn_geom(li),
11355 head_gate,
11356 }
11357 }
11358 AttnKind::LinearGdn(w) => {
11359 let cfg = self.gdn_cfg?;
11360 let (m, _, _, _) = w
11361 .in_proj_qkv
11362 .graph_weight()
11363 .or_else(|| w.in_proj_qkv.graph_weight_descriptor())?;
11364 model = Some(m.clone());
11365 crate::gpu::GraphAttn::Gdn {
11366 qkv: gw(&w.in_proj_qkv)?,
11367 z: gw(&w.in_proj_z)?,
11368 a: gw(&w.in_proj_a)?,
11369 b: gw(&w.in_proj_b)?,
11370 out: gw(&w.out_proj)?,
11371 conv1d: &w.conv1d,
11372 a_log: &w.a_log,
11373 dt_bias: &w.dt_bias,
11374 norm: &w.norm,
11375 nv: cfg.num_v_heads,
11376 nk: cfg.num_k_heads,
11377 dk: cfg.key_head_dim,
11378 dv: cfg.value_head_dim,
11379 kk: cfg.conv_kernel,
11380 cpu_state: &self.kv_cache.layers[self.phys_layer(li)].linear_state,
11381 }
11382 }
11383 AttnKind::ShortConv(w) => {
11384 let cfg = self.short_conv_cfg?;
11385 let (m, _, _, _) = w
11386 .in_proj
11387 .graph_weight()
11388 .or_else(|| w.in_proj.graph_weight_descriptor())?;
11389 model = Some(m.clone());
11390 crate::gpu::GraphAttn::ShortConv {
11391 inp: gw(&w.in_proj)?,
11392 out: gw(&w.out_proj)?,
11393 taps: &w.conv,
11394 kernel: cfg.kernel,
11395 cpu_state: &self.kv_cache.layers[self.phys_layer(li)].linear_state,
11396 }
11397 }
11398 _ => return None,
11399 };
11400 layers.push(crate::gpu::GraphLayer {
11401 input_norm: &lw.input_norm,
11402 attn,
11403 post_norm: &lw.post_norm,
11404 ffn: gffn,
11405 });
11406 }
11407 let model = model?;
11408 let lm_gw = if upto_excl == self.num_layers
11414 && self.graph_want_logits
11415 && std::env::var("CMF_GPU_LMHEAD")
11416 .map(|v| v != "0")
11417 .unwrap_or(true)
11418 {
11419 self.weights
11420 .lm_head
11421 .graph_weight()
11422 .or_else(|| self.weights.lm_head.graph_weight_descriptor())
11423 .map(|(m, i, kind, rs)| {
11424 let name = &m.tensors[i].name;
11425 let prism = if crate::prism::is_inverse_embedding(m, name) {
11426 crate::gpu::GraphPrismOp::InverseEmbedding
11427 } else if crate::prism::is_forward_weight(m, name) {
11428 crate::gpu::GraphPrismOp::Forward
11429 } else {
11430 crate::gpu::GraphPrismOp::None
11431 };
11432 (
11433 crate::gpu::GraphW {
11434 idx: i,
11435 kind,
11436 row_scale: rs,
11437 data: &[],
11438 prism,
11439 affine: crate::prism::is_affine_target(m, name),
11440 },
11441 self.weights.lm_head.rows(),
11442 )
11443 })
11444 } else {
11445 None
11446 };
11447 let lm = lm_gw.as_ref().map(|(gw, rows)| (gw, *rows));
11448 let emb_gw = if steps > 1 {
11450 self.weights
11451 .embed_tokens
11452 .graph_weight()
11453 .or_else(|| self.weights.embed_tokens.graph_weight_descriptor())
11454 .map(|(m, i, kind, rs)| {
11455 let name = &m.tensors[i].name;
11456 let prism = if crate::prism::is_inverse_embedding(m, name) {
11457 crate::gpu::GraphPrismOp::InverseEmbedding
11458 } else if crate::prism::is_forward_weight(m, name) {
11459 crate::gpu::GraphPrismOp::Forward
11460 } else {
11461 crate::gpu::GraphPrismOp::None
11462 };
11463 (
11464 crate::gpu::GraphW {
11465 idx: i,
11466 kind,
11467 row_scale: rs,
11468 data: &[],
11469 prism,
11470 affine: crate::prism::is_affine_target(m, name),
11471 },
11472 self.weights.embed_tokens.rows(),
11473 self.embed_multiplier,
11474 )
11475 })
11476 } else {
11477 None
11478 };
11479
11480 let loop_norm_at: Vec<usize> = if self.loop_final_norm {
11486 (from..upto_excl.min(self.num_layers - 1))
11487 .filter(|&li| (li + 1) % self.physical_layers == 0)
11488 .map(|li| li - from)
11489 .collect()
11490 } else {
11491 Vec::new()
11492 };
11493 let mut h = hidden.to_vec();
11494 let dump_hidden = std::env::var_os("CMF_LOGIT_DUMP").is_some();
11500 let outcome = crate::gpu::forward_token_graph(
11501 &model,
11502 self.graph_kv_id,
11503 &layers,
11504 &o1_views,
11505 self.o1_epoch,
11506 &self.inv_freq,
11507 &mut h,
11508 nh,
11509 nkv,
11510 hd,
11511 self.attn_scale,
11512 rd,
11513 self.hidden_size,
11514 self.intermediate_size,
11515 position,
11516 self.kv_cache.max_seq_len,
11517 gemma,
11518 self.rms_eps as f32,
11519 lm,
11520 &self.weights.final_norm,
11521 logits_out,
11522 &loop_norm_at,
11523 steps,
11524 emb_gw.as_ref().map(|(gw, rows, m)| (gw, *rows, *m)),
11525 ids_out,
11526 layers_run,
11527 from,
11528 dump_hidden,
11529 );
11530 match outcome {
11531 crate::gpu::TokenGraphOutcome::Completed => Some(Ok(h)),
11532 crate::gpu::TokenGraphOutcome::Failed => Some(Err(())),
11533 crate::gpu::TokenGraphOutcome::Declined => None,
11534 }
11535 }
11536
11537 #[cfg(target_os = "macos")]
11546 #[allow(clippy::type_complexity)]
11547 fn metal_rows_plan(
11548 &self,
11549 ) -> Option<(
11550 Vec<MetalRowsItem<'_>>,
11551 std::sync::Arc<cortiq_core::CmfModel>,
11552 Option<crate::gpu_metal::GdnGpuCfg>,
11553 )> {
11554 use crate::gpu_metal::{AttnGpuLayer, GdnGpuCfg, GdnGpuLayer, MetalFfn};
11555 let graph_force = crate::gpu::q1_force() || crate::gpu::q2tp_gpu_opt_in();
11556 if !graph_force
11557 || !crate::gpu::enabled_here()
11558 || std::env::var("CMF_GPU_BLOCK")
11559 .map(|v| v == "0")
11560 .unwrap_or(false)
11561 || self.attn_softcap > 0.0
11562 || self.o1_active()
11563 || self.swa.is_some()
11564 || self.global_attn.is_some()
11565 || self.attention_heads_per_layer.is_some()
11566 || self.graph_attn_decline_reason().is_some()
11568 || self.attn_v_norm
11569 || self.loop_final_norm
11570 {
11571 return None;
11572 }
11573 let attend_contract = self.head_dim % 4 == 0
11574 && self.head_dim <= 256
11575 && self.rotary_dim >= 2
11576 && self.rotary_dim <= self.head_dim
11577 && (self.rotary_dim / 2) % 32 == 0
11578 && self.num_kv_heads > 0
11579 && self.num_heads % self.num_kv_heads == 0;
11580 if !attend_contract {
11581 return None;
11582 }
11583 let mut plan: Vec<MetalRowsItem> = Vec::new();
11584 let mut model_ref: Option<std::sync::Arc<cortiq_core::CmfModel>> = None;
11585 for li in 0..self.num_layers {
11586 let lw = &self.weights.layers[self.phys_layer(li)];
11587 if lw.attn_out_norm.is_some() || lw.ffn_out_norm.is_some() || lw.layer_scale.is_some() {
11588 return None;
11589 }
11590 let ffn = match &lw.ffn {
11591 FfnKind::Dense(d) if d.act == Act::Silu && d.segs.is_empty() => {
11592 let (Some(g), Some(u), Some(dn)) = (
11593 d.gate_proj.metal_graph_parts(),
11594 d.up_proj.metal_graph_parts(),
11595 d.down_proj.metal_graph_parts(),
11596 ) else {
11597 return None;
11598 };
11599 MetalFfn::Dense {
11600 gate: g,
11601 up: u,
11602 down: dn,
11603 gelu: false, }
11605 }
11606 _ => return None,
11607 };
11608 match &lw.attn {
11609 AttnKind::LinearGdn(w) if self.gdn_cfg.is_some() => {
11610 let (Some(qkv), Some(z), Some(a), Some(bb), Some(out)) = (
11611 w.in_proj_qkv.metal_graph_parts(),
11612 w.in_proj_z.metal_graph_parts(),
11613 w.in_proj_a.f32_parts(),
11614 w.in_proj_b.f32_parts(),
11615 w.out_proj.metal_graph_parts(),
11616 ) else {
11617 return None;
11618 };
11619 if let QTensor::Mapped { model, .. } = &w.in_proj_qkv {
11620 model_ref.get_or_insert_with(|| model.clone());
11621 }
11622 let gl = GdnGpuLayer {
11623 attn_norm: &lw.input_norm,
11624 post_norm: &lw.post_norm,
11625 qkv,
11626 z,
11627 a,
11628 b: bb,
11629 out,
11630 ffn,
11631 conv1d: &w.conv1d,
11632 a_log: &w.a_log,
11633 dt_bias: &w.dt_bias,
11634 gnorm: &w.norm,
11635 };
11636 match plan.last_mut() {
11637 Some(MetalRowsItem::Gdn { run, .. }) => run.push(gl),
11638 _ => plan.push(MetalRowsItem::Gdn {
11639 run: vec![gl],
11640 first: li,
11641 }),
11642 }
11643 }
11644 AttnKind::Full {
11645 wq,
11646 wk,
11647 wv,
11648 wo,
11649 q_norm,
11650 k_norm,
11651 output_gate,
11652 softplus_gate: None,
11653 bias: None,
11654 } => {
11655 let (Some(pq), Some(pk), Some(pv), Some(po)) =
11656 (
11657 wq.metal_graph_parts(),
11658 wk.metal_graph_parts(),
11659 wv.metal_graph_parts(),
11660 wo.metal_graph_parts(),
11661 )
11662 else {
11663 return None;
11664 };
11665 if let QTensor::Mapped { model, .. } = wq {
11666 model_ref.get_or_insert_with(|| model.clone());
11667 }
11668 let cache = &self.kv_cache.layers[li];
11669 if cache.mode != crate::kv_cache::KvMode::F32 || cache.o1.is_some() {
11670 return None;
11671 }
11672 plan.push(MetalRowsItem::Attn {
11673 l: AttnGpuLayer {
11674 attn_norm: &lw.input_norm,
11675 post_norm: &lw.post_norm,
11676 wq: pq,
11677 wk: pk,
11678 wv: pv,
11679 wo: po,
11680 ffn,
11681 },
11682 li,
11683 q_norm: q_norm.as_deref(),
11684 k_norm: k_norm.as_deref(),
11685 output_gate: *output_gate,
11686 });
11687 }
11688 _ => return None,
11689 }
11690 }
11691 let model = model_ref?;
11692 let gcfg = self.gdn_cfg.map(|cfg| GdnGpuCfg {
11693 nv: cfg.num_v_heads,
11694 nk: cfg.num_k_heads,
11695 dk: cfg.key_head_dim,
11696 dv: cfg.value_head_dim,
11697 kk: cfg.conv_kernel,
11698 hidden: self.hidden_size,
11699 inter: self.intermediate_size,
11700 c_dim: cfg.conv_dim(),
11701 eps: cfg.rms_eps as f32,
11702 gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
11703 });
11704 Some((plan, model, gcfg))
11705 }
11706
11707 #[cfg(target_os = "macos")]
11709 #[allow(clippy::too_many_arguments)]
11710 fn metal_attn_params<'a>(
11711 li: usize,
11712 cache: &'a crate::kv_cache::LayerKvCache,
11713 q_norm: Option<&'a [f32]>,
11714 k_norm: Option<&'a [f32]>,
11715 output_gate: bool,
11716 inv_freq: &'a [f32],
11717 geom: (usize, usize, usize, usize),
11718 pos0: usize,
11719 kv_id: u64,
11720 scale: f32,
11721 eps: f32,
11722 gemma: bool,
11723 late_qk_norm: bool,
11724 ) -> (crate::gpu_metal::AttnDeviceParams<'a>, usize) {
11725 let (nh, nkv, hd, rd) = geom;
11726 let cpu_k: Vec<&[f32]> = (0..nkv).map(|g| cache.head_keys(g)).collect();
11727 let cpu_v: Vec<&[f32]> = (0..nkv).map(|g| cache.head_values(g)).collect();
11728 let cpu_stored = cpu_k[0].len() / hd;
11729 (
11730 crate::gpu_metal::AttnDeviceParams {
11731 kv_id,
11732 layer: li,
11733 nh,
11734 nkv,
11735 hd,
11736 rd,
11737 position: pos0,
11738 scale,
11739 eps,
11740 gemma,
11741 late_qk_norm,
11742 output_gate,
11743 q_norm,
11744 k_norm,
11745 inv_freq,
11746 cpu_k,
11747 cpu_v,
11748 cpu_stored,
11749 cpu_gen: cache.generation(),
11750 o1: None,
11751 window: None,
11752 head_gate: None,
11753 },
11754 cpu_stored,
11755 )
11756 }
11757
11758 #[cfg(target_os = "macos")]
11763 #[allow(clippy::type_complexity)]
11764 fn metal_rows_run(
11765 &mut self,
11766 hiddens: &mut [f32],
11767 pos0: usize,
11768 b: usize,
11769 prefill: bool,
11770 spec: Option<((usize, usize, usize), &[f32], &mut Vec<f32>)>,
11771 mut argmax_out: Option<(usize, &mut Vec<u32>)>,
11775 ) -> MetalRowsRun {
11776 use crate::gpu_metal::{GraphDims, VerifyGraph};
11777 if !crate::gpu_metal::wait_replay() {
11783 tracing::error!("Metal rows graph: the pending async replay failed");
11784 return MetalRowsRun::Failed;
11785 }
11786 spec_stamp("v.wait");
11787 let want = self.gdn_cfg.map(|c| c.state_len()).unwrap_or(0);
11793 if want > 0 {
11794 let phys = self.physical_layers.max(1);
11795 for (li, l) in self.kv_cache.layers.iter_mut().enumerate() {
11796 let is_gdn = self
11797 .weights
11798 .layers
11799 .get(li % phys)
11800 .is_some_and(|lw| matches!(lw.attn, AttnKind::LinearGdn(_)));
11801 if is_gdn && l.linear_state.len() != want {
11802 l.linear_state = vec![0f32; want];
11803 }
11804 }
11805 }
11806 let Some((plan, model, gcfg)) = self.metal_rows_plan() else {
11807 return MetalRowsRun::Declined;
11808 };
11809 spec_stamp("v.plan");
11810 let dims = GraphDims {
11811 hidden: self.hidden_size,
11812 eps: self.rms_eps as f32,
11813 gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
11814 };
11815 let Some(mut graph) = (if prefill {
11816 VerifyGraph::new_prefill(&model, dims, hiddens, b)
11817 } else {
11818 VerifyGraph::new(&model, dims, hiddens, b)
11819 }) else {
11820 return MetalRowsRun::Declined;
11821 };
11822 let geom = (
11823 self.num_heads,
11824 self.num_kv_heads,
11825 self.head_dim,
11826 self.rotary_dim,
11827 );
11828 let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
11829 let eps = self.rms_eps as f32;
11830 let kv_id = self.graph_kv_id;
11831 let inv_freq = self.inv_freq.clone();
11832 for item in &plan {
11833 let ok = match item {
11834 MetalRowsItem::Gdn { run, .. } => gcfg
11835 .as_ref()
11836 .map(|gc| run.iter().all(|l| graph.gdn_ok(l, gc)))
11837 .unwrap_or(false),
11838 MetalRowsItem::Attn {
11839 l,
11840 li,
11841 q_norm,
11842 k_norm,
11843 output_gate,
11844 } => {
11845 let (p, _) = Self::metal_attn_params(
11846 *li,
11847 &self.kv_cache.layers[*li],
11848 *q_norm,
11849 *k_norm,
11850 *output_gate,
11851 &inv_freq,
11852 geom,
11853 pos0,
11854 kv_id,
11855 self.attn_scale,
11856 eps,
11857 gemma,
11858 self.qk_norm_after_rope,
11859 );
11860 graph.attn_ok(l, &p)
11861 }
11862 };
11863 if !ok {
11864 use std::sync::atomic::{AtomicBool, Ordering};
11865 static SAID: AtomicBool = AtomicBool::new(false);
11866 if !SAID.swap(true, Ordering::Relaxed) {
11867 tracing::warn!("metal rows graph: a layer failed preflight — declining");
11868 }
11869 return MetalRowsRun::Declined;
11870 }
11871 }
11872 let lm = match &spec {
11873 Some((lm, _, _)) => {
11874 if !graph.lm_head_ok(*lm) {
11875 return MetalRowsRun::Declined;
11876 }
11877 Some(*lm)
11878 }
11879 None => None,
11880 };
11881 let mut gdn_layers = Vec::new();
11882 let mut attn_layers = Vec::new();
11883 for item in &plan {
11884 match item {
11885 MetalRowsItem::Gdn { run, first } => {
11886 let ro: Vec<&[f32]> = self.kv_cache.layers[*first..*first + run.len()]
11887 .iter()
11888 .map(|l| l.linear_state.as_slice())
11889 .collect();
11890 if !graph.encode_gdn_run_b(run, &ro, gcfg.as_ref().unwrap()) {
11891 return MetalRowsRun::Declined;
11892 }
11893 gdn_layers.extend(*first..*first + run.len());
11894 }
11895 MetalRowsItem::Attn {
11896 l,
11897 li,
11898 q_norm,
11899 k_norm,
11900 output_gate,
11901 } => {
11902 let (p, cpu_stored) = Self::metal_attn_params(
11903 *li,
11904 &self.kv_cache.layers[*li],
11905 *q_norm,
11906 *k_norm,
11907 *output_gate,
11908 &inv_freq,
11909 geom,
11910 pos0,
11911 kv_id,
11912 self.attn_scale,
11913 eps,
11914 gemma,
11915 self.qk_norm_after_rope,
11916 );
11917 if !graph.encode_attn_b(l, &p) {
11918 return MetalRowsRun::Declined;
11919 }
11920 attn_layers.push((*li, cpu_stored));
11921 }
11922 }
11923 }
11924 if let (Some(lm), Some((_, final_norm, _))) = (lm, spec.as_ref()) {
11925 if !graph.encode_lm_head_b(final_norm, lm) {
11926 return MetalRowsRun::Declined;
11927 }
11928 if let Some((n, _)) = argmax_out.as_ref() {
11933 if !graph.encode_argmax_b(*n) {
11934 argmax_out = None;
11935 }
11936 }
11937 }
11938 spec_stamp("v.enc");
11939 if !graph.sync() {
11940 return MetalRowsRun::Failed;
11941 }
11942 spec_stamp("v.gpu");
11943 match (spec, argmax_out) {
11944 (Some(_), Some((_, ids))) => {
11945 ids.resize(b, 0);
11946 if !graph.read_argmax(ids) {
11947 return MetalRowsRun::Failed;
11948 }
11949 spec_stamp("v.am");
11950 }
11951 (Some((lm, _, logits)), None) => {
11952 logits.resize(b * lm.1, 0.0);
11953 if !graph.read_logits(logits) {
11954 return MetalRowsRun::Failed;
11955 }
11956 spec_stamp("v.lg");
11957 }
11958 (None, _) => {}
11959 }
11960 if !graph.read_hidden(hiddens) {
11961 return MetalRowsRun::Failed;
11962 }
11963 spec_stamp("v.hid");
11964 MetalRowsRun::Completed(MetalVerifyPending {
11965 graph,
11966 gdn_layers,
11967 attn_layers,
11968 })
11969 }
11970
11971 #[cfg(target_os = "macos")]
11977 fn try_batch_graph_metal(
11978 &mut self,
11979 hiddens: &mut [f32],
11980 positions: &[usize],
11981 b: usize,
11982 spec: Option<((usize, usize, usize), &[f32], &mut Vec<f32>)>,
11983 argmax_out: Option<(usize, &mut Vec<u32>)>,
11984 ) -> crate::gpu::BatchGraphOutcome {
11985 let _t0 = std::time::Instant::now();
11986 if positions.len() != b
11987 || positions.windows(2).any(|w| w[1] != w[0] + 1)
11988 || hiddens.len() != b * self.hidden_size
11989 {
11990 return crate::gpu::BatchGraphOutcome::Declined;
11991 }
11992 let pending = match self.metal_rows_run(hiddens, positions[0], b, false, spec, argmax_out) {
11993 MetalRowsRun::Declined => return crate::gpu::BatchGraphOutcome::Declined,
11994 MetalRowsRun::Failed => return crate::gpu::BatchGraphOutcome::Failed,
11995 MetalRowsRun::Completed(pending) => pending,
11996 };
11997 if std::env::var("CMF_GRAPH_SPEC_TIME").is_ok() {
11998 eprintln!(
11999 "metal-verify: {:.1} ms | b={b}",
12000 _t0.elapsed().as_secs_f64() * 1e3
12001 );
12002 }
12003 self.metal_verify = Some(pending);
12004 crate::gpu::BatchGraphOutcome::Completed
12005 }
12006
12007 #[cfg(target_os = "macos")]
12012 fn prefill_rows_metal(
12013 &mut self,
12014 ids: &[u32],
12015 start_pos: usize,
12016 spec: Option<((usize, usize, usize), &[f32], &mut Vec<f32>)>,
12017 ) -> MetalPrefillOutcome {
12018 let b = ids.len();
12019 if b == 0 || b > 512 {
12020 return MetalPrefillOutcome::Declined;
12021 }
12022 METAL_PREFILL_CHUNKS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
12023 let with_head = spec.is_some();
12024 let hs = self.hidden_size;
12025 let mut hiddens = vec![0f32; b * hs];
12026 for (j, &id) in ids.iter().enumerate() {
12027 let e = self.embed_single(id);
12028 hiddens[j * hs..(j + 1) * hs].copy_from_slice(&e);
12029 }
12030 let mut pending = match self.metal_rows_run(&mut hiddens, start_pos, b, true, spec, None) {
12031 MetalRowsRun::Declined => return MetalPrefillOutcome::Declined,
12032 MetalRowsRun::Failed => {
12033 METAL_PREFILL_ERRORS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
12034 return MetalPrefillOutcome::Failed;
12035 }
12036 MetalRowsRun::Completed(pending) => pending,
12037 };
12038 let idxs = pending.gdn_layers.clone();
12040 let mut outs: Vec<&mut [f32]> = self
12041 .kv_cache
12042 .layers
12043 .iter_mut()
12044 .enumerate()
12045 .filter(|(i, _)| idxs.binary_search(i).is_ok())
12046 .map(|(_, l)| l.linear_state.as_mut_slice())
12047 .collect();
12048 if !pending.graph.finish_states(&mut outs) {
12049 METAL_PREFILL_ERRORS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
12050 return MetalPrefillOutcome::Failed;
12051 }
12052 let (nkv, hd) = (self.num_kv_heads, self.head_dim);
12053 let mut rows = Vec::with_capacity(pending.attn_layers.len());
12057 for (li, cpu_stored) in &pending.attn_layers {
12058 let mut kbuf = vec![0f32; b * nkv * hd];
12059 let mut vbuf = vec![0f32; b * nkv * hd];
12060 if !crate::gpu_metal::kv_mirror_read_rows(
12061 self.graph_kv_id,
12062 *li,
12063 nkv,
12064 hd,
12065 *cpu_stored,
12066 b,
12067 &mut kbuf,
12068 &mut vbuf,
12069 ) {
12070 METAL_PREFILL_ERRORS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
12071 return MetalPrefillOutcome::Failed;
12072 }
12073 rows.push((*li, *cpu_stored, kbuf, vbuf));
12074 }
12075 for (li, cpu_stored, kbuf, vbuf) in rows {
12076 let cache = &mut self.kv_cache.layers[li];
12077 for r in 0..b {
12078 cache.append(
12079 &kbuf[r * nkv * hd..(r + 1) * nkv * hd],
12080 &vbuf[r * nkv * hd..(r + 1) * nkv * hd],
12081 &[],
12082 );
12083 }
12084 crate::gpu_metal::kv_mirror_set_stored(self.graph_kv_id, li, cpu_stored + b);
12085 }
12086 METAL_PREFILL_ROWS.fetch_add(b as u64, std::sync::atomic::Ordering::Relaxed);
12087 if with_head {
12088 METAL_PREFILL_HEAD_ROWS.fetch_add(b as u64, std::sync::atomic::Ordering::Relaxed);
12089 }
12090 MetalPrefillOutcome::Completed(hiddens)
12091 }
12092
12093 #[cfg(target_os = "macos")]
12094 fn prefill_batch_metal(&mut self, ids: &[u32], start_pos: usize) -> MetalPrefillOutcome {
12095 self.prefill_rows_metal(ids, start_pos, None)
12096 }
12097
12098 #[cfg(target_os = "macos")]
12103 fn nll_batch_metal(&mut self, ids: &[u32], start: usize) -> MetalBatchNllOutcome {
12104 if ids.len() < 2 || self.o1_active() || self.head_clusters.is_some() {
12105 return MetalBatchNllOutcome::Declined;
12106 }
12107 let Some(lm) = self.weights.lm_head.metal_graph_parts() else {
12108 return MetalBatchNllOutcome::Declined;
12109 };
12110 let chunk = std::env::var("CMF_METAL_PREFILL_CHUNK")
12111 .ok()
12112 .and_then(|v| v.parse::<usize>().ok())
12113 .filter(|&v| (1..=512).contains(&v))
12114 .unwrap_or(32);
12115 let final_norm = self.weights.final_norm.clone();
12116 let mut nll = 0.0f64;
12117 let mut count = 0usize;
12118 let mut pos = 0usize;
12119 let mut completed = 0usize;
12120 while pos < ids.len() {
12121 let end = (pos + chunk).min(ids.len());
12122 let mut logits = Vec::new();
12123 let outcome = self.prefill_rows_metal(
12124 &ids[pos..end],
12125 pos,
12126 Some((lm, &final_norm, &mut logits)),
12127 );
12128 match outcome {
12129 MetalPrefillOutcome::Declined => {
12130 return if completed == 0 {
12131 MetalBatchNllOutcome::Declined
12132 } else {
12133 MetalBatchNllOutcome::Failed(format!(
12134 "ordinary Metal NLL batch declined after {completed} chunks"
12135 ))
12136 };
12137 }
12138 MetalPrefillOutcome::Failed => {
12139 return MetalBatchNllOutcome::Failed(
12140 "ordinary Metal NLL batch failed after admission".to_string(),
12141 );
12142 }
12143 MetalPrefillOutcome::Completed(_) => {}
12144 }
12145 completed += 1;
12146 let vocab = self.vocab_size.min(lm.1);
12147 if logits.len() != (end - pos) * lm.1 || vocab == 0 {
12148 return MetalBatchNllOutcome::Failed(
12149 "ordinary Metal NLL head returned an invalid shape".to_string(),
12150 );
12151 }
12152 for row in 0..(end - pos) {
12153 let absolute = pos + row;
12154 if absolute < start || absolute + 1 >= ids.len() {
12155 continue;
12156 }
12157 let lg = &mut logits[row * lm.1..row * lm.1 + vocab];
12158 if let Some(mu) = self.logit_multiplier {
12159 for v in lg.iter_mut() {
12160 *v *= mu;
12161 }
12162 }
12163 if let Some(c) = self.final_softcap {
12164 for v in lg.iter_mut() {
12165 *v = c * (*v / c).tanh();
12166 }
12167 }
12168 let target = ids[absolute + 1] as usize;
12169 if target >= vocab {
12170 return MetalBatchNllOutcome::Failed(format!(
12171 "target token {target} exceeds Metal head rows {vocab}"
12172 ));
12173 }
12174 let max = lg.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
12175 let lse: f64 = lg
12176 .iter()
12177 .map(|&v| ((v - max) as f64).exp())
12178 .sum::<f64>()
12179 .ln()
12180 + max as f64;
12181 nll += lse - lg[target] as f64;
12182 count += 1;
12183 }
12184 pos = end;
12185 }
12186 MetalBatchNllOutcome::Completed(nll, count)
12187 }
12188
12189 #[cfg(target_os = "macos")]
12193 fn metal_verify_commit(&mut self, a: usize) -> bool {
12194 let Some(mut pending) = self.metal_verify.take() else {
12195 return false;
12196 };
12197 let n = a + 1;
12198 let idxs = pending.gdn_layers.clone();
12200 let mut outs: Vec<&mut [f32]> = self
12201 .kv_cache
12202 .layers
12203 .iter_mut()
12204 .enumerate()
12205 .filter(|(i, _)| idxs.binary_search(i).is_ok())
12206 .map(|(_, l)| l.linear_state.as_mut_slice())
12207 .collect();
12208 if !pending.graph.commit(n, &mut outs) {
12209 return false;
12210 }
12211 spec_stamp("c.replay");
12212 let (nkv, hd) = (self.num_kv_heads, self.head_dim);
12213 let mut rows = Vec::with_capacity(pending.attn_layers.len());
12217 for (li, cpu_stored) in &pending.attn_layers {
12218 let mut kbuf = vec![0f32; n * nkv * hd];
12219 let mut vbuf = vec![0f32; n * nkv * hd];
12220 if !crate::gpu_metal::kv_mirror_read_rows(
12221 self.graph_kv_id,
12222 *li,
12223 nkv,
12224 hd,
12225 *cpu_stored,
12226 n,
12227 &mut kbuf,
12228 &mut vbuf,
12229 ) {
12230 return false;
12231 }
12232 rows.push((*li, *cpu_stored, kbuf, vbuf));
12233 }
12234 for (li, cpu_stored, kbuf, vbuf) in rows {
12235 let cache = &mut self.kv_cache.layers[li];
12236 for r in 0..n {
12237 cache.append(
12238 &kbuf[r * nkv * hd..(r + 1) * nkv * hd],
12239 &vbuf[r * nkv * hd..(r + 1) * nkv * hd],
12240 &[],
12241 );
12242 }
12243 crate::gpu_metal::kv_mirror_set_stored(self.graph_kv_id, li, cpu_stored + n);
12244 }
12245 spec_stamp("c.kv");
12246 true
12247 }
12248
12249 #[cfg(target_os = "macos")]
12256 fn mtp_warm_batch_submit(
12257 &mut self,
12258 m: &mut MtpModule,
12259 pairs: &[(&[f32], u32)],
12260 first_pos: usize,
12261 ) -> Option<MetalWarmPending> {
12262 use crate::gpu_metal::{AttnDeviceParams, AttnGpuLayer, GraphDims, MetalFfn, VerifyGraph};
12263 let b = pairs.len();
12264 if b == 0 || b > 512 || m.kv.mode != crate::kv_cache::KvMode::F32 || m.kv.o1.is_some() {
12265 return None;
12266 }
12267 let AttnKind::Full {
12268 wq,
12269 wk,
12270 wv,
12271 wo,
12272 q_norm,
12273 k_norm,
12274 output_gate,
12275 softplus_gate: None,
12276 bias: None,
12277 } = &m.layer.attn
12278 else {
12279 return None;
12280 };
12281 let FfnKind::Dense(d) = &m.layer.ffn else {
12282 return None;
12283 };
12284 if !d.segs.is_empty() {
12285 return None;
12286 }
12287 let (Some(pq), Some(pk), Some(pv), Some(po)) =
12288 (wq.q1_parts(), wk.q1_parts(), wv.q1_parts(), wo.q1_parts())
12289 else {
12290 return None;
12291 };
12292 let (Some(g), Some(u), Some(dn)) = (
12293 d.gate_proj.q1_parts(),
12294 d.up_proj.q1_parts(),
12295 d.down_proj.q1_parts(),
12296 ) else {
12297 return None;
12298 };
12299 let Some(eh) = m.eh_proj.q1_parts() else {
12300 return None;
12301 };
12302 let QTensor::Mapped { model, .. } = wq else {
12303 return None;
12304 };
12305 let model = model.clone();
12306 let hs = self.hidden_size;
12307 let mut cat = vec![0f32; b * 2 * hs];
12309 for (j, (h, tok)) in pairs.iter().enumerate() {
12310 let e = self.embed_single(*tok);
12311 let (ce, ch) = cat[j * 2 * hs..(j + 1) * 2 * hs].split_at_mut(hs);
12312 inference::rms_norm_into(&e, &m.enorm, self.rms_eps, self.norm_style, ce);
12313 inference::rms_norm_into(h, &m.hnorm, self.rms_eps, self.norm_style, ch);
12314 }
12315 let dims = GraphDims {
12316 hidden: hs,
12317 eps: self.rms_eps as f32,
12318 gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
12319 };
12320 spec_stamp("w.cat");
12321 let Some(mut graph) = VerifyGraph::new_via_proj(&model, dims, eh, &cat, b) else {
12322 return None;
12323 };
12324 spec_stamp("w.new");
12325 let l = AttnGpuLayer {
12326 attn_norm: &m.layer.input_norm,
12327 post_norm: &m.layer.post_norm,
12328 wq: pq,
12329 wk: pk,
12330 wv: pv,
12331 wo: po,
12332 ffn: MetalFfn::Dense {
12333 gate: g,
12334 up: u,
12335 down: dn,
12336 gelu: d.act == Act::Gelu,
12337 },
12338 };
12339 let (nh, nkv, hd, rd) = (
12340 self.num_heads,
12341 self.num_kv_heads,
12342 self.head_dim,
12343 self.rotary_dim,
12344 );
12345 let inv_freq = self.inv_freq.clone();
12346 let cpu_stored;
12347 {
12348 let cache = &m.kv;
12349 let cpu_k: Vec<&[f32]> = (0..nkv).map(|g| cache.head_keys(g)).collect();
12350 let cpu_v: Vec<&[f32]> = (0..nkv).map(|g| cache.head_values(g)).collect();
12351 cpu_stored = cpu_k[0].len() / hd;
12352 if cpu_stored > first_pos {
12357 spec_stamp("w.decl");
12358 return None;
12359 }
12360 let p = AttnDeviceParams {
12361 kv_id: self.mtp_kv_id(),
12362 layer: Self::MTP_LAYER_BASE,
12363 nh,
12364 nkv,
12365 hd,
12366 rd,
12367 position: first_pos,
12368 scale: self.attn_scale,
12369 eps: self.rms_eps as f32,
12370 gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
12371 late_qk_norm: self.qk_norm_after_rope,
12372 output_gate: *output_gate,
12373 q_norm: q_norm.as_deref(),
12374 k_norm: k_norm.as_deref(),
12375 inv_freq: &inv_freq,
12376 cpu_k,
12377 cpu_v,
12378 cpu_stored,
12379 cpu_gen: cache.generation(),
12380 o1: None,
12381 window: None,
12382 head_gate: None,
12383 };
12384 if !graph.attn_ok(&l, &p) || !graph.encode_attn_b(&l, &p) {
12385 return None;
12386 }
12387 }
12388 spec_stamp("w.enc");
12389 if !graph.submit() {
12390 return None;
12391 }
12392 spec_stamp("w.sub");
12393 Some(MetalWarmPending {
12394 graph,
12395 cpu_stored,
12396 b,
12397 })
12398 }
12399
12400 #[cfg(target_os = "macos")]
12403 fn mtp_warm_batch_metal(
12404 &mut self,
12405 m: &mut MtpModule,
12406 pairs: &[(&[f32], u32)],
12407 first_pos: usize,
12408 ) -> bool {
12409 match self.mtp_warm_batch_submit(m, pairs, first_pos) {
12410 Some(p) => self.mtp_warm_batch_finish(m, p),
12411 None => false,
12412 }
12413 }
12414
12415 #[cfg(target_os = "macos")]
12420 fn mtp_warm_batch_finish(&mut self, m: &mut MtpModule, pending: MetalWarmPending) -> bool {
12421 let MetalWarmPending {
12422 mut graph,
12423 cpu_stored,
12424 b,
12425 } = pending;
12426 let (nkv, hd) = (self.num_kv_heads, self.head_dim);
12427 if !graph.sync() {
12428 return false;
12429 }
12430 spec_stamp("w.gpu");
12431 let mut kbuf = vec![0f32; b * nkv * hd];
12432 let mut vbuf = vec![0f32; b * nkv * hd];
12433 if !crate::gpu_metal::kv_mirror_read_rows(
12434 self.mtp_kv_id(),
12435 Self::MTP_LAYER_BASE,
12436 nkv,
12437 hd,
12438 cpu_stored,
12439 b,
12440 &mut kbuf,
12441 &mut vbuf,
12442 ) {
12443 return false;
12444 }
12445 for r in 0..b {
12446 m.kv.append(
12447 &kbuf[r * nkv * hd..(r + 1) * nkv * hd],
12448 &vbuf[r * nkv * hd..(r + 1) * nkv * hd],
12449 &[],
12450 );
12451 }
12452 crate::gpu_metal::kv_mirror_set_stored(
12453 self.mtp_kv_id(),
12454 Self::MTP_LAYER_BASE,
12455 cpu_stored + b,
12456 );
12457 spec_stamp("w.kv");
12458 true
12459 }
12460
12461 pub(crate) fn note_draft_id(&mut self, id: u32) {
12468 let cut = Self::draft_vocab_rows(usize::MAX).max(131_072);
12469 if (id as usize) >= cut {
12470 self.draft_full_streak = 16;
12471 } else {
12472 self.draft_full_streak = self.draft_full_streak.saturating_sub(1);
12473 }
12474 }
12475
12476 fn draft_head_rows(&self, head_rows: usize) -> usize {
12479 if self.draft_full_streak > 0 {
12480 head_rows
12481 } else {
12482 Self::draft_vocab_rows(head_rows)
12483 }
12484 }
12485
12486 fn draft_vocab_rows(head_rows: usize) -> usize {
12489 static N: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
12490 let n = *N.get_or_init(|| {
12491 std::env::var("CMF_DRAFT_VOCAB")
12492 .ok()
12493 .and_then(|v| v.parse().ok())
12494 .unwrap_or(65536)
12495 });
12496 if n == 0 { head_rows } else { n.min(head_rows) }
12497 }
12498
12499 #[cfg(target_os = "macos")]
12504 fn mtp_step_metal(
12505 &mut self,
12506 m: &mut MtpModule,
12507 hidden: &[f32],
12508 next_token: u32,
12509 position: usize,
12510 want_logits: bool,
12511 ) -> Option<(Vec<f32>, Vec<f32>)> {
12512 use crate::gpu_metal::{AttnDeviceParams, AttnGpuLayer, GraphDims, MetalFfn, TokenGraph};
12513 if std::env::var("CMF_MTP_GRAPH").as_deref() == Ok("0")
12514 || !crate::gpu::q1_force()
12515 || !crate::gpu::enabled_here()
12516 || self.attn_softcap > 0.0
12517 || self.attention_heads_per_layer.is_some()
12518 || m.kv.mode != crate::kv_cache::KvMode::F32
12519 || m.kv.o1.is_some()
12520 {
12521 return None;
12522 }
12523 let AttnKind::Full {
12524 wq,
12525 wk,
12526 wv,
12527 wo,
12528 q_norm,
12529 k_norm,
12530 output_gate,
12531 softplus_gate: None,
12532 bias: None,
12533 } = &m.layer.attn
12534 else {
12535 return None;
12536 };
12537 let FfnKind::Dense(d) = &m.layer.ffn else {
12538 return None;
12539 };
12540 if d.act != Act::Silu || !d.segs.is_empty() {
12541 return None;
12542 }
12543 let (pq, pk, pv, po) = (
12544 wq.q1_parts()?,
12545 wk.q1_parts()?,
12546 wv.q1_parts()?,
12547 wo.q1_parts()?,
12548 );
12549 let (g, u, dn) = (
12550 d.gate_proj.q1_parts()?,
12551 d.up_proj.q1_parts()?,
12552 d.down_proj.q1_parts()?,
12553 );
12554 let QTensor::Mapped { model, .. } = wq else {
12555 return None;
12556 };
12557 let model = model.clone();
12558 let lm = if want_logits {
12559 Some(self.weights.lm_head.q1_parts()?)
12560 } else {
12561 None
12562 };
12563 let dims = GraphDims {
12564 hidden: self.hidden_size,
12565 eps: self.rms_eps as f32,
12566 gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
12567 };
12568 let hs = self.hidden_size;
12571 let mut x = vec![0f32; hs];
12572 let mut graph = TokenGraph::new(&model, dims, &x)?;
12573 let mut folded = false;
12574 if let Some(eh) = m.eh_proj.q1_parts() {
12575 let e = self.embed_single(next_token);
12576 let mut cat = vec![0.0f32; 2 * hs];
12577 let (cat_e, cat_h) = cat.split_at_mut(hs);
12578 inference::rms_norm_into(&e, &m.enorm, self.rms_eps, self.norm_style, cat_e);
12579 inference::rms_norm_into(hidden, &m.hnorm, self.rms_eps, self.norm_style, cat_h);
12580 folded = graph.encode_input_proj(eh, &cat);
12581 }
12582 if !folded {
12583 x = self.mtp_block_input(m, hidden, next_token);
12584 graph = TokenGraph::new(&model, dims, &x)?;
12585 }
12586 spec_stamp("d.in");
12587 let l = AttnGpuLayer {
12588 attn_norm: &m.layer.input_norm,
12589 post_norm: &m.layer.post_norm,
12590 wq: pq,
12591 wk: pk,
12592 wv: pv,
12593 wo: po,
12594 ffn: MetalFfn::Dense {
12595 gate: g,
12596 up: u,
12597 down: dn,
12598 gelu: d.act == Act::Gelu,
12599 },
12600 };
12601 let (nh, nkv, hd, rd) = (
12602 self.num_heads,
12603 self.num_kv_heads,
12604 self.head_dim,
12605 self.rotary_dim,
12606 );
12607 let inv_freq = self.inv_freq.clone();
12608 {
12609 let cache = &m.kv;
12610 let cpu_k: Vec<&[f32]> = (0..nkv).map(|g| cache.head_keys(g)).collect();
12611 let cpu_v: Vec<&[f32]> = (0..nkv).map(|g| cache.head_values(g)).collect();
12612 let cpu_stored = cpu_k[0].len() / hd;
12613 let p = AttnDeviceParams {
12614 kv_id: self.mtp_kv_id(),
12615 layer: Self::MTP_LAYER_BASE,
12616 nh,
12617 nkv,
12618 hd,
12619 rd,
12620 position,
12621 scale: self.attn_scale,
12622 eps: self.rms_eps as f32,
12623 gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
12624 late_qk_norm: self.qk_norm_after_rope,
12625 output_gate: *output_gate,
12626 q_norm: q_norm.as_deref(),
12627 k_norm: k_norm.as_deref(),
12628 inv_freq: &inv_freq,
12629 cpu_k,
12630 cpu_v,
12631 cpu_stored,
12632 cpu_gen: cache.generation(),
12633 o1: None,
12634 window: None,
12635 head_gate: None,
12636 };
12637 if !graph.attn_device_ok(&l, &p) || !graph.encode_attn_device(&l, &p) {
12638 return None;
12639 }
12640 }
12641 let draft_rows = if let Some(lm) = lm {
12647 self.draft_head_rows(lm.1)
12648 } else {
12649 0
12650 };
12651 if let Some(lm) = lm {
12652 if !graph.lm_head_ok(lm) {
12653 return None;
12654 }
12655 if draft_rows < lm.1 {
12656 if !graph.encode_lm_head_part(&m.final_norm, lm, draft_rows) {
12657 return None;
12658 }
12659 } else {
12660 graph.encode_lm_head(&m.final_norm, lm);
12661 }
12662 }
12663 spec_stamp("d.enc");
12664 if graph.sync_checked().is_err() {
12665 return None;
12666 }
12667 spec_stamp("d.gpu");
12668 let mut logits = Vec::new();
12669 if let Some(lm) = lm {
12670 let n_read = draft_rows.min(lm.1).min(self.vocab_size);
12671 logits = attention::take_buf(n_read);
12672 graph.read_logits(&mut logits);
12673 logits.resize(self.vocab_size, f32::NEG_INFINITY);
12675 }
12676 graph.finish(&mut x);
12677 let mut krow = attention::take_buf(nkv * hd);
12678 let mut vrow = attention::take_buf(nkv * hd);
12679 if crate::gpu_metal::kv_mirror_read_last(
12680 self.mtp_kv_id(),
12681 Self::MTP_LAYER_BASE,
12682 nkv,
12683 hd,
12684 &mut krow,
12685 &mut vrow,
12686 ) {
12687 m.kv.append(&krow, &vrow, &[]);
12688 }
12689 attention::recycle_buf(&mut krow);
12690 attention::recycle_buf(&mut vrow);
12691 spec_stamp("d.rd");
12692 Some((logits, x))
12693 }
12694
12695 fn mtp_chain_on() -> bool {
12708 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
12709 *ON.get_or_init(|| std::env::var("CMF_MTP_CHAIN").as_deref() != Ok("0"))
12710 }
12711
12712 #[cfg(target_os = "macos")]
12724 fn mtp_draft_chain_metal(
12725 &mut self,
12726 m: &mut MtpModule,
12727 hidden: &[f32],
12728 t_next: u32,
12729 position: usize,
12730 k: usize,
12731 ) -> Result<Vec<u32>, bool> {
12732 use crate::gpu_metal::{AttnDeviceParams, AttnGpuLayer, GraphDims, MetalFfn, TokenGraph};
12733 if k == 0
12734 || k > 64
12735 || !Self::mtp_chain_on()
12736 || std::env::var("CMF_MTP_GRAPH").as_deref() == Ok("0")
12737 || !crate::gpu::q1_force()
12738 || !crate::gpu::enabled_here()
12739 || self.attn_softcap > 0.0
12740 || self.attention_heads_per_layer.is_some()
12741 || m.kv.mode != crate::kv_cache::KvMode::F32
12742 || m.kv.o1.is_some()
12743 || self.dsv4.is_some()
12745 || self.dsv41.is_some()
12746 || self.qwen4_exp.is_some()
12747 || self.g3n.is_some()
12748 {
12749 return Err(false);
12750 }
12751 let AttnKind::Full {
12752 wq,
12753 wk,
12754 wv,
12755 wo,
12756 q_norm,
12757 k_norm,
12758 output_gate,
12759 softplus_gate: None,
12760 bias: None,
12761 } = &m.layer.attn
12762 else {
12763 return Err(false);
12764 };
12765 let FfnKind::Dense(d) = &m.layer.ffn else {
12766 return Err(false);
12767 };
12768 if d.act != Act::Silu || !d.segs.is_empty() {
12769 return Err(false);
12770 }
12771 let (Some(pq), Some(pk), Some(pv), Some(po)) =
12772 (wq.q1_parts(), wk.q1_parts(), wv.q1_parts(), wo.q1_parts())
12773 else {
12774 return Err(false);
12775 };
12776 let (Some(g), Some(u), Some(dn)) = (
12777 d.gate_proj.q1_parts(),
12778 d.up_proj.q1_parts(),
12779 d.down_proj.q1_parts(),
12780 ) else {
12781 return Err(false);
12782 };
12783 let (Some(eh), Some(lm)) = (m.eh_proj.q1_parts(), self.weights.lm_head.q1_parts()) else {
12784 return Err(false);
12785 };
12786 let QTensor::Mapped { model, .. } = wq else {
12787 return Err(false);
12788 };
12789 let model = model.clone();
12790 let QTensor::Mapped {
12793 model: em,
12794 idx: eidx,
12795 dtype: cortiq_core::TensorDtype::Q4TiledP,
12796 ..
12797 } = &self.weights.embed_tokens
12798 else {
12799 return Err(false);
12800 };
12801 if !std::sync::Arc::ptr_eq(em, &model)
12802 || crate::prism::is_inverse_embedding(&model, &model.tensors[*eidx].name)
12803 {
12804 return Err(false);
12805 }
12806 let embed = (
12807 *eidx,
12808 self.weights.embed_tokens.rows(),
12809 self.weights.embed_tokens.cols(),
12810 );
12811 if embed.2 != self.hidden_size || hidden.len() != self.hidden_size {
12812 return Err(false);
12813 }
12814 let dims = GraphDims {
12815 hidden: self.hidden_size,
12816 eps: self.rms_eps as f32,
12817 gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
12818 };
12819 let Some(mut graph) = TokenGraph::new(&model, dims, hidden) else {
12820 return Err(false);
12821 };
12822 if !graph.chain_embed_ok(embed) || !graph.lm_head_ok(lm) {
12823 return Err(false);
12824 }
12825 let l = AttnGpuLayer {
12826 attn_norm: &m.layer.input_norm,
12827 post_norm: &m.layer.post_norm,
12828 wq: pq,
12829 wk: pk,
12830 wv: pv,
12831 wo: po,
12832 ffn: MetalFfn::Dense {
12833 gate: g,
12834 up: u,
12835 down: dn,
12836 gelu: d.act == Act::Gelu,
12837 },
12838 };
12839 let (nh, nkv, hd, rd) = (
12840 self.num_heads,
12841 self.num_kv_heads,
12842 self.head_dim,
12843 self.rotary_dim,
12844 );
12845 let inv_freq = self.inv_freq.clone();
12846 let draft_rows = self.draft_head_rows(lm.1);
12847 let n_arg = draft_rows.min(lm.1).min(self.vocab_size);
12848 if n_arg == 0 {
12849 return Err(false);
12850 }
12851 let split = std::env::var("CMF_MTP_CHAIN_SPLIT").as_deref() == Ok("1");
12858 let t_chain = std::time::Instant::now();
12859 graph.chain_ids_init(t_next, k);
12860 let cpu_stored;
12861 {
12862 let cache = &m.kv;
12863 let cpu_k: Vec<&[f32]> = (0..nkv).map(|g| cache.head_keys(g)).collect();
12864 let cpu_v: Vec<&[f32]> = (0..nkv).map(|g| cache.head_values(g)).collect();
12865 cpu_stored = cpu_k[0].len() / hd;
12866 for j in 0..k {
12867 if !graph.encode_chain_input(
12868 embed,
12869 j as u32,
12870 &m.enorm,
12871 &m.hnorm,
12872 self.embed_multiplier,
12873 eh,
12874 ) {
12875 return Err(false);
12876 }
12877 let p = AttnDeviceParams {
12881 kv_id: self.mtp_kv_id(),
12882 layer: Self::MTP_LAYER_BASE,
12883 nh,
12884 nkv,
12885 hd,
12886 rd,
12887 position: position + j,
12888 scale: self.attn_scale,
12889 eps: self.rms_eps as f32,
12890 gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
12891 late_qk_norm: self.qk_norm_after_rope,
12892 output_gate: *output_gate,
12893 q_norm: q_norm.as_deref(),
12894 k_norm: k_norm.as_deref(),
12895 inv_freq: &inv_freq,
12896 cpu_k: cpu_k.clone(),
12897 cpu_v: cpu_v.clone(),
12898 cpu_stored: cpu_stored + j,
12899 cpu_gen: cache.generation(),
12900 o1: None,
12901 window: None,
12902 head_gate: None,
12903 };
12904 if !graph.attn_device_ok(&l, &p) || !graph.encode_attn_device(&l, &p) {
12905 return Err(false);
12906 }
12907 if draft_rows < lm.1 {
12908 if !graph.encode_lm_head_part(&m.final_norm, lm, draft_rows) {
12909 return Err(false);
12910 }
12911 } else {
12912 graph.encode_lm_head(&m.final_norm, lm);
12913 }
12914 if !graph.encode_argmax(n_arg, j as u32 + 1) {
12915 return Err(false);
12916 }
12917 if split {
12918 graph.commit();
12921 }
12922 }
12923 }
12924 let t_enc = t_chain.elapsed();
12925 if graph.sync_checked().is_err() {
12926 return Err(true);
12927 }
12928 if std::env::var_os("CMF_GRAPH_SPEC_TIME").is_some() {
12929 eprintln!(
12930 "mtp-chain: encode {:.1} ms | wait {:.1} ms (k={k}, head rows {draft_rows}{})",
12931 t_enc.as_secs_f64() * 1e3,
12932 (t_chain.elapsed() - t_enc).as_secs_f64() * 1e3,
12933 if split { ", split" } else { "" }
12934 );
12935 }
12936 let mut ids = vec![0u32; k];
12937 if !graph.chain_ids_read(&mut ids) {
12938 return Err(true);
12939 }
12940 let mut kbuf = vec![0f32; k * nkv * hd];
12941 let mut vbuf = vec![0f32; k * nkv * hd];
12942 if !crate::gpu_metal::kv_mirror_read_rows(
12943 self.mtp_kv_id(),
12944 Self::MTP_LAYER_BASE,
12945 nkv,
12946 hd,
12947 cpu_stored,
12948 k,
12949 &mut kbuf,
12950 &mut vbuf,
12951 ) {
12952 return Err(true);
12953 }
12954 for r in 0..k {
12955 m.kv.append(
12956 &kbuf[r * nkv * hd..(r + 1) * nkv * hd],
12957 &vbuf[r * nkv * hd..(r + 1) * nkv * hd],
12958 &[],
12959 );
12960 }
12961 Ok(ids)
12962 }
12963
12964 fn try_batch_graph_wgpu(
12965 &self,
12966 hiddens: &mut [f32],
12967 positions: &[usize],
12968 k: usize,
12969 spec: Option<crate::gpu::SpecTail<'_>>,
12970 ) -> crate::gpu::BatchGraphOutcome {
12971 self.try_batch_graph_wgpu_prefix(hiddens, positions, k, spec, None)
12972 }
12973
12974 fn try_batch_graph_wgpu_prefix(
12979 &self,
12980 hiddens: &mut [f32],
12981 positions: &[usize],
12982 k: usize,
12983 spec: Option<crate::gpu::SpecTail<'_>>,
12984 layers_run: Option<&mut usize>,
12985 ) -> crate::gpu::BatchGraphOutcome {
12986 let graph_end = match self.mimo_moe.graph_prefix_end() {
12987 Some(end) if end < self.num_layers => {
12988 if layers_run.is_none() || spec.is_some() || end == 0 {
12989 return crate::gpu::BatchGraphOutcome::Declined;
12990 }
12991 end
12992 }
12993 _ => self.num_layers,
12994 };
12995 let _tb = std::time::Instant::now();
12996 let batch_debug = std::env::var_os("CMF_BATCH_DEBUG").is_some();
12997 if self.attn_softcap > 0.0 {
12998 return crate::gpu::BatchGraphOutcome::Declined; }
13000 if let Some(reason) = self.wgpu_graph_attn_decline() {
13003 self.note_graph_decline("wgpu batch graph", reason);
13004 return crate::gpu::BatchGraphOutcome::Declined;
13005 }
13006 let nh = self.num_heads;
13007 let (nkv, hd, rd) = self.layer_geom(0);
13008 let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
13009 fn gw(t: &QTensor) -> Option<crate::gpu::GraphW<'_>> {
13010 if let Some((m, i, kind, rs)) = t
13011 .graph_weight()
13012 .or_else(|| t.graph_weight_descriptor())
13013 {
13014 let name = &m.tensors[i].name;
13015 let prism = if crate::prism::is_inverse_embedding(m, name) {
13016 crate::gpu::GraphPrismOp::InverseEmbedding
13017 } else if crate::prism::is_forward_weight(m, name) {
13018 crate::gpu::GraphPrismOp::Forward
13019 } else {
13020 crate::gpu::GraphPrismOp::None
13021 };
13022 return Some(crate::gpu::GraphW {
13023 idx: i,
13024 kind,
13025 row_scale: rs,
13026 data: &[],
13027 prism,
13028 affine: crate::prism::is_affine_target(m, name),
13029 });
13030 }
13031 if std::env::var_os("CMF_BATCH_DEBUG").is_some() {
13032 eprintln!(
13033 "batch graph: tensor has no graph descriptor/f32 fallback rows={} cols={}",
13034 t.rows(),
13035 t.cols()
13036 );
13037 }
13038 t.as_f32().map(|d| crate::gpu::GraphW {
13039 idx: 0,
13040 kind: 4,
13041 row_scale: &[],
13042 data: d,
13043 prism: crate::gpu::GraphPrismOp::None,
13044 affine: false,
13045 })
13046 }
13047 let built: Option<(
13048 Vec<crate::gpu::GraphLayer<'_>>,
13049 std::sync::Arc<cortiq_core::CmfModel>,
13050 )> = (|| {
13051 let mut layers = Vec::with_capacity(graph_end);
13052 let mut model = None;
13053 for li in 0..graph_end {
13054 let lw = &self.weights.layers[self.phys_layer(li)];
13055 let gffn = match &lw.ffn {
13062 FfnKind::Dense(d) if !d.segs.is_empty() => {
13063 if batch_debug {
13064 eprintln!("batch graph: dense segmented FFN at layer {li}");
13065 }
13066 return None;
13067 }
13068 FfnKind::Dense(d) => {
13069 let Some(act) = d.act.graph_act() else {
13070 if batch_debug {
13071 eprintln!(
13072 "batch graph: dense FFN activation {:?} without a graph kernel at layer {li}",
13073 d.act
13074 );
13075 }
13076 return None;
13077 };
13078 crate::gpu::GraphFfn::Dense {
13079 gate: gw(&d.gate_proj)?,
13080 up: gw(&d.up_proj)?,
13081 down: gw(&d.down_proj)?,
13082 act,
13083 }
13084 }
13085 FfnKind::Moe(m) => {
13086 if m.route_tau.is_some() || m.mask.is_some() {
13093 return None;
13094 }
13095 let shared = m.shared.as_ref();
13099 let has_shared = shared.is_some();
13100 let shared_gated = matches!(shared, Some((_, Some(_))));
13101 let sgate = match shared {
13102 Some((_, Some(sg))) => gw(sg)?,
13103 _ => gw(&m.router)?,
13107 };
13108 let router = gw(&m.router)?;
13109 if router.prism != crate::gpu::GraphPrismOp::None
13115 || router.affine
13116 || sgate.prism != crate::gpu::GraphPrismOp::None
13117 || sgate.affine
13118 {
13119 return None;
13120 }
13121 let inter = m.experts.first()?.gate_proj.rows();
13122 let mut experts = Vec::with_capacity(m.experts.len() + 1);
13123 let mut q4tp: Option<bool> = None;
13124 let mut gu_q2: Option<bool> = None;
13125 for e in m.experts.iter().chain(shared.map(|(se, _)| se)) {
13126 if !matches!(e.act, Act::Silu)
13127 || e.gate_proj.rows() != inter
13128 || e.up_proj.rows() != inter
13129 {
13130 return None;
13131 }
13132 let (mm, gi, ui, di, is_p, is_q2) = match e.gate_proj.mapped_q4t() {
13136 Some((mm, gi)) => (
13137 mm,
13138 gi,
13139 e.up_proj.mapped_q4t()?.1,
13140 e.down_proj.mapped_q4t()?.1,
13141 false,
13142 false,
13143 ),
13144 None => match e.gate_proj.mapped_q2tp() {
13145 Some((mm, gi)) => (
13146 mm,
13147 gi,
13148 e.up_proj.mapped_q2tp()?.1,
13149 e.down_proj.mapped_q4tp()?.1,
13150 true,
13151 true,
13152 ),
13153 None => {
13154 let (mm, gi) = e.gate_proj.mapped_q4tp()?;
13155 (
13156 mm,
13157 gi,
13158 e.up_proj.mapped_q4tp()?.1,
13159 e.down_proj.mapped_q4tp()?.1,
13160 true,
13161 false,
13162 )
13163 }
13164 },
13165 };
13166 if *q4tp.get_or_insert(is_p) != is_p
13167 || *gu_q2.get_or_insert(is_q2) != is_q2
13168 {
13169 return None;
13170 }
13171 if [gi, ui, di].into_iter().any(|idx| {
13172 mm.tensors
13173 .get(idx)
13174 .is_some_and(|t| {
13175 crate::prism::is_forward_weight(mm, &t.name)
13176 || crate::prism::is_affine_target(mm, &t.name)
13177 })
13178 }) {
13179 return None;
13180 }
13181 model.get_or_insert_with(|| mm.clone());
13182 experts.push((gi, ui, di));
13183 }
13184 crate::gpu::GraphFfn::Moe {
13185 router,
13186 shared_gate: sgate,
13187 experts,
13188 n_exp: m.experts.len(),
13189 top_k: m.top_k,
13190 inter,
13191 norm_topk: m.norm_topk_prob,
13192 q4tp: q4tp?,
13193 gu_q2: gu_q2.unwrap_or(false),
13194 sigmoid: m.router_sigmoid,
13195 bias: m.expert_bias.as_deref(),
13196 has_shared,
13197 shared_gated,
13198 route_scale: m.routed_scaling,
13199 }
13200 }
13201 _ => return None,
13202 };
13203 let attn = match &lw.attn {
13204 AttnKind::Full {
13205 wq,
13206 wk,
13207 wv,
13208 wo,
13209 q_norm,
13210 k_norm,
13211 output_gate,
13212 softplus_gate,
13213 bias,
13214 } => {
13215 if softplus_gate.is_some() || self.attention_heads_per_layer.is_some() {
13216 if batch_debug {
13217 eprintln!(
13218 "batch graph: unsupported Full attention gate at layer {li} softplus={} heads={}",
13219 softplus_gate.is_some(),
13220 self.attention_heads_per_layer.is_some()
13221 );
13222 }
13223 return None;
13224 }
13225 let (m, _, _, _) = wq
13226 .graph_weight()
13227 .or_else(|| wq.graph_weight_descriptor())?;
13228 model = Some(m.clone());
13229 crate::gpu::GraphAttn::Full {
13230 wq: gw(wq)?,
13231 wk: gw(wk)?,
13232 wv: gw(wv)?,
13233 wo: gw(wo)?,
13234 q_norm: q_norm.as_deref(),
13235 k_norm: k_norm.as_deref(),
13236 late_qk_norm: self.qk_norm_after_rope,
13237 bias: bias
13238 .as_ref()
13239 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
13240 output_gate: *output_gate,
13241 cpu_k: self.kv_cache.layers[li].k_heads(),
13242 cpu_v: self.kv_cache.layers[li].v_heads(),
13243 cpu_base: self.kv_cache.layers[li].base(),
13244 geom: self.graph_attn_geom(li),
13245 head_gate: None,
13248 }
13249 }
13250 AttnKind::LinearGdn(w) => {
13251 let Some(cfg) = self.gdn_cfg else {
13252 if batch_debug {
13253 eprintln!("batch graph: no GDN config at layer {li}");
13254 }
13255 return None;
13256 };
13257 let (m, _, _, _) = w
13258 .in_proj_qkv
13259 .graph_weight()
13260 .or_else(|| w.in_proj_qkv.graph_weight_descriptor())?;
13261 model = Some(m.clone());
13262 crate::gpu::GraphAttn::Gdn {
13263 qkv: gw(&w.in_proj_qkv)?,
13264 z: gw(&w.in_proj_z)?,
13265 a: gw(&w.in_proj_a)?,
13266 b: gw(&w.in_proj_b)?,
13267 out: gw(&w.out_proj)?,
13268 conv1d: &w.conv1d,
13269 a_log: &w.a_log,
13270 dt_bias: &w.dt_bias,
13271 norm: &w.norm,
13272 nv: cfg.num_v_heads,
13273 nk: cfg.num_k_heads,
13274 dk: cfg.key_head_dim,
13275 dv: cfg.value_head_dim,
13276 kk: cfg.conv_kernel,
13277 cpu_state: &self.kv_cache.layers[self.phys_layer(li)].linear_state,
13278 }
13279 }
13280 _ => return None,
13281 };
13282 layers.push(crate::gpu::GraphLayer {
13283 input_norm: &lw.input_norm,
13284 attn,
13285 post_norm: &lw.post_norm,
13286 ffn: gffn,
13287 });
13288 }
13289 Some((layers, model?))
13290 })();
13291 let Some((layers, model)) = built else {
13292 {
13293 use std::sync::atomic::{AtomicBool, Ordering};
13294 static SAID: AtomicBool = AtomicBool::new(false);
13295 if !SAID.swap(true, Ordering::Relaxed) {
13296 tracing::warn!("batch graph: BUILDER refused (layer weights/kinds)");
13297 }
13298 }
13299 return crate::gpu::BatchGraphOutcome::Declined;
13300 };
13301 if std::env::var("CMF_GRAPH_SPEC_TIME").is_ok() {
13302 eprintln!("batch-build: {:.1} ms", _tb.elapsed().as_secs_f64() * 1e3);
13303 }
13304 crate::gpu::forward_batch_graph(
13305 &model,
13306 self.graph_kv_id,
13307 &layers,
13308 &self.inv_freq,
13309 hiddens,
13310 nh,
13311 nkv,
13312 hd,
13313 rd,
13314 self.hidden_size,
13315 self.intermediate_size,
13316 positions,
13317 self.kv_cache.max_seq_len,
13318 gemma,
13319 self.rms_eps as f32,
13320 self.attn_scale,
13321 k,
13322 &(0..graph_end)
13323 .map(|li| self.kv_cache.layers[self.phys_layer(li)].o1_views())
13324 .collect::<Vec<_>>(),
13325 self.o1_epoch,
13326 spec,
13327 layers_run,
13328 )
13329 }
13330
13331 fn draft_probe() -> bool {
13335 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
13336 *ON.get_or_init(|| std::env::var("CMF_DSV4_DRAFT_PROBE").is_ok_and(|v| v != "0"))
13337 }
13338
13339 #[cfg(feature = "gpu")]
13351 fn dsv4_spec_on() -> bool {
13352 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
13353 *ON.get_or_init(|| {
13354 if let Ok(v) = std::env::var("CMF_DSV4_SPEC_RUN") {
13358 return v != "0";
13359 }
13360 std::env::var("CMF_DSV4_SPEC")
13367 .map(|v| v != "0")
13368 .unwrap_or_else(|_| {
13369 crate::gpu_wgpu::DRAFT_RESERVE.load(std::sync::atomic::Ordering::Relaxed) > 0
13370 })
13371 })
13372 }
13373
13374 #[cfg(feature = "gpu")]
13381 fn dsv4_spec_step(
13382 &mut self,
13383 tip_token: u32,
13384 t_next: u32,
13385 next_pos: usize,
13386 max_extra: usize,
13387 drafted: &mut usize,
13388 accepted_ctr: &mut usize,
13389 ) -> Option<(Vec<u32>, usize)> {
13390 let t_all = std::time::Instant::now();
13391 if std::env::var("CMF_DSV4_SPEC_TIME").is_ok() {
13392 thread_local! {
13393 static LAST: std::cell::Cell<Option<std::time::Instant>> =
13394 const { std::cell::Cell::new(None) };
13395 }
13396 LAST.with(|l| {
13397 if let Some(prev) = l.get() {
13398 eprintln!(
13399 "между раундами {:.1} мс",
13400 prev.elapsed().as_secs_f64() * 1e3
13401 );
13402 }
13403 l.set(Some(std::time::Instant::now()));
13404 });
13405 }
13406 if std::env::var("CMF_DSV4_SPEC_DEBUG").is_ok() {
13407 eprintln!("spec_step: вход pos={next_pos}");
13408 }
13409 let n_layers = self.dsv4.as_ref().map(|b| b.1.len())?;
13410 let cfg = self.dsv4.as_ref().map(|b| b.2)?;
13411 if self.dspark.is_none() {
13413 let t = crate::dsv4::dspark_targets(&self.dsv4_mtp, &cfg, n_layers);
13414 if t.is_empty() {
13415 return None;
13416 }
13417 crate::dsv4::dspark_arm(&t, cfg.dim);
13418 self.dspark = Some(crate::dsv4::DsparkState::new(
13419 self.dsv4_mtp.len(),
13420 &cfg,
13421 t.len(),
13422 ));
13423 }
13424 let targets = crate::dsv4::dspark_targets(&self.dsv4_mtp, &cfg, n_layers);
13425 let pack = crate::dsv4::dspark_pack_get(&self.dsv4_mtp, &cfg);
13426 if pack.is_none() && std::env::var("CMF_DSV4_SPEC_DEBUG").is_ok() {
13427 eprintln!("spec_step: пак не построился (targets {targets:?})");
13428 }
13429 let pack = pack?;
13430 let block = crate::dsv4::dspark_block();
13431 let b_box = self.dsv4.as_mut()?;
13432 let (g, layers, st) = (&b_box.0, &b_box.1, &mut b_box.3);
13433 let ds = self.dspark.as_mut()?;
13434 let dbg = std::env::var("CMF_DSV4_SPEC_DEBUG").is_ok();
13437 if !crate::dsv4::dspark_take(&mut ds.main_hidden) && !ds.have_hidden {
13438 if dbg {
13439 eprintln!("spec_step: нет захвата");
13440 }
13441 return None;
13442 }
13443 ds.have_hidden = true;
13444 let tip_pos = next_pos.checked_sub(1)?;
13445 let draft_started = std::time::Instant::now();
13446 let mut conf = Vec::new();
13447 let props = crate::dsv4::dspark_draft_gpu(
13448 g,
13449 &self.dsv4_mtp,
13450 &cfg,
13451 ds,
13452 pack,
13453 st.kv_id,
13454 tip_token,
13455 tip_pos,
13456 self.pool.as_deref(),
13457 &mut conf,
13458 );
13459 self.dspark_draft_ns += draft_started.elapsed().as_nanos();
13460 *drafted += block;
13461 if props.is_empty() || props[0] != t_next {
13462 if dbg {
13463 eprintln!(
13464 "spec_step: черновик {} (props0={:?} t_next={t_next})",
13465 if props.is_empty() {
13466 "пуст"
13467 } else {
13468 "мимо"
13469 },
13470 props.first()
13471 );
13472 }
13473 return None;
13474 }
13475 let mut k_verify = crate::dsv4::dspark_verify_k()
13482 .min(props.len())
13483 .min(max_extra.saturating_add(1));
13484 let conf_min = {
13490 static M: std::sync::OnceLock<f32> = std::sync::OnceLock::new();
13491 *M.get_or_init(|| {
13492 std::env::var("CMF_DSPARK_CONF_MIN")
13493 .ok()
13494 .and_then(|v| v.parse().ok())
13495 .unwrap_or(0.0)
13496 })
13497 };
13498 if conf_min > 0.0 && conf.len() >= props.len() {
13499 let mut keep = 1usize;
13500 while keep < k_verify && conf.get(keep).copied().unwrap_or(0.0) >= conf_min {
13501 keep += 1;
13502 }
13503 k_verify = k_verify.min(keep.max(2));
13504 }
13505 if k_verify < 2 {
13506 return None;
13507 }
13508 let mut fed = Vec::with_capacity(k_verify);
13509 fed.push(t_next);
13510 fed.extend_from_slice(&props[1..k_verify]);
13511 let mut argmax = Vec::new();
13512 let mut logits_all = Vec::new();
13513 let mut walked = Vec::new();
13514 let txn = crate::dsv4::dsv4_verify_chunk(
13515 g,
13516 layers,
13517 &cfg,
13518 st,
13519 &fed,
13520 next_pos,
13521 &self.inv_freq,
13522 self.pool.as_deref(),
13523 &targets,
13524 &mut argmax,
13525 &mut logits_all,
13526 &mut walked,
13527 );
13528 if txn.is_none() && dbg {
13529 eprintln!("spec_step: verify отказал");
13530 }
13531 let txn = txn?;
13532 let spec_gpu_end = txn.gpu_end;
13533 let b = fed.len();
13534 let mut accepted = 1usize;
13535 while accepted < b && fed[accepted] == argmax[accepted - 1] {
13536 accepted += 1;
13537 }
13538 if std::env::var("CMF_DSV4_SPEC_FORCE_REJECT").is_ok_and(|v| v != "0") {
13543 accepted = 1;
13544 }
13545 if std::env::var("CMF_DSV4_SPEC_TRACE").is_ok() {
13546 eprintln!("spec@{next_pos}: fed={fed:?} argmax={argmax:?} accepted={accepted}");
13547 }
13548 let t_fin = std::time::Instant::now();
13549 if !crate::dsv4::dsv4_spec_finish(
13550 g,
13551 layers,
13552 &cfg,
13553 st,
13554 txn,
13555 accepted,
13556 &fed,
13557 &self.inv_freq,
13558 self.pool.as_deref(),
13559 ) {
13560 tracing::warn!("dsv4: спекулятивный откат не удался — состояние подозрительно");
13561 return None;
13562 }
13563 if std::env::var("CMF_DSV4_SPEC_TIME").is_ok() {
13564 eprintln!(
13565 "finish(k={accepted}): {:.1} мс",
13566 t_fin.elapsed().as_secs_f64() * 1e3
13567 );
13568 }
13569 *accepted_ctr += accepted - 1;
13570 let (hc, dim) = (cfg.hc_mult, cfg.dim);
13575 let dev_caps: Vec<usize> = targets
13580 .iter()
13581 .copied()
13582 .filter(|&t| t < spec_gpu_end)
13583 .collect();
13584 let mut caps_all = vec![0.0f32; dev_caps.len() * b * hc * dim];
13585 if !crate::gpu_wgpu::dsv4_spec_cap_read_all(b, dev_caps.len(), hc * dim, &mut caps_all) {
13586 return None;
13587 }
13588 for t in 0..accepted {
13589 let tip = t + 1 == accepted;
13590 for (slot, &tl) in targets.iter().enumerate() {
13591 if let Some(di) = dev_caps.iter().position(|&d| d == tl) {
13592 let lo = (di * b + t) * hc * dim;
13593 crate::dsv4::dspark_capture(
13594 &caps_all[lo..lo + hc * dim],
13595 &cfg,
13596 slot,
13597 &mut ds.main_hidden,
13598 );
13599 } else if tip
13600 && crate::dsv4::dspark_peek_slot(slot, dim, {
13601 let lo = slot * dim;
13602 &mut ds.main_hidden[lo..lo + dim]
13603 })
13604 {
13605 } else {
13610 crate::dsv4::dspark_capture(
13614 &walked[t * hc * dim..(t + 1) * hc * dim],
13615 &cfg,
13616 slot,
13617 &mut ds.main_hidden,
13618 );
13619 }
13620 }
13621 crate::dsv4::dspark_ring_append(
13622 g,
13623 &self.dsv4_mtp,
13624 &cfg,
13625 ds,
13626 next_pos + t,
13627 self.pool.as_deref(),
13628 );
13629 }
13630 let row = logits_all[(accepted - 1) * cfg.vocab..accepted * cfg.vocab].to_vec();
13631 self.graph_logits = Some(row);
13632 if std::env::var("CMF_DSV4_TRUNK_PICK_DUMP").is_ok() {
13637 crate::dsv4::trunk_freq_note(&crate::dsv4::pick_tally_take());
13638 crate::dsv4::pick_tally_arm();
13639 }
13640 if std::env::var("CMF_DSV4_SPEC_TIME").is_ok() {
13641 eprintln!(
13642 "spec_step total {:.1} мс (k={accepted})",
13643 t_all.elapsed().as_secs_f64() * 1e3
13644 );
13645 }
13646 Some((fed[1..accepted].to_vec(), next_pos + accepted))
13647 }
13648
13649 fn dspark_probe(&mut self, position: usize, token_id: u32) {
13650 if self.dsv4_mtp.is_empty() || !Self::draft_probe() {
13651 return;
13652 }
13653 let trunk_now = crate::dsv4::pick_tally_take();
13655 crate::dsv4::trunk_freq_note(&trunk_now);
13656 if !trunk_now.is_empty() {
13657 self.dspark_trunk_picks.push(trunk_now);
13658 let keep = crate::dsv4::dspark_block();
13659 if self.dspark_trunk_picks.len() > keep {
13660 self.dspark_trunk_picks.remove(0);
13661 }
13662 }
13663 for p in std::mem::take(&mut self.dspark_pending) {
13666 let Some(i) = position.checked_sub(p.0 + 1) else {
13667 continue;
13668 };
13669 let mut p = p;
13670 if i < p.1.len() {
13671 if p.2 && p.1[i] == token_id {
13672 p.3 = i + 1;
13673 } else {
13674 p.2 = false;
13675 }
13676 if i + 1 < p.1.len() {
13677 self.dspark_pending.push(p);
13678 continue;
13679 }
13680 }
13681 self.dspark_hist.push(p.3);
13682 self.dspark_real.push(token_id);
13683 }
13684 let Some(b) = &mut self.dsv4 else { return };
13685 let (g, layers, cfg) = (&b.0, &b.1, b.2);
13686 let n_layers = layers.len();
13687 if self.dspark.is_none() {
13688 let t = crate::dsv4::dspark_targets(&self.dsv4_mtp, &cfg, n_layers);
13689 if t.is_empty() {
13690 return;
13691 }
13692 eprintln!(
13693 "DSpark: захват со слоёв {t:?}, блок {}",
13694 crate::dsv4::dspark_block()
13695 );
13696 crate::dsv4::dspark_arm(&t, cfg.dim);
13697 self.dspark = Some(crate::dsv4::DsparkState::new(
13698 self.dsv4_mtp.len(),
13699 &cfg,
13700 t.len(),
13701 ));
13702 }
13703 let ds = self.dspark.as_mut().unwrap();
13704 if !crate::dsv4::dspark_take(&mut ds.main_hidden) {
13705 return; }
13707 let mut conf = Vec::new();
13708 crate::dsv4::pick_tally_arm();
13709 let draft_started = std::time::Instant::now();
13714 #[cfg(feature = "gpu")]
13715 let gpu_draft = crate::dsv4::dspark_gpu_on();
13716 #[cfg(not(feature = "gpu"))]
13717 let gpu_draft = false;
13718 let props = if gpu_draft {
13719 #[cfg(feature = "gpu")]
13720 {
13721 let kv_id = b.3.kv_id;
13722 match crate::dsv4::dspark_pack_get(&self.dsv4_mtp, &cfg) {
13723 Some(pk) => crate::dsv4::dspark_draft_gpu(
13724 g,
13725 &self.dsv4_mtp,
13726 &cfg,
13727 ds,
13728 pk,
13729 kv_id,
13730 token_id,
13731 position,
13732 self.pool.as_deref(),
13733 &mut conf,
13734 ),
13735 None => Vec::new(),
13736 }
13737 }
13738 #[cfg(not(feature = "gpu"))]
13739 Vec::new()
13740 } else {
13741 crate::gpu::cpu_scope(|| {
13742 crate::dsv4::dspark_draft(
13743 g,
13744 &self.dsv4_mtp,
13745 &cfg,
13746 ds,
13747 token_id,
13748 position,
13749 self.pool.as_deref(),
13750 &mut conf,
13751 )
13752 })
13753 };
13754 self.dspark_draft_ns += draft_started.elapsed().as_nanos();
13755 let draft_picks = crate::dsv4::pick_tally_take();
13756 crate::dsv4::dspark_freq_note(&draft_picks);
13757 crate::dsv4::pick_tally_arm();
13760 if !props.is_empty() {
13761 let (tu, tt) = {
13765 let flat: Vec<(usize, Vec<usize>)> = self
13766 .dspark_trunk_picks
13767 .iter()
13768 .flat_map(|v| v.iter().cloned())
13769 .collect();
13770 let mut per: std::collections::HashMap<usize, Vec<usize>> =
13772 std::collections::HashMap::new();
13773 for (li, picks) in flat {
13774 per.entry(li).or_default().extend(picks);
13775 }
13776 let n = per.len().max(1);
13777 let mut u = 0usize;
13778 let mut t = 0usize;
13779 for (_, v) in per {
13780 t += v.len();
13781 u += v.iter().collect::<std::collections::HashSet<_>>().len();
13782 }
13783 (u / n, t / n)
13784 };
13785 let (du, dt) = crate::dsv4::tally_unique(&draft_picks);
13786 self.dspark_exp.push((tu, tt, du, dt));
13787 self.dspark_pending.push((position, props, true, 0));
13788 }
13789 if self.dspark_hist.len() >= 8 && self.dspark_hist.len() % 8 == 0 {
13790 let n = self.dspark_hist.len() as f32;
13791 let mean: f32 = self.dspark_hist.iter().sum::<usize>() as f32 / n;
13792 let block = crate::dsv4::dspark_block();
13793 let mut at = vec![0usize; block + 1];
13794 for &k in &self.dspark_hist {
13795 at[k] += 1;
13796 }
13797 let mut surv = Vec::with_capacity(block);
13799 for i in 1..=block {
13800 let k = at[i..].iter().sum::<usize>() as f32 / n;
13801 surv.push(format!("{k:.2}"));
13802 }
13803 let distinct = self
13804 .dspark_real
13805 .iter()
13806 .collect::<std::collections::HashSet<_>>()
13807 .len();
13808 let (tu, tt, du, dt) = self.dspark_exp.iter().fold((0, 0, 0, 0), |a, b| {
13809 (a.0 + b.0, a.1 + b.1, a.2 + b.2, a.3 + b.3)
13810 });
13811 let m = self.dspark_exp.len().max(1);
13812 eprintln!(
13813 "DSpark: черновиков {}, принято в среднем {mean:.2} из {block} \
13814 (токенов за проход {:.2}), распределение {at:?}, выживание [{}]",
13815 self.dspark_hist.len(),
13816 mean + 1.0,
13817 surv.join(" ")
13818 );
13819 eprintln!(
13820 "DSpark: разных токенов {distinct} из {} (вырожденность), \
13821 эксперты ствол {}/{} на слой за {block} токенов, \
13822 черновик {}/{} за блок, draft {:.2} мс/блок",
13823 self.dspark_real.len(),
13824 tu / m,
13825 tt / m,
13826 du / m,
13827 dt / m,
13828 self.dspark_draft_ns as f64 / self.dspark_exp.len().max(1) as f64 / 1e6
13829 );
13830 }
13831 }
13832
13833 fn forward_layers_upto(
13834 &mut self,
13835 hidden: &[f32],
13836 position: usize,
13837 task_mask: Option<&TaskMask>,
13838 upto: Option<usize>,
13839 ) -> Vec<f32> {
13840 if let Some(plan) = self.gpu_plan.clone() {
13846 if upto.is_none() && plan.len() > 1 {
13847 let mut h = hidden.to_vec();
13848 for &(dev, from, upto_incl) in plan.iter() {
13849 h = crate::gpu::with_device(dev, || {
13850 self.forward_layers_span(&h, position, task_mask, from, Some(upto_incl))
13851 });
13852 }
13853 return h;
13854 }
13855 }
13856 self.forward_layers_span(hidden, position, task_mask, 0, upto)
13857 }
13858
13859 pub fn set_gpu_plan(&mut self, devices: Option<&[usize]>) -> Result<(), String> {
13864 self.set_gpu_plan_at(devices, None)
13865 }
13866
13867 pub fn set_gpu_plan_at(
13871 &mut self,
13872 devices: Option<&[usize]>,
13873 at: Option<usize>,
13874 ) -> Result<(), String> {
13875 let Some(devs) = devices.filter(|d| d.len() > 1) else {
13876 self.gpu_plan = None;
13877 return Ok(());
13878 };
13879 self.split_supported()?;
13880 let n = self.num_layers;
13881 if devs.len() > n {
13882 return Err(format!("{} devices for {n} layers", devs.len()));
13883 }
13884 if let Some(k) = at {
13885 if k == 0 || k >= n {
13886 return Err(format!("split at {k}: the model has {n} layers"));
13887 }
13888 if devs.len() == 2 {
13889 self.gpu_plan = Some(std::sync::Arc::new(vec![
13890 (devs[0], 0, k - 1),
13891 (devs[1], k, n - 1),
13892 ]));
13893 return Ok(());
13894 }
13895 return Err(format!(
13896 "an explicit split point takes exactly 2 devices, got {}",
13897 devs.len()
13898 ));
13899 }
13900 let per = n.div_ceil(devs.len());
13901 let mut plan = Vec::with_capacity(devs.len());
13902 let mut from = 0usize;
13903 for &d in devs {
13904 if from >= n {
13905 break;
13906 }
13907 let upto = (from + per - 1).min(n - 1);
13908 plan.push((d, from, upto));
13909 from = upto + 1;
13910 }
13911 self.gpu_plan = Some(std::sync::Arc::new(plan));
13912 Ok(())
13913 }
13914
13915 pub fn gpu_plan(&self) -> Option<Vec<(usize, usize, usize)>> {
13917 self.gpu_plan.as_ref().map(|p| p.as_ref().clone())
13918 }
13919
13920 fn embryo_qtensor_f32(t: &QTensor) -> Vec<f32> {
13926 if let Some(x) = t.as_f32() {
13927 return x.to_vec();
13928 }
13929 let mut out = vec![0.0; t.rows() * t.cols()];
13930 for r in 0..t.rows() {
13931 t.row_f32(r, &mut out[r * t.cols()..(r + 1) * t.cols()]);
13932 }
13933 out
13934 }
13935
13936 fn embryo_resident_eligible(&self) -> bool {
13937 if (self.vmf_cfg.is_none() && self.gdn_cfg.is_none())
13940 || self.num_layers != self.physical_layers
13941 || self.loop_final_norm
13942 || self.weights.layers.len() != self.num_layers
13943 || self.head_clusters.is_none()
13944 || self.final_softcap.is_some()
13945 || self.logit_multiplier.is_some()
13946 || self.attn_softcap != 0.0
13947 || self.mtp.is_some()
13948 || self.g3n.is_some()
13949 || self.dsv4.is_some()
13950 || self.dsv41.is_some()
13951 || self.qwen4_exp.is_some()
13952 || self.dyn_router.is_some()
13959 || self.dyn_phi_layer.is_some()
13960 || self.dyn_blend_loaded
13961 || self.o1_cfg.is_some()
13962 || self.swa.is_some()
13963 || self.sliding_layers.is_some()
13964 || self.global_attn.is_some()
13965 || self.attention_heads_per_layer.is_some()
13966 || self.attn_v_norm
13967 || self
13968 .kv_cache
13969 .layers
13970 .iter()
13971 .any(|l| l.mode != crate::kv_cache::KvMode::F32)
13972 || self.rope_scale != 1.0
13973 || self.rope_scale_local != 1.0
13974 || self.attn_scale != 1.0 / (self.head_dim as f32).sqrt()
13975 || self.hidden_size == 0
13976 || self.hidden_size > 1024
13977 || self.intermediate_size > 1024
13978 || self.num_heads == 0
13979 || self.num_kv_heads == 0
13980 || self.num_heads % self.num_kv_heads != 0
13981 || self.num_heads.saturating_mul(self.head_dim) > 1024
13982 || self.num_kv_heads.saturating_mul(self.head_dim) > 1024
13983 || self.vocab_size == 0
13984 || self.kv_cache.max_seq_len == 0
13985 || self.rotary_dim == 0
13986 || self.rotary_dim > self.head_dim
13987 || self.rotary_dim % 2 != 0
13988 || self.inv_freq.len() < self.rotary_dim / 2
13989 {
13990 return false;
13991 }
13992 if self.weights.lm_head.as_f32().is_none()
14003 || self.weights.embed_tokens.as_f32().is_none()
14004 || self.weights.lm_head.rows() < self.vocab_size
14005 || self.weights.lm_head.cols() != self.hidden_size
14006 || self.weights.embed_tokens.rows() < self.vocab_size
14007 || self.weights.embed_tokens.cols() != self.hidden_size
14008 || self.weights.final_norm.len() != self.hidden_size
14009 {
14010 return false;
14011 }
14012 if let Some(cfg) = self.vmf_cfg {
14013 if cfg.num_heads.saturating_mul(cfg.nphase) > 1024
14014 || cfg.num_heads.saturating_mul(cfg.nphase.saturating_add(1)) > 1024
14015 || cfg.num_heads.saturating_mul(cfg.value_head_dim) > 1024
14016 || cfg.state_len() == 0
14017 {
14018 return false;
14019 }
14020 }
14021 if let Some(g) = self.gdn_cfg {
14022 if g.num_v_heads == 0
14027 || g.num_k_heads == 0
14028 || g.num_v_heads % g.num_k_heads != 0
14029 || g.key_head_dim == 0
14030 || g.key_head_dim > 128
14031 || g.value_head_dim == 0
14032 || g.value_head_dim > 256
14033 || g.value_head_dim % 4 != 0
14034 || g.conv_kernel == 0
14035 || g.num_v_heads > 512
14036 || g.num_v_heads.saturating_mul(g.value_head_dim) > 1024
14037 || g.conv_dim() > 2048
14038 || g.conv_dim() % 4 != 0
14039 || g.hidden_size != self.hidden_size
14040 || g.output_gate_sigmoid
14041 || g.rms_eps != self.rms_eps
14042 || g.state_len() == 0
14043 {
14044 return false;
14045 }
14046 }
14047 let mut full_seen = false;
14048 for lw in &self.weights.layers {
14049 if lw.attn_out_norm.is_some() || lw.ffn_out_norm.is_some() || lw.layer_scale.is_some() {
14050 return false;
14051 }
14052 match &lw.attn {
14053 AttnKind::LinearGdn(w) => {
14054 let Some(g) = self.gdn_cfg else {
14055 return false;
14056 };
14057 let (nv, dv, kk) = (g.num_v_heads, g.value_head_dim, g.conv_kernel);
14058 if w.in_proj_qkv.rows() != g.conv_dim()
14059 || w.in_proj_qkv.cols() != self.hidden_size
14060 || w.in_proj_qkv.as_f32().is_none()
14061 || w.in_proj_z.rows() != nv * dv
14062 || w.in_proj_z.cols() != self.hidden_size
14063 || w.in_proj_z.as_f32().is_none()
14064 || w.in_proj_a.rows() != nv
14065 || w.in_proj_a.cols() != self.hidden_size
14066 || w.in_proj_a.as_f32().is_none()
14067 || w.in_proj_b.rows() != nv
14068 || w.in_proj_b.cols() != self.hidden_size
14069 || w.in_proj_b.as_f32().is_none()
14070 || w.conv1d.len() != g.conv_dim() * kk
14071 || w.a_log.len() != nv
14072 || w.dt_bias.len() != nv
14073 || w.norm.len() != dv
14074 || w.out_proj.rows() != self.hidden_size
14075 || w.out_proj.cols() != nv * dv
14076 || w.out_proj.as_f32().is_none()
14077 {
14078 return false;
14079 }
14080 }
14081 AttnKind::Linear(w) => {
14082 let Some(cfg) = self.vmf_cfg else {
14083 return false;
14084 };
14085 if w.thq.rows() != cfg.num_heads * cfg.nphase
14086 || w.thq.cols() != self.hidden_size
14087 || w.thq.as_f32().is_none()
14088 || w.thk.rows() != cfg.num_heads * cfg.nphase
14089 || w.thk.cols() != self.hidden_size
14090 || w.thk.as_f32().is_none()
14091 || w.v_proj.rows() != cfg.num_heads * cfg.value_head_dim
14092 || w.v_proj.cols() != self.hidden_size
14093 || w.v_proj.as_f32().is_none()
14094 || w.out_proj.rows() != self.hidden_size
14095 || w.out_proj.cols() != cfg.num_heads * cfg.value_head_dim
14096 || w.out_proj.as_f32().is_none()
14097 || w.decay.len() != cfg.num_heads * 2 * cfg.nphase
14098 {
14099 return false;
14100 }
14101 if let Some((kg, kb)) = &w.k_gate {
14102 if kg.rows() != cfg.num_heads
14103 || kg.cols() != self.hidden_size
14104 || kg.as_f32().is_none()
14105 || kb.len() != cfg.num_heads
14106 {
14107 return false;
14108 }
14109 }
14110 if let Some(conv) = &w.conv {
14111 if conv.len() % self.hidden_size != 0 || conv.len() / self.hidden_size < 2 {
14112 return false;
14113 }
14114 }
14115 }
14116 AttnKind::Full {
14117 wq,
14118 wk,
14119 wv,
14120 wo,
14121 q_norm,
14122 k_norm,
14123 output_gate,
14124 softplus_gate,
14125 bias,
14126 } => {
14127 if full_seen
14128 || q_norm.is_some()
14129 || k_norm.is_some()
14130 || *output_gate
14131 || softplus_gate.is_some()
14132 || bias.is_some()
14133 || wq.as_f32().is_none()
14134 || wk.as_f32().is_none()
14135 || wv.as_f32().is_none()
14136 || wo.as_f32().is_none()
14137 || wq.rows() != self.num_heads * self.head_dim
14138 || wk.rows() != self.num_kv_heads * self.head_dim
14139 || wv.rows() != self.num_kv_heads * self.head_dim
14140 || wq.cols() != self.hidden_size
14141 || wk.cols() != self.hidden_size
14142 || wv.cols() != self.hidden_size
14143 || wo.rows() != self.hidden_size
14144 || wo.cols() != self.num_heads * self.head_dim
14145 {
14146 return false;
14147 }
14148 full_seen = true;
14149 }
14150 AttnKind::Bounded(w) => {
14151 let Some(ac) = self.anchor_core.as_ref() else {
14154 return false;
14155 };
14156 if self.bounded_rope.is_none()
14157 || w.window != ac.window
14158 || w.sink != ac.sink
14159 || w.window == 0
14160 || w.window + w.sink > 256
14161 || w.sink_k.len() != self.num_kv_heads * w.sink * self.head_dim
14162 || w.sink_v.len() != self.num_kv_heads * w.sink * self.head_dim
14163 || w.wq.as_f32().is_none()
14164 || w.wk.as_f32().is_none()
14165 || w.wv.as_f32().is_none()
14166 || w.wo.as_f32().is_none()
14167 || w.wq.rows() != self.num_heads * self.head_dim
14168 || w.wk.rows() != self.num_kv_heads * self.head_dim
14169 || w.wv.rows() != self.num_kv_heads * self.head_dim
14170 || w.wq.cols() != self.hidden_size
14171 || w.wk.cols() != self.hidden_size
14172 || w.wv.cols() != self.hidden_size
14173 || w.wo.rows() != self.hidden_size
14174 || w.wo.cols() != self.num_heads * self.head_dim
14175 {
14176 return false;
14177 }
14178 }
14179 _ => return false,
14180 }
14181 match &lw.ffn {
14182 FfnKind::Dense(d) => {
14183 if d.act != Act::Silu
14184 || !d.segs.is_empty()
14185 || d.gate_proj.as_f32().is_none()
14186 || d.up_proj.as_f32().is_none()
14187 || d.down_proj.as_f32().is_none()
14188 || d.gate_proj.rows() != self.intermediate_size
14189 || d.gate_proj.cols() != self.hidden_size
14190 || d.up_proj.rows() != self.intermediate_size
14191 || d.up_proj.cols() != self.hidden_size
14192 || d.down_proj.rows() != self.hidden_size
14193 || d.down_proj.cols() != self.intermediate_size
14194 {
14195 return false;
14196 }
14197 }
14198 FfnKind::Moe(m) => {
14199 if m.resonance.is_none()
14200 || m.top_k != 1
14201 || m.router_sigmoid
14202 || !m.norm_topk_prob
14203 || m.expert_bias.is_some()
14204 || m.routed_scaling != 1.0
14205 || m.route_tau.is_some()
14206 || m.shared.is_none()
14207 || m.mask.is_some()
14208 || m.per_expert_scale.is_some()
14209 || m.router_input_norm
14210 || m.experts.is_empty()
14211 || m.experts.len() > 8
14212 {
14213 return false;
14214 }
14215 let r = m.resonance.as_ref().unwrap();
14216 if r.mu.len() != m.experts.len() * self.hidden_size
14217 || r.bias.len() != m.experts.len()
14218 || r.u.len() != m.experts.len() * r.k * self.hidden_size
14219 || r.k > 128
14220 {
14221 return false;
14222 }
14223 let Some((shared, gate)) = &m.shared else {
14224 return false;
14225 };
14226 if gate.is_some() || shared.act != Act::Silu || !shared.segs.is_empty() {
14227 return false;
14228 }
14229 if shared.gate_proj.as_f32().is_none()
14230 || shared.up_proj.as_f32().is_none()
14231 || shared.down_proj.as_f32().is_none()
14232 || shared.gate_proj.rows() != self.intermediate_size
14233 || shared.gate_proj.cols() != self.hidden_size
14234 || shared.up_proj.rows() != self.intermediate_size
14235 || shared.up_proj.cols() != self.hidden_size
14236 || shared.down_proj.rows() != self.hidden_size
14237 || shared.down_proj.cols() != self.intermediate_size
14238 {
14239 return false;
14240 }
14241 for e in &m.experts {
14242 if e.act != Act::Silu
14243 || !e.segs.is_empty()
14244 || e.gate_proj.as_f32().is_none()
14245 || e.up_proj.as_f32().is_none()
14246 || e.down_proj.as_f32().is_none()
14247 || e.gate_proj.rows() != self.intermediate_size
14248 || e.gate_proj.cols() != self.hidden_size
14249 || e.up_proj.rows() != self.intermediate_size
14250 || e.up_proj.cols() != self.hidden_size
14251 || e.down_proj.rows() != self.hidden_size
14252 || e.down_proj.cols() != self.intermediate_size
14253 {
14254 return false;
14255 }
14256 }
14257 }
14258 FfnKind::DenseMoe(_) => return false,
14259 }
14260 if lw.input_norm.len() != self.hidden_size || lw.post_norm.len() != self.hidden_size {
14261 return false;
14262 }
14263 }
14264 if full_seen && self.anchor_core.is_some() {
14265 return false;
14266 }
14267 full_seen || self.num_layers > 0
14268 }
14269
14270 fn embryo_resident_wanted(&self) -> bool {
14275 crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Decode)
14276 && crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Prefill)
14277 && matches!(
14278 std::env::var("CMF_EMBRYO_RESIDENT").as_deref(),
14279 Ok("1") | Ok("parallel")
14280 )
14281 && crate::gpu::enabled_here()
14282 && !self.graph_refused()
14283 && self.embryo_resident_eligible()
14284 }
14285
14286 fn embryo_prefill_chunked(&mut self, ids: &[u32], start: usize) -> Option<Vec<f32>> {
14294 if ids.len() < 2
14295 || std::env::var("CMF_EMBRYO_CHUNK").as_deref() == Ok("0")
14296 || !self.embryo_resident_wanted()
14297 {
14298 return None;
14299 }
14300 let model = self.ensure_embryo_graph()?;
14301 let cmax = std::env::var("CMF_EMBRYO_CHUNK")
14302 .ok()
14303 .and_then(|v| v.parse::<usize>().ok())
14304 .filter(|&v| v >= 1)
14305 .unwrap_or(crate::gpu::EMBRYO_CHUNK_MAX)
14306 .min(crate::gpu::EMBRYO_CHUNK_MAX);
14307 let hs = self.hidden_size;
14308 let n = ids.len();
14309 let mut pos = start;
14310 let mut last = None;
14311 let mut rows = Vec::with_capacity(cmax * hs);
14312 while pos < n {
14313 let end = (pos + cmax).min(n);
14314 rows.clear();
14315 for &id in &ids[pos..end] {
14316 rows.extend_from_slice(&self.embed_single(id));
14317 }
14318 let mut lg = Vec::new();
14319 if !crate::gpu::forward_embryo_graph_chunk(
14320 &model,
14321 self.graph_kv_id,
14322 &rows,
14323 pos,
14324 end - pos,
14325 &mut lg,
14326 ) {
14327 if pos == start {
14328 if std::env::var("CMF_EMBRYO_DBG").as_deref() == Ok("1") {
14329 eprintln!("embryo-dbg: chunk prefill refused at position {pos}");
14330 }
14331 return None;
14332 }
14333 self.kv_cache.clear();
14337 self.clear_history();
14338 crate::gpu::graph_kv_reset(self.graph_kv_id);
14339 panic!(
14340 "resident Embryo chunk prefill refused at position {pos}; refusing an in-flight host fallback"
14341 );
14342 }
14343 last = Some(lg);
14344 pos = end;
14345 }
14346 if std::env::var("CMF_EMBRYO_DBG").as_deref() == Ok("1") {
14347 eprintln!(
14348 "embryo-dbg: chunk prefill {} ids from {start} in {} submits (chunk {cmax})",
14349 n - start,
14350 (n - start).div_ceil(cmax)
14351 );
14352 }
14353 last
14354 }
14355
14356 fn ensure_embryo_graph(&mut self) -> Option<std::sync::Arc<crate::gpu::EmbryoGraphModel>> {
14357 if self.embryo_graph.is_none() && self.embryo_resident_eligible() {
14358 const UMAX: u32 = u32::MAX;
14359 const HEADER: usize = crate::gpu::EMBRYO_META_HEADER;
14360 const REC: usize = 64;
14361 struct Pack {
14362 data: Vec<f32>,
14363 }
14364 impl Pack {
14365 fn put(&mut self, x: &[f32]) -> u32 {
14366 if x.is_empty() {
14367 return u32::MAX;
14368 }
14369 let off = self.data.len();
14370 self.data.extend_from_slice(x);
14371 off as u32
14372 }
14373 }
14374 let vmf = self.vmf_cfg;
14378 let gdn = self.gdn_cfg;
14379 let mut pack = Pack { data: Vec::new() };
14380 let mut meta = vec![0u32; HEADER];
14381 meta[0] = self.hidden_size as u32;
14382 meta[1] = self.intermediate_size as u32;
14383 meta[2] = self.vocab_size as u32;
14384 meta[3] = self.num_layers as u32;
14385 meta[4] = vmf.map(|c| c.num_heads).unwrap_or(0) as u32;
14386 meta[5] = vmf.map(|c| c.nphase).unwrap_or(0) as u32;
14387 meta[6] = vmf.map(|c| c.value_head_dim).unwrap_or(0) as u32;
14388 if let Some(g) = gdn {
14389 meta[24] = g.num_v_heads as u32;
14390 meta[25] = g.num_k_heads as u32;
14391 meta[26] = g.key_head_dim as u32;
14392 meta[27] = g.value_head_dim as u32;
14393 meta[28] = g.conv_kernel as u32;
14394 meta[29] = g.conv_dim() as u32;
14395 }
14396 meta[7] = self.num_heads as u32;
14397 meta[8] = self.num_kv_heads as u32;
14398 meta[9] = self.head_dim as u32;
14399 meta[10] = self.kv_cache.max_seq_len as u32;
14400 let clusters = self.head_clusters.as_ref().unwrap();
14401 let cluster_count = clusters.len() / self.hidden_size;
14402 if clusters.len() % self.hidden_size != 0
14403 || cluster_count == 0
14404 || cluster_count > 1024
14405 || self.vocab_size % cluster_count != 0
14406 || self.weights.lm_head.rows() < self.vocab_size
14407 || self.weights.final_norm.len() != self.hidden_size
14408 {
14409 return None;
14410 }
14411 meta[11] = cluster_count as u32;
14412 meta[12] = (self.vocab_size / cluster_count) as u32;
14413 meta[13] = vmf.map(|c| c.state_len()).unwrap_or(0) as u32;
14414 meta[16] = self.rotary_dim as u32;
14415 meta[17] = matches!(self.norm_style, NormStyle::Gemma) as u32;
14416 meta[19] = (self.rms_eps as f32).to_bits();
14417 let max_conv = self
14418 .weights
14419 .layers
14420 .iter()
14421 .filter_map(|lw| match &lw.attn {
14422 AttnKind::Linear(w) => w.conv.as_ref().map(|v| v.len() / self.hidden_size),
14423 _ => None,
14424 })
14425 .max()
14426 .unwrap_or(1);
14427 let phase_stride = vmf
14433 .map(|c| c.state_len() + max_conv.saturating_sub(1) * self.hidden_size)
14434 .unwrap_or(0);
14435 let gdn_stride = gdn.map(|g| g.state_len()).unwrap_or(0);
14436 let state_stride = phase_stride.max(gdn_stride);
14437 let bounded = self.anchor_core.clone();
14441 let (anchor_window, anchor_sink) = bounded
14442 .as_ref()
14443 .map(|ac| (ac.window, ac.sink))
14444 .unwrap_or((0, 0));
14445 let kv_stride = if bounded.is_some() {
14446 2usize
14447 .saturating_mul(self.num_kv_heads)
14448 .saturating_mul(anchor_window)
14449 .saturating_mul(self.head_dim)
14450 } else {
14451 2usize
14452 .saturating_mul(self.num_kv_heads)
14453 .saturating_mul(self.kv_cache.max_seq_len)
14454 .saturating_mul(self.head_dim)
14455 };
14456 meta[14] = state_stride as u32;
14457 meta[15] = kv_stride as u32;
14458 meta[18] = anchor_window as u32;
14459 meta[20] = anchor_sink as u32;
14460 meta[21] = match &self.bounded_rope {
14461 Some(rope) => {
14462 let off = pack.put(&rope.cos);
14464 let _ = pack.put(&rope.sin);
14465 off
14466 }
14467 None => UMAX,
14468 };
14469 let mut full_seen = false;
14470 let mut bounded_seen = 0usize;
14471 let mut phase_seen = 0usize;
14475 let mut gdn_seen = 0usize;
14476 for (li, lw) in self.weights.layers.iter().enumerate() {
14477 let base = meta.len();
14478 meta.resize(base + REC, UMAX);
14479 meta[base] = match &lw.attn {
14480 AttnKind::Linear(w) if w.phase_delta => 1,
14481 AttnKind::Linear(_) => 0,
14482 AttnKind::Full { .. } => 2,
14483 AttnKind::Bounded(_) => 3,
14484 AttnKind::LinearGdn(_) => 4,
14485 _ => UMAX,
14486 };
14487 meta[base + 1] = pack.put(&lw.input_norm);
14488 meta[base + 2] = pack.put(&lw.post_norm);
14489 meta[base + 25] = match &lw.attn {
14490 AttnKind::Linear(_) => {
14491 let off = ((phase_seen + gdn_seen) * state_stride) as u32;
14492 phase_seen += 1;
14493 off
14494 }
14495 AttnKind::LinearGdn(_) => {
14496 let off = ((phase_seen + gdn_seen) * state_stride) as u32;
14497 gdn_seen += 1;
14498 off
14499 }
14500 _ => UMAX,
14501 };
14502 match &lw.attn {
14503 AttnKind::LinearGdn(w) => {
14504 meta[base + 56] = pack.put(&Self::embryo_qtensor_f32(&w.in_proj_qkv));
14507 meta[base + 57] = pack.put(&Self::embryo_qtensor_f32(&w.in_proj_z));
14508 meta[base + 58] = pack.put(&Self::embryo_qtensor_f32(&w.in_proj_a));
14509 meta[base + 59] = pack.put(&Self::embryo_qtensor_f32(&w.in_proj_b));
14510 meta[base + 60] = pack.put(&w.conv1d);
14511 meta[base + 61] = pack.put(&w.a_log);
14512 meta[base + 62] = pack.put(&w.dt_bias);
14513 meta[base + 63] = pack.put(&w.norm);
14514 meta[base + 29] = pack.put(&Self::embryo_qtensor_f32(&w.out_proj));
14515 meta[base + 24] = 0;
14516 }
14517 AttnKind::Linear(w) => {
14518 meta[base + 3] = pack.put(&Self::embryo_qtensor_f32(&w.thq));
14519 meta[base + 4] = pack.put(&Self::embryo_qtensor_f32(&w.thk));
14520 meta[base + 5] = pack.put(&Self::embryo_qtensor_f32(&w.v_proj));
14521 meta[base + 6] = pack.put(&Self::embryo_qtensor_f32(&w.out_proj));
14522 let decay: Vec<f32> = w.decay.iter().map(|&x| x as f32).collect();
14523 meta[base + 7] = pack.put(&decay);
14524 if let Some((kg, kb)) = &w.k_gate {
14525 meta[base + 8] = pack.put(&Self::embryo_qtensor_f32(kg));
14526 meta[base + 9] = pack.put(kb);
14527 }
14528 if let Some(conv) = &w.conv {
14529 meta[base + 10] = pack.put(conv);
14530 meta[base + 24] = (conv.len() / self.hidden_size) as u32;
14531 } else {
14532 meta[base + 24] = 0;
14533 }
14534 }
14535 AttnKind::Full { wq, wk, wv, wo, .. } => {
14536 full_seen = true;
14537 meta[base + 11] = pack.put(&Self::embryo_qtensor_f32(wq));
14538 meta[base + 12] = pack.put(&Self::embryo_qtensor_f32(wk));
14539 meta[base + 13] = pack.put(&Self::embryo_qtensor_f32(wv));
14540 meta[base + 14] = pack.put(&Self::embryo_qtensor_f32(wo));
14541 meta[base + 26] = (li * kv_stride) as u32;
14542 }
14543 AttnKind::Bounded(w) => {
14544 meta[base + 11] = pack.put(&Self::embryo_qtensor_f32(&w.wq));
14545 meta[base + 12] = pack.put(&Self::embryo_qtensor_f32(&w.wk));
14546 meta[base + 13] = pack.put(&Self::embryo_qtensor_f32(&w.wv));
14547 meta[base + 14] = pack.put(&Self::embryo_qtensor_f32(&w.wo));
14548 meta[base + 26] = (bounded_seen * kv_stride) as u32;
14550 meta[base + 27] = pack.put(&w.sink_k);
14551 meta[base + 28] = pack.put(&w.sink_v);
14552 bounded_seen += 1;
14553 }
14554 _ => return None,
14555 }
14556 match &lw.ffn {
14557 FfnKind::Dense(d) => {
14558 meta[base + 15] = 0;
14559 meta[base + 16] = 0;
14560 meta[base + 21] = pack.put(&Self::embryo_qtensor_f32(&d.gate_proj));
14561 meta[base + 22] = pack.put(&Self::embryo_qtensor_f32(&d.up_proj));
14562 meta[base + 23] = pack.put(&Self::embryo_qtensor_f32(&d.down_proj));
14563 }
14564 FfnKind::Moe(m) => {
14565 let r = m.resonance.as_ref().unwrap();
14566 let (shared, _) = m.shared.as_ref().unwrap();
14567 meta[base + 15] = 1;
14568 meta[base + 16] = m.experts.len() as u32;
14569 meta[base + 17] = pack.put(&r.mu);
14570 meta[base + 18] = pack.put(&r.u);
14571 meta[base + 19] = pack.put(&r.bias);
14572 meta[base + 20] = r.k as u32;
14573 let mut shell = r.effective_shell(m.experts.len());
14582 shell.push(f32::NEG_INFINITY);
14583 meta[base + 30] = pack.put(&shell);
14584 meta[base + 21] = pack.put(&Self::embryo_qtensor_f32(&shared.gate_proj));
14585 meta[base + 22] = pack.put(&Self::embryo_qtensor_f32(&shared.up_proj));
14586 meta[base + 23] = pack.put(&Self::embryo_qtensor_f32(&shared.down_proj));
14587 for (e, ex) in m.experts.iter().enumerate() {
14588 meta[base + 32 + e * 3] =
14589 pack.put(&Self::embryo_qtensor_f32(&ex.gate_proj));
14590 meta[base + 33 + e * 3] =
14591 pack.put(&Self::embryo_qtensor_f32(&ex.up_proj));
14592 meta[base + 34 + e * 3] =
14593 pack.put(&Self::embryo_qtensor_f32(&ex.down_proj));
14594 }
14595 }
14596 FfnKind::DenseMoe(_) => return None,
14597 }
14598 }
14599 if !full_seen && self.num_layers == 0 {
14600 return None;
14601 }
14602 let id = {
14603 static NEXT: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(1);
14604 NEXT.fetch_add(1, std::sync::atomic::Ordering::Relaxed)
14605 };
14606 let model = crate::gpu::EmbryoGraphModel {
14607 id,
14608 hidden: self.hidden_size,
14609 intermediate: self.intermediate_size,
14610 vocab: self.vocab_size,
14611 layers: self.num_layers,
14612 phase_heads: vmf.map(|c| c.num_heads).unwrap_or(0),
14613 nphase: vmf.map(|c| c.nphase).unwrap_or(0),
14614 phase_dv: vmf.map(|c| c.value_head_dim).unwrap_or(0),
14615 anchor_q_heads: self.num_heads,
14616 anchor_kv_heads: self.num_kv_heads,
14617 anchor_head_dim: self.head_dim,
14618 rotary_dim: self.rotary_dim,
14619 max_seq: self.kv_cache.max_seq_len,
14620 cluster_count,
14621 cluster_size: self.vocab_size / cluster_count,
14622 phase_state_len: vmf.map(|c| c.state_len()).unwrap_or(0),
14623 state_stride,
14624 kv_stride,
14625 norm_gemma: matches!(self.norm_style, NormStyle::Gemma),
14626 phase_mass: vmf.map(|c| c.phase_mass).unwrap_or(0.0),
14627 weights: pack.data,
14628 meta,
14629 lm_head: Self::embryo_qtensor_f32(&self.weights.lm_head),
14630 clusters: clusters.as_ref().clone(),
14631 final_norm: self.weights.final_norm.clone(),
14632 inv_freq: self.inv_freq.as_ref().clone(),
14633 bounded: bounded.is_some(),
14634 kv_layers: if bounded.is_some() {
14635 bounded_seen
14636 } else {
14637 self.num_layers
14638 },
14639 state_layers: phase_seen + gdn_seen,
14640 anchor_window,
14641 anchor_sink,
14642 phase_layers: phase_seen,
14643 gdn_layers: gdn_seen,
14644 gdn_heads: gdn.map(|g| g.num_v_heads).unwrap_or(0),
14645 gdn_k_heads: gdn.map(|g| g.num_k_heads).unwrap_or(0),
14646 gdn_dk: gdn.map(|g| g.key_head_dim).unwrap_or(0),
14647 gdn_dv: gdn.map(|g| g.value_head_dim).unwrap_or(0),
14648 gdn_kk: gdn.map(|g| g.conv_kernel).unwrap_or(0),
14649 };
14650 self.embryo_graph = Some(std::sync::Arc::new(model));
14651 }
14652 self.embryo_graph.clone()
14653 }
14654
14655 fn forward_layers_span(
14656 &mut self,
14657 hidden: &[f32],
14658 position: usize,
14659 task_mask: Option<&TaskMask>,
14660 from: usize,
14661 upto: Option<usize>,
14662 ) -> Vec<f32> {
14663 debug_assert!(
14664 from == 0
14665 || (self.dsv4.is_none()
14666 && self.dsv41.is_none()
14667 && self.qwen4_exp.is_none()
14668 && self.g3n.is_none())
14669 );
14670 #[cfg(target_os = "macos")]
14676 if !crate::gpu_metal::wait_replay() {
14677 self.fail_metal_graph("the pending async replay failed before a plain forward");
14678 return vec![0.0; self.hidden_size];
14679 }
14680 if let Some(b) = &mut self.qwen4_exp {
14681 let _ = (task_mask, upto);
14682 let token_id = hidden.first().copied().unwrap_or(0.0) as u32;
14683 let mut logits = Vec::new();
14684 crate::qwen4_exp::forward_token(
14685 &b.0,
14686 &b.1,
14687 &b.2,
14688 &mut b.3,
14689 token_id,
14690 position,
14691 &self.inv_freq,
14692 self.pool.as_deref(),
14693 &mut logits,
14694 true,
14695 );
14696 self.graph_logits = Some(logits);
14697 return vec![0.0; self.hidden_size];
14698 }
14699 if let Some(b) = &mut self.dsv4 {
14705 let _ = (task_mask, upto);
14706 let token_id = hidden.first().copied().unwrap_or(0.0) as u32;
14707 let (g, layers, cfg, st) = (&b.0, &b.1, b.2, &mut b.3);
14708 st.pos = position;
14709 let mut logits = Vec::new();
14710 crate::dsv4::forward_token(
14711 g,
14712 layers,
14713 &cfg,
14714 st,
14715 token_id,
14716 &self.inv_freq,
14717 self.pool.as_deref(),
14718 &mut logits,
14719 );
14720 self.graph_logits = Some(logits);
14721 self.dspark_probe(position, token_id);
14722 return vec![0.0; self.hidden_size];
14725 }
14726 if let Some(b) = &mut self.dsv41 {
14728 let _ = (task_mask, upto);
14729 let token_id = hidden.first().copied().unwrap_or(0.0) as u32;
14730 let mut logits = Vec::new();
14731 crate::dsv41::forward_token(
14732 &b.0,
14733 &b.1,
14734 &b.2,
14735 &mut b.3,
14736 token_id,
14737 position,
14738 self.pool.as_deref(),
14739 &mut logits,
14740 );
14741 self.graph_logits = Some(logits);
14742 return vec![0.0; self.hidden_size];
14743 }
14744 if let Some(b) = &self.g3n {
14747 let _ = (task_mask, upto);
14748 return crate::g3n::g3n_forward(
14749 &b.0,
14750 &b.1,
14751 hidden,
14752 position,
14753 &mut self.kv_cache.layers,
14754 self.num_heads,
14755 self.num_kv_heads,
14756 self.head_dim,
14757 self.pool.as_deref(),
14758 );
14759 }
14760 if from == 0
14767 && upto.is_none()
14768 && task_mask.is_none()
14769 && self.anchor_core.is_some()
14770 && std::env::var("CMF_EMBRYO_DBG").as_deref() == Ok("1")
14771 {
14772 static ONCE: std::sync::Once = std::sync::Once::new();
14773 ONCE.call_once(|| {
14774 eprintln!(
14775 "embryo-dbg: graph_on decode={} prefill={} resident_env={:?} enabled_here={} \
14776 unsupported={} eligible={}",
14777 crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Decode),
14778 crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Prefill),
14779 std::env::var("CMF_EMBRYO_RESIDENT").ok(),
14780 crate::gpu::enabled_here(),
14781 self.graph_refused(),
14782 self.embryo_resident_eligible(),
14783 );
14784 });
14785 }
14786 if from == 0
14787 && upto.is_none()
14788 && task_mask.is_none()
14789 && crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Decode)
14793 && crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Prefill)
14794 && matches!(
14798 std::env::var("CMF_EMBRYO_RESIDENT").as_deref(),
14799 Ok("1") | Ok("parallel")
14800 )
14801 && crate::gpu::enabled_here()
14802 && !self.graph_refused()
14803 && (position == 0 || self.device_sequence_position().is_some())
14809 && self.embryo_resident_eligible()
14810 && let Some(model) = self.ensure_embryo_graph()
14811 {
14812 let mut lg = Vec::new();
14813 if crate::gpu::forward_embryo_graph(&model, self.graph_kv_id, hidden, position, &mut lg)
14814 {
14815 self.graph_logits = Some(lg);
14816 return vec![0.0; self.hidden_size];
14817 }
14818 if std::env::var("CMF_EMBRYO_DBG").as_deref() == Ok("1") {
14819 eprintln!("embryo-dbg: forward_embryo_graph refused at position {position}");
14820 }
14821 self.mark_graph_refused();
14827 if position != 0 {
14828 self.kv_cache.clear();
14832 self.clear_history();
14833 crate::gpu::graph_kv_reset(self.graph_kv_id);
14834 panic!(
14835 "resident Embryo graph refused at position {position}; refusing an in-flight host fallback"
14836 );
14837 }
14838 }
14839 let mut h = hidden.to_vec();
14840 self.mimo_moe_prepare();
14843 let _mimo_q8 = self.mimo_moe.is_on()
14844 .then(crate::qtensor::enter_full_gpu_q8_scope);
14845 let (nh, _nkv, _hd, hs, _rd, eps) = (
14848 self.num_heads,
14849 self.num_kv_heads,
14850 self.head_dim,
14851 self.hidden_size,
14852 self.rotary_dim,
14853 self.rms_eps,
14854 );
14855 let pool = self.pool.clone();
14856 let graph_env = std::env::var("CMF_GPU_WGPU_GRAPH").ok();
14868 let graph_on = match graph_env.as_deref() {
14869 Some("0") => false,
14870 Some("prefill") => false, Some(_) => true,
14872 None => crate::gpu::wgpu_graph_default(),
14878 };
14879 let graph_trusted =
14880 graph_env.is_some() || crate::gpu::wgpu_graph_default() || self.gdn_cfg.is_some();
14881 let race_eligible = graph_on
14882 && upto.is_none()
14883 && task_mask.is_none()
14884 && from == 0
14885 && !self.graph_refused();
14886 let mut tail_start = 0usize;
14887 if race_eligible && crate::gpu::graph_race_use_graph(graph_trusted) {
14888 let t_graph = std::time::Instant::now();
14889 let mut lg = Vec::new();
14890 let mut gl = 0usize;
14891 let built = self.try_token_graph_wgpu(hidden, position, &mut lg, &mut gl);
14892 let declined = built.is_none();
14893 let built = match built {
14894 Some(Ok(hh)) => Some(hh),
14895 Some(Err(())) => {
14896 self.clear_sequence_state();
14900 self.graph_failed
14901 .store(true, std::sync::atomic::Ordering::Relaxed);
14902 self.cancel
14903 .store(true, std::sync::atomic::Ordering::Relaxed);
14904 tracing::error!("token graph failed after admission; sequence state cleared");
14905 return vec![0.0; self.hidden_size];
14906 }
14907 None => None,
14908 };
14909 if declined && !self.o1_active() && self.attn_softcap == 0.0 {
14914 self.mark_graph_refused();
14915 }
14916 graph_note(built.is_some(), gl, self.num_layers);
14917 if let Some(hh) = built {
14918 let dur = t_graph.elapsed();
14919 if std::env::var("CMF_GRAPH_PROF").is_ok() {
14920 eprintln!("graph-call: {:.2} ms total", dur.as_secs_f64() * 1000.0);
14921 }
14922 if gl > 0 && gl < self.num_layers {
14923 h = hh;
14929 tail_start = gl;
14930 } else if graph_trusted || !crate::gpu::graph_race_first_token_hopeless(dur) {
14931 if !graph_trusted {
14932 crate::gpu::graph_race_record(true, dur);
14933 }
14934 if !lg.is_empty() {
14935 lg.resize(self.vocab_size, 0.0);
14938 if let Some(c) = self.final_softcap {
14939 for l in lg.iter_mut() {
14940 *l = c * (*l / c).tanh();
14941 }
14942 }
14943 self.graph_logits = Some(lg);
14944 }
14945 return hh;
14946 }
14947 }
14953 }
14954 let span = from > 0 || upto.is_some();
14978 if span && graph_on && task_mask.is_none() && graph_trusted {
14979 let upto_excl = upto.map_or(self.num_layers, |u| u + 1);
14980 let mut lg = Vec::new();
14981 let mut gl = 0usize;
14982 let span_res =
14983 self.try_token_graph_wgpu_span(hidden, position, &mut lg, from, upto_excl, &mut gl);
14984 let span_res = match span_res {
14985 Some(Ok(hh)) => Some(hh),
14986 Some(Err(())) => {
14987 self.clear_sequence_state();
14988 self.graph_failed
14989 .store(true, std::sync::atomic::Ordering::Relaxed);
14990 self.cancel
14991 .store(true, std::sync::atomic::Ordering::Relaxed);
14992 tracing::error!(
14993 "span token graph failed after admission; sequence state cleared"
14994 );
14995 return vec![0.0; self.hidden_size];
14996 }
14997 None => None,
14998 };
14999 graph_note(span_res.is_some(), gl, upto_excl - from);
15000 if std::env::var("CMF_GPU_DEBUG").is_ok() {
15001 static SEEN: std::sync::atomic::AtomicU32 = std::sync::atomic::AtomicU32::new(0);
15005 if SEEN.fetch_add(1, std::sync::atomic::Ordering::Relaxed) < 4 {
15006 eprintln!(
15007 "span graph: covered {gl} of {} layers [{from}..{upto_excl}) res={}",
15008 upto_excl - from,
15009 span_res.is_some()
15010 );
15011 }
15012 }
15013 if let Some(hh) = span_res {
15014 if gl == upto_excl - from {
15015 if !lg.is_empty() {
15016 lg.resize(self.vocab_size, 0.0);
15017 if let Some(c) = self.final_softcap {
15018 for l in lg.iter_mut() {
15019 *l = c * (*l / c).tanh();
15020 }
15021 }
15022 self.graph_logits = Some(lg);
15023 }
15024 crate::gpu::set_layer(-1);
15025 return hh;
15026 }
15027 h = hh;
15029 tail_start = from + gl;
15030 }
15031 }
15032 let t_race_cpu = (race_eligible && !graph_trusted).then(std::time::Instant::now);
15037
15038 let host_tail = tail_start > from;
15048 let _host_tail = (host_tail && !self.mimo_moe.is_on()).then(crate::gpu::enter_cpu_scope);
15049 let automatic_gpu_prefix = self.automatic_gpu_prefix();
15050
15051 let _prof_layers = crate::cpuprof::time(crate::cpuprof::Slot::Layers);
15052 #[cfg(target_os = "macos")]
15053 let mut gpu_skip_until = 0usize;
15054 for li in tail_start.max(from)..self.num_layers {
15055 let _capacity_tail = automatic_gpu_prefix
15056 .filter(|&prefix| li >= prefix && !self.mimo_moe.is_dynamic(li, host_tail))
15057 .map(|_| crate::gpu::enter_cpu_scope());
15058 crate::gpu::set_layer(li as i64); if let Some(u) = upto {
15060 if li > u {
15061 break;
15062 }
15063 }
15064 if let Some(mask) = task_mask {
15065 if !mask.layer_alive(li) {
15066 continue; }
15068 }
15069 #[cfg(target_os = "macos")]
15073 {
15074 if li < gpu_skip_until {
15075 continue;
15076 }
15077 if task_mask.is_none() {
15078 let end = self.q1_graph_gpu(li, upto, position, &mut h);
15079 if self
15080 .graph_failed
15081 .load(std::sync::atomic::Ordering::Relaxed)
15082 {
15083 return vec![0.0; self.hidden_size];
15087 }
15088 if end > li {
15089 gpu_skip_until = end;
15090 if self.is_loop_end(end - 1) && end < self.num_layers {
15093 h = inference::rms_norm(
15094 &h,
15095 &self.weights.final_norm,
15096 self.rms_eps,
15097 self.norm_style,
15098 );
15099 }
15100 continue;
15101 }
15102 }
15103 }
15104
15105 if task_mask.is_none() {
15106 match self.mimo_graph_layer_rows(li, &mut h, &[position]) {
15107 crate::gpu::BatchGraphOutcome::Completed => continue,
15108 crate::gpu::BatchGraphOutcome::Failed => return vec![0.0; self.hidden_size],
15109 crate::gpu::BatchGraphOutcome::Declined => {},
15110 }
15111 }
15112 #[cfg(feature = "gpu")]
15113 self.pull_lagging_host_kv(li, li + 1, position);
15114 let lw = &self.weights.layers[self.phys_layer(li)];
15115 if let Ok(tp) = std::env::var("CMF_TRACE_POS") {
15116 if tp.parse::<usize>().ok() == Some(position) {
15117 let n: f32 = h.iter().map(|x| x * x).sum::<f32>().sqrt();
15118 eprintln!(
15119 "TRACE pos {position} layer {li}: |h| = {n:.6} h0 {:.6} h1 {:.6}",
15120 h[0], h[1]
15121 );
15122 }
15123 }
15124 let prof = crate::cpuprof::time(crate::cpuprof::Slot::Norms);
15127 inference::rms_norm_into(
15128 &h,
15129 &lw.input_norm,
15130 self.rms_eps,
15131 self.norm_style,
15132 &mut self.ws.n1,
15133 );
15134 drop(prof);
15135
15136 let attn_out = match &lw.attn {
15137 AttnKind::Mla(w) => {
15138 let inv_freq_l = self.layer_inv_freq(li);
15139 let rs = self.layer_rope_scale(li);
15140 let eps = self.rms_eps;
15141 let pool = self.pool.clone();
15142 mla_attention(
15143 w,
15144 &self.ws.n1,
15145 &mut self.kv_cache.layers[li],
15146 position,
15147 &inv_freq_l,
15148 rs,
15149 eps,
15150 pool.as_deref(),
15151 )
15152 }
15153 AttnKind::Linear(w) => {
15154 let cfg = self.vmf_cfg.expect("linear layer without vmf_cfg");
15155 vmf_phase_forward(
15156 &self.ws.n1,
15157 w,
15158 &cfg,
15159 &mut self.kv_cache.layers[li].linear_state,
15160 self.pool.as_deref(),
15161 )
15162 }
15163 AttnKind::Kda(w) => {
15164 let cfg = self.kda_cfg.expect("kda layer without kda_cfg");
15165 crate::linear_core::kda_forward(
15166 &self.ws.n1,
15167 w,
15168 &cfg,
15169 &mut self.kv_cache.layers[li].linear_state,
15170 self.pool.as_deref(),
15171 )
15172 }
15173 AttnKind::LinearGdn(w) => {
15174 let cfg = self.gdn_cfg.expect("gdn layer without gdn_cfg");
15175 gdn_forward(
15176 &self.ws.n1,
15177 w,
15178 &cfg,
15179 &mut self.kv_cache.layers[li].linear_state,
15180 self.pool.as_deref(),
15181 )
15182 }
15183 AttnKind::ShortConv(w) => {
15184 let cfg = self
15185 .short_conv_cfg
15186 .expect("short-conv layer without short_conv_cfg");
15187 short_conv_forward(
15188 &self.ws.n1,
15189 w,
15190 &cfg,
15191 &mut self.kv_cache.layers[li].linear_state,
15192 self.pool.as_deref(),
15193 )
15194 }
15195 AttnKind::Bounded(w) => {
15196 let rope = self
15199 .bounded_rope
15200 .clone()
15201 .expect("bounded layer without an installed rotation table");
15202 let cfg = crate::bounded::BoundedAttnCfg {
15203 num_heads: self.num_heads,
15204 num_kv_heads: self.num_kv_heads,
15205 head_dim: self.head_dim,
15206 hidden_size: hs,
15207 scale: self.attn_scale,
15208 rope: &rope,
15209 pool: pool.as_deref(),
15210 };
15211 crate::bounded::bounded_attention(
15212 &self.ws.n1,
15213 w,
15214 &mut self.kv_cache.layers[li],
15215 &cfg,
15216 )
15217 }
15218 AttnKind::Full {
15219 wq,
15220 wk,
15221 wv,
15222 wo,
15223 q_norm,
15224 k_norm,
15225 output_gate,
15226 softplus_gate,
15227 bias,
15228 } if self.kv_cache.layers[li].o1_sealed() => {
15229 let inv_freq_l = self.layer_inv_freq(li);
15232 let (nkv_l, hd_l, rd_l) = self.layer_geom(li);
15233 let cfg = QwenAttnCfg {
15234 num_heads: self.layer_num_heads(li),
15235 num_kv_heads: nkv_l,
15236 head_dim: hd_l,
15237 hidden_size: hs,
15238 position,
15239 inv_freq: &inv_freq_l,
15240 rotary_dim: rd_l,
15241 scale: self.attn_scale,
15242 softcap: self.attn_softcap,
15243 window: None,
15244 v_norm: self.attn_v_norm,
15245 qk_norm_after_rope: self.qk_norm_after_rope,
15246 gate_sigmoid: self.proj_gate_sigmoid,
15247 q_norm: q_norm.as_deref(),
15248 k_norm: k_norm.as_deref(),
15249 output_gate: *output_gate,
15250 softplus_gate: softplus_gate
15251 .as_ref()
15252 .map(|(gate, per_head)| (gate, *per_head)),
15253 rope_scale: self.layer_rope_scale(li),
15254 bias: bias
15255 .as_ref()
15256 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
15257 rms_eps: eps,
15258 norm_style: self.norm_style,
15259 pool: pool.as_deref(),
15260 v_head_dim: self.layer_v_dim(li),
15261 };
15262 attention::qwen_attention_nystrom(
15263 &self.ws.n1,
15264 wq,
15265 wk,
15266 wv,
15267 wo,
15268 &mut self.kv_cache.layers[li],
15269 &cfg,
15270 )
15271 }
15272 AttnKind::Full {
15273 wq,
15274 wk,
15275 wv,
15276 wo,
15277 q_norm,
15278 k_norm,
15279 output_gate,
15280 softplus_gate,
15281 bias,
15282 } => 'attn: {
15283 let dropin_reason =
15288 graph_on.then(|| self.graph_attn_decline_reason()).flatten();
15289 if let Some(reason) = dropin_reason {
15290 self.note_graph_decline("wgpu attn dropin", reason);
15291 }
15292 if graph_on
15293 && dropin_reason.is_none()
15294 && !*output_gate
15295 && softplus_gate.is_none()
15296 && self.attention_heads_per_layer.is_none()
15297 && bias.is_none()
15298 && task_mask.is_none()
15299 {
15300 let inv_freq_l = self.layer_inv_freq(li);
15301 let (nkv_l, hd_l, rd_l) = self.layer_geom(li);
15302 let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
15303 if let (Some((gm, qi)), Some((_, ki)), Some((_, vi)), Some((_, oi))) = (
15304 wq.mapped_q1(),
15305 wk.mapped_q1(),
15306 wv.mapped_q1(),
15307 wo.mapped_q1(),
15308 ) {
15309 let gm = gm.clone();
15310 let mut out = vec![0f32; hs];
15311 let cache = &self.kv_cache.layers[li];
15312 if crate::gpu::attn_dropin(
15313 &gm,
15314 self.graph_kv_id,
15315 li,
15316 &self.ws.n1,
15317 qi,
15318 ki,
15319 vi,
15320 oi,
15321 q_norm.as_deref(),
15322 k_norm.as_deref(),
15323 self.qk_norm_after_rope,
15324 &inv_freq_l,
15325 nh,
15326 nkv_l,
15327 hd_l,
15328 rd_l,
15329 hs,
15330 position,
15331 self.kv_cache.max_seq_len,
15332 gemma,
15333 eps as f32,
15334 cache.k_heads(),
15335 cache.v_heads(),
15336 &mut out,
15337 ) {
15338 break 'attn out;
15339 }
15340 }
15341 }
15342 let masked = task_mask
15343 .map(|m| m.head_flags(li, self.num_heads).iter().any(|&a| !a))
15344 .unwrap_or(false);
15345 let f32_view = (wq.as_f32(), wk.as_f32(), wv.as_f32(), wo.as_f32());
15346 let plain = self.layer_attn_plain(li);
15349 match (masked, f32_view) {
15350 (true, (Some(q), Some(k), Some(v), Some(o))) if plain => {
15353 let active_heads = task_mask.unwrap().head_flags(li, self.num_heads);
15354 attention::multi_head_attention(
15355 &self.ws.n1,
15356 q,
15357 k,
15358 v,
15359 o,
15360 &mut self.kv_cache.layers[li],
15361 self.num_heads,
15362 self.num_kv_heads,
15363 self.head_dim,
15364 self.hidden_size,
15365 position,
15366 &active_heads,
15367 &self.inv_freq,
15368 )
15369 }
15370 (masked, _) => {
15371 if masked {
15372 tracing::warn!(
15373 "layer {li}: head mask on quantized weights or on a \
15374 window/sink/per-layer-geometry layer not supported \
15375 yet — executing dense"
15376 );
15377 }
15378 let inv_freq_l = self.layer_inv_freq(li);
15379 let (nkv_l, hd_l, rd_l) = self.layer_geom(li);
15380 let cfg = QwenAttnCfg {
15381 num_heads: self.layer_num_heads(li),
15382 num_kv_heads: nkv_l,
15383 head_dim: hd_l,
15384 hidden_size: hs,
15385 position,
15386 inv_freq: &inv_freq_l,
15387 rotary_dim: rd_l,
15388 scale: self.attn_scale,
15389 softcap: self.attn_softcap,
15390 window: self.layer_window(li),
15391 v_norm: self.attn_v_norm,
15392 qk_norm_after_rope: self.qk_norm_after_rope,
15393 gate_sigmoid: self.proj_gate_sigmoid,
15394 q_norm: q_norm.as_deref(),
15395 k_norm: k_norm.as_deref(),
15396 output_gate: *output_gate,
15397 softplus_gate: softplus_gate
15398 .as_ref()
15399 .map(|(gate, per_head)| (gate, *per_head)),
15400 rope_scale: self.layer_rope_scale(li),
15401 bias: bias
15402 .as_ref()
15403 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
15404 rms_eps: eps,
15405 norm_style: self.norm_style,
15406 pool: pool.as_deref(),
15407 v_head_dim: self.layer_v_dim(li),
15408 };
15409 attention::qwen_attention(
15410 &self.ws.n1,
15411 wq,
15412 wk,
15413 wv,
15414 wo,
15415 &mut self.kv_cache.layers[li],
15416 &cfg,
15417 )
15418 }
15419 }
15420 }
15421 };
15422 let attn_out = match &self.weights.layers[self.phys_layer(li)].attn_out_norm {
15425 Some(w) => inference::rms_norm(&attn_out, w, self.rms_eps, self.norm_style),
15426 None => attn_out,
15427 };
15428 let lw = &self.weights.layers[self.phys_layer(li)];
15429 let prof = crate::cpuprof::time(crate::cpuprof::Slot::Norms);
15430 inference::add_rmsnorm_fused_into(
15431 &mut h,
15432 &attn_out,
15433 &lw.post_norm,
15434 self.rms_eps,
15435 self.norm_style,
15436 &mut self.ws.p1,
15437 );
15438 drop(prof);
15439 let mut attn_out = attn_out;
15440 attention::recycle_buf(&mut attn_out);
15441 let post_normed = &self.ws.p1;
15442
15443 let ffn_masked = task_mask
15444 .map(|m| m.ffn_active_count(li) < self.intermediate_size)
15445 .unwrap_or(false);
15446 let ffn_out = match (ffn_masked, &lw.ffn) {
15458 (_, FfnKind::Dense(d)) if !d.segs.is_empty() => {
15462 let row = task_mask
15463 .and_then(|tm| tm.ffn_masks.get(li))
15464 .map(|v| v.as_slice());
15465 tube_ffn(d, post_normed, 1, self.pool.as_deref(), row)
15466 }
15467 (true, FfnKind::Dense(d)) => {
15468 let tm = task_mask.unwrap();
15469 let alive = tm.ffn_active_count(li);
15470 let deep = alive * 2 <= self.intermediate_size;
15471 if deep && d.down_proj.sparse_col_ok() && !d.gate_proj.has_prism_contract() {
15472 let active = tm.ffn_active_indices(li);
15473 sparse_ffn_quant(
15474 d,
15475 post_normed,
15476 &active,
15477 self.hidden_size,
15478 self.pool.as_deref(),
15479 )
15480 } else if deep
15481 && let (Some(g), Some(u), Some(dn)) = (
15482 d.gate_proj.as_f32(),
15483 d.up_proj.as_f32(),
15484 d.down_proj.as_f32(),
15485 )
15486 {
15487 let active = tm.ffn_active_indices(li);
15488 inference::sparse_ffn_forward(
15489 post_normed,
15490 g,
15491 u,
15492 dn,
15493 self.hidden_size,
15494 self.intermediate_size,
15495 &active,
15496 self.pool.as_deref(),
15497 )
15498 } else {
15499 let row = tm.ffn_masks.get(li).map(|v| v.as_slice());
15500 dense_ffn_batch(d, post_normed, 1, self.pool.as_deref(), row)
15501 }
15502 }
15503 (true, FfnKind::Moe(m)) => {
15504 let allowed = task_mask.and_then(|tm| tm.expert_flags(li, m.experts.len()));
15508 ffn_forward(
15509 &lw.ffn,
15510 post_normed,
15511 self.pool.as_deref(),
15512 allowed.as_deref(),
15513 )
15514 }
15515 (true, FfnKind::DenseMoe(dm)) => dense_moe_ffn(
15516 dm,
15517 post_normed,
15518 &h,
15519 self.rms_eps,
15520 self.norm_style,
15521 self.pool.as_deref(),
15522 ),
15523 (false, _) => match &lw.ffn {
15524 FfnKind::DenseMoe(dm) => dense_moe_ffn(
15525 dm,
15526 post_normed,
15527 &h,
15528 self.rms_eps,
15529 self.norm_style,
15530 self.pool.as_deref(),
15531 ),
15532 FfnKind::Moe(m)
15533 if task_mask.is_none() && self.mimo_moe.is_dynamic(li, host_tail) =>
15534 {
15535 moe_ffn_banked(&mut self.mimo_moe, li, m, post_normed, self.pool.as_deref())
15536 }
15537 _ => {
15538 let allowed = match (&lw.ffn, task_mask) {
15539 (FfnKind::Moe(m), Some(tm)) => tm.expert_flags(li, m.experts.len()),
15540 _ => None,
15541 };
15542 ffn_forward(
15543 &lw.ffn,
15544 post_normed,
15545 self.pool.as_deref(),
15546 allowed.as_deref(),
15547 )
15548 }
15549 },
15550 };
15551 let ffn_out = match &self.weights.layers[self.phys_layer(li)].ffn_out_norm {
15552 Some(w) => inference::rms_norm(&ffn_out, w, self.rms_eps, self.norm_style),
15553 None => ffn_out,
15554 };
15555 for (i, &f) in ffn_out.iter().enumerate() {
15556 h[i] += f;
15557 }
15558 let mut ffn_out = ffn_out;
15559 attention::recycle_buf(&mut ffn_out);
15560
15561 if let Some(sc) = self.weights.layers[self.phys_layer(li)].layer_scale {
15563 for v in h.iter_mut() {
15564 *v *= sc;
15565 }
15566 }
15567 if self.layer_dump.is_some() {
15569 self.dump_layer_row(position, li, &h);
15570 }
15571
15572 if self.is_loop_end(li) && li + 1 < self.num_layers {
15575 h = inference::rms_norm(
15576 &h,
15577 &self.weights.final_norm,
15578 self.rms_eps,
15579 self.norm_style,
15580 );
15581 }
15582
15583 if self.dyn_phi_layer == Some(li) {
15587 self.update_dyn_phi(&h);
15588 }
15589 }
15590 crate::gpu::set_layer(-1); if let Some(t) = t_race_cpu {
15592 crate::gpu::graph_race_record(false, t.elapsed());
15593 }
15594
15595 h
15596 }
15597
15598 fn update_dyn_phi(&mut self, h: &[f32]) {
15601 const A: f32 = 0.2;
15602 if self.dyn_phi_ema.len() != h.len() {
15603 self.dyn_phi_ema = vec![0.0; h.len()];
15604 self.dyn_phi_seen = 0;
15605 }
15606 if self.dyn_phi_seen == 0 {
15607 self.dyn_phi_ema.copy_from_slice(h);
15608 } else {
15609 for (e, &v) in self.dyn_phi_ema.iter_mut().zip(h) {
15610 *e = (1.0 - A) * *e + A * v;
15611 }
15612 }
15613 self.dyn_phi_seen += 1;
15614 }
15615
15616 pub fn dyn_phi(&self) -> &[f32] {
15618 &self.dyn_phi_ema
15619 }
15620
15621 pub fn set_dyn_phi_layer(&mut self, layer: Option<usize>) {
15623 self.dyn_phi_layer = layer;
15624 self.dyn_phi_ema.clear();
15625 self.dyn_phi_seen = 0;
15626 }
15627
15628 pub fn dynamic_skills(&self) -> Vec<(usize, String, usize)> {
15630 let Some(model) = &self.model else {
15631 return Vec::new();
15632 };
15633 model
15634 .header
15635 .skills
15636 .iter()
15637 .enumerate()
15638 .filter_map(|(i, sk)| {
15639 if sk.is_v2() {
15643 return None;
15644 }
15645 let ok = matches!(self.dyn_skill_layers.get(i), Some(Some(_)));
15646 let sel = sk.selection.as_ref()?;
15647 (ok).then(|| (i, sk.id.clone(), sel.phi_layer))
15648 })
15649 .collect()
15650 }
15651
15652 pub fn active_skill(&self) -> Option<usize> {
15654 self.dyn_active
15655 }
15656
15657 pub fn enable_dynamic_routing(&mut self) -> usize {
15662 use crate::swarm::{DynRouter, RoutableSkill};
15663 let Some(model) = self.model.clone() else {
15664 return 0;
15665 };
15666 if let Some(r) = &model.header.router {
15673 tracing::warn!(
15674 "dynamic routing disabled: this file declares router policy '{}' with \
15675 granularity \"{}\" — the request-level decision applies instead",
15676 r.policy,
15677 r.granularity
15678 );
15679 return 0;
15680 }
15681 if model.required_features & cortiq_core::format::features::SKILLS_V2 != 0
15687 || model.header.skills.iter().any(|s| s.is_v2())
15688 {
15689 tracing::warn!(
15690 "dynamic routing disabled: this file carries format-v2 skill records \
15691 (SKILLS_V2) — they route per request through a router policy only"
15692 );
15693 return 0;
15694 }
15695 if self.dyn_blend_loaded {
15698 tracing::warn!("dynamic routing unavailable on a blend-loaded pipeline");
15699 return 0;
15700 }
15701 if let Some(a) = self.dyn_active {
15705 if !matches!(self.dyn_skill_layers.get(a), Some(Some(_))) {
15706 tracing::warn!("loaded skill is not FFN-eligible — dynamic routing unavailable");
15707 return 0;
15708 }
15709 }
15710 let hidden = self.hidden_size;
15711 let mut skills = Vec::new();
15712 for (idx, id, _phi) in self.dynamic_skills() {
15713 if let Some(sel) = model.header.skills[idx].selection.as_ref() {
15714 if let Some(rs) = RoutableSkill::from_descriptor(idx, id, sel, hidden) {
15715 skills.push(rs);
15716 }
15717 }
15718 }
15719 if skills.is_empty() {
15720 return 0;
15721 }
15722 let phi = skills[0].phi_layer;
15724 if skills.iter().any(|s| s.phi_layer != phi) {
15725 tracing::warn!("routable skills disagree on phi_layer; using {phi}");
15726 }
15727 let n = skills.len();
15728 self.set_dyn_phi_layer(Some(phi));
15729 self.dyn_router = Some(DynRouter::new(skills));
15730 n
15731 }
15732
15733 pub fn route_switches(&self) -> Vec<(usize, Option<String>, Option<String>)> {
15735 self.dyn_router
15736 .as_ref()
15737 .map(|r| r.switches.clone())
15738 .unwrap_or_default()
15739 }
15740
15741 fn lm_head_forward(&self, hidden: &[f32]) -> Vec<f32> {
15744 let _mimo_q8 = self.mimo_moe.is_on()
15745 .then(crate::qtensor::enter_full_gpu_q8_scope);
15746 let rows = self.weights.lm_head.rows();
15747 let mut logits = attention::take_buf(rows.min(self.vocab_size));
15748 let served = self.mimo_moe.is_on() && crate::gpu::mimo_q8_short_enabled()
15752 && rows == self.vocab_size && !self.weights.lm_head.has_prism_contract()
15753 && self.weights.lm_head.graph_weight().is_some_and(|(model, idx, kind, _)| {
15754 kind == 7 && crate::gpu::q82_short_rows(model, idx, hidden, 1,
15755 rows, self.hidden_size, &mut logits)
15756 });
15757 if !served {
15758 self.weights.lm_head.matvec(hidden, &mut logits, self.pool.as_deref());
15759 }
15760 logits.resize(self.vocab_size, 0.0);
15761 if let Some(m) = self.logit_multiplier {
15762 for l in logits.iter_mut() {
15763 *l *= m;
15764 }
15765 }
15766 if let Some(c) = self.final_softcap {
15767 for l in logits.iter_mut() {
15768 *l = c * (*l / c).tanh();
15769 }
15770 }
15771 if let Some(cm) = self.head_clusters.as_ref() {
15772 self.hierarchical_head_logprobs(hidden, cm, &mut logits);
15773 }
15774 logits
15775 }
15776
15777 fn hierarchical_head_logprobs(&self, hidden: &[f32], cm: &[f32], logits: &mut [f32]) {
15780 let h = hidden.len();
15781 let ncl = cm.len() / h.max(1);
15782 if ncl == 0 || logits.len() % ncl != 0 {
15783 return;
15784 }
15785 let cs = logits.len() / ncl;
15786 let mut lc = vec![0.0f32; ncl];
15788 for c in 0..ncl {
15789 let row = &cm[c * h..(c + 1) * h];
15790 let mut s = 0.0f32;
15791 for j in 0..h {
15792 s += row[j] * hidden[j];
15793 }
15794 lc[c] = s;
15795 }
15796 let mx = lc.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
15797 let lse: f32 = mx + lc.iter().map(|v| (v - mx).exp()).sum::<f32>().ln();
15798 for c in 0..ncl {
15799 let blk = &mut logits[c * cs..(c + 1) * cs];
15800 let bm = blk.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
15801 let bl: f32 = bm + blk.iter().map(|v| (v - bm).exp()).sum::<f32>().ln();
15802 let add = lc[c] - lse - bl;
15803 for v in blk.iter_mut() {
15804 *v += add;
15805 }
15806 }
15807 }
15808
15809 pub fn prefill_next_logits(&mut self, ids: &[u32], task_mask: Option<&TaskMask>) -> Vec<f32> {
15814 #[cfg(target_os = "macos")]
15815 crate::gpu_metal::set_io_namespace(self.graph_kv_id);
15816 self.clear_sequence_state();
15817 crate::gpu::graph_race_begin_generation();
15821 if task_mask.is_none() {
15822 self.o1_begin();
15823 }
15824 let mut hidden = vec![0.0f32; self.hidden_size];
15825 for (pos, &id) in ids.iter().enumerate() {
15826 let emb = self.embed_single(id);
15827 hidden = self.forward_layers(&emb, pos, task_mask);
15828 }
15829 if let Err(err) = self.o1_seal_checked() {
15830 self.o1_fail(err);
15831 }
15832 if let Some(logits) = self.graph_logits.take() {
15835 return logits;
15836 }
15837 inference::rms_norm_into(
15838 &hidden,
15839 &self.weights.final_norm,
15840 self.rms_eps,
15841 self.norm_style,
15842 &mut self.ws.n1,
15843 );
15844 self.lm_head_forward(&self.ws.n1)
15845 }
15846}
15847
15848pub fn create_test_pipeline(
15850 hidden_size: usize,
15851 intermediate_size: usize,
15852 num_heads: usize,
15853 num_kv_heads: usize,
15854 head_dim: usize,
15855 num_layers: usize,
15856 vocab_size: usize,
15857) -> Pipeline {
15858 let synth = |n: usize, salt: usize| -> Vec<f32> {
15861 (0..n)
15862 .map(|i| (((i * 31 + salt * 17 + 7) % 97) as f32 / 97.0 - 0.5) * 0.2)
15863 .collect()
15864 };
15865 let qt = |rows: usize, cols: usize, salt: usize| -> QTensor {
15866 QTensor::from_f32(synth(rows * cols, salt), rows, cols)
15867 };
15868 let layer_weights: Vec<LayerWeights> = (0..num_layers)
15869 .map(|li| LayerWeights {
15870 input_norm: vec![1.0; hidden_size],
15871 post_norm: vec![1.0; hidden_size],
15872 attn_out_norm: None,
15873 ffn_out_norm: None,
15874 layer_scale: None,
15875 ffn: FfnKind::Dense(DenseFfn {
15876 gate_proj: qt(intermediate_size, hidden_size, li * 10 + 5),
15877 up_proj: qt(intermediate_size, hidden_size, li * 10 + 6),
15878 down_proj: qt(hidden_size, intermediate_size, li * 10 + 7),
15879 act: Act::Silu,
15880 down_t: None,
15881 segs: Vec::new(),
15882 }),
15883 attn: AttnKind::Full {
15884 bias: None,
15885 wq: qt(num_heads * head_dim, hidden_size, li * 10 + 1),
15886 wk: qt(num_kv_heads * head_dim, hidden_size, li * 10 + 2),
15887 wv: qt(num_kv_heads * head_dim, hidden_size, li * 10 + 3),
15888 wo: qt(hidden_size, num_heads * head_dim, li * 10 + 4),
15889 q_norm: None,
15890 k_norm: None,
15891 output_gate: false,
15892 softplus_gate: None,
15893 },
15894 })
15895 .collect();
15896
15897 Pipeline::new(
15898 Tokenizer::byte_level(),
15899 PipelineWeights {
15900 embed_tokens: qt(vocab_size, hidden_size, 100),
15901 layers: layer_weights,
15902 lm_head: qt(vocab_size, hidden_size, 200),
15903 final_norm: vec![1.0; hidden_size],
15904 },
15905 hidden_size,
15906 intermediate_size,
15907 num_heads,
15908 num_kv_heads,
15909 head_dim,
15910 num_layers,
15911 num_layers, false, vocab_size,
15914 1e-6,
15915 10_000.0,
15916 NormStyle::Qwen,
15917 4096,
15918 SamplerConfig {
15919 seed: Some(42),
15920 ..Default::default()
15921 },
15922 )
15923}
15924
15925#[inline]
15930fn mask_bit(row: &[u8], j: usize) -> bool {
15931 (row.get(j >> 3).copied().unwrap_or(0) >> (j & 7)) & 1 != 0
15932}
15933
15934fn mask_gain() -> f32 {
15945 static G: std::sync::OnceLock<f32> = std::sync::OnceLock::new();
15946 *G.get_or_init(|| {
15947 std::env::var("CMF_FFN_MASK_GAIN")
15948 .ok()
15949 .and_then(|v| v.parse().ok())
15950 .unwrap_or(1.0)
15951 })
15952}
15953
15954fn zero_masked_cols(g: &mut [f32], rows: usize, inter: usize, row: &[u8]) {
15955 let fill = meanfill().and_then(|(i, v)| {
15958 let li = crate::gpu::cur_layer();
15959 (*i == inter && li >= 0).then(|| &v[li as usize * inter..(li as usize + 1) * inter])
15960 });
15961 for r in 0..rows {
15962 let base = r * inter;
15963 for (bi, &byte) in row.iter().enumerate() {
15964 if byte == 0xFF {
15965 continue;
15966 }
15967 let j0 = bi * 8;
15968 for bit in 0..8 {
15969 let j = j0 + bit;
15970 if j < inter && byte & (1 << bit) == 0 {
15971 g[base + j] = fill.map_or(0.0, |f| f[j]);
15972 }
15973 }
15974 }
15975 }
15976 let gain = mask_gain();
15977 if gain != 1.0 {
15978 for v in g[..rows * inter].iter_mut() {
15979 *v *= gain;
15980 }
15981 }
15982}
15983
15984#[inline]
15986fn tube_bit(row: Option<&[u8]>, i: usize) -> bool {
15987 row.is_none_or(|r| mask_bit(r, i))
15988}
15989
15990fn all_bits_on(row: &[u8], n: usize) -> bool {
15993 (0..n).all(|i| mask_bit(row, i))
15994}
15995
15996fn tube_topk() -> usize {
16004 static K: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
16005 *K.get_or_init(|| {
16006 std::env::var("CMF_TUBE_TOPK")
16007 .ok()
16008 .and_then(|v| v.parse().ok())
16009 .unwrap_or(0)
16010 })
16011}
16012
16013fn tube_score_oracle() -> bool {
16014 static O: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
16015 *O.get_or_init(|| std::env::var("CMF_TUBE_SCORE").is_ok_and(|v| v == "oracle"))
16016}
16017
16018fn tube_ffn_routed(
16025 d: &DenseFfn,
16026 xs: &[f32],
16027 b: usize,
16028 pool: Option<&Pool>,
16029 mask_row: Option<&[u8]>,
16030 k: usize,
16031) -> Vec<f32> {
16032 let hidden = d.down_proj.rows();
16033 let core = d.gate_proj.rows();
16034 let core_full = mask_row.is_none_or(|r| all_bits_on(r, core));
16035 let mut out = match (b, core_full, mask_row) {
16036 (1, true, _) => dense_ffn(d, xs, pool),
16037 (1, false, Some(row)) => dense_ffn_masked(d, xs, pool, row),
16038 (_, true, _) => dense_ffn_batch(d, xs, b, pool, None),
16039 (_, false, row) => dense_ffn_batch(d, xs, b, pool, row),
16040 };
16041 let cand: Vec<usize> = (0..d.segs.len())
16042 .filter(|&i| tube_bit(mask_row, d.segs[i].start))
16043 .collect();
16044 if cand.is_empty() {
16045 return out;
16046 }
16047 let oracle = tube_score_oracle();
16051 let mut acts: Vec<Vec<f32>> = Vec::with_capacity(cand.len());
16052 let mut scores = vec![0f32; b * cand.len()];
16053 for (ci, &i) in cand.iter().enumerate() {
16054 let seg = &d.segs[i];
16055 let w = seg.width;
16056 let mut g = vec![0.0f32; b * w];
16057 if b == 1 {
16058 seg.gate.matvec(xs, &mut g, pool);
16059 } else {
16060 seg.gate.matmat(xs, b, &mut g, pool);
16061 }
16062 for v in g.iter_mut() {
16063 *v = Act::Silu.combine(*v, 1.0);
16064 }
16065 if !oracle {
16066 for t in 0..b {
16067 scores[t * cand.len() + ci] =
16068 g[t * w..(t + 1) * w].iter().map(|v| v * v).sum::<f32>();
16069 }
16070 }
16071 if oracle || b > 1 {
16072 let mut u = vec![0.0f32; b * w];
16073 if b == 1 {
16074 seg.up.matvec(xs, &mut u, pool);
16075 } else {
16076 seg.up.matmat(xs, b, &mut u, pool);
16077 }
16078 for (a, &v) in g.iter_mut().zip(u.iter()) {
16079 *a *= v;
16080 }
16081 if oracle {
16082 for t in 0..b {
16083 scores[t * cand.len() + ci] =
16084 g[t * w..(t + 1) * w].iter().map(|v| v * v).sum::<f32>();
16085 }
16086 }
16087 }
16088 acts.push(g);
16089 }
16090 let keep = k.min(cand.len());
16092 let mut scratch: Vec<f32> = Vec::new();
16093 for t in 0..b {
16094 let mut sc: Vec<(f32, usize)> = (0..cand.len())
16095 .map(|ci| (scores[t * cand.len() + ci], ci))
16096 .collect();
16097 sc.sort_unstable_by(|x, y| y.0.total_cmp(&x.0));
16098 let mut alive = vec![false; cand.len()];
16099 for &(_, ci) in sc.iter().take(keep) {
16100 alive[ci] = true;
16101 }
16102 if b > 1 {
16103 for (ci, a) in acts.iter_mut().enumerate() {
16104 if !alive[ci] {
16105 let w = d.segs[cand[ci]].width;
16106 a[t * w..(t + 1) * w].fill(0.0);
16107 }
16108 }
16109 } else {
16110 for (ci, &i) in cand.iter().enumerate() {
16114 if !alive[ci] {
16115 continue;
16116 }
16117 let seg = &d.segs[i];
16118 let w = seg.width;
16119 let g = &mut acts[ci];
16120 if !tube_score_oracle() {
16121 scratch.clear();
16122 scratch.resize(w, 0.0);
16123 seg.up.matvec(xs, &mut scratch, pool);
16124 for (a, &v) in g.iter_mut().zip(scratch.iter()) {
16125 *a *= v;
16126 }
16127 }
16128 let mut acc = vec![0.0f32; hidden];
16129 seg.down.matvec(g, &mut acc, pool);
16130 for (o, a) in out.iter_mut().zip(&acc) {
16131 *o += *a;
16132 }
16133 }
16134 }
16135 }
16136 if b > 1 {
16137 for (ci, &i) in cand.iter().enumerate() {
16138 let seg = &d.segs[i];
16139 let mut acc = vec![0.0f32; b * hidden];
16140 seg.down.matmat(&acts[ci], b, &mut acc, pool);
16141 for (o, a) in out.iter_mut().zip(&acc) {
16142 *o += *a;
16143 }
16144 }
16145 }
16146 out
16147}
16148
16149fn tube_ffn(
16155 d: &DenseFfn,
16156 xs: &[f32],
16157 b: usize,
16158 pool: Option<&Pool>,
16159 mask_row: Option<&[u8]>,
16160) -> Vec<f32> {
16161 if tube_topk() > 0 {
16162 return tube_ffn_routed(d, xs, b, pool, mask_row, tube_topk());
16163 }
16164 let hidden = d.down_proj.rows();
16165 let core = d.gate_proj.rows();
16166 let core_full = mask_row.is_none_or(|r| all_bits_on(r, core));
16167 let mut out = match (b, core_full, mask_row) {
16168 (1, true, _) => dense_ffn(d, xs, pool),
16169 (1, false, Some(row)) => dense_ffn_masked(d, xs, pool, row),
16170 (_, true, _) => dense_ffn_batch(d, xs, b, pool, None),
16171 (_, false, row) => dense_ffn_batch(d, xs, b, pool, row),
16172 };
16173 TUBE_SCRATCH.with(|sc| {
16174 let mut sc = sc.borrow_mut();
16175 let [g, u, acc] = &mut *sc;
16176 for seg in &d.segs {
16177 if !tube_bit(mask_row, seg.start) {
16178 continue;
16179 }
16180 let w = seg.width;
16181 g.resize(b * w, 0.0);
16182 if b == 1
16183 && d.act == Act::Silu
16184 && QTensor::matvec_silu_mul(&seg.gate, &seg.up, xs, g, pool)
16185 {
16186 } else {
16188 u.resize(b * w, 0.0);
16189 if b == 1 {
16190 QTensor::matvec_many([&seg.gate, &seg.up], xs, [g, u], pool);
16191 } else {
16192 seg.gate.matmat(xs, b, g, pool);
16193 seg.up.matmat(xs, b, u, pool);
16194 }
16195 for i in 0..b * w {
16196 g[i] = d.act.combine(g[i], u[i]);
16197 }
16198 }
16199 acc.resize(b * hidden, 0.0);
16200 acc.fill(0.0);
16201 if b == 1 {
16202 seg.down.matvec(g, acc, pool);
16203 } else {
16204 seg.down.matmat(g, b, acc, pool);
16205 }
16206 for (o, a) in out.iter_mut().zip(acc.iter()) {
16207 *o += *a;
16208 }
16209 }
16210 out
16211 })
16212}
16213
16214thread_local! {
16215 static TUBE_SCRATCH: std::cell::RefCell<[Vec<f32>; 3]> =
16219 const { std::cell::RefCell::new([Vec::new(), Vec::new(), Vec::new()]) };
16220}
16221
16222static PREFILL_SPLIT: [std::sync::atomic::AtomicU64; 2] =
16225 [std::sync::atomic::AtomicU64::new(0), std::sync::atomic::AtomicU64::new(0)];
16226
16227fn wgpu_prefill_counters() -> (f64, f64, f64, u64) {
16231 #[cfg(feature = "gpu")]
16232 {
16233 use std::sync::atomic::Ordering::Relaxed;
16234 let ms = |a: &std::sync::atomic::AtomicU64| a.load(Relaxed) as f64 / 1e6;
16235 (
16236 ms(&crate::gpu_wgpu::GEMM_MANY_NS[0]),
16237 ms(&crate::gpu_wgpu::GEMM_MANY_NS[1]),
16238 ms(&crate::gpu_wgpu::CHUNK_ATTEND_WAIT_NS),
16239 crate::gpu_wgpu::ZI_GEMM_CALLS.load(Relaxed),
16240 )
16241 }
16242 #[cfg(not(feature = "gpu"))]
16243 (0.0, 0.0, 0.0, 0)
16244}
16245
16246fn wgpu_mirror_counters() -> [u64; 6] {
16249 #[cfg(feature = "gpu")]
16250 {
16251 crate::gpu_wgpu::MIRROR_EVENTS
16252 .each_ref()
16253 .map(|a| a.load(std::sync::atomic::Ordering::Relaxed))
16254 }
16255 #[cfg(not(feature = "gpu"))]
16256 [0; 6]
16257}
16258
16259fn prefill_prof_on() -> bool {
16260 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
16261 *ON.get_or_init(|| std::env::var_os("CMF_PREFILL_PROF").is_some())
16262}
16263
16264fn dense_ffn_batch(
16265 d: &DenseFfn,
16266 xs: &[f32],
16267 b: usize,
16268 pool: Option<&Pool>,
16269 mask_row: Option<&[u8]>,
16270) -> Vec<f32> {
16271 let inter = d.gate_proj.rows();
16272 let hidden = d.down_proj.rows();
16273 let fused_act = d.act.graph_act();
16283 if mask_row.is_none()
16284 && fused_act.is_some()
16285 && b >= 32
16286 && crate::gpu::enabled_here()
16287 && !crate::gpu::mm_killed()
16288 && refit_dir().is_none()
16293 && !ffn_probe_active()
16298 {
16299 if let (Some((model, w1)), Some((_, w3)), Some((_, w2))) = (
16300 d.gate_proj.mapped_q4t(),
16301 d.up_proj.mapped_q4t(),
16302 d.down_proj.mapped_q4t(),
16303 ) {
16304 let mut out = vec![0.0f32; b * hidden];
16305 let act = fused_act.expect("checked above");
16306 if crate::gpu::q4_ffn_act(model, w1, w3, w2, xs, b, hidden, inter, false, act, &mut out)
16307 {
16308 return out;
16309 }
16310 }
16311 if let (Some((model, w1)), Some((_, w3)), Some((_, w2))) = (
16316 d.gate_proj.mapped_q4tp(),
16317 d.up_proj.mapped_q4tp(),
16318 d.down_proj.mapped_q4tp(),
16319 ) {
16320 let mut out = vec![0.0f32; b * hidden];
16321 let act = fused_act.expect("checked above");
16322 if crate::gpu::q4_ffn_act(model, w1, w3, w2, xs, b, hidden, inter, true, act, &mut out)
16323 {
16324 return out;
16325 }
16326 }
16327 if let (true, Some((model, w1)), Some((_, w3)), Some((_, w2))) = (
16333 std::env::var("CMF_FFN_KEEP").as_deref() != Ok("0"),
16334 d.gate_proj.mapped_device_gemm(),
16335 d.up_proj.mapped_device_gemm(),
16336 d.down_proj.mapped_device_gemm(),
16337 ) {
16338 let mut out = vec![0.0f32; b * hidden];
16339 let act = fused_act.expect("checked above");
16340 if crate::gpu::ffn_act_keep(model, w1, w3, w2, xs, b, hidden, inter, act, &mut out) {
16341 return out;
16342 }
16343 }
16344 }
16345 let mut g = vec![0.0f32; b * inter];
16346 d.gate_proj.matmat(xs, b, &mut g, pool);
16347 let mut u = vec![0.0f32; b * inter];
16348 d.up_proj.matmat(xs, b, &mut u, pool);
16349 if gate_topk() > 0 && d.act == Act::Silu {
16350 for t in 0..b {
16351 let row = &mut g[t * inter..(t + 1) * inter];
16352 for v in row.iter_mut() {
16353 *v = Act::Silu.combine(*v, 1.0);
16354 }
16355 keep_top_k(row, gate_topk());
16356 }
16357 for i in 0..b * inter {
16358 g[i] *= u[i];
16359 }
16360 } else {
16361 for i in 0..b * inter {
16362 g[i] = d.act.combine(g[i], u[i]);
16363 }
16364 }
16365 if let Some(row) = mask_row {
16366 zero_masked_cols(&mut g, b, inter, row);
16367 }
16368 if oracle_topk() > 0 {
16369 for t in 0..b {
16370 keep_top_k(&mut g[t * inter..(t + 1) * inter], oracle_topk());
16371 }
16372 }
16373 let mut out = vec![0.0f32; b * hidden];
16374 d.down_proj.matmat(&g, b, &mut out, pool);
16375 if refit_dir().is_some() {
16376 let li = crate::gpu::cur_layer();
16377 if li >= 0 {
16378 refit_accumulate(li as usize, &g, b, inter, &out, hidden, pool);
16379 }
16380 }
16381 FFN_PROBE.with(|pr| {
16385 if let Some(acc) = pr.borrow_mut().as_mut() {
16386 let li = crate::gpu::cur_layer();
16387 if li < 0 {
16388 return;
16389 }
16390 let Some(row) = acc.get_mut(li as usize) else {
16391 return;
16392 };
16393 let sq = probe_sq();
16394 for t in 0..b {
16395 for (a, &v) in row.iter_mut().zip(&g[t * inter..(t + 1) * inter]) {
16396 *a += if sq {
16397 (v as f64) * (v as f64)
16398 } else {
16399 (v as f64).abs()
16400 };
16401 }
16402 }
16403 }
16404 });
16405 out
16406}
16407
16408fn accumulate_act(m: &MoeFfn, xs: &[f32], b: usize) {
16413 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
16414 static DUMP: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
16415 let on = *ON.get_or_init(|| std::env::var("CMF_RMS_TRACE").is_ok());
16416 let dump = *DUMP.get_or_init(|| std::env::var("CMF_ACT_DUMP").is_ok());
16417 if (!on && !dump) || b == 0 {
16418 return;
16419 }
16420 let hidden = xs.len() / b;
16421 if on {
16422 let mut acc = m.act_sq.borrow_mut();
16423 if acc.len() < hidden {
16424 acc.resize(hidden, 0.0);
16425 }
16426 for t in 0..b {
16427 let row = &xs[t * hidden..(t + 1) * hidden];
16428 for (a, &v) in acc.iter_mut().zip(row) {
16429 *a += (v as f64) * (v as f64);
16430 }
16431 }
16432 }
16433 if dump {
16434 let cap: usize = std::env::var("CMF_ACT_DUMP_ROWS")
16437 .ok()
16438 .and_then(|v| v.parse().ok())
16439 .unwrap_or(4096);
16440 let mut rows = m.act_rows.borrow_mut();
16441 if rows.len() < cap * hidden {
16442 let take = b.min((cap * hidden - rows.len()) / hidden.max(1));
16443 rows.extend_from_slice(&xs[..take * hidden]);
16444 }
16445 }
16446}
16447
16448#[derive(Clone, Copy)]
16451struct SendVecs(*mut Vec<f32>);
16452unsafe impl Send for SendVecs {}
16453unsafe impl Sync for SendVecs {}
16454impl SendVecs {
16455 #[inline]
16456 fn at(self, i: usize) -> *mut Vec<f32> {
16457 unsafe { self.0.add(i) }
16458 }
16459}
16460
16461fn moe_ffn_batch(
16462 m: &MoeFfn,
16463 xs: &[f32],
16464 b: usize,
16465 hidden: usize,
16466 pool: Option<&Pool>,
16467 allowed: Option<&[bool]>,
16468) -> Vec<f32> {
16469 accumulate_act(m, xs, b);
16470 let ne = m.experts.len();
16471 let mut logits = vec![0.0f32; b * ne];
16472 match &m.resonance {
16473 Some(r) => {
16474 let hdim = xs.len() / b.max(1);
16475 for bi in 0..b {
16476 r.scores(
16477 &xs[bi * hdim..(bi + 1) * hdim],
16478 &mut logits[bi * ne..(bi + 1) * ne],
16479 );
16480 }
16481 }
16482 None => m.router.matmat(xs, b, &mut logits, pool),
16483 }
16484
16485 let mut assign: Vec<Vec<(usize, f32)>> = vec![Vec::new(); ne];
16488 {
16489 let mut st = m.stats.borrow_mut();
16490 if st.len() < ne {
16491 st.resize(ne, 0);
16492 }
16493 for bi in 0..b {
16494 let (idx, p, wsum) = moe_route(&logits[bi * ne..(bi + 1) * ne], m, allowed);
16495 for &e in &idx {
16496 st[e] += 1;
16497 assign[e].push((bi, p[e] / wsum));
16498 }
16499 }
16500 }
16501
16502 let mut out = vec![0.0f32; b * hidden];
16503 let cols = m.experts[0].gate_proj.cols();
16504 let run_expert = |d: &DenseFfn, list: &[(usize, f32)], out: &mut [f32]| {
16505 let sb = list.len();
16506 let mut sub = vec![0.0f32; sb * cols];
16507 for (k, &(bi, _)) in list.iter().enumerate() {
16508 sub[k * cols..(k + 1) * cols].copy_from_slice(&xs[bi * cols..(bi + 1) * cols]);
16509 }
16510 let eo = dense_ffn_batch(d, &sub, sb, pool, None);
16511 for (k, &(bi, w)) in list.iter().enumerate() {
16512 for i in 0..hidden {
16513 out[bi * hidden + i] += w * eo[k * hidden + i];
16514 }
16515 }
16516 };
16517 let active: Vec<usize> = (0..ne).filter(|&e| !assign[e].is_empty()).collect();
16523 if pool.is_some() && active.len() >= 8 {
16524 let mut panels: Vec<Vec<f32>> = vec![Vec::new(); active.len()];
16525 {
16526 let panel_ptr = SendVecs(panels.as_mut_ptr());
16527 let experts = &m.experts;
16530 let (active_r, assign_r) = (&active, &assign);
16531 let inherit_cpu = crate::gpu::inherit_cpu_scope();
16532 let run = |start: usize, end: usize| {
16533 let _cpu_scope = inherit_cpu();
16534 for ai in start..end {
16535 let e = active_r[ai];
16536 let list = &assign_r[e];
16537 let sb = list.len();
16538 let mut sub = vec![0.0f32; sb * cols];
16539 for (k, &(bi, _)) in list.iter().enumerate() {
16540 sub[k * cols..(k + 1) * cols]
16541 .copy_from_slice(&xs[bi * cols..(bi + 1) * cols]);
16542 }
16543 unsafe {
16545 *panel_ptr.at(ai) = dense_ffn_batch(&experts[e], &sub, sb, None, None);
16546 }
16547 }
16548 };
16549 match pool {
16550 Some(p) => p.run_rows(active.len(), &run),
16551 None => run(0, active.len()),
16552 }
16553 }
16554 for (ai, &e) in active.iter().enumerate() {
16555 for (k, &(bi, w)) in assign[e].iter().enumerate() {
16556 let eo = &panels[ai][k * hidden..(k + 1) * hidden];
16557 for i in 0..hidden {
16558 out[bi * hidden + i] += w * eo[i];
16559 }
16560 }
16561 }
16562 } else {
16563 for &e in &active {
16564 run_expert(&m.experts[e], &assign[e], &mut out);
16565 }
16566 }
16567 if let Some((se, gate)) = &m.shared {
16568 let all: Vec<(usize, f32)> = if let Some(gate) = gate {
16569 let mut gl = vec![0.0f32; b];
16570 gate.matmat(xs, b, &mut gl, pool);
16571 (0..b)
16572 .map(|bi| (bi, 1.0 / (1.0 + (-gl[bi]).exp())))
16573 .collect()
16574 } else {
16575 (0..b).map(|bi| (bi, 1.0)).collect()
16576 };
16577 run_expert(se, &all, &mut out);
16578 }
16579 out
16580}
16581
16582fn moe_ffn_rows_exact(
16594 m: &MoeFfn,
16595 xs: &[f32],
16596 b: usize,
16597 hidden: usize,
16598 pool: Option<&Pool>,
16599) -> Vec<f32> {
16600 let mut out = vec![0.0f32; b * hidden];
16601 let per_row = |out: &mut [f32]| {
16602 for r in 0..b {
16603 let o = moe_ffn(m, &xs[r * hidden..(r + 1) * hidden], pool, None);
16604 out[r * hidden..(r + 1) * hidden].copy_from_slice(&o);
16605 }
16606 };
16607 let covered = !crate::gpu::enabled_here()
16608 && moe_batch_enabled()
16609 && m.shared.is_none()
16610 && m.resonance.is_none()
16611 && FFN_PROBE.with(|pr| pr.borrow().is_none())
16612 && m.experts.iter().all(|d| d.act == Act::Silu);
16613 if !covered {
16614 per_row(&mut out);
16615 return out;
16616 }
16617 let ne = m.experts.len();
16618 let mut routes: Vec<(Vec<usize>, Vec<f32>)> = Vec::with_capacity(b);
16620 for r in 0..b {
16621 let x = &xs[r * hidden..(r + 1) * hidden];
16622 accumulate_act(m, x, 1);
16623 let mut logits = vec![0.0f32; ne];
16624 m.router.matvec(x, &mut logits, pool);
16625 let (idx, p, wsum) = moe_route(&logits, m, None);
16626 {
16627 let mut st = m.stats.borrow_mut();
16628 if st.len() < ne {
16629 st.resize(ne, 0);
16630 }
16631 for &e in &idx {
16632 st[e] += 1;
16633 }
16634 }
16635 let w: Vec<f32> = idx
16636 .iter()
16637 .map(|&e| p[e] / wsum * m.per_expert_scale.as_ref().map_or(1.0, |v| v[e]))
16638 .collect();
16639 routes.push((idx, w));
16640 }
16641 if routes.iter().any(|(idx, _)| idx.is_empty()) {
16642 per_row(&mut out);
16643 return out;
16644 }
16645 let mut experts: Vec<usize> = Vec::new();
16647 let mut groups: Vec<Vec<usize>> = Vec::new();
16648 for (r, (idx, _)) in routes.iter().enumerate() {
16649 for &e in idx {
16650 match experts.iter().position(|&x| x == e) {
16651 Some(g) => groups[g].push(r),
16652 None => {
16653 experts.push(e);
16654 groups.push(vec![r]);
16655 }
16656 }
16657 }
16658 }
16659 let n_pairs: usize = groups.iter().map(|g| g.len()).sum();
16660 let inter = m.experts[experts[0]].gate_proj.rows();
16661 let pairs: Vec<(&QTensor, &QTensor)> = experts
16662 .iter()
16663 .map(|&e| (&m.experts[e].gate_proj, &m.experts[e].up_proj))
16664 .collect();
16665 let mut gs: Vec<Vec<f32>> = (0..n_pairs).map(|_| vec![0f32; inter]).collect();
16666 if !QTensor::moe_gate_up_rows(&pairs, &groups, xs, &mut gs, pool) {
16667 per_row(&mut out);
16668 return out;
16669 }
16670 let downs: Vec<&QTensor> = experts.iter().map(|&e| &m.experts[e].down_proj).collect();
16671 let lens: Vec<usize> = groups.iter().map(|g| g.len()).collect();
16672 let mut ds: Vec<Vec<f32>> = (0..n_pairs).map(|_| vec![0f32; hidden]).collect();
16673 if !QTensor::moe_down_rows(&downs, &lens, &gs, &mut ds, pool) {
16674 per_row(&mut out);
16675 return out;
16676 }
16677 let mut slot = std::collections::HashMap::with_capacity(n_pairs);
16679 let mut p = 0usize;
16680 for (g, &e) in experts.iter().enumerate() {
16681 for &r in &groups[g] {
16682 slot.insert((r, e), p);
16683 p += 1;
16684 }
16685 }
16686 for (r, (idx, w)) in routes.iter().enumerate() {
16687 let terms: Vec<(&[f32], f32)> = idx
16688 .iter()
16689 .zip(w)
16690 .map(|(&e, &we)| (ds[slot[&(r, e)]].as_slice(), we))
16691 .collect();
16692 let row = &mut out[r * hidden..(r + 1) * hidden];
16693 for (i, dst) in row.iter_mut().enumerate() {
16694 let mut acc = 0f32;
16696 for (d, we) in &terms {
16697 acc += we * d[i];
16698 }
16699 *dst = acc;
16700 }
16701 }
16702 out
16703}
16704
16705thread_local! {
16706 static FFN_SCRATCH: std::cell::RefCell<[Vec<f32>; 4]> =
16710 const { std::cell::RefCell::new([Vec::new(), Vec::new(), Vec::new(), Vec::new()]) };
16711}
16712
16713fn dense_ffn(d: &DenseFfn, x: &[f32], pool: Option<&Pool>) -> Vec<f32> {
16715 if gate_topk() > 0
16718 && let Some(out) = dense_ffn_dynamic(d, x, pool, gate_topk())
16719 {
16720 return out;
16721 }
16722 let prism_body = d.gate_proj.has_prism_contract()
16738 || d.up_proj.has_prism_contract()
16739 || d.down_proj.has_prism_contract();
16740 if !prism_body
16741 && crate::gpu::enabled_here()
16742 && (d.gate_proj.rows() >= crate::gpu::min_rows() || d.gate_proj.is_q1())
16743 {
16744 let arm = if d.gate_proj.is_q1() && crate::gpu::q1_force() {
16745 crate::gpu::ProbeArm::Gpu
16746 } else {
16747 crate::gpu::probe_arm(crate::gpu::OpClass::Ffn)
16748 };
16749 match arm {
16750 crate::gpu::ProbeArm::Gpu => {
16751 let t0 = std::time::Instant::now();
16752 if let Some(out) = dense_ffn_gpu(d, x, pool) {
16753 crate::gpu::probe_record(crate::gpu::OpClass::Ffn, true, t0.elapsed());
16754 return out;
16755 }
16756 crate::gpu::probe_note_decline(crate::gpu::OpClass::Ffn);
16760 }
16761 crate::gpu::ProbeArm::CpuTimed => {
16762 let t0 = std::time::Instant::now();
16763 let out = crate::gpu::cpu_scope(|| dense_ffn_cpu(d, x, pool));
16764 crate::gpu::probe_record(crate::gpu::OpClass::Ffn, false, t0.elapsed());
16765 return out;
16766 }
16767 crate::gpu::ProbeArm::Cpu => {
16768 return crate::gpu::cpu_scope(|| dense_ffn_cpu(d, x, pool));
16769 }
16770 }
16771 }
16772 dense_ffn_cpu(d, x, pool)
16773}
16774
16775fn dense_ffn_cpu(d: &DenseFfn, x: &[f32], pool: Option<&Pool>) -> Vec<f32> {
16777 let inter = d.gate_proj.rows();
16778 FFN_SCRATCH.with(|s| {
16779 let mut s = s.borrow_mut();
16780 let [g, u, ..] = &mut *s;
16781 g.resize(inter, 0.0);
16782 if gate_topk() > 0 {
16785 u.resize(inter, 0.0);
16789 QTensor::matvec_many([&d.gate_proj, &d.up_proj], x, [g, u], pool);
16790 for i in 0..inter {
16791 g[i] = Act::Silu.combine(g[i], 1.0);
16792 }
16793 keep_top_k(g, gate_topk());
16794 for i in 0..inter {
16795 g[i] *= u[i];
16796 }
16797 } else if d.act == Act::Silu && {
16798 let _prof = crate::cpuprof::time(crate::cpuprof::Slot::FfnGateUp);
16799 QTensor::matvec_silu_mul(&d.gate_proj, &d.up_proj, x, g, pool)
16800 } {
16801 } else {
16803 u.resize(inter, 0.0);
16804 let _prof = crate::cpuprof::time(crate::cpuprof::Slot::FfnGateUp);
16806 QTensor::matvec_many([&d.gate_proj, &d.up_proj], x, [g, u], pool);
16807 for i in 0..inter {
16808 g[i] = d.act.combine(g[i], u[i]);
16809 }
16810 }
16811 FFN_PROBE.with(|pr| {
16819 if let Some(acc) = pr.borrow_mut().as_mut() {
16820 let li = crate::gpu::cur_layer();
16821 if li >= 0 {
16822 if let Some(row) = acc.get_mut(li as usize) {
16823 match probe_topk() {
16824 0 if probe_sq() => {
16825 for (a, &v) in row.iter_mut().zip(g.iter()) {
16826 *a += (v as f64) * (v as f64);
16827 }
16828 }
16829 0 if probe_signed() => {
16830 for (a, &v) in row.iter_mut().zip(g.iter()) {
16831 *a += v as f64;
16832 }
16833 }
16834 0 => {
16835 for (a, &v) in row.iter_mut().zip(g.iter()) {
16836 *a += (v as f64).abs();
16837 }
16838 }
16839 k => {
16840 let n = g.len();
16841 let k = k.min(n);
16842 let mut mag: Vec<f32> = g.iter().map(|v| v.abs()).collect();
16843 let (_, kth, _) = mag.select_nth_unstable_by(k - 1, |a, b| {
16844 b.partial_cmp(a).unwrap_or(std::cmp::Ordering::Equal)
16845 });
16846 let thr = *kth;
16847 for (a, &v) in row.iter_mut().zip(g.iter()) {
16848 if v.abs() >= thr {
16849 *a += 1.0;
16850 }
16851 }
16852 }
16853 }
16854 }
16855 }
16856 }
16857 });
16858 if oracle_topk() > 0 {
16859 keep_top_k(g, oracle_topk());
16860 }
16861 {
16862 let li = crate::gpu::cur_layer();
16863 if li >= 0 {
16864 adump_row(li as usize, g);
16865 }
16866 }
16867 let mut out = attention::take_buf(d.down_proj.rows());
16868 let _prof = crate::cpuprof::time(crate::cpuprof::Slot::FfnDown);
16869 d.down_proj.matvec(g, &mut out, pool);
16870 out
16871 })
16872}
16873
16874pub struct RefitAcc {
16887 pub support: Vec<u32>,
16888 pub gss: Vec<f32>,
16889 pub ya: Vec<f32>,
16890 pub hidden: usize,
16891 pub tokens: u64,
16892 pub buf_g: Vec<f32>,
16898 pub buf_o: Vec<f32>,
16899 pub buf_t: usize,
16900}
16901
16902type RefitState = (std::collections::HashMap<usize, RefitAcc>, Vec<f32>);
16906
16907static REFIT: std::sync::OnceLock<Option<(String, std::sync::Mutex<RefitState>)>> =
16908 std::sync::OnceLock::new();
16909
16910fn ffn_probe_active() -> bool {
16913 FFN_PROBE.with(|p| p.borrow().is_some())
16914}
16915
16916fn refit_dir() -> Option<&'static (String, std::sync::Mutex<RefitState>)> {
16917 REFIT
16918 .get_or_init(|| {
16919 std::env::var("CMF_FFN_REFIT").ok().map(|d| {
16920 (
16921 d,
16922 std::sync::Mutex::new((std::collections::HashMap::new(), Vec::new())),
16923 )
16924 })
16925 })
16926 .as_ref()
16927}
16928
16929fn refit_accumulate(
16931 li: usize,
16932 g: &[f32],
16933 b: usize,
16934 inter: usize,
16935 out: &[f32],
16936 hidden: usize,
16937 pool: Option<&Pool>,
16938) {
16939 let Some((dir, map)) = refit_dir() else {
16940 return;
16941 };
16942 static SPAN: std::sync::OnceLock<(usize, usize)> = std::sync::OnceLock::new();
16943 let (from, to) = *SPAN.get_or_init(|| {
16944 let g = |k: &str, d: usize| {
16945 std::env::var(k)
16946 .ok()
16947 .and_then(|v| v.parse().ok())
16948 .unwrap_or(d)
16949 };
16950 (
16951 g("CMF_FFN_REFIT_FROM", 0),
16952 g("CMF_FFN_REFIT_TO", usize::MAX),
16953 )
16954 });
16955 if li < from || li > to {
16956 return;
16957 }
16958 let mut guard = map.lock().unwrap();
16959 let (map, shared) = &mut *guard;
16960 let acc = match map.entry(li) {
16961 std::collections::hash_map::Entry::Occupied(e) => e.into_mut(),
16962 std::collections::hash_map::Entry::Vacant(e) => {
16963 let path = format!("{dir}/support.{li}.u32");
16964 let Ok(bytes) = std::fs::read(&path) else {
16965 eprintln!("refit: no {path} — layer {li} skipped");
16966 return;
16967 };
16968 let n = u32::from_le_bytes(bytes[0..4].try_into().unwrap()) as usize;
16969 let support: Vec<u32> = bytes[4..4 + n * 4]
16970 .chunks_exact(4)
16971 .map(|c| u32::from_le_bytes([c[0], c[1], c[2], c[3]]))
16972 .collect();
16973 eprintln!(
16974 "refit: layer {li} support {n} ({:.0} MB of accumulator)",
16975 (n * n + hidden * n) as f64 * 4.0 / 1e6
16976 );
16977 e.insert(RefitAcc {
16978 gss: vec![0.0; n * n],
16979 ya: vec![0.0; hidden * n],
16980 buf_g: Vec::new(),
16981 buf_o: Vec::new(),
16982 buf_t: 0,
16983 support,
16984 hidden,
16985 tokens: 0,
16986 })
16987 }
16988 };
16989 let ns = acc.support.len();
16990 let cap = refit_batch();
16992 if acc.buf_g.is_empty() {
16993 acc.buf_g = vec![0.0; ns * cap];
16994 acc.buf_o = vec![0.0; hidden * cap];
16995 }
16996 let take = b.min(cap - acc.buf_t);
16997 for t in 0..take {
16998 let col = acc.buf_t + t;
16999 for (j, &n) in acc.support.iter().enumerate() {
17000 acc.buf_g[j * cap + col] = g[t * inter + n as usize];
17001 }
17002 for h in 0..hidden {
17003 acc.buf_o[h * cap + col] = out[t * hidden + h];
17004 }
17005 }
17006 acc.buf_t += take;
17007 acc.tokens += take as u64;
17008 if acc.buf_t < cap {
17009 return;
17010 }
17011 let bt = acc.buf_t;
17012 acc.buf_t = 0;
17013 let RefitAcc {
17023 gss,
17024 ya,
17025 buf_g,
17026 buf_o,
17027 ..
17028 } = acc;
17029 let need = (ns * ns).max(hidden * ns);
17030 if shared.len() < need {
17031 shared.resize(need, 0.0);
17032 }
17033 let scratch = &mut shared[..];
17034 let _ = bt;
17035 if crate::gpu::gemm_nt_f32_transient(buf_g, buf_g, &mut scratch[..ns * ns], ns, cap, ns) {
17036 add_into(gss, &scratch[..ns * ns], pool);
17037 if crate::gpu::gemm_nt_f32_transient(
17038 buf_o,
17039 buf_g,
17040 &mut scratch[..hidden * ns],
17041 hidden,
17042 cap,
17043 ns,
17044 ) {
17045 add_into(ya, &scratch[..hidden * ns], pool);
17046 } else {
17047 accum_outer_t(ya, hidden, ns, cap, buf_o, buf_g, pool);
17048 }
17049 } else {
17050 accum_outer_t(gss, ns, ns, cap, buf_g, buf_g, pool);
17051 accum_outer_t(ya, hidden, ns, cap, buf_o, buf_g, pool);
17052 }
17053 }
17057
17058fn refit_batch() -> usize {
17060 static B: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
17061 *B.get_or_init(|| {
17062 std::env::var("CMF_FFN_REFIT_BATCH")
17063 .ok()
17064 .and_then(|v| v.parse().ok())
17065 .unwrap_or(4096)
17066 })
17067}
17068
17069fn accum_outer_t(
17072 c: &mut [f32],
17073 m: usize,
17074 n: usize,
17075 b: usize,
17076 left: &[f32],
17077 right: &[f32],
17078 pool: Option<&Pool>,
17079) {
17080 let ptr = SendMut(c.as_mut_ptr());
17081 let body = |i: usize| {
17082 let ptr = &ptr;
17083 let row = unsafe { std::slice::from_raw_parts_mut(ptr.0.add(i * n), n) };
17084 for t in 0..b {
17085 let a = left[i * b + t];
17086 if a == 0.0 {
17087 continue;
17088 }
17089 for (j, o) in row.iter_mut().enumerate() {
17090 *o += a * right[j * b + t];
17091 }
17092 }
17093 };
17094 match pool {
17095 Some(p) if m > 1 => p.run_rows(m, &|s, e| {
17096 for i in s..e {
17097 body(i);
17098 }
17099 }),
17100 _ => {
17101 for i in 0..m {
17102 body(i);
17103 }
17104 }
17105 }
17106}
17107
17108fn add_into(dst: &mut [f32], src: &[f32], pool: Option<&Pool>) {
17111 let n = dst.len().min(src.len());
17112 match pool {
17113 Some(p) if n >= 1 << 16 => {
17114 let ptr = SendMut(dst.as_mut_ptr());
17115 let f = |s: usize, e: usize| {
17116 let ptr = &ptr;
17117 for blk in s..e {
17118 let (a, b) = (blk * 4096, ((blk + 1) * 4096).min(n));
17119 for i in a..b {
17120 unsafe { *ptr.0.add(i) += src[i] };
17121 }
17122 }
17123 };
17124 p.run_rows(n.div_ceil(4096), &f);
17125 }
17126 _ => {
17127 for (d, v) in dst.iter_mut().zip(&src[..n]) {
17128 *d += *v;
17129 }
17130 }
17131 }
17132}
17133
17134fn accum_outer(
17139 c: &mut [f32],
17140 m: usize,
17141 n: usize,
17142 b: usize,
17143 left: &[f32],
17144 right: &[f32],
17145 pool: Option<&Pool>,
17146) {
17147 const TILE: usize = 32;
17148 let tiles = m.div_ceil(TILE);
17149 let cp = SendMut(c.as_mut_ptr());
17150 let body = |ti: usize| {
17151 let cp = &cp;
17152 let i0 = ti * TILE;
17153 let i1 = (i0 + TILE).min(m);
17154 for t in 0..b {
17155 let r = &right[t * n..t * n + n];
17156 for i in i0..i1 {
17157 let a = left[i * b + t];
17158 if a == 0.0 {
17159 continue;
17160 }
17161 let row = unsafe { std::slice::from_raw_parts_mut(cp.0.add(i * n), n) };
17163 for (o, v) in row.iter_mut().zip(r) {
17164 *o += a * *v;
17165 }
17166 }
17167 }
17168 };
17169 match pool {
17170 Some(p) if tiles > 1 => p.run_rows(tiles, &|s, e| {
17171 for ti in s..e {
17172 body(ti);
17173 }
17174 }),
17175 _ => {
17176 for ti in 0..tiles {
17177 body(ti);
17178 }
17179 }
17180 }
17181}
17182
17183pub fn refit_flush() -> usize {
17185 let Some((dir, map)) = refit_dir() else {
17186 return 0;
17187 };
17188 let guard = map.lock().unwrap();
17189 let mut n = 0;
17190 for (li, acc) in guard.0.iter() {
17191 let w = |name: &str, v: &[f32]| {
17194 let path = format!("{dir}/{name}.{li}.f32");
17195 let bytes: Vec<u8> = v.iter().flat_map(|x| x.to_le_bytes()).collect();
17196 match std::fs::write(&path, &bytes) {
17197 Ok(()) => {}
17198 Err(e) => eprintln!(
17199 "refit: FAILED to write {path} ({} MB): {e}",
17200 bytes.len() / 1_000_000
17201 ),
17202 }
17203 };
17204 w("gss", &acc.gss);
17205 w("ya", &acc.ya);
17206 println!(
17207 "refit L{li}: {} support, {} tokens, hidden {}",
17208 acc.support.len(),
17209 acc.tokens,
17210 acc.hidden
17211 );
17212 n += 1;
17213 }
17214 n
17215}
17216
17217fn adump_row(li: usize, g: &[f32]) {
17222 use std::io::Write as _;
17223 static FILES: std::sync::OnceLock<
17224 Option<(
17225 String,
17226 std::sync::Mutex<std::collections::HashMap<usize, std::fs::File>>,
17227 )>,
17228 > = std::sync::OnceLock::new();
17229 let Some((prefix, map)) = FILES
17230 .get_or_init(|| {
17231 std::env::var("CMF_FFN_ADUMP")
17232 .ok()
17233 .map(|p| (p, std::sync::Mutex::new(std::collections::HashMap::new())))
17234 })
17235 .as_ref()
17236 else {
17237 return;
17238 };
17239 static SPAN: std::sync::OnceLock<(usize, usize)> = std::sync::OnceLock::new();
17242 let (from, to) = *SPAN.get_or_init(|| {
17243 let g = |k: &str, d: usize| {
17244 std::env::var(k)
17245 .ok()
17246 .and_then(|v| v.parse().ok())
17247 .unwrap_or(d)
17248 };
17249 (
17250 g("CMF_FFN_ADUMP_FROM", 0),
17251 g("CMF_FFN_ADUMP_TO", usize::MAX),
17252 )
17253 });
17254 if li < from || li > to {
17255 return;
17256 }
17257 let mut map = map.lock().unwrap();
17258 let f = map.entry(li).or_insert_with(|| {
17259 std::fs::File::create(format!("{prefix}.{li}.f16")).expect("adump file")
17260 });
17261 let mut bytes = Vec::with_capacity(g.len() * 2);
17262 for v in g {
17263 bytes.extend_from_slice(&cortiq_core::quant::f32_to_f16(*v).to_le_bytes());
17264 }
17265 let _ = f.write_all(&bytes);
17266}
17267
17268fn oracle_topk() -> usize {
17274 static K: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
17275 *K.get_or_init(|| {
17276 std::env::var("CMF_FFN_ORACLE_TOPK")
17277 .ok()
17278 .and_then(|v| v.parse().ok())
17279 .unwrap_or(0)
17280 })
17281}
17282
17283fn gate_topk() -> usize {
17289 static K: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
17290 *K.get_or_init(|| {
17291 std::env::var("CMF_FFN_GATE_TOPK")
17292 .ok()
17293 .and_then(|v| v.parse().ok())
17294 .unwrap_or(0)
17295 })
17296}
17297
17298fn gate_block() -> usize {
17305 static B: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
17306 *B.get_or_init(|| {
17307 std::env::var("CMF_FFN_GATE_BLOCK")
17308 .ok()
17309 .and_then(|v| v.parse().ok())
17310 .unwrap_or(1)
17311 })
17312}
17313
17314fn keep_top_blocks(g: &mut [f32], keep_n: usize, block: usize) {
17316 let n = g.len();
17317 let nb = n.div_ceil(block);
17318 let kb = (keep_n.div_ceil(block)).clamp(1, nb);
17319 if kb >= nb {
17320 return;
17321 }
17322 let mut score: Vec<f32> = (0..nb)
17323 .map(|b| {
17324 g[b * block..((b + 1) * block).min(n)]
17325 .iter()
17326 .map(|v| v * v)
17327 .sum::<f32>()
17328 })
17329 .collect();
17330 let mut ord = score.clone();
17331 let (_, kth, _) = ord.select_nth_unstable_by(kb - 1, |a, b| {
17332 b.partial_cmp(a).unwrap_or(std::cmp::Ordering::Equal)
17333 });
17334 let thr = *kth;
17335 for b in 0..nb {
17336 if score[b] < thr {
17337 g[b * block..((b + 1) * block).min(n)].fill(0.0);
17338 }
17339 }
17340 score.clear();
17341}
17342
17343fn keep_top_k(g: &mut [f32], k: usize) {
17345 if gate_block() > 1 {
17346 return keep_top_blocks(g, k, gate_block());
17347 }
17348 let n = g.len();
17349 if k == 0 || k >= n {
17350 return;
17351 }
17352 let mut mag: Vec<f32> = g.iter().map(|v| v.abs()).collect();
17353 let (_, kth, _) = mag.select_nth_unstable_by(k - 1, |a, b| {
17354 b.partial_cmp(a).unwrap_or(std::cmp::Ordering::Equal)
17355 });
17356 let thr = *kth;
17357 for v in g.iter_mut() {
17358 if v.abs() < thr {
17359 *v = 0.0;
17360 }
17361 }
17362}
17363
17364fn probe_sq() -> bool {
17368 static S: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
17369 *S.get_or_init(|| std::env::var("CMF_FFN_PROBE_SQ").is_ok())
17370}
17371
17372fn probe_signed() -> bool {
17376 static S: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
17377 *S.get_or_init(|| std::env::var("CMF_FFN_PROBE_SIGNED").is_ok())
17378}
17379
17380fn meanfill() -> Option<&'static (usize, Vec<f32>)> {
17388 static M: std::sync::OnceLock<Option<(usize, Vec<f32>)>> = std::sync::OnceLock::new();
17389 M.get_or_init(|| {
17390 let p = std::env::var("CMF_FFN_MEANFILL").ok()?;
17391 let b = std::fs::read(&p).ok()?;
17392 let inter = u32::from_le_bytes(b[4..8].try_into().ok()?) as usize;
17393 let vals: Vec<f32> = b[8..]
17394 .chunks_exact(4)
17395 .map(|c| f32::from_le_bytes([c[0], c[1], c[2], c[3]]))
17396 .collect();
17397 eprintln!("meanfill: {} value(s), inter {inter}", vals.len());
17398 Some((inter, vals))
17399 })
17400 .as_ref()
17401}
17402
17403fn probe_topk() -> usize {
17406 static K: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
17407 *K.get_or_init(|| {
17408 std::env::var("CMF_FFN_PROBE_TOPK")
17409 .ok()
17410 .and_then(|v| v.parse().ok())
17411 .unwrap_or(0)
17412 })
17413}
17414
17415thread_local! {
17416 static FFN_PROBE: std::cell::RefCell<Option<Vec<Vec<f64>>>> =
17419 const { std::cell::RefCell::new(None) };
17420}
17421
17422fn dense_ffn_dynamic(d: &DenseFfn, x: &[f32], pool: Option<&Pool>, k: usize) -> Option<Vec<f32>> {
17435 if d.gate_proj.has_prism_contract()
17440 || d.up_proj.has_prism_contract()
17441 || d.down_proj.has_prism_contract()
17442 {
17443 return None;
17444 }
17445 let dt = d.down_t.as_ref()?;
17446 let inter = d.gate_proj.rows();
17447 let hidden = dt.cols();
17448 if k == 0 || k >= inter || d.act != Act::Silu {
17449 return None;
17450 }
17451 DYN_SCRATCH.with(|sc| {
17452 let mut sc = sc.borrow_mut();
17453 let DynScratch {
17454 g,
17455 mag,
17456 live,
17457 parts,
17458 } = &mut *sc;
17459 g.resize(inter, 0.0);
17460 d.gate_proj.matvec(x, g, pool);
17461 for v in g.iter_mut() {
17462 *v = inference::silu(*v);
17463 }
17464 mag.clear();
17467 mag.extend(g.iter().map(|v| v.abs()));
17468 let (_, kth, _) = mag.select_nth_unstable_by(k - 1, |a, b| {
17469 b.partial_cmp(a).unwrap_or(std::cmp::Ordering::Equal)
17470 });
17471 let thr = *kth;
17472 live.clear();
17473 live.extend((0..inter as u32).filter(|&n| g[n as usize].abs() >= thr));
17474 let mut out = vec![0.0f32; hidden];
17475 match pool {
17476 Some(p) if live.len() >= 64 => {
17477 let nw = p.n_workers() + 1;
17478 parts.clear();
17479 parts.resize(nw * hidden, 0.0);
17480 let ptr = SendMut(parts.as_mut_ptr());
17481 let n = live.len();
17482 let live_ref: &[u32] = live;
17483 let g_ref: &[f32] = g;
17484 p.run(&|w, workers| {
17485 let chunk = n.div_ceil(workers);
17486 let (s, e) = (w * chunk, ((w + 1) * chunk).min(n));
17487 if s >= e {
17488 return;
17489 }
17490 WORKER_SCRATCH.with(|ws| {
17491 let mut ws = ws.borrow_mut();
17492 let [scratch, acc] = &mut *ws;
17493 scratch.resize(hidden.max(x.len()), 0.0);
17494 acc.clear();
17495 acc.resize(hidden, 0.0);
17496 for (o, &nrm) in live_ref[s..e].iter().enumerate() {
17497 if let Some(&nx) = live_ref[s..e].get(o + 1) {
17500 d.up_proj.prefetch_row(nx as usize);
17501 dt.prefetch_row(nx as usize);
17502 }
17503 let idx = nrm as usize;
17504 let up = d.up_proj.row_dot(idx, x, scratch);
17505 let a = g_ref[idx] * up;
17506 if a != 0.0 {
17507 dt.add_row_scaled(idx, a, acc, scratch);
17508 }
17509 }
17510 for (j, v) in acc.iter().enumerate() {
17511 unsafe { *ptr.at(w * hidden + j) = *v };
17512 }
17513 });
17514 });
17515 for w in 0..nw {
17516 for (j, o) in out.iter_mut().enumerate() {
17517 *o += parts[w * hidden + j];
17518 }
17519 }
17520 }
17521 _ => {
17522 WORKER_SCRATCH.with(|ws| {
17523 let mut ws = ws.borrow_mut();
17524 let [scratch, _acc] = &mut *ws;
17525 scratch.resize(hidden.max(x.len()), 0.0);
17526 for &nrm in live.iter() {
17527 let idx = nrm as usize;
17528 let up = d.up_proj.row_dot(idx, x, scratch);
17529 let a = g[idx] * up;
17530 if a != 0.0 {
17531 dt.add_row_scaled(idx, a, &mut out, scratch);
17532 }
17533 }
17534 });
17535 }
17536 }
17537 Some(out)
17538 })
17539}
17540
17541struct DynScratch {
17544 g: Vec<f32>,
17545 mag: Vec<f32>,
17546 live: Vec<u32>,
17547 parts: Vec<f32>,
17548}
17549
17550thread_local! {
17551 static DYN_SCRATCH: std::cell::RefCell<DynScratch> = const {
17552 std::cell::RefCell::new(DynScratch {
17553 g: Vec::new(),
17554 mag: Vec::new(),
17555 live: Vec::new(),
17556 parts: Vec::new(),
17557 })
17558 };
17559 static WORKER_SCRATCH: std::cell::RefCell<[Vec<f32>; 2]> =
17561 const { std::cell::RefCell::new([Vec::new(), Vec::new()]) };
17562}
17563
17564fn dense_ffn_masked(d: &DenseFfn, x: &[f32], pool: Option<&Pool>, mask_row: &[u8]) -> Vec<f32> {
17569 let inter = d.gate_proj.rows();
17570 FFN_SCRATCH.with(|s| {
17571 let mut s = s.borrow_mut();
17572 let [g, u, ..] = &mut *s;
17573 g.resize(inter, 0.0);
17574 if d.act == Act::Silu && QTensor::matvec_silu_mul(&d.gate_proj, &d.up_proj, x, g, pool) {
17575 } else {
17577 u.resize(inter, 0.0);
17578 QTensor::matvec_many([&d.gate_proj, &d.up_proj], x, [g, u], pool);
17579 for i in 0..inter {
17580 g[i] = d.act.combine(g[i], u[i]);
17581 }
17582 }
17583 zero_masked_cols(g, 1, inter, mask_row);
17584 let mut out = attention::take_buf(d.down_proj.rows());
17585 d.down_proj.matvec(g, &mut out, pool);
17586 out
17587 })
17588}
17589
17590fn dense_ffn_gpu(d: &DenseFfn, x: &[f32], _pool: Option<&Pool>) -> Option<Vec<f32>> {
17596 if d.gate_proj.has_prism_contract()
17597 || d.up_proj.has_prism_contract()
17598 || d.down_proj.has_prism_contract()
17599 {
17600 return None;
17601 }
17602 if d.act != Act::Silu {
17604 return None;
17605 }
17606 if d.gate_proj.rows() < crate::gpu::min_rows() && !d.gate_proj.is_q1() {
17609 return None;
17610 }
17611 let mut jobs: Vec<crate::gpu::MoeJob> = Vec::with_capacity(1);
17612 let mut model_ref = None;
17613 moe_push_job(d, x, 1.0, &mut jobs, &mut model_ref)?;
17614 let model = model_ref?;
17615 let hidden = jobs[0].down.1;
17616 let mut out = attention::take_buf(hidden);
17617 if crate::gpu::moe_block(&model, &jobs, &mut out) {
17618 Some(out)
17619 } else {
17620 let mut out = out;
17621 attention::recycle_buf(&mut out);
17622 None
17623 }
17624}
17625
17626#[allow(clippy::type_complexity)]
17631#[allow(clippy::type_complexity)]
17632pub(crate) fn moe_parts(
17633 t: &QTensor,
17634) -> Option<(
17635 &std::sync::Arc<cortiq_core::CmfModel>,
17636 usize,
17637 usize,
17638 usize,
17639 &[f32],
17640 &[f32],
17641 bool,
17642 bool,
17643 bool,
17644)> {
17645 match t {
17646 QTensor::Mapped {
17647 model,
17648 idx,
17649 dtype: dt @ (cortiq_core::TensorDtype::Q8_2f | cortiq_core::TensorDtype::Q8Row),
17650 rows,
17651 cols,
17652 row_scale,
17653 col_field,
17654 ..
17655 } if (*dt == cortiq_core::TensorDtype::Q8Row) || !col_field.is_empty() => Some((
17656 model, *idx, *rows, *cols, row_scale, col_field, false, false, false,
17657 )),
17658 QTensor::Mapped {
17660 model,
17661 idx,
17662 dtype: cortiq_core::TensorDtype::Q1,
17663 rows,
17664 cols,
17665 ..
17666 } => Some((
17667 model,
17668 *idx,
17669 *rows,
17670 *cols,
17671 &[][..],
17672 &[][..],
17673 true,
17674 false,
17675 false,
17676 )),
17677 QTensor::Mapped {
17679 model,
17680 idx,
17681 dtype: cortiq_core::TensorDtype::Q4Tiled,
17682 rows,
17683 cols,
17684 ..
17685 } => Some((
17686 model,
17687 *idx,
17688 *rows,
17689 *cols,
17690 &[][..],
17691 &[][..],
17692 false,
17693 true,
17694 false,
17695 )),
17696 QTensor::Mapped {
17698 model,
17699 idx,
17700 dtype: cortiq_core::TensorDtype::Q4TiledP,
17701 rows,
17702 cols,
17703 ..
17704 } => Some((
17705 model,
17706 *idx,
17707 *rows,
17708 *cols,
17709 &[][..],
17710 &[][..],
17711 false,
17712 true,
17713 false,
17714 )),
17715 QTensor::Mapped {
17719 model,
17720 idx,
17721 dtype: cortiq_core::TensorDtype::Q2TiledP,
17722 rows,
17723 cols,
17724 ..
17725 } => Some((
17726 model,
17727 *idx,
17728 *rows,
17729 *cols,
17730 &[][..],
17731 &[][..],
17732 false,
17733 true,
17734 true,
17735 )),
17736 _ => None,
17737 }
17738}
17739
17740#[cfg(target_os = "macos")]
17748fn metal_moe_graph_parts(m: &MoeFfn, hidden: usize) -> Option<crate::gpu::GpuMoe<'_>> {
17749 if m.router_input_norm
17750 || m.route_tau.is_some()
17751 || m.mask.is_some()
17752 || m.per_expert_scale.is_some()
17753 || m.experts.is_empty()
17754 || m.top_k == 0
17755 || m.resonance.is_some()
17756 {
17757 return None;
17758 }
17759 let (sh, sg) = match &m.shared {
17762 Some((sh, sg)) => (sh, sg.as_ref()),
17763 None => return None,
17764 };
17765 let (rf, rr, rc) = m.router.f32_parts()?;
17766 if rr != m.experts.len() || rc != hidden {
17767 return None;
17768 }
17769 let shared_gated = sg.is_some();
17770 let sf = match sg {
17771 Some(sg) => {
17772 let (sf, sr, sc) = sg.f32_parts()?;
17773 if sr * sc != hidden {
17774 return None;
17775 }
17776 sf
17777 }
17778 None => &rf[..hidden],
17781 };
17782 if let Some(b) = &m.expert_bias {
17783 if b.len() != m.experts.len() {
17784 return None;
17785 }
17786 }
17787 let inter = m.experts[0].gate_proj.rows();
17788 let gu_q2 = m.experts[0].gate_proj.mapped_q2tp().is_some();
17791 let trio = |e: &DenseFfn| -> Option<(usize, usize, usize)> {
17792 if e.act != Act::Silu
17793 || e.gate_proj.rows() != inter
17794 || e.gate_proj.cols() != hidden
17795 || e.up_proj.rows() != inter
17796 || e.up_proj.cols() != hidden
17797 || e.down_proj.rows() != hidden
17798 || e.down_proj.cols() != inter
17799 {
17800 return None;
17801 }
17802 let pick = |t: &QTensor| -> Option<usize> {
17803 if gu_q2 {
17804 t.mapped_q2tp().map(|(_, i)| i)
17805 } else {
17806 t.mapped_q4tp().map(|(_, i)| i)
17807 }
17808 };
17809 Some((
17810 pick(&e.gate_proj)?,
17811 pick(&e.up_proj)?,
17812 e.down_proj.mapped_q4tp().map(|(_, i)| i)?,
17813 ))
17814 };
17815 let experts = m.experts.iter().map(trio).collect::<Option<Vec<_>>>()?;
17816 let shared = trio(sh)?;
17817 Some(crate::gpu::GpuMoe {
17818 router: rf,
17819 sgate: sf,
17820 experts,
17821 shared,
17822 n_exp: m.experts.len(),
17823 top_k: m.top_k,
17824 inter,
17825 norm_topk: m.norm_topk_prob,
17826 route_scale: m.routed_scaling,
17827 gu_q2,
17828 sigmoid: m.router_sigmoid,
17829 bias: m.expert_bias.as_deref(),
17830 shared_gated,
17831 })
17832}
17833
17834pub(crate) fn moe_push_job_parts<'a>(
17838 gate: &'a QTensor,
17839 up: &'a QTensor,
17840 down: &'a QTensor,
17841 x: &[f32],
17842 w: f32,
17843 swiglu_limit: f32,
17844 jobs: &mut Vec<crate::gpu::MoeJob<'a>>,
17845 model_ref: &mut Option<std::sync::Arc<cortiq_core::CmfModel>>,
17846) -> Option<()> {
17847 use crate::qtensor::prescale;
17848 let (gm, gi, gr, gc, grs, gcf, gq1, gq4, gq2) = moe_parts(gate)?;
17849 let (_, ui, ur, uc, urs, ucf, uq1, uq4, uq2) = moe_parts(up)?;
17850 let (_, di, dr, dc, drs, dcf, dq1, dq4, dq2) = moe_parts(down)?;
17851 if gq1 != uq1 || uq1 != dq1 || gq4 != uq4 || uq4 != dq4 || gq2 != uq2 {
17852 return None; }
17854 if gq2 && (dq2 || !dq4 || down.mapped_q4tp().is_none()) {
17857 return None;
17858 }
17859 if !gq2 && dq2 {
17860 return None;
17861 }
17862 model_ref.get_or_insert_with(|| gm.clone());
17863 let dt = |cf: &[f32]| {
17864 if cf.is_empty() {
17865 cortiq_core::TensorDtype::Q8Row
17866 } else {
17867 cortiq_core::TensorDtype::Q8_2f
17868 }
17869 };
17870 jobs.push(crate::gpu::MoeJob {
17871 gate: (gi, gr, gc, grs),
17872 up: (ui, ur, uc, urs),
17873 down: (di, dr, dc, drs),
17874 xs_gate: prescale(x, gcf, dt(gcf)).into_owned(),
17875 xs_up: prescale(x, ucf, dt(ucf)).into_owned(),
17876 down_col: dcf,
17877 w,
17878 q1: gq1,
17879 q4t: gq4 && !gq2 && gate.mapped_q4tp().is_none(),
17880 q4tp: gq4 && (gq2 || gate.mapped_q4tp().is_some()),
17881 gu_q2: gq2,
17882 swiglu_limit,
17883 });
17884 Some(())
17885}
17886
17887fn moe_push_job<'a>(
17889 d: &'a DenseFfn,
17890 x: &[f32],
17891 w: f32,
17892 jobs: &mut Vec<crate::gpu::MoeJob<'a>>,
17893 model_ref: &mut Option<std::sync::Arc<cortiq_core::CmfModel>>,
17894) -> Option<()> {
17895 use crate::qtensor::prescale;
17896 if d.act != Act::Silu {
17897 return None; }
17899 let (gm, gi, gr, gc, grs, gcf, gq1, gq4, gq2) = moe_parts(&d.gate_proj)?;
17900 let (_, ui, ur, uc, urs, ucf, uq1, uq4, uq2) = moe_parts(&d.up_proj)?;
17901 let (_, di, dr, dc, drs, dcf, dq1, dq4, dq2) = moe_parts(&d.down_proj)?;
17902 if gq1 != uq1 || uq1 != dq1 || gq4 != uq4 || uq4 != dq4 || gq2 != uq2 {
17903 return None; }
17905 if gq2 && (dq2 || !dq4 || d.down_proj.mapped_q4tp().is_none()) {
17906 return None;
17907 }
17908 if !gq2 && dq2 {
17909 return None;
17910 }
17911 model_ref.get_or_insert_with(|| gm.clone());
17912 let gdt = if gcf.is_empty() {
17913 cortiq_core::TensorDtype::Q8Row
17914 } else {
17915 cortiq_core::TensorDtype::Q8_2f
17916 };
17917 let udt = if ucf.is_empty() {
17918 cortiq_core::TensorDtype::Q8Row
17919 } else {
17920 cortiq_core::TensorDtype::Q8_2f
17921 };
17922 jobs.push(crate::gpu::MoeJob {
17923 gate: (gi, gr, gc, grs),
17924 up: (ui, ur, uc, urs),
17925 down: (di, dr, dc, drs),
17926 xs_gate: prescale(x, gcf, gdt).into_owned(),
17927 xs_up: prescale(x, ucf, udt).into_owned(),
17928 down_col: dcf,
17929 w,
17930 q1: gq1,
17931 q4t: gq4 && !gq2 && d.gate_proj.mapped_q4tp().is_none(),
17932 q4tp: gq4 && (gq2 || d.gate_proj.mapped_q4tp().is_some()),
17933 gu_q2: gq2,
17934 swiglu_limit: 0.0,
17935 });
17936 Some(())
17937}
17938
17939fn sparse_ffn_quant(
17946 d: &DenseFfn,
17947 x: &[f32],
17948 active: &[u16],
17949 hidden: usize,
17950 pool: Option<&Pool>,
17951) -> Vec<f32> {
17952 let n = active.len();
17953 let inter = d.gate_proj.rows();
17954 let mut act = vec![0.0f32; n];
17955 let need_scratch = !(d.gate_proj.sparse_col_ok() && d.up_proj.sparse_col_ok());
17958 let compute = |ai: usize| -> f32 {
17959 let idx = active[ai] as usize;
17960 if idx >= inter {
17961 return 0.0; }
17963 let mut s = if need_scratch {
17964 vec![0.0f32; hidden]
17965 } else {
17966 Vec::new()
17967 };
17968 let gate = d.gate_proj.row_dot(idx, x, &mut s);
17969 let up = d.up_proj.row_dot(idx, x, &mut s);
17970 d.act.combine(gate, up)
17971 };
17972 match pool {
17973 Some(p) if n >= 256 => {
17974 let ptr = SendMut(act.as_mut_ptr());
17975 p.run(&|widx, nw| {
17976 let chunk = n.div_ceil(nw);
17977 let (s, e) = (widx * chunk, ((widx + 1) * chunk).min(n));
17978 for ai in s..e {
17979 unsafe { *ptr.at(ai) = compute(ai) };
17980 }
17981 });
17982 }
17983 _ => {
17984 for (ai, a) in act.iter_mut().enumerate() {
17985 *a = compute(ai);
17986 }
17987 }
17988 }
17989 let mut out = vec![0.0f32; hidden];
17991 for (ai, &idx) in active.iter().enumerate() {
17992 let w = act[ai];
17993 if w.abs() >= 1e-12 && (idx as usize) < inter {
17994 d.down_proj.add_col_scaled(idx as usize, w, &mut out);
17995 }
17996 }
17997 out
17998}
17999
18000#[doc(hidden)]
18002pub fn sparse_ffn_quant_for_test(
18003 d: &DenseFfn,
18004 x: &[f32],
18005 active: &[u16],
18006 hidden: usize,
18007) -> Vec<f32> {
18008 sparse_ffn_quant(d, x, active, hidden, None)
18009}
18010
18011fn dequant_dense_f32(d: &DenseFfn) -> (Vec<f32>, Vec<f32>, Vec<f32>) {
18015 let deq = |t: &QTensor| -> Vec<f32> {
18016 let (rows, cols) = (t.rows(), t.cols());
18017 let mut out = vec![0.0f32; rows * cols];
18018 for r in 0..rows {
18019 t.row_f32(r, &mut out[r * cols..(r + 1) * cols]);
18020 }
18021 out
18022 };
18023 (deq(&d.gate_proj), deq(&d.up_proj), deq(&d.down_proj))
18024}
18025
18026struct SendMut(*mut f32);
18028unsafe impl Send for SendMut {}
18029unsafe impl Sync for SendMut {}
18030impl SendMut {
18031 #[inline]
18032 #[allow(clippy::mut_from_ref)]
18035 unsafe fn at(&self, i: usize) -> &mut f32 {
18036 unsafe { &mut *self.0.add(i) }
18037 }
18038}
18039
18040pub(crate) fn moe_route(
18053 logits: &[f32],
18054 m: &MoeFfn,
18055 allowed: Option<&[bool]>,
18056) -> (Vec<usize>, Vec<f32>, f32) {
18057 moe_route_with_eps(logits, m, allowed, 1e-6)
18058}
18059
18060pub(crate) fn moe_route_with_eps(
18069 logits: &[f32],
18070 m: &MoeFfn,
18071 allowed: Option<&[bool]>,
18072 sigmoid_denom_eps: f32,
18073) -> (Vec<usize>, Vec<f32>, f32) {
18074 let ne = logits.len();
18075 let admit = |e: usize| {
18081 m.mask.as_ref().is_none_or(|mk| mk[e])
18082 && allowed.is_none_or(|a| a.get(e).copied().unwrap_or(false))
18083 };
18084 if m.resonance.is_some() && m.top_k == 1 {
18097 let mut best: Option<usize> = None;
18098 for e in (0..ne).filter(|&e| admit(e)) {
18099 let l = logits[e];
18100 if l == f32::NEG_INFINITY || l.is_nan() {
18101 continue;
18102 }
18103 if best.is_none_or(|b| l > logits[b]) {
18104 best = Some(e);
18105 }
18106 }
18107 if let Some(b) = best {
18108 let mut p = vec![0.0f32; ne];
18109 p[b] = 1.0;
18110 return (vec![b], p, 1.0 / m.routed_scaling);
18111 }
18112 }
18113 let p: Vec<f32> = if m.router_sigmoid {
18119 logits.iter().map(|&l| 1.0 / (1.0 + (-l).exp())).collect()
18120 } else {
18121 let mx = logits.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
18122 if mx == f32::NEG_INFINITY {
18123 vec![1.0 / ne.max(1) as f32; ne]
18124 } else {
18125 let mut e: Vec<f32> = logits.iter().map(|&l| (l - mx).exp()).collect();
18126 let s: f32 = e.iter().sum();
18127 for v in &mut e {
18128 *v /= s;
18129 }
18130 e
18131 }
18132 };
18133 let mut idx: Vec<usize> = (0..ne).filter(|&e| admit(e)).collect();
18134 match &m.expert_bias {
18136 Some(b) => idx.sort_unstable_by(|&x, &y| {
18137 (p[y] + b[y])
18138 .partial_cmp(&(p[x] + b[x]))
18139 .unwrap()
18140 .then(x.cmp(&y))
18141 }),
18142 None => idx.sort_unstable_by(|&x, &y| p[y].partial_cmp(&p[x]).unwrap().then(x.cmp(&y))),
18143 }
18144 idx.truncate(m.top_k);
18145 if let Some(tau) = m.route_tau {
18149 let total: f32 = idx.iter().map(|&e| p[e]).sum();
18150 if total > 0.0 {
18151 let mut acc = 0.0f32;
18152 let mut keep = idx.len();
18153 for (i, &e) in idx.iter().enumerate() {
18154 acc += p[e];
18155 if acc >= tau * total {
18156 keep = i + 1;
18157 break;
18158 }
18159 }
18160 idx.truncate(keep);
18161 }
18162 }
18163 let wsum: f32 = if m.norm_topk_prob {
18164 let s: f32 = idx.iter().map(|&e| p[e]).sum();
18165 (if m.router_sigmoid {
18169 s + sigmoid_denom_eps
18170 } else {
18171 s
18172 }) / m.routed_scaling
18173 } else {
18174 1.0 / m.routed_scaling
18175 };
18176 (idx, p, wsum)
18177}
18178
18179fn moe_trace(idx: &[usize]) {
18181 moe_trace_at(crate::gpu::cur_layer() as i32, idx)
18182}
18183
18184pub(crate) fn moe_trace_at(li: i32, idx: &[usize]) {
18187 use std::io::Write;
18188 static F: std::sync::OnceLock<Option<std::sync::Mutex<std::fs::File>>> =
18189 std::sync::OnceLock::new();
18190 let Some(f) = F.get_or_init(|| {
18191 let p = std::env::var("CMF_MOE_TRACE").ok()?;
18192 Some(std::sync::Mutex::new(
18193 std::fs::OpenOptions::new()
18194 .create(true)
18195 .append(true)
18196 .open(p)
18197 .ok()?,
18198 ))
18199 }) else {
18200 return;
18201 };
18202 let ids: Vec<String> = idx.iter().map(|e| e.to_string()).collect();
18203 let _ = writeln!(f.lock().unwrap(), "{li}:{}", ids.join(","));
18204}
18205
18206pub(crate) fn moe_ffn(
18209 m: &MoeFfn,
18210 x: &[f32],
18211 pool: Option<&Pool>,
18212 allowed: Option<&[bool]>,
18213) -> Vec<f32> {
18214 let r = moe_ffn_route(m, x, pool, allowed);
18215 moe_ffn_experts(m, x, &r, pool)
18216}
18217
18218pub(crate) struct MoeRoute {
18222 pub idx: Vec<usize>,
18223 pub p: Vec<f32>,
18224 pub wsum: f32,
18225 pub logits: Vec<f32>,
18226}
18227
18228pub(crate) fn moe_ffn_route(
18234 m: &MoeFfn,
18235 x: &[f32],
18236 pool: Option<&Pool>,
18237 allowed: Option<&[bool]>,
18238) -> MoeRoute {
18239 accumulate_act(m, x, 1);
18240 let ne = m.experts.len();
18241 let mut logits = vec![0.0f32; ne];
18242 match &m.resonance {
18243 Some(r) => r.scores(x, &mut logits),
18244 None => m.router.matvec(x, &mut logits, pool),
18245 }
18246 let (idx, p, wsum) = moe_route(&logits, m, allowed);
18247 {
18248 let mut st = m.stats.borrow_mut();
18249 if st.len() < ne {
18250 st.resize(ne, 0);
18251 }
18252 for &e in &idx {
18253 st[e] += 1;
18254 }
18255 }
18256 moe_trace(&idx);
18262 MoeRoute {
18263 idx,
18264 p,
18265 wsum,
18266 logits,
18267 }
18268}
18269
18270pub(crate) fn moe_ffn_experts(
18273 m: &MoeFfn,
18274 x: &[f32],
18275 r: &MoeRoute,
18276 pool: Option<&Pool>,
18277) -> Vec<f32> {
18278 let (idx, p, wsum) = (&r.idx, &r.p, r.wsum);
18279 if crate::gpu::enabled_here() {
18284 match crate::gpu::probe_arm(crate::gpu::OpClass::Ffn) {
18285 crate::gpu::ProbeArm::Gpu => {
18286 let t0 = std::time::Instant::now();
18287 if let Some(out) = moe_ffn_gpu(m, x, &idx, &p, wsum, pool) {
18288 crate::gpu::probe_record(crate::gpu::OpClass::Ffn, true, t0.elapsed());
18289 return out;
18290 }
18291 }
18292 crate::gpu::ProbeArm::CpuTimed => {
18293 let t0 = std::time::Instant::now();
18294 let out = crate::gpu::cpu_scope(|| moe_ffn_cpu(m, x, &idx, &p, wsum, pool));
18295 crate::gpu::probe_record(crate::gpu::OpClass::Ffn, false, t0.elapsed());
18296 return out;
18297 }
18298 crate::gpu::ProbeArm::Cpu => {
18299 return crate::gpu::cpu_scope(|| moe_ffn_cpu(m, x, &idx, &p, wsum, pool));
18300 }
18301 }
18302 }
18303 moe_ffn_cpu(m, x, &idx, &p, wsum, pool)
18304}
18305
18306fn moe_ffn_banked(
18310 slot: &mut crate::mimo_moe::Slot,
18311 li: usize,
18312 m: &MoeFfn,
18313 x: &[f32],
18314 pool: Option<&Pool>,
18315) -> Vec<f32> {
18316 let t0 = std::time::Instant::now();
18317 let r = moe_ffn_route(m, x, pool, None);
18318 slot.note_route(t0.elapsed().as_nanos() as u64);
18319 match slot.forward(li, m, x, &r, pool) {
18320 Some(out) => out,
18321 None => crate::qtensor::float_activations_scope(|| {
18322 crate::gpu::cpu_scope(|| moe_ffn_experts(m, x, &r, pool))
18323 }),
18324 }
18325}
18326
18327fn moe_ffn_banked_rows(
18329 slot: &mut crate::mimo_moe::Slot,
18330 li: usize,
18331 m: &MoeFfn,
18332 xs: &[f32],
18333 b: usize,
18334 hidden: usize,
18335 pool: Option<&Pool>,
18336) -> Vec<f32> {
18337 let t0 = std::time::Instant::now();
18338 let routes: Vec<_> = xs
18339 .chunks_exact(hidden)
18340 .map(|x| moe_ffn_route(m, x, pool, None))
18341 .collect();
18342 slot.note_route(t0.elapsed().as_nanos() as u64);
18343 if let Some(out) = slot.forward_rows(li, m, xs, &routes, pool) {
18344 return out;
18345 }
18346 let mut out = Vec::with_capacity(b * hidden);
18347 for (x, r) in xs.chunks_exact(hidden).zip(&routes) {
18348 let row = slot.forward(li, m, x, r, pool).unwrap_or_else(|| {
18349 crate::qtensor::float_activations_scope(|| {
18351 crate::gpu::cpu_scope(|| moe_ffn_experts(m, x, r, pool))
18352 })
18353 });
18354 out.extend(row);
18355 }
18356 out
18357}
18358
18359fn graph_note(built: bool, layers_run: usize, total_layers: usize) {
18364 use std::sync::atomic::{AtomicBool, Ordering};
18365 if built {
18366 GRAPH_TOK_OK.fetch_add(1, Ordering::Relaxed);
18367 if total_layers > 0 && layers_run < total_layers {
18368 GRAPH_TOK_PREFIX.fetch_add(1, Ordering::Relaxed);
18369 } else {
18370 GRAPH_TOK_FULL.fetch_add(1, Ordering::Relaxed);
18371 }
18372 } else {
18373 GRAPH_TOK_MISS.fetch_add(1, Ordering::Relaxed);
18374 }
18375 static SAID: AtomicBool = AtomicBool::new(false);
18376 if !SAID.swap(true, Ordering::Relaxed) {
18377 if built {
18378 tracing::info!("wgpu whole-token graph: ACTIVE");
18379 } else {
18380 tracing::warn!("wgpu whole-token graph refused — per-op path");
18381 }
18382 }
18383}
18384
18385pub static GRAPH_TOK_OK: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
18389pub static GRAPH_TOK_MISS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
18390pub static GRAPH_TOK_PREFIX: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
18394pub static GRAPH_TOK_FULL: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
18396
18397pub static METAL_GRAPH_TOK_OK: std::sync::atomic::AtomicU64 =
18401 std::sync::atomic::AtomicU64::new(0);
18402pub static METAL_GRAPH_HEAD_OK: std::sync::atomic::AtomicU64 =
18403 std::sync::atomic::AtomicU64::new(0);
18404pub static METAL_GRAPH_HEAD_MISS: std::sync::atomic::AtomicU64 =
18405 std::sync::atomic::AtomicU64::new(0);
18406pub static METAL_GRAPH_LAYERS: std::sync::atomic::AtomicU64 =
18407 std::sync::atomic::AtomicU64::new(0);
18408pub static METAL_GRAPH_ERRORS: std::sync::atomic::AtomicU64 =
18409 std::sync::atomic::AtomicU64::new(0);
18410pub static METAL_PREFILL_CHUNKS: std::sync::atomic::AtomicU64 =
18414 std::sync::atomic::AtomicU64::new(0);
18415pub static METAL_PREFILL_ROWS: std::sync::atomic::AtomicU64 =
18416 std::sync::atomic::AtomicU64::new(0);
18417pub static METAL_PREFILL_HEAD_ROWS: std::sync::atomic::AtomicU64 =
18418 std::sync::atomic::AtomicU64::new(0);
18419pub static METAL_PREFILL_ERRORS: std::sync::atomic::AtomicU64 =
18420 std::sync::atomic::AtomicU64::new(0);
18421
18422fn moe_batch_enabled() -> bool {
18425 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
18426 *ON.get_or_init(|| std::env::var("CMF_MOE_BATCH").as_deref() != Ok("0"))
18427}
18428
18429fn moe_ffn_cpu_batched(
18435 m: &MoeFfn,
18436 x: &[f32],
18437 idx: &[usize],
18438 p: &[f32],
18439 wsum: f32,
18440 pool: Option<&Pool>,
18441) -> Option<Vec<f32>> {
18442 if idx.is_empty() || !moe_batch_enabled() {
18443 return None;
18444 }
18445 if FFN_PROBE.with(|pr| pr.borrow().is_some()) {
18449 return None;
18450 }
18451 let n = idx.len() + usize::from(m.shared.is_some());
18452 let mut pairs = Vec::with_capacity(n);
18453 let mut downs = Vec::with_capacity(n);
18454 let mut ws = Vec::with_capacity(n);
18455 for &e in idx {
18456 let d = &m.experts[e];
18457 if d.act != Act::Silu {
18458 return None;
18459 }
18460 pairs.push((&d.gate_proj, &d.up_proj));
18461 downs.push(&d.down_proj);
18462 ws.push(p[e] / wsum * m.per_expert_scale.as_ref().map_or(1.0, |v| v[e]));
18463 }
18464 if let Some((se, gate)) = &m.shared {
18467 if se.act != Act::Silu {
18468 return None;
18469 }
18470 let g = gate.as_ref().map_or(1.0, |gate| {
18471 let mut gl = [0.0f32; 1];
18472 gate.matvec(x, &mut gl, pool);
18473 1.0 / (1.0 + (-gl[0]).exp())
18474 });
18475 pairs.push((&se.gate_proj, &se.up_proj));
18476 downs.push(&se.down_proj);
18477 ws.push(g);
18478 }
18479 let inter = pairs[0].0.rows();
18480 let mut gs: Vec<Vec<f32>> = (0..pairs.len()).map(|_| vec![0f32; inter]).collect();
18481 if !QTensor::moe_gate_up_many(&pairs, x, &mut gs, pool) {
18482 return None;
18483 }
18484 let mut out = attention::take_buf(x.len());
18485 if !QTensor::moe_down_many(&downs, &gs, &ws, &mut out, pool) {
18486 attention::recycle_buf(&mut out);
18487 return None;
18488 }
18489 Some(out)
18490}
18491
18492pub(crate) fn moe_cold_experts_cpu(
18498 experts: &[(&DenseFfn, f32)],
18499 x: &[f32],
18500 pool: Option<&Pool>,
18501) -> Vec<f32> {
18502 let mut out = attention::take_buf(x.len());
18503 if experts.is_empty() {
18504 return out;
18505 }
18506 let pairs: Vec<_> = experts
18507 .iter()
18508 .map(|(e, _)| (&e.gate_proj, &e.up_proj))
18509 .collect();
18510 let downs: Vec<_> = experts.iter().map(|(e, _)| &e.down_proj).collect();
18511 let weights: Vec<_> = experts.iter().map(|(_, w)| *w).collect();
18512 let inter = experts[0].0.gate_proj.rows();
18513 let mut activations: Vec<Vec<f32>> = (0..experts.len()).map(|_| vec![0.0; inter]).collect();
18514 if QTensor::moe_gate_up_many(&pairs, x, &mut activations, pool)
18515 && QTensor::moe_down_many(&downs, &activations, &weights, &mut out, pool)
18516 {
18517 return out;
18518 }
18519 out.fill(0.0);
18520 for &(expert, weight) in experts {
18521 let mut one = dense_ffn(expert, x, pool);
18522 for (o, v) in out.iter_mut().zip(&one) {
18523 *o += weight * v;
18524 }
18525 attention::recycle_buf(&mut one);
18526 }
18527 out
18528}
18529
18530pub(crate) fn moe_cold_experts_rows_cpu(
18534 jobs: &[Vec<(&DenseFfn, f32)>],
18535 xs: &[f32],
18536 hidden: usize,
18537 pool: Option<&Pool>,
18538) -> Vec<f32> {
18539 let mut out = vec![0.0; xs.len()];
18540 let mut experts: Vec<&DenseFfn> = Vec::new();
18541 let mut groups: Vec<Vec<usize>> = Vec::new();
18542 let mut terms = vec![Vec::new(); jobs.len()];
18543 for (r, row) in jobs.iter().enumerate() {
18544 for &(e, w) in row {
18545 let g = match experts.iter().position(|&d| std::ptr::eq(d, e)) {
18546 Some(g) => g,
18547 None => {
18548 experts.push(e);
18549 groups.push(Vec::new());
18550 groups.len() - 1
18551 }
18552 };
18553 terms[r].push((g, groups[g].len(), w));
18554 groups[g].push(r);
18555 }
18556 }
18557 if experts.is_empty() {
18558 return out;
18559 }
18560 let pairs: Vec<_> = experts.iter().map(|e| (&e.gate_proj, &e.up_proj)).collect();
18561 let downs: Vec<_> = experts.iter().map(|e| &e.down_proj).collect();
18562 let lens: Vec<_> = groups.iter().map(Vec::len).collect();
18563 let count: usize = lens.iter().sum();
18564 let mut acts = vec![vec![0.0; experts[0].gate_proj.rows()]; count];
18565 let mut ds = vec![vec![0.0; hidden]; count];
18566 if QTensor::moe_gate_up_rows(&pairs, &groups, xs, &mut acts, pool)
18567 && QTensor::moe_down_rows(&downs, &lens, &acts, &mut ds, pool)
18568 {
18569 let mut offset = 0;
18570 let offsets: Vec<_> = lens
18571 .iter()
18572 .map(|&n| {
18573 let start = offset;
18574 offset += n;
18575 start
18576 })
18577 .collect();
18578 for (r, terms) in terms.iter().enumerate() {
18579 for &(g, slot, w) in terms {
18580 for (o, &v) in out[r * hidden..(r + 1) * hidden]
18581 .iter_mut()
18582 .zip(&ds[offsets[g] + slot])
18583 {
18584 *o += w * v;
18585 }
18586 }
18587 }
18588 } else {
18589 for (r, jobs) in jobs.iter().enumerate() {
18590 if !jobs.is_empty() {
18591 let mut row = moe_cold_experts_cpu(jobs, &xs[r * hidden..(r + 1) * hidden], pool);
18592 out[r * hidden..(r + 1) * hidden].copy_from_slice(&row);
18593 attention::recycle_buf(&mut row);
18594 }
18595 }
18596 }
18597 out
18598}
18599
18600fn moe_ffn_cpu(
18602 m: &MoeFfn,
18603 x: &[f32],
18604 idx: &[usize],
18605 p: &[f32],
18606 wsum: f32,
18607 pool: Option<&Pool>,
18608) -> Vec<f32> {
18609 if let Some(out) = moe_ffn_cpu_batched(m, x, idx, p, wsum, pool) {
18610 return out;
18611 }
18612 let mut out = attention::take_buf(x.len());
18613 for &e in idx {
18614 let mut eo = dense_ffn(&m.experts[e], x, pool);
18615 let w = p[e] / wsum * m.per_expert_scale.as_ref().map_or(1.0, |v| v[e]);
18616 for i in 0..out.len() {
18617 out[i] += w * eo[i];
18618 }
18619 attention::recycle_buf(&mut eo);
18620 }
18621 if let Some((se, gate)) = &m.shared {
18622 let mut so = dense_ffn(se, x, pool);
18623 let g = gate.as_ref().map_or(1.0, |gate| {
18624 let mut gl = [0.0f32; 1];
18625 gate.matvec(x, &mut gl, pool);
18626 1.0 / (1.0 + (-gl[0]).exp())
18627 });
18628 for i in 0..out.len() {
18629 out[i] += g * so[i];
18630 }
18631 attention::recycle_buf(&mut so);
18632 }
18633 out
18634}
18635
18636#[allow(clippy::too_many_arguments)]
18644pub(crate) fn mla_attention(
18645 w: &MlaWeights,
18646 normed: &[f32],
18647 cache: &mut crate::kv_cache::LayerKvCache,
18648 position: usize,
18649 inv_freq: &[f32],
18650 rope_scale: f32,
18651 eps: f64,
18652 pool: Option<&Pool>,
18653) -> Vec<f32> {
18654 let (nh, dr, dn, dv, lora) = (w.nh, w.qk_rope, w.qk_nope, w.v_dim, w.lora);
18655 let hd = dr + dn;
18656 let mut q = vec![0.0f32; nh * hd];
18657 match (&w.q_a, &w.q_a_norm) {
18658 (Some(qa), Some(qn)) => {
18659 let mut t = vec![0.0f32; qa.rows()];
18660 qa.matvec(normed, &mut t, pool);
18661 let tn = inference::rms_norm(&t, qn, eps, NormStyle::Qwen);
18662 w.q_proj.matvec(&tn, &mut q, pool);
18663 }
18664 _ => w.q_proj.matvec(normed, &mut q, pool),
18665 }
18666 let mut ca = vec![0.0f32; lora + dr];
18667 w.kv_a.matvec(normed, &mut ca, pool);
18668 let (c_lat, k_rope) = ca.split_at_mut(lora);
18669 let latn = inference::rms_norm(c_lat, &w.kv_a_norm, eps, NormStyle::Qwen);
18670 let mut kvb = vec![0.0f32; nh * (dn + dv)];
18671 w.kv_b.matvec(&latn, &mut kvb, pool);
18672 if !w.nope {
18673 attention::rope_rotate_scaled(k_rope, position, inv_freq, rope_scale);
18674 }
18675 for h in 0..nh {
18676 if !w.nope {
18677 attention::rope_rotate_scaled(
18678 &mut q[h * hd..h * hd + dr],
18679 position,
18680 inv_freq,
18681 rope_scale,
18682 );
18683 }
18684 }
18685 let mut k = vec![0.0f32; nh * hd];
18686 let mut v = vec![0.0f32; nh * hd];
18687 for h in 0..nh {
18688 k[h * hd..h * hd + dr].copy_from_slice(k_rope);
18689 k[h * hd + dr..(h + 1) * hd].copy_from_slice(&kvb[h * (dn + dv)..h * (dn + dv) + dn]);
18690 v[h * hd..h * hd + dv].copy_from_slice(&kvb[h * (dn + dv) + dn..(h + 1) * (dn + dv)]);
18691 }
18692 cache.append(&k, &v, &vec![true; nh]);
18693 let (ao, mut imp) = attention::attend_all_heads(&q, cache, nh, 1, hd, w.scale, None, 0.0);
18694 attention::recycle_buf(&mut imp);
18695 let mut ov = vec![0.0f32; nh * dv];
18696 for h in 0..nh {
18697 ov[h * dv..(h + 1) * dv].copy_from_slice(&ao[h * hd..h * hd + dv]);
18698 }
18699 let mut out = vec![0.0f32; w.o_proj.rows()];
18700 w.o_proj.matvec(&ov, &mut out, pool);
18701 out
18702}
18703
18704fn dense_moe_ffn(
18711 dm: &DenseMoeFfn,
18712 x_normed: &[f32],
18713 h_raw: &[f32],
18714 eps: f64,
18715 norm_style: NormStyle,
18716 pool: Option<&Pool>,
18717) -> Vec<f32> {
18718 let mut d = dense_ffn(&dm.dense, x_normed, pool);
18719 d = inference::rms_norm(&d, &dm.post_norm_1, eps, norm_style);
18720 let m = &dm.moe;
18721 let ne = m.experts.len();
18722 let mut logits = vec![0.0f32; ne];
18723 if m.router_input_norm {
18724 let ss: f32 = h_raw.iter().map(|v| v * v).sum::<f32>() / h_raw.len() as f32;
18725 let inv = 1.0 / (ss + eps as f32).sqrt();
18726 let xr: Vec<f32> = h_raw.iter().map(|v| v * inv).collect();
18727 m.router.matvec(&xr, &mut logits, pool);
18728 } else {
18729 m.router.matvec(h_raw, &mut logits, pool);
18730 }
18731 let (idx, p, wsum) = moe_route(&logits, m, None);
18732 {
18733 let mut st = m.stats.borrow_mut();
18734 if st.len() < ne {
18735 st.resize(ne, 0);
18736 }
18737 for &e in &idx {
18738 st[e] += 1;
18739 }
18740 }
18741 let x2 = inference::rms_norm(h_raw, &dm.pre_norm_2, eps, norm_style);
18742 let mo = moe_ffn_cpu(m, &x2, &idx, &p, wsum, pool);
18743 let mo = inference::rms_norm(&mo, &dm.post_norm_2, eps, norm_style);
18744 for (di, mi) in d.iter_mut().zip(&mo) {
18745 *di += mi;
18746 }
18747 d
18748}
18749
18750fn moe_gpu_refused(why: &'static str) {
18757 use std::sync::atomic::{AtomicBool, Ordering};
18758 static SAID: AtomicBool = AtomicBool::new(false);
18759 if !SAID.swap(true, Ordering::Relaxed) {
18760 tracing::warn!("MoE GPU block refused ({why}) — experts run on the CPU");
18761 }
18762}
18763
18764fn moe_ffn_gpu(
18765 m: &MoeFfn,
18766 x: &[f32],
18767 idx: &[usize],
18768 p: &[f32],
18769 wsum: f32,
18770 pool: Option<&Pool>,
18771) -> Option<Vec<f32>> {
18772 use crate::gpu::MoeJob;
18773
18774 let mut jobs: Vec<MoeJob> = Vec::with_capacity(idx.len() + 1);
18775 let mut model_ref = None;
18776 for &e in idx {
18777 if moe_push_job(&m.experts[e], x, p[e] / wsum, &mut jobs, &mut model_ref).is_none() {
18778 moe_gpu_refused("push_job(expert)");
18779 return None;
18780 }
18781 }
18782 if let Some((se, gate)) = &m.shared {
18783 let g = gate.as_ref().map_or(1.0, |gate| {
18784 let mut gl = [0.0f32; 1];
18785 gate.matvec(x, &mut gl, pool);
18786 1.0 / (1.0 + (-gl[0]).exp())
18787 });
18788 if moe_push_job(se, x, g, &mut jobs, &mut model_ref).is_none() {
18789 moe_gpu_refused("push_job(shared)");
18790 return None;
18791 }
18792 }
18793 let Some(model) = model_ref else {
18794 moe_gpu_refused("no model_ref");
18795 return None;
18796 };
18797 let hidden = jobs[0].down.1;
18798 let mut out = vec![0.0f32; hidden];
18799 if crate::gpu::moe_block(&model, &jobs, &mut out) {
18800 Some(out)
18801 } else {
18802 moe_gpu_refused("gpu::moe_block");
18803 None
18804 }
18805}
18806
18807fn ffn_forward(
18809 ffn: &FfnKind,
18810 x: &[f32],
18811 pool: Option<&Pool>,
18812 experts_allowed: Option<&[bool]>,
18813) -> Vec<f32> {
18814 match ffn {
18815 FfnKind::Dense(d) if !d.segs.is_empty() => tube_ffn(d, x, 1, pool, None),
18816 FfnKind::Dense(d) => dense_ffn(d, x, pool),
18817 FfnKind::Moe(m) => moe_ffn(m, x, pool, experts_allowed),
18818 FfnKind::DenseMoe(_) => unreachable!("DenseMoe dispatches via dense_moe_ffn"),
18822 }
18823}
18824
18825fn ffn_forward_pair(
18829 ffn: &FfnKind,
18830 x1: &[f32],
18831 x2: &[f32],
18832 pool: Option<&Pool>,
18833 experts_allowed: Option<&[bool]>,
18834) -> (Vec<f32>, Vec<f32>) {
18835 let d = match ffn {
18836 FfnKind::Dense(d) if !d.segs.is_empty() => {
18839 return (
18840 tube_ffn(d, x1, 1, pool, None),
18841 tube_ffn(d, x2, 1, pool, None),
18842 );
18843 }
18844 FfnKind::Dense(d) => d,
18845 FfnKind::Moe(m) => {
18846 return (
18847 moe_ffn(m, x1, pool, experts_allowed),
18848 moe_ffn(m, x2, pool, experts_allowed),
18849 );
18850 }
18851 FfnKind::DenseMoe(_) => unreachable!("DenseMoe dispatches via dense_moe_ffn"),
18852 };
18853 let inter = d.gate_proj.rows();
18854 FFN_SCRATCH.with(|s| {
18855 let mut s = s.borrow_mut();
18856 let [g1, g2, u1, u2] = &mut *s;
18857 g1.resize(inter, 0.0);
18858 g2.resize(inter, 0.0);
18859 u1.resize(inter, 0.0);
18860 u2.resize(inter, 0.0);
18861 QTensor::matvec2_many(
18864 [&d.gate_proj, &d.up_proj],
18865 x1,
18866 x2,
18867 [g1.as_mut_slice(), u1.as_mut_slice()],
18868 [g2.as_mut_slice(), u2.as_mut_slice()],
18869 pool,
18870 );
18871 for i in 0..inter {
18872 g1[i] = d.act.combine(g1[i], u1[i]);
18873 g2[i] = d.act.combine(g2[i], u2[i]);
18874 }
18875 let mut o1 = attention::take_buf(d.down_proj.rows());
18876 let mut o2 = attention::take_buf(d.down_proj.rows());
18877 d.down_proj.matvec2(g1, g2, &mut o1, &mut o2, pool);
18878 (o1, o2)
18879 })
18880}
18881
18882#[cfg(test)]
18883mod tests {
18884
18885 #[test]
18890 fn prefill_chunk_rule_widens_only_dense_on_discrete() {
18891 use super::{
18892 prefill_chunk_rule, ChunkHost, ChunkStackFacts, DISCRETE_DENSE_PREFILL_CHUNK,
18893 };
18894 let dense_card = ChunkStackFacts {
18895 plain_dense: true,
18896 discrete: true,
18897 gpu_on: true,
18898 ..Default::default()
18899 };
18900 assert!(dense_card.dense_on_discrete());
18901 assert_eq!(
18903 prefill_chunk_rule(None, ChunkHost::Other, dense_card.dense_on_discrete()),
18904 DISCRETE_DENSE_PREFILL_CHUNK
18905 );
18906 assert!(DISCRETE_DENSE_PREFILL_CHUNK > 48);
18907 for (label, facts) in [
18908 ("GDN hybrid / MoE / DeepSeek stack", ChunkStackFacts { plain_dense: false, ..dense_card }),
18909 ("integrated GPU", ChunkStackFacts { discrete: false, ..dense_card }),
18910 ("CPU only", ChunkStackFacts { gpu_on: false, discrete: false, ..dense_card }),
18911 ("capacity split", ChunkStackFacts { capacity_split: true, ..dense_card }),
18912 ("multi-GPU plan", ChunkStackFacts { multi_gpu: true, ..dense_card }),
18913 ("O(1) layers", ChunkStackFacts { o1: true, ..dense_card }),
18914 ] {
18915 assert!(!facts.dense_on_discrete(), "{label}");
18916 assert_eq!(
18917 prefill_chunk_rule(None, ChunkHost::Other, facts.dense_on_discrete()),
18918 48,
18919 "{label} keeps the historical x86 chunk"
18920 );
18921 }
18922 for dense in [false, true] {
18924 assert_eq!(prefill_chunk_rule(None, ChunkHost::Macos, dense), 512);
18925 assert_eq!(prefill_chunk_rule(None, ChunkHost::Aarch64, dense), 256);
18926 }
18927 for host in [ChunkHost::Macos, ChunkHost::Aarch64, ChunkHost::Other] {
18929 for dense in [false, true] {
18930 assert_eq!(prefill_chunk_rule(Some(48), host, dense), 48);
18931 assert_eq!(prefill_chunk_rule(Some(0), host, dense), 1);
18932 }
18933 }
18934 }
18935
18936 #[test]
18937 fn kv_reuse_plan_pulls_rows_decode_wrote_only_on_the_device() {
18938 use super::{ReuseLayer, ReusePlan, kv_reuse_plan};
18939 let full = |host_rows, device_rows| ReuseLayer {
18940 full: true,
18941 host_rows,
18942 device_rows,
18943 device_state: false,
18944 };
18945 assert_eq!(
18948 kv_reuse_plan(339, &[full(300, Some(339)), full(300, Some(339))]),
18949 ReusePlan::Pull(vec![(0, 300, 339), (1, 300, 339)])
18950 );
18951 assert_eq!(kv_reuse_plan(339, &[full(339, None)]), ReusePlan::Ready);
18953 assert_eq!(kv_reuse_plan(339, &[full(339, Some(345))]), ReusePlan::Ready);
18955 assert_eq!(
18957 kv_reuse_plan(339, &[full(300, Some(339)), full(339, None)]),
18958 ReusePlan::Pull(vec![(0, 300, 339)])
18959 );
18960 assert_eq!(kv_reuse_plan(339, &[full(300, Some(320))]), ReusePlan::Fresh);
18962 assert_eq!(kv_reuse_plan(339, &[full(300, None)]), ReusePlan::Fresh);
18963 assert_eq!(kv_reuse_plan(339, &[full(350, None)]), ReusePlan::Fresh);
18964 let conv = |device_state| ReuseLayer {
18967 full: false,
18968 host_rows: 0,
18969 device_rows: None,
18970 device_state,
18971 };
18972 assert_eq!(
18973 kv_reuse_plan(339, &[conv(true), full(300, Some(339))]),
18974 ReusePlan::Fresh
18975 );
18976 assert_eq!(kv_reuse_plan(339, &[conv(false), full(339, None)]), ReusePlan::Ready);
18977 }
18978
18979 #[test]
18980 fn nll_graph_policy_scopes_only_the_fused_head() {
18981 for (label, unmasked, prefer_graph, native_metal, want_graph, want_head) in [
18982 ("vulkan graph", true, true, false, true, false),
18984 ("native Metal graph", true, true, true, true, true),
18986 ("masked", false, true, false, false, false),
18988 ("graph disabled", true, false, true, false, false),
18989 ] {
18990 let (graph_quality, graph_head_required) =
18991 super::nll_graph_policy(unmasked, prefer_graph, native_metal);
18992 assert_eq!(graph_quality, want_graph, "{label}: graph quality");
18993 assert_eq!(graph_head_required, want_head, "{label}: fused head");
18994 }
18995 }
18996
18997 #[test]
18998 fn mtp_prefill_pair_boundaries_skip_only_final_prompt_row() {
18999 assert_eq!(mtp_prefill_pair_count(0, 128, 256), 128);
19000 assert_eq!(mtp_prefill_pair_count(128, 256, 256), 127);
19001 assert_eq!(mtp_prefill_pair_count(0, 256, 256), 255);
19002 assert_eq!(mtp_prefill_pair_count(256, 256, 256), 0);
19003 assert_eq!(mtp_prefill_pair_count(300, 320, 256), 0);
19004 }
19005
19006 #[test]
19007 fn cancel_flag_stops_generation() {
19008 let mut p = create_test_pipeline(16, 32, 2, 2, 8, 2, 32);
19009 p.cancel.store(true, std::sync::atomic::Ordering::Relaxed);
19012 let r = p.generate_from_ids(&[1, 2, 3], 8, None, None).unwrap();
19013 assert_eq!(r.finish_reason, "cancelled");
19014 assert!(
19015 r.token_ids.is_empty(),
19016 "no tokens after cancel: {:?}",
19017 r.token_ids
19018 );
19019 assert_eq!(p.kv_cache.seq_len(), 0);
19020 assert!(p.kv_history.is_empty());
19021 assert!(!p.graph_want_logits);
19022 assert!(p.graph_logits.is_none());
19023 let r2 = p.generate_from_ids(&[1, 2, 3], 4, None, None).unwrap();
19025 assert_ne!(r2.finish_reason, "cancelled");
19026 }
19027 use super::*;
19028
19029 #[test]
19037 fn dynamic_ffn_equals_the_zeroing_arm() {
19038 let (hidden, inter) = (8usize, 32usize);
19039 let synth = |n: usize, salt: usize| -> Vec<f32> {
19040 (0..n)
19041 .map(|i| (((i * 29 + salt * 13 + 7) % 89) as f32 / 89.0 - 0.5) * 0.6)
19042 .collect()
19043 };
19044 let down = synth(hidden * inter, 3);
19045 let mut down_t = vec![0.0f32; inter * hidden];
19046 for r in 0..hidden {
19047 for c in 0..inter {
19048 down_t[c * hidden + r] = down[r * inter + c];
19049 }
19050 }
19051 let d = DenseFfn {
19052 gate_proj: QTensor::from_f32(synth(inter * hidden, 1), inter, hidden),
19053 up_proj: QTensor::from_f32(synth(inter * hidden, 2), inter, hidden),
19054 down_proj: QTensor::from_f32(down.clone(), hidden, inter),
19055 act: Act::Silu,
19056 down_t: Some(QTensor::from_f32(down_t, inter, hidden)),
19057 segs: Vec::new(),
19058 };
19059 let x = synth(hidden, 11);
19060 let k = 12usize;
19061 let got = dense_ffn_dynamic(&d, &x, None, k).expect("down_t present");
19062 let mut g = vec![0.0f32; inter];
19064 d.gate_proj.matvec(&x, &mut g, None);
19065 let mut u = vec![0.0f32; inter];
19066 d.up_proj.matvec(&x, &mut u, None);
19067 for v in g.iter_mut() {
19068 *v = inference::silu(*v);
19069 }
19070 keep_top_k(&mut g, k);
19071 for i in 0..inter {
19072 g[i] *= u[i];
19073 }
19074 let mut want = vec![0.0f32; hidden];
19075 d.down_proj.matvec(&g, &mut want, None);
19076 for (a, b) in want.iter().zip(&got) {
19077 assert!((a - b).abs() < 1e-5, "dynamic {b} vs reference {a}");
19078 }
19079 }
19080
19081 #[test]
19087 fn tube_ffn_open_equals_dense_and_closed_equals_masked() {
19088 let (hidden, core, tube) = (8usize, 12usize, 8usize);
19089 let inter = core + tube;
19090 let synth = |n: usize, salt: usize| -> Vec<f32> {
19091 (0..n)
19092 .map(|i| (((i * 41 + salt * 17 + 5) % 97) as f32 / 97.0 - 0.5) * 0.5)
19093 .collect()
19094 };
19095 let (g_all, u_all) = (synth(inter * hidden, 1), synth(inter * hidden, 2));
19096 let d_all = synth(hidden * inter, 3);
19097 let dense = DenseFfn {
19099 gate_proj: QTensor::from_f32(g_all.clone(), inter, hidden),
19100 up_proj: QTensor::from_f32(u_all.clone(), inter, hidden),
19101 down_proj: QTensor::from_f32(d_all.clone(), hidden, inter),
19102 act: Act::Silu,
19103 down_t: None,
19104 segs: Vec::new(),
19105 };
19106 let rows =
19107 |v: &[f32], a: usize, b: usize| -> Vec<f32> { v[a * hidden..b * hidden].to_vec() };
19108 let cols = |v: &[f32], a: usize, b: usize| -> Vec<f32> {
19109 let mut o = Vec::with_capacity(hidden * (b - a));
19110 for r in 0..hidden {
19111 o.extend_from_slice(&v[r * inter + a..r * inter + b]);
19112 }
19113 o
19114 };
19115 let tubed = DenseFfn {
19116 down_t: None,
19117 gate_proj: QTensor::from_f32(rows(&g_all, 0, core), core, hidden),
19118 up_proj: QTensor::from_f32(rows(&u_all, 0, core), core, hidden),
19119 down_proj: QTensor::from_f32(cols(&d_all, 0, core), hidden, core),
19120 act: Act::Silu,
19121 segs: vec![FfnSeg {
19122 gate: QTensor::from_f32(rows(&g_all, core, inter), tube, hidden),
19123 up: QTensor::from_f32(rows(&u_all, core, inter), tube, hidden),
19124 down: QTensor::from_f32(cols(&d_all, core, inter), hidden, tube),
19125 start: core,
19126 width: tube,
19127 }],
19128 };
19129 let x = synth(hidden, 7);
19130 let want = dense_ffn(&dense, &x, None);
19131 let got = tube_ffn(&tubed, &x, 1, None, None);
19132 for (a, b) in want.iter().zip(&got) {
19133 assert!((a - b).abs() < 1e-5, "open tube: {a} vs {b}");
19134 }
19135 let mut bits = vec![0u8; inter.div_ceil(8)];
19137 for n in 0..core {
19138 bits[n / 8] |= 1 << (n % 8);
19139 }
19140 let closed = tube_ffn(&tubed, &x, 1, None, Some(&bits));
19141 let masked = dense_ffn_masked(&dense, &x, None, &bits);
19142 for (a, b) in masked.iter().zip(&closed) {
19143 assert!((a - b).abs() < 1e-5, "closed tube: {a} vs {b}");
19144 }
19145 let batch = tube_ffn(&tubed, &x, 1, None, Some(&bits));
19147 for (a, b) in closed.iter().zip(&batch) {
19148 assert_eq!(a, b, "batch arm disagrees with decode arm");
19149 }
19150 }
19151
19152 #[test]
19154 fn sparse_ffn_quant_equals_dense_with_inactive_zeroed() {
19155 let (hidden, inter) = (16usize, 40usize);
19156 let synth = |n: usize, salt: usize| -> Vec<f32> {
19157 (0..n)
19158 .map(|i| (((i * 37 + salt * 11 + 3) % 101) as f32 / 101.0 - 0.5) * 0.4)
19159 .collect()
19160 };
19161 let d = DenseFfn {
19162 gate_proj: QTensor::from_f32(synth(inter * hidden, 1), inter, hidden),
19163 up_proj: QTensor::from_f32(synth(inter * hidden, 2), inter, hidden),
19164 down_proj: QTensor::from_f32(synth(hidden * inter, 3), hidden, inter),
19165 act: Act::Silu,
19166 down_t: None,
19167 segs: Vec::new(),
19168 };
19169 let x = synth(hidden, 9);
19170 let active: Vec<u16> = (0..inter as u16).filter(|i| i % 3 == 0).collect();
19172
19173 let sparse = sparse_ffn_quant(&d, &x, &active, hidden, None);
19174
19175 let mut g = vec![0.0f32; inter];
19177 d.gate_proj.matvec(&x, &mut g, None);
19178 let mut u = vec![0.0f32; inter];
19179 d.up_proj.matvec(&x, &mut u, None);
19180 let act_set: std::collections::HashSet<u16> = active.iter().copied().collect();
19181 for i in 0..inter {
19182 g[i] = if act_set.contains(&(i as u16)) {
19183 inference::silu(g[i]) * u[i]
19184 } else {
19185 0.0
19186 };
19187 }
19188 let mut reference = vec![0.0f32; hidden];
19189 d.down_proj.matvec(&g, &mut reference, None);
19190
19191 let max_d = sparse
19192 .iter()
19193 .zip(&reference)
19194 .map(|(a, b)| (a - b).abs())
19195 .fold(0.0f32, f32::max);
19196 assert!(max_d < 1e-5, "sparse != dense-zeroed: max|Δ| = {max_d}");
19197 }
19198
19199 fn attach_test_mtp(p: &mut Pipeline) {
19201 let (h, inter, heads, kv, hd) = (
19202 p.hidden_size,
19203 p.intermediate_size,
19204 p.num_heads,
19205 p.num_kv_heads,
19206 p.head_dim,
19207 );
19208 let synth = |n: usize, salt: usize| -> Vec<f32> {
19209 (0..n)
19210 .map(|i| (((i * 29 + salt * 23 + 5) % 101) as f32 / 101.0 - 0.5) * 0.2)
19211 .collect()
19212 };
19213 let qt = |rows: usize, cols: usize, salt: usize| -> QTensor {
19214 QTensor::from_f32(synth(rows * cols, salt), rows, cols)
19215 };
19216 p.mtp = Some(MtpModule {
19217 enorm: vec![1.0; h],
19218 hnorm: vec![1.0; h],
19219 eh_proj: qt(h, 2 * h, 301),
19220 layer: LayerWeights {
19221 input_norm: vec![1.0; h],
19222 post_norm: vec![1.0; h],
19223 attn_out_norm: None,
19224 ffn_out_norm: None,
19225 layer_scale: None,
19226 ffn: FfnKind::Dense(DenseFfn {
19227 gate_proj: qt(inter, h, 315),
19228 up_proj: qt(inter, h, 316),
19229 down_proj: qt(h, inter, 317),
19230 act: Act::Silu,
19231 down_t: None,
19232 segs: Vec::new(),
19233 }),
19234 attn: AttnKind::Full {
19235 bias: None,
19236 wq: qt(heads * hd, h, 311),
19237 wk: qt(kv * hd, h, 312),
19238 wv: qt(kv * hd, h, 313),
19239 wo: qt(h, heads * hd, 314),
19240 q_norm: None,
19241 k_norm: None,
19242 output_gate: false,
19243 softplus_gate: None,
19244 },
19245 },
19246 final_norm: vec![1.0; h],
19247 kv: crate::kv_cache::LayerKvCache::new(kv, hd),
19248 });
19249 }
19250
19251 #[test]
19252 fn speculative_equals_vanilla_greedy() {
19253 unsafe { std::env::set_var("CMF_GPU_WGPU_GRAPH", "0") };
19257 let run = |spec: bool| {
19258 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
19259 p.sampler_config.temperature = 0.0;
19260 attach_test_mtp(&mut p);
19261 p.speculative = spec;
19262 let r = p.generate("abcdef", 12, None, None).unwrap();
19263 (r.token_ids, r.mtp_drafted, r.mtp_accepted)
19264 };
19265 let (vanilla, d0, _) = run(false);
19266 let (spec, d1, a1) = run(true);
19267 assert_eq!(d0, 0, "vanilla path must not draft");
19268 assert!(d1 > 0, "speculative path must draft");
19269 assert_eq!(
19270 vanilla, spec,
19271 "speculative must reproduce the exact greedy sequence (accepted {a1}/{d1})"
19272 );
19273 }
19274
19275 #[test]
19276 fn speculative_accepts_constant_oracle() {
19277 unsafe { std::env::set_var("CMF_GPU_WGPU_GRAPH", "0") };
19279 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
19280 p.sampler_config.temperature = 0.0;
19281 p.sampler_config.repetition_penalty = 1.0;
19282 p.weights.lm_head = QTensor::from_f32(vec![0.01; 64 * 8], 64, 8);
19285 attach_test_mtp(&mut p);
19286 p.speculative = true;
19287 let r = p.generate("abcd", 10, None, None).unwrap();
19288 assert!(r.mtp_drafted > 0);
19289 assert_eq!(
19290 r.mtp_accepted, r.mtp_drafted,
19291 "constant logits → every draft accepted"
19292 );
19293 assert!(r.token_ids.windows(2).all(|w| w[0] == w[1]));
19296 }
19297
19298 #[test]
19299 fn empty_prompt_is_an_error_not_a_panic() {
19300 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 260);
19301 let r = p.generate("", 4, None, None);
19302 assert!(r.is_err(), "empty prompt must be a clean error");
19303 }
19304
19305 #[test]
19306 fn every_token_enters_kv_exactly_once() {
19307 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
19308 p.sampler_config.temperature = 0.0;
19310 let r = p.generate("abc", 2, None, None).unwrap();
19311 assert_eq!(r.prompt_tokens, 3);
19312 assert_eq!(
19316 p.kv_cache.seq_len(),
19317 3 + r.tokens_generated - 1,
19318 "each token must be cached exactly once (v1 cached the last prompt token twice)"
19319 );
19320 }
19321
19322 #[test]
19323 fn generation_is_reproducible_with_seed() {
19324 let run = || {
19325 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
19326 p.generate("hello", 8, None, None).unwrap().token_ids
19327 };
19328 assert_eq!(run(), run());
19329 }
19330
19331 #[test]
19332 fn resetting_sampler_restarts_the_seeded_stream() {
19333 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
19334 let config = SamplerConfig {
19335 seed: Some(1234),
19336 ..SamplerConfig::default()
19337 };
19338 p.set_sampler_config(config.clone());
19339 let first = p.generate("hello", 8, None, None).unwrap().token_ids;
19340 p.set_sampler_config(config);
19341 let second = p.generate("hello", 8, None, None).unwrap().token_ids;
19342 assert_eq!(first, second);
19343 }
19344
19345 #[test]
19346 fn eviction_bounds_the_cache() {
19347 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 260);
19348 p.kv_cache.max_seq_len = 6;
19349 p.sampler_config.temperature = 0.0;
19350 let _ = p.generate("abcd", 12, None, None).unwrap();
19351 assert!(
19352 p.kv_cache.seq_len() <= 6 + 1,
19353 "cache must stay bounded by max_seq_len (got {})",
19354 p.kv_cache.seq_len()
19355 );
19356 }
19357
19358 #[test]
19359 fn confidence_matches_tokens_and_is_a_probability() {
19360 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
19361 p.sampler_config.temperature = 0.0;
19362 p.sampler_config.repetition_penalty = 1.0;
19363 let r = p.generate("abcd", 10, None, None).unwrap();
19364 assert_eq!(
19365 r.token_confidence.len(),
19366 r.token_ids.len(),
19367 "one confidence per emitted token"
19368 );
19369 for &c in &r.token_confidence {
19370 assert!((0.0..=1.0).contains(&c), "confidence out of [0,1]: {c}");
19371 }
19372 let logits = [1.0f32, 3.0, 0.5, 3.0];
19374 let p0 = top1_prob_t(&logits, 1, 1.0);
19375 let p1 = top1_prob_t(&logits, 3, 1.0);
19376 assert!((p0 - p1).abs() < 1e-6, "equal logits → equal prob");
19377 assert!(p0 > 0.0 && p0 < 1.0);
19378 let sharp = top1_prob_t(&logits, 1, 1.0);
19380 let soft = top1_prob_t(&logits, 1, 2.0);
19381 assert!(soft < sharp, "higher temperature lowers peak confidence");
19382 }
19383
19384 #[test]
19385 fn trace_is_opt_in_and_parallels_the_output() {
19386 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
19388 p.sampler_config.temperature = 0.0;
19389 p.sampler_config.repetition_penalty = 1.0;
19390 let r = p.generate("abcd", 10, None, None).unwrap();
19391 assert!(r.traces.is_empty(), "trace must be empty unless enabled");
19392
19393 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
19395 p.sampler_config.temperature = 0.0;
19396 p.sampler_config.repetition_penalty = 1.0;
19397 p.set_trace(true);
19398 let r = p.generate("abcd", 10, None, None).unwrap();
19399 assert_eq!(r.traces.len(), r.token_ids.len(), "one trace row per token");
19400 for (i, tr) in r.traces.iter().enumerate() {
19401 assert_eq!(tr.t, i, "trace index is sequential");
19402 assert_eq!(tr.token_id, r.token_ids[i], "trace token_id matches output");
19403 assert_eq!(
19404 tr.confidence, r.token_confidence[i],
19405 "trace confidence matches the confidence channel"
19406 );
19407 assert!(tr.active_skill.is_none() && tr.recon.is_none() && !tr.switched);
19409 }
19410 }
19411
19412 #[test]
19413 fn explain_prefill_logits_match_greedy_first_token() {
19414 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
19418 p.sampler_config.temperature = 0.0;
19419 p.sampler_config.repetition_penalty = 1.0;
19420 let ids = p.tokenizer.encode("abcd");
19421 let logits = p.prefill_next_logits(&ids, None);
19422 let argmax = logits
19423 .iter()
19424 .enumerate()
19425 .max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
19426 .unwrap()
19427 .0 as u32;
19428 let r = p.generate("abcd", 1, None, None).unwrap();
19429 assert_eq!(
19430 argmax, r.token_ids[0],
19431 "explain preview must match greedy emit"
19432 );
19433 }
19434
19435 #[test]
19436 fn laguna_shared_expert_is_unconditionally_added() {
19437 let matrix = |values: Vec<f32>| QTensor::from_f32(values, 2, 2);
19438 let identity = || matrix(vec![1.0, 0.0, 0.0, 1.0]);
19439 let zero_dense = || DenseFfn {
19440 gate_proj: matrix(vec![0.0; 4]),
19441 up_proj: matrix(vec![0.0; 4]),
19442 down_proj: matrix(vec![0.0; 4]),
19443 act: Act::Silu,
19444 down_t: None,
19445 segs: Vec::new(),
19446 };
19447 let shared = DenseFfn {
19448 gate_proj: identity(),
19449 up_proj: identity(),
19450 down_proj: identity(),
19451 act: Act::Silu,
19452 down_t: None,
19453 segs: Vec::new(),
19454 };
19455 let x = [1.0, 2.0];
19456 let expected = dense_ffn(&shared, &x, None);
19457 let moe = MoeFfn {
19458 router: QTensor::from_f32(vec![0.0, 0.0], 1, 2),
19459 experts: vec![zero_dense()],
19460 top_k: 1,
19461 norm_topk_prob: true,
19462 router_sigmoid: true,
19463 expert_bias: None,
19464 routed_scaling: 1.0,
19465 route_tau: None,
19466 shared: Some((shared, None)),
19467 stats: std::cell::RefCell::new(Vec::new()),
19468 act_sq: std::cell::RefCell::new(Vec::new()),
19469 act_rows: std::cell::RefCell::new(Vec::new()),
19470 mask: None,
19471 per_expert_scale: None,
19472 router_input_norm: false,
19473 resonance: None,
19474 grown: Vec::new(),
19475 };
19476 let actual = moe_ffn_cpu(&moe, &x, &[0], &[0.0], 1.0, None);
19477 for (actual, expected) in actual.iter().zip(expected) {
19478 assert!((actual - expected).abs() < 1e-6);
19479 }
19480 }
19481
19482 fn mimo_test_pipeline() -> Pipeline {
19491 let (hs, inter, nh, hd, vd, vocab) = (16usize, 24usize, 4usize, 8usize, 4usize, 64usize);
19492 let kvh = [1usize, 2, 2, 1];
19493 let synth = |n: usize, salt: usize| -> Vec<f32> {
19494 (0..n)
19495 .map(|i| (((i * 31 + salt * 17 + 7) % 97) as f32 / 97.0 - 0.5) * 0.4)
19496 .collect()
19497 };
19498 let qt = |rows: usize, cols: usize, salt: usize| {
19499 QTensor::from_f32(synth(rows * cols, salt), rows, cols)
19500 };
19501 let dense = |inter: usize, salt: usize| DenseFfn {
19502 gate_proj: qt(inter, hs, salt),
19503 up_proj: qt(inter, hs, salt + 1),
19504 down_proj: qt(hs, inter, salt + 2),
19505 act: Act::Silu,
19506 down_t: None,
19507 segs: Vec::new(),
19508 };
19509 let layers: Vec<LayerWeights> = (0..4)
19510 .map(|li| LayerWeights {
19511 input_norm: vec![1.0; hs],
19512 post_norm: vec![1.0; hs],
19513 attn_out_norm: None,
19514 ffn_out_norm: None,
19515 layer_scale: None,
19516 ffn: if li == 0 {
19517 FfnKind::Dense(dense(inter, 50))
19518 } else {
19519 FfnKind::Moe(MoeFfn {
19520 router: qt(4, hs, 60 + li),
19521 experts: (0..4).map(|e| dense(8, 70 + li * 10 + e * 3)).collect(),
19522 top_k: 2,
19523 norm_topk_prob: true,
19524 router_sigmoid: true,
19525 expert_bias: Some(vec![0.02, -0.03, 0.01, 0.0]),
19526 routed_scaling: 1.0,
19527 route_tau: None,
19528 shared: None,
19529 stats: std::cell::RefCell::new(Vec::new()),
19530 act_sq: std::cell::RefCell::new(Vec::new()),
19531 act_rows: std::cell::RefCell::new(Vec::new()),
19532 mask: None,
19533 per_expert_scale: None,
19534 router_input_norm: false,
19535 resonance: None,
19536 grown: Vec::new(),
19537 })
19538 },
19539 attn: AttnKind::Full {
19540 wq: qt(nh * hd, hs, li * 10 + 1),
19541 wk: qt(kvh[li] * hd, hs, li * 10 + 2),
19542 wv: qt(kvh[li] * vd, hs, li * 10 + 3),
19543 wo: qt(hs, nh * vd, li * 10 + 4),
19544 q_norm: None,
19545 k_norm: None,
19546 output_gate: false,
19547 softplus_gate: None,
19548 bias: None,
19549 },
19550 })
19551 .collect();
19552 let mut p = Pipeline::new(
19553 Tokenizer::byte_level(),
19554 PipelineWeights {
19555 embed_tokens: qt(vocab, hs, 100),
19556 layers,
19557 lm_head: qt(vocab, hs, 200),
19558 final_norm: vec![1.0; hs],
19559 },
19560 hs,
19561 inter,
19562 nh,
19563 1, hd,
19565 4,
19566 4,
19567 false,
19568 vocab,
19569 1e-6,
19570 1e7,
19571 NormStyle::Qwen,
19572 4096,
19573 SamplerConfig {
19574 seed: Some(7),
19575 ..Default::default()
19576 },
19577 );
19578 p.layer_dump = None;
19580 p.set_rotary(4, 1e7);
19581 p.sliding_layers = Some(vec![false, true, true, false]);
19582 p.swa = Some((3, usize::MAX));
19583 p.rotary_dim_local = Some(4);
19584 p.inv_freq_local = Some(std::sync::Arc::new(attention::rope_inv_freq(4, 1e4)));
19585 p.set_attn_geometry(Some(kvh.to_vec()), Some(vd)).unwrap();
19586 p.set_layer_sinks(1, vec![0.5, -1.0, 1.5, 0.0]).unwrap();
19587 p.set_layer_sinks(2, vec![-0.25, 2.0, 0.75, -1.5]).unwrap();
19588 p
19589 }
19590
19591 #[test]
19592 fn mimo_embedded_prompt_uses_rows_and_never_reuses_token_only_kv() {
19593 let mut p = mimo_test_pipeline();
19594 p.speculative = false;
19595 p.ignore_eos = true;
19596 p.sampler_config.temperature = 0.0;
19597 p.sampler_config.repetition_penalty = 1.0;
19598 let a = vec![3, 5, 7, 9, 11, 13];
19599 let b = vec![4, 8, 12, 16, 20, 24];
19600 let rows: Vec<_> = b.iter().flat_map(|&id| p.embed_id(id)).collect();
19601 let expected = p.generate_from_ids(&b, 8, None, None).unwrap().token_ids;
19602 let actual = p.generate_from_embeds(&a, &rows, 8, None, None).unwrap().token_ids;
19606 assert_eq!(actual, expected);
19607 assert!(p.kv_history.is_empty());
19608 let mut extended = a.clone();
19609 extended.push(17);
19610 let after_media = p.generate_from_ids(&extended, 8, None, None).unwrap().token_ids;
19611 p.reset_session();
19612 let fresh = p.generate_from_ids(&extended, 8, None, None).unwrap().token_ids;
19613 assert_eq!(after_media, fresh);
19614 assert!(p.generate_from_embeds(&a, &rows[..rows.len()-1], 1, None, None).is_err());
19615 p.reset_session();
19619 p.generate_from_ids(&a, 1, None, None).unwrap();
19620 let mut media_ids = p.kv_history.clone();
19621 assert!(!media_ids.is_empty());
19622 media_ids.extend_from_slice(&[19, 21, 23]);
19623 let source_ids: Vec<_> = (0..media_ids.len()).map(|i| b[i % b.len()]).collect();
19624 let source_rows: Vec<_> = source_ids.iter().flat_map(|&id| p.embed_id(id)).collect();
19625 let mut oracle = mimo_test_pipeline();
19626 oracle.speculative = false;
19627 oracle.ignore_eos = true;
19628 oracle.sampler_config.temperature = 0.0;
19629 oracle.sampler_config.repetition_penalty = 1.0;
19630 let expected = oracle.generate_from_ids(&source_ids, 8, None, None).unwrap().token_ids;
19631 assert_eq!(p.generate_from_embeds(&media_ids, &source_rows, 8, None, None).unwrap().token_ids, expected);
19632 assert!(p.kv_history.is_empty());
19633 let mut bad = rows;
19634 bad[0] = f32::NAN;
19635 assert!(p.generate_from_embeds(&a, &bad, 1, None, None).is_err());
19636 }
19637
19638 fn f32_bits(v: &[f32]) -> Vec<u32> {
19639 v.iter().map(|x| x.to_bits()).collect()
19640 }
19641
19642 #[test]
19649 fn mimo_shaped_decode_matches_prefill_batch_bitwise() {
19650 let mut p = mimo_test_pipeline();
19651 let kv: Vec<usize> = p.kv_cache.layers.iter().map(|l| l.num_kv_heads).collect();
19652 assert_eq!(kv, vec![1, 2, 2, 1]);
19653 assert!(p.kv_cache.layers[0].sinks.is_none() && p.kv_cache.layers[3].sinks.is_none());
19654 assert!(p.kv_cache.layers[1].sinks.is_some() && p.kv_cache.layers[2].sinks.is_some());
19655 let ids: Vec<u32> = (0..12u32).map(|i| (i * 7 + 3) % 64).collect();
19656 let hs = p.hidden_size;
19657 let mut decode = Vec::new();
19658 for (pos, &id) in ids.iter().enumerate() {
19659 let e = p.embed_single(id);
19660 let h = p.forward_layers(&e, pos, None);
19661 decode.push(p.logits_from_hidden(&h));
19662 }
19663 for l in &p.kv_cache.layers {
19664 assert_eq!(l.seq_len, 12);
19665 assert_eq!(l.head_values(0).len(), 12 * 8);
19667 }
19668 assert!(decode.iter().flatten().all(|v| v.is_finite()));
19669
19670 p.clear_sequence_state();
19671 let hb = p.prefill_batch_span(PrefillIn::Ids(&ids), 0, None, 0, p.num_layers);
19672 for pos in 0..ids.len() {
19673 let lg = p.logits_from_hidden(&hb[pos * hs..(pos + 1) * hs]);
19674 assert_eq!(
19675 f32_bits(&decode[pos]),
19676 f32_bits(&lg),
19677 "whole prompt, pos {pos}"
19678 );
19679 }
19680
19681 p.clear_sequence_state();
19682 let a = p.prefill_batch_span(PrefillIn::Ids(&ids[..5]), 0, None, 0, p.num_layers);
19683 let b = p.prefill_batch_span(PrefillIn::Ids(&ids[5..]), 5, None, 0, p.num_layers);
19684 for pos in 0..ids.len() {
19685 let row = if pos < 5 {
19686 &a[pos * hs..(pos + 1) * hs]
19687 } else {
19688 &b[(pos - 5) * hs..(pos - 4) * hs]
19689 };
19690 let lg = p.logits_from_hidden(row);
19691 assert_eq!(
19692 f32_bits(&decode[pos]),
19693 f32_bits(&lg),
19694 "two chunks, pos {pos}"
19695 );
19696 }
19697
19698 let last = |p: &mut Pipeline| {
19701 p.clear_sequence_state();
19702 let hb = p.prefill_batch_span(PrefillIn::Ids(&ids), 0, None, 0, p.num_layers);
19703 p.logits_from_hidden(&hb[11 * hs..12 * hs])
19704 };
19705 let base = last(&mut p);
19706 let mut no_sinks = mimo_test_pipeline();
19707 for l in &mut no_sinks.kv_cache.layers {
19708 l.sinks = None;
19709 }
19710 assert_ne!(
19711 f32_bits(&last(&mut no_sinks)),
19712 f32_bits(&base),
19713 "sinks are live"
19714 );
19715 let mut wide = mimo_test_pipeline();
19716 wide.swa = Some((64, usize::MAX));
19717 assert_ne!(
19718 f32_bits(&last(&mut wide)),
19719 f32_bits(&base),
19720 "window is live"
19721 );
19722
19723 p.clear_sequence_state();
19725 p.ignore_eos = true;
19726 let r = p.generate_from_ids(&ids, 4, None, None).unwrap();
19727 assert_eq!(r.token_ids.len(), 4);
19728 }
19729
19730 fn spark_test_pipeline(trim: Option<(usize, usize)>) -> Pipeline {
19733 let (hs, nh) = (16usize, 4usize);
19734 let mut p = create_test_pipeline(hs, 24, nh, 2, 8, 4, 64);
19735 p.layer_dump = None;
19736 p.swa = Some((6, 4));
19737 p.proj_gate_sigmoid = true;
19738 for (li, lw) in p.weights.layers.iter_mut().enumerate() {
19739 if let AttnKind::Full { softplus_gate, .. } = &mut lw.attn {
19740 let g: Vec<f32> = (0..nh * hs)
19741 .map(|i| (((i * 29 + li * 7) % 83) as f32 / 83.0 - 0.5) * 0.8)
19742 .collect();
19743 *softplus_gate = Some((QTensor::from_f32(g, nh, hs), true));
19744 }
19745 }
19746 p.swa_trim = trim;
19747 p
19748 }
19749
19750 #[test]
19756 fn swa_trim_pipeline_matches_untrimmed_bitwise() {
19757 let mut a = spark_test_pipeline(None);
19758 let mut b = spark_test_pipeline(Some((2, 4)));
19759 assert_eq!(
19760 (0..4).map(|li| b.layer_window(li)).collect::<Vec<_>>(),
19761 vec![Some(6), Some(6), Some(6), None]
19762 );
19763 let ids: Vec<u32> = (0..41u32).map(|i| (i * 11 + 5) % 64).collect();
19764 let mut pos = 0usize;
19765 for &n in [5usize, 13, 7, 16].iter().cycle() {
19766 if pos >= ids.len() {
19767 break;
19768 }
19769 let end = (pos + n).min(ids.len());
19770 let ha = a.prefill_batch_span(PrefillIn::Ids(&ids[pos..end]), pos, None, 0, 4);
19771 let hb = b.prefill_batch_span(PrefillIn::Ids(&ids[pos..end]), pos, None, 0, 4);
19772 assert_eq!(f32_bits(&ha), f32_bits(&hb), "prefill chunk at {pos}");
19773 pos = end;
19774 }
19775 for _ in 0..23 {
19776 let e = a.embed_single(((pos * 7) % 64) as u32);
19777 let ha = a.forward_layers(&e, pos, None);
19778 let hb = b.forward_layers(&e, pos, None);
19779 assert_eq!(f32_bits(&ha), f32_bits(&hb), "decode at {pos}");
19780 for li in 0..3 {
19781 assert!(b.kv_cache.layers[li].seq_len <= 12, "layer {li} bounded");
19782 }
19783 pos += 1;
19784 }
19785 for round in 0..9 {
19786 let (e1, e2) = (a.embed_single(round * 3 + 1), a.embed_single(round * 5 + 2));
19787 let (a1, a2) = a.forward_pair(&e1, &e2, pos);
19788 let (b1, b2) = b.forward_pair(&e1, &e2, pos);
19789 assert_eq!(f32_bits(&a1), f32_bits(&b1), "pair lane 1 round {round}");
19790 assert_eq!(f32_bits(&a2), f32_bits(&b2), "pair lane 2 round {round}");
19791 if round % 2 == 1 {
19792 for p in [&mut a, &mut b] {
19794 for l in &mut p.kv_cache.layers {
19795 l.truncate_last(2);
19796 }
19797 }
19798 } else {
19799 pos += 2;
19800 }
19801 }
19802 let e = a.embed_single(9);
19803 let (ha, hb) = (
19804 a.forward_layers(&e, pos, None),
19805 b.forward_layers(&e, pos, None),
19806 );
19807 assert_eq!(f32_bits(&ha), f32_bits(&hb), "after the pairs");
19808 for li in 0..3 {
19809 let (la, lb) = (&a.kv_cache.layers[li], &b.kv_cache.layers[li]);
19810 assert!(lb.base() > 0, "layer {li} trimmed");
19811 assert_eq!(la.base(), 0);
19812 assert_eq!(lb.pos_len(), la.seq_len);
19813 assert_eq!(lb.head_keys(0), &la.head_keys(0)[lb.base() * 8..]);
19814 }
19815 let (ga, gb) = (&a.kv_cache.layers[3], &b.kv_cache.layers[3]);
19816 assert_eq!(
19817 (gb.base(), gb.seq_len),
19818 (0, ga.seq_len),
19819 "the global layer keeps all"
19820 );
19821 assert_eq!(b.kv_cache.seq_len(), a.kv_cache.seq_len());
19822 assert!(b.kv_cache.total_memory_bytes() < a.kv_cache.total_memory_bytes());
19823 }
19824
19825 #[test]
19830 fn swa_trim_wire_handoff_continues_bitwise() {
19831 let mut a = spark_test_pipeline(None);
19832 let mut b = spark_test_pipeline(Some((2, 4)));
19833 let ids: Vec<u32> = (0..29u32).map(|i| (i * 13 + 7) % 64).collect();
19834 let ha = a.prefill_batch_span(PrefillIn::Ids(&ids), 0, None, 0, 4);
19835 let hb = b.prefill_batch_span(PrefillIn::Ids(&ids), 0, None, 0, 4);
19836 assert_eq!(f32_bits(&ha), f32_bits(&hb));
19837 let mut pos = ids.len();
19838 for _ in 0..5 {
19839 let e = a.embed_single(((pos * 5) % 64) as u32);
19840 assert_eq!(
19841 f32_bits(&a.forward_layers(&e, pos, None)),
19842 f32_bits(&b.forward_layers(&e, pos, None))
19843 );
19844 pos += 1;
19845 }
19846 let mut c = spark_test_pipeline(Some((2, 4)));
19847 for li in 0..4 {
19848 let bytes = b.kv_cache.layers[li].export_wire(false).unwrap();
19849 c.kv_cache.layers[li].import_wire(&bytes).unwrap();
19850 }
19851 assert!(c.kv_cache.layers[0].base() > 0, "a tail travelled");
19852 assert_eq!(c.kv_cache.seq_len(), pos);
19853 for step in 0..17 {
19854 let e = a.embed_single(((pos * 3 + 1) % 64) as u32);
19855 let ha = a.forward_layers(&e, pos, None);
19856 let hc = c.forward_layers(&e, pos, None);
19857 assert_eq!(
19858 f32_bits(&ha),
19859 f32_bits(&hc),
19860 "after the hand-off, step {step}"
19861 );
19862 pos += 1;
19863 }
19864 assert!(c.kv_cache.layers[0].seq_len <= 12, "the receiver trims on");
19865 }
19866
19867 fn mimo_test_mtp(n: usize, gain: f32) -> mimo_mtp::MimoMtp {
19870 let (hs, inter, nh, hd, vd, nkv) = (16usize, 24usize, 4usize, 8usize, 4usize, 2usize);
19871 let synth = |len: usize, salt: usize| -> Vec<f32> {
19872 (0..len)
19873 .map(|i| (((i * 37 + salt * 13 + 3) % 89) as f32 / 89.0 - 0.5) * 0.6 * gain)
19874 .collect()
19875 };
19876 let qt = |rows: usize, cols: usize, salt: usize| {
19877 QTensor::from_f32(synth(rows * cols, salt), rows, cols)
19878 };
19879 let layers = (0..n)
19880 .map(|k| {
19881 let s = 500 + k * 40;
19882 let mut kv = crate::kv_cache::LayerKvCache::new(nkv, hd);
19883 kv.sinks = Some(vec![0.3, -0.7, 1.1, 0.0]);
19884 MtpModule {
19885 enorm: vec![1.0; hs],
19886 hnorm: vec![1.0; hs],
19887 eh_proj: qt(hs, 2 * hs, s),
19888 layer: LayerWeights {
19889 input_norm: vec![1.0; hs],
19890 post_norm: vec![1.0; hs],
19891 attn_out_norm: None,
19892 ffn_out_norm: None,
19893 layer_scale: None,
19894 attn: AttnKind::Full {
19895 wq: qt(nh * hd, hs, s + 1),
19896 wk: qt(nkv * hd, hs, s + 2),
19897 wv: qt(nkv * vd, hs, s + 3),
19898 wo: qt(hs, nh * vd, s + 4),
19899 q_norm: None,
19900 k_norm: None,
19901 output_gate: false,
19902 softplus_gate: None,
19903 bias: None,
19904 },
19905 ffn: FfnKind::Dense(DenseFfn {
19906 gate_proj: qt(inter, hs, s + 5),
19907 up_proj: qt(inter, hs, s + 6),
19908 down_proj: qt(hs, inter, s + 7),
19909 act: Act::Silu,
19910 down_t: None,
19911 segs: Vec::new(),
19912 }),
19913 },
19914 final_norm: vec![1.0; hs],
19915 kv,
19916 }
19917 })
19918 .collect();
19919 mimo_mtp::MimoMtp::from_layers(layers)
19920 }
19921
19922 fn mimo_greedy(p: &mut Pipeline, ids: &[u32], n: usize, spec: bool) -> GenerateResult {
19923 p.clear_sequence_state();
19924 p.speculative = spec;
19925 p.ignore_eos = true;
19926 p.sampler_config.temperature = 0.0;
19927 p.generate_from_ids(ids, n, None, None).unwrap()
19928 }
19929
19930 #[test]
19936 fn mimo_mtp_incremental_rounds_equal_the_teacher_forced_table() {
19937 for post in [false, true] {
19940 let mut p = mimo_test_pipeline();
19941 p.weights.final_norm = (0..p.hidden_size).map(|i| 0.5 + 0.1 * i as f32).collect();
19943 let mut st0 = mimo_test_mtp(3, 1.0);
19944 st0.post_norm_hidden = post;
19945 p.mimo_mtp = Some(st0);
19946 let ids: Vec<u32> = (0..14u32).map(|i| (i * 11 + 5) % 64).collect();
19947 let hs = p.hidden_size;
19948 let hb = p.prefill_batch_span(PrefillIn::Ids(&ids), 0, None, 0, p.num_layers);
19949 p.mimo_note_rows(&hb, 0);
19950 let mut st = p.mimo_mtp.take().unwrap();
19951 let k = 3;
19954 let mut inc = Vec::new();
19955 for t in 0..ids.len() - k - 1 {
19956 inc.push(p.mimo_mtp_draft(&mut st, t, &ids, k));
19957 }
19958 let s = ids.len();
19961 let mut reference = vec![vec![0u32; k]; s - k - 1];
19962 let mut fresh = mimo_test_mtp(3, 1.0);
19963 for (layer, m) in fresh.layers.iter_mut().enumerate() {
19964 let n = s - layer - 1;
19965 let mut cats = vec![0.0f32; n * 2 * hs];
19966 for j in 0..n {
19967 let e = p.embed_single(ids[j + layer + 1]);
19968 let raw = &hb[j * hs..(j + 1) * hs];
19969 let g = if post {
19970 inference::rms_norm(raw, &p.weights.final_norm, p.rms_eps, p.norm_style)
19971 } else {
19972 raw.to_vec()
19973 };
19974 let (ce, ch) = cats[j * 2 * hs..(j + 1) * 2 * hs].split_at_mut(hs);
19975 inference::rms_norm_into(&e, &m.enorm, p.rms_eps, p.norm_style, ce);
19976 inference::rms_norm_into(&g, &m.hnorm, p.rms_eps, p.norm_style, ch);
19977 }
19978 let mut x = vec![0.0f32; n * hs];
19979 m.eh_proj.matmat(&cats, n, &mut x, None);
19980 p.mimo_mtp_block(m, &mut x, n, 0);
19981 for (t, row) in reference.iter_mut().enumerate() {
19982 let y = inference::rms_norm(
19983 &x[t * hs..(t + 1) * hs],
19984 &m.final_norm,
19985 p.rms_eps,
19986 p.norm_style,
19987 );
19988 row[layer] = sampler::argmax(&p.lm_head_forward(&y));
19989 }
19990 }
19991 assert_eq!(inc, reference, "post_norm_hidden = {post}");
19992 let distinct: std::collections::HashSet<u32> =
19994 inc.iter().flatten().copied().collect();
19995 assert!(distinct.len() > 3, "{inc:?}");
19996 let last_t = ids.len() - k - 2;
19998 for m in &st.layers {
19999 assert_eq!(m.kv.seq_len, last_t + 1);
20000 }
20001 }
20002 }
20003
20004 #[test]
20011 fn mimo_speculative_greedy_equals_plain_greedy() {
20012 unsafe { std::env::set_var("CMF_GPU_WGPU_GRAPH", "0") };
20013 let ids: Vec<u32> = (0..9u32).map(|i| (i * 7 + 3) % 64).collect();
20014 let n = 24;
20015 let mut p = mimo_test_pipeline();
20016 let plain = mimo_greedy(&mut p, &ids, n, false);
20017 assert_eq!(plain.mtp_drafted, 0);
20018 assert_eq!(plain.token_ids.len(), n);
20019 let plain_kv = p.kv_cache.layers[0].seq_len;
20020
20021 p.mimo_mtp = Some(mimo_test_mtp(3, 1.0));
20023 let spec = mimo_greedy(&mut p, &ids, n, true);
20024 assert!(spec.mtp_drafted > 0, "the round must draft");
20025 assert_eq!(spec.token_ids, plain.token_ids);
20026 assert_eq!(p.kv_cache.layers[0].seq_len, plain_kv);
20027
20028 let mut truth: Vec<u32> = ids.clone();
20031 truth.extend(&plain.token_ids);
20032 let mut noisy = truth.clone();
20033 for (i, t) in noisy.iter_mut().enumerate() {
20034 if i % 5 == 0 {
20035 *t = (*t + 1) % 64;
20036 }
20037 }
20038 let mut st = mimo_test_mtp(3, 1.0);
20039 st.draft_override = Some(noisy);
20040 p.mimo_mtp = Some(st);
20041 let spec = mimo_greedy(&mut p, &ids, n, true);
20042 assert_eq!(spec.token_ids, plain.token_ids);
20043 assert_eq!(p.kv_cache.layers[0].seq_len, plain_kv);
20044 let stats = p.mimo_mtp.as_ref().unwrap().stats.clone();
20045 assert_eq!(stats.accepted as usize, spec.mtp_accepted);
20046 assert!(spec.mtp_accepted > 0 && spec.mtp_accepted < spec.mtp_drafted);
20047 assert!(
20048 stats.accept_hist.iter().filter(|&&c| c > 0).count() >= 3,
20049 "{:?}",
20050 stats.accept_hist
20051 );
20052 assert!(stats.tokens_per_round() > 1.5, "{}", stats.line());
20053
20054 let mut st = mimo_test_mtp(3, 1.0);
20057 st.draft_override = Some(truth);
20058 p.mimo_mtp = Some(st);
20059 let spec = mimo_greedy(&mut p, &ids, n, true);
20060 assert_eq!(spec.token_ids, plain.token_ids);
20061 assert_eq!(spec.mtp_accepted, spec.mtp_drafted);
20062 assert_eq!(p.kv_cache.layers[0].seq_len, plain_kv);
20063
20064 let off = mimo_greedy(&mut p, &ids, n, false);
20066 assert_eq!(off.token_ids, plain.token_ids);
20067 assert_eq!(off.mtp_drafted, 0);
20068 }
20069
20070 #[test]
20079 fn mimo_shaped_model_rides_the_wgpu_graph_geometry() {
20080 let p = mimo_test_pipeline();
20081 assert_eq!(
20082 p.graph_attn_decline_reason(),
20083 Some("per-layer KV head counts")
20084 );
20085 assert_eq!(p.wgpu_graph_attn_decline(), None);
20086 let g0 = p.graph_attn_geom(0).expect("full layer geometry");
20087 assert_eq!(
20088 (g0.nkv, g0.dv, g0.rd, g0.window, g0.sink.is_some()),
20089 (1, 4, 4, None, false)
20090 );
20091 assert_eq!(g0.invf, p.inv_freq.as_slice());
20092 let g1 = p.graph_attn_geom(1).expect("sliding layer geometry");
20093 assert_eq!((g1.nkv, g1.dv, g1.rd, g1.window), (2, 4, 4, Some(3)));
20094 assert_eq!(g1.sink, Some(&[0.5f32, -1.0, 1.5, 0.0][..]));
20095 assert_eq!(g1.invf, p.inv_freq_local.as_ref().unwrap().as_slice());
20096 assert_ne!(g0.invf, g1.invf, "two RoPE tables");
20097 let g3 = p.graph_attn_geom(3).expect("full layer geometry");
20098 assert_eq!((g3.nkv, g3.window, g3.sink.is_some()), (1, None, false));
20099
20100 let emb = p.embed_single(3);
20103 let mut lg = Vec::new();
20104 assert!(
20105 p.try_token_graph_wgpu_steps(&emb, 0, &mut lg, 1, None, None, 0, p.num_layers)
20106 .is_none()
20107 );
20108 let mut hid = emb.clone();
20109 assert_eq!(
20110 p.try_batch_graph_wgpu(&mut hid, &[0], 1, None),
20111 crate::gpu::BatchGraphOutcome::Declined
20112 );
20113 assert_eq!(hid, emb, "a declined batch graph leaves the rows untouched");
20114 assert!(p.try_multi_burst(3, 0, 4).is_none());
20115 assert!(
20116 p.graph_declines().is_empty(),
20117 "no attention decline logged: {:?}",
20118 p.graph_declines()
20119 );
20120 let plain = || create_test_pipeline(8, 16, 2, 1, 4, 2, 32);
20125 assert_eq!(plain().graph_attn_decline_reason(), None);
20126 assert_eq!(plain().wgpu_graph_attn_decline(), None);
20127 assert!(
20128 plain().graph_attn_geom(0).is_none(),
20129 "uniform models keep the historical arms"
20130 );
20131 let mut q = plain();
20132 q.set_layer_sinks(1, vec![0.25, -0.25]).unwrap();
20133 assert_eq!(
20134 q.graph_attn_decline_reason(),
20135 Some("learned attention sinks")
20136 );
20137 assert_eq!(
20138 q.graph_attn_geom(1).unwrap().sink,
20139 Some(&[0.25f32, -0.25][..])
20140 );
20141 let mut q = plain();
20142 q.set_attn_geometry(None, Some(2)).unwrap();
20143 assert_eq!(
20144 q.graph_attn_decline_reason(),
20145 Some("V heads narrower than Q/K heads")
20146 );
20147 assert_eq!(q.graph_attn_geom(0).unwrap().dv, 2);
20148 let mut q = plain();
20149 q.sliding_layers = Some(vec![true, false]);
20150 q.swa = Some((4, usize::MAX));
20151 assert_eq!(q.graph_attn_decline_reason(), Some("sliding-window layers"));
20152 assert_eq!(q.graph_attn_geom(0).unwrap().window, Some(4));
20153 assert_eq!(q.graph_attn_geom(1).unwrap().window, None);
20154
20155 let mut q = mimo_test_pipeline();
20158 q.rope_scale = 2.0;
20159 assert_eq!(
20160 q.wgpu_graph_attn_decline(),
20161 Some("scaled RoPE positions with per-layer geometry")
20162 );
20163 let emb = q.embed_single(3);
20164 assert!(
20165 q.try_token_graph_wgpu_steps(&emb, 0, &mut lg, 1, None, None, 0, q.num_layers)
20166 .is_none()
20167 );
20168 let _ = q.try_token_graph_wgpu_steps(&emb, 1, &mut lg, 1, None, None, 0, q.num_layers);
20169 let lines = q.graph_declines();
20170 assert_eq!(
20171 lines
20172 .iter()
20173 .filter(|l| l.starts_with("wgpu token graph") && l.contains("scaled RoPE"))
20174 .count(),
20175 1,
20176 "{lines:?}"
20177 );
20178 }
20179
20180 #[test]
20181 fn mimo_verify_rewind_preserves_lagging_host_caches() {
20182 let mut p = mimo_test_pipeline();
20183 for (li, layer) in p.kv_cache.layers.iter_mut().enumerate() {
20184 let row = vec![0.0; layer.num_kv_heads * layer.head_dim];
20185 for _ in 0..if li == 0 { 2 } else { 12 } {
20186 layer.append(&row, &row, &[]);
20187 }
20188 }
20189 p.mimo_verify_rewind(9).unwrap();
20190 assert_eq!(p.kv_cache.layers[0].seq_len, 2);
20191 for layer in &p.kv_cache.layers[1..] {
20192 assert_eq!(layer.seq_len, 9);
20193 }
20194 }
20195
20196 #[test]
20200 fn layer_dump_covers_every_position_and_layer_on_both_walks() {
20201 let dir = std::env::temp_dir().join(format!("cmf-layer-dump-{}", std::process::id()));
20202 let _ = std::fs::remove_dir_all(&dir);
20203 let mut p = mimo_test_pipeline();
20204 let hs = p.hidden_size;
20205 let ids = [5u32, 9, 11, 2, 40];
20206 p.layer_dump = Some(dir.join("decode"));
20207 for (pos, &id) in ids.iter().enumerate() {
20208 let e = p.embed_single(id);
20209 let _ = p.forward_layers(&e, pos, None);
20210 }
20211 p.clear_sequence_state();
20212 p.layer_dump = Some(dir.join("prefill"));
20213 let hb = p.prefill_batch_span(PrefillIn::Ids(&ids), 0, None, 0, p.num_layers);
20214 for pos in 0..ids.len() {
20215 for li in 0..p.num_layers {
20216 let name = format!("p{pos:06}_l{li:02}.f32");
20217 let a = std::fs::read(dir.join("decode").join(&name)).unwrap();
20218 let b = std::fs::read(dir.join("prefill").join(&name)).unwrap();
20219 assert_eq!(a.len(), hs * 4, "{name}");
20220 assert_eq!(a, b, "{name}");
20221 }
20222 }
20223 let last = std::fs::read(dir.join("prefill").join("p000004_l03.f32")).unwrap();
20224 let vals: Vec<f32> = last
20225 .chunks(4)
20226 .map(|c| f32::from_le_bytes(c.try_into().unwrap()))
20227 .collect();
20228 assert_eq!(f32_bits(&vals), f32_bits(&hb[4 * hs..5 * hs]));
20229 let _ = std::fs::remove_dir_all(&dir);
20230 }
20231
20232 #[test]
20233 fn attn_geometry_and_sinks_are_validated() {
20234 let mut p = create_test_pipeline(8, 16, 4, 2, 4, 2, 32);
20235 assert!(
20236 p.set_attn_geometry(Some(vec![2]), None).is_err(),
20237 "one entry per layer"
20238 );
20239 assert!(
20240 p.set_attn_geometry(Some(vec![2, 3]), None).is_err(),
20241 "3 does not divide 4"
20242 );
20243 assert!(p.set_attn_geometry(Some(vec![2, 0]), None).is_err());
20244 assert!(p.set_attn_geometry(None, Some(0)).is_err());
20245 assert!(
20246 p.set_attn_geometry(None, Some(5)).is_err(),
20247 "V wider than the head"
20248 );
20249 p.set_attn_geometry(None, Some(4)).unwrap();
20250 assert_eq!(
20251 p.v_head_dim, None,
20252 "v_head_dim == head_dim is the uniform case"
20253 );
20254 p.set_layer_sinks(1, vec![0.1; 4]).unwrap();
20255 p.set_attn_geometry(Some(vec![1, 4]), None).unwrap();
20256 assert_eq!(p.kv_cache.layers[0].num_kv_heads, 1);
20257 assert_eq!(p.kv_cache.layers[1].num_kv_heads, 4);
20258 assert!(
20259 p.kv_cache.layers[1].sinks.is_some(),
20260 "a reshape keeps the layer's sinks"
20261 );
20262 assert_eq!(p.layer_geom(1).0, 4);
20263 assert!(
20264 p.set_layer_sinks(0, vec![0.0; 3]).is_err(),
20265 "one sink per Q head"
20266 );
20267 assert!(p.set_layer_sinks(7, vec![0.0; 4]).is_err());
20268 assert!(p.set_layer_sinks(0, vec![f32::NAN, 0.0, 0.0, 0.0]).is_err());
20269 }
20270
20271 #[test]
20274 fn o1_is_never_armed_on_sink_window_or_narrow_v_layers() {
20275 let cfg = || {
20276 Some(crate::nystrom::O1Cfg {
20277 layers: crate::nystrom::O1Layers::All,
20278 m: 4,
20279 w: 8,
20280 sink: 2,
20281 rect: crate::nystrom::O1Rect::Aggregate,
20282 })
20283 };
20284 let mut p = mimo_test_pipeline();
20285 p.set_o1(cfg());
20286 assert!(!p.o1_active(), "every MiMo-shaped layer is ineligible");
20287 let mut q = create_test_pipeline(8, 16, 2, 1, 4, 3, 64);
20288 q.set_layer_sinks(1, vec![0.0, 0.0]).unwrap();
20289 q.sliding_layers = Some(vec![false, false, true]);
20290 q.swa = Some((4, usize::MAX));
20291 q.set_o1(cfg());
20292 assert_eq!(q.o1_flags, vec![true, false, false]);
20293 }
20294
20295 #[test]
20296 fn o1_batch_transition_publishes_one_epoch_before_serial_handoff() {
20297 const B: usize = 19;
20298 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
20299 p.set_o1(Some(crate::nystrom::O1Cfg {
20300 layers: crate::nystrom::O1Layers::All,
20301 m: 4,
20302 w: 8,
20303 sink: 2,
20304 rect: crate::nystrom::O1Rect::Aggregate,
20305 }));
20306 p.o1_begin_with_prefix(Some(B));
20307 let ids: Vec<u32> = (0..B as u32).collect();
20308 let _ = p.prefill_batch_span(PrefillIn::Ids(&ids), 0, None, 0, p.num_layers);
20309
20310 assert_eq!(p.o1_epoch, 1, "all layers publish one completed transition");
20311 assert!(p.kv_cache.layers.iter().all(|l| l.o1_sealed()));
20312 let next = p.embed_single(B as u32);
20313 let _ = p.forward_layers(&next, B, None);
20314 assert_eq!(p.o1_epoch, 1, "sealed handoff must not republish the epoch");
20315 }
20316
20317 #[test]
20318 fn o1_pair_transition_commits_scratch_before_epoch_publication() {
20319 const B: usize = 19;
20320 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
20321 let gdn_cfg = crate::linear_core::GdnCfg {
20325 num_v_heads: 2,
20326 num_k_heads: 1,
20327 key_head_dim: 2,
20328 value_head_dim: 4,
20329 conv_kernel: 3,
20330 hidden_size: 8,
20331 rms_eps: 1e-6,
20332 output_gate_sigmoid: false,
20333 };
20334 let synth = |n: usize, salt: usize| -> Vec<f32> {
20335 (0..n)
20336 .map(|i| (((i * 13 + salt * 7) % 97) as f32 / 97.0 - 0.5) * 0.4)
20337 .collect()
20338 };
20339 let qt = |rows: usize, cols: usize, salt: usize| {
20340 crate::qtensor::QTensor::from_f32(synth(rows * cols, salt), rows, cols)
20341 };
20342 let c_dim = gdn_cfg.conv_dim();
20343 let vd = gdn_cfg.num_v_heads * gdn_cfg.value_head_dim;
20344 p.weights.layers[0].attn = AttnKind::LinearGdn(crate::linear_core::GdnWeights {
20345 in_proj_qkv: qt(c_dim, 8, 1),
20346 in_proj_z: qt(vd, 8, 2),
20347 in_proj_a: qt(gdn_cfg.num_v_heads, 8, 3),
20348 in_proj_b: qt(gdn_cfg.num_v_heads, 8, 4),
20349 conv1d: synth(c_dim * gdn_cfg.conv_kernel, 5),
20350 a_log: vec![0.2, 0.5],
20351 dt_bias: synth(gdn_cfg.num_v_heads, 6),
20352 norm: vec![1.0; gdn_cfg.value_head_dim],
20353 out_proj: qt(8, vd, 7),
20354 });
20355 p.gdn_cfg = Some(gdn_cfg);
20356 p.set_o1(Some(crate::nystrom::O1Cfg {
20357 layers: crate::nystrom::O1Layers::All,
20358 m: 4,
20359 w: 8,
20360 sink: 2,
20361 rect: crate::nystrom::O1Rect::Aggregate,
20362 }));
20363 p.o1_begin_with_prefix(Some(B));
20364 for pos in 0..B - 2 {
20365 let emb = p.embed_single(pos as u32);
20366 let _ = p.forward_layers(&emb, pos, None);
20367 }
20368 let lane1_state = p.kv_cache.layers[0].linear_state.clone();
20369
20370 let e1 = p.embed_single((B - 2) as u32);
20371 let e2 = p.embed_single((B - 1) as u32);
20372 let _ = p.forward_pair(&e1, &e2, B - 2);
20373
20374 assert_eq!(p.o1_epoch, 1, "pair crossing B publishes one epoch");
20375 assert!(
20376 p.kv_cache
20377 .layers
20378 .iter()
20379 .enumerate()
20380 .all(|(li, l)| !p.o1_flags[li] || l.o1_sealed())
20381 );
20382 assert!(!p.kv_cache.layers[0].linear_state.is_empty());
20383 assert_ne!(
20384 p.kv_cache.layers[0].linear_state, lane1_state,
20385 "real pair must commit GDN lane 2 before returning"
20386 );
20387 assert!(p.kv_cache.layers[0].linear_scratch.is_empty());
20388 let next = p.embed_single(B as u32);
20389 let _ = p.forward_layers(&next, B, None);
20390 assert_eq!(p.o1_epoch, 1, "serial continuation must reuse the epoch");
20391 }
20392
20393 #[test]
20394 fn o1_error_observation_stays_terminal_until_reset() {
20395 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
20396 p.set_o1(Some(crate::nystrom::O1Cfg {
20397 layers: crate::nystrom::O1Layers::All,
20398 m: 4,
20399 w: 8,
20400 sink: 2,
20401 rect: crate::nystrom::O1Rect::Aggregate,
20402 }));
20403 p.o1_begin();
20404 p.kv_cache.layers[0].o1_abort("synthetic transition failure".into());
20405
20406 assert!(p.o1_seal_checked().is_err());
20407 assert!(
20408 p.o1_seal_checked().is_err(),
20409 "retry must see the sticky error"
20410 );
20411 let k = vec![0.2f32; 4];
20412 let v = vec![0.3f32; 4];
20413 p.kv_cache.layers[0].append(&k, &v, &[]);
20414 assert_eq!(p.kv_cache.layers[0].seq_len, 0);
20415
20416 p.reset_session();
20417 p.o1_begin();
20418 p.kv_cache.layers[0].append(&k, &v, &[]);
20419 assert_eq!(p.kv_cache.layers[0].seq_len, 1);
20420 }
20421
20422 #[test]
20423 fn nll_graph_failure_is_terminal_and_request_is_reusable() {
20424 let ids = vec![1u32, 2, 3, 4, 5, 6];
20425 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
20426 p.graph_logits = Some(vec![123.0]);
20427 p.graph_want_logits = true;
20428 p.graph_failed
20429 .store(true, std::sync::atomic::Ordering::Relaxed);
20430 p.cancel.store(true, std::sync::atomic::Ordering::Relaxed);
20431 let err = p.nll_ids_from(&ids, 0).expect_err("prior graph failure");
20432 assert!(err.contains("before NLL"));
20433 assert!(p.graph_logits.is_none());
20434 assert!(!p.graph_want_logits);
20435 assert!(!p.graph_failed.load(std::sync::atomic::Ordering::Relaxed));
20436 assert!(!p.cancel.load(std::sync::atomic::Ordering::Relaxed));
20437
20438 let mut fresh = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
20439 let expected = fresh.nll_ids_from(&ids, 0).expect("fresh NLL");
20440 let actual = p.nll_ids_from(&ids, 0).expect("reused NLL");
20441 assert_eq!(actual.1, expected.1);
20442 assert!((actual.0 - expected.0).abs() < 1e-9);
20443 }
20444
20445 #[test]
20446 fn nll_forward_failure_discards_partial_score_and_clears_sidechannels() {
20447 let ids = vec![1u32, 2, 3, 4, 5, 6];
20448 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
20449 p.nll_test_fail_at = Some(1);
20450 let err = p
20451 .nll_ids_from(&ids, 0)
20452 .expect_err("one-shot forward failure");
20453 assert!(err.contains("forward") || err.contains("score row"));
20454 assert!(!p.graph_failed.load(std::sync::atomic::Ordering::Relaxed));
20455 assert!(!p.graph_want_logits);
20456 assert!(p.graph_logits.is_none());
20457 assert!(p.kv_history.is_empty());
20458
20459 let mut fresh = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
20460 let expected = fresh.nll_ids_from(&ids, 0).expect("fresh NLL");
20461 let actual = p.nll_ids_from(&ids, 0).expect("reused NLL");
20462 assert_eq!(actual.1, expected.1);
20463 assert!((actual.0 - expected.0).abs() < 1e-9);
20464 }
20465
20466 #[test]
20467 fn nll_serial_failure_before_first_row_is_reported() {
20468 let ids = vec![1u32, 2, 3, 4];
20469 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
20470 p.nll_test_force_serial = true;
20471 p.nll_test_fail_at = Some(0);
20472 let err = p.nll_ids_from(&ids, 0).expect_err("serial forward failure");
20473 assert!(err.contains("serial forward"));
20474 assert!(p.kv_history.is_empty());
20475 assert!(!p.graph_failed.load(std::sync::atomic::Ordering::Relaxed));
20476 assert!(!p.cancel.load(std::sync::atomic::Ordering::Relaxed));
20477 }
20478
20479 #[test]
20480 fn ffn_probe_failure_discards_recorder_and_state() {
20481 let ids = vec![1u32, 2, 3, 4];
20482 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
20483 p.nll_test_fail_at = Some(0);
20484 let err = p
20485 .probe_ffn_mass_batch(&ids)
20486 .expect_err("probe forward failure");
20487 assert!(err.contains("NLL"));
20488 assert!(FFN_PROBE.with(|probe| probe.borrow().is_none()));
20489 assert!(p.kv_history.is_empty());
20490 assert!(!p.graph_failed.load(std::sync::atomic::Ordering::Relaxed));
20491 }
20492
20493 #[test]
20494 fn nll_test_controls_are_pipeline_scoped() {
20495 let ids = vec![1u32, 2, 3, 4];
20496 let mut failing = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
20497 let mut unaffected = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
20498 failing.nll_test_force_serial = true;
20499 failing.nll_test_fail_at = Some(0);
20500
20501 assert!(!failing.can_prefill_batched());
20502 assert!(unaffected.can_prefill_batched());
20503 let expected = unaffected
20504 .nll_ids_from(&ids, 0)
20505 .expect("unaffected pipeline remains usable");
20506 let err = failing
20507 .nll_ids_from(&ids, 0)
20508 .expect_err("failure injection belongs to failing pipeline");
20509 assert!(err.contains("serial forward"));
20510 assert!(failing.nll_test_fail_at.is_none());
20511 assert!(unaffected.can_prefill_batched());
20512 let actual = unaffected
20513 .nll_ids_from(&ids, 0)
20514 .expect("unaffected pipeline remains reusable");
20515 assert_eq!(actual.1, expected.1);
20516 assert!((actual.0 - expected.0).abs() < 1e-9);
20517 }
20518
20519 #[test]
20520 fn forward_ids_failure_channel_is_terminal_and_reusable() {
20521 let ids = vec![1u32, 2, 3, 4, 5, 6];
20522 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
20523 p.graph_logits = Some(vec![123.0]);
20524 p.graph_want_logits = true;
20525 p.graph_failed
20526 .store(true, std::sync::atomic::Ordering::Relaxed);
20527 p.cancel.store(true, std::sync::atomic::Ordering::Relaxed);
20528
20529 let err = p
20530 .forward_ids(&ids, None)
20531 .expect_err("a failed forward must not become a valid head result");
20532 assert!(err.contains("forward_ids setup"));
20533 assert!(p.graph_logits.is_none());
20534 assert!(!p.graph_want_logits);
20535 assert!(!p.graph_failed.load(std::sync::atomic::Ordering::Relaxed));
20536 assert!(!p.cancel.load(std::sync::atomic::Ordering::Relaxed));
20537 assert_eq!(p.kv_cache.seq_len(), 0);
20538
20539 let expected = create_test_pipeline(8, 16, 2, 1, 4, 1, 64)
20540 .forward_ids(&ids, None)
20541 .expect("fresh forward_ids");
20542 let actual = p
20543 .forward_ids(&ids, None)
20544 .expect("pipeline remains reusable after a failed forward");
20545 assert_eq!(actual.len(), expected.len());
20546 assert!(
20547 actual
20548 .iter()
20549 .zip(expected)
20550 .all(|(a, b)| (a - b).abs() < 1e-9)
20551 );
20552 assert_eq!(p.kv_cache.seq_len(), ids.len());
20553 }
20554
20555 #[test]
20556 fn sigmoid_router_floor_is_explicit_per_architecture() {
20557 let zero = || DenseFfn {
20563 gate_proj: QTensor::from_f32(vec![0.0; 4], 2, 2),
20564 up_proj: QTensor::from_f32(vec![0.0; 4], 2, 2),
20565 down_proj: QTensor::from_f32(vec![0.0; 4], 2, 2),
20566 act: Act::Silu,
20567 down_t: None,
20568 segs: Vec::new(),
20569 };
20570 let m = MoeFfn {
20571 router: QTensor::from_f32(vec![0.0; 4], 2, 2),
20572 experts: vec![zero(), zero()],
20573 top_k: 1,
20574 norm_topk_prob: true,
20575 router_sigmoid: true,
20576 expert_bias: None,
20577 routed_scaling: 2.5,
20578 route_tau: None,
20579 shared: None,
20580 stats: std::cell::RefCell::new(Vec::new()),
20581 act_sq: std::cell::RefCell::new(Vec::new()),
20582 act_rows: std::cell::RefCell::new(Vec::new()),
20583 mask: None,
20584 per_expert_scale: None,
20585 router_input_norm: false,
20586 resonance: None,
20587 grown: Vec::new(),
20588 };
20589 let logits = [-20.0f32, -20.0];
20590 let (_, p, glm_wsum) = moe_route_with_eps(&logits, &m, None, 1e-20);
20591 let (_, _, generic_wsum) = moe_route(&logits, &m, None);
20592 let expected = (p[0] + 1e-20) / m.routed_scaling;
20593 assert!((glm_wsum - expected).abs() < 1e-15);
20594 assert!(generic_wsum > glm_wsum * 100.0);
20595 }
20596
20597 #[test]
20598 fn resonance_scores_match_formula_and_stable_tie() {
20599 let r = Resonance {
20600 mu: vec![0.0, 0.0, 1.0, 0.0, 0.0, 1.0],
20602 u: vec![1.0, 0.0, 0.0, 1.0, 1.0, 1.0],
20603 k: 1,
20604 bias: vec![1.5, 0.5, 0.0],
20605 shell: Vec::new(),
20606 };
20607 let x = [1.0f32, 1.0];
20608 let mut got = vec![0.0; 3];
20609 r.scores(&x, &mut got);
20610 assert!((got[0] - 0.5).abs() < 1e-6);
20614 assert!((got[1] - 0.5).abs() < 1e-6);
20615 assert!(got[2].abs() < 1e-6);
20616 let best = got
20617 .iter()
20618 .enumerate()
20619 .max_by(|(ia, a), (ib, b)| a.partial_cmp(b).unwrap().then(ib.cmp(ia)))
20620 .map(|(i, _)| i);
20621 assert_eq!(best, Some(0));
20622 assert!(got.iter().all(|v| v.is_finite()));
20623 }
20624
20625 #[test]
20630 fn resonance_shell_masks_outside_keeps_inside_and_trunk_bits() {
20631 let plain = Resonance {
20637 mu: vec![0.0, 0.0, 1.0, 0.0, 0.0, 1.0, 0.0, 1.0],
20638 u: vec![1.0, 0.0, 0.0, 1.0, 1.0, 1.0, 1.0, 1.0],
20639 k: 1,
20640 bias: vec![1.5, 0.5, 0.0, 0.0],
20641 shell: Vec::new(),
20642 };
20643 let shelled = Resonance {
20644 mu: plain.mu.clone(),
20645 u: plain.u.clone(),
20646 k: 1,
20647 bias: plain.bias.clone(),
20648 shell: vec![f32::INFINITY, f32::INFINITY, 6.0, 0.25],
20649 };
20650 assert!(!plain.has_shell());
20651 assert!(shelled.has_shell());
20652 let bits = |v: &[f32]| v.iter().map(|x| x.to_bits()).collect::<Vec<_>>();
20653 let (mut a, mut b) = (vec![0.0; 4], vec![0.0; 4]);
20654 set_growth_shell(Some(true));
20655 assert!(growth_shell_enabled());
20656 let x = [1.0f32, 1.0];
20659 plain.scores(&x, &mut a);
20660 shelled.scores(&x, &mut b);
20661 assert_eq!(bits(&a), bits(&b), "inside every shell: unchanged");
20662 assert!(a[2] == 0.0 && a[3] == 0.0);
20663 let xo = [3.0f32, 0.0];
20666 plain.scores(&xo, &mut a);
20667 shelled.scores(&xo, &mut b);
20668 assert_eq!(a[2], -6.0);
20669 assert_eq!(a[3], -6.0);
20670 assert_eq!(bits(&a[..3]), bits(&b[..3]), "trunk rows + the boundary row");
20671 assert_eq!(b[3], f32::NEG_INFINITY, "outside the shell: −∞");
20672 assert_eq!(shelled.effective_shell(4), shelled.shell);
20673 set_growth_shell(Some(false));
20676 assert!(!growth_shell_enabled());
20677 assert_eq!(shelled.effective_shell(4), vec![f32::INFINITY; 4]);
20678 shelled.scores(&xo, &mut b);
20679 assert_eq!(bits(&a), bits(&b));
20680 set_growth_shell(None);
20681 let short = Resonance {
20684 shell: vec![f32::INFINITY, f32::INFINITY],
20685 ..shelled
20686 };
20687 set_growth_shell(Some(true));
20688 short.scores(&xo, &mut b);
20689 assert_eq!(bits(&a), bits(&b));
20690 set_growth_shell(None);
20691 }
20692
20693 #[test]
20697 fn moe_route_neg_inf_logits_pick_the_best_finite_expert_with_weight_one() {
20698 let zero = || DenseFfn {
20699 gate_proj: QTensor::from_f32(vec![0.0; 4], 2, 2),
20700 up_proj: QTensor::from_f32(vec![0.0; 4], 2, 2),
20701 down_proj: QTensor::from_f32(vec![0.0; 4], 2, 2),
20702 act: Act::Silu,
20703 down_t: None,
20704 segs: Vec::new(),
20705 };
20706 let moe = |sigmoid: bool| MoeFfn {
20707 router: QTensor::from_f32(vec![0.0; 8], 4, 2),
20708 experts: vec![zero(), zero(), zero(), zero()],
20709 top_k: 1,
20710 norm_topk_prob: true,
20711 router_sigmoid: sigmoid,
20712 expert_bias: None,
20713 routed_scaling: 1.0,
20714 route_tau: None,
20715 shared: None,
20716 stats: std::cell::RefCell::new(Vec::new()),
20717 act_sq: std::cell::RefCell::new(Vec::new()),
20718 act_rows: std::cell::RefCell::new(Vec::new()),
20719 mask: None,
20720 per_expert_scale: None,
20721 router_input_norm: false,
20722 resonance: None,
20723 grown: Vec::new(),
20724 };
20725 let logits = [-1.0f32, f32::NEG_INFINITY, -0.5, f32::NEG_INFINITY];
20726 for sigmoid in [false, true] {
20727 let m = moe(sigmoid);
20728 let (idx, p, wsum) = moe_route(&logits, &m, None);
20729 assert_eq!(idx, vec![2], "sigmoid {sigmoid}");
20730 assert_eq!(p[1], 0.0);
20731 assert_eq!(p[3], 0.0);
20732 assert!(p[2] > p[0] && p[0] > 0.0);
20733 assert!(p.iter().all(|v| v.is_finite()));
20734 let w = p[2] / wsum;
20735 if sigmoid {
20736 assert!((w - 1.0).abs() < 1e-5, "sigmoid: weight {w}");
20738 } else {
20739 assert_eq!(w, 1.0, "softmax: the renormalized top-1 weight is exactly 1");
20740 }
20741 let (idx, _, _) = moe_route(&logits, &m, Some(&[true, true, true, true]));
20744 assert_eq!(idx, vec![2]);
20745 }
20746 let m = moe(false);
20748 let (idx, _, _) = moe_route(&[f32::NEG_INFINITY, f32::NEG_INFINITY, f32::NEG_INFINITY, -9.0], &m, None);
20749 assert_eq!(idx, vec![3]);
20750 let (idx, p, wsum) = moe_route(&[f32::NEG_INFINITY; 4], &m, None);
20753 assert_eq!(idx, vec![0]);
20754 assert!(p.iter().all(|&v| v == 0.25));
20755 assert!(wsum.is_finite() && wsum > 0.0);
20756 }
20757
20758 #[test]
20763 fn resonance_route_picks_the_raw_argmax_not_the_softmax_tie() {
20764 let zero = || DenseFfn {
20765 gate_proj: QTensor::from_f32(vec![0.0; 4], 2, 2),
20766 up_proj: QTensor::from_f32(vec![0.0; 4], 2, 2),
20767 down_proj: QTensor::from_f32(vec![0.0; 4], 2, 2),
20768 act: Act::Silu,
20769 down_t: None,
20770 segs: Vec::new(),
20771 };
20772 let moe = |resonant: bool, norm_topk: bool| MoeFfn {
20773 router: QTensor::from_f32(vec![0.0; 6], 3, 2),
20774 experts: vec![zero(), zero(), zero()],
20775 top_k: 1,
20776 norm_topk_prob: norm_topk,
20777 router_sigmoid: false,
20778 expert_bias: None,
20779 routed_scaling: 1.0,
20780 route_tau: None,
20781 shared: None,
20782 stats: std::cell::RefCell::new(Vec::new()),
20783 act_sq: std::cell::RefCell::new(Vec::new()),
20784 act_rows: std::cell::RefCell::new(Vec::new()),
20785 mask: None,
20786 per_expert_scale: None,
20787 router_input_norm: false,
20788 resonance: resonant.then(|| Resonance {
20789 mu: vec![0.0; 6],
20790 u: Vec::new(),
20791 k: 0,
20792 bias: vec![0.0; 3],
20793 shell: Vec::new(),
20794 }),
20795 grown: Vec::new(),
20796 };
20797 let lo = -0.1f32;
20800 let hi = f32::from_bits(lo.to_bits() - 1);
20801 assert!(hi > lo && hi - lo < 2f32.powi(-25));
20802 assert_eq!((lo - hi).exp(), 1.0, "the tie this test is about");
20803 let (idx, _, _) = moe_route(&[lo, hi], &moe(false, true), None);
20805 assert_eq!(idx, vec![0], "softmax tie → lower index (the gated contract)");
20806 for norm in [true, false] {
20809 let m = moe(true, norm);
20810 let (idx, p, wsum) = moe_route(&[lo, hi, f32::NEG_INFINITY], &m, None);
20811 assert_eq!(idx, vec![1], "norm_topk {norm}");
20812 assert_eq!(p, vec![0.0, 1.0, 0.0]);
20813 assert_eq!(p[1] / wsum, 1.0);
20814 let (idx, _, _) = moe_route(&[hi, hi, lo], &m, None);
20817 assert_eq!(idx, vec![0]);
20818 let (idx, _, _) = moe_route(&[f32::NEG_INFINITY, lo, hi], &m, None);
20820 assert_eq!(idx, vec![2]);
20821 let (idx, p, wsum) = moe_route(&[lo, hi, hi], &m, Some(&[true, false, false]));
20822 assert_eq!(idx, vec![0]);
20823 assert_eq!(p[0] / wsum, 1.0);
20824 let (idx, p, wsum) = moe_route(&[f32::NEG_INFINITY; 3], &m, None);
20827 assert_eq!(idx, vec![0]);
20828 assert!(p.iter().all(|v| v.is_finite()) && wsum.is_finite() && wsum > 0.0);
20829 }
20830 }
20831}