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 pub tokenizer: std::sync::Arc<Tokenizer>,
53 pub kv_cache: KvCache,
54 pub sampler_config: SamplerConfig,
55 pub weights: PipelineWeights,
56 pub hidden_size: usize,
57 pub intermediate_size: usize,
58 pub num_heads: usize,
59 pub num_kv_heads: usize,
60 pub head_dim: usize,
61 pub num_layers: usize,
63 pub physical_layers: usize,
65 pub loop_final_norm: bool,
67 pub vocab_size: usize,
68 pub rms_eps: f64,
69 pub rope_base: f32,
70 pub norm_style: NormStyle,
71 pub rotary_dim: usize,
73 pub attention_heads_per_layer: Option<Vec<usize>>,
75 pub vmf_cfg: Option<VmfPhaseCfg>,
77 pub gdn_cfg: Option<GdnCfg>,
79 pub logit_multiplier: Option<f32>,
81 pub cancel: std::sync::Arc<std::sync::atomic::AtomicBool>,
86 pub kv_history: Vec<u32>,
91 pub kda_cfg: Option<crate::linear_core::KdaCfg>,
93 pub g3n: Option<Box<(crate::g3n::G3nGlobals, Vec<crate::g3n::G3nLayer>)>>,
96 pub dsv4: Option<
100 Box<(
101 crate::dsv4::Dsv4Globals,
102 Vec<crate::dsv4::Dsv4Layer>,
103 crate::dsv4::Dsv4Cfg,
104 crate::dsv4::Dsv4State,
105 )>,
106 >,
107 pub dsv4_mtp: Vec<crate::dsv4::Dsv4Mtp>,
111 pub dspark: Option<crate::dsv4::DsparkState>,
113 pub dspark_pending: Vec<(usize, Vec<u32>, bool, usize)>,
116 pub dspark_hist: Vec<usize>,
118 pub dspark_real: Vec<u32>,
122 pub dspark_trunk_picks: Vec<Vec<(usize, Vec<usize>)>>,
126 pub dspark_exp: Vec<(usize, usize, usize, usize)>,
128 pub dspark_draft_ns: u128,
132 pub short_conv_cfg: Option<ShortConvCfg>,
135 pub mtp: Option<MtpModule>,
137 pub speculative: bool,
139 rng: SplitMix64,
140 sampler_scratch: SamplerScratch,
141 pub(crate) inv_freq: std::sync::Arc<Vec<f32>>,
145 ws: ForwardScratch,
149 pool: Option<std::sync::Arc<Pool>>,
151 pub(crate) model: Option<std::sync::Arc<cortiq_core::CmfModel>>,
155 pub(crate) dyn_force_f32: bool,
157 pub(crate) dyn_skill_layers: Vec<Option<Vec<usize>>>,
162 pub(crate) dyn_active: Option<usize>,
168 pub(crate) dyn_blend_loaded: bool,
172 pub(crate) dyn_phi_layer: Option<usize>,
175 dyn_phi_ema: Vec<f32>,
177 dyn_phi_seen: usize,
178 pub dyn_router: Option<crate::swarm::DynRouter>,
181 o1_cfg: Option<crate::nystrom::O1Cfg>,
184 o1_epoch: u64,
187 o1_flags: Vec<bool>,
189 trace: bool,
192 calib_temp: f32,
195 #[cfg_attr(not(target_os = "macos"), allow(dead_code))]
197 graph_kv_id: u64,
198 #[cfg_attr(not(target_os = "macos"), allow(dead_code))]
201 graph_want_logits: bool,
202 graph_logits: Option<Vec<f32>>,
205 pub embed_multiplier: f32,
207 pub attn_scale: f32,
210 pub swa: Option<(usize, usize)>,
213 pub sliding_layers: Option<Vec<bool>>,
216 pub inv_freq_local: Option<std::sync::Arc<Vec<f32>>>,
219 pub rotary_dim_local: Option<usize>,
220 pub rope_scale: f32,
221 pub rope_scale_local: f32,
222 pub global_attn: Option<(usize, usize)>,
225 pub inv_freq_global: Option<std::sync::Arc<Vec<f32>>>,
228 pub attn_v_norm: bool,
230 pub final_softcap: Option<f32>,
232 pub attn_softcap: f32,
234 confidence_on: bool,
238}
239
240#[cfg(target_os = "macos")]
241impl Drop for Pipeline {
242 fn drop(&mut self) {
243 crate::gpu::kv_mirror_drop(self.graph_kv_id);
244 }
245}
246
247pub struct PipelineWeights {
252 pub embed_tokens: QTensor,
254 pub layers: Vec<LayerWeights>,
256 pub lm_head: QTensor,
258 pub final_norm: Vec<f32>,
260}
261
262pub struct LayerWeights {
264 pub input_norm: Vec<f32>,
265 pub post_norm: Vec<f32>,
268 pub attn_out_norm: Option<Vec<f32>>,
271 pub layer_scale: Option<f32>,
273 pub ffn_out_norm: Option<Vec<f32>>,
276 pub ffn: FfnKind,
277 pub attn: AttnKind,
278}
279
280#[derive(Clone, Copy, PartialEq, Debug, Default)]
283pub enum Act {
284 #[default]
285 Silu,
286 GeluTanh,
287 Situ {
290 beta: f32,
291 linear_beta: f32,
292 },
293}
294
295impl Act {
296 pub fn from_arch(name: &str) -> Self {
297 if name == "gelu_tanh" {
298 Self::GeluTanh
299 } else {
300 Self::Silu
301 }
302 }
303
304 pub fn from_arch_full(arch: &cortiq_core::ModelArch) -> Self {
306 match arch.hidden_act.as_str() {
307 "situ" => Self::Situ {
308 beta: arch.activation_situ_beta.unwrap_or(1.0) as f32,
309 linear_beta: arch.activation_situ_linear_beta.unwrap_or(0.0) as f32,
310 },
311 other => Self::from_arch(other),
312 }
313 }
314
315 #[inline]
316 pub fn apply(self, x: f32) -> f32 {
317 match self {
318 Self::Silu => inference::silu(x),
319 Self::GeluTanh => inference::gelu_tanh(x),
320 Self::Situ { beta, .. } => beta * (x / beta).tanh() * (1.0 / (1.0 + (-x).exp())),
321 }
322 }
323
324 #[inline]
327 pub fn combine(self, g: f32, u: f32) -> f32 {
328 match self {
329 Self::Situ { linear_beta, .. } if linear_beta > 0.0 => {
330 self.apply(g) * (linear_beta * (u / linear_beta).tanh())
331 }
332 _ => self.apply(g) * u,
333 }
334 }
335}
336
337pub struct DenseFfn {
339 pub gate_proj: QTensor,
340 pub up_proj: QTensor,
341 pub down_proj: QTensor,
342 pub act: Act,
344}
345
346pub enum FfnKind {
349 Dense(DenseFfn),
350 Moe(MoeFfn),
354 DenseMoe(Box<DenseMoeFfn>),
361}
362
363pub struct DenseMoeFfn {
365 pub dense: DenseFfn,
366 pub moe: MoeFfn,
367 pub post_norm_1: Vec<f32>,
369 pub pre_norm_2: Vec<f32>,
372 pub post_norm_2: Vec<f32>,
374}
375
376pub struct MoeFfn {
377 pub router: QTensor,
379 pub experts: Vec<DenseFfn>,
380 pub top_k: usize,
381 pub norm_topk_prob: bool,
382 pub router_sigmoid: bool,
385 pub expert_bias: Option<Vec<f32>>,
389 pub routed_scaling: f32,
392 pub route_tau: Option<f32>,
398 pub shared: Option<(DenseFfn, Option<QTensor>)>,
401 pub stats: std::cell::RefCell<Vec<u64>>,
405 pub act_sq: std::cell::RefCell<Vec<f64>>,
412 pub act_rows: std::cell::RefCell<Vec<f32>>,
418 pub mask: Option<Vec<bool>>,
423 pub per_expert_scale: Option<Vec<f32>>,
426 pub router_input_norm: bool,
430}
431
432pub enum AttnKind {
435 Full {
437 wq: QTensor,
438 wk: QTensor,
439 wv: QTensor,
440 wo: QTensor,
441 q_norm: Option<Vec<f32>>,
442 k_norm: Option<Vec<f32>>,
443 output_gate: bool,
444 softplus_gate: Option<(QTensor, bool)>,
448 bias: Option<(Vec<f32>, Vec<f32>, Vec<f32>)>,
450 },
451 Linear(VmfPhaseWeights),
453 LinearGdn(GdnWeights),
455 ShortConv(ShortConvWeights),
458 Mla(Box<MlaWeights>),
466 Kda(Box<crate::linear_core::KdaWeights>),
470}
471
472pub struct MlaWeights {
474 pub q_proj: QTensor,
478 pub q_a: Option<QTensor>,
481 pub q_a_norm: Option<Vec<f32>>,
482 pub kv_a: QTensor,
484 pub kv_a_norm: Vec<f32>,
486 pub kv_b: QTensor,
488 pub o_proj: QTensor,
490 pub nh: usize,
491 pub qk_rope: usize,
492 pub qk_nope: usize,
493 pub v_dim: usize,
494 pub lora: usize,
495 pub scale: f32,
497 pub nope: bool,
499}
500
501pub struct MtpModule {
506 pub enorm: Vec<f32>,
507 pub hnorm: Vec<f32>,
508 pub eh_proj: QTensor,
510 pub layer: LayerWeights,
511 pub final_norm: Vec<f32>,
512 pub kv: crate::kv_cache::LayerKvCache,
513}
514
515pub struct GenerateResult {
517 pub text: String,
518 pub token_ids: Vec<u32>,
519 pub prompt_tokens: usize,
520 pub tokens_generated: usize,
521 pub finish_reason: String,
522 pub mtp_drafted: usize,
524 pub mtp_accepted: usize,
525 pub token_confidence: Vec<f32>,
530 pub traces: Vec<TokenTrace>,
533}
534
535#[derive(Clone, Debug)]
540pub struct TokenTrace {
541 pub t: usize,
543 pub token_id: u32,
545 pub confidence: f32,
547 pub active_skill: Option<String>,
549 pub recon: Option<f32>,
553 pub switched: bool,
556}
557
558fn top1_prob_t(logits: &[f32], id: u32, temp: f32) -> f32 {
563 let t = if temp > 1e-3 { temp } else { 1.0 };
564 let max = logits.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
565 let sum: f32 = logits.iter().map(|&v| ((v - max) / t).exp()).sum();
566 if sum > 0.0 {
567 (((logits[id as usize] - max) / t).exp()) / sum
568 } else {
569 0.0
570 }
571}
572
573fn prefill_batched() -> bool {
576 std::env::var("CMF_PREFILL")
577 .map(|v| v != "seq")
578 .unwrap_or(true)
579}
580
581impl Pipeline {
588 fn can_prefill_batched(&self) -> bool {
589 prefill_batched() && !self.weights.layers.is_empty()
590 }
591}
592
593fn prefill_chunk() -> usize {
598 if let Some(n) = std::env::var("CMF_PREFILL_CHUNK")
599 .ok()
600 .and_then(|v| v.parse::<usize>().ok())
601 {
602 return n.max(1);
603 }
604 if cfg!(target_os = "macos") {
605 512
606 } else if cfg!(target_arch = "aarch64") {
607 256
610 } else {
611 48
612 }
613}
614
615pub type TokenCallback = Box<dyn FnMut(&str) -> bool + Send>;
617
618impl Pipeline {
619 #[inline]
623 pub fn phys_layer(&self, virtual_idx: usize) -> usize {
624 virtual_idx % self.physical_layers
625 }
626
627 #[inline]
630 pub fn is_loop_end(&self, virtual_idx: usize) -> bool {
631 self.loop_final_norm && (virtual_idx + 1) % self.physical_layers == 0
632 }
633
634 #[allow(clippy::too_many_arguments)]
636
637 #[cfg(target_os = "macos")]
656 fn graph_prefill_preferred(&self) -> bool {
657 if !crate::gpu::enabled_here()
658 || !crate::gpu::q1_force()
659 || std::env::var("CMF_GPU_BLOCK")
660 .map(|v| v == "0")
661 .unwrap_or(false)
662 {
663 return false;
664 }
665 self.weights
666 .layers
667 .iter()
668 .any(|lw| matches!(&lw.attn, AttnKind::LinearGdn(w) if w.in_proj_qkv.is_q1()))
669 }
670
671 #[cfg(not(target_os = "macos"))]
672 fn graph_prefill_preferred(&self) -> bool {
673 let graph_on = std::env::var("CMF_GPU_WGPU_GRAPH")
681 .map(|v| v != "0")
682 .unwrap_or_else(|_| {
683 crate::gpu::wgpu_graph_default()
687 });
688 if !graph_on || !crate::gpu::enabled_here() {
689 return false;
690 }
691 self.weights
692 .layers
693 .iter()
694 .any(|lw| matches!(&lw.attn, AttnKind::LinearGdn(_)))
695 }
696
697 #[cfg(target_os = "macos")]
698 fn q1_graph_gpu(
699 &mut self,
700 start: usize,
701 upto: Option<usize>,
702 position: usize,
703 h: &mut [f32],
704 ) -> usize {
705 use crate::gpu::{AttnGpuLayer, GdnGpuCfg, GdnGpuLayer, GraphDims, TokenGraph};
706 if self.attn_softcap > 0.0 || !crate::gpu::enabled_here()
708 || !crate::gpu::q1_force()
709 || std::env::var("CMF_GPU_BLOCK")
710 .map(|v| v == "0")
711 .unwrap_or(false)
712 {
713 return start;
714 }
715 if self.swa.is_some()
720 || self.global_attn.is_some()
721 || self.attention_heads_per_layer.is_some()
722 || self.attn_v_norm
723 || (self.attn_scale - 1.0 / (self.head_dim as f32).sqrt()).abs() > 1e-9
724 || self.weights.layers.iter().any(|lw| {
725 lw.attn_out_norm.is_some()
726 || lw.ffn_out_norm.is_some()
727 || lw.layer_scale.is_some()
728 || matches!(&lw.ffn, FfnKind::Dense(d) if d.act != Act::Silu)
729 })
730 {
731 return start;
732 }
733 let limit = upto
736 .map(|u| u + 1)
737 .unwrap_or(self.num_layers)
738 .min(self.num_layers);
739
740 enum Item<'a> {
741 Gdn {
742 run: Vec<GdnGpuLayer<'a>>,
743 first: usize,
744 },
745 Attn {
746 l: AttnGpuLayer<'a>,
747 li: usize,
748 q_norm: Option<&'a [f32]>,
749 k_norm: Option<&'a [f32]>,
750 output_gate: bool,
751 bias: Option<(&'a [f32], &'a [f32], &'a [f32])>,
752 full_gpu: bool,
755 },
756 }
757
758 let attend_mode = std::env::var("CMF_GPU_ATTEND").unwrap_or_else(|_| "auto".into());
760 let dev_attend = attend_mode != "0"
761 && attend_mode != "off"
762 && (self.head_dim <= 128 || attend_mode == "force" || attend_mode == "256")
766 && self.head_dim % 4 == 0
767 && self.head_dim <= 256
768 && self.rotary_dim >= 2
769 && self.rotary_dim <= self.head_dim
770 && (self.rotary_dim / 2) % 32 == 0
771 && self.num_kv_heads > 0
772 && self.num_heads % self.num_kv_heads == 0;
773
774 let mut plan: Vec<Item> = Vec::new();
775 let mut model_ref: Option<std::sync::Arc<cortiq_core::CmfModel>> = None;
776 let mut scan = start;
777 while scan < limit {
778 let lw = &self.weights.layers[self.phys_layer(scan)];
779 let FfnKind::Dense(d) = &lw.ffn else { break };
780 let (Some(g), Some(u), Some(dn)) = (
781 d.gate_proj.q1_parts(),
782 d.up_proj.q1_parts(),
783 d.down_proj.q1_parts(),
784 ) else {
785 break;
786 };
787 match &lw.attn {
788 AttnKind::LinearGdn(w) if self.gdn_cfg.is_some() => {
789 let parts = (
790 w.in_proj_qkv.q1_parts(),
791 w.in_proj_z.q1_parts(),
792 w.in_proj_a.f32_parts(),
793 w.in_proj_b.f32_parts(),
794 w.out_proj.q1_parts(),
795 );
796 let (Some(qkv), Some(z), Some(a), Some(b), Some(out)) = parts else {
797 break;
798 };
799 if let QTensor::Mapped { model, .. } = &w.in_proj_qkv {
800 model_ref.get_or_insert_with(|| model.clone());
801 }
802 let gl = GdnGpuLayer {
803 attn_norm: &lw.input_norm,
804 post_norm: &lw.post_norm,
805 qkv,
806 z,
807 a,
808 b,
809 out,
810 gate: g,
811 up: u,
812 down: dn,
813 conv1d: &w.conv1d,
814 a_log: &w.a_log,
815 dt_bias: &w.dt_bias,
816 gnorm: &w.norm,
817 };
818 match plan.last_mut() {
819 Some(Item::Gdn { run, .. }) => run.push(gl),
820 _ => plan.push(Item::Gdn {
821 run: vec![gl],
822 first: scan,
823 }),
824 }
825 }
826 AttnKind::Full {
827 wq,
828 wk,
829 wv,
830 wo,
831 q_norm,
832 k_norm,
833 output_gate,
834 softplus_gate: None,
835 bias,
836 } if !self.kv_cache.layers[scan].o1_sealed() => {
837 let parts = (wq.q1_parts(), wk.q1_parts(), wv.q1_parts(), wo.q1_parts());
838 let (Some(pq), Some(pk), Some(pv), Some(po)) = parts else {
839 break;
840 };
841 if let QTensor::Mapped { model, .. } = wq {
842 model_ref.get_or_insert_with(|| model.clone());
843 }
844 let cache = &self.kv_cache.layers[scan];
845 let full_gpu = dev_attend
846 && cache.mode == crate::kv_cache::KvMode::F32
847 && cache.o1.is_none()
848 && bias.is_none()
849 && pq.1 == self.num_heads * self.head_dim * (1 + *output_gate as usize)
850 && pk.1 == self.num_kv_heads * self.head_dim
851 && pv.1 == self.num_kv_heads * self.head_dim
852 && po.2 == self.num_heads * self.head_dim;
853 plan.push(Item::Attn {
854 l: AttnGpuLayer {
855 attn_norm: &lw.input_norm,
856 post_norm: &lw.post_norm,
857 wq: pq,
858 wk: pk,
859 wv: pv,
860 wo: po,
861 gate: g,
862 up: u,
863 down: dn,
864 },
865 li: scan,
866 q_norm: q_norm.as_deref(),
867 k_norm: k_norm.as_deref(),
868 output_gate: *output_gate,
869 bias: bias
870 .as_ref()
871 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
872 full_gpu,
873 });
874 }
875 _ => break,
876 }
877 scan += 1;
878 }
879 let Some(model) = model_ref else { return start };
880 if plan.is_empty() {
881 return start;
882 }
883 let dims = GraphDims {
884 hidden: self.hidden_size,
885 eps: self.rms_eps as f32,
886 gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
887 };
888 let Some(mut graph) = TokenGraph::new(&model, dims, h) else {
889 return start;
890 };
891 let gcfg = self.gdn_cfg.map(|cfg| GdnGpuCfg {
892 nv: cfg.num_v_heads,
893 nk: cfg.num_k_heads,
894 dk: cfg.key_head_dim,
895 dv: cfg.value_head_dim,
896 kk: cfg.conv_kernel,
897 hidden: self.hidden_size,
898 inter: self.intermediate_size,
899 c_dim: cfg.conv_dim(),
900 eps: cfg.rms_eps as f32,
901 gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
902 });
903 let mut valid = 0usize;
907 let mut end = start;
908 for item in &plan {
909 let ok = match item {
910 Item::Gdn { run, .. } => gcfg
911 .as_ref()
912 .map(|gc| run.iter().all(|l| graph.gdn_ok(l, gc)))
913 .unwrap_or(false),
914 Item::Attn { l, .. } => graph.attn_ok(l),
915 };
916 if !ok {
917 break;
918 }
919 valid += 1;
920 end += match item {
921 Item::Gdn { run, .. } => run.len(),
922 Item::Attn { .. } => 1,
923 };
924 }
925 plan.truncate(valid);
926 if plan.is_empty() {
927 return start;
928 }
929
930 let inv_freq = self.inv_freq.clone();
931 let pool = self.pool.clone();
932 let (nh, nkv, hd, hs, rd, eps) = (
933 self.num_heads,
934 self.num_kv_heads,
935 self.head_dim,
936 self.hidden_size,
937 self.rotary_dim,
938 self.rms_eps,
939 );
940 let norm_style = self.norm_style;
941 let gemma = norm_style == cortiq_core::NormStyle::Gemma;
942 let want = self.gdn_cfg.map(|c| c.state_len()).unwrap_or(0);
943 let kv_id = self.graph_kv_id;
944 let mut pending: Vec<(usize, usize)> = Vec::new();
947 let mut dev_attn: Vec<usize> = Vec::new();
950 for item in &plan {
951 if self.loop_final_norm {
953 let item_start = match item {
954 Item::Gdn { first, .. } => *first,
955 Item::Attn { li, .. } => *li,
956 };
957 if item_start > start && self.is_loop_end(item_start - 1) {
958 graph.encode_loop_norm(&self.weights.final_norm);
959 }
960 }
961 match item {
962 Item::Gdn { run, first } => {
963 for l in &mut self.kv_cache.layers[*first..*first + run.len()] {
964 if l.linear_state.len() != want {
965 l.linear_state = vec![0f32; want];
966 }
967 }
968 let ro: Vec<&[f32]> = self.kv_cache.layers[*first..*first + run.len()]
969 .iter()
970 .map(|l| l.linear_state.as_slice())
971 .collect();
972 if !graph.encode_gdn_run(run, &ro, gcfg.as_ref().unwrap()) {
973 tracing::error!("q1 graph: GDN run refused after validation");
975 return start;
976 }
977 graph.commit();
980 pending.push((*first, run.len()));
981 }
982 Item::Attn {
983 l,
984 li,
985 q_norm,
986 k_norm,
987 output_gate,
988 bias,
989 full_gpu,
990 } => {
991 if *full_gpu {
993 let cache = &self.kv_cache.layers[*li];
994 let cpu_k: Vec<&[f32]> = (0..nkv).map(|g| cache.head_keys(g)).collect();
995 let cpu_v: Vec<&[f32]> = (0..nkv).map(|g| cache.head_values(g)).collect();
996 let cpu_stored = cpu_k[0].len() / hd;
997 let p = crate::gpu::AttnDeviceParams {
998 kv_id,
999 layer: *li,
1000 nh,
1001 nkv,
1002 hd,
1003 rd,
1004 position,
1005 eps: eps as f32,
1006 gemma,
1007 output_gate: *output_gate,
1008 q_norm: *q_norm,
1009 k_norm: *k_norm,
1010 inv_freq: &inv_freq,
1011 cpu_k,
1012 cpu_v,
1013 cpu_stored,
1014 };
1015 if graph.attn_device_ok(l, &p) && graph.encode_attn_device(l, &p) {
1016 graph.commit();
1017 dev_attn.push(*li);
1018 continue;
1019 }
1020 }
1022 graph.encode_attn_prefix(l);
1023 graph.sync();
1024 if !pending.is_empty() {
1025 let idxs: Vec<usize> =
1026 pending.drain(..).flat_map(|(f, n)| f..f + n).collect();
1027 let mut outs: Vec<&mut [f32]> = self
1028 .kv_cache
1029 .layers
1030 .iter_mut()
1031 .enumerate()
1032 .filter(|(i, _)| idxs.binary_search(i).is_ok())
1033 .map(|(_, s)| s.linear_state.as_mut_slice())
1034 .collect();
1035 graph.read_states(&mut outs);
1036 }
1037 let mut q_raw = attention::take_buf(l.wq.1);
1038 let mut k = attention::take_buf(l.wk.1);
1039 let mut v = attention::take_buf(l.wv.1);
1040 graph.read_qkv(&mut q_raw, &mut k, &mut v);
1041 let cfg = QwenAttnCfg {
1042 num_heads: nh,
1043 num_kv_heads: nkv,
1044 head_dim: hd,
1045 hidden_size: hs,
1046 position,
1047 inv_freq: &inv_freq,
1048 rotary_dim: rd,
1049 scale: self.attn_scale,
1050 softcap: self.attn_softcap,
1051 window: None,
1052 v_norm: false,
1053 q_norm: *q_norm,
1054 k_norm: *k_norm,
1055 output_gate: *output_gate,
1056 softplus_gate: None,
1057 rope_scale: 1.0,
1058 bias: *bias,
1059 rms_eps: eps,
1060 norm_style,
1061 pool: pool.as_deref(),
1062 };
1063 let mut ao = attention::qwen_attention_core(
1064 q_raw,
1065 k,
1066 v,
1067 &mut self.kv_cache.layers[*li],
1068 &cfg,
1069 );
1070 graph.encode_attn_suffix(l, &ao);
1071 graph.commit();
1074 attention::recycle_buf(&mut ao);
1075 }
1076 }
1077 }
1078 let mut lm_rows = None;
1083 if self.graph_want_logits
1084 && upto.is_none()
1085 && end == self.num_layers
1086 && std::env::var("CMF_GPU_LMHEAD")
1087 .map(|v| v != "0")
1088 .unwrap_or(true)
1089 {
1090 if let Some(lm) = self.weights.lm_head.q1_parts() {
1091 if graph.lm_head_ok(lm) {
1092 graph.encode_lm_head(&self.weights.final_norm, lm);
1093 lm_rows = Some(lm.1);
1094 }
1095 }
1096 }
1097 graph.sync();
1098 if !pending.is_empty() {
1099 let idxs: Vec<usize> = pending.drain(..).flat_map(|(f, n)| f..f + n).collect();
1100 let mut outs: Vec<&mut [f32]> = self
1101 .kv_cache
1102 .layers
1103 .iter_mut()
1104 .enumerate()
1105 .filter(|(i, _)| idxs.binary_search(i).is_ok())
1106 .map(|(_, s)| s.linear_state.as_mut_slice())
1107 .collect();
1108 graph.read_states(&mut outs);
1109 }
1110 if let Some(rows) = lm_rows {
1111 let mut lg = attention::take_buf(rows.min(self.vocab_size));
1112 graph.read_logits(&mut lg);
1113 lg.resize(self.vocab_size, 0.0);
1114 if let Some(c) = self.final_softcap {
1115 for l in lg.iter_mut() {
1116 *l = c * (*l / c).tanh();
1117 }
1118 }
1119 self.graph_logits = Some(lg);
1120 }
1121 graph.finish(h);
1122 for li in dev_attn {
1126 let mut krow = attention::take_buf(nkv * hd);
1127 let mut vrow = attention::take_buf(nkv * hd);
1128 if crate::gpu::kv_mirror_read_last(kv_id, li, nkv, hd, &mut krow, &mut vrow) {
1129 let cache = &mut self.kv_cache.layers[li];
1130 cache.append(&krow, &vrow, &[]);
1131 let n = cache.seq_len;
1132 let mut imp = attention::take_buf(n);
1133 crate::gpu::kv_mirror_take_imp(kv_id, li, &mut imp);
1134 cache.accumulate_imp(&imp);
1135 attention::recycle_buf(&mut imp);
1136 }
1137 attention::recycle_buf(&mut krow);
1138 attention::recycle_buf(&mut vrow);
1139 }
1140 end
1141 }
1142
1143 pub fn new(
1144 tokenizer: Tokenizer,
1145 weights: PipelineWeights,
1146 hidden_size: usize,
1147 intermediate_size: usize,
1148 num_heads: usize,
1149 num_kv_heads: usize,
1150 head_dim: usize,
1151 num_layers: usize,
1152 physical_layers: usize,
1153 loop_final_norm: bool,
1154 vocab_size: usize,
1155 rms_eps: f64,
1156 rope_base: f32,
1157 norm_style: NormStyle,
1158 max_seq_len: usize,
1159 sampler_config: SamplerConfig,
1160 ) -> Self {
1161 let rng = match sampler_config.seed {
1162 Some(s) => SplitMix64::new(s),
1163 None => SplitMix64::from_entropy(),
1164 };
1165 let inv_freq = std::sync::Arc::new(attention::rope_inv_freq(head_dim, rope_base));
1166 let pool = Pool::from_env();
1167 if let Some(p) = &pool {
1168 tracing::info!("worker pool: {} threads", p.n_workers());
1169 }
1170 Self {
1171 tokenizer: std::sync::Arc::new(tokenizer),
1172 kv_cache: KvCache::new(num_layers, num_kv_heads, head_dim, max_seq_len),
1173 sampler_config,
1174 weights,
1175 hidden_size,
1176 intermediate_size,
1177 num_heads,
1178 num_kv_heads,
1179 head_dim,
1180 num_layers,
1181 physical_layers,
1182 loop_final_norm,
1183 vocab_size,
1184 rms_eps,
1185 rope_base,
1186 norm_style,
1187 rotary_dim: head_dim,
1188 attention_heads_per_layer: None,
1189 vmf_cfg: None,
1190 gdn_cfg: None,
1191 kda_cfg: None,
1192 g3n: None,
1193 dsv4: None,
1194 dsv4_mtp: Vec::new(),
1195 dspark: None,
1196 dspark_pending: Vec::new(),
1197 dspark_hist: Vec::new(),
1198 dspark_real: Vec::new(),
1199 dspark_trunk_picks: Vec::new(),
1200 dspark_exp: Vec::new(),
1201 dspark_draft_ns: 0,
1202 logit_multiplier: None,
1203 cancel: std::sync::Arc::new(std::sync::atomic::AtomicBool::new(false)),
1204 kv_history: Vec::new(),
1205 short_conv_cfg: None,
1206 mtp: None,
1207 speculative: std::env::var("CMF_MTP").map(|v| v != "0").unwrap_or(true),
1208 rng,
1209 sampler_scratch: SamplerScratch::default(),
1210 inv_freq,
1211 ws: ForwardScratch::new(hidden_size),
1212 pool,
1213 model: None,
1214 dyn_force_f32: false,
1215 dyn_skill_layers: Vec::new(),
1216 dyn_active: None,
1217 dyn_blend_loaded: false,
1218 dyn_phi_layer: None,
1219 dyn_phi_ema: Vec::new(),
1220 dyn_phi_seen: 0,
1221 dyn_router: None,
1222 o1_cfg: None,
1223 o1_epoch: 0,
1224 o1_flags: Vec::new(),
1225 trace: false,
1226 calib_temp: 1.0,
1227 confidence_on: true,
1228 embed_multiplier: 1.0,
1229 attn_scale: 1.0 / (head_dim as f32).sqrt(),
1230 swa: None,
1231 sliding_layers: None,
1232 inv_freq_local: None,
1233 rotary_dim_local: None,
1234 rope_scale: 1.0,
1235 rope_scale_local: 1.0,
1236 global_attn: None,
1237 inv_freq_global: None,
1238 attn_v_norm: false,
1239 final_softcap: None,
1240 attn_softcap: 0.0,
1241 graph_want_logits: false,
1242 graph_logits: None,
1243 graph_kv_id: {
1244 static NEXT: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(1);
1245 NEXT.fetch_add(1, std::sync::atomic::Ordering::Relaxed)
1246 },
1247 }
1248 }
1249
1250 pub fn set_o1(&mut self, cfg: Option<crate::nystrom::O1Cfg>) {
1257 self.o1_flags = match &cfg {
1258 Some(c) => {
1259 let mut flags = c.layer_flags(self.num_layers);
1260 for (li, f) in flags.iter_mut().enumerate() {
1261 if *f
1262 && !matches!(
1263 self.weights.layers[self.phys_layer(li)].attn,
1264 AttnKind::Full { .. }
1265 )
1266 {
1267 *f = false;
1268 }
1269 }
1270 flags
1271 }
1272 None => Vec::new(),
1273 };
1274 if let Some(c) = &cfg {
1275 let n = self.o1_flags.iter().filter(|&&f| f).count();
1276 tracing::info!(
1277 "o1 nystrom attention: {n}/{} layer(s), m={} w={} sink={} rect={:?}",
1278 self.num_layers,
1279 c.m,
1280 c.w,
1281 c.sink,
1282 c.rect
1283 );
1284 }
1285 self.o1_cfg = cfg;
1286 }
1287
1288 pub fn o1_active(&self) -> bool {
1290 self.o1_cfg.is_some() && self.o1_flags.iter().any(|&f| f)
1291 }
1292
1293 fn o1_begin(&mut self) {
1295 if let Some(c) = &self.o1_cfg {
1296 let (m, w, sink, rect) = (c.m, c.w, c.sink, c.rect);
1297 for (li, &f) in self.o1_flags.iter().enumerate() {
1298 if f {
1299 self.kv_cache.layers[li].o1_begin(m, w, sink, rect);
1300 }
1301 }
1302 }
1303 }
1304
1305 fn o1_seal(&mut self) {
1308 self.o1_epoch = self.o1_epoch.wrapping_add(1);
1309 if self.o1_cfg.is_none() {
1310 return;
1311 }
1312 for li in 0..self.num_layers {
1313 if self.o1_flags.get(li).copied().unwrap_or(false) {
1314 self.kv_cache.layers[li].o1_seal(self.num_heads);
1315 }
1316 }
1317 }
1318
1319 pub fn set_trace(&mut self, on: bool) {
1321 self.trace = on;
1322 }
1323
1324 pub fn set_sampler_config(&mut self, config: SamplerConfig) {
1327 self.rng = match config.seed {
1328 Some(seed) => SplitMix64::new(seed),
1329 None => SplitMix64::from_entropy(),
1330 };
1331 self.sampler_config = config;
1332 }
1333
1334 pub fn set_confidence(&mut self, on: bool) {
1339 self.confidence_on = on;
1340 }
1341
1342 pub fn set_calib_temp(&mut self, t: f32) {
1345 self.calib_temp = if t > 1e-3 { t } else { 1.0 };
1346 }
1347
1348 pub fn calib_temp(&self) -> f32 {
1350 self.calib_temp
1351 }
1352
1353 pub fn set_rotary(&mut self, rotary_dim: usize, base: f32) {
1356 self.rotary_dim = rotary_dim.min(self.head_dim);
1357 self.inv_freq = std::sync::Arc::new(attention::rope_inv_freq(self.rotary_dim, base));
1358 }
1359
1360 fn attn_cfg(&self, position: usize) -> QwenAttnCfg<'_> {
1361 QwenAttnCfg {
1362 num_heads: self.num_heads,
1363 num_kv_heads: self.num_kv_heads,
1364 head_dim: self.head_dim,
1365 hidden_size: self.hidden_size,
1366 position,
1367 inv_freq: &self.inv_freq,
1368 rotary_dim: self.rotary_dim,
1369 scale: self.attn_scale,
1370 softcap: self.attn_softcap,
1371 window: None,
1372 v_norm: false,
1373 q_norm: None,
1374 k_norm: None,
1375 output_gate: false,
1376 softplus_gate: None,
1377 rope_scale: self.rope_scale,
1378 bias: None,
1379 rms_eps: self.rms_eps,
1380 norm_style: self.norm_style,
1381 pool: self.pool.as_deref(),
1382 }
1383 }
1384
1385 pub fn generate(
1387 &mut self,
1388 prompt: &str,
1389 max_tokens: usize,
1390 task_mask: Option<&TaskMask>,
1391 on_token: Option<TokenCallback>,
1392 ) -> Result<GenerateResult, String> {
1393 let input_ids = self.tokenizer.with_bos(self.tokenizer.encode(prompt));
1394 self.generate_from_ids(&input_ids, max_tokens, task_mask, on_token)
1395 }
1396
1397 pub fn generate_from_ids(
1405 &mut self,
1406 input_ids: &[u32],
1407 max_tokens: usize,
1408 task_mask: Option<&TaskMask>,
1409 mut on_token: Option<TokenCallback>,
1410 ) -> Result<GenerateResult, String> {
1411 if std::env::var("CMF_TRACE_H").is_ok() {
1412 eprintln!("input_ids: {input_ids:?}");
1413 }
1414 if input_ids.is_empty() {
1415 return Err("empty prompt: nothing to generate from".to_string());
1416 }
1417
1418 let reuse_from = {
1426 let on = !std::env::var("CMF_KV_REUSE").is_ok_and(|v| v == "0");
1427 let h = &self.kv_history;
1428 if on
1429 && task_mask.is_none()
1430 && self.mtp.is_none()
1431 && self.o1_cfg.is_none()
1432 && !h.is_empty()
1433 && h.len() < input_ids.len()
1434 && input_ids[..h.len()] == h[..]
1435 {
1436 h.len()
1437 } else {
1438 0
1439 }
1440 };
1441 if reuse_from == 0 {
1442 self.kv_cache.clear();
1444 self.kv_history.clear();
1445 crate::gpu::graph_kv_reset(self.graph_kv_id);
1446 } else if std::env::var("CMF_PREFILL_PROF").is_ok() {
1447 eprintln!(
1448 "kv-reuse: {} of {} prompt positions already cached",
1449 reuse_from,
1450 input_ids.len()
1451 );
1452 }
1453 crate::gpu::graph_race_begin_generation();
1454 self.o1_begin();
1455
1456 let graph_on = std::env::var("CMF_GPU_WGPU_GRAPH")
1462 .map(|v| v != "0")
1463 .unwrap_or_else(|_| {
1464 crate::gpu::wgpu_graph_default()
1468 });
1469 let graph_spec = self.speculative
1476 && graph_on
1477 && self.mtp.is_some()
1478 && task_mask.is_none()
1479 && !self.o1_active()
1480 && self.sampler_config.temperature < 1e-6
1481 && self.sampler_config.repetition_penalty == 1.0
1482 && std::env::var("CMF_GRAPH_SPEC").is_ok_and(|v| v != "0");
1483 let spec_active = self.speculative
1484 && self.mtp.is_some()
1485 && task_mask.is_none()
1486 && !self.o1_active()
1487 && (!graph_on || graph_spec)
1488 && self.sampler_config.temperature < 1e-6;
1489 let mut mtp = if spec_active { self.mtp.take() } else { None };
1492 if std::env::var("CMF_MTP_CHAIN_PROBE").is_ok() {
1493 eprintln!(
1494 "mtp-probe gate: spec_active={spec_active} mtp={} speculative={} graph_on={graph_on} temp_ok={}",
1495 mtp.is_some(),
1496 self.speculative,
1497 self.sampler_config.temperature < 1e-6,
1498 );
1499 }
1500 if let Some(m) = &mut mtp {
1501 m.kv.clear();
1502 }
1503 let mut router = if mtp.is_none() {
1507 self.dyn_router.take()
1508 } else {
1509 None
1510 };
1511 if let Some(r) = &mut router {
1512 r.reset(); self.dyn_phi_seen = 0; let _ = self.set_active_skill(None);
1515 }
1516
1517 let mut all_ids = input_ids.to_vec();
1518 let mut generated = 0usize;
1519 let mut finish_reason = "max_tokens".to_string();
1520 let mut drafted = 0usize;
1521 let mut accepted = 0usize;
1522 let mut confidence: Vec<f32> = Vec::new();
1523 let trace_on = self.trace;
1524 let calib_temp = self.calib_temp;
1525 let mut traces: Vec<TokenTrace> = Vec::new();
1526
1527 let mut hidden = vec![0.0f32; self.hidden_size];
1533 let mut pos = reuse_from;
1534 let fuse_lm = mtp.is_none()
1543 && router.is_none()
1544 && std::env::var("CMF_GPU_LMHEAD").as_deref() != Ok("0");
1545 self.graph_logits = None;
1546 self.graph_want_logits = false;
1547 let _tpf = std::time::Instant::now();
1548 let batch_k = std::env::var("CMF_BATCH_K")
1549 .ok()
1550 .and_then(|v| v.parse::<usize>().ok())
1551 .unwrap_or(0);
1552 while self.dsv4.is_some()
1563 && mtp.is_none()
1564 && pos < input_ids.len()
1565 && !self.cancel.load(std::sync::atomic::Ordering::Relaxed)
1566 {
1567 let end = (pos + prefill_chunk()).min(input_ids.len());
1568 let ids: Vec<u32> = input_ids[pos..end].to_vec();
1569 let mut lg = Vec::new();
1570 if let Some(b) = &mut self.dsv4 {
1571 let (g, layers, cfg, st) = (&b.0, &b.1, b.2, &mut b.3);
1572 crate::dsv4::forward_chunk(
1573 g,
1574 layers,
1575 &cfg,
1576 st,
1577 &ids,
1578 pos,
1579 &self.inv_freq,
1580 self.pool.as_deref(),
1581 &mut lg,
1582 end == input_ids.len(),
1583 );
1584 }
1585 if end == input_ids.len() {
1586 self.graph_logits = Some(lg);
1587 }
1588 pos = end;
1589 hidden = vec![0.0; self.hidden_size];
1590 }
1591 let dyn_prefill = router.is_some();
1596 let graph_prefill = self.graph_prefill_preferred();
1602 if task_mask.is_none()
1603 && !dyn_prefill
1604 && !graph_prefill
1605 && self.can_prefill_batched()
1606 && self.g3n.is_none()
1607 && input_ids.len() > 2
1608 {
1609 let chunk = prefill_chunk();
1615 let hs = self.hidden_size;
1616 while pos < input_ids.len() && !self.cancel.load(std::sync::atomic::Ordering::Relaxed) {
1617 let end = (pos + chunk).min(input_ids.len());
1618 let hb = self.prefill_batch(&input_ids[pos..end], pos);
1619 if let Some(m) = &mut mtp {
1620 let probe: usize = std::env::var("CMF_MTP_CHAIN_PROBE")
1621 .ok()
1622 .and_then(|v| v.parse().ok())
1623 .unwrap_or(0);
1624 for p in pos..end {
1625 if p + 1 < input_ids.len() {
1626 if probe >= 1 && p + 2 < input_ids.len() {
1627 let (d1, mut hx) = self.mtp_step_h(
1631 m,
1632 &hb[(p - pos) * hs..(p - pos + 1) * hs],
1633 input_ids[p + 1],
1634 p,
1635 );
1636 let mut ok = d1 == input_ids[p + 2];
1637 Self::chain_probe_note(0, ok);
1638 let mut d_prev = d1;
1639 let mut extra = 0usize;
1640 for j in 1..probe {
1641 if p + 2 + j >= input_ids.len() {
1642 break;
1643 }
1644 let (dj, hj) =
1645 self.mtp_step_h(m, &hx, d_prev, p + 1 + j);
1646 extra += 1;
1647 ok = ok && dj == input_ids[p + 2 + j];
1648 Self::chain_probe_note(j, ok);
1649 d_prev = dj;
1650 hx = hj;
1651 }
1652 m.kv.truncate_last(extra);
1653 } else {
1654 let _ = self.mtp_step(
1655 m,
1656 &hb[(p - pos) * hs..(p - pos + 1) * hs],
1657 input_ids[p + 1],
1658 p,
1659 );
1660 }
1661 }
1662 }
1663 }
1664 hidden.copy_from_slice(&hb[(end - pos - 1) * hs..]);
1665 pos = end;
1666 }
1667 }
1668 let pair_off = std::env::var("CMF_PAIR").is_ok_and(|v| v == "0");
1669 if task_mask.is_none()
1670 && !dyn_prefill
1671 && !graph_prefill
1672 && !pair_off
1673 && self.pair_supported()
1674 {
1675 while pos + 1 < input_ids.len()
1676 && !self.cancel.load(std::sync::atomic::Ordering::Relaxed)
1677 {
1678 let e1 = self.embed_single(input_ids[pos]);
1679 let e2 = self.embed_single(input_ids[pos + 1]);
1680 let (h1, h2) = self.forward_pair(&e1, &e2, pos);
1681 self.commit_linear_scratch();
1683 if let Some(m) = &mut mtp {
1684 let _ = self.mtp_step(m, &h1, input_ids[pos + 1], pos);
1685 if pos + 2 < input_ids.len() {
1686 let probe: usize = std::env::var("CMF_MTP_CHAIN_PROBE")
1687 .ok()
1688 .and_then(|v| v.parse().ok())
1689 .unwrap_or(0);
1690 if probe >= 1 && pos + 3 < input_ids.len() {
1691 let (d1, mut hx) =
1695 self.mtp_step_h(m, &h2, input_ids[pos + 2], pos + 1);
1696 let mut ok = d1 == input_ids[pos + 3];
1697 Self::chain_probe_note(0, ok);
1698 let mut d_prev = d1;
1699 let mut extra = 0usize;
1700 for j in 1..probe {
1701 if pos + 3 + j >= input_ids.len() {
1702 break;
1703 }
1704 let (dj, hj) =
1705 self.mtp_step_h(m, &hx, d_prev, pos + 2 + j);
1706 extra += 1;
1707 ok = ok && dj == input_ids[pos + 3 + j];
1708 Self::chain_probe_note(j, ok);
1709 d_prev = dj;
1710 hx = hj;
1711 }
1712 m.kv.truncate_last(extra);
1713 } else {
1714 let _ = self.mtp_step(m, &h2, input_ids[pos + 2], pos + 1);
1715 }
1716 }
1717 }
1718 hidden = h2;
1719 pos += 2;
1720 }
1721 }
1722 if batch_k > 0
1731 && graph_prefill
1732 && task_mask.is_none()
1733 && !self.o1_active()
1734 && mtp.is_none()
1735 && !dyn_prefill
1736 && pos + 1 < input_ids.len()
1737 {
1738 let hs = self.hidden_size;
1739 let chunk = batch_k;
1740 while pos < input_ids.len() {
1741 let end = (pos + chunk).min(input_ids.len());
1742 let bk = end - pos;
1743 let mut hiddens = vec![0f32; bk * hs];
1744 for (j, &id) in input_ids[pos..end].iter().enumerate() {
1745 hiddens[j * hs..(j + 1) * hs].copy_from_slice(&self.embed_single(id));
1746 }
1747 let positions: Vec<usize> = (pos..end).collect();
1748 let t_chunk = std::time::Instant::now();
1749 let ok_b = self.try_batch_graph_wgpu(&mut hiddens, &positions, bk, None);
1750 if std::env::var("CMF_GRAPH_PROF").is_ok() {
1751 let ms = t_chunk.elapsed().as_secs_f64() * 1000.0;
1752 eprintln!(
1753 "batch-chunk: k={bk} ok={ok_b} {ms:.1} ms ({:.1} tok/s)",
1754 bk as f64 / (ms / 1000.0)
1755 );
1756 }
1757 {
1758 use std::sync::atomic::{AtomicBool, Ordering};
1759 static SAID: AtomicBool = AtomicBool::new(false);
1760 if !SAID.swap(true, Ordering::Relaxed) {
1761 if ok_b {
1762 tracing::info!("batched prefill: ACTIVE (k={bk})");
1763 } else {
1764 tracing::warn!("batched prefill declined — per-position graph");
1765 }
1766 }
1767 }
1768 if ok_b {
1769 hidden.copy_from_slice(&hiddens[(bk - 1) * hs..]);
1770 pos = end;
1771 } else {
1772 break; }
1774 }
1775 }
1776 while pos < input_ids.len() && !self.cancel.load(std::sync::atomic::Ordering::Relaxed) {
1777 self.graph_want_logits = fuse_lm && pos + 1 == input_ids.len();
1778 hidden = self.forward_layers(&self.embed_single(input_ids[pos]), pos, task_mask);
1779 if let Some(m) = &mut mtp {
1780 if pos + 1 < input_ids.len() {
1781 let probe: usize = std::env::var("CMF_MTP_CHAIN_PROBE")
1787 .ok()
1788 .and_then(|v| v.parse().ok())
1789 .unwrap_or(0);
1790 if probe >= 1 && pos + 2 < input_ids.len() {
1791 let (d1, mut hx) = self.mtp_step_h(m, &hidden, input_ids[pos + 1], pos);
1792 let mut ok = d1 == input_ids[pos + 2];
1793 Self::chain_probe_note(0, ok);
1794 let mut d_prev = d1;
1795 let mut extra = 0usize;
1796 for j in 1..probe {
1797 if pos + 2 + j >= input_ids.len() {
1798 break;
1799 }
1800 let (dj, hj) =
1801 self.mtp_step_h(m, &hx, d_prev, pos + 1 + j);
1802 extra += 1;
1803 ok = ok && dj == input_ids[pos + 2 + j];
1804 Self::chain_probe_note(j, ok);
1805 d_prev = dj;
1806 hx = hj;
1807 }
1808 m.kv.truncate_last(extra);
1811 } else {
1812 let _ = self.mtp_step(m, &hidden, input_ids[pos + 1], pos);
1813 }
1814 }
1815 }
1816 pos += 1;
1817 }
1818 if std::env::var("CMF_PREFILL_PROF").is_ok() {
1819 eprintln!(
1820 "prefill: {} tokens in {:.1} ms (batch_k={batch_k})",
1821 input_ids.len(),
1822 _tpf.elapsed().as_secs_f64() * 1000.0
1823 );
1824 }
1825 if self
1828 .cancel
1829 .swap(false, std::sync::atomic::Ordering::Relaxed)
1830 {
1831 self.kv_history.clear();
1832 if let Some(m) = mtp {
1833 self.mtp = Some(m);
1834 }
1835 return Ok(GenerateResult {
1836 text: String::new(),
1837 token_ids: Vec::new(),
1838 prompt_tokens: input_ids.len(),
1839 tokens_generated: 0,
1840 finish_reason: "cancelled".to_string(),
1841 mtp_drafted: 0,
1842 mtp_accepted: 0,
1843 token_confidence: Vec::new(),
1844 traces: Vec::new(),
1845 });
1846 }
1847
1848 self.o1_seal();
1851
1852 macro_rules! commit {
1854 ($id:expr) => {{
1855 all_ids.push($id);
1856 generated += 1;
1857 if self.tokenizer.is_eos($id) {
1858 finish_reason = "stop".to_string();
1859 false
1860 } else {
1861 let token_text = self.tokenizer.decode_token($id);
1862 let mut go = true;
1863 if let Some(ref mut cb) = on_token {
1864 if !cb(&token_text) {
1865 finish_reason = "cancelled".to_string();
1866 go = false;
1867 }
1868 }
1869 go
1870 }
1871 }};
1872 }
1873
1874 let mut next_pos = input_ids.len();
1876 'decode: while generated < max_tokens {
1877 if self
1878 .cancel
1879 .swap(false, std::sync::atomic::Ordering::Relaxed)
1880 {
1881 finish_reason = "cancelled".to_string();
1882 break 'decode;
1883 }
1884 let mut logits = match self.graph_logits.take() {
1885 Some(lg) => lg,
1886 None => {
1887 inference::rms_norm_into(
1888 &hidden,
1889 &self.weights.final_norm,
1890 self.rms_eps,
1891 self.norm_style,
1892 &mut self.ws.n1,
1893 );
1894 self.lm_head_forward(&self.ws.n1)
1895 }
1896 };
1897 let t_next = sampler::sample_with_scratch(
1898 &logits,
1899 &self.sampler_config,
1900 &all_ids,
1901 &mut self.rng,
1902 &mut self.sampler_scratch,
1903 );
1904 if self.confidence_on {
1905 confidence.push(top1_prob_t(&logits, t_next, calib_temp));
1906 }
1907 attention::recycle_buf(&mut logits);
1908 if trace_on {
1909 let skill = router.as_ref().and_then(|r| r.active_id());
1913 traces.push(TokenTrace {
1914 t: generated,
1915 token_id: t_next,
1916 confidence: confidence.last().copied().unwrap_or(0.0),
1917 active_skill: skill,
1918 recon: None,
1919 switched: false,
1920 });
1921 }
1922 if !commit!(t_next) {
1923 break 'decode;
1924 }
1925 if generated >= max_tokens {
1926 break 'decode;
1927 }
1928
1929 if self.kv_cache.needs_eviction() {
1930 let keep = (self.kv_cache.max_seq_len / 2).max(1);
1931 self.kv_cache.evict(keep);
1932 }
1933
1934 match &mut mtp {
1935 #[cfg(feature = "gpu")]
1937 Some(m) if graph_spec && generated + 1 < max_tokens && next_pos > 0 => {
1938 if let Some((extra, n_pos, new_h)) = self.graph_spec_step(
1939 m,
1940 &hidden,
1941 t_next,
1942 next_pos,
1943 &mut drafted,
1944 &mut accepted,
1945 ) {
1946 next_pos = n_pos;
1947 hidden = new_h;
1948 let mut stopped = false;
1949 for &id in &extra {
1950 if self.confidence_on {
1951 confidence.push(0.0);
1952 }
1953 if !commit!(id) {
1954 stopped = true;
1955 break;
1956 }
1957 }
1958 if stopped {
1959 break 'decode;
1960 }
1961 continue 'decode;
1962 }
1963 hidden = self.forward_layers(&self.embed_single(t_next), next_pos, task_mask);
1965 next_pos += 1;
1966 continue 'decode;
1967 }
1968 Some(m) if !graph_spec && generated + 1 < max_tokens => {
1970 let draft = self.mtp_step(m, &hidden, t_next, next_pos - 1);
1971 drafted += 1;
1972 let emb1 = self.embed_single(t_next);
1973 let emb2 = self.embed_single(draft);
1974 let (h1, h2) = self.forward_pair(&emb1, &emb2, next_pos);
1975
1976 inference::rms_norm_into(
1977 &h1,
1978 &self.weights.final_norm,
1979 self.rms_eps,
1980 self.norm_style,
1981 &mut self.ws.n1,
1982 );
1983 let mut logits1 = self.lm_head_forward(&self.ws.n1);
1984 let t_after = sampler::sample_with_scratch(
1985 &logits1,
1986 &self.sampler_config,
1987 &all_ids,
1988 &mut self.rng,
1989 &mut self.sampler_scratch,
1990 );
1991 if self.confidence_on {
1992 confidence.push(top1_prob_t(&logits1, t_after, calib_temp));
1993 }
1994 attention::recycle_buf(&mut logits1);
1995 if trace_on {
1996 traces.push(TokenTrace {
1999 t: generated,
2000 token_id: t_after,
2001 confidence: confidence.last().copied().unwrap_or(0.0),
2002 active_skill: None,
2003 recon: None,
2004 switched: false,
2005 });
2006 }
2007 let stop = !commit!(t_after);
2008
2009 if t_after == draft {
2010 accepted += 1;
2011 self.commit_linear_scratch();
2012 let _ = self.mtp_step(m, &h1, t_after, next_pos);
2013 hidden = h2;
2014 next_pos += 2;
2015 } else {
2016 for layer in &mut self.kv_cache.layers {
2018 layer.truncate_last(1);
2019 }
2020 if !stop {
2021 let _ = self.mtp_step(m, &h1, t_after, next_pos);
2022 hidden = self.forward_layers(
2023 &self.embed_single(t_after),
2024 next_pos + 1,
2025 None,
2026 );
2027 }
2028 next_pos += 2;
2029 }
2030 if stop {
2031 break 'decode;
2032 }
2033 }
2034 _ => {
2036 #[cfg(feature = "gpu")]
2041 if Self::dsv4_spec_on() && self.dsv4.is_some() {
2042 static SAID: std::sync::Once = std::sync::Once::new();
2043 SAID.call_once(|| {
2044 eprintln!(
2045 "dsv4-spec гейт: mtp={} mask={} router={} trace={} temp={} rep={} ",
2046 !self.dsv4_mtp.is_empty(),
2047 task_mask.is_none(),
2048 router.is_none(),
2049 !trace_on,
2050 self.sampler_config.temperature < 1e-6,
2051 self.sampler_config.repetition_penalty == 1.0,
2052 );
2053 });
2054 }
2055 #[cfg(feature = "gpu")]
2056 if Self::dsv4_spec_on()
2057 && self.dsv4.is_some()
2058 && !self.dsv4_mtp.is_empty()
2059 && task_mask.is_none()
2060 && router.is_none()
2061 && !trace_on
2062 && self.sampler_config.temperature < 1e-6
2063 && self.sampler_config.repetition_penalty == 1.0
2064 && generated + 1 < max_tokens
2065 && all_ids.len() >= 2
2066 {
2067 let tip_token = all_ids[all_ids.len() - 2];
2068 if let Some((extra, n_pos)) = self.dsv4_spec_step(
2069 tip_token,
2070 t_next,
2071 next_pos,
2072 &mut drafted,
2073 &mut accepted,
2074 ) {
2075 next_pos = n_pos;
2076 let mut stopped = false;
2077 for &id in &extra {
2078 if self.confidence_on {
2079 confidence.push(0.0);
2080 }
2081 if !commit!(id) {
2082 stopped = true;
2083 break;
2084 }
2085 }
2086 if stopped {
2087 break 'decode;
2088 }
2089 continue 'decode;
2090 }
2091 }
2092 self.graph_want_logits = fuse_lm;
2093 let mut t_fwd = t_next;
2099 let pure_greedy = self.sampler_config.temperature < 1e-6
2100 && self.sampler_config.repetition_penalty == 1.0
2101 && self.sampler_config.suppress_tokens.is_empty();
2102 let burst_k = std::env::var("CMF_MULTISTEP")
2107 .ok()
2108 .and_then(|v| v.parse::<usize>().ok())
2109 .unwrap_or(0);
2110 if pure_greedy
2111 && burst_k >= 1
2112 && fuse_lm
2113 && task_mask.is_none()
2114 && router.is_none()
2115 && !trace_on
2116 && !self.confidence_on
2117 {
2118 let mut stopped = false;
2119 loop {
2120 let room = max_tokens.saturating_sub(generated);
2121 if room <= 2 {
2122 break;
2123 }
2124 let k = burst_k.min(room - 1);
2125 if k < 1 {
2126 break;
2127 }
2128 let Some(ids) = self.try_multi_burst(t_fwd, next_pos, k) else {
2129 break;
2130 };
2131 next_pos += k;
2132 for &id in &ids {
2133 if !commit!(id) {
2134 stopped = true;
2135 break;
2136 }
2137 }
2138 if stopped {
2139 break;
2140 }
2141 t_fwd = *ids.last().unwrap();
2142 }
2143 if stopped {
2144 break 'decode;
2145 }
2146 }
2147 hidden = self.forward_layers(&self.embed_single(t_fwd), next_pos, task_mask);
2148 next_pos += 1;
2149 if let Some(r) = &mut router {
2152 let phi = self.dyn_phi_ema.clone();
2153 let decision = r.step(&phi, generated);
2154 if let Some(new_active) = decision {
2155 let _ = self.set_active_skill(new_active);
2156 }
2157 if trace_on {
2160 if let Some(last) = traces.last_mut() {
2161 let e = r.last_best_e();
2162 last.recon = e.is_finite().then_some(e);
2163 last.switched = decision.is_some();
2164 }
2165 }
2166 }
2167 }
2168 }
2169 }
2170
2171 self.graph_want_logits = false;
2172 self.graph_logits = None;
2173 if router.is_some() {
2175 let _ = self.set_active_skill(None);
2176 }
2177 self.dyn_router = router.or(self.dyn_router.take());
2178 self.mtp = mtp.or(self.mtp.take());
2179
2180 let output_ids = &all_ids[input_ids.len()..];
2181 let forwarded = input_ids.len() + output_ids.len().saturating_sub(1);
2185 self.kv_history = all_ids[..forwarded.min(all_ids.len())].to_vec();
2186 confidence.truncate(output_ids.len()); traces.truncate(output_ids.len());
2188 Ok(GenerateResult {
2189 text: self.tokenizer.decode(output_ids),
2190 token_ids: output_ids.to_vec(),
2191 prompt_tokens: input_ids.len(),
2192 tokens_generated: generated,
2193 finish_reason,
2194 mtp_drafted: drafted,
2195 mtp_accepted: accepted,
2196 token_confidence: confidence,
2197 traces,
2198 })
2199 }
2200
2201 fn mtp_step(
2205 &mut self,
2206 m: &mut MtpModule,
2207 hidden: &[f32],
2208 next_token: u32,
2209 position: usize,
2210 ) -> u32 {
2211 self.mtp_step_h(m, hidden, next_token, position).0
2212 }
2213
2214 fn chain_probe_note(depth: usize, prefix_ok: bool) {
2218 use std::sync::Mutex;
2219 static T: Mutex<Vec<(u64, u64)>> = Mutex::new(Vec::new());
2220 let mut t = T.lock().unwrap();
2221 if t.len() <= depth {
2222 t.resize(depth + 1, (0, 0));
2223 }
2224 t[depth].0 += 1;
2225 t[depth].1 += prefix_ok as u64;
2226 if depth == 0 && t[0].0 % 128 == 0 {
2227 let line: Vec<String> = t
2228 .iter()
2229 .enumerate()
2230 .map(|(d, (n, k))| format!("d{}={:.0}%({n})", d + 1, 100.0 * *k as f64 / (*n).max(1) as f64))
2231 .collect();
2232 eprintln!("mtp-chain: {}", line.join(" "));
2233 }
2234 }
2235
2236 fn mtp_step_h(
2240 &mut self,
2241 m: &mut MtpModule,
2242 hidden: &[f32],
2243 next_token: u32,
2244 position: usize,
2245 ) -> (u32, Vec<f32>) {
2246 let e = self.embed_single(next_token);
2250 let mut cat = vec![0.0f32; 2 * self.hidden_size];
2251 let (cat_e, cat_h) = cat.split_at_mut(self.hidden_size);
2252 inference::rms_norm_into(&e, &m.enorm, self.rms_eps, self.norm_style, cat_e);
2253 inference::rms_norm_into(hidden, &m.hnorm, self.rms_eps, self.norm_style, cat_h);
2254 let mut x = vec![0.0f32; self.hidden_size];
2255 m.eh_proj.matvec(&cat, &mut x, self.pool.as_deref());
2256
2257 let lw = &m.layer;
2259 inference::rms_norm_into(
2260 &x,
2261 &lw.input_norm,
2262 self.rms_eps,
2263 self.norm_style,
2264 &mut self.ws.n1,
2265 );
2266 let attn = match &lw.attn {
2267 AttnKind::Mla(_) => unreachable!("MLA has no MTP/pair path"),
2269 AttnKind::Kda(_) => unreachable!("KDA has no MTP/pair path"),
2270 AttnKind::Full {
2271 wq,
2272 wk,
2273 wv,
2274 wo,
2275 q_norm,
2276 k_norm,
2277 output_gate,
2278 softplus_gate,
2279 bias,
2280 } => {
2281 let mut cfg = self.attn_cfg(position);
2282 cfg.q_norm = q_norm.as_deref();
2283 cfg.k_norm = k_norm.as_deref();
2284 cfg.output_gate = *output_gate;
2285 cfg.softplus_gate = softplus_gate
2286 .as_ref()
2287 .map(|(gate, per_head)| (gate, *per_head));
2288 cfg.bias = bias
2289 .as_ref()
2290 .map(|(q, k, v)| (q.as_slice(), k.as_slice(), v.as_slice()));
2291 attention::qwen_attention(&self.ws.n1, wq, wk, wv, wo, &mut m.kv, &cfg)
2292 }
2293 AttnKind::Linear(_) | AttnKind::LinearGdn(_) | AttnKind::ShortConv(_) => {
2294 unreachable!("MTP block is full attention")
2295 }
2296 };
2297 for (i, &a) in attn.iter().enumerate() {
2298 x[i] += a;
2299 }
2300 inference::rms_norm_into(
2301 &x,
2302 &lw.post_norm,
2303 self.rms_eps,
2304 self.norm_style,
2305 &mut self.ws.p1,
2306 );
2307 let ffn = ffn_forward(&lw.ffn, &self.ws.p1, self.pool.as_deref(), None);
2308 for (i, &f) in ffn.iter().enumerate() {
2309 x[i] += f;
2310 }
2311
2312 inference::rms_norm_into(
2313 &x,
2314 &m.final_norm,
2315 self.rms_eps,
2316 self.norm_style,
2317 &mut self.ws.n1,
2318 );
2319 let mut lg = self.lm_head_forward(&self.ws.n1);
2320 let draft = sampler::argmax(&lg);
2321 attention::recycle_buf(&mut lg);
2322 (draft, x)
2323 }
2324
2325 fn mtp_warm(&mut self, m: &mut MtpModule, hidden: &[f32], next_token: u32, position: usize) {
2329 let e = self.embed_single(next_token);
2330 let mut cat = vec![0.0f32; 2 * self.hidden_size];
2331 let (cat_e, cat_h) = cat.split_at_mut(self.hidden_size);
2332 inference::rms_norm_into(&e, &m.enorm, self.rms_eps, self.norm_style, cat_e);
2333 inference::rms_norm_into(hidden, &m.hnorm, self.rms_eps, self.norm_style, cat_h);
2334 let mut x = vec![0.0f32; self.hidden_size];
2335 m.eh_proj.matvec(&cat, &mut x, self.pool.as_deref());
2336 inference::rms_norm_into(&x, &m.layer.input_norm, self.rms_eps, self.norm_style, &mut self.ws.n1);
2337 let attn = match &m.layer.attn {
2338 AttnKind::Full { wq, wk, wv, wo, q_norm, k_norm, output_gate, softplus_gate, bias } => {
2339 let mut cfg = self.attn_cfg(position);
2340 cfg.q_norm = q_norm.as_deref();
2341 cfg.k_norm = k_norm.as_deref();
2342 cfg.output_gate = *output_gate;
2343 cfg.softplus_gate = softplus_gate.as_ref().map(|(g, p)| (g, *p));
2344 cfg.bias = bias.as_ref().map(|(q, k, v)| (q.as_slice(), k.as_slice(), v.as_slice()));
2345 attention::qwen_attention(&self.ws.n1, wq, wk, wv, wo, &mut m.kv, &cfg)
2346 }
2347 _ => return,
2348 };
2349 let _ = attn;
2350 }
2351
2352 #[cfg(feature = "gpu")]
2359 #[allow(clippy::too_many_arguments)]
2360 fn graph_spec_step(
2361 &mut self,
2362 m: &mut MtpModule,
2363 hidden: &[f32],
2364 t_next: u32,
2365 next_pos: usize,
2366 drafted: &mut usize,
2367 accepted: &mut usize,
2368 ) -> Option<(Vec<u32>, usize, Vec<f32>)> {
2369 let k_spec: usize = std::env::var("CMF_GRAPH_SPEC_K")
2370 .ok()
2371 .and_then(|v| v.parse().ok())
2372 .filter(|&v| (1..=8).contains(&v))
2373 .unwrap_or(2);
2374 if next_pos == 0 {
2375 return None;
2376 }
2377 let t_round = std::time::Instant::now();
2378 let mut drafts = Vec::with_capacity(k_spec);
2383 let (d1, mut hx) = self.mtp_step_h(m, hidden, t_next, next_pos - 1);
2384 drafts.push(d1);
2385 for j in 1..k_spec {
2386 let (dj, hj) = self.mtp_step_h(m, &hx, drafts[j - 1], next_pos - 1 + j);
2387 drafts.push(dj);
2388 hx = hj;
2389 }
2390 *drafted += k_spec;
2391 let t_draft = t_round.elapsed();
2392 let b = k_spec + 1;
2395 let mut hiddens = vec![0.0f32; b * self.hidden_size];
2396 for (i, &t) in std::iter::once(&t_next).chain(drafts.iter()).enumerate() {
2397 let e = self.embed_single(t);
2398 hiddens[i * self.hidden_size..(i + 1) * self.hidden_size].copy_from_slice(&e);
2399 }
2400 let positions: Vec<usize> = (next_pos..next_pos + b).collect();
2401 let (lm_gw, lm_rows) = {
2402 let (_, i, kind, rs) = self.weights.lm_head.graph_weight()?;
2403 (
2404 crate::gpu::GraphW { idx: i, kind, row_scale: rs, data: &[] },
2405 self.weights.lm_head.rows(),
2406 )
2407 };
2408 let mut logits = Vec::new();
2409 let final_norm = self.weights.final_norm.clone();
2410 let ok = self.try_batch_graph_wgpu(
2411 &mut hiddens,
2412 &positions,
2413 b,
2414 Some(crate::gpu::SpecTail {
2415 lm: lm_gw,
2416 lm_rows,
2417 final_norm: &final_norm,
2418 logits_out: &mut logits,
2419 }),
2420 );
2421 if !ok {
2422 m.kv.truncate_last(k_spec);
2425 return None;
2426 }
2427 if std::env::var("CMF_GRAPH_SPEC_TIME").is_ok() {
2428 eprintln!(
2429 "spec-round: draft {:.1} ms | verify {:.1} ms",
2430 t_draft.as_secs_f64() * 1e3,
2431 (t_round.elapsed() - t_draft).as_secs_f64() * 1e3,
2432 );
2433 }
2434 let ids: Vec<u32> = (0..b)
2436 .map(|i| sampler::argmax(&logits[i * lm_rows..(i + 1) * lm_rows]))
2437 .collect();
2438 let mut a = 0usize;
2439 while a < k_spec && ids[a] == drafts[a] {
2440 a += 1;
2441 }
2442 if a + 1 < b {
2444 crate::gpu::gdn_spec_restore(self.graph_kv_id, a);
2445 }
2446 *accepted += a;
2447 m.kv.truncate_last(k_spec.saturating_sub(1));
2450 for j in 0..a {
2451 let row = &hiddens[j * self.hidden_size..(j + 1) * self.hidden_size];
2452 let row = row.to_vec();
2453 self.mtp_warm(m, &row, ids[j], next_pos + j);
2454 }
2455 let mut row = logits[a * lm_rows..(a + 1) * lm_rows].to_vec();
2457 row.resize(self.vocab_size, 0.0);
2458 if let Some(c) = self.final_softcap {
2459 for l in row.iter_mut() {
2460 *l = c * (*l / c).tanh();
2461 }
2462 }
2463 self.graph_logits = Some(row);
2464 let new_hidden = hiddens[a * self.hidden_size..(a + 1) * self.hidden_size].to_vec();
2465 Some((drafts[..a].to_vec(), next_pos + a + 1, new_hidden))
2466 }
2467
2468 pub fn measure_pair_fusion(&mut self, iters: usize) -> (f64, f64) {
2477 if !self.pair_supported() {
2478 return (0.0, 0.0);
2479 }
2480 let emb1 = self.embed_single(1);
2481 let emb2 = self.embed_single(2);
2482 let pos = self.kv_cache.seq_len();
2483
2484 let t0 = std::time::Instant::now();
2485 for _ in 0..iters {
2486 let _ = self.forward_layers(&emb1, pos, None);
2487 let _ = self.forward_layers(&emb2, pos + 1, None);
2488 for l in &mut self.kv_cache.layers {
2489 l.truncate_last(2);
2490 }
2491 }
2492 let singles_ms = t0.elapsed().as_secs_f64() * 1000.0 / iters as f64;
2493
2494 let t1 = std::time::Instant::now();
2495 for _ in 0..iters {
2496 let _ = self.forward_pair(&emb1, &emb2, pos);
2497 for l in &mut self.kv_cache.layers {
2498 l.truncate_last(2);
2499 }
2500 }
2501 let pair_ms = t1.elapsed().as_secs_f64() * 1000.0 / iters as f64;
2502 (singles_ms, pair_ms)
2503 }
2504
2505 fn pair_supported(&self) -> bool {
2513 !self.weights.layers.is_empty()
2520 && self.g3n.is_none()
2521 && !self
2522 .weights
2523 .layers
2524 .iter()
2525 .any(|lw| matches!(&lw.attn, AttnKind::Mla(_) | AttnKind::Kda(_)))
2526 }
2527
2528 fn forward_pair(
2529 &mut self,
2530 emb1: &[f32],
2531 emb2: &[f32],
2532 position: usize,
2533 ) -> (Vec<f32>, Vec<f32>) {
2534 let mut h1 = emb1.to_vec();
2535 let mut h2 = emb2.to_vec();
2536 let (_nkv, _hd, hs, _rd, eps) = (
2537 self.num_kv_heads,
2538 self.head_dim,
2539 self.hidden_size,
2540 self.rotary_dim,
2541 self.rms_eps,
2542 );
2543 let pool = self.pool.clone();
2544
2545 for li in 0..self.num_layers {
2546 let lw = &self.weights.layers[self.phys_layer(li)];
2547 inference::rms_norm_into(
2550 &h1,
2551 &lw.input_norm,
2552 self.rms_eps,
2553 self.norm_style,
2554 &mut self.ws.n1,
2555 );
2556 inference::rms_norm_into(
2557 &h2,
2558 &lw.input_norm,
2559 self.rms_eps,
2560 self.norm_style,
2561 &mut self.ws.n2,
2562 );
2563
2564 let (a1, a2) = match &lw.attn {
2565 AttnKind::Mla(_) => unreachable!("MLA has no MTP/pair path"),
2566 AttnKind::Kda(_) => unreachable!("KDA has no MTP/pair path"),
2567 AttnKind::Linear(w) => {
2568 let cfg = self.vmf_cfg.expect("linear layer without vmf_cfg");
2569 let layer = &mut self.kv_cache.layers[li];
2570 let (state, scratch) = (&mut layer.linear_state, &mut layer.linear_scratch);
2571 vmf_phase_pair(
2572 &self.ws.n1,
2573 &self.ws.n2,
2574 w,
2575 &cfg,
2576 state,
2577 scratch,
2578 self.pool.as_deref(),
2579 )
2580 }
2581 AttnKind::LinearGdn(w) => {
2582 let cfg = self.gdn_cfg.expect("gdn layer without gdn_cfg");
2583 let layer = &mut self.kv_cache.layers[li];
2584 let (state, scratch) = (&mut layer.linear_state, &mut layer.linear_scratch);
2585 gdn_pair(
2586 &self.ws.n1,
2587 &self.ws.n2,
2588 w,
2589 &cfg,
2590 state,
2591 scratch,
2592 self.pool.as_deref(),
2593 )
2594 }
2595 AttnKind::ShortConv(w) => {
2596 let cfg = self
2597 .short_conv_cfg
2598 .expect("short-conv layer without short_conv_cfg");
2599 let layer = &mut self.kv_cache.layers[li];
2600 let (state, scratch) = (&mut layer.linear_state, &mut layer.linear_scratch);
2601 short_conv_pair(
2602 &self.ws.n1,
2603 &self.ws.n2,
2604 w,
2605 &cfg,
2606 state,
2607 scratch,
2608 self.pool.as_deref(),
2609 )
2610 }
2611 AttnKind::Full {
2612 wq,
2613 wk,
2614 wv,
2615 wo,
2616 q_norm,
2617 k_norm,
2618 output_gate,
2619 softplus_gate,
2620 bias,
2621 } => {
2622 let inv_freq_l = self.layer_inv_freq(li);
2623 let (nkv_l, hd_l, rd_l) = self.layer_geom(li);
2624 let cfg = QwenAttnCfg {
2625 num_heads: self.layer_num_heads(li),
2626 num_kv_heads: nkv_l,
2627 head_dim: hd_l,
2628 hidden_size: hs,
2629 position,
2630 inv_freq: &inv_freq_l,
2631 rotary_dim: rd_l,
2632 scale: self.attn_scale,
2633 softcap: self.attn_softcap,
2634 window: self.layer_window(li),
2635 v_norm: self.attn_v_norm,
2636 q_norm: q_norm.as_deref(),
2637 k_norm: k_norm.as_deref(),
2638 output_gate: *output_gate,
2639 softplus_gate: softplus_gate
2640 .as_ref()
2641 .map(|(gate, per_head)| (gate, *per_head)),
2642 rope_scale: self.layer_rope_scale(li),
2643 bias: bias
2644 .as_ref()
2645 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
2646 rms_eps: eps,
2647 norm_style: self.norm_style,
2648 pool: pool.as_deref(),
2649 };
2650 attention::qwen_attention_pair(
2651 &self.ws.n1,
2652 &self.ws.n2,
2653 wq,
2654 wk,
2655 wv,
2656 wo,
2657 &mut self.kv_cache.layers[li],
2658 &cfg,
2659 )
2660 }
2661 };
2662 let (a1, a2) = match &self.weights.layers[self.phys_layer(li)].attn_out_norm {
2663 Some(w) => (
2664 inference::rms_norm(&a1, w, self.rms_eps, self.norm_style),
2665 inference::rms_norm(&a2, w, self.rms_eps, self.norm_style),
2666 ),
2667 None => (a1, a2),
2668 };
2669 for i in 0..self.hidden_size {
2670 h1[i] += a1[i];
2671 h2[i] += a2[i];
2672 }
2673 let (mut a1, mut a2) = (a1, a2);
2674 attention::recycle_buf(&mut a1);
2675 attention::recycle_buf(&mut a2);
2676
2677 let lw = &self.weights.layers[self.phys_layer(li)];
2678 inference::rms_norm_into(
2679 &h1,
2680 &lw.post_norm,
2681 self.rms_eps,
2682 self.norm_style,
2683 &mut self.ws.p1,
2684 );
2685 inference::rms_norm_into(
2686 &h2,
2687 &lw.post_norm,
2688 self.rms_eps,
2689 self.norm_style,
2690 &mut self.ws.p2,
2691 );
2692 let (f1, f2) = match &lw.ffn {
2693 FfnKind::DenseMoe(dm) => (
2696 dense_moe_ffn(
2697 dm,
2698 &self.ws.p1,
2699 &h1,
2700 self.rms_eps,
2701 self.norm_style,
2702 self.pool.as_deref(),
2703 ),
2704 dense_moe_ffn(
2705 dm,
2706 &self.ws.p2,
2707 &h2,
2708 self.rms_eps,
2709 self.norm_style,
2710 self.pool.as_deref(),
2711 ),
2712 ),
2713 _ => ffn_forward_pair(
2714 &lw.ffn,
2715 &self.ws.p1,
2716 &self.ws.p2,
2717 self.pool.as_deref(),
2718 None,
2719 ),
2720 };
2721 let (f1, f2) = match &self.weights.layers[self.phys_layer(li)].ffn_out_norm {
2722 Some(w) => (
2723 inference::rms_norm(&f1, w, self.rms_eps, self.norm_style),
2724 inference::rms_norm(&f2, w, self.rms_eps, self.norm_style),
2725 ),
2726 None => (f1, f2),
2727 };
2728 for i in 0..self.hidden_size {
2729 h1[i] += f1[i];
2730 h2[i] += f2[i];
2731 }
2732 let (mut f1, mut f2) = (f1, f2);
2733 attention::recycle_buf(&mut f1);
2734 attention::recycle_buf(&mut f2);
2735 if let Some(sc) = self.weights.layers[self.phys_layer(li)].layer_scale {
2736 for i in 0..self.hidden_size {
2737 h1[i] *= sc;
2738 h2[i] *= sc;
2739 }
2740 }
2741 if self.is_loop_end(li) && li + 1 < self.num_layers {
2743 h1 = inference::rms_norm(
2744 &h1,
2745 &self.weights.final_norm,
2746 self.rms_eps,
2747 self.norm_style,
2748 );
2749 h2 = inference::rms_norm(
2750 &h2,
2751 &self.weights.final_norm,
2752 self.rms_eps,
2753 self.norm_style,
2754 );
2755 }
2756 }
2757 (h1, h2)
2758 }
2759
2760 fn commit_linear_scratch(&mut self) {
2762 for layer in &mut self.kv_cache.layers {
2763 if !layer.linear_scratch.is_empty() {
2764 std::mem::swap(&mut layer.linear_state, &mut layer.linear_scratch);
2765 layer.linear_scratch.clear();
2766 }
2767 }
2768 }
2769
2770 pub fn forward_ids(
2773 &mut self,
2774 ids: &[u32],
2775 task_mask: Option<&TaskMask>,
2776 ) -> Result<Vec<f32>, String> {
2777 if ids.is_empty() {
2778 return Err("empty id sequence".to_string());
2779 }
2780 self.kv_cache.clear();
2781 self.kv_history.clear();
2782 self.o1_begin();
2783 let mut hidden = vec![0.0f32; self.hidden_size];
2784 let mut pos = 0usize;
2785 if task_mask.is_none() && self.can_prefill_batched() && ids.len() > 2 {
2786 let chunk = prefill_chunk();
2790 let hs = self.hidden_size;
2791 while pos < ids.len() {
2792 let end = (pos + chunk).min(ids.len());
2793 let hb = self.prefill_batch(&ids[pos..end], pos);
2794 hidden.copy_from_slice(&hb[(end - pos - 1) * hs..]);
2795 pos = end;
2796 }
2797 }
2798 if task_mask.is_none()
2802 && !std::env::var("CMF_PAIR").is_ok_and(|v| v == "0")
2803 && self.pair_supported()
2804 {
2805 while pos + 1 < ids.len() {
2806 let e1 = self.embed_single(ids[pos]);
2807 let e2 = self.embed_single(ids[pos + 1]);
2808 let (_, h2) = self.forward_pair(&e1, &e2, pos);
2809 self.commit_linear_scratch();
2810 hidden = h2;
2811 pos += 2;
2812 }
2813 }
2814 while pos < ids.len() {
2815 hidden = self.forward_layers(&self.embed_single(ids[pos]), pos, task_mask);
2816 pos += 1;
2817 }
2818 self.o1_seal();
2822 let normed = inference::rms_norm(
2823 &hidden,
2824 &self.weights.final_norm,
2825 self.rms_eps,
2826 self.norm_style,
2827 );
2828 Ok(self.lm_head_forward(&normed))
2829 }
2830
2831 pub fn ppl_ids(&mut self, ids: &[u32]) -> f64 {
2838 let (nll, cnt) = self.nll_ids_from(ids, 0);
2839 (nll / cnt.max(1) as f64).exp()
2840 }
2841
2842 pub fn probe_ffn_mass(&mut self, ids: &[u32]) -> Vec<Vec<f64>> {
2847 self.kv_cache.clear();
2848 self.kv_history.clear();
2849 FFN_PROBE.with(|p| {
2850 *p.borrow_mut() = Some(vec![vec![0f64; self.intermediate_size]; self.num_layers]);
2851 });
2852 crate::gpu::cpu_scope(|| {
2853 for (pos, &id) in ids.iter().enumerate() {
2854 let emb = self.embed_single(id);
2855 let _ = self.forward_layers(&emb, pos, None);
2856 }
2857 });
2858 self.kv_cache.clear();
2859 self.kv_history.clear();
2860 FFN_PROBE
2861 .with(|p| p.borrow_mut().take())
2862 .unwrap_or_default()
2863 }
2864
2865 pub fn ppl_ids_masked(&mut self, ids: &[u32], mask: &TaskMask) -> f64 {
2869 self.kv_cache.clear();
2870 self.kv_history.clear();
2871 let mut nll = 0f64;
2872 let mut cnt = 0usize;
2873 let mut hidden = vec![0f32; self.hidden_size];
2874 for (pos, &id) in ids.iter().enumerate() {
2875 if pos > 0 {
2876 inference::rms_norm_into(
2877 &hidden,
2878 &self.weights.final_norm,
2879 self.rms_eps,
2880 self.norm_style,
2881 &mut self.ws.n1,
2882 );
2883 let mut logits = self.lm_head_forward(&self.ws.n1);
2884 let max = logits.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
2885 let sum: f64 = logits.iter().map(|&v| ((v - max) as f64).exp()).sum();
2886 let p = ((logits[id as usize] - max) as f64).exp() / sum.max(1e-300);
2887 nll -= p.max(1e-300).ln();
2888 cnt += 1;
2889 attention::recycle_buf(&mut logits);
2890 }
2891 let emb = self.embed_single(id);
2892 hidden = self.forward_layers(&emb, pos, Some(mask));
2893 }
2894 self.kv_cache.clear();
2895 self.kv_history.clear();
2896 (nll / cnt.max(1) as f64).exp()
2897 }
2898
2899 pub fn nll_ids_from(&mut self, ids: &[u32], start: usize) -> (f64, usize) {
2908 self.kv_cache.clear();
2909 self.kv_history.clear();
2910 let mut nll = 0f64;
2911 let mut cnt = 0usize;
2912 if self.can_prefill_batched() {
2913 const CHUNK: usize = 128;
2919 const LM_SUB: usize = 32;
2920 let n = ids.len().saturating_sub(1);
2921 let hs = self.hidden_size;
2922 let rows = self.weights.lm_head.rows();
2923 let mut pos = 0usize;
2924 while pos < n {
2925 let end = (pos + CHUNK).min(n);
2926 let bsz = end - pos;
2927 let hb = self.prefill_batch(&ids[pos..end], pos);
2928 let mut k0 = 0usize;
2929 while k0 < bsz {
2930 let k1 = (k0 + LM_SUB).min(bsz);
2931 let sb = k1 - k0;
2932 if pos + k1 <= start {
2935 k0 = k1;
2936 continue;
2937 }
2938 let mut normed = vec![0.0f32; sb * hs];
2939 for k in 0..sb {
2940 let r = inference::rms_norm(
2941 &hb[(k0 + k) * hs..(k0 + k + 1) * hs],
2942 &self.weights.final_norm,
2943 self.rms_eps,
2944 self.norm_style,
2945 );
2946 normed[k * hs..(k + 1) * hs].copy_from_slice(&r);
2947 }
2948 let mut logits = vec![0.0f32; sb * rows];
2949 self.weights
2950 .lm_head
2951 .matmat(&normed, sb, &mut logits, self.pool.as_deref());
2952 for k in 0..sb {
2953 if pos + k0 + k < start {
2954 continue;
2955 }
2956 let lg = &mut logits[k * rows..k * rows + self.vocab_size.min(rows)];
2957 if let Some(mu) = self.logit_multiplier {
2958 for v in lg.iter_mut() {
2959 *v *= mu;
2960 }
2961 }
2962 if let Some(c) = self.final_softcap {
2966 for v in lg.iter_mut() {
2967 *v = c * (*v / c).tanh();
2968 }
2969 }
2970 let lg = &logits[k * rows..k * rows + self.vocab_size.min(rows)];
2971 let target = ids[pos + k0 + k + 1] as usize;
2972 let max = lg.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
2973 let lse: f64 = lg
2974 .iter()
2975 .map(|&v| ((v - max) as f64).exp())
2976 .sum::<f64>()
2977 .ln()
2978 + max as f64;
2979 nll += lse - lg[target] as f64;
2980 cnt += 1;
2981 if std::env::var("CMF_PPL_TRACE").is_ok() {
2982 let top = lg
2983 .iter()
2984 .enumerate()
2985 .max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
2986 .map(|(i, _)| i)
2987 .unwrap_or(0);
2988 eprintln!(
2989 "BTRACE pos {} target {} nll {:.4} top {} lg_t {:.3} lg_top {:.3}",
2990 pos + k0 + k,
2991 target,
2992 lse - lg[target] as f64,
2993 top,
2994 lg[target],
2995 lg[top]
2996 );
2997 }
2998 }
2999 k0 = k1;
3000 }
3001 pos = end;
3002 }
3003 self.kv_cache.clear();
3004 self.kv_history.clear();
3005 return (nll, cnt);
3006 }
3007 for pos in 0..ids.len().saturating_sub(1) {
3008 let hidden = self.forward_layers(&self.embed_single(ids[pos]), pos, None);
3009 let out_of_band = self.graph_logits.take();
3017 if pos < start {
3018 continue;
3019 }
3020 let logits = match out_of_band {
3021 Some(lg) => lg,
3022 None => {
3023 let normed = inference::rms_norm(
3024 &hidden,
3025 &self.weights.final_norm,
3026 self.rms_eps,
3027 self.norm_style,
3028 );
3029 self.lm_head_forward(&normed)
3033 }
3034 };
3035 let target = ids[pos + 1] as usize;
3036 let max = logits.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
3037 let lse: f64 = logits
3038 .iter()
3039 .map(|&v| ((v - max) as f64).exp())
3040 .sum::<f64>()
3041 .ln()
3042 + max as f64;
3043 let tok_nll = lse - logits[target] as f64;
3044 if std::env::var("CMF_PPL_TRACE").is_ok() && pos < 48 {
3045 let top = logits
3046 .iter()
3047 .enumerate()
3048 .max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
3049 .map(|(i, _)| i)
3050 .unwrap_or(0);
3051 eprintln!(
3052 "pos {pos:3} tgt {target:6} nll {tok_nll:7.3} | top1 {top:6} lg[t]={:.2} lg[top]={:.2}",
3053 logits[target], logits[top]
3054 );
3055 }
3056 nll += tok_nll;
3057 cnt += 1;
3058 }
3059 self.kv_cache.clear();
3060 self.kv_history.clear();
3061 (nll, cnt)
3062 }
3063
3064 pub fn nll_ids_o1(&mut self, ids: &[u32], prefill: usize) -> (f64, usize) {
3080 self.kv_cache.clear();
3081 self.kv_history.clear();
3082 self.o1_begin();
3083 let n = ids.len().saturating_sub(1);
3084 let p = prefill.min(n);
3085 let mut pos = 0usize;
3087 if self.can_prefill_batched() {
3088 const CHUNK: usize = 128;
3089 while pos < p {
3090 let end = (pos + CHUNK).min(p);
3091 let _ = self.prefill_batch(&ids[pos..end], pos);
3092 pos = end;
3093 }
3094 } else {
3095 while pos < p {
3096 let _ = self.forward_layers(&self.embed_single(ids[pos]), pos, None);
3097 pos += 1;
3098 }
3099 }
3100 self.o1_seal();
3101
3102 let mut nll = 0f64;
3103 let mut cnt = 0usize;
3104 for pos in p..n {
3105 let hidden = self.forward_layers(&self.embed_single(ids[pos]), pos, None);
3106 let normed = inference::rms_norm(
3107 &hidden,
3108 &self.weights.final_norm,
3109 self.rms_eps,
3110 self.norm_style,
3111 );
3112 let logits = self.lm_head_forward(&normed);
3116 let target = ids[pos + 1] as usize;
3117 let max = logits.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
3118 let lse: f64 = logits
3119 .iter()
3120 .map(|&v| ((v - max) as f64).exp())
3121 .sum::<f64>()
3122 .ln()
3123 + max as f64;
3124 let tok_nll = lse - logits[target] as f64;
3125 if std::env::var("CMF_PPL_TRACE").is_ok() && pos < 48 {
3126 let top = logits
3127 .iter()
3128 .enumerate()
3129 .max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
3130 .map(|(i, _)| i)
3131 .unwrap_or(0);
3132 eprintln!(
3133 "pos {pos:3} tgt {target:6} nll {tok_nll:7.3} | top1 {top:6} lg[t]={:.2} lg[top]={:.2}",
3134 logits[target], logits[top]
3135 );
3136 }
3137 nll += tok_nll;
3138 cnt += 1;
3139 }
3140 self.kv_cache.clear();
3141 self.kv_history.clear();
3142 (nll, cnt)
3143 }
3144
3145 pub fn calib_ids(&mut self, ids: &[u32], temps: &[f32]) -> (Vec<bool>, Vec<Vec<f32>>) {
3153 self.kv_cache.clear();
3154 self.kv_history.clear();
3155 let n = ids.len().saturating_sub(1);
3156 let mut correct = Vec::with_capacity(n);
3157 let mut pmax = Vec::with_capacity(n);
3158 for pos in 0..n {
3159 let emb = self.embed_single(ids[pos]);
3160 let hidden = self.forward_layers(&emb, pos, None);
3161 let normed = inference::rms_norm(
3162 &hidden,
3163 &self.weights.final_norm,
3164 self.rms_eps,
3165 self.norm_style,
3166 );
3167 let logits = self.lm_head_forward(&normed);
3171 let target = ids[pos + 1] as usize;
3172 let (mut amax, mut mval) = (0usize, f32::NEG_INFINITY);
3173 for (i, &v) in logits.iter().enumerate() {
3174 if v > mval {
3175 mval = v;
3176 amax = i;
3177 }
3178 }
3179 correct.push(amax == target);
3180 let row: Vec<f32> = temps
3181 .iter()
3182 .map(|&t| {
3183 let tt = t.max(1e-3);
3184 let s: f32 = logits.iter().map(|&v| ((v - mval) / tt).exp()).sum();
3185 1.0 / s.max(1e-12) })
3187 .collect();
3188 pmax.push(row);
3189 }
3190 self.kv_cache.clear();
3191 self.kv_history.clear();
3192 (correct, pmax)
3193 }
3194
3195 pub fn ppl_ids_dynamic(&mut self, ids: &[u32]) -> (f64, usize) {
3202 let mut router = match self.dyn_router.take() {
3203 Some(r) => r,
3204 None => return (self.ppl_ids(ids), 0),
3205 };
3206 router.reset();
3207 self.dyn_phi_seen = 0;
3208 let _ = self.set_active_skill(None);
3209
3210 self.kv_cache.clear();
3211
3212 self.kv_history.clear();
3213 let mut nll = 0f64;
3214 let mut cnt = 0usize;
3215 for pos in 0..ids.len().saturating_sub(1) {
3216 let hidden = self.forward_layers(&self.embed_single(ids[pos]), pos, None);
3217 let normed = inference::rms_norm(
3218 &hidden,
3219 &self.weights.final_norm,
3220 self.rms_eps,
3221 self.norm_style,
3222 );
3223 let logits = self.lm_head_forward(&normed);
3227 let target = ids[pos + 1] as usize;
3228 let max = logits.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
3229 let lse: f64 = logits
3230 .iter()
3231 .map(|&v| ((v - max) as f64).exp())
3232 .sum::<f64>()
3233 .ln()
3234 + max as f64;
3235 let tok_nll = lse - logits[target] as f64;
3236 if std::env::var("CMF_PPL_TRACE").is_ok() && pos < 48 {
3237 let top = logits
3238 .iter()
3239 .enumerate()
3240 .max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
3241 .map(|(i, _)| i)
3242 .unwrap_or(0);
3243 eprintln!(
3244 "pos {pos:3} tgt {target:6} nll {tok_nll:7.3} | top1 {top:6} lg[t]={:.2} lg[top]={:.2}",
3245 logits[target], logits[top]
3246 );
3247 }
3248 nll += tok_nll;
3249 cnt += 1;
3250 let phi = self.dyn_phi_ema.clone();
3252 if let Some(new_active) = router.step(&phi, pos) {
3253 let _ = self.set_active_skill(new_active);
3254 }
3255 }
3256 let switches = router.switches.len();
3257 let _ = self.set_active_skill(None);
3258 self.dyn_router = Some(router);
3259 self.kv_cache.clear();
3260 self.kv_history.clear();
3261 ((nll / cnt.max(1) as f64).exp(), switches)
3262 }
3263
3264 pub fn probe_phi(&mut self, ids: &[u32], layer: usize) -> Vec<f32> {
3266 self.kv_cache.clear();
3267 self.kv_history.clear();
3268 let mut acc = vec![0f32; self.hidden_size];
3269 for (pos, &id) in ids.iter().enumerate() {
3270 let h = self.forward_layers_upto(&self.embed_single(id), pos, None, Some(layer));
3271 for (a, v) in acc.iter_mut().zip(&h) {
3272 *a += v;
3273 }
3274 }
3275 let n = ids.len().max(1) as f32;
3276 for a in acc.iter_mut() {
3277 *a /= n;
3278 }
3279 self.kv_cache.clear();
3280 self.kv_history.clear();
3281 acc
3282 }
3283
3284 fn prefill_batch(&mut self, ids: &[u32], start_pos: usize) -> Vec<f32> {
3290 let b = ids.len();
3291 let hs = self.hidden_size;
3292 let mut h: Vec<f32> = vec![0.0; b * hs];
3295 let mut h_ready = false;
3296 let fill_h = |h: &mut Vec<f32>, me: &Self| {
3297 for (bi, &id) in ids.iter().enumerate() {
3298 let e = me.embed_single(id);
3299 h[bi * hs..(bi + 1) * hs].copy_from_slice(&e);
3300 }
3301 };
3302 let (_nkv, _hd, _rd, eps) = (
3303 self.num_kv_heads,
3304 self.head_dim,
3305 self.rotary_dim,
3306 self.rms_eps,
3307 );
3308 let pool = self.pool.clone();
3309 let norm_style = self.norm_style;
3310
3311 #[cfg(target_os = "macos")]
3312 let mut chunk_skip_until = 0usize;
3313 for li in 0..self.num_layers {
3314 crate::gpu::set_layer(li as i64); #[cfg(target_os = "macos")]
3321 {
3322 if li < chunk_skip_until {
3323 continue;
3324 }
3325 if !h_ready && li == 0 && self.weights.embed_tokens.q8_row_parts().is_none() {
3331 fill_h(&mut h, self);
3332 h_ready = true;
3333 }
3334 let ids_for_embed = (!h_ready && li == 0).then_some(ids);
3335 let end = self.chunk_run_gpu(li, &mut h, b, start_pos, ids_for_embed);
3336 if end > li {
3337 h_ready = true;
3338 chunk_skip_until = end;
3339 if self.is_loop_end(end - 1) && end < self.num_layers {
3342 for bi in 0..b {
3343 let normed = inference::rms_norm(
3344 &h[bi * hs..(bi + 1) * hs],
3345 &self.weights.final_norm,
3346 eps,
3347 norm_style,
3348 );
3349 h[bi * hs..(bi + 1) * hs].copy_from_slice(&normed);
3350 }
3351 }
3352 continue;
3353 }
3354 }
3355 if !h_ready {
3356 fill_h(&mut h, self);
3357 h_ready = true;
3358 }
3359 let lw = &self.weights.layers[self.phys_layer(li)];
3360 match &lw.attn {
3362 AttnKind::Kda(w) => {
3363 let cfg = self.kda_cfg.expect("kda layer without kda_cfg");
3365 let mut normed = vec![0.0f32; b * hs];
3366 for bi in 0..b {
3367 inference::rms_norm_into(
3368 &h[bi * hs..(bi + 1) * hs],
3369 &lw.input_norm,
3370 eps,
3371 norm_style,
3372 &mut normed[bi * hs..(bi + 1) * hs],
3373 );
3374 }
3375 let attn = crate::linear_core::kda_forward_batch(
3376 &normed,
3377 b,
3378 w,
3379 &cfg,
3380 &mut self.kv_cache.layers[li].linear_state,
3381 pool.as_deref(),
3382 );
3383 for (dst, &a) in h.iter_mut().zip(&attn) {
3384 *dst += a;
3385 }
3386 }
3387 AttnKind::LinearGdn(w) => {
3388 let cfg = self.gdn_cfg.expect("gdn layer without gdn_cfg");
3390 let mut normed = vec![0.0f32; b * hs];
3391 for bi in 0..b {
3392 let r = inference::rms_norm(
3393 &h[bi * hs..(bi + 1) * hs],
3394 &lw.input_norm,
3395 eps,
3396 norm_style,
3397 );
3398 normed[bi * hs..(bi + 1) * hs].copy_from_slice(&r);
3399 }
3400 let attn = crate::linear_core::gdn_forward_batch(
3401 &normed,
3402 b,
3403 w,
3404 &cfg,
3405 &mut self.kv_cache.layers[li].linear_state,
3406 pool.as_deref(),
3407 );
3408 for (dst, &a) in h.iter_mut().zip(&attn) {
3409 *dst += a;
3410 }
3411 }
3412 AttnKind::ShortConv(w) => {
3413 let cfg = self
3416 .short_conv_cfg
3417 .expect("short-conv layer without short_conv_cfg");
3418 let mut normed = vec![0.0f32; b * hs];
3419 for bi in 0..b {
3420 inference::rms_norm_into(
3421 &h[bi * hs..(bi + 1) * hs],
3422 &lw.input_norm,
3423 eps,
3424 norm_style,
3425 &mut normed[bi * hs..(bi + 1) * hs],
3426 );
3427 }
3428 let attn = short_conv_forward_batch(
3429 &normed,
3430 b,
3431 w,
3432 &cfg,
3433 &mut self.kv_cache.layers[li].linear_state,
3434 pool.as_deref(),
3435 );
3436 for (dst, &a) in h.iter_mut().zip(&attn) {
3437 *dst += a;
3438 }
3439 }
3440 AttnKind::Mla(w) => {
3441 let inv_freq_l = self.layer_inv_freq(li);
3444 let rs = self.layer_rope_scale(li);
3445 let mut normed = vec![0.0f32; hs];
3446 for bi in 0..b {
3447 inference::rms_norm_into(
3448 &h[bi * hs..(bi + 1) * hs],
3449 &lw.input_norm,
3450 eps,
3451 norm_style,
3452 &mut normed,
3453 );
3454 let ao = mla_attention(
3455 w,
3456 &normed,
3457 &mut self.kv_cache.layers[li],
3458 start_pos + bi,
3459 &inv_freq_l,
3460 rs,
3461 eps,
3462 pool.as_deref(),
3463 );
3464 for (dst, &a) in h[bi * hs..(bi + 1) * hs].iter_mut().zip(&ao) {
3465 *dst += a;
3466 }
3467 }
3468 }
3469 AttnKind::Full {
3470 wq,
3471 wk,
3472 wv,
3473 wo,
3474 q_norm,
3475 k_norm,
3476 output_gate,
3477 softplus_gate,
3478 bias,
3479 } => {
3480 let mut normed = vec![0.0f32; b * hs];
3484 for bi in 0..b {
3485 inference::rms_norm_into(
3486 &h[bi * hs..(bi + 1) * hs],
3487 &lw.input_norm,
3488 eps,
3489 norm_style,
3490 &mut normed[bi * hs..(bi + 1) * hs],
3491 );
3492 }
3493 let inv_freq_l = self.layer_inv_freq(li);
3494 let (nkv_l, hd_l, rd_l) = self.layer_geom(li);
3495 let cfg = QwenAttnCfg {
3496 num_heads: self.layer_num_heads(li),
3497 num_kv_heads: nkv_l,
3498 head_dim: hd_l,
3499 hidden_size: hs,
3500 position: start_pos,
3501 inv_freq: &inv_freq_l,
3502 rotary_dim: rd_l,
3503 scale: self.attn_scale,
3504 softcap: self.attn_softcap,
3505 window: self.layer_window(li),
3506 v_norm: self.attn_v_norm,
3507 q_norm: q_norm.as_deref(),
3508 k_norm: k_norm.as_deref(),
3509 output_gate: *output_gate,
3510 softplus_gate: softplus_gate
3511 .as_ref()
3512 .map(|(gate, per_head)| (gate, *per_head)),
3513 rope_scale: self.layer_rope_scale(li),
3514 bias: bias
3515 .as_ref()
3516 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
3517 rms_eps: eps,
3518 norm_style,
3519 pool: pool.as_deref(),
3520 };
3521 let mut attn = attention::qwen_attention_batch(
3522 &normed,
3523 b,
3524 wq,
3525 wk,
3526 wv,
3527 wo,
3528 &mut self.kv_cache.layers[li],
3529 &cfg,
3530 );
3531 if let Some(w) = &lw.attn_out_norm {
3532 for bi in 0..b {
3533 inference::rms_norm_into(
3534 &attn[bi * hs..(bi + 1) * hs],
3535 w,
3536 eps,
3537 norm_style,
3538 &mut normed[bi * hs..(bi + 1) * hs],
3539 );
3540 }
3541 attn.copy_from_slice(&normed);
3542 }
3543 for (dst, &a) in h.iter_mut().zip(&attn) {
3544 *dst += a;
3545 }
3546 }
3547 AttnKind::Linear(w) => {
3548 for bi in 0..b {
3549 let normed = inference::rms_norm(
3550 &h[bi * hs..(bi + 1) * hs],
3551 &lw.input_norm,
3552 eps,
3553 norm_style,
3554 );
3555 vmf_phase_forward(
3556 &normed,
3557 w,
3558 &self.vmf_cfg.expect("linear layer without vmf_cfg"),
3559 &mut self.kv_cache.layers[li].linear_state,
3560 pool.as_deref(),
3561 )
3562 .iter()
3563 .enumerate()
3564 .for_each(|(i, &a)| h[bi * hs + i] += a);
3565 }
3566 }
3567 }
3568
3569 let lw = &self.weights.layers[self.phys_layer(li)];
3571 let mut post = vec![0.0f32; b * hs];
3572 for bi in 0..b {
3573 let r =
3574 inference::rms_norm(&h[bi * hs..(bi + 1) * hs], &lw.post_norm, eps, norm_style);
3575 post[bi * hs..(bi + 1) * hs].copy_from_slice(&r);
3576 }
3577 let mut ffn = match &lw.ffn {
3578 FfnKind::Dense(d) => dense_ffn_batch(d, &post, b, pool.as_deref()),
3579 FfnKind::Moe(m) => moe_ffn_batch(m, &post, b, hs, pool.as_deref(), None),
3580 FfnKind::DenseMoe(dm) => {
3583 let mut out = vec![0.0f32; b * hs];
3584 for bi in 0..b {
3585 let r = dense_moe_ffn(
3586 dm,
3587 &post[bi * hs..(bi + 1) * hs],
3588 &h[bi * hs..(bi + 1) * hs],
3589 eps,
3590 norm_style,
3591 pool.as_deref(),
3592 );
3593 out[bi * hs..(bi + 1) * hs].copy_from_slice(&r);
3594 }
3595 out
3596 }
3597 };
3598 if let Some(w) = &lw.ffn_out_norm {
3599 for bi in 0..b {
3600 inference::rms_norm_into(
3601 &ffn[bi * hs..(bi + 1) * hs],
3602 w,
3603 eps,
3604 norm_style,
3605 &mut post[bi * hs..(bi + 1) * hs],
3606 );
3607 }
3608 ffn.copy_from_slice(&post);
3609 }
3610 for (dst, &f) in h.iter_mut().zip(&ffn) {
3611 *dst += f;
3612 }
3613 if let Some(sc) = lw.layer_scale {
3614 for v in h.iter_mut() {
3615 *v *= sc;
3616 }
3617 }
3618 if let Ok(tp) = std::env::var("CMF_TRACE_POS") {
3619 if let Some(t) = tp.parse::<usize>().ok() {
3620 if t >= start_pos && t < start_pos + b {
3621 let bi = t - start_pos;
3622 let row = &h[bi * hs..(bi + 1) * hs];
3623 let n: f32 = row.iter().map(|x| x * x).sum::<f32>().sqrt();
3624 eprintln!(
3625 "BATCH pos {t} after layer {li}: |h| = {n:.6} h0 {:.6} h1 {:.6}",
3626 row[0], row[1]
3627 );
3628 }
3629 }
3630 }
3631 if std::env::var("CMF_DEBUG_LAYERS").is_ok() {
3635 let row = &h[(b - 1) * hs..b * hs];
3636 let rms =
3637 (row.iter().map(|&v| (v as f64) * (v as f64)).sum::<f64>() / hs as f64).sqrt();
3638 let mx = row.iter().fold(0f32, |m, &v| m.max(v.abs()));
3639 eprintln!(
3640 "layer {li:>3} {:>10} ffn={:<5} rms={rms:>12.4} max={mx:>12.4}",
3641 match &self.weights.layers[self.phys_layer(li)].attn {
3642 AttnKind::LinearGdn(_) => "gdn",
3643 AttnKind::Linear(_) => "vmf",
3644 AttnKind::ShortConv(_) => "conv",
3645 _ => "attn",
3646 },
3647 match &lw.ffn {
3648 FfnKind::Moe(_) => "moe",
3649 FfnKind::Dense(_) => "dense",
3650 FfnKind::DenseMoe(_) => "dense+moe",
3651 },
3652 );
3653 }
3654 if self.is_loop_end(li) && li + 1 < self.num_layers {
3656 for bi in 0..b {
3657 let normed = inference::rms_norm(
3658 &h[bi * hs..(bi + 1) * hs],
3659 &self.weights.final_norm,
3660 eps,
3661 norm_style,
3662 );
3663 h[bi * hs..(bi + 1) * hs].copy_from_slice(&normed);
3664 }
3665 }
3666 if std::env::var("CMF_TRACE_H").is_ok() {
3667 let n = h[..hs].iter().map(|v| v.abs()).sum::<f32>() / hs as f32;
3668 let mx = h[..hs].iter().fold(0.0f32, |a, &v| a.max(v.abs()));
3669 eprintln!(
3670 "layer {li}: mean|h|={n:.4} max|h|={mx:.2} scale={:?}",
3671 lw.layer_scale
3672 );
3673 }
3674 }
3675 crate::gpu::set_layer(-1); h
3677 }
3678
3679 fn embed_single(&self, id: u32) -> Vec<f32> {
3681 let mut out = vec![0.0f32; self.hidden_size];
3682 if (id as usize) < self.weights.embed_tokens.rows() {
3683 self.weights.embed_tokens.row_f32(id as usize, &mut out);
3684 }
3685 if self.embed_multiplier != 1.0 {
3686 for v in out.iter_mut() {
3687 *v *= self.embed_multiplier;
3688 }
3689 }
3690 if self.dsv4.is_some() {
3694 let mut v = vec![0.0f32; self.hidden_size.max(1)];
3695 v[0] = id as f32;
3696 return v;
3697 }
3698 if let Some(b) = &self.g3n {
3701 return b.0.extend_embedding(id, &out, self.pool.as_deref());
3702 }
3703 out
3704 }
3705
3706 #[cfg(target_os = "macos")]
3712 fn chunk_run_gpu(
3713 &mut self,
3714 li0: usize,
3715 h: &mut [f32],
3716 b: usize,
3717 pos0: usize,
3718 embed_ids: Option<&[u32]>,
3719 ) -> usize {
3720 if !crate::gpu::enabled_here()
3724 || std::env::var("CMF_GPU_CHUNK")
3725 .map(|v| v == "0")
3726 .unwrap_or(false)
3727 || b < 32
3728 || self.swa.is_some()
3729 || self.global_attn.is_some()
3730 || self.attn_v_norm
3731 || (self.attn_scale - 1.0 / (self.head_dim as f32).sqrt()).abs() > 1e-9
3732 {
3733 return li0;
3734 }
3735 let Some(model) = self.model.clone() else {
3736 return li0;
3737 };
3738 let inv_freq = self.inv_freq.clone();
3739 let (nh, nkv, hd, hs) = (
3740 self.num_heads,
3741 self.num_kv_heads,
3742 self.head_dim,
3743 self.hidden_size,
3744 );
3745 let loop_end = if self.loop_final_norm {
3749 ((li0 / self.physical_layers) + 1) * self.physical_layers
3750 } else {
3751 self.num_layers
3752 };
3753 let mut layers: Vec<crate::gpu_metal::ChunkLayer> = Vec::new();
3754 let mut stored_at: Vec<usize> = Vec::new();
3755 for li in li0..self.num_layers.min(loop_end) {
3756 let lw = &self.weights.layers[self.phys_layer(li)];
3757 if lw.attn_out_norm.is_some() || lw.ffn_out_norm.is_some() || lw.layer_scale.is_some() {
3758 break;
3759 }
3760 let AttnKind::Full {
3761 wq,
3762 wk,
3763 wv,
3764 wo,
3765 q_norm,
3766 k_norm,
3767 output_gate: false,
3768 softplus_gate: None,
3769 bias,
3770 } = &lw.attn
3771 else {
3772 break;
3773 };
3774 let FfnKind::Dense(d) = &lw.ffn else { break };
3775 if d.act != Act::Silu {
3776 break;
3777 }
3778 fn cw(t: &QTensor) -> Option<(usize, usize, usize, &[f32])> {
3783 t.q8_row_parts()
3784 .or_else(|| t.q4t_parts().map(|(i, r, c)| (i, r, c, &[][..])))
3785 .or_else(|| t.q4tp_parts().map(|(i, r, c)| (i, r, c, &[][..])))
3786 }
3787 let parts = (
3788 cw(wq),
3789 cw(wk),
3790 cw(wv),
3791 cw(wo),
3792 cw(&d.gate_proj),
3793 cw(&d.up_proj),
3794 cw(&d.down_proj),
3795 );
3796 let (Some(pq), Some(pk), Some(pv), Some(po), Some(pg), Some(pu), Some(pd)) = parts
3797 else {
3798 break;
3799 };
3800 let layer = &self.kv_cache.layers[li];
3801 if layer.mode != crate::kv_cache::KvMode::F32 || layer.o1.is_some() {
3802 break;
3803 }
3804 stored_at.push(layer.head_len(0));
3805 layers.push(crate::gpu_metal::ChunkLayer {
3806 model: &model,
3807 kv_id: self.graph_kv_id,
3808 layer: li,
3809 wq: pq,
3810 wk: pk,
3811 wv: pv,
3812 wo: po,
3813 gate: pg,
3814 up: pu,
3815 down: pd,
3816 input_norm: &lw.input_norm,
3817 post_norm: &lw.post_norm,
3818 bias: bias
3819 .as_ref()
3820 .map(|(a, bb, cc)| (a.as_slice(), bb.as_slice(), cc.as_slice())),
3821 q_norm: q_norm.as_deref(),
3822 k_norm: k_norm.as_deref(),
3823 inv_freq: &inv_freq,
3824 rd: self.rotary_dim,
3825 nh,
3826 nkv,
3827 hd,
3828 hs,
3829 inter: d.gate_proj.rows(),
3830 gemma: matches!(self.norm_style, cortiq_core::NormStyle::Gemma),
3831 eps: self.rms_eps as f32,
3832 });
3833 }
3834 if layers.is_empty() {
3835 return li0;
3836 }
3837 let row = nkv * hd;
3838 let mut store: Vec<(Vec<f32>, Vec<f32>, Vec<f32>)> = stored_at
3839 .iter()
3840 .map(|&st| (vec![0f32; b * row], vec![0f32; b * row], vec![0f32; st + b]))
3841 .collect();
3842 let mut io: Vec<crate::gpu_metal::ChunkIo> = Vec::with_capacity(layers.len());
3843 for (i, (ok, ov, oi)) in store.iter_mut().enumerate() {
3844 let li = layers[i].layer;
3845 let layer = &self.kv_cache.layers[li];
3846 io.push(crate::gpu_metal::ChunkIo {
3847 cpu_stored: stored_at[i],
3848 cpu_k: (0..nkv).map(|g| layer.head_keys(g)).collect(),
3849 cpu_v: (0..nkv).map(|g| layer.head_values(g)).collect(),
3850 out_k: ok,
3851 out_v: ov,
3852 imp: oi,
3853 });
3854 }
3855 let n_run = layers.len();
3856 let last = layers.last().map(|l| l.layer + 1).unwrap_or(li0);
3857 let ep = embed_ids.and_then(|ids| {
3860 self.weights
3861 .embed_tokens
3862 .q8_row_parts()
3863 .map(|(idx, rows, _c, rs)| crate::gpu_metal::ChunkEmbed {
3864 idx,
3865 rows,
3866 row_scale: rs,
3867 ids,
3868 mult: self.embed_multiplier,
3869 })
3870 });
3871 if embed_ids.is_some() && ep.is_none() {
3872 return li0;
3873 }
3874 if !crate::gpu_metal::chunk_run_gpu(&layers, &mut io, h, b, pos0, ep.as_ref()) {
3875 return li0;
3876 }
3877 drop(io);
3878 drop(layers);
3879 for (i, (ok, ov, oi)) in store.iter().enumerate().take(n_run) {
3882 let li = li0 + i;
3883 let layer = &mut self.kv_cache.layers[li];
3884 for bi in 0..b {
3885 layer.append(
3886 &ok[bi * row..(bi + 1) * row],
3887 &ov[bi * row..(bi + 1) * row],
3888 &[],
3889 );
3890 }
3891 layer.accumulate_imp(oi);
3892 }
3893 last
3894 }
3895
3896 fn layer_is_local(&self, li: usize) -> bool {
3899 if let Some(layers) = &self.sliding_layers {
3900 return layers.get(li).copied().unwrap_or(false);
3901 }
3902 match self.swa {
3903 Some((_, pattern)) => (li + 1) % pattern.max(1) != 0,
3904 None => false,
3905 }
3906 }
3907
3908 fn layer_inv_freq(&self, li: usize) -> std::sync::Arc<Vec<f32>> {
3911 if self.layer_is_local(li) {
3912 if let Some(f) = &self.inv_freq_local {
3913 return f.clone();
3914 }
3915 } else if let Some(f) = &self.inv_freq_global {
3916 return f.clone();
3917 }
3918 self.inv_freq.clone()
3919 }
3920
3921 fn layer_window(&self, li: usize) -> Option<usize> {
3923 self.swa
3924 .and_then(|(w, _)| self.layer_is_local(li).then_some(w))
3925 }
3926
3927 fn layer_num_heads(&self, li: usize) -> usize {
3928 self.attention_heads_per_layer
3929 .as_ref()
3930 .and_then(|v| v.get(li).copied())
3931 .unwrap_or(self.num_heads)
3932 }
3933
3934 fn layer_rope_scale(&self, li: usize) -> f32 {
3935 if self.layer_is_local(li) {
3936 self.rope_scale_local
3937 } else {
3938 self.rope_scale
3939 }
3940 }
3941
3942 fn layer_geom(&self, li: usize) -> (usize, usize, usize) {
3945 if !self.layer_is_local(li) {
3946 if let Some((ghd, gkv)) = self.global_attn {
3947 return (gkv, ghd, ghd);
3948 }
3949 }
3950 (
3951 self.num_kv_heads,
3952 self.head_dim,
3953 if self.layer_is_local(li) {
3954 self.rotary_dim_local.unwrap_or(self.rotary_dim)
3955 } else {
3956 self.rotary_dim
3957 },
3958 )
3959 }
3960
3961 fn forward_layers(
3963 &mut self,
3964 hidden: &[f32],
3965 position: usize,
3966 task_mask: Option<&TaskMask>,
3967 ) -> Vec<f32> {
3968 self.forward_layers_upto(hidden, position, task_mask, None)
3969 }
3970
3971 fn try_token_graph_wgpu(
3975 &self,
3976 hidden: &[f32],
3977 position: usize,
3978 logits_out: &mut Vec<f32>,
3979 layers_run: &mut usize,
3980 ) -> Option<Vec<f32>> {
3981 self.try_token_graph_wgpu_steps(hidden, position, logits_out, 1, None, Some(layers_run))
3982 }
3983
3984 fn try_multi_burst(&self, t_next: u32, position: usize, k: usize) -> Option<Vec<u32>> {
3988 if self.o1_active() || self.attn_softcap > 0.0 {
3989 return None;
3990 }
3991 let graph_on = match std::env::var("CMF_GPU_WGPU_GRAPH").ok().as_deref() {
3992 Some("0") => return None,
3993 Some(_) => true,
3994 None => crate::gpu::wgpu_graph_default(),
3995 };
3996 if !graph_on {
3997 return None;
3998 }
3999 let emb = self.embed_single(t_next);
4000 let mut lg = Vec::new();
4001 let mut ids = Vec::new();
4002 self.try_token_graph_wgpu_steps(&emb, position, &mut lg, k, Some(&mut ids), None)?;
4003 (ids.len() == k).then_some(ids)
4004 }
4005
4006 fn try_token_graph_wgpu_steps(
4010 &self,
4011 hidden: &[f32],
4012 position: usize,
4013 logits_out: &mut Vec<f32>,
4014 steps: usize,
4015 ids_out: Option<&mut Vec<u32>>,
4016 layers_run: Option<&mut usize>,
4017 ) -> Option<Vec<f32>> {
4018 let o1_gpu = std::env::var("CMF_O1_GPU").as_deref() == Ok("1");
4021 if (self.o1_active() && !o1_gpu) || self.attn_softcap > 0.0 {
4022 return None;
4026 }
4027 let o1_views: Vec<Option<Vec<crate::nystrom::O1DeviceView<'_>>>> = (0..self.num_layers)
4032 .map(|li| {
4033 if !o1_gpu {
4034 return None;
4035 }
4036 self.kv_cache.layers[self.phys_layer(li)].o1_views()
4037 })
4038 .collect();
4039 if self.o1_active() && o1_gpu {
4040 let want: usize = (0..self.num_layers)
4043 .filter(|li| !matches!(self.kv_cache.layers[self.phys_layer(*li)].o1, None))
4044 .count();
4045 let have = o1_views.iter().filter(|v| v.is_some()).count();
4046 if want == 0 || have != want {
4047 return None;
4048 }
4049 }
4050 let nh = self.num_heads;
4051 let (nkv, hd, rd) = self.layer_geom(0);
4052 let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
4053 let mut layers = Vec::with_capacity(self.num_layers);
4054 let mut model = None;
4055 let dbg = std::env::var("CMF_GRAPH_DEBUG").is_ok();
4056 fn gw(t: &QTensor) -> Option<crate::gpu::GraphW<'_>> {
4057 if let Some((_, i, kind, rs)) = t.graph_weight() {
4058 return Some(crate::gpu::GraphW {
4059 idx: i,
4060 kind,
4061 row_scale: rs,
4062 data: &[],
4063 });
4064 }
4065 t.as_f32().map(|d| crate::gpu::GraphW {
4067 idx: 0,
4068 kind: 4,
4069 row_scale: &[],
4070 data: d,
4071 })
4072 }
4073 for li in 0..self.num_layers {
4074 let lw = &self.weights.layers[self.phys_layer(li)];
4075 if dbg {
4076 let ak = match &lw.attn {
4077 AttnKind::Mla(_) => "Mla".into(),
4078 AttnKind::Full {
4079 output_gate, bias, ..
4080 } => format!("Full gate={output_gate} bias={}", bias.is_some()),
4081 AttnKind::LinearGdn(_) => "LinearGdn".into(),
4082 AttnKind::Kda(_) => "Kda".into(),
4083 AttnKind::Linear(_) => "Linear".into(),
4084 AttnKind::ShortConv(_) => "ShortConv".into(),
4085 };
4086 let fk = match &lw.ffn {
4087 FfnKind::Dense(_) => "Dense",
4088 FfnKind::Moe(_) => "Moe",
4089 FfnKind::DenseMoe(_) => "DenseMoe",
4090 };
4091 eprintln!("graph L{li}: attn={ak} ffn={fk}");
4092 }
4093 let gffn = match &lw.ffn {
4094 FfnKind::DenseMoe(_) => return None, FfnKind::Dense(d) => crate::gpu::GraphFfn::Dense {
4096 gate: gw(&d.gate_proj)?,
4097 up: gw(&d.up_proj)?,
4098 down: gw(&d.down_proj)?,
4099 },
4100 FfnKind::Moe(m) => {
4101 if m.router_sigmoid
4106 || m.expert_bias.is_some()
4107 || m.route_tau.is_some()
4108 || m.mask.is_some()
4109 {
4110 return None;
4111 }
4112 let (se, sg) = m.shared.as_ref()?;
4113 let sgate = gw(sg.as_ref()?)?;
4114 let router = gw(&m.router)?;
4115 let inter = m.experts.first()?.gate_proj.rows();
4116 let mut experts = Vec::with_capacity(m.experts.len() + 1);
4117 let mut q4tp: Option<bool> = None;
4120 let mut gu_q2: Option<bool> = None;
4123 for e in m.experts.iter().chain(std::iter::once(se)) {
4124 if !matches!(e.act, Act::Silu)
4125 || e.gate_proj.rows() != inter
4126 || e.up_proj.rows() != inter
4127 {
4128 return None;
4129 }
4130 let (mm, gi, ui, di, is_p, is_q2) = match e.gate_proj.mapped_q4t() {
4131 Some((mm, gi)) => (
4132 mm,
4133 gi,
4134 e.up_proj.mapped_q4t()?.1,
4135 e.down_proj.mapped_q4t()?.1,
4136 false,
4137 false,
4138 ),
4139 None => match e.gate_proj.mapped_q2tp() {
4140 Some((mm, gi)) => (
4141 mm,
4142 gi,
4143 e.up_proj.mapped_q2tp()?.1,
4144 e.down_proj.mapped_q4tp()?.1,
4145 true,
4146 true,
4147 ),
4148 None => {
4149 let (mm, gi) = e.gate_proj.mapped_q4tp()?;
4150 (
4151 mm,
4152 gi,
4153 e.up_proj.mapped_q4tp()?.1,
4154 e.down_proj.mapped_q4tp()?.1,
4155 true,
4156 false,
4157 )
4158 }
4159 },
4160 };
4161 if *q4tp.get_or_insert(is_p) != is_p || *gu_q2.get_or_insert(is_q2) != is_q2
4162 {
4163 tracing::warn!(
4169 "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."
4170 );
4171 return None;
4172 }
4173 model.get_or_insert_with(|| mm.clone());
4174 experts.push((gi, ui, di));
4175 }
4176 crate::gpu::GraphFfn::Moe {
4177 router,
4178 shared_gate: sgate,
4179 experts,
4180 n_exp: m.experts.len(),
4181 top_k: std::env::var("CMF_TOPK_PROBE")
4187 .ok()
4188 .and_then(|v| v.parse::<usize>().ok())
4189 .filter(|k| *k > 0 && *k <= m.top_k)
4190 .unwrap_or(m.top_k),
4191 inter,
4192 norm_topk: m.norm_topk_prob,
4193 q4tp: q4tp?,
4194 gu_q2: gu_q2.unwrap_or(false),
4195 }
4196 }
4197 };
4198 let attn = match &lw.attn {
4199 AttnKind::Full {
4200 wq,
4201 wk,
4202 wv,
4203 wo,
4204 q_norm,
4205 k_norm,
4206 output_gate,
4207 softplus_gate,
4208 bias,
4209 } => {
4210 if softplus_gate.is_some() || self.attention_heads_per_layer.is_some() {
4211 return None;
4212 }
4213 let (m, _, _, _) = wq.graph_weight()?;
4214 model = Some(m.clone());
4215 crate::gpu::GraphAttn::Full {
4216 wq: gw(wq)?,
4217 wk: gw(wk)?,
4218 wv: gw(wv)?,
4219 wo: gw(wo)?,
4220 q_norm: q_norm.as_deref(),
4221 k_norm: k_norm.as_deref(),
4222 bias: bias
4223 .as_ref()
4224 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
4225 output_gate: *output_gate,
4226 cpu_k: self.kv_cache.layers[li].k_heads(),
4227 cpu_v: self.kv_cache.layers[li].v_heads(),
4228 }
4229 }
4230 AttnKind::LinearGdn(w) => {
4231 let cfg = self.gdn_cfg?;
4232 let (m, _, _, _) = w.in_proj_qkv.graph_weight()?;
4233 model = Some(m.clone());
4234 crate::gpu::GraphAttn::Gdn {
4235 qkv: gw(&w.in_proj_qkv)?,
4236 z: gw(&w.in_proj_z)?,
4237 a: gw(&w.in_proj_a)?,
4238 b: gw(&w.in_proj_b)?,
4239 out: gw(&w.out_proj)?,
4240 conv1d: &w.conv1d,
4241 a_log: &w.a_log,
4242 dt_bias: &w.dt_bias,
4243 norm: &w.norm,
4244 nv: cfg.num_v_heads,
4245 nk: cfg.num_k_heads,
4246 dk: cfg.key_head_dim,
4247 dv: cfg.value_head_dim,
4248 kk: cfg.conv_kernel,
4249 cpu_state: &self.kv_cache.layers[self.phys_layer(li)].linear_state,
4250 }
4251 }
4252 _ => return None,
4253 };
4254 layers.push(crate::gpu::GraphLayer {
4255 input_norm: &lw.input_norm,
4256 attn,
4257 post_norm: &lw.post_norm,
4258 ffn: gffn,
4259 });
4260 }
4261 let model = model?;
4262 let lm_gw = if self.graph_want_logits
4268 && std::env::var("CMF_GPU_LMHEAD")
4269 .map(|v| v != "0")
4270 .unwrap_or(true)
4271 {
4272 self.weights.lm_head.graph_weight().map(|(_, i, kind, rs)| {
4273 (
4274 crate::gpu::GraphW {
4275 idx: i,
4276 kind,
4277 row_scale: rs,
4278 data: &[],
4279 },
4280 self.weights.lm_head.rows(),
4281 )
4282 })
4283 } else {
4284 None
4285 };
4286 let lm = lm_gw.as_ref().map(|(gw, rows)| (gw, *rows));
4287 let emb_gw = if steps > 1 {
4289 self.weights
4290 .embed_tokens
4291 .graph_weight()
4292 .map(|(_, i, kind, rs)| {
4293 (
4294 crate::gpu::GraphW {
4295 idx: i,
4296 kind,
4297 row_scale: rs,
4298 data: &[],
4299 },
4300 self.weights.embed_tokens.rows(),
4301 self.embed_multiplier as f32,
4302 )
4303 })
4304 } else {
4305 None
4306 };
4307
4308 let loop_norm_at: Vec<usize> = if self.loop_final_norm {
4311 (0..self.num_layers - 1)
4312 .filter(|&li| (li + 1) % self.physical_layers == 0)
4313 .collect()
4314 } else {
4315 Vec::new()
4316 };
4317 let mut h = hidden.to_vec();
4318 crate::gpu::forward_token_graph(
4319 &model,
4320 self.graph_kv_id,
4321 &layers,
4322 &o1_views,
4323 self.o1_epoch,
4324 &self.inv_freq,
4325 &mut h,
4326 nh,
4327 nkv,
4328 hd,
4329 rd,
4330 self.hidden_size,
4331 self.intermediate_size,
4332 position,
4333 self.kv_cache.max_seq_len,
4334 gemma,
4335 self.rms_eps as f32,
4336 lm,
4337 &self.weights.final_norm,
4338 logits_out,
4339 &loop_norm_at,
4340 steps,
4341 emb_gw.as_ref().map(|(gw, rows, m)| (gw, *rows, *m)),
4342 ids_out,
4343 layers_run,
4344 )
4345 .then_some(h)
4346 }
4347
4348 fn try_batch_graph_wgpu(
4353 &self,
4354 hiddens: &mut [f32],
4355 positions: &[usize],
4356 k: usize,
4357 spec: Option<crate::gpu::SpecTail<'_>>,
4358 ) -> bool {
4359 let _tb = std::time::Instant::now();
4360 if self.attn_softcap > 0.0 {
4361 return false; }
4363 if self.o1_active() {
4364 return false;
4365 }
4366 let nh = self.num_heads;
4367 let (nkv, hd, rd) = self.layer_geom(0);
4368 let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
4369 fn gw(t: &QTensor) -> Option<crate::gpu::GraphW<'_>> {
4370 if let Some((_, i, kind, rs)) = t.graph_weight() {
4371 return Some(crate::gpu::GraphW {
4372 idx: i,
4373 kind,
4374 row_scale: rs,
4375 data: &[],
4376 });
4377 }
4378 t.as_f32().map(|d| crate::gpu::GraphW {
4379 idx: 0,
4380 kind: 4,
4381 row_scale: &[],
4382 data: d,
4383 })
4384 }
4385 let built: Option<(
4386 Vec<crate::gpu::GraphLayer<'_>>,
4387 std::sync::Arc<cortiq_core::CmfModel>,
4388 )> = (|| {
4389 let mut layers = Vec::with_capacity(self.num_layers);
4390 let mut model = None;
4391 for li in 0..self.num_layers {
4392 let lw = &self.weights.layers[self.phys_layer(li)];
4393 let gffn = match &lw.ffn {
4400 FfnKind::Dense(d) => crate::gpu::GraphFfn::Dense {
4401 gate: gw(&d.gate_proj)?,
4402 up: gw(&d.up_proj)?,
4403 down: gw(&d.down_proj)?,
4404 },
4405 FfnKind::Moe(m) => {
4406 if m.router_sigmoid
4407 || m.expert_bias.is_some()
4408 || m.route_tau.is_some()
4409 || m.mask.is_some()
4410 {
4411 return None;
4412 }
4413 let (se, sg) = m.shared.as_ref()?;
4414 let sgate = gw(sg.as_ref()?)?;
4415 let router = gw(&m.router)?;
4416 let inter = m.experts.first()?.gate_proj.rows();
4417 let mut experts = Vec::with_capacity(m.experts.len() + 1);
4418 let mut q4tp: Option<bool> = None;
4419 for e in m.experts.iter().chain(std::iter::once(se)) {
4420 if !matches!(e.act, Act::Silu)
4421 || e.gate_proj.rows() != inter
4422 || e.up_proj.rows() != inter
4423 {
4424 return None;
4425 }
4426 let (mm, gi, ui, di, is_p) = match e.gate_proj.mapped_q4t() {
4427 Some((mm, gi)) => (
4428 mm,
4429 gi,
4430 e.up_proj.mapped_q4t()?.1,
4431 e.down_proj.mapped_q4t()?.1,
4432 false,
4433 ),
4434 None => {
4435 let (mm, gi) = e.gate_proj.mapped_q4tp()?;
4436 (
4437 mm,
4438 gi,
4439 e.up_proj.mapped_q4tp()?.1,
4440 e.down_proj.mapped_q4tp()?.1,
4441 true,
4442 )
4443 }
4444 };
4445 if *q4tp.get_or_insert(is_p) != is_p {
4446 return None;
4447 }
4448 model.get_or_insert_with(|| mm.clone());
4449 experts.push((gi, ui, di));
4450 }
4451 crate::gpu::GraphFfn::Moe {
4452 router,
4453 shared_gate: sgate,
4454 experts,
4455 n_exp: m.experts.len(),
4456 top_k: m.top_k,
4457 inter,
4458 norm_topk: m.norm_topk_prob,
4459 q4tp: q4tp?,
4460 gu_q2: false,
4463 }
4464 }
4465 _ => return None,
4466 };
4467 let attn = match &lw.attn {
4468 AttnKind::Full {
4469 wq,
4470 wk,
4471 wv,
4472 wo,
4473 q_norm,
4474 k_norm,
4475 output_gate,
4476 softplus_gate,
4477 bias,
4478 } => {
4479 if softplus_gate.is_some() || self.attention_heads_per_layer.is_some() {
4480 return None;
4481 }
4482 let (m, _, _, _) = wq.graph_weight()?;
4483 model = Some(m.clone());
4484 crate::gpu::GraphAttn::Full {
4485 wq: gw(wq)?,
4486 wk: gw(wk)?,
4487 wv: gw(wv)?,
4488 wo: gw(wo)?,
4489 q_norm: q_norm.as_deref(),
4490 k_norm: k_norm.as_deref(),
4491 bias: bias
4492 .as_ref()
4493 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
4494 output_gate: *output_gate,
4495 cpu_k: self.kv_cache.layers[li].k_heads(),
4496 cpu_v: self.kv_cache.layers[li].v_heads(),
4497 }
4498 }
4499 AttnKind::LinearGdn(w) => {
4500 let cfg = self.gdn_cfg?;
4501 let (m, _, _, _) = w.in_proj_qkv.graph_weight()?;
4502 model = Some(m.clone());
4503 crate::gpu::GraphAttn::Gdn {
4504 qkv: gw(&w.in_proj_qkv)?,
4505 z: gw(&w.in_proj_z)?,
4506 a: gw(&w.in_proj_a)?,
4507 b: gw(&w.in_proj_b)?,
4508 out: gw(&w.out_proj)?,
4509 conv1d: &w.conv1d,
4510 a_log: &w.a_log,
4511 dt_bias: &w.dt_bias,
4512 norm: &w.norm,
4513 nv: cfg.num_v_heads,
4514 nk: cfg.num_k_heads,
4515 dk: cfg.key_head_dim,
4516 dv: cfg.value_head_dim,
4517 kk: cfg.conv_kernel,
4518 cpu_state: &self.kv_cache.layers[self.phys_layer(li)].linear_state,
4519 }
4520 }
4521 _ => return None,
4522 };
4523 layers.push(crate::gpu::GraphLayer {
4524 input_norm: &lw.input_norm,
4525 attn,
4526 post_norm: &lw.post_norm,
4527 ffn: gffn,
4528 });
4529 }
4530 Some((layers, model?))
4531 })();
4532 let Some((layers, model)) = built else {
4533 {
4534 use std::sync::atomic::{AtomicBool, Ordering};
4535 static SAID: AtomicBool = AtomicBool::new(false);
4536 if !SAID.swap(true, Ordering::Relaxed) {
4537 tracing::warn!("batch graph: BUILDER refused (layer weights/kinds)");
4538 }
4539 }
4540 return false;
4541 };
4542 if std::env::var("CMF_GRAPH_SPEC_TIME").is_ok() {
4543 eprintln!("batch-build: {:.1} ms", _tb.elapsed().as_secs_f64() * 1e3);
4544 }
4545 crate::gpu::forward_batch_graph(
4546 &model,
4547 self.graph_kv_id,
4548 &layers,
4549 &self.inv_freq,
4550 hiddens,
4551 nh,
4552 nkv,
4553 hd,
4554 rd,
4555 self.hidden_size,
4556 self.intermediate_size,
4557 positions,
4558 self.kv_cache.max_seq_len,
4559 gemma,
4560 self.rms_eps as f32,
4561 k,
4562 spec,
4563 )
4564 }
4565
4566 fn draft_probe() -> bool {
4570 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
4571 *ON.get_or_init(|| std::env::var("CMF_DSV4_DRAFT_PROBE").is_ok_and(|v| v != "0"))
4572}
4573
4574 #[cfg(feature = "gpu")]
4586 fn dsv4_spec_on() -> bool {
4587 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
4588 *ON.get_or_init(|| std::env::var("CMF_DSV4_SPEC").map(|v| v != "0").unwrap_or(true))
4589 }
4590
4591 #[cfg(feature = "gpu")]
4598 fn dsv4_spec_step(
4599 &mut self,
4600 tip_token: u32,
4601 t_next: u32,
4602 next_pos: usize,
4603 drafted: &mut usize,
4604 accepted_ctr: &mut usize,
4605 ) -> Option<(Vec<u32>, usize)> {
4606 let t_all = std::time::Instant::now();
4607 if std::env::var("CMF_DSV4_SPEC_TIME").is_ok() {
4608 thread_local! {
4609 static LAST: std::cell::Cell<Option<std::time::Instant>> =
4610 const { std::cell::Cell::new(None) };
4611 }
4612 LAST.with(|l| {
4613 if let Some(prev) = l.get() {
4614 eprintln!("между раундами {:.1} мс", prev.elapsed().as_secs_f64() * 1e3);
4615 }
4616 l.set(Some(std::time::Instant::now()));
4617 });
4618 }
4619 if std::env::var("CMF_DSV4_SPEC_DEBUG").is_ok() {
4620 eprintln!("spec_step: вход pos={next_pos}");
4621 }
4622 let n_layers = self.dsv4.as_ref().map(|b| b.1.len())?;
4623 let cfg = self.dsv4.as_ref().map(|b| b.2)?;
4624 if self.dspark.is_none() {
4626 let t = crate::dsv4::dspark_targets(&self.dsv4_mtp, &cfg, n_layers);
4627 if t.is_empty() {
4628 return None;
4629 }
4630 crate::dsv4::dspark_arm(&t, cfg.dim);
4631 self.dspark = Some(crate::dsv4::DsparkState::new(
4632 self.dsv4_mtp.len(),
4633 &cfg,
4634 t.len(),
4635 ));
4636 }
4637 let targets = crate::dsv4::dspark_targets(&self.dsv4_mtp, &cfg, n_layers);
4638 let pack = crate::dsv4::dspark_pack_get(&self.dsv4_mtp, &cfg);
4639 if pack.is_none() && std::env::var("CMF_DSV4_SPEC_DEBUG").is_ok() {
4640 eprintln!("spec_step: пак не построился (targets {targets:?})");
4641 }
4642 let pack = pack?;
4643 let block = crate::dsv4::dspark_block();
4644 let b_box = self.dsv4.as_mut()?;
4645 let (g, layers, st) = (&b_box.0, &b_box.1, &mut b_box.3);
4646 let ds = self.dspark.as_mut()?;
4647 let dbg = std::env::var("CMF_DSV4_SPEC_DEBUG").is_ok();
4650 if !crate::dsv4::dspark_take(&mut ds.main_hidden) && !ds.have_hidden {
4651 if dbg {
4652 eprintln!("spec_step: нет захвата");
4653 }
4654 return None;
4655 }
4656 ds.have_hidden = true;
4657 let tip_pos = next_pos.checked_sub(1)?;
4658 let draft_started = std::time::Instant::now();
4659 let mut conf = Vec::new();
4660 let props = crate::dsv4::dspark_draft_gpu(
4661 g,
4662 &self.dsv4_mtp,
4663 &cfg,
4664 ds,
4665 pack,
4666 st.kv_id,
4667 tip_token,
4668 tip_pos,
4669 self.pool.as_deref(),
4670 &mut conf,
4671 );
4672 self.dspark_draft_ns += draft_started.elapsed().as_nanos();
4673 *drafted += block;
4674 if props.is_empty() || props[0] != t_next {
4675 if dbg {
4676 eprintln!(
4677 "spec_step: черновик {} (props0={:?} t_next={t_next})",
4678 if props.is_empty() { "пуст" } else { "мимо" },
4679 props.first()
4680 );
4681 }
4682 return None;
4683 }
4684 let mut k_verify = crate::dsv4::dspark_verify_k().min(props.len());
4685 let conf_min = {
4691 static M: std::sync::OnceLock<f32> = std::sync::OnceLock::new();
4692 *M.get_or_init(|| {
4693 std::env::var("CMF_DSPARK_CONF_MIN")
4694 .ok()
4695 .and_then(|v| v.parse().ok())
4696 .unwrap_or(0.0)
4697 })
4698 };
4699 if conf_min > 0.0 && conf.len() >= props.len() {
4700 let mut keep = 1usize;
4701 while keep < k_verify && conf.get(keep).copied().unwrap_or(0.0) >= conf_min {
4702 keep += 1;
4703 }
4704 k_verify = k_verify.min(keep.max(2));
4705 }
4706 if k_verify < 2 {
4707 return None;
4708 }
4709 let mut fed = Vec::with_capacity(k_verify);
4710 fed.push(t_next);
4711 fed.extend_from_slice(&props[1..k_verify]);
4712 let mut argmax = Vec::new();
4713 let mut logits_all = Vec::new();
4714 let mut walked = Vec::new();
4715 let txn = crate::dsv4::dsv4_verify_chunk(
4716 g,
4717 layers,
4718 &cfg,
4719 st,
4720 &fed,
4721 next_pos,
4722 &self.inv_freq,
4723 self.pool.as_deref(),
4724 &targets,
4725 &mut argmax,
4726 &mut logits_all,
4727 &mut walked,
4728 );
4729 if txn.is_none() && dbg {
4730 eprintln!("spec_step: verify отказал");
4731 }
4732 let txn = txn?;
4733 let b = fed.len();
4734 let mut accepted = 1usize;
4735 while accepted < b && fed[accepted] == argmax[accepted - 1] {
4736 accepted += 1;
4737 }
4738 if std::env::var("CMF_DSV4_SPEC_FORCE_REJECT").is_ok_and(|v| v != "0") {
4743 accepted = 1;
4744 }
4745 if std::env::var("CMF_DSV4_SPEC_TRACE").is_ok() {
4746 eprintln!(
4747 "spec@{next_pos}: fed={fed:?} argmax={argmax:?} accepted={accepted}"
4748 );
4749 }
4750 let t_fin = std::time::Instant::now();
4751 if !crate::dsv4::dsv4_spec_finish(
4752 g,
4753 layers,
4754 &cfg,
4755 st,
4756 txn,
4757 accepted,
4758 &fed,
4759 &self.inv_freq,
4760 self.pool.as_deref(),
4761 ) {
4762 tracing::warn!("dsv4: спекулятивный откат не удался — состояние подозрительно");
4763 return None;
4764 }
4765 if std::env::var("CMF_DSV4_SPEC_TIME").is_ok() {
4766 eprintln!("finish(k={accepted}): {:.1} мс", t_fin.elapsed().as_secs_f64() * 1e3);
4767 }
4768 *accepted_ctr += accepted - 1;
4769 let (hc, dim) = (cfg.hc_mult, cfg.dim);
4774 let dev_caps: Vec<usize> = targets
4781 .iter()
4782 .copied()
4783 .filter(|&t| {
4784 st.dev_set.get(t).copied().unwrap_or(false)
4785 && !st.partial_set.get(t).copied().unwrap_or(false)
4786 })
4787 .collect();
4788 let mut caps_all = vec![0.0f32; dev_caps.len() * b * hc * dim];
4789 if !crate::gpu_wgpu::dsv4_spec_cap_read_all(b, dev_caps.len(), hc * dim, &mut caps_all) {
4790 return None;
4791 }
4792 for t in 0..accepted {
4793 let tip = t + 1 == accepted;
4794 for (slot, &tl) in targets.iter().enumerate() {
4795 if let Some(di) = dev_caps.iter().position(|&d| d == tl) {
4796 let lo = (di * b + t) * hc * dim;
4797 crate::dsv4::dspark_capture(
4798 &caps_all[lo..lo + hc * dim],
4799 &cfg,
4800 slot,
4801 &mut ds.main_hidden,
4802 );
4803 } else if tip
4804 && crate::dsv4::dspark_peek_slot(slot, dim, {
4805 let lo = slot * dim;
4806 &mut ds.main_hidden[lo..lo + dim]
4807 })
4808 {
4809 } else {
4814 crate::dsv4::dspark_capture(
4818 &walked[t * hc * dim..(t + 1) * hc * dim],
4819 &cfg,
4820 slot,
4821 &mut ds.main_hidden,
4822 );
4823 }
4824 }
4825 crate::dsv4::dspark_ring_append(g, &self.dsv4_mtp, &cfg, ds, next_pos + t, self.pool.as_deref());
4826 }
4827 let row = logits_all[(accepted - 1) * cfg.vocab..accepted * cfg.vocab].to_vec();
4828 self.graph_logits = Some(row);
4829 if std::env::var("CMF_DSV4_TRUNK_PICK_DUMP").is_ok() {
4834 crate::dsv4::trunk_freq_note(&crate::dsv4::pick_tally_take());
4835 crate::dsv4::pick_tally_arm();
4836 }
4837 if std::env::var("CMF_DSV4_SPEC_TIME").is_ok() {
4838 eprintln!("spec_step total {:.1} мс (k={accepted})", t_all.elapsed().as_secs_f64() * 1e3);
4839 }
4840 Some((fed[1..accepted].to_vec(), next_pos + accepted))
4841 }
4842
4843 fn dspark_probe(&mut self, position: usize, token_id: u32) {
4844 if self.dsv4_mtp.is_empty() || !Self::draft_probe() {
4845 return;
4846 }
4847 let trunk_now = crate::dsv4::pick_tally_take();
4849 crate::dsv4::trunk_freq_note(&trunk_now);
4850 if !trunk_now.is_empty() {
4851 self.dspark_trunk_picks.push(trunk_now);
4852 let keep = crate::dsv4::dspark_block();
4853 if self.dspark_trunk_picks.len() > keep {
4854 self.dspark_trunk_picks.remove(0);
4855 }
4856 }
4857 for p in std::mem::take(&mut self.dspark_pending) {
4860 let Some(i) = position.checked_sub(p.0 + 1) else {
4861 continue;
4862 };
4863 let mut p = p;
4864 if i < p.1.len() {
4865 if p.2 && p.1[i] == token_id {
4866 p.3 = i + 1;
4867 } else {
4868 p.2 = false;
4869 }
4870 if i + 1 < p.1.len() {
4871 self.dspark_pending.push(p);
4872 continue;
4873 }
4874 }
4875 self.dspark_hist.push(p.3);
4876 self.dspark_real.push(token_id);
4877 }
4878 let Some(b) = &mut self.dsv4 else { return };
4879 let (g, layers, cfg) = (&b.0, &b.1, b.2);
4880 let n_layers = layers.len();
4881 if self.dspark.is_none() {
4882 let t = crate::dsv4::dspark_targets(&self.dsv4_mtp, &cfg, n_layers);
4883 if t.is_empty() {
4884 return;
4885 }
4886 eprintln!("DSpark: захват со слоёв {t:?}, блок {}", crate::dsv4::dspark_block());
4887 crate::dsv4::dspark_arm(&t, cfg.dim);
4888 self.dspark = Some(crate::dsv4::DsparkState::new(
4889 self.dsv4_mtp.len(),
4890 &cfg,
4891 t.len(),
4892 ));
4893 }
4894 let ds = self.dspark.as_mut().unwrap();
4895 if !crate::dsv4::dspark_take(&mut ds.main_hidden) {
4896 return; }
4898 let mut conf = Vec::new();
4899 crate::dsv4::pick_tally_arm();
4900 let draft_started = std::time::Instant::now();
4905 #[cfg(feature = "gpu")]
4906 let gpu_draft = crate::dsv4::dspark_gpu_on();
4907 #[cfg(not(feature = "gpu"))]
4908 let gpu_draft = false;
4909 let props = if gpu_draft {
4910 #[cfg(feature = "gpu")]
4911 {
4912 let kv_id = b.3.kv_id;
4913 match crate::dsv4::dspark_pack_get(&self.dsv4_mtp, &cfg) {
4914 Some(pk) => crate::dsv4::dspark_draft_gpu(
4915 g,
4916 &self.dsv4_mtp,
4917 &cfg,
4918 ds,
4919 pk,
4920 kv_id,
4921 token_id,
4922 position,
4923 self.pool.as_deref(),
4924 &mut conf,
4925 ),
4926 None => Vec::new(),
4927 }
4928 }
4929 #[cfg(not(feature = "gpu"))]
4930 Vec::new()
4931 } else {
4932 crate::gpu::cpu_scope(|| {
4933 crate::dsv4::dspark_draft(
4934 g,
4935 &self.dsv4_mtp,
4936 &cfg,
4937 ds,
4938 token_id,
4939 position,
4940 self.pool.as_deref(),
4941 &mut conf,
4942 )
4943 })
4944 };
4945 self.dspark_draft_ns += draft_started.elapsed().as_nanos();
4946 let draft_picks = crate::dsv4::pick_tally_take();
4947 crate::dsv4::dspark_freq_note(&draft_picks);
4948 crate::dsv4::pick_tally_arm();
4951 if !props.is_empty() {
4952 let (tu, tt) = {
4956 let flat: Vec<(usize, Vec<usize>)> = self
4957 .dspark_trunk_picks
4958 .iter()
4959 .flat_map(|v| v.iter().cloned())
4960 .collect();
4961 let mut per: std::collections::HashMap<usize, Vec<usize>> =
4963 std::collections::HashMap::new();
4964 for (li, picks) in flat {
4965 per.entry(li).or_default().extend(picks);
4966 }
4967 let n = per.len().max(1);
4968 let mut u = 0usize;
4969 let mut t = 0usize;
4970 for (_, v) in per {
4971 t += v.len();
4972 u += v.iter().collect::<std::collections::HashSet<_>>().len();
4973 }
4974 (u / n, t / n)
4975 };
4976 let (du, dt) = crate::dsv4::tally_unique(&draft_picks);
4977 self.dspark_exp.push((tu, tt, du, dt));
4978 self.dspark_pending.push((position, props, true, 0));
4979 }
4980 if self.dspark_hist.len() >= 8 && self.dspark_hist.len() % 8 == 0 {
4981 let n = self.dspark_hist.len() as f32;
4982 let mean: f32 = self.dspark_hist.iter().sum::<usize>() as f32 / n;
4983 let block = crate::dsv4::dspark_block();
4984 let mut at = vec![0usize; block + 1];
4985 for &k in &self.dspark_hist {
4986 at[k] += 1;
4987 }
4988 let mut surv = Vec::with_capacity(block);
4990 for i in 1..=block {
4991 let k = at[i..].iter().sum::<usize>() as f32 / n;
4992 surv.push(format!("{k:.2}"));
4993 }
4994 let distinct = self
4995 .dspark_real
4996 .iter()
4997 .collect::<std::collections::HashSet<_>>()
4998 .len();
4999 let (tu, tt, du, dt) = self.dspark_exp.iter().fold((0, 0, 0, 0), |a, b| {
5000 (a.0 + b.0, a.1 + b.1, a.2 + b.2, a.3 + b.3)
5001 });
5002 let m = self.dspark_exp.len().max(1);
5003 eprintln!(
5004 "DSpark: черновиков {}, принято в среднем {mean:.2} из {block} \
5005 (токенов за проход {:.2}), распределение {at:?}, выживание [{}]",
5006 self.dspark_hist.len(),
5007 mean + 1.0,
5008 surv.join(" ")
5009 );
5010 eprintln!(
5011 "DSpark: разных токенов {distinct} из {} (вырожденность), \
5012 эксперты ствол {}/{} на слой за {block} токенов, \
5013 черновик {}/{} за блок, draft {:.2} мс/блок",
5014 self.dspark_real.len(),
5015 tu / m,
5016 tt / m,
5017 du / m,
5018 dt / m,
5019 self.dspark_draft_ns as f64 / self.dspark_exp.len().max(1) as f64 / 1e6
5020 );
5021 }
5022 }
5023
5024 fn forward_layers_upto(
5025 &mut self,
5026 hidden: &[f32],
5027 position: usize,
5028 task_mask: Option<&TaskMask>,
5029 upto: Option<usize>,
5030 ) -> Vec<f32> {
5031 if let Some(b) = &mut self.dsv4 {
5037 let _ = (task_mask, upto);
5038 let token_id = hidden.first().copied().unwrap_or(0.0) as u32;
5039 let (g, layers, cfg, st) = (&b.0, &b.1, b.2, &mut b.3);
5040 st.pos = position;
5041 let mut logits = Vec::new();
5042 crate::dsv4::forward_token(
5043 g,
5044 layers,
5045 &cfg,
5046 st,
5047 token_id,
5048 &self.inv_freq,
5049 self.pool.as_deref(),
5050 &mut logits,
5051 );
5052 self.graph_logits = Some(logits);
5053 self.dspark_probe(position, token_id);
5054 return vec![0.0; self.hidden_size];
5057 }
5058 if let Some(b) = &self.g3n {
5061 let _ = (task_mask, upto);
5062 return crate::g3n::g3n_forward(
5063 &b.0,
5064 &b.1,
5065 hidden,
5066 position,
5067 &mut self.kv_cache.layers,
5068 self.num_heads,
5069 self.num_kv_heads,
5070 self.head_dim,
5071 self.pool.as_deref(),
5072 );
5073 }
5074 let mut h = hidden.to_vec();
5075 let (nh, _nkv, _hd, hs, _rd, eps) = (
5078 self.num_heads,
5079 self.num_kv_heads,
5080 self.head_dim,
5081 self.hidden_size,
5082 self.rotary_dim,
5083 self.rms_eps,
5084 );
5085 let pool = self.pool.clone();
5086 let graph_env = std::env::var("CMF_GPU_WGPU_GRAPH").ok();
5098 let graph_on = match graph_env.as_deref() {
5099 Some("0") => false,
5100 Some(_) => true,
5101 None => crate::gpu::wgpu_graph_default(),
5107 };
5108 let graph_trusted =
5109 graph_env.is_some() || crate::gpu::wgpu_graph_default() || self.gdn_cfg.is_some();
5110 let race_eligible = graph_on && upto.is_none() && task_mask.is_none();
5111 let mut tail_start = 0usize;
5112 if race_eligible && crate::gpu::graph_race_use_graph(graph_trusted) {
5113 let t_graph = std::time::Instant::now();
5114 let mut lg = Vec::new();
5115 let mut gl = 0usize;
5116 let built = self.try_token_graph_wgpu(hidden, position, &mut lg, &mut gl);
5117 graph_note(built.is_some());
5118 if let Some(hh) = built {
5119 let dur = t_graph.elapsed();
5120 if std::env::var("CMF_GRAPH_PROF").is_ok() {
5121 eprintln!("graph-call: {:.2} ms total", dur.as_secs_f64() * 1000.0);
5122 }
5123 if gl > 0 && gl < self.num_layers {
5124 h = hh;
5130 tail_start = gl;
5131 } else if graph_trusted || !crate::gpu::graph_race_first_token_hopeless(dur) {
5132 if !graph_trusted {
5133 crate::gpu::graph_race_record(true, dur);
5134 }
5135 if !lg.is_empty() {
5136 lg.resize(self.vocab_size, 0.0);
5139 if let Some(c) = self.final_softcap {
5140 for l in lg.iter_mut() {
5141 *l = c * (*l / c).tanh();
5142 }
5143 }
5144 self.graph_logits = Some(lg);
5145 }
5146 return hh;
5147 }
5148 }
5154 }
5155 let t_race_cpu = (race_eligible && !graph_trusted).then(std::time::Instant::now);
5156
5157 #[cfg(target_os = "macos")]
5158 let mut gpu_skip_until = 0usize;
5159 for li in tail_start..self.num_layers {
5160 crate::gpu::set_layer(li as i64); if let Some(u) = upto {
5162 if li > u {
5163 break;
5164 }
5165 }
5166 if let Some(mask) = task_mask {
5167 if !mask.layer_alive(li) {
5168 continue; }
5170 }
5171 #[cfg(target_os = "macos")]
5175 {
5176 if li < gpu_skip_until {
5177 continue;
5178 }
5179 if task_mask.is_none() {
5180 let end = self.q1_graph_gpu(li, upto, position, &mut h);
5181 if end > li {
5182 gpu_skip_until = end;
5183 if self.is_loop_end(end - 1) && end < self.num_layers {
5186 h = inference::rms_norm(
5187 &h,
5188 &self.weights.final_norm,
5189 self.rms_eps,
5190 self.norm_style,
5191 );
5192 }
5193 continue;
5194 }
5195 }
5196 }
5197
5198 let lw = &self.weights.layers[self.phys_layer(li)];
5199 if let Ok(tp) = std::env::var("CMF_TRACE_POS") {
5200 if tp.parse::<usize>().ok() == Some(position) {
5201 let n: f32 = h.iter().map(|x| x * x).sum::<f32>().sqrt();
5202 eprintln!(
5203 "TRACE pos {position} layer {li}: |h| = {n:.6} h0 {:.6} h1 {:.6}",
5204 h[0], h[1]
5205 );
5206 }
5207 }
5208 inference::rms_norm_into(
5211 &h,
5212 &lw.input_norm,
5213 self.rms_eps,
5214 self.norm_style,
5215 &mut self.ws.n1,
5216 );
5217
5218 let attn_out = match &lw.attn {
5219 AttnKind::Mla(w) => {
5220 let inv_freq_l = self.layer_inv_freq(li);
5221 let rs = self.layer_rope_scale(li);
5222 let eps = self.rms_eps;
5223 let pool = self.pool.clone();
5224 mla_attention(
5225 w,
5226 &self.ws.n1,
5227 &mut self.kv_cache.layers[li],
5228 position,
5229 &inv_freq_l,
5230 rs,
5231 eps,
5232 pool.as_deref(),
5233 )
5234 }
5235 AttnKind::Linear(w) => {
5236 let cfg = self.vmf_cfg.expect("linear layer without vmf_cfg");
5237 vmf_phase_forward(
5238 &self.ws.n1,
5239 w,
5240 &cfg,
5241 &mut self.kv_cache.layers[li].linear_state,
5242 self.pool.as_deref(),
5243 )
5244 }
5245 AttnKind::Kda(w) => {
5246 let cfg = self.kda_cfg.expect("kda layer without kda_cfg");
5247 crate::linear_core::kda_forward(
5248 &self.ws.n1,
5249 w,
5250 &cfg,
5251 &mut self.kv_cache.layers[li].linear_state,
5252 self.pool.as_deref(),
5253 )
5254 }
5255 AttnKind::LinearGdn(w) => {
5256 let cfg = self.gdn_cfg.expect("gdn layer without gdn_cfg");
5257 gdn_forward(
5258 &self.ws.n1,
5259 w,
5260 &cfg,
5261 &mut self.kv_cache.layers[li].linear_state,
5262 self.pool.as_deref(),
5263 )
5264 }
5265 AttnKind::ShortConv(w) => {
5266 let cfg = self
5267 .short_conv_cfg
5268 .expect("short-conv layer without short_conv_cfg");
5269 short_conv_forward(
5270 &self.ws.n1,
5271 w,
5272 &cfg,
5273 &mut self.kv_cache.layers[li].linear_state,
5274 self.pool.as_deref(),
5275 )
5276 }
5277 AttnKind::Full {
5278 wq,
5279 wk,
5280 wv,
5281 wo,
5282 q_norm,
5283 k_norm,
5284 output_gate,
5285 softplus_gate,
5286 bias,
5287 } if self.kv_cache.layers[li].o1_sealed() => {
5288 let inv_freq_l = self.layer_inv_freq(li);
5291 let (nkv_l, hd_l, rd_l) = self.layer_geom(li);
5292 let cfg = QwenAttnCfg {
5293 num_heads: self.layer_num_heads(li),
5294 num_kv_heads: nkv_l,
5295 head_dim: hd_l,
5296 hidden_size: hs,
5297 position,
5298 inv_freq: &inv_freq_l,
5299 rotary_dim: rd_l,
5300 scale: self.attn_scale,
5301 softcap: self.attn_softcap,
5302 window: None,
5303 v_norm: self.attn_v_norm,
5304 q_norm: q_norm.as_deref(),
5305 k_norm: k_norm.as_deref(),
5306 output_gate: *output_gate,
5307 softplus_gate: softplus_gate
5308 .as_ref()
5309 .map(|(gate, per_head)| (gate, *per_head)),
5310 rope_scale: self.layer_rope_scale(li),
5311 bias: bias
5312 .as_ref()
5313 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
5314 rms_eps: eps,
5315 norm_style: self.norm_style,
5316 pool: pool.as_deref(),
5317 };
5318 attention::qwen_attention_nystrom(
5319 &self.ws.n1,
5320 wq,
5321 wk,
5322 wv,
5323 wo,
5324 &mut self.kv_cache.layers[li],
5325 &cfg,
5326 )
5327 }
5328 AttnKind::Full {
5329 wq,
5330 wk,
5331 wv,
5332 wo,
5333 q_norm,
5334 k_norm,
5335 output_gate,
5336 softplus_gate,
5337 bias,
5338 } => 'attn: {
5339 if graph_on
5342 && !*output_gate
5343 && softplus_gate.is_none()
5344 && self.attention_heads_per_layer.is_none()
5345 && bias.is_none()
5346 && task_mask.is_none()
5347 {
5348 let inv_freq_l = self.layer_inv_freq(li);
5349 let (nkv_l, hd_l, rd_l) = self.layer_geom(li);
5350 let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
5351 if let (Some((gm, qi)), Some((_, ki)), Some((_, vi)), Some((_, oi))) = (
5352 wq.mapped_q1(),
5353 wk.mapped_q1(),
5354 wv.mapped_q1(),
5355 wo.mapped_q1(),
5356 ) {
5357 let gm = gm.clone();
5358 let mut out = vec![0f32; hs];
5359 let cache = &self.kv_cache.layers[li];
5360 if crate::gpu::attn_dropin(
5361 &gm,
5362 self.graph_kv_id,
5363 li,
5364 &self.ws.n1,
5365 qi,
5366 ki,
5367 vi,
5368 oi,
5369 q_norm.as_deref(),
5370 k_norm.as_deref(),
5371 &inv_freq_l,
5372 nh,
5373 nkv_l,
5374 hd_l,
5375 rd_l,
5376 hs,
5377 position,
5378 self.kv_cache.max_seq_len,
5379 gemma,
5380 eps as f32,
5381 cache.k_heads(),
5382 cache.v_heads(),
5383 &mut out,
5384 ) {
5385 break 'attn out;
5386 }
5387 }
5388 }
5389 let masked = task_mask
5390 .map(|m| m.head_flags(li, self.num_heads).iter().any(|&a| !a))
5391 .unwrap_or(false);
5392 let f32_view = (wq.as_f32(), wk.as_f32(), wv.as_f32(), wo.as_f32());
5393 match (masked, f32_view) {
5394 (true, (Some(q), Some(k), Some(v), Some(o))) => {
5397 let active_heads = task_mask.unwrap().head_flags(li, self.num_heads);
5398 attention::multi_head_attention(
5399 &self.ws.n1,
5400 q,
5401 k,
5402 v,
5403 o,
5404 &mut self.kv_cache.layers[li],
5405 self.num_heads,
5406 self.num_kv_heads,
5407 self.head_dim,
5408 self.hidden_size,
5409 position,
5410 &active_heads,
5411 &self.inv_freq,
5412 )
5413 }
5414 (masked, _) => {
5415 if masked {
5416 tracing::warn!(
5417 "layer {li}: head mask on quantized weights not \
5418 supported yet — executing dense"
5419 );
5420 }
5421 let inv_freq_l = self.layer_inv_freq(li);
5422 let (nkv_l, hd_l, rd_l) = self.layer_geom(li);
5423 let cfg = QwenAttnCfg {
5424 num_heads: self.layer_num_heads(li),
5425 num_kv_heads: nkv_l,
5426 head_dim: hd_l,
5427 hidden_size: hs,
5428 position,
5429 inv_freq: &inv_freq_l,
5430 rotary_dim: rd_l,
5431 scale: self.attn_scale,
5432 softcap: self.attn_softcap,
5433 window: self.layer_window(li),
5434 v_norm: self.attn_v_norm,
5435 q_norm: q_norm.as_deref(),
5436 k_norm: k_norm.as_deref(),
5437 output_gate: *output_gate,
5438 softplus_gate: softplus_gate
5439 .as_ref()
5440 .map(|(gate, per_head)| (gate, *per_head)),
5441 rope_scale: self.layer_rope_scale(li),
5442 bias: bias
5443 .as_ref()
5444 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
5445 rms_eps: eps,
5446 norm_style: self.norm_style,
5447 pool: pool.as_deref(),
5448 };
5449 attention::qwen_attention(
5450 &self.ws.n1,
5451 wq,
5452 wk,
5453 wv,
5454 wo,
5455 &mut self.kv_cache.layers[li],
5456 &cfg,
5457 )
5458 }
5459 }
5460 }
5461 };
5462 let attn_out = match &self.weights.layers[self.phys_layer(li)].attn_out_norm {
5465 Some(w) => inference::rms_norm(&attn_out, w, self.rms_eps, self.norm_style),
5466 None => attn_out,
5467 };
5468 let lw = &self.weights.layers[self.phys_layer(li)];
5469 inference::add_rmsnorm_fused_into(
5470 &mut h,
5471 &attn_out,
5472 &lw.post_norm,
5473 self.rms_eps,
5474 self.norm_style,
5475 &mut self.ws.p1,
5476 );
5477 let mut attn_out = attn_out;
5478 attention::recycle_buf(&mut attn_out);
5479 let post_normed = &self.ws.p1;
5480
5481 let ffn_masked = task_mask
5482 .map(|m| m.ffn_active_count(li) < self.intermediate_size)
5483 .unwrap_or(false);
5484 let f32_ffn = match &lw.ffn {
5487 FfnKind::Dense(d) => (
5488 d.gate_proj.as_f32(),
5489 d.up_proj.as_f32(),
5490 d.down_proj.as_f32(),
5491 ),
5492 FfnKind::Moe(_) | FfnKind::DenseMoe(_) => (None, None, None),
5493 };
5494 let ffn_out = match (ffn_masked, f32_ffn) {
5495 (true, (Some(g), Some(u), Some(d))) => {
5496 let active = task_mask.unwrap().ffn_active_indices(li);
5497 inference::sparse_ffn_forward(
5498 post_normed,
5499 g,
5500 u,
5501 d,
5502 self.hidden_size,
5503 self.intermediate_size,
5504 &active,
5505 self.pool.as_deref(),
5506 )
5507 }
5508 (true, _) => match &lw.ffn {
5512 FfnKind::Dense(d) if d.down_proj.sparse_col_ok() => {
5513 let active = task_mask.unwrap().ffn_active_indices(li);
5514 sparse_ffn_quant(
5515 d,
5516 post_normed,
5517 &active,
5518 self.hidden_size,
5519 self.pool.as_deref(),
5520 )
5521 }
5522 FfnKind::Dense(d) => {
5527 let active = task_mask.unwrap().ffn_active_indices(li);
5528 let (gf, uf, df) = dequant_dense_f32(d);
5529 inference::sparse_ffn_forward(
5530 post_normed,
5531 &gf,
5532 &uf,
5533 &df,
5534 self.hidden_size,
5535 self.intermediate_size,
5536 &active,
5537 self.pool.as_deref(),
5538 )
5539 }
5540 FfnKind::Moe(m) => {
5541 let allowed = task_mask.and_then(|tm| tm.expert_flags(li, m.experts.len()));
5545 ffn_forward(
5546 &lw.ffn,
5547 post_normed,
5548 self.pool.as_deref(),
5549 allowed.as_deref(),
5550 )
5551 }
5552 FfnKind::DenseMoe(dm) => dense_moe_ffn(
5553 dm,
5554 post_normed,
5555 &h,
5556 self.rms_eps,
5557 self.norm_style,
5558 self.pool.as_deref(),
5559 ),
5560 },
5561 (false, _) => match &lw.ffn {
5562 FfnKind::DenseMoe(dm) => dense_moe_ffn(
5563 dm,
5564 post_normed,
5565 &h,
5566 self.rms_eps,
5567 self.norm_style,
5568 self.pool.as_deref(),
5569 ),
5570 _ => {
5571 let allowed = match (&lw.ffn, task_mask) {
5572 (FfnKind::Moe(m), Some(tm)) => tm.expert_flags(li, m.experts.len()),
5573 _ => None,
5574 };
5575 ffn_forward(
5576 &lw.ffn,
5577 post_normed,
5578 self.pool.as_deref(),
5579 allowed.as_deref(),
5580 )
5581 }
5582 },
5583 };
5584 let ffn_out = match &self.weights.layers[self.phys_layer(li)].ffn_out_norm {
5585 Some(w) => inference::rms_norm(&ffn_out, w, self.rms_eps, self.norm_style),
5586 None => ffn_out,
5587 };
5588 for (i, &f) in ffn_out.iter().enumerate() {
5589 h[i] += f;
5590 }
5591 let mut ffn_out = ffn_out;
5592 attention::recycle_buf(&mut ffn_out);
5593
5594 if let Some(sc) = self.weights.layers[self.phys_layer(li)].layer_scale {
5596 for v in h.iter_mut() {
5597 *v *= sc;
5598 }
5599 }
5600
5601 if self.is_loop_end(li) && li + 1 < self.num_layers {
5604 h = inference::rms_norm(
5605 &h,
5606 &self.weights.final_norm,
5607 self.rms_eps,
5608 self.norm_style,
5609 );
5610 }
5611
5612 if self.dyn_phi_layer == Some(li) {
5616 self.update_dyn_phi(&h);
5617 }
5618 }
5619 crate::gpu::set_layer(-1); if let Some(t) = t_race_cpu {
5621 crate::gpu::graph_race_record(false, t.elapsed());
5622 }
5623
5624 h
5625 }
5626
5627 fn update_dyn_phi(&mut self, h: &[f32]) {
5630 const A: f32 = 0.2;
5631 if self.dyn_phi_ema.len() != h.len() {
5632 self.dyn_phi_ema = vec![0.0; h.len()];
5633 self.dyn_phi_seen = 0;
5634 }
5635 if self.dyn_phi_seen == 0 {
5636 self.dyn_phi_ema.copy_from_slice(h);
5637 } else {
5638 for (e, &v) in self.dyn_phi_ema.iter_mut().zip(h) {
5639 *e = (1.0 - A) * *e + A * v;
5640 }
5641 }
5642 self.dyn_phi_seen += 1;
5643 }
5644
5645 pub fn dyn_phi(&self) -> &[f32] {
5647 &self.dyn_phi_ema
5648 }
5649
5650 pub fn set_dyn_phi_layer(&mut self, layer: Option<usize>) {
5652 self.dyn_phi_layer = layer;
5653 self.dyn_phi_ema.clear();
5654 self.dyn_phi_seen = 0;
5655 }
5656
5657 pub fn dynamic_skills(&self) -> Vec<(usize, String, usize)> {
5659 let Some(model) = &self.model else {
5660 return Vec::new();
5661 };
5662 model
5663 .header
5664 .skills
5665 .iter()
5666 .enumerate()
5667 .filter_map(|(i, sk)| {
5668 let ok = matches!(self.dyn_skill_layers.get(i), Some(Some(_)));
5669 let sel = sk.selection.as_ref()?;
5670 (ok).then(|| (i, sk.id.clone(), sel.phi_layer))
5671 })
5672 .collect()
5673 }
5674
5675 pub fn active_skill(&self) -> Option<usize> {
5677 self.dyn_active
5678 }
5679
5680 pub fn enable_dynamic_routing(&mut self) -> usize {
5685 use crate::swarm::{DynRouter, RoutableSkill};
5686 let Some(model) = self.model.clone() else {
5687 return 0;
5688 };
5689 if self.dyn_blend_loaded {
5692 tracing::warn!("dynamic routing unavailable on a blend-loaded pipeline");
5693 return 0;
5694 }
5695 if let Some(a) = self.dyn_active {
5699 if !matches!(self.dyn_skill_layers.get(a), Some(Some(_))) {
5700 tracing::warn!("loaded skill is not FFN-eligible — dynamic routing unavailable");
5701 return 0;
5702 }
5703 }
5704 let hidden = self.hidden_size;
5705 let mut skills = Vec::new();
5706 for (idx, id, _phi) in self.dynamic_skills() {
5707 if let Some(sel) = model.header.skills[idx].selection.as_ref() {
5708 if let Some(rs) = RoutableSkill::from_descriptor(idx, id, sel, hidden) {
5709 skills.push(rs);
5710 }
5711 }
5712 }
5713 if skills.is_empty() {
5714 return 0;
5715 }
5716 let phi = skills[0].phi_layer;
5718 if skills.iter().any(|s| s.phi_layer != phi) {
5719 tracing::warn!("routable skills disagree on phi_layer; using {phi}");
5720 }
5721 let n = skills.len();
5722 self.set_dyn_phi_layer(Some(phi));
5723 self.dyn_router = Some(DynRouter::new(skills));
5724 n
5725 }
5726
5727 pub fn route_switches(&self) -> Vec<(usize, Option<String>, Option<String>)> {
5729 self.dyn_router
5730 .as_ref()
5731 .map(|r| r.switches.clone())
5732 .unwrap_or_default()
5733 }
5734
5735 fn lm_head_forward(&self, hidden: &[f32]) -> Vec<f32> {
5738 let rows = self.weights.lm_head.rows();
5739 let mut logits = attention::take_buf(rows.min(self.vocab_size));
5740 self.weights
5741 .lm_head
5742 .matvec(hidden, &mut logits, self.pool.as_deref());
5743 logits.resize(self.vocab_size, 0.0);
5744 if let Some(m) = self.logit_multiplier {
5745 for l in logits.iter_mut() {
5746 *l *= m;
5747 }
5748 }
5749 if let Some(c) = self.final_softcap {
5750 for l in logits.iter_mut() {
5751 *l = c * (*l / c).tanh();
5752 }
5753 }
5754 logits
5755 }
5756
5757 pub fn prefill_next_logits(&mut self, ids: &[u32], task_mask: Option<&TaskMask>) -> Vec<f32> {
5762 self.kv_cache.clear();
5763 self.kv_history.clear();
5764 let mut hidden = vec![0.0f32; self.hidden_size];
5765 for (pos, &id) in ids.iter().enumerate() {
5766 let emb = self.embed_single(id);
5767 hidden = self.forward_layers(&emb, pos, task_mask);
5768 }
5769 inference::rms_norm_into(
5770 &hidden,
5771 &self.weights.final_norm,
5772 self.rms_eps,
5773 self.norm_style,
5774 &mut self.ws.n1,
5775 );
5776 self.lm_head_forward(&self.ws.n1)
5777 }
5778}
5779
5780pub fn create_test_pipeline(
5782 hidden_size: usize,
5783 intermediate_size: usize,
5784 num_heads: usize,
5785 num_kv_heads: usize,
5786 head_dim: usize,
5787 num_layers: usize,
5788 vocab_size: usize,
5789) -> Pipeline {
5790 let synth = |n: usize, salt: usize| -> Vec<f32> {
5793 (0..n)
5794 .map(|i| (((i * 31 + salt * 17 + 7) % 97) as f32 / 97.0 - 0.5) * 0.2)
5795 .collect()
5796 };
5797 let qt = |rows: usize, cols: usize, salt: usize| -> QTensor {
5798 QTensor::from_f32(synth(rows * cols, salt), rows, cols)
5799 };
5800 let layer_weights: Vec<LayerWeights> = (0..num_layers)
5801 .map(|li| LayerWeights {
5802 input_norm: vec![1.0; hidden_size],
5803 post_norm: vec![1.0; hidden_size],
5804 attn_out_norm: None,
5805 ffn_out_norm: None,
5806 layer_scale: None,
5807 ffn: FfnKind::Dense(DenseFfn {
5808 gate_proj: qt(intermediate_size, hidden_size, li * 10 + 5),
5809 up_proj: qt(intermediate_size, hidden_size, li * 10 + 6),
5810 down_proj: qt(hidden_size, intermediate_size, li * 10 + 7),
5811 act: Act::Silu,
5812 }),
5813 attn: AttnKind::Full {
5814 bias: None,
5815 wq: qt(num_heads * head_dim, hidden_size, li * 10 + 1),
5816 wk: qt(num_kv_heads * head_dim, hidden_size, li * 10 + 2),
5817 wv: qt(num_kv_heads * head_dim, hidden_size, li * 10 + 3),
5818 wo: qt(hidden_size, num_heads * head_dim, li * 10 + 4),
5819 q_norm: None,
5820 k_norm: None,
5821 output_gate: false,
5822 softplus_gate: None,
5823 },
5824 })
5825 .collect();
5826
5827 Pipeline::new(
5828 Tokenizer::byte_level(),
5829 PipelineWeights {
5830 embed_tokens: qt(vocab_size, hidden_size, 100),
5831 layers: layer_weights,
5832 lm_head: qt(vocab_size, hidden_size, 200),
5833 final_norm: vec![1.0; hidden_size],
5834 },
5835 hidden_size,
5836 intermediate_size,
5837 num_heads,
5838 num_kv_heads,
5839 head_dim,
5840 num_layers,
5841 num_layers, false, vocab_size,
5844 1e-6,
5845 10_000.0,
5846 NormStyle::Qwen,
5847 4096,
5848 SamplerConfig {
5849 seed: Some(42),
5850 ..Default::default()
5851 },
5852 )
5853}
5854
5855fn dense_ffn_batch(d: &DenseFfn, xs: &[f32], b: usize, pool: Option<&Pool>) -> Vec<f32> {
5858 let inter = d.gate_proj.rows();
5859 let hidden = d.down_proj.rows();
5860 if d.act == Act::Silu && b >= 32 && crate::gpu::enabled_here() && !crate::gpu::mm_killed() {
5866 if let (Some((model, w1)), Some((_, w3)), Some((_, w2))) = (
5867 d.gate_proj.mapped_q4t(),
5868 d.up_proj.mapped_q4t(),
5869 d.down_proj.mapped_q4t(),
5870 ) {
5871 let mut out = vec![0.0f32; b * hidden];
5872 if crate::gpu::q4t_ffn(model, w1, w3, w2, xs, b, hidden, inter, &mut out) {
5873 return out;
5874 }
5875 }
5876 }
5877 let mut g = vec![0.0f32; b * inter];
5878 d.gate_proj.matmat(xs, b, &mut g, pool);
5879 let mut u = vec![0.0f32; b * inter];
5880 d.up_proj.matmat(xs, b, &mut u, pool);
5881 for i in 0..b * inter {
5882 g[i] = d.act.combine(g[i], u[i]);
5883 }
5884 let mut out = vec![0.0f32; b * hidden];
5885 d.down_proj.matmat(&g, b, &mut out, pool);
5886 out
5887}
5888
5889fn accumulate_act(m: &MoeFfn, xs: &[f32], b: usize) {
5894 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
5895 static DUMP: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
5896 let on = *ON.get_or_init(|| std::env::var("CMF_RMS_TRACE").is_ok());
5897 let dump = *DUMP.get_or_init(|| std::env::var("CMF_ACT_DUMP").is_ok());
5898 if (!on && !dump) || b == 0 {
5899 return;
5900 }
5901 let hidden = xs.len() / b;
5902 if on {
5903 let mut acc = m.act_sq.borrow_mut();
5904 if acc.len() < hidden {
5905 acc.resize(hidden, 0.0);
5906 }
5907 for t in 0..b {
5908 let row = &xs[t * hidden..(t + 1) * hidden];
5909 for (a, &v) in acc.iter_mut().zip(row) {
5910 *a += (v as f64) * (v as f64);
5911 }
5912 }
5913 }
5914 if dump {
5915 let cap: usize = std::env::var("CMF_ACT_DUMP_ROWS")
5918 .ok()
5919 .and_then(|v| v.parse().ok())
5920 .unwrap_or(4096);
5921 let mut rows = m.act_rows.borrow_mut();
5922 if rows.len() < cap * hidden {
5923 let take = b.min((cap * hidden - rows.len()) / hidden.max(1));
5924 rows.extend_from_slice(&xs[..take * hidden]);
5925 }
5926 }
5927}
5928
5929fn moe_ffn_batch(
5930 m: &MoeFfn,
5931 xs: &[f32],
5932 b: usize,
5933 hidden: usize,
5934 pool: Option<&Pool>,
5935 allowed: Option<&[bool]>,
5936) -> Vec<f32> {
5937 accumulate_act(m, xs, b);
5938 let ne = m.experts.len();
5939 let mut logits = vec![0.0f32; b * ne];
5940 m.router.matmat(xs, b, &mut logits, pool);
5941
5942 let mut assign: Vec<Vec<(usize, f32)>> = vec![Vec::new(); ne];
5945 {
5946 let mut st = m.stats.borrow_mut();
5947 if st.len() < ne {
5948 st.resize(ne, 0);
5949 }
5950 for bi in 0..b {
5951 let (idx, p, wsum) = moe_route(&logits[bi * ne..(bi + 1) * ne], m, allowed);
5952 for &e in &idx {
5953 st[e] += 1;
5954 assign[e].push((bi, p[e] / wsum));
5955 }
5956 }
5957 }
5958
5959 let mut out = vec![0.0f32; b * hidden];
5960 let cols = m.experts[0].gate_proj.cols();
5961 let mut run_expert = |d: &DenseFfn, list: &[(usize, f32)]| {
5962 let sb = list.len();
5963 let mut sub = vec![0.0f32; sb * cols];
5964 for (k, &(bi, _)) in list.iter().enumerate() {
5965 sub[k * cols..(k + 1) * cols].copy_from_slice(&xs[bi * cols..(bi + 1) * cols]);
5966 }
5967 let eo = dense_ffn_batch(d, &sub, sb, pool);
5968 for (k, &(bi, w)) in list.iter().enumerate() {
5969 for i in 0..hidden {
5970 out[bi * hidden + i] += w * eo[k * hidden + i];
5971 }
5972 }
5973 };
5974 for (e, a) in assign.iter().enumerate().take(ne) {
5975 if !a.is_empty() {
5976 run_expert(&m.experts[e], a);
5977 }
5978 }
5979 if let Some((se, gate)) = &m.shared {
5980 let all: Vec<(usize, f32)> = if let Some(gate) = gate {
5981 let mut gl = vec![0.0f32; b];
5982 gate.matmat(xs, b, &mut gl, pool);
5983 (0..b)
5984 .map(|bi| (bi, 1.0 / (1.0 + (-gl[bi]).exp())))
5985 .collect()
5986 } else {
5987 (0..b).map(|bi| (bi, 1.0)).collect()
5988 };
5989 run_expert(se, &all);
5990 }
5991 out
5992}
5993
5994thread_local! {
5995 static FFN_SCRATCH: std::cell::RefCell<[Vec<f32>; 4]> =
5999 const { std::cell::RefCell::new([Vec::new(), Vec::new(), Vec::new(), Vec::new()]) };
6000}
6001
6002fn dense_ffn(d: &DenseFfn, x: &[f32], pool: Option<&Pool>) -> Vec<f32> {
6004 if crate::gpu::enabled_here()
6015 && (d.gate_proj.rows() >= crate::gpu::min_rows() || d.gate_proj.is_q1())
6016 {
6017 let arm = if d.gate_proj.is_q1() && crate::gpu::q1_force() {
6018 crate::gpu::ProbeArm::Gpu
6019 } else {
6020 crate::gpu::probe_arm(crate::gpu::OpClass::Ffn)
6021 };
6022 match arm {
6023 crate::gpu::ProbeArm::Gpu => {
6024 let t0 = std::time::Instant::now();
6025 if let Some(out) = dense_ffn_gpu(d, x, pool) {
6026 crate::gpu::probe_record(crate::gpu::OpClass::Ffn, true, t0.elapsed());
6027 return out;
6028 }
6029 }
6030 crate::gpu::ProbeArm::CpuTimed => {
6031 let t0 = std::time::Instant::now();
6032 let out = crate::gpu::cpu_scope(|| dense_ffn_cpu(d, x, pool));
6033 crate::gpu::probe_record(crate::gpu::OpClass::Ffn, false, t0.elapsed());
6034 return out;
6035 }
6036 crate::gpu::ProbeArm::Cpu => {
6037 return crate::gpu::cpu_scope(|| dense_ffn_cpu(d, x, pool));
6038 }
6039 }
6040 }
6041 dense_ffn_cpu(d, x, pool)
6042}
6043
6044fn dense_ffn_cpu(d: &DenseFfn, x: &[f32], pool: Option<&Pool>) -> Vec<f32> {
6046 let inter = d.gate_proj.rows();
6047 FFN_SCRATCH.with(|s| {
6048 let mut s = s.borrow_mut();
6049 let [g, u, ..] = &mut *s;
6050 g.resize(inter, 0.0);
6051 if d.act == Act::Silu && QTensor::matvec_silu_mul(&d.gate_proj, &d.up_proj, x, g, pool) {
6054 } else {
6056 u.resize(inter, 0.0);
6057 QTensor::matvec_many([&d.gate_proj, &d.up_proj], x, [g, u], pool);
6059 for i in 0..inter {
6060 g[i] = d.act.combine(g[i], u[i]);
6061 }
6062 }
6063 FFN_PROBE.with(|pr| {
6066 if let Some(acc) = pr.borrow_mut().as_mut() {
6067 let li = crate::gpu::cur_layer();
6068 if li >= 0 {
6069 if let Some(row) = acc.get_mut(li as usize) {
6070 for (a, &v) in row.iter_mut().zip(g.iter()) {
6071 *a += (v as f64).abs();
6072 }
6073 }
6074 }
6075 }
6076 });
6077 let mut out = attention::take_buf(d.down_proj.rows());
6078 d.down_proj.matvec(g, &mut out, pool);
6079 out
6080 })
6081}
6082
6083thread_local! {
6084 static FFN_PROBE: std::cell::RefCell<Option<Vec<Vec<f64>>>> =
6087 const { std::cell::RefCell::new(None) };
6088}
6089
6090fn dense_ffn_gpu(d: &DenseFfn, x: &[f32], _pool: Option<&Pool>) -> Option<Vec<f32>> {
6096 if d.act != Act::Silu {
6098 return None;
6099 }
6100 if d.gate_proj.rows() < crate::gpu::min_rows() && !d.gate_proj.is_q1() {
6103 return None;
6104 }
6105 let mut jobs: Vec<crate::gpu::MoeJob> = Vec::with_capacity(1);
6106 let mut model_ref = None;
6107 moe_push_job(d, x, 1.0, &mut jobs, &mut model_ref)?;
6108 let model = model_ref?;
6109 let hidden = jobs[0].down.1;
6110 let mut out = attention::take_buf(hidden);
6111 if crate::gpu::moe_block(&model, &jobs, &mut out) {
6112 Some(out)
6113 } else {
6114 let mut out = out;
6115 attention::recycle_buf(&mut out);
6116 None
6117 }
6118}
6119
6120#[allow(clippy::type_complexity)]
6125#[allow(clippy::type_complexity)]
6126pub(crate) fn moe_parts(
6127 t: &QTensor,
6128) -> Option<(
6129 &std::sync::Arc<cortiq_core::CmfModel>,
6130 usize,
6131 usize,
6132 usize,
6133 &[f32],
6134 &[f32],
6135 bool,
6136 bool,
6137)> {
6138 match t {
6139 QTensor::Mapped {
6140 model,
6141 idx,
6142 dtype: dt @ (cortiq_core::TensorDtype::Q8_2f | cortiq_core::TensorDtype::Q8Row),
6143 rows,
6144 cols,
6145 row_scale,
6146 col_field,
6147 ..
6148 } if (*dt == cortiq_core::TensorDtype::Q8Row) || !col_field.is_empty() => Some((
6149 model, *idx, *rows, *cols, row_scale, col_field, false, false,
6150 )),
6151 QTensor::Mapped {
6153 model,
6154 idx,
6155 dtype: cortiq_core::TensorDtype::Q1,
6156 rows,
6157 cols,
6158 ..
6159 } => Some((model, *idx, *rows, *cols, &[][..], &[][..], true, false)),
6160 QTensor::Mapped {
6162 model,
6163 idx,
6164 dtype: cortiq_core::TensorDtype::Q4Tiled,
6165 rows,
6166 cols,
6167 ..
6168 } => Some((model, *idx, *rows, *cols, &[][..], &[][..], false, true)),
6169 QTensor::Mapped {
6171 model,
6172 idx,
6173 dtype: cortiq_core::TensorDtype::Q4TiledP,
6174 rows,
6175 cols,
6176 ..
6177 } => Some((model, *idx, *rows, *cols, &[][..], &[][..], false, true)),
6178 _ => None,
6179 }
6180}
6181
6182pub(crate) fn moe_push_job_parts<'a>(
6186 gate: &'a QTensor,
6187 up: &'a QTensor,
6188 down: &'a QTensor,
6189 x: &[f32],
6190 w: f32,
6191 swiglu_limit: f32,
6192 jobs: &mut Vec<crate::gpu::MoeJob<'a>>,
6193 model_ref: &mut Option<std::sync::Arc<cortiq_core::CmfModel>>,
6194) -> Option<()> {
6195 use crate::qtensor::prescale;
6196 let (gm, gi, gr, gc, grs, gcf, gq1, gq4) = moe_parts(gate)?;
6197 let (_, ui, ur, uc, urs, ucf, uq1, uq4) = moe_parts(up)?;
6198 let (_, di, dr, dc, drs, dcf, dq1, dq4) = moe_parts(down)?;
6199 if gq1 != uq1 || uq1 != dq1 || gq4 != uq4 || uq4 != dq4 {
6200 return None; }
6202 model_ref.get_or_insert_with(|| gm.clone());
6203 let dt = |cf: &[f32]| {
6204 if cf.is_empty() {
6205 cortiq_core::TensorDtype::Q8Row
6206 } else {
6207 cortiq_core::TensorDtype::Q8_2f
6208 }
6209 };
6210 jobs.push(crate::gpu::MoeJob {
6211 gate: (gi, gr, gc, grs),
6212 up: (ui, ur, uc, urs),
6213 down: (di, dr, dc, drs),
6214 xs_gate: prescale(x, gcf, dt(gcf)).into_owned(),
6215 xs_up: prescale(x, ucf, dt(ucf)).into_owned(),
6216 down_col: dcf,
6217 w,
6218 q1: gq1,
6219 q4t: gq4 && gate.mapped_q4tp().is_none(),
6220 q4tp: gq4 && gate.mapped_q4tp().is_some(),
6221 swiglu_limit,
6222 });
6223 Some(())
6224}
6225
6226fn moe_push_job<'a>(
6228 d: &'a DenseFfn,
6229 x: &[f32],
6230 w: f32,
6231 jobs: &mut Vec<crate::gpu::MoeJob<'a>>,
6232 model_ref: &mut Option<std::sync::Arc<cortiq_core::CmfModel>>,
6233) -> Option<()> {
6234 use crate::qtensor::prescale;
6235 if d.act != Act::Silu {
6236 return None; }
6238 let (gm, gi, gr, gc, grs, gcf, gq1, gq4) = moe_parts(&d.gate_proj)?;
6239 let (_, ui, ur, uc, urs, ucf, uq1, uq4) = moe_parts(&d.up_proj)?;
6240 let (_, di, dr, dc, drs, dcf, dq1, dq4) = moe_parts(&d.down_proj)?;
6241 if gq1 != uq1 || uq1 != dq1 || gq4 != uq4 || uq4 != dq4 {
6242 return None; }
6244 model_ref.get_or_insert_with(|| gm.clone());
6245 let gdt = if gcf.is_empty() {
6246 cortiq_core::TensorDtype::Q8Row
6247 } else {
6248 cortiq_core::TensorDtype::Q8_2f
6249 };
6250 let udt = if ucf.is_empty() {
6251 cortiq_core::TensorDtype::Q8Row
6252 } else {
6253 cortiq_core::TensorDtype::Q8_2f
6254 };
6255 jobs.push(crate::gpu::MoeJob {
6256 gate: (gi, gr, gc, grs),
6257 up: (ui, ur, uc, urs),
6258 down: (di, dr, dc, drs),
6259 xs_gate: prescale(x, gcf, gdt).into_owned(),
6260 xs_up: prescale(x, ucf, udt).into_owned(),
6261 down_col: dcf,
6262 w,
6263 q1: gq1,
6264 q4t: gq4 && d.gate_proj.mapped_q4tp().is_none(),
6265 q4tp: gq4 && d.gate_proj.mapped_q4tp().is_some(),
6266 swiglu_limit: 0.0,
6267 });
6268 Some(())
6269}
6270
6271fn sparse_ffn_quant(
6278 d: &DenseFfn,
6279 x: &[f32],
6280 active: &[u16],
6281 hidden: usize,
6282 pool: Option<&Pool>,
6283) -> Vec<f32> {
6284 let n = active.len();
6285 let inter = d.gate_proj.rows();
6286 let mut act = vec![0.0f32; n];
6287 let need_scratch = !(d.gate_proj.sparse_col_ok() && d.up_proj.sparse_col_ok());
6290 let compute = |ai: usize| -> f32 {
6291 let idx = active[ai] as usize;
6292 if idx >= inter {
6293 return 0.0; }
6295 let mut s = if need_scratch {
6296 vec![0.0f32; hidden]
6297 } else {
6298 Vec::new()
6299 };
6300 let gate = d.gate_proj.row_dot(idx, x, &mut s);
6301 let up = d.up_proj.row_dot(idx, x, &mut s);
6302 d.act.combine(gate, up)
6303 };
6304 match pool {
6305 Some(p) if n >= 256 => {
6306 let ptr = SendMut(act.as_mut_ptr());
6307 p.run(&|widx, nw| {
6308 let chunk = n.div_ceil(nw);
6309 let (s, e) = (widx * chunk, ((widx + 1) * chunk).min(n));
6310 for ai in s..e {
6311 unsafe { *ptr.at(ai) = compute(ai) };
6312 }
6313 });
6314 }
6315 _ => {
6316 for (ai, a) in act.iter_mut().enumerate() {
6317 *a = compute(ai);
6318 }
6319 }
6320 }
6321 let mut out = vec![0.0f32; hidden];
6323 for (ai, &idx) in active.iter().enumerate() {
6324 let w = act[ai];
6325 if w.abs() >= 1e-12 && (idx as usize) < inter {
6326 d.down_proj.add_col_scaled(idx as usize, w, &mut out);
6327 }
6328 }
6329 out
6330}
6331
6332#[doc(hidden)]
6334pub fn sparse_ffn_quant_for_test(
6335 d: &DenseFfn,
6336 x: &[f32],
6337 active: &[u16],
6338 hidden: usize,
6339) -> Vec<f32> {
6340 sparse_ffn_quant(d, x, active, hidden, None)
6341}
6342
6343fn dequant_dense_f32(d: &DenseFfn) -> (Vec<f32>, Vec<f32>, Vec<f32>) {
6347 let deq = |t: &QTensor| -> Vec<f32> {
6348 let (rows, cols) = (t.rows(), t.cols());
6349 let mut out = vec![0.0f32; rows * cols];
6350 for r in 0..rows {
6351 t.row_f32(r, &mut out[r * cols..(r + 1) * cols]);
6352 }
6353 out
6354 };
6355 (deq(&d.gate_proj), deq(&d.up_proj), deq(&d.down_proj))
6356}
6357
6358struct SendMut(*mut f32);
6360unsafe impl Send for SendMut {}
6361unsafe impl Sync for SendMut {}
6362impl SendMut {
6363 #[inline]
6364 #[allow(clippy::mut_from_ref)]
6367 unsafe fn at(&self, i: usize) -> &mut f32 {
6368 unsafe { &mut *self.0.add(i) }
6369 }
6370}
6371
6372fn moe_route(logits: &[f32], m: &MoeFfn, allowed: Option<&[bool]>) -> (Vec<usize>, Vec<f32>, f32) {
6382 let ne = logits.len();
6383 let p: Vec<f32> = if m.router_sigmoid {
6384 logits.iter().map(|&l| 1.0 / (1.0 + (-l).exp())).collect()
6385 } else {
6386 let mx = logits.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
6387 let mut e: Vec<f32> = logits.iter().map(|&l| (l - mx).exp()).collect();
6388 let s: f32 = e.iter().sum();
6389 for v in &mut e {
6390 *v /= s;
6391 }
6392 e
6393 };
6394 let admit = |e: usize| {
6400 m.mask.as_ref().is_none_or(|mk| mk[e])
6401 && allowed.is_none_or(|a| a.get(e).copied().unwrap_or(false))
6402 };
6403 let mut idx: Vec<usize> = (0..ne).filter(|&e| admit(e)).collect();
6404 match &m.expert_bias {
6406 Some(b) => idx.sort_unstable_by(|&x, &y| {
6407 (p[y] + b[y])
6408 .partial_cmp(&(p[x] + b[x]))
6409 .unwrap()
6410 .then(x.cmp(&y))
6411 }),
6412 None => idx.sort_unstable_by(|&x, &y| p[y].partial_cmp(&p[x]).unwrap().then(x.cmp(&y))),
6413 }
6414 idx.truncate(m.top_k);
6415 if let Some(tau) = m.route_tau {
6419 let total: f32 = idx.iter().map(|&e| p[e]).sum();
6420 if total > 0.0 {
6421 let mut acc = 0.0f32;
6422 let mut keep = idx.len();
6423 for (i, &e) in idx.iter().enumerate() {
6424 acc += p[e];
6425 if acc >= tau * total {
6426 keep = i + 1;
6427 break;
6428 }
6429 }
6430 idx.truncate(keep);
6431 }
6432 }
6433 let wsum: f32 = if m.norm_topk_prob {
6434 let s: f32 = idx.iter().map(|&e| p[e]).sum();
6435 (if m.router_sigmoid { s + 1e-6 } else { s }) / m.routed_scaling
6438 } else {
6439 1.0 / m.routed_scaling
6440 };
6441 (idx, p, wsum)
6442}
6443
6444fn moe_ffn(m: &MoeFfn, x: &[f32], pool: Option<&Pool>, allowed: Option<&[bool]>) -> Vec<f32> {
6447 accumulate_act(m, x, 1);
6448 let ne = m.experts.len();
6449 let mut logits = vec![0.0f32; ne];
6450 m.router.matvec(x, &mut logits, pool);
6451 let (idx, p, wsum) = moe_route(&logits, m, allowed);
6452 {
6453 let mut st = m.stats.borrow_mut();
6454 if st.len() < ne {
6455 st.resize(ne, 0);
6456 }
6457 for &e in &idx {
6458 st[e] += 1;
6459 }
6460 }
6461 if crate::gpu::enabled_here() {
6466 match crate::gpu::probe_arm(crate::gpu::OpClass::Ffn) {
6467 crate::gpu::ProbeArm::Gpu => {
6468 let t0 = std::time::Instant::now();
6469 if let Some(out) = moe_ffn_gpu(m, x, &idx, &p, wsum, pool) {
6470 crate::gpu::probe_record(crate::gpu::OpClass::Ffn, true, t0.elapsed());
6471 return out;
6472 }
6473 }
6474 crate::gpu::ProbeArm::CpuTimed => {
6475 let t0 = std::time::Instant::now();
6476 let out = crate::gpu::cpu_scope(|| moe_ffn_cpu(m, x, &idx, &p, wsum, pool));
6477 crate::gpu::probe_record(crate::gpu::OpClass::Ffn, false, t0.elapsed());
6478 return out;
6479 }
6480 crate::gpu::ProbeArm::Cpu => {
6481 return crate::gpu::cpu_scope(|| moe_ffn_cpu(m, x, &idx, &p, wsum, pool));
6482 }
6483 }
6484 }
6485 moe_ffn_cpu(m, x, &idx, &p, wsum, pool)
6486}
6487
6488fn graph_note(built: bool) {
6492 use std::sync::atomic::{AtomicBool, Ordering};
6493 static SAID: AtomicBool = AtomicBool::new(false);
6494 if !SAID.swap(true, Ordering::Relaxed) {
6495 if built {
6496 tracing::info!("wgpu whole-token graph: ACTIVE");
6497 } else {
6498 tracing::warn!("wgpu whole-token graph refused — per-op path");
6499 }
6500 }
6501}
6502
6503fn moe_batch_enabled() -> bool {
6506 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
6507 *ON.get_or_init(|| std::env::var("CMF_MOE_BATCH").as_deref() != Ok("0"))
6508}
6509
6510fn moe_ffn_cpu_batched(
6516 m: &MoeFfn,
6517 x: &[f32],
6518 idx: &[usize],
6519 p: &[f32],
6520 wsum: f32,
6521 pool: Option<&Pool>,
6522) -> Option<Vec<f32>> {
6523 if idx.is_empty() || !moe_batch_enabled() {
6524 return None;
6525 }
6526 if FFN_PROBE.with(|pr| pr.borrow().is_some()) {
6530 return None;
6531 }
6532 let n = idx.len() + usize::from(m.shared.is_some());
6533 let mut pairs = Vec::with_capacity(n);
6534 let mut downs = Vec::with_capacity(n);
6535 let mut ws = Vec::with_capacity(n);
6536 for &e in idx {
6537 let d = &m.experts[e];
6538 if d.act != Act::Silu {
6539 return None;
6540 }
6541 pairs.push((&d.gate_proj, &d.up_proj));
6542 downs.push(&d.down_proj);
6543 ws.push(p[e] / wsum * m.per_expert_scale.as_ref().map_or(1.0, |v| v[e]));
6544 }
6545 if let Some((se, gate)) = &m.shared {
6548 if se.act != Act::Silu {
6549 return None;
6550 }
6551 let g = gate.as_ref().map_or(1.0, |gate| {
6552 let mut gl = [0.0f32; 1];
6553 gate.matvec(x, &mut gl, pool);
6554 1.0 / (1.0 + (-gl[0]).exp())
6555 });
6556 pairs.push((&se.gate_proj, &se.up_proj));
6557 downs.push(&se.down_proj);
6558 ws.push(g);
6559 }
6560 let inter = pairs[0].0.rows();
6561 let mut gs: Vec<Vec<f32>> = (0..pairs.len()).map(|_| vec![0f32; inter]).collect();
6562 if !QTensor::moe_gate_up_many(&pairs, x, &mut gs, pool) {
6563 return None;
6564 }
6565 let mut out = attention::take_buf(x.len());
6566 if !QTensor::moe_down_many(&downs, &gs, &ws, &mut out, pool) {
6567 attention::recycle_buf(&mut out);
6568 return None;
6569 }
6570 Some(out)
6571}
6572
6573fn moe_ffn_cpu(
6575 m: &MoeFfn,
6576 x: &[f32],
6577 idx: &[usize],
6578 p: &[f32],
6579 wsum: f32,
6580 pool: Option<&Pool>,
6581) -> Vec<f32> {
6582 if let Some(out) = moe_ffn_cpu_batched(m, x, idx, p, wsum, pool) {
6583 return out;
6584 }
6585 let mut out = attention::take_buf(x.len());
6586 for &e in idx {
6587 let mut eo = dense_ffn(&m.experts[e], x, pool);
6588 let w = p[e] / wsum * m.per_expert_scale.as_ref().map_or(1.0, |v| v[e]);
6589 for i in 0..out.len() {
6590 out[i] += w * eo[i];
6591 }
6592 attention::recycle_buf(&mut eo);
6593 }
6594 if let Some((se, gate)) = &m.shared {
6595 let mut so = dense_ffn(se, x, pool);
6596 let g = gate.as_ref().map_or(1.0, |gate| {
6597 let mut gl = [0.0f32; 1];
6598 gate.matvec(x, &mut gl, pool);
6599 1.0 / (1.0 + (-gl[0]).exp())
6600 });
6601 for i in 0..out.len() {
6602 out[i] += g * so[i];
6603 }
6604 attention::recycle_buf(&mut so);
6605 }
6606 out
6607}
6608
6609#[allow(clippy::too_many_arguments)]
6617fn mla_attention(
6618 w: &MlaWeights,
6619 normed: &[f32],
6620 cache: &mut crate::kv_cache::LayerKvCache,
6621 position: usize,
6622 inv_freq: &[f32],
6623 rope_scale: f32,
6624 eps: f64,
6625 pool: Option<&Pool>,
6626) -> Vec<f32> {
6627 let (nh, dr, dn, dv, lora) = (w.nh, w.qk_rope, w.qk_nope, w.v_dim, w.lora);
6628 let hd = dr + dn;
6629 let mut q = vec![0.0f32; nh * hd];
6630 match (&w.q_a, &w.q_a_norm) {
6631 (Some(qa), Some(qn)) => {
6632 let mut t = vec![0.0f32; qa.rows()];
6633 qa.matvec(normed, &mut t, pool);
6634 let tn = inference::rms_norm(&t, qn, eps, NormStyle::Qwen);
6635 w.q_proj.matvec(&tn, &mut q, pool);
6636 }
6637 _ => w.q_proj.matvec(normed, &mut q, pool),
6638 }
6639 let mut ca = vec![0.0f32; lora + dr];
6640 w.kv_a.matvec(normed, &mut ca, pool);
6641 let (c_lat, k_rope) = ca.split_at_mut(lora);
6642 let latn = inference::rms_norm(c_lat, &w.kv_a_norm, eps, NormStyle::Qwen);
6643 let mut kvb = vec![0.0f32; nh * (dn + dv)];
6644 w.kv_b.matvec(&latn, &mut kvb, pool);
6645 if !w.nope {
6646 attention::rope_rotate_scaled(k_rope, position, inv_freq, rope_scale);
6647 }
6648 for h in 0..nh {
6649 if !w.nope {
6650 attention::rope_rotate_scaled(
6651 &mut q[h * hd..h * hd + dr],
6652 position,
6653 inv_freq,
6654 rope_scale,
6655 );
6656 }
6657 }
6658 let mut k = vec![0.0f32; nh * hd];
6659 let mut v = vec![0.0f32; nh * hd];
6660 for h in 0..nh {
6661 k[h * hd..h * hd + dr].copy_from_slice(k_rope);
6662 k[h * hd + dr..(h + 1) * hd].copy_from_slice(&kvb[h * (dn + dv)..h * (dn + dv) + dn]);
6663 v[h * hd..h * hd + dv].copy_from_slice(&kvb[h * (dn + dv) + dn..(h + 1) * (dn + dv)]);
6664 }
6665 cache.append(&k, &v, &vec![true; nh]);
6666 let (ao, mut imp) = attention::attend_all_heads(&q, cache, nh, 1, hd, w.scale, None, 0.0);
6667 attention::recycle_buf(&mut imp);
6668 let mut ov = vec![0.0f32; nh * dv];
6669 for h in 0..nh {
6670 ov[h * dv..(h + 1) * dv].copy_from_slice(&ao[h * hd..h * hd + dv]);
6671 }
6672 let mut out = vec![0.0f32; w.o_proj.rows()];
6673 w.o_proj.matvec(&ov, &mut out, pool);
6674 out
6675}
6676
6677fn dense_moe_ffn(
6684 dm: &DenseMoeFfn,
6685 x_normed: &[f32],
6686 h_raw: &[f32],
6687 eps: f64,
6688 norm_style: NormStyle,
6689 pool: Option<&Pool>,
6690) -> Vec<f32> {
6691 let mut d = dense_ffn(&dm.dense, x_normed, pool);
6692 d = inference::rms_norm(&d, &dm.post_norm_1, eps, norm_style);
6693 let m = &dm.moe;
6694 let ne = m.experts.len();
6695 let mut logits = vec![0.0f32; ne];
6696 if m.router_input_norm {
6697 let ss: f32 = h_raw.iter().map(|v| v * v).sum::<f32>() / h_raw.len() as f32;
6698 let inv = 1.0 / (ss + eps as f32).sqrt();
6699 let xr: Vec<f32> = h_raw.iter().map(|v| v * inv).collect();
6700 m.router.matvec(&xr, &mut logits, pool);
6701 } else {
6702 m.router.matvec(h_raw, &mut logits, pool);
6703 }
6704 let (idx, p, wsum) = moe_route(&logits, m, None);
6705 {
6706 let mut st = m.stats.borrow_mut();
6707 if st.len() < ne {
6708 st.resize(ne, 0);
6709 }
6710 for &e in &idx {
6711 st[e] += 1;
6712 }
6713 }
6714 let x2 = inference::rms_norm(h_raw, &dm.pre_norm_2, eps, norm_style);
6715 let mo = moe_ffn_cpu(m, &x2, &idx, &p, wsum, pool);
6716 let mo = inference::rms_norm(&mo, &dm.post_norm_2, eps, norm_style);
6717 for (di, mi) in d.iter_mut().zip(&mo) {
6718 *di += mi;
6719 }
6720 d
6721}
6722
6723fn moe_gpu_refused(why: &'static str) {
6730 use std::sync::atomic::{AtomicBool, Ordering};
6731 static SAID: AtomicBool = AtomicBool::new(false);
6732 if !SAID.swap(true, Ordering::Relaxed) {
6733 tracing::warn!("MoE GPU block refused ({why}) — experts run on the CPU");
6734 }
6735}
6736
6737fn moe_ffn_gpu(
6738 m: &MoeFfn,
6739 x: &[f32],
6740 idx: &[usize],
6741 p: &[f32],
6742 wsum: f32,
6743 pool: Option<&Pool>,
6744) -> Option<Vec<f32>> {
6745 use crate::gpu::MoeJob;
6746
6747 let mut jobs: Vec<MoeJob> = Vec::with_capacity(idx.len() + 1);
6748 let mut model_ref = None;
6749 for &e in idx {
6750 if moe_push_job(&m.experts[e], x, p[e] / wsum, &mut jobs, &mut model_ref).is_none() {
6751 moe_gpu_refused("push_job(expert)");
6752 return None;
6753 }
6754 }
6755 if let Some((se, gate)) = &m.shared {
6756 let g = gate.as_ref().map_or(1.0, |gate| {
6757 let mut gl = [0.0f32; 1];
6758 gate.matvec(x, &mut gl, pool);
6759 1.0 / (1.0 + (-gl[0]).exp())
6760 });
6761 if moe_push_job(se, x, g, &mut jobs, &mut model_ref).is_none() {
6762 moe_gpu_refused("push_job(shared)");
6763 return None;
6764 }
6765 }
6766 let Some(model) = model_ref else {
6767 moe_gpu_refused("no model_ref");
6768 return None;
6769 };
6770 let hidden = jobs[0].down.1;
6771 let mut out = vec![0.0f32; hidden];
6772 if crate::gpu::moe_block(&model, &jobs, &mut out) {
6773 Some(out)
6774 } else {
6775 moe_gpu_refused("gpu::moe_block");
6776 None
6777 }
6778}
6779
6780fn ffn_forward(
6782 ffn: &FfnKind,
6783 x: &[f32],
6784 pool: Option<&Pool>,
6785 experts_allowed: Option<&[bool]>,
6786) -> Vec<f32> {
6787 match ffn {
6788 FfnKind::Dense(d) => dense_ffn(d, x, pool),
6789 FfnKind::Moe(m) => moe_ffn(m, x, pool, experts_allowed),
6790 FfnKind::DenseMoe(_) => unreachable!("DenseMoe dispatches via dense_moe_ffn"),
6794 }
6795}
6796
6797fn ffn_forward_pair(
6801 ffn: &FfnKind,
6802 x1: &[f32],
6803 x2: &[f32],
6804 pool: Option<&Pool>,
6805 experts_allowed: Option<&[bool]>,
6806) -> (Vec<f32>, Vec<f32>) {
6807 let d = match ffn {
6808 FfnKind::Dense(d) => d,
6809 FfnKind::Moe(m) => {
6810 return (
6811 moe_ffn(m, x1, pool, experts_allowed),
6812 moe_ffn(m, x2, pool, experts_allowed),
6813 );
6814 }
6815 FfnKind::DenseMoe(_) => unreachable!("DenseMoe dispatches via dense_moe_ffn"),
6816 };
6817 let inter = d.gate_proj.rows();
6818 FFN_SCRATCH.with(|s| {
6819 let mut s = s.borrow_mut();
6820 let [g1, g2, u1, u2] = &mut *s;
6821 g1.resize(inter, 0.0);
6822 g2.resize(inter, 0.0);
6823 u1.resize(inter, 0.0);
6824 u2.resize(inter, 0.0);
6825 QTensor::matvec2_many(
6828 [&d.gate_proj, &d.up_proj],
6829 x1,
6830 x2,
6831 [g1.as_mut_slice(), u1.as_mut_slice()],
6832 [g2.as_mut_slice(), u2.as_mut_slice()],
6833 pool,
6834 );
6835 for i in 0..inter {
6836 g1[i] = d.act.combine(g1[i], u1[i]);
6837 g2[i] = d.act.combine(g2[i], u2[i]);
6838 }
6839 let mut o1 = attention::take_buf(d.down_proj.rows());
6840 let mut o2 = attention::take_buf(d.down_proj.rows());
6841 d.down_proj.matvec2(g1, g2, &mut o1, &mut o2, pool);
6842 (o1, o2)
6843 })
6844}
6845
6846#[cfg(test)]
6847mod tests {
6848
6849 #[test]
6850 fn cancel_flag_stops_generation() {
6851 let mut p = create_test_pipeline(16, 32, 2, 2, 8, 2, 32);
6852 p.cancel.store(true, std::sync::atomic::Ordering::Relaxed);
6855 let r = p.generate_from_ids(&[1, 2, 3], 8, None, None).unwrap();
6856 assert_eq!(r.finish_reason, "cancelled");
6857 assert!(
6858 r.token_ids.is_empty(),
6859 "no tokens after cancel: {:?}",
6860 r.token_ids
6861 );
6862 let r2 = p.generate_from_ids(&[1, 2, 3], 4, None, None).unwrap();
6864 assert_ne!(r2.finish_reason, "cancelled");
6865 }
6866 use super::*;
6867
6868 #[test]
6874 fn sparse_ffn_quant_equals_dense_with_inactive_zeroed() {
6875 let (hidden, inter) = (16usize, 40usize);
6876 let synth = |n: usize, salt: usize| -> Vec<f32> {
6877 (0..n)
6878 .map(|i| (((i * 37 + salt * 11 + 3) % 101) as f32 / 101.0 - 0.5) * 0.4)
6879 .collect()
6880 };
6881 let d = DenseFfn {
6882 gate_proj: QTensor::from_f32(synth(inter * hidden, 1), inter, hidden),
6883 up_proj: QTensor::from_f32(synth(inter * hidden, 2), inter, hidden),
6884 down_proj: QTensor::from_f32(synth(hidden * inter, 3), hidden, inter),
6885 act: Act::Silu,
6886 };
6887 let x = synth(hidden, 9);
6888 let active: Vec<u16> = (0..inter as u16).filter(|i| i % 3 == 0).collect();
6890
6891 let sparse = sparse_ffn_quant(&d, &x, &active, hidden, None);
6892
6893 let mut g = vec![0.0f32; inter];
6895 d.gate_proj.matvec(&x, &mut g, None);
6896 let mut u = vec![0.0f32; inter];
6897 d.up_proj.matvec(&x, &mut u, None);
6898 let act_set: std::collections::HashSet<u16> = active.iter().copied().collect();
6899 for i in 0..inter {
6900 g[i] = if act_set.contains(&(i as u16)) {
6901 inference::silu(g[i]) * u[i]
6902 } else {
6903 0.0
6904 };
6905 }
6906 let mut reference = vec![0.0f32; hidden];
6907 d.down_proj.matvec(&g, &mut reference, None);
6908
6909 let max_d = sparse
6910 .iter()
6911 .zip(&reference)
6912 .map(|(a, b)| (a - b).abs())
6913 .fold(0.0f32, f32::max);
6914 assert!(max_d < 1e-5, "sparse != dense-zeroed: max|Δ| = {max_d}");
6915 }
6916
6917 fn attach_test_mtp(p: &mut Pipeline) {
6919 let (h, inter, heads, kv, hd) = (
6920 p.hidden_size,
6921 p.intermediate_size,
6922 p.num_heads,
6923 p.num_kv_heads,
6924 p.head_dim,
6925 );
6926 let synth = |n: usize, salt: usize| -> Vec<f32> {
6927 (0..n)
6928 .map(|i| (((i * 29 + salt * 23 + 5) % 101) as f32 / 101.0 - 0.5) * 0.2)
6929 .collect()
6930 };
6931 let qt = |rows: usize, cols: usize, salt: usize| -> QTensor {
6932 QTensor::from_f32(synth(rows * cols, salt), rows, cols)
6933 };
6934 p.mtp = Some(MtpModule {
6935 enorm: vec![1.0; h],
6936 hnorm: vec![1.0; h],
6937 eh_proj: qt(h, 2 * h, 301),
6938 layer: LayerWeights {
6939 input_norm: vec![1.0; h],
6940 post_norm: vec![1.0; h],
6941 attn_out_norm: None,
6942 ffn_out_norm: None,
6943 layer_scale: None,
6944 ffn: FfnKind::Dense(DenseFfn {
6945 gate_proj: qt(inter, h, 315),
6946 up_proj: qt(inter, h, 316),
6947 down_proj: qt(h, inter, 317),
6948 act: Act::Silu,
6949 }),
6950 attn: AttnKind::Full {
6951 bias: None,
6952 wq: qt(heads * hd, h, 311),
6953 wk: qt(kv * hd, h, 312),
6954 wv: qt(kv * hd, h, 313),
6955 wo: qt(h, heads * hd, 314),
6956 q_norm: None,
6957 k_norm: None,
6958 output_gate: false,
6959 softplus_gate: None,
6960 },
6961 },
6962 final_norm: vec![1.0; h],
6963 kv: crate::kv_cache::LayerKvCache::new(kv, hd),
6964 });
6965 }
6966
6967 #[test]
6968 fn speculative_equals_vanilla_greedy() {
6969 unsafe { std::env::set_var("CMF_GPU_WGPU_GRAPH", "0") };
6973 let run = |spec: bool| {
6974 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
6975 p.sampler_config.temperature = 0.0;
6976 attach_test_mtp(&mut p);
6977 p.speculative = spec;
6978 let r = p.generate("abcdef", 12, None, None).unwrap();
6979 (r.token_ids, r.mtp_drafted, r.mtp_accepted)
6980 };
6981 let (vanilla, d0, _) = run(false);
6982 let (spec, d1, a1) = run(true);
6983 assert_eq!(d0, 0, "vanilla path must not draft");
6984 assert!(d1 > 0, "speculative path must draft");
6985 assert_eq!(
6986 vanilla, spec,
6987 "speculative must reproduce the exact greedy sequence (accepted {a1}/{d1})"
6988 );
6989 }
6990
6991 #[test]
6992 fn speculative_accepts_constant_oracle() {
6993 unsafe { std::env::set_var("CMF_GPU_WGPU_GRAPH", "0") };
6995 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
6996 p.sampler_config.temperature = 0.0;
6997 p.sampler_config.repetition_penalty = 1.0;
6998 p.weights.lm_head = QTensor::from_f32(vec![0.01; 64 * 8], 64, 8);
7001 attach_test_mtp(&mut p);
7002 p.speculative = true;
7003 let r = p.generate("abcd", 10, None, None).unwrap();
7004 assert!(r.mtp_drafted > 0);
7005 assert_eq!(
7006 r.mtp_accepted, r.mtp_drafted,
7007 "constant logits → every draft accepted"
7008 );
7009 assert!(r.token_ids.windows(2).all(|w| w[0] == w[1]));
7012 }
7013
7014 #[test]
7015 fn empty_prompt_is_an_error_not_a_panic() {
7016 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 260);
7017 let r = p.generate("", 4, None, None);
7018 assert!(r.is_err(), "empty prompt must be a clean error");
7019 }
7020
7021 #[test]
7022 fn every_token_enters_kv_exactly_once() {
7023 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
7024 p.sampler_config.temperature = 0.0;
7026 let r = p.generate("abc", 2, None, None).unwrap();
7027 assert_eq!(r.prompt_tokens, 3);
7028 assert_eq!(
7032 p.kv_cache.seq_len(),
7033 3 + r.tokens_generated - 1,
7034 "each token must be cached exactly once (v1 cached the last prompt token twice)"
7035 );
7036 }
7037
7038 #[test]
7039 fn generation_is_reproducible_with_seed() {
7040 let run = || {
7041 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
7042 p.generate("hello", 8, None, None).unwrap().token_ids
7043 };
7044 assert_eq!(run(), run());
7045 }
7046
7047 #[test]
7048 fn resetting_sampler_restarts_the_seeded_stream() {
7049 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
7050 let config = SamplerConfig {
7051 seed: Some(1234),
7052 ..SamplerConfig::default()
7053 };
7054 p.set_sampler_config(config.clone());
7055 let first = p.generate("hello", 8, None, None).unwrap().token_ids;
7056 p.set_sampler_config(config);
7057 let second = p.generate("hello", 8, None, None).unwrap().token_ids;
7058 assert_eq!(first, second);
7059 }
7060
7061 #[test]
7062 fn eviction_bounds_the_cache() {
7063 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 260);
7064 p.kv_cache.max_seq_len = 6;
7065 p.sampler_config.temperature = 0.0;
7066 let _ = p.generate("abcd", 12, None, None).unwrap();
7067 assert!(
7068 p.kv_cache.seq_len() <= 6 + 1,
7069 "cache must stay bounded by max_seq_len (got {})",
7070 p.kv_cache.seq_len()
7071 );
7072 }
7073
7074 #[test]
7075 fn confidence_matches_tokens_and_is_a_probability() {
7076 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
7077 p.sampler_config.temperature = 0.0;
7078 p.sampler_config.repetition_penalty = 1.0;
7079 let r = p.generate("abcd", 10, None, None).unwrap();
7080 assert_eq!(
7081 r.token_confidence.len(),
7082 r.token_ids.len(),
7083 "one confidence per emitted token"
7084 );
7085 for &c in &r.token_confidence {
7086 assert!((0.0..=1.0).contains(&c), "confidence out of [0,1]: {c}");
7087 }
7088 let logits = [1.0f32, 3.0, 0.5, 3.0];
7090 let p0 = top1_prob_t(&logits, 1, 1.0);
7091 let p1 = top1_prob_t(&logits, 3, 1.0);
7092 assert!((p0 - p1).abs() < 1e-6, "equal logits → equal prob");
7093 assert!(p0 > 0.0 && p0 < 1.0);
7094 let sharp = top1_prob_t(&logits, 1, 1.0);
7096 let soft = top1_prob_t(&logits, 1, 2.0);
7097 assert!(soft < sharp, "higher temperature lowers peak confidence");
7098 }
7099
7100 #[test]
7101 fn trace_is_opt_in_and_parallels_the_output() {
7102 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
7104 p.sampler_config.temperature = 0.0;
7105 p.sampler_config.repetition_penalty = 1.0;
7106 let r = p.generate("abcd", 10, None, None).unwrap();
7107 assert!(r.traces.is_empty(), "trace must be empty unless enabled");
7108
7109 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
7111 p.sampler_config.temperature = 0.0;
7112 p.sampler_config.repetition_penalty = 1.0;
7113 p.set_trace(true);
7114 let r = p.generate("abcd", 10, None, None).unwrap();
7115 assert_eq!(r.traces.len(), r.token_ids.len(), "one trace row per token");
7116 for (i, tr) in r.traces.iter().enumerate() {
7117 assert_eq!(tr.t, i, "trace index is sequential");
7118 assert_eq!(tr.token_id, r.token_ids[i], "trace token_id matches output");
7119 assert_eq!(
7120 tr.confidence, r.token_confidence[i],
7121 "trace confidence matches the confidence channel"
7122 );
7123 assert!(tr.active_skill.is_none() && tr.recon.is_none() && !tr.switched);
7125 }
7126 }
7127
7128 #[test]
7129 fn explain_prefill_logits_match_greedy_first_token() {
7130 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
7134 p.sampler_config.temperature = 0.0;
7135 p.sampler_config.repetition_penalty = 1.0;
7136 let ids = p.tokenizer.encode("abcd");
7137 let logits = p.prefill_next_logits(&ids, None);
7138 let argmax = logits
7139 .iter()
7140 .enumerate()
7141 .max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
7142 .unwrap()
7143 .0 as u32;
7144 let r = p.generate("abcd", 1, None, None).unwrap();
7145 assert_eq!(
7146 argmax, r.token_ids[0],
7147 "explain preview must match greedy emit"
7148 );
7149 }
7150
7151 #[test]
7152 fn laguna_shared_expert_is_unconditionally_added() {
7153 let matrix = |values: Vec<f32>| QTensor::from_f32(values, 2, 2);
7154 let identity = || matrix(vec![1.0, 0.0, 0.0, 1.0]);
7155 let zero_dense = || DenseFfn {
7156 gate_proj: matrix(vec![0.0; 4]),
7157 up_proj: matrix(vec![0.0; 4]),
7158 down_proj: matrix(vec![0.0; 4]),
7159 act: Act::Silu,
7160 };
7161 let shared = DenseFfn {
7162 gate_proj: identity(),
7163 up_proj: identity(),
7164 down_proj: identity(),
7165 act: Act::Silu,
7166 };
7167 let x = [1.0, 2.0];
7168 let expected = dense_ffn(&shared, &x, None);
7169 let moe = MoeFfn {
7170 router: QTensor::from_f32(vec![0.0, 0.0], 1, 2),
7171 experts: vec![zero_dense()],
7172 top_k: 1,
7173 norm_topk_prob: true,
7174 router_sigmoid: true,
7175 expert_bias: None,
7176 routed_scaling: 1.0,
7177 route_tau: None,
7178 shared: Some((shared, None)),
7179 stats: std::cell::RefCell::new(Vec::new()),
7180 act_sq: std::cell::RefCell::new(Vec::new()),
7181 act_rows: std::cell::RefCell::new(Vec::new()),
7182 mask: None,
7183 per_expert_scale: None,
7184 router_input_norm: false,
7185 };
7186 let actual = moe_ffn_cpu(&moe, &x, &[0], &[0.0], 1.0, None);
7187 for (actual, expected) in actual.iter().zip(expected) {
7188 assert!((actual - expected).abs() < 1e-6);
7189 }
7190 }
7191}