1use crate::attention::{self, QwenAttnCfg};
10use crate::inference;
11use crate::kv_cache::KvCache;
12use crate::linear_core::{
13 GdnCfg, GdnWeights, ShortConvCfg, ShortConvWeights, VmfPhaseCfg, VmfPhaseWeights, gdn_forward,
14 gdn_pair, short_conv_forward, short_conv_forward_batch, short_conv_pair, vmf_phase_forward,
15 vmf_phase_pair,
16};
17use crate::pool::Pool;
18use crate::qtensor::QTensor;
19use crate::sampler::{self, SamplerConfig, SamplerScratch, SplitMix64};
20use crate::tokenizer::Tokenizer;
21use cortiq_core::mask::TaskMask;
22use cortiq_core::types::NormStyle;
23
24pub static GLOBAL_USE_GPU: std::sync::atomic::AtomicBool =
25 std::sync::atomic::AtomicBool::new(false);
26
27struct ForwardScratch {
31 n1: Vec<f32>,
32 n2: Vec<f32>,
33 p1: Vec<f32>,
34 p2: Vec<f32>,
35}
36
37impl ForwardScratch {
38 fn new(hidden: usize) -> Self {
39 Self {
40 n1: vec![0.0; hidden],
41 n2: vec![0.0; hidden],
42 p1: vec![0.0; hidden],
43 p2: vec![0.0; hidden],
44 }
45 }
46}
47
48pub struct Pipeline {
50 gpu_plan: Option<std::sync::Arc<Vec<(usize, usize, usize)>>>,
55 pub tokenizer: std::sync::Arc<Tokenizer>,
58 pub kv_cache: KvCache,
59 pub sampler_config: SamplerConfig,
60 pub weights: PipelineWeights,
61 pub hidden_size: usize,
62 pub intermediate_size: usize,
63 pub num_heads: usize,
64 pub num_kv_heads: usize,
65 pub head_dim: usize,
66 pub num_layers: usize,
68 pub physical_layers: usize,
70 pub loop_final_norm: bool,
72 pub vocab_size: usize,
73 pub rms_eps: f64,
74 pub rope_base: f32,
75 pub norm_style: NormStyle,
76 pub rotary_dim: usize,
78 pub attention_heads_per_layer: Option<Vec<usize>>,
80 pub vmf_cfg: Option<VmfPhaseCfg>,
82 pub gdn_cfg: Option<GdnCfg>,
84 pub logit_multiplier: Option<f32>,
86 pub cancel: std::sync::Arc<std::sync::atomic::AtomicBool>,
91 pub kv_history: Vec<u32>,
96 pub kda_cfg: Option<crate::linear_core::KdaCfg>,
98 pub g3n: Option<Box<(crate::g3n::G3nGlobals, Vec<crate::g3n::G3nLayer>)>>,
101 pub dsv4: Option<
105 Box<(
106 crate::dsv4::Dsv4Globals,
107 Vec<crate::dsv4::Dsv4Layer>,
108 crate::dsv4::Dsv4Cfg,
109 crate::dsv4::Dsv4State,
110 )>,
111 >,
112 pub dsv4_mtp: Vec<crate::dsv4::Dsv4Mtp>,
116 pub dspark: Option<crate::dsv4::DsparkState>,
118 pub dspark_pending: Vec<(usize, Vec<u32>, bool, usize)>,
121 pub dspark_hist: Vec<usize>,
123 pub dspark_real: Vec<u32>,
127 pub dspark_trunk_picks: Vec<Vec<(usize, Vec<usize>)>>,
131 pub dspark_exp: Vec<(usize, usize, usize, usize)>,
133 pub dspark_draft_ns: u128,
137 pub short_conv_cfg: Option<ShortConvCfg>,
140 pub mtp: Option<MtpModule>,
142 pub speculative: bool,
144 rng: SplitMix64,
145 sampler_scratch: SamplerScratch,
146 pub(crate) inv_freq: std::sync::Arc<Vec<f32>>,
150 ws: ForwardScratch,
154 pool: Option<std::sync::Arc<Pool>>,
156 pub(crate) model: Option<std::sync::Arc<cortiq_core::CmfModel>>,
160 pub(crate) dyn_force_f32: bool,
162 pub(crate) dyn_skill_layers: Vec<Option<Vec<usize>>>,
167 pub(crate) dyn_active: Option<usize>,
173 pub(crate) dyn_blend_loaded: bool,
177 pub(crate) dyn_phi_layer: Option<usize>,
180 dyn_phi_ema: Vec<f32>,
182 dyn_phi_seen: usize,
183 pub dyn_router: Option<crate::swarm::DynRouter>,
186 o1_cfg: Option<crate::nystrom::O1Cfg>,
189 o1_epoch: u64,
192 o1_flags: Vec<bool>,
194 trace: bool,
197 calib_temp: f32,
200 #[cfg_attr(not(target_os = "macos"), allow(dead_code))]
202 graph_kv_id: u64,
203 #[cfg_attr(not(target_os = "macos"), allow(dead_code))]
206 graph_want_logits: bool,
207 graph_logits: Option<Vec<f32>>,
210 pub embed_multiplier: f32,
212 pub attn_scale: f32,
215 pub swa: Option<(usize, usize)>,
218 pub sliding_layers: Option<Vec<bool>>,
221 pub inv_freq_local: Option<std::sync::Arc<Vec<f32>>>,
224 pub rotary_dim_local: Option<usize>,
225 pub rope_scale: f32,
226 pub rope_scale_local: f32,
227 pub global_attn: Option<(usize, usize)>,
230 pub inv_freq_global: Option<std::sync::Arc<Vec<f32>>>,
233 pub attn_v_norm: bool,
235 pub final_softcap: Option<f32>,
237 pub attn_softcap: f32,
239 confidence_on: bool,
243}
244
245#[cfg(target_os = "macos")]
246impl Drop for Pipeline {
247 fn drop(&mut self) {
248 crate::gpu::kv_mirror_drop(self.graph_kv_id);
249 }
250}
251
252pub struct PipelineWeights {
257 pub embed_tokens: QTensor,
259 pub layers: Vec<LayerWeights>,
261 pub lm_head: QTensor,
263 pub final_norm: Vec<f32>,
265}
266
267pub struct LayerWeights {
269 pub input_norm: Vec<f32>,
270 pub post_norm: Vec<f32>,
273 pub attn_out_norm: Option<Vec<f32>>,
276 pub layer_scale: Option<f32>,
278 pub ffn_out_norm: Option<Vec<f32>>,
281 pub ffn: FfnKind,
282 pub attn: AttnKind,
283}
284
285#[derive(Clone, Copy, PartialEq, Debug, Default)]
288pub enum Act {
289 #[default]
290 Silu,
291 GeluTanh,
292 Situ {
295 beta: f32,
296 linear_beta: f32,
297 },
298}
299
300impl Act {
301 pub fn from_arch(name: &str) -> Self {
302 if name == "gelu_tanh" {
303 Self::GeluTanh
304 } else {
305 Self::Silu
306 }
307 }
308
309 pub fn from_arch_full(arch: &cortiq_core::ModelArch) -> Self {
311 match arch.hidden_act.as_str() {
312 "situ" => Self::Situ {
313 beta: arch.activation_situ_beta.unwrap_or(1.0) as f32,
314 linear_beta: arch.activation_situ_linear_beta.unwrap_or(0.0) as f32,
315 },
316 other => Self::from_arch(other),
317 }
318 }
319
320 #[inline]
321 pub fn apply(self, x: f32) -> f32 {
322 match self {
323 Self::Silu => inference::silu(x),
324 Self::GeluTanh => inference::gelu_tanh(x),
325 Self::Situ { beta, .. } => beta * (x / beta).tanh() * (1.0 / (1.0 + (-x).exp())),
326 }
327 }
328
329 #[inline]
332 pub fn combine(self, g: f32, u: f32) -> f32 {
333 match self {
334 Self::Situ { linear_beta, .. } if linear_beta > 0.0 => {
335 self.apply(g) * (linear_beta * (u / linear_beta).tanh())
336 }
337 _ => self.apply(g) * u,
338 }
339 }
340}
341
342pub struct DenseFfn {
344 pub gate_proj: QTensor,
345 pub up_proj: QTensor,
346 pub down_proj: QTensor,
347 pub act: Act,
349}
350
351pub enum FfnKind {
354 Dense(DenseFfn),
355 Moe(MoeFfn),
359 DenseMoe(Box<DenseMoeFfn>),
366}
367
368pub struct DenseMoeFfn {
370 pub dense: DenseFfn,
371 pub moe: MoeFfn,
372 pub post_norm_1: Vec<f32>,
374 pub pre_norm_2: Vec<f32>,
377 pub post_norm_2: Vec<f32>,
379}
380
381pub struct MoeFfn {
382 pub router: QTensor,
384 pub experts: Vec<DenseFfn>,
385 pub top_k: usize,
386 pub norm_topk_prob: bool,
387 pub router_sigmoid: bool,
390 pub expert_bias: Option<Vec<f32>>,
394 pub routed_scaling: f32,
397 pub route_tau: Option<f32>,
403 pub shared: Option<(DenseFfn, Option<QTensor>)>,
406 pub stats: std::cell::RefCell<Vec<u64>>,
410 pub act_sq: std::cell::RefCell<Vec<f64>>,
417 pub act_rows: std::cell::RefCell<Vec<f32>>,
423 pub mask: Option<Vec<bool>>,
428 pub per_expert_scale: Option<Vec<f32>>,
431 pub router_input_norm: bool,
435}
436
437pub enum AttnKind {
440 Full {
442 wq: QTensor,
443 wk: QTensor,
444 wv: QTensor,
445 wo: QTensor,
446 q_norm: Option<Vec<f32>>,
447 k_norm: Option<Vec<f32>>,
448 output_gate: bool,
449 softplus_gate: Option<(QTensor, bool)>,
453 bias: Option<(Vec<f32>, Vec<f32>, Vec<f32>)>,
455 },
456 Linear(VmfPhaseWeights),
458 LinearGdn(GdnWeights),
460 ShortConv(ShortConvWeights),
463 Mla(Box<MlaWeights>),
471 Kda(Box<crate::linear_core::KdaWeights>),
475}
476
477pub struct MlaWeights {
479 pub q_proj: QTensor,
483 pub q_a: Option<QTensor>,
486 pub q_a_norm: Option<Vec<f32>>,
487 pub kv_a: QTensor,
489 pub kv_a_norm: Vec<f32>,
491 pub kv_b: QTensor,
493 pub o_proj: QTensor,
495 pub nh: usize,
496 pub qk_rope: usize,
497 pub qk_nope: usize,
498 pub v_dim: usize,
499 pub lora: usize,
500 pub scale: f32,
502 pub nope: bool,
504}
505
506pub struct MtpModule {
511 pub enorm: Vec<f32>,
512 pub hnorm: Vec<f32>,
513 pub eh_proj: QTensor,
515 pub layer: LayerWeights,
516 pub final_norm: Vec<f32>,
517 pub kv: crate::kv_cache::LayerKvCache,
518}
519
520pub struct GenerateResult {
522 pub text: String,
523 pub token_ids: Vec<u32>,
524 pub prompt_tokens: usize,
525 pub tokens_generated: usize,
526 pub finish_reason: String,
527 pub mtp_drafted: usize,
529 pub mtp_accepted: usize,
530 pub token_confidence: Vec<f32>,
535 pub traces: Vec<TokenTrace>,
538}
539
540#[derive(Clone, Debug)]
545pub struct TokenTrace {
546 pub t: usize,
548 pub token_id: u32,
550 pub confidence: f32,
552 pub active_skill: Option<String>,
554 pub recon: Option<f32>,
558 pub switched: bool,
561}
562
563fn top1_prob_t(logits: &[f32], id: u32, temp: f32) -> f32 {
568 let t = if temp > 1e-3 { temp } else { 1.0 };
569 let max = logits.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
570 let sum: f32 = logits.iter().map(|&v| ((v - max) / t).exp()).sum();
571 if sum > 0.0 {
572 (((logits[id as usize] - max) / t).exp()) / sum
573 } else {
574 0.0
575 }
576}
577
578fn prefill_batched() -> bool {
581 std::env::var("CMF_PREFILL")
582 .map(|v| v != "seq")
583 .unwrap_or(true)
584}
585
586#[derive(Clone, Copy)]
590enum PrefillIn<'a> {
591 Ids(&'a [u32]),
592 Hidden(&'a [f32]),
593}
594
595impl Pipeline {
602 fn can_prefill_batched(&self) -> bool {
603 prefill_batched() && !self.weights.layers.is_empty()
604 }
605}
606
607pub fn prefill_chunk() -> usize {
614 if let Some(n) = std::env::var("CMF_PREFILL_CHUNK")
615 .ok()
616 .and_then(|v| v.parse::<usize>().ok())
617 {
618 return n.max(1);
619 }
620 if cfg!(target_os = "macos") {
621 512
622 } else if cfg!(target_arch = "aarch64") {
623 256
626 } else {
627 48
628 }
629}
630
631pub type TokenCallback = Box<dyn FnMut(&str) -> bool + Send>;
633
634impl Pipeline {
635 #[inline]
639 pub fn phys_layer(&self, virtual_idx: usize) -> usize {
640 virtual_idx % self.physical_layers
641 }
642
643 #[inline]
646 pub fn is_loop_end(&self, virtual_idx: usize) -> bool {
647 self.loop_final_norm && (virtual_idx + 1) % self.physical_layers == 0
648 }
649
650 #[allow(clippy::too_many_arguments)]
652
653 #[cfg(target_os = "macos")]
672 fn graph_prefill_preferred(&self) -> bool {
673 if !crate::gpu::enabled_here()
674 || !crate::gpu::q1_force()
675 || std::env::var("CMF_GPU_BLOCK")
676 .map(|v| v == "0")
677 .unwrap_or(false)
678 {
679 return false;
680 }
681 self.weights
682 .layers
683 .iter()
684 .any(|lw| matches!(&lw.attn, AttnKind::LinearGdn(w) if w.in_proj_qkv.is_q1()))
685 }
686
687 #[cfg(not(target_os = "macos"))]
688 fn graph_prefill_preferred(&self) -> bool {
689 let graph_on = std::env::var("CMF_GPU_WGPU_GRAPH")
697 .map(|v| v != "0")
698 .unwrap_or_else(|_| {
699 crate::gpu::wgpu_graph_default()
703 });
704 if !graph_on || !crate::gpu::enabled_here() {
705 return false;
706 }
707 self.weights
708 .layers
709 .iter()
710 .any(|lw| matches!(&lw.attn, AttnKind::LinearGdn(_)))
711 }
712
713 #[cfg(target_os = "macos")]
714 fn q1_graph_gpu(
715 &mut self,
716 start: usize,
717 upto: Option<usize>,
718 position: usize,
719 h: &mut [f32],
720 ) -> usize {
721 use crate::gpu::{AttnGpuLayer, GdnGpuCfg, GdnGpuLayer, GraphDims, MetalFfn, TokenGraph};
722 if self.attn_softcap > 0.0 || !crate::gpu::enabled_here()
724 || !crate::gpu::q1_force()
725 || std::env::var("CMF_GPU_BLOCK")
726 .map(|v| v == "0")
727 .unwrap_or(false)
728 {
729 if std::env::var("CMF_GRAPH_DBG").is_ok() {
730 eprintln!(
731 "block-graph: front gate (softcap={} enabled_here={} q1_force={})",
732 self.attn_softcap > 0.0,
733 crate::gpu::enabled_here(),
734 crate::gpu::q1_force(),
735 );
736 }
737 return start;
738 }
739 if self.swa.is_some()
744 || self.global_attn.is_some()
745 || self.attention_heads_per_layer.is_some()
746 || self.attn_v_norm
747 || (self.attn_scale - 1.0 / (self.head_dim as f32).sqrt()).abs() > 1e-9
748 || self.weights.layers.iter().any(|lw| {
749 lw.attn_out_norm.is_some()
750 || lw.ffn_out_norm.is_some()
751 || lw.layer_scale.is_some()
752 || matches!(&lw.ffn, FfnKind::Dense(d) if d.act != Act::Silu)
753 })
754 {
755 if std::env::var("CMF_GRAPH_DBG").is_ok() {
756 eprintln!(
757 "block-graph: arch ineligible (swa={} gattn={} hpl={} vnorm={} scale_delta={:.2e})",
758 self.swa.is_some(),
759 self.global_attn.is_some(),
760 self.attention_heads_per_layer.is_some(),
761 self.attn_v_norm,
762 (self.attn_scale - 1.0 / (self.head_dim as f32).sqrt()).abs(),
763 );
764 }
765 return start;
766 }
767 let limit = upto
770 .map(|u| u + 1)
771 .unwrap_or(self.num_layers)
772 .min(self.num_layers);
773
774 enum Item<'a> {
775 Gdn {
776 run: Vec<GdnGpuLayer<'a>>,
777 first: usize,
778 },
779 Attn {
780 l: AttnGpuLayer<'a>,
781 li: usize,
782 q_norm: Option<&'a [f32]>,
783 k_norm: Option<&'a [f32]>,
784 output_gate: bool,
785 bias: Option<(&'a [f32], &'a [f32], &'a [f32])>,
786 full_gpu: bool,
789 },
790 }
791
792 let attend_mode = std::env::var("CMF_GPU_ATTEND").unwrap_or_else(|_| "auto".into());
799 let attend_contract = attend_mode != "0"
800 && attend_mode != "off"
801 && self.head_dim % 4 == 0
802 && self.head_dim <= 256
803 && self.rotary_dim >= 2
804 && self.rotary_dim <= self.head_dim
805 && (self.rotary_dim / 2) % 32 == 0
806 && self.num_kv_heads > 0
807 && self.num_heads % self.num_kv_heads == 0;
808
809 let mut plan: Vec<Item> = Vec::new();
810 let mut model_ref: Option<std::sync::Arc<cortiq_core::CmfModel>> = None;
811 let block_diag = std::env::var("CMF_GRAPH_DBG").is_ok();
813 let mut scan = start;
814 while scan < limit {
815 let lw = &self.weights.layers[self.phys_layer(scan)];
816 let ffn = match &lw.ffn {
817 FfnKind::Dense(d) => {
818 let (Some(g), Some(u), Some(dn)) = (
819 d.gate_proj.q1_parts(),
820 d.up_proj.q1_parts(),
821 d.down_proj.q1_parts(),
822 ) else {
823 if block_diag {
824 eprintln!(
825 "block-graph: L{scan} FFN trio not graph-mappable — run ends"
826 );
827 }
828 break;
829 };
830 MetalFfn::Dense {
831 gate: g,
832 up: u,
833 down: dn,
834 }
835 }
836 FfnKind::Moe(m) => {
837 let Some(moe) = metal_moe_graph_parts(m, self.hidden_size) else {
838 if block_diag {
839 eprintln!(
840 "block-graph: L{scan} MoE outside the graph contract — run ends"
841 );
842 }
843 break;
844 };
845 if let QTensor::Mapped { model, .. } = &m.experts[0].gate_proj {
846 model_ref.get_or_insert_with(|| model.clone());
847 }
848 MetalFfn::Moe(moe)
849 }
850 _ => {
851 if block_diag {
852 eprintln!("block-graph: L{scan} non-graph FFN — run ends");
853 }
854 break;
855 }
856 };
857 match &lw.attn {
858 AttnKind::LinearGdn(w) if self.gdn_cfg.is_some() => {
859 let parts = (
860 w.in_proj_qkv.q1_parts(),
861 w.in_proj_z.q1_parts(),
862 w.in_proj_a.f32_parts(),
863 w.in_proj_b.f32_parts(),
864 w.out_proj.q1_parts(),
865 );
866 let (Some(qkv), Some(z), Some(a), Some(b), Some(out)) = parts else {
867 if block_diag {
868 eprintln!(
869 "block-graph: L{scan} GDN parts refused (qkv={} z={} a_f32={} b_f32={} out={})",
870 w.in_proj_qkv.q1_parts().is_some(),
871 w.in_proj_z.q1_parts().is_some(),
872 w.in_proj_a.f32_parts().is_some(),
873 w.in_proj_b.f32_parts().is_some(),
874 w.out_proj.q1_parts().is_some(),
875 );
876 }
877 break;
878 };
879 if let QTensor::Mapped { model, .. } = &w.in_proj_qkv {
880 model_ref.get_or_insert_with(|| model.clone());
881 }
882 let gl = GdnGpuLayer {
883 attn_norm: &lw.input_norm,
884 post_norm: &lw.post_norm,
885 qkv,
886 z,
887 a,
888 b,
889 out,
890 ffn,
891 conv1d: &w.conv1d,
892 a_log: &w.a_log,
893 dt_bias: &w.dt_bias,
894 gnorm: &w.norm,
895 };
896 match plan.last_mut() {
897 Some(Item::Gdn { run, .. }) => run.push(gl),
898 _ => plan.push(Item::Gdn {
899 run: vec![gl],
900 first: scan,
901 }),
902 }
903 }
904 AttnKind::Full {
905 wq,
906 wk,
907 wv,
908 wo,
909 q_norm,
910 k_norm,
911 output_gate,
912 softplus_gate: None,
913 bias,
914 } if !self.kv_cache.layers[scan].o1_sealed() => {
915 let parts = (wq.q1_parts(), wk.q1_parts(), wv.q1_parts(), wo.q1_parts());
916 let (Some(pq), Some(pk), Some(pv), Some(po)) = parts else {
917 break;
918 };
919 if let QTensor::Mapped { model, .. } = wq {
920 model_ref.get_or_insert_with(|| model.clone());
921 }
922 let cache = &self.kv_cache.layers[scan];
923 let full_gpu = attend_contract
924 && cache.mode == crate::kv_cache::KvMode::F32
925 && cache.o1.is_none()
926 && bias.is_none()
927 && pq.1 == self.num_heads * self.head_dim * (1 + *output_gate as usize)
928 && pk.1 == self.num_kv_heads * self.head_dim
929 && pv.1 == self.num_kv_heads * self.head_dim
930 && po.2 == self.num_heads * self.head_dim;
931 plan.push(Item::Attn {
932 l: AttnGpuLayer {
933 attn_norm: &lw.input_norm,
934 post_norm: &lw.post_norm,
935 wq: pq,
936 wk: pk,
937 wv: pv,
938 wo: po,
939 ffn,
940 },
941 li: scan,
942 q_norm: q_norm.as_deref(),
943 k_norm: k_norm.as_deref(),
944 output_gate: *output_gate,
945 bias: bias
946 .as_ref()
947 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
948 full_gpu,
949 });
950 }
951 _ => break,
952 }
953 scan += 1;
954 }
955 let Some(model) = model_ref else {
956 if std::env::var("CMF_GRAPH_DBG").is_ok() {
957 eprintln!("q1-graph: no model ref (start {start}, scanned to {scan})");
958 }
959 return start;
960 };
961 if plan.is_empty() {
962 if std::env::var("CMF_GRAPH_DBG").is_ok() {
963 eprintln!("q1-graph: empty plan at layer {start}");
964 }
965 return start;
966 }
967 let has_moe = plan.iter().any(|it| match it {
968 Item::Gdn { run, .. } => run.iter().any(|l| matches!(l.ffn, MetalFfn::Moe(_))),
969 Item::Attn { l, .. } => matches!(l.ffn, MetalFfn::Moe(_)),
970 });
971 let dev_attend = attend_contract
972 && (self.head_dim <= 128
973 || has_moe
974 || attend_mode == "force"
975 || attend_mode == "256");
976 if !dev_attend {
977 for it in &mut plan {
978 if let Item::Attn { full_gpu, .. } = it {
979 *full_gpu = false;
980 }
981 }
982 }
983 if std::env::var("CMF_GRAPH_DBG").is_ok() {
984 use std::sync::atomic::{AtomicBool, Ordering};
985 static SAID: AtomicBool = AtomicBool::new(false);
986 if !SAID.swap(true, Ordering::Relaxed) {
987 let fg = plan
988 .iter()
989 .filter(|it| matches!(it, Item::Attn { full_gpu: true, .. }))
990 .count();
991 let att = plan
992 .iter()
993 .filter(|it| matches!(it, Item::Attn { .. }))
994 .count();
995 eprintln!(
996 "q1-graph: plan of {} items from layer {start} to {scan} | dev_attend={dev_attend} full_gpu {fg}/{att} | hd={} rd={} nkv={} nh={}",
997 plan.len(),
998 self.head_dim,
999 self.rotary_dim,
1000 self.num_kv_heads,
1001 self.num_heads,
1002 );
1003 }
1004 }
1005 let dims = GraphDims {
1006 hidden: self.hidden_size,
1007 eps: self.rms_eps as f32,
1008 gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
1009 };
1010 let Some(mut graph) = TokenGraph::new(&model, dims, h) else {
1011 return start;
1012 };
1013 let gcfg = self.gdn_cfg.map(|cfg| GdnGpuCfg {
1014 nv: cfg.num_v_heads,
1015 nk: cfg.num_k_heads,
1016 dk: cfg.key_head_dim,
1017 dv: cfg.value_head_dim,
1018 kk: cfg.conv_kernel,
1019 hidden: self.hidden_size,
1020 inter: self.intermediate_size,
1021 c_dim: cfg.conv_dim(),
1022 eps: cfg.rms_eps as f32,
1023 gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
1024 });
1025 let mut valid = 0usize;
1029 let mut end = start;
1030 for item in &plan {
1031 let ok = match item {
1032 Item::Gdn { run, .. } => gcfg
1033 .as_ref()
1034 .map(|gc| run.iter().all(|l| graph.gdn_ok(l, gc)))
1035 .unwrap_or(false),
1036 Item::Attn { l, .. } => graph.attn_ok(l),
1037 };
1038 if !ok {
1039 if block_diag {
1040 eprintln!(
1041 "block-graph: plan item {} ({}) failed graph preflight",
1042 valid,
1043 match item {
1044 Item::Gdn { run, first } =>
1045 format!("GDN run L{first}+{}", run.len()),
1046 Item::Attn { li, .. } => format!("Attn L{li}"),
1047 }
1048 );
1049 }
1050 break;
1051 }
1052 valid += 1;
1053 end += match item {
1054 Item::Gdn { run, .. } => run.len(),
1055 Item::Attn { .. } => 1,
1056 };
1057 }
1058 plan.truncate(valid);
1059 if plan.is_empty() {
1060 return start;
1061 }
1062
1063 let inv_freq = self.inv_freq.clone();
1064 let pool = self.pool.clone();
1065 let (nh, nkv, hd, hs, rd, eps) = (
1066 self.num_heads,
1067 self.num_kv_heads,
1068 self.head_dim,
1069 self.hidden_size,
1070 self.rotary_dim,
1071 self.rms_eps,
1072 );
1073 let norm_style = self.norm_style;
1074 let gemma = norm_style == cortiq_core::NormStyle::Gemma;
1075 let want = self.gdn_cfg.map(|c| c.state_len()).unwrap_or(0);
1076 let kv_id = self.graph_kv_id;
1077 let mut pending: Vec<(usize, usize)> = Vec::new();
1080 let mut dev_attn: Vec<usize> = Vec::new();
1083 for item in &plan {
1084 if self.loop_final_norm {
1086 let item_start = match item {
1087 Item::Gdn { first, .. } => *first,
1088 Item::Attn { li, .. } => *li,
1089 };
1090 if item_start > start && self.is_loop_end(item_start - 1) {
1091 graph.encode_loop_norm(&self.weights.final_norm);
1092 }
1093 }
1094 match item {
1095 Item::Gdn { run, first } => {
1096 for l in &mut self.kv_cache.layers[*first..*first + run.len()] {
1097 if l.linear_state.len() != want {
1098 l.linear_state = vec![0f32; want];
1099 }
1100 }
1101 let ro: Vec<&[f32]> = self.kv_cache.layers[*first..*first + run.len()]
1102 .iter()
1103 .map(|l| l.linear_state.as_slice())
1104 .collect();
1105 if !graph.encode_gdn_run(run, &ro, gcfg.as_ref().unwrap()) {
1106 tracing::error!("q1 graph: GDN run refused after validation");
1108 return start;
1109 }
1110 graph.commit();
1113 pending.push((*first, run.len()));
1114 }
1115 Item::Attn {
1116 l,
1117 li,
1118 q_norm,
1119 k_norm,
1120 output_gate,
1121 bias,
1122 full_gpu,
1123 } => {
1124 if *full_gpu {
1126 let cache = &self.kv_cache.layers[*li];
1127 let cpu_k: Vec<&[f32]> = (0..nkv).map(|g| cache.head_keys(g)).collect();
1128 let cpu_v: Vec<&[f32]> = (0..nkv).map(|g| cache.head_values(g)).collect();
1129 let cpu_stored = cpu_k[0].len() / hd;
1130 let p = crate::gpu::AttnDeviceParams {
1131 kv_id,
1132 layer: *li,
1133 nh,
1134 nkv,
1135 hd,
1136 rd,
1137 position,
1138 eps: eps as f32,
1139 gemma,
1140 output_gate: *output_gate,
1141 q_norm: *q_norm,
1142 k_norm: *k_norm,
1143 inv_freq: &inv_freq,
1144 cpu_k,
1145 cpu_v,
1146 cpu_stored,
1147 };
1148 if graph.attn_device_ok(l, &p) && graph.encode_attn_device(l, &p) {
1149 graph.commit();
1150 dev_attn.push(*li);
1151 continue;
1152 }
1153 }
1155 graph.encode_attn_prefix(l);
1156 graph.sync();
1157 if !pending.is_empty() {
1158 let idxs: Vec<usize> =
1159 pending.drain(..).flat_map(|(f, n)| f..f + n).collect();
1160 let mut outs: Vec<&mut [f32]> = self
1161 .kv_cache
1162 .layers
1163 .iter_mut()
1164 .enumerate()
1165 .filter(|(i, _)| idxs.binary_search(i).is_ok())
1166 .map(|(_, s)| s.linear_state.as_mut_slice())
1167 .collect();
1168 graph.read_states(&mut outs);
1169 }
1170 let mut q_raw = attention::take_buf(l.wq.1);
1171 let mut k = attention::take_buf(l.wk.1);
1172 let mut v = attention::take_buf(l.wv.1);
1173 graph.read_qkv(&mut q_raw, &mut k, &mut v);
1174 let cfg = QwenAttnCfg {
1175 num_heads: nh,
1176 num_kv_heads: nkv,
1177 head_dim: hd,
1178 hidden_size: hs,
1179 position,
1180 inv_freq: &inv_freq,
1181 rotary_dim: rd,
1182 scale: self.attn_scale,
1183 softcap: self.attn_softcap,
1184 window: None,
1185 v_norm: false,
1186 q_norm: *q_norm,
1187 k_norm: *k_norm,
1188 output_gate: *output_gate,
1189 softplus_gate: None,
1190 rope_scale: 1.0,
1191 bias: *bias,
1192 rms_eps: eps,
1193 norm_style,
1194 pool: pool.as_deref(),
1195 };
1196 let mut ao = attention::qwen_attention_core(
1197 q_raw,
1198 k,
1199 v,
1200 &mut self.kv_cache.layers[*li],
1201 &cfg,
1202 );
1203 graph.encode_attn_suffix(l, &ao);
1204 graph.commit();
1207 attention::recycle_buf(&mut ao);
1208 }
1209 }
1210 }
1211 let mut lm_rows = None;
1216 if self.graph_want_logits
1217 && upto.is_none()
1218 && end == self.num_layers
1219 && std::env::var("CMF_GPU_LMHEAD")
1220 .map(|v| v != "0")
1221 .unwrap_or(true)
1222 {
1223 if let Some(lm) = self.weights.lm_head.q1_parts() {
1224 if graph.lm_head_ok(lm) {
1225 graph.encode_lm_head(&self.weights.final_norm, lm);
1226 lm_rows = Some(lm.1);
1227 }
1228 }
1229 }
1230 graph.sync();
1231 if !pending.is_empty() {
1232 let idxs: Vec<usize> = pending.drain(..).flat_map(|(f, n)| f..f + n).collect();
1233 let mut outs: Vec<&mut [f32]> = self
1234 .kv_cache
1235 .layers
1236 .iter_mut()
1237 .enumerate()
1238 .filter(|(i, _)| idxs.binary_search(i).is_ok())
1239 .map(|(_, s)| s.linear_state.as_mut_slice())
1240 .collect();
1241 graph.read_states(&mut outs);
1242 }
1243 if let Some(rows) = lm_rows {
1244 let mut lg = attention::take_buf(rows.min(self.vocab_size));
1245 graph.read_logits(&mut lg);
1246 lg.resize(self.vocab_size, 0.0);
1247 if let Some(c) = self.final_softcap {
1248 for l in lg.iter_mut() {
1249 *l = c * (*l / c).tanh();
1250 }
1251 }
1252 self.graph_logits = Some(lg);
1253 }
1254 graph.finish(h);
1255 for li in dev_attn {
1259 let mut krow = attention::take_buf(nkv * hd);
1260 let mut vrow = attention::take_buf(nkv * hd);
1261 if crate::gpu::kv_mirror_read_last(kv_id, li, nkv, hd, &mut krow, &mut vrow) {
1262 let cache = &mut self.kv_cache.layers[li];
1263 cache.append(&krow, &vrow, &[]);
1264 let n = cache.seq_len;
1265 let mut imp = attention::take_buf(n);
1266 crate::gpu::kv_mirror_take_imp(kv_id, li, &mut imp);
1267 cache.accumulate_imp(&imp);
1268 attention::recycle_buf(&mut imp);
1269 }
1270 attention::recycle_buf(&mut krow);
1271 attention::recycle_buf(&mut vrow);
1272 }
1273 end
1274 }
1275
1276 pub fn new(
1277 tokenizer: Tokenizer,
1278 weights: PipelineWeights,
1279 hidden_size: usize,
1280 intermediate_size: usize,
1281 num_heads: usize,
1282 num_kv_heads: usize,
1283 head_dim: usize,
1284 num_layers: usize,
1285 physical_layers: usize,
1286 loop_final_norm: bool,
1287 vocab_size: usize,
1288 rms_eps: f64,
1289 rope_base: f32,
1290 norm_style: NormStyle,
1291 max_seq_len: usize,
1292 sampler_config: SamplerConfig,
1293 ) -> Self {
1294 let rng = match sampler_config.seed {
1295 Some(s) => SplitMix64::new(s),
1296 None => SplitMix64::from_entropy(),
1297 };
1298 let inv_freq = std::sync::Arc::new(attention::rope_inv_freq(head_dim, rope_base));
1299 let pool = Pool::from_env();
1300 if let Some(p) = &pool {
1301 tracing::info!("worker pool: {} threads", p.n_workers());
1302 }
1303 Self {
1304 gpu_plan: None,
1305 tokenizer: std::sync::Arc::new(tokenizer),
1306 kv_cache: KvCache::new(num_layers, num_kv_heads, head_dim, max_seq_len),
1307 sampler_config,
1308 weights,
1309 hidden_size,
1310 intermediate_size,
1311 num_heads,
1312 num_kv_heads,
1313 head_dim,
1314 num_layers,
1315 physical_layers,
1316 loop_final_norm,
1317 vocab_size,
1318 rms_eps,
1319 rope_base,
1320 norm_style,
1321 rotary_dim: head_dim,
1322 attention_heads_per_layer: None,
1323 vmf_cfg: None,
1324 gdn_cfg: None,
1325 kda_cfg: None,
1326 g3n: None,
1327 dsv4: None,
1328 dsv4_mtp: Vec::new(),
1329 dspark: None,
1330 dspark_pending: Vec::new(),
1331 dspark_hist: Vec::new(),
1332 dspark_real: Vec::new(),
1333 dspark_trunk_picks: Vec::new(),
1334 dspark_exp: Vec::new(),
1335 dspark_draft_ns: 0,
1336 logit_multiplier: None,
1337 cancel: std::sync::Arc::new(std::sync::atomic::AtomicBool::new(false)),
1338 kv_history: Vec::new(),
1339 short_conv_cfg: None,
1340 mtp: None,
1341 speculative: std::env::var("CMF_MTP").map(|v| v != "0").unwrap_or(true),
1342 rng,
1343 sampler_scratch: SamplerScratch::default(),
1344 inv_freq,
1345 ws: ForwardScratch::new(hidden_size),
1346 pool,
1347 model: None,
1348 dyn_force_f32: false,
1349 dyn_skill_layers: Vec::new(),
1350 dyn_active: None,
1351 dyn_blend_loaded: false,
1352 dyn_phi_layer: None,
1353 dyn_phi_ema: Vec::new(),
1354 dyn_phi_seen: 0,
1355 dyn_router: None,
1356 o1_cfg: None,
1357 o1_epoch: 0,
1358 o1_flags: Vec::new(),
1359 trace: false,
1360 calib_temp: 1.0,
1361 confidence_on: true,
1362 embed_multiplier: 1.0,
1363 attn_scale: 1.0 / (head_dim as f32).sqrt(),
1364 swa: None,
1365 sliding_layers: None,
1366 inv_freq_local: None,
1367 rotary_dim_local: None,
1368 rope_scale: 1.0,
1369 rope_scale_local: 1.0,
1370 global_attn: None,
1371 inv_freq_global: None,
1372 attn_v_norm: false,
1373 final_softcap: None,
1374 attn_softcap: 0.0,
1375 graph_want_logits: false,
1376 graph_logits: None,
1377 graph_kv_id: {
1378 static NEXT: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(1);
1379 NEXT.fetch_add(1, std::sync::atomic::Ordering::Relaxed)
1380 },
1381 }
1382 }
1383
1384 pub fn set_o1(&mut self, cfg: Option<crate::nystrom::O1Cfg>) {
1391 self.o1_flags = match &cfg {
1392 Some(c) => {
1393 let mut flags = c.layer_flags(self.num_layers);
1394 for (li, f) in flags.iter_mut().enumerate() {
1395 if *f
1396 && !matches!(
1397 self.weights.layers[self.phys_layer(li)].attn,
1398 AttnKind::Full { .. }
1399 )
1400 {
1401 *f = false;
1402 }
1403 }
1404 flags
1405 }
1406 None => Vec::new(),
1407 };
1408 if let Some(c) = &cfg {
1409 let n = self.o1_flags.iter().filter(|&&f| f).count();
1410 tracing::info!(
1411 "o1 nystrom attention: {n}/{} layer(s), m={} w={} sink={} rect={:?}",
1412 self.num_layers,
1413 c.m,
1414 c.w,
1415 c.sink,
1416 c.rect
1417 );
1418 }
1419 self.o1_cfg = cfg;
1420 }
1421
1422 pub fn o1_active(&self) -> bool {
1424 self.o1_cfg.is_some() && self.o1_flags.iter().any(|&f| f)
1425 }
1426
1427 pub fn o1_begin(&mut self) {
1432 if let Some(c) = &self.o1_cfg {
1433 let (m, w, sink, rect) = (c.m, c.w, c.sink, c.rect);
1434 for (li, &f) in self.o1_flags.iter().enumerate() {
1435 if f {
1436 self.kv_cache.layers[li].o1_begin(m, w, sink, rect);
1437 }
1438 }
1439 }
1440 }
1441
1442 pub fn o1_seal(&mut self) {
1446 self.o1_epoch = self.o1_epoch.wrapping_add(1);
1447 if self.o1_cfg.is_none() {
1448 return;
1449 }
1450 for li in 0..self.num_layers {
1451 if self.o1_flags.get(li).copied().unwrap_or(false) {
1452 self.kv_cache.layers[li].o1_seal(self.num_heads);
1453 }
1454 }
1455 }
1456
1457 pub fn set_trace(&mut self, on: bool) {
1459 self.trace = on;
1460 }
1461
1462 pub fn set_sampler_config(&mut self, config: SamplerConfig) {
1465 self.rng = match config.seed {
1466 Some(seed) => SplitMix64::new(seed),
1467 None => SplitMix64::from_entropy(),
1468 };
1469 self.sampler_config = config;
1470 }
1471
1472 pub fn set_confidence(&mut self, on: bool) {
1477 self.confidence_on = on;
1478 }
1479
1480 pub fn set_calib_temp(&mut self, t: f32) {
1483 self.calib_temp = if t > 1e-3 { t } else { 1.0 };
1484 }
1485
1486 pub fn calib_temp(&self) -> f32 {
1488 self.calib_temp
1489 }
1490
1491 pub fn set_rotary(&mut self, rotary_dim: usize, base: f32) {
1494 self.rotary_dim = rotary_dim.min(self.head_dim);
1495 self.inv_freq = std::sync::Arc::new(attention::rope_inv_freq(self.rotary_dim, base));
1496 }
1497
1498 fn attn_cfg(&self, position: usize) -> QwenAttnCfg<'_> {
1499 QwenAttnCfg {
1500 num_heads: self.num_heads,
1501 num_kv_heads: self.num_kv_heads,
1502 head_dim: self.head_dim,
1503 hidden_size: self.hidden_size,
1504 position,
1505 inv_freq: &self.inv_freq,
1506 rotary_dim: self.rotary_dim,
1507 scale: self.attn_scale,
1508 softcap: self.attn_softcap,
1509 window: None,
1510 v_norm: false,
1511 q_norm: None,
1512 k_norm: None,
1513 output_gate: false,
1514 softplus_gate: None,
1515 rope_scale: self.rope_scale,
1516 bias: None,
1517 rms_eps: self.rms_eps,
1518 norm_style: self.norm_style,
1519 pool: self.pool.as_deref(),
1520 }
1521 }
1522
1523 pub fn generate(
1525 &mut self,
1526 prompt: &str,
1527 max_tokens: usize,
1528 task_mask: Option<&TaskMask>,
1529 on_token: Option<TokenCallback>,
1530 ) -> Result<GenerateResult, String> {
1531 let input_ids = self.tokenizer.with_bos(self.tokenizer.encode(prompt));
1532 self.generate_from_ids(&input_ids, max_tokens, task_mask, on_token)
1533 }
1534
1535 pub fn generate_from_ids(
1543 &mut self,
1544 input_ids: &[u32],
1545 max_tokens: usize,
1546 task_mask: Option<&TaskMask>,
1547 mut on_token: Option<TokenCallback>,
1548 ) -> Result<GenerateResult, String> {
1549 if std::env::var("CMF_TRACE_H").is_ok() {
1550 eprintln!("input_ids: {input_ids:?}");
1551 }
1552 if input_ids.is_empty() {
1553 return Err("empty prompt: nothing to generate from".to_string());
1554 }
1555
1556 let reuse_from = {
1564 let on = !std::env::var("CMF_KV_REUSE").is_ok_and(|v| v == "0");
1565 let h = &self.kv_history;
1566 if on
1567 && task_mask.is_none()
1568 && self.mtp.is_none()
1569 && self.o1_cfg.is_none()
1570 && !h.is_empty()
1571 && h.len() < input_ids.len()
1572 && input_ids[..h.len()] == h[..]
1573 {
1574 h.len()
1575 } else {
1576 0
1577 }
1578 };
1579 if reuse_from == 0 {
1580 self.kv_cache.clear();
1582 self.kv_history.clear();
1583 crate::gpu::graph_kv_reset(self.graph_kv_id);
1584 } else if std::env::var("CMF_PREFILL_PROF").is_ok() {
1585 eprintln!(
1586 "kv-reuse: {} of {} prompt positions already cached",
1587 reuse_from,
1588 input_ids.len()
1589 );
1590 }
1591 crate::gpu::graph_race_begin_generation();
1592 self.o1_begin();
1593
1594 let graph_on = std::env::var("CMF_GPU_WGPU_GRAPH")
1600 .map(|v| v != "0")
1601 .unwrap_or_else(|_| {
1602 crate::gpu::wgpu_graph_default()
1606 });
1607 let graph_spec = self.speculative
1623 && graph_on
1624 && self.mtp.is_some()
1625 && task_mask.is_none()
1626 && !self.o1_active()
1627 && self.sampler_config.temperature < 1e-6
1628 && self.sampler_config.repetition_penalty == 1.0
1629 && std::env::var("CMF_GRAPH_SPEC").is_ok_and(|v| v != "0");
1630 let pair_pays = self.gdn_cfg.is_none()
1637 || std::env::var("CMF_MTP").as_deref() == Ok("1");
1638 let spec_active = self.speculative
1639 && self.mtp.is_some()
1640 && task_mask.is_none()
1641 && !self.o1_active()
1642 && ((!graph_on && pair_pays) || graph_spec)
1643 && self.sampler_config.temperature < 1e-6;
1644 let mut mtp = if spec_active { self.mtp.take() } else { None };
1647 if std::env::var("CMF_MTP_CHAIN_PROBE").is_ok() {
1648 eprintln!(
1649 "mtp-probe gate: spec_active={spec_active} mtp={} speculative={} graph_on={graph_on} temp_ok={}",
1650 mtp.is_some(),
1651 self.speculative,
1652 self.sampler_config.temperature < 1e-6,
1653 );
1654 }
1655 if let Some(m) = &mut mtp {
1656 m.kv.clear();
1657 }
1658 let mut router = if mtp.is_none() {
1662 self.dyn_router.take()
1663 } else {
1664 None
1665 };
1666 if let Some(r) = &mut router {
1667 r.reset(); self.dyn_phi_seen = 0; let _ = self.set_active_skill(None);
1670 }
1671
1672 let mut all_ids = input_ids.to_vec();
1673 let mut generated = 0usize;
1674 let mut finish_reason = "max_tokens".to_string();
1675 let mut drafted = 0usize;
1676 let mut accepted = 0usize;
1677 let mut confidence: Vec<f32> = Vec::new();
1678 let trace_on = self.trace;
1679 let calib_temp = self.calib_temp;
1680 let mut traces: Vec<TokenTrace> = Vec::new();
1681
1682 let mut hidden = vec![0.0f32; self.hidden_size];
1688 let mut pos = reuse_from;
1689 let fuse_lm = mtp.is_none()
1698 && router.is_none()
1699 && std::env::var("CMF_GPU_LMHEAD").as_deref() != Ok("0");
1700 self.graph_logits = None;
1701 self.graph_want_logits = false;
1702 let _tpf = std::time::Instant::now();
1703 let batch_k = std::env::var("CMF_BATCH_K")
1704 .ok()
1705 .and_then(|v| v.parse::<usize>().ok())
1706 .unwrap_or(0);
1707 while self.dsv4.is_some()
1718 && mtp.is_none()
1719 && pos < input_ids.len()
1720 && !self.cancel.load(std::sync::atomic::Ordering::Relaxed)
1721 {
1722 let end = (pos + prefill_chunk()).min(input_ids.len());
1723 let ids: Vec<u32> = input_ids[pos..end].to_vec();
1724 let mut lg = Vec::new();
1725 if let Some(b) = &mut self.dsv4 {
1726 let (g, layers, cfg, st) = (&b.0, &b.1, b.2, &mut b.3);
1727 crate::dsv4::forward_chunk(
1728 g,
1729 layers,
1730 &cfg,
1731 st,
1732 &ids,
1733 pos,
1734 &self.inv_freq,
1735 self.pool.as_deref(),
1736 &mut lg,
1737 end == input_ids.len(),
1738 );
1739 }
1740 if end == input_ids.len() {
1741 self.graph_logits = Some(lg);
1742 }
1743 pos = end;
1744 hidden = vec![0.0; self.hidden_size];
1745 }
1746 let dyn_prefill = router.is_some();
1751 let graph_prefill = self.graph_prefill_preferred();
1757 if task_mask.is_none()
1758 && !dyn_prefill
1759 && !graph_prefill
1760 && self.can_prefill_batched()
1761 && self.g3n.is_none()
1762 && input_ids.len() > 2
1763 {
1764 let chunk = prefill_chunk();
1770 let hs = self.hidden_size;
1771 while pos < input_ids.len() && !self.cancel.load(std::sync::atomic::Ordering::Relaxed) {
1772 let end = (pos + chunk).min(input_ids.len());
1773 let hb = self.prefill_batch(&input_ids[pos..end], pos);
1774 if let Some(m) = &mut mtp {
1775 let probe: usize = std::env::var("CMF_MTP_CHAIN_PROBE")
1776 .ok()
1777 .and_then(|v| v.parse().ok())
1778 .unwrap_or(0);
1779 for p in pos..end {
1780 if p + 1 < input_ids.len() {
1781 if probe >= 1 && p + 2 < input_ids.len() {
1782 let (d1, mut hx) = self.mtp_step_h(
1786 m,
1787 &hb[(p - pos) * hs..(p - pos + 1) * hs],
1788 input_ids[p + 1],
1789 p,
1790 );
1791 let mut ok = d1 == input_ids[p + 2];
1792 Self::chain_probe_note(0, ok);
1793 let mut d_prev = d1;
1794 let mut extra = 0usize;
1795 for j in 1..probe {
1796 if p + 2 + j >= input_ids.len() {
1797 break;
1798 }
1799 let (dj, hj) =
1800 self.mtp_step_h(m, &hx, d_prev, p + 1 + j);
1801 extra += 1;
1802 ok = ok && dj == input_ids[p + 2 + j];
1803 Self::chain_probe_note(j, ok);
1804 d_prev = dj;
1805 hx = hj;
1806 }
1807 m.kv.truncate_last(extra);
1808 } else {
1809 let _ = self.mtp_step(
1810 m,
1811 &hb[(p - pos) * hs..(p - pos + 1) * hs],
1812 input_ids[p + 1],
1813 p,
1814 );
1815 }
1816 }
1817 }
1818 }
1819 hidden.copy_from_slice(&hb[(end - pos - 1) * hs..]);
1820 pos = end;
1821 }
1822 }
1823 let pair_off = std::env::var("CMF_PAIR").is_ok_and(|v| v == "0");
1824 if task_mask.is_none()
1825 && !dyn_prefill
1826 && !graph_prefill
1827 && !pair_off
1828 && self.pair_supported()
1829 {
1830 while pos + 1 < input_ids.len()
1831 && !self.cancel.load(std::sync::atomic::Ordering::Relaxed)
1832 {
1833 let e1 = self.embed_single(input_ids[pos]);
1834 let e2 = self.embed_single(input_ids[pos + 1]);
1835 let (h1, h2) = self.forward_pair(&e1, &e2, pos);
1836 self.commit_linear_scratch();
1838 if let Some(m) = &mut mtp {
1839 let _ = self.mtp_step(m, &h1, input_ids[pos + 1], pos);
1840 if pos + 2 < input_ids.len() {
1841 let probe: usize = std::env::var("CMF_MTP_CHAIN_PROBE")
1842 .ok()
1843 .and_then(|v| v.parse().ok())
1844 .unwrap_or(0);
1845 if probe >= 1 && pos + 3 < input_ids.len() {
1846 let (d1, mut hx) =
1850 self.mtp_step_h(m, &h2, input_ids[pos + 2], pos + 1);
1851 let mut ok = d1 == input_ids[pos + 3];
1852 Self::chain_probe_note(0, ok);
1853 let mut d_prev = d1;
1854 let mut extra = 0usize;
1855 for j in 1..probe {
1856 if pos + 3 + j >= input_ids.len() {
1857 break;
1858 }
1859 let (dj, hj) =
1860 self.mtp_step_h(m, &hx, d_prev, pos + 2 + j);
1861 extra += 1;
1862 ok = ok && dj == input_ids[pos + 3 + j];
1863 Self::chain_probe_note(j, ok);
1864 d_prev = dj;
1865 hx = hj;
1866 }
1867 m.kv.truncate_last(extra);
1868 } else {
1869 let _ = self.mtp_step(m, &h2, input_ids[pos + 2], pos + 1);
1870 }
1871 }
1872 }
1873 hidden = h2;
1874 pos += 2;
1875 }
1876 }
1877 if batch_k > 0
1886 && graph_prefill
1887 && task_mask.is_none()
1888 && !self.o1_active()
1889 && mtp.is_none()
1890 && !dyn_prefill
1891 && pos + 1 < input_ids.len()
1892 {
1893 let hs = self.hidden_size;
1894 let chunk = batch_k;
1895 while pos < input_ids.len() {
1896 let end = (pos + chunk).min(input_ids.len());
1897 let bk = end - pos;
1898 let mut hiddens = vec![0f32; bk * hs];
1899 for (j, &id) in input_ids[pos..end].iter().enumerate() {
1900 hiddens[j * hs..(j + 1) * hs].copy_from_slice(&self.embed_single(id));
1901 }
1902 let positions: Vec<usize> = (pos..end).collect();
1903 let t_chunk = std::time::Instant::now();
1904 let ok_b = self.try_batch_graph_wgpu(&mut hiddens, &positions, bk, None);
1905 if std::env::var("CMF_GRAPH_PROF").is_ok() {
1906 let ms = t_chunk.elapsed().as_secs_f64() * 1000.0;
1907 eprintln!(
1908 "batch-chunk: k={bk} ok={ok_b} {ms:.1} ms ({:.1} tok/s)",
1909 bk as f64 / (ms / 1000.0)
1910 );
1911 }
1912 {
1913 use std::sync::atomic::{AtomicBool, Ordering};
1914 static SAID: AtomicBool = AtomicBool::new(false);
1915 if !SAID.swap(true, Ordering::Relaxed) {
1916 if ok_b {
1917 tracing::info!("batched prefill: ACTIVE (k={bk})");
1918 } else {
1919 tracing::warn!("batched prefill declined — per-position graph");
1920 }
1921 }
1922 }
1923 if ok_b {
1924 hidden.copy_from_slice(&hiddens[(bk - 1) * hs..]);
1925 pos = end;
1926 } else {
1927 break; }
1929 }
1930 }
1931 while pos < input_ids.len() && !self.cancel.load(std::sync::atomic::Ordering::Relaxed) {
1932 self.graph_want_logits = fuse_lm && pos + 1 == input_ids.len();
1933 hidden = self.forward_layers(&self.embed_single(input_ids[pos]), pos, task_mask);
1934 if let Some(m) = &mut mtp {
1935 if pos + 1 < input_ids.len() {
1936 let probe: usize = std::env::var("CMF_MTP_CHAIN_PROBE")
1942 .ok()
1943 .and_then(|v| v.parse().ok())
1944 .unwrap_or(0);
1945 if probe >= 1 && pos + 2 < input_ids.len() {
1946 let (d1, mut hx) = self.mtp_step_h(m, &hidden, input_ids[pos + 1], pos);
1947 let mut ok = d1 == input_ids[pos + 2];
1948 Self::chain_probe_note(0, ok);
1949 let mut d_prev = d1;
1950 let mut extra = 0usize;
1951 for j in 1..probe {
1952 if pos + 2 + j >= input_ids.len() {
1953 break;
1954 }
1955 let (dj, hj) =
1956 self.mtp_step_h(m, &hx, d_prev, pos + 1 + j);
1957 extra += 1;
1958 ok = ok && dj == input_ids[pos + 2 + j];
1959 Self::chain_probe_note(j, ok);
1960 d_prev = dj;
1961 hx = hj;
1962 }
1963 m.kv.truncate_last(extra);
1966 } else {
1967 let _ = self.mtp_step(m, &hidden, input_ids[pos + 1], pos);
1968 }
1969 }
1970 }
1971 pos += 1;
1972 }
1973 if std::env::var("CMF_PREFILL_PROF").is_ok() {
1974 eprintln!(
1975 "prefill: {} tokens in {:.1} ms (batch_k={batch_k})",
1976 input_ids.len(),
1977 _tpf.elapsed().as_secs_f64() * 1000.0
1978 );
1979 }
1980 if self
1983 .cancel
1984 .swap(false, std::sync::atomic::Ordering::Relaxed)
1985 {
1986 self.kv_history.clear();
1987 if let Some(m) = mtp {
1988 self.mtp = Some(m);
1989 }
1990 return Ok(GenerateResult {
1991 text: String::new(),
1992 token_ids: Vec::new(),
1993 prompt_tokens: input_ids.len(),
1994 tokens_generated: 0,
1995 finish_reason: "cancelled".to_string(),
1996 mtp_drafted: 0,
1997 mtp_accepted: 0,
1998 token_confidence: Vec::new(),
1999 traces: Vec::new(),
2000 });
2001 }
2002
2003 self.o1_seal();
2006
2007 macro_rules! commit {
2009 ($id:expr) => {{
2010 all_ids.push($id);
2011 generated += 1;
2012 if self.tokenizer.is_eos($id) {
2013 finish_reason = "stop".to_string();
2014 false
2015 } else {
2016 let token_text = self.tokenizer.decode_token($id);
2017 let mut go = true;
2018 if let Some(ref mut cb) = on_token {
2019 if !cb(&token_text) {
2020 finish_reason = "cancelled".to_string();
2021 go = false;
2022 }
2023 }
2024 go
2025 }
2026 }};
2027 }
2028
2029 let mut next_pos = input_ids.len();
2031 'decode: while generated < max_tokens {
2032 if self
2033 .cancel
2034 .swap(false, std::sync::atomic::Ordering::Relaxed)
2035 {
2036 finish_reason = "cancelled".to_string();
2037 break 'decode;
2038 }
2039 let mut logits = match self.graph_logits.take() {
2040 Some(lg) => lg,
2041 None => {
2042 inference::rms_norm_into(
2043 &hidden,
2044 &self.weights.final_norm,
2045 self.rms_eps,
2046 self.norm_style,
2047 &mut self.ws.n1,
2048 );
2049 self.lm_head_forward(&self.ws.n1)
2050 }
2051 };
2052 let t_next = sampler::sample_with_scratch(
2053 &logits,
2054 &self.sampler_config,
2055 &all_ids,
2056 &mut self.rng,
2057 &mut self.sampler_scratch,
2058 );
2059 if self.confidence_on {
2060 confidence.push(top1_prob_t(&logits, t_next, calib_temp));
2061 }
2062 attention::recycle_buf(&mut logits);
2063 if trace_on {
2064 let skill = router.as_ref().and_then(|r| r.active_id());
2068 traces.push(TokenTrace {
2069 t: generated,
2070 token_id: t_next,
2071 confidence: confidence.last().copied().unwrap_or(0.0),
2072 active_skill: skill,
2073 recon: None,
2074 switched: false,
2075 });
2076 }
2077 if !commit!(t_next) {
2078 break 'decode;
2079 }
2080 if generated >= max_tokens {
2081 break 'decode;
2082 }
2083
2084 if self.kv_cache.needs_eviction() {
2085 let keep = (self.kv_cache.max_seq_len / 2).max(1);
2086 self.kv_cache.evict(keep);
2087 }
2088
2089 match &mut mtp {
2090 #[cfg(feature = "gpu")]
2092 Some(m) if graph_spec && generated + 1 < max_tokens && next_pos > 0 => {
2093 if let Some((extra, n_pos, new_h)) = self.graph_spec_step(
2094 m,
2095 &hidden,
2096 t_next,
2097 next_pos,
2098 &mut drafted,
2099 &mut accepted,
2100 ) {
2101 next_pos = n_pos;
2102 hidden = new_h;
2103 let mut stopped = false;
2104 for &id in &extra {
2105 if self.confidence_on {
2106 confidence.push(0.0);
2107 }
2108 if !commit!(id) {
2109 stopped = true;
2110 break;
2111 }
2112 }
2113 if stopped {
2114 break 'decode;
2115 }
2116 continue 'decode;
2117 }
2118 hidden = self.forward_layers(&self.embed_single(t_next), next_pos, task_mask);
2120 next_pos += 1;
2121 continue 'decode;
2122 }
2123 Some(m) if !graph_spec && generated + 1 < max_tokens => {
2125 let draft = self.mtp_step(m, &hidden, t_next, next_pos - 1);
2126 drafted += 1;
2127 let emb1 = self.embed_single(t_next);
2128 let emb2 = self.embed_single(draft);
2129 let (h1, h2) = self.forward_pair(&emb1, &emb2, next_pos);
2130
2131 inference::rms_norm_into(
2132 &h1,
2133 &self.weights.final_norm,
2134 self.rms_eps,
2135 self.norm_style,
2136 &mut self.ws.n1,
2137 );
2138 let mut logits1 = self.lm_head_forward(&self.ws.n1);
2139 let t_after = sampler::sample_with_scratch(
2140 &logits1,
2141 &self.sampler_config,
2142 &all_ids,
2143 &mut self.rng,
2144 &mut self.sampler_scratch,
2145 );
2146 if self.confidence_on {
2147 confidence.push(top1_prob_t(&logits1, t_after, calib_temp));
2148 }
2149 attention::recycle_buf(&mut logits1);
2150 if trace_on {
2151 traces.push(TokenTrace {
2154 t: generated,
2155 token_id: t_after,
2156 confidence: confidence.last().copied().unwrap_or(0.0),
2157 active_skill: None,
2158 recon: None,
2159 switched: false,
2160 });
2161 }
2162 let stop = !commit!(t_after);
2163
2164 if t_after == draft {
2165 accepted += 1;
2166 self.commit_linear_scratch();
2167 let _ = self.mtp_step(m, &h1, t_after, next_pos);
2168 hidden = h2;
2169 next_pos += 2;
2170 } else {
2171 for layer in &mut self.kv_cache.layers {
2173 layer.truncate_last(1);
2174 }
2175 if !stop {
2176 let _ = self.mtp_step(m, &h1, t_after, next_pos);
2177 hidden = self.forward_layers(
2178 &self.embed_single(t_after),
2179 next_pos + 1,
2180 None,
2181 );
2182 }
2183 next_pos += 2;
2184 }
2185 if stop {
2186 break 'decode;
2187 }
2188 }
2189 _ => {
2191 #[cfg(feature = "gpu")]
2196 if Self::dsv4_spec_on() && self.dsv4.is_some() {
2197 static SAID: std::sync::Once = std::sync::Once::new();
2198 SAID.call_once(|| {
2199 eprintln!(
2200 "dsv4-spec гейт: mtp={} mask={} router={} trace={} temp={} rep={} ",
2201 !self.dsv4_mtp.is_empty(),
2202 task_mask.is_none(),
2203 router.is_none(),
2204 !trace_on,
2205 self.sampler_config.temperature < 1e-6,
2206 self.sampler_config.repetition_penalty == 1.0,
2207 );
2208 });
2209 }
2210 #[cfg(feature = "gpu")]
2211 if Self::dsv4_spec_on()
2212 && self.dsv4.is_some()
2213 && !self.dsv4_mtp.is_empty()
2214 && task_mask.is_none()
2215 && router.is_none()
2216 && !trace_on
2217 && self.sampler_config.temperature < 1e-6
2218 && self.sampler_config.repetition_penalty == 1.0
2219 && generated + 1 < max_tokens
2220 && all_ids.len() >= 2
2221 {
2222 let tip_token = all_ids[all_ids.len() - 2];
2223 if let Some((extra, n_pos)) = self.dsv4_spec_step(
2224 tip_token,
2225 t_next,
2226 next_pos,
2227 &mut drafted,
2228 &mut accepted,
2229 ) {
2230 next_pos = n_pos;
2231 let mut stopped = false;
2232 for &id in &extra {
2233 if self.confidence_on {
2234 confidence.push(0.0);
2235 }
2236 if !commit!(id) {
2237 stopped = true;
2238 break;
2239 }
2240 }
2241 if stopped {
2242 break 'decode;
2243 }
2244 continue 'decode;
2245 }
2246 }
2247 self.graph_want_logits = fuse_lm;
2248 let mut t_fwd = t_next;
2254 let pure_greedy = self.sampler_config.temperature < 1e-6
2255 && self.sampler_config.repetition_penalty == 1.0
2256 && self.sampler_config.suppress_tokens.is_empty();
2257 let burst_k = std::env::var("CMF_MULTISTEP")
2262 .ok()
2263 .and_then(|v| v.parse::<usize>().ok())
2264 .unwrap_or(0);
2265 if pure_greedy
2266 && burst_k >= 1
2267 && fuse_lm
2268 && task_mask.is_none()
2269 && router.is_none()
2270 && !trace_on
2271 && !self.confidence_on
2272 {
2273 let mut stopped = false;
2274 loop {
2275 let room = max_tokens.saturating_sub(generated);
2276 if room <= 2 {
2277 break;
2278 }
2279 let k = burst_k.min(room - 1);
2280 if k < 1 {
2281 break;
2282 }
2283 let Some(ids) = self.try_multi_burst(t_fwd, next_pos, k) else {
2284 break;
2285 };
2286 next_pos += k;
2287 for &id in &ids {
2288 if !commit!(id) {
2289 stopped = true;
2290 break;
2291 }
2292 }
2293 if stopped {
2294 break;
2295 }
2296 t_fwd = *ids.last().unwrap();
2297 }
2298 if stopped {
2299 break 'decode;
2300 }
2301 }
2302 hidden = self.forward_layers(&self.embed_single(t_fwd), next_pos, task_mask);
2303 next_pos += 1;
2304 if let Some(r) = &mut router {
2307 let phi = self.dyn_phi_ema.clone();
2308 let decision = r.step(&phi, generated);
2309 if let Some(new_active) = decision {
2310 let _ = self.set_active_skill(new_active);
2311 }
2312 if trace_on {
2315 if let Some(last) = traces.last_mut() {
2316 let e = r.last_best_e();
2317 last.recon = e.is_finite().then_some(e);
2318 last.switched = decision.is_some();
2319 }
2320 }
2321 }
2322 }
2323 }
2324 }
2325
2326 self.graph_want_logits = false;
2327 self.graph_logits = None;
2328 if router.is_some() {
2330 let _ = self.set_active_skill(None);
2331 }
2332 self.dyn_router = router.or(self.dyn_router.take());
2333 self.mtp = mtp.or(self.mtp.take());
2334
2335 let output_ids = &all_ids[input_ids.len()..];
2336 let forwarded = input_ids.len() + output_ids.len().saturating_sub(1);
2340 self.kv_history = all_ids[..forwarded.min(all_ids.len())].to_vec();
2341 confidence.truncate(output_ids.len()); traces.truncate(output_ids.len());
2343 Ok(GenerateResult {
2344 text: self.tokenizer.decode(output_ids),
2345 token_ids: output_ids.to_vec(),
2346 prompt_tokens: input_ids.len(),
2347 tokens_generated: generated,
2348 finish_reason,
2349 mtp_drafted: drafted,
2350 mtp_accepted: accepted,
2351 token_confidence: confidence,
2352 traces,
2353 })
2354 }
2355
2356 fn mtp_step(
2360 &mut self,
2361 m: &mut MtpModule,
2362 hidden: &[f32],
2363 next_token: u32,
2364 position: usize,
2365 ) -> u32 {
2366 self.mtp_step_h(m, hidden, next_token, position).0
2367 }
2368
2369 fn chain_probe_note(depth: usize, prefix_ok: bool) {
2373 use std::sync::Mutex;
2374 static T: Mutex<Vec<(u64, u64)>> = Mutex::new(Vec::new());
2375 let mut t = T.lock().unwrap();
2376 if t.len() <= depth {
2377 t.resize(depth + 1, (0, 0));
2378 }
2379 t[depth].0 += 1;
2380 t[depth].1 += prefix_ok as u64;
2381 if depth == 0 && t[0].0 % 128 == 0 {
2382 let line: Vec<String> = t
2383 .iter()
2384 .enumerate()
2385 .map(|(d, (n, k))| format!("d{}={:.0}%({n})", d + 1, 100.0 * *k as f64 / (*n).max(1) as f64))
2386 .collect();
2387 eprintln!("mtp-chain: {}", line.join(" "));
2388 }
2389 }
2390
2391 fn mtp_step_h(
2395 &mut self,
2396 m: &mut MtpModule,
2397 hidden: &[f32],
2398 next_token: u32,
2399 position: usize,
2400 ) -> (u32, Vec<f32>) {
2401 let e = self.embed_single(next_token);
2405 let mut cat = vec![0.0f32; 2 * self.hidden_size];
2406 let (cat_e, cat_h) = cat.split_at_mut(self.hidden_size);
2407 inference::rms_norm_into(&e, &m.enorm, self.rms_eps, self.norm_style, cat_e);
2408 inference::rms_norm_into(hidden, &m.hnorm, self.rms_eps, self.norm_style, cat_h);
2409 let mut x = vec![0.0f32; self.hidden_size];
2410 m.eh_proj.matvec(&cat, &mut x, self.pool.as_deref());
2411
2412 let lw = &m.layer;
2414 inference::rms_norm_into(
2415 &x,
2416 &lw.input_norm,
2417 self.rms_eps,
2418 self.norm_style,
2419 &mut self.ws.n1,
2420 );
2421 let attn = match &lw.attn {
2422 AttnKind::Mla(_) => unreachable!("MLA has no MTP/pair path"),
2424 AttnKind::Kda(_) => unreachable!("KDA has no MTP/pair path"),
2425 AttnKind::Full {
2426 wq,
2427 wk,
2428 wv,
2429 wo,
2430 q_norm,
2431 k_norm,
2432 output_gate,
2433 softplus_gate,
2434 bias,
2435 } => {
2436 let mut cfg = self.attn_cfg(position);
2437 cfg.q_norm = q_norm.as_deref();
2438 cfg.k_norm = k_norm.as_deref();
2439 cfg.output_gate = *output_gate;
2440 cfg.softplus_gate = softplus_gate
2441 .as_ref()
2442 .map(|(gate, per_head)| (gate, *per_head));
2443 cfg.bias = bias
2444 .as_ref()
2445 .map(|(q, k, v)| (q.as_slice(), k.as_slice(), v.as_slice()));
2446 attention::qwen_attention(&self.ws.n1, wq, wk, wv, wo, &mut m.kv, &cfg)
2447 }
2448 AttnKind::Linear(_) | AttnKind::LinearGdn(_) | AttnKind::ShortConv(_) => {
2449 unreachable!("MTP block is full attention")
2450 }
2451 };
2452 for (i, &a) in attn.iter().enumerate() {
2453 x[i] += a;
2454 }
2455 inference::rms_norm_into(
2456 &x,
2457 &lw.post_norm,
2458 self.rms_eps,
2459 self.norm_style,
2460 &mut self.ws.p1,
2461 );
2462 let ffn = ffn_forward(&lw.ffn, &self.ws.p1, self.pool.as_deref(), None);
2463 for (i, &f) in ffn.iter().enumerate() {
2464 x[i] += f;
2465 }
2466
2467 inference::rms_norm_into(
2468 &x,
2469 &m.final_norm,
2470 self.rms_eps,
2471 self.norm_style,
2472 &mut self.ws.n1,
2473 );
2474 let mut lg = self.lm_head_forward(&self.ws.n1);
2475 let draft = sampler::argmax(&lg);
2476 attention::recycle_buf(&mut lg);
2477 (draft, x)
2478 }
2479
2480 fn mtp_warm(&mut self, m: &mut MtpModule, hidden: &[f32], next_token: u32, position: usize) {
2484 let e = self.embed_single(next_token);
2485 let mut cat = vec![0.0f32; 2 * self.hidden_size];
2486 let (cat_e, cat_h) = cat.split_at_mut(self.hidden_size);
2487 inference::rms_norm_into(&e, &m.enorm, self.rms_eps, self.norm_style, cat_e);
2488 inference::rms_norm_into(hidden, &m.hnorm, self.rms_eps, self.norm_style, cat_h);
2489 let mut x = vec![0.0f32; self.hidden_size];
2490 m.eh_proj.matvec(&cat, &mut x, self.pool.as_deref());
2491 inference::rms_norm_into(&x, &m.layer.input_norm, self.rms_eps, self.norm_style, &mut self.ws.n1);
2492 let attn = match &m.layer.attn {
2493 AttnKind::Full { wq, wk, wv, wo, q_norm, k_norm, output_gate, softplus_gate, bias } => {
2494 let mut cfg = self.attn_cfg(position);
2495 cfg.q_norm = q_norm.as_deref();
2496 cfg.k_norm = k_norm.as_deref();
2497 cfg.output_gate = *output_gate;
2498 cfg.softplus_gate = softplus_gate.as_ref().map(|(g, p)| (g, *p));
2499 cfg.bias = bias.as_ref().map(|(q, k, v)| (q.as_slice(), k.as_slice(), v.as_slice()));
2500 attention::qwen_attention(&self.ws.n1, wq, wk, wv, wo, &mut m.kv, &cfg)
2501 }
2502 _ => return,
2503 };
2504 let _ = attn;
2505 }
2506
2507 #[cfg(feature = "gpu")]
2514 #[allow(clippy::too_many_arguments)]
2515 fn graph_spec_step(
2516 &mut self,
2517 m: &mut MtpModule,
2518 hidden: &[f32],
2519 t_next: u32,
2520 next_pos: usize,
2521 drafted: &mut usize,
2522 accepted: &mut usize,
2523 ) -> Option<(Vec<u32>, usize, Vec<f32>)> {
2524 let k_spec: usize = std::env::var("CMF_GRAPH_SPEC_K")
2530 .ok()
2531 .and_then(|v| v.parse().ok())
2532 .filter(|&v| (1..=8).contains(&v))
2533 .unwrap_or(3);
2534 if next_pos == 0 {
2535 return None;
2536 }
2537 let t_round = std::time::Instant::now();
2538 let subs = || crate::gpu_wgpu::SUBMITS.load(std::sync::atomic::Ordering::Relaxed);
2554 let sub0 = subs();
2555 let mut drafts = Vec::with_capacity(k_spec);
2560 let (d1, mut hx) = self.mtp_step_h(m, hidden, t_next, next_pos - 1);
2561 drafts.push(d1);
2562 for j in 1..k_spec {
2563 let (dj, hj) = self.mtp_step_h(m, &hx, drafts[j - 1], next_pos - 1 + j);
2564 drafts.push(dj);
2565 hx = hj;
2566 }
2567 *drafted += k_spec;
2568 let t_draft = t_round.elapsed();
2569 let sub_draft = subs();
2570 let b = k_spec + 1;
2573 let mut hiddens = vec![0.0f32; b * self.hidden_size];
2574 for (i, &t) in std::iter::once(&t_next).chain(drafts.iter()).enumerate() {
2575 let e = self.embed_single(t);
2576 hiddens[i * self.hidden_size..(i + 1) * self.hidden_size].copy_from_slice(&e);
2577 }
2578 let positions: Vec<usize> = (next_pos..next_pos + b).collect();
2579 let (lm_gw, lm_rows) = {
2580 let (_, i, kind, rs) = self.weights.lm_head.graph_weight()?;
2581 (
2582 crate::gpu::GraphW { idx: i, kind, row_scale: rs, data: &[] },
2583 self.weights.lm_head.rows(),
2584 )
2585 };
2586 let mut logits = Vec::new();
2587 let final_norm = self.weights.final_norm.clone();
2588 let ok = self.try_batch_graph_wgpu(
2589 &mut hiddens,
2590 &positions,
2591 b,
2592 Some(crate::gpu::SpecTail {
2593 lm: lm_gw,
2594 lm_rows,
2595 final_norm: &final_norm,
2596 logits_out: &mut logits,
2597 }),
2598 );
2599 if !ok {
2600 m.kv.truncate_last(k_spec);
2603 return None;
2604 }
2605 let t_verify = t_round.elapsed();
2606 let sub_verify = subs();
2607 let ids: Vec<u32> = (0..b)
2609 .map(|i| sampler::argmax(&logits[i * lm_rows..(i + 1) * lm_rows]))
2610 .collect();
2611 let mut a = 0usize;
2612 while a < k_spec && ids[a] == drafts[a] {
2613 a += 1;
2614 }
2615 if a + 1 < b {
2617 crate::gpu::gdn_spec_restore(self.graph_kv_id, a);
2618 }
2619 *accepted += a;
2620 m.kv.truncate_last(k_spec.saturating_sub(1));
2631 let warm_off = std::env::var("CMF_SPEC_WARM").is_ok_and(|v| v == "0");
2632 if !warm_off {
2633 for j in 0..a {
2634 let row = &hiddens[j * self.hidden_size..(j + 1) * self.hidden_size];
2635 let row = row.to_vec();
2636 self.mtp_warm(m, &row, ids[j], next_pos + j);
2637 }
2638 }
2639 let mut row = logits[a * lm_rows..(a + 1) * lm_rows].to_vec();
2641 row.resize(self.vocab_size, 0.0);
2642 if let Some(c) = self.final_softcap {
2643 for l in row.iter_mut() {
2644 *l = c * (*l / c).tanh();
2645 }
2646 }
2647 self.graph_logits = Some(row);
2648 let new_hidden = hiddens[a * self.hidden_size..(a + 1) * self.hidden_size].to_vec();
2649 if std::env::var("CMF_GRAPH_SPEC_TIME").is_ok() {
2655 let end = subs();
2656 eprintln!(
2657 "spec-round: draft {:.1} ms/{} sub | verify {:.1} ms/{} sub | \
2658 commit {:.1} ms/{} sub (accepted {a} of {k_spec})",
2659 t_draft.as_secs_f64() * 1e3,
2660 sub_draft - sub0,
2661 (t_verify - t_draft).as_secs_f64() * 1e3,
2662 sub_verify - sub_draft,
2663 (t_round.elapsed() - t_verify).as_secs_f64() * 1e3,
2664 end - sub_verify,
2665 );
2666 }
2667 Some((drafts[..a].to_vec(), next_pos + a + 1, new_hidden))
2668 }
2669
2670 pub fn measure_pair_fusion(&mut self, iters: usize) -> (f64, f64) {
2679 if !self.pair_supported() {
2680 return (0.0, 0.0);
2681 }
2682 let emb1 = self.embed_single(1);
2683 let emb2 = self.embed_single(2);
2684 let pos = self.kv_cache.seq_len();
2685
2686 let t0 = std::time::Instant::now();
2687 for _ in 0..iters {
2688 let _ = self.forward_layers(&emb1, pos, None);
2689 let _ = self.forward_layers(&emb2, pos + 1, None);
2690 for l in &mut self.kv_cache.layers {
2691 l.truncate_last(2);
2692 }
2693 }
2694 let singles_ms = t0.elapsed().as_secs_f64() * 1000.0 / iters as f64;
2695
2696 let t1 = std::time::Instant::now();
2697 for _ in 0..iters {
2698 let _ = self.forward_pair(&emb1, &emb2, pos);
2699 for l in &mut self.kv_cache.layers {
2700 l.truncate_last(2);
2701 }
2702 }
2703 let pair_ms = t1.elapsed().as_secs_f64() * 1000.0 / iters as f64;
2704 (singles_ms, pair_ms)
2705 }
2706
2707 fn pair_supported(&self) -> bool {
2715 !self.weights.layers.is_empty()
2722 && self.g3n.is_none()
2723 && !self
2724 .weights
2725 .layers
2726 .iter()
2727 .any(|lw| matches!(&lw.attn, AttnKind::Mla(_) | AttnKind::Kda(_)))
2728 }
2729
2730 fn forward_pair(
2731 &mut self,
2732 emb1: &[f32],
2733 emb2: &[f32],
2734 position: usize,
2735 ) -> (Vec<f32>, Vec<f32>) {
2736 let mut h1 = emb1.to_vec();
2737 let mut h2 = emb2.to_vec();
2738 let (_nkv, _hd, hs, _rd, eps) = (
2739 self.num_kv_heads,
2740 self.head_dim,
2741 self.hidden_size,
2742 self.rotary_dim,
2743 self.rms_eps,
2744 );
2745 let pool = self.pool.clone();
2746
2747 for li in 0..self.num_layers {
2748 let lw = &self.weights.layers[self.phys_layer(li)];
2749 inference::rms_norm_into(
2752 &h1,
2753 &lw.input_norm,
2754 self.rms_eps,
2755 self.norm_style,
2756 &mut self.ws.n1,
2757 );
2758 inference::rms_norm_into(
2759 &h2,
2760 &lw.input_norm,
2761 self.rms_eps,
2762 self.norm_style,
2763 &mut self.ws.n2,
2764 );
2765
2766 let (a1, a2) = match &lw.attn {
2767 AttnKind::Mla(_) => unreachable!("MLA has no MTP/pair path"),
2768 AttnKind::Kda(_) => unreachable!("KDA has no MTP/pair path"),
2769 AttnKind::Linear(w) => {
2770 let cfg = self.vmf_cfg.expect("linear layer without vmf_cfg");
2771 let layer = &mut self.kv_cache.layers[li];
2772 let (state, scratch) = (&mut layer.linear_state, &mut layer.linear_scratch);
2773 vmf_phase_pair(
2774 &self.ws.n1,
2775 &self.ws.n2,
2776 w,
2777 &cfg,
2778 state,
2779 scratch,
2780 self.pool.as_deref(),
2781 )
2782 }
2783 AttnKind::LinearGdn(w) => {
2784 let cfg = self.gdn_cfg.expect("gdn layer without gdn_cfg");
2785 let layer = &mut self.kv_cache.layers[li];
2786 let (state, scratch) = (&mut layer.linear_state, &mut layer.linear_scratch);
2787 gdn_pair(
2788 &self.ws.n1,
2789 &self.ws.n2,
2790 w,
2791 &cfg,
2792 state,
2793 scratch,
2794 self.pool.as_deref(),
2795 )
2796 }
2797 AttnKind::ShortConv(w) => {
2798 let cfg = self
2799 .short_conv_cfg
2800 .expect("short-conv layer without short_conv_cfg");
2801 let layer = &mut self.kv_cache.layers[li];
2802 let (state, scratch) = (&mut layer.linear_state, &mut layer.linear_scratch);
2803 short_conv_pair(
2804 &self.ws.n1,
2805 &self.ws.n2,
2806 w,
2807 &cfg,
2808 state,
2809 scratch,
2810 self.pool.as_deref(),
2811 )
2812 }
2813 AttnKind::Full {
2814 wq,
2815 wk,
2816 wv,
2817 wo,
2818 q_norm,
2819 k_norm,
2820 output_gate,
2821 softplus_gate,
2822 bias,
2823 } => {
2824 let inv_freq_l = self.layer_inv_freq(li);
2825 let (nkv_l, hd_l, rd_l) = self.layer_geom(li);
2826 let cfg = QwenAttnCfg {
2827 num_heads: self.layer_num_heads(li),
2828 num_kv_heads: nkv_l,
2829 head_dim: hd_l,
2830 hidden_size: hs,
2831 position,
2832 inv_freq: &inv_freq_l,
2833 rotary_dim: rd_l,
2834 scale: self.attn_scale,
2835 softcap: self.attn_softcap,
2836 window: self.layer_window(li),
2837 v_norm: self.attn_v_norm,
2838 q_norm: q_norm.as_deref(),
2839 k_norm: k_norm.as_deref(),
2840 output_gate: *output_gate,
2841 softplus_gate: softplus_gate
2842 .as_ref()
2843 .map(|(gate, per_head)| (gate, *per_head)),
2844 rope_scale: self.layer_rope_scale(li),
2845 bias: bias
2846 .as_ref()
2847 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
2848 rms_eps: eps,
2849 norm_style: self.norm_style,
2850 pool: pool.as_deref(),
2851 };
2852 attention::qwen_attention_pair(
2853 &self.ws.n1,
2854 &self.ws.n2,
2855 wq,
2856 wk,
2857 wv,
2858 wo,
2859 &mut self.kv_cache.layers[li],
2860 &cfg,
2861 )
2862 }
2863 };
2864 let (a1, a2) = match &self.weights.layers[self.phys_layer(li)].attn_out_norm {
2865 Some(w) => (
2866 inference::rms_norm(&a1, w, self.rms_eps, self.norm_style),
2867 inference::rms_norm(&a2, w, self.rms_eps, self.norm_style),
2868 ),
2869 None => (a1, a2),
2870 };
2871 for i in 0..self.hidden_size {
2872 h1[i] += a1[i];
2873 h2[i] += a2[i];
2874 }
2875 let (mut a1, mut a2) = (a1, a2);
2876 attention::recycle_buf(&mut a1);
2877 attention::recycle_buf(&mut a2);
2878
2879 let lw = &self.weights.layers[self.phys_layer(li)];
2880 inference::rms_norm_into(
2881 &h1,
2882 &lw.post_norm,
2883 self.rms_eps,
2884 self.norm_style,
2885 &mut self.ws.p1,
2886 );
2887 inference::rms_norm_into(
2888 &h2,
2889 &lw.post_norm,
2890 self.rms_eps,
2891 self.norm_style,
2892 &mut self.ws.p2,
2893 );
2894 let (f1, f2) = match &lw.ffn {
2895 FfnKind::DenseMoe(dm) => (
2898 dense_moe_ffn(
2899 dm,
2900 &self.ws.p1,
2901 &h1,
2902 self.rms_eps,
2903 self.norm_style,
2904 self.pool.as_deref(),
2905 ),
2906 dense_moe_ffn(
2907 dm,
2908 &self.ws.p2,
2909 &h2,
2910 self.rms_eps,
2911 self.norm_style,
2912 self.pool.as_deref(),
2913 ),
2914 ),
2915 _ => ffn_forward_pair(
2916 &lw.ffn,
2917 &self.ws.p1,
2918 &self.ws.p2,
2919 self.pool.as_deref(),
2920 None,
2921 ),
2922 };
2923 let (f1, f2) = match &self.weights.layers[self.phys_layer(li)].ffn_out_norm {
2924 Some(w) => (
2925 inference::rms_norm(&f1, w, self.rms_eps, self.norm_style),
2926 inference::rms_norm(&f2, w, self.rms_eps, self.norm_style),
2927 ),
2928 None => (f1, f2),
2929 };
2930 for i in 0..self.hidden_size {
2931 h1[i] += f1[i];
2932 h2[i] += f2[i];
2933 }
2934 let (mut f1, mut f2) = (f1, f2);
2935 attention::recycle_buf(&mut f1);
2936 attention::recycle_buf(&mut f2);
2937 if let Some(sc) = self.weights.layers[self.phys_layer(li)].layer_scale {
2938 for i in 0..self.hidden_size {
2939 h1[i] *= sc;
2940 h2[i] *= sc;
2941 }
2942 }
2943 if self.is_loop_end(li) && li + 1 < self.num_layers {
2945 h1 = inference::rms_norm(
2946 &h1,
2947 &self.weights.final_norm,
2948 self.rms_eps,
2949 self.norm_style,
2950 );
2951 h2 = inference::rms_norm(
2952 &h2,
2953 &self.weights.final_norm,
2954 self.rms_eps,
2955 self.norm_style,
2956 );
2957 }
2958 }
2959 (h1, h2)
2960 }
2961
2962 fn commit_linear_scratch(&mut self) {
2964 for layer in &mut self.kv_cache.layers {
2965 if !layer.linear_scratch.is_empty() {
2966 std::mem::swap(&mut layer.linear_state, &mut layer.linear_scratch);
2967 layer.linear_scratch.clear();
2968 }
2969 }
2970 }
2971
2972 pub fn forward_ids(
2975 &mut self,
2976 ids: &[u32],
2977 task_mask: Option<&TaskMask>,
2978 ) -> Result<Vec<f32>, String> {
2979 if ids.is_empty() {
2980 return Err("empty id sequence".to_string());
2981 }
2982 self.kv_cache.clear();
2983 self.kv_history.clear();
2984 self.o1_begin();
2985 let mut hidden = vec![0.0f32; self.hidden_size];
2986 let mut pos = 0usize;
2987 if self.can_prefill_batched() && !self.graph_prefill_preferred() && ids.len() > 2 {
2995 let chunk = prefill_chunk();
2999 let hs = self.hidden_size;
3000 while pos < ids.len() {
3001 let end = (pos + chunk).min(ids.len());
3002 let hb = self.prefill_batch_masked(&ids[pos..end], pos, task_mask);
3003 hidden.copy_from_slice(&hb[(end - pos - 1) * hs..]);
3004 pos = end;
3005 }
3006 }
3007 if task_mask.is_none()
3016 && !self.graph_prefill_preferred()
3017 && !std::env::var("CMF_PAIR").is_ok_and(|v| v == "0")
3018 && self.pair_supported()
3019 {
3020 while pos + 1 < ids.len() {
3021 let e1 = self.embed_single(ids[pos]);
3022 let e2 = self.embed_single(ids[pos + 1]);
3023 let (_, h2) = self.forward_pair(&e1, &e2, pos);
3024 self.commit_linear_scratch();
3025 hidden = h2;
3026 pos += 2;
3027 }
3028 }
3029 while pos < ids.len() {
3030 hidden = self.forward_layers(&self.embed_single(ids[pos]), pos, task_mask);
3031 pos += 1;
3032 }
3033 self.o1_seal();
3037 let normed = inference::rms_norm(
3038 &hidden,
3039 &self.weights.final_norm,
3040 self.rms_eps,
3041 self.norm_style,
3042 );
3043 Ok(self.lm_head_forward(&normed))
3044 }
3045
3046 pub fn ppl_ids(&mut self, ids: &[u32]) -> f64 {
3053 let (nll, cnt) = self.nll_ids_from(ids, 0);
3054 (nll / cnt.max(1) as f64).exp()
3055 }
3056
3057 pub fn probe_ffn_mass(&mut self, ids: &[u32]) -> Vec<Vec<f64>> {
3062 self.kv_cache.clear();
3063 self.kv_history.clear();
3064 FFN_PROBE.with(|p| {
3065 *p.borrow_mut() = Some(vec![vec![0f64; self.intermediate_size]; self.num_layers]);
3066 });
3067 crate::gpu::cpu_scope(|| {
3068 for (pos, &id) in ids.iter().enumerate() {
3069 let emb = self.embed_single(id);
3070 let _ = self.forward_layers(&emb, pos, None);
3071 }
3072 });
3073 self.kv_cache.clear();
3074 self.kv_history.clear();
3075 FFN_PROBE
3076 .with(|p| p.borrow_mut().take())
3077 .unwrap_or_default()
3078 }
3079
3080 pub fn ppl_ids_masked(&mut self, ids: &[u32], mask: &TaskMask) -> f64 {
3084 self.kv_cache.clear();
3085 self.kv_history.clear();
3086 let mut nll = 0f64;
3087 let mut cnt = 0usize;
3088 let mut hidden = vec![0f32; self.hidden_size];
3089 for (pos, &id) in ids.iter().enumerate() {
3090 if pos > 0 {
3091 inference::rms_norm_into(
3092 &hidden,
3093 &self.weights.final_norm,
3094 self.rms_eps,
3095 self.norm_style,
3096 &mut self.ws.n1,
3097 );
3098 let mut logits = self.lm_head_forward(&self.ws.n1);
3099 let max = logits.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
3100 let sum: f64 = logits.iter().map(|&v| ((v - max) as f64).exp()).sum();
3101 let p = ((logits[id as usize] - max) as f64).exp() / sum.max(1e-300);
3102 nll -= p.max(1e-300).ln();
3103 cnt += 1;
3104 attention::recycle_buf(&mut logits);
3105 }
3106 let emb = self.embed_single(id);
3107 hidden = self.forward_layers(&emb, pos, Some(mask));
3108 }
3109 self.kv_cache.clear();
3110 self.kv_history.clear();
3111 (nll / cnt.max(1) as f64).exp()
3112 }
3113
3114 pub fn nll_ids_masked(
3133 &mut self,
3134 ids: &[u32],
3135 start: usize,
3136 task_mask: Option<&TaskMask>,
3137 ) -> (f64, usize) {
3138 self.nll_ids_inner(ids, start, task_mask)
3139 }
3140
3141 pub fn nll_ids_from(&mut self, ids: &[u32], start: usize) -> (f64, usize) {
3142 self.nll_ids_inner(ids, start, None)
3143 }
3144
3145 fn nll_ids_inner(
3146 &mut self,
3147 ids: &[u32],
3148 start: usize,
3149 task_mask: Option<&TaskMask>,
3150 ) -> (f64, usize) {
3151 self.kv_cache.clear();
3152 self.kv_history.clear();
3153 let mut nll = 0f64;
3154 let mut cnt = 0usize;
3155 if self.can_prefill_batched() {
3156 const CHUNK: usize = 128;
3162 const LM_SUB: usize = 32;
3163 let n = ids.len().saturating_sub(1);
3164 let hs = self.hidden_size;
3165 let rows = self.weights.lm_head.rows();
3166 let mut pos = 0usize;
3167 while pos < n {
3168 let end = (pos + CHUNK).min(n);
3169 let bsz = end - pos;
3170 let hb = self.prefill_batch_masked(&ids[pos..end], pos, task_mask);
3171 let mut k0 = 0usize;
3172 while k0 < bsz {
3173 let k1 = (k0 + LM_SUB).min(bsz);
3174 let sb = k1 - k0;
3175 if pos + k1 <= start {
3178 k0 = k1;
3179 continue;
3180 }
3181 let mut normed = vec![0.0f32; sb * hs];
3182 for k in 0..sb {
3183 let r = inference::rms_norm(
3184 &hb[(k0 + k) * hs..(k0 + k + 1) * hs],
3185 &self.weights.final_norm,
3186 self.rms_eps,
3187 self.norm_style,
3188 );
3189 normed[k * hs..(k + 1) * hs].copy_from_slice(&r);
3190 }
3191 let mut logits = vec![0.0f32; sb * rows];
3192 self.weights
3193 .lm_head
3194 .matmat(&normed, sb, &mut logits, self.pool.as_deref());
3195 for k in 0..sb {
3196 if pos + k0 + k < start {
3197 continue;
3198 }
3199 let lg = &mut logits[k * rows..k * rows + self.vocab_size.min(rows)];
3200 if let Some(mu) = self.logit_multiplier {
3201 for v in lg.iter_mut() {
3202 *v *= mu;
3203 }
3204 }
3205 if let Some(c) = self.final_softcap {
3209 for v in lg.iter_mut() {
3210 *v = c * (*v / c).tanh();
3211 }
3212 }
3213 let lg = &logits[k * rows..k * rows + self.vocab_size.min(rows)];
3214 let target = ids[pos + k0 + k + 1] as usize;
3215 let max = lg.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
3216 let lse: f64 = lg
3217 .iter()
3218 .map(|&v| ((v - max) as f64).exp())
3219 .sum::<f64>()
3220 .ln()
3221 + max as f64;
3222 nll += lse - lg[target] as f64;
3223 cnt += 1;
3224 if std::env::var("CMF_PPL_TRACE").is_ok() {
3225 let top = lg
3226 .iter()
3227 .enumerate()
3228 .max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
3229 .map(|(i, _)| i)
3230 .unwrap_or(0);
3231 eprintln!(
3232 "BTRACE pos {} target {} nll {:.4} top {} lg_t {:.3} lg_top {:.3}",
3233 pos + k0 + k,
3234 target,
3235 lse - lg[target] as f64,
3236 top,
3237 lg[target],
3238 lg[top]
3239 );
3240 }
3241 }
3242 k0 = k1;
3243 }
3244 pos = end;
3245 }
3246 self.kv_cache.clear();
3247 self.kv_history.clear();
3248 return (nll, cnt);
3249 }
3250 for pos in 0..ids.len().saturating_sub(1) {
3251 let hidden = self.forward_layers(&self.embed_single(ids[pos]), pos, task_mask);
3252 let out_of_band = self.graph_logits.take();
3260 if pos < start {
3261 continue;
3262 }
3263 let logits = match out_of_band {
3264 Some(lg) => lg,
3265 None => {
3266 let normed = inference::rms_norm(
3267 &hidden,
3268 &self.weights.final_norm,
3269 self.rms_eps,
3270 self.norm_style,
3271 );
3272 self.lm_head_forward(&normed)
3276 }
3277 };
3278 let target = ids[pos + 1] as usize;
3279 let max = logits.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
3280 let lse: f64 = logits
3281 .iter()
3282 .map(|&v| ((v - max) as f64).exp())
3283 .sum::<f64>()
3284 .ln()
3285 + max as f64;
3286 let tok_nll = lse - logits[target] as f64;
3287 if std::env::var("CMF_PPL_TRACE").is_ok() && pos < 48 {
3288 let top = logits
3289 .iter()
3290 .enumerate()
3291 .max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
3292 .map(|(i, _)| i)
3293 .unwrap_or(0);
3294 eprintln!(
3295 "pos {pos:3} tgt {target:6} nll {tok_nll:7.3} | top1 {top:6} lg[t]={:.2} lg[top]={:.2}",
3296 logits[target], logits[top]
3297 );
3298 }
3299 nll += tok_nll;
3300 cnt += 1;
3301 }
3302 self.kv_cache.clear();
3303 self.kv_history.clear();
3304 (nll, cnt)
3305 }
3306
3307 pub fn nll_ids_o1(&mut self, ids: &[u32], prefill: usize) -> (f64, usize) {
3323 self.kv_cache.clear();
3324 self.kv_history.clear();
3325 self.o1_begin();
3326 let n = ids.len().saturating_sub(1);
3327 let p = prefill.min(n);
3328 let mut pos = 0usize;
3330 if self.can_prefill_batched() {
3331 const CHUNK: usize = 128;
3332 while pos < p {
3333 let end = (pos + CHUNK).min(p);
3334 let _ = self.prefill_batch(&ids[pos..end], pos);
3335 pos = end;
3336 }
3337 } else {
3338 while pos < p {
3339 let _ = self.forward_layers(&self.embed_single(ids[pos]), pos, None);
3340 pos += 1;
3341 }
3342 }
3343 self.o1_seal();
3344
3345 let mut nll = 0f64;
3346 let mut cnt = 0usize;
3347 for pos in p..n {
3348 let hidden = self.forward_layers(&self.embed_single(ids[pos]), pos, None);
3349 let normed = inference::rms_norm(
3350 &hidden,
3351 &self.weights.final_norm,
3352 self.rms_eps,
3353 self.norm_style,
3354 );
3355 let logits = self.lm_head_forward(&normed);
3359 let target = ids[pos + 1] as usize;
3360 let max = logits.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
3361 let lse: f64 = logits
3362 .iter()
3363 .map(|&v| ((v - max) as f64).exp())
3364 .sum::<f64>()
3365 .ln()
3366 + max as f64;
3367 let tok_nll = lse - logits[target] as f64;
3368 if std::env::var("CMF_PPL_TRACE").is_ok() && pos < 48 {
3369 let top = logits
3370 .iter()
3371 .enumerate()
3372 .max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
3373 .map(|(i, _)| i)
3374 .unwrap_or(0);
3375 eprintln!(
3376 "pos {pos:3} tgt {target:6} nll {tok_nll:7.3} | top1 {top:6} lg[t]={:.2} lg[top]={:.2}",
3377 logits[target], logits[top]
3378 );
3379 }
3380 nll += tok_nll;
3381 cnt += 1;
3382 }
3383 self.kv_cache.clear();
3384 self.kv_history.clear();
3385 (nll, cnt)
3386 }
3387
3388 pub fn calib_ids(&mut self, ids: &[u32], temps: &[f32]) -> (Vec<bool>, Vec<Vec<f32>>) {
3396 self.kv_cache.clear();
3397 self.kv_history.clear();
3398 let n = ids.len().saturating_sub(1);
3399 let mut correct = Vec::with_capacity(n);
3400 let mut pmax = Vec::with_capacity(n);
3401 for pos in 0..n {
3402 let emb = self.embed_single(ids[pos]);
3403 let hidden = self.forward_layers(&emb, pos, None);
3404 let normed = inference::rms_norm(
3405 &hidden,
3406 &self.weights.final_norm,
3407 self.rms_eps,
3408 self.norm_style,
3409 );
3410 let logits = self.lm_head_forward(&normed);
3414 let target = ids[pos + 1] as usize;
3415 let (mut amax, mut mval) = (0usize, f32::NEG_INFINITY);
3416 for (i, &v) in logits.iter().enumerate() {
3417 if v > mval {
3418 mval = v;
3419 amax = i;
3420 }
3421 }
3422 correct.push(amax == target);
3423 let row: Vec<f32> = temps
3424 .iter()
3425 .map(|&t| {
3426 let tt = t.max(1e-3);
3427 let s: f32 = logits.iter().map(|&v| ((v - mval) / tt).exp()).sum();
3428 1.0 / s.max(1e-12) })
3430 .collect();
3431 pmax.push(row);
3432 }
3433 self.kv_cache.clear();
3434 self.kv_history.clear();
3435 (correct, pmax)
3436 }
3437
3438 pub fn ppl_ids_dynamic(&mut self, ids: &[u32]) -> (f64, usize) {
3445 let mut router = match self.dyn_router.take() {
3446 Some(r) => r,
3447 None => return (self.ppl_ids(ids), 0),
3448 };
3449 router.reset();
3450 self.dyn_phi_seen = 0;
3451 let _ = self.set_active_skill(None);
3452
3453 self.kv_cache.clear();
3454
3455 self.kv_history.clear();
3456 let mut nll = 0f64;
3457 let mut cnt = 0usize;
3458 for pos in 0..ids.len().saturating_sub(1) {
3459 let hidden = self.forward_layers(&self.embed_single(ids[pos]), pos, None);
3460 let normed = inference::rms_norm(
3461 &hidden,
3462 &self.weights.final_norm,
3463 self.rms_eps,
3464 self.norm_style,
3465 );
3466 let logits = self.lm_head_forward(&normed);
3470 let target = ids[pos + 1] as usize;
3471 let max = logits.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
3472 let lse: f64 = logits
3473 .iter()
3474 .map(|&v| ((v - max) as f64).exp())
3475 .sum::<f64>()
3476 .ln()
3477 + max as f64;
3478 let tok_nll = lse - logits[target] as f64;
3479 if std::env::var("CMF_PPL_TRACE").is_ok() && pos < 48 {
3480 let top = logits
3481 .iter()
3482 .enumerate()
3483 .max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
3484 .map(|(i, _)| i)
3485 .unwrap_or(0);
3486 eprintln!(
3487 "pos {pos:3} tgt {target:6} nll {tok_nll:7.3} | top1 {top:6} lg[t]={:.2} lg[top]={:.2}",
3488 logits[target], logits[top]
3489 );
3490 }
3491 nll += tok_nll;
3492 cnt += 1;
3493 let phi = self.dyn_phi_ema.clone();
3495 if let Some(new_active) = router.step(&phi, pos) {
3496 let _ = self.set_active_skill(new_active);
3497 }
3498 }
3499 let switches = router.switches.len();
3500 let _ = self.set_active_skill(None);
3501 self.dyn_router = Some(router);
3502 self.kv_cache.clear();
3503 self.kv_history.clear();
3504 ((nll / cnt.max(1) as f64).exp(), switches)
3505 }
3506
3507 pub fn probe_phi(&mut self, ids: &[u32], layer: usize) -> Vec<f32> {
3509 self.kv_cache.clear();
3510 self.kv_history.clear();
3511 let mut acc = vec![0f32; self.hidden_size];
3512 for (pos, &id) in ids.iter().enumerate() {
3513 let h = self.forward_layers_upto(&self.embed_single(id), pos, None, Some(layer));
3514 for (a, v) in acc.iter_mut().zip(&h) {
3515 *a += v;
3516 }
3517 }
3518 let n = ids.len().max(1) as f32;
3519 for a in acc.iter_mut() {
3520 *a /= n;
3521 }
3522 self.kv_cache.clear();
3523 self.kv_history.clear();
3524 acc
3525 }
3526
3527 fn prefill_batch(&mut self, ids: &[u32], start_pos: usize) -> Vec<f32> {
3533 self.prefill_batch_masked(ids, start_pos, None)
3534 }
3535
3536 fn prefill_batch_masked(
3542 &mut self,
3543 ids: &[u32],
3544 start_pos: usize,
3545 task_mask: Option<&TaskMask>,
3546 ) -> Vec<f32> {
3547 self.prefill_batch_span(PrefillIn::Ids(ids), start_pos, task_mask, 0, usize::MAX)
3548 }
3549
3550 fn prefill_batch_span(
3556 &mut self,
3557 input: PrefillIn<'_>,
3558 start_pos: usize,
3559 task_mask: Option<&TaskMask>,
3560 from: usize,
3561 upto_excl: usize,
3562 ) -> Vec<f32> {
3563 let hs = self.hidden_size;
3564 let b = match input {
3565 PrefillIn::Ids(ids) => ids.len(),
3566 PrefillIn::Hidden(hb) => hb.len() / hs,
3567 };
3568 let upto_excl = upto_excl.min(self.num_layers);
3569 let mut h: Vec<f32>;
3573 let mut h_ready;
3574 match input {
3575 PrefillIn::Ids(_) => {
3576 h = vec![0.0; b * hs];
3577 h_ready = false;
3578 }
3579 PrefillIn::Hidden(hb) => {
3580 h = hb.to_vec();
3581 h_ready = true;
3582 }
3583 }
3584 let fill_h = |h: &mut Vec<f32>, me: &Self| {
3585 if let PrefillIn::Ids(ids) = input {
3586 for (bi, &id) in ids.iter().enumerate() {
3587 let e = me.embed_single(id);
3588 h[bi * hs..(bi + 1) * hs].copy_from_slice(&e);
3589 }
3590 }
3591 };
3592 let (_nkv, _hd, _rd, eps) = (
3593 self.num_kv_heads,
3594 self.head_dim,
3595 self.rotary_dim,
3596 self.rms_eps,
3597 );
3598 let pool = self.pool.clone();
3599 let norm_style = self.norm_style;
3600
3601 #[cfg(target_os = "macos")]
3602 let mut chunk_skip_until = 0usize;
3603 for li in from..upto_excl {
3604 crate::gpu::set_layer(li as i64); #[cfg(target_os = "macos")]
3611 if task_mask.is_none() {
3612 if li < chunk_skip_until {
3613 continue;
3614 }
3615 if !h_ready && li == 0 && self.weights.embed_tokens.q8_row_parts().is_none() {
3621 fill_h(&mut h, self);
3622 h_ready = true;
3623 }
3624 let ids_for_embed = match input {
3625 PrefillIn::Ids(ids) => (!h_ready && li == 0).then_some(ids),
3626 PrefillIn::Hidden(_) => None,
3627 };
3628 let end = self.chunk_run_gpu(li, &mut h, b, start_pos, ids_for_embed, upto_excl);
3629 if end > li {
3630 h_ready = true;
3631 chunk_skip_until = end;
3632 if self.is_loop_end(end - 1) && end < self.num_layers {
3635 for bi in 0..b {
3636 let normed = inference::rms_norm(
3637 &h[bi * hs..(bi + 1) * hs],
3638 &self.weights.final_norm,
3639 eps,
3640 norm_style,
3641 );
3642 h[bi * hs..(bi + 1) * hs].copy_from_slice(&normed);
3643 }
3644 }
3645 continue;
3646 }
3647 }
3648 if !h_ready {
3649 fill_h(&mut h, self);
3650 h_ready = true;
3651 }
3652 let lw = &self.weights.layers[self.phys_layer(li)];
3653 match &lw.attn {
3655 AttnKind::Kda(w) => {
3656 let cfg = self.kda_cfg.expect("kda layer without kda_cfg");
3658 let mut normed = vec![0.0f32; b * hs];
3659 for bi in 0..b {
3660 inference::rms_norm_into(
3661 &h[bi * hs..(bi + 1) * hs],
3662 &lw.input_norm,
3663 eps,
3664 norm_style,
3665 &mut normed[bi * hs..(bi + 1) * hs],
3666 );
3667 }
3668 let attn = crate::linear_core::kda_forward_batch(
3669 &normed,
3670 b,
3671 w,
3672 &cfg,
3673 &mut self.kv_cache.layers[li].linear_state,
3674 pool.as_deref(),
3675 );
3676 for (dst, &a) in h.iter_mut().zip(&attn) {
3677 *dst += a;
3678 }
3679 }
3680 AttnKind::LinearGdn(w) => {
3681 let cfg = self.gdn_cfg.expect("gdn layer without gdn_cfg");
3683 let mut normed = vec![0.0f32; b * hs];
3684 for bi in 0..b {
3685 let r = inference::rms_norm(
3686 &h[bi * hs..(bi + 1) * hs],
3687 &lw.input_norm,
3688 eps,
3689 norm_style,
3690 );
3691 normed[bi * hs..(bi + 1) * hs].copy_from_slice(&r);
3692 }
3693 let attn = crate::linear_core::gdn_forward_batch(
3694 &normed,
3695 b,
3696 w,
3697 &cfg,
3698 &mut self.kv_cache.layers[li].linear_state,
3699 pool.as_deref(),
3700 );
3701 for (dst, &a) in h.iter_mut().zip(&attn) {
3702 *dst += a;
3703 }
3704 }
3705 AttnKind::ShortConv(w) => {
3706 let cfg = self
3709 .short_conv_cfg
3710 .expect("short-conv layer without short_conv_cfg");
3711 let mut normed = vec![0.0f32; b * hs];
3712 for bi in 0..b {
3713 inference::rms_norm_into(
3714 &h[bi * hs..(bi + 1) * hs],
3715 &lw.input_norm,
3716 eps,
3717 norm_style,
3718 &mut normed[bi * hs..(bi + 1) * hs],
3719 );
3720 }
3721 let attn = short_conv_forward_batch(
3722 &normed,
3723 b,
3724 w,
3725 &cfg,
3726 &mut self.kv_cache.layers[li].linear_state,
3727 pool.as_deref(),
3728 );
3729 for (dst, &a) in h.iter_mut().zip(&attn) {
3730 *dst += a;
3731 }
3732 }
3733 AttnKind::Mla(w) => {
3734 let inv_freq_l = self.layer_inv_freq(li);
3737 let rs = self.layer_rope_scale(li);
3738 let mut normed = vec![0.0f32; hs];
3739 for bi in 0..b {
3740 inference::rms_norm_into(
3741 &h[bi * hs..(bi + 1) * hs],
3742 &lw.input_norm,
3743 eps,
3744 norm_style,
3745 &mut normed,
3746 );
3747 let ao = mla_attention(
3748 w,
3749 &normed,
3750 &mut self.kv_cache.layers[li],
3751 start_pos + bi,
3752 &inv_freq_l,
3753 rs,
3754 eps,
3755 pool.as_deref(),
3756 );
3757 for (dst, &a) in h[bi * hs..(bi + 1) * hs].iter_mut().zip(&ao) {
3758 *dst += a;
3759 }
3760 }
3761 }
3762 AttnKind::Full {
3763 wq,
3764 wk,
3765 wv,
3766 wo,
3767 q_norm,
3768 k_norm,
3769 output_gate,
3770 softplus_gate,
3771 bias,
3772 } => {
3773 let mut normed = vec![0.0f32; b * hs];
3777 for bi in 0..b {
3778 inference::rms_norm_into(
3779 &h[bi * hs..(bi + 1) * hs],
3780 &lw.input_norm,
3781 eps,
3782 norm_style,
3783 &mut normed[bi * hs..(bi + 1) * hs],
3784 );
3785 }
3786 let inv_freq_l = self.layer_inv_freq(li);
3787 let (nkv_l, hd_l, rd_l) = self.layer_geom(li);
3788 let cfg = QwenAttnCfg {
3789 num_heads: self.layer_num_heads(li),
3790 num_kv_heads: nkv_l,
3791 head_dim: hd_l,
3792 hidden_size: hs,
3793 position: start_pos,
3794 inv_freq: &inv_freq_l,
3795 rotary_dim: rd_l,
3796 scale: self.attn_scale,
3797 softcap: self.attn_softcap,
3798 window: self.layer_window(li),
3799 v_norm: self.attn_v_norm,
3800 q_norm: q_norm.as_deref(),
3801 k_norm: k_norm.as_deref(),
3802 output_gate: *output_gate,
3803 softplus_gate: softplus_gate
3804 .as_ref()
3805 .map(|(gate, per_head)| (gate, *per_head)),
3806 rope_scale: self.layer_rope_scale(li),
3807 bias: bias
3808 .as_ref()
3809 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
3810 rms_eps: eps,
3811 norm_style,
3812 pool: pool.as_deref(),
3813 };
3814 let mut attn = attention::qwen_attention_batch(
3815 &normed,
3816 b,
3817 wq,
3818 wk,
3819 wv,
3820 wo,
3821 &mut self.kv_cache.layers[li],
3822 &cfg,
3823 );
3824 if let Some(w) = &lw.attn_out_norm {
3825 for bi in 0..b {
3826 inference::rms_norm_into(
3827 &attn[bi * hs..(bi + 1) * hs],
3828 w,
3829 eps,
3830 norm_style,
3831 &mut normed[bi * hs..(bi + 1) * hs],
3832 );
3833 }
3834 attn.copy_from_slice(&normed);
3835 }
3836 for (dst, &a) in h.iter_mut().zip(&attn) {
3837 *dst += a;
3838 }
3839 }
3840 AttnKind::Linear(w) => {
3841 for bi in 0..b {
3842 let normed = inference::rms_norm(
3843 &h[bi * hs..(bi + 1) * hs],
3844 &lw.input_norm,
3845 eps,
3846 norm_style,
3847 );
3848 vmf_phase_forward(
3849 &normed,
3850 w,
3851 &self.vmf_cfg.expect("linear layer without vmf_cfg"),
3852 &mut self.kv_cache.layers[li].linear_state,
3853 pool.as_deref(),
3854 )
3855 .iter()
3856 .enumerate()
3857 .for_each(|(i, &a)| h[bi * hs + i] += a);
3858 }
3859 }
3860 }
3861
3862 let lw = &self.weights.layers[self.phys_layer(li)];
3864 let mut post = vec![0.0f32; b * hs];
3865 for bi in 0..b {
3866 let r =
3867 inference::rms_norm(&h[bi * hs..(bi + 1) * hs], &lw.post_norm, eps, norm_style);
3868 post[bi * hs..(bi + 1) * hs].copy_from_slice(&r);
3869 }
3870 let mask_row = task_mask
3873 .filter(|m| m.ffn_active_count(li) < self.intermediate_size)
3874 .and_then(|m| m.ffn_masks.get(li))
3875 .map(|v| v.as_slice());
3876 let mut ffn = match &lw.ffn {
3877 FfnKind::Dense(d) => dense_ffn_batch(d, &post, b, pool.as_deref(), mask_row),
3878 FfnKind::Moe(m) => moe_ffn_batch(m, &post, b, hs, pool.as_deref(), None),
3879 FfnKind::DenseMoe(dm) => {
3882 let mut out = vec![0.0f32; b * hs];
3883 for bi in 0..b {
3884 let r = dense_moe_ffn(
3885 dm,
3886 &post[bi * hs..(bi + 1) * hs],
3887 &h[bi * hs..(bi + 1) * hs],
3888 eps,
3889 norm_style,
3890 pool.as_deref(),
3891 );
3892 out[bi * hs..(bi + 1) * hs].copy_from_slice(&r);
3893 }
3894 out
3895 }
3896 };
3897 if let Some(w) = &lw.ffn_out_norm {
3898 for bi in 0..b {
3899 inference::rms_norm_into(
3900 &ffn[bi * hs..(bi + 1) * hs],
3901 w,
3902 eps,
3903 norm_style,
3904 &mut post[bi * hs..(bi + 1) * hs],
3905 );
3906 }
3907 ffn.copy_from_slice(&post);
3908 }
3909 for (dst, &f) in h.iter_mut().zip(&ffn) {
3910 *dst += f;
3911 }
3912 if let Some(sc) = lw.layer_scale {
3913 for v in h.iter_mut() {
3914 *v *= sc;
3915 }
3916 }
3917 if let Ok(tp) = std::env::var("CMF_TRACE_POS") {
3918 if let Some(t) = tp.parse::<usize>().ok() {
3919 if t >= start_pos && t < start_pos + b {
3920 let bi = t - start_pos;
3921 let row = &h[bi * hs..(bi + 1) * hs];
3922 let n: f32 = row.iter().map(|x| x * x).sum::<f32>().sqrt();
3923 eprintln!(
3924 "BATCH pos {t} after layer {li}: |h| = {n:.6} h0 {:.6} h1 {:.6}",
3925 row[0], row[1]
3926 );
3927 }
3928 }
3929 }
3930 if std::env::var("CMF_DEBUG_LAYERS").is_ok() {
3934 let row = &h[(b - 1) * hs..b * hs];
3935 let rms =
3936 (row.iter().map(|&v| (v as f64) * (v as f64)).sum::<f64>() / hs as f64).sqrt();
3937 let mx = row.iter().fold(0f32, |m, &v| m.max(v.abs()));
3938 eprintln!(
3939 "layer {li:>3} {:>10} ffn={:<5} rms={rms:>12.4} max={mx:>12.4}",
3940 match &self.weights.layers[self.phys_layer(li)].attn {
3941 AttnKind::LinearGdn(_) => "gdn",
3942 AttnKind::Linear(_) => "vmf",
3943 AttnKind::ShortConv(_) => "conv",
3944 _ => "attn",
3945 },
3946 match &lw.ffn {
3947 FfnKind::Moe(_) => "moe",
3948 FfnKind::Dense(_) => "dense",
3949 FfnKind::DenseMoe(_) => "dense+moe",
3950 },
3951 );
3952 }
3953 if self.is_loop_end(li) && li + 1 < self.num_layers {
3955 for bi in 0..b {
3956 let normed = inference::rms_norm(
3957 &h[bi * hs..(bi + 1) * hs],
3958 &self.weights.final_norm,
3959 eps,
3960 norm_style,
3961 );
3962 h[bi * hs..(bi + 1) * hs].copy_from_slice(&normed);
3963 }
3964 }
3965 if std::env::var("CMF_TRACE_H").is_ok() {
3966 let n = h[..hs].iter().map(|v| v.abs()).sum::<f32>() / hs as f32;
3967 let mx = h[..hs].iter().fold(0.0f32, |a, &v| a.max(v.abs()));
3968 eprintln!(
3969 "layer {li}: mean|h|={n:.4} max|h|={mx:.2} scale={:?}",
3970 lw.layer_scale
3971 );
3972 }
3973 }
3974 crate::gpu::set_layer(-1); h
3976 }
3977
3978 fn embed_single(&self, id: u32) -> Vec<f32> {
3980 let mut out = vec![0.0f32; self.hidden_size];
3981 if (id as usize) < self.weights.embed_tokens.rows() {
3982 self.weights.embed_tokens.row_f32(id as usize, &mut out);
3983 }
3984 if self.embed_multiplier != 1.0 {
3985 for v in out.iter_mut() {
3986 *v *= self.embed_multiplier;
3987 }
3988 }
3989 if self.dsv4.is_some() {
3993 let mut v = vec![0.0f32; self.hidden_size.max(1)];
3994 v[0] = id as f32;
3995 return v;
3996 }
3997 if let Some(b) = &self.g3n {
4000 return b.0.extend_embedding(id, &out, self.pool.as_deref());
4001 }
4002 out
4003 }
4004
4005 #[cfg(target_os = "macos")]
4011 fn chunk_run_gpu(
4012 &mut self,
4013 li0: usize,
4014 h: &mut [f32],
4015 b: usize,
4016 pos0: usize,
4017 embed_ids: Option<&[u32]>,
4018 cap: usize,
4019 ) -> usize {
4020 if !crate::gpu::enabled_here()
4024 || std::env::var("CMF_GPU_CHUNK")
4025 .map(|v| v == "0")
4026 .unwrap_or(false)
4027 || b < 32
4028 || self.swa.is_some()
4029 || self.global_attn.is_some()
4030 || self.attn_v_norm
4031 || (self.attn_scale - 1.0 / (self.head_dim as f32).sqrt()).abs() > 1e-9
4032 {
4033 return li0;
4034 }
4035 let Some(model) = self.model.clone() else {
4036 return li0;
4037 };
4038 let inv_freq = self.inv_freq.clone();
4039 let (nh, nkv, hd, hs) = (
4040 self.num_heads,
4041 self.num_kv_heads,
4042 self.head_dim,
4043 self.hidden_size,
4044 );
4045 let loop_end = if self.loop_final_norm {
4049 ((li0 / self.physical_layers) + 1) * self.physical_layers
4050 } else {
4051 self.num_layers
4052 };
4053 let mut layers: Vec<crate::gpu_metal::ChunkLayer> = Vec::new();
4054 let mut stored_at: Vec<usize> = Vec::new();
4055 for li in li0..self.num_layers.min(loop_end).min(cap) {
4056 let lw = &self.weights.layers[self.phys_layer(li)];
4057 if lw.attn_out_norm.is_some() || lw.ffn_out_norm.is_some() || lw.layer_scale.is_some() {
4058 break;
4059 }
4060 let AttnKind::Full {
4061 wq,
4062 wk,
4063 wv,
4064 wo,
4065 q_norm,
4066 k_norm,
4067 output_gate: false,
4068 softplus_gate: None,
4069 bias,
4070 } = &lw.attn
4071 else {
4072 break;
4073 };
4074 let FfnKind::Dense(d) = &lw.ffn else { break };
4075 if d.act != Act::Silu {
4076 break;
4077 }
4078 fn cw(t: &QTensor) -> Option<(usize, usize, usize, &[f32])> {
4083 t.q8_row_parts()
4084 .or_else(|| t.q4t_parts().map(|(i, r, c)| (i, r, c, &[][..])))
4085 .or_else(|| t.q4tp_parts().map(|(i, r, c)| (i, r, c, &[][..])))
4086 }
4087 let parts = (
4088 cw(wq),
4089 cw(wk),
4090 cw(wv),
4091 cw(wo),
4092 cw(&d.gate_proj),
4093 cw(&d.up_proj),
4094 cw(&d.down_proj),
4095 );
4096 let (Some(pq), Some(pk), Some(pv), Some(po), Some(pg), Some(pu), Some(pd)) = parts
4097 else {
4098 break;
4099 };
4100 let layer = &self.kv_cache.layers[li];
4101 if layer.mode != crate::kv_cache::KvMode::F32 || layer.o1.is_some() {
4102 break;
4103 }
4104 stored_at.push(layer.head_len(0));
4105 layers.push(crate::gpu_metal::ChunkLayer {
4106 model: &model,
4107 kv_id: self.graph_kv_id,
4108 layer: li,
4109 wq: pq,
4110 wk: pk,
4111 wv: pv,
4112 wo: po,
4113 gate: pg,
4114 up: pu,
4115 down: pd,
4116 input_norm: &lw.input_norm,
4117 post_norm: &lw.post_norm,
4118 bias: bias
4119 .as_ref()
4120 .map(|(a, bb, cc)| (a.as_slice(), bb.as_slice(), cc.as_slice())),
4121 q_norm: q_norm.as_deref(),
4122 k_norm: k_norm.as_deref(),
4123 inv_freq: &inv_freq,
4124 rd: self.rotary_dim,
4125 nh,
4126 nkv,
4127 hd,
4128 hs,
4129 inter: d.gate_proj.rows(),
4130 gemma: matches!(self.norm_style, cortiq_core::NormStyle::Gemma),
4131 eps: self.rms_eps as f32,
4132 });
4133 }
4134 if layers.is_empty() {
4135 return li0;
4136 }
4137 let row = nkv * hd;
4138 let mut store: Vec<(Vec<f32>, Vec<f32>, Vec<f32>)> = stored_at
4139 .iter()
4140 .map(|&st| (vec![0f32; b * row], vec![0f32; b * row], vec![0f32; st + b]))
4141 .collect();
4142 let mut io: Vec<crate::gpu_metal::ChunkIo> = Vec::with_capacity(layers.len());
4143 for (i, (ok, ov, oi)) in store.iter_mut().enumerate() {
4144 let li = layers[i].layer;
4145 let layer = &self.kv_cache.layers[li];
4146 io.push(crate::gpu_metal::ChunkIo {
4147 cpu_stored: stored_at[i],
4148 cpu_k: (0..nkv).map(|g| layer.head_keys(g)).collect(),
4149 cpu_v: (0..nkv).map(|g| layer.head_values(g)).collect(),
4150 out_k: ok,
4151 out_v: ov,
4152 imp: oi,
4153 });
4154 }
4155 let n_run = layers.len();
4156 let last = layers.last().map(|l| l.layer + 1).unwrap_or(li0);
4157 let ep = embed_ids.and_then(|ids| {
4160 self.weights
4161 .embed_tokens
4162 .q8_row_parts()
4163 .map(|(idx, rows, _c, rs)| crate::gpu_metal::ChunkEmbed {
4164 idx,
4165 rows,
4166 row_scale: rs,
4167 ids,
4168 mult: self.embed_multiplier,
4169 })
4170 });
4171 if embed_ids.is_some() && ep.is_none() {
4172 return li0;
4173 }
4174 if !crate::gpu_metal::chunk_run_gpu(&layers, &mut io, h, b, pos0, ep.as_ref()) {
4175 return li0;
4176 }
4177 drop(io);
4178 drop(layers);
4179 for (i, (ok, ov, oi)) in store.iter().enumerate().take(n_run) {
4182 let li = li0 + i;
4183 let layer = &mut self.kv_cache.layers[li];
4184 for bi in 0..b {
4185 layer.append(
4186 &ok[bi * row..(bi + 1) * row],
4187 &ov[bi * row..(bi + 1) * row],
4188 &[],
4189 );
4190 }
4191 layer.accumulate_imp(oi);
4192 }
4193 last
4194 }
4195
4196 fn layer_is_local(&self, li: usize) -> bool {
4199 if let Some(layers) = &self.sliding_layers {
4200 return layers.get(li).copied().unwrap_or(false);
4201 }
4202 match self.swa {
4203 Some((_, pattern)) => (li + 1) % pattern.max(1) != 0,
4204 None => false,
4205 }
4206 }
4207
4208 fn layer_inv_freq(&self, li: usize) -> std::sync::Arc<Vec<f32>> {
4211 if self.layer_is_local(li) {
4212 if let Some(f) = &self.inv_freq_local {
4213 return f.clone();
4214 }
4215 } else if let Some(f) = &self.inv_freq_global {
4216 return f.clone();
4217 }
4218 self.inv_freq.clone()
4219 }
4220
4221 fn layer_window(&self, li: usize) -> Option<usize> {
4223 self.swa
4224 .and_then(|(w, _)| self.layer_is_local(li).then_some(w))
4225 }
4226
4227 fn layer_num_heads(&self, li: usize) -> usize {
4228 self.attention_heads_per_layer
4229 .as_ref()
4230 .and_then(|v| v.get(li).copied())
4231 .unwrap_or(self.num_heads)
4232 }
4233
4234 fn layer_rope_scale(&self, li: usize) -> f32 {
4235 if self.layer_is_local(li) {
4236 self.rope_scale_local
4237 } else {
4238 self.rope_scale
4239 }
4240 }
4241
4242 fn layer_geom(&self, li: usize) -> (usize, usize, usize) {
4245 if !self.layer_is_local(li) {
4246 if let Some((ghd, gkv)) = self.global_attn {
4247 return (gkv, ghd, ghd);
4248 }
4249 }
4250 (
4251 self.num_kv_heads,
4252 self.head_dim,
4253 if self.layer_is_local(li) {
4254 self.rotary_dim_local.unwrap_or(self.rotary_dim)
4255 } else {
4256 self.rotary_dim
4257 },
4258 )
4259 }
4260
4261 fn forward_layers(
4263 &mut self,
4264 hidden: &[f32],
4265 position: usize,
4266 task_mask: Option<&TaskMask>,
4267 ) -> Vec<f32> {
4268 self.forward_layers_upto(hidden, position, task_mask, None)
4269 }
4270
4271 pub fn embed_id(&self, id: u32) -> Vec<f32> {
4279 self.embed_single(id)
4280 }
4281
4282 pub fn split_supported(&self) -> Result<(), String> {
4286 if self.dsv4.is_some() {
4287 return Err(
4288 "network split: DeepSeek-V4 runs its own fused stack (not splittable yet)".into(),
4289 );
4290 }
4291 if self.g3n.is_some() {
4292 return Err(
4293 "network split: Gemma-3n runs its own AltUp stack (not splittable yet)".into(),
4294 );
4295 }
4296 Ok(())
4297 }
4298
4299 pub fn forward_span(
4304 &mut self,
4305 hidden: &[f32],
4306 position: usize,
4307 from: usize,
4308 upto: usize,
4309 task_mask: Option<&TaskMask>,
4310 ) -> Result<Vec<f32>, String> {
4311 self.split_supported()?;
4312 if from > upto || upto >= self.num_layers {
4313 return Err(format!(
4314 "forward_span: layer range {from}..={upto} outside 0..{}",
4315 self.num_layers
4316 ));
4317 }
4318 if hidden.len() != self.hidden_size {
4319 return Err(format!(
4320 "forward_span: hidden len {} ≠ hidden_size {}",
4321 hidden.len(),
4322 self.hidden_size
4323 ));
4324 }
4325 Ok(self.forward_layers_span(hidden, position, task_mask, from, Some(upto)))
4326 }
4327
4328 pub fn logits_from_hidden(&mut self, hidden: &[f32]) -> Vec<f32> {
4331 let normed = inference::rms_norm(
4332 hidden,
4333 &self.weights.final_norm,
4334 self.rms_eps,
4335 self.norm_style,
4336 );
4337 self.lm_head_forward(&normed)
4338 }
4339
4340 pub fn sample_next(&mut self, logits: &[f32], past_tokens: &[u32]) -> u32 {
4342 sampler::sample_with_scratch(
4343 logits,
4344 &self.sampler_config,
4345 past_tokens,
4346 &mut self.rng,
4347 &mut self.sampler_scratch,
4348 )
4349 }
4350
4351 pub fn reset_session(&mut self) {
4353 self.kv_cache.clear();
4354 self.kv_history.clear();
4355 crate::gpu::graph_kv_reset(self.graph_kv_id);
4356 }
4357
4358 pub fn prefill_span_ids(
4364 &mut self,
4365 ids: &[u32],
4366 start_pos: usize,
4367 upto: usize,
4368 task_mask: Option<&TaskMask>,
4369 ) -> Result<Vec<f32>, String> {
4370 self.split_supported()?;
4371 if upto >= self.num_layers {
4372 return Err(format!(
4373 "prefill_span_ids: upto {upto} outside 0..{}",
4374 self.num_layers
4375 ));
4376 }
4377 if self.can_prefill_batched() && !self.graph_prefill_preferred() {
4381 Ok(self.prefill_batch_span(PrefillIn::Ids(ids), start_pos, task_mask, 0, upto + 1))
4382 } else {
4383 let hs = self.hidden_size;
4384 let mut out = Vec::with_capacity(ids.len() * hs);
4385 for (i, &id) in ids.iter().enumerate() {
4386 let emb = self.embed_id(id);
4387 out.extend_from_slice(&self.forward_span(
4388 &emb,
4389 start_pos + i,
4390 0,
4391 upto,
4392 task_mask,
4393 )?);
4394 }
4395 Ok(out)
4396 }
4397 }
4398
4399 pub fn prefill_span_hidden(
4402 &mut self,
4403 hidden: &[f32],
4404 start_pos: usize,
4405 from: usize,
4406 upto: usize,
4407 task_mask: Option<&TaskMask>,
4408 ) -> Result<Vec<f32>, String> {
4409 self.split_supported()?;
4410 let hs = self.hidden_size;
4411 if hidden.is_empty() || hidden.len() % hs != 0 {
4412 return Err(format!(
4413 "prefill_span_hidden: {} floats is not a multiple of hidden {hs}",
4414 hidden.len()
4415 ));
4416 }
4417 if from > upto || upto >= self.num_layers {
4418 return Err(format!(
4419 "prefill_span_hidden: layer range {from}..={upto} outside 0..{}",
4420 self.num_layers
4421 ));
4422 }
4423 if self.can_prefill_batched() && !self.graph_prefill_preferred() {
4424 Ok(self.prefill_batch_span(
4425 PrefillIn::Hidden(hidden),
4426 start_pos,
4427 task_mask,
4428 from,
4429 upto + 1,
4430 ))
4431 } else {
4432 let b = hidden.len() / hs;
4433 let mut out = Vec::with_capacity(hidden.len());
4434 for i in 0..b {
4435 let h = self.forward_span(
4436 &hidden[i * hs..(i + 1) * hs],
4437 start_pos + i,
4438 from,
4439 upto,
4440 task_mask,
4441 )?;
4442 out.extend_from_slice(&h);
4443 }
4444 Ok(out)
4445 }
4446 }
4447
4448 fn try_token_graph_wgpu(
4452 &self,
4453 hidden: &[f32],
4454 position: usize,
4455 logits_out: &mut Vec<f32>,
4456 layers_run: &mut usize,
4457 ) -> Option<Vec<f32>> {
4458 self.try_token_graph_wgpu_steps(
4459 hidden,
4460 position,
4461 logits_out,
4462 1,
4463 None,
4464 Some(layers_run),
4465 0,
4466 self.num_layers,
4467 )
4468 }
4469
4470 fn try_token_graph_wgpu_span(
4474 &self,
4475 hidden: &[f32],
4476 position: usize,
4477 logits_out: &mut Vec<f32>,
4478 from: usize,
4479 upto_excl: usize,
4480 layers_run: &mut usize,
4481 ) -> Option<Vec<f32>> {
4482 self.try_token_graph_wgpu_steps(
4483 hidden,
4484 position,
4485 logits_out,
4486 1,
4487 None,
4488 Some(layers_run),
4489 from,
4490 upto_excl,
4491 )
4492 }
4493
4494 fn try_multi_burst(&self, t_next: u32, position: usize, k: usize) -> Option<Vec<u32>> {
4498 if self.o1_active() || self.attn_softcap > 0.0 {
4499 return None;
4500 }
4501 let graph_on = match std::env::var("CMF_GPU_WGPU_GRAPH").ok().as_deref() {
4502 Some("0") => return None,
4503 Some(_) => true,
4504 None => crate::gpu::wgpu_graph_default(),
4505 };
4506 if !graph_on {
4507 return None;
4508 }
4509 let emb = self.embed_single(t_next);
4510 let mut lg = Vec::new();
4511 let mut ids = Vec::new();
4512 self.try_token_graph_wgpu_steps(
4513 &emb,
4514 position,
4515 &mut lg,
4516 k,
4517 Some(&mut ids),
4518 None,
4519 0,
4520 self.num_layers,
4521 )?;
4522 (ids.len() == k).then_some(ids)
4523 }
4524
4525 fn try_token_graph_wgpu_steps(
4529 &self,
4530 hidden: &[f32],
4531 position: usize,
4532 logits_out: &mut Vec<f32>,
4533 steps: usize,
4534 ids_out: Option<&mut Vec<u32>>,
4535 layers_run: Option<&mut usize>,
4536 from: usize,
4537 upto_excl: usize,
4538 ) -> Option<Vec<f32>> {
4539 let o1_gpu = std::env::var("CMF_O1_GPU").as_deref() == Ok("1");
4542 if (self.o1_active() && !o1_gpu) || self.attn_softcap > 0.0 {
4543 return None;
4547 }
4548 let o1_views: Vec<Option<Vec<crate::nystrom::O1DeviceView<'_>>>> = (from..upto_excl)
4553 .map(|li| {
4554 if !o1_gpu {
4555 return None;
4556 }
4557 self.kv_cache.layers[self.phys_layer(li)].o1_views()
4558 })
4559 .collect();
4560 if self.o1_active() && o1_gpu {
4561 let want: usize = (from..upto_excl)
4564 .filter(|li| !matches!(self.kv_cache.layers[self.phys_layer(*li)].o1, None))
4565 .count();
4566 let have = o1_views.iter().filter(|v| v.is_some()).count();
4567 if want == 0 || have != want {
4568 return None;
4569 }
4570 }
4571 let nh = self.num_heads;
4572 let (nkv, hd, rd) = self.layer_geom(0);
4573 let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
4574 let mut layers = Vec::with_capacity(upto_excl - from);
4575 let mut model = None;
4576 let dbg = std::env::var("CMF_GRAPH_DEBUG").is_ok();
4577 fn gw(t: &QTensor) -> Option<crate::gpu::GraphW<'_>> {
4578 if let Some((_, i, kind, rs)) = t.graph_weight() {
4579 return Some(crate::gpu::GraphW {
4580 idx: i,
4581 kind,
4582 row_scale: rs,
4583 data: &[],
4584 });
4585 }
4586 t.as_f32().map(|d| crate::gpu::GraphW {
4588 idx: 0,
4589 kind: 4,
4590 row_scale: &[],
4591 data: d,
4592 })
4593 }
4594 for li in from..upto_excl {
4595 let lw = &self.weights.layers[self.phys_layer(li)];
4596 if dbg {
4597 let ak = match &lw.attn {
4598 AttnKind::Mla(_) => "Mla".into(),
4599 AttnKind::Full {
4600 output_gate, bias, ..
4601 } => format!("Full gate={output_gate} bias={}", bias.is_some()),
4602 AttnKind::LinearGdn(_) => "LinearGdn".into(),
4603 AttnKind::Kda(_) => "Kda".into(),
4604 AttnKind::Linear(_) => "Linear".into(),
4605 AttnKind::ShortConv(_) => "ShortConv".into(),
4606 };
4607 let fk = match &lw.ffn {
4608 FfnKind::Dense(_) => "Dense",
4609 FfnKind::Moe(_) => "Moe",
4610 FfnKind::DenseMoe(_) => "DenseMoe",
4611 };
4612 eprintln!("graph L{li}: attn={ak} ffn={fk}");
4613 }
4614 let gffn = match &lw.ffn {
4615 FfnKind::DenseMoe(_) => return None, FfnKind::Dense(d) => crate::gpu::GraphFfn::Dense {
4617 gate: gw(&d.gate_proj)?,
4618 up: gw(&d.up_proj)?,
4619 down: gw(&d.down_proj)?,
4620 },
4621 FfnKind::Moe(m) => {
4622 if m.router_sigmoid
4627 || m.expert_bias.is_some()
4628 || m.route_tau.is_some()
4629 || m.mask.is_some()
4630 {
4631 return None;
4632 }
4633 let (se, sg) = m.shared.as_ref()?;
4634 let sgate = gw(sg.as_ref()?)?;
4635 let router = gw(&m.router)?;
4636 let inter = m.experts.first()?.gate_proj.rows();
4637 let mut experts = Vec::with_capacity(m.experts.len() + 1);
4638 let mut q4tp: Option<bool> = None;
4641 let mut gu_q2: Option<bool> = None;
4644 for e in m.experts.iter().chain(std::iter::once(se)) {
4645 if !matches!(e.act, Act::Silu)
4646 || e.gate_proj.rows() != inter
4647 || e.up_proj.rows() != inter
4648 {
4649 return None;
4650 }
4651 let (mm, gi, ui, di, is_p, is_q2) = match e.gate_proj.mapped_q4t() {
4652 Some((mm, gi)) => (
4653 mm,
4654 gi,
4655 e.up_proj.mapped_q4t()?.1,
4656 e.down_proj.mapped_q4t()?.1,
4657 false,
4658 false,
4659 ),
4660 None => match e.gate_proj.mapped_q2tp() {
4661 Some((mm, gi)) => (
4662 mm,
4663 gi,
4664 e.up_proj.mapped_q2tp()?.1,
4665 e.down_proj.mapped_q4tp()?.1,
4666 true,
4667 true,
4668 ),
4669 None => {
4670 let (mm, gi) = e.gate_proj.mapped_q4tp()?;
4671 (
4672 mm,
4673 gi,
4674 e.up_proj.mapped_q4tp()?.1,
4675 e.down_proj.mapped_q4tp()?.1,
4676 true,
4677 false,
4678 )
4679 }
4680 },
4681 };
4682 if *q4tp.get_or_insert(is_p) != is_p || *gu_q2.get_or_insert(is_q2) != is_q2
4683 {
4684 tracing::warn!(
4690 "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."
4691 );
4692 return None;
4693 }
4694 model.get_or_insert_with(|| mm.clone());
4695 experts.push((gi, ui, di));
4696 }
4697 crate::gpu::GraphFfn::Moe {
4698 router,
4699 shared_gate: sgate,
4700 experts,
4701 n_exp: m.experts.len(),
4702 top_k: std::env::var("CMF_TOPK_PROBE")
4708 .ok()
4709 .and_then(|v| v.parse::<usize>().ok())
4710 .filter(|k| *k > 0 && *k <= m.top_k)
4711 .unwrap_or(m.top_k),
4712 inter,
4713 norm_topk: m.norm_topk_prob,
4714 q4tp: q4tp?,
4715 gu_q2: gu_q2.unwrap_or(false),
4716 }
4717 }
4718 };
4719 let attn = match &lw.attn {
4720 AttnKind::Full {
4721 wq,
4722 wk,
4723 wv,
4724 wo,
4725 q_norm,
4726 k_norm,
4727 output_gate,
4728 softplus_gate,
4729 bias,
4730 } => {
4731 if softplus_gate.is_some() || self.attention_heads_per_layer.is_some() {
4732 return None;
4733 }
4734 let (m, _, _, _) = wq.graph_weight()?;
4735 model = Some(m.clone());
4736 crate::gpu::GraphAttn::Full {
4737 wq: gw(wq)?,
4738 wk: gw(wk)?,
4739 wv: gw(wv)?,
4740 wo: gw(wo)?,
4741 q_norm: q_norm.as_deref(),
4742 k_norm: k_norm.as_deref(),
4743 bias: bias
4744 .as_ref()
4745 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
4746 output_gate: *output_gate,
4747 cpu_k: self.kv_cache.layers[li].k_heads(),
4748 cpu_v: self.kv_cache.layers[li].v_heads(),
4749 }
4750 }
4751 AttnKind::LinearGdn(w) => {
4752 let cfg = self.gdn_cfg?;
4753 let (m, _, _, _) = w.in_proj_qkv.graph_weight()?;
4754 model = Some(m.clone());
4755 crate::gpu::GraphAttn::Gdn {
4756 qkv: gw(&w.in_proj_qkv)?,
4757 z: gw(&w.in_proj_z)?,
4758 a: gw(&w.in_proj_a)?,
4759 b: gw(&w.in_proj_b)?,
4760 out: gw(&w.out_proj)?,
4761 conv1d: &w.conv1d,
4762 a_log: &w.a_log,
4763 dt_bias: &w.dt_bias,
4764 norm: &w.norm,
4765 nv: cfg.num_v_heads,
4766 nk: cfg.num_k_heads,
4767 dk: cfg.key_head_dim,
4768 dv: cfg.value_head_dim,
4769 kk: cfg.conv_kernel,
4770 cpu_state: &self.kv_cache.layers[self.phys_layer(li)].linear_state,
4771 }
4772 }
4773 _ => return None,
4774 };
4775 layers.push(crate::gpu::GraphLayer {
4776 input_norm: &lw.input_norm,
4777 attn,
4778 post_norm: &lw.post_norm,
4779 ffn: gffn,
4780 });
4781 }
4782 let model = model?;
4783 let lm_gw = if upto_excl == self.num_layers
4789 && self.graph_want_logits
4790 && std::env::var("CMF_GPU_LMHEAD")
4791 .map(|v| v != "0")
4792 .unwrap_or(true)
4793 {
4794 self.weights.lm_head.graph_weight().map(|(_, i, kind, rs)| {
4795 (
4796 crate::gpu::GraphW {
4797 idx: i,
4798 kind,
4799 row_scale: rs,
4800 data: &[],
4801 },
4802 self.weights.lm_head.rows(),
4803 )
4804 })
4805 } else {
4806 None
4807 };
4808 let lm = lm_gw.as_ref().map(|(gw, rows)| (gw, *rows));
4809 let emb_gw = if steps > 1 {
4811 self.weights
4812 .embed_tokens
4813 .graph_weight()
4814 .map(|(_, i, kind, rs)| {
4815 (
4816 crate::gpu::GraphW {
4817 idx: i,
4818 kind,
4819 row_scale: rs,
4820 data: &[],
4821 },
4822 self.weights.embed_tokens.rows(),
4823 self.embed_multiplier as f32,
4824 )
4825 })
4826 } else {
4827 None
4828 };
4829
4830 let loop_norm_at: Vec<usize> = if self.loop_final_norm {
4836 (from..upto_excl.min(self.num_layers - 1))
4837 .filter(|&li| (li + 1) % self.physical_layers == 0)
4838 .map(|li| li - from)
4839 .collect()
4840 } else {
4841 Vec::new()
4842 };
4843 let mut h = hidden.to_vec();
4844 crate::gpu::forward_token_graph(
4845 &model,
4846 self.graph_kv_id,
4847 &layers,
4848 &o1_views,
4849 self.o1_epoch,
4850 &self.inv_freq,
4851 &mut h,
4852 nh,
4853 nkv,
4854 hd,
4855 rd,
4856 self.hidden_size,
4857 self.intermediate_size,
4858 position,
4859 self.kv_cache.max_seq_len,
4860 gemma,
4861 self.rms_eps as f32,
4862 lm,
4863 &self.weights.final_norm,
4864 logits_out,
4865 &loop_norm_at,
4866 steps,
4867 emb_gw.as_ref().map(|(gw, rows, m)| (gw, *rows, *m)),
4868 ids_out,
4869 layers_run,
4870 from,
4871 )
4872 .then_some(h)
4873 }
4874
4875 fn try_batch_graph_wgpu(
4880 &self,
4881 hiddens: &mut [f32],
4882 positions: &[usize],
4883 k: usize,
4884 spec: Option<crate::gpu::SpecTail<'_>>,
4885 ) -> bool {
4886 let _tb = std::time::Instant::now();
4887 if self.attn_softcap > 0.0 {
4888 return false; }
4890 if self.o1_active() {
4891 return false;
4892 }
4893 let nh = self.num_heads;
4894 let (nkv, hd, rd) = self.layer_geom(0);
4895 let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
4896 fn gw(t: &QTensor) -> Option<crate::gpu::GraphW<'_>> {
4897 if let Some((_, i, kind, rs)) = t.graph_weight() {
4898 return Some(crate::gpu::GraphW {
4899 idx: i,
4900 kind,
4901 row_scale: rs,
4902 data: &[],
4903 });
4904 }
4905 t.as_f32().map(|d| crate::gpu::GraphW {
4906 idx: 0,
4907 kind: 4,
4908 row_scale: &[],
4909 data: d,
4910 })
4911 }
4912 let built: Option<(
4913 Vec<crate::gpu::GraphLayer<'_>>,
4914 std::sync::Arc<cortiq_core::CmfModel>,
4915 )> = (|| {
4916 let mut layers = Vec::with_capacity(self.num_layers);
4917 let mut model = None;
4918 for li in 0..self.num_layers {
4919 let lw = &self.weights.layers[self.phys_layer(li)];
4920 let gffn = match &lw.ffn {
4927 FfnKind::Dense(d) => crate::gpu::GraphFfn::Dense {
4928 gate: gw(&d.gate_proj)?,
4929 up: gw(&d.up_proj)?,
4930 down: gw(&d.down_proj)?,
4931 },
4932 FfnKind::Moe(m) => {
4933 if m.router_sigmoid
4934 || m.expert_bias.is_some()
4935 || m.route_tau.is_some()
4936 || m.mask.is_some()
4937 {
4938 return None;
4939 }
4940 let (se, sg) = m.shared.as_ref()?;
4941 let sgate = gw(sg.as_ref()?)?;
4942 let router = gw(&m.router)?;
4943 let inter = m.experts.first()?.gate_proj.rows();
4944 let mut experts = Vec::with_capacity(m.experts.len() + 1);
4945 let mut q4tp: Option<bool> = None;
4946 let mut gu_q2: Option<bool> = None;
4947 for e in m.experts.iter().chain(std::iter::once(se)) {
4948 if !matches!(e.act, Act::Silu)
4949 || e.gate_proj.rows() != inter
4950 || e.up_proj.rows() != inter
4951 {
4952 return None;
4953 }
4954 let (mm, gi, ui, di, is_p, is_q2) = match e.gate_proj.mapped_q4t() {
4958 Some((mm, gi)) => (
4959 mm,
4960 gi,
4961 e.up_proj.mapped_q4t()?.1,
4962 e.down_proj.mapped_q4t()?.1,
4963 false,
4964 false,
4965 ),
4966 None => match e.gate_proj.mapped_q2tp() {
4967 Some((mm, gi)) => (
4968 mm,
4969 gi,
4970 e.up_proj.mapped_q2tp()?.1,
4971 e.down_proj.mapped_q4tp()?.1,
4972 true,
4973 true,
4974 ),
4975 None => {
4976 let (mm, gi) = e.gate_proj.mapped_q4tp()?;
4977 (
4978 mm,
4979 gi,
4980 e.up_proj.mapped_q4tp()?.1,
4981 e.down_proj.mapped_q4tp()?.1,
4982 true,
4983 false,
4984 )
4985 }
4986 },
4987 };
4988 if *q4tp.get_or_insert(is_p) != is_p
4989 || *gu_q2.get_or_insert(is_q2) != is_q2
4990 {
4991 return None;
4992 }
4993 model.get_or_insert_with(|| mm.clone());
4994 experts.push((gi, ui, di));
4995 }
4996 crate::gpu::GraphFfn::Moe {
4997 router,
4998 shared_gate: sgate,
4999 experts,
5000 n_exp: m.experts.len(),
5001 top_k: m.top_k,
5002 inter,
5003 norm_topk: m.norm_topk_prob,
5004 q4tp: q4tp?,
5005 gu_q2: gu_q2.unwrap_or(false),
5006 }
5007 }
5008 _ => return None,
5009 };
5010 let attn = match &lw.attn {
5011 AttnKind::Full {
5012 wq,
5013 wk,
5014 wv,
5015 wo,
5016 q_norm,
5017 k_norm,
5018 output_gate,
5019 softplus_gate,
5020 bias,
5021 } => {
5022 if softplus_gate.is_some() || self.attention_heads_per_layer.is_some() {
5023 return None;
5024 }
5025 let (m, _, _, _) = wq.graph_weight()?;
5026 model = Some(m.clone());
5027 crate::gpu::GraphAttn::Full {
5028 wq: gw(wq)?,
5029 wk: gw(wk)?,
5030 wv: gw(wv)?,
5031 wo: gw(wo)?,
5032 q_norm: q_norm.as_deref(),
5033 k_norm: k_norm.as_deref(),
5034 bias: bias
5035 .as_ref()
5036 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
5037 output_gate: *output_gate,
5038 cpu_k: self.kv_cache.layers[li].k_heads(),
5039 cpu_v: self.kv_cache.layers[li].v_heads(),
5040 }
5041 }
5042 AttnKind::LinearGdn(w) => {
5043 let cfg = self.gdn_cfg?;
5044 let (m, _, _, _) = w.in_proj_qkv.graph_weight()?;
5045 model = Some(m.clone());
5046 crate::gpu::GraphAttn::Gdn {
5047 qkv: gw(&w.in_proj_qkv)?,
5048 z: gw(&w.in_proj_z)?,
5049 a: gw(&w.in_proj_a)?,
5050 b: gw(&w.in_proj_b)?,
5051 out: gw(&w.out_proj)?,
5052 conv1d: &w.conv1d,
5053 a_log: &w.a_log,
5054 dt_bias: &w.dt_bias,
5055 norm: &w.norm,
5056 nv: cfg.num_v_heads,
5057 nk: cfg.num_k_heads,
5058 dk: cfg.key_head_dim,
5059 dv: cfg.value_head_dim,
5060 kk: cfg.conv_kernel,
5061 cpu_state: &self.kv_cache.layers[self.phys_layer(li)].linear_state,
5062 }
5063 }
5064 _ => return None,
5065 };
5066 layers.push(crate::gpu::GraphLayer {
5067 input_norm: &lw.input_norm,
5068 attn,
5069 post_norm: &lw.post_norm,
5070 ffn: gffn,
5071 });
5072 }
5073 Some((layers, model?))
5074 })();
5075 let Some((layers, model)) = built else {
5076 {
5077 use std::sync::atomic::{AtomicBool, Ordering};
5078 static SAID: AtomicBool = AtomicBool::new(false);
5079 if !SAID.swap(true, Ordering::Relaxed) {
5080 tracing::warn!("batch graph: BUILDER refused (layer weights/kinds)");
5081 }
5082 }
5083 return false;
5084 };
5085 if std::env::var("CMF_GRAPH_SPEC_TIME").is_ok() {
5086 eprintln!("batch-build: {:.1} ms", _tb.elapsed().as_secs_f64() * 1e3);
5087 }
5088 crate::gpu::forward_batch_graph(
5089 &model,
5090 self.graph_kv_id,
5091 &layers,
5092 &self.inv_freq,
5093 hiddens,
5094 nh,
5095 nkv,
5096 hd,
5097 rd,
5098 self.hidden_size,
5099 self.intermediate_size,
5100 positions,
5101 self.kv_cache.max_seq_len,
5102 gemma,
5103 self.rms_eps as f32,
5104 k,
5105 spec,
5106 )
5107 }
5108
5109 fn draft_probe() -> bool {
5113 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
5114 *ON.get_or_init(|| std::env::var("CMF_DSV4_DRAFT_PROBE").is_ok_and(|v| v != "0"))
5115}
5116
5117 #[cfg(feature = "gpu")]
5129 fn dsv4_spec_on() -> bool {
5130 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
5131 *ON.get_or_init(|| std::env::var("CMF_DSV4_SPEC").map(|v| v != "0").unwrap_or(true))
5132 }
5133
5134 #[cfg(feature = "gpu")]
5141 fn dsv4_spec_step(
5142 &mut self,
5143 tip_token: u32,
5144 t_next: u32,
5145 next_pos: usize,
5146 drafted: &mut usize,
5147 accepted_ctr: &mut usize,
5148 ) -> Option<(Vec<u32>, usize)> {
5149 let t_all = std::time::Instant::now();
5150 if std::env::var("CMF_DSV4_SPEC_TIME").is_ok() {
5151 thread_local! {
5152 static LAST: std::cell::Cell<Option<std::time::Instant>> =
5153 const { std::cell::Cell::new(None) };
5154 }
5155 LAST.with(|l| {
5156 if let Some(prev) = l.get() {
5157 eprintln!("между раундами {:.1} мс", prev.elapsed().as_secs_f64() * 1e3);
5158 }
5159 l.set(Some(std::time::Instant::now()));
5160 });
5161 }
5162 if std::env::var("CMF_DSV4_SPEC_DEBUG").is_ok() {
5163 eprintln!("spec_step: вход pos={next_pos}");
5164 }
5165 let n_layers = self.dsv4.as_ref().map(|b| b.1.len())?;
5166 let cfg = self.dsv4.as_ref().map(|b| b.2)?;
5167 if self.dspark.is_none() {
5169 let t = crate::dsv4::dspark_targets(&self.dsv4_mtp, &cfg, n_layers);
5170 if t.is_empty() {
5171 return None;
5172 }
5173 crate::dsv4::dspark_arm(&t, cfg.dim);
5174 self.dspark = Some(crate::dsv4::DsparkState::new(
5175 self.dsv4_mtp.len(),
5176 &cfg,
5177 t.len(),
5178 ));
5179 }
5180 let targets = crate::dsv4::dspark_targets(&self.dsv4_mtp, &cfg, n_layers);
5181 let pack = crate::dsv4::dspark_pack_get(&self.dsv4_mtp, &cfg);
5182 if pack.is_none() && std::env::var("CMF_DSV4_SPEC_DEBUG").is_ok() {
5183 eprintln!("spec_step: пак не построился (targets {targets:?})");
5184 }
5185 let pack = pack?;
5186 let block = crate::dsv4::dspark_block();
5187 let b_box = self.dsv4.as_mut()?;
5188 let (g, layers, st) = (&b_box.0, &b_box.1, &mut b_box.3);
5189 let ds = self.dspark.as_mut()?;
5190 let dbg = std::env::var("CMF_DSV4_SPEC_DEBUG").is_ok();
5193 if !crate::dsv4::dspark_take(&mut ds.main_hidden) && !ds.have_hidden {
5194 if dbg {
5195 eprintln!("spec_step: нет захвата");
5196 }
5197 return None;
5198 }
5199 ds.have_hidden = true;
5200 let tip_pos = next_pos.checked_sub(1)?;
5201 let draft_started = std::time::Instant::now();
5202 let mut conf = Vec::new();
5203 let props = crate::dsv4::dspark_draft_gpu(
5204 g,
5205 &self.dsv4_mtp,
5206 &cfg,
5207 ds,
5208 pack,
5209 st.kv_id,
5210 tip_token,
5211 tip_pos,
5212 self.pool.as_deref(),
5213 &mut conf,
5214 );
5215 self.dspark_draft_ns += draft_started.elapsed().as_nanos();
5216 *drafted += block;
5217 if props.is_empty() || props[0] != t_next {
5218 if dbg {
5219 eprintln!(
5220 "spec_step: черновик {} (props0={:?} t_next={t_next})",
5221 if props.is_empty() { "пуст" } else { "мимо" },
5222 props.first()
5223 );
5224 }
5225 return None;
5226 }
5227 let mut k_verify = crate::dsv4::dspark_verify_k().min(props.len());
5228 let conf_min = {
5234 static M: std::sync::OnceLock<f32> = std::sync::OnceLock::new();
5235 *M.get_or_init(|| {
5236 std::env::var("CMF_DSPARK_CONF_MIN")
5237 .ok()
5238 .and_then(|v| v.parse().ok())
5239 .unwrap_or(0.0)
5240 })
5241 };
5242 if conf_min > 0.0 && conf.len() >= props.len() {
5243 let mut keep = 1usize;
5244 while keep < k_verify && conf.get(keep).copied().unwrap_or(0.0) >= conf_min {
5245 keep += 1;
5246 }
5247 k_verify = k_verify.min(keep.max(2));
5248 }
5249 if k_verify < 2 {
5250 return None;
5251 }
5252 let mut fed = Vec::with_capacity(k_verify);
5253 fed.push(t_next);
5254 fed.extend_from_slice(&props[1..k_verify]);
5255 let mut argmax = Vec::new();
5256 let mut logits_all = Vec::new();
5257 let mut walked = Vec::new();
5258 let txn = crate::dsv4::dsv4_verify_chunk(
5259 g,
5260 layers,
5261 &cfg,
5262 st,
5263 &fed,
5264 next_pos,
5265 &self.inv_freq,
5266 self.pool.as_deref(),
5267 &targets,
5268 &mut argmax,
5269 &mut logits_all,
5270 &mut walked,
5271 );
5272 if txn.is_none() && dbg {
5273 eprintln!("spec_step: verify отказал");
5274 }
5275 let txn = txn?;
5276 let b = fed.len();
5277 let mut accepted = 1usize;
5278 while accepted < b && fed[accepted] == argmax[accepted - 1] {
5279 accepted += 1;
5280 }
5281 if std::env::var("CMF_DSV4_SPEC_FORCE_REJECT").is_ok_and(|v| v != "0") {
5286 accepted = 1;
5287 }
5288 if std::env::var("CMF_DSV4_SPEC_TRACE").is_ok() {
5289 eprintln!(
5290 "spec@{next_pos}: fed={fed:?} argmax={argmax:?} accepted={accepted}"
5291 );
5292 }
5293 let t_fin = std::time::Instant::now();
5294 if !crate::dsv4::dsv4_spec_finish(
5295 g,
5296 layers,
5297 &cfg,
5298 st,
5299 txn,
5300 accepted,
5301 &fed,
5302 &self.inv_freq,
5303 self.pool.as_deref(),
5304 ) {
5305 tracing::warn!("dsv4: спекулятивный откат не удался — состояние подозрительно");
5306 return None;
5307 }
5308 if std::env::var("CMF_DSV4_SPEC_TIME").is_ok() {
5309 eprintln!("finish(k={accepted}): {:.1} мс", t_fin.elapsed().as_secs_f64() * 1e3);
5310 }
5311 *accepted_ctr += accepted - 1;
5312 let (hc, dim) = (cfg.hc_mult, cfg.dim);
5317 let dev_caps: Vec<usize> = targets
5324 .iter()
5325 .copied()
5326 .filter(|&t| {
5327 st.dev_set.get(t).copied().unwrap_or(false)
5328 && !st.partial_set.get(t).copied().unwrap_or(false)
5329 })
5330 .collect();
5331 let mut caps_all = vec![0.0f32; dev_caps.len() * b * hc * dim];
5332 if !crate::gpu_wgpu::dsv4_spec_cap_read_all(b, dev_caps.len(), hc * dim, &mut caps_all) {
5333 return None;
5334 }
5335 for t in 0..accepted {
5336 let tip = t + 1 == accepted;
5337 for (slot, &tl) in targets.iter().enumerate() {
5338 if let Some(di) = dev_caps.iter().position(|&d| d == tl) {
5339 let lo = (di * b + t) * hc * dim;
5340 crate::dsv4::dspark_capture(
5341 &caps_all[lo..lo + hc * dim],
5342 &cfg,
5343 slot,
5344 &mut ds.main_hidden,
5345 );
5346 } else if tip
5347 && crate::dsv4::dspark_peek_slot(slot, dim, {
5348 let lo = slot * dim;
5349 &mut ds.main_hidden[lo..lo + dim]
5350 })
5351 {
5352 } else {
5357 crate::dsv4::dspark_capture(
5361 &walked[t * hc * dim..(t + 1) * hc * dim],
5362 &cfg,
5363 slot,
5364 &mut ds.main_hidden,
5365 );
5366 }
5367 }
5368 crate::dsv4::dspark_ring_append(g, &self.dsv4_mtp, &cfg, ds, next_pos + t, self.pool.as_deref());
5369 }
5370 let row = logits_all[(accepted - 1) * cfg.vocab..accepted * cfg.vocab].to_vec();
5371 self.graph_logits = Some(row);
5372 if std::env::var("CMF_DSV4_TRUNK_PICK_DUMP").is_ok() {
5377 crate::dsv4::trunk_freq_note(&crate::dsv4::pick_tally_take());
5378 crate::dsv4::pick_tally_arm();
5379 }
5380 if std::env::var("CMF_DSV4_SPEC_TIME").is_ok() {
5381 eprintln!("spec_step total {:.1} мс (k={accepted})", t_all.elapsed().as_secs_f64() * 1e3);
5382 }
5383 Some((fed[1..accepted].to_vec(), next_pos + accepted))
5384 }
5385
5386 fn dspark_probe(&mut self, position: usize, token_id: u32) {
5387 if self.dsv4_mtp.is_empty() || !Self::draft_probe() {
5388 return;
5389 }
5390 let trunk_now = crate::dsv4::pick_tally_take();
5392 crate::dsv4::trunk_freq_note(&trunk_now);
5393 if !trunk_now.is_empty() {
5394 self.dspark_trunk_picks.push(trunk_now);
5395 let keep = crate::dsv4::dspark_block();
5396 if self.dspark_trunk_picks.len() > keep {
5397 self.dspark_trunk_picks.remove(0);
5398 }
5399 }
5400 for p in std::mem::take(&mut self.dspark_pending) {
5403 let Some(i) = position.checked_sub(p.0 + 1) else {
5404 continue;
5405 };
5406 let mut p = p;
5407 if i < p.1.len() {
5408 if p.2 && p.1[i] == token_id {
5409 p.3 = i + 1;
5410 } else {
5411 p.2 = false;
5412 }
5413 if i + 1 < p.1.len() {
5414 self.dspark_pending.push(p);
5415 continue;
5416 }
5417 }
5418 self.dspark_hist.push(p.3);
5419 self.dspark_real.push(token_id);
5420 }
5421 let Some(b) = &mut self.dsv4 else { return };
5422 let (g, layers, cfg) = (&b.0, &b.1, b.2);
5423 let n_layers = layers.len();
5424 if self.dspark.is_none() {
5425 let t = crate::dsv4::dspark_targets(&self.dsv4_mtp, &cfg, n_layers);
5426 if t.is_empty() {
5427 return;
5428 }
5429 eprintln!("DSpark: захват со слоёв {t:?}, блок {}", crate::dsv4::dspark_block());
5430 crate::dsv4::dspark_arm(&t, cfg.dim);
5431 self.dspark = Some(crate::dsv4::DsparkState::new(
5432 self.dsv4_mtp.len(),
5433 &cfg,
5434 t.len(),
5435 ));
5436 }
5437 let ds = self.dspark.as_mut().unwrap();
5438 if !crate::dsv4::dspark_take(&mut ds.main_hidden) {
5439 return; }
5441 let mut conf = Vec::new();
5442 crate::dsv4::pick_tally_arm();
5443 let draft_started = std::time::Instant::now();
5448 #[cfg(feature = "gpu")]
5449 let gpu_draft = crate::dsv4::dspark_gpu_on();
5450 #[cfg(not(feature = "gpu"))]
5451 let gpu_draft = false;
5452 let props = if gpu_draft {
5453 #[cfg(feature = "gpu")]
5454 {
5455 let kv_id = b.3.kv_id;
5456 match crate::dsv4::dspark_pack_get(&self.dsv4_mtp, &cfg) {
5457 Some(pk) => crate::dsv4::dspark_draft_gpu(
5458 g,
5459 &self.dsv4_mtp,
5460 &cfg,
5461 ds,
5462 pk,
5463 kv_id,
5464 token_id,
5465 position,
5466 self.pool.as_deref(),
5467 &mut conf,
5468 ),
5469 None => Vec::new(),
5470 }
5471 }
5472 #[cfg(not(feature = "gpu"))]
5473 Vec::new()
5474 } else {
5475 crate::gpu::cpu_scope(|| {
5476 crate::dsv4::dspark_draft(
5477 g,
5478 &self.dsv4_mtp,
5479 &cfg,
5480 ds,
5481 token_id,
5482 position,
5483 self.pool.as_deref(),
5484 &mut conf,
5485 )
5486 })
5487 };
5488 self.dspark_draft_ns += draft_started.elapsed().as_nanos();
5489 let draft_picks = crate::dsv4::pick_tally_take();
5490 crate::dsv4::dspark_freq_note(&draft_picks);
5491 crate::dsv4::pick_tally_arm();
5494 if !props.is_empty() {
5495 let (tu, tt) = {
5499 let flat: Vec<(usize, Vec<usize>)> = self
5500 .dspark_trunk_picks
5501 .iter()
5502 .flat_map(|v| v.iter().cloned())
5503 .collect();
5504 let mut per: std::collections::HashMap<usize, Vec<usize>> =
5506 std::collections::HashMap::new();
5507 for (li, picks) in flat {
5508 per.entry(li).or_default().extend(picks);
5509 }
5510 let n = per.len().max(1);
5511 let mut u = 0usize;
5512 let mut t = 0usize;
5513 for (_, v) in per {
5514 t += v.len();
5515 u += v.iter().collect::<std::collections::HashSet<_>>().len();
5516 }
5517 (u / n, t / n)
5518 };
5519 let (du, dt) = crate::dsv4::tally_unique(&draft_picks);
5520 self.dspark_exp.push((tu, tt, du, dt));
5521 self.dspark_pending.push((position, props, true, 0));
5522 }
5523 if self.dspark_hist.len() >= 8 && self.dspark_hist.len() % 8 == 0 {
5524 let n = self.dspark_hist.len() as f32;
5525 let mean: f32 = self.dspark_hist.iter().sum::<usize>() as f32 / n;
5526 let block = crate::dsv4::dspark_block();
5527 let mut at = vec![0usize; block + 1];
5528 for &k in &self.dspark_hist {
5529 at[k] += 1;
5530 }
5531 let mut surv = Vec::with_capacity(block);
5533 for i in 1..=block {
5534 let k = at[i..].iter().sum::<usize>() as f32 / n;
5535 surv.push(format!("{k:.2}"));
5536 }
5537 let distinct = self
5538 .dspark_real
5539 .iter()
5540 .collect::<std::collections::HashSet<_>>()
5541 .len();
5542 let (tu, tt, du, dt) = self.dspark_exp.iter().fold((0, 0, 0, 0), |a, b| {
5543 (a.0 + b.0, a.1 + b.1, a.2 + b.2, a.3 + b.3)
5544 });
5545 let m = self.dspark_exp.len().max(1);
5546 eprintln!(
5547 "DSpark: черновиков {}, принято в среднем {mean:.2} из {block} \
5548 (токенов за проход {:.2}), распределение {at:?}, выживание [{}]",
5549 self.dspark_hist.len(),
5550 mean + 1.0,
5551 surv.join(" ")
5552 );
5553 eprintln!(
5554 "DSpark: разных токенов {distinct} из {} (вырожденность), \
5555 эксперты ствол {}/{} на слой за {block} токенов, \
5556 черновик {}/{} за блок, draft {:.2} мс/блок",
5557 self.dspark_real.len(),
5558 tu / m,
5559 tt / m,
5560 du / m,
5561 dt / m,
5562 self.dspark_draft_ns as f64 / self.dspark_exp.len().max(1) as f64 / 1e6
5563 );
5564 }
5565 }
5566
5567 fn forward_layers_upto(
5568 &mut self,
5569 hidden: &[f32],
5570 position: usize,
5571 task_mask: Option<&TaskMask>,
5572 upto: Option<usize>,
5573 ) -> Vec<f32> {
5574 if let Some(plan) = self.gpu_plan.clone() {
5580 if upto.is_none() && plan.len() > 1 {
5581 let mut h = hidden.to_vec();
5582 for &(dev, from, upto_incl) in plan.iter() {
5583 h = crate::gpu::with_device(dev, || {
5584 self.forward_layers_span(&h, position, task_mask, from, Some(upto_incl))
5585 });
5586 }
5587 return h;
5588 }
5589 }
5590 self.forward_layers_span(hidden, position, task_mask, 0, upto)
5591 }
5592
5593 pub fn set_gpu_plan(&mut self, devices: Option<&[usize]>) -> Result<(), String> {
5598 self.set_gpu_plan_at(devices, None)
5599 }
5600
5601 pub fn set_gpu_plan_at(
5605 &mut self,
5606 devices: Option<&[usize]>,
5607 at: Option<usize>,
5608 ) -> Result<(), String> {
5609 let Some(devs) = devices.filter(|d| d.len() > 1) else {
5610 self.gpu_plan = None;
5611 return Ok(());
5612 };
5613 self.split_supported()?;
5614 let n = self.num_layers;
5615 if devs.len() > n {
5616 return Err(format!("{} devices for {n} layers", devs.len()));
5617 }
5618 if let Some(k) = at {
5619 if k == 0 || k >= n {
5620 return Err(format!("split at {k}: the model has {n} layers"));
5621 }
5622 if devs.len() == 2 {
5623 self.gpu_plan = Some(std::sync::Arc::new(vec![
5624 (devs[0], 0, k - 1),
5625 (devs[1], k, n - 1),
5626 ]));
5627 return Ok(());
5628 }
5629 return Err(format!(
5630 "an explicit split point takes exactly 2 devices, got {}",
5631 devs.len()
5632 ));
5633 }
5634 let per = n.div_ceil(devs.len());
5635 let mut plan = Vec::with_capacity(devs.len());
5636 let mut from = 0usize;
5637 for &d in devs {
5638 if from >= n {
5639 break;
5640 }
5641 let upto = (from + per - 1).min(n - 1);
5642 plan.push((d, from, upto));
5643 from = upto + 1;
5644 }
5645 self.gpu_plan = Some(std::sync::Arc::new(plan));
5646 Ok(())
5647 }
5648
5649 pub fn gpu_plan(&self) -> Option<Vec<(usize, usize, usize)>> {
5651 self.gpu_plan.as_ref().map(|p| p.as_ref().clone())
5652 }
5653
5654 fn forward_layers_span(
5660 &mut self,
5661 hidden: &[f32],
5662 position: usize,
5663 task_mask: Option<&TaskMask>,
5664 from: usize,
5665 upto: Option<usize>,
5666 ) -> Vec<f32> {
5667 debug_assert!(from == 0 || (self.dsv4.is_none() && self.g3n.is_none()));
5668 if let Some(b) = &mut self.dsv4 {
5674 let _ = (task_mask, upto);
5675 let token_id = hidden.first().copied().unwrap_or(0.0) as u32;
5676 let (g, layers, cfg, st) = (&b.0, &b.1, b.2, &mut b.3);
5677 st.pos = position;
5678 let mut logits = Vec::new();
5679 crate::dsv4::forward_token(
5680 g,
5681 layers,
5682 &cfg,
5683 st,
5684 token_id,
5685 &self.inv_freq,
5686 self.pool.as_deref(),
5687 &mut logits,
5688 );
5689 self.graph_logits = Some(logits);
5690 self.dspark_probe(position, token_id);
5691 return vec![0.0; self.hidden_size];
5694 }
5695 if let Some(b) = &self.g3n {
5698 let _ = (task_mask, upto);
5699 return crate::g3n::g3n_forward(
5700 &b.0,
5701 &b.1,
5702 hidden,
5703 position,
5704 &mut self.kv_cache.layers,
5705 self.num_heads,
5706 self.num_kv_heads,
5707 self.head_dim,
5708 self.pool.as_deref(),
5709 );
5710 }
5711 let mut h = hidden.to_vec();
5712 let (nh, _nkv, _hd, hs, _rd, eps) = (
5715 self.num_heads,
5716 self.num_kv_heads,
5717 self.head_dim,
5718 self.hidden_size,
5719 self.rotary_dim,
5720 self.rms_eps,
5721 );
5722 let pool = self.pool.clone();
5723 let graph_env = std::env::var("CMF_GPU_WGPU_GRAPH").ok();
5735 let graph_on = match graph_env.as_deref() {
5736 Some("0") => false,
5737 Some(_) => true,
5738 None => crate::gpu::wgpu_graph_default(),
5744 };
5745 let graph_trusted =
5746 graph_env.is_some() || crate::gpu::wgpu_graph_default() || self.gdn_cfg.is_some();
5747 let race_eligible = graph_on && upto.is_none() && task_mask.is_none() && from == 0;
5748 let mut tail_start = 0usize;
5749 if race_eligible && crate::gpu::graph_race_use_graph(graph_trusted) {
5750 let t_graph = std::time::Instant::now();
5751 let mut lg = Vec::new();
5752 let mut gl = 0usize;
5753 let built = self.try_token_graph_wgpu(hidden, position, &mut lg, &mut gl);
5754 graph_note(built.is_some());
5755 if let Some(hh) = built {
5756 let dur = t_graph.elapsed();
5757 if std::env::var("CMF_GRAPH_PROF").is_ok() {
5758 eprintln!("graph-call: {:.2} ms total", dur.as_secs_f64() * 1000.0);
5759 }
5760 if gl > 0 && gl < self.num_layers {
5761 h = hh;
5767 tail_start = gl;
5768 } else if graph_trusted || !crate::gpu::graph_race_first_token_hopeless(dur) {
5769 if !graph_trusted {
5770 crate::gpu::graph_race_record(true, dur);
5771 }
5772 if !lg.is_empty() {
5773 lg.resize(self.vocab_size, 0.0);
5776 if let Some(c) = self.final_softcap {
5777 for l in lg.iter_mut() {
5778 *l = c * (*l / c).tanh();
5779 }
5780 }
5781 self.graph_logits = Some(lg);
5782 }
5783 return hh;
5784 }
5785 }
5791 }
5792 let span = from > 0 || upto.is_some();
5816 if span && graph_on && task_mask.is_none() && graph_trusted {
5817 let upto_excl = upto.map_or(self.num_layers, |u| u + 1);
5818 let mut lg = Vec::new();
5819 let mut gl = 0usize;
5820 let span_res =
5821 self.try_token_graph_wgpu_span(hidden, position, &mut lg, from, upto_excl, &mut gl);
5822 graph_note(span_res.is_some() && gl == upto_excl - from);
5823 if std::env::var("CMF_GPU_DEBUG").is_ok() {
5824 static SEEN: std::sync::atomic::AtomicU32 =
5828 std::sync::atomic::AtomicU32::new(0);
5829 if SEEN.fetch_add(1, std::sync::atomic::Ordering::Relaxed) < 4 {
5830 eprintln!(
5831 "span graph: covered {gl} of {} layers [{from}..{upto_excl}) res={}",
5832 upto_excl - from,
5833 span_res.is_some()
5834 );
5835 }
5836 }
5837 if let Some(hh) = span_res {
5838 if gl == upto_excl - from {
5839 if !lg.is_empty() {
5840 lg.resize(self.vocab_size, 0.0);
5841 if let Some(c) = self.final_softcap {
5842 for l in lg.iter_mut() {
5843 *l = c * (*l / c).tanh();
5844 }
5845 }
5846 self.graph_logits = Some(lg);
5847 }
5848 crate::gpu::set_layer(-1);
5849 return hh;
5850 }
5851 h = hh;
5853 tail_start = from + gl;
5854 }
5855 }
5856 let t_race_cpu = (race_eligible && !graph_trusted).then(std::time::Instant::now);
5857
5858 #[cfg(target_os = "macos")]
5859 let mut gpu_skip_until = 0usize;
5860 for li in tail_start.max(from)..self.num_layers {
5861 crate::gpu::set_layer(li as i64); if let Some(u) = upto {
5863 if li > u {
5864 break;
5865 }
5866 }
5867 if let Some(mask) = task_mask {
5868 if !mask.layer_alive(li) {
5869 continue; }
5871 }
5872 #[cfg(target_os = "macos")]
5876 {
5877 if li < gpu_skip_until {
5878 continue;
5879 }
5880 if task_mask.is_none() {
5881 let end = self.q1_graph_gpu(li, upto, position, &mut h);
5882 if end > li {
5883 gpu_skip_until = end;
5884 if self.is_loop_end(end - 1) && end < self.num_layers {
5887 h = inference::rms_norm(
5888 &h,
5889 &self.weights.final_norm,
5890 self.rms_eps,
5891 self.norm_style,
5892 );
5893 }
5894 continue;
5895 }
5896 }
5897 }
5898
5899 let lw = &self.weights.layers[self.phys_layer(li)];
5900 if let Ok(tp) = std::env::var("CMF_TRACE_POS") {
5901 if tp.parse::<usize>().ok() == Some(position) {
5902 let n: f32 = h.iter().map(|x| x * x).sum::<f32>().sqrt();
5903 eprintln!(
5904 "TRACE pos {position} layer {li}: |h| = {n:.6} h0 {:.6} h1 {:.6}",
5905 h[0], h[1]
5906 );
5907 }
5908 }
5909 inference::rms_norm_into(
5912 &h,
5913 &lw.input_norm,
5914 self.rms_eps,
5915 self.norm_style,
5916 &mut self.ws.n1,
5917 );
5918
5919 let attn_out = match &lw.attn {
5920 AttnKind::Mla(w) => {
5921 let inv_freq_l = self.layer_inv_freq(li);
5922 let rs = self.layer_rope_scale(li);
5923 let eps = self.rms_eps;
5924 let pool = self.pool.clone();
5925 mla_attention(
5926 w,
5927 &self.ws.n1,
5928 &mut self.kv_cache.layers[li],
5929 position,
5930 &inv_freq_l,
5931 rs,
5932 eps,
5933 pool.as_deref(),
5934 )
5935 }
5936 AttnKind::Linear(w) => {
5937 let cfg = self.vmf_cfg.expect("linear layer without vmf_cfg");
5938 vmf_phase_forward(
5939 &self.ws.n1,
5940 w,
5941 &cfg,
5942 &mut self.kv_cache.layers[li].linear_state,
5943 self.pool.as_deref(),
5944 )
5945 }
5946 AttnKind::Kda(w) => {
5947 let cfg = self.kda_cfg.expect("kda layer without kda_cfg");
5948 crate::linear_core::kda_forward(
5949 &self.ws.n1,
5950 w,
5951 &cfg,
5952 &mut self.kv_cache.layers[li].linear_state,
5953 self.pool.as_deref(),
5954 )
5955 }
5956 AttnKind::LinearGdn(w) => {
5957 let cfg = self.gdn_cfg.expect("gdn layer without gdn_cfg");
5958 gdn_forward(
5959 &self.ws.n1,
5960 w,
5961 &cfg,
5962 &mut self.kv_cache.layers[li].linear_state,
5963 self.pool.as_deref(),
5964 )
5965 }
5966 AttnKind::ShortConv(w) => {
5967 let cfg = self
5968 .short_conv_cfg
5969 .expect("short-conv layer without short_conv_cfg");
5970 short_conv_forward(
5971 &self.ws.n1,
5972 w,
5973 &cfg,
5974 &mut self.kv_cache.layers[li].linear_state,
5975 self.pool.as_deref(),
5976 )
5977 }
5978 AttnKind::Full {
5979 wq,
5980 wk,
5981 wv,
5982 wo,
5983 q_norm,
5984 k_norm,
5985 output_gate,
5986 softplus_gate,
5987 bias,
5988 } if self.kv_cache.layers[li].o1_sealed() => {
5989 let inv_freq_l = self.layer_inv_freq(li);
5992 let (nkv_l, hd_l, rd_l) = self.layer_geom(li);
5993 let cfg = QwenAttnCfg {
5994 num_heads: self.layer_num_heads(li),
5995 num_kv_heads: nkv_l,
5996 head_dim: hd_l,
5997 hidden_size: hs,
5998 position,
5999 inv_freq: &inv_freq_l,
6000 rotary_dim: rd_l,
6001 scale: self.attn_scale,
6002 softcap: self.attn_softcap,
6003 window: None,
6004 v_norm: self.attn_v_norm,
6005 q_norm: q_norm.as_deref(),
6006 k_norm: k_norm.as_deref(),
6007 output_gate: *output_gate,
6008 softplus_gate: softplus_gate
6009 .as_ref()
6010 .map(|(gate, per_head)| (gate, *per_head)),
6011 rope_scale: self.layer_rope_scale(li),
6012 bias: bias
6013 .as_ref()
6014 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
6015 rms_eps: eps,
6016 norm_style: self.norm_style,
6017 pool: pool.as_deref(),
6018 };
6019 attention::qwen_attention_nystrom(
6020 &self.ws.n1,
6021 wq,
6022 wk,
6023 wv,
6024 wo,
6025 &mut self.kv_cache.layers[li],
6026 &cfg,
6027 )
6028 }
6029 AttnKind::Full {
6030 wq,
6031 wk,
6032 wv,
6033 wo,
6034 q_norm,
6035 k_norm,
6036 output_gate,
6037 softplus_gate,
6038 bias,
6039 } => 'attn: {
6040 if graph_on
6043 && !*output_gate
6044 && softplus_gate.is_none()
6045 && self.attention_heads_per_layer.is_none()
6046 && bias.is_none()
6047 && task_mask.is_none()
6048 {
6049 let inv_freq_l = self.layer_inv_freq(li);
6050 let (nkv_l, hd_l, rd_l) = self.layer_geom(li);
6051 let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
6052 if let (Some((gm, qi)), Some((_, ki)), Some((_, vi)), Some((_, oi))) = (
6053 wq.mapped_q1(),
6054 wk.mapped_q1(),
6055 wv.mapped_q1(),
6056 wo.mapped_q1(),
6057 ) {
6058 let gm = gm.clone();
6059 let mut out = vec![0f32; hs];
6060 let cache = &self.kv_cache.layers[li];
6061 if crate::gpu::attn_dropin(
6062 &gm,
6063 self.graph_kv_id,
6064 li,
6065 &self.ws.n1,
6066 qi,
6067 ki,
6068 vi,
6069 oi,
6070 q_norm.as_deref(),
6071 k_norm.as_deref(),
6072 &inv_freq_l,
6073 nh,
6074 nkv_l,
6075 hd_l,
6076 rd_l,
6077 hs,
6078 position,
6079 self.kv_cache.max_seq_len,
6080 gemma,
6081 eps as f32,
6082 cache.k_heads(),
6083 cache.v_heads(),
6084 &mut out,
6085 ) {
6086 break 'attn out;
6087 }
6088 }
6089 }
6090 let masked = task_mask
6091 .map(|m| m.head_flags(li, self.num_heads).iter().any(|&a| !a))
6092 .unwrap_or(false);
6093 let f32_view = (wq.as_f32(), wk.as_f32(), wv.as_f32(), wo.as_f32());
6094 match (masked, f32_view) {
6095 (true, (Some(q), Some(k), Some(v), Some(o))) => {
6098 let active_heads = task_mask.unwrap().head_flags(li, self.num_heads);
6099 attention::multi_head_attention(
6100 &self.ws.n1,
6101 q,
6102 k,
6103 v,
6104 o,
6105 &mut self.kv_cache.layers[li],
6106 self.num_heads,
6107 self.num_kv_heads,
6108 self.head_dim,
6109 self.hidden_size,
6110 position,
6111 &active_heads,
6112 &self.inv_freq,
6113 )
6114 }
6115 (masked, _) => {
6116 if masked {
6117 tracing::warn!(
6118 "layer {li}: head mask on quantized weights not \
6119 supported yet — executing dense"
6120 );
6121 }
6122 let inv_freq_l = self.layer_inv_freq(li);
6123 let (nkv_l, hd_l, rd_l) = self.layer_geom(li);
6124 let cfg = QwenAttnCfg {
6125 num_heads: self.layer_num_heads(li),
6126 num_kv_heads: nkv_l,
6127 head_dim: hd_l,
6128 hidden_size: hs,
6129 position,
6130 inv_freq: &inv_freq_l,
6131 rotary_dim: rd_l,
6132 scale: self.attn_scale,
6133 softcap: self.attn_softcap,
6134 window: self.layer_window(li),
6135 v_norm: self.attn_v_norm,
6136 q_norm: q_norm.as_deref(),
6137 k_norm: k_norm.as_deref(),
6138 output_gate: *output_gate,
6139 softplus_gate: softplus_gate
6140 .as_ref()
6141 .map(|(gate, per_head)| (gate, *per_head)),
6142 rope_scale: self.layer_rope_scale(li),
6143 bias: bias
6144 .as_ref()
6145 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
6146 rms_eps: eps,
6147 norm_style: self.norm_style,
6148 pool: pool.as_deref(),
6149 };
6150 attention::qwen_attention(
6151 &self.ws.n1,
6152 wq,
6153 wk,
6154 wv,
6155 wo,
6156 &mut self.kv_cache.layers[li],
6157 &cfg,
6158 )
6159 }
6160 }
6161 }
6162 };
6163 let attn_out = match &self.weights.layers[self.phys_layer(li)].attn_out_norm {
6166 Some(w) => inference::rms_norm(&attn_out, w, self.rms_eps, self.norm_style),
6167 None => attn_out,
6168 };
6169 let lw = &self.weights.layers[self.phys_layer(li)];
6170 inference::add_rmsnorm_fused_into(
6171 &mut h,
6172 &attn_out,
6173 &lw.post_norm,
6174 self.rms_eps,
6175 self.norm_style,
6176 &mut self.ws.p1,
6177 );
6178 let mut attn_out = attn_out;
6179 attention::recycle_buf(&mut attn_out);
6180 let post_normed = &self.ws.p1;
6181
6182 let ffn_masked = task_mask
6183 .map(|m| m.ffn_active_count(li) < self.intermediate_size)
6184 .unwrap_or(false);
6185 let ffn_out = match (ffn_masked, &lw.ffn) {
6197 (true, FfnKind::Dense(d)) => {
6198 let tm = task_mask.unwrap();
6199 let alive = tm.ffn_active_count(li);
6200 let deep = alive * 2 <= self.intermediate_size;
6201 if deep && d.down_proj.sparse_col_ok() {
6202 let active = tm.ffn_active_indices(li);
6203 sparse_ffn_quant(
6204 d,
6205 post_normed,
6206 &active,
6207 self.hidden_size,
6208 self.pool.as_deref(),
6209 )
6210 } else if deep
6211 && let (Some(g), Some(u), Some(dn)) = (
6212 d.gate_proj.as_f32(),
6213 d.up_proj.as_f32(),
6214 d.down_proj.as_f32(),
6215 )
6216 {
6217 let active = tm.ffn_active_indices(li);
6218 inference::sparse_ffn_forward(
6219 post_normed,
6220 g,
6221 u,
6222 dn,
6223 self.hidden_size,
6224 self.intermediate_size,
6225 &active,
6226 self.pool.as_deref(),
6227 )
6228 } else {
6229 let row = tm.ffn_masks.get(li).map(|v| v.as_slice());
6230 dense_ffn_batch(d, post_normed, 1, self.pool.as_deref(), row)
6231 }
6232 }
6233 (true, FfnKind::Moe(m)) => {
6234 let allowed = task_mask.and_then(|tm| tm.expert_flags(li, m.experts.len()));
6238 ffn_forward(
6239 &lw.ffn,
6240 post_normed,
6241 self.pool.as_deref(),
6242 allowed.as_deref(),
6243 )
6244 }
6245 (true, FfnKind::DenseMoe(dm)) => dense_moe_ffn(
6246 dm,
6247 post_normed,
6248 &h,
6249 self.rms_eps,
6250 self.norm_style,
6251 self.pool.as_deref(),
6252 ),
6253 (false, _) => match &lw.ffn {
6254 FfnKind::DenseMoe(dm) => dense_moe_ffn(
6255 dm,
6256 post_normed,
6257 &h,
6258 self.rms_eps,
6259 self.norm_style,
6260 self.pool.as_deref(),
6261 ),
6262 _ => {
6263 let allowed = match (&lw.ffn, task_mask) {
6264 (FfnKind::Moe(m), Some(tm)) => tm.expert_flags(li, m.experts.len()),
6265 _ => None,
6266 };
6267 ffn_forward(
6268 &lw.ffn,
6269 post_normed,
6270 self.pool.as_deref(),
6271 allowed.as_deref(),
6272 )
6273 }
6274 },
6275 };
6276 let ffn_out = match &self.weights.layers[self.phys_layer(li)].ffn_out_norm {
6277 Some(w) => inference::rms_norm(&ffn_out, w, self.rms_eps, self.norm_style),
6278 None => ffn_out,
6279 };
6280 for (i, &f) in ffn_out.iter().enumerate() {
6281 h[i] += f;
6282 }
6283 let mut ffn_out = ffn_out;
6284 attention::recycle_buf(&mut ffn_out);
6285
6286 if let Some(sc) = self.weights.layers[self.phys_layer(li)].layer_scale {
6288 for v in h.iter_mut() {
6289 *v *= sc;
6290 }
6291 }
6292
6293 if self.is_loop_end(li) && li + 1 < self.num_layers {
6296 h = inference::rms_norm(
6297 &h,
6298 &self.weights.final_norm,
6299 self.rms_eps,
6300 self.norm_style,
6301 );
6302 }
6303
6304 if self.dyn_phi_layer == Some(li) {
6308 self.update_dyn_phi(&h);
6309 }
6310 }
6311 crate::gpu::set_layer(-1); if let Some(t) = t_race_cpu {
6313 crate::gpu::graph_race_record(false, t.elapsed());
6314 }
6315
6316 h
6317 }
6318
6319 fn update_dyn_phi(&mut self, h: &[f32]) {
6322 const A: f32 = 0.2;
6323 if self.dyn_phi_ema.len() != h.len() {
6324 self.dyn_phi_ema = vec![0.0; h.len()];
6325 self.dyn_phi_seen = 0;
6326 }
6327 if self.dyn_phi_seen == 0 {
6328 self.dyn_phi_ema.copy_from_slice(h);
6329 } else {
6330 for (e, &v) in self.dyn_phi_ema.iter_mut().zip(h) {
6331 *e = (1.0 - A) * *e + A * v;
6332 }
6333 }
6334 self.dyn_phi_seen += 1;
6335 }
6336
6337 pub fn dyn_phi(&self) -> &[f32] {
6339 &self.dyn_phi_ema
6340 }
6341
6342 pub fn set_dyn_phi_layer(&mut self, layer: Option<usize>) {
6344 self.dyn_phi_layer = layer;
6345 self.dyn_phi_ema.clear();
6346 self.dyn_phi_seen = 0;
6347 }
6348
6349 pub fn dynamic_skills(&self) -> Vec<(usize, String, usize)> {
6351 let Some(model) = &self.model else {
6352 return Vec::new();
6353 };
6354 model
6355 .header
6356 .skills
6357 .iter()
6358 .enumerate()
6359 .filter_map(|(i, sk)| {
6360 let ok = matches!(self.dyn_skill_layers.get(i), Some(Some(_)));
6361 let sel = sk.selection.as_ref()?;
6362 (ok).then(|| (i, sk.id.clone(), sel.phi_layer))
6363 })
6364 .collect()
6365 }
6366
6367 pub fn active_skill(&self) -> Option<usize> {
6369 self.dyn_active
6370 }
6371
6372 pub fn enable_dynamic_routing(&mut self) -> usize {
6377 use crate::swarm::{DynRouter, RoutableSkill};
6378 let Some(model) = self.model.clone() else {
6379 return 0;
6380 };
6381 if self.dyn_blend_loaded {
6384 tracing::warn!("dynamic routing unavailable on a blend-loaded pipeline");
6385 return 0;
6386 }
6387 if let Some(a) = self.dyn_active {
6391 if !matches!(self.dyn_skill_layers.get(a), Some(Some(_))) {
6392 tracing::warn!("loaded skill is not FFN-eligible — dynamic routing unavailable");
6393 return 0;
6394 }
6395 }
6396 let hidden = self.hidden_size;
6397 let mut skills = Vec::new();
6398 for (idx, id, _phi) in self.dynamic_skills() {
6399 if let Some(sel) = model.header.skills[idx].selection.as_ref() {
6400 if let Some(rs) = RoutableSkill::from_descriptor(idx, id, sel, hidden) {
6401 skills.push(rs);
6402 }
6403 }
6404 }
6405 if skills.is_empty() {
6406 return 0;
6407 }
6408 let phi = skills[0].phi_layer;
6410 if skills.iter().any(|s| s.phi_layer != phi) {
6411 tracing::warn!("routable skills disagree on phi_layer; using {phi}");
6412 }
6413 let n = skills.len();
6414 self.set_dyn_phi_layer(Some(phi));
6415 self.dyn_router = Some(DynRouter::new(skills));
6416 n
6417 }
6418
6419 pub fn route_switches(&self) -> Vec<(usize, Option<String>, Option<String>)> {
6421 self.dyn_router
6422 .as_ref()
6423 .map(|r| r.switches.clone())
6424 .unwrap_or_default()
6425 }
6426
6427 fn lm_head_forward(&self, hidden: &[f32]) -> Vec<f32> {
6430 let rows = self.weights.lm_head.rows();
6431 let mut logits = attention::take_buf(rows.min(self.vocab_size));
6432 self.weights
6433 .lm_head
6434 .matvec(hidden, &mut logits, self.pool.as_deref());
6435 logits.resize(self.vocab_size, 0.0);
6436 if let Some(m) = self.logit_multiplier {
6437 for l in logits.iter_mut() {
6438 *l *= m;
6439 }
6440 }
6441 if let Some(c) = self.final_softcap {
6442 for l in logits.iter_mut() {
6443 *l = c * (*l / c).tanh();
6444 }
6445 }
6446 logits
6447 }
6448
6449 pub fn prefill_next_logits(&mut self, ids: &[u32], task_mask: Option<&TaskMask>) -> Vec<f32> {
6454 self.kv_cache.clear();
6455 self.kv_history.clear();
6456 let mut hidden = vec![0.0f32; self.hidden_size];
6457 for (pos, &id) in ids.iter().enumerate() {
6458 let emb = self.embed_single(id);
6459 hidden = self.forward_layers(&emb, pos, task_mask);
6460 }
6461 inference::rms_norm_into(
6462 &hidden,
6463 &self.weights.final_norm,
6464 self.rms_eps,
6465 self.norm_style,
6466 &mut self.ws.n1,
6467 );
6468 self.lm_head_forward(&self.ws.n1)
6469 }
6470}
6471
6472pub fn create_test_pipeline(
6474 hidden_size: usize,
6475 intermediate_size: usize,
6476 num_heads: usize,
6477 num_kv_heads: usize,
6478 head_dim: usize,
6479 num_layers: usize,
6480 vocab_size: usize,
6481) -> Pipeline {
6482 let synth = |n: usize, salt: usize| -> Vec<f32> {
6485 (0..n)
6486 .map(|i| (((i * 31 + salt * 17 + 7) % 97) as f32 / 97.0 - 0.5) * 0.2)
6487 .collect()
6488 };
6489 let qt = |rows: usize, cols: usize, salt: usize| -> QTensor {
6490 QTensor::from_f32(synth(rows * cols, salt), rows, cols)
6491 };
6492 let layer_weights: Vec<LayerWeights> = (0..num_layers)
6493 .map(|li| LayerWeights {
6494 input_norm: vec![1.0; hidden_size],
6495 post_norm: vec![1.0; hidden_size],
6496 attn_out_norm: None,
6497 ffn_out_norm: None,
6498 layer_scale: None,
6499 ffn: FfnKind::Dense(DenseFfn {
6500 gate_proj: qt(intermediate_size, hidden_size, li * 10 + 5),
6501 up_proj: qt(intermediate_size, hidden_size, li * 10 + 6),
6502 down_proj: qt(hidden_size, intermediate_size, li * 10 + 7),
6503 act: Act::Silu,
6504 }),
6505 attn: AttnKind::Full {
6506 bias: None,
6507 wq: qt(num_heads * head_dim, hidden_size, li * 10 + 1),
6508 wk: qt(num_kv_heads * head_dim, hidden_size, li * 10 + 2),
6509 wv: qt(num_kv_heads * head_dim, hidden_size, li * 10 + 3),
6510 wo: qt(hidden_size, num_heads * head_dim, li * 10 + 4),
6511 q_norm: None,
6512 k_norm: None,
6513 output_gate: false,
6514 softplus_gate: None,
6515 },
6516 })
6517 .collect();
6518
6519 Pipeline::new(
6520 Tokenizer::byte_level(),
6521 PipelineWeights {
6522 embed_tokens: qt(vocab_size, hidden_size, 100),
6523 layers: layer_weights,
6524 lm_head: qt(vocab_size, hidden_size, 200),
6525 final_norm: vec![1.0; hidden_size],
6526 },
6527 hidden_size,
6528 intermediate_size,
6529 num_heads,
6530 num_kv_heads,
6531 head_dim,
6532 num_layers,
6533 num_layers, false, vocab_size,
6536 1e-6,
6537 10_000.0,
6538 NormStyle::Qwen,
6539 4096,
6540 SamplerConfig {
6541 seed: Some(42),
6542 ..Default::default()
6543 },
6544 )
6545}
6546
6547#[inline]
6552fn mask_bit(row: &[u8], j: usize) -> bool {
6553 (row.get(j >> 3).copied().unwrap_or(0) >> (j & 7)) & 1 != 0
6554}
6555
6556fn zero_masked_cols(g: &mut [f32], rows: usize, inter: usize, row: &[u8]) {
6562 for r in 0..rows {
6563 let base = r * inter;
6564 for (bi, &byte) in row.iter().enumerate() {
6565 if byte == 0xFF {
6566 continue;
6567 }
6568 let j0 = bi * 8;
6569 for bit in 0..8 {
6570 let j = j0 + bit;
6571 if j < inter && byte & (1 << bit) == 0 {
6572 g[base + j] = 0.0;
6573 }
6574 }
6575 }
6576 }
6577}
6578
6579fn dense_ffn_batch(
6580 d: &DenseFfn,
6581 xs: &[f32],
6582 b: usize,
6583 pool: Option<&Pool>,
6584 mask_row: Option<&[u8]>,
6585) -> Vec<f32> {
6586 let inter = d.gate_proj.rows();
6587 let hidden = d.down_proj.rows();
6588 if mask_row.is_none()
6596 && d.act == Act::Silu
6597 && b >= 32
6598 && crate::gpu::enabled_here()
6599 && !crate::gpu::mm_killed()
6600 {
6601 if let (Some((model, w1)), Some((_, w3)), Some((_, w2))) = (
6602 d.gate_proj.mapped_q4t(),
6603 d.up_proj.mapped_q4t(),
6604 d.down_proj.mapped_q4t(),
6605 ) {
6606 let mut out = vec![0.0f32; b * hidden];
6607 if crate::gpu::q4t_ffn(model, w1, w3, w2, xs, b, hidden, inter, &mut out) {
6608 return out;
6609 }
6610 }
6611 if let (Some((model, w1)), Some((_, w3)), Some((_, w2))) = (
6616 d.gate_proj.mapped_q4tp(),
6617 d.up_proj.mapped_q4tp(),
6618 d.down_proj.mapped_q4tp(),
6619 ) {
6620 let mut out = vec![0.0f32; b * hidden];
6621 if crate::gpu::q4tp_ffn(model, w1, w3, w2, xs, b, hidden, inter, &mut out) {
6622 return out;
6623 }
6624 }
6625 }
6626 let mut g = vec![0.0f32; b * inter];
6627 d.gate_proj.matmat(xs, b, &mut g, pool);
6628 let mut u = vec![0.0f32; b * inter];
6629 d.up_proj.matmat(xs, b, &mut u, pool);
6630 for i in 0..b * inter {
6631 g[i] = d.act.combine(g[i], u[i]);
6632 }
6633 if let Some(row) = mask_row {
6634 zero_masked_cols(&mut g, b, inter, row);
6635 }
6636 let mut out = vec![0.0f32; b * hidden];
6637 d.down_proj.matmat(&g, b, &mut out, pool);
6638 out
6639}
6640
6641fn accumulate_act(m: &MoeFfn, xs: &[f32], b: usize) {
6646 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
6647 static DUMP: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
6648 let on = *ON.get_or_init(|| std::env::var("CMF_RMS_TRACE").is_ok());
6649 let dump = *DUMP.get_or_init(|| std::env::var("CMF_ACT_DUMP").is_ok());
6650 if (!on && !dump) || b == 0 {
6651 return;
6652 }
6653 let hidden = xs.len() / b;
6654 if on {
6655 let mut acc = m.act_sq.borrow_mut();
6656 if acc.len() < hidden {
6657 acc.resize(hidden, 0.0);
6658 }
6659 for t in 0..b {
6660 let row = &xs[t * hidden..(t + 1) * hidden];
6661 for (a, &v) in acc.iter_mut().zip(row) {
6662 *a += (v as f64) * (v as f64);
6663 }
6664 }
6665 }
6666 if dump {
6667 let cap: usize = std::env::var("CMF_ACT_DUMP_ROWS")
6670 .ok()
6671 .and_then(|v| v.parse().ok())
6672 .unwrap_or(4096);
6673 let mut rows = m.act_rows.borrow_mut();
6674 if rows.len() < cap * hidden {
6675 let take = b.min((cap * hidden - rows.len()) / hidden.max(1));
6676 rows.extend_from_slice(&xs[..take * hidden]);
6677 }
6678 }
6679}
6680
6681#[derive(Clone, Copy)]
6684struct SendVecs(*mut Vec<f32>);
6685unsafe impl Send for SendVecs {}
6686unsafe impl Sync for SendVecs {}
6687impl SendVecs {
6688 #[inline]
6689 fn at(self, i: usize) -> *mut Vec<f32> {
6690 unsafe { self.0.add(i) }
6691 }
6692}
6693
6694fn moe_ffn_batch(
6695 m: &MoeFfn,
6696 xs: &[f32],
6697 b: usize,
6698 hidden: usize,
6699 pool: Option<&Pool>,
6700 allowed: Option<&[bool]>,
6701) -> Vec<f32> {
6702 accumulate_act(m, xs, b);
6703 let ne = m.experts.len();
6704 let mut logits = vec![0.0f32; b * ne];
6705 m.router.matmat(xs, b, &mut logits, pool);
6706
6707 let mut assign: Vec<Vec<(usize, f32)>> = vec![Vec::new(); ne];
6710 {
6711 let mut st = m.stats.borrow_mut();
6712 if st.len() < ne {
6713 st.resize(ne, 0);
6714 }
6715 for bi in 0..b {
6716 let (idx, p, wsum) = moe_route(&logits[bi * ne..(bi + 1) * ne], m, allowed);
6717 for &e in &idx {
6718 st[e] += 1;
6719 assign[e].push((bi, p[e] / wsum));
6720 }
6721 }
6722 }
6723
6724 let mut out = vec![0.0f32; b * hidden];
6725 let cols = m.experts[0].gate_proj.cols();
6726 let run_expert = |d: &DenseFfn, list: &[(usize, f32)], out: &mut [f32]| {
6727 let sb = list.len();
6728 let mut sub = vec![0.0f32; sb * cols];
6729 for (k, &(bi, _)) in list.iter().enumerate() {
6730 sub[k * cols..(k + 1) * cols].copy_from_slice(&xs[bi * cols..(bi + 1) * cols]);
6731 }
6732 let eo = dense_ffn_batch(d, &sub, sb, pool, None);
6733 for (k, &(bi, w)) in list.iter().enumerate() {
6734 for i in 0..hidden {
6735 out[bi * hidden + i] += w * eo[k * hidden + i];
6736 }
6737 }
6738 };
6739 let active: Vec<usize> = (0..ne).filter(|&e| !assign[e].is_empty()).collect();
6745 if pool.is_some() && active.len() >= 8 {
6746 let mut panels: Vec<Vec<f32>> = vec![Vec::new(); active.len()];
6747 {
6748 let panel_ptr = SendVecs(panels.as_mut_ptr());
6749 let experts = &m.experts;
6752 let (active_r, assign_r) = (&active, &assign);
6753 let run = |start: usize, end: usize| {
6754 for ai in start..end {
6755 let e = active_r[ai];
6756 let list = &assign_r[e];
6757 let sb = list.len();
6758 let mut sub = vec![0.0f32; sb * cols];
6759 for (k, &(bi, _)) in list.iter().enumerate() {
6760 sub[k * cols..(k + 1) * cols]
6761 .copy_from_slice(&xs[bi * cols..(bi + 1) * cols]);
6762 }
6763 unsafe {
6765 *panel_ptr.at(ai) =
6766 dense_ffn_batch(&experts[e], &sub, sb, None, None);
6767 }
6768 }
6769 };
6770 match pool {
6771 Some(p) => p.run_rows(active.len(), &run),
6772 None => run(0, active.len()),
6773 }
6774 }
6775 for (ai, &e) in active.iter().enumerate() {
6776 for (k, &(bi, w)) in assign[e].iter().enumerate() {
6777 let eo = &panels[ai][k * hidden..(k + 1) * hidden];
6778 for i in 0..hidden {
6779 out[bi * hidden + i] += w * eo[i];
6780 }
6781 }
6782 }
6783 } else {
6784 for &e in &active {
6785 run_expert(&m.experts[e], &assign[e], &mut out);
6786 }
6787 }
6788 if let Some((se, gate)) = &m.shared {
6789 let all: Vec<(usize, f32)> = if let Some(gate) = gate {
6790 let mut gl = vec![0.0f32; b];
6791 gate.matmat(xs, b, &mut gl, pool);
6792 (0..b)
6793 .map(|bi| (bi, 1.0 / (1.0 + (-gl[bi]).exp())))
6794 .collect()
6795 } else {
6796 (0..b).map(|bi| (bi, 1.0)).collect()
6797 };
6798 run_expert(se, &all, &mut out);
6799 }
6800 out
6801}
6802
6803thread_local! {
6804 static FFN_SCRATCH: std::cell::RefCell<[Vec<f32>; 4]> =
6808 const { std::cell::RefCell::new([Vec::new(), Vec::new(), Vec::new(), Vec::new()]) };
6809}
6810
6811fn dense_ffn(d: &DenseFfn, x: &[f32], pool: Option<&Pool>) -> Vec<f32> {
6813 if crate::gpu::enabled_here()
6824 && (d.gate_proj.rows() >= crate::gpu::min_rows() || d.gate_proj.is_q1())
6825 {
6826 let arm = if d.gate_proj.is_q1() && crate::gpu::q1_force() {
6827 crate::gpu::ProbeArm::Gpu
6828 } else {
6829 crate::gpu::probe_arm(crate::gpu::OpClass::Ffn)
6830 };
6831 match arm {
6832 crate::gpu::ProbeArm::Gpu => {
6833 let t0 = std::time::Instant::now();
6834 if let Some(out) = dense_ffn_gpu(d, x, pool) {
6835 crate::gpu::probe_record(crate::gpu::OpClass::Ffn, true, t0.elapsed());
6836 return out;
6837 }
6838 }
6839 crate::gpu::ProbeArm::CpuTimed => {
6840 let t0 = std::time::Instant::now();
6841 let out = crate::gpu::cpu_scope(|| dense_ffn_cpu(d, x, pool));
6842 crate::gpu::probe_record(crate::gpu::OpClass::Ffn, false, t0.elapsed());
6843 return out;
6844 }
6845 crate::gpu::ProbeArm::Cpu => {
6846 return crate::gpu::cpu_scope(|| dense_ffn_cpu(d, x, pool));
6847 }
6848 }
6849 }
6850 dense_ffn_cpu(d, x, pool)
6851}
6852
6853fn dense_ffn_cpu(d: &DenseFfn, x: &[f32], pool: Option<&Pool>) -> Vec<f32> {
6855 let inter = d.gate_proj.rows();
6856 FFN_SCRATCH.with(|s| {
6857 let mut s = s.borrow_mut();
6858 let [g, u, ..] = &mut *s;
6859 g.resize(inter, 0.0);
6860 if d.act == Act::Silu && QTensor::matvec_silu_mul(&d.gate_proj, &d.up_proj, x, g, pool) {
6863 } else {
6865 u.resize(inter, 0.0);
6866 QTensor::matvec_many([&d.gate_proj, &d.up_proj], x, [g, u], pool);
6868 for i in 0..inter {
6869 g[i] = d.act.combine(g[i], u[i]);
6870 }
6871 }
6872 FFN_PROBE.with(|pr| {
6875 if let Some(acc) = pr.borrow_mut().as_mut() {
6876 let li = crate::gpu::cur_layer();
6877 if li >= 0 {
6878 if let Some(row) = acc.get_mut(li as usize) {
6879 for (a, &v) in row.iter_mut().zip(g.iter()) {
6880 *a += (v as f64).abs();
6881 }
6882 }
6883 }
6884 }
6885 });
6886 let mut out = attention::take_buf(d.down_proj.rows());
6887 d.down_proj.matvec(g, &mut out, pool);
6888 out
6889 })
6890}
6891
6892thread_local! {
6893 static FFN_PROBE: std::cell::RefCell<Option<Vec<Vec<f64>>>> =
6896 const { std::cell::RefCell::new(None) };
6897}
6898
6899fn dense_ffn_masked(
6904 d: &DenseFfn,
6905 x: &[f32],
6906 pool: Option<&Pool>,
6907 mask_row: &[u8],
6908) -> Vec<f32> {
6909 let inter = d.gate_proj.rows();
6910 FFN_SCRATCH.with(|s| {
6911 let mut s = s.borrow_mut();
6912 let [g, u, ..] = &mut *s;
6913 g.resize(inter, 0.0);
6914 if d.act == Act::Silu && QTensor::matvec_silu_mul(&d.gate_proj, &d.up_proj, x, g, pool) {
6915 } else {
6917 u.resize(inter, 0.0);
6918 QTensor::matvec_many([&d.gate_proj, &d.up_proj], x, [g, u], pool);
6919 for i in 0..inter {
6920 g[i] = d.act.combine(g[i], u[i]);
6921 }
6922 }
6923 zero_masked_cols(g, 1, inter, mask_row);
6924 let mut out = attention::take_buf(d.down_proj.rows());
6925 d.down_proj.matvec(g, &mut out, pool);
6926 out
6927 })
6928}
6929
6930fn dense_ffn_gpu(d: &DenseFfn, x: &[f32], _pool: Option<&Pool>) -> Option<Vec<f32>> {
6936 if d.act != Act::Silu {
6938 return None;
6939 }
6940 if d.gate_proj.rows() < crate::gpu::min_rows() && !d.gate_proj.is_q1() {
6943 return None;
6944 }
6945 let mut jobs: Vec<crate::gpu::MoeJob> = Vec::with_capacity(1);
6946 let mut model_ref = None;
6947 moe_push_job(d, x, 1.0, &mut jobs, &mut model_ref)?;
6948 let model = model_ref?;
6949 let hidden = jobs[0].down.1;
6950 let mut out = attention::take_buf(hidden);
6951 if crate::gpu::moe_block(&model, &jobs, &mut out) {
6952 Some(out)
6953 } else {
6954 let mut out = out;
6955 attention::recycle_buf(&mut out);
6956 None
6957 }
6958}
6959
6960#[allow(clippy::type_complexity)]
6965#[allow(clippy::type_complexity)]
6966pub(crate) fn moe_parts(
6967 t: &QTensor,
6968) -> Option<(
6969 &std::sync::Arc<cortiq_core::CmfModel>,
6970 usize,
6971 usize,
6972 usize,
6973 &[f32],
6974 &[f32],
6975 bool,
6976 bool,
6977 bool,
6978)> {
6979 match t {
6980 QTensor::Mapped {
6981 model,
6982 idx,
6983 dtype: dt @ (cortiq_core::TensorDtype::Q8_2f | cortiq_core::TensorDtype::Q8Row),
6984 rows,
6985 cols,
6986 row_scale,
6987 col_field,
6988 ..
6989 } if (*dt == cortiq_core::TensorDtype::Q8Row) || !col_field.is_empty() => Some((
6990 model, *idx, *rows, *cols, row_scale, col_field, false, false, false,
6991 )),
6992 QTensor::Mapped {
6994 model,
6995 idx,
6996 dtype: cortiq_core::TensorDtype::Q1,
6997 rows,
6998 cols,
6999 ..
7000 } => Some((model, *idx, *rows, *cols, &[][..], &[][..], true, false, false)),
7001 QTensor::Mapped {
7003 model,
7004 idx,
7005 dtype: cortiq_core::TensorDtype::Q4Tiled,
7006 rows,
7007 cols,
7008 ..
7009 } => Some((model, *idx, *rows, *cols, &[][..], &[][..], false, true, false)),
7010 QTensor::Mapped {
7012 model,
7013 idx,
7014 dtype: cortiq_core::TensorDtype::Q4TiledP,
7015 rows,
7016 cols,
7017 ..
7018 } => Some((model, *idx, *rows, *cols, &[][..], &[][..], false, true, false)),
7019 QTensor::Mapped {
7023 model,
7024 idx,
7025 dtype: cortiq_core::TensorDtype::Q2TiledP,
7026 rows,
7027 cols,
7028 ..
7029 } => Some((model, *idx, *rows, *cols, &[][..], &[][..], false, true, true)),
7030 _ => None,
7031 }
7032}
7033
7034#[cfg(target_os = "macos")]
7040fn metal_moe_graph_parts(m: &MoeFfn, hidden: usize) -> Option<crate::gpu::GpuMoe<'_>> {
7041 if m.router_sigmoid
7042 || m.router_input_norm
7043 || m.expert_bias.is_some()
7044 || m.route_tau.is_some()
7045 || m.mask.is_some()
7046 || m.per_expert_scale.is_some()
7047 || m.experts.is_empty()
7048 || m.top_k == 0
7049 {
7050 return None;
7051 }
7052 let (sh, sg) = match &m.shared {
7055 Some((sh, Some(sg))) => (sh, sg),
7056 _ => return None,
7057 };
7058 let (rf, rr, rc) = m.router.f32_parts()?;
7059 if rr != m.experts.len() || rc != hidden {
7060 return None;
7061 }
7062 let (sf, sr, sc) = sg.f32_parts()?;
7063 if sr * sc != hidden {
7064 return None;
7065 }
7066 let inter = m.experts[0].gate_proj.rows();
7067 let gu_q2 = m.experts[0].gate_proj.mapped_q2tp().is_some();
7070 let trio = |e: &DenseFfn| -> Option<(usize, usize, usize)> {
7071 if e.act != Act::Silu
7072 || e.gate_proj.rows() != inter
7073 || e.gate_proj.cols() != hidden
7074 || e.up_proj.rows() != inter
7075 || e.up_proj.cols() != hidden
7076 || e.down_proj.rows() != hidden
7077 || e.down_proj.cols() != inter
7078 {
7079 return None;
7080 }
7081 let pick = |t: &QTensor| -> Option<usize> {
7082 if gu_q2 {
7083 t.mapped_q2tp().map(|(_, i)| i)
7084 } else {
7085 t.mapped_q4tp().map(|(_, i)| i)
7086 }
7087 };
7088 Some((
7089 pick(&e.gate_proj)?,
7090 pick(&e.up_proj)?,
7091 e.down_proj.mapped_q4tp().map(|(_, i)| i)?,
7092 ))
7093 };
7094 let experts = m
7095 .experts
7096 .iter()
7097 .map(trio)
7098 .collect::<Option<Vec<_>>>()?;
7099 let shared = trio(sh)?;
7100 Some(crate::gpu::GpuMoe {
7101 router: rf,
7102 sgate: sf,
7103 experts,
7104 shared,
7105 n_exp: m.experts.len(),
7106 top_k: m.top_k,
7107 inter,
7108 norm_topk: m.norm_topk_prob,
7109 route_scale: m.routed_scaling,
7110 gu_q2,
7111 })
7112}
7113
7114pub(crate) fn moe_push_job_parts<'a>(
7118 gate: &'a QTensor,
7119 up: &'a QTensor,
7120 down: &'a QTensor,
7121 x: &[f32],
7122 w: f32,
7123 swiglu_limit: f32,
7124 jobs: &mut Vec<crate::gpu::MoeJob<'a>>,
7125 model_ref: &mut Option<std::sync::Arc<cortiq_core::CmfModel>>,
7126) -> Option<()> {
7127 use crate::qtensor::prescale;
7128 let (gm, gi, gr, gc, grs, gcf, gq1, gq4, gq2) = moe_parts(gate)?;
7129 let (_, ui, ur, uc, urs, ucf, uq1, uq4, uq2) = moe_parts(up)?;
7130 let (_, di, dr, dc, drs, dcf, dq1, dq4, dq2) = moe_parts(down)?;
7131 if gq1 != uq1 || uq1 != dq1 || gq4 != uq4 || uq4 != dq4 || gq2 != uq2 {
7132 return None; }
7134 if gq2 && (dq2 || !dq4 || down.mapped_q4tp().is_none()) {
7137 return None;
7138 }
7139 if !gq2 && dq2 {
7140 return None;
7141 }
7142 model_ref.get_or_insert_with(|| gm.clone());
7143 let dt = |cf: &[f32]| {
7144 if cf.is_empty() {
7145 cortiq_core::TensorDtype::Q8Row
7146 } else {
7147 cortiq_core::TensorDtype::Q8_2f
7148 }
7149 };
7150 jobs.push(crate::gpu::MoeJob {
7151 gate: (gi, gr, gc, grs),
7152 up: (ui, ur, uc, urs),
7153 down: (di, dr, dc, drs),
7154 xs_gate: prescale(x, gcf, dt(gcf)).into_owned(),
7155 xs_up: prescale(x, ucf, dt(ucf)).into_owned(),
7156 down_col: dcf,
7157 w,
7158 q1: gq1,
7159 q4t: gq4 && !gq2 && gate.mapped_q4tp().is_none(),
7160 q4tp: gq4 && (gq2 || gate.mapped_q4tp().is_some()),
7161 gu_q2: gq2,
7162 swiglu_limit,
7163 });
7164 Some(())
7165}
7166
7167fn moe_push_job<'a>(
7169 d: &'a DenseFfn,
7170 x: &[f32],
7171 w: f32,
7172 jobs: &mut Vec<crate::gpu::MoeJob<'a>>,
7173 model_ref: &mut Option<std::sync::Arc<cortiq_core::CmfModel>>,
7174) -> Option<()> {
7175 use crate::qtensor::prescale;
7176 if d.act != Act::Silu {
7177 return None; }
7179 let (gm, gi, gr, gc, grs, gcf, gq1, gq4, gq2) = moe_parts(&d.gate_proj)?;
7180 let (_, ui, ur, uc, urs, ucf, uq1, uq4, uq2) = moe_parts(&d.up_proj)?;
7181 let (_, di, dr, dc, drs, dcf, dq1, dq4, dq2) = moe_parts(&d.down_proj)?;
7182 if gq1 != uq1 || uq1 != dq1 || gq4 != uq4 || uq4 != dq4 || gq2 != uq2 {
7183 return None; }
7185 if gq2 && (dq2 || !dq4 || d.down_proj.mapped_q4tp().is_none()) {
7186 return None;
7187 }
7188 if !gq2 && dq2 {
7189 return None;
7190 }
7191 model_ref.get_or_insert_with(|| gm.clone());
7192 let gdt = if gcf.is_empty() {
7193 cortiq_core::TensorDtype::Q8Row
7194 } else {
7195 cortiq_core::TensorDtype::Q8_2f
7196 };
7197 let udt = if ucf.is_empty() {
7198 cortiq_core::TensorDtype::Q8Row
7199 } else {
7200 cortiq_core::TensorDtype::Q8_2f
7201 };
7202 jobs.push(crate::gpu::MoeJob {
7203 gate: (gi, gr, gc, grs),
7204 up: (ui, ur, uc, urs),
7205 down: (di, dr, dc, drs),
7206 xs_gate: prescale(x, gcf, gdt).into_owned(),
7207 xs_up: prescale(x, ucf, udt).into_owned(),
7208 down_col: dcf,
7209 w,
7210 q1: gq1,
7211 q4t: gq4 && !gq2 && d.gate_proj.mapped_q4tp().is_none(),
7212 q4tp: gq4 && (gq2 || d.gate_proj.mapped_q4tp().is_some()),
7213 gu_q2: gq2,
7214 swiglu_limit: 0.0,
7215 });
7216 Some(())
7217}
7218
7219fn sparse_ffn_quant(
7226 d: &DenseFfn,
7227 x: &[f32],
7228 active: &[u16],
7229 hidden: usize,
7230 pool: Option<&Pool>,
7231) -> Vec<f32> {
7232 let n = active.len();
7233 let inter = d.gate_proj.rows();
7234 let mut act = vec![0.0f32; n];
7235 let need_scratch = !(d.gate_proj.sparse_col_ok() && d.up_proj.sparse_col_ok());
7238 let compute = |ai: usize| -> f32 {
7239 let idx = active[ai] as usize;
7240 if idx >= inter {
7241 return 0.0; }
7243 let mut s = if need_scratch {
7244 vec![0.0f32; hidden]
7245 } else {
7246 Vec::new()
7247 };
7248 let gate = d.gate_proj.row_dot(idx, x, &mut s);
7249 let up = d.up_proj.row_dot(idx, x, &mut s);
7250 d.act.combine(gate, up)
7251 };
7252 match pool {
7253 Some(p) if n >= 256 => {
7254 let ptr = SendMut(act.as_mut_ptr());
7255 p.run(&|widx, nw| {
7256 let chunk = n.div_ceil(nw);
7257 let (s, e) = (widx * chunk, ((widx + 1) * chunk).min(n));
7258 for ai in s..e {
7259 unsafe { *ptr.at(ai) = compute(ai) };
7260 }
7261 });
7262 }
7263 _ => {
7264 for (ai, a) in act.iter_mut().enumerate() {
7265 *a = compute(ai);
7266 }
7267 }
7268 }
7269 let mut out = vec![0.0f32; hidden];
7271 for (ai, &idx) in active.iter().enumerate() {
7272 let w = act[ai];
7273 if w.abs() >= 1e-12 && (idx as usize) < inter {
7274 d.down_proj.add_col_scaled(idx as usize, w, &mut out);
7275 }
7276 }
7277 out
7278}
7279
7280#[doc(hidden)]
7282pub fn sparse_ffn_quant_for_test(
7283 d: &DenseFfn,
7284 x: &[f32],
7285 active: &[u16],
7286 hidden: usize,
7287) -> Vec<f32> {
7288 sparse_ffn_quant(d, x, active, hidden, None)
7289}
7290
7291fn dequant_dense_f32(d: &DenseFfn) -> (Vec<f32>, Vec<f32>, Vec<f32>) {
7295 let deq = |t: &QTensor| -> Vec<f32> {
7296 let (rows, cols) = (t.rows(), t.cols());
7297 let mut out = vec![0.0f32; rows * cols];
7298 for r in 0..rows {
7299 t.row_f32(r, &mut out[r * cols..(r + 1) * cols]);
7300 }
7301 out
7302 };
7303 (deq(&d.gate_proj), deq(&d.up_proj), deq(&d.down_proj))
7304}
7305
7306struct SendMut(*mut f32);
7308unsafe impl Send for SendMut {}
7309unsafe impl Sync for SendMut {}
7310impl SendMut {
7311 #[inline]
7312 #[allow(clippy::mut_from_ref)]
7315 unsafe fn at(&self, i: usize) -> &mut f32 {
7316 unsafe { &mut *self.0.add(i) }
7317 }
7318}
7319
7320fn moe_route(logits: &[f32], m: &MoeFfn, allowed: Option<&[bool]>) -> (Vec<usize>, Vec<f32>, f32) {
7330 let ne = logits.len();
7331 let p: Vec<f32> = if m.router_sigmoid {
7332 logits.iter().map(|&l| 1.0 / (1.0 + (-l).exp())).collect()
7333 } else {
7334 let mx = logits.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
7335 let mut e: Vec<f32> = logits.iter().map(|&l| (l - mx).exp()).collect();
7336 let s: f32 = e.iter().sum();
7337 for v in &mut e {
7338 *v /= s;
7339 }
7340 e
7341 };
7342 let admit = |e: usize| {
7348 m.mask.as_ref().is_none_or(|mk| mk[e])
7349 && allowed.is_none_or(|a| a.get(e).copied().unwrap_or(false))
7350 };
7351 let mut idx: Vec<usize> = (0..ne).filter(|&e| admit(e)).collect();
7352 match &m.expert_bias {
7354 Some(b) => idx.sort_unstable_by(|&x, &y| {
7355 (p[y] + b[y])
7356 .partial_cmp(&(p[x] + b[x]))
7357 .unwrap()
7358 .then(x.cmp(&y))
7359 }),
7360 None => idx.sort_unstable_by(|&x, &y| p[y].partial_cmp(&p[x]).unwrap().then(x.cmp(&y))),
7361 }
7362 idx.truncate(m.top_k);
7363 if let Some(tau) = m.route_tau {
7367 let total: f32 = idx.iter().map(|&e| p[e]).sum();
7368 if total > 0.0 {
7369 let mut acc = 0.0f32;
7370 let mut keep = idx.len();
7371 for (i, &e) in idx.iter().enumerate() {
7372 acc += p[e];
7373 if acc >= tau * total {
7374 keep = i + 1;
7375 break;
7376 }
7377 }
7378 idx.truncate(keep);
7379 }
7380 }
7381 let wsum: f32 = if m.norm_topk_prob {
7382 let s: f32 = idx.iter().map(|&e| p[e]).sum();
7383 (if m.router_sigmoid { s + 1e-6 } else { s }) / m.routed_scaling
7386 } else {
7387 1.0 / m.routed_scaling
7388 };
7389 (idx, p, wsum)
7390}
7391
7392fn moe_ffn(m: &MoeFfn, x: &[f32], pool: Option<&Pool>, allowed: Option<&[bool]>) -> Vec<f32> {
7395 accumulate_act(m, x, 1);
7396 let ne = m.experts.len();
7397 let mut logits = vec![0.0f32; ne];
7398 m.router.matvec(x, &mut logits, pool);
7399 let (idx, p, wsum) = moe_route(&logits, m, allowed);
7400 {
7401 let mut st = m.stats.borrow_mut();
7402 if st.len() < ne {
7403 st.resize(ne, 0);
7404 }
7405 for &e in &idx {
7406 st[e] += 1;
7407 }
7408 }
7409 if crate::gpu::enabled_here() {
7414 match crate::gpu::probe_arm(crate::gpu::OpClass::Ffn) {
7415 crate::gpu::ProbeArm::Gpu => {
7416 let t0 = std::time::Instant::now();
7417 if let Some(out) = moe_ffn_gpu(m, x, &idx, &p, wsum, pool) {
7418 crate::gpu::probe_record(crate::gpu::OpClass::Ffn, true, t0.elapsed());
7419 return out;
7420 }
7421 }
7422 crate::gpu::ProbeArm::CpuTimed => {
7423 let t0 = std::time::Instant::now();
7424 let out = crate::gpu::cpu_scope(|| moe_ffn_cpu(m, x, &idx, &p, wsum, pool));
7425 crate::gpu::probe_record(crate::gpu::OpClass::Ffn, false, t0.elapsed());
7426 return out;
7427 }
7428 crate::gpu::ProbeArm::Cpu => {
7429 return crate::gpu::cpu_scope(|| moe_ffn_cpu(m, x, &idx, &p, wsum, pool));
7430 }
7431 }
7432 }
7433 moe_ffn_cpu(m, x, &idx, &p, wsum, pool)
7434}
7435
7436fn graph_note(built: bool) {
7440 use std::sync::atomic::{AtomicBool, Ordering};
7441 if built {
7442 GRAPH_TOK_OK.fetch_add(1, Ordering::Relaxed);
7443 } else {
7444 GRAPH_TOK_MISS.fetch_add(1, Ordering::Relaxed);
7445 }
7446 static SAID: AtomicBool = AtomicBool::new(false);
7447 if !SAID.swap(true, Ordering::Relaxed) {
7448 if built {
7449 tracing::info!("wgpu whole-token graph: ACTIVE");
7450 } else {
7451 tracing::warn!("wgpu whole-token graph refused — per-op path");
7452 }
7453 }
7454}
7455
7456pub static GRAPH_TOK_OK: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
7460pub static GRAPH_TOK_MISS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
7461
7462fn moe_batch_enabled() -> bool {
7465 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
7466 *ON.get_or_init(|| std::env::var("CMF_MOE_BATCH").as_deref() != Ok("0"))
7467}
7468
7469fn moe_ffn_cpu_batched(
7475 m: &MoeFfn,
7476 x: &[f32],
7477 idx: &[usize],
7478 p: &[f32],
7479 wsum: f32,
7480 pool: Option<&Pool>,
7481) -> Option<Vec<f32>> {
7482 if idx.is_empty() || !moe_batch_enabled() {
7483 return None;
7484 }
7485 if FFN_PROBE.with(|pr| pr.borrow().is_some()) {
7489 return None;
7490 }
7491 let n = idx.len() + usize::from(m.shared.is_some());
7492 let mut pairs = Vec::with_capacity(n);
7493 let mut downs = Vec::with_capacity(n);
7494 let mut ws = Vec::with_capacity(n);
7495 for &e in idx {
7496 let d = &m.experts[e];
7497 if d.act != Act::Silu {
7498 return None;
7499 }
7500 pairs.push((&d.gate_proj, &d.up_proj));
7501 downs.push(&d.down_proj);
7502 ws.push(p[e] / wsum * m.per_expert_scale.as_ref().map_or(1.0, |v| v[e]));
7503 }
7504 if let Some((se, gate)) = &m.shared {
7507 if se.act != Act::Silu {
7508 return None;
7509 }
7510 let g = gate.as_ref().map_or(1.0, |gate| {
7511 let mut gl = [0.0f32; 1];
7512 gate.matvec(x, &mut gl, pool);
7513 1.0 / (1.0 + (-gl[0]).exp())
7514 });
7515 pairs.push((&se.gate_proj, &se.up_proj));
7516 downs.push(&se.down_proj);
7517 ws.push(g);
7518 }
7519 let inter = pairs[0].0.rows();
7520 let mut gs: Vec<Vec<f32>> = (0..pairs.len()).map(|_| vec![0f32; inter]).collect();
7521 if !QTensor::moe_gate_up_many(&pairs, x, &mut gs, pool) {
7522 return None;
7523 }
7524 let mut out = attention::take_buf(x.len());
7525 if !QTensor::moe_down_many(&downs, &gs, &ws, &mut out, pool) {
7526 attention::recycle_buf(&mut out);
7527 return None;
7528 }
7529 Some(out)
7530}
7531
7532fn moe_ffn_cpu(
7534 m: &MoeFfn,
7535 x: &[f32],
7536 idx: &[usize],
7537 p: &[f32],
7538 wsum: f32,
7539 pool: Option<&Pool>,
7540) -> Vec<f32> {
7541 if let Some(out) = moe_ffn_cpu_batched(m, x, idx, p, wsum, pool) {
7542 return out;
7543 }
7544 let mut out = attention::take_buf(x.len());
7545 for &e in idx {
7546 let mut eo = dense_ffn(&m.experts[e], x, pool);
7547 let w = p[e] / wsum * m.per_expert_scale.as_ref().map_or(1.0, |v| v[e]);
7548 for i in 0..out.len() {
7549 out[i] += w * eo[i];
7550 }
7551 attention::recycle_buf(&mut eo);
7552 }
7553 if let Some((se, gate)) = &m.shared {
7554 let mut so = dense_ffn(se, x, pool);
7555 let g = gate.as_ref().map_or(1.0, |gate| {
7556 let mut gl = [0.0f32; 1];
7557 gate.matvec(x, &mut gl, pool);
7558 1.0 / (1.0 + (-gl[0]).exp())
7559 });
7560 for i in 0..out.len() {
7561 out[i] += g * so[i];
7562 }
7563 attention::recycle_buf(&mut so);
7564 }
7565 out
7566}
7567
7568#[allow(clippy::too_many_arguments)]
7576fn mla_attention(
7577 w: &MlaWeights,
7578 normed: &[f32],
7579 cache: &mut crate::kv_cache::LayerKvCache,
7580 position: usize,
7581 inv_freq: &[f32],
7582 rope_scale: f32,
7583 eps: f64,
7584 pool: Option<&Pool>,
7585) -> Vec<f32> {
7586 let (nh, dr, dn, dv, lora) = (w.nh, w.qk_rope, w.qk_nope, w.v_dim, w.lora);
7587 let hd = dr + dn;
7588 let mut q = vec![0.0f32; nh * hd];
7589 match (&w.q_a, &w.q_a_norm) {
7590 (Some(qa), Some(qn)) => {
7591 let mut t = vec![0.0f32; qa.rows()];
7592 qa.matvec(normed, &mut t, pool);
7593 let tn = inference::rms_norm(&t, qn, eps, NormStyle::Qwen);
7594 w.q_proj.matvec(&tn, &mut q, pool);
7595 }
7596 _ => w.q_proj.matvec(normed, &mut q, pool),
7597 }
7598 let mut ca = vec![0.0f32; lora + dr];
7599 w.kv_a.matvec(normed, &mut ca, pool);
7600 let (c_lat, k_rope) = ca.split_at_mut(lora);
7601 let latn = inference::rms_norm(c_lat, &w.kv_a_norm, eps, NormStyle::Qwen);
7602 let mut kvb = vec![0.0f32; nh * (dn + dv)];
7603 w.kv_b.matvec(&latn, &mut kvb, pool);
7604 if !w.nope {
7605 attention::rope_rotate_scaled(k_rope, position, inv_freq, rope_scale);
7606 }
7607 for h in 0..nh {
7608 if !w.nope {
7609 attention::rope_rotate_scaled(
7610 &mut q[h * hd..h * hd + dr],
7611 position,
7612 inv_freq,
7613 rope_scale,
7614 );
7615 }
7616 }
7617 let mut k = vec![0.0f32; nh * hd];
7618 let mut v = vec![0.0f32; nh * hd];
7619 for h in 0..nh {
7620 k[h * hd..h * hd + dr].copy_from_slice(k_rope);
7621 k[h * hd + dr..(h + 1) * hd].copy_from_slice(&kvb[h * (dn + dv)..h * (dn + dv) + dn]);
7622 v[h * hd..h * hd + dv].copy_from_slice(&kvb[h * (dn + dv) + dn..(h + 1) * (dn + dv)]);
7623 }
7624 cache.append(&k, &v, &vec![true; nh]);
7625 let (ao, mut imp) = attention::attend_all_heads(&q, cache, nh, 1, hd, w.scale, None, 0.0);
7626 attention::recycle_buf(&mut imp);
7627 let mut ov = vec![0.0f32; nh * dv];
7628 for h in 0..nh {
7629 ov[h * dv..(h + 1) * dv].copy_from_slice(&ao[h * hd..h * hd + dv]);
7630 }
7631 let mut out = vec![0.0f32; w.o_proj.rows()];
7632 w.o_proj.matvec(&ov, &mut out, pool);
7633 out
7634}
7635
7636fn dense_moe_ffn(
7643 dm: &DenseMoeFfn,
7644 x_normed: &[f32],
7645 h_raw: &[f32],
7646 eps: f64,
7647 norm_style: NormStyle,
7648 pool: Option<&Pool>,
7649) -> Vec<f32> {
7650 let mut d = dense_ffn(&dm.dense, x_normed, pool);
7651 d = inference::rms_norm(&d, &dm.post_norm_1, eps, norm_style);
7652 let m = &dm.moe;
7653 let ne = m.experts.len();
7654 let mut logits = vec![0.0f32; ne];
7655 if m.router_input_norm {
7656 let ss: f32 = h_raw.iter().map(|v| v * v).sum::<f32>() / h_raw.len() as f32;
7657 let inv = 1.0 / (ss + eps as f32).sqrt();
7658 let xr: Vec<f32> = h_raw.iter().map(|v| v * inv).collect();
7659 m.router.matvec(&xr, &mut logits, pool);
7660 } else {
7661 m.router.matvec(h_raw, &mut logits, pool);
7662 }
7663 let (idx, p, wsum) = moe_route(&logits, m, None);
7664 {
7665 let mut st = m.stats.borrow_mut();
7666 if st.len() < ne {
7667 st.resize(ne, 0);
7668 }
7669 for &e in &idx {
7670 st[e] += 1;
7671 }
7672 }
7673 let x2 = inference::rms_norm(h_raw, &dm.pre_norm_2, eps, norm_style);
7674 let mo = moe_ffn_cpu(m, &x2, &idx, &p, wsum, pool);
7675 let mo = inference::rms_norm(&mo, &dm.post_norm_2, eps, norm_style);
7676 for (di, mi) in d.iter_mut().zip(&mo) {
7677 *di += mi;
7678 }
7679 d
7680}
7681
7682fn moe_gpu_refused(why: &'static str) {
7689 use std::sync::atomic::{AtomicBool, Ordering};
7690 static SAID: AtomicBool = AtomicBool::new(false);
7691 if !SAID.swap(true, Ordering::Relaxed) {
7692 tracing::warn!("MoE GPU block refused ({why}) — experts run on the CPU");
7693 }
7694}
7695
7696fn moe_ffn_gpu(
7697 m: &MoeFfn,
7698 x: &[f32],
7699 idx: &[usize],
7700 p: &[f32],
7701 wsum: f32,
7702 pool: Option<&Pool>,
7703) -> Option<Vec<f32>> {
7704 use crate::gpu::MoeJob;
7705
7706 let mut jobs: Vec<MoeJob> = Vec::with_capacity(idx.len() + 1);
7707 let mut model_ref = None;
7708 for &e in idx {
7709 if moe_push_job(&m.experts[e], x, p[e] / wsum, &mut jobs, &mut model_ref).is_none() {
7710 moe_gpu_refused("push_job(expert)");
7711 return None;
7712 }
7713 }
7714 if let Some((se, gate)) = &m.shared {
7715 let g = gate.as_ref().map_or(1.0, |gate| {
7716 let mut gl = [0.0f32; 1];
7717 gate.matvec(x, &mut gl, pool);
7718 1.0 / (1.0 + (-gl[0]).exp())
7719 });
7720 if moe_push_job(se, x, g, &mut jobs, &mut model_ref).is_none() {
7721 moe_gpu_refused("push_job(shared)");
7722 return None;
7723 }
7724 }
7725 let Some(model) = model_ref else {
7726 moe_gpu_refused("no model_ref");
7727 return None;
7728 };
7729 let hidden = jobs[0].down.1;
7730 let mut out = vec![0.0f32; hidden];
7731 if crate::gpu::moe_block(&model, &jobs, &mut out) {
7732 Some(out)
7733 } else {
7734 moe_gpu_refused("gpu::moe_block");
7735 None
7736 }
7737}
7738
7739fn ffn_forward(
7741 ffn: &FfnKind,
7742 x: &[f32],
7743 pool: Option<&Pool>,
7744 experts_allowed: Option<&[bool]>,
7745) -> Vec<f32> {
7746 match ffn {
7747 FfnKind::Dense(d) => dense_ffn(d, x, pool),
7748 FfnKind::Moe(m) => moe_ffn(m, x, pool, experts_allowed),
7749 FfnKind::DenseMoe(_) => unreachable!("DenseMoe dispatches via dense_moe_ffn"),
7753 }
7754}
7755
7756fn ffn_forward_pair(
7760 ffn: &FfnKind,
7761 x1: &[f32],
7762 x2: &[f32],
7763 pool: Option<&Pool>,
7764 experts_allowed: Option<&[bool]>,
7765) -> (Vec<f32>, Vec<f32>) {
7766 let d = match ffn {
7767 FfnKind::Dense(d) => d,
7768 FfnKind::Moe(m) => {
7769 return (
7770 moe_ffn(m, x1, pool, experts_allowed),
7771 moe_ffn(m, x2, pool, experts_allowed),
7772 );
7773 }
7774 FfnKind::DenseMoe(_) => unreachable!("DenseMoe dispatches via dense_moe_ffn"),
7775 };
7776 let inter = d.gate_proj.rows();
7777 FFN_SCRATCH.with(|s| {
7778 let mut s = s.borrow_mut();
7779 let [g1, g2, u1, u2] = &mut *s;
7780 g1.resize(inter, 0.0);
7781 g2.resize(inter, 0.0);
7782 u1.resize(inter, 0.0);
7783 u2.resize(inter, 0.0);
7784 QTensor::matvec2_many(
7787 [&d.gate_proj, &d.up_proj],
7788 x1,
7789 x2,
7790 [g1.as_mut_slice(), u1.as_mut_slice()],
7791 [g2.as_mut_slice(), u2.as_mut_slice()],
7792 pool,
7793 );
7794 for i in 0..inter {
7795 g1[i] = d.act.combine(g1[i], u1[i]);
7796 g2[i] = d.act.combine(g2[i], u2[i]);
7797 }
7798 let mut o1 = attention::take_buf(d.down_proj.rows());
7799 let mut o2 = attention::take_buf(d.down_proj.rows());
7800 d.down_proj.matvec2(g1, g2, &mut o1, &mut o2, pool);
7801 (o1, o2)
7802 })
7803}
7804
7805#[cfg(test)]
7806mod tests {
7807
7808 #[test]
7809 fn cancel_flag_stops_generation() {
7810 let mut p = create_test_pipeline(16, 32, 2, 2, 8, 2, 32);
7811 p.cancel.store(true, std::sync::atomic::Ordering::Relaxed);
7814 let r = p.generate_from_ids(&[1, 2, 3], 8, None, None).unwrap();
7815 assert_eq!(r.finish_reason, "cancelled");
7816 assert!(
7817 r.token_ids.is_empty(),
7818 "no tokens after cancel: {:?}",
7819 r.token_ids
7820 );
7821 let r2 = p.generate_from_ids(&[1, 2, 3], 4, None, None).unwrap();
7823 assert_ne!(r2.finish_reason, "cancelled");
7824 }
7825 use super::*;
7826
7827 #[test]
7833 fn sparse_ffn_quant_equals_dense_with_inactive_zeroed() {
7834 let (hidden, inter) = (16usize, 40usize);
7835 let synth = |n: usize, salt: usize| -> Vec<f32> {
7836 (0..n)
7837 .map(|i| (((i * 37 + salt * 11 + 3) % 101) as f32 / 101.0 - 0.5) * 0.4)
7838 .collect()
7839 };
7840 let d = DenseFfn {
7841 gate_proj: QTensor::from_f32(synth(inter * hidden, 1), inter, hidden),
7842 up_proj: QTensor::from_f32(synth(inter * hidden, 2), inter, hidden),
7843 down_proj: QTensor::from_f32(synth(hidden * inter, 3), hidden, inter),
7844 act: Act::Silu,
7845 };
7846 let x = synth(hidden, 9);
7847 let active: Vec<u16> = (0..inter as u16).filter(|i| i % 3 == 0).collect();
7849
7850 let sparse = sparse_ffn_quant(&d, &x, &active, hidden, None);
7851
7852 let mut g = vec![0.0f32; inter];
7854 d.gate_proj.matvec(&x, &mut g, None);
7855 let mut u = vec![0.0f32; inter];
7856 d.up_proj.matvec(&x, &mut u, None);
7857 let act_set: std::collections::HashSet<u16> = active.iter().copied().collect();
7858 for i in 0..inter {
7859 g[i] = if act_set.contains(&(i as u16)) {
7860 inference::silu(g[i]) * u[i]
7861 } else {
7862 0.0
7863 };
7864 }
7865 let mut reference = vec![0.0f32; hidden];
7866 d.down_proj.matvec(&g, &mut reference, None);
7867
7868 let max_d = sparse
7869 .iter()
7870 .zip(&reference)
7871 .map(|(a, b)| (a - b).abs())
7872 .fold(0.0f32, f32::max);
7873 assert!(max_d < 1e-5, "sparse != dense-zeroed: max|Δ| = {max_d}");
7874 }
7875
7876 fn attach_test_mtp(p: &mut Pipeline) {
7878 let (h, inter, heads, kv, hd) = (
7879 p.hidden_size,
7880 p.intermediate_size,
7881 p.num_heads,
7882 p.num_kv_heads,
7883 p.head_dim,
7884 );
7885 let synth = |n: usize, salt: usize| -> Vec<f32> {
7886 (0..n)
7887 .map(|i| (((i * 29 + salt * 23 + 5) % 101) as f32 / 101.0 - 0.5) * 0.2)
7888 .collect()
7889 };
7890 let qt = |rows: usize, cols: usize, salt: usize| -> QTensor {
7891 QTensor::from_f32(synth(rows * cols, salt), rows, cols)
7892 };
7893 p.mtp = Some(MtpModule {
7894 enorm: vec![1.0; h],
7895 hnorm: vec![1.0; h],
7896 eh_proj: qt(h, 2 * h, 301),
7897 layer: LayerWeights {
7898 input_norm: vec![1.0; h],
7899 post_norm: vec![1.0; h],
7900 attn_out_norm: None,
7901 ffn_out_norm: None,
7902 layer_scale: None,
7903 ffn: FfnKind::Dense(DenseFfn {
7904 gate_proj: qt(inter, h, 315),
7905 up_proj: qt(inter, h, 316),
7906 down_proj: qt(h, inter, 317),
7907 act: Act::Silu,
7908 }),
7909 attn: AttnKind::Full {
7910 bias: None,
7911 wq: qt(heads * hd, h, 311),
7912 wk: qt(kv * hd, h, 312),
7913 wv: qt(kv * hd, h, 313),
7914 wo: qt(h, heads * hd, 314),
7915 q_norm: None,
7916 k_norm: None,
7917 output_gate: false,
7918 softplus_gate: None,
7919 },
7920 },
7921 final_norm: vec![1.0; h],
7922 kv: crate::kv_cache::LayerKvCache::new(kv, hd),
7923 });
7924 }
7925
7926 #[test]
7927 fn speculative_equals_vanilla_greedy() {
7928 unsafe { std::env::set_var("CMF_GPU_WGPU_GRAPH", "0") };
7932 let run = |spec: bool| {
7933 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
7934 p.sampler_config.temperature = 0.0;
7935 attach_test_mtp(&mut p);
7936 p.speculative = spec;
7937 let r = p.generate("abcdef", 12, None, None).unwrap();
7938 (r.token_ids, r.mtp_drafted, r.mtp_accepted)
7939 };
7940 let (vanilla, d0, _) = run(false);
7941 let (spec, d1, a1) = run(true);
7942 assert_eq!(d0, 0, "vanilla path must not draft");
7943 assert!(d1 > 0, "speculative path must draft");
7944 assert_eq!(
7945 vanilla, spec,
7946 "speculative must reproduce the exact greedy sequence (accepted {a1}/{d1})"
7947 );
7948 }
7949
7950 #[test]
7951 fn speculative_accepts_constant_oracle() {
7952 unsafe { std::env::set_var("CMF_GPU_WGPU_GRAPH", "0") };
7954 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
7955 p.sampler_config.temperature = 0.0;
7956 p.sampler_config.repetition_penalty = 1.0;
7957 p.weights.lm_head = QTensor::from_f32(vec![0.01; 64 * 8], 64, 8);
7960 attach_test_mtp(&mut p);
7961 p.speculative = true;
7962 let r = p.generate("abcd", 10, None, None).unwrap();
7963 assert!(r.mtp_drafted > 0);
7964 assert_eq!(
7965 r.mtp_accepted, r.mtp_drafted,
7966 "constant logits → every draft accepted"
7967 );
7968 assert!(r.token_ids.windows(2).all(|w| w[0] == w[1]));
7971 }
7972
7973 #[test]
7974 fn empty_prompt_is_an_error_not_a_panic() {
7975 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 260);
7976 let r = p.generate("", 4, None, None);
7977 assert!(r.is_err(), "empty prompt must be a clean error");
7978 }
7979
7980 #[test]
7981 fn every_token_enters_kv_exactly_once() {
7982 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
7983 p.sampler_config.temperature = 0.0;
7985 let r = p.generate("abc", 2, None, None).unwrap();
7986 assert_eq!(r.prompt_tokens, 3);
7987 assert_eq!(
7991 p.kv_cache.seq_len(),
7992 3 + r.tokens_generated - 1,
7993 "each token must be cached exactly once (v1 cached the last prompt token twice)"
7994 );
7995 }
7996
7997 #[test]
7998 fn generation_is_reproducible_with_seed() {
7999 let run = || {
8000 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
8001 p.generate("hello", 8, None, None).unwrap().token_ids
8002 };
8003 assert_eq!(run(), run());
8004 }
8005
8006 #[test]
8007 fn resetting_sampler_restarts_the_seeded_stream() {
8008 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
8009 let config = SamplerConfig {
8010 seed: Some(1234),
8011 ..SamplerConfig::default()
8012 };
8013 p.set_sampler_config(config.clone());
8014 let first = p.generate("hello", 8, None, None).unwrap().token_ids;
8015 p.set_sampler_config(config);
8016 let second = p.generate("hello", 8, None, None).unwrap().token_ids;
8017 assert_eq!(first, second);
8018 }
8019
8020 #[test]
8021 fn eviction_bounds_the_cache() {
8022 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 260);
8023 p.kv_cache.max_seq_len = 6;
8024 p.sampler_config.temperature = 0.0;
8025 let _ = p.generate("abcd", 12, None, None).unwrap();
8026 assert!(
8027 p.kv_cache.seq_len() <= 6 + 1,
8028 "cache must stay bounded by max_seq_len (got {})",
8029 p.kv_cache.seq_len()
8030 );
8031 }
8032
8033 #[test]
8034 fn confidence_matches_tokens_and_is_a_probability() {
8035 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
8036 p.sampler_config.temperature = 0.0;
8037 p.sampler_config.repetition_penalty = 1.0;
8038 let r = p.generate("abcd", 10, None, None).unwrap();
8039 assert_eq!(
8040 r.token_confidence.len(),
8041 r.token_ids.len(),
8042 "one confidence per emitted token"
8043 );
8044 for &c in &r.token_confidence {
8045 assert!((0.0..=1.0).contains(&c), "confidence out of [0,1]: {c}");
8046 }
8047 let logits = [1.0f32, 3.0, 0.5, 3.0];
8049 let p0 = top1_prob_t(&logits, 1, 1.0);
8050 let p1 = top1_prob_t(&logits, 3, 1.0);
8051 assert!((p0 - p1).abs() < 1e-6, "equal logits → equal prob");
8052 assert!(p0 > 0.0 && p0 < 1.0);
8053 let sharp = top1_prob_t(&logits, 1, 1.0);
8055 let soft = top1_prob_t(&logits, 1, 2.0);
8056 assert!(soft < sharp, "higher temperature lowers peak confidence");
8057 }
8058
8059 #[test]
8060 fn trace_is_opt_in_and_parallels_the_output() {
8061 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
8063 p.sampler_config.temperature = 0.0;
8064 p.sampler_config.repetition_penalty = 1.0;
8065 let r = p.generate("abcd", 10, None, None).unwrap();
8066 assert!(r.traces.is_empty(), "trace must be empty unless enabled");
8067
8068 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
8070 p.sampler_config.temperature = 0.0;
8071 p.sampler_config.repetition_penalty = 1.0;
8072 p.set_trace(true);
8073 let r = p.generate("abcd", 10, None, None).unwrap();
8074 assert_eq!(r.traces.len(), r.token_ids.len(), "one trace row per token");
8075 for (i, tr) in r.traces.iter().enumerate() {
8076 assert_eq!(tr.t, i, "trace index is sequential");
8077 assert_eq!(tr.token_id, r.token_ids[i], "trace token_id matches output");
8078 assert_eq!(
8079 tr.confidence, r.token_confidence[i],
8080 "trace confidence matches the confidence channel"
8081 );
8082 assert!(tr.active_skill.is_none() && tr.recon.is_none() && !tr.switched);
8084 }
8085 }
8086
8087 #[test]
8088 fn explain_prefill_logits_match_greedy_first_token() {
8089 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
8093 p.sampler_config.temperature = 0.0;
8094 p.sampler_config.repetition_penalty = 1.0;
8095 let ids = p.tokenizer.encode("abcd");
8096 let logits = p.prefill_next_logits(&ids, None);
8097 let argmax = logits
8098 .iter()
8099 .enumerate()
8100 .max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
8101 .unwrap()
8102 .0 as u32;
8103 let r = p.generate("abcd", 1, None, None).unwrap();
8104 assert_eq!(
8105 argmax, r.token_ids[0],
8106 "explain preview must match greedy emit"
8107 );
8108 }
8109
8110 #[test]
8111 fn laguna_shared_expert_is_unconditionally_added() {
8112 let matrix = |values: Vec<f32>| QTensor::from_f32(values, 2, 2);
8113 let identity = || matrix(vec![1.0, 0.0, 0.0, 1.0]);
8114 let zero_dense = || DenseFfn {
8115 gate_proj: matrix(vec![0.0; 4]),
8116 up_proj: matrix(vec![0.0; 4]),
8117 down_proj: matrix(vec![0.0; 4]),
8118 act: Act::Silu,
8119 };
8120 let shared = DenseFfn {
8121 gate_proj: identity(),
8122 up_proj: identity(),
8123 down_proj: identity(),
8124 act: Act::Silu,
8125 };
8126 let x = [1.0, 2.0];
8127 let expected = dense_ffn(&shared, &x, None);
8128 let moe = MoeFfn {
8129 router: QTensor::from_f32(vec![0.0, 0.0], 1, 2),
8130 experts: vec![zero_dense()],
8131 top_k: 1,
8132 norm_topk_prob: true,
8133 router_sigmoid: true,
8134 expert_bias: None,
8135 routed_scaling: 1.0,
8136 route_tau: None,
8137 shared: Some((shared, None)),
8138 stats: std::cell::RefCell::new(Vec::new()),
8139 act_sq: std::cell::RefCell::new(Vec::new()),
8140 act_rows: std::cell::RefCell::new(Vec::new()),
8141 mask: None,
8142 per_expert_scale: None,
8143 router_input_norm: false,
8144 };
8145 let actual = moe_ffn_cpu(&moe, &x, &[0], &[0.0], 1.0, None);
8146 for (actual, expected) in actual.iter().zip(expected) {
8147 assert!((actual - expected).abs() < 1e-6);
8148 }
8149 }
8150}