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 = crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Prefill);
697 if !graph_on || !crate::gpu::enabled_here() {
698 return false;
699 }
700 self.weights
701 .layers
702 .iter()
703 .any(|lw| matches!(&lw.attn, AttnKind::LinearGdn(_)))
704 }
705
706 #[cfg(target_os = "macos")]
707 fn q1_graph_gpu(
708 &mut self,
709 start: usize,
710 upto: Option<usize>,
711 position: usize,
712 h: &mut [f32],
713 ) -> usize {
714 use crate::gpu::{AttnGpuLayer, GdnGpuCfg, GdnGpuLayer, GraphDims, MetalFfn, TokenGraph};
715 if self.attn_softcap > 0.0 || !crate::gpu::enabled_here()
717 || !crate::gpu::q1_force()
718 || std::env::var("CMF_GPU_BLOCK")
719 .map(|v| v == "0")
720 .unwrap_or(false)
721 {
722 if std::env::var("CMF_GRAPH_DBG").is_ok() {
723 eprintln!(
724 "block-graph: front gate (softcap={} enabled_here={} q1_force={})",
725 self.attn_softcap > 0.0,
726 crate::gpu::enabled_here(),
727 crate::gpu::q1_force(),
728 );
729 }
730 return start;
731 }
732 if self.swa.is_some()
737 || self.global_attn.is_some()
738 || self.attention_heads_per_layer.is_some()
739 || self.attn_v_norm
740 || (self.attn_scale - 1.0 / (self.head_dim as f32).sqrt()).abs() > 1e-9
741 || self.weights.layers.iter().any(|lw| {
742 lw.attn_out_norm.is_some()
743 || lw.ffn_out_norm.is_some()
744 || lw.layer_scale.is_some()
745 || matches!(&lw.ffn, FfnKind::Dense(d) if d.act != Act::Silu)
746 })
747 {
748 if std::env::var("CMF_GRAPH_DBG").is_ok() {
749 eprintln!(
750 "block-graph: arch ineligible (swa={} gattn={} hpl={} vnorm={} scale_delta={:.2e})",
751 self.swa.is_some(),
752 self.global_attn.is_some(),
753 self.attention_heads_per_layer.is_some(),
754 self.attn_v_norm,
755 (self.attn_scale - 1.0 / (self.head_dim as f32).sqrt()).abs(),
756 );
757 }
758 return start;
759 }
760 let limit = upto
763 .map(|u| u + 1)
764 .unwrap_or(self.num_layers)
765 .min(self.num_layers);
766
767 enum Item<'a> {
768 Gdn {
769 run: Vec<GdnGpuLayer<'a>>,
770 first: usize,
771 },
772 Attn {
773 l: AttnGpuLayer<'a>,
774 li: usize,
775 q_norm: Option<&'a [f32]>,
776 k_norm: Option<&'a [f32]>,
777 output_gate: bool,
778 bias: Option<(&'a [f32], &'a [f32], &'a [f32])>,
779 full_gpu: bool,
782 },
783 }
784
785 let attend_mode = std::env::var("CMF_GPU_ATTEND").unwrap_or_else(|_| "auto".into());
792 let attend_contract = attend_mode != "0"
793 && attend_mode != "off"
794 && self.head_dim % 4 == 0
795 && self.head_dim <= 256
796 && self.rotary_dim >= 2
797 && self.rotary_dim <= self.head_dim
798 && (self.rotary_dim / 2) % 32 == 0
799 && self.num_kv_heads > 0
800 && self.num_heads % self.num_kv_heads == 0;
801
802 let mut plan: Vec<Item> = Vec::new();
803 let mut model_ref: Option<std::sync::Arc<cortiq_core::CmfModel>> = None;
804 let block_diag = std::env::var("CMF_GRAPH_DBG").is_ok();
806 let mut scan = start;
807 while scan < limit {
808 let lw = &self.weights.layers[self.phys_layer(scan)];
809 let ffn = match &lw.ffn {
810 FfnKind::Dense(d) => {
811 let (Some(g), Some(u), Some(dn)) = (
812 d.gate_proj.q1_parts(),
813 d.up_proj.q1_parts(),
814 d.down_proj.q1_parts(),
815 ) else {
816 if block_diag {
817 eprintln!(
818 "block-graph: L{scan} FFN trio not graph-mappable — run ends"
819 );
820 }
821 break;
822 };
823 MetalFfn::Dense {
824 gate: g,
825 up: u,
826 down: dn,
827 }
828 }
829 FfnKind::Moe(m) => {
830 let Some(moe) = metal_moe_graph_parts(m, self.hidden_size) else {
831 if block_diag {
832 eprintln!(
833 "block-graph: L{scan} MoE outside the graph contract — run ends"
834 );
835 }
836 break;
837 };
838 if let QTensor::Mapped { model, .. } = &m.experts[0].gate_proj {
839 model_ref.get_or_insert_with(|| model.clone());
840 }
841 MetalFfn::Moe(moe)
842 }
843 _ => {
844 if block_diag {
845 eprintln!("block-graph: L{scan} non-graph FFN — run ends");
846 }
847 break;
848 }
849 };
850 match &lw.attn {
851 AttnKind::LinearGdn(w) if self.gdn_cfg.is_some() => {
852 let parts = (
853 w.in_proj_qkv.q1_parts(),
854 w.in_proj_z.q1_parts(),
855 w.in_proj_a.f32_parts(),
856 w.in_proj_b.f32_parts(),
857 w.out_proj.q1_parts(),
858 );
859 let (Some(qkv), Some(z), Some(a), Some(b), Some(out)) = parts else {
860 if block_diag {
861 eprintln!(
862 "block-graph: L{scan} GDN parts refused (qkv={} z={} a_f32={} b_f32={} out={})",
863 w.in_proj_qkv.q1_parts().is_some(),
864 w.in_proj_z.q1_parts().is_some(),
865 w.in_proj_a.f32_parts().is_some(),
866 w.in_proj_b.f32_parts().is_some(),
867 w.out_proj.q1_parts().is_some(),
868 );
869 }
870 break;
871 };
872 if let QTensor::Mapped { model, .. } = &w.in_proj_qkv {
873 model_ref.get_or_insert_with(|| model.clone());
874 }
875 let gl = GdnGpuLayer {
876 attn_norm: &lw.input_norm,
877 post_norm: &lw.post_norm,
878 qkv,
879 z,
880 a,
881 b,
882 out,
883 ffn,
884 conv1d: &w.conv1d,
885 a_log: &w.a_log,
886 dt_bias: &w.dt_bias,
887 gnorm: &w.norm,
888 };
889 match plan.last_mut() {
890 Some(Item::Gdn { run, .. }) => run.push(gl),
891 _ => plan.push(Item::Gdn {
892 run: vec![gl],
893 first: scan,
894 }),
895 }
896 }
897 AttnKind::Full {
898 wq,
899 wk,
900 wv,
901 wo,
902 q_norm,
903 k_norm,
904 output_gate,
905 softplus_gate: None,
906 bias,
907 } if !self.kv_cache.layers[scan].o1_sealed() => {
908 let parts = (wq.q1_parts(), wk.q1_parts(), wv.q1_parts(), wo.q1_parts());
909 let (Some(pq), Some(pk), Some(pv), Some(po)) = parts else {
910 break;
911 };
912 if let QTensor::Mapped { model, .. } = wq {
913 model_ref.get_or_insert_with(|| model.clone());
914 }
915 let cache = &self.kv_cache.layers[scan];
916 let full_gpu = attend_contract
917 && cache.mode == crate::kv_cache::KvMode::F32
918 && cache.o1.is_none()
919 && bias.is_none()
920 && pq.1 == self.num_heads * self.head_dim * (1 + *output_gate as usize)
921 && pk.1 == self.num_kv_heads * self.head_dim
922 && pv.1 == self.num_kv_heads * self.head_dim
923 && po.2 == self.num_heads * self.head_dim;
924 plan.push(Item::Attn {
925 l: AttnGpuLayer {
926 attn_norm: &lw.input_norm,
927 post_norm: &lw.post_norm,
928 wq: pq,
929 wk: pk,
930 wv: pv,
931 wo: po,
932 ffn,
933 },
934 li: scan,
935 q_norm: q_norm.as_deref(),
936 k_norm: k_norm.as_deref(),
937 output_gate: *output_gate,
938 bias: bias
939 .as_ref()
940 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
941 full_gpu,
942 });
943 }
944 _ => break,
945 }
946 scan += 1;
947 }
948 let Some(model) = model_ref else {
949 if std::env::var("CMF_GRAPH_DBG").is_ok() {
950 eprintln!("q1-graph: no model ref (start {start}, scanned to {scan})");
951 }
952 return start;
953 };
954 if plan.is_empty() {
955 if std::env::var("CMF_GRAPH_DBG").is_ok() {
956 eprintln!("q1-graph: empty plan at layer {start}");
957 }
958 return start;
959 }
960 let has_moe = plan.iter().any(|it| match it {
961 Item::Gdn { run, .. } => run.iter().any(|l| matches!(l.ffn, MetalFfn::Moe(_))),
962 Item::Attn { l, .. } => matches!(l.ffn, MetalFfn::Moe(_)),
963 });
964 let dev_attend = attend_contract
965 && (self.head_dim <= 128
966 || has_moe
967 || attend_mode == "force"
968 || attend_mode == "256");
969 if !dev_attend {
970 for it in &mut plan {
971 if let Item::Attn { full_gpu, .. } = it {
972 *full_gpu = false;
973 }
974 }
975 }
976 if std::env::var("CMF_GRAPH_DBG").is_ok() {
977 use std::sync::atomic::{AtomicBool, Ordering};
978 static SAID: AtomicBool = AtomicBool::new(false);
979 if !SAID.swap(true, Ordering::Relaxed) {
980 let fg = plan
981 .iter()
982 .filter(|it| matches!(it, Item::Attn { full_gpu: true, .. }))
983 .count();
984 let att = plan
985 .iter()
986 .filter(|it| matches!(it, Item::Attn { .. }))
987 .count();
988 eprintln!(
989 "q1-graph: plan of {} items from layer {start} to {scan} | dev_attend={dev_attend} full_gpu {fg}/{att} | hd={} rd={} nkv={} nh={}",
990 plan.len(),
991 self.head_dim,
992 self.rotary_dim,
993 self.num_kv_heads,
994 self.num_heads,
995 );
996 }
997 }
998 let dims = GraphDims {
999 hidden: self.hidden_size,
1000 eps: self.rms_eps as f32,
1001 gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
1002 };
1003 let Some(mut graph) = TokenGraph::new(&model, dims, h) else {
1004 return start;
1005 };
1006 let gcfg = self.gdn_cfg.map(|cfg| GdnGpuCfg {
1007 nv: cfg.num_v_heads,
1008 nk: cfg.num_k_heads,
1009 dk: cfg.key_head_dim,
1010 dv: cfg.value_head_dim,
1011 kk: cfg.conv_kernel,
1012 hidden: self.hidden_size,
1013 inter: self.intermediate_size,
1014 c_dim: cfg.conv_dim(),
1015 eps: cfg.rms_eps as f32,
1016 gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
1017 });
1018 let mut valid = 0usize;
1022 let mut end = start;
1023 for item in &plan {
1024 let ok = match item {
1025 Item::Gdn { run, .. } => gcfg
1026 .as_ref()
1027 .map(|gc| run.iter().all(|l| graph.gdn_ok(l, gc)))
1028 .unwrap_or(false),
1029 Item::Attn { l, .. } => graph.attn_ok(l),
1030 };
1031 if !ok {
1032 if block_diag {
1033 eprintln!(
1034 "block-graph: plan item {} ({}) failed graph preflight",
1035 valid,
1036 match item {
1037 Item::Gdn { run, first } =>
1038 format!("GDN run L{first}+{}", run.len()),
1039 Item::Attn { li, .. } => format!("Attn L{li}"),
1040 }
1041 );
1042 }
1043 break;
1044 }
1045 valid += 1;
1046 end += match item {
1047 Item::Gdn { run, .. } => run.len(),
1048 Item::Attn { .. } => 1,
1049 };
1050 }
1051 plan.truncate(valid);
1052 if plan.is_empty() {
1053 return start;
1054 }
1055
1056 let inv_freq = self.inv_freq.clone();
1057 let pool = self.pool.clone();
1058 let (nh, nkv, hd, hs, rd, eps) = (
1059 self.num_heads,
1060 self.num_kv_heads,
1061 self.head_dim,
1062 self.hidden_size,
1063 self.rotary_dim,
1064 self.rms_eps,
1065 );
1066 let norm_style = self.norm_style;
1067 let gemma = norm_style == cortiq_core::NormStyle::Gemma;
1068 let want = self.gdn_cfg.map(|c| c.state_len()).unwrap_or(0);
1069 let kv_id = self.graph_kv_id;
1070 let mut pending: Vec<(usize, usize)> = Vec::new();
1073 let mut dev_attn: Vec<usize> = Vec::new();
1076 for item in &plan {
1077 if self.loop_final_norm {
1079 let item_start = match item {
1080 Item::Gdn { first, .. } => *first,
1081 Item::Attn { li, .. } => *li,
1082 };
1083 if item_start > start && self.is_loop_end(item_start - 1) {
1084 graph.encode_loop_norm(&self.weights.final_norm);
1085 }
1086 }
1087 match item {
1088 Item::Gdn { run, first } => {
1089 for l in &mut self.kv_cache.layers[*first..*first + run.len()] {
1090 if l.linear_state.len() != want {
1091 l.linear_state = vec![0f32; want];
1092 }
1093 }
1094 let ro: Vec<&[f32]> = self.kv_cache.layers[*first..*first + run.len()]
1095 .iter()
1096 .map(|l| l.linear_state.as_slice())
1097 .collect();
1098 if !graph.encode_gdn_run(run, &ro, gcfg.as_ref().unwrap()) {
1099 tracing::error!("q1 graph: GDN run refused after validation");
1101 return start;
1102 }
1103 graph.commit();
1106 pending.push((*first, run.len()));
1107 }
1108 Item::Attn {
1109 l,
1110 li,
1111 q_norm,
1112 k_norm,
1113 output_gate,
1114 bias,
1115 full_gpu,
1116 } => {
1117 if *full_gpu {
1119 let cache = &self.kv_cache.layers[*li];
1120 let cpu_k: Vec<&[f32]> = (0..nkv).map(|g| cache.head_keys(g)).collect();
1121 let cpu_v: Vec<&[f32]> = (0..nkv).map(|g| cache.head_values(g)).collect();
1122 let cpu_stored = cpu_k[0].len() / hd;
1123 let p = crate::gpu::AttnDeviceParams {
1124 kv_id,
1125 layer: *li,
1126 nh,
1127 nkv,
1128 hd,
1129 rd,
1130 position,
1131 eps: eps as f32,
1132 gemma,
1133 output_gate: *output_gate,
1134 q_norm: *q_norm,
1135 k_norm: *k_norm,
1136 inv_freq: &inv_freq,
1137 cpu_k,
1138 cpu_v,
1139 cpu_stored,
1140 };
1141 if graph.attn_device_ok(l, &p) && graph.encode_attn_device(l, &p) {
1142 graph.commit();
1143 dev_attn.push(*li);
1144 continue;
1145 }
1146 }
1148 graph.encode_attn_prefix(l);
1149 graph.sync();
1150 if !pending.is_empty() {
1151 let idxs: Vec<usize> =
1152 pending.drain(..).flat_map(|(f, n)| f..f + n).collect();
1153 let mut outs: Vec<&mut [f32]> = self
1154 .kv_cache
1155 .layers
1156 .iter_mut()
1157 .enumerate()
1158 .filter(|(i, _)| idxs.binary_search(i).is_ok())
1159 .map(|(_, s)| s.linear_state.as_mut_slice())
1160 .collect();
1161 graph.read_states(&mut outs);
1162 }
1163 let mut q_raw = attention::take_buf(l.wq.1);
1164 let mut k = attention::take_buf(l.wk.1);
1165 let mut v = attention::take_buf(l.wv.1);
1166 graph.read_qkv(&mut q_raw, &mut k, &mut v);
1167 let cfg = QwenAttnCfg {
1168 num_heads: nh,
1169 num_kv_heads: nkv,
1170 head_dim: hd,
1171 hidden_size: hs,
1172 position,
1173 inv_freq: &inv_freq,
1174 rotary_dim: rd,
1175 scale: self.attn_scale,
1176 softcap: self.attn_softcap,
1177 window: None,
1178 v_norm: false,
1179 q_norm: *q_norm,
1180 k_norm: *k_norm,
1181 output_gate: *output_gate,
1182 softplus_gate: None,
1183 rope_scale: 1.0,
1184 bias: *bias,
1185 rms_eps: eps,
1186 norm_style,
1187 pool: pool.as_deref(),
1188 };
1189 let mut ao = attention::qwen_attention_core(
1190 q_raw,
1191 k,
1192 v,
1193 &mut self.kv_cache.layers[*li],
1194 &cfg,
1195 );
1196 graph.encode_attn_suffix(l, &ao);
1197 graph.commit();
1200 attention::recycle_buf(&mut ao);
1201 }
1202 }
1203 }
1204 let mut lm_rows = None;
1209 if self.graph_want_logits
1210 && upto.is_none()
1211 && end == self.num_layers
1212 && std::env::var("CMF_GPU_LMHEAD")
1213 .map(|v| v != "0")
1214 .unwrap_or(true)
1215 {
1216 if let Some(lm) = self.weights.lm_head.q1_parts() {
1217 if graph.lm_head_ok(lm) {
1218 graph.encode_lm_head(&self.weights.final_norm, lm);
1219 lm_rows = Some(lm.1);
1220 }
1221 }
1222 }
1223 graph.sync();
1224 if !pending.is_empty() {
1225 let idxs: Vec<usize> = pending.drain(..).flat_map(|(f, n)| f..f + n).collect();
1226 let mut outs: Vec<&mut [f32]> = self
1227 .kv_cache
1228 .layers
1229 .iter_mut()
1230 .enumerate()
1231 .filter(|(i, _)| idxs.binary_search(i).is_ok())
1232 .map(|(_, s)| s.linear_state.as_mut_slice())
1233 .collect();
1234 graph.read_states(&mut outs);
1235 }
1236 if let Some(rows) = lm_rows {
1237 let mut lg = attention::take_buf(rows.min(self.vocab_size));
1238 graph.read_logits(&mut lg);
1239 lg.resize(self.vocab_size, 0.0);
1240 if let Some(c) = self.final_softcap {
1241 for l in lg.iter_mut() {
1242 *l = c * (*l / c).tanh();
1243 }
1244 }
1245 self.graph_logits = Some(lg);
1246 }
1247 graph.finish(h);
1248 for li in dev_attn {
1252 let mut krow = attention::take_buf(nkv * hd);
1253 let mut vrow = attention::take_buf(nkv * hd);
1254 if crate::gpu::kv_mirror_read_last(kv_id, li, nkv, hd, &mut krow, &mut vrow) {
1255 let cache = &mut self.kv_cache.layers[li];
1256 cache.append(&krow, &vrow, &[]);
1257 let n = cache.seq_len;
1258 let mut imp = attention::take_buf(n);
1259 crate::gpu::kv_mirror_take_imp(kv_id, li, &mut imp);
1260 cache.accumulate_imp(&imp);
1261 attention::recycle_buf(&mut imp);
1262 }
1263 attention::recycle_buf(&mut krow);
1264 attention::recycle_buf(&mut vrow);
1265 }
1266 end
1267 }
1268
1269 pub fn new(
1270 tokenizer: Tokenizer,
1271 weights: PipelineWeights,
1272 hidden_size: usize,
1273 intermediate_size: usize,
1274 num_heads: usize,
1275 num_kv_heads: usize,
1276 head_dim: usize,
1277 num_layers: usize,
1278 physical_layers: usize,
1279 loop_final_norm: bool,
1280 vocab_size: usize,
1281 rms_eps: f64,
1282 rope_base: f32,
1283 norm_style: NormStyle,
1284 max_seq_len: usize,
1285 sampler_config: SamplerConfig,
1286 ) -> Self {
1287 let rng = match sampler_config.seed {
1288 Some(s) => SplitMix64::new(s),
1289 None => SplitMix64::from_entropy(),
1290 };
1291 let inv_freq = std::sync::Arc::new(attention::rope_inv_freq(head_dim, rope_base));
1292 let pool = Pool::from_env();
1293 if let Some(p) = &pool {
1294 tracing::info!("worker pool: {} threads", p.n_workers());
1295 }
1296 Self {
1297 gpu_plan: None,
1298 tokenizer: std::sync::Arc::new(tokenizer),
1299 kv_cache: KvCache::new(num_layers, num_kv_heads, head_dim, max_seq_len),
1300 sampler_config,
1301 weights,
1302 hidden_size,
1303 intermediate_size,
1304 num_heads,
1305 num_kv_heads,
1306 head_dim,
1307 num_layers,
1308 physical_layers,
1309 loop_final_norm,
1310 vocab_size,
1311 rms_eps,
1312 rope_base,
1313 norm_style,
1314 rotary_dim: head_dim,
1315 attention_heads_per_layer: None,
1316 vmf_cfg: None,
1317 gdn_cfg: None,
1318 kda_cfg: None,
1319 g3n: None,
1320 dsv4: None,
1321 dsv4_mtp: Vec::new(),
1322 dspark: None,
1323 dspark_pending: Vec::new(),
1324 dspark_hist: Vec::new(),
1325 dspark_real: Vec::new(),
1326 dspark_trunk_picks: Vec::new(),
1327 dspark_exp: Vec::new(),
1328 dspark_draft_ns: 0,
1329 logit_multiplier: None,
1330 cancel: std::sync::Arc::new(std::sync::atomic::AtomicBool::new(false)),
1331 kv_history: Vec::new(),
1332 short_conv_cfg: None,
1333 mtp: None,
1334 speculative: std::env::var("CMF_MTP").map(|v| v != "0").unwrap_or(true),
1335 rng,
1336 sampler_scratch: SamplerScratch::default(),
1337 inv_freq,
1338 ws: ForwardScratch::new(hidden_size),
1339 pool,
1340 model: None,
1341 dyn_force_f32: false,
1342 dyn_skill_layers: Vec::new(),
1343 dyn_active: None,
1344 dyn_blend_loaded: false,
1345 dyn_phi_layer: None,
1346 dyn_phi_ema: Vec::new(),
1347 dyn_phi_seen: 0,
1348 dyn_router: None,
1349 o1_cfg: None,
1350 o1_epoch: 0,
1351 o1_flags: Vec::new(),
1352 trace: false,
1353 calib_temp: 1.0,
1354 confidence_on: true,
1355 embed_multiplier: 1.0,
1356 attn_scale: 1.0 / (head_dim as f32).sqrt(),
1357 swa: None,
1358 sliding_layers: None,
1359 inv_freq_local: None,
1360 rotary_dim_local: None,
1361 rope_scale: 1.0,
1362 rope_scale_local: 1.0,
1363 global_attn: None,
1364 inv_freq_global: None,
1365 attn_v_norm: false,
1366 final_softcap: None,
1367 attn_softcap: 0.0,
1368 graph_want_logits: false,
1369 graph_logits: None,
1370 graph_kv_id: {
1371 static NEXT: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(1);
1372 NEXT.fetch_add(1, std::sync::atomic::Ordering::Relaxed)
1373 },
1374 }
1375 }
1376
1377 pub fn set_o1(&mut self, cfg: Option<crate::nystrom::O1Cfg>) {
1384 self.o1_flags = match &cfg {
1385 Some(c) => {
1386 let mut flags = c.layer_flags(self.num_layers);
1387 for (li, f) in flags.iter_mut().enumerate() {
1388 if *f
1389 && !matches!(
1390 self.weights.layers[self.phys_layer(li)].attn,
1391 AttnKind::Full { .. }
1392 )
1393 {
1394 *f = false;
1395 }
1396 }
1397 flags
1398 }
1399 None => Vec::new(),
1400 };
1401 if let Some(c) = &cfg {
1402 let n = self.o1_flags.iter().filter(|&&f| f).count();
1403 tracing::info!(
1404 "o1 nystrom attention: {n}/{} layer(s), m={} w={} sink={} rect={:?}",
1405 self.num_layers,
1406 c.m,
1407 c.w,
1408 c.sink,
1409 c.rect
1410 );
1411 }
1412 self.o1_cfg = cfg;
1413 }
1414
1415 pub fn o1_active(&self) -> bool {
1417 self.o1_cfg.is_some() && self.o1_flags.iter().any(|&f| f)
1418 }
1419
1420 pub fn o1_begin(&mut self) {
1425 if let Some(c) = &self.o1_cfg {
1426 let (m, w, sink, rect) = (c.m, c.w, c.sink, c.rect);
1427 for (li, &f) in self.o1_flags.iter().enumerate() {
1428 if f {
1429 self.kv_cache.layers[li].o1_begin(m, w, sink, rect);
1430 }
1431 }
1432 }
1433 }
1434
1435 pub fn o1_seal(&mut self) {
1439 self.o1_epoch = self.o1_epoch.wrapping_add(1);
1440 if self.o1_cfg.is_none() {
1441 return;
1442 }
1443 for li in 0..self.num_layers {
1444 if self.o1_flags.get(li).copied().unwrap_or(false) {
1445 self.kv_cache.layers[li].o1_seal(self.num_heads);
1446 }
1447 }
1448 }
1449
1450 pub fn set_trace(&mut self, on: bool) {
1452 self.trace = on;
1453 }
1454
1455 pub fn set_sampler_config(&mut self, config: SamplerConfig) {
1458 self.rng = match config.seed {
1459 Some(seed) => SplitMix64::new(seed),
1460 None => SplitMix64::from_entropy(),
1461 };
1462 self.sampler_config = config;
1463 }
1464
1465 pub fn set_confidence(&mut self, on: bool) {
1470 self.confidence_on = on;
1471 }
1472
1473 pub fn set_calib_temp(&mut self, t: f32) {
1476 self.calib_temp = if t > 1e-3 { t } else { 1.0 };
1477 }
1478
1479 pub fn calib_temp(&self) -> f32 {
1481 self.calib_temp
1482 }
1483
1484 pub fn set_rotary(&mut self, rotary_dim: usize, base: f32) {
1487 self.rotary_dim = rotary_dim.min(self.head_dim);
1488 self.inv_freq = std::sync::Arc::new(attention::rope_inv_freq(self.rotary_dim, base));
1489 }
1490
1491 fn attn_cfg(&self, position: usize) -> QwenAttnCfg<'_> {
1492 QwenAttnCfg {
1493 num_heads: self.num_heads,
1494 num_kv_heads: self.num_kv_heads,
1495 head_dim: self.head_dim,
1496 hidden_size: self.hidden_size,
1497 position,
1498 inv_freq: &self.inv_freq,
1499 rotary_dim: self.rotary_dim,
1500 scale: self.attn_scale,
1501 softcap: self.attn_softcap,
1502 window: None,
1503 v_norm: false,
1504 q_norm: None,
1505 k_norm: None,
1506 output_gate: false,
1507 softplus_gate: None,
1508 rope_scale: self.rope_scale,
1509 bias: None,
1510 rms_eps: self.rms_eps,
1511 norm_style: self.norm_style,
1512 pool: self.pool.as_deref(),
1513 }
1514 }
1515
1516 pub fn generate(
1518 &mut self,
1519 prompt: &str,
1520 max_tokens: usize,
1521 task_mask: Option<&TaskMask>,
1522 on_token: Option<TokenCallback>,
1523 ) -> Result<GenerateResult, String> {
1524 let input_ids = self.tokenizer.with_bos(self.tokenizer.encode(prompt));
1525 self.generate_from_ids(&input_ids, max_tokens, task_mask, on_token)
1526 }
1527
1528 pub fn generate_from_ids(
1536 &mut self,
1537 input_ids: &[u32],
1538 max_tokens: usize,
1539 task_mask: Option<&TaskMask>,
1540 mut on_token: Option<TokenCallback>,
1541 ) -> Result<GenerateResult, String> {
1542 if std::env::var("CMF_TRACE_H").is_ok() {
1543 eprintln!("input_ids: {input_ids:?}");
1544 }
1545 if input_ids.is_empty() {
1546 return Err("empty prompt: nothing to generate from".to_string());
1547 }
1548
1549 let reuse_from = {
1557 let on = !std::env::var("CMF_KV_REUSE").is_ok_and(|v| v == "0");
1558 let h = &self.kv_history;
1559 if on
1560 && task_mask.is_none()
1561 && self.mtp.is_none()
1562 && self.o1_cfg.is_none()
1563 && !h.is_empty()
1564 && h.len() < input_ids.len()
1565 && input_ids[..h.len()] == h[..]
1566 {
1567 h.len()
1568 } else {
1569 0
1570 }
1571 };
1572 if reuse_from == 0 {
1573 self.kv_cache.clear();
1575 self.kv_history.clear();
1576 crate::gpu::graph_kv_reset(self.graph_kv_id);
1577 } else if std::env::var("CMF_PREFILL_PROF").is_ok() {
1578 eprintln!(
1579 "kv-reuse: {} of {} prompt positions already cached",
1580 reuse_from,
1581 input_ids.len()
1582 );
1583 }
1584 crate::gpu::graph_race_begin_generation();
1585 self.o1_begin();
1586
1587 let graph_on = crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Decode);
1593 let graph_spec = self.speculative
1609 && graph_on
1610 && self.mtp.is_some()
1611 && task_mask.is_none()
1612 && !self.o1_active()
1613 && self.sampler_config.temperature < 1e-6
1614 && self.sampler_config.repetition_penalty == 1.0
1615 && std::env::var("CMF_GRAPH_SPEC").is_ok_and(|v| v != "0");
1616 let pair_pays = self.gdn_cfg.is_none()
1623 || std::env::var("CMF_MTP").as_deref() == Ok("1");
1624 let spec_active = self.speculative
1625 && self.mtp.is_some()
1626 && task_mask.is_none()
1627 && !self.o1_active()
1628 && ((!graph_on && pair_pays) || graph_spec)
1629 && self.sampler_config.temperature < 1e-6;
1630 let mut mtp = if spec_active { self.mtp.take() } else { None };
1633 if std::env::var("CMF_MTP_CHAIN_PROBE").is_ok() {
1634 eprintln!(
1635 "mtp-probe gate: spec_active={spec_active} mtp={} speculative={} graph_on={graph_on} temp_ok={}",
1636 mtp.is_some(),
1637 self.speculative,
1638 self.sampler_config.temperature < 1e-6,
1639 );
1640 }
1641 if let Some(m) = &mut mtp {
1642 m.kv.clear();
1643 }
1644 let mut router = if mtp.is_none() {
1648 self.dyn_router.take()
1649 } else {
1650 None
1651 };
1652 if let Some(r) = &mut router {
1653 r.reset(); self.dyn_phi_seen = 0; let _ = self.set_active_skill(None);
1656 }
1657
1658 let mut all_ids = input_ids.to_vec();
1659 let mut generated = 0usize;
1660 let mut finish_reason = "max_tokens".to_string();
1661 let mut drafted = 0usize;
1662 let mut accepted = 0usize;
1663 let mut confidence: Vec<f32> = Vec::new();
1664 let trace_on = self.trace;
1665 let calib_temp = self.calib_temp;
1666 let mut traces: Vec<TokenTrace> = Vec::new();
1667
1668 let mut hidden = vec![0.0f32; self.hidden_size];
1674 let mut pos = reuse_from;
1675 let fuse_lm = mtp.is_none()
1684 && router.is_none()
1685 && std::env::var("CMF_GPU_LMHEAD").as_deref() != Ok("0");
1686 self.graph_logits = None;
1687 self.graph_want_logits = false;
1688 let _tpf = std::time::Instant::now();
1689 let batch_k = std::env::var("CMF_BATCH_K")
1690 .ok()
1691 .and_then(|v| v.parse::<usize>().ok())
1692 .unwrap_or(0);
1693 while self.dsv4.is_some()
1704 && mtp.is_none()
1705 && pos < input_ids.len()
1706 && !self.cancel.load(std::sync::atomic::Ordering::Relaxed)
1707 {
1708 let end = (pos + prefill_chunk()).min(input_ids.len());
1709 let ids: Vec<u32> = input_ids[pos..end].to_vec();
1710 let mut lg = Vec::new();
1711 if let Some(b) = &mut self.dsv4 {
1712 let (g, layers, cfg, st) = (&b.0, &b.1, b.2, &mut b.3);
1713 crate::dsv4::forward_chunk(
1714 g,
1715 layers,
1716 &cfg,
1717 st,
1718 &ids,
1719 pos,
1720 &self.inv_freq,
1721 self.pool.as_deref(),
1722 &mut lg,
1723 end == input_ids.len(),
1724 );
1725 }
1726 if end == input_ids.len() {
1727 self.graph_logits = Some(lg);
1728 }
1729 pos = end;
1730 hidden = vec![0.0; self.hidden_size];
1731 }
1732 let dyn_prefill = router.is_some();
1737 let graph_prefill = self.graph_prefill_preferred();
1743 if task_mask.is_none()
1744 && !dyn_prefill
1745 && !graph_prefill
1746 && self.can_prefill_batched()
1747 && self.g3n.is_none()
1748 && input_ids.len() > 2
1749 {
1750 let chunk = prefill_chunk();
1756 let hs = self.hidden_size;
1757 while pos < input_ids.len() && !self.cancel.load(std::sync::atomic::Ordering::Relaxed) {
1758 let end = (pos + chunk).min(input_ids.len());
1759 let hb = self.prefill_batch(&input_ids[pos..end], pos);
1760 if let Some(m) = &mut mtp {
1761 let probe: usize = std::env::var("CMF_MTP_CHAIN_PROBE")
1762 .ok()
1763 .and_then(|v| v.parse().ok())
1764 .unwrap_or(0);
1765 for p in pos..end {
1766 if p + 1 < input_ids.len() {
1767 if probe >= 1 && p + 2 < input_ids.len() {
1768 let (d1, mut hx) = self.mtp_step_h(
1772 m,
1773 &hb[(p - pos) * hs..(p - pos + 1) * hs],
1774 input_ids[p + 1],
1775 p,
1776 );
1777 let mut ok = d1 == input_ids[p + 2];
1778 Self::chain_probe_note(0, ok);
1779 let mut d_prev = d1;
1780 let mut extra = 0usize;
1781 for j in 1..probe {
1782 if p + 2 + j >= input_ids.len() {
1783 break;
1784 }
1785 let (dj, hj) =
1786 self.mtp_step_h(m, &hx, d_prev, p + 1 + j);
1787 extra += 1;
1788 ok = ok && dj == input_ids[p + 2 + j];
1789 Self::chain_probe_note(j, ok);
1790 d_prev = dj;
1791 hx = hj;
1792 }
1793 m.kv.truncate_last(extra);
1794 } else {
1795 let _ = self.mtp_step(
1796 m,
1797 &hb[(p - pos) * hs..(p - pos + 1) * hs],
1798 input_ids[p + 1],
1799 p,
1800 );
1801 }
1802 }
1803 }
1804 }
1805 hidden.copy_from_slice(&hb[(end - pos - 1) * hs..]);
1806 pos = end;
1807 }
1808 }
1809 let pair_off = std::env::var("CMF_PAIR").is_ok_and(|v| v == "0");
1810 if task_mask.is_none()
1811 && !dyn_prefill
1812 && !graph_prefill
1813 && !pair_off
1814 && self.pair_supported()
1815 {
1816 while pos + 1 < input_ids.len()
1817 && !self.cancel.load(std::sync::atomic::Ordering::Relaxed)
1818 {
1819 let e1 = self.embed_single(input_ids[pos]);
1820 let e2 = self.embed_single(input_ids[pos + 1]);
1821 let (h1, h2) = self.forward_pair(&e1, &e2, pos);
1822 self.commit_linear_scratch();
1824 if let Some(m) = &mut mtp {
1825 let _ = self.mtp_step(m, &h1, input_ids[pos + 1], pos);
1826 if pos + 2 < input_ids.len() {
1827 let probe: usize = std::env::var("CMF_MTP_CHAIN_PROBE")
1828 .ok()
1829 .and_then(|v| v.parse().ok())
1830 .unwrap_or(0);
1831 if probe >= 1 && pos + 3 < input_ids.len() {
1832 let (d1, mut hx) =
1836 self.mtp_step_h(m, &h2, input_ids[pos + 2], pos + 1);
1837 let mut ok = d1 == input_ids[pos + 3];
1838 Self::chain_probe_note(0, ok);
1839 let mut d_prev = d1;
1840 let mut extra = 0usize;
1841 for j in 1..probe {
1842 if pos + 3 + j >= input_ids.len() {
1843 break;
1844 }
1845 let (dj, hj) =
1846 self.mtp_step_h(m, &hx, d_prev, pos + 2 + j);
1847 extra += 1;
1848 ok = ok && dj == input_ids[pos + 3 + j];
1849 Self::chain_probe_note(j, ok);
1850 d_prev = dj;
1851 hx = hj;
1852 }
1853 m.kv.truncate_last(extra);
1854 } else {
1855 let _ = self.mtp_step(m, &h2, input_ids[pos + 2], pos + 1);
1856 }
1857 }
1858 }
1859 hidden = h2;
1860 pos += 2;
1861 }
1862 }
1863 if batch_k > 0
1872 && graph_prefill
1873 && task_mask.is_none()
1874 && !self.o1_active()
1875 && mtp.is_none()
1876 && !dyn_prefill
1877 && pos + 1 < input_ids.len()
1878 {
1879 let hs = self.hidden_size;
1880 let chunk = batch_k;
1881 while pos < input_ids.len() {
1882 let end = (pos + chunk).min(input_ids.len());
1883 let bk = end - pos;
1884 let mut hiddens = vec![0f32; bk * hs];
1885 for (j, &id) in input_ids[pos..end].iter().enumerate() {
1886 hiddens[j * hs..(j + 1) * hs].copy_from_slice(&self.embed_single(id));
1887 }
1888 let positions: Vec<usize> = (pos..end).collect();
1889 let t_chunk = std::time::Instant::now();
1890 let ok_b = self.try_batch_graph_wgpu(&mut hiddens, &positions, bk, None);
1891 if std::env::var("CMF_GRAPH_PROF").is_ok() {
1892 let ms = t_chunk.elapsed().as_secs_f64() * 1000.0;
1893 eprintln!(
1894 "batch-chunk: k={bk} ok={ok_b} {ms:.1} ms ({:.1} tok/s)",
1895 bk as f64 / (ms / 1000.0)
1896 );
1897 }
1898 {
1899 use std::sync::atomic::{AtomicBool, Ordering};
1900 static SAID: AtomicBool = AtomicBool::new(false);
1901 if !SAID.swap(true, Ordering::Relaxed) {
1902 if ok_b {
1903 tracing::info!("batched prefill: ACTIVE (k={bk})");
1904 } else {
1905 tracing::warn!("batched prefill declined — per-position graph");
1906 }
1907 }
1908 }
1909 if ok_b {
1910 hidden.copy_from_slice(&hiddens[(bk - 1) * hs..]);
1911 pos = end;
1912 } else {
1913 break; }
1915 }
1916 }
1917 while pos < input_ids.len() && !self.cancel.load(std::sync::atomic::Ordering::Relaxed) {
1918 self.graph_want_logits = fuse_lm && pos + 1 == input_ids.len();
1919 hidden = self.forward_layers(&self.embed_single(input_ids[pos]), pos, task_mask);
1920 if let Some(m) = &mut mtp {
1921 if pos + 1 < input_ids.len() {
1922 let probe: usize = std::env::var("CMF_MTP_CHAIN_PROBE")
1928 .ok()
1929 .and_then(|v| v.parse().ok())
1930 .unwrap_or(0);
1931 if probe >= 1 && pos + 2 < input_ids.len() {
1932 let (d1, mut hx) = self.mtp_step_h(m, &hidden, input_ids[pos + 1], pos);
1933 let mut ok = d1 == input_ids[pos + 2];
1934 Self::chain_probe_note(0, ok);
1935 let mut d_prev = d1;
1936 let mut extra = 0usize;
1937 for j in 1..probe {
1938 if pos + 2 + j >= input_ids.len() {
1939 break;
1940 }
1941 let (dj, hj) =
1942 self.mtp_step_h(m, &hx, d_prev, pos + 1 + j);
1943 extra += 1;
1944 ok = ok && dj == input_ids[pos + 2 + j];
1945 Self::chain_probe_note(j, ok);
1946 d_prev = dj;
1947 hx = hj;
1948 }
1949 m.kv.truncate_last(extra);
1952 } else {
1953 let _ = self.mtp_step(m, &hidden, input_ids[pos + 1], pos);
1954 }
1955 }
1956 }
1957 pos += 1;
1958 }
1959 if std::env::var("CMF_PREFILL_PROF").is_ok() {
1960 eprintln!(
1961 "prefill: {} tokens in {:.1} ms (batch_k={batch_k})",
1962 input_ids.len(),
1963 _tpf.elapsed().as_secs_f64() * 1000.0
1964 );
1965 }
1966 if self
1969 .cancel
1970 .swap(false, std::sync::atomic::Ordering::Relaxed)
1971 {
1972 self.kv_history.clear();
1973 if let Some(m) = mtp {
1974 self.mtp = Some(m);
1975 }
1976 return Ok(GenerateResult {
1977 text: String::new(),
1978 token_ids: Vec::new(),
1979 prompt_tokens: input_ids.len(),
1980 tokens_generated: 0,
1981 finish_reason: "cancelled".to_string(),
1982 mtp_drafted: 0,
1983 mtp_accepted: 0,
1984 token_confidence: Vec::new(),
1985 traces: Vec::new(),
1986 });
1987 }
1988
1989 self.o1_seal();
1992
1993 macro_rules! commit {
1995 ($id:expr) => {{
1996 all_ids.push($id);
1997 generated += 1;
1998 if self.tokenizer.is_eos($id) {
1999 finish_reason = "stop".to_string();
2000 false
2001 } else {
2002 let token_text = self.tokenizer.decode_token($id);
2003 let mut go = true;
2004 if let Some(ref mut cb) = on_token {
2005 if !cb(&token_text) {
2006 finish_reason = "cancelled".to_string();
2007 go = false;
2008 }
2009 }
2010 go
2011 }
2012 }};
2013 }
2014
2015 let mut next_pos = input_ids.len();
2017 'decode: while generated < max_tokens {
2018 if self
2019 .cancel
2020 .swap(false, std::sync::atomic::Ordering::Relaxed)
2021 {
2022 finish_reason = "cancelled".to_string();
2023 break 'decode;
2024 }
2025 let mut logits = match self.graph_logits.take() {
2026 Some(lg) => lg,
2027 None => {
2028 inference::rms_norm_into(
2029 &hidden,
2030 &self.weights.final_norm,
2031 self.rms_eps,
2032 self.norm_style,
2033 &mut self.ws.n1,
2034 );
2035 self.lm_head_forward(&self.ws.n1)
2036 }
2037 };
2038 let t_next = sampler::sample_with_scratch(
2039 &logits,
2040 &self.sampler_config,
2041 &all_ids,
2042 &mut self.rng,
2043 &mut self.sampler_scratch,
2044 );
2045 if self.confidence_on {
2046 confidence.push(top1_prob_t(&logits, t_next, calib_temp));
2047 }
2048 attention::recycle_buf(&mut logits);
2049 if trace_on {
2050 let skill = router.as_ref().and_then(|r| r.active_id());
2054 traces.push(TokenTrace {
2055 t: generated,
2056 token_id: t_next,
2057 confidence: confidence.last().copied().unwrap_or(0.0),
2058 active_skill: skill,
2059 recon: None,
2060 switched: false,
2061 });
2062 }
2063 if !commit!(t_next) {
2064 break 'decode;
2065 }
2066 if generated >= max_tokens {
2067 break 'decode;
2068 }
2069
2070 if self.kv_cache.needs_eviction() {
2071 let keep = (self.kv_cache.max_seq_len / 2).max(1);
2072 self.kv_cache.evict(keep);
2073 }
2074
2075 match &mut mtp {
2076 #[cfg(feature = "gpu")]
2078 Some(m) if graph_spec && generated + 1 < max_tokens && next_pos > 0 => {
2079 if let Some((extra, n_pos, new_h)) = self.graph_spec_step(
2080 m,
2081 &hidden,
2082 t_next,
2083 next_pos,
2084 &mut drafted,
2085 &mut accepted,
2086 ) {
2087 next_pos = n_pos;
2088 hidden = new_h;
2089 let mut stopped = false;
2090 for &id in &extra {
2091 if self.confidence_on {
2092 confidence.push(0.0);
2093 }
2094 if !commit!(id) {
2095 stopped = true;
2096 break;
2097 }
2098 }
2099 if stopped {
2100 break 'decode;
2101 }
2102 continue 'decode;
2103 }
2104 hidden = self.forward_layers(&self.embed_single(t_next), next_pos, task_mask);
2106 next_pos += 1;
2107 continue 'decode;
2108 }
2109 Some(m) if !graph_spec && generated + 1 < max_tokens => {
2111 let draft = self.mtp_step(m, &hidden, t_next, next_pos - 1);
2112 drafted += 1;
2113 let emb1 = self.embed_single(t_next);
2114 let emb2 = self.embed_single(draft);
2115 let (h1, h2) = self.forward_pair(&emb1, &emb2, next_pos);
2116
2117 inference::rms_norm_into(
2118 &h1,
2119 &self.weights.final_norm,
2120 self.rms_eps,
2121 self.norm_style,
2122 &mut self.ws.n1,
2123 );
2124 let mut logits1 = self.lm_head_forward(&self.ws.n1);
2125 let t_after = sampler::sample_with_scratch(
2126 &logits1,
2127 &self.sampler_config,
2128 &all_ids,
2129 &mut self.rng,
2130 &mut self.sampler_scratch,
2131 );
2132 if self.confidence_on {
2133 confidence.push(top1_prob_t(&logits1, t_after, calib_temp));
2134 }
2135 attention::recycle_buf(&mut logits1);
2136 if trace_on {
2137 traces.push(TokenTrace {
2140 t: generated,
2141 token_id: t_after,
2142 confidence: confidence.last().copied().unwrap_or(0.0),
2143 active_skill: None,
2144 recon: None,
2145 switched: false,
2146 });
2147 }
2148 let stop = !commit!(t_after);
2149
2150 if t_after == draft {
2151 accepted += 1;
2152 self.commit_linear_scratch();
2153 let _ = self.mtp_step(m, &h1, t_after, next_pos);
2154 hidden = h2;
2155 next_pos += 2;
2156 } else {
2157 for layer in &mut self.kv_cache.layers {
2159 layer.truncate_last(1);
2160 }
2161 if !stop {
2162 let _ = self.mtp_step(m, &h1, t_after, next_pos);
2163 hidden = self.forward_layers(
2164 &self.embed_single(t_after),
2165 next_pos + 1,
2166 None,
2167 );
2168 }
2169 next_pos += 2;
2170 }
2171 if stop {
2172 break 'decode;
2173 }
2174 }
2175 _ => {
2177 #[cfg(feature = "gpu")]
2182 if Self::dsv4_spec_on() && self.dsv4.is_some() {
2183 static SAID: std::sync::Once = std::sync::Once::new();
2184 SAID.call_once(|| {
2185 eprintln!(
2186 "dsv4-spec гейт: mtp={} mask={} router={} trace={} temp={} rep={} ",
2187 !self.dsv4_mtp.is_empty(),
2188 task_mask.is_none(),
2189 router.is_none(),
2190 !trace_on,
2191 self.sampler_config.temperature < 1e-6,
2192 self.sampler_config.repetition_penalty == 1.0,
2193 );
2194 });
2195 }
2196 #[cfg(feature = "gpu")]
2197 if Self::dsv4_spec_on()
2198 && self.dsv4.is_some()
2199 && !self.dsv4_mtp.is_empty()
2200 && task_mask.is_none()
2201 && router.is_none()
2202 && !trace_on
2203 && self.sampler_config.temperature < 1e-6
2204 && self.sampler_config.repetition_penalty == 1.0
2205 && generated + 1 < max_tokens
2206 && all_ids.len() >= 2
2207 {
2208 let tip_token = all_ids[all_ids.len() - 2];
2209 if let Some((extra, n_pos)) = self.dsv4_spec_step(
2210 tip_token,
2211 t_next,
2212 next_pos,
2213 &mut drafted,
2214 &mut accepted,
2215 ) {
2216 next_pos = n_pos;
2217 let mut stopped = false;
2218 for &id in &extra {
2219 if self.confidence_on {
2220 confidence.push(0.0);
2221 }
2222 if !commit!(id) {
2223 stopped = true;
2224 break;
2225 }
2226 }
2227 if stopped {
2228 break 'decode;
2229 }
2230 continue 'decode;
2231 }
2232 }
2233 self.graph_want_logits = fuse_lm;
2234 let mut t_fwd = t_next;
2240 let pure_greedy = self.sampler_config.temperature < 1e-6
2241 && self.sampler_config.repetition_penalty == 1.0
2242 && self.sampler_config.suppress_tokens.is_empty();
2243 let burst_k = std::env::var("CMF_MULTISTEP")
2248 .ok()
2249 .and_then(|v| v.parse::<usize>().ok())
2250 .unwrap_or(0);
2251 if pure_greedy
2252 && burst_k >= 1
2253 && fuse_lm
2254 && task_mask.is_none()
2255 && router.is_none()
2256 && !trace_on
2257 && !self.confidence_on
2258 {
2259 let mut stopped = false;
2260 loop {
2261 let room = max_tokens.saturating_sub(generated);
2262 if room <= 2 {
2263 break;
2264 }
2265 let k = burst_k.min(room - 1);
2266 if k < 1 {
2267 break;
2268 }
2269 let Some(ids) = self.try_multi_burst(t_fwd, next_pos, k) else {
2270 break;
2271 };
2272 next_pos += k;
2273 for &id in &ids {
2274 if !commit!(id) {
2275 stopped = true;
2276 break;
2277 }
2278 }
2279 if stopped {
2280 break;
2281 }
2282 t_fwd = *ids.last().unwrap();
2283 }
2284 if stopped {
2285 break 'decode;
2286 }
2287 }
2288 hidden = self.forward_layers(&self.embed_single(t_fwd), next_pos, task_mask);
2289 next_pos += 1;
2290 if let Some(r) = &mut router {
2293 let phi = self.dyn_phi_ema.clone();
2294 let decision = r.step(&phi, generated);
2295 if let Some(new_active) = decision {
2296 let _ = self.set_active_skill(new_active);
2297 }
2298 if trace_on {
2301 if let Some(last) = traces.last_mut() {
2302 let e = r.last_best_e();
2303 last.recon = e.is_finite().then_some(e);
2304 last.switched = decision.is_some();
2305 }
2306 }
2307 }
2308 }
2309 }
2310 }
2311
2312 self.graph_want_logits = false;
2313 self.graph_logits = None;
2314 if router.is_some() {
2316 let _ = self.set_active_skill(None);
2317 }
2318 self.dyn_router = router.or(self.dyn_router.take());
2319 self.mtp = mtp.or(self.mtp.take());
2320
2321 let output_ids = &all_ids[input_ids.len()..];
2322 let forwarded = input_ids.len() + output_ids.len().saturating_sub(1);
2326 self.kv_history = all_ids[..forwarded.min(all_ids.len())].to_vec();
2327 confidence.truncate(output_ids.len()); traces.truncate(output_ids.len());
2329 Ok(GenerateResult {
2330 text: self.tokenizer.decode(output_ids),
2331 token_ids: output_ids.to_vec(),
2332 prompt_tokens: input_ids.len(),
2333 tokens_generated: generated,
2334 finish_reason,
2335 mtp_drafted: drafted,
2336 mtp_accepted: accepted,
2337 token_confidence: confidence,
2338 traces,
2339 })
2340 }
2341
2342 fn mtp_step(
2346 &mut self,
2347 m: &mut MtpModule,
2348 hidden: &[f32],
2349 next_token: u32,
2350 position: usize,
2351 ) -> u32 {
2352 self.mtp_step_h(m, hidden, next_token, position).0
2353 }
2354
2355 fn chain_probe_note(depth: usize, prefix_ok: bool) {
2359 use std::sync::Mutex;
2360 static T: Mutex<Vec<(u64, u64)>> = Mutex::new(Vec::new());
2361 let mut t = T.lock().unwrap();
2362 if t.len() <= depth {
2363 t.resize(depth + 1, (0, 0));
2364 }
2365 t[depth].0 += 1;
2366 t[depth].1 += prefix_ok as u64;
2367 if depth == 0 && t[0].0 % 128 == 0 {
2368 let line: Vec<String> = t
2369 .iter()
2370 .enumerate()
2371 .map(|(d, (n, k))| format!("d{}={:.0}%({n})", d + 1, 100.0 * *k as f64 / (*n).max(1) as f64))
2372 .collect();
2373 eprintln!("mtp-chain: {}", line.join(" "));
2374 }
2375 }
2376
2377 fn mtp_step_h(
2381 &mut self,
2382 m: &mut MtpModule,
2383 hidden: &[f32],
2384 next_token: u32,
2385 position: usize,
2386 ) -> (u32, Vec<f32>) {
2387 let e = self.embed_single(next_token);
2391 let mut cat = vec![0.0f32; 2 * self.hidden_size];
2392 let (cat_e, cat_h) = cat.split_at_mut(self.hidden_size);
2393 inference::rms_norm_into(&e, &m.enorm, self.rms_eps, self.norm_style, cat_e);
2394 inference::rms_norm_into(hidden, &m.hnorm, self.rms_eps, self.norm_style, cat_h);
2395 let mut x = vec![0.0f32; self.hidden_size];
2396 m.eh_proj.matvec(&cat, &mut x, self.pool.as_deref());
2397
2398 let lw = &m.layer;
2400 inference::rms_norm_into(
2401 &x,
2402 &lw.input_norm,
2403 self.rms_eps,
2404 self.norm_style,
2405 &mut self.ws.n1,
2406 );
2407 let attn = match &lw.attn {
2408 AttnKind::Mla(_) => unreachable!("MLA has no MTP/pair path"),
2410 AttnKind::Kda(_) => unreachable!("KDA has no MTP/pair path"),
2411 AttnKind::Full {
2412 wq,
2413 wk,
2414 wv,
2415 wo,
2416 q_norm,
2417 k_norm,
2418 output_gate,
2419 softplus_gate,
2420 bias,
2421 } => {
2422 let mut cfg = self.attn_cfg(position);
2423 cfg.q_norm = q_norm.as_deref();
2424 cfg.k_norm = k_norm.as_deref();
2425 cfg.output_gate = *output_gate;
2426 cfg.softplus_gate = softplus_gate
2427 .as_ref()
2428 .map(|(gate, per_head)| (gate, *per_head));
2429 cfg.bias = bias
2430 .as_ref()
2431 .map(|(q, k, v)| (q.as_slice(), k.as_slice(), v.as_slice()));
2432 attention::qwen_attention(&self.ws.n1, wq, wk, wv, wo, &mut m.kv, &cfg)
2433 }
2434 AttnKind::Linear(_) | AttnKind::LinearGdn(_) | AttnKind::ShortConv(_) => {
2435 unreachable!("MTP block is full attention")
2436 }
2437 };
2438 for (i, &a) in attn.iter().enumerate() {
2439 x[i] += a;
2440 }
2441 inference::rms_norm_into(
2442 &x,
2443 &lw.post_norm,
2444 self.rms_eps,
2445 self.norm_style,
2446 &mut self.ws.p1,
2447 );
2448 let ffn = ffn_forward(&lw.ffn, &self.ws.p1, self.pool.as_deref(), None);
2449 for (i, &f) in ffn.iter().enumerate() {
2450 x[i] += f;
2451 }
2452
2453 inference::rms_norm_into(
2454 &x,
2455 &m.final_norm,
2456 self.rms_eps,
2457 self.norm_style,
2458 &mut self.ws.n1,
2459 );
2460 let mut lg = self.lm_head_forward(&self.ws.n1);
2461 let draft = sampler::argmax(&lg);
2462 attention::recycle_buf(&mut lg);
2463 (draft, x)
2464 }
2465
2466 fn mtp_warm(&mut self, m: &mut MtpModule, hidden: &[f32], next_token: u32, position: usize) {
2470 let e = self.embed_single(next_token);
2471 let mut cat = vec![0.0f32; 2 * self.hidden_size];
2472 let (cat_e, cat_h) = cat.split_at_mut(self.hidden_size);
2473 inference::rms_norm_into(&e, &m.enorm, self.rms_eps, self.norm_style, cat_e);
2474 inference::rms_norm_into(hidden, &m.hnorm, self.rms_eps, self.norm_style, cat_h);
2475 let mut x = vec![0.0f32; self.hidden_size];
2476 m.eh_proj.matvec(&cat, &mut x, self.pool.as_deref());
2477 inference::rms_norm_into(&x, &m.layer.input_norm, self.rms_eps, self.norm_style, &mut self.ws.n1);
2478 let attn = match &m.layer.attn {
2479 AttnKind::Full { wq, wk, wv, wo, q_norm, k_norm, output_gate, softplus_gate, bias } => {
2480 let mut cfg = self.attn_cfg(position);
2481 cfg.q_norm = q_norm.as_deref();
2482 cfg.k_norm = k_norm.as_deref();
2483 cfg.output_gate = *output_gate;
2484 cfg.softplus_gate = softplus_gate.as_ref().map(|(g, p)| (g, *p));
2485 cfg.bias = bias.as_ref().map(|(q, k, v)| (q.as_slice(), k.as_slice(), v.as_slice()));
2486 attention::qwen_attention(&self.ws.n1, wq, wk, wv, wo, &mut m.kv, &cfg)
2487 }
2488 _ => return,
2489 };
2490 let _ = attn;
2491 }
2492
2493 #[cfg(feature = "gpu")]
2500 #[allow(clippy::too_many_arguments)]
2501 fn graph_spec_step(
2502 &mut self,
2503 m: &mut MtpModule,
2504 hidden: &[f32],
2505 t_next: u32,
2506 next_pos: usize,
2507 drafted: &mut usize,
2508 accepted: &mut usize,
2509 ) -> Option<(Vec<u32>, usize, Vec<f32>)> {
2510 let k_spec: usize = std::env::var("CMF_GRAPH_SPEC_K")
2516 .ok()
2517 .and_then(|v| v.parse().ok())
2518 .filter(|&v| (1..=8).contains(&v))
2519 .unwrap_or(3);
2520 if next_pos == 0 {
2521 return None;
2522 }
2523 let t_round = std::time::Instant::now();
2524 let subs = || crate::gpu_wgpu::SUBMITS.load(std::sync::atomic::Ordering::Relaxed);
2540 let sub0 = subs();
2541 let mut drafts = Vec::with_capacity(k_spec);
2546 let (d1, mut hx) = self.mtp_step_h(m, hidden, t_next, next_pos - 1);
2547 drafts.push(d1);
2548 for j in 1..k_spec {
2549 let (dj, hj) = self.mtp_step_h(m, &hx, drafts[j - 1], next_pos - 1 + j);
2550 drafts.push(dj);
2551 hx = hj;
2552 }
2553 *drafted += k_spec;
2554 let t_draft = t_round.elapsed();
2555 let sub_draft = subs();
2556 let b = k_spec + 1;
2559 let mut hiddens = vec![0.0f32; b * self.hidden_size];
2560 for (i, &t) in std::iter::once(&t_next).chain(drafts.iter()).enumerate() {
2561 let e = self.embed_single(t);
2562 hiddens[i * self.hidden_size..(i + 1) * self.hidden_size].copy_from_slice(&e);
2563 }
2564 let positions: Vec<usize> = (next_pos..next_pos + b).collect();
2565 let (lm_gw, lm_rows) = {
2566 let (_, i, kind, rs) = self.weights.lm_head.graph_weight()?;
2567 (
2568 crate::gpu::GraphW { idx: i, kind, row_scale: rs, data: &[] },
2569 self.weights.lm_head.rows(),
2570 )
2571 };
2572 let mut logits = Vec::new();
2573 let final_norm = self.weights.final_norm.clone();
2574 let ok = self.try_batch_graph_wgpu(
2575 &mut hiddens,
2576 &positions,
2577 b,
2578 Some(crate::gpu::SpecTail {
2579 lm: lm_gw,
2580 lm_rows,
2581 final_norm: &final_norm,
2582 logits_out: &mut logits,
2583 }),
2584 );
2585 if !ok {
2586 m.kv.truncate_last(k_spec);
2589 return None;
2590 }
2591 let t_verify = t_round.elapsed();
2592 let sub_verify = subs();
2593 let ids: Vec<u32> = (0..b)
2595 .map(|i| sampler::argmax(&logits[i * lm_rows..(i + 1) * lm_rows]))
2596 .collect();
2597 let mut a = 0usize;
2598 while a < k_spec && ids[a] == drafts[a] {
2599 a += 1;
2600 }
2601 if a + 1 < b {
2603 crate::gpu::gdn_spec_restore(self.graph_kv_id, a);
2604 }
2605 *accepted += a;
2606 m.kv.truncate_last(k_spec.saturating_sub(1));
2617 let warm_off = std::env::var("CMF_SPEC_WARM").is_ok_and(|v| v == "0");
2618 if !warm_off {
2619 for j in 0..a {
2620 let row = &hiddens[j * self.hidden_size..(j + 1) * self.hidden_size];
2621 let row = row.to_vec();
2622 self.mtp_warm(m, &row, ids[j], next_pos + j);
2623 }
2624 }
2625 let mut row = logits[a * lm_rows..(a + 1) * lm_rows].to_vec();
2627 row.resize(self.vocab_size, 0.0);
2628 if let Some(c) = self.final_softcap {
2629 for l in row.iter_mut() {
2630 *l = c * (*l / c).tanh();
2631 }
2632 }
2633 self.graph_logits = Some(row);
2634 let new_hidden = hiddens[a * self.hidden_size..(a + 1) * self.hidden_size].to_vec();
2635 if std::env::var("CMF_GRAPH_SPEC_TIME").is_ok() {
2641 let end = subs();
2642 eprintln!(
2643 "spec-round: draft {:.1} ms/{} sub | verify {:.1} ms/{} sub | \
2644 commit {:.1} ms/{} sub (accepted {a} of {k_spec})",
2645 t_draft.as_secs_f64() * 1e3,
2646 sub_draft - sub0,
2647 (t_verify - t_draft).as_secs_f64() * 1e3,
2648 sub_verify - sub_draft,
2649 (t_round.elapsed() - t_verify).as_secs_f64() * 1e3,
2650 end - sub_verify,
2651 );
2652 }
2653 Some((drafts[..a].to_vec(), next_pos + a + 1, new_hidden))
2654 }
2655
2656 pub fn measure_pair_fusion(&mut self, iters: usize) -> (f64, f64) {
2665 if !self.pair_supported() {
2666 return (0.0, 0.0);
2667 }
2668 let emb1 = self.embed_single(1);
2669 let emb2 = self.embed_single(2);
2670 let pos = self.kv_cache.seq_len();
2671
2672 let t0 = std::time::Instant::now();
2673 for _ in 0..iters {
2674 let _ = self.forward_layers(&emb1, pos, None);
2675 let _ = self.forward_layers(&emb2, pos + 1, None);
2676 for l in &mut self.kv_cache.layers {
2677 l.truncate_last(2);
2678 }
2679 }
2680 let singles_ms = t0.elapsed().as_secs_f64() * 1000.0 / iters as f64;
2681
2682 let t1 = std::time::Instant::now();
2683 for _ in 0..iters {
2684 let _ = self.forward_pair(&emb1, &emb2, pos);
2685 for l in &mut self.kv_cache.layers {
2686 l.truncate_last(2);
2687 }
2688 }
2689 let pair_ms = t1.elapsed().as_secs_f64() * 1000.0 / iters as f64;
2690 (singles_ms, pair_ms)
2691 }
2692
2693 fn pair_supported(&self) -> bool {
2701 !self.weights.layers.is_empty()
2708 && self.g3n.is_none()
2709 && !self
2710 .weights
2711 .layers
2712 .iter()
2713 .any(|lw| matches!(&lw.attn, AttnKind::Mla(_) | AttnKind::Kda(_)))
2714 }
2715
2716 fn forward_pair(
2717 &mut self,
2718 emb1: &[f32],
2719 emb2: &[f32],
2720 position: usize,
2721 ) -> (Vec<f32>, Vec<f32>) {
2722 let mut h1 = emb1.to_vec();
2723 let mut h2 = emb2.to_vec();
2724 let (_nkv, _hd, hs, _rd, eps) = (
2725 self.num_kv_heads,
2726 self.head_dim,
2727 self.hidden_size,
2728 self.rotary_dim,
2729 self.rms_eps,
2730 );
2731 let pool = self.pool.clone();
2732
2733 for li in 0..self.num_layers {
2734 let lw = &self.weights.layers[self.phys_layer(li)];
2735 inference::rms_norm_into(
2738 &h1,
2739 &lw.input_norm,
2740 self.rms_eps,
2741 self.norm_style,
2742 &mut self.ws.n1,
2743 );
2744 inference::rms_norm_into(
2745 &h2,
2746 &lw.input_norm,
2747 self.rms_eps,
2748 self.norm_style,
2749 &mut self.ws.n2,
2750 );
2751
2752 let (a1, a2) = match &lw.attn {
2753 AttnKind::Mla(_) => unreachable!("MLA has no MTP/pair path"),
2754 AttnKind::Kda(_) => unreachable!("KDA has no MTP/pair path"),
2755 AttnKind::Linear(w) => {
2756 let cfg = self.vmf_cfg.expect("linear layer without vmf_cfg");
2757 let layer = &mut self.kv_cache.layers[li];
2758 let (state, scratch) = (&mut layer.linear_state, &mut layer.linear_scratch);
2759 vmf_phase_pair(
2760 &self.ws.n1,
2761 &self.ws.n2,
2762 w,
2763 &cfg,
2764 state,
2765 scratch,
2766 self.pool.as_deref(),
2767 )
2768 }
2769 AttnKind::LinearGdn(w) => {
2770 let cfg = self.gdn_cfg.expect("gdn layer without gdn_cfg");
2771 let layer = &mut self.kv_cache.layers[li];
2772 let (state, scratch) = (&mut layer.linear_state, &mut layer.linear_scratch);
2773 gdn_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::ShortConv(w) => {
2784 let cfg = self
2785 .short_conv_cfg
2786 .expect("short-conv layer without short_conv_cfg");
2787 let layer = &mut self.kv_cache.layers[li];
2788 let (state, scratch) = (&mut layer.linear_state, &mut layer.linear_scratch);
2789 short_conv_pair(
2790 &self.ws.n1,
2791 &self.ws.n2,
2792 w,
2793 &cfg,
2794 state,
2795 scratch,
2796 self.pool.as_deref(),
2797 )
2798 }
2799 AttnKind::Full {
2800 wq,
2801 wk,
2802 wv,
2803 wo,
2804 q_norm,
2805 k_norm,
2806 output_gate,
2807 softplus_gate,
2808 bias,
2809 } => {
2810 let inv_freq_l = self.layer_inv_freq(li);
2811 let (nkv_l, hd_l, rd_l) = self.layer_geom(li);
2812 let cfg = QwenAttnCfg {
2813 num_heads: self.layer_num_heads(li),
2814 num_kv_heads: nkv_l,
2815 head_dim: hd_l,
2816 hidden_size: hs,
2817 position,
2818 inv_freq: &inv_freq_l,
2819 rotary_dim: rd_l,
2820 scale: self.attn_scale,
2821 softcap: self.attn_softcap,
2822 window: self.layer_window(li),
2823 v_norm: self.attn_v_norm,
2824 q_norm: q_norm.as_deref(),
2825 k_norm: k_norm.as_deref(),
2826 output_gate: *output_gate,
2827 softplus_gate: softplus_gate
2828 .as_ref()
2829 .map(|(gate, per_head)| (gate, *per_head)),
2830 rope_scale: self.layer_rope_scale(li),
2831 bias: bias
2832 .as_ref()
2833 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
2834 rms_eps: eps,
2835 norm_style: self.norm_style,
2836 pool: pool.as_deref(),
2837 };
2838 attention::qwen_attention_pair(
2839 &self.ws.n1,
2840 &self.ws.n2,
2841 wq,
2842 wk,
2843 wv,
2844 wo,
2845 &mut self.kv_cache.layers[li],
2846 &cfg,
2847 )
2848 }
2849 };
2850 let (a1, a2) = match &self.weights.layers[self.phys_layer(li)].attn_out_norm {
2851 Some(w) => (
2852 inference::rms_norm(&a1, w, self.rms_eps, self.norm_style),
2853 inference::rms_norm(&a2, w, self.rms_eps, self.norm_style),
2854 ),
2855 None => (a1, a2),
2856 };
2857 for i in 0..self.hidden_size {
2858 h1[i] += a1[i];
2859 h2[i] += a2[i];
2860 }
2861 let (mut a1, mut a2) = (a1, a2);
2862 attention::recycle_buf(&mut a1);
2863 attention::recycle_buf(&mut a2);
2864
2865 let lw = &self.weights.layers[self.phys_layer(li)];
2866 inference::rms_norm_into(
2867 &h1,
2868 &lw.post_norm,
2869 self.rms_eps,
2870 self.norm_style,
2871 &mut self.ws.p1,
2872 );
2873 inference::rms_norm_into(
2874 &h2,
2875 &lw.post_norm,
2876 self.rms_eps,
2877 self.norm_style,
2878 &mut self.ws.p2,
2879 );
2880 let (f1, f2) = match &lw.ffn {
2881 FfnKind::DenseMoe(dm) => (
2884 dense_moe_ffn(
2885 dm,
2886 &self.ws.p1,
2887 &h1,
2888 self.rms_eps,
2889 self.norm_style,
2890 self.pool.as_deref(),
2891 ),
2892 dense_moe_ffn(
2893 dm,
2894 &self.ws.p2,
2895 &h2,
2896 self.rms_eps,
2897 self.norm_style,
2898 self.pool.as_deref(),
2899 ),
2900 ),
2901 _ => ffn_forward_pair(
2902 &lw.ffn,
2903 &self.ws.p1,
2904 &self.ws.p2,
2905 self.pool.as_deref(),
2906 None,
2907 ),
2908 };
2909 let (f1, f2) = match &self.weights.layers[self.phys_layer(li)].ffn_out_norm {
2910 Some(w) => (
2911 inference::rms_norm(&f1, w, self.rms_eps, self.norm_style),
2912 inference::rms_norm(&f2, w, self.rms_eps, self.norm_style),
2913 ),
2914 None => (f1, f2),
2915 };
2916 for i in 0..self.hidden_size {
2917 h1[i] += f1[i];
2918 h2[i] += f2[i];
2919 }
2920 let (mut f1, mut f2) = (f1, f2);
2921 attention::recycle_buf(&mut f1);
2922 attention::recycle_buf(&mut f2);
2923 if let Some(sc) = self.weights.layers[self.phys_layer(li)].layer_scale {
2924 for i in 0..self.hidden_size {
2925 h1[i] *= sc;
2926 h2[i] *= sc;
2927 }
2928 }
2929 if self.is_loop_end(li) && li + 1 < self.num_layers {
2931 h1 = inference::rms_norm(
2932 &h1,
2933 &self.weights.final_norm,
2934 self.rms_eps,
2935 self.norm_style,
2936 );
2937 h2 = inference::rms_norm(
2938 &h2,
2939 &self.weights.final_norm,
2940 self.rms_eps,
2941 self.norm_style,
2942 );
2943 }
2944 }
2945 (h1, h2)
2946 }
2947
2948 fn commit_linear_scratch(&mut self) {
2950 for layer in &mut self.kv_cache.layers {
2951 if !layer.linear_scratch.is_empty() {
2952 std::mem::swap(&mut layer.linear_state, &mut layer.linear_scratch);
2953 layer.linear_scratch.clear();
2954 }
2955 }
2956 }
2957
2958 pub fn forward_ids(
2961 &mut self,
2962 ids: &[u32],
2963 task_mask: Option<&TaskMask>,
2964 ) -> Result<Vec<f32>, String> {
2965 if ids.is_empty() {
2966 return Err("empty id sequence".to_string());
2967 }
2968 self.kv_cache.clear();
2969 self.kv_history.clear();
2970 self.o1_begin();
2971 let mut hidden = vec![0.0f32; self.hidden_size];
2972 let mut pos = 0usize;
2973 if self.can_prefill_batched() && !self.graph_prefill_preferred() && ids.len() > 2 {
2981 let chunk = prefill_chunk();
2985 let hs = self.hidden_size;
2986 while pos < ids.len() {
2987 let end = (pos + chunk).min(ids.len());
2988 let hb = self.prefill_batch_masked(&ids[pos..end], pos, task_mask);
2989 hidden.copy_from_slice(&hb[(end - pos - 1) * hs..]);
2990 pos = end;
2991 }
2992 }
2993 if task_mask.is_none()
3002 && !self.graph_prefill_preferred()
3003 && !std::env::var("CMF_PAIR").is_ok_and(|v| v == "0")
3004 && self.pair_supported()
3005 {
3006 while pos + 1 < ids.len() {
3007 let e1 = self.embed_single(ids[pos]);
3008 let e2 = self.embed_single(ids[pos + 1]);
3009 let (_, h2) = self.forward_pair(&e1, &e2, pos);
3010 self.commit_linear_scratch();
3011 hidden = h2;
3012 pos += 2;
3013 }
3014 }
3015 while pos < ids.len() {
3016 hidden = self.forward_layers(&self.embed_single(ids[pos]), pos, task_mask);
3017 pos += 1;
3018 }
3019 self.o1_seal();
3023 let normed = inference::rms_norm(
3024 &hidden,
3025 &self.weights.final_norm,
3026 self.rms_eps,
3027 self.norm_style,
3028 );
3029 Ok(self.lm_head_forward(&normed))
3030 }
3031
3032 pub fn ppl_ids(&mut self, ids: &[u32]) -> f64 {
3039 let (nll, cnt) = self.nll_ids_from(ids, 0);
3040 (nll / cnt.max(1) as f64).exp()
3041 }
3042
3043 pub fn probe_ffn_mass(&mut self, ids: &[u32]) -> Vec<Vec<f64>> {
3048 self.kv_cache.clear();
3049 self.kv_history.clear();
3050 FFN_PROBE.with(|p| {
3051 *p.borrow_mut() = Some(vec![vec![0f64; self.intermediate_size]; self.num_layers]);
3052 });
3053 crate::gpu::cpu_scope(|| {
3054 for (pos, &id) in ids.iter().enumerate() {
3055 let emb = self.embed_single(id);
3056 let _ = self.forward_layers(&emb, pos, None);
3057 }
3058 });
3059 self.kv_cache.clear();
3060 self.kv_history.clear();
3061 FFN_PROBE
3062 .with(|p| p.borrow_mut().take())
3063 .unwrap_or_default()
3064 }
3065
3066 pub fn ppl_ids_masked(&mut self, ids: &[u32], mask: &TaskMask) -> f64 {
3070 self.kv_cache.clear();
3071 self.kv_history.clear();
3072 let mut nll = 0f64;
3073 let mut cnt = 0usize;
3074 let mut hidden = vec![0f32; self.hidden_size];
3075 for (pos, &id) in ids.iter().enumerate() {
3076 if pos > 0 {
3077 inference::rms_norm_into(
3078 &hidden,
3079 &self.weights.final_norm,
3080 self.rms_eps,
3081 self.norm_style,
3082 &mut self.ws.n1,
3083 );
3084 let mut logits = self.lm_head_forward(&self.ws.n1);
3085 let max = logits.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
3086 let sum: f64 = logits.iter().map(|&v| ((v - max) as f64).exp()).sum();
3087 let p = ((logits[id as usize] - max) as f64).exp() / sum.max(1e-300);
3088 nll -= p.max(1e-300).ln();
3089 cnt += 1;
3090 attention::recycle_buf(&mut logits);
3091 }
3092 let emb = self.embed_single(id);
3093 hidden = self.forward_layers(&emb, pos, Some(mask));
3094 }
3095 self.kv_cache.clear();
3096 self.kv_history.clear();
3097 (nll / cnt.max(1) as f64).exp()
3098 }
3099
3100 pub fn nll_ids_masked(
3119 &mut self,
3120 ids: &[u32],
3121 start: usize,
3122 task_mask: Option<&TaskMask>,
3123 ) -> (f64, usize) {
3124 self.nll_ids_inner(ids, start, task_mask)
3125 }
3126
3127 pub fn nll_ids_from(&mut self, ids: &[u32], start: usize) -> (f64, usize) {
3128 self.nll_ids_inner(ids, start, None)
3129 }
3130
3131 fn nll_ids_inner(
3132 &mut self,
3133 ids: &[u32],
3134 start: usize,
3135 task_mask: Option<&TaskMask>,
3136 ) -> (f64, usize) {
3137 self.kv_cache.clear();
3138 self.kv_history.clear();
3139 let mut nll = 0f64;
3140 let mut cnt = 0usize;
3141 if self.can_prefill_batched() {
3142 const CHUNK: usize = 128;
3148 const LM_SUB: usize = 32;
3149 let n = ids.len().saturating_sub(1);
3150 let hs = self.hidden_size;
3151 let rows = self.weights.lm_head.rows();
3152 let mut pos = 0usize;
3153 while pos < n {
3154 let end = (pos + CHUNK).min(n);
3155 let bsz = end - pos;
3156 let hb = self.prefill_batch_masked(&ids[pos..end], pos, task_mask);
3157 let mut k0 = 0usize;
3158 while k0 < bsz {
3159 let k1 = (k0 + LM_SUB).min(bsz);
3160 let sb = k1 - k0;
3161 if pos + k1 <= start {
3164 k0 = k1;
3165 continue;
3166 }
3167 let mut normed = vec![0.0f32; sb * hs];
3168 for k in 0..sb {
3169 let r = inference::rms_norm(
3170 &hb[(k0 + k) * hs..(k0 + k + 1) * hs],
3171 &self.weights.final_norm,
3172 self.rms_eps,
3173 self.norm_style,
3174 );
3175 normed[k * hs..(k + 1) * hs].copy_from_slice(&r);
3176 }
3177 let mut logits = vec![0.0f32; sb * rows];
3178 self.weights
3179 .lm_head
3180 .matmat(&normed, sb, &mut logits, self.pool.as_deref());
3181 for k in 0..sb {
3182 if pos + k0 + k < start {
3183 continue;
3184 }
3185 let lg = &mut logits[k * rows..k * rows + self.vocab_size.min(rows)];
3186 if let Some(mu) = self.logit_multiplier {
3187 for v in lg.iter_mut() {
3188 *v *= mu;
3189 }
3190 }
3191 if let Some(c) = self.final_softcap {
3195 for v in lg.iter_mut() {
3196 *v = c * (*v / c).tanh();
3197 }
3198 }
3199 let lg = &logits[k * rows..k * rows + self.vocab_size.min(rows)];
3200 let target = ids[pos + k0 + k + 1] as usize;
3201 let max = lg.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
3202 let lse: f64 = lg
3203 .iter()
3204 .map(|&v| ((v - max) as f64).exp())
3205 .sum::<f64>()
3206 .ln()
3207 + max as f64;
3208 nll += lse - lg[target] as f64;
3209 cnt += 1;
3210 if std::env::var("CMF_PPL_TRACE").is_ok() {
3211 let top = lg
3212 .iter()
3213 .enumerate()
3214 .max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
3215 .map(|(i, _)| i)
3216 .unwrap_or(0);
3217 eprintln!(
3218 "BTRACE pos {} target {} nll {:.4} top {} lg_t {:.3} lg_top {:.3}",
3219 pos + k0 + k,
3220 target,
3221 lse - lg[target] as f64,
3222 top,
3223 lg[target],
3224 lg[top]
3225 );
3226 }
3227 }
3228 k0 = k1;
3229 }
3230 pos = end;
3231 }
3232 self.kv_cache.clear();
3233 self.kv_history.clear();
3234 return (nll, cnt);
3235 }
3236 for pos in 0..ids.len().saturating_sub(1) {
3237 let hidden = self.forward_layers(&self.embed_single(ids[pos]), pos, task_mask);
3238 let out_of_band = self.graph_logits.take();
3246 if pos < start {
3247 continue;
3248 }
3249 let logits = match out_of_band {
3250 Some(lg) => lg,
3251 None => {
3252 let normed = inference::rms_norm(
3253 &hidden,
3254 &self.weights.final_norm,
3255 self.rms_eps,
3256 self.norm_style,
3257 );
3258 self.lm_head_forward(&normed)
3262 }
3263 };
3264 let target = ids[pos + 1] as usize;
3265 let max = logits.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
3266 let lse: f64 = logits
3267 .iter()
3268 .map(|&v| ((v - max) as f64).exp())
3269 .sum::<f64>()
3270 .ln()
3271 + max as f64;
3272 let tok_nll = lse - logits[target] as f64;
3273 if std::env::var("CMF_PPL_TRACE").is_ok() && pos < 48 {
3274 let top = logits
3275 .iter()
3276 .enumerate()
3277 .max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
3278 .map(|(i, _)| i)
3279 .unwrap_or(0);
3280 eprintln!(
3281 "pos {pos:3} tgt {target:6} nll {tok_nll:7.3} | top1 {top:6} lg[t]={:.2} lg[top]={:.2}",
3282 logits[target], logits[top]
3283 );
3284 }
3285 nll += tok_nll;
3286 cnt += 1;
3287 }
3288 self.kv_cache.clear();
3289 self.kv_history.clear();
3290 (nll, cnt)
3291 }
3292
3293 pub fn nll_ids_o1(&mut self, ids: &[u32], prefill: usize) -> (f64, usize) {
3309 self.kv_cache.clear();
3310 self.kv_history.clear();
3311 self.o1_begin();
3312 let n = ids.len().saturating_sub(1);
3313 let p = prefill.min(n);
3314 let mut pos = 0usize;
3316 if self.can_prefill_batched() {
3317 const CHUNK: usize = 128;
3318 while pos < p {
3319 let end = (pos + CHUNK).min(p);
3320 let _ = self.prefill_batch(&ids[pos..end], pos);
3321 pos = end;
3322 }
3323 } else {
3324 while pos < p {
3325 let _ = self.forward_layers(&self.embed_single(ids[pos]), pos, None);
3326 pos += 1;
3327 }
3328 }
3329 self.o1_seal();
3330
3331 let mut nll = 0f64;
3332 let mut cnt = 0usize;
3333 for pos in p..n {
3334 let hidden = self.forward_layers(&self.embed_single(ids[pos]), pos, None);
3335 let normed = inference::rms_norm(
3336 &hidden,
3337 &self.weights.final_norm,
3338 self.rms_eps,
3339 self.norm_style,
3340 );
3341 let logits = self.lm_head_forward(&normed);
3345 let target = ids[pos + 1] as usize;
3346 let max = logits.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
3347 let lse: f64 = logits
3348 .iter()
3349 .map(|&v| ((v - max) as f64).exp())
3350 .sum::<f64>()
3351 .ln()
3352 + max as f64;
3353 let tok_nll = lse - logits[target] as f64;
3354 if std::env::var("CMF_PPL_TRACE").is_ok() && pos < 48 {
3355 let top = logits
3356 .iter()
3357 .enumerate()
3358 .max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
3359 .map(|(i, _)| i)
3360 .unwrap_or(0);
3361 eprintln!(
3362 "pos {pos:3} tgt {target:6} nll {tok_nll:7.3} | top1 {top:6} lg[t]={:.2} lg[top]={:.2}",
3363 logits[target], logits[top]
3364 );
3365 }
3366 nll += tok_nll;
3367 cnt += 1;
3368 }
3369 self.kv_cache.clear();
3370 self.kv_history.clear();
3371 (nll, cnt)
3372 }
3373
3374 pub fn calib_ids(&mut self, ids: &[u32], temps: &[f32]) -> (Vec<bool>, Vec<Vec<f32>>) {
3382 self.kv_cache.clear();
3383 self.kv_history.clear();
3384 let n = ids.len().saturating_sub(1);
3385 let mut correct = Vec::with_capacity(n);
3386 let mut pmax = Vec::with_capacity(n);
3387 for pos in 0..n {
3388 let emb = self.embed_single(ids[pos]);
3389 let hidden = self.forward_layers(&emb, pos, None);
3390 let normed = inference::rms_norm(
3391 &hidden,
3392 &self.weights.final_norm,
3393 self.rms_eps,
3394 self.norm_style,
3395 );
3396 let logits = self.lm_head_forward(&normed);
3400 let target = ids[pos + 1] as usize;
3401 let (mut amax, mut mval) = (0usize, f32::NEG_INFINITY);
3402 for (i, &v) in logits.iter().enumerate() {
3403 if v > mval {
3404 mval = v;
3405 amax = i;
3406 }
3407 }
3408 correct.push(amax == target);
3409 let row: Vec<f32> = temps
3410 .iter()
3411 .map(|&t| {
3412 let tt = t.max(1e-3);
3413 let s: f32 = logits.iter().map(|&v| ((v - mval) / tt).exp()).sum();
3414 1.0 / s.max(1e-12) })
3416 .collect();
3417 pmax.push(row);
3418 }
3419 self.kv_cache.clear();
3420 self.kv_history.clear();
3421 (correct, pmax)
3422 }
3423
3424 pub fn ppl_ids_dynamic(&mut self, ids: &[u32]) -> (f64, usize) {
3431 let mut router = match self.dyn_router.take() {
3432 Some(r) => r,
3433 None => return (self.ppl_ids(ids), 0),
3434 };
3435 router.reset();
3436 self.dyn_phi_seen = 0;
3437 let _ = self.set_active_skill(None);
3438
3439 self.kv_cache.clear();
3440
3441 self.kv_history.clear();
3442 let mut nll = 0f64;
3443 let mut cnt = 0usize;
3444 for pos in 0..ids.len().saturating_sub(1) {
3445 let hidden = self.forward_layers(&self.embed_single(ids[pos]), pos, None);
3446 let normed = inference::rms_norm(
3447 &hidden,
3448 &self.weights.final_norm,
3449 self.rms_eps,
3450 self.norm_style,
3451 );
3452 let logits = self.lm_head_forward(&normed);
3456 let target = ids[pos + 1] as usize;
3457 let max = logits.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
3458 let lse: f64 = logits
3459 .iter()
3460 .map(|&v| ((v - max) as f64).exp())
3461 .sum::<f64>()
3462 .ln()
3463 + max as f64;
3464 let tok_nll = lse - logits[target] as f64;
3465 if std::env::var("CMF_PPL_TRACE").is_ok() && pos < 48 {
3466 let top = logits
3467 .iter()
3468 .enumerate()
3469 .max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
3470 .map(|(i, _)| i)
3471 .unwrap_or(0);
3472 eprintln!(
3473 "pos {pos:3} tgt {target:6} nll {tok_nll:7.3} | top1 {top:6} lg[t]={:.2} lg[top]={:.2}",
3474 logits[target], logits[top]
3475 );
3476 }
3477 nll += tok_nll;
3478 cnt += 1;
3479 let phi = self.dyn_phi_ema.clone();
3481 if let Some(new_active) = router.step(&phi, pos) {
3482 let _ = self.set_active_skill(new_active);
3483 }
3484 }
3485 let switches = router.switches.len();
3486 let _ = self.set_active_skill(None);
3487 self.dyn_router = Some(router);
3488 self.kv_cache.clear();
3489 self.kv_history.clear();
3490 ((nll / cnt.max(1) as f64).exp(), switches)
3491 }
3492
3493 pub fn probe_phi(&mut self, ids: &[u32], layer: usize) -> Vec<f32> {
3495 self.kv_cache.clear();
3496 self.kv_history.clear();
3497 let mut acc = vec![0f32; self.hidden_size];
3498 for (pos, &id) in ids.iter().enumerate() {
3499 let h = self.forward_layers_upto(&self.embed_single(id), pos, None, Some(layer));
3500 for (a, v) in acc.iter_mut().zip(&h) {
3501 *a += v;
3502 }
3503 }
3504 let n = ids.len().max(1) as f32;
3505 for a in acc.iter_mut() {
3506 *a /= n;
3507 }
3508 self.kv_cache.clear();
3509 self.kv_history.clear();
3510 acc
3511 }
3512
3513 fn prefill_batch(&mut self, ids: &[u32], start_pos: usize) -> Vec<f32> {
3519 self.prefill_batch_masked(ids, start_pos, None)
3520 }
3521
3522 fn prefill_batch_masked(
3528 &mut self,
3529 ids: &[u32],
3530 start_pos: usize,
3531 task_mask: Option<&TaskMask>,
3532 ) -> Vec<f32> {
3533 self.prefill_batch_span(PrefillIn::Ids(ids), start_pos, task_mask, 0, usize::MAX)
3534 }
3535
3536 fn prefill_batch_span(
3542 &mut self,
3543 input: PrefillIn<'_>,
3544 start_pos: usize,
3545 task_mask: Option<&TaskMask>,
3546 from: usize,
3547 upto_excl: usize,
3548 ) -> Vec<f32> {
3549 let hs = self.hidden_size;
3550 let b = match input {
3551 PrefillIn::Ids(ids) => ids.len(),
3552 PrefillIn::Hidden(hb) => hb.len() / hs,
3553 };
3554 let upto_excl = upto_excl.min(self.num_layers);
3555 let mut h: Vec<f32>;
3559 let mut h_ready;
3560 match input {
3561 PrefillIn::Ids(_) => {
3562 h = vec![0.0; b * hs];
3563 h_ready = false;
3564 }
3565 PrefillIn::Hidden(hb) => {
3566 h = hb.to_vec();
3567 h_ready = true;
3568 }
3569 }
3570 let fill_h = |h: &mut Vec<f32>, me: &Self| {
3571 if let PrefillIn::Ids(ids) = input {
3572 for (bi, &id) in ids.iter().enumerate() {
3573 let e = me.embed_single(id);
3574 h[bi * hs..(bi + 1) * hs].copy_from_slice(&e);
3575 }
3576 }
3577 };
3578 let (_nkv, _hd, _rd, eps) = (
3579 self.num_kv_heads,
3580 self.head_dim,
3581 self.rotary_dim,
3582 self.rms_eps,
3583 );
3584 let pool = self.pool.clone();
3585 let norm_style = self.norm_style;
3586
3587 #[cfg(target_os = "macos")]
3588 let mut chunk_skip_until = 0usize;
3589 for li in from..upto_excl {
3590 crate::gpu::set_layer(li as i64); #[cfg(target_os = "macos")]
3597 if task_mask.is_none() {
3598 if li < chunk_skip_until {
3599 continue;
3600 }
3601 if !h_ready && li == 0 && self.weights.embed_tokens.q8_row_parts().is_none() {
3607 fill_h(&mut h, self);
3608 h_ready = true;
3609 }
3610 let ids_for_embed = match input {
3611 PrefillIn::Ids(ids) => (!h_ready && li == 0).then_some(ids),
3612 PrefillIn::Hidden(_) => None,
3613 };
3614 let end = self.chunk_run_gpu(li, &mut h, b, start_pos, ids_for_embed, upto_excl);
3615 if end > li {
3616 h_ready = true;
3617 chunk_skip_until = end;
3618 if self.is_loop_end(end - 1) && end < self.num_layers {
3621 for bi in 0..b {
3622 let normed = inference::rms_norm(
3623 &h[bi * hs..(bi + 1) * hs],
3624 &self.weights.final_norm,
3625 eps,
3626 norm_style,
3627 );
3628 h[bi * hs..(bi + 1) * hs].copy_from_slice(&normed);
3629 }
3630 }
3631 continue;
3632 }
3633 }
3634 if !h_ready {
3635 fill_h(&mut h, self);
3636 h_ready = true;
3637 }
3638 let lw = &self.weights.layers[self.phys_layer(li)];
3639 match &lw.attn {
3641 AttnKind::Kda(w) => {
3642 let cfg = self.kda_cfg.expect("kda layer without kda_cfg");
3644 let mut normed = vec![0.0f32; b * hs];
3645 for bi in 0..b {
3646 inference::rms_norm_into(
3647 &h[bi * hs..(bi + 1) * hs],
3648 &lw.input_norm,
3649 eps,
3650 norm_style,
3651 &mut normed[bi * hs..(bi + 1) * hs],
3652 );
3653 }
3654 let attn = crate::linear_core::kda_forward_batch(
3655 &normed,
3656 b,
3657 w,
3658 &cfg,
3659 &mut self.kv_cache.layers[li].linear_state,
3660 pool.as_deref(),
3661 );
3662 for (dst, &a) in h.iter_mut().zip(&attn) {
3663 *dst += a;
3664 }
3665 }
3666 AttnKind::LinearGdn(w) => {
3667 let cfg = self.gdn_cfg.expect("gdn layer without gdn_cfg");
3669 let mut normed = vec![0.0f32; b * hs];
3670 for bi in 0..b {
3671 let r = inference::rms_norm(
3672 &h[bi * hs..(bi + 1) * hs],
3673 &lw.input_norm,
3674 eps,
3675 norm_style,
3676 );
3677 normed[bi * hs..(bi + 1) * hs].copy_from_slice(&r);
3678 }
3679 let attn = crate::linear_core::gdn_forward_batch(
3680 &normed,
3681 b,
3682 w,
3683 &cfg,
3684 &mut self.kv_cache.layers[li].linear_state,
3685 pool.as_deref(),
3686 );
3687 for (dst, &a) in h.iter_mut().zip(&attn) {
3688 *dst += a;
3689 }
3690 }
3691 AttnKind::ShortConv(w) => {
3692 let cfg = self
3695 .short_conv_cfg
3696 .expect("short-conv layer without short_conv_cfg");
3697 let mut normed = vec![0.0f32; b * hs];
3698 for bi in 0..b {
3699 inference::rms_norm_into(
3700 &h[bi * hs..(bi + 1) * hs],
3701 &lw.input_norm,
3702 eps,
3703 norm_style,
3704 &mut normed[bi * hs..(bi + 1) * hs],
3705 );
3706 }
3707 let attn = short_conv_forward_batch(
3708 &normed,
3709 b,
3710 w,
3711 &cfg,
3712 &mut self.kv_cache.layers[li].linear_state,
3713 pool.as_deref(),
3714 );
3715 for (dst, &a) in h.iter_mut().zip(&attn) {
3716 *dst += a;
3717 }
3718 }
3719 AttnKind::Mla(w) => {
3720 let inv_freq_l = self.layer_inv_freq(li);
3723 let rs = self.layer_rope_scale(li);
3724 let mut normed = vec![0.0f32; hs];
3725 for bi in 0..b {
3726 inference::rms_norm_into(
3727 &h[bi * hs..(bi + 1) * hs],
3728 &lw.input_norm,
3729 eps,
3730 norm_style,
3731 &mut normed,
3732 );
3733 let ao = mla_attention(
3734 w,
3735 &normed,
3736 &mut self.kv_cache.layers[li],
3737 start_pos + bi,
3738 &inv_freq_l,
3739 rs,
3740 eps,
3741 pool.as_deref(),
3742 );
3743 for (dst, &a) in h[bi * hs..(bi + 1) * hs].iter_mut().zip(&ao) {
3744 *dst += a;
3745 }
3746 }
3747 }
3748 AttnKind::Full {
3749 wq,
3750 wk,
3751 wv,
3752 wo,
3753 q_norm,
3754 k_norm,
3755 output_gate,
3756 softplus_gate,
3757 bias,
3758 } => {
3759 let mut normed = vec![0.0f32; b * hs];
3763 for bi in 0..b {
3764 inference::rms_norm_into(
3765 &h[bi * hs..(bi + 1) * hs],
3766 &lw.input_norm,
3767 eps,
3768 norm_style,
3769 &mut normed[bi * hs..(bi + 1) * hs],
3770 );
3771 }
3772 let inv_freq_l = self.layer_inv_freq(li);
3773 let (nkv_l, hd_l, rd_l) = self.layer_geom(li);
3774 let cfg = QwenAttnCfg {
3775 num_heads: self.layer_num_heads(li),
3776 num_kv_heads: nkv_l,
3777 head_dim: hd_l,
3778 hidden_size: hs,
3779 position: start_pos,
3780 inv_freq: &inv_freq_l,
3781 rotary_dim: rd_l,
3782 scale: self.attn_scale,
3783 softcap: self.attn_softcap,
3784 window: self.layer_window(li),
3785 v_norm: self.attn_v_norm,
3786 q_norm: q_norm.as_deref(),
3787 k_norm: k_norm.as_deref(),
3788 output_gate: *output_gate,
3789 softplus_gate: softplus_gate
3790 .as_ref()
3791 .map(|(gate, per_head)| (gate, *per_head)),
3792 rope_scale: self.layer_rope_scale(li),
3793 bias: bias
3794 .as_ref()
3795 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
3796 rms_eps: eps,
3797 norm_style,
3798 pool: pool.as_deref(),
3799 };
3800 let mut attn = attention::qwen_attention_batch(
3801 &normed,
3802 b,
3803 wq,
3804 wk,
3805 wv,
3806 wo,
3807 &mut self.kv_cache.layers[li],
3808 &cfg,
3809 );
3810 if let Some(w) = &lw.attn_out_norm {
3811 for bi in 0..b {
3812 inference::rms_norm_into(
3813 &attn[bi * hs..(bi + 1) * hs],
3814 w,
3815 eps,
3816 norm_style,
3817 &mut normed[bi * hs..(bi + 1) * hs],
3818 );
3819 }
3820 attn.copy_from_slice(&normed);
3821 }
3822 for (dst, &a) in h.iter_mut().zip(&attn) {
3823 *dst += a;
3824 }
3825 }
3826 AttnKind::Linear(w) => {
3827 for bi in 0..b {
3828 let normed = inference::rms_norm(
3829 &h[bi * hs..(bi + 1) * hs],
3830 &lw.input_norm,
3831 eps,
3832 norm_style,
3833 );
3834 vmf_phase_forward(
3835 &normed,
3836 w,
3837 &self.vmf_cfg.expect("linear layer without vmf_cfg"),
3838 &mut self.kv_cache.layers[li].linear_state,
3839 pool.as_deref(),
3840 )
3841 .iter()
3842 .enumerate()
3843 .for_each(|(i, &a)| h[bi * hs + i] += a);
3844 }
3845 }
3846 }
3847
3848 let lw = &self.weights.layers[self.phys_layer(li)];
3850 let mut post = vec![0.0f32; b * hs];
3851 for bi in 0..b {
3852 let r =
3853 inference::rms_norm(&h[bi * hs..(bi + 1) * hs], &lw.post_norm, eps, norm_style);
3854 post[bi * hs..(bi + 1) * hs].copy_from_slice(&r);
3855 }
3856 let mask_row = task_mask
3859 .filter(|m| m.ffn_active_count(li) < self.intermediate_size)
3860 .and_then(|m| m.ffn_masks.get(li))
3861 .map(|v| v.as_slice());
3862 let mut ffn = match &lw.ffn {
3863 FfnKind::Dense(d) => dense_ffn_batch(d, &post, b, pool.as_deref(), mask_row),
3864 FfnKind::Moe(m) => moe_ffn_batch(m, &post, b, hs, pool.as_deref(), None),
3865 FfnKind::DenseMoe(dm) => {
3868 let mut out = vec![0.0f32; b * hs];
3869 for bi in 0..b {
3870 let r = dense_moe_ffn(
3871 dm,
3872 &post[bi * hs..(bi + 1) * hs],
3873 &h[bi * hs..(bi + 1) * hs],
3874 eps,
3875 norm_style,
3876 pool.as_deref(),
3877 );
3878 out[bi * hs..(bi + 1) * hs].copy_from_slice(&r);
3879 }
3880 out
3881 }
3882 };
3883 if let Some(w) = &lw.ffn_out_norm {
3884 for bi in 0..b {
3885 inference::rms_norm_into(
3886 &ffn[bi * hs..(bi + 1) * hs],
3887 w,
3888 eps,
3889 norm_style,
3890 &mut post[bi * hs..(bi + 1) * hs],
3891 );
3892 }
3893 ffn.copy_from_slice(&post);
3894 }
3895 for (dst, &f) in h.iter_mut().zip(&ffn) {
3896 *dst += f;
3897 }
3898 if let Some(sc) = lw.layer_scale {
3899 for v in h.iter_mut() {
3900 *v *= sc;
3901 }
3902 }
3903 if let Ok(tp) = std::env::var("CMF_TRACE_POS") {
3904 if let Some(t) = tp.parse::<usize>().ok() {
3905 if t >= start_pos && t < start_pos + b {
3906 let bi = t - start_pos;
3907 let row = &h[bi * hs..(bi + 1) * hs];
3908 let n: f32 = row.iter().map(|x| x * x).sum::<f32>().sqrt();
3909 eprintln!(
3910 "BATCH pos {t} after layer {li}: |h| = {n:.6} h0 {:.6} h1 {:.6}",
3911 row[0], row[1]
3912 );
3913 }
3914 }
3915 }
3916 if std::env::var("CMF_DEBUG_LAYERS").is_ok() {
3920 let row = &h[(b - 1) * hs..b * hs];
3921 let rms =
3922 (row.iter().map(|&v| (v as f64) * (v as f64)).sum::<f64>() / hs as f64).sqrt();
3923 let mx = row.iter().fold(0f32, |m, &v| m.max(v.abs()));
3924 eprintln!(
3925 "layer {li:>3} {:>10} ffn={:<5} rms={rms:>12.4} max={mx:>12.4}",
3926 match &self.weights.layers[self.phys_layer(li)].attn {
3927 AttnKind::LinearGdn(_) => "gdn",
3928 AttnKind::Linear(_) => "vmf",
3929 AttnKind::ShortConv(_) => "conv",
3930 _ => "attn",
3931 },
3932 match &lw.ffn {
3933 FfnKind::Moe(_) => "moe",
3934 FfnKind::Dense(_) => "dense",
3935 FfnKind::DenseMoe(_) => "dense+moe",
3936 },
3937 );
3938 }
3939 if self.is_loop_end(li) && li + 1 < self.num_layers {
3941 for bi in 0..b {
3942 let normed = inference::rms_norm(
3943 &h[bi * hs..(bi + 1) * hs],
3944 &self.weights.final_norm,
3945 eps,
3946 norm_style,
3947 );
3948 h[bi * hs..(bi + 1) * hs].copy_from_slice(&normed);
3949 }
3950 }
3951 if std::env::var("CMF_TRACE_H").is_ok() {
3952 let n = h[..hs].iter().map(|v| v.abs()).sum::<f32>() / hs as f32;
3953 let mx = h[..hs].iter().fold(0.0f32, |a, &v| a.max(v.abs()));
3954 eprintln!(
3955 "layer {li}: mean|h|={n:.4} max|h|={mx:.2} scale={:?}",
3956 lw.layer_scale
3957 );
3958 }
3959 }
3960 crate::gpu::set_layer(-1); h
3962 }
3963
3964 fn embed_single(&self, id: u32) -> Vec<f32> {
3966 let mut out = vec![0.0f32; self.hidden_size];
3967 if (id as usize) < self.weights.embed_tokens.rows() {
3968 self.weights.embed_tokens.row_f32(id as usize, &mut out);
3969 }
3970 if self.embed_multiplier != 1.0 {
3971 for v in out.iter_mut() {
3972 *v *= self.embed_multiplier;
3973 }
3974 }
3975 if self.dsv4.is_some() {
3979 let mut v = vec![0.0f32; self.hidden_size.max(1)];
3980 v[0] = id as f32;
3981 return v;
3982 }
3983 if let Some(b) = &self.g3n {
3986 return b.0.extend_embedding(id, &out, self.pool.as_deref());
3987 }
3988 out
3989 }
3990
3991 #[cfg(target_os = "macos")]
3997 fn chunk_run_gpu(
3998 &mut self,
3999 li0: usize,
4000 h: &mut [f32],
4001 b: usize,
4002 pos0: usize,
4003 embed_ids: Option<&[u32]>,
4004 cap: usize,
4005 ) -> usize {
4006 if !crate::gpu::enabled_here()
4010 || std::env::var("CMF_GPU_CHUNK")
4011 .map(|v| v == "0")
4012 .unwrap_or(false)
4013 || b < 32
4014 || self.swa.is_some()
4015 || self.global_attn.is_some()
4016 || self.attn_v_norm
4017 || (self.attn_scale - 1.0 / (self.head_dim as f32).sqrt()).abs() > 1e-9
4018 {
4019 return li0;
4020 }
4021 let Some(model) = self.model.clone() else {
4022 return li0;
4023 };
4024 let inv_freq = self.inv_freq.clone();
4025 let (nh, nkv, hd, hs) = (
4026 self.num_heads,
4027 self.num_kv_heads,
4028 self.head_dim,
4029 self.hidden_size,
4030 );
4031 let loop_end = if self.loop_final_norm {
4035 ((li0 / self.physical_layers) + 1) * self.physical_layers
4036 } else {
4037 self.num_layers
4038 };
4039 let mut layers: Vec<crate::gpu_metal::ChunkLayer> = Vec::new();
4040 let mut stored_at: Vec<usize> = Vec::new();
4041 for li in li0..self.num_layers.min(loop_end).min(cap) {
4042 let lw = &self.weights.layers[self.phys_layer(li)];
4043 if lw.attn_out_norm.is_some() || lw.ffn_out_norm.is_some() || lw.layer_scale.is_some() {
4044 break;
4045 }
4046 let AttnKind::Full {
4047 wq,
4048 wk,
4049 wv,
4050 wo,
4051 q_norm,
4052 k_norm,
4053 output_gate: false,
4054 softplus_gate: None,
4055 bias,
4056 } = &lw.attn
4057 else {
4058 break;
4059 };
4060 let FfnKind::Dense(d) = &lw.ffn else { break };
4061 if d.act != Act::Silu {
4062 break;
4063 }
4064 fn cw(t: &QTensor) -> Option<(usize, usize, usize, &[f32])> {
4069 t.q8_row_parts()
4070 .or_else(|| t.q4t_parts().map(|(i, r, c)| (i, r, c, &[][..])))
4071 .or_else(|| t.q4tp_parts().map(|(i, r, c)| (i, r, c, &[][..])))
4072 }
4073 let parts = (
4074 cw(wq),
4075 cw(wk),
4076 cw(wv),
4077 cw(wo),
4078 cw(&d.gate_proj),
4079 cw(&d.up_proj),
4080 cw(&d.down_proj),
4081 );
4082 let (Some(pq), Some(pk), Some(pv), Some(po), Some(pg), Some(pu), Some(pd)) = parts
4083 else {
4084 break;
4085 };
4086 let layer = &self.kv_cache.layers[li];
4087 if layer.mode != crate::kv_cache::KvMode::F32 || layer.o1.is_some() {
4088 break;
4089 }
4090 stored_at.push(layer.head_len(0));
4091 layers.push(crate::gpu_metal::ChunkLayer {
4092 model: &model,
4093 kv_id: self.graph_kv_id,
4094 layer: li,
4095 wq: pq,
4096 wk: pk,
4097 wv: pv,
4098 wo: po,
4099 gate: pg,
4100 up: pu,
4101 down: pd,
4102 input_norm: &lw.input_norm,
4103 post_norm: &lw.post_norm,
4104 bias: bias
4105 .as_ref()
4106 .map(|(a, bb, cc)| (a.as_slice(), bb.as_slice(), cc.as_slice())),
4107 q_norm: q_norm.as_deref(),
4108 k_norm: k_norm.as_deref(),
4109 inv_freq: &inv_freq,
4110 rd: self.rotary_dim,
4111 nh,
4112 nkv,
4113 hd,
4114 hs,
4115 inter: d.gate_proj.rows(),
4116 gemma: matches!(self.norm_style, cortiq_core::NormStyle::Gemma),
4117 eps: self.rms_eps as f32,
4118 });
4119 }
4120 if layers.is_empty() {
4121 return li0;
4122 }
4123 let row = nkv * hd;
4124 let mut store: Vec<(Vec<f32>, Vec<f32>, Vec<f32>)> = stored_at
4125 .iter()
4126 .map(|&st| (vec![0f32; b * row], vec![0f32; b * row], vec![0f32; st + b]))
4127 .collect();
4128 let mut io: Vec<crate::gpu_metal::ChunkIo> = Vec::with_capacity(layers.len());
4129 for (i, (ok, ov, oi)) in store.iter_mut().enumerate() {
4130 let li = layers[i].layer;
4131 let layer = &self.kv_cache.layers[li];
4132 io.push(crate::gpu_metal::ChunkIo {
4133 cpu_stored: stored_at[i],
4134 cpu_k: (0..nkv).map(|g| layer.head_keys(g)).collect(),
4135 cpu_v: (0..nkv).map(|g| layer.head_values(g)).collect(),
4136 out_k: ok,
4137 out_v: ov,
4138 imp: oi,
4139 });
4140 }
4141 let n_run = layers.len();
4142 let last = layers.last().map(|l| l.layer + 1).unwrap_or(li0);
4143 let ep = embed_ids.and_then(|ids| {
4146 self.weights
4147 .embed_tokens
4148 .q8_row_parts()
4149 .map(|(idx, rows, _c, rs)| crate::gpu_metal::ChunkEmbed {
4150 idx,
4151 rows,
4152 row_scale: rs,
4153 ids,
4154 mult: self.embed_multiplier,
4155 })
4156 });
4157 if embed_ids.is_some() && ep.is_none() {
4158 return li0;
4159 }
4160 if !crate::gpu_metal::chunk_run_gpu(&layers, &mut io, h, b, pos0, ep.as_ref()) {
4161 return li0;
4162 }
4163 drop(io);
4164 drop(layers);
4165 for (i, (ok, ov, oi)) in store.iter().enumerate().take(n_run) {
4168 let li = li0 + i;
4169 let layer = &mut self.kv_cache.layers[li];
4170 for bi in 0..b {
4171 layer.append(
4172 &ok[bi * row..(bi + 1) * row],
4173 &ov[bi * row..(bi + 1) * row],
4174 &[],
4175 );
4176 }
4177 layer.accumulate_imp(oi);
4178 }
4179 last
4180 }
4181
4182 fn layer_is_local(&self, li: usize) -> bool {
4185 if let Some(layers) = &self.sliding_layers {
4186 return layers.get(li).copied().unwrap_or(false);
4187 }
4188 match self.swa {
4189 Some((_, pattern)) => (li + 1) % pattern.max(1) != 0,
4190 None => false,
4191 }
4192 }
4193
4194 fn layer_inv_freq(&self, li: usize) -> std::sync::Arc<Vec<f32>> {
4197 if self.layer_is_local(li) {
4198 if let Some(f) = &self.inv_freq_local {
4199 return f.clone();
4200 }
4201 } else if let Some(f) = &self.inv_freq_global {
4202 return f.clone();
4203 }
4204 self.inv_freq.clone()
4205 }
4206
4207 fn layer_window(&self, li: usize) -> Option<usize> {
4209 self.swa
4210 .and_then(|(w, _)| self.layer_is_local(li).then_some(w))
4211 }
4212
4213 fn layer_num_heads(&self, li: usize) -> usize {
4214 self.attention_heads_per_layer
4215 .as_ref()
4216 .and_then(|v| v.get(li).copied())
4217 .unwrap_or(self.num_heads)
4218 }
4219
4220 fn layer_rope_scale(&self, li: usize) -> f32 {
4221 if self.layer_is_local(li) {
4222 self.rope_scale_local
4223 } else {
4224 self.rope_scale
4225 }
4226 }
4227
4228 fn layer_geom(&self, li: usize) -> (usize, usize, usize) {
4231 if !self.layer_is_local(li) {
4232 if let Some((ghd, gkv)) = self.global_attn {
4233 return (gkv, ghd, ghd);
4234 }
4235 }
4236 (
4237 self.num_kv_heads,
4238 self.head_dim,
4239 if self.layer_is_local(li) {
4240 self.rotary_dim_local.unwrap_or(self.rotary_dim)
4241 } else {
4242 self.rotary_dim
4243 },
4244 )
4245 }
4246
4247 fn forward_layers(
4249 &mut self,
4250 hidden: &[f32],
4251 position: usize,
4252 task_mask: Option<&TaskMask>,
4253 ) -> Vec<f32> {
4254 self.forward_layers_upto(hidden, position, task_mask, None)
4255 }
4256
4257 pub fn embed_id(&self, id: u32) -> Vec<f32> {
4265 self.embed_single(id)
4266 }
4267
4268 pub fn split_supported(&self) -> Result<(), String> {
4272 if self.dsv4.is_some() {
4273 return Err(
4274 "network split: DeepSeek-V4 runs its own fused stack (not splittable yet)".into(),
4275 );
4276 }
4277 if self.g3n.is_some() {
4278 return Err(
4279 "network split: Gemma-3n runs its own AltUp stack (not splittable yet)".into(),
4280 );
4281 }
4282 Ok(())
4283 }
4284
4285 pub fn forward_span(
4290 &mut self,
4291 hidden: &[f32],
4292 position: usize,
4293 from: usize,
4294 upto: usize,
4295 task_mask: Option<&TaskMask>,
4296 ) -> Result<Vec<f32>, String> {
4297 self.split_supported()?;
4298 if from > upto || upto >= self.num_layers {
4299 return Err(format!(
4300 "forward_span: layer range {from}..={upto} outside 0..{}",
4301 self.num_layers
4302 ));
4303 }
4304 if hidden.len() != self.hidden_size {
4305 return Err(format!(
4306 "forward_span: hidden len {} ≠ hidden_size {}",
4307 hidden.len(),
4308 self.hidden_size
4309 ));
4310 }
4311 Ok(self.forward_layers_span(hidden, position, task_mask, from, Some(upto)))
4312 }
4313
4314 pub fn logits_from_hidden(&mut self, hidden: &[f32]) -> Vec<f32> {
4317 let normed = inference::rms_norm(
4318 hidden,
4319 &self.weights.final_norm,
4320 self.rms_eps,
4321 self.norm_style,
4322 );
4323 self.lm_head_forward(&normed)
4324 }
4325
4326 pub fn sample_next(&mut self, logits: &[f32], past_tokens: &[u32]) -> u32 {
4328 sampler::sample_with_scratch(
4329 logits,
4330 &self.sampler_config,
4331 past_tokens,
4332 &mut self.rng,
4333 &mut self.sampler_scratch,
4334 )
4335 }
4336
4337 pub fn reset_session(&mut self) {
4339 self.kv_cache.clear();
4340 self.kv_history.clear();
4341 crate::gpu::graph_kv_reset(self.graph_kv_id);
4342 }
4343
4344 pub fn prefill_span_ids(
4350 &mut self,
4351 ids: &[u32],
4352 start_pos: usize,
4353 upto: usize,
4354 task_mask: Option<&TaskMask>,
4355 ) -> Result<Vec<f32>, String> {
4356 self.split_supported()?;
4357 if upto >= self.num_layers {
4358 return Err(format!(
4359 "prefill_span_ids: upto {upto} outside 0..{}",
4360 self.num_layers
4361 ));
4362 }
4363 if self.can_prefill_batched() && !self.graph_prefill_preferred() {
4367 Ok(self.prefill_batch_span(PrefillIn::Ids(ids), start_pos, task_mask, 0, upto + 1))
4368 } else {
4369 let hs = self.hidden_size;
4370 let mut out = Vec::with_capacity(ids.len() * hs);
4371 for (i, &id) in ids.iter().enumerate() {
4372 let emb = self.embed_id(id);
4373 out.extend_from_slice(&self.forward_span(
4374 &emb,
4375 start_pos + i,
4376 0,
4377 upto,
4378 task_mask,
4379 )?);
4380 }
4381 Ok(out)
4382 }
4383 }
4384
4385 pub fn prefill_span_hidden(
4388 &mut self,
4389 hidden: &[f32],
4390 start_pos: usize,
4391 from: usize,
4392 upto: usize,
4393 task_mask: Option<&TaskMask>,
4394 ) -> Result<Vec<f32>, String> {
4395 self.split_supported()?;
4396 let hs = self.hidden_size;
4397 if hidden.is_empty() || hidden.len() % hs != 0 {
4398 return Err(format!(
4399 "prefill_span_hidden: {} floats is not a multiple of hidden {hs}",
4400 hidden.len()
4401 ));
4402 }
4403 if from > upto || upto >= self.num_layers {
4404 return Err(format!(
4405 "prefill_span_hidden: layer range {from}..={upto} outside 0..{}",
4406 self.num_layers
4407 ));
4408 }
4409 if self.can_prefill_batched() && !self.graph_prefill_preferred() {
4410 Ok(self.prefill_batch_span(
4411 PrefillIn::Hidden(hidden),
4412 start_pos,
4413 task_mask,
4414 from,
4415 upto + 1,
4416 ))
4417 } else {
4418 let b = hidden.len() / hs;
4419 let mut out = Vec::with_capacity(hidden.len());
4420 for i in 0..b {
4421 let h = self.forward_span(
4422 &hidden[i * hs..(i + 1) * hs],
4423 start_pos + i,
4424 from,
4425 upto,
4426 task_mask,
4427 )?;
4428 out.extend_from_slice(&h);
4429 }
4430 Ok(out)
4431 }
4432 }
4433
4434 fn try_token_graph_wgpu(
4438 &self,
4439 hidden: &[f32],
4440 position: usize,
4441 logits_out: &mut Vec<f32>,
4442 layers_run: &mut usize,
4443 ) -> Option<Vec<f32>> {
4444 self.try_token_graph_wgpu_steps(
4445 hidden,
4446 position,
4447 logits_out,
4448 1,
4449 None,
4450 Some(layers_run),
4451 0,
4452 self.num_layers,
4453 )
4454 }
4455
4456 fn try_token_graph_wgpu_span(
4460 &self,
4461 hidden: &[f32],
4462 position: usize,
4463 logits_out: &mut Vec<f32>,
4464 from: usize,
4465 upto_excl: usize,
4466 layers_run: &mut usize,
4467 ) -> Option<Vec<f32>> {
4468 self.try_token_graph_wgpu_steps(
4469 hidden,
4470 position,
4471 logits_out,
4472 1,
4473 None,
4474 Some(layers_run),
4475 from,
4476 upto_excl,
4477 )
4478 }
4479
4480 fn try_multi_burst(&self, t_next: u32, position: usize, k: usize) -> Option<Vec<u32>> {
4484 if self.o1_active() || self.attn_softcap > 0.0 {
4485 return None;
4486 }
4487 let graph_on = crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Decode);
4488 if !graph_on || crate::gpu::graph_unsupported() {
4489 return None;
4496 }
4497 let emb = self.embed_single(t_next);
4498 let mut lg = Vec::new();
4499 let mut ids = Vec::new();
4500 self.try_token_graph_wgpu_steps(
4501 &emb,
4502 position,
4503 &mut lg,
4504 k,
4505 Some(&mut ids),
4506 None,
4507 0,
4508 self.num_layers,
4509 )?;
4510 (ids.len() == k).then_some(ids)
4511 }
4512
4513 fn try_token_graph_wgpu_steps(
4517 &self,
4518 hidden: &[f32],
4519 position: usize,
4520 logits_out: &mut Vec<f32>,
4521 steps: usize,
4522 ids_out: Option<&mut Vec<u32>>,
4523 layers_run: Option<&mut usize>,
4524 from: usize,
4525 upto_excl: usize,
4526 ) -> Option<Vec<f32>> {
4527 let o1_gpu = std::env::var("CMF_O1_GPU").as_deref() == Ok("1");
4530 if (self.o1_active() && !o1_gpu) || self.attn_softcap > 0.0 {
4531 return None;
4535 }
4536 let o1_views: Vec<Option<Vec<crate::nystrom::O1DeviceView<'_>>>> = (from..upto_excl)
4541 .map(|li| {
4542 if !o1_gpu {
4543 return None;
4544 }
4545 self.kv_cache.layers[self.phys_layer(li)].o1_views()
4546 })
4547 .collect();
4548 if self.o1_active() && o1_gpu {
4549 let want: usize = (from..upto_excl)
4552 .filter(|li| !matches!(self.kv_cache.layers[self.phys_layer(*li)].o1, None))
4553 .count();
4554 let have = o1_views.iter().filter(|v| v.is_some()).count();
4555 if want == 0 || have != want {
4556 return None;
4557 }
4558 }
4559 let nh = self.num_heads;
4560 let (nkv, hd, rd) = self.layer_geom(0);
4561 let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
4562 let mut layers = Vec::with_capacity(upto_excl - from);
4563 let mut model = None;
4564 let dbg = std::env::var("CMF_GRAPH_DEBUG").is_ok();
4565 fn gw(t: &QTensor) -> Option<crate::gpu::GraphW<'_>> {
4566 if let Some((_, i, kind, rs)) = t.graph_weight() {
4567 return Some(crate::gpu::GraphW {
4568 idx: i,
4569 kind,
4570 row_scale: rs,
4571 data: &[],
4572 });
4573 }
4574 t.as_f32().map(|d| crate::gpu::GraphW {
4576 idx: 0,
4577 kind: 4,
4578 row_scale: &[],
4579 data: d,
4580 })
4581 }
4582 for li in from..upto_excl {
4583 let lw = &self.weights.layers[self.phys_layer(li)];
4584 if dbg {
4585 let ak = match &lw.attn {
4586 AttnKind::Mla(_) => "Mla".into(),
4587 AttnKind::Full {
4588 output_gate, bias, ..
4589 } => format!("Full gate={output_gate} bias={}", bias.is_some()),
4590 AttnKind::LinearGdn(_) => "LinearGdn".into(),
4591 AttnKind::Kda(_) => "Kda".into(),
4592 AttnKind::Linear(_) => "Linear".into(),
4593 AttnKind::ShortConv(_) => "ShortConv".into(),
4594 };
4595 let fk = match &lw.ffn {
4596 FfnKind::Dense(_) => "Dense",
4597 FfnKind::Moe(_) => "Moe",
4598 FfnKind::DenseMoe(_) => "DenseMoe",
4599 };
4600 eprintln!("graph L{li}: attn={ak} ffn={fk}");
4601 }
4602 let gffn = match &lw.ffn {
4603 FfnKind::DenseMoe(_) => return None, FfnKind::Dense(d) => crate::gpu::GraphFfn::Dense {
4605 gate: gw(&d.gate_proj)?,
4606 up: gw(&d.up_proj)?,
4607 down: gw(&d.down_proj)?,
4608 },
4609 FfnKind::Moe(m) => {
4610 if m.router_sigmoid
4615 || m.expert_bias.is_some()
4616 || m.route_tau.is_some()
4617 || m.mask.is_some()
4618 {
4619 return None;
4620 }
4621 let (se, sg) = m.shared.as_ref()?;
4622 let sgate = gw(sg.as_ref()?)?;
4623 let router = gw(&m.router)?;
4624 let inter = m.experts.first()?.gate_proj.rows();
4625 let mut experts = Vec::with_capacity(m.experts.len() + 1);
4626 let mut q4tp: Option<bool> = None;
4629 let mut gu_q2: Option<bool> = None;
4632 for e in m.experts.iter().chain(std::iter::once(se)) {
4633 if !matches!(e.act, Act::Silu)
4634 || e.gate_proj.rows() != inter
4635 || e.up_proj.rows() != inter
4636 {
4637 return None;
4638 }
4639 let (mm, gi, ui, di, is_p, is_q2) = match e.gate_proj.mapped_q4t() {
4640 Some((mm, gi)) => (
4641 mm,
4642 gi,
4643 e.up_proj.mapped_q4t()?.1,
4644 e.down_proj.mapped_q4t()?.1,
4645 false,
4646 false,
4647 ),
4648 None => match e.gate_proj.mapped_q2tp() {
4649 Some((mm, gi)) => (
4650 mm,
4651 gi,
4652 e.up_proj.mapped_q2tp()?.1,
4653 e.down_proj.mapped_q4tp()?.1,
4654 true,
4655 true,
4656 ),
4657 None => {
4658 let (mm, gi) = e.gate_proj.mapped_q4tp()?;
4659 (
4660 mm,
4661 gi,
4662 e.up_proj.mapped_q4tp()?.1,
4663 e.down_proj.mapped_q4tp()?.1,
4664 true,
4665 false,
4666 )
4667 }
4668 },
4669 };
4670 if *q4tp.get_or_insert(is_p) != is_p || *gu_q2.get_or_insert(is_q2) != is_q2
4671 {
4672 tracing::warn!(
4678 "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."
4679 );
4680 return None;
4681 }
4682 model.get_or_insert_with(|| mm.clone());
4683 experts.push((gi, ui, di));
4684 }
4685 crate::gpu::GraphFfn::Moe {
4686 router,
4687 shared_gate: sgate,
4688 experts,
4689 n_exp: m.experts.len(),
4690 top_k: std::env::var("CMF_TOPK_PROBE")
4696 .ok()
4697 .and_then(|v| v.parse::<usize>().ok())
4698 .filter(|k| *k > 0 && *k <= m.top_k)
4699 .unwrap_or(m.top_k),
4700 inter,
4701 norm_topk: m.norm_topk_prob,
4702 q4tp: q4tp?,
4703 gu_q2: gu_q2.unwrap_or(false),
4704 }
4705 }
4706 };
4707 let attn = match &lw.attn {
4708 AttnKind::Full {
4709 wq,
4710 wk,
4711 wv,
4712 wo,
4713 q_norm,
4714 k_norm,
4715 output_gate,
4716 softplus_gate,
4717 bias,
4718 } => {
4719 if softplus_gate.is_some() || self.attention_heads_per_layer.is_some() {
4720 return None;
4721 }
4722 let (m, _, _, _) = wq.graph_weight()?;
4723 model = Some(m.clone());
4724 crate::gpu::GraphAttn::Full {
4725 wq: gw(wq)?,
4726 wk: gw(wk)?,
4727 wv: gw(wv)?,
4728 wo: gw(wo)?,
4729 q_norm: q_norm.as_deref(),
4730 k_norm: k_norm.as_deref(),
4731 bias: bias
4732 .as_ref()
4733 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
4734 output_gate: *output_gate,
4735 cpu_k: self.kv_cache.layers[li].k_heads(),
4736 cpu_v: self.kv_cache.layers[li].v_heads(),
4737 }
4738 }
4739 AttnKind::LinearGdn(w) => {
4740 let cfg = self.gdn_cfg?;
4741 let (m, _, _, _) = w.in_proj_qkv.graph_weight()?;
4742 model = Some(m.clone());
4743 crate::gpu::GraphAttn::Gdn {
4744 qkv: gw(&w.in_proj_qkv)?,
4745 z: gw(&w.in_proj_z)?,
4746 a: gw(&w.in_proj_a)?,
4747 b: gw(&w.in_proj_b)?,
4748 out: gw(&w.out_proj)?,
4749 conv1d: &w.conv1d,
4750 a_log: &w.a_log,
4751 dt_bias: &w.dt_bias,
4752 norm: &w.norm,
4753 nv: cfg.num_v_heads,
4754 nk: cfg.num_k_heads,
4755 dk: cfg.key_head_dim,
4756 dv: cfg.value_head_dim,
4757 kk: cfg.conv_kernel,
4758 cpu_state: &self.kv_cache.layers[self.phys_layer(li)].linear_state,
4759 }
4760 }
4761 _ => return None,
4762 };
4763 layers.push(crate::gpu::GraphLayer {
4764 input_norm: &lw.input_norm,
4765 attn,
4766 post_norm: &lw.post_norm,
4767 ffn: gffn,
4768 });
4769 }
4770 let model = model?;
4771 let lm_gw = if upto_excl == self.num_layers
4777 && self.graph_want_logits
4778 && std::env::var("CMF_GPU_LMHEAD")
4779 .map(|v| v != "0")
4780 .unwrap_or(true)
4781 {
4782 self.weights.lm_head.graph_weight().map(|(_, i, kind, rs)| {
4783 (
4784 crate::gpu::GraphW {
4785 idx: i,
4786 kind,
4787 row_scale: rs,
4788 data: &[],
4789 },
4790 self.weights.lm_head.rows(),
4791 )
4792 })
4793 } else {
4794 None
4795 };
4796 let lm = lm_gw.as_ref().map(|(gw, rows)| (gw, *rows));
4797 let emb_gw = if steps > 1 {
4799 self.weights
4800 .embed_tokens
4801 .graph_weight()
4802 .map(|(_, i, kind, rs)| {
4803 (
4804 crate::gpu::GraphW {
4805 idx: i,
4806 kind,
4807 row_scale: rs,
4808 data: &[],
4809 },
4810 self.weights.embed_tokens.rows(),
4811 self.embed_multiplier as f32,
4812 )
4813 })
4814 } else {
4815 None
4816 };
4817
4818 let loop_norm_at: Vec<usize> = if self.loop_final_norm {
4824 (from..upto_excl.min(self.num_layers - 1))
4825 .filter(|&li| (li + 1) % self.physical_layers == 0)
4826 .map(|li| li - from)
4827 .collect()
4828 } else {
4829 Vec::new()
4830 };
4831 let mut h = hidden.to_vec();
4832 crate::gpu::forward_token_graph(
4833 &model,
4834 self.graph_kv_id,
4835 &layers,
4836 &o1_views,
4837 self.o1_epoch,
4838 &self.inv_freq,
4839 &mut h,
4840 nh,
4841 nkv,
4842 hd,
4843 rd,
4844 self.hidden_size,
4845 self.intermediate_size,
4846 position,
4847 self.kv_cache.max_seq_len,
4848 gemma,
4849 self.rms_eps as f32,
4850 lm,
4851 &self.weights.final_norm,
4852 logits_out,
4853 &loop_norm_at,
4854 steps,
4855 emb_gw.as_ref().map(|(gw, rows, m)| (gw, *rows, *m)),
4856 ids_out,
4857 layers_run,
4858 from,
4859 )
4860 .then_some(h)
4861 }
4862
4863 fn try_batch_graph_wgpu(
4868 &self,
4869 hiddens: &mut [f32],
4870 positions: &[usize],
4871 k: usize,
4872 spec: Option<crate::gpu::SpecTail<'_>>,
4873 ) -> bool {
4874 let _tb = std::time::Instant::now();
4875 if self.attn_softcap > 0.0 {
4876 return false; }
4878 if self.o1_active() {
4879 return false;
4880 }
4881 let nh = self.num_heads;
4882 let (nkv, hd, rd) = self.layer_geom(0);
4883 let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
4884 fn gw(t: &QTensor) -> Option<crate::gpu::GraphW<'_>> {
4885 if let Some((_, i, kind, rs)) = t.graph_weight() {
4886 return Some(crate::gpu::GraphW {
4887 idx: i,
4888 kind,
4889 row_scale: rs,
4890 data: &[],
4891 });
4892 }
4893 t.as_f32().map(|d| crate::gpu::GraphW {
4894 idx: 0,
4895 kind: 4,
4896 row_scale: &[],
4897 data: d,
4898 })
4899 }
4900 let built: Option<(
4901 Vec<crate::gpu::GraphLayer<'_>>,
4902 std::sync::Arc<cortiq_core::CmfModel>,
4903 )> = (|| {
4904 let mut layers = Vec::with_capacity(self.num_layers);
4905 let mut model = None;
4906 for li in 0..self.num_layers {
4907 let lw = &self.weights.layers[self.phys_layer(li)];
4908 let gffn = match &lw.ffn {
4915 FfnKind::Dense(d) => crate::gpu::GraphFfn::Dense {
4916 gate: gw(&d.gate_proj)?,
4917 up: gw(&d.up_proj)?,
4918 down: gw(&d.down_proj)?,
4919 },
4920 FfnKind::Moe(m) => {
4921 if m.router_sigmoid
4922 || m.expert_bias.is_some()
4923 || m.route_tau.is_some()
4924 || m.mask.is_some()
4925 {
4926 return None;
4927 }
4928 let (se, sg) = m.shared.as_ref()?;
4929 let sgate = gw(sg.as_ref()?)?;
4930 let router = gw(&m.router)?;
4931 let inter = m.experts.first()?.gate_proj.rows();
4932 let mut experts = Vec::with_capacity(m.experts.len() + 1);
4933 let mut q4tp: Option<bool> = None;
4934 let mut gu_q2: Option<bool> = None;
4935 for e in m.experts.iter().chain(std::iter::once(se)) {
4936 if !matches!(e.act, Act::Silu)
4937 || e.gate_proj.rows() != inter
4938 || e.up_proj.rows() != inter
4939 {
4940 return None;
4941 }
4942 let (mm, gi, ui, di, is_p, is_q2) = match e.gate_proj.mapped_q4t() {
4946 Some((mm, gi)) => (
4947 mm,
4948 gi,
4949 e.up_proj.mapped_q4t()?.1,
4950 e.down_proj.mapped_q4t()?.1,
4951 false,
4952 false,
4953 ),
4954 None => match e.gate_proj.mapped_q2tp() {
4955 Some((mm, gi)) => (
4956 mm,
4957 gi,
4958 e.up_proj.mapped_q2tp()?.1,
4959 e.down_proj.mapped_q4tp()?.1,
4960 true,
4961 true,
4962 ),
4963 None => {
4964 let (mm, gi) = e.gate_proj.mapped_q4tp()?;
4965 (
4966 mm,
4967 gi,
4968 e.up_proj.mapped_q4tp()?.1,
4969 e.down_proj.mapped_q4tp()?.1,
4970 true,
4971 false,
4972 )
4973 }
4974 },
4975 };
4976 if *q4tp.get_or_insert(is_p) != is_p
4977 || *gu_q2.get_or_insert(is_q2) != is_q2
4978 {
4979 return None;
4980 }
4981 model.get_or_insert_with(|| mm.clone());
4982 experts.push((gi, ui, di));
4983 }
4984 crate::gpu::GraphFfn::Moe {
4985 router,
4986 shared_gate: sgate,
4987 experts,
4988 n_exp: m.experts.len(),
4989 top_k: m.top_k,
4990 inter,
4991 norm_topk: m.norm_topk_prob,
4992 q4tp: q4tp?,
4993 gu_q2: gu_q2.unwrap_or(false),
4994 }
4995 }
4996 _ => return None,
4997 };
4998 let attn = match &lw.attn {
4999 AttnKind::Full {
5000 wq,
5001 wk,
5002 wv,
5003 wo,
5004 q_norm,
5005 k_norm,
5006 output_gate,
5007 softplus_gate,
5008 bias,
5009 } => {
5010 if softplus_gate.is_some() || self.attention_heads_per_layer.is_some() {
5011 return None;
5012 }
5013 let (m, _, _, _) = wq.graph_weight()?;
5014 model = Some(m.clone());
5015 crate::gpu::GraphAttn::Full {
5016 wq: gw(wq)?,
5017 wk: gw(wk)?,
5018 wv: gw(wv)?,
5019 wo: gw(wo)?,
5020 q_norm: q_norm.as_deref(),
5021 k_norm: k_norm.as_deref(),
5022 bias: bias
5023 .as_ref()
5024 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
5025 output_gate: *output_gate,
5026 cpu_k: self.kv_cache.layers[li].k_heads(),
5027 cpu_v: self.kv_cache.layers[li].v_heads(),
5028 }
5029 }
5030 AttnKind::LinearGdn(w) => {
5031 let cfg = self.gdn_cfg?;
5032 let (m, _, _, _) = w.in_proj_qkv.graph_weight()?;
5033 model = Some(m.clone());
5034 crate::gpu::GraphAttn::Gdn {
5035 qkv: gw(&w.in_proj_qkv)?,
5036 z: gw(&w.in_proj_z)?,
5037 a: gw(&w.in_proj_a)?,
5038 b: gw(&w.in_proj_b)?,
5039 out: gw(&w.out_proj)?,
5040 conv1d: &w.conv1d,
5041 a_log: &w.a_log,
5042 dt_bias: &w.dt_bias,
5043 norm: &w.norm,
5044 nv: cfg.num_v_heads,
5045 nk: cfg.num_k_heads,
5046 dk: cfg.key_head_dim,
5047 dv: cfg.value_head_dim,
5048 kk: cfg.conv_kernel,
5049 cpu_state: &self.kv_cache.layers[self.phys_layer(li)].linear_state,
5050 }
5051 }
5052 _ => return None,
5053 };
5054 layers.push(crate::gpu::GraphLayer {
5055 input_norm: &lw.input_norm,
5056 attn,
5057 post_norm: &lw.post_norm,
5058 ffn: gffn,
5059 });
5060 }
5061 Some((layers, model?))
5062 })();
5063 let Some((layers, model)) = built else {
5064 {
5065 use std::sync::atomic::{AtomicBool, Ordering};
5066 static SAID: AtomicBool = AtomicBool::new(false);
5067 if !SAID.swap(true, Ordering::Relaxed) {
5068 tracing::warn!("batch graph: BUILDER refused (layer weights/kinds)");
5069 }
5070 }
5071 return false;
5072 };
5073 if std::env::var("CMF_GRAPH_SPEC_TIME").is_ok() {
5074 eprintln!("batch-build: {:.1} ms", _tb.elapsed().as_secs_f64() * 1e3);
5075 }
5076 crate::gpu::forward_batch_graph(
5077 &model,
5078 self.graph_kv_id,
5079 &layers,
5080 &self.inv_freq,
5081 hiddens,
5082 nh,
5083 nkv,
5084 hd,
5085 rd,
5086 self.hidden_size,
5087 self.intermediate_size,
5088 positions,
5089 self.kv_cache.max_seq_len,
5090 gemma,
5091 self.rms_eps as f32,
5092 k,
5093 spec,
5094 )
5095 }
5096
5097 fn draft_probe() -> bool {
5101 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
5102 *ON.get_or_init(|| std::env::var("CMF_DSV4_DRAFT_PROBE").is_ok_and(|v| v != "0"))
5103}
5104
5105 #[cfg(feature = "gpu")]
5117 fn dsv4_spec_on() -> bool {
5118 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
5119 *ON.get_or_init(|| std::env::var("CMF_DSV4_SPEC").map(|v| v != "0").unwrap_or(true))
5120 }
5121
5122 #[cfg(feature = "gpu")]
5129 fn dsv4_spec_step(
5130 &mut self,
5131 tip_token: u32,
5132 t_next: u32,
5133 next_pos: usize,
5134 drafted: &mut usize,
5135 accepted_ctr: &mut usize,
5136 ) -> Option<(Vec<u32>, usize)> {
5137 let t_all = std::time::Instant::now();
5138 if std::env::var("CMF_DSV4_SPEC_TIME").is_ok() {
5139 thread_local! {
5140 static LAST: std::cell::Cell<Option<std::time::Instant>> =
5141 const { std::cell::Cell::new(None) };
5142 }
5143 LAST.with(|l| {
5144 if let Some(prev) = l.get() {
5145 eprintln!("между раундами {:.1} мс", prev.elapsed().as_secs_f64() * 1e3);
5146 }
5147 l.set(Some(std::time::Instant::now()));
5148 });
5149 }
5150 if std::env::var("CMF_DSV4_SPEC_DEBUG").is_ok() {
5151 eprintln!("spec_step: вход pos={next_pos}");
5152 }
5153 let n_layers = self.dsv4.as_ref().map(|b| b.1.len())?;
5154 let cfg = self.dsv4.as_ref().map(|b| b.2)?;
5155 if self.dspark.is_none() {
5157 let t = crate::dsv4::dspark_targets(&self.dsv4_mtp, &cfg, n_layers);
5158 if t.is_empty() {
5159 return None;
5160 }
5161 crate::dsv4::dspark_arm(&t, cfg.dim);
5162 self.dspark = Some(crate::dsv4::DsparkState::new(
5163 self.dsv4_mtp.len(),
5164 &cfg,
5165 t.len(),
5166 ));
5167 }
5168 let targets = crate::dsv4::dspark_targets(&self.dsv4_mtp, &cfg, n_layers);
5169 let pack = crate::dsv4::dspark_pack_get(&self.dsv4_mtp, &cfg);
5170 if pack.is_none() && std::env::var("CMF_DSV4_SPEC_DEBUG").is_ok() {
5171 eprintln!("spec_step: пак не построился (targets {targets:?})");
5172 }
5173 let pack = pack?;
5174 let block = crate::dsv4::dspark_block();
5175 let b_box = self.dsv4.as_mut()?;
5176 let (g, layers, st) = (&b_box.0, &b_box.1, &mut b_box.3);
5177 let ds = self.dspark.as_mut()?;
5178 let dbg = std::env::var("CMF_DSV4_SPEC_DEBUG").is_ok();
5181 if !crate::dsv4::dspark_take(&mut ds.main_hidden) && !ds.have_hidden {
5182 if dbg {
5183 eprintln!("spec_step: нет захвата");
5184 }
5185 return None;
5186 }
5187 ds.have_hidden = true;
5188 let tip_pos = next_pos.checked_sub(1)?;
5189 let draft_started = std::time::Instant::now();
5190 let mut conf = Vec::new();
5191 let props = crate::dsv4::dspark_draft_gpu(
5192 g,
5193 &self.dsv4_mtp,
5194 &cfg,
5195 ds,
5196 pack,
5197 st.kv_id,
5198 tip_token,
5199 tip_pos,
5200 self.pool.as_deref(),
5201 &mut conf,
5202 );
5203 self.dspark_draft_ns += draft_started.elapsed().as_nanos();
5204 *drafted += block;
5205 if props.is_empty() || props[0] != t_next {
5206 if dbg {
5207 eprintln!(
5208 "spec_step: черновик {} (props0={:?} t_next={t_next})",
5209 if props.is_empty() { "пуст" } else { "мимо" },
5210 props.first()
5211 );
5212 }
5213 return None;
5214 }
5215 let mut k_verify = crate::dsv4::dspark_verify_k().min(props.len());
5216 let conf_min = {
5222 static M: std::sync::OnceLock<f32> = std::sync::OnceLock::new();
5223 *M.get_or_init(|| {
5224 std::env::var("CMF_DSPARK_CONF_MIN")
5225 .ok()
5226 .and_then(|v| v.parse().ok())
5227 .unwrap_or(0.0)
5228 })
5229 };
5230 if conf_min > 0.0 && conf.len() >= props.len() {
5231 let mut keep = 1usize;
5232 while keep < k_verify && conf.get(keep).copied().unwrap_or(0.0) >= conf_min {
5233 keep += 1;
5234 }
5235 k_verify = k_verify.min(keep.max(2));
5236 }
5237 if k_verify < 2 {
5238 return None;
5239 }
5240 let mut fed = Vec::with_capacity(k_verify);
5241 fed.push(t_next);
5242 fed.extend_from_slice(&props[1..k_verify]);
5243 let mut argmax = Vec::new();
5244 let mut logits_all = Vec::new();
5245 let mut walked = Vec::new();
5246 let txn = crate::dsv4::dsv4_verify_chunk(
5247 g,
5248 layers,
5249 &cfg,
5250 st,
5251 &fed,
5252 next_pos,
5253 &self.inv_freq,
5254 self.pool.as_deref(),
5255 &targets,
5256 &mut argmax,
5257 &mut logits_all,
5258 &mut walked,
5259 );
5260 if txn.is_none() && dbg {
5261 eprintln!("spec_step: verify отказал");
5262 }
5263 let txn = txn?;
5264 let b = fed.len();
5265 let mut accepted = 1usize;
5266 while accepted < b && fed[accepted] == argmax[accepted - 1] {
5267 accepted += 1;
5268 }
5269 if std::env::var("CMF_DSV4_SPEC_FORCE_REJECT").is_ok_and(|v| v != "0") {
5274 accepted = 1;
5275 }
5276 if std::env::var("CMF_DSV4_SPEC_TRACE").is_ok() {
5277 eprintln!(
5278 "spec@{next_pos}: fed={fed:?} argmax={argmax:?} accepted={accepted}"
5279 );
5280 }
5281 let t_fin = std::time::Instant::now();
5282 if !crate::dsv4::dsv4_spec_finish(
5283 g,
5284 layers,
5285 &cfg,
5286 st,
5287 txn,
5288 accepted,
5289 &fed,
5290 &self.inv_freq,
5291 self.pool.as_deref(),
5292 ) {
5293 tracing::warn!("dsv4: спекулятивный откат не удался — состояние подозрительно");
5294 return None;
5295 }
5296 if std::env::var("CMF_DSV4_SPEC_TIME").is_ok() {
5297 eprintln!("finish(k={accepted}): {:.1} мс", t_fin.elapsed().as_secs_f64() * 1e3);
5298 }
5299 *accepted_ctr += accepted - 1;
5300 let (hc, dim) = (cfg.hc_mult, cfg.dim);
5305 let dev_caps: Vec<usize> = targets
5312 .iter()
5313 .copied()
5314 .filter(|&t| {
5315 st.dev_set.get(t).copied().unwrap_or(false)
5316 && !st.partial_set.get(t).copied().unwrap_or(false)
5317 })
5318 .collect();
5319 let mut caps_all = vec![0.0f32; dev_caps.len() * b * hc * dim];
5320 if !crate::gpu_wgpu::dsv4_spec_cap_read_all(b, dev_caps.len(), hc * dim, &mut caps_all) {
5321 return None;
5322 }
5323 for t in 0..accepted {
5324 let tip = t + 1 == accepted;
5325 for (slot, &tl) in targets.iter().enumerate() {
5326 if let Some(di) = dev_caps.iter().position(|&d| d == tl) {
5327 let lo = (di * b + t) * hc * dim;
5328 crate::dsv4::dspark_capture(
5329 &caps_all[lo..lo + hc * dim],
5330 &cfg,
5331 slot,
5332 &mut ds.main_hidden,
5333 );
5334 } else if tip
5335 && crate::dsv4::dspark_peek_slot(slot, dim, {
5336 let lo = slot * dim;
5337 &mut ds.main_hidden[lo..lo + dim]
5338 })
5339 {
5340 } else {
5345 crate::dsv4::dspark_capture(
5349 &walked[t * hc * dim..(t + 1) * hc * dim],
5350 &cfg,
5351 slot,
5352 &mut ds.main_hidden,
5353 );
5354 }
5355 }
5356 crate::dsv4::dspark_ring_append(g, &self.dsv4_mtp, &cfg, ds, next_pos + t, self.pool.as_deref());
5357 }
5358 let row = logits_all[(accepted - 1) * cfg.vocab..accepted * cfg.vocab].to_vec();
5359 self.graph_logits = Some(row);
5360 if std::env::var("CMF_DSV4_TRUNK_PICK_DUMP").is_ok() {
5365 crate::dsv4::trunk_freq_note(&crate::dsv4::pick_tally_take());
5366 crate::dsv4::pick_tally_arm();
5367 }
5368 if std::env::var("CMF_DSV4_SPEC_TIME").is_ok() {
5369 eprintln!("spec_step total {:.1} мс (k={accepted})", t_all.elapsed().as_secs_f64() * 1e3);
5370 }
5371 Some((fed[1..accepted].to_vec(), next_pos + accepted))
5372 }
5373
5374 fn dspark_probe(&mut self, position: usize, token_id: u32) {
5375 if self.dsv4_mtp.is_empty() || !Self::draft_probe() {
5376 return;
5377 }
5378 let trunk_now = crate::dsv4::pick_tally_take();
5380 crate::dsv4::trunk_freq_note(&trunk_now);
5381 if !trunk_now.is_empty() {
5382 self.dspark_trunk_picks.push(trunk_now);
5383 let keep = crate::dsv4::dspark_block();
5384 if self.dspark_trunk_picks.len() > keep {
5385 self.dspark_trunk_picks.remove(0);
5386 }
5387 }
5388 for p in std::mem::take(&mut self.dspark_pending) {
5391 let Some(i) = position.checked_sub(p.0 + 1) else {
5392 continue;
5393 };
5394 let mut p = p;
5395 if i < p.1.len() {
5396 if p.2 && p.1[i] == token_id {
5397 p.3 = i + 1;
5398 } else {
5399 p.2 = false;
5400 }
5401 if i + 1 < p.1.len() {
5402 self.dspark_pending.push(p);
5403 continue;
5404 }
5405 }
5406 self.dspark_hist.push(p.3);
5407 self.dspark_real.push(token_id);
5408 }
5409 let Some(b) = &mut self.dsv4 else { return };
5410 let (g, layers, cfg) = (&b.0, &b.1, b.2);
5411 let n_layers = layers.len();
5412 if self.dspark.is_none() {
5413 let t = crate::dsv4::dspark_targets(&self.dsv4_mtp, &cfg, n_layers);
5414 if t.is_empty() {
5415 return;
5416 }
5417 eprintln!("DSpark: захват со слоёв {t:?}, блок {}", crate::dsv4::dspark_block());
5418 crate::dsv4::dspark_arm(&t, cfg.dim);
5419 self.dspark = Some(crate::dsv4::DsparkState::new(
5420 self.dsv4_mtp.len(),
5421 &cfg,
5422 t.len(),
5423 ));
5424 }
5425 let ds = self.dspark.as_mut().unwrap();
5426 if !crate::dsv4::dspark_take(&mut ds.main_hidden) {
5427 return; }
5429 let mut conf = Vec::new();
5430 crate::dsv4::pick_tally_arm();
5431 let draft_started = std::time::Instant::now();
5436 #[cfg(feature = "gpu")]
5437 let gpu_draft = crate::dsv4::dspark_gpu_on();
5438 #[cfg(not(feature = "gpu"))]
5439 let gpu_draft = false;
5440 let props = if gpu_draft {
5441 #[cfg(feature = "gpu")]
5442 {
5443 let kv_id = b.3.kv_id;
5444 match crate::dsv4::dspark_pack_get(&self.dsv4_mtp, &cfg) {
5445 Some(pk) => crate::dsv4::dspark_draft_gpu(
5446 g,
5447 &self.dsv4_mtp,
5448 &cfg,
5449 ds,
5450 pk,
5451 kv_id,
5452 token_id,
5453 position,
5454 self.pool.as_deref(),
5455 &mut conf,
5456 ),
5457 None => Vec::new(),
5458 }
5459 }
5460 #[cfg(not(feature = "gpu"))]
5461 Vec::new()
5462 } else {
5463 crate::gpu::cpu_scope(|| {
5464 crate::dsv4::dspark_draft(
5465 g,
5466 &self.dsv4_mtp,
5467 &cfg,
5468 ds,
5469 token_id,
5470 position,
5471 self.pool.as_deref(),
5472 &mut conf,
5473 )
5474 })
5475 };
5476 self.dspark_draft_ns += draft_started.elapsed().as_nanos();
5477 let draft_picks = crate::dsv4::pick_tally_take();
5478 crate::dsv4::dspark_freq_note(&draft_picks);
5479 crate::dsv4::pick_tally_arm();
5482 if !props.is_empty() {
5483 let (tu, tt) = {
5487 let flat: Vec<(usize, Vec<usize>)> = self
5488 .dspark_trunk_picks
5489 .iter()
5490 .flat_map(|v| v.iter().cloned())
5491 .collect();
5492 let mut per: std::collections::HashMap<usize, Vec<usize>> =
5494 std::collections::HashMap::new();
5495 for (li, picks) in flat {
5496 per.entry(li).or_default().extend(picks);
5497 }
5498 let n = per.len().max(1);
5499 let mut u = 0usize;
5500 let mut t = 0usize;
5501 for (_, v) in per {
5502 t += v.len();
5503 u += v.iter().collect::<std::collections::HashSet<_>>().len();
5504 }
5505 (u / n, t / n)
5506 };
5507 let (du, dt) = crate::dsv4::tally_unique(&draft_picks);
5508 self.dspark_exp.push((tu, tt, du, dt));
5509 self.dspark_pending.push((position, props, true, 0));
5510 }
5511 if self.dspark_hist.len() >= 8 && self.dspark_hist.len() % 8 == 0 {
5512 let n = self.dspark_hist.len() as f32;
5513 let mean: f32 = self.dspark_hist.iter().sum::<usize>() as f32 / n;
5514 let block = crate::dsv4::dspark_block();
5515 let mut at = vec![0usize; block + 1];
5516 for &k in &self.dspark_hist {
5517 at[k] += 1;
5518 }
5519 let mut surv = Vec::with_capacity(block);
5521 for i in 1..=block {
5522 let k = at[i..].iter().sum::<usize>() as f32 / n;
5523 surv.push(format!("{k:.2}"));
5524 }
5525 let distinct = self
5526 .dspark_real
5527 .iter()
5528 .collect::<std::collections::HashSet<_>>()
5529 .len();
5530 let (tu, tt, du, dt) = self.dspark_exp.iter().fold((0, 0, 0, 0), |a, b| {
5531 (a.0 + b.0, a.1 + b.1, a.2 + b.2, a.3 + b.3)
5532 });
5533 let m = self.dspark_exp.len().max(1);
5534 eprintln!(
5535 "DSpark: черновиков {}, принято в среднем {mean:.2} из {block} \
5536 (токенов за проход {:.2}), распределение {at:?}, выживание [{}]",
5537 self.dspark_hist.len(),
5538 mean + 1.0,
5539 surv.join(" ")
5540 );
5541 eprintln!(
5542 "DSpark: разных токенов {distinct} из {} (вырожденность), \
5543 эксперты ствол {}/{} на слой за {block} токенов, \
5544 черновик {}/{} за блок, draft {:.2} мс/блок",
5545 self.dspark_real.len(),
5546 tu / m,
5547 tt / m,
5548 du / m,
5549 dt / m,
5550 self.dspark_draft_ns as f64 / self.dspark_exp.len().max(1) as f64 / 1e6
5551 );
5552 }
5553 }
5554
5555 fn forward_layers_upto(
5556 &mut self,
5557 hidden: &[f32],
5558 position: usize,
5559 task_mask: Option<&TaskMask>,
5560 upto: Option<usize>,
5561 ) -> Vec<f32> {
5562 if let Some(plan) = self.gpu_plan.clone() {
5568 if upto.is_none() && plan.len() > 1 {
5569 let mut h = hidden.to_vec();
5570 for &(dev, from, upto_incl) in plan.iter() {
5571 h = crate::gpu::with_device(dev, || {
5572 self.forward_layers_span(&h, position, task_mask, from, Some(upto_incl))
5573 });
5574 }
5575 return h;
5576 }
5577 }
5578 self.forward_layers_span(hidden, position, task_mask, 0, upto)
5579 }
5580
5581 pub fn set_gpu_plan(&mut self, devices: Option<&[usize]>) -> Result<(), String> {
5586 self.set_gpu_plan_at(devices, None)
5587 }
5588
5589 pub fn set_gpu_plan_at(
5593 &mut self,
5594 devices: Option<&[usize]>,
5595 at: Option<usize>,
5596 ) -> Result<(), String> {
5597 let Some(devs) = devices.filter(|d| d.len() > 1) else {
5598 self.gpu_plan = None;
5599 return Ok(());
5600 };
5601 self.split_supported()?;
5602 let n = self.num_layers;
5603 if devs.len() > n {
5604 return Err(format!("{} devices for {n} layers", devs.len()));
5605 }
5606 if let Some(k) = at {
5607 if k == 0 || k >= n {
5608 return Err(format!("split at {k}: the model has {n} layers"));
5609 }
5610 if devs.len() == 2 {
5611 self.gpu_plan = Some(std::sync::Arc::new(vec![
5612 (devs[0], 0, k - 1),
5613 (devs[1], k, n - 1),
5614 ]));
5615 return Ok(());
5616 }
5617 return Err(format!(
5618 "an explicit split point takes exactly 2 devices, got {}",
5619 devs.len()
5620 ));
5621 }
5622 let per = n.div_ceil(devs.len());
5623 let mut plan = Vec::with_capacity(devs.len());
5624 let mut from = 0usize;
5625 for &d in devs {
5626 if from >= n {
5627 break;
5628 }
5629 let upto = (from + per - 1).min(n - 1);
5630 plan.push((d, from, upto));
5631 from = upto + 1;
5632 }
5633 self.gpu_plan = Some(std::sync::Arc::new(plan));
5634 Ok(())
5635 }
5636
5637 pub fn gpu_plan(&self) -> Option<Vec<(usize, usize, usize)>> {
5639 self.gpu_plan.as_ref().map(|p| p.as_ref().clone())
5640 }
5641
5642 fn forward_layers_span(
5648 &mut self,
5649 hidden: &[f32],
5650 position: usize,
5651 task_mask: Option<&TaskMask>,
5652 from: usize,
5653 upto: Option<usize>,
5654 ) -> Vec<f32> {
5655 debug_assert!(from == 0 || (self.dsv4.is_none() && self.g3n.is_none()));
5656 if let Some(b) = &mut self.dsv4 {
5662 let _ = (task_mask, upto);
5663 let token_id = hidden.first().copied().unwrap_or(0.0) as u32;
5664 let (g, layers, cfg, st) = (&b.0, &b.1, b.2, &mut b.3);
5665 st.pos = position;
5666 let mut logits = Vec::new();
5667 crate::dsv4::forward_token(
5668 g,
5669 layers,
5670 &cfg,
5671 st,
5672 token_id,
5673 &self.inv_freq,
5674 self.pool.as_deref(),
5675 &mut logits,
5676 );
5677 self.graph_logits = Some(logits);
5678 self.dspark_probe(position, token_id);
5679 return vec![0.0; self.hidden_size];
5682 }
5683 if let Some(b) = &self.g3n {
5686 let _ = (task_mask, upto);
5687 return crate::g3n::g3n_forward(
5688 &b.0,
5689 &b.1,
5690 hidden,
5691 position,
5692 &mut self.kv_cache.layers,
5693 self.num_heads,
5694 self.num_kv_heads,
5695 self.head_dim,
5696 self.pool.as_deref(),
5697 );
5698 }
5699 let mut h = hidden.to_vec();
5700 let (nh, _nkv, _hd, hs, _rd, eps) = (
5703 self.num_heads,
5704 self.num_kv_heads,
5705 self.head_dim,
5706 self.hidden_size,
5707 self.rotary_dim,
5708 self.rms_eps,
5709 );
5710 let pool = self.pool.clone();
5711 let graph_env = std::env::var("CMF_GPU_WGPU_GRAPH").ok();
5723 let graph_on = match graph_env.as_deref() {
5724 Some("0") => false,
5725 Some("prefill") => false, Some(_) => true,
5727 None => crate::gpu::wgpu_graph_default(),
5733 };
5734 let graph_trusted =
5735 graph_env.is_some() || crate::gpu::wgpu_graph_default() || self.gdn_cfg.is_some();
5736 let race_eligible = graph_on
5737 && upto.is_none()
5738 && task_mask.is_none()
5739 && from == 0
5740 && !crate::gpu::graph_unsupported();
5741 let mut tail_start = 0usize;
5742 if race_eligible && crate::gpu::graph_race_use_graph(graph_trusted) {
5743 let t_graph = std::time::Instant::now();
5744 let mut lg = Vec::new();
5745 let mut gl = 0usize;
5746 let built = self.try_token_graph_wgpu(hidden, position, &mut lg, &mut gl);
5747 if built.is_none() && !self.o1_active() && self.attn_softcap == 0.0 {
5752 crate::gpu::graph_mark_unsupported();
5753 }
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}