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 AttnKind::ShortConv(w) => {
6385 let cfg = self.short_conv_cfg?;
6386 let (m, _, _, _) = w.in_proj.graph_weight()?;
6387 model = Some(m.clone());
6388 crate::gpu::GraphAttn::ShortConv {
6389 inp: gw(&w.in_proj)?,
6390 out: gw(&w.out_proj)?,
6391 taps: &w.conv,
6392 kernel: cfg.kernel,
6393 cpu_state: &self.kv_cache.layers[self.phys_layer(li)].linear_state,
6394 }
6395 }
6396 _ => return None,
6397 };
6398 layers.push(crate::gpu::GraphLayer {
6399 input_norm: &lw.input_norm,
6400 attn,
6401 post_norm: &lw.post_norm,
6402 ffn: gffn,
6403 });
6404 }
6405 let model = model?;
6406 let lm_gw = if upto_excl == self.num_layers
6412 && self.graph_want_logits
6413 && std::env::var("CMF_GPU_LMHEAD")
6414 .map(|v| v != "0")
6415 .unwrap_or(true)
6416 {
6417 self.weights.lm_head.graph_weight().map(|(_, i, kind, rs)| {
6418 (
6419 crate::gpu::GraphW {
6420 idx: i,
6421 kind,
6422 row_scale: rs,
6423 data: &[],
6424 },
6425 self.weights.lm_head.rows(),
6426 )
6427 })
6428 } else {
6429 None
6430 };
6431 let lm = lm_gw.as_ref().map(|(gw, rows)| (gw, *rows));
6432 let emb_gw = if steps > 1 {
6434 self.weights
6435 .embed_tokens
6436 .graph_weight()
6437 .map(|(_, i, kind, rs)| {
6438 (
6439 crate::gpu::GraphW {
6440 idx: i,
6441 kind,
6442 row_scale: rs,
6443 data: &[],
6444 },
6445 self.weights.embed_tokens.rows(),
6446 self.embed_multiplier,
6447 )
6448 })
6449 } else {
6450 None
6451 };
6452
6453 let loop_norm_at: Vec<usize> = if self.loop_final_norm {
6459 (from..upto_excl.min(self.num_layers - 1))
6460 .filter(|&li| (li + 1) % self.physical_layers == 0)
6461 .map(|li| li - from)
6462 .collect()
6463 } else {
6464 Vec::new()
6465 };
6466 let mut h = hidden.to_vec();
6467 crate::gpu::forward_token_graph(
6468 &model,
6469 self.graph_kv_id,
6470 &layers,
6471 &o1_views,
6472 self.o1_epoch,
6473 &self.inv_freq,
6474 &mut h,
6475 nh,
6476 nkv,
6477 hd,
6478 rd,
6479 self.hidden_size,
6480 self.intermediate_size,
6481 position,
6482 self.kv_cache.max_seq_len,
6483 gemma,
6484 self.rms_eps as f32,
6485 lm,
6486 &self.weights.final_norm,
6487 logits_out,
6488 &loop_norm_at,
6489 steps,
6490 emb_gw.as_ref().map(|(gw, rows, m)| (gw, *rows, *m)),
6491 ids_out,
6492 layers_run,
6493 from,
6494 false,
6495 )
6496 .then_some(h)
6497 }
6498
6499 #[cfg(target_os = "macos")]
6508 #[allow(clippy::type_complexity)]
6509 fn metal_rows_plan(&self) -> Option<(Vec<MetalRowsItem<'_>>, std::sync::Arc<cortiq_core::CmfModel>, Option<crate::gpu_metal::GdnGpuCfg>)> {
6510 use crate::gpu_metal::{AttnGpuLayer, GdnGpuCfg, GdnGpuLayer, MetalFfn};
6511 if !crate::gpu::q1_force()
6512 || !crate::gpu::enabled_here()
6513 || std::env::var("CMF_GPU_BLOCK").map(|v| v == "0").unwrap_or(false)
6514 || self.attn_softcap > 0.0
6515 || self.o1_active()
6516 || self.swa.is_some()
6517 || self.global_attn.is_some()
6518 || self.attention_heads_per_layer.is_some()
6519 || self.attn_v_norm
6520 || self.loop_final_norm
6521 || (self.attn_scale - 1.0 / (self.head_dim as f32).sqrt()).abs() > 1e-9
6522 {
6523 return None;
6524 }
6525 let attend_contract = self.head_dim % 4 == 0
6526 && self.head_dim <= 256
6527 && self.rotary_dim >= 2
6528 && self.rotary_dim <= self.head_dim
6529 && (self.rotary_dim / 2) % 32 == 0
6530 && self.num_kv_heads > 0
6531 && self.num_heads % self.num_kv_heads == 0;
6532 if !attend_contract {
6533 return None;
6534 }
6535 let mut plan: Vec<MetalRowsItem> = Vec::new();
6536 let mut model_ref: Option<std::sync::Arc<cortiq_core::CmfModel>> = None;
6537 for li in 0..self.num_layers {
6538 let lw = &self.weights.layers[self.phys_layer(li)];
6539 if lw.attn_out_norm.is_some() || lw.ffn_out_norm.is_some() || lw.layer_scale.is_some() {
6540 return None;
6541 }
6542 let ffn = match &lw.ffn {
6543 FfnKind::Dense(d) if d.act == Act::Silu && d.segs.is_empty() => {
6544 let (Some(g), Some(u), Some(dn)) =
6545 (d.gate_proj.q1_parts(), d.up_proj.q1_parts(), d.down_proj.q1_parts())
6546 else {
6547 return None;
6548 };
6549 MetalFfn::Dense { gate: g, up: u, down: dn }
6550 }
6551 _ => return None,
6552 };
6553 match &lw.attn {
6554 AttnKind::LinearGdn(w) if self.gdn_cfg.is_some() => {
6555 let (Some(qkv), Some(z), Some(a), Some(bb), Some(out)) = (
6556 w.in_proj_qkv.q1_parts(),
6557 w.in_proj_z.q1_parts(),
6558 w.in_proj_a.f32_parts(),
6559 w.in_proj_b.f32_parts(),
6560 w.out_proj.q1_parts(),
6561 ) else {
6562 return None;
6563 };
6564 if let QTensor::Mapped { model, .. } = &w.in_proj_qkv {
6565 model_ref.get_or_insert_with(|| model.clone());
6566 }
6567 let gl = GdnGpuLayer {
6568 attn_norm: &lw.input_norm,
6569 post_norm: &lw.post_norm,
6570 qkv,
6571 z,
6572 a,
6573 b: bb,
6574 out,
6575 ffn,
6576 conv1d: &w.conv1d,
6577 a_log: &w.a_log,
6578 dt_bias: &w.dt_bias,
6579 gnorm: &w.norm,
6580 };
6581 match plan.last_mut() {
6582 Some(MetalRowsItem::Gdn { run, .. }) => run.push(gl),
6583 _ => plan.push(MetalRowsItem::Gdn { run: vec![gl], first: li }),
6584 }
6585 }
6586 AttnKind::Full {
6587 wq,
6588 wk,
6589 wv,
6590 wo,
6591 q_norm,
6592 k_norm,
6593 output_gate,
6594 softplus_gate: None,
6595 bias: None,
6596 } => {
6597 let (Some(pq), Some(pk), Some(pv), Some(po)) =
6598 (wq.q1_parts(), wk.q1_parts(), wv.q1_parts(), wo.q1_parts())
6599 else {
6600 return None;
6601 };
6602 if let QTensor::Mapped { model, .. } = wq {
6603 model_ref.get_or_insert_with(|| model.clone());
6604 }
6605 let cache = &self.kv_cache.layers[li];
6606 if cache.mode != crate::kv_cache::KvMode::F32 || cache.o1.is_some() {
6607 return None;
6608 }
6609 plan.push(MetalRowsItem::Attn {
6610 l: AttnGpuLayer {
6611 attn_norm: &lw.input_norm,
6612 post_norm: &lw.post_norm,
6613 wq: pq,
6614 wk: pk,
6615 wv: pv,
6616 wo: po,
6617 ffn,
6618 },
6619 li,
6620 q_norm: q_norm.as_deref(),
6621 k_norm: k_norm.as_deref(),
6622 output_gate: *output_gate,
6623 });
6624 }
6625 _ => return None,
6626 }
6627 }
6628 let model = model_ref?;
6629 let gcfg = self.gdn_cfg.map(|cfg| GdnGpuCfg {
6630 nv: cfg.num_v_heads,
6631 nk: cfg.num_k_heads,
6632 dk: cfg.key_head_dim,
6633 dv: cfg.value_head_dim,
6634 kk: cfg.conv_kernel,
6635 hidden: self.hidden_size,
6636 inter: self.intermediate_size,
6637 c_dim: cfg.conv_dim(),
6638 eps: cfg.rms_eps as f32,
6639 gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
6640 });
6641 Some((plan, model, gcfg))
6642 }
6643
6644 #[cfg(target_os = "macos")]
6646 #[allow(clippy::too_many_arguments)]
6647 fn metal_attn_params<'a>(
6648 li: usize,
6649 cache: &'a crate::kv_cache::LayerKvCache,
6650 q_norm: Option<&'a [f32]>,
6651 k_norm: Option<&'a [f32]>,
6652 output_gate: bool,
6653 inv_freq: &'a [f32],
6654 geom: (usize, usize, usize, usize),
6655 pos0: usize,
6656 kv_id: u64,
6657 eps: f32,
6658 gemma: bool,
6659 ) -> (crate::gpu_metal::AttnDeviceParams<'a>, usize) {
6660 let (nh, nkv, hd, rd) = geom;
6661 let cpu_k: Vec<&[f32]> = (0..nkv).map(|g| cache.head_keys(g)).collect();
6662 let cpu_v: Vec<&[f32]> = (0..nkv).map(|g| cache.head_values(g)).collect();
6663 let cpu_stored = cpu_k[0].len() / hd;
6664 (
6665 crate::gpu_metal::AttnDeviceParams {
6666 kv_id,
6667 layer: li,
6668 nh,
6669 nkv,
6670 hd,
6671 rd,
6672 position: pos0,
6673 eps,
6674 gemma,
6675 output_gate,
6676 q_norm,
6677 k_norm,
6678 inv_freq,
6679 cpu_k,
6680 cpu_v,
6681 cpu_stored,
6682 o1: None,
6683 },
6684 cpu_stored,
6685 )
6686 }
6687
6688 #[cfg(target_os = "macos")]
6693 #[allow(clippy::type_complexity)]
6694 fn metal_rows_run(
6695 &mut self,
6696 hiddens: &mut [f32],
6697 pos0: usize,
6698 b: usize,
6699 prefill: bool,
6700 spec: Option<((usize, usize, usize), &[f32], &mut Vec<f32>)>,
6701 ) -> Option<MetalVerifyPending> {
6702 use crate::gpu_metal::{GraphDims, VerifyGraph};
6703 let want = self.gdn_cfg.map(|c| c.state_len()).unwrap_or(0);
6704 for l in &mut self.kv_cache.layers {
6705 if l.linear_state.len() != want && want > 0 {
6706 l.linear_state = vec![0f32; want];
6707 }
6708 }
6709 let (plan, model, gcfg) = self.metal_rows_plan()?;
6710 let dims = GraphDims {
6711 hidden: self.hidden_size,
6712 eps: self.rms_eps as f32,
6713 gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
6714 };
6715 let mut graph = if prefill {
6716 VerifyGraph::new_prefill(&model, dims, hiddens, b)?
6717 } else {
6718 VerifyGraph::new(&model, dims, hiddens, b)?
6719 };
6720 let geom = (self.num_heads, self.num_kv_heads, self.head_dim, self.rotary_dim);
6721 let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
6722 let eps = self.rms_eps as f32;
6723 let kv_id = self.graph_kv_id;
6724 let inv_freq = self.inv_freq.clone();
6725 for item in &plan {
6726 let ok = match item {
6727 MetalRowsItem::Gdn { run, .. } => gcfg
6728 .as_ref()
6729 .map(|gc| run.iter().all(|l| graph.gdn_ok(l, gc)))
6730 .unwrap_or(false),
6731 MetalRowsItem::Attn { l, li, q_norm, k_norm, output_gate } => {
6732 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);
6733 graph.attn_ok(l, &p)
6734 }
6735 };
6736 if !ok {
6737 use std::sync::atomic::{AtomicBool, Ordering};
6738 static SAID: AtomicBool = AtomicBool::new(false);
6739 if !SAID.swap(true, Ordering::Relaxed) {
6740 tracing::warn!("metal rows graph: a layer failed preflight — declining");
6741 }
6742 return None;
6743 }
6744 }
6745 let lm = match &spec {
6746 Some((lm, _, _)) => {
6747 if !graph.lm_head_ok(*lm) {
6748 return None;
6749 }
6750 Some(*lm)
6751 }
6752 None => None,
6753 };
6754 let mut gdn_layers = Vec::new();
6755 let mut attn_layers = Vec::new();
6756 for item in &plan {
6757 match item {
6758 MetalRowsItem::Gdn { run, first } => {
6759 let ro: Vec<&[f32]> = self.kv_cache.layers[*first..*first + run.len()]
6760 .iter()
6761 .map(|l| l.linear_state.as_slice())
6762 .collect();
6763 if !graph.encode_gdn_run_b(run, &ro, gcfg.as_ref().unwrap()) {
6764 return None;
6765 }
6766 gdn_layers.extend(*first..*first + run.len());
6767 }
6768 MetalRowsItem::Attn { l, li, q_norm, k_norm, output_gate } => {
6769 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);
6770 if !graph.encode_attn_b(l, &p) {
6771 return None;
6772 }
6773 attn_layers.push((*li, cpu_stored));
6774 }
6775 }
6776 }
6777 if let (Some(lm), Some((_, final_norm, _))) = (lm, spec.as_ref()) {
6778 if !graph.encode_lm_head_b(final_norm, lm) {
6779 return None;
6780 }
6781 }
6782 graph.sync();
6783 if let Some((lm, _, logits)) = spec {
6784 logits.resize(b * lm.1, 0.0);
6785 graph.read_logits(logits);
6786 }
6787 graph.read_hidden(hiddens);
6788 Some(MetalVerifyPending { graph, gdn_layers, attn_layers })
6789 }
6790
6791 #[cfg(target_os = "macos")]
6797 fn try_batch_graph_metal(
6798 &mut self,
6799 hiddens: &mut [f32],
6800 positions: &[usize],
6801 b: usize,
6802 spec: Option<((usize, usize, usize), &[f32], &mut Vec<f32>)>,
6803 ) -> bool {
6804 let _t0 = std::time::Instant::now();
6805 if positions.len() != b
6806 || positions.windows(2).any(|w| w[1] != w[0] + 1)
6807 || hiddens.len() != b * self.hidden_size
6808 {
6809 return false;
6810 }
6811 let Some(pending) = self.metal_rows_run(hiddens, positions[0], b, false, spec) else {
6812 return false;
6813 };
6814 if std::env::var("CMF_GRAPH_SPEC_TIME").is_ok() {
6815 eprintln!("metal-verify: {:.1} ms | b={b}", _t0.elapsed().as_secs_f64() * 1e3);
6816 }
6817 self.metal_verify = Some(pending);
6818 true
6819 }
6820
6821 #[cfg(target_os = "macos")]
6826 fn prefill_batch_metal(&mut self, ids: &[u32], start_pos: usize) -> Option<Vec<f32>> {
6827 let b = ids.len();
6828 if b == 0 || b > 512 {
6829 return None;
6830 }
6831 let hs = self.hidden_size;
6832 let mut hiddens = vec![0f32; b * hs];
6833 for (j, &id) in ids.iter().enumerate() {
6834 let e = self.embed_single(id);
6835 hiddens[j * hs..(j + 1) * hs].copy_from_slice(&e);
6836 }
6837 let mut pending = self.metal_rows_run(&mut hiddens, start_pos, b, true, None)?;
6838 let idxs = pending.gdn_layers.clone();
6840 let mut outs: Vec<&mut [f32]> = self
6841 .kv_cache
6842 .layers
6843 .iter_mut()
6844 .enumerate()
6845 .filter(|(i, _)| idxs.binary_search(i).is_ok())
6846 .map(|(_, l)| l.linear_state.as_mut_slice())
6847 .collect();
6848 pending.graph.finish_states(&mut outs);
6849 let (nkv, hd) = (self.num_kv_heads, self.head_dim);
6850 let mut kbuf = vec![0f32; b * nkv * hd];
6851 let mut vbuf = vec![0f32; b * nkv * hd];
6852 for (li, cpu_stored) in &pending.attn_layers {
6853 if crate::gpu_metal::kv_mirror_read_rows(self.graph_kv_id, *li, nkv, hd, *cpu_stored, b, &mut kbuf, &mut vbuf) {
6854 let cache = &mut self.kv_cache.layers[*li];
6855 for r in 0..b {
6856 cache.append(&kbuf[r * nkv * hd..(r + 1) * nkv * hd], &vbuf[r * nkv * hd..(r + 1) * nkv * hd], &[]);
6857 }
6858 crate::gpu_metal::kv_mirror_set_stored(self.graph_kv_id, *li, cpu_stored + b);
6859 }
6860 }
6861 Some(hiddens)
6862 }
6863
6864 #[cfg(target_os = "macos")]
6868 fn metal_verify_commit(&mut self, a: usize) -> bool {
6869 let Some(mut pending) = self.metal_verify.take() else {
6870 return false;
6871 };
6872 let n = a + 1;
6873 let idxs = pending.gdn_layers.clone();
6875 let mut outs: Vec<&mut [f32]> = self
6876 .kv_cache
6877 .layers
6878 .iter_mut()
6879 .enumerate()
6880 .filter(|(i, _)| idxs.binary_search(i).is_ok())
6881 .map(|(_, l)| l.linear_state.as_mut_slice())
6882 .collect();
6883 if !pending.graph.commit(n, &mut outs) {
6884 return false;
6885 }
6886 let (nkv, hd) = (self.num_kv_heads, self.head_dim);
6887 let mut kbuf = vec![0f32; n * nkv * hd];
6888 let mut vbuf = vec![0f32; n * nkv * hd];
6889 for (li, cpu_stored) in &pending.attn_layers {
6890 if crate::gpu_metal::kv_mirror_read_rows(self.graph_kv_id, *li, nkv, hd, *cpu_stored, n, &mut kbuf, &mut vbuf) {
6891 let cache = &mut self.kv_cache.layers[*li];
6892 for r in 0..n {
6893 cache.append(&kbuf[r * nkv * hd..(r + 1) * nkv * hd], &vbuf[r * nkv * hd..(r + 1) * nkv * hd], &[]);
6894 }
6895 crate::gpu_metal::kv_mirror_set_stored(self.graph_kv_id, *li, cpu_stored + n);
6896 }
6897 }
6898 true
6899 }
6900
6901 #[cfg(target_os = "macos")]
6907 fn mtp_warm_batch_metal(&mut self, m: &mut MtpModule, pairs: &[(&[f32], u32)], first_pos: usize) -> bool {
6908 use crate::gpu_metal::{AttnDeviceParams, AttnGpuLayer, GraphDims, MetalFfn, VerifyGraph};
6909 let b = pairs.len();
6910 if b == 0 || b > 512 || m.kv.mode != crate::kv_cache::KvMode::F32 || m.kv.o1.is_some() {
6911 return false;
6912 }
6913 let AttnKind::Full { wq, wk, wv, wo, q_norm, k_norm, output_gate, softplus_gate: None, bias: None } = &m.layer.attn else {
6914 return false;
6915 };
6916 let FfnKind::Dense(d) = &m.layer.ffn else { return false };
6917 if !d.segs.is_empty() {
6918 return false;
6919 }
6920 let (Some(pq), Some(pk), Some(pv), Some(po)) = (wq.q1_parts(), wk.q1_parts(), wv.q1_parts(), wo.q1_parts()) else {
6921 return false;
6922 };
6923 let (Some(g), Some(u), Some(dn)) = (d.gate_proj.q1_parts(), d.up_proj.q1_parts(), d.down_proj.q1_parts()) else {
6924 return false;
6925 };
6926 let Some(eh) = m.eh_proj.q1_parts() else { return false };
6927 let QTensor::Mapped { model, .. } = wq else { return false };
6928 let model = model.clone();
6929 let hs = self.hidden_size;
6930 let mut cat = vec![0f32; b * 2 * hs];
6932 for (j, (h, tok)) in pairs.iter().enumerate() {
6933 let e = self.embed_single(*tok);
6934 let (ce, ch) = cat[j * 2 * hs..(j + 1) * 2 * hs].split_at_mut(hs);
6935 inference::rms_norm_into(&e, &m.enorm, self.rms_eps, self.norm_style, ce);
6936 inference::rms_norm_into(h, &m.hnorm, self.rms_eps, self.norm_style, ch);
6937 }
6938 let dims = GraphDims { hidden: hs, eps: self.rms_eps as f32, gemma: self.norm_style == cortiq_core::NormStyle::Gemma };
6939 let Some(mut graph) = VerifyGraph::new_via_proj(&model, dims, eh, &cat, b) else {
6940 return false;
6941 };
6942 let l = AttnGpuLayer {
6943 attn_norm: &m.layer.input_norm,
6944 post_norm: &m.layer.post_norm,
6945 wq: pq,
6946 wk: pk,
6947 wv: pv,
6948 wo: po,
6949 ffn: MetalFfn::Dense { gate: g, up: u, down: dn },
6950 };
6951 let (nh, nkv, hd, rd) = (self.num_heads, self.num_kv_heads, self.head_dim, self.rotary_dim);
6952 let inv_freq = self.inv_freq.clone();
6953 let cpu_stored;
6954 {
6955 let cache = &m.kv;
6956 let cpu_k: Vec<&[f32]> = (0..nkv).map(|g| cache.head_keys(g)).collect();
6957 let cpu_v: Vec<&[f32]> = (0..nkv).map(|g| cache.head_values(g)).collect();
6958 cpu_stored = cpu_k[0].len() / hd;
6959 if cpu_stored != first_pos {
6960 return false;
6961 }
6962 let p = AttnDeviceParams {
6963 kv_id: self.mtp_kv_id(),
6964 layer: Self::MTP_LAYER_BASE,
6965 nh,
6966 nkv,
6967 hd,
6968 rd,
6969 position: first_pos,
6970 eps: self.rms_eps as f32,
6971 gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
6972 output_gate: *output_gate,
6973 q_norm: q_norm.as_deref(),
6974 k_norm: k_norm.as_deref(),
6975 inv_freq: &inv_freq,
6976 cpu_k,
6977 cpu_v,
6978 cpu_stored,
6979 o1: None,
6980 };
6981 if !graph.attn_ok(&l, &p) || !graph.encode_attn_b(&l, &p) {
6982 return false;
6983 }
6984 }
6985 graph.sync();
6986 let mut kbuf = vec![0f32; b * nkv * hd];
6987 let mut vbuf = vec![0f32; b * nkv * hd];
6988 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) {
6989 return false;
6990 }
6991 for r in 0..b {
6992 m.kv.append(&kbuf[r * nkv * hd..(r + 1) * nkv * hd], &vbuf[r * nkv * hd..(r + 1) * nkv * hd], &[]);
6993 }
6994 crate::gpu_metal::kv_mirror_set_stored(self.mtp_kv_id(), Self::MTP_LAYER_BASE, cpu_stored + b);
6995 true
6996 }
6997
6998 fn draft_vocab_rows(head_rows: usize) -> usize {
7001 static N: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
7002 let n = *N.get_or_init(|| {
7003 std::env::var("CMF_DRAFT_VOCAB")
7004 .ok()
7005 .and_then(|v| v.parse().ok())
7006 .unwrap_or(65536)
7007 });
7008 if n == 0 { head_rows } else { n.min(head_rows) }
7009 }
7010
7011 #[cfg(target_os = "macos")]
7016 fn mtp_step_metal(
7017 &mut self,
7018 m: &mut MtpModule,
7019 hidden: &[f32],
7020 next_token: u32,
7021 position: usize,
7022 want_logits: bool,
7023 ) -> Option<(Vec<f32>, Vec<f32>)> {
7024 use crate::gpu_metal::{AttnDeviceParams, AttnGpuLayer, GraphDims, MetalFfn, TokenGraph};
7025 if std::env::var("CMF_MTP_GRAPH").as_deref() == Ok("0")
7026 || !crate::gpu::q1_force()
7027 || !crate::gpu::enabled_here()
7028 || self.attn_softcap > 0.0
7029 || self.attention_heads_per_layer.is_some()
7030 || m.kv.mode != crate::kv_cache::KvMode::F32
7031 || m.kv.o1.is_some()
7032 {
7033 return None;
7034 }
7035 let AttnKind::Full {
7036 wq,
7037 wk,
7038 wv,
7039 wo,
7040 q_norm,
7041 k_norm,
7042 output_gate,
7043 softplus_gate: None,
7044 bias: None,
7045 } = &m.layer.attn
7046 else {
7047 return None;
7048 };
7049 let FfnKind::Dense(d) = &m.layer.ffn else { return None };
7050 if d.act != Act::Silu || !d.segs.is_empty() {
7051 return None;
7052 }
7053 let (pq, pk, pv, po) = (wq.q1_parts()?, wk.q1_parts()?, wv.q1_parts()?, wo.q1_parts()?);
7054 let (g, u, dn) = (d.gate_proj.q1_parts()?, d.up_proj.q1_parts()?, d.down_proj.q1_parts()?);
7055 let QTensor::Mapped { model, .. } = wq else { return None };
7056 let model = model.clone();
7057 let lm = if want_logits { Some(self.weights.lm_head.q1_parts()?) } else { None };
7058 let dims = GraphDims {
7059 hidden: self.hidden_size,
7060 eps: self.rms_eps as f32,
7061 gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
7062 };
7063 let hs = self.hidden_size;
7066 let mut x = vec![0f32; hs];
7067 let mut graph = TokenGraph::new(&model, dims, &x)?;
7068 let mut folded = false;
7069 if let Some(eh) = m.eh_proj.q1_parts() {
7070 let e = self.embed_single(next_token);
7071 let mut cat = vec![0.0f32; 2 * hs];
7072 let (cat_e, cat_h) = cat.split_at_mut(hs);
7073 inference::rms_norm_into(&e, &m.enorm, self.rms_eps, self.norm_style, cat_e);
7074 inference::rms_norm_into(hidden, &m.hnorm, self.rms_eps, self.norm_style, cat_h);
7075 folded = graph.encode_input_proj(eh, &cat);
7076 }
7077 if !folded {
7078 x = self.mtp_block_input(m, hidden, next_token);
7079 graph = TokenGraph::new(&model, dims, &x)?;
7080 }
7081 let l = AttnGpuLayer {
7082 attn_norm: &m.layer.input_norm,
7083 post_norm: &m.layer.post_norm,
7084 wq: pq,
7085 wk: pk,
7086 wv: pv,
7087 wo: po,
7088 ffn: MetalFfn::Dense { gate: g, up: u, down: dn },
7089 };
7090 let (nh, nkv, hd, rd) = (self.num_heads, self.num_kv_heads, self.head_dim, self.rotary_dim);
7091 let inv_freq = self.inv_freq.clone();
7092 {
7093 let cache = &m.kv;
7094 let cpu_k: Vec<&[f32]> = (0..nkv).map(|g| cache.head_keys(g)).collect();
7095 let cpu_v: Vec<&[f32]> = (0..nkv).map(|g| cache.head_values(g)).collect();
7096 let cpu_stored = cpu_k[0].len() / hd;
7097 let p = AttnDeviceParams {
7098 kv_id: self.mtp_kv_id(),
7099 layer: Self::MTP_LAYER_BASE,
7100 nh,
7101 nkv,
7102 hd,
7103 rd,
7104 position,
7105 eps: self.rms_eps as f32,
7106 gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
7107 output_gate: *output_gate,
7108 q_norm: q_norm.as_deref(),
7109 k_norm: k_norm.as_deref(),
7110 inv_freq: &inv_freq,
7111 cpu_k,
7112 cpu_v,
7113 cpu_stored,
7114 o1: None,
7115 };
7116 if !graph.attn_device_ok(&l, &p) || !graph.encode_attn_device(&l, &p) {
7117 return None;
7118 }
7119 }
7120 let draft_rows = if let Some(lm) = lm { Self::draft_vocab_rows(lm.1) } else { 0 };
7126 if let Some(lm) = lm {
7127 if !graph.lm_head_ok(lm) {
7128 return None;
7129 }
7130 if draft_rows < lm.1 {
7131 if !graph.encode_lm_head_part(&m.final_norm, lm, draft_rows) {
7132 return None;
7133 }
7134 } else {
7135 graph.encode_lm_head(&m.final_norm, lm);
7136 }
7137 }
7138 graph.sync();
7139 let mut logits = Vec::new();
7140 if let Some(lm) = lm {
7141 let n_read = draft_rows.min(lm.1).min(self.vocab_size);
7142 logits = attention::take_buf(n_read);
7143 graph.read_logits(&mut logits);
7144 logits.resize(self.vocab_size, f32::NEG_INFINITY);
7146 }
7147 graph.finish(&mut x);
7148 let mut krow = attention::take_buf(nkv * hd);
7149 let mut vrow = attention::take_buf(nkv * hd);
7150 if crate::gpu_metal::kv_mirror_read_last(self.mtp_kv_id(), Self::MTP_LAYER_BASE, nkv, hd, &mut krow, &mut vrow) {
7151 m.kv.append(&krow, &vrow, &[]);
7152 }
7153 attention::recycle_buf(&mut krow);
7154 attention::recycle_buf(&mut vrow);
7155 Some((logits, x))
7156 }
7157
7158 fn try_batch_graph_wgpu(
7159 &self,
7160 hiddens: &mut [f32],
7161 positions: &[usize],
7162 k: usize,
7163 spec: Option<crate::gpu::SpecTail<'_>>,
7164 ) -> bool {
7165 let _tb = std::time::Instant::now();
7166 if self.attn_softcap > 0.0 {
7167 return false; }
7169 if self.o1_active() {
7170 return false;
7171 }
7172 let nh = self.num_heads;
7173 let (nkv, hd, rd) = self.layer_geom(0);
7174 let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
7175 fn gw(t: &QTensor) -> Option<crate::gpu::GraphW<'_>> {
7176 if let Some((_, i, kind, rs)) = t.graph_weight() {
7177 return Some(crate::gpu::GraphW {
7178 idx: i,
7179 kind,
7180 row_scale: rs,
7181 data: &[],
7182 });
7183 }
7184 t.as_f32().map(|d| crate::gpu::GraphW {
7185 idx: 0,
7186 kind: 4,
7187 row_scale: &[],
7188 data: d,
7189 })
7190 }
7191 let built: Option<(
7192 Vec<crate::gpu::GraphLayer<'_>>,
7193 std::sync::Arc<cortiq_core::CmfModel>,
7194 )> = (|| {
7195 let mut layers = Vec::with_capacity(self.num_layers);
7196 let mut model = None;
7197 for li in 0..self.num_layers {
7198 let lw = &self.weights.layers[self.phys_layer(li)];
7199 let gffn = match &lw.ffn {
7206 FfnKind::Dense(d) if !d.segs.is_empty() => return None,
7207 FfnKind::Dense(d) => crate::gpu::GraphFfn::Dense {
7208 gate: gw(&d.gate_proj)?,
7209 up: gw(&d.up_proj)?,
7210 down: gw(&d.down_proj)?,
7211 },
7212 FfnKind::Moe(m) => {
7213 if m.router_sigmoid
7214 || m.expert_bias.is_some()
7215 || m.route_tau.is_some()
7216 || m.mask.is_some()
7217 {
7218 return None;
7219 }
7220 let (se, sg) = m.shared.as_ref()?;
7221 let sgate = gw(sg.as_ref()?)?;
7222 let router = gw(&m.router)?;
7223 let inter = m.experts.first()?.gate_proj.rows();
7224 let mut experts = Vec::with_capacity(m.experts.len() + 1);
7225 let mut q4tp: Option<bool> = None;
7226 let mut gu_q2: Option<bool> = None;
7227 for e in m.experts.iter().chain(std::iter::once(se)) {
7228 if !matches!(e.act, Act::Silu)
7229 || e.gate_proj.rows() != inter
7230 || e.up_proj.rows() != inter
7231 {
7232 return None;
7233 }
7234 let (mm, gi, ui, di, is_p, is_q2) = match e.gate_proj.mapped_q4t() {
7238 Some((mm, gi)) => (
7239 mm,
7240 gi,
7241 e.up_proj.mapped_q4t()?.1,
7242 e.down_proj.mapped_q4t()?.1,
7243 false,
7244 false,
7245 ),
7246 None => match e.gate_proj.mapped_q2tp() {
7247 Some((mm, gi)) => (
7248 mm,
7249 gi,
7250 e.up_proj.mapped_q2tp()?.1,
7251 e.down_proj.mapped_q4tp()?.1,
7252 true,
7253 true,
7254 ),
7255 None => {
7256 let (mm, gi) = e.gate_proj.mapped_q4tp()?;
7257 (
7258 mm,
7259 gi,
7260 e.up_proj.mapped_q4tp()?.1,
7261 e.down_proj.mapped_q4tp()?.1,
7262 true,
7263 false,
7264 )
7265 }
7266 },
7267 };
7268 if *q4tp.get_or_insert(is_p) != is_p
7269 || *gu_q2.get_or_insert(is_q2) != is_q2
7270 {
7271 return None;
7272 }
7273 model.get_or_insert_with(|| mm.clone());
7274 experts.push((gi, ui, di));
7275 }
7276 crate::gpu::GraphFfn::Moe {
7277 router,
7278 shared_gate: sgate,
7279 experts,
7280 n_exp: m.experts.len(),
7281 top_k: m.top_k,
7282 inter,
7283 norm_topk: m.norm_topk_prob,
7284 q4tp: q4tp?,
7285 gu_q2: gu_q2.unwrap_or(false),
7286 }
7287 }
7288 _ => return None,
7289 };
7290 let attn = match &lw.attn {
7291 AttnKind::Full {
7292 wq,
7293 wk,
7294 wv,
7295 wo,
7296 q_norm,
7297 k_norm,
7298 output_gate,
7299 softplus_gate,
7300 bias,
7301 } => {
7302 if softplus_gate.is_some() || self.attention_heads_per_layer.is_some() {
7303 return None;
7304 }
7305 let (m, _, _, _) = wq.graph_weight()?;
7306 model = Some(m.clone());
7307 crate::gpu::GraphAttn::Full {
7308 wq: gw(wq)?,
7309 wk: gw(wk)?,
7310 wv: gw(wv)?,
7311 wo: gw(wo)?,
7312 q_norm: q_norm.as_deref(),
7313 k_norm: k_norm.as_deref(),
7314 bias: bias
7315 .as_ref()
7316 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
7317 output_gate: *output_gate,
7318 cpu_k: self.kv_cache.layers[li].k_heads(),
7319 cpu_v: self.kv_cache.layers[li].v_heads(),
7320 }
7321 }
7322 AttnKind::LinearGdn(w) => {
7323 let cfg = self.gdn_cfg?;
7324 let (m, _, _, _) = w.in_proj_qkv.graph_weight()?;
7325 model = Some(m.clone());
7326 crate::gpu::GraphAttn::Gdn {
7327 qkv: gw(&w.in_proj_qkv)?,
7328 z: gw(&w.in_proj_z)?,
7329 a: gw(&w.in_proj_a)?,
7330 b: gw(&w.in_proj_b)?,
7331 out: gw(&w.out_proj)?,
7332 conv1d: &w.conv1d,
7333 a_log: &w.a_log,
7334 dt_bias: &w.dt_bias,
7335 norm: &w.norm,
7336 nv: cfg.num_v_heads,
7337 nk: cfg.num_k_heads,
7338 dk: cfg.key_head_dim,
7339 dv: cfg.value_head_dim,
7340 kk: cfg.conv_kernel,
7341 cpu_state: &self.kv_cache.layers[self.phys_layer(li)].linear_state,
7342 }
7343 }
7344 _ => return None,
7345 };
7346 layers.push(crate::gpu::GraphLayer {
7347 input_norm: &lw.input_norm,
7348 attn,
7349 post_norm: &lw.post_norm,
7350 ffn: gffn,
7351 });
7352 }
7353 Some((layers, model?))
7354 })();
7355 let Some((layers, model)) = built else {
7356 {
7357 use std::sync::atomic::{AtomicBool, Ordering};
7358 static SAID: AtomicBool = AtomicBool::new(false);
7359 if !SAID.swap(true, Ordering::Relaxed) {
7360 tracing::warn!("batch graph: BUILDER refused (layer weights/kinds)");
7361 }
7362 }
7363 return false;
7364 };
7365 if std::env::var("CMF_GRAPH_SPEC_TIME").is_ok() {
7366 eprintln!("batch-build: {:.1} ms", _tb.elapsed().as_secs_f64() * 1e3);
7367 }
7368 crate::gpu::forward_batch_graph(
7369 &model,
7370 self.graph_kv_id,
7371 &layers,
7372 &self.inv_freq,
7373 hiddens,
7374 nh,
7375 nkv,
7376 hd,
7377 rd,
7378 self.hidden_size,
7379 self.intermediate_size,
7380 positions,
7381 self.kv_cache.max_seq_len,
7382 gemma,
7383 self.rms_eps as f32,
7384 k,
7385 spec,
7386 )
7387 }
7388
7389 fn draft_probe() -> bool {
7393 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
7394 *ON.get_or_init(|| std::env::var("CMF_DSV4_DRAFT_PROBE").is_ok_and(|v| v != "0"))
7395 }
7396
7397 #[cfg(feature = "gpu")]
7409 fn dsv4_spec_on() -> bool {
7410 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
7411 *ON.get_or_init(|| {
7412 std::env::var("CMF_DSV4_SPEC")
7413 .map(|v| v != "0")
7414 .unwrap_or(true)
7415 })
7416 }
7417
7418 #[cfg(feature = "gpu")]
7425 fn dsv4_spec_step(
7426 &mut self,
7427 tip_token: u32,
7428 t_next: u32,
7429 next_pos: usize,
7430 drafted: &mut usize,
7431 accepted_ctr: &mut usize,
7432 ) -> Option<(Vec<u32>, usize)> {
7433 let t_all = std::time::Instant::now();
7434 if std::env::var("CMF_DSV4_SPEC_TIME").is_ok() {
7435 thread_local! {
7436 static LAST: std::cell::Cell<Option<std::time::Instant>> =
7437 const { std::cell::Cell::new(None) };
7438 }
7439 LAST.with(|l| {
7440 if let Some(prev) = l.get() {
7441 eprintln!(
7442 "между раундами {:.1} мс",
7443 prev.elapsed().as_secs_f64() * 1e3
7444 );
7445 }
7446 l.set(Some(std::time::Instant::now()));
7447 });
7448 }
7449 if std::env::var("CMF_DSV4_SPEC_DEBUG").is_ok() {
7450 eprintln!("spec_step: вход pos={next_pos}");
7451 }
7452 let n_layers = self.dsv4.as_ref().map(|b| b.1.len())?;
7453 let cfg = self.dsv4.as_ref().map(|b| b.2)?;
7454 if self.dspark.is_none() {
7456 let t = crate::dsv4::dspark_targets(&self.dsv4_mtp, &cfg, n_layers);
7457 if t.is_empty() {
7458 return None;
7459 }
7460 crate::dsv4::dspark_arm(&t, cfg.dim);
7461 self.dspark = Some(crate::dsv4::DsparkState::new(
7462 self.dsv4_mtp.len(),
7463 &cfg,
7464 t.len(),
7465 ));
7466 }
7467 let targets = crate::dsv4::dspark_targets(&self.dsv4_mtp, &cfg, n_layers);
7468 let pack = crate::dsv4::dspark_pack_get(&self.dsv4_mtp, &cfg);
7469 if pack.is_none() && std::env::var("CMF_DSV4_SPEC_DEBUG").is_ok() {
7470 eprintln!("spec_step: пак не построился (targets {targets:?})");
7471 }
7472 let pack = pack?;
7473 let block = crate::dsv4::dspark_block();
7474 let b_box = self.dsv4.as_mut()?;
7475 let (g, layers, st) = (&b_box.0, &b_box.1, &mut b_box.3);
7476 let ds = self.dspark.as_mut()?;
7477 let dbg = std::env::var("CMF_DSV4_SPEC_DEBUG").is_ok();
7480 if !crate::dsv4::dspark_take(&mut ds.main_hidden) && !ds.have_hidden {
7481 if dbg {
7482 eprintln!("spec_step: нет захвата");
7483 }
7484 return None;
7485 }
7486 ds.have_hidden = true;
7487 let tip_pos = next_pos.checked_sub(1)?;
7488 let draft_started = std::time::Instant::now();
7489 let mut conf = Vec::new();
7490 let props = crate::dsv4::dspark_draft_gpu(
7491 g,
7492 &self.dsv4_mtp,
7493 &cfg,
7494 ds,
7495 pack,
7496 st.kv_id,
7497 tip_token,
7498 tip_pos,
7499 self.pool.as_deref(),
7500 &mut conf,
7501 );
7502 self.dspark_draft_ns += draft_started.elapsed().as_nanos();
7503 *drafted += block;
7504 if props.is_empty() || props[0] != t_next {
7505 if dbg {
7506 eprintln!(
7507 "spec_step: черновик {} (props0={:?} t_next={t_next})",
7508 if props.is_empty() {
7509 "пуст"
7510 } else {
7511 "мимо"
7512 },
7513 props.first()
7514 );
7515 }
7516 return None;
7517 }
7518 let mut k_verify = crate::dsv4::dspark_verify_k().min(props.len());
7519 let conf_min = {
7525 static M: std::sync::OnceLock<f32> = std::sync::OnceLock::new();
7526 *M.get_or_init(|| {
7527 std::env::var("CMF_DSPARK_CONF_MIN")
7528 .ok()
7529 .and_then(|v| v.parse().ok())
7530 .unwrap_or(0.0)
7531 })
7532 };
7533 if conf_min > 0.0 && conf.len() >= props.len() {
7534 let mut keep = 1usize;
7535 while keep < k_verify && conf.get(keep).copied().unwrap_or(0.0) >= conf_min {
7536 keep += 1;
7537 }
7538 k_verify = k_verify.min(keep.max(2));
7539 }
7540 if k_verify < 2 {
7541 return None;
7542 }
7543 let mut fed = Vec::with_capacity(k_verify);
7544 fed.push(t_next);
7545 fed.extend_from_slice(&props[1..k_verify]);
7546 let mut argmax = Vec::new();
7547 let mut logits_all = Vec::new();
7548 let mut walked = Vec::new();
7549 let txn = crate::dsv4::dsv4_verify_chunk(
7550 g,
7551 layers,
7552 &cfg,
7553 st,
7554 &fed,
7555 next_pos,
7556 &self.inv_freq,
7557 self.pool.as_deref(),
7558 &targets,
7559 &mut argmax,
7560 &mut logits_all,
7561 &mut walked,
7562 );
7563 if txn.is_none() && dbg {
7564 eprintln!("spec_step: verify отказал");
7565 }
7566 let txn = txn?;
7567 let b = fed.len();
7568 let mut accepted = 1usize;
7569 while accepted < b && fed[accepted] == argmax[accepted - 1] {
7570 accepted += 1;
7571 }
7572 if std::env::var("CMF_DSV4_SPEC_FORCE_REJECT").is_ok_and(|v| v != "0") {
7577 accepted = 1;
7578 }
7579 if std::env::var("CMF_DSV4_SPEC_TRACE").is_ok() {
7580 eprintln!("spec@{next_pos}: fed={fed:?} argmax={argmax:?} accepted={accepted}");
7581 }
7582 let t_fin = std::time::Instant::now();
7583 if !crate::dsv4::dsv4_spec_finish(
7584 g,
7585 layers,
7586 &cfg,
7587 st,
7588 txn,
7589 accepted,
7590 &fed,
7591 &self.inv_freq,
7592 self.pool.as_deref(),
7593 ) {
7594 tracing::warn!("dsv4: спекулятивный откат не удался — состояние подозрительно");
7595 return None;
7596 }
7597 if std::env::var("CMF_DSV4_SPEC_TIME").is_ok() {
7598 eprintln!(
7599 "finish(k={accepted}): {:.1} мс",
7600 t_fin.elapsed().as_secs_f64() * 1e3
7601 );
7602 }
7603 *accepted_ctr += accepted - 1;
7604 let (hc, dim) = (cfg.hc_mult, cfg.dim);
7609 let dev_caps: Vec<usize> = targets
7616 .iter()
7617 .copied()
7618 .filter(|&t| {
7619 st.dev_set.get(t).copied().unwrap_or(false)
7620 && !st.partial_set.get(t).copied().unwrap_or(false)
7621 })
7622 .collect();
7623 let mut caps_all = vec![0.0f32; dev_caps.len() * b * hc * dim];
7624 if !crate::gpu_wgpu::dsv4_spec_cap_read_all(b, dev_caps.len(), hc * dim, &mut caps_all) {
7625 return None;
7626 }
7627 for t in 0..accepted {
7628 let tip = t + 1 == accepted;
7629 for (slot, &tl) in targets.iter().enumerate() {
7630 if let Some(di) = dev_caps.iter().position(|&d| d == tl) {
7631 let lo = (di * b + t) * hc * dim;
7632 crate::dsv4::dspark_capture(
7633 &caps_all[lo..lo + hc * dim],
7634 &cfg,
7635 slot,
7636 &mut ds.main_hidden,
7637 );
7638 } else if tip
7639 && crate::dsv4::dspark_peek_slot(slot, dim, {
7640 let lo = slot * dim;
7641 &mut ds.main_hidden[lo..lo + dim]
7642 })
7643 {
7644 } else {
7649 crate::dsv4::dspark_capture(
7653 &walked[t * hc * dim..(t + 1) * hc * dim],
7654 &cfg,
7655 slot,
7656 &mut ds.main_hidden,
7657 );
7658 }
7659 }
7660 crate::dsv4::dspark_ring_append(
7661 g,
7662 &self.dsv4_mtp,
7663 &cfg,
7664 ds,
7665 next_pos + t,
7666 self.pool.as_deref(),
7667 );
7668 }
7669 let row = logits_all[(accepted - 1) * cfg.vocab..accepted * cfg.vocab].to_vec();
7670 self.graph_logits = Some(row);
7671 if std::env::var("CMF_DSV4_TRUNK_PICK_DUMP").is_ok() {
7676 crate::dsv4::trunk_freq_note(&crate::dsv4::pick_tally_take());
7677 crate::dsv4::pick_tally_arm();
7678 }
7679 if std::env::var("CMF_DSV4_SPEC_TIME").is_ok() {
7680 eprintln!(
7681 "spec_step total {:.1} мс (k={accepted})",
7682 t_all.elapsed().as_secs_f64() * 1e3
7683 );
7684 }
7685 Some((fed[1..accepted].to_vec(), next_pos + accepted))
7686 }
7687
7688 fn dspark_probe(&mut self, position: usize, token_id: u32) {
7689 if self.dsv4_mtp.is_empty() || !Self::draft_probe() {
7690 return;
7691 }
7692 let trunk_now = crate::dsv4::pick_tally_take();
7694 crate::dsv4::trunk_freq_note(&trunk_now);
7695 if !trunk_now.is_empty() {
7696 self.dspark_trunk_picks.push(trunk_now);
7697 let keep = crate::dsv4::dspark_block();
7698 if self.dspark_trunk_picks.len() > keep {
7699 self.dspark_trunk_picks.remove(0);
7700 }
7701 }
7702 for p in std::mem::take(&mut self.dspark_pending) {
7705 let Some(i) = position.checked_sub(p.0 + 1) else {
7706 continue;
7707 };
7708 let mut p = p;
7709 if i < p.1.len() {
7710 if p.2 && p.1[i] == token_id {
7711 p.3 = i + 1;
7712 } else {
7713 p.2 = false;
7714 }
7715 if i + 1 < p.1.len() {
7716 self.dspark_pending.push(p);
7717 continue;
7718 }
7719 }
7720 self.dspark_hist.push(p.3);
7721 self.dspark_real.push(token_id);
7722 }
7723 let Some(b) = &mut self.dsv4 else { return };
7724 let (g, layers, cfg) = (&b.0, &b.1, b.2);
7725 let n_layers = layers.len();
7726 if self.dspark.is_none() {
7727 let t = crate::dsv4::dspark_targets(&self.dsv4_mtp, &cfg, n_layers);
7728 if t.is_empty() {
7729 return;
7730 }
7731 eprintln!(
7732 "DSpark: захват со слоёв {t:?}, блок {}",
7733 crate::dsv4::dspark_block()
7734 );
7735 crate::dsv4::dspark_arm(&t, cfg.dim);
7736 self.dspark = Some(crate::dsv4::DsparkState::new(
7737 self.dsv4_mtp.len(),
7738 &cfg,
7739 t.len(),
7740 ));
7741 }
7742 let ds = self.dspark.as_mut().unwrap();
7743 if !crate::dsv4::dspark_take(&mut ds.main_hidden) {
7744 return; }
7746 let mut conf = Vec::new();
7747 crate::dsv4::pick_tally_arm();
7748 let draft_started = std::time::Instant::now();
7753 #[cfg(feature = "gpu")]
7754 let gpu_draft = crate::dsv4::dspark_gpu_on();
7755 #[cfg(not(feature = "gpu"))]
7756 let gpu_draft = false;
7757 let props = if gpu_draft {
7758 #[cfg(feature = "gpu")]
7759 {
7760 let kv_id = b.3.kv_id;
7761 match crate::dsv4::dspark_pack_get(&self.dsv4_mtp, &cfg) {
7762 Some(pk) => crate::dsv4::dspark_draft_gpu(
7763 g,
7764 &self.dsv4_mtp,
7765 &cfg,
7766 ds,
7767 pk,
7768 kv_id,
7769 token_id,
7770 position,
7771 self.pool.as_deref(),
7772 &mut conf,
7773 ),
7774 None => Vec::new(),
7775 }
7776 }
7777 #[cfg(not(feature = "gpu"))]
7778 Vec::new()
7779 } else {
7780 crate::gpu::cpu_scope(|| {
7781 crate::dsv4::dspark_draft(
7782 g,
7783 &self.dsv4_mtp,
7784 &cfg,
7785 ds,
7786 token_id,
7787 position,
7788 self.pool.as_deref(),
7789 &mut conf,
7790 )
7791 })
7792 };
7793 self.dspark_draft_ns += draft_started.elapsed().as_nanos();
7794 let draft_picks = crate::dsv4::pick_tally_take();
7795 crate::dsv4::dspark_freq_note(&draft_picks);
7796 crate::dsv4::pick_tally_arm();
7799 if !props.is_empty() {
7800 let (tu, tt) = {
7804 let flat: Vec<(usize, Vec<usize>)> = self
7805 .dspark_trunk_picks
7806 .iter()
7807 .flat_map(|v| v.iter().cloned())
7808 .collect();
7809 let mut per: std::collections::HashMap<usize, Vec<usize>> =
7811 std::collections::HashMap::new();
7812 for (li, picks) in flat {
7813 per.entry(li).or_default().extend(picks);
7814 }
7815 let n = per.len().max(1);
7816 let mut u = 0usize;
7817 let mut t = 0usize;
7818 for (_, v) in per {
7819 t += v.len();
7820 u += v.iter().collect::<std::collections::HashSet<_>>().len();
7821 }
7822 (u / n, t / n)
7823 };
7824 let (du, dt) = crate::dsv4::tally_unique(&draft_picks);
7825 self.dspark_exp.push((tu, tt, du, dt));
7826 self.dspark_pending.push((position, props, true, 0));
7827 }
7828 if self.dspark_hist.len() >= 8 && self.dspark_hist.len() % 8 == 0 {
7829 let n = self.dspark_hist.len() as f32;
7830 let mean: f32 = self.dspark_hist.iter().sum::<usize>() as f32 / n;
7831 let block = crate::dsv4::dspark_block();
7832 let mut at = vec![0usize; block + 1];
7833 for &k in &self.dspark_hist {
7834 at[k] += 1;
7835 }
7836 let mut surv = Vec::with_capacity(block);
7838 for i in 1..=block {
7839 let k = at[i..].iter().sum::<usize>() as f32 / n;
7840 surv.push(format!("{k:.2}"));
7841 }
7842 let distinct = self
7843 .dspark_real
7844 .iter()
7845 .collect::<std::collections::HashSet<_>>()
7846 .len();
7847 let (tu, tt, du, dt) = self.dspark_exp.iter().fold((0, 0, 0, 0), |a, b| {
7848 (a.0 + b.0, a.1 + b.1, a.2 + b.2, a.3 + b.3)
7849 });
7850 let m = self.dspark_exp.len().max(1);
7851 eprintln!(
7852 "DSpark: черновиков {}, принято в среднем {mean:.2} из {block} \
7853 (токенов за проход {:.2}), распределение {at:?}, выживание [{}]",
7854 self.dspark_hist.len(),
7855 mean + 1.0,
7856 surv.join(" ")
7857 );
7858 eprintln!(
7859 "DSpark: разных токенов {distinct} из {} (вырожденность), \
7860 эксперты ствол {}/{} на слой за {block} токенов, \
7861 черновик {}/{} за блок, draft {:.2} мс/блок",
7862 self.dspark_real.len(),
7863 tu / m,
7864 tt / m,
7865 du / m,
7866 dt / m,
7867 self.dspark_draft_ns as f64 / self.dspark_exp.len().max(1) as f64 / 1e6
7868 );
7869 }
7870 }
7871
7872 fn forward_layers_upto(
7873 &mut self,
7874 hidden: &[f32],
7875 position: usize,
7876 task_mask: Option<&TaskMask>,
7877 upto: Option<usize>,
7878 ) -> Vec<f32> {
7879 if let Some(plan) = self.gpu_plan.clone() {
7885 if upto.is_none() && plan.len() > 1 {
7886 let mut h = hidden.to_vec();
7887 for &(dev, from, upto_incl) in plan.iter() {
7888 h = crate::gpu::with_device(dev, || {
7889 self.forward_layers_span(&h, position, task_mask, from, Some(upto_incl))
7890 });
7891 }
7892 return h;
7893 }
7894 }
7895 self.forward_layers_span(hidden, position, task_mask, 0, upto)
7896 }
7897
7898 pub fn set_gpu_plan(&mut self, devices: Option<&[usize]>) -> Result<(), String> {
7903 self.set_gpu_plan_at(devices, None)
7904 }
7905
7906 pub fn set_gpu_plan_at(
7910 &mut self,
7911 devices: Option<&[usize]>,
7912 at: Option<usize>,
7913 ) -> Result<(), String> {
7914 let Some(devs) = devices.filter(|d| d.len() > 1) else {
7915 self.gpu_plan = None;
7916 return Ok(());
7917 };
7918 self.split_supported()?;
7919 let n = self.num_layers;
7920 if devs.len() > n {
7921 return Err(format!("{} devices for {n} layers", devs.len()));
7922 }
7923 if let Some(k) = at {
7924 if k == 0 || k >= n {
7925 return Err(format!("split at {k}: the model has {n} layers"));
7926 }
7927 if devs.len() == 2 {
7928 self.gpu_plan = Some(std::sync::Arc::new(vec![
7929 (devs[0], 0, k - 1),
7930 (devs[1], k, n - 1),
7931 ]));
7932 return Ok(());
7933 }
7934 return Err(format!(
7935 "an explicit split point takes exactly 2 devices, got {}",
7936 devs.len()
7937 ));
7938 }
7939 let per = n.div_ceil(devs.len());
7940 let mut plan = Vec::with_capacity(devs.len());
7941 let mut from = 0usize;
7942 for &d in devs {
7943 if from >= n {
7944 break;
7945 }
7946 let upto = (from + per - 1).min(n - 1);
7947 plan.push((d, from, upto));
7948 from = upto + 1;
7949 }
7950 self.gpu_plan = Some(std::sync::Arc::new(plan));
7951 Ok(())
7952 }
7953
7954 pub fn gpu_plan(&self) -> Option<Vec<(usize, usize, usize)>> {
7956 self.gpu_plan.as_ref().map(|p| p.as_ref().clone())
7957 }
7958
7959 fn forward_layers_span(
7965 &mut self,
7966 hidden: &[f32],
7967 position: usize,
7968 task_mask: Option<&TaskMask>,
7969 from: usize,
7970 upto: Option<usize>,
7971 ) -> Vec<f32> {
7972 debug_assert!(from == 0 || (self.dsv4.is_none() && self.g3n.is_none()));
7973 if let Some(b) = &mut self.dsv4 {
7979 let _ = (task_mask, upto);
7980 let token_id = hidden.first().copied().unwrap_or(0.0) as u32;
7981 let (g, layers, cfg, st) = (&b.0, &b.1, b.2, &mut b.3);
7982 st.pos = position;
7983 let mut logits = Vec::new();
7984 crate::dsv4::forward_token(
7985 g,
7986 layers,
7987 &cfg,
7988 st,
7989 token_id,
7990 &self.inv_freq,
7991 self.pool.as_deref(),
7992 &mut logits,
7993 );
7994 self.graph_logits = Some(logits);
7995 self.dspark_probe(position, token_id);
7996 return vec![0.0; self.hidden_size];
7999 }
8000 if let Some(b) = &self.g3n {
8003 let _ = (task_mask, upto);
8004 return crate::g3n::g3n_forward(
8005 &b.0,
8006 &b.1,
8007 hidden,
8008 position,
8009 &mut self.kv_cache.layers,
8010 self.num_heads,
8011 self.num_kv_heads,
8012 self.head_dim,
8013 self.pool.as_deref(),
8014 );
8015 }
8016 let mut h = hidden.to_vec();
8017 let (nh, _nkv, _hd, hs, _rd, eps) = (
8020 self.num_heads,
8021 self.num_kv_heads,
8022 self.head_dim,
8023 self.hidden_size,
8024 self.rotary_dim,
8025 self.rms_eps,
8026 );
8027 let pool = self.pool.clone();
8028 let graph_env = std::env::var("CMF_GPU_WGPU_GRAPH").ok();
8040 let graph_on = match graph_env.as_deref() {
8041 Some("0") => false,
8042 Some("prefill") => false, Some(_) => true,
8044 None => crate::gpu::wgpu_graph_default(),
8050 };
8051 let graph_trusted =
8052 graph_env.is_some() || crate::gpu::wgpu_graph_default() || self.gdn_cfg.is_some();
8053 let race_eligible = graph_on
8054 && upto.is_none()
8055 && task_mask.is_none()
8056 && from == 0
8057 && !crate::gpu::graph_unsupported();
8058 let mut tail_start = 0usize;
8059 if race_eligible && crate::gpu::graph_race_use_graph(graph_trusted) {
8060 let t_graph = std::time::Instant::now();
8061 let mut lg = Vec::new();
8062 let mut gl = 0usize;
8063 let built = self.try_token_graph_wgpu(hidden, position, &mut lg, &mut gl);
8064 if built.is_none() && !self.o1_active() && self.attn_softcap == 0.0 {
8069 crate::gpu::graph_mark_unsupported();
8070 }
8071 graph_note(built.is_some());
8072 if let Some(hh) = built {
8073 let dur = t_graph.elapsed();
8074 if std::env::var("CMF_GRAPH_PROF").is_ok() {
8075 eprintln!("graph-call: {:.2} ms total", dur.as_secs_f64() * 1000.0);
8076 }
8077 if gl > 0 && gl < self.num_layers {
8078 h = hh;
8084 tail_start = gl;
8085 } else if graph_trusted || !crate::gpu::graph_race_first_token_hopeless(dur) {
8086 if !graph_trusted {
8087 crate::gpu::graph_race_record(true, dur);
8088 }
8089 if !lg.is_empty() {
8090 lg.resize(self.vocab_size, 0.0);
8093 if let Some(c) = self.final_softcap {
8094 for l in lg.iter_mut() {
8095 *l = c * (*l / c).tanh();
8096 }
8097 }
8098 self.graph_logits = Some(lg);
8099 }
8100 return hh;
8101 }
8102 }
8108 }
8109 let span = from > 0 || upto.is_some();
8133 if span && graph_on && task_mask.is_none() && graph_trusted {
8134 let upto_excl = upto.map_or(self.num_layers, |u| u + 1);
8135 let mut lg = Vec::new();
8136 let mut gl = 0usize;
8137 let span_res =
8138 self.try_token_graph_wgpu_span(hidden, position, &mut lg, from, upto_excl, &mut gl);
8139 graph_note(span_res.is_some() && gl == upto_excl - from);
8140 if std::env::var("CMF_GPU_DEBUG").is_ok() {
8141 static SEEN: std::sync::atomic::AtomicU32 = std::sync::atomic::AtomicU32::new(0);
8145 if SEEN.fetch_add(1, std::sync::atomic::Ordering::Relaxed) < 4 {
8146 eprintln!(
8147 "span graph: covered {gl} of {} layers [{from}..{upto_excl}) res={}",
8148 upto_excl - from,
8149 span_res.is_some()
8150 );
8151 }
8152 }
8153 if let Some(hh) = span_res {
8154 if gl == upto_excl - from {
8155 if !lg.is_empty() {
8156 lg.resize(self.vocab_size, 0.0);
8157 if let Some(c) = self.final_softcap {
8158 for l in lg.iter_mut() {
8159 *l = c * (*l / c).tanh();
8160 }
8161 }
8162 self.graph_logits = Some(lg);
8163 }
8164 crate::gpu::set_layer(-1);
8165 return hh;
8166 }
8167 h = hh;
8169 tail_start = from + gl;
8170 }
8171 }
8172 let t_race_cpu = (race_eligible && !graph_trusted).then(std::time::Instant::now);
8173
8174 #[cfg(target_os = "macos")]
8175 let mut gpu_skip_until = 0usize;
8176 for li in tail_start.max(from)..self.num_layers {
8177 crate::gpu::set_layer(li as i64); if let Some(u) = upto {
8179 if li > u {
8180 break;
8181 }
8182 }
8183 if let Some(mask) = task_mask {
8184 if !mask.layer_alive(li) {
8185 continue; }
8187 }
8188 #[cfg(target_os = "macos")]
8192 {
8193 if li < gpu_skip_until {
8194 continue;
8195 }
8196 if task_mask.is_none() {
8197 let end = self.q1_graph_gpu(li, upto, position, &mut h);
8198 if end > li {
8199 gpu_skip_until = end;
8200 if self.is_loop_end(end - 1) && end < self.num_layers {
8203 h = inference::rms_norm(
8204 &h,
8205 &self.weights.final_norm,
8206 self.rms_eps,
8207 self.norm_style,
8208 );
8209 }
8210 continue;
8211 }
8212 }
8213 }
8214
8215 let lw = &self.weights.layers[self.phys_layer(li)];
8216 if let Ok(tp) = std::env::var("CMF_TRACE_POS") {
8217 if tp.parse::<usize>().ok() == Some(position) {
8218 let n: f32 = h.iter().map(|x| x * x).sum::<f32>().sqrt();
8219 eprintln!(
8220 "TRACE pos {position} layer {li}: |h| = {n:.6} h0 {:.6} h1 {:.6}",
8221 h[0], h[1]
8222 );
8223 }
8224 }
8225 inference::rms_norm_into(
8228 &h,
8229 &lw.input_norm,
8230 self.rms_eps,
8231 self.norm_style,
8232 &mut self.ws.n1,
8233 );
8234
8235 let attn_out = match &lw.attn {
8236 AttnKind::Mla(w) => {
8237 let inv_freq_l = self.layer_inv_freq(li);
8238 let rs = self.layer_rope_scale(li);
8239 let eps = self.rms_eps;
8240 let pool = self.pool.clone();
8241 mla_attention(
8242 w,
8243 &self.ws.n1,
8244 &mut self.kv_cache.layers[li],
8245 position,
8246 &inv_freq_l,
8247 rs,
8248 eps,
8249 pool.as_deref(),
8250 )
8251 }
8252 AttnKind::Linear(w) => {
8253 let cfg = self.vmf_cfg.expect("linear layer without vmf_cfg");
8254 vmf_phase_forward(
8255 &self.ws.n1,
8256 w,
8257 &cfg,
8258 &mut self.kv_cache.layers[li].linear_state,
8259 self.pool.as_deref(),
8260 )
8261 }
8262 AttnKind::Kda(w) => {
8263 let cfg = self.kda_cfg.expect("kda layer without kda_cfg");
8264 crate::linear_core::kda_forward(
8265 &self.ws.n1,
8266 w,
8267 &cfg,
8268 &mut self.kv_cache.layers[li].linear_state,
8269 self.pool.as_deref(),
8270 )
8271 }
8272 AttnKind::LinearGdn(w) => {
8273 let cfg = self.gdn_cfg.expect("gdn layer without gdn_cfg");
8274 gdn_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::ShortConv(w) => {
8283 let cfg = self
8284 .short_conv_cfg
8285 .expect("short-conv layer without short_conv_cfg");
8286 short_conv_forward(
8287 &self.ws.n1,
8288 w,
8289 &cfg,
8290 &mut self.kv_cache.layers[li].linear_state,
8291 self.pool.as_deref(),
8292 )
8293 }
8294 AttnKind::Full {
8295 wq,
8296 wk,
8297 wv,
8298 wo,
8299 q_norm,
8300 k_norm,
8301 output_gate,
8302 softplus_gate,
8303 bias,
8304 } if self.kv_cache.layers[li].o1_sealed() => {
8305 let inv_freq_l = self.layer_inv_freq(li);
8308 let (nkv_l, hd_l, rd_l) = self.layer_geom(li);
8309 let cfg = QwenAttnCfg {
8310 num_heads: self.layer_num_heads(li),
8311 num_kv_heads: nkv_l,
8312 head_dim: hd_l,
8313 hidden_size: hs,
8314 position,
8315 inv_freq: &inv_freq_l,
8316 rotary_dim: rd_l,
8317 scale: self.attn_scale,
8318 softcap: self.attn_softcap,
8319 window: None,
8320 v_norm: self.attn_v_norm,
8321 q_norm: q_norm.as_deref(),
8322 k_norm: k_norm.as_deref(),
8323 output_gate: *output_gate,
8324 softplus_gate: softplus_gate
8325 .as_ref()
8326 .map(|(gate, per_head)| (gate, *per_head)),
8327 rope_scale: self.layer_rope_scale(li),
8328 bias: bias
8329 .as_ref()
8330 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
8331 rms_eps: eps,
8332 norm_style: self.norm_style,
8333 pool: pool.as_deref(),
8334 };
8335 attention::qwen_attention_nystrom(
8336 &self.ws.n1,
8337 wq,
8338 wk,
8339 wv,
8340 wo,
8341 &mut self.kv_cache.layers[li],
8342 &cfg,
8343 )
8344 }
8345 AttnKind::Full {
8346 wq,
8347 wk,
8348 wv,
8349 wo,
8350 q_norm,
8351 k_norm,
8352 output_gate,
8353 softplus_gate,
8354 bias,
8355 } => 'attn: {
8356 if graph_on
8359 && !*output_gate
8360 && softplus_gate.is_none()
8361 && self.attention_heads_per_layer.is_none()
8362 && bias.is_none()
8363 && task_mask.is_none()
8364 {
8365 let inv_freq_l = self.layer_inv_freq(li);
8366 let (nkv_l, hd_l, rd_l) = self.layer_geom(li);
8367 let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
8368 if let (Some((gm, qi)), Some((_, ki)), Some((_, vi)), Some((_, oi))) = (
8369 wq.mapped_q1(),
8370 wk.mapped_q1(),
8371 wv.mapped_q1(),
8372 wo.mapped_q1(),
8373 ) {
8374 let gm = gm.clone();
8375 let mut out = vec![0f32; hs];
8376 let cache = &self.kv_cache.layers[li];
8377 if crate::gpu::attn_dropin(
8378 &gm,
8379 self.graph_kv_id,
8380 li,
8381 &self.ws.n1,
8382 qi,
8383 ki,
8384 vi,
8385 oi,
8386 q_norm.as_deref(),
8387 k_norm.as_deref(),
8388 &inv_freq_l,
8389 nh,
8390 nkv_l,
8391 hd_l,
8392 rd_l,
8393 hs,
8394 position,
8395 self.kv_cache.max_seq_len,
8396 gemma,
8397 eps as f32,
8398 cache.k_heads(),
8399 cache.v_heads(),
8400 &mut out,
8401 ) {
8402 break 'attn out;
8403 }
8404 }
8405 }
8406 let masked = task_mask
8407 .map(|m| m.head_flags(li, self.num_heads).iter().any(|&a| !a))
8408 .unwrap_or(false);
8409 let f32_view = (wq.as_f32(), wk.as_f32(), wv.as_f32(), wo.as_f32());
8410 match (masked, f32_view) {
8411 (true, (Some(q), Some(k), Some(v), Some(o))) => {
8414 let active_heads = task_mask.unwrap().head_flags(li, self.num_heads);
8415 attention::multi_head_attention(
8416 &self.ws.n1,
8417 q,
8418 k,
8419 v,
8420 o,
8421 &mut self.kv_cache.layers[li],
8422 self.num_heads,
8423 self.num_kv_heads,
8424 self.head_dim,
8425 self.hidden_size,
8426 position,
8427 &active_heads,
8428 &self.inv_freq,
8429 )
8430 }
8431 (masked, _) => {
8432 if masked {
8433 tracing::warn!(
8434 "layer {li}: head mask on quantized weights not \
8435 supported yet — executing dense"
8436 );
8437 }
8438 let inv_freq_l = self.layer_inv_freq(li);
8439 let (nkv_l, hd_l, rd_l) = self.layer_geom(li);
8440 let cfg = QwenAttnCfg {
8441 num_heads: self.layer_num_heads(li),
8442 num_kv_heads: nkv_l,
8443 head_dim: hd_l,
8444 hidden_size: hs,
8445 position,
8446 inv_freq: &inv_freq_l,
8447 rotary_dim: rd_l,
8448 scale: self.attn_scale,
8449 softcap: self.attn_softcap,
8450 window: self.layer_window(li),
8451 v_norm: self.attn_v_norm,
8452 q_norm: q_norm.as_deref(),
8453 k_norm: k_norm.as_deref(),
8454 output_gate: *output_gate,
8455 softplus_gate: softplus_gate
8456 .as_ref()
8457 .map(|(gate, per_head)| (gate, *per_head)),
8458 rope_scale: self.layer_rope_scale(li),
8459 bias: bias
8460 .as_ref()
8461 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
8462 rms_eps: eps,
8463 norm_style: self.norm_style,
8464 pool: pool.as_deref(),
8465 };
8466 attention::qwen_attention(
8467 &self.ws.n1,
8468 wq,
8469 wk,
8470 wv,
8471 wo,
8472 &mut self.kv_cache.layers[li],
8473 &cfg,
8474 )
8475 }
8476 }
8477 }
8478 };
8479 let attn_out = match &self.weights.layers[self.phys_layer(li)].attn_out_norm {
8482 Some(w) => inference::rms_norm(&attn_out, w, self.rms_eps, self.norm_style),
8483 None => attn_out,
8484 };
8485 let lw = &self.weights.layers[self.phys_layer(li)];
8486 inference::add_rmsnorm_fused_into(
8487 &mut h,
8488 &attn_out,
8489 &lw.post_norm,
8490 self.rms_eps,
8491 self.norm_style,
8492 &mut self.ws.p1,
8493 );
8494 let mut attn_out = attn_out;
8495 attention::recycle_buf(&mut attn_out);
8496 let post_normed = &self.ws.p1;
8497
8498 let ffn_masked = task_mask
8499 .map(|m| m.ffn_active_count(li) < self.intermediate_size)
8500 .unwrap_or(false);
8501 let ffn_out = match (ffn_masked, &lw.ffn) {
8513 (_, FfnKind::Dense(d)) if !d.segs.is_empty() => {
8517 let row = task_mask
8518 .and_then(|tm| tm.ffn_masks.get(li))
8519 .map(|v| v.as_slice());
8520 tube_ffn(d, post_normed, 1, self.pool.as_deref(), row)
8521 }
8522 (true, FfnKind::Dense(d)) => {
8523 let tm = task_mask.unwrap();
8524 let alive = tm.ffn_active_count(li);
8525 let deep = alive * 2 <= self.intermediate_size;
8526 if deep && d.down_proj.sparse_col_ok() {
8527 let active = tm.ffn_active_indices(li);
8528 sparse_ffn_quant(
8529 d,
8530 post_normed,
8531 &active,
8532 self.hidden_size,
8533 self.pool.as_deref(),
8534 )
8535 } else if deep
8536 && let (Some(g), Some(u), Some(dn)) = (
8537 d.gate_proj.as_f32(),
8538 d.up_proj.as_f32(),
8539 d.down_proj.as_f32(),
8540 )
8541 {
8542 let active = tm.ffn_active_indices(li);
8543 inference::sparse_ffn_forward(
8544 post_normed,
8545 g,
8546 u,
8547 dn,
8548 self.hidden_size,
8549 self.intermediate_size,
8550 &active,
8551 self.pool.as_deref(),
8552 )
8553 } else {
8554 let row = tm.ffn_masks.get(li).map(|v| v.as_slice());
8555 dense_ffn_batch(d, post_normed, 1, self.pool.as_deref(), row)
8556 }
8557 }
8558 (true, FfnKind::Moe(m)) => {
8559 let allowed = task_mask.and_then(|tm| tm.expert_flags(li, m.experts.len()));
8563 ffn_forward(
8564 &lw.ffn,
8565 post_normed,
8566 self.pool.as_deref(),
8567 allowed.as_deref(),
8568 )
8569 }
8570 (true, FfnKind::DenseMoe(dm)) => dense_moe_ffn(
8571 dm,
8572 post_normed,
8573 &h,
8574 self.rms_eps,
8575 self.norm_style,
8576 self.pool.as_deref(),
8577 ),
8578 (false, _) => match &lw.ffn {
8579 FfnKind::DenseMoe(dm) => dense_moe_ffn(
8580 dm,
8581 post_normed,
8582 &h,
8583 self.rms_eps,
8584 self.norm_style,
8585 self.pool.as_deref(),
8586 ),
8587 _ => {
8588 let allowed = match (&lw.ffn, task_mask) {
8589 (FfnKind::Moe(m), Some(tm)) => tm.expert_flags(li, m.experts.len()),
8590 _ => None,
8591 };
8592 ffn_forward(
8593 &lw.ffn,
8594 post_normed,
8595 self.pool.as_deref(),
8596 allowed.as_deref(),
8597 )
8598 }
8599 },
8600 };
8601 let ffn_out = match &self.weights.layers[self.phys_layer(li)].ffn_out_norm {
8602 Some(w) => inference::rms_norm(&ffn_out, w, self.rms_eps, self.norm_style),
8603 None => ffn_out,
8604 };
8605 for (i, &f) in ffn_out.iter().enumerate() {
8606 h[i] += f;
8607 }
8608 let mut ffn_out = ffn_out;
8609 attention::recycle_buf(&mut ffn_out);
8610
8611 if let Some(sc) = self.weights.layers[self.phys_layer(li)].layer_scale {
8613 for v in h.iter_mut() {
8614 *v *= sc;
8615 }
8616 }
8617
8618 if self.is_loop_end(li) && li + 1 < self.num_layers {
8621 h = inference::rms_norm(
8622 &h,
8623 &self.weights.final_norm,
8624 self.rms_eps,
8625 self.norm_style,
8626 );
8627 }
8628
8629 if self.dyn_phi_layer == Some(li) {
8633 self.update_dyn_phi(&h);
8634 }
8635 }
8636 crate::gpu::set_layer(-1); if let Some(t) = t_race_cpu {
8638 crate::gpu::graph_race_record(false, t.elapsed());
8639 }
8640
8641 h
8642 }
8643
8644 fn update_dyn_phi(&mut self, h: &[f32]) {
8647 const A: f32 = 0.2;
8648 if self.dyn_phi_ema.len() != h.len() {
8649 self.dyn_phi_ema = vec![0.0; h.len()];
8650 self.dyn_phi_seen = 0;
8651 }
8652 if self.dyn_phi_seen == 0 {
8653 self.dyn_phi_ema.copy_from_slice(h);
8654 } else {
8655 for (e, &v) in self.dyn_phi_ema.iter_mut().zip(h) {
8656 *e = (1.0 - A) * *e + A * v;
8657 }
8658 }
8659 self.dyn_phi_seen += 1;
8660 }
8661
8662 pub fn dyn_phi(&self) -> &[f32] {
8664 &self.dyn_phi_ema
8665 }
8666
8667 pub fn set_dyn_phi_layer(&mut self, layer: Option<usize>) {
8669 self.dyn_phi_layer = layer;
8670 self.dyn_phi_ema.clear();
8671 self.dyn_phi_seen = 0;
8672 }
8673
8674 pub fn dynamic_skills(&self) -> Vec<(usize, String, usize)> {
8676 let Some(model) = &self.model else {
8677 return Vec::new();
8678 };
8679 model
8680 .header
8681 .skills
8682 .iter()
8683 .enumerate()
8684 .filter_map(|(i, sk)| {
8685 let ok = matches!(self.dyn_skill_layers.get(i), Some(Some(_)));
8686 let sel = sk.selection.as_ref()?;
8687 (ok).then(|| (i, sk.id.clone(), sel.phi_layer))
8688 })
8689 .collect()
8690 }
8691
8692 pub fn active_skill(&self) -> Option<usize> {
8694 self.dyn_active
8695 }
8696
8697 pub fn enable_dynamic_routing(&mut self) -> usize {
8702 use crate::swarm::{DynRouter, RoutableSkill};
8703 let Some(model) = self.model.clone() else {
8704 return 0;
8705 };
8706 if self.dyn_blend_loaded {
8709 tracing::warn!("dynamic routing unavailable on a blend-loaded pipeline");
8710 return 0;
8711 }
8712 if let Some(a) = self.dyn_active {
8716 if !matches!(self.dyn_skill_layers.get(a), Some(Some(_))) {
8717 tracing::warn!("loaded skill is not FFN-eligible — dynamic routing unavailable");
8718 return 0;
8719 }
8720 }
8721 let hidden = self.hidden_size;
8722 let mut skills = Vec::new();
8723 for (idx, id, _phi) in self.dynamic_skills() {
8724 if let Some(sel) = model.header.skills[idx].selection.as_ref() {
8725 if let Some(rs) = RoutableSkill::from_descriptor(idx, id, sel, hidden) {
8726 skills.push(rs);
8727 }
8728 }
8729 }
8730 if skills.is_empty() {
8731 return 0;
8732 }
8733 let phi = skills[0].phi_layer;
8735 if skills.iter().any(|s| s.phi_layer != phi) {
8736 tracing::warn!("routable skills disagree on phi_layer; using {phi}");
8737 }
8738 let n = skills.len();
8739 self.set_dyn_phi_layer(Some(phi));
8740 self.dyn_router = Some(DynRouter::new(skills));
8741 n
8742 }
8743
8744 pub fn route_switches(&self) -> Vec<(usize, Option<String>, Option<String>)> {
8746 self.dyn_router
8747 .as_ref()
8748 .map(|r| r.switches.clone())
8749 .unwrap_or_default()
8750 }
8751
8752 fn lm_head_forward(&self, hidden: &[f32]) -> Vec<f32> {
8755 let rows = self.weights.lm_head.rows();
8756 let mut logits = attention::take_buf(rows.min(self.vocab_size));
8757 self.weights
8758 .lm_head
8759 .matvec(hidden, &mut logits, self.pool.as_deref());
8760 logits.resize(self.vocab_size, 0.0);
8761 if let Some(m) = self.logit_multiplier {
8762 for l in logits.iter_mut() {
8763 *l *= m;
8764 }
8765 }
8766 if let Some(c) = self.final_softcap {
8767 for l in logits.iter_mut() {
8768 *l = c * (*l / c).tanh();
8769 }
8770 }
8771 if let Some(cm) = self.head_clusters.as_ref() {
8772 self.hierarchical_head_logprobs(hidden, cm, &mut logits);
8773 }
8774 logits
8775 }
8776
8777 fn hierarchical_head_logprobs(&self, hidden: &[f32], cm: &[f32], logits: &mut [f32]) {
8780 let h = hidden.len();
8781 let ncl = cm.len() / h.max(1);
8782 if ncl == 0 || logits.len() % ncl != 0 {
8783 return;
8784 }
8785 let cs = logits.len() / ncl;
8786 let mut lc = vec![0.0f32; ncl];
8788 for c in 0..ncl {
8789 let row = &cm[c * h..(c + 1) * h];
8790 let mut s = 0.0f32;
8791 for j in 0..h {
8792 s += row[j] * hidden[j];
8793 }
8794 lc[c] = s;
8795 }
8796 let mx = lc.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
8797 let lse: f32 = mx + lc.iter().map(|v| (v - mx).exp()).sum::<f32>().ln();
8798 for c in 0..ncl {
8799 let blk = &mut logits[c * cs..(c + 1) * cs];
8800 let bm = blk.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
8801 let bl: f32 = bm + blk.iter().map(|v| (v - bm).exp()).sum::<f32>().ln();
8802 let add = lc[c] - lse - bl;
8803 for v in blk.iter_mut() {
8804 *v += add;
8805 }
8806 }
8807 }
8808
8809 pub fn prefill_next_logits(&mut self, ids: &[u32], task_mask: Option<&TaskMask>) -> Vec<f32> {
8814 self.kv_cache.clear();
8815 self.kv_history.clear();
8816 let mut hidden = vec![0.0f32; self.hidden_size];
8817 for (pos, &id) in ids.iter().enumerate() {
8818 let emb = self.embed_single(id);
8819 hidden = self.forward_layers(&emb, pos, task_mask);
8820 }
8821 inference::rms_norm_into(
8822 &hidden,
8823 &self.weights.final_norm,
8824 self.rms_eps,
8825 self.norm_style,
8826 &mut self.ws.n1,
8827 );
8828 self.lm_head_forward(&self.ws.n1)
8829 }
8830}
8831
8832pub fn create_test_pipeline(
8834 hidden_size: usize,
8835 intermediate_size: usize,
8836 num_heads: usize,
8837 num_kv_heads: usize,
8838 head_dim: usize,
8839 num_layers: usize,
8840 vocab_size: usize,
8841) -> Pipeline {
8842 let synth = |n: usize, salt: usize| -> Vec<f32> {
8845 (0..n)
8846 .map(|i| (((i * 31 + salt * 17 + 7) % 97) as f32 / 97.0 - 0.5) * 0.2)
8847 .collect()
8848 };
8849 let qt = |rows: usize, cols: usize, salt: usize| -> QTensor {
8850 QTensor::from_f32(synth(rows * cols, salt), rows, cols)
8851 };
8852 let layer_weights: Vec<LayerWeights> = (0..num_layers)
8853 .map(|li| LayerWeights {
8854 input_norm: vec![1.0; hidden_size],
8855 post_norm: vec![1.0; hidden_size],
8856 attn_out_norm: None,
8857 ffn_out_norm: None,
8858 layer_scale: None,
8859 ffn: FfnKind::Dense(DenseFfn {
8860 gate_proj: qt(intermediate_size, hidden_size, li * 10 + 5),
8861 up_proj: qt(intermediate_size, hidden_size, li * 10 + 6),
8862 down_proj: qt(hidden_size, intermediate_size, li * 10 + 7),
8863 act: Act::Silu,
8864 down_t: None,
8865 segs: Vec::new(),
8866 }),
8867 attn: AttnKind::Full {
8868 bias: None,
8869 wq: qt(num_heads * head_dim, hidden_size, li * 10 + 1),
8870 wk: qt(num_kv_heads * head_dim, hidden_size, li * 10 + 2),
8871 wv: qt(num_kv_heads * head_dim, hidden_size, li * 10 + 3),
8872 wo: qt(hidden_size, num_heads * head_dim, li * 10 + 4),
8873 q_norm: None,
8874 k_norm: None,
8875 output_gate: false,
8876 softplus_gate: None,
8877 },
8878 })
8879 .collect();
8880
8881 Pipeline::new(
8882 Tokenizer::byte_level(),
8883 PipelineWeights {
8884 embed_tokens: qt(vocab_size, hidden_size, 100),
8885 layers: layer_weights,
8886 lm_head: qt(vocab_size, hidden_size, 200),
8887 final_norm: vec![1.0; hidden_size],
8888 },
8889 hidden_size,
8890 intermediate_size,
8891 num_heads,
8892 num_kv_heads,
8893 head_dim,
8894 num_layers,
8895 num_layers, false, vocab_size,
8898 1e-6,
8899 10_000.0,
8900 NormStyle::Qwen,
8901 4096,
8902 SamplerConfig {
8903 seed: Some(42),
8904 ..Default::default()
8905 },
8906 )
8907}
8908
8909#[inline]
8914fn mask_bit(row: &[u8], j: usize) -> bool {
8915 (row.get(j >> 3).copied().unwrap_or(0) >> (j & 7)) & 1 != 0
8916}
8917
8918fn mask_gain() -> f32 {
8929 static G: std::sync::OnceLock<f32> = std::sync::OnceLock::new();
8930 *G.get_or_init(|| {
8931 std::env::var("CMF_FFN_MASK_GAIN")
8932 .ok()
8933 .and_then(|v| v.parse().ok())
8934 .unwrap_or(1.0)
8935 })
8936}
8937
8938fn zero_masked_cols(g: &mut [f32], rows: usize, inter: usize, row: &[u8]) {
8939 let fill = meanfill().and_then(|(i, v)| {
8942 let li = crate::gpu::cur_layer();
8943 (*i == inter && li >= 0).then(|| &v[li as usize * inter..(li as usize + 1) * inter])
8944 });
8945 for r in 0..rows {
8946 let base = r * inter;
8947 for (bi, &byte) in row.iter().enumerate() {
8948 if byte == 0xFF {
8949 continue;
8950 }
8951 let j0 = bi * 8;
8952 for bit in 0..8 {
8953 let j = j0 + bit;
8954 if j < inter && byte & (1 << bit) == 0 {
8955 g[base + j] = fill.map_or(0.0, |f| f[j]);
8956 }
8957 }
8958 }
8959 }
8960 let gain = mask_gain();
8961 if gain != 1.0 {
8962 for v in g[..rows * inter].iter_mut() {
8963 *v *= gain;
8964 }
8965 }
8966}
8967
8968#[inline]
8970fn tube_bit(row: Option<&[u8]>, i: usize) -> bool {
8971 row.is_none_or(|r| mask_bit(r, i))
8972}
8973
8974fn all_bits_on(row: &[u8], n: usize) -> bool {
8977 (0..n).all(|i| mask_bit(row, i))
8978}
8979
8980fn tube_topk() -> usize {
8988 static K: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
8989 *K.get_or_init(|| {
8990 std::env::var("CMF_TUBE_TOPK")
8991 .ok()
8992 .and_then(|v| v.parse().ok())
8993 .unwrap_or(0)
8994 })
8995}
8996
8997fn tube_score_oracle() -> bool {
8998 static O: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
8999 *O.get_or_init(|| {
9000 std::env::var("CMF_TUBE_SCORE").is_ok_and(|v| v == "oracle")
9001 })
9002}
9003
9004fn tube_ffn_routed(
9011 d: &DenseFfn,
9012 xs: &[f32],
9013 b: usize,
9014 pool: Option<&Pool>,
9015 mask_row: Option<&[u8]>,
9016 k: usize,
9017) -> Vec<f32> {
9018 let hidden = d.down_proj.rows();
9019 let core = d.gate_proj.rows();
9020 let core_full = mask_row.is_none_or(|r| all_bits_on(r, core));
9021 let mut out = match (b, core_full, mask_row) {
9022 (1, true, _) => dense_ffn(d, xs, pool),
9023 (1, false, Some(row)) => dense_ffn_masked(d, xs, pool, row),
9024 (_, true, _) => dense_ffn_batch(d, xs, b, pool, None),
9025 (_, false, row) => dense_ffn_batch(d, xs, b, pool, row),
9026 };
9027 let cand: Vec<usize> = (0..d.segs.len())
9028 .filter(|&i| tube_bit(mask_row, d.segs[i].start))
9029 .collect();
9030 if cand.is_empty() {
9031 return out;
9032 }
9033 let oracle = tube_score_oracle();
9037 let mut acts: Vec<Vec<f32>> = Vec::with_capacity(cand.len());
9038 let mut scores = vec![0f32; b * cand.len()];
9039 for (ci, &i) in cand.iter().enumerate() {
9040 let seg = &d.segs[i];
9041 let w = seg.width;
9042 let mut g = vec![0.0f32; b * w];
9043 if b == 1 {
9044 seg.gate.matvec(xs, &mut g, pool);
9045 } else {
9046 seg.gate.matmat(xs, b, &mut g, pool);
9047 }
9048 for v in g.iter_mut() {
9049 *v = Act::Silu.combine(*v, 1.0);
9050 }
9051 if !oracle {
9052 for t in 0..b {
9053 scores[t * cand.len() + ci] =
9054 g[t * w..(t + 1) * w].iter().map(|v| v * v).sum::<f32>();
9055 }
9056 }
9057 if oracle || b > 1 {
9058 let mut u = vec![0.0f32; b * w];
9059 if b == 1 {
9060 seg.up.matvec(xs, &mut u, pool);
9061 } else {
9062 seg.up.matmat(xs, b, &mut u, pool);
9063 }
9064 for (a, &v) in g.iter_mut().zip(u.iter()) {
9065 *a *= v;
9066 }
9067 if oracle {
9068 for t in 0..b {
9069 scores[t * cand.len() + ci] =
9070 g[t * w..(t + 1) * w].iter().map(|v| v * v).sum::<f32>();
9071 }
9072 }
9073 }
9074 acts.push(g);
9075 }
9076 let keep = k.min(cand.len());
9078 let mut scratch: Vec<f32> = Vec::new();
9079 for t in 0..b {
9080 let mut sc: Vec<(f32, usize)> = (0..cand.len())
9081 .map(|ci| (scores[t * cand.len() + ci], ci))
9082 .collect();
9083 sc.sort_unstable_by(|x, y| y.0.total_cmp(&x.0));
9084 let mut alive = vec![false; cand.len()];
9085 for &(_, ci) in sc.iter().take(keep) {
9086 alive[ci] = true;
9087 }
9088 if b > 1 {
9089 for (ci, a) in acts.iter_mut().enumerate() {
9090 if !alive[ci] {
9091 let w = d.segs[cand[ci]].width;
9092 a[t * w..(t + 1) * w].fill(0.0);
9093 }
9094 }
9095 } else {
9096 for (ci, &i) in cand.iter().enumerate() {
9100 if !alive[ci] {
9101 continue;
9102 }
9103 let seg = &d.segs[i];
9104 let w = seg.width;
9105 let g = &mut acts[ci];
9106 if !tube_score_oracle() {
9107 scratch.clear();
9108 scratch.resize(w, 0.0);
9109 seg.up.matvec(xs, &mut scratch, pool);
9110 for (a, &v) in g.iter_mut().zip(scratch.iter()) {
9111 *a *= v;
9112 }
9113 }
9114 let mut acc = vec![0.0f32; hidden];
9115 seg.down.matvec(g, &mut acc, pool);
9116 for (o, a) in out.iter_mut().zip(&acc) {
9117 *o += *a;
9118 }
9119 }
9120 }
9121 }
9122 if b > 1 {
9123 for (ci, &i) in cand.iter().enumerate() {
9124 let seg = &d.segs[i];
9125 let mut acc = vec![0.0f32; b * hidden];
9126 seg.down.matmat(&acts[ci], b, &mut acc, pool);
9127 for (o, a) in out.iter_mut().zip(&acc) {
9128 *o += *a;
9129 }
9130 }
9131 }
9132 out
9133}
9134
9135fn tube_ffn(
9141 d: &DenseFfn,
9142 xs: &[f32],
9143 b: usize,
9144 pool: Option<&Pool>,
9145 mask_row: Option<&[u8]>,
9146) -> Vec<f32> {
9147 if tube_topk() > 0 {
9148 return tube_ffn_routed(d, xs, b, pool, mask_row, tube_topk());
9149 }
9150 let hidden = d.down_proj.rows();
9151 let core = d.gate_proj.rows();
9152 let core_full = mask_row.is_none_or(|r| all_bits_on(r, core));
9153 let mut out = match (b, core_full, mask_row) {
9154 (1, true, _) => dense_ffn(d, xs, pool),
9155 (1, false, Some(row)) => dense_ffn_masked(d, xs, pool, row),
9156 (_, true, _) => dense_ffn_batch(d, xs, b, pool, None),
9157 (_, false, row) => dense_ffn_batch(d, xs, b, pool, row),
9158 };
9159 TUBE_SCRATCH.with(|sc| {
9160 let mut sc = sc.borrow_mut();
9161 let [g, u, acc] = &mut *sc;
9162 for seg in &d.segs {
9163 if !tube_bit(mask_row, seg.start) {
9164 continue;
9165 }
9166 let w = seg.width;
9167 g.resize(b * w, 0.0);
9168 if b == 1 && d.act == Act::Silu && QTensor::matvec_silu_mul(&seg.gate, &seg.up, xs, g, pool)
9169 {
9170 } else {
9172 u.resize(b * w, 0.0);
9173 if b == 1 {
9174 QTensor::matvec_many([&seg.gate, &seg.up], xs, [g, u], pool);
9175 } else {
9176 seg.gate.matmat(xs, b, g, pool);
9177 seg.up.matmat(xs, b, u, pool);
9178 }
9179 for i in 0..b * w {
9180 g[i] = d.act.combine(g[i], u[i]);
9181 }
9182 }
9183 acc.resize(b * hidden, 0.0);
9184 acc.fill(0.0);
9185 if b == 1 {
9186 seg.down.matvec(g, acc, pool);
9187 } else {
9188 seg.down.matmat(g, b, acc, pool);
9189 }
9190 for (o, a) in out.iter_mut().zip(acc.iter()) {
9191 *o += *a;
9192 }
9193 }
9194 out
9195 })
9196}
9197
9198thread_local! {
9199 static TUBE_SCRATCH: std::cell::RefCell<[Vec<f32>; 3]> =
9203 const { std::cell::RefCell::new([Vec::new(), Vec::new(), Vec::new()]) };
9204}
9205
9206fn dense_ffn_batch(
9207 d: &DenseFfn,
9208 xs: &[f32],
9209 b: usize,
9210 pool: Option<&Pool>,
9211 mask_row: Option<&[u8]>,
9212) -> Vec<f32> {
9213 let inter = d.gate_proj.rows();
9214 let hidden = d.down_proj.rows();
9215 if mask_row.is_none()
9223 && d.act == Act::Silu
9224 && b >= 32
9225 && crate::gpu::enabled_here()
9226 && !crate::gpu::mm_killed()
9227 && refit_dir().is_none()
9232 && !ffn_probe_active()
9237 {
9238 if let (Some((model, w1)), Some((_, w3)), Some((_, w2))) = (
9239 d.gate_proj.mapped_q4t(),
9240 d.up_proj.mapped_q4t(),
9241 d.down_proj.mapped_q4t(),
9242 ) {
9243 let mut out = vec![0.0f32; b * hidden];
9244 if crate::gpu::q4t_ffn(model, w1, w3, w2, xs, b, hidden, inter, &mut out) {
9245 return out;
9246 }
9247 }
9248 if let (Some((model, w1)), Some((_, w3)), Some((_, w2))) = (
9253 d.gate_proj.mapped_q4tp(),
9254 d.up_proj.mapped_q4tp(),
9255 d.down_proj.mapped_q4tp(),
9256 ) {
9257 let mut out = vec![0.0f32; b * hidden];
9258 if crate::gpu::q4tp_ffn(model, w1, w3, w2, xs, b, hidden, inter, &mut out) {
9259 return out;
9260 }
9261 }
9262 }
9263 let mut g = vec![0.0f32; b * inter];
9264 d.gate_proj.matmat(xs, b, &mut g, pool);
9265 let mut u = vec![0.0f32; b * inter];
9266 d.up_proj.matmat(xs, b, &mut u, pool);
9267 if gate_topk() > 0 && d.act == Act::Silu {
9268 for t in 0..b {
9269 let row = &mut g[t * inter..(t + 1) * inter];
9270 for v in row.iter_mut() {
9271 *v = Act::Silu.combine(*v, 1.0);
9272 }
9273 keep_top_k(row, gate_topk());
9274 }
9275 for i in 0..b * inter {
9276 g[i] *= u[i];
9277 }
9278 } else {
9279 for i in 0..b * inter {
9280 g[i] = d.act.combine(g[i], u[i]);
9281 }
9282 }
9283 if let Some(row) = mask_row {
9284 zero_masked_cols(&mut g, b, inter, row);
9285 }
9286 if oracle_topk() > 0 {
9287 for t in 0..b {
9288 keep_top_k(&mut g[t * inter..(t + 1) * inter], oracle_topk());
9289 }
9290 }
9291 let mut out = vec![0.0f32; b * hidden];
9292 d.down_proj.matmat(&g, b, &mut out, pool);
9293 if refit_dir().is_some() {
9294 let li = crate::gpu::cur_layer();
9295 if li >= 0 {
9296 refit_accumulate(li as usize, &g, b, inter, &out, hidden, pool);
9297 }
9298 }
9299 FFN_PROBE.with(|pr| {
9303 if let Some(acc) = pr.borrow_mut().as_mut() {
9304 let li = crate::gpu::cur_layer();
9305 if li < 0 {
9306 return;
9307 }
9308 let Some(row) = acc.get_mut(li as usize) else {
9309 return;
9310 };
9311 let sq = probe_sq();
9312 for t in 0..b {
9313 for (a, &v) in row.iter_mut().zip(&g[t * inter..(t + 1) * inter]) {
9314 *a += if sq { (v as f64) * (v as f64) } else { (v as f64).abs() };
9315 }
9316 }
9317 }
9318 });
9319 out
9320}
9321
9322fn accumulate_act(m: &MoeFfn, xs: &[f32], b: usize) {
9327 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
9328 static DUMP: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
9329 let on = *ON.get_or_init(|| std::env::var("CMF_RMS_TRACE").is_ok());
9330 let dump = *DUMP.get_or_init(|| std::env::var("CMF_ACT_DUMP").is_ok());
9331 if (!on && !dump) || b == 0 {
9332 return;
9333 }
9334 let hidden = xs.len() / b;
9335 if on {
9336 let mut acc = m.act_sq.borrow_mut();
9337 if acc.len() < hidden {
9338 acc.resize(hidden, 0.0);
9339 }
9340 for t in 0..b {
9341 let row = &xs[t * hidden..(t + 1) * hidden];
9342 for (a, &v) in acc.iter_mut().zip(row) {
9343 *a += (v as f64) * (v as f64);
9344 }
9345 }
9346 }
9347 if dump {
9348 let cap: usize = std::env::var("CMF_ACT_DUMP_ROWS")
9351 .ok()
9352 .and_then(|v| v.parse().ok())
9353 .unwrap_or(4096);
9354 let mut rows = m.act_rows.borrow_mut();
9355 if rows.len() < cap * hidden {
9356 let take = b.min((cap * hidden - rows.len()) / hidden.max(1));
9357 rows.extend_from_slice(&xs[..take * hidden]);
9358 }
9359 }
9360}
9361
9362#[derive(Clone, Copy)]
9365struct SendVecs(*mut Vec<f32>);
9366unsafe impl Send for SendVecs {}
9367unsafe impl Sync for SendVecs {}
9368impl SendVecs {
9369 #[inline]
9370 fn at(self, i: usize) -> *mut Vec<f32> {
9371 unsafe { self.0.add(i) }
9372 }
9373}
9374
9375fn moe_ffn_batch(
9376 m: &MoeFfn,
9377 xs: &[f32],
9378 b: usize,
9379 hidden: usize,
9380 pool: Option<&Pool>,
9381 allowed: Option<&[bool]>,
9382) -> Vec<f32> {
9383 accumulate_act(m, xs, b);
9384 let ne = m.experts.len();
9385 let mut logits = vec![0.0f32; b * ne];
9386 match &m.resonance {
9387 Some(r) => {
9388 let hdim = xs.len() / b.max(1);
9389 for bi in 0..b {
9390 r.scores(&xs[bi * hdim..(bi + 1) * hdim], &mut logits[bi * ne..(bi + 1) * ne]);
9391 }
9392 }
9393 None => m.router.matmat(xs, b, &mut logits, pool),
9394 }
9395
9396 let mut assign: Vec<Vec<(usize, f32)>> = vec![Vec::new(); ne];
9399 {
9400 let mut st = m.stats.borrow_mut();
9401 if st.len() < ne {
9402 st.resize(ne, 0);
9403 }
9404 for bi in 0..b {
9405 let (idx, p, wsum) = moe_route(&logits[bi * ne..(bi + 1) * ne], m, allowed);
9406 for &e in &idx {
9407 st[e] += 1;
9408 assign[e].push((bi, p[e] / wsum));
9409 }
9410 }
9411 }
9412
9413 let mut out = vec![0.0f32; b * hidden];
9414 let cols = m.experts[0].gate_proj.cols();
9415 let run_expert = |d: &DenseFfn, list: &[(usize, f32)], out: &mut [f32]| {
9416 let sb = list.len();
9417 let mut sub = vec![0.0f32; sb * cols];
9418 for (k, &(bi, _)) in list.iter().enumerate() {
9419 sub[k * cols..(k + 1) * cols].copy_from_slice(&xs[bi * cols..(bi + 1) * cols]);
9420 }
9421 let eo = dense_ffn_batch(d, &sub, sb, pool, None);
9422 for (k, &(bi, w)) in list.iter().enumerate() {
9423 for i in 0..hidden {
9424 out[bi * hidden + i] += w * eo[k * hidden + i];
9425 }
9426 }
9427 };
9428 let active: Vec<usize> = (0..ne).filter(|&e| !assign[e].is_empty()).collect();
9434 if pool.is_some() && active.len() >= 8 {
9435 let mut panels: Vec<Vec<f32>> = vec![Vec::new(); active.len()];
9436 {
9437 let panel_ptr = SendVecs(panels.as_mut_ptr());
9438 let experts = &m.experts;
9441 let (active_r, assign_r) = (&active, &assign);
9442 let run = |start: usize, end: usize| {
9443 for ai in start..end {
9444 let e = active_r[ai];
9445 let list = &assign_r[e];
9446 let sb = list.len();
9447 let mut sub = vec![0.0f32; sb * cols];
9448 for (k, &(bi, _)) in list.iter().enumerate() {
9449 sub[k * cols..(k + 1) * cols]
9450 .copy_from_slice(&xs[bi * cols..(bi + 1) * cols]);
9451 }
9452 unsafe {
9454 *panel_ptr.at(ai) = dense_ffn_batch(&experts[e], &sub, sb, None, None);
9455 }
9456 }
9457 };
9458 match pool {
9459 Some(p) => p.run_rows(active.len(), &run),
9460 None => run(0, active.len()),
9461 }
9462 }
9463 for (ai, &e) in active.iter().enumerate() {
9464 for (k, &(bi, w)) in assign[e].iter().enumerate() {
9465 let eo = &panels[ai][k * hidden..(k + 1) * hidden];
9466 for i in 0..hidden {
9467 out[bi * hidden + i] += w * eo[i];
9468 }
9469 }
9470 }
9471 } else {
9472 for &e in &active {
9473 run_expert(&m.experts[e], &assign[e], &mut out);
9474 }
9475 }
9476 if let Some((se, gate)) = &m.shared {
9477 let all: Vec<(usize, f32)> = if let Some(gate) = gate {
9478 let mut gl = vec![0.0f32; b];
9479 gate.matmat(xs, b, &mut gl, pool);
9480 (0..b)
9481 .map(|bi| (bi, 1.0 / (1.0 + (-gl[bi]).exp())))
9482 .collect()
9483 } else {
9484 (0..b).map(|bi| (bi, 1.0)).collect()
9485 };
9486 run_expert(se, &all, &mut out);
9487 }
9488 out
9489}
9490
9491thread_local! {
9492 static FFN_SCRATCH: std::cell::RefCell<[Vec<f32>; 4]> =
9496 const { std::cell::RefCell::new([Vec::new(), Vec::new(), Vec::new(), Vec::new()]) };
9497}
9498
9499fn dense_ffn(d: &DenseFfn, x: &[f32], pool: Option<&Pool>) -> Vec<f32> {
9501 if gate_topk() > 0
9504 && let Some(out) = dense_ffn_dynamic(d, x, pool, gate_topk())
9505 {
9506 return out;
9507 }
9508 if crate::gpu::enabled_here()
9519 && (d.gate_proj.rows() >= crate::gpu::min_rows() || d.gate_proj.is_q1())
9520 {
9521 let arm = if d.gate_proj.is_q1() && crate::gpu::q1_force() {
9522 crate::gpu::ProbeArm::Gpu
9523 } else {
9524 crate::gpu::probe_arm(crate::gpu::OpClass::Ffn)
9525 };
9526 match arm {
9527 crate::gpu::ProbeArm::Gpu => {
9528 let t0 = std::time::Instant::now();
9529 if let Some(out) = dense_ffn_gpu(d, x, pool) {
9530 crate::gpu::probe_record(crate::gpu::OpClass::Ffn, true, t0.elapsed());
9531 return out;
9532 }
9533 crate::gpu::probe_note_decline(crate::gpu::OpClass::Ffn);
9537 }
9538 crate::gpu::ProbeArm::CpuTimed => {
9539 let t0 = std::time::Instant::now();
9540 let out = crate::gpu::cpu_scope(|| dense_ffn_cpu(d, x, pool));
9541 crate::gpu::probe_record(crate::gpu::OpClass::Ffn, false, t0.elapsed());
9542 return out;
9543 }
9544 crate::gpu::ProbeArm::Cpu => {
9545 return crate::gpu::cpu_scope(|| dense_ffn_cpu(d, x, pool));
9546 }
9547 }
9548 }
9549 dense_ffn_cpu(d, x, pool)
9550}
9551
9552fn dense_ffn_cpu(d: &DenseFfn, x: &[f32], pool: Option<&Pool>) -> Vec<f32> {
9554 let inter = d.gate_proj.rows();
9555 FFN_SCRATCH.with(|s| {
9556 let mut s = s.borrow_mut();
9557 let [g, u, ..] = &mut *s;
9558 g.resize(inter, 0.0);
9559 if gate_topk() > 0 {
9562 u.resize(inter, 0.0);
9566 QTensor::matvec_many([&d.gate_proj, &d.up_proj], x, [g, u], pool);
9567 for i in 0..inter {
9568 g[i] = Act::Silu.combine(g[i], 1.0);
9569 }
9570 keep_top_k(g, gate_topk());
9571 for i in 0..inter {
9572 g[i] *= u[i];
9573 }
9574 } else if d.act == Act::Silu && QTensor::matvec_silu_mul(&d.gate_proj, &d.up_proj, x, g, pool) {
9575 } else {
9577 u.resize(inter, 0.0);
9578 QTensor::matvec_many([&d.gate_proj, &d.up_proj], x, [g, u], pool);
9580 for i in 0..inter {
9581 g[i] = d.act.combine(g[i], u[i]);
9582 }
9583 }
9584 FFN_PROBE.with(|pr| {
9592 if let Some(acc) = pr.borrow_mut().as_mut() {
9593 let li = crate::gpu::cur_layer();
9594 if li >= 0 {
9595 if let Some(row) = acc.get_mut(li as usize) {
9596 match probe_topk() {
9597 0 if probe_sq() => {
9598 for (a, &v) in row.iter_mut().zip(g.iter()) {
9599 *a += (v as f64) * (v as f64);
9600 }
9601 }
9602 0 if probe_signed() => {
9603 for (a, &v) in row.iter_mut().zip(g.iter()) {
9604 *a += v as f64;
9605 }
9606 }
9607 0 => {
9608 for (a, &v) in row.iter_mut().zip(g.iter()) {
9609 *a += (v as f64).abs();
9610 }
9611 }
9612 k => {
9613 let n = g.len();
9614 let k = k.min(n);
9615 let mut mag: Vec<f32> = g.iter().map(|v| v.abs()).collect();
9616 let (_, kth, _) = mag.select_nth_unstable_by(k - 1, |a, b| {
9617 b.partial_cmp(a).unwrap_or(std::cmp::Ordering::Equal)
9618 });
9619 let thr = *kth;
9620 for (a, &v) in row.iter_mut().zip(g.iter()) {
9621 if v.abs() >= thr {
9622 *a += 1.0;
9623 }
9624 }
9625 }
9626 }
9627 }
9628 }
9629 }
9630 });
9631 if oracle_topk() > 0 {
9632 keep_top_k(g, oracle_topk());
9633 }
9634 {
9635 let li = crate::gpu::cur_layer();
9636 if li >= 0 {
9637 adump_row(li as usize, g);
9638 }
9639 }
9640 let mut out = attention::take_buf(d.down_proj.rows());
9641 d.down_proj.matvec(g, &mut out, pool);
9642 out
9643 })
9644}
9645
9646pub struct RefitAcc {
9659 pub support: Vec<u32>,
9660 pub gss: Vec<f32>,
9661 pub ya: Vec<f32>,
9662 pub hidden: usize,
9663 pub tokens: u64,
9664 pub buf_g: Vec<f32>,
9670 pub buf_o: Vec<f32>,
9671 pub buf_t: usize,
9672}
9673
9674type RefitState = (
9678 std::collections::HashMap<usize, RefitAcc>,
9679 Vec<f32>,
9680);
9681
9682static REFIT: std::sync::OnceLock<Option<(String, std::sync::Mutex<RefitState>)>> =
9683 std::sync::OnceLock::new();
9684
9685fn ffn_probe_active() -> bool {
9688 FFN_PROBE.with(|p| p.borrow().is_some())
9689}
9690
9691fn refit_dir() -> Option<&'static (String, std::sync::Mutex<RefitState>)> {
9692 REFIT
9693 .get_or_init(|| {
9694 std::env::var("CMF_FFN_REFIT").ok().map(|d| {
9695 (
9696 d,
9697 std::sync::Mutex::new((std::collections::HashMap::new(), Vec::new())),
9698 )
9699 })
9700 })
9701 .as_ref()
9702}
9703
9704fn refit_accumulate(
9706 li: usize,
9707 g: &[f32],
9708 b: usize,
9709 inter: usize,
9710 out: &[f32],
9711 hidden: usize,
9712 pool: Option<&Pool>,
9713) {
9714 let Some((dir, map)) = refit_dir() else {
9715 return;
9716 };
9717 static SPAN: std::sync::OnceLock<(usize, usize)> = std::sync::OnceLock::new();
9718 let (from, to) = *SPAN.get_or_init(|| {
9719 let g = |k: &str, d: usize| {
9720 std::env::var(k)
9721 .ok()
9722 .and_then(|v| v.parse().ok())
9723 .unwrap_or(d)
9724 };
9725 (g("CMF_FFN_REFIT_FROM", 0), g("CMF_FFN_REFIT_TO", usize::MAX))
9726 });
9727 if li < from || li > to {
9728 return;
9729 }
9730 let mut guard = map.lock().unwrap();
9731 let (map, shared) = &mut *guard;
9732 let acc = match map.entry(li) {
9733 std::collections::hash_map::Entry::Occupied(e) => e.into_mut(),
9734 std::collections::hash_map::Entry::Vacant(e) => {
9735 let path = format!("{dir}/support.{li}.u32");
9736 let Ok(bytes) = std::fs::read(&path) else {
9737 eprintln!("refit: no {path} — layer {li} skipped");
9738 return;
9739 };
9740 let n = u32::from_le_bytes(bytes[0..4].try_into().unwrap()) as usize;
9741 let support: Vec<u32> = bytes[4..4 + n * 4]
9742 .chunks_exact(4)
9743 .map(|c| u32::from_le_bytes([c[0], c[1], c[2], c[3]]))
9744 .collect();
9745 eprintln!(
9746 "refit: layer {li} support {n} ({:.0} MB of accumulator)",
9747 (n * n + hidden * n) as f64 * 4.0 / 1e6
9748 );
9749 e.insert(RefitAcc {
9750 gss: vec![0.0; n * n],
9751 ya: vec![0.0; hidden * n],
9752 buf_g: Vec::new(),
9753 buf_o: Vec::new(),
9754 buf_t: 0,
9755 support,
9756 hidden,
9757 tokens: 0,
9758 })
9759 }
9760 };
9761 let ns = acc.support.len();
9762 let cap = refit_batch();
9764 if acc.buf_g.is_empty() {
9765 acc.buf_g = vec![0.0; ns * cap];
9766 acc.buf_o = vec![0.0; hidden * cap];
9767 }
9768 let take = b.min(cap - acc.buf_t);
9769 for t in 0..take {
9770 let col = acc.buf_t + t;
9771 for (j, &n) in acc.support.iter().enumerate() {
9772 acc.buf_g[j * cap + col] = g[t * inter + n as usize];
9773 }
9774 for h in 0..hidden {
9775 acc.buf_o[h * cap + col] = out[t * hidden + h];
9776 }
9777 }
9778 acc.buf_t += take;
9779 acc.tokens += take as u64;
9780 if acc.buf_t < cap {
9781 return;
9782 }
9783 let bt = acc.buf_t;
9784 acc.buf_t = 0;
9785 let RefitAcc {
9795 gss, ya, buf_g, buf_o, ..
9796 } = acc;
9797 let need = (ns * ns).max(hidden * ns);
9798 if shared.len() < need {
9799 shared.resize(need, 0.0);
9800 }
9801 let scratch = &mut shared[..];
9802 let _ = bt;
9803 if crate::gpu::gemm_nt_f32_transient(buf_g, buf_g, &mut scratch[..ns * ns], ns, cap, ns) {
9804 add_into(gss, &scratch[..ns * ns], pool);
9805 if crate::gpu::gemm_nt_f32_transient(buf_o, buf_g, &mut scratch[..hidden * ns], hidden, cap, ns) {
9806 add_into(ya, &scratch[..hidden * ns], pool);
9807 } else {
9808 accum_outer_t(ya, hidden, ns, cap, buf_o, buf_g, pool);
9809 }
9810 } else {
9811 accum_outer_t(gss, ns, ns, cap, buf_g, buf_g, pool);
9812 accum_outer_t(ya, hidden, ns, cap, buf_o, buf_g, pool);
9813 }
9814 }
9818
9819fn refit_batch() -> usize {
9821 static B: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
9822 *B.get_or_init(|| {
9823 std::env::var("CMF_FFN_REFIT_BATCH")
9824 .ok()
9825 .and_then(|v| v.parse().ok())
9826 .unwrap_or(4096)
9827 })
9828}
9829
9830fn accum_outer_t(
9833 c: &mut [f32],
9834 m: usize,
9835 n: usize,
9836 b: usize,
9837 left: &[f32],
9838 right: &[f32],
9839 pool: Option<&Pool>,
9840) {
9841 let ptr = SendMut(c.as_mut_ptr());
9842 let body = |i: usize| {
9843 let ptr = &ptr;
9844 let row = unsafe { std::slice::from_raw_parts_mut(ptr.0.add(i * n), n) };
9845 for t in 0..b {
9846 let a = left[i * b + t];
9847 if a == 0.0 {
9848 continue;
9849 }
9850 for (j, o) in row.iter_mut().enumerate() {
9851 *o += a * right[j * b + t];
9852 }
9853 }
9854 };
9855 match pool {
9856 Some(p) if m > 1 => p.run_rows(m, &|s, e| {
9857 for i in s..e {
9858 body(i);
9859 }
9860 }),
9861 _ => {
9862 for i in 0..m {
9863 body(i);
9864 }
9865 }
9866 }
9867}
9868
9869fn add_into(dst: &mut [f32], src: &[f32], pool: Option<&Pool>) {
9872 let n = dst.len().min(src.len());
9873 match pool {
9874 Some(p) if n >= 1 << 16 => {
9875 let ptr = SendMut(dst.as_mut_ptr());
9876 let f = |s: usize, e: usize| {
9877 let ptr = &ptr;
9878 for blk in s..e {
9879 let (a, b) = (blk * 4096, ((blk + 1) * 4096).min(n));
9880 for i in a..b {
9881 unsafe { *ptr.0.add(i) += src[i] };
9882 }
9883 }
9884 };
9885 p.run_rows(n.div_ceil(4096), &f);
9886 }
9887 _ => {
9888 for (d, v) in dst.iter_mut().zip(&src[..n]) {
9889 *d += *v;
9890 }
9891 }
9892 }
9893}
9894
9895fn accum_outer(
9900 c: &mut [f32],
9901 m: usize,
9902 n: usize,
9903 b: usize,
9904 left: &[f32],
9905 right: &[f32],
9906 pool: Option<&Pool>,
9907) {
9908 const TILE: usize = 32;
9909 let tiles = m.div_ceil(TILE);
9910 let cp = SendMut(c.as_mut_ptr());
9911 let body = |ti: usize| {
9912 let cp = &cp;
9913 let i0 = ti * TILE;
9914 let i1 = (i0 + TILE).min(m);
9915 for t in 0..b {
9916 let r = &right[t * n..t * n + n];
9917 for i in i0..i1 {
9918 let a = left[i * b + t];
9919 if a == 0.0 {
9920 continue;
9921 }
9922 let row = unsafe { std::slice::from_raw_parts_mut(cp.0.add(i * n), n) };
9924 for (o, v) in row.iter_mut().zip(r) {
9925 *o += a * *v;
9926 }
9927 }
9928 }
9929 };
9930 match pool {
9931 Some(p) if tiles > 1 => p.run_rows(tiles, &|s, e| {
9932 for ti in s..e {
9933 body(ti);
9934 }
9935 }),
9936 _ => {
9937 for ti in 0..tiles {
9938 body(ti);
9939 }
9940 }
9941 }
9942}
9943
9944pub fn refit_flush() -> usize {
9946 let Some((dir, map)) = refit_dir() else {
9947 return 0;
9948 };
9949 let guard = map.lock().unwrap();
9950 let mut n = 0;
9951 for (li, acc) in guard.0.iter() {
9952 let w = |name: &str, v: &[f32]| {
9955 let path = format!("{dir}/{name}.{li}.f32");
9956 let bytes: Vec<u8> = v.iter().flat_map(|x| x.to_le_bytes()).collect();
9957 match std::fs::write(&path, &bytes) {
9958 Ok(()) => {}
9959 Err(e) => eprintln!("refit: FAILED to write {path} ({} MB): {e}", bytes.len() / 1_000_000),
9960 }
9961 };
9962 w("gss", &acc.gss);
9963 w("ya", &acc.ya);
9964 println!(
9965 "refit L{li}: {} support, {} tokens, hidden {}",
9966 acc.support.len(),
9967 acc.tokens,
9968 acc.hidden
9969 );
9970 n += 1;
9971 }
9972 n
9973}
9974
9975fn adump_row(li: usize, g: &[f32]) {
9980 use std::io::Write as _;
9981 static FILES: std::sync::OnceLock<
9982 Option<(String, std::sync::Mutex<std::collections::HashMap<usize, std::fs::File>>)>,
9983 > = std::sync::OnceLock::new();
9984 let Some((prefix, map)) = FILES
9985 .get_or_init(|| {
9986 std::env::var("CMF_FFN_ADUMP")
9987 .ok()
9988 .map(|p| (p, std::sync::Mutex::new(std::collections::HashMap::new())))
9989 })
9990 .as_ref()
9991 else {
9992 return;
9993 };
9994 static SPAN: std::sync::OnceLock<(usize, usize)> = std::sync::OnceLock::new();
9997 let (from, to) = *SPAN.get_or_init(|| {
9998 let g = |k: &str, d: usize| {
9999 std::env::var(k)
10000 .ok()
10001 .and_then(|v| v.parse().ok())
10002 .unwrap_or(d)
10003 };
10004 (g("CMF_FFN_ADUMP_FROM", 0), g("CMF_FFN_ADUMP_TO", usize::MAX))
10005 });
10006 if li < from || li > to {
10007 return;
10008 }
10009 let mut map = map.lock().unwrap();
10010 let f = map.entry(li).or_insert_with(|| {
10011 std::fs::File::create(format!("{prefix}.{li}.f16")).expect("adump file")
10012 });
10013 let mut bytes = Vec::with_capacity(g.len() * 2);
10014 for v in g {
10015 bytes.extend_from_slice(&cortiq_core::quant::f32_to_f16(*v).to_le_bytes());
10016 }
10017 let _ = f.write_all(&bytes);
10018}
10019
10020fn oracle_topk() -> usize {
10026 static K: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
10027 *K.get_or_init(|| {
10028 std::env::var("CMF_FFN_ORACLE_TOPK")
10029 .ok()
10030 .and_then(|v| v.parse().ok())
10031 .unwrap_or(0)
10032 })
10033}
10034
10035fn gate_topk() -> usize {
10041 static K: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
10042 *K.get_or_init(|| {
10043 std::env::var("CMF_FFN_GATE_TOPK")
10044 .ok()
10045 .and_then(|v| v.parse().ok())
10046 .unwrap_or(0)
10047 })
10048}
10049
10050fn gate_block() -> usize {
10057 static B: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
10058 *B.get_or_init(|| {
10059 std::env::var("CMF_FFN_GATE_BLOCK")
10060 .ok()
10061 .and_then(|v| v.parse().ok())
10062 .unwrap_or(1)
10063 })
10064}
10065
10066fn keep_top_blocks(g: &mut [f32], keep_n: usize, block: usize) {
10068 let n = g.len();
10069 let nb = n.div_ceil(block);
10070 let kb = (keep_n.div_ceil(block)).clamp(1, nb);
10071 if kb >= nb {
10072 return;
10073 }
10074 let mut score: Vec<f32> = (0..nb)
10075 .map(|b| {
10076 g[b * block..((b + 1) * block).min(n)]
10077 .iter()
10078 .map(|v| v * v)
10079 .sum::<f32>()
10080 })
10081 .collect();
10082 let mut ord = score.clone();
10083 let (_, kth, _) = ord.select_nth_unstable_by(kb - 1, |a, b| {
10084 b.partial_cmp(a).unwrap_or(std::cmp::Ordering::Equal)
10085 });
10086 let thr = *kth;
10087 for b in 0..nb {
10088 if score[b] < thr {
10089 g[b * block..((b + 1) * block).min(n)].fill(0.0);
10090 }
10091 }
10092 score.clear();
10093}
10094
10095fn keep_top_k(g: &mut [f32], k: usize) {
10097 if gate_block() > 1 {
10098 return keep_top_blocks(g, k, gate_block());
10099 }
10100 let n = g.len();
10101 if k == 0 || k >= n {
10102 return;
10103 }
10104 let mut mag: Vec<f32> = g.iter().map(|v| v.abs()).collect();
10105 let (_, kth, _) = mag.select_nth_unstable_by(k - 1, |a, b| {
10106 b.partial_cmp(a).unwrap_or(std::cmp::Ordering::Equal)
10107 });
10108 let thr = *kth;
10109 for v in g.iter_mut() {
10110 if v.abs() < thr {
10111 *v = 0.0;
10112 }
10113 }
10114}
10115
10116fn probe_sq() -> bool {
10120 static S: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
10121 *S.get_or_init(|| std::env::var("CMF_FFN_PROBE_SQ").is_ok())
10122}
10123
10124fn probe_signed() -> bool {
10128 static S: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
10129 *S.get_or_init(|| std::env::var("CMF_FFN_PROBE_SIGNED").is_ok())
10130}
10131
10132fn meanfill() -> Option<&'static (usize, Vec<f32>)> {
10140 static M: std::sync::OnceLock<Option<(usize, Vec<f32>)>> = std::sync::OnceLock::new();
10141 M.get_or_init(|| {
10142 let p = std::env::var("CMF_FFN_MEANFILL").ok()?;
10143 let b = std::fs::read(&p).ok()?;
10144 let inter = u32::from_le_bytes(b[4..8].try_into().ok()?) as usize;
10145 let vals: Vec<f32> = b[8..]
10146 .chunks_exact(4)
10147 .map(|c| f32::from_le_bytes([c[0], c[1], c[2], c[3]]))
10148 .collect();
10149 eprintln!("meanfill: {} value(s), inter {inter}", vals.len());
10150 Some((inter, vals))
10151 })
10152 .as_ref()
10153}
10154
10155fn probe_topk() -> usize {
10158 static K: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
10159 *K.get_or_init(|| {
10160 std::env::var("CMF_FFN_PROBE_TOPK")
10161 .ok()
10162 .and_then(|v| v.parse().ok())
10163 .unwrap_or(0)
10164 })
10165}
10166
10167thread_local! {
10168 static FFN_PROBE: std::cell::RefCell<Option<Vec<Vec<f64>>>> =
10171 const { std::cell::RefCell::new(None) };
10172}
10173
10174fn dense_ffn_dynamic(d: &DenseFfn, x: &[f32], pool: Option<&Pool>, k: usize) -> Option<Vec<f32>> {
10187 let dt = d.down_t.as_ref()?;
10188 let inter = d.gate_proj.rows();
10189 let hidden = dt.cols();
10190 if k == 0 || k >= inter || d.act != Act::Silu {
10191 return None;
10192 }
10193 DYN_SCRATCH.with(|sc| {
10194 let mut sc = sc.borrow_mut();
10195 let DynScratch { g, mag, live, parts } = &mut *sc;
10196 g.resize(inter, 0.0);
10197 d.gate_proj.matvec(x, g, pool);
10198 for v in g.iter_mut() {
10199 *v = inference::silu(*v);
10200 }
10201 mag.clear();
10204 mag.extend(g.iter().map(|v| v.abs()));
10205 let (_, kth, _) = mag.select_nth_unstable_by(k - 1, |a, b| {
10206 b.partial_cmp(a).unwrap_or(std::cmp::Ordering::Equal)
10207 });
10208 let thr = *kth;
10209 live.clear();
10210 live.extend((0..inter as u32).filter(|&n| g[n as usize].abs() >= thr));
10211 let mut out = vec![0.0f32; hidden];
10212 match pool {
10213 Some(p) if live.len() >= 64 => {
10214 let nw = p.n_workers() + 1;
10215 parts.clear();
10216 parts.resize(nw * hidden, 0.0);
10217 let ptr = SendMut(parts.as_mut_ptr());
10218 let n = live.len();
10219 let live_ref: &[u32] = live;
10220 let g_ref: &[f32] = g;
10221 p.run(&|w, workers| {
10222 let chunk = n.div_ceil(workers);
10223 let (s, e) = (w * chunk, ((w + 1) * chunk).min(n));
10224 if s >= e {
10225 return;
10226 }
10227 WORKER_SCRATCH.with(|ws| {
10228 let mut ws = ws.borrow_mut();
10229 let [scratch, acc] = &mut *ws;
10230 scratch.resize(hidden.max(x.len()), 0.0);
10231 acc.clear();
10232 acc.resize(hidden, 0.0);
10233 for (o, &nrm) in live_ref[s..e].iter().enumerate() {
10234 if let Some(&nx) = live_ref[s..e].get(o + 1) {
10237 d.up_proj.prefetch_row(nx as usize);
10238 dt.prefetch_row(nx as usize);
10239 }
10240 let idx = nrm as usize;
10241 let up = d.up_proj.row_dot(idx, x, scratch);
10242 let a = g_ref[idx] * up;
10243 if a != 0.0 {
10244 dt.add_row_scaled(idx, a, acc, scratch);
10245 }
10246 }
10247 for (j, v) in acc.iter().enumerate() {
10248 unsafe { *ptr.at(w * hidden + j) = *v };
10249 }
10250 });
10251 });
10252 for w in 0..nw {
10253 for (j, o) in out.iter_mut().enumerate() {
10254 *o += parts[w * hidden + j];
10255 }
10256 }
10257 }
10258 _ => {
10259 WORKER_SCRATCH.with(|ws| {
10260 let mut ws = ws.borrow_mut();
10261 let [scratch, _acc] = &mut *ws;
10262 scratch.resize(hidden.max(x.len()), 0.0);
10263 for &nrm in live.iter() {
10264 let idx = nrm as usize;
10265 let up = d.up_proj.row_dot(idx, x, scratch);
10266 let a = g[idx] * up;
10267 if a != 0.0 {
10268 dt.add_row_scaled(idx, a, &mut out, scratch);
10269 }
10270 }
10271 });
10272 }
10273 }
10274 Some(out)
10275 })
10276}
10277
10278struct DynScratch {
10281 g: Vec<f32>,
10282 mag: Vec<f32>,
10283 live: Vec<u32>,
10284 parts: Vec<f32>,
10285}
10286
10287thread_local! {
10288 static DYN_SCRATCH: std::cell::RefCell<DynScratch> = const {
10289 std::cell::RefCell::new(DynScratch {
10290 g: Vec::new(),
10291 mag: Vec::new(),
10292 live: Vec::new(),
10293 parts: Vec::new(),
10294 })
10295 };
10296 static WORKER_SCRATCH: std::cell::RefCell<[Vec<f32>; 2]> =
10298 const { std::cell::RefCell::new([Vec::new(), Vec::new()]) };
10299}
10300
10301fn dense_ffn_masked(d: &DenseFfn, x: &[f32], pool: Option<&Pool>, mask_row: &[u8]) -> Vec<f32> {
10306 let inter = d.gate_proj.rows();
10307 FFN_SCRATCH.with(|s| {
10308 let mut s = s.borrow_mut();
10309 let [g, u, ..] = &mut *s;
10310 g.resize(inter, 0.0);
10311 if d.act == Act::Silu && QTensor::matvec_silu_mul(&d.gate_proj, &d.up_proj, x, g, pool) {
10312 } else {
10314 u.resize(inter, 0.0);
10315 QTensor::matvec_many([&d.gate_proj, &d.up_proj], x, [g, u], pool);
10316 for i in 0..inter {
10317 g[i] = d.act.combine(g[i], u[i]);
10318 }
10319 }
10320 zero_masked_cols(g, 1, inter, mask_row);
10321 let mut out = attention::take_buf(d.down_proj.rows());
10322 d.down_proj.matvec(g, &mut out, pool);
10323 out
10324 })
10325}
10326
10327fn dense_ffn_gpu(d: &DenseFfn, x: &[f32], _pool: Option<&Pool>) -> Option<Vec<f32>> {
10333 if d.act != Act::Silu {
10335 return None;
10336 }
10337 if d.gate_proj.rows() < crate::gpu::min_rows() && !d.gate_proj.is_q1() {
10340 return None;
10341 }
10342 let mut jobs: Vec<crate::gpu::MoeJob> = Vec::with_capacity(1);
10343 let mut model_ref = None;
10344 moe_push_job(d, x, 1.0, &mut jobs, &mut model_ref)?;
10345 let model = model_ref?;
10346 let hidden = jobs[0].down.1;
10347 let mut out = attention::take_buf(hidden);
10348 if crate::gpu::moe_block(&model, &jobs, &mut out) {
10349 Some(out)
10350 } else {
10351 let mut out = out;
10352 attention::recycle_buf(&mut out);
10353 None
10354 }
10355}
10356
10357#[allow(clippy::type_complexity)]
10362#[allow(clippy::type_complexity)]
10363pub(crate) fn moe_parts(
10364 t: &QTensor,
10365) -> Option<(
10366 &std::sync::Arc<cortiq_core::CmfModel>,
10367 usize,
10368 usize,
10369 usize,
10370 &[f32],
10371 &[f32],
10372 bool,
10373 bool,
10374 bool,
10375)> {
10376 match t {
10377 QTensor::Mapped {
10378 model,
10379 idx,
10380 dtype: dt @ (cortiq_core::TensorDtype::Q8_2f | cortiq_core::TensorDtype::Q8Row),
10381 rows,
10382 cols,
10383 row_scale,
10384 col_field,
10385 ..
10386 } if (*dt == cortiq_core::TensorDtype::Q8Row) || !col_field.is_empty() => Some((
10387 model, *idx, *rows, *cols, row_scale, col_field, false, false, false,
10388 )),
10389 QTensor::Mapped {
10391 model,
10392 idx,
10393 dtype: cortiq_core::TensorDtype::Q1,
10394 rows,
10395 cols,
10396 ..
10397 } => Some((
10398 model,
10399 *idx,
10400 *rows,
10401 *cols,
10402 &[][..],
10403 &[][..],
10404 true,
10405 false,
10406 false,
10407 )),
10408 QTensor::Mapped {
10410 model,
10411 idx,
10412 dtype: cortiq_core::TensorDtype::Q4Tiled,
10413 rows,
10414 cols,
10415 ..
10416 } => Some((
10417 model,
10418 *idx,
10419 *rows,
10420 *cols,
10421 &[][..],
10422 &[][..],
10423 false,
10424 true,
10425 false,
10426 )),
10427 QTensor::Mapped {
10429 model,
10430 idx,
10431 dtype: cortiq_core::TensorDtype::Q4TiledP,
10432 rows,
10433 cols,
10434 ..
10435 } => Some((
10436 model,
10437 *idx,
10438 *rows,
10439 *cols,
10440 &[][..],
10441 &[][..],
10442 false,
10443 true,
10444 false,
10445 )),
10446 QTensor::Mapped {
10450 model,
10451 idx,
10452 dtype: cortiq_core::TensorDtype::Q2TiledP,
10453 rows,
10454 cols,
10455 ..
10456 } => Some((
10457 model,
10458 *idx,
10459 *rows,
10460 *cols,
10461 &[][..],
10462 &[][..],
10463 false,
10464 true,
10465 true,
10466 )),
10467 _ => None,
10468 }
10469}
10470
10471#[cfg(target_os = "macos")]
10477fn metal_moe_graph_parts(m: &MoeFfn, hidden: usize) -> Option<crate::gpu::GpuMoe<'_>> {
10478 if m.router_sigmoid
10479 || m.router_input_norm
10480 || m.expert_bias.is_some()
10481 || m.route_tau.is_some()
10482 || m.mask.is_some()
10483 || m.per_expert_scale.is_some()
10484 || m.experts.is_empty()
10485 || m.top_k == 0
10486 || m.resonance.is_some()
10487 {
10488 return None;
10489 }
10490 let (sh, sg) = match &m.shared {
10493 Some((sh, Some(sg))) => (sh, sg),
10494 _ => return None,
10495 };
10496 let (rf, rr, rc) = m.router.f32_parts()?;
10497 if rr != m.experts.len() || rc != hidden {
10498 return None;
10499 }
10500 let (sf, sr, sc) = sg.f32_parts()?;
10501 if sr * sc != hidden {
10502 return None;
10503 }
10504 let inter = m.experts[0].gate_proj.rows();
10505 let gu_q2 = m.experts[0].gate_proj.mapped_q2tp().is_some();
10508 let trio = |e: &DenseFfn| -> Option<(usize, usize, usize)> {
10509 if e.act != Act::Silu
10510 || e.gate_proj.rows() != inter
10511 || e.gate_proj.cols() != hidden
10512 || e.up_proj.rows() != inter
10513 || e.up_proj.cols() != hidden
10514 || e.down_proj.rows() != hidden
10515 || e.down_proj.cols() != inter
10516 {
10517 return None;
10518 }
10519 let pick = |t: &QTensor| -> Option<usize> {
10520 if gu_q2 {
10521 t.mapped_q2tp().map(|(_, i)| i)
10522 } else {
10523 t.mapped_q4tp().map(|(_, i)| i)
10524 }
10525 };
10526 Some((
10527 pick(&e.gate_proj)?,
10528 pick(&e.up_proj)?,
10529 e.down_proj.mapped_q4tp().map(|(_, i)| i)?,
10530 ))
10531 };
10532 let experts = m.experts.iter().map(trio).collect::<Option<Vec<_>>>()?;
10533 let shared = trio(sh)?;
10534 Some(crate::gpu::GpuMoe {
10535 router: rf,
10536 sgate: sf,
10537 experts,
10538 shared,
10539 n_exp: m.experts.len(),
10540 top_k: m.top_k,
10541 inter,
10542 norm_topk: m.norm_topk_prob,
10543 route_scale: m.routed_scaling,
10544 gu_q2,
10545 })
10546}
10547
10548pub(crate) fn moe_push_job_parts<'a>(
10552 gate: &'a QTensor,
10553 up: &'a QTensor,
10554 down: &'a QTensor,
10555 x: &[f32],
10556 w: f32,
10557 swiglu_limit: f32,
10558 jobs: &mut Vec<crate::gpu::MoeJob<'a>>,
10559 model_ref: &mut Option<std::sync::Arc<cortiq_core::CmfModel>>,
10560) -> Option<()> {
10561 use crate::qtensor::prescale;
10562 let (gm, gi, gr, gc, grs, gcf, gq1, gq4, gq2) = moe_parts(gate)?;
10563 let (_, ui, ur, uc, urs, ucf, uq1, uq4, uq2) = moe_parts(up)?;
10564 let (_, di, dr, dc, drs, dcf, dq1, dq4, dq2) = moe_parts(down)?;
10565 if gq1 != uq1 || uq1 != dq1 || gq4 != uq4 || uq4 != dq4 || gq2 != uq2 {
10566 return None; }
10568 if gq2 && (dq2 || !dq4 || down.mapped_q4tp().is_none()) {
10571 return None;
10572 }
10573 if !gq2 && dq2 {
10574 return None;
10575 }
10576 model_ref.get_or_insert_with(|| gm.clone());
10577 let dt = |cf: &[f32]| {
10578 if cf.is_empty() {
10579 cortiq_core::TensorDtype::Q8Row
10580 } else {
10581 cortiq_core::TensorDtype::Q8_2f
10582 }
10583 };
10584 jobs.push(crate::gpu::MoeJob {
10585 gate: (gi, gr, gc, grs),
10586 up: (ui, ur, uc, urs),
10587 down: (di, dr, dc, drs),
10588 xs_gate: prescale(x, gcf, dt(gcf)).into_owned(),
10589 xs_up: prescale(x, ucf, dt(ucf)).into_owned(),
10590 down_col: dcf,
10591 w,
10592 q1: gq1,
10593 q4t: gq4 && !gq2 && gate.mapped_q4tp().is_none(),
10594 q4tp: gq4 && (gq2 || gate.mapped_q4tp().is_some()),
10595 gu_q2: gq2,
10596 swiglu_limit,
10597 });
10598 Some(())
10599}
10600
10601fn moe_push_job<'a>(
10603 d: &'a DenseFfn,
10604 x: &[f32],
10605 w: f32,
10606 jobs: &mut Vec<crate::gpu::MoeJob<'a>>,
10607 model_ref: &mut Option<std::sync::Arc<cortiq_core::CmfModel>>,
10608) -> Option<()> {
10609 use crate::qtensor::prescale;
10610 if d.act != Act::Silu {
10611 return None; }
10613 let (gm, gi, gr, gc, grs, gcf, gq1, gq4, gq2) = moe_parts(&d.gate_proj)?;
10614 let (_, ui, ur, uc, urs, ucf, uq1, uq4, uq2) = moe_parts(&d.up_proj)?;
10615 let (_, di, dr, dc, drs, dcf, dq1, dq4, dq2) = moe_parts(&d.down_proj)?;
10616 if gq1 != uq1 || uq1 != dq1 || gq4 != uq4 || uq4 != dq4 || gq2 != uq2 {
10617 return None; }
10619 if gq2 && (dq2 || !dq4 || d.down_proj.mapped_q4tp().is_none()) {
10620 return None;
10621 }
10622 if !gq2 && dq2 {
10623 return None;
10624 }
10625 model_ref.get_or_insert_with(|| gm.clone());
10626 let gdt = if gcf.is_empty() {
10627 cortiq_core::TensorDtype::Q8Row
10628 } else {
10629 cortiq_core::TensorDtype::Q8_2f
10630 };
10631 let udt = if ucf.is_empty() {
10632 cortiq_core::TensorDtype::Q8Row
10633 } else {
10634 cortiq_core::TensorDtype::Q8_2f
10635 };
10636 jobs.push(crate::gpu::MoeJob {
10637 gate: (gi, gr, gc, grs),
10638 up: (ui, ur, uc, urs),
10639 down: (di, dr, dc, drs),
10640 xs_gate: prescale(x, gcf, gdt).into_owned(),
10641 xs_up: prescale(x, ucf, udt).into_owned(),
10642 down_col: dcf,
10643 w,
10644 q1: gq1,
10645 q4t: gq4 && !gq2 && d.gate_proj.mapped_q4tp().is_none(),
10646 q4tp: gq4 && (gq2 || d.gate_proj.mapped_q4tp().is_some()),
10647 gu_q2: gq2,
10648 swiglu_limit: 0.0,
10649 });
10650 Some(())
10651}
10652
10653fn sparse_ffn_quant(
10660 d: &DenseFfn,
10661 x: &[f32],
10662 active: &[u16],
10663 hidden: usize,
10664 pool: Option<&Pool>,
10665) -> Vec<f32> {
10666 let n = active.len();
10667 let inter = d.gate_proj.rows();
10668 let mut act = vec![0.0f32; n];
10669 let need_scratch = !(d.gate_proj.sparse_col_ok() && d.up_proj.sparse_col_ok());
10672 let compute = |ai: usize| -> f32 {
10673 let idx = active[ai] as usize;
10674 if idx >= inter {
10675 return 0.0; }
10677 let mut s = if need_scratch {
10678 vec![0.0f32; hidden]
10679 } else {
10680 Vec::new()
10681 };
10682 let gate = d.gate_proj.row_dot(idx, x, &mut s);
10683 let up = d.up_proj.row_dot(idx, x, &mut s);
10684 d.act.combine(gate, up)
10685 };
10686 match pool {
10687 Some(p) if n >= 256 => {
10688 let ptr = SendMut(act.as_mut_ptr());
10689 p.run(&|widx, nw| {
10690 let chunk = n.div_ceil(nw);
10691 let (s, e) = (widx * chunk, ((widx + 1) * chunk).min(n));
10692 for ai in s..e {
10693 unsafe { *ptr.at(ai) = compute(ai) };
10694 }
10695 });
10696 }
10697 _ => {
10698 for (ai, a) in act.iter_mut().enumerate() {
10699 *a = compute(ai);
10700 }
10701 }
10702 }
10703 let mut out = vec![0.0f32; hidden];
10705 for (ai, &idx) in active.iter().enumerate() {
10706 let w = act[ai];
10707 if w.abs() >= 1e-12 && (idx as usize) < inter {
10708 d.down_proj.add_col_scaled(idx as usize, w, &mut out);
10709 }
10710 }
10711 out
10712}
10713
10714#[doc(hidden)]
10716pub fn sparse_ffn_quant_for_test(
10717 d: &DenseFfn,
10718 x: &[f32],
10719 active: &[u16],
10720 hidden: usize,
10721) -> Vec<f32> {
10722 sparse_ffn_quant(d, x, active, hidden, None)
10723}
10724
10725fn dequant_dense_f32(d: &DenseFfn) -> (Vec<f32>, Vec<f32>, Vec<f32>) {
10729 let deq = |t: &QTensor| -> Vec<f32> {
10730 let (rows, cols) = (t.rows(), t.cols());
10731 let mut out = vec![0.0f32; rows * cols];
10732 for r in 0..rows {
10733 t.row_f32(r, &mut out[r * cols..(r + 1) * cols]);
10734 }
10735 out
10736 };
10737 (deq(&d.gate_proj), deq(&d.up_proj), deq(&d.down_proj))
10738}
10739
10740struct SendMut(*mut f32);
10742unsafe impl Send for SendMut {}
10743unsafe impl Sync for SendMut {}
10744impl SendMut {
10745 #[inline]
10746 #[allow(clippy::mut_from_ref)]
10749 unsafe fn at(&self, i: usize) -> &mut f32 {
10750 unsafe { &mut *self.0.add(i) }
10751 }
10752}
10753
10754fn moe_route(logits: &[f32], m: &MoeFfn, allowed: Option<&[bool]>) -> (Vec<usize>, Vec<f32>, f32) {
10764 let ne = logits.len();
10765 let p: Vec<f32> = if m.router_sigmoid {
10766 logits.iter().map(|&l| 1.0 / (1.0 + (-l).exp())).collect()
10767 } else {
10768 let mx = logits.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
10769 let mut e: Vec<f32> = logits.iter().map(|&l| (l - mx).exp()).collect();
10770 let s: f32 = e.iter().sum();
10771 for v in &mut e {
10772 *v /= s;
10773 }
10774 e
10775 };
10776 let admit = |e: usize| {
10782 m.mask.as_ref().is_none_or(|mk| mk[e])
10783 && allowed.is_none_or(|a| a.get(e).copied().unwrap_or(false))
10784 };
10785 let mut idx: Vec<usize> = (0..ne).filter(|&e| admit(e)).collect();
10786 match &m.expert_bias {
10788 Some(b) => idx.sort_unstable_by(|&x, &y| {
10789 (p[y] + b[y])
10790 .partial_cmp(&(p[x] + b[x]))
10791 .unwrap()
10792 .then(x.cmp(&y))
10793 }),
10794 None => idx.sort_unstable_by(|&x, &y| p[y].partial_cmp(&p[x]).unwrap().then(x.cmp(&y))),
10795 }
10796 idx.truncate(m.top_k);
10797 if let Some(tau) = m.route_tau {
10801 let total: f32 = idx.iter().map(|&e| p[e]).sum();
10802 if total > 0.0 {
10803 let mut acc = 0.0f32;
10804 let mut keep = idx.len();
10805 for (i, &e) in idx.iter().enumerate() {
10806 acc += p[e];
10807 if acc >= tau * total {
10808 keep = i + 1;
10809 break;
10810 }
10811 }
10812 idx.truncate(keep);
10813 }
10814 }
10815 let wsum: f32 = if m.norm_topk_prob {
10816 let s: f32 = idx.iter().map(|&e| p[e]).sum();
10817 (if m.router_sigmoid { s + 1e-6 } else { s }) / m.routed_scaling
10820 } else {
10821 1.0 / m.routed_scaling
10822 };
10823 (idx, p, wsum)
10824}
10825
10826fn moe_ffn(m: &MoeFfn, x: &[f32], pool: Option<&Pool>, allowed: Option<&[bool]>) -> Vec<f32> {
10829 accumulate_act(m, x, 1);
10830 let ne = m.experts.len();
10831 let mut logits = vec![0.0f32; ne];
10832 match &m.resonance {
10833 Some(r) => r.scores(x, &mut logits),
10834 None => m.router.matvec(x, &mut logits, pool),
10835 }
10836 let (idx, p, wsum) = moe_route(&logits, m, allowed);
10837 {
10838 let mut st = m.stats.borrow_mut();
10839 if st.len() < ne {
10840 st.resize(ne, 0);
10841 }
10842 for &e in &idx {
10843 st[e] += 1;
10844 }
10845 }
10846 if crate::gpu::enabled_here() {
10851 match crate::gpu::probe_arm(crate::gpu::OpClass::Ffn) {
10852 crate::gpu::ProbeArm::Gpu => {
10853 let t0 = std::time::Instant::now();
10854 if let Some(out) = moe_ffn_gpu(m, x, &idx, &p, wsum, pool) {
10855 crate::gpu::probe_record(crate::gpu::OpClass::Ffn, true, t0.elapsed());
10856 return out;
10857 }
10858 }
10859 crate::gpu::ProbeArm::CpuTimed => {
10860 let t0 = std::time::Instant::now();
10861 let out = crate::gpu::cpu_scope(|| moe_ffn_cpu(m, x, &idx, &p, wsum, pool));
10862 crate::gpu::probe_record(crate::gpu::OpClass::Ffn, false, t0.elapsed());
10863 return out;
10864 }
10865 crate::gpu::ProbeArm::Cpu => {
10866 return crate::gpu::cpu_scope(|| moe_ffn_cpu(m, x, &idx, &p, wsum, pool));
10867 }
10868 }
10869 }
10870 moe_ffn_cpu(m, x, &idx, &p, wsum, pool)
10871}
10872
10873fn graph_note(built: bool) {
10877 use std::sync::atomic::{AtomicBool, Ordering};
10878 if built {
10879 GRAPH_TOK_OK.fetch_add(1, Ordering::Relaxed);
10880 } else {
10881 GRAPH_TOK_MISS.fetch_add(1, Ordering::Relaxed);
10882 }
10883 static SAID: AtomicBool = AtomicBool::new(false);
10884 if !SAID.swap(true, Ordering::Relaxed) {
10885 if built {
10886 tracing::info!("wgpu whole-token graph: ACTIVE");
10887 } else {
10888 tracing::warn!("wgpu whole-token graph refused — per-op path");
10889 }
10890 }
10891}
10892
10893pub static GRAPH_TOK_OK: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
10897pub static GRAPH_TOK_MISS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
10898
10899fn moe_batch_enabled() -> bool {
10902 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
10903 *ON.get_or_init(|| std::env::var("CMF_MOE_BATCH").as_deref() != Ok("0"))
10904}
10905
10906fn moe_ffn_cpu_batched(
10912 m: &MoeFfn,
10913 x: &[f32],
10914 idx: &[usize],
10915 p: &[f32],
10916 wsum: f32,
10917 pool: Option<&Pool>,
10918) -> Option<Vec<f32>> {
10919 if idx.is_empty() || !moe_batch_enabled() {
10920 return None;
10921 }
10922 if FFN_PROBE.with(|pr| pr.borrow().is_some()) {
10926 return None;
10927 }
10928 let n = idx.len() + usize::from(m.shared.is_some());
10929 let mut pairs = Vec::with_capacity(n);
10930 let mut downs = Vec::with_capacity(n);
10931 let mut ws = Vec::with_capacity(n);
10932 for &e in idx {
10933 let d = &m.experts[e];
10934 if d.act != Act::Silu {
10935 return None;
10936 }
10937 pairs.push((&d.gate_proj, &d.up_proj));
10938 downs.push(&d.down_proj);
10939 ws.push(p[e] / wsum * m.per_expert_scale.as_ref().map_or(1.0, |v| v[e]));
10940 }
10941 if let Some((se, gate)) = &m.shared {
10944 if se.act != Act::Silu {
10945 return None;
10946 }
10947 let g = gate.as_ref().map_or(1.0, |gate| {
10948 let mut gl = [0.0f32; 1];
10949 gate.matvec(x, &mut gl, pool);
10950 1.0 / (1.0 + (-gl[0]).exp())
10951 });
10952 pairs.push((&se.gate_proj, &se.up_proj));
10953 downs.push(&se.down_proj);
10954 ws.push(g);
10955 }
10956 let inter = pairs[0].0.rows();
10957 let mut gs: Vec<Vec<f32>> = (0..pairs.len()).map(|_| vec![0f32; inter]).collect();
10958 if !QTensor::moe_gate_up_many(&pairs, x, &mut gs, pool) {
10959 return None;
10960 }
10961 let mut out = attention::take_buf(x.len());
10962 if !QTensor::moe_down_many(&downs, &gs, &ws, &mut out, pool) {
10963 attention::recycle_buf(&mut out);
10964 return None;
10965 }
10966 Some(out)
10967}
10968
10969fn moe_ffn_cpu(
10971 m: &MoeFfn,
10972 x: &[f32],
10973 idx: &[usize],
10974 p: &[f32],
10975 wsum: f32,
10976 pool: Option<&Pool>,
10977) -> Vec<f32> {
10978 if let Some(out) = moe_ffn_cpu_batched(m, x, idx, p, wsum, pool) {
10979 return out;
10980 }
10981 let mut out = attention::take_buf(x.len());
10982 for &e in idx {
10983 let mut eo = dense_ffn(&m.experts[e], x, pool);
10984 let w = p[e] / wsum * m.per_expert_scale.as_ref().map_or(1.0, |v| v[e]);
10985 for i in 0..out.len() {
10986 out[i] += w * eo[i];
10987 }
10988 attention::recycle_buf(&mut eo);
10989 }
10990 if let Some((se, gate)) = &m.shared {
10991 let mut so = dense_ffn(se, x, pool);
10992 let g = gate.as_ref().map_or(1.0, |gate| {
10993 let mut gl = [0.0f32; 1];
10994 gate.matvec(x, &mut gl, pool);
10995 1.0 / (1.0 + (-gl[0]).exp())
10996 });
10997 for i in 0..out.len() {
10998 out[i] += g * so[i];
10999 }
11000 attention::recycle_buf(&mut so);
11001 }
11002 out
11003}
11004
11005#[allow(clippy::too_many_arguments)]
11013fn mla_attention(
11014 w: &MlaWeights,
11015 normed: &[f32],
11016 cache: &mut crate::kv_cache::LayerKvCache,
11017 position: usize,
11018 inv_freq: &[f32],
11019 rope_scale: f32,
11020 eps: f64,
11021 pool: Option<&Pool>,
11022) -> Vec<f32> {
11023 let (nh, dr, dn, dv, lora) = (w.nh, w.qk_rope, w.qk_nope, w.v_dim, w.lora);
11024 let hd = dr + dn;
11025 let mut q = vec![0.0f32; nh * hd];
11026 match (&w.q_a, &w.q_a_norm) {
11027 (Some(qa), Some(qn)) => {
11028 let mut t = vec![0.0f32; qa.rows()];
11029 qa.matvec(normed, &mut t, pool);
11030 let tn = inference::rms_norm(&t, qn, eps, NormStyle::Qwen);
11031 w.q_proj.matvec(&tn, &mut q, pool);
11032 }
11033 _ => w.q_proj.matvec(normed, &mut q, pool),
11034 }
11035 let mut ca = vec![0.0f32; lora + dr];
11036 w.kv_a.matvec(normed, &mut ca, pool);
11037 let (c_lat, k_rope) = ca.split_at_mut(lora);
11038 let latn = inference::rms_norm(c_lat, &w.kv_a_norm, eps, NormStyle::Qwen);
11039 let mut kvb = vec![0.0f32; nh * (dn + dv)];
11040 w.kv_b.matvec(&latn, &mut kvb, pool);
11041 if !w.nope {
11042 attention::rope_rotate_scaled(k_rope, position, inv_freq, rope_scale);
11043 }
11044 for h in 0..nh {
11045 if !w.nope {
11046 attention::rope_rotate_scaled(
11047 &mut q[h * hd..h * hd + dr],
11048 position,
11049 inv_freq,
11050 rope_scale,
11051 );
11052 }
11053 }
11054 let mut k = vec![0.0f32; nh * hd];
11055 let mut v = vec![0.0f32; nh * hd];
11056 for h in 0..nh {
11057 k[h * hd..h * hd + dr].copy_from_slice(k_rope);
11058 k[h * hd + dr..(h + 1) * hd].copy_from_slice(&kvb[h * (dn + dv)..h * (dn + dv) + dn]);
11059 v[h * hd..h * hd + dv].copy_from_slice(&kvb[h * (dn + dv) + dn..(h + 1) * (dn + dv)]);
11060 }
11061 cache.append(&k, &v, &vec![true; nh]);
11062 let (ao, mut imp) = attention::attend_all_heads(&q, cache, nh, 1, hd, w.scale, None, 0.0);
11063 attention::recycle_buf(&mut imp);
11064 let mut ov = vec![0.0f32; nh * dv];
11065 for h in 0..nh {
11066 ov[h * dv..(h + 1) * dv].copy_from_slice(&ao[h * hd..h * hd + dv]);
11067 }
11068 let mut out = vec![0.0f32; w.o_proj.rows()];
11069 w.o_proj.matvec(&ov, &mut out, pool);
11070 out
11071}
11072
11073fn dense_moe_ffn(
11080 dm: &DenseMoeFfn,
11081 x_normed: &[f32],
11082 h_raw: &[f32],
11083 eps: f64,
11084 norm_style: NormStyle,
11085 pool: Option<&Pool>,
11086) -> Vec<f32> {
11087 let mut d = dense_ffn(&dm.dense, x_normed, pool);
11088 d = inference::rms_norm(&d, &dm.post_norm_1, eps, norm_style);
11089 let m = &dm.moe;
11090 let ne = m.experts.len();
11091 let mut logits = vec![0.0f32; ne];
11092 if m.router_input_norm {
11093 let ss: f32 = h_raw.iter().map(|v| v * v).sum::<f32>() / h_raw.len() as f32;
11094 let inv = 1.0 / (ss + eps as f32).sqrt();
11095 let xr: Vec<f32> = h_raw.iter().map(|v| v * inv).collect();
11096 m.router.matvec(&xr, &mut logits, pool);
11097 } else {
11098 m.router.matvec(h_raw, &mut logits, pool);
11099 }
11100 let (idx, p, wsum) = moe_route(&logits, m, None);
11101 {
11102 let mut st = m.stats.borrow_mut();
11103 if st.len() < ne {
11104 st.resize(ne, 0);
11105 }
11106 for &e in &idx {
11107 st[e] += 1;
11108 }
11109 }
11110 let x2 = inference::rms_norm(h_raw, &dm.pre_norm_2, eps, norm_style);
11111 let mo = moe_ffn_cpu(m, &x2, &idx, &p, wsum, pool);
11112 let mo = inference::rms_norm(&mo, &dm.post_norm_2, eps, norm_style);
11113 for (di, mi) in d.iter_mut().zip(&mo) {
11114 *di += mi;
11115 }
11116 d
11117}
11118
11119fn moe_gpu_refused(why: &'static str) {
11126 use std::sync::atomic::{AtomicBool, Ordering};
11127 static SAID: AtomicBool = AtomicBool::new(false);
11128 if !SAID.swap(true, Ordering::Relaxed) {
11129 tracing::warn!("MoE GPU block refused ({why}) — experts run on the CPU");
11130 }
11131}
11132
11133fn moe_ffn_gpu(
11134 m: &MoeFfn,
11135 x: &[f32],
11136 idx: &[usize],
11137 p: &[f32],
11138 wsum: f32,
11139 pool: Option<&Pool>,
11140) -> Option<Vec<f32>> {
11141 use crate::gpu::MoeJob;
11142
11143 let mut jobs: Vec<MoeJob> = Vec::with_capacity(idx.len() + 1);
11144 let mut model_ref = None;
11145 for &e in idx {
11146 if moe_push_job(&m.experts[e], x, p[e] / wsum, &mut jobs, &mut model_ref).is_none() {
11147 moe_gpu_refused("push_job(expert)");
11148 return None;
11149 }
11150 }
11151 if let Some((se, gate)) = &m.shared {
11152 let g = gate.as_ref().map_or(1.0, |gate| {
11153 let mut gl = [0.0f32; 1];
11154 gate.matvec(x, &mut gl, pool);
11155 1.0 / (1.0 + (-gl[0]).exp())
11156 });
11157 if moe_push_job(se, x, g, &mut jobs, &mut model_ref).is_none() {
11158 moe_gpu_refused("push_job(shared)");
11159 return None;
11160 }
11161 }
11162 let Some(model) = model_ref else {
11163 moe_gpu_refused("no model_ref");
11164 return None;
11165 };
11166 let hidden = jobs[0].down.1;
11167 let mut out = vec![0.0f32; hidden];
11168 if crate::gpu::moe_block(&model, &jobs, &mut out) {
11169 Some(out)
11170 } else {
11171 moe_gpu_refused("gpu::moe_block");
11172 None
11173 }
11174}
11175
11176fn ffn_forward(
11178 ffn: &FfnKind,
11179 x: &[f32],
11180 pool: Option<&Pool>,
11181 experts_allowed: Option<&[bool]>,
11182) -> Vec<f32> {
11183 match ffn {
11184 FfnKind::Dense(d) if !d.segs.is_empty() => tube_ffn(d, x, 1, pool, None),
11185 FfnKind::Dense(d) => dense_ffn(d, x, pool),
11186 FfnKind::Moe(m) => moe_ffn(m, x, pool, experts_allowed),
11187 FfnKind::DenseMoe(_) => unreachable!("DenseMoe dispatches via dense_moe_ffn"),
11191 }
11192}
11193
11194fn ffn_forward_pair(
11198 ffn: &FfnKind,
11199 x1: &[f32],
11200 x2: &[f32],
11201 pool: Option<&Pool>,
11202 experts_allowed: Option<&[bool]>,
11203) -> (Vec<f32>, Vec<f32>) {
11204 let d = match ffn {
11205 FfnKind::Dense(d) if !d.segs.is_empty() => {
11208 return (
11209 tube_ffn(d, x1, 1, pool, None),
11210 tube_ffn(d, x2, 1, pool, None),
11211 );
11212 }
11213 FfnKind::Dense(d) => d,
11214 FfnKind::Moe(m) => {
11215 return (
11216 moe_ffn(m, x1, pool, experts_allowed),
11217 moe_ffn(m, x2, pool, experts_allowed),
11218 );
11219 }
11220 FfnKind::DenseMoe(_) => unreachable!("DenseMoe dispatches via dense_moe_ffn"),
11221 };
11222 let inter = d.gate_proj.rows();
11223 FFN_SCRATCH.with(|s| {
11224 let mut s = s.borrow_mut();
11225 let [g1, g2, u1, u2] = &mut *s;
11226 g1.resize(inter, 0.0);
11227 g2.resize(inter, 0.0);
11228 u1.resize(inter, 0.0);
11229 u2.resize(inter, 0.0);
11230 QTensor::matvec2_many(
11233 [&d.gate_proj, &d.up_proj],
11234 x1,
11235 x2,
11236 [g1.as_mut_slice(), u1.as_mut_slice()],
11237 [g2.as_mut_slice(), u2.as_mut_slice()],
11238 pool,
11239 );
11240 for i in 0..inter {
11241 g1[i] = d.act.combine(g1[i], u1[i]);
11242 g2[i] = d.act.combine(g2[i], u2[i]);
11243 }
11244 let mut o1 = attention::take_buf(d.down_proj.rows());
11245 let mut o2 = attention::take_buf(d.down_proj.rows());
11246 d.down_proj.matvec2(g1, g2, &mut o1, &mut o2, pool);
11247 (o1, o2)
11248 })
11249}
11250
11251#[cfg(test)]
11252mod tests {
11253
11254 #[test]
11255 fn cancel_flag_stops_generation() {
11256 let mut p = create_test_pipeline(16, 32, 2, 2, 8, 2, 32);
11257 p.cancel.store(true, std::sync::atomic::Ordering::Relaxed);
11260 let r = p.generate_from_ids(&[1, 2, 3], 8, None, None).unwrap();
11261 assert_eq!(r.finish_reason, "cancelled");
11262 assert!(
11263 r.token_ids.is_empty(),
11264 "no tokens after cancel: {:?}",
11265 r.token_ids
11266 );
11267 let r2 = p.generate_from_ids(&[1, 2, 3], 4, None, None).unwrap();
11269 assert_ne!(r2.finish_reason, "cancelled");
11270 }
11271 use super::*;
11272
11273 #[test]
11281 fn dynamic_ffn_equals_the_zeroing_arm() {
11282 let (hidden, inter) = (8usize, 32usize);
11283 let synth = |n: usize, salt: usize| -> Vec<f32> {
11284 (0..n)
11285 .map(|i| (((i * 29 + salt * 13 + 7) % 89) as f32 / 89.0 - 0.5) * 0.6)
11286 .collect()
11287 };
11288 let down = synth(hidden * inter, 3);
11289 let mut down_t = vec![0.0f32; inter * hidden];
11290 for r in 0..hidden {
11291 for c in 0..inter {
11292 down_t[c * hidden + r] = down[r * inter + c];
11293 }
11294 }
11295 let d = DenseFfn {
11296 gate_proj: QTensor::from_f32(synth(inter * hidden, 1), inter, hidden),
11297 up_proj: QTensor::from_f32(synth(inter * hidden, 2), inter, hidden),
11298 down_proj: QTensor::from_f32(down.clone(), hidden, inter),
11299 act: Act::Silu,
11300 down_t: Some(QTensor::from_f32(down_t, inter, hidden)),
11301 segs: Vec::new(),
11302 };
11303 let x = synth(hidden, 11);
11304 let k = 12usize;
11305 let got = dense_ffn_dynamic(&d, &x, None, k).expect("down_t present");
11306 let mut g = vec![0.0f32; inter];
11308 d.gate_proj.matvec(&x, &mut g, None);
11309 let mut u = vec![0.0f32; inter];
11310 d.up_proj.matvec(&x, &mut u, None);
11311 for v in g.iter_mut() {
11312 *v = inference::silu(*v);
11313 }
11314 keep_top_k(&mut g, k);
11315 for i in 0..inter {
11316 g[i] *= u[i];
11317 }
11318 let mut want = vec![0.0f32; hidden];
11319 d.down_proj.matvec(&g, &mut want, None);
11320 for (a, b) in want.iter().zip(&got) {
11321 assert!((a - b).abs() < 1e-5, "dynamic {b} vs reference {a}");
11322 }
11323 }
11324
11325 #[test]
11331 fn tube_ffn_open_equals_dense_and_closed_equals_masked() {
11332 let (hidden, core, tube) = (8usize, 12usize, 8usize);
11333 let inter = core + tube;
11334 let synth = |n: usize, salt: usize| -> Vec<f32> {
11335 (0..n)
11336 .map(|i| (((i * 41 + salt * 17 + 5) % 97) as f32 / 97.0 - 0.5) * 0.5)
11337 .collect()
11338 };
11339 let (g_all, u_all) = (synth(inter * hidden, 1), synth(inter * hidden, 2));
11340 let d_all = synth(hidden * inter, 3);
11341 let dense = DenseFfn {
11343 gate_proj: QTensor::from_f32(g_all.clone(), inter, hidden),
11344 up_proj: QTensor::from_f32(u_all.clone(), inter, hidden),
11345 down_proj: QTensor::from_f32(d_all.clone(), hidden, inter),
11346 act: Act::Silu,
11347 down_t: None,
11348 segs: Vec::new(),
11349 };
11350 let rows = |v: &[f32], a: usize, b: usize| -> Vec<f32> {
11351 v[a * hidden..b * hidden].to_vec()
11352 };
11353 let cols = |v: &[f32], a: usize, b: usize| -> Vec<f32> {
11354 let mut o = Vec::with_capacity(hidden * (b - a));
11355 for r in 0..hidden {
11356 o.extend_from_slice(&v[r * inter + a..r * inter + b]);
11357 }
11358 o
11359 };
11360 let tubed = DenseFfn {
11361 down_t: None,
11362 gate_proj: QTensor::from_f32(rows(&g_all, 0, core), core, hidden),
11363 up_proj: QTensor::from_f32(rows(&u_all, 0, core), core, hidden),
11364 down_proj: QTensor::from_f32(cols(&d_all, 0, core), hidden, core),
11365 act: Act::Silu,
11366 segs: vec![FfnSeg {
11367 gate: QTensor::from_f32(rows(&g_all, core, inter), tube, hidden),
11368 up: QTensor::from_f32(rows(&u_all, core, inter), tube, hidden),
11369 down: QTensor::from_f32(cols(&d_all, core, inter), hidden, tube),
11370 start: core,
11371 width: tube,
11372 }],
11373 };
11374 let x = synth(hidden, 7);
11375 let want = dense_ffn(&dense, &x, None);
11376 let got = tube_ffn(&tubed, &x, 1, None, None);
11377 for (a, b) in want.iter().zip(&got) {
11378 assert!((a - b).abs() < 1e-5, "open tube: {a} vs {b}");
11379 }
11380 let mut bits = vec![0u8; inter.div_ceil(8)];
11382 for n in 0..core {
11383 bits[n / 8] |= 1 << (n % 8);
11384 }
11385 let closed = tube_ffn(&tubed, &x, 1, None, Some(&bits));
11386 let masked = dense_ffn_masked(&dense, &x, None, &bits);
11387 for (a, b) in masked.iter().zip(&closed) {
11388 assert!((a - b).abs() < 1e-5, "closed tube: {a} vs {b}");
11389 }
11390 let batch = tube_ffn(&tubed, &x, 1, None, Some(&bits));
11392 for (a, b) in closed.iter().zip(&batch) {
11393 assert_eq!(a, b, "batch arm disagrees with decode arm");
11394 }
11395 }
11396
11397 #[test]
11399 fn sparse_ffn_quant_equals_dense_with_inactive_zeroed() {
11400 let (hidden, inter) = (16usize, 40usize);
11401 let synth = |n: usize, salt: usize| -> Vec<f32> {
11402 (0..n)
11403 .map(|i| (((i * 37 + salt * 11 + 3) % 101) as f32 / 101.0 - 0.5) * 0.4)
11404 .collect()
11405 };
11406 let d = DenseFfn {
11407 gate_proj: QTensor::from_f32(synth(inter * hidden, 1), inter, hidden),
11408 up_proj: QTensor::from_f32(synth(inter * hidden, 2), inter, hidden),
11409 down_proj: QTensor::from_f32(synth(hidden * inter, 3), hidden, inter),
11410 act: Act::Silu,
11411 down_t: None,
11412 segs: Vec::new(),
11413 };
11414 let x = synth(hidden, 9);
11415 let active: Vec<u16> = (0..inter as u16).filter(|i| i % 3 == 0).collect();
11417
11418 let sparse = sparse_ffn_quant(&d, &x, &active, hidden, None);
11419
11420 let mut g = vec![0.0f32; inter];
11422 d.gate_proj.matvec(&x, &mut g, None);
11423 let mut u = vec![0.0f32; inter];
11424 d.up_proj.matvec(&x, &mut u, None);
11425 let act_set: std::collections::HashSet<u16> = active.iter().copied().collect();
11426 for i in 0..inter {
11427 g[i] = if act_set.contains(&(i as u16)) {
11428 inference::silu(g[i]) * u[i]
11429 } else {
11430 0.0
11431 };
11432 }
11433 let mut reference = vec![0.0f32; hidden];
11434 d.down_proj.matvec(&g, &mut reference, None);
11435
11436 let max_d = sparse
11437 .iter()
11438 .zip(&reference)
11439 .map(|(a, b)| (a - b).abs())
11440 .fold(0.0f32, f32::max);
11441 assert!(max_d < 1e-5, "sparse != dense-zeroed: max|Δ| = {max_d}");
11442 }
11443
11444 fn attach_test_mtp(p: &mut Pipeline) {
11446 let (h, inter, heads, kv, hd) = (
11447 p.hidden_size,
11448 p.intermediate_size,
11449 p.num_heads,
11450 p.num_kv_heads,
11451 p.head_dim,
11452 );
11453 let synth = |n: usize, salt: usize| -> Vec<f32> {
11454 (0..n)
11455 .map(|i| (((i * 29 + salt * 23 + 5) % 101) as f32 / 101.0 - 0.5) * 0.2)
11456 .collect()
11457 };
11458 let qt = |rows: usize, cols: usize, salt: usize| -> QTensor {
11459 QTensor::from_f32(synth(rows * cols, salt), rows, cols)
11460 };
11461 p.mtp = Some(MtpModule {
11462 enorm: vec![1.0; h],
11463 hnorm: vec![1.0; h],
11464 eh_proj: qt(h, 2 * h, 301),
11465 layer: LayerWeights {
11466 input_norm: vec![1.0; h],
11467 post_norm: vec![1.0; h],
11468 attn_out_norm: None,
11469 ffn_out_norm: None,
11470 layer_scale: None,
11471 ffn: FfnKind::Dense(DenseFfn {
11472 gate_proj: qt(inter, h, 315),
11473 up_proj: qt(inter, h, 316),
11474 down_proj: qt(h, inter, 317),
11475 act: Act::Silu,
11476 down_t: None,
11477 segs: Vec::new(),
11478 }),
11479 attn: AttnKind::Full {
11480 bias: None,
11481 wq: qt(heads * hd, h, 311),
11482 wk: qt(kv * hd, h, 312),
11483 wv: qt(kv * hd, h, 313),
11484 wo: qt(h, heads * hd, 314),
11485 q_norm: None,
11486 k_norm: None,
11487 output_gate: false,
11488 softplus_gate: None,
11489 },
11490 },
11491 final_norm: vec![1.0; h],
11492 kv: crate::kv_cache::LayerKvCache::new(kv, hd),
11493 });
11494 }
11495
11496 #[test]
11497 fn speculative_equals_vanilla_greedy() {
11498 unsafe { std::env::set_var("CMF_GPU_WGPU_GRAPH", "0") };
11502 let run = |spec: bool| {
11503 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
11504 p.sampler_config.temperature = 0.0;
11505 attach_test_mtp(&mut p);
11506 p.speculative = spec;
11507 let r = p.generate("abcdef", 12, None, None).unwrap();
11508 (r.token_ids, r.mtp_drafted, r.mtp_accepted)
11509 };
11510 let (vanilla, d0, _) = run(false);
11511 let (spec, d1, a1) = run(true);
11512 assert_eq!(d0, 0, "vanilla path must not draft");
11513 assert!(d1 > 0, "speculative path must draft");
11514 assert_eq!(
11515 vanilla, spec,
11516 "speculative must reproduce the exact greedy sequence (accepted {a1}/{d1})"
11517 );
11518 }
11519
11520 #[test]
11521 fn speculative_accepts_constant_oracle() {
11522 unsafe { std::env::set_var("CMF_GPU_WGPU_GRAPH", "0") };
11524 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
11525 p.sampler_config.temperature = 0.0;
11526 p.sampler_config.repetition_penalty = 1.0;
11527 p.weights.lm_head = QTensor::from_f32(vec![0.01; 64 * 8], 64, 8);
11530 attach_test_mtp(&mut p);
11531 p.speculative = true;
11532 let r = p.generate("abcd", 10, None, None).unwrap();
11533 assert!(r.mtp_drafted > 0);
11534 assert_eq!(
11535 r.mtp_accepted, r.mtp_drafted,
11536 "constant logits → every draft accepted"
11537 );
11538 assert!(r.token_ids.windows(2).all(|w| w[0] == w[1]));
11541 }
11542
11543 #[test]
11544 fn empty_prompt_is_an_error_not_a_panic() {
11545 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 260);
11546 let r = p.generate("", 4, None, None);
11547 assert!(r.is_err(), "empty prompt must be a clean error");
11548 }
11549
11550 #[test]
11551 fn every_token_enters_kv_exactly_once() {
11552 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
11553 p.sampler_config.temperature = 0.0;
11555 let r = p.generate("abc", 2, None, None).unwrap();
11556 assert_eq!(r.prompt_tokens, 3);
11557 assert_eq!(
11561 p.kv_cache.seq_len(),
11562 3 + r.tokens_generated - 1,
11563 "each token must be cached exactly once (v1 cached the last prompt token twice)"
11564 );
11565 }
11566
11567 #[test]
11568 fn generation_is_reproducible_with_seed() {
11569 let run = || {
11570 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
11571 p.generate("hello", 8, None, None).unwrap().token_ids
11572 };
11573 assert_eq!(run(), run());
11574 }
11575
11576 #[test]
11577 fn resetting_sampler_restarts_the_seeded_stream() {
11578 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
11579 let config = SamplerConfig {
11580 seed: Some(1234),
11581 ..SamplerConfig::default()
11582 };
11583 p.set_sampler_config(config.clone());
11584 let first = p.generate("hello", 8, None, None).unwrap().token_ids;
11585 p.set_sampler_config(config);
11586 let second = p.generate("hello", 8, None, None).unwrap().token_ids;
11587 assert_eq!(first, second);
11588 }
11589
11590 #[test]
11591 fn eviction_bounds_the_cache() {
11592 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 260);
11593 p.kv_cache.max_seq_len = 6;
11594 p.sampler_config.temperature = 0.0;
11595 let _ = p.generate("abcd", 12, None, None).unwrap();
11596 assert!(
11597 p.kv_cache.seq_len() <= 6 + 1,
11598 "cache must stay bounded by max_seq_len (got {})",
11599 p.kv_cache.seq_len()
11600 );
11601 }
11602
11603 #[test]
11604 fn confidence_matches_tokens_and_is_a_probability() {
11605 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
11606 p.sampler_config.temperature = 0.0;
11607 p.sampler_config.repetition_penalty = 1.0;
11608 let r = p.generate("abcd", 10, None, None).unwrap();
11609 assert_eq!(
11610 r.token_confidence.len(),
11611 r.token_ids.len(),
11612 "one confidence per emitted token"
11613 );
11614 for &c in &r.token_confidence {
11615 assert!((0.0..=1.0).contains(&c), "confidence out of [0,1]: {c}");
11616 }
11617 let logits = [1.0f32, 3.0, 0.5, 3.0];
11619 let p0 = top1_prob_t(&logits, 1, 1.0);
11620 let p1 = top1_prob_t(&logits, 3, 1.0);
11621 assert!((p0 - p1).abs() < 1e-6, "equal logits → equal prob");
11622 assert!(p0 > 0.0 && p0 < 1.0);
11623 let sharp = top1_prob_t(&logits, 1, 1.0);
11625 let soft = top1_prob_t(&logits, 1, 2.0);
11626 assert!(soft < sharp, "higher temperature lowers peak confidence");
11627 }
11628
11629 #[test]
11630 fn trace_is_opt_in_and_parallels_the_output() {
11631 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
11633 p.sampler_config.temperature = 0.0;
11634 p.sampler_config.repetition_penalty = 1.0;
11635 let r = p.generate("abcd", 10, None, None).unwrap();
11636 assert!(r.traces.is_empty(), "trace must be empty unless enabled");
11637
11638 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
11640 p.sampler_config.temperature = 0.0;
11641 p.sampler_config.repetition_penalty = 1.0;
11642 p.set_trace(true);
11643 let r = p.generate("abcd", 10, None, None).unwrap();
11644 assert_eq!(r.traces.len(), r.token_ids.len(), "one trace row per token");
11645 for (i, tr) in r.traces.iter().enumerate() {
11646 assert_eq!(tr.t, i, "trace index is sequential");
11647 assert_eq!(tr.token_id, r.token_ids[i], "trace token_id matches output");
11648 assert_eq!(
11649 tr.confidence, r.token_confidence[i],
11650 "trace confidence matches the confidence channel"
11651 );
11652 assert!(tr.active_skill.is_none() && tr.recon.is_none() && !tr.switched);
11654 }
11655 }
11656
11657 #[test]
11658 fn explain_prefill_logits_match_greedy_first_token() {
11659 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
11663 p.sampler_config.temperature = 0.0;
11664 p.sampler_config.repetition_penalty = 1.0;
11665 let ids = p.tokenizer.encode("abcd");
11666 let logits = p.prefill_next_logits(&ids, None);
11667 let argmax = logits
11668 .iter()
11669 .enumerate()
11670 .max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
11671 .unwrap()
11672 .0 as u32;
11673 let r = p.generate("abcd", 1, None, None).unwrap();
11674 assert_eq!(
11675 argmax, r.token_ids[0],
11676 "explain preview must match greedy emit"
11677 );
11678 }
11679
11680 #[test]
11681 fn laguna_shared_expert_is_unconditionally_added() {
11682 let matrix = |values: Vec<f32>| QTensor::from_f32(values, 2, 2);
11683 let identity = || matrix(vec![1.0, 0.0, 0.0, 1.0]);
11684 let zero_dense = || DenseFfn {
11685 gate_proj: matrix(vec![0.0; 4]),
11686 up_proj: matrix(vec![0.0; 4]),
11687 down_proj: matrix(vec![0.0; 4]),
11688 act: Act::Silu,
11689 down_t: None,
11690 segs: Vec::new(),
11691 };
11692 let shared = DenseFfn {
11693 gate_proj: identity(),
11694 up_proj: identity(),
11695 down_proj: identity(),
11696 act: Act::Silu,
11697 down_t: None,
11698 segs: Vec::new(),
11699 };
11700 let x = [1.0, 2.0];
11701 let expected = dense_ffn(&shared, &x, None);
11702 let moe = MoeFfn {
11703 router: QTensor::from_f32(vec![0.0, 0.0], 1, 2),
11704 experts: vec![zero_dense()],
11705 top_k: 1,
11706 norm_topk_prob: true,
11707 router_sigmoid: true,
11708 expert_bias: None,
11709 routed_scaling: 1.0,
11710 route_tau: None,
11711 shared: Some((shared, None)),
11712 stats: std::cell::RefCell::new(Vec::new()),
11713 act_sq: std::cell::RefCell::new(Vec::new()),
11714 act_rows: std::cell::RefCell::new(Vec::new()),
11715 mask: None,
11716 per_expert_scale: None,
11717 router_input_norm: false,
11718 resonance: None,
11719 };
11720 let actual = moe_ffn_cpu(&moe, &x, &[0], &[0.0], 1.0, None);
11721 for (actual, expected) in actual.iter().zip(expected) {
11722 assert!((actual - expected).abs() < 1e-6);
11723 }
11724 }
11725}