1use crate::attention::{self, QwenAttnCfg};
10use crate::inference;
11use crate::kv_cache::KvCache;
12use crate::linear_core::{
13 GdnCfg, GdnWeights, ShortConvCfg, ShortConvWeights, VmfPhaseCfg, VmfPhaseWeights, gdn_forward,
14 gdn_pair, short_conv_forward, short_conv_forward_batch, short_conv_pair, vmf_phase_forward,
15 vmf_phase_pair,
16};
17use crate::pool::Pool;
18use crate::qtensor::QTensor;
19use crate::sampler::{self, SamplerConfig, SamplerScratch, SplitMix64};
20use crate::tokenizer::Tokenizer;
21use cortiq_core::mask::TaskMask;
22use cortiq_core::types::NormStyle;
23
24pub static GLOBAL_USE_GPU: std::sync::atomic::AtomicBool =
25 std::sync::atomic::AtomicBool::new(false);
26
27struct ForwardScratch {
31 n1: Vec<f32>,
32 n2: Vec<f32>,
33 p1: Vec<f32>,
34 p2: Vec<f32>,
35}
36
37impl ForwardScratch {
38 fn new(hidden: usize) -> Self {
39 Self {
40 n1: vec![0.0; hidden],
41 n2: vec![0.0; hidden],
42 p1: vec![0.0; hidden],
43 p2: vec![0.0; hidden],
44 }
45 }
46}
47
48pub struct Pipeline {
50 gpu_plan: Option<std::sync::Arc<Vec<(usize, usize, usize)>>>,
55 pub tokenizer: std::sync::Arc<Tokenizer>,
58 pub kv_cache: KvCache,
59 pub sampler_config: SamplerConfig,
60 pub weights: PipelineWeights,
61 pub hidden_size: usize,
62 pub intermediate_size: usize,
63 pub num_heads: usize,
64 pub num_kv_heads: usize,
65 pub head_dim: usize,
66 pub num_layers: usize,
68 pub physical_layers: usize,
70 pub loop_final_norm: bool,
72 pub vocab_size: usize,
73 pub rms_eps: f64,
74 pub rope_base: f32,
75 pub norm_style: NormStyle,
76 pub rotary_dim: usize,
78 pub attention_heads_per_layer: Option<Vec<usize>>,
80 pub vmf_cfg: Option<VmfPhaseCfg>,
82 pub gdn_cfg: Option<GdnCfg>,
84 pub logit_multiplier: Option<f32>,
86 pub cancel: std::sync::Arc<std::sync::atomic::AtomicBool>,
91 graph_failed: std::sync::atomic::AtomicBool,
96 pub kv_history: Vec<u32>,
101 pub kda_cfg: Option<crate::linear_core::KdaCfg>,
103 pub g3n: Option<Box<(crate::g3n::G3nGlobals, Vec<crate::g3n::G3nLayer>)>>,
106 pub dsv4: Option<
110 Box<(
111 crate::dsv4::Dsv4Globals,
112 Vec<crate::dsv4::Dsv4Layer>,
113 crate::dsv4::Dsv4Cfg,
114 crate::dsv4::Dsv4State,
115 )>,
116 >,
117 pub dsv41: Option<
121 Box<(
122 crate::dsv41::Dsv41Globals,
123 Vec<crate::dsv41::Dsv41Layer>,
124 crate::dsv41::Dsv41Cfg,
125 crate::dsv41::Dsv41State,
126 )>,
127 >,
128 pub dsv41_vision: Option<crate::dsv41_vision::VisionModel>,
130 dsv41_prefill: Option<(Vec<Option<Vec<f32>>>, Vec<bool>)>,
132 pub qwen4_exp: Option<
135 Box<(
136 crate::qwen4_exp::Globals,
137 Vec<crate::qwen4_exp::Layer>,
138 crate::qwen4_exp::Cfg,
139 crate::qwen4_exp::State,
140 )>,
141 >,
142 pub dsv4_mtp: Vec<crate::dsv4::Dsv4Mtp>,
146 pub dspark: Option<crate::dsv4::DsparkState>,
148 pub dspark_pending: Vec<(usize, Vec<u32>, bool, usize)>,
151 pub dspark_hist: Vec<usize>,
153 pub dspark_real: Vec<u32>,
157 pub dspark_trunk_picks: Vec<Vec<(usize, Vec<usize>)>>,
161 pub dspark_exp: Vec<(usize, usize, usize, usize)>,
163 pub dspark_draft_ns: u128,
167 pub short_conv_cfg: Option<ShortConvCfg>,
170 pub mtp: Option<MtpModule>,
172 pub speculative: bool,
174 pub ignore_eos: bool,
180 pub draft_full_streak: u32,
186 pub spec_k_adapt: Option<usize>,
194 pub spec_acc_ewma: f32,
196 rng: SplitMix64,
197 sampler_scratch: SamplerScratch,
198 spec_forced: Option<u32>,
204 spec_q: Vec<Vec<f32>>,
205 spec_p: Vec<f32>,
206 spec_res: Vec<f32>,
207 spec_qs: Vec<sampler::Sparse>,
209 spec_ps: sampler::Sparse,
210 spec_ress: sampler::Sparse,
211 mtp_graph_mode: Option<bool>,
218 #[cfg(target_os = "macos")]
221 metal_verify: Option<MetalVerifyPending>,
222 pub(crate) inv_freq: std::sync::Arc<Vec<f32>>,
226 ws: ForwardScratch,
230 pool: Option<std::sync::Arc<Pool>>,
232 pub(crate) model: Option<std::sync::Arc<cortiq_core::CmfModel>>,
236 pub(crate) dyn_force_f32: bool,
238 pub(crate) dyn_skill_layers: Vec<Option<Vec<usize>>>,
243 pub(crate) dyn_active: Option<usize>,
249 pub(crate) dyn_blend_loaded: bool,
253 pub(crate) dyn_phi_layer: Option<usize>,
256 dyn_phi_ema: Vec<f32>,
258 dyn_phi_seen: usize,
259 pub dyn_router: Option<crate::swarm::DynRouter>,
262 o1_cfg: Option<crate::nystrom::O1Cfg>,
265 o1_epoch: u64,
268 o1_flags: Vec<bool>,
270 trace: bool,
273 calib_temp: f32,
276 #[cfg_attr(not(target_os = "macos"), allow(dead_code))]
278 graph_kv_id: u64,
279 #[cfg_attr(not(target_os = "macos"), allow(dead_code))]
282 graph_want_logits: bool,
283 #[cfg_attr(not(target_os = "macos"), allow(dead_code))]
287 graph_head_required: bool,
288 graph_logits: Option<Vec<f32>>,
291 pub embed_multiplier: f32,
293 pub attn_scale: f32,
296 pub swa: Option<(usize, usize)>,
299 pub sliding_layers: Option<Vec<bool>>,
302 pub inv_freq_local: Option<std::sync::Arc<Vec<f32>>>,
305 pub rotary_dim_local: Option<usize>,
306 pub rope_scale: f32,
307 pub rope_scale_local: f32,
308 pub global_attn: Option<(usize, usize)>,
311 pub inv_freq_global: Option<std::sync::Arc<Vec<f32>>>,
314 pub attn_v_norm: bool,
316 pub qk_norm_after_rope: bool,
318 pub final_softcap: Option<f32>,
320 pub head_clusters: Option<std::sync::Arc<Vec<f32>>>,
324 pub attn_softcap: f32,
326 confidence_on: bool,
330 #[cfg(test)]
333 nll_test_fail_at: Option<usize>,
334 #[cfg(test)]
337 nll_test_force_serial: bool,
338}
339
340#[cfg(target_os = "macos")]
341impl Drop for Pipeline {
342 fn drop(&mut self) {
343 let _ = crate::gpu_metal::wait_replay();
345 crate::gpu::kv_mirror_drop(self.graph_kv_id);
346 }
347}
348
349pub struct PipelineWeights {
354 pub embed_tokens: QTensor,
356 pub layers: Vec<LayerWeights>,
358 pub lm_head: QTensor,
360 pub final_norm: Vec<f32>,
362}
363
364pub struct LayerWeights {
366 pub input_norm: Vec<f32>,
367 pub post_norm: Vec<f32>,
370 pub attn_out_norm: Option<Vec<f32>>,
373 pub layer_scale: Option<f32>,
375 pub ffn_out_norm: Option<Vec<f32>>,
378 pub ffn: FfnKind,
379 pub attn: AttnKind,
380}
381
382#[derive(Clone, Copy, PartialEq, Debug, Default)]
385pub enum Act {
386 #[default]
387 Silu,
388 GeluTanh,
389 Situ {
392 beta: f32,
393 linear_beta: f32,
394 },
395}
396
397impl Act {
398 pub fn from_arch(name: &str) -> Self {
399 if name == "gelu_tanh" {
400 Self::GeluTanh
401 } else {
402 Self::Silu
403 }
404 }
405
406 pub fn from_arch_full(arch: &cortiq_core::ModelArch) -> Self {
408 match arch.hidden_act.as_str() {
409 "situ" => Self::Situ {
410 beta: arch.activation_situ_beta.unwrap_or(1.0) as f32,
411 linear_beta: arch.activation_situ_linear_beta.unwrap_or(0.0) as f32,
412 },
413 other => Self::from_arch(other),
414 }
415 }
416
417 #[inline]
418 pub fn apply(self, x: f32) -> f32 {
419 match self {
420 Self::Silu => inference::silu(x),
421 Self::GeluTanh => inference::gelu_tanh(x),
422 Self::Situ { beta, .. } => beta * (x / beta).tanh() * (1.0 / (1.0 + (-x).exp())),
423 }
424 }
425
426 #[inline]
429 pub fn combine(self, g: f32, u: f32) -> f32 {
430 match self {
431 Self::Situ { linear_beta, .. } if linear_beta > 0.0 => {
432 self.apply(g) * (linear_beta * (u / linear_beta).tanh())
433 }
434 _ => self.apply(g) * u,
435 }
436 }
437}
438
439pub struct DenseFfn {
441 pub gate_proj: QTensor,
442 pub up_proj: QTensor,
443 pub down_proj: QTensor,
444 pub act: Act,
446 pub down_t: Option<QTensor>,
452 pub segs: Vec<FfnSeg>,
459}
460
461pub struct FfnSeg {
466 pub gate: QTensor,
467 pub up: QTensor,
468 pub down: QTensor,
469 pub start: usize,
470 pub width: usize,
471}
472
473pub enum FfnKind {
476 Dense(DenseFfn),
477 Moe(MoeFfn),
481 DenseMoe(Box<DenseMoeFfn>),
488}
489
490pub struct DenseMoeFfn {
492 pub dense: DenseFfn,
493 pub moe: MoeFfn,
494 pub post_norm_1: Vec<f32>,
496 pub pre_norm_2: Vec<f32>,
499 pub post_norm_2: Vec<f32>,
501}
502
503pub struct MoeFfn {
504 pub router: QTensor,
506 pub experts: Vec<DenseFfn>,
507 pub top_k: usize,
508 pub norm_topk_prob: bool,
509 pub router_sigmoid: bool,
512 pub expert_bias: Option<Vec<f32>>,
516 pub routed_scaling: f32,
519 pub route_tau: Option<f32>,
525 pub shared: Option<(DenseFfn, Option<QTensor>)>,
528 pub stats: std::cell::RefCell<Vec<u64>>,
532 pub act_sq: std::cell::RefCell<Vec<f64>>,
539 pub act_rows: std::cell::RefCell<Vec<f32>>,
545 pub mask: Option<Vec<bool>>,
550 pub per_expert_scale: Option<Vec<f32>>,
553 pub router_input_norm: bool,
557 pub resonance: Option<Resonance>,
561}
562
563pub struct Resonance {
565 pub mu: Vec<f32>,
567 pub u: Vec<f32>,
569 pub k: usize,
570 pub bias: Vec<f32>,
572}
573
574impl Resonance {
575 pub fn scores(&self, x: &[f32], out: &mut [f32]) {
577 let h = x.len();
578 let ne = out.len();
579 for e in 0..ne {
580 let mu = &self.mu[e * h..(e + 1) * h];
581 let mut d2 = 0.0f32;
582 for j in 0..h {
583 let d = x[j] - mu[j];
584 d2 += d * d;
585 }
586 let mut proj = 0.0f32;
587 for i in 0..self.k {
588 let u = &self.u[(e * self.k + i) * h..(e * self.k + i + 1) * h];
589 let mut p = 0.0f32;
590 for j in 0..h {
591 p += (x[j] - mu[j]) * u[j];
592 }
593 proj += p * p;
594 }
595 out[e] = self.bias.get(e).copied().unwrap_or(0.0) - (d2 - proj);
596 }
597 }
598}
599
600pub enum AttnKind {
603 Full {
605 wq: QTensor,
606 wk: QTensor,
607 wv: QTensor,
608 wo: QTensor,
609 q_norm: Option<Vec<f32>>,
610 k_norm: Option<Vec<f32>>,
611 output_gate: bool,
612 softplus_gate: Option<(QTensor, bool)>,
616 bias: Option<(Vec<f32>, Vec<f32>, Vec<f32>)>,
618 },
619 Linear(VmfPhaseWeights),
621 LinearGdn(GdnWeights),
623 ShortConv(ShortConvWeights),
626 Mla(Box<MlaWeights>),
634 Kda(Box<crate::linear_core::KdaWeights>),
638}
639
640pub struct MlaWeights {
642 pub q_proj: QTensor,
646 pub q_a: Option<QTensor>,
649 pub q_a_norm: Option<Vec<f32>>,
650 pub kv_a: QTensor,
652 pub kv_a_norm: Vec<f32>,
654 pub kv_b: QTensor,
656 pub o_proj: QTensor,
658 pub nh: usize,
659 pub qk_rope: usize,
660 pub qk_nope: usize,
661 pub v_dim: usize,
662 pub lora: usize,
663 pub scale: f32,
665 pub nope: bool,
667}
668
669pub struct MtpModule {
674 pub enorm: Vec<f32>,
675 pub hnorm: Vec<f32>,
676 pub eh_proj: QTensor,
678 pub layer: LayerWeights,
679 pub final_norm: Vec<f32>,
680 pub kv: crate::kv_cache::LayerKvCache,
681}
682
683#[cfg(target_os = "macos")]
690enum MetalRowsItem<'a> {
691 Gdn {
692 run: Vec<crate::gpu_metal::GdnGpuLayer<'a>>,
693 first: usize,
694 },
695 Attn {
696 l: crate::gpu_metal::AttnGpuLayer<'a>,
697 li: usize,
698 q_norm: Option<&'a [f32]>,
699 k_norm: Option<&'a [f32]>,
700 output_gate: bool,
701 },
702}
703
704#[cfg(target_os = "macos")]
705struct MetalVerifyPending {
706 graph: crate::gpu_metal::VerifyGraph,
707 gdn_layers: Vec<usize>,
708 attn_layers: Vec<(usize, usize)>,
709}
710
711#[cfg(target_os = "macos")]
715struct MetalWarmPending {
716 graph: crate::gpu_metal::VerifyGraph,
717 cpu_stored: usize,
718 b: usize,
719}
720
721#[cfg(target_os = "macos")]
722enum MetalRowsRun {
723 Declined,
725 Failed,
728 Completed(MetalVerifyPending),
729}
730
731#[cfg(target_os = "macos")]
732enum MetalPrefillOutcome {
733 Declined,
734 Failed,
735 Completed(Vec<f32>),
736}
737
738#[cfg(target_os = "macos")]
739enum MetalBatchNllOutcome {
740 Declined,
741 Failed(String),
742 Completed(f64, usize),
743}
744
745#[derive(Clone, Copy)]
749enum SpecTrial {
750 Spec {
751 t0: std::time::Instant,
752 gen0: usize,
753 rounds: usize,
754 },
755 Plain {
756 t0: std::time::Instant,
757 gen0: usize,
758 },
759 Decided {
760 spec: bool,
761 recheck_at: usize,
762 },
763}
764
765pub(crate) fn spec_time_level() -> u8 {
769 static L: std::sync::OnceLock<u8> = std::sync::OnceLock::new();
770 *L.get_or_init(|| match std::env::var("CMF_GRAPH_SPEC_TIME") {
771 Ok(v) => v.trim().parse::<u8>().map(|n| n.max(1)).unwrap_or(1),
772 Err(_) => 0,
773 })
774}
775
776struct SpecStampLog {
782 t_last: std::time::Instant,
783 items: Vec<(&'static str, f32)>,
784}
785
786static SPEC_STAMPS: std::sync::Mutex<Option<SpecStampLog>> = std::sync::Mutex::new(None);
787
788pub(crate) fn spec_stamp(name: &'static str) {
789 if spec_time_level() == 0 {
790 return;
791 }
792 if let Ok(mut g) = SPEC_STAMPS.lock() {
793 if let Some(log) = g.as_mut() {
794 let now = std::time::Instant::now();
795 log.items
796 .push((name, (now - log.t_last).as_secs_f32() * 1e3));
797 log.t_last = now;
798 }
799 }
800}
801
802fn spec_stamps_begin() {
803 if spec_time_level() == 0 {
804 return;
805 }
806 if let Ok(mut g) = SPEC_STAMPS.lock() {
807 *g = Some(SpecStampLog {
808 t_last: std::time::Instant::now(),
809 items: Vec::with_capacity(64),
810 });
811 }
812}
813
814fn spec_stamps_take() -> Vec<(&'static str, f32)> {
815 SPEC_STAMPS
816 .lock()
817 .ok()
818 .and_then(|mut g| g.take())
819 .map(|l| l.items)
820 .unwrap_or_default()
821}
822
823fn spec_stamps_format(items: &[(&'static str, f32)]) -> String {
826 let mut agg: Vec<(&'static str, f32, u32)> = Vec::with_capacity(items.len());
827 for &(n, ms) in items {
828 match agg.iter_mut().find(|e| e.0 == n) {
829 Some(e) => {
830 e.1 += ms;
831 e.2 += 1;
832 }
833 None => agg.push((n, ms, 1)),
834 }
835 }
836 let mut s = String::with_capacity(agg.len() * 16);
837 for (n, ms, k) in agg {
838 if k > 1 {
839 s.push_str(&format!("{n} {ms:.1}/{k} "));
840 } else {
841 s.push_str(&format!("{n} {ms:.1} "));
842 }
843 }
844 s
845}
846
847#[derive(Default, Clone, Copy)]
869struct SpecMon {
870 round_ms: f64,
871 tokens: f64,
872 plain_ms: f64,
873 n: u32,
874 fails: u32,
875 metal: bool,
876}
877
878const SPEC_PROXY_TOKENS: f64 = 3.5;
881const SPEC_PLAIN_MIN_MS: f64 = 200.0;
884
885impl SpecMon {
886 fn round(&mut self, dt_ms: f64, produced: usize) {
887 self.n += 1;
888 if self.n == 1 {
889 return; }
891 let a = if self.n == 2 { 1.0 } else { 0.3 };
892 self.round_ms += a * (dt_ms - self.round_ms);
893 self.tokens += a * (produced as f64 - self.tokens);
894 }
895 fn pays(&self) -> bool {
896 if self.plain_ms > 0.0 {
897 self.tokens * self.plain_ms > self.round_ms * 1.03
898 } else {
899 self.metal && self.tokens >= SPEC_PROXY_TOKENS
900 }
901 }
902 fn plain_done(&self, t0: std::time::Instant, gen0: usize, generated: usize) -> bool {
904 let n = generated.saturating_sub(gen0);
905 if n >= 8 {
906 return true;
907 }
908 self.metal && n >= 2 && t0.elapsed().as_secs_f64() * 1e3 >= SPEC_PLAIN_MIN_MS
909 }
910}
911
912pub struct GenerateResult {
914 pub text: String,
915 pub token_ids: Vec<u32>,
916 pub prompt_tokens: usize,
917 pub tokens_generated: usize,
918 pub finish_reason: String,
919 pub mtp_drafted: usize,
921 pub mtp_accepted: usize,
922 pub token_confidence: Vec<f32>,
927 pub traces: Vec<TokenTrace>,
930}
931
932#[derive(Clone, Debug)]
937pub struct TokenTrace {
938 pub t: usize,
940 pub token_id: u32,
942 pub confidence: f32,
944 pub active_skill: Option<String>,
946 pub recon: Option<f32>,
950 pub switched: bool,
953}
954
955#[cfg_attr(not(test), allow(dead_code))]
960fn top1_prob_t(logits: &[f32], id: u32, temp: f32) -> f32 {
961 let t = if temp > 1e-3 { temp } else { 1.0 };
962 let max = logits.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
963 let sum: f32 = logits.iter().map(|&v| ((v - max) / t).exp()).sum();
964 if sum > 0.0 {
965 (((logits[id as usize] - max) / t).exp()) / sum
966 } else {
967 0.0
968 }
969}
970
971fn prefill_batched() -> bool {
974 std::env::var("CMF_PREFILL")
975 .map(|v| v != "seq")
976 .unwrap_or(true)
977}
978
979#[inline]
983fn nll_graph_policy(
984 unmasked: bool,
985 prefer_graph: bool,
986 native_metal: bool,
987) -> (bool, bool) {
988 let graph_quality = unmasked && prefer_graph;
989 let fused_head_quality = graph_quality && native_metal;
990 (graph_quality, fused_head_quality)
991}
992
993#[derive(Clone, Copy)]
997enum PrefillIn<'a> {
998 Ids(&'a [u32]),
999 Hidden(&'a [f32]),
1000}
1001
1002impl Pipeline {
1009 fn can_prefill_batched(&self) -> bool {
1010 #[cfg(test)]
1011 let force_serial = self.nll_test_force_serial;
1012 #[cfg(not(test))]
1013 let force_serial = false;
1014 prefill_batched() && !force_serial && !self.weights.layers.is_empty()
1015 }
1016
1017 fn automatic_gpu_prefix(&self) -> Option<usize> {
1020 let (model, _, _, _) = self.weights.embed_tokens.graph_weight()?;
1021 crate::gpu::automatic_layer_prefix(&model, self.num_layers, self.physical_layers)
1022 }
1023}
1024
1025pub fn prefill_chunk() -> usize {
1032 if let Some(n) = std::env::var("CMF_PREFILL_CHUNK")
1033 .ok()
1034 .and_then(|v| v.parse::<usize>().ok())
1035 {
1036 return n.max(1);
1037 }
1038 if cfg!(target_os = "macos") {
1039 512
1040 } else if cfg!(target_arch = "aarch64") {
1041 256
1044 } else {
1045 48
1046 }
1047}
1048
1049#[inline]
1055fn mtp_prefill_pair_count(start: usize, end: usize, input_len: usize) -> usize {
1056 if end <= start || start >= input_len {
1057 return 0;
1058 }
1059 let rows = (end.min(input_len) - start).min(input_len - start);
1060 if end < input_len {
1061 rows
1062 } else {
1063 rows.saturating_sub(1)
1064 }
1065}
1066
1067pub type TokenCallback = Box<dyn FnMut(&str) -> bool + Send>;
1069
1070impl Pipeline {
1071 fn clear_sequence_state(&mut self) {
1079 #[cfg(target_os = "macos")]
1082 let _ = crate::gpu_metal::wait_replay();
1083 self.kv_cache.clear();
1084 self.kv_history.clear();
1085 if let Some(b) = &mut self.dsv41 {
1086 b.3.clear();
1087 }
1088 crate::gpu::graph_kv_reset(self.graph_kv_id);
1089 crate::gpu::graph_kv_reset(self.mtp_kv_id());
1094 }
1095
1096 fn finish_generation(
1102 &mut self,
1103 mtp: &mut Option<MtpModule>,
1104 router: &mut Option<crate::swarm::DynRouter>,
1105 clear_sequence: bool,
1106 ) {
1107 if router.is_some() {
1111 let _ = self.set_active_skill(None);
1112 }
1113 #[cfg(target_os = "macos")]
1120 let clear_sequence = clear_sequence || !crate::gpu_metal::wait_replay();
1121 if clear_sequence {
1122 self.clear_sequence_state();
1123 if let Some(m) = mtp.as_mut() {
1124 m.kv.clear();
1130 }
1131 if let Some(m) = self.mtp.as_mut() {
1132 m.kv.clear();
1136 }
1137 }
1138 self.graph_want_logits = false;
1139 self.graph_head_required = false;
1140 self.graph_logits = None;
1141 self.graph_failed
1142 .store(false, std::sync::atomic::Ordering::Relaxed);
1143 self.cancel
1144 .store(false, std::sync::atomic::Ordering::Relaxed);
1145 self.dyn_router = router.take().or(self.dyn_router.take());
1146 self.mtp = mtp.take().or(self.mtp.take());
1147 self.mtp_graph_mode = None;
1148 self.spec_forced = None;
1149 }
1150
1151 fn check_forward_graph(&mut self, phase: &str, pos: usize) -> Result<(), String> {
1155 if self
1156 .graph_failed
1157 .swap(false, std::sync::atomic::Ordering::Relaxed)
1158 {
1159 self.cancel
1160 .store(false, std::sync::atomic::Ordering::Relaxed);
1161 self.clear_sequence_state();
1162 self.graph_logits = None;
1163 self.graph_want_logits = false;
1164 self.graph_head_required = false;
1165 return Err(format!("GPU graph failed during {phase} at position {pos}"));
1166 }
1167 Ok(())
1168 }
1169
1170 #[cfg(target_os = "macos")]
1171 fn fail_metal_graph(&mut self, reason: &str) {
1172 crate::pipeline::METAL_GRAPH_ERRORS
1173 .fetch_add(1, std::sync::atomic::Ordering::Relaxed);
1174 self.clear_sequence_state();
1175 self.graph_logits = None;
1176 self.graph_failed
1177 .store(true, std::sync::atomic::Ordering::Relaxed);
1178 self.cancel
1179 .store(true, std::sync::atomic::Ordering::Relaxed);
1180 tracing::error!("native Metal TokenGraph failed closed: {reason}");
1181 }
1182
1183 fn nll_begin(&mut self) -> Result<(), String> {
1188 if self
1189 .graph_failed
1190 .swap(false, std::sync::atomic::Ordering::Relaxed)
1191 {
1192 self.cancel
1193 .store(false, std::sync::atomic::Ordering::Relaxed);
1194 self.clear_sequence_state();
1195 self.graph_logits = None;
1196 self.graph_want_logits = false;
1197 self.graph_head_required = false;
1198 return Err("GPU graph failed before NLL scoring".to_string());
1199 }
1200 self.clear_sequence_state();
1201 self.graph_logits = None;
1202 self.graph_want_logits = false;
1203 self.graph_head_required = false;
1204 Ok(())
1205 }
1206
1207 fn nll_end(&mut self) {
1211 self.clear_sequence_state();
1212 self.graph_logits = None;
1213 self.graph_want_logits = false;
1214 self.graph_head_required = false;
1215 self.graph_failed
1216 .store(false, std::sync::atomic::Ordering::Relaxed);
1217 }
1218
1219 fn nll_check_graph(&mut self, phase: &str, pos: usize) -> Result<(), String> {
1222 #[cfg(test)]
1223 if self.nll_test_fail_at == Some(pos) {
1224 self.nll_test_fail_at = None;
1225 self.graph_failed
1226 .store(true, std::sync::atomic::Ordering::Relaxed);
1227 self.cancel
1228 .store(true, std::sync::atomic::Ordering::Relaxed);
1229 }
1230 if self
1231 .graph_failed
1232 .swap(false, std::sync::atomic::Ordering::Relaxed)
1233 {
1234 self.cancel
1235 .store(false, std::sync::atomic::Ordering::Relaxed);
1236 self.clear_sequence_state();
1237 self.graph_logits = None;
1238 self.graph_want_logits = false;
1239 return Err(format!(
1240 "GPU graph failed during NLL {phase} at position {pos}"
1241 ));
1242 }
1243 Ok(())
1244 }
1245
1246 #[inline]
1250 pub fn phys_layer(&self, virtual_idx: usize) -> usize {
1251 virtual_idx % self.physical_layers
1252 }
1253
1254 #[inline]
1257 pub fn is_loop_end(&self, virtual_idx: usize) -> bool {
1258 self.loop_final_norm && (virtual_idx + 1) % self.physical_layers == 0
1259 }
1260
1261 #[allow(clippy::too_many_arguments)]
1263
1264 #[cfg(target_os = "macos")]
1283 fn graph_prefill_preferred(&self) -> bool {
1284 let graph_force = crate::gpu::q1_force() || crate::gpu::q2tp_gpu_opt_in();
1285 if !crate::gpu::enabled_here()
1286 || !graph_force
1287 || std::env::var("CMF_GPU_BLOCK")
1288 .map(|v| v == "0")
1289 .unwrap_or(false)
1290 || std::env::var("CMF_PREFILL_GRAPH").as_deref() == Ok("0")
1293 {
1294 return false;
1295 }
1296 self.weights
1297 .layers
1298 .iter()
1299 .any(|lw| {
1300 matches!(&lw.attn, AttnKind::LinearGdn(w) if w.in_proj_qkv.metal_graph_parts().is_some())
1301 })
1302 }
1303
1304 #[cfg(not(target_os = "macos"))]
1305 fn graph_prefill_preferred(&self) -> bool {
1306 let graph_on = crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Prefill);
1314 if !graph_on || !crate::gpu::enabled_here() {
1315 return false;
1316 }
1317 if self.o1_active() {
1331 return false;
1332 }
1333 if self
1334 .weights
1335 .layers
1336 .iter()
1337 .any(|lw| matches!(&lw.attn, AttnKind::LinearGdn(_)))
1338 {
1339 return true;
1340 }
1341 self.weights
1352 .layers
1353 .iter()
1354 .any(|lw| matches!(&lw.ffn, FfnKind::Moe(_)))
1355 && self.automatic_gpu_prefix().is_none()
1356 }
1357
1358 #[cfg(target_os = "macos")]
1359 fn q1_graph_gpu(
1360 &mut self,
1361 start: usize,
1362 upto: Option<usize>,
1363 position: usize,
1364 h: &mut [f32],
1365 ) -> usize {
1366 let _mt0 = std::time::Instant::now(); use crate::gpu::{AttnGpuLayer, GdnGpuCfg, GdnGpuLayer, GraphDims, MetalFfn, TokenGraph};
1368 let graph_force = crate::gpu::q1_force() || crate::gpu::q2tp_gpu_opt_in();
1369 if self.attn_softcap > 0.0 || !crate::gpu::enabled_here()
1371 || !graph_force
1372 || std::env::var("CMF_GPU_BLOCK")
1373 .map(|v| v == "0")
1374 .unwrap_or(false)
1375 {
1376 if std::env::var("CMF_GRAPH_DBG").is_ok() {
1377 eprintln!(
1378 "block-graph: front gate (softcap={} enabled_here={} graph_force={})",
1379 self.attn_softcap > 0.0,
1380 crate::gpu::enabled_here(),
1381 graph_force,
1382 );
1383 }
1384 if self.graph_head_required {
1385 self.fail_metal_graph("native graph front gate refused");
1386 }
1387 return start;
1388 }
1389 if self.swa.is_some()
1393 || self.global_attn.is_some()
1394 || self.attention_heads_per_layer.is_some()
1395 || self.attn_v_norm
1396 || self.weights.layers.iter().any(|lw| {
1397 lw.attn_out_norm.is_some()
1398 || lw.ffn_out_norm.is_some()
1399 || lw.layer_scale.is_some()
1400 || matches!(&lw.ffn, FfnKind::Dense(d) if d.act != Act::Silu)
1401 })
1402 {
1403 if std::env::var("CMF_GRAPH_DBG").is_ok() {
1404 eprintln!(
1405 "block-graph: arch ineligible (swa={} gattn={} hpl={} vnorm={} scale_delta={:.2e})",
1406 self.swa.is_some(),
1407 self.global_attn.is_some(),
1408 self.attention_heads_per_layer.is_some(),
1409 self.attn_v_norm,
1410 (self.attn_scale - 1.0 / (self.head_dim as f32).sqrt()).abs(),
1411 );
1412 }
1413 if self.graph_head_required {
1414 self.fail_metal_graph("native graph architecture gate refused");
1415 }
1416 return start;
1417 }
1418 let limit = upto
1421 .map(|u| u + 1)
1422 .unwrap_or(self.num_layers)
1423 .min(self.num_layers);
1424
1425 enum Item<'a> {
1426 Gdn {
1427 run: Vec<GdnGpuLayer<'a>>,
1428 first: usize,
1429 },
1430 Attn {
1431 l: AttnGpuLayer<'a>,
1432 li: usize,
1433 q_norm: Option<&'a [f32]>,
1434 k_norm: Option<&'a [f32]>,
1435 output_gate: bool,
1436 bias: Option<(&'a [f32], &'a [f32], &'a [f32])>,
1437 full_gpu: bool,
1440 },
1441 }
1442
1443 let attend_mode = std::env::var("CMF_GPU_ATTEND").unwrap_or_else(|_| "auto".into());
1450 let attend_contract = attend_mode != "0"
1451 && attend_mode != "off"
1452 && self.head_dim % 4 == 0
1453 && self.head_dim <= 256
1454 && self.rotary_dim >= 2
1455 && self.rotary_dim <= self.head_dim
1456 && (self.rotary_dim / 2) % 32 == 0
1457 && self.num_kv_heads > 0
1458 && self.num_heads % self.num_kv_heads == 0;
1459
1460 let mut plan: Vec<Item> = Vec::new();
1461 let mut model_ref: Option<std::sync::Arc<cortiq_core::CmfModel>> = None;
1462 let block_diag = std::env::var("CMF_GRAPH_DBG").is_ok();
1464 let mut scan = start;
1465 while scan < limit {
1466 let lw = &self.weights.layers[self.phys_layer(scan)];
1467 let ffn = match &lw.ffn {
1468 FfnKind::Dense(d) if d.segs.is_empty() => {
1469 let (Some(g), Some(u), Some(dn)) = (
1470 d.gate_proj.metal_graph_parts(),
1471 d.up_proj.metal_graph_parts(),
1472 d.down_proj.metal_graph_parts(),
1473 ) else {
1474 if block_diag {
1475 eprintln!(
1476 "block-graph: L{scan} FFN trio not graph-mappable — run ends"
1477 );
1478 }
1479 break;
1480 };
1481 MetalFfn::Dense {
1482 gate: g,
1483 up: u,
1484 down: dn,
1485 }
1486 }
1487 FfnKind::Moe(m) => {
1488 let Some(moe) = metal_moe_graph_parts(m, self.hidden_size) else {
1489 if block_diag {
1490 eprintln!(
1491 "block-graph: L{scan} MoE outside the graph contract — run ends"
1492 );
1493 }
1494 break;
1495 };
1496 if let QTensor::Mapped { model, .. } = &m.experts[0].gate_proj {
1497 model_ref.get_or_insert_with(|| model.clone());
1498 }
1499 MetalFfn::Moe(moe)
1500 }
1501 _ => {
1502 if block_diag {
1503 eprintln!("block-graph: L{scan} non-graph FFN — run ends");
1504 }
1505 break;
1506 }
1507 };
1508 match &lw.attn {
1509 AttnKind::LinearGdn(w) if self.gdn_cfg.is_some() => {
1510 let parts = (
1511 w.in_proj_qkv.metal_graph_parts(),
1512 w.in_proj_z.metal_graph_parts(),
1513 w.in_proj_a.f32_parts(),
1514 w.in_proj_b.f32_parts(),
1515 w.out_proj.metal_graph_parts(),
1516 );
1517 let (Some(qkv), Some(z), Some(a), Some(b), Some(out)) = parts else {
1518 if block_diag {
1519 eprintln!(
1520 "block-graph: L{scan} GDN parts refused (qkv={} z={} a_f32={} b_f32={} out={})",
1521 w.in_proj_qkv.metal_graph_parts().is_some(),
1522 w.in_proj_z.metal_graph_parts().is_some(),
1523 w.in_proj_a.f32_parts().is_some(),
1524 w.in_proj_b.f32_parts().is_some(),
1525 w.out_proj.metal_graph_parts().is_some(),
1526 );
1527 }
1528 break;
1529 };
1530 if let QTensor::Mapped { model, .. } = &w.in_proj_qkv {
1531 model_ref.get_or_insert_with(|| model.clone());
1532 }
1533 let gl = GdnGpuLayer {
1534 attn_norm: &lw.input_norm,
1535 post_norm: &lw.post_norm,
1536 qkv,
1537 z,
1538 a,
1539 b,
1540 out,
1541 ffn,
1542 conv1d: &w.conv1d,
1543 a_log: &w.a_log,
1544 dt_bias: &w.dt_bias,
1545 gnorm: &w.norm,
1546 };
1547 match plan.last_mut() {
1548 Some(Item::Gdn { run, .. }) => run.push(gl),
1549 _ => plan.push(Item::Gdn {
1550 run: vec![gl],
1551 first: scan,
1552 }),
1553 }
1554 }
1555 AttnKind::Full {
1556 wq,
1557 wk,
1558 wv,
1559 wo,
1560 q_norm,
1561 k_norm,
1562 output_gate,
1563 softplus_gate: None,
1564 bias,
1565 } if !self.kv_cache.layers[scan].o1_sealed()
1566 || std::env::var("CMF_O1_METAL").as_deref() == Ok("1") =>
1571 {
1572 let parts = (
1573 wq.metal_graph_parts(),
1574 wk.metal_graph_parts(),
1575 wv.metal_graph_parts(),
1576 wo.metal_graph_parts(),
1577 );
1578 let (Some(pq), Some(pk), Some(pv), Some(po)) = parts else {
1579 break;
1580 };
1581 if let QTensor::Mapped { model, .. } = wq {
1582 model_ref.get_or_insert_with(|| model.clone());
1583 }
1584 let cache = &self.kv_cache.layers[scan];
1585 let o1_metal = cache.o1.is_some()
1589 && std::env::var("CMF_O1_METAL").as_deref() == Ok("1")
1590 && cache.o1_views().is_some();
1591 let full_gpu = attend_contract
1592 && cache.mode == crate::kv_cache::KvMode::F32
1593 && (cache.o1.is_none() || o1_metal)
1594 && bias.is_none()
1595 && pq.1 == self.num_heads * self.head_dim * (1 + *output_gate as usize)
1596 && pk.1 == self.num_kv_heads * self.head_dim
1597 && pv.1 == self.num_kv_heads * self.head_dim
1598 && po.2 == self.num_heads * self.head_dim;
1599 plan.push(Item::Attn {
1600 l: AttnGpuLayer {
1601 attn_norm: &lw.input_norm,
1602 post_norm: &lw.post_norm,
1603 wq: pq,
1604 wk: pk,
1605 wv: pv,
1606 wo: po,
1607 ffn,
1608 },
1609 li: scan,
1610 q_norm: q_norm.as_deref(),
1611 k_norm: k_norm.as_deref(),
1612 output_gate: *output_gate,
1613 bias: bias
1614 .as_ref()
1615 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
1616 full_gpu,
1617 });
1618 }
1619 _ => break,
1620 }
1621 scan += 1;
1622 }
1623 let Some(model) = model_ref else {
1624 if std::env::var("CMF_GRAPH_DBG").is_ok() {
1625 eprintln!("q1-graph: no model ref (start {start}, scanned to {scan})");
1626 }
1627 if self.graph_head_required {
1628 self.fail_metal_graph("native graph has no mapped model reference");
1629 }
1630 return start;
1631 };
1632 if plan.is_empty() {
1633 if std::env::var("CMF_GRAPH_DBG").is_ok() {
1634 eprintln!("q1-graph: empty plan at layer {start}");
1635 }
1636 if self.graph_head_required {
1637 self.fail_metal_graph("native graph plan is empty");
1638 }
1639 return start;
1640 }
1641 let has_moe = plan.iter().any(|it| match it {
1642 Item::Gdn { run, .. } => run.iter().any(|l| matches!(l.ffn, MetalFfn::Moe(_))),
1643 Item::Attn { l, .. } => matches!(l.ffn, MetalFfn::Moe(_)),
1644 });
1645 let has_gdn = plan.iter().any(|it| matches!(it, Item::Gdn { .. }));
1646 let dev_attend = attend_contract
1647 && (self.head_dim <= 128
1648 || has_moe
1649 || (self.head_dim <= 256 && has_gdn)
1655 || attend_mode == "force"
1656 || attend_mode == "256");
1657 if !dev_attend {
1658 for it in &mut plan {
1659 if let Item::Attn { li, full_gpu, .. } = it {
1660 let keep_o1 = self.kv_cache.layers[*li].o1.is_some()
1663 && std::env::var("CMF_O1_METAL").as_deref() == Ok("1");
1664 if !keep_o1 {
1665 *full_gpu = false;
1666 }
1667 }
1668 }
1669 }
1670 if std::env::var("CMF_GRAPH_DBG").is_ok() {
1671 use std::sync::atomic::{AtomicBool, Ordering};
1672 static SAID: AtomicBool = AtomicBool::new(false);
1673 if !SAID.swap(true, Ordering::Relaxed) {
1674 let fg = plan
1675 .iter()
1676 .filter(|it| matches!(it, Item::Attn { full_gpu: true, .. }))
1677 .count();
1678 let att = plan
1679 .iter()
1680 .filter(|it| matches!(it, Item::Attn { .. }))
1681 .count();
1682 eprintln!(
1683 "q1-graph: plan of {} items from layer {start} to {scan} | dev_attend={dev_attend} full_gpu {fg}/{att} | hd={} rd={} nkv={} nh={}",
1684 plan.len(),
1685 self.head_dim,
1686 self.rotary_dim,
1687 self.num_kv_heads,
1688 self.num_heads,
1689 );
1690 }
1691 }
1692 let dims = GraphDims {
1693 hidden: self.hidden_size,
1694 eps: self.rms_eps as f32,
1695 gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
1696 };
1697 let Some(mut graph) = TokenGraph::new(&model, dims, h) else {
1698 if self.graph_head_required {
1699 self.fail_metal_graph("native TokenGraph allocation refused");
1700 }
1701 return start;
1702 };
1703 let gcfg = self.gdn_cfg.map(|cfg| GdnGpuCfg {
1704 nv: cfg.num_v_heads,
1705 nk: cfg.num_k_heads,
1706 dk: cfg.key_head_dim,
1707 dv: cfg.value_head_dim,
1708 kk: cfg.conv_kernel,
1709 hidden: self.hidden_size,
1710 inter: self.intermediate_size,
1711 c_dim: cfg.conv_dim(),
1712 eps: cfg.rms_eps as f32,
1713 gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
1714 });
1715 let mut valid = 0usize;
1719 let mut end = start;
1720 crate::gpu::stageprof(1, _mt0.elapsed()); if std::env::var("CMF_PLAN_DUMP").is_ok() {
1722 static ONCE: std::sync::Once = std::sync::Once::new();
1723 ONCE.call_once(|| {
1724 for it in &plan {
1725 match it {
1726 Item::Gdn { first, run } => {
1727 eprintln!("plan: Gdn first={first} len={}", run.len())
1728 }
1729 Item::Attn { li, full_gpu, .. } => {
1730 eprintln!("plan: Attn li={li} full_gpu={full_gpu}")
1731 }
1732 }
1733 }
1734 });
1735 }
1736 for item in &plan {
1737 let ok = match item {
1738 Item::Gdn { run, .. } => gcfg
1739 .as_ref()
1740 .map(|gc| run.iter().all(|l| graph.gdn_ok(l, gc)))
1741 .unwrap_or(false),
1742 Item::Attn { l, .. } => graph.attn_ok(l),
1743 };
1744 if !ok {
1745 if block_diag {
1746 eprintln!(
1747 "block-graph: plan item {} ({}) failed graph preflight",
1748 valid,
1749 match item {
1750 Item::Gdn { run, first } => format!("GDN run L{first}+{}", run.len()),
1751 Item::Attn { li, .. } => format!("Attn L{li}"),
1752 }
1753 );
1754 }
1755 break;
1756 }
1757 valid += 1;
1758 end += match item {
1759 Item::Gdn { run, .. } => run.len(),
1760 Item::Attn { .. } => 1,
1761 };
1762 }
1763 plan.truncate(valid);
1764 if plan.is_empty() {
1765 if self.graph_head_required {
1766 self.fail_metal_graph("native graph preflight produced no valid items");
1767 }
1768 return start;
1769 }
1770
1771 if self.graph_head_required && (upto.is_some() || end != self.num_layers) {
1772 self.fail_metal_graph("fused-head NLL requires a complete 64-layer graph");
1773 return start;
1774 }
1775
1776 let inv_freq = self.inv_freq.clone();
1777 let pool = self.pool.clone();
1778 let (nh, nkv, hd, hs, rd, eps) = (
1779 self.num_heads,
1780 self.num_kv_heads,
1781 self.head_dim,
1782 self.hidden_size,
1783 self.rotary_dim,
1784 self.rms_eps,
1785 );
1786 let norm_style = self.norm_style;
1787 let gemma = norm_style == cortiq_core::NormStyle::Gemma;
1788 let want = self.gdn_cfg.map(|c| c.state_len()).unwrap_or(0);
1789 let kv_id = self.graph_kv_id;
1790 let mut pending: Vec<(usize, usize)> = Vec::new();
1793 let mut dev_attn: Vec<usize> = Vec::new();
1796 for item in &plan {
1797 let _xt0 = std::time::Instant::now();
1798 let _xkind: u32 = match item {
1799 Item::Gdn { .. } => 2,
1800 Item::Attn { .. } => 3,
1801 };
1802 if self.loop_final_norm {
1804 let item_start = match item {
1805 Item::Gdn { first, .. } => *first,
1806 Item::Attn { li, .. } => *li,
1807 };
1808 if item_start > start && self.is_loop_end(item_start - 1) {
1809 graph.encode_loop_norm(&self.weights.final_norm);
1810 }
1811 }
1812 match item {
1813 Item::Gdn { run, first } => {
1814 for l in &mut self.kv_cache.layers[*first..*first + run.len()] {
1815 if l.linear_state.len() != want {
1816 l.linear_state = vec![0f32; want];
1817 }
1818 }
1819 let ro: Vec<&[f32]> = self.kv_cache.layers[*first..*first + run.len()]
1820 .iter()
1821 .map(|l| l.linear_state.as_slice())
1822 .collect();
1823 let _ig = std::time::Instant::now();
1824 if !graph.encode_gdn_run(run, &ro, gcfg.as_ref().unwrap()) {
1825 tracing::error!("q1 graph: GDN run refused after validation");
1827 return start;
1828 }
1829 graph.commit_kind = 2;
1832 graph.commit();
1833 crate::gpu::stageprof(0, _ig.elapsed());
1834 pending.push((*first, run.len()));
1835 }
1836 Item::Attn {
1837 l,
1838 li,
1839 q_norm,
1840 k_norm,
1841 output_gate,
1842 bias,
1843 full_gpu,
1844 } => {
1845 let _ia = std::time::Instant::now();
1846 if *full_gpu {
1848 let cache = &self.kv_cache.layers[*li];
1849 let o1p = if cache.o1.is_some() {
1850 match cache.o1_views() {
1851 Some(views) => Some(crate::gpu::O1AttnParams {
1852 views,
1853 epoch: self.o1_epoch,
1854 }),
1855 None => None,
1857 }
1858 } else {
1859 None
1860 };
1861 let o1_layer = cache.o1.is_some();
1862 if o1_layer && o1p.is_none() {
1863 }
1865 let cpu_k: Vec<&[f32]> = (0..nkv).map(|g| cache.head_keys(g)).collect();
1866 let cpu_v: Vec<&[f32]> = (0..nkv).map(|g| cache.head_values(g)).collect();
1867 let cpu_stored = if o1_layer { 0 } else { cpu_k[0].len() / hd };
1868 let p = crate::gpu::AttnDeviceParams {
1869 kv_id,
1870 layer: *li,
1871 nh,
1872 nkv,
1873 hd,
1874 rd,
1875 position,
1876 scale: self.attn_scale,
1877 eps: eps as f32,
1878 gemma,
1879 late_qk_norm: self.qk_norm_after_rope,
1880 output_gate: *output_gate,
1881 q_norm: *q_norm,
1882 k_norm: *k_norm,
1883 inv_freq: &inv_freq,
1884 cpu_k,
1885 cpu_v,
1886 cpu_stored,
1887 o1: o1p,
1888 };
1889 let o1_bad = o1_layer && p.o1.is_none();
1890 if !o1_bad && graph.attn_device_ok(l, &p) && graph.encode_attn_device(l, &p)
1891 {
1892 if p.o1.is_none() {
1894 dev_attn.push(*li);
1895 }
1896 graph.commit_kind = 3;
1897 graph.commit();
1898 crate::gpu::stageprof(_xkind, _xt0.elapsed());
1902 continue;
1903 }
1904 }
1906 graph.encode_attn_prefix(l);
1907 if let Err(err) = graph.sync_checked() {
1908 self.fail_metal_graph(&err);
1909 return start;
1910 }
1911 if !pending.is_empty() {
1912 let idxs: Vec<usize> =
1913 pending.drain(..).flat_map(|(f, n)| f..f + n).collect();
1914 let mut outs: Vec<&mut [f32]> = self
1915 .kv_cache
1916 .layers
1917 .iter_mut()
1918 .enumerate()
1919 .filter(|(i, _)| idxs.binary_search(i).is_ok())
1920 .map(|(_, s)| s.linear_state.as_mut_slice())
1921 .collect();
1922 graph.read_states(&mut outs);
1923 }
1924 let mut q_raw = attention::take_buf(l.wq.1);
1925 let mut k = attention::take_buf(l.wk.1);
1926 let mut v = attention::take_buf(l.wv.1);
1927 graph.read_qkv(&mut q_raw, &mut k, &mut v);
1928 let cfg = QwenAttnCfg {
1929 num_heads: nh,
1930 num_kv_heads: nkv,
1931 head_dim: hd,
1932 hidden_size: hs,
1933 position,
1934 inv_freq: &inv_freq,
1935 rotary_dim: rd,
1936 scale: self.attn_scale,
1937 softcap: self.attn_softcap,
1938 window: None,
1939 v_norm: false,
1940 qk_norm_after_rope: self.qk_norm_after_rope,
1941 q_norm: *q_norm,
1942 k_norm: *k_norm,
1943 output_gate: *output_gate,
1944 softplus_gate: None,
1945 rope_scale: 1.0,
1946 bias: *bias,
1947 rms_eps: eps,
1948 norm_style,
1949 pool: pool.as_deref(),
1950 };
1951 let oracle = std::env::var("CMF_ATTN_ORACLE").as_deref() == Ok("1")
1954 || std::env::var("CMF_ATTN_DUMP").is_ok();
1955 let _ = full_gpu;
1956 let oracle_in = oracle.then(|| (q_raw.clone(), k.clone(), v.clone()));
1957 let mut ao = attention::qwen_attention_core(
1958 q_raw,
1959 k,
1960 v,
1961 &mut self.kv_cache.layers[*li],
1962 &cfg,
1963 );
1964 if let Ok(dir) = std::env::var("CMF_ATTN_DUMP") {
1968 if let Some((qr0, k0, v0)) = oracle_in.clone() {
1969 let (cq, _cg, _ck, _cv) =
1970 attention::finish_projection_debug(qr0, k0, v0, &cfg, position);
1971 let cache = &self.kv_cache.layers[*li];
1972 let n = cache.head_keys(0).len() / hd;
1973 let mut bytes: Vec<u8> = Vec::new();
1974 for v in [nh as u32, nkv as u32, hd as u32, n as u32, position as u32] {
1975 bytes.extend_from_slice(&v.to_le_bytes());
1976 }
1977 for v in &cq {
1978 bytes.extend_from_slice(&v.to_le_bytes());
1979 }
1980 for g in 0..nkv {
1981 for v in cache.head_keys(g) {
1982 bytes.extend_from_slice(&v.to_le_bytes());
1983 }
1984 }
1985 for g in 0..nkv {
1986 for v in cache.head_values(g) {
1987 bytes.extend_from_slice(&v.to_le_bytes());
1988 }
1989 }
1990 let _ =
1991 std::fs::write(format!("{dir}/L{li}_pos{position}.bin"), &bytes);
1992 }
1993 }
1994 if let Some((qr0, k0, v0)) =
1995 oracle_in.filter(|_| std::env::var("CMF_ATTN_ORACLE").as_deref() == Ok("1"))
1996 {
1997 let (cq, _cg, ck, cv) =
1998 attention::finish_projection_debug(qr0, k0, v0, &cfg, position);
1999 let mut h_now = vec![0f32; hs];
2000 graph.read_h(&mut h_now);
2001 let cache = &self.kv_cache.layers[*li];
2002 let n_after = cache.head_keys(0).len() / hd;
2003 let stored = n_after.saturating_sub(1);
2007 let cpu_k: Vec<&[f32]> = (0..nkv)
2008 .map(|g| &cache.head_keys(g)[..stored * hd])
2009 .collect();
2010 let cpu_v: Vec<&[f32]> = (0..nkv)
2011 .map(|g| &cache.head_values(g)[..stored * hd])
2012 .collect();
2013 let p = crate::gpu::AttnDeviceParams {
2014 kv_id,
2015 layer: *li,
2016 nh,
2017 nkv,
2018 hd,
2019 rd,
2020 position,
2021 scale: self.attn_scale,
2022 eps: eps as f32,
2023 gemma,
2024 late_qk_norm: self.qk_norm_after_rope,
2025 output_gate: *output_gate,
2026 q_norm: *q_norm,
2027 k_norm: *k_norm,
2028 inv_freq: &inv_freq,
2029 cpu_k,
2030 cpu_v,
2031 cpu_stored: stored,
2032 o1: None,
2033 };
2034 if let Some((dq, dk, dv, dao)) = graph.debug_attn_device(l, &p, &h_now) {
2035 let md = |a: &[f32], b: &[f32]| {
2036 a.iter().zip(b).fold(0f32, |m, (x, y)| m.max((x - y).abs()))
2037 };
2038 let nn = |a: &[f32]| a.iter().map(|x| x * x).sum::<f32>().sqrt();
2039 eprintln!(
2040 "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}",
2041 nn(&cq),
2042 md(&cq, &dq),
2043 nn(&ck),
2044 md(&ck, &dk),
2045 nn(&cv),
2046 md(&cv, &dv),
2047 nn(&ao),
2048 md(&ao, &dao)
2049 );
2050 } else {
2051 eprintln!("attn-oracle L{li}: device probe declined");
2052 }
2053 }
2054 graph.encode_attn_suffix(l, &ao);
2055 graph.commit();
2058 attention::recycle_buf(&mut ao);
2059 }
2060 }
2061
2062 crate::gpu::stageprof(_xkind, _xt0.elapsed());
2063 }
2064 let mut lm_rows = None;
2069 if self.graph_want_logits
2070 && upto.is_none()
2071 && end == self.num_layers
2072 && std::env::var("CMF_GPU_LMHEAD")
2073 .map(|v| v != "0")
2074 .unwrap_or(true)
2075 {
2076 if let Some(lm) = self.weights.lm_head.metal_graph_parts() {
2077 if graph.lm_head_ok(lm) {
2078 graph.encode_lm_head(&self.weights.final_norm, lm);
2079 lm_rows = Some(lm.1);
2080 }
2081 }
2082 }
2083 if self.graph_head_required && lm_rows.is_none() {
2084 METAL_GRAPH_HEAD_MISS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
2085 self.fail_metal_graph("fused graph head was requested but not encodable");
2086 return start;
2087 }
2088 let _sy0 = std::time::Instant::now();
2089 if let Err(err) = graph.sync_checked() {
2090 self.fail_metal_graph(&err);
2091 return start;
2092 }
2093 let _rs0 = std::time::Instant::now();
2094 if !pending.is_empty() {
2095 let idxs: Vec<usize> = pending.drain(..).flat_map(|(f, n)| f..f + n).collect();
2096 let mut outs: Vec<&mut [f32]> = self
2097 .kv_cache
2098 .layers
2099 .iter_mut()
2100 .enumerate()
2101 .filter(|(i, _)| idxs.binary_search(i).is_ok())
2102 .map(|(_, s)| s.linear_state.as_mut_slice())
2103 .collect();
2104 graph.read_states(&mut outs);
2105 }
2106 if std::env::var("CMF_GRAPH_HOSTPROF").as_deref() == Ok("1") {
2107 use std::sync::atomic::{AtomicU64, Ordering};
2108 static SY: AtomicU64 = AtomicU64::new(0);
2109 static RS: AtomicU64 = AtomicU64::new(0);
2110 static N: AtomicU64 = AtomicU64::new(0);
2111 SY.fetch_add((_rs0 - _sy0).as_nanos() as u64, Ordering::Relaxed);
2112 RS.fetch_add(_rs0.elapsed().as_nanos() as u64, Ordering::Relaxed);
2113 let n = N.fetch_add(1, Ordering::Relaxed) + 1;
2114 if n % 100 == 0 {
2115 eprintln!(
2116 "postprof: sync-wait {:.1} ms/ток | read_states {:.1} ms/ток ({n})",
2117 SY.load(Ordering::Relaxed) as f64 / n as f64 / 1e6,
2118 RS.load(Ordering::Relaxed) as f64 / n as f64 / 1e6
2119 );
2120 }
2121 }
2122 if let Some(rows) = lm_rows {
2123 crate::gpu::hostprof_encode_done(_mt0);
2124 let mut lg = attention::take_buf(rows.min(self.vocab_size));
2125 graph.read_logits(&mut lg);
2126 crate::gpu::hostprof_total(_mt0);
2127 lg.resize(self.vocab_size, 0.0);
2128 if let Some(c) = self.final_softcap {
2129 for l in lg.iter_mut() {
2130 *l = c * (*l / c).tanh();
2131 }
2132 }
2133 self.graph_logits = Some(lg);
2134 }
2135 graph.read_h(h);
2136 if self.graph_head_required && self.graph_logits.is_none() {
2137 METAL_GRAPH_HEAD_MISS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
2138 self.fail_metal_graph("fused graph head completed without logits readback");
2139 return start;
2140 }
2141 METAL_GRAPH_TOK_OK.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
2142 METAL_GRAPH_LAYERS.fetch_add(
2143 end.saturating_sub(start) as u64,
2144 std::sync::atomic::Ordering::Relaxed,
2145 );
2146 if self.graph_head_required {
2147 METAL_GRAPH_HEAD_OK.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
2148 }
2149 for li in dev_attn {
2153 let mut krow = attention::take_buf(nkv * hd);
2154 let mut vrow = attention::take_buf(nkv * hd);
2155 if crate::gpu::kv_mirror_read_last(kv_id, li, nkv, hd, &mut krow, &mut vrow) {
2156 let cache = &mut self.kv_cache.layers[li];
2157 cache.append(&krow, &vrow, &[]);
2158 let n = cache.seq_len;
2159 let mut imp = attention::take_buf(n);
2160 crate::gpu::kv_mirror_take_imp(kv_id, li, &mut imp);
2161 cache.accumulate_imp(&imp);
2162 attention::recycle_buf(&mut imp);
2163 }
2164 attention::recycle_buf(&mut krow);
2165 attention::recycle_buf(&mut vrow);
2166 }
2167 end
2168 }
2169
2170 pub fn new(
2171 tokenizer: Tokenizer,
2172 weights: PipelineWeights,
2173 hidden_size: usize,
2174 intermediate_size: usize,
2175 num_heads: usize,
2176 num_kv_heads: usize,
2177 head_dim: usize,
2178 num_layers: usize,
2179 physical_layers: usize,
2180 loop_final_norm: bool,
2181 vocab_size: usize,
2182 rms_eps: f64,
2183 rope_base: f32,
2184 norm_style: NormStyle,
2185 max_seq_len: usize,
2186 sampler_config: SamplerConfig,
2187 ) -> Self {
2188 let rng = match sampler_config.seed {
2189 Some(s) => SplitMix64::new(s),
2190 None => SplitMix64::from_entropy(),
2191 };
2192 let inv_freq = std::sync::Arc::new(attention::rope_inv_freq(head_dim, rope_base));
2193 let pool = Pool::from_env();
2194 if let Some(p) = &pool {
2195 tracing::info!("worker pool: {} threads", p.n_workers());
2196 }
2197 Self {
2198 gpu_plan: None,
2199 tokenizer: std::sync::Arc::new(tokenizer),
2200 kv_cache: KvCache::new(num_layers, num_kv_heads, head_dim, max_seq_len),
2201 sampler_config,
2202 weights,
2203 hidden_size,
2204 intermediate_size,
2205 num_heads,
2206 num_kv_heads,
2207 head_dim,
2208 num_layers,
2209 physical_layers,
2210 loop_final_norm,
2211 vocab_size,
2212 rms_eps,
2213 rope_base,
2214 norm_style,
2215 rotary_dim: head_dim,
2216 attention_heads_per_layer: None,
2217 vmf_cfg: None,
2218 gdn_cfg: None,
2219 kda_cfg: None,
2220 g3n: None,
2221 dsv4: None,
2222 dsv41: None,
2223 dsv41_vision: None,
2224 dsv41_prefill: None,
2225 qwen4_exp: None,
2226 dsv4_mtp: Vec::new(),
2227 dspark: None,
2228 dspark_pending: Vec::new(),
2229 dspark_hist: Vec::new(),
2230 dspark_real: Vec::new(),
2231 dspark_trunk_picks: Vec::new(),
2232 dspark_exp: Vec::new(),
2233 dspark_draft_ns: 0,
2234 logit_multiplier: None,
2235 cancel: std::sync::Arc::new(std::sync::atomic::AtomicBool::new(false)),
2236 graph_failed: std::sync::atomic::AtomicBool::new(false),
2237 kv_history: Vec::new(),
2238 short_conv_cfg: None,
2239 mtp: None,
2240 speculative: std::env::var("CMF_MTP").map(|v| v != "0").unwrap_or(true),
2241 ignore_eos: false,
2242 draft_full_streak: 0,
2243 spec_k_adapt: None,
2244 spec_acc_ewma: 0.7,
2245 rng,
2246 sampler_scratch: SamplerScratch::default(),
2247 spec_forced: None,
2248 spec_q: Vec::new(),
2249 spec_p: Vec::new(),
2250 spec_res: Vec::new(),
2251 spec_qs: Vec::new(),
2252 spec_ps: Vec::new(),
2253 spec_ress: Vec::new(),
2254 mtp_graph_mode: None,
2255 #[cfg(target_os = "macos")]
2256 metal_verify: None,
2257 inv_freq,
2258 ws: ForwardScratch::new(hidden_size),
2259 pool,
2260 model: None,
2261 dyn_force_f32: false,
2262 dyn_skill_layers: Vec::new(),
2263 dyn_active: None,
2264 dyn_blend_loaded: false,
2265 dyn_phi_layer: None,
2266 dyn_phi_ema: Vec::new(),
2267 dyn_phi_seen: 0,
2268 dyn_router: None,
2269 o1_cfg: None,
2270 o1_epoch: 0,
2271 o1_flags: Vec::new(),
2272 trace: false,
2273 calib_temp: 1.0,
2274 confidence_on: true,
2275 embed_multiplier: 1.0,
2276 attn_scale: 1.0 / (head_dim as f32).sqrt(),
2277 swa: None,
2278 sliding_layers: None,
2279 inv_freq_local: None,
2280 rotary_dim_local: None,
2281 rope_scale: 1.0,
2282 rope_scale_local: 1.0,
2283 global_attn: None,
2284 inv_freq_global: None,
2285 attn_v_norm: false,
2286 qk_norm_after_rope: false,
2287 final_softcap: None,
2288 head_clusters: None,
2289 attn_softcap: 0.0,
2290 graph_want_logits: false,
2291 graph_head_required: false,
2292 graph_logits: None,
2293 graph_kv_id: {
2294 static NEXT: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(1);
2295 NEXT.fetch_add(1, std::sync::atomic::Ordering::Relaxed)
2296 },
2297 #[cfg(test)]
2298 nll_test_fail_at: None,
2299 #[cfg(test)]
2300 nll_test_force_serial: false,
2301 }
2302 }
2303
2304 pub fn set_o1(&mut self, cfg: Option<crate::nystrom::O1Cfg>) {
2312 if let Some(c) = &cfg {
2313 if crate::nystrom::o1_deferred_boundary(c.w, c.sink).is_none() {
2314 tracing::error!(
2315 "o1 disabled: w + sink + slack + 1 overflows usize (w={}, sink={})",
2316 c.w,
2317 c.sink
2318 );
2319 self.o1_flags.clear();
2320 self.o1_cfg = None;
2321 return;
2322 }
2323 }
2324 self.o1_flags = match &cfg {
2325 Some(c) => {
2326 let mut flags = c.layer_flags(self.num_layers);
2327 for (li, f) in flags.iter_mut().enumerate() {
2328 if *f
2329 && !matches!(
2330 self.weights.layers[self.phys_layer(li)].attn,
2331 AttnKind::Full { .. }
2332 )
2333 {
2334 *f = false;
2335 }
2336 }
2337 flags
2338 }
2339 None => Vec::new(),
2340 };
2341 if let Some(c) = &cfg {
2342 let n = self.o1_flags.iter().filter(|&&f| f).count();
2343 tracing::info!(
2344 "o1 nystrom attention: {n}/{} layer(s), m={} w={} sink={} rect={:?}",
2345 self.num_layers,
2346 c.m,
2347 c.w,
2348 c.sink,
2349 c.rect
2350 );
2351 }
2352 self.o1_cfg = cfg;
2353 }
2354
2355 pub fn o1_active(&self) -> bool {
2357 self.o1_cfg.is_some() && self.o1_flags.iter().any(|&f| f)
2358 }
2359
2360 pub fn generation_batch_k(&self) -> usize {
2372 if let Some(k) = std::env::var("CMF_BATCH_K")
2373 .ok()
2374 .and_then(|v| v.parse::<usize>().ok())
2375 {
2376 return k;
2377 }
2378 #[cfg(not(target_os = "macos"))]
2379 if self.graph_prefill_preferred() && !self.o1_active() {
2380 return 32;
2381 }
2382 0
2383 }
2384
2385 pub fn generation_graph_prefill(&self) -> bool {
2386 let graph = self.graph_prefill_preferred();
2387 #[cfg(not(target_os = "macos"))]
2398 if graph
2399 && self.generation_batch_k() > 0
2400 && std::env::var("CMF_MTP_CHAIN_PROBE").is_err()
2401 {
2402 return false;
2403 }
2404 graph
2405 }
2406
2407 pub fn o1_device_stats(&self) -> (usize, u64) {
2412 crate::gpu::o1_device_stats(self.graph_kv_id)
2413 }
2414
2415 pub fn o1_begin(&mut self) {
2420 self.o1_begin_with_prefix(None);
2421 }
2422
2423 pub fn o1_begin_with_prefix(&mut self, requested_prefix: Option<usize>) {
2427 if let Some(c) = &self.o1_cfg {
2428 let (m, w, sink, rect) = (c.m, c.w, c.sink, c.rect);
2429 let boundary = requested_prefix.map(|p| {
2430 p.max(
2431 crate::nystrom::o1_deferred_boundary(w, sink)
2432 .expect("o1 config boundary validated in set_o1"),
2433 )
2434 });
2435 for (li, &f) in self.o1_flags.iter().enumerate() {
2436 if f {
2437 self.kv_cache.layers[li].o1_begin_with_boundary(m, w, sink, rect, boundary);
2438 }
2439 }
2440 }
2441 }
2442
2443 fn o1_effective_boundary(&self, requested_prefix: usize) -> Option<usize> {
2445 self.o1_cfg.as_ref().and_then(|c| {
2446 crate::nystrom::o1_deferred_boundary(c.w, c.sink)
2447 .map(|floor| requested_prefix.max(floor))
2448 })
2449 }
2450
2451 fn o1_note_transition(&mut self) {
2452 let mut transitioned = false;
2456 for (li, &flagged) in self.o1_flags.iter().enumerate() {
2457 if flagged {
2458 transitioned |= self.kv_cache.layers[li].take_o1_transition();
2459 }
2460 }
2461 if transitioned {
2462 self.o1_epoch = self.o1_epoch.wrapping_add(1);
2463 }
2464 }
2465
2466 fn o1_pending(&self) -> bool {
2467 self.o1_flags.iter().enumerate().any(|(li, &f)| {
2468 f && self.kv_cache.layers[li].seq_len > 0
2469 && self.kv_cache.layers[li].o1_pending_boundary().is_some()
2470 })
2471 }
2472
2473 fn o1_fail(&mut self, err: String) {
2474 tracing::error!("o1 deferred seal failed; terminating sequence: {err}");
2475 self.clear_sequence_state();
2476 self.graph_failed
2477 .store(true, std::sync::atomic::Ordering::Relaxed);
2478 self.cancel
2479 .store(true, std::sync::atomic::Ordering::Relaxed);
2480 }
2481
2482 pub fn o1_seal_checked(&mut self) -> Result<bool, String> {
2487 if self.o1_cfg.is_none() {
2488 return Ok(false);
2489 }
2490 let mut participating = false;
2491 for li in 0..self.num_layers {
2492 if !self.o1_flags.get(li).copied().unwrap_or(false) {
2493 continue;
2494 }
2495 if let Some(err) = self.kv_cache.layers[li].take_o1_error() {
2496 return Err(err);
2497 }
2498 if self.kv_cache.layers[li].seq_len == 0 {
2499 continue;
2500 }
2501 participating = true;
2502 let num_heads = self.layer_num_heads(li);
2503 self.kv_cache.layers[li].o1_seal_checked(num_heads)?;
2504 }
2505 self.o1_note_transition();
2506 for li in 0..self.num_layers {
2507 if self.o1_flags.get(li).copied().unwrap_or(false) {
2508 if let Some(err) = self.kv_cache.layers[li].take_o1_error() {
2509 return Err(err);
2510 }
2511 }
2512 }
2513 Ok(participating
2514 && (0..self.num_layers).all(|li| {
2515 !self.o1_flags.get(li).copied().unwrap_or(false)
2516 || self.kv_cache.layers[li].seq_len == 0
2517 || self.kv_cache.layers[li].o1_sealed()
2518 }))
2519 }
2520
2521 fn o1_progress(&mut self) {
2524 if !self.o1_active() {
2525 return;
2526 }
2527 for li in 0..self.num_layers {
2528 if self.o1_flags.get(li).copied().unwrap_or(false) {
2529 if let Some(err) = self.kv_cache.layers[li].take_o1_error() {
2530 self.o1_fail(err);
2531 return;
2532 }
2533 }
2534 }
2535 self.o1_note_transition();
2539 if !self.o1_pending() {
2540 return;
2541 }
2542 if let Err(err) = self.o1_seal_checked() {
2543 self.o1_fail(err);
2544 }
2545 }
2546
2547 fn check_o1_progress_failure(&mut self, phase: &str) -> Result<(), String> {
2552 if self
2553 .graph_failed
2554 .swap(false, std::sync::atomic::Ordering::Relaxed)
2555 {
2556 self.cancel
2557 .store(false, std::sync::atomic::Ordering::Relaxed);
2558 self.clear_sequence_state();
2559 return Err(format!("{phase}: deferred O(1) transition failed"));
2560 }
2561 Ok(())
2562 }
2563
2564 pub fn o1_seal(&mut self) {
2568 if let Err(err) = self.o1_seal_checked() {
2569 self.o1_fail(err);
2570 }
2571 }
2572
2573 pub fn set_trace(&mut self, on: bool) {
2575 self.trace = on;
2576 }
2577
2578 pub fn set_sampler_config(&mut self, config: SamplerConfig) {
2581 self.rng = match config.seed {
2582 Some(seed) => SplitMix64::new(seed),
2583 None => SplitMix64::from_entropy(),
2584 };
2585 self.sampler_config = config;
2586 }
2587
2588 pub fn set_confidence(&mut self, on: bool) {
2593 self.confidence_on = on;
2594 }
2595
2596 pub fn set_calib_temp(&mut self, t: f32) {
2599 self.calib_temp = if t > 1e-3 { t } else { 1.0 };
2600 }
2601
2602 pub fn calib_temp(&self) -> f32 {
2604 self.calib_temp
2605 }
2606
2607 pub fn set_rotary(&mut self, rotary_dim: usize, base: f32) {
2610 self.rotary_dim = rotary_dim.min(self.head_dim);
2611 self.inv_freq = std::sync::Arc::new(attention::rope_inv_freq(self.rotary_dim, base));
2612 }
2613
2614 fn attn_cfg(&self, position: usize) -> QwenAttnCfg<'_> {
2615 QwenAttnCfg {
2616 num_heads: self.num_heads,
2617 num_kv_heads: self.num_kv_heads,
2618 head_dim: self.head_dim,
2619 hidden_size: self.hidden_size,
2620 position,
2621 inv_freq: &self.inv_freq,
2622 rotary_dim: self.rotary_dim,
2623 scale: self.attn_scale,
2624 softcap: self.attn_softcap,
2625 window: None,
2626 v_norm: false,
2627 qk_norm_after_rope: self.qk_norm_after_rope,
2628 q_norm: None,
2629 k_norm: None,
2630 output_gate: false,
2631 softplus_gate: None,
2632 rope_scale: self.rope_scale,
2633 bias: None,
2634 rms_eps: self.rms_eps,
2635 norm_style: self.norm_style,
2636 pool: self.pool.as_deref(),
2637 }
2638 }
2639
2640 pub fn generate(
2642 &mut self,
2643 prompt: &str,
2644 max_tokens: usize,
2645 task_mask: Option<&TaskMask>,
2646 on_token: Option<TokenCallback>,
2647 ) -> Result<GenerateResult, String> {
2648 let input_ids = self.tokenizer.with_bos(self.tokenizer.encode(prompt));
2649 self.generate_from_ids(&input_ids, max_tokens, task_mask, on_token)
2650 }
2651
2652 pub fn generate_from_vl(
2655 &mut self,
2656 input: &crate::dsv41_vision::PreparedVlInputs,
2657 max_tokens: usize,
2658 task_mask: Option<&TaskMask>,
2659 on_token: Option<TokenCallback>,
2660 ) -> Result<GenerateResult, String> {
2661 let Some(dsv41) = &self.dsv41 else {
2662 return Err("V4.1 multimodal input requires a DeepSeek-V4.1 pipeline".into());
2663 };
2664 if input.token_ids.is_empty() {
2665 return Err("empty V4.1 multimodal prompt".into());
2666 }
2667 if input.token_types.len() != input.token_ids.len() {
2668 return Err(format!(
2669 "V4.1 token type count {} != token count {}",
2670 input.token_types.len(),
2671 input.token_ids.len()
2672 ));
2673 }
2674 let dim = dsv41.2.dim;
2675 let mut embeddings = vec![None; input.token_ids.len()];
2676 let mut participates = vec![true; input.token_ids.len()];
2677 if !input.images.is_empty() {
2678 let vision = self
2679 .dsv41_vision
2680 .as_ref()
2681 .ok_or_else(|| "V4.1 image prompt has no loaded vision tower".to_string())?;
2682 for image in &input.images {
2683 let end = image.start.saturating_add(image.types.len());
2684 if end > input.token_ids.len() {
2685 return Err(format!(
2686 "V4.1 image span {}..{} exceeds prompt length {}",
2687 image.start,
2688 end,
2689 input.token_ids.len()
2690 ));
2691 }
2692 let mut span = vec![0.0f32; image.types.len() * dim];
2693 vision.fill_image_span(image, &mut span, self.pool.as_deref())?;
2694 for (offset, &kind) in image.types.iter().enumerate() {
2695 let pos = image.start + offset;
2696 if input.token_types[pos] != kind {
2697 return Err(format!(
2698 "V4.1 image type mismatch at position {pos}: {} != {kind}",
2699 input.token_types[pos]
2700 ));
2701 }
2702 embeddings[pos] = Some(span[offset * dim..(offset + 1) * dim].to_vec());
2703 participates[pos] = false;
2704 }
2705 }
2706 }
2707 for (pos, &kind) in input.token_types.iter().enumerate() {
2708 if kind == crate::dsv41_vision::TEXT && embeddings[pos].is_some() {
2709 return Err(format!("V4.1 text position {pos} has an image embedding"));
2710 }
2711 if kind != crate::dsv41_vision::TEXT && embeddings[pos].is_none() {
2712 return Err(format!("V4.1 image position {pos} has no image embedding"));
2713 }
2714 }
2715 self.dsv41_prefill = Some((embeddings, participates));
2716 let result = self.generate_from_ids(&input.token_ids, max_tokens, task_mask, on_token);
2717 self.dsv41_prefill = None;
2718 result
2719 }
2720
2721 fn drop_open_mask<'m>(&self, m: Option<&'m TaskMask>) -> Option<&'m TaskMask> {
2723 m.filter(|m| !m.fully_open(self.intermediate_size, self.num_heads))
2724 }
2725
2726 pub fn generate_from_ids(
2734 &mut self,
2735 input_ids: &[u32],
2736 max_tokens: usize,
2737 task_mask: Option<&TaskMask>,
2738 mut on_token: Option<TokenCallback>,
2739 ) -> Result<GenerateResult, String> {
2740 if std::env::var("CMF_TRACE_H").is_ok() {
2741 eprintln!("input_ids: {input_ids:?}");
2742 }
2743 if input_ids.is_empty() {
2744 return Err("empty prompt: nothing to generate from".to_string());
2745 }
2746 self.graph_failed
2750 .store(false, std::sync::atomic::Ordering::Relaxed);
2751 let task_mask = self.drop_open_mask(task_mask);
2756
2757 let reuse_from = {
2765 let on = !std::env::var("CMF_KV_REUSE").is_ok_and(|v| v == "0");
2766 let h = &self.kv_history;
2767 if on
2768 && task_mask.is_none()
2769 && self.mtp.is_none()
2770 && self.o1_cfg.is_none()
2771 && self.dsv41.is_none()
2772 && !h.is_empty()
2773 && h.len() < input_ids.len()
2774 && input_ids[..h.len()] == h[..]
2775 {
2776 h.len()
2777 } else {
2778 0
2779 }
2780 };
2781 if reuse_from == 0 {
2782 self.clear_sequence_state();
2784 } else if std::env::var("CMF_PREFILL_PROF").is_ok() {
2785 eprintln!(
2786 "kv-reuse: {} of {} prompt positions already cached",
2787 reuse_from,
2788 input_ids.len()
2789 );
2790 }
2791 crate::gpu::graph_race_begin_generation();
2792 let o1_prefill = if self.o1_active() && task_mask.is_none() {
2796 std::env::var("CMF_O1_PREFILL")
2797 .ok()
2798 .and_then(|v| v.parse::<usize>().ok())
2799 .filter(|&p| p > 0)
2800 } else {
2801 None
2802 };
2803 if task_mask.is_none() {
2804 self.o1_begin_with_prefix(o1_prefill);
2805 }
2806
2807 let graph_on = crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Decode);
2813 #[cfg(target_os = "macos")]
2852 let metal_graph = crate::gpu::q1_force()
2853 && crate::gpu::enabled_here()
2854 && std::env::var("CMF_GPU_BLOCK")
2855 .map(|v| v != "0")
2856 .unwrap_or(true);
2857 #[cfg(not(target_os = "macos"))]
2858 let metal_graph = false;
2859 let spec_sample_env = std::env::var("CMF_GRAPH_SPEC_SAMPLE").ok();
2860 let spec_cheap_round = self.sampler_config.temperature < 1e-6
2864 || sampler::sparse_ok(&self.sampler_config);
2865 let spec_sampling_ok = self.sampler_config.temperature < 1e-6
2866 || match spec_sample_env.as_deref() {
2867 Some("1") => true,
2868 Some(_) => false,
2869 None => metal_graph && spec_cheap_round,
2870 };
2871 let (mut dense_n, mut dense_q4tp) = (0usize, 0usize);
2887 for lw in &self.weights.layers {
2888 if let FfnKind::Dense(d) = &lw.ffn {
2889 dense_n += 1;
2890 if matches!(d.gate_proj.graph_weight(), Some((_, _, 6, _)))
2891 && matches!(d.up_proj.graph_weight(), Some((_, _, 6, _)))
2892 && matches!(d.down_proj.graph_weight(), Some((_, _, 6, _)))
2893 {
2894 dense_q4tp += 1;
2895 }
2896 }
2897 }
2898 let spec_default_ok = dense_n == 0 || dense_q4tp * 10 >= dense_n * 9;
2899 let penalized = !metal_graph
2911 && (self.sampler_config.repetition_penalty != 1.0
2912 || self.sampler_config.presence_penalty != 0.0
2913 || !self.sampler_config.suppress_tokens.is_empty());
2914 #[cfg(feature = "gpu")]
2919 let metal_wgpu = graph_on && crate::gpu_wgpu::wgpu_backend_is_metal();
2920 #[cfg(not(feature = "gpu"))]
2921 let metal_wgpu = false;
2922 let spec_env = std::env::var("CMF_GRAPH_SPEC").ok();
2923 let spec_wanted = match spec_env.as_deref() {
2924 Some("0") => false,
2925 Some(_) => {
2926 if metal_wgpu {
2927 tracing::warn!(
2928 "CMF_GRAPH_SPEC forced on wgpu/Metal: the batched verify graph is not \
2929 verified on this backend (garbage measured on Qwen3.5-0.8B)"
2930 );
2931 }
2932 true
2933 }
2934 None => spec_default_ok && !penalized && !metal_wgpu,
2935 };
2936 let graph_spec = self.speculative
2940 && (graph_on || metal_graph)
2941 && self.mtp.is_some()
2942 && task_mask.is_none()
2943 && !self.o1_active()
2944 && spec_sampling_ok
2945 && spec_wanted;
2946 #[cfg(target_os = "macos")]
2950 if metal_graph {
2951 static SAID: std::sync::Once = std::sync::Once::new();
2952 SAID.call_once(|| {
2953 let spec = if graph_spec {
2954 let k = std::env::var("CMF_GRAPH_SPEC_K")
2955 .ok()
2956 .and_then(|v| v.parse::<usize>().ok())
2957 .filter(|&v| (1..=8).contains(&v))
2958 .unwrap_or(7);
2959 let arm = if self.sampler_config.temperature < 1e-6 {
2960 "greedy"
2961 } else {
2962 "sampling"
2963 };
2964 format!(
2965 "spec k={k} {arm} (batched verify, draft shortlist {}, trial: proxy)",
2966 Self::draft_vocab_rows(usize::MAX)
2967 )
2968 } else if !self.speculative {
2969 "spec off (CMF_MTP=0)".to_string()
2970 } else if self.mtp.is_none() {
2971 "spec off (no MTP head)".to_string()
2972 } else if !spec_sampling_ok {
2973 if spec_cheap_round {
2974 "spec off (CMF_GRAPH_SPEC_SAMPLE=0)".to_string()
2975 } else {
2976 "spec off (sampling without a top-k: the dense chain \
2977 costs more than it saves)"
2978 .to_string()
2979 }
2980 } else if !spec_wanted {
2981 "spec off (CMF_GRAPH_SPEC=0 or non-q4tp FFNs)".to_string()
2982 } else if task_mask.is_some() {
2983 "spec off (task mask)".to_string()
2984 } else {
2985 "spec off (O(1) attention)".to_string()
2986 };
2987 let on = |var: &str| {
2988 if std::env::var(var).as_deref() == Ok("0") {
2989 "off"
2990 } else {
2991 "on"
2992 }
2993 };
2994 tracing::info!(
2995 "metal native: {spec}, state4 {}, async replay {}, prefill graph {}, \
2996 MTP graph {}, attend {}, probe {}",
2997 if crate::gpu_metal::state4_on() { "on" } else { "off" },
2998 if crate::gpu_metal::async_replay_on() { "on" } else { "off" },
2999 on("CMF_METAL_PREFILL"),
3000 on("CMF_MTP_GRAPH"),
3001 std::env::var("CMF_GPU_ATTEND").unwrap_or_else(|_| "auto".into()),
3002 if crate::gpu::probe_enabled() { "bypassed (q1 force)" } else { "off" },
3003 );
3004 });
3005 }
3006 let pair_pays = self.gdn_cfg.is_none() || std::env::var("CMF_MTP").as_deref() == Ok("1");
3013 let spec_active = self.speculative
3014 && self.mtp.is_some()
3015 && task_mask.is_none()
3016 && !self.o1_active()
3017 && ((!graph_on && pair_pays && self.sampler_config.temperature < 1e-6) || graph_spec);
3018 let mut mtp = if spec_active { self.mtp.take() } else { None };
3021 if std::env::var("CMF_MTP_CHAIN_PROBE").is_ok() {
3022 eprintln!(
3023 "mtp-probe gate: spec_active={spec_active} mtp={} speculative={} graph_on={graph_on} temp_ok={}",
3024 mtp.is_some(),
3025 self.speculative,
3026 self.sampler_config.temperature < 1e-6,
3027 );
3028 }
3029 if let Some(m) = &mut mtp {
3030 m.kv.clear();
3031 crate::gpu::graph_kv_reset(self.mtp_kv_id());
3033 self.mtp_graph_mode = None;
3034 }
3035 let mut router = if mtp.is_none() {
3039 self.dyn_router.take()
3040 } else {
3041 None
3042 };
3043 if let Some(r) = &mut router {
3044 r.reset(); self.dyn_phi_seen = 0; let _ = self.set_active_skill(None);
3047 }
3048
3049 let mut all_ids = input_ids.to_vec();
3050 let mut generated = 0usize;
3051 let mut finish_reason = "max_tokens".to_string();
3052 let mut drafted = 0usize;
3053 let mut accepted = 0usize;
3054 let mut dsv4_spec_bad = 0usize;
3061 let mut dsv4_spec_retry_at = 0usize;
3062 let mut confidence: Vec<f32> = Vec::new();
3063 let trace_on = self.trace;
3064 let calib_temp = self.calib_temp;
3065 let mut traces: Vec<TokenTrace> = Vec::new();
3066
3067 let mut hidden = vec![0.0f32; self.hidden_size];
3073 let mut pos = reuse_from;
3074 let fuse_lm = mtp.is_none()
3083 && router.is_none()
3084 && std::env::var("CMF_GPU_LMHEAD").as_deref() != Ok("0");
3085 self.graph_logits = None;
3086 self.graph_want_logits = false;
3087 let _tpf = std::time::Instant::now();
3088 let batch_k = self.generation_batch_k();
3089 while self.qwen4_exp.is_some()
3100 && mtp.is_none()
3101 && pos < input_ids.len()
3102 && !self.cancel.load(std::sync::atomic::Ordering::Relaxed)
3103 {
3104 let token_id = input_ids[pos];
3105 let want_logits = pos + 1 == input_ids.len();
3106 let mut lg = Vec::new();
3107 if let Some(b) = &mut self.qwen4_exp {
3108 crate::qwen4_exp::forward_token(
3109 &b.0,
3110 &b.1,
3111 &b.2,
3112 &mut b.3,
3113 token_id,
3114 pos,
3115 &self.inv_freq,
3116 self.pool.as_deref(),
3117 &mut lg,
3118 want_logits,
3119 );
3120 }
3121 if want_logits {
3122 self.graph_logits = Some(lg);
3123 }
3124 pos += 1;
3125 hidden.fill(0.0);
3126 }
3127 while self.dsv4.is_some()
3128 && mtp.is_none()
3129 && pos < input_ids.len()
3130 && !self.cancel.load(std::sync::atomic::Ordering::Relaxed)
3131 {
3132 let end = (pos + prefill_chunk()).min(input_ids.len());
3133 let ids: Vec<u32> = input_ids[pos..end].to_vec();
3134 let mut lg = Vec::new();
3135 if let Some(b) = &mut self.dsv4 {
3136 let (g, layers, cfg, st) = (&b.0, &b.1, b.2, &mut b.3);
3137 crate::dsv4::forward_chunk(
3138 g,
3139 layers,
3140 &cfg,
3141 st,
3142 &ids,
3143 pos,
3144 &self.inv_freq,
3145 self.pool.as_deref(),
3146 &mut lg,
3147 end == input_ids.len(),
3148 );
3149 }
3150 if end == input_ids.len() {
3151 self.graph_logits = Some(lg);
3152 }
3153 pos = end;
3154 hidden = vec![0.0; self.hidden_size];
3155 }
3156 let dsv41_prefill = self.dsv41_prefill.take();
3157 while self.dsv41.is_some()
3158 && mtp.is_none()
3159 && pos < input_ids.len()
3160 && !self.cancel.load(std::sync::atomic::Ordering::Relaxed)
3161 {
3162 let end = (pos + prefill_chunk()).min(input_ids.len());
3163 let ids: Vec<u32> = input_ids[pos..end].to_vec();
3164 let mut lg = Vec::new();
3165 if let Some(b) = &mut self.dsv41 {
3166 let (g, layers, cfg, st) = (&b.0, &b.1, &b.2, &mut b.3);
3167 if let Some((embeddings, participates)) = dsv41_prefill.as_ref() {
3168 crate::dsv41::forward_chunk_masked_with_embeddings(
3169 g,
3170 layers,
3171 cfg,
3172 st,
3173 &ids,
3174 pos,
3175 &embeddings[pos..end],
3176 &participates[pos..end],
3177 self.pool.as_deref(),
3178 &mut lg,
3179 );
3180 } else {
3181 crate::dsv41::forward_chunk(
3182 g,
3183 layers,
3184 cfg,
3185 st,
3186 &ids,
3187 pos,
3188 self.pool.as_deref(),
3189 &mut lg,
3190 );
3191 }
3192 }
3193 if end == input_ids.len() {
3194 self.graph_logits = Some(lg);
3195 }
3196 pos = end;
3197 hidden = vec![0.0; self.hidden_size];
3198 }
3199 let dyn_prefill = router.is_some();
3204 let o1_prefill_limit = o1_prefill
3212 .and_then(|requested| self.o1_effective_boundary(requested))
3213 .map(|boundary| boundary.min(input_ids.len()));
3214 let mut o1_sealed = false;
3215 if let Some(limit) = o1_prefill_limit {
3216 if self.can_prefill_batched() && limit > 2 {
3219 let chunk = prefill_chunk();
3220 let hs = self.hidden_size;
3221 while pos < limit && !self.cancel.load(std::sync::atomic::Ordering::Relaxed) {
3222 let end = (pos + chunk).min(limit);
3223 let hb = self.prefill_batch(&input_ids[pos..end], pos);
3224 hidden.copy_from_slice(&hb[(end - pos - 1) * hs..]);
3225 pos = end;
3226 }
3227 } else {
3228 while pos < limit && !self.cancel.load(std::sync::atomic::Ordering::Relaxed) {
3229 hidden = self.forward_layers(&self.embed_single(input_ids[pos]), pos, None);
3230 pos += 1;
3231 }
3232 }
3233 if pos >= limit {
3234 o1_sealed = match self.o1_seal_checked() {
3235 Ok(sealed) => sealed,
3236 Err(err) => {
3237 self.finish_generation(&mut mtp, &mut router, true);
3238 return Err(err);
3239 }
3240 };
3241 tracing::info!(
3242 "o1 bounded prompt prefix: requested={} effective={} processed={} of {} token(s)",
3243 o1_prefill.unwrap_or(0),
3244 self.o1_effective_boundary(o1_prefill.unwrap_or(0))
3245 .unwrap_or(limit),
3246 limit,
3247 input_ids.len()
3248 );
3249 }
3250 }
3251 let graph_prefill = self.graph_prefill_preferred();
3257 #[cfg(target_os = "macos")]
3265 if task_mask.is_none()
3266 && !dyn_prefill
3267 && (crate::gpu::q1_force() || crate::gpu::q2tp_gpu_opt_in())
3268 && crate::gpu::enabled_here()
3269 && self.gdn_cfg.is_some()
3270 && self.g3n.is_none()
3271 && input_ids.len() > 8
3272 && std::env::var("CMF_MTP_CHAIN_PROBE").is_err()
3273 && std::env::var("CMF_METAL_PREFILL").as_deref() != Ok("0")
3274 {
3275 let chunk: usize = std::env::var("CMF_METAL_PREFILL_CHUNK")
3276 .ok()
3277 .and_then(|v| v.parse().ok())
3278 .filter(|&v| (16..=512).contains(&v))
3279 .unwrap_or(256);
3280 let hs = self.hidden_size;
3281 let _tp = std::time::Instant::now();
3282 while pos < input_ids.len() && !self.cancel.load(std::sync::atomic::Ordering::Relaxed) {
3283 let end = (pos + chunk).min(input_ids.len());
3284 let hb = match self.prefill_batch_metal(&input_ids[pos..end], pos) {
3285 MetalPrefillOutcome::Completed(hb) => hb,
3286 MetalPrefillOutcome::Declined => break,
3287 MetalPrefillOutcome::Failed => {
3288 self.finish_generation(&mut mtp, &mut router, true);
3289 return Err("ordinary Metal prefill failed after admission".into());
3290 }
3291 };
3292 if let Some(m) = &mut mtp {
3293 let n_pairs = if end < input_ids.len() {
3294 end - pos
3295 } else {
3296 end - pos - 1
3297 };
3298 if n_pairs > 0 {
3299 let pairs: Vec<(&[f32], u32)> = (0..n_pairs)
3300 .map(|j| (&hb[j * hs..(j + 1) * hs], input_ids[pos + j + 1]))
3301 .collect();
3302 if !self.mtp_warm_batch_metal(m, &pairs, pos) {
3303 for (j, (h, t)) in pairs.iter().enumerate() {
3304 let h = h.to_vec();
3305 let _ = self.mtp_step(m, &h, *t, pos + j);
3306 }
3307 }
3308 }
3309 }
3310 hidden.copy_from_slice(&hb[(end - pos - 1) * hs..]);
3311 pos = end;
3312 }
3313 if std::env::var("CMF_PREFILL_PROF").is_ok() {
3314 eprintln!(
3315 "metal-prefill: {} of {} tokens in {:.1} ms",
3316 pos,
3317 input_ids.len(),
3318 _tp.elapsed().as_secs_f64() * 1e3
3319 );
3320 }
3321 }
3322 if task_mask.is_none()
3323 && !dyn_prefill
3324 && !graph_prefill
3325 && self.can_prefill_batched()
3326 && self.g3n.is_none()
3327 && o1_prefill.is_none()
3328 && input_ids.len() > 2
3329 {
3330 let chunk = prefill_chunk();
3336 let hs = self.hidden_size;
3337 while pos < input_ids.len() && !self.cancel.load(std::sync::atomic::Ordering::Relaxed) {
3338 let end = (pos + chunk).min(input_ids.len());
3339 let hb = self.prefill_batch(&input_ids[pos..end], pos);
3340 if let Some(m) = &mut mtp {
3341 let probe: usize = std::env::var("CMF_MTP_CHAIN_PROBE")
3342 .ok()
3343 .and_then(|v| v.parse().ok())
3344 .unwrap_or(0);
3345 for p in pos..end {
3346 if p + 1 < input_ids.len() {
3347 if probe >= 1 && p + 2 < input_ids.len() {
3348 let (d1, mut hx) = self.mtp_step_h(
3352 m,
3353 &hb[(p - pos) * hs..(p - pos + 1) * hs],
3354 input_ids[p + 1],
3355 p,
3356 );
3357 let mut ok = d1 == input_ids[p + 2];
3358 Self::chain_probe_note(0, ok);
3359 let mut d_prev = d1;
3360 let mut extra = 0usize;
3361 for j in 1..probe {
3362 if p + 2 + j >= input_ids.len() {
3363 break;
3364 }
3365 let (dj, hj) = self.mtp_step_h(m, &hx, d_prev, p + 1 + j);
3366 extra += 1;
3367 ok = ok && dj == input_ids[p + 2 + j];
3368 Self::chain_probe_note(j, ok);
3369 d_prev = dj;
3370 hx = hj;
3371 }
3372 m.kv.truncate_last(extra);
3373 } else {
3374 let _ = self.mtp_step(
3375 m,
3376 &hb[(p - pos) * hs..(p - pos + 1) * hs],
3377 input_ids[p + 1],
3378 p,
3379 );
3380 }
3381 }
3382 }
3383 }
3384 hidden.copy_from_slice(&hb[(end - pos - 1) * hs..]);
3385 pos = end;
3386 }
3387 }
3388 let pair_off = std::env::var("CMF_PAIR").is_ok_and(|v| v == "0");
3389 if task_mask.is_none()
3390 && !dyn_prefill
3391 && !graph_prefill
3392 && !pair_off
3393 && self.pair_supported()
3394 && o1_prefill.is_none()
3395 {
3396 while pos + 1 < input_ids.len()
3397 && !self.cancel.load(std::sync::atomic::Ordering::Relaxed)
3398 {
3399 let e1 = self.embed_single(input_ids[pos]);
3400 let e2 = self.embed_single(input_ids[pos + 1]);
3401 let (h1, h2) = self.forward_pair(&e1, &e2, pos);
3402 self.commit_linear_scratch();
3404 if let Some(m) = &mut mtp {
3405 let _ = self.mtp_step(m, &h1, input_ids[pos + 1], pos);
3406 if pos + 2 < input_ids.len() {
3407 let probe: usize = std::env::var("CMF_MTP_CHAIN_PROBE")
3408 .ok()
3409 .and_then(|v| v.parse().ok())
3410 .unwrap_or(0);
3411 if probe >= 1 && pos + 3 < input_ids.len() {
3412 let (d1, mut hx) = self.mtp_step_h(m, &h2, input_ids[pos + 2], pos + 1);
3416 let mut ok = d1 == input_ids[pos + 3];
3417 Self::chain_probe_note(0, ok);
3418 let mut d_prev = d1;
3419 let mut extra = 0usize;
3420 for j in 1..probe {
3421 if pos + 3 + j >= input_ids.len() {
3422 break;
3423 }
3424 let (dj, hj) = self.mtp_step_h(m, &hx, d_prev, pos + 2 + j);
3425 extra += 1;
3426 ok = ok && dj == input_ids[pos + 3 + j];
3427 Self::chain_probe_note(j, ok);
3428 d_prev = dj;
3429 hx = hj;
3430 }
3431 m.kv.truncate_last(extra);
3432 } else {
3433 let _ = self.mtp_step(m, &h2, input_ids[pos + 2], pos + 1);
3434 }
3435 }
3436 }
3437 hidden = h2;
3438 pos += 2;
3439 }
3440 }
3441 let o1_batch_ready = o1_sealed
3454 && o1_prefill.is_some()
3455 && mtp.is_none()
3456 && std::env::var("CMF_O1_GPU").as_deref() == Ok("1")
3457 && (0..self.num_layers).all(|li| {
3458 let cache = &self.kv_cache.layers[self.phys_layer(li)];
3459 cache.o1.is_none() || cache.o1_views().is_some()
3460 });
3461 let mtp_batch_prefill = mtp.is_some()
3466 && graph_prefill
3467 && task_mask.is_none()
3468 && !dyn_prefill
3469 && !self.o1_active()
3470 && std::env::var("CMF_MTP_CHAIN_PROBE").is_err();
3471 if batch_k > 0
3472 && (graph_prefill || o1_batch_ready)
3473 && task_mask.is_none()
3474 && (!self.o1_active() || o1_batch_ready)
3475 && (mtp.is_none() || mtp_batch_prefill)
3476 && !dyn_prefill
3477 && pos + 1 < input_ids.len()
3478 {
3479 let hs = self.hidden_size;
3480 let chunk = batch_k;
3481 while pos < input_ids.len() {
3482 let end = (pos + chunk).min(input_ids.len());
3483 let bk = end - pos;
3484 let mut hiddens = vec![0f32; bk * hs];
3485 for (j, &id) in input_ids[pos..end].iter().enumerate() {
3486 hiddens[j * hs..(j + 1) * hs].copy_from_slice(&self.embed_single(id));
3487 }
3488 let positions: Vec<usize> = (pos..end).collect();
3489 let t_chunk = std::time::Instant::now();
3490 let outcome = self.try_batch_graph_wgpu(&mut hiddens, &positions, bk, None);
3491 let ok_b = outcome == crate::gpu::BatchGraphOutcome::Completed;
3492 if std::env::var("CMF_GRAPH_PROF").is_ok() {
3493 let ms = t_chunk.elapsed().as_secs_f64() * 1000.0;
3494 eprintln!(
3495 "batch-chunk: phase=prompt mode={} k={bk} outcome={outcome:?} {ms:.1} ms ({:.1} tok/s)",
3496 if o1_batch_ready {
3497 "o1"
3498 } else if mtp_batch_prefill {
3499 "ordinary_mtp"
3500 } else {
3501 "ordinary"
3502 },
3503 bk as f64 / (ms / 1000.0)
3504 );
3505 }
3506 {
3507 use std::sync::atomic::{AtomicBool, Ordering};
3508 static SAID: AtomicBool = AtomicBool::new(false);
3509 if !SAID.swap(true, Ordering::Relaxed) {
3510 if ok_b {
3511 tracing::info!(
3512 "batched prefill: ACTIVE mode={} (k={bk})",
3513 if o1_batch_ready {
3514 "o1"
3515 } else if mtp_batch_prefill {
3516 "ordinary_mtp"
3517 } else {
3518 "ordinary"
3519 }
3520 );
3521 } else {
3522 tracing::warn!("batched prefill {:?} — per-position graph", outcome);
3523 }
3524 }
3525 }
3526 if ok_b {
3527 if mtp_batch_prefill {
3528 let n_pairs = mtp_prefill_pair_count(pos, end, input_ids.len());
3529 if n_pairs > 0 {
3530 let rows: Vec<Vec<f32>> = (0..n_pairs)
3536 .map(|j| hiddens[j * hs..(j + 1) * hs].to_vec())
3537 .collect();
3538 let pairs: Vec<(&[f32], u32)> = rows
3539 .iter()
3540 .enumerate()
3541 .map(|(j, row)| (row.as_slice(), input_ids[pos + j + 1]))
3542 .collect();
3543 if std::env::var("CMF_GRAPH_PROF").is_ok() {
3544 eprintln!(
3545 "mtp-warm: phase=prompt mode=ordinary_mtp first_pos={} pairs={} last_pos={}",
3546 pos,
3547 n_pairs,
3548 pos + n_pairs - 1,
3549 );
3550 }
3551 let warm_error = if let Some(m) = mtp.as_mut() {
3552 self.mtp_warm_prefill_pairs(m, &pairs, pos).err()
3553 } else {
3554 None
3555 };
3556 if let Some(err) = warm_error {
3557 self.finish_generation(&mut mtp, &mut router, true);
3562 return Err(err.to_string());
3563 }
3564 }
3565 }
3566 hidden.copy_from_slice(&hiddens[(bk - 1) * hs..]);
3567 pos = end;
3568 } else if outcome == crate::gpu::BatchGraphOutcome::Failed {
3569 self.finish_generation(&mut mtp, &mut router, true);
3574 return Err(if o1_batch_ready {
3575 "sealed O(1) batch graph failed after admission".to_string()
3576 } else {
3577 "ordinary recurrent batch graph failed after admission".to_string()
3578 });
3579 } else {
3580 break; }
3582 }
3583 }
3584 while pos < input_ids.len() && !self.cancel.load(std::sync::atomic::Ordering::Relaxed) {
3585 self.graph_want_logits = fuse_lm && pos + 1 == input_ids.len();
3586 hidden = self.forward_layers(&self.embed_single(input_ids[pos]), pos, task_mask);
3587 if let Some(m) = &mut mtp {
3588 if pos + 1 < input_ids.len() {
3589 let probe: usize = std::env::var("CMF_MTP_CHAIN_PROBE")
3595 .ok()
3596 .and_then(|v| v.parse().ok())
3597 .unwrap_or(0);
3598 if probe >= 1 && pos + 2 < input_ids.len() {
3599 let (d1, mut hx) = self.mtp_step_h(m, &hidden, input_ids[pos + 1], pos);
3600 let mut ok = d1 == input_ids[pos + 2];
3601 Self::chain_probe_note(0, ok);
3602 let mut d_prev = d1;
3603 let mut extra = 0usize;
3604 for j in 1..probe {
3605 if pos + 2 + j >= input_ids.len() {
3606 break;
3607 }
3608 let (dj, hj) = self.mtp_step_h(m, &hx, d_prev, pos + 1 + j);
3609 extra += 1;
3610 ok = ok && dj == input_ids[pos + 2 + j];
3611 Self::chain_probe_note(j, ok);
3612 d_prev = dj;
3613 hx = hj;
3614 }
3615 m.kv.truncate_last(extra);
3618 } else {
3619 let _ = self.mtp_step(m, &hidden, input_ids[pos + 1], pos);
3620 }
3621 }
3622 }
3623 pos += 1;
3624 }
3625 if std::env::var("CMF_PREFILL_PROF").is_ok() {
3626 eprintln!(
3627 "prefill: {} tokens in {:.1} ms (batch_k={batch_k})",
3628 input_ids.len(),
3629 _tpf.elapsed().as_secs_f64() * 1000.0
3630 );
3631 }
3632 if self
3633 .graph_failed
3634 .swap(false, std::sync::atomic::Ordering::Relaxed)
3635 {
3636 self.finish_generation(&mut mtp, &mut router, true);
3641 return Err("GPU token graph failed during prefill".to_string());
3642 }
3643 if self
3646 .cancel
3647 .swap(false, std::sync::atomic::Ordering::Relaxed)
3648 {
3649 self.finish_generation(&mut mtp, &mut router, true);
3653 return Ok(GenerateResult {
3654 text: String::new(),
3655 token_ids: Vec::new(),
3656 prompt_tokens: input_ids.len(),
3657 tokens_generated: 0,
3658 finish_reason: "cancelled".to_string(),
3659 mtp_drafted: 0,
3660 mtp_accepted: 0,
3661 token_confidence: Vec::new(),
3662 traces: Vec::new(),
3663 });
3664 }
3665
3666 if !o1_sealed {
3669 match self.o1_seal_checked() {
3670 Ok(_) => {}
3671 Err(err) => {
3672 self.finish_generation(&mut mtp, &mut router, true);
3673 return Err(err);
3674 }
3675 }
3676 }
3677
3678 macro_rules! commit {
3680 ($id:expr) => {{
3681 all_ids.push($id);
3682 generated += 1;
3683 self.note_draft_id($id);
3684 if self.tokenizer.is_eos($id) && !self.ignore_eos {
3685 finish_reason = "stop".to_string();
3686 false
3687 } else {
3688 let token_text = self.tokenizer.decode_token($id);
3689 let mut go = true;
3690 if let Some(ref mut cb) = on_token {
3691 if !cb(&token_text) {
3692 finish_reason = "cancelled".to_string();
3693 go = false;
3694 }
3695 }
3696 go
3697 }
3698 }};
3699 }
3700
3701 let mut spec_trial = SpecTrial::Spec {
3712 t0: std::time::Instant::now(),
3713 gen0: generated,
3714 rounds: 0,
3715 };
3716 let mut spec_mon = SpecMon {
3722 metal: graph_spec && crate::gpu::q1_force() && spec_cheap_round,
3723 ..SpecMon::default()
3724 };
3725 let mut spec_watchdog_off = false;
3726 let mut spec_walls: Vec<f32> = Vec::new();
3729 let mut spec_round_end: Option<std::time::Instant> = None;
3732 let mut next_pos = input_ids.len();
3734 'decode: while generated < max_tokens {
3735 if self
3736 .graph_failed
3737 .swap(false, std::sync::atomic::Ordering::Relaxed)
3738 {
3739 self.finish_generation(&mut mtp, &mut router, true);
3744 return Err("GPU token graph failed during decode".to_string());
3745 }
3746 if self
3747 .cancel
3748 .swap(false, std::sync::atomic::Ordering::Relaxed)
3749 {
3750 finish_reason = "cancelled".to_string();
3751 break 'decode;
3752 }
3753 let forced = self.spec_forced.take();
3758 let mut logits = match (forced, self.graph_logits.take()) {
3759 (Some(_), _) => Vec::new(),
3760 (None, Some(lg)) => lg,
3761 (None, None) => {
3762 inference::rms_norm_into(
3763 &hidden,
3764 &self.weights.final_norm,
3765 self.rms_eps,
3766 self.norm_style,
3767 &mut self.ws.n1,
3768 );
3769 self.lm_head_forward(&self.ws.n1)
3770 }
3771 };
3772 if generated
3775 == std::env::var("CMF_LOGIT_DUMP_STEP")
3776 .ok()
3777 .and_then(|v| v.parse().ok())
3778 .unwrap_or(0)
3779 {
3780 if let Ok(path) = std::env::var("CMF_LOGIT_DUMP") {
3781 let mut bytes: Vec<u8> = Vec::with_capacity((hidden.len() + logits.len()) * 4);
3782 for v in hidden.iter().chain(logits.iter()) {
3783 bytes.extend_from_slice(&v.to_le_bytes());
3784 }
3785 if let Err(e) = std::fs::write(&path, &bytes) {
3786 eprintln!("logit dump: failed to write {path}: {e}");
3787 self.finish_generation(&mut mtp, &mut router, true);
3788 return Err(format!("logit dump write failed: {e}"));
3789 }
3790 }
3791 }
3792 let t_next = match forced {
3793 Some(c) => c,
3794 None => sampler::sample_with_scratch_pool(
3795 &logits,
3796 &self.sampler_config,
3797 &all_ids,
3798 &mut self.rng,
3799 &mut self.sampler_scratch,
3800 self.pool.as_deref(),
3801 ),
3802 };
3803 if self.confidence_on {
3804 confidence.push(if logits.is_empty() {
3805 0.0
3806 } else {
3807 sampler::top1_prob_pool(
3808 self.pool.as_deref(),
3809 &mut self.sampler_scratch,
3810 &logits,
3811 t_next,
3812 calib_temp,
3813 )
3814 });
3815 }
3816 if !logits.is_empty() {
3817 attention::recycle_buf(&mut logits);
3818 }
3819 if trace_on {
3820 let skill = router.as_ref().and_then(|r| r.active_id());
3824 traces.push(TokenTrace {
3825 t: generated,
3826 token_id: t_next,
3827 confidence: confidence.last().copied().unwrap_or(0.0),
3828 active_skill: skill,
3829 recon: None,
3830 switched: false,
3831 });
3832 }
3833 if !commit!(t_next) {
3834 break 'decode;
3835 }
3836 if generated >= max_tokens {
3837 break 'decode;
3838 }
3839
3840 if self.dsv41.is_none() && self.kv_cache.needs_eviction() {
3841 static SAID: std::sync::Once = std::sync::Once::new();
3847 SAID.call_once(|| {
3848 tracing::warn!(
3849 "KV cache full at {} positions — evicting half; quality \
3850 will degrade. Raise CMF_MAX_SEQ.",
3851 self.kv_cache.max_seq_len,
3852 );
3853 });
3854 let keep = (self.kv_cache.max_seq_len / 2).max(1);
3855 self.kv_cache.evict(keep);
3856 }
3857
3858 if graph_spec {
3861 match spec_trial {
3862 SpecTrial::Plain { t0, gen0 } if spec_mon.plain_done(t0, gen0, generated) => {
3863 spec_mon.plain_ms =
3864 t0.elapsed().as_secs_f64() * 1e3 / (generated - gen0) as f64;
3865 let keep = spec_mon.pays();
3866 tracing::info!(
3867 "speculation trial: {:.2} tok/round in {:.1} ms vs plain {:.1} ms/tok — {}",
3868 spec_mon.tokens,
3869 spec_mon.round_ms,
3870 spec_mon.plain_ms,
3871 if keep { "speculating" } else { "plain" }
3872 );
3873 spec_mon.fails = 0;
3874 spec_trial = SpecTrial::Decided {
3875 spec: keep,
3876 recheck_at: if keep { usize::MAX } else { generated + 128 },
3877 };
3878 }
3879 SpecTrial::Decided { recheck_at, .. } if generated >= recheck_at => {
3880 spec_mon.n = 0;
3881 spec_trial = SpecTrial::Spec {
3882 t0: std::time::Instant::now(),
3883 gen0: generated,
3884 rounds: 0,
3885 };
3886 }
3887 _ => {}
3888 }
3889 spec_watchdog_off = matches!(
3890 spec_trial,
3891 SpecTrial::Plain { .. } | SpecTrial::Decided { spec: false, .. }
3892 );
3893 }
3894 match &mut mtp {
3895 #[cfg(feature = "gpu")]
3897 Some(m)
3898 if graph_spec
3899 && !spec_watchdog_off
3900 && generated + 1 < max_tokens
3901 && next_pos > 0 =>
3902 {
3903 let t_round = std::time::Instant::now();
3904 if spec_time_level() >= 2 {
3905 if let Some(t) = spec_round_end.take() {
3906 eprintln!(
3907 "spec-gap {:.2} ms (host between rounds)",
3908 t.elapsed().as_secs_f64() * 1e3
3909 );
3910 }
3911 }
3912 spec_stamps_begin();
3913 #[cfg(target_os = "macos")]
3918 let allocs0 = crate::gpu_metal::IO_BUF_ALLOCS
3919 .load(std::sync::atomic::Ordering::Relaxed);
3920 #[cfg(not(target_os = "macos"))]
3921 let allocs0 = 0u64;
3922 if let Some((extra, n_pos, new_h)) = self.graph_spec_step(
3923 m,
3924 &hidden,
3925 t_next,
3926 next_pos,
3927 &mut drafted,
3928 &mut accepted,
3929 &mut all_ids,
3930 max_tokens - generated,
3931 ) {
3932 next_pos = n_pos;
3933 hidden = new_h;
3934 let level = spec_time_level();
3935 if level > 0 {
3936 let wall = t_round.elapsed().as_secs_f32() * 1e3;
3937 let stamps = spec_stamps_take();
3938 let median = if spec_walls.len() >= 3 {
3941 let mut s = spec_walls.clone();
3942 s.sort_by(|a, b| a.partial_cmp(b).unwrap());
3943 Some(s[s.len() / 2])
3944 } else {
3945 None
3946 };
3947 let outlier = median.is_some_and(|m| wall > 1.4 * m);
3948 #[cfg(target_os = "macos")]
3949 let allocs = crate::gpu_metal::IO_BUF_ALLOCS
3950 .load(std::sync::atomic::Ordering::Relaxed)
3951 - allocs0;
3952 #[cfg(not(target_os = "macos"))]
3953 let allocs = allocs0;
3954 eprintln!(
3955 "spec-round wall {wall:.1} ms → {} tokens{}{}",
3956 extra.len() + 1,
3957 if allocs > 0 {
3958 format!(" [{allocs} new device buffers]")
3959 } else {
3960 String::new()
3961 },
3962 match (outlier, median) {
3963 (true, Some(m)) => format!(" OUTLIER (median {m:.1})"),
3964 _ => String::new(),
3965 }
3966 );
3967 if level >= 2 || outlier {
3968 let sum: f32 = stamps.iter().map(|s| s.1).sum();
3969 eprintln!(
3970 "spec-stamps: {}| untracked {:.1}",
3971 spec_stamps_format(&stamps),
3972 wall - sum
3973 );
3974 }
3975 if spec_mon.n >= 1 {
3976 spec_walls.push(wall);
3977 }
3978 }
3979 spec_mon.round(t_round.elapsed().as_secs_f64() * 1e3, extra.len() + 1);
3983 spec_trial = Self::spec_trial_round(
3986 spec_trial,
3987 &mut spec_mon,
3988 generated + extra.len() + 1,
3989 );
3990 let mut stopped = false;
3991 for &id in &extra {
3992 if self.confidence_on {
3993 confidence.push(0.0);
3994 }
3995 if !commit!(id) {
3996 stopped = true;
3997 break;
3998 }
3999 }
4000 if stopped {
4001 break 'decode;
4002 }
4003 if spec_time_level() >= 2 {
4004 spec_round_end = Some(std::time::Instant::now());
4005 }
4006 continue 'decode;
4007 }
4008 if self
4009 .graph_failed
4010 .swap(false, std::sync::atomic::Ordering::Relaxed)
4011 {
4012 self.finish_generation(&mut mtp, &mut router, true);
4018 return Err("GPU MTP graph failed during speculative decode".to_string());
4019 }
4020 spec_mon.round(t_round.elapsed().as_secs_f64() * 1e3, 1);
4031 spec_mon.tokens = 0.0;
4032 spec_mon.fails = 3;
4033 spec_trial = Self::spec_trial_round(spec_trial, &mut spec_mon, generated + 1);
4034 hidden = self.forward_layers(&self.embed_single(t_next), next_pos, task_mask);
4035 next_pos += 1;
4036 continue 'decode;
4037 }
4038 Some(m) if !graph_spec && generated + 1 < max_tokens => {
4040 let draft = self.mtp_step(m, &hidden, t_next, next_pos - 1);
4041 drafted += 1;
4042 let emb1 = self.embed_single(t_next);
4043 let emb2 = self.embed_single(draft);
4044 let (h1, h2) = self.forward_pair(&emb1, &emb2, next_pos);
4045
4046 inference::rms_norm_into(
4047 &h1,
4048 &self.weights.final_norm,
4049 self.rms_eps,
4050 self.norm_style,
4051 &mut self.ws.n1,
4052 );
4053 let mut logits1 = self.lm_head_forward(&self.ws.n1);
4054 let t_after = sampler::sample_with_scratch_pool(
4055 &logits1,
4056 &self.sampler_config,
4057 &all_ids,
4058 &mut self.rng,
4059 &mut self.sampler_scratch,
4060 self.pool.as_deref(),
4061 );
4062 if self.confidence_on {
4063 confidence.push(sampler::top1_prob_pool(
4064 self.pool.as_deref(),
4065 &mut self.sampler_scratch,
4066 &logits1,
4067 t_after,
4068 calib_temp,
4069 ));
4070 }
4071 attention::recycle_buf(&mut logits1);
4072 if trace_on {
4073 traces.push(TokenTrace {
4076 t: generated,
4077 token_id: t_after,
4078 confidence: confidence.last().copied().unwrap_or(0.0),
4079 active_skill: None,
4080 recon: None,
4081 switched: false,
4082 });
4083 }
4084 let stop = !commit!(t_after);
4085
4086 if t_after == draft {
4087 accepted += 1;
4088 self.commit_linear_scratch();
4089 let _ = self.mtp_step(m, &h1, t_after, next_pos);
4090 hidden = h2;
4091 next_pos += 2;
4092 } else {
4093 for layer in &mut self.kv_cache.layers {
4095 layer.truncate_last(1);
4096 }
4097 if !stop {
4098 let _ = self.mtp_step(m, &h1, t_after, next_pos);
4099 hidden = self.forward_layers(
4100 &self.embed_single(t_after),
4101 next_pos + 1,
4102 None,
4103 );
4104 }
4105 next_pos += 2;
4106 }
4107 if stop {
4108 break 'decode;
4109 }
4110 }
4111 _ => {
4113 #[cfg(feature = "gpu")]
4118 if Self::dsv4_spec_on() && self.dsv4.is_some() {
4119 static SAID: std::sync::Once = std::sync::Once::new();
4120 SAID.call_once(|| {
4121 eprintln!(
4122 "dsv4-spec гейт: mtp={} mask={} router={} trace={} temp={} rep={} ",
4123 !self.dsv4_mtp.is_empty(),
4124 task_mask.is_none(),
4125 router.is_none(),
4126 !trace_on,
4127 self.sampler_config.temperature < 1e-6,
4128 self.sampler_config.repetition_penalty == 1.0,
4129 );
4130 });
4131 }
4132 #[cfg(feature = "gpu")]
4133 if Self::dsv4_spec_on()
4134 && self.dsv4.is_some()
4135 && !self.dsv4_mtp.is_empty()
4136 && task_mask.is_none()
4137 && router.is_none()
4138 && !trace_on
4139 && self.sampler_config.temperature < 1e-6
4140 && self.sampler_config.repetition_penalty == 1.0
4141 && generated + 1 < max_tokens
4142 && all_ids.len() >= 2
4143 && generated >= dsv4_spec_retry_at
4144 {
4145 let tip_token = all_ids[all_ids.len() - 2];
4146 let drafted0 = drafted;
4147 let round = self.dsv4_spec_step(
4148 tip_token,
4149 t_next,
4150 next_pos,
4151 max_tokens.saturating_sub(generated),
4152 &mut drafted,
4153 &mut accepted,
4154 );
4155 if drafted > drafted0 {
4156 let useful = round.as_ref().is_some_and(|(extra, _)| !extra.is_empty());
4157 if useful {
4158 dsv4_spec_bad = 0;
4159 } else {
4160 dsv4_spec_bad += 1;
4161 if dsv4_spec_bad >= 2 {
4162 dsv4_spec_bad = 0;
4163 dsv4_spec_retry_at = generated.saturating_add(32);
4164 tracing::info!(
4165 "dsv4: draft не окупился дважды — точный walk на 32 токена"
4166 );
4167 }
4168 }
4169 }
4170 if let Some((extra, n_pos)) = round {
4171 next_pos = n_pos;
4172 let mut stopped = false;
4173 for &id in &extra {
4174 if self.confidence_on {
4175 confidence.push(0.0);
4176 }
4177 if !commit!(id) {
4178 stopped = true;
4179 break;
4180 }
4181 }
4182 if stopped {
4183 break 'decode;
4184 }
4185 continue 'decode;
4186 }
4187 }
4188 self.graph_want_logits = fuse_lm;
4189 let mut t_fwd = t_next;
4195 let pure_greedy = self.sampler_config.temperature < 1e-6
4196 && self.sampler_config.repetition_penalty == 1.0
4197 && self.sampler_config.suppress_tokens.is_empty();
4198 let burst_k = std::env::var("CMF_MULTISTEP")
4203 .ok()
4204 .and_then(|v| v.parse::<usize>().ok())
4205 .unwrap_or(0);
4206 if pure_greedy
4207 && burst_k >= 1
4208 && fuse_lm
4209 && task_mask.is_none()
4210 && router.is_none()
4211 && !trace_on
4212 && !self.confidence_on
4213 {
4214 let mut stopped = false;
4215 loop {
4216 let room = max_tokens.saturating_sub(generated);
4217 if room <= 2 {
4218 break;
4219 }
4220 let k = burst_k.min(room - 1);
4221 if k < 1 {
4222 break;
4223 }
4224 let Some(ids) = self.try_multi_burst(t_fwd, next_pos, k) else {
4225 if self
4226 .graph_failed
4227 .swap(false, std::sync::atomic::Ordering::Relaxed)
4228 {
4229 self.finish_generation(&mut mtp, &mut router, true);
4230 return Err(
4231 "GPU token graph failed during greedy burst".to_string()
4232 );
4233 }
4234 break;
4235 };
4236 next_pos += k;
4237 for &id in &ids {
4238 if !commit!(id) {
4239 stopped = true;
4240 break;
4241 }
4242 }
4243 if stopped {
4244 break;
4245 }
4246 t_fwd = *ids.last().unwrap();
4247 }
4248 if stopped {
4249 break 'decode;
4250 }
4251 }
4252 #[cfg(target_os = "macos")]
4262 if graph_spec
4263 && spec_watchdog_off
4264 && next_pos > 0
4265 && self.mtp_graph_mode == Some(true)
4266 && crate::gpu::q1_force()
4267 {
4268 if let Some(m) = mtp.as_mut() {
4269 let _ = self.mtp_step_metal(m, &hidden, t_fwd, next_pos - 1, false);
4270 }
4271 }
4272 hidden = self.forward_layers(&self.embed_single(t_fwd), next_pos, task_mask);
4273 next_pos += 1;
4274 if let Some(r) = &mut router {
4277 let phi = self.dyn_phi_ema.clone();
4278 let decision = r.step(&phi, generated);
4279 if let Some(new_active) = decision {
4280 let _ = self.set_active_skill(new_active);
4281 }
4282 if trace_on {
4285 if let Some(last) = traces.last_mut() {
4286 let e = r.last_best_e();
4287 last.recon = e.is_finite().then_some(e);
4288 last.switched = decision.is_some();
4289 }
4290 }
4291 }
4292 }
4293 }
4294 }
4295
4296 let cancelled = finish_reason == "cancelled";
4297 self.finish_generation(&mut mtp, &mut router, cancelled);
4298
4299 let output_ids = &all_ids[input_ids.len()..];
4300 let forwarded = input_ids.len() + output_ids.len().saturating_sub(1);
4304 if cancelled {
4305 self.kv_history.clear();
4306 } else {
4307 self.kv_history = all_ids[..forwarded.min(all_ids.len())].to_vec();
4308 }
4309 confidence.truncate(output_ids.len()); traces.truncate(output_ids.len());
4311 Ok(GenerateResult {
4312 text: self.tokenizer.decode(output_ids),
4313 token_ids: output_ids.to_vec(),
4314 prompt_tokens: input_ids.len(),
4315 tokens_generated: generated,
4316 finish_reason,
4317 mtp_drafted: drafted,
4318 mtp_accepted: accepted,
4319 token_confidence: confidence,
4320 traces,
4321 })
4322 }
4323
4324 fn mtp_step(
4328 &mut self,
4329 m: &mut MtpModule,
4330 hidden: &[f32],
4331 next_token: u32,
4332 position: usize,
4333 ) -> u32 {
4334 self.mtp_step_h(m, hidden, next_token, position).0
4335 }
4336
4337 fn chain_probe_note(depth: usize, prefix_ok: bool) {
4341 use std::sync::Mutex;
4342 static T: Mutex<Vec<(u64, u64)>> = Mutex::new(Vec::new());
4343 let mut t = T.lock().unwrap();
4344 if t.len() <= depth {
4345 t.resize(depth + 1, (0, 0));
4346 }
4347 t[depth].0 += 1;
4348 t[depth].1 += prefix_ok as u64;
4349 if depth == 0 && t[0].0 % 128 == 0 {
4350 let line: Vec<String> = t
4351 .iter()
4352 .enumerate()
4353 .map(|(d, (n, k))| {
4354 format!(
4355 "d{}={:.0}%({n})",
4356 d + 1,
4357 100.0 * *k as f64 / (*n).max(1) as f64
4358 )
4359 })
4360 .collect();
4361 eprintln!("mtp-chain: {}", line.join(" "));
4362 }
4363 }
4364
4365 fn mtp_step_hl(
4373 &mut self,
4374 m: &mut MtpModule,
4375 hidden: &[f32],
4376 next_token: u32,
4377 position: usize,
4378 ) -> (Vec<f32>, Vec<f32>) {
4379 #[cfg(target_os = "macos")]
4384 if self.mtp_graph_mode != Some(false) && crate::gpu::q1_force() {
4385 if let Some(r) = self.mtp_step_metal(m, hidden, next_token, position, true) {
4386 self.mtp_graph_mode = Some(true);
4387 return r;
4388 }
4389 if self.mtp_graph_mode == Some(true) {
4390 tracing::error!("mtp Metal graph failed after admission");
4391 self.clear_sequence_state();
4392 self.graph_failed
4393 .store(true, std::sync::atomic::Ordering::Relaxed);
4394 self.cancel
4395 .store(true, std::sync::atomic::Ordering::Relaxed);
4396 return (Vec::new(), Vec::new());
4397 }
4398 self.mtp_graph_mode = Some(false);
4399 }
4400 #[cfg(feature = "gpu")]
4401 if self.mtp_graph_mode != Some(false) {
4402 if !self.mtp_graph_ok(m) {
4403 if self.mtp_graph_mode == Some(true) {
4404 tracing::error!("mtp graph became unavailable after admission");
4409 self.clear_sequence_state();
4410 self.graph_failed
4411 .store(true, std::sync::atomic::Ordering::Relaxed);
4412 self.cancel
4413 .store(true, std::sync::atomic::Ordering::Relaxed);
4414 return (Vec::new(), Vec::new());
4415 }
4416 self.mtp_graph_mode = Some(false);
4417 } else {
4418 if let Some(r) = self.mtp_step_graph(m, hidden, next_token, position) {
4419 self.mtp_graph_mode = Some(true);
4420 return r;
4421 }
4422 if self.graph_failed.load(std::sync::atomic::Ordering::Relaxed) {
4423 return (Vec::new(), Vec::new());
4430 }
4431 tracing::error!("mtp graph failed or declined after admission");
4435 self.clear_sequence_state();
4436 self.graph_failed
4437 .store(true, std::sync::atomic::Ordering::Relaxed);
4438 self.cancel
4439 .store(true, std::sync::atomic::Ordering::Relaxed);
4440 return (Vec::new(), Vec::new());
4441 }
4442 }
4443 let e = self.embed_single(next_token);
4447 let mut cat = vec![0.0f32; 2 * self.hidden_size];
4448 let (cat_e, cat_h) = cat.split_at_mut(self.hidden_size);
4449 inference::rms_norm_into(&e, &m.enorm, self.rms_eps, self.norm_style, cat_e);
4450 inference::rms_norm_into(hidden, &m.hnorm, self.rms_eps, self.norm_style, cat_h);
4451 let mut x = vec![0.0f32; self.hidden_size];
4452 m.eh_proj.matvec(&cat, &mut x, self.pool.as_deref());
4453
4454 let lw = &m.layer;
4456 inference::rms_norm_into(
4457 &x,
4458 &lw.input_norm,
4459 self.rms_eps,
4460 self.norm_style,
4461 &mut self.ws.n1,
4462 );
4463 let attn = match &lw.attn {
4464 AttnKind::Mla(_) => unreachable!("MLA has no MTP/pair path"),
4466 AttnKind::Kda(_) => unreachable!("KDA has no MTP/pair path"),
4467 AttnKind::Full {
4468 wq,
4469 wk,
4470 wv,
4471 wo,
4472 q_norm,
4473 k_norm,
4474 output_gate,
4475 softplus_gate,
4476 bias,
4477 } => {
4478 let mut cfg = self.attn_cfg(position);
4479 cfg.q_norm = q_norm.as_deref();
4480 cfg.k_norm = k_norm.as_deref();
4481 cfg.output_gate = *output_gate;
4482 cfg.softplus_gate = softplus_gate
4483 .as_ref()
4484 .map(|(gate, per_head)| (gate, *per_head));
4485 cfg.bias = bias
4486 .as_ref()
4487 .map(|(q, k, v)| (q.as_slice(), k.as_slice(), v.as_slice()));
4488 attention::qwen_attention(&self.ws.n1, wq, wk, wv, wo, &mut m.kv, &cfg)
4489 }
4490 AttnKind::Linear(_) | AttnKind::LinearGdn(_) | AttnKind::ShortConv(_) => {
4491 unreachable!("MTP block is full attention")
4492 }
4493 };
4494 for (i, &a) in attn.iter().enumerate() {
4495 x[i] += a;
4496 }
4497 inference::rms_norm_into(
4498 &x,
4499 &lw.post_norm,
4500 self.rms_eps,
4501 self.norm_style,
4502 &mut self.ws.p1,
4503 );
4504 let ffn = ffn_forward(&lw.ffn, &self.ws.p1, self.pool.as_deref(), None);
4505 for (i, &f) in ffn.iter().enumerate() {
4506 x[i] += f;
4507 }
4508
4509 inference::rms_norm_into(
4510 &x,
4511 &m.final_norm,
4512 self.rms_eps,
4513 self.norm_style,
4514 &mut self.ws.n1,
4515 );
4516 let lg = self.lm_head_forward(&self.ws.n1);
4517 (lg, x)
4518 }
4519
4520 fn mtp_step_h(
4522 &mut self,
4523 m: &mut MtpModule,
4524 hidden: &[f32],
4525 next_token: u32,
4526 position: usize,
4527 ) -> (u32, Vec<f32>) {
4528 let (mut lg, x) = self.mtp_step_hl(m, hidden, next_token, position);
4529 let draft = sampler::argmax(&lg);
4530 attention::recycle_buf(&mut lg);
4531 (draft, x)
4532 }
4533
4534 fn spec_trial_round(trial: SpecTrial, mon: &mut SpecMon, generated: usize) -> SpecTrial {
4540 match trial {
4541 SpecTrial::Spec { t0, gen0, rounds } => {
4542 let rounds = rounds + 1;
4543 if rounds >= 5 {
4544 if mon.plain_ms > 0.0 {
4545 let keep = mon.pays();
4546 mon.fails = 0;
4547 tracing::info!(
4548 "speculation re-check: {:.2} tok/round in {:.1} ms vs plain {:.1} ms/tok — {}",
4549 mon.tokens,
4550 mon.round_ms,
4551 mon.plain_ms,
4552 if keep { "speculating" } else { "plain" }
4553 );
4554 SpecTrial::Decided {
4555 spec: keep,
4556 recheck_at: if keep { usize::MAX } else { generated + 128 },
4557 }
4558 } else if mon.pays() {
4559 mon.fails = 0;
4564 tracing::info!(
4565 "speculation trial: {:.2} tok/round in {:.1} ms — speculating (plain not timed)",
4566 mon.tokens,
4567 mon.round_ms,
4568 );
4569 SpecTrial::Decided {
4570 spec: true,
4571 recheck_at: usize::MAX,
4572 }
4573 } else {
4574 SpecTrial::Plain {
4575 t0: std::time::Instant::now(),
4576 gen0: generated,
4577 }
4578 }
4579 } else {
4580 SpecTrial::Spec { t0, gen0, rounds }
4581 }
4582 }
4583 SpecTrial::Decided { spec: true, .. } => {
4584 if mon.pays() {
4585 mon.fails = 0;
4586 trial
4587 } else {
4588 mon.fails += 1;
4589 if mon.fails >= 4 {
4590 if mon.plain_ms <= 0.0 {
4591 tracing::info!(
4595 "speculation doubtful: {:.2} tok/round in {:.1} ms — timing plain",
4596 mon.tokens,
4597 mon.round_ms,
4598 );
4599 return SpecTrial::Plain {
4600 t0: std::time::Instant::now(),
4601 gen0: generated,
4602 };
4603 }
4604 tracing::info!(
4605 "speculation stopped: {:.2} tok/round in {:.1} ms vs plain {:.1} ms/tok",
4606 mon.tokens,
4607 mon.round_ms,
4608 mon.plain_ms
4609 );
4610 SpecTrial::Decided {
4611 spec: false,
4612 recheck_at: generated + 128,
4613 }
4614 } else {
4615 trial
4616 }
4617 }
4618 }
4619 other => other,
4620 }
4621 }
4622
4623 fn mtp_kv_id(&self) -> u64 {
4626 self.graph_kv_id | (1u64 << 40)
4627 }
4628
4629 const MTP_LAYER_BASE: usize = 0;
4634
4635 #[cfg(feature = "gpu")]
4642 fn rewind_mtp_graph_mirror(&self, stored: usize) -> bool {
4643 self.mtp_graph_mode != Some(true)
4644 || crate::gpu::graph_kv_set_stored(self.mtp_kv_id(), Self::MTP_LAYER_BASE, stored)
4645 }
4646
4647 #[cfg(feature = "gpu")]
4653 fn rewind_trunk_graph_mirrors(&self, stored: usize) -> bool {
4654 let mut ok = true;
4655 let mut expected = false;
4656 for li in 0..self.num_layers {
4657 if matches!(
4658 self.weights.layers[self.phys_layer(li)].attn,
4659 AttnKind::Full { .. }
4660 ) {
4661 expected = true;
4662 ok &= crate::gpu::graph_kv_set_stored(self.graph_kv_id, li, stored);
4663 }
4664 }
4665 !expected || ok
4666 }
4667
4668 fn graph_gdn_layer_count(&self) -> usize {
4672 (0..self.num_layers)
4673 .filter(|&li| {
4674 matches!(
4675 &self.weights.layers[self.phys_layer(li)].attn,
4676 AttnKind::LinearGdn(_)
4677 )
4678 })
4679 .count()
4680 }
4681
4682 fn mtp_block_input(&mut self, m: &MtpModule, hidden: &[f32], next_token: u32) -> Vec<f32> {
4685 let e = self.embed_single(next_token);
4686 let mut cat = vec![0.0f32; 2 * self.hidden_size];
4687 let (cat_e, cat_h) = cat.split_at_mut(self.hidden_size);
4688 inference::rms_norm_into(&e, &m.enorm, self.rms_eps, self.norm_style, cat_e);
4689 inference::rms_norm_into(hidden, &m.hnorm, self.rms_eps, self.norm_style, cat_h);
4690 let mut x = vec![0.0f32; self.hidden_size];
4691 m.eh_proj.matvec(&cat, &mut x, self.pool.as_deref());
4692 x
4693 }
4694
4695 #[cfg(feature = "gpu")]
4698 fn mtp_block_graph_ok(&self, m: &MtpModule) -> bool {
4699 if std::env::var("CMF_MTP_GRAPH").as_deref() == Ok("0") {
4700 return false;
4701 }
4702 if !crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Decode)
4703 || !crate::gpu::enabled_here()
4704 || self.attn_softcap > 0.0
4705 || self.attention_heads_per_layer.is_some()
4706 {
4707 return false;
4708 }
4709 matches!(
4710 &m.layer.attn,
4711 AttnKind::Full {
4712 softplus_gate: None,
4713 ..
4714 }
4715 ) && matches!(&m.layer.ffn, FfnKind::Dense(_))
4716 }
4717
4718 #[cfg(feature = "gpu")]
4722 fn mtp_graph_ok(&self, m: &MtpModule) -> bool {
4723 if !self.mtp_block_graph_ok(m) {
4724 return false;
4725 }
4726 let AttnKind::Full { wq, wk, wv, wo, .. } = &m.layer.attn else {
4727 return false;
4728 };
4729 let FfnKind::Dense(d) = &m.layer.ffn else {
4730 return false;
4731 };
4732 d.segs.is_empty()
4733 && wq.graph_weight().is_some()
4734 && wk.graph_weight().is_some()
4735 && wv.graph_weight().is_some()
4736 && wo.graph_weight().is_some()
4737 && d.gate_proj.graph_weight().is_some()
4738 && d.up_proj.graph_weight().is_some()
4739 && d.down_proj.graph_weight().is_some()
4740 && self.weights.lm_head.graph_weight().is_some()
4741 }
4742
4743 #[cfg(feature = "gpu")]
4749 fn mtp_step_graph(
4750 &mut self,
4751 m: &mut MtpModule,
4752 hidden: &[f32],
4753 next_token: u32,
4754 position: usize,
4755 ) -> Option<(Vec<f32>, Vec<f32>)> {
4756 if !self.mtp_graph_ok(m) {
4757 return None;
4758 }
4759 let lw = &m.layer;
4760 let AttnKind::Full {
4761 wq,
4762 wk,
4763 wv,
4764 wo,
4765 q_norm,
4766 k_norm,
4767 output_gate,
4768 softplus_gate,
4769 bias,
4770 } = &lw.attn
4771 else {
4772 return None;
4773 };
4774 if softplus_gate.is_some() {
4775 return None;
4776 }
4777 let FfnKind::Dense(d) = &lw.ffn else {
4778 return None;
4779 };
4780 if !d.segs.is_empty() {
4781 return None; }
4783 let mut x = self.mtp_block_input(m, hidden, next_token);
4786 fn gw(t: &QTensor) -> Option<crate::gpu::GraphW<'_>> {
4787 let (_, i, kind, rs) = t.graph_weight()?;
4788 Some(crate::gpu::GraphW {
4789 idx: i,
4790 kind,
4791 row_scale: rs,
4792 data: &[],
4793 prism: crate::gpu::GraphPrismOp::None,
4794 affine: false,
4795 })
4796 }
4797 let (model, _, _, _) = wq.graph_weight()?;
4798 let model = model.clone();
4799 let (lm_gw, lm_rows) = {
4800 let (_, i, kind, rs) = self.weights.lm_head.graph_weight()?;
4801 let rows = if kind == 6 {
4805 self.draft_head_rows(self.weights.lm_head.rows())
4806 } else {
4807 self.weights.lm_head.rows()
4808 };
4809 (
4810 crate::gpu::GraphW {
4811 idx: i,
4812 kind,
4813 row_scale: rs,
4814 data: &[],
4815 prism: crate::gpu::GraphPrismOp::None,
4816 affine: false,
4817 },
4818 rows,
4819 )
4820 };
4821 let layer = crate::gpu::GraphLayer {
4822 input_norm: &lw.input_norm,
4823 attn: crate::gpu::GraphAttn::Full {
4824 wq: gw(wq)?,
4825 wk: gw(wk)?,
4826 wv: gw(wv)?,
4827 wo: gw(wo)?,
4828 q_norm: q_norm.as_deref(),
4829 k_norm: k_norm.as_deref(),
4830 late_qk_norm: self.qk_norm_after_rope,
4831 bias: bias
4832 .as_ref()
4833 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
4834 output_gate: *output_gate,
4835 cpu_k: m.kv.k_heads(),
4836 cpu_v: m.kv.v_heads(),
4837 },
4838 post_norm: &lw.post_norm,
4839 ffn: crate::gpu::GraphFfn::Dense {
4840 gate: gw(&d.gate_proj)?,
4841 up: gw(&d.up_proj)?,
4842 down: gw(&d.down_proj)?,
4843 },
4844 };
4845 let nh = self.num_heads;
4846 let (nkv, hd, rd) = self.layer_geom(0);
4847 let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
4848 let mut logits = Vec::new();
4849 let ok = crate::gpu::forward_token_graph(
4850 &model,
4851 self.mtp_kv_id(),
4852 std::slice::from_ref(&layer),
4853 &[None],
4854 self.o1_epoch,
4855 &self.inv_freq,
4856 &mut x,
4857 nh,
4858 nkv,
4859 hd,
4860 self.attn_scale,
4861 rd,
4862 self.hidden_size,
4863 self.intermediate_size,
4864 position,
4865 self.kv_cache.max_seq_len,
4866 gemma,
4867 self.rms_eps as f32,
4868 Some((&lm_gw, lm_rows)),
4869 &m.final_norm,
4870 &mut logits,
4871 &[],
4872 1,
4873 None,
4874 None,
4875 None,
4876 Self::MTP_LAYER_BASE,
4877 true,
4878 );
4879 match ok {
4880 crate::gpu::TokenGraphOutcome::Completed => {}
4881 crate::gpu::TokenGraphOutcome::Declined => return None,
4882 crate::gpu::TokenGraphOutcome::Failed => {
4883 self.clear_sequence_state();
4887 self.graph_failed
4888 .store(true, std::sync::atomic::Ordering::Relaxed);
4889 self.cancel
4890 .store(true, std::sync::atomic::Ordering::Relaxed);
4891 return None;
4892 }
4893 }
4894 logits.resize(self.vocab_size, 0.0);
4895 Some((logits, x))
4896 }
4897
4898 #[cfg(feature = "gpu")]
4906 fn mtp_warm_graph(
4907 &mut self,
4908 m: &mut MtpModule,
4909 pairs: &[(&[f32], u32)],
4910 first_pos: usize,
4911 ) -> crate::gpu::BatchGraphOutcome {
4912 if pairs.is_empty() {
4913 return crate::gpu::BatchGraphOutcome::Completed;
4914 }
4915 if !self.mtp_block_graph_ok(m) {
4916 return crate::gpu::BatchGraphOutcome::Declined;
4917 }
4918 let hs = self.hidden_size;
4919 let mut hiddens = Vec::with_capacity(pairs.len() * hs);
4922 for (h, t) in pairs {
4923 hiddens.extend_from_slice(&self.mtp_block_input(m, h, *t));
4924 }
4925 let lw = &m.layer;
4926 let AttnKind::Full {
4927 wq,
4928 wk,
4929 wv,
4930 wo,
4931 q_norm,
4932 k_norm,
4933 output_gate,
4934 bias,
4935 ..
4936 } = &lw.attn
4937 else {
4938 return crate::gpu::BatchGraphOutcome::Declined;
4939 };
4940 let FfnKind::Dense(d) = &lw.ffn else {
4941 return crate::gpu::BatchGraphOutcome::Declined;
4942 };
4943 if !d.segs.is_empty() {
4944 return crate::gpu::BatchGraphOutcome::Declined; }
4946 fn gw(t: &QTensor) -> Option<crate::gpu::GraphW<'_>> {
4947 let (_, i, kind, rs) = t.graph_weight()?;
4948 Some(crate::gpu::GraphW {
4949 idx: i,
4950 kind,
4951 row_scale: rs,
4952 data: &[],
4953 prism: crate::gpu::GraphPrismOp::None,
4954 affine: false,
4955 })
4956 }
4957 let Some((model, _, _, _)) = wq.graph_weight() else {
4958 return crate::gpu::BatchGraphOutcome::Declined;
4959 };
4960 let model = model.clone();
4961 let (Some(gwq), Some(gwk), Some(gwv), Some(gwo), Some(gg), Some(gu), Some(gd)) = (
4962 gw(wq),
4963 gw(wk),
4964 gw(wv),
4965 gw(wo),
4966 gw(&d.gate_proj),
4967 gw(&d.up_proj),
4968 gw(&d.down_proj),
4969 ) else {
4970 return crate::gpu::BatchGraphOutcome::Declined;
4971 };
4972 let layer = crate::gpu::GraphLayer {
4973 input_norm: &lw.input_norm,
4974 attn: crate::gpu::GraphAttn::Full {
4975 wq: gwq,
4976 wk: gwk,
4977 wv: gwv,
4978 wo: gwo,
4979 q_norm: q_norm.as_deref(),
4980 k_norm: k_norm.as_deref(),
4981 late_qk_norm: self.qk_norm_after_rope,
4982 bias: bias
4983 .as_ref()
4984 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
4985 output_gate: *output_gate,
4986 cpu_k: m.kv.k_heads(),
4987 cpu_v: m.kv.v_heads(),
4988 },
4989 post_norm: &lw.post_norm,
4990 ffn: crate::gpu::GraphFfn::Dense {
4991 gate: gg,
4992 up: gu,
4993 down: gd,
4994 },
4995 };
4996 let positions: Vec<usize> = (first_pos..first_pos + pairs.len()).collect();
4997 let nh = self.num_heads;
4998 let (nkv, hd, rd) = self.layer_geom(0);
4999 let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
5000 crate::gpu::forward_batch_graph(
5001 &model,
5002 self.mtp_kv_id(),
5003 std::slice::from_ref(&layer),
5004 &self.inv_freq,
5005 &mut hiddens,
5006 nh,
5007 nkv,
5008 hd,
5009 rd,
5010 hs,
5011 self.intermediate_size,
5012 &positions,
5013 self.kv_cache.max_seq_len,
5014 gemma,
5015 self.rms_eps as f32,
5016 self.attn_scale,
5017 pairs.len(),
5018 &[],
5019 0,
5020 None,
5021 )
5022 }
5023
5024 #[cfg(feature = "gpu")]
5031 fn mtp_warm_graph_fallback(
5032 &mut self,
5033 m: &mut MtpModule,
5034 pairs: &[(&[f32], u32)],
5035 first_pos: usize,
5036 ) -> bool {
5037 if pairs.is_empty() {
5038 return true;
5039 }
5040 let graphable = self.mtp_block_graph_ok(m);
5041 if !graphable {
5042 if self.mtp_graph_mode == Some(true) {
5046 return false;
5047 }
5048 self.mtp_graph_mode = Some(false);
5049 for (j, (h, t)) in pairs.iter().enumerate() {
5050 self.mtp_warm(m, h, *t, first_pos + j);
5051 }
5052 return true;
5053 }
5054
5055 for (j, (h, t)) in pairs.iter().enumerate() {
5060 if self.mtp_step_graph(m, h, *t, first_pos + j).is_none() {
5061 return false;
5062 }
5063 }
5064 self.mtp_graph_mode = Some(true);
5065 true
5066 }
5067
5068 #[cfg(feature = "gpu")]
5073 fn mtp_warm_prefill_pairs(
5074 &mut self,
5075 m: &mut MtpModule,
5076 pairs: &[(&[f32], u32)],
5077 first_pos: usize,
5078 ) -> Result<(), &'static str> {
5079 if self.mtp_graph_mode == Some(false) || !self.mtp_graph_ok(m) {
5084 if self.mtp_graph_mode == Some(true) {
5085 return Err("MTP token graph became unavailable after admission");
5086 }
5087 self.mtp_graph_mode = Some(false);
5088 for (j, (h, t)) in pairs.iter().enumerate() {
5089 self.mtp_warm(m, h, *t, first_pos + j);
5090 }
5091 return Ok(());
5092 }
5093 match self.mtp_warm_graph(m, pairs, first_pos) {
5094 crate::gpu::BatchGraphOutcome::Completed => {
5095 if !pairs.is_empty() {
5096 self.mtp_graph_mode = Some(true);
5097 }
5098 Ok(())
5099 }
5100 crate::gpu::BatchGraphOutcome::Declined => {
5101 if self.mtp_warm_graph_fallback(m, pairs, first_pos) {
5102 Ok(())
5103 } else {
5104 Err("MTP warm-up fallback failed after device admission")
5105 }
5106 }
5107 crate::gpu::BatchGraphOutcome::Failed => {
5108 Err("MTP warm batch graph failed after admission")
5109 }
5110 }
5111 }
5112
5113 #[cfg(not(feature = "gpu"))]
5114 fn mtp_warm_prefill_pairs(
5115 &mut self,
5116 m: &mut MtpModule,
5117 pairs: &[(&[f32], u32)],
5118 first_pos: usize,
5119 ) -> Result<(), &'static str> {
5120 for (j, (h, t)) in pairs.iter().enumerate() {
5121 self.mtp_warm(m, h, *t, first_pos + j);
5122 }
5123 Ok(())
5124 }
5125
5126 fn mtp_warm(&mut self, m: &mut MtpModule, hidden: &[f32], next_token: u32, position: usize) {
5130 let e = self.embed_single(next_token);
5131 let mut cat = vec![0.0f32; 2 * self.hidden_size];
5132 let (cat_e, cat_h) = cat.split_at_mut(self.hidden_size);
5133 inference::rms_norm_into(&e, &m.enorm, self.rms_eps, self.norm_style, cat_e);
5134 inference::rms_norm_into(hidden, &m.hnorm, self.rms_eps, self.norm_style, cat_h);
5135 let mut x = vec![0.0f32; self.hidden_size];
5136 m.eh_proj.matvec(&cat, &mut x, self.pool.as_deref());
5137 inference::rms_norm_into(
5138 &x,
5139 &m.layer.input_norm,
5140 self.rms_eps,
5141 self.norm_style,
5142 &mut self.ws.n1,
5143 );
5144 let attn = match &m.layer.attn {
5145 AttnKind::Full {
5146 wq,
5147 wk,
5148 wv,
5149 wo,
5150 q_norm,
5151 k_norm,
5152 output_gate,
5153 softplus_gate,
5154 bias,
5155 } => {
5156 let mut cfg = self.attn_cfg(position);
5157 cfg.q_norm = q_norm.as_deref();
5158 cfg.k_norm = k_norm.as_deref();
5159 cfg.output_gate = *output_gate;
5160 cfg.softplus_gate = softplus_gate.as_ref().map(|(g, p)| (g, *p));
5161 cfg.bias = bias
5162 .as_ref()
5163 .map(|(q, k, v)| (q.as_slice(), k.as_slice(), v.as_slice()));
5164 attention::qwen_attention(&self.ws.n1, wq, wk, wv, wo, &mut m.kv, &cfg)
5165 }
5166 _ => return,
5167 };
5168 let _ = attn;
5169 }
5170
5171 #[cfg(feature = "gpu")]
5178 #[allow(clippy::too_many_arguments)]
5179 fn graph_spec_step(
5180 &mut self,
5181 m: &mut MtpModule,
5182 hidden: &[f32],
5183 t_next: u32,
5184 next_pos: usize,
5185 drafted: &mut usize,
5186 accepted: &mut usize,
5187 all_ids: &mut Vec<u32>,
5191 room: usize,
5196 ) -> Option<(Vec<u32>, usize, Vec<f32>)> {
5197 #[cfg(target_os = "macos")]
5208 let metal_native = crate::gpu::q1_force();
5209 #[cfg(not(target_os = "macos"))]
5210 let metal_native = false;
5211 #[cfg(feature = "gpu")]
5212 let k_default = if metal_native {
5213 7
5216 } else if crate::gpu_wgpu::verify_i8_on() {
5217 5
5218 } else {
5219 4
5220 };
5221 #[cfg(not(feature = "gpu"))]
5222 let k_default = 4;
5223 let k_env: Option<usize> = std::env::var("CMF_GRAPH_SPEC_K")
5224 .ok()
5225 .and_then(|v| v.parse().ok())
5226 .filter(|&v| (1..=8).contains(&v));
5227 let (k_start, k_max) = if metal_native { (7, 7) } else { (3, k_default.max(5)) };
5232 let k_full: usize = k_env.unwrap_or_else(|| self.spec_k_adapt.unwrap_or(k_start));
5233 let k_spec = k_full.min(room).max(1);
5234 let k_capped = k_spec < k_full;
5237 if next_pos == 0 {
5238 return None;
5239 }
5240 let t_round = std::time::Instant::now();
5241 let subs = || crate::gpu_wgpu::SUBMITS.load(std::sync::atomic::Ordering::Relaxed);
5257 let sub0 = subs();
5258 let cfg = self.sampler_config.clone();
5263 let penalized = !(cfg.repetition_penalty == 1.0
5264 && cfg.presence_penalty == 0.0
5265 && cfg.suppress_tokens.is_empty());
5266 let greedy_pen = cfg.temperature < 1e-6 && penalized;
5271 let sampling = cfg.temperature >= 1e-6;
5272 let sparse = sampling && sampler::sparse_ok(&cfg);
5278 let base_len = all_ids.len();
5279 if sampling && !sparse && self.spec_q.len() < k_spec {
5280 self.spec_q.resize_with(k_spec, Vec::new);
5281 }
5282 if sparse && self.spec_qs.len() < k_spec {
5283 self.spec_qs.resize_with(k_spec, Vec::new);
5284 }
5285 let mut drafts = Vec::with_capacity(k_spec);
5290 let mut hx = hidden.to_vec();
5291 let spec_dbg = std::env::var("CMF_SPEC_DBG").is_ok();
5294 spec_stamp("pro");
5295 #[cfg(target_os = "macos")]
5301 if metal_native && !sampling && !greedy_pen && self.mtp_graph_mode != Some(false) {
5302 match self.mtp_draft_chain_metal(m, hidden, t_next, next_pos - 1, k_spec) {
5303 Ok(ids) => {
5304 self.mtp_graph_mode = Some(true);
5305 drafts = ids;
5306 }
5307 Err(true) => {
5308 tracing::error!("mtp Metal draft chain failed after commit");
5309 self.clear_sequence_state();
5310 self.graph_failed
5311 .store(true, std::sync::atomic::Ordering::Relaxed);
5312 self.cancel
5313 .store(true, std::sync::atomic::Ordering::Relaxed);
5314 return None;
5315 }
5316 Err(false) => {}
5317 }
5318 }
5319 for j in drafts.len()..k_spec {
5320 let tok_in = if j == 0 { t_next } else { drafts[j - 1] };
5321 let mut dbg_ref: Option<(Vec<f32>, Vec<f32>)> = None;
5322 if spec_dbg {
5323 let saved = self.mtp_graph_mode;
5324 self.mtp_graph_mode = Some(false);
5325 let r = self.mtp_step_hl(m, &hx, tok_in, next_pos - 1 + j);
5326 self.mtp_graph_mode = saved;
5327 if self.graph_failed.load(std::sync::atomic::Ordering::Relaxed) {
5328 return None;
5329 }
5330 m.kv.truncate_last(1);
5331 dbg_ref = Some(r);
5332 }
5333 let (mut lg, hj) = self.mtp_step_hl(m, &hx, tok_in, next_pos - 1 + j);
5334 if self.graph_failed.load(std::sync::atomic::Ordering::Relaxed) {
5335 return None;
5336 }
5337 if let Some((lg_cpu, h_cpu)) = dbg_ref {
5338 let n = |v: &[f32]| v.iter().map(|x| x * x).sum::<f32>().sqrt();
5339 let dl = lg
5340 .iter()
5341 .zip(&lg_cpu)
5342 .fold(0f32, |m, (a, b)| m.max((a - b).abs()));
5343 let dh = hj
5344 .iter()
5345 .zip(&h_cpu)
5346 .fold(0f32, |m, (a, b)| m.max((a - b).abs()));
5347 eprintln!(
5348 "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 {}",
5349 next_pos - 1 + j,
5350 sampler::argmax(&lg_cpu),
5351 sampler::argmax(&lg),
5352 n(&h_cpu),
5353 n(&hj),
5354 m.kv.seq_len
5355 );
5356 }
5357 let dj = if sparse {
5358 let mut q = std::mem::take(&mut self.spec_qs[j]);
5359 let ok = sampler::sparse_distribution_into(
5360 &lg,
5361 &cfg,
5362 all_ids,
5363 &mut self.sampler_scratch,
5364 self.pool.as_deref(),
5365 &mut q,
5366 );
5367 let d = if ok {
5368 sampler::draw_sparse(&q, &mut self.rng)
5369 } else {
5370 let t = sampler::argmax(&lg);
5372 q.clear();
5373 q.push((t, 1.0));
5374 t
5375 };
5376 self.spec_qs[j] = q;
5377 all_ids.push(d);
5378 d
5379 } else if sampling {
5380 let mut q = std::mem::take(&mut self.spec_q[j]);
5381 sampler::distribution_into(
5382 &lg,
5383 &cfg,
5384 all_ids,
5385 &mut self.sampler_scratch,
5386 self.pool.as_deref(),
5387 &mut q,
5388 );
5389 let d = sampler::draw(&q, &mut self.rng);
5390 self.spec_q[j] = q;
5391 all_ids.push(d); d
5393 } else if greedy_pen {
5394 let d = sampler::argmax_penalized(
5395 &lg,
5396 &cfg,
5397 all_ids,
5398 &mut self.sampler_scratch,
5399 self.pool.as_deref(),
5400 );
5401 all_ids.push(d);
5402 d
5403 } else {
5404 sampler::argmax(&lg)
5405 };
5406 attention::recycle_buf(&mut lg);
5407 drafts.push(dj);
5408 hx = hj;
5409 spec_stamp("d.pick");
5410 }
5411 all_ids.truncate(base_len);
5412 *drafted += k_spec;
5413 let t_draft = t_round.elapsed();
5414 let sub_draft = subs();
5415 let b = k_spec + 1;
5418 let mut hiddens = vec![0.0f32; b * self.hidden_size];
5419 for (i, &t) in std::iter::once(&t_next).chain(drafts.iter()).enumerate() {
5420 let e = self.embed_single(t);
5421 hiddens[i * self.hidden_size..(i + 1) * self.hidden_size].copy_from_slice(&e);
5422 }
5423 let positions: Vec<usize> = (next_pos..next_pos + b).collect();
5424 spec_stamp("v.emb");
5425 let (lm_gw, lm_rows) = {
5426 let (_, i, kind, rs) = self.weights.lm_head.graph_weight()?;
5427 (
5428 crate::gpu::GraphW {
5429 idx: i,
5430 kind,
5431 row_scale: rs,
5432 data: &[],
5433 prism: crate::gpu::GraphPrismOp::None,
5434 affine: false,
5435 },
5436 self.weights.lm_head.rows(),
5437 )
5438 };
5439 let mut logits = Vec::new();
5440 let final_norm = self.weights.final_norm.clone();
5441 #[cfg(target_os = "macos")]
5450 let greedy_dev = metal_native
5451 && !sampling
5452 && !greedy_pen
5453 && !self.confidence_on
5454 && self.final_softcap.is_none()
5455 && self.vocab_size == lm_rows
5461 && std::env::var_os("CMF_METAL_VERIFY_CHECK").is_none()
5462 && std::env::var_os("CMF_LOGIT_DUMP").is_none()
5463 && std::env::var("CMF_METAL_DEV_ARGMAX").as_deref() != Ok("0");
5464 #[cfg(not(target_os = "macos"))]
5465 let greedy_dev = false;
5466 let mut dev_ids: Vec<u32> = Vec::new();
5467 #[cfg(target_os = "macos")]
5468 let verify_outcome = if metal_native {
5469 let lm = self.weights.lm_head.q1_parts()?;
5470 let n_score = self.vocab_size.min(lm_rows);
5471 self.try_batch_graph_metal(
5472 &mut hiddens,
5473 &positions,
5474 b,
5475 Some((lm, &final_norm, &mut logits)),
5476 if greedy_dev {
5477 Some((n_score, &mut dev_ids))
5478 } else {
5479 None
5480 },
5481 )
5482 } else {
5483 self.try_batch_graph_wgpu(
5484 &mut hiddens,
5485 &positions,
5486 b,
5487 Some(crate::gpu::SpecTail {
5488 lm: lm_gw,
5489 lm_rows,
5490 final_norm: &final_norm,
5491 logits_out: &mut logits,
5492 }),
5493 )
5494 };
5495 #[cfg(not(target_os = "macos"))]
5496 let verify_outcome = self.try_batch_graph_wgpu(
5497 &mut hiddens,
5498 &positions,
5499 b,
5500 Some(crate::gpu::SpecTail {
5501 lm: lm_gw,
5502 lm_rows,
5503 final_norm: &final_norm,
5504 logits_out: &mut logits,
5505 }),
5506 );
5507 match verify_outcome {
5508 crate::gpu::BatchGraphOutcome::Completed => {}
5509 crate::gpu::BatchGraphOutcome::Declined => {
5510 m.kv.truncate_last(k_spec);
5514 if !metal_native && !self.rewind_mtp_graph_mirror(next_pos) {
5515 self.clear_sequence_state();
5516 self.graph_failed
5517 .store(true, std::sync::atomic::Ordering::Relaxed);
5518 self.cancel
5519 .store(true, std::sync::atomic::Ordering::Relaxed);
5520 tracing::error!("MTP graph mirror rewind failed after verify decline");
5521 }
5522 return None;
5523 }
5524 crate::gpu::BatchGraphOutcome::Failed => {
5525 self.clear_sequence_state();
5529 self.graph_failed
5530 .store(true, std::sync::atomic::Ordering::Relaxed);
5531 self.cancel
5532 .store(true, std::sync::atomic::Ordering::Relaxed);
5533 tracing::error!("MTP verify batch graph failed after admission");
5534 return None;
5535 }
5536 }
5537 #[cfg(target_os = "macos")]
5543 if metal_native && std::env::var("CMF_METAL_VERIFY_CHECK").as_deref() == Ok("1") {
5544 let snap: Vec<Vec<f32>> = self
5545 .kv_cache
5546 .layers
5547 .iter()
5548 .map(|l| l.linear_state.clone())
5549 .collect();
5550 let attn_lens: Vec<usize> = self.kv_cache.layers.iter().map(|l| l.seq_len).collect();
5551 let toks: Vec<u32> = std::iter::once(t_next)
5552 .chain(drafts.iter().copied())
5553 .collect();
5554 let want_save = self.graph_want_logits;
5555 self.graph_want_logits = false;
5556 for (i, &t) in toks.iter().enumerate() {
5557 let hi = self.forward_layers(&self.embed_single(t), next_pos + i, None);
5558 let _ = self.graph_logits.take();
5559 if std::env::var("CMF_SPEC_PLAIN_HIDDEN").as_deref() == Ok("1") {
5563 hiddens[i * self.hidden_size..(i + 1) * self.hidden_size].copy_from_slice(&hi);
5564 }
5565 let ref_lg = self.logits_from_hidden(&hi);
5566 let row = &logits[i * lm_rows..(i + 1) * lm_rows];
5567 let ra = sampler::argmax(&ref_lg);
5568 let va = sampler::argmax(row);
5569 let mut md = 0f32;
5570 let mut rms = 0f64;
5571 for j in 0..lm_rows.min(ref_lg.len()) {
5572 let d = (ref_lg[j] - row[j]).abs();
5573 md = md.max(d);
5574 rms += (d as f64) * (d as f64);
5575 }
5576 let mut hd = 0f32;
5577 for j in 0..self.hidden_size {
5578 hd = hd.max((hi[j] - hiddens[i * self.hidden_size + j]).abs());
5579 }
5580 eprintln!(
5581 "verify-check row {i} tok {t} pos {}: ref argmax {ra} verify argmax {va} {} | max|dlogit| {md:.3} rms {:.4} | max|dhidden| {hd:.4}",
5582 next_pos + i,
5583 if ra == va { "OK" } else { "MISMATCH" },
5584 (rms / lm_rows as f64).sqrt()
5585 );
5586 }
5587 self.graph_want_logits = want_save;
5588 for (l, st) in self.kv_cache.layers.iter_mut().zip(snap) {
5591 if l.linear_state.len() == st.len() {
5592 l.linear_state.copy_from_slice(&st);
5593 } else {
5594 l.linear_state = st;
5595 }
5596 }
5597 for (li, (l, n0)) in self.kv_cache.layers.iter_mut().zip(attn_lens).enumerate() {
5598 let extra = l.seq_len.saturating_sub(n0);
5599 if extra > 0 {
5600 l.truncate_last(extra);
5601 crate::gpu_metal::kv_mirror_set_stored(self.graph_kv_id, li, n0);
5602 }
5603 }
5604 }
5605 let t_verify = t_round.elapsed();
5606 let sub_verify = subs();
5607 let mut a = 0usize;
5612 let mut forced: Option<u32> = None;
5613 let ids: Vec<u32> = if sparse {
5614 let mut p = std::mem::take(&mut self.spec_ps);
5615 let mut res = std::mem::take(&mut self.spec_ress);
5616 while a < k_spec {
5617 let ok = sampler::sparse_distribution_into(
5618 &logits[a * lm_rows..(a + 1) * lm_rows],
5619 &cfg,
5620 all_ids,
5621 &mut self.sampler_scratch,
5622 self.pool.as_deref(),
5623 &mut p,
5624 );
5625 if !ok {
5626 let t = sampler::argmax(&logits[a * lm_rows..(a + 1) * lm_rows]);
5627 p.clear();
5628 p.push((t, 1.0));
5629 }
5630 match sampler::spec_accept_or_correct_sparse(
5631 &p,
5632 &self.spec_qs[a],
5633 drafts[a],
5634 &mut self.rng,
5635 &mut res,
5636 ) {
5637 None => {
5638 all_ids.push(drafts[a]);
5639 a += 1;
5640 }
5641 Some(c) => {
5642 forced = Some(c);
5643 break;
5644 }
5645 }
5646 }
5647 all_ids.truncate(base_len);
5648 self.spec_ps = p;
5649 self.spec_ress = res;
5650 drafts.clone()
5651 } else if sampling {
5652 let mut p = std::mem::take(&mut self.spec_p);
5653 let mut res = std::mem::take(&mut self.spec_res);
5654 while a < k_spec {
5655 sampler::distribution_into(
5656 &logits[a * lm_rows..(a + 1) * lm_rows],
5657 &cfg,
5658 all_ids,
5659 &mut self.sampler_scratch,
5660 self.pool.as_deref(),
5661 &mut p,
5662 );
5663 match sampler::spec_accept_or_correct(
5664 &p,
5665 &self.spec_q[a],
5666 drafts[a],
5667 &mut self.rng,
5668 &mut res,
5669 self.pool.as_deref(),
5670 ) {
5671 None => {
5672 all_ids.push(drafts[a]);
5673 a += 1;
5674 }
5675 Some(c) => {
5676 forced = Some(c);
5677 break;
5678 }
5679 }
5680 }
5681 all_ids.truncate(base_len);
5682 self.spec_p = p;
5683 self.spec_res = res;
5684 drafts.clone()
5686 } else if greedy_pen {
5687 let mut ids: Vec<u32> = Vec::with_capacity(b);
5691 for i in 0..b {
5692 let t = sampler::argmax_penalized(
5693 &logits[i * lm_rows..(i + 1) * lm_rows],
5694 &cfg,
5695 all_ids,
5696 &mut self.sampler_scratch,
5697 self.pool.as_deref(),
5698 );
5699 ids.push(t);
5700 if i < k_spec && t == drafts[i] {
5701 all_ids.push(t);
5702 } else {
5703 break;
5704 }
5705 }
5706 all_ids.truncate(base_len);
5707 while a < k_spec && a < ids.len() && ids[a] == drafts[a] {
5708 a += 1;
5709 }
5710 ids
5713 } else if greedy_dev && dev_ids.len() == b {
5714 let ids = std::mem::take(&mut dev_ids);
5715 while a < k_spec && ids[a] == drafts[a] {
5716 a += 1;
5717 }
5718 ids
5719 } else {
5720 if logits.len() < b * lm_rows {
5721 self.clear_sequence_state();
5724 self.graph_failed
5725 .store(true, std::sync::atomic::Ordering::Relaxed);
5726 self.cancel
5727 .store(true, std::sync::atomic::Ordering::Relaxed);
5728 tracing::error!("Metal verify returned neither logits nor argmax ids");
5729 return None;
5730 }
5731 let ids: Vec<u32> = (0..b)
5732 .map(|i| sampler::argmax(&logits[i * lm_rows..(i + 1) * lm_rows]))
5733 .collect();
5734 while a < k_spec && ids[a] == drafts[a] {
5735 a += 1;
5736 }
5737 ids
5738 };
5739 spec_stamp("acc");
5740 if spec_dbg {
5741 eprintln!(
5742 "spec-dbg round: t_next {t_next} drafts {:?} verified {:?} accepted {a}",
5743 drafts, ids
5744 );
5745 }
5746 #[cfg(target_os = "macos")]
5750 let commit_ref: Option<(Vec<Vec<f32>>, Vec<(usize, Vec<f32>, Vec<f32>)>)> = if metal_native
5751 && std::env::var("CMF_METAL_VERIFY_CHECK").as_deref() == Ok("2")
5752 {
5753 let snap: Vec<Vec<f32>> = self
5754 .kv_cache
5755 .layers
5756 .iter()
5757 .map(|l| l.linear_state.clone())
5758 .collect();
5759 let attn_lens: Vec<usize> = self.kv_cache.layers.iter().map(|l| l.seq_len).collect();
5760 let toks: Vec<u32> = std::iter::once(t_next)
5761 .chain(drafts.iter().copied())
5762 .collect();
5763 let want_save = self.graph_want_logits;
5764 self.graph_want_logits = false;
5765 for (i, &t) in toks.iter().take(a + 1).enumerate() {
5766 let _ = self.forward_layers(&self.embed_single(t), next_pos + i, None);
5767 let _ = self.graph_logits.take();
5768 }
5769 self.graph_want_logits = want_save;
5770 let plain_states: Vec<Vec<f32>> = self
5771 .kv_cache
5772 .layers
5773 .iter()
5774 .map(|l| l.linear_state.clone())
5775 .collect();
5776 let (nkv, hd) = (self.num_kv_heads, self.head_dim);
5777 let mut rows = Vec::new();
5778 for (li, (l, n0)) in self
5779 .kv_cache
5780 .layers
5781 .iter_mut()
5782 .zip(attn_lens.iter())
5783 .enumerate()
5784 {
5785 let extra = l.seq_len.saturating_sub(*n0);
5786 if extra > 0 {
5787 let mut kk = Vec::new();
5788 let mut vv = Vec::new();
5789 for g in 0..nkv {
5790 kk.extend_from_slice(&l.head_keys(g)[n0 * hd..]);
5791 vv.extend_from_slice(&l.head_values(g)[n0 * hd..]);
5792 }
5793 rows.push((li, kk, vv));
5794 l.truncate_last(extra);
5795 crate::gpu_metal::kv_mirror_set_stored(self.graph_kv_id, li, *n0);
5796 }
5797 }
5798 for (l, st) in self.kv_cache.layers.iter_mut().zip(snap) {
5799 if l.linear_state.len() == st.len() {
5800 l.linear_state.copy_from_slice(&st);
5801 } else {
5802 l.linear_state = st;
5803 }
5804 }
5805 Some((plain_states, rows))
5806 } else {
5807 None
5808 };
5809 let warm_off = std::env::var("CMF_SPEC_WARM").is_ok_and(|v| v == "0");
5810 #[cfg(target_os = "macos")]
5819 let mut warm_pending: Option<MetalWarmPending> = None;
5820 #[cfg(target_os = "macos")]
5821 if metal_native {
5822 m.kv.truncate_last(k_spec.saturating_sub(1));
5823 if self.mtp_graph_mode == Some(true) {
5824 crate::gpu_metal::kv_mirror_set_stored(
5827 self.mtp_kv_id(),
5828 Self::MTP_LAYER_BASE,
5829 m.kv.seq_len,
5830 );
5831 if !warm_off && a > 0 {
5832 let pairs: Vec<(&[f32], u32)> = (0..a)
5833 .map(|j| {
5834 (
5835 &hiddens[j * self.hidden_size..(j + 1) * self.hidden_size],
5836 ids[j],
5837 )
5838 })
5839 .collect();
5840 warm_pending = self.mtp_warm_batch_submit(m, &pairs, next_pos);
5841 }
5842 }
5843 spec_stamp("c.wsub");
5844 }
5845 #[cfg(target_os = "macos")]
5847 if metal_native {
5848 if !self.metal_verify_commit(a) {
5851 self.clear_sequence_state();
5852 self.graph_failed
5853 .store(true, std::sync::atomic::Ordering::Relaxed);
5854 self.cancel
5855 .store(true, std::sync::atomic::Ordering::Relaxed);
5856 tracing::error!("Metal verify state/KV handoff failed after admission");
5857 return None;
5858 }
5859 if let Some((plain_states, rows)) = commit_ref {
5860 crate::gpu_metal::queue_fence();
5861 let _ = crate::gpu_metal::wait_replay();
5864 let (nkv, hd) = (self.num_kv_heads, self.head_dim);
5865 let mut worst_s = 0f32;
5866 let mut worst_li = 0usize;
5867 for (li, (l, ps)) in self.kv_cache.layers.iter().zip(&plain_states).enumerate() {
5868 if l.linear_state.len() != ps.len() || ps.is_empty() {
5869 continue;
5870 }
5871 let d = l
5872 .linear_state
5873 .iter()
5874 .zip(ps)
5875 .fold(0f32, |m, (x, y)| m.max((x - y).abs()));
5876 let n = ps.iter().fold(0f32, |m, y| m.max(y.abs()));
5877 let rel = d / n.max(1e-6);
5878 if rel > worst_s {
5879 worst_s = rel;
5880 worst_li = li;
5881 }
5882 }
5883 let mut worst_k = 0f32;
5884 for (li, kk, vv) in &rows {
5885 let l = &self.kv_cache.layers[*li];
5886 let n0 = l.seq_len - (kk.len() / (nkv * hd));
5887 let mut ck = Vec::new();
5888 let mut cv = Vec::new();
5889 for g in 0..nkv {
5890 ck.extend_from_slice(&l.head_keys(g)[n0 * hd..]);
5891 cv.extend_from_slice(&l.head_values(g)[n0 * hd..]);
5892 }
5893 if ck.len() == kk.len() {
5894 let dk = ck
5895 .iter()
5896 .zip(kk)
5897 .fold(0f32, |m, (x, y)| m.max((x - y).abs()));
5898 let dv = cv
5899 .iter()
5900 .zip(vv)
5901 .fold(0f32, |m, (x, y)| m.max((x - y).abs()));
5902 worst_k = worst_k.max(dk).max(dv);
5903 } else {
5904 eprintln!(
5905 "commit-check L{li}: kv row count mismatch {} vs {}",
5906 ck.len(),
5907 kk.len()
5908 );
5909 }
5910 }
5911 eprintln!(
5912 "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}"
5913 );
5914 }
5915 }
5916 if !metal_native && a + 1 < b {
5917 let expected_gdn_layers = self.graph_gdn_layer_count();
5918 if expected_gdn_layers > 0
5919 && !crate::gpu::gdn_spec_restore(self.graph_kv_id, a, next_pos, expected_gdn_layers)
5920 {
5921 self.clear_sequence_state();
5922 self.graph_failed
5923 .store(true, std::sync::atomic::Ordering::Relaxed);
5924 self.cancel
5925 .store(true, std::sync::atomic::Ordering::Relaxed);
5926 tracing::error!("GDN speculative restore failed after verify");
5927 return None;
5928 }
5929 }
5930 if !metal_native && !self.rewind_trunk_graph_mirrors(next_pos + a + 1) {
5931 self.clear_sequence_state();
5936 self.graph_failed
5937 .store(true, std::sync::atomic::Ordering::Relaxed);
5938 self.cancel
5939 .store(true, std::sync::atomic::Ordering::Relaxed);
5940 tracing::error!("trunk graph KV rewind failed after speculative verify");
5941 return None;
5942 }
5943 *accepted += a;
5944 if !metal_native {
5955 m.kv.truncate_last(k_spec.saturating_sub(1));
5957 }
5958 spec_stamp("c.trunc");
5959 if !metal_native
5960 && self.mtp_graph_mode == Some(true)
5961 && !self.rewind_mtp_graph_mirror(next_pos)
5962 {
5963 self.clear_sequence_state();
5967 self.graph_failed
5968 .store(true, std::sync::atomic::Ordering::Relaxed);
5969 self.cancel
5970 .store(true, std::sync::atomic::Ordering::Relaxed);
5971 tracing::error!("MTP graph mirror rewind failed after verify commit");
5972 return None;
5973 }
5974 if !warm_off && a > 0 {
5975 let mut warmed = false;
5978 #[cfg(target_os = "macos")]
5979 if metal_native && self.mtp_graph_mode == Some(true) {
5980 warmed = match warm_pending.take() {
5984 Some(p) => self.mtp_warm_batch_finish(m, p),
5985 None => false,
5986 };
5987 if !warmed {
5988 warmed = true;
5989 for j in 0..a {
5990 let row =
5991 hiddens[j * self.hidden_size..(j + 1) * self.hidden_size].to_vec();
5992 if self
5993 .mtp_step_metal(m, &row, ids[j], next_pos + j, false)
5994 .is_none()
5995 {
5996 warmed = false;
5997 break;
5998 }
5999 }
6000 }
6001 }
6002 if !warmed && self.mtp_graph_mode != Some(false) && !metal_native {
6003 let rows: Vec<Vec<f32>> = (0..a)
6004 .map(|j| hiddens[j * self.hidden_size..(j + 1) * self.hidden_size].to_vec())
6005 .collect();
6006 let pairs: Vec<(&[f32], u32)> = rows
6007 .iter()
6008 .zip(ids.iter())
6009 .map(|(r, &t)| (r.as_slice(), t))
6010 .collect();
6011 match self.mtp_warm_prefill_pairs(m, &pairs, next_pos) {
6012 Ok(()) => warmed = true,
6013 Err(err) => {
6014 tracing::error!("{err}");
6020 self.clear_sequence_state();
6021 self.graph_failed
6022 .store(true, std::sync::atomic::Ordering::Relaxed);
6023 self.cancel
6024 .store(true, std::sync::atomic::Ordering::Relaxed);
6025 return None;
6026 }
6027 }
6028 }
6029 if !warmed {
6030 for j in 0..a {
6031 let row = &hiddens[j * self.hidden_size..(j + 1) * self.hidden_size];
6032 let row = row.to_vec();
6033 self.mtp_warm(m, &row, ids[j], next_pos + j);
6034 }
6035 }
6036 }
6037 spec_stamp("c.warm");
6041 if let Some(c) = forced {
6042 self.spec_forced = Some(c);
6043 self.graph_logits = None;
6044 } else if greedy_dev && logits.is_empty() {
6045 self.spec_forced = Some(ids[a]);
6048 self.graph_logits = None;
6049 } else {
6050 let mut row = logits[a * lm_rows..(a + 1) * lm_rows].to_vec();
6051 row.resize(self.vocab_size, 0.0);
6052 if let Some(c) = self.final_softcap {
6053 for l in row.iter_mut() {
6054 *l = c * (*l / c).tanh();
6055 }
6056 }
6057 self.graph_logits = Some(row);
6058 }
6059 let new_hidden = hiddens[a * self.hidden_size..(a + 1) * self.hidden_size].to_vec();
6060 spec_stamp("c.row");
6061 if std::env::var("CMF_GRAPH_SPEC_TIME").is_ok() {
6067 let end = subs();
6068 eprintln!(
6069 "spec-round: draft {:.1} ms/{} sub | verify {:.1} ms/{} sub | \
6070 commit {:.1} ms/{} sub (accepted {a} of {k_spec}, full-head streak {})",
6071 t_draft.as_secs_f64() * 1e3,
6072 sub_draft - sub0,
6073 (t_verify - t_draft).as_secs_f64() * 1e3,
6074 sub_verify - sub_draft,
6075 (t_round.elapsed() - t_verify).as_secs_f64() * 1e3,
6076 end - sub_verify,
6077 self.draft_full_streak,
6078 );
6079 }
6080 if k_env.is_none() && !metal_native && !k_capped {
6085 let f = a as f32 / k_spec.max(1) as f32;
6089 self.spec_acc_ewma += 0.2 * (f - self.spec_acc_ewma);
6090 let mut k_next = k_spec;
6091 if self.spec_acc_ewma >= 0.75 && k_spec < k_max {
6092 k_next = k_spec + 1;
6093 } else if self.spec_acc_ewma < 0.4 && k_spec > 2 {
6094 k_next = k_spec - 1;
6095 }
6096 if k_next != k_spec {
6097 self.spec_acc_ewma = 0.6;
6098 if std::env::var("CMF_GRAPH_SPEC_TIME").is_ok() {
6099 eprintln!("spec-k: {k_spec} → {k_next}");
6100 }
6101 }
6102 self.spec_k_adapt = Some(k_next);
6103 }
6104 spec_stamp("end");
6105 Some((drafts[..a].to_vec(), next_pos + a + 1, new_hidden))
6106 }
6107
6108 pub fn measure_pair_fusion(&mut self, iters: usize) -> (f64, f64) {
6117 if !self.pair_supported() {
6118 return (0.0, 0.0);
6119 }
6120 let graph_env = std::env::var_os("CMF_GPU_WGPU_GRAPH");
6127 unsafe { std::env::set_var("CMF_GPU_WGPU_GRAPH", "0") };
6128 let emb1 = self.embed_single(1);
6129 let emb2 = self.embed_single(2);
6130 let pos = self.kv_cache.seq_len();
6131
6132 let t0 = std::time::Instant::now();
6133 for _ in 0..iters {
6134 let _ = self.forward_layers(&emb1, pos, None);
6135 let _ = self.forward_layers(&emb2, pos + 1, None);
6136 for l in &mut self.kv_cache.layers {
6137 l.truncate_last(2);
6138 }
6139 }
6140 let singles_ms = t0.elapsed().as_secs_f64() * 1000.0 / iters as f64;
6141
6142 let t1 = std::time::Instant::now();
6143 for _ in 0..iters {
6144 let _ = self.forward_pair(&emb1, &emb2, pos);
6145 for l in &mut self.kv_cache.layers {
6146 l.truncate_last(2);
6147 }
6148 }
6149 let pair_ms = t1.elapsed().as_secs_f64() * 1000.0 / iters as f64;
6150 match graph_env {
6151 Some(value) => unsafe { std::env::set_var("CMF_GPU_WGPU_GRAPH", value) },
6152 None => unsafe { std::env::remove_var("CMF_GPU_WGPU_GRAPH") },
6153 }
6154 (singles_ms, pair_ms)
6155 }
6156
6157 fn pair_supported(&self) -> bool {
6165 !self.weights.layers.is_empty()
6172 && self.g3n.is_none()
6173 && !self
6174 .weights
6175 .layers
6176 .iter()
6177 .any(|lw| matches!(&lw.attn, AttnKind::Mla(_) | AttnKind::Kda(_)))
6178 }
6179
6180 fn forward_pair(
6181 &mut self,
6182 emb1: &[f32],
6183 emb2: &[f32],
6184 position: usize,
6185 ) -> (Vec<f32>, Vec<f32>) {
6186 let mut h1 = emb1.to_vec();
6187 let mut h2 = emb2.to_vec();
6188 let (_nkv, _hd, hs, _rd, eps) = (
6189 self.num_kv_heads,
6190 self.head_dim,
6191 self.hidden_size,
6192 self.rotary_dim,
6193 self.rms_eps,
6194 );
6195 let pool = self.pool.clone();
6196
6197 for li in 0..self.num_layers {
6198 let lw = &self.weights.layers[self.phys_layer(li)];
6199 inference::rms_norm_into(
6202 &h1,
6203 &lw.input_norm,
6204 self.rms_eps,
6205 self.norm_style,
6206 &mut self.ws.n1,
6207 );
6208 inference::rms_norm_into(
6209 &h2,
6210 &lw.input_norm,
6211 self.rms_eps,
6212 self.norm_style,
6213 &mut self.ws.n2,
6214 );
6215
6216 let (a1, a2) = match &lw.attn {
6217 AttnKind::Mla(_) => unreachable!("MLA has no MTP/pair path"),
6218 AttnKind::Kda(_) => unreachable!("KDA has no MTP/pair path"),
6219 AttnKind::Linear(w) => {
6220 let cfg = self.vmf_cfg.expect("linear layer without vmf_cfg");
6221 let layer = &mut self.kv_cache.layers[li];
6222 let (state, scratch) = (&mut layer.linear_state, &mut layer.linear_scratch);
6223 vmf_phase_pair(
6224 &self.ws.n1,
6225 &self.ws.n2,
6226 w,
6227 &cfg,
6228 state,
6229 scratch,
6230 self.pool.as_deref(),
6231 )
6232 }
6233 AttnKind::LinearGdn(w) => {
6234 let cfg = self.gdn_cfg.expect("gdn layer without gdn_cfg");
6235 let layer = &mut self.kv_cache.layers[li];
6236 let (state, scratch) = (&mut layer.linear_state, &mut layer.linear_scratch);
6237 gdn_pair(
6238 &self.ws.n1,
6239 &self.ws.n2,
6240 w,
6241 &cfg,
6242 state,
6243 scratch,
6244 self.pool.as_deref(),
6245 )
6246 }
6247 AttnKind::ShortConv(w) => {
6248 let cfg = self
6249 .short_conv_cfg
6250 .expect("short-conv layer without short_conv_cfg");
6251 let layer = &mut self.kv_cache.layers[li];
6252 let (state, scratch) = (&mut layer.linear_state, &mut layer.linear_scratch);
6253 short_conv_pair(
6254 &self.ws.n1,
6255 &self.ws.n2,
6256 w,
6257 &cfg,
6258 state,
6259 scratch,
6260 self.pool.as_deref(),
6261 )
6262 }
6263 AttnKind::Full {
6264 wq,
6265 wk,
6266 wv,
6267 wo,
6268 q_norm,
6269 k_norm,
6270 output_gate,
6271 softplus_gate,
6272 bias,
6273 } => {
6274 let inv_freq_l = self.layer_inv_freq(li);
6275 let (nkv_l, hd_l, rd_l) = self.layer_geom(li);
6276 let cfg = QwenAttnCfg {
6277 num_heads: self.layer_num_heads(li),
6278 num_kv_heads: nkv_l,
6279 head_dim: hd_l,
6280 hidden_size: hs,
6281 position,
6282 inv_freq: &inv_freq_l,
6283 rotary_dim: rd_l,
6284 scale: self.attn_scale,
6285 softcap: self.attn_softcap,
6286 window: self.layer_window(li),
6287 v_norm: self.attn_v_norm,
6288 qk_norm_after_rope: self.qk_norm_after_rope,
6289 q_norm: q_norm.as_deref(),
6290 k_norm: k_norm.as_deref(),
6291 output_gate: *output_gate,
6292 softplus_gate: softplus_gate
6293 .as_ref()
6294 .map(|(gate, per_head)| (gate, *per_head)),
6295 rope_scale: self.layer_rope_scale(li),
6296 bias: bias
6297 .as_ref()
6298 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
6299 rms_eps: eps,
6300 norm_style: self.norm_style,
6301 pool: pool.as_deref(),
6302 };
6303 attention::qwen_attention_pair(
6304 &self.ws.n1,
6305 &self.ws.n2,
6306 wq,
6307 wk,
6308 wv,
6309 wo,
6310 &mut self.kv_cache.layers[li],
6311 &cfg,
6312 )
6313 }
6314 };
6315 let (a1, a2) = match &self.weights.layers[self.phys_layer(li)].attn_out_norm {
6316 Some(w) => (
6317 inference::rms_norm(&a1, w, self.rms_eps, self.norm_style),
6318 inference::rms_norm(&a2, w, self.rms_eps, self.norm_style),
6319 ),
6320 None => (a1, a2),
6321 };
6322 for i in 0..self.hidden_size {
6323 h1[i] += a1[i];
6324 h2[i] += a2[i];
6325 }
6326 let (mut a1, mut a2) = (a1, a2);
6327 attention::recycle_buf(&mut a1);
6328 attention::recycle_buf(&mut a2);
6329
6330 let lw = &self.weights.layers[self.phys_layer(li)];
6331 inference::rms_norm_into(
6332 &h1,
6333 &lw.post_norm,
6334 self.rms_eps,
6335 self.norm_style,
6336 &mut self.ws.p1,
6337 );
6338 inference::rms_norm_into(
6339 &h2,
6340 &lw.post_norm,
6341 self.rms_eps,
6342 self.norm_style,
6343 &mut self.ws.p2,
6344 );
6345 let (f1, f2) = match &lw.ffn {
6346 FfnKind::DenseMoe(dm) => (
6349 dense_moe_ffn(
6350 dm,
6351 &self.ws.p1,
6352 &h1,
6353 self.rms_eps,
6354 self.norm_style,
6355 self.pool.as_deref(),
6356 ),
6357 dense_moe_ffn(
6358 dm,
6359 &self.ws.p2,
6360 &h2,
6361 self.rms_eps,
6362 self.norm_style,
6363 self.pool.as_deref(),
6364 ),
6365 ),
6366 _ => ffn_forward_pair(
6367 &lw.ffn,
6368 &self.ws.p1,
6369 &self.ws.p2,
6370 self.pool.as_deref(),
6371 None,
6372 ),
6373 };
6374 let (f1, f2) = match &self.weights.layers[self.phys_layer(li)].ffn_out_norm {
6375 Some(w) => (
6376 inference::rms_norm(&f1, w, self.rms_eps, self.norm_style),
6377 inference::rms_norm(&f2, w, self.rms_eps, self.norm_style),
6378 ),
6379 None => (f1, f2),
6380 };
6381 for i in 0..self.hidden_size {
6382 h1[i] += f1[i];
6383 h2[i] += f2[i];
6384 }
6385 let (mut f1, mut f2) = (f1, f2);
6386 attention::recycle_buf(&mut f1);
6387 attention::recycle_buf(&mut f2);
6388 if let Some(sc) = self.weights.layers[self.phys_layer(li)].layer_scale {
6389 for i in 0..self.hidden_size {
6390 h1[i] *= sc;
6391 h2[i] *= sc;
6392 }
6393 }
6394 if self.is_loop_end(li) && li + 1 < self.num_layers {
6396 h1 = inference::rms_norm(
6397 &h1,
6398 &self.weights.final_norm,
6399 self.rms_eps,
6400 self.norm_style,
6401 );
6402 h2 = inference::rms_norm(
6403 &h2,
6404 &self.weights.final_norm,
6405 self.rms_eps,
6406 self.norm_style,
6407 );
6408 }
6409 }
6410 if self.o1_active() {
6416 self.commit_linear_scratch();
6417 }
6418 self.o1_progress();
6419 (h1, h2)
6420 }
6421
6422 fn commit_linear_scratch(&mut self) {
6424 for layer in &mut self.kv_cache.layers {
6425 if !layer.linear_scratch.is_empty() {
6426 std::mem::swap(&mut layer.linear_state, &mut layer.linear_scratch);
6427 layer.linear_scratch.clear();
6428 }
6429 }
6430 }
6431
6432 pub fn forward_ids(
6435 &mut self,
6436 ids: &[u32],
6437 task_mask: Option<&TaskMask>,
6438 ) -> Result<Vec<f32>, String> {
6439 if ids.is_empty() {
6440 return Err("empty id sequence".to_string());
6441 }
6442 self.clear_sequence_state();
6443 self.check_forward_graph("forward_ids setup", 0)?;
6444 if task_mask.is_none() {
6445 self.o1_begin();
6446 }
6447 let mut hidden = vec![0.0f32; self.hidden_size];
6448 let mut pos = 0usize;
6449 if let Some(b) = &mut self.dsv41 {
6450 let pool = self.pool.clone();
6451 let mut logits = Vec::new();
6452 crate::dsv41::forward_chunk(
6453 &b.0,
6454 &b.1,
6455 &b.2,
6456 &mut b.3,
6457 ids,
6458 0,
6459 pool.as_deref(),
6460 &mut logits,
6461 );
6462 if let Err(err) = self.o1_seal_checked() {
6463 self.clear_sequence_state();
6464 return Err(err);
6465 }
6466 return Ok(logits);
6467 }
6468 if self.can_prefill_batched() && !self.graph_prefill_preferred() && ids.len() > 2 {
6476 let chunk = prefill_chunk();
6480 let hs = self.hidden_size;
6481 while pos < ids.len() {
6482 let end = (pos + chunk).min(ids.len());
6483 let hb = self.prefill_batch_masked(&ids[pos..end], pos, task_mask);
6484 self.check_forward_graph("forward_ids batched prefill", end - 1)?;
6485 hidden.copy_from_slice(&hb[(end - pos - 1) * hs..]);
6486 pos = end;
6487 }
6488 }
6489 if task_mask.is_none()
6498 && !self.graph_prefill_preferred()
6499 && !std::env::var("CMF_PAIR").is_ok_and(|v| v == "0")
6500 && self.pair_supported()
6501 {
6502 while pos + 1 < ids.len() {
6503 let e1 = self.embed_single(ids[pos]);
6504 let e2 = self.embed_single(ids[pos + 1]);
6505 let (_, h2) = self.forward_pair(&e1, &e2, pos);
6506 self.check_forward_graph("forward_ids pair", pos + 1)?;
6507 self.commit_linear_scratch();
6508 hidden = h2;
6509 pos += 2;
6510 }
6511 }
6512 while pos < ids.len() {
6513 hidden = self.forward_layers(&self.embed_single(ids[pos]), pos, task_mask);
6514 self.check_forward_graph("forward_ids", pos)?;
6515 pos += 1;
6516 }
6517 if let Err(err) = self.o1_seal_checked() {
6521 self.clear_sequence_state();
6522 return Err(err);
6523 }
6524 let normed = inference::rms_norm(
6525 &hidden,
6526 &self.weights.final_norm,
6527 self.rms_eps,
6528 self.norm_style,
6529 );
6530 Ok(self.lm_head_forward(&normed))
6531 }
6532
6533 #[doc(hidden)]
6537 pub fn dsv41_serial_logits(&mut self, ids: &[u32]) -> Result<Vec<Vec<f32>>, String> {
6538 #[cfg(target_os = "macos")]
6539 crate::gpu_metal::set_io_namespace(self.graph_kv_id);
6540 if ids.is_empty() {
6541 return Err("empty id sequence".to_string());
6542 }
6543 self.clear_sequence_state();
6544 self.dsv41
6545 .as_ref()
6546 .ok_or_else(|| "dsv41 serial logits require a DeepSeek-V4.1 model".to_string())?;
6547 self.o1_begin();
6548 let rows = {
6549 let pool = self.pool.clone();
6550 let b = self
6551 .dsv41
6552 .as_mut()
6553 .expect("dsv41 checked above; state cannot change during forward");
6554 let mut rows = Vec::with_capacity(ids.len());
6555 for (position, &id) in ids.iter().enumerate() {
6556 let mut logits = Vec::new();
6557 crate::dsv41::forward_token(
6558 &b.0,
6559 &b.1,
6560 &b.2,
6561 &mut b.3,
6562 id,
6563 position,
6564 pool.as_deref(),
6565 &mut logits,
6566 );
6567 rows.push(logits);
6568 }
6569 rows
6570 };
6571 self.o1_seal();
6572 Ok(rows)
6573 }
6574
6575 pub fn ppl_ids(&mut self, ids: &[u32]) -> Result<f64, String> {
6582 let (nll, cnt) = self.nll_ids_from(ids, 0)?;
6583 Ok((nll / cnt.max(1) as f64).exp())
6584 }
6585
6586 pub fn probe_ffn_mass(&mut self, ids: &[u32]) -> Vec<Vec<f64>> {
6591 self.clear_sequence_state();
6592 FFN_PROBE.with(|p| {
6593 *p.borrow_mut() = Some(vec![vec![0f64; self.intermediate_size]; self.num_layers]);
6594 });
6595 crate::gpu::cpu_scope(|| {
6596 for (pos, &id) in ids.iter().enumerate() {
6597 let emb = self.embed_single(id);
6598 let _ = self.forward_layers(&emb, pos, None);
6599 }
6600 });
6601 self.clear_sequence_state();
6602 FFN_PROBE
6603 .with(|p| p.borrow_mut().take())
6604 .unwrap_or_default()
6605 }
6606
6607 pub fn probe_ffn_mass_batch(&mut self, ids: &[u32]) -> Result<Vec<Vec<f64>>, String> {
6611 if let Err(err) = self.nll_begin() {
6612 let _ = FFN_PROBE.with(|p| p.borrow_mut().take());
6616 self.nll_end();
6617 return Err(err);
6618 }
6619 FFN_PROBE.with(|p| {
6620 *p.borrow_mut() = Some(vec![vec![0f64; self.intermediate_size]; self.num_layers]);
6621 });
6622 let result: Result<(), String> = (|| {
6623 for chunk in ids.chunks(256) {
6624 if chunk.len() < 2 {
6625 continue;
6626 }
6627 self.nll_ids_masked(chunk, 0, None)?;
6628 }
6629 Ok(())
6630 })();
6631 self.nll_end();
6632 let probe = FFN_PROBE
6633 .with(|p| p.borrow_mut().take())
6634 .unwrap_or_default();
6635 match result {
6636 Ok(()) => Ok(probe),
6637 Err(err) => {
6638 drop(probe);
6639 Err(err)
6640 }
6641 }
6642 }
6643
6644 pub fn ppl_ids_masked(&mut self, ids: &[u32], mask: &TaskMask) -> Result<f64, String> {
6648 self.nll_begin()?;
6649 let result: Result<f64, String> = (|| {
6650 let mut nll = 0f64;
6651 let mut cnt = 0usize;
6652 let mut hidden = vec![0f32; self.hidden_size];
6653 for (pos, &id) in ids.iter().enumerate() {
6654 if pos > 0 {
6655 inference::rms_norm_into(
6656 &hidden,
6657 &self.weights.final_norm,
6658 self.rms_eps,
6659 self.norm_style,
6660 &mut self.ws.n1,
6661 );
6662 let mut logits = self.lm_head_forward(&self.ws.n1);
6663 let max = logits.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
6664 let sum: f64 = logits.iter().map(|&v| ((v - max) as f64).exp()).sum();
6665 let p = ((logits[id as usize] - max) as f64).exp() / sum.max(1e-300);
6666 nll -= p.max(1e-300).ln();
6667 cnt += 1;
6668 attention::recycle_buf(&mut logits);
6669 }
6670 let emb = self.embed_single(id);
6671 hidden = self.forward_layers(&emb, pos, Some(mask));
6672 self.nll_check_graph("masked serial forward", pos)?;
6673 let _ = self.graph_logits.take();
6677 }
6678 Ok((nll / cnt.max(1) as f64).exp())
6679 })();
6680 self.nll_end();
6681 result
6682 }
6683
6684 pub fn nll_ids_masked(
6703 &mut self,
6704 ids: &[u32],
6705 start: usize,
6706 task_mask: Option<&TaskMask>,
6707 ) -> Result<(f64, usize), String> {
6708 let task_mask = self.drop_open_mask(task_mask);
6709 self.nll_ids_inner(ids, start, task_mask)
6710 }
6711
6712 pub fn nll_ids_from(&mut self, ids: &[u32], start: usize) -> Result<(f64, usize), String> {
6713 self.nll_ids_inner(ids, start, None)
6714 }
6715
6716 fn nll_ids_inner(
6717 &mut self,
6718 ids: &[u32],
6719 start: usize,
6720 task_mask: Option<&TaskMask>,
6721 ) -> Result<(f64, usize), String> {
6722 self.nll_begin()?;
6723 let result: Result<(f64, usize), String> = (|| {
6724 let mut nll = 0f64;
6725 let mut cnt = 0usize;
6726 let (graph_quality, fused_head_quality) = nll_graph_policy(
6739 task_mask.is_none(),
6740 self.graph_prefill_preferred(),
6741 crate::gpu::q1_force(),
6742 );
6743 self.graph_head_required = fused_head_quality;
6744 self.graph_want_logits = fused_head_quality;
6745 #[cfg(target_os = "macos")]
6746 if graph_quality && std::env::var("CMF_METAL_BATCH_NLL").as_deref() != Ok("0") {
6747 match self.nll_batch_metal(ids, start) {
6748 MetalBatchNllOutcome::Completed(nll, count) => {
6749 return Ok((nll, count));
6750 }
6751 MetalBatchNllOutcome::Declined => {}
6752 MetalBatchNllOutcome::Failed(err) => return Err(err),
6753 }
6754 }
6755 if self.can_prefill_batched() && !graph_quality {
6756 const CHUNK: usize = 128;
6762 const LM_SUB: usize = 32;
6763 let n = ids.len().saturating_sub(1);
6764 let hs = self.hidden_size;
6765 let rows = self.weights.lm_head.rows();
6766 let mut pos = 0usize;
6767 while pos < n {
6768 let end = (pos + CHUNK).min(n);
6769 let bsz = end - pos;
6770 let hb = self.prefill_batch_masked(&ids[pos..end], pos, task_mask);
6771 self.nll_check_graph("batched prefill", pos)?;
6772 let mut k0 = 0usize;
6773 while k0 < bsz {
6774 let k1 = (k0 + LM_SUB).min(bsz);
6775 let sb = k1 - k0;
6776 if pos + k1 <= start {
6779 k0 = k1;
6780 continue;
6781 }
6782 let mut normed = vec![0.0f32; sb * hs];
6783 for k in 0..sb {
6784 let r = inference::rms_norm(
6785 &hb[(k0 + k) * hs..(k0 + k + 1) * hs],
6786 &self.weights.final_norm,
6787 self.rms_eps,
6788 self.norm_style,
6789 );
6790 normed[k * hs..(k + 1) * hs].copy_from_slice(&r);
6791 }
6792 let mut logits = vec![0.0f32; sb * rows];
6793 self.weights
6794 .lm_head
6795 .matmat(&normed, sb, &mut logits, self.pool.as_deref());
6796 for k in 0..sb {
6797 if pos + k0 + k < start {
6798 continue;
6799 }
6800 self.nll_check_graph("batched score row", pos + k0 + k)?;
6801 let lg = &mut logits[k * rows..k * rows + self.vocab_size.min(rows)];
6802 if let Some(mu) = self.logit_multiplier {
6803 for v in lg.iter_mut() {
6804 *v *= mu;
6805 }
6806 }
6807 if let Some(c) = self.final_softcap {
6811 for v in lg.iter_mut() {
6812 *v = c * (*v / c).tanh();
6813 }
6814 }
6815 if let Some(cm) = self.head_clusters.clone() {
6818 self.hierarchical_head_logprobs(
6819 &normed[k * hs..(k + 1) * hs],
6820 &cm,
6821 lg,
6822 );
6823 }
6824 let lg = &logits[k * rows..k * rows + self.vocab_size.min(rows)];
6825 let target = ids[pos + k0 + k + 1] as usize;
6826 let max = lg.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
6827 let lse: f64 = lg
6828 .iter()
6829 .map(|&v| ((v - max) as f64).exp())
6830 .sum::<f64>()
6831 .ln()
6832 + max as f64;
6833 nll += lse - lg[target] as f64;
6834 cnt += 1;
6835 if std::env::var("CMF_PPL_TRACE").is_ok() {
6836 let top = lg
6837 .iter()
6838 .enumerate()
6839 .max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
6840 .map(|(i, _)| i)
6841 .unwrap_or(0);
6842 eprintln!(
6843 "BTRACE pos {} target {} nll {:.4} top {} lg_t {:.3} lg_top {:.3}",
6844 pos + k0 + k,
6845 target,
6846 lse - lg[target] as f64,
6847 top,
6848 lg[target],
6849 lg[top]
6850 );
6851 }
6852 }
6853 k0 = k1;
6854 }
6855 pos = end;
6856 }
6857 return Ok((nll, cnt));
6858 }
6859 for pos in 0..ids.len().saturating_sub(1) {
6860 let hidden = self.forward_layers(&self.embed_single(ids[pos]), pos, task_mask);
6861 self.nll_check_graph("serial forward", pos)?;
6862 let out_of_band = self.graph_logits.take();
6870 if self.graph_head_required && out_of_band.is_none() {
6871 METAL_GRAPH_HEAD_MISS.fetch_add(
6872 1,
6873 std::sync::atomic::Ordering::Relaxed,
6874 );
6875 return Err(format!(
6876 "fused Metal graph head did not complete at NLL position {pos}"
6877 ));
6878 }
6879 if pos < start {
6880 continue;
6881 }
6882 let logits = match out_of_band {
6883 Some(lg) => lg,
6884 None => {
6885 let normed = inference::rms_norm(
6886 &hidden,
6887 &self.weights.final_norm,
6888 self.rms_eps,
6889 self.norm_style,
6890 );
6891 self.lm_head_forward(&normed)
6895 }
6896 };
6897 let target = ids[pos + 1] as usize;
6898 let max = logits.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
6899 let lse: f64 = logits
6900 .iter()
6901 .map(|&v| ((v - max) as f64).exp())
6902 .sum::<f64>()
6903 .ln()
6904 + max as f64;
6905 let tok_nll = lse - logits[target] as f64;
6906 if std::env::var("CMF_PPL_TRACE").is_ok() && pos < 48 {
6907 let top = logits
6908 .iter()
6909 .enumerate()
6910 .max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
6911 .map(|(i, _)| i)
6912 .unwrap_or(0);
6913 eprintln!(
6914 "pos {pos:3} tgt {target:6} nll {tok_nll:7.3} | top1 {top:6} lg[t]={:.2} lg[top]={:.2}",
6915 logits[target], logits[top]
6916 );
6917 }
6918 nll += tok_nll;
6919 cnt += 1;
6920 }
6921 Ok((nll, cnt))
6922 })();
6923 self.nll_end();
6924 result
6925 }
6926
6927 fn nll_from_hidden(&mut self, hidden: &[f32], target: u32, pos: usize) -> f64 {
6932 let normed = inference::rms_norm(
6933 hidden,
6934 &self.weights.final_norm,
6935 self.rms_eps,
6936 self.norm_style,
6937 );
6938 let mut logits = self.lm_head_forward(&normed);
6941 let target = target as usize;
6942 let max = logits.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
6943 let lse: f64 = logits
6944 .iter()
6945 .map(|&v| ((v - max) as f64).exp())
6946 .sum::<f64>()
6947 .ln()
6948 + max as f64;
6949 let tok_nll = lse - logits[target] as f64;
6950 if std::env::var("CMF_PPL_TRACE").is_ok() && pos < 48 {
6951 let top = logits
6952 .iter()
6953 .enumerate()
6954 .max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
6955 .map(|(i, _)| i)
6956 .unwrap_or(0);
6957 eprintln!(
6958 "pos {pos:3} tgt {target:6} nll {tok_nll:7.3} | top1 {top:6} lg[t]={:.2} lg[top]={:.2}",
6959 logits[target], logits[top]
6960 );
6961 }
6962 attention::recycle_buf(&mut logits);
6963 tok_nll
6964 }
6965
6966 pub fn nll_ids_o1(&mut self, ids: &[u32], prefill: usize) -> Result<(f64, usize), String> {
6984 self.nll_begin()?;
6989 let requested_prefix = (prefill > 0).then_some(prefill);
6990 self.o1_begin_with_prefix(requested_prefix);
6991 let n = ids.len().saturating_sub(1);
6992 let requested_start = prefill.min(n);
6993 let exact_end = if self.o1_active() {
6998 match requested_prefix {
6999 Some(requested) => self.o1_effective_boundary(requested),
7000 None => self
7001 .o1_cfg
7002 .as_ref()
7003 .and_then(|c| crate::nystrom::o1_deferred_boundary(c.w, c.sink)),
7004 }
7005 .unwrap_or(requested_start)
7006 .min(n)
7007 } else {
7008 requested_start
7009 };
7010 let mut nll = 0f64;
7011 let mut cnt = 0usize;
7012
7013 let mut pos = 0usize;
7017 if self.can_prefill_batched() {
7018 const CHUNK: usize = 128;
7019 while pos < exact_end {
7020 let end = (pos + CHUNK).min(exact_end);
7021 let hiddens = self.prefill_batch(&ids[pos..end], pos);
7022 if self
7023 .graph_failed
7024 .swap(false, std::sync::atomic::Ordering::Relaxed)
7025 {
7026 self.cancel
7027 .store(false, std::sync::atomic::Ordering::Relaxed);
7028 self.nll_end();
7029 return Err("GPU graph failed during O(1) NLL prefix".into());
7030 }
7031 for row in 0..end - pos {
7032 let score_pos = pos + row;
7033 if score_pos >= requested_start && score_pos < n {
7034 nll += self.nll_from_hidden(
7035 &hiddens[row * self.hidden_size..(row + 1) * self.hidden_size],
7036 ids[score_pos + 1],
7037 score_pos,
7038 );
7039 cnt += 1;
7040 }
7041 }
7042 pos = end;
7043 }
7044 } else {
7045 while pos < exact_end {
7046 let hidden = self.forward_layers(&self.embed_single(ids[pos]), pos, None);
7047 if self
7048 .graph_failed
7049 .swap(false, std::sync::atomic::Ordering::Relaxed)
7050 {
7051 self.cancel
7052 .store(false, std::sync::atomic::Ordering::Relaxed);
7053 self.nll_end();
7054 return Err("GPU graph failed during O(1) NLL prefix".into());
7055 }
7056 if pos >= requested_start && pos < n {
7057 nll += self.nll_from_hidden(&hidden, ids[pos + 1], pos);
7058 cnt += 1;
7059 }
7060 pos += 1;
7061 }
7062 }
7063 self.o1_seal_checked().map_err(|err| {
7064 self.nll_end();
7065 err
7066 })?;
7067
7068 let batch_k = std::env::var("CMF_BATCH_K")
7077 .ok()
7078 .and_then(|v| v.parse::<usize>().ok())
7079 .unwrap_or(0);
7080 let batch_admitted = batch_k > 0
7081 && self.can_prefill_batched()
7082 && self.o1_active()
7083 && std::env::var("CMF_O1_GPU").as_deref() == Ok("1")
7084 && (0..self.num_layers).all(|li| {
7085 let cache = &self.kv_cache.layers[self.phys_layer(li)];
7086 cache.o1.is_none() || cache.o1_views().is_some()
7087 });
7088 if std::env::var("CMF_GRAPH_PROF").is_ok() {
7089 eprintln!(
7090 "nll-batch: phase=post-seal admission={} requested_k={} scored_rows={}",
7091 batch_admitted,
7092 batch_k,
7093 n.saturating_sub(exact_end),
7094 );
7095 }
7096 let mut batch_completed = false;
7097 if batch_admitted && exact_end < n {
7098 let hs = self.hidden_size;
7099 let mut batch_pos = exact_end;
7100 while batch_pos < n {
7101 let end = (batch_pos + batch_k).min(n);
7102 let bk = end - batch_pos;
7103 let mut hiddens = vec![0.0f32; bk * hs];
7104 for (row, &id) in ids[batch_pos..end].iter().enumerate() {
7105 hiddens[row * hs..(row + 1) * hs].copy_from_slice(&self.embed_single(id));
7106 }
7107 let positions: Vec<usize> = (batch_pos..end).collect();
7108 let t_batch = std::time::Instant::now();
7109 let outcome = self.try_batch_graph_wgpu(&mut hiddens, &positions, bk, None);
7110 if std::env::var("CMF_GRAPH_PROF").is_ok() {
7111 let ms = t_batch.elapsed().as_secs_f64() * 1000.0;
7112 eprintln!(
7113 "nll-batch: phase=post-seal mode=o1 k={bk} pos={}..{} outcome={outcome:?} {ms:.1} ms ({:.1} tok/s)",
7114 batch_pos,
7115 end.saturating_sub(1),
7116 bk as f64 / (ms / 1000.0),
7117 );
7118 }
7119 if let Err(err) = self.nll_check_graph("batch graph", batch_pos) {
7120 self.nll_end();
7121 return Err(err);
7122 }
7123 match outcome {
7124 crate::gpu::BatchGraphOutcome::Completed => {
7125 batch_completed = true;
7126 for row in 0..bk {
7127 nll += self.nll_from_hidden(
7128 &hiddens[row * hs..(row + 1) * hs],
7129 ids[batch_pos + row + 1],
7130 batch_pos + row,
7131 );
7132 cnt += 1;
7133 }
7134 batch_pos = end;
7135 }
7136 crate::gpu::BatchGraphOutcome::Declined => {
7137 if batch_completed {
7138 self.nll_end();
7139 return Err(format!(
7140 "O(1) NLL batch declined after completed chunk at position {batch_pos}"
7141 ));
7142 }
7143 break;
7144 }
7145 crate::gpu::BatchGraphOutcome::Failed => {
7146 self.nll_end();
7147 return Err(format!(
7148 "O(1) NLL batch graph failed after admission at position {batch_pos}"
7149 ));
7150 }
7151 }
7152 }
7153 if batch_completed && cnt == n.saturating_sub(requested_start) {
7154 self.nll_end();
7155 return Ok((nll, cnt));
7156 }
7157 }
7158
7159 for pos in exact_end..n {
7164 let hidden = self.forward_layers(&self.embed_single(ids[pos]), pos, None);
7165 if self
7166 .graph_failed
7167 .swap(false, std::sync::atomic::Ordering::Relaxed)
7168 {
7169 self.cancel
7170 .store(false, std::sync::atomic::Ordering::Relaxed);
7171 self.nll_end();
7172 return Err(format!(
7173 "GPU graph failed during O(1) NLL serial scoring at position {pos}"
7174 ));
7175 }
7176 nll += self.nll_from_hidden(&hidden, ids[pos + 1], pos);
7177 cnt += 1;
7178 }
7179 self.nll_end();
7180 Ok((nll, cnt))
7181 }
7182
7183 pub fn calib_ids(&mut self, ids: &[u32], temps: &[f32]) -> (Vec<bool>, Vec<Vec<f32>>) {
7191 self.clear_sequence_state();
7192 let n = ids.len().saturating_sub(1);
7193 let mut correct = Vec::with_capacity(n);
7194 let mut pmax = Vec::with_capacity(n);
7195 for pos in 0..n {
7196 let emb = self.embed_single(ids[pos]);
7197 let hidden = self.forward_layers(&emb, pos, None);
7198 let normed = inference::rms_norm(
7199 &hidden,
7200 &self.weights.final_norm,
7201 self.rms_eps,
7202 self.norm_style,
7203 );
7204 let logits = self.lm_head_forward(&normed);
7208 let target = ids[pos + 1] as usize;
7209 let (mut amax, mut mval) = (0usize, f32::NEG_INFINITY);
7210 for (i, &v) in logits.iter().enumerate() {
7211 if v > mval {
7212 mval = v;
7213 amax = i;
7214 }
7215 }
7216 correct.push(amax == target);
7217 let row: Vec<f32> = temps
7218 .iter()
7219 .map(|&t| {
7220 let tt = t.max(1e-3);
7221 let s: f32 = logits.iter().map(|&v| ((v - mval) / tt).exp()).sum();
7222 1.0 / s.max(1e-12) })
7224 .collect();
7225 pmax.push(row);
7226 }
7227 self.clear_sequence_state();
7228 (correct, pmax)
7229 }
7230
7231 pub fn ppl_ids_dynamic(&mut self, ids: &[u32]) -> Result<(f64, usize), String> {
7238 if self.dyn_router.is_none() {
7239 return Ok((self.ppl_ids(ids)?, 0));
7240 }
7241 self.nll_begin()?;
7242 let saved_active = self.dyn_active;
7243 let mut router = self
7244 .dyn_router
7245 .take()
7246 .ok_or_else(|| "dynamic router disappeared before PPL scoring".to_string())?;
7247 router.reset();
7248 self.dyn_phi_seen = 0;
7249 let _ = self.set_active_skill(None);
7250
7251 let result: Result<(f64, usize), String> = (|| {
7252 let mut nll = 0f64;
7253 let mut cnt = 0usize;
7254 for pos in 0..ids.len().saturating_sub(1) {
7255 let hidden = self.forward_layers(&self.embed_single(ids[pos]), pos, None);
7256 self.nll_check_graph("dynamic serial forward", pos)?;
7257 let out_of_band = self.graph_logits.take();
7258 let mut logits = match out_of_band {
7259 Some(lg) => lg,
7260 None => {
7261 let normed = inference::rms_norm(
7262 &hidden,
7263 &self.weights.final_norm,
7264 self.rms_eps,
7265 self.norm_style,
7266 );
7267 self.lm_head_forward(&normed)
7271 }
7272 };
7273 let target = ids[pos + 1] as usize;
7274 let max = logits.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
7275 let lse: f64 = logits
7276 .iter()
7277 .map(|&v| ((v - max) as f64).exp())
7278 .sum::<f64>()
7279 .ln()
7280 + max as f64;
7281 let tok_nll = lse - logits[target] as f64;
7282 if std::env::var("CMF_PPL_TRACE").is_ok() && pos < 48 {
7283 let top = logits
7284 .iter()
7285 .enumerate()
7286 .max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
7287 .map(|(i, _)| i)
7288 .unwrap_or(0);
7289 eprintln!(
7290 "pos {pos:3} tgt {target:6} nll {tok_nll:7.3} | top1 {top:6} lg[t]={:.2} lg[top]={:.2}",
7291 logits[target], logits[top]
7292 );
7293 }
7294 nll += tok_nll;
7295 cnt += 1;
7296 attention::recycle_buf(&mut logits);
7297 let phi = self.dyn_phi_ema.clone();
7299 if let Some(new_active) = router.step(&phi, pos) {
7300 let _ = self.set_active_skill(new_active);
7301 }
7302 }
7303 Ok(((nll / cnt.max(1) as f64).exp(), router.switches.len()))
7304 })();
7305
7306 let _ = self.set_active_skill(saved_active);
7309 self.dyn_router = Some(router);
7310 self.nll_end();
7311 result
7312 }
7313
7314 pub fn probe_phi(&mut self, ids: &[u32], layer: usize) -> Vec<f32> {
7316 self.clear_sequence_state();
7317 let mut acc = vec![0f32; self.hidden_size];
7318 for (pos, &id) in ids.iter().enumerate() {
7319 let h = self.forward_layers_upto(&self.embed_single(id), pos, None, Some(layer));
7320 for (a, v) in acc.iter_mut().zip(&h) {
7321 *a += v;
7322 }
7323 }
7324 let n = ids.len().max(1) as f32;
7325 for a in acc.iter_mut() {
7326 *a /= n;
7327 }
7328 self.clear_sequence_state();
7329 acc
7330 }
7331
7332 fn prefill_batch(&mut self, ids: &[u32], start_pos: usize) -> Vec<f32> {
7338 self.prefill_batch_masked(ids, start_pos, None)
7339 }
7340
7341 fn prefill_batch_masked(
7347 &mut self,
7348 ids: &[u32],
7349 start_pos: usize,
7350 task_mask: Option<&TaskMask>,
7351 ) -> Vec<f32> {
7352 self.prefill_batch_span(PrefillIn::Ids(ids), start_pos, task_mask, 0, usize::MAX)
7353 }
7354
7355 fn prefill_batch_span(
7361 &mut self,
7362 input: PrefillIn<'_>,
7363 start_pos: usize,
7364 task_mask: Option<&TaskMask>,
7365 from: usize,
7366 upto_excl: usize,
7367 ) -> Vec<f32> {
7368 let hs = self.hidden_size;
7369 let b = match input {
7370 PrefillIn::Ids(ids) => ids.len(),
7371 PrefillIn::Hidden(hb) => hb.len() / hs,
7372 };
7373 let upto_excl = upto_excl.min(self.num_layers);
7374 let mut h: Vec<f32>;
7378 let mut h_ready;
7379 match input {
7380 PrefillIn::Ids(_) => {
7381 h = vec![0.0; b * hs];
7382 h_ready = false;
7383 }
7384 PrefillIn::Hidden(hb) => {
7385 h = hb.to_vec();
7386 h_ready = true;
7387 }
7388 }
7389 let fill_h = |h: &mut Vec<f32>, me: &Self| {
7390 if let PrefillIn::Ids(ids) = input {
7391 for (bi, &id) in ids.iter().enumerate() {
7392 let e = me.embed_single(id);
7393 h[bi * hs..(bi + 1) * hs].copy_from_slice(&e);
7394 }
7395 if let Ok(tp) = std::env::var("CMF_TRACE_POS") {
7396 if let Ok(t) = tp.parse::<usize>() {
7397 if t >= start_pos && t < start_pos + ids.len() {
7398 let bi = t - start_pos;
7399 let row = &h[bi * hs..(bi + 1) * hs];
7400 let n: f32 = row.iter().map(|x| x * x).sum::<f32>().sqrt();
7401 eprintln!(
7402 "BATCH pos {t} embed: id {} |h| = {n:.6} h0 {:.6} h1 {:.6} | b={} start={start_pos} ids[..8]={:?}",
7403 ids[bi],
7404 row[0],
7405 row[1],
7406 ids.len(),
7407 &ids[..ids.len().min(8)]
7408 );
7409 }
7410 }
7411 }
7412 }
7413 };
7414 let (_nkv, _hd, _rd, eps) = (
7415 self.num_kv_heads,
7416 self.head_dim,
7417 self.rotary_dim,
7418 self.rms_eps,
7419 );
7420 let pool = self.pool.clone();
7421 let norm_style = self.norm_style;
7422 let automatic_gpu_prefix = self.automatic_gpu_prefix();
7423
7424 #[cfg(target_os = "macos")]
7425 let mut chunk_skip_until = 0usize;
7426 for li in from..upto_excl {
7427 let _capacity_tail = automatic_gpu_prefix
7428 .filter(|&prefix| li >= prefix)
7429 .map(|_| crate::gpu::enter_cpu_scope());
7430 crate::gpu::set_layer(li as i64); #[cfg(target_os = "macos")]
7437 if task_mask.is_none() {
7438 if li < chunk_skip_until {
7439 continue;
7440 }
7441 if !h_ready && li == 0 && self.weights.embed_tokens.q8_row_parts().is_none() {
7447 fill_h(&mut h, self);
7448 h_ready = true;
7449 }
7450 let ids_for_embed = match input {
7451 PrefillIn::Ids(ids) => (!h_ready && li == 0).then_some(ids),
7452 PrefillIn::Hidden(_) => None,
7453 };
7454 let end = self.chunk_run_gpu(li, &mut h, b, start_pos, ids_for_embed, upto_excl);
7455 if end > li {
7456 h_ready = true;
7457 chunk_skip_until = end;
7458 if self.is_loop_end(end - 1) && end < self.num_layers {
7461 for bi in 0..b {
7462 let normed = inference::rms_norm(
7463 &h[bi * hs..(bi + 1) * hs],
7464 &self.weights.final_norm,
7465 eps,
7466 norm_style,
7467 );
7468 h[bi * hs..(bi + 1) * hs].copy_from_slice(&normed);
7469 }
7470 }
7471 continue;
7472 }
7473 }
7474 if !h_ready {
7475 fill_h(&mut h, self);
7476 h_ready = true;
7477 }
7478 let lw = &self.weights.layers[self.phys_layer(li)];
7479 match &lw.attn {
7481 AttnKind::Kda(w) => {
7482 let cfg = self.kda_cfg.expect("kda layer without kda_cfg");
7484 let mut normed = vec![0.0f32; b * hs];
7485 for bi in 0..b {
7486 inference::rms_norm_into(
7487 &h[bi * hs..(bi + 1) * hs],
7488 &lw.input_norm,
7489 eps,
7490 norm_style,
7491 &mut normed[bi * hs..(bi + 1) * hs],
7492 );
7493 }
7494 let attn = crate::linear_core::kda_forward_batch(
7495 &normed,
7496 b,
7497 w,
7498 &cfg,
7499 &mut self.kv_cache.layers[li].linear_state,
7500 pool.as_deref(),
7501 );
7502 for (dst, &a) in h.iter_mut().zip(&attn) {
7503 *dst += a;
7504 }
7505 }
7506 AttnKind::LinearGdn(w) => {
7507 let cfg = self.gdn_cfg.expect("gdn layer without gdn_cfg");
7509 let mut normed = vec![0.0f32; b * hs];
7510 for bi in 0..b {
7511 let r = inference::rms_norm(
7512 &h[bi * hs..(bi + 1) * hs],
7513 &lw.input_norm,
7514 eps,
7515 norm_style,
7516 );
7517 normed[bi * hs..(bi + 1) * hs].copy_from_slice(&r);
7518 }
7519 let attn = crate::linear_core::gdn_forward_batch(
7520 &normed,
7521 b,
7522 w,
7523 &cfg,
7524 &mut self.kv_cache.layers[li].linear_state,
7525 pool.as_deref(),
7526 );
7527 for (dst, &a) in h.iter_mut().zip(&attn) {
7528 *dst += a;
7529 }
7530 }
7531 AttnKind::ShortConv(w) => {
7532 let cfg = self
7535 .short_conv_cfg
7536 .expect("short-conv layer without short_conv_cfg");
7537 let mut normed = vec![0.0f32; b * hs];
7538 for bi in 0..b {
7539 inference::rms_norm_into(
7540 &h[bi * hs..(bi + 1) * hs],
7541 &lw.input_norm,
7542 eps,
7543 norm_style,
7544 &mut normed[bi * hs..(bi + 1) * hs],
7545 );
7546 }
7547 let attn = short_conv_forward_batch(
7548 &normed,
7549 b,
7550 w,
7551 &cfg,
7552 &mut self.kv_cache.layers[li].linear_state,
7553 pool.as_deref(),
7554 );
7555 for (dst, &a) in h.iter_mut().zip(&attn) {
7556 *dst += a;
7557 }
7558 }
7559 AttnKind::Mla(w) => {
7560 let inv_freq_l = self.layer_inv_freq(li);
7563 let rs = self.layer_rope_scale(li);
7564 let mut normed = vec![0.0f32; hs];
7565 for bi in 0..b {
7566 inference::rms_norm_into(
7567 &h[bi * hs..(bi + 1) * hs],
7568 &lw.input_norm,
7569 eps,
7570 norm_style,
7571 &mut normed,
7572 );
7573 let ao = mla_attention(
7574 w,
7575 &normed,
7576 &mut self.kv_cache.layers[li],
7577 start_pos + bi,
7578 &inv_freq_l,
7579 rs,
7580 eps,
7581 pool.as_deref(),
7582 );
7583 for (dst, &a) in h[bi * hs..(bi + 1) * hs].iter_mut().zip(&ao) {
7584 *dst += a;
7585 }
7586 }
7587 }
7588 AttnKind::Full {
7589 wq,
7590 wk,
7591 wv,
7592 wo,
7593 q_norm,
7594 k_norm,
7595 output_gate,
7596 softplus_gate,
7597 bias,
7598 } => {
7599 let mut normed = vec![0.0f32; b * hs];
7603 for bi in 0..b {
7604 inference::rms_norm_into(
7605 &h[bi * hs..(bi + 1) * hs],
7606 &lw.input_norm,
7607 eps,
7608 norm_style,
7609 &mut normed[bi * hs..(bi + 1) * hs],
7610 );
7611 }
7612 let inv_freq_l = self.layer_inv_freq(li);
7613 let (nkv_l, hd_l, rd_l) = self.layer_geom(li);
7614 let cfg = QwenAttnCfg {
7615 num_heads: self.layer_num_heads(li),
7616 num_kv_heads: nkv_l,
7617 head_dim: hd_l,
7618 hidden_size: hs,
7619 position: start_pos,
7620 inv_freq: &inv_freq_l,
7621 rotary_dim: rd_l,
7622 scale: self.attn_scale,
7623 softcap: self.attn_softcap,
7624 window: self.layer_window(li),
7625 v_norm: self.attn_v_norm,
7626 qk_norm_after_rope: self.qk_norm_after_rope,
7627 q_norm: q_norm.as_deref(),
7628 k_norm: k_norm.as_deref(),
7629 output_gate: *output_gate,
7630 softplus_gate: softplus_gate
7631 .as_ref()
7632 .map(|(gate, per_head)| (gate, *per_head)),
7633 rope_scale: self.layer_rope_scale(li),
7634 bias: bias
7635 .as_ref()
7636 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
7637 rms_eps: eps,
7638 norm_style,
7639 pool: pool.as_deref(),
7640 };
7641 let mut attn = attention::qwen_attention_batch(
7642 &normed,
7643 b,
7644 wq,
7645 wk,
7646 wv,
7647 wo,
7648 &mut self.kv_cache.layers[li],
7649 &cfg,
7650 );
7651 if let Some(w) = &lw.attn_out_norm {
7652 for bi in 0..b {
7653 inference::rms_norm_into(
7654 &attn[bi * hs..(bi + 1) * hs],
7655 w,
7656 eps,
7657 norm_style,
7658 &mut normed[bi * hs..(bi + 1) * hs],
7659 );
7660 }
7661 attn.copy_from_slice(&normed);
7662 }
7663 for (dst, &a) in h.iter_mut().zip(&attn) {
7664 *dst += a;
7665 }
7666 }
7667 AttnKind::Linear(w) => {
7668 for bi in 0..b {
7669 let normed = inference::rms_norm(
7670 &h[bi * hs..(bi + 1) * hs],
7671 &lw.input_norm,
7672 eps,
7673 norm_style,
7674 );
7675 vmf_phase_forward(
7676 &normed,
7677 w,
7678 &self.vmf_cfg.expect("linear layer without vmf_cfg"),
7679 &mut self.kv_cache.layers[li].linear_state,
7680 pool.as_deref(),
7681 )
7682 .iter()
7683 .enumerate()
7684 .for_each(|(i, &a)| h[bi * hs + i] += a);
7685 }
7686 }
7687 }
7688
7689 let lw = &self.weights.layers[self.phys_layer(li)];
7691 let mut post = vec![0.0f32; b * hs];
7692 for bi in 0..b {
7693 let r =
7694 inference::rms_norm(&h[bi * hs..(bi + 1) * hs], &lw.post_norm, eps, norm_style);
7695 post[bi * hs..(bi + 1) * hs].copy_from_slice(&r);
7696 }
7697 let mask_row = task_mask
7700 .filter(|m| m.ffn_active_count(li) < self.intermediate_size)
7701 .and_then(|m| m.ffn_masks.get(li))
7702 .map(|v| v.as_slice());
7703 let mut ffn = match &lw.ffn {
7704 FfnKind::Dense(d) if !d.segs.is_empty() => {
7705 tube_ffn(d, &post, b, pool.as_deref(), mask_row)
7706 }
7707 FfnKind::Dense(d) => dense_ffn_batch(d, &post, b, pool.as_deref(), mask_row),
7708 FfnKind::Moe(m) => moe_ffn_batch(m, &post, b, hs, pool.as_deref(), None),
7709 FfnKind::DenseMoe(dm) => {
7712 let mut out = vec![0.0f32; b * hs];
7713 for bi in 0..b {
7714 let r = dense_moe_ffn(
7715 dm,
7716 &post[bi * hs..(bi + 1) * hs],
7717 &h[bi * hs..(bi + 1) * hs],
7718 eps,
7719 norm_style,
7720 pool.as_deref(),
7721 );
7722 out[bi * hs..(bi + 1) * hs].copy_from_slice(&r);
7723 }
7724 out
7725 }
7726 };
7727 if let Some(w) = &lw.ffn_out_norm {
7728 for bi in 0..b {
7729 inference::rms_norm_into(
7730 &ffn[bi * hs..(bi + 1) * hs],
7731 w,
7732 eps,
7733 norm_style,
7734 &mut post[bi * hs..(bi + 1) * hs],
7735 );
7736 }
7737 ffn.copy_from_slice(&post);
7738 }
7739 for (dst, &f) in h.iter_mut().zip(&ffn) {
7740 *dst += f;
7741 }
7742 if let Some(sc) = lw.layer_scale {
7743 for v in h.iter_mut() {
7744 *v *= sc;
7745 }
7746 }
7747 if let Ok(tp) = std::env::var("CMF_TRACE_POS") {
7748 if let Ok(t) = tp.parse::<usize>() {
7749 if t >= start_pos && t < start_pos + b {
7750 let bi = t - start_pos;
7751 let row = &h[bi * hs..(bi + 1) * hs];
7752 let n: f32 = row.iter().map(|x| x * x).sum::<f32>().sqrt();
7753 eprintln!(
7754 "BATCH pos {t} after layer {li}: |h| = {n:.6} h0 {:.6} h1 {:.6}",
7755 row[0], row[1]
7756 );
7757 }
7758 }
7759 }
7760 if std::env::var("CMF_DEBUG_LAYERS").is_ok() {
7764 let row = &h[(b - 1) * hs..b * hs];
7765 let rms =
7766 (row.iter().map(|&v| (v as f64) * (v as f64)).sum::<f64>() / hs as f64).sqrt();
7767 let mx = row.iter().fold(0f32, |m, &v| m.max(v.abs()));
7768 eprintln!(
7769 "layer {li:>3} {:>10} ffn={:<5} rms={rms:>12.4} max={mx:>12.4}",
7770 match &self.weights.layers[self.phys_layer(li)].attn {
7771 AttnKind::LinearGdn(_) => "gdn",
7772 AttnKind::Linear(_) => "vmf",
7773 AttnKind::ShortConv(_) => "conv",
7774 _ => "attn",
7775 },
7776 match &lw.ffn {
7777 FfnKind::Moe(_) => "moe",
7778 FfnKind::Dense(_) => "dense",
7779 FfnKind::DenseMoe(_) => "dense+moe",
7780 },
7781 );
7782 }
7783 if self.is_loop_end(li) && li + 1 < self.num_layers {
7785 for bi in 0..b {
7786 let normed = inference::rms_norm(
7787 &h[bi * hs..(bi + 1) * hs],
7788 &self.weights.final_norm,
7789 eps,
7790 norm_style,
7791 );
7792 h[bi * hs..(bi + 1) * hs].copy_from_slice(&normed);
7793 }
7794 }
7795 if std::env::var("CMF_TRACE_H").is_ok() {
7796 let n = h[..hs].iter().map(|v| v.abs()).sum::<f32>() / hs as f32;
7797 let mx = h[..hs].iter().fold(0.0f32, |a, &v| a.max(v.abs()));
7798 eprintln!(
7799 "layer {li}: mean|h|={n:.4} max|h|={mx:.2} scale={:?}",
7800 lw.layer_scale
7801 );
7802 }
7803 }
7804 crate::gpu::set_layer(-1); self.o1_progress();
7810 h
7811 }
7812
7813 fn embed_single(&self, id: u32) -> Vec<f32> {
7815 let mut out = vec![0.0f32; self.hidden_size];
7816 if (id as usize) < self.weights.embed_tokens.rows() {
7817 self.weights.embed_tokens.row_f32(id as usize, &mut out);
7818 }
7819 if self.embed_multiplier != 1.0 {
7820 for v in out.iter_mut() {
7821 *v *= self.embed_multiplier;
7822 }
7823 }
7824 if self.dsv4.is_some() || self.dsv41.is_some() || self.qwen4_exp.is_some() {
7828 let mut v = vec![0.0f32; self.hidden_size.max(1)];
7829 v[0] = id as f32;
7830 return v;
7831 }
7832 if let Some(b) = &self.g3n {
7835 return b.0.extend_embedding(id, &out, self.pool.as_deref());
7836 }
7837 out
7838 }
7839
7840 #[cfg(target_os = "macos")]
7846 fn chunk_run_gpu(
7847 &mut self,
7848 li0: usize,
7849 h: &mut [f32],
7850 b: usize,
7851 pos0: usize,
7852 embed_ids: Option<&[u32]>,
7853 cap: usize,
7854 ) -> usize {
7855 if !crate::gpu::enabled_here()
7859 || std::env::var("CMF_GPU_CHUNK")
7860 .map(|v| v == "0")
7861 .unwrap_or(false)
7862 || b < 32
7863 || self.swa.is_some()
7864 || self.global_attn.is_some()
7865 || self.o1_active()
7868 || self.attn_v_norm
7869 || (self.attn_scale - 1.0 / (self.head_dim as f32).sqrt()).abs() > 1e-9
7870 {
7871 return li0;
7872 }
7873 let Some(model) = self.model.clone() else {
7874 return li0;
7875 };
7876 let inv_freq = self.inv_freq.clone();
7877 let (nh, nkv, hd, hs) = (
7878 self.num_heads,
7879 self.num_kv_heads,
7880 self.head_dim,
7881 self.hidden_size,
7882 );
7883 let loop_end = if self.loop_final_norm {
7887 ((li0 / self.physical_layers) + 1) * self.physical_layers
7888 } else {
7889 self.num_layers
7890 };
7891 let mut layers: Vec<crate::gpu_metal::ChunkLayer> = Vec::new();
7892 let mut stored_at: Vec<usize> = Vec::new();
7893 for li in li0..self.num_layers.min(loop_end).min(cap) {
7894 let lw = &self.weights.layers[self.phys_layer(li)];
7895 if lw.attn_out_norm.is_some() || lw.ffn_out_norm.is_some() || lw.layer_scale.is_some() {
7896 break;
7897 }
7898 let AttnKind::Full {
7899 wq,
7900 wk,
7901 wv,
7902 wo,
7903 q_norm,
7904 k_norm,
7905 output_gate: false,
7906 softplus_gate: None,
7907 bias,
7908 } = &lw.attn
7909 else {
7910 break;
7911 };
7912 let FfnKind::Dense(d) = &lw.ffn else { break };
7913 if d.act != Act::Silu || !d.segs.is_empty() {
7914 break;
7915 }
7916 fn cw(t: &QTensor) -> Option<(usize, usize, usize, &[f32])> {
7921 t.q8_row_parts()
7922 .or_else(|| t.q4t_parts().map(|(i, r, c)| (i, r, c, &[][..])))
7923 .or_else(|| t.q4tp_parts().map(|(i, r, c)| (i, r, c, &[][..])))
7924 }
7925 let parts = (
7926 cw(wq),
7927 cw(wk),
7928 cw(wv),
7929 cw(wo),
7930 cw(&d.gate_proj),
7931 cw(&d.up_proj),
7932 cw(&d.down_proj),
7933 );
7934 let (Some(pq), Some(pk), Some(pv), Some(po), Some(pg), Some(pu), Some(pd)) = parts
7935 else {
7936 break;
7937 };
7938 let layer = &self.kv_cache.layers[li];
7939 if layer.mode != crate::kv_cache::KvMode::F32 || layer.o1.is_some() {
7940 break;
7941 }
7942 stored_at.push(layer.head_len(0));
7943 layers.push(crate::gpu_metal::ChunkLayer {
7944 model: &model,
7945 kv_id: self.graph_kv_id,
7946 layer: li,
7947 wq: pq,
7948 wk: pk,
7949 wv: pv,
7950 wo: po,
7951 gate: pg,
7952 up: pu,
7953 down: pd,
7954 input_norm: &lw.input_norm,
7955 post_norm: &lw.post_norm,
7956 bias: bias
7957 .as_ref()
7958 .map(|(a, bb, cc)| (a.as_slice(), bb.as_slice(), cc.as_slice())),
7959 q_norm: q_norm.as_deref(),
7960 k_norm: k_norm.as_deref(),
7961 inv_freq: &inv_freq,
7962 rd: self.rotary_dim,
7963 nh,
7964 nkv,
7965 hd,
7966 hs,
7967 inter: d.gate_proj.rows(),
7968 gemma: matches!(self.norm_style, cortiq_core::NormStyle::Gemma),
7969 late_qk_norm: self.qk_norm_after_rope,
7970 eps: self.rms_eps as f32,
7971 });
7972 }
7973 if layers.is_empty() {
7974 return li0;
7975 }
7976 let row = nkv * hd;
7977 let mut store: Vec<(Vec<f32>, Vec<f32>, Vec<f32>)> = stored_at
7978 .iter()
7979 .map(|&st| (vec![0f32; b * row], vec![0f32; b * row], vec![0f32; st + b]))
7980 .collect();
7981 let mut io: Vec<crate::gpu_metal::ChunkIo> = Vec::with_capacity(layers.len());
7982 for (i, (ok, ov, oi)) in store.iter_mut().enumerate() {
7983 let li = layers[i].layer;
7984 let layer = &self.kv_cache.layers[li];
7985 io.push(crate::gpu_metal::ChunkIo {
7986 cpu_stored: stored_at[i],
7987 cpu_k: (0..nkv).map(|g| layer.head_keys(g)).collect(),
7988 cpu_v: (0..nkv).map(|g| layer.head_values(g)).collect(),
7989 out_k: ok,
7990 out_v: ov,
7991 imp: oi,
7992 });
7993 }
7994 let n_run = layers.len();
7995 let last = layers.last().map(|l| l.layer + 1).unwrap_or(li0);
7996 let ep = embed_ids.and_then(|ids| {
7999 self.weights
8000 .embed_tokens
8001 .q8_row_parts()
8002 .map(|(idx, rows, _c, rs)| crate::gpu_metal::ChunkEmbed {
8003 idx,
8004 rows,
8005 row_scale: rs,
8006 ids,
8007 mult: self.embed_multiplier,
8008 })
8009 });
8010 if embed_ids.is_some() && ep.is_none() {
8011 return li0;
8012 }
8013 if !crate::gpu_metal::chunk_run_gpu(&layers, &mut io, h, b, pos0, ep.as_ref()) {
8014 return li0;
8015 }
8016 drop(io);
8017 drop(layers);
8018 for (i, (ok, ov, oi)) in store.iter().enumerate().take(n_run) {
8021 let li = li0 + i;
8022 let layer = &mut self.kv_cache.layers[li];
8023 for bi in 0..b {
8024 layer.append(
8025 &ok[bi * row..(bi + 1) * row],
8026 &ov[bi * row..(bi + 1) * row],
8027 &[],
8028 );
8029 }
8030 layer.accumulate_imp(oi);
8031 }
8032 last
8033 }
8034
8035 fn layer_is_local(&self, li: usize) -> bool {
8038 if let Some(layers) = &self.sliding_layers {
8039 return layers.get(li).copied().unwrap_or(false);
8040 }
8041 match self.swa {
8042 Some((_, pattern)) => (li + 1) % pattern.max(1) != 0,
8043 None => false,
8044 }
8045 }
8046
8047 fn layer_inv_freq(&self, li: usize) -> std::sync::Arc<Vec<f32>> {
8050 if self.layer_is_local(li) {
8051 if let Some(f) = &self.inv_freq_local {
8052 return f.clone();
8053 }
8054 } else if let Some(f) = &self.inv_freq_global {
8055 return f.clone();
8056 }
8057 self.inv_freq.clone()
8058 }
8059
8060 fn layer_window(&self, li: usize) -> Option<usize> {
8062 self.swa
8063 .and_then(|(w, _)| self.layer_is_local(li).then_some(w))
8064 }
8065
8066 fn layer_num_heads(&self, li: usize) -> usize {
8067 self.attention_heads_per_layer
8068 .as_ref()
8069 .and_then(|v| v.get(li).copied())
8070 .unwrap_or(self.num_heads)
8071 }
8072
8073 fn layer_rope_scale(&self, li: usize) -> f32 {
8074 if self.layer_is_local(li) {
8075 self.rope_scale_local
8076 } else {
8077 self.rope_scale
8078 }
8079 }
8080
8081 fn layer_geom(&self, li: usize) -> (usize, usize, usize) {
8084 if !self.layer_is_local(li) {
8085 if let Some((ghd, gkv)) = self.global_attn {
8086 return (gkv, ghd, ghd);
8087 }
8088 }
8089 (
8090 self.num_kv_heads,
8091 self.head_dim,
8092 if self.layer_is_local(li) {
8093 self.rotary_dim_local.unwrap_or(self.rotary_dim)
8094 } else {
8095 self.rotary_dim
8096 },
8097 )
8098 }
8099
8100 fn forward_layers(
8102 &mut self,
8103 hidden: &[f32],
8104 position: usize,
8105 task_mask: Option<&TaskMask>,
8106 ) -> Vec<f32> {
8107 let out = self.forward_layers_upto(hidden, position, task_mask, None);
8108 self.o1_progress();
8109 out
8110 }
8111
8112 pub fn embed_id(&self, id: u32) -> Vec<f32> {
8120 self.embed_single(id)
8121 }
8122
8123 pub fn split_supported(&self) -> Result<(), String> {
8127 if self.dsv4.is_some() {
8128 return Err(
8129 "network split: DeepSeek-V4 runs its own fused stack (not splittable yet)".into(),
8130 );
8131 }
8132 if self.dsv41.is_some() {
8133 return Err(
8134 "network split: DeepSeek-V4.1 owns the shared CED/CSA2 state (not splittable)"
8135 .into(),
8136 );
8137 }
8138 if self.qwen4_exp.is_some() {
8139 return Err(
8140 "network split: Qwen3.8-Flash-Next hyper/QSA stack is not splittable yet".into(),
8141 );
8142 }
8143 if self.g3n.is_some() {
8144 return Err(
8145 "network split: Gemma-3n runs its own AltUp stack (not splittable yet)".into(),
8146 );
8147 }
8148 Ok(())
8149 }
8150
8151 pub fn forward_span(
8156 &mut self,
8157 hidden: &[f32],
8158 position: usize,
8159 from: usize,
8160 upto: usize,
8161 task_mask: Option<&TaskMask>,
8162 ) -> Result<Vec<f32>, String> {
8163 self.split_supported()?;
8164 if from > upto || upto >= self.num_layers {
8165 return Err(format!(
8166 "forward_span: layer range {from}..={upto} outside 0..{}",
8167 self.num_layers
8168 ));
8169 }
8170 if hidden.len() != self.hidden_size {
8171 return Err(format!(
8172 "forward_span: hidden len {} ≠ hidden_size {}",
8173 hidden.len(),
8174 self.hidden_size
8175 ));
8176 }
8177 let out = self.forward_layers_span(hidden, position, task_mask, from, Some(upto));
8178 self.o1_progress();
8179 if self
8180 .graph_failed
8181 .swap(false, std::sync::atomic::Ordering::Relaxed)
8182 {
8183 self.cancel
8184 .store(false, std::sync::atomic::Ordering::Relaxed);
8185 self.clear_sequence_state();
8186 return Err("forward_span: deferred O(1) transition failed".into());
8187 }
8188 Ok(out)
8189 }
8190
8191 pub fn logits_from_hidden(&mut self, hidden: &[f32]) -> Vec<f32> {
8194 let normed = inference::rms_norm(
8195 hidden,
8196 &self.weights.final_norm,
8197 self.rms_eps,
8198 self.norm_style,
8199 );
8200 self.lm_head_forward(&normed)
8201 }
8202
8203 pub fn sample_next(&mut self, logits: &[f32], past_tokens: &[u32]) -> u32 {
8205 sampler::sample_with_scratch(
8206 logits,
8207 &self.sampler_config,
8208 past_tokens,
8209 &mut self.rng,
8210 &mut self.sampler_scratch,
8211 )
8212 }
8213
8214 pub fn reset_session(&mut self) {
8216 self.clear_sequence_state();
8217 }
8218
8219 pub fn prefill_span_ids(
8225 &mut self,
8226 ids: &[u32],
8227 start_pos: usize,
8228 upto: usize,
8229 task_mask: Option<&TaskMask>,
8230 ) -> Result<Vec<f32>, String> {
8231 self.split_supported()?;
8232 if upto >= self.num_layers {
8233 return Err(format!(
8234 "prefill_span_ids: upto {upto} outside 0..{}",
8235 self.num_layers
8236 ));
8237 }
8238 if self.can_prefill_batched() && !self.graph_prefill_preferred() {
8242 let out =
8243 self.prefill_batch_span(PrefillIn::Ids(ids), start_pos, task_mask, 0, upto + 1);
8244 self.check_o1_progress_failure("prefill_span_ids")?;
8245 Ok(out)
8246 } else {
8247 let hs = self.hidden_size;
8248 let mut out = Vec::with_capacity(ids.len() * hs);
8249 for (i, &id) in ids.iter().enumerate() {
8250 let emb = self.embed_id(id);
8251 out.extend_from_slice(&self.forward_span(
8252 &emb,
8253 start_pos + i,
8254 0,
8255 upto,
8256 task_mask,
8257 )?);
8258 }
8259 Ok(out)
8260 }
8261 }
8262
8263 pub fn prefill_span_hidden(
8266 &mut self,
8267 hidden: &[f32],
8268 start_pos: usize,
8269 from: usize,
8270 upto: usize,
8271 task_mask: Option<&TaskMask>,
8272 ) -> Result<Vec<f32>, String> {
8273 self.split_supported()?;
8274 let hs = self.hidden_size;
8275 if hidden.is_empty() || hidden.len() % hs != 0 {
8276 return Err(format!(
8277 "prefill_span_hidden: {} floats is not a multiple of hidden {hs}",
8278 hidden.len()
8279 ));
8280 }
8281 if from > upto || upto >= self.num_layers {
8282 return Err(format!(
8283 "prefill_span_hidden: layer range {from}..={upto} outside 0..{}",
8284 self.num_layers
8285 ));
8286 }
8287 if self.can_prefill_batched() && !self.graph_prefill_preferred() {
8288 let out = self.prefill_batch_span(
8289 PrefillIn::Hidden(hidden),
8290 start_pos,
8291 task_mask,
8292 from,
8293 upto + 1,
8294 );
8295 self.check_o1_progress_failure("prefill_span_hidden")?;
8296 Ok(out)
8297 } else {
8298 let b = hidden.len() / hs;
8299 let mut out = Vec::with_capacity(hidden.len());
8300 for i in 0..b {
8301 let h = self.forward_span(
8302 &hidden[i * hs..(i + 1) * hs],
8303 start_pos + i,
8304 from,
8305 upto,
8306 task_mask,
8307 )?;
8308 out.extend_from_slice(&h);
8309 }
8310 Ok(out)
8311 }
8312 }
8313
8314 fn try_token_graph_wgpu(
8318 &self,
8319 hidden: &[f32],
8320 position: usize,
8321 logits_out: &mut Vec<f32>,
8322 layers_run: &mut usize,
8323 ) -> Option<Result<Vec<f32>, ()>> {
8324 self.try_token_graph_wgpu_steps(
8325 hidden,
8326 position,
8327 logits_out,
8328 1,
8329 None,
8330 Some(layers_run),
8331 0,
8332 self.num_layers,
8333 )
8334 }
8335
8336 fn try_token_graph_wgpu_span(
8340 &self,
8341 hidden: &[f32],
8342 position: usize,
8343 logits_out: &mut Vec<f32>,
8344 from: usize,
8345 upto_excl: usize,
8346 layers_run: &mut usize,
8347 ) -> Option<Result<Vec<f32>, ()>> {
8348 self.try_token_graph_wgpu_steps(
8349 hidden,
8350 position,
8351 logits_out,
8352 1,
8353 None,
8354 Some(layers_run),
8355 from,
8356 upto_excl,
8357 )
8358 }
8359
8360 fn try_multi_burst(&self, t_next: u32, position: usize, k: usize) -> Option<Vec<u32>> {
8364 if self.o1_active() || self.attn_softcap > 0.0 {
8365 return None;
8366 }
8367 let graph_on = crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Decode);
8368 if !graph_on || crate::gpu::graph_unsupported() {
8369 return None;
8376 }
8377 let emb = self.embed_single(t_next);
8378 let mut lg = Vec::new();
8379 let mut ids = Vec::new();
8380 match self.try_token_graph_wgpu_steps(
8381 &emb,
8382 position,
8383 &mut lg,
8384 k,
8385 Some(&mut ids),
8386 None,
8387 0,
8388 self.num_layers,
8389 ) {
8390 Some(Ok(_)) => {}
8391 Some(Err(())) => {
8392 self.graph_failed
8397 .store(true, std::sync::atomic::Ordering::Relaxed);
8398 return None;
8399 }
8400 None => return None,
8401 }
8402 (ids.len() == k).then_some(ids)
8403 }
8404
8405 fn try_token_graph_wgpu_steps(
8409 &self,
8410 hidden: &[f32],
8411 position: usize,
8412 logits_out: &mut Vec<f32>,
8413 steps: usize,
8414 ids_out: Option<&mut Vec<u32>>,
8415 layers_run: Option<&mut usize>,
8416 from: usize,
8417 upto_excl: usize,
8418 ) -> Option<Result<Vec<f32>, ()>> {
8419 let o1_gpu = std::env::var("CMF_O1_GPU").as_deref() == Ok("1");
8422 if (self.o1_active() && !o1_gpu) || self.attn_softcap > 0.0 {
8423 return None;
8427 }
8428 let o1_views: Vec<Option<Vec<crate::nystrom::O1DeviceView<'_>>>> = (from..upto_excl)
8433 .map(|li| {
8434 if !o1_gpu {
8435 return None;
8436 }
8437 self.kv_cache.layers[self.phys_layer(li)].o1_views()
8438 })
8439 .collect();
8440 if self.o1_active() && o1_gpu {
8441 let want: usize = (from..upto_excl)
8444 .filter(|li| self.kv_cache.layers[self.phys_layer(*li)].o1.is_some())
8445 .count();
8446 let have = o1_views.iter().filter(|v| v.is_some()).count();
8447 if want == 0 || have != want {
8448 use std::sync::atomic::{AtomicUsize, Ordering};
8458 static LAST: AtomicUsize = AtomicUsize::new(usize::MAX);
8459 let code = have * 1000 + want;
8460 if LAST.swap(code, Ordering::Relaxed) != code {
8461 tracing::warn!(
8462 "o1 graph: {have} of {want} layers sealed — per-op until all seal"
8463 );
8464 }
8465 return None;
8466 }
8467 }
8468 let nh = self.num_heads;
8469 let (nkv, hd, rd) = self.layer_geom(0);
8470 let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
8471 let mut layers = Vec::with_capacity(upto_excl - from);
8472 let mut model = None;
8473 let dbg = std::env::var("CMF_GRAPH_DEBUG").is_ok();
8474 fn gw(t: &QTensor) -> Option<crate::gpu::GraphW<'_>> {
8475 if let Some((m, i, kind, rs)) = t
8476 .graph_weight()
8477 .or_else(|| t.graph_weight_descriptor())
8478 {
8479 let name = &m.tensors[i].name;
8480 let prism = if crate::prism::is_inverse_embedding(m, name) {
8481 crate::gpu::GraphPrismOp::InverseEmbedding
8482 } else if crate::prism::is_forward_weight(m, name) {
8483 crate::gpu::GraphPrismOp::Forward
8484 } else {
8485 crate::gpu::GraphPrismOp::None
8486 };
8487 return Some(crate::gpu::GraphW {
8488 idx: i,
8489 kind,
8490 row_scale: rs,
8491 data: &[],
8492 prism,
8493 affine: crate::prism::is_affine_target(m, name),
8494 });
8495 }
8496 match t.as_f32() {
8498 Some(d) => Some(crate::gpu::GraphW {
8499 idx: 0,
8500 kind: 4,
8501 row_scale: &[],
8502 data: d,
8503 prism: crate::gpu::GraphPrismOp::None,
8504 affine: false,
8505 }),
8506 None => {
8507 if std::env::var_os("CMF_BATCH_DEBUG").is_some() {
8508 eprintln!("batch graph: weight has no graph/f32 representation");
8509 }
8510 None
8511 }
8512 }
8513 }
8514 for li in from..upto_excl {
8515 let lw = &self.weights.layers[self.phys_layer(li)];
8516 if dbg {
8517 let ak = match &lw.attn {
8518 AttnKind::Mla(_) => "Mla".into(),
8519 AttnKind::Full {
8520 output_gate, bias, ..
8521 } => format!("Full gate={output_gate} bias={}", bias.is_some()),
8522 AttnKind::LinearGdn(_) => "LinearGdn".into(),
8523 AttnKind::Kda(_) => "Kda".into(),
8524 AttnKind::Linear(_) => "Linear".into(),
8525 AttnKind::ShortConv(_) => "ShortConv".into(),
8526 };
8527 let fk = match &lw.ffn {
8528 FfnKind::Dense(_) => "Dense",
8529 FfnKind::Moe(_) => "Moe",
8530 FfnKind::DenseMoe(_) => "DenseMoe",
8531 };
8532 eprintln!("graph L{li}: attn={ak} ffn={fk}");
8533 }
8534 let gffn = match &lw.ffn {
8535 FfnKind::DenseMoe(_) => return None, FfnKind::Dense(d) if !d.segs.is_empty() => return None,
8539 FfnKind::Dense(d) => crate::gpu::GraphFfn::Dense {
8540 gate: gw(&d.gate_proj)?,
8541 up: gw(&d.up_proj)?,
8542 down: gw(&d.down_proj)?,
8543 },
8544 FfnKind::Moe(m) => {
8545 if m.route_tau.is_some() || m.mask.is_some() {
8553 return None;
8554 }
8555 let shared = m.shared.as_ref();
8556 let has_shared = shared.is_some();
8557 let shared_gated = matches!(shared, Some((_, Some(_))));
8558 let sgate = match shared {
8559 Some((_, Some(sg))) => gw(sg)?,
8560 _ => gw(&m.router)?,
8564 };
8565 let router = gw(&m.router)?;
8566 if router.prism != crate::gpu::GraphPrismOp::None
8572 || sgate.prism != crate::gpu::GraphPrismOp::None
8573 || router.affine
8574 || sgate.affine
8575 {
8576 tracing::warn!(
8577 "resident MoE declined: Prism/affine router or shared gate transform is not implemented"
8578 );
8579 return None;
8580 }
8581 let inter = m.experts.first()?.gate_proj.rows();
8582 let mut experts = Vec::with_capacity(m.experts.len() + 1);
8583 let mut q4tp: Option<bool> = None;
8586 let mut gu_q2: Option<bool> = None;
8589 for e in m.experts.iter().chain(shared.map(|(se, _)| se)) {
8590 if !matches!(e.act, Act::Silu)
8591 || e.gate_proj.rows() != inter
8592 || e.up_proj.rows() != inter
8593 {
8594 return None;
8595 }
8596 for expert_weight in [&e.gate_proj, &e.up_proj, &e.down_proj] {
8601 let Some((em, ei, _, _)) = expert_weight
8602 .graph_weight()
8603 .or_else(|| expert_weight.graph_weight_descriptor())
8604 else {
8605 return None;
8606 };
8607 let name = &em.tensors[ei].name;
8608 if crate::prism::is_forward_weight(em, name)
8609 || crate::prism::is_inverse_embedding(em, name)
8610 || crate::prism::is_affine_target(em, name)
8611 {
8612 tracing::warn!(
8613 "resident MoE declined: expert Prism/affine transform is not implemented"
8614 );
8615 return None;
8616 }
8617 }
8618 let (mm, gi, ui, di, is_p, is_q2) = match e.gate_proj.mapped_q4t() {
8619 Some((mm, gi)) => (
8620 mm,
8621 gi,
8622 e.up_proj.mapped_q4t()?.1,
8623 e.down_proj.mapped_q4t()?.1,
8624 false,
8625 false,
8626 ),
8627 None => match e.gate_proj.mapped_q2tp() {
8628 Some((mm, gi)) => (
8629 mm,
8630 gi,
8631 e.up_proj.mapped_q2tp()?.1,
8632 e.down_proj.mapped_q4tp()?.1,
8633 true,
8634 true,
8635 ),
8636 None => {
8637 let (mm, gi) = e.gate_proj.mapped_q4tp()?;
8638 (
8639 mm,
8640 gi,
8641 e.up_proj.mapped_q4tp()?.1,
8642 e.down_proj.mapped_q4tp()?.1,
8643 true,
8644 false,
8645 )
8646 }
8647 },
8648 };
8649 if *q4tp.get_or_insert(is_p) != is_p || *gu_q2.get_or_insert(is_q2) != is_q2
8650 {
8651 tracing::warn!(
8657 "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."
8658 );
8659 return None;
8660 }
8661 model.get_or_insert_with(|| mm.clone());
8662 experts.push((gi, ui, di));
8663 }
8664 crate::gpu::GraphFfn::Moe {
8665 router,
8666 shared_gate: sgate,
8667 experts,
8668 n_exp: m.experts.len(),
8669 top_k: std::env::var("CMF_TOPK_PROBE")
8675 .ok()
8676 .and_then(|v| v.parse::<usize>().ok())
8677 .filter(|k| *k > 0 && *k <= m.top_k)
8678 .unwrap_or(m.top_k),
8679 inter,
8680 norm_topk: m.norm_topk_prob,
8681 q4tp: q4tp?,
8682 gu_q2: gu_q2.unwrap_or(false),
8683 sigmoid: m.router_sigmoid,
8684 bias: m.expert_bias.as_deref(),
8685 has_shared,
8686 shared_gated,
8687 route_scale: m.routed_scaling,
8688 }
8689 }
8690 };
8691 let attn = match &lw.attn {
8692 AttnKind::Full {
8693 wq,
8694 wk,
8695 wv,
8696 wo,
8697 q_norm,
8698 k_norm,
8699 output_gate,
8700 softplus_gate,
8701 bias,
8702 } => {
8703 if softplus_gate.is_some() || self.attention_heads_per_layer.is_some() {
8704 return None;
8705 }
8706 let (m, _, _, _) = wq
8707 .graph_weight()
8708 .or_else(|| wq.graph_weight_descriptor())?;
8709 model = Some(m.clone());
8710 crate::gpu::GraphAttn::Full {
8711 wq: gw(wq)?,
8712 wk: gw(wk)?,
8713 wv: gw(wv)?,
8714 wo: gw(wo)?,
8715 q_norm: q_norm.as_deref(),
8716 k_norm: k_norm.as_deref(),
8717 late_qk_norm: self.qk_norm_after_rope,
8718 bias: bias
8719 .as_ref()
8720 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
8721 output_gate: *output_gate,
8722 cpu_k: self.kv_cache.layers[li].k_heads(),
8723 cpu_v: self.kv_cache.layers[li].v_heads(),
8724 }
8725 }
8726 AttnKind::LinearGdn(w) => {
8727 let cfg = self.gdn_cfg?;
8728 let (m, _, _, _) = w
8729 .in_proj_qkv
8730 .graph_weight()
8731 .or_else(|| w.in_proj_qkv.graph_weight_descriptor())?;
8732 model = Some(m.clone());
8733 crate::gpu::GraphAttn::Gdn {
8734 qkv: gw(&w.in_proj_qkv)?,
8735 z: gw(&w.in_proj_z)?,
8736 a: gw(&w.in_proj_a)?,
8737 b: gw(&w.in_proj_b)?,
8738 out: gw(&w.out_proj)?,
8739 conv1d: &w.conv1d,
8740 a_log: &w.a_log,
8741 dt_bias: &w.dt_bias,
8742 norm: &w.norm,
8743 nv: cfg.num_v_heads,
8744 nk: cfg.num_k_heads,
8745 dk: cfg.key_head_dim,
8746 dv: cfg.value_head_dim,
8747 kk: cfg.conv_kernel,
8748 cpu_state: &self.kv_cache.layers[self.phys_layer(li)].linear_state,
8749 }
8750 }
8751 AttnKind::ShortConv(w) => {
8752 let cfg = self.short_conv_cfg?;
8753 let (m, _, _, _) = w
8754 .in_proj
8755 .graph_weight()
8756 .or_else(|| w.in_proj.graph_weight_descriptor())?;
8757 model = Some(m.clone());
8758 crate::gpu::GraphAttn::ShortConv {
8759 inp: gw(&w.in_proj)?,
8760 out: gw(&w.out_proj)?,
8761 taps: &w.conv,
8762 kernel: cfg.kernel,
8763 cpu_state: &self.kv_cache.layers[self.phys_layer(li)].linear_state,
8764 }
8765 }
8766 _ => return None,
8767 };
8768 layers.push(crate::gpu::GraphLayer {
8769 input_norm: &lw.input_norm,
8770 attn,
8771 post_norm: &lw.post_norm,
8772 ffn: gffn,
8773 });
8774 }
8775 let model = model?;
8776 let lm_gw = if upto_excl == self.num_layers
8782 && self.graph_want_logits
8783 && std::env::var("CMF_GPU_LMHEAD")
8784 .map(|v| v != "0")
8785 .unwrap_or(true)
8786 {
8787 self.weights
8788 .lm_head
8789 .graph_weight()
8790 .or_else(|| self.weights.lm_head.graph_weight_descriptor())
8791 .map(|(m, i, kind, rs)| {
8792 let name = &m.tensors[i].name;
8793 let prism = if crate::prism::is_inverse_embedding(m, name) {
8794 crate::gpu::GraphPrismOp::InverseEmbedding
8795 } else if crate::prism::is_forward_weight(m, name) {
8796 crate::gpu::GraphPrismOp::Forward
8797 } else {
8798 crate::gpu::GraphPrismOp::None
8799 };
8800 (
8801 crate::gpu::GraphW {
8802 idx: i,
8803 kind,
8804 row_scale: rs,
8805 data: &[],
8806 prism,
8807 affine: crate::prism::is_affine_target(m, name),
8808 },
8809 self.weights.lm_head.rows(),
8810 )
8811 })
8812 } else {
8813 None
8814 };
8815 let lm = lm_gw.as_ref().map(|(gw, rows)| (gw, *rows));
8816 let emb_gw = if steps > 1 {
8818 self.weights
8819 .embed_tokens
8820 .graph_weight()
8821 .or_else(|| self.weights.embed_tokens.graph_weight_descriptor())
8822 .map(|(m, i, kind, rs)| {
8823 let name = &m.tensors[i].name;
8824 let prism = if crate::prism::is_inverse_embedding(m, name) {
8825 crate::gpu::GraphPrismOp::InverseEmbedding
8826 } else if crate::prism::is_forward_weight(m, name) {
8827 crate::gpu::GraphPrismOp::Forward
8828 } else {
8829 crate::gpu::GraphPrismOp::None
8830 };
8831 (
8832 crate::gpu::GraphW {
8833 idx: i,
8834 kind,
8835 row_scale: rs,
8836 data: &[],
8837 prism,
8838 affine: crate::prism::is_affine_target(m, name),
8839 },
8840 self.weights.embed_tokens.rows(),
8841 self.embed_multiplier,
8842 )
8843 })
8844 } else {
8845 None
8846 };
8847
8848 let loop_norm_at: Vec<usize> = if self.loop_final_norm {
8854 (from..upto_excl.min(self.num_layers - 1))
8855 .filter(|&li| (li + 1) % self.physical_layers == 0)
8856 .map(|li| li - from)
8857 .collect()
8858 } else {
8859 Vec::new()
8860 };
8861 let mut h = hidden.to_vec();
8862 let dump_hidden = std::env::var_os("CMF_LOGIT_DUMP").is_some();
8868 let outcome = crate::gpu::forward_token_graph(
8869 &model,
8870 self.graph_kv_id,
8871 &layers,
8872 &o1_views,
8873 self.o1_epoch,
8874 &self.inv_freq,
8875 &mut h,
8876 nh,
8877 nkv,
8878 hd,
8879 self.attn_scale,
8880 rd,
8881 self.hidden_size,
8882 self.intermediate_size,
8883 position,
8884 self.kv_cache.max_seq_len,
8885 gemma,
8886 self.rms_eps as f32,
8887 lm,
8888 &self.weights.final_norm,
8889 logits_out,
8890 &loop_norm_at,
8891 steps,
8892 emb_gw.as_ref().map(|(gw, rows, m)| (gw, *rows, *m)),
8893 ids_out,
8894 layers_run,
8895 from,
8896 dump_hidden,
8897 );
8898 match outcome {
8899 crate::gpu::TokenGraphOutcome::Completed => Some(Ok(h)),
8900 crate::gpu::TokenGraphOutcome::Failed => Some(Err(())),
8901 crate::gpu::TokenGraphOutcome::Declined => None,
8902 }
8903 }
8904
8905 #[cfg(target_os = "macos")]
8914 #[allow(clippy::type_complexity)]
8915 fn metal_rows_plan(
8916 &self,
8917 ) -> Option<(
8918 Vec<MetalRowsItem<'_>>,
8919 std::sync::Arc<cortiq_core::CmfModel>,
8920 Option<crate::gpu_metal::GdnGpuCfg>,
8921 )> {
8922 use crate::gpu_metal::{AttnGpuLayer, GdnGpuCfg, GdnGpuLayer, MetalFfn};
8923 let graph_force = crate::gpu::q1_force() || crate::gpu::q2tp_gpu_opt_in();
8924 if !graph_force
8925 || !crate::gpu::enabled_here()
8926 || std::env::var("CMF_GPU_BLOCK")
8927 .map(|v| v == "0")
8928 .unwrap_or(false)
8929 || self.attn_softcap > 0.0
8930 || self.o1_active()
8931 || self.swa.is_some()
8932 || self.global_attn.is_some()
8933 || self.attention_heads_per_layer.is_some()
8934 || self.attn_v_norm
8935 || self.loop_final_norm
8936 {
8937 return None;
8938 }
8939 let attend_contract = self.head_dim % 4 == 0
8940 && self.head_dim <= 256
8941 && self.rotary_dim >= 2
8942 && self.rotary_dim <= self.head_dim
8943 && (self.rotary_dim / 2) % 32 == 0
8944 && self.num_kv_heads > 0
8945 && self.num_heads % self.num_kv_heads == 0;
8946 if !attend_contract {
8947 return None;
8948 }
8949 let mut plan: Vec<MetalRowsItem> = Vec::new();
8950 let mut model_ref: Option<std::sync::Arc<cortiq_core::CmfModel>> = None;
8951 for li in 0..self.num_layers {
8952 let lw = &self.weights.layers[self.phys_layer(li)];
8953 if lw.attn_out_norm.is_some() || lw.ffn_out_norm.is_some() || lw.layer_scale.is_some() {
8954 return None;
8955 }
8956 let ffn = match &lw.ffn {
8957 FfnKind::Dense(d) if d.act == Act::Silu && d.segs.is_empty() => {
8958 let (Some(g), Some(u), Some(dn)) = (
8959 d.gate_proj.metal_graph_parts(),
8960 d.up_proj.metal_graph_parts(),
8961 d.down_proj.metal_graph_parts(),
8962 ) else {
8963 return None;
8964 };
8965 MetalFfn::Dense {
8966 gate: g,
8967 up: u,
8968 down: dn,
8969 }
8970 }
8971 _ => return None,
8972 };
8973 match &lw.attn {
8974 AttnKind::LinearGdn(w) if self.gdn_cfg.is_some() => {
8975 let (Some(qkv), Some(z), Some(a), Some(bb), Some(out)) = (
8976 w.in_proj_qkv.metal_graph_parts(),
8977 w.in_proj_z.metal_graph_parts(),
8978 w.in_proj_a.f32_parts(),
8979 w.in_proj_b.f32_parts(),
8980 w.out_proj.metal_graph_parts(),
8981 ) else {
8982 return None;
8983 };
8984 if let QTensor::Mapped { model, .. } = &w.in_proj_qkv {
8985 model_ref.get_or_insert_with(|| model.clone());
8986 }
8987 let gl = GdnGpuLayer {
8988 attn_norm: &lw.input_norm,
8989 post_norm: &lw.post_norm,
8990 qkv,
8991 z,
8992 a,
8993 b: bb,
8994 out,
8995 ffn,
8996 conv1d: &w.conv1d,
8997 a_log: &w.a_log,
8998 dt_bias: &w.dt_bias,
8999 gnorm: &w.norm,
9000 };
9001 match plan.last_mut() {
9002 Some(MetalRowsItem::Gdn { run, .. }) => run.push(gl),
9003 _ => plan.push(MetalRowsItem::Gdn {
9004 run: vec![gl],
9005 first: li,
9006 }),
9007 }
9008 }
9009 AttnKind::Full {
9010 wq,
9011 wk,
9012 wv,
9013 wo,
9014 q_norm,
9015 k_norm,
9016 output_gate,
9017 softplus_gate: None,
9018 bias: None,
9019 } => {
9020 let (Some(pq), Some(pk), Some(pv), Some(po)) =
9021 (
9022 wq.metal_graph_parts(),
9023 wk.metal_graph_parts(),
9024 wv.metal_graph_parts(),
9025 wo.metal_graph_parts(),
9026 )
9027 else {
9028 return None;
9029 };
9030 if let QTensor::Mapped { model, .. } = wq {
9031 model_ref.get_or_insert_with(|| model.clone());
9032 }
9033 let cache = &self.kv_cache.layers[li];
9034 if cache.mode != crate::kv_cache::KvMode::F32 || cache.o1.is_some() {
9035 return None;
9036 }
9037 plan.push(MetalRowsItem::Attn {
9038 l: AttnGpuLayer {
9039 attn_norm: &lw.input_norm,
9040 post_norm: &lw.post_norm,
9041 wq: pq,
9042 wk: pk,
9043 wv: pv,
9044 wo: po,
9045 ffn,
9046 },
9047 li,
9048 q_norm: q_norm.as_deref(),
9049 k_norm: k_norm.as_deref(),
9050 output_gate: *output_gate,
9051 });
9052 }
9053 _ => return None,
9054 }
9055 }
9056 let model = model_ref?;
9057 let gcfg = self.gdn_cfg.map(|cfg| GdnGpuCfg {
9058 nv: cfg.num_v_heads,
9059 nk: cfg.num_k_heads,
9060 dk: cfg.key_head_dim,
9061 dv: cfg.value_head_dim,
9062 kk: cfg.conv_kernel,
9063 hidden: self.hidden_size,
9064 inter: self.intermediate_size,
9065 c_dim: cfg.conv_dim(),
9066 eps: cfg.rms_eps as f32,
9067 gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
9068 });
9069 Some((plan, model, gcfg))
9070 }
9071
9072 #[cfg(target_os = "macos")]
9074 #[allow(clippy::too_many_arguments)]
9075 fn metal_attn_params<'a>(
9076 li: usize,
9077 cache: &'a crate::kv_cache::LayerKvCache,
9078 q_norm: Option<&'a [f32]>,
9079 k_norm: Option<&'a [f32]>,
9080 output_gate: bool,
9081 inv_freq: &'a [f32],
9082 geom: (usize, usize, usize, usize),
9083 pos0: usize,
9084 kv_id: u64,
9085 scale: f32,
9086 eps: f32,
9087 gemma: bool,
9088 late_qk_norm: bool,
9089 ) -> (crate::gpu_metal::AttnDeviceParams<'a>, usize) {
9090 let (nh, nkv, hd, rd) = geom;
9091 let cpu_k: Vec<&[f32]> = (0..nkv).map(|g| cache.head_keys(g)).collect();
9092 let cpu_v: Vec<&[f32]> = (0..nkv).map(|g| cache.head_values(g)).collect();
9093 let cpu_stored = cpu_k[0].len() / hd;
9094 (
9095 crate::gpu_metal::AttnDeviceParams {
9096 kv_id,
9097 layer: li,
9098 nh,
9099 nkv,
9100 hd,
9101 rd,
9102 position: pos0,
9103 scale,
9104 eps,
9105 gemma,
9106 late_qk_norm,
9107 output_gate,
9108 q_norm,
9109 k_norm,
9110 inv_freq,
9111 cpu_k,
9112 cpu_v,
9113 cpu_stored,
9114 o1: None,
9115 },
9116 cpu_stored,
9117 )
9118 }
9119
9120 #[cfg(target_os = "macos")]
9125 #[allow(clippy::type_complexity)]
9126 fn metal_rows_run(
9127 &mut self,
9128 hiddens: &mut [f32],
9129 pos0: usize,
9130 b: usize,
9131 prefill: bool,
9132 spec: Option<((usize, usize, usize), &[f32], &mut Vec<f32>)>,
9133 mut argmax_out: Option<(usize, &mut Vec<u32>)>,
9137 ) -> MetalRowsRun {
9138 use crate::gpu_metal::{GraphDims, VerifyGraph};
9139 if !crate::gpu_metal::wait_replay() {
9145 tracing::error!("Metal rows graph: the pending async replay failed");
9146 return MetalRowsRun::Failed;
9147 }
9148 spec_stamp("v.wait");
9149 let want = self.gdn_cfg.map(|c| c.state_len()).unwrap_or(0);
9150 for l in &mut self.kv_cache.layers {
9151 if l.linear_state.len() != want && want > 0 {
9152 l.linear_state = vec![0f32; want];
9153 }
9154 }
9155 let Some((plan, model, gcfg)) = self.metal_rows_plan() else {
9156 return MetalRowsRun::Declined;
9157 };
9158 spec_stamp("v.plan");
9159 let dims = GraphDims {
9160 hidden: self.hidden_size,
9161 eps: self.rms_eps as f32,
9162 gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
9163 };
9164 let Some(mut graph) = (if prefill {
9165 VerifyGraph::new_prefill(&model, dims, hiddens, b)
9166 } else {
9167 VerifyGraph::new(&model, dims, hiddens, b)
9168 }) else {
9169 return MetalRowsRun::Declined;
9170 };
9171 let geom = (
9172 self.num_heads,
9173 self.num_kv_heads,
9174 self.head_dim,
9175 self.rotary_dim,
9176 );
9177 let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
9178 let eps = self.rms_eps as f32;
9179 let kv_id = self.graph_kv_id;
9180 let inv_freq = self.inv_freq.clone();
9181 for item in &plan {
9182 let ok = match item {
9183 MetalRowsItem::Gdn { run, .. } => gcfg
9184 .as_ref()
9185 .map(|gc| run.iter().all(|l| graph.gdn_ok(l, gc)))
9186 .unwrap_or(false),
9187 MetalRowsItem::Attn {
9188 l,
9189 li,
9190 q_norm,
9191 k_norm,
9192 output_gate,
9193 } => {
9194 let (p, _) = Self::metal_attn_params(
9195 *li,
9196 &self.kv_cache.layers[*li],
9197 *q_norm,
9198 *k_norm,
9199 *output_gate,
9200 &inv_freq,
9201 geom,
9202 pos0,
9203 kv_id,
9204 self.attn_scale,
9205 eps,
9206 gemma,
9207 self.qk_norm_after_rope,
9208 );
9209 graph.attn_ok(l, &p)
9210 }
9211 };
9212 if !ok {
9213 use std::sync::atomic::{AtomicBool, Ordering};
9214 static SAID: AtomicBool = AtomicBool::new(false);
9215 if !SAID.swap(true, Ordering::Relaxed) {
9216 tracing::warn!("metal rows graph: a layer failed preflight — declining");
9217 }
9218 return MetalRowsRun::Declined;
9219 }
9220 }
9221 let lm = match &spec {
9222 Some((lm, _, _)) => {
9223 if !graph.lm_head_ok(*lm) {
9224 return MetalRowsRun::Declined;
9225 }
9226 Some(*lm)
9227 }
9228 None => None,
9229 };
9230 let mut gdn_layers = Vec::new();
9231 let mut attn_layers = Vec::new();
9232 for item in &plan {
9233 match item {
9234 MetalRowsItem::Gdn { run, first } => {
9235 let ro: Vec<&[f32]> = self.kv_cache.layers[*first..*first + run.len()]
9236 .iter()
9237 .map(|l| l.linear_state.as_slice())
9238 .collect();
9239 if !graph.encode_gdn_run_b(run, &ro, gcfg.as_ref().unwrap()) {
9240 return MetalRowsRun::Declined;
9241 }
9242 gdn_layers.extend(*first..*first + run.len());
9243 }
9244 MetalRowsItem::Attn {
9245 l,
9246 li,
9247 q_norm,
9248 k_norm,
9249 output_gate,
9250 } => {
9251 let (p, cpu_stored) = Self::metal_attn_params(
9252 *li,
9253 &self.kv_cache.layers[*li],
9254 *q_norm,
9255 *k_norm,
9256 *output_gate,
9257 &inv_freq,
9258 geom,
9259 pos0,
9260 kv_id,
9261 self.attn_scale,
9262 eps,
9263 gemma,
9264 self.qk_norm_after_rope,
9265 );
9266 if !graph.encode_attn_b(l, &p) {
9267 return MetalRowsRun::Declined;
9268 }
9269 attn_layers.push((*li, cpu_stored));
9270 }
9271 }
9272 }
9273 if let (Some(lm), Some((_, final_norm, _))) = (lm, spec.as_ref()) {
9274 if !graph.encode_lm_head_b(final_norm, lm) {
9275 return MetalRowsRun::Declined;
9276 }
9277 if let Some((n, _)) = argmax_out.as_ref() {
9282 if !graph.encode_argmax_b(*n) {
9283 argmax_out = None;
9284 }
9285 }
9286 }
9287 spec_stamp("v.enc");
9288 if !graph.sync() {
9289 return MetalRowsRun::Failed;
9290 }
9291 spec_stamp("v.gpu");
9292 match (spec, argmax_out) {
9293 (Some(_), Some((_, ids))) => {
9294 ids.resize(b, 0);
9295 if !graph.read_argmax(ids) {
9296 return MetalRowsRun::Failed;
9297 }
9298 spec_stamp("v.am");
9299 }
9300 (Some((lm, _, logits)), None) => {
9301 logits.resize(b * lm.1, 0.0);
9302 if !graph.read_logits(logits) {
9303 return MetalRowsRun::Failed;
9304 }
9305 spec_stamp("v.lg");
9306 }
9307 (None, _) => {}
9308 }
9309 if !graph.read_hidden(hiddens) {
9310 return MetalRowsRun::Failed;
9311 }
9312 spec_stamp("v.hid");
9313 MetalRowsRun::Completed(MetalVerifyPending {
9314 graph,
9315 gdn_layers,
9316 attn_layers,
9317 })
9318 }
9319
9320 #[cfg(target_os = "macos")]
9326 fn try_batch_graph_metal(
9327 &mut self,
9328 hiddens: &mut [f32],
9329 positions: &[usize],
9330 b: usize,
9331 spec: Option<((usize, usize, usize), &[f32], &mut Vec<f32>)>,
9332 argmax_out: Option<(usize, &mut Vec<u32>)>,
9333 ) -> crate::gpu::BatchGraphOutcome {
9334 let _t0 = std::time::Instant::now();
9335 if positions.len() != b
9336 || positions.windows(2).any(|w| w[1] != w[0] + 1)
9337 || hiddens.len() != b * self.hidden_size
9338 {
9339 return crate::gpu::BatchGraphOutcome::Declined;
9340 }
9341 let pending = match self.metal_rows_run(hiddens, positions[0], b, false, spec, argmax_out) {
9342 MetalRowsRun::Declined => return crate::gpu::BatchGraphOutcome::Declined,
9343 MetalRowsRun::Failed => return crate::gpu::BatchGraphOutcome::Failed,
9344 MetalRowsRun::Completed(pending) => pending,
9345 };
9346 if std::env::var("CMF_GRAPH_SPEC_TIME").is_ok() {
9347 eprintln!(
9348 "metal-verify: {:.1} ms | b={b}",
9349 _t0.elapsed().as_secs_f64() * 1e3
9350 );
9351 }
9352 self.metal_verify = Some(pending);
9353 crate::gpu::BatchGraphOutcome::Completed
9354 }
9355
9356 #[cfg(target_os = "macos")]
9361 fn prefill_rows_metal(
9362 &mut self,
9363 ids: &[u32],
9364 start_pos: usize,
9365 spec: Option<((usize, usize, usize), &[f32], &mut Vec<f32>)>,
9366 ) -> MetalPrefillOutcome {
9367 let b = ids.len();
9368 if b == 0 || b > 512 {
9369 return MetalPrefillOutcome::Declined;
9370 }
9371 METAL_PREFILL_CHUNKS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
9372 let with_head = spec.is_some();
9373 let hs = self.hidden_size;
9374 let mut hiddens = vec![0f32; b * hs];
9375 for (j, &id) in ids.iter().enumerate() {
9376 let e = self.embed_single(id);
9377 hiddens[j * hs..(j + 1) * hs].copy_from_slice(&e);
9378 }
9379 let mut pending = match self.metal_rows_run(&mut hiddens, start_pos, b, true, spec, None) {
9380 MetalRowsRun::Declined => return MetalPrefillOutcome::Declined,
9381 MetalRowsRun::Failed => {
9382 METAL_PREFILL_ERRORS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
9383 return MetalPrefillOutcome::Failed;
9384 }
9385 MetalRowsRun::Completed(pending) => pending,
9386 };
9387 let idxs = pending.gdn_layers.clone();
9389 let mut outs: Vec<&mut [f32]> = self
9390 .kv_cache
9391 .layers
9392 .iter_mut()
9393 .enumerate()
9394 .filter(|(i, _)| idxs.binary_search(i).is_ok())
9395 .map(|(_, l)| l.linear_state.as_mut_slice())
9396 .collect();
9397 if !pending.graph.finish_states(&mut outs) {
9398 METAL_PREFILL_ERRORS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
9399 return MetalPrefillOutcome::Failed;
9400 }
9401 let (nkv, hd) = (self.num_kv_heads, self.head_dim);
9402 let mut rows = Vec::with_capacity(pending.attn_layers.len());
9406 for (li, cpu_stored) in &pending.attn_layers {
9407 let mut kbuf = vec![0f32; b * nkv * hd];
9408 let mut vbuf = vec![0f32; b * nkv * hd];
9409 if !crate::gpu_metal::kv_mirror_read_rows(
9410 self.graph_kv_id,
9411 *li,
9412 nkv,
9413 hd,
9414 *cpu_stored,
9415 b,
9416 &mut kbuf,
9417 &mut vbuf,
9418 ) {
9419 METAL_PREFILL_ERRORS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
9420 return MetalPrefillOutcome::Failed;
9421 }
9422 rows.push((*li, *cpu_stored, kbuf, vbuf));
9423 }
9424 for (li, cpu_stored, kbuf, vbuf) in rows {
9425 let cache = &mut self.kv_cache.layers[li];
9426 for r in 0..b {
9427 cache.append(
9428 &kbuf[r * nkv * hd..(r + 1) * nkv * hd],
9429 &vbuf[r * nkv * hd..(r + 1) * nkv * hd],
9430 &[],
9431 );
9432 }
9433 crate::gpu_metal::kv_mirror_set_stored(self.graph_kv_id, li, cpu_stored + b);
9434 }
9435 METAL_PREFILL_ROWS.fetch_add(b as u64, std::sync::atomic::Ordering::Relaxed);
9436 if with_head {
9437 METAL_PREFILL_HEAD_ROWS.fetch_add(b as u64, std::sync::atomic::Ordering::Relaxed);
9438 }
9439 MetalPrefillOutcome::Completed(hiddens)
9440 }
9441
9442 #[cfg(target_os = "macos")]
9443 fn prefill_batch_metal(&mut self, ids: &[u32], start_pos: usize) -> MetalPrefillOutcome {
9444 self.prefill_rows_metal(ids, start_pos, None)
9445 }
9446
9447 #[cfg(target_os = "macos")]
9452 fn nll_batch_metal(&mut self, ids: &[u32], start: usize) -> MetalBatchNllOutcome {
9453 if ids.len() < 2 || self.o1_active() || self.head_clusters.is_some() {
9454 return MetalBatchNllOutcome::Declined;
9455 }
9456 let Some(lm) = self.weights.lm_head.metal_graph_parts() else {
9457 return MetalBatchNllOutcome::Declined;
9458 };
9459 let chunk = std::env::var("CMF_METAL_PREFILL_CHUNK")
9460 .ok()
9461 .and_then(|v| v.parse::<usize>().ok())
9462 .filter(|&v| (1..=512).contains(&v))
9463 .unwrap_or(32);
9464 let final_norm = self.weights.final_norm.clone();
9465 let mut nll = 0.0f64;
9466 let mut count = 0usize;
9467 let mut pos = 0usize;
9468 let mut completed = 0usize;
9469 while pos < ids.len() {
9470 let end = (pos + chunk).min(ids.len());
9471 let mut logits = Vec::new();
9472 let outcome = self.prefill_rows_metal(
9473 &ids[pos..end],
9474 pos,
9475 Some((lm, &final_norm, &mut logits)),
9476 );
9477 match outcome {
9478 MetalPrefillOutcome::Declined => {
9479 return if completed == 0 {
9480 MetalBatchNllOutcome::Declined
9481 } else {
9482 MetalBatchNllOutcome::Failed(format!(
9483 "ordinary Metal NLL batch declined after {completed} chunks"
9484 ))
9485 };
9486 }
9487 MetalPrefillOutcome::Failed => {
9488 return MetalBatchNllOutcome::Failed(
9489 "ordinary Metal NLL batch failed after admission".to_string(),
9490 );
9491 }
9492 MetalPrefillOutcome::Completed(_) => {}
9493 }
9494 completed += 1;
9495 let vocab = self.vocab_size.min(lm.1);
9496 if logits.len() != (end - pos) * lm.1 || vocab == 0 {
9497 return MetalBatchNllOutcome::Failed(
9498 "ordinary Metal NLL head returned an invalid shape".to_string(),
9499 );
9500 }
9501 for row in 0..(end - pos) {
9502 let absolute = pos + row;
9503 if absolute < start || absolute + 1 >= ids.len() {
9504 continue;
9505 }
9506 let lg = &mut logits[row * lm.1..row * lm.1 + vocab];
9507 if let Some(mu) = self.logit_multiplier {
9508 for v in lg.iter_mut() {
9509 *v *= mu;
9510 }
9511 }
9512 if let Some(c) = self.final_softcap {
9513 for v in lg.iter_mut() {
9514 *v = c * (*v / c).tanh();
9515 }
9516 }
9517 let target = ids[absolute + 1] as usize;
9518 if target >= vocab {
9519 return MetalBatchNllOutcome::Failed(format!(
9520 "target token {target} exceeds Metal head rows {vocab}"
9521 ));
9522 }
9523 let max = lg.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
9524 let lse: f64 = lg
9525 .iter()
9526 .map(|&v| ((v - max) as f64).exp())
9527 .sum::<f64>()
9528 .ln()
9529 + max as f64;
9530 nll += lse - lg[target] as f64;
9531 count += 1;
9532 }
9533 pos = end;
9534 }
9535 MetalBatchNllOutcome::Completed(nll, count)
9536 }
9537
9538 #[cfg(target_os = "macos")]
9542 fn metal_verify_commit(&mut self, a: usize) -> bool {
9543 let Some(mut pending) = self.metal_verify.take() else {
9544 return false;
9545 };
9546 let n = a + 1;
9547 let idxs = pending.gdn_layers.clone();
9549 let mut outs: Vec<&mut [f32]> = self
9550 .kv_cache
9551 .layers
9552 .iter_mut()
9553 .enumerate()
9554 .filter(|(i, _)| idxs.binary_search(i).is_ok())
9555 .map(|(_, l)| l.linear_state.as_mut_slice())
9556 .collect();
9557 if !pending.graph.commit(n, &mut outs) {
9558 return false;
9559 }
9560 spec_stamp("c.replay");
9561 let (nkv, hd) = (self.num_kv_heads, self.head_dim);
9562 let mut rows = Vec::with_capacity(pending.attn_layers.len());
9566 for (li, cpu_stored) in &pending.attn_layers {
9567 let mut kbuf = vec![0f32; n * nkv * hd];
9568 let mut vbuf = vec![0f32; n * nkv * hd];
9569 if !crate::gpu_metal::kv_mirror_read_rows(
9570 self.graph_kv_id,
9571 *li,
9572 nkv,
9573 hd,
9574 *cpu_stored,
9575 n,
9576 &mut kbuf,
9577 &mut vbuf,
9578 ) {
9579 return false;
9580 }
9581 rows.push((*li, *cpu_stored, kbuf, vbuf));
9582 }
9583 for (li, cpu_stored, kbuf, vbuf) in rows {
9584 let cache = &mut self.kv_cache.layers[li];
9585 for r in 0..n {
9586 cache.append(
9587 &kbuf[r * nkv * hd..(r + 1) * nkv * hd],
9588 &vbuf[r * nkv * hd..(r + 1) * nkv * hd],
9589 &[],
9590 );
9591 }
9592 crate::gpu_metal::kv_mirror_set_stored(self.graph_kv_id, li, cpu_stored + n);
9593 }
9594 spec_stamp("c.kv");
9595 true
9596 }
9597
9598 #[cfg(target_os = "macos")]
9605 fn mtp_warm_batch_submit(
9606 &mut self,
9607 m: &mut MtpModule,
9608 pairs: &[(&[f32], u32)],
9609 first_pos: usize,
9610 ) -> Option<MetalWarmPending> {
9611 use crate::gpu_metal::{AttnDeviceParams, AttnGpuLayer, GraphDims, MetalFfn, VerifyGraph};
9612 let b = pairs.len();
9613 if b == 0 || b > 512 || m.kv.mode != crate::kv_cache::KvMode::F32 || m.kv.o1.is_some() {
9614 return None;
9615 }
9616 let AttnKind::Full {
9617 wq,
9618 wk,
9619 wv,
9620 wo,
9621 q_norm,
9622 k_norm,
9623 output_gate,
9624 softplus_gate: None,
9625 bias: None,
9626 } = &m.layer.attn
9627 else {
9628 return None;
9629 };
9630 let FfnKind::Dense(d) = &m.layer.ffn else {
9631 return None;
9632 };
9633 if !d.segs.is_empty() {
9634 return None;
9635 }
9636 let (Some(pq), Some(pk), Some(pv), Some(po)) =
9637 (wq.q1_parts(), wk.q1_parts(), wv.q1_parts(), wo.q1_parts())
9638 else {
9639 return None;
9640 };
9641 let (Some(g), Some(u), Some(dn)) = (
9642 d.gate_proj.q1_parts(),
9643 d.up_proj.q1_parts(),
9644 d.down_proj.q1_parts(),
9645 ) else {
9646 return None;
9647 };
9648 let Some(eh) = m.eh_proj.q1_parts() else {
9649 return None;
9650 };
9651 let QTensor::Mapped { model, .. } = wq else {
9652 return None;
9653 };
9654 let model = model.clone();
9655 let hs = self.hidden_size;
9656 let mut cat = vec![0f32; b * 2 * hs];
9658 for (j, (h, tok)) in pairs.iter().enumerate() {
9659 let e = self.embed_single(*tok);
9660 let (ce, ch) = cat[j * 2 * hs..(j + 1) * 2 * hs].split_at_mut(hs);
9661 inference::rms_norm_into(&e, &m.enorm, self.rms_eps, self.norm_style, ce);
9662 inference::rms_norm_into(h, &m.hnorm, self.rms_eps, self.norm_style, ch);
9663 }
9664 let dims = GraphDims {
9665 hidden: hs,
9666 eps: self.rms_eps as f32,
9667 gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
9668 };
9669 spec_stamp("w.cat");
9670 let Some(mut graph) = VerifyGraph::new_via_proj(&model, dims, eh, &cat, b) else {
9671 return None;
9672 };
9673 spec_stamp("w.new");
9674 let l = AttnGpuLayer {
9675 attn_norm: &m.layer.input_norm,
9676 post_norm: &m.layer.post_norm,
9677 wq: pq,
9678 wk: pk,
9679 wv: pv,
9680 wo: po,
9681 ffn: MetalFfn::Dense {
9682 gate: g,
9683 up: u,
9684 down: dn,
9685 },
9686 };
9687 let (nh, nkv, hd, rd) = (
9688 self.num_heads,
9689 self.num_kv_heads,
9690 self.head_dim,
9691 self.rotary_dim,
9692 );
9693 let inv_freq = self.inv_freq.clone();
9694 let cpu_stored;
9695 {
9696 let cache = &m.kv;
9697 let cpu_k: Vec<&[f32]> = (0..nkv).map(|g| cache.head_keys(g)).collect();
9698 let cpu_v: Vec<&[f32]> = (0..nkv).map(|g| cache.head_values(g)).collect();
9699 cpu_stored = cpu_k[0].len() / hd;
9700 if cpu_stored > first_pos {
9705 spec_stamp("w.decl");
9706 return None;
9707 }
9708 let p = AttnDeviceParams {
9709 kv_id: self.mtp_kv_id(),
9710 layer: Self::MTP_LAYER_BASE,
9711 nh,
9712 nkv,
9713 hd,
9714 rd,
9715 position: first_pos,
9716 scale: self.attn_scale,
9717 eps: self.rms_eps as f32,
9718 gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
9719 late_qk_norm: self.qk_norm_after_rope,
9720 output_gate: *output_gate,
9721 q_norm: q_norm.as_deref(),
9722 k_norm: k_norm.as_deref(),
9723 inv_freq: &inv_freq,
9724 cpu_k,
9725 cpu_v,
9726 cpu_stored,
9727 o1: None,
9728 };
9729 if !graph.attn_ok(&l, &p) || !graph.encode_attn_b(&l, &p) {
9730 return None;
9731 }
9732 }
9733 spec_stamp("w.enc");
9734 if !graph.submit() {
9735 return None;
9736 }
9737 spec_stamp("w.sub");
9738 Some(MetalWarmPending {
9739 graph,
9740 cpu_stored,
9741 b,
9742 })
9743 }
9744
9745 #[cfg(target_os = "macos")]
9748 fn mtp_warm_batch_metal(
9749 &mut self,
9750 m: &mut MtpModule,
9751 pairs: &[(&[f32], u32)],
9752 first_pos: usize,
9753 ) -> bool {
9754 match self.mtp_warm_batch_submit(m, pairs, first_pos) {
9755 Some(p) => self.mtp_warm_batch_finish(m, p),
9756 None => false,
9757 }
9758 }
9759
9760 #[cfg(target_os = "macos")]
9765 fn mtp_warm_batch_finish(&mut self, m: &mut MtpModule, pending: MetalWarmPending) -> bool {
9766 let MetalWarmPending {
9767 mut graph,
9768 cpu_stored,
9769 b,
9770 } = pending;
9771 let (nkv, hd) = (self.num_kv_heads, self.head_dim);
9772 if !graph.sync() {
9773 return false;
9774 }
9775 spec_stamp("w.gpu");
9776 let mut kbuf = vec![0f32; b * nkv * hd];
9777 let mut vbuf = vec![0f32; b * nkv * hd];
9778 if !crate::gpu_metal::kv_mirror_read_rows(
9779 self.mtp_kv_id(),
9780 Self::MTP_LAYER_BASE,
9781 nkv,
9782 hd,
9783 cpu_stored,
9784 b,
9785 &mut kbuf,
9786 &mut vbuf,
9787 ) {
9788 return false;
9789 }
9790 for r in 0..b {
9791 m.kv.append(
9792 &kbuf[r * nkv * hd..(r + 1) * nkv * hd],
9793 &vbuf[r * nkv * hd..(r + 1) * nkv * hd],
9794 &[],
9795 );
9796 }
9797 crate::gpu_metal::kv_mirror_set_stored(
9798 self.mtp_kv_id(),
9799 Self::MTP_LAYER_BASE,
9800 cpu_stored + b,
9801 );
9802 spec_stamp("w.kv");
9803 true
9804 }
9805
9806 pub(crate) fn note_draft_id(&mut self, id: u32) {
9813 let cut = Self::draft_vocab_rows(usize::MAX).max(131_072);
9814 if (id as usize) >= cut {
9815 self.draft_full_streak = 16;
9816 } else {
9817 self.draft_full_streak = self.draft_full_streak.saturating_sub(1);
9818 }
9819 }
9820
9821 fn draft_head_rows(&self, head_rows: usize) -> usize {
9824 if self.draft_full_streak > 0 {
9825 head_rows
9826 } else {
9827 Self::draft_vocab_rows(head_rows)
9828 }
9829 }
9830
9831 fn draft_vocab_rows(head_rows: usize) -> usize {
9834 static N: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
9835 let n = *N.get_or_init(|| {
9836 std::env::var("CMF_DRAFT_VOCAB")
9837 .ok()
9838 .and_then(|v| v.parse().ok())
9839 .unwrap_or(65536)
9840 });
9841 if n == 0 { head_rows } else { n.min(head_rows) }
9842 }
9843
9844 #[cfg(target_os = "macos")]
9849 fn mtp_step_metal(
9850 &mut self,
9851 m: &mut MtpModule,
9852 hidden: &[f32],
9853 next_token: u32,
9854 position: usize,
9855 want_logits: bool,
9856 ) -> Option<(Vec<f32>, Vec<f32>)> {
9857 use crate::gpu_metal::{AttnDeviceParams, AttnGpuLayer, GraphDims, MetalFfn, TokenGraph};
9858 if std::env::var("CMF_MTP_GRAPH").as_deref() == Ok("0")
9859 || !crate::gpu::q1_force()
9860 || !crate::gpu::enabled_here()
9861 || self.attn_softcap > 0.0
9862 || self.attention_heads_per_layer.is_some()
9863 || m.kv.mode != crate::kv_cache::KvMode::F32
9864 || m.kv.o1.is_some()
9865 {
9866 return None;
9867 }
9868 let AttnKind::Full {
9869 wq,
9870 wk,
9871 wv,
9872 wo,
9873 q_norm,
9874 k_norm,
9875 output_gate,
9876 softplus_gate: None,
9877 bias: None,
9878 } = &m.layer.attn
9879 else {
9880 return None;
9881 };
9882 let FfnKind::Dense(d) = &m.layer.ffn else {
9883 return None;
9884 };
9885 if d.act != Act::Silu || !d.segs.is_empty() {
9886 return None;
9887 }
9888 let (pq, pk, pv, po) = (
9889 wq.q1_parts()?,
9890 wk.q1_parts()?,
9891 wv.q1_parts()?,
9892 wo.q1_parts()?,
9893 );
9894 let (g, u, dn) = (
9895 d.gate_proj.q1_parts()?,
9896 d.up_proj.q1_parts()?,
9897 d.down_proj.q1_parts()?,
9898 );
9899 let QTensor::Mapped { model, .. } = wq else {
9900 return None;
9901 };
9902 let model = model.clone();
9903 let lm = if want_logits {
9904 Some(self.weights.lm_head.q1_parts()?)
9905 } else {
9906 None
9907 };
9908 let dims = GraphDims {
9909 hidden: self.hidden_size,
9910 eps: self.rms_eps as f32,
9911 gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
9912 };
9913 let hs = self.hidden_size;
9916 let mut x = vec![0f32; hs];
9917 let mut graph = TokenGraph::new(&model, dims, &x)?;
9918 let mut folded = false;
9919 if let Some(eh) = m.eh_proj.q1_parts() {
9920 let e = self.embed_single(next_token);
9921 let mut cat = vec![0.0f32; 2 * hs];
9922 let (cat_e, cat_h) = cat.split_at_mut(hs);
9923 inference::rms_norm_into(&e, &m.enorm, self.rms_eps, self.norm_style, cat_e);
9924 inference::rms_norm_into(hidden, &m.hnorm, self.rms_eps, self.norm_style, cat_h);
9925 folded = graph.encode_input_proj(eh, &cat);
9926 }
9927 if !folded {
9928 x = self.mtp_block_input(m, hidden, next_token);
9929 graph = TokenGraph::new(&model, dims, &x)?;
9930 }
9931 spec_stamp("d.in");
9932 let l = AttnGpuLayer {
9933 attn_norm: &m.layer.input_norm,
9934 post_norm: &m.layer.post_norm,
9935 wq: pq,
9936 wk: pk,
9937 wv: pv,
9938 wo: po,
9939 ffn: MetalFfn::Dense {
9940 gate: g,
9941 up: u,
9942 down: dn,
9943 },
9944 };
9945 let (nh, nkv, hd, rd) = (
9946 self.num_heads,
9947 self.num_kv_heads,
9948 self.head_dim,
9949 self.rotary_dim,
9950 );
9951 let inv_freq = self.inv_freq.clone();
9952 {
9953 let cache = &m.kv;
9954 let cpu_k: Vec<&[f32]> = (0..nkv).map(|g| cache.head_keys(g)).collect();
9955 let cpu_v: Vec<&[f32]> = (0..nkv).map(|g| cache.head_values(g)).collect();
9956 let cpu_stored = cpu_k[0].len() / hd;
9957 let p = AttnDeviceParams {
9958 kv_id: self.mtp_kv_id(),
9959 layer: Self::MTP_LAYER_BASE,
9960 nh,
9961 nkv,
9962 hd,
9963 rd,
9964 position,
9965 scale: self.attn_scale,
9966 eps: self.rms_eps as f32,
9967 gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
9968 late_qk_norm: self.qk_norm_after_rope,
9969 output_gate: *output_gate,
9970 q_norm: q_norm.as_deref(),
9971 k_norm: k_norm.as_deref(),
9972 inv_freq: &inv_freq,
9973 cpu_k,
9974 cpu_v,
9975 cpu_stored,
9976 o1: None,
9977 };
9978 if !graph.attn_device_ok(&l, &p) || !graph.encode_attn_device(&l, &p) {
9979 return None;
9980 }
9981 }
9982 let draft_rows = if let Some(lm) = lm {
9988 self.draft_head_rows(lm.1)
9989 } else {
9990 0
9991 };
9992 if let Some(lm) = lm {
9993 if !graph.lm_head_ok(lm) {
9994 return None;
9995 }
9996 if draft_rows < lm.1 {
9997 if !graph.encode_lm_head_part(&m.final_norm, lm, draft_rows) {
9998 return None;
9999 }
10000 } else {
10001 graph.encode_lm_head(&m.final_norm, lm);
10002 }
10003 }
10004 spec_stamp("d.enc");
10005 if graph.sync_checked().is_err() {
10006 return None;
10007 }
10008 spec_stamp("d.gpu");
10009 let mut logits = Vec::new();
10010 if let Some(lm) = lm {
10011 let n_read = draft_rows.min(lm.1).min(self.vocab_size);
10012 logits = attention::take_buf(n_read);
10013 graph.read_logits(&mut logits);
10014 logits.resize(self.vocab_size, f32::NEG_INFINITY);
10016 }
10017 graph.finish(&mut x);
10018 let mut krow = attention::take_buf(nkv * hd);
10019 let mut vrow = attention::take_buf(nkv * hd);
10020 if crate::gpu_metal::kv_mirror_read_last(
10021 self.mtp_kv_id(),
10022 Self::MTP_LAYER_BASE,
10023 nkv,
10024 hd,
10025 &mut krow,
10026 &mut vrow,
10027 ) {
10028 m.kv.append(&krow, &vrow, &[]);
10029 }
10030 attention::recycle_buf(&mut krow);
10031 attention::recycle_buf(&mut vrow);
10032 spec_stamp("d.rd");
10033 Some((logits, x))
10034 }
10035
10036 fn mtp_chain_on() -> bool {
10049 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
10050 *ON.get_or_init(|| std::env::var("CMF_MTP_CHAIN").as_deref() != Ok("0"))
10051 }
10052
10053 #[cfg(target_os = "macos")]
10065 fn mtp_draft_chain_metal(
10066 &mut self,
10067 m: &mut MtpModule,
10068 hidden: &[f32],
10069 t_next: u32,
10070 position: usize,
10071 k: usize,
10072 ) -> Result<Vec<u32>, bool> {
10073 use crate::gpu_metal::{AttnDeviceParams, AttnGpuLayer, GraphDims, MetalFfn, TokenGraph};
10074 if k == 0
10075 || k > 64
10076 || !Self::mtp_chain_on()
10077 || std::env::var("CMF_MTP_GRAPH").as_deref() == Ok("0")
10078 || !crate::gpu::q1_force()
10079 || !crate::gpu::enabled_here()
10080 || self.attn_softcap > 0.0
10081 || self.attention_heads_per_layer.is_some()
10082 || m.kv.mode != crate::kv_cache::KvMode::F32
10083 || m.kv.o1.is_some()
10084 || self.dsv4.is_some()
10086 || self.dsv41.is_some()
10087 || self.qwen4_exp.is_some()
10088 || self.g3n.is_some()
10089 {
10090 return Err(false);
10091 }
10092 let AttnKind::Full {
10093 wq,
10094 wk,
10095 wv,
10096 wo,
10097 q_norm,
10098 k_norm,
10099 output_gate,
10100 softplus_gate: None,
10101 bias: None,
10102 } = &m.layer.attn
10103 else {
10104 return Err(false);
10105 };
10106 let FfnKind::Dense(d) = &m.layer.ffn else {
10107 return Err(false);
10108 };
10109 if d.act != Act::Silu || !d.segs.is_empty() {
10110 return Err(false);
10111 }
10112 let (Some(pq), Some(pk), Some(pv), Some(po)) =
10113 (wq.q1_parts(), wk.q1_parts(), wv.q1_parts(), wo.q1_parts())
10114 else {
10115 return Err(false);
10116 };
10117 let (Some(g), Some(u), Some(dn)) = (
10118 d.gate_proj.q1_parts(),
10119 d.up_proj.q1_parts(),
10120 d.down_proj.q1_parts(),
10121 ) else {
10122 return Err(false);
10123 };
10124 let (Some(eh), Some(lm)) = (m.eh_proj.q1_parts(), self.weights.lm_head.q1_parts()) else {
10125 return Err(false);
10126 };
10127 let QTensor::Mapped { model, .. } = wq else {
10128 return Err(false);
10129 };
10130 let model = model.clone();
10131 let QTensor::Mapped {
10134 model: em,
10135 idx: eidx,
10136 dtype: cortiq_core::TensorDtype::Q4TiledP,
10137 ..
10138 } = &self.weights.embed_tokens
10139 else {
10140 return Err(false);
10141 };
10142 if !std::sync::Arc::ptr_eq(em, &model)
10143 || crate::prism::is_inverse_embedding(&model, &model.tensors[*eidx].name)
10144 {
10145 return Err(false);
10146 }
10147 let embed = (
10148 *eidx,
10149 self.weights.embed_tokens.rows(),
10150 self.weights.embed_tokens.cols(),
10151 );
10152 if embed.2 != self.hidden_size || hidden.len() != self.hidden_size {
10153 return Err(false);
10154 }
10155 let dims = GraphDims {
10156 hidden: self.hidden_size,
10157 eps: self.rms_eps as f32,
10158 gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
10159 };
10160 let Some(mut graph) = TokenGraph::new(&model, dims, hidden) else {
10161 return Err(false);
10162 };
10163 if !graph.chain_embed_ok(embed) || !graph.lm_head_ok(lm) {
10164 return Err(false);
10165 }
10166 let l = AttnGpuLayer {
10167 attn_norm: &m.layer.input_norm,
10168 post_norm: &m.layer.post_norm,
10169 wq: pq,
10170 wk: pk,
10171 wv: pv,
10172 wo: po,
10173 ffn: MetalFfn::Dense {
10174 gate: g,
10175 up: u,
10176 down: dn,
10177 },
10178 };
10179 let (nh, nkv, hd, rd) = (
10180 self.num_heads,
10181 self.num_kv_heads,
10182 self.head_dim,
10183 self.rotary_dim,
10184 );
10185 let inv_freq = self.inv_freq.clone();
10186 let draft_rows = self.draft_head_rows(lm.1);
10187 let n_arg = draft_rows.min(lm.1).min(self.vocab_size);
10188 if n_arg == 0 {
10189 return Err(false);
10190 }
10191 let split = std::env::var("CMF_MTP_CHAIN_SPLIT").as_deref() == Ok("1");
10198 let t_chain = std::time::Instant::now();
10199 graph.chain_ids_init(t_next, k);
10200 let cpu_stored;
10201 {
10202 let cache = &m.kv;
10203 let cpu_k: Vec<&[f32]> = (0..nkv).map(|g| cache.head_keys(g)).collect();
10204 let cpu_v: Vec<&[f32]> = (0..nkv).map(|g| cache.head_values(g)).collect();
10205 cpu_stored = cpu_k[0].len() / hd;
10206 for j in 0..k {
10207 if !graph.encode_chain_input(
10208 embed,
10209 j as u32,
10210 &m.enorm,
10211 &m.hnorm,
10212 self.embed_multiplier,
10213 eh,
10214 ) {
10215 return Err(false);
10216 }
10217 let p = AttnDeviceParams {
10221 kv_id: self.mtp_kv_id(),
10222 layer: Self::MTP_LAYER_BASE,
10223 nh,
10224 nkv,
10225 hd,
10226 rd,
10227 position: position + j,
10228 scale: self.attn_scale,
10229 eps: self.rms_eps as f32,
10230 gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
10231 late_qk_norm: self.qk_norm_after_rope,
10232 output_gate: *output_gate,
10233 q_norm: q_norm.as_deref(),
10234 k_norm: k_norm.as_deref(),
10235 inv_freq: &inv_freq,
10236 cpu_k: cpu_k.clone(),
10237 cpu_v: cpu_v.clone(),
10238 cpu_stored: cpu_stored + j,
10239 o1: None,
10240 };
10241 if !graph.attn_device_ok(&l, &p) || !graph.encode_attn_device(&l, &p) {
10242 return Err(false);
10243 }
10244 if draft_rows < lm.1 {
10245 if !graph.encode_lm_head_part(&m.final_norm, lm, draft_rows) {
10246 return Err(false);
10247 }
10248 } else {
10249 graph.encode_lm_head(&m.final_norm, lm);
10250 }
10251 if !graph.encode_argmax(n_arg, j as u32 + 1) {
10252 return Err(false);
10253 }
10254 if split {
10255 graph.commit();
10258 }
10259 }
10260 }
10261 let t_enc = t_chain.elapsed();
10262 if graph.sync_checked().is_err() {
10263 return Err(true);
10264 }
10265 if std::env::var_os("CMF_GRAPH_SPEC_TIME").is_some() {
10266 eprintln!(
10267 "mtp-chain: encode {:.1} ms | wait {:.1} ms (k={k}, head rows {draft_rows}{})",
10268 t_enc.as_secs_f64() * 1e3,
10269 (t_chain.elapsed() - t_enc).as_secs_f64() * 1e3,
10270 if split { ", split" } else { "" }
10271 );
10272 }
10273 let mut ids = vec![0u32; k];
10274 if !graph.chain_ids_read(&mut ids) {
10275 return Err(true);
10276 }
10277 let mut kbuf = vec![0f32; k * nkv * hd];
10278 let mut vbuf = vec![0f32; k * nkv * hd];
10279 if !crate::gpu_metal::kv_mirror_read_rows(
10280 self.mtp_kv_id(),
10281 Self::MTP_LAYER_BASE,
10282 nkv,
10283 hd,
10284 cpu_stored,
10285 k,
10286 &mut kbuf,
10287 &mut vbuf,
10288 ) {
10289 return Err(true);
10290 }
10291 for r in 0..k {
10292 m.kv.append(
10293 &kbuf[r * nkv * hd..(r + 1) * nkv * hd],
10294 &vbuf[r * nkv * hd..(r + 1) * nkv * hd],
10295 &[],
10296 );
10297 }
10298 Ok(ids)
10299 }
10300
10301 fn try_batch_graph_wgpu(
10302 &self,
10303 hiddens: &mut [f32],
10304 positions: &[usize],
10305 k: usize,
10306 spec: Option<crate::gpu::SpecTail<'_>>,
10307 ) -> crate::gpu::BatchGraphOutcome {
10308 let _tb = std::time::Instant::now();
10309 let batch_debug = std::env::var_os("CMF_BATCH_DEBUG").is_some();
10310 if self.attn_softcap > 0.0 {
10311 return crate::gpu::BatchGraphOutcome::Declined; }
10313 let nh = self.num_heads;
10314 let (nkv, hd, rd) = self.layer_geom(0);
10315 let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
10316 fn gw(t: &QTensor) -> Option<crate::gpu::GraphW<'_>> {
10317 if let Some((m, i, kind, rs)) = t
10318 .graph_weight()
10319 .or_else(|| t.graph_weight_descriptor())
10320 {
10321 let name = &m.tensors[i].name;
10322 let prism = if crate::prism::is_inverse_embedding(m, name) {
10323 crate::gpu::GraphPrismOp::InverseEmbedding
10324 } else if crate::prism::is_forward_weight(m, name) {
10325 crate::gpu::GraphPrismOp::Forward
10326 } else {
10327 crate::gpu::GraphPrismOp::None
10328 };
10329 return Some(crate::gpu::GraphW {
10330 idx: i,
10331 kind,
10332 row_scale: rs,
10333 data: &[],
10334 prism,
10335 affine: crate::prism::is_affine_target(m, name),
10336 });
10337 }
10338 if std::env::var_os("CMF_BATCH_DEBUG").is_some() {
10339 eprintln!(
10340 "batch graph: tensor has no graph descriptor/f32 fallback rows={} cols={}",
10341 t.rows(),
10342 t.cols()
10343 );
10344 }
10345 t.as_f32().map(|d| crate::gpu::GraphW {
10346 idx: 0,
10347 kind: 4,
10348 row_scale: &[],
10349 data: d,
10350 prism: crate::gpu::GraphPrismOp::None,
10351 affine: false,
10352 })
10353 }
10354 let built: Option<(
10355 Vec<crate::gpu::GraphLayer<'_>>,
10356 std::sync::Arc<cortiq_core::CmfModel>,
10357 )> = (|| {
10358 let mut layers = Vec::with_capacity(self.num_layers);
10359 let mut model = None;
10360 for li in 0..self.num_layers {
10361 let lw = &self.weights.layers[self.phys_layer(li)];
10362 let gffn = match &lw.ffn {
10369 FfnKind::Dense(d) if !d.segs.is_empty() => {
10370 if batch_debug {
10371 eprintln!("batch graph: dense segmented FFN at layer {li}");
10372 }
10373 return None;
10374 }
10375 FfnKind::Dense(d) => crate::gpu::GraphFfn::Dense {
10376 gate: gw(&d.gate_proj)?,
10377 up: gw(&d.up_proj)?,
10378 down: gw(&d.down_proj)?,
10379 },
10380 FfnKind::Moe(m) => {
10381 if m.route_tau.is_some() || m.mask.is_some() {
10388 return None;
10389 }
10390 let (se, sg) = m.shared.as_ref()?;
10393 let shared_gated = sg.is_some();
10394 let sgate = match sg {
10395 Some(sg) => gw(sg)?,
10396 None => gw(&m.router)?,
10399 };
10400 let router = gw(&m.router)?;
10401 if router.prism != crate::gpu::GraphPrismOp::None
10407 || router.affine
10408 || sgate.prism != crate::gpu::GraphPrismOp::None
10409 || sgate.affine
10410 {
10411 return None;
10412 }
10413 let inter = m.experts.first()?.gate_proj.rows();
10414 let mut experts = Vec::with_capacity(m.experts.len() + 1);
10415 let mut q4tp: Option<bool> = None;
10416 let mut gu_q2: Option<bool> = None;
10417 for e in m.experts.iter().chain(std::iter::once(se)) {
10418 if !matches!(e.act, Act::Silu)
10419 || e.gate_proj.rows() != inter
10420 || e.up_proj.rows() != inter
10421 {
10422 return None;
10423 }
10424 let (mm, gi, ui, di, is_p, is_q2) = match e.gate_proj.mapped_q4t() {
10428 Some((mm, gi)) => (
10429 mm,
10430 gi,
10431 e.up_proj.mapped_q4t()?.1,
10432 e.down_proj.mapped_q4t()?.1,
10433 false,
10434 false,
10435 ),
10436 None => match e.gate_proj.mapped_q2tp() {
10437 Some((mm, gi)) => (
10438 mm,
10439 gi,
10440 e.up_proj.mapped_q2tp()?.1,
10441 e.down_proj.mapped_q4tp()?.1,
10442 true,
10443 true,
10444 ),
10445 None => {
10446 let (mm, gi) = e.gate_proj.mapped_q4tp()?;
10447 (
10448 mm,
10449 gi,
10450 e.up_proj.mapped_q4tp()?.1,
10451 e.down_proj.mapped_q4tp()?.1,
10452 true,
10453 false,
10454 )
10455 }
10456 },
10457 };
10458 if *q4tp.get_or_insert(is_p) != is_p
10459 || *gu_q2.get_or_insert(is_q2) != is_q2
10460 {
10461 return None;
10462 }
10463 if [gi, ui, di].into_iter().any(|idx| {
10464 mm.tensors
10465 .get(idx)
10466 .is_some_and(|t| {
10467 crate::prism::is_forward_weight(mm, &t.name)
10468 || crate::prism::is_affine_target(mm, &t.name)
10469 })
10470 }) {
10471 return None;
10472 }
10473 model.get_or_insert_with(|| mm.clone());
10474 experts.push((gi, ui, di));
10475 }
10476 crate::gpu::GraphFfn::Moe {
10477 router,
10478 shared_gate: sgate,
10479 experts,
10480 n_exp: m.experts.len(),
10481 top_k: m.top_k,
10482 inter,
10483 norm_topk: m.norm_topk_prob,
10484 q4tp: q4tp?,
10485 gu_q2: gu_q2.unwrap_or(false),
10486 sigmoid: m.router_sigmoid,
10487 bias: m.expert_bias.as_deref(),
10488 has_shared: true,
10489 shared_gated,
10490 route_scale: m.routed_scaling,
10491 }
10492 }
10493 _ => return None,
10494 };
10495 let attn = match &lw.attn {
10496 AttnKind::Full {
10497 wq,
10498 wk,
10499 wv,
10500 wo,
10501 q_norm,
10502 k_norm,
10503 output_gate,
10504 softplus_gate,
10505 bias,
10506 } => {
10507 if softplus_gate.is_some() || self.attention_heads_per_layer.is_some() {
10508 if batch_debug {
10509 eprintln!(
10510 "batch graph: unsupported Full attention gate at layer {li} softplus={} heads={}",
10511 softplus_gate.is_some(),
10512 self.attention_heads_per_layer.is_some()
10513 );
10514 }
10515 return None;
10516 }
10517 let (m, _, _, _) = wq
10518 .graph_weight()
10519 .or_else(|| wq.graph_weight_descriptor())?;
10520 model = Some(m.clone());
10521 crate::gpu::GraphAttn::Full {
10522 wq: gw(wq)?,
10523 wk: gw(wk)?,
10524 wv: gw(wv)?,
10525 wo: gw(wo)?,
10526 q_norm: q_norm.as_deref(),
10527 k_norm: k_norm.as_deref(),
10528 late_qk_norm: self.qk_norm_after_rope,
10529 bias: bias
10530 .as_ref()
10531 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
10532 output_gate: *output_gate,
10533 cpu_k: self.kv_cache.layers[li].k_heads(),
10534 cpu_v: self.kv_cache.layers[li].v_heads(),
10535 }
10536 }
10537 AttnKind::LinearGdn(w) => {
10538 let Some(cfg) = self.gdn_cfg else {
10539 if batch_debug {
10540 eprintln!("batch graph: no GDN config at layer {li}");
10541 }
10542 return None;
10543 };
10544 let (m, _, _, _) = w
10545 .in_proj_qkv
10546 .graph_weight()
10547 .or_else(|| w.in_proj_qkv.graph_weight_descriptor())?;
10548 model = Some(m.clone());
10549 crate::gpu::GraphAttn::Gdn {
10550 qkv: gw(&w.in_proj_qkv)?,
10551 z: gw(&w.in_proj_z)?,
10552 a: gw(&w.in_proj_a)?,
10553 b: gw(&w.in_proj_b)?,
10554 out: gw(&w.out_proj)?,
10555 conv1d: &w.conv1d,
10556 a_log: &w.a_log,
10557 dt_bias: &w.dt_bias,
10558 norm: &w.norm,
10559 nv: cfg.num_v_heads,
10560 nk: cfg.num_k_heads,
10561 dk: cfg.key_head_dim,
10562 dv: cfg.value_head_dim,
10563 kk: cfg.conv_kernel,
10564 cpu_state: &self.kv_cache.layers[self.phys_layer(li)].linear_state,
10565 }
10566 }
10567 _ => return None,
10568 };
10569 layers.push(crate::gpu::GraphLayer {
10570 input_norm: &lw.input_norm,
10571 attn,
10572 post_norm: &lw.post_norm,
10573 ffn: gffn,
10574 });
10575 }
10576 Some((layers, model?))
10577 })();
10578 let Some((layers, model)) = built else {
10579 {
10580 use std::sync::atomic::{AtomicBool, Ordering};
10581 static SAID: AtomicBool = AtomicBool::new(false);
10582 if !SAID.swap(true, Ordering::Relaxed) {
10583 tracing::warn!("batch graph: BUILDER refused (layer weights/kinds)");
10584 }
10585 }
10586 return crate::gpu::BatchGraphOutcome::Declined;
10587 };
10588 if std::env::var("CMF_GRAPH_SPEC_TIME").is_ok() {
10589 eprintln!("batch-build: {:.1} ms", _tb.elapsed().as_secs_f64() * 1e3);
10590 }
10591 crate::gpu::forward_batch_graph(
10592 &model,
10593 self.graph_kv_id,
10594 &layers,
10595 &self.inv_freq,
10596 hiddens,
10597 nh,
10598 nkv,
10599 hd,
10600 rd,
10601 self.hidden_size,
10602 self.intermediate_size,
10603 positions,
10604 self.kv_cache.max_seq_len,
10605 gemma,
10606 self.rms_eps as f32,
10607 self.attn_scale,
10608 k,
10609 &(0..self.num_layers)
10610 .map(|li| self.kv_cache.layers[self.phys_layer(li)].o1_views())
10611 .collect::<Vec<_>>(),
10612 self.o1_epoch,
10613 spec,
10614 )
10615 }
10616
10617 fn draft_probe() -> bool {
10621 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
10622 *ON.get_or_init(|| std::env::var("CMF_DSV4_DRAFT_PROBE").is_ok_and(|v| v != "0"))
10623 }
10624
10625 #[cfg(feature = "gpu")]
10637 fn dsv4_spec_on() -> bool {
10638 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
10639 *ON.get_or_init(|| {
10640 if let Ok(v) = std::env::var("CMF_DSV4_SPEC_RUN") {
10644 return v != "0";
10645 }
10646 std::env::var("CMF_DSV4_SPEC")
10653 .map(|v| v != "0")
10654 .unwrap_or_else(|_| {
10655 crate::gpu_wgpu::DRAFT_RESERVE.load(std::sync::atomic::Ordering::Relaxed) > 0
10656 })
10657 })
10658 }
10659
10660 #[cfg(feature = "gpu")]
10667 fn dsv4_spec_step(
10668 &mut self,
10669 tip_token: u32,
10670 t_next: u32,
10671 next_pos: usize,
10672 max_extra: usize,
10673 drafted: &mut usize,
10674 accepted_ctr: &mut usize,
10675 ) -> Option<(Vec<u32>, usize)> {
10676 let t_all = std::time::Instant::now();
10677 if std::env::var("CMF_DSV4_SPEC_TIME").is_ok() {
10678 thread_local! {
10679 static LAST: std::cell::Cell<Option<std::time::Instant>> =
10680 const { std::cell::Cell::new(None) };
10681 }
10682 LAST.with(|l| {
10683 if let Some(prev) = l.get() {
10684 eprintln!(
10685 "между раундами {:.1} мс",
10686 prev.elapsed().as_secs_f64() * 1e3
10687 );
10688 }
10689 l.set(Some(std::time::Instant::now()));
10690 });
10691 }
10692 if std::env::var("CMF_DSV4_SPEC_DEBUG").is_ok() {
10693 eprintln!("spec_step: вход pos={next_pos}");
10694 }
10695 let n_layers = self.dsv4.as_ref().map(|b| b.1.len())?;
10696 let cfg = self.dsv4.as_ref().map(|b| b.2)?;
10697 if self.dspark.is_none() {
10699 let t = crate::dsv4::dspark_targets(&self.dsv4_mtp, &cfg, n_layers);
10700 if t.is_empty() {
10701 return None;
10702 }
10703 crate::dsv4::dspark_arm(&t, cfg.dim);
10704 self.dspark = Some(crate::dsv4::DsparkState::new(
10705 self.dsv4_mtp.len(),
10706 &cfg,
10707 t.len(),
10708 ));
10709 }
10710 let targets = crate::dsv4::dspark_targets(&self.dsv4_mtp, &cfg, n_layers);
10711 let pack = crate::dsv4::dspark_pack_get(&self.dsv4_mtp, &cfg);
10712 if pack.is_none() && std::env::var("CMF_DSV4_SPEC_DEBUG").is_ok() {
10713 eprintln!("spec_step: пак не построился (targets {targets:?})");
10714 }
10715 let pack = pack?;
10716 let block = crate::dsv4::dspark_block();
10717 let b_box = self.dsv4.as_mut()?;
10718 let (g, layers, st) = (&b_box.0, &b_box.1, &mut b_box.3);
10719 let ds = self.dspark.as_mut()?;
10720 let dbg = std::env::var("CMF_DSV4_SPEC_DEBUG").is_ok();
10723 if !crate::dsv4::dspark_take(&mut ds.main_hidden) && !ds.have_hidden {
10724 if dbg {
10725 eprintln!("spec_step: нет захвата");
10726 }
10727 return None;
10728 }
10729 ds.have_hidden = true;
10730 let tip_pos = next_pos.checked_sub(1)?;
10731 let draft_started = std::time::Instant::now();
10732 let mut conf = Vec::new();
10733 let props = crate::dsv4::dspark_draft_gpu(
10734 g,
10735 &self.dsv4_mtp,
10736 &cfg,
10737 ds,
10738 pack,
10739 st.kv_id,
10740 tip_token,
10741 tip_pos,
10742 self.pool.as_deref(),
10743 &mut conf,
10744 );
10745 self.dspark_draft_ns += draft_started.elapsed().as_nanos();
10746 *drafted += block;
10747 if props.is_empty() || props[0] != t_next {
10748 if dbg {
10749 eprintln!(
10750 "spec_step: черновик {} (props0={:?} t_next={t_next})",
10751 if props.is_empty() {
10752 "пуст"
10753 } else {
10754 "мимо"
10755 },
10756 props.first()
10757 );
10758 }
10759 return None;
10760 }
10761 let mut k_verify = crate::dsv4::dspark_verify_k()
10768 .min(props.len())
10769 .min(max_extra.saturating_add(1));
10770 let conf_min = {
10776 static M: std::sync::OnceLock<f32> = std::sync::OnceLock::new();
10777 *M.get_or_init(|| {
10778 std::env::var("CMF_DSPARK_CONF_MIN")
10779 .ok()
10780 .and_then(|v| v.parse().ok())
10781 .unwrap_or(0.0)
10782 })
10783 };
10784 if conf_min > 0.0 && conf.len() >= props.len() {
10785 let mut keep = 1usize;
10786 while keep < k_verify && conf.get(keep).copied().unwrap_or(0.0) >= conf_min {
10787 keep += 1;
10788 }
10789 k_verify = k_verify.min(keep.max(2));
10790 }
10791 if k_verify < 2 {
10792 return None;
10793 }
10794 let mut fed = Vec::with_capacity(k_verify);
10795 fed.push(t_next);
10796 fed.extend_from_slice(&props[1..k_verify]);
10797 let mut argmax = Vec::new();
10798 let mut logits_all = Vec::new();
10799 let mut walked = Vec::new();
10800 let txn = crate::dsv4::dsv4_verify_chunk(
10801 g,
10802 layers,
10803 &cfg,
10804 st,
10805 &fed,
10806 next_pos,
10807 &self.inv_freq,
10808 self.pool.as_deref(),
10809 &targets,
10810 &mut argmax,
10811 &mut logits_all,
10812 &mut walked,
10813 );
10814 if txn.is_none() && dbg {
10815 eprintln!("spec_step: verify отказал");
10816 }
10817 let txn = txn?;
10818 let spec_gpu_end = txn.gpu_end;
10819 let b = fed.len();
10820 let mut accepted = 1usize;
10821 while accepted < b && fed[accepted] == argmax[accepted - 1] {
10822 accepted += 1;
10823 }
10824 if std::env::var("CMF_DSV4_SPEC_FORCE_REJECT").is_ok_and(|v| v != "0") {
10829 accepted = 1;
10830 }
10831 if std::env::var("CMF_DSV4_SPEC_TRACE").is_ok() {
10832 eprintln!("spec@{next_pos}: fed={fed:?} argmax={argmax:?} accepted={accepted}");
10833 }
10834 let t_fin = std::time::Instant::now();
10835 if !crate::dsv4::dsv4_spec_finish(
10836 g,
10837 layers,
10838 &cfg,
10839 st,
10840 txn,
10841 accepted,
10842 &fed,
10843 &self.inv_freq,
10844 self.pool.as_deref(),
10845 ) {
10846 tracing::warn!("dsv4: спекулятивный откат не удался — состояние подозрительно");
10847 return None;
10848 }
10849 if std::env::var("CMF_DSV4_SPEC_TIME").is_ok() {
10850 eprintln!(
10851 "finish(k={accepted}): {:.1} мс",
10852 t_fin.elapsed().as_secs_f64() * 1e3
10853 );
10854 }
10855 *accepted_ctr += accepted - 1;
10856 let (hc, dim) = (cfg.hc_mult, cfg.dim);
10861 let dev_caps: Vec<usize> = targets
10866 .iter()
10867 .copied()
10868 .filter(|&t| t < spec_gpu_end)
10869 .collect();
10870 let mut caps_all = vec![0.0f32; dev_caps.len() * b * hc * dim];
10871 if !crate::gpu_wgpu::dsv4_spec_cap_read_all(b, dev_caps.len(), hc * dim, &mut caps_all) {
10872 return None;
10873 }
10874 for t in 0..accepted {
10875 let tip = t + 1 == accepted;
10876 for (slot, &tl) in targets.iter().enumerate() {
10877 if let Some(di) = dev_caps.iter().position(|&d| d == tl) {
10878 let lo = (di * b + t) * hc * dim;
10879 crate::dsv4::dspark_capture(
10880 &caps_all[lo..lo + hc * dim],
10881 &cfg,
10882 slot,
10883 &mut ds.main_hidden,
10884 );
10885 } else if tip
10886 && crate::dsv4::dspark_peek_slot(slot, dim, {
10887 let lo = slot * dim;
10888 &mut ds.main_hidden[lo..lo + dim]
10889 })
10890 {
10891 } else {
10896 crate::dsv4::dspark_capture(
10900 &walked[t * hc * dim..(t + 1) * hc * dim],
10901 &cfg,
10902 slot,
10903 &mut ds.main_hidden,
10904 );
10905 }
10906 }
10907 crate::dsv4::dspark_ring_append(
10908 g,
10909 &self.dsv4_mtp,
10910 &cfg,
10911 ds,
10912 next_pos + t,
10913 self.pool.as_deref(),
10914 );
10915 }
10916 let row = logits_all[(accepted - 1) * cfg.vocab..accepted * cfg.vocab].to_vec();
10917 self.graph_logits = Some(row);
10918 if std::env::var("CMF_DSV4_TRUNK_PICK_DUMP").is_ok() {
10923 crate::dsv4::trunk_freq_note(&crate::dsv4::pick_tally_take());
10924 crate::dsv4::pick_tally_arm();
10925 }
10926 if std::env::var("CMF_DSV4_SPEC_TIME").is_ok() {
10927 eprintln!(
10928 "spec_step total {:.1} мс (k={accepted})",
10929 t_all.elapsed().as_secs_f64() * 1e3
10930 );
10931 }
10932 Some((fed[1..accepted].to_vec(), next_pos + accepted))
10933 }
10934
10935 fn dspark_probe(&mut self, position: usize, token_id: u32) {
10936 if self.dsv4_mtp.is_empty() || !Self::draft_probe() {
10937 return;
10938 }
10939 let trunk_now = crate::dsv4::pick_tally_take();
10941 crate::dsv4::trunk_freq_note(&trunk_now);
10942 if !trunk_now.is_empty() {
10943 self.dspark_trunk_picks.push(trunk_now);
10944 let keep = crate::dsv4::dspark_block();
10945 if self.dspark_trunk_picks.len() > keep {
10946 self.dspark_trunk_picks.remove(0);
10947 }
10948 }
10949 for p in std::mem::take(&mut self.dspark_pending) {
10952 let Some(i) = position.checked_sub(p.0 + 1) else {
10953 continue;
10954 };
10955 let mut p = p;
10956 if i < p.1.len() {
10957 if p.2 && p.1[i] == token_id {
10958 p.3 = i + 1;
10959 } else {
10960 p.2 = false;
10961 }
10962 if i + 1 < p.1.len() {
10963 self.dspark_pending.push(p);
10964 continue;
10965 }
10966 }
10967 self.dspark_hist.push(p.3);
10968 self.dspark_real.push(token_id);
10969 }
10970 let Some(b) = &mut self.dsv4 else { return };
10971 let (g, layers, cfg) = (&b.0, &b.1, b.2);
10972 let n_layers = layers.len();
10973 if self.dspark.is_none() {
10974 let t = crate::dsv4::dspark_targets(&self.dsv4_mtp, &cfg, n_layers);
10975 if t.is_empty() {
10976 return;
10977 }
10978 eprintln!(
10979 "DSpark: захват со слоёв {t:?}, блок {}",
10980 crate::dsv4::dspark_block()
10981 );
10982 crate::dsv4::dspark_arm(&t, cfg.dim);
10983 self.dspark = Some(crate::dsv4::DsparkState::new(
10984 self.dsv4_mtp.len(),
10985 &cfg,
10986 t.len(),
10987 ));
10988 }
10989 let ds = self.dspark.as_mut().unwrap();
10990 if !crate::dsv4::dspark_take(&mut ds.main_hidden) {
10991 return; }
10993 let mut conf = Vec::new();
10994 crate::dsv4::pick_tally_arm();
10995 let draft_started = std::time::Instant::now();
11000 #[cfg(feature = "gpu")]
11001 let gpu_draft = crate::dsv4::dspark_gpu_on();
11002 #[cfg(not(feature = "gpu"))]
11003 let gpu_draft = false;
11004 let props = if gpu_draft {
11005 #[cfg(feature = "gpu")]
11006 {
11007 let kv_id = b.3.kv_id;
11008 match crate::dsv4::dspark_pack_get(&self.dsv4_mtp, &cfg) {
11009 Some(pk) => crate::dsv4::dspark_draft_gpu(
11010 g,
11011 &self.dsv4_mtp,
11012 &cfg,
11013 ds,
11014 pk,
11015 kv_id,
11016 token_id,
11017 position,
11018 self.pool.as_deref(),
11019 &mut conf,
11020 ),
11021 None => Vec::new(),
11022 }
11023 }
11024 #[cfg(not(feature = "gpu"))]
11025 Vec::new()
11026 } else {
11027 crate::gpu::cpu_scope(|| {
11028 crate::dsv4::dspark_draft(
11029 g,
11030 &self.dsv4_mtp,
11031 &cfg,
11032 ds,
11033 token_id,
11034 position,
11035 self.pool.as_deref(),
11036 &mut conf,
11037 )
11038 })
11039 };
11040 self.dspark_draft_ns += draft_started.elapsed().as_nanos();
11041 let draft_picks = crate::dsv4::pick_tally_take();
11042 crate::dsv4::dspark_freq_note(&draft_picks);
11043 crate::dsv4::pick_tally_arm();
11046 if !props.is_empty() {
11047 let (tu, tt) = {
11051 let flat: Vec<(usize, Vec<usize>)> = self
11052 .dspark_trunk_picks
11053 .iter()
11054 .flat_map(|v| v.iter().cloned())
11055 .collect();
11056 let mut per: std::collections::HashMap<usize, Vec<usize>> =
11058 std::collections::HashMap::new();
11059 for (li, picks) in flat {
11060 per.entry(li).or_default().extend(picks);
11061 }
11062 let n = per.len().max(1);
11063 let mut u = 0usize;
11064 let mut t = 0usize;
11065 for (_, v) in per {
11066 t += v.len();
11067 u += v.iter().collect::<std::collections::HashSet<_>>().len();
11068 }
11069 (u / n, t / n)
11070 };
11071 let (du, dt) = crate::dsv4::tally_unique(&draft_picks);
11072 self.dspark_exp.push((tu, tt, du, dt));
11073 self.dspark_pending.push((position, props, true, 0));
11074 }
11075 if self.dspark_hist.len() >= 8 && self.dspark_hist.len() % 8 == 0 {
11076 let n = self.dspark_hist.len() as f32;
11077 let mean: f32 = self.dspark_hist.iter().sum::<usize>() as f32 / n;
11078 let block = crate::dsv4::dspark_block();
11079 let mut at = vec![0usize; block + 1];
11080 for &k in &self.dspark_hist {
11081 at[k] += 1;
11082 }
11083 let mut surv = Vec::with_capacity(block);
11085 for i in 1..=block {
11086 let k = at[i..].iter().sum::<usize>() as f32 / n;
11087 surv.push(format!("{k:.2}"));
11088 }
11089 let distinct = self
11090 .dspark_real
11091 .iter()
11092 .collect::<std::collections::HashSet<_>>()
11093 .len();
11094 let (tu, tt, du, dt) = self.dspark_exp.iter().fold((0, 0, 0, 0), |a, b| {
11095 (a.0 + b.0, a.1 + b.1, a.2 + b.2, a.3 + b.3)
11096 });
11097 let m = self.dspark_exp.len().max(1);
11098 eprintln!(
11099 "DSpark: черновиков {}, принято в среднем {mean:.2} из {block} \
11100 (токенов за проход {:.2}), распределение {at:?}, выживание [{}]",
11101 self.dspark_hist.len(),
11102 mean + 1.0,
11103 surv.join(" ")
11104 );
11105 eprintln!(
11106 "DSpark: разных токенов {distinct} из {} (вырожденность), \
11107 эксперты ствол {}/{} на слой за {block} токенов, \
11108 черновик {}/{} за блок, draft {:.2} мс/блок",
11109 self.dspark_real.len(),
11110 tu / m,
11111 tt / m,
11112 du / m,
11113 dt / m,
11114 self.dspark_draft_ns as f64 / self.dspark_exp.len().max(1) as f64 / 1e6
11115 );
11116 }
11117 }
11118
11119 fn forward_layers_upto(
11120 &mut self,
11121 hidden: &[f32],
11122 position: usize,
11123 task_mask: Option<&TaskMask>,
11124 upto: Option<usize>,
11125 ) -> Vec<f32> {
11126 if let Some(plan) = self.gpu_plan.clone() {
11132 if upto.is_none() && plan.len() > 1 {
11133 let mut h = hidden.to_vec();
11134 for &(dev, from, upto_incl) in plan.iter() {
11135 h = crate::gpu::with_device(dev, || {
11136 self.forward_layers_span(&h, position, task_mask, from, Some(upto_incl))
11137 });
11138 }
11139 return h;
11140 }
11141 }
11142 self.forward_layers_span(hidden, position, task_mask, 0, upto)
11143 }
11144
11145 pub fn set_gpu_plan(&mut self, devices: Option<&[usize]>) -> Result<(), String> {
11150 self.set_gpu_plan_at(devices, None)
11151 }
11152
11153 pub fn set_gpu_plan_at(
11157 &mut self,
11158 devices: Option<&[usize]>,
11159 at: Option<usize>,
11160 ) -> Result<(), String> {
11161 let Some(devs) = devices.filter(|d| d.len() > 1) else {
11162 self.gpu_plan = None;
11163 return Ok(());
11164 };
11165 self.split_supported()?;
11166 let n = self.num_layers;
11167 if devs.len() > n {
11168 return Err(format!("{} devices for {n} layers", devs.len()));
11169 }
11170 if let Some(k) = at {
11171 if k == 0 || k >= n {
11172 return Err(format!("split at {k}: the model has {n} layers"));
11173 }
11174 if devs.len() == 2 {
11175 self.gpu_plan = Some(std::sync::Arc::new(vec![
11176 (devs[0], 0, k - 1),
11177 (devs[1], k, n - 1),
11178 ]));
11179 return Ok(());
11180 }
11181 return Err(format!(
11182 "an explicit split point takes exactly 2 devices, got {}",
11183 devs.len()
11184 ));
11185 }
11186 let per = n.div_ceil(devs.len());
11187 let mut plan = Vec::with_capacity(devs.len());
11188 let mut from = 0usize;
11189 for &d in devs {
11190 if from >= n {
11191 break;
11192 }
11193 let upto = (from + per - 1).min(n - 1);
11194 plan.push((d, from, upto));
11195 from = upto + 1;
11196 }
11197 self.gpu_plan = Some(std::sync::Arc::new(plan));
11198 Ok(())
11199 }
11200
11201 pub fn gpu_plan(&self) -> Option<Vec<(usize, usize, usize)>> {
11203 self.gpu_plan.as_ref().map(|p| p.as_ref().clone())
11204 }
11205
11206 fn forward_layers_span(
11212 &mut self,
11213 hidden: &[f32],
11214 position: usize,
11215 task_mask: Option<&TaskMask>,
11216 from: usize,
11217 upto: Option<usize>,
11218 ) -> Vec<f32> {
11219 debug_assert!(
11220 from == 0
11221 || (self.dsv4.is_none()
11222 && self.dsv41.is_none()
11223 && self.qwen4_exp.is_none()
11224 && self.g3n.is_none())
11225 );
11226 #[cfg(target_os = "macos")]
11232 if !crate::gpu_metal::wait_replay() {
11233 self.fail_metal_graph("the pending async replay failed before a plain forward");
11234 return vec![0.0; self.hidden_size];
11235 }
11236 if let Some(b) = &mut self.qwen4_exp {
11237 let _ = (task_mask, upto);
11238 let token_id = hidden.first().copied().unwrap_or(0.0) as u32;
11239 let mut logits = Vec::new();
11240 crate::qwen4_exp::forward_token(
11241 &b.0,
11242 &b.1,
11243 &b.2,
11244 &mut b.3,
11245 token_id,
11246 position,
11247 &self.inv_freq,
11248 self.pool.as_deref(),
11249 &mut logits,
11250 true,
11251 );
11252 self.graph_logits = Some(logits);
11253 return vec![0.0; self.hidden_size];
11254 }
11255 if let Some(b) = &mut self.dsv4 {
11261 let _ = (task_mask, upto);
11262 let token_id = hidden.first().copied().unwrap_or(0.0) as u32;
11263 let (g, layers, cfg, st) = (&b.0, &b.1, b.2, &mut b.3);
11264 st.pos = position;
11265 let mut logits = Vec::new();
11266 crate::dsv4::forward_token(
11267 g,
11268 layers,
11269 &cfg,
11270 st,
11271 token_id,
11272 &self.inv_freq,
11273 self.pool.as_deref(),
11274 &mut logits,
11275 );
11276 self.graph_logits = Some(logits);
11277 self.dspark_probe(position, token_id);
11278 return vec![0.0; self.hidden_size];
11281 }
11282 if let Some(b) = &mut self.dsv41 {
11284 let _ = (task_mask, upto);
11285 let token_id = hidden.first().copied().unwrap_or(0.0) as u32;
11286 let mut logits = Vec::new();
11287 crate::dsv41::forward_token(
11288 &b.0,
11289 &b.1,
11290 &b.2,
11291 &mut b.3,
11292 token_id,
11293 position,
11294 self.pool.as_deref(),
11295 &mut logits,
11296 );
11297 self.graph_logits = Some(logits);
11298 return vec![0.0; self.hidden_size];
11299 }
11300 if let Some(b) = &self.g3n {
11303 let _ = (task_mask, upto);
11304 return crate::g3n::g3n_forward(
11305 &b.0,
11306 &b.1,
11307 hidden,
11308 position,
11309 &mut self.kv_cache.layers,
11310 self.num_heads,
11311 self.num_kv_heads,
11312 self.head_dim,
11313 self.pool.as_deref(),
11314 );
11315 }
11316 let mut h = hidden.to_vec();
11317 let (nh, _nkv, _hd, hs, _rd, eps) = (
11320 self.num_heads,
11321 self.num_kv_heads,
11322 self.head_dim,
11323 self.hidden_size,
11324 self.rotary_dim,
11325 self.rms_eps,
11326 );
11327 let pool = self.pool.clone();
11328 let graph_env = std::env::var("CMF_GPU_WGPU_GRAPH").ok();
11340 let graph_on = match graph_env.as_deref() {
11341 Some("0") => false,
11342 Some("prefill") => false, Some(_) => true,
11344 None => crate::gpu::wgpu_graph_default(),
11350 };
11351 let graph_trusted =
11352 graph_env.is_some() || crate::gpu::wgpu_graph_default() || self.gdn_cfg.is_some();
11353 let race_eligible = graph_on
11354 && upto.is_none()
11355 && task_mask.is_none()
11356 && from == 0
11357 && !crate::gpu::graph_unsupported();
11358 let mut tail_start = 0usize;
11359 if race_eligible && crate::gpu::graph_race_use_graph(graph_trusted) {
11360 let t_graph = std::time::Instant::now();
11361 let mut lg = Vec::new();
11362 let mut gl = 0usize;
11363 let built = self.try_token_graph_wgpu(hidden, position, &mut lg, &mut gl);
11364 let declined = built.is_none();
11365 let built = match built {
11366 Some(Ok(hh)) => Some(hh),
11367 Some(Err(())) => {
11368 self.clear_sequence_state();
11372 self.graph_failed
11373 .store(true, std::sync::atomic::Ordering::Relaxed);
11374 self.cancel
11375 .store(true, std::sync::atomic::Ordering::Relaxed);
11376 tracing::error!("token graph failed after admission; sequence state cleared");
11377 return vec![0.0; self.hidden_size];
11378 }
11379 None => None,
11380 };
11381 if declined && !self.o1_active() && self.attn_softcap == 0.0 {
11386 crate::gpu::graph_mark_unsupported();
11387 }
11388 graph_note(built.is_some(), gl, self.num_layers);
11389 if let Some(hh) = built {
11390 let dur = t_graph.elapsed();
11391 if std::env::var("CMF_GRAPH_PROF").is_ok() {
11392 eprintln!("graph-call: {:.2} ms total", dur.as_secs_f64() * 1000.0);
11393 }
11394 if gl > 0 && gl < self.num_layers {
11395 h = hh;
11401 tail_start = gl;
11402 } else if graph_trusted || !crate::gpu::graph_race_first_token_hopeless(dur) {
11403 if !graph_trusted {
11404 crate::gpu::graph_race_record(true, dur);
11405 }
11406 if !lg.is_empty() {
11407 lg.resize(self.vocab_size, 0.0);
11410 if let Some(c) = self.final_softcap {
11411 for l in lg.iter_mut() {
11412 *l = c * (*l / c).tanh();
11413 }
11414 }
11415 self.graph_logits = Some(lg);
11416 }
11417 return hh;
11418 }
11419 }
11425 }
11426 let span = from > 0 || upto.is_some();
11450 if span && graph_on && task_mask.is_none() && graph_trusted {
11451 let upto_excl = upto.map_or(self.num_layers, |u| u + 1);
11452 let mut lg = Vec::new();
11453 let mut gl = 0usize;
11454 let span_res =
11455 self.try_token_graph_wgpu_span(hidden, position, &mut lg, from, upto_excl, &mut gl);
11456 let span_res = match span_res {
11457 Some(Ok(hh)) => Some(hh),
11458 Some(Err(())) => {
11459 self.clear_sequence_state();
11460 self.graph_failed
11461 .store(true, std::sync::atomic::Ordering::Relaxed);
11462 self.cancel
11463 .store(true, std::sync::atomic::Ordering::Relaxed);
11464 tracing::error!(
11465 "span token graph failed after admission; sequence state cleared"
11466 );
11467 return vec![0.0; self.hidden_size];
11468 }
11469 None => None,
11470 };
11471 graph_note(span_res.is_some(), gl, upto_excl - from);
11472 if std::env::var("CMF_GPU_DEBUG").is_ok() {
11473 static SEEN: std::sync::atomic::AtomicU32 = std::sync::atomic::AtomicU32::new(0);
11477 if SEEN.fetch_add(1, std::sync::atomic::Ordering::Relaxed) < 4 {
11478 eprintln!(
11479 "span graph: covered {gl} of {} layers [{from}..{upto_excl}) res={}",
11480 upto_excl - from,
11481 span_res.is_some()
11482 );
11483 }
11484 }
11485 if let Some(hh) = span_res {
11486 if gl == upto_excl - from {
11487 if !lg.is_empty() {
11488 lg.resize(self.vocab_size, 0.0);
11489 if let Some(c) = self.final_softcap {
11490 for l in lg.iter_mut() {
11491 *l = c * (*l / c).tanh();
11492 }
11493 }
11494 self.graph_logits = Some(lg);
11495 }
11496 crate::gpu::set_layer(-1);
11497 return hh;
11498 }
11499 h = hh;
11501 tail_start = from + gl;
11502 }
11503 }
11504 let t_race_cpu = (race_eligible && !graph_trusted).then(std::time::Instant::now);
11505
11506 let _host_tail = (tail_start > from).then(crate::gpu::enter_cpu_scope);
11512 let automatic_gpu_prefix = self.automatic_gpu_prefix();
11513
11514 #[cfg(target_os = "macos")]
11515 let mut gpu_skip_until = 0usize;
11516 for li in tail_start.max(from)..self.num_layers {
11517 let _capacity_tail = automatic_gpu_prefix
11518 .filter(|&prefix| li >= prefix)
11519 .map(|_| crate::gpu::enter_cpu_scope());
11520 crate::gpu::set_layer(li as i64); if let Some(u) = upto {
11522 if li > u {
11523 break;
11524 }
11525 }
11526 if let Some(mask) = task_mask {
11527 if !mask.layer_alive(li) {
11528 continue; }
11530 }
11531 #[cfg(target_os = "macos")]
11535 {
11536 if li < gpu_skip_until {
11537 continue;
11538 }
11539 if task_mask.is_none() {
11540 let end = self.q1_graph_gpu(li, upto, position, &mut h);
11541 if self
11542 .graph_failed
11543 .load(std::sync::atomic::Ordering::Relaxed)
11544 {
11545 return vec![0.0; self.hidden_size];
11549 }
11550 if end > li {
11551 gpu_skip_until = end;
11552 if self.is_loop_end(end - 1) && end < self.num_layers {
11555 h = inference::rms_norm(
11556 &h,
11557 &self.weights.final_norm,
11558 self.rms_eps,
11559 self.norm_style,
11560 );
11561 }
11562 continue;
11563 }
11564 }
11565 }
11566
11567 let lw = &self.weights.layers[self.phys_layer(li)];
11568 if let Ok(tp) = std::env::var("CMF_TRACE_POS") {
11569 if tp.parse::<usize>().ok() == Some(position) {
11570 let n: f32 = h.iter().map(|x| x * x).sum::<f32>().sqrt();
11571 eprintln!(
11572 "TRACE pos {position} layer {li}: |h| = {n:.6} h0 {:.6} h1 {:.6}",
11573 h[0], h[1]
11574 );
11575 }
11576 }
11577 inference::rms_norm_into(
11580 &h,
11581 &lw.input_norm,
11582 self.rms_eps,
11583 self.norm_style,
11584 &mut self.ws.n1,
11585 );
11586
11587 let attn_out = match &lw.attn {
11588 AttnKind::Mla(w) => {
11589 let inv_freq_l = self.layer_inv_freq(li);
11590 let rs = self.layer_rope_scale(li);
11591 let eps = self.rms_eps;
11592 let pool = self.pool.clone();
11593 mla_attention(
11594 w,
11595 &self.ws.n1,
11596 &mut self.kv_cache.layers[li],
11597 position,
11598 &inv_freq_l,
11599 rs,
11600 eps,
11601 pool.as_deref(),
11602 )
11603 }
11604 AttnKind::Linear(w) => {
11605 let cfg = self.vmf_cfg.expect("linear layer without vmf_cfg");
11606 vmf_phase_forward(
11607 &self.ws.n1,
11608 w,
11609 &cfg,
11610 &mut self.kv_cache.layers[li].linear_state,
11611 self.pool.as_deref(),
11612 )
11613 }
11614 AttnKind::Kda(w) => {
11615 let cfg = self.kda_cfg.expect("kda layer without kda_cfg");
11616 crate::linear_core::kda_forward(
11617 &self.ws.n1,
11618 w,
11619 &cfg,
11620 &mut self.kv_cache.layers[li].linear_state,
11621 self.pool.as_deref(),
11622 )
11623 }
11624 AttnKind::LinearGdn(w) => {
11625 let cfg = self.gdn_cfg.expect("gdn layer without gdn_cfg");
11626 gdn_forward(
11627 &self.ws.n1,
11628 w,
11629 &cfg,
11630 &mut self.kv_cache.layers[li].linear_state,
11631 self.pool.as_deref(),
11632 )
11633 }
11634 AttnKind::ShortConv(w) => {
11635 let cfg = self
11636 .short_conv_cfg
11637 .expect("short-conv layer without short_conv_cfg");
11638 short_conv_forward(
11639 &self.ws.n1,
11640 w,
11641 &cfg,
11642 &mut self.kv_cache.layers[li].linear_state,
11643 self.pool.as_deref(),
11644 )
11645 }
11646 AttnKind::Full {
11647 wq,
11648 wk,
11649 wv,
11650 wo,
11651 q_norm,
11652 k_norm,
11653 output_gate,
11654 softplus_gate,
11655 bias,
11656 } if self.kv_cache.layers[li].o1_sealed() => {
11657 let inv_freq_l = self.layer_inv_freq(li);
11660 let (nkv_l, hd_l, rd_l) = self.layer_geom(li);
11661 let cfg = QwenAttnCfg {
11662 num_heads: self.layer_num_heads(li),
11663 num_kv_heads: nkv_l,
11664 head_dim: hd_l,
11665 hidden_size: hs,
11666 position,
11667 inv_freq: &inv_freq_l,
11668 rotary_dim: rd_l,
11669 scale: self.attn_scale,
11670 softcap: self.attn_softcap,
11671 window: None,
11672 v_norm: self.attn_v_norm,
11673 qk_norm_after_rope: self.qk_norm_after_rope,
11674 q_norm: q_norm.as_deref(),
11675 k_norm: k_norm.as_deref(),
11676 output_gate: *output_gate,
11677 softplus_gate: softplus_gate
11678 .as_ref()
11679 .map(|(gate, per_head)| (gate, *per_head)),
11680 rope_scale: self.layer_rope_scale(li),
11681 bias: bias
11682 .as_ref()
11683 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
11684 rms_eps: eps,
11685 norm_style: self.norm_style,
11686 pool: pool.as_deref(),
11687 };
11688 attention::qwen_attention_nystrom(
11689 &self.ws.n1,
11690 wq,
11691 wk,
11692 wv,
11693 wo,
11694 &mut self.kv_cache.layers[li],
11695 &cfg,
11696 )
11697 }
11698 AttnKind::Full {
11699 wq,
11700 wk,
11701 wv,
11702 wo,
11703 q_norm,
11704 k_norm,
11705 output_gate,
11706 softplus_gate,
11707 bias,
11708 } => 'attn: {
11709 if graph_on
11712 && !*output_gate
11713 && softplus_gate.is_none()
11714 && self.attention_heads_per_layer.is_none()
11715 && bias.is_none()
11716 && task_mask.is_none()
11717 {
11718 let inv_freq_l = self.layer_inv_freq(li);
11719 let (nkv_l, hd_l, rd_l) = self.layer_geom(li);
11720 let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
11721 if let (Some((gm, qi)), Some((_, ki)), Some((_, vi)), Some((_, oi))) = (
11722 wq.mapped_q1(),
11723 wk.mapped_q1(),
11724 wv.mapped_q1(),
11725 wo.mapped_q1(),
11726 ) {
11727 let gm = gm.clone();
11728 let mut out = vec![0f32; hs];
11729 let cache = &self.kv_cache.layers[li];
11730 if crate::gpu::attn_dropin(
11731 &gm,
11732 self.graph_kv_id,
11733 li,
11734 &self.ws.n1,
11735 qi,
11736 ki,
11737 vi,
11738 oi,
11739 q_norm.as_deref(),
11740 k_norm.as_deref(),
11741 self.qk_norm_after_rope,
11742 &inv_freq_l,
11743 nh,
11744 nkv_l,
11745 hd_l,
11746 rd_l,
11747 hs,
11748 position,
11749 self.kv_cache.max_seq_len,
11750 gemma,
11751 eps as f32,
11752 cache.k_heads(),
11753 cache.v_heads(),
11754 &mut out,
11755 ) {
11756 break 'attn out;
11757 }
11758 }
11759 }
11760 let masked = task_mask
11761 .map(|m| m.head_flags(li, self.num_heads).iter().any(|&a| !a))
11762 .unwrap_or(false);
11763 let f32_view = (wq.as_f32(), wk.as_f32(), wv.as_f32(), wo.as_f32());
11764 match (masked, f32_view) {
11765 (true, (Some(q), Some(k), Some(v), Some(o))) => {
11768 let active_heads = task_mask.unwrap().head_flags(li, self.num_heads);
11769 attention::multi_head_attention(
11770 &self.ws.n1,
11771 q,
11772 k,
11773 v,
11774 o,
11775 &mut self.kv_cache.layers[li],
11776 self.num_heads,
11777 self.num_kv_heads,
11778 self.head_dim,
11779 self.hidden_size,
11780 position,
11781 &active_heads,
11782 &self.inv_freq,
11783 )
11784 }
11785 (masked, _) => {
11786 if masked {
11787 tracing::warn!(
11788 "layer {li}: head mask on quantized weights not \
11789 supported yet — executing dense"
11790 );
11791 }
11792 let inv_freq_l = self.layer_inv_freq(li);
11793 let (nkv_l, hd_l, rd_l) = self.layer_geom(li);
11794 let cfg = QwenAttnCfg {
11795 num_heads: self.layer_num_heads(li),
11796 num_kv_heads: nkv_l,
11797 head_dim: hd_l,
11798 hidden_size: hs,
11799 position,
11800 inv_freq: &inv_freq_l,
11801 rotary_dim: rd_l,
11802 scale: self.attn_scale,
11803 softcap: self.attn_softcap,
11804 window: self.layer_window(li),
11805 v_norm: self.attn_v_norm,
11806 qk_norm_after_rope: self.qk_norm_after_rope,
11807 q_norm: q_norm.as_deref(),
11808 k_norm: k_norm.as_deref(),
11809 output_gate: *output_gate,
11810 softplus_gate: softplus_gate
11811 .as_ref()
11812 .map(|(gate, per_head)| (gate, *per_head)),
11813 rope_scale: self.layer_rope_scale(li),
11814 bias: bias
11815 .as_ref()
11816 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
11817 rms_eps: eps,
11818 norm_style: self.norm_style,
11819 pool: pool.as_deref(),
11820 };
11821 attention::qwen_attention(
11822 &self.ws.n1,
11823 wq,
11824 wk,
11825 wv,
11826 wo,
11827 &mut self.kv_cache.layers[li],
11828 &cfg,
11829 )
11830 }
11831 }
11832 }
11833 };
11834 let attn_out = match &self.weights.layers[self.phys_layer(li)].attn_out_norm {
11837 Some(w) => inference::rms_norm(&attn_out, w, self.rms_eps, self.norm_style),
11838 None => attn_out,
11839 };
11840 let lw = &self.weights.layers[self.phys_layer(li)];
11841 inference::add_rmsnorm_fused_into(
11842 &mut h,
11843 &attn_out,
11844 &lw.post_norm,
11845 self.rms_eps,
11846 self.norm_style,
11847 &mut self.ws.p1,
11848 );
11849 let mut attn_out = attn_out;
11850 attention::recycle_buf(&mut attn_out);
11851 let post_normed = &self.ws.p1;
11852
11853 let ffn_masked = task_mask
11854 .map(|m| m.ffn_active_count(li) < self.intermediate_size)
11855 .unwrap_or(false);
11856 let ffn_out = match (ffn_masked, &lw.ffn) {
11868 (_, FfnKind::Dense(d)) if !d.segs.is_empty() => {
11872 let row = task_mask
11873 .and_then(|tm| tm.ffn_masks.get(li))
11874 .map(|v| v.as_slice());
11875 tube_ffn(d, post_normed, 1, self.pool.as_deref(), row)
11876 }
11877 (true, FfnKind::Dense(d)) => {
11878 let tm = task_mask.unwrap();
11879 let alive = tm.ffn_active_count(li);
11880 let deep = alive * 2 <= self.intermediate_size;
11881 if deep && d.down_proj.sparse_col_ok() && !d.gate_proj.has_prism_contract() {
11882 let active = tm.ffn_active_indices(li);
11883 sparse_ffn_quant(
11884 d,
11885 post_normed,
11886 &active,
11887 self.hidden_size,
11888 self.pool.as_deref(),
11889 )
11890 } else if deep
11891 && let (Some(g), Some(u), Some(dn)) = (
11892 d.gate_proj.as_f32(),
11893 d.up_proj.as_f32(),
11894 d.down_proj.as_f32(),
11895 )
11896 {
11897 let active = tm.ffn_active_indices(li);
11898 inference::sparse_ffn_forward(
11899 post_normed,
11900 g,
11901 u,
11902 dn,
11903 self.hidden_size,
11904 self.intermediate_size,
11905 &active,
11906 self.pool.as_deref(),
11907 )
11908 } else {
11909 let row = tm.ffn_masks.get(li).map(|v| v.as_slice());
11910 dense_ffn_batch(d, post_normed, 1, self.pool.as_deref(), row)
11911 }
11912 }
11913 (true, FfnKind::Moe(m)) => {
11914 let allowed = task_mask.and_then(|tm| tm.expert_flags(li, m.experts.len()));
11918 ffn_forward(
11919 &lw.ffn,
11920 post_normed,
11921 self.pool.as_deref(),
11922 allowed.as_deref(),
11923 )
11924 }
11925 (true, FfnKind::DenseMoe(dm)) => dense_moe_ffn(
11926 dm,
11927 post_normed,
11928 &h,
11929 self.rms_eps,
11930 self.norm_style,
11931 self.pool.as_deref(),
11932 ),
11933 (false, _) => match &lw.ffn {
11934 FfnKind::DenseMoe(dm) => dense_moe_ffn(
11935 dm,
11936 post_normed,
11937 &h,
11938 self.rms_eps,
11939 self.norm_style,
11940 self.pool.as_deref(),
11941 ),
11942 _ => {
11943 let allowed = match (&lw.ffn, task_mask) {
11944 (FfnKind::Moe(m), Some(tm)) => tm.expert_flags(li, m.experts.len()),
11945 _ => None,
11946 };
11947 ffn_forward(
11948 &lw.ffn,
11949 post_normed,
11950 self.pool.as_deref(),
11951 allowed.as_deref(),
11952 )
11953 }
11954 },
11955 };
11956 let ffn_out = match &self.weights.layers[self.phys_layer(li)].ffn_out_norm {
11957 Some(w) => inference::rms_norm(&ffn_out, w, self.rms_eps, self.norm_style),
11958 None => ffn_out,
11959 };
11960 for (i, &f) in ffn_out.iter().enumerate() {
11961 h[i] += f;
11962 }
11963 let mut ffn_out = ffn_out;
11964 attention::recycle_buf(&mut ffn_out);
11965
11966 if let Some(sc) = self.weights.layers[self.phys_layer(li)].layer_scale {
11968 for v in h.iter_mut() {
11969 *v *= sc;
11970 }
11971 }
11972
11973 if self.is_loop_end(li) && li + 1 < self.num_layers {
11976 h = inference::rms_norm(
11977 &h,
11978 &self.weights.final_norm,
11979 self.rms_eps,
11980 self.norm_style,
11981 );
11982 }
11983
11984 if self.dyn_phi_layer == Some(li) {
11988 self.update_dyn_phi(&h);
11989 }
11990 }
11991 crate::gpu::set_layer(-1); if let Some(t) = t_race_cpu {
11993 crate::gpu::graph_race_record(false, t.elapsed());
11994 }
11995
11996 h
11997 }
11998
11999 fn update_dyn_phi(&mut self, h: &[f32]) {
12002 const A: f32 = 0.2;
12003 if self.dyn_phi_ema.len() != h.len() {
12004 self.dyn_phi_ema = vec![0.0; h.len()];
12005 self.dyn_phi_seen = 0;
12006 }
12007 if self.dyn_phi_seen == 0 {
12008 self.dyn_phi_ema.copy_from_slice(h);
12009 } else {
12010 for (e, &v) in self.dyn_phi_ema.iter_mut().zip(h) {
12011 *e = (1.0 - A) * *e + A * v;
12012 }
12013 }
12014 self.dyn_phi_seen += 1;
12015 }
12016
12017 pub fn dyn_phi(&self) -> &[f32] {
12019 &self.dyn_phi_ema
12020 }
12021
12022 pub fn set_dyn_phi_layer(&mut self, layer: Option<usize>) {
12024 self.dyn_phi_layer = layer;
12025 self.dyn_phi_ema.clear();
12026 self.dyn_phi_seen = 0;
12027 }
12028
12029 pub fn dynamic_skills(&self) -> Vec<(usize, String, usize)> {
12031 let Some(model) = &self.model else {
12032 return Vec::new();
12033 };
12034 model
12035 .header
12036 .skills
12037 .iter()
12038 .enumerate()
12039 .filter_map(|(i, sk)| {
12040 let ok = matches!(self.dyn_skill_layers.get(i), Some(Some(_)));
12041 let sel = sk.selection.as_ref()?;
12042 (ok).then(|| (i, sk.id.clone(), sel.phi_layer))
12043 })
12044 .collect()
12045 }
12046
12047 pub fn active_skill(&self) -> Option<usize> {
12049 self.dyn_active
12050 }
12051
12052 pub fn enable_dynamic_routing(&mut self) -> usize {
12057 use crate::swarm::{DynRouter, RoutableSkill};
12058 let Some(model) = self.model.clone() else {
12059 return 0;
12060 };
12061 if self.dyn_blend_loaded {
12064 tracing::warn!("dynamic routing unavailable on a blend-loaded pipeline");
12065 return 0;
12066 }
12067 if let Some(a) = self.dyn_active {
12071 if !matches!(self.dyn_skill_layers.get(a), Some(Some(_))) {
12072 tracing::warn!("loaded skill is not FFN-eligible — dynamic routing unavailable");
12073 return 0;
12074 }
12075 }
12076 let hidden = self.hidden_size;
12077 let mut skills = Vec::new();
12078 for (idx, id, _phi) in self.dynamic_skills() {
12079 if let Some(sel) = model.header.skills[idx].selection.as_ref() {
12080 if let Some(rs) = RoutableSkill::from_descriptor(idx, id, sel, hidden) {
12081 skills.push(rs);
12082 }
12083 }
12084 }
12085 if skills.is_empty() {
12086 return 0;
12087 }
12088 let phi = skills[0].phi_layer;
12090 if skills.iter().any(|s| s.phi_layer != phi) {
12091 tracing::warn!("routable skills disagree on phi_layer; using {phi}");
12092 }
12093 let n = skills.len();
12094 self.set_dyn_phi_layer(Some(phi));
12095 self.dyn_router = Some(DynRouter::new(skills));
12096 n
12097 }
12098
12099 pub fn route_switches(&self) -> Vec<(usize, Option<String>, Option<String>)> {
12101 self.dyn_router
12102 .as_ref()
12103 .map(|r| r.switches.clone())
12104 .unwrap_or_default()
12105 }
12106
12107 fn lm_head_forward(&self, hidden: &[f32]) -> Vec<f32> {
12110 let rows = self.weights.lm_head.rows();
12111 let mut logits = attention::take_buf(rows.min(self.vocab_size));
12112 self.weights
12113 .lm_head
12114 .matvec(hidden, &mut logits, self.pool.as_deref());
12115 logits.resize(self.vocab_size, 0.0);
12116 if let Some(m) = self.logit_multiplier {
12117 for l in logits.iter_mut() {
12118 *l *= m;
12119 }
12120 }
12121 if let Some(c) = self.final_softcap {
12122 for l in logits.iter_mut() {
12123 *l = c * (*l / c).tanh();
12124 }
12125 }
12126 if let Some(cm) = self.head_clusters.as_ref() {
12127 self.hierarchical_head_logprobs(hidden, cm, &mut logits);
12128 }
12129 logits
12130 }
12131
12132 fn hierarchical_head_logprobs(&self, hidden: &[f32], cm: &[f32], logits: &mut [f32]) {
12135 let h = hidden.len();
12136 let ncl = cm.len() / h.max(1);
12137 if ncl == 0 || logits.len() % ncl != 0 {
12138 return;
12139 }
12140 let cs = logits.len() / ncl;
12141 let mut lc = vec![0.0f32; ncl];
12143 for c in 0..ncl {
12144 let row = &cm[c * h..(c + 1) * h];
12145 let mut s = 0.0f32;
12146 for j in 0..h {
12147 s += row[j] * hidden[j];
12148 }
12149 lc[c] = s;
12150 }
12151 let mx = lc.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
12152 let lse: f32 = mx + lc.iter().map(|v| (v - mx).exp()).sum::<f32>().ln();
12153 for c in 0..ncl {
12154 let blk = &mut logits[c * cs..(c + 1) * cs];
12155 let bm = blk.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
12156 let bl: f32 = bm + blk.iter().map(|v| (v - bm).exp()).sum::<f32>().ln();
12157 let add = lc[c] - lse - bl;
12158 for v in blk.iter_mut() {
12159 *v += add;
12160 }
12161 }
12162 }
12163
12164 pub fn prefill_next_logits(&mut self, ids: &[u32], task_mask: Option<&TaskMask>) -> Vec<f32> {
12169 self.clear_sequence_state();
12170 crate::gpu::graph_race_begin_generation();
12174 if task_mask.is_none() {
12175 self.o1_begin();
12176 }
12177 let mut hidden = vec![0.0f32; self.hidden_size];
12178 for (pos, &id) in ids.iter().enumerate() {
12179 let emb = self.embed_single(id);
12180 hidden = self.forward_layers(&emb, pos, task_mask);
12181 }
12182 if let Err(err) = self.o1_seal_checked() {
12183 self.o1_fail(err);
12184 }
12185 inference::rms_norm_into(
12186 &hidden,
12187 &self.weights.final_norm,
12188 self.rms_eps,
12189 self.norm_style,
12190 &mut self.ws.n1,
12191 );
12192 self.lm_head_forward(&self.ws.n1)
12193 }
12194}
12195
12196pub fn create_test_pipeline(
12198 hidden_size: usize,
12199 intermediate_size: usize,
12200 num_heads: usize,
12201 num_kv_heads: usize,
12202 head_dim: usize,
12203 num_layers: usize,
12204 vocab_size: usize,
12205) -> Pipeline {
12206 let synth = |n: usize, salt: usize| -> Vec<f32> {
12209 (0..n)
12210 .map(|i| (((i * 31 + salt * 17 + 7) % 97) as f32 / 97.0 - 0.5) * 0.2)
12211 .collect()
12212 };
12213 let qt = |rows: usize, cols: usize, salt: usize| -> QTensor {
12214 QTensor::from_f32(synth(rows * cols, salt), rows, cols)
12215 };
12216 let layer_weights: Vec<LayerWeights> = (0..num_layers)
12217 .map(|li| LayerWeights {
12218 input_norm: vec![1.0; hidden_size],
12219 post_norm: vec![1.0; hidden_size],
12220 attn_out_norm: None,
12221 ffn_out_norm: None,
12222 layer_scale: None,
12223 ffn: FfnKind::Dense(DenseFfn {
12224 gate_proj: qt(intermediate_size, hidden_size, li * 10 + 5),
12225 up_proj: qt(intermediate_size, hidden_size, li * 10 + 6),
12226 down_proj: qt(hidden_size, intermediate_size, li * 10 + 7),
12227 act: Act::Silu,
12228 down_t: None,
12229 segs: Vec::new(),
12230 }),
12231 attn: AttnKind::Full {
12232 bias: None,
12233 wq: qt(num_heads * head_dim, hidden_size, li * 10 + 1),
12234 wk: qt(num_kv_heads * head_dim, hidden_size, li * 10 + 2),
12235 wv: qt(num_kv_heads * head_dim, hidden_size, li * 10 + 3),
12236 wo: qt(hidden_size, num_heads * head_dim, li * 10 + 4),
12237 q_norm: None,
12238 k_norm: None,
12239 output_gate: false,
12240 softplus_gate: None,
12241 },
12242 })
12243 .collect();
12244
12245 Pipeline::new(
12246 Tokenizer::byte_level(),
12247 PipelineWeights {
12248 embed_tokens: qt(vocab_size, hidden_size, 100),
12249 layers: layer_weights,
12250 lm_head: qt(vocab_size, hidden_size, 200),
12251 final_norm: vec![1.0; hidden_size],
12252 },
12253 hidden_size,
12254 intermediate_size,
12255 num_heads,
12256 num_kv_heads,
12257 head_dim,
12258 num_layers,
12259 num_layers, false, vocab_size,
12262 1e-6,
12263 10_000.0,
12264 NormStyle::Qwen,
12265 4096,
12266 SamplerConfig {
12267 seed: Some(42),
12268 ..Default::default()
12269 },
12270 )
12271}
12272
12273#[inline]
12278fn mask_bit(row: &[u8], j: usize) -> bool {
12279 (row.get(j >> 3).copied().unwrap_or(0) >> (j & 7)) & 1 != 0
12280}
12281
12282fn mask_gain() -> f32 {
12293 static G: std::sync::OnceLock<f32> = std::sync::OnceLock::new();
12294 *G.get_or_init(|| {
12295 std::env::var("CMF_FFN_MASK_GAIN")
12296 .ok()
12297 .and_then(|v| v.parse().ok())
12298 .unwrap_or(1.0)
12299 })
12300}
12301
12302fn zero_masked_cols(g: &mut [f32], rows: usize, inter: usize, row: &[u8]) {
12303 let fill = meanfill().and_then(|(i, v)| {
12306 let li = crate::gpu::cur_layer();
12307 (*i == inter && li >= 0).then(|| &v[li as usize * inter..(li as usize + 1) * inter])
12308 });
12309 for r in 0..rows {
12310 let base = r * inter;
12311 for (bi, &byte) in row.iter().enumerate() {
12312 if byte == 0xFF {
12313 continue;
12314 }
12315 let j0 = bi * 8;
12316 for bit in 0..8 {
12317 let j = j0 + bit;
12318 if j < inter && byte & (1 << bit) == 0 {
12319 g[base + j] = fill.map_or(0.0, |f| f[j]);
12320 }
12321 }
12322 }
12323 }
12324 let gain = mask_gain();
12325 if gain != 1.0 {
12326 for v in g[..rows * inter].iter_mut() {
12327 *v *= gain;
12328 }
12329 }
12330}
12331
12332#[inline]
12334fn tube_bit(row: Option<&[u8]>, i: usize) -> bool {
12335 row.is_none_or(|r| mask_bit(r, i))
12336}
12337
12338fn all_bits_on(row: &[u8], n: usize) -> bool {
12341 (0..n).all(|i| mask_bit(row, i))
12342}
12343
12344fn tube_topk() -> usize {
12352 static K: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
12353 *K.get_or_init(|| {
12354 std::env::var("CMF_TUBE_TOPK")
12355 .ok()
12356 .and_then(|v| v.parse().ok())
12357 .unwrap_or(0)
12358 })
12359}
12360
12361fn tube_score_oracle() -> bool {
12362 static O: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
12363 *O.get_or_init(|| std::env::var("CMF_TUBE_SCORE").is_ok_and(|v| v == "oracle"))
12364}
12365
12366fn tube_ffn_routed(
12373 d: &DenseFfn,
12374 xs: &[f32],
12375 b: usize,
12376 pool: Option<&Pool>,
12377 mask_row: Option<&[u8]>,
12378 k: usize,
12379) -> Vec<f32> {
12380 let hidden = d.down_proj.rows();
12381 let core = d.gate_proj.rows();
12382 let core_full = mask_row.is_none_or(|r| all_bits_on(r, core));
12383 let mut out = match (b, core_full, mask_row) {
12384 (1, true, _) => dense_ffn(d, xs, pool),
12385 (1, false, Some(row)) => dense_ffn_masked(d, xs, pool, row),
12386 (_, true, _) => dense_ffn_batch(d, xs, b, pool, None),
12387 (_, false, row) => dense_ffn_batch(d, xs, b, pool, row),
12388 };
12389 let cand: Vec<usize> = (0..d.segs.len())
12390 .filter(|&i| tube_bit(mask_row, d.segs[i].start))
12391 .collect();
12392 if cand.is_empty() {
12393 return out;
12394 }
12395 let oracle = tube_score_oracle();
12399 let mut acts: Vec<Vec<f32>> = Vec::with_capacity(cand.len());
12400 let mut scores = vec![0f32; b * cand.len()];
12401 for (ci, &i) in cand.iter().enumerate() {
12402 let seg = &d.segs[i];
12403 let w = seg.width;
12404 let mut g = vec![0.0f32; b * w];
12405 if b == 1 {
12406 seg.gate.matvec(xs, &mut g, pool);
12407 } else {
12408 seg.gate.matmat(xs, b, &mut g, pool);
12409 }
12410 for v in g.iter_mut() {
12411 *v = Act::Silu.combine(*v, 1.0);
12412 }
12413 if !oracle {
12414 for t in 0..b {
12415 scores[t * cand.len() + ci] =
12416 g[t * w..(t + 1) * w].iter().map(|v| v * v).sum::<f32>();
12417 }
12418 }
12419 if oracle || b > 1 {
12420 let mut u = vec![0.0f32; b * w];
12421 if b == 1 {
12422 seg.up.matvec(xs, &mut u, pool);
12423 } else {
12424 seg.up.matmat(xs, b, &mut u, pool);
12425 }
12426 for (a, &v) in g.iter_mut().zip(u.iter()) {
12427 *a *= v;
12428 }
12429 if oracle {
12430 for t in 0..b {
12431 scores[t * cand.len() + ci] =
12432 g[t * w..(t + 1) * w].iter().map(|v| v * v).sum::<f32>();
12433 }
12434 }
12435 }
12436 acts.push(g);
12437 }
12438 let keep = k.min(cand.len());
12440 let mut scratch: Vec<f32> = Vec::new();
12441 for t in 0..b {
12442 let mut sc: Vec<(f32, usize)> = (0..cand.len())
12443 .map(|ci| (scores[t * cand.len() + ci], ci))
12444 .collect();
12445 sc.sort_unstable_by(|x, y| y.0.total_cmp(&x.0));
12446 let mut alive = vec![false; cand.len()];
12447 for &(_, ci) in sc.iter().take(keep) {
12448 alive[ci] = true;
12449 }
12450 if b > 1 {
12451 for (ci, a) in acts.iter_mut().enumerate() {
12452 if !alive[ci] {
12453 let w = d.segs[cand[ci]].width;
12454 a[t * w..(t + 1) * w].fill(0.0);
12455 }
12456 }
12457 } else {
12458 for (ci, &i) in cand.iter().enumerate() {
12462 if !alive[ci] {
12463 continue;
12464 }
12465 let seg = &d.segs[i];
12466 let w = seg.width;
12467 let g = &mut acts[ci];
12468 if !tube_score_oracle() {
12469 scratch.clear();
12470 scratch.resize(w, 0.0);
12471 seg.up.matvec(xs, &mut scratch, pool);
12472 for (a, &v) in g.iter_mut().zip(scratch.iter()) {
12473 *a *= v;
12474 }
12475 }
12476 let mut acc = vec![0.0f32; hidden];
12477 seg.down.matvec(g, &mut acc, pool);
12478 for (o, a) in out.iter_mut().zip(&acc) {
12479 *o += *a;
12480 }
12481 }
12482 }
12483 }
12484 if b > 1 {
12485 for (ci, &i) in cand.iter().enumerate() {
12486 let seg = &d.segs[i];
12487 let mut acc = vec![0.0f32; b * hidden];
12488 seg.down.matmat(&acts[ci], b, &mut acc, pool);
12489 for (o, a) in out.iter_mut().zip(&acc) {
12490 *o += *a;
12491 }
12492 }
12493 }
12494 out
12495}
12496
12497fn tube_ffn(
12503 d: &DenseFfn,
12504 xs: &[f32],
12505 b: usize,
12506 pool: Option<&Pool>,
12507 mask_row: Option<&[u8]>,
12508) -> Vec<f32> {
12509 if tube_topk() > 0 {
12510 return tube_ffn_routed(d, xs, b, pool, mask_row, tube_topk());
12511 }
12512 let hidden = d.down_proj.rows();
12513 let core = d.gate_proj.rows();
12514 let core_full = mask_row.is_none_or(|r| all_bits_on(r, core));
12515 let mut out = match (b, core_full, mask_row) {
12516 (1, true, _) => dense_ffn(d, xs, pool),
12517 (1, false, Some(row)) => dense_ffn_masked(d, xs, pool, row),
12518 (_, true, _) => dense_ffn_batch(d, xs, b, pool, None),
12519 (_, false, row) => dense_ffn_batch(d, xs, b, pool, row),
12520 };
12521 TUBE_SCRATCH.with(|sc| {
12522 let mut sc = sc.borrow_mut();
12523 let [g, u, acc] = &mut *sc;
12524 for seg in &d.segs {
12525 if !tube_bit(mask_row, seg.start) {
12526 continue;
12527 }
12528 let w = seg.width;
12529 g.resize(b * w, 0.0);
12530 if b == 1
12531 && d.act == Act::Silu
12532 && QTensor::matvec_silu_mul(&seg.gate, &seg.up, xs, g, pool)
12533 {
12534 } else {
12536 u.resize(b * w, 0.0);
12537 if b == 1 {
12538 QTensor::matvec_many([&seg.gate, &seg.up], xs, [g, u], pool);
12539 } else {
12540 seg.gate.matmat(xs, b, g, pool);
12541 seg.up.matmat(xs, b, u, pool);
12542 }
12543 for i in 0..b * w {
12544 g[i] = d.act.combine(g[i], u[i]);
12545 }
12546 }
12547 acc.resize(b * hidden, 0.0);
12548 acc.fill(0.0);
12549 if b == 1 {
12550 seg.down.matvec(g, acc, pool);
12551 } else {
12552 seg.down.matmat(g, b, acc, pool);
12553 }
12554 for (o, a) in out.iter_mut().zip(acc.iter()) {
12555 *o += *a;
12556 }
12557 }
12558 out
12559 })
12560}
12561
12562thread_local! {
12563 static TUBE_SCRATCH: std::cell::RefCell<[Vec<f32>; 3]> =
12567 const { std::cell::RefCell::new([Vec::new(), Vec::new(), Vec::new()]) };
12568}
12569
12570fn dense_ffn_batch(
12571 d: &DenseFfn,
12572 xs: &[f32],
12573 b: usize,
12574 pool: Option<&Pool>,
12575 mask_row: Option<&[u8]>,
12576) -> Vec<f32> {
12577 let inter = d.gate_proj.rows();
12578 let hidden = d.down_proj.rows();
12579 if mask_row.is_none()
12587 && d.act == Act::Silu
12588 && b >= 32
12589 && crate::gpu::enabled_here()
12590 && !crate::gpu::mm_killed()
12591 && refit_dir().is_none()
12596 && !ffn_probe_active()
12601 {
12602 if let (Some((model, w1)), Some((_, w3)), Some((_, w2))) = (
12603 d.gate_proj.mapped_q4t(),
12604 d.up_proj.mapped_q4t(),
12605 d.down_proj.mapped_q4t(),
12606 ) {
12607 let mut out = vec![0.0f32; b * hidden];
12608 if crate::gpu::q4t_ffn(model, w1, w3, w2, xs, b, hidden, inter, &mut out) {
12609 return out;
12610 }
12611 }
12612 if let (Some((model, w1)), Some((_, w3)), Some((_, w2))) = (
12617 d.gate_proj.mapped_q4tp(),
12618 d.up_proj.mapped_q4tp(),
12619 d.down_proj.mapped_q4tp(),
12620 ) {
12621 let mut out = vec![0.0f32; b * hidden];
12622 if crate::gpu::q4tp_ffn(model, w1, w3, w2, xs, b, hidden, inter, &mut out) {
12623 return out;
12624 }
12625 }
12626 }
12627 let mut g = vec![0.0f32; b * inter];
12628 d.gate_proj.matmat(xs, b, &mut g, pool);
12629 let mut u = vec![0.0f32; b * inter];
12630 d.up_proj.matmat(xs, b, &mut u, pool);
12631 if gate_topk() > 0 && d.act == Act::Silu {
12632 for t in 0..b {
12633 let row = &mut g[t * inter..(t + 1) * inter];
12634 for v in row.iter_mut() {
12635 *v = Act::Silu.combine(*v, 1.0);
12636 }
12637 keep_top_k(row, gate_topk());
12638 }
12639 for i in 0..b * inter {
12640 g[i] *= u[i];
12641 }
12642 } else {
12643 for i in 0..b * inter {
12644 g[i] = d.act.combine(g[i], u[i]);
12645 }
12646 }
12647 if let Some(row) = mask_row {
12648 zero_masked_cols(&mut g, b, inter, row);
12649 }
12650 if oracle_topk() > 0 {
12651 for t in 0..b {
12652 keep_top_k(&mut g[t * inter..(t + 1) * inter], oracle_topk());
12653 }
12654 }
12655 let mut out = vec![0.0f32; b * hidden];
12656 d.down_proj.matmat(&g, b, &mut out, pool);
12657 if refit_dir().is_some() {
12658 let li = crate::gpu::cur_layer();
12659 if li >= 0 {
12660 refit_accumulate(li as usize, &g, b, inter, &out, hidden, pool);
12661 }
12662 }
12663 FFN_PROBE.with(|pr| {
12667 if let Some(acc) = pr.borrow_mut().as_mut() {
12668 let li = crate::gpu::cur_layer();
12669 if li < 0 {
12670 return;
12671 }
12672 let Some(row) = acc.get_mut(li as usize) else {
12673 return;
12674 };
12675 let sq = probe_sq();
12676 for t in 0..b {
12677 for (a, &v) in row.iter_mut().zip(&g[t * inter..(t + 1) * inter]) {
12678 *a += if sq {
12679 (v as f64) * (v as f64)
12680 } else {
12681 (v as f64).abs()
12682 };
12683 }
12684 }
12685 }
12686 });
12687 out
12688}
12689
12690fn accumulate_act(m: &MoeFfn, xs: &[f32], b: usize) {
12695 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
12696 static DUMP: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
12697 let on = *ON.get_or_init(|| std::env::var("CMF_RMS_TRACE").is_ok());
12698 let dump = *DUMP.get_or_init(|| std::env::var("CMF_ACT_DUMP").is_ok());
12699 if (!on && !dump) || b == 0 {
12700 return;
12701 }
12702 let hidden = xs.len() / b;
12703 if on {
12704 let mut acc = m.act_sq.borrow_mut();
12705 if acc.len() < hidden {
12706 acc.resize(hidden, 0.0);
12707 }
12708 for t in 0..b {
12709 let row = &xs[t * hidden..(t + 1) * hidden];
12710 for (a, &v) in acc.iter_mut().zip(row) {
12711 *a += (v as f64) * (v as f64);
12712 }
12713 }
12714 }
12715 if dump {
12716 let cap: usize = std::env::var("CMF_ACT_DUMP_ROWS")
12719 .ok()
12720 .and_then(|v| v.parse().ok())
12721 .unwrap_or(4096);
12722 let mut rows = m.act_rows.borrow_mut();
12723 if rows.len() < cap * hidden {
12724 let take = b.min((cap * hidden - rows.len()) / hidden.max(1));
12725 rows.extend_from_slice(&xs[..take * hidden]);
12726 }
12727 }
12728}
12729
12730#[derive(Clone, Copy)]
12733struct SendVecs(*mut Vec<f32>);
12734unsafe impl Send for SendVecs {}
12735unsafe impl Sync for SendVecs {}
12736impl SendVecs {
12737 #[inline]
12738 fn at(self, i: usize) -> *mut Vec<f32> {
12739 unsafe { self.0.add(i) }
12740 }
12741}
12742
12743fn moe_ffn_batch(
12744 m: &MoeFfn,
12745 xs: &[f32],
12746 b: usize,
12747 hidden: usize,
12748 pool: Option<&Pool>,
12749 allowed: Option<&[bool]>,
12750) -> Vec<f32> {
12751 accumulate_act(m, xs, b);
12752 let ne = m.experts.len();
12753 let mut logits = vec![0.0f32; b * ne];
12754 match &m.resonance {
12755 Some(r) => {
12756 let hdim = xs.len() / b.max(1);
12757 for bi in 0..b {
12758 r.scores(
12759 &xs[bi * hdim..(bi + 1) * hdim],
12760 &mut logits[bi * ne..(bi + 1) * ne],
12761 );
12762 }
12763 }
12764 None => m.router.matmat(xs, b, &mut logits, pool),
12765 }
12766
12767 let mut assign: Vec<Vec<(usize, f32)>> = vec![Vec::new(); ne];
12770 {
12771 let mut st = m.stats.borrow_mut();
12772 if st.len() < ne {
12773 st.resize(ne, 0);
12774 }
12775 for bi in 0..b {
12776 let (idx, p, wsum) = moe_route(&logits[bi * ne..(bi + 1) * ne], m, allowed);
12777 for &e in &idx {
12778 st[e] += 1;
12779 assign[e].push((bi, p[e] / wsum));
12780 }
12781 }
12782 }
12783
12784 let mut out = vec![0.0f32; b * hidden];
12785 let cols = m.experts[0].gate_proj.cols();
12786 let run_expert = |d: &DenseFfn, list: &[(usize, f32)], out: &mut [f32]| {
12787 let sb = list.len();
12788 let mut sub = vec![0.0f32; sb * cols];
12789 for (k, &(bi, _)) in list.iter().enumerate() {
12790 sub[k * cols..(k + 1) * cols].copy_from_slice(&xs[bi * cols..(bi + 1) * cols]);
12791 }
12792 let eo = dense_ffn_batch(d, &sub, sb, pool, None);
12793 for (k, &(bi, w)) in list.iter().enumerate() {
12794 for i in 0..hidden {
12795 out[bi * hidden + i] += w * eo[k * hidden + i];
12796 }
12797 }
12798 };
12799 let active: Vec<usize> = (0..ne).filter(|&e| !assign[e].is_empty()).collect();
12805 if pool.is_some() && active.len() >= 8 {
12806 let mut panels: Vec<Vec<f32>> = vec![Vec::new(); active.len()];
12807 {
12808 let panel_ptr = SendVecs(panels.as_mut_ptr());
12809 let experts = &m.experts;
12812 let (active_r, assign_r) = (&active, &assign);
12813 let run = |start: usize, end: usize| {
12814 for ai in start..end {
12815 let e = active_r[ai];
12816 let list = &assign_r[e];
12817 let sb = list.len();
12818 let mut sub = vec![0.0f32; sb * cols];
12819 for (k, &(bi, _)) in list.iter().enumerate() {
12820 sub[k * cols..(k + 1) * cols]
12821 .copy_from_slice(&xs[bi * cols..(bi + 1) * cols]);
12822 }
12823 unsafe {
12825 *panel_ptr.at(ai) = dense_ffn_batch(&experts[e], &sub, sb, None, None);
12826 }
12827 }
12828 };
12829 match pool {
12830 Some(p) => p.run_rows(active.len(), &run),
12831 None => run(0, active.len()),
12832 }
12833 }
12834 for (ai, &e) in active.iter().enumerate() {
12835 for (k, &(bi, w)) in assign[e].iter().enumerate() {
12836 let eo = &panels[ai][k * hidden..(k + 1) * hidden];
12837 for i in 0..hidden {
12838 out[bi * hidden + i] += w * eo[i];
12839 }
12840 }
12841 }
12842 } else {
12843 for &e in &active {
12844 run_expert(&m.experts[e], &assign[e], &mut out);
12845 }
12846 }
12847 if let Some((se, gate)) = &m.shared {
12848 let all: Vec<(usize, f32)> = if let Some(gate) = gate {
12849 let mut gl = vec![0.0f32; b];
12850 gate.matmat(xs, b, &mut gl, pool);
12851 (0..b)
12852 .map(|bi| (bi, 1.0 / (1.0 + (-gl[bi]).exp())))
12853 .collect()
12854 } else {
12855 (0..b).map(|bi| (bi, 1.0)).collect()
12856 };
12857 run_expert(se, &all, &mut out);
12858 }
12859 out
12860}
12861
12862thread_local! {
12863 static FFN_SCRATCH: std::cell::RefCell<[Vec<f32>; 4]> =
12867 const { std::cell::RefCell::new([Vec::new(), Vec::new(), Vec::new(), Vec::new()]) };
12868}
12869
12870fn dense_ffn(d: &DenseFfn, x: &[f32], pool: Option<&Pool>) -> Vec<f32> {
12872 if gate_topk() > 0
12875 && let Some(out) = dense_ffn_dynamic(d, x, pool, gate_topk())
12876 {
12877 return out;
12878 }
12879 let prism_body = d.gate_proj.has_prism_contract()
12895 || d.up_proj.has_prism_contract()
12896 || d.down_proj.has_prism_contract();
12897 if !prism_body
12898 && crate::gpu::enabled_here()
12899 && (d.gate_proj.rows() >= crate::gpu::min_rows() || d.gate_proj.is_q1())
12900 {
12901 let arm = if d.gate_proj.is_q1() && crate::gpu::q1_force() {
12902 crate::gpu::ProbeArm::Gpu
12903 } else {
12904 crate::gpu::probe_arm(crate::gpu::OpClass::Ffn)
12905 };
12906 match arm {
12907 crate::gpu::ProbeArm::Gpu => {
12908 let t0 = std::time::Instant::now();
12909 if let Some(out) = dense_ffn_gpu(d, x, pool) {
12910 crate::gpu::probe_record(crate::gpu::OpClass::Ffn, true, t0.elapsed());
12911 return out;
12912 }
12913 crate::gpu::probe_note_decline(crate::gpu::OpClass::Ffn);
12917 }
12918 crate::gpu::ProbeArm::CpuTimed => {
12919 let t0 = std::time::Instant::now();
12920 let out = crate::gpu::cpu_scope(|| dense_ffn_cpu(d, x, pool));
12921 crate::gpu::probe_record(crate::gpu::OpClass::Ffn, false, t0.elapsed());
12922 return out;
12923 }
12924 crate::gpu::ProbeArm::Cpu => {
12925 return crate::gpu::cpu_scope(|| dense_ffn_cpu(d, x, pool));
12926 }
12927 }
12928 }
12929 dense_ffn_cpu(d, x, pool)
12930}
12931
12932fn dense_ffn_cpu(d: &DenseFfn, x: &[f32], pool: Option<&Pool>) -> Vec<f32> {
12934 let inter = d.gate_proj.rows();
12935 FFN_SCRATCH.with(|s| {
12936 let mut s = s.borrow_mut();
12937 let [g, u, ..] = &mut *s;
12938 g.resize(inter, 0.0);
12939 if gate_topk() > 0 {
12942 u.resize(inter, 0.0);
12946 QTensor::matvec_many([&d.gate_proj, &d.up_proj], x, [g, u], pool);
12947 for i in 0..inter {
12948 g[i] = Act::Silu.combine(g[i], 1.0);
12949 }
12950 keep_top_k(g, gate_topk());
12951 for i in 0..inter {
12952 g[i] *= u[i];
12953 }
12954 } else if d.act == Act::Silu
12955 && QTensor::matvec_silu_mul(&d.gate_proj, &d.up_proj, x, g, pool)
12956 {
12957 } else {
12959 u.resize(inter, 0.0);
12960 QTensor::matvec_many([&d.gate_proj, &d.up_proj], x, [g, u], pool);
12962 for i in 0..inter {
12963 g[i] = d.act.combine(g[i], u[i]);
12964 }
12965 }
12966 FFN_PROBE.with(|pr| {
12974 if let Some(acc) = pr.borrow_mut().as_mut() {
12975 let li = crate::gpu::cur_layer();
12976 if li >= 0 {
12977 if let Some(row) = acc.get_mut(li as usize) {
12978 match probe_topk() {
12979 0 if probe_sq() => {
12980 for (a, &v) in row.iter_mut().zip(g.iter()) {
12981 *a += (v as f64) * (v as f64);
12982 }
12983 }
12984 0 if probe_signed() => {
12985 for (a, &v) in row.iter_mut().zip(g.iter()) {
12986 *a += v as f64;
12987 }
12988 }
12989 0 => {
12990 for (a, &v) in row.iter_mut().zip(g.iter()) {
12991 *a += (v as f64).abs();
12992 }
12993 }
12994 k => {
12995 let n = g.len();
12996 let k = k.min(n);
12997 let mut mag: Vec<f32> = g.iter().map(|v| v.abs()).collect();
12998 let (_, kth, _) = mag.select_nth_unstable_by(k - 1, |a, b| {
12999 b.partial_cmp(a).unwrap_or(std::cmp::Ordering::Equal)
13000 });
13001 let thr = *kth;
13002 for (a, &v) in row.iter_mut().zip(g.iter()) {
13003 if v.abs() >= thr {
13004 *a += 1.0;
13005 }
13006 }
13007 }
13008 }
13009 }
13010 }
13011 }
13012 });
13013 if oracle_topk() > 0 {
13014 keep_top_k(g, oracle_topk());
13015 }
13016 {
13017 let li = crate::gpu::cur_layer();
13018 if li >= 0 {
13019 adump_row(li as usize, g);
13020 }
13021 }
13022 let mut out = attention::take_buf(d.down_proj.rows());
13023 d.down_proj.matvec(g, &mut out, pool);
13024 out
13025 })
13026}
13027
13028pub struct RefitAcc {
13041 pub support: Vec<u32>,
13042 pub gss: Vec<f32>,
13043 pub ya: Vec<f32>,
13044 pub hidden: usize,
13045 pub tokens: u64,
13046 pub buf_g: Vec<f32>,
13052 pub buf_o: Vec<f32>,
13053 pub buf_t: usize,
13054}
13055
13056type RefitState = (std::collections::HashMap<usize, RefitAcc>, Vec<f32>);
13060
13061static REFIT: std::sync::OnceLock<Option<(String, std::sync::Mutex<RefitState>)>> =
13062 std::sync::OnceLock::new();
13063
13064fn ffn_probe_active() -> bool {
13067 FFN_PROBE.with(|p| p.borrow().is_some())
13068}
13069
13070fn refit_dir() -> Option<&'static (String, std::sync::Mutex<RefitState>)> {
13071 REFIT
13072 .get_or_init(|| {
13073 std::env::var("CMF_FFN_REFIT").ok().map(|d| {
13074 (
13075 d,
13076 std::sync::Mutex::new((std::collections::HashMap::new(), Vec::new())),
13077 )
13078 })
13079 })
13080 .as_ref()
13081}
13082
13083fn refit_accumulate(
13085 li: usize,
13086 g: &[f32],
13087 b: usize,
13088 inter: usize,
13089 out: &[f32],
13090 hidden: usize,
13091 pool: Option<&Pool>,
13092) {
13093 let Some((dir, map)) = refit_dir() else {
13094 return;
13095 };
13096 static SPAN: std::sync::OnceLock<(usize, usize)> = std::sync::OnceLock::new();
13097 let (from, to) = *SPAN.get_or_init(|| {
13098 let g = |k: &str, d: usize| {
13099 std::env::var(k)
13100 .ok()
13101 .and_then(|v| v.parse().ok())
13102 .unwrap_or(d)
13103 };
13104 (
13105 g("CMF_FFN_REFIT_FROM", 0),
13106 g("CMF_FFN_REFIT_TO", usize::MAX),
13107 )
13108 });
13109 if li < from || li > to {
13110 return;
13111 }
13112 let mut guard = map.lock().unwrap();
13113 let (map, shared) = &mut *guard;
13114 let acc = match map.entry(li) {
13115 std::collections::hash_map::Entry::Occupied(e) => e.into_mut(),
13116 std::collections::hash_map::Entry::Vacant(e) => {
13117 let path = format!("{dir}/support.{li}.u32");
13118 let Ok(bytes) = std::fs::read(&path) else {
13119 eprintln!("refit: no {path} — layer {li} skipped");
13120 return;
13121 };
13122 let n = u32::from_le_bytes(bytes[0..4].try_into().unwrap()) as usize;
13123 let support: Vec<u32> = bytes[4..4 + n * 4]
13124 .chunks_exact(4)
13125 .map(|c| u32::from_le_bytes([c[0], c[1], c[2], c[3]]))
13126 .collect();
13127 eprintln!(
13128 "refit: layer {li} support {n} ({:.0} MB of accumulator)",
13129 (n * n + hidden * n) as f64 * 4.0 / 1e6
13130 );
13131 e.insert(RefitAcc {
13132 gss: vec![0.0; n * n],
13133 ya: vec![0.0; hidden * n],
13134 buf_g: Vec::new(),
13135 buf_o: Vec::new(),
13136 buf_t: 0,
13137 support,
13138 hidden,
13139 tokens: 0,
13140 })
13141 }
13142 };
13143 let ns = acc.support.len();
13144 let cap = refit_batch();
13146 if acc.buf_g.is_empty() {
13147 acc.buf_g = vec![0.0; ns * cap];
13148 acc.buf_o = vec![0.0; hidden * cap];
13149 }
13150 let take = b.min(cap - acc.buf_t);
13151 for t in 0..take {
13152 let col = acc.buf_t + t;
13153 for (j, &n) in acc.support.iter().enumerate() {
13154 acc.buf_g[j * cap + col] = g[t * inter + n as usize];
13155 }
13156 for h in 0..hidden {
13157 acc.buf_o[h * cap + col] = out[t * hidden + h];
13158 }
13159 }
13160 acc.buf_t += take;
13161 acc.tokens += take as u64;
13162 if acc.buf_t < cap {
13163 return;
13164 }
13165 let bt = acc.buf_t;
13166 acc.buf_t = 0;
13167 let RefitAcc {
13177 gss,
13178 ya,
13179 buf_g,
13180 buf_o,
13181 ..
13182 } = acc;
13183 let need = (ns * ns).max(hidden * ns);
13184 if shared.len() < need {
13185 shared.resize(need, 0.0);
13186 }
13187 let scratch = &mut shared[..];
13188 let _ = bt;
13189 if crate::gpu::gemm_nt_f32_transient(buf_g, buf_g, &mut scratch[..ns * ns], ns, cap, ns) {
13190 add_into(gss, &scratch[..ns * ns], pool);
13191 if crate::gpu::gemm_nt_f32_transient(
13192 buf_o,
13193 buf_g,
13194 &mut scratch[..hidden * ns],
13195 hidden,
13196 cap,
13197 ns,
13198 ) {
13199 add_into(ya, &scratch[..hidden * ns], pool);
13200 } else {
13201 accum_outer_t(ya, hidden, ns, cap, buf_o, buf_g, pool);
13202 }
13203 } else {
13204 accum_outer_t(gss, ns, ns, cap, buf_g, buf_g, pool);
13205 accum_outer_t(ya, hidden, ns, cap, buf_o, buf_g, pool);
13206 }
13207 }
13211
13212fn refit_batch() -> usize {
13214 static B: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
13215 *B.get_or_init(|| {
13216 std::env::var("CMF_FFN_REFIT_BATCH")
13217 .ok()
13218 .and_then(|v| v.parse().ok())
13219 .unwrap_or(4096)
13220 })
13221}
13222
13223fn accum_outer_t(
13226 c: &mut [f32],
13227 m: usize,
13228 n: usize,
13229 b: usize,
13230 left: &[f32],
13231 right: &[f32],
13232 pool: Option<&Pool>,
13233) {
13234 let ptr = SendMut(c.as_mut_ptr());
13235 let body = |i: usize| {
13236 let ptr = &ptr;
13237 let row = unsafe { std::slice::from_raw_parts_mut(ptr.0.add(i * n), n) };
13238 for t in 0..b {
13239 let a = left[i * b + t];
13240 if a == 0.0 {
13241 continue;
13242 }
13243 for (j, o) in row.iter_mut().enumerate() {
13244 *o += a * right[j * b + t];
13245 }
13246 }
13247 };
13248 match pool {
13249 Some(p) if m > 1 => p.run_rows(m, &|s, e| {
13250 for i in s..e {
13251 body(i);
13252 }
13253 }),
13254 _ => {
13255 for i in 0..m {
13256 body(i);
13257 }
13258 }
13259 }
13260}
13261
13262fn add_into(dst: &mut [f32], src: &[f32], pool: Option<&Pool>) {
13265 let n = dst.len().min(src.len());
13266 match pool {
13267 Some(p) if n >= 1 << 16 => {
13268 let ptr = SendMut(dst.as_mut_ptr());
13269 let f = |s: usize, e: usize| {
13270 let ptr = &ptr;
13271 for blk in s..e {
13272 let (a, b) = (blk * 4096, ((blk + 1) * 4096).min(n));
13273 for i in a..b {
13274 unsafe { *ptr.0.add(i) += src[i] };
13275 }
13276 }
13277 };
13278 p.run_rows(n.div_ceil(4096), &f);
13279 }
13280 _ => {
13281 for (d, v) in dst.iter_mut().zip(&src[..n]) {
13282 *d += *v;
13283 }
13284 }
13285 }
13286}
13287
13288fn accum_outer(
13293 c: &mut [f32],
13294 m: usize,
13295 n: usize,
13296 b: usize,
13297 left: &[f32],
13298 right: &[f32],
13299 pool: Option<&Pool>,
13300) {
13301 const TILE: usize = 32;
13302 let tiles = m.div_ceil(TILE);
13303 let cp = SendMut(c.as_mut_ptr());
13304 let body = |ti: usize| {
13305 let cp = &cp;
13306 let i0 = ti * TILE;
13307 let i1 = (i0 + TILE).min(m);
13308 for t in 0..b {
13309 let r = &right[t * n..t * n + n];
13310 for i in i0..i1 {
13311 let a = left[i * b + t];
13312 if a == 0.0 {
13313 continue;
13314 }
13315 let row = unsafe { std::slice::from_raw_parts_mut(cp.0.add(i * n), n) };
13317 for (o, v) in row.iter_mut().zip(r) {
13318 *o += a * *v;
13319 }
13320 }
13321 }
13322 };
13323 match pool {
13324 Some(p) if tiles > 1 => p.run_rows(tiles, &|s, e| {
13325 for ti in s..e {
13326 body(ti);
13327 }
13328 }),
13329 _ => {
13330 for ti in 0..tiles {
13331 body(ti);
13332 }
13333 }
13334 }
13335}
13336
13337pub fn refit_flush() -> usize {
13339 let Some((dir, map)) = refit_dir() else {
13340 return 0;
13341 };
13342 let guard = map.lock().unwrap();
13343 let mut n = 0;
13344 for (li, acc) in guard.0.iter() {
13345 let w = |name: &str, v: &[f32]| {
13348 let path = format!("{dir}/{name}.{li}.f32");
13349 let bytes: Vec<u8> = v.iter().flat_map(|x| x.to_le_bytes()).collect();
13350 match std::fs::write(&path, &bytes) {
13351 Ok(()) => {}
13352 Err(e) => eprintln!(
13353 "refit: FAILED to write {path} ({} MB): {e}",
13354 bytes.len() / 1_000_000
13355 ),
13356 }
13357 };
13358 w("gss", &acc.gss);
13359 w("ya", &acc.ya);
13360 println!(
13361 "refit L{li}: {} support, {} tokens, hidden {}",
13362 acc.support.len(),
13363 acc.tokens,
13364 acc.hidden
13365 );
13366 n += 1;
13367 }
13368 n
13369}
13370
13371fn adump_row(li: usize, g: &[f32]) {
13376 use std::io::Write as _;
13377 static FILES: std::sync::OnceLock<
13378 Option<(
13379 String,
13380 std::sync::Mutex<std::collections::HashMap<usize, std::fs::File>>,
13381 )>,
13382 > = std::sync::OnceLock::new();
13383 let Some((prefix, map)) = FILES
13384 .get_or_init(|| {
13385 std::env::var("CMF_FFN_ADUMP")
13386 .ok()
13387 .map(|p| (p, std::sync::Mutex::new(std::collections::HashMap::new())))
13388 })
13389 .as_ref()
13390 else {
13391 return;
13392 };
13393 static SPAN: std::sync::OnceLock<(usize, usize)> = std::sync::OnceLock::new();
13396 let (from, to) = *SPAN.get_or_init(|| {
13397 let g = |k: &str, d: usize| {
13398 std::env::var(k)
13399 .ok()
13400 .and_then(|v| v.parse().ok())
13401 .unwrap_or(d)
13402 };
13403 (
13404 g("CMF_FFN_ADUMP_FROM", 0),
13405 g("CMF_FFN_ADUMP_TO", usize::MAX),
13406 )
13407 });
13408 if li < from || li > to {
13409 return;
13410 }
13411 let mut map = map.lock().unwrap();
13412 let f = map.entry(li).or_insert_with(|| {
13413 std::fs::File::create(format!("{prefix}.{li}.f16")).expect("adump file")
13414 });
13415 let mut bytes = Vec::with_capacity(g.len() * 2);
13416 for v in g {
13417 bytes.extend_from_slice(&cortiq_core::quant::f32_to_f16(*v).to_le_bytes());
13418 }
13419 let _ = f.write_all(&bytes);
13420}
13421
13422fn oracle_topk() -> usize {
13428 static K: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
13429 *K.get_or_init(|| {
13430 std::env::var("CMF_FFN_ORACLE_TOPK")
13431 .ok()
13432 .and_then(|v| v.parse().ok())
13433 .unwrap_or(0)
13434 })
13435}
13436
13437fn gate_topk() -> usize {
13443 static K: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
13444 *K.get_or_init(|| {
13445 std::env::var("CMF_FFN_GATE_TOPK")
13446 .ok()
13447 .and_then(|v| v.parse().ok())
13448 .unwrap_or(0)
13449 })
13450}
13451
13452fn gate_block() -> usize {
13459 static B: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
13460 *B.get_or_init(|| {
13461 std::env::var("CMF_FFN_GATE_BLOCK")
13462 .ok()
13463 .and_then(|v| v.parse().ok())
13464 .unwrap_or(1)
13465 })
13466}
13467
13468fn keep_top_blocks(g: &mut [f32], keep_n: usize, block: usize) {
13470 let n = g.len();
13471 let nb = n.div_ceil(block);
13472 let kb = (keep_n.div_ceil(block)).clamp(1, nb);
13473 if kb >= nb {
13474 return;
13475 }
13476 let mut score: Vec<f32> = (0..nb)
13477 .map(|b| {
13478 g[b * block..((b + 1) * block).min(n)]
13479 .iter()
13480 .map(|v| v * v)
13481 .sum::<f32>()
13482 })
13483 .collect();
13484 let mut ord = score.clone();
13485 let (_, kth, _) = ord.select_nth_unstable_by(kb - 1, |a, b| {
13486 b.partial_cmp(a).unwrap_or(std::cmp::Ordering::Equal)
13487 });
13488 let thr = *kth;
13489 for b in 0..nb {
13490 if score[b] < thr {
13491 g[b * block..((b + 1) * block).min(n)].fill(0.0);
13492 }
13493 }
13494 score.clear();
13495}
13496
13497fn keep_top_k(g: &mut [f32], k: usize) {
13499 if gate_block() > 1 {
13500 return keep_top_blocks(g, k, gate_block());
13501 }
13502 let n = g.len();
13503 if k == 0 || k >= n {
13504 return;
13505 }
13506 let mut mag: Vec<f32> = g.iter().map(|v| v.abs()).collect();
13507 let (_, kth, _) = mag.select_nth_unstable_by(k - 1, |a, b| {
13508 b.partial_cmp(a).unwrap_or(std::cmp::Ordering::Equal)
13509 });
13510 let thr = *kth;
13511 for v in g.iter_mut() {
13512 if v.abs() < thr {
13513 *v = 0.0;
13514 }
13515 }
13516}
13517
13518fn probe_sq() -> bool {
13522 static S: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
13523 *S.get_or_init(|| std::env::var("CMF_FFN_PROBE_SQ").is_ok())
13524}
13525
13526fn probe_signed() -> bool {
13530 static S: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
13531 *S.get_or_init(|| std::env::var("CMF_FFN_PROBE_SIGNED").is_ok())
13532}
13533
13534fn meanfill() -> Option<&'static (usize, Vec<f32>)> {
13542 static M: std::sync::OnceLock<Option<(usize, Vec<f32>)>> = std::sync::OnceLock::new();
13543 M.get_or_init(|| {
13544 let p = std::env::var("CMF_FFN_MEANFILL").ok()?;
13545 let b = std::fs::read(&p).ok()?;
13546 let inter = u32::from_le_bytes(b[4..8].try_into().ok()?) as usize;
13547 let vals: Vec<f32> = b[8..]
13548 .chunks_exact(4)
13549 .map(|c| f32::from_le_bytes([c[0], c[1], c[2], c[3]]))
13550 .collect();
13551 eprintln!("meanfill: {} value(s), inter {inter}", vals.len());
13552 Some((inter, vals))
13553 })
13554 .as_ref()
13555}
13556
13557fn probe_topk() -> usize {
13560 static K: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
13561 *K.get_or_init(|| {
13562 std::env::var("CMF_FFN_PROBE_TOPK")
13563 .ok()
13564 .and_then(|v| v.parse().ok())
13565 .unwrap_or(0)
13566 })
13567}
13568
13569thread_local! {
13570 static FFN_PROBE: std::cell::RefCell<Option<Vec<Vec<f64>>>> =
13573 const { std::cell::RefCell::new(None) };
13574}
13575
13576fn dense_ffn_dynamic(d: &DenseFfn, x: &[f32], pool: Option<&Pool>, k: usize) -> Option<Vec<f32>> {
13589 if d.gate_proj.has_prism_contract()
13594 || d.up_proj.has_prism_contract()
13595 || d.down_proj.has_prism_contract()
13596 {
13597 return None;
13598 }
13599 let dt = d.down_t.as_ref()?;
13600 let inter = d.gate_proj.rows();
13601 let hidden = dt.cols();
13602 if k == 0 || k >= inter || d.act != Act::Silu {
13603 return None;
13604 }
13605 DYN_SCRATCH.with(|sc| {
13606 let mut sc = sc.borrow_mut();
13607 let DynScratch {
13608 g,
13609 mag,
13610 live,
13611 parts,
13612 } = &mut *sc;
13613 g.resize(inter, 0.0);
13614 d.gate_proj.matvec(x, g, pool);
13615 for v in g.iter_mut() {
13616 *v = inference::silu(*v);
13617 }
13618 mag.clear();
13621 mag.extend(g.iter().map(|v| v.abs()));
13622 let (_, kth, _) = mag.select_nth_unstable_by(k - 1, |a, b| {
13623 b.partial_cmp(a).unwrap_or(std::cmp::Ordering::Equal)
13624 });
13625 let thr = *kth;
13626 live.clear();
13627 live.extend((0..inter as u32).filter(|&n| g[n as usize].abs() >= thr));
13628 let mut out = vec![0.0f32; hidden];
13629 match pool {
13630 Some(p) if live.len() >= 64 => {
13631 let nw = p.n_workers() + 1;
13632 parts.clear();
13633 parts.resize(nw * hidden, 0.0);
13634 let ptr = SendMut(parts.as_mut_ptr());
13635 let n = live.len();
13636 let live_ref: &[u32] = live;
13637 let g_ref: &[f32] = g;
13638 p.run(&|w, workers| {
13639 let chunk = n.div_ceil(workers);
13640 let (s, e) = (w * chunk, ((w + 1) * chunk).min(n));
13641 if s >= e {
13642 return;
13643 }
13644 WORKER_SCRATCH.with(|ws| {
13645 let mut ws = ws.borrow_mut();
13646 let [scratch, acc] = &mut *ws;
13647 scratch.resize(hidden.max(x.len()), 0.0);
13648 acc.clear();
13649 acc.resize(hidden, 0.0);
13650 for (o, &nrm) in live_ref[s..e].iter().enumerate() {
13651 if let Some(&nx) = live_ref[s..e].get(o + 1) {
13654 d.up_proj.prefetch_row(nx as usize);
13655 dt.prefetch_row(nx as usize);
13656 }
13657 let idx = nrm as usize;
13658 let up = d.up_proj.row_dot(idx, x, scratch);
13659 let a = g_ref[idx] * up;
13660 if a != 0.0 {
13661 dt.add_row_scaled(idx, a, acc, scratch);
13662 }
13663 }
13664 for (j, v) in acc.iter().enumerate() {
13665 unsafe { *ptr.at(w * hidden + j) = *v };
13666 }
13667 });
13668 });
13669 for w in 0..nw {
13670 for (j, o) in out.iter_mut().enumerate() {
13671 *o += parts[w * hidden + j];
13672 }
13673 }
13674 }
13675 _ => {
13676 WORKER_SCRATCH.with(|ws| {
13677 let mut ws = ws.borrow_mut();
13678 let [scratch, _acc] = &mut *ws;
13679 scratch.resize(hidden.max(x.len()), 0.0);
13680 for &nrm in live.iter() {
13681 let idx = nrm as usize;
13682 let up = d.up_proj.row_dot(idx, x, scratch);
13683 let a = g[idx] * up;
13684 if a != 0.0 {
13685 dt.add_row_scaled(idx, a, &mut out, scratch);
13686 }
13687 }
13688 });
13689 }
13690 }
13691 Some(out)
13692 })
13693}
13694
13695struct DynScratch {
13698 g: Vec<f32>,
13699 mag: Vec<f32>,
13700 live: Vec<u32>,
13701 parts: Vec<f32>,
13702}
13703
13704thread_local! {
13705 static DYN_SCRATCH: std::cell::RefCell<DynScratch> = const {
13706 std::cell::RefCell::new(DynScratch {
13707 g: Vec::new(),
13708 mag: Vec::new(),
13709 live: Vec::new(),
13710 parts: Vec::new(),
13711 })
13712 };
13713 static WORKER_SCRATCH: std::cell::RefCell<[Vec<f32>; 2]> =
13715 const { std::cell::RefCell::new([Vec::new(), Vec::new()]) };
13716}
13717
13718fn dense_ffn_masked(d: &DenseFfn, x: &[f32], pool: Option<&Pool>, mask_row: &[u8]) -> Vec<f32> {
13723 let inter = d.gate_proj.rows();
13724 FFN_SCRATCH.with(|s| {
13725 let mut s = s.borrow_mut();
13726 let [g, u, ..] = &mut *s;
13727 g.resize(inter, 0.0);
13728 if d.act == Act::Silu && QTensor::matvec_silu_mul(&d.gate_proj, &d.up_proj, x, g, pool) {
13729 } else {
13731 u.resize(inter, 0.0);
13732 QTensor::matvec_many([&d.gate_proj, &d.up_proj], x, [g, u], pool);
13733 for i in 0..inter {
13734 g[i] = d.act.combine(g[i], u[i]);
13735 }
13736 }
13737 zero_masked_cols(g, 1, inter, mask_row);
13738 let mut out = attention::take_buf(d.down_proj.rows());
13739 d.down_proj.matvec(g, &mut out, pool);
13740 out
13741 })
13742}
13743
13744fn dense_ffn_gpu(d: &DenseFfn, x: &[f32], _pool: Option<&Pool>) -> Option<Vec<f32>> {
13750 if d.gate_proj.has_prism_contract()
13751 || d.up_proj.has_prism_contract()
13752 || d.down_proj.has_prism_contract()
13753 {
13754 return None;
13755 }
13756 if d.act != Act::Silu {
13758 return None;
13759 }
13760 if d.gate_proj.rows() < crate::gpu::min_rows() && !d.gate_proj.is_q1() {
13763 return None;
13764 }
13765 let mut jobs: Vec<crate::gpu::MoeJob> = Vec::with_capacity(1);
13766 let mut model_ref = None;
13767 moe_push_job(d, x, 1.0, &mut jobs, &mut model_ref)?;
13768 let model = model_ref?;
13769 let hidden = jobs[0].down.1;
13770 let mut out = attention::take_buf(hidden);
13771 if crate::gpu::moe_block(&model, &jobs, &mut out) {
13772 Some(out)
13773 } else {
13774 let mut out = out;
13775 attention::recycle_buf(&mut out);
13776 None
13777 }
13778}
13779
13780#[allow(clippy::type_complexity)]
13785#[allow(clippy::type_complexity)]
13786pub(crate) fn moe_parts(
13787 t: &QTensor,
13788) -> Option<(
13789 &std::sync::Arc<cortiq_core::CmfModel>,
13790 usize,
13791 usize,
13792 usize,
13793 &[f32],
13794 &[f32],
13795 bool,
13796 bool,
13797 bool,
13798)> {
13799 match t {
13800 QTensor::Mapped {
13801 model,
13802 idx,
13803 dtype: dt @ (cortiq_core::TensorDtype::Q8_2f | cortiq_core::TensorDtype::Q8Row),
13804 rows,
13805 cols,
13806 row_scale,
13807 col_field,
13808 ..
13809 } if (*dt == cortiq_core::TensorDtype::Q8Row) || !col_field.is_empty() => Some((
13810 model, *idx, *rows, *cols, row_scale, col_field, false, false, false,
13811 )),
13812 QTensor::Mapped {
13814 model,
13815 idx,
13816 dtype: cortiq_core::TensorDtype::Q1,
13817 rows,
13818 cols,
13819 ..
13820 } => Some((
13821 model,
13822 *idx,
13823 *rows,
13824 *cols,
13825 &[][..],
13826 &[][..],
13827 true,
13828 false,
13829 false,
13830 )),
13831 QTensor::Mapped {
13833 model,
13834 idx,
13835 dtype: cortiq_core::TensorDtype::Q4Tiled,
13836 rows,
13837 cols,
13838 ..
13839 } => Some((
13840 model,
13841 *idx,
13842 *rows,
13843 *cols,
13844 &[][..],
13845 &[][..],
13846 false,
13847 true,
13848 false,
13849 )),
13850 QTensor::Mapped {
13852 model,
13853 idx,
13854 dtype: cortiq_core::TensorDtype::Q4TiledP,
13855 rows,
13856 cols,
13857 ..
13858 } => Some((
13859 model,
13860 *idx,
13861 *rows,
13862 *cols,
13863 &[][..],
13864 &[][..],
13865 false,
13866 true,
13867 false,
13868 )),
13869 QTensor::Mapped {
13873 model,
13874 idx,
13875 dtype: cortiq_core::TensorDtype::Q2TiledP,
13876 rows,
13877 cols,
13878 ..
13879 } => Some((
13880 model,
13881 *idx,
13882 *rows,
13883 *cols,
13884 &[][..],
13885 &[][..],
13886 false,
13887 true,
13888 true,
13889 )),
13890 _ => None,
13891 }
13892}
13893
13894#[cfg(target_os = "macos")]
13902fn metal_moe_graph_parts(m: &MoeFfn, hidden: usize) -> Option<crate::gpu::GpuMoe<'_>> {
13903 if m.router_input_norm
13904 || m.route_tau.is_some()
13905 || m.mask.is_some()
13906 || m.per_expert_scale.is_some()
13907 || m.experts.is_empty()
13908 || m.top_k == 0
13909 || m.resonance.is_some()
13910 {
13911 return None;
13912 }
13913 let (sh, sg) = match &m.shared {
13916 Some((sh, sg)) => (sh, sg.as_ref()),
13917 None => return None,
13918 };
13919 let (rf, rr, rc) = m.router.f32_parts()?;
13920 if rr != m.experts.len() || rc != hidden {
13921 return None;
13922 }
13923 let shared_gated = sg.is_some();
13924 let sf = match sg {
13925 Some(sg) => {
13926 let (sf, sr, sc) = sg.f32_parts()?;
13927 if sr * sc != hidden {
13928 return None;
13929 }
13930 sf
13931 }
13932 None => &rf[..hidden],
13935 };
13936 if let Some(b) = &m.expert_bias {
13937 if b.len() != m.experts.len() {
13938 return None;
13939 }
13940 }
13941 let inter = m.experts[0].gate_proj.rows();
13942 let gu_q2 = m.experts[0].gate_proj.mapped_q2tp().is_some();
13945 let trio = |e: &DenseFfn| -> Option<(usize, usize, usize)> {
13946 if e.act != Act::Silu
13947 || e.gate_proj.rows() != inter
13948 || e.gate_proj.cols() != hidden
13949 || e.up_proj.rows() != inter
13950 || e.up_proj.cols() != hidden
13951 || e.down_proj.rows() != hidden
13952 || e.down_proj.cols() != inter
13953 {
13954 return None;
13955 }
13956 let pick = |t: &QTensor| -> Option<usize> {
13957 if gu_q2 {
13958 t.mapped_q2tp().map(|(_, i)| i)
13959 } else {
13960 t.mapped_q4tp().map(|(_, i)| i)
13961 }
13962 };
13963 Some((
13964 pick(&e.gate_proj)?,
13965 pick(&e.up_proj)?,
13966 e.down_proj.mapped_q4tp().map(|(_, i)| i)?,
13967 ))
13968 };
13969 let experts = m.experts.iter().map(trio).collect::<Option<Vec<_>>>()?;
13970 let shared = trio(sh)?;
13971 Some(crate::gpu::GpuMoe {
13972 router: rf,
13973 sgate: sf,
13974 experts,
13975 shared,
13976 n_exp: m.experts.len(),
13977 top_k: m.top_k,
13978 inter,
13979 norm_topk: m.norm_topk_prob,
13980 route_scale: m.routed_scaling,
13981 gu_q2,
13982 sigmoid: m.router_sigmoid,
13983 bias: m.expert_bias.as_deref(),
13984 shared_gated,
13985 })
13986}
13987
13988pub(crate) fn moe_push_job_parts<'a>(
13992 gate: &'a QTensor,
13993 up: &'a QTensor,
13994 down: &'a QTensor,
13995 x: &[f32],
13996 w: f32,
13997 swiglu_limit: f32,
13998 jobs: &mut Vec<crate::gpu::MoeJob<'a>>,
13999 model_ref: &mut Option<std::sync::Arc<cortiq_core::CmfModel>>,
14000) -> Option<()> {
14001 use crate::qtensor::prescale;
14002 let (gm, gi, gr, gc, grs, gcf, gq1, gq4, gq2) = moe_parts(gate)?;
14003 let (_, ui, ur, uc, urs, ucf, uq1, uq4, uq2) = moe_parts(up)?;
14004 let (_, di, dr, dc, drs, dcf, dq1, dq4, dq2) = moe_parts(down)?;
14005 if gq1 != uq1 || uq1 != dq1 || gq4 != uq4 || uq4 != dq4 || gq2 != uq2 {
14006 return None; }
14008 if gq2 && (dq2 || !dq4 || down.mapped_q4tp().is_none()) {
14011 return None;
14012 }
14013 if !gq2 && dq2 {
14014 return None;
14015 }
14016 model_ref.get_or_insert_with(|| gm.clone());
14017 let dt = |cf: &[f32]| {
14018 if cf.is_empty() {
14019 cortiq_core::TensorDtype::Q8Row
14020 } else {
14021 cortiq_core::TensorDtype::Q8_2f
14022 }
14023 };
14024 jobs.push(crate::gpu::MoeJob {
14025 gate: (gi, gr, gc, grs),
14026 up: (ui, ur, uc, urs),
14027 down: (di, dr, dc, drs),
14028 xs_gate: prescale(x, gcf, dt(gcf)).into_owned(),
14029 xs_up: prescale(x, ucf, dt(ucf)).into_owned(),
14030 down_col: dcf,
14031 w,
14032 q1: gq1,
14033 q4t: gq4 && !gq2 && gate.mapped_q4tp().is_none(),
14034 q4tp: gq4 && (gq2 || gate.mapped_q4tp().is_some()),
14035 gu_q2: gq2,
14036 swiglu_limit,
14037 });
14038 Some(())
14039}
14040
14041fn moe_push_job<'a>(
14043 d: &'a DenseFfn,
14044 x: &[f32],
14045 w: f32,
14046 jobs: &mut Vec<crate::gpu::MoeJob<'a>>,
14047 model_ref: &mut Option<std::sync::Arc<cortiq_core::CmfModel>>,
14048) -> Option<()> {
14049 use crate::qtensor::prescale;
14050 if d.act != Act::Silu {
14051 return None; }
14053 let (gm, gi, gr, gc, grs, gcf, gq1, gq4, gq2) = moe_parts(&d.gate_proj)?;
14054 let (_, ui, ur, uc, urs, ucf, uq1, uq4, uq2) = moe_parts(&d.up_proj)?;
14055 let (_, di, dr, dc, drs, dcf, dq1, dq4, dq2) = moe_parts(&d.down_proj)?;
14056 if gq1 != uq1 || uq1 != dq1 || gq4 != uq4 || uq4 != dq4 || gq2 != uq2 {
14057 return None; }
14059 if gq2 && (dq2 || !dq4 || d.down_proj.mapped_q4tp().is_none()) {
14060 return None;
14061 }
14062 if !gq2 && dq2 {
14063 return None;
14064 }
14065 model_ref.get_or_insert_with(|| gm.clone());
14066 let gdt = if gcf.is_empty() {
14067 cortiq_core::TensorDtype::Q8Row
14068 } else {
14069 cortiq_core::TensorDtype::Q8_2f
14070 };
14071 let udt = if ucf.is_empty() {
14072 cortiq_core::TensorDtype::Q8Row
14073 } else {
14074 cortiq_core::TensorDtype::Q8_2f
14075 };
14076 jobs.push(crate::gpu::MoeJob {
14077 gate: (gi, gr, gc, grs),
14078 up: (ui, ur, uc, urs),
14079 down: (di, dr, dc, drs),
14080 xs_gate: prescale(x, gcf, gdt).into_owned(),
14081 xs_up: prescale(x, ucf, udt).into_owned(),
14082 down_col: dcf,
14083 w,
14084 q1: gq1,
14085 q4t: gq4 && !gq2 && d.gate_proj.mapped_q4tp().is_none(),
14086 q4tp: gq4 && (gq2 || d.gate_proj.mapped_q4tp().is_some()),
14087 gu_q2: gq2,
14088 swiglu_limit: 0.0,
14089 });
14090 Some(())
14091}
14092
14093fn sparse_ffn_quant(
14100 d: &DenseFfn,
14101 x: &[f32],
14102 active: &[u16],
14103 hidden: usize,
14104 pool: Option<&Pool>,
14105) -> Vec<f32> {
14106 let n = active.len();
14107 let inter = d.gate_proj.rows();
14108 let mut act = vec![0.0f32; n];
14109 let need_scratch = !(d.gate_proj.sparse_col_ok() && d.up_proj.sparse_col_ok());
14112 let compute = |ai: usize| -> f32 {
14113 let idx = active[ai] as usize;
14114 if idx >= inter {
14115 return 0.0; }
14117 let mut s = if need_scratch {
14118 vec![0.0f32; hidden]
14119 } else {
14120 Vec::new()
14121 };
14122 let gate = d.gate_proj.row_dot(idx, x, &mut s);
14123 let up = d.up_proj.row_dot(idx, x, &mut s);
14124 d.act.combine(gate, up)
14125 };
14126 match pool {
14127 Some(p) if n >= 256 => {
14128 let ptr = SendMut(act.as_mut_ptr());
14129 p.run(&|widx, nw| {
14130 let chunk = n.div_ceil(nw);
14131 let (s, e) = (widx * chunk, ((widx + 1) * chunk).min(n));
14132 for ai in s..e {
14133 unsafe { *ptr.at(ai) = compute(ai) };
14134 }
14135 });
14136 }
14137 _ => {
14138 for (ai, a) in act.iter_mut().enumerate() {
14139 *a = compute(ai);
14140 }
14141 }
14142 }
14143 let mut out = vec![0.0f32; hidden];
14145 for (ai, &idx) in active.iter().enumerate() {
14146 let w = act[ai];
14147 if w.abs() >= 1e-12 && (idx as usize) < inter {
14148 d.down_proj.add_col_scaled(idx as usize, w, &mut out);
14149 }
14150 }
14151 out
14152}
14153
14154#[doc(hidden)]
14156pub fn sparse_ffn_quant_for_test(
14157 d: &DenseFfn,
14158 x: &[f32],
14159 active: &[u16],
14160 hidden: usize,
14161) -> Vec<f32> {
14162 sparse_ffn_quant(d, x, active, hidden, None)
14163}
14164
14165fn dequant_dense_f32(d: &DenseFfn) -> (Vec<f32>, Vec<f32>, Vec<f32>) {
14169 let deq = |t: &QTensor| -> Vec<f32> {
14170 let (rows, cols) = (t.rows(), t.cols());
14171 let mut out = vec![0.0f32; rows * cols];
14172 for r in 0..rows {
14173 t.row_f32(r, &mut out[r * cols..(r + 1) * cols]);
14174 }
14175 out
14176 };
14177 (deq(&d.gate_proj), deq(&d.up_proj), deq(&d.down_proj))
14178}
14179
14180struct SendMut(*mut f32);
14182unsafe impl Send for SendMut {}
14183unsafe impl Sync for SendMut {}
14184impl SendMut {
14185 #[inline]
14186 #[allow(clippy::mut_from_ref)]
14189 unsafe fn at(&self, i: usize) -> &mut f32 {
14190 unsafe { &mut *self.0.add(i) }
14191 }
14192}
14193
14194pub(crate) fn moe_route(
14204 logits: &[f32],
14205 m: &MoeFfn,
14206 allowed: Option<&[bool]>,
14207) -> (Vec<usize>, Vec<f32>, f32) {
14208 let ne = logits.len();
14209 let p: Vec<f32> = if m.router_sigmoid {
14210 logits.iter().map(|&l| 1.0 / (1.0 + (-l).exp())).collect()
14211 } else {
14212 let mx = logits.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
14213 let mut e: Vec<f32> = logits.iter().map(|&l| (l - mx).exp()).collect();
14214 let s: f32 = e.iter().sum();
14215 for v in &mut e {
14216 *v /= s;
14217 }
14218 e
14219 };
14220 let admit = |e: usize| {
14226 m.mask.as_ref().is_none_or(|mk| mk[e])
14227 && allowed.is_none_or(|a| a.get(e).copied().unwrap_or(false))
14228 };
14229 let mut idx: Vec<usize> = (0..ne).filter(|&e| admit(e)).collect();
14230 match &m.expert_bias {
14232 Some(b) => idx.sort_unstable_by(|&x, &y| {
14233 (p[y] + b[y])
14234 .partial_cmp(&(p[x] + b[x]))
14235 .unwrap()
14236 .then(x.cmp(&y))
14237 }),
14238 None => idx.sort_unstable_by(|&x, &y| p[y].partial_cmp(&p[x]).unwrap().then(x.cmp(&y))),
14239 }
14240 idx.truncate(m.top_k);
14241 if let Some(tau) = m.route_tau {
14245 let total: f32 = idx.iter().map(|&e| p[e]).sum();
14246 if total > 0.0 {
14247 let mut acc = 0.0f32;
14248 let mut keep = idx.len();
14249 for (i, &e) in idx.iter().enumerate() {
14250 acc += p[e];
14251 if acc >= tau * total {
14252 keep = i + 1;
14253 break;
14254 }
14255 }
14256 idx.truncate(keep);
14257 }
14258 }
14259 let wsum: f32 = if m.norm_topk_prob {
14260 let s: f32 = idx.iter().map(|&e| p[e]).sum();
14261 (if m.router_sigmoid { s + 1e-6 } else { s }) / m.routed_scaling
14264 } else {
14265 1.0 / m.routed_scaling
14266 };
14267 (idx, p, wsum)
14268}
14269
14270fn moe_trace(idx: &[usize]) {
14272 moe_trace_at(crate::gpu::cur_layer() as i32, idx)
14273}
14274
14275pub(crate) fn moe_trace_at(li: i32, idx: &[usize]) {
14278 use std::io::Write;
14279 static F: std::sync::OnceLock<Option<std::sync::Mutex<std::fs::File>>> =
14280 std::sync::OnceLock::new();
14281 let Some(f) = F.get_or_init(|| {
14282 let p = std::env::var("CMF_MOE_TRACE").ok()?;
14283 Some(std::sync::Mutex::new(
14284 std::fs::OpenOptions::new()
14285 .create(true)
14286 .append(true)
14287 .open(p)
14288 .ok()?,
14289 ))
14290 }) else {
14291 return;
14292 };
14293 let ids: Vec<String> = idx.iter().map(|e| e.to_string()).collect();
14294 let _ = writeln!(f.lock().unwrap(), "{li}:{}", ids.join(","));
14295}
14296
14297pub(crate) fn moe_ffn(
14300 m: &MoeFfn,
14301 x: &[f32],
14302 pool: Option<&Pool>,
14303 allowed: Option<&[bool]>,
14304) -> Vec<f32> {
14305 accumulate_act(m, x, 1);
14306 let ne = m.experts.len();
14307 let mut logits = vec![0.0f32; ne];
14308 match &m.resonance {
14309 Some(r) => r.scores(x, &mut logits),
14310 None => m.router.matvec(x, &mut logits, pool),
14311 }
14312 let (idx, p, wsum) = moe_route(&logits, m, allowed);
14313 {
14314 let mut st = m.stats.borrow_mut();
14315 if st.len() < ne {
14316 st.resize(ne, 0);
14317 }
14318 for &e in &idx {
14319 st[e] += 1;
14320 }
14321 }
14322 moe_trace(&idx);
14328 if crate::gpu::enabled_here() {
14333 match crate::gpu::probe_arm(crate::gpu::OpClass::Ffn) {
14334 crate::gpu::ProbeArm::Gpu => {
14335 let t0 = std::time::Instant::now();
14336 if let Some(out) = moe_ffn_gpu(m, x, &idx, &p, wsum, pool) {
14337 crate::gpu::probe_record(crate::gpu::OpClass::Ffn, true, t0.elapsed());
14338 return out;
14339 }
14340 }
14341 crate::gpu::ProbeArm::CpuTimed => {
14342 let t0 = std::time::Instant::now();
14343 let out = crate::gpu::cpu_scope(|| moe_ffn_cpu(m, x, &idx, &p, wsum, pool));
14344 crate::gpu::probe_record(crate::gpu::OpClass::Ffn, false, t0.elapsed());
14345 return out;
14346 }
14347 crate::gpu::ProbeArm::Cpu => {
14348 return crate::gpu::cpu_scope(|| moe_ffn_cpu(m, x, &idx, &p, wsum, pool));
14349 }
14350 }
14351 }
14352 moe_ffn_cpu(m, x, &idx, &p, wsum, pool)
14353}
14354
14355fn graph_note(built: bool, layers_run: usize, total_layers: usize) {
14360 use std::sync::atomic::{AtomicBool, Ordering};
14361 if built {
14362 GRAPH_TOK_OK.fetch_add(1, Ordering::Relaxed);
14363 if total_layers > 0 && layers_run < total_layers {
14364 GRAPH_TOK_PREFIX.fetch_add(1, Ordering::Relaxed);
14365 } else {
14366 GRAPH_TOK_FULL.fetch_add(1, Ordering::Relaxed);
14367 }
14368 } else {
14369 GRAPH_TOK_MISS.fetch_add(1, Ordering::Relaxed);
14370 }
14371 static SAID: AtomicBool = AtomicBool::new(false);
14372 if !SAID.swap(true, Ordering::Relaxed) {
14373 if built {
14374 tracing::info!("wgpu whole-token graph: ACTIVE");
14375 } else {
14376 tracing::warn!("wgpu whole-token graph refused — per-op path");
14377 }
14378 }
14379}
14380
14381pub static GRAPH_TOK_OK: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
14385pub static GRAPH_TOK_MISS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
14386pub static GRAPH_TOK_PREFIX: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
14390pub static GRAPH_TOK_FULL: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
14392
14393pub static METAL_GRAPH_TOK_OK: std::sync::atomic::AtomicU64 =
14397 std::sync::atomic::AtomicU64::new(0);
14398pub static METAL_GRAPH_HEAD_OK: std::sync::atomic::AtomicU64 =
14399 std::sync::atomic::AtomicU64::new(0);
14400pub static METAL_GRAPH_HEAD_MISS: std::sync::atomic::AtomicU64 =
14401 std::sync::atomic::AtomicU64::new(0);
14402pub static METAL_GRAPH_LAYERS: std::sync::atomic::AtomicU64 =
14403 std::sync::atomic::AtomicU64::new(0);
14404pub static METAL_GRAPH_ERRORS: std::sync::atomic::AtomicU64 =
14405 std::sync::atomic::AtomicU64::new(0);
14406pub static METAL_PREFILL_CHUNKS: std::sync::atomic::AtomicU64 =
14410 std::sync::atomic::AtomicU64::new(0);
14411pub static METAL_PREFILL_ROWS: std::sync::atomic::AtomicU64 =
14412 std::sync::atomic::AtomicU64::new(0);
14413pub static METAL_PREFILL_HEAD_ROWS: std::sync::atomic::AtomicU64 =
14414 std::sync::atomic::AtomicU64::new(0);
14415pub static METAL_PREFILL_ERRORS: std::sync::atomic::AtomicU64 =
14416 std::sync::atomic::AtomicU64::new(0);
14417
14418fn moe_batch_enabled() -> bool {
14421 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
14422 *ON.get_or_init(|| std::env::var("CMF_MOE_BATCH").as_deref() != Ok("0"))
14423}
14424
14425fn moe_ffn_cpu_batched(
14431 m: &MoeFfn,
14432 x: &[f32],
14433 idx: &[usize],
14434 p: &[f32],
14435 wsum: f32,
14436 pool: Option<&Pool>,
14437) -> Option<Vec<f32>> {
14438 if idx.is_empty() || !moe_batch_enabled() {
14439 return None;
14440 }
14441 if FFN_PROBE.with(|pr| pr.borrow().is_some()) {
14445 return None;
14446 }
14447 let n = idx.len() + usize::from(m.shared.is_some());
14448 let mut pairs = Vec::with_capacity(n);
14449 let mut downs = Vec::with_capacity(n);
14450 let mut ws = Vec::with_capacity(n);
14451 for &e in idx {
14452 let d = &m.experts[e];
14453 if d.act != Act::Silu {
14454 return None;
14455 }
14456 pairs.push((&d.gate_proj, &d.up_proj));
14457 downs.push(&d.down_proj);
14458 ws.push(p[e] / wsum * m.per_expert_scale.as_ref().map_or(1.0, |v| v[e]));
14459 }
14460 if let Some((se, gate)) = &m.shared {
14463 if se.act != Act::Silu {
14464 return None;
14465 }
14466 let g = gate.as_ref().map_or(1.0, |gate| {
14467 let mut gl = [0.0f32; 1];
14468 gate.matvec(x, &mut gl, pool);
14469 1.0 / (1.0 + (-gl[0]).exp())
14470 });
14471 pairs.push((&se.gate_proj, &se.up_proj));
14472 downs.push(&se.down_proj);
14473 ws.push(g);
14474 }
14475 let inter = pairs[0].0.rows();
14476 let mut gs: Vec<Vec<f32>> = (0..pairs.len()).map(|_| vec![0f32; inter]).collect();
14477 if !QTensor::moe_gate_up_many(&pairs, x, &mut gs, pool) {
14478 return None;
14479 }
14480 let mut out = attention::take_buf(x.len());
14481 if !QTensor::moe_down_many(&downs, &gs, &ws, &mut out, pool) {
14482 attention::recycle_buf(&mut out);
14483 return None;
14484 }
14485 Some(out)
14486}
14487
14488pub(crate) fn moe_cold_experts_cpu(
14494 experts: &[(&DenseFfn, f32)],
14495 x: &[f32],
14496 pool: Option<&Pool>,
14497) -> Vec<f32> {
14498 let mut out = attention::take_buf(x.len());
14499 if experts.is_empty() {
14500 return out;
14501 }
14502 let pairs: Vec<_> = experts
14503 .iter()
14504 .map(|(e, _)| (&e.gate_proj, &e.up_proj))
14505 .collect();
14506 let downs: Vec<_> = experts.iter().map(|(e, _)| &e.down_proj).collect();
14507 let weights: Vec<_> = experts.iter().map(|(_, w)| *w).collect();
14508 let inter = experts[0].0.gate_proj.rows();
14509 let mut activations: Vec<Vec<f32>> = (0..experts.len()).map(|_| vec![0.0; inter]).collect();
14510 if QTensor::moe_gate_up_many(&pairs, x, &mut activations, pool)
14511 && QTensor::moe_down_many(&downs, &activations, &weights, &mut out, pool)
14512 {
14513 return out;
14514 }
14515 out.fill(0.0);
14516 for &(expert, weight) in experts {
14517 let mut one = dense_ffn(expert, x, pool);
14518 for (o, v) in out.iter_mut().zip(&one) {
14519 *o += weight * v;
14520 }
14521 attention::recycle_buf(&mut one);
14522 }
14523 out
14524}
14525
14526fn moe_ffn_cpu(
14528 m: &MoeFfn,
14529 x: &[f32],
14530 idx: &[usize],
14531 p: &[f32],
14532 wsum: f32,
14533 pool: Option<&Pool>,
14534) -> Vec<f32> {
14535 if let Some(out) = moe_ffn_cpu_batched(m, x, idx, p, wsum, pool) {
14536 return out;
14537 }
14538 let mut out = attention::take_buf(x.len());
14539 for &e in idx {
14540 let mut eo = dense_ffn(&m.experts[e], x, pool);
14541 let w = p[e] / wsum * m.per_expert_scale.as_ref().map_or(1.0, |v| v[e]);
14542 for i in 0..out.len() {
14543 out[i] += w * eo[i];
14544 }
14545 attention::recycle_buf(&mut eo);
14546 }
14547 if let Some((se, gate)) = &m.shared {
14548 let mut so = dense_ffn(se, x, pool);
14549 let g = gate.as_ref().map_or(1.0, |gate| {
14550 let mut gl = [0.0f32; 1];
14551 gate.matvec(x, &mut gl, pool);
14552 1.0 / (1.0 + (-gl[0]).exp())
14553 });
14554 for i in 0..out.len() {
14555 out[i] += g * so[i];
14556 }
14557 attention::recycle_buf(&mut so);
14558 }
14559 out
14560}
14561
14562#[allow(clippy::too_many_arguments)]
14570fn mla_attention(
14571 w: &MlaWeights,
14572 normed: &[f32],
14573 cache: &mut crate::kv_cache::LayerKvCache,
14574 position: usize,
14575 inv_freq: &[f32],
14576 rope_scale: f32,
14577 eps: f64,
14578 pool: Option<&Pool>,
14579) -> Vec<f32> {
14580 let (nh, dr, dn, dv, lora) = (w.nh, w.qk_rope, w.qk_nope, w.v_dim, w.lora);
14581 let hd = dr + dn;
14582 let mut q = vec![0.0f32; nh * hd];
14583 match (&w.q_a, &w.q_a_norm) {
14584 (Some(qa), Some(qn)) => {
14585 let mut t = vec![0.0f32; qa.rows()];
14586 qa.matvec(normed, &mut t, pool);
14587 let tn = inference::rms_norm(&t, qn, eps, NormStyle::Qwen);
14588 w.q_proj.matvec(&tn, &mut q, pool);
14589 }
14590 _ => w.q_proj.matvec(normed, &mut q, pool),
14591 }
14592 let mut ca = vec![0.0f32; lora + dr];
14593 w.kv_a.matvec(normed, &mut ca, pool);
14594 let (c_lat, k_rope) = ca.split_at_mut(lora);
14595 let latn = inference::rms_norm(c_lat, &w.kv_a_norm, eps, NormStyle::Qwen);
14596 let mut kvb = vec![0.0f32; nh * (dn + dv)];
14597 w.kv_b.matvec(&latn, &mut kvb, pool);
14598 if !w.nope {
14599 attention::rope_rotate_scaled(k_rope, position, inv_freq, rope_scale);
14600 }
14601 for h in 0..nh {
14602 if !w.nope {
14603 attention::rope_rotate_scaled(
14604 &mut q[h * hd..h * hd + dr],
14605 position,
14606 inv_freq,
14607 rope_scale,
14608 );
14609 }
14610 }
14611 let mut k = vec![0.0f32; nh * hd];
14612 let mut v = vec![0.0f32; nh * hd];
14613 for h in 0..nh {
14614 k[h * hd..h * hd + dr].copy_from_slice(k_rope);
14615 k[h * hd + dr..(h + 1) * hd].copy_from_slice(&kvb[h * (dn + dv)..h * (dn + dv) + dn]);
14616 v[h * hd..h * hd + dv].copy_from_slice(&kvb[h * (dn + dv) + dn..(h + 1) * (dn + dv)]);
14617 }
14618 cache.append(&k, &v, &vec![true; nh]);
14619 let (ao, mut imp) = attention::attend_all_heads(&q, cache, nh, 1, hd, w.scale, None, 0.0);
14620 attention::recycle_buf(&mut imp);
14621 let mut ov = vec![0.0f32; nh * dv];
14622 for h in 0..nh {
14623 ov[h * dv..(h + 1) * dv].copy_from_slice(&ao[h * hd..h * hd + dv]);
14624 }
14625 let mut out = vec![0.0f32; w.o_proj.rows()];
14626 w.o_proj.matvec(&ov, &mut out, pool);
14627 out
14628}
14629
14630fn dense_moe_ffn(
14637 dm: &DenseMoeFfn,
14638 x_normed: &[f32],
14639 h_raw: &[f32],
14640 eps: f64,
14641 norm_style: NormStyle,
14642 pool: Option<&Pool>,
14643) -> Vec<f32> {
14644 let mut d = dense_ffn(&dm.dense, x_normed, pool);
14645 d = inference::rms_norm(&d, &dm.post_norm_1, eps, norm_style);
14646 let m = &dm.moe;
14647 let ne = m.experts.len();
14648 let mut logits = vec![0.0f32; ne];
14649 if m.router_input_norm {
14650 let ss: f32 = h_raw.iter().map(|v| v * v).sum::<f32>() / h_raw.len() as f32;
14651 let inv = 1.0 / (ss + eps as f32).sqrt();
14652 let xr: Vec<f32> = h_raw.iter().map(|v| v * inv).collect();
14653 m.router.matvec(&xr, &mut logits, pool);
14654 } else {
14655 m.router.matvec(h_raw, &mut logits, pool);
14656 }
14657 let (idx, p, wsum) = moe_route(&logits, m, None);
14658 {
14659 let mut st = m.stats.borrow_mut();
14660 if st.len() < ne {
14661 st.resize(ne, 0);
14662 }
14663 for &e in &idx {
14664 st[e] += 1;
14665 }
14666 }
14667 let x2 = inference::rms_norm(h_raw, &dm.pre_norm_2, eps, norm_style);
14668 let mo = moe_ffn_cpu(m, &x2, &idx, &p, wsum, pool);
14669 let mo = inference::rms_norm(&mo, &dm.post_norm_2, eps, norm_style);
14670 for (di, mi) in d.iter_mut().zip(&mo) {
14671 *di += mi;
14672 }
14673 d
14674}
14675
14676fn moe_gpu_refused(why: &'static str) {
14683 use std::sync::atomic::{AtomicBool, Ordering};
14684 static SAID: AtomicBool = AtomicBool::new(false);
14685 if !SAID.swap(true, Ordering::Relaxed) {
14686 tracing::warn!("MoE GPU block refused ({why}) — experts run on the CPU");
14687 }
14688}
14689
14690fn moe_ffn_gpu(
14691 m: &MoeFfn,
14692 x: &[f32],
14693 idx: &[usize],
14694 p: &[f32],
14695 wsum: f32,
14696 pool: Option<&Pool>,
14697) -> Option<Vec<f32>> {
14698 use crate::gpu::MoeJob;
14699
14700 let mut jobs: Vec<MoeJob> = Vec::with_capacity(idx.len() + 1);
14701 let mut model_ref = None;
14702 for &e in idx {
14703 if moe_push_job(&m.experts[e], x, p[e] / wsum, &mut jobs, &mut model_ref).is_none() {
14704 moe_gpu_refused("push_job(expert)");
14705 return None;
14706 }
14707 }
14708 if let Some((se, gate)) = &m.shared {
14709 let g = gate.as_ref().map_or(1.0, |gate| {
14710 let mut gl = [0.0f32; 1];
14711 gate.matvec(x, &mut gl, pool);
14712 1.0 / (1.0 + (-gl[0]).exp())
14713 });
14714 if moe_push_job(se, x, g, &mut jobs, &mut model_ref).is_none() {
14715 moe_gpu_refused("push_job(shared)");
14716 return None;
14717 }
14718 }
14719 let Some(model) = model_ref else {
14720 moe_gpu_refused("no model_ref");
14721 return None;
14722 };
14723 let hidden = jobs[0].down.1;
14724 let mut out = vec![0.0f32; hidden];
14725 if crate::gpu::moe_block(&model, &jobs, &mut out) {
14726 Some(out)
14727 } else {
14728 moe_gpu_refused("gpu::moe_block");
14729 None
14730 }
14731}
14732
14733fn ffn_forward(
14735 ffn: &FfnKind,
14736 x: &[f32],
14737 pool: Option<&Pool>,
14738 experts_allowed: Option<&[bool]>,
14739) -> Vec<f32> {
14740 match ffn {
14741 FfnKind::Dense(d) if !d.segs.is_empty() => tube_ffn(d, x, 1, pool, None),
14742 FfnKind::Dense(d) => dense_ffn(d, x, pool),
14743 FfnKind::Moe(m) => moe_ffn(m, x, pool, experts_allowed),
14744 FfnKind::DenseMoe(_) => unreachable!("DenseMoe dispatches via dense_moe_ffn"),
14748 }
14749}
14750
14751fn ffn_forward_pair(
14755 ffn: &FfnKind,
14756 x1: &[f32],
14757 x2: &[f32],
14758 pool: Option<&Pool>,
14759 experts_allowed: Option<&[bool]>,
14760) -> (Vec<f32>, Vec<f32>) {
14761 let d = match ffn {
14762 FfnKind::Dense(d) if !d.segs.is_empty() => {
14765 return (
14766 tube_ffn(d, x1, 1, pool, None),
14767 tube_ffn(d, x2, 1, pool, None),
14768 );
14769 }
14770 FfnKind::Dense(d) => d,
14771 FfnKind::Moe(m) => {
14772 return (
14773 moe_ffn(m, x1, pool, experts_allowed),
14774 moe_ffn(m, x2, pool, experts_allowed),
14775 );
14776 }
14777 FfnKind::DenseMoe(_) => unreachable!("DenseMoe dispatches via dense_moe_ffn"),
14778 };
14779 let inter = d.gate_proj.rows();
14780 FFN_SCRATCH.with(|s| {
14781 let mut s = s.borrow_mut();
14782 let [g1, g2, u1, u2] = &mut *s;
14783 g1.resize(inter, 0.0);
14784 g2.resize(inter, 0.0);
14785 u1.resize(inter, 0.0);
14786 u2.resize(inter, 0.0);
14787 QTensor::matvec2_many(
14790 [&d.gate_proj, &d.up_proj],
14791 x1,
14792 x2,
14793 [g1.as_mut_slice(), u1.as_mut_slice()],
14794 [g2.as_mut_slice(), u2.as_mut_slice()],
14795 pool,
14796 );
14797 for i in 0..inter {
14798 g1[i] = d.act.combine(g1[i], u1[i]);
14799 g2[i] = d.act.combine(g2[i], u2[i]);
14800 }
14801 let mut o1 = attention::take_buf(d.down_proj.rows());
14802 let mut o2 = attention::take_buf(d.down_proj.rows());
14803 d.down_proj.matvec2(g1, g2, &mut o1, &mut o2, pool);
14804 (o1, o2)
14805 })
14806}
14807
14808#[cfg(test)]
14809mod tests {
14810
14811 #[test]
14812 fn nll_graph_policy_scopes_only_the_fused_head() {
14813 for (label, unmasked, prefer_graph, native_metal, want_graph, want_head) in [
14814 ("vulkan graph", true, true, false, true, false),
14816 ("native Metal graph", true, true, true, true, true),
14818 ("masked", false, true, false, false, false),
14820 ("graph disabled", true, false, true, false, false),
14821 ] {
14822 let (graph_quality, graph_head_required) =
14823 super::nll_graph_policy(unmasked, prefer_graph, native_metal);
14824 assert_eq!(graph_quality, want_graph, "{label}: graph quality");
14825 assert_eq!(graph_head_required, want_head, "{label}: fused head");
14826 }
14827 }
14828
14829 #[test]
14830 fn mtp_prefill_pair_boundaries_skip_only_final_prompt_row() {
14831 assert_eq!(mtp_prefill_pair_count(0, 128, 256), 128);
14832 assert_eq!(mtp_prefill_pair_count(128, 256, 256), 127);
14833 assert_eq!(mtp_prefill_pair_count(0, 256, 256), 255);
14834 assert_eq!(mtp_prefill_pair_count(256, 256, 256), 0);
14835 assert_eq!(mtp_prefill_pair_count(300, 320, 256), 0);
14836 }
14837
14838 #[test]
14839 fn cancel_flag_stops_generation() {
14840 let mut p = create_test_pipeline(16, 32, 2, 2, 8, 2, 32);
14841 p.cancel.store(true, std::sync::atomic::Ordering::Relaxed);
14844 let r = p.generate_from_ids(&[1, 2, 3], 8, None, None).unwrap();
14845 assert_eq!(r.finish_reason, "cancelled");
14846 assert!(
14847 r.token_ids.is_empty(),
14848 "no tokens after cancel: {:?}",
14849 r.token_ids
14850 );
14851 assert_eq!(p.kv_cache.seq_len(), 0);
14852 assert!(p.kv_history.is_empty());
14853 assert!(!p.graph_want_logits);
14854 assert!(p.graph_logits.is_none());
14855 let r2 = p.generate_from_ids(&[1, 2, 3], 4, None, None).unwrap();
14857 assert_ne!(r2.finish_reason, "cancelled");
14858 }
14859 use super::*;
14860
14861 #[test]
14869 fn dynamic_ffn_equals_the_zeroing_arm() {
14870 let (hidden, inter) = (8usize, 32usize);
14871 let synth = |n: usize, salt: usize| -> Vec<f32> {
14872 (0..n)
14873 .map(|i| (((i * 29 + salt * 13 + 7) % 89) as f32 / 89.0 - 0.5) * 0.6)
14874 .collect()
14875 };
14876 let down = synth(hidden * inter, 3);
14877 let mut down_t = vec![0.0f32; inter * hidden];
14878 for r in 0..hidden {
14879 for c in 0..inter {
14880 down_t[c * hidden + r] = down[r * inter + c];
14881 }
14882 }
14883 let d = DenseFfn {
14884 gate_proj: QTensor::from_f32(synth(inter * hidden, 1), inter, hidden),
14885 up_proj: QTensor::from_f32(synth(inter * hidden, 2), inter, hidden),
14886 down_proj: QTensor::from_f32(down.clone(), hidden, inter),
14887 act: Act::Silu,
14888 down_t: Some(QTensor::from_f32(down_t, inter, hidden)),
14889 segs: Vec::new(),
14890 };
14891 let x = synth(hidden, 11);
14892 let k = 12usize;
14893 let got = dense_ffn_dynamic(&d, &x, None, k).expect("down_t present");
14894 let mut g = vec![0.0f32; inter];
14896 d.gate_proj.matvec(&x, &mut g, None);
14897 let mut u = vec![0.0f32; inter];
14898 d.up_proj.matvec(&x, &mut u, None);
14899 for v in g.iter_mut() {
14900 *v = inference::silu(*v);
14901 }
14902 keep_top_k(&mut g, k);
14903 for i in 0..inter {
14904 g[i] *= u[i];
14905 }
14906 let mut want = vec![0.0f32; hidden];
14907 d.down_proj.matvec(&g, &mut want, None);
14908 for (a, b) in want.iter().zip(&got) {
14909 assert!((a - b).abs() < 1e-5, "dynamic {b} vs reference {a}");
14910 }
14911 }
14912
14913 #[test]
14919 fn tube_ffn_open_equals_dense_and_closed_equals_masked() {
14920 let (hidden, core, tube) = (8usize, 12usize, 8usize);
14921 let inter = core + tube;
14922 let synth = |n: usize, salt: usize| -> Vec<f32> {
14923 (0..n)
14924 .map(|i| (((i * 41 + salt * 17 + 5) % 97) as f32 / 97.0 - 0.5) * 0.5)
14925 .collect()
14926 };
14927 let (g_all, u_all) = (synth(inter * hidden, 1), synth(inter * hidden, 2));
14928 let d_all = synth(hidden * inter, 3);
14929 let dense = DenseFfn {
14931 gate_proj: QTensor::from_f32(g_all.clone(), inter, hidden),
14932 up_proj: QTensor::from_f32(u_all.clone(), inter, hidden),
14933 down_proj: QTensor::from_f32(d_all.clone(), hidden, inter),
14934 act: Act::Silu,
14935 down_t: None,
14936 segs: Vec::new(),
14937 };
14938 let rows =
14939 |v: &[f32], a: usize, b: usize| -> Vec<f32> { v[a * hidden..b * hidden].to_vec() };
14940 let cols = |v: &[f32], a: usize, b: usize| -> Vec<f32> {
14941 let mut o = Vec::with_capacity(hidden * (b - a));
14942 for r in 0..hidden {
14943 o.extend_from_slice(&v[r * inter + a..r * inter + b]);
14944 }
14945 o
14946 };
14947 let tubed = DenseFfn {
14948 down_t: None,
14949 gate_proj: QTensor::from_f32(rows(&g_all, 0, core), core, hidden),
14950 up_proj: QTensor::from_f32(rows(&u_all, 0, core), core, hidden),
14951 down_proj: QTensor::from_f32(cols(&d_all, 0, core), hidden, core),
14952 act: Act::Silu,
14953 segs: vec![FfnSeg {
14954 gate: QTensor::from_f32(rows(&g_all, core, inter), tube, hidden),
14955 up: QTensor::from_f32(rows(&u_all, core, inter), tube, hidden),
14956 down: QTensor::from_f32(cols(&d_all, core, inter), hidden, tube),
14957 start: core,
14958 width: tube,
14959 }],
14960 };
14961 let x = synth(hidden, 7);
14962 let want = dense_ffn(&dense, &x, None);
14963 let got = tube_ffn(&tubed, &x, 1, None, None);
14964 for (a, b) in want.iter().zip(&got) {
14965 assert!((a - b).abs() < 1e-5, "open tube: {a} vs {b}");
14966 }
14967 let mut bits = vec![0u8; inter.div_ceil(8)];
14969 for n in 0..core {
14970 bits[n / 8] |= 1 << (n % 8);
14971 }
14972 let closed = tube_ffn(&tubed, &x, 1, None, Some(&bits));
14973 let masked = dense_ffn_masked(&dense, &x, None, &bits);
14974 for (a, b) in masked.iter().zip(&closed) {
14975 assert!((a - b).abs() < 1e-5, "closed tube: {a} vs {b}");
14976 }
14977 let batch = tube_ffn(&tubed, &x, 1, None, Some(&bits));
14979 for (a, b) in closed.iter().zip(&batch) {
14980 assert_eq!(a, b, "batch arm disagrees with decode arm");
14981 }
14982 }
14983
14984 #[test]
14986 fn sparse_ffn_quant_equals_dense_with_inactive_zeroed() {
14987 let (hidden, inter) = (16usize, 40usize);
14988 let synth = |n: usize, salt: usize| -> Vec<f32> {
14989 (0..n)
14990 .map(|i| (((i * 37 + salt * 11 + 3) % 101) as f32 / 101.0 - 0.5) * 0.4)
14991 .collect()
14992 };
14993 let d = DenseFfn {
14994 gate_proj: QTensor::from_f32(synth(inter * hidden, 1), inter, hidden),
14995 up_proj: QTensor::from_f32(synth(inter * hidden, 2), inter, hidden),
14996 down_proj: QTensor::from_f32(synth(hidden * inter, 3), hidden, inter),
14997 act: Act::Silu,
14998 down_t: None,
14999 segs: Vec::new(),
15000 };
15001 let x = synth(hidden, 9);
15002 let active: Vec<u16> = (0..inter as u16).filter(|i| i % 3 == 0).collect();
15004
15005 let sparse = sparse_ffn_quant(&d, &x, &active, hidden, None);
15006
15007 let mut g = vec![0.0f32; inter];
15009 d.gate_proj.matvec(&x, &mut g, None);
15010 let mut u = vec![0.0f32; inter];
15011 d.up_proj.matvec(&x, &mut u, None);
15012 let act_set: std::collections::HashSet<u16> = active.iter().copied().collect();
15013 for i in 0..inter {
15014 g[i] = if act_set.contains(&(i as u16)) {
15015 inference::silu(g[i]) * u[i]
15016 } else {
15017 0.0
15018 };
15019 }
15020 let mut reference = vec![0.0f32; hidden];
15021 d.down_proj.matvec(&g, &mut reference, None);
15022
15023 let max_d = sparse
15024 .iter()
15025 .zip(&reference)
15026 .map(|(a, b)| (a - b).abs())
15027 .fold(0.0f32, f32::max);
15028 assert!(max_d < 1e-5, "sparse != dense-zeroed: max|Δ| = {max_d}");
15029 }
15030
15031 fn attach_test_mtp(p: &mut Pipeline) {
15033 let (h, inter, heads, kv, hd) = (
15034 p.hidden_size,
15035 p.intermediate_size,
15036 p.num_heads,
15037 p.num_kv_heads,
15038 p.head_dim,
15039 );
15040 let synth = |n: usize, salt: usize| -> Vec<f32> {
15041 (0..n)
15042 .map(|i| (((i * 29 + salt * 23 + 5) % 101) as f32 / 101.0 - 0.5) * 0.2)
15043 .collect()
15044 };
15045 let qt = |rows: usize, cols: usize, salt: usize| -> QTensor {
15046 QTensor::from_f32(synth(rows * cols, salt), rows, cols)
15047 };
15048 p.mtp = Some(MtpModule {
15049 enorm: vec![1.0; h],
15050 hnorm: vec![1.0; h],
15051 eh_proj: qt(h, 2 * h, 301),
15052 layer: LayerWeights {
15053 input_norm: vec![1.0; h],
15054 post_norm: vec![1.0; h],
15055 attn_out_norm: None,
15056 ffn_out_norm: None,
15057 layer_scale: None,
15058 ffn: FfnKind::Dense(DenseFfn {
15059 gate_proj: qt(inter, h, 315),
15060 up_proj: qt(inter, h, 316),
15061 down_proj: qt(h, inter, 317),
15062 act: Act::Silu,
15063 down_t: None,
15064 segs: Vec::new(),
15065 }),
15066 attn: AttnKind::Full {
15067 bias: None,
15068 wq: qt(heads * hd, h, 311),
15069 wk: qt(kv * hd, h, 312),
15070 wv: qt(kv * hd, h, 313),
15071 wo: qt(h, heads * hd, 314),
15072 q_norm: None,
15073 k_norm: None,
15074 output_gate: false,
15075 softplus_gate: None,
15076 },
15077 },
15078 final_norm: vec![1.0; h],
15079 kv: crate::kv_cache::LayerKvCache::new(kv, hd),
15080 });
15081 }
15082
15083 #[test]
15084 fn speculative_equals_vanilla_greedy() {
15085 unsafe { std::env::set_var("CMF_GPU_WGPU_GRAPH", "0") };
15089 let run = |spec: bool| {
15090 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
15091 p.sampler_config.temperature = 0.0;
15092 attach_test_mtp(&mut p);
15093 p.speculative = spec;
15094 let r = p.generate("abcdef", 12, None, None).unwrap();
15095 (r.token_ids, r.mtp_drafted, r.mtp_accepted)
15096 };
15097 let (vanilla, d0, _) = run(false);
15098 let (spec, d1, a1) = run(true);
15099 assert_eq!(d0, 0, "vanilla path must not draft");
15100 assert!(d1 > 0, "speculative path must draft");
15101 assert_eq!(
15102 vanilla, spec,
15103 "speculative must reproduce the exact greedy sequence (accepted {a1}/{d1})"
15104 );
15105 }
15106
15107 #[test]
15108 fn speculative_accepts_constant_oracle() {
15109 unsafe { std::env::set_var("CMF_GPU_WGPU_GRAPH", "0") };
15111 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
15112 p.sampler_config.temperature = 0.0;
15113 p.sampler_config.repetition_penalty = 1.0;
15114 p.weights.lm_head = QTensor::from_f32(vec![0.01; 64 * 8], 64, 8);
15117 attach_test_mtp(&mut p);
15118 p.speculative = true;
15119 let r = p.generate("abcd", 10, None, None).unwrap();
15120 assert!(r.mtp_drafted > 0);
15121 assert_eq!(
15122 r.mtp_accepted, r.mtp_drafted,
15123 "constant logits → every draft accepted"
15124 );
15125 assert!(r.token_ids.windows(2).all(|w| w[0] == w[1]));
15128 }
15129
15130 #[test]
15131 fn empty_prompt_is_an_error_not_a_panic() {
15132 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 260);
15133 let r = p.generate("", 4, None, None);
15134 assert!(r.is_err(), "empty prompt must be a clean error");
15135 }
15136
15137 #[test]
15138 fn every_token_enters_kv_exactly_once() {
15139 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
15140 p.sampler_config.temperature = 0.0;
15142 let r = p.generate("abc", 2, None, None).unwrap();
15143 assert_eq!(r.prompt_tokens, 3);
15144 assert_eq!(
15148 p.kv_cache.seq_len(),
15149 3 + r.tokens_generated - 1,
15150 "each token must be cached exactly once (v1 cached the last prompt token twice)"
15151 );
15152 }
15153
15154 #[test]
15155 fn generation_is_reproducible_with_seed() {
15156 let run = || {
15157 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
15158 p.generate("hello", 8, None, None).unwrap().token_ids
15159 };
15160 assert_eq!(run(), run());
15161 }
15162
15163 #[test]
15164 fn resetting_sampler_restarts_the_seeded_stream() {
15165 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
15166 let config = SamplerConfig {
15167 seed: Some(1234),
15168 ..SamplerConfig::default()
15169 };
15170 p.set_sampler_config(config.clone());
15171 let first = p.generate("hello", 8, None, None).unwrap().token_ids;
15172 p.set_sampler_config(config);
15173 let second = p.generate("hello", 8, None, None).unwrap().token_ids;
15174 assert_eq!(first, second);
15175 }
15176
15177 #[test]
15178 fn eviction_bounds_the_cache() {
15179 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 260);
15180 p.kv_cache.max_seq_len = 6;
15181 p.sampler_config.temperature = 0.0;
15182 let _ = p.generate("abcd", 12, None, None).unwrap();
15183 assert!(
15184 p.kv_cache.seq_len() <= 6 + 1,
15185 "cache must stay bounded by max_seq_len (got {})",
15186 p.kv_cache.seq_len()
15187 );
15188 }
15189
15190 #[test]
15191 fn confidence_matches_tokens_and_is_a_probability() {
15192 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
15193 p.sampler_config.temperature = 0.0;
15194 p.sampler_config.repetition_penalty = 1.0;
15195 let r = p.generate("abcd", 10, None, None).unwrap();
15196 assert_eq!(
15197 r.token_confidence.len(),
15198 r.token_ids.len(),
15199 "one confidence per emitted token"
15200 );
15201 for &c in &r.token_confidence {
15202 assert!((0.0..=1.0).contains(&c), "confidence out of [0,1]: {c}");
15203 }
15204 let logits = [1.0f32, 3.0, 0.5, 3.0];
15206 let p0 = top1_prob_t(&logits, 1, 1.0);
15207 let p1 = top1_prob_t(&logits, 3, 1.0);
15208 assert!((p0 - p1).abs() < 1e-6, "equal logits → equal prob");
15209 assert!(p0 > 0.0 && p0 < 1.0);
15210 let sharp = top1_prob_t(&logits, 1, 1.0);
15212 let soft = top1_prob_t(&logits, 1, 2.0);
15213 assert!(soft < sharp, "higher temperature lowers peak confidence");
15214 }
15215
15216 #[test]
15217 fn trace_is_opt_in_and_parallels_the_output() {
15218 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
15220 p.sampler_config.temperature = 0.0;
15221 p.sampler_config.repetition_penalty = 1.0;
15222 let r = p.generate("abcd", 10, None, None).unwrap();
15223 assert!(r.traces.is_empty(), "trace must be empty unless enabled");
15224
15225 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
15227 p.sampler_config.temperature = 0.0;
15228 p.sampler_config.repetition_penalty = 1.0;
15229 p.set_trace(true);
15230 let r = p.generate("abcd", 10, None, None).unwrap();
15231 assert_eq!(r.traces.len(), r.token_ids.len(), "one trace row per token");
15232 for (i, tr) in r.traces.iter().enumerate() {
15233 assert_eq!(tr.t, i, "trace index is sequential");
15234 assert_eq!(tr.token_id, r.token_ids[i], "trace token_id matches output");
15235 assert_eq!(
15236 tr.confidence, r.token_confidence[i],
15237 "trace confidence matches the confidence channel"
15238 );
15239 assert!(tr.active_skill.is_none() && tr.recon.is_none() && !tr.switched);
15241 }
15242 }
15243
15244 #[test]
15245 fn explain_prefill_logits_match_greedy_first_token() {
15246 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
15250 p.sampler_config.temperature = 0.0;
15251 p.sampler_config.repetition_penalty = 1.0;
15252 let ids = p.tokenizer.encode("abcd");
15253 let logits = p.prefill_next_logits(&ids, None);
15254 let argmax = logits
15255 .iter()
15256 .enumerate()
15257 .max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
15258 .unwrap()
15259 .0 as u32;
15260 let r = p.generate("abcd", 1, None, None).unwrap();
15261 assert_eq!(
15262 argmax, r.token_ids[0],
15263 "explain preview must match greedy emit"
15264 );
15265 }
15266
15267 #[test]
15268 fn laguna_shared_expert_is_unconditionally_added() {
15269 let matrix = |values: Vec<f32>| QTensor::from_f32(values, 2, 2);
15270 let identity = || matrix(vec![1.0, 0.0, 0.0, 1.0]);
15271 let zero_dense = || DenseFfn {
15272 gate_proj: matrix(vec![0.0; 4]),
15273 up_proj: matrix(vec![0.0; 4]),
15274 down_proj: matrix(vec![0.0; 4]),
15275 act: Act::Silu,
15276 down_t: None,
15277 segs: Vec::new(),
15278 };
15279 let shared = DenseFfn {
15280 gate_proj: identity(),
15281 up_proj: identity(),
15282 down_proj: identity(),
15283 act: Act::Silu,
15284 down_t: None,
15285 segs: Vec::new(),
15286 };
15287 let x = [1.0, 2.0];
15288 let expected = dense_ffn(&shared, &x, None);
15289 let moe = MoeFfn {
15290 router: QTensor::from_f32(vec![0.0, 0.0], 1, 2),
15291 experts: vec![zero_dense()],
15292 top_k: 1,
15293 norm_topk_prob: true,
15294 router_sigmoid: true,
15295 expert_bias: None,
15296 routed_scaling: 1.0,
15297 route_tau: None,
15298 shared: Some((shared, None)),
15299 stats: std::cell::RefCell::new(Vec::new()),
15300 act_sq: std::cell::RefCell::new(Vec::new()),
15301 act_rows: std::cell::RefCell::new(Vec::new()),
15302 mask: None,
15303 per_expert_scale: None,
15304 router_input_norm: false,
15305 resonance: None,
15306 };
15307 let actual = moe_ffn_cpu(&moe, &x, &[0], &[0.0], 1.0, None);
15308 for (actual, expected) in actual.iter().zip(expected) {
15309 assert!((actual - expected).abs() < 1e-6);
15310 }
15311 }
15312
15313 #[test]
15314 fn o1_batch_transition_publishes_one_epoch_before_serial_handoff() {
15315 const B: usize = 19;
15316 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
15317 p.set_o1(Some(crate::nystrom::O1Cfg {
15318 layers: crate::nystrom::O1Layers::All,
15319 m: 4,
15320 w: 8,
15321 sink: 2,
15322 rect: crate::nystrom::O1Rect::Aggregate,
15323 }));
15324 p.o1_begin_with_prefix(Some(B));
15325 let ids: Vec<u32> = (0..B as u32).collect();
15326 let _ = p.prefill_batch_span(PrefillIn::Ids(&ids), 0, None, 0, p.num_layers);
15327
15328 assert_eq!(p.o1_epoch, 1, "all layers publish one completed transition");
15329 assert!(p.kv_cache.layers.iter().all(|l| l.o1_sealed()));
15330 let next = p.embed_single(B as u32);
15331 let _ = p.forward_layers(&next, B, None);
15332 assert_eq!(p.o1_epoch, 1, "sealed handoff must not republish the epoch");
15333 }
15334
15335 #[test]
15336 fn o1_pair_transition_commits_scratch_before_epoch_publication() {
15337 const B: usize = 19;
15338 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
15339 let gdn_cfg = crate::linear_core::GdnCfg {
15343 num_v_heads: 2,
15344 num_k_heads: 1,
15345 key_head_dim: 2,
15346 value_head_dim: 4,
15347 conv_kernel: 3,
15348 hidden_size: 8,
15349 rms_eps: 1e-6,
15350 output_gate_sigmoid: false,
15351 };
15352 let synth = |n: usize, salt: usize| -> Vec<f32> {
15353 (0..n)
15354 .map(|i| (((i * 13 + salt * 7) % 97) as f32 / 97.0 - 0.5) * 0.4)
15355 .collect()
15356 };
15357 let qt = |rows: usize, cols: usize, salt: usize| {
15358 crate::qtensor::QTensor::from_f32(synth(rows * cols, salt), rows, cols)
15359 };
15360 let c_dim = gdn_cfg.conv_dim();
15361 let vd = gdn_cfg.num_v_heads * gdn_cfg.value_head_dim;
15362 p.weights.layers[0].attn = AttnKind::LinearGdn(crate::linear_core::GdnWeights {
15363 in_proj_qkv: qt(c_dim, 8, 1),
15364 in_proj_z: qt(vd, 8, 2),
15365 in_proj_a: qt(gdn_cfg.num_v_heads, 8, 3),
15366 in_proj_b: qt(gdn_cfg.num_v_heads, 8, 4),
15367 conv1d: synth(c_dim * gdn_cfg.conv_kernel, 5),
15368 a_log: vec![0.2, 0.5],
15369 dt_bias: synth(gdn_cfg.num_v_heads, 6),
15370 norm: vec![1.0; gdn_cfg.value_head_dim],
15371 out_proj: qt(8, vd, 7),
15372 });
15373 p.gdn_cfg = Some(gdn_cfg);
15374 p.set_o1(Some(crate::nystrom::O1Cfg {
15375 layers: crate::nystrom::O1Layers::All,
15376 m: 4,
15377 w: 8,
15378 sink: 2,
15379 rect: crate::nystrom::O1Rect::Aggregate,
15380 }));
15381 p.o1_begin_with_prefix(Some(B));
15382 for pos in 0..B - 2 {
15383 let emb = p.embed_single(pos as u32);
15384 let _ = p.forward_layers(&emb, pos, None);
15385 }
15386 let lane1_state = p.kv_cache.layers[0].linear_state.clone();
15387
15388 let e1 = p.embed_single((B - 2) as u32);
15389 let e2 = p.embed_single((B - 1) as u32);
15390 let _ = p.forward_pair(&e1, &e2, B - 2);
15391
15392 assert_eq!(p.o1_epoch, 1, "pair crossing B publishes one epoch");
15393 assert!(
15394 p.kv_cache
15395 .layers
15396 .iter()
15397 .enumerate()
15398 .all(|(li, l)| !p.o1_flags[li] || l.o1_sealed())
15399 );
15400 assert!(!p.kv_cache.layers[0].linear_state.is_empty());
15401 assert_ne!(
15402 p.kv_cache.layers[0].linear_state, lane1_state,
15403 "real pair must commit GDN lane 2 before returning"
15404 );
15405 assert!(p.kv_cache.layers[0].linear_scratch.is_empty());
15406 let next = p.embed_single(B as u32);
15407 let _ = p.forward_layers(&next, B, None);
15408 assert_eq!(p.o1_epoch, 1, "serial continuation must reuse the epoch");
15409 }
15410
15411 #[test]
15412 fn o1_error_observation_stays_terminal_until_reset() {
15413 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
15414 p.set_o1(Some(crate::nystrom::O1Cfg {
15415 layers: crate::nystrom::O1Layers::All,
15416 m: 4,
15417 w: 8,
15418 sink: 2,
15419 rect: crate::nystrom::O1Rect::Aggregate,
15420 }));
15421 p.o1_begin();
15422 p.kv_cache.layers[0].o1_abort("synthetic transition failure".into());
15423
15424 assert!(p.o1_seal_checked().is_err());
15425 assert!(
15426 p.o1_seal_checked().is_err(),
15427 "retry must see the sticky error"
15428 );
15429 let k = vec![0.2f32; 4];
15430 let v = vec![0.3f32; 4];
15431 p.kv_cache.layers[0].append(&k, &v, &[]);
15432 assert_eq!(p.kv_cache.layers[0].seq_len, 0);
15433
15434 p.reset_session();
15435 p.o1_begin();
15436 p.kv_cache.layers[0].append(&k, &v, &[]);
15437 assert_eq!(p.kv_cache.layers[0].seq_len, 1);
15438 }
15439
15440 #[test]
15441 fn nll_graph_failure_is_terminal_and_request_is_reusable() {
15442 let ids = vec![1u32, 2, 3, 4, 5, 6];
15443 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
15444 p.graph_logits = Some(vec![123.0]);
15445 p.graph_want_logits = true;
15446 p.graph_failed
15447 .store(true, std::sync::atomic::Ordering::Relaxed);
15448 p.cancel.store(true, std::sync::atomic::Ordering::Relaxed);
15449 let err = p.nll_ids_from(&ids, 0).expect_err("prior graph failure");
15450 assert!(err.contains("before NLL"));
15451 assert!(p.graph_logits.is_none());
15452 assert!(!p.graph_want_logits);
15453 assert!(!p.graph_failed.load(std::sync::atomic::Ordering::Relaxed));
15454 assert!(!p.cancel.load(std::sync::atomic::Ordering::Relaxed));
15455
15456 let mut fresh = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
15457 let expected = fresh.nll_ids_from(&ids, 0).expect("fresh NLL");
15458 let actual = p.nll_ids_from(&ids, 0).expect("reused NLL");
15459 assert_eq!(actual.1, expected.1);
15460 assert!((actual.0 - expected.0).abs() < 1e-9);
15461 }
15462
15463 #[test]
15464 fn nll_forward_failure_discards_partial_score_and_clears_sidechannels() {
15465 let ids = vec![1u32, 2, 3, 4, 5, 6];
15466 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
15467 p.nll_test_fail_at = Some(1);
15468 let err = p
15469 .nll_ids_from(&ids, 0)
15470 .expect_err("one-shot forward failure");
15471 assert!(err.contains("forward") || err.contains("score row"));
15472 assert!(!p.graph_failed.load(std::sync::atomic::Ordering::Relaxed));
15473 assert!(!p.graph_want_logits);
15474 assert!(p.graph_logits.is_none());
15475 assert!(p.kv_history.is_empty());
15476
15477 let mut fresh = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
15478 let expected = fresh.nll_ids_from(&ids, 0).expect("fresh NLL");
15479 let actual = p.nll_ids_from(&ids, 0).expect("reused NLL");
15480 assert_eq!(actual.1, expected.1);
15481 assert!((actual.0 - expected.0).abs() < 1e-9);
15482 }
15483
15484 #[test]
15485 fn nll_serial_failure_before_first_row_is_reported() {
15486 let ids = vec![1u32, 2, 3, 4];
15487 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
15488 p.nll_test_force_serial = true;
15489 p.nll_test_fail_at = Some(0);
15490 let err = p.nll_ids_from(&ids, 0).expect_err("serial forward failure");
15491 assert!(err.contains("serial forward"));
15492 assert!(p.kv_history.is_empty());
15493 assert!(!p.graph_failed.load(std::sync::atomic::Ordering::Relaxed));
15494 assert!(!p.cancel.load(std::sync::atomic::Ordering::Relaxed));
15495 }
15496
15497 #[test]
15498 fn ffn_probe_failure_discards_recorder_and_state() {
15499 let ids = vec![1u32, 2, 3, 4];
15500 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
15501 p.nll_test_fail_at = Some(0);
15502 let err = p
15503 .probe_ffn_mass_batch(&ids)
15504 .expect_err("probe forward failure");
15505 assert!(err.contains("NLL"));
15506 assert!(FFN_PROBE.with(|probe| probe.borrow().is_none()));
15507 assert!(p.kv_history.is_empty());
15508 assert!(!p.graph_failed.load(std::sync::atomic::Ordering::Relaxed));
15509 }
15510
15511 #[test]
15512 fn nll_test_controls_are_pipeline_scoped() {
15513 let ids = vec![1u32, 2, 3, 4];
15514 let mut failing = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
15515 let mut unaffected = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
15516 failing.nll_test_force_serial = true;
15517 failing.nll_test_fail_at = Some(0);
15518
15519 assert!(!failing.can_prefill_batched());
15520 assert!(unaffected.can_prefill_batched());
15521 let expected = unaffected
15522 .nll_ids_from(&ids, 0)
15523 .expect("unaffected pipeline remains usable");
15524 let err = failing
15525 .nll_ids_from(&ids, 0)
15526 .expect_err("failure injection belongs to failing pipeline");
15527 assert!(err.contains("serial forward"));
15528 assert!(failing.nll_test_fail_at.is_none());
15529 assert!(unaffected.can_prefill_batched());
15530 let actual = unaffected
15531 .nll_ids_from(&ids, 0)
15532 .expect("unaffected pipeline remains reusable");
15533 assert_eq!(actual.1, expected.1);
15534 assert!((actual.0 - expected.0).abs() < 1e-9);
15535 }
15536
15537 #[test]
15538 fn forward_ids_failure_channel_is_terminal_and_reusable() {
15539 let ids = vec![1u32, 2, 3, 4, 5, 6];
15540 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
15541 p.graph_logits = Some(vec![123.0]);
15542 p.graph_want_logits = true;
15543 p.graph_failed
15544 .store(true, std::sync::atomic::Ordering::Relaxed);
15545 p.cancel.store(true, std::sync::atomic::Ordering::Relaxed);
15546
15547 let err = p
15548 .forward_ids(&ids, None)
15549 .expect_err("a failed forward must not become a valid head result");
15550 assert!(err.contains("forward_ids setup"));
15551 assert!(p.graph_logits.is_none());
15552 assert!(!p.graph_want_logits);
15553 assert!(!p.graph_failed.load(std::sync::atomic::Ordering::Relaxed));
15554 assert!(!p.cancel.load(std::sync::atomic::Ordering::Relaxed));
15555 assert_eq!(p.kv_cache.seq_len(), 0);
15556
15557 let expected = create_test_pipeline(8, 16, 2, 1, 4, 1, 64)
15558 .forward_ids(&ids, None)
15559 .expect("fresh forward_ids");
15560 let actual = p
15561 .forward_ids(&ids, None)
15562 .expect("pipeline remains reusable after a failed forward");
15563 assert_eq!(actual.len(), expected.len());
15564 assert!(
15565 actual
15566 .iter()
15567 .zip(expected)
15568 .all(|(a, b)| (a - b).abs() < 1e-9)
15569 );
15570 assert_eq!(p.kv_cache.seq_len(), ids.len());
15571 }
15572}