1use crate::attention::{self, QwenAttnCfg};
10use crate::inference;
11
12#[path = "mimo_mtp.rs"]
15pub mod mimo_mtp;
16use crate::kv_cache::KvCache;
17use crate::linear_core::{
18 GdnCfg, GdnWeights, ShortConvCfg, ShortConvWeights, VmfPhaseCfg, VmfPhaseWeights, gdn_forward,
19 gdn_pair, short_conv_forward, short_conv_forward_batch, short_conv_pair, vmf_phase_forward,
20 vmf_phase_pair,
21};
22use crate::pool::Pool;
23use crate::qtensor::QTensor;
24use crate::sampler::{self, SamplerConfig, SamplerScratch, SplitMix64};
25use crate::tokenizer::Tokenizer;
26use cortiq_core::mask::TaskMask;
27use cortiq_core::types::NormStyle;
28
29pub static GLOBAL_USE_GPU: std::sync::atomic::AtomicBool =
30 std::sync::atomic::AtomicBool::new(false);
31
32struct ForwardScratch {
36 n1: Vec<f32>,
37 n2: Vec<f32>,
38 p1: Vec<f32>,
39 p2: Vec<f32>,
40}
41
42impl ForwardScratch {
43 fn new(hidden: usize) -> Self {
44 Self {
45 n1: vec![0.0; hidden],
46 n2: vec![0.0; hidden],
47 p1: vec![0.0; hidden],
48 p2: vec![0.0; hidden],
49 }
50 }
51}
52
53pub struct Pipeline {
55 gpu_plan: Option<std::sync::Arc<Vec<(usize, usize, usize)>>>,
60 pub tokenizer: std::sync::Arc<Tokenizer>,
63 pub kv_cache: KvCache,
64 pub sampler_config: SamplerConfig,
65 pub weights: PipelineWeights,
66 pub hidden_size: usize,
67 pub intermediate_size: usize,
68 pub num_heads: usize,
69 pub num_kv_heads: usize,
70 pub head_dim: usize,
71 pub num_layers: usize,
73 pub physical_layers: usize,
75 pub loop_final_norm: bool,
77 pub vocab_size: usize,
78 pub rms_eps: f64,
79 pub rope_base: f32,
80 pub norm_style: NormStyle,
81 pub rotary_dim: usize,
83 pub attention_heads_per_layer: Option<Vec<usize>>,
85 pub kv_heads_per_layer: Option<Vec<usize>>,
90 pub v_head_dim: Option<usize>,
95 pub layer_dump: Option<std::path::PathBuf>,
108 graph_declines: std::cell::RefCell<Vec<(&'static str, &'static str)>>,
111 pub(crate) mimo_moe: crate::mimo_moe::Slot,
114 pub vmf_cfg: Option<VmfPhaseCfg>,
116 pub gdn_cfg: Option<GdnCfg>,
118 pub logit_multiplier: Option<f32>,
120 pub cancel: std::sync::Arc<std::sync::atomic::AtomicBool>,
125 graph_failed: std::sync::atomic::AtomicBool,
130 pub kv_history: Vec<u32>,
135 pub kv_history_device: bool,
138 pub kda_cfg: Option<crate::linear_core::KdaCfg>,
140 pub g3n: Option<Box<(crate::g3n::G3nGlobals, Vec<crate::g3n::G3nLayer>)>>,
143 pub dsv4: Option<
147 Box<(
148 crate::dsv4::Dsv4Globals,
149 Vec<crate::dsv4::Dsv4Layer>,
150 crate::dsv4::Dsv4Cfg,
151 crate::dsv4::Dsv4State,
152 )>,
153 >,
154 pub dsv41: Option<
158 Box<(
159 crate::dsv41::Dsv41Globals,
160 Vec<crate::dsv41::Dsv41Layer>,
161 crate::dsv41::Dsv41Cfg,
162 crate::dsv41::Dsv41State,
163 )>,
164 >,
165 pub dsv41_vision: Option<crate::dsv41_vision::VisionModel>,
167 dsv41_prefill: Option<(Vec<Option<Vec<f32>>>, Vec<bool>)>,
169 pub qwen4_exp: Option<
172 Box<(
173 crate::qwen4_exp::Globals,
174 Vec<crate::qwen4_exp::Layer>,
175 crate::qwen4_exp::Cfg,
176 crate::qwen4_exp::State,
177 )>,
178 >,
179 pub dsv4_mtp: Vec<crate::dsv4::Dsv4Mtp>,
183 pub dspark: Option<crate::dsv4::DsparkState>,
185 pub dspark_pending: Vec<(usize, Vec<u32>, bool, usize)>,
188 pub dspark_hist: Vec<usize>,
190 pub dspark_real: Vec<u32>,
194 pub dspark_trunk_picks: Vec<Vec<(usize, Vec<usize>)>>,
198 pub dspark_exp: Vec<(usize, usize, usize, usize)>,
200 pub dspark_draft_ns: u128,
204 pub short_conv_cfg: Option<ShortConvCfg>,
207 pub mtp: Option<MtpModule>,
209 pub mimo_mtp: Option<mimo_mtp::MimoMtp>,
213 verify_exact_moe: bool,
216 pub speculative: bool,
218 pub ignore_eos: bool,
224 pub draft_full_streak: u32,
230 pub spec_k_adapt: Option<usize>,
238 pub spec_acc_ewma: f32,
240 rng: SplitMix64,
241 sampler_scratch: SamplerScratch,
242 spec_forced: Option<u32>,
248 spec_q: Vec<Vec<f32>>,
249 spec_p: Vec<f32>,
250 spec_res: Vec<f32>,
251 spec_qs: Vec<sampler::Sparse>,
253 spec_ps: sampler::Sparse,
254 spec_ress: sampler::Sparse,
255 mtp_graph_mode: Option<bool>,
262 #[cfg(target_os = "macos")]
265 metal_verify: Option<MetalVerifyPending>,
266 pub(crate) inv_freq: std::sync::Arc<Vec<f32>>,
270 ws: ForwardScratch,
274 pool: Option<std::sync::Arc<Pool>>,
276 pub(crate) model: Option<std::sync::Arc<cortiq_core::CmfModel>>,
280 pub(crate) dyn_force_f32: bool,
282 pub(crate) dyn_skill_layers: Vec<Option<Vec<usize>>>,
287 pub(crate) dyn_active: Option<usize>,
293 pub(crate) dyn_blend_loaded: bool,
297 pub(crate) dyn_phi_layer: Option<usize>,
300 dyn_phi_ema: Vec<f32>,
302 dyn_phi_seen: usize,
303 pub dyn_router: Option<crate::swarm::DynRouter>,
306 o1_cfg: Option<crate::nystrom::O1Cfg>,
309 o1_epoch: u64,
312 o1_flags: Vec<bool>,
314 trace: bool,
317 calib_temp: f32,
320 #[cfg_attr(not(target_os = "macos"), allow(dead_code))]
322 graph_kv_id: u64,
323 #[cfg_attr(not(target_os = "macos"), allow(dead_code))]
326 graph_want_logits: bool,
327 #[cfg_attr(not(target_os = "macos"), allow(dead_code))]
331 graph_head_required: bool,
332 graph_logits: Option<Vec<f32>>,
335 embryo_graph: Option<std::sync::Arc<crate::gpu::EmbryoGraphModel>>,
338 graph_refused: std::sync::atomic::AtomicBool,
344 pub embed_multiplier: f32,
346 pub attn_scale: f32,
349 pub swa: Option<(usize, usize)>,
352 pub sliding_layers: Option<Vec<bool>>,
355 pub anchor_core: Option<cortiq_core::AnchorCoreConfig>,
361 bounded_rope: Option<std::sync::Arc<crate::bounded::BoundedRope>>,
364 pub kv_prefix: KvPrefix,
368 pub last_prefill_tokens: usize,
371 pub inv_freq_local: Option<std::sync::Arc<Vec<f32>>>,
374 pub rotary_dim_local: Option<usize>,
375 pub rope_scale: f32,
376 pub rope_scale_local: f32,
377 pub global_attn: Option<(usize, usize)>,
380 pub inv_freq_global: Option<std::sync::Arc<Vec<f32>>>,
383 pub attn_v_norm: bool,
385 pub qk_norm_after_rope: bool,
387 pub proj_gate_sigmoid: bool,
390 pub final_softcap: Option<f32>,
392 pub head_clusters: Option<std::sync::Arc<Vec<f32>>>,
396 pub attn_softcap: f32,
398 confidence_on: bool,
402 #[cfg(test)]
405 nll_test_fail_at: Option<usize>,
406 #[cfg(test)]
409 nll_test_force_serial: bool,
410}
411
412#[cfg(target_os = "macos")]
413impl Drop for Pipeline {
414 fn drop(&mut self) {
415 let _ = crate::gpu_metal::wait_replay();
417 crate::gpu::kv_mirror_drop(self.graph_kv_id);
418 }
419}
420
421#[cfg(not(target_os = "macos"))]
422impl Drop for Pipeline {
423 fn drop(&mut self) {
424 crate::gpu::graph_kv_reset(self.graph_kv_id);
428 }
429}
430
431pub struct PipelineWeights {
436 pub embed_tokens: QTensor,
438 pub layers: Vec<LayerWeights>,
440 pub lm_head: QTensor,
442 pub final_norm: Vec<f32>,
444}
445
446pub struct LayerWeights {
448 pub input_norm: Vec<f32>,
449 pub post_norm: Vec<f32>,
452 pub attn_out_norm: Option<Vec<f32>>,
455 pub layer_scale: Option<f32>,
457 pub ffn_out_norm: Option<Vec<f32>>,
460 pub ffn: FfnKind,
461 pub attn: AttnKind,
462}
463
464#[derive(Clone, Copy, PartialEq, Debug, Default)]
467pub enum Act {
468 #[default]
469 Silu,
470 GeluTanh,
471 Gelu,
473 Situ {
476 beta: f32,
477 linear_beta: f32,
478 },
479}
480
481impl Act {
482 pub fn from_arch(name: &str) -> Self {
483 if name == "gelu_tanh" {
484 Self::GeluTanh
485 } else if name == "gelu" {
486 Self::Gelu
487 } else {
488 Self::Silu
489 }
490 }
491
492 pub fn from_arch_full(arch: &cortiq_core::ModelArch) -> Self {
494 match arch.hidden_act.as_str() {
495 "situ" => Self::Situ {
496 beta: arch.activation_situ_beta.unwrap_or(1.0) as f32,
497 linear_beta: arch.activation_situ_linear_beta.unwrap_or(0.0) as f32,
498 },
499 other => Self::from_arch(other),
500 }
501 }
502
503 #[inline]
504 pub fn apply(self, x: f32) -> f32 {
505 match self {
506 Self::Silu => inference::silu(x),
507 Self::GeluTanh => inference::gelu_tanh(x),
508 Self::Gelu => inference::gelu_erf(x),
509 Self::Situ { beta, .. } => beta * (x / beta).tanh() * (1.0 / (1.0 + (-x).exp())),
510 }
511 }
512
513 #[inline]
516 pub fn combine(self, g: f32, u: f32) -> f32 {
517 match self {
518 Self::Situ { linear_beta, .. } if linear_beta > 0.0 => {
519 self.apply(g) * (linear_beta * (u / linear_beta).tanh())
520 }
521 _ => self.apply(g) * u,
522 }
523 }
524
525 pub fn graph_act(self) -> Option<crate::gpu::GraphAct> {
529 match self {
530 Self::Silu => Some(crate::gpu::GraphAct::Silu),
531 Self::Gelu => Some(crate::gpu::GraphAct::GeluErf),
532 Self::GeluTanh | Self::Situ { .. } => None,
533 }
534 }
535}
536
537pub struct DenseFfn {
539 pub gate_proj: QTensor,
540 pub up_proj: QTensor,
541 pub down_proj: QTensor,
542 pub act: Act,
544 pub down_t: Option<QTensor>,
550 pub segs: Vec<FfnSeg>,
557}
558
559pub struct FfnSeg {
564 pub gate: QTensor,
565 pub up: QTensor,
566 pub down: QTensor,
567 pub start: usize,
568 pub width: usize,
569}
570
571pub enum FfnKind {
574 Dense(DenseFfn),
575 Moe(MoeFfn),
579 DenseMoe(Box<DenseMoeFfn>),
586}
587
588pub struct DenseMoeFfn {
590 pub dense: DenseFfn,
591 pub moe: MoeFfn,
592 pub post_norm_1: Vec<f32>,
594 pub pre_norm_2: Vec<f32>,
597 pub post_norm_2: Vec<f32>,
599}
600
601pub struct MoeFfn {
602 pub router: QTensor,
604 pub experts: Vec<DenseFfn>,
605 pub top_k: usize,
606 pub norm_topk_prob: bool,
607 pub router_sigmoid: bool,
610 pub expert_bias: Option<Vec<f32>>,
614 pub routed_scaling: f32,
617 pub route_tau: Option<f32>,
623 pub shared: Option<(DenseFfn, Option<QTensor>)>,
626 pub stats: std::cell::RefCell<Vec<u64>>,
630 pub act_sq: std::cell::RefCell<Vec<f64>>,
637 pub act_rows: std::cell::RefCell<Vec<f32>>,
643 pub mask: Option<Vec<bool>>,
648 pub per_expert_scale: Option<Vec<f32>>,
651 pub router_input_norm: bool,
655 pub resonance: Option<Resonance>,
659 pub grown: Vec<GrownExpert>,
665}
666
667#[derive(Debug, Clone, PartialEq, Eq)]
669pub struct GrownExpert {
670 pub record: String,
672 pub record_index: usize,
674 pub layer: usize,
675 pub expert: usize,
679}
680
681pub struct Resonance {
683 pub mu: Vec<f32>,
685 pub u: Vec<f32>,
687 pub k: usize,
688 pub bias: Vec<f32>,
690 pub shell: Vec<f32>,
696}
697
698static GROWTH_SHELL: std::sync::atomic::AtomicU8 = std::sync::atomic::AtomicU8::new(0);
701
702pub fn growth_shell_enabled() -> bool {
707 use std::sync::atomic::Ordering;
708 match GROWTH_SHELL.load(Ordering::Relaxed) {
709 1 => true,
710 2 => false,
711 _ => {
712 let off = std::env::var("CMF_GROWTH_SHELL")
713 .map(|v| v.eq_ignore_ascii_case("off") || v == "0")
714 .unwrap_or(false);
715 GROWTH_SHELL.store(if off { 2 } else { 1 }, Ordering::Relaxed);
716 !off
717 }
718 }
719}
720
721pub fn set_growth_shell(on: Option<bool>) {
726 GROWTH_SHELL.store(
727 match on {
728 Some(true) => 1,
729 Some(false) => 2,
730 None => 0,
731 },
732 std::sync::atomic::Ordering::Relaxed,
733 );
734}
735
736impl Resonance {
737 pub fn has_shell(&self) -> bool {
739 self.shell.iter().any(|s| s.is_finite())
740 }
741
742 pub fn effective_shell(&self, ne: usize) -> Vec<f32> {
745 let mut out = vec![f32::INFINITY; ne];
746 if growth_shell_enabled() {
747 for (o, s) in out.iter_mut().zip(&self.shell) {
748 *o = *s;
749 }
750 }
751 out
752 }
753
754 pub fn scores(&self, x: &[f32], out: &mut [f32]) {
759 let h = x.len();
760 let ne = out.len();
761 let shell_on = growth_shell_enabled() && !self.shell.is_empty();
762 for e in 0..ne {
763 let mu = &self.mu[e * h..(e + 1) * h];
764 let mut d2 = 0.0f32;
765 for j in 0..h {
766 let d = x[j] - mu[j];
767 d2 += d * d;
768 }
769 let mut proj = 0.0f32;
770 for i in 0..self.k {
771 let u = &self.u[(e * self.k + i) * h..(e * self.k + i + 1) * h];
772 let mut p = 0.0f32;
773 for j in 0..h {
774 p += (x[j] - mu[j]) * u[j];
775 }
776 proj += p * p;
777 }
778 let err = d2 - proj;
779 out[e] = self.bias.get(e).copied().unwrap_or(0.0) - err;
780 if shell_on && err > self.shell.get(e).copied().unwrap_or(f32::INFINITY) {
781 out[e] = f32::NEG_INFINITY;
782 }
783 }
784 }
785}
786
787pub enum AttnKind {
790 Full {
792 wq: QTensor,
793 wk: QTensor,
794 wv: QTensor,
795 wo: QTensor,
796 q_norm: Option<Vec<f32>>,
797 k_norm: Option<Vec<f32>>,
798 output_gate: bool,
799 softplus_gate: Option<(QTensor, bool)>,
803 bias: Option<(Vec<f32>, Vec<f32>, Vec<f32>)>,
805 },
806 Linear(VmfPhaseWeights),
808 LinearGdn(GdnWeights),
810 ShortConv(ShortConvWeights),
813 Mla(Box<MlaWeights>),
821 Kda(Box<crate::linear_core::KdaWeights>),
825 Bounded(Box<crate::bounded::BoundedWeights>),
830}
831
832pub struct MlaWeights {
834 pub q_proj: QTensor,
838 pub q_a: Option<QTensor>,
841 pub q_a_norm: Option<Vec<f32>>,
842 pub kv_a: QTensor,
844 pub kv_a_norm: Vec<f32>,
846 pub kv_b: QTensor,
848 pub o_proj: QTensor,
850 pub nh: usize,
851 pub qk_rope: usize,
852 pub qk_nope: usize,
853 pub v_dim: usize,
854 pub lora: usize,
855 pub scale: f32,
857 pub nope: bool,
859}
860
861pub struct MtpModule {
866 pub enorm: Vec<f32>,
867 pub hnorm: Vec<f32>,
868 pub eh_proj: QTensor,
870 pub layer: LayerWeights,
871 pub final_norm: Vec<f32>,
872 pub kv: crate::kv_cache::LayerKvCache,
873}
874
875#[cfg(target_os = "macos")]
882enum MetalRowsItem<'a> {
883 Gdn {
884 run: Vec<crate::gpu_metal::GdnGpuLayer<'a>>,
885 first: usize,
886 },
887 Attn {
888 l: crate::gpu_metal::AttnGpuLayer<'a>,
889 li: usize,
890 q_norm: Option<&'a [f32]>,
891 k_norm: Option<&'a [f32]>,
892 output_gate: bool,
893 },
894}
895
896#[cfg(target_os = "macos")]
897struct MetalVerifyPending {
898 graph: crate::gpu_metal::VerifyGraph,
899 gdn_layers: Vec<usize>,
900 attn_layers: Vec<(usize, usize)>,
901}
902
903#[cfg(target_os = "macos")]
907struct MetalWarmPending {
908 graph: crate::gpu_metal::VerifyGraph,
909 cpu_stored: usize,
910 b: usize,
911}
912
913#[cfg(target_os = "macos")]
914enum MetalRowsRun {
915 Declined,
917 Failed,
920 Completed(MetalVerifyPending),
921}
922
923#[cfg(target_os = "macos")]
924enum MetalPrefillOutcome {
925 Declined,
926 Failed,
927 Completed(Vec<f32>),
928}
929
930#[cfg(target_os = "macos")]
931enum MetalBatchNllOutcome {
932 Declined,
933 Failed(String),
934 Completed(f64, usize),
935}
936
937#[derive(Clone, Copy)]
941enum SpecTrial {
942 Spec {
943 t0: std::time::Instant,
944 gen0: usize,
945 rounds: usize,
946 },
947 Plain {
948 t0: std::time::Instant,
949 gen0: usize,
950 },
951 Decided {
952 spec: bool,
953 recheck_at: usize,
954 },
955}
956
957pub(crate) fn spec_time_level() -> u8 {
961 static L: std::sync::OnceLock<u8> = std::sync::OnceLock::new();
962 *L.get_or_init(|| match std::env::var("CMF_GRAPH_SPEC_TIME") {
963 Ok(v) => v.trim().parse::<u8>().map(|n| n.max(1)).unwrap_or(1),
964 Err(_) => 0,
965 })
966}
967
968struct SpecStampLog {
974 t_last: std::time::Instant,
975 items: Vec<(&'static str, f32)>,
976}
977
978static SPEC_STAMPS: std::sync::Mutex<Option<SpecStampLog>> = std::sync::Mutex::new(None);
979
980pub(crate) fn spec_stamp(name: &'static str) {
981 if spec_time_level() == 0 {
982 return;
983 }
984 if let Ok(mut g) = SPEC_STAMPS.lock() {
985 if let Some(log) = g.as_mut() {
986 let now = std::time::Instant::now();
987 log.items
988 .push((name, (now - log.t_last).as_secs_f32() * 1e3));
989 log.t_last = now;
990 }
991 }
992}
993
994fn spec_stamps_begin() {
995 if spec_time_level() == 0 {
996 return;
997 }
998 if let Ok(mut g) = SPEC_STAMPS.lock() {
999 *g = Some(SpecStampLog {
1000 t_last: std::time::Instant::now(),
1001 items: Vec::with_capacity(64),
1002 });
1003 }
1004}
1005
1006fn spec_stamps_take() -> Vec<(&'static str, f32)> {
1007 SPEC_STAMPS
1008 .lock()
1009 .ok()
1010 .and_then(|mut g| g.take())
1011 .map(|l| l.items)
1012 .unwrap_or_default()
1013}
1014
1015fn spec_stamps_format(items: &[(&'static str, f32)]) -> String {
1018 let mut agg: Vec<(&'static str, f32, u32)> = Vec::with_capacity(items.len());
1019 for &(n, ms) in items {
1020 match agg.iter_mut().find(|e| e.0 == n) {
1021 Some(e) => {
1022 e.1 += ms;
1023 e.2 += 1;
1024 }
1025 None => agg.push((n, ms, 1)),
1026 }
1027 }
1028 let mut s = String::with_capacity(agg.len() * 16);
1029 for (n, ms, k) in agg {
1030 if k > 1 {
1031 s.push_str(&format!("{n} {ms:.1}/{k} "));
1032 } else {
1033 s.push_str(&format!("{n} {ms:.1} "));
1034 }
1035 }
1036 s
1037}
1038
1039#[derive(Default, Clone, Copy)]
1061struct SpecMon {
1062 round_ms: f64,
1063 tokens: f64,
1064 plain_ms: f64,
1065 n: u32,
1066 fails: u32,
1067 metal: bool,
1068}
1069
1070const SPEC_PROXY_TOKENS: f64 = 3.5;
1073const SPEC_PLAIN_MIN_MS: f64 = 200.0;
1076
1077impl SpecMon {
1078 fn round(&mut self, dt_ms: f64, produced: usize) {
1079 self.n += 1;
1080 if self.n == 1 {
1081 return; }
1083 let a = if self.n == 2 { 1.0 } else { 0.3 };
1084 self.round_ms += a * (dt_ms - self.round_ms);
1085 self.tokens += a * (produced as f64 - self.tokens);
1086 }
1087 fn pays(&self) -> bool {
1088 if self.plain_ms > 0.0 {
1089 self.tokens * self.plain_ms > self.round_ms * 1.03
1090 } else {
1091 self.metal && self.tokens >= SPEC_PROXY_TOKENS
1092 }
1093 }
1094 fn plain_done(&self, t0: std::time::Instant, gen0: usize, generated: usize) -> bool {
1096 let n = generated.saturating_sub(gen0);
1097 if n >= 8 {
1098 return true;
1099 }
1100 self.metal && n >= 2 && t0.elapsed().as_secs_f64() * 1e3 >= SPEC_PLAIN_MIN_MS
1101 }
1102}
1103
1104pub const KV_PREFIX_TAIL: usize = 128;
1107
1108#[derive(Debug, Clone, Default)]
1114pub struct KvPrefix {
1115 len: usize,
1116 hash: u64,
1117 tail: Vec<u32>,
1118 device: bool,
1123}
1124
1125impl KvPrefix {
1126 #[inline]
1127 fn fold(mut h: u64, ids: &[u32]) -> u64 {
1128 for &id in ids {
1129 h ^= id as u64;
1130 h = h.wrapping_mul(0x100000001b3);
1131 h ^= h >> 29;
1132 }
1133 h
1134 }
1135
1136 pub fn clear(&mut self) {
1137 self.len = 0;
1138 self.hash = 0xcbf29ce484222325;
1139 self.tail.clear();
1140 self.device = false;
1141 }
1142
1143 pub fn on_device(&self) -> bool {
1145 self.device
1146 }
1147
1148 pub fn set_on_device(&mut self, device: bool) {
1150 self.device = device;
1151 }
1152
1153 pub fn len(&self) -> usize {
1155 self.len
1156 }
1157
1158 pub fn is_empty(&self) -> bool {
1159 self.len == 0
1160 }
1161
1162 pub fn tail_len(&self) -> usize {
1164 self.tail.len()
1165 }
1166
1167 pub fn set(&mut self, ids: &[u32]) {
1169 self.clear();
1170 self.extend(ids);
1171 }
1172
1173 pub fn extend(&mut self, more: &[u32]) {
1175 if self.len == 0 && self.hash == 0 {
1176 self.hash = 0xcbf29ce484222325;
1177 }
1178 self.hash = Self::fold(self.hash, more);
1179 self.len += more.len();
1180 if more.len() >= KV_PREFIX_TAIL {
1181 self.tail.clear();
1182 self.tail.extend_from_slice(&more[more.len() - KV_PREFIX_TAIL..]);
1183 } else {
1184 let drop = (self.tail.len() + more.len()).saturating_sub(KV_PREFIX_TAIL);
1185 self.tail.drain(..drop);
1186 self.tail.extend_from_slice(more);
1187 }
1188 }
1189
1190 pub fn extension(&self, ids: &[u32]) -> usize {
1194 if self.len == 0 || ids.len() <= self.len {
1195 return 0;
1196 }
1197 let t = self.tail.len();
1198 if ids[self.len - t..self.len] != self.tail[..] {
1199 return 0;
1200 }
1201 if Self::fold(0xcbf29ce484222325, &ids[..self.len]) != self.hash {
1202 return 0;
1203 }
1204 self.len
1205 }
1206}
1207
1208pub struct GenerateResult {
1210 pub text: String,
1211 pub token_ids: Vec<u32>,
1212 pub prompt_tokens: usize,
1213 pub tokens_generated: usize,
1214 pub finish_reason: String,
1215 pub mtp_drafted: usize,
1217 pub mtp_accepted: usize,
1218 pub token_confidence: Vec<f32>,
1223 pub traces: Vec<TokenTrace>,
1226}
1227
1228#[derive(Clone, Debug)]
1233pub struct TokenTrace {
1234 pub t: usize,
1236 pub token_id: u32,
1238 pub confidence: f32,
1240 pub active_skill: Option<String>,
1242 pub recon: Option<f32>,
1246 pub switched: bool,
1249}
1250
1251#[cfg_attr(not(test), allow(dead_code))]
1256fn top1_prob_t(logits: &[f32], id: u32, temp: f32) -> f32 {
1257 let t = if temp > 1e-3 { temp } else { 1.0 };
1258 let max = logits.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
1259 let sum: f32 = logits.iter().map(|&v| ((v - max) / t).exp()).sum();
1260 if sum > 0.0 {
1261 (((logits[id as usize] - max) / t).exp()) / sum
1262 } else {
1263 0.0
1264 }
1265}
1266
1267fn prefill_batched() -> bool {
1270 std::env::var("CMF_PREFILL")
1271 .map(|v| v != "seq")
1272 .unwrap_or(true)
1273}
1274
1275#[inline]
1279fn nll_graph_policy(
1280 unmasked: bool,
1281 prefer_graph: bool,
1282 native_metal: bool,
1283) -> (bool, bool) {
1284 let graph_quality = unmasked && prefer_graph;
1285 let fused_head_quality = graph_quality && native_metal;
1286 (graph_quality, fused_head_quality)
1287}
1288
1289#[derive(Clone, Copy)]
1293enum PrefillIn<'a> {
1294 Ids(&'a [u32]),
1295 Hidden(&'a [f32]),
1296}
1297
1298impl Pipeline {
1305 fn can_prefill_batched(&self) -> bool {
1306 #[cfg(test)]
1307 let force_serial = self.nll_test_force_serial;
1308 #[cfg(not(test))]
1309 let force_serial = false;
1310 prefill_batched() && !force_serial && !self.weights.layers.is_empty()
1311 }
1312
1313 fn automatic_gpu_prefix(&self) -> Option<usize> {
1316 let (model, _, _, _) = self.weights.embed_tokens.graph_weight()?;
1317 crate::gpu::automatic_layer_prefix(&model, self.num_layers, self.physical_layers)
1318 }
1319
1320 pub fn prefill_chunk(&self) -> usize {
1324 let env = env_prefill_chunk();
1325 if env.is_some() || ChunkHost::here() != ChunkHost::Other {
1326 return prefill_chunk_rule(env, ChunkHost::here(), false);
1327 }
1328 prefill_chunk_rule(None, ChunkHost::Other, self.chunk_stack_facts().dense_on_discrete())
1329 }
1330
1331 fn chunk_stack_facts(&self) -> ChunkStackFacts {
1332 let plain_dense = !self.weights.layers.is_empty()
1333 && self.g3n.is_none()
1334 && self.dsv4.is_none()
1335 && self.dsv41.is_none()
1336 && self.qwen4_exp.is_none()
1337 && self.weights.layers.iter().all(|lw| {
1338 matches!(lw.attn, AttnKind::Full { .. }) && matches!(lw.ffn, FfnKind::Dense(_))
1339 });
1340 let gpu_on = crate::gpu::enabled();
1341 ChunkStackFacts {
1342 plain_dense,
1343 discrete: gpu_on && crate::gpu::discrete(),
1344 gpu_on,
1345 capacity_split: std::env::var_os("CMF_GPU_LAYERS").is_some()
1348 || (plain_dense && gpu_on && self.automatic_gpu_prefix().is_some()),
1349 multi_gpu: self.gpu_plan.is_some(),
1350 o1: self.o1_active(),
1351 }
1352 }
1353}
1354
1355pub fn prefill_chunk() -> usize {
1364 prefill_chunk_rule(env_prefill_chunk(), ChunkHost::here(), false)
1365}
1366
1367fn env_prefill_chunk() -> Option<usize> {
1368 std::env::var("CMF_PREFILL_CHUNK")
1369 .ok()
1370 .and_then(|v| v.parse::<usize>().ok())
1371}
1372
1373#[derive(Clone, Copy, Debug, PartialEq, Eq)]
1375enum ChunkHost {
1376 Macos,
1377 Aarch64,
1379 Other,
1381}
1382
1383impl ChunkHost {
1384 fn here() -> Self {
1385 if cfg!(target_os = "macos") {
1386 ChunkHost::Macos
1387 } else if cfg!(target_arch = "aarch64") {
1388 ChunkHost::Aarch64
1389 } else {
1390 ChunkHost::Other
1391 }
1392 }
1393}
1394
1395const DISCRETE_DENSE_PREFILL_CHUNK: usize = 512;
1402
1403fn prefill_chunk_rule(env: Option<usize>, host: ChunkHost, dense_on_discrete: bool) -> usize {
1409 if let Some(n) = env {
1410 return n.max(1);
1411 }
1412 match host {
1413 ChunkHost::Macos => 512,
1414 ChunkHost::Aarch64 => 256,
1417 ChunkHost::Other if dense_on_discrete => DISCRETE_DENSE_PREFILL_CHUNK,
1418 ChunkHost::Other => 48,
1419 }
1420}
1421
1422#[derive(Clone, Copy, Debug, Default)]
1424struct ChunkStackFacts {
1425 plain_dense: bool,
1428 discrete: bool,
1430 gpu_on: bool,
1432 capacity_split: bool,
1434 multi_gpu: bool,
1436 o1: bool,
1438}
1439
1440impl ChunkStackFacts {
1441 fn dense_on_discrete(self) -> bool {
1442 self.plain_dense
1443 && self.discrete
1444 && self.gpu_on
1445 && !self.capacity_split
1446 && !self.multi_gpu
1447 && !self.o1
1448 }
1449}
1450
1451#[inline]
1457fn mtp_prefill_pair_count(start: usize, end: usize, input_len: usize) -> usize {
1458 if end <= start || start >= input_len {
1459 return 0;
1460 }
1461 let rows = (end.min(input_len) - start).min(input_len - start);
1462 if end < input_len {
1463 rows
1464 } else {
1465 rows.saturating_sub(1)
1466 }
1467}
1468
1469pub type TokenCallback = Box<dyn FnMut(&str) -> bool + Send>;
1471
1472#[derive(Debug, Clone, Copy, PartialEq, Eq)]
1474pub(crate) struct ReuseLayer {
1475 pub full: bool,
1478 pub host_rows: usize,
1480 pub device_rows: Option<usize>,
1482 pub device_state: bool,
1484}
1485
1486#[derive(Debug, Clone, PartialEq, Eq)]
1488pub(crate) enum ReusePlan {
1489 Ready,
1491 Pull(Vec<(usize, usize, usize)>),
1494 Fresh,
1496}
1497
1498pub(crate) fn kv_reuse_plan(reuse_from: usize, layers: &[ReuseLayer]) -> ReusePlan {
1508 let mut pulls = Vec::new();
1509 for (li, l) in layers.iter().enumerate() {
1510 if !l.full {
1511 if l.device_state {
1512 return ReusePlan::Fresh;
1513 }
1514 continue;
1515 }
1516 if l.host_rows == reuse_from {
1517 continue;
1518 }
1519 if l.host_rows < reuse_from && l.device_rows.is_some_and(|d| d >= reuse_from) {
1520 pulls.push((li, l.host_rows, reuse_from));
1521 continue;
1522 }
1523 return ReusePlan::Fresh;
1524 }
1525 if pulls.is_empty() {
1526 ReusePlan::Ready
1527 } else {
1528 ReusePlan::Pull(pulls)
1529 }
1530}
1531
1532impl Pipeline {
1533 fn clear_sequence_state(&mut self) {
1541 #[cfg(target_os = "macos")]
1544 let _ = crate::gpu_metal::wait_replay();
1545 self.kv_cache.clear();
1546 self.clear_history();
1549 self.graph_logits = None;
1550 if let Some(b) = &mut self.dsv41 {
1551 b.3.clear();
1552 }
1553 crate::gpu::graph_kv_reset(self.graph_kv_id);
1554 crate::gpu::graph_kv_reset(self.mtp_kv_id());
1559 }
1560
1561 fn prepare_kv_reuse(&mut self, reuse_from: usize) -> bool {
1568 if self.graph_prefill_preferred() {
1569 return true;
1570 }
1571 let kv_id = self.graph_kv_id;
1572 let layers: Vec<ReuseLayer> = (0..self.num_layers)
1573 .map(|li| {
1574 let full = matches!(
1575 self.weights.layers[self.phys_layer(li)].attn,
1576 AttnKind::Full { .. }
1577 );
1578 ReuseLayer {
1579 full,
1580 host_rows: self.kv_cache.layers[li].seq_len,
1581 device_rows: crate::gpu::graph_kv_stored(kv_id, li),
1582 device_state: crate::gpu::graph_state_resident(kv_id, li),
1583 }
1584 })
1585 .collect();
1586 if layers
1590 .iter()
1591 .all(|l| l.device_rows.is_none() && !l.device_state)
1592 {
1593 return true;
1594 }
1595 let plan = kv_reuse_plan(reuse_from, &layers);
1596 let (what, rows, n) = match &plan {
1597 ReusePlan::Ready => ("host ready", 0, 0),
1598 ReusePlan::Fresh => ("fresh", 0, 0),
1599 ReusePlan::Pull(p) => (
1600 "pull",
1601 p.iter().map(|&(_, a, b)| b - a).max().unwrap_or(0),
1602 p.len(),
1603 ),
1604 };
1605 let t0 = std::time::Instant::now();
1606 let ok = self.apply_kv_reuse_plan(reuse_from, plan, &layers);
1607 if std::env::var("CMF_PREFILL_PROF").is_ok() {
1608 eprintln!(
1609 "kv-reuse: {what}{}: {rows} device row(s) × {n} layer(s) to the host in {:.2} ms",
1610 if ok { "" } else { " (failed → fresh)" },
1611 t0.elapsed().as_secs_f64() * 1e3
1612 );
1613 }
1614 ok
1615 }
1616
1617 fn apply_kv_reuse_plan(
1618 &mut self,
1619 reuse_from: usize,
1620 plan: ReusePlan,
1621 layers: &[ReuseLayer],
1622 ) -> bool {
1623 let kv_id = self.graph_kv_id;
1624 match plan {
1625 ReusePlan::Fresh => return false,
1626 ReusePlan::Ready => {}
1627 ReusePlan::Pull(pulls) => {
1628 let (nkv, hd) = {
1633 let c = &self.kv_cache.layers[pulls[0].0];
1634 (c.num_kv_heads, c.head_dim)
1635 };
1636 let uniform = pulls.iter().all(|&(li, _, _)| {
1637 let c = &self.kv_cache.layers[li];
1638 (c.num_kv_heads, c.head_dim) == (nkv, hd)
1639 });
1640 let batched = if uniform {
1641 crate::gpu::graph_kv_read_rows(kv_id, &pulls, nkv, hd)
1642 } else {
1643 None
1644 };
1645 let rows: Vec<(Vec<f32>, Vec<f32>)> = match batched {
1646 Some(rows) => rows,
1647 None => {
1648 let mut rows = Vec::with_capacity(pulls.len());
1649 for &(li, from, to) in &pulls {
1650 let (lnkv, lhd) = {
1651 let c = &self.kv_cache.layers[li];
1652 (c.num_kv_heads, c.head_dim)
1653 };
1654 let Some((k, v, first_valid)) =
1655 crate::gpu::graph_kv_pull_host(kv_id, li, from, to, lnkv, lhd)
1656 else {
1657 return false;
1658 };
1659 let need_from = match self.layer_window(li) {
1663 Some(w) => from.max((to + 1).saturating_sub(w)),
1664 None => from,
1665 };
1666 if first_valid > need_from {
1667 return false;
1668 }
1669 rows.push((k, v));
1670 }
1671 rows
1672 }
1673 };
1674 for ((li, from, to), (k, v)) in pulls.into_iter().zip(rows) {
1675 let cache = &mut self.kv_cache.layers[li];
1676 let row = cache.num_kv_heads * cache.head_dim;
1677 for p in 0..to - from {
1678 cache.append(&k[p * row..(p + 1) * row], &v[p * row..(p + 1) * row], &[]);
1679 }
1680 if cache.seq_len != to {
1681 return false;
1682 }
1683 }
1684 }
1685 }
1686 for (li, l) in layers.iter().enumerate() {
1690 if l.full
1691 && l.device_rows.is_some_and(|d| d > reuse_from)
1692 && !crate::gpu::graph_kv_set_stored(kv_id, li, reuse_from)
1693 {
1694 return false;
1695 }
1696 }
1697 true
1698 }
1699
1700 fn finish_generation(
1706 &mut self,
1707 mtp: &mut Option<MtpModule>,
1708 router: &mut Option<crate::swarm::DynRouter>,
1709 clear_sequence: bool,
1710 ) {
1711 if router.is_some() {
1715 let _ = self.set_active_skill(None);
1716 }
1717 #[cfg(target_os = "macos")]
1724 let clear_sequence = clear_sequence || !crate::gpu_metal::wait_replay();
1725 if clear_sequence {
1726 self.clear_sequence_state();
1727 if let Some(m) = mtp.as_mut() {
1728 m.kv.clear();
1734 }
1735 if let Some(m) = self.mtp.as_mut() {
1736 m.kv.clear();
1740 }
1741 }
1742 self.graph_want_logits = false;
1743 self.graph_head_required = false;
1744 self.graph_logits = None;
1745 self.graph_failed
1746 .store(false, std::sync::atomic::Ordering::Relaxed);
1747 self.cancel
1748 .store(false, std::sync::atomic::Ordering::Relaxed);
1749 self.dyn_router = router.take().or(self.dyn_router.take());
1750 self.mtp = mtp.take().or(self.mtp.take());
1751 self.mtp_graph_mode = None;
1752 self.spec_forced = None;
1753 }
1754
1755 fn check_forward_graph(&mut self, phase: &str, pos: usize) -> Result<(), String> {
1759 if self
1760 .graph_failed
1761 .swap(false, std::sync::atomic::Ordering::Relaxed)
1762 {
1763 self.cancel
1764 .store(false, std::sync::atomic::Ordering::Relaxed);
1765 self.clear_sequence_state();
1766 self.graph_logits = None;
1767 self.graph_want_logits = false;
1768 self.graph_head_required = false;
1769 return Err(format!("GPU graph failed during {phase} at position {pos}"));
1770 }
1771 Ok(())
1772 }
1773
1774 #[cfg(target_os = "macos")]
1775 fn fail_metal_graph(&mut self, reason: &str) {
1776 crate::pipeline::METAL_GRAPH_ERRORS
1777 .fetch_add(1, std::sync::atomic::Ordering::Relaxed);
1778 self.clear_sequence_state();
1779 self.graph_logits = None;
1780 self.graph_failed
1781 .store(true, std::sync::atomic::Ordering::Relaxed);
1782 self.cancel
1783 .store(true, std::sync::atomic::Ordering::Relaxed);
1784 tracing::error!("native Metal TokenGraph failed closed: {reason}");
1785 }
1786
1787 fn nll_begin(&mut self) -> Result<(), String> {
1792 if self
1793 .graph_failed
1794 .swap(false, std::sync::atomic::Ordering::Relaxed)
1795 {
1796 self.cancel
1797 .store(false, std::sync::atomic::Ordering::Relaxed);
1798 self.clear_sequence_state();
1799 self.graph_logits = None;
1800 self.graph_want_logits = false;
1801 self.graph_head_required = false;
1802 return Err("GPU graph failed before NLL scoring".to_string());
1803 }
1804 self.clear_sequence_state();
1805 self.graph_logits = None;
1806 self.graph_want_logits = false;
1807 self.graph_head_required = false;
1808 Ok(())
1809 }
1810
1811 fn nll_end(&mut self) {
1815 self.clear_sequence_state();
1816 self.graph_logits = None;
1817 self.graph_want_logits = false;
1818 self.graph_head_required = false;
1819 self.graph_failed
1820 .store(false, std::sync::atomic::Ordering::Relaxed);
1821 }
1822
1823 fn nll_check_graph(&mut self, phase: &str, pos: usize) -> Result<(), String> {
1826 #[cfg(test)]
1827 if self.nll_test_fail_at == Some(pos) {
1828 self.nll_test_fail_at = None;
1829 self.graph_failed
1830 .store(true, std::sync::atomic::Ordering::Relaxed);
1831 self.cancel
1832 .store(true, std::sync::atomic::Ordering::Relaxed);
1833 }
1834 if self
1835 .graph_failed
1836 .swap(false, std::sync::atomic::Ordering::Relaxed)
1837 {
1838 self.cancel
1839 .store(false, std::sync::atomic::Ordering::Relaxed);
1840 self.clear_sequence_state();
1841 self.graph_logits = None;
1842 self.graph_want_logits = false;
1843 return Err(format!(
1844 "GPU graph failed during NLL {phase} at position {pos}"
1845 ));
1846 }
1847 Ok(())
1848 }
1849
1850 #[inline]
1854 pub fn phys_layer(&self, virtual_idx: usize) -> usize {
1855 virtual_idx % self.physical_layers
1856 }
1857
1858 #[inline]
1861 pub fn is_loop_end(&self, virtual_idx: usize) -> bool {
1862 self.loop_final_norm && (virtual_idx + 1) % self.physical_layers == 0
1863 }
1864
1865 #[allow(clippy::too_many_arguments)]
1867
1868 #[cfg(target_os = "macos")]
1887 fn graph_prefill_preferred(&self) -> bool {
1888 let graph_force = crate::gpu::q1_force() || crate::gpu::q2tp_gpu_opt_in();
1889 if !crate::gpu::enabled_here()
1890 || !graph_force
1891 || std::env::var("CMF_GPU_BLOCK")
1892 .map(|v| v == "0")
1893 .unwrap_or(false)
1894 || std::env::var("CMF_PREFILL_GRAPH").as_deref() == Ok("0")
1897 {
1898 return false;
1899 }
1900 self.weights
1901 .layers
1902 .iter()
1903 .any(|lw| {
1904 matches!(&lw.attn, AttnKind::LinearGdn(w) if w.in_proj_qkv.metal_graph_parts().is_some())
1905 })
1906 }
1907
1908 #[cfg(not(target_os = "macos"))]
1915 fn batch_prefix_prefill(&self) -> bool {
1916 let forced = match std::env::var("CMF_BATCH_PREFIX").as_deref() {
1917 Ok("0") => return false,
1918 Ok("1") => true,
1919 _ => false,
1920 };
1921 crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Prefill)
1922 && crate::gpu::enabled_here()
1923 && !self.graph_refused()
1924 && (forced || self.graph_attn_decline_reason().is_some())
1925 && self.wgpu_graph_attn_decline().is_none()
1926 && self.attn_softcap == 0.0
1927 && self
1928 .weights
1929 .layers
1930 .iter()
1931 .any(|lw| matches!(&lw.ffn, FfnKind::Moe(_)))
1932 && self.automatic_gpu_prefix().is_some()
1933 }
1934
1935 #[cfg(not(target_os = "macos"))]
1936 fn graph_prefill_preferred(&self) -> bool {
1937 let graph_on = crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Prefill);
1945 if !graph_on || !crate::gpu::enabled_here() {
1946 return false;
1947 }
1948 if self.embryo_resident_eligible() {
1952 return crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Decode);
1957 }
1958 if self.o1_active() {
1972 return false;
1973 }
1974 if self.wgpu_graph_attn_decline().is_some() {
1978 return false;
1979 }
1980 if self
1981 .weights
1982 .layers
1983 .iter()
1984 .any(|lw| matches!(&lw.attn, AttnKind::LinearGdn(_)))
1985 {
1986 return true;
1987 }
1988 self.weights
1999 .layers
2000 .iter()
2001 .any(|lw| matches!(&lw.ffn, FfnKind::Moe(_)))
2002 && self.automatic_gpu_prefix().is_none()
2003 }
2004
2005 #[cfg(target_os = "macos")]
2006 fn q1_graph_gpu(
2007 &mut self,
2008 start: usize,
2009 upto: Option<usize>,
2010 position: usize,
2011 h: &mut [f32],
2012 ) -> usize {
2013 let _mt0 = std::time::Instant::now(); use crate::gpu::{AttnGpuLayer, GdnGpuCfg, GdnGpuLayer, GraphDims, MetalFfn, TokenGraph};
2015 let graph_force = crate::gpu::q1_force() || crate::gpu::q2tp_gpu_opt_in();
2016 if self.attn_softcap > 0.0 || !crate::gpu::enabled_here()
2018 || !graph_force
2019 || std::env::var("CMF_GPU_BLOCK")
2020 .map(|v| v == "0")
2021 .unwrap_or(false)
2022 {
2023 if std::env::var("CMF_GRAPH_DBG").is_ok() {
2024 eprintln!(
2025 "block-graph: front gate (softcap={} enabled_here={} graph_force={})",
2026 self.attn_softcap > 0.0,
2027 crate::gpu::enabled_here(),
2028 graph_force,
2029 );
2030 }
2031 if self.graph_head_required {
2032 self.fail_metal_graph("native graph front gate refused");
2033 }
2034 return start;
2035 }
2036 let swa_graph = self.metal_graph_swa();
2043 if (self.swa.is_some() && !swa_graph)
2044 || self.global_attn.is_some()
2045 || self.attention_heads_per_layer.is_some()
2046 || self.attn_v_norm
2047 || (self.graph_attn_decline_reason().is_some() && !swa_graph)
2049 || self.weights.layers.iter().any(|lw| {
2050 lw.attn_out_norm.is_some()
2051 || lw.ffn_out_norm.is_some()
2052 || lw.layer_scale.is_some()
2053 || matches!(&lw.ffn, FfnKind::Dense(d) if !matches!(d.act, Act::Silu | Act::Gelu))
2054 })
2055 {
2056 if let Some(reason) = self.graph_attn_decline_reason().filter(|_| !swa_graph) {
2059 self.note_graph_decline("metal block graph", reason);
2060 }
2061 if std::env::var("CMF_GRAPH_DBG").is_ok() {
2062 eprintln!(
2063 "block-graph: arch ineligible (swa={} gattn={} hpl={} vnorm={} scale_delta={:.2e})",
2064 self.swa.is_some(),
2065 self.global_attn.is_some(),
2066 self.attention_heads_per_layer.is_some(),
2067 self.attn_v_norm,
2068 (self.attn_scale - 1.0 / (self.head_dim as f32).sqrt()).abs(),
2069 );
2070 }
2071 if self.graph_head_required {
2072 self.fail_metal_graph("native graph architecture gate refused");
2073 }
2074 return start;
2075 }
2076 let limit = upto
2079 .map(|u| u + 1)
2080 .unwrap_or(self.num_layers)
2081 .min(self.num_layers);
2082
2083 enum Item<'a> {
2084 Gdn {
2085 run: Vec<GdnGpuLayer<'a>>,
2086 first: usize,
2087 },
2088 Attn {
2089 l: AttnGpuLayer<'a>,
2090 li: usize,
2091 q_norm: Option<&'a [f32]>,
2092 k_norm: Option<&'a [f32]>,
2093 output_gate: bool,
2094 proj_gate: Option<(&'a QTensor, bool)>,
2098 head_gate_w: Option<&'a [f32]>,
2101 bias: Option<(&'a [f32], &'a [f32], &'a [f32])>,
2102 full_gpu: bool,
2105 },
2106 }
2107
2108 let attend_mode = std::env::var("CMF_GPU_ATTEND").unwrap_or_else(|_| "auto".into());
2115 let attend_contract = attend_mode != "0"
2116 && attend_mode != "off"
2117 && self.head_dim % 4 == 0
2118 && self.head_dim <= 256
2119 && self.rotary_dim >= 2
2120 && self.rotary_dim <= self.head_dim
2121 && (self.rotary_dim / 2) % 32 == 0
2122 && self.num_kv_heads > 0
2123 && self.num_heads % self.num_kv_heads == 0;
2124
2125 let mut plan: Vec<Item> = Vec::new();
2126 let mut model_ref: Option<std::sync::Arc<cortiq_core::CmfModel>> = None;
2127 let block_diag = std::env::var("CMF_GRAPH_DBG").is_ok();
2129 let mut scan = start;
2130 while scan < limit {
2131 let lw = &self.weights.layers[self.phys_layer(scan)];
2132 let ffn = match &lw.ffn {
2133 FfnKind::Dense(d) if d.segs.is_empty() => {
2134 let (Some(g), Some(u), Some(dn)) = (
2135 d.gate_proj.metal_graph_parts(),
2136 d.up_proj.metal_graph_parts(),
2137 d.down_proj.metal_graph_parts(),
2138 ) else {
2139 if block_diag {
2140 eprintln!(
2141 "block-graph: L{scan} FFN trio not graph-mappable — run ends"
2142 );
2143 }
2144 break;
2145 };
2146 MetalFfn::Dense {
2147 gate: g,
2148 up: u,
2149 down: dn,
2150 gelu: d.act == Act::Gelu,
2152 }
2153 }
2154 FfnKind::Moe(m) => {
2155 let Some(moe) = metal_moe_graph_parts(m, self.hidden_size) else {
2156 if block_diag {
2157 eprintln!(
2158 "block-graph: L{scan} MoE outside the graph contract — run ends"
2159 );
2160 }
2161 break;
2162 };
2163 if let QTensor::Mapped { model, .. } = &m.experts[0].gate_proj {
2164 model_ref.get_or_insert_with(|| model.clone());
2165 }
2166 MetalFfn::Moe(moe)
2167 }
2168 _ => {
2169 if block_diag {
2170 eprintln!("block-graph: L{scan} non-graph FFN — run ends");
2171 }
2172 break;
2173 }
2174 };
2175 match &lw.attn {
2176 AttnKind::LinearGdn(w) if self.gdn_cfg.is_some() => {
2177 let parts = (
2178 w.in_proj_qkv.metal_graph_parts(),
2179 w.in_proj_z.metal_graph_parts(),
2180 w.in_proj_a.f32_parts(),
2181 w.in_proj_b.f32_parts(),
2182 w.out_proj.metal_graph_parts(),
2183 );
2184 let (Some(qkv), Some(z), Some(a), Some(b), Some(out)) = parts else {
2185 if block_diag {
2186 eprintln!(
2187 "block-graph: L{scan} GDN parts refused (qkv={} z={} a_f32={} b_f32={} out={})",
2188 w.in_proj_qkv.metal_graph_parts().is_some(),
2189 w.in_proj_z.metal_graph_parts().is_some(),
2190 w.in_proj_a.f32_parts().is_some(),
2191 w.in_proj_b.f32_parts().is_some(),
2192 w.out_proj.metal_graph_parts().is_some(),
2193 );
2194 }
2195 break;
2196 };
2197 if let QTensor::Mapped { model, .. } = &w.in_proj_qkv {
2198 model_ref.get_or_insert_with(|| model.clone());
2199 }
2200 let gl = GdnGpuLayer {
2201 attn_norm: &lw.input_norm,
2202 post_norm: &lw.post_norm,
2203 qkv,
2204 z,
2205 a,
2206 b,
2207 out,
2208 ffn,
2209 conv1d: &w.conv1d,
2210 a_log: &w.a_log,
2211 dt_bias: &w.dt_bias,
2212 gnorm: &w.norm,
2213 };
2214 match plan.last_mut() {
2215 Some(Item::Gdn { run, .. }) => run.push(gl),
2216 _ => plan.push(Item::Gdn {
2217 run: vec![gl],
2218 first: scan,
2219 }),
2220 }
2221 }
2222 AttnKind::Full {
2223 wq,
2224 wk,
2225 wv,
2226 wo,
2227 q_norm,
2228 k_norm,
2229 output_gate,
2230 softplus_gate,
2231 bias,
2232 } if (!self.kv_cache.layers[scan].o1_sealed()
2233 || std::env::var("CMF_O1_METAL").as_deref() == Ok("1"))
2238 && softplus_gate
2242 .as_ref()
2243 .is_none_or(|(_, per_head)| *per_head && self.proj_gate_sigmoid) =>
2244 {
2245 let parts = (
2246 wq.metal_graph_parts(),
2247 wk.metal_graph_parts(),
2248 wv.metal_graph_parts(),
2249 wo.metal_graph_parts(),
2250 );
2251 let (Some(pq), Some(pk), Some(pv), Some(po)) = parts else {
2252 break;
2253 };
2254 if let QTensor::Mapped { model, .. } = wq {
2255 model_ref.get_or_insert_with(|| model.clone());
2256 }
2257 let cache = &self.kv_cache.layers[scan];
2258 let o1_metal = cache.o1.is_some()
2262 && std::env::var("CMF_O1_METAL").as_deref() == Ok("1")
2263 && cache.o1_views().is_some();
2264 let head_gate_w = softplus_gate
2268 .as_ref()
2269 .and_then(|(g, _)| g.f32_parts())
2270 .filter(|&(_, r, c)| r == self.num_heads && c == self.hidden_size)
2271 .map(|(d, _, _)| d);
2272 let full_gpu = attend_contract
2273 && softplus_gate.is_none() == head_gate_w.is_none()
2274 && cache.mode == crate::kv_cache::KvMode::F32
2275 && (cache.o1.is_none() || o1_metal)
2276 && bias.is_none()
2277 && pq.1 == self.num_heads * self.head_dim * (1 + *output_gate as usize)
2278 && pk.1 == self.num_kv_heads * self.head_dim
2279 && pv.1 == self.num_kv_heads * self.head_dim
2280 && po.2 == self.num_heads * self.head_dim;
2281 plan.push(Item::Attn {
2282 l: AttnGpuLayer {
2283 attn_norm: &lw.input_norm,
2284 post_norm: &lw.post_norm,
2285 wq: pq,
2286 wk: pk,
2287 wv: pv,
2288 wo: po,
2289 ffn,
2290 },
2291 li: scan,
2292 q_norm: q_norm.as_deref(),
2293 k_norm: k_norm.as_deref(),
2294 output_gate: *output_gate,
2295 proj_gate: softplus_gate.as_ref().map(|(g, per_head)| (g, *per_head)),
2296 head_gate_w,
2297 bias: bias
2298 .as_ref()
2299 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
2300 full_gpu,
2301 });
2302 }
2303 _ => break,
2304 }
2305 scan += 1;
2306 }
2307 let Some(model) = model_ref else {
2308 if std::env::var("CMF_GRAPH_DBG").is_ok() {
2309 eprintln!("q1-graph: no model ref (start {start}, scanned to {scan})");
2310 }
2311 if self.graph_head_required {
2312 self.fail_metal_graph("native graph has no mapped model reference");
2313 }
2314 return start;
2315 };
2316 if plan.is_empty() {
2317 if std::env::var("CMF_GRAPH_DBG").is_ok() {
2318 eprintln!("q1-graph: empty plan at layer {start}");
2319 }
2320 if self.graph_head_required {
2321 self.fail_metal_graph("native graph plan is empty");
2322 }
2323 return start;
2324 }
2325 let has_moe = plan.iter().any(|it| match it {
2326 Item::Gdn { run, .. } => run.iter().any(|l| matches!(l.ffn, MetalFfn::Moe(_))),
2327 Item::Attn { l, .. } => matches!(l.ffn, MetalFfn::Moe(_)),
2328 });
2329 let has_gdn = plan.iter().any(|it| matches!(it, Item::Gdn { .. }));
2330 let dev_attend = attend_contract
2331 && (self.head_dim <= 128
2332 || has_moe
2333 || (self.head_dim <= 256 && has_gdn)
2339 || (self.head_dim <= 256 && swa_graph)
2343 || attend_mode == "force"
2344 || attend_mode == "256");
2345 if !dev_attend {
2346 for it in &mut plan {
2347 if let Item::Attn { li, full_gpu, .. } = it {
2348 let keep_o1 = self.kv_cache.layers[*li].o1.is_some()
2351 && std::env::var("CMF_O1_METAL").as_deref() == Ok("1");
2352 if !keep_o1 {
2353 *full_gpu = false;
2354 }
2355 }
2356 }
2357 }
2358 if std::env::var("CMF_GRAPH_DBG").is_ok() {
2359 use std::sync::atomic::{AtomicBool, Ordering};
2360 static SAID: AtomicBool = AtomicBool::new(false);
2361 if !SAID.swap(true, Ordering::Relaxed) {
2362 let fg = plan
2363 .iter()
2364 .filter(|it| matches!(it, Item::Attn { full_gpu: true, .. }))
2365 .count();
2366 let att = plan
2367 .iter()
2368 .filter(|it| matches!(it, Item::Attn { .. }))
2369 .count();
2370 eprintln!(
2371 "q1-graph: plan of {} items from layer {start} to {scan} | dev_attend={dev_attend} full_gpu {fg}/{att} | hd={} rd={} nkv={} nh={}",
2372 plan.len(),
2373 self.head_dim,
2374 self.rotary_dim,
2375 self.num_kv_heads,
2376 self.num_heads,
2377 );
2378 }
2379 }
2380 let dims = GraphDims {
2381 hidden: self.hidden_size,
2382 eps: self.rms_eps as f32,
2383 gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
2384 };
2385 let Some(mut graph) = TokenGraph::new(&model, dims, h) else {
2386 if self.graph_head_required {
2387 self.fail_metal_graph("native TokenGraph allocation refused");
2388 }
2389 return start;
2390 };
2391 if swa_graph && self.head_dim > 128 {
2392 graph.set_attend_blk_from(64);
2398 }
2399 let gcfg = self.gdn_cfg.map(|cfg| GdnGpuCfg {
2400 nv: cfg.num_v_heads,
2401 nk: cfg.num_k_heads,
2402 dk: cfg.key_head_dim,
2403 dv: cfg.value_head_dim,
2404 kk: cfg.conv_kernel,
2405 hidden: self.hidden_size,
2406 inter: self.intermediate_size,
2407 c_dim: cfg.conv_dim(),
2408 eps: cfg.rms_eps as f32,
2409 gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
2410 });
2411 let mut valid = 0usize;
2415 let mut end = start;
2416 crate::gpu::stageprof(1, _mt0.elapsed()); if std::env::var("CMF_PLAN_DUMP").is_ok() {
2418 static ONCE: std::sync::Once = std::sync::Once::new();
2419 ONCE.call_once(|| {
2420 for it in &plan {
2421 match it {
2422 Item::Gdn { first, run } => {
2423 eprintln!("plan: Gdn first={first} len={}", run.len())
2424 }
2425 Item::Attn { li, full_gpu, .. } => {
2426 eprintln!("plan: Attn li={li} full_gpu={full_gpu}")
2427 }
2428 }
2429 }
2430 });
2431 }
2432 for item in &plan {
2433 let ok = match item {
2434 Item::Gdn { run, .. } => gcfg
2435 .as_ref()
2436 .map(|gc| run.iter().all(|l| graph.gdn_ok(l, gc)))
2437 .unwrap_or(false),
2438 Item::Attn { l, .. } => graph.attn_ok(l),
2439 };
2440 if !ok {
2441 if block_diag {
2442 eprintln!(
2443 "block-graph: plan item {} ({}) failed graph preflight",
2444 valid,
2445 match item {
2446 Item::Gdn { run, first } => format!("GDN run L{first}+{}", run.len()),
2447 Item::Attn { li, .. } => format!("Attn L{li}"),
2448 }
2449 );
2450 }
2451 break;
2452 }
2453 valid += 1;
2454 end += match item {
2455 Item::Gdn { run, .. } => run.len(),
2456 Item::Attn { .. } => 1,
2457 };
2458 }
2459 plan.truncate(valid);
2460 if plan.is_empty() {
2461 if self.graph_head_required {
2462 self.fail_metal_graph("native graph preflight produced no valid items");
2463 }
2464 return start;
2465 }
2466
2467 if self.graph_head_required && (upto.is_some() || end != self.num_layers) {
2468 self.fail_metal_graph("fused-head NLL requires a complete 64-layer graph");
2469 return start;
2470 }
2471
2472 let one_pass = |t: (usize, usize, usize)| {
2481 use cortiq_core::TensorDtype as D;
2482 matches!(
2483 model.tensors[t.0].dtype,
2484 D::Q4TiledP | D::Q4Tiled | D::Q4Block | D::Q8Row | D::Q8_2f | D::Q1
2485 )
2486 };
2487 let dense_fast = plan.iter().all(|it| match it {
2488 Item::Attn {
2489 l, li, full_gpu, ..
2490 } => {
2491 *full_gpu
2492 && self.kv_cache.layers[*li].o1.is_none()
2493 && [l.wq, l.wk, l.wv, l.wo].into_iter().all(one_pass)
2494 && match l.ffn {
2495 MetalFfn::Dense { gate, up, down, .. } => {
2496 one_pass(gate) && one_pass(up) && one_pass(down)
2497 }
2498 _ => false,
2499 }
2500 }
2501 Item::Gdn { .. } => false,
2502 });
2503 let ab = crate::gpu_metal::dense_ab_arm().filter(|_| dense_fast);
2504 let _mv_fast = match ab {
2505 Some((bits, _)) => {
2506 graph.set_dense_concurrent_raw(bits & crate::gpu_metal::DENSE_CONC != 0);
2507 crate::gpu_metal::MvFastGuard::set_raw(bits)
2508 }
2509 None => {
2510 graph.set_dense_concurrent(dense_fast);
2511 let q8r4 = if swa_graph {
2514 crate::gpu_metal::DENSE_Q8R4
2515 } else {
2516 0
2517 };
2518 crate::gpu_metal::MvFastGuard::set_bits(if dense_fast {
2519 crate::gpu_metal::DENSE_MV | crate::gpu_metal::DENSE_FUSE | q8r4
2520 } else {
2521 0
2522 })
2523 }
2524 };
2525
2526 let pool = self.pool.clone();
2529 let (nh, nkv, hd, hs, eps) = (
2530 self.num_heads,
2531 self.num_kv_heads,
2532 self.head_dim,
2533 self.hidden_size,
2534 self.rms_eps,
2535 );
2536 let norm_style = self.norm_style;
2537 let gemma = norm_style == cortiq_core::NormStyle::Gemma;
2538 let want = self.gdn_cfg.map(|c| c.state_len()).unwrap_or(0);
2539 let kv_id = self.graph_kv_id;
2540 let mut pending: Vec<(usize, usize)> = Vec::new();
2543 let mut dev_attn: Vec<usize> = Vec::new();
2546 for item in &plan {
2547 let _xt0 = std::time::Instant::now();
2548 let _xkind: u32 = match item {
2549 Item::Gdn { .. } => 2,
2550 Item::Attn { .. } => 3,
2551 };
2552 if self.loop_final_norm {
2554 let item_start = match item {
2555 Item::Gdn { first, .. } => *first,
2556 Item::Attn { li, .. } => *li,
2557 };
2558 if item_start > start && self.is_loop_end(item_start - 1) {
2559 graph.encode_loop_norm(&self.weights.final_norm);
2560 }
2561 }
2562 match item {
2563 Item::Gdn { run, first } => {
2564 for l in &mut self.kv_cache.layers[*first..*first + run.len()] {
2565 if l.linear_state.len() != want {
2566 l.linear_state = vec![0f32; want];
2567 }
2568 }
2569 let ro: Vec<&[f32]> = self.kv_cache.layers[*first..*first + run.len()]
2570 .iter()
2571 .map(|l| l.linear_state.as_slice())
2572 .collect();
2573 let _ig = std::time::Instant::now();
2574 if !graph.encode_gdn_run(run, &ro, gcfg.as_ref().unwrap()) {
2575 tracing::error!("q1 graph: GDN run refused after validation");
2577 return start;
2578 }
2579 graph.commit_kind = 2;
2582 graph.commit();
2583 crate::gpu::stageprof(0, _ig.elapsed());
2584 pending.push((*first, run.len()));
2585 }
2586 Item::Attn {
2587 l,
2588 li,
2589 q_norm,
2590 k_norm,
2591 output_gate,
2592 proj_gate,
2593 head_gate_w,
2594 bias,
2595 full_gpu,
2596 } => {
2597 let _ia = std::time::Instant::now();
2598 let inv_freq_l = self.layer_inv_freq(*li);
2604 let rd_l = self.layer_geom(*li).2;
2605 let window_l = self.layer_window(*li);
2606 let head_gate_w = *head_gate_w;
2607 if *full_gpu {
2609 let cache = &self.kv_cache.layers[*li];
2610 let o1p = if cache.o1.is_some() {
2611 match cache.o1_views() {
2612 Some(views) => Some(crate::gpu::O1AttnParams {
2613 views,
2614 epoch: self.o1_epoch,
2615 }),
2616 None => None,
2618 }
2619 } else {
2620 None
2621 };
2622 let o1_layer = cache.o1.is_some();
2623 if o1_layer && o1p.is_none() {
2624 }
2626 let cpu_k: Vec<&[f32]> = (0..nkv).map(|g| cache.head_keys(g)).collect();
2627 let cpu_v: Vec<&[f32]> = (0..nkv).map(|g| cache.head_values(g)).collect();
2628 let cpu_stored = if o1_layer { 0 } else { cpu_k[0].len() / hd };
2629 let p = crate::gpu::AttnDeviceParams {
2630 kv_id,
2631 layer: *li,
2632 nh,
2633 nkv,
2634 hd,
2635 rd: rd_l,
2636 position,
2637 scale: self.attn_scale,
2638 eps: eps as f32,
2639 gemma,
2640 late_qk_norm: self.qk_norm_after_rope,
2641 output_gate: *output_gate,
2642 q_norm: *q_norm,
2643 k_norm: *k_norm,
2644 inv_freq: &inv_freq_l,
2645 cpu_k,
2646 cpu_v,
2647 cpu_stored,
2648 o1: o1p,
2649 window: window_l,
2650 head_gate: head_gate_w,
2651 };
2652 let o1_bad = o1_layer && p.o1.is_none();
2653 if !o1_bad && graph.attn_device_ok(l, &p) && graph.encode_attn_device(l, &p)
2654 {
2655 if p.o1.is_none() {
2657 dev_attn.push(*li);
2658 }
2659 graph.commit_kind = 3;
2660 graph.commit();
2661 crate::gpu::stageprof(_xkind, _xt0.elapsed());
2665 continue;
2666 }
2667 }
2669 graph.encode_attn_prefix(l);
2670 if let Err(err) = graph.sync_checked() {
2671 self.fail_metal_graph(&err);
2672 return start;
2673 }
2674 if !pending.is_empty() {
2675 let idxs: Vec<usize> =
2676 pending.drain(..).flat_map(|(f, n)| f..f + n).collect();
2677 let mut outs: Vec<&mut [f32]> = self
2678 .kv_cache
2679 .layers
2680 .iter_mut()
2681 .enumerate()
2682 .filter(|(i, _)| idxs.binary_search(i).is_ok())
2683 .map(|(_, s)| s.linear_state.as_mut_slice())
2684 .collect();
2685 graph.read_states(&mut outs);
2686 }
2687 let mut q_raw = attention::take_buf(l.wq.1);
2688 let mut k = attention::take_buf(l.wk.1);
2689 let mut v = attention::take_buf(l.wv.1);
2690 graph.read_qkv(&mut q_raw, &mut k, &mut v);
2691 let mut gate_raw = proj_gate.map(|(gp, _)| {
2694 let mut normed = attention::take_buf(hs);
2695 graph.read_normed(&mut normed);
2696 let mut raw = attention::take_buf(gp.rows());
2697 gp.matvec(&normed, &mut raw, pool.as_deref());
2698 attention::recycle_buf(&mut normed);
2699 raw
2700 });
2701 let cfg = QwenAttnCfg {
2702 num_heads: nh,
2703 num_kv_heads: nkv,
2704 head_dim: hd,
2705 hidden_size: hs,
2706 position,
2707 inv_freq: &inv_freq_l,
2708 rotary_dim: rd_l,
2709 scale: self.attn_scale,
2710 softcap: self.attn_softcap,
2711 window: window_l,
2712 v_norm: false,
2713 qk_norm_after_rope: self.qk_norm_after_rope,
2714 gate_sigmoid: self.proj_gate_sigmoid,
2715 q_norm: *q_norm,
2716 k_norm: *k_norm,
2717 output_gate: *output_gate,
2718 softplus_gate: None,
2719 rope_scale: 1.0,
2720 bias: *bias,
2721 rms_eps: eps,
2722 norm_style,
2723 pool: pool.as_deref(),
2724 v_head_dim: hd,
2725 };
2726 let oracle = std::env::var("CMF_ATTN_ORACLE").as_deref() == Ok("1")
2729 || std::env::var("CMF_ATTN_DUMP").is_ok();
2730 let _ = full_gpu;
2731 let oracle_in = oracle.then(|| (q_raw.clone(), k.clone(), v.clone()));
2732 let mut ao = attention::qwen_attention_core(
2733 q_raw,
2734 k,
2735 v,
2736 &mut self.kv_cache.layers[*li],
2737 &cfg,
2738 );
2739 if let Ok(dir) = std::env::var("CMF_ATTN_DUMP") {
2743 if let Some((qr0, k0, v0)) = oracle_in.clone() {
2744 let (cq, _cg, _ck, _cv) =
2745 attention::finish_projection_debug(qr0, k0, v0, &cfg, position);
2746 let cache = &self.kv_cache.layers[*li];
2747 let n = cache.head_keys(0).len() / hd;
2748 let mut bytes: Vec<u8> = Vec::new();
2749 for v in [nh as u32, nkv as u32, hd as u32, n as u32, position as u32] {
2750 bytes.extend_from_slice(&v.to_le_bytes());
2751 }
2752 for v in &cq {
2753 bytes.extend_from_slice(&v.to_le_bytes());
2754 }
2755 for g in 0..nkv {
2756 for v in cache.head_keys(g) {
2757 bytes.extend_from_slice(&v.to_le_bytes());
2758 }
2759 }
2760 for g in 0..nkv {
2761 for v in cache.head_values(g) {
2762 bytes.extend_from_slice(&v.to_le_bytes());
2763 }
2764 }
2765 let _ =
2766 std::fs::write(format!("{dir}/L{li}_pos{position}.bin"), &bytes);
2767 }
2768 }
2769 if let Some((qr0, k0, v0)) =
2770 oracle_in.filter(|_| std::env::var("CMF_ATTN_ORACLE").as_deref() == Ok("1"))
2771 {
2772 let (cq, _cg, ck, cv) =
2773 attention::finish_projection_debug(qr0, k0, v0, &cfg, position);
2774 let mut h_now = vec![0f32; hs];
2775 graph.read_h(&mut h_now);
2776 let cache = &self.kv_cache.layers[*li];
2777 let n_after = cache.head_keys(0).len() / hd;
2778 let stored = n_after.saturating_sub(1);
2782 let cpu_k: Vec<&[f32]> = (0..nkv)
2783 .map(|g| &cache.head_keys(g)[..stored * hd])
2784 .collect();
2785 let cpu_v: Vec<&[f32]> = (0..nkv)
2786 .map(|g| &cache.head_values(g)[..stored * hd])
2787 .collect();
2788 let p = crate::gpu::AttnDeviceParams {
2789 kv_id,
2790 layer: *li,
2791 nh,
2792 nkv,
2793 hd,
2794 rd: rd_l,
2798 position,
2799 scale: self.attn_scale,
2800 eps: eps as f32,
2801 gemma,
2802 late_qk_norm: self.qk_norm_after_rope,
2803 output_gate: *output_gate,
2804 q_norm: *q_norm,
2805 k_norm: *k_norm,
2806 inv_freq: &inv_freq_l,
2807 cpu_k,
2808 cpu_v,
2809 cpu_stored: stored,
2810 o1: None,
2811 window: window_l,
2812 head_gate: None,
2813 };
2814 if let Some((dq, dk, dv, dao)) = graph.debug_attn_device(l, &p, &h_now) {
2815 let md = |a: &[f32], b: &[f32]| {
2816 a.iter().zip(b).fold(0f32, |m, (x, y)| m.max((x - y).abs()))
2817 };
2818 let nn = |a: &[f32]| a.iter().map(|x| x * x).sum::<f32>().sqrt();
2819 eprintln!(
2820 "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}",
2821 nn(&cq),
2822 md(&cq, &dq),
2823 nn(&ck),
2824 md(&ck, &dk),
2825 nn(&cv),
2826 md(&cv, &dv),
2827 nn(&ao),
2828 md(&ao, &dao)
2829 );
2830 } else {
2831 eprintln!("attn-oracle L{li}: device probe declined");
2832 }
2833 }
2834 if let (Some(raw), Some((_, per_head))) = (gate_raw.as_deref(), *proj_gate) {
2835 attention::apply_projected_gate(
2838 &mut ao,
2839 raw,
2840 per_head,
2841 hd,
2842 self.proj_gate_sigmoid,
2843 );
2844 }
2845 if let Some(mut raw) = gate_raw.take() {
2846 attention::recycle_buf(&mut raw);
2847 }
2848 graph.encode_attn_suffix(l, &ao);
2849 graph.commit();
2852 attention::recycle_buf(&mut ao);
2853 }
2854 }
2855
2856 crate::gpu::stageprof(_xkind, _xt0.elapsed());
2857 }
2858 let mut lm_rows = None;
2863 if self.graph_want_logits
2864 && upto.is_none()
2865 && end == self.num_layers
2866 && std::env::var("CMF_GPU_LMHEAD")
2867 .map(|v| v != "0")
2868 .unwrap_or(true)
2869 {
2870 if let Some(lm) = self.weights.lm_head.metal_graph_parts() {
2871 if graph.lm_head_ok(lm) {
2872 graph.encode_lm_head(&self.weights.final_norm, lm);
2873 lm_rows = Some(lm.1);
2874 }
2875 }
2876 }
2877 if self.graph_head_required && lm_rows.is_none() {
2878 METAL_GRAPH_HEAD_MISS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
2879 self.fail_metal_graph("fused graph head was requested but not encodable");
2880 return start;
2881 }
2882 let _sy0 = std::time::Instant::now();
2883 if let Err(err) = graph.sync_checked() {
2884 self.fail_metal_graph(&err);
2885 return start;
2886 }
2887 let _rs0 = std::time::Instant::now();
2888 if !pending.is_empty() {
2889 let idxs: Vec<usize> = pending.drain(..).flat_map(|(f, n)| f..f + n).collect();
2890 let mut outs: Vec<&mut [f32]> = self
2891 .kv_cache
2892 .layers
2893 .iter_mut()
2894 .enumerate()
2895 .filter(|(i, _)| idxs.binary_search(i).is_ok())
2896 .map(|(_, s)| s.linear_state.as_mut_slice())
2897 .collect();
2898 graph.read_states(&mut outs);
2899 }
2900 if std::env::var("CMF_GRAPH_HOSTPROF").as_deref() == Ok("1") {
2901 use std::sync::atomic::{AtomicU64, Ordering};
2902 static SY: AtomicU64 = AtomicU64::new(0);
2903 static RS: AtomicU64 = AtomicU64::new(0);
2904 static N: AtomicU64 = AtomicU64::new(0);
2905 SY.fetch_add((_rs0 - _sy0).as_nanos() as u64, Ordering::Relaxed);
2906 RS.fetch_add(_rs0.elapsed().as_nanos() as u64, Ordering::Relaxed);
2907 let n = N.fetch_add(1, Ordering::Relaxed) + 1;
2908 if n % 100 == 0 {
2909 eprintln!(
2910 "postprof: sync-wait {:.1} ms/ток | read_states {:.1} ms/ток ({n})",
2911 SY.load(Ordering::Relaxed) as f64 / n as f64 / 1e6,
2912 RS.load(Ordering::Relaxed) as f64 / n as f64 / 1e6
2913 );
2914 }
2915 }
2916 if let Some(rows) = lm_rows {
2917 crate::gpu::hostprof_encode_done(_mt0);
2918 let mut lg = attention::take_buf(rows.min(self.vocab_size));
2919 graph.read_logits(&mut lg);
2920 crate::gpu::hostprof_total(_mt0);
2921 lg.resize(self.vocab_size, 0.0);
2922 if let Some(c) = self.final_softcap {
2923 for l in lg.iter_mut() {
2924 *l = c * (*l / c).tanh();
2925 }
2926 }
2927 self.graph_logits = Some(lg);
2928 }
2929 graph.read_h(h);
2930 if self.graph_head_required && self.graph_logits.is_none() {
2931 METAL_GRAPH_HEAD_MISS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
2932 self.fail_metal_graph("fused graph head completed without logits readback");
2933 return start;
2934 }
2935 METAL_GRAPH_TOK_OK.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
2936 METAL_GRAPH_LAYERS.fetch_add(
2937 end.saturating_sub(start) as u64,
2938 std::sync::atomic::Ordering::Relaxed,
2939 );
2940 if self.graph_head_required {
2941 METAL_GRAPH_HEAD_OK.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
2942 }
2943 for li in dev_attn {
2947 let mut krow = attention::take_buf(nkv * hd);
2948 let mut vrow = attention::take_buf(nkv * hd);
2949 if crate::gpu::kv_mirror_read_last(kv_id, li, nkv, hd, &mut krow, &mut vrow) {
2950 let cache = &mut self.kv_cache.layers[li];
2951 cache.append(&krow, &vrow, &[]);
2952 let n = cache.seq_len;
2953 let mut imp = attention::take_buf(n);
2954 crate::gpu::kv_mirror_take_imp(kv_id, li, &mut imp);
2955 cache.accumulate_imp(&imp);
2956 attention::recycle_buf(&mut imp);
2957 }
2958 attention::recycle_buf(&mut krow);
2959 attention::recycle_buf(&mut vrow);
2960 }
2961 if let Some((_, arm)) = ab {
2962 crate::gpu_metal::dense_ab_record(arm, _mt0.elapsed());
2963 }
2964 end
2965 }
2966
2967 pub fn new(
2968 tokenizer: Tokenizer,
2969 weights: PipelineWeights,
2970 hidden_size: usize,
2971 intermediate_size: usize,
2972 num_heads: usize,
2973 num_kv_heads: usize,
2974 head_dim: usize,
2975 num_layers: usize,
2976 physical_layers: usize,
2977 loop_final_norm: bool,
2978 vocab_size: usize,
2979 rms_eps: f64,
2980 rope_base: f32,
2981 norm_style: NormStyle,
2982 max_seq_len: usize,
2983 sampler_config: SamplerConfig,
2984 ) -> Self {
2985 let rng = match sampler_config.seed {
2986 Some(s) => SplitMix64::new(s),
2987 None => SplitMix64::from_entropy(),
2988 };
2989 let inv_freq = std::sync::Arc::new(attention::rope_inv_freq(head_dim, rope_base));
2990 let pool = Pool::from_env();
2991 if let Some(p) = &pool {
2992 tracing::info!("worker pool: {} threads", p.n_workers());
2993 if let Some(model) = weights
2995 .lm_head
2996 .model_arc()
2997 .or_else(|| weights.embed_tokens.model_arc())
2998 {
2999 let regions: Vec<&[u8]> =
3000 model.tensors.iter().map(|t| model.entry_bytes(t)).collect();
3001 p.bind_numa(®ions);
3002 }
3003 }
3004 Self {
3005 gpu_plan: None,
3006 tokenizer: std::sync::Arc::new(tokenizer),
3007 kv_cache: KvCache::new(num_layers, num_kv_heads, head_dim, max_seq_len),
3008 sampler_config,
3009 weights,
3010 hidden_size,
3011 intermediate_size,
3012 num_heads,
3013 num_kv_heads,
3014 head_dim,
3015 num_layers,
3016 physical_layers,
3017 loop_final_norm,
3018 vocab_size,
3019 rms_eps,
3020 rope_base,
3021 norm_style,
3022 rotary_dim: head_dim,
3023 attention_heads_per_layer: None,
3024 kv_heads_per_layer: None,
3025 v_head_dim: None,
3026 layer_dump: std::env::var_os("CMF_LAYER_DUMP")
3027 .filter(|v| !v.is_empty())
3028 .map(std::path::PathBuf::from),
3029 graph_declines: std::cell::RefCell::new(Vec::new()),
3030 mimo_moe: Default::default(),
3031 vmf_cfg: None,
3032 gdn_cfg: None,
3033 kda_cfg: None,
3034 g3n: None,
3035 dsv4: None,
3036 dsv41: None,
3037 dsv41_vision: None,
3038 dsv41_prefill: None,
3039 qwen4_exp: None,
3040 dsv4_mtp: Vec::new(),
3041 dspark: None,
3042 dspark_pending: Vec::new(),
3043 dspark_hist: Vec::new(),
3044 dspark_real: Vec::new(),
3045 dspark_trunk_picks: Vec::new(),
3046 dspark_exp: Vec::new(),
3047 dspark_draft_ns: 0,
3048 logit_multiplier: None,
3049 cancel: std::sync::Arc::new(std::sync::atomic::AtomicBool::new(false)),
3050 graph_failed: std::sync::atomic::AtomicBool::new(false),
3051 kv_history: Vec::new(),
3052 kv_history_device: false,
3053 short_conv_cfg: None,
3054 mtp: None,
3055 mimo_mtp: None,
3056 verify_exact_moe: false,
3057 speculative: std::env::var("CMF_MTP").map(|v| v != "0").unwrap_or(true),
3058 ignore_eos: false,
3059 draft_full_streak: 0,
3060 spec_k_adapt: None,
3061 spec_acc_ewma: 0.7,
3062 rng,
3063 sampler_scratch: SamplerScratch::default(),
3064 spec_forced: None,
3065 spec_q: Vec::new(),
3066 spec_p: Vec::new(),
3067 spec_res: Vec::new(),
3068 spec_qs: Vec::new(),
3069 spec_ps: Vec::new(),
3070 spec_ress: Vec::new(),
3071 mtp_graph_mode: None,
3072 #[cfg(target_os = "macos")]
3073 metal_verify: None,
3074 inv_freq,
3075 ws: ForwardScratch::new(hidden_size),
3076 pool,
3077 model: None,
3078 dyn_force_f32: false,
3079 dyn_skill_layers: Vec::new(),
3080 dyn_active: None,
3081 dyn_blend_loaded: false,
3082 dyn_phi_layer: None,
3083 dyn_phi_ema: Vec::new(),
3084 dyn_phi_seen: 0,
3085 dyn_router: None,
3086 o1_cfg: None,
3087 o1_epoch: 0,
3088 o1_flags: Vec::new(),
3089 trace: false,
3090 calib_temp: 1.0,
3091 confidence_on: true,
3092 embed_multiplier: 1.0,
3093 attn_scale: 1.0 / (head_dim as f32).sqrt(),
3094 swa: None,
3095 sliding_layers: None,
3096 anchor_core: None,
3097 bounded_rope: None,
3098 kv_prefix: KvPrefix::default(),
3099 last_prefill_tokens: 0,
3100 inv_freq_local: None,
3101 rotary_dim_local: None,
3102 rope_scale: 1.0,
3103 rope_scale_local: 1.0,
3104 global_attn: None,
3105 inv_freq_global: None,
3106 attn_v_norm: false,
3107 qk_norm_after_rope: false,
3108 proj_gate_sigmoid: false,
3109 final_softcap: None,
3110 head_clusters: None,
3111 attn_softcap: 0.0,
3112 graph_want_logits: false,
3113 graph_head_required: false,
3114 graph_logits: None,
3115 embryo_graph: None,
3116 graph_refused: std::sync::atomic::AtomicBool::new(false),
3117 graph_kv_id: {
3118 static NEXT: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(1);
3119 NEXT.fetch_add(1, std::sync::atomic::Ordering::Relaxed)
3120 },
3121 #[cfg(test)]
3122 nll_test_fail_at: None,
3123 #[cfg(test)]
3124 nll_test_force_serial: false,
3125 }
3126 }
3127
3128 pub fn set_o1(&mut self, cfg: Option<crate::nystrom::O1Cfg>) {
3136 if let Err(e) = self.try_set_o1(cfg) {
3137 tracing::error!("{e}");
3138 }
3139 }
3140
3141 pub fn bounded_native(&self) -> bool {
3145 self.anchor_core.is_some()
3146 }
3147
3148 pub fn device_state_bytes(&self) -> Option<(u64, u64)> {
3151 crate::gpu::embryo_device_state_bytes(self.graph_kv_id)
3152 }
3153
3154 pub fn o1_refusal(&self) -> Option<String> {
3156 self.anchor_core.as_ref().map(|ac| {
3157 format!(
3158 "--o1 / CMF_O1 refused: the anchor is native bounded \
3159 (anchor_core kind={} window={} sink={}); the file's operator \
3160 is executed as-is and no post-hoc Nyström overlay applies",
3161 ac.kind, ac.window, ac.sink
3162 )
3163 })
3164 }
3165
3166 pub fn try_set_o1(&mut self, cfg: Option<crate::nystrom::O1Cfg>) -> Result<(), String> {
3168 if let Some(c) = &cfg {
3169 if let Some(why) = self.o1_refusal() {
3170 self.o1_flags = Vec::new();
3171 self.o1_cfg = None;
3172 return Err(why);
3173 }
3174 if crate::nystrom::o1_deferred_boundary(c.w, c.sink).is_none() {
3175 self.o1_flags.clear();
3176 self.o1_cfg = None;
3177 return Err(format!(
3178 "o1 disabled: w + sink + slack + 1 overflows usize (w={}, sink={})",
3179 c.w, c.sink
3180 ));
3181 }
3182 }
3183 self.o1_flags = match &cfg {
3184 Some(c) => {
3185 let mut flags = c.layer_flags(self.num_layers);
3186 for (li, f) in flags.iter_mut().enumerate() {
3187 if *f
3193 && (!matches!(
3194 self.weights.layers[self.phys_layer(li)].attn,
3195 AttnKind::Full { .. }
3196 ) || self.layer_window(li).is_some()
3197 || self.kv_cache.layers[li].sinks.is_some()
3198 || self.layer_v_dim(li) != self.layer_geom(li).1)
3199 {
3200 *f = false;
3201 }
3202 }
3203 flags
3204 }
3205 None => Vec::new(),
3206 };
3207 if let Some(c) = &cfg {
3208 let n = self.o1_flags.iter().filter(|&&f| f).count();
3209 tracing::info!(
3210 "o1 nystrom attention: {n}/{} layer(s), m={} w={} sink={} rect={:?}",
3211 self.num_layers,
3212 c.m,
3213 c.w,
3214 c.sink,
3215 c.rect
3216 );
3217 }
3218 self.o1_cfg = cfg;
3219 Ok(())
3220 }
3221
3222 pub fn install_bounded(
3228 &mut self,
3229 cfg: &cortiq_core::AnchorCoreConfig,
3230 ) -> Result<(), String> {
3231 if !cortiq_core::AnchorCoreConfig::KINDS.contains(&cfg.kind.as_str()) {
3232 return Err(format!(
3233 "anchor_core kind '{}' is not executable by this runtime",
3234 cfg.kind
3235 ));
3236 }
3237 if cfg.window == 0 {
3238 return Err("anchor_core.window must be >= 1".into());
3239 }
3240 let mut n = 0usize;
3241 for li in 0..self.num_layers {
3242 let pl = self.phys_layer(li);
3243 if let AttnKind::Bounded(w) = &self.weights.layers[pl].attn {
3244 if w.window != cfg.window || w.sink != cfg.sink {
3245 return Err(format!(
3246 "layer {li}: bounded weights (window {} sink {}) disagree with \
3247 anchor_core (window {} sink {})",
3248 w.window, w.sink, cfg.window, cfg.sink
3249 ));
3250 }
3251 self.kv_cache.layers[li].install_bounded(cfg.window);
3252 n += 1;
3253 }
3254 }
3255 if n == 0 {
3256 return Err("anchor_core is present but no layer executes it".into());
3257 }
3258 let rope = crate::bounded::BoundedRope::new(cfg.window, &self.inv_freq, self.rope_scale);
3259 self.bounded_rope = Some(std::sync::Arc::new(rope));
3260 self.anchor_core = Some(cfg.clone());
3261 self.embryo_graph = None;
3262 tracing::info!(
3263 "bounded anchor {}: {n} layer(s), window {} sink {} — {} B of ring per layer",
3264 cfg.kind,
3265 cfg.window,
3266 cfg.sink,
3267 self.kv_cache.layers.iter().map(|l| l.bounded_state_bytes()).max().unwrap_or(0)
3268 );
3269 Ok(())
3270 }
3271
3272 pub fn install_wire_identity(&mut self, identity: u64) {
3276 for li in 0..self.kv_cache.layers.len() {
3277 let pl = self.phys_layer(li);
3278 let kind = match self.weights.layers.get(pl).map(|l| &l.attn) {
3279 Some(AttnKind::Bounded(_)) => crate::kv_cache::WireKind::Bounded,
3280 Some(AttnKind::Linear(_))
3281 | Some(AttnKind::LinearGdn(_))
3282 | Some(AttnKind::ShortConv(_))
3283 | Some(AttnKind::Kda(_)) => crate::kv_cache::WireKind::Linear,
3284 _ => crate::kv_cache::WireKind::Full,
3285 };
3286 let l = &mut self.kv_cache.layers[li];
3287 l.wire_kind = kind;
3288 l.wire_identity = identity;
3289 l.wire_layer = li as u32;
3292 }
3293 }
3294
3295 pub fn clear_history(&mut self) {
3297 self.kv_history.clear();
3298 self.kv_history_device = false;
3299 self.kv_prefix.clear();
3300 }
3301
3302 pub fn graph_refused(&self) -> bool {
3306 self.graph_refused
3307 .load(std::sync::atomic::Ordering::Relaxed)
3308 }
3309
3310 pub fn mark_graph_refused(&self) {
3312 if !self
3313 .graph_refused
3314 .swap(true, std::sync::atomic::Ordering::Relaxed)
3315 {
3316 tracing::info!(
3317 "token graph: unsupported for this pipeline (seq {}) — not retrying",
3318 self.graph_kv_id
3319 );
3320 }
3321 }
3322
3323 pub fn device_sequence_position(&self) -> Option<usize> {
3327 crate::gpu::embryo_device_next_position(self.graph_kv_id)
3328 }
3329
3330 fn prefix_owner_matches(&self, n: usize, recorded_on_device: bool) -> bool {
3335 let dev = self.device_sequence_position();
3336 if recorded_on_device {
3337 dev == Some(n) && self.embryo_resident_wanted()
3338 } else {
3339 dev.is_none()
3340 }
3341 }
3342
3343 pub(crate) fn invalidate_for_weight_change(&mut self) {
3351 self.clear_sequence_state();
3352 self.embryo_graph = None;
3353 }
3354
3355 fn cached_prefix_len(&self, input_ids: &[u32]) -> usize {
3361 let (n, on_device) = if self.bounded_native() {
3362 (self.kv_prefix.extension(input_ids), self.kv_prefix.on_device())
3363 } else {
3364 let h = &self.kv_history;
3365 if !h.is_empty() && h.len() < input_ids.len() && input_ids[..h.len()] == h[..] {
3366 (h.len(), self.kv_history_device)
3367 } else {
3368 (0, false)
3369 }
3370 };
3371 if n > 0 && !self.prefix_owner_matches(n, on_device) {
3374 tracing::warn!(
3375 "kv-reuse refused: the cached prefix ({n} positions) was built on the {} path, \
3376 the device now holds {:?} — re-prefilling from zero",
3377 if on_device { "resident device" } else { "host" },
3378 self.device_sequence_position()
3379 );
3380 return 0;
3381 }
3382 n
3383 }
3384
3385 pub fn reusable_prefix_len(&self, input_ids: &[u32]) -> usize {
3388 self.cached_prefix_len(input_ids)
3389 }
3390
3391 fn record_consumed_prefix(&mut self, consumed: &[u32], reused: usize) {
3397 let on_device = self.device_sequence_position().is_some();
3398 if self.bounded_native() {
3399 let keep = reused > 0 && reused == self.kv_prefix.len() && reused <= consumed.len();
3400 let prev_device = self.kv_prefix.on_device();
3401 self.kv_history.clear();
3402 self.kv_history_device = false;
3403 if keep && prev_device == on_device {
3404 self.kv_prefix.extend(&consumed[reused..]);
3405 } else {
3406 self.kv_prefix.set(consumed);
3407 }
3408 self.kv_prefix.set_on_device(on_device);
3409 } else {
3410 self.kv_history = consumed.to_vec();
3411 self.kv_history_device = on_device;
3412 }
3413 }
3414
3415 pub fn o1_active(&self) -> bool {
3417 self.o1_cfg.is_some() && self.o1_flags.iter().any(|&f| f)
3418 }
3419
3420 pub fn generation_batch_k(&self) -> usize {
3432 if let Some(k) = std::env::var("CMF_BATCH_K")
3433 .ok()
3434 .and_then(|v| v.parse::<usize>().ok())
3435 {
3436 return k;
3437 }
3438 #[cfg(not(target_os = "macos"))]
3439 if self.graph_prefill_preferred() && !self.o1_active() {
3440 return 32;
3441 }
3442 0
3443 }
3444
3445 pub fn generation_graph_prefill(&self) -> bool {
3446 let graph = self.graph_prefill_preferred();
3447 #[cfg(not(target_os = "macos"))]
3458 if graph
3459 && self.generation_batch_k() > 0
3460 && std::env::var("CMF_MTP_CHAIN_PROBE").is_err()
3461 {
3462 return false;
3463 }
3464 graph
3465 }
3466
3467 pub fn o1_device_stats(&self) -> (usize, u64) {
3472 crate::gpu::o1_device_stats(self.graph_kv_id)
3473 }
3474
3475 pub fn o1_begin(&mut self) {
3480 self.o1_begin_with_prefix(None);
3481 }
3482
3483 pub fn o1_begin_with_prefix(&mut self, requested_prefix: Option<usize>) {
3487 if let Some(c) = &self.o1_cfg {
3488 let (m, w, sink, rect) = (c.m, c.w, c.sink, c.rect);
3489 let boundary = requested_prefix.map(|p| {
3490 p.max(
3491 crate::nystrom::o1_deferred_boundary(w, sink)
3492 .expect("o1 config boundary validated in set_o1"),
3493 )
3494 });
3495 for (li, &f) in self.o1_flags.iter().enumerate() {
3496 if f {
3497 self.kv_cache.layers[li].o1_begin_with_boundary(m, w, sink, rect, boundary);
3498 }
3499 }
3500 }
3501 }
3502
3503 fn o1_effective_boundary(&self, requested_prefix: usize) -> Option<usize> {
3505 self.o1_cfg.as_ref().and_then(|c| {
3506 crate::nystrom::o1_deferred_boundary(c.w, c.sink)
3507 .map(|floor| requested_prefix.max(floor))
3508 })
3509 }
3510
3511 fn o1_note_transition(&mut self) {
3512 let mut transitioned = false;
3516 for (li, &flagged) in self.o1_flags.iter().enumerate() {
3517 if flagged {
3518 transitioned |= self.kv_cache.layers[li].take_o1_transition();
3519 }
3520 }
3521 if transitioned {
3522 self.o1_epoch = self.o1_epoch.wrapping_add(1);
3523 }
3524 }
3525
3526 fn o1_pending(&self) -> bool {
3527 self.o1_flags.iter().enumerate().any(|(li, &f)| {
3528 f && self.kv_cache.layers[li].seq_len > 0
3529 && self.kv_cache.layers[li].o1_pending_boundary().is_some()
3530 })
3531 }
3532
3533 fn o1_fail(&mut self, err: String) {
3534 tracing::error!("o1 deferred seal failed; terminating sequence: {err}");
3535 self.clear_sequence_state();
3536 self.graph_failed
3537 .store(true, std::sync::atomic::Ordering::Relaxed);
3538 self.cancel
3539 .store(true, std::sync::atomic::Ordering::Relaxed);
3540 }
3541
3542 pub fn o1_seal_checked(&mut self) -> Result<bool, String> {
3547 if self.o1_cfg.is_none() {
3548 return Ok(false);
3549 }
3550 let mut participating = false;
3551 for li in 0..self.num_layers {
3552 if !self.o1_flags.get(li).copied().unwrap_or(false) {
3553 continue;
3554 }
3555 if let Some(err) = self.kv_cache.layers[li].take_o1_error() {
3556 return Err(err);
3557 }
3558 if self.kv_cache.layers[li].seq_len == 0 {
3559 continue;
3560 }
3561 participating = true;
3562 let num_heads = self.layer_num_heads(li);
3563 self.kv_cache.layers[li].o1_seal_checked(num_heads)?;
3564 }
3565 self.o1_note_transition();
3566 for li in 0..self.num_layers {
3567 if self.o1_flags.get(li).copied().unwrap_or(false) {
3568 if let Some(err) = self.kv_cache.layers[li].take_o1_error() {
3569 return Err(err);
3570 }
3571 }
3572 }
3573 Ok(participating
3574 && (0..self.num_layers).all(|li| {
3575 !self.o1_flags.get(li).copied().unwrap_or(false)
3576 || self.kv_cache.layers[li].seq_len == 0
3577 || self.kv_cache.layers[li].o1_sealed()
3578 }))
3579 }
3580
3581 fn o1_progress(&mut self) {
3584 if !self.o1_active() {
3585 return;
3586 }
3587 for li in 0..self.num_layers {
3588 if self.o1_flags.get(li).copied().unwrap_or(false) {
3589 if let Some(err) = self.kv_cache.layers[li].take_o1_error() {
3590 self.o1_fail(err);
3591 return;
3592 }
3593 }
3594 }
3595 self.o1_note_transition();
3599 if !self.o1_pending() {
3600 return;
3601 }
3602 if let Err(err) = self.o1_seal_checked() {
3603 self.o1_fail(err);
3604 }
3605 }
3606
3607 fn check_o1_progress_failure(&mut self, phase: &str) -> Result<(), String> {
3612 if self
3613 .graph_failed
3614 .swap(false, std::sync::atomic::Ordering::Relaxed)
3615 {
3616 self.cancel
3617 .store(false, std::sync::atomic::Ordering::Relaxed);
3618 self.clear_sequence_state();
3619 return Err(format!("{phase}: deferred O(1) transition failed"));
3620 }
3621 Ok(())
3622 }
3623
3624 pub fn o1_seal(&mut self) {
3628 if let Err(err) = self.o1_seal_checked() {
3629 self.o1_fail(err);
3630 }
3631 }
3632
3633 pub fn set_trace(&mut self, on: bool) {
3635 self.trace = on;
3636 }
3637
3638 pub fn set_sampler_config(&mut self, config: SamplerConfig) {
3641 self.rng = match config.seed {
3642 Some(seed) => SplitMix64::new(seed),
3643 None => SplitMix64::from_entropy(),
3644 };
3645 self.sampler_config = config;
3646 }
3647
3648 pub fn set_confidence(&mut self, on: bool) {
3653 self.confidence_on = on;
3654 }
3655
3656 pub fn set_calib_temp(&mut self, t: f32) {
3659 self.calib_temp = if t > 1e-3 { t } else { 1.0 };
3660 }
3661
3662 pub fn calib_temp(&self) -> f32 {
3664 self.calib_temp
3665 }
3666
3667 pub fn set_rotary(&mut self, rotary_dim: usize, base: f32) {
3670 self.rotary_dim = rotary_dim.min(self.head_dim);
3671 self.inv_freq = std::sync::Arc::new(attention::rope_inv_freq(self.rotary_dim, base));
3672 self.embryo_graph = None;
3676 }
3677
3678 fn attn_cfg(&self, position: usize) -> QwenAttnCfg<'_> {
3679 QwenAttnCfg {
3680 num_heads: self.num_heads,
3681 num_kv_heads: self.num_kv_heads,
3682 head_dim: self.head_dim,
3683 hidden_size: self.hidden_size,
3684 position,
3685 inv_freq: &self.inv_freq,
3686 rotary_dim: self.rotary_dim,
3687 scale: self.attn_scale,
3688 softcap: self.attn_softcap,
3689 window: None,
3690 v_norm: false,
3691 qk_norm_after_rope: self.qk_norm_after_rope,
3692 gate_sigmoid: self.proj_gate_sigmoid,
3693 q_norm: None,
3694 k_norm: None,
3695 output_gate: false,
3696 softplus_gate: None,
3697 rope_scale: self.rope_scale,
3698 bias: None,
3699 rms_eps: self.rms_eps,
3700 norm_style: self.norm_style,
3701 pool: self.pool.as_deref(),
3702 v_head_dim: self.v_head_dim.unwrap_or(self.head_dim),
3703 }
3704 }
3705
3706 pub fn generate(
3708 &mut self,
3709 prompt: &str,
3710 max_tokens: usize,
3711 task_mask: Option<&TaskMask>,
3712 on_token: Option<TokenCallback>,
3713 ) -> Result<GenerateResult, String> {
3714 let input_ids = self.tokenizer.with_bos(self.tokenizer.encode(prompt));
3715 self.generate_from_ids(&input_ids, max_tokens, task_mask, on_token)
3716 }
3717
3718 pub fn generate_from_vl(
3721 &mut self,
3722 input: &crate::dsv41_vision::PreparedVlInputs,
3723 max_tokens: usize,
3724 task_mask: Option<&TaskMask>,
3725 on_token: Option<TokenCallback>,
3726 ) -> Result<GenerateResult, String> {
3727 let Some(dsv41) = &self.dsv41 else {
3728 return Err("V4.1 multimodal input requires a DeepSeek-V4.1 pipeline".into());
3729 };
3730 if input.token_ids.is_empty() {
3731 return Err("empty V4.1 multimodal prompt".into());
3732 }
3733 if input.token_types.len() != input.token_ids.len() {
3734 return Err(format!(
3735 "V4.1 token type count {} != token count {}",
3736 input.token_types.len(),
3737 input.token_ids.len()
3738 ));
3739 }
3740 let dim = dsv41.2.dim;
3741 let mut embeddings = vec![None; input.token_ids.len()];
3742 let mut participates = vec![true; input.token_ids.len()];
3743 if !input.images.is_empty() {
3744 let vision = self
3745 .dsv41_vision
3746 .as_ref()
3747 .ok_or_else(|| "V4.1 image prompt has no loaded vision tower".to_string())?;
3748 for image in &input.images {
3749 let end = image.start.saturating_add(image.types.len());
3750 if end > input.token_ids.len() {
3751 return Err(format!(
3752 "V4.1 image span {}..{} exceeds prompt length {}",
3753 image.start,
3754 end,
3755 input.token_ids.len()
3756 ));
3757 }
3758 let mut span = vec![0.0f32; image.types.len() * dim];
3759 vision.fill_image_span(image, &mut span, self.pool.as_deref())?;
3760 for (offset, &kind) in image.types.iter().enumerate() {
3761 let pos = image.start + offset;
3762 if input.token_types[pos] != kind {
3763 return Err(format!(
3764 "V4.1 image type mismatch at position {pos}: {} != {kind}",
3765 input.token_types[pos]
3766 ));
3767 }
3768 embeddings[pos] = Some(span[offset * dim..(offset + 1) * dim].to_vec());
3769 participates[pos] = false;
3770 }
3771 }
3772 }
3773 for (pos, &kind) in input.token_types.iter().enumerate() {
3774 if kind == crate::dsv41_vision::TEXT && embeddings[pos].is_some() {
3775 return Err(format!("V4.1 text position {pos} has an image embedding"));
3776 }
3777 if kind != crate::dsv41_vision::TEXT && embeddings[pos].is_none() {
3778 return Err(format!("V4.1 image position {pos} has no image embedding"));
3779 }
3780 }
3781 self.dsv41_prefill = Some((embeddings, participates));
3782 let result = self.generate_from_ids(&input.token_ids, max_tokens, task_mask, on_token);
3783 self.dsv41_prefill = None;
3784 result
3785 }
3786
3787 fn drop_open_mask<'m>(&self, m: Option<&'m TaskMask>) -> Option<&'m TaskMask> {
3789 m.filter(|m| !m.fully_open(self.intermediate_size, self.num_heads))
3790 }
3791
3792 pub fn generate_from_ids(
3800 &mut self,
3801 input_ids: &[u32],
3802 max_tokens: usize,
3803 task_mask: Option<&TaskMask>,
3804 on_token: Option<TokenCallback>,
3805 ) -> Result<GenerateResult, String> {
3806 self.generate_with_prompt_rows(input_ids, None, max_tokens, task_mask, on_token)
3807 }
3808
3809 pub fn generate_from_embeds(
3815 &mut self,
3816 input_ids: &[u32],
3817 prompt_rows: &[f32],
3818 max_tokens: usize,
3819 task_mask: Option<&TaskMask>,
3820 on_token: Option<TokenCallback>,
3821 ) -> Result<GenerateResult, String> {
3822 if input_ids.is_empty()
3823 || input_ids.len().checked_mul(self.hidden_size) != Some(prompt_rows.len())
3824 {
3825 return Err("embedded prompt dimensions must be [tokens, hidden_size]".into());
3826 }
3827 if prompt_rows.iter().any(|x| !x.is_finite()) {
3828 return Err("embedded prompt contains non-finite values".into());
3829 }
3830 if !self.can_prefill_batched() || self.dyn_router.is_some()
3831 || self.o1_active() || self.mtp.is_some() || self.gpu_plan.is_some()
3832 {
3833 return Err("embedded prompts require the ordinary transformer path without O(1), dynamic routing, GPU splitting or a generic MTP head".into());
3834 }
3835 self.generate_with_prompt_rows(input_ids, Some(prompt_rows), max_tokens, task_mask, on_token)
3836 }
3837
3838 fn generate_with_prompt_rows(
3839 &mut self,
3840 input_ids: &[u32],
3841 prompt_rows: Option<&[f32]>,
3842 max_tokens: usize,
3843 task_mask: Option<&TaskMask>,
3844 mut on_token: Option<TokenCallback>,
3845 ) -> Result<GenerateResult, String> {
3846 #[cfg(target_os = "macos")]
3847 crate::gpu_metal::set_io_namespace(self.graph_kv_id);
3848 if std::env::var("CMF_TRACE_H").is_ok() {
3849 eprintln!("input_ids: {input_ids:?}");
3850 }
3851 if input_ids.is_empty() {
3852 return Err("empty prompt: nothing to generate from".to_string());
3853 }
3854 self.graph_failed
3858 .store(false, std::sync::atomic::Ordering::Relaxed);
3859 let task_mask = self.drop_open_mask(task_mask);
3864
3865 let mut reuse_from = {
3873 let on = !std::env::var("CMF_KV_REUSE").is_ok_and(|v| v == "0");
3874 if on
3875 && prompt_rows.is_none()
3876 && task_mask.is_none()
3877 && self.mtp.is_none()
3878 && !(self.mimo_mtp.is_some() && self.speculative)
3879 && self.o1_cfg.is_none()
3880 && self.dsv41.is_none()
3881 {
3882 self.cached_prefix_len(input_ids)
3883 } else {
3884 0
3885 }
3886 };
3887 if reuse_from > 0 && !self.prepare_kv_reuse(reuse_from) {
3890 reuse_from = 0;
3891 }
3892 self.last_prefill_tokens = input_ids.len() - reuse_from;
3893 let bounded_native = self.bounded_native();
3894 if reuse_from == 0 {
3895 self.clear_sequence_state();
3897 } else if std::env::var("CMF_PREFILL_PROF").is_ok() {
3898 eprintln!(
3899 "kv-reuse: {} of {} prompt positions already cached",
3900 reuse_from,
3901 input_ids.len()
3902 );
3903 }
3904 crate::gpu::graph_race_begin_generation();
3905 let o1_prefill = if self.o1_active() && task_mask.is_none() {
3909 std::env::var("CMF_O1_PREFILL")
3910 .ok()
3911 .and_then(|v| v.parse::<usize>().ok())
3912 .filter(|&p| p > 0)
3913 } else {
3914 None
3915 };
3916 if task_mask.is_none() {
3917 self.o1_begin_with_prefix(o1_prefill);
3918 }
3919
3920 let graph_on = crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Decode);
3926 #[cfg(target_os = "macos")]
3965 let metal_graph = crate::gpu::q1_force()
3966 && crate::gpu::enabled_here()
3967 && std::env::var("CMF_GPU_BLOCK")
3968 .map(|v| v != "0")
3969 .unwrap_or(true);
3970 #[cfg(not(target_os = "macos"))]
3971 let metal_graph = false;
3972 let spec_sample_env = std::env::var("CMF_GRAPH_SPEC_SAMPLE").ok();
3973 let spec_cheap_round = self.sampler_config.temperature < 1e-6
3977 || sampler::sparse_ok(&self.sampler_config);
3978 let spec_sampling_ok = self.sampler_config.temperature < 1e-6
3979 || match spec_sample_env.as_deref() {
3980 Some("1") => true,
3981 Some(_) => false,
3982 None => metal_graph && spec_cheap_round,
3983 };
3984 let (mut dense_n, mut dense_q4tp) = (0usize, 0usize);
4000 for lw in &self.weights.layers {
4001 if let FfnKind::Dense(d) = &lw.ffn {
4002 dense_n += 1;
4003 if matches!(d.gate_proj.graph_weight(), Some((_, _, 6, _)))
4004 && matches!(d.up_proj.graph_weight(), Some((_, _, 6, _)))
4005 && matches!(d.down_proj.graph_weight(), Some((_, _, 6, _)))
4006 {
4007 dense_q4tp += 1;
4008 }
4009 }
4010 }
4011 let spec_default_ok = dense_n == 0 || dense_q4tp * 10 >= dense_n * 9;
4012 let penalized = !metal_graph
4024 && (self.sampler_config.repetition_penalty != 1.0
4025 || self.sampler_config.presence_penalty != 0.0
4026 || !self.sampler_config.suppress_tokens.is_empty());
4027 #[cfg(feature = "gpu")]
4032 let metal_wgpu = graph_on && crate::gpu_wgpu::wgpu_backend_is_metal();
4033 #[cfg(not(feature = "gpu"))]
4034 let metal_wgpu = false;
4035 let spec_env = std::env::var("CMF_GRAPH_SPEC").ok();
4036 let spec_wanted = match spec_env.as_deref() {
4037 Some("0") => false,
4038 Some(_) => {
4039 if metal_wgpu {
4040 tracing::warn!(
4041 "CMF_GRAPH_SPEC forced on wgpu/Metal: the batched verify graph is not \
4042 verified on this backend (garbage measured on Qwen3.5-0.8B)"
4043 );
4044 }
4045 true
4046 }
4047 None => spec_default_ok && !penalized && !metal_wgpu,
4048 };
4049 let graph_spec = self.speculative
4053 && (graph_on || metal_graph)
4054 && self.mtp.is_some()
4055 && task_mask.is_none()
4056 && !self.o1_active()
4057 && spec_sampling_ok
4058 && spec_wanted;
4059 #[cfg(target_os = "macos")]
4063 if metal_graph {
4064 static SAID: std::sync::Once = std::sync::Once::new();
4065 SAID.call_once(|| {
4066 let spec = if graph_spec {
4067 let k = std::env::var("CMF_GRAPH_SPEC_K")
4068 .ok()
4069 .and_then(|v| v.parse::<usize>().ok())
4070 .filter(|&v| (1..=8).contains(&v))
4071 .unwrap_or(7);
4072 let arm = if self.sampler_config.temperature < 1e-6 {
4073 "greedy"
4074 } else {
4075 "sampling"
4076 };
4077 format!(
4078 "spec k={k} {arm} (batched verify, draft shortlist {}, trial: proxy)",
4079 Self::draft_vocab_rows(usize::MAX)
4080 )
4081 } else if !self.speculative {
4082 "spec off (CMF_MTP=0)".to_string()
4083 } else if self.mtp.is_none() {
4084 "spec off (no MTP head)".to_string()
4085 } else if !spec_sampling_ok {
4086 if spec_cheap_round {
4087 "spec off (CMF_GRAPH_SPEC_SAMPLE=0)".to_string()
4088 } else {
4089 "spec off (sampling without a top-k: the dense chain \
4090 costs more than it saves)"
4091 .to_string()
4092 }
4093 } else if !spec_wanted {
4094 "spec off (CMF_GRAPH_SPEC=0 or non-q4tp FFNs)".to_string()
4095 } else if task_mask.is_some() {
4096 "spec off (task mask)".to_string()
4097 } else {
4098 "spec off (O(1) attention)".to_string()
4099 };
4100 let on = |var: &str| {
4101 if std::env::var(var).as_deref() == Ok("0") {
4102 "off"
4103 } else {
4104 "on"
4105 }
4106 };
4107 tracing::info!(
4108 "metal native: {spec}, state4 {}, async replay {}, prefill graph {}, \
4109 MTP graph {}, attend {}, probe {}",
4110 if crate::gpu_metal::state4_on() { "on" } else { "off" },
4111 if crate::gpu_metal::async_replay_on() { "on" } else { "off" },
4112 on("CMF_METAL_PREFILL"),
4113 on("CMF_MTP_GRAPH"),
4114 std::env::var("CMF_GPU_ATTEND").unwrap_or_else(|_| "auto".into()),
4115 if crate::gpu::probe_enabled() { "bypassed (q1 force)" } else { "off" },
4116 );
4117 });
4118 }
4119 let pair_pays = self.gdn_cfg.is_none() || std::env::var("CMF_MTP").as_deref() == Ok("1");
4126 let spec_active = self.speculative
4127 && self.mtp.is_some()
4128 && task_mask.is_none()
4129 && !self.o1_active()
4130 && ((!graph_on && pair_pays && self.sampler_config.temperature < 1e-6) || graph_spec);
4131 let mut mtp = if spec_active { self.mtp.take() } else { None };
4134 if std::env::var("CMF_MTP_CHAIN_PROBE").is_ok() {
4135 eprintln!(
4136 "mtp-probe gate: spec_active={spec_active} mtp={} speculative={} graph_on={graph_on} temp_ok={}",
4137 mtp.is_some(),
4138 self.speculative,
4139 self.sampler_config.temperature < 1e-6,
4140 );
4141 }
4142 if let Some(m) = &mut mtp {
4143 m.kv.clear();
4144 crate::gpu::graph_kv_reset(self.mtp_kv_id());
4146 self.mtp_graph_mode = None;
4147 }
4148 let mimo_spec = self.speculative
4152 && self.mimo_mtp.is_some()
4153 && task_mask.is_none()
4154 && !self.o1_active()
4155 && self.dyn_router.is_none()
4156 && self.sampler_config.temperature < 1e-6
4157 && std::env::var("CMF_MIMO_MTP").as_deref() != Ok("0");
4158 if let Some(st) = self.mimo_mtp.as_mut() {
4159 st.reset();
4160 if mimo_spec && std::env::var_os("CMF_MIMO_MTP_PROBE").is_some() {
4161 Self::mimo_mtp_hist_cap(st, input_ids.len());
4162 }
4163 }
4164 let mut router = if mtp.is_none() {
4168 self.dyn_router.take()
4169 } else {
4170 None
4171 };
4172 let mut reuse_from = reuse_from;
4173 if let Some(r) = &mut router {
4174 r.reset(); self.dyn_phi_seen = 0; if self.dyn_active.is_some() {
4177 let _ = self.set_active_skill(None);
4180 reuse_from = 0;
4181 self.last_prefill_tokens = input_ids.len();
4182 }
4183 }
4184
4185 let mut all_ids = input_ids.to_vec();
4186 let mut generated = 0usize;
4187 let mut finish_reason = "max_tokens".to_string();
4188 let mut drafted = 0usize;
4189 let mut accepted = 0usize;
4190 let mut dsv4_spec_bad = 0usize;
4197 let mut dsv4_spec_retry_at = 0usize;
4198 let mut confidence: Vec<f32> = Vec::new();
4199 let trace_on = self.trace;
4200 let calib_temp = self.calib_temp;
4201 let mut traces: Vec<TokenTrace> = Vec::new();
4202
4203 let mut hidden = vec![0.0f32; self.hidden_size];
4209 let mut pos = reuse_from;
4210 let fuse_lm = mtp.is_none()
4219 && router.is_none()
4220 && std::env::var("CMF_GPU_LMHEAD").as_deref() != Ok("0");
4221 self.graph_logits = None;
4222 self.graph_want_logits = false;
4223 let _tpf = std::time::Instant::now();
4224 let batch_k = self.generation_batch_k();
4225 if let Some(rows) = prompt_rows {
4226 let hs = self.hidden_size;
4227 let chunk = self.prefill_chunk().max(1);
4228 while pos < input_ids.len() && !self.cancel.load(std::sync::atomic::Ordering::Relaxed) {
4229 let end = (pos + chunk).min(input_ids.len());
4230 let hb = match self.prefill_input_rows(
4231 PrefillIn::Hidden(&rows[pos * hs..end * hs]), pos, task_mask,
4232 ) {
4233 Ok(hb) => hb,
4234 Err(err) => {
4235 self.finish_generation(&mut mtp, &mut router, true);
4236 return Err(err);
4237 }
4238 };
4239 if mimo_spec { self.mimo_note_rows(&hb, pos); }
4240 hidden.copy_from_slice(&hb[hb.len() - hs..]);
4241 pos = end;
4242 }
4243 }
4244 while self.qwen4_exp.is_some()
4255 && mtp.is_none()
4256 && pos < input_ids.len()
4257 && !self.cancel.load(std::sync::atomic::Ordering::Relaxed)
4258 {
4259 let end = (pos + crate::qwen4_exp::prefill_chunk()).min(input_ids.len());
4262 let want_logits = end == input_ids.len();
4263 let mut lg = Vec::new();
4264 if let Some(b) = &mut self.qwen4_exp {
4265 crate::qwen4_exp::forward_tokens(
4266 &b.0,
4267 &b.1,
4268 &b.2,
4269 &mut b.3,
4270 &input_ids[pos..end],
4271 pos,
4272 &self.inv_freq,
4273 self.pool.as_deref(),
4274 &mut lg,
4275 want_logits,
4276 );
4277 }
4278 if want_logits {
4279 self.graph_logits = Some(lg);
4280 }
4281 pos = end;
4282 hidden.fill(0.0);
4283 }
4284 while self.dsv4.is_some()
4285 && mtp.is_none()
4286 && pos < input_ids.len()
4287 && !self.cancel.load(std::sync::atomic::Ordering::Relaxed)
4288 {
4289 let end = (pos + prefill_chunk()).min(input_ids.len());
4290 let ids: Vec<u32> = input_ids[pos..end].to_vec();
4291 let mut lg = Vec::new();
4292 if let Some(b) = &mut self.dsv4 {
4293 let (g, layers, cfg, st) = (&b.0, &b.1, b.2, &mut b.3);
4294 crate::dsv4::forward_chunk(
4295 g,
4296 layers,
4297 &cfg,
4298 st,
4299 &ids,
4300 pos,
4301 &self.inv_freq,
4302 self.pool.as_deref(),
4303 &mut lg,
4304 end == input_ids.len(),
4305 );
4306 }
4307 if end == input_ids.len() {
4308 self.graph_logits = Some(lg);
4309 }
4310 pos = end;
4311 hidden = vec![0.0; self.hidden_size];
4312 }
4313 let dsv41_prefill = self.dsv41_prefill.take();
4314 while self.dsv41.is_some()
4315 && mtp.is_none()
4316 && pos < input_ids.len()
4317 && !self.cancel.load(std::sync::atomic::Ordering::Relaxed)
4318 {
4319 let end = (pos + prefill_chunk()).min(input_ids.len());
4320 let ids: Vec<u32> = input_ids[pos..end].to_vec();
4321 let mut lg = Vec::new();
4322 if let Some(b) = &mut self.dsv41 {
4323 let (g, layers, cfg, st) = (&b.0, &b.1, &b.2, &mut b.3);
4324 if let Some((embeddings, participates)) = dsv41_prefill.as_ref() {
4325 crate::dsv41::forward_chunk_masked_with_embeddings(
4326 g,
4327 layers,
4328 cfg,
4329 st,
4330 &ids,
4331 pos,
4332 &embeddings[pos..end],
4333 &participates[pos..end],
4334 self.pool.as_deref(),
4335 &mut lg,
4336 );
4337 } else {
4338 crate::dsv41::forward_chunk(
4339 g,
4340 layers,
4341 cfg,
4342 st,
4343 &ids,
4344 pos,
4345 self.pool.as_deref(),
4346 &mut lg,
4347 );
4348 }
4349 }
4350 if end == input_ids.len() {
4351 self.graph_logits = Some(lg);
4352 }
4353 pos = end;
4354 hidden = vec![0.0; self.hidden_size];
4355 }
4356 let dyn_prefill = router.is_some();
4361 let o1_prefill_limit = o1_prefill
4369 .and_then(|requested| self.o1_effective_boundary(requested))
4370 .map(|boundary| boundary.min(input_ids.len()));
4371 let mut o1_sealed = false;
4372 if let Some(limit) = o1_prefill_limit {
4373 if self.can_prefill_batched() && limit > 2 {
4376 let chunk = self.prefill_chunk();
4377 let hs = self.hidden_size;
4378 while pos < limit && !self.cancel.load(std::sync::atomic::Ordering::Relaxed) {
4379 let end = (pos + chunk).min(limit);
4380 let hb = self.prefill_batch(&input_ids[pos..end], pos);
4381 hidden.copy_from_slice(&hb[(end - pos - 1) * hs..]);
4382 pos = end;
4383 }
4384 } else {
4385 while pos < limit && !self.cancel.load(std::sync::atomic::Ordering::Relaxed) {
4386 hidden = self.forward_layers(&self.embed_single(input_ids[pos]), pos, None);
4387 pos += 1;
4388 }
4389 }
4390 if pos >= limit {
4391 o1_sealed = match self.o1_seal_checked() {
4392 Ok(sealed) => sealed,
4393 Err(err) => {
4394 self.finish_generation(&mut mtp, &mut router, true);
4395 return Err(err);
4396 }
4397 };
4398 tracing::info!(
4399 "o1 bounded prompt prefix: requested={} effective={} processed={} of {} token(s)",
4400 o1_prefill.unwrap_or(0),
4401 self.o1_effective_boundary(o1_prefill.unwrap_or(0))
4402 .unwrap_or(limit),
4403 limit,
4404 input_ids.len()
4405 );
4406 }
4407 }
4408 let graph_prefill = self.graph_prefill_preferred();
4414 #[cfg(target_os = "macos")]
4422 if task_mask.is_none()
4423 && !dyn_prefill
4424 && (crate::gpu::q1_force() || crate::gpu::q2tp_gpu_opt_in())
4425 && crate::gpu::enabled_here()
4426 && self.gdn_cfg.is_some()
4427 && self.g3n.is_none()
4428 && input_ids.len() > 8
4429 && std::env::var("CMF_MTP_CHAIN_PROBE").is_err()
4430 && std::env::var("CMF_METAL_PREFILL").as_deref() != Ok("0")
4431 {
4432 let chunk: usize = std::env::var("CMF_METAL_PREFILL_CHUNK")
4433 .ok()
4434 .and_then(|v| v.parse().ok())
4435 .filter(|&v| (16..=512).contains(&v))
4436 .unwrap_or(256);
4437 let hs = self.hidden_size;
4438 let _tp = std::time::Instant::now();
4439 while pos < input_ids.len() && !self.cancel.load(std::sync::atomic::Ordering::Relaxed) {
4440 let end = (pos + chunk).min(input_ids.len());
4441 let hb = match self.prefill_batch_metal(&input_ids[pos..end], pos) {
4442 MetalPrefillOutcome::Completed(hb) => hb,
4443 MetalPrefillOutcome::Declined => break,
4444 MetalPrefillOutcome::Failed => {
4445 self.finish_generation(&mut mtp, &mut router, true);
4446 return Err("ordinary Metal prefill failed after admission".into());
4447 }
4448 };
4449 if let Some(m) = &mut mtp {
4450 let n_pairs = if end < input_ids.len() {
4451 end - pos
4452 } else {
4453 end - pos - 1
4454 };
4455 if n_pairs > 0 {
4456 let pairs: Vec<(&[f32], u32)> = (0..n_pairs)
4457 .map(|j| (&hb[j * hs..(j + 1) * hs], input_ids[pos + j + 1]))
4458 .collect();
4459 if !self.mtp_warm_batch_metal(m, &pairs, pos) {
4460 for (j, (h, t)) in pairs.iter().enumerate() {
4461 let h = h.to_vec();
4462 let _ = self.mtp_step(m, &h, *t, pos + j);
4463 }
4464 }
4465 }
4466 }
4467 hidden.copy_from_slice(&hb[(end - pos - 1) * hs..]);
4468 pos = end;
4469 }
4470 if std::env::var("CMF_PREFILL_PROF").is_ok() {
4471 eprintln!(
4472 "metal-prefill: {} of {} tokens in {:.1} ms",
4473 pos,
4474 input_ids.len(),
4475 _tp.elapsed().as_secs_f64() * 1e3
4476 );
4477 }
4478 }
4479 self.mimo_moe_prepare();
4480 #[cfg(not(target_os = "macos"))]
4486 if task_mask.is_none()
4487 && !dyn_prefill
4488 && !graph_prefill
4489 && mtp.is_none()
4490 && o1_prefill.is_none()
4491 && !self.o1_active()
4492 && input_ids.len() > 2
4493 && self.batch_prefix_prefill()
4494 {
4495 let chunk = self.prefill_chunk().max(1);
4496 let hs = self.hidden_size;
4497 let t_bp = std::time::Instant::now();
4498 let pos0 = pos;
4499 while pos < input_ids.len() && !self.cancel.load(std::sync::atomic::Ordering::Relaxed) {
4500 let end = (pos + chunk).min(input_ids.len());
4501 let bk = end - pos;
4502 let mut hiddens = vec![0f32; bk * hs];
4503 for (j, &id) in input_ids[pos..end].iter().enumerate() {
4504 hiddens[j * hs..(j + 1) * hs].copy_from_slice(&self.embed_single(id));
4505 }
4506 let positions: Vec<usize> = (pos..end).collect();
4507 let mut run = 0usize;
4508 let outcome = self.try_batch_graph_wgpu_prefix(
4509 &mut hiddens,
4510 &positions,
4511 bk,
4512 None,
4513 Some(&mut run),
4514 );
4515 match outcome {
4516 crate::gpu::BatchGraphOutcome::Completed => {
4517 let hb = if run < self.num_layers {
4518 self.prefill_batch_span(
4519 PrefillIn::Hidden(&hiddens),
4520 pos,
4521 None,
4522 run,
4523 self.num_layers,
4524 )
4525 } else {
4526 hiddens
4527 };
4528 if mimo_spec {
4529 self.mimo_note_rows(&hb, pos);
4530 }
4531 hidden.copy_from_slice(&hb[(bk - 1) * hs..]);
4532 pos = end;
4533 }
4534 crate::gpu::BatchGraphOutcome::Failed => {
4535 self.finish_generation(&mut mtp, &mut router, true);
4536 return Err("batched prefix prefill failed after admission".into());
4537 }
4538 crate::gpu::BatchGraphOutcome::Declined => {
4539 #[cfg(feature = "gpu")]
4542 if pos > pos0 {
4543 self.pull_lagging_host_kv(0, self.num_layers, pos);
4544 }
4545 break;
4546 }
4547 }
4548 }
4549 if std::env::var("CMF_PREFILL_PROF").is_ok() {
4550 eprintln!(
4551 "batch-prefix prefill: {} of {} tokens in {:.1} ms",
4552 pos - pos0,
4553 input_ids.len(),
4554 t_bp.elapsed().as_secs_f64() * 1e3
4555 );
4556 }
4557 }
4558 if task_mask.is_none()
4559 && !dyn_prefill
4560 && !graph_prefill
4561 && self.can_prefill_batched()
4562 && self.g3n.is_none()
4563 && o1_prefill.is_none()
4564 && input_ids.len() > 2
4565 {
4566 let chunk = self.prefill_chunk();
4572 let hs = self.hidden_size;
4573 while pos < input_ids.len() && !self.cancel.load(std::sync::atomic::Ordering::Relaxed) {
4574 let end = (pos + chunk).min(input_ids.len());
4575 let hb = self.prefill_batch(&input_ids[pos..end], pos);
4576 if mimo_spec {
4577 self.mimo_note_rows(&hb, pos);
4578 }
4579 if let Some(m) = &mut mtp {
4580 let probe: usize = std::env::var("CMF_MTP_CHAIN_PROBE")
4581 .ok()
4582 .and_then(|v| v.parse().ok())
4583 .unwrap_or(0);
4584 for p in pos..end {
4585 if p + 1 < input_ids.len() {
4586 if probe >= 1 && p + 2 < input_ids.len() {
4587 let (d1, mut hx) = self.mtp_step_h(
4591 m,
4592 &hb[(p - pos) * hs..(p - pos + 1) * hs],
4593 input_ids[p + 1],
4594 p,
4595 );
4596 let mut ok = d1 == input_ids[p + 2];
4597 Self::chain_probe_note(0, ok);
4598 let mut d_prev = d1;
4599 let mut extra = 0usize;
4600 for j in 1..probe {
4601 if p + 2 + j >= input_ids.len() {
4602 break;
4603 }
4604 let (dj, hj) = self.mtp_step_h(m, &hx, d_prev, p + 1 + j);
4605 extra += 1;
4606 ok = ok && dj == input_ids[p + 2 + j];
4607 Self::chain_probe_note(j, ok);
4608 d_prev = dj;
4609 hx = hj;
4610 }
4611 m.kv.truncate_last(extra);
4612 } else {
4613 let _ = self.mtp_step(
4614 m,
4615 &hb[(p - pos) * hs..(p - pos + 1) * hs],
4616 input_ids[p + 1],
4617 p,
4618 );
4619 }
4620 }
4621 }
4622 }
4623 hidden.copy_from_slice(&hb[(end - pos - 1) * hs..]);
4624 pos = end;
4625 }
4626 }
4627 let pair_off = std::env::var("CMF_PAIR").is_ok_and(|v| v == "0");
4628 if task_mask.is_none()
4629 && !dyn_prefill
4630 && !graph_prefill
4631 && !pair_off
4632 && self.pair_supported()
4633 && o1_prefill.is_none()
4634 {
4635 while pos + 1 < input_ids.len()
4636 && !self.cancel.load(std::sync::atomic::Ordering::Relaxed)
4637 {
4638 let e1 = self.embed_single(input_ids[pos]);
4639 let e2 = self.embed_single(input_ids[pos + 1]);
4640 let (h1, h2) = self.forward_pair(&e1, &e2, pos);
4641 if mimo_spec {
4642 self.mimo_note_rows(&h1, pos);
4643 self.mimo_note_rows(&h2, pos + 1);
4644 }
4645 self.commit_linear_scratch();
4647 if let Some(m) = &mut mtp {
4648 let _ = self.mtp_step(m, &h1, input_ids[pos + 1], pos);
4649 if pos + 2 < input_ids.len() {
4650 let probe: usize = std::env::var("CMF_MTP_CHAIN_PROBE")
4651 .ok()
4652 .and_then(|v| v.parse().ok())
4653 .unwrap_or(0);
4654 if probe >= 1 && pos + 3 < input_ids.len() {
4655 let (d1, mut hx) = self.mtp_step_h(m, &h2, input_ids[pos + 2], pos + 1);
4659 let mut ok = d1 == input_ids[pos + 3];
4660 Self::chain_probe_note(0, ok);
4661 let mut d_prev = d1;
4662 let mut extra = 0usize;
4663 for j in 1..probe {
4664 if pos + 3 + j >= input_ids.len() {
4665 break;
4666 }
4667 let (dj, hj) = self.mtp_step_h(m, &hx, d_prev, pos + 2 + j);
4668 extra += 1;
4669 ok = ok && dj == input_ids[pos + 3 + j];
4670 Self::chain_probe_note(j, ok);
4671 d_prev = dj;
4672 hx = hj;
4673 }
4674 m.kv.truncate_last(extra);
4675 } else {
4676 let _ = self.mtp_step(m, &h2, input_ids[pos + 2], pos + 1);
4677 }
4678 }
4679 }
4680 hidden = h2;
4681 pos += 2;
4682 }
4683 }
4684 let o1_batch_ready = o1_sealed
4697 && o1_prefill.is_some()
4698 && mtp.is_none()
4699 && std::env::var("CMF_O1_GPU").as_deref() == Ok("1")
4700 && (0..self.num_layers).all(|li| {
4701 let cache = &self.kv_cache.layers[self.phys_layer(li)];
4702 cache.o1.is_none() || cache.o1_views().is_some()
4703 });
4704 let mtp_batch_prefill = mtp.is_some()
4709 && graph_prefill
4710 && task_mask.is_none()
4711 && !dyn_prefill
4712 && !self.o1_active()
4713 && std::env::var("CMF_MTP_CHAIN_PROBE").is_err();
4714 if batch_k > 0
4715 && (graph_prefill || o1_batch_ready)
4716 && task_mask.is_none()
4717 && (!self.o1_active() || o1_batch_ready)
4718 && (mtp.is_none() || mtp_batch_prefill)
4719 && !dyn_prefill
4720 && pos + 1 < input_ids.len()
4721 {
4722 let hs = self.hidden_size;
4723 let chunk = batch_k;
4724 while pos < input_ids.len() {
4725 let end = (pos + chunk).min(input_ids.len());
4726 let bk = end - pos;
4727 let mut hiddens = vec![0f32; bk * hs];
4728 for (j, &id) in input_ids[pos..end].iter().enumerate() {
4729 hiddens[j * hs..(j + 1) * hs].copy_from_slice(&self.embed_single(id));
4730 }
4731 let positions: Vec<usize> = (pos..end).collect();
4732 let t_chunk = std::time::Instant::now();
4733 let outcome = self.try_batch_graph_wgpu(&mut hiddens, &positions, bk, None);
4734 let ok_b = outcome == crate::gpu::BatchGraphOutcome::Completed;
4735 if std::env::var("CMF_GRAPH_PROF").is_ok() {
4736 let ms = t_chunk.elapsed().as_secs_f64() * 1000.0;
4737 eprintln!(
4738 "batch-chunk: phase=prompt mode={} k={bk} outcome={outcome:?} {ms:.1} ms ({:.1} tok/s)",
4739 if o1_batch_ready {
4740 "o1"
4741 } else if mtp_batch_prefill {
4742 "ordinary_mtp"
4743 } else {
4744 "ordinary"
4745 },
4746 bk as f64 / (ms / 1000.0)
4747 );
4748 }
4749 {
4750 use std::sync::atomic::{AtomicBool, Ordering};
4751 static SAID: AtomicBool = AtomicBool::new(false);
4752 if !SAID.swap(true, Ordering::Relaxed) {
4753 if ok_b {
4754 tracing::info!(
4755 "batched prefill: ACTIVE mode={} (k={bk})",
4756 if o1_batch_ready {
4757 "o1"
4758 } else if mtp_batch_prefill {
4759 "ordinary_mtp"
4760 } else {
4761 "ordinary"
4762 }
4763 );
4764 } else {
4765 tracing::warn!("batched prefill {:?} — per-position graph", outcome);
4766 }
4767 }
4768 }
4769 if ok_b {
4770 if mimo_spec {
4771 self.mimo_note_rows(&hiddens, pos);
4772 }
4773 if mtp_batch_prefill {
4774 let n_pairs = mtp_prefill_pair_count(pos, end, input_ids.len());
4775 if n_pairs > 0 {
4776 let rows: Vec<Vec<f32>> = (0..n_pairs)
4782 .map(|j| hiddens[j * hs..(j + 1) * hs].to_vec())
4783 .collect();
4784 let pairs: Vec<(&[f32], u32)> = rows
4785 .iter()
4786 .enumerate()
4787 .map(|(j, row)| (row.as_slice(), input_ids[pos + j + 1]))
4788 .collect();
4789 if std::env::var("CMF_GRAPH_PROF").is_ok() {
4790 eprintln!(
4791 "mtp-warm: phase=prompt mode=ordinary_mtp first_pos={} pairs={} last_pos={}",
4792 pos,
4793 n_pairs,
4794 pos + n_pairs - 1,
4795 );
4796 }
4797 let warm_error = if let Some(m) = mtp.as_mut() {
4798 self.mtp_warm_prefill_pairs(m, &pairs, pos).err()
4799 } else {
4800 None
4801 };
4802 if let Some(err) = warm_error {
4803 self.finish_generation(&mut mtp, &mut router, true);
4808 return Err(err.to_string());
4809 }
4810 }
4811 }
4812 hidden.copy_from_slice(&hiddens[(bk - 1) * hs..]);
4813 pos = end;
4814 } else if outcome == crate::gpu::BatchGraphOutcome::Failed {
4815 self.finish_generation(&mut mtp, &mut router, true);
4820 return Err(if o1_batch_ready {
4821 "sealed O(1) batch graph failed after admission".to_string()
4822 } else {
4823 "ordinary recurrent batch graph failed after admission".to_string()
4824 });
4825 } else {
4826 break; }
4828 }
4829 }
4830 if graph_prefill
4834 && task_mask.is_none()
4835 && mtp.is_none()
4836 && !dyn_prefill
4837 && pos == 0
4838 && input_ids.len() > 1
4839 && !self.cancel.load(std::sync::atomic::Ordering::Relaxed)
4840 {
4841 if let Some(lg) = self.embryo_prefill_chunked(input_ids, 0) {
4842 self.graph_logits = Some(lg);
4843 hidden = vec![0.0; self.hidden_size];
4844 pos = input_ids.len();
4845 }
4846 }
4847 while pos < input_ids.len() && !self.cancel.load(std::sync::atomic::Ordering::Relaxed) {
4848 self.graph_want_logits = fuse_lm && pos + 1 == input_ids.len();
4849 hidden = self.forward_layers(&self.embed_single(input_ids[pos]), pos, task_mask);
4850 if mimo_spec {
4851 self.mimo_note_rows(&hidden, pos);
4852 }
4853 if let Some(m) = &mut mtp {
4854 if pos + 1 < input_ids.len() {
4855 let probe: usize = std::env::var("CMF_MTP_CHAIN_PROBE")
4861 .ok()
4862 .and_then(|v| v.parse().ok())
4863 .unwrap_or(0);
4864 if probe >= 1 && pos + 2 < input_ids.len() {
4865 let (d1, mut hx) = self.mtp_step_h(m, &hidden, input_ids[pos + 1], pos);
4866 let mut ok = d1 == input_ids[pos + 2];
4867 Self::chain_probe_note(0, ok);
4868 let mut d_prev = d1;
4869 let mut extra = 0usize;
4870 for j in 1..probe {
4871 if pos + 2 + j >= input_ids.len() {
4872 break;
4873 }
4874 let (dj, hj) = self.mtp_step_h(m, &hx, d_prev, pos + 1 + j);
4875 extra += 1;
4876 ok = ok && dj == input_ids[pos + 2 + j];
4877 Self::chain_probe_note(j, ok);
4878 d_prev = dj;
4879 hx = hj;
4880 }
4881 m.kv.truncate_last(extra);
4884 } else {
4885 let _ = self.mtp_step(m, &hidden, input_ids[pos + 1], pos);
4886 }
4887 }
4888 }
4889 pos += 1;
4890 }
4891 if std::env::var("CMF_PREFILL_PROF").is_ok() {
4892 eprintln!(
4893 "prefill: {} tokens in {:.1} ms (batch_k={batch_k})",
4894 input_ids.len(),
4895 _tpf.elapsed().as_secs_f64() * 1000.0
4896 );
4897 }
4898 if self
4899 .graph_failed
4900 .swap(false, std::sync::atomic::Ordering::Relaxed)
4901 {
4902 self.finish_generation(&mut mtp, &mut router, true);
4907 return Err("GPU token graph failed during prefill".to_string());
4908 }
4909 if self
4912 .cancel
4913 .swap(false, std::sync::atomic::Ordering::Relaxed)
4914 {
4915 self.finish_generation(&mut mtp, &mut router, true);
4919 return Ok(GenerateResult {
4920 text: String::new(),
4921 token_ids: Vec::new(),
4922 prompt_tokens: input_ids.len(),
4923 tokens_generated: 0,
4924 finish_reason: "cancelled".to_string(),
4925 mtp_drafted: 0,
4926 mtp_accepted: 0,
4927 token_confidence: Vec::new(),
4928 traces: Vec::new(),
4929 });
4930 }
4931
4932 if !o1_sealed {
4935 match self.o1_seal_checked() {
4936 Ok(_) => {}
4937 Err(err) => {
4938 self.finish_generation(&mut mtp, &mut router, true);
4939 return Err(err);
4940 }
4941 }
4942 }
4943
4944 macro_rules! commit {
4946 ($id:expr) => {{
4947 all_ids.push($id);
4948 generated += 1;
4949 self.note_draft_id($id);
4950 if self.tokenizer.is_eos($id) && !self.ignore_eos {
4951 finish_reason = "stop".to_string();
4952 false
4953 } else {
4954 let token_text = self.tokenizer.decode_token($id);
4955 let mut go = true;
4956 if let Some(ref mut cb) = on_token {
4957 if !cb(&token_text) {
4958 finish_reason = "cancelled".to_string();
4959 go = false;
4960 }
4961 }
4962 go
4963 }
4964 }};
4965 }
4966
4967 let mut spec_trial = SpecTrial::Spec {
4978 t0: std::time::Instant::now(),
4979 gen0: generated,
4980 rounds: 0,
4981 };
4982 let mut spec_mon = SpecMon {
4988 metal: graph_spec && crate::gpu::q1_force() && spec_cheap_round,
4989 ..SpecMon::default()
4990 };
4991 let mut spec_watchdog_off = false;
4992 let mut spec_walls: Vec<f32> = Vec::new();
4995 let mut spec_round_end: Option<std::time::Instant> = None;
4998 if mimo_spec {
4999 if let Ok(path) = std::env::var("CMF_MIMO_MTP_PROBE") {
5000 if let Some(mut st) = self.mimo_mtp.take() {
5001 self.mimo_mtp_probe(&mut st, input_ids, &path);
5002 self.mimo_mtp = Some(st);
5003 }
5004 }
5005 }
5006 let mut next_pos = input_ids.len();
5008 'decode: while generated < max_tokens {
5009 if self
5010 .graph_failed
5011 .swap(false, std::sync::atomic::Ordering::Relaxed)
5012 {
5013 self.finish_generation(&mut mtp, &mut router, true);
5018 return Err("GPU token graph failed during decode".to_string());
5019 }
5020 if self
5021 .cancel
5022 .swap(false, std::sync::atomic::Ordering::Relaxed)
5023 {
5024 finish_reason = "cancelled".to_string();
5025 break 'decode;
5026 }
5027 if mimo_spec && next_pos > 0 {
5032 self.mimo_note_rows(&hidden, next_pos - 1);
5035 }
5036 let forced = self.spec_forced.take();
5037 let mut logits = match (forced, self.graph_logits.take()) {
5038 (Some(_), _) => Vec::new(),
5039 (None, Some(lg)) => lg,
5040 (None, None) => {
5041 let _prof = crate::cpuprof::time(crate::cpuprof::Slot::Head);
5042 inference::rms_norm_into(
5043 &hidden,
5044 &self.weights.final_norm,
5045 self.rms_eps,
5046 self.norm_style,
5047 &mut self.ws.n1,
5048 );
5049 self.lm_head_forward(&self.ws.n1)
5050 }
5051 };
5052 if generated
5055 == std::env::var("CMF_LOGIT_DUMP_STEP")
5056 .ok()
5057 .and_then(|v| v.parse().ok())
5058 .unwrap_or(0)
5059 {
5060 if let Ok(path) = std::env::var("CMF_LOGIT_DUMP") {
5061 let mut bytes: Vec<u8> = Vec::with_capacity((hidden.len() + logits.len()) * 4);
5062 for v in hidden.iter().chain(logits.iter()) {
5063 bytes.extend_from_slice(&v.to_le_bytes());
5064 }
5065 if let Err(e) = std::fs::write(&path, &bytes) {
5066 eprintln!("logit dump: failed to write {path}: {e}");
5067 self.finish_generation(&mut mtp, &mut router, true);
5068 return Err(format!("logit dump write failed: {e}"));
5069 }
5070 }
5071 }
5072 if let Ok(dir) = std::env::var("CMF_LOGIT_DUMP_ALL") {
5076 if !logits.is_empty() {
5077 let path = std::path::Path::new(&dir).join(format!("step{generated:05}.f32"));
5078 let bytes: Vec<u8> = logits.iter().flat_map(|v| v.to_le_bytes()).collect();
5079 if let Err(e) =
5080 std::fs::create_dir_all(&dir).and_then(|_| std::fs::write(&path, &bytes))
5081 {
5082 eprintln!("logit dump: failed to write {}: {e}", path.display());
5083 }
5084 }
5085 }
5086 let t_next = match forced {
5087 Some(c) => c,
5088 None => {
5089 let _prof = crate::cpuprof::time(crate::cpuprof::Slot::Sampler);
5090 sampler::sample_with_scratch_pool(
5091 &logits,
5092 &self.sampler_config,
5093 self.sampler_config.penalty_past(&all_ids, bounded_native),
5094 &mut self.rng,
5095 &mut self.sampler_scratch,
5096 self.pool.as_deref(),
5097 )
5098 }
5099 };
5100 if self.confidence_on {
5101 confidence.push(if logits.is_empty() {
5102 0.0
5103 } else {
5104 sampler::top1_prob_pool(
5105 self.pool.as_deref(),
5106 &mut self.sampler_scratch,
5107 &logits,
5108 t_next,
5109 calib_temp,
5110 )
5111 });
5112 }
5113 if !logits.is_empty() {
5114 attention::recycle_buf(&mut logits);
5115 }
5116 if trace_on {
5117 let skill = router.as_ref().and_then(|r| r.active_id());
5121 traces.push(TokenTrace {
5122 t: generated,
5123 token_id: t_next,
5124 confidence: confidence.last().copied().unwrap_or(0.0),
5125 active_skill: skill,
5126 recon: None,
5127 switched: false,
5128 });
5129 }
5130 if !commit!(t_next) {
5131 break 'decode;
5132 }
5133 if generated >= max_tokens {
5134 break 'decode;
5135 }
5136
5137 if self.dsv41.is_none() && self.kv_cache.needs_eviction() {
5138 static SAID: std::sync::Once = std::sync::Once::new();
5144 SAID.call_once(|| {
5145 tracing::warn!(
5146 "KV cache full at {} positions — evicting half; quality \
5147 will degrade. Raise CMF_MAX_SEQ.",
5148 self.kv_cache.max_seq_len,
5149 );
5150 });
5151 let keep = (self.kv_cache.max_seq_len / 2).max(1);
5152 self.kv_cache.evict(keep);
5153 }
5154
5155 if graph_spec {
5158 match spec_trial {
5159 SpecTrial::Plain { t0, gen0 } if spec_mon.plain_done(t0, gen0, generated) => {
5160 spec_mon.plain_ms =
5161 t0.elapsed().as_secs_f64() * 1e3 / (generated - gen0) as f64;
5162 let keep = spec_mon.pays();
5163 tracing::info!(
5164 "speculation trial: {:.2} tok/round in {:.1} ms vs plain {:.1} ms/tok — {}",
5165 spec_mon.tokens,
5166 spec_mon.round_ms,
5167 spec_mon.plain_ms,
5168 if keep { "speculating" } else { "plain" }
5169 );
5170 spec_mon.fails = 0;
5171 spec_trial = SpecTrial::Decided {
5172 spec: keep,
5173 recheck_at: if keep { usize::MAX } else { generated + 128 },
5174 };
5175 }
5176 SpecTrial::Decided { recheck_at, .. } if generated >= recheck_at => {
5177 spec_mon.n = 0;
5178 spec_trial = SpecTrial::Spec {
5179 t0: std::time::Instant::now(),
5180 gen0: generated,
5181 rounds: 0,
5182 };
5183 }
5184 _ => {}
5185 }
5186 spec_watchdog_off = matches!(
5187 spec_trial,
5188 SpecTrial::Plain { .. } | SpecTrial::Decided { spec: false, .. }
5189 );
5190 }
5191 if mimo_spec && generated + 1 < max_tokens && next_pos > 0 {
5193 let budget = max_tokens - generated - 1;
5194 if let Some(mut st) = self.mimo_mtp.take() {
5195 let k = st.depth.min(budget);
5196 let r = self.mimo_spec_round(&mut st, next_pos, &all_ids, k);
5197 self.mimo_mtp = Some(st);
5198 let r = match r {
5199 Ok(r) => r,
5200 Err(err) => {
5201 self.finish_generation(&mut mtp, &mut router, true);
5202 return Err(err);
5203 }
5204 };
5205 if let Some(r) = r {
5206 drafted += r.drafted;
5207 accepted += r.accepted.len();
5208 let mut stopped = false;
5209 for &id in &r.accepted {
5210 if self.confidence_on {
5211 confidence.push(0.0);
5212 }
5213 if !commit!(id) {
5214 stopped = true;
5215 break;
5216 }
5217 }
5218 if stopped {
5219 break 'decode;
5220 }
5221 next_pos += r.accepted.len() + 1;
5222 hidden = r.hidden;
5223 self.graph_logits = Some(r.logits);
5226 continue 'decode;
5227 }
5228 }
5229 }
5230 #[cfg(feature = "gpu")]
5233 if self.speculative
5234 && self.qwen4_exp.is_some()
5235 && task_mask.is_none()
5236 && self.sampler_config.temperature < 1e-6
5237 && generated + 1 < max_tokens
5238 && next_pos > 0
5239 && std::env::var("CMF_QWEN_MTP").as_deref() != Ok("0")
5240 {
5241 let r = match &mut self.qwen4_exp {
5242 Some(b) => crate::qwen4_exp::spec_round(
5243 &b.0,
5244 &b.1,
5245 &b.2,
5246 &mut b.3,
5247 next_pos,
5248 &all_ids,
5249 &self.inv_freq,
5250 self.pool.as_deref(),
5251 ),
5252 None => None,
5253 };
5254 if let Some(r) = r {
5255 drafted += r.drafted;
5256 accepted += r.accepted.len();
5257 let mut stopped = false;
5258 for &id in &r.accepted {
5259 if self.confidence_on {
5260 confidence.push(0.0);
5261 }
5262 if !commit!(id) {
5263 stopped = true;
5264 break;
5265 }
5266 }
5267 if stopped {
5268 break 'decode;
5269 }
5270 next_pos += r.accepted.len() + 1;
5271 hidden.fill(0.0);
5272 self.graph_logits = Some(r.logits);
5273 continue 'decode;
5274 }
5275 }
5276 match &mut mtp {
5277 #[cfg(feature = "gpu")]
5279 Some(m)
5280 if graph_spec
5281 && !spec_watchdog_off
5282 && generated + 1 < max_tokens
5283 && next_pos > 0 =>
5284 {
5285 let t_round = std::time::Instant::now();
5286 if spec_time_level() >= 2 {
5287 if let Some(t) = spec_round_end.take() {
5288 eprintln!(
5289 "spec-gap {:.2} ms (host between rounds)",
5290 t.elapsed().as_secs_f64() * 1e3
5291 );
5292 }
5293 }
5294 spec_stamps_begin();
5295 #[cfg(target_os = "macos")]
5300 let allocs0 = crate::gpu_metal::IO_BUF_ALLOCS
5301 .load(std::sync::atomic::Ordering::Relaxed);
5302 #[cfg(not(target_os = "macos"))]
5303 let allocs0 = 0u64;
5304 if let Some((extra, n_pos, new_h)) = self.graph_spec_step(
5305 m,
5306 &hidden,
5307 t_next,
5308 next_pos,
5309 &mut drafted,
5310 &mut accepted,
5311 &mut all_ids,
5312 max_tokens - generated,
5313 ) {
5314 next_pos = n_pos;
5315 hidden = new_h;
5316 let level = spec_time_level();
5317 if level > 0 {
5318 let wall = t_round.elapsed().as_secs_f32() * 1e3;
5319 let stamps = spec_stamps_take();
5320 let median = if spec_walls.len() >= 3 {
5323 let mut s = spec_walls.clone();
5324 s.sort_by(|a, b| a.partial_cmp(b).unwrap());
5325 Some(s[s.len() / 2])
5326 } else {
5327 None
5328 };
5329 let outlier = median.is_some_and(|m| wall > 1.4 * m);
5330 #[cfg(target_os = "macos")]
5331 let allocs = crate::gpu_metal::IO_BUF_ALLOCS
5332 .load(std::sync::atomic::Ordering::Relaxed)
5333 - allocs0;
5334 #[cfg(not(target_os = "macos"))]
5335 let allocs = allocs0;
5336 eprintln!(
5337 "spec-round wall {wall:.1} ms → {} tokens{}{}",
5338 extra.len() + 1,
5339 if allocs > 0 {
5340 format!(" [{allocs} new device buffers]")
5341 } else {
5342 String::new()
5343 },
5344 match (outlier, median) {
5345 (true, Some(m)) => format!(" OUTLIER (median {m:.1})"),
5346 _ => String::new(),
5347 }
5348 );
5349 if level >= 2 || outlier {
5350 let sum: f32 = stamps.iter().map(|s| s.1).sum();
5351 eprintln!(
5352 "spec-stamps: {}| untracked {:.1}",
5353 spec_stamps_format(&stamps),
5354 wall - sum
5355 );
5356 }
5357 if spec_mon.n >= 1 {
5358 spec_walls.push(wall);
5359 }
5360 }
5361 spec_mon.round(t_round.elapsed().as_secs_f64() * 1e3, extra.len() + 1);
5365 spec_trial = Self::spec_trial_round(
5368 spec_trial,
5369 &mut spec_mon,
5370 generated + extra.len() + 1,
5371 );
5372 let mut stopped = false;
5373 for &id in &extra {
5374 if self.confidence_on {
5375 confidence.push(0.0);
5376 }
5377 if !commit!(id) {
5378 stopped = true;
5379 break;
5380 }
5381 }
5382 if stopped {
5383 break 'decode;
5384 }
5385 if spec_time_level() >= 2 {
5386 spec_round_end = Some(std::time::Instant::now());
5387 }
5388 continue 'decode;
5389 }
5390 if self
5391 .graph_failed
5392 .swap(false, std::sync::atomic::Ordering::Relaxed)
5393 {
5394 self.finish_generation(&mut mtp, &mut router, true);
5400 return Err("GPU MTP graph failed during speculative decode".to_string());
5401 }
5402 spec_mon.round(t_round.elapsed().as_secs_f64() * 1e3, 1);
5413 spec_mon.tokens = 0.0;
5414 spec_mon.fails = 3;
5415 spec_trial = Self::spec_trial_round(spec_trial, &mut spec_mon, generated + 1);
5416 hidden = self.forward_layers(&self.embed_single(t_next), next_pos, task_mask);
5417 next_pos += 1;
5418 continue 'decode;
5419 }
5420 Some(m) if !graph_spec && generated + 1 < max_tokens => {
5422 let draft = self.mtp_step(m, &hidden, t_next, next_pos - 1);
5423 drafted += 1;
5424 let emb1 = self.embed_single(t_next);
5425 let emb2 = self.embed_single(draft);
5426 let (h1, h2) = self.forward_pair(&emb1, &emb2, next_pos);
5427
5428 inference::rms_norm_into(
5429 &h1,
5430 &self.weights.final_norm,
5431 self.rms_eps,
5432 self.norm_style,
5433 &mut self.ws.n1,
5434 );
5435 let mut logits1 = self.lm_head_forward(&self.ws.n1);
5436 let t_after = sampler::sample_with_scratch_pool(
5437 &logits1,
5438 &self.sampler_config,
5439 self.sampler_config.penalty_past(&all_ids, bounded_native),
5440 &mut self.rng,
5441 &mut self.sampler_scratch,
5442 self.pool.as_deref(),
5443 );
5444 if self.confidence_on {
5445 confidence.push(sampler::top1_prob_pool(
5446 self.pool.as_deref(),
5447 &mut self.sampler_scratch,
5448 &logits1,
5449 t_after,
5450 calib_temp,
5451 ));
5452 }
5453 attention::recycle_buf(&mut logits1);
5454 if trace_on {
5455 traces.push(TokenTrace {
5458 t: generated,
5459 token_id: t_after,
5460 confidence: confidence.last().copied().unwrap_or(0.0),
5461 active_skill: None,
5462 recon: None,
5463 switched: false,
5464 });
5465 }
5466 let stop = !commit!(t_after);
5467
5468 if t_after == draft {
5469 accepted += 1;
5470 self.commit_linear_scratch();
5471 let _ = self.mtp_step(m, &h1, t_after, next_pos);
5472 hidden = h2;
5473 next_pos += 2;
5474 } else {
5475 for layer in &mut self.kv_cache.layers {
5477 layer.truncate_last(1);
5478 }
5479 if !stop {
5480 let _ = self.mtp_step(m, &h1, t_after, next_pos);
5481 hidden = self.forward_layers(
5482 &self.embed_single(t_after),
5483 next_pos + 1,
5484 None,
5485 );
5486 }
5487 next_pos += 2;
5488 }
5489 if stop {
5490 break 'decode;
5491 }
5492 }
5493 _ => {
5495 #[cfg(feature = "gpu")]
5500 if Self::dsv4_spec_on() && self.dsv4.is_some() {
5501 static SAID: std::sync::Once = std::sync::Once::new();
5502 SAID.call_once(|| {
5503 eprintln!(
5504 "dsv4-spec гейт: mtp={} mask={} router={} trace={} temp={} rep={} ",
5505 !self.dsv4_mtp.is_empty(),
5506 task_mask.is_none(),
5507 router.is_none(),
5508 !trace_on,
5509 self.sampler_config.temperature < 1e-6,
5510 self.sampler_config.repetition_penalty == 1.0,
5511 );
5512 });
5513 }
5514 #[cfg(feature = "gpu")]
5515 if Self::dsv4_spec_on()
5516 && self.dsv4.is_some()
5517 && !self.dsv4_mtp.is_empty()
5518 && task_mask.is_none()
5519 && router.is_none()
5520 && !trace_on
5521 && self.sampler_config.temperature < 1e-6
5522 && self.sampler_config.repetition_penalty == 1.0
5523 && generated + 1 < max_tokens
5524 && all_ids.len() >= 2
5525 && generated >= dsv4_spec_retry_at
5526 {
5527 let tip_token = all_ids[all_ids.len() - 2];
5528 let drafted0 = drafted;
5529 let round = self.dsv4_spec_step(
5530 tip_token,
5531 t_next,
5532 next_pos,
5533 max_tokens.saturating_sub(generated),
5534 &mut drafted,
5535 &mut accepted,
5536 );
5537 if drafted > drafted0 {
5538 let useful = round.as_ref().is_some_and(|(extra, _)| !extra.is_empty());
5539 if useful {
5540 dsv4_spec_bad = 0;
5541 } else {
5542 dsv4_spec_bad += 1;
5543 if dsv4_spec_bad >= 2 {
5544 dsv4_spec_bad = 0;
5545 dsv4_spec_retry_at = generated.saturating_add(32);
5546 tracing::info!(
5547 "dsv4: draft не окупился дважды — точный walk на 32 токена"
5548 );
5549 }
5550 }
5551 }
5552 if let Some((extra, n_pos)) = round {
5553 next_pos = n_pos;
5554 let mut stopped = false;
5555 for &id in &extra {
5556 if self.confidence_on {
5557 confidence.push(0.0);
5558 }
5559 if !commit!(id) {
5560 stopped = true;
5561 break;
5562 }
5563 }
5564 if stopped {
5565 break 'decode;
5566 }
5567 continue 'decode;
5568 }
5569 }
5570 self.graph_want_logits = fuse_lm;
5571 let mut t_fwd = t_next;
5577 let pure_greedy = self.sampler_config.temperature < 1e-6
5578 && self.sampler_config.repetition_penalty == 1.0
5579 && self.sampler_config.suppress_tokens.is_empty();
5580 let burst_k = std::env::var("CMF_MULTISTEP")
5585 .ok()
5586 .and_then(|v| v.parse::<usize>().ok())
5587 .unwrap_or(0);
5588 if pure_greedy
5589 && burst_k >= 1
5590 && fuse_lm
5591 && task_mask.is_none()
5592 && router.is_none()
5593 && !trace_on
5594 && !self.confidence_on
5595 {
5596 let mut stopped = false;
5597 loop {
5598 let room = max_tokens.saturating_sub(generated);
5599 if room <= 2 {
5600 break;
5601 }
5602 let k = burst_k.min(room - 1);
5603 if k < 1 {
5604 break;
5605 }
5606 let Some(ids) = self.try_multi_burst(t_fwd, next_pos, k) else {
5607 if self
5608 .graph_failed
5609 .swap(false, std::sync::atomic::Ordering::Relaxed)
5610 {
5611 self.finish_generation(&mut mtp, &mut router, true);
5612 return Err(
5613 "GPU token graph failed during greedy burst".to_string()
5614 );
5615 }
5616 break;
5617 };
5618 next_pos += k;
5619 for &id in &ids {
5620 if !commit!(id) {
5621 stopped = true;
5622 break;
5623 }
5624 }
5625 if stopped {
5626 break;
5627 }
5628 t_fwd = *ids.last().unwrap();
5629 }
5630 if stopped {
5631 break 'decode;
5632 }
5633 }
5634 #[cfg(target_os = "macos")]
5644 if graph_spec
5645 && spec_watchdog_off
5646 && next_pos > 0
5647 && self.mtp_graph_mode == Some(true)
5648 && crate::gpu::q1_force()
5649 {
5650 if let Some(m) = mtp.as_mut() {
5651 let _ = self.mtp_step_metal(m, &hidden, t_fwd, next_pos - 1, false);
5652 }
5653 }
5654 hidden = self.forward_layers(&self.embed_single(t_fwd), next_pos, task_mask);
5655 next_pos += 1;
5656 if let Some(r) = &mut router {
5659 let phi = self.dyn_phi_ema.clone();
5660 let decision = r.step(&phi, generated);
5661 if let Some(new_active) = decision {
5662 let _ = self.set_active_skill(new_active);
5663 }
5664 if trace_on {
5667 if let Some(last) = traces.last_mut() {
5668 let e = r.last_best_e();
5669 last.recon = e.is_finite().then_some(e);
5670 last.switched = decision.is_some();
5671 }
5672 }
5673 }
5674 }
5675 }
5676 }
5677
5678 let cancelled = finish_reason == "cancelled";
5679 let dyn_switched = router.as_ref().is_some_and(|r| !r.switches.is_empty());
5683 if mimo_spec {
5684 if let Some(st) = self.mimo_mtp.as_ref() {
5685 let line = st.stats.line();
5686 tracing::info!("{line}");
5687 if std::env::var_os("CMF_MIMO_MTP_STATS").is_some() {
5688 eprintln!("{line}");
5689 }
5690 }
5691 }
5692 self.finish_generation(&mut mtp, &mut router, cancelled);
5693
5694 let output_ids = &all_ids[input_ids.len()..];
5695 let forwarded = input_ids.len() + output_ids.len().saturating_sub(1);
5699 let consumed = std::mem::take(&mut all_ids);
5704 if dyn_switched {
5705 self.clear_sequence_state();
5706 } else if cancelled || mimo_spec || prompt_rows.is_some() {
5707 self.clear_history();
5708 } else {
5709 self.record_consumed_prefix(&consumed[..forwarded.min(consumed.len())], reuse_from);
5710 }
5711 all_ids = consumed;
5712 let output_ids = &all_ids[input_ids.len()..];
5713 confidence.truncate(output_ids.len()); traces.truncate(output_ids.len());
5715 Ok(GenerateResult {
5716 text: self.tokenizer.decode(output_ids),
5717 token_ids: output_ids.to_vec(),
5718 prompt_tokens: input_ids.len(),
5719 tokens_generated: generated,
5720 finish_reason,
5721 mtp_drafted: drafted,
5722 mtp_accepted: accepted,
5723 token_confidence: confidence,
5724 traces,
5725 })
5726 }
5727
5728 fn mtp_step(
5732 &mut self,
5733 m: &mut MtpModule,
5734 hidden: &[f32],
5735 next_token: u32,
5736 position: usize,
5737 ) -> u32 {
5738 self.mtp_step_h(m, hidden, next_token, position).0
5739 }
5740
5741 fn chain_probe_note(depth: usize, prefix_ok: bool) {
5745 use std::sync::Mutex;
5746 static T: Mutex<Vec<(u64, u64)>> = Mutex::new(Vec::new());
5747 let mut t = T.lock().unwrap();
5748 if t.len() <= depth {
5749 t.resize(depth + 1, (0, 0));
5750 }
5751 t[depth].0 += 1;
5752 t[depth].1 += prefix_ok as u64;
5753 if depth == 0 && t[0].0 % 128 == 0 {
5754 let line: Vec<String> = t
5755 .iter()
5756 .enumerate()
5757 .map(|(d, (n, k))| {
5758 format!(
5759 "d{}={:.0}%({n})",
5760 d + 1,
5761 100.0 * *k as f64 / (*n).max(1) as f64
5762 )
5763 })
5764 .collect();
5765 eprintln!("mtp-chain: {}", line.join(" "));
5766 }
5767 }
5768
5769 fn mtp_step_hl(
5777 &mut self,
5778 m: &mut MtpModule,
5779 hidden: &[f32],
5780 next_token: u32,
5781 position: usize,
5782 ) -> (Vec<f32>, Vec<f32>) {
5783 #[cfg(target_os = "macos")]
5788 if self.mtp_graph_mode != Some(false) && crate::gpu::q1_force() {
5789 if let Some(r) = self.mtp_step_metal(m, hidden, next_token, position, true) {
5790 self.mtp_graph_mode = Some(true);
5791 return r;
5792 }
5793 if self.mtp_graph_mode == Some(true) {
5794 tracing::error!("mtp Metal graph failed after admission");
5795 self.clear_sequence_state();
5796 self.graph_failed
5797 .store(true, std::sync::atomic::Ordering::Relaxed);
5798 self.cancel
5799 .store(true, std::sync::atomic::Ordering::Relaxed);
5800 return (Vec::new(), Vec::new());
5801 }
5802 self.mtp_graph_mode = Some(false);
5803 }
5804 #[cfg(feature = "gpu")]
5805 if self.mtp_graph_mode != Some(false) {
5806 if !self.mtp_graph_ok(m) {
5807 if self.mtp_graph_mode == Some(true) {
5808 tracing::error!("mtp graph became unavailable after admission");
5813 self.clear_sequence_state();
5814 self.graph_failed
5815 .store(true, std::sync::atomic::Ordering::Relaxed);
5816 self.cancel
5817 .store(true, std::sync::atomic::Ordering::Relaxed);
5818 return (Vec::new(), Vec::new());
5819 }
5820 self.mtp_graph_mode = Some(false);
5821 } else {
5822 if let Some(r) = self.mtp_step_graph(m, hidden, next_token, position) {
5823 self.mtp_graph_mode = Some(true);
5824 return r;
5825 }
5826 if self.graph_failed.load(std::sync::atomic::Ordering::Relaxed) {
5827 return (Vec::new(), Vec::new());
5834 }
5835 tracing::error!("mtp graph failed or declined after admission");
5839 self.clear_sequence_state();
5840 self.graph_failed
5841 .store(true, std::sync::atomic::Ordering::Relaxed);
5842 self.cancel
5843 .store(true, std::sync::atomic::Ordering::Relaxed);
5844 return (Vec::new(), Vec::new());
5845 }
5846 }
5847 let e = self.embed_single(next_token);
5851 let mut cat = vec![0.0f32; 2 * self.hidden_size];
5852 let (cat_e, cat_h) = cat.split_at_mut(self.hidden_size);
5853 inference::rms_norm_into(&e, &m.enorm, self.rms_eps, self.norm_style, cat_e);
5854 inference::rms_norm_into(hidden, &m.hnorm, self.rms_eps, self.norm_style, cat_h);
5855 let mut x = vec![0.0f32; self.hidden_size];
5856 m.eh_proj.matvec(&cat, &mut x, self.pool.as_deref());
5857
5858 let lw = &m.layer;
5860 inference::rms_norm_into(
5861 &x,
5862 &lw.input_norm,
5863 self.rms_eps,
5864 self.norm_style,
5865 &mut self.ws.n1,
5866 );
5867 let attn = match &lw.attn {
5868 AttnKind::Mla(_) => unreachable!("MLA has no MTP/pair path"),
5870 AttnKind::Kda(_) => unreachable!("KDA has no MTP/pair path"),
5871 AttnKind::Full {
5872 wq,
5873 wk,
5874 wv,
5875 wo,
5876 q_norm,
5877 k_norm,
5878 output_gate,
5879 softplus_gate,
5880 bias,
5881 } => {
5882 let mut cfg = self.attn_cfg(position);
5883 cfg.q_norm = q_norm.as_deref();
5884 cfg.k_norm = k_norm.as_deref();
5885 cfg.output_gate = *output_gate;
5886 cfg.softplus_gate = softplus_gate
5887 .as_ref()
5888 .map(|(gate, per_head)| (gate, *per_head));
5889 cfg.bias = bias
5890 .as_ref()
5891 .map(|(q, k, v)| (q.as_slice(), k.as_slice(), v.as_slice()));
5892 attention::qwen_attention(&self.ws.n1, wq, wk, wv, wo, &mut m.kv, &cfg)
5893 }
5894 AttnKind::Linear(_)
5895 | AttnKind::LinearGdn(_)
5896 | AttnKind::ShortConv(_)
5897 | AttnKind::Bounded(_) => {
5898 unreachable!("MTP block is full attention")
5899 }
5900 };
5901 for (i, &a) in attn.iter().enumerate() {
5902 x[i] += a;
5903 }
5904 inference::rms_norm_into(
5905 &x,
5906 &lw.post_norm,
5907 self.rms_eps,
5908 self.norm_style,
5909 &mut self.ws.p1,
5910 );
5911 let ffn = ffn_forward(&lw.ffn, &self.ws.p1, self.pool.as_deref(), None);
5912 for (i, &f) in ffn.iter().enumerate() {
5913 x[i] += f;
5914 }
5915
5916 inference::rms_norm_into(
5917 &x,
5918 &m.final_norm,
5919 self.rms_eps,
5920 self.norm_style,
5921 &mut self.ws.n1,
5922 );
5923 let lg = self.lm_head_forward(&self.ws.n1);
5924 (lg, x)
5925 }
5926
5927 fn mtp_step_h(
5929 &mut self,
5930 m: &mut MtpModule,
5931 hidden: &[f32],
5932 next_token: u32,
5933 position: usize,
5934 ) -> (u32, Vec<f32>) {
5935 let (mut lg, x) = self.mtp_step_hl(m, hidden, next_token, position);
5936 let draft = sampler::argmax(&lg);
5937 attention::recycle_buf(&mut lg);
5938 (draft, x)
5939 }
5940
5941 fn spec_trial_round(trial: SpecTrial, mon: &mut SpecMon, generated: usize) -> SpecTrial {
5947 match trial {
5948 SpecTrial::Spec { t0, gen0, rounds } => {
5949 let rounds = rounds + 1;
5950 if rounds >= 5 {
5951 if mon.plain_ms > 0.0 {
5952 let keep = mon.pays();
5953 mon.fails = 0;
5954 tracing::info!(
5955 "speculation re-check: {:.2} tok/round in {:.1} ms vs plain {:.1} ms/tok — {}",
5956 mon.tokens,
5957 mon.round_ms,
5958 mon.plain_ms,
5959 if keep { "speculating" } else { "plain" }
5960 );
5961 SpecTrial::Decided {
5962 spec: keep,
5963 recheck_at: if keep { usize::MAX } else { generated + 128 },
5964 }
5965 } else if mon.pays() {
5966 mon.fails = 0;
5971 tracing::info!(
5972 "speculation trial: {:.2} tok/round in {:.1} ms — speculating (plain not timed)",
5973 mon.tokens,
5974 mon.round_ms,
5975 );
5976 SpecTrial::Decided {
5977 spec: true,
5978 recheck_at: usize::MAX,
5979 }
5980 } else {
5981 SpecTrial::Plain {
5982 t0: std::time::Instant::now(),
5983 gen0: generated,
5984 }
5985 }
5986 } else {
5987 SpecTrial::Spec { t0, gen0, rounds }
5988 }
5989 }
5990 SpecTrial::Decided { spec: true, .. } => {
5991 if mon.pays() {
5992 mon.fails = 0;
5993 trial
5994 } else {
5995 mon.fails += 1;
5996 if mon.fails >= 4 {
5997 if mon.plain_ms <= 0.0 {
5998 tracing::info!(
6002 "speculation doubtful: {:.2} tok/round in {:.1} ms — timing plain",
6003 mon.tokens,
6004 mon.round_ms,
6005 );
6006 return SpecTrial::Plain {
6007 t0: std::time::Instant::now(),
6008 gen0: generated,
6009 };
6010 }
6011 tracing::info!(
6012 "speculation stopped: {:.2} tok/round in {:.1} ms vs plain {:.1} ms/tok",
6013 mon.tokens,
6014 mon.round_ms,
6015 mon.plain_ms
6016 );
6017 SpecTrial::Decided {
6018 spec: false,
6019 recheck_at: generated + 128,
6020 }
6021 } else {
6022 trial
6023 }
6024 }
6025 }
6026 other => other,
6027 }
6028 }
6029
6030 fn mtp_kv_id(&self) -> u64 {
6033 self.graph_kv_id | (1u64 << 40)
6034 }
6035
6036 const MTP_LAYER_BASE: usize = 0;
6041
6042 #[cfg(feature = "gpu")]
6049 fn rewind_mtp_graph_mirror(&self, stored: usize) -> bool {
6050 self.mtp_graph_mode != Some(true)
6051 || crate::gpu::graph_kv_set_stored(self.mtp_kv_id(), Self::MTP_LAYER_BASE, stored)
6052 }
6053
6054 #[cfg(feature = "gpu")]
6060 fn rewind_trunk_graph_mirrors(&self, stored: usize) -> bool {
6061 let mut ok = true;
6062 let mut expected = false;
6063 for li in 0..self.num_layers {
6064 if matches!(
6065 self.weights.layers[self.phys_layer(li)].attn,
6066 AttnKind::Full { .. }
6067 ) {
6068 expected = true;
6069 ok &= crate::gpu::graph_kv_set_stored(self.graph_kv_id, li, stored);
6070 }
6071 }
6072 !expected || ok
6073 }
6074
6075 fn graph_gdn_layer_count(&self) -> usize {
6079 (0..self.num_layers)
6080 .filter(|&li| {
6081 matches!(
6082 &self.weights.layers[self.phys_layer(li)].attn,
6083 AttnKind::LinearGdn(_)
6084 )
6085 })
6086 .count()
6087 }
6088
6089 fn mtp_block_input(&mut self, m: &MtpModule, hidden: &[f32], next_token: u32) -> Vec<f32> {
6092 let e = self.embed_single(next_token);
6093 let mut cat = vec![0.0f32; 2 * self.hidden_size];
6094 let (cat_e, cat_h) = cat.split_at_mut(self.hidden_size);
6095 inference::rms_norm_into(&e, &m.enorm, self.rms_eps, self.norm_style, cat_e);
6096 inference::rms_norm_into(hidden, &m.hnorm, self.rms_eps, self.norm_style, cat_h);
6097 let mut x = vec![0.0f32; self.hidden_size];
6098 m.eh_proj.matvec(&cat, &mut x, self.pool.as_deref());
6099 x
6100 }
6101
6102 #[cfg(feature = "gpu")]
6105 fn mtp_block_graph_ok(&self, m: &MtpModule) -> bool {
6106 if std::env::var("CMF_MTP_GRAPH").as_deref() == Ok("0") {
6107 return false;
6108 }
6109 if !crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Decode)
6110 || !crate::gpu::enabled_here()
6111 || self.attn_softcap > 0.0
6112 || self.attention_heads_per_layer.is_some()
6113 || self.v_head_dim.is_some()
6116 {
6117 return false;
6118 }
6119 matches!(
6120 &m.layer.attn,
6121 AttnKind::Full {
6122 softplus_gate: None,
6123 ..
6124 }
6125 ) && matches!(&m.layer.ffn, FfnKind::Dense(_))
6126 }
6127
6128 #[cfg(feature = "gpu")]
6132 fn mtp_graph_ok(&self, m: &MtpModule) -> bool {
6133 if !self.mtp_block_graph_ok(m) {
6134 return false;
6135 }
6136 let AttnKind::Full { wq, wk, wv, wo, .. } = &m.layer.attn else {
6137 return false;
6138 };
6139 let FfnKind::Dense(d) = &m.layer.ffn else {
6140 return false;
6141 };
6142 d.segs.is_empty()
6143 && wq.graph_weight().is_some()
6144 && wk.graph_weight().is_some()
6145 && wv.graph_weight().is_some()
6146 && wo.graph_weight().is_some()
6147 && d.gate_proj.graph_weight().is_some()
6148 && d.up_proj.graph_weight().is_some()
6149 && d.down_proj.graph_weight().is_some()
6150 && self.weights.lm_head.graph_weight().is_some()
6151 }
6152
6153 #[cfg(feature = "gpu")]
6159 fn mtp_step_graph(
6160 &mut self,
6161 m: &mut MtpModule,
6162 hidden: &[f32],
6163 next_token: u32,
6164 position: usize,
6165 ) -> Option<(Vec<f32>, Vec<f32>)> {
6166 if !self.mtp_graph_ok(m) {
6167 return None;
6168 }
6169 let lw = &m.layer;
6170 let AttnKind::Full {
6171 wq,
6172 wk,
6173 wv,
6174 wo,
6175 q_norm,
6176 k_norm,
6177 output_gate,
6178 softplus_gate,
6179 bias,
6180 } = &lw.attn
6181 else {
6182 return None;
6183 };
6184 if softplus_gate.is_some() {
6185 return None;
6186 }
6187 let FfnKind::Dense(d) = &lw.ffn else {
6188 return None;
6189 };
6190 if !d.segs.is_empty() {
6191 return None; }
6193 let mut x = self.mtp_block_input(m, hidden, next_token);
6196 fn gw(t: &QTensor) -> Option<crate::gpu::GraphW<'_>> {
6197 let (_, i, kind, rs) = t.graph_weight()?;
6198 Some(crate::gpu::GraphW {
6199 idx: i,
6200 kind,
6201 row_scale: rs,
6202 data: &[],
6203 prism: crate::gpu::GraphPrismOp::None,
6204 affine: false,
6205 })
6206 }
6207 let (model, _, _, _) = wq.graph_weight()?;
6208 let model = model.clone();
6209 let (lm_gw, lm_rows) = {
6210 let (_, i, kind, rs) = self.weights.lm_head.graph_weight()?;
6211 let rows = if kind == 6 {
6215 self.draft_head_rows(self.weights.lm_head.rows())
6216 } else {
6217 self.weights.lm_head.rows()
6218 };
6219 (
6220 crate::gpu::GraphW {
6221 idx: i,
6222 kind,
6223 row_scale: rs,
6224 data: &[],
6225 prism: crate::gpu::GraphPrismOp::None,
6226 affine: false,
6227 },
6228 rows,
6229 )
6230 };
6231 let layer = crate::gpu::GraphLayer {
6232 input_norm: &lw.input_norm,
6233 attn: crate::gpu::GraphAttn::Full {
6234 wq: gw(wq)?,
6235 wk: gw(wk)?,
6236 wv: gw(wv)?,
6237 wo: gw(wo)?,
6238 q_norm: q_norm.as_deref(),
6239 k_norm: k_norm.as_deref(),
6240 late_qk_norm: self.qk_norm_after_rope,
6241 bias: bias
6242 .as_ref()
6243 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
6244 output_gate: *output_gate,
6245 cpu_k: m.kv.k_heads(),
6246 cpu_v: m.kv.v_heads(),
6247 geom: None,
6248 head_gate: None,
6249 },
6250 post_norm: &lw.post_norm,
6251 ffn: crate::gpu::GraphFfn::Dense {
6252 gate: gw(&d.gate_proj)?,
6253 up: gw(&d.up_proj)?,
6254 down: gw(&d.down_proj)?,
6255 act: d.act.graph_act()?,
6256 },
6257 };
6258 let nh = self.num_heads;
6259 let (nkv, hd, rd) = self.layer_geom(0);
6260 let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
6261 let mut logits = Vec::new();
6262 let ok = crate::gpu::forward_token_graph(
6263 &model,
6264 self.mtp_kv_id(),
6265 std::slice::from_ref(&layer),
6266 &[None],
6267 self.o1_epoch,
6268 &self.inv_freq,
6269 &mut x,
6270 nh,
6271 nkv,
6272 hd,
6273 self.attn_scale,
6274 rd,
6275 self.hidden_size,
6276 self.intermediate_size,
6277 position,
6278 self.kv_cache.max_seq_len,
6279 gemma,
6280 self.rms_eps as f32,
6281 Some((&lm_gw, lm_rows)),
6282 &m.final_norm,
6283 &mut logits,
6284 &[],
6285 1,
6286 None,
6287 None,
6288 None,
6289 Self::MTP_LAYER_BASE,
6290 true,
6291 );
6292 match ok {
6293 crate::gpu::TokenGraphOutcome::Completed => {}
6294 crate::gpu::TokenGraphOutcome::Declined => return None,
6295 crate::gpu::TokenGraphOutcome::Failed => {
6296 self.clear_sequence_state();
6300 self.graph_failed
6301 .store(true, std::sync::atomic::Ordering::Relaxed);
6302 self.cancel
6303 .store(true, std::sync::atomic::Ordering::Relaxed);
6304 return None;
6305 }
6306 }
6307 logits.resize(self.vocab_size, 0.0);
6308 Some((logits, x))
6309 }
6310
6311 #[cfg(feature = "gpu")]
6319 fn mtp_warm_graph(
6320 &mut self,
6321 m: &mut MtpModule,
6322 pairs: &[(&[f32], u32)],
6323 first_pos: usize,
6324 ) -> crate::gpu::BatchGraphOutcome {
6325 if pairs.is_empty() {
6326 return crate::gpu::BatchGraphOutcome::Completed;
6327 }
6328 if !self.mtp_block_graph_ok(m) {
6329 return crate::gpu::BatchGraphOutcome::Declined;
6330 }
6331 let hs = self.hidden_size;
6332 let mut hiddens = Vec::with_capacity(pairs.len() * hs);
6335 for (h, t) in pairs {
6336 hiddens.extend_from_slice(&self.mtp_block_input(m, h, *t));
6337 }
6338 let lw = &m.layer;
6339 let AttnKind::Full {
6340 wq,
6341 wk,
6342 wv,
6343 wo,
6344 q_norm,
6345 k_norm,
6346 output_gate,
6347 bias,
6348 ..
6349 } = &lw.attn
6350 else {
6351 return crate::gpu::BatchGraphOutcome::Declined;
6352 };
6353 let FfnKind::Dense(d) = &lw.ffn else {
6354 return crate::gpu::BatchGraphOutcome::Declined;
6355 };
6356 if !d.segs.is_empty() {
6357 return crate::gpu::BatchGraphOutcome::Declined; }
6359 let (AttnKind::Full { softplus_gate: None, .. }, Some(gact)) =
6362 (&lw.attn, d.act.graph_act())
6363 else {
6364 return crate::gpu::BatchGraphOutcome::Declined;
6365 };
6366 fn gw(t: &QTensor) -> Option<crate::gpu::GraphW<'_>> {
6367 let (_, i, kind, rs) = t.graph_weight()?;
6368 Some(crate::gpu::GraphW {
6369 idx: i,
6370 kind,
6371 row_scale: rs,
6372 data: &[],
6373 prism: crate::gpu::GraphPrismOp::None,
6374 affine: false,
6375 })
6376 }
6377 let Some((model, _, _, _)) = wq.graph_weight() else {
6378 return crate::gpu::BatchGraphOutcome::Declined;
6379 };
6380 let model = model.clone();
6381 let (Some(gwq), Some(gwk), Some(gwv), Some(gwo), Some(gg), Some(gu), Some(gd)) = (
6382 gw(wq),
6383 gw(wk),
6384 gw(wv),
6385 gw(wo),
6386 gw(&d.gate_proj),
6387 gw(&d.up_proj),
6388 gw(&d.down_proj),
6389 ) else {
6390 return crate::gpu::BatchGraphOutcome::Declined;
6391 };
6392 let layer = crate::gpu::GraphLayer {
6393 input_norm: &lw.input_norm,
6394 attn: crate::gpu::GraphAttn::Full {
6395 wq: gwq,
6396 wk: gwk,
6397 wv: gwv,
6398 wo: gwo,
6399 q_norm: q_norm.as_deref(),
6400 k_norm: k_norm.as_deref(),
6401 late_qk_norm: self.qk_norm_after_rope,
6402 bias: bias
6403 .as_ref()
6404 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
6405 output_gate: *output_gate,
6406 cpu_k: m.kv.k_heads(),
6407 cpu_v: m.kv.v_heads(),
6408 geom: None,
6409 head_gate: None,
6410 },
6411 post_norm: &lw.post_norm,
6412 ffn: crate::gpu::GraphFfn::Dense {
6413 gate: gg,
6414 up: gu,
6415 down: gd,
6416 act: gact,
6417 },
6418 };
6419 let positions: Vec<usize> = (first_pos..first_pos + pairs.len()).collect();
6420 let nh = self.num_heads;
6421 let (nkv, hd, rd) = self.layer_geom(0);
6422 let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
6423 crate::gpu::forward_batch_graph(
6424 &model,
6425 self.mtp_kv_id(),
6426 std::slice::from_ref(&layer),
6427 &self.inv_freq,
6428 &mut hiddens,
6429 nh,
6430 nkv,
6431 hd,
6432 rd,
6433 hs,
6434 self.intermediate_size,
6435 &positions,
6436 self.kv_cache.max_seq_len,
6437 gemma,
6438 self.rms_eps as f32,
6439 self.attn_scale,
6440 pairs.len(),
6441 &[],
6442 0,
6443 None,
6444 None,
6445 )
6446 }
6447
6448 #[cfg(feature = "gpu")]
6455 fn mtp_warm_graph_fallback(
6456 &mut self,
6457 m: &mut MtpModule,
6458 pairs: &[(&[f32], u32)],
6459 first_pos: usize,
6460 ) -> bool {
6461 if pairs.is_empty() {
6462 return true;
6463 }
6464 let graphable = self.mtp_block_graph_ok(m);
6465 if !graphable {
6466 if self.mtp_graph_mode == Some(true) {
6470 return false;
6471 }
6472 self.mtp_graph_mode = Some(false);
6473 for (j, (h, t)) in pairs.iter().enumerate() {
6474 self.mtp_warm(m, h, *t, first_pos + j);
6475 }
6476 return true;
6477 }
6478
6479 for (j, (h, t)) in pairs.iter().enumerate() {
6484 if self.mtp_step_graph(m, h, *t, first_pos + j).is_none() {
6485 return false;
6486 }
6487 }
6488 self.mtp_graph_mode = Some(true);
6489 true
6490 }
6491
6492 #[cfg(feature = "gpu")]
6497 fn mtp_warm_prefill_pairs(
6498 &mut self,
6499 m: &mut MtpModule,
6500 pairs: &[(&[f32], u32)],
6501 first_pos: usize,
6502 ) -> Result<(), &'static str> {
6503 if self.mtp_graph_mode == Some(false) || !self.mtp_graph_ok(m) {
6508 if self.mtp_graph_mode == Some(true) {
6509 return Err("MTP token graph became unavailable after admission");
6510 }
6511 self.mtp_graph_mode = Some(false);
6512 for (j, (h, t)) in pairs.iter().enumerate() {
6513 self.mtp_warm(m, h, *t, first_pos + j);
6514 }
6515 return Ok(());
6516 }
6517 match self.mtp_warm_graph(m, pairs, first_pos) {
6518 crate::gpu::BatchGraphOutcome::Completed => {
6519 if !pairs.is_empty() {
6520 self.mtp_graph_mode = Some(true);
6521 }
6522 Ok(())
6523 }
6524 crate::gpu::BatchGraphOutcome::Declined => {
6525 if self.mtp_warm_graph_fallback(m, pairs, first_pos) {
6526 Ok(())
6527 } else {
6528 Err("MTP warm-up fallback failed after device admission")
6529 }
6530 }
6531 crate::gpu::BatchGraphOutcome::Failed => {
6532 Err("MTP warm batch graph failed after admission")
6533 }
6534 }
6535 }
6536
6537 #[cfg(not(feature = "gpu"))]
6538 fn mtp_warm_prefill_pairs(
6539 &mut self,
6540 m: &mut MtpModule,
6541 pairs: &[(&[f32], u32)],
6542 first_pos: usize,
6543 ) -> Result<(), &'static str> {
6544 for (j, (h, t)) in pairs.iter().enumerate() {
6545 self.mtp_warm(m, h, *t, first_pos + j);
6546 }
6547 Ok(())
6548 }
6549
6550 fn mtp_warm(&mut self, m: &mut MtpModule, hidden: &[f32], next_token: u32, position: usize) {
6554 let e = self.embed_single(next_token);
6555 let mut cat = vec![0.0f32; 2 * self.hidden_size];
6556 let (cat_e, cat_h) = cat.split_at_mut(self.hidden_size);
6557 inference::rms_norm_into(&e, &m.enorm, self.rms_eps, self.norm_style, cat_e);
6558 inference::rms_norm_into(hidden, &m.hnorm, self.rms_eps, self.norm_style, cat_h);
6559 let mut x = vec![0.0f32; self.hidden_size];
6560 m.eh_proj.matvec(&cat, &mut x, self.pool.as_deref());
6561 inference::rms_norm_into(
6562 &x,
6563 &m.layer.input_norm,
6564 self.rms_eps,
6565 self.norm_style,
6566 &mut self.ws.n1,
6567 );
6568 let attn = match &m.layer.attn {
6569 AttnKind::Full {
6570 wq,
6571 wk,
6572 wv,
6573 wo,
6574 q_norm,
6575 k_norm,
6576 output_gate,
6577 softplus_gate,
6578 bias,
6579 } => {
6580 let mut cfg = self.attn_cfg(position);
6581 cfg.q_norm = q_norm.as_deref();
6582 cfg.k_norm = k_norm.as_deref();
6583 cfg.output_gate = *output_gate;
6584 cfg.softplus_gate = softplus_gate.as_ref().map(|(g, p)| (g, *p));
6585 cfg.bias = bias
6586 .as_ref()
6587 .map(|(q, k, v)| (q.as_slice(), k.as_slice(), v.as_slice()));
6588 attention::qwen_attention(&self.ws.n1, wq, wk, wv, wo, &mut m.kv, &cfg)
6589 }
6590 _ => return,
6591 };
6592 let _ = attn;
6593 }
6594
6595 #[cfg(feature = "gpu")]
6602 #[allow(clippy::too_many_arguments)]
6603 fn graph_spec_step(
6604 &mut self,
6605 m: &mut MtpModule,
6606 hidden: &[f32],
6607 t_next: u32,
6608 next_pos: usize,
6609 drafted: &mut usize,
6610 accepted: &mut usize,
6611 all_ids: &mut Vec<u32>,
6615 room: usize,
6620 ) -> Option<(Vec<u32>, usize, Vec<f32>)> {
6621 #[cfg(target_os = "macos")]
6632 let metal_native = crate::gpu::q1_force();
6633 #[cfg(not(target_os = "macos"))]
6634 let metal_native = false;
6635 #[cfg(feature = "gpu")]
6636 let k_default = if metal_native {
6637 7
6640 } else if crate::gpu_wgpu::verify_i8_on() {
6641 5
6642 } else {
6643 4
6644 };
6645 #[cfg(not(feature = "gpu"))]
6646 let k_default = 4;
6647 let k_env: Option<usize> = std::env::var("CMF_GRAPH_SPEC_K")
6648 .ok()
6649 .and_then(|v| v.parse().ok())
6650 .filter(|&v| (1..=8).contains(&v));
6651 let (k_start, k_max) = if metal_native { (7, 7) } else { (3, k_default.max(5)) };
6656 let k_full: usize = k_env.unwrap_or_else(|| self.spec_k_adapt.unwrap_or(k_start));
6657 let k_spec = k_full.min(room).max(1);
6658 let k_capped = k_spec < k_full;
6661 if next_pos == 0 {
6662 return None;
6663 }
6664 let t_round = std::time::Instant::now();
6665 let subs = || crate::gpu_wgpu::SUBMITS.load(std::sync::atomic::Ordering::Relaxed);
6681 let sub0 = subs();
6682 let cfg = self.sampler_config.clone();
6687 let penalized = !(cfg.repetition_penalty == 1.0
6688 && cfg.presence_penalty == 0.0
6689 && cfg.suppress_tokens.is_empty());
6690 let greedy_pen = cfg.temperature < 1e-6 && penalized;
6695 let sampling = cfg.temperature >= 1e-6;
6696 let sparse = sampling && sampler::sparse_ok(&cfg);
6702 let base_len = all_ids.len();
6703 if sampling && !sparse && self.spec_q.len() < k_spec {
6704 self.spec_q.resize_with(k_spec, Vec::new);
6705 }
6706 if sparse && self.spec_qs.len() < k_spec {
6707 self.spec_qs.resize_with(k_spec, Vec::new);
6708 }
6709 let mut drafts = Vec::with_capacity(k_spec);
6714 let mut hx = hidden.to_vec();
6715 let spec_dbg = std::env::var("CMF_SPEC_DBG").is_ok();
6718 spec_stamp("pro");
6719 #[cfg(target_os = "macos")]
6725 if metal_native && !sampling && !greedy_pen && self.mtp_graph_mode != Some(false) {
6726 match self.mtp_draft_chain_metal(m, hidden, t_next, next_pos - 1, k_spec) {
6727 Ok(ids) => {
6728 self.mtp_graph_mode = Some(true);
6729 drafts = ids;
6730 }
6731 Err(true) => {
6732 tracing::error!("mtp Metal draft chain failed after commit");
6733 self.clear_sequence_state();
6734 self.graph_failed
6735 .store(true, std::sync::atomic::Ordering::Relaxed);
6736 self.cancel
6737 .store(true, std::sync::atomic::Ordering::Relaxed);
6738 return None;
6739 }
6740 Err(false) => {}
6741 }
6742 }
6743 for j in drafts.len()..k_spec {
6744 let tok_in = if j == 0 { t_next } else { drafts[j - 1] };
6745 let mut dbg_ref: Option<(Vec<f32>, Vec<f32>)> = None;
6746 if spec_dbg {
6747 let saved = self.mtp_graph_mode;
6748 self.mtp_graph_mode = Some(false);
6749 let r = self.mtp_step_hl(m, &hx, tok_in, next_pos - 1 + j);
6750 self.mtp_graph_mode = saved;
6751 if self.graph_failed.load(std::sync::atomic::Ordering::Relaxed) {
6752 return None;
6753 }
6754 m.kv.truncate_last(1);
6755 dbg_ref = Some(r);
6756 }
6757 let (mut lg, hj) = self.mtp_step_hl(m, &hx, tok_in, next_pos - 1 + j);
6758 if self.graph_failed.load(std::sync::atomic::Ordering::Relaxed) {
6759 return None;
6760 }
6761 if let Some((lg_cpu, h_cpu)) = dbg_ref {
6762 let n = |v: &[f32]| v.iter().map(|x| x * x).sum::<f32>().sqrt();
6763 let dl = lg
6764 .iter()
6765 .zip(&lg_cpu)
6766 .fold(0f32, |m, (a, b)| m.max((a - b).abs()));
6767 let dh = hj
6768 .iter()
6769 .zip(&h_cpu)
6770 .fold(0f32, |m, (a, b)| m.max((a - b).abs()));
6771 eprintln!(
6772 "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 {}",
6773 next_pos - 1 + j,
6774 sampler::argmax(&lg_cpu),
6775 sampler::argmax(&lg),
6776 n(&h_cpu),
6777 n(&hj),
6778 m.kv.seq_len
6779 );
6780 }
6781 let dj = if sparse {
6782 let mut q = std::mem::take(&mut self.spec_qs[j]);
6783 let ok = sampler::sparse_distribution_into(
6784 &lg,
6785 &cfg,
6786 all_ids,
6787 &mut self.sampler_scratch,
6788 self.pool.as_deref(),
6789 &mut q,
6790 );
6791 let d = if ok {
6792 sampler::draw_sparse(&q, &mut self.rng)
6793 } else {
6794 let t = sampler::argmax(&lg);
6796 q.clear();
6797 q.push((t, 1.0));
6798 t
6799 };
6800 self.spec_qs[j] = q;
6801 all_ids.push(d);
6802 d
6803 } else if sampling {
6804 let mut q = std::mem::take(&mut self.spec_q[j]);
6805 sampler::distribution_into(
6806 &lg,
6807 &cfg,
6808 all_ids,
6809 &mut self.sampler_scratch,
6810 self.pool.as_deref(),
6811 &mut q,
6812 );
6813 let d = sampler::draw(&q, &mut self.rng);
6814 self.spec_q[j] = q;
6815 all_ids.push(d); d
6817 } else if greedy_pen {
6818 let d = sampler::argmax_penalized(
6819 &lg,
6820 &cfg,
6821 all_ids,
6822 &mut self.sampler_scratch,
6823 self.pool.as_deref(),
6824 );
6825 all_ids.push(d);
6826 d
6827 } else {
6828 sampler::argmax(&lg)
6829 };
6830 attention::recycle_buf(&mut lg);
6831 drafts.push(dj);
6832 hx = hj;
6833 spec_stamp("d.pick");
6834 }
6835 all_ids.truncate(base_len);
6836 *drafted += k_spec;
6837 let t_draft = t_round.elapsed();
6838 let sub_draft = subs();
6839 let b = k_spec + 1;
6842 let mut hiddens = vec![0.0f32; b * self.hidden_size];
6843 for (i, &t) in std::iter::once(&t_next).chain(drafts.iter()).enumerate() {
6844 let e = self.embed_single(t);
6845 hiddens[i * self.hidden_size..(i + 1) * self.hidden_size].copy_from_slice(&e);
6846 }
6847 let positions: Vec<usize> = (next_pos..next_pos + b).collect();
6848 spec_stamp("v.emb");
6849 let (lm_gw, lm_rows) = {
6850 let (_, i, kind, rs) = self.weights.lm_head.graph_weight()?;
6851 (
6852 crate::gpu::GraphW {
6853 idx: i,
6854 kind,
6855 row_scale: rs,
6856 data: &[],
6857 prism: crate::gpu::GraphPrismOp::None,
6858 affine: false,
6859 },
6860 self.weights.lm_head.rows(),
6861 )
6862 };
6863 let mut logits = Vec::new();
6864 let final_norm = self.weights.final_norm.clone();
6865 #[cfg(target_os = "macos")]
6874 let greedy_dev = metal_native
6875 && !sampling
6876 && !greedy_pen
6877 && !self.confidence_on
6878 && self.final_softcap.is_none()
6879 && self.vocab_size == lm_rows
6885 && std::env::var_os("CMF_METAL_VERIFY_CHECK").is_none()
6886 && std::env::var_os("CMF_LOGIT_DUMP").is_none()
6887 && std::env::var("CMF_METAL_DEV_ARGMAX").as_deref() != Ok("0");
6888 #[cfg(not(target_os = "macos"))]
6889 let greedy_dev = false;
6890 let mut dev_ids: Vec<u32> = Vec::new();
6891 #[cfg(target_os = "macos")]
6892 let verify_outcome = if metal_native {
6893 let lm = self.weights.lm_head.q1_parts()?;
6894 let n_score = self.vocab_size.min(lm_rows);
6895 self.try_batch_graph_metal(
6896 &mut hiddens,
6897 &positions,
6898 b,
6899 Some((lm, &final_norm, &mut logits)),
6900 if greedy_dev {
6901 Some((n_score, &mut dev_ids))
6902 } else {
6903 None
6904 },
6905 )
6906 } else {
6907 self.try_batch_graph_wgpu(
6908 &mut hiddens,
6909 &positions,
6910 b,
6911 Some(crate::gpu::SpecTail {
6912 lm: lm_gw,
6913 lm_rows,
6914 final_norm: &final_norm,
6915 logits_out: &mut logits,
6916 }),
6917 )
6918 };
6919 #[cfg(not(target_os = "macos"))]
6920 let verify_outcome = self.try_batch_graph_wgpu(
6921 &mut hiddens,
6922 &positions,
6923 b,
6924 Some(crate::gpu::SpecTail {
6925 lm: lm_gw,
6926 lm_rows,
6927 final_norm: &final_norm,
6928 logits_out: &mut logits,
6929 }),
6930 );
6931 match verify_outcome {
6932 crate::gpu::BatchGraphOutcome::Completed => {}
6933 crate::gpu::BatchGraphOutcome::Declined => {
6934 m.kv.truncate_last(k_spec);
6938 if !metal_native && !self.rewind_mtp_graph_mirror(next_pos) {
6939 self.clear_sequence_state();
6940 self.graph_failed
6941 .store(true, std::sync::atomic::Ordering::Relaxed);
6942 self.cancel
6943 .store(true, std::sync::atomic::Ordering::Relaxed);
6944 tracing::error!("MTP graph mirror rewind failed after verify decline");
6945 }
6946 return None;
6947 }
6948 crate::gpu::BatchGraphOutcome::Failed => {
6949 self.clear_sequence_state();
6953 self.graph_failed
6954 .store(true, std::sync::atomic::Ordering::Relaxed);
6955 self.cancel
6956 .store(true, std::sync::atomic::Ordering::Relaxed);
6957 tracing::error!("MTP verify batch graph failed after admission");
6958 return None;
6959 }
6960 }
6961 #[cfg(target_os = "macos")]
6967 if metal_native && std::env::var("CMF_METAL_VERIFY_CHECK").as_deref() == Ok("1") {
6968 let snap: Vec<Vec<f32>> = self
6969 .kv_cache
6970 .layers
6971 .iter()
6972 .map(|l| l.linear_state.clone())
6973 .collect();
6974 let attn_lens: Vec<usize> = self.kv_cache.layers.iter().map(|l| l.seq_len).collect();
6975 let toks: Vec<u32> = std::iter::once(t_next)
6976 .chain(drafts.iter().copied())
6977 .collect();
6978 let want_save = self.graph_want_logits;
6979 self.graph_want_logits = false;
6980 for (i, &t) in toks.iter().enumerate() {
6981 let hi = self.forward_layers(&self.embed_single(t), next_pos + i, None);
6982 let _ = self.graph_logits.take();
6983 if std::env::var("CMF_SPEC_PLAIN_HIDDEN").as_deref() == Ok("1") {
6987 hiddens[i * self.hidden_size..(i + 1) * self.hidden_size].copy_from_slice(&hi);
6988 }
6989 let ref_lg = self.logits_from_hidden(&hi);
6990 let row = &logits[i * lm_rows..(i + 1) * lm_rows];
6991 let ra = sampler::argmax(&ref_lg);
6992 let va = sampler::argmax(row);
6993 let mut md = 0f32;
6994 let mut rms = 0f64;
6995 for j in 0..lm_rows.min(ref_lg.len()) {
6996 let d = (ref_lg[j] - row[j]).abs();
6997 md = md.max(d);
6998 rms += (d as f64) * (d as f64);
6999 }
7000 let mut hd = 0f32;
7001 for j in 0..self.hidden_size {
7002 hd = hd.max((hi[j] - hiddens[i * self.hidden_size + j]).abs());
7003 }
7004 eprintln!(
7005 "verify-check row {i} tok {t} pos {}: ref argmax {ra} verify argmax {va} {} | max|dlogit| {md:.3} rms {:.4} | max|dhidden| {hd:.4}",
7006 next_pos + i,
7007 if ra == va { "OK" } else { "MISMATCH" },
7008 (rms / lm_rows as f64).sqrt()
7009 );
7010 }
7011 self.graph_want_logits = want_save;
7012 for (l, st) in self.kv_cache.layers.iter_mut().zip(snap) {
7015 if l.linear_state.len() == st.len() {
7016 l.linear_state.copy_from_slice(&st);
7017 } else {
7018 l.linear_state = st;
7019 }
7020 }
7021 for (li, (l, n0)) in self.kv_cache.layers.iter_mut().zip(attn_lens).enumerate() {
7022 let extra = l.seq_len.saturating_sub(n0);
7023 if extra > 0 {
7024 l.truncate_last(extra);
7025 crate::gpu_metal::kv_mirror_set_stored(self.graph_kv_id, li, n0);
7026 }
7027 }
7028 }
7029 let t_verify = t_round.elapsed();
7030 let sub_verify = subs();
7031 let mut a = 0usize;
7036 let mut forced: Option<u32> = None;
7037 let ids: Vec<u32> = if sparse {
7038 let mut p = std::mem::take(&mut self.spec_ps);
7039 let mut res = std::mem::take(&mut self.spec_ress);
7040 while a < k_spec {
7041 let ok = sampler::sparse_distribution_into(
7042 &logits[a * lm_rows..(a + 1) * lm_rows],
7043 &cfg,
7044 all_ids,
7045 &mut self.sampler_scratch,
7046 self.pool.as_deref(),
7047 &mut p,
7048 );
7049 if !ok {
7050 let t = sampler::argmax(&logits[a * lm_rows..(a + 1) * lm_rows]);
7051 p.clear();
7052 p.push((t, 1.0));
7053 }
7054 match sampler::spec_accept_or_correct_sparse(
7055 &p,
7056 &self.spec_qs[a],
7057 drafts[a],
7058 &mut self.rng,
7059 &mut res,
7060 ) {
7061 None => {
7062 all_ids.push(drafts[a]);
7063 a += 1;
7064 }
7065 Some(c) => {
7066 forced = Some(c);
7067 break;
7068 }
7069 }
7070 }
7071 all_ids.truncate(base_len);
7072 self.spec_ps = p;
7073 self.spec_ress = res;
7074 drafts.clone()
7075 } else if sampling {
7076 let mut p = std::mem::take(&mut self.spec_p);
7077 let mut res = std::mem::take(&mut self.spec_res);
7078 while a < k_spec {
7079 sampler::distribution_into(
7080 &logits[a * lm_rows..(a + 1) * lm_rows],
7081 &cfg,
7082 all_ids,
7083 &mut self.sampler_scratch,
7084 self.pool.as_deref(),
7085 &mut p,
7086 );
7087 match sampler::spec_accept_or_correct(
7088 &p,
7089 &self.spec_q[a],
7090 drafts[a],
7091 &mut self.rng,
7092 &mut res,
7093 self.pool.as_deref(),
7094 ) {
7095 None => {
7096 all_ids.push(drafts[a]);
7097 a += 1;
7098 }
7099 Some(c) => {
7100 forced = Some(c);
7101 break;
7102 }
7103 }
7104 }
7105 all_ids.truncate(base_len);
7106 self.spec_p = p;
7107 self.spec_res = res;
7108 drafts.clone()
7110 } else if greedy_pen {
7111 let mut ids: Vec<u32> = Vec::with_capacity(b);
7115 for i in 0..b {
7116 let t = sampler::argmax_penalized(
7117 &logits[i * lm_rows..(i + 1) * lm_rows],
7118 &cfg,
7119 all_ids,
7120 &mut self.sampler_scratch,
7121 self.pool.as_deref(),
7122 );
7123 ids.push(t);
7124 if i < k_spec && t == drafts[i] {
7125 all_ids.push(t);
7126 } else {
7127 break;
7128 }
7129 }
7130 all_ids.truncate(base_len);
7131 while a < k_spec && a < ids.len() && ids[a] == drafts[a] {
7132 a += 1;
7133 }
7134 ids
7137 } else if greedy_dev && dev_ids.len() == b {
7138 let ids = std::mem::take(&mut dev_ids);
7139 while a < k_spec && ids[a] == drafts[a] {
7140 a += 1;
7141 }
7142 ids
7143 } else {
7144 if logits.len() < b * lm_rows {
7145 self.clear_sequence_state();
7148 self.graph_failed
7149 .store(true, std::sync::atomic::Ordering::Relaxed);
7150 self.cancel
7151 .store(true, std::sync::atomic::Ordering::Relaxed);
7152 tracing::error!("Metal verify returned neither logits nor argmax ids");
7153 return None;
7154 }
7155 let ids: Vec<u32> = (0..b)
7156 .map(|i| sampler::argmax(&logits[i * lm_rows..(i + 1) * lm_rows]))
7157 .collect();
7158 while a < k_spec && ids[a] == drafts[a] {
7159 a += 1;
7160 }
7161 ids
7162 };
7163 spec_stamp("acc");
7164 if spec_dbg {
7165 eprintln!(
7166 "spec-dbg round: t_next {t_next} drafts {:?} verified {:?} accepted {a}",
7167 drafts, ids
7168 );
7169 }
7170 #[cfg(target_os = "macos")]
7174 let commit_ref: Option<(Vec<Vec<f32>>, Vec<(usize, Vec<f32>, Vec<f32>)>)> = if metal_native
7175 && std::env::var("CMF_METAL_VERIFY_CHECK").as_deref() == Ok("2")
7176 {
7177 let snap: Vec<Vec<f32>> = self
7178 .kv_cache
7179 .layers
7180 .iter()
7181 .map(|l| l.linear_state.clone())
7182 .collect();
7183 let attn_lens: Vec<usize> = self.kv_cache.layers.iter().map(|l| l.seq_len).collect();
7184 let toks: Vec<u32> = std::iter::once(t_next)
7185 .chain(drafts.iter().copied())
7186 .collect();
7187 let want_save = self.graph_want_logits;
7188 self.graph_want_logits = false;
7189 for (i, &t) in toks.iter().take(a + 1).enumerate() {
7190 let _ = self.forward_layers(&self.embed_single(t), next_pos + i, None);
7191 let _ = self.graph_logits.take();
7192 }
7193 self.graph_want_logits = want_save;
7194 let plain_states: Vec<Vec<f32>> = self
7195 .kv_cache
7196 .layers
7197 .iter()
7198 .map(|l| l.linear_state.clone())
7199 .collect();
7200 let (nkv, hd) = (self.num_kv_heads, self.head_dim);
7201 let mut rows = Vec::new();
7202 for (li, (l, n0)) in self
7203 .kv_cache
7204 .layers
7205 .iter_mut()
7206 .zip(attn_lens.iter())
7207 .enumerate()
7208 {
7209 let extra = l.seq_len.saturating_sub(*n0);
7210 if extra > 0 {
7211 let mut kk = Vec::new();
7212 let mut vv = Vec::new();
7213 for g in 0..nkv {
7214 kk.extend_from_slice(&l.head_keys(g)[n0 * hd..]);
7215 vv.extend_from_slice(&l.head_values(g)[n0 * hd..]);
7216 }
7217 rows.push((li, kk, vv));
7218 l.truncate_last(extra);
7219 crate::gpu_metal::kv_mirror_set_stored(self.graph_kv_id, li, *n0);
7220 }
7221 }
7222 for (l, st) in self.kv_cache.layers.iter_mut().zip(snap) {
7223 if l.linear_state.len() == st.len() {
7224 l.linear_state.copy_from_slice(&st);
7225 } else {
7226 l.linear_state = st;
7227 }
7228 }
7229 Some((plain_states, rows))
7230 } else {
7231 None
7232 };
7233 let warm_off = std::env::var("CMF_SPEC_WARM").is_ok_and(|v| v == "0");
7234 #[cfg(target_os = "macos")]
7243 let mut warm_pending: Option<MetalWarmPending> = None;
7244 #[cfg(target_os = "macos")]
7245 if metal_native {
7246 m.kv.truncate_last(k_spec.saturating_sub(1));
7247 if self.mtp_graph_mode == Some(true) {
7248 crate::gpu_metal::kv_mirror_set_stored(
7251 self.mtp_kv_id(),
7252 Self::MTP_LAYER_BASE,
7253 m.kv.seq_len,
7254 );
7255 if !warm_off && a > 0 {
7256 let pairs: Vec<(&[f32], u32)> = (0..a)
7257 .map(|j| {
7258 (
7259 &hiddens[j * self.hidden_size..(j + 1) * self.hidden_size],
7260 ids[j],
7261 )
7262 })
7263 .collect();
7264 warm_pending = self.mtp_warm_batch_submit(m, &pairs, next_pos);
7265 }
7266 }
7267 spec_stamp("c.wsub");
7268 }
7269 #[cfg(target_os = "macos")]
7271 if metal_native {
7272 if !self.metal_verify_commit(a) {
7275 self.clear_sequence_state();
7276 self.graph_failed
7277 .store(true, std::sync::atomic::Ordering::Relaxed);
7278 self.cancel
7279 .store(true, std::sync::atomic::Ordering::Relaxed);
7280 tracing::error!("Metal verify state/KV handoff failed after admission");
7281 return None;
7282 }
7283 if let Some((plain_states, rows)) = commit_ref {
7284 crate::gpu_metal::queue_fence();
7285 let _ = crate::gpu_metal::wait_replay();
7288 let (nkv, hd) = (self.num_kv_heads, self.head_dim);
7289 let mut worst_s = 0f32;
7290 let mut worst_li = 0usize;
7291 for (li, (l, ps)) in self.kv_cache.layers.iter().zip(&plain_states).enumerate() {
7292 if l.linear_state.len() != ps.len() || ps.is_empty() {
7293 continue;
7294 }
7295 let d = l
7296 .linear_state
7297 .iter()
7298 .zip(ps)
7299 .fold(0f32, |m, (x, y)| m.max((x - y).abs()));
7300 let n = ps.iter().fold(0f32, |m, y| m.max(y.abs()));
7301 let rel = d / n.max(1e-6);
7302 if rel > worst_s {
7303 worst_s = rel;
7304 worst_li = li;
7305 }
7306 }
7307 let mut worst_k = 0f32;
7308 for (li, kk, vv) in &rows {
7309 let l = &self.kv_cache.layers[*li];
7310 let n0 = l.seq_len - (kk.len() / (nkv * hd));
7311 let mut ck = Vec::new();
7312 let mut cv = Vec::new();
7313 for g in 0..nkv {
7314 ck.extend_from_slice(&l.head_keys(g)[n0 * hd..]);
7315 cv.extend_from_slice(&l.head_values(g)[n0 * hd..]);
7316 }
7317 if ck.len() == kk.len() {
7318 let dk = ck
7319 .iter()
7320 .zip(kk)
7321 .fold(0f32, |m, (x, y)| m.max((x - y).abs()));
7322 let dv = cv
7323 .iter()
7324 .zip(vv)
7325 .fold(0f32, |m, (x, y)| m.max((x - y).abs()));
7326 worst_k = worst_k.max(dk).max(dv);
7327 } else {
7328 eprintln!(
7329 "commit-check L{li}: kv row count mismatch {} vs {}",
7330 ck.len(),
7331 kk.len()
7332 );
7333 }
7334 }
7335 eprintln!(
7336 "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}"
7337 );
7338 }
7339 }
7340 if !metal_native && a + 1 < b {
7341 let expected_gdn_layers = self.graph_gdn_layer_count();
7342 if expected_gdn_layers > 0
7343 && !crate::gpu::gdn_spec_restore(self.graph_kv_id, a, next_pos, expected_gdn_layers)
7344 {
7345 self.clear_sequence_state();
7346 self.graph_failed
7347 .store(true, std::sync::atomic::Ordering::Relaxed);
7348 self.cancel
7349 .store(true, std::sync::atomic::Ordering::Relaxed);
7350 tracing::error!("GDN speculative restore failed after verify");
7351 return None;
7352 }
7353 }
7354 if !metal_native && !self.rewind_trunk_graph_mirrors(next_pos + a + 1) {
7355 self.clear_sequence_state();
7360 self.graph_failed
7361 .store(true, std::sync::atomic::Ordering::Relaxed);
7362 self.cancel
7363 .store(true, std::sync::atomic::Ordering::Relaxed);
7364 tracing::error!("trunk graph KV rewind failed after speculative verify");
7365 return None;
7366 }
7367 *accepted += a;
7368 if !metal_native {
7379 m.kv.truncate_last(k_spec.saturating_sub(1));
7381 }
7382 spec_stamp("c.trunc");
7383 if !metal_native
7384 && self.mtp_graph_mode == Some(true)
7385 && !self.rewind_mtp_graph_mirror(next_pos)
7386 {
7387 self.clear_sequence_state();
7391 self.graph_failed
7392 .store(true, std::sync::atomic::Ordering::Relaxed);
7393 self.cancel
7394 .store(true, std::sync::atomic::Ordering::Relaxed);
7395 tracing::error!("MTP graph mirror rewind failed after verify commit");
7396 return None;
7397 }
7398 if !warm_off && a > 0 {
7399 let mut warmed = false;
7402 #[cfg(target_os = "macos")]
7403 if metal_native && self.mtp_graph_mode == Some(true) {
7404 warmed = match warm_pending.take() {
7408 Some(p) => self.mtp_warm_batch_finish(m, p),
7409 None => false,
7410 };
7411 if !warmed {
7412 warmed = true;
7413 for j in 0..a {
7414 let row =
7415 hiddens[j * self.hidden_size..(j + 1) * self.hidden_size].to_vec();
7416 if self
7417 .mtp_step_metal(m, &row, ids[j], next_pos + j, false)
7418 .is_none()
7419 {
7420 warmed = false;
7421 break;
7422 }
7423 }
7424 }
7425 }
7426 if !warmed && self.mtp_graph_mode != Some(false) && !metal_native {
7427 let rows: Vec<Vec<f32>> = (0..a)
7428 .map(|j| hiddens[j * self.hidden_size..(j + 1) * self.hidden_size].to_vec())
7429 .collect();
7430 let pairs: Vec<(&[f32], u32)> = rows
7431 .iter()
7432 .zip(ids.iter())
7433 .map(|(r, &t)| (r.as_slice(), t))
7434 .collect();
7435 match self.mtp_warm_prefill_pairs(m, &pairs, next_pos) {
7436 Ok(()) => warmed = true,
7437 Err(err) => {
7438 tracing::error!("{err}");
7444 self.clear_sequence_state();
7445 self.graph_failed
7446 .store(true, std::sync::atomic::Ordering::Relaxed);
7447 self.cancel
7448 .store(true, std::sync::atomic::Ordering::Relaxed);
7449 return None;
7450 }
7451 }
7452 }
7453 if !warmed {
7454 for j in 0..a {
7455 let row = &hiddens[j * self.hidden_size..(j + 1) * self.hidden_size];
7456 let row = row.to_vec();
7457 self.mtp_warm(m, &row, ids[j], next_pos + j);
7458 }
7459 }
7460 }
7461 spec_stamp("c.warm");
7465 if let Some(c) = forced {
7466 self.spec_forced = Some(c);
7467 self.graph_logits = None;
7468 } else if greedy_dev && logits.is_empty() {
7469 self.spec_forced = Some(ids[a]);
7472 self.graph_logits = None;
7473 } else {
7474 let mut row = logits[a * lm_rows..(a + 1) * lm_rows].to_vec();
7475 row.resize(self.vocab_size, 0.0);
7476 if let Some(c) = self.final_softcap {
7477 for l in row.iter_mut() {
7478 *l = c * (*l / c).tanh();
7479 }
7480 }
7481 self.graph_logits = Some(row);
7482 }
7483 let new_hidden = hiddens[a * self.hidden_size..(a + 1) * self.hidden_size].to_vec();
7484 spec_stamp("c.row");
7485 if std::env::var("CMF_GRAPH_SPEC_TIME").is_ok() {
7491 let end = subs();
7492 eprintln!(
7493 "spec-round: draft {:.1} ms/{} sub | verify {:.1} ms/{} sub | \
7494 commit {:.1} ms/{} sub (accepted {a} of {k_spec}, full-head streak {})",
7495 t_draft.as_secs_f64() * 1e3,
7496 sub_draft - sub0,
7497 (t_verify - t_draft).as_secs_f64() * 1e3,
7498 sub_verify - sub_draft,
7499 (t_round.elapsed() - t_verify).as_secs_f64() * 1e3,
7500 end - sub_verify,
7501 self.draft_full_streak,
7502 );
7503 }
7504 if k_env.is_none() && !metal_native && !k_capped {
7509 let f = a as f32 / k_spec.max(1) as f32;
7513 self.spec_acc_ewma += 0.2 * (f - self.spec_acc_ewma);
7514 let mut k_next = k_spec;
7515 if self.spec_acc_ewma >= 0.75 && k_spec < k_max {
7516 k_next = k_spec + 1;
7517 } else if self.spec_acc_ewma < 0.4 && k_spec > 2 {
7518 k_next = k_spec - 1;
7519 }
7520 if k_next != k_spec {
7521 self.spec_acc_ewma = 0.6;
7522 if std::env::var("CMF_GRAPH_SPEC_TIME").is_ok() {
7523 eprintln!("spec-k: {k_spec} → {k_next}");
7524 }
7525 }
7526 self.spec_k_adapt = Some(k_next);
7527 }
7528 spec_stamp("end");
7529 Some((drafts[..a].to_vec(), next_pos + a + 1, new_hidden))
7530 }
7531
7532 pub fn measure_pair_fusion(&mut self, iters: usize) -> (f64, f64) {
7541 if !self.pair_supported() {
7542 return (0.0, 0.0);
7543 }
7544 let graph_env = std::env::var_os("CMF_GPU_WGPU_GRAPH");
7551 unsafe { std::env::set_var("CMF_GPU_WGPU_GRAPH", "0") };
7552 let emb1 = self.embed_single(1);
7553 let emb2 = self.embed_single(2);
7554 let pos = self.kv_cache.seq_len();
7555
7556 let t0 = std::time::Instant::now();
7557 for _ in 0..iters {
7558 let _ = self.forward_layers(&emb1, pos, None);
7559 let _ = self.forward_layers(&emb2, pos + 1, None);
7560 for l in &mut self.kv_cache.layers {
7561 l.truncate_last(2);
7562 }
7563 }
7564 let singles_ms = t0.elapsed().as_secs_f64() * 1000.0 / iters as f64;
7565
7566 let t1 = std::time::Instant::now();
7567 for _ in 0..iters {
7568 let _ = self.forward_pair(&emb1, &emb2, pos);
7569 for l in &mut self.kv_cache.layers {
7570 l.truncate_last(2);
7571 }
7572 }
7573 let pair_ms = t1.elapsed().as_secs_f64() * 1000.0 / iters as f64;
7574 match graph_env {
7575 Some(value) => unsafe { std::env::set_var("CMF_GPU_WGPU_GRAPH", value) },
7576 None => unsafe { std::env::remove_var("CMF_GPU_WGPU_GRAPH") },
7577 }
7578 (singles_ms, pair_ms)
7579 }
7580
7581 fn pair_supported(&self) -> bool {
7589 !self.weights.layers.is_empty()
7596 && self.g3n.is_none()
7597 && !self
7598 .weights
7599 .layers
7600 .iter()
7601 .any(|lw| matches!(&lw.attn, AttnKind::Mla(_) | AttnKind::Kda(_)))
7602 }
7603
7604 fn forward_pair(
7605 &mut self,
7606 emb1: &[f32],
7607 emb2: &[f32],
7608 position: usize,
7609 ) -> (Vec<f32>, Vec<f32>) {
7610 self.mimo_moe_prepare();
7613 let mut h1 = emb1.to_vec();
7614 let mut h2 = emb2.to_vec();
7615 let (_nkv, _hd, hs, _rd, eps) = (
7616 self.num_kv_heads,
7617 self.head_dim,
7618 self.hidden_size,
7619 self.rotary_dim,
7620 self.rms_eps,
7621 );
7622 let pool = self.pool.clone();
7623
7624 for li in 0..self.num_layers {
7625 let lw = &self.weights.layers[self.phys_layer(li)];
7626 inference::rms_norm_into(
7629 &h1,
7630 &lw.input_norm,
7631 self.rms_eps,
7632 self.norm_style,
7633 &mut self.ws.n1,
7634 );
7635 inference::rms_norm_into(
7636 &h2,
7637 &lw.input_norm,
7638 self.rms_eps,
7639 self.norm_style,
7640 &mut self.ws.n2,
7641 );
7642
7643 let (a1, a2) = match &lw.attn {
7644 AttnKind::Mla(_) => unreachable!("MLA has no MTP/pair path"),
7645 AttnKind::Kda(_) => unreachable!("KDA has no MTP/pair path"),
7646 AttnKind::Bounded(w) => {
7647 let rope = self
7650 .bounded_rope
7651 .clone()
7652 .expect("bounded layer without an installed rotation table");
7653 let cfg = crate::bounded::BoundedAttnCfg {
7654 num_heads: self.num_heads,
7655 num_kv_heads: self.num_kv_heads,
7656 head_dim: self.head_dim,
7657 hidden_size: hs,
7658 scale: self.attn_scale,
7659 rope: &rope,
7660 pool: pool.as_deref(),
7661 };
7662 let a1 = crate::bounded::bounded_attention(
7663 &self.ws.n1,
7664 w,
7665 &mut self.kv_cache.layers[li],
7666 &cfg,
7667 );
7668 let a2 = crate::bounded::bounded_attention(
7669 &self.ws.n2,
7670 w,
7671 &mut self.kv_cache.layers[li],
7672 &cfg,
7673 );
7674 (a1, a2)
7675 }
7676 AttnKind::Linear(w) => {
7677 let cfg = self.vmf_cfg.expect("linear layer without vmf_cfg");
7678 let layer = &mut self.kv_cache.layers[li];
7679 let (state, scratch) = (&mut layer.linear_state, &mut layer.linear_scratch);
7680 vmf_phase_pair(
7681 &self.ws.n1,
7682 &self.ws.n2,
7683 w,
7684 &cfg,
7685 state,
7686 scratch,
7687 self.pool.as_deref(),
7688 )
7689 }
7690 AttnKind::LinearGdn(w) => {
7691 let cfg = self.gdn_cfg.expect("gdn layer without gdn_cfg");
7692 let layer = &mut self.kv_cache.layers[li];
7693 let (state, scratch) = (&mut layer.linear_state, &mut layer.linear_scratch);
7694 gdn_pair(
7695 &self.ws.n1,
7696 &self.ws.n2,
7697 w,
7698 &cfg,
7699 state,
7700 scratch,
7701 self.pool.as_deref(),
7702 )
7703 }
7704 AttnKind::ShortConv(w) => {
7705 let cfg = self
7706 .short_conv_cfg
7707 .expect("short-conv layer without short_conv_cfg");
7708 let layer = &mut self.kv_cache.layers[li];
7709 let (state, scratch) = (&mut layer.linear_state, &mut layer.linear_scratch);
7710 short_conv_pair(
7711 &self.ws.n1,
7712 &self.ws.n2,
7713 w,
7714 &cfg,
7715 state,
7716 scratch,
7717 self.pool.as_deref(),
7718 )
7719 }
7720 AttnKind::Full {
7721 wq,
7722 wk,
7723 wv,
7724 wo,
7725 q_norm,
7726 k_norm,
7727 output_gate,
7728 softplus_gate,
7729 bias,
7730 } => {
7731 let inv_freq_l = self.layer_inv_freq(li);
7732 let (nkv_l, hd_l, rd_l) = self.layer_geom(li);
7733 let cfg = QwenAttnCfg {
7734 num_heads: self.layer_num_heads(li),
7735 num_kv_heads: nkv_l,
7736 head_dim: hd_l,
7737 hidden_size: hs,
7738 position,
7739 inv_freq: &inv_freq_l,
7740 rotary_dim: rd_l,
7741 scale: self.attn_scale,
7742 softcap: self.attn_softcap,
7743 window: self.layer_window(li),
7744 v_norm: self.attn_v_norm,
7745 qk_norm_after_rope: self.qk_norm_after_rope,
7746 gate_sigmoid: self.proj_gate_sigmoid,
7747 q_norm: q_norm.as_deref(),
7748 k_norm: k_norm.as_deref(),
7749 output_gate: *output_gate,
7750 softplus_gate: softplus_gate
7751 .as_ref()
7752 .map(|(gate, per_head)| (gate, *per_head)),
7753 rope_scale: self.layer_rope_scale(li),
7754 bias: bias
7755 .as_ref()
7756 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
7757 rms_eps: eps,
7758 norm_style: self.norm_style,
7759 pool: pool.as_deref(),
7760 v_head_dim: self.layer_v_dim(li),
7761 };
7762 attention::qwen_attention_pair(
7763 &self.ws.n1,
7764 &self.ws.n2,
7765 wq,
7766 wk,
7767 wv,
7768 wo,
7769 &mut self.kv_cache.layers[li],
7770 &cfg,
7771 )
7772 }
7773 };
7774 let (a1, a2) = match &self.weights.layers[self.phys_layer(li)].attn_out_norm {
7775 Some(w) => (
7776 inference::rms_norm(&a1, w, self.rms_eps, self.norm_style),
7777 inference::rms_norm(&a2, w, self.rms_eps, self.norm_style),
7778 ),
7779 None => (a1, a2),
7780 };
7781 for i in 0..self.hidden_size {
7782 h1[i] += a1[i];
7783 h2[i] += a2[i];
7784 }
7785 let (mut a1, mut a2) = (a1, a2);
7786 attention::recycle_buf(&mut a1);
7787 attention::recycle_buf(&mut a2);
7788
7789 let lw = &self.weights.layers[self.phys_layer(li)];
7790 inference::rms_norm_into(
7791 &h1,
7792 &lw.post_norm,
7793 self.rms_eps,
7794 self.norm_style,
7795 &mut self.ws.p1,
7796 );
7797 inference::rms_norm_into(
7798 &h2,
7799 &lw.post_norm,
7800 self.rms_eps,
7801 self.norm_style,
7802 &mut self.ws.p2,
7803 );
7804 let (f1, f2) = match &lw.ffn {
7805 FfnKind::DenseMoe(dm) => (
7808 dense_moe_ffn(
7809 dm,
7810 &self.ws.p1,
7811 &h1,
7812 self.rms_eps,
7813 self.norm_style,
7814 self.pool.as_deref(),
7815 ),
7816 dense_moe_ffn(
7817 dm,
7818 &self.ws.p2,
7819 &h2,
7820 self.rms_eps,
7821 self.norm_style,
7822 self.pool.as_deref(),
7823 ),
7824 ),
7825 FfnKind::Moe(m) if self.mimo_moe.is_dynamic(li, false) => (
7826 moe_ffn_banked(&mut self.mimo_moe, li, m, &self.ws.p1, self.pool.as_deref()),
7827 moe_ffn_banked(&mut self.mimo_moe, li, m, &self.ws.p2, self.pool.as_deref()),
7828 ),
7829 _ => ffn_forward_pair(
7830 &lw.ffn,
7831 &self.ws.p1,
7832 &self.ws.p2,
7833 self.pool.as_deref(),
7834 None,
7835 ),
7836 };
7837 let (f1, f2) = match &self.weights.layers[self.phys_layer(li)].ffn_out_norm {
7838 Some(w) => (
7839 inference::rms_norm(&f1, w, self.rms_eps, self.norm_style),
7840 inference::rms_norm(&f2, w, self.rms_eps, self.norm_style),
7841 ),
7842 None => (f1, f2),
7843 };
7844 for i in 0..self.hidden_size {
7845 h1[i] += f1[i];
7846 h2[i] += f2[i];
7847 }
7848 let (mut f1, mut f2) = (f1, f2);
7849 attention::recycle_buf(&mut f1);
7850 attention::recycle_buf(&mut f2);
7851 if let Some(sc) = self.weights.layers[self.phys_layer(li)].layer_scale {
7852 for i in 0..self.hidden_size {
7853 h1[i] *= sc;
7854 h2[i] *= sc;
7855 }
7856 }
7857 if self.is_loop_end(li) && li + 1 < self.num_layers {
7859 h1 = inference::rms_norm(
7860 &h1,
7861 &self.weights.final_norm,
7862 self.rms_eps,
7863 self.norm_style,
7864 );
7865 h2 = inference::rms_norm(
7866 &h2,
7867 &self.weights.final_norm,
7868 self.rms_eps,
7869 self.norm_style,
7870 );
7871 }
7872 }
7873 if self.o1_active() {
7879 self.commit_linear_scratch();
7880 }
7881 self.o1_progress();
7882 (h1, h2)
7883 }
7884
7885 fn commit_linear_scratch(&mut self) {
7887 for layer in &mut self.kv_cache.layers {
7888 if !layer.linear_scratch.is_empty() {
7889 std::mem::swap(&mut layer.linear_state, &mut layer.linear_scratch);
7890 layer.linear_scratch.clear();
7891 }
7892 }
7893 }
7894
7895 pub fn forward_ids(
7898 &mut self,
7899 ids: &[u32],
7900 task_mask: Option<&TaskMask>,
7901 ) -> Result<Vec<f32>, String> {
7902 #[cfg(target_os = "macos")]
7903 crate::gpu_metal::set_io_namespace(self.graph_kv_id);
7904 if ids.is_empty() {
7905 return Err("empty id sequence".to_string());
7906 }
7907 self.clear_sequence_state();
7908 self.check_forward_graph("forward_ids setup", 0)?;
7909 if task_mask.is_none() {
7910 self.o1_begin();
7911 }
7912 let mut hidden = vec![0.0f32; self.hidden_size];
7913 let mut pos = 0usize;
7914 if let Some(b) = &mut self.dsv41 {
7915 let pool = self.pool.clone();
7916 let mut logits = Vec::new();
7917 crate::dsv41::forward_chunk(
7918 &b.0,
7919 &b.1,
7920 &b.2,
7921 &mut b.3,
7922 ids,
7923 0,
7924 pool.as_deref(),
7925 &mut logits,
7926 );
7927 if let Err(err) = self.o1_seal_checked() {
7928 self.clear_sequence_state();
7929 return Err(err);
7930 }
7931 return Ok(logits);
7932 }
7933 if self.can_prefill_batched() && !self.graph_prefill_preferred() && ids.len() > 2 {
7941 let chunk = self.prefill_chunk();
7945 let hs = self.hidden_size;
7946 while pos < ids.len() {
7947 let end = (pos + chunk).min(ids.len());
7948 let hb = self.prefill_rows(&ids[pos..end], pos, task_mask)?;
7949 self.check_forward_graph("forward_ids batched prefill", end - 1)?;
7950 hidden.copy_from_slice(&hb[(end - pos - 1) * hs..]);
7951 pos = end;
7952 }
7953 }
7954 if task_mask.is_none()
7963 && !self.graph_prefill_preferred()
7964 && !std::env::var("CMF_PAIR").is_ok_and(|v| v == "0")
7965 && self.pair_supported()
7966 {
7967 while pos + 1 < ids.len() {
7968 let e1 = self.embed_single(ids[pos]);
7969 let e2 = self.embed_single(ids[pos + 1]);
7970 let (_, h2) = self.forward_pair(&e1, &e2, pos);
7971 self.check_forward_graph("forward_ids pair", pos + 1)?;
7972 self.commit_linear_scratch();
7973 hidden = h2;
7974 pos += 2;
7975 }
7976 }
7977 if task_mask.is_none() && pos == 0 && ids.len() > 1 {
7980 if let Some(lg) = self.embryo_prefill_chunked(ids, 0) {
7981 self.graph_logits = Some(lg);
7982 hidden = vec![0.0; self.hidden_size];
7983 pos = ids.len();
7984 }
7985 }
7986 while pos < ids.len() {
7987 hidden = self.forward_layers(&self.embed_single(ids[pos]), pos, task_mask);
7988 self.check_forward_graph("forward_ids", pos)?;
7989 pos += 1;
7990 }
7991 if let Some(logits) = self.graph_logits.take() {
7992 if let Err(err) = self.o1_seal_checked() {
7996 self.clear_sequence_state();
7997 return Err(err);
7998 }
7999 return Ok(logits);
8000 }
8001 if let Err(err) = self.o1_seal_checked() {
8005 self.clear_sequence_state();
8006 return Err(err);
8007 }
8008 let normed = inference::rms_norm(
8009 &hidden,
8010 &self.weights.final_norm,
8011 self.rms_eps,
8012 self.norm_style,
8013 );
8014 Ok(self.lm_head_forward(&normed))
8015 }
8016
8017 #[doc(hidden)]
8021 pub fn dsv41_serial_logits(&mut self, ids: &[u32]) -> Result<Vec<Vec<f32>>, String> {
8022 #[cfg(target_os = "macos")]
8023 crate::gpu_metal::set_io_namespace(self.graph_kv_id);
8024 if ids.is_empty() {
8025 return Err("empty id sequence".to_string());
8026 }
8027 self.clear_sequence_state();
8028 self.dsv41
8029 .as_ref()
8030 .ok_or_else(|| "dsv41 serial logits require a DeepSeek-V4.1 model".to_string())?;
8031 self.o1_begin();
8032 let rows = {
8033 let pool = self.pool.clone();
8034 let b = self
8035 .dsv41
8036 .as_mut()
8037 .expect("dsv41 checked above; state cannot change during forward");
8038 let mut rows = Vec::with_capacity(ids.len());
8039 for (position, &id) in ids.iter().enumerate() {
8040 let mut logits = Vec::new();
8041 crate::dsv41::forward_token(
8042 &b.0,
8043 &b.1,
8044 &b.2,
8045 &mut b.3,
8046 id,
8047 position,
8048 pool.as_deref(),
8049 &mut logits,
8050 );
8051 rows.push(logits);
8052 }
8053 rows
8054 };
8055 self.o1_seal();
8056 Ok(rows)
8057 }
8058
8059 pub fn ppl_ids(&mut self, ids: &[u32]) -> Result<f64, String> {
8066 let (nll, cnt) = self.nll_ids_from(ids, 0)?;
8067 Ok((nll / cnt.max(1) as f64).exp())
8068 }
8069
8070 pub fn probe_ffn_mass(&mut self, ids: &[u32]) -> Vec<Vec<f64>> {
8075 self.clear_sequence_state();
8076 FFN_PROBE.with(|p| {
8077 *p.borrow_mut() = Some(vec![vec![0f64; self.intermediate_size]; self.num_layers]);
8078 });
8079 crate::gpu::cpu_scope(|| {
8080 for (pos, &id) in ids.iter().enumerate() {
8081 let emb = self.embed_single(id);
8082 let _ = self.forward_layers(&emb, pos, None);
8083 }
8084 });
8085 self.clear_sequence_state();
8086 FFN_PROBE
8087 .with(|p| p.borrow_mut().take())
8088 .unwrap_or_default()
8089 }
8090
8091 pub fn probe_ffn_mass_batch(&mut self, ids: &[u32]) -> Result<Vec<Vec<f64>>, String> {
8095 if let Err(err) = self.nll_begin() {
8096 let _ = FFN_PROBE.with(|p| p.borrow_mut().take());
8100 self.nll_end();
8101 return Err(err);
8102 }
8103 FFN_PROBE.with(|p| {
8104 *p.borrow_mut() = Some(vec![vec![0f64; self.intermediate_size]; self.num_layers]);
8105 });
8106 let result: Result<(), String> = (|| {
8107 for chunk in ids.chunks(256) {
8108 if chunk.len() < 2 {
8109 continue;
8110 }
8111 self.nll_ids_masked(chunk, 0, None)?;
8112 }
8113 Ok(())
8114 })();
8115 self.nll_end();
8116 let probe = FFN_PROBE
8117 .with(|p| p.borrow_mut().take())
8118 .unwrap_or_default();
8119 match result {
8120 Ok(()) => Ok(probe),
8121 Err(err) => {
8122 drop(probe);
8123 Err(err)
8124 }
8125 }
8126 }
8127
8128 pub fn ppl_ids_masked(&mut self, ids: &[u32], mask: &TaskMask) -> Result<f64, String> {
8132 self.nll_begin()?;
8133 let result: Result<f64, String> = (|| {
8134 let mut nll = 0f64;
8135 let mut cnt = 0usize;
8136 let mut hidden = vec![0f32; self.hidden_size];
8137 for (pos, &id) in ids.iter().enumerate() {
8138 if pos > 0 {
8139 inference::rms_norm_into(
8140 &hidden,
8141 &self.weights.final_norm,
8142 self.rms_eps,
8143 self.norm_style,
8144 &mut self.ws.n1,
8145 );
8146 let mut logits = self.lm_head_forward(&self.ws.n1);
8147 let max = logits.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
8148 let sum: f64 = logits.iter().map(|&v| ((v - max) as f64).exp()).sum();
8149 let p = ((logits[id as usize] - max) as f64).exp() / sum.max(1e-300);
8150 nll -= p.max(1e-300).ln();
8151 cnt += 1;
8152 attention::recycle_buf(&mut logits);
8153 }
8154 let emb = self.embed_single(id);
8155 hidden = self.forward_layers(&emb, pos, Some(mask));
8156 self.nll_check_graph("masked serial forward", pos)?;
8157 let _ = self.graph_logits.take();
8161 }
8162 Ok((nll / cnt.max(1) as f64).exp())
8163 })();
8164 self.nll_end();
8165 result
8166 }
8167
8168 pub fn nll_ids_masked(
8187 &mut self,
8188 ids: &[u32],
8189 start: usize,
8190 task_mask: Option<&TaskMask>,
8191 ) -> Result<(f64, usize), String> {
8192 let task_mask = self.drop_open_mask(task_mask);
8193 self.nll_ids_inner(ids, start, task_mask)
8194 }
8195
8196 pub fn nll_ids_from(&mut self, ids: &[u32], start: usize) -> Result<(f64, usize), String> {
8197 self.nll_ids_inner(ids, start, None)
8198 }
8199
8200 fn nll_ids_inner(
8201 &mut self,
8202 ids: &[u32],
8203 start: usize,
8204 task_mask: Option<&TaskMask>,
8205 ) -> Result<(f64, usize), String> {
8206 self.nll_begin()?;
8207 let result: Result<(f64, usize), String> = (|| {
8208 let mut nll = 0f64;
8209 let mut cnt = 0usize;
8210 let (graph_quality, fused_head_quality) = nll_graph_policy(
8223 task_mask.is_none(),
8224 self.graph_prefill_preferred(),
8225 crate::gpu::q1_force(),
8226 );
8227 self.graph_head_required = fused_head_quality;
8228 self.graph_want_logits = fused_head_quality;
8229 #[cfg(target_os = "macos")]
8230 if graph_quality && std::env::var("CMF_METAL_BATCH_NLL").as_deref() != Ok("0") {
8231 match self.nll_batch_metal(ids, start) {
8232 MetalBatchNllOutcome::Completed(nll, count) => {
8233 return Ok((nll, count));
8234 }
8235 MetalBatchNllOutcome::Declined => {}
8236 MetalBatchNllOutcome::Failed(err) => return Err(err),
8237 }
8238 }
8239 let force_serial = std::env::var("CMF_NLL_SERIAL").as_deref() == Ok("1");
8244 if self.can_prefill_batched() && !graph_quality && !force_serial {
8245 const CHUNK: usize = 128;
8251 const LM_SUB: usize = 32;
8252 let n = ids.len().saturating_sub(1);
8253 let hs = self.hidden_size;
8254 let rows = self.weights.lm_head.rows();
8255 let mut pos = 0usize;
8256 let state_trace = std::env::var("CMF_STATE_TRACE").is_ok();
8257 while pos < n {
8258 let end = (pos + CHUNK).min(n);
8259 let bsz = end - pos;
8260 let hb = self.prefill_rows(&ids[pos..end], pos, task_mask)?;
8261 self.nll_check_graph("batched prefill", pos)?;
8262 if state_trace && end % 256 == 0 {
8263 self.trace_recurrent_state(end);
8264 }
8265 let mut k0 = 0usize;
8266 while k0 < bsz {
8267 let k1 = (k0 + LM_SUB).min(bsz);
8268 let sb = k1 - k0;
8269 if pos + k1 <= start {
8272 k0 = k1;
8273 continue;
8274 }
8275 let mut normed = vec![0.0f32; sb * hs];
8276 for k in 0..sb {
8277 let r = inference::rms_norm(
8278 &hb[(k0 + k) * hs..(k0 + k + 1) * hs],
8279 &self.weights.final_norm,
8280 self.rms_eps,
8281 self.norm_style,
8282 );
8283 normed[k * hs..(k + 1) * hs].copy_from_slice(&r);
8284 }
8285 let mut logits = vec![0.0f32; sb * rows];
8286 self.weights
8287 .lm_head
8288 .matmat(&normed, sb, &mut logits, self.pool.as_deref());
8289 for k in 0..sb {
8290 if pos + k0 + k < start {
8291 continue;
8292 }
8293 self.nll_check_graph("batched score row", pos + k0 + k)?;
8294 let lg = &mut logits[k * rows..k * rows + self.vocab_size.min(rows)];
8295 if let Some(mu) = self.logit_multiplier {
8296 for v in lg.iter_mut() {
8297 *v *= mu;
8298 }
8299 }
8300 if let Some(c) = self.final_softcap {
8304 for v in lg.iter_mut() {
8305 *v = c * (*v / c).tanh();
8306 }
8307 }
8308 if let Some(cm) = self.head_clusters.clone() {
8311 self.hierarchical_head_logprobs(
8312 &normed[k * hs..(k + 1) * hs],
8313 &cm,
8314 lg,
8315 );
8316 }
8317 let lg = &logits[k * rows..k * rows + self.vocab_size.min(rows)];
8318 let target = ids[pos + k0 + k + 1] as usize;
8319 let max = lg.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
8320 let lse: f64 = lg
8321 .iter()
8322 .map(|&v| ((v - max) as f64).exp())
8323 .sum::<f64>()
8324 .ln()
8325 + max as f64;
8326 nll += lse - lg[target] as f64;
8327 cnt += 1;
8328 if std::env::var("CMF_PPL_TRACE").is_ok() {
8329 let top = lg
8330 .iter()
8331 .enumerate()
8332 .max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
8333 .map(|(i, _)| i)
8334 .unwrap_or(0);
8335 eprintln!(
8336 "BTRACE pos {} target {} nll {:.4} top {} lg_t {:.3} lg_top {:.3}",
8337 pos + k0 + k,
8338 target,
8339 lse - lg[target] as f64,
8340 top,
8341 lg[target],
8342 lg[top]
8343 );
8344 }
8345 }
8346 k0 = k1;
8347 }
8348 pos = end;
8349 }
8350 return Ok((nll, cnt));
8351 }
8352 for pos in 0..ids.len().saturating_sub(1) {
8353 let hidden = self.forward_layers(&self.embed_single(ids[pos]), pos, task_mask);
8354 self.nll_check_graph("serial forward", pos)?;
8355 let out_of_band = self.graph_logits.take();
8363 if self.graph_head_required && out_of_band.is_none() {
8364 METAL_GRAPH_HEAD_MISS.fetch_add(
8365 1,
8366 std::sync::atomic::Ordering::Relaxed,
8367 );
8368 return Err(format!(
8369 "fused Metal graph head did not complete at NLL position {pos}"
8370 ));
8371 }
8372 if pos < start {
8373 continue;
8374 }
8375 let logits = match out_of_band {
8376 Some(lg) => lg,
8377 None => {
8378 let normed = inference::rms_norm(
8379 &hidden,
8380 &self.weights.final_norm,
8381 self.rms_eps,
8382 self.norm_style,
8383 );
8384 self.lm_head_forward(&normed)
8388 }
8389 };
8390 let target = ids[pos + 1] as usize;
8391 let max = logits.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
8392 let lse: f64 = logits
8393 .iter()
8394 .map(|&v| ((v - max) as f64).exp())
8395 .sum::<f64>()
8396 .ln()
8397 + max as f64;
8398 let tok_nll = lse - logits[target] as f64;
8399 if std::env::var("CMF_PPL_TRACE").is_ok() && pos < 48 {
8400 let top = logits
8401 .iter()
8402 .enumerate()
8403 .max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
8404 .map(|(i, _)| i)
8405 .unwrap_or(0);
8406 eprintln!(
8407 "pos {pos:3} tgt {target:6} nll {tok_nll:7.3} | top1 {top:6} lg[t]={:.2} lg[top]={:.2}",
8408 logits[target], logits[top]
8409 );
8410 }
8411 nll += tok_nll;
8412 cnt += 1;
8413 }
8414 Ok((nll, cnt))
8415 })();
8416 self.nll_end();
8417 result
8418 }
8419
8420 fn nll_from_hidden(&mut self, hidden: &[f32], target: u32, pos: usize) -> f64 {
8425 let normed = inference::rms_norm(
8426 hidden,
8427 &self.weights.final_norm,
8428 self.rms_eps,
8429 self.norm_style,
8430 );
8431 let mut logits = self.lm_head_forward(&normed);
8434 let target = target as usize;
8435 let max = logits.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
8436 let lse: f64 = logits
8437 .iter()
8438 .map(|&v| ((v - max) as f64).exp())
8439 .sum::<f64>()
8440 .ln()
8441 + max as f64;
8442 let tok_nll = lse - logits[target] as f64;
8443 if std::env::var("CMF_PPL_TRACE").is_ok() && pos < 48 {
8444 let top = logits
8445 .iter()
8446 .enumerate()
8447 .max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
8448 .map(|(i, _)| i)
8449 .unwrap_or(0);
8450 eprintln!(
8451 "pos {pos:3} tgt {target:6} nll {tok_nll:7.3} | top1 {top:6} lg[t]={:.2} lg[top]={:.2}",
8452 logits[target], logits[top]
8453 );
8454 }
8455 attention::recycle_buf(&mut logits);
8456 tok_nll
8457 }
8458
8459 fn trace_recurrent_state(&self, pos: usize) {
8467 let stats = |v: &[f32]| -> (f64, f64) {
8468 if v.is_empty() {
8469 return (0.0, 0.0);
8470 }
8471 let (mut ss, mut mx) = (0f64, 0f64);
8472 for &x in v {
8473 ss += (x as f64) * (x as f64);
8474 mx = mx.max((x as f64).abs());
8475 }
8476 ((ss / v.len() as f64).sqrt(), mx)
8477 };
8478 for (li, l) in self.kv_cache.layers.iter().enumerate() {
8479 let lw = &self.weights.layers[self.phys_layer(li)];
8480 let (kind, s_len) = match &lw.attn {
8481 AttnKind::Linear(_) => (
8482 "vmf",
8483 self.vmf_cfg.map(|c| c.state_len()).unwrap_or(0),
8484 ),
8485 AttnKind::LinearGdn(_) => (
8486 "gdn",
8487 self.gdn_cfg.map(|c| c.state_len()).unwrap_or(0),
8488 ),
8489 AttnKind::Bounded(_) => ("bounded", 0),
8490 AttnKind::Full { .. } => ("full", 0),
8491 _ => ("other", 0),
8492 };
8493 let (rms, max) = stats(&l.linear_state);
8494 let s_part = if kind == "vmf" {
8495 &l.linear_state[..s_len.min(l.linear_state.len())]
8496 } else if kind == "gdn" {
8497 let ring = l.linear_state.len().saturating_sub(s_len).min(l.linear_state.len());
8498 &l.linear_state[ring..]
8499 } else {
8500 &l.linear_state[..0]
8501 };
8502 let (s_rms, s_max) = stats(s_part);
8503 let (ring_rms, ring_len) = match &l.bounded {
8504 Some(b) => (stats(&b.ring_k).0, b.ring_k.len()),
8505 None => (0.0, 0),
8506 };
8507 eprintln!(
8508 "STATE pos={pos} layer={li} kind={kind} state_len={} rms={rms:.5} max={max:.4} \
8509 S_rms={s_rms:.5} S_max={s_max:.4} ring_k_rms={ring_rms:.5} ring_len={ring_len} kv_rows={}",
8510 l.linear_state.len(),
8511 l.seq_len
8512 );
8513 }
8514 }
8515
8516 pub fn nll_ids_o1(&mut self, ids: &[u32], prefill: usize) -> Result<(f64, usize), String> {
8534 self.nll_begin()?;
8539 let requested_prefix = (prefill > 0).then_some(prefill);
8540 self.o1_begin_with_prefix(requested_prefix);
8541 let n = ids.len().saturating_sub(1);
8542 let requested_start = prefill.min(n);
8543 let exact_end = if self.o1_active() {
8548 match requested_prefix {
8549 Some(requested) => self.o1_effective_boundary(requested),
8550 None => self
8551 .o1_cfg
8552 .as_ref()
8553 .and_then(|c| crate::nystrom::o1_deferred_boundary(c.w, c.sink)),
8554 }
8555 .unwrap_or(requested_start)
8556 .min(n)
8557 } else {
8558 requested_start
8559 };
8560 let mut nll = 0f64;
8561 let mut cnt = 0usize;
8562
8563 let mut pos = 0usize;
8567 if self.can_prefill_batched() {
8568 const CHUNK: usize = 128;
8569 while pos < exact_end {
8570 let end = (pos + CHUNK).min(exact_end);
8571 let hiddens = self.prefill_batch(&ids[pos..end], pos);
8572 if self
8573 .graph_failed
8574 .swap(false, std::sync::atomic::Ordering::Relaxed)
8575 {
8576 self.cancel
8577 .store(false, std::sync::atomic::Ordering::Relaxed);
8578 self.nll_end();
8579 return Err("GPU graph failed during O(1) NLL prefix".into());
8580 }
8581 for row in 0..end - pos {
8582 let score_pos = pos + row;
8583 if score_pos >= requested_start && score_pos < n {
8584 nll += self.nll_from_hidden(
8585 &hiddens[row * self.hidden_size..(row + 1) * self.hidden_size],
8586 ids[score_pos + 1],
8587 score_pos,
8588 );
8589 cnt += 1;
8590 }
8591 }
8592 pos = end;
8593 }
8594 } else {
8595 while pos < exact_end {
8596 let hidden = self.forward_layers(&self.embed_single(ids[pos]), pos, None);
8597 if self
8598 .graph_failed
8599 .swap(false, std::sync::atomic::Ordering::Relaxed)
8600 {
8601 self.cancel
8602 .store(false, std::sync::atomic::Ordering::Relaxed);
8603 self.nll_end();
8604 return Err("GPU graph failed during O(1) NLL prefix".into());
8605 }
8606 if pos >= requested_start && pos < n {
8607 nll += self.nll_from_hidden(&hidden, ids[pos + 1], pos);
8608 cnt += 1;
8609 }
8610 pos += 1;
8611 }
8612 }
8613 self.o1_seal_checked().map_err(|err| {
8614 self.nll_end();
8615 err
8616 })?;
8617
8618 let batch_k = std::env::var("CMF_BATCH_K")
8627 .ok()
8628 .and_then(|v| v.parse::<usize>().ok())
8629 .unwrap_or(0);
8630 let batch_admitted = batch_k > 0
8631 && self.can_prefill_batched()
8632 && self.o1_active()
8633 && std::env::var("CMF_O1_GPU").as_deref() == Ok("1")
8634 && (0..self.num_layers).all(|li| {
8635 let cache = &self.kv_cache.layers[self.phys_layer(li)];
8636 cache.o1.is_none() || cache.o1_views().is_some()
8637 });
8638 if std::env::var("CMF_GRAPH_PROF").is_ok() {
8639 eprintln!(
8640 "nll-batch: phase=post-seal admission={} requested_k={} scored_rows={}",
8641 batch_admitted,
8642 batch_k,
8643 n.saturating_sub(exact_end),
8644 );
8645 }
8646 let mut batch_completed = false;
8647 if batch_admitted && exact_end < n {
8648 let hs = self.hidden_size;
8649 let mut batch_pos = exact_end;
8650 while batch_pos < n {
8651 let end = (batch_pos + batch_k).min(n);
8652 let bk = end - batch_pos;
8653 let mut hiddens = vec![0.0f32; bk * hs];
8654 for (row, &id) in ids[batch_pos..end].iter().enumerate() {
8655 hiddens[row * hs..(row + 1) * hs].copy_from_slice(&self.embed_single(id));
8656 }
8657 let positions: Vec<usize> = (batch_pos..end).collect();
8658 let t_batch = std::time::Instant::now();
8659 let outcome = self.try_batch_graph_wgpu(&mut hiddens, &positions, bk, None);
8660 if std::env::var("CMF_GRAPH_PROF").is_ok() {
8661 let ms = t_batch.elapsed().as_secs_f64() * 1000.0;
8662 eprintln!(
8663 "nll-batch: phase=post-seal mode=o1 k={bk} pos={}..{} outcome={outcome:?} {ms:.1} ms ({:.1} tok/s)",
8664 batch_pos,
8665 end.saturating_sub(1),
8666 bk as f64 / (ms / 1000.0),
8667 );
8668 }
8669 if let Err(err) = self.nll_check_graph("batch graph", batch_pos) {
8670 self.nll_end();
8671 return Err(err);
8672 }
8673 match outcome {
8674 crate::gpu::BatchGraphOutcome::Completed => {
8675 batch_completed = true;
8676 for row in 0..bk {
8677 nll += self.nll_from_hidden(
8678 &hiddens[row * hs..(row + 1) * hs],
8679 ids[batch_pos + row + 1],
8680 batch_pos + row,
8681 );
8682 cnt += 1;
8683 }
8684 batch_pos = end;
8685 }
8686 crate::gpu::BatchGraphOutcome::Declined => {
8687 if batch_completed {
8688 self.nll_end();
8689 return Err(format!(
8690 "O(1) NLL batch declined after completed chunk at position {batch_pos}"
8691 ));
8692 }
8693 break;
8694 }
8695 crate::gpu::BatchGraphOutcome::Failed => {
8696 self.nll_end();
8697 return Err(format!(
8698 "O(1) NLL batch graph failed after admission at position {batch_pos}"
8699 ));
8700 }
8701 }
8702 }
8703 if batch_completed && cnt == n.saturating_sub(requested_start) {
8704 self.nll_end();
8705 return Ok((nll, cnt));
8706 }
8707 }
8708
8709 for pos in exact_end..n {
8714 let hidden = self.forward_layers(&self.embed_single(ids[pos]), pos, None);
8715 if self
8716 .graph_failed
8717 .swap(false, std::sync::atomic::Ordering::Relaxed)
8718 {
8719 self.cancel
8720 .store(false, std::sync::atomic::Ordering::Relaxed);
8721 self.nll_end();
8722 return Err(format!(
8723 "GPU graph failed during O(1) NLL serial scoring at position {pos}"
8724 ));
8725 }
8726 nll += self.nll_from_hidden(&hidden, ids[pos + 1], pos);
8727 cnt += 1;
8728 }
8729 self.nll_end();
8730 Ok((nll, cnt))
8731 }
8732
8733 pub fn calib_ids(&mut self, ids: &[u32], temps: &[f32]) -> (Vec<bool>, Vec<Vec<f32>>) {
8741 self.clear_sequence_state();
8742 let n = ids.len().saturating_sub(1);
8743 let mut correct = Vec::with_capacity(n);
8744 let mut pmax = Vec::with_capacity(n);
8745 for pos in 0..n {
8746 let emb = self.embed_single(ids[pos]);
8747 let hidden = self.forward_layers(&emb, pos, None);
8748 let logits = if let Some(logits) = self.graph_logits.take() {
8749 logits
8750 } else {
8751 let normed = inference::rms_norm(
8752 &hidden,
8753 &self.weights.final_norm,
8754 self.rms_eps,
8755 self.norm_style,
8756 );
8757 self.lm_head_forward(&normed)
8761 };
8762 let target = ids[pos + 1] as usize;
8763 let (mut amax, mut mval) = (0usize, f32::NEG_INFINITY);
8764 for (i, &v) in logits.iter().enumerate() {
8765 if v > mval {
8766 mval = v;
8767 amax = i;
8768 }
8769 }
8770 correct.push(amax == target);
8771 let row: Vec<f32> = temps
8772 .iter()
8773 .map(|&t| {
8774 let tt = t.max(1e-3);
8775 let s: f32 = logits.iter().map(|&v| ((v - mval) / tt).exp()).sum();
8776 1.0 / s.max(1e-12) })
8778 .collect();
8779 pmax.push(row);
8780 }
8781 self.clear_sequence_state();
8782 (correct, pmax)
8783 }
8784
8785 pub fn ppl_ids_dynamic(&mut self, ids: &[u32]) -> Result<(f64, usize), String> {
8792 if self.dyn_router.is_none() {
8793 return Ok((self.ppl_ids(ids)?, 0));
8794 }
8795 self.nll_begin()?;
8796 let saved_active = self.dyn_active;
8797 let mut router = self
8798 .dyn_router
8799 .take()
8800 .ok_or_else(|| "dynamic router disappeared before PPL scoring".to_string())?;
8801 router.reset();
8802 self.dyn_phi_seen = 0;
8803 let _ = self.set_active_skill(None);
8804
8805 let result: Result<(f64, usize), String> = (|| {
8806 let mut nll = 0f64;
8807 let mut cnt = 0usize;
8808 for pos in 0..ids.len().saturating_sub(1) {
8809 let hidden = self.forward_layers(&self.embed_single(ids[pos]), pos, None);
8810 self.nll_check_graph("dynamic serial forward", pos)?;
8811 let out_of_band = self.graph_logits.take();
8812 let mut logits = match out_of_band {
8813 Some(lg) => lg,
8814 None => {
8815 let normed = inference::rms_norm(
8816 &hidden,
8817 &self.weights.final_norm,
8818 self.rms_eps,
8819 self.norm_style,
8820 );
8821 self.lm_head_forward(&normed)
8825 }
8826 };
8827 let target = ids[pos + 1] as usize;
8828 let max = logits.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
8829 let lse: f64 = logits
8830 .iter()
8831 .map(|&v| ((v - max) as f64).exp())
8832 .sum::<f64>()
8833 .ln()
8834 + max as f64;
8835 let tok_nll = lse - logits[target] as f64;
8836 if std::env::var("CMF_PPL_TRACE").is_ok() && pos < 48 {
8837 let top = logits
8838 .iter()
8839 .enumerate()
8840 .max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
8841 .map(|(i, _)| i)
8842 .unwrap_or(0);
8843 eprintln!(
8844 "pos {pos:3} tgt {target:6} nll {tok_nll:7.3} | top1 {top:6} lg[t]={:.2} lg[top]={:.2}",
8845 logits[target], logits[top]
8846 );
8847 }
8848 nll += tok_nll;
8849 cnt += 1;
8850 attention::recycle_buf(&mut logits);
8851 let phi = self.dyn_phi_ema.clone();
8853 if let Some(new_active) = router.step(&phi, pos) {
8854 let _ = self.set_active_skill(new_active);
8855 }
8856 }
8857 Ok(((nll / cnt.max(1) as f64).exp(), router.switches.len()))
8858 })();
8859
8860 let _ = self.set_active_skill(saved_active);
8863 self.dyn_router = Some(router);
8864 self.nll_end();
8865 result
8866 }
8867
8868 pub fn probe_phi(&mut self, ids: &[u32], layer: usize) -> Vec<f32> {
8870 self.clear_sequence_state();
8871 let mut acc = vec![0f32; self.hidden_size];
8872 for (pos, &id) in ids.iter().enumerate() {
8873 let h = self.forward_layers_upto(&self.embed_single(id), pos, None, Some(layer));
8874 for (a, v) in acc.iter_mut().zip(&h) {
8875 *a += v;
8876 }
8877 }
8878 let n = ids.len().max(1) as f32;
8879 for a in acc.iter_mut() {
8880 *a /= n;
8881 }
8882 self.clear_sequence_state();
8883 acc
8884 }
8885
8886 pub fn probe_phi_span(
8900 &mut self,
8901 ids: &[u32],
8902 layer: usize,
8903 span: std::ops::Range<usize>,
8904 ) -> Vec<f32> {
8905 #[cfg(target_os = "macos")]
8906 crate::gpu_metal::set_io_namespace(self.graph_kv_id);
8907 let end = span.end.min(ids.len());
8908 let start = span.start.min(end);
8909 let reset = |p: &mut Self| p.clear_sequence_state();
8910 reset(self);
8911 let mut acc = vec![0f32; self.hidden_size];
8912 for (pos, &id) in ids[..end].iter().enumerate() {
8913 let h = self.forward_layers_upto(&self.embed_single(id), pos, None, Some(layer));
8914 if pos >= start {
8915 for (a, v) in acc.iter_mut().zip(&h) {
8916 *a += v;
8917 }
8918 }
8919 }
8920 let n = end - start;
8921 if n > 0 {
8922 let n = n as f32;
8923 for a in acc.iter_mut() {
8924 *a /= n;
8925 }
8926 }
8927 reset(self);
8928 acc
8929 }
8930
8931 pub fn decode_step_logits(&mut self, token: u32, position: usize) -> Vec<f32> {
8938 #[cfg(target_os = "macos")]
8939 crate::gpu_metal::set_io_namespace(self.graph_kv_id);
8940 self.graph_logits = None;
8941 let hidden = self.forward_layers(&self.embed_single(token), position, None);
8942 if let Some(logits) = self.graph_logits.take() {
8943 return logits;
8944 }
8945 inference::rms_norm_into(
8946 &hidden,
8947 &self.weights.final_norm,
8948 self.rms_eps,
8949 self.norm_style,
8950 &mut self.ws.n1,
8951 );
8952 self.lm_head_forward(&self.ws.n1)
8953 }
8954
8955 fn prefill_batch(&mut self, ids: &[u32], start_pos: usize) -> Vec<f32> {
8961 self.prefill_batch_masked(ids, start_pos, None)
8962 }
8963
8964 fn prefill_batch_masked(
8970 &mut self,
8971 ids: &[u32],
8972 start_pos: usize,
8973 task_mask: Option<&TaskMask>,
8974 ) -> Vec<f32> {
8975 self.prefill_batch_span(PrefillIn::Ids(ids), start_pos, task_mask, 0, usize::MAX)
8976 }
8977
8978 fn prefill_rows(
8986 &mut self,
8987 ids: &[u32],
8988 pos: usize,
8989 task_mask: Option<&TaskMask>,
8990 ) -> Result<Vec<f32>, String> {
8991 self.prefill_input_rows(PrefillIn::Ids(ids), pos, task_mask)
8992 }
8993
8994 fn prefill_input_rows(
8995 &mut self,
8996 input: PrefillIn<'_>,
8997 pos: usize,
8998 task_mask: Option<&TaskMask>,
8999 ) -> Result<Vec<f32>, String> {
9000 self.mimo_moe_prepare();
9001 let hs = self.hidden_size;
9002 let bk = match input {
9003 PrefillIn::Ids(ids) => ids.len(),
9004 PrefillIn::Hidden(rows) => rows.len() / hs,
9005 };
9006 #[cfg(not(target_os = "macos"))]
9007 if task_mask.is_none()
9008 && !self.o1_active()
9009 && bk > 1
9010 && (self.batch_prefix_prefill()
9011 || (self.verify_exact_moe
9012 && crate::gpu::enabled_here()
9013 && crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Decode)))
9014 {
9015 let mut hiddens = match input {
9016 PrefillIn::Hidden(rows) => rows.to_vec(),
9017 PrefillIn::Ids(ids) => ids.iter().flat_map(|&id| self.embed_single(id)).collect(),
9018 };
9019 let positions: Vec<usize> = (pos..pos + bk).collect();
9020 let mut run = 0usize;
9021 match self.try_batch_graph_wgpu_prefix(
9022 &mut hiddens,
9023 &positions,
9024 bk,
9025 None,
9026 Some(&mut run),
9027 ) {
9028 crate::gpu::BatchGraphOutcome::Completed => {
9029 let out = if run < self.num_layers {
9030 self.prefill_batch_span(
9031 PrefillIn::Hidden(&hiddens),
9032 pos,
9033 None,
9034 run,
9035 self.num_layers,
9036 )
9037 } else {
9038 hiddens
9039 };
9040 return if self.graph_failed.load(std::sync::atomic::Ordering::Relaxed) {
9041 Err("MiMo attention graph failed after admission".into())
9042 } else { Ok(out) };
9043 }
9044 crate::gpu::BatchGraphOutcome::Failed => {
9045 return Err("batched prefix prefill failed after admission".into());
9046 }
9047 crate::gpu::BatchGraphOutcome::Declined => {
9048 #[cfg(feature = "gpu")]
9050 self.pull_lagging_host_kv(0, self.num_layers, pos);
9051 }
9052 }
9053 }
9054 let out = self.prefill_batch_span(input, pos, task_mask, 0, usize::MAX);
9055 if self.graph_failed.load(std::sync::atomic::Ordering::Relaxed) {
9056 Err("batch tail graph failed after admission".into())
9057 } else { Ok(out) }
9058 }
9059
9060 fn prefill_batch_span(
9066 &mut self,
9067 input: PrefillIn<'_>,
9068 start_pos: usize,
9069 task_mask: Option<&TaskMask>,
9070 from: usize,
9071 upto_excl: usize,
9072 ) -> Vec<f32> {
9073 let hs = self.hidden_size;
9074 let b = match input {
9075 PrefillIn::Ids(ids) => ids.len(),
9076 PrefillIn::Hidden(hb) => hb.len() / hs,
9077 };
9078 let upto_excl = upto_excl.min(self.num_layers);
9079 let mut h: Vec<f32>;
9083 let mut h_ready;
9084 match input {
9085 PrefillIn::Ids(_) => {
9086 h = vec![0.0; b * hs];
9087 h_ready = false;
9088 }
9089 PrefillIn::Hidden(hb) => {
9090 h = hb.to_vec();
9091 h_ready = true;
9092 }
9093 }
9094 let fill_h = |h: &mut Vec<f32>, me: &Self| {
9095 if let PrefillIn::Ids(ids) = input {
9096 for (bi, &id) in ids.iter().enumerate() {
9097 let e = me.embed_single(id);
9098 h[bi * hs..(bi + 1) * hs].copy_from_slice(&e);
9099 }
9100 if let Ok(tp) = std::env::var("CMF_TRACE_POS") {
9101 if let Ok(t) = tp.parse::<usize>() {
9102 if t >= start_pos && t < start_pos + ids.len() {
9103 let bi = t - start_pos;
9104 let row = &h[bi * hs..(bi + 1) * hs];
9105 let n: f32 = row.iter().map(|x| x * x).sum::<f32>().sqrt();
9106 eprintln!(
9107 "BATCH pos {t} embed: id {} |h| = {n:.6} h0 {:.6} h1 {:.6} | b={} start={start_pos} ids[..8]={:?}",
9108 ids[bi],
9109 row[0],
9110 row[1],
9111 ids.len(),
9112 &ids[..ids.len().min(8)]
9113 );
9114 }
9115 }
9116 }
9117 }
9118 };
9119 let (_nkv, _hd, _rd, eps) = (
9120 self.num_kv_heads,
9121 self.head_dim,
9122 self.rotary_dim,
9123 self.rms_eps,
9124 );
9125 let pool = self.pool.clone();
9126 let norm_style = self.norm_style;
9127 self.mimo_moe_prepare();
9128 let automatic_gpu_prefix = self.automatic_gpu_prefix();
9129
9130 #[cfg(target_os = "macos")]
9131 let mut chunk_skip_until = 0usize;
9132 for li in from..upto_excl {
9133 let _capacity_tail = automatic_gpu_prefix
9134 .filter(|&prefix| {
9135 li >= prefix && !(self.verify_exact_moe && self.mimo_moe.is_dynamic(li, false))
9136 })
9137 .map(|_| crate::gpu::enter_cpu_scope());
9138 crate::gpu::set_layer(li as i64); #[cfg(target_os = "macos")]
9145 if task_mask.is_none() {
9146 if li < chunk_skip_until {
9147 continue;
9148 }
9149 if !h_ready && li == 0 && self.weights.embed_tokens.q8_row_parts().is_none() {
9155 fill_h(&mut h, self);
9156 h_ready = true;
9157 }
9158 let ids_for_embed = match input {
9159 PrefillIn::Ids(ids) => (!h_ready && li == 0).then_some(ids),
9160 PrefillIn::Hidden(_) => None,
9161 };
9162 let end = self.chunk_run_gpu(li, &mut h, b, start_pos, ids_for_embed, upto_excl);
9163 if end > li {
9164 h_ready = true;
9165 chunk_skip_until = end;
9166 if self.is_loop_end(end - 1) && end < self.num_layers {
9169 for bi in 0..b {
9170 let normed = inference::rms_norm(
9171 &h[bi * hs..(bi + 1) * hs],
9172 &self.weights.final_norm,
9173 eps,
9174 norm_style,
9175 );
9176 h[bi * hs..(bi + 1) * hs].copy_from_slice(&normed);
9177 }
9178 }
9179 continue;
9180 }
9181 }
9182 if !h_ready {
9183 fill_h(&mut h, self);
9184 h_ready = true;
9185 }
9186 if task_mask.is_none() && self.verify_exact_moe {
9187 let positions: Vec<_> = (start_pos..start_pos + b).collect();
9188 match self.mimo_graph_layer_rows(li, &mut h, &positions) {
9189 crate::gpu::BatchGraphOutcome::Completed => continue,
9190 crate::gpu::BatchGraphOutcome::Failed => return h,
9191 crate::gpu::BatchGraphOutcome::Declined => {},
9192 }
9193 }
9194 #[cfg(feature = "gpu")]
9195 self.pull_lagging_host_kv(li, li + 1, start_pos);
9196 let lw = &self.weights.layers[self.phys_layer(li)];
9197 let t_attn = std::time::Instant::now();
9198 match &lw.attn {
9200 AttnKind::Kda(w) => {
9201 let cfg = self.kda_cfg.expect("kda layer without kda_cfg");
9203 let mut normed = vec![0.0f32; b * hs];
9204 for bi in 0..b {
9205 inference::rms_norm_into(
9206 &h[bi * hs..(bi + 1) * hs],
9207 &lw.input_norm,
9208 eps,
9209 norm_style,
9210 &mut normed[bi * hs..(bi + 1) * hs],
9211 );
9212 }
9213 let attn = crate::linear_core::kda_forward_batch(
9214 &normed,
9215 b,
9216 w,
9217 &cfg,
9218 &mut self.kv_cache.layers[li].linear_state,
9219 pool.as_deref(),
9220 );
9221 for (dst, &a) in h.iter_mut().zip(&attn) {
9222 *dst += a;
9223 }
9224 }
9225 AttnKind::LinearGdn(w) => {
9226 let cfg = self.gdn_cfg.expect("gdn layer without gdn_cfg");
9228 let mut normed = vec![0.0f32; b * hs];
9229 for bi in 0..b {
9230 let r = inference::rms_norm(
9231 &h[bi * hs..(bi + 1) * hs],
9232 &lw.input_norm,
9233 eps,
9234 norm_style,
9235 );
9236 normed[bi * hs..(bi + 1) * hs].copy_from_slice(&r);
9237 }
9238 let attn = crate::linear_core::gdn_forward_batch(
9239 &normed,
9240 b,
9241 w,
9242 &cfg,
9243 &mut self.kv_cache.layers[li].linear_state,
9244 pool.as_deref(),
9245 );
9246 for (dst, &a) in h.iter_mut().zip(&attn) {
9247 *dst += a;
9248 }
9249 }
9250 AttnKind::ShortConv(w) => {
9251 let cfg = self
9254 .short_conv_cfg
9255 .expect("short-conv layer without short_conv_cfg");
9256 let mut normed = vec![0.0f32; b * hs];
9257 for bi in 0..b {
9258 inference::rms_norm_into(
9259 &h[bi * hs..(bi + 1) * hs],
9260 &lw.input_norm,
9261 eps,
9262 norm_style,
9263 &mut normed[bi * hs..(bi + 1) * hs],
9264 );
9265 }
9266 let attn = short_conv_forward_batch(
9267 &normed,
9268 b,
9269 w,
9270 &cfg,
9271 &mut self.kv_cache.layers[li].linear_state,
9272 pool.as_deref(),
9273 );
9274 for (dst, &a) in h.iter_mut().zip(&attn) {
9275 *dst += a;
9276 }
9277 }
9278 AttnKind::Mla(w) => {
9279 let inv_freq_l = self.layer_inv_freq(li);
9282 let rs = self.layer_rope_scale(li);
9283 let mut normed = vec![0.0f32; hs];
9284 for bi in 0..b {
9285 inference::rms_norm_into(
9286 &h[bi * hs..(bi + 1) * hs],
9287 &lw.input_norm,
9288 eps,
9289 norm_style,
9290 &mut normed,
9291 );
9292 let ao = mla_attention(
9293 w,
9294 &normed,
9295 &mut self.kv_cache.layers[li],
9296 start_pos + bi,
9297 &inv_freq_l,
9298 rs,
9299 eps,
9300 pool.as_deref(),
9301 );
9302 for (dst, &a) in h[bi * hs..(bi + 1) * hs].iter_mut().zip(&ao) {
9303 *dst += a;
9304 }
9305 }
9306 }
9307 AttnKind::Full {
9308 wq,
9309 wk,
9310 wv,
9311 wo,
9312 q_norm,
9313 k_norm,
9314 output_gate,
9315 softplus_gate,
9316 bias,
9317 } => {
9318 let mut normed = vec![0.0f32; b * hs];
9322 for bi in 0..b {
9323 inference::rms_norm_into(
9324 &h[bi * hs..(bi + 1) * hs],
9325 &lw.input_norm,
9326 eps,
9327 norm_style,
9328 &mut normed[bi * hs..(bi + 1) * hs],
9329 );
9330 }
9331 let inv_freq_l = self.layer_inv_freq(li);
9332 let (nkv_l, hd_l, rd_l) = self.layer_geom(li);
9333 let cfg = QwenAttnCfg {
9334 num_heads: self.layer_num_heads(li),
9335 num_kv_heads: nkv_l,
9336 head_dim: hd_l,
9337 hidden_size: hs,
9338 position: start_pos,
9339 inv_freq: &inv_freq_l,
9340 rotary_dim: rd_l,
9341 scale: self.attn_scale,
9342 softcap: self.attn_softcap,
9343 window: self.layer_window(li),
9344 v_norm: self.attn_v_norm,
9345 qk_norm_after_rope: self.qk_norm_after_rope,
9346 gate_sigmoid: self.proj_gate_sigmoid,
9347 q_norm: q_norm.as_deref(),
9348 k_norm: k_norm.as_deref(),
9349 output_gate: *output_gate,
9350 softplus_gate: softplus_gate
9351 .as_ref()
9352 .map(|(gate, per_head)| (gate, *per_head)),
9353 rope_scale: self.layer_rope_scale(li),
9354 bias: bias
9355 .as_ref()
9356 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
9357 rms_eps: eps,
9358 norm_style,
9359 pool: pool.as_deref(),
9360 v_head_dim: self.layer_v_dim(li),
9361 };
9362 let mut attn = attention::qwen_attention_batch(
9363 &normed,
9364 b,
9365 wq,
9366 wk,
9367 wv,
9368 wo,
9369 &mut self.kv_cache.layers[li],
9370 &cfg,
9371 );
9372 if let Some(w) = &lw.attn_out_norm {
9373 for bi in 0..b {
9374 inference::rms_norm_into(
9375 &attn[bi * hs..(bi + 1) * hs],
9376 w,
9377 eps,
9378 norm_style,
9379 &mut normed[bi * hs..(bi + 1) * hs],
9380 );
9381 }
9382 attn.copy_from_slice(&normed);
9383 }
9384 for (dst, &a) in h.iter_mut().zip(&attn) {
9385 *dst += a;
9386 }
9387 }
9388 AttnKind::Bounded(w) => {
9389 let mut normed = vec![0.0f32; b * hs];
9392 for bi in 0..b {
9393 inference::rms_norm_into(
9394 &h[bi * hs..(bi + 1) * hs],
9395 &lw.input_norm,
9396 eps,
9397 norm_style,
9398 &mut normed[bi * hs..(bi + 1) * hs],
9399 );
9400 }
9401 let rope = self
9402 .bounded_rope
9403 .clone()
9404 .expect("bounded layer without an installed rotation table");
9405 let cfg = crate::bounded::BoundedAttnCfg {
9406 num_heads: self.num_heads,
9407 num_kv_heads: self.num_kv_heads,
9408 head_dim: self.head_dim,
9409 hidden_size: hs,
9410 scale: self.attn_scale,
9411 rope: &rope,
9412 pool: pool.as_deref(),
9413 };
9414 let mut attn = crate::bounded::bounded_attention_batch(
9415 &normed,
9416 b,
9417 w,
9418 &mut self.kv_cache.layers[li],
9419 &cfg,
9420 );
9421 if let Some(wn) = &lw.attn_out_norm {
9422 for bi in 0..b {
9423 inference::rms_norm_into(
9424 &attn[bi * hs..(bi + 1) * hs],
9425 wn,
9426 eps,
9427 norm_style,
9428 &mut normed[bi * hs..(bi + 1) * hs],
9429 );
9430 }
9431 attn.copy_from_slice(&normed);
9432 }
9433 for (dst, &a) in h.iter_mut().zip(&attn) {
9434 *dst += a;
9435 }
9436 attention::recycle_buf(&mut attn);
9437 }
9438 AttnKind::Linear(w) => {
9439 for bi in 0..b {
9440 let normed = inference::rms_norm(
9441 &h[bi * hs..(bi + 1) * hs],
9442 &lw.input_norm,
9443 eps,
9444 norm_style,
9445 );
9446 vmf_phase_forward(
9447 &normed,
9448 w,
9449 &self.vmf_cfg.expect("linear layer without vmf_cfg"),
9450 &mut self.kv_cache.layers[li].linear_state,
9451 pool.as_deref(),
9452 )
9453 .iter()
9454 .enumerate()
9455 .for_each(|(i, &a)| h[bi * hs + i] += a);
9456 }
9457 }
9458 }
9459
9460 let lw = &self.weights.layers[self.phys_layer(li)];
9462 let mut post = vec![0.0f32; b * hs];
9463 for bi in 0..b {
9464 let r =
9465 inference::rms_norm(&h[bi * hs..(bi + 1) * hs], &lw.post_norm, eps, norm_style);
9466 post[bi * hs..(bi + 1) * hs].copy_from_slice(&r);
9467 }
9468 let mask_row = task_mask
9471 .filter(|m| m.ffn_active_count(li) < self.intermediate_size)
9472 .and_then(|m| m.ffn_masks.get(li))
9473 .map(|v| v.as_slice());
9474 let attn_ns = t_attn.elapsed().as_nanos() as u64;
9475 let t_ffn = std::time::Instant::now();
9476 let mut ffn = match &lw.ffn {
9477 FfnKind::Dense(d) if !d.segs.is_empty() => {
9478 tube_ffn(d, &post, b, pool.as_deref(), mask_row)
9479 }
9480 FfnKind::Dense(d) => dense_ffn_batch(d, &post, b, pool.as_deref(), mask_row),
9481 FfnKind::Moe(m) if self.verify_exact_moe && self.mimo_moe.is_dynamic(li, false) => {
9482 moe_ffn_banked_rows(&mut self.mimo_moe, li, m, &post, b, hs, pool.as_deref())
9483 }
9484 FfnKind::Moe(m) if self.verify_exact_moe => {
9485 moe_ffn_rows_exact(m, &post, b, hs, pool.as_deref())
9486 }
9487 FfnKind::Moe(m) if self.mimo_moe.is_dynamic(li, false) => {
9490 let before = m.stats.borrow().clone();
9491 let out = crate::gpu::cpu_scope(|| {
9492 moe_ffn_batch(m, &post, b, hs, pool.as_deref(), None)
9493 });
9494 self.mimo_moe.prime(li, m, &before);
9495 out
9496 }
9497 FfnKind::Moe(m) => moe_ffn_batch(m, &post, b, hs, pool.as_deref(), None),
9498 FfnKind::DenseMoe(dm) => {
9501 let mut out = vec![0.0f32; b * hs];
9502 for bi in 0..b {
9503 let r = dense_moe_ffn(
9504 dm,
9505 &post[bi * hs..(bi + 1) * hs],
9506 &h[bi * hs..(bi + 1) * hs],
9507 eps,
9508 norm_style,
9509 pool.as_deref(),
9510 );
9511 out[bi * hs..(bi + 1) * hs].copy_from_slice(&r);
9512 }
9513 out
9514 }
9515 };
9516 if prefill_prof_on() {
9517 PREFILL_SPLIT[0].fetch_add(attn_ns, std::sync::atomic::Ordering::Relaxed);
9518 PREFILL_SPLIT[1].fetch_add(t_ffn.elapsed().as_nanos() as u64, std::sync::atomic::Ordering::Relaxed);
9519 }
9520 if let Some(w) = &lw.ffn_out_norm {
9521 for bi in 0..b {
9522 inference::rms_norm_into(
9523 &ffn[bi * hs..(bi + 1) * hs],
9524 w,
9525 eps,
9526 norm_style,
9527 &mut post[bi * hs..(bi + 1) * hs],
9528 );
9529 }
9530 ffn.copy_from_slice(&post);
9531 }
9532 for (dst, &f) in h.iter_mut().zip(&ffn) {
9533 *dst += f;
9534 }
9535 if let Some(sc) = lw.layer_scale {
9536 for v in h.iter_mut() {
9537 *v *= sc;
9538 }
9539 }
9540 if self.layer_dump.is_some() {
9542 for bi in 0..b {
9543 self.dump_layer_row(start_pos + bi, li, &h[bi * hs..(bi + 1) * hs]);
9544 }
9545 }
9546 if let Ok(tp) = std::env::var("CMF_TRACE_POS") {
9547 if let Ok(t) = tp.parse::<usize>() {
9548 if t >= start_pos && t < start_pos + b {
9549 let bi = t - start_pos;
9550 let row = &h[bi * hs..(bi + 1) * hs];
9551 let n: f32 = row.iter().map(|x| x * x).sum::<f32>().sqrt();
9552 eprintln!(
9553 "BATCH pos {t} after layer {li}: |h| = {n:.6} h0 {:.6} h1 {:.6}",
9554 row[0], row[1]
9555 );
9556 }
9557 }
9558 }
9559 if std::env::var("CMF_DEBUG_LAYERS").is_ok() {
9563 let row = &h[(b - 1) * hs..b * hs];
9564 let rms =
9565 (row.iter().map(|&v| (v as f64) * (v as f64)).sum::<f64>() / hs as f64).sqrt();
9566 let mx = row.iter().fold(0f32, |m, &v| m.max(v.abs()));
9567 eprintln!(
9568 "layer {li:>3} {:>10} ffn={:<5} rms={rms:>12.4} max={mx:>12.4}",
9569 match &self.weights.layers[self.phys_layer(li)].attn {
9570 AttnKind::LinearGdn(_) => "gdn",
9571 AttnKind::Linear(_) => "vmf",
9572 AttnKind::ShortConv(_) => "conv",
9573 _ => "attn",
9574 },
9575 match &lw.ffn {
9576 FfnKind::Moe(_) => "moe",
9577 FfnKind::Dense(_) => "dense",
9578 FfnKind::DenseMoe(_) => "dense+moe",
9579 },
9580 );
9581 }
9582 if self.is_loop_end(li) && li + 1 < self.num_layers {
9584 for bi in 0..b {
9585 let normed = inference::rms_norm(
9586 &h[bi * hs..(bi + 1) * hs],
9587 &self.weights.final_norm,
9588 eps,
9589 norm_style,
9590 );
9591 h[bi * hs..(bi + 1) * hs].copy_from_slice(&normed);
9592 }
9593 }
9594 if std::env::var("CMF_TRACE_H").is_ok() {
9595 let n = h[..hs].iter().map(|v| v.abs()).sum::<f32>() / hs as f32;
9596 let mx = h[..hs].iter().fold(0.0f32, |a, &v| a.max(v.abs()));
9597 eprintln!(
9598 "layer {li}: mean|h|={n:.4} max|h|={mx:.2} scale={:?}",
9599 lw.layer_scale
9600 );
9601 }
9602 }
9603 crate::gpu::set_layer(-1); if prefill_prof_on() {
9605 eprintln!(
9606 "prefill-split: attention {:.1} ms, ffn {:.1} ms (cumulative)",
9607 PREFILL_SPLIT[0].load(std::sync::atomic::Ordering::Relaxed) as f64 / 1e6,
9608 PREFILL_SPLIT[1].load(std::sync::atomic::Ordering::Relaxed) as f64 / 1e6
9609 );
9610 }
9611 self.o1_progress();
9616 h
9617 }
9618
9619 fn embed_single(&self, id: u32) -> Vec<f32> {
9621 let mut out = vec![0.0f32; self.hidden_size];
9622 if (id as usize) < self.weights.embed_tokens.rows() {
9623 self.weights.embed_tokens.row_f32(id as usize, &mut out);
9624 }
9625 if self.embed_multiplier != 1.0 {
9626 for v in out.iter_mut() {
9627 *v *= self.embed_multiplier;
9628 }
9629 }
9630 if self.dsv4.is_some()
9634 || self.dsv41.is_some()
9635 || self.qwen4_exp.is_some()
9636 {
9637 let mut v = vec![0.0f32; self.hidden_size.max(1)];
9638 v[0] = id as f32;
9639 return v;
9640 }
9641 if let Some(b) = &self.g3n {
9644 return b.0.extend_embedding(id, &out, self.pool.as_deref());
9645 }
9646 out
9647 }
9648
9649 #[cfg(target_os = "macos")]
9655 fn chunk_run_gpu(
9656 &mut self,
9657 li0: usize,
9658 h: &mut [f32],
9659 b: usize,
9660 pos0: usize,
9661 embed_ids: Option<&[u32]>,
9662 cap: usize,
9663 ) -> usize {
9664 if !crate::gpu::enabled_here()
9668 || std::env::var("CMF_GPU_CHUNK")
9669 .map(|v| v == "0")
9670 .unwrap_or(false)
9671 || b < 32
9672 || (self.swa.is_some() && !self.metal_graph_swa())
9675 || self.global_attn.is_some()
9676 || (self.graph_attn_decline_reason().is_some() && !self.metal_graph_swa())
9678 || self.o1_active()
9681 || self.attn_v_norm
9682 || (self.attn_scale - 1.0 / (self.head_dim as f32).sqrt()).abs() > 1e-9
9683 {
9684 return li0;
9685 }
9686 let Some(model) = self.model.clone() else {
9687 return li0;
9688 };
9689 let (nh, nkv, hd, hs) = (
9690 self.num_heads,
9691 self.num_kv_heads,
9692 self.head_dim,
9693 self.hidden_size,
9694 );
9695 let loop_end = if self.loop_final_norm {
9699 ((li0 / self.physical_layers) + 1) * self.physical_layers
9700 } else {
9701 self.num_layers
9702 };
9703 let mut layers: Vec<crate::gpu_metal::ChunkLayer> = Vec::new();
9704 let mut stored_at: Vec<usize> = Vec::new();
9705 let run_end = self.num_layers.min(loop_end).min(cap);
9706 let tables: Vec<std::sync::Arc<Vec<f32>>> =
9709 (li0..run_end).map(|li| self.layer_inv_freq(li)).collect();
9710 for li in li0..run_end {
9711 let lw = &self.weights.layers[self.phys_layer(li)];
9712 if lw.attn_out_norm.is_some() || lw.ffn_out_norm.is_some() || lw.layer_scale.is_some() {
9713 break;
9714 }
9715 let AttnKind::Full {
9716 wq,
9717 wk,
9718 wv,
9719 wo,
9720 q_norm,
9721 k_norm,
9722 output_gate: false,
9723 softplus_gate,
9724 bias,
9725 } = &lw.attn
9726 else {
9727 break;
9728 };
9729 let head_gate = match softplus_gate {
9731 None => None,
9732 Some((g, true)) if self.proj_gate_sigmoid => match g.f32_parts() {
9733 Some((d, r, c)) if r == nh && c == hs => Some(d),
9734 _ => break,
9735 },
9736 Some(_) => break,
9737 };
9738 let FfnKind::Dense(d) = &lw.ffn else { break };
9739 if !matches!(d.act, Act::Silu | Act::Gelu) || !d.segs.is_empty() {
9740 break;
9741 }
9742 fn cw(t: &QTensor) -> Option<(usize, usize, usize, &[f32])> {
9747 t.q8_row_parts()
9748 .or_else(|| t.q4t_parts().map(|(i, r, c)| (i, r, c, &[][..])))
9749 .or_else(|| t.q4tp_parts().map(|(i, r, c)| (i, r, c, &[][..])))
9750 }
9751 let parts = (
9752 cw(wq),
9753 cw(wk),
9754 cw(wv),
9755 cw(wo),
9756 cw(&d.gate_proj),
9757 cw(&d.up_proj),
9758 cw(&d.down_proj),
9759 );
9760 let (Some(pq), Some(pk), Some(pv), Some(po), Some(pg), Some(pu), Some(pd)) = parts
9761 else {
9762 break;
9763 };
9764 let layer = &self.kv_cache.layers[li];
9765 if layer.mode != crate::kv_cache::KvMode::F32 || layer.o1.is_some() {
9766 break;
9767 }
9768 stored_at.push(layer.head_len(0));
9769 layers.push(crate::gpu_metal::ChunkLayer {
9770 model: &model,
9771 kv_id: self.graph_kv_id,
9772 layer: li,
9773 wq: pq,
9774 wk: pk,
9775 wv: pv,
9776 wo: po,
9777 gate: pg,
9778 up: pu,
9779 down: pd,
9780 input_norm: &lw.input_norm,
9781 post_norm: &lw.post_norm,
9782 bias: bias
9783 .as_ref()
9784 .map(|(a, bb, cc)| (a.as_slice(), bb.as_slice(), cc.as_slice())),
9785 q_norm: q_norm.as_deref(),
9786 k_norm: k_norm.as_deref(),
9787 inv_freq: &tables[li - li0],
9788 rd: self.layer_geom(li).2,
9789 nh,
9790 nkv,
9791 hd,
9792 hs,
9793 inter: d.gate_proj.rows(),
9794 gemma: matches!(self.norm_style, cortiq_core::NormStyle::Gemma),
9795 late_qk_norm: self.qk_norm_after_rope,
9796 eps: self.rms_eps as f32,
9797 window: self.layer_window(li),
9798 head_gate,
9799 gelu: d.act == Act::Gelu,
9800 });
9801 }
9802 if layers.is_empty() {
9803 return li0;
9804 }
9805 let row = nkv * hd;
9806 let mut store: Vec<(Vec<f32>, Vec<f32>, Vec<f32>)> = stored_at
9807 .iter()
9808 .map(|&st| (vec![0f32; b * row], vec![0f32; b * row], vec![0f32; st + b]))
9809 .collect();
9810 let mut io: Vec<crate::gpu_metal::ChunkIo> = Vec::with_capacity(layers.len());
9811 for (i, (ok, ov, oi)) in store.iter_mut().enumerate() {
9812 let li = layers[i].layer;
9813 let layer = &self.kv_cache.layers[li];
9814 io.push(crate::gpu_metal::ChunkIo {
9815 cpu_stored: stored_at[i],
9816 cpu_k: (0..nkv).map(|g| layer.head_keys(g)).collect(),
9817 cpu_v: (0..nkv).map(|g| layer.head_values(g)).collect(),
9818 out_k: ok,
9819 out_v: ov,
9820 imp: oi,
9821 });
9822 }
9823 let n_run = layers.len();
9824 let last = layers.last().map(|l| l.layer + 1).unwrap_or(li0);
9825 let ep = embed_ids.and_then(|ids| {
9828 self.weights
9829 .embed_tokens
9830 .q8_row_parts()
9831 .map(|(idx, rows, _c, rs)| crate::gpu_metal::ChunkEmbed {
9832 idx,
9833 rows,
9834 row_scale: rs,
9835 ids,
9836 mult: self.embed_multiplier,
9837 })
9838 });
9839 if embed_ids.is_some() && ep.is_none() {
9840 return li0;
9841 }
9842 if !crate::gpu_metal::chunk_run_gpu(&layers, &mut io, h, b, pos0, ep.as_ref()) {
9843 return li0;
9844 }
9845 drop(io);
9846 drop(layers);
9847 for (i, (ok, ov, oi)) in store.iter().enumerate().take(n_run) {
9850 let li = li0 + i;
9851 let layer = &mut self.kv_cache.layers[li];
9852 for bi in 0..b {
9853 layer.append(
9854 &ok[bi * row..(bi + 1) * row],
9855 &ov[bi * row..(bi + 1) * row],
9856 &[],
9857 );
9858 }
9859 layer.accumulate_imp(oi);
9860 }
9861 last
9862 }
9863
9864 fn layer_is_local(&self, li: usize) -> bool {
9867 if let Some(layers) = &self.sliding_layers {
9868 return layers.get(li).copied().unwrap_or(false);
9869 }
9870 match self.swa {
9871 Some((_, pattern)) => (li + 1) % pattern.max(1) != 0,
9872 None => false,
9873 }
9874 }
9875
9876 fn layer_inv_freq(&self, li: usize) -> std::sync::Arc<Vec<f32>> {
9879 if self.layer_is_local(li) {
9880 if let Some(f) = &self.inv_freq_local {
9881 return f.clone();
9882 }
9883 } else if let Some(f) = &self.inv_freq_global {
9884 return f.clone();
9885 }
9886 self.inv_freq.clone()
9887 }
9888
9889 fn layer_window(&self, li: usize) -> Option<usize> {
9891 self.swa
9892 .and_then(|(w, _)| self.layer_is_local(li).then_some(w))
9893 }
9894
9895 fn layer_num_heads(&self, li: usize) -> usize {
9896 self.attention_heads_per_layer
9897 .as_ref()
9898 .and_then(|v| v.get(li).copied())
9899 .unwrap_or(self.num_heads)
9900 }
9901
9902 fn layer_rope_scale(&self, li: usize) -> f32 {
9903 if self.layer_is_local(li) {
9904 self.rope_scale_local
9905 } else {
9906 self.rope_scale
9907 }
9908 }
9909
9910 fn layer_geom(&self, li: usize) -> (usize, usize, usize) {
9913 if !self.layer_is_local(li) {
9914 if let Some((ghd, gkv)) = self.global_attn {
9915 return (gkv, ghd, ghd);
9916 }
9917 }
9918 (
9919 self.layer_num_kv_heads(li),
9920 self.head_dim,
9921 if self.layer_is_local(li) {
9922 self.rotary_dim_local.unwrap_or(self.rotary_dim)
9923 } else {
9924 self.rotary_dim
9925 },
9926 )
9927 }
9928
9929 fn layer_num_kv_heads(&self, li: usize) -> usize {
9932 self.kv_heads_per_layer
9933 .as_ref()
9934 .and_then(|v| v.get(self.phys_layer(li)).copied())
9935 .unwrap_or(self.num_kv_heads)
9936 }
9937
9938 fn layer_v_dim(&self, li: usize) -> usize {
9940 let (_, hd, _) = self.layer_geom(li);
9941 self.v_head_dim.unwrap_or(hd).min(hd)
9942 }
9943
9944 pub fn set_attn_geometry(
9953 &mut self,
9954 kv_heads_per_layer: Option<Vec<usize>>,
9955 v_head_dim: Option<usize>,
9956 ) -> Result<(), String> {
9957 if kv_heads_per_layer.is_some() || v_head_dim.is_some() {
9958 if self.global_attn.is_some() {
9959 return Err(
9960 "per-layer KV heads / v_head_dim cannot combine with Gemma-4 global \
9961 attention geometry"
9962 .into(),
9963 );
9964 }
9965 if self
9966 .weights
9967 .layers
9968 .iter()
9969 .any(|lw| matches!(lw.attn, AttnKind::Mla(_)))
9970 {
9971 return Err("per-layer KV heads / v_head_dim cannot combine with MLA".into());
9972 }
9973 }
9974 if let Some(vd) = v_head_dim {
9975 if vd == 0 || vd > self.head_dim {
9976 return Err(format!(
9977 "v_head_dim {vd} must be in 1..={} (head_dim)",
9978 self.head_dim
9979 ));
9980 }
9981 }
9982 if let Some(v) = &kv_heads_per_layer {
9983 if v.len() != self.physical_layers {
9984 return Err(format!(
9985 "kv_heads_per_layer has {} entries, expected {} layers",
9986 v.len(),
9987 self.physical_layers
9988 ));
9989 }
9990 for (li, &nkv) in v.iter().enumerate() {
9991 let is_attn = matches!(
9992 self.weights.layers.get(li).map(|lw| &lw.attn),
9993 Some(AttnKind::Full { .. }) | None
9994 );
9995 if !is_attn {
9996 continue;
9997 }
9998 let nh = self
9999 .attention_heads_per_layer
10000 .as_ref()
10001 .and_then(|h| h.get(li).copied())
10002 .unwrap_or(self.num_heads);
10003 if nkv == 0 || nh % nkv != 0 {
10004 return Err(format!(
10005 "layer {li}: {nkv} KV heads must be nonzero and divide {nh} Q heads"
10006 ));
10007 }
10008 }
10009 }
10010 self.kv_heads_per_layer = kv_heads_per_layer;
10011 self.v_head_dim = v_head_dim.filter(|&vd| vd != self.head_dim);
10012 if self.kv_heads_per_layer.is_some() {
10013 for li in 0..self.kv_cache.layers.len() {
10014 let full = matches!(
10015 self.weights
10016 .layers
10017 .get(self.phys_layer(li))
10018 .map(|lw| &lw.attn),
10019 Some(AttnKind::Full { .. })
10020 );
10021 let nkv = self.layer_num_kv_heads(li);
10022 let cache = &self.kv_cache.layers[li];
10023 if full && (cache.num_kv_heads != nkv || cache.head_dim != self.head_dim) {
10024 let sinks = cache.sinks.clone();
10025 self.kv_cache.layers[li] =
10026 crate::kv_cache::LayerKvCache::new(nkv, self.head_dim);
10027 self.kv_cache.layers[li].sinks = sinks;
10028 }
10029 }
10030 }
10031 Ok(())
10032 }
10033
10034 pub fn set_layer_sinks(&mut self, phys: usize, sinks: Vec<f32>) -> Result<(), String> {
10038 let Some(lw) = self.weights.layers.get(phys) else {
10039 return Err(format!("sinks for layer {phys}: no such layer"));
10040 };
10041 if !matches!(lw.attn, AttnKind::Full { .. }) {
10042 return Err(format!(
10043 "sinks for layer {phys}: only softmax (Full) attention layers take sinks"
10044 ));
10045 }
10046 let nh = self
10047 .attention_heads_per_layer
10048 .as_ref()
10049 .and_then(|h| h.get(phys).copied())
10050 .unwrap_or(self.num_heads);
10051 if sinks.len() != nh {
10052 return Err(format!(
10053 "sinks for layer {phys}: {} values, expected one per Q head ({nh})",
10054 sinks.len()
10055 ));
10056 }
10057 if let Some(bad) = sinks.iter().find(|v| !v.is_finite()) {
10058 return Err(format!("sinks for layer {phys}: non-finite value {bad}"));
10059 }
10060 for li in 0..self.kv_cache.layers.len() {
10061 if self.phys_layer(li) == phys {
10062 self.kv_cache.layers[li].sinks = Some(sinks.clone());
10063 }
10064 }
10065 Ok(())
10066 }
10067
10068 pub fn graph_attn_decline_reason(&self) -> Option<&'static str> {
10078 if self.kv_heads_per_layer.is_some() {
10079 return Some("per-layer KV head counts");
10080 }
10081 if self.v_head_dim.is_some_and(|vd| vd != self.head_dim) {
10082 return Some("V heads narrower than Q/K heads");
10083 }
10084 if self.kv_cache.layers.iter().any(|l| l.sinks.is_some()) {
10085 return Some("learned attention sinks");
10086 }
10087 if self.swa.is_some() || self.sliding_layers.is_some() {
10088 return Some("sliding-window layers");
10089 }
10090 None
10091 }
10092
10093 #[cfg(target_os = "macos")]
10107 fn metal_graph_swa(&self) -> bool {
10108 (self.swa.is_some() || self.sliding_layers.is_some())
10109 && self.attention_heads_per_layer.is_none()
10110 && !self.attn_v_norm
10111 && self.attn_softcap == 0.0
10112 && self.kv_heads_per_layer.is_none()
10113 && self.v_head_dim.map_or(true, |vd| vd == self.head_dim)
10114 && !self.kv_cache.layers.iter().any(|l| l.sinks.is_some())
10115 && self.global_attn.is_none()
10116 && self.inv_freq_global.is_none()
10117 && (0..self.num_layers).all(|li| self.layer_rope_scale(li) == 1.0)
10118 && !(0..self.num_layers).any(|li| {
10119 self.layer_is_local(li)
10120 && self.inv_freq_local.is_none()
10121 && self.rotary_dim_local.is_some_and(|r| r != self.rotary_dim)
10122 })
10123 }
10124
10125 pub fn wgpu_graph_attn_decline(&self) -> Option<&'static str> {
10132 self.graph_attn_decline_reason()?;
10133 if self.global_attn.is_some() {
10134 return Some("per-layer head width (Gemma-4 global layers) with per-layer geometry");
10135 }
10136 if self.attention_heads_per_layer.is_some() {
10137 return Some("per-layer Q head counts with per-layer geometry");
10138 }
10139 if self.attn_v_norm {
10140 return Some("V norm with per-layer geometry");
10141 }
10142 if (0..self.num_layers).any(|li| self.layer_rope_scale(li) != 1.0) {
10143 return Some("scaled RoPE positions with per-layer geometry");
10144 }
10145 if self.weights.layers.iter().any(|lw| {
10146 lw.attn_out_norm.is_some() || lw.ffn_out_norm.is_some() || lw.layer_scale.is_some()
10147 }) {
10148 return Some("sandwich norms / layer scale with per-layer geometry");
10149 }
10150 if self.weights.layers.iter().any(|lw| {
10151 matches!(
10152 &lw.attn,
10153 AttnKind::Full {
10154 output_gate: true,
10155 ..
10156 }
10157 )
10158 }) && self.v_head_dim.is_some()
10159 {
10160 return Some("gated attention with V narrower than K");
10161 }
10162 if (0..self.num_layers).any(|li| {
10163 self.layer_is_local(li)
10164 && self.inv_freq_local.is_none()
10165 && self.rotary_dim_local.is_some_and(|r| r != self.rotary_dim)
10166 }) {
10167 return Some("local rotary width without a local RoPE table");
10168 }
10169 None
10170 }
10171
10172 fn graph_attn_geom(&self, li: usize) -> Option<crate::gpu::GraphAttnGeom<'_>> {
10177 self.graph_attn_decline_reason()?;
10178 let (nkv, _hd, rd) = self.layer_geom(li);
10179 let invf: &[f32] = if self.layer_is_local(li) {
10180 match &self.inv_freq_local {
10181 Some(f) => f.as_slice(),
10182 None => self.inv_freq.as_slice(),
10183 }
10184 } else {
10185 match &self.inv_freq_global {
10186 Some(f) => f.as_slice(),
10187 None => self.inv_freq.as_slice(),
10188 }
10189 };
10190 Some(crate::gpu::GraphAttnGeom {
10191 nkv,
10192 dv: self.layer_v_dim(li),
10193 rd,
10194 invf,
10195 window: self.layer_window(li),
10196 sink: self.kv_cache.layers[li].sinks.as_deref(),
10197 })
10198 }
10199
10200 #[cfg(feature = "gpu")]
10209 fn pull_lagging_host_kv(&mut self, from: usize, upto: usize, position: usize) {
10210 let kv_id = self.graph_kv_id;
10211 for li in from..upto.min(self.num_layers) {
10212 if !matches!(
10213 self.weights.layers[self.phys_layer(li)].attn,
10214 AttnKind::Full { .. }
10215 ) {
10216 continue;
10217 }
10218 let host = self.kv_cache.layers[li].seq_len;
10219 if host >= position {
10220 continue;
10221 }
10222 let Some(dev) = crate::gpu::graph_kv_stored(kv_id, li) else {
10223 continue;
10224 };
10225 let to = dev.min(position);
10226 if to <= host {
10227 continue;
10228 }
10229 let (nkv, hd) = {
10230 let c = &self.kv_cache.layers[li];
10231 (c.num_kv_heads, c.head_dim)
10232 };
10233 let Some((k, v, first_valid)) =
10234 crate::gpu::graph_kv_pull_host(kv_id, li, host, to, nkv, hd)
10235 else {
10236 continue;
10237 };
10238 let need_from = match self.layer_window(li) {
10241 Some(w) => host.max((position + 1).saturating_sub(w)),
10242 None => host,
10243 };
10244 if first_valid > need_from {
10245 tracing::warn!(
10246 "layer {li}: device KV rows {host}..{to} no longer resident \
10247 (from {first_valid}); host attention will miss them"
10248 );
10249 }
10250 let row = nkv * hd;
10251 let cache = &mut self.kv_cache.layers[li];
10252 for p in 0..to - host {
10253 cache.append(&k[p * row..(p + 1) * row], &v[p * row..(p + 1) * row], &[]);
10254 }
10255 }
10256 }
10257
10258 fn note_graph_decline(&self, site: &'static str, reason: &'static str) {
10261 let mut seen = self.graph_declines.borrow_mut();
10262 if !seen.iter().any(|&(s, r)| s == site && r == reason) {
10263 tracing::warn!("{site} declined: {reason} (CPU attention path)");
10264 seen.push((site, reason));
10265 }
10266 }
10267
10268 pub fn graph_declines(&self) -> Vec<String> {
10271 self.graph_declines
10272 .borrow()
10273 .iter()
10274 .map(|(site, reason)| format!("{site} declined: {reason} (CPU attention path)"))
10275 .collect()
10276 }
10277
10278 fn dump_layer_row(&self, pos: usize, li: usize, row: &[f32]) {
10287 let Some(dir) = &self.layer_dump else {
10288 return;
10289 };
10290 let mut bytes = Vec::with_capacity(row.len() * 4);
10291 for v in row {
10292 bytes.extend_from_slice(&v.to_le_bytes());
10293 }
10294 let path = dir.join(format!("p{pos:06}_l{li:02}.f32"));
10295 if let Err(e) = std::fs::create_dir_all(dir).and_then(|_| std::fs::write(&path, &bytes)) {
10296 use std::sync::atomic::{AtomicBool, Ordering};
10297 static SAID: AtomicBool = AtomicBool::new(false);
10298 if !SAID.swap(true, Ordering::Relaxed) {
10299 tracing::error!("CMF_LAYER_DUMP: cannot write {}: {e}", path.display());
10300 }
10301 }
10302 }
10303
10304 fn mimo_moe_prepare(&mut self) {
10307 if !self.mimo_moe.is_undecided() {
10308 return;
10309 }
10310 let slot = {
10311 let layers: Vec<(usize, &MoeFfn)> = (0..self.num_layers)
10312 .filter_map(
10313 |li| match &self.weights.layers.get(self.phys_layer(li))?.ffn {
10314 FfnKind::Moe(m) => Some((li, m)),
10315 _ => None,
10316 },
10317 )
10318 .collect();
10319 if layers.is_empty()
10322 || self.physical_layers != self.num_layers
10323 || self.gpu_plan.is_some()
10324 {
10325 crate::mimo_moe::Slot::Off
10326 } else {
10327 let graph_prefix =
10331 self.wgpu_graph_attn_decline().is_none() && crate::gpu::wgpu_graph_default();
10332 crate::mimo_moe::Slot::decide(&layers, self.num_layers, graph_prefix)
10333 }
10334 };
10335 self.mimo_moe = slot;
10336 }
10337
10338 #[cfg(test)]
10339 pub(crate) fn test_graph_kv_id(&self) -> u64 {
10340 self.graph_kv_id
10341 }
10342
10343 pub(crate) fn mimo_graph_layer_rows(
10347 &mut self,
10348 li: usize,
10349 h: &mut [f32],
10350 positions: &[usize],
10351 ) -> crate::gpu::BatchGraphOutcome {
10352 use crate::gpu::BatchGraphOutcome as Out;
10353 let b = positions.len();
10354 if !(1..=4).contains(&b)
10355 || h.len() != b * self.hidden_size
10356 || !self.mimo_moe.is_dynamic(li, true)
10357 || !crate::gpu::enabled_here()
10358 || !crate::gpu::wgpu_active()
10359 || self.o1_active()
10360 || self.physical_layers != self.num_layers
10361 || std::env::var("CMF_GPU_WGPU_GRAPH").as_deref() == Ok("0")
10365 || std::env::var("CMF_MIMO_ATTN_GRAPH").as_deref() == Ok("0")
10366 || self.wgpu_graph_attn_decline().is_some()
10367 {
10368 return Out::Declined;
10369 }
10370 let attn_started = std::time::Instant::now();
10371 let outcome = {
10372 let lw = &self.weights.layers[li];
10373 if lw.attn_out_norm.is_some() || lw.ffn_out_norm.is_some() || lw.layer_scale.is_some() {
10374 return Out::Declined;
10375 }
10376 let FfnKind::Moe(m) = &lw.ffn else {
10377 return Out::Declined;
10378 };
10379 let AttnKind::Full {
10380 wq,
10381 wk,
10382 wv,
10383 wo,
10384 q_norm,
10385 k_norm,
10386 output_gate,
10387 softplus_gate,
10388 bias,
10389 } = &lw.attn
10390 else {
10391 return Out::Declined;
10392 };
10393 if *output_gate || softplus_gate.is_some() {
10394 return Out::Declined;
10395 }
10396 let Some(model) = wq.graph_weight().map(|(m, ..)| m.clone()).or_else(|| {
10397 m.experts
10398 .first()?
10399 .gate_proj
10400 .mapped_q4tp()
10401 .map(|(m, _)| m.clone())
10402 }) else {
10403 return Out::Declined;
10404 };
10405 fn gw<'a>(
10406 t: &'a QTensor,
10407 owner: &std::sync::Arc<cortiq_core::CmfModel>,
10408 ) -> Option<crate::gpu::GraphW<'a>> {
10409 if let Some((m, idx, kind, rs)) = t.graph_weight() {
10410 if m.uid() != owner.uid() || t.has_prism_contract() {
10411 return None;
10412 }
10413 return Some(crate::gpu::GraphW {
10414 idx,
10415 kind,
10416 row_scale: rs,
10417 data: &[],
10418 prism: crate::gpu::GraphPrismOp::None,
10419 affine: false,
10420 });
10421 }
10422 t.as_f32().map(|data| crate::gpu::GraphW {
10423 idx: 0,
10424 kind: 4,
10425 row_scale: &[],
10426 data,
10427 prism: crate::gpu::GraphPrismOp::None,
10428 affine: false,
10429 })
10430 }
10431 let (Some(q), Some(k), Some(v), Some(o)) = (
10432 gw(wq, &model),
10433 gw(wk, &model),
10434 gw(wv, &model),
10435 gw(wo, &model),
10436 ) else {
10437 return Out::Declined;
10438 };
10439 let layer = crate::gpu::GraphLayer {
10440 input_norm: &lw.input_norm,
10441 post_norm: &lw.post_norm,
10442 ffn: crate::gpu::GraphFfn::AttentionOnly,
10443 attn: crate::gpu::GraphAttn::Full {
10444 wq: q,
10445 wk: k,
10446 wv: v,
10447 wo: o,
10448 q_norm: q_norm.as_deref(),
10449 k_norm: k_norm.as_deref(),
10450 late_qk_norm: self.qk_norm_after_rope,
10451 bias: bias
10452 .as_ref()
10453 .map(|(q, k, v)| (q.as_slice(), k.as_slice(), v.as_slice())),
10454 output_gate: false,
10455 cpu_k: self.kv_cache.layers[li].k_heads(),
10456 cpu_v: self.kv_cache.layers[li].v_heads(),
10457 geom: self.graph_attn_geom(li),
10458 head_gate: None,
10459 },
10460 };
10461 let (nkv, hd, rd) = self.layer_geom(li);
10462 crate::gpu::forward_batch_graph_at(
10463 &model,
10464 self.graph_kv_id,
10465 li,
10466 &[layer],
10467 &self.inv_freq,
10468 h,
10469 self.layer_num_heads(li),
10470 nkv,
10471 hd,
10472 rd,
10473 self.hidden_size,
10474 1,
10475 positions,
10476 self.kv_cache.max_seq_len,
10477 self.norm_style == cortiq_core::NormStyle::Gemma,
10478 self.rms_eps as f32,
10479 self.attn_scale,
10480 b,
10481 &[],
10482 self.o1_epoch,
10483 None,
10484 None,
10485 )
10486 };
10487 match outcome {
10488 Out::Completed => {}
10489 Out::Declined => return Out::Declined,
10490 Out::Failed => {
10491 self.graph_failed
10492 .store(true, std::sync::atomic::Ordering::Relaxed);
10493 return Out::Failed;
10494 }
10495 }
10496 let attn_ns = attn_started.elapsed().as_nanos() as u64;
10497 let hs = self.hidden_size;
10498 let lw = &self.weights.layers[li];
10499 let FfnKind::Moe(m) = &lw.ffn else {
10500 unreachable!()
10501 };
10502 let mut post = vec![0.0; h.len()];
10503 for (x, y) in h.chunks_exact(hs).zip(post.chunks_exact_mut(hs)) {
10504 inference::rms_norm_into(x, &lw.post_norm, self.rms_eps, self.norm_style, y);
10505 }
10506 let mut ffn = if b == 1 {
10507 moe_ffn_banked(&mut self.mimo_moe, li, m, &post, self.pool.as_deref())
10508 } else {
10509 moe_ffn_banked_rows(
10510 &mut self.mimo_moe,
10511 li,
10512 m,
10513 &post,
10514 b,
10515 hs,
10516 self.pool.as_deref(),
10517 )
10518 };
10519 for (x, &f) in h.iter_mut().zip(&ffn) {
10520 *x += f;
10521 }
10522 attention::recycle_buf(&mut ffn);
10523 if self.layer_dump.is_some() {
10524 for (&pos, row) in positions.iter().zip(h.chunks_exact(hs)) {
10525 self.dump_layer_row(pos, li, row);
10526 }
10527 }
10528 crate::mimo_moe::note_attention_graph(b, attn_ns);
10529 Out::Completed
10530 }
10531
10532 fn layer_attn_plain(&self, li: usize) -> bool {
10533 self.kv_heads_per_layer.is_none()
10534 && self.v_head_dim.is_none()
10535 && self.global_attn.is_none()
10536 && self.layer_window(li).is_none()
10537 && self.kv_cache.layers[li].sinks.is_none()
10538 }
10539
10540 fn forward_layers(
10542 &mut self,
10543 hidden: &[f32],
10544 position: usize,
10545 task_mask: Option<&TaskMask>,
10546 ) -> Vec<f32> {
10547 let out = self.forward_layers_upto(hidden, position, task_mask, None);
10548 self.o1_progress();
10549 out
10550 }
10551
10552 pub fn embed_id(&self, id: u32) -> Vec<f32> {
10560 self.embed_single(id)
10561 }
10562
10563 pub fn split_supported(&self) -> Result<(), String> {
10567 if self.dsv4.is_some() {
10568 return Err(
10569 "network split: DeepSeek-V4 runs its own fused stack (not splittable yet)".into(),
10570 );
10571 }
10572 if self.dsv41.is_some() {
10573 return Err(
10574 "network split: DeepSeek-V4.1 owns the shared CED/CSA2 state (not splittable)"
10575 .into(),
10576 );
10577 }
10578 if self.qwen4_exp.is_some() {
10579 return Err(
10580 "network split: Qwen3.8-Flash-Next hyper/QSA stack is not splittable yet".into(),
10581 );
10582 }
10583 if self.g3n.is_some() {
10584 return Err(
10585 "network split: Gemma-3n runs its own AltUp stack (not splittable yet)".into(),
10586 );
10587 }
10588 Ok(())
10589 }
10590
10591 pub fn forward_span(
10596 &mut self,
10597 hidden: &[f32],
10598 position: usize,
10599 from: usize,
10600 upto: usize,
10601 task_mask: Option<&TaskMask>,
10602 ) -> Result<Vec<f32>, String> {
10603 self.split_supported()?;
10604 if from > upto || upto >= self.num_layers {
10605 return Err(format!(
10606 "forward_span: layer range {from}..={upto} outside 0..{}",
10607 self.num_layers
10608 ));
10609 }
10610 if hidden.len() != self.hidden_size {
10611 return Err(format!(
10612 "forward_span: hidden len {} ≠ hidden_size {}",
10613 hidden.len(),
10614 self.hidden_size
10615 ));
10616 }
10617 let out = self.forward_layers_span(hidden, position, task_mask, from, Some(upto));
10618 self.o1_progress();
10619 if self
10620 .graph_failed
10621 .swap(false, std::sync::atomic::Ordering::Relaxed)
10622 {
10623 self.cancel
10624 .store(false, std::sync::atomic::Ordering::Relaxed);
10625 self.clear_sequence_state();
10626 return Err("forward_span: deferred O(1) transition failed".into());
10627 }
10628 Ok(out)
10629 }
10630
10631 pub fn logits_from_hidden(&mut self, hidden: &[f32]) -> Vec<f32> {
10634 let normed = inference::rms_norm(
10635 hidden,
10636 &self.weights.final_norm,
10637 self.rms_eps,
10638 self.norm_style,
10639 );
10640 self.lm_head_forward(&normed)
10641 }
10642
10643 pub fn sample_next(&mut self, logits: &[f32], past_tokens: &[u32]) -> u32 {
10645 sampler::sample_with_scratch(
10646 logits,
10647 &self.sampler_config,
10648 past_tokens,
10649 &mut self.rng,
10650 &mut self.sampler_scratch,
10651 )
10652 }
10653
10654 pub fn reset_session(&mut self) {
10656 self.clear_sequence_state();
10657 }
10658
10659 pub fn prefill_span_ids(
10665 &mut self,
10666 ids: &[u32],
10667 start_pos: usize,
10668 upto: usize,
10669 task_mask: Option<&TaskMask>,
10670 ) -> Result<Vec<f32>, String> {
10671 self.split_supported()?;
10672 if upto >= self.num_layers {
10673 return Err(format!(
10674 "prefill_span_ids: upto {upto} outside 0..{}",
10675 self.num_layers
10676 ));
10677 }
10678 if self.can_prefill_batched() && !self.graph_prefill_preferred() {
10682 let out =
10683 self.prefill_batch_span(PrefillIn::Ids(ids), start_pos, task_mask, 0, upto + 1);
10684 self.check_o1_progress_failure("prefill_span_ids")?;
10685 Ok(out)
10686 } else {
10687 let hs = self.hidden_size;
10688 let mut out = Vec::with_capacity(ids.len() * hs);
10689 for (i, &id) in ids.iter().enumerate() {
10690 let emb = self.embed_id(id);
10691 out.extend_from_slice(&self.forward_span(
10692 &emb,
10693 start_pos + i,
10694 0,
10695 upto,
10696 task_mask,
10697 )?);
10698 }
10699 Ok(out)
10700 }
10701 }
10702
10703 pub fn prefill_span_hidden(
10706 &mut self,
10707 hidden: &[f32],
10708 start_pos: usize,
10709 from: usize,
10710 upto: usize,
10711 task_mask: Option<&TaskMask>,
10712 ) -> Result<Vec<f32>, String> {
10713 self.split_supported()?;
10714 let hs = self.hidden_size;
10715 if hidden.is_empty() || hidden.len() % hs != 0 {
10716 return Err(format!(
10717 "prefill_span_hidden: {} floats is not a multiple of hidden {hs}",
10718 hidden.len()
10719 ));
10720 }
10721 if from > upto || upto >= self.num_layers {
10722 return Err(format!(
10723 "prefill_span_hidden: layer range {from}..={upto} outside 0..{}",
10724 self.num_layers
10725 ));
10726 }
10727 if self.can_prefill_batched() && !self.graph_prefill_preferred() {
10728 let out = self.prefill_batch_span(
10729 PrefillIn::Hidden(hidden),
10730 start_pos,
10731 task_mask,
10732 from,
10733 upto + 1,
10734 );
10735 self.check_o1_progress_failure("prefill_span_hidden")?;
10736 Ok(out)
10737 } else {
10738 let b = hidden.len() / hs;
10739 let mut out = Vec::with_capacity(hidden.len());
10740 for i in 0..b {
10741 let h = self.forward_span(
10742 &hidden[i * hs..(i + 1) * hs],
10743 start_pos + i,
10744 from,
10745 upto,
10746 task_mask,
10747 )?;
10748 out.extend_from_slice(&h);
10749 }
10750 Ok(out)
10751 }
10752 }
10753
10754 fn try_token_graph_wgpu(
10758 &self,
10759 hidden: &[f32],
10760 position: usize,
10761 logits_out: &mut Vec<f32>,
10762 layers_run: &mut usize,
10763 ) -> Option<Result<Vec<f32>, ()>> {
10764 self.try_token_graph_wgpu_steps(
10765 hidden,
10766 position,
10767 logits_out,
10768 1,
10769 None,
10770 Some(layers_run),
10771 0,
10772 self.num_layers,
10773 )
10774 }
10775
10776 fn try_token_graph_wgpu_span(
10780 &self,
10781 hidden: &[f32],
10782 position: usize,
10783 logits_out: &mut Vec<f32>,
10784 from: usize,
10785 upto_excl: usize,
10786 layers_run: &mut usize,
10787 ) -> Option<Result<Vec<f32>, ()>> {
10788 self.try_token_graph_wgpu_steps(
10789 hidden,
10790 position,
10791 logits_out,
10792 1,
10793 None,
10794 Some(layers_run),
10795 from,
10796 upto_excl,
10797 )
10798 }
10799
10800 fn try_multi_burst(&self, t_next: u32, position: usize, k: usize) -> Option<Vec<u32>> {
10804 if self.o1_active() || self.attn_softcap > 0.0 {
10805 return None;
10806 }
10807 if let Some(reason) = self.wgpu_graph_attn_decline() {
10810 self.note_graph_decline("wgpu multi-burst", reason);
10811 return None;
10812 }
10813 let graph_on = crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Decode);
10814 if !graph_on || self.graph_refused() {
10815 return None;
10822 }
10823 let emb = self.embed_single(t_next);
10824 let mut lg = Vec::new();
10825 let mut ids = Vec::new();
10826 match self.try_token_graph_wgpu_steps(
10827 &emb,
10828 position,
10829 &mut lg,
10830 k,
10831 Some(&mut ids),
10832 None,
10833 0,
10834 self.num_layers,
10835 ) {
10836 Some(Ok(_)) => {}
10837 Some(Err(())) => {
10838 self.graph_failed
10843 .store(true, std::sync::atomic::Ordering::Relaxed);
10844 return None;
10845 }
10846 None => return None,
10847 }
10848 (ids.len() == k).then_some(ids)
10849 }
10850
10851 fn try_token_graph_wgpu_steps(
10855 &self,
10856 hidden: &[f32],
10857 position: usize,
10858 logits_out: &mut Vec<f32>,
10859 steps: usize,
10860 ids_out: Option<&mut Vec<u32>>,
10861 layers_run: Option<&mut usize>,
10862 from: usize,
10863 upto_excl: usize,
10864 ) -> Option<Result<Vec<f32>, ()>> {
10865 let upto_excl = match self.mimo_moe.graph_prefix_end() {
10868 Some(end) if end < upto_excl => {
10869 if steps != 1 || layers_run.is_none() || from >= end {
10870 return None;
10871 }
10872 end
10873 }
10874 _ => upto_excl,
10875 };
10876 let o1_gpu = std::env::var("CMF_O1_GPU").as_deref() == Ok("1");
10879 if (self.o1_active() && !o1_gpu) || self.attn_softcap > 0.0 {
10880 return None;
10884 }
10885 if let Some(reason) = self.wgpu_graph_attn_decline() {
10892 self.note_graph_decline("wgpu token graph", reason);
10893 return None;
10894 }
10895 let o1_views: Vec<Option<Vec<crate::nystrom::O1DeviceView<'_>>>> = (from..upto_excl)
10900 .map(|li| {
10901 if !o1_gpu {
10902 return None;
10903 }
10904 self.kv_cache.layers[self.phys_layer(li)].o1_views()
10905 })
10906 .collect();
10907 if self.o1_active() && o1_gpu {
10908 let want: usize = (from..upto_excl)
10911 .filter(|li| self.kv_cache.layers[self.phys_layer(*li)].o1.is_some())
10912 .count();
10913 let have = o1_views.iter().filter(|v| v.is_some()).count();
10914 if want == 0 || have != want {
10915 use std::sync::atomic::{AtomicUsize, Ordering};
10925 static LAST: AtomicUsize = AtomicUsize::new(usize::MAX);
10926 let code = have * 1000 + want;
10927 if LAST.swap(code, Ordering::Relaxed) != code {
10928 tracing::warn!(
10929 "o1 graph: {have} of {want} layers sealed — per-op until all seal"
10930 );
10931 }
10932 return None;
10933 }
10934 }
10935 let nh = self.num_heads;
10936 let (nkv, hd, rd) = self.layer_geom(0);
10937 let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
10938 let mut layers = Vec::with_capacity(upto_excl - from);
10939 let mut model = None;
10940 let dbg = std::env::var("CMF_GRAPH_DEBUG").is_ok();
10941 fn gw(t: &QTensor) -> Option<crate::gpu::GraphW<'_>> {
10942 if let Some((m, i, kind, rs)) = t
10943 .graph_weight()
10944 .or_else(|| t.graph_weight_descriptor())
10945 {
10946 let name = &m.tensors[i].name;
10947 let prism = if crate::prism::is_inverse_embedding(m, name) {
10948 crate::gpu::GraphPrismOp::InverseEmbedding
10949 } else if crate::prism::is_forward_weight(m, name) {
10950 crate::gpu::GraphPrismOp::Forward
10951 } else {
10952 crate::gpu::GraphPrismOp::None
10953 };
10954 return Some(crate::gpu::GraphW {
10955 idx: i,
10956 kind,
10957 row_scale: rs,
10958 data: &[],
10959 prism,
10960 affine: crate::prism::is_affine_target(m, name),
10961 });
10962 }
10963 match t.as_f32() {
10965 Some(d) => Some(crate::gpu::GraphW {
10966 idx: 0,
10967 kind: 4,
10968 row_scale: &[],
10969 data: d,
10970 prism: crate::gpu::GraphPrismOp::None,
10971 affine: false,
10972 }),
10973 None => {
10974 if std::env::var_os("CMF_BATCH_DEBUG").is_some() {
10975 eprintln!("batch graph: weight has no graph/f32 representation");
10976 }
10977 None
10978 }
10979 }
10980 }
10981 for li in from..upto_excl {
10982 let lw = &self.weights.layers[self.phys_layer(li)];
10983 if dbg {
10984 let ak = match &lw.attn {
10985 AttnKind::Mla(_) => "Mla".into(),
10986 AttnKind::Full {
10987 output_gate, bias, ..
10988 } => format!("Full gate={output_gate} bias={}", bias.is_some()),
10989 AttnKind::LinearGdn(_) => "LinearGdn".into(),
10990 AttnKind::Kda(_) => "Kda".into(),
10991 AttnKind::Linear(_) => "Linear".into(),
10992 AttnKind::ShortConv(_) => "ShortConv".into(),
10993 AttnKind::Bounded(_) => "Bounded".into(),
10994 };
10995 let fk = match &lw.ffn {
10996 FfnKind::Dense(_) => "Dense",
10997 FfnKind::Moe(_) => "Moe",
10998 FfnKind::DenseMoe(_) => "DenseMoe",
10999 };
11000 eprintln!("graph L{li}: attn={ak} ffn={fk}");
11001 }
11002 let gffn = match &lw.ffn {
11003 FfnKind::DenseMoe(_) => return None, FfnKind::Dense(d) if !d.segs.is_empty() => return None,
11007 FfnKind::Dense(d) => {
11008 let Some(act) = d.act.graph_act() else {
11012 self.note_graph_decline(
11013 "wgpu token graph",
11014 "dense FFN activation without a graph kernel",
11015 );
11016 return None;
11017 };
11018 crate::gpu::GraphFfn::Dense {
11019 gate: gw(&d.gate_proj)?,
11020 up: gw(&d.up_proj)?,
11021 down: gw(&d.down_proj)?,
11022 act,
11023 }
11024 }
11025 FfnKind::Moe(m) => {
11026 if m.route_tau.is_some() || m.mask.is_some() {
11034 return None;
11035 }
11036 let shared = m.shared.as_ref();
11037 let has_shared = shared.is_some();
11038 let shared_gated = matches!(shared, Some((_, Some(_))));
11039 let sgate = match shared {
11040 Some((_, Some(sg))) => gw(sg)?,
11041 _ => gw(&m.router)?,
11045 };
11046 let router = gw(&m.router)?;
11047 if router.prism != crate::gpu::GraphPrismOp::None
11053 || sgate.prism != crate::gpu::GraphPrismOp::None
11054 || router.affine
11055 || sgate.affine
11056 {
11057 tracing::warn!(
11058 "resident MoE declined: Prism/affine router or shared gate transform is not implemented"
11059 );
11060 return None;
11061 }
11062 let inter = m.experts.first()?.gate_proj.rows();
11063 let mut experts = Vec::with_capacity(m.experts.len() + 1);
11064 let mut q4tp: Option<bool> = None;
11067 let mut gu_q2: Option<bool> = None;
11070 for e in m.experts.iter().chain(shared.map(|(se, _)| se)) {
11071 if !matches!(e.act, Act::Silu)
11072 || e.gate_proj.rows() != inter
11073 || e.up_proj.rows() != inter
11074 {
11075 return None;
11076 }
11077 for expert_weight in [&e.gate_proj, &e.up_proj, &e.down_proj] {
11082 let Some((em, ei, _, _)) = expert_weight
11083 .graph_weight()
11084 .or_else(|| expert_weight.graph_weight_descriptor())
11085 else {
11086 return None;
11087 };
11088 let name = &em.tensors[ei].name;
11089 if crate::prism::is_forward_weight(em, name)
11090 || crate::prism::is_inverse_embedding(em, name)
11091 || crate::prism::is_affine_target(em, name)
11092 {
11093 tracing::warn!(
11094 "resident MoE declined: expert Prism/affine transform is not implemented"
11095 );
11096 return None;
11097 }
11098 }
11099 let (mm, gi, ui, di, is_p, is_q2) = match e.gate_proj.mapped_q4t() {
11100 Some((mm, gi)) => (
11101 mm,
11102 gi,
11103 e.up_proj.mapped_q4t()?.1,
11104 e.down_proj.mapped_q4t()?.1,
11105 false,
11106 false,
11107 ),
11108 None => match e.gate_proj.mapped_q2tp() {
11109 Some((mm, gi)) => (
11110 mm,
11111 gi,
11112 e.up_proj.mapped_q2tp()?.1,
11113 e.down_proj.mapped_q4tp()?.1,
11114 true,
11115 true,
11116 ),
11117 None => {
11118 let (mm, gi) = e.gate_proj.mapped_q4tp()?;
11119 (
11120 mm,
11121 gi,
11122 e.up_proj.mapped_q4tp()?.1,
11123 e.down_proj.mapped_q4tp()?.1,
11124 true,
11125 false,
11126 )
11127 }
11128 },
11129 };
11130 if *q4tp.get_or_insert(is_p) != is_p || *gu_q2.get_or_insert(is_q2) != is_q2
11131 {
11132 tracing::warn!(
11138 "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."
11139 );
11140 return None;
11141 }
11142 model.get_or_insert_with(|| mm.clone());
11143 experts.push((gi, ui, di));
11144 }
11145 crate::gpu::GraphFfn::Moe {
11146 router,
11147 shared_gate: sgate,
11148 experts,
11149 n_exp: m.experts.len(),
11150 top_k: std::env::var("CMF_TOPK_PROBE")
11156 .ok()
11157 .and_then(|v| v.parse::<usize>().ok())
11158 .filter(|k| *k > 0 && *k <= m.top_k)
11159 .unwrap_or(m.top_k),
11160 inter,
11161 norm_topk: m.norm_topk_prob,
11162 q4tp: q4tp?,
11163 gu_q2: gu_q2.unwrap_or(false),
11164 sigmoid: m.router_sigmoid,
11165 bias: m.expert_bias.as_deref(),
11166 has_shared,
11167 shared_gated,
11168 route_scale: m.routed_scaling,
11169 }
11170 }
11171 };
11172 let attn = match &lw.attn {
11173 AttnKind::Full {
11174 wq,
11175 wk,
11176 wv,
11177 wo,
11178 q_norm,
11179 k_norm,
11180 output_gate,
11181 softplus_gate,
11182 bias,
11183 } => {
11184 if self.attention_heads_per_layer.is_some() {
11185 return None;
11186 }
11187 let head_gate = match softplus_gate {
11191 None => None,
11192 Some((g, true)) if self.proj_gate_sigmoid && !*output_gate => {
11193 Some(gw(g)?)
11194 }
11195 Some(_) => {
11196 self.note_graph_decline(
11197 "wgpu token graph",
11198 "projected softplus / per-element output gate",
11199 );
11200 return None;
11201 }
11202 };
11203 let (m, _, _, _) = wq
11204 .graph_weight()
11205 .or_else(|| wq.graph_weight_descriptor())?;
11206 model = Some(m.clone());
11207 crate::gpu::GraphAttn::Full {
11208 wq: gw(wq)?,
11209 wk: gw(wk)?,
11210 wv: gw(wv)?,
11211 wo: gw(wo)?,
11212 q_norm: q_norm.as_deref(),
11213 k_norm: k_norm.as_deref(),
11214 late_qk_norm: self.qk_norm_after_rope,
11215 bias: bias
11216 .as_ref()
11217 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
11218 output_gate: *output_gate,
11219 cpu_k: self.kv_cache.layers[li].k_heads(),
11220 cpu_v: self.kv_cache.layers[li].v_heads(),
11221 geom: self.graph_attn_geom(li),
11222 head_gate,
11223 }
11224 }
11225 AttnKind::LinearGdn(w) => {
11226 let cfg = self.gdn_cfg?;
11227 let (m, _, _, _) = w
11228 .in_proj_qkv
11229 .graph_weight()
11230 .or_else(|| w.in_proj_qkv.graph_weight_descriptor())?;
11231 model = Some(m.clone());
11232 crate::gpu::GraphAttn::Gdn {
11233 qkv: gw(&w.in_proj_qkv)?,
11234 z: gw(&w.in_proj_z)?,
11235 a: gw(&w.in_proj_a)?,
11236 b: gw(&w.in_proj_b)?,
11237 out: gw(&w.out_proj)?,
11238 conv1d: &w.conv1d,
11239 a_log: &w.a_log,
11240 dt_bias: &w.dt_bias,
11241 norm: &w.norm,
11242 nv: cfg.num_v_heads,
11243 nk: cfg.num_k_heads,
11244 dk: cfg.key_head_dim,
11245 dv: cfg.value_head_dim,
11246 kk: cfg.conv_kernel,
11247 cpu_state: &self.kv_cache.layers[self.phys_layer(li)].linear_state,
11248 }
11249 }
11250 AttnKind::ShortConv(w) => {
11251 let cfg = self.short_conv_cfg?;
11252 let (m, _, _, _) = w
11253 .in_proj
11254 .graph_weight()
11255 .or_else(|| w.in_proj.graph_weight_descriptor())?;
11256 model = Some(m.clone());
11257 crate::gpu::GraphAttn::ShortConv {
11258 inp: gw(&w.in_proj)?,
11259 out: gw(&w.out_proj)?,
11260 taps: &w.conv,
11261 kernel: cfg.kernel,
11262 cpu_state: &self.kv_cache.layers[self.phys_layer(li)].linear_state,
11263 }
11264 }
11265 _ => return None,
11266 };
11267 layers.push(crate::gpu::GraphLayer {
11268 input_norm: &lw.input_norm,
11269 attn,
11270 post_norm: &lw.post_norm,
11271 ffn: gffn,
11272 });
11273 }
11274 let model = model?;
11275 let lm_gw = if upto_excl == self.num_layers
11281 && self.graph_want_logits
11282 && std::env::var("CMF_GPU_LMHEAD")
11283 .map(|v| v != "0")
11284 .unwrap_or(true)
11285 {
11286 self.weights
11287 .lm_head
11288 .graph_weight()
11289 .or_else(|| self.weights.lm_head.graph_weight_descriptor())
11290 .map(|(m, i, kind, rs)| {
11291 let name = &m.tensors[i].name;
11292 let prism = if crate::prism::is_inverse_embedding(m, name) {
11293 crate::gpu::GraphPrismOp::InverseEmbedding
11294 } else if crate::prism::is_forward_weight(m, name) {
11295 crate::gpu::GraphPrismOp::Forward
11296 } else {
11297 crate::gpu::GraphPrismOp::None
11298 };
11299 (
11300 crate::gpu::GraphW {
11301 idx: i,
11302 kind,
11303 row_scale: rs,
11304 data: &[],
11305 prism,
11306 affine: crate::prism::is_affine_target(m, name),
11307 },
11308 self.weights.lm_head.rows(),
11309 )
11310 })
11311 } else {
11312 None
11313 };
11314 let lm = lm_gw.as_ref().map(|(gw, rows)| (gw, *rows));
11315 let emb_gw = if steps > 1 {
11317 self.weights
11318 .embed_tokens
11319 .graph_weight()
11320 .or_else(|| self.weights.embed_tokens.graph_weight_descriptor())
11321 .map(|(m, i, kind, rs)| {
11322 let name = &m.tensors[i].name;
11323 let prism = if crate::prism::is_inverse_embedding(m, name) {
11324 crate::gpu::GraphPrismOp::InverseEmbedding
11325 } else if crate::prism::is_forward_weight(m, name) {
11326 crate::gpu::GraphPrismOp::Forward
11327 } else {
11328 crate::gpu::GraphPrismOp::None
11329 };
11330 (
11331 crate::gpu::GraphW {
11332 idx: i,
11333 kind,
11334 row_scale: rs,
11335 data: &[],
11336 prism,
11337 affine: crate::prism::is_affine_target(m, name),
11338 },
11339 self.weights.embed_tokens.rows(),
11340 self.embed_multiplier,
11341 )
11342 })
11343 } else {
11344 None
11345 };
11346
11347 let loop_norm_at: Vec<usize> = if self.loop_final_norm {
11353 (from..upto_excl.min(self.num_layers - 1))
11354 .filter(|&li| (li + 1) % self.physical_layers == 0)
11355 .map(|li| li - from)
11356 .collect()
11357 } else {
11358 Vec::new()
11359 };
11360 let mut h = hidden.to_vec();
11361 let dump_hidden = std::env::var_os("CMF_LOGIT_DUMP").is_some();
11367 let outcome = crate::gpu::forward_token_graph(
11368 &model,
11369 self.graph_kv_id,
11370 &layers,
11371 &o1_views,
11372 self.o1_epoch,
11373 &self.inv_freq,
11374 &mut h,
11375 nh,
11376 nkv,
11377 hd,
11378 self.attn_scale,
11379 rd,
11380 self.hidden_size,
11381 self.intermediate_size,
11382 position,
11383 self.kv_cache.max_seq_len,
11384 gemma,
11385 self.rms_eps as f32,
11386 lm,
11387 &self.weights.final_norm,
11388 logits_out,
11389 &loop_norm_at,
11390 steps,
11391 emb_gw.as_ref().map(|(gw, rows, m)| (gw, *rows, *m)),
11392 ids_out,
11393 layers_run,
11394 from,
11395 dump_hidden,
11396 );
11397 match outcome {
11398 crate::gpu::TokenGraphOutcome::Completed => Some(Ok(h)),
11399 crate::gpu::TokenGraphOutcome::Failed => Some(Err(())),
11400 crate::gpu::TokenGraphOutcome::Declined => None,
11401 }
11402 }
11403
11404 #[cfg(target_os = "macos")]
11413 #[allow(clippy::type_complexity)]
11414 fn metal_rows_plan(
11415 &self,
11416 ) -> Option<(
11417 Vec<MetalRowsItem<'_>>,
11418 std::sync::Arc<cortiq_core::CmfModel>,
11419 Option<crate::gpu_metal::GdnGpuCfg>,
11420 )> {
11421 use crate::gpu_metal::{AttnGpuLayer, GdnGpuCfg, GdnGpuLayer, MetalFfn};
11422 let graph_force = crate::gpu::q1_force() || crate::gpu::q2tp_gpu_opt_in();
11423 if !graph_force
11424 || !crate::gpu::enabled_here()
11425 || std::env::var("CMF_GPU_BLOCK")
11426 .map(|v| v == "0")
11427 .unwrap_or(false)
11428 || self.attn_softcap > 0.0
11429 || self.o1_active()
11430 || self.swa.is_some()
11431 || self.global_attn.is_some()
11432 || self.attention_heads_per_layer.is_some()
11433 || self.graph_attn_decline_reason().is_some()
11435 || self.attn_v_norm
11436 || self.loop_final_norm
11437 {
11438 return None;
11439 }
11440 let attend_contract = self.head_dim % 4 == 0
11441 && self.head_dim <= 256
11442 && self.rotary_dim >= 2
11443 && self.rotary_dim <= self.head_dim
11444 && (self.rotary_dim / 2) % 32 == 0
11445 && self.num_kv_heads > 0
11446 && self.num_heads % self.num_kv_heads == 0;
11447 if !attend_contract {
11448 return None;
11449 }
11450 let mut plan: Vec<MetalRowsItem> = Vec::new();
11451 let mut model_ref: Option<std::sync::Arc<cortiq_core::CmfModel>> = None;
11452 for li in 0..self.num_layers {
11453 let lw = &self.weights.layers[self.phys_layer(li)];
11454 if lw.attn_out_norm.is_some() || lw.ffn_out_norm.is_some() || lw.layer_scale.is_some() {
11455 return None;
11456 }
11457 let ffn = match &lw.ffn {
11458 FfnKind::Dense(d) if d.act == Act::Silu && d.segs.is_empty() => {
11459 let (Some(g), Some(u), Some(dn)) = (
11460 d.gate_proj.metal_graph_parts(),
11461 d.up_proj.metal_graph_parts(),
11462 d.down_proj.metal_graph_parts(),
11463 ) else {
11464 return None;
11465 };
11466 MetalFfn::Dense {
11467 gate: g,
11468 up: u,
11469 down: dn,
11470 gelu: false, }
11472 }
11473 _ => return None,
11474 };
11475 match &lw.attn {
11476 AttnKind::LinearGdn(w) if self.gdn_cfg.is_some() => {
11477 let (Some(qkv), Some(z), Some(a), Some(bb), Some(out)) = (
11478 w.in_proj_qkv.metal_graph_parts(),
11479 w.in_proj_z.metal_graph_parts(),
11480 w.in_proj_a.f32_parts(),
11481 w.in_proj_b.f32_parts(),
11482 w.out_proj.metal_graph_parts(),
11483 ) else {
11484 return None;
11485 };
11486 if let QTensor::Mapped { model, .. } = &w.in_proj_qkv {
11487 model_ref.get_or_insert_with(|| model.clone());
11488 }
11489 let gl = GdnGpuLayer {
11490 attn_norm: &lw.input_norm,
11491 post_norm: &lw.post_norm,
11492 qkv,
11493 z,
11494 a,
11495 b: bb,
11496 out,
11497 ffn,
11498 conv1d: &w.conv1d,
11499 a_log: &w.a_log,
11500 dt_bias: &w.dt_bias,
11501 gnorm: &w.norm,
11502 };
11503 match plan.last_mut() {
11504 Some(MetalRowsItem::Gdn { run, .. }) => run.push(gl),
11505 _ => plan.push(MetalRowsItem::Gdn {
11506 run: vec![gl],
11507 first: li,
11508 }),
11509 }
11510 }
11511 AttnKind::Full {
11512 wq,
11513 wk,
11514 wv,
11515 wo,
11516 q_norm,
11517 k_norm,
11518 output_gate,
11519 softplus_gate: None,
11520 bias: None,
11521 } => {
11522 let (Some(pq), Some(pk), Some(pv), Some(po)) =
11523 (
11524 wq.metal_graph_parts(),
11525 wk.metal_graph_parts(),
11526 wv.metal_graph_parts(),
11527 wo.metal_graph_parts(),
11528 )
11529 else {
11530 return None;
11531 };
11532 if let QTensor::Mapped { model, .. } = wq {
11533 model_ref.get_or_insert_with(|| model.clone());
11534 }
11535 let cache = &self.kv_cache.layers[li];
11536 if cache.mode != crate::kv_cache::KvMode::F32 || cache.o1.is_some() {
11537 return None;
11538 }
11539 plan.push(MetalRowsItem::Attn {
11540 l: AttnGpuLayer {
11541 attn_norm: &lw.input_norm,
11542 post_norm: &lw.post_norm,
11543 wq: pq,
11544 wk: pk,
11545 wv: pv,
11546 wo: po,
11547 ffn,
11548 },
11549 li,
11550 q_norm: q_norm.as_deref(),
11551 k_norm: k_norm.as_deref(),
11552 output_gate: *output_gate,
11553 });
11554 }
11555 _ => return None,
11556 }
11557 }
11558 let model = model_ref?;
11559 let gcfg = self.gdn_cfg.map(|cfg| GdnGpuCfg {
11560 nv: cfg.num_v_heads,
11561 nk: cfg.num_k_heads,
11562 dk: cfg.key_head_dim,
11563 dv: cfg.value_head_dim,
11564 kk: cfg.conv_kernel,
11565 hidden: self.hidden_size,
11566 inter: self.intermediate_size,
11567 c_dim: cfg.conv_dim(),
11568 eps: cfg.rms_eps as f32,
11569 gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
11570 });
11571 Some((plan, model, gcfg))
11572 }
11573
11574 #[cfg(target_os = "macos")]
11576 #[allow(clippy::too_many_arguments)]
11577 fn metal_attn_params<'a>(
11578 li: usize,
11579 cache: &'a crate::kv_cache::LayerKvCache,
11580 q_norm: Option<&'a [f32]>,
11581 k_norm: Option<&'a [f32]>,
11582 output_gate: bool,
11583 inv_freq: &'a [f32],
11584 geom: (usize, usize, usize, usize),
11585 pos0: usize,
11586 kv_id: u64,
11587 scale: f32,
11588 eps: f32,
11589 gemma: bool,
11590 late_qk_norm: bool,
11591 ) -> (crate::gpu_metal::AttnDeviceParams<'a>, usize) {
11592 let (nh, nkv, hd, rd) = geom;
11593 let cpu_k: Vec<&[f32]> = (0..nkv).map(|g| cache.head_keys(g)).collect();
11594 let cpu_v: Vec<&[f32]> = (0..nkv).map(|g| cache.head_values(g)).collect();
11595 let cpu_stored = cpu_k[0].len() / hd;
11596 (
11597 crate::gpu_metal::AttnDeviceParams {
11598 kv_id,
11599 layer: li,
11600 nh,
11601 nkv,
11602 hd,
11603 rd,
11604 position: pos0,
11605 scale,
11606 eps,
11607 gemma,
11608 late_qk_norm,
11609 output_gate,
11610 q_norm,
11611 k_norm,
11612 inv_freq,
11613 cpu_k,
11614 cpu_v,
11615 cpu_stored,
11616 o1: None,
11617 window: None,
11618 head_gate: None,
11619 },
11620 cpu_stored,
11621 )
11622 }
11623
11624 #[cfg(target_os = "macos")]
11629 #[allow(clippy::type_complexity)]
11630 fn metal_rows_run(
11631 &mut self,
11632 hiddens: &mut [f32],
11633 pos0: usize,
11634 b: usize,
11635 prefill: bool,
11636 spec: Option<((usize, usize, usize), &[f32], &mut Vec<f32>)>,
11637 mut argmax_out: Option<(usize, &mut Vec<u32>)>,
11641 ) -> MetalRowsRun {
11642 use crate::gpu_metal::{GraphDims, VerifyGraph};
11643 if !crate::gpu_metal::wait_replay() {
11649 tracing::error!("Metal rows graph: the pending async replay failed");
11650 return MetalRowsRun::Failed;
11651 }
11652 spec_stamp("v.wait");
11653 let want = self.gdn_cfg.map(|c| c.state_len()).unwrap_or(0);
11659 if want > 0 {
11660 let phys = self.physical_layers.max(1);
11661 for (li, l) in self.kv_cache.layers.iter_mut().enumerate() {
11662 let is_gdn = self
11663 .weights
11664 .layers
11665 .get(li % phys)
11666 .is_some_and(|lw| matches!(lw.attn, AttnKind::LinearGdn(_)));
11667 if is_gdn && l.linear_state.len() != want {
11668 l.linear_state = vec![0f32; want];
11669 }
11670 }
11671 }
11672 let Some((plan, model, gcfg)) = self.metal_rows_plan() else {
11673 return MetalRowsRun::Declined;
11674 };
11675 spec_stamp("v.plan");
11676 let dims = GraphDims {
11677 hidden: self.hidden_size,
11678 eps: self.rms_eps as f32,
11679 gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
11680 };
11681 let Some(mut graph) = (if prefill {
11682 VerifyGraph::new_prefill(&model, dims, hiddens, b)
11683 } else {
11684 VerifyGraph::new(&model, dims, hiddens, b)
11685 }) else {
11686 return MetalRowsRun::Declined;
11687 };
11688 let geom = (
11689 self.num_heads,
11690 self.num_kv_heads,
11691 self.head_dim,
11692 self.rotary_dim,
11693 );
11694 let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
11695 let eps = self.rms_eps as f32;
11696 let kv_id = self.graph_kv_id;
11697 let inv_freq = self.inv_freq.clone();
11698 for item in &plan {
11699 let ok = match item {
11700 MetalRowsItem::Gdn { run, .. } => gcfg
11701 .as_ref()
11702 .map(|gc| run.iter().all(|l| graph.gdn_ok(l, gc)))
11703 .unwrap_or(false),
11704 MetalRowsItem::Attn {
11705 l,
11706 li,
11707 q_norm,
11708 k_norm,
11709 output_gate,
11710 } => {
11711 let (p, _) = Self::metal_attn_params(
11712 *li,
11713 &self.kv_cache.layers[*li],
11714 *q_norm,
11715 *k_norm,
11716 *output_gate,
11717 &inv_freq,
11718 geom,
11719 pos0,
11720 kv_id,
11721 self.attn_scale,
11722 eps,
11723 gemma,
11724 self.qk_norm_after_rope,
11725 );
11726 graph.attn_ok(l, &p)
11727 }
11728 };
11729 if !ok {
11730 use std::sync::atomic::{AtomicBool, Ordering};
11731 static SAID: AtomicBool = AtomicBool::new(false);
11732 if !SAID.swap(true, Ordering::Relaxed) {
11733 tracing::warn!("metal rows graph: a layer failed preflight — declining");
11734 }
11735 return MetalRowsRun::Declined;
11736 }
11737 }
11738 let lm = match &spec {
11739 Some((lm, _, _)) => {
11740 if !graph.lm_head_ok(*lm) {
11741 return MetalRowsRun::Declined;
11742 }
11743 Some(*lm)
11744 }
11745 None => None,
11746 };
11747 let mut gdn_layers = Vec::new();
11748 let mut attn_layers = Vec::new();
11749 for item in &plan {
11750 match item {
11751 MetalRowsItem::Gdn { run, first } => {
11752 let ro: Vec<&[f32]> = self.kv_cache.layers[*first..*first + run.len()]
11753 .iter()
11754 .map(|l| l.linear_state.as_slice())
11755 .collect();
11756 if !graph.encode_gdn_run_b(run, &ro, gcfg.as_ref().unwrap()) {
11757 return MetalRowsRun::Declined;
11758 }
11759 gdn_layers.extend(*first..*first + run.len());
11760 }
11761 MetalRowsItem::Attn {
11762 l,
11763 li,
11764 q_norm,
11765 k_norm,
11766 output_gate,
11767 } => {
11768 let (p, cpu_stored) = Self::metal_attn_params(
11769 *li,
11770 &self.kv_cache.layers[*li],
11771 *q_norm,
11772 *k_norm,
11773 *output_gate,
11774 &inv_freq,
11775 geom,
11776 pos0,
11777 kv_id,
11778 self.attn_scale,
11779 eps,
11780 gemma,
11781 self.qk_norm_after_rope,
11782 );
11783 if !graph.encode_attn_b(l, &p) {
11784 return MetalRowsRun::Declined;
11785 }
11786 attn_layers.push((*li, cpu_stored));
11787 }
11788 }
11789 }
11790 if let (Some(lm), Some((_, final_norm, _))) = (lm, spec.as_ref()) {
11791 if !graph.encode_lm_head_b(final_norm, lm) {
11792 return MetalRowsRun::Declined;
11793 }
11794 if let Some((n, _)) = argmax_out.as_ref() {
11799 if !graph.encode_argmax_b(*n) {
11800 argmax_out = None;
11801 }
11802 }
11803 }
11804 spec_stamp("v.enc");
11805 if !graph.sync() {
11806 return MetalRowsRun::Failed;
11807 }
11808 spec_stamp("v.gpu");
11809 match (spec, argmax_out) {
11810 (Some(_), Some((_, ids))) => {
11811 ids.resize(b, 0);
11812 if !graph.read_argmax(ids) {
11813 return MetalRowsRun::Failed;
11814 }
11815 spec_stamp("v.am");
11816 }
11817 (Some((lm, _, logits)), None) => {
11818 logits.resize(b * lm.1, 0.0);
11819 if !graph.read_logits(logits) {
11820 return MetalRowsRun::Failed;
11821 }
11822 spec_stamp("v.lg");
11823 }
11824 (None, _) => {}
11825 }
11826 if !graph.read_hidden(hiddens) {
11827 return MetalRowsRun::Failed;
11828 }
11829 spec_stamp("v.hid");
11830 MetalRowsRun::Completed(MetalVerifyPending {
11831 graph,
11832 gdn_layers,
11833 attn_layers,
11834 })
11835 }
11836
11837 #[cfg(target_os = "macos")]
11843 fn try_batch_graph_metal(
11844 &mut self,
11845 hiddens: &mut [f32],
11846 positions: &[usize],
11847 b: usize,
11848 spec: Option<((usize, usize, usize), &[f32], &mut Vec<f32>)>,
11849 argmax_out: Option<(usize, &mut Vec<u32>)>,
11850 ) -> crate::gpu::BatchGraphOutcome {
11851 let _t0 = std::time::Instant::now();
11852 if positions.len() != b
11853 || positions.windows(2).any(|w| w[1] != w[0] + 1)
11854 || hiddens.len() != b * self.hidden_size
11855 {
11856 return crate::gpu::BatchGraphOutcome::Declined;
11857 }
11858 let pending = match self.metal_rows_run(hiddens, positions[0], b, false, spec, argmax_out) {
11859 MetalRowsRun::Declined => return crate::gpu::BatchGraphOutcome::Declined,
11860 MetalRowsRun::Failed => return crate::gpu::BatchGraphOutcome::Failed,
11861 MetalRowsRun::Completed(pending) => pending,
11862 };
11863 if std::env::var("CMF_GRAPH_SPEC_TIME").is_ok() {
11864 eprintln!(
11865 "metal-verify: {:.1} ms | b={b}",
11866 _t0.elapsed().as_secs_f64() * 1e3
11867 );
11868 }
11869 self.metal_verify = Some(pending);
11870 crate::gpu::BatchGraphOutcome::Completed
11871 }
11872
11873 #[cfg(target_os = "macos")]
11878 fn prefill_rows_metal(
11879 &mut self,
11880 ids: &[u32],
11881 start_pos: usize,
11882 spec: Option<((usize, usize, usize), &[f32], &mut Vec<f32>)>,
11883 ) -> MetalPrefillOutcome {
11884 let b = ids.len();
11885 if b == 0 || b > 512 {
11886 return MetalPrefillOutcome::Declined;
11887 }
11888 METAL_PREFILL_CHUNKS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
11889 let with_head = spec.is_some();
11890 let hs = self.hidden_size;
11891 let mut hiddens = vec![0f32; b * hs];
11892 for (j, &id) in ids.iter().enumerate() {
11893 let e = self.embed_single(id);
11894 hiddens[j * hs..(j + 1) * hs].copy_from_slice(&e);
11895 }
11896 let mut pending = match self.metal_rows_run(&mut hiddens, start_pos, b, true, spec, None) {
11897 MetalRowsRun::Declined => return MetalPrefillOutcome::Declined,
11898 MetalRowsRun::Failed => {
11899 METAL_PREFILL_ERRORS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
11900 return MetalPrefillOutcome::Failed;
11901 }
11902 MetalRowsRun::Completed(pending) => pending,
11903 };
11904 let idxs = pending.gdn_layers.clone();
11906 let mut outs: Vec<&mut [f32]> = self
11907 .kv_cache
11908 .layers
11909 .iter_mut()
11910 .enumerate()
11911 .filter(|(i, _)| idxs.binary_search(i).is_ok())
11912 .map(|(_, l)| l.linear_state.as_mut_slice())
11913 .collect();
11914 if !pending.graph.finish_states(&mut outs) {
11915 METAL_PREFILL_ERRORS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
11916 return MetalPrefillOutcome::Failed;
11917 }
11918 let (nkv, hd) = (self.num_kv_heads, self.head_dim);
11919 let mut rows = Vec::with_capacity(pending.attn_layers.len());
11923 for (li, cpu_stored) in &pending.attn_layers {
11924 let mut kbuf = vec![0f32; b * nkv * hd];
11925 let mut vbuf = vec![0f32; b * nkv * hd];
11926 if !crate::gpu_metal::kv_mirror_read_rows(
11927 self.graph_kv_id,
11928 *li,
11929 nkv,
11930 hd,
11931 *cpu_stored,
11932 b,
11933 &mut kbuf,
11934 &mut vbuf,
11935 ) {
11936 METAL_PREFILL_ERRORS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
11937 return MetalPrefillOutcome::Failed;
11938 }
11939 rows.push((*li, *cpu_stored, kbuf, vbuf));
11940 }
11941 for (li, cpu_stored, kbuf, vbuf) in rows {
11942 let cache = &mut self.kv_cache.layers[li];
11943 for r in 0..b {
11944 cache.append(
11945 &kbuf[r * nkv * hd..(r + 1) * nkv * hd],
11946 &vbuf[r * nkv * hd..(r + 1) * nkv * hd],
11947 &[],
11948 );
11949 }
11950 crate::gpu_metal::kv_mirror_set_stored(self.graph_kv_id, li, cpu_stored + b);
11951 }
11952 METAL_PREFILL_ROWS.fetch_add(b as u64, std::sync::atomic::Ordering::Relaxed);
11953 if with_head {
11954 METAL_PREFILL_HEAD_ROWS.fetch_add(b as u64, std::sync::atomic::Ordering::Relaxed);
11955 }
11956 MetalPrefillOutcome::Completed(hiddens)
11957 }
11958
11959 #[cfg(target_os = "macos")]
11960 fn prefill_batch_metal(&mut self, ids: &[u32], start_pos: usize) -> MetalPrefillOutcome {
11961 self.prefill_rows_metal(ids, start_pos, None)
11962 }
11963
11964 #[cfg(target_os = "macos")]
11969 fn nll_batch_metal(&mut self, ids: &[u32], start: usize) -> MetalBatchNllOutcome {
11970 if ids.len() < 2 || self.o1_active() || self.head_clusters.is_some() {
11971 return MetalBatchNllOutcome::Declined;
11972 }
11973 let Some(lm) = self.weights.lm_head.metal_graph_parts() else {
11974 return MetalBatchNllOutcome::Declined;
11975 };
11976 let chunk = std::env::var("CMF_METAL_PREFILL_CHUNK")
11977 .ok()
11978 .and_then(|v| v.parse::<usize>().ok())
11979 .filter(|&v| (1..=512).contains(&v))
11980 .unwrap_or(32);
11981 let final_norm = self.weights.final_norm.clone();
11982 let mut nll = 0.0f64;
11983 let mut count = 0usize;
11984 let mut pos = 0usize;
11985 let mut completed = 0usize;
11986 while pos < ids.len() {
11987 let end = (pos + chunk).min(ids.len());
11988 let mut logits = Vec::new();
11989 let outcome = self.prefill_rows_metal(
11990 &ids[pos..end],
11991 pos,
11992 Some((lm, &final_norm, &mut logits)),
11993 );
11994 match outcome {
11995 MetalPrefillOutcome::Declined => {
11996 return if completed == 0 {
11997 MetalBatchNllOutcome::Declined
11998 } else {
11999 MetalBatchNllOutcome::Failed(format!(
12000 "ordinary Metal NLL batch declined after {completed} chunks"
12001 ))
12002 };
12003 }
12004 MetalPrefillOutcome::Failed => {
12005 return MetalBatchNllOutcome::Failed(
12006 "ordinary Metal NLL batch failed after admission".to_string(),
12007 );
12008 }
12009 MetalPrefillOutcome::Completed(_) => {}
12010 }
12011 completed += 1;
12012 let vocab = self.vocab_size.min(lm.1);
12013 if logits.len() != (end - pos) * lm.1 || vocab == 0 {
12014 return MetalBatchNllOutcome::Failed(
12015 "ordinary Metal NLL head returned an invalid shape".to_string(),
12016 );
12017 }
12018 for row in 0..(end - pos) {
12019 let absolute = pos + row;
12020 if absolute < start || absolute + 1 >= ids.len() {
12021 continue;
12022 }
12023 let lg = &mut logits[row * lm.1..row * lm.1 + vocab];
12024 if let Some(mu) = self.logit_multiplier {
12025 for v in lg.iter_mut() {
12026 *v *= mu;
12027 }
12028 }
12029 if let Some(c) = self.final_softcap {
12030 for v in lg.iter_mut() {
12031 *v = c * (*v / c).tanh();
12032 }
12033 }
12034 let target = ids[absolute + 1] as usize;
12035 if target >= vocab {
12036 return MetalBatchNllOutcome::Failed(format!(
12037 "target token {target} exceeds Metal head rows {vocab}"
12038 ));
12039 }
12040 let max = lg.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
12041 let lse: f64 = lg
12042 .iter()
12043 .map(|&v| ((v - max) as f64).exp())
12044 .sum::<f64>()
12045 .ln()
12046 + max as f64;
12047 nll += lse - lg[target] as f64;
12048 count += 1;
12049 }
12050 pos = end;
12051 }
12052 MetalBatchNllOutcome::Completed(nll, count)
12053 }
12054
12055 #[cfg(target_os = "macos")]
12059 fn metal_verify_commit(&mut self, a: usize) -> bool {
12060 let Some(mut pending) = self.metal_verify.take() else {
12061 return false;
12062 };
12063 let n = a + 1;
12064 let idxs = pending.gdn_layers.clone();
12066 let mut outs: Vec<&mut [f32]> = self
12067 .kv_cache
12068 .layers
12069 .iter_mut()
12070 .enumerate()
12071 .filter(|(i, _)| idxs.binary_search(i).is_ok())
12072 .map(|(_, l)| l.linear_state.as_mut_slice())
12073 .collect();
12074 if !pending.graph.commit(n, &mut outs) {
12075 return false;
12076 }
12077 spec_stamp("c.replay");
12078 let (nkv, hd) = (self.num_kv_heads, self.head_dim);
12079 let mut rows = Vec::with_capacity(pending.attn_layers.len());
12083 for (li, cpu_stored) in &pending.attn_layers {
12084 let mut kbuf = vec![0f32; n * nkv * hd];
12085 let mut vbuf = vec![0f32; n * nkv * hd];
12086 if !crate::gpu_metal::kv_mirror_read_rows(
12087 self.graph_kv_id,
12088 *li,
12089 nkv,
12090 hd,
12091 *cpu_stored,
12092 n,
12093 &mut kbuf,
12094 &mut vbuf,
12095 ) {
12096 return false;
12097 }
12098 rows.push((*li, *cpu_stored, kbuf, vbuf));
12099 }
12100 for (li, cpu_stored, kbuf, vbuf) in rows {
12101 let cache = &mut self.kv_cache.layers[li];
12102 for r in 0..n {
12103 cache.append(
12104 &kbuf[r * nkv * hd..(r + 1) * nkv * hd],
12105 &vbuf[r * nkv * hd..(r + 1) * nkv * hd],
12106 &[],
12107 );
12108 }
12109 crate::gpu_metal::kv_mirror_set_stored(self.graph_kv_id, li, cpu_stored + n);
12110 }
12111 spec_stamp("c.kv");
12112 true
12113 }
12114
12115 #[cfg(target_os = "macos")]
12122 fn mtp_warm_batch_submit(
12123 &mut self,
12124 m: &mut MtpModule,
12125 pairs: &[(&[f32], u32)],
12126 first_pos: usize,
12127 ) -> Option<MetalWarmPending> {
12128 use crate::gpu_metal::{AttnDeviceParams, AttnGpuLayer, GraphDims, MetalFfn, VerifyGraph};
12129 let b = pairs.len();
12130 if b == 0 || b > 512 || m.kv.mode != crate::kv_cache::KvMode::F32 || m.kv.o1.is_some() {
12131 return None;
12132 }
12133 let AttnKind::Full {
12134 wq,
12135 wk,
12136 wv,
12137 wo,
12138 q_norm,
12139 k_norm,
12140 output_gate,
12141 softplus_gate: None,
12142 bias: None,
12143 } = &m.layer.attn
12144 else {
12145 return None;
12146 };
12147 let FfnKind::Dense(d) = &m.layer.ffn else {
12148 return None;
12149 };
12150 if !d.segs.is_empty() {
12151 return None;
12152 }
12153 let (Some(pq), Some(pk), Some(pv), Some(po)) =
12154 (wq.q1_parts(), wk.q1_parts(), wv.q1_parts(), wo.q1_parts())
12155 else {
12156 return None;
12157 };
12158 let (Some(g), Some(u), Some(dn)) = (
12159 d.gate_proj.q1_parts(),
12160 d.up_proj.q1_parts(),
12161 d.down_proj.q1_parts(),
12162 ) else {
12163 return None;
12164 };
12165 let Some(eh) = m.eh_proj.q1_parts() else {
12166 return None;
12167 };
12168 let QTensor::Mapped { model, .. } = wq else {
12169 return None;
12170 };
12171 let model = model.clone();
12172 let hs = self.hidden_size;
12173 let mut cat = vec![0f32; b * 2 * hs];
12175 for (j, (h, tok)) in pairs.iter().enumerate() {
12176 let e = self.embed_single(*tok);
12177 let (ce, ch) = cat[j * 2 * hs..(j + 1) * 2 * hs].split_at_mut(hs);
12178 inference::rms_norm_into(&e, &m.enorm, self.rms_eps, self.norm_style, ce);
12179 inference::rms_norm_into(h, &m.hnorm, self.rms_eps, self.norm_style, ch);
12180 }
12181 let dims = GraphDims {
12182 hidden: hs,
12183 eps: self.rms_eps as f32,
12184 gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
12185 };
12186 spec_stamp("w.cat");
12187 let Some(mut graph) = VerifyGraph::new_via_proj(&model, dims, eh, &cat, b) else {
12188 return None;
12189 };
12190 spec_stamp("w.new");
12191 let l = AttnGpuLayer {
12192 attn_norm: &m.layer.input_norm,
12193 post_norm: &m.layer.post_norm,
12194 wq: pq,
12195 wk: pk,
12196 wv: pv,
12197 wo: po,
12198 ffn: MetalFfn::Dense {
12199 gate: g,
12200 up: u,
12201 down: dn,
12202 gelu: d.act == Act::Gelu,
12203 },
12204 };
12205 let (nh, nkv, hd, rd) = (
12206 self.num_heads,
12207 self.num_kv_heads,
12208 self.head_dim,
12209 self.rotary_dim,
12210 );
12211 let inv_freq = self.inv_freq.clone();
12212 let cpu_stored;
12213 {
12214 let cache = &m.kv;
12215 let cpu_k: Vec<&[f32]> = (0..nkv).map(|g| cache.head_keys(g)).collect();
12216 let cpu_v: Vec<&[f32]> = (0..nkv).map(|g| cache.head_values(g)).collect();
12217 cpu_stored = cpu_k[0].len() / hd;
12218 if cpu_stored > first_pos {
12223 spec_stamp("w.decl");
12224 return None;
12225 }
12226 let p = AttnDeviceParams {
12227 kv_id: self.mtp_kv_id(),
12228 layer: Self::MTP_LAYER_BASE,
12229 nh,
12230 nkv,
12231 hd,
12232 rd,
12233 position: first_pos,
12234 scale: self.attn_scale,
12235 eps: self.rms_eps as f32,
12236 gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
12237 late_qk_norm: self.qk_norm_after_rope,
12238 output_gate: *output_gate,
12239 q_norm: q_norm.as_deref(),
12240 k_norm: k_norm.as_deref(),
12241 inv_freq: &inv_freq,
12242 cpu_k,
12243 cpu_v,
12244 cpu_stored,
12245 o1: None,
12246 window: None,
12247 head_gate: None,
12248 };
12249 if !graph.attn_ok(&l, &p) || !graph.encode_attn_b(&l, &p) {
12250 return None;
12251 }
12252 }
12253 spec_stamp("w.enc");
12254 if !graph.submit() {
12255 return None;
12256 }
12257 spec_stamp("w.sub");
12258 Some(MetalWarmPending {
12259 graph,
12260 cpu_stored,
12261 b,
12262 })
12263 }
12264
12265 #[cfg(target_os = "macos")]
12268 fn mtp_warm_batch_metal(
12269 &mut self,
12270 m: &mut MtpModule,
12271 pairs: &[(&[f32], u32)],
12272 first_pos: usize,
12273 ) -> bool {
12274 match self.mtp_warm_batch_submit(m, pairs, first_pos) {
12275 Some(p) => self.mtp_warm_batch_finish(m, p),
12276 None => false,
12277 }
12278 }
12279
12280 #[cfg(target_os = "macos")]
12285 fn mtp_warm_batch_finish(&mut self, m: &mut MtpModule, pending: MetalWarmPending) -> bool {
12286 let MetalWarmPending {
12287 mut graph,
12288 cpu_stored,
12289 b,
12290 } = pending;
12291 let (nkv, hd) = (self.num_kv_heads, self.head_dim);
12292 if !graph.sync() {
12293 return false;
12294 }
12295 spec_stamp("w.gpu");
12296 let mut kbuf = vec![0f32; b * nkv * hd];
12297 let mut vbuf = vec![0f32; b * nkv * hd];
12298 if !crate::gpu_metal::kv_mirror_read_rows(
12299 self.mtp_kv_id(),
12300 Self::MTP_LAYER_BASE,
12301 nkv,
12302 hd,
12303 cpu_stored,
12304 b,
12305 &mut kbuf,
12306 &mut vbuf,
12307 ) {
12308 return false;
12309 }
12310 for r in 0..b {
12311 m.kv.append(
12312 &kbuf[r * nkv * hd..(r + 1) * nkv * hd],
12313 &vbuf[r * nkv * hd..(r + 1) * nkv * hd],
12314 &[],
12315 );
12316 }
12317 crate::gpu_metal::kv_mirror_set_stored(
12318 self.mtp_kv_id(),
12319 Self::MTP_LAYER_BASE,
12320 cpu_stored + b,
12321 );
12322 spec_stamp("w.kv");
12323 true
12324 }
12325
12326 pub(crate) fn note_draft_id(&mut self, id: u32) {
12333 let cut = Self::draft_vocab_rows(usize::MAX).max(131_072);
12334 if (id as usize) >= cut {
12335 self.draft_full_streak = 16;
12336 } else {
12337 self.draft_full_streak = self.draft_full_streak.saturating_sub(1);
12338 }
12339 }
12340
12341 fn draft_head_rows(&self, head_rows: usize) -> usize {
12344 if self.draft_full_streak > 0 {
12345 head_rows
12346 } else {
12347 Self::draft_vocab_rows(head_rows)
12348 }
12349 }
12350
12351 fn draft_vocab_rows(head_rows: usize) -> usize {
12354 static N: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
12355 let n = *N.get_or_init(|| {
12356 std::env::var("CMF_DRAFT_VOCAB")
12357 .ok()
12358 .and_then(|v| v.parse().ok())
12359 .unwrap_or(65536)
12360 });
12361 if n == 0 { head_rows } else { n.min(head_rows) }
12362 }
12363
12364 #[cfg(target_os = "macos")]
12369 fn mtp_step_metal(
12370 &mut self,
12371 m: &mut MtpModule,
12372 hidden: &[f32],
12373 next_token: u32,
12374 position: usize,
12375 want_logits: bool,
12376 ) -> Option<(Vec<f32>, Vec<f32>)> {
12377 use crate::gpu_metal::{AttnDeviceParams, AttnGpuLayer, GraphDims, MetalFfn, TokenGraph};
12378 if std::env::var("CMF_MTP_GRAPH").as_deref() == Ok("0")
12379 || !crate::gpu::q1_force()
12380 || !crate::gpu::enabled_here()
12381 || self.attn_softcap > 0.0
12382 || self.attention_heads_per_layer.is_some()
12383 || m.kv.mode != crate::kv_cache::KvMode::F32
12384 || m.kv.o1.is_some()
12385 {
12386 return None;
12387 }
12388 let AttnKind::Full {
12389 wq,
12390 wk,
12391 wv,
12392 wo,
12393 q_norm,
12394 k_norm,
12395 output_gate,
12396 softplus_gate: None,
12397 bias: None,
12398 } = &m.layer.attn
12399 else {
12400 return None;
12401 };
12402 let FfnKind::Dense(d) = &m.layer.ffn else {
12403 return None;
12404 };
12405 if d.act != Act::Silu || !d.segs.is_empty() {
12406 return None;
12407 }
12408 let (pq, pk, pv, po) = (
12409 wq.q1_parts()?,
12410 wk.q1_parts()?,
12411 wv.q1_parts()?,
12412 wo.q1_parts()?,
12413 );
12414 let (g, u, dn) = (
12415 d.gate_proj.q1_parts()?,
12416 d.up_proj.q1_parts()?,
12417 d.down_proj.q1_parts()?,
12418 );
12419 let QTensor::Mapped { model, .. } = wq else {
12420 return None;
12421 };
12422 let model = model.clone();
12423 let lm = if want_logits {
12424 Some(self.weights.lm_head.q1_parts()?)
12425 } else {
12426 None
12427 };
12428 let dims = GraphDims {
12429 hidden: self.hidden_size,
12430 eps: self.rms_eps as f32,
12431 gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
12432 };
12433 let hs = self.hidden_size;
12436 let mut x = vec![0f32; hs];
12437 let mut graph = TokenGraph::new(&model, dims, &x)?;
12438 let mut folded = false;
12439 if let Some(eh) = m.eh_proj.q1_parts() {
12440 let e = self.embed_single(next_token);
12441 let mut cat = vec![0.0f32; 2 * hs];
12442 let (cat_e, cat_h) = cat.split_at_mut(hs);
12443 inference::rms_norm_into(&e, &m.enorm, self.rms_eps, self.norm_style, cat_e);
12444 inference::rms_norm_into(hidden, &m.hnorm, self.rms_eps, self.norm_style, cat_h);
12445 folded = graph.encode_input_proj(eh, &cat);
12446 }
12447 if !folded {
12448 x = self.mtp_block_input(m, hidden, next_token);
12449 graph = TokenGraph::new(&model, dims, &x)?;
12450 }
12451 spec_stamp("d.in");
12452 let l = AttnGpuLayer {
12453 attn_norm: &m.layer.input_norm,
12454 post_norm: &m.layer.post_norm,
12455 wq: pq,
12456 wk: pk,
12457 wv: pv,
12458 wo: po,
12459 ffn: MetalFfn::Dense {
12460 gate: g,
12461 up: u,
12462 down: dn,
12463 gelu: d.act == Act::Gelu,
12464 },
12465 };
12466 let (nh, nkv, hd, rd) = (
12467 self.num_heads,
12468 self.num_kv_heads,
12469 self.head_dim,
12470 self.rotary_dim,
12471 );
12472 let inv_freq = self.inv_freq.clone();
12473 {
12474 let cache = &m.kv;
12475 let cpu_k: Vec<&[f32]> = (0..nkv).map(|g| cache.head_keys(g)).collect();
12476 let cpu_v: Vec<&[f32]> = (0..nkv).map(|g| cache.head_values(g)).collect();
12477 let cpu_stored = cpu_k[0].len() / hd;
12478 let p = AttnDeviceParams {
12479 kv_id: self.mtp_kv_id(),
12480 layer: Self::MTP_LAYER_BASE,
12481 nh,
12482 nkv,
12483 hd,
12484 rd,
12485 position,
12486 scale: self.attn_scale,
12487 eps: self.rms_eps as f32,
12488 gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
12489 late_qk_norm: self.qk_norm_after_rope,
12490 output_gate: *output_gate,
12491 q_norm: q_norm.as_deref(),
12492 k_norm: k_norm.as_deref(),
12493 inv_freq: &inv_freq,
12494 cpu_k,
12495 cpu_v,
12496 cpu_stored,
12497 o1: None,
12498 window: None,
12499 head_gate: None,
12500 };
12501 if !graph.attn_device_ok(&l, &p) || !graph.encode_attn_device(&l, &p) {
12502 return None;
12503 }
12504 }
12505 let draft_rows = if let Some(lm) = lm {
12511 self.draft_head_rows(lm.1)
12512 } else {
12513 0
12514 };
12515 if let Some(lm) = lm {
12516 if !graph.lm_head_ok(lm) {
12517 return None;
12518 }
12519 if draft_rows < lm.1 {
12520 if !graph.encode_lm_head_part(&m.final_norm, lm, draft_rows) {
12521 return None;
12522 }
12523 } else {
12524 graph.encode_lm_head(&m.final_norm, lm);
12525 }
12526 }
12527 spec_stamp("d.enc");
12528 if graph.sync_checked().is_err() {
12529 return None;
12530 }
12531 spec_stamp("d.gpu");
12532 let mut logits = Vec::new();
12533 if let Some(lm) = lm {
12534 let n_read = draft_rows.min(lm.1).min(self.vocab_size);
12535 logits = attention::take_buf(n_read);
12536 graph.read_logits(&mut logits);
12537 logits.resize(self.vocab_size, f32::NEG_INFINITY);
12539 }
12540 graph.finish(&mut x);
12541 let mut krow = attention::take_buf(nkv * hd);
12542 let mut vrow = attention::take_buf(nkv * hd);
12543 if crate::gpu_metal::kv_mirror_read_last(
12544 self.mtp_kv_id(),
12545 Self::MTP_LAYER_BASE,
12546 nkv,
12547 hd,
12548 &mut krow,
12549 &mut vrow,
12550 ) {
12551 m.kv.append(&krow, &vrow, &[]);
12552 }
12553 attention::recycle_buf(&mut krow);
12554 attention::recycle_buf(&mut vrow);
12555 spec_stamp("d.rd");
12556 Some((logits, x))
12557 }
12558
12559 fn mtp_chain_on() -> bool {
12572 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
12573 *ON.get_or_init(|| std::env::var("CMF_MTP_CHAIN").as_deref() != Ok("0"))
12574 }
12575
12576 #[cfg(target_os = "macos")]
12588 fn mtp_draft_chain_metal(
12589 &mut self,
12590 m: &mut MtpModule,
12591 hidden: &[f32],
12592 t_next: u32,
12593 position: usize,
12594 k: usize,
12595 ) -> Result<Vec<u32>, bool> {
12596 use crate::gpu_metal::{AttnDeviceParams, AttnGpuLayer, GraphDims, MetalFfn, TokenGraph};
12597 if k == 0
12598 || k > 64
12599 || !Self::mtp_chain_on()
12600 || std::env::var("CMF_MTP_GRAPH").as_deref() == Ok("0")
12601 || !crate::gpu::q1_force()
12602 || !crate::gpu::enabled_here()
12603 || self.attn_softcap > 0.0
12604 || self.attention_heads_per_layer.is_some()
12605 || m.kv.mode != crate::kv_cache::KvMode::F32
12606 || m.kv.o1.is_some()
12607 || self.dsv4.is_some()
12609 || self.dsv41.is_some()
12610 || self.qwen4_exp.is_some()
12611 || self.g3n.is_some()
12612 {
12613 return Err(false);
12614 }
12615 let AttnKind::Full {
12616 wq,
12617 wk,
12618 wv,
12619 wo,
12620 q_norm,
12621 k_norm,
12622 output_gate,
12623 softplus_gate: None,
12624 bias: None,
12625 } = &m.layer.attn
12626 else {
12627 return Err(false);
12628 };
12629 let FfnKind::Dense(d) = &m.layer.ffn else {
12630 return Err(false);
12631 };
12632 if d.act != Act::Silu || !d.segs.is_empty() {
12633 return Err(false);
12634 }
12635 let (Some(pq), Some(pk), Some(pv), Some(po)) =
12636 (wq.q1_parts(), wk.q1_parts(), wv.q1_parts(), wo.q1_parts())
12637 else {
12638 return Err(false);
12639 };
12640 let (Some(g), Some(u), Some(dn)) = (
12641 d.gate_proj.q1_parts(),
12642 d.up_proj.q1_parts(),
12643 d.down_proj.q1_parts(),
12644 ) else {
12645 return Err(false);
12646 };
12647 let (Some(eh), Some(lm)) = (m.eh_proj.q1_parts(), self.weights.lm_head.q1_parts()) else {
12648 return Err(false);
12649 };
12650 let QTensor::Mapped { model, .. } = wq else {
12651 return Err(false);
12652 };
12653 let model = model.clone();
12654 let QTensor::Mapped {
12657 model: em,
12658 idx: eidx,
12659 dtype: cortiq_core::TensorDtype::Q4TiledP,
12660 ..
12661 } = &self.weights.embed_tokens
12662 else {
12663 return Err(false);
12664 };
12665 if !std::sync::Arc::ptr_eq(em, &model)
12666 || crate::prism::is_inverse_embedding(&model, &model.tensors[*eidx].name)
12667 {
12668 return Err(false);
12669 }
12670 let embed = (
12671 *eidx,
12672 self.weights.embed_tokens.rows(),
12673 self.weights.embed_tokens.cols(),
12674 );
12675 if embed.2 != self.hidden_size || hidden.len() != self.hidden_size {
12676 return Err(false);
12677 }
12678 let dims = GraphDims {
12679 hidden: self.hidden_size,
12680 eps: self.rms_eps as f32,
12681 gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
12682 };
12683 let Some(mut graph) = TokenGraph::new(&model, dims, hidden) else {
12684 return Err(false);
12685 };
12686 if !graph.chain_embed_ok(embed) || !graph.lm_head_ok(lm) {
12687 return Err(false);
12688 }
12689 let l = AttnGpuLayer {
12690 attn_norm: &m.layer.input_norm,
12691 post_norm: &m.layer.post_norm,
12692 wq: pq,
12693 wk: pk,
12694 wv: pv,
12695 wo: po,
12696 ffn: MetalFfn::Dense {
12697 gate: g,
12698 up: u,
12699 down: dn,
12700 gelu: d.act == Act::Gelu,
12701 },
12702 };
12703 let (nh, nkv, hd, rd) = (
12704 self.num_heads,
12705 self.num_kv_heads,
12706 self.head_dim,
12707 self.rotary_dim,
12708 );
12709 let inv_freq = self.inv_freq.clone();
12710 let draft_rows = self.draft_head_rows(lm.1);
12711 let n_arg = draft_rows.min(lm.1).min(self.vocab_size);
12712 if n_arg == 0 {
12713 return Err(false);
12714 }
12715 let split = std::env::var("CMF_MTP_CHAIN_SPLIT").as_deref() == Ok("1");
12722 let t_chain = std::time::Instant::now();
12723 graph.chain_ids_init(t_next, k);
12724 let cpu_stored;
12725 {
12726 let cache = &m.kv;
12727 let cpu_k: Vec<&[f32]> = (0..nkv).map(|g| cache.head_keys(g)).collect();
12728 let cpu_v: Vec<&[f32]> = (0..nkv).map(|g| cache.head_values(g)).collect();
12729 cpu_stored = cpu_k[0].len() / hd;
12730 for j in 0..k {
12731 if !graph.encode_chain_input(
12732 embed,
12733 j as u32,
12734 &m.enorm,
12735 &m.hnorm,
12736 self.embed_multiplier,
12737 eh,
12738 ) {
12739 return Err(false);
12740 }
12741 let p = AttnDeviceParams {
12745 kv_id: self.mtp_kv_id(),
12746 layer: Self::MTP_LAYER_BASE,
12747 nh,
12748 nkv,
12749 hd,
12750 rd,
12751 position: position + j,
12752 scale: self.attn_scale,
12753 eps: self.rms_eps as f32,
12754 gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
12755 late_qk_norm: self.qk_norm_after_rope,
12756 output_gate: *output_gate,
12757 q_norm: q_norm.as_deref(),
12758 k_norm: k_norm.as_deref(),
12759 inv_freq: &inv_freq,
12760 cpu_k: cpu_k.clone(),
12761 cpu_v: cpu_v.clone(),
12762 cpu_stored: cpu_stored + j,
12763 o1: None,
12764 window: None,
12765 head_gate: None,
12766 };
12767 if !graph.attn_device_ok(&l, &p) || !graph.encode_attn_device(&l, &p) {
12768 return Err(false);
12769 }
12770 if draft_rows < lm.1 {
12771 if !graph.encode_lm_head_part(&m.final_norm, lm, draft_rows) {
12772 return Err(false);
12773 }
12774 } else {
12775 graph.encode_lm_head(&m.final_norm, lm);
12776 }
12777 if !graph.encode_argmax(n_arg, j as u32 + 1) {
12778 return Err(false);
12779 }
12780 if split {
12781 graph.commit();
12784 }
12785 }
12786 }
12787 let t_enc = t_chain.elapsed();
12788 if graph.sync_checked().is_err() {
12789 return Err(true);
12790 }
12791 if std::env::var_os("CMF_GRAPH_SPEC_TIME").is_some() {
12792 eprintln!(
12793 "mtp-chain: encode {:.1} ms | wait {:.1} ms (k={k}, head rows {draft_rows}{})",
12794 t_enc.as_secs_f64() * 1e3,
12795 (t_chain.elapsed() - t_enc).as_secs_f64() * 1e3,
12796 if split { ", split" } else { "" }
12797 );
12798 }
12799 let mut ids = vec![0u32; k];
12800 if !graph.chain_ids_read(&mut ids) {
12801 return Err(true);
12802 }
12803 let mut kbuf = vec![0f32; k * nkv * hd];
12804 let mut vbuf = vec![0f32; k * nkv * hd];
12805 if !crate::gpu_metal::kv_mirror_read_rows(
12806 self.mtp_kv_id(),
12807 Self::MTP_LAYER_BASE,
12808 nkv,
12809 hd,
12810 cpu_stored,
12811 k,
12812 &mut kbuf,
12813 &mut vbuf,
12814 ) {
12815 return Err(true);
12816 }
12817 for r in 0..k {
12818 m.kv.append(
12819 &kbuf[r * nkv * hd..(r + 1) * nkv * hd],
12820 &vbuf[r * nkv * hd..(r + 1) * nkv * hd],
12821 &[],
12822 );
12823 }
12824 Ok(ids)
12825 }
12826
12827 fn try_batch_graph_wgpu(
12828 &self,
12829 hiddens: &mut [f32],
12830 positions: &[usize],
12831 k: usize,
12832 spec: Option<crate::gpu::SpecTail<'_>>,
12833 ) -> crate::gpu::BatchGraphOutcome {
12834 self.try_batch_graph_wgpu_prefix(hiddens, positions, k, spec, None)
12835 }
12836
12837 fn try_batch_graph_wgpu_prefix(
12842 &self,
12843 hiddens: &mut [f32],
12844 positions: &[usize],
12845 k: usize,
12846 spec: Option<crate::gpu::SpecTail<'_>>,
12847 layers_run: Option<&mut usize>,
12848 ) -> crate::gpu::BatchGraphOutcome {
12849 let graph_end = match self.mimo_moe.graph_prefix_end() {
12850 Some(end) if end < self.num_layers => {
12851 if layers_run.is_none() || spec.is_some() || end == 0 {
12852 return crate::gpu::BatchGraphOutcome::Declined;
12853 }
12854 end
12855 }
12856 _ => self.num_layers,
12857 };
12858 let _tb = std::time::Instant::now();
12859 let batch_debug = std::env::var_os("CMF_BATCH_DEBUG").is_some();
12860 if self.attn_softcap > 0.0 {
12861 return crate::gpu::BatchGraphOutcome::Declined; }
12863 if let Some(reason) = self.wgpu_graph_attn_decline() {
12866 self.note_graph_decline("wgpu batch graph", reason);
12867 return crate::gpu::BatchGraphOutcome::Declined;
12868 }
12869 let nh = self.num_heads;
12870 let (nkv, hd, rd) = self.layer_geom(0);
12871 let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
12872 fn gw(t: &QTensor) -> Option<crate::gpu::GraphW<'_>> {
12873 if let Some((m, i, kind, rs)) = t
12874 .graph_weight()
12875 .or_else(|| t.graph_weight_descriptor())
12876 {
12877 let name = &m.tensors[i].name;
12878 let prism = if crate::prism::is_inverse_embedding(m, name) {
12879 crate::gpu::GraphPrismOp::InverseEmbedding
12880 } else if crate::prism::is_forward_weight(m, name) {
12881 crate::gpu::GraphPrismOp::Forward
12882 } else {
12883 crate::gpu::GraphPrismOp::None
12884 };
12885 return Some(crate::gpu::GraphW {
12886 idx: i,
12887 kind,
12888 row_scale: rs,
12889 data: &[],
12890 prism,
12891 affine: crate::prism::is_affine_target(m, name),
12892 });
12893 }
12894 if std::env::var_os("CMF_BATCH_DEBUG").is_some() {
12895 eprintln!(
12896 "batch graph: tensor has no graph descriptor/f32 fallback rows={} cols={}",
12897 t.rows(),
12898 t.cols()
12899 );
12900 }
12901 t.as_f32().map(|d| crate::gpu::GraphW {
12902 idx: 0,
12903 kind: 4,
12904 row_scale: &[],
12905 data: d,
12906 prism: crate::gpu::GraphPrismOp::None,
12907 affine: false,
12908 })
12909 }
12910 let built: Option<(
12911 Vec<crate::gpu::GraphLayer<'_>>,
12912 std::sync::Arc<cortiq_core::CmfModel>,
12913 )> = (|| {
12914 let mut layers = Vec::with_capacity(graph_end);
12915 let mut model = None;
12916 for li in 0..graph_end {
12917 let lw = &self.weights.layers[self.phys_layer(li)];
12918 let gffn = match &lw.ffn {
12925 FfnKind::Dense(d) if !d.segs.is_empty() => {
12926 if batch_debug {
12927 eprintln!("batch graph: dense segmented FFN at layer {li}");
12928 }
12929 return None;
12930 }
12931 FfnKind::Dense(d) => {
12932 let Some(act) = d.act.graph_act() else {
12933 if batch_debug {
12934 eprintln!(
12935 "batch graph: dense FFN activation {:?} without a graph kernel at layer {li}",
12936 d.act
12937 );
12938 }
12939 return None;
12940 };
12941 crate::gpu::GraphFfn::Dense {
12942 gate: gw(&d.gate_proj)?,
12943 up: gw(&d.up_proj)?,
12944 down: gw(&d.down_proj)?,
12945 act,
12946 }
12947 }
12948 FfnKind::Moe(m) => {
12949 if m.route_tau.is_some() || m.mask.is_some() {
12956 return None;
12957 }
12958 let shared = m.shared.as_ref();
12962 let has_shared = shared.is_some();
12963 let shared_gated = matches!(shared, Some((_, Some(_))));
12964 let sgate = match shared {
12965 Some((_, Some(sg))) => gw(sg)?,
12966 _ => gw(&m.router)?,
12970 };
12971 let router = gw(&m.router)?;
12972 if router.prism != crate::gpu::GraphPrismOp::None
12978 || router.affine
12979 || sgate.prism != crate::gpu::GraphPrismOp::None
12980 || sgate.affine
12981 {
12982 return None;
12983 }
12984 let inter = m.experts.first()?.gate_proj.rows();
12985 let mut experts = Vec::with_capacity(m.experts.len() + 1);
12986 let mut q4tp: Option<bool> = None;
12987 let mut gu_q2: Option<bool> = None;
12988 for e in m.experts.iter().chain(shared.map(|(se, _)| se)) {
12989 if !matches!(e.act, Act::Silu)
12990 || e.gate_proj.rows() != inter
12991 || e.up_proj.rows() != inter
12992 {
12993 return None;
12994 }
12995 let (mm, gi, ui, di, is_p, is_q2) = match e.gate_proj.mapped_q4t() {
12999 Some((mm, gi)) => (
13000 mm,
13001 gi,
13002 e.up_proj.mapped_q4t()?.1,
13003 e.down_proj.mapped_q4t()?.1,
13004 false,
13005 false,
13006 ),
13007 None => match e.gate_proj.mapped_q2tp() {
13008 Some((mm, gi)) => (
13009 mm,
13010 gi,
13011 e.up_proj.mapped_q2tp()?.1,
13012 e.down_proj.mapped_q4tp()?.1,
13013 true,
13014 true,
13015 ),
13016 None => {
13017 let (mm, gi) = e.gate_proj.mapped_q4tp()?;
13018 (
13019 mm,
13020 gi,
13021 e.up_proj.mapped_q4tp()?.1,
13022 e.down_proj.mapped_q4tp()?.1,
13023 true,
13024 false,
13025 )
13026 }
13027 },
13028 };
13029 if *q4tp.get_or_insert(is_p) != is_p
13030 || *gu_q2.get_or_insert(is_q2) != is_q2
13031 {
13032 return None;
13033 }
13034 if [gi, ui, di].into_iter().any(|idx| {
13035 mm.tensors
13036 .get(idx)
13037 .is_some_and(|t| {
13038 crate::prism::is_forward_weight(mm, &t.name)
13039 || crate::prism::is_affine_target(mm, &t.name)
13040 })
13041 }) {
13042 return None;
13043 }
13044 model.get_or_insert_with(|| mm.clone());
13045 experts.push((gi, ui, di));
13046 }
13047 crate::gpu::GraphFfn::Moe {
13048 router,
13049 shared_gate: sgate,
13050 experts,
13051 n_exp: m.experts.len(),
13052 top_k: m.top_k,
13053 inter,
13054 norm_topk: m.norm_topk_prob,
13055 q4tp: q4tp?,
13056 gu_q2: gu_q2.unwrap_or(false),
13057 sigmoid: m.router_sigmoid,
13058 bias: m.expert_bias.as_deref(),
13059 has_shared,
13060 shared_gated,
13061 route_scale: m.routed_scaling,
13062 }
13063 }
13064 _ => return None,
13065 };
13066 let attn = match &lw.attn {
13067 AttnKind::Full {
13068 wq,
13069 wk,
13070 wv,
13071 wo,
13072 q_norm,
13073 k_norm,
13074 output_gate,
13075 softplus_gate,
13076 bias,
13077 } => {
13078 if softplus_gate.is_some() || self.attention_heads_per_layer.is_some() {
13079 if batch_debug {
13080 eprintln!(
13081 "batch graph: unsupported Full attention gate at layer {li} softplus={} heads={}",
13082 softplus_gate.is_some(),
13083 self.attention_heads_per_layer.is_some()
13084 );
13085 }
13086 return None;
13087 }
13088 let (m, _, _, _) = wq
13089 .graph_weight()
13090 .or_else(|| wq.graph_weight_descriptor())?;
13091 model = Some(m.clone());
13092 crate::gpu::GraphAttn::Full {
13093 wq: gw(wq)?,
13094 wk: gw(wk)?,
13095 wv: gw(wv)?,
13096 wo: gw(wo)?,
13097 q_norm: q_norm.as_deref(),
13098 k_norm: k_norm.as_deref(),
13099 late_qk_norm: self.qk_norm_after_rope,
13100 bias: bias
13101 .as_ref()
13102 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
13103 output_gate: *output_gate,
13104 cpu_k: self.kv_cache.layers[li].k_heads(),
13105 cpu_v: self.kv_cache.layers[li].v_heads(),
13106 geom: self.graph_attn_geom(li),
13107 head_gate: None,
13110 }
13111 }
13112 AttnKind::LinearGdn(w) => {
13113 let Some(cfg) = self.gdn_cfg else {
13114 if batch_debug {
13115 eprintln!("batch graph: no GDN config at layer {li}");
13116 }
13117 return None;
13118 };
13119 let (m, _, _, _) = w
13120 .in_proj_qkv
13121 .graph_weight()
13122 .or_else(|| w.in_proj_qkv.graph_weight_descriptor())?;
13123 model = Some(m.clone());
13124 crate::gpu::GraphAttn::Gdn {
13125 qkv: gw(&w.in_proj_qkv)?,
13126 z: gw(&w.in_proj_z)?,
13127 a: gw(&w.in_proj_a)?,
13128 b: gw(&w.in_proj_b)?,
13129 out: gw(&w.out_proj)?,
13130 conv1d: &w.conv1d,
13131 a_log: &w.a_log,
13132 dt_bias: &w.dt_bias,
13133 norm: &w.norm,
13134 nv: cfg.num_v_heads,
13135 nk: cfg.num_k_heads,
13136 dk: cfg.key_head_dim,
13137 dv: cfg.value_head_dim,
13138 kk: cfg.conv_kernel,
13139 cpu_state: &self.kv_cache.layers[self.phys_layer(li)].linear_state,
13140 }
13141 }
13142 _ => return None,
13143 };
13144 layers.push(crate::gpu::GraphLayer {
13145 input_norm: &lw.input_norm,
13146 attn,
13147 post_norm: &lw.post_norm,
13148 ffn: gffn,
13149 });
13150 }
13151 Some((layers, model?))
13152 })();
13153 let Some((layers, model)) = built else {
13154 {
13155 use std::sync::atomic::{AtomicBool, Ordering};
13156 static SAID: AtomicBool = AtomicBool::new(false);
13157 if !SAID.swap(true, Ordering::Relaxed) {
13158 tracing::warn!("batch graph: BUILDER refused (layer weights/kinds)");
13159 }
13160 }
13161 return crate::gpu::BatchGraphOutcome::Declined;
13162 };
13163 if std::env::var("CMF_GRAPH_SPEC_TIME").is_ok() {
13164 eprintln!("batch-build: {:.1} ms", _tb.elapsed().as_secs_f64() * 1e3);
13165 }
13166 crate::gpu::forward_batch_graph(
13167 &model,
13168 self.graph_kv_id,
13169 &layers,
13170 &self.inv_freq,
13171 hiddens,
13172 nh,
13173 nkv,
13174 hd,
13175 rd,
13176 self.hidden_size,
13177 self.intermediate_size,
13178 positions,
13179 self.kv_cache.max_seq_len,
13180 gemma,
13181 self.rms_eps as f32,
13182 self.attn_scale,
13183 k,
13184 &(0..graph_end)
13185 .map(|li| self.kv_cache.layers[self.phys_layer(li)].o1_views())
13186 .collect::<Vec<_>>(),
13187 self.o1_epoch,
13188 spec,
13189 layers_run,
13190 )
13191 }
13192
13193 fn draft_probe() -> bool {
13197 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
13198 *ON.get_or_init(|| std::env::var("CMF_DSV4_DRAFT_PROBE").is_ok_and(|v| v != "0"))
13199 }
13200
13201 #[cfg(feature = "gpu")]
13213 fn dsv4_spec_on() -> bool {
13214 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
13215 *ON.get_or_init(|| {
13216 if let Ok(v) = std::env::var("CMF_DSV4_SPEC_RUN") {
13220 return v != "0";
13221 }
13222 std::env::var("CMF_DSV4_SPEC")
13229 .map(|v| v != "0")
13230 .unwrap_or_else(|_| {
13231 crate::gpu_wgpu::DRAFT_RESERVE.load(std::sync::atomic::Ordering::Relaxed) > 0
13232 })
13233 })
13234 }
13235
13236 #[cfg(feature = "gpu")]
13243 fn dsv4_spec_step(
13244 &mut self,
13245 tip_token: u32,
13246 t_next: u32,
13247 next_pos: usize,
13248 max_extra: usize,
13249 drafted: &mut usize,
13250 accepted_ctr: &mut usize,
13251 ) -> Option<(Vec<u32>, usize)> {
13252 let t_all = std::time::Instant::now();
13253 if std::env::var("CMF_DSV4_SPEC_TIME").is_ok() {
13254 thread_local! {
13255 static LAST: std::cell::Cell<Option<std::time::Instant>> =
13256 const { std::cell::Cell::new(None) };
13257 }
13258 LAST.with(|l| {
13259 if let Some(prev) = l.get() {
13260 eprintln!(
13261 "между раундами {:.1} мс",
13262 prev.elapsed().as_secs_f64() * 1e3
13263 );
13264 }
13265 l.set(Some(std::time::Instant::now()));
13266 });
13267 }
13268 if std::env::var("CMF_DSV4_SPEC_DEBUG").is_ok() {
13269 eprintln!("spec_step: вход pos={next_pos}");
13270 }
13271 let n_layers = self.dsv4.as_ref().map(|b| b.1.len())?;
13272 let cfg = self.dsv4.as_ref().map(|b| b.2)?;
13273 if self.dspark.is_none() {
13275 let t = crate::dsv4::dspark_targets(&self.dsv4_mtp, &cfg, n_layers);
13276 if t.is_empty() {
13277 return None;
13278 }
13279 crate::dsv4::dspark_arm(&t, cfg.dim);
13280 self.dspark = Some(crate::dsv4::DsparkState::new(
13281 self.dsv4_mtp.len(),
13282 &cfg,
13283 t.len(),
13284 ));
13285 }
13286 let targets = crate::dsv4::dspark_targets(&self.dsv4_mtp, &cfg, n_layers);
13287 let pack = crate::dsv4::dspark_pack_get(&self.dsv4_mtp, &cfg);
13288 if pack.is_none() && std::env::var("CMF_DSV4_SPEC_DEBUG").is_ok() {
13289 eprintln!("spec_step: пак не построился (targets {targets:?})");
13290 }
13291 let pack = pack?;
13292 let block = crate::dsv4::dspark_block();
13293 let b_box = self.dsv4.as_mut()?;
13294 let (g, layers, st) = (&b_box.0, &b_box.1, &mut b_box.3);
13295 let ds = self.dspark.as_mut()?;
13296 let dbg = std::env::var("CMF_DSV4_SPEC_DEBUG").is_ok();
13299 if !crate::dsv4::dspark_take(&mut ds.main_hidden) && !ds.have_hidden {
13300 if dbg {
13301 eprintln!("spec_step: нет захвата");
13302 }
13303 return None;
13304 }
13305 ds.have_hidden = true;
13306 let tip_pos = next_pos.checked_sub(1)?;
13307 let draft_started = std::time::Instant::now();
13308 let mut conf = Vec::new();
13309 let props = crate::dsv4::dspark_draft_gpu(
13310 g,
13311 &self.dsv4_mtp,
13312 &cfg,
13313 ds,
13314 pack,
13315 st.kv_id,
13316 tip_token,
13317 tip_pos,
13318 self.pool.as_deref(),
13319 &mut conf,
13320 );
13321 self.dspark_draft_ns += draft_started.elapsed().as_nanos();
13322 *drafted += block;
13323 if props.is_empty() || props[0] != t_next {
13324 if dbg {
13325 eprintln!(
13326 "spec_step: черновик {} (props0={:?} t_next={t_next})",
13327 if props.is_empty() {
13328 "пуст"
13329 } else {
13330 "мимо"
13331 },
13332 props.first()
13333 );
13334 }
13335 return None;
13336 }
13337 let mut k_verify = crate::dsv4::dspark_verify_k()
13344 .min(props.len())
13345 .min(max_extra.saturating_add(1));
13346 let conf_min = {
13352 static M: std::sync::OnceLock<f32> = std::sync::OnceLock::new();
13353 *M.get_or_init(|| {
13354 std::env::var("CMF_DSPARK_CONF_MIN")
13355 .ok()
13356 .and_then(|v| v.parse().ok())
13357 .unwrap_or(0.0)
13358 })
13359 };
13360 if conf_min > 0.0 && conf.len() >= props.len() {
13361 let mut keep = 1usize;
13362 while keep < k_verify && conf.get(keep).copied().unwrap_or(0.0) >= conf_min {
13363 keep += 1;
13364 }
13365 k_verify = k_verify.min(keep.max(2));
13366 }
13367 if k_verify < 2 {
13368 return None;
13369 }
13370 let mut fed = Vec::with_capacity(k_verify);
13371 fed.push(t_next);
13372 fed.extend_from_slice(&props[1..k_verify]);
13373 let mut argmax = Vec::new();
13374 let mut logits_all = Vec::new();
13375 let mut walked = Vec::new();
13376 let txn = crate::dsv4::dsv4_verify_chunk(
13377 g,
13378 layers,
13379 &cfg,
13380 st,
13381 &fed,
13382 next_pos,
13383 &self.inv_freq,
13384 self.pool.as_deref(),
13385 &targets,
13386 &mut argmax,
13387 &mut logits_all,
13388 &mut walked,
13389 );
13390 if txn.is_none() && dbg {
13391 eprintln!("spec_step: verify отказал");
13392 }
13393 let txn = txn?;
13394 let spec_gpu_end = txn.gpu_end;
13395 let b = fed.len();
13396 let mut accepted = 1usize;
13397 while accepted < b && fed[accepted] == argmax[accepted - 1] {
13398 accepted += 1;
13399 }
13400 if std::env::var("CMF_DSV4_SPEC_FORCE_REJECT").is_ok_and(|v| v != "0") {
13405 accepted = 1;
13406 }
13407 if std::env::var("CMF_DSV4_SPEC_TRACE").is_ok() {
13408 eprintln!("spec@{next_pos}: fed={fed:?} argmax={argmax:?} accepted={accepted}");
13409 }
13410 let t_fin = std::time::Instant::now();
13411 if !crate::dsv4::dsv4_spec_finish(
13412 g,
13413 layers,
13414 &cfg,
13415 st,
13416 txn,
13417 accepted,
13418 &fed,
13419 &self.inv_freq,
13420 self.pool.as_deref(),
13421 ) {
13422 tracing::warn!("dsv4: спекулятивный откат не удался — состояние подозрительно");
13423 return None;
13424 }
13425 if std::env::var("CMF_DSV4_SPEC_TIME").is_ok() {
13426 eprintln!(
13427 "finish(k={accepted}): {:.1} мс",
13428 t_fin.elapsed().as_secs_f64() * 1e3
13429 );
13430 }
13431 *accepted_ctr += accepted - 1;
13432 let (hc, dim) = (cfg.hc_mult, cfg.dim);
13437 let dev_caps: Vec<usize> = targets
13442 .iter()
13443 .copied()
13444 .filter(|&t| t < spec_gpu_end)
13445 .collect();
13446 let mut caps_all = vec![0.0f32; dev_caps.len() * b * hc * dim];
13447 if !crate::gpu_wgpu::dsv4_spec_cap_read_all(b, dev_caps.len(), hc * dim, &mut caps_all) {
13448 return None;
13449 }
13450 for t in 0..accepted {
13451 let tip = t + 1 == accepted;
13452 for (slot, &tl) in targets.iter().enumerate() {
13453 if let Some(di) = dev_caps.iter().position(|&d| d == tl) {
13454 let lo = (di * b + t) * hc * dim;
13455 crate::dsv4::dspark_capture(
13456 &caps_all[lo..lo + hc * dim],
13457 &cfg,
13458 slot,
13459 &mut ds.main_hidden,
13460 );
13461 } else if tip
13462 && crate::dsv4::dspark_peek_slot(slot, dim, {
13463 let lo = slot * dim;
13464 &mut ds.main_hidden[lo..lo + dim]
13465 })
13466 {
13467 } else {
13472 crate::dsv4::dspark_capture(
13476 &walked[t * hc * dim..(t + 1) * hc * dim],
13477 &cfg,
13478 slot,
13479 &mut ds.main_hidden,
13480 );
13481 }
13482 }
13483 crate::dsv4::dspark_ring_append(
13484 g,
13485 &self.dsv4_mtp,
13486 &cfg,
13487 ds,
13488 next_pos + t,
13489 self.pool.as_deref(),
13490 );
13491 }
13492 let row = logits_all[(accepted - 1) * cfg.vocab..accepted * cfg.vocab].to_vec();
13493 self.graph_logits = Some(row);
13494 if std::env::var("CMF_DSV4_TRUNK_PICK_DUMP").is_ok() {
13499 crate::dsv4::trunk_freq_note(&crate::dsv4::pick_tally_take());
13500 crate::dsv4::pick_tally_arm();
13501 }
13502 if std::env::var("CMF_DSV4_SPEC_TIME").is_ok() {
13503 eprintln!(
13504 "spec_step total {:.1} мс (k={accepted})",
13505 t_all.elapsed().as_secs_f64() * 1e3
13506 );
13507 }
13508 Some((fed[1..accepted].to_vec(), next_pos + accepted))
13509 }
13510
13511 fn dspark_probe(&mut self, position: usize, token_id: u32) {
13512 if self.dsv4_mtp.is_empty() || !Self::draft_probe() {
13513 return;
13514 }
13515 let trunk_now = crate::dsv4::pick_tally_take();
13517 crate::dsv4::trunk_freq_note(&trunk_now);
13518 if !trunk_now.is_empty() {
13519 self.dspark_trunk_picks.push(trunk_now);
13520 let keep = crate::dsv4::dspark_block();
13521 if self.dspark_trunk_picks.len() > keep {
13522 self.dspark_trunk_picks.remove(0);
13523 }
13524 }
13525 for p in std::mem::take(&mut self.dspark_pending) {
13528 let Some(i) = position.checked_sub(p.0 + 1) else {
13529 continue;
13530 };
13531 let mut p = p;
13532 if i < p.1.len() {
13533 if p.2 && p.1[i] == token_id {
13534 p.3 = i + 1;
13535 } else {
13536 p.2 = false;
13537 }
13538 if i + 1 < p.1.len() {
13539 self.dspark_pending.push(p);
13540 continue;
13541 }
13542 }
13543 self.dspark_hist.push(p.3);
13544 self.dspark_real.push(token_id);
13545 }
13546 let Some(b) = &mut self.dsv4 else { return };
13547 let (g, layers, cfg) = (&b.0, &b.1, b.2);
13548 let n_layers = layers.len();
13549 if self.dspark.is_none() {
13550 let t = crate::dsv4::dspark_targets(&self.dsv4_mtp, &cfg, n_layers);
13551 if t.is_empty() {
13552 return;
13553 }
13554 eprintln!(
13555 "DSpark: захват со слоёв {t:?}, блок {}",
13556 crate::dsv4::dspark_block()
13557 );
13558 crate::dsv4::dspark_arm(&t, cfg.dim);
13559 self.dspark = Some(crate::dsv4::DsparkState::new(
13560 self.dsv4_mtp.len(),
13561 &cfg,
13562 t.len(),
13563 ));
13564 }
13565 let ds = self.dspark.as_mut().unwrap();
13566 if !crate::dsv4::dspark_take(&mut ds.main_hidden) {
13567 return; }
13569 let mut conf = Vec::new();
13570 crate::dsv4::pick_tally_arm();
13571 let draft_started = std::time::Instant::now();
13576 #[cfg(feature = "gpu")]
13577 let gpu_draft = crate::dsv4::dspark_gpu_on();
13578 #[cfg(not(feature = "gpu"))]
13579 let gpu_draft = false;
13580 let props = if gpu_draft {
13581 #[cfg(feature = "gpu")]
13582 {
13583 let kv_id = b.3.kv_id;
13584 match crate::dsv4::dspark_pack_get(&self.dsv4_mtp, &cfg) {
13585 Some(pk) => crate::dsv4::dspark_draft_gpu(
13586 g,
13587 &self.dsv4_mtp,
13588 &cfg,
13589 ds,
13590 pk,
13591 kv_id,
13592 token_id,
13593 position,
13594 self.pool.as_deref(),
13595 &mut conf,
13596 ),
13597 None => Vec::new(),
13598 }
13599 }
13600 #[cfg(not(feature = "gpu"))]
13601 Vec::new()
13602 } else {
13603 crate::gpu::cpu_scope(|| {
13604 crate::dsv4::dspark_draft(
13605 g,
13606 &self.dsv4_mtp,
13607 &cfg,
13608 ds,
13609 token_id,
13610 position,
13611 self.pool.as_deref(),
13612 &mut conf,
13613 )
13614 })
13615 };
13616 self.dspark_draft_ns += draft_started.elapsed().as_nanos();
13617 let draft_picks = crate::dsv4::pick_tally_take();
13618 crate::dsv4::dspark_freq_note(&draft_picks);
13619 crate::dsv4::pick_tally_arm();
13622 if !props.is_empty() {
13623 let (tu, tt) = {
13627 let flat: Vec<(usize, Vec<usize>)> = self
13628 .dspark_trunk_picks
13629 .iter()
13630 .flat_map(|v| v.iter().cloned())
13631 .collect();
13632 let mut per: std::collections::HashMap<usize, Vec<usize>> =
13634 std::collections::HashMap::new();
13635 for (li, picks) in flat {
13636 per.entry(li).or_default().extend(picks);
13637 }
13638 let n = per.len().max(1);
13639 let mut u = 0usize;
13640 let mut t = 0usize;
13641 for (_, v) in per {
13642 t += v.len();
13643 u += v.iter().collect::<std::collections::HashSet<_>>().len();
13644 }
13645 (u / n, t / n)
13646 };
13647 let (du, dt) = crate::dsv4::tally_unique(&draft_picks);
13648 self.dspark_exp.push((tu, tt, du, dt));
13649 self.dspark_pending.push((position, props, true, 0));
13650 }
13651 if self.dspark_hist.len() >= 8 && self.dspark_hist.len() % 8 == 0 {
13652 let n = self.dspark_hist.len() as f32;
13653 let mean: f32 = self.dspark_hist.iter().sum::<usize>() as f32 / n;
13654 let block = crate::dsv4::dspark_block();
13655 let mut at = vec![0usize; block + 1];
13656 for &k in &self.dspark_hist {
13657 at[k] += 1;
13658 }
13659 let mut surv = Vec::with_capacity(block);
13661 for i in 1..=block {
13662 let k = at[i..].iter().sum::<usize>() as f32 / n;
13663 surv.push(format!("{k:.2}"));
13664 }
13665 let distinct = self
13666 .dspark_real
13667 .iter()
13668 .collect::<std::collections::HashSet<_>>()
13669 .len();
13670 let (tu, tt, du, dt) = self.dspark_exp.iter().fold((0, 0, 0, 0), |a, b| {
13671 (a.0 + b.0, a.1 + b.1, a.2 + b.2, a.3 + b.3)
13672 });
13673 let m = self.dspark_exp.len().max(1);
13674 eprintln!(
13675 "DSpark: черновиков {}, принято в среднем {mean:.2} из {block} \
13676 (токенов за проход {:.2}), распределение {at:?}, выживание [{}]",
13677 self.dspark_hist.len(),
13678 mean + 1.0,
13679 surv.join(" ")
13680 );
13681 eprintln!(
13682 "DSpark: разных токенов {distinct} из {} (вырожденность), \
13683 эксперты ствол {}/{} на слой за {block} токенов, \
13684 черновик {}/{} за блок, draft {:.2} мс/блок",
13685 self.dspark_real.len(),
13686 tu / m,
13687 tt / m,
13688 du / m,
13689 dt / m,
13690 self.dspark_draft_ns as f64 / self.dspark_exp.len().max(1) as f64 / 1e6
13691 );
13692 }
13693 }
13694
13695 fn forward_layers_upto(
13696 &mut self,
13697 hidden: &[f32],
13698 position: usize,
13699 task_mask: Option<&TaskMask>,
13700 upto: Option<usize>,
13701 ) -> Vec<f32> {
13702 if let Some(plan) = self.gpu_plan.clone() {
13708 if upto.is_none() && plan.len() > 1 {
13709 let mut h = hidden.to_vec();
13710 for &(dev, from, upto_incl) in plan.iter() {
13711 h = crate::gpu::with_device(dev, || {
13712 self.forward_layers_span(&h, position, task_mask, from, Some(upto_incl))
13713 });
13714 }
13715 return h;
13716 }
13717 }
13718 self.forward_layers_span(hidden, position, task_mask, 0, upto)
13719 }
13720
13721 pub fn set_gpu_plan(&mut self, devices: Option<&[usize]>) -> Result<(), String> {
13726 self.set_gpu_plan_at(devices, None)
13727 }
13728
13729 pub fn set_gpu_plan_at(
13733 &mut self,
13734 devices: Option<&[usize]>,
13735 at: Option<usize>,
13736 ) -> Result<(), String> {
13737 let Some(devs) = devices.filter(|d| d.len() > 1) else {
13738 self.gpu_plan = None;
13739 return Ok(());
13740 };
13741 self.split_supported()?;
13742 let n = self.num_layers;
13743 if devs.len() > n {
13744 return Err(format!("{} devices for {n} layers", devs.len()));
13745 }
13746 if let Some(k) = at {
13747 if k == 0 || k >= n {
13748 return Err(format!("split at {k}: the model has {n} layers"));
13749 }
13750 if devs.len() == 2 {
13751 self.gpu_plan = Some(std::sync::Arc::new(vec![
13752 (devs[0], 0, k - 1),
13753 (devs[1], k, n - 1),
13754 ]));
13755 return Ok(());
13756 }
13757 return Err(format!(
13758 "an explicit split point takes exactly 2 devices, got {}",
13759 devs.len()
13760 ));
13761 }
13762 let per = n.div_ceil(devs.len());
13763 let mut plan = Vec::with_capacity(devs.len());
13764 let mut from = 0usize;
13765 for &d in devs {
13766 if from >= n {
13767 break;
13768 }
13769 let upto = (from + per - 1).min(n - 1);
13770 plan.push((d, from, upto));
13771 from = upto + 1;
13772 }
13773 self.gpu_plan = Some(std::sync::Arc::new(plan));
13774 Ok(())
13775 }
13776
13777 pub fn gpu_plan(&self) -> Option<Vec<(usize, usize, usize)>> {
13779 self.gpu_plan.as_ref().map(|p| p.as_ref().clone())
13780 }
13781
13782 fn embryo_qtensor_f32(t: &QTensor) -> Vec<f32> {
13788 if let Some(x) = t.as_f32() {
13789 return x.to_vec();
13790 }
13791 let mut out = vec![0.0; t.rows() * t.cols()];
13792 for r in 0..t.rows() {
13793 t.row_f32(r, &mut out[r * t.cols()..(r + 1) * t.cols()]);
13794 }
13795 out
13796 }
13797
13798 fn embryo_resident_eligible(&self) -> bool {
13799 if (self.vmf_cfg.is_none() && self.gdn_cfg.is_none())
13802 || self.num_layers != self.physical_layers
13803 || self.loop_final_norm
13804 || self.weights.layers.len() != self.num_layers
13805 || self.head_clusters.is_none()
13806 || self.final_softcap.is_some()
13807 || self.logit_multiplier.is_some()
13808 || self.attn_softcap != 0.0
13809 || self.mtp.is_some()
13810 || self.g3n.is_some()
13811 || self.dsv4.is_some()
13812 || self.dsv41.is_some()
13813 || self.qwen4_exp.is_some()
13814 || self.dyn_router.is_some()
13821 || self.dyn_phi_layer.is_some()
13822 || self.dyn_blend_loaded
13823 || self.o1_cfg.is_some()
13824 || self.swa.is_some()
13825 || self.sliding_layers.is_some()
13826 || self.global_attn.is_some()
13827 || self.attention_heads_per_layer.is_some()
13828 || self.attn_v_norm
13829 || self
13830 .kv_cache
13831 .layers
13832 .iter()
13833 .any(|l| l.mode != crate::kv_cache::KvMode::F32)
13834 || self.rope_scale != 1.0
13835 || self.rope_scale_local != 1.0
13836 || self.attn_scale != 1.0 / (self.head_dim as f32).sqrt()
13837 || self.hidden_size == 0
13838 || self.hidden_size > 1024
13839 || self.intermediate_size > 1024
13840 || self.num_heads == 0
13841 || self.num_kv_heads == 0
13842 || self.num_heads % self.num_kv_heads != 0
13843 || self.num_heads.saturating_mul(self.head_dim) > 1024
13844 || self.num_kv_heads.saturating_mul(self.head_dim) > 1024
13845 || self.vocab_size == 0
13846 || self.kv_cache.max_seq_len == 0
13847 || self.rotary_dim == 0
13848 || self.rotary_dim > self.head_dim
13849 || self.rotary_dim % 2 != 0
13850 || self.inv_freq.len() < self.rotary_dim / 2
13851 {
13852 return false;
13853 }
13854 if self.weights.lm_head.as_f32().is_none()
13865 || self.weights.embed_tokens.as_f32().is_none()
13866 || self.weights.lm_head.rows() < self.vocab_size
13867 || self.weights.lm_head.cols() != self.hidden_size
13868 || self.weights.embed_tokens.rows() < self.vocab_size
13869 || self.weights.embed_tokens.cols() != self.hidden_size
13870 || self.weights.final_norm.len() != self.hidden_size
13871 {
13872 return false;
13873 }
13874 if let Some(cfg) = self.vmf_cfg {
13875 if cfg.num_heads.saturating_mul(cfg.nphase) > 1024
13876 || cfg.num_heads.saturating_mul(cfg.nphase.saturating_add(1)) > 1024
13877 || cfg.num_heads.saturating_mul(cfg.value_head_dim) > 1024
13878 || cfg.state_len() == 0
13879 {
13880 return false;
13881 }
13882 }
13883 if let Some(g) = self.gdn_cfg {
13884 if g.num_v_heads == 0
13889 || g.num_k_heads == 0
13890 || g.num_v_heads % g.num_k_heads != 0
13891 || g.key_head_dim == 0
13892 || g.key_head_dim > 128
13893 || g.value_head_dim == 0
13894 || g.value_head_dim > 256
13895 || g.value_head_dim % 4 != 0
13896 || g.conv_kernel == 0
13897 || g.num_v_heads > 512
13898 || g.num_v_heads.saturating_mul(g.value_head_dim) > 1024
13899 || g.conv_dim() > 2048
13900 || g.conv_dim() % 4 != 0
13901 || g.hidden_size != self.hidden_size
13902 || g.output_gate_sigmoid
13903 || g.rms_eps != self.rms_eps
13904 || g.state_len() == 0
13905 {
13906 return false;
13907 }
13908 }
13909 let mut full_seen = false;
13910 for lw in &self.weights.layers {
13911 if lw.attn_out_norm.is_some() || lw.ffn_out_norm.is_some() || lw.layer_scale.is_some() {
13912 return false;
13913 }
13914 match &lw.attn {
13915 AttnKind::LinearGdn(w) => {
13916 let Some(g) = self.gdn_cfg else {
13917 return false;
13918 };
13919 let (nv, dv, kk) = (g.num_v_heads, g.value_head_dim, g.conv_kernel);
13920 if w.in_proj_qkv.rows() != g.conv_dim()
13921 || w.in_proj_qkv.cols() != self.hidden_size
13922 || w.in_proj_qkv.as_f32().is_none()
13923 || w.in_proj_z.rows() != nv * dv
13924 || w.in_proj_z.cols() != self.hidden_size
13925 || w.in_proj_z.as_f32().is_none()
13926 || w.in_proj_a.rows() != nv
13927 || w.in_proj_a.cols() != self.hidden_size
13928 || w.in_proj_a.as_f32().is_none()
13929 || w.in_proj_b.rows() != nv
13930 || w.in_proj_b.cols() != self.hidden_size
13931 || w.in_proj_b.as_f32().is_none()
13932 || w.conv1d.len() != g.conv_dim() * kk
13933 || w.a_log.len() != nv
13934 || w.dt_bias.len() != nv
13935 || w.norm.len() != dv
13936 || w.out_proj.rows() != self.hidden_size
13937 || w.out_proj.cols() != nv * dv
13938 || w.out_proj.as_f32().is_none()
13939 {
13940 return false;
13941 }
13942 }
13943 AttnKind::Linear(w) => {
13944 let Some(cfg) = self.vmf_cfg else {
13945 return false;
13946 };
13947 if w.thq.rows() != cfg.num_heads * cfg.nphase
13948 || w.thq.cols() != self.hidden_size
13949 || w.thq.as_f32().is_none()
13950 || w.thk.rows() != cfg.num_heads * cfg.nphase
13951 || w.thk.cols() != self.hidden_size
13952 || w.thk.as_f32().is_none()
13953 || w.v_proj.rows() != cfg.num_heads * cfg.value_head_dim
13954 || w.v_proj.cols() != self.hidden_size
13955 || w.v_proj.as_f32().is_none()
13956 || w.out_proj.rows() != self.hidden_size
13957 || w.out_proj.cols() != cfg.num_heads * cfg.value_head_dim
13958 || w.out_proj.as_f32().is_none()
13959 || w.decay.len() != cfg.num_heads * 2 * cfg.nphase
13960 {
13961 return false;
13962 }
13963 if let Some((kg, kb)) = &w.k_gate {
13964 if kg.rows() != cfg.num_heads
13965 || kg.cols() != self.hidden_size
13966 || kg.as_f32().is_none()
13967 || kb.len() != cfg.num_heads
13968 {
13969 return false;
13970 }
13971 }
13972 if let Some(conv) = &w.conv {
13973 if conv.len() % self.hidden_size != 0 || conv.len() / self.hidden_size < 2 {
13974 return false;
13975 }
13976 }
13977 }
13978 AttnKind::Full {
13979 wq,
13980 wk,
13981 wv,
13982 wo,
13983 q_norm,
13984 k_norm,
13985 output_gate,
13986 softplus_gate,
13987 bias,
13988 } => {
13989 if full_seen
13990 || q_norm.is_some()
13991 || k_norm.is_some()
13992 || *output_gate
13993 || softplus_gate.is_some()
13994 || bias.is_some()
13995 || wq.as_f32().is_none()
13996 || wk.as_f32().is_none()
13997 || wv.as_f32().is_none()
13998 || wo.as_f32().is_none()
13999 || wq.rows() != self.num_heads * self.head_dim
14000 || wk.rows() != self.num_kv_heads * self.head_dim
14001 || wv.rows() != self.num_kv_heads * self.head_dim
14002 || wq.cols() != self.hidden_size
14003 || wk.cols() != self.hidden_size
14004 || wv.cols() != self.hidden_size
14005 || wo.rows() != self.hidden_size
14006 || wo.cols() != self.num_heads * self.head_dim
14007 {
14008 return false;
14009 }
14010 full_seen = true;
14011 }
14012 AttnKind::Bounded(w) => {
14013 let Some(ac) = self.anchor_core.as_ref() else {
14016 return false;
14017 };
14018 if self.bounded_rope.is_none()
14019 || w.window != ac.window
14020 || w.sink != ac.sink
14021 || w.window == 0
14022 || w.window + w.sink > 256
14023 || w.sink_k.len() != self.num_kv_heads * w.sink * self.head_dim
14024 || w.sink_v.len() != self.num_kv_heads * w.sink * self.head_dim
14025 || w.wq.as_f32().is_none()
14026 || w.wk.as_f32().is_none()
14027 || w.wv.as_f32().is_none()
14028 || w.wo.as_f32().is_none()
14029 || w.wq.rows() != self.num_heads * self.head_dim
14030 || w.wk.rows() != self.num_kv_heads * self.head_dim
14031 || w.wv.rows() != self.num_kv_heads * self.head_dim
14032 || w.wq.cols() != self.hidden_size
14033 || w.wk.cols() != self.hidden_size
14034 || w.wv.cols() != self.hidden_size
14035 || w.wo.rows() != self.hidden_size
14036 || w.wo.cols() != self.num_heads * self.head_dim
14037 {
14038 return false;
14039 }
14040 }
14041 _ => return false,
14042 }
14043 match &lw.ffn {
14044 FfnKind::Dense(d) => {
14045 if d.act != Act::Silu
14046 || !d.segs.is_empty()
14047 || d.gate_proj.as_f32().is_none()
14048 || d.up_proj.as_f32().is_none()
14049 || d.down_proj.as_f32().is_none()
14050 || d.gate_proj.rows() != self.intermediate_size
14051 || d.gate_proj.cols() != self.hidden_size
14052 || d.up_proj.rows() != self.intermediate_size
14053 || d.up_proj.cols() != self.hidden_size
14054 || d.down_proj.rows() != self.hidden_size
14055 || d.down_proj.cols() != self.intermediate_size
14056 {
14057 return false;
14058 }
14059 }
14060 FfnKind::Moe(m) => {
14061 if m.resonance.is_none()
14062 || m.top_k != 1
14063 || m.router_sigmoid
14064 || !m.norm_topk_prob
14065 || m.expert_bias.is_some()
14066 || m.routed_scaling != 1.0
14067 || m.route_tau.is_some()
14068 || m.shared.is_none()
14069 || m.mask.is_some()
14070 || m.per_expert_scale.is_some()
14071 || m.router_input_norm
14072 || m.experts.is_empty()
14073 || m.experts.len() > 8
14074 {
14075 return false;
14076 }
14077 let r = m.resonance.as_ref().unwrap();
14078 if r.mu.len() != m.experts.len() * self.hidden_size
14079 || r.bias.len() != m.experts.len()
14080 || r.u.len() != m.experts.len() * r.k * self.hidden_size
14081 || r.k > 128
14082 {
14083 return false;
14084 }
14085 let Some((shared, gate)) = &m.shared else {
14086 return false;
14087 };
14088 if gate.is_some() || shared.act != Act::Silu || !shared.segs.is_empty() {
14089 return false;
14090 }
14091 if shared.gate_proj.as_f32().is_none()
14092 || shared.up_proj.as_f32().is_none()
14093 || shared.down_proj.as_f32().is_none()
14094 || shared.gate_proj.rows() != self.intermediate_size
14095 || shared.gate_proj.cols() != self.hidden_size
14096 || shared.up_proj.rows() != self.intermediate_size
14097 || shared.up_proj.cols() != self.hidden_size
14098 || shared.down_proj.rows() != self.hidden_size
14099 || shared.down_proj.cols() != self.intermediate_size
14100 {
14101 return false;
14102 }
14103 for e in &m.experts {
14104 if e.act != Act::Silu
14105 || !e.segs.is_empty()
14106 || e.gate_proj.as_f32().is_none()
14107 || e.up_proj.as_f32().is_none()
14108 || e.down_proj.as_f32().is_none()
14109 || e.gate_proj.rows() != self.intermediate_size
14110 || e.gate_proj.cols() != self.hidden_size
14111 || e.up_proj.rows() != self.intermediate_size
14112 || e.up_proj.cols() != self.hidden_size
14113 || e.down_proj.rows() != self.hidden_size
14114 || e.down_proj.cols() != self.intermediate_size
14115 {
14116 return false;
14117 }
14118 }
14119 }
14120 FfnKind::DenseMoe(_) => return false,
14121 }
14122 if lw.input_norm.len() != self.hidden_size || lw.post_norm.len() != self.hidden_size {
14123 return false;
14124 }
14125 }
14126 if full_seen && self.anchor_core.is_some() {
14127 return false;
14128 }
14129 full_seen || self.num_layers > 0
14130 }
14131
14132 fn embryo_resident_wanted(&self) -> bool {
14137 crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Decode)
14138 && crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Prefill)
14139 && matches!(
14140 std::env::var("CMF_EMBRYO_RESIDENT").as_deref(),
14141 Ok("1") | Ok("parallel")
14142 )
14143 && crate::gpu::enabled_here()
14144 && !self.graph_refused()
14145 && self.embryo_resident_eligible()
14146 }
14147
14148 fn embryo_prefill_chunked(&mut self, ids: &[u32], start: usize) -> Option<Vec<f32>> {
14156 if ids.len() < 2
14157 || std::env::var("CMF_EMBRYO_CHUNK").as_deref() == Ok("0")
14158 || !self.embryo_resident_wanted()
14159 {
14160 return None;
14161 }
14162 let model = self.ensure_embryo_graph()?;
14163 let cmax = std::env::var("CMF_EMBRYO_CHUNK")
14164 .ok()
14165 .and_then(|v| v.parse::<usize>().ok())
14166 .filter(|&v| v >= 1)
14167 .unwrap_or(crate::gpu::EMBRYO_CHUNK_MAX)
14168 .min(crate::gpu::EMBRYO_CHUNK_MAX);
14169 let hs = self.hidden_size;
14170 let n = ids.len();
14171 let mut pos = start;
14172 let mut last = None;
14173 let mut rows = Vec::with_capacity(cmax * hs);
14174 while pos < n {
14175 let end = (pos + cmax).min(n);
14176 rows.clear();
14177 for &id in &ids[pos..end] {
14178 rows.extend_from_slice(&self.embed_single(id));
14179 }
14180 let mut lg = Vec::new();
14181 if !crate::gpu::forward_embryo_graph_chunk(
14182 &model,
14183 self.graph_kv_id,
14184 &rows,
14185 pos,
14186 end - pos,
14187 &mut lg,
14188 ) {
14189 if pos == start {
14190 if std::env::var("CMF_EMBRYO_DBG").as_deref() == Ok("1") {
14191 eprintln!("embryo-dbg: chunk prefill refused at position {pos}");
14192 }
14193 return None;
14194 }
14195 self.kv_cache.clear();
14199 self.clear_history();
14200 crate::gpu::graph_kv_reset(self.graph_kv_id);
14201 panic!(
14202 "resident Embryo chunk prefill refused at position {pos}; refusing an in-flight host fallback"
14203 );
14204 }
14205 last = Some(lg);
14206 pos = end;
14207 }
14208 if std::env::var("CMF_EMBRYO_DBG").as_deref() == Ok("1") {
14209 eprintln!(
14210 "embryo-dbg: chunk prefill {} ids from {start} in {} submits (chunk {cmax})",
14211 n - start,
14212 (n - start).div_ceil(cmax)
14213 );
14214 }
14215 last
14216 }
14217
14218 fn ensure_embryo_graph(&mut self) -> Option<std::sync::Arc<crate::gpu::EmbryoGraphModel>> {
14219 if self.embryo_graph.is_none() && self.embryo_resident_eligible() {
14220 const UMAX: u32 = u32::MAX;
14221 const HEADER: usize = crate::gpu::EMBRYO_META_HEADER;
14222 const REC: usize = 64;
14223 struct Pack {
14224 data: Vec<f32>,
14225 }
14226 impl Pack {
14227 fn put(&mut self, x: &[f32]) -> u32 {
14228 if x.is_empty() {
14229 return u32::MAX;
14230 }
14231 let off = self.data.len();
14232 self.data.extend_from_slice(x);
14233 off as u32
14234 }
14235 }
14236 let vmf = self.vmf_cfg;
14240 let gdn = self.gdn_cfg;
14241 let mut pack = Pack { data: Vec::new() };
14242 let mut meta = vec![0u32; HEADER];
14243 meta[0] = self.hidden_size as u32;
14244 meta[1] = self.intermediate_size as u32;
14245 meta[2] = self.vocab_size as u32;
14246 meta[3] = self.num_layers as u32;
14247 meta[4] = vmf.map(|c| c.num_heads).unwrap_or(0) as u32;
14248 meta[5] = vmf.map(|c| c.nphase).unwrap_or(0) as u32;
14249 meta[6] = vmf.map(|c| c.value_head_dim).unwrap_or(0) as u32;
14250 if let Some(g) = gdn {
14251 meta[24] = g.num_v_heads as u32;
14252 meta[25] = g.num_k_heads as u32;
14253 meta[26] = g.key_head_dim as u32;
14254 meta[27] = g.value_head_dim as u32;
14255 meta[28] = g.conv_kernel as u32;
14256 meta[29] = g.conv_dim() as u32;
14257 }
14258 meta[7] = self.num_heads as u32;
14259 meta[8] = self.num_kv_heads as u32;
14260 meta[9] = self.head_dim as u32;
14261 meta[10] = self.kv_cache.max_seq_len as u32;
14262 let clusters = self.head_clusters.as_ref().unwrap();
14263 let cluster_count = clusters.len() / self.hidden_size;
14264 if clusters.len() % self.hidden_size != 0
14265 || cluster_count == 0
14266 || cluster_count > 1024
14267 || self.vocab_size % cluster_count != 0
14268 || self.weights.lm_head.rows() < self.vocab_size
14269 || self.weights.final_norm.len() != self.hidden_size
14270 {
14271 return None;
14272 }
14273 meta[11] = cluster_count as u32;
14274 meta[12] = (self.vocab_size / cluster_count) as u32;
14275 meta[13] = vmf.map(|c| c.state_len()).unwrap_or(0) as u32;
14276 meta[16] = self.rotary_dim as u32;
14277 meta[17] = matches!(self.norm_style, NormStyle::Gemma) as u32;
14278 meta[19] = (self.rms_eps as f32).to_bits();
14279 let max_conv = self
14280 .weights
14281 .layers
14282 .iter()
14283 .filter_map(|lw| match &lw.attn {
14284 AttnKind::Linear(w) => w.conv.as_ref().map(|v| v.len() / self.hidden_size),
14285 _ => None,
14286 })
14287 .max()
14288 .unwrap_or(1);
14289 let phase_stride = vmf
14295 .map(|c| c.state_len() + max_conv.saturating_sub(1) * self.hidden_size)
14296 .unwrap_or(0);
14297 let gdn_stride = gdn.map(|g| g.state_len()).unwrap_or(0);
14298 let state_stride = phase_stride.max(gdn_stride);
14299 let bounded = self.anchor_core.clone();
14303 let (anchor_window, anchor_sink) = bounded
14304 .as_ref()
14305 .map(|ac| (ac.window, ac.sink))
14306 .unwrap_or((0, 0));
14307 let kv_stride = if bounded.is_some() {
14308 2usize
14309 .saturating_mul(self.num_kv_heads)
14310 .saturating_mul(anchor_window)
14311 .saturating_mul(self.head_dim)
14312 } else {
14313 2usize
14314 .saturating_mul(self.num_kv_heads)
14315 .saturating_mul(self.kv_cache.max_seq_len)
14316 .saturating_mul(self.head_dim)
14317 };
14318 meta[14] = state_stride as u32;
14319 meta[15] = kv_stride as u32;
14320 meta[18] = anchor_window as u32;
14321 meta[20] = anchor_sink as u32;
14322 meta[21] = match &self.bounded_rope {
14323 Some(rope) => {
14324 let off = pack.put(&rope.cos);
14326 let _ = pack.put(&rope.sin);
14327 off
14328 }
14329 None => UMAX,
14330 };
14331 let mut full_seen = false;
14332 let mut bounded_seen = 0usize;
14333 let mut phase_seen = 0usize;
14337 let mut gdn_seen = 0usize;
14338 for (li, lw) in self.weights.layers.iter().enumerate() {
14339 let base = meta.len();
14340 meta.resize(base + REC, UMAX);
14341 meta[base] = match &lw.attn {
14342 AttnKind::Linear(w) if w.phase_delta => 1,
14343 AttnKind::Linear(_) => 0,
14344 AttnKind::Full { .. } => 2,
14345 AttnKind::Bounded(_) => 3,
14346 AttnKind::LinearGdn(_) => 4,
14347 _ => UMAX,
14348 };
14349 meta[base + 1] = pack.put(&lw.input_norm);
14350 meta[base + 2] = pack.put(&lw.post_norm);
14351 meta[base + 25] = match &lw.attn {
14352 AttnKind::Linear(_) => {
14353 let off = ((phase_seen + gdn_seen) * state_stride) as u32;
14354 phase_seen += 1;
14355 off
14356 }
14357 AttnKind::LinearGdn(_) => {
14358 let off = ((phase_seen + gdn_seen) * state_stride) as u32;
14359 gdn_seen += 1;
14360 off
14361 }
14362 _ => UMAX,
14363 };
14364 match &lw.attn {
14365 AttnKind::LinearGdn(w) => {
14366 meta[base + 56] = pack.put(&Self::embryo_qtensor_f32(&w.in_proj_qkv));
14369 meta[base + 57] = pack.put(&Self::embryo_qtensor_f32(&w.in_proj_z));
14370 meta[base + 58] = pack.put(&Self::embryo_qtensor_f32(&w.in_proj_a));
14371 meta[base + 59] = pack.put(&Self::embryo_qtensor_f32(&w.in_proj_b));
14372 meta[base + 60] = pack.put(&w.conv1d);
14373 meta[base + 61] = pack.put(&w.a_log);
14374 meta[base + 62] = pack.put(&w.dt_bias);
14375 meta[base + 63] = pack.put(&w.norm);
14376 meta[base + 29] = pack.put(&Self::embryo_qtensor_f32(&w.out_proj));
14377 meta[base + 24] = 0;
14378 }
14379 AttnKind::Linear(w) => {
14380 meta[base + 3] = pack.put(&Self::embryo_qtensor_f32(&w.thq));
14381 meta[base + 4] = pack.put(&Self::embryo_qtensor_f32(&w.thk));
14382 meta[base + 5] = pack.put(&Self::embryo_qtensor_f32(&w.v_proj));
14383 meta[base + 6] = pack.put(&Self::embryo_qtensor_f32(&w.out_proj));
14384 let decay: Vec<f32> = w.decay.iter().map(|&x| x as f32).collect();
14385 meta[base + 7] = pack.put(&decay);
14386 if let Some((kg, kb)) = &w.k_gate {
14387 meta[base + 8] = pack.put(&Self::embryo_qtensor_f32(kg));
14388 meta[base + 9] = pack.put(kb);
14389 }
14390 if let Some(conv) = &w.conv {
14391 meta[base + 10] = pack.put(conv);
14392 meta[base + 24] = (conv.len() / self.hidden_size) as u32;
14393 } else {
14394 meta[base + 24] = 0;
14395 }
14396 }
14397 AttnKind::Full { wq, wk, wv, wo, .. } => {
14398 full_seen = true;
14399 meta[base + 11] = pack.put(&Self::embryo_qtensor_f32(wq));
14400 meta[base + 12] = pack.put(&Self::embryo_qtensor_f32(wk));
14401 meta[base + 13] = pack.put(&Self::embryo_qtensor_f32(wv));
14402 meta[base + 14] = pack.put(&Self::embryo_qtensor_f32(wo));
14403 meta[base + 26] = (li * kv_stride) as u32;
14404 }
14405 AttnKind::Bounded(w) => {
14406 meta[base + 11] = pack.put(&Self::embryo_qtensor_f32(&w.wq));
14407 meta[base + 12] = pack.put(&Self::embryo_qtensor_f32(&w.wk));
14408 meta[base + 13] = pack.put(&Self::embryo_qtensor_f32(&w.wv));
14409 meta[base + 14] = pack.put(&Self::embryo_qtensor_f32(&w.wo));
14410 meta[base + 26] = (bounded_seen * kv_stride) as u32;
14412 meta[base + 27] = pack.put(&w.sink_k);
14413 meta[base + 28] = pack.put(&w.sink_v);
14414 bounded_seen += 1;
14415 }
14416 _ => return None,
14417 }
14418 match &lw.ffn {
14419 FfnKind::Dense(d) => {
14420 meta[base + 15] = 0;
14421 meta[base + 16] = 0;
14422 meta[base + 21] = pack.put(&Self::embryo_qtensor_f32(&d.gate_proj));
14423 meta[base + 22] = pack.put(&Self::embryo_qtensor_f32(&d.up_proj));
14424 meta[base + 23] = pack.put(&Self::embryo_qtensor_f32(&d.down_proj));
14425 }
14426 FfnKind::Moe(m) => {
14427 let r = m.resonance.as_ref().unwrap();
14428 let (shared, _) = m.shared.as_ref().unwrap();
14429 meta[base + 15] = 1;
14430 meta[base + 16] = m.experts.len() as u32;
14431 meta[base + 17] = pack.put(&r.mu);
14432 meta[base + 18] = pack.put(&r.u);
14433 meta[base + 19] = pack.put(&r.bias);
14434 meta[base + 20] = r.k as u32;
14435 let mut shell = r.effective_shell(m.experts.len());
14444 shell.push(f32::NEG_INFINITY);
14445 meta[base + 30] = pack.put(&shell);
14446 meta[base + 21] = pack.put(&Self::embryo_qtensor_f32(&shared.gate_proj));
14447 meta[base + 22] = pack.put(&Self::embryo_qtensor_f32(&shared.up_proj));
14448 meta[base + 23] = pack.put(&Self::embryo_qtensor_f32(&shared.down_proj));
14449 for (e, ex) in m.experts.iter().enumerate() {
14450 meta[base + 32 + e * 3] =
14451 pack.put(&Self::embryo_qtensor_f32(&ex.gate_proj));
14452 meta[base + 33 + e * 3] =
14453 pack.put(&Self::embryo_qtensor_f32(&ex.up_proj));
14454 meta[base + 34 + e * 3] =
14455 pack.put(&Self::embryo_qtensor_f32(&ex.down_proj));
14456 }
14457 }
14458 FfnKind::DenseMoe(_) => return None,
14459 }
14460 }
14461 if !full_seen && self.num_layers == 0 {
14462 return None;
14463 }
14464 let id = {
14465 static NEXT: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(1);
14466 NEXT.fetch_add(1, std::sync::atomic::Ordering::Relaxed)
14467 };
14468 let model = crate::gpu::EmbryoGraphModel {
14469 id,
14470 hidden: self.hidden_size,
14471 intermediate: self.intermediate_size,
14472 vocab: self.vocab_size,
14473 layers: self.num_layers,
14474 phase_heads: vmf.map(|c| c.num_heads).unwrap_or(0),
14475 nphase: vmf.map(|c| c.nphase).unwrap_or(0),
14476 phase_dv: vmf.map(|c| c.value_head_dim).unwrap_or(0),
14477 anchor_q_heads: self.num_heads,
14478 anchor_kv_heads: self.num_kv_heads,
14479 anchor_head_dim: self.head_dim,
14480 rotary_dim: self.rotary_dim,
14481 max_seq: self.kv_cache.max_seq_len,
14482 cluster_count,
14483 cluster_size: self.vocab_size / cluster_count,
14484 phase_state_len: vmf.map(|c| c.state_len()).unwrap_or(0),
14485 state_stride,
14486 kv_stride,
14487 norm_gemma: matches!(self.norm_style, NormStyle::Gemma),
14488 phase_mass: vmf.map(|c| c.phase_mass).unwrap_or(0.0),
14489 weights: pack.data,
14490 meta,
14491 lm_head: Self::embryo_qtensor_f32(&self.weights.lm_head),
14492 clusters: clusters.as_ref().clone(),
14493 final_norm: self.weights.final_norm.clone(),
14494 inv_freq: self.inv_freq.as_ref().clone(),
14495 bounded: bounded.is_some(),
14496 kv_layers: if bounded.is_some() {
14497 bounded_seen
14498 } else {
14499 self.num_layers
14500 },
14501 state_layers: phase_seen + gdn_seen,
14502 anchor_window,
14503 anchor_sink,
14504 phase_layers: phase_seen,
14505 gdn_layers: gdn_seen,
14506 gdn_heads: gdn.map(|g| g.num_v_heads).unwrap_or(0),
14507 gdn_k_heads: gdn.map(|g| g.num_k_heads).unwrap_or(0),
14508 gdn_dk: gdn.map(|g| g.key_head_dim).unwrap_or(0),
14509 gdn_dv: gdn.map(|g| g.value_head_dim).unwrap_or(0),
14510 gdn_kk: gdn.map(|g| g.conv_kernel).unwrap_or(0),
14511 };
14512 self.embryo_graph = Some(std::sync::Arc::new(model));
14513 }
14514 self.embryo_graph.clone()
14515 }
14516
14517 fn forward_layers_span(
14518 &mut self,
14519 hidden: &[f32],
14520 position: usize,
14521 task_mask: Option<&TaskMask>,
14522 from: usize,
14523 upto: Option<usize>,
14524 ) -> Vec<f32> {
14525 debug_assert!(
14526 from == 0
14527 || (self.dsv4.is_none()
14528 && self.dsv41.is_none()
14529 && self.qwen4_exp.is_none()
14530 && self.g3n.is_none())
14531 );
14532 #[cfg(target_os = "macos")]
14538 if !crate::gpu_metal::wait_replay() {
14539 self.fail_metal_graph("the pending async replay failed before a plain forward");
14540 return vec![0.0; self.hidden_size];
14541 }
14542 if let Some(b) = &mut self.qwen4_exp {
14543 let _ = (task_mask, upto);
14544 let token_id = hidden.first().copied().unwrap_or(0.0) as u32;
14545 let mut logits = Vec::new();
14546 crate::qwen4_exp::forward_token(
14547 &b.0,
14548 &b.1,
14549 &b.2,
14550 &mut b.3,
14551 token_id,
14552 position,
14553 &self.inv_freq,
14554 self.pool.as_deref(),
14555 &mut logits,
14556 true,
14557 );
14558 self.graph_logits = Some(logits);
14559 return vec![0.0; self.hidden_size];
14560 }
14561 if let Some(b) = &mut self.dsv4 {
14567 let _ = (task_mask, upto);
14568 let token_id = hidden.first().copied().unwrap_or(0.0) as u32;
14569 let (g, layers, cfg, st) = (&b.0, &b.1, b.2, &mut b.3);
14570 st.pos = position;
14571 let mut logits = Vec::new();
14572 crate::dsv4::forward_token(
14573 g,
14574 layers,
14575 &cfg,
14576 st,
14577 token_id,
14578 &self.inv_freq,
14579 self.pool.as_deref(),
14580 &mut logits,
14581 );
14582 self.graph_logits = Some(logits);
14583 self.dspark_probe(position, token_id);
14584 return vec![0.0; self.hidden_size];
14587 }
14588 if let Some(b) = &mut self.dsv41 {
14590 let _ = (task_mask, upto);
14591 let token_id = hidden.first().copied().unwrap_or(0.0) as u32;
14592 let mut logits = Vec::new();
14593 crate::dsv41::forward_token(
14594 &b.0,
14595 &b.1,
14596 &b.2,
14597 &mut b.3,
14598 token_id,
14599 position,
14600 self.pool.as_deref(),
14601 &mut logits,
14602 );
14603 self.graph_logits = Some(logits);
14604 return vec![0.0; self.hidden_size];
14605 }
14606 if let Some(b) = &self.g3n {
14609 let _ = (task_mask, upto);
14610 return crate::g3n::g3n_forward(
14611 &b.0,
14612 &b.1,
14613 hidden,
14614 position,
14615 &mut self.kv_cache.layers,
14616 self.num_heads,
14617 self.num_kv_heads,
14618 self.head_dim,
14619 self.pool.as_deref(),
14620 );
14621 }
14622 if from == 0
14629 && upto.is_none()
14630 && task_mask.is_none()
14631 && self.anchor_core.is_some()
14632 && std::env::var("CMF_EMBRYO_DBG").as_deref() == Ok("1")
14633 {
14634 static ONCE: std::sync::Once = std::sync::Once::new();
14635 ONCE.call_once(|| {
14636 eprintln!(
14637 "embryo-dbg: graph_on decode={} prefill={} resident_env={:?} enabled_here={} \
14638 unsupported={} eligible={}",
14639 crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Decode),
14640 crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Prefill),
14641 std::env::var("CMF_EMBRYO_RESIDENT").ok(),
14642 crate::gpu::enabled_here(),
14643 self.graph_refused(),
14644 self.embryo_resident_eligible(),
14645 );
14646 });
14647 }
14648 if from == 0
14649 && upto.is_none()
14650 && task_mask.is_none()
14651 && crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Decode)
14655 && crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Prefill)
14656 && matches!(
14660 std::env::var("CMF_EMBRYO_RESIDENT").as_deref(),
14661 Ok("1") | Ok("parallel")
14662 )
14663 && crate::gpu::enabled_here()
14664 && !self.graph_refused()
14665 && (position == 0 || self.device_sequence_position().is_some())
14671 && self.embryo_resident_eligible()
14672 && let Some(model) = self.ensure_embryo_graph()
14673 {
14674 let mut lg = Vec::new();
14675 if crate::gpu::forward_embryo_graph(&model, self.graph_kv_id, hidden, position, &mut lg)
14676 {
14677 self.graph_logits = Some(lg);
14678 return vec![0.0; self.hidden_size];
14679 }
14680 if std::env::var("CMF_EMBRYO_DBG").as_deref() == Ok("1") {
14681 eprintln!("embryo-dbg: forward_embryo_graph refused at position {position}");
14682 }
14683 self.mark_graph_refused();
14689 if position != 0 {
14690 self.kv_cache.clear();
14694 self.clear_history();
14695 crate::gpu::graph_kv_reset(self.graph_kv_id);
14696 panic!(
14697 "resident Embryo graph refused at position {position}; refusing an in-flight host fallback"
14698 );
14699 }
14700 }
14701 let mut h = hidden.to_vec();
14702 self.mimo_moe_prepare();
14705 let _mimo_q8 = self.mimo_moe.is_on()
14706 .then(crate::qtensor::enter_full_gpu_q8_scope);
14707 let (nh, _nkv, _hd, hs, _rd, eps) = (
14710 self.num_heads,
14711 self.num_kv_heads,
14712 self.head_dim,
14713 self.hidden_size,
14714 self.rotary_dim,
14715 self.rms_eps,
14716 );
14717 let pool = self.pool.clone();
14718 let graph_env = std::env::var("CMF_GPU_WGPU_GRAPH").ok();
14730 let graph_on = match graph_env.as_deref() {
14731 Some("0") => false,
14732 Some("prefill") => false, Some(_) => true,
14734 None => crate::gpu::wgpu_graph_default(),
14740 };
14741 let graph_trusted =
14742 graph_env.is_some() || crate::gpu::wgpu_graph_default() || self.gdn_cfg.is_some();
14743 let race_eligible = graph_on
14744 && upto.is_none()
14745 && task_mask.is_none()
14746 && from == 0
14747 && !self.graph_refused();
14748 let mut tail_start = 0usize;
14749 if race_eligible && crate::gpu::graph_race_use_graph(graph_trusted) {
14750 let t_graph = std::time::Instant::now();
14751 let mut lg = Vec::new();
14752 let mut gl = 0usize;
14753 let built = self.try_token_graph_wgpu(hidden, position, &mut lg, &mut gl);
14754 let declined = built.is_none();
14755 let built = match built {
14756 Some(Ok(hh)) => Some(hh),
14757 Some(Err(())) => {
14758 self.clear_sequence_state();
14762 self.graph_failed
14763 .store(true, std::sync::atomic::Ordering::Relaxed);
14764 self.cancel
14765 .store(true, std::sync::atomic::Ordering::Relaxed);
14766 tracing::error!("token graph failed after admission; sequence state cleared");
14767 return vec![0.0; self.hidden_size];
14768 }
14769 None => None,
14770 };
14771 if declined && !self.o1_active() && self.attn_softcap == 0.0 {
14776 self.mark_graph_refused();
14777 }
14778 graph_note(built.is_some(), gl, self.num_layers);
14779 if let Some(hh) = built {
14780 let dur = t_graph.elapsed();
14781 if std::env::var("CMF_GRAPH_PROF").is_ok() {
14782 eprintln!("graph-call: {:.2} ms total", dur.as_secs_f64() * 1000.0);
14783 }
14784 if gl > 0 && gl < self.num_layers {
14785 h = hh;
14791 tail_start = gl;
14792 } else if graph_trusted || !crate::gpu::graph_race_first_token_hopeless(dur) {
14793 if !graph_trusted {
14794 crate::gpu::graph_race_record(true, dur);
14795 }
14796 if !lg.is_empty() {
14797 lg.resize(self.vocab_size, 0.0);
14800 if let Some(c) = self.final_softcap {
14801 for l in lg.iter_mut() {
14802 *l = c * (*l / c).tanh();
14803 }
14804 }
14805 self.graph_logits = Some(lg);
14806 }
14807 return hh;
14808 }
14809 }
14815 }
14816 let span = from > 0 || upto.is_some();
14840 if span && graph_on && task_mask.is_none() && graph_trusted {
14841 let upto_excl = upto.map_or(self.num_layers, |u| u + 1);
14842 let mut lg = Vec::new();
14843 let mut gl = 0usize;
14844 let span_res =
14845 self.try_token_graph_wgpu_span(hidden, position, &mut lg, from, upto_excl, &mut gl);
14846 let span_res = match span_res {
14847 Some(Ok(hh)) => Some(hh),
14848 Some(Err(())) => {
14849 self.clear_sequence_state();
14850 self.graph_failed
14851 .store(true, std::sync::atomic::Ordering::Relaxed);
14852 self.cancel
14853 .store(true, std::sync::atomic::Ordering::Relaxed);
14854 tracing::error!(
14855 "span token graph failed after admission; sequence state cleared"
14856 );
14857 return vec![0.0; self.hidden_size];
14858 }
14859 None => None,
14860 };
14861 graph_note(span_res.is_some(), gl, upto_excl - from);
14862 if std::env::var("CMF_GPU_DEBUG").is_ok() {
14863 static SEEN: std::sync::atomic::AtomicU32 = std::sync::atomic::AtomicU32::new(0);
14867 if SEEN.fetch_add(1, std::sync::atomic::Ordering::Relaxed) < 4 {
14868 eprintln!(
14869 "span graph: covered {gl} of {} layers [{from}..{upto_excl}) res={}",
14870 upto_excl - from,
14871 span_res.is_some()
14872 );
14873 }
14874 }
14875 if let Some(hh) = span_res {
14876 if gl == upto_excl - from {
14877 if !lg.is_empty() {
14878 lg.resize(self.vocab_size, 0.0);
14879 if let Some(c) = self.final_softcap {
14880 for l in lg.iter_mut() {
14881 *l = c * (*l / c).tanh();
14882 }
14883 }
14884 self.graph_logits = Some(lg);
14885 }
14886 crate::gpu::set_layer(-1);
14887 return hh;
14888 }
14889 h = hh;
14891 tail_start = from + gl;
14892 }
14893 }
14894 let t_race_cpu = (race_eligible && !graph_trusted).then(std::time::Instant::now);
14899
14900 let host_tail = tail_start > from;
14910 let _host_tail = (host_tail && !self.mimo_moe.is_on()).then(crate::gpu::enter_cpu_scope);
14911 let automatic_gpu_prefix = self.automatic_gpu_prefix();
14912
14913 let _prof_layers = crate::cpuprof::time(crate::cpuprof::Slot::Layers);
14914 #[cfg(target_os = "macos")]
14915 let mut gpu_skip_until = 0usize;
14916 for li in tail_start.max(from)..self.num_layers {
14917 let _capacity_tail = automatic_gpu_prefix
14918 .filter(|&prefix| li >= prefix && !self.mimo_moe.is_dynamic(li, host_tail))
14919 .map(|_| crate::gpu::enter_cpu_scope());
14920 crate::gpu::set_layer(li as i64); if let Some(u) = upto {
14922 if li > u {
14923 break;
14924 }
14925 }
14926 if let Some(mask) = task_mask {
14927 if !mask.layer_alive(li) {
14928 continue; }
14930 }
14931 #[cfg(target_os = "macos")]
14935 {
14936 if li < gpu_skip_until {
14937 continue;
14938 }
14939 if task_mask.is_none() {
14940 let end = self.q1_graph_gpu(li, upto, position, &mut h);
14941 if self
14942 .graph_failed
14943 .load(std::sync::atomic::Ordering::Relaxed)
14944 {
14945 return vec![0.0; self.hidden_size];
14949 }
14950 if end > li {
14951 gpu_skip_until = end;
14952 if self.is_loop_end(end - 1) && end < self.num_layers {
14955 h = inference::rms_norm(
14956 &h,
14957 &self.weights.final_norm,
14958 self.rms_eps,
14959 self.norm_style,
14960 );
14961 }
14962 continue;
14963 }
14964 }
14965 }
14966
14967 if task_mask.is_none() {
14968 match self.mimo_graph_layer_rows(li, &mut h, &[position]) {
14969 crate::gpu::BatchGraphOutcome::Completed => continue,
14970 crate::gpu::BatchGraphOutcome::Failed => return vec![0.0; self.hidden_size],
14971 crate::gpu::BatchGraphOutcome::Declined => {},
14972 }
14973 }
14974 #[cfg(feature = "gpu")]
14975 self.pull_lagging_host_kv(li, li + 1, position);
14976 let lw = &self.weights.layers[self.phys_layer(li)];
14977 if let Ok(tp) = std::env::var("CMF_TRACE_POS") {
14978 if tp.parse::<usize>().ok() == Some(position) {
14979 let n: f32 = h.iter().map(|x| x * x).sum::<f32>().sqrt();
14980 eprintln!(
14981 "TRACE pos {position} layer {li}: |h| = {n:.6} h0 {:.6} h1 {:.6}",
14982 h[0], h[1]
14983 );
14984 }
14985 }
14986 let prof = crate::cpuprof::time(crate::cpuprof::Slot::Norms);
14989 inference::rms_norm_into(
14990 &h,
14991 &lw.input_norm,
14992 self.rms_eps,
14993 self.norm_style,
14994 &mut self.ws.n1,
14995 );
14996 drop(prof);
14997
14998 let attn_out = match &lw.attn {
14999 AttnKind::Mla(w) => {
15000 let inv_freq_l = self.layer_inv_freq(li);
15001 let rs = self.layer_rope_scale(li);
15002 let eps = self.rms_eps;
15003 let pool = self.pool.clone();
15004 mla_attention(
15005 w,
15006 &self.ws.n1,
15007 &mut self.kv_cache.layers[li],
15008 position,
15009 &inv_freq_l,
15010 rs,
15011 eps,
15012 pool.as_deref(),
15013 )
15014 }
15015 AttnKind::Linear(w) => {
15016 let cfg = self.vmf_cfg.expect("linear layer without vmf_cfg");
15017 vmf_phase_forward(
15018 &self.ws.n1,
15019 w,
15020 &cfg,
15021 &mut self.kv_cache.layers[li].linear_state,
15022 self.pool.as_deref(),
15023 )
15024 }
15025 AttnKind::Kda(w) => {
15026 let cfg = self.kda_cfg.expect("kda layer without kda_cfg");
15027 crate::linear_core::kda_forward(
15028 &self.ws.n1,
15029 w,
15030 &cfg,
15031 &mut self.kv_cache.layers[li].linear_state,
15032 self.pool.as_deref(),
15033 )
15034 }
15035 AttnKind::LinearGdn(w) => {
15036 let cfg = self.gdn_cfg.expect("gdn layer without gdn_cfg");
15037 gdn_forward(
15038 &self.ws.n1,
15039 w,
15040 &cfg,
15041 &mut self.kv_cache.layers[li].linear_state,
15042 self.pool.as_deref(),
15043 )
15044 }
15045 AttnKind::ShortConv(w) => {
15046 let cfg = self
15047 .short_conv_cfg
15048 .expect("short-conv layer without short_conv_cfg");
15049 short_conv_forward(
15050 &self.ws.n1,
15051 w,
15052 &cfg,
15053 &mut self.kv_cache.layers[li].linear_state,
15054 self.pool.as_deref(),
15055 )
15056 }
15057 AttnKind::Bounded(w) => {
15058 let rope = self
15061 .bounded_rope
15062 .clone()
15063 .expect("bounded layer without an installed rotation table");
15064 let cfg = crate::bounded::BoundedAttnCfg {
15065 num_heads: self.num_heads,
15066 num_kv_heads: self.num_kv_heads,
15067 head_dim: self.head_dim,
15068 hidden_size: hs,
15069 scale: self.attn_scale,
15070 rope: &rope,
15071 pool: pool.as_deref(),
15072 };
15073 crate::bounded::bounded_attention(
15074 &self.ws.n1,
15075 w,
15076 &mut self.kv_cache.layers[li],
15077 &cfg,
15078 )
15079 }
15080 AttnKind::Full {
15081 wq,
15082 wk,
15083 wv,
15084 wo,
15085 q_norm,
15086 k_norm,
15087 output_gate,
15088 softplus_gate,
15089 bias,
15090 } if self.kv_cache.layers[li].o1_sealed() => {
15091 let inv_freq_l = self.layer_inv_freq(li);
15094 let (nkv_l, hd_l, rd_l) = self.layer_geom(li);
15095 let cfg = QwenAttnCfg {
15096 num_heads: self.layer_num_heads(li),
15097 num_kv_heads: nkv_l,
15098 head_dim: hd_l,
15099 hidden_size: hs,
15100 position,
15101 inv_freq: &inv_freq_l,
15102 rotary_dim: rd_l,
15103 scale: self.attn_scale,
15104 softcap: self.attn_softcap,
15105 window: None,
15106 v_norm: self.attn_v_norm,
15107 qk_norm_after_rope: self.qk_norm_after_rope,
15108 gate_sigmoid: self.proj_gate_sigmoid,
15109 q_norm: q_norm.as_deref(),
15110 k_norm: k_norm.as_deref(),
15111 output_gate: *output_gate,
15112 softplus_gate: softplus_gate
15113 .as_ref()
15114 .map(|(gate, per_head)| (gate, *per_head)),
15115 rope_scale: self.layer_rope_scale(li),
15116 bias: bias
15117 .as_ref()
15118 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
15119 rms_eps: eps,
15120 norm_style: self.norm_style,
15121 pool: pool.as_deref(),
15122 v_head_dim: self.layer_v_dim(li),
15123 };
15124 attention::qwen_attention_nystrom(
15125 &self.ws.n1,
15126 wq,
15127 wk,
15128 wv,
15129 wo,
15130 &mut self.kv_cache.layers[li],
15131 &cfg,
15132 )
15133 }
15134 AttnKind::Full {
15135 wq,
15136 wk,
15137 wv,
15138 wo,
15139 q_norm,
15140 k_norm,
15141 output_gate,
15142 softplus_gate,
15143 bias,
15144 } => 'attn: {
15145 let dropin_reason =
15150 graph_on.then(|| self.graph_attn_decline_reason()).flatten();
15151 if let Some(reason) = dropin_reason {
15152 self.note_graph_decline("wgpu attn dropin", reason);
15153 }
15154 if graph_on
15155 && dropin_reason.is_none()
15156 && !*output_gate
15157 && softplus_gate.is_none()
15158 && self.attention_heads_per_layer.is_none()
15159 && bias.is_none()
15160 && task_mask.is_none()
15161 {
15162 let inv_freq_l = self.layer_inv_freq(li);
15163 let (nkv_l, hd_l, rd_l) = self.layer_geom(li);
15164 let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
15165 if let (Some((gm, qi)), Some((_, ki)), Some((_, vi)), Some((_, oi))) = (
15166 wq.mapped_q1(),
15167 wk.mapped_q1(),
15168 wv.mapped_q1(),
15169 wo.mapped_q1(),
15170 ) {
15171 let gm = gm.clone();
15172 let mut out = vec![0f32; hs];
15173 let cache = &self.kv_cache.layers[li];
15174 if crate::gpu::attn_dropin(
15175 &gm,
15176 self.graph_kv_id,
15177 li,
15178 &self.ws.n1,
15179 qi,
15180 ki,
15181 vi,
15182 oi,
15183 q_norm.as_deref(),
15184 k_norm.as_deref(),
15185 self.qk_norm_after_rope,
15186 &inv_freq_l,
15187 nh,
15188 nkv_l,
15189 hd_l,
15190 rd_l,
15191 hs,
15192 position,
15193 self.kv_cache.max_seq_len,
15194 gemma,
15195 eps as f32,
15196 cache.k_heads(),
15197 cache.v_heads(),
15198 &mut out,
15199 ) {
15200 break 'attn out;
15201 }
15202 }
15203 }
15204 let masked = task_mask
15205 .map(|m| m.head_flags(li, self.num_heads).iter().any(|&a| !a))
15206 .unwrap_or(false);
15207 let f32_view = (wq.as_f32(), wk.as_f32(), wv.as_f32(), wo.as_f32());
15208 let plain = self.layer_attn_plain(li);
15211 match (masked, f32_view) {
15212 (true, (Some(q), Some(k), Some(v), Some(o))) if plain => {
15215 let active_heads = task_mask.unwrap().head_flags(li, self.num_heads);
15216 attention::multi_head_attention(
15217 &self.ws.n1,
15218 q,
15219 k,
15220 v,
15221 o,
15222 &mut self.kv_cache.layers[li],
15223 self.num_heads,
15224 self.num_kv_heads,
15225 self.head_dim,
15226 self.hidden_size,
15227 position,
15228 &active_heads,
15229 &self.inv_freq,
15230 )
15231 }
15232 (masked, _) => {
15233 if masked {
15234 tracing::warn!(
15235 "layer {li}: head mask on quantized weights or on a \
15236 window/sink/per-layer-geometry layer not supported \
15237 yet — executing dense"
15238 );
15239 }
15240 let inv_freq_l = self.layer_inv_freq(li);
15241 let (nkv_l, hd_l, rd_l) = self.layer_geom(li);
15242 let cfg = QwenAttnCfg {
15243 num_heads: self.layer_num_heads(li),
15244 num_kv_heads: nkv_l,
15245 head_dim: hd_l,
15246 hidden_size: hs,
15247 position,
15248 inv_freq: &inv_freq_l,
15249 rotary_dim: rd_l,
15250 scale: self.attn_scale,
15251 softcap: self.attn_softcap,
15252 window: self.layer_window(li),
15253 v_norm: self.attn_v_norm,
15254 qk_norm_after_rope: self.qk_norm_after_rope,
15255 gate_sigmoid: self.proj_gate_sigmoid,
15256 q_norm: q_norm.as_deref(),
15257 k_norm: k_norm.as_deref(),
15258 output_gate: *output_gate,
15259 softplus_gate: softplus_gate
15260 .as_ref()
15261 .map(|(gate, per_head)| (gate, *per_head)),
15262 rope_scale: self.layer_rope_scale(li),
15263 bias: bias
15264 .as_ref()
15265 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
15266 rms_eps: eps,
15267 norm_style: self.norm_style,
15268 pool: pool.as_deref(),
15269 v_head_dim: self.layer_v_dim(li),
15270 };
15271 attention::qwen_attention(
15272 &self.ws.n1,
15273 wq,
15274 wk,
15275 wv,
15276 wo,
15277 &mut self.kv_cache.layers[li],
15278 &cfg,
15279 )
15280 }
15281 }
15282 }
15283 };
15284 let attn_out = match &self.weights.layers[self.phys_layer(li)].attn_out_norm {
15287 Some(w) => inference::rms_norm(&attn_out, w, self.rms_eps, self.norm_style),
15288 None => attn_out,
15289 };
15290 let lw = &self.weights.layers[self.phys_layer(li)];
15291 let prof = crate::cpuprof::time(crate::cpuprof::Slot::Norms);
15292 inference::add_rmsnorm_fused_into(
15293 &mut h,
15294 &attn_out,
15295 &lw.post_norm,
15296 self.rms_eps,
15297 self.norm_style,
15298 &mut self.ws.p1,
15299 );
15300 drop(prof);
15301 let mut attn_out = attn_out;
15302 attention::recycle_buf(&mut attn_out);
15303 let post_normed = &self.ws.p1;
15304
15305 let ffn_masked = task_mask
15306 .map(|m| m.ffn_active_count(li) < self.intermediate_size)
15307 .unwrap_or(false);
15308 let ffn_out = match (ffn_masked, &lw.ffn) {
15320 (_, FfnKind::Dense(d)) if !d.segs.is_empty() => {
15324 let row = task_mask
15325 .and_then(|tm| tm.ffn_masks.get(li))
15326 .map(|v| v.as_slice());
15327 tube_ffn(d, post_normed, 1, self.pool.as_deref(), row)
15328 }
15329 (true, FfnKind::Dense(d)) => {
15330 let tm = task_mask.unwrap();
15331 let alive = tm.ffn_active_count(li);
15332 let deep = alive * 2 <= self.intermediate_size;
15333 if deep && d.down_proj.sparse_col_ok() && !d.gate_proj.has_prism_contract() {
15334 let active = tm.ffn_active_indices(li);
15335 sparse_ffn_quant(
15336 d,
15337 post_normed,
15338 &active,
15339 self.hidden_size,
15340 self.pool.as_deref(),
15341 )
15342 } else if deep
15343 && let (Some(g), Some(u), Some(dn)) = (
15344 d.gate_proj.as_f32(),
15345 d.up_proj.as_f32(),
15346 d.down_proj.as_f32(),
15347 )
15348 {
15349 let active = tm.ffn_active_indices(li);
15350 inference::sparse_ffn_forward(
15351 post_normed,
15352 g,
15353 u,
15354 dn,
15355 self.hidden_size,
15356 self.intermediate_size,
15357 &active,
15358 self.pool.as_deref(),
15359 )
15360 } else {
15361 let row = tm.ffn_masks.get(li).map(|v| v.as_slice());
15362 dense_ffn_batch(d, post_normed, 1, self.pool.as_deref(), row)
15363 }
15364 }
15365 (true, FfnKind::Moe(m)) => {
15366 let allowed = task_mask.and_then(|tm| tm.expert_flags(li, m.experts.len()));
15370 ffn_forward(
15371 &lw.ffn,
15372 post_normed,
15373 self.pool.as_deref(),
15374 allowed.as_deref(),
15375 )
15376 }
15377 (true, FfnKind::DenseMoe(dm)) => dense_moe_ffn(
15378 dm,
15379 post_normed,
15380 &h,
15381 self.rms_eps,
15382 self.norm_style,
15383 self.pool.as_deref(),
15384 ),
15385 (false, _) => match &lw.ffn {
15386 FfnKind::DenseMoe(dm) => dense_moe_ffn(
15387 dm,
15388 post_normed,
15389 &h,
15390 self.rms_eps,
15391 self.norm_style,
15392 self.pool.as_deref(),
15393 ),
15394 FfnKind::Moe(m)
15395 if task_mask.is_none() && self.mimo_moe.is_dynamic(li, host_tail) =>
15396 {
15397 moe_ffn_banked(&mut self.mimo_moe, li, m, post_normed, self.pool.as_deref())
15398 }
15399 _ => {
15400 let allowed = match (&lw.ffn, task_mask) {
15401 (FfnKind::Moe(m), Some(tm)) => tm.expert_flags(li, m.experts.len()),
15402 _ => None,
15403 };
15404 ffn_forward(
15405 &lw.ffn,
15406 post_normed,
15407 self.pool.as_deref(),
15408 allowed.as_deref(),
15409 )
15410 }
15411 },
15412 };
15413 let ffn_out = match &self.weights.layers[self.phys_layer(li)].ffn_out_norm {
15414 Some(w) => inference::rms_norm(&ffn_out, w, self.rms_eps, self.norm_style),
15415 None => ffn_out,
15416 };
15417 for (i, &f) in ffn_out.iter().enumerate() {
15418 h[i] += f;
15419 }
15420 let mut ffn_out = ffn_out;
15421 attention::recycle_buf(&mut ffn_out);
15422
15423 if let Some(sc) = self.weights.layers[self.phys_layer(li)].layer_scale {
15425 for v in h.iter_mut() {
15426 *v *= sc;
15427 }
15428 }
15429 if self.layer_dump.is_some() {
15431 self.dump_layer_row(position, li, &h);
15432 }
15433
15434 if self.is_loop_end(li) && li + 1 < self.num_layers {
15437 h = inference::rms_norm(
15438 &h,
15439 &self.weights.final_norm,
15440 self.rms_eps,
15441 self.norm_style,
15442 );
15443 }
15444
15445 if self.dyn_phi_layer == Some(li) {
15449 self.update_dyn_phi(&h);
15450 }
15451 }
15452 crate::gpu::set_layer(-1); if let Some(t) = t_race_cpu {
15454 crate::gpu::graph_race_record(false, t.elapsed());
15455 }
15456
15457 h
15458 }
15459
15460 fn update_dyn_phi(&mut self, h: &[f32]) {
15463 const A: f32 = 0.2;
15464 if self.dyn_phi_ema.len() != h.len() {
15465 self.dyn_phi_ema = vec![0.0; h.len()];
15466 self.dyn_phi_seen = 0;
15467 }
15468 if self.dyn_phi_seen == 0 {
15469 self.dyn_phi_ema.copy_from_slice(h);
15470 } else {
15471 for (e, &v) in self.dyn_phi_ema.iter_mut().zip(h) {
15472 *e = (1.0 - A) * *e + A * v;
15473 }
15474 }
15475 self.dyn_phi_seen += 1;
15476 }
15477
15478 pub fn dyn_phi(&self) -> &[f32] {
15480 &self.dyn_phi_ema
15481 }
15482
15483 pub fn set_dyn_phi_layer(&mut self, layer: Option<usize>) {
15485 self.dyn_phi_layer = layer;
15486 self.dyn_phi_ema.clear();
15487 self.dyn_phi_seen = 0;
15488 }
15489
15490 pub fn dynamic_skills(&self) -> Vec<(usize, String, usize)> {
15492 let Some(model) = &self.model else {
15493 return Vec::new();
15494 };
15495 model
15496 .header
15497 .skills
15498 .iter()
15499 .enumerate()
15500 .filter_map(|(i, sk)| {
15501 if sk.is_v2() {
15505 return None;
15506 }
15507 let ok = matches!(self.dyn_skill_layers.get(i), Some(Some(_)));
15508 let sel = sk.selection.as_ref()?;
15509 (ok).then(|| (i, sk.id.clone(), sel.phi_layer))
15510 })
15511 .collect()
15512 }
15513
15514 pub fn active_skill(&self) -> Option<usize> {
15516 self.dyn_active
15517 }
15518
15519 pub fn enable_dynamic_routing(&mut self) -> usize {
15524 use crate::swarm::{DynRouter, RoutableSkill};
15525 let Some(model) = self.model.clone() else {
15526 return 0;
15527 };
15528 if let Some(r) = &model.header.router {
15535 tracing::warn!(
15536 "dynamic routing disabled: this file declares router policy '{}' with \
15537 granularity \"{}\" — the request-level decision applies instead",
15538 r.policy,
15539 r.granularity
15540 );
15541 return 0;
15542 }
15543 if model.required_features & cortiq_core::format::features::SKILLS_V2 != 0
15549 || model.header.skills.iter().any(|s| s.is_v2())
15550 {
15551 tracing::warn!(
15552 "dynamic routing disabled: this file carries format-v2 skill records \
15553 (SKILLS_V2) — they route per request through a router policy only"
15554 );
15555 return 0;
15556 }
15557 if self.dyn_blend_loaded {
15560 tracing::warn!("dynamic routing unavailable on a blend-loaded pipeline");
15561 return 0;
15562 }
15563 if let Some(a) = self.dyn_active {
15567 if !matches!(self.dyn_skill_layers.get(a), Some(Some(_))) {
15568 tracing::warn!("loaded skill is not FFN-eligible — dynamic routing unavailable");
15569 return 0;
15570 }
15571 }
15572 let hidden = self.hidden_size;
15573 let mut skills = Vec::new();
15574 for (idx, id, _phi) in self.dynamic_skills() {
15575 if let Some(sel) = model.header.skills[idx].selection.as_ref() {
15576 if let Some(rs) = RoutableSkill::from_descriptor(idx, id, sel, hidden) {
15577 skills.push(rs);
15578 }
15579 }
15580 }
15581 if skills.is_empty() {
15582 return 0;
15583 }
15584 let phi = skills[0].phi_layer;
15586 if skills.iter().any(|s| s.phi_layer != phi) {
15587 tracing::warn!("routable skills disagree on phi_layer; using {phi}");
15588 }
15589 let n = skills.len();
15590 self.set_dyn_phi_layer(Some(phi));
15591 self.dyn_router = Some(DynRouter::new(skills));
15592 n
15593 }
15594
15595 pub fn route_switches(&self) -> Vec<(usize, Option<String>, Option<String>)> {
15597 self.dyn_router
15598 .as_ref()
15599 .map(|r| r.switches.clone())
15600 .unwrap_or_default()
15601 }
15602
15603 fn lm_head_forward(&self, hidden: &[f32]) -> Vec<f32> {
15606 let _mimo_q8 = self.mimo_moe.is_on()
15607 .then(crate::qtensor::enter_full_gpu_q8_scope);
15608 let rows = self.weights.lm_head.rows();
15609 let mut logits = attention::take_buf(rows.min(self.vocab_size));
15610 let served = self.mimo_moe.is_on() && crate::gpu::mimo_q8_short_enabled()
15614 && rows == self.vocab_size && !self.weights.lm_head.has_prism_contract()
15615 && self.weights.lm_head.graph_weight().is_some_and(|(model, idx, kind, _)| {
15616 kind == 7 && crate::gpu::q82_short_rows(model, idx, hidden, 1,
15617 rows, self.hidden_size, &mut logits)
15618 });
15619 if !served {
15620 self.weights.lm_head.matvec(hidden, &mut logits, self.pool.as_deref());
15621 }
15622 logits.resize(self.vocab_size, 0.0);
15623 if let Some(m) = self.logit_multiplier {
15624 for l in logits.iter_mut() {
15625 *l *= m;
15626 }
15627 }
15628 if let Some(c) = self.final_softcap {
15629 for l in logits.iter_mut() {
15630 *l = c * (*l / c).tanh();
15631 }
15632 }
15633 if let Some(cm) = self.head_clusters.as_ref() {
15634 self.hierarchical_head_logprobs(hidden, cm, &mut logits);
15635 }
15636 logits
15637 }
15638
15639 fn hierarchical_head_logprobs(&self, hidden: &[f32], cm: &[f32], logits: &mut [f32]) {
15642 let h = hidden.len();
15643 let ncl = cm.len() / h.max(1);
15644 if ncl == 0 || logits.len() % ncl != 0 {
15645 return;
15646 }
15647 let cs = logits.len() / ncl;
15648 let mut lc = vec![0.0f32; ncl];
15650 for c in 0..ncl {
15651 let row = &cm[c * h..(c + 1) * h];
15652 let mut s = 0.0f32;
15653 for j in 0..h {
15654 s += row[j] * hidden[j];
15655 }
15656 lc[c] = s;
15657 }
15658 let mx = lc.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
15659 let lse: f32 = mx + lc.iter().map(|v| (v - mx).exp()).sum::<f32>().ln();
15660 for c in 0..ncl {
15661 let blk = &mut logits[c * cs..(c + 1) * cs];
15662 let bm = blk.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
15663 let bl: f32 = bm + blk.iter().map(|v| (v - bm).exp()).sum::<f32>().ln();
15664 let add = lc[c] - lse - bl;
15665 for v in blk.iter_mut() {
15666 *v += add;
15667 }
15668 }
15669 }
15670
15671 pub fn prefill_next_logits(&mut self, ids: &[u32], task_mask: Option<&TaskMask>) -> Vec<f32> {
15676 #[cfg(target_os = "macos")]
15677 crate::gpu_metal::set_io_namespace(self.graph_kv_id);
15678 self.clear_sequence_state();
15679 crate::gpu::graph_race_begin_generation();
15683 if task_mask.is_none() {
15684 self.o1_begin();
15685 }
15686 let mut hidden = vec![0.0f32; self.hidden_size];
15687 for (pos, &id) in ids.iter().enumerate() {
15688 let emb = self.embed_single(id);
15689 hidden = self.forward_layers(&emb, pos, task_mask);
15690 }
15691 if let Err(err) = self.o1_seal_checked() {
15692 self.o1_fail(err);
15693 }
15694 if let Some(logits) = self.graph_logits.take() {
15697 return logits;
15698 }
15699 inference::rms_norm_into(
15700 &hidden,
15701 &self.weights.final_norm,
15702 self.rms_eps,
15703 self.norm_style,
15704 &mut self.ws.n1,
15705 );
15706 self.lm_head_forward(&self.ws.n1)
15707 }
15708}
15709
15710pub fn create_test_pipeline(
15712 hidden_size: usize,
15713 intermediate_size: usize,
15714 num_heads: usize,
15715 num_kv_heads: usize,
15716 head_dim: usize,
15717 num_layers: usize,
15718 vocab_size: usize,
15719) -> Pipeline {
15720 let synth = |n: usize, salt: usize| -> Vec<f32> {
15723 (0..n)
15724 .map(|i| (((i * 31 + salt * 17 + 7) % 97) as f32 / 97.0 - 0.5) * 0.2)
15725 .collect()
15726 };
15727 let qt = |rows: usize, cols: usize, salt: usize| -> QTensor {
15728 QTensor::from_f32(synth(rows * cols, salt), rows, cols)
15729 };
15730 let layer_weights: Vec<LayerWeights> = (0..num_layers)
15731 .map(|li| LayerWeights {
15732 input_norm: vec![1.0; hidden_size],
15733 post_norm: vec![1.0; hidden_size],
15734 attn_out_norm: None,
15735 ffn_out_norm: None,
15736 layer_scale: None,
15737 ffn: FfnKind::Dense(DenseFfn {
15738 gate_proj: qt(intermediate_size, hidden_size, li * 10 + 5),
15739 up_proj: qt(intermediate_size, hidden_size, li * 10 + 6),
15740 down_proj: qt(hidden_size, intermediate_size, li * 10 + 7),
15741 act: Act::Silu,
15742 down_t: None,
15743 segs: Vec::new(),
15744 }),
15745 attn: AttnKind::Full {
15746 bias: None,
15747 wq: qt(num_heads * head_dim, hidden_size, li * 10 + 1),
15748 wk: qt(num_kv_heads * head_dim, hidden_size, li * 10 + 2),
15749 wv: qt(num_kv_heads * head_dim, hidden_size, li * 10 + 3),
15750 wo: qt(hidden_size, num_heads * head_dim, li * 10 + 4),
15751 q_norm: None,
15752 k_norm: None,
15753 output_gate: false,
15754 softplus_gate: None,
15755 },
15756 })
15757 .collect();
15758
15759 Pipeline::new(
15760 Tokenizer::byte_level(),
15761 PipelineWeights {
15762 embed_tokens: qt(vocab_size, hidden_size, 100),
15763 layers: layer_weights,
15764 lm_head: qt(vocab_size, hidden_size, 200),
15765 final_norm: vec![1.0; hidden_size],
15766 },
15767 hidden_size,
15768 intermediate_size,
15769 num_heads,
15770 num_kv_heads,
15771 head_dim,
15772 num_layers,
15773 num_layers, false, vocab_size,
15776 1e-6,
15777 10_000.0,
15778 NormStyle::Qwen,
15779 4096,
15780 SamplerConfig {
15781 seed: Some(42),
15782 ..Default::default()
15783 },
15784 )
15785}
15786
15787#[inline]
15792fn mask_bit(row: &[u8], j: usize) -> bool {
15793 (row.get(j >> 3).copied().unwrap_or(0) >> (j & 7)) & 1 != 0
15794}
15795
15796fn mask_gain() -> f32 {
15807 static G: std::sync::OnceLock<f32> = std::sync::OnceLock::new();
15808 *G.get_or_init(|| {
15809 std::env::var("CMF_FFN_MASK_GAIN")
15810 .ok()
15811 .and_then(|v| v.parse().ok())
15812 .unwrap_or(1.0)
15813 })
15814}
15815
15816fn zero_masked_cols(g: &mut [f32], rows: usize, inter: usize, row: &[u8]) {
15817 let fill = meanfill().and_then(|(i, v)| {
15820 let li = crate::gpu::cur_layer();
15821 (*i == inter && li >= 0).then(|| &v[li as usize * inter..(li as usize + 1) * inter])
15822 });
15823 for r in 0..rows {
15824 let base = r * inter;
15825 for (bi, &byte) in row.iter().enumerate() {
15826 if byte == 0xFF {
15827 continue;
15828 }
15829 let j0 = bi * 8;
15830 for bit in 0..8 {
15831 let j = j0 + bit;
15832 if j < inter && byte & (1 << bit) == 0 {
15833 g[base + j] = fill.map_or(0.0, |f| f[j]);
15834 }
15835 }
15836 }
15837 }
15838 let gain = mask_gain();
15839 if gain != 1.0 {
15840 for v in g[..rows * inter].iter_mut() {
15841 *v *= gain;
15842 }
15843 }
15844}
15845
15846#[inline]
15848fn tube_bit(row: Option<&[u8]>, i: usize) -> bool {
15849 row.is_none_or(|r| mask_bit(r, i))
15850}
15851
15852fn all_bits_on(row: &[u8], n: usize) -> bool {
15855 (0..n).all(|i| mask_bit(row, i))
15856}
15857
15858fn tube_topk() -> usize {
15866 static K: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
15867 *K.get_or_init(|| {
15868 std::env::var("CMF_TUBE_TOPK")
15869 .ok()
15870 .and_then(|v| v.parse().ok())
15871 .unwrap_or(0)
15872 })
15873}
15874
15875fn tube_score_oracle() -> bool {
15876 static O: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
15877 *O.get_or_init(|| std::env::var("CMF_TUBE_SCORE").is_ok_and(|v| v == "oracle"))
15878}
15879
15880fn tube_ffn_routed(
15887 d: &DenseFfn,
15888 xs: &[f32],
15889 b: usize,
15890 pool: Option<&Pool>,
15891 mask_row: Option<&[u8]>,
15892 k: usize,
15893) -> Vec<f32> {
15894 let hidden = d.down_proj.rows();
15895 let core = d.gate_proj.rows();
15896 let core_full = mask_row.is_none_or(|r| all_bits_on(r, core));
15897 let mut out = match (b, core_full, mask_row) {
15898 (1, true, _) => dense_ffn(d, xs, pool),
15899 (1, false, Some(row)) => dense_ffn_masked(d, xs, pool, row),
15900 (_, true, _) => dense_ffn_batch(d, xs, b, pool, None),
15901 (_, false, row) => dense_ffn_batch(d, xs, b, pool, row),
15902 };
15903 let cand: Vec<usize> = (0..d.segs.len())
15904 .filter(|&i| tube_bit(mask_row, d.segs[i].start))
15905 .collect();
15906 if cand.is_empty() {
15907 return out;
15908 }
15909 let oracle = tube_score_oracle();
15913 let mut acts: Vec<Vec<f32>> = Vec::with_capacity(cand.len());
15914 let mut scores = vec![0f32; b * cand.len()];
15915 for (ci, &i) in cand.iter().enumerate() {
15916 let seg = &d.segs[i];
15917 let w = seg.width;
15918 let mut g = vec![0.0f32; b * w];
15919 if b == 1 {
15920 seg.gate.matvec(xs, &mut g, pool);
15921 } else {
15922 seg.gate.matmat(xs, b, &mut g, pool);
15923 }
15924 for v in g.iter_mut() {
15925 *v = Act::Silu.combine(*v, 1.0);
15926 }
15927 if !oracle {
15928 for t in 0..b {
15929 scores[t * cand.len() + ci] =
15930 g[t * w..(t + 1) * w].iter().map(|v| v * v).sum::<f32>();
15931 }
15932 }
15933 if oracle || b > 1 {
15934 let mut u = vec![0.0f32; b * w];
15935 if b == 1 {
15936 seg.up.matvec(xs, &mut u, pool);
15937 } else {
15938 seg.up.matmat(xs, b, &mut u, pool);
15939 }
15940 for (a, &v) in g.iter_mut().zip(u.iter()) {
15941 *a *= v;
15942 }
15943 if oracle {
15944 for t in 0..b {
15945 scores[t * cand.len() + ci] =
15946 g[t * w..(t + 1) * w].iter().map(|v| v * v).sum::<f32>();
15947 }
15948 }
15949 }
15950 acts.push(g);
15951 }
15952 let keep = k.min(cand.len());
15954 let mut scratch: Vec<f32> = Vec::new();
15955 for t in 0..b {
15956 let mut sc: Vec<(f32, usize)> = (0..cand.len())
15957 .map(|ci| (scores[t * cand.len() + ci], ci))
15958 .collect();
15959 sc.sort_unstable_by(|x, y| y.0.total_cmp(&x.0));
15960 let mut alive = vec![false; cand.len()];
15961 for &(_, ci) in sc.iter().take(keep) {
15962 alive[ci] = true;
15963 }
15964 if b > 1 {
15965 for (ci, a) in acts.iter_mut().enumerate() {
15966 if !alive[ci] {
15967 let w = d.segs[cand[ci]].width;
15968 a[t * w..(t + 1) * w].fill(0.0);
15969 }
15970 }
15971 } else {
15972 for (ci, &i) in cand.iter().enumerate() {
15976 if !alive[ci] {
15977 continue;
15978 }
15979 let seg = &d.segs[i];
15980 let w = seg.width;
15981 let g = &mut acts[ci];
15982 if !tube_score_oracle() {
15983 scratch.clear();
15984 scratch.resize(w, 0.0);
15985 seg.up.matvec(xs, &mut scratch, pool);
15986 for (a, &v) in g.iter_mut().zip(scratch.iter()) {
15987 *a *= v;
15988 }
15989 }
15990 let mut acc = vec![0.0f32; hidden];
15991 seg.down.matvec(g, &mut acc, pool);
15992 for (o, a) in out.iter_mut().zip(&acc) {
15993 *o += *a;
15994 }
15995 }
15996 }
15997 }
15998 if b > 1 {
15999 for (ci, &i) in cand.iter().enumerate() {
16000 let seg = &d.segs[i];
16001 let mut acc = vec![0.0f32; b * hidden];
16002 seg.down.matmat(&acts[ci], b, &mut acc, pool);
16003 for (o, a) in out.iter_mut().zip(&acc) {
16004 *o += *a;
16005 }
16006 }
16007 }
16008 out
16009}
16010
16011fn tube_ffn(
16017 d: &DenseFfn,
16018 xs: &[f32],
16019 b: usize,
16020 pool: Option<&Pool>,
16021 mask_row: Option<&[u8]>,
16022) -> Vec<f32> {
16023 if tube_topk() > 0 {
16024 return tube_ffn_routed(d, xs, b, pool, mask_row, tube_topk());
16025 }
16026 let hidden = d.down_proj.rows();
16027 let core = d.gate_proj.rows();
16028 let core_full = mask_row.is_none_or(|r| all_bits_on(r, core));
16029 let mut out = match (b, core_full, mask_row) {
16030 (1, true, _) => dense_ffn(d, xs, pool),
16031 (1, false, Some(row)) => dense_ffn_masked(d, xs, pool, row),
16032 (_, true, _) => dense_ffn_batch(d, xs, b, pool, None),
16033 (_, false, row) => dense_ffn_batch(d, xs, b, pool, row),
16034 };
16035 TUBE_SCRATCH.with(|sc| {
16036 let mut sc = sc.borrow_mut();
16037 let [g, u, acc] = &mut *sc;
16038 for seg in &d.segs {
16039 if !tube_bit(mask_row, seg.start) {
16040 continue;
16041 }
16042 let w = seg.width;
16043 g.resize(b * w, 0.0);
16044 if b == 1
16045 && d.act == Act::Silu
16046 && QTensor::matvec_silu_mul(&seg.gate, &seg.up, xs, g, pool)
16047 {
16048 } else {
16050 u.resize(b * w, 0.0);
16051 if b == 1 {
16052 QTensor::matvec_many([&seg.gate, &seg.up], xs, [g, u], pool);
16053 } else {
16054 seg.gate.matmat(xs, b, g, pool);
16055 seg.up.matmat(xs, b, u, pool);
16056 }
16057 for i in 0..b * w {
16058 g[i] = d.act.combine(g[i], u[i]);
16059 }
16060 }
16061 acc.resize(b * hidden, 0.0);
16062 acc.fill(0.0);
16063 if b == 1 {
16064 seg.down.matvec(g, acc, pool);
16065 } else {
16066 seg.down.matmat(g, b, acc, pool);
16067 }
16068 for (o, a) in out.iter_mut().zip(acc.iter()) {
16069 *o += *a;
16070 }
16071 }
16072 out
16073 })
16074}
16075
16076thread_local! {
16077 static TUBE_SCRATCH: std::cell::RefCell<[Vec<f32>; 3]> =
16081 const { std::cell::RefCell::new([Vec::new(), Vec::new(), Vec::new()]) };
16082}
16083
16084static PREFILL_SPLIT: [std::sync::atomic::AtomicU64; 2] =
16087 [std::sync::atomic::AtomicU64::new(0), std::sync::atomic::AtomicU64::new(0)];
16088
16089fn prefill_prof_on() -> bool {
16090 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
16091 *ON.get_or_init(|| std::env::var_os("CMF_PREFILL_PROF").is_some())
16092}
16093
16094fn dense_ffn_batch(
16095 d: &DenseFfn,
16096 xs: &[f32],
16097 b: usize,
16098 pool: Option<&Pool>,
16099 mask_row: Option<&[u8]>,
16100) -> Vec<f32> {
16101 let inter = d.gate_proj.rows();
16102 let hidden = d.down_proj.rows();
16103 let fused_act = d.act.graph_act();
16113 if mask_row.is_none()
16114 && fused_act.is_some()
16115 && b >= 32
16116 && crate::gpu::enabled_here()
16117 && !crate::gpu::mm_killed()
16118 && refit_dir().is_none()
16123 && !ffn_probe_active()
16128 {
16129 if let (Some((model, w1)), Some((_, w3)), Some((_, w2))) = (
16130 d.gate_proj.mapped_q4t(),
16131 d.up_proj.mapped_q4t(),
16132 d.down_proj.mapped_q4t(),
16133 ) {
16134 let mut out = vec![0.0f32; b * hidden];
16135 let act = fused_act.expect("checked above");
16136 if crate::gpu::q4_ffn_act(model, w1, w3, w2, xs, b, hidden, inter, false, act, &mut out)
16137 {
16138 return out;
16139 }
16140 }
16141 if let (Some((model, w1)), Some((_, w3)), Some((_, w2))) = (
16146 d.gate_proj.mapped_q4tp(),
16147 d.up_proj.mapped_q4tp(),
16148 d.down_proj.mapped_q4tp(),
16149 ) {
16150 let mut out = vec![0.0f32; b * hidden];
16151 let act = fused_act.expect("checked above");
16152 if crate::gpu::q4_ffn_act(model, w1, w3, w2, xs, b, hidden, inter, true, act, &mut out)
16153 {
16154 return out;
16155 }
16156 }
16157 if let (true, Some((model, w1)), Some((_, w3)), Some((_, w2))) = (
16163 std::env::var("CMF_FFN_KEEP").as_deref() != Ok("0"),
16164 d.gate_proj.mapped_device_gemm(),
16165 d.up_proj.mapped_device_gemm(),
16166 d.down_proj.mapped_device_gemm(),
16167 ) {
16168 let mut out = vec![0.0f32; b * hidden];
16169 let act = fused_act.expect("checked above");
16170 if crate::gpu::ffn_act_keep(model, w1, w3, w2, xs, b, hidden, inter, act, &mut out) {
16171 return out;
16172 }
16173 }
16174 }
16175 let mut g = vec![0.0f32; b * inter];
16176 d.gate_proj.matmat(xs, b, &mut g, pool);
16177 let mut u = vec![0.0f32; b * inter];
16178 d.up_proj.matmat(xs, b, &mut u, pool);
16179 if gate_topk() > 0 && d.act == Act::Silu {
16180 for t in 0..b {
16181 let row = &mut g[t * inter..(t + 1) * inter];
16182 for v in row.iter_mut() {
16183 *v = Act::Silu.combine(*v, 1.0);
16184 }
16185 keep_top_k(row, gate_topk());
16186 }
16187 for i in 0..b * inter {
16188 g[i] *= u[i];
16189 }
16190 } else {
16191 for i in 0..b * inter {
16192 g[i] = d.act.combine(g[i], u[i]);
16193 }
16194 }
16195 if let Some(row) = mask_row {
16196 zero_masked_cols(&mut g, b, inter, row);
16197 }
16198 if oracle_topk() > 0 {
16199 for t in 0..b {
16200 keep_top_k(&mut g[t * inter..(t + 1) * inter], oracle_topk());
16201 }
16202 }
16203 let mut out = vec![0.0f32; b * hidden];
16204 d.down_proj.matmat(&g, b, &mut out, pool);
16205 if refit_dir().is_some() {
16206 let li = crate::gpu::cur_layer();
16207 if li >= 0 {
16208 refit_accumulate(li as usize, &g, b, inter, &out, hidden, pool);
16209 }
16210 }
16211 FFN_PROBE.with(|pr| {
16215 if let Some(acc) = pr.borrow_mut().as_mut() {
16216 let li = crate::gpu::cur_layer();
16217 if li < 0 {
16218 return;
16219 }
16220 let Some(row) = acc.get_mut(li as usize) else {
16221 return;
16222 };
16223 let sq = probe_sq();
16224 for t in 0..b {
16225 for (a, &v) in row.iter_mut().zip(&g[t * inter..(t + 1) * inter]) {
16226 *a += if sq {
16227 (v as f64) * (v as f64)
16228 } else {
16229 (v as f64).abs()
16230 };
16231 }
16232 }
16233 }
16234 });
16235 out
16236}
16237
16238fn accumulate_act(m: &MoeFfn, xs: &[f32], b: usize) {
16243 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
16244 static DUMP: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
16245 let on = *ON.get_or_init(|| std::env::var("CMF_RMS_TRACE").is_ok());
16246 let dump = *DUMP.get_or_init(|| std::env::var("CMF_ACT_DUMP").is_ok());
16247 if (!on && !dump) || b == 0 {
16248 return;
16249 }
16250 let hidden = xs.len() / b;
16251 if on {
16252 let mut acc = m.act_sq.borrow_mut();
16253 if acc.len() < hidden {
16254 acc.resize(hidden, 0.0);
16255 }
16256 for t in 0..b {
16257 let row = &xs[t * hidden..(t + 1) * hidden];
16258 for (a, &v) in acc.iter_mut().zip(row) {
16259 *a += (v as f64) * (v as f64);
16260 }
16261 }
16262 }
16263 if dump {
16264 let cap: usize = std::env::var("CMF_ACT_DUMP_ROWS")
16267 .ok()
16268 .and_then(|v| v.parse().ok())
16269 .unwrap_or(4096);
16270 let mut rows = m.act_rows.borrow_mut();
16271 if rows.len() < cap * hidden {
16272 let take = b.min((cap * hidden - rows.len()) / hidden.max(1));
16273 rows.extend_from_slice(&xs[..take * hidden]);
16274 }
16275 }
16276}
16277
16278#[derive(Clone, Copy)]
16281struct SendVecs(*mut Vec<f32>);
16282unsafe impl Send for SendVecs {}
16283unsafe impl Sync for SendVecs {}
16284impl SendVecs {
16285 #[inline]
16286 fn at(self, i: usize) -> *mut Vec<f32> {
16287 unsafe { self.0.add(i) }
16288 }
16289}
16290
16291fn moe_ffn_batch(
16292 m: &MoeFfn,
16293 xs: &[f32],
16294 b: usize,
16295 hidden: usize,
16296 pool: Option<&Pool>,
16297 allowed: Option<&[bool]>,
16298) -> Vec<f32> {
16299 accumulate_act(m, xs, b);
16300 let ne = m.experts.len();
16301 let mut logits = vec![0.0f32; b * ne];
16302 match &m.resonance {
16303 Some(r) => {
16304 let hdim = xs.len() / b.max(1);
16305 for bi in 0..b {
16306 r.scores(
16307 &xs[bi * hdim..(bi + 1) * hdim],
16308 &mut logits[bi * ne..(bi + 1) * ne],
16309 );
16310 }
16311 }
16312 None => m.router.matmat(xs, b, &mut logits, pool),
16313 }
16314
16315 let mut assign: Vec<Vec<(usize, f32)>> = vec![Vec::new(); ne];
16318 {
16319 let mut st = m.stats.borrow_mut();
16320 if st.len() < ne {
16321 st.resize(ne, 0);
16322 }
16323 for bi in 0..b {
16324 let (idx, p, wsum) = moe_route(&logits[bi * ne..(bi + 1) * ne], m, allowed);
16325 for &e in &idx {
16326 st[e] += 1;
16327 assign[e].push((bi, p[e] / wsum));
16328 }
16329 }
16330 }
16331
16332 let mut out = vec![0.0f32; b * hidden];
16333 let cols = m.experts[0].gate_proj.cols();
16334 let run_expert = |d: &DenseFfn, list: &[(usize, f32)], out: &mut [f32]| {
16335 let sb = list.len();
16336 let mut sub = vec![0.0f32; sb * cols];
16337 for (k, &(bi, _)) in list.iter().enumerate() {
16338 sub[k * cols..(k + 1) * cols].copy_from_slice(&xs[bi * cols..(bi + 1) * cols]);
16339 }
16340 let eo = dense_ffn_batch(d, &sub, sb, pool, None);
16341 for (k, &(bi, w)) in list.iter().enumerate() {
16342 for i in 0..hidden {
16343 out[bi * hidden + i] += w * eo[k * hidden + i];
16344 }
16345 }
16346 };
16347 let active: Vec<usize> = (0..ne).filter(|&e| !assign[e].is_empty()).collect();
16353 if pool.is_some() && active.len() >= 8 {
16354 let mut panels: Vec<Vec<f32>> = vec![Vec::new(); active.len()];
16355 {
16356 let panel_ptr = SendVecs(panels.as_mut_ptr());
16357 let experts = &m.experts;
16360 let (active_r, assign_r) = (&active, &assign);
16361 let inherit_cpu = crate::gpu::inherit_cpu_scope();
16362 let run = |start: usize, end: usize| {
16363 let _cpu_scope = inherit_cpu();
16364 for ai in start..end {
16365 let e = active_r[ai];
16366 let list = &assign_r[e];
16367 let sb = list.len();
16368 let mut sub = vec![0.0f32; sb * cols];
16369 for (k, &(bi, _)) in list.iter().enumerate() {
16370 sub[k * cols..(k + 1) * cols]
16371 .copy_from_slice(&xs[bi * cols..(bi + 1) * cols]);
16372 }
16373 unsafe {
16375 *panel_ptr.at(ai) = dense_ffn_batch(&experts[e], &sub, sb, None, None);
16376 }
16377 }
16378 };
16379 match pool {
16380 Some(p) => p.run_rows(active.len(), &run),
16381 None => run(0, active.len()),
16382 }
16383 }
16384 for (ai, &e) in active.iter().enumerate() {
16385 for (k, &(bi, w)) in assign[e].iter().enumerate() {
16386 let eo = &panels[ai][k * hidden..(k + 1) * hidden];
16387 for i in 0..hidden {
16388 out[bi * hidden + i] += w * eo[i];
16389 }
16390 }
16391 }
16392 } else {
16393 for &e in &active {
16394 run_expert(&m.experts[e], &assign[e], &mut out);
16395 }
16396 }
16397 if let Some((se, gate)) = &m.shared {
16398 let all: Vec<(usize, f32)> = if let Some(gate) = gate {
16399 let mut gl = vec![0.0f32; b];
16400 gate.matmat(xs, b, &mut gl, pool);
16401 (0..b)
16402 .map(|bi| (bi, 1.0 / (1.0 + (-gl[bi]).exp())))
16403 .collect()
16404 } else {
16405 (0..b).map(|bi| (bi, 1.0)).collect()
16406 };
16407 run_expert(se, &all, &mut out);
16408 }
16409 out
16410}
16411
16412fn moe_ffn_rows_exact(
16424 m: &MoeFfn,
16425 xs: &[f32],
16426 b: usize,
16427 hidden: usize,
16428 pool: Option<&Pool>,
16429) -> Vec<f32> {
16430 let mut out = vec![0.0f32; b * hidden];
16431 let per_row = |out: &mut [f32]| {
16432 for r in 0..b {
16433 let o = moe_ffn(m, &xs[r * hidden..(r + 1) * hidden], pool, None);
16434 out[r * hidden..(r + 1) * hidden].copy_from_slice(&o);
16435 }
16436 };
16437 let covered = !crate::gpu::enabled_here()
16438 && moe_batch_enabled()
16439 && m.shared.is_none()
16440 && m.resonance.is_none()
16441 && FFN_PROBE.with(|pr| pr.borrow().is_none())
16442 && m.experts.iter().all(|d| d.act == Act::Silu);
16443 if !covered {
16444 per_row(&mut out);
16445 return out;
16446 }
16447 let ne = m.experts.len();
16448 let mut routes: Vec<(Vec<usize>, Vec<f32>)> = Vec::with_capacity(b);
16450 for r in 0..b {
16451 let x = &xs[r * hidden..(r + 1) * hidden];
16452 accumulate_act(m, x, 1);
16453 let mut logits = vec![0.0f32; ne];
16454 m.router.matvec(x, &mut logits, pool);
16455 let (idx, p, wsum) = moe_route(&logits, m, None);
16456 {
16457 let mut st = m.stats.borrow_mut();
16458 if st.len() < ne {
16459 st.resize(ne, 0);
16460 }
16461 for &e in &idx {
16462 st[e] += 1;
16463 }
16464 }
16465 let w: Vec<f32> = idx
16466 .iter()
16467 .map(|&e| p[e] / wsum * m.per_expert_scale.as_ref().map_or(1.0, |v| v[e]))
16468 .collect();
16469 routes.push((idx, w));
16470 }
16471 if routes.iter().any(|(idx, _)| idx.is_empty()) {
16472 per_row(&mut out);
16473 return out;
16474 }
16475 let mut experts: Vec<usize> = Vec::new();
16477 let mut groups: Vec<Vec<usize>> = Vec::new();
16478 for (r, (idx, _)) in routes.iter().enumerate() {
16479 for &e in idx {
16480 match experts.iter().position(|&x| x == e) {
16481 Some(g) => groups[g].push(r),
16482 None => {
16483 experts.push(e);
16484 groups.push(vec![r]);
16485 }
16486 }
16487 }
16488 }
16489 let n_pairs: usize = groups.iter().map(|g| g.len()).sum();
16490 let inter = m.experts[experts[0]].gate_proj.rows();
16491 let pairs: Vec<(&QTensor, &QTensor)> = experts
16492 .iter()
16493 .map(|&e| (&m.experts[e].gate_proj, &m.experts[e].up_proj))
16494 .collect();
16495 let mut gs: Vec<Vec<f32>> = (0..n_pairs).map(|_| vec![0f32; inter]).collect();
16496 if !QTensor::moe_gate_up_rows(&pairs, &groups, xs, &mut gs, pool) {
16497 per_row(&mut out);
16498 return out;
16499 }
16500 let downs: Vec<&QTensor> = experts.iter().map(|&e| &m.experts[e].down_proj).collect();
16501 let lens: Vec<usize> = groups.iter().map(|g| g.len()).collect();
16502 let mut ds: Vec<Vec<f32>> = (0..n_pairs).map(|_| vec![0f32; hidden]).collect();
16503 if !QTensor::moe_down_rows(&downs, &lens, &gs, &mut ds, pool) {
16504 per_row(&mut out);
16505 return out;
16506 }
16507 let mut slot = std::collections::HashMap::with_capacity(n_pairs);
16509 let mut p = 0usize;
16510 for (g, &e) in experts.iter().enumerate() {
16511 for &r in &groups[g] {
16512 slot.insert((r, e), p);
16513 p += 1;
16514 }
16515 }
16516 for (r, (idx, w)) in routes.iter().enumerate() {
16517 let terms: Vec<(&[f32], f32)> = idx
16518 .iter()
16519 .zip(w)
16520 .map(|(&e, &we)| (ds[slot[&(r, e)]].as_slice(), we))
16521 .collect();
16522 let row = &mut out[r * hidden..(r + 1) * hidden];
16523 for (i, dst) in row.iter_mut().enumerate() {
16524 let mut acc = 0f32;
16526 for (d, we) in &terms {
16527 acc += we * d[i];
16528 }
16529 *dst = acc;
16530 }
16531 }
16532 out
16533}
16534
16535thread_local! {
16536 static FFN_SCRATCH: std::cell::RefCell<[Vec<f32>; 4]> =
16540 const { std::cell::RefCell::new([Vec::new(), Vec::new(), Vec::new(), Vec::new()]) };
16541}
16542
16543fn dense_ffn(d: &DenseFfn, x: &[f32], pool: Option<&Pool>) -> Vec<f32> {
16545 if gate_topk() > 0
16548 && let Some(out) = dense_ffn_dynamic(d, x, pool, gate_topk())
16549 {
16550 return out;
16551 }
16552 let prism_body = d.gate_proj.has_prism_contract()
16568 || d.up_proj.has_prism_contract()
16569 || d.down_proj.has_prism_contract();
16570 if !prism_body
16571 && crate::gpu::enabled_here()
16572 && (d.gate_proj.rows() >= crate::gpu::min_rows() || d.gate_proj.is_q1())
16573 {
16574 let arm = if d.gate_proj.is_q1() && crate::gpu::q1_force() {
16575 crate::gpu::ProbeArm::Gpu
16576 } else {
16577 crate::gpu::probe_arm(crate::gpu::OpClass::Ffn)
16578 };
16579 match arm {
16580 crate::gpu::ProbeArm::Gpu => {
16581 let t0 = std::time::Instant::now();
16582 if let Some(out) = dense_ffn_gpu(d, x, pool) {
16583 crate::gpu::probe_record(crate::gpu::OpClass::Ffn, true, t0.elapsed());
16584 return out;
16585 }
16586 crate::gpu::probe_note_decline(crate::gpu::OpClass::Ffn);
16590 }
16591 crate::gpu::ProbeArm::CpuTimed => {
16592 let t0 = std::time::Instant::now();
16593 let out = crate::gpu::cpu_scope(|| dense_ffn_cpu(d, x, pool));
16594 crate::gpu::probe_record(crate::gpu::OpClass::Ffn, false, t0.elapsed());
16595 return out;
16596 }
16597 crate::gpu::ProbeArm::Cpu => {
16598 return crate::gpu::cpu_scope(|| dense_ffn_cpu(d, x, pool));
16599 }
16600 }
16601 }
16602 dense_ffn_cpu(d, x, pool)
16603}
16604
16605fn dense_ffn_cpu(d: &DenseFfn, x: &[f32], pool: Option<&Pool>) -> Vec<f32> {
16607 let inter = d.gate_proj.rows();
16608 FFN_SCRATCH.with(|s| {
16609 let mut s = s.borrow_mut();
16610 let [g, u, ..] = &mut *s;
16611 g.resize(inter, 0.0);
16612 if gate_topk() > 0 {
16615 u.resize(inter, 0.0);
16619 QTensor::matvec_many([&d.gate_proj, &d.up_proj], x, [g, u], pool);
16620 for i in 0..inter {
16621 g[i] = Act::Silu.combine(g[i], 1.0);
16622 }
16623 keep_top_k(g, gate_topk());
16624 for i in 0..inter {
16625 g[i] *= u[i];
16626 }
16627 } else if d.act == Act::Silu && {
16628 let _prof = crate::cpuprof::time(crate::cpuprof::Slot::FfnGateUp);
16629 QTensor::matvec_silu_mul(&d.gate_proj, &d.up_proj, x, g, pool)
16630 } {
16631 } else {
16633 u.resize(inter, 0.0);
16634 let _prof = crate::cpuprof::time(crate::cpuprof::Slot::FfnGateUp);
16636 QTensor::matvec_many([&d.gate_proj, &d.up_proj], x, [g, u], pool);
16637 for i in 0..inter {
16638 g[i] = d.act.combine(g[i], u[i]);
16639 }
16640 }
16641 FFN_PROBE.with(|pr| {
16649 if let Some(acc) = pr.borrow_mut().as_mut() {
16650 let li = crate::gpu::cur_layer();
16651 if li >= 0 {
16652 if let Some(row) = acc.get_mut(li as usize) {
16653 match probe_topk() {
16654 0 if probe_sq() => {
16655 for (a, &v) in row.iter_mut().zip(g.iter()) {
16656 *a += (v as f64) * (v as f64);
16657 }
16658 }
16659 0 if probe_signed() => {
16660 for (a, &v) in row.iter_mut().zip(g.iter()) {
16661 *a += v as f64;
16662 }
16663 }
16664 0 => {
16665 for (a, &v) in row.iter_mut().zip(g.iter()) {
16666 *a += (v as f64).abs();
16667 }
16668 }
16669 k => {
16670 let n = g.len();
16671 let k = k.min(n);
16672 let mut mag: Vec<f32> = g.iter().map(|v| v.abs()).collect();
16673 let (_, kth, _) = mag.select_nth_unstable_by(k - 1, |a, b| {
16674 b.partial_cmp(a).unwrap_or(std::cmp::Ordering::Equal)
16675 });
16676 let thr = *kth;
16677 for (a, &v) in row.iter_mut().zip(g.iter()) {
16678 if v.abs() >= thr {
16679 *a += 1.0;
16680 }
16681 }
16682 }
16683 }
16684 }
16685 }
16686 }
16687 });
16688 if oracle_topk() > 0 {
16689 keep_top_k(g, oracle_topk());
16690 }
16691 {
16692 let li = crate::gpu::cur_layer();
16693 if li >= 0 {
16694 adump_row(li as usize, g);
16695 }
16696 }
16697 let mut out = attention::take_buf(d.down_proj.rows());
16698 let _prof = crate::cpuprof::time(crate::cpuprof::Slot::FfnDown);
16699 d.down_proj.matvec(g, &mut out, pool);
16700 out
16701 })
16702}
16703
16704pub struct RefitAcc {
16717 pub support: Vec<u32>,
16718 pub gss: Vec<f32>,
16719 pub ya: Vec<f32>,
16720 pub hidden: usize,
16721 pub tokens: u64,
16722 pub buf_g: Vec<f32>,
16728 pub buf_o: Vec<f32>,
16729 pub buf_t: usize,
16730}
16731
16732type RefitState = (std::collections::HashMap<usize, RefitAcc>, Vec<f32>);
16736
16737static REFIT: std::sync::OnceLock<Option<(String, std::sync::Mutex<RefitState>)>> =
16738 std::sync::OnceLock::new();
16739
16740fn ffn_probe_active() -> bool {
16743 FFN_PROBE.with(|p| p.borrow().is_some())
16744}
16745
16746fn refit_dir() -> Option<&'static (String, std::sync::Mutex<RefitState>)> {
16747 REFIT
16748 .get_or_init(|| {
16749 std::env::var("CMF_FFN_REFIT").ok().map(|d| {
16750 (
16751 d,
16752 std::sync::Mutex::new((std::collections::HashMap::new(), Vec::new())),
16753 )
16754 })
16755 })
16756 .as_ref()
16757}
16758
16759fn refit_accumulate(
16761 li: usize,
16762 g: &[f32],
16763 b: usize,
16764 inter: usize,
16765 out: &[f32],
16766 hidden: usize,
16767 pool: Option<&Pool>,
16768) {
16769 let Some((dir, map)) = refit_dir() else {
16770 return;
16771 };
16772 static SPAN: std::sync::OnceLock<(usize, usize)> = std::sync::OnceLock::new();
16773 let (from, to) = *SPAN.get_or_init(|| {
16774 let g = |k: &str, d: usize| {
16775 std::env::var(k)
16776 .ok()
16777 .and_then(|v| v.parse().ok())
16778 .unwrap_or(d)
16779 };
16780 (
16781 g("CMF_FFN_REFIT_FROM", 0),
16782 g("CMF_FFN_REFIT_TO", usize::MAX),
16783 )
16784 });
16785 if li < from || li > to {
16786 return;
16787 }
16788 let mut guard = map.lock().unwrap();
16789 let (map, shared) = &mut *guard;
16790 let acc = match map.entry(li) {
16791 std::collections::hash_map::Entry::Occupied(e) => e.into_mut(),
16792 std::collections::hash_map::Entry::Vacant(e) => {
16793 let path = format!("{dir}/support.{li}.u32");
16794 let Ok(bytes) = std::fs::read(&path) else {
16795 eprintln!("refit: no {path} — layer {li} skipped");
16796 return;
16797 };
16798 let n = u32::from_le_bytes(bytes[0..4].try_into().unwrap()) as usize;
16799 let support: Vec<u32> = bytes[4..4 + n * 4]
16800 .chunks_exact(4)
16801 .map(|c| u32::from_le_bytes([c[0], c[1], c[2], c[3]]))
16802 .collect();
16803 eprintln!(
16804 "refit: layer {li} support {n} ({:.0} MB of accumulator)",
16805 (n * n + hidden * n) as f64 * 4.0 / 1e6
16806 );
16807 e.insert(RefitAcc {
16808 gss: vec![0.0; n * n],
16809 ya: vec![0.0; hidden * n],
16810 buf_g: Vec::new(),
16811 buf_o: Vec::new(),
16812 buf_t: 0,
16813 support,
16814 hidden,
16815 tokens: 0,
16816 })
16817 }
16818 };
16819 let ns = acc.support.len();
16820 let cap = refit_batch();
16822 if acc.buf_g.is_empty() {
16823 acc.buf_g = vec![0.0; ns * cap];
16824 acc.buf_o = vec![0.0; hidden * cap];
16825 }
16826 let take = b.min(cap - acc.buf_t);
16827 for t in 0..take {
16828 let col = acc.buf_t + t;
16829 for (j, &n) in acc.support.iter().enumerate() {
16830 acc.buf_g[j * cap + col] = g[t * inter + n as usize];
16831 }
16832 for h in 0..hidden {
16833 acc.buf_o[h * cap + col] = out[t * hidden + h];
16834 }
16835 }
16836 acc.buf_t += take;
16837 acc.tokens += take as u64;
16838 if acc.buf_t < cap {
16839 return;
16840 }
16841 let bt = acc.buf_t;
16842 acc.buf_t = 0;
16843 let RefitAcc {
16853 gss,
16854 ya,
16855 buf_g,
16856 buf_o,
16857 ..
16858 } = acc;
16859 let need = (ns * ns).max(hidden * ns);
16860 if shared.len() < need {
16861 shared.resize(need, 0.0);
16862 }
16863 let scratch = &mut shared[..];
16864 let _ = bt;
16865 if crate::gpu::gemm_nt_f32_transient(buf_g, buf_g, &mut scratch[..ns * ns], ns, cap, ns) {
16866 add_into(gss, &scratch[..ns * ns], pool);
16867 if crate::gpu::gemm_nt_f32_transient(
16868 buf_o,
16869 buf_g,
16870 &mut scratch[..hidden * ns],
16871 hidden,
16872 cap,
16873 ns,
16874 ) {
16875 add_into(ya, &scratch[..hidden * ns], pool);
16876 } else {
16877 accum_outer_t(ya, hidden, ns, cap, buf_o, buf_g, pool);
16878 }
16879 } else {
16880 accum_outer_t(gss, ns, ns, cap, buf_g, buf_g, pool);
16881 accum_outer_t(ya, hidden, ns, cap, buf_o, buf_g, pool);
16882 }
16883 }
16887
16888fn refit_batch() -> usize {
16890 static B: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
16891 *B.get_or_init(|| {
16892 std::env::var("CMF_FFN_REFIT_BATCH")
16893 .ok()
16894 .and_then(|v| v.parse().ok())
16895 .unwrap_or(4096)
16896 })
16897}
16898
16899fn accum_outer_t(
16902 c: &mut [f32],
16903 m: usize,
16904 n: usize,
16905 b: usize,
16906 left: &[f32],
16907 right: &[f32],
16908 pool: Option<&Pool>,
16909) {
16910 let ptr = SendMut(c.as_mut_ptr());
16911 let body = |i: usize| {
16912 let ptr = &ptr;
16913 let row = unsafe { std::slice::from_raw_parts_mut(ptr.0.add(i * n), n) };
16914 for t in 0..b {
16915 let a = left[i * b + t];
16916 if a == 0.0 {
16917 continue;
16918 }
16919 for (j, o) in row.iter_mut().enumerate() {
16920 *o += a * right[j * b + t];
16921 }
16922 }
16923 };
16924 match pool {
16925 Some(p) if m > 1 => p.run_rows(m, &|s, e| {
16926 for i in s..e {
16927 body(i);
16928 }
16929 }),
16930 _ => {
16931 for i in 0..m {
16932 body(i);
16933 }
16934 }
16935 }
16936}
16937
16938fn add_into(dst: &mut [f32], src: &[f32], pool: Option<&Pool>) {
16941 let n = dst.len().min(src.len());
16942 match pool {
16943 Some(p) if n >= 1 << 16 => {
16944 let ptr = SendMut(dst.as_mut_ptr());
16945 let f = |s: usize, e: usize| {
16946 let ptr = &ptr;
16947 for blk in s..e {
16948 let (a, b) = (blk * 4096, ((blk + 1) * 4096).min(n));
16949 for i in a..b {
16950 unsafe { *ptr.0.add(i) += src[i] };
16951 }
16952 }
16953 };
16954 p.run_rows(n.div_ceil(4096), &f);
16955 }
16956 _ => {
16957 for (d, v) in dst.iter_mut().zip(&src[..n]) {
16958 *d += *v;
16959 }
16960 }
16961 }
16962}
16963
16964fn accum_outer(
16969 c: &mut [f32],
16970 m: usize,
16971 n: usize,
16972 b: usize,
16973 left: &[f32],
16974 right: &[f32],
16975 pool: Option<&Pool>,
16976) {
16977 const TILE: usize = 32;
16978 let tiles = m.div_ceil(TILE);
16979 let cp = SendMut(c.as_mut_ptr());
16980 let body = |ti: usize| {
16981 let cp = &cp;
16982 let i0 = ti * TILE;
16983 let i1 = (i0 + TILE).min(m);
16984 for t in 0..b {
16985 let r = &right[t * n..t * n + n];
16986 for i in i0..i1 {
16987 let a = left[i * b + t];
16988 if a == 0.0 {
16989 continue;
16990 }
16991 let row = unsafe { std::slice::from_raw_parts_mut(cp.0.add(i * n), n) };
16993 for (o, v) in row.iter_mut().zip(r) {
16994 *o += a * *v;
16995 }
16996 }
16997 }
16998 };
16999 match pool {
17000 Some(p) if tiles > 1 => p.run_rows(tiles, &|s, e| {
17001 for ti in s..e {
17002 body(ti);
17003 }
17004 }),
17005 _ => {
17006 for ti in 0..tiles {
17007 body(ti);
17008 }
17009 }
17010 }
17011}
17012
17013pub fn refit_flush() -> usize {
17015 let Some((dir, map)) = refit_dir() else {
17016 return 0;
17017 };
17018 let guard = map.lock().unwrap();
17019 let mut n = 0;
17020 for (li, acc) in guard.0.iter() {
17021 let w = |name: &str, v: &[f32]| {
17024 let path = format!("{dir}/{name}.{li}.f32");
17025 let bytes: Vec<u8> = v.iter().flat_map(|x| x.to_le_bytes()).collect();
17026 match std::fs::write(&path, &bytes) {
17027 Ok(()) => {}
17028 Err(e) => eprintln!(
17029 "refit: FAILED to write {path} ({} MB): {e}",
17030 bytes.len() / 1_000_000
17031 ),
17032 }
17033 };
17034 w("gss", &acc.gss);
17035 w("ya", &acc.ya);
17036 println!(
17037 "refit L{li}: {} support, {} tokens, hidden {}",
17038 acc.support.len(),
17039 acc.tokens,
17040 acc.hidden
17041 );
17042 n += 1;
17043 }
17044 n
17045}
17046
17047fn adump_row(li: usize, g: &[f32]) {
17052 use std::io::Write as _;
17053 static FILES: std::sync::OnceLock<
17054 Option<(
17055 String,
17056 std::sync::Mutex<std::collections::HashMap<usize, std::fs::File>>,
17057 )>,
17058 > = std::sync::OnceLock::new();
17059 let Some((prefix, map)) = FILES
17060 .get_or_init(|| {
17061 std::env::var("CMF_FFN_ADUMP")
17062 .ok()
17063 .map(|p| (p, std::sync::Mutex::new(std::collections::HashMap::new())))
17064 })
17065 .as_ref()
17066 else {
17067 return;
17068 };
17069 static SPAN: std::sync::OnceLock<(usize, usize)> = std::sync::OnceLock::new();
17072 let (from, to) = *SPAN.get_or_init(|| {
17073 let g = |k: &str, d: usize| {
17074 std::env::var(k)
17075 .ok()
17076 .and_then(|v| v.parse().ok())
17077 .unwrap_or(d)
17078 };
17079 (
17080 g("CMF_FFN_ADUMP_FROM", 0),
17081 g("CMF_FFN_ADUMP_TO", usize::MAX),
17082 )
17083 });
17084 if li < from || li > to {
17085 return;
17086 }
17087 let mut map = map.lock().unwrap();
17088 let f = map.entry(li).or_insert_with(|| {
17089 std::fs::File::create(format!("{prefix}.{li}.f16")).expect("adump file")
17090 });
17091 let mut bytes = Vec::with_capacity(g.len() * 2);
17092 for v in g {
17093 bytes.extend_from_slice(&cortiq_core::quant::f32_to_f16(*v).to_le_bytes());
17094 }
17095 let _ = f.write_all(&bytes);
17096}
17097
17098fn oracle_topk() -> usize {
17104 static K: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
17105 *K.get_or_init(|| {
17106 std::env::var("CMF_FFN_ORACLE_TOPK")
17107 .ok()
17108 .and_then(|v| v.parse().ok())
17109 .unwrap_or(0)
17110 })
17111}
17112
17113fn gate_topk() -> usize {
17119 static K: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
17120 *K.get_or_init(|| {
17121 std::env::var("CMF_FFN_GATE_TOPK")
17122 .ok()
17123 .and_then(|v| v.parse().ok())
17124 .unwrap_or(0)
17125 })
17126}
17127
17128fn gate_block() -> usize {
17135 static B: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
17136 *B.get_or_init(|| {
17137 std::env::var("CMF_FFN_GATE_BLOCK")
17138 .ok()
17139 .and_then(|v| v.parse().ok())
17140 .unwrap_or(1)
17141 })
17142}
17143
17144fn keep_top_blocks(g: &mut [f32], keep_n: usize, block: usize) {
17146 let n = g.len();
17147 let nb = n.div_ceil(block);
17148 let kb = (keep_n.div_ceil(block)).clamp(1, nb);
17149 if kb >= nb {
17150 return;
17151 }
17152 let mut score: Vec<f32> = (0..nb)
17153 .map(|b| {
17154 g[b * block..((b + 1) * block).min(n)]
17155 .iter()
17156 .map(|v| v * v)
17157 .sum::<f32>()
17158 })
17159 .collect();
17160 let mut ord = score.clone();
17161 let (_, kth, _) = ord.select_nth_unstable_by(kb - 1, |a, b| {
17162 b.partial_cmp(a).unwrap_or(std::cmp::Ordering::Equal)
17163 });
17164 let thr = *kth;
17165 for b in 0..nb {
17166 if score[b] < thr {
17167 g[b * block..((b + 1) * block).min(n)].fill(0.0);
17168 }
17169 }
17170 score.clear();
17171}
17172
17173fn keep_top_k(g: &mut [f32], k: usize) {
17175 if gate_block() > 1 {
17176 return keep_top_blocks(g, k, gate_block());
17177 }
17178 let n = g.len();
17179 if k == 0 || k >= n {
17180 return;
17181 }
17182 let mut mag: Vec<f32> = g.iter().map(|v| v.abs()).collect();
17183 let (_, kth, _) = mag.select_nth_unstable_by(k - 1, |a, b| {
17184 b.partial_cmp(a).unwrap_or(std::cmp::Ordering::Equal)
17185 });
17186 let thr = *kth;
17187 for v in g.iter_mut() {
17188 if v.abs() < thr {
17189 *v = 0.0;
17190 }
17191 }
17192}
17193
17194fn probe_sq() -> bool {
17198 static S: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
17199 *S.get_or_init(|| std::env::var("CMF_FFN_PROBE_SQ").is_ok())
17200}
17201
17202fn probe_signed() -> bool {
17206 static S: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
17207 *S.get_or_init(|| std::env::var("CMF_FFN_PROBE_SIGNED").is_ok())
17208}
17209
17210fn meanfill() -> Option<&'static (usize, Vec<f32>)> {
17218 static M: std::sync::OnceLock<Option<(usize, Vec<f32>)>> = std::sync::OnceLock::new();
17219 M.get_or_init(|| {
17220 let p = std::env::var("CMF_FFN_MEANFILL").ok()?;
17221 let b = std::fs::read(&p).ok()?;
17222 let inter = u32::from_le_bytes(b[4..8].try_into().ok()?) as usize;
17223 let vals: Vec<f32> = b[8..]
17224 .chunks_exact(4)
17225 .map(|c| f32::from_le_bytes([c[0], c[1], c[2], c[3]]))
17226 .collect();
17227 eprintln!("meanfill: {} value(s), inter {inter}", vals.len());
17228 Some((inter, vals))
17229 })
17230 .as_ref()
17231}
17232
17233fn probe_topk() -> usize {
17236 static K: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
17237 *K.get_or_init(|| {
17238 std::env::var("CMF_FFN_PROBE_TOPK")
17239 .ok()
17240 .and_then(|v| v.parse().ok())
17241 .unwrap_or(0)
17242 })
17243}
17244
17245thread_local! {
17246 static FFN_PROBE: std::cell::RefCell<Option<Vec<Vec<f64>>>> =
17249 const { std::cell::RefCell::new(None) };
17250}
17251
17252fn dense_ffn_dynamic(d: &DenseFfn, x: &[f32], pool: Option<&Pool>, k: usize) -> Option<Vec<f32>> {
17265 if d.gate_proj.has_prism_contract()
17270 || d.up_proj.has_prism_contract()
17271 || d.down_proj.has_prism_contract()
17272 {
17273 return None;
17274 }
17275 let dt = d.down_t.as_ref()?;
17276 let inter = d.gate_proj.rows();
17277 let hidden = dt.cols();
17278 if k == 0 || k >= inter || d.act != Act::Silu {
17279 return None;
17280 }
17281 DYN_SCRATCH.with(|sc| {
17282 let mut sc = sc.borrow_mut();
17283 let DynScratch {
17284 g,
17285 mag,
17286 live,
17287 parts,
17288 } = &mut *sc;
17289 g.resize(inter, 0.0);
17290 d.gate_proj.matvec(x, g, pool);
17291 for v in g.iter_mut() {
17292 *v = inference::silu(*v);
17293 }
17294 mag.clear();
17297 mag.extend(g.iter().map(|v| v.abs()));
17298 let (_, kth, _) = mag.select_nth_unstable_by(k - 1, |a, b| {
17299 b.partial_cmp(a).unwrap_or(std::cmp::Ordering::Equal)
17300 });
17301 let thr = *kth;
17302 live.clear();
17303 live.extend((0..inter as u32).filter(|&n| g[n as usize].abs() >= thr));
17304 let mut out = vec![0.0f32; hidden];
17305 match pool {
17306 Some(p) if live.len() >= 64 => {
17307 let nw = p.n_workers() + 1;
17308 parts.clear();
17309 parts.resize(nw * hidden, 0.0);
17310 let ptr = SendMut(parts.as_mut_ptr());
17311 let n = live.len();
17312 let live_ref: &[u32] = live;
17313 let g_ref: &[f32] = g;
17314 p.run(&|w, workers| {
17315 let chunk = n.div_ceil(workers);
17316 let (s, e) = (w * chunk, ((w + 1) * chunk).min(n));
17317 if s >= e {
17318 return;
17319 }
17320 WORKER_SCRATCH.with(|ws| {
17321 let mut ws = ws.borrow_mut();
17322 let [scratch, acc] = &mut *ws;
17323 scratch.resize(hidden.max(x.len()), 0.0);
17324 acc.clear();
17325 acc.resize(hidden, 0.0);
17326 for (o, &nrm) in live_ref[s..e].iter().enumerate() {
17327 if let Some(&nx) = live_ref[s..e].get(o + 1) {
17330 d.up_proj.prefetch_row(nx as usize);
17331 dt.prefetch_row(nx as usize);
17332 }
17333 let idx = nrm as usize;
17334 let up = d.up_proj.row_dot(idx, x, scratch);
17335 let a = g_ref[idx] * up;
17336 if a != 0.0 {
17337 dt.add_row_scaled(idx, a, acc, scratch);
17338 }
17339 }
17340 for (j, v) in acc.iter().enumerate() {
17341 unsafe { *ptr.at(w * hidden + j) = *v };
17342 }
17343 });
17344 });
17345 for w in 0..nw {
17346 for (j, o) in out.iter_mut().enumerate() {
17347 *o += parts[w * hidden + j];
17348 }
17349 }
17350 }
17351 _ => {
17352 WORKER_SCRATCH.with(|ws| {
17353 let mut ws = ws.borrow_mut();
17354 let [scratch, _acc] = &mut *ws;
17355 scratch.resize(hidden.max(x.len()), 0.0);
17356 for &nrm in live.iter() {
17357 let idx = nrm as usize;
17358 let up = d.up_proj.row_dot(idx, x, scratch);
17359 let a = g[idx] * up;
17360 if a != 0.0 {
17361 dt.add_row_scaled(idx, a, &mut out, scratch);
17362 }
17363 }
17364 });
17365 }
17366 }
17367 Some(out)
17368 })
17369}
17370
17371struct DynScratch {
17374 g: Vec<f32>,
17375 mag: Vec<f32>,
17376 live: Vec<u32>,
17377 parts: Vec<f32>,
17378}
17379
17380thread_local! {
17381 static DYN_SCRATCH: std::cell::RefCell<DynScratch> = const {
17382 std::cell::RefCell::new(DynScratch {
17383 g: Vec::new(),
17384 mag: Vec::new(),
17385 live: Vec::new(),
17386 parts: Vec::new(),
17387 })
17388 };
17389 static WORKER_SCRATCH: std::cell::RefCell<[Vec<f32>; 2]> =
17391 const { std::cell::RefCell::new([Vec::new(), Vec::new()]) };
17392}
17393
17394fn dense_ffn_masked(d: &DenseFfn, x: &[f32], pool: Option<&Pool>, mask_row: &[u8]) -> Vec<f32> {
17399 let inter = d.gate_proj.rows();
17400 FFN_SCRATCH.with(|s| {
17401 let mut s = s.borrow_mut();
17402 let [g, u, ..] = &mut *s;
17403 g.resize(inter, 0.0);
17404 if d.act == Act::Silu && QTensor::matvec_silu_mul(&d.gate_proj, &d.up_proj, x, g, pool) {
17405 } else {
17407 u.resize(inter, 0.0);
17408 QTensor::matvec_many([&d.gate_proj, &d.up_proj], x, [g, u], pool);
17409 for i in 0..inter {
17410 g[i] = d.act.combine(g[i], u[i]);
17411 }
17412 }
17413 zero_masked_cols(g, 1, inter, mask_row);
17414 let mut out = attention::take_buf(d.down_proj.rows());
17415 d.down_proj.matvec(g, &mut out, pool);
17416 out
17417 })
17418}
17419
17420fn dense_ffn_gpu(d: &DenseFfn, x: &[f32], _pool: Option<&Pool>) -> Option<Vec<f32>> {
17426 if d.gate_proj.has_prism_contract()
17427 || d.up_proj.has_prism_contract()
17428 || d.down_proj.has_prism_contract()
17429 {
17430 return None;
17431 }
17432 if d.act != Act::Silu {
17434 return None;
17435 }
17436 if d.gate_proj.rows() < crate::gpu::min_rows() && !d.gate_proj.is_q1() {
17439 return None;
17440 }
17441 let mut jobs: Vec<crate::gpu::MoeJob> = Vec::with_capacity(1);
17442 let mut model_ref = None;
17443 moe_push_job(d, x, 1.0, &mut jobs, &mut model_ref)?;
17444 let model = model_ref?;
17445 let hidden = jobs[0].down.1;
17446 let mut out = attention::take_buf(hidden);
17447 if crate::gpu::moe_block(&model, &jobs, &mut out) {
17448 Some(out)
17449 } else {
17450 let mut out = out;
17451 attention::recycle_buf(&mut out);
17452 None
17453 }
17454}
17455
17456#[allow(clippy::type_complexity)]
17461#[allow(clippy::type_complexity)]
17462pub(crate) fn moe_parts(
17463 t: &QTensor,
17464) -> Option<(
17465 &std::sync::Arc<cortiq_core::CmfModel>,
17466 usize,
17467 usize,
17468 usize,
17469 &[f32],
17470 &[f32],
17471 bool,
17472 bool,
17473 bool,
17474)> {
17475 match t {
17476 QTensor::Mapped {
17477 model,
17478 idx,
17479 dtype: dt @ (cortiq_core::TensorDtype::Q8_2f | cortiq_core::TensorDtype::Q8Row),
17480 rows,
17481 cols,
17482 row_scale,
17483 col_field,
17484 ..
17485 } if (*dt == cortiq_core::TensorDtype::Q8Row) || !col_field.is_empty() => Some((
17486 model, *idx, *rows, *cols, row_scale, col_field, false, false, false,
17487 )),
17488 QTensor::Mapped {
17490 model,
17491 idx,
17492 dtype: cortiq_core::TensorDtype::Q1,
17493 rows,
17494 cols,
17495 ..
17496 } => Some((
17497 model,
17498 *idx,
17499 *rows,
17500 *cols,
17501 &[][..],
17502 &[][..],
17503 true,
17504 false,
17505 false,
17506 )),
17507 QTensor::Mapped {
17509 model,
17510 idx,
17511 dtype: cortiq_core::TensorDtype::Q4Tiled,
17512 rows,
17513 cols,
17514 ..
17515 } => Some((
17516 model,
17517 *idx,
17518 *rows,
17519 *cols,
17520 &[][..],
17521 &[][..],
17522 false,
17523 true,
17524 false,
17525 )),
17526 QTensor::Mapped {
17528 model,
17529 idx,
17530 dtype: cortiq_core::TensorDtype::Q4TiledP,
17531 rows,
17532 cols,
17533 ..
17534 } => Some((
17535 model,
17536 *idx,
17537 *rows,
17538 *cols,
17539 &[][..],
17540 &[][..],
17541 false,
17542 true,
17543 false,
17544 )),
17545 QTensor::Mapped {
17549 model,
17550 idx,
17551 dtype: cortiq_core::TensorDtype::Q2TiledP,
17552 rows,
17553 cols,
17554 ..
17555 } => Some((
17556 model,
17557 *idx,
17558 *rows,
17559 *cols,
17560 &[][..],
17561 &[][..],
17562 false,
17563 true,
17564 true,
17565 )),
17566 _ => None,
17567 }
17568}
17569
17570#[cfg(target_os = "macos")]
17578fn metal_moe_graph_parts(m: &MoeFfn, hidden: usize) -> Option<crate::gpu::GpuMoe<'_>> {
17579 if m.router_input_norm
17580 || m.route_tau.is_some()
17581 || m.mask.is_some()
17582 || m.per_expert_scale.is_some()
17583 || m.experts.is_empty()
17584 || m.top_k == 0
17585 || m.resonance.is_some()
17586 {
17587 return None;
17588 }
17589 let (sh, sg) = match &m.shared {
17592 Some((sh, sg)) => (sh, sg.as_ref()),
17593 None => return None,
17594 };
17595 let (rf, rr, rc) = m.router.f32_parts()?;
17596 if rr != m.experts.len() || rc != hidden {
17597 return None;
17598 }
17599 let shared_gated = sg.is_some();
17600 let sf = match sg {
17601 Some(sg) => {
17602 let (sf, sr, sc) = sg.f32_parts()?;
17603 if sr * sc != hidden {
17604 return None;
17605 }
17606 sf
17607 }
17608 None => &rf[..hidden],
17611 };
17612 if let Some(b) = &m.expert_bias {
17613 if b.len() != m.experts.len() {
17614 return None;
17615 }
17616 }
17617 let inter = m.experts[0].gate_proj.rows();
17618 let gu_q2 = m.experts[0].gate_proj.mapped_q2tp().is_some();
17621 let trio = |e: &DenseFfn| -> Option<(usize, usize, usize)> {
17622 if e.act != Act::Silu
17623 || e.gate_proj.rows() != inter
17624 || e.gate_proj.cols() != hidden
17625 || e.up_proj.rows() != inter
17626 || e.up_proj.cols() != hidden
17627 || e.down_proj.rows() != hidden
17628 || e.down_proj.cols() != inter
17629 {
17630 return None;
17631 }
17632 let pick = |t: &QTensor| -> Option<usize> {
17633 if gu_q2 {
17634 t.mapped_q2tp().map(|(_, i)| i)
17635 } else {
17636 t.mapped_q4tp().map(|(_, i)| i)
17637 }
17638 };
17639 Some((
17640 pick(&e.gate_proj)?,
17641 pick(&e.up_proj)?,
17642 e.down_proj.mapped_q4tp().map(|(_, i)| i)?,
17643 ))
17644 };
17645 let experts = m.experts.iter().map(trio).collect::<Option<Vec<_>>>()?;
17646 let shared = trio(sh)?;
17647 Some(crate::gpu::GpuMoe {
17648 router: rf,
17649 sgate: sf,
17650 experts,
17651 shared,
17652 n_exp: m.experts.len(),
17653 top_k: m.top_k,
17654 inter,
17655 norm_topk: m.norm_topk_prob,
17656 route_scale: m.routed_scaling,
17657 gu_q2,
17658 sigmoid: m.router_sigmoid,
17659 bias: m.expert_bias.as_deref(),
17660 shared_gated,
17661 })
17662}
17663
17664pub(crate) fn moe_push_job_parts<'a>(
17668 gate: &'a QTensor,
17669 up: &'a QTensor,
17670 down: &'a QTensor,
17671 x: &[f32],
17672 w: f32,
17673 swiglu_limit: f32,
17674 jobs: &mut Vec<crate::gpu::MoeJob<'a>>,
17675 model_ref: &mut Option<std::sync::Arc<cortiq_core::CmfModel>>,
17676) -> Option<()> {
17677 use crate::qtensor::prescale;
17678 let (gm, gi, gr, gc, grs, gcf, gq1, gq4, gq2) = moe_parts(gate)?;
17679 let (_, ui, ur, uc, urs, ucf, uq1, uq4, uq2) = moe_parts(up)?;
17680 let (_, di, dr, dc, drs, dcf, dq1, dq4, dq2) = moe_parts(down)?;
17681 if gq1 != uq1 || uq1 != dq1 || gq4 != uq4 || uq4 != dq4 || gq2 != uq2 {
17682 return None; }
17684 if gq2 && (dq2 || !dq4 || down.mapped_q4tp().is_none()) {
17687 return None;
17688 }
17689 if !gq2 && dq2 {
17690 return None;
17691 }
17692 model_ref.get_or_insert_with(|| gm.clone());
17693 let dt = |cf: &[f32]| {
17694 if cf.is_empty() {
17695 cortiq_core::TensorDtype::Q8Row
17696 } else {
17697 cortiq_core::TensorDtype::Q8_2f
17698 }
17699 };
17700 jobs.push(crate::gpu::MoeJob {
17701 gate: (gi, gr, gc, grs),
17702 up: (ui, ur, uc, urs),
17703 down: (di, dr, dc, drs),
17704 xs_gate: prescale(x, gcf, dt(gcf)).into_owned(),
17705 xs_up: prescale(x, ucf, dt(ucf)).into_owned(),
17706 down_col: dcf,
17707 w,
17708 q1: gq1,
17709 q4t: gq4 && !gq2 && gate.mapped_q4tp().is_none(),
17710 q4tp: gq4 && (gq2 || gate.mapped_q4tp().is_some()),
17711 gu_q2: gq2,
17712 swiglu_limit,
17713 });
17714 Some(())
17715}
17716
17717fn moe_push_job<'a>(
17719 d: &'a DenseFfn,
17720 x: &[f32],
17721 w: f32,
17722 jobs: &mut Vec<crate::gpu::MoeJob<'a>>,
17723 model_ref: &mut Option<std::sync::Arc<cortiq_core::CmfModel>>,
17724) -> Option<()> {
17725 use crate::qtensor::prescale;
17726 if d.act != Act::Silu {
17727 return None; }
17729 let (gm, gi, gr, gc, grs, gcf, gq1, gq4, gq2) = moe_parts(&d.gate_proj)?;
17730 let (_, ui, ur, uc, urs, ucf, uq1, uq4, uq2) = moe_parts(&d.up_proj)?;
17731 let (_, di, dr, dc, drs, dcf, dq1, dq4, dq2) = moe_parts(&d.down_proj)?;
17732 if gq1 != uq1 || uq1 != dq1 || gq4 != uq4 || uq4 != dq4 || gq2 != uq2 {
17733 return None; }
17735 if gq2 && (dq2 || !dq4 || d.down_proj.mapped_q4tp().is_none()) {
17736 return None;
17737 }
17738 if !gq2 && dq2 {
17739 return None;
17740 }
17741 model_ref.get_or_insert_with(|| gm.clone());
17742 let gdt = if gcf.is_empty() {
17743 cortiq_core::TensorDtype::Q8Row
17744 } else {
17745 cortiq_core::TensorDtype::Q8_2f
17746 };
17747 let udt = if ucf.is_empty() {
17748 cortiq_core::TensorDtype::Q8Row
17749 } else {
17750 cortiq_core::TensorDtype::Q8_2f
17751 };
17752 jobs.push(crate::gpu::MoeJob {
17753 gate: (gi, gr, gc, grs),
17754 up: (ui, ur, uc, urs),
17755 down: (di, dr, dc, drs),
17756 xs_gate: prescale(x, gcf, gdt).into_owned(),
17757 xs_up: prescale(x, ucf, udt).into_owned(),
17758 down_col: dcf,
17759 w,
17760 q1: gq1,
17761 q4t: gq4 && !gq2 && d.gate_proj.mapped_q4tp().is_none(),
17762 q4tp: gq4 && (gq2 || d.gate_proj.mapped_q4tp().is_some()),
17763 gu_q2: gq2,
17764 swiglu_limit: 0.0,
17765 });
17766 Some(())
17767}
17768
17769fn sparse_ffn_quant(
17776 d: &DenseFfn,
17777 x: &[f32],
17778 active: &[u16],
17779 hidden: usize,
17780 pool: Option<&Pool>,
17781) -> Vec<f32> {
17782 let n = active.len();
17783 let inter = d.gate_proj.rows();
17784 let mut act = vec![0.0f32; n];
17785 let need_scratch = !(d.gate_proj.sparse_col_ok() && d.up_proj.sparse_col_ok());
17788 let compute = |ai: usize| -> f32 {
17789 let idx = active[ai] as usize;
17790 if idx >= inter {
17791 return 0.0; }
17793 let mut s = if need_scratch {
17794 vec![0.0f32; hidden]
17795 } else {
17796 Vec::new()
17797 };
17798 let gate = d.gate_proj.row_dot(idx, x, &mut s);
17799 let up = d.up_proj.row_dot(idx, x, &mut s);
17800 d.act.combine(gate, up)
17801 };
17802 match pool {
17803 Some(p) if n >= 256 => {
17804 let ptr = SendMut(act.as_mut_ptr());
17805 p.run(&|widx, nw| {
17806 let chunk = n.div_ceil(nw);
17807 let (s, e) = (widx * chunk, ((widx + 1) * chunk).min(n));
17808 for ai in s..e {
17809 unsafe { *ptr.at(ai) = compute(ai) };
17810 }
17811 });
17812 }
17813 _ => {
17814 for (ai, a) in act.iter_mut().enumerate() {
17815 *a = compute(ai);
17816 }
17817 }
17818 }
17819 let mut out = vec![0.0f32; hidden];
17821 for (ai, &idx) in active.iter().enumerate() {
17822 let w = act[ai];
17823 if w.abs() >= 1e-12 && (idx as usize) < inter {
17824 d.down_proj.add_col_scaled(idx as usize, w, &mut out);
17825 }
17826 }
17827 out
17828}
17829
17830#[doc(hidden)]
17832pub fn sparse_ffn_quant_for_test(
17833 d: &DenseFfn,
17834 x: &[f32],
17835 active: &[u16],
17836 hidden: usize,
17837) -> Vec<f32> {
17838 sparse_ffn_quant(d, x, active, hidden, None)
17839}
17840
17841fn dequant_dense_f32(d: &DenseFfn) -> (Vec<f32>, Vec<f32>, Vec<f32>) {
17845 let deq = |t: &QTensor| -> Vec<f32> {
17846 let (rows, cols) = (t.rows(), t.cols());
17847 let mut out = vec![0.0f32; rows * cols];
17848 for r in 0..rows {
17849 t.row_f32(r, &mut out[r * cols..(r + 1) * cols]);
17850 }
17851 out
17852 };
17853 (deq(&d.gate_proj), deq(&d.up_proj), deq(&d.down_proj))
17854}
17855
17856struct SendMut(*mut f32);
17858unsafe impl Send for SendMut {}
17859unsafe impl Sync for SendMut {}
17860impl SendMut {
17861 #[inline]
17862 #[allow(clippy::mut_from_ref)]
17865 unsafe fn at(&self, i: usize) -> &mut f32 {
17866 unsafe { &mut *self.0.add(i) }
17867 }
17868}
17869
17870pub(crate) fn moe_route(
17883 logits: &[f32],
17884 m: &MoeFfn,
17885 allowed: Option<&[bool]>,
17886) -> (Vec<usize>, Vec<f32>, f32) {
17887 moe_route_with_eps(logits, m, allowed, 1e-6)
17888}
17889
17890pub(crate) fn moe_route_with_eps(
17899 logits: &[f32],
17900 m: &MoeFfn,
17901 allowed: Option<&[bool]>,
17902 sigmoid_denom_eps: f32,
17903) -> (Vec<usize>, Vec<f32>, f32) {
17904 let ne = logits.len();
17905 let admit = |e: usize| {
17911 m.mask.as_ref().is_none_or(|mk| mk[e])
17912 && allowed.is_none_or(|a| a.get(e).copied().unwrap_or(false))
17913 };
17914 if m.resonance.is_some() && m.top_k == 1 {
17927 let mut best: Option<usize> = None;
17928 for e in (0..ne).filter(|&e| admit(e)) {
17929 let l = logits[e];
17930 if l == f32::NEG_INFINITY || l.is_nan() {
17931 continue;
17932 }
17933 if best.is_none_or(|b| l > logits[b]) {
17934 best = Some(e);
17935 }
17936 }
17937 if let Some(b) = best {
17938 let mut p = vec![0.0f32; ne];
17939 p[b] = 1.0;
17940 return (vec![b], p, 1.0 / m.routed_scaling);
17941 }
17942 }
17943 let p: Vec<f32> = if m.router_sigmoid {
17949 logits.iter().map(|&l| 1.0 / (1.0 + (-l).exp())).collect()
17950 } else {
17951 let mx = logits.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
17952 if mx == f32::NEG_INFINITY {
17953 vec![1.0 / ne.max(1) as f32; ne]
17954 } else {
17955 let mut e: Vec<f32> = logits.iter().map(|&l| (l - mx).exp()).collect();
17956 let s: f32 = e.iter().sum();
17957 for v in &mut e {
17958 *v /= s;
17959 }
17960 e
17961 }
17962 };
17963 let mut idx: Vec<usize> = (0..ne).filter(|&e| admit(e)).collect();
17964 match &m.expert_bias {
17966 Some(b) => idx.sort_unstable_by(|&x, &y| {
17967 (p[y] + b[y])
17968 .partial_cmp(&(p[x] + b[x]))
17969 .unwrap()
17970 .then(x.cmp(&y))
17971 }),
17972 None => idx.sort_unstable_by(|&x, &y| p[y].partial_cmp(&p[x]).unwrap().then(x.cmp(&y))),
17973 }
17974 idx.truncate(m.top_k);
17975 if let Some(tau) = m.route_tau {
17979 let total: f32 = idx.iter().map(|&e| p[e]).sum();
17980 if total > 0.0 {
17981 let mut acc = 0.0f32;
17982 let mut keep = idx.len();
17983 for (i, &e) in idx.iter().enumerate() {
17984 acc += p[e];
17985 if acc >= tau * total {
17986 keep = i + 1;
17987 break;
17988 }
17989 }
17990 idx.truncate(keep);
17991 }
17992 }
17993 let wsum: f32 = if m.norm_topk_prob {
17994 let s: f32 = idx.iter().map(|&e| p[e]).sum();
17995 (if m.router_sigmoid {
17999 s + sigmoid_denom_eps
18000 } else {
18001 s
18002 }) / m.routed_scaling
18003 } else {
18004 1.0 / m.routed_scaling
18005 };
18006 (idx, p, wsum)
18007}
18008
18009fn moe_trace(idx: &[usize]) {
18011 moe_trace_at(crate::gpu::cur_layer() as i32, idx)
18012}
18013
18014pub(crate) fn moe_trace_at(li: i32, idx: &[usize]) {
18017 use std::io::Write;
18018 static F: std::sync::OnceLock<Option<std::sync::Mutex<std::fs::File>>> =
18019 std::sync::OnceLock::new();
18020 let Some(f) = F.get_or_init(|| {
18021 let p = std::env::var("CMF_MOE_TRACE").ok()?;
18022 Some(std::sync::Mutex::new(
18023 std::fs::OpenOptions::new()
18024 .create(true)
18025 .append(true)
18026 .open(p)
18027 .ok()?,
18028 ))
18029 }) else {
18030 return;
18031 };
18032 let ids: Vec<String> = idx.iter().map(|e| e.to_string()).collect();
18033 let _ = writeln!(f.lock().unwrap(), "{li}:{}", ids.join(","));
18034}
18035
18036pub(crate) fn moe_ffn(
18039 m: &MoeFfn,
18040 x: &[f32],
18041 pool: Option<&Pool>,
18042 allowed: Option<&[bool]>,
18043) -> Vec<f32> {
18044 let r = moe_ffn_route(m, x, pool, allowed);
18045 moe_ffn_experts(m, x, &r, pool)
18046}
18047
18048pub(crate) struct MoeRoute {
18052 pub idx: Vec<usize>,
18053 pub p: Vec<f32>,
18054 pub wsum: f32,
18055 pub logits: Vec<f32>,
18056}
18057
18058pub(crate) fn moe_ffn_route(
18064 m: &MoeFfn,
18065 x: &[f32],
18066 pool: Option<&Pool>,
18067 allowed: Option<&[bool]>,
18068) -> MoeRoute {
18069 accumulate_act(m, x, 1);
18070 let ne = m.experts.len();
18071 let mut logits = vec![0.0f32; ne];
18072 match &m.resonance {
18073 Some(r) => r.scores(x, &mut logits),
18074 None => m.router.matvec(x, &mut logits, pool),
18075 }
18076 let (idx, p, wsum) = moe_route(&logits, m, allowed);
18077 {
18078 let mut st = m.stats.borrow_mut();
18079 if st.len() < ne {
18080 st.resize(ne, 0);
18081 }
18082 for &e in &idx {
18083 st[e] += 1;
18084 }
18085 }
18086 moe_trace(&idx);
18092 MoeRoute {
18093 idx,
18094 p,
18095 wsum,
18096 logits,
18097 }
18098}
18099
18100pub(crate) fn moe_ffn_experts(
18103 m: &MoeFfn,
18104 x: &[f32],
18105 r: &MoeRoute,
18106 pool: Option<&Pool>,
18107) -> Vec<f32> {
18108 let (idx, p, wsum) = (&r.idx, &r.p, r.wsum);
18109 if crate::gpu::enabled_here() {
18114 match crate::gpu::probe_arm(crate::gpu::OpClass::Ffn) {
18115 crate::gpu::ProbeArm::Gpu => {
18116 let t0 = std::time::Instant::now();
18117 if let Some(out) = moe_ffn_gpu(m, x, &idx, &p, wsum, pool) {
18118 crate::gpu::probe_record(crate::gpu::OpClass::Ffn, true, t0.elapsed());
18119 return out;
18120 }
18121 }
18122 crate::gpu::ProbeArm::CpuTimed => {
18123 let t0 = std::time::Instant::now();
18124 let out = crate::gpu::cpu_scope(|| moe_ffn_cpu(m, x, &idx, &p, wsum, pool));
18125 crate::gpu::probe_record(crate::gpu::OpClass::Ffn, false, t0.elapsed());
18126 return out;
18127 }
18128 crate::gpu::ProbeArm::Cpu => {
18129 return crate::gpu::cpu_scope(|| moe_ffn_cpu(m, x, &idx, &p, wsum, pool));
18130 }
18131 }
18132 }
18133 moe_ffn_cpu(m, x, &idx, &p, wsum, pool)
18134}
18135
18136fn moe_ffn_banked(
18140 slot: &mut crate::mimo_moe::Slot,
18141 li: usize,
18142 m: &MoeFfn,
18143 x: &[f32],
18144 pool: Option<&Pool>,
18145) -> Vec<f32> {
18146 let t0 = std::time::Instant::now();
18147 let r = moe_ffn_route(m, x, pool, None);
18148 slot.note_route(t0.elapsed().as_nanos() as u64);
18149 match slot.forward(li, m, x, &r, pool) {
18150 Some(out) => out,
18151 None => crate::qtensor::float_activations_scope(|| {
18152 crate::gpu::cpu_scope(|| moe_ffn_experts(m, x, &r, pool))
18153 }),
18154 }
18155}
18156
18157fn moe_ffn_banked_rows(
18159 slot: &mut crate::mimo_moe::Slot,
18160 li: usize,
18161 m: &MoeFfn,
18162 xs: &[f32],
18163 b: usize,
18164 hidden: usize,
18165 pool: Option<&Pool>,
18166) -> Vec<f32> {
18167 let t0 = std::time::Instant::now();
18168 let routes: Vec<_> = xs
18169 .chunks_exact(hidden)
18170 .map(|x| moe_ffn_route(m, x, pool, None))
18171 .collect();
18172 slot.note_route(t0.elapsed().as_nanos() as u64);
18173 if let Some(out) = slot.forward_rows(li, m, xs, &routes, pool) {
18174 return out;
18175 }
18176 let mut out = Vec::with_capacity(b * hidden);
18177 for (x, r) in xs.chunks_exact(hidden).zip(&routes) {
18178 let row = slot.forward(li, m, x, r, pool).unwrap_or_else(|| {
18179 crate::qtensor::float_activations_scope(|| {
18181 crate::gpu::cpu_scope(|| moe_ffn_experts(m, x, r, pool))
18182 })
18183 });
18184 out.extend(row);
18185 }
18186 out
18187}
18188
18189fn graph_note(built: bool, layers_run: usize, total_layers: usize) {
18194 use std::sync::atomic::{AtomicBool, Ordering};
18195 if built {
18196 GRAPH_TOK_OK.fetch_add(1, Ordering::Relaxed);
18197 if total_layers > 0 && layers_run < total_layers {
18198 GRAPH_TOK_PREFIX.fetch_add(1, Ordering::Relaxed);
18199 } else {
18200 GRAPH_TOK_FULL.fetch_add(1, Ordering::Relaxed);
18201 }
18202 } else {
18203 GRAPH_TOK_MISS.fetch_add(1, Ordering::Relaxed);
18204 }
18205 static SAID: AtomicBool = AtomicBool::new(false);
18206 if !SAID.swap(true, Ordering::Relaxed) {
18207 if built {
18208 tracing::info!("wgpu whole-token graph: ACTIVE");
18209 } else {
18210 tracing::warn!("wgpu whole-token graph refused — per-op path");
18211 }
18212 }
18213}
18214
18215pub static GRAPH_TOK_OK: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
18219pub static GRAPH_TOK_MISS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
18220pub static GRAPH_TOK_PREFIX: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
18224pub static GRAPH_TOK_FULL: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
18226
18227pub static METAL_GRAPH_TOK_OK: std::sync::atomic::AtomicU64 =
18231 std::sync::atomic::AtomicU64::new(0);
18232pub static METAL_GRAPH_HEAD_OK: std::sync::atomic::AtomicU64 =
18233 std::sync::atomic::AtomicU64::new(0);
18234pub static METAL_GRAPH_HEAD_MISS: std::sync::atomic::AtomicU64 =
18235 std::sync::atomic::AtomicU64::new(0);
18236pub static METAL_GRAPH_LAYERS: std::sync::atomic::AtomicU64 =
18237 std::sync::atomic::AtomicU64::new(0);
18238pub static METAL_GRAPH_ERRORS: std::sync::atomic::AtomicU64 =
18239 std::sync::atomic::AtomicU64::new(0);
18240pub static METAL_PREFILL_CHUNKS: std::sync::atomic::AtomicU64 =
18244 std::sync::atomic::AtomicU64::new(0);
18245pub static METAL_PREFILL_ROWS: std::sync::atomic::AtomicU64 =
18246 std::sync::atomic::AtomicU64::new(0);
18247pub static METAL_PREFILL_HEAD_ROWS: std::sync::atomic::AtomicU64 =
18248 std::sync::atomic::AtomicU64::new(0);
18249pub static METAL_PREFILL_ERRORS: std::sync::atomic::AtomicU64 =
18250 std::sync::atomic::AtomicU64::new(0);
18251
18252fn moe_batch_enabled() -> bool {
18255 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
18256 *ON.get_or_init(|| std::env::var("CMF_MOE_BATCH").as_deref() != Ok("0"))
18257}
18258
18259fn moe_ffn_cpu_batched(
18265 m: &MoeFfn,
18266 x: &[f32],
18267 idx: &[usize],
18268 p: &[f32],
18269 wsum: f32,
18270 pool: Option<&Pool>,
18271) -> Option<Vec<f32>> {
18272 if idx.is_empty() || !moe_batch_enabled() {
18273 return None;
18274 }
18275 if FFN_PROBE.with(|pr| pr.borrow().is_some()) {
18279 return None;
18280 }
18281 let n = idx.len() + usize::from(m.shared.is_some());
18282 let mut pairs = Vec::with_capacity(n);
18283 let mut downs = Vec::with_capacity(n);
18284 let mut ws = Vec::with_capacity(n);
18285 for &e in idx {
18286 let d = &m.experts[e];
18287 if d.act != Act::Silu {
18288 return None;
18289 }
18290 pairs.push((&d.gate_proj, &d.up_proj));
18291 downs.push(&d.down_proj);
18292 ws.push(p[e] / wsum * m.per_expert_scale.as_ref().map_or(1.0, |v| v[e]));
18293 }
18294 if let Some((se, gate)) = &m.shared {
18297 if se.act != Act::Silu {
18298 return None;
18299 }
18300 let g = gate.as_ref().map_or(1.0, |gate| {
18301 let mut gl = [0.0f32; 1];
18302 gate.matvec(x, &mut gl, pool);
18303 1.0 / (1.0 + (-gl[0]).exp())
18304 });
18305 pairs.push((&se.gate_proj, &se.up_proj));
18306 downs.push(&se.down_proj);
18307 ws.push(g);
18308 }
18309 let inter = pairs[0].0.rows();
18310 let mut gs: Vec<Vec<f32>> = (0..pairs.len()).map(|_| vec![0f32; inter]).collect();
18311 if !QTensor::moe_gate_up_many(&pairs, x, &mut gs, pool) {
18312 return None;
18313 }
18314 let mut out = attention::take_buf(x.len());
18315 if !QTensor::moe_down_many(&downs, &gs, &ws, &mut out, pool) {
18316 attention::recycle_buf(&mut out);
18317 return None;
18318 }
18319 Some(out)
18320}
18321
18322pub(crate) fn moe_cold_experts_cpu(
18328 experts: &[(&DenseFfn, f32)],
18329 x: &[f32],
18330 pool: Option<&Pool>,
18331) -> Vec<f32> {
18332 let mut out = attention::take_buf(x.len());
18333 if experts.is_empty() {
18334 return out;
18335 }
18336 let pairs: Vec<_> = experts
18337 .iter()
18338 .map(|(e, _)| (&e.gate_proj, &e.up_proj))
18339 .collect();
18340 let downs: Vec<_> = experts.iter().map(|(e, _)| &e.down_proj).collect();
18341 let weights: Vec<_> = experts.iter().map(|(_, w)| *w).collect();
18342 let inter = experts[0].0.gate_proj.rows();
18343 let mut activations: Vec<Vec<f32>> = (0..experts.len()).map(|_| vec![0.0; inter]).collect();
18344 if QTensor::moe_gate_up_many(&pairs, x, &mut activations, pool)
18345 && QTensor::moe_down_many(&downs, &activations, &weights, &mut out, pool)
18346 {
18347 return out;
18348 }
18349 out.fill(0.0);
18350 for &(expert, weight) in experts {
18351 let mut one = dense_ffn(expert, x, pool);
18352 for (o, v) in out.iter_mut().zip(&one) {
18353 *o += weight * v;
18354 }
18355 attention::recycle_buf(&mut one);
18356 }
18357 out
18358}
18359
18360pub(crate) fn moe_cold_experts_rows_cpu(
18364 jobs: &[Vec<(&DenseFfn, f32)>],
18365 xs: &[f32],
18366 hidden: usize,
18367 pool: Option<&Pool>,
18368) -> Vec<f32> {
18369 let mut out = vec![0.0; xs.len()];
18370 let mut experts: Vec<&DenseFfn> = Vec::new();
18371 let mut groups: Vec<Vec<usize>> = Vec::new();
18372 let mut terms = vec![Vec::new(); jobs.len()];
18373 for (r, row) in jobs.iter().enumerate() {
18374 for &(e, w) in row {
18375 let g = match experts.iter().position(|&d| std::ptr::eq(d, e)) {
18376 Some(g) => g,
18377 None => {
18378 experts.push(e);
18379 groups.push(Vec::new());
18380 groups.len() - 1
18381 }
18382 };
18383 terms[r].push((g, groups[g].len(), w));
18384 groups[g].push(r);
18385 }
18386 }
18387 if experts.is_empty() {
18388 return out;
18389 }
18390 let pairs: Vec<_> = experts.iter().map(|e| (&e.gate_proj, &e.up_proj)).collect();
18391 let downs: Vec<_> = experts.iter().map(|e| &e.down_proj).collect();
18392 let lens: Vec<_> = groups.iter().map(Vec::len).collect();
18393 let count: usize = lens.iter().sum();
18394 let mut acts = vec![vec![0.0; experts[0].gate_proj.rows()]; count];
18395 let mut ds = vec![vec![0.0; hidden]; count];
18396 if QTensor::moe_gate_up_rows(&pairs, &groups, xs, &mut acts, pool)
18397 && QTensor::moe_down_rows(&downs, &lens, &acts, &mut ds, pool)
18398 {
18399 let mut offset = 0;
18400 let offsets: Vec<_> = lens
18401 .iter()
18402 .map(|&n| {
18403 let start = offset;
18404 offset += n;
18405 start
18406 })
18407 .collect();
18408 for (r, terms) in terms.iter().enumerate() {
18409 for &(g, slot, w) in terms {
18410 for (o, &v) in out[r * hidden..(r + 1) * hidden]
18411 .iter_mut()
18412 .zip(&ds[offsets[g] + slot])
18413 {
18414 *o += w * v;
18415 }
18416 }
18417 }
18418 } else {
18419 for (r, jobs) in jobs.iter().enumerate() {
18420 if !jobs.is_empty() {
18421 let mut row = moe_cold_experts_cpu(jobs, &xs[r * hidden..(r + 1) * hidden], pool);
18422 out[r * hidden..(r + 1) * hidden].copy_from_slice(&row);
18423 attention::recycle_buf(&mut row);
18424 }
18425 }
18426 }
18427 out
18428}
18429
18430fn moe_ffn_cpu(
18432 m: &MoeFfn,
18433 x: &[f32],
18434 idx: &[usize],
18435 p: &[f32],
18436 wsum: f32,
18437 pool: Option<&Pool>,
18438) -> Vec<f32> {
18439 if let Some(out) = moe_ffn_cpu_batched(m, x, idx, p, wsum, pool) {
18440 return out;
18441 }
18442 let mut out = attention::take_buf(x.len());
18443 for &e in idx {
18444 let mut eo = dense_ffn(&m.experts[e], x, pool);
18445 let w = p[e] / wsum * m.per_expert_scale.as_ref().map_or(1.0, |v| v[e]);
18446 for i in 0..out.len() {
18447 out[i] += w * eo[i];
18448 }
18449 attention::recycle_buf(&mut eo);
18450 }
18451 if let Some((se, gate)) = &m.shared {
18452 let mut so = dense_ffn(se, x, pool);
18453 let g = gate.as_ref().map_or(1.0, |gate| {
18454 let mut gl = [0.0f32; 1];
18455 gate.matvec(x, &mut gl, pool);
18456 1.0 / (1.0 + (-gl[0]).exp())
18457 });
18458 for i in 0..out.len() {
18459 out[i] += g * so[i];
18460 }
18461 attention::recycle_buf(&mut so);
18462 }
18463 out
18464}
18465
18466#[allow(clippy::too_many_arguments)]
18474pub(crate) fn mla_attention(
18475 w: &MlaWeights,
18476 normed: &[f32],
18477 cache: &mut crate::kv_cache::LayerKvCache,
18478 position: usize,
18479 inv_freq: &[f32],
18480 rope_scale: f32,
18481 eps: f64,
18482 pool: Option<&Pool>,
18483) -> Vec<f32> {
18484 let (nh, dr, dn, dv, lora) = (w.nh, w.qk_rope, w.qk_nope, w.v_dim, w.lora);
18485 let hd = dr + dn;
18486 let mut q = vec![0.0f32; nh * hd];
18487 match (&w.q_a, &w.q_a_norm) {
18488 (Some(qa), Some(qn)) => {
18489 let mut t = vec![0.0f32; qa.rows()];
18490 qa.matvec(normed, &mut t, pool);
18491 let tn = inference::rms_norm(&t, qn, eps, NormStyle::Qwen);
18492 w.q_proj.matvec(&tn, &mut q, pool);
18493 }
18494 _ => w.q_proj.matvec(normed, &mut q, pool),
18495 }
18496 let mut ca = vec![0.0f32; lora + dr];
18497 w.kv_a.matvec(normed, &mut ca, pool);
18498 let (c_lat, k_rope) = ca.split_at_mut(lora);
18499 let latn = inference::rms_norm(c_lat, &w.kv_a_norm, eps, NormStyle::Qwen);
18500 let mut kvb = vec![0.0f32; nh * (dn + dv)];
18501 w.kv_b.matvec(&latn, &mut kvb, pool);
18502 if !w.nope {
18503 attention::rope_rotate_scaled(k_rope, position, inv_freq, rope_scale);
18504 }
18505 for h in 0..nh {
18506 if !w.nope {
18507 attention::rope_rotate_scaled(
18508 &mut q[h * hd..h * hd + dr],
18509 position,
18510 inv_freq,
18511 rope_scale,
18512 );
18513 }
18514 }
18515 let mut k = vec![0.0f32; nh * hd];
18516 let mut v = vec![0.0f32; nh * hd];
18517 for h in 0..nh {
18518 k[h * hd..h * hd + dr].copy_from_slice(k_rope);
18519 k[h * hd + dr..(h + 1) * hd].copy_from_slice(&kvb[h * (dn + dv)..h * (dn + dv) + dn]);
18520 v[h * hd..h * hd + dv].copy_from_slice(&kvb[h * (dn + dv) + dn..(h + 1) * (dn + dv)]);
18521 }
18522 cache.append(&k, &v, &vec![true; nh]);
18523 let (ao, mut imp) = attention::attend_all_heads(&q, cache, nh, 1, hd, w.scale, None, 0.0);
18524 attention::recycle_buf(&mut imp);
18525 let mut ov = vec![0.0f32; nh * dv];
18526 for h in 0..nh {
18527 ov[h * dv..(h + 1) * dv].copy_from_slice(&ao[h * hd..h * hd + dv]);
18528 }
18529 let mut out = vec![0.0f32; w.o_proj.rows()];
18530 w.o_proj.matvec(&ov, &mut out, pool);
18531 out
18532}
18533
18534fn dense_moe_ffn(
18541 dm: &DenseMoeFfn,
18542 x_normed: &[f32],
18543 h_raw: &[f32],
18544 eps: f64,
18545 norm_style: NormStyle,
18546 pool: Option<&Pool>,
18547) -> Vec<f32> {
18548 let mut d = dense_ffn(&dm.dense, x_normed, pool);
18549 d = inference::rms_norm(&d, &dm.post_norm_1, eps, norm_style);
18550 let m = &dm.moe;
18551 let ne = m.experts.len();
18552 let mut logits = vec![0.0f32; ne];
18553 if m.router_input_norm {
18554 let ss: f32 = h_raw.iter().map(|v| v * v).sum::<f32>() / h_raw.len() as f32;
18555 let inv = 1.0 / (ss + eps as f32).sqrt();
18556 let xr: Vec<f32> = h_raw.iter().map(|v| v * inv).collect();
18557 m.router.matvec(&xr, &mut logits, pool);
18558 } else {
18559 m.router.matvec(h_raw, &mut logits, pool);
18560 }
18561 let (idx, p, wsum) = moe_route(&logits, m, None);
18562 {
18563 let mut st = m.stats.borrow_mut();
18564 if st.len() < ne {
18565 st.resize(ne, 0);
18566 }
18567 for &e in &idx {
18568 st[e] += 1;
18569 }
18570 }
18571 let x2 = inference::rms_norm(h_raw, &dm.pre_norm_2, eps, norm_style);
18572 let mo = moe_ffn_cpu(m, &x2, &idx, &p, wsum, pool);
18573 let mo = inference::rms_norm(&mo, &dm.post_norm_2, eps, norm_style);
18574 for (di, mi) in d.iter_mut().zip(&mo) {
18575 *di += mi;
18576 }
18577 d
18578}
18579
18580fn moe_gpu_refused(why: &'static str) {
18587 use std::sync::atomic::{AtomicBool, Ordering};
18588 static SAID: AtomicBool = AtomicBool::new(false);
18589 if !SAID.swap(true, Ordering::Relaxed) {
18590 tracing::warn!("MoE GPU block refused ({why}) — experts run on the CPU");
18591 }
18592}
18593
18594fn moe_ffn_gpu(
18595 m: &MoeFfn,
18596 x: &[f32],
18597 idx: &[usize],
18598 p: &[f32],
18599 wsum: f32,
18600 pool: Option<&Pool>,
18601) -> Option<Vec<f32>> {
18602 use crate::gpu::MoeJob;
18603
18604 let mut jobs: Vec<MoeJob> = Vec::with_capacity(idx.len() + 1);
18605 let mut model_ref = None;
18606 for &e in idx {
18607 if moe_push_job(&m.experts[e], x, p[e] / wsum, &mut jobs, &mut model_ref).is_none() {
18608 moe_gpu_refused("push_job(expert)");
18609 return None;
18610 }
18611 }
18612 if let Some((se, gate)) = &m.shared {
18613 let g = gate.as_ref().map_or(1.0, |gate| {
18614 let mut gl = [0.0f32; 1];
18615 gate.matvec(x, &mut gl, pool);
18616 1.0 / (1.0 + (-gl[0]).exp())
18617 });
18618 if moe_push_job(se, x, g, &mut jobs, &mut model_ref).is_none() {
18619 moe_gpu_refused("push_job(shared)");
18620 return None;
18621 }
18622 }
18623 let Some(model) = model_ref else {
18624 moe_gpu_refused("no model_ref");
18625 return None;
18626 };
18627 let hidden = jobs[0].down.1;
18628 let mut out = vec![0.0f32; hidden];
18629 if crate::gpu::moe_block(&model, &jobs, &mut out) {
18630 Some(out)
18631 } else {
18632 moe_gpu_refused("gpu::moe_block");
18633 None
18634 }
18635}
18636
18637fn ffn_forward(
18639 ffn: &FfnKind,
18640 x: &[f32],
18641 pool: Option<&Pool>,
18642 experts_allowed: Option<&[bool]>,
18643) -> Vec<f32> {
18644 match ffn {
18645 FfnKind::Dense(d) if !d.segs.is_empty() => tube_ffn(d, x, 1, pool, None),
18646 FfnKind::Dense(d) => dense_ffn(d, x, pool),
18647 FfnKind::Moe(m) => moe_ffn(m, x, pool, experts_allowed),
18648 FfnKind::DenseMoe(_) => unreachable!("DenseMoe dispatches via dense_moe_ffn"),
18652 }
18653}
18654
18655fn ffn_forward_pair(
18659 ffn: &FfnKind,
18660 x1: &[f32],
18661 x2: &[f32],
18662 pool: Option<&Pool>,
18663 experts_allowed: Option<&[bool]>,
18664) -> (Vec<f32>, Vec<f32>) {
18665 let d = match ffn {
18666 FfnKind::Dense(d) if !d.segs.is_empty() => {
18669 return (
18670 tube_ffn(d, x1, 1, pool, None),
18671 tube_ffn(d, x2, 1, pool, None),
18672 );
18673 }
18674 FfnKind::Dense(d) => d,
18675 FfnKind::Moe(m) => {
18676 return (
18677 moe_ffn(m, x1, pool, experts_allowed),
18678 moe_ffn(m, x2, pool, experts_allowed),
18679 );
18680 }
18681 FfnKind::DenseMoe(_) => unreachable!("DenseMoe dispatches via dense_moe_ffn"),
18682 };
18683 let inter = d.gate_proj.rows();
18684 FFN_SCRATCH.with(|s| {
18685 let mut s = s.borrow_mut();
18686 let [g1, g2, u1, u2] = &mut *s;
18687 g1.resize(inter, 0.0);
18688 g2.resize(inter, 0.0);
18689 u1.resize(inter, 0.0);
18690 u2.resize(inter, 0.0);
18691 QTensor::matvec2_many(
18694 [&d.gate_proj, &d.up_proj],
18695 x1,
18696 x2,
18697 [g1.as_mut_slice(), u1.as_mut_slice()],
18698 [g2.as_mut_slice(), u2.as_mut_slice()],
18699 pool,
18700 );
18701 for i in 0..inter {
18702 g1[i] = d.act.combine(g1[i], u1[i]);
18703 g2[i] = d.act.combine(g2[i], u2[i]);
18704 }
18705 let mut o1 = attention::take_buf(d.down_proj.rows());
18706 let mut o2 = attention::take_buf(d.down_proj.rows());
18707 d.down_proj.matvec2(g1, g2, &mut o1, &mut o2, pool);
18708 (o1, o2)
18709 })
18710}
18711
18712#[cfg(test)]
18713mod tests {
18714
18715 #[test]
18720 fn prefill_chunk_rule_widens_only_dense_on_discrete() {
18721 use super::{
18722 prefill_chunk_rule, ChunkHost, ChunkStackFacts, DISCRETE_DENSE_PREFILL_CHUNK,
18723 };
18724 let dense_card = ChunkStackFacts {
18725 plain_dense: true,
18726 discrete: true,
18727 gpu_on: true,
18728 ..Default::default()
18729 };
18730 assert!(dense_card.dense_on_discrete());
18731 assert_eq!(
18733 prefill_chunk_rule(None, ChunkHost::Other, dense_card.dense_on_discrete()),
18734 DISCRETE_DENSE_PREFILL_CHUNK
18735 );
18736 assert!(DISCRETE_DENSE_PREFILL_CHUNK > 48);
18737 for (label, facts) in [
18738 ("GDN hybrid / MoE / DeepSeek stack", ChunkStackFacts { plain_dense: false, ..dense_card }),
18739 ("integrated GPU", ChunkStackFacts { discrete: false, ..dense_card }),
18740 ("CPU only", ChunkStackFacts { gpu_on: false, discrete: false, ..dense_card }),
18741 ("capacity split", ChunkStackFacts { capacity_split: true, ..dense_card }),
18742 ("multi-GPU plan", ChunkStackFacts { multi_gpu: true, ..dense_card }),
18743 ("O(1) layers", ChunkStackFacts { o1: true, ..dense_card }),
18744 ] {
18745 assert!(!facts.dense_on_discrete(), "{label}");
18746 assert_eq!(
18747 prefill_chunk_rule(None, ChunkHost::Other, facts.dense_on_discrete()),
18748 48,
18749 "{label} keeps the historical x86 chunk"
18750 );
18751 }
18752 for dense in [false, true] {
18754 assert_eq!(prefill_chunk_rule(None, ChunkHost::Macos, dense), 512);
18755 assert_eq!(prefill_chunk_rule(None, ChunkHost::Aarch64, dense), 256);
18756 }
18757 for host in [ChunkHost::Macos, ChunkHost::Aarch64, ChunkHost::Other] {
18759 for dense in [false, true] {
18760 assert_eq!(prefill_chunk_rule(Some(48), host, dense), 48);
18761 assert_eq!(prefill_chunk_rule(Some(0), host, dense), 1);
18762 }
18763 }
18764 }
18765
18766 #[test]
18767 fn kv_reuse_plan_pulls_rows_decode_wrote_only_on_the_device() {
18768 use super::{ReuseLayer, ReusePlan, kv_reuse_plan};
18769 let full = |host_rows, device_rows| ReuseLayer {
18770 full: true,
18771 host_rows,
18772 device_rows,
18773 device_state: false,
18774 };
18775 assert_eq!(
18778 kv_reuse_plan(339, &[full(300, Some(339)), full(300, Some(339))]),
18779 ReusePlan::Pull(vec![(0, 300, 339), (1, 300, 339)])
18780 );
18781 assert_eq!(kv_reuse_plan(339, &[full(339, None)]), ReusePlan::Ready);
18783 assert_eq!(kv_reuse_plan(339, &[full(339, Some(345))]), ReusePlan::Ready);
18785 assert_eq!(
18787 kv_reuse_plan(339, &[full(300, Some(339)), full(339, None)]),
18788 ReusePlan::Pull(vec![(0, 300, 339)])
18789 );
18790 assert_eq!(kv_reuse_plan(339, &[full(300, Some(320))]), ReusePlan::Fresh);
18792 assert_eq!(kv_reuse_plan(339, &[full(300, None)]), ReusePlan::Fresh);
18793 assert_eq!(kv_reuse_plan(339, &[full(350, None)]), ReusePlan::Fresh);
18794 let conv = |device_state| ReuseLayer {
18797 full: false,
18798 host_rows: 0,
18799 device_rows: None,
18800 device_state,
18801 };
18802 assert_eq!(
18803 kv_reuse_plan(339, &[conv(true), full(300, Some(339))]),
18804 ReusePlan::Fresh
18805 );
18806 assert_eq!(kv_reuse_plan(339, &[conv(false), full(339, None)]), ReusePlan::Ready);
18807 }
18808
18809 #[test]
18810 fn nll_graph_policy_scopes_only_the_fused_head() {
18811 for (label, unmasked, prefer_graph, native_metal, want_graph, want_head) in [
18812 ("vulkan graph", true, true, false, true, false),
18814 ("native Metal graph", true, true, true, true, true),
18816 ("masked", false, true, false, false, false),
18818 ("graph disabled", true, false, true, false, false),
18819 ] {
18820 let (graph_quality, graph_head_required) =
18821 super::nll_graph_policy(unmasked, prefer_graph, native_metal);
18822 assert_eq!(graph_quality, want_graph, "{label}: graph quality");
18823 assert_eq!(graph_head_required, want_head, "{label}: fused head");
18824 }
18825 }
18826
18827 #[test]
18828 fn mtp_prefill_pair_boundaries_skip_only_final_prompt_row() {
18829 assert_eq!(mtp_prefill_pair_count(0, 128, 256), 128);
18830 assert_eq!(mtp_prefill_pair_count(128, 256, 256), 127);
18831 assert_eq!(mtp_prefill_pair_count(0, 256, 256), 255);
18832 assert_eq!(mtp_prefill_pair_count(256, 256, 256), 0);
18833 assert_eq!(mtp_prefill_pair_count(300, 320, 256), 0);
18834 }
18835
18836 #[test]
18837 fn cancel_flag_stops_generation() {
18838 let mut p = create_test_pipeline(16, 32, 2, 2, 8, 2, 32);
18839 p.cancel.store(true, std::sync::atomic::Ordering::Relaxed);
18842 let r = p.generate_from_ids(&[1, 2, 3], 8, None, None).unwrap();
18843 assert_eq!(r.finish_reason, "cancelled");
18844 assert!(
18845 r.token_ids.is_empty(),
18846 "no tokens after cancel: {:?}",
18847 r.token_ids
18848 );
18849 assert_eq!(p.kv_cache.seq_len(), 0);
18850 assert!(p.kv_history.is_empty());
18851 assert!(!p.graph_want_logits);
18852 assert!(p.graph_logits.is_none());
18853 let r2 = p.generate_from_ids(&[1, 2, 3], 4, None, None).unwrap();
18855 assert_ne!(r2.finish_reason, "cancelled");
18856 }
18857 use super::*;
18858
18859 #[test]
18867 fn dynamic_ffn_equals_the_zeroing_arm() {
18868 let (hidden, inter) = (8usize, 32usize);
18869 let synth = |n: usize, salt: usize| -> Vec<f32> {
18870 (0..n)
18871 .map(|i| (((i * 29 + salt * 13 + 7) % 89) as f32 / 89.0 - 0.5) * 0.6)
18872 .collect()
18873 };
18874 let down = synth(hidden * inter, 3);
18875 let mut down_t = vec![0.0f32; inter * hidden];
18876 for r in 0..hidden {
18877 for c in 0..inter {
18878 down_t[c * hidden + r] = down[r * inter + c];
18879 }
18880 }
18881 let d = DenseFfn {
18882 gate_proj: QTensor::from_f32(synth(inter * hidden, 1), inter, hidden),
18883 up_proj: QTensor::from_f32(synth(inter * hidden, 2), inter, hidden),
18884 down_proj: QTensor::from_f32(down.clone(), hidden, inter),
18885 act: Act::Silu,
18886 down_t: Some(QTensor::from_f32(down_t, inter, hidden)),
18887 segs: Vec::new(),
18888 };
18889 let x = synth(hidden, 11);
18890 let k = 12usize;
18891 let got = dense_ffn_dynamic(&d, &x, None, k).expect("down_t present");
18892 let mut g = vec![0.0f32; inter];
18894 d.gate_proj.matvec(&x, &mut g, None);
18895 let mut u = vec![0.0f32; inter];
18896 d.up_proj.matvec(&x, &mut u, None);
18897 for v in g.iter_mut() {
18898 *v = inference::silu(*v);
18899 }
18900 keep_top_k(&mut g, k);
18901 for i in 0..inter {
18902 g[i] *= u[i];
18903 }
18904 let mut want = vec![0.0f32; hidden];
18905 d.down_proj.matvec(&g, &mut want, None);
18906 for (a, b) in want.iter().zip(&got) {
18907 assert!((a - b).abs() < 1e-5, "dynamic {b} vs reference {a}");
18908 }
18909 }
18910
18911 #[test]
18917 fn tube_ffn_open_equals_dense_and_closed_equals_masked() {
18918 let (hidden, core, tube) = (8usize, 12usize, 8usize);
18919 let inter = core + tube;
18920 let synth = |n: usize, salt: usize| -> Vec<f32> {
18921 (0..n)
18922 .map(|i| (((i * 41 + salt * 17 + 5) % 97) as f32 / 97.0 - 0.5) * 0.5)
18923 .collect()
18924 };
18925 let (g_all, u_all) = (synth(inter * hidden, 1), synth(inter * hidden, 2));
18926 let d_all = synth(hidden * inter, 3);
18927 let dense = DenseFfn {
18929 gate_proj: QTensor::from_f32(g_all.clone(), inter, hidden),
18930 up_proj: QTensor::from_f32(u_all.clone(), inter, hidden),
18931 down_proj: QTensor::from_f32(d_all.clone(), hidden, inter),
18932 act: Act::Silu,
18933 down_t: None,
18934 segs: Vec::new(),
18935 };
18936 let rows =
18937 |v: &[f32], a: usize, b: usize| -> Vec<f32> { v[a * hidden..b * hidden].to_vec() };
18938 let cols = |v: &[f32], a: usize, b: usize| -> Vec<f32> {
18939 let mut o = Vec::with_capacity(hidden * (b - a));
18940 for r in 0..hidden {
18941 o.extend_from_slice(&v[r * inter + a..r * inter + b]);
18942 }
18943 o
18944 };
18945 let tubed = DenseFfn {
18946 down_t: None,
18947 gate_proj: QTensor::from_f32(rows(&g_all, 0, core), core, hidden),
18948 up_proj: QTensor::from_f32(rows(&u_all, 0, core), core, hidden),
18949 down_proj: QTensor::from_f32(cols(&d_all, 0, core), hidden, core),
18950 act: Act::Silu,
18951 segs: vec![FfnSeg {
18952 gate: QTensor::from_f32(rows(&g_all, core, inter), tube, hidden),
18953 up: QTensor::from_f32(rows(&u_all, core, inter), tube, hidden),
18954 down: QTensor::from_f32(cols(&d_all, core, inter), hidden, tube),
18955 start: core,
18956 width: tube,
18957 }],
18958 };
18959 let x = synth(hidden, 7);
18960 let want = dense_ffn(&dense, &x, None);
18961 let got = tube_ffn(&tubed, &x, 1, None, None);
18962 for (a, b) in want.iter().zip(&got) {
18963 assert!((a - b).abs() < 1e-5, "open tube: {a} vs {b}");
18964 }
18965 let mut bits = vec![0u8; inter.div_ceil(8)];
18967 for n in 0..core {
18968 bits[n / 8] |= 1 << (n % 8);
18969 }
18970 let closed = tube_ffn(&tubed, &x, 1, None, Some(&bits));
18971 let masked = dense_ffn_masked(&dense, &x, None, &bits);
18972 for (a, b) in masked.iter().zip(&closed) {
18973 assert!((a - b).abs() < 1e-5, "closed tube: {a} vs {b}");
18974 }
18975 let batch = tube_ffn(&tubed, &x, 1, None, Some(&bits));
18977 for (a, b) in closed.iter().zip(&batch) {
18978 assert_eq!(a, b, "batch arm disagrees with decode arm");
18979 }
18980 }
18981
18982 #[test]
18984 fn sparse_ffn_quant_equals_dense_with_inactive_zeroed() {
18985 let (hidden, inter) = (16usize, 40usize);
18986 let synth = |n: usize, salt: usize| -> Vec<f32> {
18987 (0..n)
18988 .map(|i| (((i * 37 + salt * 11 + 3) % 101) as f32 / 101.0 - 0.5) * 0.4)
18989 .collect()
18990 };
18991 let d = DenseFfn {
18992 gate_proj: QTensor::from_f32(synth(inter * hidden, 1), inter, hidden),
18993 up_proj: QTensor::from_f32(synth(inter * hidden, 2), inter, hidden),
18994 down_proj: QTensor::from_f32(synth(hidden * inter, 3), hidden, inter),
18995 act: Act::Silu,
18996 down_t: None,
18997 segs: Vec::new(),
18998 };
18999 let x = synth(hidden, 9);
19000 let active: Vec<u16> = (0..inter as u16).filter(|i| i % 3 == 0).collect();
19002
19003 let sparse = sparse_ffn_quant(&d, &x, &active, hidden, None);
19004
19005 let mut g = vec![0.0f32; inter];
19007 d.gate_proj.matvec(&x, &mut g, None);
19008 let mut u = vec![0.0f32; inter];
19009 d.up_proj.matvec(&x, &mut u, None);
19010 let act_set: std::collections::HashSet<u16> = active.iter().copied().collect();
19011 for i in 0..inter {
19012 g[i] = if act_set.contains(&(i as u16)) {
19013 inference::silu(g[i]) * u[i]
19014 } else {
19015 0.0
19016 };
19017 }
19018 let mut reference = vec![0.0f32; hidden];
19019 d.down_proj.matvec(&g, &mut reference, None);
19020
19021 let max_d = sparse
19022 .iter()
19023 .zip(&reference)
19024 .map(|(a, b)| (a - b).abs())
19025 .fold(0.0f32, f32::max);
19026 assert!(max_d < 1e-5, "sparse != dense-zeroed: max|Δ| = {max_d}");
19027 }
19028
19029 fn attach_test_mtp(p: &mut Pipeline) {
19031 let (h, inter, heads, kv, hd) = (
19032 p.hidden_size,
19033 p.intermediate_size,
19034 p.num_heads,
19035 p.num_kv_heads,
19036 p.head_dim,
19037 );
19038 let synth = |n: usize, salt: usize| -> Vec<f32> {
19039 (0..n)
19040 .map(|i| (((i * 29 + salt * 23 + 5) % 101) as f32 / 101.0 - 0.5) * 0.2)
19041 .collect()
19042 };
19043 let qt = |rows: usize, cols: usize, salt: usize| -> QTensor {
19044 QTensor::from_f32(synth(rows * cols, salt), rows, cols)
19045 };
19046 p.mtp = Some(MtpModule {
19047 enorm: vec![1.0; h],
19048 hnorm: vec![1.0; h],
19049 eh_proj: qt(h, 2 * h, 301),
19050 layer: LayerWeights {
19051 input_norm: vec![1.0; h],
19052 post_norm: vec![1.0; h],
19053 attn_out_norm: None,
19054 ffn_out_norm: None,
19055 layer_scale: None,
19056 ffn: FfnKind::Dense(DenseFfn {
19057 gate_proj: qt(inter, h, 315),
19058 up_proj: qt(inter, h, 316),
19059 down_proj: qt(h, inter, 317),
19060 act: Act::Silu,
19061 down_t: None,
19062 segs: Vec::new(),
19063 }),
19064 attn: AttnKind::Full {
19065 bias: None,
19066 wq: qt(heads * hd, h, 311),
19067 wk: qt(kv * hd, h, 312),
19068 wv: qt(kv * hd, h, 313),
19069 wo: qt(h, heads * hd, 314),
19070 q_norm: None,
19071 k_norm: None,
19072 output_gate: false,
19073 softplus_gate: None,
19074 },
19075 },
19076 final_norm: vec![1.0; h],
19077 kv: crate::kv_cache::LayerKvCache::new(kv, hd),
19078 });
19079 }
19080
19081 #[test]
19082 fn speculative_equals_vanilla_greedy() {
19083 unsafe { std::env::set_var("CMF_GPU_WGPU_GRAPH", "0") };
19087 let run = |spec: bool| {
19088 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
19089 p.sampler_config.temperature = 0.0;
19090 attach_test_mtp(&mut p);
19091 p.speculative = spec;
19092 let r = p.generate("abcdef", 12, None, None).unwrap();
19093 (r.token_ids, r.mtp_drafted, r.mtp_accepted)
19094 };
19095 let (vanilla, d0, _) = run(false);
19096 let (spec, d1, a1) = run(true);
19097 assert_eq!(d0, 0, "vanilla path must not draft");
19098 assert!(d1 > 0, "speculative path must draft");
19099 assert_eq!(
19100 vanilla, spec,
19101 "speculative must reproduce the exact greedy sequence (accepted {a1}/{d1})"
19102 );
19103 }
19104
19105 #[test]
19106 fn speculative_accepts_constant_oracle() {
19107 unsafe { std::env::set_var("CMF_GPU_WGPU_GRAPH", "0") };
19109 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
19110 p.sampler_config.temperature = 0.0;
19111 p.sampler_config.repetition_penalty = 1.0;
19112 p.weights.lm_head = QTensor::from_f32(vec![0.01; 64 * 8], 64, 8);
19115 attach_test_mtp(&mut p);
19116 p.speculative = true;
19117 let r = p.generate("abcd", 10, None, None).unwrap();
19118 assert!(r.mtp_drafted > 0);
19119 assert_eq!(
19120 r.mtp_accepted, r.mtp_drafted,
19121 "constant logits → every draft accepted"
19122 );
19123 assert!(r.token_ids.windows(2).all(|w| w[0] == w[1]));
19126 }
19127
19128 #[test]
19129 fn empty_prompt_is_an_error_not_a_panic() {
19130 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 260);
19131 let r = p.generate("", 4, None, None);
19132 assert!(r.is_err(), "empty prompt must be a clean error");
19133 }
19134
19135 #[test]
19136 fn every_token_enters_kv_exactly_once() {
19137 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
19138 p.sampler_config.temperature = 0.0;
19140 let r = p.generate("abc", 2, None, None).unwrap();
19141 assert_eq!(r.prompt_tokens, 3);
19142 assert_eq!(
19146 p.kv_cache.seq_len(),
19147 3 + r.tokens_generated - 1,
19148 "each token must be cached exactly once (v1 cached the last prompt token twice)"
19149 );
19150 }
19151
19152 #[test]
19153 fn generation_is_reproducible_with_seed() {
19154 let run = || {
19155 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
19156 p.generate("hello", 8, None, None).unwrap().token_ids
19157 };
19158 assert_eq!(run(), run());
19159 }
19160
19161 #[test]
19162 fn resetting_sampler_restarts_the_seeded_stream() {
19163 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
19164 let config = SamplerConfig {
19165 seed: Some(1234),
19166 ..SamplerConfig::default()
19167 };
19168 p.set_sampler_config(config.clone());
19169 let first = p.generate("hello", 8, None, None).unwrap().token_ids;
19170 p.set_sampler_config(config);
19171 let second = p.generate("hello", 8, None, None).unwrap().token_ids;
19172 assert_eq!(first, second);
19173 }
19174
19175 #[test]
19176 fn eviction_bounds_the_cache() {
19177 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 260);
19178 p.kv_cache.max_seq_len = 6;
19179 p.sampler_config.temperature = 0.0;
19180 let _ = p.generate("abcd", 12, None, None).unwrap();
19181 assert!(
19182 p.kv_cache.seq_len() <= 6 + 1,
19183 "cache must stay bounded by max_seq_len (got {})",
19184 p.kv_cache.seq_len()
19185 );
19186 }
19187
19188 #[test]
19189 fn confidence_matches_tokens_and_is_a_probability() {
19190 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
19191 p.sampler_config.temperature = 0.0;
19192 p.sampler_config.repetition_penalty = 1.0;
19193 let r = p.generate("abcd", 10, None, None).unwrap();
19194 assert_eq!(
19195 r.token_confidence.len(),
19196 r.token_ids.len(),
19197 "one confidence per emitted token"
19198 );
19199 for &c in &r.token_confidence {
19200 assert!((0.0..=1.0).contains(&c), "confidence out of [0,1]: {c}");
19201 }
19202 let logits = [1.0f32, 3.0, 0.5, 3.0];
19204 let p0 = top1_prob_t(&logits, 1, 1.0);
19205 let p1 = top1_prob_t(&logits, 3, 1.0);
19206 assert!((p0 - p1).abs() < 1e-6, "equal logits → equal prob");
19207 assert!(p0 > 0.0 && p0 < 1.0);
19208 let sharp = top1_prob_t(&logits, 1, 1.0);
19210 let soft = top1_prob_t(&logits, 1, 2.0);
19211 assert!(soft < sharp, "higher temperature lowers peak confidence");
19212 }
19213
19214 #[test]
19215 fn trace_is_opt_in_and_parallels_the_output() {
19216 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
19218 p.sampler_config.temperature = 0.0;
19219 p.sampler_config.repetition_penalty = 1.0;
19220 let r = p.generate("abcd", 10, None, None).unwrap();
19221 assert!(r.traces.is_empty(), "trace must be empty unless enabled");
19222
19223 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
19225 p.sampler_config.temperature = 0.0;
19226 p.sampler_config.repetition_penalty = 1.0;
19227 p.set_trace(true);
19228 let r = p.generate("abcd", 10, None, None).unwrap();
19229 assert_eq!(r.traces.len(), r.token_ids.len(), "one trace row per token");
19230 for (i, tr) in r.traces.iter().enumerate() {
19231 assert_eq!(tr.t, i, "trace index is sequential");
19232 assert_eq!(tr.token_id, r.token_ids[i], "trace token_id matches output");
19233 assert_eq!(
19234 tr.confidence, r.token_confidence[i],
19235 "trace confidence matches the confidence channel"
19236 );
19237 assert!(tr.active_skill.is_none() && tr.recon.is_none() && !tr.switched);
19239 }
19240 }
19241
19242 #[test]
19243 fn explain_prefill_logits_match_greedy_first_token() {
19244 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
19248 p.sampler_config.temperature = 0.0;
19249 p.sampler_config.repetition_penalty = 1.0;
19250 let ids = p.tokenizer.encode("abcd");
19251 let logits = p.prefill_next_logits(&ids, None);
19252 let argmax = logits
19253 .iter()
19254 .enumerate()
19255 .max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
19256 .unwrap()
19257 .0 as u32;
19258 let r = p.generate("abcd", 1, None, None).unwrap();
19259 assert_eq!(
19260 argmax, r.token_ids[0],
19261 "explain preview must match greedy emit"
19262 );
19263 }
19264
19265 #[test]
19266 fn laguna_shared_expert_is_unconditionally_added() {
19267 let matrix = |values: Vec<f32>| QTensor::from_f32(values, 2, 2);
19268 let identity = || matrix(vec![1.0, 0.0, 0.0, 1.0]);
19269 let zero_dense = || DenseFfn {
19270 gate_proj: matrix(vec![0.0; 4]),
19271 up_proj: matrix(vec![0.0; 4]),
19272 down_proj: matrix(vec![0.0; 4]),
19273 act: Act::Silu,
19274 down_t: None,
19275 segs: Vec::new(),
19276 };
19277 let shared = DenseFfn {
19278 gate_proj: identity(),
19279 up_proj: identity(),
19280 down_proj: identity(),
19281 act: Act::Silu,
19282 down_t: None,
19283 segs: Vec::new(),
19284 };
19285 let x = [1.0, 2.0];
19286 let expected = dense_ffn(&shared, &x, None);
19287 let moe = MoeFfn {
19288 router: QTensor::from_f32(vec![0.0, 0.0], 1, 2),
19289 experts: vec![zero_dense()],
19290 top_k: 1,
19291 norm_topk_prob: true,
19292 router_sigmoid: true,
19293 expert_bias: None,
19294 routed_scaling: 1.0,
19295 route_tau: None,
19296 shared: Some((shared, None)),
19297 stats: std::cell::RefCell::new(Vec::new()),
19298 act_sq: std::cell::RefCell::new(Vec::new()),
19299 act_rows: std::cell::RefCell::new(Vec::new()),
19300 mask: None,
19301 per_expert_scale: None,
19302 router_input_norm: false,
19303 resonance: None,
19304 grown: Vec::new(),
19305 };
19306 let actual = moe_ffn_cpu(&moe, &x, &[0], &[0.0], 1.0, None);
19307 for (actual, expected) in actual.iter().zip(expected) {
19308 assert!((actual - expected).abs() < 1e-6);
19309 }
19310 }
19311
19312 fn mimo_test_pipeline() -> Pipeline {
19321 let (hs, inter, nh, hd, vd, vocab) = (16usize, 24usize, 4usize, 8usize, 4usize, 64usize);
19322 let kvh = [1usize, 2, 2, 1];
19323 let synth = |n: usize, salt: usize| -> Vec<f32> {
19324 (0..n)
19325 .map(|i| (((i * 31 + salt * 17 + 7) % 97) as f32 / 97.0 - 0.5) * 0.4)
19326 .collect()
19327 };
19328 let qt = |rows: usize, cols: usize, salt: usize| {
19329 QTensor::from_f32(synth(rows * cols, salt), rows, cols)
19330 };
19331 let dense = |inter: usize, salt: usize| DenseFfn {
19332 gate_proj: qt(inter, hs, salt),
19333 up_proj: qt(inter, hs, salt + 1),
19334 down_proj: qt(hs, inter, salt + 2),
19335 act: Act::Silu,
19336 down_t: None,
19337 segs: Vec::new(),
19338 };
19339 let layers: Vec<LayerWeights> = (0..4)
19340 .map(|li| LayerWeights {
19341 input_norm: vec![1.0; hs],
19342 post_norm: vec![1.0; hs],
19343 attn_out_norm: None,
19344 ffn_out_norm: None,
19345 layer_scale: None,
19346 ffn: if li == 0 {
19347 FfnKind::Dense(dense(inter, 50))
19348 } else {
19349 FfnKind::Moe(MoeFfn {
19350 router: qt(4, hs, 60 + li),
19351 experts: (0..4).map(|e| dense(8, 70 + li * 10 + e * 3)).collect(),
19352 top_k: 2,
19353 norm_topk_prob: true,
19354 router_sigmoid: true,
19355 expert_bias: Some(vec![0.02, -0.03, 0.01, 0.0]),
19356 routed_scaling: 1.0,
19357 route_tau: None,
19358 shared: None,
19359 stats: std::cell::RefCell::new(Vec::new()),
19360 act_sq: std::cell::RefCell::new(Vec::new()),
19361 act_rows: std::cell::RefCell::new(Vec::new()),
19362 mask: None,
19363 per_expert_scale: None,
19364 router_input_norm: false,
19365 resonance: None,
19366 grown: Vec::new(),
19367 })
19368 },
19369 attn: AttnKind::Full {
19370 wq: qt(nh * hd, hs, li * 10 + 1),
19371 wk: qt(kvh[li] * hd, hs, li * 10 + 2),
19372 wv: qt(kvh[li] * vd, hs, li * 10 + 3),
19373 wo: qt(hs, nh * vd, li * 10 + 4),
19374 q_norm: None,
19375 k_norm: None,
19376 output_gate: false,
19377 softplus_gate: None,
19378 bias: None,
19379 },
19380 })
19381 .collect();
19382 let mut p = Pipeline::new(
19383 Tokenizer::byte_level(),
19384 PipelineWeights {
19385 embed_tokens: qt(vocab, hs, 100),
19386 layers,
19387 lm_head: qt(vocab, hs, 200),
19388 final_norm: vec![1.0; hs],
19389 },
19390 hs,
19391 inter,
19392 nh,
19393 1, hd,
19395 4,
19396 4,
19397 false,
19398 vocab,
19399 1e-6,
19400 1e7,
19401 NormStyle::Qwen,
19402 4096,
19403 SamplerConfig {
19404 seed: Some(7),
19405 ..Default::default()
19406 },
19407 );
19408 p.layer_dump = None;
19410 p.set_rotary(4, 1e7);
19411 p.sliding_layers = Some(vec![false, true, true, false]);
19412 p.swa = Some((3, usize::MAX));
19413 p.rotary_dim_local = Some(4);
19414 p.inv_freq_local = Some(std::sync::Arc::new(attention::rope_inv_freq(4, 1e4)));
19415 p.set_attn_geometry(Some(kvh.to_vec()), Some(vd)).unwrap();
19416 p.set_layer_sinks(1, vec![0.5, -1.0, 1.5, 0.0]).unwrap();
19417 p.set_layer_sinks(2, vec![-0.25, 2.0, 0.75, -1.5]).unwrap();
19418 p
19419 }
19420
19421 #[test]
19422 fn mimo_embedded_prompt_uses_rows_and_never_reuses_token_only_kv() {
19423 let mut p = mimo_test_pipeline();
19424 p.speculative = false;
19425 p.ignore_eos = true;
19426 p.sampler_config.temperature = 0.0;
19427 p.sampler_config.repetition_penalty = 1.0;
19428 let a = vec![3, 5, 7, 9, 11, 13];
19429 let b = vec![4, 8, 12, 16, 20, 24];
19430 let rows: Vec<_> = b.iter().flat_map(|&id| p.embed_id(id)).collect();
19431 let expected = p.generate_from_ids(&b, 8, None, None).unwrap().token_ids;
19432 let actual = p.generate_from_embeds(&a, &rows, 8, None, None).unwrap().token_ids;
19436 assert_eq!(actual, expected);
19437 assert!(p.kv_history.is_empty());
19438 let mut extended = a.clone();
19439 extended.push(17);
19440 let after_media = p.generate_from_ids(&extended, 8, None, None).unwrap().token_ids;
19441 p.reset_session();
19442 let fresh = p.generate_from_ids(&extended, 8, None, None).unwrap().token_ids;
19443 assert_eq!(after_media, fresh);
19444 assert!(p.generate_from_embeds(&a, &rows[..rows.len()-1], 1, None, None).is_err());
19445 p.reset_session();
19449 p.generate_from_ids(&a, 1, None, None).unwrap();
19450 let mut media_ids = p.kv_history.clone();
19451 assert!(!media_ids.is_empty());
19452 media_ids.extend_from_slice(&[19, 21, 23]);
19453 let source_ids: Vec<_> = (0..media_ids.len()).map(|i| b[i % b.len()]).collect();
19454 let source_rows: Vec<_> = source_ids.iter().flat_map(|&id| p.embed_id(id)).collect();
19455 let mut oracle = mimo_test_pipeline();
19456 oracle.speculative = false;
19457 oracle.ignore_eos = true;
19458 oracle.sampler_config.temperature = 0.0;
19459 oracle.sampler_config.repetition_penalty = 1.0;
19460 let expected = oracle.generate_from_ids(&source_ids, 8, None, None).unwrap().token_ids;
19461 assert_eq!(p.generate_from_embeds(&media_ids, &source_rows, 8, None, None).unwrap().token_ids, expected);
19462 assert!(p.kv_history.is_empty());
19463 let mut bad = rows;
19464 bad[0] = f32::NAN;
19465 assert!(p.generate_from_embeds(&a, &bad, 1, None, None).is_err());
19466 }
19467
19468 fn f32_bits(v: &[f32]) -> Vec<u32> {
19469 v.iter().map(|x| x.to_bits()).collect()
19470 }
19471
19472 #[test]
19479 fn mimo_shaped_decode_matches_prefill_batch_bitwise() {
19480 let mut p = mimo_test_pipeline();
19481 let kv: Vec<usize> = p.kv_cache.layers.iter().map(|l| l.num_kv_heads).collect();
19482 assert_eq!(kv, vec![1, 2, 2, 1]);
19483 assert!(p.kv_cache.layers[0].sinks.is_none() && p.kv_cache.layers[3].sinks.is_none());
19484 assert!(p.kv_cache.layers[1].sinks.is_some() && p.kv_cache.layers[2].sinks.is_some());
19485 let ids: Vec<u32> = (0..12u32).map(|i| (i * 7 + 3) % 64).collect();
19486 let hs = p.hidden_size;
19487 let mut decode = Vec::new();
19488 for (pos, &id) in ids.iter().enumerate() {
19489 let e = p.embed_single(id);
19490 let h = p.forward_layers(&e, pos, None);
19491 decode.push(p.logits_from_hidden(&h));
19492 }
19493 for l in &p.kv_cache.layers {
19494 assert_eq!(l.seq_len, 12);
19495 assert_eq!(l.head_values(0).len(), 12 * 8);
19497 }
19498 assert!(decode.iter().flatten().all(|v| v.is_finite()));
19499
19500 p.clear_sequence_state();
19501 let hb = p.prefill_batch_span(PrefillIn::Ids(&ids), 0, None, 0, p.num_layers);
19502 for pos in 0..ids.len() {
19503 let lg = p.logits_from_hidden(&hb[pos * hs..(pos + 1) * hs]);
19504 assert_eq!(
19505 f32_bits(&decode[pos]),
19506 f32_bits(&lg),
19507 "whole prompt, pos {pos}"
19508 );
19509 }
19510
19511 p.clear_sequence_state();
19512 let a = p.prefill_batch_span(PrefillIn::Ids(&ids[..5]), 0, None, 0, p.num_layers);
19513 let b = p.prefill_batch_span(PrefillIn::Ids(&ids[5..]), 5, None, 0, p.num_layers);
19514 for pos in 0..ids.len() {
19515 let row = if pos < 5 {
19516 &a[pos * hs..(pos + 1) * hs]
19517 } else {
19518 &b[(pos - 5) * hs..(pos - 4) * hs]
19519 };
19520 let lg = p.logits_from_hidden(row);
19521 assert_eq!(
19522 f32_bits(&decode[pos]),
19523 f32_bits(&lg),
19524 "two chunks, pos {pos}"
19525 );
19526 }
19527
19528 let last = |p: &mut Pipeline| {
19531 p.clear_sequence_state();
19532 let hb = p.prefill_batch_span(PrefillIn::Ids(&ids), 0, None, 0, p.num_layers);
19533 p.logits_from_hidden(&hb[11 * hs..12 * hs])
19534 };
19535 let base = last(&mut p);
19536 let mut no_sinks = mimo_test_pipeline();
19537 for l in &mut no_sinks.kv_cache.layers {
19538 l.sinks = None;
19539 }
19540 assert_ne!(
19541 f32_bits(&last(&mut no_sinks)),
19542 f32_bits(&base),
19543 "sinks are live"
19544 );
19545 let mut wide = mimo_test_pipeline();
19546 wide.swa = Some((64, usize::MAX));
19547 assert_ne!(
19548 f32_bits(&last(&mut wide)),
19549 f32_bits(&base),
19550 "window is live"
19551 );
19552
19553 p.clear_sequence_state();
19555 p.ignore_eos = true;
19556 let r = p.generate_from_ids(&ids, 4, None, None).unwrap();
19557 assert_eq!(r.token_ids.len(), 4);
19558 }
19559
19560 fn mimo_test_mtp(n: usize, gain: f32) -> mimo_mtp::MimoMtp {
19563 let (hs, inter, nh, hd, vd, nkv) = (16usize, 24usize, 4usize, 8usize, 4usize, 2usize);
19564 let synth = |len: usize, salt: usize| -> Vec<f32> {
19565 (0..len)
19566 .map(|i| (((i * 37 + salt * 13 + 3) % 89) as f32 / 89.0 - 0.5) * 0.6 * gain)
19567 .collect()
19568 };
19569 let qt = |rows: usize, cols: usize, salt: usize| {
19570 QTensor::from_f32(synth(rows * cols, salt), rows, cols)
19571 };
19572 let layers = (0..n)
19573 .map(|k| {
19574 let s = 500 + k * 40;
19575 let mut kv = crate::kv_cache::LayerKvCache::new(nkv, hd);
19576 kv.sinks = Some(vec![0.3, -0.7, 1.1, 0.0]);
19577 MtpModule {
19578 enorm: vec![1.0; hs],
19579 hnorm: vec![1.0; hs],
19580 eh_proj: qt(hs, 2 * hs, s),
19581 layer: LayerWeights {
19582 input_norm: vec![1.0; hs],
19583 post_norm: vec![1.0; hs],
19584 attn_out_norm: None,
19585 ffn_out_norm: None,
19586 layer_scale: None,
19587 attn: AttnKind::Full {
19588 wq: qt(nh * hd, hs, s + 1),
19589 wk: qt(nkv * hd, hs, s + 2),
19590 wv: qt(nkv * vd, hs, s + 3),
19591 wo: qt(hs, nh * vd, s + 4),
19592 q_norm: None,
19593 k_norm: None,
19594 output_gate: false,
19595 softplus_gate: None,
19596 bias: None,
19597 },
19598 ffn: FfnKind::Dense(DenseFfn {
19599 gate_proj: qt(inter, hs, s + 5),
19600 up_proj: qt(inter, hs, s + 6),
19601 down_proj: qt(hs, inter, s + 7),
19602 act: Act::Silu,
19603 down_t: None,
19604 segs: Vec::new(),
19605 }),
19606 },
19607 final_norm: vec![1.0; hs],
19608 kv,
19609 }
19610 })
19611 .collect();
19612 mimo_mtp::MimoMtp::from_layers(layers)
19613 }
19614
19615 fn mimo_greedy(p: &mut Pipeline, ids: &[u32], n: usize, spec: bool) -> GenerateResult {
19616 p.clear_sequence_state();
19617 p.speculative = spec;
19618 p.ignore_eos = true;
19619 p.sampler_config.temperature = 0.0;
19620 p.generate_from_ids(ids, n, None, None).unwrap()
19621 }
19622
19623 #[test]
19629 fn mimo_mtp_incremental_rounds_equal_the_teacher_forced_table() {
19630 for post in [false, true] {
19633 let mut p = mimo_test_pipeline();
19634 p.weights.final_norm = (0..p.hidden_size).map(|i| 0.5 + 0.1 * i as f32).collect();
19636 let mut st0 = mimo_test_mtp(3, 1.0);
19637 st0.post_norm_hidden = post;
19638 p.mimo_mtp = Some(st0);
19639 let ids: Vec<u32> = (0..14u32).map(|i| (i * 11 + 5) % 64).collect();
19640 let hs = p.hidden_size;
19641 let hb = p.prefill_batch_span(PrefillIn::Ids(&ids), 0, None, 0, p.num_layers);
19642 p.mimo_note_rows(&hb, 0);
19643 let mut st = p.mimo_mtp.take().unwrap();
19644 let k = 3;
19647 let mut inc = Vec::new();
19648 for t in 0..ids.len() - k - 1 {
19649 inc.push(p.mimo_mtp_draft(&mut st, t, &ids, k));
19650 }
19651 let s = ids.len();
19654 let mut reference = vec![vec![0u32; k]; s - k - 1];
19655 let mut fresh = mimo_test_mtp(3, 1.0);
19656 for (layer, m) in fresh.layers.iter_mut().enumerate() {
19657 let n = s - layer - 1;
19658 let mut cats = vec![0.0f32; n * 2 * hs];
19659 for j in 0..n {
19660 let e = p.embed_single(ids[j + layer + 1]);
19661 let raw = &hb[j * hs..(j + 1) * hs];
19662 let g = if post {
19663 inference::rms_norm(raw, &p.weights.final_norm, p.rms_eps, p.norm_style)
19664 } else {
19665 raw.to_vec()
19666 };
19667 let (ce, ch) = cats[j * 2 * hs..(j + 1) * 2 * hs].split_at_mut(hs);
19668 inference::rms_norm_into(&e, &m.enorm, p.rms_eps, p.norm_style, ce);
19669 inference::rms_norm_into(&g, &m.hnorm, p.rms_eps, p.norm_style, ch);
19670 }
19671 let mut x = vec![0.0f32; n * hs];
19672 m.eh_proj.matmat(&cats, n, &mut x, None);
19673 p.mimo_mtp_block(m, &mut x, n, 0);
19674 for (t, row) in reference.iter_mut().enumerate() {
19675 let y = inference::rms_norm(
19676 &x[t * hs..(t + 1) * hs],
19677 &m.final_norm,
19678 p.rms_eps,
19679 p.norm_style,
19680 );
19681 row[layer] = sampler::argmax(&p.lm_head_forward(&y));
19682 }
19683 }
19684 assert_eq!(inc, reference, "post_norm_hidden = {post}");
19685 let distinct: std::collections::HashSet<u32> =
19687 inc.iter().flatten().copied().collect();
19688 assert!(distinct.len() > 3, "{inc:?}");
19689 let last_t = ids.len() - k - 2;
19691 for m in &st.layers {
19692 assert_eq!(m.kv.seq_len, last_t + 1);
19693 }
19694 }
19695 }
19696
19697 #[test]
19704 fn mimo_speculative_greedy_equals_plain_greedy() {
19705 unsafe { std::env::set_var("CMF_GPU_WGPU_GRAPH", "0") };
19706 let ids: Vec<u32> = (0..9u32).map(|i| (i * 7 + 3) % 64).collect();
19707 let n = 24;
19708 let mut p = mimo_test_pipeline();
19709 let plain = mimo_greedy(&mut p, &ids, n, false);
19710 assert_eq!(plain.mtp_drafted, 0);
19711 assert_eq!(plain.token_ids.len(), n);
19712 let plain_kv = p.kv_cache.layers[0].seq_len;
19713
19714 p.mimo_mtp = Some(mimo_test_mtp(3, 1.0));
19716 let spec = mimo_greedy(&mut p, &ids, n, true);
19717 assert!(spec.mtp_drafted > 0, "the round must draft");
19718 assert_eq!(spec.token_ids, plain.token_ids);
19719 assert_eq!(p.kv_cache.layers[0].seq_len, plain_kv);
19720
19721 let mut truth: Vec<u32> = ids.clone();
19724 truth.extend(&plain.token_ids);
19725 let mut noisy = truth.clone();
19726 for (i, t) in noisy.iter_mut().enumerate() {
19727 if i % 5 == 0 {
19728 *t = (*t + 1) % 64;
19729 }
19730 }
19731 let mut st = mimo_test_mtp(3, 1.0);
19732 st.draft_override = Some(noisy);
19733 p.mimo_mtp = Some(st);
19734 let spec = mimo_greedy(&mut p, &ids, n, true);
19735 assert_eq!(spec.token_ids, plain.token_ids);
19736 assert_eq!(p.kv_cache.layers[0].seq_len, plain_kv);
19737 let stats = p.mimo_mtp.as_ref().unwrap().stats.clone();
19738 assert_eq!(stats.accepted as usize, spec.mtp_accepted);
19739 assert!(spec.mtp_accepted > 0 && spec.mtp_accepted < spec.mtp_drafted);
19740 assert!(
19741 stats.accept_hist.iter().filter(|&&c| c > 0).count() >= 3,
19742 "{:?}",
19743 stats.accept_hist
19744 );
19745 assert!(stats.tokens_per_round() > 1.5, "{}", stats.line());
19746
19747 let mut st = mimo_test_mtp(3, 1.0);
19750 st.draft_override = Some(truth);
19751 p.mimo_mtp = Some(st);
19752 let spec = mimo_greedy(&mut p, &ids, n, true);
19753 assert_eq!(spec.token_ids, plain.token_ids);
19754 assert_eq!(spec.mtp_accepted, spec.mtp_drafted);
19755 assert_eq!(p.kv_cache.layers[0].seq_len, plain_kv);
19756
19757 let off = mimo_greedy(&mut p, &ids, n, false);
19759 assert_eq!(off.token_ids, plain.token_ids);
19760 assert_eq!(off.mtp_drafted, 0);
19761 }
19762
19763 #[test]
19772 fn mimo_shaped_model_rides_the_wgpu_graph_geometry() {
19773 let p = mimo_test_pipeline();
19774 assert_eq!(
19775 p.graph_attn_decline_reason(),
19776 Some("per-layer KV head counts")
19777 );
19778 assert_eq!(p.wgpu_graph_attn_decline(), None);
19779 let g0 = p.graph_attn_geom(0).expect("full layer geometry");
19780 assert_eq!(
19781 (g0.nkv, g0.dv, g0.rd, g0.window, g0.sink.is_some()),
19782 (1, 4, 4, None, false)
19783 );
19784 assert_eq!(g0.invf, p.inv_freq.as_slice());
19785 let g1 = p.graph_attn_geom(1).expect("sliding layer geometry");
19786 assert_eq!((g1.nkv, g1.dv, g1.rd, g1.window), (2, 4, 4, Some(3)));
19787 assert_eq!(g1.sink, Some(&[0.5f32, -1.0, 1.5, 0.0][..]));
19788 assert_eq!(g1.invf, p.inv_freq_local.as_ref().unwrap().as_slice());
19789 assert_ne!(g0.invf, g1.invf, "two RoPE tables");
19790 let g3 = p.graph_attn_geom(3).expect("full layer geometry");
19791 assert_eq!((g3.nkv, g3.window, g3.sink.is_some()), (1, None, false));
19792
19793 let emb = p.embed_single(3);
19796 let mut lg = Vec::new();
19797 assert!(
19798 p.try_token_graph_wgpu_steps(&emb, 0, &mut lg, 1, None, None, 0, p.num_layers)
19799 .is_none()
19800 );
19801 let mut hid = emb.clone();
19802 assert_eq!(
19803 p.try_batch_graph_wgpu(&mut hid, &[0], 1, None),
19804 crate::gpu::BatchGraphOutcome::Declined
19805 );
19806 assert_eq!(hid, emb, "a declined batch graph leaves the rows untouched");
19807 assert!(p.try_multi_burst(3, 0, 4).is_none());
19808 assert!(
19809 p.graph_declines().is_empty(),
19810 "no attention decline logged: {:?}",
19811 p.graph_declines()
19812 );
19813 let plain = || create_test_pipeline(8, 16, 2, 1, 4, 2, 32);
19818 assert_eq!(plain().graph_attn_decline_reason(), None);
19819 assert_eq!(plain().wgpu_graph_attn_decline(), None);
19820 assert!(
19821 plain().graph_attn_geom(0).is_none(),
19822 "uniform models keep the historical arms"
19823 );
19824 let mut q = plain();
19825 q.set_layer_sinks(1, vec![0.25, -0.25]).unwrap();
19826 assert_eq!(
19827 q.graph_attn_decline_reason(),
19828 Some("learned attention sinks")
19829 );
19830 assert_eq!(
19831 q.graph_attn_geom(1).unwrap().sink,
19832 Some(&[0.25f32, -0.25][..])
19833 );
19834 let mut q = plain();
19835 q.set_attn_geometry(None, Some(2)).unwrap();
19836 assert_eq!(
19837 q.graph_attn_decline_reason(),
19838 Some("V heads narrower than Q/K heads")
19839 );
19840 assert_eq!(q.graph_attn_geom(0).unwrap().dv, 2);
19841 let mut q = plain();
19842 q.sliding_layers = Some(vec![true, false]);
19843 q.swa = Some((4, usize::MAX));
19844 assert_eq!(q.graph_attn_decline_reason(), Some("sliding-window layers"));
19845 assert_eq!(q.graph_attn_geom(0).unwrap().window, Some(4));
19846 assert_eq!(q.graph_attn_geom(1).unwrap().window, None);
19847
19848 let mut q = mimo_test_pipeline();
19851 q.rope_scale = 2.0;
19852 assert_eq!(
19853 q.wgpu_graph_attn_decline(),
19854 Some("scaled RoPE positions with per-layer geometry")
19855 );
19856 let emb = q.embed_single(3);
19857 assert!(
19858 q.try_token_graph_wgpu_steps(&emb, 0, &mut lg, 1, None, None, 0, q.num_layers)
19859 .is_none()
19860 );
19861 let _ = q.try_token_graph_wgpu_steps(&emb, 1, &mut lg, 1, None, None, 0, q.num_layers);
19862 let lines = q.graph_declines();
19863 assert_eq!(
19864 lines
19865 .iter()
19866 .filter(|l| l.starts_with("wgpu token graph") && l.contains("scaled RoPE"))
19867 .count(),
19868 1,
19869 "{lines:?}"
19870 );
19871 }
19872
19873 #[test]
19874 fn mimo_verify_rewind_preserves_lagging_host_caches() {
19875 let mut p = mimo_test_pipeline();
19876 for (li, layer) in p.kv_cache.layers.iter_mut().enumerate() {
19877 let row = vec![0.0; layer.num_kv_heads * layer.head_dim];
19878 for _ in 0..if li == 0 { 2 } else { 12 } {
19879 layer.append(&row, &row, &[]);
19880 }
19881 }
19882 p.mimo_verify_rewind(9).unwrap();
19883 assert_eq!(p.kv_cache.layers[0].seq_len, 2);
19884 for layer in &p.kv_cache.layers[1..] {
19885 assert_eq!(layer.seq_len, 9);
19886 }
19887 }
19888
19889 #[test]
19893 fn layer_dump_covers_every_position_and_layer_on_both_walks() {
19894 let dir = std::env::temp_dir().join(format!("cmf-layer-dump-{}", std::process::id()));
19895 let _ = std::fs::remove_dir_all(&dir);
19896 let mut p = mimo_test_pipeline();
19897 let hs = p.hidden_size;
19898 let ids = [5u32, 9, 11, 2, 40];
19899 p.layer_dump = Some(dir.join("decode"));
19900 for (pos, &id) in ids.iter().enumerate() {
19901 let e = p.embed_single(id);
19902 let _ = p.forward_layers(&e, pos, None);
19903 }
19904 p.clear_sequence_state();
19905 p.layer_dump = Some(dir.join("prefill"));
19906 let hb = p.prefill_batch_span(PrefillIn::Ids(&ids), 0, None, 0, p.num_layers);
19907 for pos in 0..ids.len() {
19908 for li in 0..p.num_layers {
19909 let name = format!("p{pos:06}_l{li:02}.f32");
19910 let a = std::fs::read(dir.join("decode").join(&name)).unwrap();
19911 let b = std::fs::read(dir.join("prefill").join(&name)).unwrap();
19912 assert_eq!(a.len(), hs * 4, "{name}");
19913 assert_eq!(a, b, "{name}");
19914 }
19915 }
19916 let last = std::fs::read(dir.join("prefill").join("p000004_l03.f32")).unwrap();
19917 let vals: Vec<f32> = last
19918 .chunks(4)
19919 .map(|c| f32::from_le_bytes(c.try_into().unwrap()))
19920 .collect();
19921 assert_eq!(f32_bits(&vals), f32_bits(&hb[4 * hs..5 * hs]));
19922 let _ = std::fs::remove_dir_all(&dir);
19923 }
19924
19925 #[test]
19926 fn attn_geometry_and_sinks_are_validated() {
19927 let mut p = create_test_pipeline(8, 16, 4, 2, 4, 2, 32);
19928 assert!(
19929 p.set_attn_geometry(Some(vec![2]), None).is_err(),
19930 "one entry per layer"
19931 );
19932 assert!(
19933 p.set_attn_geometry(Some(vec![2, 3]), None).is_err(),
19934 "3 does not divide 4"
19935 );
19936 assert!(p.set_attn_geometry(Some(vec![2, 0]), None).is_err());
19937 assert!(p.set_attn_geometry(None, Some(0)).is_err());
19938 assert!(
19939 p.set_attn_geometry(None, Some(5)).is_err(),
19940 "V wider than the head"
19941 );
19942 p.set_attn_geometry(None, Some(4)).unwrap();
19943 assert_eq!(
19944 p.v_head_dim, None,
19945 "v_head_dim == head_dim is the uniform case"
19946 );
19947 p.set_layer_sinks(1, vec![0.1; 4]).unwrap();
19948 p.set_attn_geometry(Some(vec![1, 4]), None).unwrap();
19949 assert_eq!(p.kv_cache.layers[0].num_kv_heads, 1);
19950 assert_eq!(p.kv_cache.layers[1].num_kv_heads, 4);
19951 assert!(
19952 p.kv_cache.layers[1].sinks.is_some(),
19953 "a reshape keeps the layer's sinks"
19954 );
19955 assert_eq!(p.layer_geom(1).0, 4);
19956 assert!(
19957 p.set_layer_sinks(0, vec![0.0; 3]).is_err(),
19958 "one sink per Q head"
19959 );
19960 assert!(p.set_layer_sinks(7, vec![0.0; 4]).is_err());
19961 assert!(p.set_layer_sinks(0, vec![f32::NAN, 0.0, 0.0, 0.0]).is_err());
19962 }
19963
19964 #[test]
19967 fn o1_is_never_armed_on_sink_window_or_narrow_v_layers() {
19968 let cfg = || {
19969 Some(crate::nystrom::O1Cfg {
19970 layers: crate::nystrom::O1Layers::All,
19971 m: 4,
19972 w: 8,
19973 sink: 2,
19974 rect: crate::nystrom::O1Rect::Aggregate,
19975 })
19976 };
19977 let mut p = mimo_test_pipeline();
19978 p.set_o1(cfg());
19979 assert!(!p.o1_active(), "every MiMo-shaped layer is ineligible");
19980 let mut q = create_test_pipeline(8, 16, 2, 1, 4, 3, 64);
19981 q.set_layer_sinks(1, vec![0.0, 0.0]).unwrap();
19982 q.sliding_layers = Some(vec![false, false, true]);
19983 q.swa = Some((4, usize::MAX));
19984 q.set_o1(cfg());
19985 assert_eq!(q.o1_flags, vec![true, false, false]);
19986 }
19987
19988 #[test]
19989 fn o1_batch_transition_publishes_one_epoch_before_serial_handoff() {
19990 const B: usize = 19;
19991 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
19992 p.set_o1(Some(crate::nystrom::O1Cfg {
19993 layers: crate::nystrom::O1Layers::All,
19994 m: 4,
19995 w: 8,
19996 sink: 2,
19997 rect: crate::nystrom::O1Rect::Aggregate,
19998 }));
19999 p.o1_begin_with_prefix(Some(B));
20000 let ids: Vec<u32> = (0..B as u32).collect();
20001 let _ = p.prefill_batch_span(PrefillIn::Ids(&ids), 0, None, 0, p.num_layers);
20002
20003 assert_eq!(p.o1_epoch, 1, "all layers publish one completed transition");
20004 assert!(p.kv_cache.layers.iter().all(|l| l.o1_sealed()));
20005 let next = p.embed_single(B as u32);
20006 let _ = p.forward_layers(&next, B, None);
20007 assert_eq!(p.o1_epoch, 1, "sealed handoff must not republish the epoch");
20008 }
20009
20010 #[test]
20011 fn o1_pair_transition_commits_scratch_before_epoch_publication() {
20012 const B: usize = 19;
20013 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
20014 let gdn_cfg = crate::linear_core::GdnCfg {
20018 num_v_heads: 2,
20019 num_k_heads: 1,
20020 key_head_dim: 2,
20021 value_head_dim: 4,
20022 conv_kernel: 3,
20023 hidden_size: 8,
20024 rms_eps: 1e-6,
20025 output_gate_sigmoid: false,
20026 };
20027 let synth = |n: usize, salt: usize| -> Vec<f32> {
20028 (0..n)
20029 .map(|i| (((i * 13 + salt * 7) % 97) as f32 / 97.0 - 0.5) * 0.4)
20030 .collect()
20031 };
20032 let qt = |rows: usize, cols: usize, salt: usize| {
20033 crate::qtensor::QTensor::from_f32(synth(rows * cols, salt), rows, cols)
20034 };
20035 let c_dim = gdn_cfg.conv_dim();
20036 let vd = gdn_cfg.num_v_heads * gdn_cfg.value_head_dim;
20037 p.weights.layers[0].attn = AttnKind::LinearGdn(crate::linear_core::GdnWeights {
20038 in_proj_qkv: qt(c_dim, 8, 1),
20039 in_proj_z: qt(vd, 8, 2),
20040 in_proj_a: qt(gdn_cfg.num_v_heads, 8, 3),
20041 in_proj_b: qt(gdn_cfg.num_v_heads, 8, 4),
20042 conv1d: synth(c_dim * gdn_cfg.conv_kernel, 5),
20043 a_log: vec![0.2, 0.5],
20044 dt_bias: synth(gdn_cfg.num_v_heads, 6),
20045 norm: vec![1.0; gdn_cfg.value_head_dim],
20046 out_proj: qt(8, vd, 7),
20047 });
20048 p.gdn_cfg = Some(gdn_cfg);
20049 p.set_o1(Some(crate::nystrom::O1Cfg {
20050 layers: crate::nystrom::O1Layers::All,
20051 m: 4,
20052 w: 8,
20053 sink: 2,
20054 rect: crate::nystrom::O1Rect::Aggregate,
20055 }));
20056 p.o1_begin_with_prefix(Some(B));
20057 for pos in 0..B - 2 {
20058 let emb = p.embed_single(pos as u32);
20059 let _ = p.forward_layers(&emb, pos, None);
20060 }
20061 let lane1_state = p.kv_cache.layers[0].linear_state.clone();
20062
20063 let e1 = p.embed_single((B - 2) as u32);
20064 let e2 = p.embed_single((B - 1) as u32);
20065 let _ = p.forward_pair(&e1, &e2, B - 2);
20066
20067 assert_eq!(p.o1_epoch, 1, "pair crossing B publishes one epoch");
20068 assert!(
20069 p.kv_cache
20070 .layers
20071 .iter()
20072 .enumerate()
20073 .all(|(li, l)| !p.o1_flags[li] || l.o1_sealed())
20074 );
20075 assert!(!p.kv_cache.layers[0].linear_state.is_empty());
20076 assert_ne!(
20077 p.kv_cache.layers[0].linear_state, lane1_state,
20078 "real pair must commit GDN lane 2 before returning"
20079 );
20080 assert!(p.kv_cache.layers[0].linear_scratch.is_empty());
20081 let next = p.embed_single(B as u32);
20082 let _ = p.forward_layers(&next, B, None);
20083 assert_eq!(p.o1_epoch, 1, "serial continuation must reuse the epoch");
20084 }
20085
20086 #[test]
20087 fn o1_error_observation_stays_terminal_until_reset() {
20088 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
20089 p.set_o1(Some(crate::nystrom::O1Cfg {
20090 layers: crate::nystrom::O1Layers::All,
20091 m: 4,
20092 w: 8,
20093 sink: 2,
20094 rect: crate::nystrom::O1Rect::Aggregate,
20095 }));
20096 p.o1_begin();
20097 p.kv_cache.layers[0].o1_abort("synthetic transition failure".into());
20098
20099 assert!(p.o1_seal_checked().is_err());
20100 assert!(
20101 p.o1_seal_checked().is_err(),
20102 "retry must see the sticky error"
20103 );
20104 let k = vec![0.2f32; 4];
20105 let v = vec![0.3f32; 4];
20106 p.kv_cache.layers[0].append(&k, &v, &[]);
20107 assert_eq!(p.kv_cache.layers[0].seq_len, 0);
20108
20109 p.reset_session();
20110 p.o1_begin();
20111 p.kv_cache.layers[0].append(&k, &v, &[]);
20112 assert_eq!(p.kv_cache.layers[0].seq_len, 1);
20113 }
20114
20115 #[test]
20116 fn nll_graph_failure_is_terminal_and_request_is_reusable() {
20117 let ids = vec![1u32, 2, 3, 4, 5, 6];
20118 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
20119 p.graph_logits = Some(vec![123.0]);
20120 p.graph_want_logits = true;
20121 p.graph_failed
20122 .store(true, std::sync::atomic::Ordering::Relaxed);
20123 p.cancel.store(true, std::sync::atomic::Ordering::Relaxed);
20124 let err = p.nll_ids_from(&ids, 0).expect_err("prior graph failure");
20125 assert!(err.contains("before NLL"));
20126 assert!(p.graph_logits.is_none());
20127 assert!(!p.graph_want_logits);
20128 assert!(!p.graph_failed.load(std::sync::atomic::Ordering::Relaxed));
20129 assert!(!p.cancel.load(std::sync::atomic::Ordering::Relaxed));
20130
20131 let mut fresh = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
20132 let expected = fresh.nll_ids_from(&ids, 0).expect("fresh NLL");
20133 let actual = p.nll_ids_from(&ids, 0).expect("reused NLL");
20134 assert_eq!(actual.1, expected.1);
20135 assert!((actual.0 - expected.0).abs() < 1e-9);
20136 }
20137
20138 #[test]
20139 fn nll_forward_failure_discards_partial_score_and_clears_sidechannels() {
20140 let ids = vec![1u32, 2, 3, 4, 5, 6];
20141 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
20142 p.nll_test_fail_at = Some(1);
20143 let err = p
20144 .nll_ids_from(&ids, 0)
20145 .expect_err("one-shot forward failure");
20146 assert!(err.contains("forward") || err.contains("score row"));
20147 assert!(!p.graph_failed.load(std::sync::atomic::Ordering::Relaxed));
20148 assert!(!p.graph_want_logits);
20149 assert!(p.graph_logits.is_none());
20150 assert!(p.kv_history.is_empty());
20151
20152 let mut fresh = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
20153 let expected = fresh.nll_ids_from(&ids, 0).expect("fresh NLL");
20154 let actual = p.nll_ids_from(&ids, 0).expect("reused NLL");
20155 assert_eq!(actual.1, expected.1);
20156 assert!((actual.0 - expected.0).abs() < 1e-9);
20157 }
20158
20159 #[test]
20160 fn nll_serial_failure_before_first_row_is_reported() {
20161 let ids = vec![1u32, 2, 3, 4];
20162 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
20163 p.nll_test_force_serial = true;
20164 p.nll_test_fail_at = Some(0);
20165 let err = p.nll_ids_from(&ids, 0).expect_err("serial forward failure");
20166 assert!(err.contains("serial forward"));
20167 assert!(p.kv_history.is_empty());
20168 assert!(!p.graph_failed.load(std::sync::atomic::Ordering::Relaxed));
20169 assert!(!p.cancel.load(std::sync::atomic::Ordering::Relaxed));
20170 }
20171
20172 #[test]
20173 fn ffn_probe_failure_discards_recorder_and_state() {
20174 let ids = vec![1u32, 2, 3, 4];
20175 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
20176 p.nll_test_fail_at = Some(0);
20177 let err = p
20178 .probe_ffn_mass_batch(&ids)
20179 .expect_err("probe forward failure");
20180 assert!(err.contains("NLL"));
20181 assert!(FFN_PROBE.with(|probe| probe.borrow().is_none()));
20182 assert!(p.kv_history.is_empty());
20183 assert!(!p.graph_failed.load(std::sync::atomic::Ordering::Relaxed));
20184 }
20185
20186 #[test]
20187 fn nll_test_controls_are_pipeline_scoped() {
20188 let ids = vec![1u32, 2, 3, 4];
20189 let mut failing = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
20190 let mut unaffected = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
20191 failing.nll_test_force_serial = true;
20192 failing.nll_test_fail_at = Some(0);
20193
20194 assert!(!failing.can_prefill_batched());
20195 assert!(unaffected.can_prefill_batched());
20196 let expected = unaffected
20197 .nll_ids_from(&ids, 0)
20198 .expect("unaffected pipeline remains usable");
20199 let err = failing
20200 .nll_ids_from(&ids, 0)
20201 .expect_err("failure injection belongs to failing pipeline");
20202 assert!(err.contains("serial forward"));
20203 assert!(failing.nll_test_fail_at.is_none());
20204 assert!(unaffected.can_prefill_batched());
20205 let actual = unaffected
20206 .nll_ids_from(&ids, 0)
20207 .expect("unaffected pipeline remains reusable");
20208 assert_eq!(actual.1, expected.1);
20209 assert!((actual.0 - expected.0).abs() < 1e-9);
20210 }
20211
20212 #[test]
20213 fn forward_ids_failure_channel_is_terminal_and_reusable() {
20214 let ids = vec![1u32, 2, 3, 4, 5, 6];
20215 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
20216 p.graph_logits = Some(vec![123.0]);
20217 p.graph_want_logits = true;
20218 p.graph_failed
20219 .store(true, std::sync::atomic::Ordering::Relaxed);
20220 p.cancel.store(true, std::sync::atomic::Ordering::Relaxed);
20221
20222 let err = p
20223 .forward_ids(&ids, None)
20224 .expect_err("a failed forward must not become a valid head result");
20225 assert!(err.contains("forward_ids setup"));
20226 assert!(p.graph_logits.is_none());
20227 assert!(!p.graph_want_logits);
20228 assert!(!p.graph_failed.load(std::sync::atomic::Ordering::Relaxed));
20229 assert!(!p.cancel.load(std::sync::atomic::Ordering::Relaxed));
20230 assert_eq!(p.kv_cache.seq_len(), 0);
20231
20232 let expected = create_test_pipeline(8, 16, 2, 1, 4, 1, 64)
20233 .forward_ids(&ids, None)
20234 .expect("fresh forward_ids");
20235 let actual = p
20236 .forward_ids(&ids, None)
20237 .expect("pipeline remains reusable after a failed forward");
20238 assert_eq!(actual.len(), expected.len());
20239 assert!(
20240 actual
20241 .iter()
20242 .zip(expected)
20243 .all(|(a, b)| (a - b).abs() < 1e-9)
20244 );
20245 assert_eq!(p.kv_cache.seq_len(), ids.len());
20246 }
20247
20248 #[test]
20249 fn sigmoid_router_floor_is_explicit_per_architecture() {
20250 let zero = || DenseFfn {
20256 gate_proj: QTensor::from_f32(vec![0.0; 4], 2, 2),
20257 up_proj: QTensor::from_f32(vec![0.0; 4], 2, 2),
20258 down_proj: QTensor::from_f32(vec![0.0; 4], 2, 2),
20259 act: Act::Silu,
20260 down_t: None,
20261 segs: Vec::new(),
20262 };
20263 let m = MoeFfn {
20264 router: QTensor::from_f32(vec![0.0; 4], 2, 2),
20265 experts: vec![zero(), zero()],
20266 top_k: 1,
20267 norm_topk_prob: true,
20268 router_sigmoid: true,
20269 expert_bias: None,
20270 routed_scaling: 2.5,
20271 route_tau: None,
20272 shared: None,
20273 stats: std::cell::RefCell::new(Vec::new()),
20274 act_sq: std::cell::RefCell::new(Vec::new()),
20275 act_rows: std::cell::RefCell::new(Vec::new()),
20276 mask: None,
20277 per_expert_scale: None,
20278 router_input_norm: false,
20279 resonance: None,
20280 grown: Vec::new(),
20281 };
20282 let logits = [-20.0f32, -20.0];
20283 let (_, p, glm_wsum) = moe_route_with_eps(&logits, &m, None, 1e-20);
20284 let (_, _, generic_wsum) = moe_route(&logits, &m, None);
20285 let expected = (p[0] + 1e-20) / m.routed_scaling;
20286 assert!((glm_wsum - expected).abs() < 1e-15);
20287 assert!(generic_wsum > glm_wsum * 100.0);
20288 }
20289
20290 #[test]
20291 fn resonance_scores_match_formula_and_stable_tie() {
20292 let r = Resonance {
20293 mu: vec![0.0, 0.0, 1.0, 0.0, 0.0, 1.0],
20295 u: vec![1.0, 0.0, 0.0, 1.0, 1.0, 1.0],
20296 k: 1,
20297 bias: vec![1.5, 0.5, 0.0],
20298 shell: Vec::new(),
20299 };
20300 let x = [1.0f32, 1.0];
20301 let mut got = vec![0.0; 3];
20302 r.scores(&x, &mut got);
20303 assert!((got[0] - 0.5).abs() < 1e-6);
20307 assert!((got[1] - 0.5).abs() < 1e-6);
20308 assert!(got[2].abs() < 1e-6);
20309 let best = got
20310 .iter()
20311 .enumerate()
20312 .max_by(|(ia, a), (ib, b)| a.partial_cmp(b).unwrap().then(ib.cmp(ia)))
20313 .map(|(i, _)| i);
20314 assert_eq!(best, Some(0));
20315 assert!(got.iter().all(|v| v.is_finite()));
20316 }
20317
20318 #[test]
20323 fn resonance_shell_masks_outside_keeps_inside_and_trunk_bits() {
20324 let plain = Resonance {
20330 mu: vec![0.0, 0.0, 1.0, 0.0, 0.0, 1.0, 0.0, 1.0],
20331 u: vec![1.0, 0.0, 0.0, 1.0, 1.0, 1.0, 1.0, 1.0],
20332 k: 1,
20333 bias: vec![1.5, 0.5, 0.0, 0.0],
20334 shell: Vec::new(),
20335 };
20336 let shelled = Resonance {
20337 mu: plain.mu.clone(),
20338 u: plain.u.clone(),
20339 k: 1,
20340 bias: plain.bias.clone(),
20341 shell: vec![f32::INFINITY, f32::INFINITY, 6.0, 0.25],
20342 };
20343 assert!(!plain.has_shell());
20344 assert!(shelled.has_shell());
20345 let bits = |v: &[f32]| v.iter().map(|x| x.to_bits()).collect::<Vec<_>>();
20346 let (mut a, mut b) = (vec![0.0; 4], vec![0.0; 4]);
20347 set_growth_shell(Some(true));
20348 assert!(growth_shell_enabled());
20349 let x = [1.0f32, 1.0];
20352 plain.scores(&x, &mut a);
20353 shelled.scores(&x, &mut b);
20354 assert_eq!(bits(&a), bits(&b), "inside every shell: unchanged");
20355 assert!(a[2] == 0.0 && a[3] == 0.0);
20356 let xo = [3.0f32, 0.0];
20359 plain.scores(&xo, &mut a);
20360 shelled.scores(&xo, &mut b);
20361 assert_eq!(a[2], -6.0);
20362 assert_eq!(a[3], -6.0);
20363 assert_eq!(bits(&a[..3]), bits(&b[..3]), "trunk rows + the boundary row");
20364 assert_eq!(b[3], f32::NEG_INFINITY, "outside the shell: −∞");
20365 assert_eq!(shelled.effective_shell(4), shelled.shell);
20366 set_growth_shell(Some(false));
20369 assert!(!growth_shell_enabled());
20370 assert_eq!(shelled.effective_shell(4), vec![f32::INFINITY; 4]);
20371 shelled.scores(&xo, &mut b);
20372 assert_eq!(bits(&a), bits(&b));
20373 set_growth_shell(None);
20374 let short = Resonance {
20377 shell: vec![f32::INFINITY, f32::INFINITY],
20378 ..shelled
20379 };
20380 set_growth_shell(Some(true));
20381 short.scores(&xo, &mut b);
20382 assert_eq!(bits(&a), bits(&b));
20383 set_growth_shell(None);
20384 }
20385
20386 #[test]
20390 fn moe_route_neg_inf_logits_pick_the_best_finite_expert_with_weight_one() {
20391 let zero = || DenseFfn {
20392 gate_proj: QTensor::from_f32(vec![0.0; 4], 2, 2),
20393 up_proj: QTensor::from_f32(vec![0.0; 4], 2, 2),
20394 down_proj: QTensor::from_f32(vec![0.0; 4], 2, 2),
20395 act: Act::Silu,
20396 down_t: None,
20397 segs: Vec::new(),
20398 };
20399 let moe = |sigmoid: bool| MoeFfn {
20400 router: QTensor::from_f32(vec![0.0; 8], 4, 2),
20401 experts: vec![zero(), zero(), zero(), zero()],
20402 top_k: 1,
20403 norm_topk_prob: true,
20404 router_sigmoid: sigmoid,
20405 expert_bias: None,
20406 routed_scaling: 1.0,
20407 route_tau: None,
20408 shared: None,
20409 stats: std::cell::RefCell::new(Vec::new()),
20410 act_sq: std::cell::RefCell::new(Vec::new()),
20411 act_rows: std::cell::RefCell::new(Vec::new()),
20412 mask: None,
20413 per_expert_scale: None,
20414 router_input_norm: false,
20415 resonance: None,
20416 grown: Vec::new(),
20417 };
20418 let logits = [-1.0f32, f32::NEG_INFINITY, -0.5, f32::NEG_INFINITY];
20419 for sigmoid in [false, true] {
20420 let m = moe(sigmoid);
20421 let (idx, p, wsum) = moe_route(&logits, &m, None);
20422 assert_eq!(idx, vec![2], "sigmoid {sigmoid}");
20423 assert_eq!(p[1], 0.0);
20424 assert_eq!(p[3], 0.0);
20425 assert!(p[2] > p[0] && p[0] > 0.0);
20426 assert!(p.iter().all(|v| v.is_finite()));
20427 let w = p[2] / wsum;
20428 if sigmoid {
20429 assert!((w - 1.0).abs() < 1e-5, "sigmoid: weight {w}");
20431 } else {
20432 assert_eq!(w, 1.0, "softmax: the renormalized top-1 weight is exactly 1");
20433 }
20434 let (idx, _, _) = moe_route(&logits, &m, Some(&[true, true, true, true]));
20437 assert_eq!(idx, vec![2]);
20438 }
20439 let m = moe(false);
20441 let (idx, _, _) = moe_route(&[f32::NEG_INFINITY, f32::NEG_INFINITY, f32::NEG_INFINITY, -9.0], &m, None);
20442 assert_eq!(idx, vec![3]);
20443 let (idx, p, wsum) = moe_route(&[f32::NEG_INFINITY; 4], &m, None);
20446 assert_eq!(idx, vec![0]);
20447 assert!(p.iter().all(|&v| v == 0.25));
20448 assert!(wsum.is_finite() && wsum > 0.0);
20449 }
20450
20451 #[test]
20456 fn resonance_route_picks_the_raw_argmax_not_the_softmax_tie() {
20457 let zero = || DenseFfn {
20458 gate_proj: QTensor::from_f32(vec![0.0; 4], 2, 2),
20459 up_proj: QTensor::from_f32(vec![0.0; 4], 2, 2),
20460 down_proj: QTensor::from_f32(vec![0.0; 4], 2, 2),
20461 act: Act::Silu,
20462 down_t: None,
20463 segs: Vec::new(),
20464 };
20465 let moe = |resonant: bool, norm_topk: bool| MoeFfn {
20466 router: QTensor::from_f32(vec![0.0; 6], 3, 2),
20467 experts: vec![zero(), zero(), zero()],
20468 top_k: 1,
20469 norm_topk_prob: norm_topk,
20470 router_sigmoid: false,
20471 expert_bias: None,
20472 routed_scaling: 1.0,
20473 route_tau: None,
20474 shared: None,
20475 stats: std::cell::RefCell::new(Vec::new()),
20476 act_sq: std::cell::RefCell::new(Vec::new()),
20477 act_rows: std::cell::RefCell::new(Vec::new()),
20478 mask: None,
20479 per_expert_scale: None,
20480 router_input_norm: false,
20481 resonance: resonant.then(|| Resonance {
20482 mu: vec![0.0; 6],
20483 u: Vec::new(),
20484 k: 0,
20485 bias: vec![0.0; 3],
20486 shell: Vec::new(),
20487 }),
20488 grown: Vec::new(),
20489 };
20490 let lo = -0.1f32;
20493 let hi = f32::from_bits(lo.to_bits() - 1);
20494 assert!(hi > lo && hi - lo < 2f32.powi(-25));
20495 assert_eq!((lo - hi).exp(), 1.0, "the tie this test is about");
20496 let (idx, _, _) = moe_route(&[lo, hi], &moe(false, true), None);
20498 assert_eq!(idx, vec![0], "softmax tie → lower index (the gated contract)");
20499 for norm in [true, false] {
20502 let m = moe(true, norm);
20503 let (idx, p, wsum) = moe_route(&[lo, hi, f32::NEG_INFINITY], &m, None);
20504 assert_eq!(idx, vec![1], "norm_topk {norm}");
20505 assert_eq!(p, vec![0.0, 1.0, 0.0]);
20506 assert_eq!(p[1] / wsum, 1.0);
20507 let (idx, _, _) = moe_route(&[hi, hi, lo], &m, None);
20510 assert_eq!(idx, vec![0]);
20511 let (idx, _, _) = moe_route(&[f32::NEG_INFINITY, lo, hi], &m, None);
20513 assert_eq!(idx, vec![2]);
20514 let (idx, p, wsum) = moe_route(&[lo, hi, hi], &m, Some(&[true, false, false]));
20515 assert_eq!(idx, vec![0]);
20516 assert_eq!(p[0] / wsum, 1.0);
20517 let (idx, p, wsum) = moe_route(&[f32::NEG_INFINITY; 3], &m, None);
20520 assert_eq!(idx, vec![0]);
20521 assert!(p.iter().all(|v| v.is_finite()) && wsum.is_finite() && wsum > 0.0);
20522 }
20523 }
20524}