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 qwen4_exp: Option<
115 Box<(
116 crate::qwen4_exp::Globals,
117 Vec<crate::qwen4_exp::Layer>,
118 crate::qwen4_exp::Cfg,
119 crate::qwen4_exp::State,
120 )>,
121 >,
122 pub dsv4_mtp: Vec<crate::dsv4::Dsv4Mtp>,
126 pub dspark: Option<crate::dsv4::DsparkState>,
128 pub dspark_pending: Vec<(usize, Vec<u32>, bool, usize)>,
131 pub dspark_hist: Vec<usize>,
133 pub dspark_real: Vec<u32>,
137 pub dspark_trunk_picks: Vec<Vec<(usize, Vec<usize>)>>,
141 pub dspark_exp: Vec<(usize, usize, usize, usize)>,
143 pub dspark_draft_ns: u128,
147 pub short_conv_cfg: Option<ShortConvCfg>,
150 pub mtp: Option<MtpModule>,
152 pub speculative: bool,
154 rng: SplitMix64,
155 sampler_scratch: SamplerScratch,
156 spec_forced: Option<u32>,
162 spec_q: Vec<Vec<f32>>,
163 spec_p: Vec<f32>,
164 spec_res: Vec<f32>,
165 spec_qs: Vec<sampler::Sparse>,
167 spec_ps: sampler::Sparse,
168 spec_ress: sampler::Sparse,
169 mtp_graph_mode: Option<bool>,
176 #[cfg(target_os = "macos")]
179 metal_verify: Option<MetalVerifyPending>,
180 pub(crate) inv_freq: std::sync::Arc<Vec<f32>>,
184 ws: ForwardScratch,
188 pool: Option<std::sync::Arc<Pool>>,
190 pub(crate) model: Option<std::sync::Arc<cortiq_core::CmfModel>>,
194 pub(crate) dyn_force_f32: bool,
196 pub(crate) dyn_skill_layers: Vec<Option<Vec<usize>>>,
201 pub(crate) dyn_active: Option<usize>,
207 pub(crate) dyn_blend_loaded: bool,
211 pub(crate) dyn_phi_layer: Option<usize>,
214 dyn_phi_ema: Vec<f32>,
216 dyn_phi_seen: usize,
217 pub dyn_router: Option<crate::swarm::DynRouter>,
220 o1_cfg: Option<crate::nystrom::O1Cfg>,
223 o1_epoch: u64,
226 o1_flags: Vec<bool>,
228 trace: bool,
231 calib_temp: f32,
234 #[cfg_attr(not(target_os = "macos"), allow(dead_code))]
236 graph_kv_id: u64,
237 #[cfg_attr(not(target_os = "macos"), allow(dead_code))]
240 graph_want_logits: bool,
241 graph_logits: Option<Vec<f32>>,
244 pub embed_multiplier: f32,
246 pub attn_scale: f32,
249 pub swa: Option<(usize, usize)>,
252 pub sliding_layers: Option<Vec<bool>>,
255 pub inv_freq_local: Option<std::sync::Arc<Vec<f32>>>,
258 pub rotary_dim_local: Option<usize>,
259 pub rope_scale: f32,
260 pub rope_scale_local: f32,
261 pub global_attn: Option<(usize, usize)>,
264 pub inv_freq_global: Option<std::sync::Arc<Vec<f32>>>,
267 pub attn_v_norm: bool,
269 pub final_softcap: Option<f32>,
271 pub head_clusters: Option<std::sync::Arc<Vec<f32>>>,
275 pub attn_softcap: f32,
277 confidence_on: bool,
281}
282
283#[cfg(target_os = "macos")]
284impl Drop for Pipeline {
285 fn drop(&mut self) {
286 crate::gpu::kv_mirror_drop(self.graph_kv_id);
287 }
288}
289
290pub struct PipelineWeights {
295 pub embed_tokens: QTensor,
297 pub layers: Vec<LayerWeights>,
299 pub lm_head: QTensor,
301 pub final_norm: Vec<f32>,
303}
304
305pub struct LayerWeights {
307 pub input_norm: Vec<f32>,
308 pub post_norm: Vec<f32>,
311 pub attn_out_norm: Option<Vec<f32>>,
314 pub layer_scale: Option<f32>,
316 pub ffn_out_norm: Option<Vec<f32>>,
319 pub ffn: FfnKind,
320 pub attn: AttnKind,
321}
322
323#[derive(Clone, Copy, PartialEq, Debug, Default)]
326pub enum Act {
327 #[default]
328 Silu,
329 GeluTanh,
330 Situ {
333 beta: f32,
334 linear_beta: f32,
335 },
336}
337
338impl Act {
339 pub fn from_arch(name: &str) -> Self {
340 if name == "gelu_tanh" {
341 Self::GeluTanh
342 } else {
343 Self::Silu
344 }
345 }
346
347 pub fn from_arch_full(arch: &cortiq_core::ModelArch) -> Self {
349 match arch.hidden_act.as_str() {
350 "situ" => Self::Situ {
351 beta: arch.activation_situ_beta.unwrap_or(1.0) as f32,
352 linear_beta: arch.activation_situ_linear_beta.unwrap_or(0.0) as f32,
353 },
354 other => Self::from_arch(other),
355 }
356 }
357
358 #[inline]
359 pub fn apply(self, x: f32) -> f32 {
360 match self {
361 Self::Silu => inference::silu(x),
362 Self::GeluTanh => inference::gelu_tanh(x),
363 Self::Situ { beta, .. } => beta * (x / beta).tanh() * (1.0 / (1.0 + (-x).exp())),
364 }
365 }
366
367 #[inline]
370 pub fn combine(self, g: f32, u: f32) -> f32 {
371 match self {
372 Self::Situ { linear_beta, .. } if linear_beta > 0.0 => {
373 self.apply(g) * (linear_beta * (u / linear_beta).tanh())
374 }
375 _ => self.apply(g) * u,
376 }
377 }
378}
379
380pub struct DenseFfn {
382 pub gate_proj: QTensor,
383 pub up_proj: QTensor,
384 pub down_proj: QTensor,
385 pub act: Act,
387 pub down_t: Option<QTensor>,
393 pub segs: Vec<FfnSeg>,
400}
401
402pub struct FfnSeg {
407 pub gate: QTensor,
408 pub up: QTensor,
409 pub down: QTensor,
410 pub start: usize,
411 pub width: usize,
412}
413
414pub enum FfnKind {
417 Dense(DenseFfn),
418 Moe(MoeFfn),
422 DenseMoe(Box<DenseMoeFfn>),
429}
430
431pub struct DenseMoeFfn {
433 pub dense: DenseFfn,
434 pub moe: MoeFfn,
435 pub post_norm_1: Vec<f32>,
437 pub pre_norm_2: Vec<f32>,
440 pub post_norm_2: Vec<f32>,
442}
443
444pub struct MoeFfn {
445 pub router: QTensor,
447 pub experts: Vec<DenseFfn>,
448 pub top_k: usize,
449 pub norm_topk_prob: bool,
450 pub router_sigmoid: bool,
453 pub expert_bias: Option<Vec<f32>>,
457 pub routed_scaling: f32,
460 pub route_tau: Option<f32>,
466 pub shared: Option<(DenseFfn, Option<QTensor>)>,
469 pub stats: std::cell::RefCell<Vec<u64>>,
473 pub act_sq: std::cell::RefCell<Vec<f64>>,
480 pub act_rows: std::cell::RefCell<Vec<f32>>,
486 pub mask: Option<Vec<bool>>,
491 pub per_expert_scale: Option<Vec<f32>>,
494 pub router_input_norm: bool,
498 pub resonance: Option<Resonance>,
502}
503
504pub struct Resonance {
506 pub mu: Vec<f32>,
508 pub u: Vec<f32>,
510 pub k: usize,
511 pub bias: Vec<f32>,
513}
514
515impl Resonance {
516 pub fn scores(&self, x: &[f32], out: &mut [f32]) {
518 let h = x.len();
519 let ne = out.len();
520 for e in 0..ne {
521 let mu = &self.mu[e * h..(e + 1) * h];
522 let mut d2 = 0.0f32;
523 for j in 0..h {
524 let d = x[j] - mu[j];
525 d2 += d * d;
526 }
527 let mut proj = 0.0f32;
528 for i in 0..self.k {
529 let u = &self.u[(e * self.k + i) * h..(e * self.k + i + 1) * h];
530 let mut p = 0.0f32;
531 for j in 0..h {
532 p += (x[j] - mu[j]) * u[j];
533 }
534 proj += p * p;
535 }
536 out[e] = self.bias.get(e).copied().unwrap_or(0.0) - (d2 - proj);
537 }
538 }
539}
540
541pub enum AttnKind {
544 Full {
546 wq: QTensor,
547 wk: QTensor,
548 wv: QTensor,
549 wo: QTensor,
550 q_norm: Option<Vec<f32>>,
551 k_norm: Option<Vec<f32>>,
552 output_gate: bool,
553 softplus_gate: Option<(QTensor, bool)>,
557 bias: Option<(Vec<f32>, Vec<f32>, Vec<f32>)>,
559 },
560 Linear(VmfPhaseWeights),
562 LinearGdn(GdnWeights),
564 ShortConv(ShortConvWeights),
567 Mla(Box<MlaWeights>),
575 Kda(Box<crate::linear_core::KdaWeights>),
579}
580
581pub struct MlaWeights {
583 pub q_proj: QTensor,
587 pub q_a: Option<QTensor>,
590 pub q_a_norm: Option<Vec<f32>>,
591 pub kv_a: QTensor,
593 pub kv_a_norm: Vec<f32>,
595 pub kv_b: QTensor,
597 pub o_proj: QTensor,
599 pub nh: usize,
600 pub qk_rope: usize,
601 pub qk_nope: usize,
602 pub v_dim: usize,
603 pub lora: usize,
604 pub scale: f32,
606 pub nope: bool,
608}
609
610pub struct MtpModule {
615 pub enorm: Vec<f32>,
616 pub hnorm: Vec<f32>,
617 pub eh_proj: QTensor,
619 pub layer: LayerWeights,
620 pub final_norm: Vec<f32>,
621 pub kv: crate::kv_cache::LayerKvCache,
622}
623
624#[cfg(target_os = "macos")]
631enum MetalRowsItem<'a> {
632 Gdn {
633 run: Vec<crate::gpu_metal::GdnGpuLayer<'a>>,
634 first: usize,
635 },
636 Attn {
637 l: crate::gpu_metal::AttnGpuLayer<'a>,
638 li: usize,
639 q_norm: Option<&'a [f32]>,
640 k_norm: Option<&'a [f32]>,
641 output_gate: bool,
642 },
643}
644
645#[cfg(target_os = "macos")]
646struct MetalVerifyPending {
647 graph: crate::gpu_metal::VerifyGraph,
648 gdn_layers: Vec<usize>,
649 attn_layers: Vec<(usize, usize)>,
650}
651
652#[derive(Clone, Copy)]
656enum SpecTrial {
657 Spec {
658 t0: std::time::Instant,
659 gen0: usize,
660 rounds: usize,
661 },
662 Plain {
663 t0: std::time::Instant,
664 gen0: usize,
665 },
666 Decided {
667 spec: bool,
668 recheck_at: usize,
669 },
670}
671
672#[derive(Default, Clone, Copy)]
683struct SpecMon {
684 round_ms: f64,
685 tokens: f64,
686 plain_ms: f64,
687 n: u32,
688 fails: u32,
689}
690
691impl SpecMon {
692 fn round(&mut self, dt_ms: f64, produced: usize) {
693 self.n += 1;
694 if self.n == 1 {
695 return; }
697 let a = if self.n == 2 { 1.0 } else { 0.3 };
698 self.round_ms += a * (dt_ms - self.round_ms);
699 self.tokens += a * (produced as f64 - self.tokens);
700 }
701 fn pays(&self) -> bool {
702 self.plain_ms > 0.0 && self.tokens * self.plain_ms > self.round_ms * 1.03
703 }
704}
705
706pub struct GenerateResult {
708 pub text: String,
709 pub token_ids: Vec<u32>,
710 pub prompt_tokens: usize,
711 pub tokens_generated: usize,
712 pub finish_reason: String,
713 pub mtp_drafted: usize,
715 pub mtp_accepted: usize,
716 pub token_confidence: Vec<f32>,
721 pub traces: Vec<TokenTrace>,
724}
725
726#[derive(Clone, Debug)]
731pub struct TokenTrace {
732 pub t: usize,
734 pub token_id: u32,
736 pub confidence: f32,
738 pub active_skill: Option<String>,
740 pub recon: Option<f32>,
744 pub switched: bool,
747}
748
749#[cfg_attr(not(test), allow(dead_code))]
754fn top1_prob_t(logits: &[f32], id: u32, temp: f32) -> f32 {
755 let t = if temp > 1e-3 { temp } else { 1.0 };
756 let max = logits.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
757 let sum: f32 = logits.iter().map(|&v| ((v - max) / t).exp()).sum();
758 if sum > 0.0 {
759 (((logits[id as usize] - max) / t).exp()) / sum
760 } else {
761 0.0
762 }
763}
764
765fn prefill_batched() -> bool {
768 std::env::var("CMF_PREFILL")
769 .map(|v| v != "seq")
770 .unwrap_or(true)
771}
772
773#[derive(Clone, Copy)]
777enum PrefillIn<'a> {
778 Ids(&'a [u32]),
779 Hidden(&'a [f32]),
780}
781
782impl Pipeline {
789 fn can_prefill_batched(&self) -> bool {
790 prefill_batched() && !self.weights.layers.is_empty()
791 }
792}
793
794pub fn prefill_chunk() -> usize {
801 if let Some(n) = std::env::var("CMF_PREFILL_CHUNK")
802 .ok()
803 .and_then(|v| v.parse::<usize>().ok())
804 {
805 return n.max(1);
806 }
807 if cfg!(target_os = "macos") {
808 512
809 } else if cfg!(target_arch = "aarch64") {
810 256
813 } else {
814 48
815 }
816}
817
818pub type TokenCallback = Box<dyn FnMut(&str) -> bool + Send>;
820
821impl Pipeline {
822 #[inline]
826 pub fn phys_layer(&self, virtual_idx: usize) -> usize {
827 virtual_idx % self.physical_layers
828 }
829
830 #[inline]
833 pub fn is_loop_end(&self, virtual_idx: usize) -> bool {
834 self.loop_final_norm && (virtual_idx + 1) % self.physical_layers == 0
835 }
836
837 #[allow(clippy::too_many_arguments)]
839
840 #[cfg(target_os = "macos")]
859 fn graph_prefill_preferred(&self) -> bool {
860 if !crate::gpu::enabled_here()
861 || !crate::gpu::q1_force()
862 || std::env::var("CMF_GPU_BLOCK")
863 .map(|v| v == "0")
864 .unwrap_or(false)
865 || std::env::var("CMF_PREFILL_GRAPH").as_deref() == Ok("0")
868 {
869 return false;
870 }
871 self.weights
872 .layers
873 .iter()
874 .any(|lw| matches!(&lw.attn, AttnKind::LinearGdn(w) if w.in_proj_qkv.is_q1()))
875 }
876
877 #[cfg(not(target_os = "macos"))]
878 fn graph_prefill_preferred(&self) -> bool {
879 let graph_on = crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Prefill);
887 if !graph_on || !crate::gpu::enabled_here() {
888 return false;
889 }
890 if self.o1_active() {
899 return false;
900 }
901 self.weights
902 .layers
903 .iter()
904 .any(|lw| matches!(&lw.attn, AttnKind::LinearGdn(_)))
905 }
906
907 #[cfg(target_os = "macos")]
908 fn q1_graph_gpu(
909 &mut self,
910 start: usize,
911 upto: Option<usize>,
912 position: usize,
913 h: &mut [f32],
914 ) -> usize {
915 let _mt0 = std::time::Instant::now(); use crate::gpu::{AttnGpuLayer, GdnGpuCfg, GdnGpuLayer, GraphDims, MetalFfn, TokenGraph};
917 if self.attn_softcap > 0.0 || !crate::gpu::enabled_here()
919 || !crate::gpu::q1_force()
920 || std::env::var("CMF_GPU_BLOCK")
921 .map(|v| v == "0")
922 .unwrap_or(false)
923 {
924 if std::env::var("CMF_GRAPH_DBG").is_ok() {
925 eprintln!(
926 "block-graph: front gate (softcap={} enabled_here={} q1_force={})",
927 self.attn_softcap > 0.0,
928 crate::gpu::enabled_here(),
929 crate::gpu::q1_force(),
930 );
931 }
932 return start;
933 }
934 if self.swa.is_some()
939 || self.global_attn.is_some()
940 || self.attention_heads_per_layer.is_some()
941 || self.attn_v_norm
942 || (self.attn_scale - 1.0 / (self.head_dim as f32).sqrt()).abs() > 1e-9
943 || self.weights.layers.iter().any(|lw| {
944 lw.attn_out_norm.is_some()
945 || lw.ffn_out_norm.is_some()
946 || lw.layer_scale.is_some()
947 || matches!(&lw.ffn, FfnKind::Dense(d) if d.act != Act::Silu)
948 })
949 {
950 if std::env::var("CMF_GRAPH_DBG").is_ok() {
951 eprintln!(
952 "block-graph: arch ineligible (swa={} gattn={} hpl={} vnorm={} scale_delta={:.2e})",
953 self.swa.is_some(),
954 self.global_attn.is_some(),
955 self.attention_heads_per_layer.is_some(),
956 self.attn_v_norm,
957 (self.attn_scale - 1.0 / (self.head_dim as f32).sqrt()).abs(),
958 );
959 }
960 return start;
961 }
962 let limit = upto
965 .map(|u| u + 1)
966 .unwrap_or(self.num_layers)
967 .min(self.num_layers);
968
969 enum Item<'a> {
970 Gdn {
971 run: Vec<GdnGpuLayer<'a>>,
972 first: usize,
973 },
974 Attn {
975 l: AttnGpuLayer<'a>,
976 li: usize,
977 q_norm: Option<&'a [f32]>,
978 k_norm: Option<&'a [f32]>,
979 output_gate: bool,
980 bias: Option<(&'a [f32], &'a [f32], &'a [f32])>,
981 full_gpu: bool,
984 },
985 }
986
987 let attend_mode = std::env::var("CMF_GPU_ATTEND").unwrap_or_else(|_| "auto".into());
994 let attend_contract = attend_mode != "0"
995 && attend_mode != "off"
996 && self.head_dim % 4 == 0
997 && self.head_dim <= 256
998 && self.rotary_dim >= 2
999 && self.rotary_dim <= self.head_dim
1000 && (self.rotary_dim / 2) % 32 == 0
1001 && self.num_kv_heads > 0
1002 && self.num_heads % self.num_kv_heads == 0;
1003
1004 let mut plan: Vec<Item> = Vec::new();
1005 let mut model_ref: Option<std::sync::Arc<cortiq_core::CmfModel>> = None;
1006 let block_diag = std::env::var("CMF_GRAPH_DBG").is_ok();
1008 let mut scan = start;
1009 while scan < limit {
1010 let lw = &self.weights.layers[self.phys_layer(scan)];
1011 let ffn = match &lw.ffn {
1012 FfnKind::Dense(d) if d.segs.is_empty() => {
1013 let (Some(g), Some(u), Some(dn)) = (
1014 d.gate_proj.q1_parts(),
1015 d.up_proj.q1_parts(),
1016 d.down_proj.q1_parts(),
1017 ) else {
1018 if block_diag {
1019 eprintln!(
1020 "block-graph: L{scan} FFN trio not graph-mappable — run ends"
1021 );
1022 }
1023 break;
1024 };
1025 MetalFfn::Dense {
1026 gate: g,
1027 up: u,
1028 down: dn,
1029 }
1030 }
1031 FfnKind::Moe(m) => {
1032 let Some(moe) = metal_moe_graph_parts(m, self.hidden_size) else {
1033 if block_diag {
1034 eprintln!(
1035 "block-graph: L{scan} MoE outside the graph contract — run ends"
1036 );
1037 }
1038 break;
1039 };
1040 if let QTensor::Mapped { model, .. } = &m.experts[0].gate_proj {
1041 model_ref.get_or_insert_with(|| model.clone());
1042 }
1043 MetalFfn::Moe(moe)
1044 }
1045 _ => {
1046 if block_diag {
1047 eprintln!("block-graph: L{scan} non-graph FFN — run ends");
1048 }
1049 break;
1050 }
1051 };
1052 match &lw.attn {
1053 AttnKind::LinearGdn(w) if self.gdn_cfg.is_some() => {
1054 let parts = (
1055 w.in_proj_qkv.q1_parts(),
1056 w.in_proj_z.q1_parts(),
1057 w.in_proj_a.f32_parts(),
1058 w.in_proj_b.f32_parts(),
1059 w.out_proj.q1_parts(),
1060 );
1061 let (Some(qkv), Some(z), Some(a), Some(b), Some(out)) = parts else {
1062 if block_diag {
1063 eprintln!(
1064 "block-graph: L{scan} GDN parts refused (qkv={} z={} a_f32={} b_f32={} out={})",
1065 w.in_proj_qkv.q1_parts().is_some(),
1066 w.in_proj_z.q1_parts().is_some(),
1067 w.in_proj_a.f32_parts().is_some(),
1068 w.in_proj_b.f32_parts().is_some(),
1069 w.out_proj.q1_parts().is_some(),
1070 );
1071 }
1072 break;
1073 };
1074 if let QTensor::Mapped { model, .. } = &w.in_proj_qkv {
1075 model_ref.get_or_insert_with(|| model.clone());
1076 }
1077 let gl = GdnGpuLayer {
1078 attn_norm: &lw.input_norm,
1079 post_norm: &lw.post_norm,
1080 qkv,
1081 z,
1082 a,
1083 b,
1084 out,
1085 ffn,
1086 conv1d: &w.conv1d,
1087 a_log: &w.a_log,
1088 dt_bias: &w.dt_bias,
1089 gnorm: &w.norm,
1090 };
1091 match plan.last_mut() {
1092 Some(Item::Gdn { run, .. }) => run.push(gl),
1093 _ => plan.push(Item::Gdn {
1094 run: vec![gl],
1095 first: scan,
1096 }),
1097 }
1098 }
1099 AttnKind::Full {
1100 wq,
1101 wk,
1102 wv,
1103 wo,
1104 q_norm,
1105 k_norm,
1106 output_gate,
1107 softplus_gate: None,
1108 bias,
1109 } if !self.kv_cache.layers[scan].o1_sealed()
1110 || std::env::var("CMF_O1_METAL").as_deref() == Ok("1") =>
1115 {
1116 let parts = (wq.q1_parts(), wk.q1_parts(), wv.q1_parts(), wo.q1_parts());
1117 let (Some(pq), Some(pk), Some(pv), Some(po)) = parts else {
1118 break;
1119 };
1120 if let QTensor::Mapped { model, .. } = wq {
1121 model_ref.get_or_insert_with(|| model.clone());
1122 }
1123 let cache = &self.kv_cache.layers[scan];
1124 let o1_metal = cache.o1.is_some()
1128 && std::env::var("CMF_O1_METAL").as_deref() == Ok("1")
1129 && cache.o1_views().is_some();
1130 let full_gpu = attend_contract
1131 && cache.mode == crate::kv_cache::KvMode::F32
1132 && (cache.o1.is_none() || o1_metal)
1133 && bias.is_none()
1134 && pq.1 == self.num_heads * self.head_dim * (1 + *output_gate as usize)
1135 && pk.1 == self.num_kv_heads * self.head_dim
1136 && pv.1 == self.num_kv_heads * self.head_dim
1137 && po.2 == self.num_heads * self.head_dim;
1138 plan.push(Item::Attn {
1139 l: AttnGpuLayer {
1140 attn_norm: &lw.input_norm,
1141 post_norm: &lw.post_norm,
1142 wq: pq,
1143 wk: pk,
1144 wv: pv,
1145 wo: po,
1146 ffn,
1147 },
1148 li: scan,
1149 q_norm: q_norm.as_deref(),
1150 k_norm: k_norm.as_deref(),
1151 output_gate: *output_gate,
1152 bias: bias
1153 .as_ref()
1154 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
1155 full_gpu,
1156 });
1157 }
1158 _ => break,
1159 }
1160 scan += 1;
1161 }
1162 let Some(model) = model_ref else {
1163 if std::env::var("CMF_GRAPH_DBG").is_ok() {
1164 eprintln!("q1-graph: no model ref (start {start}, scanned to {scan})");
1165 }
1166 return start;
1167 };
1168 if plan.is_empty() {
1169 if std::env::var("CMF_GRAPH_DBG").is_ok() {
1170 eprintln!("q1-graph: empty plan at layer {start}");
1171 }
1172 return start;
1173 }
1174 let has_moe = plan.iter().any(|it| match it {
1175 Item::Gdn { run, .. } => run.iter().any(|l| matches!(l.ffn, MetalFfn::Moe(_))),
1176 Item::Attn { l, .. } => matches!(l.ffn, MetalFfn::Moe(_)),
1177 });
1178 let has_gdn = plan.iter().any(|it| matches!(it, Item::Gdn { .. }));
1179 let dev_attend = attend_contract
1180 && (self.head_dim <= 128
1181 || has_moe
1182 || (self.head_dim <= 256 && has_gdn)
1188 || attend_mode == "force"
1189 || attend_mode == "256");
1190 if !dev_attend {
1191 for it in &mut plan {
1192 if let Item::Attn { li, full_gpu, .. } = it {
1193 let keep_o1 = self.kv_cache.layers[*li].o1.is_some()
1196 && std::env::var("CMF_O1_METAL").as_deref() == Ok("1");
1197 if !keep_o1 {
1198 *full_gpu = false;
1199 }
1200 }
1201 }
1202 }
1203 if std::env::var("CMF_GRAPH_DBG").is_ok() {
1204 use std::sync::atomic::{AtomicBool, Ordering};
1205 static SAID: AtomicBool = AtomicBool::new(false);
1206 if !SAID.swap(true, Ordering::Relaxed) {
1207 let fg = plan
1208 .iter()
1209 .filter(|it| matches!(it, Item::Attn { full_gpu: true, .. }))
1210 .count();
1211 let att = plan
1212 .iter()
1213 .filter(|it| matches!(it, Item::Attn { .. }))
1214 .count();
1215 eprintln!(
1216 "q1-graph: plan of {} items from layer {start} to {scan} | dev_attend={dev_attend} full_gpu {fg}/{att} | hd={} rd={} nkv={} nh={}",
1217 plan.len(),
1218 self.head_dim,
1219 self.rotary_dim,
1220 self.num_kv_heads,
1221 self.num_heads,
1222 );
1223 }
1224 }
1225 let dims = GraphDims {
1226 hidden: self.hidden_size,
1227 eps: self.rms_eps as f32,
1228 gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
1229 };
1230 let Some(mut graph) = TokenGraph::new(&model, dims, h) else {
1231 return start;
1232 };
1233 let gcfg = self.gdn_cfg.map(|cfg| GdnGpuCfg {
1234 nv: cfg.num_v_heads,
1235 nk: cfg.num_k_heads,
1236 dk: cfg.key_head_dim,
1237 dv: cfg.value_head_dim,
1238 kk: cfg.conv_kernel,
1239 hidden: self.hidden_size,
1240 inter: self.intermediate_size,
1241 c_dim: cfg.conv_dim(),
1242 eps: cfg.rms_eps as f32,
1243 gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
1244 });
1245 let mut valid = 0usize;
1249 let mut end = start;
1250 crate::gpu::stageprof(1, _mt0.elapsed()); if std::env::var("CMF_PLAN_DUMP").is_ok() {
1252 static ONCE: std::sync::Once = std::sync::Once::new();
1253 ONCE.call_once(|| {
1254 for it in &plan {
1255 match it {
1256 Item::Gdn { first, run } => {
1257 eprintln!("plan: Gdn first={first} len={}", run.len())
1258 }
1259 Item::Attn { li, full_gpu, .. } => {
1260 eprintln!("plan: Attn li={li} full_gpu={full_gpu}")
1261 }
1262 }
1263 }
1264 });
1265 }
1266 for item in &plan {
1267 let ok = match item {
1268 Item::Gdn { run, .. } => gcfg
1269 .as_ref()
1270 .map(|gc| run.iter().all(|l| graph.gdn_ok(l, gc)))
1271 .unwrap_or(false),
1272 Item::Attn { l, .. } => graph.attn_ok(l),
1273 };
1274 if !ok {
1275 if block_diag {
1276 eprintln!(
1277 "block-graph: plan item {} ({}) failed graph preflight",
1278 valid,
1279 match item {
1280 Item::Gdn { run, first } => format!("GDN run L{first}+{}", run.len()),
1281 Item::Attn { li, .. } => format!("Attn L{li}"),
1282 }
1283 );
1284 }
1285 break;
1286 }
1287 valid += 1;
1288 end += match item {
1289 Item::Gdn { run, .. } => run.len(),
1290 Item::Attn { .. } => 1,
1291 };
1292 }
1293 plan.truncate(valid);
1294 if plan.is_empty() {
1295 return start;
1296 }
1297
1298 let inv_freq = self.inv_freq.clone();
1299 let pool = self.pool.clone();
1300 let (nh, nkv, hd, hs, rd, eps) = (
1301 self.num_heads,
1302 self.num_kv_heads,
1303 self.head_dim,
1304 self.hidden_size,
1305 self.rotary_dim,
1306 self.rms_eps,
1307 );
1308 let norm_style = self.norm_style;
1309 let gemma = norm_style == cortiq_core::NormStyle::Gemma;
1310 let want = self.gdn_cfg.map(|c| c.state_len()).unwrap_or(0);
1311 let kv_id = self.graph_kv_id;
1312 let mut pending: Vec<(usize, usize)> = Vec::new();
1315 let mut dev_attn: Vec<usize> = Vec::new();
1318 for item in &plan {
1319 let _xt0 = std::time::Instant::now();
1320 let _xkind: u32 = match item {
1321 Item::Gdn { .. } => 2,
1322 Item::Attn { .. } => 3,
1323 };
1324 if self.loop_final_norm {
1326 let item_start = match item {
1327 Item::Gdn { first, .. } => *first,
1328 Item::Attn { li, .. } => *li,
1329 };
1330 if item_start > start && self.is_loop_end(item_start - 1) {
1331 graph.encode_loop_norm(&self.weights.final_norm);
1332 }
1333 }
1334 match item {
1335 Item::Gdn { run, first } => {
1336 for l in &mut self.kv_cache.layers[*first..*first + run.len()] {
1337 if l.linear_state.len() != want {
1338 l.linear_state = vec![0f32; want];
1339 }
1340 }
1341 let ro: Vec<&[f32]> = self.kv_cache.layers[*first..*first + run.len()]
1342 .iter()
1343 .map(|l| l.linear_state.as_slice())
1344 .collect();
1345 let _ig = std::time::Instant::now();
1346 if !graph.encode_gdn_run(run, &ro, gcfg.as_ref().unwrap()) {
1347 tracing::error!("q1 graph: GDN run refused after validation");
1349 return start;
1350 }
1351 graph.commit_kind = 2;
1354 graph.commit();
1355 crate::gpu::stageprof(0, _ig.elapsed());
1356 pending.push((*first, run.len()));
1357 }
1358 Item::Attn {
1359 l,
1360 li,
1361 q_norm,
1362 k_norm,
1363 output_gate,
1364 bias,
1365 full_gpu,
1366 } => {
1367 let _ia = std::time::Instant::now();
1368 if *full_gpu {
1370 let cache = &self.kv_cache.layers[*li];
1371 let o1p = if cache.o1.is_some() {
1372 match cache.o1_views() {
1373 Some(views) => Some(crate::gpu::O1AttnParams {
1374 views,
1375 epoch: self.o1_epoch,
1376 }),
1377 None => None,
1379 }
1380 } else {
1381 None
1382 };
1383 let o1_layer = cache.o1.is_some();
1384 if o1_layer && o1p.is_none() {
1385 }
1387 let cpu_k: Vec<&[f32]> = (0..nkv).map(|g| cache.head_keys(g)).collect();
1388 let cpu_v: Vec<&[f32]> = (0..nkv).map(|g| cache.head_values(g)).collect();
1389 let cpu_stored = if o1_layer { 0 } else { cpu_k[0].len() / hd };
1390 let p = crate::gpu::AttnDeviceParams {
1391 kv_id,
1392 layer: *li,
1393 nh,
1394 nkv,
1395 hd,
1396 rd,
1397 position,
1398 eps: eps as f32,
1399 gemma,
1400 output_gate: *output_gate,
1401 q_norm: *q_norm,
1402 k_norm: *k_norm,
1403 inv_freq: &inv_freq,
1404 cpu_k,
1405 cpu_v,
1406 cpu_stored,
1407 o1: o1p,
1408 };
1409 let o1_bad = o1_layer && p.o1.is_none();
1410 if !o1_bad && graph.attn_device_ok(l, &p) && graph.encode_attn_device(l, &p)
1411 {
1412 if p.o1.is_none() {
1414 dev_attn.push(*li);
1415 }
1416 graph.commit_kind = 3;
1417 graph.commit();
1418 crate::gpu::stageprof(_xkind, _xt0.elapsed());
1422 continue;
1423 }
1424 }
1426 graph.encode_attn_prefix(l);
1427 graph.sync();
1428 if !pending.is_empty() {
1429 let idxs: Vec<usize> =
1430 pending.drain(..).flat_map(|(f, n)| f..f + n).collect();
1431 let mut outs: Vec<&mut [f32]> = self
1432 .kv_cache
1433 .layers
1434 .iter_mut()
1435 .enumerate()
1436 .filter(|(i, _)| idxs.binary_search(i).is_ok())
1437 .map(|(_, s)| s.linear_state.as_mut_slice())
1438 .collect();
1439 graph.read_states(&mut outs);
1440 }
1441 let mut q_raw = attention::take_buf(l.wq.1);
1442 let mut k = attention::take_buf(l.wk.1);
1443 let mut v = attention::take_buf(l.wv.1);
1444 graph.read_qkv(&mut q_raw, &mut k, &mut v);
1445 let cfg = QwenAttnCfg {
1446 num_heads: nh,
1447 num_kv_heads: nkv,
1448 head_dim: hd,
1449 hidden_size: hs,
1450 position,
1451 inv_freq: &inv_freq,
1452 rotary_dim: rd,
1453 scale: self.attn_scale,
1454 softcap: self.attn_softcap,
1455 window: None,
1456 v_norm: false,
1457 q_norm: *q_norm,
1458 k_norm: *k_norm,
1459 output_gate: *output_gate,
1460 softplus_gate: None,
1461 rope_scale: 1.0,
1462 bias: *bias,
1463 rms_eps: eps,
1464 norm_style,
1465 pool: pool.as_deref(),
1466 };
1467 let oracle = std::env::var("CMF_ATTN_ORACLE").as_deref() == Ok("1")
1470 || std::env::var("CMF_ATTN_DUMP").is_ok();
1471 let _ = full_gpu;
1472 let oracle_in = oracle.then(|| (q_raw.clone(), k.clone(), v.clone()));
1473 let mut ao = attention::qwen_attention_core(
1474 q_raw,
1475 k,
1476 v,
1477 &mut self.kv_cache.layers[*li],
1478 &cfg,
1479 );
1480 if let Ok(dir) = std::env::var("CMF_ATTN_DUMP") {
1484 if let Some((qr0, k0, v0)) = oracle_in.clone() {
1485 let (cq, _cg, _ck, _cv) =
1486 attention::finish_projection_debug(qr0, k0, v0, &cfg, position);
1487 let cache = &self.kv_cache.layers[*li];
1488 let n = cache.head_keys(0).len() / hd;
1489 let mut bytes: Vec<u8> = Vec::new();
1490 for v in [nh as u32, nkv as u32, hd as u32, n as u32, position as u32] {
1491 bytes.extend_from_slice(&v.to_le_bytes());
1492 }
1493 for v in &cq {
1494 bytes.extend_from_slice(&v.to_le_bytes());
1495 }
1496 for g in 0..nkv {
1497 for v in cache.head_keys(g) {
1498 bytes.extend_from_slice(&v.to_le_bytes());
1499 }
1500 }
1501 for g in 0..nkv {
1502 for v in cache.head_values(g) {
1503 bytes.extend_from_slice(&v.to_le_bytes());
1504 }
1505 }
1506 let _ =
1507 std::fs::write(format!("{dir}/L{li}_pos{position}.bin"), &bytes);
1508 }
1509 }
1510 if let Some((qr0, k0, v0)) =
1511 oracle_in.filter(|_| std::env::var("CMF_ATTN_ORACLE").as_deref() == Ok("1"))
1512 {
1513 let (cq, _cg, ck, cv) =
1514 attention::finish_projection_debug(qr0, k0, v0, &cfg, position);
1515 let mut h_now = vec![0f32; hs];
1516 graph.read_h(&mut h_now);
1517 let cache = &self.kv_cache.layers[*li];
1518 let n_after = cache.head_keys(0).len() / hd;
1519 let cpu_k: Vec<&[f32]> = (0..nkv)
1520 .map(|g| &cache.head_keys(g)[..(n_after - 1) * hd])
1521 .collect();
1522 let cpu_v: Vec<&[f32]> = (0..nkv)
1523 .map(|g| &cache.head_values(g)[..(n_after - 1) * hd])
1524 .collect();
1525 let p = crate::gpu::AttnDeviceParams {
1526 kv_id,
1527 layer: *li,
1528 nh,
1529 nkv,
1530 hd,
1531 rd,
1532 position,
1533 eps: eps as f32,
1534 gemma,
1535 output_gate: *output_gate,
1536 q_norm: *q_norm,
1537 k_norm: *k_norm,
1538 inv_freq: &inv_freq,
1539 cpu_k,
1540 cpu_v,
1541 cpu_stored: n_after - 1,
1542 o1: None,
1543 };
1544 if let Some((dq, dk, dv, dao)) = graph.debug_attn_device(l, &p, &h_now) {
1545 let md = |a: &[f32], b: &[f32]| {
1546 a.iter().zip(b).fold(0f32, |m, (x, y)| m.max((x - y).abs()))
1547 };
1548 let nn = |a: &[f32]| a.iter().map(|x| x * x).sum::<f32>().sqrt();
1549 eprintln!(
1550 "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}",
1551 nn(&cq),
1552 md(&cq, &dq),
1553 nn(&ck),
1554 md(&ck, &dk),
1555 nn(&cv),
1556 md(&cv, &dv),
1557 nn(&ao),
1558 md(&ao, &dao)
1559 );
1560 } else {
1561 eprintln!("attn-oracle L{li}: device probe declined");
1562 }
1563 }
1564 graph.encode_attn_suffix(l, &ao);
1565 graph.commit();
1568 attention::recycle_buf(&mut ao);
1569 }
1570 }
1571
1572 crate::gpu::stageprof(_xkind, _xt0.elapsed());
1573 }
1574 let mut lm_rows = None;
1579 if self.graph_want_logits
1580 && upto.is_none()
1581 && end == self.num_layers
1582 && std::env::var("CMF_GPU_LMHEAD")
1583 .map(|v| v != "0")
1584 .unwrap_or(true)
1585 {
1586 if let Some(lm) = self.weights.lm_head.q1_parts() {
1587 if graph.lm_head_ok(lm) {
1588 graph.encode_lm_head(&self.weights.final_norm, lm);
1589 lm_rows = Some(lm.1);
1590 }
1591 }
1592 }
1593 let _sy0 = std::time::Instant::now();
1594 graph.sync();
1595 let _rs0 = std::time::Instant::now();
1596 if !pending.is_empty() {
1597 let idxs: Vec<usize> = pending.drain(..).flat_map(|(f, n)| f..f + n).collect();
1598 let mut outs: Vec<&mut [f32]> = self
1599 .kv_cache
1600 .layers
1601 .iter_mut()
1602 .enumerate()
1603 .filter(|(i, _)| idxs.binary_search(i).is_ok())
1604 .map(|(_, s)| s.linear_state.as_mut_slice())
1605 .collect();
1606 graph.read_states(&mut outs);
1607 }
1608 if std::env::var("CMF_GRAPH_HOSTPROF").as_deref() == Ok("1") {
1609 use std::sync::atomic::{AtomicU64, Ordering};
1610 static SY: AtomicU64 = AtomicU64::new(0);
1611 static RS: AtomicU64 = AtomicU64::new(0);
1612 static N: AtomicU64 = AtomicU64::new(0);
1613 SY.fetch_add((_rs0 - _sy0).as_nanos() as u64, Ordering::Relaxed);
1614 RS.fetch_add(_rs0.elapsed().as_nanos() as u64, Ordering::Relaxed);
1615 let n = N.fetch_add(1, Ordering::Relaxed) + 1;
1616 if n % 100 == 0 {
1617 eprintln!(
1618 "postprof: sync-wait {:.1} ms/ток | read_states {:.1} ms/ток ({n})",
1619 SY.load(Ordering::Relaxed) as f64 / n as f64 / 1e6,
1620 RS.load(Ordering::Relaxed) as f64 / n as f64 / 1e6
1621 );
1622 }
1623 }
1624 if let Some(rows) = lm_rows {
1625 crate::gpu::hostprof_encode_done(_mt0);
1626 let mut lg = attention::take_buf(rows.min(self.vocab_size));
1627 graph.read_logits(&mut lg);
1628 crate::gpu::hostprof_total(_mt0);
1629 lg.resize(self.vocab_size, 0.0);
1630 if let Some(c) = self.final_softcap {
1631 for l in lg.iter_mut() {
1632 *l = c * (*l / c).tanh();
1633 }
1634 }
1635 self.graph_logits = Some(lg);
1636 }
1637 graph.finish(h);
1638 for li in dev_attn {
1642 let mut krow = attention::take_buf(nkv * hd);
1643 let mut vrow = attention::take_buf(nkv * hd);
1644 if crate::gpu::kv_mirror_read_last(kv_id, li, nkv, hd, &mut krow, &mut vrow) {
1645 let cache = &mut self.kv_cache.layers[li];
1646 cache.append(&krow, &vrow, &[]);
1647 let n = cache.seq_len;
1648 let mut imp = attention::take_buf(n);
1649 crate::gpu::kv_mirror_take_imp(kv_id, li, &mut imp);
1650 cache.accumulate_imp(&imp);
1651 attention::recycle_buf(&mut imp);
1652 }
1653 attention::recycle_buf(&mut krow);
1654 attention::recycle_buf(&mut vrow);
1655 }
1656 end
1657 }
1658
1659 pub fn new(
1660 tokenizer: Tokenizer,
1661 weights: PipelineWeights,
1662 hidden_size: usize,
1663 intermediate_size: usize,
1664 num_heads: usize,
1665 num_kv_heads: usize,
1666 head_dim: usize,
1667 num_layers: usize,
1668 physical_layers: usize,
1669 loop_final_norm: bool,
1670 vocab_size: usize,
1671 rms_eps: f64,
1672 rope_base: f32,
1673 norm_style: NormStyle,
1674 max_seq_len: usize,
1675 sampler_config: SamplerConfig,
1676 ) -> Self {
1677 let rng = match sampler_config.seed {
1678 Some(s) => SplitMix64::new(s),
1679 None => SplitMix64::from_entropy(),
1680 };
1681 let inv_freq = std::sync::Arc::new(attention::rope_inv_freq(head_dim, rope_base));
1682 let pool = Pool::from_env();
1683 if let Some(p) = &pool {
1684 tracing::info!("worker pool: {} threads", p.n_workers());
1685 }
1686 Self {
1687 gpu_plan: None,
1688 tokenizer: std::sync::Arc::new(tokenizer),
1689 kv_cache: KvCache::new(num_layers, num_kv_heads, head_dim, max_seq_len),
1690 sampler_config,
1691 weights,
1692 hidden_size,
1693 intermediate_size,
1694 num_heads,
1695 num_kv_heads,
1696 head_dim,
1697 num_layers,
1698 physical_layers,
1699 loop_final_norm,
1700 vocab_size,
1701 rms_eps,
1702 rope_base,
1703 norm_style,
1704 rotary_dim: head_dim,
1705 attention_heads_per_layer: None,
1706 vmf_cfg: None,
1707 gdn_cfg: None,
1708 kda_cfg: None,
1709 g3n: None,
1710 dsv4: None,
1711 qwen4_exp: None,
1712 dsv4_mtp: Vec::new(),
1713 dspark: None,
1714 dspark_pending: Vec::new(),
1715 dspark_hist: Vec::new(),
1716 dspark_real: Vec::new(),
1717 dspark_trunk_picks: Vec::new(),
1718 dspark_exp: Vec::new(),
1719 dspark_draft_ns: 0,
1720 logit_multiplier: None,
1721 cancel: std::sync::Arc::new(std::sync::atomic::AtomicBool::new(false)),
1722 kv_history: Vec::new(),
1723 short_conv_cfg: None,
1724 mtp: None,
1725 speculative: std::env::var("CMF_MTP").map(|v| v != "0").unwrap_or(true),
1726 rng,
1727 sampler_scratch: SamplerScratch::default(),
1728 spec_forced: None,
1729 spec_q: Vec::new(),
1730 spec_p: Vec::new(),
1731 spec_res: Vec::new(),
1732 spec_qs: Vec::new(),
1733 spec_ps: Vec::new(),
1734 spec_ress: Vec::new(),
1735 mtp_graph_mode: None,
1736 #[cfg(target_os = "macos")]
1737 metal_verify: None,
1738 inv_freq,
1739 ws: ForwardScratch::new(hidden_size),
1740 pool,
1741 model: None,
1742 dyn_force_f32: false,
1743 dyn_skill_layers: Vec::new(),
1744 dyn_active: None,
1745 dyn_blend_loaded: false,
1746 dyn_phi_layer: None,
1747 dyn_phi_ema: Vec::new(),
1748 dyn_phi_seen: 0,
1749 dyn_router: None,
1750 o1_cfg: None,
1751 o1_epoch: 0,
1752 o1_flags: Vec::new(),
1753 trace: false,
1754 calib_temp: 1.0,
1755 confidence_on: true,
1756 embed_multiplier: 1.0,
1757 attn_scale: 1.0 / (head_dim as f32).sqrt(),
1758 swa: None,
1759 sliding_layers: None,
1760 inv_freq_local: None,
1761 rotary_dim_local: None,
1762 rope_scale: 1.0,
1763 rope_scale_local: 1.0,
1764 global_attn: None,
1765 inv_freq_global: None,
1766 attn_v_norm: false,
1767 final_softcap: None,
1768 head_clusters: None,
1769 attn_softcap: 0.0,
1770 graph_want_logits: false,
1771 graph_logits: None,
1772 graph_kv_id: {
1773 static NEXT: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(1);
1774 NEXT.fetch_add(1, std::sync::atomic::Ordering::Relaxed)
1775 },
1776 }
1777 }
1778
1779 pub fn set_o1(&mut self, cfg: Option<crate::nystrom::O1Cfg>) {
1786 self.o1_flags = match &cfg {
1787 Some(c) => {
1788 let mut flags = c.layer_flags(self.num_layers);
1789 for (li, f) in flags.iter_mut().enumerate() {
1790 if *f
1791 && !matches!(
1792 self.weights.layers[self.phys_layer(li)].attn,
1793 AttnKind::Full { .. }
1794 )
1795 {
1796 *f = false;
1797 }
1798 }
1799 flags
1800 }
1801 None => Vec::new(),
1802 };
1803 if let Some(c) = &cfg {
1804 let n = self.o1_flags.iter().filter(|&&f| f).count();
1805 tracing::info!(
1806 "o1 nystrom attention: {n}/{} layer(s), m={} w={} sink={} rect={:?}",
1807 self.num_layers,
1808 c.m,
1809 c.w,
1810 c.sink,
1811 c.rect
1812 );
1813 }
1814 self.o1_cfg = cfg;
1815 }
1816
1817 pub fn o1_active(&self) -> bool {
1819 self.o1_cfg.is_some() && self.o1_flags.iter().any(|&f| f)
1820 }
1821
1822 pub fn o1_begin(&mut self) {
1827 if let Some(c) = &self.o1_cfg {
1828 let (m, w, sink, rect) = (c.m, c.w, c.sink, c.rect);
1829 for (li, &f) in self.o1_flags.iter().enumerate() {
1830 if f {
1831 self.kv_cache.layers[li].o1_begin(m, w, sink, rect);
1832 }
1833 }
1834 }
1835 }
1836
1837 pub fn o1_seal(&mut self) {
1841 self.o1_epoch = self.o1_epoch.wrapping_add(1);
1842 if self.o1_cfg.is_none() {
1843 return;
1844 }
1845 for li in 0..self.num_layers {
1846 if self.o1_flags.get(li).copied().unwrap_or(false) {
1847 self.kv_cache.layers[li].o1_seal(self.num_heads);
1848 }
1849 }
1850 }
1851
1852 pub fn set_trace(&mut self, on: bool) {
1854 self.trace = on;
1855 }
1856
1857 pub fn set_sampler_config(&mut self, config: SamplerConfig) {
1860 self.rng = match config.seed {
1861 Some(seed) => SplitMix64::new(seed),
1862 None => SplitMix64::from_entropy(),
1863 };
1864 self.sampler_config = config;
1865 }
1866
1867 pub fn set_confidence(&mut self, on: bool) {
1872 self.confidence_on = on;
1873 }
1874
1875 pub fn set_calib_temp(&mut self, t: f32) {
1878 self.calib_temp = if t > 1e-3 { t } else { 1.0 };
1879 }
1880
1881 pub fn calib_temp(&self) -> f32 {
1883 self.calib_temp
1884 }
1885
1886 pub fn set_rotary(&mut self, rotary_dim: usize, base: f32) {
1889 self.rotary_dim = rotary_dim.min(self.head_dim);
1890 self.inv_freq = std::sync::Arc::new(attention::rope_inv_freq(self.rotary_dim, base));
1891 }
1892
1893 fn attn_cfg(&self, position: usize) -> QwenAttnCfg<'_> {
1894 QwenAttnCfg {
1895 num_heads: self.num_heads,
1896 num_kv_heads: self.num_kv_heads,
1897 head_dim: self.head_dim,
1898 hidden_size: self.hidden_size,
1899 position,
1900 inv_freq: &self.inv_freq,
1901 rotary_dim: self.rotary_dim,
1902 scale: self.attn_scale,
1903 softcap: self.attn_softcap,
1904 window: None,
1905 v_norm: false,
1906 q_norm: None,
1907 k_norm: None,
1908 output_gate: false,
1909 softplus_gate: None,
1910 rope_scale: self.rope_scale,
1911 bias: None,
1912 rms_eps: self.rms_eps,
1913 norm_style: self.norm_style,
1914 pool: self.pool.as_deref(),
1915 }
1916 }
1917
1918 pub fn generate(
1920 &mut self,
1921 prompt: &str,
1922 max_tokens: usize,
1923 task_mask: Option<&TaskMask>,
1924 on_token: Option<TokenCallback>,
1925 ) -> Result<GenerateResult, String> {
1926 let input_ids = self.tokenizer.with_bos(self.tokenizer.encode(prompt));
1927 self.generate_from_ids(&input_ids, max_tokens, task_mask, on_token)
1928 }
1929
1930 fn drop_open_mask<'m>(&self, m: Option<&'m TaskMask>) -> Option<&'m TaskMask> {
1932 m.filter(|m| !m.fully_open(self.intermediate_size, self.num_heads))
1933 }
1934
1935 pub fn generate_from_ids(
1943 &mut self,
1944 input_ids: &[u32],
1945 max_tokens: usize,
1946 task_mask: Option<&TaskMask>,
1947 mut on_token: Option<TokenCallback>,
1948 ) -> Result<GenerateResult, String> {
1949 if std::env::var("CMF_TRACE_H").is_ok() {
1950 eprintln!("input_ids: {input_ids:?}");
1951 }
1952 if input_ids.is_empty() {
1953 return Err("empty prompt: nothing to generate from".to_string());
1954 }
1955 let task_mask = self.drop_open_mask(task_mask);
1960
1961 let reuse_from = {
1969 let on = !std::env::var("CMF_KV_REUSE").is_ok_and(|v| v == "0");
1970 let h = &self.kv_history;
1971 if on
1972 && task_mask.is_none()
1973 && self.mtp.is_none()
1974 && self.o1_cfg.is_none()
1975 && !h.is_empty()
1976 && h.len() < input_ids.len()
1977 && input_ids[..h.len()] == h[..]
1978 {
1979 h.len()
1980 } else {
1981 0
1982 }
1983 };
1984 if reuse_from == 0 {
1985 self.kv_cache.clear();
1987 self.kv_history.clear();
1988 crate::gpu::graph_kv_reset(self.graph_kv_id);
1989 } else if std::env::var("CMF_PREFILL_PROF").is_ok() {
1990 eprintln!(
1991 "kv-reuse: {} of {} prompt positions already cached",
1992 reuse_from,
1993 input_ids.len()
1994 );
1995 }
1996 crate::gpu::graph_race_begin_generation();
1997 self.o1_begin();
1998
1999 let graph_on = crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Decode);
2005 let spec_sampling_ok = self.sampler_config.temperature < 1e-6
2032 || std::env::var("CMF_GRAPH_SPEC_SAMPLE").as_deref() == Ok("1");
2033 let (mut dense_n, mut dense_q4tp) = (0usize, 0usize);
2049 for lw in &self.weights.layers {
2050 if let FfnKind::Dense(d) = &lw.ffn {
2051 dense_n += 1;
2052 if matches!(d.gate_proj.graph_weight(), Some((_, _, 6, _)))
2053 && matches!(d.up_proj.graph_weight(), Some((_, _, 6, _)))
2054 && matches!(d.down_proj.graph_weight(), Some((_, _, 6, _)))
2055 {
2056 dense_q4tp += 1;
2057 }
2058 }
2059 }
2060 let spec_default_ok = dense_n == 0 || dense_q4tp * 10 >= dense_n * 9;
2061 let penalized = self.sampler_config.repetition_penalty != 1.0
2065 || self.sampler_config.presence_penalty != 0.0
2066 || !self.sampler_config.suppress_tokens.is_empty();
2067 #[cfg(feature = "gpu")]
2072 let metal_wgpu = graph_on && crate::gpu_wgpu::wgpu_backend_is_metal();
2073 #[cfg(not(feature = "gpu"))]
2074 let metal_wgpu = false;
2075 let spec_env = std::env::var("CMF_GRAPH_SPEC").ok();
2076 let spec_wanted = match spec_env.as_deref() {
2077 Some("0") => false,
2078 Some(_) => {
2079 if metal_wgpu {
2080 tracing::warn!(
2081 "CMF_GRAPH_SPEC forced on wgpu/Metal: the batched verify graph is not \
2082 verified on this backend (garbage measured on Qwen3.5-0.8B)"
2083 );
2084 }
2085 true
2086 }
2087 None => spec_default_ok && !penalized && !metal_wgpu,
2088 };
2089 #[cfg(target_os = "macos")]
2092 let metal_graph = crate::gpu::q1_force()
2093 && crate::gpu::enabled_here()
2094 && std::env::var("CMF_GPU_BLOCK")
2095 .map(|v| v != "0")
2096 .unwrap_or(true);
2097 #[cfg(not(target_os = "macos"))]
2098 let metal_graph = false;
2099 let graph_spec = self.speculative
2100 && (graph_on || metal_graph)
2101 && self.mtp.is_some()
2102 && task_mask.is_none()
2103 && !self.o1_active()
2104 && spec_sampling_ok
2105 && spec_wanted;
2106 let pair_pays = self.gdn_cfg.is_none() || std::env::var("CMF_MTP").as_deref() == Ok("1");
2113 let spec_active = self.speculative
2114 && self.mtp.is_some()
2115 && task_mask.is_none()
2116 && !self.o1_active()
2117 && ((!graph_on && pair_pays && self.sampler_config.temperature < 1e-6) || graph_spec);
2118 let mut mtp = if spec_active { self.mtp.take() } else { None };
2121 if std::env::var("CMF_MTP_CHAIN_PROBE").is_ok() {
2122 eprintln!(
2123 "mtp-probe gate: spec_active={spec_active} mtp={} speculative={} graph_on={graph_on} temp_ok={}",
2124 mtp.is_some(),
2125 self.speculative,
2126 self.sampler_config.temperature < 1e-6,
2127 );
2128 }
2129 if let Some(m) = &mut mtp {
2130 m.kv.clear();
2131 crate::gpu::graph_kv_reset(self.mtp_kv_id());
2133 self.mtp_graph_mode = None;
2134 }
2135 let mut router = if mtp.is_none() {
2139 self.dyn_router.take()
2140 } else {
2141 None
2142 };
2143 if let Some(r) = &mut router {
2144 r.reset(); self.dyn_phi_seen = 0; let _ = self.set_active_skill(None);
2147 }
2148
2149 let mut all_ids = input_ids.to_vec();
2150 let mut generated = 0usize;
2151 let mut finish_reason = "max_tokens".to_string();
2152 let mut drafted = 0usize;
2153 let mut accepted = 0usize;
2154 let mut dsv4_spec_bad = 0usize;
2161 let mut dsv4_spec_retry_at = 0usize;
2162 let mut confidence: Vec<f32> = Vec::new();
2163 let trace_on = self.trace;
2164 let calib_temp = self.calib_temp;
2165 let mut traces: Vec<TokenTrace> = Vec::new();
2166
2167 let mut hidden = vec![0.0f32; self.hidden_size];
2173 let mut pos = reuse_from;
2174 let fuse_lm = mtp.is_none()
2183 && router.is_none()
2184 && std::env::var("CMF_GPU_LMHEAD").as_deref() != Ok("0");
2185 self.graph_logits = None;
2186 self.graph_want_logits = false;
2187 let _tpf = std::time::Instant::now();
2188 let batch_k = std::env::var("CMF_BATCH_K")
2189 .ok()
2190 .and_then(|v| v.parse::<usize>().ok())
2191 .unwrap_or(0);
2192 while self.qwen4_exp.is_some()
2203 && mtp.is_none()
2204 && pos < input_ids.len()
2205 && !self.cancel.load(std::sync::atomic::Ordering::Relaxed)
2206 {
2207 let token_id = input_ids[pos];
2208 let want_logits = pos + 1 == input_ids.len();
2209 let mut lg = Vec::new();
2210 if let Some(b) = &mut self.qwen4_exp {
2211 crate::qwen4_exp::forward_token(
2212 &b.0,
2213 &b.1,
2214 &b.2,
2215 &mut b.3,
2216 token_id,
2217 pos,
2218 &self.inv_freq,
2219 self.pool.as_deref(),
2220 &mut lg,
2221 want_logits,
2222 );
2223 }
2224 if want_logits {
2225 self.graph_logits = Some(lg);
2226 }
2227 pos += 1;
2228 hidden.fill(0.0);
2229 }
2230 while self.dsv4.is_some()
2231 && mtp.is_none()
2232 && pos < input_ids.len()
2233 && !self.cancel.load(std::sync::atomic::Ordering::Relaxed)
2234 {
2235 let end = (pos + prefill_chunk()).min(input_ids.len());
2236 let ids: Vec<u32> = input_ids[pos..end].to_vec();
2237 let mut lg = Vec::new();
2238 if let Some(b) = &mut self.dsv4 {
2239 let (g, layers, cfg, st) = (&b.0, &b.1, b.2, &mut b.3);
2240 crate::dsv4::forward_chunk(
2241 g,
2242 layers,
2243 &cfg,
2244 st,
2245 &ids,
2246 pos,
2247 &self.inv_freq,
2248 self.pool.as_deref(),
2249 &mut lg,
2250 end == input_ids.len(),
2251 );
2252 }
2253 if end == input_ids.len() {
2254 self.graph_logits = Some(lg);
2255 }
2256 pos = end;
2257 hidden = vec![0.0; self.hidden_size];
2258 }
2259 let dyn_prefill = router.is_some();
2264 let graph_prefill = self.graph_prefill_preferred();
2270 #[cfg(target_os = "macos")]
2278 if task_mask.is_none()
2279 && !dyn_prefill
2280 && crate::gpu::q1_force()
2281 && crate::gpu::enabled_here()
2282 && self.gdn_cfg.is_some()
2283 && self.g3n.is_none()
2284 && input_ids.len() > 8
2285 && std::env::var("CMF_MTP_CHAIN_PROBE").is_err()
2286 && std::env::var("CMF_METAL_PREFILL").as_deref() != Ok("0")
2287 {
2288 let chunk: usize = std::env::var("CMF_METAL_PREFILL_CHUNK")
2289 .ok()
2290 .and_then(|v| v.parse().ok())
2291 .filter(|&v| (16..=512).contains(&v))
2292 .unwrap_or(256);
2293 let hs = self.hidden_size;
2294 let _tp = std::time::Instant::now();
2295 while pos < input_ids.len() && !self.cancel.load(std::sync::atomic::Ordering::Relaxed) {
2296 let end = (pos + chunk).min(input_ids.len());
2297 let Some(hb) = self.prefill_batch_metal(&input_ids[pos..end], pos) else {
2298 break;
2299 };
2300 if let Some(m) = &mut mtp {
2301 let n_pairs = if end < input_ids.len() {
2302 end - pos
2303 } else {
2304 end - pos - 1
2305 };
2306 if n_pairs > 0 {
2307 let pairs: Vec<(&[f32], u32)> = (0..n_pairs)
2308 .map(|j| (&hb[j * hs..(j + 1) * hs], input_ids[pos + j + 1]))
2309 .collect();
2310 if !self.mtp_warm_batch_metal(m, &pairs, pos) {
2311 for (j, (h, t)) in pairs.iter().enumerate() {
2312 let h = h.to_vec();
2313 let _ = self.mtp_step(m, &h, *t, pos + j);
2314 }
2315 }
2316 }
2317 }
2318 hidden.copy_from_slice(&hb[(end - pos - 1) * hs..]);
2319 pos = end;
2320 }
2321 if std::env::var("CMF_PREFILL_PROF").is_ok() {
2322 eprintln!(
2323 "metal-prefill: {} of {} tokens in {:.1} ms",
2324 pos,
2325 input_ids.len(),
2326 _tp.elapsed().as_secs_f64() * 1e3
2327 );
2328 }
2329 }
2330 if task_mask.is_none()
2331 && !dyn_prefill
2332 && !graph_prefill
2333 && self.can_prefill_batched()
2334 && self.g3n.is_none()
2335 && input_ids.len() > 2
2336 {
2337 let chunk = prefill_chunk();
2343 let hs = self.hidden_size;
2344 while pos < input_ids.len() && !self.cancel.load(std::sync::atomic::Ordering::Relaxed) {
2345 let end = (pos + chunk).min(input_ids.len());
2346 let hb = self.prefill_batch(&input_ids[pos..end], pos);
2347 if let Some(m) = &mut mtp {
2348 let probe: usize = std::env::var("CMF_MTP_CHAIN_PROBE")
2349 .ok()
2350 .and_then(|v| v.parse().ok())
2351 .unwrap_or(0);
2352 for p in pos..end {
2353 if p + 1 < input_ids.len() {
2354 if probe >= 1 && p + 2 < input_ids.len() {
2355 let (d1, mut hx) = self.mtp_step_h(
2359 m,
2360 &hb[(p - pos) * hs..(p - pos + 1) * hs],
2361 input_ids[p + 1],
2362 p,
2363 );
2364 let mut ok = d1 == input_ids[p + 2];
2365 Self::chain_probe_note(0, ok);
2366 let mut d_prev = d1;
2367 let mut extra = 0usize;
2368 for j in 1..probe {
2369 if p + 2 + j >= input_ids.len() {
2370 break;
2371 }
2372 let (dj, hj) = self.mtp_step_h(m, &hx, d_prev, p + 1 + j);
2373 extra += 1;
2374 ok = ok && dj == input_ids[p + 2 + j];
2375 Self::chain_probe_note(j, ok);
2376 d_prev = dj;
2377 hx = hj;
2378 }
2379 m.kv.truncate_last(extra);
2380 } else {
2381 let _ = self.mtp_step(
2382 m,
2383 &hb[(p - pos) * hs..(p - pos + 1) * hs],
2384 input_ids[p + 1],
2385 p,
2386 );
2387 }
2388 }
2389 }
2390 }
2391 hidden.copy_from_slice(&hb[(end - pos - 1) * hs..]);
2392 pos = end;
2393 }
2394 }
2395 let pair_off = std::env::var("CMF_PAIR").is_ok_and(|v| v == "0");
2396 if task_mask.is_none()
2397 && !dyn_prefill
2398 && !graph_prefill
2399 && !pair_off
2400 && self.pair_supported()
2401 {
2402 while pos + 1 < input_ids.len()
2403 && !self.cancel.load(std::sync::atomic::Ordering::Relaxed)
2404 {
2405 let e1 = self.embed_single(input_ids[pos]);
2406 let e2 = self.embed_single(input_ids[pos + 1]);
2407 let (h1, h2) = self.forward_pair(&e1, &e2, pos);
2408 self.commit_linear_scratch();
2410 if let Some(m) = &mut mtp {
2411 let _ = self.mtp_step(m, &h1, input_ids[pos + 1], pos);
2412 if pos + 2 < input_ids.len() {
2413 let probe: usize = std::env::var("CMF_MTP_CHAIN_PROBE")
2414 .ok()
2415 .and_then(|v| v.parse().ok())
2416 .unwrap_or(0);
2417 if probe >= 1 && pos + 3 < input_ids.len() {
2418 let (d1, mut hx) = self.mtp_step_h(m, &h2, input_ids[pos + 2], pos + 1);
2422 let mut ok = d1 == input_ids[pos + 3];
2423 Self::chain_probe_note(0, ok);
2424 let mut d_prev = d1;
2425 let mut extra = 0usize;
2426 for j in 1..probe {
2427 if pos + 3 + j >= input_ids.len() {
2428 break;
2429 }
2430 let (dj, hj) = self.mtp_step_h(m, &hx, d_prev, pos + 2 + j);
2431 extra += 1;
2432 ok = ok && dj == input_ids[pos + 3 + j];
2433 Self::chain_probe_note(j, ok);
2434 d_prev = dj;
2435 hx = hj;
2436 }
2437 m.kv.truncate_last(extra);
2438 } else {
2439 let _ = self.mtp_step(m, &h2, input_ids[pos + 2], pos + 1);
2440 }
2441 }
2442 }
2443 hidden = h2;
2444 pos += 2;
2445 }
2446 }
2447 if batch_k > 0
2456 && graph_prefill
2457 && task_mask.is_none()
2458 && !self.o1_active()
2459 && mtp.is_none()
2460 && !dyn_prefill
2461 && pos + 1 < input_ids.len()
2462 {
2463 let hs = self.hidden_size;
2464 let chunk = batch_k;
2465 while pos < input_ids.len() {
2466 let end = (pos + chunk).min(input_ids.len());
2467 let bk = end - pos;
2468 let mut hiddens = vec![0f32; bk * hs];
2469 for (j, &id) in input_ids[pos..end].iter().enumerate() {
2470 hiddens[j * hs..(j + 1) * hs].copy_from_slice(&self.embed_single(id));
2471 }
2472 let positions: Vec<usize> = (pos..end).collect();
2473 let t_chunk = std::time::Instant::now();
2474 let ok_b = self.try_batch_graph_wgpu(&mut hiddens, &positions, bk, None);
2475 if std::env::var("CMF_GRAPH_PROF").is_ok() {
2476 let ms = t_chunk.elapsed().as_secs_f64() * 1000.0;
2477 eprintln!(
2478 "batch-chunk: k={bk} ok={ok_b} {ms:.1} ms ({:.1} tok/s)",
2479 bk as f64 / (ms / 1000.0)
2480 );
2481 }
2482 {
2483 use std::sync::atomic::{AtomicBool, Ordering};
2484 static SAID: AtomicBool = AtomicBool::new(false);
2485 if !SAID.swap(true, Ordering::Relaxed) {
2486 if ok_b {
2487 tracing::info!("batched prefill: ACTIVE (k={bk})");
2488 } else {
2489 tracing::warn!("batched prefill declined — per-position graph");
2490 }
2491 }
2492 }
2493 if ok_b {
2494 hidden.copy_from_slice(&hiddens[(bk - 1) * hs..]);
2495 pos = end;
2496 } else {
2497 break; }
2499 }
2500 }
2501 while pos < input_ids.len() && !self.cancel.load(std::sync::atomic::Ordering::Relaxed) {
2502 self.graph_want_logits = fuse_lm && pos + 1 == input_ids.len();
2503 hidden = self.forward_layers(&self.embed_single(input_ids[pos]), pos, task_mask);
2504 if let Some(m) = &mut mtp {
2505 if pos + 1 < input_ids.len() {
2506 let probe: usize = std::env::var("CMF_MTP_CHAIN_PROBE")
2512 .ok()
2513 .and_then(|v| v.parse().ok())
2514 .unwrap_or(0);
2515 if probe >= 1 && pos + 2 < input_ids.len() {
2516 let (d1, mut hx) = self.mtp_step_h(m, &hidden, input_ids[pos + 1], pos);
2517 let mut ok = d1 == input_ids[pos + 2];
2518 Self::chain_probe_note(0, ok);
2519 let mut d_prev = d1;
2520 let mut extra = 0usize;
2521 for j in 1..probe {
2522 if pos + 2 + j >= input_ids.len() {
2523 break;
2524 }
2525 let (dj, hj) = self.mtp_step_h(m, &hx, d_prev, pos + 1 + j);
2526 extra += 1;
2527 ok = ok && dj == input_ids[pos + 2 + j];
2528 Self::chain_probe_note(j, ok);
2529 d_prev = dj;
2530 hx = hj;
2531 }
2532 m.kv.truncate_last(extra);
2535 } else {
2536 let _ = self.mtp_step(m, &hidden, input_ids[pos + 1], pos);
2537 }
2538 }
2539 }
2540 pos += 1;
2541 }
2542 if std::env::var("CMF_PREFILL_PROF").is_ok() {
2543 eprintln!(
2544 "prefill: {} tokens in {:.1} ms (batch_k={batch_k})",
2545 input_ids.len(),
2546 _tpf.elapsed().as_secs_f64() * 1000.0
2547 );
2548 }
2549 if self
2552 .cancel
2553 .swap(false, std::sync::atomic::Ordering::Relaxed)
2554 {
2555 self.kv_history.clear();
2556 if let Some(m) = mtp {
2557 self.mtp = Some(m);
2558 }
2559 return Ok(GenerateResult {
2560 text: String::new(),
2561 token_ids: Vec::new(),
2562 prompt_tokens: input_ids.len(),
2563 tokens_generated: 0,
2564 finish_reason: "cancelled".to_string(),
2565 mtp_drafted: 0,
2566 mtp_accepted: 0,
2567 token_confidence: Vec::new(),
2568 traces: Vec::new(),
2569 });
2570 }
2571
2572 self.o1_seal();
2575
2576 macro_rules! commit {
2578 ($id:expr) => {{
2579 all_ids.push($id);
2580 generated += 1;
2581 if self.tokenizer.is_eos($id) {
2582 finish_reason = "stop".to_string();
2583 false
2584 } else {
2585 let token_text = self.tokenizer.decode_token($id);
2586 let mut go = true;
2587 if let Some(ref mut cb) = on_token {
2588 if !cb(&token_text) {
2589 finish_reason = "cancelled".to_string();
2590 go = false;
2591 }
2592 }
2593 go
2594 }
2595 }};
2596 }
2597
2598 let mut spec_trial = SpecTrial::Spec {
2609 t0: std::time::Instant::now(),
2610 gen0: generated,
2611 rounds: 0,
2612 };
2613 let mut spec_mon = SpecMon::default();
2614 let mut spec_watchdog_off = false;
2615 let mut next_pos = input_ids.len();
2617 'decode: while generated < max_tokens {
2618 if self
2619 .cancel
2620 .swap(false, std::sync::atomic::Ordering::Relaxed)
2621 {
2622 finish_reason = "cancelled".to_string();
2623 break 'decode;
2624 }
2625 let forced = self.spec_forced.take();
2630 let mut logits = match (forced, self.graph_logits.take()) {
2631 (Some(_), _) => Vec::new(),
2632 (None, Some(lg)) => lg,
2633 (None, None) => {
2634 inference::rms_norm_into(
2635 &hidden,
2636 &self.weights.final_norm,
2637 self.rms_eps,
2638 self.norm_style,
2639 &mut self.ws.n1,
2640 );
2641 self.lm_head_forward(&self.ws.n1)
2642 }
2643 };
2644 if generated
2647 == std::env::var("CMF_LOGIT_DUMP_STEP")
2648 .ok()
2649 .and_then(|v| v.parse().ok())
2650 .unwrap_or(0)
2651 {
2652 if let Ok(path) = std::env::var("CMF_LOGIT_DUMP") {
2653 let mut bytes: Vec<u8> = Vec::with_capacity((hidden.len() + logits.len()) * 4);
2654 for v in hidden.iter().chain(logits.iter()) {
2655 bytes.extend_from_slice(&v.to_le_bytes());
2656 }
2657 let _ = std::fs::write(&path, &bytes);
2658 }
2659 }
2660 let t_next = match forced {
2661 Some(c) => c,
2662 None => sampler::sample_with_scratch_pool(
2663 &logits,
2664 &self.sampler_config,
2665 &all_ids,
2666 &mut self.rng,
2667 &mut self.sampler_scratch,
2668 self.pool.as_deref(),
2669 ),
2670 };
2671 if self.confidence_on {
2672 confidence.push(if logits.is_empty() {
2673 0.0
2674 } else {
2675 sampler::top1_prob_pool(
2676 self.pool.as_deref(),
2677 &mut self.sampler_scratch,
2678 &logits,
2679 t_next,
2680 calib_temp,
2681 )
2682 });
2683 }
2684 if !logits.is_empty() {
2685 attention::recycle_buf(&mut logits);
2686 }
2687 if trace_on {
2688 let skill = router.as_ref().and_then(|r| r.active_id());
2692 traces.push(TokenTrace {
2693 t: generated,
2694 token_id: t_next,
2695 confidence: confidence.last().copied().unwrap_or(0.0),
2696 active_skill: skill,
2697 recon: None,
2698 switched: false,
2699 });
2700 }
2701 if !commit!(t_next) {
2702 break 'decode;
2703 }
2704 if generated >= max_tokens {
2705 break 'decode;
2706 }
2707
2708 if self.kv_cache.needs_eviction() {
2709 static SAID: std::sync::Once = std::sync::Once::new();
2715 SAID.call_once(|| {
2716 tracing::warn!(
2717 "KV cache full at {} positions — evicting half; quality \
2718 will degrade. Raise CMF_MAX_SEQ.",
2719 self.kv_cache.max_seq_len,
2720 );
2721 });
2722 let keep = (self.kv_cache.max_seq_len / 2).max(1);
2723 self.kv_cache.evict(keep);
2724 }
2725
2726 if graph_spec {
2729 match spec_trial {
2730 SpecTrial::Plain { t0, gen0 } if generated >= gen0 + 8 => {
2731 spec_mon.plain_ms =
2732 t0.elapsed().as_secs_f64() * 1e3 / (generated - gen0) as f64;
2733 let keep = spec_mon.pays();
2734 tracing::info!(
2735 "speculation trial: {:.2} tok/round in {:.1} ms vs plain {:.1} ms/tok — {}",
2736 spec_mon.tokens,
2737 spec_mon.round_ms,
2738 spec_mon.plain_ms,
2739 if keep { "speculating" } else { "plain" }
2740 );
2741 spec_mon.fails = 0;
2742 spec_trial = SpecTrial::Decided {
2743 spec: keep,
2744 recheck_at: if keep { usize::MAX } else { generated + 128 },
2745 };
2746 }
2747 SpecTrial::Decided { recheck_at, .. } if generated >= recheck_at => {
2748 spec_mon.n = 0;
2749 spec_trial = SpecTrial::Spec {
2750 t0: std::time::Instant::now(),
2751 gen0: generated,
2752 rounds: 0,
2753 };
2754 }
2755 _ => {}
2756 }
2757 spec_watchdog_off = matches!(
2758 spec_trial,
2759 SpecTrial::Plain { .. } | SpecTrial::Decided { spec: false, .. }
2760 );
2761 }
2762 match &mut mtp {
2763 #[cfg(feature = "gpu")]
2765 Some(m)
2766 if graph_spec
2767 && !spec_watchdog_off
2768 && generated + 1 < max_tokens
2769 && next_pos > 0 =>
2770 {
2771 let t_round = std::time::Instant::now();
2772 if let Some((extra, n_pos, new_h)) = self.graph_spec_step(
2773 m,
2774 &hidden,
2775 t_next,
2776 next_pos,
2777 &mut drafted,
2778 &mut accepted,
2779 &mut all_ids,
2780 ) {
2781 next_pos = n_pos;
2782 hidden = new_h;
2783 if std::env::var("CMF_GRAPH_SPEC_TIME").is_ok() {
2784 eprintln!(
2785 "spec-round wall {:.1} ms → {} tokens",
2786 t_round.elapsed().as_secs_f64() * 1e3,
2787 extra.len() + 1
2788 );
2789 }
2790 spec_mon.round(t_round.elapsed().as_secs_f64() * 1e3, extra.len() + 1);
2794 spec_trial = Self::spec_trial_round(
2797 spec_trial,
2798 &mut spec_mon,
2799 generated + extra.len() + 1,
2800 );
2801 let mut stopped = false;
2802 for &id in &extra {
2803 if self.confidence_on {
2804 confidence.push(0.0);
2805 }
2806 if !commit!(id) {
2807 stopped = true;
2808 break;
2809 }
2810 }
2811 if stopped {
2812 break 'decode;
2813 }
2814 continue 'decode;
2815 }
2816 spec_mon.round(t_round.elapsed().as_secs_f64() * 1e3, 1);
2827 spec_mon.tokens = 0.0;
2828 spec_mon.fails = 3;
2829 spec_trial = Self::spec_trial_round(spec_trial, &mut spec_mon, generated + 1);
2830 hidden = self.forward_layers(&self.embed_single(t_next), next_pos, task_mask);
2831 next_pos += 1;
2832 continue 'decode;
2833 }
2834 Some(m) if !graph_spec && generated + 1 < max_tokens => {
2836 let draft = self.mtp_step(m, &hidden, t_next, next_pos - 1);
2837 drafted += 1;
2838 let emb1 = self.embed_single(t_next);
2839 let emb2 = self.embed_single(draft);
2840 let (h1, h2) = self.forward_pair(&emb1, &emb2, next_pos);
2841
2842 inference::rms_norm_into(
2843 &h1,
2844 &self.weights.final_norm,
2845 self.rms_eps,
2846 self.norm_style,
2847 &mut self.ws.n1,
2848 );
2849 let mut logits1 = self.lm_head_forward(&self.ws.n1);
2850 let t_after = sampler::sample_with_scratch_pool(
2851 &logits1,
2852 &self.sampler_config,
2853 &all_ids,
2854 &mut self.rng,
2855 &mut self.sampler_scratch,
2856 self.pool.as_deref(),
2857 );
2858 if self.confidence_on {
2859 confidence.push(sampler::top1_prob_pool(
2860 self.pool.as_deref(),
2861 &mut self.sampler_scratch,
2862 &logits1,
2863 t_after,
2864 calib_temp,
2865 ));
2866 }
2867 attention::recycle_buf(&mut logits1);
2868 if trace_on {
2869 traces.push(TokenTrace {
2872 t: generated,
2873 token_id: t_after,
2874 confidence: confidence.last().copied().unwrap_or(0.0),
2875 active_skill: None,
2876 recon: None,
2877 switched: false,
2878 });
2879 }
2880 let stop = !commit!(t_after);
2881
2882 if t_after == draft {
2883 accepted += 1;
2884 self.commit_linear_scratch();
2885 let _ = self.mtp_step(m, &h1, t_after, next_pos);
2886 hidden = h2;
2887 next_pos += 2;
2888 } else {
2889 for layer in &mut self.kv_cache.layers {
2891 layer.truncate_last(1);
2892 }
2893 if !stop {
2894 let _ = self.mtp_step(m, &h1, t_after, next_pos);
2895 hidden = self.forward_layers(
2896 &self.embed_single(t_after),
2897 next_pos + 1,
2898 None,
2899 );
2900 }
2901 next_pos += 2;
2902 }
2903 if stop {
2904 break 'decode;
2905 }
2906 }
2907 _ => {
2909 #[cfg(feature = "gpu")]
2914 if Self::dsv4_spec_on() && self.dsv4.is_some() {
2915 static SAID: std::sync::Once = std::sync::Once::new();
2916 SAID.call_once(|| {
2917 eprintln!(
2918 "dsv4-spec гейт: mtp={} mask={} router={} trace={} temp={} rep={} ",
2919 !self.dsv4_mtp.is_empty(),
2920 task_mask.is_none(),
2921 router.is_none(),
2922 !trace_on,
2923 self.sampler_config.temperature < 1e-6,
2924 self.sampler_config.repetition_penalty == 1.0,
2925 );
2926 });
2927 }
2928 #[cfg(feature = "gpu")]
2929 if Self::dsv4_spec_on()
2930 && self.dsv4.is_some()
2931 && !self.dsv4_mtp.is_empty()
2932 && task_mask.is_none()
2933 && router.is_none()
2934 && !trace_on
2935 && self.sampler_config.temperature < 1e-6
2936 && self.sampler_config.repetition_penalty == 1.0
2937 && generated + 1 < max_tokens
2938 && all_ids.len() >= 2
2939 && generated >= dsv4_spec_retry_at
2940 {
2941 let tip_token = all_ids[all_ids.len() - 2];
2942 let drafted0 = drafted;
2943 let round = self.dsv4_spec_step(
2944 tip_token,
2945 t_next,
2946 next_pos,
2947 max_tokens.saturating_sub(generated),
2948 &mut drafted,
2949 &mut accepted,
2950 );
2951 if drafted > drafted0 {
2952 let useful = round.as_ref().is_some_and(|(extra, _)| !extra.is_empty());
2953 if useful {
2954 dsv4_spec_bad = 0;
2955 } else {
2956 dsv4_spec_bad += 1;
2957 if dsv4_spec_bad >= 2 {
2958 dsv4_spec_bad = 0;
2959 dsv4_spec_retry_at = generated.saturating_add(32);
2960 tracing::info!(
2961 "dsv4: draft не окупился дважды — точный walk на 32 токена"
2962 );
2963 }
2964 }
2965 }
2966 if let Some((extra, n_pos)) = round {
2967 next_pos = n_pos;
2968 let mut stopped = false;
2969 for &id in &extra {
2970 if self.confidence_on {
2971 confidence.push(0.0);
2972 }
2973 if !commit!(id) {
2974 stopped = true;
2975 break;
2976 }
2977 }
2978 if stopped {
2979 break 'decode;
2980 }
2981 continue 'decode;
2982 }
2983 }
2984 self.graph_want_logits = fuse_lm;
2985 let mut t_fwd = t_next;
2991 let pure_greedy = self.sampler_config.temperature < 1e-6
2992 && self.sampler_config.repetition_penalty == 1.0
2993 && self.sampler_config.suppress_tokens.is_empty();
2994 let burst_k = std::env::var("CMF_MULTISTEP")
2999 .ok()
3000 .and_then(|v| v.parse::<usize>().ok())
3001 .unwrap_or(0);
3002 if pure_greedy
3003 && burst_k >= 1
3004 && fuse_lm
3005 && task_mask.is_none()
3006 && router.is_none()
3007 && !trace_on
3008 && !self.confidence_on
3009 {
3010 let mut stopped = false;
3011 loop {
3012 let room = max_tokens.saturating_sub(generated);
3013 if room <= 2 {
3014 break;
3015 }
3016 let k = burst_k.min(room - 1);
3017 if k < 1 {
3018 break;
3019 }
3020 let Some(ids) = self.try_multi_burst(t_fwd, next_pos, k) else {
3021 break;
3022 };
3023 next_pos += k;
3024 for &id in &ids {
3025 if !commit!(id) {
3026 stopped = true;
3027 break;
3028 }
3029 }
3030 if stopped {
3031 break;
3032 }
3033 t_fwd = *ids.last().unwrap();
3034 }
3035 if stopped {
3036 break 'decode;
3037 }
3038 }
3039 hidden = self.forward_layers(&self.embed_single(t_fwd), next_pos, task_mask);
3040 next_pos += 1;
3041 if let Some(r) = &mut router {
3044 let phi = self.dyn_phi_ema.clone();
3045 let decision = r.step(&phi, generated);
3046 if let Some(new_active) = decision {
3047 let _ = self.set_active_skill(new_active);
3048 }
3049 if trace_on {
3052 if let Some(last) = traces.last_mut() {
3053 let e = r.last_best_e();
3054 last.recon = e.is_finite().then_some(e);
3055 last.switched = decision.is_some();
3056 }
3057 }
3058 }
3059 }
3060 }
3061 }
3062
3063 self.graph_want_logits = false;
3064 self.graph_logits = None;
3065 if router.is_some() {
3067 let _ = self.set_active_skill(None);
3068 }
3069 self.dyn_router = router.or(self.dyn_router.take());
3070 self.mtp = mtp.or(self.mtp.take());
3071
3072 let output_ids = &all_ids[input_ids.len()..];
3073 let forwarded = input_ids.len() + output_ids.len().saturating_sub(1);
3077 self.kv_history = all_ids[..forwarded.min(all_ids.len())].to_vec();
3078 confidence.truncate(output_ids.len()); traces.truncate(output_ids.len());
3080 Ok(GenerateResult {
3081 text: self.tokenizer.decode(output_ids),
3082 token_ids: output_ids.to_vec(),
3083 prompt_tokens: input_ids.len(),
3084 tokens_generated: generated,
3085 finish_reason,
3086 mtp_drafted: drafted,
3087 mtp_accepted: accepted,
3088 token_confidence: confidence,
3089 traces,
3090 })
3091 }
3092
3093 fn mtp_step(
3097 &mut self,
3098 m: &mut MtpModule,
3099 hidden: &[f32],
3100 next_token: u32,
3101 position: usize,
3102 ) -> u32 {
3103 self.mtp_step_h(m, hidden, next_token, position).0
3104 }
3105
3106 fn chain_probe_note(depth: usize, prefix_ok: bool) {
3110 use std::sync::Mutex;
3111 static T: Mutex<Vec<(u64, u64)>> = Mutex::new(Vec::new());
3112 let mut t = T.lock().unwrap();
3113 if t.len() <= depth {
3114 t.resize(depth + 1, (0, 0));
3115 }
3116 t[depth].0 += 1;
3117 t[depth].1 += prefix_ok as u64;
3118 if depth == 0 && t[0].0 % 128 == 0 {
3119 let line: Vec<String> = t
3120 .iter()
3121 .enumerate()
3122 .map(|(d, (n, k))| {
3123 format!(
3124 "d{}={:.0}%({n})",
3125 d + 1,
3126 100.0 * *k as f64 / (*n).max(1) as f64
3127 )
3128 })
3129 .collect();
3130 eprintln!("mtp-chain: {}", line.join(" "));
3131 }
3132 }
3133
3134 fn mtp_step_hl(
3142 &mut self,
3143 m: &mut MtpModule,
3144 hidden: &[f32],
3145 next_token: u32,
3146 position: usize,
3147 ) -> (Vec<f32>, Vec<f32>) {
3148 #[cfg(target_os = "macos")]
3153 if self.mtp_graph_mode != Some(false) && crate::gpu::q1_force() {
3154 if let Some(r) = self.mtp_step_metal(m, hidden, next_token, position, true) {
3155 self.mtp_graph_mode = Some(true);
3156 return r;
3157 }
3158 self.mtp_graph_mode = Some(false);
3159 }
3160 #[cfg(feature = "gpu")]
3161 if self.mtp_graph_mode != Some(false) {
3162 if let Some(r) = self.mtp_step_graph(m, hidden, next_token, position) {
3163 self.mtp_graph_mode = Some(true);
3164 return r;
3165 }
3166 if self.mtp_graph_mode == Some(true) {
3167 tracing::warn!("mtp graph declined mid-run — draft falls to the per-op path");
3172 }
3173 self.mtp_graph_mode = Some(false);
3174 }
3175 let e = self.embed_single(next_token);
3179 let mut cat = vec![0.0f32; 2 * self.hidden_size];
3180 let (cat_e, cat_h) = cat.split_at_mut(self.hidden_size);
3181 inference::rms_norm_into(&e, &m.enorm, self.rms_eps, self.norm_style, cat_e);
3182 inference::rms_norm_into(hidden, &m.hnorm, self.rms_eps, self.norm_style, cat_h);
3183 let mut x = vec![0.0f32; self.hidden_size];
3184 m.eh_proj.matvec(&cat, &mut x, self.pool.as_deref());
3185
3186 let lw = &m.layer;
3188 inference::rms_norm_into(
3189 &x,
3190 &lw.input_norm,
3191 self.rms_eps,
3192 self.norm_style,
3193 &mut self.ws.n1,
3194 );
3195 let attn = match &lw.attn {
3196 AttnKind::Mla(_) => unreachable!("MLA has no MTP/pair path"),
3198 AttnKind::Kda(_) => unreachable!("KDA has no MTP/pair path"),
3199 AttnKind::Full {
3200 wq,
3201 wk,
3202 wv,
3203 wo,
3204 q_norm,
3205 k_norm,
3206 output_gate,
3207 softplus_gate,
3208 bias,
3209 } => {
3210 let mut cfg = self.attn_cfg(position);
3211 cfg.q_norm = q_norm.as_deref();
3212 cfg.k_norm = k_norm.as_deref();
3213 cfg.output_gate = *output_gate;
3214 cfg.softplus_gate = softplus_gate
3215 .as_ref()
3216 .map(|(gate, per_head)| (gate, *per_head));
3217 cfg.bias = bias
3218 .as_ref()
3219 .map(|(q, k, v)| (q.as_slice(), k.as_slice(), v.as_slice()));
3220 attention::qwen_attention(&self.ws.n1, wq, wk, wv, wo, &mut m.kv, &cfg)
3221 }
3222 AttnKind::Linear(_) | AttnKind::LinearGdn(_) | AttnKind::ShortConv(_) => {
3223 unreachable!("MTP block is full attention")
3224 }
3225 };
3226 for (i, &a) in attn.iter().enumerate() {
3227 x[i] += a;
3228 }
3229 inference::rms_norm_into(
3230 &x,
3231 &lw.post_norm,
3232 self.rms_eps,
3233 self.norm_style,
3234 &mut self.ws.p1,
3235 );
3236 let ffn = ffn_forward(&lw.ffn, &self.ws.p1, self.pool.as_deref(), None);
3237 for (i, &f) in ffn.iter().enumerate() {
3238 x[i] += f;
3239 }
3240
3241 inference::rms_norm_into(
3242 &x,
3243 &m.final_norm,
3244 self.rms_eps,
3245 self.norm_style,
3246 &mut self.ws.n1,
3247 );
3248 let lg = self.lm_head_forward(&self.ws.n1);
3249 (lg, x)
3250 }
3251
3252 fn mtp_step_h(
3254 &mut self,
3255 m: &mut MtpModule,
3256 hidden: &[f32],
3257 next_token: u32,
3258 position: usize,
3259 ) -> (u32, Vec<f32>) {
3260 let (mut lg, x) = self.mtp_step_hl(m, hidden, next_token, position);
3261 let draft = sampler::argmax(&lg);
3262 attention::recycle_buf(&mut lg);
3263 (draft, x)
3264 }
3265
3266 fn spec_trial_round(trial: SpecTrial, mon: &mut SpecMon, generated: usize) -> SpecTrial {
3272 match trial {
3273 SpecTrial::Spec { t0, gen0, rounds } => {
3274 let rounds = rounds + 1;
3275 if rounds >= 5 {
3276 if mon.plain_ms > 0.0 {
3277 let keep = mon.pays();
3278 mon.fails = 0;
3279 tracing::info!(
3280 "speculation re-check: {:.2} tok/round in {:.1} ms vs plain {:.1} ms/tok — {}",
3281 mon.tokens,
3282 mon.round_ms,
3283 mon.plain_ms,
3284 if keep { "speculating" } else { "plain" }
3285 );
3286 SpecTrial::Decided {
3287 spec: keep,
3288 recheck_at: if keep { usize::MAX } else { generated + 128 },
3289 }
3290 } else {
3291 SpecTrial::Plain {
3292 t0: std::time::Instant::now(),
3293 gen0: generated,
3294 }
3295 }
3296 } else {
3297 SpecTrial::Spec { t0, gen0, rounds }
3298 }
3299 }
3300 SpecTrial::Decided { spec: true, .. } => {
3301 if mon.pays() {
3302 mon.fails = 0;
3303 trial
3304 } else {
3305 mon.fails += 1;
3306 if mon.fails >= 4 {
3307 tracing::info!(
3308 "speculation stopped: {:.2} tok/round in {:.1} ms vs plain {:.1} ms/tok",
3309 mon.tokens,
3310 mon.round_ms,
3311 mon.plain_ms
3312 );
3313 SpecTrial::Decided {
3314 spec: false,
3315 recheck_at: generated + 128,
3316 }
3317 } else {
3318 trial
3319 }
3320 }
3321 }
3322 other => other,
3323 }
3324 }
3325
3326 fn mtp_kv_id(&self) -> u64 {
3329 self.graph_kv_id | (1u64 << 40)
3330 }
3331
3332 const MTP_LAYER_BASE: usize = 0;
3337
3338 fn mtp_block_input(&mut self, m: &MtpModule, hidden: &[f32], next_token: u32) -> Vec<f32> {
3341 let e = self.embed_single(next_token);
3342 let mut cat = vec![0.0f32; 2 * self.hidden_size];
3343 let (cat_e, cat_h) = cat.split_at_mut(self.hidden_size);
3344 inference::rms_norm_into(&e, &m.enorm, self.rms_eps, self.norm_style, cat_e);
3345 inference::rms_norm_into(hidden, &m.hnorm, self.rms_eps, self.norm_style, cat_h);
3346 let mut x = vec![0.0f32; self.hidden_size];
3347 m.eh_proj.matvec(&cat, &mut x, self.pool.as_deref());
3348 x
3349 }
3350
3351 #[cfg(feature = "gpu")]
3354 fn mtp_graph_ok(&self, m: &MtpModule) -> bool {
3355 if std::env::var("CMF_MTP_GRAPH").as_deref() == Ok("0") {
3356 return false;
3357 }
3358 if !crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Decode)
3359 || !crate::gpu::enabled_here()
3360 || self.attn_softcap > 0.0
3361 || self.attention_heads_per_layer.is_some()
3362 {
3363 return false;
3364 }
3365 matches!(
3366 &m.layer.attn,
3367 AttnKind::Full {
3368 softplus_gate: None,
3369 ..
3370 }
3371 ) && matches!(&m.layer.ffn, FfnKind::Dense(_))
3372 }
3373
3374 #[cfg(feature = "gpu")]
3380 fn mtp_step_graph(
3381 &mut self,
3382 m: &mut MtpModule,
3383 hidden: &[f32],
3384 next_token: u32,
3385 position: usize,
3386 ) -> Option<(Vec<f32>, Vec<f32>)> {
3387 if !self.mtp_graph_ok(m) {
3388 return None;
3389 }
3390 let lw = &m.layer;
3391 let AttnKind::Full {
3392 wq,
3393 wk,
3394 wv,
3395 wo,
3396 q_norm,
3397 k_norm,
3398 output_gate,
3399 softplus_gate,
3400 bias,
3401 } = &lw.attn
3402 else {
3403 return None;
3404 };
3405 if softplus_gate.is_some() {
3406 return None;
3407 }
3408 let FfnKind::Dense(d) = &lw.ffn else {
3409 return None;
3410 };
3411 if !d.segs.is_empty() {
3412 return None; }
3414 let mut x = self.mtp_block_input(m, hidden, next_token);
3417 fn gw(t: &QTensor) -> Option<crate::gpu::GraphW<'_>> {
3418 let (_, i, kind, rs) = t.graph_weight()?;
3419 Some(crate::gpu::GraphW {
3420 idx: i,
3421 kind,
3422 row_scale: rs,
3423 data: &[],
3424 })
3425 }
3426 let (model, _, _, _) = wq.graph_weight()?;
3427 let model = model.clone();
3428 let (lm_gw, lm_rows) = {
3429 let (_, i, kind, rs) = self.weights.lm_head.graph_weight()?;
3430 (
3431 crate::gpu::GraphW {
3432 idx: i,
3433 kind,
3434 row_scale: rs,
3435 data: &[],
3436 },
3437 self.weights.lm_head.rows(),
3438 )
3439 };
3440 let layer = crate::gpu::GraphLayer {
3441 input_norm: &lw.input_norm,
3442 attn: crate::gpu::GraphAttn::Full {
3443 wq: gw(wq)?,
3444 wk: gw(wk)?,
3445 wv: gw(wv)?,
3446 wo: gw(wo)?,
3447 q_norm: q_norm.as_deref(),
3448 k_norm: k_norm.as_deref(),
3449 bias: bias
3450 .as_ref()
3451 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
3452 output_gate: *output_gate,
3453 cpu_k: m.kv.k_heads(),
3454 cpu_v: m.kv.v_heads(),
3455 },
3456 post_norm: &lw.post_norm,
3457 ffn: crate::gpu::GraphFfn::Dense {
3458 gate: gw(&d.gate_proj)?,
3459 up: gw(&d.up_proj)?,
3460 down: gw(&d.down_proj)?,
3461 },
3462 };
3463 let nh = self.num_heads;
3464 let (nkv, hd, rd) = self.layer_geom(0);
3465 let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
3466 let mut logits = Vec::new();
3467 let ok = crate::gpu::forward_token_graph(
3468 &model,
3469 self.mtp_kv_id(),
3470 std::slice::from_ref(&layer),
3471 &[None],
3472 self.o1_epoch,
3473 &self.inv_freq,
3474 &mut x,
3475 nh,
3476 nkv,
3477 hd,
3478 rd,
3479 self.hidden_size,
3480 self.intermediate_size,
3481 position,
3482 self.kv_cache.max_seq_len,
3483 gemma,
3484 self.rms_eps as f32,
3485 Some((&lm_gw, lm_rows)),
3486 &m.final_norm,
3487 &mut logits,
3488 &[],
3489 1,
3490 None,
3491 None,
3492 None,
3493 Self::MTP_LAYER_BASE,
3494 true,
3495 );
3496 if !ok {
3497 return None;
3498 }
3499 logits.resize(self.vocab_size, 0.0);
3500 Some((logits, x))
3501 }
3502
3503 #[cfg(feature = "gpu")]
3510 fn mtp_warm_graph(
3511 &mut self,
3512 m: &mut MtpModule,
3513 pairs: &[(&[f32], u32)],
3514 first_pos: usize,
3515 ) -> bool {
3516 if pairs.is_empty() || !self.mtp_graph_ok(m) {
3517 return pairs.is_empty();
3518 }
3519 let hs = self.hidden_size;
3520 let mut hiddens = Vec::with_capacity(pairs.len() * hs);
3523 for (h, t) in pairs {
3524 hiddens.extend_from_slice(&self.mtp_block_input(m, h, *t));
3525 }
3526 let lw = &m.layer;
3527 let AttnKind::Full {
3528 wq,
3529 wk,
3530 wv,
3531 wo,
3532 q_norm,
3533 k_norm,
3534 output_gate,
3535 bias,
3536 ..
3537 } = &lw.attn
3538 else {
3539 return false;
3540 };
3541 let FfnKind::Dense(d) = &lw.ffn else {
3542 return false;
3543 };
3544 if !d.segs.is_empty() {
3545 return false; }
3547 fn gw(t: &QTensor) -> Option<crate::gpu::GraphW<'_>> {
3548 let (_, i, kind, rs) = t.graph_weight()?;
3549 Some(crate::gpu::GraphW {
3550 idx: i,
3551 kind,
3552 row_scale: rs,
3553 data: &[],
3554 })
3555 }
3556 let Some((model, _, _, _)) = wq.graph_weight() else {
3557 return false;
3558 };
3559 let model = model.clone();
3560 let (Some(gwq), Some(gwk), Some(gwv), Some(gwo), Some(gg), Some(gu), Some(gd)) = (
3561 gw(wq),
3562 gw(wk),
3563 gw(wv),
3564 gw(wo),
3565 gw(&d.gate_proj),
3566 gw(&d.up_proj),
3567 gw(&d.down_proj),
3568 ) else {
3569 return false;
3570 };
3571 let layer = crate::gpu::GraphLayer {
3572 input_norm: &lw.input_norm,
3573 attn: crate::gpu::GraphAttn::Full {
3574 wq: gwq,
3575 wk: gwk,
3576 wv: gwv,
3577 wo: gwo,
3578 q_norm: q_norm.as_deref(),
3579 k_norm: k_norm.as_deref(),
3580 bias: bias
3581 .as_ref()
3582 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
3583 output_gate: *output_gate,
3584 cpu_k: m.kv.k_heads(),
3585 cpu_v: m.kv.v_heads(),
3586 },
3587 post_norm: &lw.post_norm,
3588 ffn: crate::gpu::GraphFfn::Dense {
3589 gate: gg,
3590 up: gu,
3591 down: gd,
3592 },
3593 };
3594 let positions: Vec<usize> = (first_pos..first_pos + pairs.len()).collect();
3595 let nh = self.num_heads;
3596 let (nkv, hd, rd) = self.layer_geom(0);
3597 let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
3598 crate::gpu::forward_batch_graph(
3599 &model,
3600 self.mtp_kv_id(),
3601 std::slice::from_ref(&layer),
3602 &self.inv_freq,
3603 &mut hiddens,
3604 nh,
3605 nkv,
3606 hd,
3607 rd,
3608 hs,
3609 self.intermediate_size,
3610 &positions,
3611 self.kv_cache.max_seq_len,
3612 gemma,
3613 self.rms_eps as f32,
3614 pairs.len(),
3615 None,
3616 )
3617 }
3618
3619 fn mtp_warm(&mut self, m: &mut MtpModule, hidden: &[f32], next_token: u32, position: usize) {
3623 let e = self.embed_single(next_token);
3624 let mut cat = vec![0.0f32; 2 * self.hidden_size];
3625 let (cat_e, cat_h) = cat.split_at_mut(self.hidden_size);
3626 inference::rms_norm_into(&e, &m.enorm, self.rms_eps, self.norm_style, cat_e);
3627 inference::rms_norm_into(hidden, &m.hnorm, self.rms_eps, self.norm_style, cat_h);
3628 let mut x = vec![0.0f32; self.hidden_size];
3629 m.eh_proj.matvec(&cat, &mut x, self.pool.as_deref());
3630 inference::rms_norm_into(
3631 &x,
3632 &m.layer.input_norm,
3633 self.rms_eps,
3634 self.norm_style,
3635 &mut self.ws.n1,
3636 );
3637 let attn = match &m.layer.attn {
3638 AttnKind::Full {
3639 wq,
3640 wk,
3641 wv,
3642 wo,
3643 q_norm,
3644 k_norm,
3645 output_gate,
3646 softplus_gate,
3647 bias,
3648 } => {
3649 let mut cfg = self.attn_cfg(position);
3650 cfg.q_norm = q_norm.as_deref();
3651 cfg.k_norm = k_norm.as_deref();
3652 cfg.output_gate = *output_gate;
3653 cfg.softplus_gate = softplus_gate.as_ref().map(|(g, p)| (g, *p));
3654 cfg.bias = bias
3655 .as_ref()
3656 .map(|(q, k, v)| (q.as_slice(), k.as_slice(), v.as_slice()));
3657 attention::qwen_attention(&self.ws.n1, wq, wk, wv, wo, &mut m.kv, &cfg)
3658 }
3659 _ => return,
3660 };
3661 let _ = attn;
3662 }
3663
3664 #[cfg(feature = "gpu")]
3671 #[allow(clippy::too_many_arguments)]
3672 fn graph_spec_step(
3673 &mut self,
3674 m: &mut MtpModule,
3675 hidden: &[f32],
3676 t_next: u32,
3677 next_pos: usize,
3678 drafted: &mut usize,
3679 accepted: &mut usize,
3680 all_ids: &mut Vec<u32>,
3684 ) -> Option<(Vec<u32>, usize, Vec<f32>)> {
3685 #[cfg(target_os = "macos")]
3696 let metal_native = crate::gpu::q1_force();
3697 #[cfg(not(target_os = "macos"))]
3698 let metal_native = false;
3699 #[cfg(feature = "gpu")]
3700 let k_default = if metal_native {
3701 7
3704 } else if crate::gpu_wgpu::verify_i8_on() {
3705 5
3706 } else {
3707 4
3708 };
3709 #[cfg(not(feature = "gpu"))]
3710 let k_default = 4;
3711 let k_spec: usize = std::env::var("CMF_GRAPH_SPEC_K")
3712 .ok()
3713 .and_then(|v| v.parse().ok())
3714 .filter(|&v| (1..=8).contains(&v))
3715 .unwrap_or(k_default);
3716 if next_pos == 0 {
3717 return None;
3718 }
3719 let t_round = std::time::Instant::now();
3720 let subs = || crate::gpu_wgpu::SUBMITS.load(std::sync::atomic::Ordering::Relaxed);
3736 let sub0 = subs();
3737 let cfg = self.sampler_config.clone();
3742 let penalized = !(cfg.repetition_penalty == 1.0
3743 && cfg.presence_penalty == 0.0
3744 && cfg.suppress_tokens.is_empty());
3745 let greedy_pen = cfg.temperature < 1e-6 && penalized;
3750 let sampling = cfg.temperature >= 1e-6;
3751 let sparse = sampling && sampler::sparse_ok(&cfg);
3757 let base_len = all_ids.len();
3758 if sampling && !sparse && self.spec_q.len() < k_spec {
3759 self.spec_q.resize_with(k_spec, Vec::new);
3760 }
3761 if sparse && self.spec_qs.len() < k_spec {
3762 self.spec_qs.resize_with(k_spec, Vec::new);
3763 }
3764 let mut drafts = Vec::with_capacity(k_spec);
3769 let mut hx = hidden.to_vec();
3770 let spec_dbg = std::env::var("CMF_SPEC_DBG").is_ok();
3773 for j in 0..k_spec {
3774 let tok_in = if j == 0 { t_next } else { drafts[j - 1] };
3775 let mut dbg_ref: Option<(Vec<f32>, Vec<f32>)> = None;
3776 if spec_dbg {
3777 let saved = self.mtp_graph_mode;
3778 self.mtp_graph_mode = Some(false);
3779 let r = self.mtp_step_hl(m, &hx, tok_in, next_pos - 1 + j);
3780 self.mtp_graph_mode = saved;
3781 m.kv.truncate_last(1);
3782 dbg_ref = Some(r);
3783 }
3784 let (mut lg, hj) = self.mtp_step_hl(m, &hx, tok_in, next_pos - 1 + j);
3785 if let Some((lg_cpu, h_cpu)) = dbg_ref {
3786 let n = |v: &[f32]| v.iter().map(|x| x * x).sum::<f32>().sqrt();
3787 let dl = lg
3788 .iter()
3789 .zip(&lg_cpu)
3790 .fold(0f32, |m, (a, b)| m.max((a - b).abs()));
3791 let dh = hj
3792 .iter()
3793 .zip(&h_cpu)
3794 .fold(0f32, |m, (a, b)| m.max((a - b).abs()));
3795 eprintln!(
3796 "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 {}",
3797 next_pos - 1 + j,
3798 sampler::argmax(&lg_cpu),
3799 sampler::argmax(&lg),
3800 n(&h_cpu),
3801 n(&hj),
3802 m.kv.seq_len
3803 );
3804 }
3805 let dj = if sparse {
3806 let mut q = std::mem::take(&mut self.spec_qs[j]);
3807 let ok = sampler::sparse_distribution_into(
3808 &lg,
3809 &cfg,
3810 all_ids,
3811 &mut self.sampler_scratch,
3812 self.pool.as_deref(),
3813 &mut q,
3814 );
3815 let d = if ok {
3816 sampler::draw_sparse(&q, &mut self.rng)
3817 } else {
3818 let t = sampler::argmax(&lg);
3820 q.clear();
3821 q.push((t, 1.0));
3822 t
3823 };
3824 self.spec_qs[j] = q;
3825 all_ids.push(d);
3826 d
3827 } else if sampling {
3828 let mut q = std::mem::take(&mut self.spec_q[j]);
3829 sampler::distribution_into(
3830 &lg,
3831 &cfg,
3832 all_ids,
3833 &mut self.sampler_scratch,
3834 self.pool.as_deref(),
3835 &mut q,
3836 );
3837 let d = sampler::draw(&q, &mut self.rng);
3838 self.spec_q[j] = q;
3839 all_ids.push(d); d
3841 } else if greedy_pen {
3842 let d = sampler::argmax_penalized(
3843 &lg,
3844 &cfg,
3845 all_ids,
3846 &mut self.sampler_scratch,
3847 self.pool.as_deref(),
3848 );
3849 all_ids.push(d);
3850 d
3851 } else {
3852 sampler::argmax(&lg)
3853 };
3854 attention::recycle_buf(&mut lg);
3855 drafts.push(dj);
3856 hx = hj;
3857 }
3858 all_ids.truncate(base_len);
3859 *drafted += k_spec;
3860 let t_draft = t_round.elapsed();
3861 let sub_draft = subs();
3862 let b = k_spec + 1;
3865 let mut hiddens = vec![0.0f32; b * self.hidden_size];
3866 for (i, &t) in std::iter::once(&t_next).chain(drafts.iter()).enumerate() {
3867 let e = self.embed_single(t);
3868 hiddens[i * self.hidden_size..(i + 1) * self.hidden_size].copy_from_slice(&e);
3869 }
3870 let positions: Vec<usize> = (next_pos..next_pos + b).collect();
3871 let (lm_gw, lm_rows) = {
3872 let (_, i, kind, rs) = self.weights.lm_head.graph_weight()?;
3873 (
3874 crate::gpu::GraphW {
3875 idx: i,
3876 kind,
3877 row_scale: rs,
3878 data: &[],
3879 },
3880 self.weights.lm_head.rows(),
3881 )
3882 };
3883 let mut logits = Vec::new();
3884 let final_norm = self.weights.final_norm.clone();
3885 #[cfg(target_os = "macos")]
3886 let ok = if metal_native {
3887 let lm = self.weights.lm_head.q1_parts()?;
3888 self.try_batch_graph_metal(
3889 &mut hiddens,
3890 &positions,
3891 b,
3892 Some((lm, &final_norm, &mut logits)),
3893 )
3894 } else {
3895 self.try_batch_graph_wgpu(
3896 &mut hiddens,
3897 &positions,
3898 b,
3899 Some(crate::gpu::SpecTail {
3900 lm: lm_gw,
3901 lm_rows,
3902 final_norm: &final_norm,
3903 logits_out: &mut logits,
3904 }),
3905 )
3906 };
3907 #[cfg(not(target_os = "macos"))]
3908 let ok = self.try_batch_graph_wgpu(
3909 &mut hiddens,
3910 &positions,
3911 b,
3912 Some(crate::gpu::SpecTail {
3913 lm: lm_gw,
3914 lm_rows,
3915 final_norm: &final_norm,
3916 logits_out: &mut logits,
3917 }),
3918 );
3919 if !ok {
3920 m.kv.truncate_last(k_spec);
3923 return None;
3924 }
3925 #[cfg(target_os = "macos")]
3931 if metal_native && std::env::var("CMF_METAL_VERIFY_CHECK").as_deref() == Ok("1") {
3932 let snap: Vec<Vec<f32>> = self
3933 .kv_cache
3934 .layers
3935 .iter()
3936 .map(|l| l.linear_state.clone())
3937 .collect();
3938 let attn_lens: Vec<usize> = self.kv_cache.layers.iter().map(|l| l.seq_len).collect();
3939 let toks: Vec<u32> = std::iter::once(t_next)
3940 .chain(drafts.iter().copied())
3941 .collect();
3942 let want_save = self.graph_want_logits;
3943 self.graph_want_logits = false;
3944 for (i, &t) in toks.iter().enumerate() {
3945 let hi = self.forward_layers(&self.embed_single(t), next_pos + i, None);
3946 let _ = self.graph_logits.take();
3947 if std::env::var("CMF_SPEC_PLAIN_HIDDEN").as_deref() == Ok("1") {
3951 hiddens[i * self.hidden_size..(i + 1) * self.hidden_size].copy_from_slice(&hi);
3952 }
3953 let ref_lg = self.logits_from_hidden(&hi);
3954 let row = &logits[i * lm_rows..(i + 1) * lm_rows];
3955 let ra = sampler::argmax(&ref_lg);
3956 let va = sampler::argmax(row);
3957 let mut md = 0f32;
3958 let mut rms = 0f64;
3959 for j in 0..lm_rows.min(ref_lg.len()) {
3960 let d = (ref_lg[j] - row[j]).abs();
3961 md = md.max(d);
3962 rms += (d as f64) * (d as f64);
3963 }
3964 let mut hd = 0f32;
3965 for j in 0..self.hidden_size {
3966 hd = hd.max((hi[j] - hiddens[i * self.hidden_size + j]).abs());
3967 }
3968 eprintln!(
3969 "verify-check row {i} tok {t} pos {}: ref argmax {ra} verify argmax {va} {} | max|dlogit| {md:.3} rms {:.4} | max|dhidden| {hd:.4}",
3970 next_pos + i,
3971 if ra == va { "OK" } else { "MISMATCH" },
3972 (rms / lm_rows as f64).sqrt()
3973 );
3974 }
3975 self.graph_want_logits = want_save;
3976 for (l, st) in self.kv_cache.layers.iter_mut().zip(snap) {
3979 if l.linear_state.len() == st.len() {
3980 l.linear_state.copy_from_slice(&st);
3981 } else {
3982 l.linear_state = st;
3983 }
3984 }
3985 for (li, (l, n0)) in self.kv_cache.layers.iter_mut().zip(attn_lens).enumerate() {
3986 let extra = l.seq_len.saturating_sub(n0);
3987 if extra > 0 {
3988 l.truncate_last(extra);
3989 crate::gpu_metal::kv_mirror_set_stored(self.graph_kv_id, li, n0);
3990 }
3991 }
3992 }
3993 let t_verify = t_round.elapsed();
3994 let sub_verify = subs();
3995 let mut a = 0usize;
4000 let mut forced: Option<u32> = None;
4001 let ids: Vec<u32> = if sparse {
4002 let mut p = std::mem::take(&mut self.spec_ps);
4003 let mut res = std::mem::take(&mut self.spec_ress);
4004 while a < k_spec {
4005 let ok = sampler::sparse_distribution_into(
4006 &logits[a * lm_rows..(a + 1) * lm_rows],
4007 &cfg,
4008 all_ids,
4009 &mut self.sampler_scratch,
4010 self.pool.as_deref(),
4011 &mut p,
4012 );
4013 if !ok {
4014 let t = sampler::argmax(&logits[a * lm_rows..(a + 1) * lm_rows]);
4015 p.clear();
4016 p.push((t, 1.0));
4017 }
4018 match sampler::spec_accept_or_correct_sparse(
4019 &p,
4020 &self.spec_qs[a],
4021 drafts[a],
4022 &mut self.rng,
4023 &mut res,
4024 ) {
4025 None => {
4026 all_ids.push(drafts[a]);
4027 a += 1;
4028 }
4029 Some(c) => {
4030 forced = Some(c);
4031 break;
4032 }
4033 }
4034 }
4035 all_ids.truncate(base_len);
4036 self.spec_ps = p;
4037 self.spec_ress = res;
4038 drafts.clone()
4039 } else if sampling {
4040 let mut p = std::mem::take(&mut self.spec_p);
4041 let mut res = std::mem::take(&mut self.spec_res);
4042 while a < k_spec {
4043 sampler::distribution_into(
4044 &logits[a * lm_rows..(a + 1) * lm_rows],
4045 &cfg,
4046 all_ids,
4047 &mut self.sampler_scratch,
4048 self.pool.as_deref(),
4049 &mut p,
4050 );
4051 match sampler::spec_accept_or_correct(
4052 &p,
4053 &self.spec_q[a],
4054 drafts[a],
4055 &mut self.rng,
4056 &mut res,
4057 self.pool.as_deref(),
4058 ) {
4059 None => {
4060 all_ids.push(drafts[a]);
4061 a += 1;
4062 }
4063 Some(c) => {
4064 forced = Some(c);
4065 break;
4066 }
4067 }
4068 }
4069 all_ids.truncate(base_len);
4070 self.spec_p = p;
4071 self.spec_res = res;
4072 drafts.clone()
4074 } else if greedy_pen {
4075 let mut ids: Vec<u32> = Vec::with_capacity(b);
4079 for i in 0..b {
4080 let t = sampler::argmax_penalized(
4081 &logits[i * lm_rows..(i + 1) * lm_rows],
4082 &cfg,
4083 all_ids,
4084 &mut self.sampler_scratch,
4085 self.pool.as_deref(),
4086 );
4087 ids.push(t);
4088 if i < k_spec && t == drafts[i] {
4089 all_ids.push(t);
4090 } else {
4091 break;
4092 }
4093 }
4094 all_ids.truncate(base_len);
4095 while a < k_spec && a < ids.len() && ids[a] == drafts[a] {
4096 a += 1;
4097 }
4098 ids
4101 } else {
4102 let ids: Vec<u32> = (0..b)
4103 .map(|i| sampler::argmax(&logits[i * lm_rows..(i + 1) * lm_rows]))
4104 .collect();
4105 while a < k_spec && ids[a] == drafts[a] {
4106 a += 1;
4107 }
4108 ids
4109 };
4110 if spec_dbg {
4111 eprintln!(
4112 "spec-dbg round: t_next {t_next} drafts {:?} verified {:?} accepted {a}",
4113 drafts, ids
4114 );
4115 }
4116 #[cfg(target_os = "macos")]
4120 let commit_ref: Option<(Vec<Vec<f32>>, Vec<(usize, Vec<f32>, Vec<f32>)>)> = if metal_native
4121 && std::env::var("CMF_METAL_VERIFY_CHECK").as_deref() == Ok("2")
4122 {
4123 let snap: Vec<Vec<f32>> = self
4124 .kv_cache
4125 .layers
4126 .iter()
4127 .map(|l| l.linear_state.clone())
4128 .collect();
4129 let attn_lens: Vec<usize> = self.kv_cache.layers.iter().map(|l| l.seq_len).collect();
4130 let toks: Vec<u32> = std::iter::once(t_next)
4131 .chain(drafts.iter().copied())
4132 .collect();
4133 let want_save = self.graph_want_logits;
4134 self.graph_want_logits = false;
4135 for (i, &t) in toks.iter().take(a + 1).enumerate() {
4136 let _ = self.forward_layers(&self.embed_single(t), next_pos + i, None);
4137 let _ = self.graph_logits.take();
4138 }
4139 self.graph_want_logits = want_save;
4140 let plain_states: Vec<Vec<f32>> = self
4141 .kv_cache
4142 .layers
4143 .iter()
4144 .map(|l| l.linear_state.clone())
4145 .collect();
4146 let (nkv, hd) = (self.num_kv_heads, self.head_dim);
4147 let mut rows = Vec::new();
4148 for (li, (l, n0)) in self
4149 .kv_cache
4150 .layers
4151 .iter_mut()
4152 .zip(attn_lens.iter())
4153 .enumerate()
4154 {
4155 let extra = l.seq_len.saturating_sub(*n0);
4156 if extra > 0 {
4157 let mut kk = Vec::new();
4158 let mut vv = Vec::new();
4159 for g in 0..nkv {
4160 kk.extend_from_slice(&l.head_keys(g)[n0 * hd..]);
4161 vv.extend_from_slice(&l.head_values(g)[n0 * hd..]);
4162 }
4163 rows.push((li, kk, vv));
4164 l.truncate_last(extra);
4165 crate::gpu_metal::kv_mirror_set_stored(self.graph_kv_id, li, *n0);
4166 }
4167 }
4168 for (l, st) in self.kv_cache.layers.iter_mut().zip(snap) {
4169 if l.linear_state.len() == st.len() {
4170 l.linear_state.copy_from_slice(&st);
4171 } else {
4172 l.linear_state = st;
4173 }
4174 }
4175 Some((plain_states, rows))
4176 } else {
4177 None
4178 };
4179 #[cfg(target_os = "macos")]
4181 if metal_native {
4182 self.metal_verify_commit(a);
4185 if let Some((plain_states, rows)) = commit_ref {
4186 crate::gpu_metal::queue_fence();
4187 let (nkv, hd) = (self.num_kv_heads, self.head_dim);
4188 let mut worst_s = 0f32;
4189 let mut worst_li = 0usize;
4190 for (li, (l, ps)) in self.kv_cache.layers.iter().zip(&plain_states).enumerate() {
4191 if l.linear_state.len() != ps.len() || ps.is_empty() {
4192 continue;
4193 }
4194 let d = l
4195 .linear_state
4196 .iter()
4197 .zip(ps)
4198 .fold(0f32, |m, (x, y)| m.max((x - y).abs()));
4199 let n = ps.iter().fold(0f32, |m, y| m.max(y.abs()));
4200 let rel = d / n.max(1e-6);
4201 if rel > worst_s {
4202 worst_s = rel;
4203 worst_li = li;
4204 }
4205 }
4206 let mut worst_k = 0f32;
4207 for (li, kk, vv) in &rows {
4208 let l = &self.kv_cache.layers[*li];
4209 let n0 = l.seq_len - (kk.len() / (nkv * hd));
4210 let mut ck = Vec::new();
4211 let mut cv = Vec::new();
4212 for g in 0..nkv {
4213 ck.extend_from_slice(&l.head_keys(g)[n0 * hd..]);
4214 cv.extend_from_slice(&l.head_values(g)[n0 * hd..]);
4215 }
4216 if ck.len() == kk.len() {
4217 let dk = ck
4218 .iter()
4219 .zip(kk)
4220 .fold(0f32, |m, (x, y)| m.max((x - y).abs()));
4221 let dv = cv
4222 .iter()
4223 .zip(vv)
4224 .fold(0f32, |m, (x, y)| m.max((x - y).abs()));
4225 worst_k = worst_k.max(dk).max(dv);
4226 } else {
4227 eprintln!(
4228 "commit-check L{li}: kv row count mismatch {} vs {}",
4229 ck.len(),
4230 kk.len()
4231 );
4232 }
4233 }
4234 eprintln!(
4235 "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}"
4236 );
4237 }
4238 } else if a + 1 < b {
4239 crate::gpu::gdn_spec_restore(self.graph_kv_id, a);
4240 }
4241 #[cfg(not(target_os = "macos"))]
4242 if a + 1 < b {
4243 crate::gpu::gdn_spec_restore(self.graph_kv_id, a);
4244 }
4245 *accepted += a;
4246 m.kv.truncate_last(k_spec.saturating_sub(1));
4257 #[cfg(target_os = "macos")]
4258 if metal_native && self.mtp_graph_mode == Some(true) {
4259 crate::gpu_metal::kv_mirror_set_stored(
4262 self.mtp_kv_id(),
4263 Self::MTP_LAYER_BASE,
4264 m.kv.seq_len,
4265 );
4266 }
4267 let warm_off = std::env::var("CMF_SPEC_WARM").is_ok_and(|v| v == "0");
4268 if !warm_off && a > 0 {
4269 let mut warmed = false;
4272 #[cfg(target_os = "macos")]
4273 if metal_native && self.mtp_graph_mode == Some(true) {
4274 let pairs: Vec<(&[f32], u32)> = (0..a)
4278 .map(|j| {
4279 (
4280 &hiddens[j * self.hidden_size..(j + 1) * self.hidden_size],
4281 ids[j],
4282 )
4283 })
4284 .collect();
4285 warmed = self.mtp_warm_batch_metal(m, &pairs, next_pos);
4286 if !warmed {
4287 warmed = true;
4288 for j in 0..a {
4289 let row =
4290 hiddens[j * self.hidden_size..(j + 1) * self.hidden_size].to_vec();
4291 if self
4292 .mtp_step_metal(m, &row, ids[j], next_pos + j, false)
4293 .is_none()
4294 {
4295 warmed = false;
4296 break;
4297 }
4298 }
4299 }
4300 }
4301 if !warmed && self.mtp_graph_mode == Some(true) && !metal_native {
4302 let rows: Vec<Vec<f32>> = (0..a)
4303 .map(|j| hiddens[j * self.hidden_size..(j + 1) * self.hidden_size].to_vec())
4304 .collect();
4305 let pairs: Vec<(&[f32], u32)> = rows
4306 .iter()
4307 .zip(ids.iter())
4308 .map(|(r, &t)| (r.as_slice(), t))
4309 .collect();
4310 warmed = self.mtp_warm_graph(m, &pairs, next_pos);
4311 if !warmed {
4312 warmed = true;
4314 for j in 0..a {
4315 if self
4316 .mtp_step_graph(m, &rows[j], ids[j], next_pos + j)
4317 .is_none()
4318 {
4319 warmed = false;
4320 break;
4321 }
4322 }
4323 }
4324 }
4325 if !warmed {
4326 for j in 0..a {
4327 let row = &hiddens[j * self.hidden_size..(j + 1) * self.hidden_size];
4328 let row = row.to_vec();
4329 self.mtp_warm(m, &row, ids[j], next_pos + j);
4330 }
4331 }
4332 }
4333 if let Some(c) = forced {
4337 self.spec_forced = Some(c);
4338 self.graph_logits = None;
4339 } else {
4340 let mut row = logits[a * lm_rows..(a + 1) * lm_rows].to_vec();
4341 row.resize(self.vocab_size, 0.0);
4342 if let Some(c) = self.final_softcap {
4343 for l in row.iter_mut() {
4344 *l = c * (*l / c).tanh();
4345 }
4346 }
4347 self.graph_logits = Some(row);
4348 }
4349 let new_hidden = hiddens[a * self.hidden_size..(a + 1) * self.hidden_size].to_vec();
4350 if std::env::var("CMF_GRAPH_SPEC_TIME").is_ok() {
4356 let end = subs();
4357 eprintln!(
4358 "spec-round: draft {:.1} ms/{} sub | verify {:.1} ms/{} sub | \
4359 commit {:.1} ms/{} sub (accepted {a} of {k_spec})",
4360 t_draft.as_secs_f64() * 1e3,
4361 sub_draft - sub0,
4362 (t_verify - t_draft).as_secs_f64() * 1e3,
4363 sub_verify - sub_draft,
4364 (t_round.elapsed() - t_verify).as_secs_f64() * 1e3,
4365 end - sub_verify,
4366 );
4367 }
4368 Some((drafts[..a].to_vec(), next_pos + a + 1, new_hidden))
4369 }
4370
4371 pub fn measure_pair_fusion(&mut self, iters: usize) -> (f64, f64) {
4380 if !self.pair_supported() {
4381 return (0.0, 0.0);
4382 }
4383 let emb1 = self.embed_single(1);
4384 let emb2 = self.embed_single(2);
4385 let pos = self.kv_cache.seq_len();
4386
4387 let t0 = std::time::Instant::now();
4388 for _ in 0..iters {
4389 let _ = self.forward_layers(&emb1, pos, None);
4390 let _ = self.forward_layers(&emb2, pos + 1, None);
4391 for l in &mut self.kv_cache.layers {
4392 l.truncate_last(2);
4393 }
4394 }
4395 let singles_ms = t0.elapsed().as_secs_f64() * 1000.0 / iters as f64;
4396
4397 let t1 = std::time::Instant::now();
4398 for _ in 0..iters {
4399 let _ = self.forward_pair(&emb1, &emb2, pos);
4400 for l in &mut self.kv_cache.layers {
4401 l.truncate_last(2);
4402 }
4403 }
4404 let pair_ms = t1.elapsed().as_secs_f64() * 1000.0 / iters as f64;
4405 (singles_ms, pair_ms)
4406 }
4407
4408 fn pair_supported(&self) -> bool {
4416 !self.weights.layers.is_empty()
4423 && self.g3n.is_none()
4424 && !self
4425 .weights
4426 .layers
4427 .iter()
4428 .any(|lw| matches!(&lw.attn, AttnKind::Mla(_) | AttnKind::Kda(_)))
4429 }
4430
4431 fn forward_pair(
4432 &mut self,
4433 emb1: &[f32],
4434 emb2: &[f32],
4435 position: usize,
4436 ) -> (Vec<f32>, Vec<f32>) {
4437 let mut h1 = emb1.to_vec();
4438 let mut h2 = emb2.to_vec();
4439 let (_nkv, _hd, hs, _rd, eps) = (
4440 self.num_kv_heads,
4441 self.head_dim,
4442 self.hidden_size,
4443 self.rotary_dim,
4444 self.rms_eps,
4445 );
4446 let pool = self.pool.clone();
4447
4448 for li in 0..self.num_layers {
4449 let lw = &self.weights.layers[self.phys_layer(li)];
4450 inference::rms_norm_into(
4453 &h1,
4454 &lw.input_norm,
4455 self.rms_eps,
4456 self.norm_style,
4457 &mut self.ws.n1,
4458 );
4459 inference::rms_norm_into(
4460 &h2,
4461 &lw.input_norm,
4462 self.rms_eps,
4463 self.norm_style,
4464 &mut self.ws.n2,
4465 );
4466
4467 let (a1, a2) = match &lw.attn {
4468 AttnKind::Mla(_) => unreachable!("MLA has no MTP/pair path"),
4469 AttnKind::Kda(_) => unreachable!("KDA has no MTP/pair path"),
4470 AttnKind::Linear(w) => {
4471 let cfg = self.vmf_cfg.expect("linear layer without vmf_cfg");
4472 let layer = &mut self.kv_cache.layers[li];
4473 let (state, scratch) = (&mut layer.linear_state, &mut layer.linear_scratch);
4474 vmf_phase_pair(
4475 &self.ws.n1,
4476 &self.ws.n2,
4477 w,
4478 &cfg,
4479 state,
4480 scratch,
4481 self.pool.as_deref(),
4482 )
4483 }
4484 AttnKind::LinearGdn(w) => {
4485 let cfg = self.gdn_cfg.expect("gdn layer without gdn_cfg");
4486 let layer = &mut self.kv_cache.layers[li];
4487 let (state, scratch) = (&mut layer.linear_state, &mut layer.linear_scratch);
4488 gdn_pair(
4489 &self.ws.n1,
4490 &self.ws.n2,
4491 w,
4492 &cfg,
4493 state,
4494 scratch,
4495 self.pool.as_deref(),
4496 )
4497 }
4498 AttnKind::ShortConv(w) => {
4499 let cfg = self
4500 .short_conv_cfg
4501 .expect("short-conv layer without short_conv_cfg");
4502 let layer = &mut self.kv_cache.layers[li];
4503 let (state, scratch) = (&mut layer.linear_state, &mut layer.linear_scratch);
4504 short_conv_pair(
4505 &self.ws.n1,
4506 &self.ws.n2,
4507 w,
4508 &cfg,
4509 state,
4510 scratch,
4511 self.pool.as_deref(),
4512 )
4513 }
4514 AttnKind::Full {
4515 wq,
4516 wk,
4517 wv,
4518 wo,
4519 q_norm,
4520 k_norm,
4521 output_gate,
4522 softplus_gate,
4523 bias,
4524 } => {
4525 let inv_freq_l = self.layer_inv_freq(li);
4526 let (nkv_l, hd_l, rd_l) = self.layer_geom(li);
4527 let cfg = QwenAttnCfg {
4528 num_heads: self.layer_num_heads(li),
4529 num_kv_heads: nkv_l,
4530 head_dim: hd_l,
4531 hidden_size: hs,
4532 position,
4533 inv_freq: &inv_freq_l,
4534 rotary_dim: rd_l,
4535 scale: self.attn_scale,
4536 softcap: self.attn_softcap,
4537 window: self.layer_window(li),
4538 v_norm: self.attn_v_norm,
4539 q_norm: q_norm.as_deref(),
4540 k_norm: k_norm.as_deref(),
4541 output_gate: *output_gate,
4542 softplus_gate: softplus_gate
4543 .as_ref()
4544 .map(|(gate, per_head)| (gate, *per_head)),
4545 rope_scale: self.layer_rope_scale(li),
4546 bias: bias
4547 .as_ref()
4548 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
4549 rms_eps: eps,
4550 norm_style: self.norm_style,
4551 pool: pool.as_deref(),
4552 };
4553 attention::qwen_attention_pair(
4554 &self.ws.n1,
4555 &self.ws.n2,
4556 wq,
4557 wk,
4558 wv,
4559 wo,
4560 &mut self.kv_cache.layers[li],
4561 &cfg,
4562 )
4563 }
4564 };
4565 let (a1, a2) = match &self.weights.layers[self.phys_layer(li)].attn_out_norm {
4566 Some(w) => (
4567 inference::rms_norm(&a1, w, self.rms_eps, self.norm_style),
4568 inference::rms_norm(&a2, w, self.rms_eps, self.norm_style),
4569 ),
4570 None => (a1, a2),
4571 };
4572 for i in 0..self.hidden_size {
4573 h1[i] += a1[i];
4574 h2[i] += a2[i];
4575 }
4576 let (mut a1, mut a2) = (a1, a2);
4577 attention::recycle_buf(&mut a1);
4578 attention::recycle_buf(&mut a2);
4579
4580 let lw = &self.weights.layers[self.phys_layer(li)];
4581 inference::rms_norm_into(
4582 &h1,
4583 &lw.post_norm,
4584 self.rms_eps,
4585 self.norm_style,
4586 &mut self.ws.p1,
4587 );
4588 inference::rms_norm_into(
4589 &h2,
4590 &lw.post_norm,
4591 self.rms_eps,
4592 self.norm_style,
4593 &mut self.ws.p2,
4594 );
4595 let (f1, f2) = match &lw.ffn {
4596 FfnKind::DenseMoe(dm) => (
4599 dense_moe_ffn(
4600 dm,
4601 &self.ws.p1,
4602 &h1,
4603 self.rms_eps,
4604 self.norm_style,
4605 self.pool.as_deref(),
4606 ),
4607 dense_moe_ffn(
4608 dm,
4609 &self.ws.p2,
4610 &h2,
4611 self.rms_eps,
4612 self.norm_style,
4613 self.pool.as_deref(),
4614 ),
4615 ),
4616 _ => ffn_forward_pair(
4617 &lw.ffn,
4618 &self.ws.p1,
4619 &self.ws.p2,
4620 self.pool.as_deref(),
4621 None,
4622 ),
4623 };
4624 let (f1, f2) = match &self.weights.layers[self.phys_layer(li)].ffn_out_norm {
4625 Some(w) => (
4626 inference::rms_norm(&f1, w, self.rms_eps, self.norm_style),
4627 inference::rms_norm(&f2, w, self.rms_eps, self.norm_style),
4628 ),
4629 None => (f1, f2),
4630 };
4631 for i in 0..self.hidden_size {
4632 h1[i] += f1[i];
4633 h2[i] += f2[i];
4634 }
4635 let (mut f1, mut f2) = (f1, f2);
4636 attention::recycle_buf(&mut f1);
4637 attention::recycle_buf(&mut f2);
4638 if let Some(sc) = self.weights.layers[self.phys_layer(li)].layer_scale {
4639 for i in 0..self.hidden_size {
4640 h1[i] *= sc;
4641 h2[i] *= sc;
4642 }
4643 }
4644 if self.is_loop_end(li) && li + 1 < self.num_layers {
4646 h1 = inference::rms_norm(
4647 &h1,
4648 &self.weights.final_norm,
4649 self.rms_eps,
4650 self.norm_style,
4651 );
4652 h2 = inference::rms_norm(
4653 &h2,
4654 &self.weights.final_norm,
4655 self.rms_eps,
4656 self.norm_style,
4657 );
4658 }
4659 }
4660 (h1, h2)
4661 }
4662
4663 fn commit_linear_scratch(&mut self) {
4665 for layer in &mut self.kv_cache.layers {
4666 if !layer.linear_scratch.is_empty() {
4667 std::mem::swap(&mut layer.linear_state, &mut layer.linear_scratch);
4668 layer.linear_scratch.clear();
4669 }
4670 }
4671 }
4672
4673 pub fn forward_ids(
4676 &mut self,
4677 ids: &[u32],
4678 task_mask: Option<&TaskMask>,
4679 ) -> Result<Vec<f32>, String> {
4680 if ids.is_empty() {
4681 return Err("empty id sequence".to_string());
4682 }
4683 self.kv_cache.clear();
4684 self.kv_history.clear();
4685 self.o1_begin();
4686 let mut hidden = vec![0.0f32; self.hidden_size];
4687 let mut pos = 0usize;
4688 if self.can_prefill_batched() && !self.graph_prefill_preferred() && ids.len() > 2 {
4696 let chunk = prefill_chunk();
4700 let hs = self.hidden_size;
4701 while pos < ids.len() {
4702 let end = (pos + chunk).min(ids.len());
4703 let hb = self.prefill_batch_masked(&ids[pos..end], pos, task_mask);
4704 hidden.copy_from_slice(&hb[(end - pos - 1) * hs..]);
4705 pos = end;
4706 }
4707 }
4708 if task_mask.is_none()
4717 && !self.graph_prefill_preferred()
4718 && !std::env::var("CMF_PAIR").is_ok_and(|v| v == "0")
4719 && self.pair_supported()
4720 {
4721 while pos + 1 < ids.len() {
4722 let e1 = self.embed_single(ids[pos]);
4723 let e2 = self.embed_single(ids[pos + 1]);
4724 let (_, h2) = self.forward_pair(&e1, &e2, pos);
4725 self.commit_linear_scratch();
4726 hidden = h2;
4727 pos += 2;
4728 }
4729 }
4730 while pos < ids.len() {
4731 hidden = self.forward_layers(&self.embed_single(ids[pos]), pos, task_mask);
4732 pos += 1;
4733 }
4734 self.o1_seal();
4738 let normed = inference::rms_norm(
4739 &hidden,
4740 &self.weights.final_norm,
4741 self.rms_eps,
4742 self.norm_style,
4743 );
4744 Ok(self.lm_head_forward(&normed))
4745 }
4746
4747 pub fn ppl_ids(&mut self, ids: &[u32]) -> f64 {
4754 let (nll, cnt) = self.nll_ids_from(ids, 0);
4755 (nll / cnt.max(1) as f64).exp()
4756 }
4757
4758 pub fn probe_ffn_mass(&mut self, ids: &[u32]) -> Vec<Vec<f64>> {
4763 self.kv_cache.clear();
4764 self.kv_history.clear();
4765 FFN_PROBE.with(|p| {
4766 *p.borrow_mut() = Some(vec![vec![0f64; self.intermediate_size]; self.num_layers]);
4767 });
4768 crate::gpu::cpu_scope(|| {
4769 for (pos, &id) in ids.iter().enumerate() {
4770 let emb = self.embed_single(id);
4771 let _ = self.forward_layers(&emb, pos, None);
4772 }
4773 });
4774 self.kv_cache.clear();
4775 self.kv_history.clear();
4776 FFN_PROBE
4777 .with(|p| p.borrow_mut().take())
4778 .unwrap_or_default()
4779 }
4780
4781 pub fn probe_ffn_mass_batch(&mut self, ids: &[u32]) -> Vec<Vec<f64>> {
4785 self.kv_cache.clear();
4786 self.kv_history.clear();
4787 FFN_PROBE.with(|p| {
4788 *p.borrow_mut() = Some(vec![vec![0f64; self.intermediate_size]; self.num_layers]);
4789 });
4790 for chunk in ids.chunks(256) {
4791 if chunk.len() < 2 {
4792 continue;
4793 }
4794 let _ = self.nll_ids_masked(chunk, 0, None);
4795 }
4796 self.kv_cache.clear();
4797 self.kv_history.clear();
4798 FFN_PROBE
4799 .with(|p| p.borrow_mut().take())
4800 .unwrap_or_default()
4801 }
4802
4803 pub fn ppl_ids_masked(&mut self, ids: &[u32], mask: &TaskMask) -> f64 {
4807 self.kv_cache.clear();
4808 self.kv_history.clear();
4809 let mut nll = 0f64;
4810 let mut cnt = 0usize;
4811 let mut hidden = vec![0f32; self.hidden_size];
4812 for (pos, &id) in ids.iter().enumerate() {
4813 if pos > 0 {
4814 inference::rms_norm_into(
4815 &hidden,
4816 &self.weights.final_norm,
4817 self.rms_eps,
4818 self.norm_style,
4819 &mut self.ws.n1,
4820 );
4821 let mut logits = self.lm_head_forward(&self.ws.n1);
4822 let max = logits.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
4823 let sum: f64 = logits.iter().map(|&v| ((v - max) as f64).exp()).sum();
4824 let p = ((logits[id as usize] - max) as f64).exp() / sum.max(1e-300);
4825 nll -= p.max(1e-300).ln();
4826 cnt += 1;
4827 attention::recycle_buf(&mut logits);
4828 }
4829 let emb = self.embed_single(id);
4830 hidden = self.forward_layers(&emb, pos, Some(mask));
4831 }
4832 self.kv_cache.clear();
4833 self.kv_history.clear();
4834 (nll / cnt.max(1) as f64).exp()
4835 }
4836
4837 pub fn nll_ids_masked(
4856 &mut self,
4857 ids: &[u32],
4858 start: usize,
4859 task_mask: Option<&TaskMask>,
4860 ) -> (f64, usize) {
4861 let task_mask = self.drop_open_mask(task_mask);
4862 self.nll_ids_inner(ids, start, task_mask)
4863 }
4864
4865 pub fn nll_ids_from(&mut self, ids: &[u32], start: usize) -> (f64, usize) {
4866 self.nll_ids_inner(ids, start, None)
4867 }
4868
4869 fn nll_ids_inner(
4870 &mut self,
4871 ids: &[u32],
4872 start: usize,
4873 task_mask: Option<&TaskMask>,
4874 ) -> (f64, usize) {
4875 self.kv_cache.clear();
4876 self.kv_history.clear();
4877 let mut nll = 0f64;
4878 let mut cnt = 0usize;
4879 if self.can_prefill_batched() {
4880 const CHUNK: usize = 128;
4886 const LM_SUB: usize = 32;
4887 let n = ids.len().saturating_sub(1);
4888 let hs = self.hidden_size;
4889 let rows = self.weights.lm_head.rows();
4890 let mut pos = 0usize;
4891 while pos < n {
4892 let end = (pos + CHUNK).min(n);
4893 let bsz = end - pos;
4894 let hb = self.prefill_batch_masked(&ids[pos..end], pos, task_mask);
4895 let mut k0 = 0usize;
4896 while k0 < bsz {
4897 let k1 = (k0 + LM_SUB).min(bsz);
4898 let sb = k1 - k0;
4899 if pos + k1 <= start {
4902 k0 = k1;
4903 continue;
4904 }
4905 let mut normed = vec![0.0f32; sb * hs];
4906 for k in 0..sb {
4907 let r = inference::rms_norm(
4908 &hb[(k0 + k) * hs..(k0 + k + 1) * hs],
4909 &self.weights.final_norm,
4910 self.rms_eps,
4911 self.norm_style,
4912 );
4913 normed[k * hs..(k + 1) * hs].copy_from_slice(&r);
4914 }
4915 let mut logits = vec![0.0f32; sb * rows];
4916 self.weights
4917 .lm_head
4918 .matmat(&normed, sb, &mut logits, self.pool.as_deref());
4919 for k in 0..sb {
4920 if pos + k0 + k < start {
4921 continue;
4922 }
4923 let lg = &mut logits[k * rows..k * rows + self.vocab_size.min(rows)];
4924 if let Some(mu) = self.logit_multiplier {
4925 for v in lg.iter_mut() {
4926 *v *= mu;
4927 }
4928 }
4929 if let Some(c) = self.final_softcap {
4933 for v in lg.iter_mut() {
4934 *v = c * (*v / c).tanh();
4935 }
4936 }
4937 if let Some(cm) = self.head_clusters.clone() {
4940 self.hierarchical_head_logprobs(&normed[k * hs..(k + 1) * hs], &cm, lg);
4941 }
4942 let lg = &logits[k * rows..k * rows + self.vocab_size.min(rows)];
4943 let target = ids[pos + k0 + k + 1] as usize;
4944 let max = lg.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
4945 let lse: f64 = lg
4946 .iter()
4947 .map(|&v| ((v - max) as f64).exp())
4948 .sum::<f64>()
4949 .ln()
4950 + max as f64;
4951 nll += lse - lg[target] as f64;
4952 cnt += 1;
4953 if std::env::var("CMF_PPL_TRACE").is_ok() {
4954 let top = lg
4955 .iter()
4956 .enumerate()
4957 .max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
4958 .map(|(i, _)| i)
4959 .unwrap_or(0);
4960 eprintln!(
4961 "BTRACE pos {} target {} nll {:.4} top {} lg_t {:.3} lg_top {:.3}",
4962 pos + k0 + k,
4963 target,
4964 lse - lg[target] as f64,
4965 top,
4966 lg[target],
4967 lg[top]
4968 );
4969 }
4970 }
4971 k0 = k1;
4972 }
4973 pos = end;
4974 }
4975 self.kv_cache.clear();
4976 self.kv_history.clear();
4977 return (nll, cnt);
4978 }
4979 for pos in 0..ids.len().saturating_sub(1) {
4980 let hidden = self.forward_layers(&self.embed_single(ids[pos]), pos, task_mask);
4981 let out_of_band = self.graph_logits.take();
4989 if pos < start {
4990 continue;
4991 }
4992 let logits = match out_of_band {
4993 Some(lg) => lg,
4994 None => {
4995 let normed = inference::rms_norm(
4996 &hidden,
4997 &self.weights.final_norm,
4998 self.rms_eps,
4999 self.norm_style,
5000 );
5001 self.lm_head_forward(&normed)
5005 }
5006 };
5007 let target = ids[pos + 1] as usize;
5008 let max = logits.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
5009 let lse: f64 = logits
5010 .iter()
5011 .map(|&v| ((v - max) as f64).exp())
5012 .sum::<f64>()
5013 .ln()
5014 + max as f64;
5015 let tok_nll = lse - logits[target] as f64;
5016 if std::env::var("CMF_PPL_TRACE").is_ok() && pos < 48 {
5017 let top = logits
5018 .iter()
5019 .enumerate()
5020 .max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
5021 .map(|(i, _)| i)
5022 .unwrap_or(0);
5023 eprintln!(
5024 "pos {pos:3} tgt {target:6} nll {tok_nll:7.3} | top1 {top:6} lg[t]={:.2} lg[top]={:.2}",
5025 logits[target], logits[top]
5026 );
5027 }
5028 nll += tok_nll;
5029 cnt += 1;
5030 }
5031 self.kv_cache.clear();
5032 self.kv_history.clear();
5033 (nll, cnt)
5034 }
5035
5036 pub fn nll_ids_o1(&mut self, ids: &[u32], prefill: usize) -> (f64, usize) {
5052 self.kv_cache.clear();
5053 self.kv_history.clear();
5054 self.o1_begin();
5055 let n = ids.len().saturating_sub(1);
5056 let p = prefill.min(n);
5057 let mut pos = 0usize;
5059 if self.can_prefill_batched() {
5060 const CHUNK: usize = 128;
5061 while pos < p {
5062 let end = (pos + CHUNK).min(p);
5063 let _ = self.prefill_batch(&ids[pos..end], pos);
5064 pos = end;
5065 }
5066 } else {
5067 while pos < p {
5068 let _ = self.forward_layers(&self.embed_single(ids[pos]), pos, None);
5069 pos += 1;
5070 }
5071 }
5072 self.o1_seal();
5073
5074 let mut nll = 0f64;
5075 let mut cnt = 0usize;
5076 for pos in p..n {
5077 let hidden = self.forward_layers(&self.embed_single(ids[pos]), pos, None);
5078 let normed = inference::rms_norm(
5079 &hidden,
5080 &self.weights.final_norm,
5081 self.rms_eps,
5082 self.norm_style,
5083 );
5084 let logits = self.lm_head_forward(&normed);
5088 let target = ids[pos + 1] as usize;
5089 let max = logits.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
5090 let lse: f64 = logits
5091 .iter()
5092 .map(|&v| ((v - max) as f64).exp())
5093 .sum::<f64>()
5094 .ln()
5095 + max as f64;
5096 let tok_nll = lse - logits[target] as f64;
5097 if std::env::var("CMF_PPL_TRACE").is_ok() && pos < 48 {
5098 let top = logits
5099 .iter()
5100 .enumerate()
5101 .max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
5102 .map(|(i, _)| i)
5103 .unwrap_or(0);
5104 eprintln!(
5105 "pos {pos:3} tgt {target:6} nll {tok_nll:7.3} | top1 {top:6} lg[t]={:.2} lg[top]={:.2}",
5106 logits[target], logits[top]
5107 );
5108 }
5109 nll += tok_nll;
5110 cnt += 1;
5111 }
5112 self.kv_cache.clear();
5113 self.kv_history.clear();
5114 (nll, cnt)
5115 }
5116
5117 pub fn calib_ids(&mut self, ids: &[u32], temps: &[f32]) -> (Vec<bool>, Vec<Vec<f32>>) {
5125 self.kv_cache.clear();
5126 self.kv_history.clear();
5127 let n = ids.len().saturating_sub(1);
5128 let mut correct = Vec::with_capacity(n);
5129 let mut pmax = Vec::with_capacity(n);
5130 for pos in 0..n {
5131 let emb = self.embed_single(ids[pos]);
5132 let hidden = self.forward_layers(&emb, pos, None);
5133 let normed = inference::rms_norm(
5134 &hidden,
5135 &self.weights.final_norm,
5136 self.rms_eps,
5137 self.norm_style,
5138 );
5139 let logits = self.lm_head_forward(&normed);
5143 let target = ids[pos + 1] as usize;
5144 let (mut amax, mut mval) = (0usize, f32::NEG_INFINITY);
5145 for (i, &v) in logits.iter().enumerate() {
5146 if v > mval {
5147 mval = v;
5148 amax = i;
5149 }
5150 }
5151 correct.push(amax == target);
5152 let row: Vec<f32> = temps
5153 .iter()
5154 .map(|&t| {
5155 let tt = t.max(1e-3);
5156 let s: f32 = logits.iter().map(|&v| ((v - mval) / tt).exp()).sum();
5157 1.0 / s.max(1e-12) })
5159 .collect();
5160 pmax.push(row);
5161 }
5162 self.kv_cache.clear();
5163 self.kv_history.clear();
5164 (correct, pmax)
5165 }
5166
5167 pub fn ppl_ids_dynamic(&mut self, ids: &[u32]) -> (f64, usize) {
5174 let mut router = match self.dyn_router.take() {
5175 Some(r) => r,
5176 None => return (self.ppl_ids(ids), 0),
5177 };
5178 router.reset();
5179 self.dyn_phi_seen = 0;
5180 let _ = self.set_active_skill(None);
5181
5182 self.kv_cache.clear();
5183
5184 self.kv_history.clear();
5185 let mut nll = 0f64;
5186 let mut cnt = 0usize;
5187 for pos in 0..ids.len().saturating_sub(1) {
5188 let hidden = self.forward_layers(&self.embed_single(ids[pos]), pos, None);
5189 let normed = inference::rms_norm(
5190 &hidden,
5191 &self.weights.final_norm,
5192 self.rms_eps,
5193 self.norm_style,
5194 );
5195 let logits = self.lm_head_forward(&normed);
5199 let target = ids[pos + 1] as usize;
5200 let max = logits.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
5201 let lse: f64 = logits
5202 .iter()
5203 .map(|&v| ((v - max) as f64).exp())
5204 .sum::<f64>()
5205 .ln()
5206 + max as f64;
5207 let tok_nll = lse - logits[target] as f64;
5208 if std::env::var("CMF_PPL_TRACE").is_ok() && pos < 48 {
5209 let top = logits
5210 .iter()
5211 .enumerate()
5212 .max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
5213 .map(|(i, _)| i)
5214 .unwrap_or(0);
5215 eprintln!(
5216 "pos {pos:3} tgt {target:6} nll {tok_nll:7.3} | top1 {top:6} lg[t]={:.2} lg[top]={:.2}",
5217 logits[target], logits[top]
5218 );
5219 }
5220 nll += tok_nll;
5221 cnt += 1;
5222 let phi = self.dyn_phi_ema.clone();
5224 if let Some(new_active) = router.step(&phi, pos) {
5225 let _ = self.set_active_skill(new_active);
5226 }
5227 }
5228 let switches = router.switches.len();
5229 let _ = self.set_active_skill(None);
5230 self.dyn_router = Some(router);
5231 self.kv_cache.clear();
5232 self.kv_history.clear();
5233 ((nll / cnt.max(1) as f64).exp(), switches)
5234 }
5235
5236 pub fn probe_phi(&mut self, ids: &[u32], layer: usize) -> Vec<f32> {
5238 self.kv_cache.clear();
5239 self.kv_history.clear();
5240 let mut acc = vec![0f32; self.hidden_size];
5241 for (pos, &id) in ids.iter().enumerate() {
5242 let h = self.forward_layers_upto(&self.embed_single(id), pos, None, Some(layer));
5243 for (a, v) in acc.iter_mut().zip(&h) {
5244 *a += v;
5245 }
5246 }
5247 let n = ids.len().max(1) as f32;
5248 for a in acc.iter_mut() {
5249 *a /= n;
5250 }
5251 self.kv_cache.clear();
5252 self.kv_history.clear();
5253 acc
5254 }
5255
5256 fn prefill_batch(&mut self, ids: &[u32], start_pos: usize) -> Vec<f32> {
5262 self.prefill_batch_masked(ids, start_pos, None)
5263 }
5264
5265 fn prefill_batch_masked(
5271 &mut self,
5272 ids: &[u32],
5273 start_pos: usize,
5274 task_mask: Option<&TaskMask>,
5275 ) -> Vec<f32> {
5276 self.prefill_batch_span(PrefillIn::Ids(ids), start_pos, task_mask, 0, usize::MAX)
5277 }
5278
5279 fn prefill_batch_span(
5285 &mut self,
5286 input: PrefillIn<'_>,
5287 start_pos: usize,
5288 task_mask: Option<&TaskMask>,
5289 from: usize,
5290 upto_excl: usize,
5291 ) -> Vec<f32> {
5292 let hs = self.hidden_size;
5293 let b = match input {
5294 PrefillIn::Ids(ids) => ids.len(),
5295 PrefillIn::Hidden(hb) => hb.len() / hs,
5296 };
5297 let upto_excl = upto_excl.min(self.num_layers);
5298 let mut h: Vec<f32>;
5302 let mut h_ready;
5303 match input {
5304 PrefillIn::Ids(_) => {
5305 h = vec![0.0; b * hs];
5306 h_ready = false;
5307 }
5308 PrefillIn::Hidden(hb) => {
5309 h = hb.to_vec();
5310 h_ready = true;
5311 }
5312 }
5313 let fill_h = |h: &mut Vec<f32>, me: &Self| {
5314 if let PrefillIn::Ids(ids) = input {
5315 for (bi, &id) in ids.iter().enumerate() {
5316 let e = me.embed_single(id);
5317 h[bi * hs..(bi + 1) * hs].copy_from_slice(&e);
5318 }
5319 if let Ok(tp) = std::env::var("CMF_TRACE_POS") {
5320 if let Ok(t) = tp.parse::<usize>() {
5321 if t >= start_pos && t < start_pos + ids.len() {
5322 let bi = t - start_pos;
5323 let row = &h[bi * hs..(bi + 1) * hs];
5324 let n: f32 = row.iter().map(|x| x * x).sum::<f32>().sqrt();
5325 eprintln!(
5326 "BATCH pos {t} embed: id {} |h| = {n:.6} h0 {:.6} h1 {:.6} | b={} start={start_pos} ids[..8]={:?}",
5327 ids[bi],
5328 row[0],
5329 row[1],
5330 ids.len(),
5331 &ids[..ids.len().min(8)]
5332 );
5333 }
5334 }
5335 }
5336 }
5337 };
5338 let (_nkv, _hd, _rd, eps) = (
5339 self.num_kv_heads,
5340 self.head_dim,
5341 self.rotary_dim,
5342 self.rms_eps,
5343 );
5344 let pool = self.pool.clone();
5345 let norm_style = self.norm_style;
5346
5347 #[cfg(target_os = "macos")]
5348 let mut chunk_skip_until = 0usize;
5349 for li in from..upto_excl {
5350 crate::gpu::set_layer(li as i64); #[cfg(target_os = "macos")]
5357 if task_mask.is_none() {
5358 if li < chunk_skip_until {
5359 continue;
5360 }
5361 if !h_ready && li == 0 && self.weights.embed_tokens.q8_row_parts().is_none() {
5367 fill_h(&mut h, self);
5368 h_ready = true;
5369 }
5370 let ids_for_embed = match input {
5371 PrefillIn::Ids(ids) => (!h_ready && li == 0).then_some(ids),
5372 PrefillIn::Hidden(_) => None,
5373 };
5374 let end = self.chunk_run_gpu(li, &mut h, b, start_pos, ids_for_embed, upto_excl);
5375 if end > li {
5376 h_ready = true;
5377 chunk_skip_until = end;
5378 if self.is_loop_end(end - 1) && end < self.num_layers {
5381 for bi in 0..b {
5382 let normed = inference::rms_norm(
5383 &h[bi * hs..(bi + 1) * hs],
5384 &self.weights.final_norm,
5385 eps,
5386 norm_style,
5387 );
5388 h[bi * hs..(bi + 1) * hs].copy_from_slice(&normed);
5389 }
5390 }
5391 continue;
5392 }
5393 }
5394 if !h_ready {
5395 fill_h(&mut h, self);
5396 h_ready = true;
5397 }
5398 let lw = &self.weights.layers[self.phys_layer(li)];
5399 match &lw.attn {
5401 AttnKind::Kda(w) => {
5402 let cfg = self.kda_cfg.expect("kda layer without kda_cfg");
5404 let mut normed = vec![0.0f32; b * hs];
5405 for bi in 0..b {
5406 inference::rms_norm_into(
5407 &h[bi * hs..(bi + 1) * hs],
5408 &lw.input_norm,
5409 eps,
5410 norm_style,
5411 &mut normed[bi * hs..(bi + 1) * hs],
5412 );
5413 }
5414 let attn = crate::linear_core::kda_forward_batch(
5415 &normed,
5416 b,
5417 w,
5418 &cfg,
5419 &mut self.kv_cache.layers[li].linear_state,
5420 pool.as_deref(),
5421 );
5422 for (dst, &a) in h.iter_mut().zip(&attn) {
5423 *dst += a;
5424 }
5425 }
5426 AttnKind::LinearGdn(w) => {
5427 let cfg = self.gdn_cfg.expect("gdn layer without gdn_cfg");
5429 let mut normed = vec![0.0f32; b * hs];
5430 for bi in 0..b {
5431 let r = inference::rms_norm(
5432 &h[bi * hs..(bi + 1) * hs],
5433 &lw.input_norm,
5434 eps,
5435 norm_style,
5436 );
5437 normed[bi * hs..(bi + 1) * hs].copy_from_slice(&r);
5438 }
5439 let attn = crate::linear_core::gdn_forward_batch(
5440 &normed,
5441 b,
5442 w,
5443 &cfg,
5444 &mut self.kv_cache.layers[li].linear_state,
5445 pool.as_deref(),
5446 );
5447 for (dst, &a) in h.iter_mut().zip(&attn) {
5448 *dst += a;
5449 }
5450 }
5451 AttnKind::ShortConv(w) => {
5452 let cfg = self
5455 .short_conv_cfg
5456 .expect("short-conv layer without short_conv_cfg");
5457 let mut normed = vec![0.0f32; b * hs];
5458 for bi in 0..b {
5459 inference::rms_norm_into(
5460 &h[bi * hs..(bi + 1) * hs],
5461 &lw.input_norm,
5462 eps,
5463 norm_style,
5464 &mut normed[bi * hs..(bi + 1) * hs],
5465 );
5466 }
5467 let attn = short_conv_forward_batch(
5468 &normed,
5469 b,
5470 w,
5471 &cfg,
5472 &mut self.kv_cache.layers[li].linear_state,
5473 pool.as_deref(),
5474 );
5475 for (dst, &a) in h.iter_mut().zip(&attn) {
5476 *dst += a;
5477 }
5478 }
5479 AttnKind::Mla(w) => {
5480 let inv_freq_l = self.layer_inv_freq(li);
5483 let rs = self.layer_rope_scale(li);
5484 let mut normed = vec![0.0f32; hs];
5485 for bi in 0..b {
5486 inference::rms_norm_into(
5487 &h[bi * hs..(bi + 1) * hs],
5488 &lw.input_norm,
5489 eps,
5490 norm_style,
5491 &mut normed,
5492 );
5493 let ao = mla_attention(
5494 w,
5495 &normed,
5496 &mut self.kv_cache.layers[li],
5497 start_pos + bi,
5498 &inv_freq_l,
5499 rs,
5500 eps,
5501 pool.as_deref(),
5502 );
5503 for (dst, &a) in h[bi * hs..(bi + 1) * hs].iter_mut().zip(&ao) {
5504 *dst += a;
5505 }
5506 }
5507 }
5508 AttnKind::Full {
5509 wq,
5510 wk,
5511 wv,
5512 wo,
5513 q_norm,
5514 k_norm,
5515 output_gate,
5516 softplus_gate,
5517 bias,
5518 } => {
5519 let mut normed = vec![0.0f32; b * hs];
5523 for bi in 0..b {
5524 inference::rms_norm_into(
5525 &h[bi * hs..(bi + 1) * hs],
5526 &lw.input_norm,
5527 eps,
5528 norm_style,
5529 &mut normed[bi * hs..(bi + 1) * hs],
5530 );
5531 }
5532 let inv_freq_l = self.layer_inv_freq(li);
5533 let (nkv_l, hd_l, rd_l) = self.layer_geom(li);
5534 let cfg = QwenAttnCfg {
5535 num_heads: self.layer_num_heads(li),
5536 num_kv_heads: nkv_l,
5537 head_dim: hd_l,
5538 hidden_size: hs,
5539 position: start_pos,
5540 inv_freq: &inv_freq_l,
5541 rotary_dim: rd_l,
5542 scale: self.attn_scale,
5543 softcap: self.attn_softcap,
5544 window: self.layer_window(li),
5545 v_norm: self.attn_v_norm,
5546 q_norm: q_norm.as_deref(),
5547 k_norm: k_norm.as_deref(),
5548 output_gate: *output_gate,
5549 softplus_gate: softplus_gate
5550 .as_ref()
5551 .map(|(gate, per_head)| (gate, *per_head)),
5552 rope_scale: self.layer_rope_scale(li),
5553 bias: bias
5554 .as_ref()
5555 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
5556 rms_eps: eps,
5557 norm_style,
5558 pool: pool.as_deref(),
5559 };
5560 let mut attn = attention::qwen_attention_batch(
5561 &normed,
5562 b,
5563 wq,
5564 wk,
5565 wv,
5566 wo,
5567 &mut self.kv_cache.layers[li],
5568 &cfg,
5569 );
5570 if let Some(w) = &lw.attn_out_norm {
5571 for bi in 0..b {
5572 inference::rms_norm_into(
5573 &attn[bi * hs..(bi + 1) * hs],
5574 w,
5575 eps,
5576 norm_style,
5577 &mut normed[bi * hs..(bi + 1) * hs],
5578 );
5579 }
5580 attn.copy_from_slice(&normed);
5581 }
5582 for (dst, &a) in h.iter_mut().zip(&attn) {
5583 *dst += a;
5584 }
5585 }
5586 AttnKind::Linear(w) => {
5587 for bi in 0..b {
5588 let normed = inference::rms_norm(
5589 &h[bi * hs..(bi + 1) * hs],
5590 &lw.input_norm,
5591 eps,
5592 norm_style,
5593 );
5594 vmf_phase_forward(
5595 &normed,
5596 w,
5597 &self.vmf_cfg.expect("linear layer without vmf_cfg"),
5598 &mut self.kv_cache.layers[li].linear_state,
5599 pool.as_deref(),
5600 )
5601 .iter()
5602 .enumerate()
5603 .for_each(|(i, &a)| h[bi * hs + i] += a);
5604 }
5605 }
5606 }
5607
5608 let lw = &self.weights.layers[self.phys_layer(li)];
5610 let mut post = vec![0.0f32; b * hs];
5611 for bi in 0..b {
5612 let r =
5613 inference::rms_norm(&h[bi * hs..(bi + 1) * hs], &lw.post_norm, eps, norm_style);
5614 post[bi * hs..(bi + 1) * hs].copy_from_slice(&r);
5615 }
5616 let mask_row = task_mask
5619 .filter(|m| m.ffn_active_count(li) < self.intermediate_size)
5620 .and_then(|m| m.ffn_masks.get(li))
5621 .map(|v| v.as_slice());
5622 let mut ffn = match &lw.ffn {
5623 FfnKind::Dense(d) if !d.segs.is_empty() => {
5624 tube_ffn(d, &post, b, pool.as_deref(), mask_row)
5625 }
5626 FfnKind::Dense(d) => dense_ffn_batch(d, &post, b, pool.as_deref(), mask_row),
5627 FfnKind::Moe(m) => moe_ffn_batch(m, &post, b, hs, pool.as_deref(), None),
5628 FfnKind::DenseMoe(dm) => {
5631 let mut out = vec![0.0f32; b * hs];
5632 for bi in 0..b {
5633 let r = dense_moe_ffn(
5634 dm,
5635 &post[bi * hs..(bi + 1) * hs],
5636 &h[bi * hs..(bi + 1) * hs],
5637 eps,
5638 norm_style,
5639 pool.as_deref(),
5640 );
5641 out[bi * hs..(bi + 1) * hs].copy_from_slice(&r);
5642 }
5643 out
5644 }
5645 };
5646 if let Some(w) = &lw.ffn_out_norm {
5647 for bi in 0..b {
5648 inference::rms_norm_into(
5649 &ffn[bi * hs..(bi + 1) * hs],
5650 w,
5651 eps,
5652 norm_style,
5653 &mut post[bi * hs..(bi + 1) * hs],
5654 );
5655 }
5656 ffn.copy_from_slice(&post);
5657 }
5658 for (dst, &f) in h.iter_mut().zip(&ffn) {
5659 *dst += f;
5660 }
5661 if let Some(sc) = lw.layer_scale {
5662 for v in h.iter_mut() {
5663 *v *= sc;
5664 }
5665 }
5666 if let Ok(tp) = std::env::var("CMF_TRACE_POS") {
5667 if let Ok(t) = tp.parse::<usize>() {
5668 if t >= start_pos && t < start_pos + b {
5669 let bi = t - start_pos;
5670 let row = &h[bi * hs..(bi + 1) * hs];
5671 let n: f32 = row.iter().map(|x| x * x).sum::<f32>().sqrt();
5672 eprintln!(
5673 "BATCH pos {t} after layer {li}: |h| = {n:.6} h0 {:.6} h1 {:.6}",
5674 row[0], row[1]
5675 );
5676 }
5677 }
5678 }
5679 if std::env::var("CMF_DEBUG_LAYERS").is_ok() {
5683 let row = &h[(b - 1) * hs..b * hs];
5684 let rms =
5685 (row.iter().map(|&v| (v as f64) * (v as f64)).sum::<f64>() / hs as f64).sqrt();
5686 let mx = row.iter().fold(0f32, |m, &v| m.max(v.abs()));
5687 eprintln!(
5688 "layer {li:>3} {:>10} ffn={:<5} rms={rms:>12.4} max={mx:>12.4}",
5689 match &self.weights.layers[self.phys_layer(li)].attn {
5690 AttnKind::LinearGdn(_) => "gdn",
5691 AttnKind::Linear(_) => "vmf",
5692 AttnKind::ShortConv(_) => "conv",
5693 _ => "attn",
5694 },
5695 match &lw.ffn {
5696 FfnKind::Moe(_) => "moe",
5697 FfnKind::Dense(_) => "dense",
5698 FfnKind::DenseMoe(_) => "dense+moe",
5699 },
5700 );
5701 }
5702 if self.is_loop_end(li) && li + 1 < self.num_layers {
5704 for bi in 0..b {
5705 let normed = inference::rms_norm(
5706 &h[bi * hs..(bi + 1) * hs],
5707 &self.weights.final_norm,
5708 eps,
5709 norm_style,
5710 );
5711 h[bi * hs..(bi + 1) * hs].copy_from_slice(&normed);
5712 }
5713 }
5714 if std::env::var("CMF_TRACE_H").is_ok() {
5715 let n = h[..hs].iter().map(|v| v.abs()).sum::<f32>() / hs as f32;
5716 let mx = h[..hs].iter().fold(0.0f32, |a, &v| a.max(v.abs()));
5717 eprintln!(
5718 "layer {li}: mean|h|={n:.4} max|h|={mx:.2} scale={:?}",
5719 lw.layer_scale
5720 );
5721 }
5722 }
5723 crate::gpu::set_layer(-1); h
5725 }
5726
5727 fn embed_single(&self, id: u32) -> Vec<f32> {
5729 let mut out = vec![0.0f32; self.hidden_size];
5730 if (id as usize) < self.weights.embed_tokens.rows() {
5731 self.weights.embed_tokens.row_f32(id as usize, &mut out);
5732 }
5733 if self.embed_multiplier != 1.0 {
5734 for v in out.iter_mut() {
5735 *v *= self.embed_multiplier;
5736 }
5737 }
5738 if self.dsv4.is_some() || self.qwen4_exp.is_some() {
5742 let mut v = vec![0.0f32; self.hidden_size.max(1)];
5743 v[0] = id as f32;
5744 return v;
5745 }
5746 if let Some(b) = &self.g3n {
5749 return b.0.extend_embedding(id, &out, self.pool.as_deref());
5750 }
5751 out
5752 }
5753
5754 #[cfg(target_os = "macos")]
5760 fn chunk_run_gpu(
5761 &mut self,
5762 li0: usize,
5763 h: &mut [f32],
5764 b: usize,
5765 pos0: usize,
5766 embed_ids: Option<&[u32]>,
5767 cap: usize,
5768 ) -> usize {
5769 if !crate::gpu::enabled_here()
5773 || std::env::var("CMF_GPU_CHUNK")
5774 .map(|v| v == "0")
5775 .unwrap_or(false)
5776 || b < 32
5777 || self.swa.is_some()
5778 || self.global_attn.is_some()
5779 || self.attn_v_norm
5780 || (self.attn_scale - 1.0 / (self.head_dim as f32).sqrt()).abs() > 1e-9
5781 {
5782 return li0;
5783 }
5784 let Some(model) = self.model.clone() else {
5785 return li0;
5786 };
5787 let inv_freq = self.inv_freq.clone();
5788 let (nh, nkv, hd, hs) = (
5789 self.num_heads,
5790 self.num_kv_heads,
5791 self.head_dim,
5792 self.hidden_size,
5793 );
5794 let loop_end = if self.loop_final_norm {
5798 ((li0 / self.physical_layers) + 1) * self.physical_layers
5799 } else {
5800 self.num_layers
5801 };
5802 let mut layers: Vec<crate::gpu_metal::ChunkLayer> = Vec::new();
5803 let mut stored_at: Vec<usize> = Vec::new();
5804 for li in li0..self.num_layers.min(loop_end).min(cap) {
5805 let lw = &self.weights.layers[self.phys_layer(li)];
5806 if lw.attn_out_norm.is_some() || lw.ffn_out_norm.is_some() || lw.layer_scale.is_some() {
5807 break;
5808 }
5809 let AttnKind::Full {
5810 wq,
5811 wk,
5812 wv,
5813 wo,
5814 q_norm,
5815 k_norm,
5816 output_gate: false,
5817 softplus_gate: None,
5818 bias,
5819 } = &lw.attn
5820 else {
5821 break;
5822 };
5823 let FfnKind::Dense(d) = &lw.ffn else { break };
5824 if d.act != Act::Silu || !d.segs.is_empty() {
5825 break;
5826 }
5827 fn cw(t: &QTensor) -> Option<(usize, usize, usize, &[f32])> {
5832 t.q8_row_parts()
5833 .or_else(|| t.q4t_parts().map(|(i, r, c)| (i, r, c, &[][..])))
5834 .or_else(|| t.q4tp_parts().map(|(i, r, c)| (i, r, c, &[][..])))
5835 }
5836 let parts = (
5837 cw(wq),
5838 cw(wk),
5839 cw(wv),
5840 cw(wo),
5841 cw(&d.gate_proj),
5842 cw(&d.up_proj),
5843 cw(&d.down_proj),
5844 );
5845 let (Some(pq), Some(pk), Some(pv), Some(po), Some(pg), Some(pu), Some(pd)) = parts
5846 else {
5847 break;
5848 };
5849 let layer = &self.kv_cache.layers[li];
5850 if layer.mode != crate::kv_cache::KvMode::F32 || layer.o1.is_some() {
5851 break;
5852 }
5853 stored_at.push(layer.head_len(0));
5854 layers.push(crate::gpu_metal::ChunkLayer {
5855 model: &model,
5856 kv_id: self.graph_kv_id,
5857 layer: li,
5858 wq: pq,
5859 wk: pk,
5860 wv: pv,
5861 wo: po,
5862 gate: pg,
5863 up: pu,
5864 down: pd,
5865 input_norm: &lw.input_norm,
5866 post_norm: &lw.post_norm,
5867 bias: bias
5868 .as_ref()
5869 .map(|(a, bb, cc)| (a.as_slice(), bb.as_slice(), cc.as_slice())),
5870 q_norm: q_norm.as_deref(),
5871 k_norm: k_norm.as_deref(),
5872 inv_freq: &inv_freq,
5873 rd: self.rotary_dim,
5874 nh,
5875 nkv,
5876 hd,
5877 hs,
5878 inter: d.gate_proj.rows(),
5879 gemma: matches!(self.norm_style, cortiq_core::NormStyle::Gemma),
5880 eps: self.rms_eps as f32,
5881 });
5882 }
5883 if layers.is_empty() {
5884 return li0;
5885 }
5886 let row = nkv * hd;
5887 let mut store: Vec<(Vec<f32>, Vec<f32>, Vec<f32>)> = stored_at
5888 .iter()
5889 .map(|&st| (vec![0f32; b * row], vec![0f32; b * row], vec![0f32; st + b]))
5890 .collect();
5891 let mut io: Vec<crate::gpu_metal::ChunkIo> = Vec::with_capacity(layers.len());
5892 for (i, (ok, ov, oi)) in store.iter_mut().enumerate() {
5893 let li = layers[i].layer;
5894 let layer = &self.kv_cache.layers[li];
5895 io.push(crate::gpu_metal::ChunkIo {
5896 cpu_stored: stored_at[i],
5897 cpu_k: (0..nkv).map(|g| layer.head_keys(g)).collect(),
5898 cpu_v: (0..nkv).map(|g| layer.head_values(g)).collect(),
5899 out_k: ok,
5900 out_v: ov,
5901 imp: oi,
5902 });
5903 }
5904 let n_run = layers.len();
5905 let last = layers.last().map(|l| l.layer + 1).unwrap_or(li0);
5906 let ep = embed_ids.and_then(|ids| {
5909 self.weights
5910 .embed_tokens
5911 .q8_row_parts()
5912 .map(|(idx, rows, _c, rs)| crate::gpu_metal::ChunkEmbed {
5913 idx,
5914 rows,
5915 row_scale: rs,
5916 ids,
5917 mult: self.embed_multiplier,
5918 })
5919 });
5920 if embed_ids.is_some() && ep.is_none() {
5921 return li0;
5922 }
5923 if !crate::gpu_metal::chunk_run_gpu(&layers, &mut io, h, b, pos0, ep.as_ref()) {
5924 return li0;
5925 }
5926 drop(io);
5927 drop(layers);
5928 for (i, (ok, ov, oi)) in store.iter().enumerate().take(n_run) {
5931 let li = li0 + i;
5932 let layer = &mut self.kv_cache.layers[li];
5933 for bi in 0..b {
5934 layer.append(
5935 &ok[bi * row..(bi + 1) * row],
5936 &ov[bi * row..(bi + 1) * row],
5937 &[],
5938 );
5939 }
5940 layer.accumulate_imp(oi);
5941 }
5942 last
5943 }
5944
5945 fn layer_is_local(&self, li: usize) -> bool {
5948 if let Some(layers) = &self.sliding_layers {
5949 return layers.get(li).copied().unwrap_or(false);
5950 }
5951 match self.swa {
5952 Some((_, pattern)) => (li + 1) % pattern.max(1) != 0,
5953 None => false,
5954 }
5955 }
5956
5957 fn layer_inv_freq(&self, li: usize) -> std::sync::Arc<Vec<f32>> {
5960 if self.layer_is_local(li) {
5961 if let Some(f) = &self.inv_freq_local {
5962 return f.clone();
5963 }
5964 } else if let Some(f) = &self.inv_freq_global {
5965 return f.clone();
5966 }
5967 self.inv_freq.clone()
5968 }
5969
5970 fn layer_window(&self, li: usize) -> Option<usize> {
5972 self.swa
5973 .and_then(|(w, _)| self.layer_is_local(li).then_some(w))
5974 }
5975
5976 fn layer_num_heads(&self, li: usize) -> usize {
5977 self.attention_heads_per_layer
5978 .as_ref()
5979 .and_then(|v| v.get(li).copied())
5980 .unwrap_or(self.num_heads)
5981 }
5982
5983 fn layer_rope_scale(&self, li: usize) -> f32 {
5984 if self.layer_is_local(li) {
5985 self.rope_scale_local
5986 } else {
5987 self.rope_scale
5988 }
5989 }
5990
5991 fn layer_geom(&self, li: usize) -> (usize, usize, usize) {
5994 if !self.layer_is_local(li) {
5995 if let Some((ghd, gkv)) = self.global_attn {
5996 return (gkv, ghd, ghd);
5997 }
5998 }
5999 (
6000 self.num_kv_heads,
6001 self.head_dim,
6002 if self.layer_is_local(li) {
6003 self.rotary_dim_local.unwrap_or(self.rotary_dim)
6004 } else {
6005 self.rotary_dim
6006 },
6007 )
6008 }
6009
6010 fn forward_layers(
6012 &mut self,
6013 hidden: &[f32],
6014 position: usize,
6015 task_mask: Option<&TaskMask>,
6016 ) -> Vec<f32> {
6017 self.forward_layers_upto(hidden, position, task_mask, None)
6018 }
6019
6020 pub fn embed_id(&self, id: u32) -> Vec<f32> {
6028 self.embed_single(id)
6029 }
6030
6031 pub fn split_supported(&self) -> Result<(), String> {
6035 if self.dsv4.is_some() {
6036 return Err(
6037 "network split: DeepSeek-V4 runs its own fused stack (not splittable yet)".into(),
6038 );
6039 }
6040 if self.qwen4_exp.is_some() {
6041 return Err(
6042 "network split: Qwen3.8-Flash-Next hyper/QSA stack is not splittable yet".into(),
6043 );
6044 }
6045 if self.g3n.is_some() {
6046 return Err(
6047 "network split: Gemma-3n runs its own AltUp stack (not splittable yet)".into(),
6048 );
6049 }
6050 Ok(())
6051 }
6052
6053 pub fn forward_span(
6058 &mut self,
6059 hidden: &[f32],
6060 position: usize,
6061 from: usize,
6062 upto: usize,
6063 task_mask: Option<&TaskMask>,
6064 ) -> Result<Vec<f32>, String> {
6065 self.split_supported()?;
6066 if from > upto || upto >= self.num_layers {
6067 return Err(format!(
6068 "forward_span: layer range {from}..={upto} outside 0..{}",
6069 self.num_layers
6070 ));
6071 }
6072 if hidden.len() != self.hidden_size {
6073 return Err(format!(
6074 "forward_span: hidden len {} ≠ hidden_size {}",
6075 hidden.len(),
6076 self.hidden_size
6077 ));
6078 }
6079 Ok(self.forward_layers_span(hidden, position, task_mask, from, Some(upto)))
6080 }
6081
6082 pub fn logits_from_hidden(&mut self, hidden: &[f32]) -> Vec<f32> {
6085 let normed = inference::rms_norm(
6086 hidden,
6087 &self.weights.final_norm,
6088 self.rms_eps,
6089 self.norm_style,
6090 );
6091 self.lm_head_forward(&normed)
6092 }
6093
6094 pub fn sample_next(&mut self, logits: &[f32], past_tokens: &[u32]) -> u32 {
6096 sampler::sample_with_scratch(
6097 logits,
6098 &self.sampler_config,
6099 past_tokens,
6100 &mut self.rng,
6101 &mut self.sampler_scratch,
6102 )
6103 }
6104
6105 pub fn reset_session(&mut self) {
6107 self.kv_cache.clear();
6108 self.kv_history.clear();
6109 crate::gpu::graph_kv_reset(self.graph_kv_id);
6110 }
6111
6112 pub fn prefill_span_ids(
6118 &mut self,
6119 ids: &[u32],
6120 start_pos: usize,
6121 upto: usize,
6122 task_mask: Option<&TaskMask>,
6123 ) -> Result<Vec<f32>, String> {
6124 self.split_supported()?;
6125 if upto >= self.num_layers {
6126 return Err(format!(
6127 "prefill_span_ids: upto {upto} outside 0..{}",
6128 self.num_layers
6129 ));
6130 }
6131 if self.can_prefill_batched() && !self.graph_prefill_preferred() {
6135 Ok(self.prefill_batch_span(PrefillIn::Ids(ids), start_pos, task_mask, 0, upto + 1))
6136 } else {
6137 let hs = self.hidden_size;
6138 let mut out = Vec::with_capacity(ids.len() * hs);
6139 for (i, &id) in ids.iter().enumerate() {
6140 let emb = self.embed_id(id);
6141 out.extend_from_slice(&self.forward_span(
6142 &emb,
6143 start_pos + i,
6144 0,
6145 upto,
6146 task_mask,
6147 )?);
6148 }
6149 Ok(out)
6150 }
6151 }
6152
6153 pub fn prefill_span_hidden(
6156 &mut self,
6157 hidden: &[f32],
6158 start_pos: usize,
6159 from: usize,
6160 upto: usize,
6161 task_mask: Option<&TaskMask>,
6162 ) -> Result<Vec<f32>, String> {
6163 self.split_supported()?;
6164 let hs = self.hidden_size;
6165 if hidden.is_empty() || hidden.len() % hs != 0 {
6166 return Err(format!(
6167 "prefill_span_hidden: {} floats is not a multiple of hidden {hs}",
6168 hidden.len()
6169 ));
6170 }
6171 if from > upto || upto >= self.num_layers {
6172 return Err(format!(
6173 "prefill_span_hidden: layer range {from}..={upto} outside 0..{}",
6174 self.num_layers
6175 ));
6176 }
6177 if self.can_prefill_batched() && !self.graph_prefill_preferred() {
6178 Ok(self.prefill_batch_span(
6179 PrefillIn::Hidden(hidden),
6180 start_pos,
6181 task_mask,
6182 from,
6183 upto + 1,
6184 ))
6185 } else {
6186 let b = hidden.len() / hs;
6187 let mut out = Vec::with_capacity(hidden.len());
6188 for i in 0..b {
6189 let h = self.forward_span(
6190 &hidden[i * hs..(i + 1) * hs],
6191 start_pos + i,
6192 from,
6193 upto,
6194 task_mask,
6195 )?;
6196 out.extend_from_slice(&h);
6197 }
6198 Ok(out)
6199 }
6200 }
6201
6202 fn try_token_graph_wgpu(
6206 &self,
6207 hidden: &[f32],
6208 position: usize,
6209 logits_out: &mut Vec<f32>,
6210 layers_run: &mut usize,
6211 ) -> Option<Vec<f32>> {
6212 self.try_token_graph_wgpu_steps(
6213 hidden,
6214 position,
6215 logits_out,
6216 1,
6217 None,
6218 Some(layers_run),
6219 0,
6220 self.num_layers,
6221 )
6222 }
6223
6224 fn try_token_graph_wgpu_span(
6228 &self,
6229 hidden: &[f32],
6230 position: usize,
6231 logits_out: &mut Vec<f32>,
6232 from: usize,
6233 upto_excl: usize,
6234 layers_run: &mut usize,
6235 ) -> Option<Vec<f32>> {
6236 self.try_token_graph_wgpu_steps(
6237 hidden,
6238 position,
6239 logits_out,
6240 1,
6241 None,
6242 Some(layers_run),
6243 from,
6244 upto_excl,
6245 )
6246 }
6247
6248 fn try_multi_burst(&self, t_next: u32, position: usize, k: usize) -> Option<Vec<u32>> {
6252 if self.o1_active() || self.attn_softcap > 0.0 {
6253 return None;
6254 }
6255 let graph_on = crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Decode);
6256 if !graph_on || crate::gpu::graph_unsupported() {
6257 return None;
6264 }
6265 let emb = self.embed_single(t_next);
6266 let mut lg = Vec::new();
6267 let mut ids = Vec::new();
6268 self.try_token_graph_wgpu_steps(
6269 &emb,
6270 position,
6271 &mut lg,
6272 k,
6273 Some(&mut ids),
6274 None,
6275 0,
6276 self.num_layers,
6277 )?;
6278 (ids.len() == k).then_some(ids)
6279 }
6280
6281 fn try_token_graph_wgpu_steps(
6285 &self,
6286 hidden: &[f32],
6287 position: usize,
6288 logits_out: &mut Vec<f32>,
6289 steps: usize,
6290 ids_out: Option<&mut Vec<u32>>,
6291 layers_run: Option<&mut usize>,
6292 from: usize,
6293 upto_excl: usize,
6294 ) -> Option<Vec<f32>> {
6295 let o1_gpu = std::env::var("CMF_O1_GPU").as_deref() == Ok("1");
6298 if (self.o1_active() && !o1_gpu) || self.attn_softcap > 0.0 {
6299 return None;
6303 }
6304 let o1_views: Vec<Option<Vec<crate::nystrom::O1DeviceView<'_>>>> = (from..upto_excl)
6309 .map(|li| {
6310 if !o1_gpu {
6311 return None;
6312 }
6313 self.kv_cache.layers[self.phys_layer(li)].o1_views()
6314 })
6315 .collect();
6316 if self.o1_active() && o1_gpu {
6317 let want: usize = (from..upto_excl)
6320 .filter(|li| self.kv_cache.layers[self.phys_layer(*li)].o1.is_some())
6321 .count();
6322 let have = o1_views.iter().filter(|v| v.is_some()).count();
6323 if want == 0 || have != want {
6324 use std::sync::atomic::{AtomicUsize, Ordering};
6334 static LAST: AtomicUsize = AtomicUsize::new(usize::MAX);
6335 let code = have * 1000 + want;
6336 if LAST.swap(code, Ordering::Relaxed) != code {
6337 tracing::warn!(
6338 "o1 graph: {have} of {want} layers sealed — per-op until all seal"
6339 );
6340 }
6341 return None;
6342 }
6343 }
6344 let nh = self.num_heads;
6345 let (nkv, hd, rd) = self.layer_geom(0);
6346 let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
6347 let mut layers = Vec::with_capacity(upto_excl - from);
6348 let mut model = None;
6349 let dbg = std::env::var("CMF_GRAPH_DEBUG").is_ok();
6350 fn gw(t: &QTensor) -> Option<crate::gpu::GraphW<'_>> {
6351 if let Some((_, i, kind, rs)) = t.graph_weight() {
6352 return Some(crate::gpu::GraphW {
6353 idx: i,
6354 kind,
6355 row_scale: rs,
6356 data: &[],
6357 });
6358 }
6359 t.as_f32().map(|d| crate::gpu::GraphW {
6361 idx: 0,
6362 kind: 4,
6363 row_scale: &[],
6364 data: d,
6365 })
6366 }
6367 for li in from..upto_excl {
6368 let lw = &self.weights.layers[self.phys_layer(li)];
6369 if dbg {
6370 let ak = match &lw.attn {
6371 AttnKind::Mla(_) => "Mla".into(),
6372 AttnKind::Full {
6373 output_gate, bias, ..
6374 } => format!("Full gate={output_gate} bias={}", bias.is_some()),
6375 AttnKind::LinearGdn(_) => "LinearGdn".into(),
6376 AttnKind::Kda(_) => "Kda".into(),
6377 AttnKind::Linear(_) => "Linear".into(),
6378 AttnKind::ShortConv(_) => "ShortConv".into(),
6379 };
6380 let fk = match &lw.ffn {
6381 FfnKind::Dense(_) => "Dense",
6382 FfnKind::Moe(_) => "Moe",
6383 FfnKind::DenseMoe(_) => "DenseMoe",
6384 };
6385 eprintln!("graph L{li}: attn={ak} ffn={fk}");
6386 }
6387 let gffn = match &lw.ffn {
6388 FfnKind::DenseMoe(_) => return None, FfnKind::Dense(d) if !d.segs.is_empty() => return None,
6392 FfnKind::Dense(d) => crate::gpu::GraphFfn::Dense {
6393 gate: gw(&d.gate_proj)?,
6394 up: gw(&d.up_proj)?,
6395 down: gw(&d.down_proj)?,
6396 },
6397 FfnKind::Moe(m) => {
6398 if m.route_tau.is_some()
6405 || m.mask.is_some()
6406 || (m.routed_scaling - 1.0).abs() > 1e-9
6407 {
6408 return None;
6409 }
6410 let shared = m.shared.as_ref();
6411 let has_shared = shared.is_some();
6412 let sgate = match shared {
6413 Some((_, sg)) => gw(sg.as_ref()?)?,
6414 None => gw(&m.router)?,
6417 };
6418 let router = gw(&m.router)?;
6419 let inter = m.experts.first()?.gate_proj.rows();
6420 let mut experts = Vec::with_capacity(m.experts.len() + 1);
6421 let mut q4tp: Option<bool> = None;
6424 let mut gu_q2: Option<bool> = None;
6427 for e in m.experts.iter().chain(shared.map(|(se, _)| se)) {
6428 if !matches!(e.act, Act::Silu)
6429 || e.gate_proj.rows() != inter
6430 || e.up_proj.rows() != inter
6431 {
6432 return None;
6433 }
6434 let (mm, gi, ui, di, is_p, is_q2) = match e.gate_proj.mapped_q4t() {
6435 Some((mm, gi)) => (
6436 mm,
6437 gi,
6438 e.up_proj.mapped_q4t()?.1,
6439 e.down_proj.mapped_q4t()?.1,
6440 false,
6441 false,
6442 ),
6443 None => match e.gate_proj.mapped_q2tp() {
6444 Some((mm, gi)) => (
6445 mm,
6446 gi,
6447 e.up_proj.mapped_q2tp()?.1,
6448 e.down_proj.mapped_q4tp()?.1,
6449 true,
6450 true,
6451 ),
6452 None => {
6453 let (mm, gi) = e.gate_proj.mapped_q4tp()?;
6454 (
6455 mm,
6456 gi,
6457 e.up_proj.mapped_q4tp()?.1,
6458 e.down_proj.mapped_q4tp()?.1,
6459 true,
6460 false,
6461 )
6462 }
6463 },
6464 };
6465 if *q4tp.get_or_insert(is_p) != is_p || *gu_q2.get_or_insert(is_q2) != is_q2
6466 {
6467 tracing::warn!(
6473 "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."
6474 );
6475 return None;
6476 }
6477 model.get_or_insert_with(|| mm.clone());
6478 experts.push((gi, ui, di));
6479 }
6480 crate::gpu::GraphFfn::Moe {
6481 router,
6482 shared_gate: sgate,
6483 experts,
6484 n_exp: m.experts.len(),
6485 top_k: std::env::var("CMF_TOPK_PROBE")
6491 .ok()
6492 .and_then(|v| v.parse::<usize>().ok())
6493 .filter(|k| *k > 0 && *k <= m.top_k)
6494 .unwrap_or(m.top_k),
6495 inter,
6496 norm_topk: m.norm_topk_prob,
6497 q4tp: q4tp?,
6498 gu_q2: gu_q2.unwrap_or(false),
6499 sigmoid: m.router_sigmoid,
6500 bias: m.expert_bias.as_deref(),
6501 has_shared,
6502 }
6503 }
6504 };
6505 let attn = match &lw.attn {
6506 AttnKind::Full {
6507 wq,
6508 wk,
6509 wv,
6510 wo,
6511 q_norm,
6512 k_norm,
6513 output_gate,
6514 softplus_gate,
6515 bias,
6516 } => {
6517 if softplus_gate.is_some() || self.attention_heads_per_layer.is_some() {
6518 return None;
6519 }
6520 let (m, _, _, _) = wq.graph_weight()?;
6521 model = Some(m.clone());
6522 crate::gpu::GraphAttn::Full {
6523 wq: gw(wq)?,
6524 wk: gw(wk)?,
6525 wv: gw(wv)?,
6526 wo: gw(wo)?,
6527 q_norm: q_norm.as_deref(),
6528 k_norm: k_norm.as_deref(),
6529 bias: bias
6530 .as_ref()
6531 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
6532 output_gate: *output_gate,
6533 cpu_k: self.kv_cache.layers[li].k_heads(),
6534 cpu_v: self.kv_cache.layers[li].v_heads(),
6535 }
6536 }
6537 AttnKind::LinearGdn(w) => {
6538 let cfg = self.gdn_cfg?;
6539 let (m, _, _, _) = w.in_proj_qkv.graph_weight()?;
6540 model = Some(m.clone());
6541 crate::gpu::GraphAttn::Gdn {
6542 qkv: gw(&w.in_proj_qkv)?,
6543 z: gw(&w.in_proj_z)?,
6544 a: gw(&w.in_proj_a)?,
6545 b: gw(&w.in_proj_b)?,
6546 out: gw(&w.out_proj)?,
6547 conv1d: &w.conv1d,
6548 a_log: &w.a_log,
6549 dt_bias: &w.dt_bias,
6550 norm: &w.norm,
6551 nv: cfg.num_v_heads,
6552 nk: cfg.num_k_heads,
6553 dk: cfg.key_head_dim,
6554 dv: cfg.value_head_dim,
6555 kk: cfg.conv_kernel,
6556 cpu_state: &self.kv_cache.layers[self.phys_layer(li)].linear_state,
6557 }
6558 }
6559 AttnKind::ShortConv(w) => {
6560 let cfg = self.short_conv_cfg?;
6561 let (m, _, _, _) = w.in_proj.graph_weight()?;
6562 model = Some(m.clone());
6563 crate::gpu::GraphAttn::ShortConv {
6564 inp: gw(&w.in_proj)?,
6565 out: gw(&w.out_proj)?,
6566 taps: &w.conv,
6567 kernel: cfg.kernel,
6568 cpu_state: &self.kv_cache.layers[self.phys_layer(li)].linear_state,
6569 }
6570 }
6571 _ => return None,
6572 };
6573 layers.push(crate::gpu::GraphLayer {
6574 input_norm: &lw.input_norm,
6575 attn,
6576 post_norm: &lw.post_norm,
6577 ffn: gffn,
6578 });
6579 }
6580 let model = model?;
6581 let lm_gw = if upto_excl == self.num_layers
6587 && self.graph_want_logits
6588 && std::env::var("CMF_GPU_LMHEAD")
6589 .map(|v| v != "0")
6590 .unwrap_or(true)
6591 {
6592 self.weights.lm_head.graph_weight().map(|(_, i, kind, rs)| {
6593 (
6594 crate::gpu::GraphW {
6595 idx: i,
6596 kind,
6597 row_scale: rs,
6598 data: &[],
6599 },
6600 self.weights.lm_head.rows(),
6601 )
6602 })
6603 } else {
6604 None
6605 };
6606 let lm = lm_gw.as_ref().map(|(gw, rows)| (gw, *rows));
6607 let emb_gw = if steps > 1 {
6609 self.weights
6610 .embed_tokens
6611 .graph_weight()
6612 .map(|(_, i, kind, rs)| {
6613 (
6614 crate::gpu::GraphW {
6615 idx: i,
6616 kind,
6617 row_scale: rs,
6618 data: &[],
6619 },
6620 self.weights.embed_tokens.rows(),
6621 self.embed_multiplier,
6622 )
6623 })
6624 } else {
6625 None
6626 };
6627
6628 let loop_norm_at: Vec<usize> = if self.loop_final_norm {
6634 (from..upto_excl.min(self.num_layers - 1))
6635 .filter(|&li| (li + 1) % self.physical_layers == 0)
6636 .map(|li| li - from)
6637 .collect()
6638 } else {
6639 Vec::new()
6640 };
6641 let mut h = hidden.to_vec();
6642 crate::gpu::forward_token_graph(
6643 &model,
6644 self.graph_kv_id,
6645 &layers,
6646 &o1_views,
6647 self.o1_epoch,
6648 &self.inv_freq,
6649 &mut h,
6650 nh,
6651 nkv,
6652 hd,
6653 rd,
6654 self.hidden_size,
6655 self.intermediate_size,
6656 position,
6657 self.kv_cache.max_seq_len,
6658 gemma,
6659 self.rms_eps as f32,
6660 lm,
6661 &self.weights.final_norm,
6662 logits_out,
6663 &loop_norm_at,
6664 steps,
6665 emb_gw.as_ref().map(|(gw, rows, m)| (gw, *rows, *m)),
6666 ids_out,
6667 layers_run,
6668 from,
6669 false,
6670 )
6671 .then_some(h)
6672 }
6673
6674 #[cfg(target_os = "macos")]
6683 #[allow(clippy::type_complexity)]
6684 fn metal_rows_plan(
6685 &self,
6686 ) -> Option<(
6687 Vec<MetalRowsItem<'_>>,
6688 std::sync::Arc<cortiq_core::CmfModel>,
6689 Option<crate::gpu_metal::GdnGpuCfg>,
6690 )> {
6691 use crate::gpu_metal::{AttnGpuLayer, GdnGpuCfg, GdnGpuLayer, MetalFfn};
6692 if !crate::gpu::q1_force()
6693 || !crate::gpu::enabled_here()
6694 || std::env::var("CMF_GPU_BLOCK")
6695 .map(|v| v == "0")
6696 .unwrap_or(false)
6697 || self.attn_softcap > 0.0
6698 || self.o1_active()
6699 || self.swa.is_some()
6700 || self.global_attn.is_some()
6701 || self.attention_heads_per_layer.is_some()
6702 || self.attn_v_norm
6703 || self.loop_final_norm
6704 || (self.attn_scale - 1.0 / (self.head_dim as f32).sqrt()).abs() > 1e-9
6705 {
6706 return None;
6707 }
6708 let attend_contract = self.head_dim % 4 == 0
6709 && self.head_dim <= 256
6710 && self.rotary_dim >= 2
6711 && self.rotary_dim <= self.head_dim
6712 && (self.rotary_dim / 2) % 32 == 0
6713 && self.num_kv_heads > 0
6714 && self.num_heads % self.num_kv_heads == 0;
6715 if !attend_contract {
6716 return None;
6717 }
6718 let mut plan: Vec<MetalRowsItem> = Vec::new();
6719 let mut model_ref: Option<std::sync::Arc<cortiq_core::CmfModel>> = None;
6720 for li in 0..self.num_layers {
6721 let lw = &self.weights.layers[self.phys_layer(li)];
6722 if lw.attn_out_norm.is_some() || lw.ffn_out_norm.is_some() || lw.layer_scale.is_some() {
6723 return None;
6724 }
6725 let ffn = match &lw.ffn {
6726 FfnKind::Dense(d) if d.act == Act::Silu && d.segs.is_empty() => {
6727 let (Some(g), Some(u), Some(dn)) = (
6728 d.gate_proj.q1_parts(),
6729 d.up_proj.q1_parts(),
6730 d.down_proj.q1_parts(),
6731 ) else {
6732 return None;
6733 };
6734 MetalFfn::Dense {
6735 gate: g,
6736 up: u,
6737 down: dn,
6738 }
6739 }
6740 _ => return None,
6741 };
6742 match &lw.attn {
6743 AttnKind::LinearGdn(w) if self.gdn_cfg.is_some() => {
6744 let (Some(qkv), Some(z), Some(a), Some(bb), Some(out)) = (
6745 w.in_proj_qkv.q1_parts(),
6746 w.in_proj_z.q1_parts(),
6747 w.in_proj_a.f32_parts(),
6748 w.in_proj_b.f32_parts(),
6749 w.out_proj.q1_parts(),
6750 ) else {
6751 return None;
6752 };
6753 if let QTensor::Mapped { model, .. } = &w.in_proj_qkv {
6754 model_ref.get_or_insert_with(|| model.clone());
6755 }
6756 let gl = GdnGpuLayer {
6757 attn_norm: &lw.input_norm,
6758 post_norm: &lw.post_norm,
6759 qkv,
6760 z,
6761 a,
6762 b: bb,
6763 out,
6764 ffn,
6765 conv1d: &w.conv1d,
6766 a_log: &w.a_log,
6767 dt_bias: &w.dt_bias,
6768 gnorm: &w.norm,
6769 };
6770 match plan.last_mut() {
6771 Some(MetalRowsItem::Gdn { run, .. }) => run.push(gl),
6772 _ => plan.push(MetalRowsItem::Gdn {
6773 run: vec![gl],
6774 first: li,
6775 }),
6776 }
6777 }
6778 AttnKind::Full {
6779 wq,
6780 wk,
6781 wv,
6782 wo,
6783 q_norm,
6784 k_norm,
6785 output_gate,
6786 softplus_gate: None,
6787 bias: None,
6788 } => {
6789 let (Some(pq), Some(pk), Some(pv), Some(po)) =
6790 (wq.q1_parts(), wk.q1_parts(), wv.q1_parts(), wo.q1_parts())
6791 else {
6792 return None;
6793 };
6794 if let QTensor::Mapped { model, .. } = wq {
6795 model_ref.get_or_insert_with(|| model.clone());
6796 }
6797 let cache = &self.kv_cache.layers[li];
6798 if cache.mode != crate::kv_cache::KvMode::F32 || cache.o1.is_some() {
6799 return None;
6800 }
6801 plan.push(MetalRowsItem::Attn {
6802 l: AttnGpuLayer {
6803 attn_norm: &lw.input_norm,
6804 post_norm: &lw.post_norm,
6805 wq: pq,
6806 wk: pk,
6807 wv: pv,
6808 wo: po,
6809 ffn,
6810 },
6811 li,
6812 q_norm: q_norm.as_deref(),
6813 k_norm: k_norm.as_deref(),
6814 output_gate: *output_gate,
6815 });
6816 }
6817 _ => return None,
6818 }
6819 }
6820 let model = model_ref?;
6821 let gcfg = self.gdn_cfg.map(|cfg| GdnGpuCfg {
6822 nv: cfg.num_v_heads,
6823 nk: cfg.num_k_heads,
6824 dk: cfg.key_head_dim,
6825 dv: cfg.value_head_dim,
6826 kk: cfg.conv_kernel,
6827 hidden: self.hidden_size,
6828 inter: self.intermediate_size,
6829 c_dim: cfg.conv_dim(),
6830 eps: cfg.rms_eps as f32,
6831 gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
6832 });
6833 Some((plan, model, gcfg))
6834 }
6835
6836 #[cfg(target_os = "macos")]
6838 #[allow(clippy::too_many_arguments)]
6839 fn metal_attn_params<'a>(
6840 li: usize,
6841 cache: &'a crate::kv_cache::LayerKvCache,
6842 q_norm: Option<&'a [f32]>,
6843 k_norm: Option<&'a [f32]>,
6844 output_gate: bool,
6845 inv_freq: &'a [f32],
6846 geom: (usize, usize, usize, usize),
6847 pos0: usize,
6848 kv_id: u64,
6849 eps: f32,
6850 gemma: bool,
6851 ) -> (crate::gpu_metal::AttnDeviceParams<'a>, usize) {
6852 let (nh, nkv, hd, rd) = geom;
6853 let cpu_k: Vec<&[f32]> = (0..nkv).map(|g| cache.head_keys(g)).collect();
6854 let cpu_v: Vec<&[f32]> = (0..nkv).map(|g| cache.head_values(g)).collect();
6855 let cpu_stored = cpu_k[0].len() / hd;
6856 (
6857 crate::gpu_metal::AttnDeviceParams {
6858 kv_id,
6859 layer: li,
6860 nh,
6861 nkv,
6862 hd,
6863 rd,
6864 position: pos0,
6865 eps,
6866 gemma,
6867 output_gate,
6868 q_norm,
6869 k_norm,
6870 inv_freq,
6871 cpu_k,
6872 cpu_v,
6873 cpu_stored,
6874 o1: None,
6875 },
6876 cpu_stored,
6877 )
6878 }
6879
6880 #[cfg(target_os = "macos")]
6885 #[allow(clippy::type_complexity)]
6886 fn metal_rows_run(
6887 &mut self,
6888 hiddens: &mut [f32],
6889 pos0: usize,
6890 b: usize,
6891 prefill: bool,
6892 spec: Option<((usize, usize, usize), &[f32], &mut Vec<f32>)>,
6893 ) -> Option<MetalVerifyPending> {
6894 use crate::gpu_metal::{GraphDims, VerifyGraph};
6895 let want = self.gdn_cfg.map(|c| c.state_len()).unwrap_or(0);
6896 for l in &mut self.kv_cache.layers {
6897 if l.linear_state.len() != want && want > 0 {
6898 l.linear_state = vec![0f32; want];
6899 }
6900 }
6901 let (plan, model, gcfg) = self.metal_rows_plan()?;
6902 let dims = GraphDims {
6903 hidden: self.hidden_size,
6904 eps: self.rms_eps as f32,
6905 gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
6906 };
6907 let mut graph = if prefill {
6908 VerifyGraph::new_prefill(&model, dims, hiddens, b)?
6909 } else {
6910 VerifyGraph::new(&model, dims, hiddens, b)?
6911 };
6912 let geom = (
6913 self.num_heads,
6914 self.num_kv_heads,
6915 self.head_dim,
6916 self.rotary_dim,
6917 );
6918 let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
6919 let eps = self.rms_eps as f32;
6920 let kv_id = self.graph_kv_id;
6921 let inv_freq = self.inv_freq.clone();
6922 for item in &plan {
6923 let ok = match item {
6924 MetalRowsItem::Gdn { run, .. } => gcfg
6925 .as_ref()
6926 .map(|gc| run.iter().all(|l| graph.gdn_ok(l, gc)))
6927 .unwrap_or(false),
6928 MetalRowsItem::Attn {
6929 l,
6930 li,
6931 q_norm,
6932 k_norm,
6933 output_gate,
6934 } => {
6935 let (p, _) = Self::metal_attn_params(
6936 *li,
6937 &self.kv_cache.layers[*li],
6938 *q_norm,
6939 *k_norm,
6940 *output_gate,
6941 &inv_freq,
6942 geom,
6943 pos0,
6944 kv_id,
6945 eps,
6946 gemma,
6947 );
6948 graph.attn_ok(l, &p)
6949 }
6950 };
6951 if !ok {
6952 use std::sync::atomic::{AtomicBool, Ordering};
6953 static SAID: AtomicBool = AtomicBool::new(false);
6954 if !SAID.swap(true, Ordering::Relaxed) {
6955 tracing::warn!("metal rows graph: a layer failed preflight — declining");
6956 }
6957 return None;
6958 }
6959 }
6960 let lm = match &spec {
6961 Some((lm, _, _)) => {
6962 if !graph.lm_head_ok(*lm) {
6963 return None;
6964 }
6965 Some(*lm)
6966 }
6967 None => None,
6968 };
6969 let mut gdn_layers = Vec::new();
6970 let mut attn_layers = Vec::new();
6971 for item in &plan {
6972 match item {
6973 MetalRowsItem::Gdn { run, first } => {
6974 let ro: Vec<&[f32]> = self.kv_cache.layers[*first..*first + run.len()]
6975 .iter()
6976 .map(|l| l.linear_state.as_slice())
6977 .collect();
6978 if !graph.encode_gdn_run_b(run, &ro, gcfg.as_ref().unwrap()) {
6979 return None;
6980 }
6981 gdn_layers.extend(*first..*first + run.len());
6982 }
6983 MetalRowsItem::Attn {
6984 l,
6985 li,
6986 q_norm,
6987 k_norm,
6988 output_gate,
6989 } => {
6990 let (p, cpu_stored) = Self::metal_attn_params(
6991 *li,
6992 &self.kv_cache.layers[*li],
6993 *q_norm,
6994 *k_norm,
6995 *output_gate,
6996 &inv_freq,
6997 geom,
6998 pos0,
6999 kv_id,
7000 eps,
7001 gemma,
7002 );
7003 if !graph.encode_attn_b(l, &p) {
7004 return None;
7005 }
7006 attn_layers.push((*li, cpu_stored));
7007 }
7008 }
7009 }
7010 if let (Some(lm), Some((_, final_norm, _))) = (lm, spec.as_ref()) {
7011 if !graph.encode_lm_head_b(final_norm, lm) {
7012 return None;
7013 }
7014 }
7015 graph.sync();
7016 if let Some((lm, _, logits)) = spec {
7017 logits.resize(b * lm.1, 0.0);
7018 graph.read_logits(logits);
7019 }
7020 graph.read_hidden(hiddens);
7021 Some(MetalVerifyPending {
7022 graph,
7023 gdn_layers,
7024 attn_layers,
7025 })
7026 }
7027
7028 #[cfg(target_os = "macos")]
7034 fn try_batch_graph_metal(
7035 &mut self,
7036 hiddens: &mut [f32],
7037 positions: &[usize],
7038 b: usize,
7039 spec: Option<((usize, usize, usize), &[f32], &mut Vec<f32>)>,
7040 ) -> bool {
7041 let _t0 = std::time::Instant::now();
7042 if positions.len() != b
7043 || positions.windows(2).any(|w| w[1] != w[0] + 1)
7044 || hiddens.len() != b * self.hidden_size
7045 {
7046 return false;
7047 }
7048 let Some(pending) = self.metal_rows_run(hiddens, positions[0], b, false, spec) else {
7049 return false;
7050 };
7051 if std::env::var("CMF_GRAPH_SPEC_TIME").is_ok() {
7052 eprintln!(
7053 "metal-verify: {:.1} ms | b={b}",
7054 _t0.elapsed().as_secs_f64() * 1e3
7055 );
7056 }
7057 self.metal_verify = Some(pending);
7058 true
7059 }
7060
7061 #[cfg(target_os = "macos")]
7066 fn prefill_batch_metal(&mut self, ids: &[u32], start_pos: usize) -> Option<Vec<f32>> {
7067 let b = ids.len();
7068 if b == 0 || b > 512 {
7069 return None;
7070 }
7071 let hs = self.hidden_size;
7072 let mut hiddens = vec![0f32; b * hs];
7073 for (j, &id) in ids.iter().enumerate() {
7074 let e = self.embed_single(id);
7075 hiddens[j * hs..(j + 1) * hs].copy_from_slice(&e);
7076 }
7077 let mut pending = self.metal_rows_run(&mut hiddens, start_pos, b, true, None)?;
7078 let idxs = pending.gdn_layers.clone();
7080 let mut outs: Vec<&mut [f32]> = self
7081 .kv_cache
7082 .layers
7083 .iter_mut()
7084 .enumerate()
7085 .filter(|(i, _)| idxs.binary_search(i).is_ok())
7086 .map(|(_, l)| l.linear_state.as_mut_slice())
7087 .collect();
7088 pending.graph.finish_states(&mut outs);
7089 let (nkv, hd) = (self.num_kv_heads, self.head_dim);
7090 let mut kbuf = vec![0f32; b * nkv * hd];
7091 let mut vbuf = vec![0f32; b * nkv * hd];
7092 for (li, cpu_stored) in &pending.attn_layers {
7093 if crate::gpu_metal::kv_mirror_read_rows(
7094 self.graph_kv_id,
7095 *li,
7096 nkv,
7097 hd,
7098 *cpu_stored,
7099 b,
7100 &mut kbuf,
7101 &mut vbuf,
7102 ) {
7103 let cache = &mut self.kv_cache.layers[*li];
7104 for r in 0..b {
7105 cache.append(
7106 &kbuf[r * nkv * hd..(r + 1) * nkv * hd],
7107 &vbuf[r * nkv * hd..(r + 1) * nkv * hd],
7108 &[],
7109 );
7110 }
7111 crate::gpu_metal::kv_mirror_set_stored(self.graph_kv_id, *li, cpu_stored + b);
7112 }
7113 }
7114 Some(hiddens)
7115 }
7116
7117 #[cfg(target_os = "macos")]
7121 fn metal_verify_commit(&mut self, a: usize) -> bool {
7122 let Some(mut pending) = self.metal_verify.take() else {
7123 return false;
7124 };
7125 let n = a + 1;
7126 let idxs = pending.gdn_layers.clone();
7128 let mut outs: Vec<&mut [f32]> = self
7129 .kv_cache
7130 .layers
7131 .iter_mut()
7132 .enumerate()
7133 .filter(|(i, _)| idxs.binary_search(i).is_ok())
7134 .map(|(_, l)| l.linear_state.as_mut_slice())
7135 .collect();
7136 if !pending.graph.commit(n, &mut outs) {
7137 return false;
7138 }
7139 let (nkv, hd) = (self.num_kv_heads, self.head_dim);
7140 let mut kbuf = vec![0f32; n * nkv * hd];
7141 let mut vbuf = vec![0f32; n * nkv * hd];
7142 for (li, cpu_stored) in &pending.attn_layers {
7143 if crate::gpu_metal::kv_mirror_read_rows(
7144 self.graph_kv_id,
7145 *li,
7146 nkv,
7147 hd,
7148 *cpu_stored,
7149 n,
7150 &mut kbuf,
7151 &mut vbuf,
7152 ) {
7153 let cache = &mut self.kv_cache.layers[*li];
7154 for r in 0..n {
7155 cache.append(
7156 &kbuf[r * nkv * hd..(r + 1) * nkv * hd],
7157 &vbuf[r * nkv * hd..(r + 1) * nkv * hd],
7158 &[],
7159 );
7160 }
7161 crate::gpu_metal::kv_mirror_set_stored(self.graph_kv_id, *li, cpu_stored + n);
7162 }
7163 }
7164 true
7165 }
7166
7167 #[cfg(target_os = "macos")]
7173 fn mtp_warm_batch_metal(
7174 &mut self,
7175 m: &mut MtpModule,
7176 pairs: &[(&[f32], u32)],
7177 first_pos: usize,
7178 ) -> bool {
7179 use crate::gpu_metal::{AttnDeviceParams, AttnGpuLayer, GraphDims, MetalFfn, VerifyGraph};
7180 let b = pairs.len();
7181 if b == 0 || b > 512 || m.kv.mode != crate::kv_cache::KvMode::F32 || m.kv.o1.is_some() {
7182 return false;
7183 }
7184 let AttnKind::Full {
7185 wq,
7186 wk,
7187 wv,
7188 wo,
7189 q_norm,
7190 k_norm,
7191 output_gate,
7192 softplus_gate: None,
7193 bias: None,
7194 } = &m.layer.attn
7195 else {
7196 return false;
7197 };
7198 let FfnKind::Dense(d) = &m.layer.ffn else {
7199 return false;
7200 };
7201 if !d.segs.is_empty() {
7202 return false;
7203 }
7204 let (Some(pq), Some(pk), Some(pv), Some(po)) =
7205 (wq.q1_parts(), wk.q1_parts(), wv.q1_parts(), wo.q1_parts())
7206 else {
7207 return false;
7208 };
7209 let (Some(g), Some(u), Some(dn)) = (
7210 d.gate_proj.q1_parts(),
7211 d.up_proj.q1_parts(),
7212 d.down_proj.q1_parts(),
7213 ) else {
7214 return false;
7215 };
7216 let Some(eh) = m.eh_proj.q1_parts() else {
7217 return false;
7218 };
7219 let QTensor::Mapped { model, .. } = wq else {
7220 return false;
7221 };
7222 let model = model.clone();
7223 let hs = self.hidden_size;
7224 let mut cat = vec![0f32; b * 2 * hs];
7226 for (j, (h, tok)) in pairs.iter().enumerate() {
7227 let e = self.embed_single(*tok);
7228 let (ce, ch) = cat[j * 2 * hs..(j + 1) * 2 * hs].split_at_mut(hs);
7229 inference::rms_norm_into(&e, &m.enorm, self.rms_eps, self.norm_style, ce);
7230 inference::rms_norm_into(h, &m.hnorm, self.rms_eps, self.norm_style, ch);
7231 }
7232 let dims = GraphDims {
7233 hidden: hs,
7234 eps: self.rms_eps as f32,
7235 gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
7236 };
7237 let Some(mut graph) = VerifyGraph::new_via_proj(&model, dims, eh, &cat, b) else {
7238 return false;
7239 };
7240 let l = AttnGpuLayer {
7241 attn_norm: &m.layer.input_norm,
7242 post_norm: &m.layer.post_norm,
7243 wq: pq,
7244 wk: pk,
7245 wv: pv,
7246 wo: po,
7247 ffn: MetalFfn::Dense {
7248 gate: g,
7249 up: u,
7250 down: dn,
7251 },
7252 };
7253 let (nh, nkv, hd, rd) = (
7254 self.num_heads,
7255 self.num_kv_heads,
7256 self.head_dim,
7257 self.rotary_dim,
7258 );
7259 let inv_freq = self.inv_freq.clone();
7260 let cpu_stored;
7261 {
7262 let cache = &m.kv;
7263 let cpu_k: Vec<&[f32]> = (0..nkv).map(|g| cache.head_keys(g)).collect();
7264 let cpu_v: Vec<&[f32]> = (0..nkv).map(|g| cache.head_values(g)).collect();
7265 cpu_stored = cpu_k[0].len() / hd;
7266 if cpu_stored != first_pos {
7267 return false;
7268 }
7269 let p = AttnDeviceParams {
7270 kv_id: self.mtp_kv_id(),
7271 layer: Self::MTP_LAYER_BASE,
7272 nh,
7273 nkv,
7274 hd,
7275 rd,
7276 position: first_pos,
7277 eps: self.rms_eps as f32,
7278 gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
7279 output_gate: *output_gate,
7280 q_norm: q_norm.as_deref(),
7281 k_norm: k_norm.as_deref(),
7282 inv_freq: &inv_freq,
7283 cpu_k,
7284 cpu_v,
7285 cpu_stored,
7286 o1: None,
7287 };
7288 if !graph.attn_ok(&l, &p) || !graph.encode_attn_b(&l, &p) {
7289 return false;
7290 }
7291 }
7292 graph.sync();
7293 let mut kbuf = vec![0f32; b * nkv * hd];
7294 let mut vbuf = vec![0f32; b * nkv * hd];
7295 if !crate::gpu_metal::kv_mirror_read_rows(
7296 self.mtp_kv_id(),
7297 Self::MTP_LAYER_BASE,
7298 nkv,
7299 hd,
7300 cpu_stored,
7301 b,
7302 &mut kbuf,
7303 &mut vbuf,
7304 ) {
7305 return false;
7306 }
7307 for r in 0..b {
7308 m.kv.append(
7309 &kbuf[r * nkv * hd..(r + 1) * nkv * hd],
7310 &vbuf[r * nkv * hd..(r + 1) * nkv * hd],
7311 &[],
7312 );
7313 }
7314 crate::gpu_metal::kv_mirror_set_stored(
7315 self.mtp_kv_id(),
7316 Self::MTP_LAYER_BASE,
7317 cpu_stored + b,
7318 );
7319 true
7320 }
7321
7322 fn draft_vocab_rows(head_rows: usize) -> usize {
7325 static N: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
7326 let n = *N.get_or_init(|| {
7327 std::env::var("CMF_DRAFT_VOCAB")
7328 .ok()
7329 .and_then(|v| v.parse().ok())
7330 .unwrap_or(65536)
7331 });
7332 if n == 0 { head_rows } else { n.min(head_rows) }
7333 }
7334
7335 #[cfg(target_os = "macos")]
7340 fn mtp_step_metal(
7341 &mut self,
7342 m: &mut MtpModule,
7343 hidden: &[f32],
7344 next_token: u32,
7345 position: usize,
7346 want_logits: bool,
7347 ) -> Option<(Vec<f32>, Vec<f32>)> {
7348 use crate::gpu_metal::{AttnDeviceParams, AttnGpuLayer, GraphDims, MetalFfn, TokenGraph};
7349 if std::env::var("CMF_MTP_GRAPH").as_deref() == Ok("0")
7350 || !crate::gpu::q1_force()
7351 || !crate::gpu::enabled_here()
7352 || self.attn_softcap > 0.0
7353 || self.attention_heads_per_layer.is_some()
7354 || m.kv.mode != crate::kv_cache::KvMode::F32
7355 || m.kv.o1.is_some()
7356 {
7357 return None;
7358 }
7359 let AttnKind::Full {
7360 wq,
7361 wk,
7362 wv,
7363 wo,
7364 q_norm,
7365 k_norm,
7366 output_gate,
7367 softplus_gate: None,
7368 bias: None,
7369 } = &m.layer.attn
7370 else {
7371 return None;
7372 };
7373 let FfnKind::Dense(d) = &m.layer.ffn else {
7374 return None;
7375 };
7376 if d.act != Act::Silu || !d.segs.is_empty() {
7377 return None;
7378 }
7379 let (pq, pk, pv, po) = (
7380 wq.q1_parts()?,
7381 wk.q1_parts()?,
7382 wv.q1_parts()?,
7383 wo.q1_parts()?,
7384 );
7385 let (g, u, dn) = (
7386 d.gate_proj.q1_parts()?,
7387 d.up_proj.q1_parts()?,
7388 d.down_proj.q1_parts()?,
7389 );
7390 let QTensor::Mapped { model, .. } = wq else {
7391 return None;
7392 };
7393 let model = model.clone();
7394 let lm = if want_logits {
7395 Some(self.weights.lm_head.q1_parts()?)
7396 } else {
7397 None
7398 };
7399 let dims = GraphDims {
7400 hidden: self.hidden_size,
7401 eps: self.rms_eps as f32,
7402 gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
7403 };
7404 let hs = self.hidden_size;
7407 let mut x = vec![0f32; hs];
7408 let mut graph = TokenGraph::new(&model, dims, &x)?;
7409 let mut folded = false;
7410 if let Some(eh) = m.eh_proj.q1_parts() {
7411 let e = self.embed_single(next_token);
7412 let mut cat = vec![0.0f32; 2 * hs];
7413 let (cat_e, cat_h) = cat.split_at_mut(hs);
7414 inference::rms_norm_into(&e, &m.enorm, self.rms_eps, self.norm_style, cat_e);
7415 inference::rms_norm_into(hidden, &m.hnorm, self.rms_eps, self.norm_style, cat_h);
7416 folded = graph.encode_input_proj(eh, &cat);
7417 }
7418 if !folded {
7419 x = self.mtp_block_input(m, hidden, next_token);
7420 graph = TokenGraph::new(&model, dims, &x)?;
7421 }
7422 let l = AttnGpuLayer {
7423 attn_norm: &m.layer.input_norm,
7424 post_norm: &m.layer.post_norm,
7425 wq: pq,
7426 wk: pk,
7427 wv: pv,
7428 wo: po,
7429 ffn: MetalFfn::Dense {
7430 gate: g,
7431 up: u,
7432 down: dn,
7433 },
7434 };
7435 let (nh, nkv, hd, rd) = (
7436 self.num_heads,
7437 self.num_kv_heads,
7438 self.head_dim,
7439 self.rotary_dim,
7440 );
7441 let inv_freq = self.inv_freq.clone();
7442 {
7443 let cache = &m.kv;
7444 let cpu_k: Vec<&[f32]> = (0..nkv).map(|g| cache.head_keys(g)).collect();
7445 let cpu_v: Vec<&[f32]> = (0..nkv).map(|g| cache.head_values(g)).collect();
7446 let cpu_stored = cpu_k[0].len() / hd;
7447 let p = AttnDeviceParams {
7448 kv_id: self.mtp_kv_id(),
7449 layer: Self::MTP_LAYER_BASE,
7450 nh,
7451 nkv,
7452 hd,
7453 rd,
7454 position,
7455 eps: self.rms_eps as f32,
7456 gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
7457 output_gate: *output_gate,
7458 q_norm: q_norm.as_deref(),
7459 k_norm: k_norm.as_deref(),
7460 inv_freq: &inv_freq,
7461 cpu_k,
7462 cpu_v,
7463 cpu_stored,
7464 o1: None,
7465 };
7466 if !graph.attn_device_ok(&l, &p) || !graph.encode_attn_device(&l, &p) {
7467 return None;
7468 }
7469 }
7470 let draft_rows = if let Some(lm) = lm {
7476 Self::draft_vocab_rows(lm.1)
7477 } else {
7478 0
7479 };
7480 if let Some(lm) = lm {
7481 if !graph.lm_head_ok(lm) {
7482 return None;
7483 }
7484 if draft_rows < lm.1 {
7485 if !graph.encode_lm_head_part(&m.final_norm, lm, draft_rows) {
7486 return None;
7487 }
7488 } else {
7489 graph.encode_lm_head(&m.final_norm, lm);
7490 }
7491 }
7492 graph.sync();
7493 let mut logits = Vec::new();
7494 if let Some(lm) = lm {
7495 let n_read = draft_rows.min(lm.1).min(self.vocab_size);
7496 logits = attention::take_buf(n_read);
7497 graph.read_logits(&mut logits);
7498 logits.resize(self.vocab_size, f32::NEG_INFINITY);
7500 }
7501 graph.finish(&mut x);
7502 let mut krow = attention::take_buf(nkv * hd);
7503 let mut vrow = attention::take_buf(nkv * hd);
7504 if crate::gpu_metal::kv_mirror_read_last(
7505 self.mtp_kv_id(),
7506 Self::MTP_LAYER_BASE,
7507 nkv,
7508 hd,
7509 &mut krow,
7510 &mut vrow,
7511 ) {
7512 m.kv.append(&krow, &vrow, &[]);
7513 }
7514 attention::recycle_buf(&mut krow);
7515 attention::recycle_buf(&mut vrow);
7516 Some((logits, x))
7517 }
7518
7519 fn try_batch_graph_wgpu(
7520 &self,
7521 hiddens: &mut [f32],
7522 positions: &[usize],
7523 k: usize,
7524 spec: Option<crate::gpu::SpecTail<'_>>,
7525 ) -> bool {
7526 let _tb = std::time::Instant::now();
7527 if self.attn_softcap > 0.0 {
7528 return false; }
7530 if self.o1_active() {
7531 return false;
7532 }
7533 let nh = self.num_heads;
7534 let (nkv, hd, rd) = self.layer_geom(0);
7535 let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
7536 fn gw(t: &QTensor) -> Option<crate::gpu::GraphW<'_>> {
7537 if let Some((_, i, kind, rs)) = t.graph_weight() {
7538 return Some(crate::gpu::GraphW {
7539 idx: i,
7540 kind,
7541 row_scale: rs,
7542 data: &[],
7543 });
7544 }
7545 t.as_f32().map(|d| crate::gpu::GraphW {
7546 idx: 0,
7547 kind: 4,
7548 row_scale: &[],
7549 data: d,
7550 })
7551 }
7552 let built: Option<(
7553 Vec<crate::gpu::GraphLayer<'_>>,
7554 std::sync::Arc<cortiq_core::CmfModel>,
7555 )> = (|| {
7556 let mut layers = Vec::with_capacity(self.num_layers);
7557 let mut model = None;
7558 for li in 0..self.num_layers {
7559 let lw = &self.weights.layers[self.phys_layer(li)];
7560 let gffn = match &lw.ffn {
7567 FfnKind::Dense(d) if !d.segs.is_empty() => return None,
7568 FfnKind::Dense(d) => crate::gpu::GraphFfn::Dense {
7569 gate: gw(&d.gate_proj)?,
7570 up: gw(&d.up_proj)?,
7571 down: gw(&d.down_proj)?,
7572 },
7573 FfnKind::Moe(m) => {
7574 if m.router_sigmoid
7575 || m.expert_bias.is_some()
7576 || m.route_tau.is_some()
7577 || m.mask.is_some()
7578 {
7579 return None;
7580 }
7581 let (se, sg) = m.shared.as_ref()?;
7582 let sgate = gw(sg.as_ref()?)?;
7583 let router = gw(&m.router)?;
7584 let inter = m.experts.first()?.gate_proj.rows();
7585 let mut experts = Vec::with_capacity(m.experts.len() + 1);
7586 let mut q4tp: Option<bool> = None;
7587 let mut gu_q2: Option<bool> = None;
7588 for e in m.experts.iter().chain(std::iter::once(se)) {
7589 if !matches!(e.act, Act::Silu)
7590 || e.gate_proj.rows() != inter
7591 || e.up_proj.rows() != inter
7592 {
7593 return None;
7594 }
7595 let (mm, gi, ui, di, is_p, is_q2) = match e.gate_proj.mapped_q4t() {
7599 Some((mm, gi)) => (
7600 mm,
7601 gi,
7602 e.up_proj.mapped_q4t()?.1,
7603 e.down_proj.mapped_q4t()?.1,
7604 false,
7605 false,
7606 ),
7607 None => match e.gate_proj.mapped_q2tp() {
7608 Some((mm, gi)) => (
7609 mm,
7610 gi,
7611 e.up_proj.mapped_q2tp()?.1,
7612 e.down_proj.mapped_q4tp()?.1,
7613 true,
7614 true,
7615 ),
7616 None => {
7617 let (mm, gi) = e.gate_proj.mapped_q4tp()?;
7618 (
7619 mm,
7620 gi,
7621 e.up_proj.mapped_q4tp()?.1,
7622 e.down_proj.mapped_q4tp()?.1,
7623 true,
7624 false,
7625 )
7626 }
7627 },
7628 };
7629 if *q4tp.get_or_insert(is_p) != is_p
7630 || *gu_q2.get_or_insert(is_q2) != is_q2
7631 {
7632 return None;
7633 }
7634 model.get_or_insert_with(|| mm.clone());
7635 experts.push((gi, ui, di));
7636 }
7637 crate::gpu::GraphFfn::Moe {
7638 router,
7639 shared_gate: sgate,
7640 experts,
7641 n_exp: m.experts.len(),
7642 top_k: m.top_k,
7643 inter,
7644 norm_topk: m.norm_topk_prob,
7645 q4tp: q4tp?,
7646 gu_q2: gu_q2.unwrap_or(false),
7647 sigmoid: false,
7648 bias: None,
7649 has_shared: true,
7650 }
7651 }
7652 _ => return None,
7653 };
7654 let attn = match &lw.attn {
7655 AttnKind::Full {
7656 wq,
7657 wk,
7658 wv,
7659 wo,
7660 q_norm,
7661 k_norm,
7662 output_gate,
7663 softplus_gate,
7664 bias,
7665 } => {
7666 if softplus_gate.is_some() || self.attention_heads_per_layer.is_some() {
7667 return None;
7668 }
7669 let (m, _, _, _) = wq.graph_weight()?;
7670 model = Some(m.clone());
7671 crate::gpu::GraphAttn::Full {
7672 wq: gw(wq)?,
7673 wk: gw(wk)?,
7674 wv: gw(wv)?,
7675 wo: gw(wo)?,
7676 q_norm: q_norm.as_deref(),
7677 k_norm: k_norm.as_deref(),
7678 bias: bias
7679 .as_ref()
7680 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
7681 output_gate: *output_gate,
7682 cpu_k: self.kv_cache.layers[li].k_heads(),
7683 cpu_v: self.kv_cache.layers[li].v_heads(),
7684 }
7685 }
7686 AttnKind::LinearGdn(w) => {
7687 let cfg = self.gdn_cfg?;
7688 let (m, _, _, _) = w.in_proj_qkv.graph_weight()?;
7689 model = Some(m.clone());
7690 crate::gpu::GraphAttn::Gdn {
7691 qkv: gw(&w.in_proj_qkv)?,
7692 z: gw(&w.in_proj_z)?,
7693 a: gw(&w.in_proj_a)?,
7694 b: gw(&w.in_proj_b)?,
7695 out: gw(&w.out_proj)?,
7696 conv1d: &w.conv1d,
7697 a_log: &w.a_log,
7698 dt_bias: &w.dt_bias,
7699 norm: &w.norm,
7700 nv: cfg.num_v_heads,
7701 nk: cfg.num_k_heads,
7702 dk: cfg.key_head_dim,
7703 dv: cfg.value_head_dim,
7704 kk: cfg.conv_kernel,
7705 cpu_state: &self.kv_cache.layers[self.phys_layer(li)].linear_state,
7706 }
7707 }
7708 _ => return None,
7709 };
7710 layers.push(crate::gpu::GraphLayer {
7711 input_norm: &lw.input_norm,
7712 attn,
7713 post_norm: &lw.post_norm,
7714 ffn: gffn,
7715 });
7716 }
7717 Some((layers, model?))
7718 })();
7719 let Some((layers, model)) = built else {
7720 {
7721 use std::sync::atomic::{AtomicBool, Ordering};
7722 static SAID: AtomicBool = AtomicBool::new(false);
7723 if !SAID.swap(true, Ordering::Relaxed) {
7724 tracing::warn!("batch graph: BUILDER refused (layer weights/kinds)");
7725 }
7726 }
7727 return false;
7728 };
7729 if std::env::var("CMF_GRAPH_SPEC_TIME").is_ok() {
7730 eprintln!("batch-build: {:.1} ms", _tb.elapsed().as_secs_f64() * 1e3);
7731 }
7732 crate::gpu::forward_batch_graph(
7733 &model,
7734 self.graph_kv_id,
7735 &layers,
7736 &self.inv_freq,
7737 hiddens,
7738 nh,
7739 nkv,
7740 hd,
7741 rd,
7742 self.hidden_size,
7743 self.intermediate_size,
7744 positions,
7745 self.kv_cache.max_seq_len,
7746 gemma,
7747 self.rms_eps as f32,
7748 k,
7749 spec,
7750 )
7751 }
7752
7753 fn draft_probe() -> bool {
7757 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
7758 *ON.get_or_init(|| std::env::var("CMF_DSV4_DRAFT_PROBE").is_ok_and(|v| v != "0"))
7759 }
7760
7761 #[cfg(feature = "gpu")]
7773 fn dsv4_spec_on() -> bool {
7774 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
7775 *ON.get_or_init(|| {
7776 if let Ok(v) = std::env::var("CMF_DSV4_SPEC_RUN") {
7780 return v != "0";
7781 }
7782 std::env::var("CMF_DSV4_SPEC")
7789 .map(|v| v != "0")
7790 .unwrap_or_else(|_| {
7791 crate::gpu_wgpu::DRAFT_RESERVE.load(std::sync::atomic::Ordering::Relaxed) > 0
7792 })
7793 })
7794 }
7795
7796 #[cfg(feature = "gpu")]
7803 fn dsv4_spec_step(
7804 &mut self,
7805 tip_token: u32,
7806 t_next: u32,
7807 next_pos: usize,
7808 max_extra: usize,
7809 drafted: &mut usize,
7810 accepted_ctr: &mut usize,
7811 ) -> Option<(Vec<u32>, usize)> {
7812 let t_all = std::time::Instant::now();
7813 if std::env::var("CMF_DSV4_SPEC_TIME").is_ok() {
7814 thread_local! {
7815 static LAST: std::cell::Cell<Option<std::time::Instant>> =
7816 const { std::cell::Cell::new(None) };
7817 }
7818 LAST.with(|l| {
7819 if let Some(prev) = l.get() {
7820 eprintln!(
7821 "между раундами {:.1} мс",
7822 prev.elapsed().as_secs_f64() * 1e3
7823 );
7824 }
7825 l.set(Some(std::time::Instant::now()));
7826 });
7827 }
7828 if std::env::var("CMF_DSV4_SPEC_DEBUG").is_ok() {
7829 eprintln!("spec_step: вход pos={next_pos}");
7830 }
7831 let n_layers = self.dsv4.as_ref().map(|b| b.1.len())?;
7832 let cfg = self.dsv4.as_ref().map(|b| b.2)?;
7833 if self.dspark.is_none() {
7835 let t = crate::dsv4::dspark_targets(&self.dsv4_mtp, &cfg, n_layers);
7836 if t.is_empty() {
7837 return None;
7838 }
7839 crate::dsv4::dspark_arm(&t, cfg.dim);
7840 self.dspark = Some(crate::dsv4::DsparkState::new(
7841 self.dsv4_mtp.len(),
7842 &cfg,
7843 t.len(),
7844 ));
7845 }
7846 let targets = crate::dsv4::dspark_targets(&self.dsv4_mtp, &cfg, n_layers);
7847 let pack = crate::dsv4::dspark_pack_get(&self.dsv4_mtp, &cfg);
7848 if pack.is_none() && std::env::var("CMF_DSV4_SPEC_DEBUG").is_ok() {
7849 eprintln!("spec_step: пак не построился (targets {targets:?})");
7850 }
7851 let pack = pack?;
7852 let block = crate::dsv4::dspark_block();
7853 let b_box = self.dsv4.as_mut()?;
7854 let (g, layers, st) = (&b_box.0, &b_box.1, &mut b_box.3);
7855 let ds = self.dspark.as_mut()?;
7856 let dbg = std::env::var("CMF_DSV4_SPEC_DEBUG").is_ok();
7859 if !crate::dsv4::dspark_take(&mut ds.main_hidden) && !ds.have_hidden {
7860 if dbg {
7861 eprintln!("spec_step: нет захвата");
7862 }
7863 return None;
7864 }
7865 ds.have_hidden = true;
7866 let tip_pos = next_pos.checked_sub(1)?;
7867 let draft_started = std::time::Instant::now();
7868 let mut conf = Vec::new();
7869 let props = crate::dsv4::dspark_draft_gpu(
7870 g,
7871 &self.dsv4_mtp,
7872 &cfg,
7873 ds,
7874 pack,
7875 st.kv_id,
7876 tip_token,
7877 tip_pos,
7878 self.pool.as_deref(),
7879 &mut conf,
7880 );
7881 self.dspark_draft_ns += draft_started.elapsed().as_nanos();
7882 *drafted += block;
7883 if props.is_empty() || props[0] != t_next {
7884 if dbg {
7885 eprintln!(
7886 "spec_step: черновик {} (props0={:?} t_next={t_next})",
7887 if props.is_empty() {
7888 "пуст"
7889 } else {
7890 "мимо"
7891 },
7892 props.first()
7893 );
7894 }
7895 return None;
7896 }
7897 let mut k_verify = crate::dsv4::dspark_verify_k()
7904 .min(props.len())
7905 .min(max_extra.saturating_add(1));
7906 let conf_min = {
7912 static M: std::sync::OnceLock<f32> = std::sync::OnceLock::new();
7913 *M.get_or_init(|| {
7914 std::env::var("CMF_DSPARK_CONF_MIN")
7915 .ok()
7916 .and_then(|v| v.parse().ok())
7917 .unwrap_or(0.0)
7918 })
7919 };
7920 if conf_min > 0.0 && conf.len() >= props.len() {
7921 let mut keep = 1usize;
7922 while keep < k_verify && conf.get(keep).copied().unwrap_or(0.0) >= conf_min {
7923 keep += 1;
7924 }
7925 k_verify = k_verify.min(keep.max(2));
7926 }
7927 if k_verify < 2 {
7928 return None;
7929 }
7930 let mut fed = Vec::with_capacity(k_verify);
7931 fed.push(t_next);
7932 fed.extend_from_slice(&props[1..k_verify]);
7933 let mut argmax = Vec::new();
7934 let mut logits_all = Vec::new();
7935 let mut walked = Vec::new();
7936 let txn = crate::dsv4::dsv4_verify_chunk(
7937 g,
7938 layers,
7939 &cfg,
7940 st,
7941 &fed,
7942 next_pos,
7943 &self.inv_freq,
7944 self.pool.as_deref(),
7945 &targets,
7946 &mut argmax,
7947 &mut logits_all,
7948 &mut walked,
7949 );
7950 if txn.is_none() && dbg {
7951 eprintln!("spec_step: verify отказал");
7952 }
7953 let txn = txn?;
7954 let spec_gpu_end = txn.gpu_end;
7955 let b = fed.len();
7956 let mut accepted = 1usize;
7957 while accepted < b && fed[accepted] == argmax[accepted - 1] {
7958 accepted += 1;
7959 }
7960 if std::env::var("CMF_DSV4_SPEC_FORCE_REJECT").is_ok_and(|v| v != "0") {
7965 accepted = 1;
7966 }
7967 if std::env::var("CMF_DSV4_SPEC_TRACE").is_ok() {
7968 eprintln!("spec@{next_pos}: fed={fed:?} argmax={argmax:?} accepted={accepted}");
7969 }
7970 let t_fin = std::time::Instant::now();
7971 if !crate::dsv4::dsv4_spec_finish(
7972 g,
7973 layers,
7974 &cfg,
7975 st,
7976 txn,
7977 accepted,
7978 &fed,
7979 &self.inv_freq,
7980 self.pool.as_deref(),
7981 ) {
7982 tracing::warn!("dsv4: спекулятивный откат не удался — состояние подозрительно");
7983 return None;
7984 }
7985 if std::env::var("CMF_DSV4_SPEC_TIME").is_ok() {
7986 eprintln!(
7987 "finish(k={accepted}): {:.1} мс",
7988 t_fin.elapsed().as_secs_f64() * 1e3
7989 );
7990 }
7991 *accepted_ctr += accepted - 1;
7992 let (hc, dim) = (cfg.hc_mult, cfg.dim);
7997 let dev_caps: Vec<usize> = targets
8002 .iter()
8003 .copied()
8004 .filter(|&t| t < spec_gpu_end)
8005 .collect();
8006 let mut caps_all = vec![0.0f32; dev_caps.len() * b * hc * dim];
8007 if !crate::gpu_wgpu::dsv4_spec_cap_read_all(b, dev_caps.len(), hc * dim, &mut caps_all) {
8008 return None;
8009 }
8010 for t in 0..accepted {
8011 let tip = t + 1 == accepted;
8012 for (slot, &tl) in targets.iter().enumerate() {
8013 if let Some(di) = dev_caps.iter().position(|&d| d == tl) {
8014 let lo = (di * b + t) * hc * dim;
8015 crate::dsv4::dspark_capture(
8016 &caps_all[lo..lo + hc * dim],
8017 &cfg,
8018 slot,
8019 &mut ds.main_hidden,
8020 );
8021 } else if tip
8022 && crate::dsv4::dspark_peek_slot(slot, dim, {
8023 let lo = slot * dim;
8024 &mut ds.main_hidden[lo..lo + dim]
8025 })
8026 {
8027 } else {
8032 crate::dsv4::dspark_capture(
8036 &walked[t * hc * dim..(t + 1) * hc * dim],
8037 &cfg,
8038 slot,
8039 &mut ds.main_hidden,
8040 );
8041 }
8042 }
8043 crate::dsv4::dspark_ring_append(
8044 g,
8045 &self.dsv4_mtp,
8046 &cfg,
8047 ds,
8048 next_pos + t,
8049 self.pool.as_deref(),
8050 );
8051 }
8052 let row = logits_all[(accepted - 1) * cfg.vocab..accepted * cfg.vocab].to_vec();
8053 self.graph_logits = Some(row);
8054 if std::env::var("CMF_DSV4_TRUNK_PICK_DUMP").is_ok() {
8059 crate::dsv4::trunk_freq_note(&crate::dsv4::pick_tally_take());
8060 crate::dsv4::pick_tally_arm();
8061 }
8062 if std::env::var("CMF_DSV4_SPEC_TIME").is_ok() {
8063 eprintln!(
8064 "spec_step total {:.1} мс (k={accepted})",
8065 t_all.elapsed().as_secs_f64() * 1e3
8066 );
8067 }
8068 Some((fed[1..accepted].to_vec(), next_pos + accepted))
8069 }
8070
8071 fn dspark_probe(&mut self, position: usize, token_id: u32) {
8072 if self.dsv4_mtp.is_empty() || !Self::draft_probe() {
8073 return;
8074 }
8075 let trunk_now = crate::dsv4::pick_tally_take();
8077 crate::dsv4::trunk_freq_note(&trunk_now);
8078 if !trunk_now.is_empty() {
8079 self.dspark_trunk_picks.push(trunk_now);
8080 let keep = crate::dsv4::dspark_block();
8081 if self.dspark_trunk_picks.len() > keep {
8082 self.dspark_trunk_picks.remove(0);
8083 }
8084 }
8085 for p in std::mem::take(&mut self.dspark_pending) {
8088 let Some(i) = position.checked_sub(p.0 + 1) else {
8089 continue;
8090 };
8091 let mut p = p;
8092 if i < p.1.len() {
8093 if p.2 && p.1[i] == token_id {
8094 p.3 = i + 1;
8095 } else {
8096 p.2 = false;
8097 }
8098 if i + 1 < p.1.len() {
8099 self.dspark_pending.push(p);
8100 continue;
8101 }
8102 }
8103 self.dspark_hist.push(p.3);
8104 self.dspark_real.push(token_id);
8105 }
8106 let Some(b) = &mut self.dsv4 else { return };
8107 let (g, layers, cfg) = (&b.0, &b.1, b.2);
8108 let n_layers = layers.len();
8109 if self.dspark.is_none() {
8110 let t = crate::dsv4::dspark_targets(&self.dsv4_mtp, &cfg, n_layers);
8111 if t.is_empty() {
8112 return;
8113 }
8114 eprintln!(
8115 "DSpark: захват со слоёв {t:?}, блок {}",
8116 crate::dsv4::dspark_block()
8117 );
8118 crate::dsv4::dspark_arm(&t, cfg.dim);
8119 self.dspark = Some(crate::dsv4::DsparkState::new(
8120 self.dsv4_mtp.len(),
8121 &cfg,
8122 t.len(),
8123 ));
8124 }
8125 let ds = self.dspark.as_mut().unwrap();
8126 if !crate::dsv4::dspark_take(&mut ds.main_hidden) {
8127 return; }
8129 let mut conf = Vec::new();
8130 crate::dsv4::pick_tally_arm();
8131 let draft_started = std::time::Instant::now();
8136 #[cfg(feature = "gpu")]
8137 let gpu_draft = crate::dsv4::dspark_gpu_on();
8138 #[cfg(not(feature = "gpu"))]
8139 let gpu_draft = false;
8140 let props = if gpu_draft {
8141 #[cfg(feature = "gpu")]
8142 {
8143 let kv_id = b.3.kv_id;
8144 match crate::dsv4::dspark_pack_get(&self.dsv4_mtp, &cfg) {
8145 Some(pk) => crate::dsv4::dspark_draft_gpu(
8146 g,
8147 &self.dsv4_mtp,
8148 &cfg,
8149 ds,
8150 pk,
8151 kv_id,
8152 token_id,
8153 position,
8154 self.pool.as_deref(),
8155 &mut conf,
8156 ),
8157 None => Vec::new(),
8158 }
8159 }
8160 #[cfg(not(feature = "gpu"))]
8161 Vec::new()
8162 } else {
8163 crate::gpu::cpu_scope(|| {
8164 crate::dsv4::dspark_draft(
8165 g,
8166 &self.dsv4_mtp,
8167 &cfg,
8168 ds,
8169 token_id,
8170 position,
8171 self.pool.as_deref(),
8172 &mut conf,
8173 )
8174 })
8175 };
8176 self.dspark_draft_ns += draft_started.elapsed().as_nanos();
8177 let draft_picks = crate::dsv4::pick_tally_take();
8178 crate::dsv4::dspark_freq_note(&draft_picks);
8179 crate::dsv4::pick_tally_arm();
8182 if !props.is_empty() {
8183 let (tu, tt) = {
8187 let flat: Vec<(usize, Vec<usize>)> = self
8188 .dspark_trunk_picks
8189 .iter()
8190 .flat_map(|v| v.iter().cloned())
8191 .collect();
8192 let mut per: std::collections::HashMap<usize, Vec<usize>> =
8194 std::collections::HashMap::new();
8195 for (li, picks) in flat {
8196 per.entry(li).or_default().extend(picks);
8197 }
8198 let n = per.len().max(1);
8199 let mut u = 0usize;
8200 let mut t = 0usize;
8201 for (_, v) in per {
8202 t += v.len();
8203 u += v.iter().collect::<std::collections::HashSet<_>>().len();
8204 }
8205 (u / n, t / n)
8206 };
8207 let (du, dt) = crate::dsv4::tally_unique(&draft_picks);
8208 self.dspark_exp.push((tu, tt, du, dt));
8209 self.dspark_pending.push((position, props, true, 0));
8210 }
8211 if self.dspark_hist.len() >= 8 && self.dspark_hist.len() % 8 == 0 {
8212 let n = self.dspark_hist.len() as f32;
8213 let mean: f32 = self.dspark_hist.iter().sum::<usize>() as f32 / n;
8214 let block = crate::dsv4::dspark_block();
8215 let mut at = vec![0usize; block + 1];
8216 for &k in &self.dspark_hist {
8217 at[k] += 1;
8218 }
8219 let mut surv = Vec::with_capacity(block);
8221 for i in 1..=block {
8222 let k = at[i..].iter().sum::<usize>() as f32 / n;
8223 surv.push(format!("{k:.2}"));
8224 }
8225 let distinct = self
8226 .dspark_real
8227 .iter()
8228 .collect::<std::collections::HashSet<_>>()
8229 .len();
8230 let (tu, tt, du, dt) = self.dspark_exp.iter().fold((0, 0, 0, 0), |a, b| {
8231 (a.0 + b.0, a.1 + b.1, a.2 + b.2, a.3 + b.3)
8232 });
8233 let m = self.dspark_exp.len().max(1);
8234 eprintln!(
8235 "DSpark: черновиков {}, принято в среднем {mean:.2} из {block} \
8236 (токенов за проход {:.2}), распределение {at:?}, выживание [{}]",
8237 self.dspark_hist.len(),
8238 mean + 1.0,
8239 surv.join(" ")
8240 );
8241 eprintln!(
8242 "DSpark: разных токенов {distinct} из {} (вырожденность), \
8243 эксперты ствол {}/{} на слой за {block} токенов, \
8244 черновик {}/{} за блок, draft {:.2} мс/блок",
8245 self.dspark_real.len(),
8246 tu / m,
8247 tt / m,
8248 du / m,
8249 dt / m,
8250 self.dspark_draft_ns as f64 / self.dspark_exp.len().max(1) as f64 / 1e6
8251 );
8252 }
8253 }
8254
8255 fn forward_layers_upto(
8256 &mut self,
8257 hidden: &[f32],
8258 position: usize,
8259 task_mask: Option<&TaskMask>,
8260 upto: Option<usize>,
8261 ) -> Vec<f32> {
8262 if let Some(plan) = self.gpu_plan.clone() {
8268 if upto.is_none() && plan.len() > 1 {
8269 let mut h = hidden.to_vec();
8270 for &(dev, from, upto_incl) in plan.iter() {
8271 h = crate::gpu::with_device(dev, || {
8272 self.forward_layers_span(&h, position, task_mask, from, Some(upto_incl))
8273 });
8274 }
8275 return h;
8276 }
8277 }
8278 self.forward_layers_span(hidden, position, task_mask, 0, upto)
8279 }
8280
8281 pub fn set_gpu_plan(&mut self, devices: Option<&[usize]>) -> Result<(), String> {
8286 self.set_gpu_plan_at(devices, None)
8287 }
8288
8289 pub fn set_gpu_plan_at(
8293 &mut self,
8294 devices: Option<&[usize]>,
8295 at: Option<usize>,
8296 ) -> Result<(), String> {
8297 let Some(devs) = devices.filter(|d| d.len() > 1) else {
8298 self.gpu_plan = None;
8299 return Ok(());
8300 };
8301 self.split_supported()?;
8302 let n = self.num_layers;
8303 if devs.len() > n {
8304 return Err(format!("{} devices for {n} layers", devs.len()));
8305 }
8306 if let Some(k) = at {
8307 if k == 0 || k >= n {
8308 return Err(format!("split at {k}: the model has {n} layers"));
8309 }
8310 if devs.len() == 2 {
8311 self.gpu_plan = Some(std::sync::Arc::new(vec![
8312 (devs[0], 0, k - 1),
8313 (devs[1], k, n - 1),
8314 ]));
8315 return Ok(());
8316 }
8317 return Err(format!(
8318 "an explicit split point takes exactly 2 devices, got {}",
8319 devs.len()
8320 ));
8321 }
8322 let per = n.div_ceil(devs.len());
8323 let mut plan = Vec::with_capacity(devs.len());
8324 let mut from = 0usize;
8325 for &d in devs {
8326 if from >= n {
8327 break;
8328 }
8329 let upto = (from + per - 1).min(n - 1);
8330 plan.push((d, from, upto));
8331 from = upto + 1;
8332 }
8333 self.gpu_plan = Some(std::sync::Arc::new(plan));
8334 Ok(())
8335 }
8336
8337 pub fn gpu_plan(&self) -> Option<Vec<(usize, usize, usize)>> {
8339 self.gpu_plan.as_ref().map(|p| p.as_ref().clone())
8340 }
8341
8342 fn forward_layers_span(
8348 &mut self,
8349 hidden: &[f32],
8350 position: usize,
8351 task_mask: Option<&TaskMask>,
8352 from: usize,
8353 upto: Option<usize>,
8354 ) -> Vec<f32> {
8355 debug_assert!(
8356 from == 0 || (self.dsv4.is_none() && self.qwen4_exp.is_none() && self.g3n.is_none())
8357 );
8358 if let Some(b) = &mut self.qwen4_exp {
8359 let _ = (task_mask, upto);
8360 let token_id = hidden.first().copied().unwrap_or(0.0) as u32;
8361 let mut logits = Vec::new();
8362 crate::qwen4_exp::forward_token(
8363 &b.0,
8364 &b.1,
8365 &b.2,
8366 &mut b.3,
8367 token_id,
8368 position,
8369 &self.inv_freq,
8370 self.pool.as_deref(),
8371 &mut logits,
8372 true,
8373 );
8374 self.graph_logits = Some(logits);
8375 return vec![0.0; self.hidden_size];
8376 }
8377 if let Some(b) = &mut self.dsv4 {
8383 let _ = (task_mask, upto);
8384 let token_id = hidden.first().copied().unwrap_or(0.0) as u32;
8385 let (g, layers, cfg, st) = (&b.0, &b.1, b.2, &mut b.3);
8386 st.pos = position;
8387 let mut logits = Vec::new();
8388 crate::dsv4::forward_token(
8389 g,
8390 layers,
8391 &cfg,
8392 st,
8393 token_id,
8394 &self.inv_freq,
8395 self.pool.as_deref(),
8396 &mut logits,
8397 );
8398 self.graph_logits = Some(logits);
8399 self.dspark_probe(position, token_id);
8400 return vec![0.0; self.hidden_size];
8403 }
8404 if let Some(b) = &self.g3n {
8407 let _ = (task_mask, upto);
8408 return crate::g3n::g3n_forward(
8409 &b.0,
8410 &b.1,
8411 hidden,
8412 position,
8413 &mut self.kv_cache.layers,
8414 self.num_heads,
8415 self.num_kv_heads,
8416 self.head_dim,
8417 self.pool.as_deref(),
8418 );
8419 }
8420 let mut h = hidden.to_vec();
8421 let (nh, _nkv, _hd, hs, _rd, eps) = (
8424 self.num_heads,
8425 self.num_kv_heads,
8426 self.head_dim,
8427 self.hidden_size,
8428 self.rotary_dim,
8429 self.rms_eps,
8430 );
8431 let pool = self.pool.clone();
8432 let graph_env = std::env::var("CMF_GPU_WGPU_GRAPH").ok();
8444 let graph_on = match graph_env.as_deref() {
8445 Some("0") => false,
8446 Some("prefill") => false, Some(_) => true,
8448 None => crate::gpu::wgpu_graph_default(),
8454 };
8455 let graph_trusted =
8456 graph_env.is_some() || crate::gpu::wgpu_graph_default() || self.gdn_cfg.is_some();
8457 let race_eligible = graph_on
8458 && upto.is_none()
8459 && task_mask.is_none()
8460 && from == 0
8461 && !crate::gpu::graph_unsupported();
8462 let mut tail_start = 0usize;
8463 if race_eligible && crate::gpu::graph_race_use_graph(graph_trusted) {
8464 let t_graph = std::time::Instant::now();
8465 let mut lg = Vec::new();
8466 let mut gl = 0usize;
8467 let built = self.try_token_graph_wgpu(hidden, position, &mut lg, &mut gl);
8468 if built.is_none() && !self.o1_active() && self.attn_softcap == 0.0 {
8473 crate::gpu::graph_mark_unsupported();
8474 }
8475 graph_note(built.is_some());
8476 if let Some(hh) = built {
8477 let dur = t_graph.elapsed();
8478 if std::env::var("CMF_GRAPH_PROF").is_ok() {
8479 eprintln!("graph-call: {:.2} ms total", dur.as_secs_f64() * 1000.0);
8480 }
8481 if gl > 0 && gl < self.num_layers {
8482 h = hh;
8488 tail_start = gl;
8489 } else if graph_trusted || !crate::gpu::graph_race_first_token_hopeless(dur) {
8490 if !graph_trusted {
8491 crate::gpu::graph_race_record(true, dur);
8492 }
8493 if !lg.is_empty() {
8494 lg.resize(self.vocab_size, 0.0);
8497 if let Some(c) = self.final_softcap {
8498 for l in lg.iter_mut() {
8499 *l = c * (*l / c).tanh();
8500 }
8501 }
8502 self.graph_logits = Some(lg);
8503 }
8504 return hh;
8505 }
8506 }
8512 }
8513 let span = from > 0 || upto.is_some();
8537 if span && graph_on && task_mask.is_none() && graph_trusted {
8538 let upto_excl = upto.map_or(self.num_layers, |u| u + 1);
8539 let mut lg = Vec::new();
8540 let mut gl = 0usize;
8541 let span_res =
8542 self.try_token_graph_wgpu_span(hidden, position, &mut lg, from, upto_excl, &mut gl);
8543 graph_note(span_res.is_some() && gl == upto_excl - from);
8544 if std::env::var("CMF_GPU_DEBUG").is_ok() {
8545 static SEEN: std::sync::atomic::AtomicU32 = std::sync::atomic::AtomicU32::new(0);
8549 if SEEN.fetch_add(1, std::sync::atomic::Ordering::Relaxed) < 4 {
8550 eprintln!(
8551 "span graph: covered {gl} of {} layers [{from}..{upto_excl}) res={}",
8552 upto_excl - from,
8553 span_res.is_some()
8554 );
8555 }
8556 }
8557 if let Some(hh) = span_res {
8558 if gl == upto_excl - from {
8559 if !lg.is_empty() {
8560 lg.resize(self.vocab_size, 0.0);
8561 if let Some(c) = self.final_softcap {
8562 for l in lg.iter_mut() {
8563 *l = c * (*l / c).tanh();
8564 }
8565 }
8566 self.graph_logits = Some(lg);
8567 }
8568 crate::gpu::set_layer(-1);
8569 return hh;
8570 }
8571 h = hh;
8573 tail_start = from + gl;
8574 }
8575 }
8576 let t_race_cpu = (race_eligible && !graph_trusted).then(std::time::Instant::now);
8577
8578 #[cfg(target_os = "macos")]
8579 let mut gpu_skip_until = 0usize;
8580 for li in tail_start.max(from)..self.num_layers {
8581 crate::gpu::set_layer(li as i64); if let Some(u) = upto {
8583 if li > u {
8584 break;
8585 }
8586 }
8587 if let Some(mask) = task_mask {
8588 if !mask.layer_alive(li) {
8589 continue; }
8591 }
8592 #[cfg(target_os = "macos")]
8596 {
8597 if li < gpu_skip_until {
8598 continue;
8599 }
8600 if task_mask.is_none() {
8601 let end = self.q1_graph_gpu(li, upto, position, &mut h);
8602 if end > li {
8603 gpu_skip_until = end;
8604 if self.is_loop_end(end - 1) && end < self.num_layers {
8607 h = inference::rms_norm(
8608 &h,
8609 &self.weights.final_norm,
8610 self.rms_eps,
8611 self.norm_style,
8612 );
8613 }
8614 continue;
8615 }
8616 }
8617 }
8618
8619 let lw = &self.weights.layers[self.phys_layer(li)];
8620 if let Ok(tp) = std::env::var("CMF_TRACE_POS") {
8621 if tp.parse::<usize>().ok() == Some(position) {
8622 let n: f32 = h.iter().map(|x| x * x).sum::<f32>().sqrt();
8623 eprintln!(
8624 "TRACE pos {position} layer {li}: |h| = {n:.6} h0 {:.6} h1 {:.6}",
8625 h[0], h[1]
8626 );
8627 }
8628 }
8629 inference::rms_norm_into(
8632 &h,
8633 &lw.input_norm,
8634 self.rms_eps,
8635 self.norm_style,
8636 &mut self.ws.n1,
8637 );
8638
8639 let attn_out = match &lw.attn {
8640 AttnKind::Mla(w) => {
8641 let inv_freq_l = self.layer_inv_freq(li);
8642 let rs = self.layer_rope_scale(li);
8643 let eps = self.rms_eps;
8644 let pool = self.pool.clone();
8645 mla_attention(
8646 w,
8647 &self.ws.n1,
8648 &mut self.kv_cache.layers[li],
8649 position,
8650 &inv_freq_l,
8651 rs,
8652 eps,
8653 pool.as_deref(),
8654 )
8655 }
8656 AttnKind::Linear(w) => {
8657 let cfg = self.vmf_cfg.expect("linear layer without vmf_cfg");
8658 vmf_phase_forward(
8659 &self.ws.n1,
8660 w,
8661 &cfg,
8662 &mut self.kv_cache.layers[li].linear_state,
8663 self.pool.as_deref(),
8664 )
8665 }
8666 AttnKind::Kda(w) => {
8667 let cfg = self.kda_cfg.expect("kda layer without kda_cfg");
8668 crate::linear_core::kda_forward(
8669 &self.ws.n1,
8670 w,
8671 &cfg,
8672 &mut self.kv_cache.layers[li].linear_state,
8673 self.pool.as_deref(),
8674 )
8675 }
8676 AttnKind::LinearGdn(w) => {
8677 let cfg = self.gdn_cfg.expect("gdn layer without gdn_cfg");
8678 gdn_forward(
8679 &self.ws.n1,
8680 w,
8681 &cfg,
8682 &mut self.kv_cache.layers[li].linear_state,
8683 self.pool.as_deref(),
8684 )
8685 }
8686 AttnKind::ShortConv(w) => {
8687 let cfg = self
8688 .short_conv_cfg
8689 .expect("short-conv layer without short_conv_cfg");
8690 short_conv_forward(
8691 &self.ws.n1,
8692 w,
8693 &cfg,
8694 &mut self.kv_cache.layers[li].linear_state,
8695 self.pool.as_deref(),
8696 )
8697 }
8698 AttnKind::Full {
8699 wq,
8700 wk,
8701 wv,
8702 wo,
8703 q_norm,
8704 k_norm,
8705 output_gate,
8706 softplus_gate,
8707 bias,
8708 } if self.kv_cache.layers[li].o1_sealed() => {
8709 let inv_freq_l = self.layer_inv_freq(li);
8712 let (nkv_l, hd_l, rd_l) = self.layer_geom(li);
8713 let cfg = QwenAttnCfg {
8714 num_heads: self.layer_num_heads(li),
8715 num_kv_heads: nkv_l,
8716 head_dim: hd_l,
8717 hidden_size: hs,
8718 position,
8719 inv_freq: &inv_freq_l,
8720 rotary_dim: rd_l,
8721 scale: self.attn_scale,
8722 softcap: self.attn_softcap,
8723 window: None,
8724 v_norm: self.attn_v_norm,
8725 q_norm: q_norm.as_deref(),
8726 k_norm: k_norm.as_deref(),
8727 output_gate: *output_gate,
8728 softplus_gate: softplus_gate
8729 .as_ref()
8730 .map(|(gate, per_head)| (gate, *per_head)),
8731 rope_scale: self.layer_rope_scale(li),
8732 bias: bias
8733 .as_ref()
8734 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
8735 rms_eps: eps,
8736 norm_style: self.norm_style,
8737 pool: pool.as_deref(),
8738 };
8739 attention::qwen_attention_nystrom(
8740 &self.ws.n1,
8741 wq,
8742 wk,
8743 wv,
8744 wo,
8745 &mut self.kv_cache.layers[li],
8746 &cfg,
8747 )
8748 }
8749 AttnKind::Full {
8750 wq,
8751 wk,
8752 wv,
8753 wo,
8754 q_norm,
8755 k_norm,
8756 output_gate,
8757 softplus_gate,
8758 bias,
8759 } => 'attn: {
8760 if graph_on
8763 && !*output_gate
8764 && softplus_gate.is_none()
8765 && self.attention_heads_per_layer.is_none()
8766 && bias.is_none()
8767 && task_mask.is_none()
8768 {
8769 let inv_freq_l = self.layer_inv_freq(li);
8770 let (nkv_l, hd_l, rd_l) = self.layer_geom(li);
8771 let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
8772 if let (Some((gm, qi)), Some((_, ki)), Some((_, vi)), Some((_, oi))) = (
8773 wq.mapped_q1(),
8774 wk.mapped_q1(),
8775 wv.mapped_q1(),
8776 wo.mapped_q1(),
8777 ) {
8778 let gm = gm.clone();
8779 let mut out = vec![0f32; hs];
8780 let cache = &self.kv_cache.layers[li];
8781 if crate::gpu::attn_dropin(
8782 &gm,
8783 self.graph_kv_id,
8784 li,
8785 &self.ws.n1,
8786 qi,
8787 ki,
8788 vi,
8789 oi,
8790 q_norm.as_deref(),
8791 k_norm.as_deref(),
8792 &inv_freq_l,
8793 nh,
8794 nkv_l,
8795 hd_l,
8796 rd_l,
8797 hs,
8798 position,
8799 self.kv_cache.max_seq_len,
8800 gemma,
8801 eps as f32,
8802 cache.k_heads(),
8803 cache.v_heads(),
8804 &mut out,
8805 ) {
8806 break 'attn out;
8807 }
8808 }
8809 }
8810 let masked = task_mask
8811 .map(|m| m.head_flags(li, self.num_heads).iter().any(|&a| !a))
8812 .unwrap_or(false);
8813 let f32_view = (wq.as_f32(), wk.as_f32(), wv.as_f32(), wo.as_f32());
8814 match (masked, f32_view) {
8815 (true, (Some(q), Some(k), Some(v), Some(o))) => {
8818 let active_heads = task_mask.unwrap().head_flags(li, self.num_heads);
8819 attention::multi_head_attention(
8820 &self.ws.n1,
8821 q,
8822 k,
8823 v,
8824 o,
8825 &mut self.kv_cache.layers[li],
8826 self.num_heads,
8827 self.num_kv_heads,
8828 self.head_dim,
8829 self.hidden_size,
8830 position,
8831 &active_heads,
8832 &self.inv_freq,
8833 )
8834 }
8835 (masked, _) => {
8836 if masked {
8837 tracing::warn!(
8838 "layer {li}: head mask on quantized weights not \
8839 supported yet — executing dense"
8840 );
8841 }
8842 let inv_freq_l = self.layer_inv_freq(li);
8843 let (nkv_l, hd_l, rd_l) = self.layer_geom(li);
8844 let cfg = QwenAttnCfg {
8845 num_heads: self.layer_num_heads(li),
8846 num_kv_heads: nkv_l,
8847 head_dim: hd_l,
8848 hidden_size: hs,
8849 position,
8850 inv_freq: &inv_freq_l,
8851 rotary_dim: rd_l,
8852 scale: self.attn_scale,
8853 softcap: self.attn_softcap,
8854 window: self.layer_window(li),
8855 v_norm: self.attn_v_norm,
8856 q_norm: q_norm.as_deref(),
8857 k_norm: k_norm.as_deref(),
8858 output_gate: *output_gate,
8859 softplus_gate: softplus_gate
8860 .as_ref()
8861 .map(|(gate, per_head)| (gate, *per_head)),
8862 rope_scale: self.layer_rope_scale(li),
8863 bias: bias
8864 .as_ref()
8865 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
8866 rms_eps: eps,
8867 norm_style: self.norm_style,
8868 pool: pool.as_deref(),
8869 };
8870 attention::qwen_attention(
8871 &self.ws.n1,
8872 wq,
8873 wk,
8874 wv,
8875 wo,
8876 &mut self.kv_cache.layers[li],
8877 &cfg,
8878 )
8879 }
8880 }
8881 }
8882 };
8883 let attn_out = match &self.weights.layers[self.phys_layer(li)].attn_out_norm {
8886 Some(w) => inference::rms_norm(&attn_out, w, self.rms_eps, self.norm_style),
8887 None => attn_out,
8888 };
8889 let lw = &self.weights.layers[self.phys_layer(li)];
8890 inference::add_rmsnorm_fused_into(
8891 &mut h,
8892 &attn_out,
8893 &lw.post_norm,
8894 self.rms_eps,
8895 self.norm_style,
8896 &mut self.ws.p1,
8897 );
8898 let mut attn_out = attn_out;
8899 attention::recycle_buf(&mut attn_out);
8900 let post_normed = &self.ws.p1;
8901
8902 let ffn_masked = task_mask
8903 .map(|m| m.ffn_active_count(li) < self.intermediate_size)
8904 .unwrap_or(false);
8905 let ffn_out = match (ffn_masked, &lw.ffn) {
8917 (_, FfnKind::Dense(d)) if !d.segs.is_empty() => {
8921 let row = task_mask
8922 .and_then(|tm| tm.ffn_masks.get(li))
8923 .map(|v| v.as_slice());
8924 tube_ffn(d, post_normed, 1, self.pool.as_deref(), row)
8925 }
8926 (true, FfnKind::Dense(d)) => {
8927 let tm = task_mask.unwrap();
8928 let alive = tm.ffn_active_count(li);
8929 let deep = alive * 2 <= self.intermediate_size;
8930 if deep && d.down_proj.sparse_col_ok() {
8931 let active = tm.ffn_active_indices(li);
8932 sparse_ffn_quant(
8933 d,
8934 post_normed,
8935 &active,
8936 self.hidden_size,
8937 self.pool.as_deref(),
8938 )
8939 } else if deep
8940 && let (Some(g), Some(u), Some(dn)) = (
8941 d.gate_proj.as_f32(),
8942 d.up_proj.as_f32(),
8943 d.down_proj.as_f32(),
8944 )
8945 {
8946 let active = tm.ffn_active_indices(li);
8947 inference::sparse_ffn_forward(
8948 post_normed,
8949 g,
8950 u,
8951 dn,
8952 self.hidden_size,
8953 self.intermediate_size,
8954 &active,
8955 self.pool.as_deref(),
8956 )
8957 } else {
8958 let row = tm.ffn_masks.get(li).map(|v| v.as_slice());
8959 dense_ffn_batch(d, post_normed, 1, self.pool.as_deref(), row)
8960 }
8961 }
8962 (true, FfnKind::Moe(m)) => {
8963 let allowed = task_mask.and_then(|tm| tm.expert_flags(li, m.experts.len()));
8967 ffn_forward(
8968 &lw.ffn,
8969 post_normed,
8970 self.pool.as_deref(),
8971 allowed.as_deref(),
8972 )
8973 }
8974 (true, FfnKind::DenseMoe(dm)) => dense_moe_ffn(
8975 dm,
8976 post_normed,
8977 &h,
8978 self.rms_eps,
8979 self.norm_style,
8980 self.pool.as_deref(),
8981 ),
8982 (false, _) => match &lw.ffn {
8983 FfnKind::DenseMoe(dm) => dense_moe_ffn(
8984 dm,
8985 post_normed,
8986 &h,
8987 self.rms_eps,
8988 self.norm_style,
8989 self.pool.as_deref(),
8990 ),
8991 _ => {
8992 let allowed = match (&lw.ffn, task_mask) {
8993 (FfnKind::Moe(m), Some(tm)) => tm.expert_flags(li, m.experts.len()),
8994 _ => None,
8995 };
8996 ffn_forward(
8997 &lw.ffn,
8998 post_normed,
8999 self.pool.as_deref(),
9000 allowed.as_deref(),
9001 )
9002 }
9003 },
9004 };
9005 let ffn_out = match &self.weights.layers[self.phys_layer(li)].ffn_out_norm {
9006 Some(w) => inference::rms_norm(&ffn_out, w, self.rms_eps, self.norm_style),
9007 None => ffn_out,
9008 };
9009 for (i, &f) in ffn_out.iter().enumerate() {
9010 h[i] += f;
9011 }
9012 let mut ffn_out = ffn_out;
9013 attention::recycle_buf(&mut ffn_out);
9014
9015 if let Some(sc) = self.weights.layers[self.phys_layer(li)].layer_scale {
9017 for v in h.iter_mut() {
9018 *v *= sc;
9019 }
9020 }
9021
9022 if self.is_loop_end(li) && li + 1 < self.num_layers {
9025 h = inference::rms_norm(
9026 &h,
9027 &self.weights.final_norm,
9028 self.rms_eps,
9029 self.norm_style,
9030 );
9031 }
9032
9033 if self.dyn_phi_layer == Some(li) {
9037 self.update_dyn_phi(&h);
9038 }
9039 }
9040 crate::gpu::set_layer(-1); if let Some(t) = t_race_cpu {
9042 crate::gpu::graph_race_record(false, t.elapsed());
9043 }
9044
9045 h
9046 }
9047
9048 fn update_dyn_phi(&mut self, h: &[f32]) {
9051 const A: f32 = 0.2;
9052 if self.dyn_phi_ema.len() != h.len() {
9053 self.dyn_phi_ema = vec![0.0; h.len()];
9054 self.dyn_phi_seen = 0;
9055 }
9056 if self.dyn_phi_seen == 0 {
9057 self.dyn_phi_ema.copy_from_slice(h);
9058 } else {
9059 for (e, &v) in self.dyn_phi_ema.iter_mut().zip(h) {
9060 *e = (1.0 - A) * *e + A * v;
9061 }
9062 }
9063 self.dyn_phi_seen += 1;
9064 }
9065
9066 pub fn dyn_phi(&self) -> &[f32] {
9068 &self.dyn_phi_ema
9069 }
9070
9071 pub fn set_dyn_phi_layer(&mut self, layer: Option<usize>) {
9073 self.dyn_phi_layer = layer;
9074 self.dyn_phi_ema.clear();
9075 self.dyn_phi_seen = 0;
9076 }
9077
9078 pub fn dynamic_skills(&self) -> Vec<(usize, String, usize)> {
9080 let Some(model) = &self.model else {
9081 return Vec::new();
9082 };
9083 model
9084 .header
9085 .skills
9086 .iter()
9087 .enumerate()
9088 .filter_map(|(i, sk)| {
9089 let ok = matches!(self.dyn_skill_layers.get(i), Some(Some(_)));
9090 let sel = sk.selection.as_ref()?;
9091 (ok).then(|| (i, sk.id.clone(), sel.phi_layer))
9092 })
9093 .collect()
9094 }
9095
9096 pub fn active_skill(&self) -> Option<usize> {
9098 self.dyn_active
9099 }
9100
9101 pub fn enable_dynamic_routing(&mut self) -> usize {
9106 use crate::swarm::{DynRouter, RoutableSkill};
9107 let Some(model) = self.model.clone() else {
9108 return 0;
9109 };
9110 if self.dyn_blend_loaded {
9113 tracing::warn!("dynamic routing unavailable on a blend-loaded pipeline");
9114 return 0;
9115 }
9116 if let Some(a) = self.dyn_active {
9120 if !matches!(self.dyn_skill_layers.get(a), Some(Some(_))) {
9121 tracing::warn!("loaded skill is not FFN-eligible — dynamic routing unavailable");
9122 return 0;
9123 }
9124 }
9125 let hidden = self.hidden_size;
9126 let mut skills = Vec::new();
9127 for (idx, id, _phi) in self.dynamic_skills() {
9128 if let Some(sel) = model.header.skills[idx].selection.as_ref() {
9129 if let Some(rs) = RoutableSkill::from_descriptor(idx, id, sel, hidden) {
9130 skills.push(rs);
9131 }
9132 }
9133 }
9134 if skills.is_empty() {
9135 return 0;
9136 }
9137 let phi = skills[0].phi_layer;
9139 if skills.iter().any(|s| s.phi_layer != phi) {
9140 tracing::warn!("routable skills disagree on phi_layer; using {phi}");
9141 }
9142 let n = skills.len();
9143 self.set_dyn_phi_layer(Some(phi));
9144 self.dyn_router = Some(DynRouter::new(skills));
9145 n
9146 }
9147
9148 pub fn route_switches(&self) -> Vec<(usize, Option<String>, Option<String>)> {
9150 self.dyn_router
9151 .as_ref()
9152 .map(|r| r.switches.clone())
9153 .unwrap_or_default()
9154 }
9155
9156 fn lm_head_forward(&self, hidden: &[f32]) -> Vec<f32> {
9159 let rows = self.weights.lm_head.rows();
9160 let mut logits = attention::take_buf(rows.min(self.vocab_size));
9161 self.weights
9162 .lm_head
9163 .matvec(hidden, &mut logits, self.pool.as_deref());
9164 logits.resize(self.vocab_size, 0.0);
9165 if let Some(m) = self.logit_multiplier {
9166 for l in logits.iter_mut() {
9167 *l *= m;
9168 }
9169 }
9170 if let Some(c) = self.final_softcap {
9171 for l in logits.iter_mut() {
9172 *l = c * (*l / c).tanh();
9173 }
9174 }
9175 if let Some(cm) = self.head_clusters.as_ref() {
9176 self.hierarchical_head_logprobs(hidden, cm, &mut logits);
9177 }
9178 logits
9179 }
9180
9181 fn hierarchical_head_logprobs(&self, hidden: &[f32], cm: &[f32], logits: &mut [f32]) {
9184 let h = hidden.len();
9185 let ncl = cm.len() / h.max(1);
9186 if ncl == 0 || logits.len() % ncl != 0 {
9187 return;
9188 }
9189 let cs = logits.len() / ncl;
9190 let mut lc = vec![0.0f32; ncl];
9192 for c in 0..ncl {
9193 let row = &cm[c * h..(c + 1) * h];
9194 let mut s = 0.0f32;
9195 for j in 0..h {
9196 s += row[j] * hidden[j];
9197 }
9198 lc[c] = s;
9199 }
9200 let mx = lc.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
9201 let lse: f32 = mx + lc.iter().map(|v| (v - mx).exp()).sum::<f32>().ln();
9202 for c in 0..ncl {
9203 let blk = &mut logits[c * cs..(c + 1) * cs];
9204 let bm = blk.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
9205 let bl: f32 = bm + blk.iter().map(|v| (v - bm).exp()).sum::<f32>().ln();
9206 let add = lc[c] - lse - bl;
9207 for v in blk.iter_mut() {
9208 *v += add;
9209 }
9210 }
9211 }
9212
9213 pub fn prefill_next_logits(&mut self, ids: &[u32], task_mask: Option<&TaskMask>) -> Vec<f32> {
9218 self.kv_cache.clear();
9219 self.kv_history.clear();
9220 let mut hidden = vec![0.0f32; self.hidden_size];
9221 for (pos, &id) in ids.iter().enumerate() {
9222 let emb = self.embed_single(id);
9223 hidden = self.forward_layers(&emb, pos, task_mask);
9224 }
9225 inference::rms_norm_into(
9226 &hidden,
9227 &self.weights.final_norm,
9228 self.rms_eps,
9229 self.norm_style,
9230 &mut self.ws.n1,
9231 );
9232 self.lm_head_forward(&self.ws.n1)
9233 }
9234}
9235
9236pub fn create_test_pipeline(
9238 hidden_size: usize,
9239 intermediate_size: usize,
9240 num_heads: usize,
9241 num_kv_heads: usize,
9242 head_dim: usize,
9243 num_layers: usize,
9244 vocab_size: usize,
9245) -> Pipeline {
9246 let synth = |n: usize, salt: usize| -> Vec<f32> {
9249 (0..n)
9250 .map(|i| (((i * 31 + salt * 17 + 7) % 97) as f32 / 97.0 - 0.5) * 0.2)
9251 .collect()
9252 };
9253 let qt = |rows: usize, cols: usize, salt: usize| -> QTensor {
9254 QTensor::from_f32(synth(rows * cols, salt), rows, cols)
9255 };
9256 let layer_weights: Vec<LayerWeights> = (0..num_layers)
9257 .map(|li| LayerWeights {
9258 input_norm: vec![1.0; hidden_size],
9259 post_norm: vec![1.0; hidden_size],
9260 attn_out_norm: None,
9261 ffn_out_norm: None,
9262 layer_scale: None,
9263 ffn: FfnKind::Dense(DenseFfn {
9264 gate_proj: qt(intermediate_size, hidden_size, li * 10 + 5),
9265 up_proj: qt(intermediate_size, hidden_size, li * 10 + 6),
9266 down_proj: qt(hidden_size, intermediate_size, li * 10 + 7),
9267 act: Act::Silu,
9268 down_t: None,
9269 segs: Vec::new(),
9270 }),
9271 attn: AttnKind::Full {
9272 bias: None,
9273 wq: qt(num_heads * head_dim, hidden_size, li * 10 + 1),
9274 wk: qt(num_kv_heads * head_dim, hidden_size, li * 10 + 2),
9275 wv: qt(num_kv_heads * head_dim, hidden_size, li * 10 + 3),
9276 wo: qt(hidden_size, num_heads * head_dim, li * 10 + 4),
9277 q_norm: None,
9278 k_norm: None,
9279 output_gate: false,
9280 softplus_gate: None,
9281 },
9282 })
9283 .collect();
9284
9285 Pipeline::new(
9286 Tokenizer::byte_level(),
9287 PipelineWeights {
9288 embed_tokens: qt(vocab_size, hidden_size, 100),
9289 layers: layer_weights,
9290 lm_head: qt(vocab_size, hidden_size, 200),
9291 final_norm: vec![1.0; hidden_size],
9292 },
9293 hidden_size,
9294 intermediate_size,
9295 num_heads,
9296 num_kv_heads,
9297 head_dim,
9298 num_layers,
9299 num_layers, false, vocab_size,
9302 1e-6,
9303 10_000.0,
9304 NormStyle::Qwen,
9305 4096,
9306 SamplerConfig {
9307 seed: Some(42),
9308 ..Default::default()
9309 },
9310 )
9311}
9312
9313#[inline]
9318fn mask_bit(row: &[u8], j: usize) -> bool {
9319 (row.get(j >> 3).copied().unwrap_or(0) >> (j & 7)) & 1 != 0
9320}
9321
9322fn mask_gain() -> f32 {
9333 static G: std::sync::OnceLock<f32> = std::sync::OnceLock::new();
9334 *G.get_or_init(|| {
9335 std::env::var("CMF_FFN_MASK_GAIN")
9336 .ok()
9337 .and_then(|v| v.parse().ok())
9338 .unwrap_or(1.0)
9339 })
9340}
9341
9342fn zero_masked_cols(g: &mut [f32], rows: usize, inter: usize, row: &[u8]) {
9343 let fill = meanfill().and_then(|(i, v)| {
9346 let li = crate::gpu::cur_layer();
9347 (*i == inter && li >= 0).then(|| &v[li as usize * inter..(li as usize + 1) * inter])
9348 });
9349 for r in 0..rows {
9350 let base = r * inter;
9351 for (bi, &byte) in row.iter().enumerate() {
9352 if byte == 0xFF {
9353 continue;
9354 }
9355 let j0 = bi * 8;
9356 for bit in 0..8 {
9357 let j = j0 + bit;
9358 if j < inter && byte & (1 << bit) == 0 {
9359 g[base + j] = fill.map_or(0.0, |f| f[j]);
9360 }
9361 }
9362 }
9363 }
9364 let gain = mask_gain();
9365 if gain != 1.0 {
9366 for v in g[..rows * inter].iter_mut() {
9367 *v *= gain;
9368 }
9369 }
9370}
9371
9372#[inline]
9374fn tube_bit(row: Option<&[u8]>, i: usize) -> bool {
9375 row.is_none_or(|r| mask_bit(r, i))
9376}
9377
9378fn all_bits_on(row: &[u8], n: usize) -> bool {
9381 (0..n).all(|i| mask_bit(row, i))
9382}
9383
9384fn tube_topk() -> usize {
9392 static K: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
9393 *K.get_or_init(|| {
9394 std::env::var("CMF_TUBE_TOPK")
9395 .ok()
9396 .and_then(|v| v.parse().ok())
9397 .unwrap_or(0)
9398 })
9399}
9400
9401fn tube_score_oracle() -> bool {
9402 static O: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
9403 *O.get_or_init(|| std::env::var("CMF_TUBE_SCORE").is_ok_and(|v| v == "oracle"))
9404}
9405
9406fn tube_ffn_routed(
9413 d: &DenseFfn,
9414 xs: &[f32],
9415 b: usize,
9416 pool: Option<&Pool>,
9417 mask_row: Option<&[u8]>,
9418 k: usize,
9419) -> Vec<f32> {
9420 let hidden = d.down_proj.rows();
9421 let core = d.gate_proj.rows();
9422 let core_full = mask_row.is_none_or(|r| all_bits_on(r, core));
9423 let mut out = match (b, core_full, mask_row) {
9424 (1, true, _) => dense_ffn(d, xs, pool),
9425 (1, false, Some(row)) => dense_ffn_masked(d, xs, pool, row),
9426 (_, true, _) => dense_ffn_batch(d, xs, b, pool, None),
9427 (_, false, row) => dense_ffn_batch(d, xs, b, pool, row),
9428 };
9429 let cand: Vec<usize> = (0..d.segs.len())
9430 .filter(|&i| tube_bit(mask_row, d.segs[i].start))
9431 .collect();
9432 if cand.is_empty() {
9433 return out;
9434 }
9435 let oracle = tube_score_oracle();
9439 let mut acts: Vec<Vec<f32>> = Vec::with_capacity(cand.len());
9440 let mut scores = vec![0f32; b * cand.len()];
9441 for (ci, &i) in cand.iter().enumerate() {
9442 let seg = &d.segs[i];
9443 let w = seg.width;
9444 let mut g = vec![0.0f32; b * w];
9445 if b == 1 {
9446 seg.gate.matvec(xs, &mut g, pool);
9447 } else {
9448 seg.gate.matmat(xs, b, &mut g, pool);
9449 }
9450 for v in g.iter_mut() {
9451 *v = Act::Silu.combine(*v, 1.0);
9452 }
9453 if !oracle {
9454 for t in 0..b {
9455 scores[t * cand.len() + ci] =
9456 g[t * w..(t + 1) * w].iter().map(|v| v * v).sum::<f32>();
9457 }
9458 }
9459 if oracle || b > 1 {
9460 let mut u = vec![0.0f32; b * w];
9461 if b == 1 {
9462 seg.up.matvec(xs, &mut u, pool);
9463 } else {
9464 seg.up.matmat(xs, b, &mut u, pool);
9465 }
9466 for (a, &v) in g.iter_mut().zip(u.iter()) {
9467 *a *= v;
9468 }
9469 if oracle {
9470 for t in 0..b {
9471 scores[t * cand.len() + ci] =
9472 g[t * w..(t + 1) * w].iter().map(|v| v * v).sum::<f32>();
9473 }
9474 }
9475 }
9476 acts.push(g);
9477 }
9478 let keep = k.min(cand.len());
9480 let mut scratch: Vec<f32> = Vec::new();
9481 for t in 0..b {
9482 let mut sc: Vec<(f32, usize)> = (0..cand.len())
9483 .map(|ci| (scores[t * cand.len() + ci], ci))
9484 .collect();
9485 sc.sort_unstable_by(|x, y| y.0.total_cmp(&x.0));
9486 let mut alive = vec![false; cand.len()];
9487 for &(_, ci) in sc.iter().take(keep) {
9488 alive[ci] = true;
9489 }
9490 if b > 1 {
9491 for (ci, a) in acts.iter_mut().enumerate() {
9492 if !alive[ci] {
9493 let w = d.segs[cand[ci]].width;
9494 a[t * w..(t + 1) * w].fill(0.0);
9495 }
9496 }
9497 } else {
9498 for (ci, &i) in cand.iter().enumerate() {
9502 if !alive[ci] {
9503 continue;
9504 }
9505 let seg = &d.segs[i];
9506 let w = seg.width;
9507 let g = &mut acts[ci];
9508 if !tube_score_oracle() {
9509 scratch.clear();
9510 scratch.resize(w, 0.0);
9511 seg.up.matvec(xs, &mut scratch, pool);
9512 for (a, &v) in g.iter_mut().zip(scratch.iter()) {
9513 *a *= v;
9514 }
9515 }
9516 let mut acc = vec![0.0f32; hidden];
9517 seg.down.matvec(g, &mut acc, pool);
9518 for (o, a) in out.iter_mut().zip(&acc) {
9519 *o += *a;
9520 }
9521 }
9522 }
9523 }
9524 if b > 1 {
9525 for (ci, &i) in cand.iter().enumerate() {
9526 let seg = &d.segs[i];
9527 let mut acc = vec![0.0f32; b * hidden];
9528 seg.down.matmat(&acts[ci], b, &mut acc, pool);
9529 for (o, a) in out.iter_mut().zip(&acc) {
9530 *o += *a;
9531 }
9532 }
9533 }
9534 out
9535}
9536
9537fn tube_ffn(
9543 d: &DenseFfn,
9544 xs: &[f32],
9545 b: usize,
9546 pool: Option<&Pool>,
9547 mask_row: Option<&[u8]>,
9548) -> Vec<f32> {
9549 if tube_topk() > 0 {
9550 return tube_ffn_routed(d, xs, b, pool, mask_row, tube_topk());
9551 }
9552 let hidden = d.down_proj.rows();
9553 let core = d.gate_proj.rows();
9554 let core_full = mask_row.is_none_or(|r| all_bits_on(r, core));
9555 let mut out = match (b, core_full, mask_row) {
9556 (1, true, _) => dense_ffn(d, xs, pool),
9557 (1, false, Some(row)) => dense_ffn_masked(d, xs, pool, row),
9558 (_, true, _) => dense_ffn_batch(d, xs, b, pool, None),
9559 (_, false, row) => dense_ffn_batch(d, xs, b, pool, row),
9560 };
9561 TUBE_SCRATCH.with(|sc| {
9562 let mut sc = sc.borrow_mut();
9563 let [g, u, acc] = &mut *sc;
9564 for seg in &d.segs {
9565 if !tube_bit(mask_row, seg.start) {
9566 continue;
9567 }
9568 let w = seg.width;
9569 g.resize(b * w, 0.0);
9570 if b == 1
9571 && d.act == Act::Silu
9572 && QTensor::matvec_silu_mul(&seg.gate, &seg.up, xs, g, pool)
9573 {
9574 } else {
9576 u.resize(b * w, 0.0);
9577 if b == 1 {
9578 QTensor::matvec_many([&seg.gate, &seg.up], xs, [g, u], pool);
9579 } else {
9580 seg.gate.matmat(xs, b, g, pool);
9581 seg.up.matmat(xs, b, u, pool);
9582 }
9583 for i in 0..b * w {
9584 g[i] = d.act.combine(g[i], u[i]);
9585 }
9586 }
9587 acc.resize(b * hidden, 0.0);
9588 acc.fill(0.0);
9589 if b == 1 {
9590 seg.down.matvec(g, acc, pool);
9591 } else {
9592 seg.down.matmat(g, b, acc, pool);
9593 }
9594 for (o, a) in out.iter_mut().zip(acc.iter()) {
9595 *o += *a;
9596 }
9597 }
9598 out
9599 })
9600}
9601
9602thread_local! {
9603 static TUBE_SCRATCH: std::cell::RefCell<[Vec<f32>; 3]> =
9607 const { std::cell::RefCell::new([Vec::new(), Vec::new(), Vec::new()]) };
9608}
9609
9610fn dense_ffn_batch(
9611 d: &DenseFfn,
9612 xs: &[f32],
9613 b: usize,
9614 pool: Option<&Pool>,
9615 mask_row: Option<&[u8]>,
9616) -> Vec<f32> {
9617 let inter = d.gate_proj.rows();
9618 let hidden = d.down_proj.rows();
9619 if mask_row.is_none()
9627 && d.act == Act::Silu
9628 && b >= 32
9629 && crate::gpu::enabled_here()
9630 && !crate::gpu::mm_killed()
9631 && refit_dir().is_none()
9636 && !ffn_probe_active()
9641 {
9642 if let (Some((model, w1)), Some((_, w3)), Some((_, w2))) = (
9643 d.gate_proj.mapped_q4t(),
9644 d.up_proj.mapped_q4t(),
9645 d.down_proj.mapped_q4t(),
9646 ) {
9647 let mut out = vec![0.0f32; b * hidden];
9648 if crate::gpu::q4t_ffn(model, w1, w3, w2, xs, b, hidden, inter, &mut out) {
9649 return out;
9650 }
9651 }
9652 if let (Some((model, w1)), Some((_, w3)), Some((_, w2))) = (
9657 d.gate_proj.mapped_q4tp(),
9658 d.up_proj.mapped_q4tp(),
9659 d.down_proj.mapped_q4tp(),
9660 ) {
9661 let mut out = vec![0.0f32; b * hidden];
9662 if crate::gpu::q4tp_ffn(model, w1, w3, w2, xs, b, hidden, inter, &mut out) {
9663 return out;
9664 }
9665 }
9666 }
9667 let mut g = vec![0.0f32; b * inter];
9668 d.gate_proj.matmat(xs, b, &mut g, pool);
9669 let mut u = vec![0.0f32; b * inter];
9670 d.up_proj.matmat(xs, b, &mut u, pool);
9671 if gate_topk() > 0 && d.act == Act::Silu {
9672 for t in 0..b {
9673 let row = &mut g[t * inter..(t + 1) * inter];
9674 for v in row.iter_mut() {
9675 *v = Act::Silu.combine(*v, 1.0);
9676 }
9677 keep_top_k(row, gate_topk());
9678 }
9679 for i in 0..b * inter {
9680 g[i] *= u[i];
9681 }
9682 } else {
9683 for i in 0..b * inter {
9684 g[i] = d.act.combine(g[i], u[i]);
9685 }
9686 }
9687 if let Some(row) = mask_row {
9688 zero_masked_cols(&mut g, b, inter, row);
9689 }
9690 if oracle_topk() > 0 {
9691 for t in 0..b {
9692 keep_top_k(&mut g[t * inter..(t + 1) * inter], oracle_topk());
9693 }
9694 }
9695 let mut out = vec![0.0f32; b * hidden];
9696 d.down_proj.matmat(&g, b, &mut out, pool);
9697 if refit_dir().is_some() {
9698 let li = crate::gpu::cur_layer();
9699 if li >= 0 {
9700 refit_accumulate(li as usize, &g, b, inter, &out, hidden, pool);
9701 }
9702 }
9703 FFN_PROBE.with(|pr| {
9707 if let Some(acc) = pr.borrow_mut().as_mut() {
9708 let li = crate::gpu::cur_layer();
9709 if li < 0 {
9710 return;
9711 }
9712 let Some(row) = acc.get_mut(li as usize) else {
9713 return;
9714 };
9715 let sq = probe_sq();
9716 for t in 0..b {
9717 for (a, &v) in row.iter_mut().zip(&g[t * inter..(t + 1) * inter]) {
9718 *a += if sq {
9719 (v as f64) * (v as f64)
9720 } else {
9721 (v as f64).abs()
9722 };
9723 }
9724 }
9725 }
9726 });
9727 out
9728}
9729
9730fn accumulate_act(m: &MoeFfn, xs: &[f32], b: usize) {
9735 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
9736 static DUMP: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
9737 let on = *ON.get_or_init(|| std::env::var("CMF_RMS_TRACE").is_ok());
9738 let dump = *DUMP.get_or_init(|| std::env::var("CMF_ACT_DUMP").is_ok());
9739 if (!on && !dump) || b == 0 {
9740 return;
9741 }
9742 let hidden = xs.len() / b;
9743 if on {
9744 let mut acc = m.act_sq.borrow_mut();
9745 if acc.len() < hidden {
9746 acc.resize(hidden, 0.0);
9747 }
9748 for t in 0..b {
9749 let row = &xs[t * hidden..(t + 1) * hidden];
9750 for (a, &v) in acc.iter_mut().zip(row) {
9751 *a += (v as f64) * (v as f64);
9752 }
9753 }
9754 }
9755 if dump {
9756 let cap: usize = std::env::var("CMF_ACT_DUMP_ROWS")
9759 .ok()
9760 .and_then(|v| v.parse().ok())
9761 .unwrap_or(4096);
9762 let mut rows = m.act_rows.borrow_mut();
9763 if rows.len() < cap * hidden {
9764 let take = b.min((cap * hidden - rows.len()) / hidden.max(1));
9765 rows.extend_from_slice(&xs[..take * hidden]);
9766 }
9767 }
9768}
9769
9770#[derive(Clone, Copy)]
9773struct SendVecs(*mut Vec<f32>);
9774unsafe impl Send for SendVecs {}
9775unsafe impl Sync for SendVecs {}
9776impl SendVecs {
9777 #[inline]
9778 fn at(self, i: usize) -> *mut Vec<f32> {
9779 unsafe { self.0.add(i) }
9780 }
9781}
9782
9783fn moe_ffn_batch(
9784 m: &MoeFfn,
9785 xs: &[f32],
9786 b: usize,
9787 hidden: usize,
9788 pool: Option<&Pool>,
9789 allowed: Option<&[bool]>,
9790) -> Vec<f32> {
9791 accumulate_act(m, xs, b);
9792 let ne = m.experts.len();
9793 let mut logits = vec![0.0f32; b * ne];
9794 match &m.resonance {
9795 Some(r) => {
9796 let hdim = xs.len() / b.max(1);
9797 for bi in 0..b {
9798 r.scores(
9799 &xs[bi * hdim..(bi + 1) * hdim],
9800 &mut logits[bi * ne..(bi + 1) * ne],
9801 );
9802 }
9803 }
9804 None => m.router.matmat(xs, b, &mut logits, pool),
9805 }
9806
9807 let mut assign: Vec<Vec<(usize, f32)>> = vec![Vec::new(); ne];
9810 {
9811 let mut st = m.stats.borrow_mut();
9812 if st.len() < ne {
9813 st.resize(ne, 0);
9814 }
9815 for bi in 0..b {
9816 let (idx, p, wsum) = moe_route(&logits[bi * ne..(bi + 1) * ne], m, allowed);
9817 for &e in &idx {
9818 st[e] += 1;
9819 assign[e].push((bi, p[e] / wsum));
9820 }
9821 }
9822 }
9823
9824 let mut out = vec![0.0f32; b * hidden];
9825 let cols = m.experts[0].gate_proj.cols();
9826 let run_expert = |d: &DenseFfn, list: &[(usize, f32)], out: &mut [f32]| {
9827 let sb = list.len();
9828 let mut sub = vec![0.0f32; sb * cols];
9829 for (k, &(bi, _)) in list.iter().enumerate() {
9830 sub[k * cols..(k + 1) * cols].copy_from_slice(&xs[bi * cols..(bi + 1) * cols]);
9831 }
9832 let eo = dense_ffn_batch(d, &sub, sb, pool, None);
9833 for (k, &(bi, w)) in list.iter().enumerate() {
9834 for i in 0..hidden {
9835 out[bi * hidden + i] += w * eo[k * hidden + i];
9836 }
9837 }
9838 };
9839 let active: Vec<usize> = (0..ne).filter(|&e| !assign[e].is_empty()).collect();
9845 if pool.is_some() && active.len() >= 8 {
9846 let mut panels: Vec<Vec<f32>> = vec![Vec::new(); active.len()];
9847 {
9848 let panel_ptr = SendVecs(panels.as_mut_ptr());
9849 let experts = &m.experts;
9852 let (active_r, assign_r) = (&active, &assign);
9853 let run = |start: usize, end: usize| {
9854 for ai in start..end {
9855 let e = active_r[ai];
9856 let list = &assign_r[e];
9857 let sb = list.len();
9858 let mut sub = vec![0.0f32; sb * cols];
9859 for (k, &(bi, _)) in list.iter().enumerate() {
9860 sub[k * cols..(k + 1) * cols]
9861 .copy_from_slice(&xs[bi * cols..(bi + 1) * cols]);
9862 }
9863 unsafe {
9865 *panel_ptr.at(ai) = dense_ffn_batch(&experts[e], &sub, sb, None, None);
9866 }
9867 }
9868 };
9869 match pool {
9870 Some(p) => p.run_rows(active.len(), &run),
9871 None => run(0, active.len()),
9872 }
9873 }
9874 for (ai, &e) in active.iter().enumerate() {
9875 for (k, &(bi, w)) in assign[e].iter().enumerate() {
9876 let eo = &panels[ai][k * hidden..(k + 1) * hidden];
9877 for i in 0..hidden {
9878 out[bi * hidden + i] += w * eo[i];
9879 }
9880 }
9881 }
9882 } else {
9883 for &e in &active {
9884 run_expert(&m.experts[e], &assign[e], &mut out);
9885 }
9886 }
9887 if let Some((se, gate)) = &m.shared {
9888 let all: Vec<(usize, f32)> = if let Some(gate) = gate {
9889 let mut gl = vec![0.0f32; b];
9890 gate.matmat(xs, b, &mut gl, pool);
9891 (0..b)
9892 .map(|bi| (bi, 1.0 / (1.0 + (-gl[bi]).exp())))
9893 .collect()
9894 } else {
9895 (0..b).map(|bi| (bi, 1.0)).collect()
9896 };
9897 run_expert(se, &all, &mut out);
9898 }
9899 out
9900}
9901
9902thread_local! {
9903 static FFN_SCRATCH: std::cell::RefCell<[Vec<f32>; 4]> =
9907 const { std::cell::RefCell::new([Vec::new(), Vec::new(), Vec::new(), Vec::new()]) };
9908}
9909
9910fn dense_ffn(d: &DenseFfn, x: &[f32], pool: Option<&Pool>) -> Vec<f32> {
9912 if gate_topk() > 0
9915 && let Some(out) = dense_ffn_dynamic(d, x, pool, gate_topk())
9916 {
9917 return out;
9918 }
9919 if crate::gpu::enabled_here()
9930 && (d.gate_proj.rows() >= crate::gpu::min_rows() || d.gate_proj.is_q1())
9931 {
9932 let arm = if d.gate_proj.is_q1() && crate::gpu::q1_force() {
9933 crate::gpu::ProbeArm::Gpu
9934 } else {
9935 crate::gpu::probe_arm(crate::gpu::OpClass::Ffn)
9936 };
9937 match arm {
9938 crate::gpu::ProbeArm::Gpu => {
9939 let t0 = std::time::Instant::now();
9940 if let Some(out) = dense_ffn_gpu(d, x, pool) {
9941 crate::gpu::probe_record(crate::gpu::OpClass::Ffn, true, t0.elapsed());
9942 return out;
9943 }
9944 crate::gpu::probe_note_decline(crate::gpu::OpClass::Ffn);
9948 }
9949 crate::gpu::ProbeArm::CpuTimed => {
9950 let t0 = std::time::Instant::now();
9951 let out = crate::gpu::cpu_scope(|| dense_ffn_cpu(d, x, pool));
9952 crate::gpu::probe_record(crate::gpu::OpClass::Ffn, false, t0.elapsed());
9953 return out;
9954 }
9955 crate::gpu::ProbeArm::Cpu => {
9956 return crate::gpu::cpu_scope(|| dense_ffn_cpu(d, x, pool));
9957 }
9958 }
9959 }
9960 dense_ffn_cpu(d, x, pool)
9961}
9962
9963fn dense_ffn_cpu(d: &DenseFfn, x: &[f32], pool: Option<&Pool>) -> Vec<f32> {
9965 let inter = d.gate_proj.rows();
9966 FFN_SCRATCH.with(|s| {
9967 let mut s = s.borrow_mut();
9968 let [g, u, ..] = &mut *s;
9969 g.resize(inter, 0.0);
9970 if gate_topk() > 0 {
9973 u.resize(inter, 0.0);
9977 QTensor::matvec_many([&d.gate_proj, &d.up_proj], x, [g, u], pool);
9978 for i in 0..inter {
9979 g[i] = Act::Silu.combine(g[i], 1.0);
9980 }
9981 keep_top_k(g, gate_topk());
9982 for i in 0..inter {
9983 g[i] *= u[i];
9984 }
9985 } else if d.act == Act::Silu
9986 && QTensor::matvec_silu_mul(&d.gate_proj, &d.up_proj, x, g, pool)
9987 {
9988 } else {
9990 u.resize(inter, 0.0);
9991 QTensor::matvec_many([&d.gate_proj, &d.up_proj], x, [g, u], pool);
9993 for i in 0..inter {
9994 g[i] = d.act.combine(g[i], u[i]);
9995 }
9996 }
9997 FFN_PROBE.with(|pr| {
10005 if let Some(acc) = pr.borrow_mut().as_mut() {
10006 let li = crate::gpu::cur_layer();
10007 if li >= 0 {
10008 if let Some(row) = acc.get_mut(li as usize) {
10009 match probe_topk() {
10010 0 if probe_sq() => {
10011 for (a, &v) in row.iter_mut().zip(g.iter()) {
10012 *a += (v as f64) * (v as f64);
10013 }
10014 }
10015 0 if probe_signed() => {
10016 for (a, &v) in row.iter_mut().zip(g.iter()) {
10017 *a += v as f64;
10018 }
10019 }
10020 0 => {
10021 for (a, &v) in row.iter_mut().zip(g.iter()) {
10022 *a += (v as f64).abs();
10023 }
10024 }
10025 k => {
10026 let n = g.len();
10027 let k = k.min(n);
10028 let mut mag: Vec<f32> = g.iter().map(|v| v.abs()).collect();
10029 let (_, kth, _) = mag.select_nth_unstable_by(k - 1, |a, b| {
10030 b.partial_cmp(a).unwrap_or(std::cmp::Ordering::Equal)
10031 });
10032 let thr = *kth;
10033 for (a, &v) in row.iter_mut().zip(g.iter()) {
10034 if v.abs() >= thr {
10035 *a += 1.0;
10036 }
10037 }
10038 }
10039 }
10040 }
10041 }
10042 }
10043 });
10044 if oracle_topk() > 0 {
10045 keep_top_k(g, oracle_topk());
10046 }
10047 {
10048 let li = crate::gpu::cur_layer();
10049 if li >= 0 {
10050 adump_row(li as usize, g);
10051 }
10052 }
10053 let mut out = attention::take_buf(d.down_proj.rows());
10054 d.down_proj.matvec(g, &mut out, pool);
10055 out
10056 })
10057}
10058
10059pub struct RefitAcc {
10072 pub support: Vec<u32>,
10073 pub gss: Vec<f32>,
10074 pub ya: Vec<f32>,
10075 pub hidden: usize,
10076 pub tokens: u64,
10077 pub buf_g: Vec<f32>,
10083 pub buf_o: Vec<f32>,
10084 pub buf_t: usize,
10085}
10086
10087type RefitState = (std::collections::HashMap<usize, RefitAcc>, Vec<f32>);
10091
10092static REFIT: std::sync::OnceLock<Option<(String, std::sync::Mutex<RefitState>)>> =
10093 std::sync::OnceLock::new();
10094
10095fn ffn_probe_active() -> bool {
10098 FFN_PROBE.with(|p| p.borrow().is_some())
10099}
10100
10101fn refit_dir() -> Option<&'static (String, std::sync::Mutex<RefitState>)> {
10102 REFIT
10103 .get_or_init(|| {
10104 std::env::var("CMF_FFN_REFIT").ok().map(|d| {
10105 (
10106 d,
10107 std::sync::Mutex::new((std::collections::HashMap::new(), Vec::new())),
10108 )
10109 })
10110 })
10111 .as_ref()
10112}
10113
10114fn refit_accumulate(
10116 li: usize,
10117 g: &[f32],
10118 b: usize,
10119 inter: usize,
10120 out: &[f32],
10121 hidden: usize,
10122 pool: Option<&Pool>,
10123) {
10124 let Some((dir, map)) = refit_dir() else {
10125 return;
10126 };
10127 static SPAN: std::sync::OnceLock<(usize, usize)> = std::sync::OnceLock::new();
10128 let (from, to) = *SPAN.get_or_init(|| {
10129 let g = |k: &str, d: usize| {
10130 std::env::var(k)
10131 .ok()
10132 .and_then(|v| v.parse().ok())
10133 .unwrap_or(d)
10134 };
10135 (
10136 g("CMF_FFN_REFIT_FROM", 0),
10137 g("CMF_FFN_REFIT_TO", usize::MAX),
10138 )
10139 });
10140 if li < from || li > to {
10141 return;
10142 }
10143 let mut guard = map.lock().unwrap();
10144 let (map, shared) = &mut *guard;
10145 let acc = match map.entry(li) {
10146 std::collections::hash_map::Entry::Occupied(e) => e.into_mut(),
10147 std::collections::hash_map::Entry::Vacant(e) => {
10148 let path = format!("{dir}/support.{li}.u32");
10149 let Ok(bytes) = std::fs::read(&path) else {
10150 eprintln!("refit: no {path} — layer {li} skipped");
10151 return;
10152 };
10153 let n = u32::from_le_bytes(bytes[0..4].try_into().unwrap()) as usize;
10154 let support: Vec<u32> = bytes[4..4 + n * 4]
10155 .chunks_exact(4)
10156 .map(|c| u32::from_le_bytes([c[0], c[1], c[2], c[3]]))
10157 .collect();
10158 eprintln!(
10159 "refit: layer {li} support {n} ({:.0} MB of accumulator)",
10160 (n * n + hidden * n) as f64 * 4.0 / 1e6
10161 );
10162 e.insert(RefitAcc {
10163 gss: vec![0.0; n * n],
10164 ya: vec![0.0; hidden * n],
10165 buf_g: Vec::new(),
10166 buf_o: Vec::new(),
10167 buf_t: 0,
10168 support,
10169 hidden,
10170 tokens: 0,
10171 })
10172 }
10173 };
10174 let ns = acc.support.len();
10175 let cap = refit_batch();
10177 if acc.buf_g.is_empty() {
10178 acc.buf_g = vec![0.0; ns * cap];
10179 acc.buf_o = vec![0.0; hidden * cap];
10180 }
10181 let take = b.min(cap - acc.buf_t);
10182 for t in 0..take {
10183 let col = acc.buf_t + t;
10184 for (j, &n) in acc.support.iter().enumerate() {
10185 acc.buf_g[j * cap + col] = g[t * inter + n as usize];
10186 }
10187 for h in 0..hidden {
10188 acc.buf_o[h * cap + col] = out[t * hidden + h];
10189 }
10190 }
10191 acc.buf_t += take;
10192 acc.tokens += take as u64;
10193 if acc.buf_t < cap {
10194 return;
10195 }
10196 let bt = acc.buf_t;
10197 acc.buf_t = 0;
10198 let RefitAcc {
10208 gss,
10209 ya,
10210 buf_g,
10211 buf_o,
10212 ..
10213 } = acc;
10214 let need = (ns * ns).max(hidden * ns);
10215 if shared.len() < need {
10216 shared.resize(need, 0.0);
10217 }
10218 let scratch = &mut shared[..];
10219 let _ = bt;
10220 if crate::gpu::gemm_nt_f32_transient(buf_g, buf_g, &mut scratch[..ns * ns], ns, cap, ns) {
10221 add_into(gss, &scratch[..ns * ns], pool);
10222 if crate::gpu::gemm_nt_f32_transient(
10223 buf_o,
10224 buf_g,
10225 &mut scratch[..hidden * ns],
10226 hidden,
10227 cap,
10228 ns,
10229 ) {
10230 add_into(ya, &scratch[..hidden * ns], pool);
10231 } else {
10232 accum_outer_t(ya, hidden, ns, cap, buf_o, buf_g, pool);
10233 }
10234 } else {
10235 accum_outer_t(gss, ns, ns, cap, buf_g, buf_g, pool);
10236 accum_outer_t(ya, hidden, ns, cap, buf_o, buf_g, pool);
10237 }
10238 }
10242
10243fn refit_batch() -> usize {
10245 static B: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
10246 *B.get_or_init(|| {
10247 std::env::var("CMF_FFN_REFIT_BATCH")
10248 .ok()
10249 .and_then(|v| v.parse().ok())
10250 .unwrap_or(4096)
10251 })
10252}
10253
10254fn accum_outer_t(
10257 c: &mut [f32],
10258 m: usize,
10259 n: usize,
10260 b: usize,
10261 left: &[f32],
10262 right: &[f32],
10263 pool: Option<&Pool>,
10264) {
10265 let ptr = SendMut(c.as_mut_ptr());
10266 let body = |i: usize| {
10267 let ptr = &ptr;
10268 let row = unsafe { std::slice::from_raw_parts_mut(ptr.0.add(i * n), n) };
10269 for t in 0..b {
10270 let a = left[i * b + t];
10271 if a == 0.0 {
10272 continue;
10273 }
10274 for (j, o) in row.iter_mut().enumerate() {
10275 *o += a * right[j * b + t];
10276 }
10277 }
10278 };
10279 match pool {
10280 Some(p) if m > 1 => p.run_rows(m, &|s, e| {
10281 for i in s..e {
10282 body(i);
10283 }
10284 }),
10285 _ => {
10286 for i in 0..m {
10287 body(i);
10288 }
10289 }
10290 }
10291}
10292
10293fn add_into(dst: &mut [f32], src: &[f32], pool: Option<&Pool>) {
10296 let n = dst.len().min(src.len());
10297 match pool {
10298 Some(p) if n >= 1 << 16 => {
10299 let ptr = SendMut(dst.as_mut_ptr());
10300 let f = |s: usize, e: usize| {
10301 let ptr = &ptr;
10302 for blk in s..e {
10303 let (a, b) = (blk * 4096, ((blk + 1) * 4096).min(n));
10304 for i in a..b {
10305 unsafe { *ptr.0.add(i) += src[i] };
10306 }
10307 }
10308 };
10309 p.run_rows(n.div_ceil(4096), &f);
10310 }
10311 _ => {
10312 for (d, v) in dst.iter_mut().zip(&src[..n]) {
10313 *d += *v;
10314 }
10315 }
10316 }
10317}
10318
10319fn accum_outer(
10324 c: &mut [f32],
10325 m: usize,
10326 n: usize,
10327 b: usize,
10328 left: &[f32],
10329 right: &[f32],
10330 pool: Option<&Pool>,
10331) {
10332 const TILE: usize = 32;
10333 let tiles = m.div_ceil(TILE);
10334 let cp = SendMut(c.as_mut_ptr());
10335 let body = |ti: usize| {
10336 let cp = &cp;
10337 let i0 = ti * TILE;
10338 let i1 = (i0 + TILE).min(m);
10339 for t in 0..b {
10340 let r = &right[t * n..t * n + n];
10341 for i in i0..i1 {
10342 let a = left[i * b + t];
10343 if a == 0.0 {
10344 continue;
10345 }
10346 let row = unsafe { std::slice::from_raw_parts_mut(cp.0.add(i * n), n) };
10348 for (o, v) in row.iter_mut().zip(r) {
10349 *o += a * *v;
10350 }
10351 }
10352 }
10353 };
10354 match pool {
10355 Some(p) if tiles > 1 => p.run_rows(tiles, &|s, e| {
10356 for ti in s..e {
10357 body(ti);
10358 }
10359 }),
10360 _ => {
10361 for ti in 0..tiles {
10362 body(ti);
10363 }
10364 }
10365 }
10366}
10367
10368pub fn refit_flush() -> usize {
10370 let Some((dir, map)) = refit_dir() else {
10371 return 0;
10372 };
10373 let guard = map.lock().unwrap();
10374 let mut n = 0;
10375 for (li, acc) in guard.0.iter() {
10376 let w = |name: &str, v: &[f32]| {
10379 let path = format!("{dir}/{name}.{li}.f32");
10380 let bytes: Vec<u8> = v.iter().flat_map(|x| x.to_le_bytes()).collect();
10381 match std::fs::write(&path, &bytes) {
10382 Ok(()) => {}
10383 Err(e) => eprintln!(
10384 "refit: FAILED to write {path} ({} MB): {e}",
10385 bytes.len() / 1_000_000
10386 ),
10387 }
10388 };
10389 w("gss", &acc.gss);
10390 w("ya", &acc.ya);
10391 println!(
10392 "refit L{li}: {} support, {} tokens, hidden {}",
10393 acc.support.len(),
10394 acc.tokens,
10395 acc.hidden
10396 );
10397 n += 1;
10398 }
10399 n
10400}
10401
10402fn adump_row(li: usize, g: &[f32]) {
10407 use std::io::Write as _;
10408 static FILES: std::sync::OnceLock<
10409 Option<(
10410 String,
10411 std::sync::Mutex<std::collections::HashMap<usize, std::fs::File>>,
10412 )>,
10413 > = std::sync::OnceLock::new();
10414 let Some((prefix, map)) = FILES
10415 .get_or_init(|| {
10416 std::env::var("CMF_FFN_ADUMP")
10417 .ok()
10418 .map(|p| (p, std::sync::Mutex::new(std::collections::HashMap::new())))
10419 })
10420 .as_ref()
10421 else {
10422 return;
10423 };
10424 static SPAN: std::sync::OnceLock<(usize, usize)> = std::sync::OnceLock::new();
10427 let (from, to) = *SPAN.get_or_init(|| {
10428 let g = |k: &str, d: usize| {
10429 std::env::var(k)
10430 .ok()
10431 .and_then(|v| v.parse().ok())
10432 .unwrap_or(d)
10433 };
10434 (
10435 g("CMF_FFN_ADUMP_FROM", 0),
10436 g("CMF_FFN_ADUMP_TO", usize::MAX),
10437 )
10438 });
10439 if li < from || li > to {
10440 return;
10441 }
10442 let mut map = map.lock().unwrap();
10443 let f = map.entry(li).or_insert_with(|| {
10444 std::fs::File::create(format!("{prefix}.{li}.f16")).expect("adump file")
10445 });
10446 let mut bytes = Vec::with_capacity(g.len() * 2);
10447 for v in g {
10448 bytes.extend_from_slice(&cortiq_core::quant::f32_to_f16(*v).to_le_bytes());
10449 }
10450 let _ = f.write_all(&bytes);
10451}
10452
10453fn oracle_topk() -> usize {
10459 static K: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
10460 *K.get_or_init(|| {
10461 std::env::var("CMF_FFN_ORACLE_TOPK")
10462 .ok()
10463 .and_then(|v| v.parse().ok())
10464 .unwrap_or(0)
10465 })
10466}
10467
10468fn gate_topk() -> usize {
10474 static K: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
10475 *K.get_or_init(|| {
10476 std::env::var("CMF_FFN_GATE_TOPK")
10477 .ok()
10478 .and_then(|v| v.parse().ok())
10479 .unwrap_or(0)
10480 })
10481}
10482
10483fn gate_block() -> usize {
10490 static B: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
10491 *B.get_or_init(|| {
10492 std::env::var("CMF_FFN_GATE_BLOCK")
10493 .ok()
10494 .and_then(|v| v.parse().ok())
10495 .unwrap_or(1)
10496 })
10497}
10498
10499fn keep_top_blocks(g: &mut [f32], keep_n: usize, block: usize) {
10501 let n = g.len();
10502 let nb = n.div_ceil(block);
10503 let kb = (keep_n.div_ceil(block)).clamp(1, nb);
10504 if kb >= nb {
10505 return;
10506 }
10507 let mut score: Vec<f32> = (0..nb)
10508 .map(|b| {
10509 g[b * block..((b + 1) * block).min(n)]
10510 .iter()
10511 .map(|v| v * v)
10512 .sum::<f32>()
10513 })
10514 .collect();
10515 let mut ord = score.clone();
10516 let (_, kth, _) = ord.select_nth_unstable_by(kb - 1, |a, b| {
10517 b.partial_cmp(a).unwrap_or(std::cmp::Ordering::Equal)
10518 });
10519 let thr = *kth;
10520 for b in 0..nb {
10521 if score[b] < thr {
10522 g[b * block..((b + 1) * block).min(n)].fill(0.0);
10523 }
10524 }
10525 score.clear();
10526}
10527
10528fn keep_top_k(g: &mut [f32], k: usize) {
10530 if gate_block() > 1 {
10531 return keep_top_blocks(g, k, gate_block());
10532 }
10533 let n = g.len();
10534 if k == 0 || k >= n {
10535 return;
10536 }
10537 let mut mag: Vec<f32> = g.iter().map(|v| v.abs()).collect();
10538 let (_, kth, _) = mag.select_nth_unstable_by(k - 1, |a, b| {
10539 b.partial_cmp(a).unwrap_or(std::cmp::Ordering::Equal)
10540 });
10541 let thr = *kth;
10542 for v in g.iter_mut() {
10543 if v.abs() < thr {
10544 *v = 0.0;
10545 }
10546 }
10547}
10548
10549fn probe_sq() -> bool {
10553 static S: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
10554 *S.get_or_init(|| std::env::var("CMF_FFN_PROBE_SQ").is_ok())
10555}
10556
10557fn probe_signed() -> bool {
10561 static S: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
10562 *S.get_or_init(|| std::env::var("CMF_FFN_PROBE_SIGNED").is_ok())
10563}
10564
10565fn meanfill() -> Option<&'static (usize, Vec<f32>)> {
10573 static M: std::sync::OnceLock<Option<(usize, Vec<f32>)>> = std::sync::OnceLock::new();
10574 M.get_or_init(|| {
10575 let p = std::env::var("CMF_FFN_MEANFILL").ok()?;
10576 let b = std::fs::read(&p).ok()?;
10577 let inter = u32::from_le_bytes(b[4..8].try_into().ok()?) as usize;
10578 let vals: Vec<f32> = b[8..]
10579 .chunks_exact(4)
10580 .map(|c| f32::from_le_bytes([c[0], c[1], c[2], c[3]]))
10581 .collect();
10582 eprintln!("meanfill: {} value(s), inter {inter}", vals.len());
10583 Some((inter, vals))
10584 })
10585 .as_ref()
10586}
10587
10588fn probe_topk() -> usize {
10591 static K: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
10592 *K.get_or_init(|| {
10593 std::env::var("CMF_FFN_PROBE_TOPK")
10594 .ok()
10595 .and_then(|v| v.parse().ok())
10596 .unwrap_or(0)
10597 })
10598}
10599
10600thread_local! {
10601 static FFN_PROBE: std::cell::RefCell<Option<Vec<Vec<f64>>>> =
10604 const { std::cell::RefCell::new(None) };
10605}
10606
10607fn dense_ffn_dynamic(d: &DenseFfn, x: &[f32], pool: Option<&Pool>, k: usize) -> Option<Vec<f32>> {
10620 let dt = d.down_t.as_ref()?;
10621 let inter = d.gate_proj.rows();
10622 let hidden = dt.cols();
10623 if k == 0 || k >= inter || d.act != Act::Silu {
10624 return None;
10625 }
10626 DYN_SCRATCH.with(|sc| {
10627 let mut sc = sc.borrow_mut();
10628 let DynScratch {
10629 g,
10630 mag,
10631 live,
10632 parts,
10633 } = &mut *sc;
10634 g.resize(inter, 0.0);
10635 d.gate_proj.matvec(x, g, pool);
10636 for v in g.iter_mut() {
10637 *v = inference::silu(*v);
10638 }
10639 mag.clear();
10642 mag.extend(g.iter().map(|v| v.abs()));
10643 let (_, kth, _) = mag.select_nth_unstable_by(k - 1, |a, b| {
10644 b.partial_cmp(a).unwrap_or(std::cmp::Ordering::Equal)
10645 });
10646 let thr = *kth;
10647 live.clear();
10648 live.extend((0..inter as u32).filter(|&n| g[n as usize].abs() >= thr));
10649 let mut out = vec![0.0f32; hidden];
10650 match pool {
10651 Some(p) if live.len() >= 64 => {
10652 let nw = p.n_workers() + 1;
10653 parts.clear();
10654 parts.resize(nw * hidden, 0.0);
10655 let ptr = SendMut(parts.as_mut_ptr());
10656 let n = live.len();
10657 let live_ref: &[u32] = live;
10658 let g_ref: &[f32] = g;
10659 p.run(&|w, workers| {
10660 let chunk = n.div_ceil(workers);
10661 let (s, e) = (w * chunk, ((w + 1) * chunk).min(n));
10662 if s >= e {
10663 return;
10664 }
10665 WORKER_SCRATCH.with(|ws| {
10666 let mut ws = ws.borrow_mut();
10667 let [scratch, acc] = &mut *ws;
10668 scratch.resize(hidden.max(x.len()), 0.0);
10669 acc.clear();
10670 acc.resize(hidden, 0.0);
10671 for (o, &nrm) in live_ref[s..e].iter().enumerate() {
10672 if let Some(&nx) = live_ref[s..e].get(o + 1) {
10675 d.up_proj.prefetch_row(nx as usize);
10676 dt.prefetch_row(nx as usize);
10677 }
10678 let idx = nrm as usize;
10679 let up = d.up_proj.row_dot(idx, x, scratch);
10680 let a = g_ref[idx] * up;
10681 if a != 0.0 {
10682 dt.add_row_scaled(idx, a, acc, scratch);
10683 }
10684 }
10685 for (j, v) in acc.iter().enumerate() {
10686 unsafe { *ptr.at(w * hidden + j) = *v };
10687 }
10688 });
10689 });
10690 for w in 0..nw {
10691 for (j, o) in out.iter_mut().enumerate() {
10692 *o += parts[w * hidden + j];
10693 }
10694 }
10695 }
10696 _ => {
10697 WORKER_SCRATCH.with(|ws| {
10698 let mut ws = ws.borrow_mut();
10699 let [scratch, _acc] = &mut *ws;
10700 scratch.resize(hidden.max(x.len()), 0.0);
10701 for &nrm in live.iter() {
10702 let idx = nrm as usize;
10703 let up = d.up_proj.row_dot(idx, x, scratch);
10704 let a = g[idx] * up;
10705 if a != 0.0 {
10706 dt.add_row_scaled(idx, a, &mut out, scratch);
10707 }
10708 }
10709 });
10710 }
10711 }
10712 Some(out)
10713 })
10714}
10715
10716struct DynScratch {
10719 g: Vec<f32>,
10720 mag: Vec<f32>,
10721 live: Vec<u32>,
10722 parts: Vec<f32>,
10723}
10724
10725thread_local! {
10726 static DYN_SCRATCH: std::cell::RefCell<DynScratch> = const {
10727 std::cell::RefCell::new(DynScratch {
10728 g: Vec::new(),
10729 mag: Vec::new(),
10730 live: Vec::new(),
10731 parts: Vec::new(),
10732 })
10733 };
10734 static WORKER_SCRATCH: std::cell::RefCell<[Vec<f32>; 2]> =
10736 const { std::cell::RefCell::new([Vec::new(), Vec::new()]) };
10737}
10738
10739fn dense_ffn_masked(d: &DenseFfn, x: &[f32], pool: Option<&Pool>, mask_row: &[u8]) -> Vec<f32> {
10744 let inter = d.gate_proj.rows();
10745 FFN_SCRATCH.with(|s| {
10746 let mut s = s.borrow_mut();
10747 let [g, u, ..] = &mut *s;
10748 g.resize(inter, 0.0);
10749 if d.act == Act::Silu && QTensor::matvec_silu_mul(&d.gate_proj, &d.up_proj, x, g, pool) {
10750 } else {
10752 u.resize(inter, 0.0);
10753 QTensor::matvec_many([&d.gate_proj, &d.up_proj], x, [g, u], pool);
10754 for i in 0..inter {
10755 g[i] = d.act.combine(g[i], u[i]);
10756 }
10757 }
10758 zero_masked_cols(g, 1, inter, mask_row);
10759 let mut out = attention::take_buf(d.down_proj.rows());
10760 d.down_proj.matvec(g, &mut out, pool);
10761 out
10762 })
10763}
10764
10765fn dense_ffn_gpu(d: &DenseFfn, x: &[f32], _pool: Option<&Pool>) -> Option<Vec<f32>> {
10771 if d.act != Act::Silu {
10773 return None;
10774 }
10775 if d.gate_proj.rows() < crate::gpu::min_rows() && !d.gate_proj.is_q1() {
10778 return None;
10779 }
10780 let mut jobs: Vec<crate::gpu::MoeJob> = Vec::with_capacity(1);
10781 let mut model_ref = None;
10782 moe_push_job(d, x, 1.0, &mut jobs, &mut model_ref)?;
10783 let model = model_ref?;
10784 let hidden = jobs[0].down.1;
10785 let mut out = attention::take_buf(hidden);
10786 if crate::gpu::moe_block(&model, &jobs, &mut out) {
10787 Some(out)
10788 } else {
10789 let mut out = out;
10790 attention::recycle_buf(&mut out);
10791 None
10792 }
10793}
10794
10795#[allow(clippy::type_complexity)]
10800#[allow(clippy::type_complexity)]
10801pub(crate) fn moe_parts(
10802 t: &QTensor,
10803) -> Option<(
10804 &std::sync::Arc<cortiq_core::CmfModel>,
10805 usize,
10806 usize,
10807 usize,
10808 &[f32],
10809 &[f32],
10810 bool,
10811 bool,
10812 bool,
10813)> {
10814 match t {
10815 QTensor::Mapped {
10816 model,
10817 idx,
10818 dtype: dt @ (cortiq_core::TensorDtype::Q8_2f | cortiq_core::TensorDtype::Q8Row),
10819 rows,
10820 cols,
10821 row_scale,
10822 col_field,
10823 ..
10824 } if (*dt == cortiq_core::TensorDtype::Q8Row) || !col_field.is_empty() => Some((
10825 model, *idx, *rows, *cols, row_scale, col_field, false, false, false,
10826 )),
10827 QTensor::Mapped {
10829 model,
10830 idx,
10831 dtype: cortiq_core::TensorDtype::Q1,
10832 rows,
10833 cols,
10834 ..
10835 } => Some((
10836 model,
10837 *idx,
10838 *rows,
10839 *cols,
10840 &[][..],
10841 &[][..],
10842 true,
10843 false,
10844 false,
10845 )),
10846 QTensor::Mapped {
10848 model,
10849 idx,
10850 dtype: cortiq_core::TensorDtype::Q4Tiled,
10851 rows,
10852 cols,
10853 ..
10854 } => Some((
10855 model,
10856 *idx,
10857 *rows,
10858 *cols,
10859 &[][..],
10860 &[][..],
10861 false,
10862 true,
10863 false,
10864 )),
10865 QTensor::Mapped {
10867 model,
10868 idx,
10869 dtype: cortiq_core::TensorDtype::Q4TiledP,
10870 rows,
10871 cols,
10872 ..
10873 } => Some((
10874 model,
10875 *idx,
10876 *rows,
10877 *cols,
10878 &[][..],
10879 &[][..],
10880 false,
10881 true,
10882 false,
10883 )),
10884 QTensor::Mapped {
10888 model,
10889 idx,
10890 dtype: cortiq_core::TensorDtype::Q2TiledP,
10891 rows,
10892 cols,
10893 ..
10894 } => Some((
10895 model,
10896 *idx,
10897 *rows,
10898 *cols,
10899 &[][..],
10900 &[][..],
10901 false,
10902 true,
10903 true,
10904 )),
10905 _ => None,
10906 }
10907}
10908
10909#[cfg(target_os = "macos")]
10915fn metal_moe_graph_parts(m: &MoeFfn, hidden: usize) -> Option<crate::gpu::GpuMoe<'_>> {
10916 if m.router_sigmoid
10917 || m.router_input_norm
10918 || m.expert_bias.is_some()
10919 || m.route_tau.is_some()
10920 || m.mask.is_some()
10921 || m.per_expert_scale.is_some()
10922 || m.experts.is_empty()
10923 || m.top_k == 0
10924 || m.resonance.is_some()
10925 {
10926 return None;
10927 }
10928 let (sh, sg) = match &m.shared {
10931 Some((sh, Some(sg))) => (sh, sg),
10932 _ => return None,
10933 };
10934 let (rf, rr, rc) = m.router.f32_parts()?;
10935 if rr != m.experts.len() || rc != hidden {
10936 return None;
10937 }
10938 let (sf, sr, sc) = sg.f32_parts()?;
10939 if sr * sc != hidden {
10940 return None;
10941 }
10942 let inter = m.experts[0].gate_proj.rows();
10943 let gu_q2 = m.experts[0].gate_proj.mapped_q2tp().is_some();
10946 let trio = |e: &DenseFfn| -> Option<(usize, usize, usize)> {
10947 if e.act != Act::Silu
10948 || e.gate_proj.rows() != inter
10949 || e.gate_proj.cols() != hidden
10950 || e.up_proj.rows() != inter
10951 || e.up_proj.cols() != hidden
10952 || e.down_proj.rows() != hidden
10953 || e.down_proj.cols() != inter
10954 {
10955 return None;
10956 }
10957 let pick = |t: &QTensor| -> Option<usize> {
10958 if gu_q2 {
10959 t.mapped_q2tp().map(|(_, i)| i)
10960 } else {
10961 t.mapped_q4tp().map(|(_, i)| i)
10962 }
10963 };
10964 Some((
10965 pick(&e.gate_proj)?,
10966 pick(&e.up_proj)?,
10967 e.down_proj.mapped_q4tp().map(|(_, i)| i)?,
10968 ))
10969 };
10970 let experts = m.experts.iter().map(trio).collect::<Option<Vec<_>>>()?;
10971 let shared = trio(sh)?;
10972 Some(crate::gpu::GpuMoe {
10973 router: rf,
10974 sgate: sf,
10975 experts,
10976 shared,
10977 n_exp: m.experts.len(),
10978 top_k: m.top_k,
10979 inter,
10980 norm_topk: m.norm_topk_prob,
10981 route_scale: m.routed_scaling,
10982 gu_q2,
10983 })
10984}
10985
10986pub(crate) fn moe_push_job_parts<'a>(
10990 gate: &'a QTensor,
10991 up: &'a QTensor,
10992 down: &'a QTensor,
10993 x: &[f32],
10994 w: f32,
10995 swiglu_limit: f32,
10996 jobs: &mut Vec<crate::gpu::MoeJob<'a>>,
10997 model_ref: &mut Option<std::sync::Arc<cortiq_core::CmfModel>>,
10998) -> Option<()> {
10999 use crate::qtensor::prescale;
11000 let (gm, gi, gr, gc, grs, gcf, gq1, gq4, gq2) = moe_parts(gate)?;
11001 let (_, ui, ur, uc, urs, ucf, uq1, uq4, uq2) = moe_parts(up)?;
11002 let (_, di, dr, dc, drs, dcf, dq1, dq4, dq2) = moe_parts(down)?;
11003 if gq1 != uq1 || uq1 != dq1 || gq4 != uq4 || uq4 != dq4 || gq2 != uq2 {
11004 return None; }
11006 if gq2 && (dq2 || !dq4 || down.mapped_q4tp().is_none()) {
11009 return None;
11010 }
11011 if !gq2 && dq2 {
11012 return None;
11013 }
11014 model_ref.get_or_insert_with(|| gm.clone());
11015 let dt = |cf: &[f32]| {
11016 if cf.is_empty() {
11017 cortiq_core::TensorDtype::Q8Row
11018 } else {
11019 cortiq_core::TensorDtype::Q8_2f
11020 }
11021 };
11022 jobs.push(crate::gpu::MoeJob {
11023 gate: (gi, gr, gc, grs),
11024 up: (ui, ur, uc, urs),
11025 down: (di, dr, dc, drs),
11026 xs_gate: prescale(x, gcf, dt(gcf)).into_owned(),
11027 xs_up: prescale(x, ucf, dt(ucf)).into_owned(),
11028 down_col: dcf,
11029 w,
11030 q1: gq1,
11031 q4t: gq4 && !gq2 && gate.mapped_q4tp().is_none(),
11032 q4tp: gq4 && (gq2 || gate.mapped_q4tp().is_some()),
11033 gu_q2: gq2,
11034 swiglu_limit,
11035 });
11036 Some(())
11037}
11038
11039fn moe_push_job<'a>(
11041 d: &'a DenseFfn,
11042 x: &[f32],
11043 w: f32,
11044 jobs: &mut Vec<crate::gpu::MoeJob<'a>>,
11045 model_ref: &mut Option<std::sync::Arc<cortiq_core::CmfModel>>,
11046) -> Option<()> {
11047 use crate::qtensor::prescale;
11048 if d.act != Act::Silu {
11049 return None; }
11051 let (gm, gi, gr, gc, grs, gcf, gq1, gq4, gq2) = moe_parts(&d.gate_proj)?;
11052 let (_, ui, ur, uc, urs, ucf, uq1, uq4, uq2) = moe_parts(&d.up_proj)?;
11053 let (_, di, dr, dc, drs, dcf, dq1, dq4, dq2) = moe_parts(&d.down_proj)?;
11054 if gq1 != uq1 || uq1 != dq1 || gq4 != uq4 || uq4 != dq4 || gq2 != uq2 {
11055 return None; }
11057 if gq2 && (dq2 || !dq4 || d.down_proj.mapped_q4tp().is_none()) {
11058 return None;
11059 }
11060 if !gq2 && dq2 {
11061 return None;
11062 }
11063 model_ref.get_or_insert_with(|| gm.clone());
11064 let gdt = if gcf.is_empty() {
11065 cortiq_core::TensorDtype::Q8Row
11066 } else {
11067 cortiq_core::TensorDtype::Q8_2f
11068 };
11069 let udt = if ucf.is_empty() {
11070 cortiq_core::TensorDtype::Q8Row
11071 } else {
11072 cortiq_core::TensorDtype::Q8_2f
11073 };
11074 jobs.push(crate::gpu::MoeJob {
11075 gate: (gi, gr, gc, grs),
11076 up: (ui, ur, uc, urs),
11077 down: (di, dr, dc, drs),
11078 xs_gate: prescale(x, gcf, gdt).into_owned(),
11079 xs_up: prescale(x, ucf, udt).into_owned(),
11080 down_col: dcf,
11081 w,
11082 q1: gq1,
11083 q4t: gq4 && !gq2 && d.gate_proj.mapped_q4tp().is_none(),
11084 q4tp: gq4 && (gq2 || d.gate_proj.mapped_q4tp().is_some()),
11085 gu_q2: gq2,
11086 swiglu_limit: 0.0,
11087 });
11088 Some(())
11089}
11090
11091fn sparse_ffn_quant(
11098 d: &DenseFfn,
11099 x: &[f32],
11100 active: &[u16],
11101 hidden: usize,
11102 pool: Option<&Pool>,
11103) -> Vec<f32> {
11104 let n = active.len();
11105 let inter = d.gate_proj.rows();
11106 let mut act = vec![0.0f32; n];
11107 let need_scratch = !(d.gate_proj.sparse_col_ok() && d.up_proj.sparse_col_ok());
11110 let compute = |ai: usize| -> f32 {
11111 let idx = active[ai] as usize;
11112 if idx >= inter {
11113 return 0.0; }
11115 let mut s = if need_scratch {
11116 vec![0.0f32; hidden]
11117 } else {
11118 Vec::new()
11119 };
11120 let gate = d.gate_proj.row_dot(idx, x, &mut s);
11121 let up = d.up_proj.row_dot(idx, x, &mut s);
11122 d.act.combine(gate, up)
11123 };
11124 match pool {
11125 Some(p) if n >= 256 => {
11126 let ptr = SendMut(act.as_mut_ptr());
11127 p.run(&|widx, nw| {
11128 let chunk = n.div_ceil(nw);
11129 let (s, e) = (widx * chunk, ((widx + 1) * chunk).min(n));
11130 for ai in s..e {
11131 unsafe { *ptr.at(ai) = compute(ai) };
11132 }
11133 });
11134 }
11135 _ => {
11136 for (ai, a) in act.iter_mut().enumerate() {
11137 *a = compute(ai);
11138 }
11139 }
11140 }
11141 let mut out = vec![0.0f32; hidden];
11143 for (ai, &idx) in active.iter().enumerate() {
11144 let w = act[ai];
11145 if w.abs() >= 1e-12 && (idx as usize) < inter {
11146 d.down_proj.add_col_scaled(idx as usize, w, &mut out);
11147 }
11148 }
11149 out
11150}
11151
11152#[doc(hidden)]
11154pub fn sparse_ffn_quant_for_test(
11155 d: &DenseFfn,
11156 x: &[f32],
11157 active: &[u16],
11158 hidden: usize,
11159) -> Vec<f32> {
11160 sparse_ffn_quant(d, x, active, hidden, None)
11161}
11162
11163fn dequant_dense_f32(d: &DenseFfn) -> (Vec<f32>, Vec<f32>, Vec<f32>) {
11167 let deq = |t: &QTensor| -> Vec<f32> {
11168 let (rows, cols) = (t.rows(), t.cols());
11169 let mut out = vec![0.0f32; rows * cols];
11170 for r in 0..rows {
11171 t.row_f32(r, &mut out[r * cols..(r + 1) * cols]);
11172 }
11173 out
11174 };
11175 (deq(&d.gate_proj), deq(&d.up_proj), deq(&d.down_proj))
11176}
11177
11178struct SendMut(*mut f32);
11180unsafe impl Send for SendMut {}
11181unsafe impl Sync for SendMut {}
11182impl SendMut {
11183 #[inline]
11184 #[allow(clippy::mut_from_ref)]
11187 unsafe fn at(&self, i: usize) -> &mut f32 {
11188 unsafe { &mut *self.0.add(i) }
11189 }
11190}
11191
11192pub(crate) fn moe_route(
11202 logits: &[f32],
11203 m: &MoeFfn,
11204 allowed: Option<&[bool]>,
11205) -> (Vec<usize>, Vec<f32>, f32) {
11206 let ne = logits.len();
11207 let p: Vec<f32> = if m.router_sigmoid {
11208 logits.iter().map(|&l| 1.0 / (1.0 + (-l).exp())).collect()
11209 } else {
11210 let mx = logits.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
11211 let mut e: Vec<f32> = logits.iter().map(|&l| (l - mx).exp()).collect();
11212 let s: f32 = e.iter().sum();
11213 for v in &mut e {
11214 *v /= s;
11215 }
11216 e
11217 };
11218 let admit = |e: usize| {
11224 m.mask.as_ref().is_none_or(|mk| mk[e])
11225 && allowed.is_none_or(|a| a.get(e).copied().unwrap_or(false))
11226 };
11227 let mut idx: Vec<usize> = (0..ne).filter(|&e| admit(e)).collect();
11228 match &m.expert_bias {
11230 Some(b) => idx.sort_unstable_by(|&x, &y| {
11231 (p[y] + b[y])
11232 .partial_cmp(&(p[x] + b[x]))
11233 .unwrap()
11234 .then(x.cmp(&y))
11235 }),
11236 None => idx.sort_unstable_by(|&x, &y| p[y].partial_cmp(&p[x]).unwrap().then(x.cmp(&y))),
11237 }
11238 idx.truncate(m.top_k);
11239 if let Some(tau) = m.route_tau {
11243 let total: f32 = idx.iter().map(|&e| p[e]).sum();
11244 if total > 0.0 {
11245 let mut acc = 0.0f32;
11246 let mut keep = idx.len();
11247 for (i, &e) in idx.iter().enumerate() {
11248 acc += p[e];
11249 if acc >= tau * total {
11250 keep = i + 1;
11251 break;
11252 }
11253 }
11254 idx.truncate(keep);
11255 }
11256 }
11257 let wsum: f32 = if m.norm_topk_prob {
11258 let s: f32 = idx.iter().map(|&e| p[e]).sum();
11259 (if m.router_sigmoid { s + 1e-6 } else { s }) / m.routed_scaling
11262 } else {
11263 1.0 / m.routed_scaling
11264 };
11265 (idx, p, wsum)
11266}
11267
11268fn moe_trace(idx: &[usize]) {
11270 moe_trace_at(crate::gpu::cur_layer() as i32, idx)
11271}
11272
11273pub(crate) fn moe_trace_at(li: i32, idx: &[usize]) {
11276 use std::io::Write;
11277 static F: std::sync::OnceLock<Option<std::sync::Mutex<std::fs::File>>> =
11278 std::sync::OnceLock::new();
11279 let Some(f) = F.get_or_init(|| {
11280 let p = std::env::var("CMF_MOE_TRACE").ok()?;
11281 Some(std::sync::Mutex::new(
11282 std::fs::OpenOptions::new()
11283 .create(true)
11284 .append(true)
11285 .open(p)
11286 .ok()?,
11287 ))
11288 }) else {
11289 return;
11290 };
11291 let ids: Vec<String> = idx.iter().map(|e| e.to_string()).collect();
11292 let _ = writeln!(f.lock().unwrap(), "{li}:{}", ids.join(","));
11293}
11294
11295pub(crate) fn moe_ffn(
11298 m: &MoeFfn,
11299 x: &[f32],
11300 pool: Option<&Pool>,
11301 allowed: Option<&[bool]>,
11302) -> Vec<f32> {
11303 accumulate_act(m, x, 1);
11304 let ne = m.experts.len();
11305 let mut logits = vec![0.0f32; ne];
11306 match &m.resonance {
11307 Some(r) => r.scores(x, &mut logits),
11308 None => m.router.matvec(x, &mut logits, pool),
11309 }
11310 let (idx, p, wsum) = moe_route(&logits, m, allowed);
11311 {
11312 let mut st = m.stats.borrow_mut();
11313 if st.len() < ne {
11314 st.resize(ne, 0);
11315 }
11316 for &e in &idx {
11317 st[e] += 1;
11318 }
11319 }
11320 moe_trace(&idx);
11326 if crate::gpu::enabled_here() {
11331 match crate::gpu::probe_arm(crate::gpu::OpClass::Ffn) {
11332 crate::gpu::ProbeArm::Gpu => {
11333 let t0 = std::time::Instant::now();
11334 if let Some(out) = moe_ffn_gpu(m, x, &idx, &p, wsum, pool) {
11335 crate::gpu::probe_record(crate::gpu::OpClass::Ffn, true, t0.elapsed());
11336 return out;
11337 }
11338 }
11339 crate::gpu::ProbeArm::CpuTimed => {
11340 let t0 = std::time::Instant::now();
11341 let out = crate::gpu::cpu_scope(|| moe_ffn_cpu(m, x, &idx, &p, wsum, pool));
11342 crate::gpu::probe_record(crate::gpu::OpClass::Ffn, false, t0.elapsed());
11343 return out;
11344 }
11345 crate::gpu::ProbeArm::Cpu => {
11346 return crate::gpu::cpu_scope(|| moe_ffn_cpu(m, x, &idx, &p, wsum, pool));
11347 }
11348 }
11349 }
11350 moe_ffn_cpu(m, x, &idx, &p, wsum, pool)
11351}
11352
11353fn graph_note(built: bool) {
11357 use std::sync::atomic::{AtomicBool, Ordering};
11358 if built {
11359 GRAPH_TOK_OK.fetch_add(1, Ordering::Relaxed);
11360 } else {
11361 GRAPH_TOK_MISS.fetch_add(1, Ordering::Relaxed);
11362 }
11363 static SAID: AtomicBool = AtomicBool::new(false);
11364 if !SAID.swap(true, Ordering::Relaxed) {
11365 if built {
11366 tracing::info!("wgpu whole-token graph: ACTIVE");
11367 } else {
11368 tracing::warn!("wgpu whole-token graph refused — per-op path");
11369 }
11370 }
11371}
11372
11373pub static GRAPH_TOK_OK: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
11377pub static GRAPH_TOK_MISS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
11378
11379fn moe_batch_enabled() -> bool {
11382 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
11383 *ON.get_or_init(|| std::env::var("CMF_MOE_BATCH").as_deref() != Ok("0"))
11384}
11385
11386fn moe_ffn_cpu_batched(
11392 m: &MoeFfn,
11393 x: &[f32],
11394 idx: &[usize],
11395 p: &[f32],
11396 wsum: f32,
11397 pool: Option<&Pool>,
11398) -> Option<Vec<f32>> {
11399 if idx.is_empty() || !moe_batch_enabled() {
11400 return None;
11401 }
11402 if FFN_PROBE.with(|pr| pr.borrow().is_some()) {
11406 return None;
11407 }
11408 let n = idx.len() + usize::from(m.shared.is_some());
11409 let mut pairs = Vec::with_capacity(n);
11410 let mut downs = Vec::with_capacity(n);
11411 let mut ws = Vec::with_capacity(n);
11412 for &e in idx {
11413 let d = &m.experts[e];
11414 if d.act != Act::Silu {
11415 return None;
11416 }
11417 pairs.push((&d.gate_proj, &d.up_proj));
11418 downs.push(&d.down_proj);
11419 ws.push(p[e] / wsum * m.per_expert_scale.as_ref().map_or(1.0, |v| v[e]));
11420 }
11421 if let Some((se, gate)) = &m.shared {
11424 if se.act != Act::Silu {
11425 return None;
11426 }
11427 let g = gate.as_ref().map_or(1.0, |gate| {
11428 let mut gl = [0.0f32; 1];
11429 gate.matvec(x, &mut gl, pool);
11430 1.0 / (1.0 + (-gl[0]).exp())
11431 });
11432 pairs.push((&se.gate_proj, &se.up_proj));
11433 downs.push(&se.down_proj);
11434 ws.push(g);
11435 }
11436 let inter = pairs[0].0.rows();
11437 let mut gs: Vec<Vec<f32>> = (0..pairs.len()).map(|_| vec![0f32; inter]).collect();
11438 if !QTensor::moe_gate_up_many(&pairs, x, &mut gs, pool) {
11439 return None;
11440 }
11441 let mut out = attention::take_buf(x.len());
11442 if !QTensor::moe_down_many(&downs, &gs, &ws, &mut out, pool) {
11443 attention::recycle_buf(&mut out);
11444 return None;
11445 }
11446 Some(out)
11447}
11448
11449pub(crate) fn moe_cold_experts_cpu(
11455 experts: &[(&DenseFfn, f32)],
11456 x: &[f32],
11457 pool: Option<&Pool>,
11458) -> Vec<f32> {
11459 let mut out = attention::take_buf(x.len());
11460 if experts.is_empty() {
11461 return out;
11462 }
11463 let pairs: Vec<_> = experts
11464 .iter()
11465 .map(|(e, _)| (&e.gate_proj, &e.up_proj))
11466 .collect();
11467 let downs: Vec<_> = experts.iter().map(|(e, _)| &e.down_proj).collect();
11468 let weights: Vec<_> = experts.iter().map(|(_, w)| *w).collect();
11469 let inter = experts[0].0.gate_proj.rows();
11470 let mut activations: Vec<Vec<f32>> = (0..experts.len()).map(|_| vec![0.0; inter]).collect();
11471 if QTensor::moe_gate_up_many(&pairs, x, &mut activations, pool)
11472 && QTensor::moe_down_many(&downs, &activations, &weights, &mut out, pool)
11473 {
11474 return out;
11475 }
11476 out.fill(0.0);
11477 for &(expert, weight) in experts {
11478 let mut one = dense_ffn(expert, x, pool);
11479 for (o, v) in out.iter_mut().zip(&one) {
11480 *o += weight * v;
11481 }
11482 attention::recycle_buf(&mut one);
11483 }
11484 out
11485}
11486
11487fn moe_ffn_cpu(
11489 m: &MoeFfn,
11490 x: &[f32],
11491 idx: &[usize],
11492 p: &[f32],
11493 wsum: f32,
11494 pool: Option<&Pool>,
11495) -> Vec<f32> {
11496 if let Some(out) = moe_ffn_cpu_batched(m, x, idx, p, wsum, pool) {
11497 return out;
11498 }
11499 let mut out = attention::take_buf(x.len());
11500 for &e in idx {
11501 let mut eo = dense_ffn(&m.experts[e], x, pool);
11502 let w = p[e] / wsum * m.per_expert_scale.as_ref().map_or(1.0, |v| v[e]);
11503 for i in 0..out.len() {
11504 out[i] += w * eo[i];
11505 }
11506 attention::recycle_buf(&mut eo);
11507 }
11508 if let Some((se, gate)) = &m.shared {
11509 let mut so = dense_ffn(se, x, pool);
11510 let g = gate.as_ref().map_or(1.0, |gate| {
11511 let mut gl = [0.0f32; 1];
11512 gate.matvec(x, &mut gl, pool);
11513 1.0 / (1.0 + (-gl[0]).exp())
11514 });
11515 for i in 0..out.len() {
11516 out[i] += g * so[i];
11517 }
11518 attention::recycle_buf(&mut so);
11519 }
11520 out
11521}
11522
11523#[allow(clippy::too_many_arguments)]
11531fn mla_attention(
11532 w: &MlaWeights,
11533 normed: &[f32],
11534 cache: &mut crate::kv_cache::LayerKvCache,
11535 position: usize,
11536 inv_freq: &[f32],
11537 rope_scale: f32,
11538 eps: f64,
11539 pool: Option<&Pool>,
11540) -> Vec<f32> {
11541 let (nh, dr, dn, dv, lora) = (w.nh, w.qk_rope, w.qk_nope, w.v_dim, w.lora);
11542 let hd = dr + dn;
11543 let mut q = vec![0.0f32; nh * hd];
11544 match (&w.q_a, &w.q_a_norm) {
11545 (Some(qa), Some(qn)) => {
11546 let mut t = vec![0.0f32; qa.rows()];
11547 qa.matvec(normed, &mut t, pool);
11548 let tn = inference::rms_norm(&t, qn, eps, NormStyle::Qwen);
11549 w.q_proj.matvec(&tn, &mut q, pool);
11550 }
11551 _ => w.q_proj.matvec(normed, &mut q, pool),
11552 }
11553 let mut ca = vec![0.0f32; lora + dr];
11554 w.kv_a.matvec(normed, &mut ca, pool);
11555 let (c_lat, k_rope) = ca.split_at_mut(lora);
11556 let latn = inference::rms_norm(c_lat, &w.kv_a_norm, eps, NormStyle::Qwen);
11557 let mut kvb = vec![0.0f32; nh * (dn + dv)];
11558 w.kv_b.matvec(&latn, &mut kvb, pool);
11559 if !w.nope {
11560 attention::rope_rotate_scaled(k_rope, position, inv_freq, rope_scale);
11561 }
11562 for h in 0..nh {
11563 if !w.nope {
11564 attention::rope_rotate_scaled(
11565 &mut q[h * hd..h * hd + dr],
11566 position,
11567 inv_freq,
11568 rope_scale,
11569 );
11570 }
11571 }
11572 let mut k = vec![0.0f32; nh * hd];
11573 let mut v = vec![0.0f32; nh * hd];
11574 for h in 0..nh {
11575 k[h * hd..h * hd + dr].copy_from_slice(k_rope);
11576 k[h * hd + dr..(h + 1) * hd].copy_from_slice(&kvb[h * (dn + dv)..h * (dn + dv) + dn]);
11577 v[h * hd..h * hd + dv].copy_from_slice(&kvb[h * (dn + dv) + dn..(h + 1) * (dn + dv)]);
11578 }
11579 cache.append(&k, &v, &vec![true; nh]);
11580 let (ao, mut imp) = attention::attend_all_heads(&q, cache, nh, 1, hd, w.scale, None, 0.0);
11581 attention::recycle_buf(&mut imp);
11582 let mut ov = vec![0.0f32; nh * dv];
11583 for h in 0..nh {
11584 ov[h * dv..(h + 1) * dv].copy_from_slice(&ao[h * hd..h * hd + dv]);
11585 }
11586 let mut out = vec![0.0f32; w.o_proj.rows()];
11587 w.o_proj.matvec(&ov, &mut out, pool);
11588 out
11589}
11590
11591fn dense_moe_ffn(
11598 dm: &DenseMoeFfn,
11599 x_normed: &[f32],
11600 h_raw: &[f32],
11601 eps: f64,
11602 norm_style: NormStyle,
11603 pool: Option<&Pool>,
11604) -> Vec<f32> {
11605 let mut d = dense_ffn(&dm.dense, x_normed, pool);
11606 d = inference::rms_norm(&d, &dm.post_norm_1, eps, norm_style);
11607 let m = &dm.moe;
11608 let ne = m.experts.len();
11609 let mut logits = vec![0.0f32; ne];
11610 if m.router_input_norm {
11611 let ss: f32 = h_raw.iter().map(|v| v * v).sum::<f32>() / h_raw.len() as f32;
11612 let inv = 1.0 / (ss + eps as f32).sqrt();
11613 let xr: Vec<f32> = h_raw.iter().map(|v| v * inv).collect();
11614 m.router.matvec(&xr, &mut logits, pool);
11615 } else {
11616 m.router.matvec(h_raw, &mut logits, pool);
11617 }
11618 let (idx, p, wsum) = moe_route(&logits, m, None);
11619 {
11620 let mut st = m.stats.borrow_mut();
11621 if st.len() < ne {
11622 st.resize(ne, 0);
11623 }
11624 for &e in &idx {
11625 st[e] += 1;
11626 }
11627 }
11628 let x2 = inference::rms_norm(h_raw, &dm.pre_norm_2, eps, norm_style);
11629 let mo = moe_ffn_cpu(m, &x2, &idx, &p, wsum, pool);
11630 let mo = inference::rms_norm(&mo, &dm.post_norm_2, eps, norm_style);
11631 for (di, mi) in d.iter_mut().zip(&mo) {
11632 *di += mi;
11633 }
11634 d
11635}
11636
11637fn moe_gpu_refused(why: &'static str) {
11644 use std::sync::atomic::{AtomicBool, Ordering};
11645 static SAID: AtomicBool = AtomicBool::new(false);
11646 if !SAID.swap(true, Ordering::Relaxed) {
11647 tracing::warn!("MoE GPU block refused ({why}) — experts run on the CPU");
11648 }
11649}
11650
11651fn moe_ffn_gpu(
11652 m: &MoeFfn,
11653 x: &[f32],
11654 idx: &[usize],
11655 p: &[f32],
11656 wsum: f32,
11657 pool: Option<&Pool>,
11658) -> Option<Vec<f32>> {
11659 use crate::gpu::MoeJob;
11660
11661 let mut jobs: Vec<MoeJob> = Vec::with_capacity(idx.len() + 1);
11662 let mut model_ref = None;
11663 for &e in idx {
11664 if moe_push_job(&m.experts[e], x, p[e] / wsum, &mut jobs, &mut model_ref).is_none() {
11665 moe_gpu_refused("push_job(expert)");
11666 return None;
11667 }
11668 }
11669 if let Some((se, gate)) = &m.shared {
11670 let g = gate.as_ref().map_or(1.0, |gate| {
11671 let mut gl = [0.0f32; 1];
11672 gate.matvec(x, &mut gl, pool);
11673 1.0 / (1.0 + (-gl[0]).exp())
11674 });
11675 if moe_push_job(se, x, g, &mut jobs, &mut model_ref).is_none() {
11676 moe_gpu_refused("push_job(shared)");
11677 return None;
11678 }
11679 }
11680 let Some(model) = model_ref else {
11681 moe_gpu_refused("no model_ref");
11682 return None;
11683 };
11684 let hidden = jobs[0].down.1;
11685 let mut out = vec![0.0f32; hidden];
11686 if crate::gpu::moe_block(&model, &jobs, &mut out) {
11687 Some(out)
11688 } else {
11689 moe_gpu_refused("gpu::moe_block");
11690 None
11691 }
11692}
11693
11694fn ffn_forward(
11696 ffn: &FfnKind,
11697 x: &[f32],
11698 pool: Option<&Pool>,
11699 experts_allowed: Option<&[bool]>,
11700) -> Vec<f32> {
11701 match ffn {
11702 FfnKind::Dense(d) if !d.segs.is_empty() => tube_ffn(d, x, 1, pool, None),
11703 FfnKind::Dense(d) => dense_ffn(d, x, pool),
11704 FfnKind::Moe(m) => moe_ffn(m, x, pool, experts_allowed),
11705 FfnKind::DenseMoe(_) => unreachable!("DenseMoe dispatches via dense_moe_ffn"),
11709 }
11710}
11711
11712fn ffn_forward_pair(
11716 ffn: &FfnKind,
11717 x1: &[f32],
11718 x2: &[f32],
11719 pool: Option<&Pool>,
11720 experts_allowed: Option<&[bool]>,
11721) -> (Vec<f32>, Vec<f32>) {
11722 let d = match ffn {
11723 FfnKind::Dense(d) if !d.segs.is_empty() => {
11726 return (
11727 tube_ffn(d, x1, 1, pool, None),
11728 tube_ffn(d, x2, 1, pool, None),
11729 );
11730 }
11731 FfnKind::Dense(d) => d,
11732 FfnKind::Moe(m) => {
11733 return (
11734 moe_ffn(m, x1, pool, experts_allowed),
11735 moe_ffn(m, x2, pool, experts_allowed),
11736 );
11737 }
11738 FfnKind::DenseMoe(_) => unreachable!("DenseMoe dispatches via dense_moe_ffn"),
11739 };
11740 let inter = d.gate_proj.rows();
11741 FFN_SCRATCH.with(|s| {
11742 let mut s = s.borrow_mut();
11743 let [g1, g2, u1, u2] = &mut *s;
11744 g1.resize(inter, 0.0);
11745 g2.resize(inter, 0.0);
11746 u1.resize(inter, 0.0);
11747 u2.resize(inter, 0.0);
11748 QTensor::matvec2_many(
11751 [&d.gate_proj, &d.up_proj],
11752 x1,
11753 x2,
11754 [g1.as_mut_slice(), u1.as_mut_slice()],
11755 [g2.as_mut_slice(), u2.as_mut_slice()],
11756 pool,
11757 );
11758 for i in 0..inter {
11759 g1[i] = d.act.combine(g1[i], u1[i]);
11760 g2[i] = d.act.combine(g2[i], u2[i]);
11761 }
11762 let mut o1 = attention::take_buf(d.down_proj.rows());
11763 let mut o2 = attention::take_buf(d.down_proj.rows());
11764 d.down_proj.matvec2(g1, g2, &mut o1, &mut o2, pool);
11765 (o1, o2)
11766 })
11767}
11768
11769#[cfg(test)]
11770mod tests {
11771
11772 #[test]
11773 fn cancel_flag_stops_generation() {
11774 let mut p = create_test_pipeline(16, 32, 2, 2, 8, 2, 32);
11775 p.cancel.store(true, std::sync::atomic::Ordering::Relaxed);
11778 let r = p.generate_from_ids(&[1, 2, 3], 8, None, None).unwrap();
11779 assert_eq!(r.finish_reason, "cancelled");
11780 assert!(
11781 r.token_ids.is_empty(),
11782 "no tokens after cancel: {:?}",
11783 r.token_ids
11784 );
11785 let r2 = p.generate_from_ids(&[1, 2, 3], 4, None, None).unwrap();
11787 assert_ne!(r2.finish_reason, "cancelled");
11788 }
11789 use super::*;
11790
11791 #[test]
11799 fn dynamic_ffn_equals_the_zeroing_arm() {
11800 let (hidden, inter) = (8usize, 32usize);
11801 let synth = |n: usize, salt: usize| -> Vec<f32> {
11802 (0..n)
11803 .map(|i| (((i * 29 + salt * 13 + 7) % 89) as f32 / 89.0 - 0.5) * 0.6)
11804 .collect()
11805 };
11806 let down = synth(hidden * inter, 3);
11807 let mut down_t = vec![0.0f32; inter * hidden];
11808 for r in 0..hidden {
11809 for c in 0..inter {
11810 down_t[c * hidden + r] = down[r * inter + c];
11811 }
11812 }
11813 let d = DenseFfn {
11814 gate_proj: QTensor::from_f32(synth(inter * hidden, 1), inter, hidden),
11815 up_proj: QTensor::from_f32(synth(inter * hidden, 2), inter, hidden),
11816 down_proj: QTensor::from_f32(down.clone(), hidden, inter),
11817 act: Act::Silu,
11818 down_t: Some(QTensor::from_f32(down_t, inter, hidden)),
11819 segs: Vec::new(),
11820 };
11821 let x = synth(hidden, 11);
11822 let k = 12usize;
11823 let got = dense_ffn_dynamic(&d, &x, None, k).expect("down_t present");
11824 let mut g = vec![0.0f32; inter];
11826 d.gate_proj.matvec(&x, &mut g, None);
11827 let mut u = vec![0.0f32; inter];
11828 d.up_proj.matvec(&x, &mut u, None);
11829 for v in g.iter_mut() {
11830 *v = inference::silu(*v);
11831 }
11832 keep_top_k(&mut g, k);
11833 for i in 0..inter {
11834 g[i] *= u[i];
11835 }
11836 let mut want = vec![0.0f32; hidden];
11837 d.down_proj.matvec(&g, &mut want, None);
11838 for (a, b) in want.iter().zip(&got) {
11839 assert!((a - b).abs() < 1e-5, "dynamic {b} vs reference {a}");
11840 }
11841 }
11842
11843 #[test]
11849 fn tube_ffn_open_equals_dense_and_closed_equals_masked() {
11850 let (hidden, core, tube) = (8usize, 12usize, 8usize);
11851 let inter = core + tube;
11852 let synth = |n: usize, salt: usize| -> Vec<f32> {
11853 (0..n)
11854 .map(|i| (((i * 41 + salt * 17 + 5) % 97) as f32 / 97.0 - 0.5) * 0.5)
11855 .collect()
11856 };
11857 let (g_all, u_all) = (synth(inter * hidden, 1), synth(inter * hidden, 2));
11858 let d_all = synth(hidden * inter, 3);
11859 let dense = DenseFfn {
11861 gate_proj: QTensor::from_f32(g_all.clone(), inter, hidden),
11862 up_proj: QTensor::from_f32(u_all.clone(), inter, hidden),
11863 down_proj: QTensor::from_f32(d_all.clone(), hidden, inter),
11864 act: Act::Silu,
11865 down_t: None,
11866 segs: Vec::new(),
11867 };
11868 let rows =
11869 |v: &[f32], a: usize, b: usize| -> Vec<f32> { v[a * hidden..b * hidden].to_vec() };
11870 let cols = |v: &[f32], a: usize, b: usize| -> Vec<f32> {
11871 let mut o = Vec::with_capacity(hidden * (b - a));
11872 for r in 0..hidden {
11873 o.extend_from_slice(&v[r * inter + a..r * inter + b]);
11874 }
11875 o
11876 };
11877 let tubed = DenseFfn {
11878 down_t: None,
11879 gate_proj: QTensor::from_f32(rows(&g_all, 0, core), core, hidden),
11880 up_proj: QTensor::from_f32(rows(&u_all, 0, core), core, hidden),
11881 down_proj: QTensor::from_f32(cols(&d_all, 0, core), hidden, core),
11882 act: Act::Silu,
11883 segs: vec![FfnSeg {
11884 gate: QTensor::from_f32(rows(&g_all, core, inter), tube, hidden),
11885 up: QTensor::from_f32(rows(&u_all, core, inter), tube, hidden),
11886 down: QTensor::from_f32(cols(&d_all, core, inter), hidden, tube),
11887 start: core,
11888 width: tube,
11889 }],
11890 };
11891 let x = synth(hidden, 7);
11892 let want = dense_ffn(&dense, &x, None);
11893 let got = tube_ffn(&tubed, &x, 1, None, None);
11894 for (a, b) in want.iter().zip(&got) {
11895 assert!((a - b).abs() < 1e-5, "open tube: {a} vs {b}");
11896 }
11897 let mut bits = vec![0u8; inter.div_ceil(8)];
11899 for n in 0..core {
11900 bits[n / 8] |= 1 << (n % 8);
11901 }
11902 let closed = tube_ffn(&tubed, &x, 1, None, Some(&bits));
11903 let masked = dense_ffn_masked(&dense, &x, None, &bits);
11904 for (a, b) in masked.iter().zip(&closed) {
11905 assert!((a - b).abs() < 1e-5, "closed tube: {a} vs {b}");
11906 }
11907 let batch = tube_ffn(&tubed, &x, 1, None, Some(&bits));
11909 for (a, b) in closed.iter().zip(&batch) {
11910 assert_eq!(a, b, "batch arm disagrees with decode arm");
11911 }
11912 }
11913
11914 #[test]
11916 fn sparse_ffn_quant_equals_dense_with_inactive_zeroed() {
11917 let (hidden, inter) = (16usize, 40usize);
11918 let synth = |n: usize, salt: usize| -> Vec<f32> {
11919 (0..n)
11920 .map(|i| (((i * 37 + salt * 11 + 3) % 101) as f32 / 101.0 - 0.5) * 0.4)
11921 .collect()
11922 };
11923 let d = DenseFfn {
11924 gate_proj: QTensor::from_f32(synth(inter * hidden, 1), inter, hidden),
11925 up_proj: QTensor::from_f32(synth(inter * hidden, 2), inter, hidden),
11926 down_proj: QTensor::from_f32(synth(hidden * inter, 3), hidden, inter),
11927 act: Act::Silu,
11928 down_t: None,
11929 segs: Vec::new(),
11930 };
11931 let x = synth(hidden, 9);
11932 let active: Vec<u16> = (0..inter as u16).filter(|i| i % 3 == 0).collect();
11934
11935 let sparse = sparse_ffn_quant(&d, &x, &active, hidden, None);
11936
11937 let mut g = vec![0.0f32; inter];
11939 d.gate_proj.matvec(&x, &mut g, None);
11940 let mut u = vec![0.0f32; inter];
11941 d.up_proj.matvec(&x, &mut u, None);
11942 let act_set: std::collections::HashSet<u16> = active.iter().copied().collect();
11943 for i in 0..inter {
11944 g[i] = if act_set.contains(&(i as u16)) {
11945 inference::silu(g[i]) * u[i]
11946 } else {
11947 0.0
11948 };
11949 }
11950 let mut reference = vec![0.0f32; hidden];
11951 d.down_proj.matvec(&g, &mut reference, None);
11952
11953 let max_d = sparse
11954 .iter()
11955 .zip(&reference)
11956 .map(|(a, b)| (a - b).abs())
11957 .fold(0.0f32, f32::max);
11958 assert!(max_d < 1e-5, "sparse != dense-zeroed: max|Δ| = {max_d}");
11959 }
11960
11961 fn attach_test_mtp(p: &mut Pipeline) {
11963 let (h, inter, heads, kv, hd) = (
11964 p.hidden_size,
11965 p.intermediate_size,
11966 p.num_heads,
11967 p.num_kv_heads,
11968 p.head_dim,
11969 );
11970 let synth = |n: usize, salt: usize| -> Vec<f32> {
11971 (0..n)
11972 .map(|i| (((i * 29 + salt * 23 + 5) % 101) as f32 / 101.0 - 0.5) * 0.2)
11973 .collect()
11974 };
11975 let qt = |rows: usize, cols: usize, salt: usize| -> QTensor {
11976 QTensor::from_f32(synth(rows * cols, salt), rows, cols)
11977 };
11978 p.mtp = Some(MtpModule {
11979 enorm: vec![1.0; h],
11980 hnorm: vec![1.0; h],
11981 eh_proj: qt(h, 2 * h, 301),
11982 layer: LayerWeights {
11983 input_norm: vec![1.0; h],
11984 post_norm: vec![1.0; h],
11985 attn_out_norm: None,
11986 ffn_out_norm: None,
11987 layer_scale: None,
11988 ffn: FfnKind::Dense(DenseFfn {
11989 gate_proj: qt(inter, h, 315),
11990 up_proj: qt(inter, h, 316),
11991 down_proj: qt(h, inter, 317),
11992 act: Act::Silu,
11993 down_t: None,
11994 segs: Vec::new(),
11995 }),
11996 attn: AttnKind::Full {
11997 bias: None,
11998 wq: qt(heads * hd, h, 311),
11999 wk: qt(kv * hd, h, 312),
12000 wv: qt(kv * hd, h, 313),
12001 wo: qt(h, heads * hd, 314),
12002 q_norm: None,
12003 k_norm: None,
12004 output_gate: false,
12005 softplus_gate: None,
12006 },
12007 },
12008 final_norm: vec![1.0; h],
12009 kv: crate::kv_cache::LayerKvCache::new(kv, hd),
12010 });
12011 }
12012
12013 #[test]
12014 fn speculative_equals_vanilla_greedy() {
12015 unsafe { std::env::set_var("CMF_GPU_WGPU_GRAPH", "0") };
12019 let run = |spec: bool| {
12020 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
12021 p.sampler_config.temperature = 0.0;
12022 attach_test_mtp(&mut p);
12023 p.speculative = spec;
12024 let r = p.generate("abcdef", 12, None, None).unwrap();
12025 (r.token_ids, r.mtp_drafted, r.mtp_accepted)
12026 };
12027 let (vanilla, d0, _) = run(false);
12028 let (spec, d1, a1) = run(true);
12029 assert_eq!(d0, 0, "vanilla path must not draft");
12030 assert!(d1 > 0, "speculative path must draft");
12031 assert_eq!(
12032 vanilla, spec,
12033 "speculative must reproduce the exact greedy sequence (accepted {a1}/{d1})"
12034 );
12035 }
12036
12037 #[test]
12038 fn speculative_accepts_constant_oracle() {
12039 unsafe { std::env::set_var("CMF_GPU_WGPU_GRAPH", "0") };
12041 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
12042 p.sampler_config.temperature = 0.0;
12043 p.sampler_config.repetition_penalty = 1.0;
12044 p.weights.lm_head = QTensor::from_f32(vec![0.01; 64 * 8], 64, 8);
12047 attach_test_mtp(&mut p);
12048 p.speculative = true;
12049 let r = p.generate("abcd", 10, None, None).unwrap();
12050 assert!(r.mtp_drafted > 0);
12051 assert_eq!(
12052 r.mtp_accepted, r.mtp_drafted,
12053 "constant logits → every draft accepted"
12054 );
12055 assert!(r.token_ids.windows(2).all(|w| w[0] == w[1]));
12058 }
12059
12060 #[test]
12061 fn empty_prompt_is_an_error_not_a_panic() {
12062 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 260);
12063 let r = p.generate("", 4, None, None);
12064 assert!(r.is_err(), "empty prompt must be a clean error");
12065 }
12066
12067 #[test]
12068 fn every_token_enters_kv_exactly_once() {
12069 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
12070 p.sampler_config.temperature = 0.0;
12072 let r = p.generate("abc", 2, None, None).unwrap();
12073 assert_eq!(r.prompt_tokens, 3);
12074 assert_eq!(
12078 p.kv_cache.seq_len(),
12079 3 + r.tokens_generated - 1,
12080 "each token must be cached exactly once (v1 cached the last prompt token twice)"
12081 );
12082 }
12083
12084 #[test]
12085 fn generation_is_reproducible_with_seed() {
12086 let run = || {
12087 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
12088 p.generate("hello", 8, None, None).unwrap().token_ids
12089 };
12090 assert_eq!(run(), run());
12091 }
12092
12093 #[test]
12094 fn resetting_sampler_restarts_the_seeded_stream() {
12095 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
12096 let config = SamplerConfig {
12097 seed: Some(1234),
12098 ..SamplerConfig::default()
12099 };
12100 p.set_sampler_config(config.clone());
12101 let first = p.generate("hello", 8, None, None).unwrap().token_ids;
12102 p.set_sampler_config(config);
12103 let second = p.generate("hello", 8, None, None).unwrap().token_ids;
12104 assert_eq!(first, second);
12105 }
12106
12107 #[test]
12108 fn eviction_bounds_the_cache() {
12109 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 260);
12110 p.kv_cache.max_seq_len = 6;
12111 p.sampler_config.temperature = 0.0;
12112 let _ = p.generate("abcd", 12, None, None).unwrap();
12113 assert!(
12114 p.kv_cache.seq_len() <= 6 + 1,
12115 "cache must stay bounded by max_seq_len (got {})",
12116 p.kv_cache.seq_len()
12117 );
12118 }
12119
12120 #[test]
12121 fn confidence_matches_tokens_and_is_a_probability() {
12122 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
12123 p.sampler_config.temperature = 0.0;
12124 p.sampler_config.repetition_penalty = 1.0;
12125 let r = p.generate("abcd", 10, None, None).unwrap();
12126 assert_eq!(
12127 r.token_confidence.len(),
12128 r.token_ids.len(),
12129 "one confidence per emitted token"
12130 );
12131 for &c in &r.token_confidence {
12132 assert!((0.0..=1.0).contains(&c), "confidence out of [0,1]: {c}");
12133 }
12134 let logits = [1.0f32, 3.0, 0.5, 3.0];
12136 let p0 = top1_prob_t(&logits, 1, 1.0);
12137 let p1 = top1_prob_t(&logits, 3, 1.0);
12138 assert!((p0 - p1).abs() < 1e-6, "equal logits → equal prob");
12139 assert!(p0 > 0.0 && p0 < 1.0);
12140 let sharp = top1_prob_t(&logits, 1, 1.0);
12142 let soft = top1_prob_t(&logits, 1, 2.0);
12143 assert!(soft < sharp, "higher temperature lowers peak confidence");
12144 }
12145
12146 #[test]
12147 fn trace_is_opt_in_and_parallels_the_output() {
12148 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
12150 p.sampler_config.temperature = 0.0;
12151 p.sampler_config.repetition_penalty = 1.0;
12152 let r = p.generate("abcd", 10, None, None).unwrap();
12153 assert!(r.traces.is_empty(), "trace must be empty unless enabled");
12154
12155 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
12157 p.sampler_config.temperature = 0.0;
12158 p.sampler_config.repetition_penalty = 1.0;
12159 p.set_trace(true);
12160 let r = p.generate("abcd", 10, None, None).unwrap();
12161 assert_eq!(r.traces.len(), r.token_ids.len(), "one trace row per token");
12162 for (i, tr) in r.traces.iter().enumerate() {
12163 assert_eq!(tr.t, i, "trace index is sequential");
12164 assert_eq!(tr.token_id, r.token_ids[i], "trace token_id matches output");
12165 assert_eq!(
12166 tr.confidence, r.token_confidence[i],
12167 "trace confidence matches the confidence channel"
12168 );
12169 assert!(tr.active_skill.is_none() && tr.recon.is_none() && !tr.switched);
12171 }
12172 }
12173
12174 #[test]
12175 fn explain_prefill_logits_match_greedy_first_token() {
12176 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
12180 p.sampler_config.temperature = 0.0;
12181 p.sampler_config.repetition_penalty = 1.0;
12182 let ids = p.tokenizer.encode("abcd");
12183 let logits = p.prefill_next_logits(&ids, None);
12184 let argmax = logits
12185 .iter()
12186 .enumerate()
12187 .max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
12188 .unwrap()
12189 .0 as u32;
12190 let r = p.generate("abcd", 1, None, None).unwrap();
12191 assert_eq!(
12192 argmax, r.token_ids[0],
12193 "explain preview must match greedy emit"
12194 );
12195 }
12196
12197 #[test]
12198 fn laguna_shared_expert_is_unconditionally_added() {
12199 let matrix = |values: Vec<f32>| QTensor::from_f32(values, 2, 2);
12200 let identity = || matrix(vec![1.0, 0.0, 0.0, 1.0]);
12201 let zero_dense = || DenseFfn {
12202 gate_proj: matrix(vec![0.0; 4]),
12203 up_proj: matrix(vec![0.0; 4]),
12204 down_proj: matrix(vec![0.0; 4]),
12205 act: Act::Silu,
12206 down_t: None,
12207 segs: Vec::new(),
12208 };
12209 let shared = DenseFfn {
12210 gate_proj: identity(),
12211 up_proj: identity(),
12212 down_proj: identity(),
12213 act: Act::Silu,
12214 down_t: None,
12215 segs: Vec::new(),
12216 };
12217 let x = [1.0, 2.0];
12218 let expected = dense_ffn(&shared, &x, None);
12219 let moe = MoeFfn {
12220 router: QTensor::from_f32(vec![0.0, 0.0], 1, 2),
12221 experts: vec![zero_dense()],
12222 top_k: 1,
12223 norm_topk_prob: true,
12224 router_sigmoid: true,
12225 expert_bias: None,
12226 routed_scaling: 1.0,
12227 route_tau: None,
12228 shared: Some((shared, None)),
12229 stats: std::cell::RefCell::new(Vec::new()),
12230 act_sq: std::cell::RefCell::new(Vec::new()),
12231 act_rows: std::cell::RefCell::new(Vec::new()),
12232 mask: None,
12233 per_expert_scale: None,
12234 router_input_norm: false,
12235 resonance: None,
12236 };
12237 let actual = moe_ffn_cpu(&moe, &x, &[0], &[0.0], 1.0, None);
12238 for (actual, expected) in actual.iter().zip(expected) {
12239 assert!((actual - expected).abs() < 1e-6);
12240 }
12241 }
12242}