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 pub kv_history: Vec<u32>,
96 pub kda_cfg: Option<crate::linear_core::KdaCfg>,
98 pub g3n: Option<Box<(crate::g3n::G3nGlobals, Vec<crate::g3n::G3nLayer>)>>,
101 pub dsv4: Option<
105 Box<(
106 crate::dsv4::Dsv4Globals,
107 Vec<crate::dsv4::Dsv4Layer>,
108 crate::dsv4::Dsv4Cfg,
109 crate::dsv4::Dsv4State,
110 )>,
111 >,
112 pub dsv4_mtp: Vec<crate::dsv4::Dsv4Mtp>,
116 pub dspark: Option<crate::dsv4::DsparkState>,
118 pub dspark_pending: Vec<(usize, Vec<u32>, bool, usize)>,
121 pub dspark_hist: Vec<usize>,
123 pub dspark_real: Vec<u32>,
127 pub dspark_trunk_picks: Vec<Vec<(usize, Vec<usize>)>>,
131 pub dspark_exp: Vec<(usize, usize, usize, usize)>,
133 pub dspark_draft_ns: u128,
137 pub short_conv_cfg: Option<ShortConvCfg>,
140 pub mtp: Option<MtpModule>,
142 pub speculative: bool,
144 rng: SplitMix64,
145 sampler_scratch: SamplerScratch,
146 spec_forced: Option<u32>,
152 spec_q: Vec<Vec<f32>>,
153 spec_p: Vec<f32>,
154 spec_res: Vec<f32>,
155 spec_qs: Vec<sampler::Sparse>,
157 spec_ps: sampler::Sparse,
158 spec_ress: sampler::Sparse,
159 mtp_graph_mode: Option<bool>,
166 #[cfg(target_os = "macos")]
169 metal_verify: Option<MetalVerifyPending>,
170 pub(crate) inv_freq: std::sync::Arc<Vec<f32>>,
174 ws: ForwardScratch,
178 pool: Option<std::sync::Arc<Pool>>,
180 pub(crate) model: Option<std::sync::Arc<cortiq_core::CmfModel>>,
184 pub(crate) dyn_force_f32: bool,
186 pub(crate) dyn_skill_layers: Vec<Option<Vec<usize>>>,
191 pub(crate) dyn_active: Option<usize>,
197 pub(crate) dyn_blend_loaded: bool,
201 pub(crate) dyn_phi_layer: Option<usize>,
204 dyn_phi_ema: Vec<f32>,
206 dyn_phi_seen: usize,
207 pub dyn_router: Option<crate::swarm::DynRouter>,
210 o1_cfg: Option<crate::nystrom::O1Cfg>,
213 o1_epoch: u64,
216 o1_flags: Vec<bool>,
218 trace: bool,
221 calib_temp: f32,
224 #[cfg_attr(not(target_os = "macos"), allow(dead_code))]
226 graph_kv_id: u64,
227 #[cfg_attr(not(target_os = "macos"), allow(dead_code))]
230 graph_want_logits: bool,
231 graph_logits: Option<Vec<f32>>,
234 pub embed_multiplier: f32,
236 pub attn_scale: f32,
239 pub swa: Option<(usize, usize)>,
242 pub sliding_layers: Option<Vec<bool>>,
245 pub inv_freq_local: Option<std::sync::Arc<Vec<f32>>>,
248 pub rotary_dim_local: Option<usize>,
249 pub rope_scale: f32,
250 pub rope_scale_local: f32,
251 pub global_attn: Option<(usize, usize)>,
254 pub inv_freq_global: Option<std::sync::Arc<Vec<f32>>>,
257 pub attn_v_norm: bool,
259 pub final_softcap: Option<f32>,
261 pub head_clusters: Option<std::sync::Arc<Vec<f32>>>,
265 pub attn_softcap: f32,
267 confidence_on: bool,
271}
272
273#[cfg(target_os = "macos")]
274impl Drop for Pipeline {
275 fn drop(&mut self) {
276 crate::gpu::kv_mirror_drop(self.graph_kv_id);
277 }
278}
279
280pub struct PipelineWeights {
285 pub embed_tokens: QTensor,
287 pub layers: Vec<LayerWeights>,
289 pub lm_head: QTensor,
291 pub final_norm: Vec<f32>,
293}
294
295pub struct LayerWeights {
297 pub input_norm: Vec<f32>,
298 pub post_norm: Vec<f32>,
301 pub attn_out_norm: Option<Vec<f32>>,
304 pub layer_scale: Option<f32>,
306 pub ffn_out_norm: Option<Vec<f32>>,
309 pub ffn: FfnKind,
310 pub attn: AttnKind,
311}
312
313#[derive(Clone, Copy, PartialEq, Debug, Default)]
316pub enum Act {
317 #[default]
318 Silu,
319 GeluTanh,
320 Situ {
323 beta: f32,
324 linear_beta: f32,
325 },
326}
327
328impl Act {
329 pub fn from_arch(name: &str) -> Self {
330 if name == "gelu_tanh" {
331 Self::GeluTanh
332 } else {
333 Self::Silu
334 }
335 }
336
337 pub fn from_arch_full(arch: &cortiq_core::ModelArch) -> Self {
339 match arch.hidden_act.as_str() {
340 "situ" => Self::Situ {
341 beta: arch.activation_situ_beta.unwrap_or(1.0) as f32,
342 linear_beta: arch.activation_situ_linear_beta.unwrap_or(0.0) as f32,
343 },
344 other => Self::from_arch(other),
345 }
346 }
347
348 #[inline]
349 pub fn apply(self, x: f32) -> f32 {
350 match self {
351 Self::Silu => inference::silu(x),
352 Self::GeluTanh => inference::gelu_tanh(x),
353 Self::Situ { beta, .. } => beta * (x / beta).tanh() * (1.0 / (1.0 + (-x).exp())),
354 }
355 }
356
357 #[inline]
360 pub fn combine(self, g: f32, u: f32) -> f32 {
361 match self {
362 Self::Situ { linear_beta, .. } if linear_beta > 0.0 => {
363 self.apply(g) * (linear_beta * (u / linear_beta).tanh())
364 }
365 _ => self.apply(g) * u,
366 }
367 }
368}
369
370pub struct DenseFfn {
372 pub gate_proj: QTensor,
373 pub up_proj: QTensor,
374 pub down_proj: QTensor,
375 pub act: Act,
377 pub down_t: Option<QTensor>,
383 pub segs: Vec<FfnSeg>,
390}
391
392pub struct FfnSeg {
397 pub gate: QTensor,
398 pub up: QTensor,
399 pub down: QTensor,
400 pub start: usize,
401 pub width: usize,
402}
403
404pub enum FfnKind {
407 Dense(DenseFfn),
408 Moe(MoeFfn),
412 DenseMoe(Box<DenseMoeFfn>),
419}
420
421pub struct DenseMoeFfn {
423 pub dense: DenseFfn,
424 pub moe: MoeFfn,
425 pub post_norm_1: Vec<f32>,
427 pub pre_norm_2: Vec<f32>,
430 pub post_norm_2: Vec<f32>,
432}
433
434pub struct MoeFfn {
435 pub router: QTensor,
437 pub experts: Vec<DenseFfn>,
438 pub top_k: usize,
439 pub norm_topk_prob: bool,
440 pub router_sigmoid: bool,
443 pub expert_bias: Option<Vec<f32>>,
447 pub routed_scaling: f32,
450 pub route_tau: Option<f32>,
456 pub shared: Option<(DenseFfn, Option<QTensor>)>,
459 pub stats: std::cell::RefCell<Vec<u64>>,
463 pub act_sq: std::cell::RefCell<Vec<f64>>,
470 pub act_rows: std::cell::RefCell<Vec<f32>>,
476 pub mask: Option<Vec<bool>>,
481 pub per_expert_scale: Option<Vec<f32>>,
484 pub router_input_norm: bool,
488 pub resonance: Option<Resonance>,
492}
493
494pub struct Resonance {
496 pub mu: Vec<f32>,
498 pub u: Vec<f32>,
500 pub k: usize,
501 pub bias: Vec<f32>,
503}
504
505impl Resonance {
506 pub fn scores(&self, x: &[f32], out: &mut [f32]) {
508 let h = x.len();
509 let ne = out.len();
510 for e in 0..ne {
511 let mu = &self.mu[e * h..(e + 1) * h];
512 let mut d2 = 0.0f32;
513 for j in 0..h {
514 let d = x[j] - mu[j];
515 d2 += d * d;
516 }
517 let mut proj = 0.0f32;
518 for i in 0..self.k {
519 let u = &self.u[(e * self.k + i) * h..(e * self.k + i + 1) * h];
520 let mut p = 0.0f32;
521 for j in 0..h {
522 p += (x[j] - mu[j]) * u[j];
523 }
524 proj += p * p;
525 }
526 out[e] = self.bias.get(e).copied().unwrap_or(0.0) - (d2 - proj);
527 }
528 }
529}
530
531pub enum AttnKind {
534 Full {
536 wq: QTensor,
537 wk: QTensor,
538 wv: QTensor,
539 wo: QTensor,
540 q_norm: Option<Vec<f32>>,
541 k_norm: Option<Vec<f32>>,
542 output_gate: bool,
543 softplus_gate: Option<(QTensor, bool)>,
547 bias: Option<(Vec<f32>, Vec<f32>, Vec<f32>)>,
549 },
550 Linear(VmfPhaseWeights),
552 LinearGdn(GdnWeights),
554 ShortConv(ShortConvWeights),
557 Mla(Box<MlaWeights>),
565 Kda(Box<crate::linear_core::KdaWeights>),
569}
570
571pub struct MlaWeights {
573 pub q_proj: QTensor,
577 pub q_a: Option<QTensor>,
580 pub q_a_norm: Option<Vec<f32>>,
581 pub kv_a: QTensor,
583 pub kv_a_norm: Vec<f32>,
585 pub kv_b: QTensor,
587 pub o_proj: QTensor,
589 pub nh: usize,
590 pub qk_rope: usize,
591 pub qk_nope: usize,
592 pub v_dim: usize,
593 pub lora: usize,
594 pub scale: f32,
596 pub nope: bool,
598}
599
600pub struct MtpModule {
605 pub enorm: Vec<f32>,
606 pub hnorm: Vec<f32>,
607 pub eh_proj: QTensor,
609 pub layer: LayerWeights,
610 pub final_norm: Vec<f32>,
611 pub kv: crate::kv_cache::LayerKvCache,
612}
613
614#[cfg(target_os = "macos")]
621enum MetalRowsItem<'a> {
622 Gdn {
623 run: Vec<crate::gpu_metal::GdnGpuLayer<'a>>,
624 first: usize,
625 },
626 Attn {
627 l: crate::gpu_metal::AttnGpuLayer<'a>,
628 li: usize,
629 q_norm: Option<&'a [f32]>,
630 k_norm: Option<&'a [f32]>,
631 output_gate: bool,
632 },
633}
634
635#[cfg(target_os = "macos")]
636struct MetalVerifyPending {
637 graph: crate::gpu_metal::VerifyGraph,
638 gdn_layers: Vec<usize>,
639 attn_layers: Vec<(usize, usize)>,
640}
641
642#[derive(Clone, Copy)]
646enum SpecTrial {
647 Spec {
648 t0: std::time::Instant,
649 gen0: usize,
650 rounds: usize,
651 },
652 Plain {
653 t0: std::time::Instant,
654 gen0: usize,
655 },
656 Decided {
657 spec: bool,
658 recheck_at: usize,
659 },
660}
661
662#[derive(Default, Clone, Copy)]
673struct SpecMon {
674 round_ms: f64,
675 tokens: f64,
676 plain_ms: f64,
677 n: u32,
678 fails: u32,
679}
680
681impl SpecMon {
682 fn round(&mut self, dt_ms: f64, produced: usize) {
683 self.n += 1;
684 if self.n == 1 {
685 return; }
687 let a = if self.n == 2 { 1.0 } else { 0.3 };
688 self.round_ms += a * (dt_ms - self.round_ms);
689 self.tokens += a * (produced as f64 - self.tokens);
690 }
691 fn pays(&self) -> bool {
692 self.plain_ms > 0.0 && self.tokens * self.plain_ms > self.round_ms * 1.03
693 }
694}
695
696pub struct GenerateResult {
698 pub text: String,
699 pub token_ids: Vec<u32>,
700 pub prompt_tokens: usize,
701 pub tokens_generated: usize,
702 pub finish_reason: String,
703 pub mtp_drafted: usize,
705 pub mtp_accepted: usize,
706 pub token_confidence: Vec<f32>,
711 pub traces: Vec<TokenTrace>,
714}
715
716#[derive(Clone, Debug)]
721pub struct TokenTrace {
722 pub t: usize,
724 pub token_id: u32,
726 pub confidence: f32,
728 pub active_skill: Option<String>,
730 pub recon: Option<f32>,
734 pub switched: bool,
737}
738
739#[cfg_attr(not(test), allow(dead_code))]
744fn top1_prob_t(logits: &[f32], id: u32, temp: f32) -> f32 {
745 let t = if temp > 1e-3 { temp } else { 1.0 };
746 let max = logits.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
747 let sum: f32 = logits.iter().map(|&v| ((v - max) / t).exp()).sum();
748 if sum > 0.0 {
749 (((logits[id as usize] - max) / t).exp()) / sum
750 } else {
751 0.0
752 }
753}
754
755fn prefill_batched() -> bool {
758 std::env::var("CMF_PREFILL")
759 .map(|v| v != "seq")
760 .unwrap_or(true)
761}
762
763#[derive(Clone, Copy)]
767enum PrefillIn<'a> {
768 Ids(&'a [u32]),
769 Hidden(&'a [f32]),
770}
771
772impl Pipeline {
779 fn can_prefill_batched(&self) -> bool {
780 prefill_batched() && !self.weights.layers.is_empty()
781 }
782}
783
784pub fn prefill_chunk() -> usize {
791 if let Some(n) = std::env::var("CMF_PREFILL_CHUNK")
792 .ok()
793 .and_then(|v| v.parse::<usize>().ok())
794 {
795 return n.max(1);
796 }
797 if cfg!(target_os = "macos") {
798 512
799 } else if cfg!(target_arch = "aarch64") {
800 256
803 } else {
804 48
805 }
806}
807
808pub type TokenCallback = Box<dyn FnMut(&str) -> bool + Send>;
810
811impl Pipeline {
812 #[inline]
816 pub fn phys_layer(&self, virtual_idx: usize) -> usize {
817 virtual_idx % self.physical_layers
818 }
819
820 #[inline]
823 pub fn is_loop_end(&self, virtual_idx: usize) -> bool {
824 self.loop_final_norm && (virtual_idx + 1) % self.physical_layers == 0
825 }
826
827 #[allow(clippy::too_many_arguments)]
829
830 #[cfg(target_os = "macos")]
849 fn graph_prefill_preferred(&self) -> bool {
850 if !crate::gpu::enabled_here()
851 || !crate::gpu::q1_force()
852 || std::env::var("CMF_GPU_BLOCK")
853 .map(|v| v == "0")
854 .unwrap_or(false)
855 || std::env::var("CMF_PREFILL_GRAPH").as_deref() == Ok("0")
858 {
859 return false;
860 }
861 self.weights
862 .layers
863 .iter()
864 .any(|lw| matches!(&lw.attn, AttnKind::LinearGdn(w) if w.in_proj_qkv.is_q1()))
865 }
866
867 #[cfg(not(target_os = "macos"))]
868 fn graph_prefill_preferred(&self) -> bool {
869 let graph_on = crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Prefill);
877 if !graph_on || !crate::gpu::enabled_here() {
878 return false;
879 }
880 if self.o1_active() {
889 return false;
890 }
891 self.weights
892 .layers
893 .iter()
894 .any(|lw| matches!(&lw.attn, AttnKind::LinearGdn(_)))
895 }
896
897 #[cfg(target_os = "macos")]
898 fn q1_graph_gpu(
899 &mut self,
900 start: usize,
901 upto: Option<usize>,
902 position: usize,
903 h: &mut [f32],
904 ) -> usize {
905 let _mt0 = std::time::Instant::now(); use crate::gpu::{AttnGpuLayer, GdnGpuCfg, GdnGpuLayer, GraphDims, MetalFfn, TokenGraph};
907 if self.attn_softcap > 0.0 || !crate::gpu::enabled_here()
909 || !crate::gpu::q1_force()
910 || std::env::var("CMF_GPU_BLOCK")
911 .map(|v| v == "0")
912 .unwrap_or(false)
913 {
914 if std::env::var("CMF_GRAPH_DBG").is_ok() {
915 eprintln!(
916 "block-graph: front gate (softcap={} enabled_here={} q1_force={})",
917 self.attn_softcap > 0.0,
918 crate::gpu::enabled_here(),
919 crate::gpu::q1_force(),
920 );
921 }
922 return start;
923 }
924 if self.swa.is_some()
929 || self.global_attn.is_some()
930 || self.attention_heads_per_layer.is_some()
931 || self.attn_v_norm
932 || (self.attn_scale - 1.0 / (self.head_dim as f32).sqrt()).abs() > 1e-9
933 || self.weights.layers.iter().any(|lw| {
934 lw.attn_out_norm.is_some()
935 || lw.ffn_out_norm.is_some()
936 || lw.layer_scale.is_some()
937 || matches!(&lw.ffn, FfnKind::Dense(d) if d.act != Act::Silu)
938 })
939 {
940 if std::env::var("CMF_GRAPH_DBG").is_ok() {
941 eprintln!(
942 "block-graph: arch ineligible (swa={} gattn={} hpl={} vnorm={} scale_delta={:.2e})",
943 self.swa.is_some(),
944 self.global_attn.is_some(),
945 self.attention_heads_per_layer.is_some(),
946 self.attn_v_norm,
947 (self.attn_scale - 1.0 / (self.head_dim as f32).sqrt()).abs(),
948 );
949 }
950 return start;
951 }
952 let limit = upto
955 .map(|u| u + 1)
956 .unwrap_or(self.num_layers)
957 .min(self.num_layers);
958
959 enum Item<'a> {
960 Gdn {
961 run: Vec<GdnGpuLayer<'a>>,
962 first: usize,
963 },
964 Attn {
965 l: AttnGpuLayer<'a>,
966 li: usize,
967 q_norm: Option<&'a [f32]>,
968 k_norm: Option<&'a [f32]>,
969 output_gate: bool,
970 bias: Option<(&'a [f32], &'a [f32], &'a [f32])>,
971 full_gpu: bool,
974 },
975 }
976
977 let attend_mode = std::env::var("CMF_GPU_ATTEND").unwrap_or_else(|_| "auto".into());
984 let attend_contract = attend_mode != "0"
985 && attend_mode != "off"
986 && self.head_dim % 4 == 0
987 && self.head_dim <= 256
988 && self.rotary_dim >= 2
989 && self.rotary_dim <= self.head_dim
990 && (self.rotary_dim / 2) % 32 == 0
991 && self.num_kv_heads > 0
992 && self.num_heads % self.num_kv_heads == 0;
993
994 let mut plan: Vec<Item> = Vec::new();
995 let mut model_ref: Option<std::sync::Arc<cortiq_core::CmfModel>> = None;
996 let block_diag = std::env::var("CMF_GRAPH_DBG").is_ok();
998 let mut scan = start;
999 while scan < limit {
1000 let lw = &self.weights.layers[self.phys_layer(scan)];
1001 let ffn = match &lw.ffn {
1002 FfnKind::Dense(d) if d.segs.is_empty() => {
1003 let (Some(g), Some(u), Some(dn)) = (
1004 d.gate_proj.q1_parts(),
1005 d.up_proj.q1_parts(),
1006 d.down_proj.q1_parts(),
1007 ) else {
1008 if block_diag {
1009 eprintln!(
1010 "block-graph: L{scan} FFN trio not graph-mappable — run ends"
1011 );
1012 }
1013 break;
1014 };
1015 MetalFfn::Dense {
1016 gate: g,
1017 up: u,
1018 down: dn,
1019 }
1020 }
1021 FfnKind::Moe(m) => {
1022 let Some(moe) = metal_moe_graph_parts(m, self.hidden_size) else {
1023 if block_diag {
1024 eprintln!(
1025 "block-graph: L{scan} MoE outside the graph contract — run ends"
1026 );
1027 }
1028 break;
1029 };
1030 if let QTensor::Mapped { model, .. } = &m.experts[0].gate_proj {
1031 model_ref.get_or_insert_with(|| model.clone());
1032 }
1033 MetalFfn::Moe(moe)
1034 }
1035 _ => {
1036 if block_diag {
1037 eprintln!("block-graph: L{scan} non-graph FFN — run ends");
1038 }
1039 break;
1040 }
1041 };
1042 match &lw.attn {
1043 AttnKind::LinearGdn(w) if self.gdn_cfg.is_some() => {
1044 let parts = (
1045 w.in_proj_qkv.q1_parts(),
1046 w.in_proj_z.q1_parts(),
1047 w.in_proj_a.f32_parts(),
1048 w.in_proj_b.f32_parts(),
1049 w.out_proj.q1_parts(),
1050 );
1051 let (Some(qkv), Some(z), Some(a), Some(b), Some(out)) = parts else {
1052 if block_diag {
1053 eprintln!(
1054 "block-graph: L{scan} GDN parts refused (qkv={} z={} a_f32={} b_f32={} out={})",
1055 w.in_proj_qkv.q1_parts().is_some(),
1056 w.in_proj_z.q1_parts().is_some(),
1057 w.in_proj_a.f32_parts().is_some(),
1058 w.in_proj_b.f32_parts().is_some(),
1059 w.out_proj.q1_parts().is_some(),
1060 );
1061 }
1062 break;
1063 };
1064 if let QTensor::Mapped { model, .. } = &w.in_proj_qkv {
1065 model_ref.get_or_insert_with(|| model.clone());
1066 }
1067 let gl = GdnGpuLayer {
1068 attn_norm: &lw.input_norm,
1069 post_norm: &lw.post_norm,
1070 qkv,
1071 z,
1072 a,
1073 b,
1074 out,
1075 ffn,
1076 conv1d: &w.conv1d,
1077 a_log: &w.a_log,
1078 dt_bias: &w.dt_bias,
1079 gnorm: &w.norm,
1080 };
1081 match plan.last_mut() {
1082 Some(Item::Gdn { run, .. }) => run.push(gl),
1083 _ => plan.push(Item::Gdn {
1084 run: vec![gl],
1085 first: scan,
1086 }),
1087 }
1088 }
1089 AttnKind::Full {
1090 wq,
1091 wk,
1092 wv,
1093 wo,
1094 q_norm,
1095 k_norm,
1096 output_gate,
1097 softplus_gate: None,
1098 bias,
1099 } if !self.kv_cache.layers[scan].o1_sealed()
1100 || std::env::var("CMF_O1_METAL").as_deref() == Ok("1") =>
1105 {
1106 let parts = (wq.q1_parts(), wk.q1_parts(), wv.q1_parts(), wo.q1_parts());
1107 let (Some(pq), Some(pk), Some(pv), Some(po)) = parts else {
1108 break;
1109 };
1110 if let QTensor::Mapped { model, .. } = wq {
1111 model_ref.get_or_insert_with(|| model.clone());
1112 }
1113 let cache = &self.kv_cache.layers[scan];
1114 let o1_metal = cache.o1.is_some()
1118 && std::env::var("CMF_O1_METAL").as_deref() == Ok("1")
1119 && cache.o1_views().is_some();
1120 let full_gpu = attend_contract
1121 && cache.mode == crate::kv_cache::KvMode::F32
1122 && (cache.o1.is_none() || o1_metal)
1123 && bias.is_none()
1124 && pq.1 == self.num_heads * self.head_dim * (1 + *output_gate as usize)
1125 && pk.1 == self.num_kv_heads * self.head_dim
1126 && pv.1 == self.num_kv_heads * self.head_dim
1127 && po.2 == self.num_heads * self.head_dim;
1128 plan.push(Item::Attn {
1129 l: AttnGpuLayer {
1130 attn_norm: &lw.input_norm,
1131 post_norm: &lw.post_norm,
1132 wq: pq,
1133 wk: pk,
1134 wv: pv,
1135 wo: po,
1136 ffn,
1137 },
1138 li: scan,
1139 q_norm: q_norm.as_deref(),
1140 k_norm: k_norm.as_deref(),
1141 output_gate: *output_gate,
1142 bias: bias
1143 .as_ref()
1144 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
1145 full_gpu,
1146 });
1147 }
1148 _ => break,
1149 }
1150 scan += 1;
1151 }
1152 let Some(model) = model_ref else {
1153 if std::env::var("CMF_GRAPH_DBG").is_ok() {
1154 eprintln!("q1-graph: no model ref (start {start}, scanned to {scan})");
1155 }
1156 return start;
1157 };
1158 if plan.is_empty() {
1159 if std::env::var("CMF_GRAPH_DBG").is_ok() {
1160 eprintln!("q1-graph: empty plan at layer {start}");
1161 }
1162 return start;
1163 }
1164 let has_moe = plan.iter().any(|it| match it {
1165 Item::Gdn { run, .. } => run.iter().any(|l| matches!(l.ffn, MetalFfn::Moe(_))),
1166 Item::Attn { l, .. } => matches!(l.ffn, MetalFfn::Moe(_)),
1167 });
1168 let has_gdn = plan.iter().any(|it| matches!(it, Item::Gdn { .. }));
1169 let dev_attend = attend_contract
1170 && (self.head_dim <= 128
1171 || has_moe
1172 || (self.head_dim <= 256 && has_gdn)
1178 || attend_mode == "force"
1179 || attend_mode == "256");
1180 if !dev_attend {
1181 for it in &mut plan {
1182 if let Item::Attn { li, full_gpu, .. } = it {
1183 let keep_o1 = self.kv_cache.layers[*li].o1.is_some()
1186 && std::env::var("CMF_O1_METAL").as_deref() == Ok("1");
1187 if !keep_o1 {
1188 *full_gpu = false;
1189 }
1190 }
1191 }
1192 }
1193 if std::env::var("CMF_GRAPH_DBG").is_ok() {
1194 use std::sync::atomic::{AtomicBool, Ordering};
1195 static SAID: AtomicBool = AtomicBool::new(false);
1196 if !SAID.swap(true, Ordering::Relaxed) {
1197 let fg = plan
1198 .iter()
1199 .filter(|it| matches!(it, Item::Attn { full_gpu: true, .. }))
1200 .count();
1201 let att = plan
1202 .iter()
1203 .filter(|it| matches!(it, Item::Attn { .. }))
1204 .count();
1205 eprintln!(
1206 "q1-graph: plan of {} items from layer {start} to {scan} | dev_attend={dev_attend} full_gpu {fg}/{att} | hd={} rd={} nkv={} nh={}",
1207 plan.len(),
1208 self.head_dim,
1209 self.rotary_dim,
1210 self.num_kv_heads,
1211 self.num_heads,
1212 );
1213 }
1214 }
1215 let dims = GraphDims {
1216 hidden: self.hidden_size,
1217 eps: self.rms_eps as f32,
1218 gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
1219 };
1220 let Some(mut graph) = TokenGraph::new(&model, dims, h) else {
1221 return start;
1222 };
1223 let gcfg = self.gdn_cfg.map(|cfg| GdnGpuCfg {
1224 nv: cfg.num_v_heads,
1225 nk: cfg.num_k_heads,
1226 dk: cfg.key_head_dim,
1227 dv: cfg.value_head_dim,
1228 kk: cfg.conv_kernel,
1229 hidden: self.hidden_size,
1230 inter: self.intermediate_size,
1231 c_dim: cfg.conv_dim(),
1232 eps: cfg.rms_eps as f32,
1233 gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
1234 });
1235 let mut valid = 0usize;
1239 let mut end = start;
1240 crate::gpu::stageprof(1, _mt0.elapsed()); if std::env::var("CMF_PLAN_DUMP").is_ok() {
1242 static ONCE: std::sync::Once = std::sync::Once::new();
1243 ONCE.call_once(|| {
1244 for it in &plan {
1245 match it {
1246 Item::Gdn { first, run } => {
1247 eprintln!("plan: Gdn first={first} len={}", run.len())
1248 }
1249 Item::Attn { li, full_gpu, .. } => {
1250 eprintln!("plan: Attn li={li} full_gpu={full_gpu}")
1251 }
1252 }
1253 }
1254 });
1255 }
1256 for item in &plan {
1257 let ok = match item {
1258 Item::Gdn { run, .. } => gcfg
1259 .as_ref()
1260 .map(|gc| run.iter().all(|l| graph.gdn_ok(l, gc)))
1261 .unwrap_or(false),
1262 Item::Attn { l, .. } => graph.attn_ok(l),
1263 };
1264 if !ok {
1265 if block_diag {
1266 eprintln!(
1267 "block-graph: plan item {} ({}) failed graph preflight",
1268 valid,
1269 match item {
1270 Item::Gdn { run, first } => format!("GDN run L{first}+{}", run.len()),
1271 Item::Attn { li, .. } => format!("Attn L{li}"),
1272 }
1273 );
1274 }
1275 break;
1276 }
1277 valid += 1;
1278 end += match item {
1279 Item::Gdn { run, .. } => run.len(),
1280 Item::Attn { .. } => 1,
1281 };
1282 }
1283 plan.truncate(valid);
1284 if plan.is_empty() {
1285 return start;
1286 }
1287
1288 let inv_freq = self.inv_freq.clone();
1289 let pool = self.pool.clone();
1290 let (nh, nkv, hd, hs, rd, eps) = (
1291 self.num_heads,
1292 self.num_kv_heads,
1293 self.head_dim,
1294 self.hidden_size,
1295 self.rotary_dim,
1296 self.rms_eps,
1297 );
1298 let norm_style = self.norm_style;
1299 let gemma = norm_style == cortiq_core::NormStyle::Gemma;
1300 let want = self.gdn_cfg.map(|c| c.state_len()).unwrap_or(0);
1301 let kv_id = self.graph_kv_id;
1302 let mut pending: Vec<(usize, usize)> = Vec::new();
1305 let mut dev_attn: Vec<usize> = Vec::new();
1308 for item in &plan {
1309 let _xt0 = std::time::Instant::now();
1310 let _xkind: u32 = match item {
1311 Item::Gdn { .. } => 2,
1312 Item::Attn { .. } => 3,
1313 };
1314 if self.loop_final_norm {
1316 let item_start = match item {
1317 Item::Gdn { first, .. } => *first,
1318 Item::Attn { li, .. } => *li,
1319 };
1320 if item_start > start && self.is_loop_end(item_start - 1) {
1321 graph.encode_loop_norm(&self.weights.final_norm);
1322 }
1323 }
1324 match item {
1325 Item::Gdn { run, first } => {
1326 for l in &mut self.kv_cache.layers[*first..*first + run.len()] {
1327 if l.linear_state.len() != want {
1328 l.linear_state = vec![0f32; want];
1329 }
1330 }
1331 let ro: Vec<&[f32]> = self.kv_cache.layers[*first..*first + run.len()]
1332 .iter()
1333 .map(|l| l.linear_state.as_slice())
1334 .collect();
1335 let _ig = std::time::Instant::now();
1336 if !graph.encode_gdn_run(run, &ro, gcfg.as_ref().unwrap()) {
1337 tracing::error!("q1 graph: GDN run refused after validation");
1339 return start;
1340 }
1341 graph.commit_kind = 2;
1344 graph.commit();
1345 crate::gpu::stageprof(0, _ig.elapsed());
1346 pending.push((*first, run.len()));
1347 }
1348 Item::Attn {
1349 l,
1350 li,
1351 q_norm,
1352 k_norm,
1353 output_gate,
1354 bias,
1355 full_gpu,
1356 } => {
1357 let _ia = std::time::Instant::now();
1358 if *full_gpu {
1360 let cache = &self.kv_cache.layers[*li];
1361 let o1p = if cache.o1.is_some() {
1362 match cache.o1_views() {
1363 Some(views) => Some(crate::gpu::O1AttnParams {
1364 views,
1365 epoch: self.o1_epoch,
1366 }),
1367 None => None,
1369 }
1370 } else {
1371 None
1372 };
1373 let o1_layer = cache.o1.is_some();
1374 if o1_layer && o1p.is_none() {
1375 }
1377 let cpu_k: Vec<&[f32]> = (0..nkv).map(|g| cache.head_keys(g)).collect();
1378 let cpu_v: Vec<&[f32]> = (0..nkv).map(|g| cache.head_values(g)).collect();
1379 let cpu_stored = if o1_layer { 0 } else { cpu_k[0].len() / hd };
1380 let p = crate::gpu::AttnDeviceParams {
1381 kv_id,
1382 layer: *li,
1383 nh,
1384 nkv,
1385 hd,
1386 rd,
1387 position,
1388 eps: eps as f32,
1389 gemma,
1390 output_gate: *output_gate,
1391 q_norm: *q_norm,
1392 k_norm: *k_norm,
1393 inv_freq: &inv_freq,
1394 cpu_k,
1395 cpu_v,
1396 cpu_stored,
1397 o1: o1p,
1398 };
1399 let o1_bad = o1_layer && p.o1.is_none();
1400 if !o1_bad && graph.attn_device_ok(l, &p) && graph.encode_attn_device(l, &p)
1401 {
1402 if p.o1.is_none() {
1404 dev_attn.push(*li);
1405 }
1406 graph.commit_kind = 3;
1407 graph.commit();
1408 crate::gpu::stageprof(_xkind, _xt0.elapsed());
1412 continue;
1413 }
1414 }
1416 graph.encode_attn_prefix(l);
1417 graph.sync();
1418 if !pending.is_empty() {
1419 let idxs: Vec<usize> =
1420 pending.drain(..).flat_map(|(f, n)| f..f + n).collect();
1421 let mut outs: Vec<&mut [f32]> = self
1422 .kv_cache
1423 .layers
1424 .iter_mut()
1425 .enumerate()
1426 .filter(|(i, _)| idxs.binary_search(i).is_ok())
1427 .map(|(_, s)| s.linear_state.as_mut_slice())
1428 .collect();
1429 graph.read_states(&mut outs);
1430 }
1431 let mut q_raw = attention::take_buf(l.wq.1);
1432 let mut k = attention::take_buf(l.wk.1);
1433 let mut v = attention::take_buf(l.wv.1);
1434 graph.read_qkv(&mut q_raw, &mut k, &mut v);
1435 let cfg = QwenAttnCfg {
1436 num_heads: nh,
1437 num_kv_heads: nkv,
1438 head_dim: hd,
1439 hidden_size: hs,
1440 position,
1441 inv_freq: &inv_freq,
1442 rotary_dim: rd,
1443 scale: self.attn_scale,
1444 softcap: self.attn_softcap,
1445 window: None,
1446 v_norm: false,
1447 q_norm: *q_norm,
1448 k_norm: *k_norm,
1449 output_gate: *output_gate,
1450 softplus_gate: None,
1451 rope_scale: 1.0,
1452 bias: *bias,
1453 rms_eps: eps,
1454 norm_style,
1455 pool: pool.as_deref(),
1456 };
1457 let oracle = std::env::var("CMF_ATTN_ORACLE").as_deref() == Ok("1")
1460 || std::env::var("CMF_ATTN_DUMP").is_ok();
1461 let _ = full_gpu;
1462 let oracle_in = oracle.then(|| (q_raw.clone(), k.clone(), v.clone()));
1463 let mut ao = attention::qwen_attention_core(
1464 q_raw,
1465 k,
1466 v,
1467 &mut self.kv_cache.layers[*li],
1468 &cfg,
1469 );
1470 if let Ok(dir) = std::env::var("CMF_ATTN_DUMP") {
1474 if let Some((qr0, k0, v0)) = oracle_in.clone() {
1475 let (cq, _cg, _ck, _cv) = attention::finish_projection_debug(qr0, k0, v0, &cfg, position);
1476 let cache = &self.kv_cache.layers[*li];
1477 let n = cache.head_keys(0).len() / hd;
1478 let mut bytes: Vec<u8> = Vec::new();
1479 for v in [nh as u32, nkv as u32, hd as u32, n as u32, position as u32] {
1480 bytes.extend_from_slice(&v.to_le_bytes());
1481 }
1482 for v in &cq {
1483 bytes.extend_from_slice(&v.to_le_bytes());
1484 }
1485 for g in 0..nkv {
1486 for v in cache.head_keys(g) {
1487 bytes.extend_from_slice(&v.to_le_bytes());
1488 }
1489 }
1490 for g in 0..nkv {
1491 for v in cache.head_values(g) {
1492 bytes.extend_from_slice(&v.to_le_bytes());
1493 }
1494 }
1495 let _ = std::fs::write(format!("{dir}/L{li}_pos{position}.bin"), &bytes);
1496 }
1497 }
1498 if let Some((qr0, k0, v0)) = oracle_in.filter(|_| std::env::var("CMF_ATTN_ORACLE").as_deref() == Ok("1")) {
1499 let (cq, _cg, ck, cv) = attention::finish_projection_debug(qr0, k0, v0, &cfg, position);
1500 let mut h_now = vec![0f32; hs];
1501 graph.read_h(&mut h_now);
1502 let cache = &self.kv_cache.layers[*li];
1503 let n_after = cache.head_keys(0).len() / hd;
1504 let cpu_k: Vec<&[f32]> = (0..nkv).map(|g| &cache.head_keys(g)[..(n_after - 1) * hd]).collect();
1505 let cpu_v: Vec<&[f32]> = (0..nkv).map(|g| &cache.head_values(g)[..(n_after - 1) * hd]).collect();
1506 let p = crate::gpu::AttnDeviceParams {
1507 kv_id,
1508 layer: *li,
1509 nh,
1510 nkv,
1511 hd,
1512 rd,
1513 position,
1514 eps: eps as f32,
1515 gemma,
1516 output_gate: *output_gate,
1517 q_norm: *q_norm,
1518 k_norm: *k_norm,
1519 inv_freq: &inv_freq,
1520 cpu_k,
1521 cpu_v,
1522 cpu_stored: n_after - 1,
1523 o1: None,
1524 };
1525 if let Some((dq, dk, dv, dao)) = graph.debug_attn_device(l, &p, &h_now) {
1526 let md = |a: &[f32], b: &[f32]| a.iter().zip(b).fold(0f32, |m, (x, y)| m.max((x - y).abs()));
1527 let nn = |a: &[f32]| a.iter().map(|x| x * x).sum::<f32>().sqrt();
1528 eprintln!(
1529 "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}",
1530 nn(&cq), md(&cq, &dq), nn(&ck), md(&ck, &dk), nn(&cv), md(&cv, &dv), nn(&ao), md(&ao, &dao)
1531 );
1532 } else {
1533 eprintln!("attn-oracle L{li}: device probe declined");
1534 }
1535 }
1536 graph.encode_attn_suffix(l, &ao);
1537 graph.commit();
1540 attention::recycle_buf(&mut ao);
1541 }
1542 }
1543
1544 crate::gpu::stageprof(_xkind, _xt0.elapsed());
1545 }
1546 let mut lm_rows = None;
1551 if self.graph_want_logits
1552 && upto.is_none()
1553 && end == self.num_layers
1554 && std::env::var("CMF_GPU_LMHEAD")
1555 .map(|v| v != "0")
1556 .unwrap_or(true)
1557 {
1558 if let Some(lm) = self.weights.lm_head.q1_parts() {
1559 if graph.lm_head_ok(lm) {
1560 graph.encode_lm_head(&self.weights.final_norm, lm);
1561 lm_rows = Some(lm.1);
1562 }
1563 }
1564 }
1565 let _sy0 = std::time::Instant::now();
1566 graph.sync();
1567 let _rs0 = std::time::Instant::now();
1568 if !pending.is_empty() {
1569 let idxs: Vec<usize> = pending.drain(..).flat_map(|(f, n)| f..f + n).collect();
1570 let mut outs: Vec<&mut [f32]> = self
1571 .kv_cache
1572 .layers
1573 .iter_mut()
1574 .enumerate()
1575 .filter(|(i, _)| idxs.binary_search(i).is_ok())
1576 .map(|(_, s)| s.linear_state.as_mut_slice())
1577 .collect();
1578 graph.read_states(&mut outs);
1579 }
1580 if std::env::var("CMF_GRAPH_HOSTPROF").as_deref() == Ok("1") {
1581 use std::sync::atomic::{AtomicU64, Ordering};
1582 static SY: AtomicU64 = AtomicU64::new(0);
1583 static RS: AtomicU64 = AtomicU64::new(0);
1584 static N: AtomicU64 = AtomicU64::new(0);
1585 SY.fetch_add((_rs0 - _sy0).as_nanos() as u64, Ordering::Relaxed);
1586 RS.fetch_add(_rs0.elapsed().as_nanos() as u64, Ordering::Relaxed);
1587 let n = N.fetch_add(1, Ordering::Relaxed) + 1;
1588 if n % 100 == 0 {
1589 eprintln!(
1590 "postprof: sync-wait {:.1} ms/ток | read_states {:.1} ms/ток ({n})",
1591 SY.load(Ordering::Relaxed) as f64 / n as f64 / 1e6,
1592 RS.load(Ordering::Relaxed) as f64 / n as f64 / 1e6
1593 );
1594 }
1595 }
1596 if let Some(rows) = lm_rows {
1597 crate::gpu::hostprof_encode_done(_mt0);
1598 let mut lg = attention::take_buf(rows.min(self.vocab_size));
1599 graph.read_logits(&mut lg);
1600 crate::gpu::hostprof_total(_mt0);
1601 lg.resize(self.vocab_size, 0.0);
1602 if let Some(c) = self.final_softcap {
1603 for l in lg.iter_mut() {
1604 *l = c * (*l / c).tanh();
1605 }
1606 }
1607 self.graph_logits = Some(lg);
1608 }
1609 graph.finish(h);
1610 for li in dev_attn {
1614 let mut krow = attention::take_buf(nkv * hd);
1615 let mut vrow = attention::take_buf(nkv * hd);
1616 if crate::gpu::kv_mirror_read_last(kv_id, li, nkv, hd, &mut krow, &mut vrow) {
1617 let cache = &mut self.kv_cache.layers[li];
1618 cache.append(&krow, &vrow, &[]);
1619 let n = cache.seq_len;
1620 let mut imp = attention::take_buf(n);
1621 crate::gpu::kv_mirror_take_imp(kv_id, li, &mut imp);
1622 cache.accumulate_imp(&imp);
1623 attention::recycle_buf(&mut imp);
1624 }
1625 attention::recycle_buf(&mut krow);
1626 attention::recycle_buf(&mut vrow);
1627 }
1628 end
1629 }
1630
1631 pub fn new(
1632 tokenizer: Tokenizer,
1633 weights: PipelineWeights,
1634 hidden_size: usize,
1635 intermediate_size: usize,
1636 num_heads: usize,
1637 num_kv_heads: usize,
1638 head_dim: usize,
1639 num_layers: usize,
1640 physical_layers: usize,
1641 loop_final_norm: bool,
1642 vocab_size: usize,
1643 rms_eps: f64,
1644 rope_base: f32,
1645 norm_style: NormStyle,
1646 max_seq_len: usize,
1647 sampler_config: SamplerConfig,
1648 ) -> Self {
1649 let rng = match sampler_config.seed {
1650 Some(s) => SplitMix64::new(s),
1651 None => SplitMix64::from_entropy(),
1652 };
1653 let inv_freq = std::sync::Arc::new(attention::rope_inv_freq(head_dim, rope_base));
1654 let pool = Pool::from_env();
1655 if let Some(p) = &pool {
1656 tracing::info!("worker pool: {} threads", p.n_workers());
1657 }
1658 Self {
1659 gpu_plan: None,
1660 tokenizer: std::sync::Arc::new(tokenizer),
1661 kv_cache: KvCache::new(num_layers, num_kv_heads, head_dim, max_seq_len),
1662 sampler_config,
1663 weights,
1664 hidden_size,
1665 intermediate_size,
1666 num_heads,
1667 num_kv_heads,
1668 head_dim,
1669 num_layers,
1670 physical_layers,
1671 loop_final_norm,
1672 vocab_size,
1673 rms_eps,
1674 rope_base,
1675 norm_style,
1676 rotary_dim: head_dim,
1677 attention_heads_per_layer: None,
1678 vmf_cfg: None,
1679 gdn_cfg: None,
1680 kda_cfg: None,
1681 g3n: None,
1682 dsv4: None,
1683 dsv4_mtp: Vec::new(),
1684 dspark: None,
1685 dspark_pending: Vec::new(),
1686 dspark_hist: Vec::new(),
1687 dspark_real: Vec::new(),
1688 dspark_trunk_picks: Vec::new(),
1689 dspark_exp: Vec::new(),
1690 dspark_draft_ns: 0,
1691 logit_multiplier: None,
1692 cancel: std::sync::Arc::new(std::sync::atomic::AtomicBool::new(false)),
1693 kv_history: Vec::new(),
1694 short_conv_cfg: None,
1695 mtp: None,
1696 speculative: std::env::var("CMF_MTP").map(|v| v != "0").unwrap_or(true),
1697 rng,
1698 sampler_scratch: SamplerScratch::default(),
1699 spec_forced: None,
1700 spec_q: Vec::new(),
1701 spec_p: Vec::new(),
1702 spec_res: Vec::new(),
1703 spec_qs: Vec::new(),
1704 spec_ps: Vec::new(),
1705 spec_ress: Vec::new(),
1706 mtp_graph_mode: None,
1707 #[cfg(target_os = "macos")]
1708 metal_verify: None,
1709 inv_freq,
1710 ws: ForwardScratch::new(hidden_size),
1711 pool,
1712 model: None,
1713 dyn_force_f32: false,
1714 dyn_skill_layers: Vec::new(),
1715 dyn_active: None,
1716 dyn_blend_loaded: false,
1717 dyn_phi_layer: None,
1718 dyn_phi_ema: Vec::new(),
1719 dyn_phi_seen: 0,
1720 dyn_router: None,
1721 o1_cfg: None,
1722 o1_epoch: 0,
1723 o1_flags: Vec::new(),
1724 trace: false,
1725 calib_temp: 1.0,
1726 confidence_on: true,
1727 embed_multiplier: 1.0,
1728 attn_scale: 1.0 / (head_dim as f32).sqrt(),
1729 swa: None,
1730 sliding_layers: None,
1731 inv_freq_local: None,
1732 rotary_dim_local: None,
1733 rope_scale: 1.0,
1734 rope_scale_local: 1.0,
1735 global_attn: None,
1736 inv_freq_global: None,
1737 attn_v_norm: false,
1738 final_softcap: None,
1739 head_clusters: None,
1740 attn_softcap: 0.0,
1741 graph_want_logits: false,
1742 graph_logits: None,
1743 graph_kv_id: {
1744 static NEXT: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(1);
1745 NEXT.fetch_add(1, std::sync::atomic::Ordering::Relaxed)
1746 },
1747 }
1748 }
1749
1750 pub fn set_o1(&mut self, cfg: Option<crate::nystrom::O1Cfg>) {
1757 self.o1_flags = match &cfg {
1758 Some(c) => {
1759 let mut flags = c.layer_flags(self.num_layers);
1760 for (li, f) in flags.iter_mut().enumerate() {
1761 if *f
1762 && !matches!(
1763 self.weights.layers[self.phys_layer(li)].attn,
1764 AttnKind::Full { .. }
1765 )
1766 {
1767 *f = false;
1768 }
1769 }
1770 flags
1771 }
1772 None => Vec::new(),
1773 };
1774 if let Some(c) = &cfg {
1775 let n = self.o1_flags.iter().filter(|&&f| f).count();
1776 tracing::info!(
1777 "o1 nystrom attention: {n}/{} layer(s), m={} w={} sink={} rect={:?}",
1778 self.num_layers,
1779 c.m,
1780 c.w,
1781 c.sink,
1782 c.rect
1783 );
1784 }
1785 self.o1_cfg = cfg;
1786 }
1787
1788 pub fn o1_active(&self) -> bool {
1790 self.o1_cfg.is_some() && self.o1_flags.iter().any(|&f| f)
1791 }
1792
1793 pub fn o1_begin(&mut self) {
1798 if let Some(c) = &self.o1_cfg {
1799 let (m, w, sink, rect) = (c.m, c.w, c.sink, c.rect);
1800 for (li, &f) in self.o1_flags.iter().enumerate() {
1801 if f {
1802 self.kv_cache.layers[li].o1_begin(m, w, sink, rect);
1803 }
1804 }
1805 }
1806 }
1807
1808 pub fn o1_seal(&mut self) {
1812 self.o1_epoch = self.o1_epoch.wrapping_add(1);
1813 if self.o1_cfg.is_none() {
1814 return;
1815 }
1816 for li in 0..self.num_layers {
1817 if self.o1_flags.get(li).copied().unwrap_or(false) {
1818 self.kv_cache.layers[li].o1_seal(self.num_heads);
1819 }
1820 }
1821 }
1822
1823 pub fn set_trace(&mut self, on: bool) {
1825 self.trace = on;
1826 }
1827
1828 pub fn set_sampler_config(&mut self, config: SamplerConfig) {
1831 self.rng = match config.seed {
1832 Some(seed) => SplitMix64::new(seed),
1833 None => SplitMix64::from_entropy(),
1834 };
1835 self.sampler_config = config;
1836 }
1837
1838 pub fn set_confidence(&mut self, on: bool) {
1843 self.confidence_on = on;
1844 }
1845
1846 pub fn set_calib_temp(&mut self, t: f32) {
1849 self.calib_temp = if t > 1e-3 { t } else { 1.0 };
1850 }
1851
1852 pub fn calib_temp(&self) -> f32 {
1854 self.calib_temp
1855 }
1856
1857 pub fn set_rotary(&mut self, rotary_dim: usize, base: f32) {
1860 self.rotary_dim = rotary_dim.min(self.head_dim);
1861 self.inv_freq = std::sync::Arc::new(attention::rope_inv_freq(self.rotary_dim, base));
1862 }
1863
1864 fn attn_cfg(&self, position: usize) -> QwenAttnCfg<'_> {
1865 QwenAttnCfg {
1866 num_heads: self.num_heads,
1867 num_kv_heads: self.num_kv_heads,
1868 head_dim: self.head_dim,
1869 hidden_size: self.hidden_size,
1870 position,
1871 inv_freq: &self.inv_freq,
1872 rotary_dim: self.rotary_dim,
1873 scale: self.attn_scale,
1874 softcap: self.attn_softcap,
1875 window: None,
1876 v_norm: false,
1877 q_norm: None,
1878 k_norm: None,
1879 output_gate: false,
1880 softplus_gate: None,
1881 rope_scale: self.rope_scale,
1882 bias: None,
1883 rms_eps: self.rms_eps,
1884 norm_style: self.norm_style,
1885 pool: self.pool.as_deref(),
1886 }
1887 }
1888
1889 pub fn generate(
1891 &mut self,
1892 prompt: &str,
1893 max_tokens: usize,
1894 task_mask: Option<&TaskMask>,
1895 on_token: Option<TokenCallback>,
1896 ) -> Result<GenerateResult, String> {
1897 let input_ids = self.tokenizer.with_bos(self.tokenizer.encode(prompt));
1898 self.generate_from_ids(&input_ids, max_tokens, task_mask, on_token)
1899 }
1900
1901 fn drop_open_mask<'m>(&self, m: Option<&'m TaskMask>) -> Option<&'m TaskMask> {
1903 m.filter(|m| !m.fully_open(self.intermediate_size, self.num_heads))
1904 }
1905
1906 pub fn generate_from_ids(
1914 &mut self,
1915 input_ids: &[u32],
1916 max_tokens: usize,
1917 task_mask: Option<&TaskMask>,
1918 mut on_token: Option<TokenCallback>,
1919 ) -> Result<GenerateResult, String> {
1920 if std::env::var("CMF_TRACE_H").is_ok() {
1921 eprintln!("input_ids: {input_ids:?}");
1922 }
1923 if input_ids.is_empty() {
1924 return Err("empty prompt: nothing to generate from".to_string());
1925 }
1926 let task_mask = self.drop_open_mask(task_mask);
1931
1932 let reuse_from = {
1940 let on = !std::env::var("CMF_KV_REUSE").is_ok_and(|v| v == "0");
1941 let h = &self.kv_history;
1942 if on
1943 && task_mask.is_none()
1944 && self.mtp.is_none()
1945 && self.o1_cfg.is_none()
1946 && !h.is_empty()
1947 && h.len() < input_ids.len()
1948 && input_ids[..h.len()] == h[..]
1949 {
1950 h.len()
1951 } else {
1952 0
1953 }
1954 };
1955 if reuse_from == 0 {
1956 self.kv_cache.clear();
1958 self.kv_history.clear();
1959 crate::gpu::graph_kv_reset(self.graph_kv_id);
1960 } else if std::env::var("CMF_PREFILL_PROF").is_ok() {
1961 eprintln!(
1962 "kv-reuse: {} of {} prompt positions already cached",
1963 reuse_from,
1964 input_ids.len()
1965 );
1966 }
1967 crate::gpu::graph_race_begin_generation();
1968 self.o1_begin();
1969
1970 let graph_on = crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Decode);
1976 let spec_sampling_ok = self.sampler_config.temperature < 1e-6
2003 || std::env::var("CMF_GRAPH_SPEC_SAMPLE").as_deref() == Ok("1");
2004 let (mut dense_n, mut dense_q4tp) = (0usize, 0usize);
2020 for lw in &self.weights.layers {
2021 if let FfnKind::Dense(d) = &lw.ffn {
2022 dense_n += 1;
2023 if matches!(d.gate_proj.graph_weight(), Some((_, _, 6, _)))
2024 && matches!(d.up_proj.graph_weight(), Some((_, _, 6, _)))
2025 && matches!(d.down_proj.graph_weight(), Some((_, _, 6, _)))
2026 {
2027 dense_q4tp += 1;
2028 }
2029 }
2030 }
2031 let spec_default_ok = dense_n == 0 || dense_q4tp * 10 >= dense_n * 9;
2032 let penalized = self.sampler_config.repetition_penalty != 1.0
2036 || self.sampler_config.presence_penalty != 0.0
2037 || !self.sampler_config.suppress_tokens.is_empty();
2038 #[cfg(feature = "gpu")]
2043 let metal_wgpu = graph_on && crate::gpu_wgpu::wgpu_backend_is_metal();
2044 #[cfg(not(feature = "gpu"))]
2045 let metal_wgpu = false;
2046 let spec_env = std::env::var("CMF_GRAPH_SPEC").ok();
2047 let spec_wanted = match spec_env.as_deref() {
2048 Some("0") => false,
2049 Some(_) => {
2050 if metal_wgpu {
2051 tracing::warn!(
2052 "CMF_GRAPH_SPEC forced on wgpu/Metal: the batched verify graph is not \
2053 verified on this backend (garbage measured on Qwen3.5-0.8B)"
2054 );
2055 }
2056 true
2057 }
2058 None => spec_default_ok && !penalized && !metal_wgpu,
2059 };
2060 #[cfg(target_os = "macos")]
2063 let metal_graph = crate::gpu::q1_force()
2064 && crate::gpu::enabled_here()
2065 && std::env::var("CMF_GPU_BLOCK").map(|v| v != "0").unwrap_or(true);
2066 #[cfg(not(target_os = "macos"))]
2067 let metal_graph = false;
2068 let graph_spec = self.speculative
2069 && (graph_on || metal_graph)
2070 && self.mtp.is_some()
2071 && task_mask.is_none()
2072 && !self.o1_active()
2073 && spec_sampling_ok
2074 && spec_wanted;
2075 let pair_pays = self.gdn_cfg.is_none() || std::env::var("CMF_MTP").as_deref() == Ok("1");
2082 let spec_active = self.speculative
2083 && self.mtp.is_some()
2084 && task_mask.is_none()
2085 && !self.o1_active()
2086 && ((!graph_on && pair_pays && self.sampler_config.temperature < 1e-6) || graph_spec);
2087 let mut mtp = if spec_active { self.mtp.take() } else { None };
2090 if std::env::var("CMF_MTP_CHAIN_PROBE").is_ok() {
2091 eprintln!(
2092 "mtp-probe gate: spec_active={spec_active} mtp={} speculative={} graph_on={graph_on} temp_ok={}",
2093 mtp.is_some(),
2094 self.speculative,
2095 self.sampler_config.temperature < 1e-6,
2096 );
2097 }
2098 if let Some(m) = &mut mtp {
2099 m.kv.clear();
2100 crate::gpu::graph_kv_reset(self.mtp_kv_id());
2102 self.mtp_graph_mode = None;
2103 }
2104 let mut router = if mtp.is_none() {
2108 self.dyn_router.take()
2109 } else {
2110 None
2111 };
2112 if let Some(r) = &mut router {
2113 r.reset(); self.dyn_phi_seen = 0; let _ = self.set_active_skill(None);
2116 }
2117
2118 let mut all_ids = input_ids.to_vec();
2119 let mut generated = 0usize;
2120 let mut finish_reason = "max_tokens".to_string();
2121 let mut drafted = 0usize;
2122 let mut accepted = 0usize;
2123 let mut confidence: Vec<f32> = Vec::new();
2124 let trace_on = self.trace;
2125 let calib_temp = self.calib_temp;
2126 let mut traces: Vec<TokenTrace> = Vec::new();
2127
2128 let mut hidden = vec![0.0f32; self.hidden_size];
2134 let mut pos = reuse_from;
2135 let fuse_lm = mtp.is_none()
2144 && router.is_none()
2145 && std::env::var("CMF_GPU_LMHEAD").as_deref() != Ok("0");
2146 self.graph_logits = None;
2147 self.graph_want_logits = false;
2148 let _tpf = std::time::Instant::now();
2149 let batch_k = std::env::var("CMF_BATCH_K")
2150 .ok()
2151 .and_then(|v| v.parse::<usize>().ok())
2152 .unwrap_or(0);
2153 while self.dsv4.is_some()
2164 && mtp.is_none()
2165 && pos < input_ids.len()
2166 && !self.cancel.load(std::sync::atomic::Ordering::Relaxed)
2167 {
2168 let end = (pos + prefill_chunk()).min(input_ids.len());
2169 let ids: Vec<u32> = input_ids[pos..end].to_vec();
2170 let mut lg = Vec::new();
2171 if let Some(b) = &mut self.dsv4 {
2172 let (g, layers, cfg, st) = (&b.0, &b.1, b.2, &mut b.3);
2173 crate::dsv4::forward_chunk(
2174 g,
2175 layers,
2176 &cfg,
2177 st,
2178 &ids,
2179 pos,
2180 &self.inv_freq,
2181 self.pool.as_deref(),
2182 &mut lg,
2183 end == input_ids.len(),
2184 );
2185 }
2186 if end == input_ids.len() {
2187 self.graph_logits = Some(lg);
2188 }
2189 pos = end;
2190 hidden = vec![0.0; self.hidden_size];
2191 }
2192 let dyn_prefill = router.is_some();
2197 let graph_prefill = self.graph_prefill_preferred();
2203 #[cfg(target_os = "macos")]
2211 if task_mask.is_none()
2212 && !dyn_prefill
2213 && crate::gpu::q1_force()
2214 && crate::gpu::enabled_here()
2215 && self.gdn_cfg.is_some()
2216 && self.g3n.is_none()
2217 && input_ids.len() > 8
2218 && std::env::var("CMF_MTP_CHAIN_PROBE").is_err()
2219 && std::env::var("CMF_METAL_PREFILL").as_deref() != Ok("0")
2220 {
2221 let chunk: usize = std::env::var("CMF_METAL_PREFILL_CHUNK")
2222 .ok()
2223 .and_then(|v| v.parse().ok())
2224 .filter(|&v| (16..=512).contains(&v))
2225 .unwrap_or(256);
2226 let hs = self.hidden_size;
2227 let _tp = std::time::Instant::now();
2228 while pos < input_ids.len() && !self.cancel.load(std::sync::atomic::Ordering::Relaxed) {
2229 let end = (pos + chunk).min(input_ids.len());
2230 let Some(hb) = self.prefill_batch_metal(&input_ids[pos..end], pos) else {
2231 break;
2232 };
2233 if let Some(m) = &mut mtp {
2234 let n_pairs = if end < input_ids.len() { end - pos } else { end - pos - 1 };
2235 if n_pairs > 0 {
2236 let pairs: Vec<(&[f32], u32)> = (0..n_pairs)
2237 .map(|j| (&hb[j * hs..(j + 1) * hs], input_ids[pos + j + 1]))
2238 .collect();
2239 if !self.mtp_warm_batch_metal(m, &pairs, pos) {
2240 for (j, (h, t)) in pairs.iter().enumerate() {
2241 let h = h.to_vec();
2242 let _ = self.mtp_step(m, &h, *t, pos + j);
2243 }
2244 }
2245 }
2246 }
2247 hidden.copy_from_slice(&hb[(end - pos - 1) * hs..]);
2248 pos = end;
2249 }
2250 if std::env::var("CMF_PREFILL_PROF").is_ok() {
2251 eprintln!(
2252 "metal-prefill: {} of {} tokens in {:.1} ms",
2253 pos,
2254 input_ids.len(),
2255 _tp.elapsed().as_secs_f64() * 1e3
2256 );
2257 }
2258 }
2259 if task_mask.is_none()
2260 && !dyn_prefill
2261 && !graph_prefill
2262 && self.can_prefill_batched()
2263 && self.g3n.is_none()
2264 && input_ids.len() > 2
2265 {
2266 let chunk = prefill_chunk();
2272 let hs = self.hidden_size;
2273 while pos < input_ids.len() && !self.cancel.load(std::sync::atomic::Ordering::Relaxed) {
2274 let end = (pos + chunk).min(input_ids.len());
2275 let hb = self.prefill_batch(&input_ids[pos..end], pos);
2276 if let Some(m) = &mut mtp {
2277 let probe: usize = std::env::var("CMF_MTP_CHAIN_PROBE")
2278 .ok()
2279 .and_then(|v| v.parse().ok())
2280 .unwrap_or(0);
2281 for p in pos..end {
2282 if p + 1 < input_ids.len() {
2283 if probe >= 1 && p + 2 < input_ids.len() {
2284 let (d1, mut hx) = self.mtp_step_h(
2288 m,
2289 &hb[(p - pos) * hs..(p - pos + 1) * hs],
2290 input_ids[p + 1],
2291 p,
2292 );
2293 let mut ok = d1 == input_ids[p + 2];
2294 Self::chain_probe_note(0, ok);
2295 let mut d_prev = d1;
2296 let mut extra = 0usize;
2297 for j in 1..probe {
2298 if p + 2 + j >= input_ids.len() {
2299 break;
2300 }
2301 let (dj, hj) = self.mtp_step_h(m, &hx, d_prev, p + 1 + j);
2302 extra += 1;
2303 ok = ok && dj == input_ids[p + 2 + j];
2304 Self::chain_probe_note(j, ok);
2305 d_prev = dj;
2306 hx = hj;
2307 }
2308 m.kv.truncate_last(extra);
2309 } else {
2310 let _ = self.mtp_step(
2311 m,
2312 &hb[(p - pos) * hs..(p - pos + 1) * hs],
2313 input_ids[p + 1],
2314 p,
2315 );
2316 }
2317 }
2318 }
2319 }
2320 hidden.copy_from_slice(&hb[(end - pos - 1) * hs..]);
2321 pos = end;
2322 }
2323 }
2324 let pair_off = std::env::var("CMF_PAIR").is_ok_and(|v| v == "0");
2325 if task_mask.is_none()
2326 && !dyn_prefill
2327 && !graph_prefill
2328 && !pair_off
2329 && self.pair_supported()
2330 {
2331 while pos + 1 < input_ids.len()
2332 && !self.cancel.load(std::sync::atomic::Ordering::Relaxed)
2333 {
2334 let e1 = self.embed_single(input_ids[pos]);
2335 let e2 = self.embed_single(input_ids[pos + 1]);
2336 let (h1, h2) = self.forward_pair(&e1, &e2, pos);
2337 self.commit_linear_scratch();
2339 if let Some(m) = &mut mtp {
2340 let _ = self.mtp_step(m, &h1, input_ids[pos + 1], pos);
2341 if pos + 2 < input_ids.len() {
2342 let probe: usize = std::env::var("CMF_MTP_CHAIN_PROBE")
2343 .ok()
2344 .and_then(|v| v.parse().ok())
2345 .unwrap_or(0);
2346 if probe >= 1 && pos + 3 < input_ids.len() {
2347 let (d1, mut hx) = self.mtp_step_h(m, &h2, input_ids[pos + 2], pos + 1);
2351 let mut ok = d1 == input_ids[pos + 3];
2352 Self::chain_probe_note(0, ok);
2353 let mut d_prev = d1;
2354 let mut extra = 0usize;
2355 for j in 1..probe {
2356 if pos + 3 + j >= input_ids.len() {
2357 break;
2358 }
2359 let (dj, hj) = self.mtp_step_h(m, &hx, d_prev, pos + 2 + j);
2360 extra += 1;
2361 ok = ok && dj == input_ids[pos + 3 + j];
2362 Self::chain_probe_note(j, ok);
2363 d_prev = dj;
2364 hx = hj;
2365 }
2366 m.kv.truncate_last(extra);
2367 } else {
2368 let _ = self.mtp_step(m, &h2, input_ids[pos + 2], pos + 1);
2369 }
2370 }
2371 }
2372 hidden = h2;
2373 pos += 2;
2374 }
2375 }
2376 if batch_k > 0
2385 && graph_prefill
2386 && task_mask.is_none()
2387 && !self.o1_active()
2388 && mtp.is_none()
2389 && !dyn_prefill
2390 && pos + 1 < input_ids.len()
2391 {
2392 let hs = self.hidden_size;
2393 let chunk = batch_k;
2394 while pos < input_ids.len() {
2395 let end = (pos + chunk).min(input_ids.len());
2396 let bk = end - pos;
2397 let mut hiddens = vec![0f32; bk * hs];
2398 for (j, &id) in input_ids[pos..end].iter().enumerate() {
2399 hiddens[j * hs..(j + 1) * hs].copy_from_slice(&self.embed_single(id));
2400 }
2401 let positions: Vec<usize> = (pos..end).collect();
2402 let t_chunk = std::time::Instant::now();
2403 let ok_b = self.try_batch_graph_wgpu(&mut hiddens, &positions, bk, None);
2404 if std::env::var("CMF_GRAPH_PROF").is_ok() {
2405 let ms = t_chunk.elapsed().as_secs_f64() * 1000.0;
2406 eprintln!(
2407 "batch-chunk: k={bk} ok={ok_b} {ms:.1} ms ({:.1} tok/s)",
2408 bk as f64 / (ms / 1000.0)
2409 );
2410 }
2411 {
2412 use std::sync::atomic::{AtomicBool, Ordering};
2413 static SAID: AtomicBool = AtomicBool::new(false);
2414 if !SAID.swap(true, Ordering::Relaxed) {
2415 if ok_b {
2416 tracing::info!("batched prefill: ACTIVE (k={bk})");
2417 } else {
2418 tracing::warn!("batched prefill declined — per-position graph");
2419 }
2420 }
2421 }
2422 if ok_b {
2423 hidden.copy_from_slice(&hiddens[(bk - 1) * hs..]);
2424 pos = end;
2425 } else {
2426 break; }
2428 }
2429 }
2430 while pos < input_ids.len() && !self.cancel.load(std::sync::atomic::Ordering::Relaxed) {
2431 self.graph_want_logits = fuse_lm && pos + 1 == input_ids.len();
2432 hidden = self.forward_layers(&self.embed_single(input_ids[pos]), pos, task_mask);
2433 if let Some(m) = &mut mtp {
2434 if pos + 1 < input_ids.len() {
2435 let probe: usize = std::env::var("CMF_MTP_CHAIN_PROBE")
2441 .ok()
2442 .and_then(|v| v.parse().ok())
2443 .unwrap_or(0);
2444 if probe >= 1 && pos + 2 < input_ids.len() {
2445 let (d1, mut hx) = self.mtp_step_h(m, &hidden, input_ids[pos + 1], pos);
2446 let mut ok = d1 == input_ids[pos + 2];
2447 Self::chain_probe_note(0, ok);
2448 let mut d_prev = d1;
2449 let mut extra = 0usize;
2450 for j in 1..probe {
2451 if pos + 2 + j >= input_ids.len() {
2452 break;
2453 }
2454 let (dj, hj) = self.mtp_step_h(m, &hx, d_prev, pos + 1 + j);
2455 extra += 1;
2456 ok = ok && dj == input_ids[pos + 2 + j];
2457 Self::chain_probe_note(j, ok);
2458 d_prev = dj;
2459 hx = hj;
2460 }
2461 m.kv.truncate_last(extra);
2464 } else {
2465 let _ = self.mtp_step(m, &hidden, input_ids[pos + 1], pos);
2466 }
2467 }
2468 }
2469 pos += 1;
2470 }
2471 if std::env::var("CMF_PREFILL_PROF").is_ok() {
2472 eprintln!(
2473 "prefill: {} tokens in {:.1} ms (batch_k={batch_k})",
2474 input_ids.len(),
2475 _tpf.elapsed().as_secs_f64() * 1000.0
2476 );
2477 }
2478 if self
2481 .cancel
2482 .swap(false, std::sync::atomic::Ordering::Relaxed)
2483 {
2484 self.kv_history.clear();
2485 if let Some(m) = mtp {
2486 self.mtp = Some(m);
2487 }
2488 return Ok(GenerateResult {
2489 text: String::new(),
2490 token_ids: Vec::new(),
2491 prompt_tokens: input_ids.len(),
2492 tokens_generated: 0,
2493 finish_reason: "cancelled".to_string(),
2494 mtp_drafted: 0,
2495 mtp_accepted: 0,
2496 token_confidence: Vec::new(),
2497 traces: Vec::new(),
2498 });
2499 }
2500
2501 self.o1_seal();
2504
2505 macro_rules! commit {
2507 ($id:expr) => {{
2508 all_ids.push($id);
2509 generated += 1;
2510 if self.tokenizer.is_eos($id) {
2511 finish_reason = "stop".to_string();
2512 false
2513 } else {
2514 let token_text = self.tokenizer.decode_token($id);
2515 let mut go = true;
2516 if let Some(ref mut cb) = on_token {
2517 if !cb(&token_text) {
2518 finish_reason = "cancelled".to_string();
2519 go = false;
2520 }
2521 }
2522 go
2523 }
2524 }};
2525 }
2526
2527 let mut spec_trial = SpecTrial::Spec {
2538 t0: std::time::Instant::now(),
2539 gen0: generated,
2540 rounds: 0,
2541 };
2542 let mut spec_mon = SpecMon::default();
2543 let mut spec_watchdog_off = false;
2544 let mut next_pos = input_ids.len();
2546 'decode: while generated < max_tokens {
2547 if self
2548 .cancel
2549 .swap(false, std::sync::atomic::Ordering::Relaxed)
2550 {
2551 finish_reason = "cancelled".to_string();
2552 break 'decode;
2553 }
2554 let forced = self.spec_forced.take();
2559 let mut logits = match (forced, self.graph_logits.take()) {
2560 (Some(_), _) => Vec::new(),
2561 (None, Some(lg)) => lg,
2562 (None, None) => {
2563 inference::rms_norm_into(
2564 &hidden,
2565 &self.weights.final_norm,
2566 self.rms_eps,
2567 self.norm_style,
2568 &mut self.ws.n1,
2569 );
2570 self.lm_head_forward(&self.ws.n1)
2571 }
2572 };
2573 if generated
2576 == std::env::var("CMF_LOGIT_DUMP_STEP")
2577 .ok()
2578 .and_then(|v| v.parse().ok())
2579 .unwrap_or(0)
2580 {
2581 if let Ok(path) = std::env::var("CMF_LOGIT_DUMP") {
2582 let mut bytes: Vec<u8> = Vec::with_capacity((hidden.len() + logits.len()) * 4);
2583 for v in hidden.iter().chain(logits.iter()) {
2584 bytes.extend_from_slice(&v.to_le_bytes());
2585 }
2586 let _ = std::fs::write(&path, &bytes);
2587 }
2588 }
2589 let t_next = match forced {
2590 Some(c) => c,
2591 None => sampler::sample_with_scratch_pool(
2592 &logits,
2593 &self.sampler_config,
2594 &all_ids,
2595 &mut self.rng,
2596 &mut self.sampler_scratch,
2597 self.pool.as_deref(),
2598 ),
2599 };
2600 if self.confidence_on {
2601 confidence.push(if logits.is_empty() {
2602 0.0
2603 } else {
2604 sampler::top1_prob_pool(
2605 self.pool.as_deref(),
2606 &mut self.sampler_scratch,
2607 &logits,
2608 t_next,
2609 calib_temp,
2610 )
2611 });
2612 }
2613 if !logits.is_empty() {
2614 attention::recycle_buf(&mut logits);
2615 }
2616 if trace_on {
2617 let skill = router.as_ref().and_then(|r| r.active_id());
2621 traces.push(TokenTrace {
2622 t: generated,
2623 token_id: t_next,
2624 confidence: confidence.last().copied().unwrap_or(0.0),
2625 active_skill: skill,
2626 recon: None,
2627 switched: false,
2628 });
2629 }
2630 if !commit!(t_next) {
2631 break 'decode;
2632 }
2633 if generated >= max_tokens {
2634 break 'decode;
2635 }
2636
2637 if self.kv_cache.needs_eviction() {
2638 static SAID: std::sync::Once = std::sync::Once::new();
2644 SAID.call_once(|| {
2645 tracing::warn!(
2646 "KV cache full at {} positions — evicting half; quality \
2647 will degrade. Raise CMF_MAX_SEQ.",
2648 self.kv_cache.max_seq_len,
2649 );
2650 });
2651 let keep = (self.kv_cache.max_seq_len / 2).max(1);
2652 self.kv_cache.evict(keep);
2653 }
2654
2655 if graph_spec {
2658 match spec_trial {
2659 SpecTrial::Plain { t0, gen0 } if generated >= gen0 + 8 => {
2660 spec_mon.plain_ms =
2661 t0.elapsed().as_secs_f64() * 1e3 / (generated - gen0) as f64;
2662 let keep = spec_mon.pays();
2663 tracing::info!(
2664 "speculation trial: {:.2} tok/round in {:.1} ms vs plain {:.1} ms/tok — {}",
2665 spec_mon.tokens,
2666 spec_mon.round_ms,
2667 spec_mon.plain_ms,
2668 if keep { "speculating" } else { "plain" }
2669 );
2670 spec_mon.fails = 0;
2671 spec_trial = SpecTrial::Decided {
2672 spec: keep,
2673 recheck_at: if keep { usize::MAX } else { generated + 128 },
2674 };
2675 }
2676 SpecTrial::Decided { recheck_at, .. } if generated >= recheck_at => {
2677 spec_mon.n = 0;
2678 spec_trial = SpecTrial::Spec {
2679 t0: std::time::Instant::now(),
2680 gen0: generated,
2681 rounds: 0,
2682 };
2683 }
2684 _ => {}
2685 }
2686 spec_watchdog_off = matches!(
2687 spec_trial,
2688 SpecTrial::Plain { .. } | SpecTrial::Decided { spec: false, .. }
2689 );
2690 }
2691 match &mut mtp {
2692 #[cfg(feature = "gpu")]
2694 Some(m)
2695 if graph_spec
2696 && !spec_watchdog_off
2697 && generated + 1 < max_tokens
2698 && next_pos > 0 =>
2699 {
2700 let t_round = std::time::Instant::now();
2701 if let Some((extra, n_pos, new_h)) = self.graph_spec_step(
2702 m,
2703 &hidden,
2704 t_next,
2705 next_pos,
2706 &mut drafted,
2707 &mut accepted,
2708 &mut all_ids,
2709 ) {
2710 next_pos = n_pos;
2711 hidden = new_h;
2712 if std::env::var("CMF_GRAPH_SPEC_TIME").is_ok() {
2713 eprintln!(
2714 "spec-round wall {:.1} ms → {} tokens",
2715 t_round.elapsed().as_secs_f64() * 1e3,
2716 extra.len() + 1
2717 );
2718 }
2719 spec_mon.round(t_round.elapsed().as_secs_f64() * 1e3, extra.len() + 1);
2723 spec_trial = Self::spec_trial_round(
2726 spec_trial,
2727 &mut spec_mon,
2728 generated + extra.len() + 1,
2729 );
2730 let mut stopped = false;
2731 for &id in &extra {
2732 if self.confidence_on {
2733 confidence.push(0.0);
2734 }
2735 if !commit!(id) {
2736 stopped = true;
2737 break;
2738 }
2739 }
2740 if stopped {
2741 break 'decode;
2742 }
2743 continue 'decode;
2744 }
2745 spec_mon.round(t_round.elapsed().as_secs_f64() * 1e3, 1);
2756 spec_mon.tokens = 0.0;
2757 spec_mon.fails = 3;
2758 spec_trial = Self::spec_trial_round(spec_trial, &mut spec_mon, generated + 1);
2759 hidden = self.forward_layers(&self.embed_single(t_next), next_pos, task_mask);
2760 next_pos += 1;
2761 continue 'decode;
2762 }
2763 Some(m) if !graph_spec && generated + 1 < max_tokens => {
2765 let draft = self.mtp_step(m, &hidden, t_next, next_pos - 1);
2766 drafted += 1;
2767 let emb1 = self.embed_single(t_next);
2768 let emb2 = self.embed_single(draft);
2769 let (h1, h2) = self.forward_pair(&emb1, &emb2, next_pos);
2770
2771 inference::rms_norm_into(
2772 &h1,
2773 &self.weights.final_norm,
2774 self.rms_eps,
2775 self.norm_style,
2776 &mut self.ws.n1,
2777 );
2778 let mut logits1 = self.lm_head_forward(&self.ws.n1);
2779 let t_after = sampler::sample_with_scratch_pool(
2780 &logits1,
2781 &self.sampler_config,
2782 &all_ids,
2783 &mut self.rng,
2784 &mut self.sampler_scratch,
2785 self.pool.as_deref(),
2786 );
2787 if self.confidence_on {
2788 confidence.push(sampler::top1_prob_pool(
2789 self.pool.as_deref(),
2790 &mut self.sampler_scratch,
2791 &logits1,
2792 t_after,
2793 calib_temp,
2794 ));
2795 }
2796 attention::recycle_buf(&mut logits1);
2797 if trace_on {
2798 traces.push(TokenTrace {
2801 t: generated,
2802 token_id: t_after,
2803 confidence: confidence.last().copied().unwrap_or(0.0),
2804 active_skill: None,
2805 recon: None,
2806 switched: false,
2807 });
2808 }
2809 let stop = !commit!(t_after);
2810
2811 if t_after == draft {
2812 accepted += 1;
2813 self.commit_linear_scratch();
2814 let _ = self.mtp_step(m, &h1, t_after, next_pos);
2815 hidden = h2;
2816 next_pos += 2;
2817 } else {
2818 for layer in &mut self.kv_cache.layers {
2820 layer.truncate_last(1);
2821 }
2822 if !stop {
2823 let _ = self.mtp_step(m, &h1, t_after, next_pos);
2824 hidden = self.forward_layers(
2825 &self.embed_single(t_after),
2826 next_pos + 1,
2827 None,
2828 );
2829 }
2830 next_pos += 2;
2831 }
2832 if stop {
2833 break 'decode;
2834 }
2835 }
2836 _ => {
2838 #[cfg(feature = "gpu")]
2843 if Self::dsv4_spec_on() && self.dsv4.is_some() {
2844 static SAID: std::sync::Once = std::sync::Once::new();
2845 SAID.call_once(|| {
2846 eprintln!(
2847 "dsv4-spec гейт: mtp={} mask={} router={} trace={} temp={} rep={} ",
2848 !self.dsv4_mtp.is_empty(),
2849 task_mask.is_none(),
2850 router.is_none(),
2851 !trace_on,
2852 self.sampler_config.temperature < 1e-6,
2853 self.sampler_config.repetition_penalty == 1.0,
2854 );
2855 });
2856 }
2857 #[cfg(feature = "gpu")]
2858 if Self::dsv4_spec_on()
2859 && self.dsv4.is_some()
2860 && !self.dsv4_mtp.is_empty()
2861 && task_mask.is_none()
2862 && router.is_none()
2863 && !trace_on
2864 && self.sampler_config.temperature < 1e-6
2865 && self.sampler_config.repetition_penalty == 1.0
2866 && generated + 1 < max_tokens
2867 && all_ids.len() >= 2
2868 {
2869 let tip_token = all_ids[all_ids.len() - 2];
2870 if let Some((extra, n_pos)) = self.dsv4_spec_step(
2871 tip_token,
2872 t_next,
2873 next_pos,
2874 &mut drafted,
2875 &mut accepted,
2876 ) {
2877 next_pos = n_pos;
2878 let mut stopped = false;
2879 for &id in &extra {
2880 if self.confidence_on {
2881 confidence.push(0.0);
2882 }
2883 if !commit!(id) {
2884 stopped = true;
2885 break;
2886 }
2887 }
2888 if stopped {
2889 break 'decode;
2890 }
2891 continue 'decode;
2892 }
2893 }
2894 self.graph_want_logits = fuse_lm;
2895 let mut t_fwd = t_next;
2901 let pure_greedy = self.sampler_config.temperature < 1e-6
2902 && self.sampler_config.repetition_penalty == 1.0
2903 && self.sampler_config.suppress_tokens.is_empty();
2904 let burst_k = std::env::var("CMF_MULTISTEP")
2909 .ok()
2910 .and_then(|v| v.parse::<usize>().ok())
2911 .unwrap_or(0);
2912 if pure_greedy
2913 && burst_k >= 1
2914 && fuse_lm
2915 && task_mask.is_none()
2916 && router.is_none()
2917 && !trace_on
2918 && !self.confidence_on
2919 {
2920 let mut stopped = false;
2921 loop {
2922 let room = max_tokens.saturating_sub(generated);
2923 if room <= 2 {
2924 break;
2925 }
2926 let k = burst_k.min(room - 1);
2927 if k < 1 {
2928 break;
2929 }
2930 let Some(ids) = self.try_multi_burst(t_fwd, next_pos, k) else {
2931 break;
2932 };
2933 next_pos += k;
2934 for &id in &ids {
2935 if !commit!(id) {
2936 stopped = true;
2937 break;
2938 }
2939 }
2940 if stopped {
2941 break;
2942 }
2943 t_fwd = *ids.last().unwrap();
2944 }
2945 if stopped {
2946 break 'decode;
2947 }
2948 }
2949 hidden = self.forward_layers(&self.embed_single(t_fwd), next_pos, task_mask);
2950 next_pos += 1;
2951 if let Some(r) = &mut router {
2954 let phi = self.dyn_phi_ema.clone();
2955 let decision = r.step(&phi, generated);
2956 if let Some(new_active) = decision {
2957 let _ = self.set_active_skill(new_active);
2958 }
2959 if trace_on {
2962 if let Some(last) = traces.last_mut() {
2963 let e = r.last_best_e();
2964 last.recon = e.is_finite().then_some(e);
2965 last.switched = decision.is_some();
2966 }
2967 }
2968 }
2969 }
2970 }
2971 }
2972
2973 self.graph_want_logits = false;
2974 self.graph_logits = None;
2975 if router.is_some() {
2977 let _ = self.set_active_skill(None);
2978 }
2979 self.dyn_router = router.or(self.dyn_router.take());
2980 self.mtp = mtp.or(self.mtp.take());
2981
2982 let output_ids = &all_ids[input_ids.len()..];
2983 let forwarded = input_ids.len() + output_ids.len().saturating_sub(1);
2987 self.kv_history = all_ids[..forwarded.min(all_ids.len())].to_vec();
2988 confidence.truncate(output_ids.len()); traces.truncate(output_ids.len());
2990 Ok(GenerateResult {
2991 text: self.tokenizer.decode(output_ids),
2992 token_ids: output_ids.to_vec(),
2993 prompt_tokens: input_ids.len(),
2994 tokens_generated: generated,
2995 finish_reason,
2996 mtp_drafted: drafted,
2997 mtp_accepted: accepted,
2998 token_confidence: confidence,
2999 traces,
3000 })
3001 }
3002
3003 fn mtp_step(
3007 &mut self,
3008 m: &mut MtpModule,
3009 hidden: &[f32],
3010 next_token: u32,
3011 position: usize,
3012 ) -> u32 {
3013 self.mtp_step_h(m, hidden, next_token, position).0
3014 }
3015
3016 fn chain_probe_note(depth: usize, prefix_ok: bool) {
3020 use std::sync::Mutex;
3021 static T: Mutex<Vec<(u64, u64)>> = Mutex::new(Vec::new());
3022 let mut t = T.lock().unwrap();
3023 if t.len() <= depth {
3024 t.resize(depth + 1, (0, 0));
3025 }
3026 t[depth].0 += 1;
3027 t[depth].1 += prefix_ok as u64;
3028 if depth == 0 && t[0].0 % 128 == 0 {
3029 let line: Vec<String> = t
3030 .iter()
3031 .enumerate()
3032 .map(|(d, (n, k))| {
3033 format!(
3034 "d{}={:.0}%({n})",
3035 d + 1,
3036 100.0 * *k as f64 / (*n).max(1) as f64
3037 )
3038 })
3039 .collect();
3040 eprintln!("mtp-chain: {}", line.join(" "));
3041 }
3042 }
3043
3044 fn mtp_step_hl(
3052 &mut self,
3053 m: &mut MtpModule,
3054 hidden: &[f32],
3055 next_token: u32,
3056 position: usize,
3057 ) -> (Vec<f32>, Vec<f32>) {
3058 #[cfg(target_os = "macos")]
3063 if self.mtp_graph_mode != Some(false) && crate::gpu::q1_force() {
3064 if let Some(r) = self.mtp_step_metal(m, hidden, next_token, position, true) {
3065 self.mtp_graph_mode = Some(true);
3066 return r;
3067 }
3068 self.mtp_graph_mode = Some(false);
3069 }
3070 #[cfg(feature = "gpu")]
3071 if self.mtp_graph_mode != Some(false) {
3072 if let Some(r) = self.mtp_step_graph(m, hidden, next_token, position) {
3073 self.mtp_graph_mode = Some(true);
3074 return r;
3075 }
3076 if self.mtp_graph_mode == Some(true) {
3077 tracing::warn!("mtp graph declined mid-run — draft falls to the per-op path");
3082 }
3083 self.mtp_graph_mode = Some(false);
3084 }
3085 let e = self.embed_single(next_token);
3089 let mut cat = vec![0.0f32; 2 * self.hidden_size];
3090 let (cat_e, cat_h) = cat.split_at_mut(self.hidden_size);
3091 inference::rms_norm_into(&e, &m.enorm, self.rms_eps, self.norm_style, cat_e);
3092 inference::rms_norm_into(hidden, &m.hnorm, self.rms_eps, self.norm_style, cat_h);
3093 let mut x = vec![0.0f32; self.hidden_size];
3094 m.eh_proj.matvec(&cat, &mut x, self.pool.as_deref());
3095
3096 let lw = &m.layer;
3098 inference::rms_norm_into(
3099 &x,
3100 &lw.input_norm,
3101 self.rms_eps,
3102 self.norm_style,
3103 &mut self.ws.n1,
3104 );
3105 let attn = match &lw.attn {
3106 AttnKind::Mla(_) => unreachable!("MLA has no MTP/pair path"),
3108 AttnKind::Kda(_) => unreachable!("KDA has no MTP/pair path"),
3109 AttnKind::Full {
3110 wq,
3111 wk,
3112 wv,
3113 wo,
3114 q_norm,
3115 k_norm,
3116 output_gate,
3117 softplus_gate,
3118 bias,
3119 } => {
3120 let mut cfg = self.attn_cfg(position);
3121 cfg.q_norm = q_norm.as_deref();
3122 cfg.k_norm = k_norm.as_deref();
3123 cfg.output_gate = *output_gate;
3124 cfg.softplus_gate = softplus_gate
3125 .as_ref()
3126 .map(|(gate, per_head)| (gate, *per_head));
3127 cfg.bias = bias
3128 .as_ref()
3129 .map(|(q, k, v)| (q.as_slice(), k.as_slice(), v.as_slice()));
3130 attention::qwen_attention(&self.ws.n1, wq, wk, wv, wo, &mut m.kv, &cfg)
3131 }
3132 AttnKind::Linear(_) | AttnKind::LinearGdn(_) | AttnKind::ShortConv(_) => {
3133 unreachable!("MTP block is full attention")
3134 }
3135 };
3136 for (i, &a) in attn.iter().enumerate() {
3137 x[i] += a;
3138 }
3139 inference::rms_norm_into(
3140 &x,
3141 &lw.post_norm,
3142 self.rms_eps,
3143 self.norm_style,
3144 &mut self.ws.p1,
3145 );
3146 let ffn = ffn_forward(&lw.ffn, &self.ws.p1, self.pool.as_deref(), None);
3147 for (i, &f) in ffn.iter().enumerate() {
3148 x[i] += f;
3149 }
3150
3151 inference::rms_norm_into(
3152 &x,
3153 &m.final_norm,
3154 self.rms_eps,
3155 self.norm_style,
3156 &mut self.ws.n1,
3157 );
3158 let lg = self.lm_head_forward(&self.ws.n1);
3159 (lg, x)
3160 }
3161
3162 fn mtp_step_h(
3164 &mut self,
3165 m: &mut MtpModule,
3166 hidden: &[f32],
3167 next_token: u32,
3168 position: usize,
3169 ) -> (u32, Vec<f32>) {
3170 let (mut lg, x) = self.mtp_step_hl(m, hidden, next_token, position);
3171 let draft = sampler::argmax(&lg);
3172 attention::recycle_buf(&mut lg);
3173 (draft, x)
3174 }
3175
3176 fn spec_trial_round(trial: SpecTrial, mon: &mut SpecMon, generated: usize) -> SpecTrial {
3182 match trial {
3183 SpecTrial::Spec { t0, gen0, rounds } => {
3184 let rounds = rounds + 1;
3185 if rounds >= 5 {
3186 if mon.plain_ms > 0.0 {
3187 let keep = mon.pays();
3188 mon.fails = 0;
3189 tracing::info!(
3190 "speculation re-check: {:.2} tok/round in {:.1} ms vs plain {:.1} ms/tok — {}",
3191 mon.tokens,
3192 mon.round_ms,
3193 mon.plain_ms,
3194 if keep { "speculating" } else { "plain" }
3195 );
3196 SpecTrial::Decided {
3197 spec: keep,
3198 recheck_at: if keep { usize::MAX } else { generated + 128 },
3199 }
3200 } else {
3201 SpecTrial::Plain {
3202 t0: std::time::Instant::now(),
3203 gen0: generated,
3204 }
3205 }
3206 } else {
3207 SpecTrial::Spec { t0, gen0, rounds }
3208 }
3209 }
3210 SpecTrial::Decided { spec: true, .. } => {
3211 if mon.pays() {
3212 mon.fails = 0;
3213 trial
3214 } else {
3215 mon.fails += 1;
3216 if mon.fails >= 4 {
3217 tracing::info!(
3218 "speculation stopped: {:.2} tok/round in {:.1} ms vs plain {:.1} ms/tok",
3219 mon.tokens,
3220 mon.round_ms,
3221 mon.plain_ms
3222 );
3223 SpecTrial::Decided {
3224 spec: false,
3225 recheck_at: generated + 128,
3226 }
3227 } else {
3228 trial
3229 }
3230 }
3231 }
3232 other => other,
3233 }
3234 }
3235
3236 fn mtp_kv_id(&self) -> u64 {
3239 self.graph_kv_id | (1u64 << 40)
3240 }
3241
3242 const MTP_LAYER_BASE: usize = 0;
3247
3248 fn mtp_block_input(&mut self, m: &MtpModule, hidden: &[f32], next_token: u32) -> Vec<f32> {
3251 let e = self.embed_single(next_token);
3252 let mut cat = vec![0.0f32; 2 * self.hidden_size];
3253 let (cat_e, cat_h) = cat.split_at_mut(self.hidden_size);
3254 inference::rms_norm_into(&e, &m.enorm, self.rms_eps, self.norm_style, cat_e);
3255 inference::rms_norm_into(hidden, &m.hnorm, self.rms_eps, self.norm_style, cat_h);
3256 let mut x = vec![0.0f32; self.hidden_size];
3257 m.eh_proj.matvec(&cat, &mut x, self.pool.as_deref());
3258 x
3259 }
3260
3261 #[cfg(feature = "gpu")]
3264 fn mtp_graph_ok(&self, m: &MtpModule) -> bool {
3265 if std::env::var("CMF_MTP_GRAPH").as_deref() == Ok("0") {
3266 return false;
3267 }
3268 if !crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Decode)
3269 || !crate::gpu::enabled_here()
3270 || self.attn_softcap > 0.0
3271 || self.attention_heads_per_layer.is_some()
3272 {
3273 return false;
3274 }
3275 matches!(
3276 &m.layer.attn,
3277 AttnKind::Full {
3278 softplus_gate: None,
3279 ..
3280 }
3281 ) && matches!(&m.layer.ffn, FfnKind::Dense(_))
3282 }
3283
3284 #[cfg(feature = "gpu")]
3290 fn mtp_step_graph(
3291 &mut self,
3292 m: &mut MtpModule,
3293 hidden: &[f32],
3294 next_token: u32,
3295 position: usize,
3296 ) -> Option<(Vec<f32>, Vec<f32>)> {
3297 if !self.mtp_graph_ok(m) {
3298 return None;
3299 }
3300 let lw = &m.layer;
3301 let AttnKind::Full {
3302 wq,
3303 wk,
3304 wv,
3305 wo,
3306 q_norm,
3307 k_norm,
3308 output_gate,
3309 softplus_gate,
3310 bias,
3311 } = &lw.attn
3312 else {
3313 return None;
3314 };
3315 if softplus_gate.is_some() {
3316 return None;
3317 }
3318 let FfnKind::Dense(d) = &lw.ffn else {
3319 return None;
3320 };
3321 if !d.segs.is_empty() {
3322 return None; }
3324 let mut x = self.mtp_block_input(m, hidden, next_token);
3327 fn gw(t: &QTensor) -> Option<crate::gpu::GraphW<'_>> {
3328 let (_, i, kind, rs) = t.graph_weight()?;
3329 Some(crate::gpu::GraphW {
3330 idx: i,
3331 kind,
3332 row_scale: rs,
3333 data: &[],
3334 })
3335 }
3336 let (model, _, _, _) = wq.graph_weight()?;
3337 let model = model.clone();
3338 let (lm_gw, lm_rows) = {
3339 let (_, i, kind, rs) = self.weights.lm_head.graph_weight()?;
3340 (
3341 crate::gpu::GraphW {
3342 idx: i,
3343 kind,
3344 row_scale: rs,
3345 data: &[],
3346 },
3347 self.weights.lm_head.rows(),
3348 )
3349 };
3350 let layer = crate::gpu::GraphLayer {
3351 input_norm: &lw.input_norm,
3352 attn: crate::gpu::GraphAttn::Full {
3353 wq: gw(wq)?,
3354 wk: gw(wk)?,
3355 wv: gw(wv)?,
3356 wo: gw(wo)?,
3357 q_norm: q_norm.as_deref(),
3358 k_norm: k_norm.as_deref(),
3359 bias: bias
3360 .as_ref()
3361 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
3362 output_gate: *output_gate,
3363 cpu_k: m.kv.k_heads(),
3364 cpu_v: m.kv.v_heads(),
3365 },
3366 post_norm: &lw.post_norm,
3367 ffn: crate::gpu::GraphFfn::Dense {
3368 gate: gw(&d.gate_proj)?,
3369 up: gw(&d.up_proj)?,
3370 down: gw(&d.down_proj)?,
3371 },
3372 };
3373 let nh = self.num_heads;
3374 let (nkv, hd, rd) = self.layer_geom(0);
3375 let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
3376 let mut logits = Vec::new();
3377 let ok = crate::gpu::forward_token_graph(
3378 &model,
3379 self.mtp_kv_id(),
3380 std::slice::from_ref(&layer),
3381 &[None],
3382 self.o1_epoch,
3383 &self.inv_freq,
3384 &mut x,
3385 nh,
3386 nkv,
3387 hd,
3388 rd,
3389 self.hidden_size,
3390 self.intermediate_size,
3391 position,
3392 self.kv_cache.max_seq_len,
3393 gemma,
3394 self.rms_eps as f32,
3395 Some((&lm_gw, lm_rows)),
3396 &m.final_norm,
3397 &mut logits,
3398 &[],
3399 1,
3400 None,
3401 None,
3402 None,
3403 Self::MTP_LAYER_BASE,
3404 true,
3405 );
3406 if !ok {
3407 return None;
3408 }
3409 logits.resize(self.vocab_size, 0.0);
3410 Some((logits, x))
3411 }
3412
3413 #[cfg(feature = "gpu")]
3420 fn mtp_warm_graph(
3421 &mut self,
3422 m: &mut MtpModule,
3423 pairs: &[(&[f32], u32)],
3424 first_pos: usize,
3425 ) -> bool {
3426 if pairs.is_empty() || !self.mtp_graph_ok(m) {
3427 return pairs.is_empty();
3428 }
3429 let hs = self.hidden_size;
3430 let mut hiddens = Vec::with_capacity(pairs.len() * hs);
3433 for (h, t) in pairs {
3434 hiddens.extend_from_slice(&self.mtp_block_input(m, h, *t));
3435 }
3436 let lw = &m.layer;
3437 let AttnKind::Full {
3438 wq,
3439 wk,
3440 wv,
3441 wo,
3442 q_norm,
3443 k_norm,
3444 output_gate,
3445 bias,
3446 ..
3447 } = &lw.attn
3448 else {
3449 return false;
3450 };
3451 let FfnKind::Dense(d) = &lw.ffn else {
3452 return false;
3453 };
3454 if !d.segs.is_empty() {
3455 return false; }
3457 fn gw(t: &QTensor) -> Option<crate::gpu::GraphW<'_>> {
3458 let (_, i, kind, rs) = t.graph_weight()?;
3459 Some(crate::gpu::GraphW {
3460 idx: i,
3461 kind,
3462 row_scale: rs,
3463 data: &[],
3464 })
3465 }
3466 let Some((model, _, _, _)) = wq.graph_weight() else {
3467 return false;
3468 };
3469 let model = model.clone();
3470 let (Some(gwq), Some(gwk), Some(gwv), Some(gwo), Some(gg), Some(gu), Some(gd)) = (
3471 gw(wq),
3472 gw(wk),
3473 gw(wv),
3474 gw(wo),
3475 gw(&d.gate_proj),
3476 gw(&d.up_proj),
3477 gw(&d.down_proj),
3478 ) else {
3479 return false;
3480 };
3481 let layer = crate::gpu::GraphLayer {
3482 input_norm: &lw.input_norm,
3483 attn: crate::gpu::GraphAttn::Full {
3484 wq: gwq,
3485 wk: gwk,
3486 wv: gwv,
3487 wo: gwo,
3488 q_norm: q_norm.as_deref(),
3489 k_norm: k_norm.as_deref(),
3490 bias: bias
3491 .as_ref()
3492 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
3493 output_gate: *output_gate,
3494 cpu_k: m.kv.k_heads(),
3495 cpu_v: m.kv.v_heads(),
3496 },
3497 post_norm: &lw.post_norm,
3498 ffn: crate::gpu::GraphFfn::Dense {
3499 gate: gg,
3500 up: gu,
3501 down: gd,
3502 },
3503 };
3504 let positions: Vec<usize> = (first_pos..first_pos + pairs.len()).collect();
3505 let nh = self.num_heads;
3506 let (nkv, hd, rd) = self.layer_geom(0);
3507 let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
3508 crate::gpu::forward_batch_graph(
3509 &model,
3510 self.mtp_kv_id(),
3511 std::slice::from_ref(&layer),
3512 &self.inv_freq,
3513 &mut hiddens,
3514 nh,
3515 nkv,
3516 hd,
3517 rd,
3518 hs,
3519 self.intermediate_size,
3520 &positions,
3521 self.kv_cache.max_seq_len,
3522 gemma,
3523 self.rms_eps as f32,
3524 pairs.len(),
3525 None,
3526 )
3527 }
3528
3529 fn mtp_warm(&mut self, m: &mut MtpModule, hidden: &[f32], next_token: u32, position: usize) {
3533 let e = self.embed_single(next_token);
3534 let mut cat = vec![0.0f32; 2 * self.hidden_size];
3535 let (cat_e, cat_h) = cat.split_at_mut(self.hidden_size);
3536 inference::rms_norm_into(&e, &m.enorm, self.rms_eps, self.norm_style, cat_e);
3537 inference::rms_norm_into(hidden, &m.hnorm, self.rms_eps, self.norm_style, cat_h);
3538 let mut x = vec![0.0f32; self.hidden_size];
3539 m.eh_proj.matvec(&cat, &mut x, self.pool.as_deref());
3540 inference::rms_norm_into(
3541 &x,
3542 &m.layer.input_norm,
3543 self.rms_eps,
3544 self.norm_style,
3545 &mut self.ws.n1,
3546 );
3547 let attn = match &m.layer.attn {
3548 AttnKind::Full {
3549 wq,
3550 wk,
3551 wv,
3552 wo,
3553 q_norm,
3554 k_norm,
3555 output_gate,
3556 softplus_gate,
3557 bias,
3558 } => {
3559 let mut cfg = self.attn_cfg(position);
3560 cfg.q_norm = q_norm.as_deref();
3561 cfg.k_norm = k_norm.as_deref();
3562 cfg.output_gate = *output_gate;
3563 cfg.softplus_gate = softplus_gate.as_ref().map(|(g, p)| (g, *p));
3564 cfg.bias = bias
3565 .as_ref()
3566 .map(|(q, k, v)| (q.as_slice(), k.as_slice(), v.as_slice()));
3567 attention::qwen_attention(&self.ws.n1, wq, wk, wv, wo, &mut m.kv, &cfg)
3568 }
3569 _ => return,
3570 };
3571 let _ = attn;
3572 }
3573
3574 #[cfg(feature = "gpu")]
3581 #[allow(clippy::too_many_arguments)]
3582 fn graph_spec_step(
3583 &mut self,
3584 m: &mut MtpModule,
3585 hidden: &[f32],
3586 t_next: u32,
3587 next_pos: usize,
3588 drafted: &mut usize,
3589 accepted: &mut usize,
3590 all_ids: &mut Vec<u32>,
3594 ) -> Option<(Vec<u32>, usize, Vec<f32>)> {
3595 #[cfg(target_os = "macos")]
3606 let metal_native = crate::gpu::q1_force();
3607 #[cfg(not(target_os = "macos"))]
3608 let metal_native = false;
3609 #[cfg(feature = "gpu")]
3610 let k_default = if metal_native {
3611 7
3614 } else if crate::gpu_wgpu::verify_i8_on() {
3615 5
3616 } else {
3617 4
3618 };
3619 #[cfg(not(feature = "gpu"))]
3620 let k_default = 4;
3621 let k_spec: usize = std::env::var("CMF_GRAPH_SPEC_K")
3622 .ok()
3623 .and_then(|v| v.parse().ok())
3624 .filter(|&v| (1..=8).contains(&v))
3625 .unwrap_or(k_default);
3626 if next_pos == 0 {
3627 return None;
3628 }
3629 let t_round = std::time::Instant::now();
3630 let subs = || crate::gpu_wgpu::SUBMITS.load(std::sync::atomic::Ordering::Relaxed);
3646 let sub0 = subs();
3647 let cfg = self.sampler_config.clone();
3652 let penalized = !(cfg.repetition_penalty == 1.0
3653 && cfg.presence_penalty == 0.0
3654 && cfg.suppress_tokens.is_empty());
3655 let greedy_pen = cfg.temperature < 1e-6 && penalized;
3660 let sampling = cfg.temperature >= 1e-6;
3661 let sparse = sampling && sampler::sparse_ok(&cfg);
3667 let base_len = all_ids.len();
3668 if sampling && !sparse && self.spec_q.len() < k_spec {
3669 self.spec_q.resize_with(k_spec, Vec::new);
3670 }
3671 if sparse && self.spec_qs.len() < k_spec {
3672 self.spec_qs.resize_with(k_spec, Vec::new);
3673 }
3674 let mut drafts = Vec::with_capacity(k_spec);
3679 let mut hx = hidden.to_vec();
3680 let spec_dbg = std::env::var("CMF_SPEC_DBG").is_ok();
3683 for j in 0..k_spec {
3684 let tok_in = if j == 0 { t_next } else { drafts[j - 1] };
3685 let mut dbg_ref: Option<(Vec<f32>, Vec<f32>)> = None;
3686 if spec_dbg {
3687 let saved = self.mtp_graph_mode;
3688 self.mtp_graph_mode = Some(false);
3689 let r = self.mtp_step_hl(m, &hx, tok_in, next_pos - 1 + j);
3690 self.mtp_graph_mode = saved;
3691 m.kv.truncate_last(1);
3692 dbg_ref = Some(r);
3693 }
3694 let (mut lg, hj) = self.mtp_step_hl(m, &hx, tok_in, next_pos - 1 + j);
3695 if let Some((lg_cpu, h_cpu)) = dbg_ref {
3696 let n = |v: &[f32]| v.iter().map(|x| x * x).sum::<f32>().sqrt();
3697 let dl = lg.iter().zip(&lg_cpu).fold(0f32, |m, (a, b)| m.max((a - b).abs()));
3698 let dh = hj.iter().zip(&h_cpu).fold(0f32, |m, (a, b)| m.max((a - b).abs()));
3699 eprintln!(
3700 "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 {}",
3701 next_pos - 1 + j,
3702 sampler::argmax(&lg_cpu),
3703 sampler::argmax(&lg),
3704 n(&h_cpu),
3705 n(&hj),
3706 m.kv.seq_len
3707 );
3708 }
3709 let dj = if sparse {
3710 let mut q = std::mem::take(&mut self.spec_qs[j]);
3711 let ok = sampler::sparse_distribution_into(
3712 &lg,
3713 &cfg,
3714 all_ids,
3715 &mut self.sampler_scratch,
3716 self.pool.as_deref(),
3717 &mut q,
3718 );
3719 let d = if ok {
3720 sampler::draw_sparse(&q, &mut self.rng)
3721 } else {
3722 let t = sampler::argmax(&lg);
3724 q.clear();
3725 q.push((t, 1.0));
3726 t
3727 };
3728 self.spec_qs[j] = q;
3729 all_ids.push(d);
3730 d
3731 } else if sampling {
3732 let mut q = std::mem::take(&mut self.spec_q[j]);
3733 sampler::distribution_into(
3734 &lg,
3735 &cfg,
3736 all_ids,
3737 &mut self.sampler_scratch,
3738 self.pool.as_deref(),
3739 &mut q,
3740 );
3741 let d = sampler::draw(&q, &mut self.rng);
3742 self.spec_q[j] = q;
3743 all_ids.push(d); d
3745 } else if greedy_pen {
3746 let d = sampler::argmax_penalized(
3747 &lg,
3748 &cfg,
3749 all_ids,
3750 &mut self.sampler_scratch,
3751 self.pool.as_deref(),
3752 );
3753 all_ids.push(d);
3754 d
3755 } else {
3756 sampler::argmax(&lg)
3757 };
3758 attention::recycle_buf(&mut lg);
3759 drafts.push(dj);
3760 hx = hj;
3761 }
3762 all_ids.truncate(base_len);
3763 *drafted += k_spec;
3764 let t_draft = t_round.elapsed();
3765 let sub_draft = subs();
3766 let b = k_spec + 1;
3769 let mut hiddens = vec![0.0f32; b * self.hidden_size];
3770 for (i, &t) in std::iter::once(&t_next).chain(drafts.iter()).enumerate() {
3771 let e = self.embed_single(t);
3772 hiddens[i * self.hidden_size..(i + 1) * self.hidden_size].copy_from_slice(&e);
3773 }
3774 let positions: Vec<usize> = (next_pos..next_pos + b).collect();
3775 let (lm_gw, lm_rows) = {
3776 let (_, i, kind, rs) = self.weights.lm_head.graph_weight()?;
3777 (
3778 crate::gpu::GraphW {
3779 idx: i,
3780 kind,
3781 row_scale: rs,
3782 data: &[],
3783 },
3784 self.weights.lm_head.rows(),
3785 )
3786 };
3787 let mut logits = Vec::new();
3788 let final_norm = self.weights.final_norm.clone();
3789 #[cfg(target_os = "macos")]
3790 let ok = if metal_native {
3791 let lm = self.weights.lm_head.q1_parts()?;
3792 self.try_batch_graph_metal(&mut hiddens, &positions, b, Some((lm, &final_norm, &mut logits)))
3793 } else {
3794 self.try_batch_graph_wgpu(
3795 &mut hiddens,
3796 &positions,
3797 b,
3798 Some(crate::gpu::SpecTail {
3799 lm: lm_gw,
3800 lm_rows,
3801 final_norm: &final_norm,
3802 logits_out: &mut logits,
3803 }),
3804 )
3805 };
3806 #[cfg(not(target_os = "macos"))]
3807 let ok = self.try_batch_graph_wgpu(
3808 &mut hiddens,
3809 &positions,
3810 b,
3811 Some(crate::gpu::SpecTail {
3812 lm: lm_gw,
3813 lm_rows,
3814 final_norm: &final_norm,
3815 logits_out: &mut logits,
3816 }),
3817 );
3818 if !ok {
3819 m.kv.truncate_last(k_spec);
3822 return None;
3823 }
3824 #[cfg(target_os = "macos")]
3830 if metal_native && std::env::var("CMF_METAL_VERIFY_CHECK").as_deref() == Ok("1") {
3831 let snap: Vec<Vec<f32>> = self.kv_cache.layers.iter().map(|l| l.linear_state.clone()).collect();
3832 let attn_lens: Vec<usize> = self.kv_cache.layers.iter().map(|l| l.seq_len).collect();
3833 let toks: Vec<u32> = std::iter::once(t_next).chain(drafts.iter().copied()).collect();
3834 let want_save = self.graph_want_logits;
3835 self.graph_want_logits = false;
3836 for (i, &t) in toks.iter().enumerate() {
3837 let hi = self.forward_layers(&self.embed_single(t), next_pos + i, None);
3838 let _ = self.graph_logits.take();
3839 if std::env::var("CMF_SPEC_PLAIN_HIDDEN").as_deref() == Ok("1") {
3843 hiddens[i * self.hidden_size..(i + 1) * self.hidden_size].copy_from_slice(&hi);
3844 }
3845 let ref_lg = self.logits_from_hidden(&hi);
3846 let row = &logits[i * lm_rows..(i + 1) * lm_rows];
3847 let ra = sampler::argmax(&ref_lg);
3848 let va = sampler::argmax(row);
3849 let mut md = 0f32;
3850 let mut rms = 0f64;
3851 for j in 0..lm_rows.min(ref_lg.len()) {
3852 let d = (ref_lg[j] - row[j]).abs();
3853 md = md.max(d);
3854 rms += (d as f64) * (d as f64);
3855 }
3856 let mut hd = 0f32;
3857 for j in 0..self.hidden_size {
3858 hd = hd.max((hi[j] - hiddens[i * self.hidden_size + j]).abs());
3859 }
3860 eprintln!(
3861 "verify-check row {i} tok {t} pos {}: ref argmax {ra} verify argmax {va} {} | max|dlogit| {md:.3} rms {:.4} | max|dhidden| {hd:.4}",
3862 next_pos + i,
3863 if ra == va { "OK" } else { "MISMATCH" },
3864 (rms / lm_rows as f64).sqrt()
3865 );
3866 }
3867 self.graph_want_logits = want_save;
3868 for (l, st) in self.kv_cache.layers.iter_mut().zip(snap) {
3871 if l.linear_state.len() == st.len() {
3872 l.linear_state.copy_from_slice(&st);
3873 } else {
3874 l.linear_state = st;
3875 }
3876 }
3877 for (li, (l, n0)) in self.kv_cache.layers.iter_mut().zip(attn_lens).enumerate() {
3878 let extra = l.seq_len.saturating_sub(n0);
3879 if extra > 0 {
3880 l.truncate_last(extra);
3881 crate::gpu_metal::kv_mirror_set_stored(self.graph_kv_id, li, n0);
3882 }
3883 }
3884 }
3885 let t_verify = t_round.elapsed();
3886 let sub_verify = subs();
3887 let mut a = 0usize;
3892 let mut forced: Option<u32> = None;
3893 let ids: Vec<u32> = if sparse {
3894 let mut p = std::mem::take(&mut self.spec_ps);
3895 let mut res = std::mem::take(&mut self.spec_ress);
3896 while a < k_spec {
3897 let ok = sampler::sparse_distribution_into(
3898 &logits[a * lm_rows..(a + 1) * lm_rows],
3899 &cfg,
3900 all_ids,
3901 &mut self.sampler_scratch,
3902 self.pool.as_deref(),
3903 &mut p,
3904 );
3905 if !ok {
3906 let t = sampler::argmax(&logits[a * lm_rows..(a + 1) * lm_rows]);
3907 p.clear();
3908 p.push((t, 1.0));
3909 }
3910 match sampler::spec_accept_or_correct_sparse(
3911 &p,
3912 &self.spec_qs[a],
3913 drafts[a],
3914 &mut self.rng,
3915 &mut res,
3916 ) {
3917 None => {
3918 all_ids.push(drafts[a]);
3919 a += 1;
3920 }
3921 Some(c) => {
3922 forced = Some(c);
3923 break;
3924 }
3925 }
3926 }
3927 all_ids.truncate(base_len);
3928 self.spec_ps = p;
3929 self.spec_ress = res;
3930 drafts.clone()
3931 } else if sampling {
3932 let mut p = std::mem::take(&mut self.spec_p);
3933 let mut res = std::mem::take(&mut self.spec_res);
3934 while a < k_spec {
3935 sampler::distribution_into(
3936 &logits[a * lm_rows..(a + 1) * lm_rows],
3937 &cfg,
3938 all_ids,
3939 &mut self.sampler_scratch,
3940 self.pool.as_deref(),
3941 &mut p,
3942 );
3943 match sampler::spec_accept_or_correct(
3944 &p,
3945 &self.spec_q[a],
3946 drafts[a],
3947 &mut self.rng,
3948 &mut res,
3949 self.pool.as_deref(),
3950 ) {
3951 None => {
3952 all_ids.push(drafts[a]);
3953 a += 1;
3954 }
3955 Some(c) => {
3956 forced = Some(c);
3957 break;
3958 }
3959 }
3960 }
3961 all_ids.truncate(base_len);
3962 self.spec_p = p;
3963 self.spec_res = res;
3964 drafts.clone()
3966 } else if greedy_pen {
3967 let mut ids: Vec<u32> = Vec::with_capacity(b);
3971 for i in 0..b {
3972 let t = sampler::argmax_penalized(
3973 &logits[i * lm_rows..(i + 1) * lm_rows],
3974 &cfg,
3975 all_ids,
3976 &mut self.sampler_scratch,
3977 self.pool.as_deref(),
3978 );
3979 ids.push(t);
3980 if i < k_spec && t == drafts[i] {
3981 all_ids.push(t);
3982 } else {
3983 break;
3984 }
3985 }
3986 all_ids.truncate(base_len);
3987 while a < k_spec && a < ids.len() && ids[a] == drafts[a] {
3988 a += 1;
3989 }
3990 ids
3993 } else {
3994 let ids: Vec<u32> = (0..b)
3995 .map(|i| sampler::argmax(&logits[i * lm_rows..(i + 1) * lm_rows]))
3996 .collect();
3997 while a < k_spec && ids[a] == drafts[a] {
3998 a += 1;
3999 }
4000 ids
4001 };
4002 if spec_dbg {
4003 eprintln!("spec-dbg round: t_next {t_next} drafts {:?} verified {:?} accepted {a}", drafts, ids);
4004 }
4005 #[cfg(target_os = "macos")]
4009 let commit_ref: Option<(Vec<Vec<f32>>, Vec<(usize, Vec<f32>, Vec<f32>)>)> = if metal_native
4010 && std::env::var("CMF_METAL_VERIFY_CHECK").as_deref() == Ok("2")
4011 {
4012 let snap: Vec<Vec<f32>> = self.kv_cache.layers.iter().map(|l| l.linear_state.clone()).collect();
4013 let attn_lens: Vec<usize> = self.kv_cache.layers.iter().map(|l| l.seq_len).collect();
4014 let toks: Vec<u32> = std::iter::once(t_next).chain(drafts.iter().copied()).collect();
4015 let want_save = self.graph_want_logits;
4016 self.graph_want_logits = false;
4017 for (i, &t) in toks.iter().take(a + 1).enumerate() {
4018 let _ = self.forward_layers(&self.embed_single(t), next_pos + i, None);
4019 let _ = self.graph_logits.take();
4020 }
4021 self.graph_want_logits = want_save;
4022 let plain_states: Vec<Vec<f32>> = self.kv_cache.layers.iter().map(|l| l.linear_state.clone()).collect();
4023 let (nkv, hd) = (self.num_kv_heads, self.head_dim);
4024 let mut rows = Vec::new();
4025 for (li, (l, n0)) in self.kv_cache.layers.iter_mut().zip(attn_lens.iter()).enumerate() {
4026 let extra = l.seq_len.saturating_sub(*n0);
4027 if extra > 0 {
4028 let mut kk = Vec::new();
4029 let mut vv = Vec::new();
4030 for g in 0..nkv {
4031 kk.extend_from_slice(&l.head_keys(g)[n0 * hd..]);
4032 vv.extend_from_slice(&l.head_values(g)[n0 * hd..]);
4033 }
4034 rows.push((li, kk, vv));
4035 l.truncate_last(extra);
4036 crate::gpu_metal::kv_mirror_set_stored(self.graph_kv_id, li, *n0);
4037 }
4038 }
4039 for (l, st) in self.kv_cache.layers.iter_mut().zip(snap) {
4040 if l.linear_state.len() == st.len() {
4041 l.linear_state.copy_from_slice(&st);
4042 } else {
4043 l.linear_state = st;
4044 }
4045 }
4046 Some((plain_states, rows))
4047 } else {
4048 None
4049 };
4050 #[cfg(target_os = "macos")]
4052 if metal_native {
4053 self.metal_verify_commit(a);
4056 if let Some((plain_states, rows)) = commit_ref {
4057 crate::gpu_metal::queue_fence();
4058 let (nkv, hd) = (self.num_kv_heads, self.head_dim);
4059 let mut worst_s = 0f32;
4060 let mut worst_li = 0usize;
4061 for (li, (l, ps)) in self.kv_cache.layers.iter().zip(&plain_states).enumerate() {
4062 if l.linear_state.len() != ps.len() || ps.is_empty() {
4063 continue;
4064 }
4065 let d = l.linear_state.iter().zip(ps).fold(0f32, |m, (x, y)| m.max((x - y).abs()));
4066 let n = ps.iter().fold(0f32, |m, y| m.max(y.abs()));
4067 let rel = d / n.max(1e-6);
4068 if rel > worst_s {
4069 worst_s = rel;
4070 worst_li = li;
4071 }
4072 }
4073 let mut worst_k = 0f32;
4074 for (li, kk, vv) in &rows {
4075 let l = &self.kv_cache.layers[*li];
4076 let n0 = l.seq_len - (kk.len() / (nkv * hd));
4077 let mut ck = Vec::new();
4078 let mut cv = Vec::new();
4079 for g in 0..nkv {
4080 ck.extend_from_slice(&l.head_keys(g)[n0 * hd..]);
4081 cv.extend_from_slice(&l.head_values(g)[n0 * hd..]);
4082 }
4083 if ck.len() == kk.len() {
4084 let dk = ck.iter().zip(kk).fold(0f32, |m, (x, y)| m.max((x - y).abs()));
4085 let dv = cv.iter().zip(vv).fold(0f32, |m, (x, y)| m.max((x - y).abs()));
4086 worst_k = worst_k.max(dk).max(dv);
4087 } else {
4088 eprintln!("commit-check L{li}: kv row count mismatch {} vs {}", ck.len(), kk.len());
4089 }
4090 }
4091 eprintln!(
4092 "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}"
4093 );
4094 }
4095 } else if a + 1 < b {
4096 crate::gpu::gdn_spec_restore(self.graph_kv_id, a);
4097 }
4098 #[cfg(not(target_os = "macos"))]
4099 if a + 1 < b {
4100 crate::gpu::gdn_spec_restore(self.graph_kv_id, a);
4101 }
4102 *accepted += a;
4103 m.kv.truncate_last(k_spec.saturating_sub(1));
4114 #[cfg(target_os = "macos")]
4115 if metal_native && self.mtp_graph_mode == Some(true) {
4116 crate::gpu_metal::kv_mirror_set_stored(self.mtp_kv_id(), Self::MTP_LAYER_BASE, m.kv.seq_len);
4119 }
4120 let warm_off = std::env::var("CMF_SPEC_WARM").is_ok_and(|v| v == "0");
4121 if !warm_off && a > 0 {
4122 let mut warmed = false;
4125 #[cfg(target_os = "macos")]
4126 if metal_native && self.mtp_graph_mode == Some(true) {
4127 let pairs: Vec<(&[f32], u32)> = (0..a)
4131 .map(|j| (&hiddens[j * self.hidden_size..(j + 1) * self.hidden_size], ids[j]))
4132 .collect();
4133 warmed = self.mtp_warm_batch_metal(m, &pairs, next_pos);
4134 if !warmed {
4135 warmed = true;
4136 for j in 0..a {
4137 let row = hiddens[j * self.hidden_size..(j + 1) * self.hidden_size].to_vec();
4138 if self.mtp_step_metal(m, &row, ids[j], next_pos + j, false).is_none() {
4139 warmed = false;
4140 break;
4141 }
4142 }
4143 }
4144 }
4145 if !warmed && self.mtp_graph_mode == Some(true) && !metal_native {
4146 let rows: Vec<Vec<f32>> = (0..a)
4147 .map(|j| hiddens[j * self.hidden_size..(j + 1) * self.hidden_size].to_vec())
4148 .collect();
4149 let pairs: Vec<(&[f32], u32)> = rows
4150 .iter()
4151 .zip(ids.iter())
4152 .map(|(r, &t)| (r.as_slice(), t))
4153 .collect();
4154 warmed = self.mtp_warm_graph(m, &pairs, next_pos);
4155 if !warmed {
4156 warmed = true;
4158 for j in 0..a {
4159 if self
4160 .mtp_step_graph(m, &rows[j], ids[j], next_pos + j)
4161 .is_none()
4162 {
4163 warmed = false;
4164 break;
4165 }
4166 }
4167 }
4168 }
4169 if !warmed {
4170 for j in 0..a {
4171 let row = &hiddens[j * self.hidden_size..(j + 1) * self.hidden_size];
4172 let row = row.to_vec();
4173 self.mtp_warm(m, &row, ids[j], next_pos + j);
4174 }
4175 }
4176 }
4177 if let Some(c) = forced {
4181 self.spec_forced = Some(c);
4182 self.graph_logits = None;
4183 } else {
4184 let mut row = logits[a * lm_rows..(a + 1) * lm_rows].to_vec();
4185 row.resize(self.vocab_size, 0.0);
4186 if let Some(c) = self.final_softcap {
4187 for l in row.iter_mut() {
4188 *l = c * (*l / c).tanh();
4189 }
4190 }
4191 self.graph_logits = Some(row);
4192 }
4193 let new_hidden = hiddens[a * self.hidden_size..(a + 1) * self.hidden_size].to_vec();
4194 if std::env::var("CMF_GRAPH_SPEC_TIME").is_ok() {
4200 let end = subs();
4201 eprintln!(
4202 "spec-round: draft {:.1} ms/{} sub | verify {:.1} ms/{} sub | \
4203 commit {:.1} ms/{} sub (accepted {a} of {k_spec})",
4204 t_draft.as_secs_f64() * 1e3,
4205 sub_draft - sub0,
4206 (t_verify - t_draft).as_secs_f64() * 1e3,
4207 sub_verify - sub_draft,
4208 (t_round.elapsed() - t_verify).as_secs_f64() * 1e3,
4209 end - sub_verify,
4210 );
4211 }
4212 Some((drafts[..a].to_vec(), next_pos + a + 1, new_hidden))
4213 }
4214
4215 pub fn measure_pair_fusion(&mut self, iters: usize) -> (f64, f64) {
4224 if !self.pair_supported() {
4225 return (0.0, 0.0);
4226 }
4227 let emb1 = self.embed_single(1);
4228 let emb2 = self.embed_single(2);
4229 let pos = self.kv_cache.seq_len();
4230
4231 let t0 = std::time::Instant::now();
4232 for _ in 0..iters {
4233 let _ = self.forward_layers(&emb1, pos, None);
4234 let _ = self.forward_layers(&emb2, pos + 1, None);
4235 for l in &mut self.kv_cache.layers {
4236 l.truncate_last(2);
4237 }
4238 }
4239 let singles_ms = t0.elapsed().as_secs_f64() * 1000.0 / iters as f64;
4240
4241 let t1 = std::time::Instant::now();
4242 for _ in 0..iters {
4243 let _ = self.forward_pair(&emb1, &emb2, pos);
4244 for l in &mut self.kv_cache.layers {
4245 l.truncate_last(2);
4246 }
4247 }
4248 let pair_ms = t1.elapsed().as_secs_f64() * 1000.0 / iters as f64;
4249 (singles_ms, pair_ms)
4250 }
4251
4252 fn pair_supported(&self) -> bool {
4260 !self.weights.layers.is_empty()
4267 && self.g3n.is_none()
4268 && !self
4269 .weights
4270 .layers
4271 .iter()
4272 .any(|lw| matches!(&lw.attn, AttnKind::Mla(_) | AttnKind::Kda(_)))
4273 }
4274
4275 fn forward_pair(
4276 &mut self,
4277 emb1: &[f32],
4278 emb2: &[f32],
4279 position: usize,
4280 ) -> (Vec<f32>, Vec<f32>) {
4281 let mut h1 = emb1.to_vec();
4282 let mut h2 = emb2.to_vec();
4283 let (_nkv, _hd, hs, _rd, eps) = (
4284 self.num_kv_heads,
4285 self.head_dim,
4286 self.hidden_size,
4287 self.rotary_dim,
4288 self.rms_eps,
4289 );
4290 let pool = self.pool.clone();
4291
4292 for li in 0..self.num_layers {
4293 let lw = &self.weights.layers[self.phys_layer(li)];
4294 inference::rms_norm_into(
4297 &h1,
4298 &lw.input_norm,
4299 self.rms_eps,
4300 self.norm_style,
4301 &mut self.ws.n1,
4302 );
4303 inference::rms_norm_into(
4304 &h2,
4305 &lw.input_norm,
4306 self.rms_eps,
4307 self.norm_style,
4308 &mut self.ws.n2,
4309 );
4310
4311 let (a1, a2) = match &lw.attn {
4312 AttnKind::Mla(_) => unreachable!("MLA has no MTP/pair path"),
4313 AttnKind::Kda(_) => unreachable!("KDA has no MTP/pair path"),
4314 AttnKind::Linear(w) => {
4315 let cfg = self.vmf_cfg.expect("linear layer without vmf_cfg");
4316 let layer = &mut self.kv_cache.layers[li];
4317 let (state, scratch) = (&mut layer.linear_state, &mut layer.linear_scratch);
4318 vmf_phase_pair(
4319 &self.ws.n1,
4320 &self.ws.n2,
4321 w,
4322 &cfg,
4323 state,
4324 scratch,
4325 self.pool.as_deref(),
4326 )
4327 }
4328 AttnKind::LinearGdn(w) => {
4329 let cfg = self.gdn_cfg.expect("gdn layer without gdn_cfg");
4330 let layer = &mut self.kv_cache.layers[li];
4331 let (state, scratch) = (&mut layer.linear_state, &mut layer.linear_scratch);
4332 gdn_pair(
4333 &self.ws.n1,
4334 &self.ws.n2,
4335 w,
4336 &cfg,
4337 state,
4338 scratch,
4339 self.pool.as_deref(),
4340 )
4341 }
4342 AttnKind::ShortConv(w) => {
4343 let cfg = self
4344 .short_conv_cfg
4345 .expect("short-conv layer without short_conv_cfg");
4346 let layer = &mut self.kv_cache.layers[li];
4347 let (state, scratch) = (&mut layer.linear_state, &mut layer.linear_scratch);
4348 short_conv_pair(
4349 &self.ws.n1,
4350 &self.ws.n2,
4351 w,
4352 &cfg,
4353 state,
4354 scratch,
4355 self.pool.as_deref(),
4356 )
4357 }
4358 AttnKind::Full {
4359 wq,
4360 wk,
4361 wv,
4362 wo,
4363 q_norm,
4364 k_norm,
4365 output_gate,
4366 softplus_gate,
4367 bias,
4368 } => {
4369 let inv_freq_l = self.layer_inv_freq(li);
4370 let (nkv_l, hd_l, rd_l) = self.layer_geom(li);
4371 let cfg = QwenAttnCfg {
4372 num_heads: self.layer_num_heads(li),
4373 num_kv_heads: nkv_l,
4374 head_dim: hd_l,
4375 hidden_size: hs,
4376 position,
4377 inv_freq: &inv_freq_l,
4378 rotary_dim: rd_l,
4379 scale: self.attn_scale,
4380 softcap: self.attn_softcap,
4381 window: self.layer_window(li),
4382 v_norm: self.attn_v_norm,
4383 q_norm: q_norm.as_deref(),
4384 k_norm: k_norm.as_deref(),
4385 output_gate: *output_gate,
4386 softplus_gate: softplus_gate
4387 .as_ref()
4388 .map(|(gate, per_head)| (gate, *per_head)),
4389 rope_scale: self.layer_rope_scale(li),
4390 bias: bias
4391 .as_ref()
4392 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
4393 rms_eps: eps,
4394 norm_style: self.norm_style,
4395 pool: pool.as_deref(),
4396 };
4397 attention::qwen_attention_pair(
4398 &self.ws.n1,
4399 &self.ws.n2,
4400 wq,
4401 wk,
4402 wv,
4403 wo,
4404 &mut self.kv_cache.layers[li],
4405 &cfg,
4406 )
4407 }
4408 };
4409 let (a1, a2) = match &self.weights.layers[self.phys_layer(li)].attn_out_norm {
4410 Some(w) => (
4411 inference::rms_norm(&a1, w, self.rms_eps, self.norm_style),
4412 inference::rms_norm(&a2, w, self.rms_eps, self.norm_style),
4413 ),
4414 None => (a1, a2),
4415 };
4416 for i in 0..self.hidden_size {
4417 h1[i] += a1[i];
4418 h2[i] += a2[i];
4419 }
4420 let (mut a1, mut a2) = (a1, a2);
4421 attention::recycle_buf(&mut a1);
4422 attention::recycle_buf(&mut a2);
4423
4424 let lw = &self.weights.layers[self.phys_layer(li)];
4425 inference::rms_norm_into(
4426 &h1,
4427 &lw.post_norm,
4428 self.rms_eps,
4429 self.norm_style,
4430 &mut self.ws.p1,
4431 );
4432 inference::rms_norm_into(
4433 &h2,
4434 &lw.post_norm,
4435 self.rms_eps,
4436 self.norm_style,
4437 &mut self.ws.p2,
4438 );
4439 let (f1, f2) = match &lw.ffn {
4440 FfnKind::DenseMoe(dm) => (
4443 dense_moe_ffn(
4444 dm,
4445 &self.ws.p1,
4446 &h1,
4447 self.rms_eps,
4448 self.norm_style,
4449 self.pool.as_deref(),
4450 ),
4451 dense_moe_ffn(
4452 dm,
4453 &self.ws.p2,
4454 &h2,
4455 self.rms_eps,
4456 self.norm_style,
4457 self.pool.as_deref(),
4458 ),
4459 ),
4460 _ => ffn_forward_pair(
4461 &lw.ffn,
4462 &self.ws.p1,
4463 &self.ws.p2,
4464 self.pool.as_deref(),
4465 None,
4466 ),
4467 };
4468 let (f1, f2) = match &self.weights.layers[self.phys_layer(li)].ffn_out_norm {
4469 Some(w) => (
4470 inference::rms_norm(&f1, w, self.rms_eps, self.norm_style),
4471 inference::rms_norm(&f2, w, self.rms_eps, self.norm_style),
4472 ),
4473 None => (f1, f2),
4474 };
4475 for i in 0..self.hidden_size {
4476 h1[i] += f1[i];
4477 h2[i] += f2[i];
4478 }
4479 let (mut f1, mut f2) = (f1, f2);
4480 attention::recycle_buf(&mut f1);
4481 attention::recycle_buf(&mut f2);
4482 if let Some(sc) = self.weights.layers[self.phys_layer(li)].layer_scale {
4483 for i in 0..self.hidden_size {
4484 h1[i] *= sc;
4485 h2[i] *= sc;
4486 }
4487 }
4488 if self.is_loop_end(li) && li + 1 < self.num_layers {
4490 h1 = inference::rms_norm(
4491 &h1,
4492 &self.weights.final_norm,
4493 self.rms_eps,
4494 self.norm_style,
4495 );
4496 h2 = inference::rms_norm(
4497 &h2,
4498 &self.weights.final_norm,
4499 self.rms_eps,
4500 self.norm_style,
4501 );
4502 }
4503 }
4504 (h1, h2)
4505 }
4506
4507 fn commit_linear_scratch(&mut self) {
4509 for layer in &mut self.kv_cache.layers {
4510 if !layer.linear_scratch.is_empty() {
4511 std::mem::swap(&mut layer.linear_state, &mut layer.linear_scratch);
4512 layer.linear_scratch.clear();
4513 }
4514 }
4515 }
4516
4517 pub fn forward_ids(
4520 &mut self,
4521 ids: &[u32],
4522 task_mask: Option<&TaskMask>,
4523 ) -> Result<Vec<f32>, String> {
4524 if ids.is_empty() {
4525 return Err("empty id sequence".to_string());
4526 }
4527 self.kv_cache.clear();
4528 self.kv_history.clear();
4529 self.o1_begin();
4530 let mut hidden = vec![0.0f32; self.hidden_size];
4531 let mut pos = 0usize;
4532 if self.can_prefill_batched() && !self.graph_prefill_preferred() && ids.len() > 2 {
4540 let chunk = prefill_chunk();
4544 let hs = self.hidden_size;
4545 while pos < ids.len() {
4546 let end = (pos + chunk).min(ids.len());
4547 let hb = self.prefill_batch_masked(&ids[pos..end], pos, task_mask);
4548 hidden.copy_from_slice(&hb[(end - pos - 1) * hs..]);
4549 pos = end;
4550 }
4551 }
4552 if task_mask.is_none()
4561 && !self.graph_prefill_preferred()
4562 && !std::env::var("CMF_PAIR").is_ok_and(|v| v == "0")
4563 && self.pair_supported()
4564 {
4565 while pos + 1 < ids.len() {
4566 let e1 = self.embed_single(ids[pos]);
4567 let e2 = self.embed_single(ids[pos + 1]);
4568 let (_, h2) = self.forward_pair(&e1, &e2, pos);
4569 self.commit_linear_scratch();
4570 hidden = h2;
4571 pos += 2;
4572 }
4573 }
4574 while pos < ids.len() {
4575 hidden = self.forward_layers(&self.embed_single(ids[pos]), pos, task_mask);
4576 pos += 1;
4577 }
4578 self.o1_seal();
4582 let normed = inference::rms_norm(
4583 &hidden,
4584 &self.weights.final_norm,
4585 self.rms_eps,
4586 self.norm_style,
4587 );
4588 Ok(self.lm_head_forward(&normed))
4589 }
4590
4591 pub fn ppl_ids(&mut self, ids: &[u32]) -> f64 {
4598 let (nll, cnt) = self.nll_ids_from(ids, 0);
4599 (nll / cnt.max(1) as f64).exp()
4600 }
4601
4602 pub fn probe_ffn_mass(&mut self, ids: &[u32]) -> Vec<Vec<f64>> {
4607 self.kv_cache.clear();
4608 self.kv_history.clear();
4609 FFN_PROBE.with(|p| {
4610 *p.borrow_mut() = Some(vec![vec![0f64; self.intermediate_size]; self.num_layers]);
4611 });
4612 crate::gpu::cpu_scope(|| {
4613 for (pos, &id) in ids.iter().enumerate() {
4614 let emb = self.embed_single(id);
4615 let _ = self.forward_layers(&emb, pos, None);
4616 }
4617 });
4618 self.kv_cache.clear();
4619 self.kv_history.clear();
4620 FFN_PROBE
4621 .with(|p| p.borrow_mut().take())
4622 .unwrap_or_default()
4623 }
4624
4625 pub fn probe_ffn_mass_batch(&mut self, ids: &[u32]) -> Vec<Vec<f64>> {
4629 self.kv_cache.clear();
4630 self.kv_history.clear();
4631 FFN_PROBE.with(|p| {
4632 *p.borrow_mut() = Some(vec![vec![0f64; self.intermediate_size]; self.num_layers]);
4633 });
4634 for chunk in ids.chunks(256) {
4635 if chunk.len() < 2 {
4636 continue;
4637 }
4638 let _ = self.nll_ids_masked(chunk, 0, None);
4639 }
4640 self.kv_cache.clear();
4641 self.kv_history.clear();
4642 FFN_PROBE
4643 .with(|p| p.borrow_mut().take())
4644 .unwrap_or_default()
4645 }
4646
4647 pub fn ppl_ids_masked(&mut self, ids: &[u32], mask: &TaskMask) -> f64 {
4651 self.kv_cache.clear();
4652 self.kv_history.clear();
4653 let mut nll = 0f64;
4654 let mut cnt = 0usize;
4655 let mut hidden = vec![0f32; self.hidden_size];
4656 for (pos, &id) in ids.iter().enumerate() {
4657 if pos > 0 {
4658 inference::rms_norm_into(
4659 &hidden,
4660 &self.weights.final_norm,
4661 self.rms_eps,
4662 self.norm_style,
4663 &mut self.ws.n1,
4664 );
4665 let mut logits = self.lm_head_forward(&self.ws.n1);
4666 let max = logits.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
4667 let sum: f64 = logits.iter().map(|&v| ((v - max) as f64).exp()).sum();
4668 let p = ((logits[id as usize] - max) as f64).exp() / sum.max(1e-300);
4669 nll -= p.max(1e-300).ln();
4670 cnt += 1;
4671 attention::recycle_buf(&mut logits);
4672 }
4673 let emb = self.embed_single(id);
4674 hidden = self.forward_layers(&emb, pos, Some(mask));
4675 }
4676 self.kv_cache.clear();
4677 self.kv_history.clear();
4678 (nll / cnt.max(1) as f64).exp()
4679 }
4680
4681 pub fn nll_ids_masked(
4700 &mut self,
4701 ids: &[u32],
4702 start: usize,
4703 task_mask: Option<&TaskMask>,
4704 ) -> (f64, usize) {
4705 let task_mask = self.drop_open_mask(task_mask);
4706 self.nll_ids_inner(ids, start, task_mask)
4707 }
4708
4709 pub fn nll_ids_from(&mut self, ids: &[u32], start: usize) -> (f64, usize) {
4710 self.nll_ids_inner(ids, start, None)
4711 }
4712
4713 fn nll_ids_inner(
4714 &mut self,
4715 ids: &[u32],
4716 start: usize,
4717 task_mask: Option<&TaskMask>,
4718 ) -> (f64, usize) {
4719 self.kv_cache.clear();
4720 self.kv_history.clear();
4721 let mut nll = 0f64;
4722 let mut cnt = 0usize;
4723 if self.can_prefill_batched() {
4724 const CHUNK: usize = 128;
4730 const LM_SUB: usize = 32;
4731 let n = ids.len().saturating_sub(1);
4732 let hs = self.hidden_size;
4733 let rows = self.weights.lm_head.rows();
4734 let mut pos = 0usize;
4735 while pos < n {
4736 let end = (pos + CHUNK).min(n);
4737 let bsz = end - pos;
4738 let hb = self.prefill_batch_masked(&ids[pos..end], pos, task_mask);
4739 let mut k0 = 0usize;
4740 while k0 < bsz {
4741 let k1 = (k0 + LM_SUB).min(bsz);
4742 let sb = k1 - k0;
4743 if pos + k1 <= start {
4746 k0 = k1;
4747 continue;
4748 }
4749 let mut normed = vec![0.0f32; sb * hs];
4750 for k in 0..sb {
4751 let r = inference::rms_norm(
4752 &hb[(k0 + k) * hs..(k0 + k + 1) * hs],
4753 &self.weights.final_norm,
4754 self.rms_eps,
4755 self.norm_style,
4756 );
4757 normed[k * hs..(k + 1) * hs].copy_from_slice(&r);
4758 }
4759 let mut logits = vec![0.0f32; sb * rows];
4760 self.weights
4761 .lm_head
4762 .matmat(&normed, sb, &mut logits, self.pool.as_deref());
4763 for k in 0..sb {
4764 if pos + k0 + k < start {
4765 continue;
4766 }
4767 let lg = &mut logits[k * rows..k * rows + self.vocab_size.min(rows)];
4768 if let Some(mu) = self.logit_multiplier {
4769 for v in lg.iter_mut() {
4770 *v *= mu;
4771 }
4772 }
4773 if let Some(c) = self.final_softcap {
4777 for v in lg.iter_mut() {
4778 *v = c * (*v / c).tanh();
4779 }
4780 }
4781 if let Some(cm) = self.head_clusters.clone() {
4784 self.hierarchical_head_logprobs(&normed[k * hs..(k + 1) * hs], &cm, lg);
4785 }
4786 let lg = &logits[k * rows..k * rows + self.vocab_size.min(rows)];
4787 let target = ids[pos + k0 + k + 1] as usize;
4788 let max = lg.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
4789 let lse: f64 = lg
4790 .iter()
4791 .map(|&v| ((v - max) as f64).exp())
4792 .sum::<f64>()
4793 .ln()
4794 + max as f64;
4795 nll += lse - lg[target] as f64;
4796 cnt += 1;
4797 if std::env::var("CMF_PPL_TRACE").is_ok() {
4798 let top = lg
4799 .iter()
4800 .enumerate()
4801 .max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
4802 .map(|(i, _)| i)
4803 .unwrap_or(0);
4804 eprintln!(
4805 "BTRACE pos {} target {} nll {:.4} top {} lg_t {:.3} lg_top {:.3}",
4806 pos + k0 + k,
4807 target,
4808 lse - lg[target] as f64,
4809 top,
4810 lg[target],
4811 lg[top]
4812 );
4813 }
4814 }
4815 k0 = k1;
4816 }
4817 pos = end;
4818 }
4819 self.kv_cache.clear();
4820 self.kv_history.clear();
4821 return (nll, cnt);
4822 }
4823 for pos in 0..ids.len().saturating_sub(1) {
4824 let hidden = self.forward_layers(&self.embed_single(ids[pos]), pos, task_mask);
4825 let out_of_band = self.graph_logits.take();
4833 if pos < start {
4834 continue;
4835 }
4836 let logits = match out_of_band {
4837 Some(lg) => lg,
4838 None => {
4839 let normed = inference::rms_norm(
4840 &hidden,
4841 &self.weights.final_norm,
4842 self.rms_eps,
4843 self.norm_style,
4844 );
4845 self.lm_head_forward(&normed)
4849 }
4850 };
4851 let target = ids[pos + 1] as usize;
4852 let max = logits.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
4853 let lse: f64 = logits
4854 .iter()
4855 .map(|&v| ((v - max) as f64).exp())
4856 .sum::<f64>()
4857 .ln()
4858 + max as f64;
4859 let tok_nll = lse - logits[target] as f64;
4860 if std::env::var("CMF_PPL_TRACE").is_ok() && pos < 48 {
4861 let top = logits
4862 .iter()
4863 .enumerate()
4864 .max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
4865 .map(|(i, _)| i)
4866 .unwrap_or(0);
4867 eprintln!(
4868 "pos {pos:3} tgt {target:6} nll {tok_nll:7.3} | top1 {top:6} lg[t]={:.2} lg[top]={:.2}",
4869 logits[target], logits[top]
4870 );
4871 }
4872 nll += tok_nll;
4873 cnt += 1;
4874 }
4875 self.kv_cache.clear();
4876 self.kv_history.clear();
4877 (nll, cnt)
4878 }
4879
4880 pub fn nll_ids_o1(&mut self, ids: &[u32], prefill: usize) -> (f64, usize) {
4896 self.kv_cache.clear();
4897 self.kv_history.clear();
4898 self.o1_begin();
4899 let n = ids.len().saturating_sub(1);
4900 let p = prefill.min(n);
4901 let mut pos = 0usize;
4903 if self.can_prefill_batched() {
4904 const CHUNK: usize = 128;
4905 while pos < p {
4906 let end = (pos + CHUNK).min(p);
4907 let _ = self.prefill_batch(&ids[pos..end], pos);
4908 pos = end;
4909 }
4910 } else {
4911 while pos < p {
4912 let _ = self.forward_layers(&self.embed_single(ids[pos]), pos, None);
4913 pos += 1;
4914 }
4915 }
4916 self.o1_seal();
4917
4918 let mut nll = 0f64;
4919 let mut cnt = 0usize;
4920 for pos in p..n {
4921 let hidden = self.forward_layers(&self.embed_single(ids[pos]), pos, None);
4922 let normed = inference::rms_norm(
4923 &hidden,
4924 &self.weights.final_norm,
4925 self.rms_eps,
4926 self.norm_style,
4927 );
4928 let logits = self.lm_head_forward(&normed);
4932 let target = ids[pos + 1] as usize;
4933 let max = logits.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
4934 let lse: f64 = logits
4935 .iter()
4936 .map(|&v| ((v - max) as f64).exp())
4937 .sum::<f64>()
4938 .ln()
4939 + max as f64;
4940 let tok_nll = lse - logits[target] as f64;
4941 if std::env::var("CMF_PPL_TRACE").is_ok() && pos < 48 {
4942 let top = logits
4943 .iter()
4944 .enumerate()
4945 .max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
4946 .map(|(i, _)| i)
4947 .unwrap_or(0);
4948 eprintln!(
4949 "pos {pos:3} tgt {target:6} nll {tok_nll:7.3} | top1 {top:6} lg[t]={:.2} lg[top]={:.2}",
4950 logits[target], logits[top]
4951 );
4952 }
4953 nll += tok_nll;
4954 cnt += 1;
4955 }
4956 self.kv_cache.clear();
4957 self.kv_history.clear();
4958 (nll, cnt)
4959 }
4960
4961 pub fn calib_ids(&mut self, ids: &[u32], temps: &[f32]) -> (Vec<bool>, Vec<Vec<f32>>) {
4969 self.kv_cache.clear();
4970 self.kv_history.clear();
4971 let n = ids.len().saturating_sub(1);
4972 let mut correct = Vec::with_capacity(n);
4973 let mut pmax = Vec::with_capacity(n);
4974 for pos in 0..n {
4975 let emb = self.embed_single(ids[pos]);
4976 let hidden = self.forward_layers(&emb, pos, None);
4977 let normed = inference::rms_norm(
4978 &hidden,
4979 &self.weights.final_norm,
4980 self.rms_eps,
4981 self.norm_style,
4982 );
4983 let logits = self.lm_head_forward(&normed);
4987 let target = ids[pos + 1] as usize;
4988 let (mut amax, mut mval) = (0usize, f32::NEG_INFINITY);
4989 for (i, &v) in logits.iter().enumerate() {
4990 if v > mval {
4991 mval = v;
4992 amax = i;
4993 }
4994 }
4995 correct.push(amax == target);
4996 let row: Vec<f32> = temps
4997 .iter()
4998 .map(|&t| {
4999 let tt = t.max(1e-3);
5000 let s: f32 = logits.iter().map(|&v| ((v - mval) / tt).exp()).sum();
5001 1.0 / s.max(1e-12) })
5003 .collect();
5004 pmax.push(row);
5005 }
5006 self.kv_cache.clear();
5007 self.kv_history.clear();
5008 (correct, pmax)
5009 }
5010
5011 pub fn ppl_ids_dynamic(&mut self, ids: &[u32]) -> (f64, usize) {
5018 let mut router = match self.dyn_router.take() {
5019 Some(r) => r,
5020 None => return (self.ppl_ids(ids), 0),
5021 };
5022 router.reset();
5023 self.dyn_phi_seen = 0;
5024 let _ = self.set_active_skill(None);
5025
5026 self.kv_cache.clear();
5027
5028 self.kv_history.clear();
5029 let mut nll = 0f64;
5030 let mut cnt = 0usize;
5031 for pos in 0..ids.len().saturating_sub(1) {
5032 let hidden = self.forward_layers(&self.embed_single(ids[pos]), pos, None);
5033 let normed = inference::rms_norm(
5034 &hidden,
5035 &self.weights.final_norm,
5036 self.rms_eps,
5037 self.norm_style,
5038 );
5039 let logits = self.lm_head_forward(&normed);
5043 let target = ids[pos + 1] as usize;
5044 let max = logits.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
5045 let lse: f64 = logits
5046 .iter()
5047 .map(|&v| ((v - max) as f64).exp())
5048 .sum::<f64>()
5049 .ln()
5050 + max as f64;
5051 let tok_nll = lse - logits[target] as f64;
5052 if std::env::var("CMF_PPL_TRACE").is_ok() && pos < 48 {
5053 let top = logits
5054 .iter()
5055 .enumerate()
5056 .max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
5057 .map(|(i, _)| i)
5058 .unwrap_or(0);
5059 eprintln!(
5060 "pos {pos:3} tgt {target:6} nll {tok_nll:7.3} | top1 {top:6} lg[t]={:.2} lg[top]={:.2}",
5061 logits[target], logits[top]
5062 );
5063 }
5064 nll += tok_nll;
5065 cnt += 1;
5066 let phi = self.dyn_phi_ema.clone();
5068 if let Some(new_active) = router.step(&phi, pos) {
5069 let _ = self.set_active_skill(new_active);
5070 }
5071 }
5072 let switches = router.switches.len();
5073 let _ = self.set_active_skill(None);
5074 self.dyn_router = Some(router);
5075 self.kv_cache.clear();
5076 self.kv_history.clear();
5077 ((nll / cnt.max(1) as f64).exp(), switches)
5078 }
5079
5080 pub fn probe_phi(&mut self, ids: &[u32], layer: usize) -> Vec<f32> {
5082 self.kv_cache.clear();
5083 self.kv_history.clear();
5084 let mut acc = vec![0f32; self.hidden_size];
5085 for (pos, &id) in ids.iter().enumerate() {
5086 let h = self.forward_layers_upto(&self.embed_single(id), pos, None, Some(layer));
5087 for (a, v) in acc.iter_mut().zip(&h) {
5088 *a += v;
5089 }
5090 }
5091 let n = ids.len().max(1) as f32;
5092 for a in acc.iter_mut() {
5093 *a /= n;
5094 }
5095 self.kv_cache.clear();
5096 self.kv_history.clear();
5097 acc
5098 }
5099
5100 fn prefill_batch(&mut self, ids: &[u32], start_pos: usize) -> Vec<f32> {
5106 self.prefill_batch_masked(ids, start_pos, None)
5107 }
5108
5109 fn prefill_batch_masked(
5115 &mut self,
5116 ids: &[u32],
5117 start_pos: usize,
5118 task_mask: Option<&TaskMask>,
5119 ) -> Vec<f32> {
5120 self.prefill_batch_span(PrefillIn::Ids(ids), start_pos, task_mask, 0, usize::MAX)
5121 }
5122
5123 fn prefill_batch_span(
5129 &mut self,
5130 input: PrefillIn<'_>,
5131 start_pos: usize,
5132 task_mask: Option<&TaskMask>,
5133 from: usize,
5134 upto_excl: usize,
5135 ) -> Vec<f32> {
5136 let hs = self.hidden_size;
5137 let b = match input {
5138 PrefillIn::Ids(ids) => ids.len(),
5139 PrefillIn::Hidden(hb) => hb.len() / hs,
5140 };
5141 let upto_excl = upto_excl.min(self.num_layers);
5142 let mut h: Vec<f32>;
5146 let mut h_ready;
5147 match input {
5148 PrefillIn::Ids(_) => {
5149 h = vec![0.0; b * hs];
5150 h_ready = false;
5151 }
5152 PrefillIn::Hidden(hb) => {
5153 h = hb.to_vec();
5154 h_ready = true;
5155 }
5156 }
5157 let fill_h = |h: &mut Vec<f32>, me: &Self| {
5158 if let PrefillIn::Ids(ids) = input {
5159 for (bi, &id) in ids.iter().enumerate() {
5160 let e = me.embed_single(id);
5161 h[bi * hs..(bi + 1) * hs].copy_from_slice(&e);
5162 }
5163 if let Ok(tp) = std::env::var("CMF_TRACE_POS") {
5164 if let Ok(t) = tp.parse::<usize>() {
5165 if t >= start_pos && t < start_pos + ids.len() {
5166 let bi = t - start_pos;
5167 let row = &h[bi * hs..(bi + 1) * hs];
5168 let n: f32 = row.iter().map(|x| x * x).sum::<f32>().sqrt();
5169 eprintln!(
5170 "BATCH pos {t} embed: id {} |h| = {n:.6} h0 {:.6} h1 {:.6} | b={} start={start_pos} ids[..8]={:?}",
5171 ids[bi], row[0], row[1], ids.len(), &ids[..ids.len().min(8)]
5172 );
5173 }
5174 }
5175 }
5176 }
5177 };
5178 let (_nkv, _hd, _rd, eps) = (
5179 self.num_kv_heads,
5180 self.head_dim,
5181 self.rotary_dim,
5182 self.rms_eps,
5183 );
5184 let pool = self.pool.clone();
5185 let norm_style = self.norm_style;
5186
5187 #[cfg(target_os = "macos")]
5188 let mut chunk_skip_until = 0usize;
5189 for li in from..upto_excl {
5190 crate::gpu::set_layer(li as i64); #[cfg(target_os = "macos")]
5197 if task_mask.is_none() {
5198 if li < chunk_skip_until {
5199 continue;
5200 }
5201 if !h_ready && li == 0 && self.weights.embed_tokens.q8_row_parts().is_none() {
5207 fill_h(&mut h, self);
5208 h_ready = true;
5209 }
5210 let ids_for_embed = match input {
5211 PrefillIn::Ids(ids) => (!h_ready && li == 0).then_some(ids),
5212 PrefillIn::Hidden(_) => None,
5213 };
5214 let end = self.chunk_run_gpu(li, &mut h, b, start_pos, ids_for_embed, upto_excl);
5215 if end > li {
5216 h_ready = true;
5217 chunk_skip_until = end;
5218 if self.is_loop_end(end - 1) && end < self.num_layers {
5221 for bi in 0..b {
5222 let normed = inference::rms_norm(
5223 &h[bi * hs..(bi + 1) * hs],
5224 &self.weights.final_norm,
5225 eps,
5226 norm_style,
5227 );
5228 h[bi * hs..(bi + 1) * hs].copy_from_slice(&normed);
5229 }
5230 }
5231 continue;
5232 }
5233 }
5234 if !h_ready {
5235 fill_h(&mut h, self);
5236 h_ready = true;
5237 }
5238 let lw = &self.weights.layers[self.phys_layer(li)];
5239 match &lw.attn {
5241 AttnKind::Kda(w) => {
5242 let cfg = self.kda_cfg.expect("kda layer without kda_cfg");
5244 let mut normed = vec![0.0f32; b * hs];
5245 for bi in 0..b {
5246 inference::rms_norm_into(
5247 &h[bi * hs..(bi + 1) * hs],
5248 &lw.input_norm,
5249 eps,
5250 norm_style,
5251 &mut normed[bi * hs..(bi + 1) * hs],
5252 );
5253 }
5254 let attn = crate::linear_core::kda_forward_batch(
5255 &normed,
5256 b,
5257 w,
5258 &cfg,
5259 &mut self.kv_cache.layers[li].linear_state,
5260 pool.as_deref(),
5261 );
5262 for (dst, &a) in h.iter_mut().zip(&attn) {
5263 *dst += a;
5264 }
5265 }
5266 AttnKind::LinearGdn(w) => {
5267 let cfg = self.gdn_cfg.expect("gdn layer without gdn_cfg");
5269 let mut normed = vec![0.0f32; b * hs];
5270 for bi in 0..b {
5271 let r = inference::rms_norm(
5272 &h[bi * hs..(bi + 1) * hs],
5273 &lw.input_norm,
5274 eps,
5275 norm_style,
5276 );
5277 normed[bi * hs..(bi + 1) * hs].copy_from_slice(&r);
5278 }
5279 let attn = crate::linear_core::gdn_forward_batch(
5280 &normed,
5281 b,
5282 w,
5283 &cfg,
5284 &mut self.kv_cache.layers[li].linear_state,
5285 pool.as_deref(),
5286 );
5287 for (dst, &a) in h.iter_mut().zip(&attn) {
5288 *dst += a;
5289 }
5290 }
5291 AttnKind::ShortConv(w) => {
5292 let cfg = self
5295 .short_conv_cfg
5296 .expect("short-conv layer without short_conv_cfg");
5297 let mut normed = vec![0.0f32; b * hs];
5298 for bi in 0..b {
5299 inference::rms_norm_into(
5300 &h[bi * hs..(bi + 1) * hs],
5301 &lw.input_norm,
5302 eps,
5303 norm_style,
5304 &mut normed[bi * hs..(bi + 1) * hs],
5305 );
5306 }
5307 let attn = short_conv_forward_batch(
5308 &normed,
5309 b,
5310 w,
5311 &cfg,
5312 &mut self.kv_cache.layers[li].linear_state,
5313 pool.as_deref(),
5314 );
5315 for (dst, &a) in h.iter_mut().zip(&attn) {
5316 *dst += a;
5317 }
5318 }
5319 AttnKind::Mla(w) => {
5320 let inv_freq_l = self.layer_inv_freq(li);
5323 let rs = self.layer_rope_scale(li);
5324 let mut normed = vec![0.0f32; hs];
5325 for bi in 0..b {
5326 inference::rms_norm_into(
5327 &h[bi * hs..(bi + 1) * hs],
5328 &lw.input_norm,
5329 eps,
5330 norm_style,
5331 &mut normed,
5332 );
5333 let ao = mla_attention(
5334 w,
5335 &normed,
5336 &mut self.kv_cache.layers[li],
5337 start_pos + bi,
5338 &inv_freq_l,
5339 rs,
5340 eps,
5341 pool.as_deref(),
5342 );
5343 for (dst, &a) in h[bi * hs..(bi + 1) * hs].iter_mut().zip(&ao) {
5344 *dst += a;
5345 }
5346 }
5347 }
5348 AttnKind::Full {
5349 wq,
5350 wk,
5351 wv,
5352 wo,
5353 q_norm,
5354 k_norm,
5355 output_gate,
5356 softplus_gate,
5357 bias,
5358 } => {
5359 let mut normed = vec![0.0f32; b * hs];
5363 for bi in 0..b {
5364 inference::rms_norm_into(
5365 &h[bi * hs..(bi + 1) * hs],
5366 &lw.input_norm,
5367 eps,
5368 norm_style,
5369 &mut normed[bi * hs..(bi + 1) * hs],
5370 );
5371 }
5372 let inv_freq_l = self.layer_inv_freq(li);
5373 let (nkv_l, hd_l, rd_l) = self.layer_geom(li);
5374 let cfg = QwenAttnCfg {
5375 num_heads: self.layer_num_heads(li),
5376 num_kv_heads: nkv_l,
5377 head_dim: hd_l,
5378 hidden_size: hs,
5379 position: start_pos,
5380 inv_freq: &inv_freq_l,
5381 rotary_dim: rd_l,
5382 scale: self.attn_scale,
5383 softcap: self.attn_softcap,
5384 window: self.layer_window(li),
5385 v_norm: self.attn_v_norm,
5386 q_norm: q_norm.as_deref(),
5387 k_norm: k_norm.as_deref(),
5388 output_gate: *output_gate,
5389 softplus_gate: softplus_gate
5390 .as_ref()
5391 .map(|(gate, per_head)| (gate, *per_head)),
5392 rope_scale: self.layer_rope_scale(li),
5393 bias: bias
5394 .as_ref()
5395 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
5396 rms_eps: eps,
5397 norm_style,
5398 pool: pool.as_deref(),
5399 };
5400 let mut attn = attention::qwen_attention_batch(
5401 &normed,
5402 b,
5403 wq,
5404 wk,
5405 wv,
5406 wo,
5407 &mut self.kv_cache.layers[li],
5408 &cfg,
5409 );
5410 if let Some(w) = &lw.attn_out_norm {
5411 for bi in 0..b {
5412 inference::rms_norm_into(
5413 &attn[bi * hs..(bi + 1) * hs],
5414 w,
5415 eps,
5416 norm_style,
5417 &mut normed[bi * hs..(bi + 1) * hs],
5418 );
5419 }
5420 attn.copy_from_slice(&normed);
5421 }
5422 for (dst, &a) in h.iter_mut().zip(&attn) {
5423 *dst += a;
5424 }
5425 }
5426 AttnKind::Linear(w) => {
5427 for bi in 0..b {
5428 let normed = inference::rms_norm(
5429 &h[bi * hs..(bi + 1) * hs],
5430 &lw.input_norm,
5431 eps,
5432 norm_style,
5433 );
5434 vmf_phase_forward(
5435 &normed,
5436 w,
5437 &self.vmf_cfg.expect("linear layer without vmf_cfg"),
5438 &mut self.kv_cache.layers[li].linear_state,
5439 pool.as_deref(),
5440 )
5441 .iter()
5442 .enumerate()
5443 .for_each(|(i, &a)| h[bi * hs + i] += a);
5444 }
5445 }
5446 }
5447
5448 let lw = &self.weights.layers[self.phys_layer(li)];
5450 let mut post = vec![0.0f32; b * hs];
5451 for bi in 0..b {
5452 let r =
5453 inference::rms_norm(&h[bi * hs..(bi + 1) * hs], &lw.post_norm, eps, norm_style);
5454 post[bi * hs..(bi + 1) * hs].copy_from_slice(&r);
5455 }
5456 let mask_row = task_mask
5459 .filter(|m| m.ffn_active_count(li) < self.intermediate_size)
5460 .and_then(|m| m.ffn_masks.get(li))
5461 .map(|v| v.as_slice());
5462 let mut ffn = match &lw.ffn {
5463 FfnKind::Dense(d) if !d.segs.is_empty() => {
5464 tube_ffn(d, &post, b, pool.as_deref(), mask_row)
5465 }
5466 FfnKind::Dense(d) => dense_ffn_batch(d, &post, b, pool.as_deref(), mask_row),
5467 FfnKind::Moe(m) => moe_ffn_batch(m, &post, b, hs, pool.as_deref(), None),
5468 FfnKind::DenseMoe(dm) => {
5471 let mut out = vec![0.0f32; b * hs];
5472 for bi in 0..b {
5473 let r = dense_moe_ffn(
5474 dm,
5475 &post[bi * hs..(bi + 1) * hs],
5476 &h[bi * hs..(bi + 1) * hs],
5477 eps,
5478 norm_style,
5479 pool.as_deref(),
5480 );
5481 out[bi * hs..(bi + 1) * hs].copy_from_slice(&r);
5482 }
5483 out
5484 }
5485 };
5486 if let Some(w) = &lw.ffn_out_norm {
5487 for bi in 0..b {
5488 inference::rms_norm_into(
5489 &ffn[bi * hs..(bi + 1) * hs],
5490 w,
5491 eps,
5492 norm_style,
5493 &mut post[bi * hs..(bi + 1) * hs],
5494 );
5495 }
5496 ffn.copy_from_slice(&post);
5497 }
5498 for (dst, &f) in h.iter_mut().zip(&ffn) {
5499 *dst += f;
5500 }
5501 if let Some(sc) = lw.layer_scale {
5502 for v in h.iter_mut() {
5503 *v *= sc;
5504 }
5505 }
5506 if let Ok(tp) = std::env::var("CMF_TRACE_POS") {
5507 if let Ok(t) = tp.parse::<usize>() {
5508 if t >= start_pos && t < start_pos + b {
5509 let bi = t - start_pos;
5510 let row = &h[bi * hs..(bi + 1) * hs];
5511 let n: f32 = row.iter().map(|x| x * x).sum::<f32>().sqrt();
5512 eprintln!(
5513 "BATCH pos {t} after layer {li}: |h| = {n:.6} h0 {:.6} h1 {:.6}",
5514 row[0], row[1]
5515 );
5516 }
5517 }
5518 }
5519 if std::env::var("CMF_DEBUG_LAYERS").is_ok() {
5523 let row = &h[(b - 1) * hs..b * hs];
5524 let rms =
5525 (row.iter().map(|&v| (v as f64) * (v as f64)).sum::<f64>() / hs as f64).sqrt();
5526 let mx = row.iter().fold(0f32, |m, &v| m.max(v.abs()));
5527 eprintln!(
5528 "layer {li:>3} {:>10} ffn={:<5} rms={rms:>12.4} max={mx:>12.4}",
5529 match &self.weights.layers[self.phys_layer(li)].attn {
5530 AttnKind::LinearGdn(_) => "gdn",
5531 AttnKind::Linear(_) => "vmf",
5532 AttnKind::ShortConv(_) => "conv",
5533 _ => "attn",
5534 },
5535 match &lw.ffn {
5536 FfnKind::Moe(_) => "moe",
5537 FfnKind::Dense(_) => "dense",
5538 FfnKind::DenseMoe(_) => "dense+moe",
5539 },
5540 );
5541 }
5542 if self.is_loop_end(li) && li + 1 < self.num_layers {
5544 for bi in 0..b {
5545 let normed = inference::rms_norm(
5546 &h[bi * hs..(bi + 1) * hs],
5547 &self.weights.final_norm,
5548 eps,
5549 norm_style,
5550 );
5551 h[bi * hs..(bi + 1) * hs].copy_from_slice(&normed);
5552 }
5553 }
5554 if std::env::var("CMF_TRACE_H").is_ok() {
5555 let n = h[..hs].iter().map(|v| v.abs()).sum::<f32>() / hs as f32;
5556 let mx = h[..hs].iter().fold(0.0f32, |a, &v| a.max(v.abs()));
5557 eprintln!(
5558 "layer {li}: mean|h|={n:.4} max|h|={mx:.2} scale={:?}",
5559 lw.layer_scale
5560 );
5561 }
5562 }
5563 crate::gpu::set_layer(-1); h
5565 }
5566
5567 fn embed_single(&self, id: u32) -> Vec<f32> {
5569 let mut out = vec![0.0f32; self.hidden_size];
5570 if (id as usize) < self.weights.embed_tokens.rows() {
5571 self.weights.embed_tokens.row_f32(id as usize, &mut out);
5572 }
5573 if self.embed_multiplier != 1.0 {
5574 for v in out.iter_mut() {
5575 *v *= self.embed_multiplier;
5576 }
5577 }
5578 if self.dsv4.is_some() {
5582 let mut v = vec![0.0f32; self.hidden_size.max(1)];
5583 v[0] = id as f32;
5584 return v;
5585 }
5586 if let Some(b) = &self.g3n {
5589 return b.0.extend_embedding(id, &out, self.pool.as_deref());
5590 }
5591 out
5592 }
5593
5594 #[cfg(target_os = "macos")]
5600 fn chunk_run_gpu(
5601 &mut self,
5602 li0: usize,
5603 h: &mut [f32],
5604 b: usize,
5605 pos0: usize,
5606 embed_ids: Option<&[u32]>,
5607 cap: usize,
5608 ) -> usize {
5609 if !crate::gpu::enabled_here()
5613 || std::env::var("CMF_GPU_CHUNK")
5614 .map(|v| v == "0")
5615 .unwrap_or(false)
5616 || b < 32
5617 || self.swa.is_some()
5618 || self.global_attn.is_some()
5619 || self.attn_v_norm
5620 || (self.attn_scale - 1.0 / (self.head_dim as f32).sqrt()).abs() > 1e-9
5621 {
5622 return li0;
5623 }
5624 let Some(model) = self.model.clone() else {
5625 return li0;
5626 };
5627 let inv_freq = self.inv_freq.clone();
5628 let (nh, nkv, hd, hs) = (
5629 self.num_heads,
5630 self.num_kv_heads,
5631 self.head_dim,
5632 self.hidden_size,
5633 );
5634 let loop_end = if self.loop_final_norm {
5638 ((li0 / self.physical_layers) + 1) * self.physical_layers
5639 } else {
5640 self.num_layers
5641 };
5642 let mut layers: Vec<crate::gpu_metal::ChunkLayer> = Vec::new();
5643 let mut stored_at: Vec<usize> = Vec::new();
5644 for li in li0..self.num_layers.min(loop_end).min(cap) {
5645 let lw = &self.weights.layers[self.phys_layer(li)];
5646 if lw.attn_out_norm.is_some() || lw.ffn_out_norm.is_some() || lw.layer_scale.is_some() {
5647 break;
5648 }
5649 let AttnKind::Full {
5650 wq,
5651 wk,
5652 wv,
5653 wo,
5654 q_norm,
5655 k_norm,
5656 output_gate: false,
5657 softplus_gate: None,
5658 bias,
5659 } = &lw.attn
5660 else {
5661 break;
5662 };
5663 let FfnKind::Dense(d) = &lw.ffn else { break };
5664 if d.act != Act::Silu || !d.segs.is_empty() {
5665 break;
5666 }
5667 fn cw(t: &QTensor) -> Option<(usize, usize, usize, &[f32])> {
5672 t.q8_row_parts()
5673 .or_else(|| t.q4t_parts().map(|(i, r, c)| (i, r, c, &[][..])))
5674 .or_else(|| t.q4tp_parts().map(|(i, r, c)| (i, r, c, &[][..])))
5675 }
5676 let parts = (
5677 cw(wq),
5678 cw(wk),
5679 cw(wv),
5680 cw(wo),
5681 cw(&d.gate_proj),
5682 cw(&d.up_proj),
5683 cw(&d.down_proj),
5684 );
5685 let (Some(pq), Some(pk), Some(pv), Some(po), Some(pg), Some(pu), Some(pd)) = parts
5686 else {
5687 break;
5688 };
5689 let layer = &self.kv_cache.layers[li];
5690 if layer.mode != crate::kv_cache::KvMode::F32 || layer.o1.is_some() {
5691 break;
5692 }
5693 stored_at.push(layer.head_len(0));
5694 layers.push(crate::gpu_metal::ChunkLayer {
5695 model: &model,
5696 kv_id: self.graph_kv_id,
5697 layer: li,
5698 wq: pq,
5699 wk: pk,
5700 wv: pv,
5701 wo: po,
5702 gate: pg,
5703 up: pu,
5704 down: pd,
5705 input_norm: &lw.input_norm,
5706 post_norm: &lw.post_norm,
5707 bias: bias
5708 .as_ref()
5709 .map(|(a, bb, cc)| (a.as_slice(), bb.as_slice(), cc.as_slice())),
5710 q_norm: q_norm.as_deref(),
5711 k_norm: k_norm.as_deref(),
5712 inv_freq: &inv_freq,
5713 rd: self.rotary_dim,
5714 nh,
5715 nkv,
5716 hd,
5717 hs,
5718 inter: d.gate_proj.rows(),
5719 gemma: matches!(self.norm_style, cortiq_core::NormStyle::Gemma),
5720 eps: self.rms_eps as f32,
5721 });
5722 }
5723 if layers.is_empty() {
5724 return li0;
5725 }
5726 let row = nkv * hd;
5727 let mut store: Vec<(Vec<f32>, Vec<f32>, Vec<f32>)> = stored_at
5728 .iter()
5729 .map(|&st| (vec![0f32; b * row], vec![0f32; b * row], vec![0f32; st + b]))
5730 .collect();
5731 let mut io: Vec<crate::gpu_metal::ChunkIo> = Vec::with_capacity(layers.len());
5732 for (i, (ok, ov, oi)) in store.iter_mut().enumerate() {
5733 let li = layers[i].layer;
5734 let layer = &self.kv_cache.layers[li];
5735 io.push(crate::gpu_metal::ChunkIo {
5736 cpu_stored: stored_at[i],
5737 cpu_k: (0..nkv).map(|g| layer.head_keys(g)).collect(),
5738 cpu_v: (0..nkv).map(|g| layer.head_values(g)).collect(),
5739 out_k: ok,
5740 out_v: ov,
5741 imp: oi,
5742 });
5743 }
5744 let n_run = layers.len();
5745 let last = layers.last().map(|l| l.layer + 1).unwrap_or(li0);
5746 let ep = embed_ids.and_then(|ids| {
5749 self.weights
5750 .embed_tokens
5751 .q8_row_parts()
5752 .map(|(idx, rows, _c, rs)| crate::gpu_metal::ChunkEmbed {
5753 idx,
5754 rows,
5755 row_scale: rs,
5756 ids,
5757 mult: self.embed_multiplier,
5758 })
5759 });
5760 if embed_ids.is_some() && ep.is_none() {
5761 return li0;
5762 }
5763 if !crate::gpu_metal::chunk_run_gpu(&layers, &mut io, h, b, pos0, ep.as_ref()) {
5764 return li0;
5765 }
5766 drop(io);
5767 drop(layers);
5768 for (i, (ok, ov, oi)) in store.iter().enumerate().take(n_run) {
5771 let li = li0 + i;
5772 let layer = &mut self.kv_cache.layers[li];
5773 for bi in 0..b {
5774 layer.append(
5775 &ok[bi * row..(bi + 1) * row],
5776 &ov[bi * row..(bi + 1) * row],
5777 &[],
5778 );
5779 }
5780 layer.accumulate_imp(oi);
5781 }
5782 last
5783 }
5784
5785 fn layer_is_local(&self, li: usize) -> bool {
5788 if let Some(layers) = &self.sliding_layers {
5789 return layers.get(li).copied().unwrap_or(false);
5790 }
5791 match self.swa {
5792 Some((_, pattern)) => (li + 1) % pattern.max(1) != 0,
5793 None => false,
5794 }
5795 }
5796
5797 fn layer_inv_freq(&self, li: usize) -> std::sync::Arc<Vec<f32>> {
5800 if self.layer_is_local(li) {
5801 if let Some(f) = &self.inv_freq_local {
5802 return f.clone();
5803 }
5804 } else if let Some(f) = &self.inv_freq_global {
5805 return f.clone();
5806 }
5807 self.inv_freq.clone()
5808 }
5809
5810 fn layer_window(&self, li: usize) -> Option<usize> {
5812 self.swa
5813 .and_then(|(w, _)| self.layer_is_local(li).then_some(w))
5814 }
5815
5816 fn layer_num_heads(&self, li: usize) -> usize {
5817 self.attention_heads_per_layer
5818 .as_ref()
5819 .and_then(|v| v.get(li).copied())
5820 .unwrap_or(self.num_heads)
5821 }
5822
5823 fn layer_rope_scale(&self, li: usize) -> f32 {
5824 if self.layer_is_local(li) {
5825 self.rope_scale_local
5826 } else {
5827 self.rope_scale
5828 }
5829 }
5830
5831 fn layer_geom(&self, li: usize) -> (usize, usize, usize) {
5834 if !self.layer_is_local(li) {
5835 if let Some((ghd, gkv)) = self.global_attn {
5836 return (gkv, ghd, ghd);
5837 }
5838 }
5839 (
5840 self.num_kv_heads,
5841 self.head_dim,
5842 if self.layer_is_local(li) {
5843 self.rotary_dim_local.unwrap_or(self.rotary_dim)
5844 } else {
5845 self.rotary_dim
5846 },
5847 )
5848 }
5849
5850 fn forward_layers(
5852 &mut self,
5853 hidden: &[f32],
5854 position: usize,
5855 task_mask: Option<&TaskMask>,
5856 ) -> Vec<f32> {
5857 self.forward_layers_upto(hidden, position, task_mask, None)
5858 }
5859
5860 pub fn embed_id(&self, id: u32) -> Vec<f32> {
5868 self.embed_single(id)
5869 }
5870
5871 pub fn split_supported(&self) -> Result<(), String> {
5875 if self.dsv4.is_some() {
5876 return Err(
5877 "network split: DeepSeek-V4 runs its own fused stack (not splittable yet)".into(),
5878 );
5879 }
5880 if self.g3n.is_some() {
5881 return Err(
5882 "network split: Gemma-3n runs its own AltUp stack (not splittable yet)".into(),
5883 );
5884 }
5885 Ok(())
5886 }
5887
5888 pub fn forward_span(
5893 &mut self,
5894 hidden: &[f32],
5895 position: usize,
5896 from: usize,
5897 upto: usize,
5898 task_mask: Option<&TaskMask>,
5899 ) -> Result<Vec<f32>, String> {
5900 self.split_supported()?;
5901 if from > upto || upto >= self.num_layers {
5902 return Err(format!(
5903 "forward_span: layer range {from}..={upto} outside 0..{}",
5904 self.num_layers
5905 ));
5906 }
5907 if hidden.len() != self.hidden_size {
5908 return Err(format!(
5909 "forward_span: hidden len {} ≠ hidden_size {}",
5910 hidden.len(),
5911 self.hidden_size
5912 ));
5913 }
5914 Ok(self.forward_layers_span(hidden, position, task_mask, from, Some(upto)))
5915 }
5916
5917 pub fn logits_from_hidden(&mut self, hidden: &[f32]) -> Vec<f32> {
5920 let normed = inference::rms_norm(
5921 hidden,
5922 &self.weights.final_norm,
5923 self.rms_eps,
5924 self.norm_style,
5925 );
5926 self.lm_head_forward(&normed)
5927 }
5928
5929 pub fn sample_next(&mut self, logits: &[f32], past_tokens: &[u32]) -> u32 {
5931 sampler::sample_with_scratch(
5932 logits,
5933 &self.sampler_config,
5934 past_tokens,
5935 &mut self.rng,
5936 &mut self.sampler_scratch,
5937 )
5938 }
5939
5940 pub fn reset_session(&mut self) {
5942 self.kv_cache.clear();
5943 self.kv_history.clear();
5944 crate::gpu::graph_kv_reset(self.graph_kv_id);
5945 }
5946
5947 pub fn prefill_span_ids(
5953 &mut self,
5954 ids: &[u32],
5955 start_pos: usize,
5956 upto: usize,
5957 task_mask: Option<&TaskMask>,
5958 ) -> Result<Vec<f32>, String> {
5959 self.split_supported()?;
5960 if upto >= self.num_layers {
5961 return Err(format!(
5962 "prefill_span_ids: upto {upto} outside 0..{}",
5963 self.num_layers
5964 ));
5965 }
5966 if self.can_prefill_batched() && !self.graph_prefill_preferred() {
5970 Ok(self.prefill_batch_span(PrefillIn::Ids(ids), start_pos, task_mask, 0, upto + 1))
5971 } else {
5972 let hs = self.hidden_size;
5973 let mut out = Vec::with_capacity(ids.len() * hs);
5974 for (i, &id) in ids.iter().enumerate() {
5975 let emb = self.embed_id(id);
5976 out.extend_from_slice(&self.forward_span(
5977 &emb,
5978 start_pos + i,
5979 0,
5980 upto,
5981 task_mask,
5982 )?);
5983 }
5984 Ok(out)
5985 }
5986 }
5987
5988 pub fn prefill_span_hidden(
5991 &mut self,
5992 hidden: &[f32],
5993 start_pos: usize,
5994 from: usize,
5995 upto: usize,
5996 task_mask: Option<&TaskMask>,
5997 ) -> Result<Vec<f32>, String> {
5998 self.split_supported()?;
5999 let hs = self.hidden_size;
6000 if hidden.is_empty() || hidden.len() % hs != 0 {
6001 return Err(format!(
6002 "prefill_span_hidden: {} floats is not a multiple of hidden {hs}",
6003 hidden.len()
6004 ));
6005 }
6006 if from > upto || upto >= self.num_layers {
6007 return Err(format!(
6008 "prefill_span_hidden: layer range {from}..={upto} outside 0..{}",
6009 self.num_layers
6010 ));
6011 }
6012 if self.can_prefill_batched() && !self.graph_prefill_preferred() {
6013 Ok(self.prefill_batch_span(
6014 PrefillIn::Hidden(hidden),
6015 start_pos,
6016 task_mask,
6017 from,
6018 upto + 1,
6019 ))
6020 } else {
6021 let b = hidden.len() / hs;
6022 let mut out = Vec::with_capacity(hidden.len());
6023 for i in 0..b {
6024 let h = self.forward_span(
6025 &hidden[i * hs..(i + 1) * hs],
6026 start_pos + i,
6027 from,
6028 upto,
6029 task_mask,
6030 )?;
6031 out.extend_from_slice(&h);
6032 }
6033 Ok(out)
6034 }
6035 }
6036
6037 fn try_token_graph_wgpu(
6041 &self,
6042 hidden: &[f32],
6043 position: usize,
6044 logits_out: &mut Vec<f32>,
6045 layers_run: &mut usize,
6046 ) -> Option<Vec<f32>> {
6047 self.try_token_graph_wgpu_steps(
6048 hidden,
6049 position,
6050 logits_out,
6051 1,
6052 None,
6053 Some(layers_run),
6054 0,
6055 self.num_layers,
6056 )
6057 }
6058
6059 fn try_token_graph_wgpu_span(
6063 &self,
6064 hidden: &[f32],
6065 position: usize,
6066 logits_out: &mut Vec<f32>,
6067 from: usize,
6068 upto_excl: usize,
6069 layers_run: &mut usize,
6070 ) -> Option<Vec<f32>> {
6071 self.try_token_graph_wgpu_steps(
6072 hidden,
6073 position,
6074 logits_out,
6075 1,
6076 None,
6077 Some(layers_run),
6078 from,
6079 upto_excl,
6080 )
6081 }
6082
6083 fn try_multi_burst(&self, t_next: u32, position: usize, k: usize) -> Option<Vec<u32>> {
6087 if self.o1_active() || self.attn_softcap > 0.0 {
6088 return None;
6089 }
6090 let graph_on = crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Decode);
6091 if !graph_on || crate::gpu::graph_unsupported() {
6092 return None;
6099 }
6100 let emb = self.embed_single(t_next);
6101 let mut lg = Vec::new();
6102 let mut ids = Vec::new();
6103 self.try_token_graph_wgpu_steps(
6104 &emb,
6105 position,
6106 &mut lg,
6107 k,
6108 Some(&mut ids),
6109 None,
6110 0,
6111 self.num_layers,
6112 )?;
6113 (ids.len() == k).then_some(ids)
6114 }
6115
6116 fn try_token_graph_wgpu_steps(
6120 &self,
6121 hidden: &[f32],
6122 position: usize,
6123 logits_out: &mut Vec<f32>,
6124 steps: usize,
6125 ids_out: Option<&mut Vec<u32>>,
6126 layers_run: Option<&mut usize>,
6127 from: usize,
6128 upto_excl: usize,
6129 ) -> Option<Vec<f32>> {
6130 let o1_gpu = std::env::var("CMF_O1_GPU").as_deref() == Ok("1");
6133 if (self.o1_active() && !o1_gpu) || self.attn_softcap > 0.0 {
6134 return None;
6138 }
6139 let o1_views: Vec<Option<Vec<crate::nystrom::O1DeviceView<'_>>>> = (from..upto_excl)
6144 .map(|li| {
6145 if !o1_gpu {
6146 return None;
6147 }
6148 self.kv_cache.layers[self.phys_layer(li)].o1_views()
6149 })
6150 .collect();
6151 if self.o1_active() && o1_gpu {
6152 let want: usize = (from..upto_excl)
6155 .filter(|li| self.kv_cache.layers[self.phys_layer(*li)].o1.is_some())
6156 .count();
6157 let have = o1_views.iter().filter(|v| v.is_some()).count();
6158 if want == 0 || have != want {
6159 use std::sync::atomic::{AtomicUsize, Ordering};
6169 static LAST: AtomicUsize = AtomicUsize::new(usize::MAX);
6170 let code = have * 1000 + want;
6171 if LAST.swap(code, Ordering::Relaxed) != code {
6172 tracing::warn!(
6173 "o1 graph: {have} of {want} layers sealed — per-op until all seal"
6174 );
6175 }
6176 return None;
6177 }
6178 }
6179 let nh = self.num_heads;
6180 let (nkv, hd, rd) = self.layer_geom(0);
6181 let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
6182 let mut layers = Vec::with_capacity(upto_excl - from);
6183 let mut model = None;
6184 let dbg = std::env::var("CMF_GRAPH_DEBUG").is_ok();
6185 fn gw(t: &QTensor) -> Option<crate::gpu::GraphW<'_>> {
6186 if let Some((_, i, kind, rs)) = t.graph_weight() {
6187 return Some(crate::gpu::GraphW {
6188 idx: i,
6189 kind,
6190 row_scale: rs,
6191 data: &[],
6192 });
6193 }
6194 t.as_f32().map(|d| crate::gpu::GraphW {
6196 idx: 0,
6197 kind: 4,
6198 row_scale: &[],
6199 data: d,
6200 })
6201 }
6202 for li in from..upto_excl {
6203 let lw = &self.weights.layers[self.phys_layer(li)];
6204 if dbg {
6205 let ak = match &lw.attn {
6206 AttnKind::Mla(_) => "Mla".into(),
6207 AttnKind::Full {
6208 output_gate, bias, ..
6209 } => format!("Full gate={output_gate} bias={}", bias.is_some()),
6210 AttnKind::LinearGdn(_) => "LinearGdn".into(),
6211 AttnKind::Kda(_) => "Kda".into(),
6212 AttnKind::Linear(_) => "Linear".into(),
6213 AttnKind::ShortConv(_) => "ShortConv".into(),
6214 };
6215 let fk = match &lw.ffn {
6216 FfnKind::Dense(_) => "Dense",
6217 FfnKind::Moe(_) => "Moe",
6218 FfnKind::DenseMoe(_) => "DenseMoe",
6219 };
6220 eprintln!("graph L{li}: attn={ak} ffn={fk}");
6221 }
6222 let gffn = match &lw.ffn {
6223 FfnKind::DenseMoe(_) => return None, FfnKind::Dense(d) if !d.segs.is_empty() => return None,
6227 FfnKind::Dense(d) => crate::gpu::GraphFfn::Dense {
6228 gate: gw(&d.gate_proj)?,
6229 up: gw(&d.up_proj)?,
6230 down: gw(&d.down_proj)?,
6231 },
6232 FfnKind::Moe(m) => {
6233 if m.router_sigmoid
6238 || m.expert_bias.is_some()
6239 || m.route_tau.is_some()
6240 || m.mask.is_some()
6241 {
6242 return None;
6243 }
6244 let (se, sg) = m.shared.as_ref()?;
6245 let sgate = gw(sg.as_ref()?)?;
6246 let router = gw(&m.router)?;
6247 let inter = m.experts.first()?.gate_proj.rows();
6248 let mut experts = Vec::with_capacity(m.experts.len() + 1);
6249 let mut q4tp: Option<bool> = None;
6252 let mut gu_q2: Option<bool> = None;
6255 for e in m.experts.iter().chain(std::iter::once(se)) {
6256 if !matches!(e.act, Act::Silu)
6257 || e.gate_proj.rows() != inter
6258 || e.up_proj.rows() != inter
6259 {
6260 return None;
6261 }
6262 let (mm, gi, ui, di, is_p, is_q2) = match e.gate_proj.mapped_q4t() {
6263 Some((mm, gi)) => (
6264 mm,
6265 gi,
6266 e.up_proj.mapped_q4t()?.1,
6267 e.down_proj.mapped_q4t()?.1,
6268 false,
6269 false,
6270 ),
6271 None => match e.gate_proj.mapped_q2tp() {
6272 Some((mm, gi)) => (
6273 mm,
6274 gi,
6275 e.up_proj.mapped_q2tp()?.1,
6276 e.down_proj.mapped_q4tp()?.1,
6277 true,
6278 true,
6279 ),
6280 None => {
6281 let (mm, gi) = e.gate_proj.mapped_q4tp()?;
6282 (
6283 mm,
6284 gi,
6285 e.up_proj.mapped_q4tp()?.1,
6286 e.down_proj.mapped_q4tp()?.1,
6287 true,
6288 false,
6289 )
6290 }
6291 },
6292 };
6293 if *q4tp.get_or_insert(is_p) != is_p || *gu_q2.get_or_insert(is_q2) != is_q2
6294 {
6295 tracing::warn!(
6301 "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."
6302 );
6303 return None;
6304 }
6305 model.get_or_insert_with(|| mm.clone());
6306 experts.push((gi, ui, di));
6307 }
6308 crate::gpu::GraphFfn::Moe {
6309 router,
6310 shared_gate: sgate,
6311 experts,
6312 n_exp: m.experts.len(),
6313 top_k: std::env::var("CMF_TOPK_PROBE")
6319 .ok()
6320 .and_then(|v| v.parse::<usize>().ok())
6321 .filter(|k| *k > 0 && *k <= m.top_k)
6322 .unwrap_or(m.top_k),
6323 inter,
6324 norm_topk: m.norm_topk_prob,
6325 q4tp: q4tp?,
6326 gu_q2: gu_q2.unwrap_or(false),
6327 }
6328 }
6329 };
6330 let attn = match &lw.attn {
6331 AttnKind::Full {
6332 wq,
6333 wk,
6334 wv,
6335 wo,
6336 q_norm,
6337 k_norm,
6338 output_gate,
6339 softplus_gate,
6340 bias,
6341 } => {
6342 if softplus_gate.is_some() || self.attention_heads_per_layer.is_some() {
6343 return None;
6344 }
6345 let (m, _, _, _) = wq.graph_weight()?;
6346 model = Some(m.clone());
6347 crate::gpu::GraphAttn::Full {
6348 wq: gw(wq)?,
6349 wk: gw(wk)?,
6350 wv: gw(wv)?,
6351 wo: gw(wo)?,
6352 q_norm: q_norm.as_deref(),
6353 k_norm: k_norm.as_deref(),
6354 bias: bias
6355 .as_ref()
6356 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
6357 output_gate: *output_gate,
6358 cpu_k: self.kv_cache.layers[li].k_heads(),
6359 cpu_v: self.kv_cache.layers[li].v_heads(),
6360 }
6361 }
6362 AttnKind::LinearGdn(w) => {
6363 let cfg = self.gdn_cfg?;
6364 let (m, _, _, _) = w.in_proj_qkv.graph_weight()?;
6365 model = Some(m.clone());
6366 crate::gpu::GraphAttn::Gdn {
6367 qkv: gw(&w.in_proj_qkv)?,
6368 z: gw(&w.in_proj_z)?,
6369 a: gw(&w.in_proj_a)?,
6370 b: gw(&w.in_proj_b)?,
6371 out: gw(&w.out_proj)?,
6372 conv1d: &w.conv1d,
6373 a_log: &w.a_log,
6374 dt_bias: &w.dt_bias,
6375 norm: &w.norm,
6376 nv: cfg.num_v_heads,
6377 nk: cfg.num_k_heads,
6378 dk: cfg.key_head_dim,
6379 dv: cfg.value_head_dim,
6380 kk: cfg.conv_kernel,
6381 cpu_state: &self.kv_cache.layers[self.phys_layer(li)].linear_state,
6382 }
6383 }
6384 _ => return None,
6385 };
6386 layers.push(crate::gpu::GraphLayer {
6387 input_norm: &lw.input_norm,
6388 attn,
6389 post_norm: &lw.post_norm,
6390 ffn: gffn,
6391 });
6392 }
6393 let model = model?;
6394 let lm_gw = if upto_excl == self.num_layers
6400 && self.graph_want_logits
6401 && std::env::var("CMF_GPU_LMHEAD")
6402 .map(|v| v != "0")
6403 .unwrap_or(true)
6404 {
6405 self.weights.lm_head.graph_weight().map(|(_, i, kind, rs)| {
6406 (
6407 crate::gpu::GraphW {
6408 idx: i,
6409 kind,
6410 row_scale: rs,
6411 data: &[],
6412 },
6413 self.weights.lm_head.rows(),
6414 )
6415 })
6416 } else {
6417 None
6418 };
6419 let lm = lm_gw.as_ref().map(|(gw, rows)| (gw, *rows));
6420 let emb_gw = if steps > 1 {
6422 self.weights
6423 .embed_tokens
6424 .graph_weight()
6425 .map(|(_, i, kind, rs)| {
6426 (
6427 crate::gpu::GraphW {
6428 idx: i,
6429 kind,
6430 row_scale: rs,
6431 data: &[],
6432 },
6433 self.weights.embed_tokens.rows(),
6434 self.embed_multiplier,
6435 )
6436 })
6437 } else {
6438 None
6439 };
6440
6441 let loop_norm_at: Vec<usize> = if self.loop_final_norm {
6447 (from..upto_excl.min(self.num_layers - 1))
6448 .filter(|&li| (li + 1) % self.physical_layers == 0)
6449 .map(|li| li - from)
6450 .collect()
6451 } else {
6452 Vec::new()
6453 };
6454 let mut h = hidden.to_vec();
6455 crate::gpu::forward_token_graph(
6456 &model,
6457 self.graph_kv_id,
6458 &layers,
6459 &o1_views,
6460 self.o1_epoch,
6461 &self.inv_freq,
6462 &mut h,
6463 nh,
6464 nkv,
6465 hd,
6466 rd,
6467 self.hidden_size,
6468 self.intermediate_size,
6469 position,
6470 self.kv_cache.max_seq_len,
6471 gemma,
6472 self.rms_eps as f32,
6473 lm,
6474 &self.weights.final_norm,
6475 logits_out,
6476 &loop_norm_at,
6477 steps,
6478 emb_gw.as_ref().map(|(gw, rows, m)| (gw, *rows, *m)),
6479 ids_out,
6480 layers_run,
6481 from,
6482 false,
6483 )
6484 .then_some(h)
6485 }
6486
6487 #[cfg(target_os = "macos")]
6496 #[allow(clippy::type_complexity)]
6497 fn metal_rows_plan(&self) -> Option<(Vec<MetalRowsItem<'_>>, std::sync::Arc<cortiq_core::CmfModel>, Option<crate::gpu_metal::GdnGpuCfg>)> {
6498 use crate::gpu_metal::{AttnGpuLayer, GdnGpuCfg, GdnGpuLayer, MetalFfn};
6499 if !crate::gpu::q1_force()
6500 || !crate::gpu::enabled_here()
6501 || std::env::var("CMF_GPU_BLOCK").map(|v| v == "0").unwrap_or(false)
6502 || self.attn_softcap > 0.0
6503 || self.o1_active()
6504 || self.swa.is_some()
6505 || self.global_attn.is_some()
6506 || self.attention_heads_per_layer.is_some()
6507 || self.attn_v_norm
6508 || self.loop_final_norm
6509 || (self.attn_scale - 1.0 / (self.head_dim as f32).sqrt()).abs() > 1e-9
6510 {
6511 return None;
6512 }
6513 let attend_contract = self.head_dim % 4 == 0
6514 && self.head_dim <= 256
6515 && self.rotary_dim >= 2
6516 && self.rotary_dim <= self.head_dim
6517 && (self.rotary_dim / 2) % 32 == 0
6518 && self.num_kv_heads > 0
6519 && self.num_heads % self.num_kv_heads == 0;
6520 if !attend_contract {
6521 return None;
6522 }
6523 let mut plan: Vec<MetalRowsItem> = Vec::new();
6524 let mut model_ref: Option<std::sync::Arc<cortiq_core::CmfModel>> = None;
6525 for li in 0..self.num_layers {
6526 let lw = &self.weights.layers[self.phys_layer(li)];
6527 if lw.attn_out_norm.is_some() || lw.ffn_out_norm.is_some() || lw.layer_scale.is_some() {
6528 return None;
6529 }
6530 let ffn = match &lw.ffn {
6531 FfnKind::Dense(d) if d.act == Act::Silu && d.segs.is_empty() => {
6532 let (Some(g), Some(u), Some(dn)) =
6533 (d.gate_proj.q1_parts(), d.up_proj.q1_parts(), d.down_proj.q1_parts())
6534 else {
6535 return None;
6536 };
6537 MetalFfn::Dense { gate: g, up: u, down: dn }
6538 }
6539 _ => return None,
6540 };
6541 match &lw.attn {
6542 AttnKind::LinearGdn(w) if self.gdn_cfg.is_some() => {
6543 let (Some(qkv), Some(z), Some(a), Some(bb), Some(out)) = (
6544 w.in_proj_qkv.q1_parts(),
6545 w.in_proj_z.q1_parts(),
6546 w.in_proj_a.f32_parts(),
6547 w.in_proj_b.f32_parts(),
6548 w.out_proj.q1_parts(),
6549 ) else {
6550 return None;
6551 };
6552 if let QTensor::Mapped { model, .. } = &w.in_proj_qkv {
6553 model_ref.get_or_insert_with(|| model.clone());
6554 }
6555 let gl = GdnGpuLayer {
6556 attn_norm: &lw.input_norm,
6557 post_norm: &lw.post_norm,
6558 qkv,
6559 z,
6560 a,
6561 b: bb,
6562 out,
6563 ffn,
6564 conv1d: &w.conv1d,
6565 a_log: &w.a_log,
6566 dt_bias: &w.dt_bias,
6567 gnorm: &w.norm,
6568 };
6569 match plan.last_mut() {
6570 Some(MetalRowsItem::Gdn { run, .. }) => run.push(gl),
6571 _ => plan.push(MetalRowsItem::Gdn { run: vec![gl], first: li }),
6572 }
6573 }
6574 AttnKind::Full {
6575 wq,
6576 wk,
6577 wv,
6578 wo,
6579 q_norm,
6580 k_norm,
6581 output_gate,
6582 softplus_gate: None,
6583 bias: None,
6584 } => {
6585 let (Some(pq), Some(pk), Some(pv), Some(po)) =
6586 (wq.q1_parts(), wk.q1_parts(), wv.q1_parts(), wo.q1_parts())
6587 else {
6588 return None;
6589 };
6590 if let QTensor::Mapped { model, .. } = wq {
6591 model_ref.get_or_insert_with(|| model.clone());
6592 }
6593 let cache = &self.kv_cache.layers[li];
6594 if cache.mode != crate::kv_cache::KvMode::F32 || cache.o1.is_some() {
6595 return None;
6596 }
6597 plan.push(MetalRowsItem::Attn {
6598 l: AttnGpuLayer {
6599 attn_norm: &lw.input_norm,
6600 post_norm: &lw.post_norm,
6601 wq: pq,
6602 wk: pk,
6603 wv: pv,
6604 wo: po,
6605 ffn,
6606 },
6607 li,
6608 q_norm: q_norm.as_deref(),
6609 k_norm: k_norm.as_deref(),
6610 output_gate: *output_gate,
6611 });
6612 }
6613 _ => return None,
6614 }
6615 }
6616 let model = model_ref?;
6617 let gcfg = self.gdn_cfg.map(|cfg| GdnGpuCfg {
6618 nv: cfg.num_v_heads,
6619 nk: cfg.num_k_heads,
6620 dk: cfg.key_head_dim,
6621 dv: cfg.value_head_dim,
6622 kk: cfg.conv_kernel,
6623 hidden: self.hidden_size,
6624 inter: self.intermediate_size,
6625 c_dim: cfg.conv_dim(),
6626 eps: cfg.rms_eps as f32,
6627 gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
6628 });
6629 Some((plan, model, gcfg))
6630 }
6631
6632 #[cfg(target_os = "macos")]
6634 #[allow(clippy::too_many_arguments)]
6635 fn metal_attn_params<'a>(
6636 li: usize,
6637 cache: &'a crate::kv_cache::LayerKvCache,
6638 q_norm: Option<&'a [f32]>,
6639 k_norm: Option<&'a [f32]>,
6640 output_gate: bool,
6641 inv_freq: &'a [f32],
6642 geom: (usize, usize, usize, usize),
6643 pos0: usize,
6644 kv_id: u64,
6645 eps: f32,
6646 gemma: bool,
6647 ) -> (crate::gpu_metal::AttnDeviceParams<'a>, usize) {
6648 let (nh, nkv, hd, rd) = geom;
6649 let cpu_k: Vec<&[f32]> = (0..nkv).map(|g| cache.head_keys(g)).collect();
6650 let cpu_v: Vec<&[f32]> = (0..nkv).map(|g| cache.head_values(g)).collect();
6651 let cpu_stored = cpu_k[0].len() / hd;
6652 (
6653 crate::gpu_metal::AttnDeviceParams {
6654 kv_id,
6655 layer: li,
6656 nh,
6657 nkv,
6658 hd,
6659 rd,
6660 position: pos0,
6661 eps,
6662 gemma,
6663 output_gate,
6664 q_norm,
6665 k_norm,
6666 inv_freq,
6667 cpu_k,
6668 cpu_v,
6669 cpu_stored,
6670 o1: None,
6671 },
6672 cpu_stored,
6673 )
6674 }
6675
6676 #[cfg(target_os = "macos")]
6681 #[allow(clippy::type_complexity)]
6682 fn metal_rows_run(
6683 &mut self,
6684 hiddens: &mut [f32],
6685 pos0: usize,
6686 b: usize,
6687 prefill: bool,
6688 spec: Option<((usize, usize, usize), &[f32], &mut Vec<f32>)>,
6689 ) -> Option<MetalVerifyPending> {
6690 use crate::gpu_metal::{GraphDims, VerifyGraph};
6691 let want = self.gdn_cfg.map(|c| c.state_len()).unwrap_or(0);
6692 for l in &mut self.kv_cache.layers {
6693 if l.linear_state.len() != want && want > 0 {
6694 l.linear_state = vec![0f32; want];
6695 }
6696 }
6697 let (plan, model, gcfg) = self.metal_rows_plan()?;
6698 let dims = GraphDims {
6699 hidden: self.hidden_size,
6700 eps: self.rms_eps as f32,
6701 gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
6702 };
6703 let mut graph = if prefill {
6704 VerifyGraph::new_prefill(&model, dims, hiddens, b)?
6705 } else {
6706 VerifyGraph::new(&model, dims, hiddens, b)?
6707 };
6708 let geom = (self.num_heads, self.num_kv_heads, self.head_dim, self.rotary_dim);
6709 let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
6710 let eps = self.rms_eps as f32;
6711 let kv_id = self.graph_kv_id;
6712 let inv_freq = self.inv_freq.clone();
6713 for item in &plan {
6714 let ok = match item {
6715 MetalRowsItem::Gdn { run, .. } => gcfg
6716 .as_ref()
6717 .map(|gc| run.iter().all(|l| graph.gdn_ok(l, gc)))
6718 .unwrap_or(false),
6719 MetalRowsItem::Attn { l, li, q_norm, k_norm, output_gate } => {
6720 let (p, _) = Self::metal_attn_params(*li, &self.kv_cache.layers[*li], *q_norm, *k_norm, *output_gate, &inv_freq, geom, pos0, kv_id, eps, gemma);
6721 graph.attn_ok(l, &p)
6722 }
6723 };
6724 if !ok {
6725 use std::sync::atomic::{AtomicBool, Ordering};
6726 static SAID: AtomicBool = AtomicBool::new(false);
6727 if !SAID.swap(true, Ordering::Relaxed) {
6728 tracing::warn!("metal rows graph: a layer failed preflight — declining");
6729 }
6730 return None;
6731 }
6732 }
6733 let lm = match &spec {
6734 Some((lm, _, _)) => {
6735 if !graph.lm_head_ok(*lm) {
6736 return None;
6737 }
6738 Some(*lm)
6739 }
6740 None => None,
6741 };
6742 let mut gdn_layers = Vec::new();
6743 let mut attn_layers = Vec::new();
6744 for item in &plan {
6745 match item {
6746 MetalRowsItem::Gdn { run, first } => {
6747 let ro: Vec<&[f32]> = self.kv_cache.layers[*first..*first + run.len()]
6748 .iter()
6749 .map(|l| l.linear_state.as_slice())
6750 .collect();
6751 if !graph.encode_gdn_run_b(run, &ro, gcfg.as_ref().unwrap()) {
6752 return None;
6753 }
6754 gdn_layers.extend(*first..*first + run.len());
6755 }
6756 MetalRowsItem::Attn { l, li, q_norm, k_norm, output_gate } => {
6757 let (p, cpu_stored) = Self::metal_attn_params(*li, &self.kv_cache.layers[*li], *q_norm, *k_norm, *output_gate, &inv_freq, geom, pos0, kv_id, eps, gemma);
6758 if !graph.encode_attn_b(l, &p) {
6759 return None;
6760 }
6761 attn_layers.push((*li, cpu_stored));
6762 }
6763 }
6764 }
6765 if let (Some(lm), Some((_, final_norm, _))) = (lm, spec.as_ref()) {
6766 if !graph.encode_lm_head_b(final_norm, lm) {
6767 return None;
6768 }
6769 }
6770 graph.sync();
6771 if let Some((lm, _, logits)) = spec {
6772 logits.resize(b * lm.1, 0.0);
6773 graph.read_logits(logits);
6774 }
6775 graph.read_hidden(hiddens);
6776 Some(MetalVerifyPending { graph, gdn_layers, attn_layers })
6777 }
6778
6779 #[cfg(target_os = "macos")]
6785 fn try_batch_graph_metal(
6786 &mut self,
6787 hiddens: &mut [f32],
6788 positions: &[usize],
6789 b: usize,
6790 spec: Option<((usize, usize, usize), &[f32], &mut Vec<f32>)>,
6791 ) -> bool {
6792 let _t0 = std::time::Instant::now();
6793 if positions.len() != b
6794 || positions.windows(2).any(|w| w[1] != w[0] + 1)
6795 || hiddens.len() != b * self.hidden_size
6796 {
6797 return false;
6798 }
6799 let Some(pending) = self.metal_rows_run(hiddens, positions[0], b, false, spec) else {
6800 return false;
6801 };
6802 if std::env::var("CMF_GRAPH_SPEC_TIME").is_ok() {
6803 eprintln!("metal-verify: {:.1} ms | b={b}", _t0.elapsed().as_secs_f64() * 1e3);
6804 }
6805 self.metal_verify = Some(pending);
6806 true
6807 }
6808
6809 #[cfg(target_os = "macos")]
6814 fn prefill_batch_metal(&mut self, ids: &[u32], start_pos: usize) -> Option<Vec<f32>> {
6815 let b = ids.len();
6816 if b == 0 || b > 512 {
6817 return None;
6818 }
6819 let hs = self.hidden_size;
6820 let mut hiddens = vec![0f32; b * hs];
6821 for (j, &id) in ids.iter().enumerate() {
6822 let e = self.embed_single(id);
6823 hiddens[j * hs..(j + 1) * hs].copy_from_slice(&e);
6824 }
6825 let mut pending = self.metal_rows_run(&mut hiddens, start_pos, b, true, None)?;
6826 let idxs = pending.gdn_layers.clone();
6828 let mut outs: Vec<&mut [f32]> = self
6829 .kv_cache
6830 .layers
6831 .iter_mut()
6832 .enumerate()
6833 .filter(|(i, _)| idxs.binary_search(i).is_ok())
6834 .map(|(_, l)| l.linear_state.as_mut_slice())
6835 .collect();
6836 pending.graph.finish_states(&mut outs);
6837 let (nkv, hd) = (self.num_kv_heads, self.head_dim);
6838 let mut kbuf = vec![0f32; b * nkv * hd];
6839 let mut vbuf = vec![0f32; b * nkv * hd];
6840 for (li, cpu_stored) in &pending.attn_layers {
6841 if crate::gpu_metal::kv_mirror_read_rows(self.graph_kv_id, *li, nkv, hd, *cpu_stored, b, &mut kbuf, &mut vbuf) {
6842 let cache = &mut self.kv_cache.layers[*li];
6843 for r in 0..b {
6844 cache.append(&kbuf[r * nkv * hd..(r + 1) * nkv * hd], &vbuf[r * nkv * hd..(r + 1) * nkv * hd], &[]);
6845 }
6846 crate::gpu_metal::kv_mirror_set_stored(self.graph_kv_id, *li, cpu_stored + b);
6847 }
6848 }
6849 Some(hiddens)
6850 }
6851
6852 #[cfg(target_os = "macos")]
6856 fn metal_verify_commit(&mut self, a: usize) -> bool {
6857 let Some(mut pending) = self.metal_verify.take() else {
6858 return false;
6859 };
6860 let n = a + 1;
6861 let idxs = pending.gdn_layers.clone();
6863 let mut outs: Vec<&mut [f32]> = self
6864 .kv_cache
6865 .layers
6866 .iter_mut()
6867 .enumerate()
6868 .filter(|(i, _)| idxs.binary_search(i).is_ok())
6869 .map(|(_, l)| l.linear_state.as_mut_slice())
6870 .collect();
6871 if !pending.graph.commit(n, &mut outs) {
6872 return false;
6873 }
6874 let (nkv, hd) = (self.num_kv_heads, self.head_dim);
6875 let mut kbuf = vec![0f32; n * nkv * hd];
6876 let mut vbuf = vec![0f32; n * nkv * hd];
6877 for (li, cpu_stored) in &pending.attn_layers {
6878 if crate::gpu_metal::kv_mirror_read_rows(self.graph_kv_id, *li, nkv, hd, *cpu_stored, n, &mut kbuf, &mut vbuf) {
6879 let cache = &mut self.kv_cache.layers[*li];
6880 for r in 0..n {
6881 cache.append(&kbuf[r * nkv * hd..(r + 1) * nkv * hd], &vbuf[r * nkv * hd..(r + 1) * nkv * hd], &[]);
6882 }
6883 crate::gpu_metal::kv_mirror_set_stored(self.graph_kv_id, *li, cpu_stored + n);
6884 }
6885 }
6886 true
6887 }
6888
6889 #[cfg(target_os = "macos")]
6895 fn mtp_warm_batch_metal(&mut self, m: &mut MtpModule, pairs: &[(&[f32], u32)], first_pos: usize) -> bool {
6896 use crate::gpu_metal::{AttnDeviceParams, AttnGpuLayer, GraphDims, MetalFfn, VerifyGraph};
6897 let b = pairs.len();
6898 if b == 0 || b > 512 || m.kv.mode != crate::kv_cache::KvMode::F32 || m.kv.o1.is_some() {
6899 return false;
6900 }
6901 let AttnKind::Full { wq, wk, wv, wo, q_norm, k_norm, output_gate, softplus_gate: None, bias: None } = &m.layer.attn else {
6902 return false;
6903 };
6904 let FfnKind::Dense(d) = &m.layer.ffn else { return false };
6905 if !d.segs.is_empty() {
6906 return false;
6907 }
6908 let (Some(pq), Some(pk), Some(pv), Some(po)) = (wq.q1_parts(), wk.q1_parts(), wv.q1_parts(), wo.q1_parts()) else {
6909 return false;
6910 };
6911 let (Some(g), Some(u), Some(dn)) = (d.gate_proj.q1_parts(), d.up_proj.q1_parts(), d.down_proj.q1_parts()) else {
6912 return false;
6913 };
6914 let Some(eh) = m.eh_proj.q1_parts() else { return false };
6915 let QTensor::Mapped { model, .. } = wq else { return false };
6916 let model = model.clone();
6917 let hs = self.hidden_size;
6918 let mut cat = vec![0f32; b * 2 * hs];
6920 for (j, (h, tok)) in pairs.iter().enumerate() {
6921 let e = self.embed_single(*tok);
6922 let (ce, ch) = cat[j * 2 * hs..(j + 1) * 2 * hs].split_at_mut(hs);
6923 inference::rms_norm_into(&e, &m.enorm, self.rms_eps, self.norm_style, ce);
6924 inference::rms_norm_into(h, &m.hnorm, self.rms_eps, self.norm_style, ch);
6925 }
6926 let dims = GraphDims { hidden: hs, eps: self.rms_eps as f32, gemma: self.norm_style == cortiq_core::NormStyle::Gemma };
6927 let Some(mut graph) = VerifyGraph::new_via_proj(&model, dims, eh, &cat, b) else {
6928 return false;
6929 };
6930 let l = AttnGpuLayer {
6931 attn_norm: &m.layer.input_norm,
6932 post_norm: &m.layer.post_norm,
6933 wq: pq,
6934 wk: pk,
6935 wv: pv,
6936 wo: po,
6937 ffn: MetalFfn::Dense { gate: g, up: u, down: dn },
6938 };
6939 let (nh, nkv, hd, rd) = (self.num_heads, self.num_kv_heads, self.head_dim, self.rotary_dim);
6940 let inv_freq = self.inv_freq.clone();
6941 let cpu_stored;
6942 {
6943 let cache = &m.kv;
6944 let cpu_k: Vec<&[f32]> = (0..nkv).map(|g| cache.head_keys(g)).collect();
6945 let cpu_v: Vec<&[f32]> = (0..nkv).map(|g| cache.head_values(g)).collect();
6946 cpu_stored = cpu_k[0].len() / hd;
6947 if cpu_stored != first_pos {
6948 return false;
6949 }
6950 let p = AttnDeviceParams {
6951 kv_id: self.mtp_kv_id(),
6952 layer: Self::MTP_LAYER_BASE,
6953 nh,
6954 nkv,
6955 hd,
6956 rd,
6957 position: first_pos,
6958 eps: self.rms_eps as f32,
6959 gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
6960 output_gate: *output_gate,
6961 q_norm: q_norm.as_deref(),
6962 k_norm: k_norm.as_deref(),
6963 inv_freq: &inv_freq,
6964 cpu_k,
6965 cpu_v,
6966 cpu_stored,
6967 o1: None,
6968 };
6969 if !graph.attn_ok(&l, &p) || !graph.encode_attn_b(&l, &p) {
6970 return false;
6971 }
6972 }
6973 graph.sync();
6974 let mut kbuf = vec![0f32; b * nkv * hd];
6975 let mut vbuf = vec![0f32; b * nkv * hd];
6976 if !crate::gpu_metal::kv_mirror_read_rows(self.mtp_kv_id(), Self::MTP_LAYER_BASE, nkv, hd, cpu_stored, b, &mut kbuf, &mut vbuf) {
6977 return false;
6978 }
6979 for r in 0..b {
6980 m.kv.append(&kbuf[r * nkv * hd..(r + 1) * nkv * hd], &vbuf[r * nkv * hd..(r + 1) * nkv * hd], &[]);
6981 }
6982 crate::gpu_metal::kv_mirror_set_stored(self.mtp_kv_id(), Self::MTP_LAYER_BASE, cpu_stored + b);
6983 true
6984 }
6985
6986 fn draft_vocab_rows(head_rows: usize) -> usize {
6989 static N: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
6990 let n = *N.get_or_init(|| {
6991 std::env::var("CMF_DRAFT_VOCAB")
6992 .ok()
6993 .and_then(|v| v.parse().ok())
6994 .unwrap_or(65536)
6995 });
6996 if n == 0 { head_rows } else { n.min(head_rows) }
6997 }
6998
6999 #[cfg(target_os = "macos")]
7004 fn mtp_step_metal(
7005 &mut self,
7006 m: &mut MtpModule,
7007 hidden: &[f32],
7008 next_token: u32,
7009 position: usize,
7010 want_logits: bool,
7011 ) -> Option<(Vec<f32>, Vec<f32>)> {
7012 use crate::gpu_metal::{AttnDeviceParams, AttnGpuLayer, GraphDims, MetalFfn, TokenGraph};
7013 if std::env::var("CMF_MTP_GRAPH").as_deref() == Ok("0")
7014 || !crate::gpu::q1_force()
7015 || !crate::gpu::enabled_here()
7016 || self.attn_softcap > 0.0
7017 || self.attention_heads_per_layer.is_some()
7018 || m.kv.mode != crate::kv_cache::KvMode::F32
7019 || m.kv.o1.is_some()
7020 {
7021 return None;
7022 }
7023 let AttnKind::Full {
7024 wq,
7025 wk,
7026 wv,
7027 wo,
7028 q_norm,
7029 k_norm,
7030 output_gate,
7031 softplus_gate: None,
7032 bias: None,
7033 } = &m.layer.attn
7034 else {
7035 return None;
7036 };
7037 let FfnKind::Dense(d) = &m.layer.ffn else { return None };
7038 if d.act != Act::Silu || !d.segs.is_empty() {
7039 return None;
7040 }
7041 let (pq, pk, pv, po) = (wq.q1_parts()?, wk.q1_parts()?, wv.q1_parts()?, wo.q1_parts()?);
7042 let (g, u, dn) = (d.gate_proj.q1_parts()?, d.up_proj.q1_parts()?, d.down_proj.q1_parts()?);
7043 let QTensor::Mapped { model, .. } = wq else { return None };
7044 let model = model.clone();
7045 let lm = if want_logits { Some(self.weights.lm_head.q1_parts()?) } else { None };
7046 let dims = GraphDims {
7047 hidden: self.hidden_size,
7048 eps: self.rms_eps as f32,
7049 gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
7050 };
7051 let hs = self.hidden_size;
7054 let mut x = vec![0f32; hs];
7055 let mut graph = TokenGraph::new(&model, dims, &x)?;
7056 let mut folded = false;
7057 if let Some(eh) = m.eh_proj.q1_parts() {
7058 let e = self.embed_single(next_token);
7059 let mut cat = vec![0.0f32; 2 * hs];
7060 let (cat_e, cat_h) = cat.split_at_mut(hs);
7061 inference::rms_norm_into(&e, &m.enorm, self.rms_eps, self.norm_style, cat_e);
7062 inference::rms_norm_into(hidden, &m.hnorm, self.rms_eps, self.norm_style, cat_h);
7063 folded = graph.encode_input_proj(eh, &cat);
7064 }
7065 if !folded {
7066 x = self.mtp_block_input(m, hidden, next_token);
7067 graph = TokenGraph::new(&model, dims, &x)?;
7068 }
7069 let l = AttnGpuLayer {
7070 attn_norm: &m.layer.input_norm,
7071 post_norm: &m.layer.post_norm,
7072 wq: pq,
7073 wk: pk,
7074 wv: pv,
7075 wo: po,
7076 ffn: MetalFfn::Dense { gate: g, up: u, down: dn },
7077 };
7078 let (nh, nkv, hd, rd) = (self.num_heads, self.num_kv_heads, self.head_dim, self.rotary_dim);
7079 let inv_freq = self.inv_freq.clone();
7080 {
7081 let cache = &m.kv;
7082 let cpu_k: Vec<&[f32]> = (0..nkv).map(|g| cache.head_keys(g)).collect();
7083 let cpu_v: Vec<&[f32]> = (0..nkv).map(|g| cache.head_values(g)).collect();
7084 let cpu_stored = cpu_k[0].len() / hd;
7085 let p = AttnDeviceParams {
7086 kv_id: self.mtp_kv_id(),
7087 layer: Self::MTP_LAYER_BASE,
7088 nh,
7089 nkv,
7090 hd,
7091 rd,
7092 position,
7093 eps: self.rms_eps as f32,
7094 gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
7095 output_gate: *output_gate,
7096 q_norm: q_norm.as_deref(),
7097 k_norm: k_norm.as_deref(),
7098 inv_freq: &inv_freq,
7099 cpu_k,
7100 cpu_v,
7101 cpu_stored,
7102 o1: None,
7103 };
7104 if !graph.attn_device_ok(&l, &p) || !graph.encode_attn_device(&l, &p) {
7105 return None;
7106 }
7107 }
7108 let draft_rows = if let Some(lm) = lm { Self::draft_vocab_rows(lm.1) } else { 0 };
7114 if let Some(lm) = lm {
7115 if !graph.lm_head_ok(lm) {
7116 return None;
7117 }
7118 if draft_rows < lm.1 {
7119 if !graph.encode_lm_head_part(&m.final_norm, lm, draft_rows) {
7120 return None;
7121 }
7122 } else {
7123 graph.encode_lm_head(&m.final_norm, lm);
7124 }
7125 }
7126 graph.sync();
7127 let mut logits = Vec::new();
7128 if let Some(lm) = lm {
7129 let n_read = draft_rows.min(lm.1).min(self.vocab_size);
7130 logits = attention::take_buf(n_read);
7131 graph.read_logits(&mut logits);
7132 logits.resize(self.vocab_size, f32::NEG_INFINITY);
7134 }
7135 graph.finish(&mut x);
7136 let mut krow = attention::take_buf(nkv * hd);
7137 let mut vrow = attention::take_buf(nkv * hd);
7138 if crate::gpu_metal::kv_mirror_read_last(self.mtp_kv_id(), Self::MTP_LAYER_BASE, nkv, hd, &mut krow, &mut vrow) {
7139 m.kv.append(&krow, &vrow, &[]);
7140 }
7141 attention::recycle_buf(&mut krow);
7142 attention::recycle_buf(&mut vrow);
7143 Some((logits, x))
7144 }
7145
7146 fn try_batch_graph_wgpu(
7147 &self,
7148 hiddens: &mut [f32],
7149 positions: &[usize],
7150 k: usize,
7151 spec: Option<crate::gpu::SpecTail<'_>>,
7152 ) -> bool {
7153 let _tb = std::time::Instant::now();
7154 if self.attn_softcap > 0.0 {
7155 return false; }
7157 if self.o1_active() {
7158 return false;
7159 }
7160 let nh = self.num_heads;
7161 let (nkv, hd, rd) = self.layer_geom(0);
7162 let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
7163 fn gw(t: &QTensor) -> Option<crate::gpu::GraphW<'_>> {
7164 if let Some((_, i, kind, rs)) = t.graph_weight() {
7165 return Some(crate::gpu::GraphW {
7166 idx: i,
7167 kind,
7168 row_scale: rs,
7169 data: &[],
7170 });
7171 }
7172 t.as_f32().map(|d| crate::gpu::GraphW {
7173 idx: 0,
7174 kind: 4,
7175 row_scale: &[],
7176 data: d,
7177 })
7178 }
7179 let built: Option<(
7180 Vec<crate::gpu::GraphLayer<'_>>,
7181 std::sync::Arc<cortiq_core::CmfModel>,
7182 )> = (|| {
7183 let mut layers = Vec::with_capacity(self.num_layers);
7184 let mut model = None;
7185 for li in 0..self.num_layers {
7186 let lw = &self.weights.layers[self.phys_layer(li)];
7187 let gffn = match &lw.ffn {
7194 FfnKind::Dense(d) if !d.segs.is_empty() => return None,
7195 FfnKind::Dense(d) => crate::gpu::GraphFfn::Dense {
7196 gate: gw(&d.gate_proj)?,
7197 up: gw(&d.up_proj)?,
7198 down: gw(&d.down_proj)?,
7199 },
7200 FfnKind::Moe(m) => {
7201 if m.router_sigmoid
7202 || m.expert_bias.is_some()
7203 || m.route_tau.is_some()
7204 || m.mask.is_some()
7205 {
7206 return None;
7207 }
7208 let (se, sg) = m.shared.as_ref()?;
7209 let sgate = gw(sg.as_ref()?)?;
7210 let router = gw(&m.router)?;
7211 let inter = m.experts.first()?.gate_proj.rows();
7212 let mut experts = Vec::with_capacity(m.experts.len() + 1);
7213 let mut q4tp: Option<bool> = None;
7214 let mut gu_q2: Option<bool> = None;
7215 for e in m.experts.iter().chain(std::iter::once(se)) {
7216 if !matches!(e.act, Act::Silu)
7217 || e.gate_proj.rows() != inter
7218 || e.up_proj.rows() != inter
7219 {
7220 return None;
7221 }
7222 let (mm, gi, ui, di, is_p, is_q2) = match e.gate_proj.mapped_q4t() {
7226 Some((mm, gi)) => (
7227 mm,
7228 gi,
7229 e.up_proj.mapped_q4t()?.1,
7230 e.down_proj.mapped_q4t()?.1,
7231 false,
7232 false,
7233 ),
7234 None => match e.gate_proj.mapped_q2tp() {
7235 Some((mm, gi)) => (
7236 mm,
7237 gi,
7238 e.up_proj.mapped_q2tp()?.1,
7239 e.down_proj.mapped_q4tp()?.1,
7240 true,
7241 true,
7242 ),
7243 None => {
7244 let (mm, gi) = e.gate_proj.mapped_q4tp()?;
7245 (
7246 mm,
7247 gi,
7248 e.up_proj.mapped_q4tp()?.1,
7249 e.down_proj.mapped_q4tp()?.1,
7250 true,
7251 false,
7252 )
7253 }
7254 },
7255 };
7256 if *q4tp.get_or_insert(is_p) != is_p
7257 || *gu_q2.get_or_insert(is_q2) != is_q2
7258 {
7259 return None;
7260 }
7261 model.get_or_insert_with(|| mm.clone());
7262 experts.push((gi, ui, di));
7263 }
7264 crate::gpu::GraphFfn::Moe {
7265 router,
7266 shared_gate: sgate,
7267 experts,
7268 n_exp: m.experts.len(),
7269 top_k: m.top_k,
7270 inter,
7271 norm_topk: m.norm_topk_prob,
7272 q4tp: q4tp?,
7273 gu_q2: gu_q2.unwrap_or(false),
7274 }
7275 }
7276 _ => return None,
7277 };
7278 let attn = match &lw.attn {
7279 AttnKind::Full {
7280 wq,
7281 wk,
7282 wv,
7283 wo,
7284 q_norm,
7285 k_norm,
7286 output_gate,
7287 softplus_gate,
7288 bias,
7289 } => {
7290 if softplus_gate.is_some() || self.attention_heads_per_layer.is_some() {
7291 return None;
7292 }
7293 let (m, _, _, _) = wq.graph_weight()?;
7294 model = Some(m.clone());
7295 crate::gpu::GraphAttn::Full {
7296 wq: gw(wq)?,
7297 wk: gw(wk)?,
7298 wv: gw(wv)?,
7299 wo: gw(wo)?,
7300 q_norm: q_norm.as_deref(),
7301 k_norm: k_norm.as_deref(),
7302 bias: bias
7303 .as_ref()
7304 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
7305 output_gate: *output_gate,
7306 cpu_k: self.kv_cache.layers[li].k_heads(),
7307 cpu_v: self.kv_cache.layers[li].v_heads(),
7308 }
7309 }
7310 AttnKind::LinearGdn(w) => {
7311 let cfg = self.gdn_cfg?;
7312 let (m, _, _, _) = w.in_proj_qkv.graph_weight()?;
7313 model = Some(m.clone());
7314 crate::gpu::GraphAttn::Gdn {
7315 qkv: gw(&w.in_proj_qkv)?,
7316 z: gw(&w.in_proj_z)?,
7317 a: gw(&w.in_proj_a)?,
7318 b: gw(&w.in_proj_b)?,
7319 out: gw(&w.out_proj)?,
7320 conv1d: &w.conv1d,
7321 a_log: &w.a_log,
7322 dt_bias: &w.dt_bias,
7323 norm: &w.norm,
7324 nv: cfg.num_v_heads,
7325 nk: cfg.num_k_heads,
7326 dk: cfg.key_head_dim,
7327 dv: cfg.value_head_dim,
7328 kk: cfg.conv_kernel,
7329 cpu_state: &self.kv_cache.layers[self.phys_layer(li)].linear_state,
7330 }
7331 }
7332 _ => return None,
7333 };
7334 layers.push(crate::gpu::GraphLayer {
7335 input_norm: &lw.input_norm,
7336 attn,
7337 post_norm: &lw.post_norm,
7338 ffn: gffn,
7339 });
7340 }
7341 Some((layers, model?))
7342 })();
7343 let Some((layers, model)) = built else {
7344 {
7345 use std::sync::atomic::{AtomicBool, Ordering};
7346 static SAID: AtomicBool = AtomicBool::new(false);
7347 if !SAID.swap(true, Ordering::Relaxed) {
7348 tracing::warn!("batch graph: BUILDER refused (layer weights/kinds)");
7349 }
7350 }
7351 return false;
7352 };
7353 if std::env::var("CMF_GRAPH_SPEC_TIME").is_ok() {
7354 eprintln!("batch-build: {:.1} ms", _tb.elapsed().as_secs_f64() * 1e3);
7355 }
7356 crate::gpu::forward_batch_graph(
7357 &model,
7358 self.graph_kv_id,
7359 &layers,
7360 &self.inv_freq,
7361 hiddens,
7362 nh,
7363 nkv,
7364 hd,
7365 rd,
7366 self.hidden_size,
7367 self.intermediate_size,
7368 positions,
7369 self.kv_cache.max_seq_len,
7370 gemma,
7371 self.rms_eps as f32,
7372 k,
7373 spec,
7374 )
7375 }
7376
7377 fn draft_probe() -> bool {
7381 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
7382 *ON.get_or_init(|| std::env::var("CMF_DSV4_DRAFT_PROBE").is_ok_and(|v| v != "0"))
7383 }
7384
7385 #[cfg(feature = "gpu")]
7397 fn dsv4_spec_on() -> bool {
7398 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
7399 *ON.get_or_init(|| {
7400 std::env::var("CMF_DSV4_SPEC")
7401 .map(|v| v != "0")
7402 .unwrap_or(true)
7403 })
7404 }
7405
7406 #[cfg(feature = "gpu")]
7413 fn dsv4_spec_step(
7414 &mut self,
7415 tip_token: u32,
7416 t_next: u32,
7417 next_pos: usize,
7418 drafted: &mut usize,
7419 accepted_ctr: &mut usize,
7420 ) -> Option<(Vec<u32>, usize)> {
7421 let t_all = std::time::Instant::now();
7422 if std::env::var("CMF_DSV4_SPEC_TIME").is_ok() {
7423 thread_local! {
7424 static LAST: std::cell::Cell<Option<std::time::Instant>> =
7425 const { std::cell::Cell::new(None) };
7426 }
7427 LAST.with(|l| {
7428 if let Some(prev) = l.get() {
7429 eprintln!(
7430 "между раундами {:.1} мс",
7431 prev.elapsed().as_secs_f64() * 1e3
7432 );
7433 }
7434 l.set(Some(std::time::Instant::now()));
7435 });
7436 }
7437 if std::env::var("CMF_DSV4_SPEC_DEBUG").is_ok() {
7438 eprintln!("spec_step: вход pos={next_pos}");
7439 }
7440 let n_layers = self.dsv4.as_ref().map(|b| b.1.len())?;
7441 let cfg = self.dsv4.as_ref().map(|b| b.2)?;
7442 if self.dspark.is_none() {
7444 let t = crate::dsv4::dspark_targets(&self.dsv4_mtp, &cfg, n_layers);
7445 if t.is_empty() {
7446 return None;
7447 }
7448 crate::dsv4::dspark_arm(&t, cfg.dim);
7449 self.dspark = Some(crate::dsv4::DsparkState::new(
7450 self.dsv4_mtp.len(),
7451 &cfg,
7452 t.len(),
7453 ));
7454 }
7455 let targets = crate::dsv4::dspark_targets(&self.dsv4_mtp, &cfg, n_layers);
7456 let pack = crate::dsv4::dspark_pack_get(&self.dsv4_mtp, &cfg);
7457 if pack.is_none() && std::env::var("CMF_DSV4_SPEC_DEBUG").is_ok() {
7458 eprintln!("spec_step: пак не построился (targets {targets:?})");
7459 }
7460 let pack = pack?;
7461 let block = crate::dsv4::dspark_block();
7462 let b_box = self.dsv4.as_mut()?;
7463 let (g, layers, st) = (&b_box.0, &b_box.1, &mut b_box.3);
7464 let ds = self.dspark.as_mut()?;
7465 let dbg = std::env::var("CMF_DSV4_SPEC_DEBUG").is_ok();
7468 if !crate::dsv4::dspark_take(&mut ds.main_hidden) && !ds.have_hidden {
7469 if dbg {
7470 eprintln!("spec_step: нет захвата");
7471 }
7472 return None;
7473 }
7474 ds.have_hidden = true;
7475 let tip_pos = next_pos.checked_sub(1)?;
7476 let draft_started = std::time::Instant::now();
7477 let mut conf = Vec::new();
7478 let props = crate::dsv4::dspark_draft_gpu(
7479 g,
7480 &self.dsv4_mtp,
7481 &cfg,
7482 ds,
7483 pack,
7484 st.kv_id,
7485 tip_token,
7486 tip_pos,
7487 self.pool.as_deref(),
7488 &mut conf,
7489 );
7490 self.dspark_draft_ns += draft_started.elapsed().as_nanos();
7491 *drafted += block;
7492 if props.is_empty() || props[0] != t_next {
7493 if dbg {
7494 eprintln!(
7495 "spec_step: черновик {} (props0={:?} t_next={t_next})",
7496 if props.is_empty() {
7497 "пуст"
7498 } else {
7499 "мимо"
7500 },
7501 props.first()
7502 );
7503 }
7504 return None;
7505 }
7506 let mut k_verify = crate::dsv4::dspark_verify_k().min(props.len());
7507 let conf_min = {
7513 static M: std::sync::OnceLock<f32> = std::sync::OnceLock::new();
7514 *M.get_or_init(|| {
7515 std::env::var("CMF_DSPARK_CONF_MIN")
7516 .ok()
7517 .and_then(|v| v.parse().ok())
7518 .unwrap_or(0.0)
7519 })
7520 };
7521 if conf_min > 0.0 && conf.len() >= props.len() {
7522 let mut keep = 1usize;
7523 while keep < k_verify && conf.get(keep).copied().unwrap_or(0.0) >= conf_min {
7524 keep += 1;
7525 }
7526 k_verify = k_verify.min(keep.max(2));
7527 }
7528 if k_verify < 2 {
7529 return None;
7530 }
7531 let mut fed = Vec::with_capacity(k_verify);
7532 fed.push(t_next);
7533 fed.extend_from_slice(&props[1..k_verify]);
7534 let mut argmax = Vec::new();
7535 let mut logits_all = Vec::new();
7536 let mut walked = Vec::new();
7537 let txn = crate::dsv4::dsv4_verify_chunk(
7538 g,
7539 layers,
7540 &cfg,
7541 st,
7542 &fed,
7543 next_pos,
7544 &self.inv_freq,
7545 self.pool.as_deref(),
7546 &targets,
7547 &mut argmax,
7548 &mut logits_all,
7549 &mut walked,
7550 );
7551 if txn.is_none() && dbg {
7552 eprintln!("spec_step: verify отказал");
7553 }
7554 let txn = txn?;
7555 let b = fed.len();
7556 let mut accepted = 1usize;
7557 while accepted < b && fed[accepted] == argmax[accepted - 1] {
7558 accepted += 1;
7559 }
7560 if std::env::var("CMF_DSV4_SPEC_FORCE_REJECT").is_ok_and(|v| v != "0") {
7565 accepted = 1;
7566 }
7567 if std::env::var("CMF_DSV4_SPEC_TRACE").is_ok() {
7568 eprintln!("spec@{next_pos}: fed={fed:?} argmax={argmax:?} accepted={accepted}");
7569 }
7570 let t_fin = std::time::Instant::now();
7571 if !crate::dsv4::dsv4_spec_finish(
7572 g,
7573 layers,
7574 &cfg,
7575 st,
7576 txn,
7577 accepted,
7578 &fed,
7579 &self.inv_freq,
7580 self.pool.as_deref(),
7581 ) {
7582 tracing::warn!("dsv4: спекулятивный откат не удался — состояние подозрительно");
7583 return None;
7584 }
7585 if std::env::var("CMF_DSV4_SPEC_TIME").is_ok() {
7586 eprintln!(
7587 "finish(k={accepted}): {:.1} мс",
7588 t_fin.elapsed().as_secs_f64() * 1e3
7589 );
7590 }
7591 *accepted_ctr += accepted - 1;
7592 let (hc, dim) = (cfg.hc_mult, cfg.dim);
7597 let dev_caps: Vec<usize> = targets
7604 .iter()
7605 .copied()
7606 .filter(|&t| {
7607 st.dev_set.get(t).copied().unwrap_or(false)
7608 && !st.partial_set.get(t).copied().unwrap_or(false)
7609 })
7610 .collect();
7611 let mut caps_all = vec![0.0f32; dev_caps.len() * b * hc * dim];
7612 if !crate::gpu_wgpu::dsv4_spec_cap_read_all(b, dev_caps.len(), hc * dim, &mut caps_all) {
7613 return None;
7614 }
7615 for t in 0..accepted {
7616 let tip = t + 1 == accepted;
7617 for (slot, &tl) in targets.iter().enumerate() {
7618 if let Some(di) = dev_caps.iter().position(|&d| d == tl) {
7619 let lo = (di * b + t) * hc * dim;
7620 crate::dsv4::dspark_capture(
7621 &caps_all[lo..lo + hc * dim],
7622 &cfg,
7623 slot,
7624 &mut ds.main_hidden,
7625 );
7626 } else if tip
7627 && crate::dsv4::dspark_peek_slot(slot, dim, {
7628 let lo = slot * dim;
7629 &mut ds.main_hidden[lo..lo + dim]
7630 })
7631 {
7632 } else {
7637 crate::dsv4::dspark_capture(
7641 &walked[t * hc * dim..(t + 1) * hc * dim],
7642 &cfg,
7643 slot,
7644 &mut ds.main_hidden,
7645 );
7646 }
7647 }
7648 crate::dsv4::dspark_ring_append(
7649 g,
7650 &self.dsv4_mtp,
7651 &cfg,
7652 ds,
7653 next_pos + t,
7654 self.pool.as_deref(),
7655 );
7656 }
7657 let row = logits_all[(accepted - 1) * cfg.vocab..accepted * cfg.vocab].to_vec();
7658 self.graph_logits = Some(row);
7659 if std::env::var("CMF_DSV4_TRUNK_PICK_DUMP").is_ok() {
7664 crate::dsv4::trunk_freq_note(&crate::dsv4::pick_tally_take());
7665 crate::dsv4::pick_tally_arm();
7666 }
7667 if std::env::var("CMF_DSV4_SPEC_TIME").is_ok() {
7668 eprintln!(
7669 "spec_step total {:.1} мс (k={accepted})",
7670 t_all.elapsed().as_secs_f64() * 1e3
7671 );
7672 }
7673 Some((fed[1..accepted].to_vec(), next_pos + accepted))
7674 }
7675
7676 fn dspark_probe(&mut self, position: usize, token_id: u32) {
7677 if self.dsv4_mtp.is_empty() || !Self::draft_probe() {
7678 return;
7679 }
7680 let trunk_now = crate::dsv4::pick_tally_take();
7682 crate::dsv4::trunk_freq_note(&trunk_now);
7683 if !trunk_now.is_empty() {
7684 self.dspark_trunk_picks.push(trunk_now);
7685 let keep = crate::dsv4::dspark_block();
7686 if self.dspark_trunk_picks.len() > keep {
7687 self.dspark_trunk_picks.remove(0);
7688 }
7689 }
7690 for p in std::mem::take(&mut self.dspark_pending) {
7693 let Some(i) = position.checked_sub(p.0 + 1) else {
7694 continue;
7695 };
7696 let mut p = p;
7697 if i < p.1.len() {
7698 if p.2 && p.1[i] == token_id {
7699 p.3 = i + 1;
7700 } else {
7701 p.2 = false;
7702 }
7703 if i + 1 < p.1.len() {
7704 self.dspark_pending.push(p);
7705 continue;
7706 }
7707 }
7708 self.dspark_hist.push(p.3);
7709 self.dspark_real.push(token_id);
7710 }
7711 let Some(b) = &mut self.dsv4 else { return };
7712 let (g, layers, cfg) = (&b.0, &b.1, b.2);
7713 let n_layers = layers.len();
7714 if self.dspark.is_none() {
7715 let t = crate::dsv4::dspark_targets(&self.dsv4_mtp, &cfg, n_layers);
7716 if t.is_empty() {
7717 return;
7718 }
7719 eprintln!(
7720 "DSpark: захват со слоёв {t:?}, блок {}",
7721 crate::dsv4::dspark_block()
7722 );
7723 crate::dsv4::dspark_arm(&t, cfg.dim);
7724 self.dspark = Some(crate::dsv4::DsparkState::new(
7725 self.dsv4_mtp.len(),
7726 &cfg,
7727 t.len(),
7728 ));
7729 }
7730 let ds = self.dspark.as_mut().unwrap();
7731 if !crate::dsv4::dspark_take(&mut ds.main_hidden) {
7732 return; }
7734 let mut conf = Vec::new();
7735 crate::dsv4::pick_tally_arm();
7736 let draft_started = std::time::Instant::now();
7741 #[cfg(feature = "gpu")]
7742 let gpu_draft = crate::dsv4::dspark_gpu_on();
7743 #[cfg(not(feature = "gpu"))]
7744 let gpu_draft = false;
7745 let props = if gpu_draft {
7746 #[cfg(feature = "gpu")]
7747 {
7748 let kv_id = b.3.kv_id;
7749 match crate::dsv4::dspark_pack_get(&self.dsv4_mtp, &cfg) {
7750 Some(pk) => crate::dsv4::dspark_draft_gpu(
7751 g,
7752 &self.dsv4_mtp,
7753 &cfg,
7754 ds,
7755 pk,
7756 kv_id,
7757 token_id,
7758 position,
7759 self.pool.as_deref(),
7760 &mut conf,
7761 ),
7762 None => Vec::new(),
7763 }
7764 }
7765 #[cfg(not(feature = "gpu"))]
7766 Vec::new()
7767 } else {
7768 crate::gpu::cpu_scope(|| {
7769 crate::dsv4::dspark_draft(
7770 g,
7771 &self.dsv4_mtp,
7772 &cfg,
7773 ds,
7774 token_id,
7775 position,
7776 self.pool.as_deref(),
7777 &mut conf,
7778 )
7779 })
7780 };
7781 self.dspark_draft_ns += draft_started.elapsed().as_nanos();
7782 let draft_picks = crate::dsv4::pick_tally_take();
7783 crate::dsv4::dspark_freq_note(&draft_picks);
7784 crate::dsv4::pick_tally_arm();
7787 if !props.is_empty() {
7788 let (tu, tt) = {
7792 let flat: Vec<(usize, Vec<usize>)> = self
7793 .dspark_trunk_picks
7794 .iter()
7795 .flat_map(|v| v.iter().cloned())
7796 .collect();
7797 let mut per: std::collections::HashMap<usize, Vec<usize>> =
7799 std::collections::HashMap::new();
7800 for (li, picks) in flat {
7801 per.entry(li).or_default().extend(picks);
7802 }
7803 let n = per.len().max(1);
7804 let mut u = 0usize;
7805 let mut t = 0usize;
7806 for (_, v) in per {
7807 t += v.len();
7808 u += v.iter().collect::<std::collections::HashSet<_>>().len();
7809 }
7810 (u / n, t / n)
7811 };
7812 let (du, dt) = crate::dsv4::tally_unique(&draft_picks);
7813 self.dspark_exp.push((tu, tt, du, dt));
7814 self.dspark_pending.push((position, props, true, 0));
7815 }
7816 if self.dspark_hist.len() >= 8 && self.dspark_hist.len() % 8 == 0 {
7817 let n = self.dspark_hist.len() as f32;
7818 let mean: f32 = self.dspark_hist.iter().sum::<usize>() as f32 / n;
7819 let block = crate::dsv4::dspark_block();
7820 let mut at = vec![0usize; block + 1];
7821 for &k in &self.dspark_hist {
7822 at[k] += 1;
7823 }
7824 let mut surv = Vec::with_capacity(block);
7826 for i in 1..=block {
7827 let k = at[i..].iter().sum::<usize>() as f32 / n;
7828 surv.push(format!("{k:.2}"));
7829 }
7830 let distinct = self
7831 .dspark_real
7832 .iter()
7833 .collect::<std::collections::HashSet<_>>()
7834 .len();
7835 let (tu, tt, du, dt) = self.dspark_exp.iter().fold((0, 0, 0, 0), |a, b| {
7836 (a.0 + b.0, a.1 + b.1, a.2 + b.2, a.3 + b.3)
7837 });
7838 let m = self.dspark_exp.len().max(1);
7839 eprintln!(
7840 "DSpark: черновиков {}, принято в среднем {mean:.2} из {block} \
7841 (токенов за проход {:.2}), распределение {at:?}, выживание [{}]",
7842 self.dspark_hist.len(),
7843 mean + 1.0,
7844 surv.join(" ")
7845 );
7846 eprintln!(
7847 "DSpark: разных токенов {distinct} из {} (вырожденность), \
7848 эксперты ствол {}/{} на слой за {block} токенов, \
7849 черновик {}/{} за блок, draft {:.2} мс/блок",
7850 self.dspark_real.len(),
7851 tu / m,
7852 tt / m,
7853 du / m,
7854 dt / m,
7855 self.dspark_draft_ns as f64 / self.dspark_exp.len().max(1) as f64 / 1e6
7856 );
7857 }
7858 }
7859
7860 fn forward_layers_upto(
7861 &mut self,
7862 hidden: &[f32],
7863 position: usize,
7864 task_mask: Option<&TaskMask>,
7865 upto: Option<usize>,
7866 ) -> Vec<f32> {
7867 if let Some(plan) = self.gpu_plan.clone() {
7873 if upto.is_none() && plan.len() > 1 {
7874 let mut h = hidden.to_vec();
7875 for &(dev, from, upto_incl) in plan.iter() {
7876 h = crate::gpu::with_device(dev, || {
7877 self.forward_layers_span(&h, position, task_mask, from, Some(upto_incl))
7878 });
7879 }
7880 return h;
7881 }
7882 }
7883 self.forward_layers_span(hidden, position, task_mask, 0, upto)
7884 }
7885
7886 pub fn set_gpu_plan(&mut self, devices: Option<&[usize]>) -> Result<(), String> {
7891 self.set_gpu_plan_at(devices, None)
7892 }
7893
7894 pub fn set_gpu_plan_at(
7898 &mut self,
7899 devices: Option<&[usize]>,
7900 at: Option<usize>,
7901 ) -> Result<(), String> {
7902 let Some(devs) = devices.filter(|d| d.len() > 1) else {
7903 self.gpu_plan = None;
7904 return Ok(());
7905 };
7906 self.split_supported()?;
7907 let n = self.num_layers;
7908 if devs.len() > n {
7909 return Err(format!("{} devices for {n} layers", devs.len()));
7910 }
7911 if let Some(k) = at {
7912 if k == 0 || k >= n {
7913 return Err(format!("split at {k}: the model has {n} layers"));
7914 }
7915 if devs.len() == 2 {
7916 self.gpu_plan = Some(std::sync::Arc::new(vec![
7917 (devs[0], 0, k - 1),
7918 (devs[1], k, n - 1),
7919 ]));
7920 return Ok(());
7921 }
7922 return Err(format!(
7923 "an explicit split point takes exactly 2 devices, got {}",
7924 devs.len()
7925 ));
7926 }
7927 let per = n.div_ceil(devs.len());
7928 let mut plan = Vec::with_capacity(devs.len());
7929 let mut from = 0usize;
7930 for &d in devs {
7931 if from >= n {
7932 break;
7933 }
7934 let upto = (from + per - 1).min(n - 1);
7935 plan.push((d, from, upto));
7936 from = upto + 1;
7937 }
7938 self.gpu_plan = Some(std::sync::Arc::new(plan));
7939 Ok(())
7940 }
7941
7942 pub fn gpu_plan(&self) -> Option<Vec<(usize, usize, usize)>> {
7944 self.gpu_plan.as_ref().map(|p| p.as_ref().clone())
7945 }
7946
7947 fn forward_layers_span(
7953 &mut self,
7954 hidden: &[f32],
7955 position: usize,
7956 task_mask: Option<&TaskMask>,
7957 from: usize,
7958 upto: Option<usize>,
7959 ) -> Vec<f32> {
7960 debug_assert!(from == 0 || (self.dsv4.is_none() && self.g3n.is_none()));
7961 if let Some(b) = &mut self.dsv4 {
7967 let _ = (task_mask, upto);
7968 let token_id = hidden.first().copied().unwrap_or(0.0) as u32;
7969 let (g, layers, cfg, st) = (&b.0, &b.1, b.2, &mut b.3);
7970 st.pos = position;
7971 let mut logits = Vec::new();
7972 crate::dsv4::forward_token(
7973 g,
7974 layers,
7975 &cfg,
7976 st,
7977 token_id,
7978 &self.inv_freq,
7979 self.pool.as_deref(),
7980 &mut logits,
7981 );
7982 self.graph_logits = Some(logits);
7983 self.dspark_probe(position, token_id);
7984 return vec![0.0; self.hidden_size];
7987 }
7988 if let Some(b) = &self.g3n {
7991 let _ = (task_mask, upto);
7992 return crate::g3n::g3n_forward(
7993 &b.0,
7994 &b.1,
7995 hidden,
7996 position,
7997 &mut self.kv_cache.layers,
7998 self.num_heads,
7999 self.num_kv_heads,
8000 self.head_dim,
8001 self.pool.as_deref(),
8002 );
8003 }
8004 let mut h = hidden.to_vec();
8005 let (nh, _nkv, _hd, hs, _rd, eps) = (
8008 self.num_heads,
8009 self.num_kv_heads,
8010 self.head_dim,
8011 self.hidden_size,
8012 self.rotary_dim,
8013 self.rms_eps,
8014 );
8015 let pool = self.pool.clone();
8016 let graph_env = std::env::var("CMF_GPU_WGPU_GRAPH").ok();
8028 let graph_on = match graph_env.as_deref() {
8029 Some("0") => false,
8030 Some("prefill") => false, Some(_) => true,
8032 None => crate::gpu::wgpu_graph_default(),
8038 };
8039 let graph_trusted =
8040 graph_env.is_some() || crate::gpu::wgpu_graph_default() || self.gdn_cfg.is_some();
8041 let race_eligible = graph_on
8042 && upto.is_none()
8043 && task_mask.is_none()
8044 && from == 0
8045 && !crate::gpu::graph_unsupported();
8046 let mut tail_start = 0usize;
8047 if race_eligible && crate::gpu::graph_race_use_graph(graph_trusted) {
8048 let t_graph = std::time::Instant::now();
8049 let mut lg = Vec::new();
8050 let mut gl = 0usize;
8051 let built = self.try_token_graph_wgpu(hidden, position, &mut lg, &mut gl);
8052 if built.is_none() && !self.o1_active() && self.attn_softcap == 0.0 {
8057 crate::gpu::graph_mark_unsupported();
8058 }
8059 graph_note(built.is_some());
8060 if let Some(hh) = built {
8061 let dur = t_graph.elapsed();
8062 if std::env::var("CMF_GRAPH_PROF").is_ok() {
8063 eprintln!("graph-call: {:.2} ms total", dur.as_secs_f64() * 1000.0);
8064 }
8065 if gl > 0 && gl < self.num_layers {
8066 h = hh;
8072 tail_start = gl;
8073 } else if graph_trusted || !crate::gpu::graph_race_first_token_hopeless(dur) {
8074 if !graph_trusted {
8075 crate::gpu::graph_race_record(true, dur);
8076 }
8077 if !lg.is_empty() {
8078 lg.resize(self.vocab_size, 0.0);
8081 if let Some(c) = self.final_softcap {
8082 for l in lg.iter_mut() {
8083 *l = c * (*l / c).tanh();
8084 }
8085 }
8086 self.graph_logits = Some(lg);
8087 }
8088 return hh;
8089 }
8090 }
8096 }
8097 let span = from > 0 || upto.is_some();
8121 if span && graph_on && task_mask.is_none() && graph_trusted {
8122 let upto_excl = upto.map_or(self.num_layers, |u| u + 1);
8123 let mut lg = Vec::new();
8124 let mut gl = 0usize;
8125 let span_res =
8126 self.try_token_graph_wgpu_span(hidden, position, &mut lg, from, upto_excl, &mut gl);
8127 graph_note(span_res.is_some() && gl == upto_excl - from);
8128 if std::env::var("CMF_GPU_DEBUG").is_ok() {
8129 static SEEN: std::sync::atomic::AtomicU32 = std::sync::atomic::AtomicU32::new(0);
8133 if SEEN.fetch_add(1, std::sync::atomic::Ordering::Relaxed) < 4 {
8134 eprintln!(
8135 "span graph: covered {gl} of {} layers [{from}..{upto_excl}) res={}",
8136 upto_excl - from,
8137 span_res.is_some()
8138 );
8139 }
8140 }
8141 if let Some(hh) = span_res {
8142 if gl == upto_excl - from {
8143 if !lg.is_empty() {
8144 lg.resize(self.vocab_size, 0.0);
8145 if let Some(c) = self.final_softcap {
8146 for l in lg.iter_mut() {
8147 *l = c * (*l / c).tanh();
8148 }
8149 }
8150 self.graph_logits = Some(lg);
8151 }
8152 crate::gpu::set_layer(-1);
8153 return hh;
8154 }
8155 h = hh;
8157 tail_start = from + gl;
8158 }
8159 }
8160 let t_race_cpu = (race_eligible && !graph_trusted).then(std::time::Instant::now);
8161
8162 #[cfg(target_os = "macos")]
8163 let mut gpu_skip_until = 0usize;
8164 for li in tail_start.max(from)..self.num_layers {
8165 crate::gpu::set_layer(li as i64); if let Some(u) = upto {
8167 if li > u {
8168 break;
8169 }
8170 }
8171 if let Some(mask) = task_mask {
8172 if !mask.layer_alive(li) {
8173 continue; }
8175 }
8176 #[cfg(target_os = "macos")]
8180 {
8181 if li < gpu_skip_until {
8182 continue;
8183 }
8184 if task_mask.is_none() {
8185 let end = self.q1_graph_gpu(li, upto, position, &mut h);
8186 if end > li {
8187 gpu_skip_until = end;
8188 if self.is_loop_end(end - 1) && end < self.num_layers {
8191 h = inference::rms_norm(
8192 &h,
8193 &self.weights.final_norm,
8194 self.rms_eps,
8195 self.norm_style,
8196 );
8197 }
8198 continue;
8199 }
8200 }
8201 }
8202
8203 let lw = &self.weights.layers[self.phys_layer(li)];
8204 if let Ok(tp) = std::env::var("CMF_TRACE_POS") {
8205 if tp.parse::<usize>().ok() == Some(position) {
8206 let n: f32 = h.iter().map(|x| x * x).sum::<f32>().sqrt();
8207 eprintln!(
8208 "TRACE pos {position} layer {li}: |h| = {n:.6} h0 {:.6} h1 {:.6}",
8209 h[0], h[1]
8210 );
8211 }
8212 }
8213 inference::rms_norm_into(
8216 &h,
8217 &lw.input_norm,
8218 self.rms_eps,
8219 self.norm_style,
8220 &mut self.ws.n1,
8221 );
8222
8223 let attn_out = match &lw.attn {
8224 AttnKind::Mla(w) => {
8225 let inv_freq_l = self.layer_inv_freq(li);
8226 let rs = self.layer_rope_scale(li);
8227 let eps = self.rms_eps;
8228 let pool = self.pool.clone();
8229 mla_attention(
8230 w,
8231 &self.ws.n1,
8232 &mut self.kv_cache.layers[li],
8233 position,
8234 &inv_freq_l,
8235 rs,
8236 eps,
8237 pool.as_deref(),
8238 )
8239 }
8240 AttnKind::Linear(w) => {
8241 let cfg = self.vmf_cfg.expect("linear layer without vmf_cfg");
8242 vmf_phase_forward(
8243 &self.ws.n1,
8244 w,
8245 &cfg,
8246 &mut self.kv_cache.layers[li].linear_state,
8247 self.pool.as_deref(),
8248 )
8249 }
8250 AttnKind::Kda(w) => {
8251 let cfg = self.kda_cfg.expect("kda layer without kda_cfg");
8252 crate::linear_core::kda_forward(
8253 &self.ws.n1,
8254 w,
8255 &cfg,
8256 &mut self.kv_cache.layers[li].linear_state,
8257 self.pool.as_deref(),
8258 )
8259 }
8260 AttnKind::LinearGdn(w) => {
8261 let cfg = self.gdn_cfg.expect("gdn layer without gdn_cfg");
8262 gdn_forward(
8263 &self.ws.n1,
8264 w,
8265 &cfg,
8266 &mut self.kv_cache.layers[li].linear_state,
8267 self.pool.as_deref(),
8268 )
8269 }
8270 AttnKind::ShortConv(w) => {
8271 let cfg = self
8272 .short_conv_cfg
8273 .expect("short-conv layer without short_conv_cfg");
8274 short_conv_forward(
8275 &self.ws.n1,
8276 w,
8277 &cfg,
8278 &mut self.kv_cache.layers[li].linear_state,
8279 self.pool.as_deref(),
8280 )
8281 }
8282 AttnKind::Full {
8283 wq,
8284 wk,
8285 wv,
8286 wo,
8287 q_norm,
8288 k_norm,
8289 output_gate,
8290 softplus_gate,
8291 bias,
8292 } if self.kv_cache.layers[li].o1_sealed() => {
8293 let inv_freq_l = self.layer_inv_freq(li);
8296 let (nkv_l, hd_l, rd_l) = self.layer_geom(li);
8297 let cfg = QwenAttnCfg {
8298 num_heads: self.layer_num_heads(li),
8299 num_kv_heads: nkv_l,
8300 head_dim: hd_l,
8301 hidden_size: hs,
8302 position,
8303 inv_freq: &inv_freq_l,
8304 rotary_dim: rd_l,
8305 scale: self.attn_scale,
8306 softcap: self.attn_softcap,
8307 window: None,
8308 v_norm: self.attn_v_norm,
8309 q_norm: q_norm.as_deref(),
8310 k_norm: k_norm.as_deref(),
8311 output_gate: *output_gate,
8312 softplus_gate: softplus_gate
8313 .as_ref()
8314 .map(|(gate, per_head)| (gate, *per_head)),
8315 rope_scale: self.layer_rope_scale(li),
8316 bias: bias
8317 .as_ref()
8318 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
8319 rms_eps: eps,
8320 norm_style: self.norm_style,
8321 pool: pool.as_deref(),
8322 };
8323 attention::qwen_attention_nystrom(
8324 &self.ws.n1,
8325 wq,
8326 wk,
8327 wv,
8328 wo,
8329 &mut self.kv_cache.layers[li],
8330 &cfg,
8331 )
8332 }
8333 AttnKind::Full {
8334 wq,
8335 wk,
8336 wv,
8337 wo,
8338 q_norm,
8339 k_norm,
8340 output_gate,
8341 softplus_gate,
8342 bias,
8343 } => 'attn: {
8344 if graph_on
8347 && !*output_gate
8348 && softplus_gate.is_none()
8349 && self.attention_heads_per_layer.is_none()
8350 && bias.is_none()
8351 && task_mask.is_none()
8352 {
8353 let inv_freq_l = self.layer_inv_freq(li);
8354 let (nkv_l, hd_l, rd_l) = self.layer_geom(li);
8355 let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
8356 if let (Some((gm, qi)), Some((_, ki)), Some((_, vi)), Some((_, oi))) = (
8357 wq.mapped_q1(),
8358 wk.mapped_q1(),
8359 wv.mapped_q1(),
8360 wo.mapped_q1(),
8361 ) {
8362 let gm = gm.clone();
8363 let mut out = vec![0f32; hs];
8364 let cache = &self.kv_cache.layers[li];
8365 if crate::gpu::attn_dropin(
8366 &gm,
8367 self.graph_kv_id,
8368 li,
8369 &self.ws.n1,
8370 qi,
8371 ki,
8372 vi,
8373 oi,
8374 q_norm.as_deref(),
8375 k_norm.as_deref(),
8376 &inv_freq_l,
8377 nh,
8378 nkv_l,
8379 hd_l,
8380 rd_l,
8381 hs,
8382 position,
8383 self.kv_cache.max_seq_len,
8384 gemma,
8385 eps as f32,
8386 cache.k_heads(),
8387 cache.v_heads(),
8388 &mut out,
8389 ) {
8390 break 'attn out;
8391 }
8392 }
8393 }
8394 let masked = task_mask
8395 .map(|m| m.head_flags(li, self.num_heads).iter().any(|&a| !a))
8396 .unwrap_or(false);
8397 let f32_view = (wq.as_f32(), wk.as_f32(), wv.as_f32(), wo.as_f32());
8398 match (masked, f32_view) {
8399 (true, (Some(q), Some(k), Some(v), Some(o))) => {
8402 let active_heads = task_mask.unwrap().head_flags(li, self.num_heads);
8403 attention::multi_head_attention(
8404 &self.ws.n1,
8405 q,
8406 k,
8407 v,
8408 o,
8409 &mut self.kv_cache.layers[li],
8410 self.num_heads,
8411 self.num_kv_heads,
8412 self.head_dim,
8413 self.hidden_size,
8414 position,
8415 &active_heads,
8416 &self.inv_freq,
8417 )
8418 }
8419 (masked, _) => {
8420 if masked {
8421 tracing::warn!(
8422 "layer {li}: head mask on quantized weights not \
8423 supported yet — executing dense"
8424 );
8425 }
8426 let inv_freq_l = self.layer_inv_freq(li);
8427 let (nkv_l, hd_l, rd_l) = self.layer_geom(li);
8428 let cfg = QwenAttnCfg {
8429 num_heads: self.layer_num_heads(li),
8430 num_kv_heads: nkv_l,
8431 head_dim: hd_l,
8432 hidden_size: hs,
8433 position,
8434 inv_freq: &inv_freq_l,
8435 rotary_dim: rd_l,
8436 scale: self.attn_scale,
8437 softcap: self.attn_softcap,
8438 window: self.layer_window(li),
8439 v_norm: self.attn_v_norm,
8440 q_norm: q_norm.as_deref(),
8441 k_norm: k_norm.as_deref(),
8442 output_gate: *output_gate,
8443 softplus_gate: softplus_gate
8444 .as_ref()
8445 .map(|(gate, per_head)| (gate, *per_head)),
8446 rope_scale: self.layer_rope_scale(li),
8447 bias: bias
8448 .as_ref()
8449 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
8450 rms_eps: eps,
8451 norm_style: self.norm_style,
8452 pool: pool.as_deref(),
8453 };
8454 attention::qwen_attention(
8455 &self.ws.n1,
8456 wq,
8457 wk,
8458 wv,
8459 wo,
8460 &mut self.kv_cache.layers[li],
8461 &cfg,
8462 )
8463 }
8464 }
8465 }
8466 };
8467 let attn_out = match &self.weights.layers[self.phys_layer(li)].attn_out_norm {
8470 Some(w) => inference::rms_norm(&attn_out, w, self.rms_eps, self.norm_style),
8471 None => attn_out,
8472 };
8473 let lw = &self.weights.layers[self.phys_layer(li)];
8474 inference::add_rmsnorm_fused_into(
8475 &mut h,
8476 &attn_out,
8477 &lw.post_norm,
8478 self.rms_eps,
8479 self.norm_style,
8480 &mut self.ws.p1,
8481 );
8482 let mut attn_out = attn_out;
8483 attention::recycle_buf(&mut attn_out);
8484 let post_normed = &self.ws.p1;
8485
8486 let ffn_masked = task_mask
8487 .map(|m| m.ffn_active_count(li) < self.intermediate_size)
8488 .unwrap_or(false);
8489 let ffn_out = match (ffn_masked, &lw.ffn) {
8501 (_, FfnKind::Dense(d)) if !d.segs.is_empty() => {
8505 let row = task_mask
8506 .and_then(|tm| tm.ffn_masks.get(li))
8507 .map(|v| v.as_slice());
8508 tube_ffn(d, post_normed, 1, self.pool.as_deref(), row)
8509 }
8510 (true, FfnKind::Dense(d)) => {
8511 let tm = task_mask.unwrap();
8512 let alive = tm.ffn_active_count(li);
8513 let deep = alive * 2 <= self.intermediate_size;
8514 if deep && d.down_proj.sparse_col_ok() {
8515 let active = tm.ffn_active_indices(li);
8516 sparse_ffn_quant(
8517 d,
8518 post_normed,
8519 &active,
8520 self.hidden_size,
8521 self.pool.as_deref(),
8522 )
8523 } else if deep
8524 && let (Some(g), Some(u), Some(dn)) = (
8525 d.gate_proj.as_f32(),
8526 d.up_proj.as_f32(),
8527 d.down_proj.as_f32(),
8528 )
8529 {
8530 let active = tm.ffn_active_indices(li);
8531 inference::sparse_ffn_forward(
8532 post_normed,
8533 g,
8534 u,
8535 dn,
8536 self.hidden_size,
8537 self.intermediate_size,
8538 &active,
8539 self.pool.as_deref(),
8540 )
8541 } else {
8542 let row = tm.ffn_masks.get(li).map(|v| v.as_slice());
8543 dense_ffn_batch(d, post_normed, 1, self.pool.as_deref(), row)
8544 }
8545 }
8546 (true, FfnKind::Moe(m)) => {
8547 let allowed = task_mask.and_then(|tm| tm.expert_flags(li, m.experts.len()));
8551 ffn_forward(
8552 &lw.ffn,
8553 post_normed,
8554 self.pool.as_deref(),
8555 allowed.as_deref(),
8556 )
8557 }
8558 (true, FfnKind::DenseMoe(dm)) => dense_moe_ffn(
8559 dm,
8560 post_normed,
8561 &h,
8562 self.rms_eps,
8563 self.norm_style,
8564 self.pool.as_deref(),
8565 ),
8566 (false, _) => match &lw.ffn {
8567 FfnKind::DenseMoe(dm) => dense_moe_ffn(
8568 dm,
8569 post_normed,
8570 &h,
8571 self.rms_eps,
8572 self.norm_style,
8573 self.pool.as_deref(),
8574 ),
8575 _ => {
8576 let allowed = match (&lw.ffn, task_mask) {
8577 (FfnKind::Moe(m), Some(tm)) => tm.expert_flags(li, m.experts.len()),
8578 _ => None,
8579 };
8580 ffn_forward(
8581 &lw.ffn,
8582 post_normed,
8583 self.pool.as_deref(),
8584 allowed.as_deref(),
8585 )
8586 }
8587 },
8588 };
8589 let ffn_out = match &self.weights.layers[self.phys_layer(li)].ffn_out_norm {
8590 Some(w) => inference::rms_norm(&ffn_out, w, self.rms_eps, self.norm_style),
8591 None => ffn_out,
8592 };
8593 for (i, &f) in ffn_out.iter().enumerate() {
8594 h[i] += f;
8595 }
8596 let mut ffn_out = ffn_out;
8597 attention::recycle_buf(&mut ffn_out);
8598
8599 if let Some(sc) = self.weights.layers[self.phys_layer(li)].layer_scale {
8601 for v in h.iter_mut() {
8602 *v *= sc;
8603 }
8604 }
8605
8606 if self.is_loop_end(li) && li + 1 < self.num_layers {
8609 h = inference::rms_norm(
8610 &h,
8611 &self.weights.final_norm,
8612 self.rms_eps,
8613 self.norm_style,
8614 );
8615 }
8616
8617 if self.dyn_phi_layer == Some(li) {
8621 self.update_dyn_phi(&h);
8622 }
8623 }
8624 crate::gpu::set_layer(-1); if let Some(t) = t_race_cpu {
8626 crate::gpu::graph_race_record(false, t.elapsed());
8627 }
8628
8629 h
8630 }
8631
8632 fn update_dyn_phi(&mut self, h: &[f32]) {
8635 const A: f32 = 0.2;
8636 if self.dyn_phi_ema.len() != h.len() {
8637 self.dyn_phi_ema = vec![0.0; h.len()];
8638 self.dyn_phi_seen = 0;
8639 }
8640 if self.dyn_phi_seen == 0 {
8641 self.dyn_phi_ema.copy_from_slice(h);
8642 } else {
8643 for (e, &v) in self.dyn_phi_ema.iter_mut().zip(h) {
8644 *e = (1.0 - A) * *e + A * v;
8645 }
8646 }
8647 self.dyn_phi_seen += 1;
8648 }
8649
8650 pub fn dyn_phi(&self) -> &[f32] {
8652 &self.dyn_phi_ema
8653 }
8654
8655 pub fn set_dyn_phi_layer(&mut self, layer: Option<usize>) {
8657 self.dyn_phi_layer = layer;
8658 self.dyn_phi_ema.clear();
8659 self.dyn_phi_seen = 0;
8660 }
8661
8662 pub fn dynamic_skills(&self) -> Vec<(usize, String, usize)> {
8664 let Some(model) = &self.model else {
8665 return Vec::new();
8666 };
8667 model
8668 .header
8669 .skills
8670 .iter()
8671 .enumerate()
8672 .filter_map(|(i, sk)| {
8673 let ok = matches!(self.dyn_skill_layers.get(i), Some(Some(_)));
8674 let sel = sk.selection.as_ref()?;
8675 (ok).then(|| (i, sk.id.clone(), sel.phi_layer))
8676 })
8677 .collect()
8678 }
8679
8680 pub fn active_skill(&self) -> Option<usize> {
8682 self.dyn_active
8683 }
8684
8685 pub fn enable_dynamic_routing(&mut self) -> usize {
8690 use crate::swarm::{DynRouter, RoutableSkill};
8691 let Some(model) = self.model.clone() else {
8692 return 0;
8693 };
8694 if self.dyn_blend_loaded {
8697 tracing::warn!("dynamic routing unavailable on a blend-loaded pipeline");
8698 return 0;
8699 }
8700 if let Some(a) = self.dyn_active {
8704 if !matches!(self.dyn_skill_layers.get(a), Some(Some(_))) {
8705 tracing::warn!("loaded skill is not FFN-eligible — dynamic routing unavailable");
8706 return 0;
8707 }
8708 }
8709 let hidden = self.hidden_size;
8710 let mut skills = Vec::new();
8711 for (idx, id, _phi) in self.dynamic_skills() {
8712 if let Some(sel) = model.header.skills[idx].selection.as_ref() {
8713 if let Some(rs) = RoutableSkill::from_descriptor(idx, id, sel, hidden) {
8714 skills.push(rs);
8715 }
8716 }
8717 }
8718 if skills.is_empty() {
8719 return 0;
8720 }
8721 let phi = skills[0].phi_layer;
8723 if skills.iter().any(|s| s.phi_layer != phi) {
8724 tracing::warn!("routable skills disagree on phi_layer; using {phi}");
8725 }
8726 let n = skills.len();
8727 self.set_dyn_phi_layer(Some(phi));
8728 self.dyn_router = Some(DynRouter::new(skills));
8729 n
8730 }
8731
8732 pub fn route_switches(&self) -> Vec<(usize, Option<String>, Option<String>)> {
8734 self.dyn_router
8735 .as_ref()
8736 .map(|r| r.switches.clone())
8737 .unwrap_or_default()
8738 }
8739
8740 fn lm_head_forward(&self, hidden: &[f32]) -> Vec<f32> {
8743 let rows = self.weights.lm_head.rows();
8744 let mut logits = attention::take_buf(rows.min(self.vocab_size));
8745 self.weights
8746 .lm_head
8747 .matvec(hidden, &mut logits, self.pool.as_deref());
8748 logits.resize(self.vocab_size, 0.0);
8749 if let Some(m) = self.logit_multiplier {
8750 for l in logits.iter_mut() {
8751 *l *= m;
8752 }
8753 }
8754 if let Some(c) = self.final_softcap {
8755 for l in logits.iter_mut() {
8756 *l = c * (*l / c).tanh();
8757 }
8758 }
8759 if let Some(cm) = self.head_clusters.as_ref() {
8760 self.hierarchical_head_logprobs(hidden, cm, &mut logits);
8761 }
8762 logits
8763 }
8764
8765 fn hierarchical_head_logprobs(&self, hidden: &[f32], cm: &[f32], logits: &mut [f32]) {
8768 let h = hidden.len();
8769 let ncl = cm.len() / h.max(1);
8770 if ncl == 0 || logits.len() % ncl != 0 {
8771 return;
8772 }
8773 let cs = logits.len() / ncl;
8774 let mut lc = vec![0.0f32; ncl];
8776 for c in 0..ncl {
8777 let row = &cm[c * h..(c + 1) * h];
8778 let mut s = 0.0f32;
8779 for j in 0..h {
8780 s += row[j] * hidden[j];
8781 }
8782 lc[c] = s;
8783 }
8784 let mx = lc.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
8785 let lse: f32 = mx + lc.iter().map(|v| (v - mx).exp()).sum::<f32>().ln();
8786 for c in 0..ncl {
8787 let blk = &mut logits[c * cs..(c + 1) * cs];
8788 let bm = blk.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
8789 let bl: f32 = bm + blk.iter().map(|v| (v - bm).exp()).sum::<f32>().ln();
8790 let add = lc[c] - lse - bl;
8791 for v in blk.iter_mut() {
8792 *v += add;
8793 }
8794 }
8795 }
8796
8797 pub fn prefill_next_logits(&mut self, ids: &[u32], task_mask: Option<&TaskMask>) -> Vec<f32> {
8802 self.kv_cache.clear();
8803 self.kv_history.clear();
8804 let mut hidden = vec![0.0f32; self.hidden_size];
8805 for (pos, &id) in ids.iter().enumerate() {
8806 let emb = self.embed_single(id);
8807 hidden = self.forward_layers(&emb, pos, task_mask);
8808 }
8809 inference::rms_norm_into(
8810 &hidden,
8811 &self.weights.final_norm,
8812 self.rms_eps,
8813 self.norm_style,
8814 &mut self.ws.n1,
8815 );
8816 self.lm_head_forward(&self.ws.n1)
8817 }
8818}
8819
8820pub fn create_test_pipeline(
8822 hidden_size: usize,
8823 intermediate_size: usize,
8824 num_heads: usize,
8825 num_kv_heads: usize,
8826 head_dim: usize,
8827 num_layers: usize,
8828 vocab_size: usize,
8829) -> Pipeline {
8830 let synth = |n: usize, salt: usize| -> Vec<f32> {
8833 (0..n)
8834 .map(|i| (((i * 31 + salt * 17 + 7) % 97) as f32 / 97.0 - 0.5) * 0.2)
8835 .collect()
8836 };
8837 let qt = |rows: usize, cols: usize, salt: usize| -> QTensor {
8838 QTensor::from_f32(synth(rows * cols, salt), rows, cols)
8839 };
8840 let layer_weights: Vec<LayerWeights> = (0..num_layers)
8841 .map(|li| LayerWeights {
8842 input_norm: vec![1.0; hidden_size],
8843 post_norm: vec![1.0; hidden_size],
8844 attn_out_norm: None,
8845 ffn_out_norm: None,
8846 layer_scale: None,
8847 ffn: FfnKind::Dense(DenseFfn {
8848 gate_proj: qt(intermediate_size, hidden_size, li * 10 + 5),
8849 up_proj: qt(intermediate_size, hidden_size, li * 10 + 6),
8850 down_proj: qt(hidden_size, intermediate_size, li * 10 + 7),
8851 act: Act::Silu,
8852 down_t: None,
8853 segs: Vec::new(),
8854 }),
8855 attn: AttnKind::Full {
8856 bias: None,
8857 wq: qt(num_heads * head_dim, hidden_size, li * 10 + 1),
8858 wk: qt(num_kv_heads * head_dim, hidden_size, li * 10 + 2),
8859 wv: qt(num_kv_heads * head_dim, hidden_size, li * 10 + 3),
8860 wo: qt(hidden_size, num_heads * head_dim, li * 10 + 4),
8861 q_norm: None,
8862 k_norm: None,
8863 output_gate: false,
8864 softplus_gate: None,
8865 },
8866 })
8867 .collect();
8868
8869 Pipeline::new(
8870 Tokenizer::byte_level(),
8871 PipelineWeights {
8872 embed_tokens: qt(vocab_size, hidden_size, 100),
8873 layers: layer_weights,
8874 lm_head: qt(vocab_size, hidden_size, 200),
8875 final_norm: vec![1.0; hidden_size],
8876 },
8877 hidden_size,
8878 intermediate_size,
8879 num_heads,
8880 num_kv_heads,
8881 head_dim,
8882 num_layers,
8883 num_layers, false, vocab_size,
8886 1e-6,
8887 10_000.0,
8888 NormStyle::Qwen,
8889 4096,
8890 SamplerConfig {
8891 seed: Some(42),
8892 ..Default::default()
8893 },
8894 )
8895}
8896
8897#[inline]
8902fn mask_bit(row: &[u8], j: usize) -> bool {
8903 (row.get(j >> 3).copied().unwrap_or(0) >> (j & 7)) & 1 != 0
8904}
8905
8906fn mask_gain() -> f32 {
8917 static G: std::sync::OnceLock<f32> = std::sync::OnceLock::new();
8918 *G.get_or_init(|| {
8919 std::env::var("CMF_FFN_MASK_GAIN")
8920 .ok()
8921 .and_then(|v| v.parse().ok())
8922 .unwrap_or(1.0)
8923 })
8924}
8925
8926fn zero_masked_cols(g: &mut [f32], rows: usize, inter: usize, row: &[u8]) {
8927 let fill = meanfill().and_then(|(i, v)| {
8930 let li = crate::gpu::cur_layer();
8931 (*i == inter && li >= 0).then(|| &v[li as usize * inter..(li as usize + 1) * inter])
8932 });
8933 for r in 0..rows {
8934 let base = r * inter;
8935 for (bi, &byte) in row.iter().enumerate() {
8936 if byte == 0xFF {
8937 continue;
8938 }
8939 let j0 = bi * 8;
8940 for bit in 0..8 {
8941 let j = j0 + bit;
8942 if j < inter && byte & (1 << bit) == 0 {
8943 g[base + j] = fill.map_or(0.0, |f| f[j]);
8944 }
8945 }
8946 }
8947 }
8948 let gain = mask_gain();
8949 if gain != 1.0 {
8950 for v in g[..rows * inter].iter_mut() {
8951 *v *= gain;
8952 }
8953 }
8954}
8955
8956#[inline]
8958fn tube_bit(row: Option<&[u8]>, i: usize) -> bool {
8959 row.is_none_or(|r| mask_bit(r, i))
8960}
8961
8962fn all_bits_on(row: &[u8], n: usize) -> bool {
8965 (0..n).all(|i| mask_bit(row, i))
8966}
8967
8968fn tube_topk() -> usize {
8976 static K: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
8977 *K.get_or_init(|| {
8978 std::env::var("CMF_TUBE_TOPK")
8979 .ok()
8980 .and_then(|v| v.parse().ok())
8981 .unwrap_or(0)
8982 })
8983}
8984
8985fn tube_score_oracle() -> bool {
8986 static O: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
8987 *O.get_or_init(|| {
8988 std::env::var("CMF_TUBE_SCORE").is_ok_and(|v| v == "oracle")
8989 })
8990}
8991
8992fn tube_ffn_routed(
8999 d: &DenseFfn,
9000 xs: &[f32],
9001 b: usize,
9002 pool: Option<&Pool>,
9003 mask_row: Option<&[u8]>,
9004 k: usize,
9005) -> Vec<f32> {
9006 let hidden = d.down_proj.rows();
9007 let core = d.gate_proj.rows();
9008 let core_full = mask_row.is_none_or(|r| all_bits_on(r, core));
9009 let mut out = match (b, core_full, mask_row) {
9010 (1, true, _) => dense_ffn(d, xs, pool),
9011 (1, false, Some(row)) => dense_ffn_masked(d, xs, pool, row),
9012 (_, true, _) => dense_ffn_batch(d, xs, b, pool, None),
9013 (_, false, row) => dense_ffn_batch(d, xs, b, pool, row),
9014 };
9015 let cand: Vec<usize> = (0..d.segs.len())
9016 .filter(|&i| tube_bit(mask_row, d.segs[i].start))
9017 .collect();
9018 if cand.is_empty() {
9019 return out;
9020 }
9021 let oracle = tube_score_oracle();
9025 let mut acts: Vec<Vec<f32>> = Vec::with_capacity(cand.len());
9026 let mut scores = vec![0f32; b * cand.len()];
9027 for (ci, &i) in cand.iter().enumerate() {
9028 let seg = &d.segs[i];
9029 let w = seg.width;
9030 let mut g = vec![0.0f32; b * w];
9031 if b == 1 {
9032 seg.gate.matvec(xs, &mut g, pool);
9033 } else {
9034 seg.gate.matmat(xs, b, &mut g, pool);
9035 }
9036 for v in g.iter_mut() {
9037 *v = Act::Silu.combine(*v, 1.0);
9038 }
9039 if !oracle {
9040 for t in 0..b {
9041 scores[t * cand.len() + ci] =
9042 g[t * w..(t + 1) * w].iter().map(|v| v * v).sum::<f32>();
9043 }
9044 }
9045 if oracle || b > 1 {
9046 let mut u = vec![0.0f32; b * w];
9047 if b == 1 {
9048 seg.up.matvec(xs, &mut u, pool);
9049 } else {
9050 seg.up.matmat(xs, b, &mut u, pool);
9051 }
9052 for (a, &v) in g.iter_mut().zip(u.iter()) {
9053 *a *= v;
9054 }
9055 if oracle {
9056 for t in 0..b {
9057 scores[t * cand.len() + ci] =
9058 g[t * w..(t + 1) * w].iter().map(|v| v * v).sum::<f32>();
9059 }
9060 }
9061 }
9062 acts.push(g);
9063 }
9064 let keep = k.min(cand.len());
9066 let mut scratch: Vec<f32> = Vec::new();
9067 for t in 0..b {
9068 let mut sc: Vec<(f32, usize)> = (0..cand.len())
9069 .map(|ci| (scores[t * cand.len() + ci], ci))
9070 .collect();
9071 sc.sort_unstable_by(|x, y| y.0.total_cmp(&x.0));
9072 let mut alive = vec![false; cand.len()];
9073 for &(_, ci) in sc.iter().take(keep) {
9074 alive[ci] = true;
9075 }
9076 if b > 1 {
9077 for (ci, a) in acts.iter_mut().enumerate() {
9078 if !alive[ci] {
9079 let w = d.segs[cand[ci]].width;
9080 a[t * w..(t + 1) * w].fill(0.0);
9081 }
9082 }
9083 } else {
9084 for (ci, &i) in cand.iter().enumerate() {
9088 if !alive[ci] {
9089 continue;
9090 }
9091 let seg = &d.segs[i];
9092 let w = seg.width;
9093 let g = &mut acts[ci];
9094 if !tube_score_oracle() {
9095 scratch.clear();
9096 scratch.resize(w, 0.0);
9097 seg.up.matvec(xs, &mut scratch, pool);
9098 for (a, &v) in g.iter_mut().zip(scratch.iter()) {
9099 *a *= v;
9100 }
9101 }
9102 let mut acc = vec![0.0f32; hidden];
9103 seg.down.matvec(g, &mut acc, pool);
9104 for (o, a) in out.iter_mut().zip(&acc) {
9105 *o += *a;
9106 }
9107 }
9108 }
9109 }
9110 if b > 1 {
9111 for (ci, &i) in cand.iter().enumerate() {
9112 let seg = &d.segs[i];
9113 let mut acc = vec![0.0f32; b * hidden];
9114 seg.down.matmat(&acts[ci], b, &mut acc, pool);
9115 for (o, a) in out.iter_mut().zip(&acc) {
9116 *o += *a;
9117 }
9118 }
9119 }
9120 out
9121}
9122
9123fn tube_ffn(
9129 d: &DenseFfn,
9130 xs: &[f32],
9131 b: usize,
9132 pool: Option<&Pool>,
9133 mask_row: Option<&[u8]>,
9134) -> Vec<f32> {
9135 if tube_topk() > 0 {
9136 return tube_ffn_routed(d, xs, b, pool, mask_row, tube_topk());
9137 }
9138 let hidden = d.down_proj.rows();
9139 let core = d.gate_proj.rows();
9140 let core_full = mask_row.is_none_or(|r| all_bits_on(r, core));
9141 let mut out = match (b, core_full, mask_row) {
9142 (1, true, _) => dense_ffn(d, xs, pool),
9143 (1, false, Some(row)) => dense_ffn_masked(d, xs, pool, row),
9144 (_, true, _) => dense_ffn_batch(d, xs, b, pool, None),
9145 (_, false, row) => dense_ffn_batch(d, xs, b, pool, row),
9146 };
9147 TUBE_SCRATCH.with(|sc| {
9148 let mut sc = sc.borrow_mut();
9149 let [g, u, acc] = &mut *sc;
9150 for seg in &d.segs {
9151 if !tube_bit(mask_row, seg.start) {
9152 continue;
9153 }
9154 let w = seg.width;
9155 g.resize(b * w, 0.0);
9156 if b == 1 && d.act == Act::Silu && QTensor::matvec_silu_mul(&seg.gate, &seg.up, xs, g, pool)
9157 {
9158 } else {
9160 u.resize(b * w, 0.0);
9161 if b == 1 {
9162 QTensor::matvec_many([&seg.gate, &seg.up], xs, [g, u], pool);
9163 } else {
9164 seg.gate.matmat(xs, b, g, pool);
9165 seg.up.matmat(xs, b, u, pool);
9166 }
9167 for i in 0..b * w {
9168 g[i] = d.act.combine(g[i], u[i]);
9169 }
9170 }
9171 acc.resize(b * hidden, 0.0);
9172 acc.fill(0.0);
9173 if b == 1 {
9174 seg.down.matvec(g, acc, pool);
9175 } else {
9176 seg.down.matmat(g, b, acc, pool);
9177 }
9178 for (o, a) in out.iter_mut().zip(acc.iter()) {
9179 *o += *a;
9180 }
9181 }
9182 out
9183 })
9184}
9185
9186thread_local! {
9187 static TUBE_SCRATCH: std::cell::RefCell<[Vec<f32>; 3]> =
9191 const { std::cell::RefCell::new([Vec::new(), Vec::new(), Vec::new()]) };
9192}
9193
9194fn dense_ffn_batch(
9195 d: &DenseFfn,
9196 xs: &[f32],
9197 b: usize,
9198 pool: Option<&Pool>,
9199 mask_row: Option<&[u8]>,
9200) -> Vec<f32> {
9201 let inter = d.gate_proj.rows();
9202 let hidden = d.down_proj.rows();
9203 if mask_row.is_none()
9211 && d.act == Act::Silu
9212 && b >= 32
9213 && crate::gpu::enabled_here()
9214 && !crate::gpu::mm_killed()
9215 && refit_dir().is_none()
9220 && !ffn_probe_active()
9225 {
9226 if let (Some((model, w1)), Some((_, w3)), Some((_, w2))) = (
9227 d.gate_proj.mapped_q4t(),
9228 d.up_proj.mapped_q4t(),
9229 d.down_proj.mapped_q4t(),
9230 ) {
9231 let mut out = vec![0.0f32; b * hidden];
9232 if crate::gpu::q4t_ffn(model, w1, w3, w2, xs, b, hidden, inter, &mut out) {
9233 return out;
9234 }
9235 }
9236 if let (Some((model, w1)), Some((_, w3)), Some((_, w2))) = (
9241 d.gate_proj.mapped_q4tp(),
9242 d.up_proj.mapped_q4tp(),
9243 d.down_proj.mapped_q4tp(),
9244 ) {
9245 let mut out = vec![0.0f32; b * hidden];
9246 if crate::gpu::q4tp_ffn(model, w1, w3, w2, xs, b, hidden, inter, &mut out) {
9247 return out;
9248 }
9249 }
9250 }
9251 let mut g = vec![0.0f32; b * inter];
9252 d.gate_proj.matmat(xs, b, &mut g, pool);
9253 let mut u = vec![0.0f32; b * inter];
9254 d.up_proj.matmat(xs, b, &mut u, pool);
9255 if gate_topk() > 0 && d.act == Act::Silu {
9256 for t in 0..b {
9257 let row = &mut g[t * inter..(t + 1) * inter];
9258 for v in row.iter_mut() {
9259 *v = Act::Silu.combine(*v, 1.0);
9260 }
9261 keep_top_k(row, gate_topk());
9262 }
9263 for i in 0..b * inter {
9264 g[i] *= u[i];
9265 }
9266 } else {
9267 for i in 0..b * inter {
9268 g[i] = d.act.combine(g[i], u[i]);
9269 }
9270 }
9271 if let Some(row) = mask_row {
9272 zero_masked_cols(&mut g, b, inter, row);
9273 }
9274 if oracle_topk() > 0 {
9275 for t in 0..b {
9276 keep_top_k(&mut g[t * inter..(t + 1) * inter], oracle_topk());
9277 }
9278 }
9279 let mut out = vec![0.0f32; b * hidden];
9280 d.down_proj.matmat(&g, b, &mut out, pool);
9281 if refit_dir().is_some() {
9282 let li = crate::gpu::cur_layer();
9283 if li >= 0 {
9284 refit_accumulate(li as usize, &g, b, inter, &out, hidden, pool);
9285 }
9286 }
9287 FFN_PROBE.with(|pr| {
9291 if let Some(acc) = pr.borrow_mut().as_mut() {
9292 let li = crate::gpu::cur_layer();
9293 if li < 0 {
9294 return;
9295 }
9296 let Some(row) = acc.get_mut(li as usize) else {
9297 return;
9298 };
9299 let sq = probe_sq();
9300 for t in 0..b {
9301 for (a, &v) in row.iter_mut().zip(&g[t * inter..(t + 1) * inter]) {
9302 *a += if sq { (v as f64) * (v as f64) } else { (v as f64).abs() };
9303 }
9304 }
9305 }
9306 });
9307 out
9308}
9309
9310fn accumulate_act(m: &MoeFfn, xs: &[f32], b: usize) {
9315 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
9316 static DUMP: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
9317 let on = *ON.get_or_init(|| std::env::var("CMF_RMS_TRACE").is_ok());
9318 let dump = *DUMP.get_or_init(|| std::env::var("CMF_ACT_DUMP").is_ok());
9319 if (!on && !dump) || b == 0 {
9320 return;
9321 }
9322 let hidden = xs.len() / b;
9323 if on {
9324 let mut acc = m.act_sq.borrow_mut();
9325 if acc.len() < hidden {
9326 acc.resize(hidden, 0.0);
9327 }
9328 for t in 0..b {
9329 let row = &xs[t * hidden..(t + 1) * hidden];
9330 for (a, &v) in acc.iter_mut().zip(row) {
9331 *a += (v as f64) * (v as f64);
9332 }
9333 }
9334 }
9335 if dump {
9336 let cap: usize = std::env::var("CMF_ACT_DUMP_ROWS")
9339 .ok()
9340 .and_then(|v| v.parse().ok())
9341 .unwrap_or(4096);
9342 let mut rows = m.act_rows.borrow_mut();
9343 if rows.len() < cap * hidden {
9344 let take = b.min((cap * hidden - rows.len()) / hidden.max(1));
9345 rows.extend_from_slice(&xs[..take * hidden]);
9346 }
9347 }
9348}
9349
9350#[derive(Clone, Copy)]
9353struct SendVecs(*mut Vec<f32>);
9354unsafe impl Send for SendVecs {}
9355unsafe impl Sync for SendVecs {}
9356impl SendVecs {
9357 #[inline]
9358 fn at(self, i: usize) -> *mut Vec<f32> {
9359 unsafe { self.0.add(i) }
9360 }
9361}
9362
9363fn moe_ffn_batch(
9364 m: &MoeFfn,
9365 xs: &[f32],
9366 b: usize,
9367 hidden: usize,
9368 pool: Option<&Pool>,
9369 allowed: Option<&[bool]>,
9370) -> Vec<f32> {
9371 accumulate_act(m, xs, b);
9372 let ne = m.experts.len();
9373 let mut logits = vec![0.0f32; b * ne];
9374 match &m.resonance {
9375 Some(r) => {
9376 let hdim = xs.len() / b.max(1);
9377 for bi in 0..b {
9378 r.scores(&xs[bi * hdim..(bi + 1) * hdim], &mut logits[bi * ne..(bi + 1) * ne]);
9379 }
9380 }
9381 None => m.router.matmat(xs, b, &mut logits, pool),
9382 }
9383
9384 let mut assign: Vec<Vec<(usize, f32)>> = vec![Vec::new(); ne];
9387 {
9388 let mut st = m.stats.borrow_mut();
9389 if st.len() < ne {
9390 st.resize(ne, 0);
9391 }
9392 for bi in 0..b {
9393 let (idx, p, wsum) = moe_route(&logits[bi * ne..(bi + 1) * ne], m, allowed);
9394 for &e in &idx {
9395 st[e] += 1;
9396 assign[e].push((bi, p[e] / wsum));
9397 }
9398 }
9399 }
9400
9401 let mut out = vec![0.0f32; b * hidden];
9402 let cols = m.experts[0].gate_proj.cols();
9403 let run_expert = |d: &DenseFfn, list: &[(usize, f32)], out: &mut [f32]| {
9404 let sb = list.len();
9405 let mut sub = vec![0.0f32; sb * cols];
9406 for (k, &(bi, _)) in list.iter().enumerate() {
9407 sub[k * cols..(k + 1) * cols].copy_from_slice(&xs[bi * cols..(bi + 1) * cols]);
9408 }
9409 let eo = dense_ffn_batch(d, &sub, sb, pool, None);
9410 for (k, &(bi, w)) in list.iter().enumerate() {
9411 for i in 0..hidden {
9412 out[bi * hidden + i] += w * eo[k * hidden + i];
9413 }
9414 }
9415 };
9416 let active: Vec<usize> = (0..ne).filter(|&e| !assign[e].is_empty()).collect();
9422 if pool.is_some() && active.len() >= 8 {
9423 let mut panels: Vec<Vec<f32>> = vec![Vec::new(); active.len()];
9424 {
9425 let panel_ptr = SendVecs(panels.as_mut_ptr());
9426 let experts = &m.experts;
9429 let (active_r, assign_r) = (&active, &assign);
9430 let run = |start: usize, end: usize| {
9431 for ai in start..end {
9432 let e = active_r[ai];
9433 let list = &assign_r[e];
9434 let sb = list.len();
9435 let mut sub = vec![0.0f32; sb * cols];
9436 for (k, &(bi, _)) in list.iter().enumerate() {
9437 sub[k * cols..(k + 1) * cols]
9438 .copy_from_slice(&xs[bi * cols..(bi + 1) * cols]);
9439 }
9440 unsafe {
9442 *panel_ptr.at(ai) = dense_ffn_batch(&experts[e], &sub, sb, None, None);
9443 }
9444 }
9445 };
9446 match pool {
9447 Some(p) => p.run_rows(active.len(), &run),
9448 None => run(0, active.len()),
9449 }
9450 }
9451 for (ai, &e) in active.iter().enumerate() {
9452 for (k, &(bi, w)) in assign[e].iter().enumerate() {
9453 let eo = &panels[ai][k * hidden..(k + 1) * hidden];
9454 for i in 0..hidden {
9455 out[bi * hidden + i] += w * eo[i];
9456 }
9457 }
9458 }
9459 } else {
9460 for &e in &active {
9461 run_expert(&m.experts[e], &assign[e], &mut out);
9462 }
9463 }
9464 if let Some((se, gate)) = &m.shared {
9465 let all: Vec<(usize, f32)> = if let Some(gate) = gate {
9466 let mut gl = vec![0.0f32; b];
9467 gate.matmat(xs, b, &mut gl, pool);
9468 (0..b)
9469 .map(|bi| (bi, 1.0 / (1.0 + (-gl[bi]).exp())))
9470 .collect()
9471 } else {
9472 (0..b).map(|bi| (bi, 1.0)).collect()
9473 };
9474 run_expert(se, &all, &mut out);
9475 }
9476 out
9477}
9478
9479thread_local! {
9480 static FFN_SCRATCH: std::cell::RefCell<[Vec<f32>; 4]> =
9484 const { std::cell::RefCell::new([Vec::new(), Vec::new(), Vec::new(), Vec::new()]) };
9485}
9486
9487fn dense_ffn(d: &DenseFfn, x: &[f32], pool: Option<&Pool>) -> Vec<f32> {
9489 if gate_topk() > 0
9492 && let Some(out) = dense_ffn_dynamic(d, x, pool, gate_topk())
9493 {
9494 return out;
9495 }
9496 if crate::gpu::enabled_here()
9507 && (d.gate_proj.rows() >= crate::gpu::min_rows() || d.gate_proj.is_q1())
9508 {
9509 let arm = if d.gate_proj.is_q1() && crate::gpu::q1_force() {
9510 crate::gpu::ProbeArm::Gpu
9511 } else {
9512 crate::gpu::probe_arm(crate::gpu::OpClass::Ffn)
9513 };
9514 match arm {
9515 crate::gpu::ProbeArm::Gpu => {
9516 let t0 = std::time::Instant::now();
9517 if let Some(out) = dense_ffn_gpu(d, x, pool) {
9518 crate::gpu::probe_record(crate::gpu::OpClass::Ffn, true, t0.elapsed());
9519 return out;
9520 }
9521 }
9522 crate::gpu::ProbeArm::CpuTimed => {
9523 let t0 = std::time::Instant::now();
9524 let out = crate::gpu::cpu_scope(|| dense_ffn_cpu(d, x, pool));
9525 crate::gpu::probe_record(crate::gpu::OpClass::Ffn, false, t0.elapsed());
9526 return out;
9527 }
9528 crate::gpu::ProbeArm::Cpu => {
9529 return crate::gpu::cpu_scope(|| dense_ffn_cpu(d, x, pool));
9530 }
9531 }
9532 }
9533 dense_ffn_cpu(d, x, pool)
9534}
9535
9536fn dense_ffn_cpu(d: &DenseFfn, x: &[f32], pool: Option<&Pool>) -> Vec<f32> {
9538 let inter = d.gate_proj.rows();
9539 FFN_SCRATCH.with(|s| {
9540 let mut s = s.borrow_mut();
9541 let [g, u, ..] = &mut *s;
9542 g.resize(inter, 0.0);
9543 if gate_topk() > 0 {
9546 u.resize(inter, 0.0);
9550 QTensor::matvec_many([&d.gate_proj, &d.up_proj], x, [g, u], pool);
9551 for i in 0..inter {
9552 g[i] = Act::Silu.combine(g[i], 1.0);
9553 }
9554 keep_top_k(g, gate_topk());
9555 for i in 0..inter {
9556 g[i] *= u[i];
9557 }
9558 } else if d.act == Act::Silu && QTensor::matvec_silu_mul(&d.gate_proj, &d.up_proj, x, g, pool) {
9559 } else {
9561 u.resize(inter, 0.0);
9562 QTensor::matvec_many([&d.gate_proj, &d.up_proj], x, [g, u], pool);
9564 for i in 0..inter {
9565 g[i] = d.act.combine(g[i], u[i]);
9566 }
9567 }
9568 FFN_PROBE.with(|pr| {
9576 if let Some(acc) = pr.borrow_mut().as_mut() {
9577 let li = crate::gpu::cur_layer();
9578 if li >= 0 {
9579 if let Some(row) = acc.get_mut(li as usize) {
9580 match probe_topk() {
9581 0 if probe_sq() => {
9582 for (a, &v) in row.iter_mut().zip(g.iter()) {
9583 *a += (v as f64) * (v as f64);
9584 }
9585 }
9586 0 if probe_signed() => {
9587 for (a, &v) in row.iter_mut().zip(g.iter()) {
9588 *a += v as f64;
9589 }
9590 }
9591 0 => {
9592 for (a, &v) in row.iter_mut().zip(g.iter()) {
9593 *a += (v as f64).abs();
9594 }
9595 }
9596 k => {
9597 let n = g.len();
9598 let k = k.min(n);
9599 let mut mag: Vec<f32> = g.iter().map(|v| v.abs()).collect();
9600 let (_, kth, _) = mag.select_nth_unstable_by(k - 1, |a, b| {
9601 b.partial_cmp(a).unwrap_or(std::cmp::Ordering::Equal)
9602 });
9603 let thr = *kth;
9604 for (a, &v) in row.iter_mut().zip(g.iter()) {
9605 if v.abs() >= thr {
9606 *a += 1.0;
9607 }
9608 }
9609 }
9610 }
9611 }
9612 }
9613 }
9614 });
9615 if oracle_topk() > 0 {
9616 keep_top_k(g, oracle_topk());
9617 }
9618 {
9619 let li = crate::gpu::cur_layer();
9620 if li >= 0 {
9621 adump_row(li as usize, g);
9622 }
9623 }
9624 let mut out = attention::take_buf(d.down_proj.rows());
9625 d.down_proj.matvec(g, &mut out, pool);
9626 out
9627 })
9628}
9629
9630pub struct RefitAcc {
9643 pub support: Vec<u32>,
9644 pub gss: Vec<f32>,
9645 pub ya: Vec<f32>,
9646 pub hidden: usize,
9647 pub tokens: u64,
9648 pub buf_g: Vec<f32>,
9654 pub buf_o: Vec<f32>,
9655 pub buf_t: usize,
9656}
9657
9658type RefitState = (
9662 std::collections::HashMap<usize, RefitAcc>,
9663 Vec<f32>,
9664);
9665
9666static REFIT: std::sync::OnceLock<Option<(String, std::sync::Mutex<RefitState>)>> =
9667 std::sync::OnceLock::new();
9668
9669fn ffn_probe_active() -> bool {
9672 FFN_PROBE.with(|p| p.borrow().is_some())
9673}
9674
9675fn refit_dir() -> Option<&'static (String, std::sync::Mutex<RefitState>)> {
9676 REFIT
9677 .get_or_init(|| {
9678 std::env::var("CMF_FFN_REFIT").ok().map(|d| {
9679 (
9680 d,
9681 std::sync::Mutex::new((std::collections::HashMap::new(), Vec::new())),
9682 )
9683 })
9684 })
9685 .as_ref()
9686}
9687
9688fn refit_accumulate(
9690 li: usize,
9691 g: &[f32],
9692 b: usize,
9693 inter: usize,
9694 out: &[f32],
9695 hidden: usize,
9696 pool: Option<&Pool>,
9697) {
9698 let Some((dir, map)) = refit_dir() else {
9699 return;
9700 };
9701 static SPAN: std::sync::OnceLock<(usize, usize)> = std::sync::OnceLock::new();
9702 let (from, to) = *SPAN.get_or_init(|| {
9703 let g = |k: &str, d: usize| {
9704 std::env::var(k)
9705 .ok()
9706 .and_then(|v| v.parse().ok())
9707 .unwrap_or(d)
9708 };
9709 (g("CMF_FFN_REFIT_FROM", 0), g("CMF_FFN_REFIT_TO", usize::MAX))
9710 });
9711 if li < from || li > to {
9712 return;
9713 }
9714 let mut guard = map.lock().unwrap();
9715 let (map, shared) = &mut *guard;
9716 let acc = match map.entry(li) {
9717 std::collections::hash_map::Entry::Occupied(e) => e.into_mut(),
9718 std::collections::hash_map::Entry::Vacant(e) => {
9719 let path = format!("{dir}/support.{li}.u32");
9720 let Ok(bytes) = std::fs::read(&path) else {
9721 eprintln!("refit: no {path} — layer {li} skipped");
9722 return;
9723 };
9724 let n = u32::from_le_bytes(bytes[0..4].try_into().unwrap()) as usize;
9725 let support: Vec<u32> = bytes[4..4 + n * 4]
9726 .chunks_exact(4)
9727 .map(|c| u32::from_le_bytes([c[0], c[1], c[2], c[3]]))
9728 .collect();
9729 eprintln!(
9730 "refit: layer {li} support {n} ({:.0} MB of accumulator)",
9731 (n * n + hidden * n) as f64 * 4.0 / 1e6
9732 );
9733 e.insert(RefitAcc {
9734 gss: vec![0.0; n * n],
9735 ya: vec![0.0; hidden * n],
9736 buf_g: Vec::new(),
9737 buf_o: Vec::new(),
9738 buf_t: 0,
9739 support,
9740 hidden,
9741 tokens: 0,
9742 })
9743 }
9744 };
9745 let ns = acc.support.len();
9746 let cap = refit_batch();
9748 if acc.buf_g.is_empty() {
9749 acc.buf_g = vec![0.0; ns * cap];
9750 acc.buf_o = vec![0.0; hidden * cap];
9751 }
9752 let take = b.min(cap - acc.buf_t);
9753 for t in 0..take {
9754 let col = acc.buf_t + t;
9755 for (j, &n) in acc.support.iter().enumerate() {
9756 acc.buf_g[j * cap + col] = g[t * inter + n as usize];
9757 }
9758 for h in 0..hidden {
9759 acc.buf_o[h * cap + col] = out[t * hidden + h];
9760 }
9761 }
9762 acc.buf_t += take;
9763 acc.tokens += take as u64;
9764 if acc.buf_t < cap {
9765 return;
9766 }
9767 let bt = acc.buf_t;
9768 acc.buf_t = 0;
9769 let RefitAcc {
9779 gss, ya, buf_g, buf_o, ..
9780 } = acc;
9781 let need = (ns * ns).max(hidden * ns);
9782 if shared.len() < need {
9783 shared.resize(need, 0.0);
9784 }
9785 let scratch = &mut shared[..];
9786 let _ = bt;
9787 if crate::gpu::gemm_nt_f32_transient(buf_g, buf_g, &mut scratch[..ns * ns], ns, cap, ns) {
9788 add_into(gss, &scratch[..ns * ns], pool);
9789 if crate::gpu::gemm_nt_f32_transient(buf_o, buf_g, &mut scratch[..hidden * ns], hidden, cap, ns) {
9790 add_into(ya, &scratch[..hidden * ns], pool);
9791 } else {
9792 accum_outer_t(ya, hidden, ns, cap, buf_o, buf_g, pool);
9793 }
9794 } else {
9795 accum_outer_t(gss, ns, ns, cap, buf_g, buf_g, pool);
9796 accum_outer_t(ya, hidden, ns, cap, buf_o, buf_g, pool);
9797 }
9798 }
9802
9803fn refit_batch() -> usize {
9805 static B: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
9806 *B.get_or_init(|| {
9807 std::env::var("CMF_FFN_REFIT_BATCH")
9808 .ok()
9809 .and_then(|v| v.parse().ok())
9810 .unwrap_or(4096)
9811 })
9812}
9813
9814fn accum_outer_t(
9817 c: &mut [f32],
9818 m: usize,
9819 n: usize,
9820 b: usize,
9821 left: &[f32],
9822 right: &[f32],
9823 pool: Option<&Pool>,
9824) {
9825 let ptr = SendMut(c.as_mut_ptr());
9826 let body = |i: usize| {
9827 let ptr = &ptr;
9828 let row = unsafe { std::slice::from_raw_parts_mut(ptr.0.add(i * n), n) };
9829 for t in 0..b {
9830 let a = left[i * b + t];
9831 if a == 0.0 {
9832 continue;
9833 }
9834 for (j, o) in row.iter_mut().enumerate() {
9835 *o += a * right[j * b + t];
9836 }
9837 }
9838 };
9839 match pool {
9840 Some(p) if m > 1 => p.run_rows(m, &|s, e| {
9841 for i in s..e {
9842 body(i);
9843 }
9844 }),
9845 _ => {
9846 for i in 0..m {
9847 body(i);
9848 }
9849 }
9850 }
9851}
9852
9853fn add_into(dst: &mut [f32], src: &[f32], pool: Option<&Pool>) {
9856 let n = dst.len().min(src.len());
9857 match pool {
9858 Some(p) if n >= 1 << 16 => {
9859 let ptr = SendMut(dst.as_mut_ptr());
9860 let f = |s: usize, e: usize| {
9861 let ptr = &ptr;
9862 for blk in s..e {
9863 let (a, b) = (blk * 4096, ((blk + 1) * 4096).min(n));
9864 for i in a..b {
9865 unsafe { *ptr.0.add(i) += src[i] };
9866 }
9867 }
9868 };
9869 p.run_rows(n.div_ceil(4096), &f);
9870 }
9871 _ => {
9872 for (d, v) in dst.iter_mut().zip(&src[..n]) {
9873 *d += *v;
9874 }
9875 }
9876 }
9877}
9878
9879fn accum_outer(
9884 c: &mut [f32],
9885 m: usize,
9886 n: usize,
9887 b: usize,
9888 left: &[f32],
9889 right: &[f32],
9890 pool: Option<&Pool>,
9891) {
9892 const TILE: usize = 32;
9893 let tiles = m.div_ceil(TILE);
9894 let cp = SendMut(c.as_mut_ptr());
9895 let body = |ti: usize| {
9896 let cp = &cp;
9897 let i0 = ti * TILE;
9898 let i1 = (i0 + TILE).min(m);
9899 for t in 0..b {
9900 let r = &right[t * n..t * n + n];
9901 for i in i0..i1 {
9902 let a = left[i * b + t];
9903 if a == 0.0 {
9904 continue;
9905 }
9906 let row = unsafe { std::slice::from_raw_parts_mut(cp.0.add(i * n), n) };
9908 for (o, v) in row.iter_mut().zip(r) {
9909 *o += a * *v;
9910 }
9911 }
9912 }
9913 };
9914 match pool {
9915 Some(p) if tiles > 1 => p.run_rows(tiles, &|s, e| {
9916 for ti in s..e {
9917 body(ti);
9918 }
9919 }),
9920 _ => {
9921 for ti in 0..tiles {
9922 body(ti);
9923 }
9924 }
9925 }
9926}
9927
9928pub fn refit_flush() -> usize {
9930 let Some((dir, map)) = refit_dir() else {
9931 return 0;
9932 };
9933 let guard = map.lock().unwrap();
9934 let mut n = 0;
9935 for (li, acc) in guard.0.iter() {
9936 let w = |name: &str, v: &[f32]| {
9939 let path = format!("{dir}/{name}.{li}.f32");
9940 let bytes: Vec<u8> = v.iter().flat_map(|x| x.to_le_bytes()).collect();
9941 match std::fs::write(&path, &bytes) {
9942 Ok(()) => {}
9943 Err(e) => eprintln!("refit: FAILED to write {path} ({} MB): {e}", bytes.len() / 1_000_000),
9944 }
9945 };
9946 w("gss", &acc.gss);
9947 w("ya", &acc.ya);
9948 println!(
9949 "refit L{li}: {} support, {} tokens, hidden {}",
9950 acc.support.len(),
9951 acc.tokens,
9952 acc.hidden
9953 );
9954 n += 1;
9955 }
9956 n
9957}
9958
9959fn adump_row(li: usize, g: &[f32]) {
9964 use std::io::Write as _;
9965 static FILES: std::sync::OnceLock<
9966 Option<(String, std::sync::Mutex<std::collections::HashMap<usize, std::fs::File>>)>,
9967 > = std::sync::OnceLock::new();
9968 let Some((prefix, map)) = FILES
9969 .get_or_init(|| {
9970 std::env::var("CMF_FFN_ADUMP")
9971 .ok()
9972 .map(|p| (p, std::sync::Mutex::new(std::collections::HashMap::new())))
9973 })
9974 .as_ref()
9975 else {
9976 return;
9977 };
9978 static SPAN: std::sync::OnceLock<(usize, usize)> = std::sync::OnceLock::new();
9981 let (from, to) = *SPAN.get_or_init(|| {
9982 let g = |k: &str, d: usize| {
9983 std::env::var(k)
9984 .ok()
9985 .and_then(|v| v.parse().ok())
9986 .unwrap_or(d)
9987 };
9988 (g("CMF_FFN_ADUMP_FROM", 0), g("CMF_FFN_ADUMP_TO", usize::MAX))
9989 });
9990 if li < from || li > to {
9991 return;
9992 }
9993 let mut map = map.lock().unwrap();
9994 let f = map.entry(li).or_insert_with(|| {
9995 std::fs::File::create(format!("{prefix}.{li}.f16")).expect("adump file")
9996 });
9997 let mut bytes = Vec::with_capacity(g.len() * 2);
9998 for v in g {
9999 bytes.extend_from_slice(&cortiq_core::quant::f32_to_f16(*v).to_le_bytes());
10000 }
10001 let _ = f.write_all(&bytes);
10002}
10003
10004fn oracle_topk() -> usize {
10010 static K: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
10011 *K.get_or_init(|| {
10012 std::env::var("CMF_FFN_ORACLE_TOPK")
10013 .ok()
10014 .and_then(|v| v.parse().ok())
10015 .unwrap_or(0)
10016 })
10017}
10018
10019fn gate_topk() -> usize {
10025 static K: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
10026 *K.get_or_init(|| {
10027 std::env::var("CMF_FFN_GATE_TOPK")
10028 .ok()
10029 .and_then(|v| v.parse().ok())
10030 .unwrap_or(0)
10031 })
10032}
10033
10034fn gate_block() -> usize {
10041 static B: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
10042 *B.get_or_init(|| {
10043 std::env::var("CMF_FFN_GATE_BLOCK")
10044 .ok()
10045 .and_then(|v| v.parse().ok())
10046 .unwrap_or(1)
10047 })
10048}
10049
10050fn keep_top_blocks(g: &mut [f32], keep_n: usize, block: usize) {
10052 let n = g.len();
10053 let nb = n.div_ceil(block);
10054 let kb = (keep_n.div_ceil(block)).clamp(1, nb);
10055 if kb >= nb {
10056 return;
10057 }
10058 let mut score: Vec<f32> = (0..nb)
10059 .map(|b| {
10060 g[b * block..((b + 1) * block).min(n)]
10061 .iter()
10062 .map(|v| v * v)
10063 .sum::<f32>()
10064 })
10065 .collect();
10066 let mut ord = score.clone();
10067 let (_, kth, _) = ord.select_nth_unstable_by(kb - 1, |a, b| {
10068 b.partial_cmp(a).unwrap_or(std::cmp::Ordering::Equal)
10069 });
10070 let thr = *kth;
10071 for b in 0..nb {
10072 if score[b] < thr {
10073 g[b * block..((b + 1) * block).min(n)].fill(0.0);
10074 }
10075 }
10076 score.clear();
10077}
10078
10079fn keep_top_k(g: &mut [f32], k: usize) {
10081 if gate_block() > 1 {
10082 return keep_top_blocks(g, k, gate_block());
10083 }
10084 let n = g.len();
10085 if k == 0 || k >= n {
10086 return;
10087 }
10088 let mut mag: Vec<f32> = g.iter().map(|v| v.abs()).collect();
10089 let (_, kth, _) = mag.select_nth_unstable_by(k - 1, |a, b| {
10090 b.partial_cmp(a).unwrap_or(std::cmp::Ordering::Equal)
10091 });
10092 let thr = *kth;
10093 for v in g.iter_mut() {
10094 if v.abs() < thr {
10095 *v = 0.0;
10096 }
10097 }
10098}
10099
10100fn probe_sq() -> bool {
10104 static S: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
10105 *S.get_or_init(|| std::env::var("CMF_FFN_PROBE_SQ").is_ok())
10106}
10107
10108fn probe_signed() -> bool {
10112 static S: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
10113 *S.get_or_init(|| std::env::var("CMF_FFN_PROBE_SIGNED").is_ok())
10114}
10115
10116fn meanfill() -> Option<&'static (usize, Vec<f32>)> {
10124 static M: std::sync::OnceLock<Option<(usize, Vec<f32>)>> = std::sync::OnceLock::new();
10125 M.get_or_init(|| {
10126 let p = std::env::var("CMF_FFN_MEANFILL").ok()?;
10127 let b = std::fs::read(&p).ok()?;
10128 let inter = u32::from_le_bytes(b[4..8].try_into().ok()?) as usize;
10129 let vals: Vec<f32> = b[8..]
10130 .chunks_exact(4)
10131 .map(|c| f32::from_le_bytes([c[0], c[1], c[2], c[3]]))
10132 .collect();
10133 eprintln!("meanfill: {} value(s), inter {inter}", vals.len());
10134 Some((inter, vals))
10135 })
10136 .as_ref()
10137}
10138
10139fn probe_topk() -> usize {
10142 static K: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
10143 *K.get_or_init(|| {
10144 std::env::var("CMF_FFN_PROBE_TOPK")
10145 .ok()
10146 .and_then(|v| v.parse().ok())
10147 .unwrap_or(0)
10148 })
10149}
10150
10151thread_local! {
10152 static FFN_PROBE: std::cell::RefCell<Option<Vec<Vec<f64>>>> =
10155 const { std::cell::RefCell::new(None) };
10156}
10157
10158fn dense_ffn_dynamic(d: &DenseFfn, x: &[f32], pool: Option<&Pool>, k: usize) -> Option<Vec<f32>> {
10171 let dt = d.down_t.as_ref()?;
10172 let inter = d.gate_proj.rows();
10173 let hidden = dt.cols();
10174 if k == 0 || k >= inter || d.act != Act::Silu {
10175 return None;
10176 }
10177 DYN_SCRATCH.with(|sc| {
10178 let mut sc = sc.borrow_mut();
10179 let DynScratch { g, mag, live, parts } = &mut *sc;
10180 g.resize(inter, 0.0);
10181 d.gate_proj.matvec(x, g, pool);
10182 for v in g.iter_mut() {
10183 *v = inference::silu(*v);
10184 }
10185 mag.clear();
10188 mag.extend(g.iter().map(|v| v.abs()));
10189 let (_, kth, _) = mag.select_nth_unstable_by(k - 1, |a, b| {
10190 b.partial_cmp(a).unwrap_or(std::cmp::Ordering::Equal)
10191 });
10192 let thr = *kth;
10193 live.clear();
10194 live.extend((0..inter as u32).filter(|&n| g[n as usize].abs() >= thr));
10195 let mut out = vec![0.0f32; hidden];
10196 match pool {
10197 Some(p) if live.len() >= 64 => {
10198 let nw = p.n_workers() + 1;
10199 parts.clear();
10200 parts.resize(nw * hidden, 0.0);
10201 let ptr = SendMut(parts.as_mut_ptr());
10202 let n = live.len();
10203 let live_ref: &[u32] = live;
10204 let g_ref: &[f32] = g;
10205 p.run(&|w, workers| {
10206 let chunk = n.div_ceil(workers);
10207 let (s, e) = (w * chunk, ((w + 1) * chunk).min(n));
10208 if s >= e {
10209 return;
10210 }
10211 WORKER_SCRATCH.with(|ws| {
10212 let mut ws = ws.borrow_mut();
10213 let [scratch, acc] = &mut *ws;
10214 scratch.resize(hidden.max(x.len()), 0.0);
10215 acc.clear();
10216 acc.resize(hidden, 0.0);
10217 for (o, &nrm) in live_ref[s..e].iter().enumerate() {
10218 if let Some(&nx) = live_ref[s..e].get(o + 1) {
10221 d.up_proj.prefetch_row(nx as usize);
10222 dt.prefetch_row(nx as usize);
10223 }
10224 let idx = nrm as usize;
10225 let up = d.up_proj.row_dot(idx, x, scratch);
10226 let a = g_ref[idx] * up;
10227 if a != 0.0 {
10228 dt.add_row_scaled(idx, a, acc, scratch);
10229 }
10230 }
10231 for (j, v) in acc.iter().enumerate() {
10232 unsafe { *ptr.at(w * hidden + j) = *v };
10233 }
10234 });
10235 });
10236 for w in 0..nw {
10237 for (j, o) in out.iter_mut().enumerate() {
10238 *o += parts[w * hidden + j];
10239 }
10240 }
10241 }
10242 _ => {
10243 WORKER_SCRATCH.with(|ws| {
10244 let mut ws = ws.borrow_mut();
10245 let [scratch, _acc] = &mut *ws;
10246 scratch.resize(hidden.max(x.len()), 0.0);
10247 for &nrm in live.iter() {
10248 let idx = nrm as usize;
10249 let up = d.up_proj.row_dot(idx, x, scratch);
10250 let a = g[idx] * up;
10251 if a != 0.0 {
10252 dt.add_row_scaled(idx, a, &mut out, scratch);
10253 }
10254 }
10255 });
10256 }
10257 }
10258 Some(out)
10259 })
10260}
10261
10262struct DynScratch {
10265 g: Vec<f32>,
10266 mag: Vec<f32>,
10267 live: Vec<u32>,
10268 parts: Vec<f32>,
10269}
10270
10271thread_local! {
10272 static DYN_SCRATCH: std::cell::RefCell<DynScratch> = const {
10273 std::cell::RefCell::new(DynScratch {
10274 g: Vec::new(),
10275 mag: Vec::new(),
10276 live: Vec::new(),
10277 parts: Vec::new(),
10278 })
10279 };
10280 static WORKER_SCRATCH: std::cell::RefCell<[Vec<f32>; 2]> =
10282 const { std::cell::RefCell::new([Vec::new(), Vec::new()]) };
10283}
10284
10285fn dense_ffn_masked(d: &DenseFfn, x: &[f32], pool: Option<&Pool>, mask_row: &[u8]) -> Vec<f32> {
10290 let inter = d.gate_proj.rows();
10291 FFN_SCRATCH.with(|s| {
10292 let mut s = s.borrow_mut();
10293 let [g, u, ..] = &mut *s;
10294 g.resize(inter, 0.0);
10295 if d.act == Act::Silu && QTensor::matvec_silu_mul(&d.gate_proj, &d.up_proj, x, g, pool) {
10296 } else {
10298 u.resize(inter, 0.0);
10299 QTensor::matvec_many([&d.gate_proj, &d.up_proj], x, [g, u], pool);
10300 for i in 0..inter {
10301 g[i] = d.act.combine(g[i], u[i]);
10302 }
10303 }
10304 zero_masked_cols(g, 1, inter, mask_row);
10305 let mut out = attention::take_buf(d.down_proj.rows());
10306 d.down_proj.matvec(g, &mut out, pool);
10307 out
10308 })
10309}
10310
10311fn dense_ffn_gpu(d: &DenseFfn, x: &[f32], _pool: Option<&Pool>) -> Option<Vec<f32>> {
10317 if d.act != Act::Silu {
10319 return None;
10320 }
10321 if d.gate_proj.rows() < crate::gpu::min_rows() && !d.gate_proj.is_q1() {
10324 return None;
10325 }
10326 let mut jobs: Vec<crate::gpu::MoeJob> = Vec::with_capacity(1);
10327 let mut model_ref = None;
10328 moe_push_job(d, x, 1.0, &mut jobs, &mut model_ref)?;
10329 let model = model_ref?;
10330 let hidden = jobs[0].down.1;
10331 let mut out = attention::take_buf(hidden);
10332 if crate::gpu::moe_block(&model, &jobs, &mut out) {
10333 Some(out)
10334 } else {
10335 let mut out = out;
10336 attention::recycle_buf(&mut out);
10337 None
10338 }
10339}
10340
10341#[allow(clippy::type_complexity)]
10346#[allow(clippy::type_complexity)]
10347pub(crate) fn moe_parts(
10348 t: &QTensor,
10349) -> Option<(
10350 &std::sync::Arc<cortiq_core::CmfModel>,
10351 usize,
10352 usize,
10353 usize,
10354 &[f32],
10355 &[f32],
10356 bool,
10357 bool,
10358 bool,
10359)> {
10360 match t {
10361 QTensor::Mapped {
10362 model,
10363 idx,
10364 dtype: dt @ (cortiq_core::TensorDtype::Q8_2f | cortiq_core::TensorDtype::Q8Row),
10365 rows,
10366 cols,
10367 row_scale,
10368 col_field,
10369 ..
10370 } if (*dt == cortiq_core::TensorDtype::Q8Row) || !col_field.is_empty() => Some((
10371 model, *idx, *rows, *cols, row_scale, col_field, false, false, false,
10372 )),
10373 QTensor::Mapped {
10375 model,
10376 idx,
10377 dtype: cortiq_core::TensorDtype::Q1,
10378 rows,
10379 cols,
10380 ..
10381 } => Some((
10382 model,
10383 *idx,
10384 *rows,
10385 *cols,
10386 &[][..],
10387 &[][..],
10388 true,
10389 false,
10390 false,
10391 )),
10392 QTensor::Mapped {
10394 model,
10395 idx,
10396 dtype: cortiq_core::TensorDtype::Q4Tiled,
10397 rows,
10398 cols,
10399 ..
10400 } => Some((
10401 model,
10402 *idx,
10403 *rows,
10404 *cols,
10405 &[][..],
10406 &[][..],
10407 false,
10408 true,
10409 false,
10410 )),
10411 QTensor::Mapped {
10413 model,
10414 idx,
10415 dtype: cortiq_core::TensorDtype::Q4TiledP,
10416 rows,
10417 cols,
10418 ..
10419 } => Some((
10420 model,
10421 *idx,
10422 *rows,
10423 *cols,
10424 &[][..],
10425 &[][..],
10426 false,
10427 true,
10428 false,
10429 )),
10430 QTensor::Mapped {
10434 model,
10435 idx,
10436 dtype: cortiq_core::TensorDtype::Q2TiledP,
10437 rows,
10438 cols,
10439 ..
10440 } => Some((
10441 model,
10442 *idx,
10443 *rows,
10444 *cols,
10445 &[][..],
10446 &[][..],
10447 false,
10448 true,
10449 true,
10450 )),
10451 _ => None,
10452 }
10453}
10454
10455#[cfg(target_os = "macos")]
10461fn metal_moe_graph_parts(m: &MoeFfn, hidden: usize) -> Option<crate::gpu::GpuMoe<'_>> {
10462 if m.router_sigmoid
10463 || m.router_input_norm
10464 || m.expert_bias.is_some()
10465 || m.route_tau.is_some()
10466 || m.mask.is_some()
10467 || m.per_expert_scale.is_some()
10468 || m.experts.is_empty()
10469 || m.top_k == 0
10470 || m.resonance.is_some()
10471 {
10472 return None;
10473 }
10474 let (sh, sg) = match &m.shared {
10477 Some((sh, Some(sg))) => (sh, sg),
10478 _ => return None,
10479 };
10480 let (rf, rr, rc) = m.router.f32_parts()?;
10481 if rr != m.experts.len() || rc != hidden {
10482 return None;
10483 }
10484 let (sf, sr, sc) = sg.f32_parts()?;
10485 if sr * sc != hidden {
10486 return None;
10487 }
10488 let inter = m.experts[0].gate_proj.rows();
10489 let gu_q2 = m.experts[0].gate_proj.mapped_q2tp().is_some();
10492 let trio = |e: &DenseFfn| -> Option<(usize, usize, usize)> {
10493 if e.act != Act::Silu
10494 || e.gate_proj.rows() != inter
10495 || e.gate_proj.cols() != hidden
10496 || e.up_proj.rows() != inter
10497 || e.up_proj.cols() != hidden
10498 || e.down_proj.rows() != hidden
10499 || e.down_proj.cols() != inter
10500 {
10501 return None;
10502 }
10503 let pick = |t: &QTensor| -> Option<usize> {
10504 if gu_q2 {
10505 t.mapped_q2tp().map(|(_, i)| i)
10506 } else {
10507 t.mapped_q4tp().map(|(_, i)| i)
10508 }
10509 };
10510 Some((
10511 pick(&e.gate_proj)?,
10512 pick(&e.up_proj)?,
10513 e.down_proj.mapped_q4tp().map(|(_, i)| i)?,
10514 ))
10515 };
10516 let experts = m.experts.iter().map(trio).collect::<Option<Vec<_>>>()?;
10517 let shared = trio(sh)?;
10518 Some(crate::gpu::GpuMoe {
10519 router: rf,
10520 sgate: sf,
10521 experts,
10522 shared,
10523 n_exp: m.experts.len(),
10524 top_k: m.top_k,
10525 inter,
10526 norm_topk: m.norm_topk_prob,
10527 route_scale: m.routed_scaling,
10528 gu_q2,
10529 })
10530}
10531
10532pub(crate) fn moe_push_job_parts<'a>(
10536 gate: &'a QTensor,
10537 up: &'a QTensor,
10538 down: &'a QTensor,
10539 x: &[f32],
10540 w: f32,
10541 swiglu_limit: f32,
10542 jobs: &mut Vec<crate::gpu::MoeJob<'a>>,
10543 model_ref: &mut Option<std::sync::Arc<cortiq_core::CmfModel>>,
10544) -> Option<()> {
10545 use crate::qtensor::prescale;
10546 let (gm, gi, gr, gc, grs, gcf, gq1, gq4, gq2) = moe_parts(gate)?;
10547 let (_, ui, ur, uc, urs, ucf, uq1, uq4, uq2) = moe_parts(up)?;
10548 let (_, di, dr, dc, drs, dcf, dq1, dq4, dq2) = moe_parts(down)?;
10549 if gq1 != uq1 || uq1 != dq1 || gq4 != uq4 || uq4 != dq4 || gq2 != uq2 {
10550 return None; }
10552 if gq2 && (dq2 || !dq4 || down.mapped_q4tp().is_none()) {
10555 return None;
10556 }
10557 if !gq2 && dq2 {
10558 return None;
10559 }
10560 model_ref.get_or_insert_with(|| gm.clone());
10561 let dt = |cf: &[f32]| {
10562 if cf.is_empty() {
10563 cortiq_core::TensorDtype::Q8Row
10564 } else {
10565 cortiq_core::TensorDtype::Q8_2f
10566 }
10567 };
10568 jobs.push(crate::gpu::MoeJob {
10569 gate: (gi, gr, gc, grs),
10570 up: (ui, ur, uc, urs),
10571 down: (di, dr, dc, drs),
10572 xs_gate: prescale(x, gcf, dt(gcf)).into_owned(),
10573 xs_up: prescale(x, ucf, dt(ucf)).into_owned(),
10574 down_col: dcf,
10575 w,
10576 q1: gq1,
10577 q4t: gq4 && !gq2 && gate.mapped_q4tp().is_none(),
10578 q4tp: gq4 && (gq2 || gate.mapped_q4tp().is_some()),
10579 gu_q2: gq2,
10580 swiglu_limit,
10581 });
10582 Some(())
10583}
10584
10585fn moe_push_job<'a>(
10587 d: &'a DenseFfn,
10588 x: &[f32],
10589 w: f32,
10590 jobs: &mut Vec<crate::gpu::MoeJob<'a>>,
10591 model_ref: &mut Option<std::sync::Arc<cortiq_core::CmfModel>>,
10592) -> Option<()> {
10593 use crate::qtensor::prescale;
10594 if d.act != Act::Silu {
10595 return None; }
10597 let (gm, gi, gr, gc, grs, gcf, gq1, gq4, gq2) = moe_parts(&d.gate_proj)?;
10598 let (_, ui, ur, uc, urs, ucf, uq1, uq4, uq2) = moe_parts(&d.up_proj)?;
10599 let (_, di, dr, dc, drs, dcf, dq1, dq4, dq2) = moe_parts(&d.down_proj)?;
10600 if gq1 != uq1 || uq1 != dq1 || gq4 != uq4 || uq4 != dq4 || gq2 != uq2 {
10601 return None; }
10603 if gq2 && (dq2 || !dq4 || d.down_proj.mapped_q4tp().is_none()) {
10604 return None;
10605 }
10606 if !gq2 && dq2 {
10607 return None;
10608 }
10609 model_ref.get_or_insert_with(|| gm.clone());
10610 let gdt = if gcf.is_empty() {
10611 cortiq_core::TensorDtype::Q8Row
10612 } else {
10613 cortiq_core::TensorDtype::Q8_2f
10614 };
10615 let udt = if ucf.is_empty() {
10616 cortiq_core::TensorDtype::Q8Row
10617 } else {
10618 cortiq_core::TensorDtype::Q8_2f
10619 };
10620 jobs.push(crate::gpu::MoeJob {
10621 gate: (gi, gr, gc, grs),
10622 up: (ui, ur, uc, urs),
10623 down: (di, dr, dc, drs),
10624 xs_gate: prescale(x, gcf, gdt).into_owned(),
10625 xs_up: prescale(x, ucf, udt).into_owned(),
10626 down_col: dcf,
10627 w,
10628 q1: gq1,
10629 q4t: gq4 && !gq2 && d.gate_proj.mapped_q4tp().is_none(),
10630 q4tp: gq4 && (gq2 || d.gate_proj.mapped_q4tp().is_some()),
10631 gu_q2: gq2,
10632 swiglu_limit: 0.0,
10633 });
10634 Some(())
10635}
10636
10637fn sparse_ffn_quant(
10644 d: &DenseFfn,
10645 x: &[f32],
10646 active: &[u16],
10647 hidden: usize,
10648 pool: Option<&Pool>,
10649) -> Vec<f32> {
10650 let n = active.len();
10651 let inter = d.gate_proj.rows();
10652 let mut act = vec![0.0f32; n];
10653 let need_scratch = !(d.gate_proj.sparse_col_ok() && d.up_proj.sparse_col_ok());
10656 let compute = |ai: usize| -> f32 {
10657 let idx = active[ai] as usize;
10658 if idx >= inter {
10659 return 0.0; }
10661 let mut s = if need_scratch {
10662 vec![0.0f32; hidden]
10663 } else {
10664 Vec::new()
10665 };
10666 let gate = d.gate_proj.row_dot(idx, x, &mut s);
10667 let up = d.up_proj.row_dot(idx, x, &mut s);
10668 d.act.combine(gate, up)
10669 };
10670 match pool {
10671 Some(p) if n >= 256 => {
10672 let ptr = SendMut(act.as_mut_ptr());
10673 p.run(&|widx, nw| {
10674 let chunk = n.div_ceil(nw);
10675 let (s, e) = (widx * chunk, ((widx + 1) * chunk).min(n));
10676 for ai in s..e {
10677 unsafe { *ptr.at(ai) = compute(ai) };
10678 }
10679 });
10680 }
10681 _ => {
10682 for (ai, a) in act.iter_mut().enumerate() {
10683 *a = compute(ai);
10684 }
10685 }
10686 }
10687 let mut out = vec![0.0f32; hidden];
10689 for (ai, &idx) in active.iter().enumerate() {
10690 let w = act[ai];
10691 if w.abs() >= 1e-12 && (idx as usize) < inter {
10692 d.down_proj.add_col_scaled(idx as usize, w, &mut out);
10693 }
10694 }
10695 out
10696}
10697
10698#[doc(hidden)]
10700pub fn sparse_ffn_quant_for_test(
10701 d: &DenseFfn,
10702 x: &[f32],
10703 active: &[u16],
10704 hidden: usize,
10705) -> Vec<f32> {
10706 sparse_ffn_quant(d, x, active, hidden, None)
10707}
10708
10709fn dequant_dense_f32(d: &DenseFfn) -> (Vec<f32>, Vec<f32>, Vec<f32>) {
10713 let deq = |t: &QTensor| -> Vec<f32> {
10714 let (rows, cols) = (t.rows(), t.cols());
10715 let mut out = vec![0.0f32; rows * cols];
10716 for r in 0..rows {
10717 t.row_f32(r, &mut out[r * cols..(r + 1) * cols]);
10718 }
10719 out
10720 };
10721 (deq(&d.gate_proj), deq(&d.up_proj), deq(&d.down_proj))
10722}
10723
10724struct SendMut(*mut f32);
10726unsafe impl Send for SendMut {}
10727unsafe impl Sync for SendMut {}
10728impl SendMut {
10729 #[inline]
10730 #[allow(clippy::mut_from_ref)]
10733 unsafe fn at(&self, i: usize) -> &mut f32 {
10734 unsafe { &mut *self.0.add(i) }
10735 }
10736}
10737
10738fn moe_route(logits: &[f32], m: &MoeFfn, allowed: Option<&[bool]>) -> (Vec<usize>, Vec<f32>, f32) {
10748 let ne = logits.len();
10749 let p: Vec<f32> = if m.router_sigmoid {
10750 logits.iter().map(|&l| 1.0 / (1.0 + (-l).exp())).collect()
10751 } else {
10752 let mx = logits.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
10753 let mut e: Vec<f32> = logits.iter().map(|&l| (l - mx).exp()).collect();
10754 let s: f32 = e.iter().sum();
10755 for v in &mut e {
10756 *v /= s;
10757 }
10758 e
10759 };
10760 let admit = |e: usize| {
10766 m.mask.as_ref().is_none_or(|mk| mk[e])
10767 && allowed.is_none_or(|a| a.get(e).copied().unwrap_or(false))
10768 };
10769 let mut idx: Vec<usize> = (0..ne).filter(|&e| admit(e)).collect();
10770 match &m.expert_bias {
10772 Some(b) => idx.sort_unstable_by(|&x, &y| {
10773 (p[y] + b[y])
10774 .partial_cmp(&(p[x] + b[x]))
10775 .unwrap()
10776 .then(x.cmp(&y))
10777 }),
10778 None => idx.sort_unstable_by(|&x, &y| p[y].partial_cmp(&p[x]).unwrap().then(x.cmp(&y))),
10779 }
10780 idx.truncate(m.top_k);
10781 if let Some(tau) = m.route_tau {
10785 let total: f32 = idx.iter().map(|&e| p[e]).sum();
10786 if total > 0.0 {
10787 let mut acc = 0.0f32;
10788 let mut keep = idx.len();
10789 for (i, &e) in idx.iter().enumerate() {
10790 acc += p[e];
10791 if acc >= tau * total {
10792 keep = i + 1;
10793 break;
10794 }
10795 }
10796 idx.truncate(keep);
10797 }
10798 }
10799 let wsum: f32 = if m.norm_topk_prob {
10800 let s: f32 = idx.iter().map(|&e| p[e]).sum();
10801 (if m.router_sigmoid { s + 1e-6 } else { s }) / m.routed_scaling
10804 } else {
10805 1.0 / m.routed_scaling
10806 };
10807 (idx, p, wsum)
10808}
10809
10810fn moe_ffn(m: &MoeFfn, x: &[f32], pool: Option<&Pool>, allowed: Option<&[bool]>) -> Vec<f32> {
10813 accumulate_act(m, x, 1);
10814 let ne = m.experts.len();
10815 let mut logits = vec![0.0f32; ne];
10816 match &m.resonance {
10817 Some(r) => r.scores(x, &mut logits),
10818 None => m.router.matvec(x, &mut logits, pool),
10819 }
10820 let (idx, p, wsum) = moe_route(&logits, m, allowed);
10821 {
10822 let mut st = m.stats.borrow_mut();
10823 if st.len() < ne {
10824 st.resize(ne, 0);
10825 }
10826 for &e in &idx {
10827 st[e] += 1;
10828 }
10829 }
10830 if crate::gpu::enabled_here() {
10835 match crate::gpu::probe_arm(crate::gpu::OpClass::Ffn) {
10836 crate::gpu::ProbeArm::Gpu => {
10837 let t0 = std::time::Instant::now();
10838 if let Some(out) = moe_ffn_gpu(m, x, &idx, &p, wsum, pool) {
10839 crate::gpu::probe_record(crate::gpu::OpClass::Ffn, true, t0.elapsed());
10840 return out;
10841 }
10842 }
10843 crate::gpu::ProbeArm::CpuTimed => {
10844 let t0 = std::time::Instant::now();
10845 let out = crate::gpu::cpu_scope(|| moe_ffn_cpu(m, x, &idx, &p, wsum, pool));
10846 crate::gpu::probe_record(crate::gpu::OpClass::Ffn, false, t0.elapsed());
10847 return out;
10848 }
10849 crate::gpu::ProbeArm::Cpu => {
10850 return crate::gpu::cpu_scope(|| moe_ffn_cpu(m, x, &idx, &p, wsum, pool));
10851 }
10852 }
10853 }
10854 moe_ffn_cpu(m, x, &idx, &p, wsum, pool)
10855}
10856
10857fn graph_note(built: bool) {
10861 use std::sync::atomic::{AtomicBool, Ordering};
10862 if built {
10863 GRAPH_TOK_OK.fetch_add(1, Ordering::Relaxed);
10864 } else {
10865 GRAPH_TOK_MISS.fetch_add(1, Ordering::Relaxed);
10866 }
10867 static SAID: AtomicBool = AtomicBool::new(false);
10868 if !SAID.swap(true, Ordering::Relaxed) {
10869 if built {
10870 tracing::info!("wgpu whole-token graph: ACTIVE");
10871 } else {
10872 tracing::warn!("wgpu whole-token graph refused — per-op path");
10873 }
10874 }
10875}
10876
10877pub static GRAPH_TOK_OK: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
10881pub static GRAPH_TOK_MISS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
10882
10883fn moe_batch_enabled() -> bool {
10886 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
10887 *ON.get_or_init(|| std::env::var("CMF_MOE_BATCH").as_deref() != Ok("0"))
10888}
10889
10890fn moe_ffn_cpu_batched(
10896 m: &MoeFfn,
10897 x: &[f32],
10898 idx: &[usize],
10899 p: &[f32],
10900 wsum: f32,
10901 pool: Option<&Pool>,
10902) -> Option<Vec<f32>> {
10903 if idx.is_empty() || !moe_batch_enabled() {
10904 return None;
10905 }
10906 if FFN_PROBE.with(|pr| pr.borrow().is_some()) {
10910 return None;
10911 }
10912 let n = idx.len() + usize::from(m.shared.is_some());
10913 let mut pairs = Vec::with_capacity(n);
10914 let mut downs = Vec::with_capacity(n);
10915 let mut ws = Vec::with_capacity(n);
10916 for &e in idx {
10917 let d = &m.experts[e];
10918 if d.act != Act::Silu {
10919 return None;
10920 }
10921 pairs.push((&d.gate_proj, &d.up_proj));
10922 downs.push(&d.down_proj);
10923 ws.push(p[e] / wsum * m.per_expert_scale.as_ref().map_or(1.0, |v| v[e]));
10924 }
10925 if let Some((se, gate)) = &m.shared {
10928 if se.act != Act::Silu {
10929 return None;
10930 }
10931 let g = gate.as_ref().map_or(1.0, |gate| {
10932 let mut gl = [0.0f32; 1];
10933 gate.matvec(x, &mut gl, pool);
10934 1.0 / (1.0 + (-gl[0]).exp())
10935 });
10936 pairs.push((&se.gate_proj, &se.up_proj));
10937 downs.push(&se.down_proj);
10938 ws.push(g);
10939 }
10940 let inter = pairs[0].0.rows();
10941 let mut gs: Vec<Vec<f32>> = (0..pairs.len()).map(|_| vec![0f32; inter]).collect();
10942 if !QTensor::moe_gate_up_many(&pairs, x, &mut gs, pool) {
10943 return None;
10944 }
10945 let mut out = attention::take_buf(x.len());
10946 if !QTensor::moe_down_many(&downs, &gs, &ws, &mut out, pool) {
10947 attention::recycle_buf(&mut out);
10948 return None;
10949 }
10950 Some(out)
10951}
10952
10953fn moe_ffn_cpu(
10955 m: &MoeFfn,
10956 x: &[f32],
10957 idx: &[usize],
10958 p: &[f32],
10959 wsum: f32,
10960 pool: Option<&Pool>,
10961) -> Vec<f32> {
10962 if let Some(out) = moe_ffn_cpu_batched(m, x, idx, p, wsum, pool) {
10963 return out;
10964 }
10965 let mut out = attention::take_buf(x.len());
10966 for &e in idx {
10967 let mut eo = dense_ffn(&m.experts[e], x, pool);
10968 let w = p[e] / wsum * m.per_expert_scale.as_ref().map_or(1.0, |v| v[e]);
10969 for i in 0..out.len() {
10970 out[i] += w * eo[i];
10971 }
10972 attention::recycle_buf(&mut eo);
10973 }
10974 if let Some((se, gate)) = &m.shared {
10975 let mut so = dense_ffn(se, x, pool);
10976 let g = gate.as_ref().map_or(1.0, |gate| {
10977 let mut gl = [0.0f32; 1];
10978 gate.matvec(x, &mut gl, pool);
10979 1.0 / (1.0 + (-gl[0]).exp())
10980 });
10981 for i in 0..out.len() {
10982 out[i] += g * so[i];
10983 }
10984 attention::recycle_buf(&mut so);
10985 }
10986 out
10987}
10988
10989#[allow(clippy::too_many_arguments)]
10997fn mla_attention(
10998 w: &MlaWeights,
10999 normed: &[f32],
11000 cache: &mut crate::kv_cache::LayerKvCache,
11001 position: usize,
11002 inv_freq: &[f32],
11003 rope_scale: f32,
11004 eps: f64,
11005 pool: Option<&Pool>,
11006) -> Vec<f32> {
11007 let (nh, dr, dn, dv, lora) = (w.nh, w.qk_rope, w.qk_nope, w.v_dim, w.lora);
11008 let hd = dr + dn;
11009 let mut q = vec![0.0f32; nh * hd];
11010 match (&w.q_a, &w.q_a_norm) {
11011 (Some(qa), Some(qn)) => {
11012 let mut t = vec![0.0f32; qa.rows()];
11013 qa.matvec(normed, &mut t, pool);
11014 let tn = inference::rms_norm(&t, qn, eps, NormStyle::Qwen);
11015 w.q_proj.matvec(&tn, &mut q, pool);
11016 }
11017 _ => w.q_proj.matvec(normed, &mut q, pool),
11018 }
11019 let mut ca = vec![0.0f32; lora + dr];
11020 w.kv_a.matvec(normed, &mut ca, pool);
11021 let (c_lat, k_rope) = ca.split_at_mut(lora);
11022 let latn = inference::rms_norm(c_lat, &w.kv_a_norm, eps, NormStyle::Qwen);
11023 let mut kvb = vec![0.0f32; nh * (dn + dv)];
11024 w.kv_b.matvec(&latn, &mut kvb, pool);
11025 if !w.nope {
11026 attention::rope_rotate_scaled(k_rope, position, inv_freq, rope_scale);
11027 }
11028 for h in 0..nh {
11029 if !w.nope {
11030 attention::rope_rotate_scaled(
11031 &mut q[h * hd..h * hd + dr],
11032 position,
11033 inv_freq,
11034 rope_scale,
11035 );
11036 }
11037 }
11038 let mut k = vec![0.0f32; nh * hd];
11039 let mut v = vec![0.0f32; nh * hd];
11040 for h in 0..nh {
11041 k[h * hd..h * hd + dr].copy_from_slice(k_rope);
11042 k[h * hd + dr..(h + 1) * hd].copy_from_slice(&kvb[h * (dn + dv)..h * (dn + dv) + dn]);
11043 v[h * hd..h * hd + dv].copy_from_slice(&kvb[h * (dn + dv) + dn..(h + 1) * (dn + dv)]);
11044 }
11045 cache.append(&k, &v, &vec![true; nh]);
11046 let (ao, mut imp) = attention::attend_all_heads(&q, cache, nh, 1, hd, w.scale, None, 0.0);
11047 attention::recycle_buf(&mut imp);
11048 let mut ov = vec![0.0f32; nh * dv];
11049 for h in 0..nh {
11050 ov[h * dv..(h + 1) * dv].copy_from_slice(&ao[h * hd..h * hd + dv]);
11051 }
11052 let mut out = vec![0.0f32; w.o_proj.rows()];
11053 w.o_proj.matvec(&ov, &mut out, pool);
11054 out
11055}
11056
11057fn dense_moe_ffn(
11064 dm: &DenseMoeFfn,
11065 x_normed: &[f32],
11066 h_raw: &[f32],
11067 eps: f64,
11068 norm_style: NormStyle,
11069 pool: Option<&Pool>,
11070) -> Vec<f32> {
11071 let mut d = dense_ffn(&dm.dense, x_normed, pool);
11072 d = inference::rms_norm(&d, &dm.post_norm_1, eps, norm_style);
11073 let m = &dm.moe;
11074 let ne = m.experts.len();
11075 let mut logits = vec![0.0f32; ne];
11076 if m.router_input_norm {
11077 let ss: f32 = h_raw.iter().map(|v| v * v).sum::<f32>() / h_raw.len() as f32;
11078 let inv = 1.0 / (ss + eps as f32).sqrt();
11079 let xr: Vec<f32> = h_raw.iter().map(|v| v * inv).collect();
11080 m.router.matvec(&xr, &mut logits, pool);
11081 } else {
11082 m.router.matvec(h_raw, &mut logits, pool);
11083 }
11084 let (idx, p, wsum) = moe_route(&logits, m, None);
11085 {
11086 let mut st = m.stats.borrow_mut();
11087 if st.len() < ne {
11088 st.resize(ne, 0);
11089 }
11090 for &e in &idx {
11091 st[e] += 1;
11092 }
11093 }
11094 let x2 = inference::rms_norm(h_raw, &dm.pre_norm_2, eps, norm_style);
11095 let mo = moe_ffn_cpu(m, &x2, &idx, &p, wsum, pool);
11096 let mo = inference::rms_norm(&mo, &dm.post_norm_2, eps, norm_style);
11097 for (di, mi) in d.iter_mut().zip(&mo) {
11098 *di += mi;
11099 }
11100 d
11101}
11102
11103fn moe_gpu_refused(why: &'static str) {
11110 use std::sync::atomic::{AtomicBool, Ordering};
11111 static SAID: AtomicBool = AtomicBool::new(false);
11112 if !SAID.swap(true, Ordering::Relaxed) {
11113 tracing::warn!("MoE GPU block refused ({why}) — experts run on the CPU");
11114 }
11115}
11116
11117fn moe_ffn_gpu(
11118 m: &MoeFfn,
11119 x: &[f32],
11120 idx: &[usize],
11121 p: &[f32],
11122 wsum: f32,
11123 pool: Option<&Pool>,
11124) -> Option<Vec<f32>> {
11125 use crate::gpu::MoeJob;
11126
11127 let mut jobs: Vec<MoeJob> = Vec::with_capacity(idx.len() + 1);
11128 let mut model_ref = None;
11129 for &e in idx {
11130 if moe_push_job(&m.experts[e], x, p[e] / wsum, &mut jobs, &mut model_ref).is_none() {
11131 moe_gpu_refused("push_job(expert)");
11132 return None;
11133 }
11134 }
11135 if let Some((se, gate)) = &m.shared {
11136 let g = gate.as_ref().map_or(1.0, |gate| {
11137 let mut gl = [0.0f32; 1];
11138 gate.matvec(x, &mut gl, pool);
11139 1.0 / (1.0 + (-gl[0]).exp())
11140 });
11141 if moe_push_job(se, x, g, &mut jobs, &mut model_ref).is_none() {
11142 moe_gpu_refused("push_job(shared)");
11143 return None;
11144 }
11145 }
11146 let Some(model) = model_ref else {
11147 moe_gpu_refused("no model_ref");
11148 return None;
11149 };
11150 let hidden = jobs[0].down.1;
11151 let mut out = vec![0.0f32; hidden];
11152 if crate::gpu::moe_block(&model, &jobs, &mut out) {
11153 Some(out)
11154 } else {
11155 moe_gpu_refused("gpu::moe_block");
11156 None
11157 }
11158}
11159
11160fn ffn_forward(
11162 ffn: &FfnKind,
11163 x: &[f32],
11164 pool: Option<&Pool>,
11165 experts_allowed: Option<&[bool]>,
11166) -> Vec<f32> {
11167 match ffn {
11168 FfnKind::Dense(d) if !d.segs.is_empty() => tube_ffn(d, x, 1, pool, None),
11169 FfnKind::Dense(d) => dense_ffn(d, x, pool),
11170 FfnKind::Moe(m) => moe_ffn(m, x, pool, experts_allowed),
11171 FfnKind::DenseMoe(_) => unreachable!("DenseMoe dispatches via dense_moe_ffn"),
11175 }
11176}
11177
11178fn ffn_forward_pair(
11182 ffn: &FfnKind,
11183 x1: &[f32],
11184 x2: &[f32],
11185 pool: Option<&Pool>,
11186 experts_allowed: Option<&[bool]>,
11187) -> (Vec<f32>, Vec<f32>) {
11188 let d = match ffn {
11189 FfnKind::Dense(d) if !d.segs.is_empty() => {
11192 return (
11193 tube_ffn(d, x1, 1, pool, None),
11194 tube_ffn(d, x2, 1, pool, None),
11195 );
11196 }
11197 FfnKind::Dense(d) => d,
11198 FfnKind::Moe(m) => {
11199 return (
11200 moe_ffn(m, x1, pool, experts_allowed),
11201 moe_ffn(m, x2, pool, experts_allowed),
11202 );
11203 }
11204 FfnKind::DenseMoe(_) => unreachable!("DenseMoe dispatches via dense_moe_ffn"),
11205 };
11206 let inter = d.gate_proj.rows();
11207 FFN_SCRATCH.with(|s| {
11208 let mut s = s.borrow_mut();
11209 let [g1, g2, u1, u2] = &mut *s;
11210 g1.resize(inter, 0.0);
11211 g2.resize(inter, 0.0);
11212 u1.resize(inter, 0.0);
11213 u2.resize(inter, 0.0);
11214 QTensor::matvec2_many(
11217 [&d.gate_proj, &d.up_proj],
11218 x1,
11219 x2,
11220 [g1.as_mut_slice(), u1.as_mut_slice()],
11221 [g2.as_mut_slice(), u2.as_mut_slice()],
11222 pool,
11223 );
11224 for i in 0..inter {
11225 g1[i] = d.act.combine(g1[i], u1[i]);
11226 g2[i] = d.act.combine(g2[i], u2[i]);
11227 }
11228 let mut o1 = attention::take_buf(d.down_proj.rows());
11229 let mut o2 = attention::take_buf(d.down_proj.rows());
11230 d.down_proj.matvec2(g1, g2, &mut o1, &mut o2, pool);
11231 (o1, o2)
11232 })
11233}
11234
11235#[cfg(test)]
11236mod tests {
11237
11238 #[test]
11239 fn cancel_flag_stops_generation() {
11240 let mut p = create_test_pipeline(16, 32, 2, 2, 8, 2, 32);
11241 p.cancel.store(true, std::sync::atomic::Ordering::Relaxed);
11244 let r = p.generate_from_ids(&[1, 2, 3], 8, None, None).unwrap();
11245 assert_eq!(r.finish_reason, "cancelled");
11246 assert!(
11247 r.token_ids.is_empty(),
11248 "no tokens after cancel: {:?}",
11249 r.token_ids
11250 );
11251 let r2 = p.generate_from_ids(&[1, 2, 3], 4, None, None).unwrap();
11253 assert_ne!(r2.finish_reason, "cancelled");
11254 }
11255 use super::*;
11256
11257 #[test]
11265 fn dynamic_ffn_equals_the_zeroing_arm() {
11266 let (hidden, inter) = (8usize, 32usize);
11267 let synth = |n: usize, salt: usize| -> Vec<f32> {
11268 (0..n)
11269 .map(|i| (((i * 29 + salt * 13 + 7) % 89) as f32 / 89.0 - 0.5) * 0.6)
11270 .collect()
11271 };
11272 let down = synth(hidden * inter, 3);
11273 let mut down_t = vec![0.0f32; inter * hidden];
11274 for r in 0..hidden {
11275 for c in 0..inter {
11276 down_t[c * hidden + r] = down[r * inter + c];
11277 }
11278 }
11279 let d = DenseFfn {
11280 gate_proj: QTensor::from_f32(synth(inter * hidden, 1), inter, hidden),
11281 up_proj: QTensor::from_f32(synth(inter * hidden, 2), inter, hidden),
11282 down_proj: QTensor::from_f32(down.clone(), hidden, inter),
11283 act: Act::Silu,
11284 down_t: Some(QTensor::from_f32(down_t, inter, hidden)),
11285 segs: Vec::new(),
11286 };
11287 let x = synth(hidden, 11);
11288 let k = 12usize;
11289 let got = dense_ffn_dynamic(&d, &x, None, k).expect("down_t present");
11290 let mut g = vec![0.0f32; inter];
11292 d.gate_proj.matvec(&x, &mut g, None);
11293 let mut u = vec![0.0f32; inter];
11294 d.up_proj.matvec(&x, &mut u, None);
11295 for v in g.iter_mut() {
11296 *v = inference::silu(*v);
11297 }
11298 keep_top_k(&mut g, k);
11299 for i in 0..inter {
11300 g[i] *= u[i];
11301 }
11302 let mut want = vec![0.0f32; hidden];
11303 d.down_proj.matvec(&g, &mut want, None);
11304 for (a, b) in want.iter().zip(&got) {
11305 assert!((a - b).abs() < 1e-5, "dynamic {b} vs reference {a}");
11306 }
11307 }
11308
11309 #[test]
11315 fn tube_ffn_open_equals_dense_and_closed_equals_masked() {
11316 let (hidden, core, tube) = (8usize, 12usize, 8usize);
11317 let inter = core + tube;
11318 let synth = |n: usize, salt: usize| -> Vec<f32> {
11319 (0..n)
11320 .map(|i| (((i * 41 + salt * 17 + 5) % 97) as f32 / 97.0 - 0.5) * 0.5)
11321 .collect()
11322 };
11323 let (g_all, u_all) = (synth(inter * hidden, 1), synth(inter * hidden, 2));
11324 let d_all = synth(hidden * inter, 3);
11325 let dense = DenseFfn {
11327 gate_proj: QTensor::from_f32(g_all.clone(), inter, hidden),
11328 up_proj: QTensor::from_f32(u_all.clone(), inter, hidden),
11329 down_proj: QTensor::from_f32(d_all.clone(), hidden, inter),
11330 act: Act::Silu,
11331 down_t: None,
11332 segs: Vec::new(),
11333 };
11334 let rows = |v: &[f32], a: usize, b: usize| -> Vec<f32> {
11335 v[a * hidden..b * hidden].to_vec()
11336 };
11337 let cols = |v: &[f32], a: usize, b: usize| -> Vec<f32> {
11338 let mut o = Vec::with_capacity(hidden * (b - a));
11339 for r in 0..hidden {
11340 o.extend_from_slice(&v[r * inter + a..r * inter + b]);
11341 }
11342 o
11343 };
11344 let tubed = DenseFfn {
11345 down_t: None,
11346 gate_proj: QTensor::from_f32(rows(&g_all, 0, core), core, hidden),
11347 up_proj: QTensor::from_f32(rows(&u_all, 0, core), core, hidden),
11348 down_proj: QTensor::from_f32(cols(&d_all, 0, core), hidden, core),
11349 act: Act::Silu,
11350 segs: vec![FfnSeg {
11351 gate: QTensor::from_f32(rows(&g_all, core, inter), tube, hidden),
11352 up: QTensor::from_f32(rows(&u_all, core, inter), tube, hidden),
11353 down: QTensor::from_f32(cols(&d_all, core, inter), hidden, tube),
11354 start: core,
11355 width: tube,
11356 }],
11357 };
11358 let x = synth(hidden, 7);
11359 let want = dense_ffn(&dense, &x, None);
11360 let got = tube_ffn(&tubed, &x, 1, None, None);
11361 for (a, b) in want.iter().zip(&got) {
11362 assert!((a - b).abs() < 1e-5, "open tube: {a} vs {b}");
11363 }
11364 let mut bits = vec![0u8; inter.div_ceil(8)];
11366 for n in 0..core {
11367 bits[n / 8] |= 1 << (n % 8);
11368 }
11369 let closed = tube_ffn(&tubed, &x, 1, None, Some(&bits));
11370 let masked = dense_ffn_masked(&dense, &x, None, &bits);
11371 for (a, b) in masked.iter().zip(&closed) {
11372 assert!((a - b).abs() < 1e-5, "closed tube: {a} vs {b}");
11373 }
11374 let batch = tube_ffn(&tubed, &x, 1, None, Some(&bits));
11376 for (a, b) in closed.iter().zip(&batch) {
11377 assert_eq!(a, b, "batch arm disagrees with decode arm");
11378 }
11379 }
11380
11381 #[test]
11383 fn sparse_ffn_quant_equals_dense_with_inactive_zeroed() {
11384 let (hidden, inter) = (16usize, 40usize);
11385 let synth = |n: usize, salt: usize| -> Vec<f32> {
11386 (0..n)
11387 .map(|i| (((i * 37 + salt * 11 + 3) % 101) as f32 / 101.0 - 0.5) * 0.4)
11388 .collect()
11389 };
11390 let d = DenseFfn {
11391 gate_proj: QTensor::from_f32(synth(inter * hidden, 1), inter, hidden),
11392 up_proj: QTensor::from_f32(synth(inter * hidden, 2), inter, hidden),
11393 down_proj: QTensor::from_f32(synth(hidden * inter, 3), hidden, inter),
11394 act: Act::Silu,
11395 down_t: None,
11396 segs: Vec::new(),
11397 };
11398 let x = synth(hidden, 9);
11399 let active: Vec<u16> = (0..inter as u16).filter(|i| i % 3 == 0).collect();
11401
11402 let sparse = sparse_ffn_quant(&d, &x, &active, hidden, None);
11403
11404 let mut g = vec![0.0f32; inter];
11406 d.gate_proj.matvec(&x, &mut g, None);
11407 let mut u = vec![0.0f32; inter];
11408 d.up_proj.matvec(&x, &mut u, None);
11409 let act_set: std::collections::HashSet<u16> = active.iter().copied().collect();
11410 for i in 0..inter {
11411 g[i] = if act_set.contains(&(i as u16)) {
11412 inference::silu(g[i]) * u[i]
11413 } else {
11414 0.0
11415 };
11416 }
11417 let mut reference = vec![0.0f32; hidden];
11418 d.down_proj.matvec(&g, &mut reference, None);
11419
11420 let max_d = sparse
11421 .iter()
11422 .zip(&reference)
11423 .map(|(a, b)| (a - b).abs())
11424 .fold(0.0f32, f32::max);
11425 assert!(max_d < 1e-5, "sparse != dense-zeroed: max|Δ| = {max_d}");
11426 }
11427
11428 fn attach_test_mtp(p: &mut Pipeline) {
11430 let (h, inter, heads, kv, hd) = (
11431 p.hidden_size,
11432 p.intermediate_size,
11433 p.num_heads,
11434 p.num_kv_heads,
11435 p.head_dim,
11436 );
11437 let synth = |n: usize, salt: usize| -> Vec<f32> {
11438 (0..n)
11439 .map(|i| (((i * 29 + salt * 23 + 5) % 101) as f32 / 101.0 - 0.5) * 0.2)
11440 .collect()
11441 };
11442 let qt = |rows: usize, cols: usize, salt: usize| -> QTensor {
11443 QTensor::from_f32(synth(rows * cols, salt), rows, cols)
11444 };
11445 p.mtp = Some(MtpModule {
11446 enorm: vec![1.0; h],
11447 hnorm: vec![1.0; h],
11448 eh_proj: qt(h, 2 * h, 301),
11449 layer: LayerWeights {
11450 input_norm: vec![1.0; h],
11451 post_norm: vec![1.0; h],
11452 attn_out_norm: None,
11453 ffn_out_norm: None,
11454 layer_scale: None,
11455 ffn: FfnKind::Dense(DenseFfn {
11456 gate_proj: qt(inter, h, 315),
11457 up_proj: qt(inter, h, 316),
11458 down_proj: qt(h, inter, 317),
11459 act: Act::Silu,
11460 down_t: None,
11461 segs: Vec::new(),
11462 }),
11463 attn: AttnKind::Full {
11464 bias: None,
11465 wq: qt(heads * hd, h, 311),
11466 wk: qt(kv * hd, h, 312),
11467 wv: qt(kv * hd, h, 313),
11468 wo: qt(h, heads * hd, 314),
11469 q_norm: None,
11470 k_norm: None,
11471 output_gate: false,
11472 softplus_gate: None,
11473 },
11474 },
11475 final_norm: vec![1.0; h],
11476 kv: crate::kv_cache::LayerKvCache::new(kv, hd),
11477 });
11478 }
11479
11480 #[test]
11481 fn speculative_equals_vanilla_greedy() {
11482 unsafe { std::env::set_var("CMF_GPU_WGPU_GRAPH", "0") };
11486 let run = |spec: bool| {
11487 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
11488 p.sampler_config.temperature = 0.0;
11489 attach_test_mtp(&mut p);
11490 p.speculative = spec;
11491 let r = p.generate("abcdef", 12, None, None).unwrap();
11492 (r.token_ids, r.mtp_drafted, r.mtp_accepted)
11493 };
11494 let (vanilla, d0, _) = run(false);
11495 let (spec, d1, a1) = run(true);
11496 assert_eq!(d0, 0, "vanilla path must not draft");
11497 assert!(d1 > 0, "speculative path must draft");
11498 assert_eq!(
11499 vanilla, spec,
11500 "speculative must reproduce the exact greedy sequence (accepted {a1}/{d1})"
11501 );
11502 }
11503
11504 #[test]
11505 fn speculative_accepts_constant_oracle() {
11506 unsafe { std::env::set_var("CMF_GPU_WGPU_GRAPH", "0") };
11508 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
11509 p.sampler_config.temperature = 0.0;
11510 p.sampler_config.repetition_penalty = 1.0;
11511 p.weights.lm_head = QTensor::from_f32(vec![0.01; 64 * 8], 64, 8);
11514 attach_test_mtp(&mut p);
11515 p.speculative = true;
11516 let r = p.generate("abcd", 10, None, None).unwrap();
11517 assert!(r.mtp_drafted > 0);
11518 assert_eq!(
11519 r.mtp_accepted, r.mtp_drafted,
11520 "constant logits → every draft accepted"
11521 );
11522 assert!(r.token_ids.windows(2).all(|w| w[0] == w[1]));
11525 }
11526
11527 #[test]
11528 fn empty_prompt_is_an_error_not_a_panic() {
11529 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 260);
11530 let r = p.generate("", 4, None, None);
11531 assert!(r.is_err(), "empty prompt must be a clean error");
11532 }
11533
11534 #[test]
11535 fn every_token_enters_kv_exactly_once() {
11536 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
11537 p.sampler_config.temperature = 0.0;
11539 let r = p.generate("abc", 2, None, None).unwrap();
11540 assert_eq!(r.prompt_tokens, 3);
11541 assert_eq!(
11545 p.kv_cache.seq_len(),
11546 3 + r.tokens_generated - 1,
11547 "each token must be cached exactly once (v1 cached the last prompt token twice)"
11548 );
11549 }
11550
11551 #[test]
11552 fn generation_is_reproducible_with_seed() {
11553 let run = || {
11554 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
11555 p.generate("hello", 8, None, None).unwrap().token_ids
11556 };
11557 assert_eq!(run(), run());
11558 }
11559
11560 #[test]
11561 fn resetting_sampler_restarts_the_seeded_stream() {
11562 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
11563 let config = SamplerConfig {
11564 seed: Some(1234),
11565 ..SamplerConfig::default()
11566 };
11567 p.set_sampler_config(config.clone());
11568 let first = p.generate("hello", 8, None, None).unwrap().token_ids;
11569 p.set_sampler_config(config);
11570 let second = p.generate("hello", 8, None, None).unwrap().token_ids;
11571 assert_eq!(first, second);
11572 }
11573
11574 #[test]
11575 fn eviction_bounds_the_cache() {
11576 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 260);
11577 p.kv_cache.max_seq_len = 6;
11578 p.sampler_config.temperature = 0.0;
11579 let _ = p.generate("abcd", 12, None, None).unwrap();
11580 assert!(
11581 p.kv_cache.seq_len() <= 6 + 1,
11582 "cache must stay bounded by max_seq_len (got {})",
11583 p.kv_cache.seq_len()
11584 );
11585 }
11586
11587 #[test]
11588 fn confidence_matches_tokens_and_is_a_probability() {
11589 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
11590 p.sampler_config.temperature = 0.0;
11591 p.sampler_config.repetition_penalty = 1.0;
11592 let r = p.generate("abcd", 10, None, None).unwrap();
11593 assert_eq!(
11594 r.token_confidence.len(),
11595 r.token_ids.len(),
11596 "one confidence per emitted token"
11597 );
11598 for &c in &r.token_confidence {
11599 assert!((0.0..=1.0).contains(&c), "confidence out of [0,1]: {c}");
11600 }
11601 let logits = [1.0f32, 3.0, 0.5, 3.0];
11603 let p0 = top1_prob_t(&logits, 1, 1.0);
11604 let p1 = top1_prob_t(&logits, 3, 1.0);
11605 assert!((p0 - p1).abs() < 1e-6, "equal logits → equal prob");
11606 assert!(p0 > 0.0 && p0 < 1.0);
11607 let sharp = top1_prob_t(&logits, 1, 1.0);
11609 let soft = top1_prob_t(&logits, 1, 2.0);
11610 assert!(soft < sharp, "higher temperature lowers peak confidence");
11611 }
11612
11613 #[test]
11614 fn trace_is_opt_in_and_parallels_the_output() {
11615 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
11617 p.sampler_config.temperature = 0.0;
11618 p.sampler_config.repetition_penalty = 1.0;
11619 let r = p.generate("abcd", 10, None, None).unwrap();
11620 assert!(r.traces.is_empty(), "trace must be empty unless enabled");
11621
11622 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
11624 p.sampler_config.temperature = 0.0;
11625 p.sampler_config.repetition_penalty = 1.0;
11626 p.set_trace(true);
11627 let r = p.generate("abcd", 10, None, None).unwrap();
11628 assert_eq!(r.traces.len(), r.token_ids.len(), "one trace row per token");
11629 for (i, tr) in r.traces.iter().enumerate() {
11630 assert_eq!(tr.t, i, "trace index is sequential");
11631 assert_eq!(tr.token_id, r.token_ids[i], "trace token_id matches output");
11632 assert_eq!(
11633 tr.confidence, r.token_confidence[i],
11634 "trace confidence matches the confidence channel"
11635 );
11636 assert!(tr.active_skill.is_none() && tr.recon.is_none() && !tr.switched);
11638 }
11639 }
11640
11641 #[test]
11642 fn explain_prefill_logits_match_greedy_first_token() {
11643 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
11647 p.sampler_config.temperature = 0.0;
11648 p.sampler_config.repetition_penalty = 1.0;
11649 let ids = p.tokenizer.encode("abcd");
11650 let logits = p.prefill_next_logits(&ids, None);
11651 let argmax = logits
11652 .iter()
11653 .enumerate()
11654 .max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
11655 .unwrap()
11656 .0 as u32;
11657 let r = p.generate("abcd", 1, None, None).unwrap();
11658 assert_eq!(
11659 argmax, r.token_ids[0],
11660 "explain preview must match greedy emit"
11661 );
11662 }
11663
11664 #[test]
11665 fn laguna_shared_expert_is_unconditionally_added() {
11666 let matrix = |values: Vec<f32>| QTensor::from_f32(values, 2, 2);
11667 let identity = || matrix(vec![1.0, 0.0, 0.0, 1.0]);
11668 let zero_dense = || DenseFfn {
11669 gate_proj: matrix(vec![0.0; 4]),
11670 up_proj: matrix(vec![0.0; 4]),
11671 down_proj: matrix(vec![0.0; 4]),
11672 act: Act::Silu,
11673 down_t: None,
11674 segs: Vec::new(),
11675 };
11676 let shared = DenseFfn {
11677 gate_proj: identity(),
11678 up_proj: identity(),
11679 down_proj: identity(),
11680 act: Act::Silu,
11681 down_t: None,
11682 segs: Vec::new(),
11683 };
11684 let x = [1.0, 2.0];
11685 let expected = dense_ffn(&shared, &x, None);
11686 let moe = MoeFfn {
11687 router: QTensor::from_f32(vec![0.0, 0.0], 1, 2),
11688 experts: vec![zero_dense()],
11689 top_k: 1,
11690 norm_topk_prob: true,
11691 router_sigmoid: true,
11692 expert_bias: None,
11693 routed_scaling: 1.0,
11694 route_tau: None,
11695 shared: Some((shared, None)),
11696 stats: std::cell::RefCell::new(Vec::new()),
11697 act_sq: std::cell::RefCell::new(Vec::new()),
11698 act_rows: std::cell::RefCell::new(Vec::new()),
11699 mask: None,
11700 per_expert_scale: None,
11701 router_input_norm: false,
11702 resonance: None,
11703 };
11704 let actual = moe_ffn_cpu(&moe, &x, &[0], &[0.0], 1.0, None);
11705 for (actual, expected) in actual.iter().zip(expected) {
11706 assert!((actual - expected).abs() < 1e-6);
11707 }
11708 }
11709}