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 spec_active = self.speculative
1470 && self.mtp.is_some()
1471 && task_mask.is_none()
1472 && !self.o1_active()
1473 && !graph_on
1474 && self.sampler_config.temperature < 1e-6;
1475 let mut mtp = if spec_active { self.mtp.take() } else { None };
1478 if let Some(m) = &mut mtp {
1479 m.kv.clear();
1480 }
1481 let mut router = if mtp.is_none() {
1485 self.dyn_router.take()
1486 } else {
1487 None
1488 };
1489 if let Some(r) = &mut router {
1490 r.reset(); self.dyn_phi_seen = 0; let _ = self.set_active_skill(None);
1493 }
1494
1495 let mut all_ids = input_ids.to_vec();
1496 let mut generated = 0usize;
1497 let mut finish_reason = "max_tokens".to_string();
1498 let mut drafted = 0usize;
1499 let mut accepted = 0usize;
1500 let mut confidence: Vec<f32> = Vec::new();
1501 let trace_on = self.trace;
1502 let calib_temp = self.calib_temp;
1503 let mut traces: Vec<TokenTrace> = Vec::new();
1504
1505 let mut hidden = vec![0.0f32; self.hidden_size];
1511 let mut pos = reuse_from;
1512 let fuse_lm = mtp.is_none()
1521 && router.is_none()
1522 && std::env::var("CMF_GPU_LMHEAD").as_deref() != Ok("0");
1523 self.graph_logits = None;
1524 self.graph_want_logits = false;
1525 let _tpf = std::time::Instant::now();
1526 let batch_k = std::env::var("CMF_BATCH_K")
1527 .ok()
1528 .and_then(|v| v.parse::<usize>().ok())
1529 .unwrap_or(0);
1530 while self.dsv4.is_some()
1541 && mtp.is_none()
1542 && pos < input_ids.len()
1543 && !self.cancel.load(std::sync::atomic::Ordering::Relaxed)
1544 {
1545 let end = (pos + prefill_chunk()).min(input_ids.len());
1546 let ids: Vec<u32> = input_ids[pos..end].to_vec();
1547 let mut lg = Vec::new();
1548 if let Some(b) = &mut self.dsv4 {
1549 let (g, layers, cfg, st) = (&b.0, &b.1, b.2, &mut b.3);
1550 crate::dsv4::forward_chunk(
1551 g,
1552 layers,
1553 &cfg,
1554 st,
1555 &ids,
1556 pos,
1557 &self.inv_freq,
1558 self.pool.as_deref(),
1559 &mut lg,
1560 end == input_ids.len(),
1561 );
1562 }
1563 if end == input_ids.len() {
1564 self.graph_logits = Some(lg);
1565 }
1566 pos = end;
1567 hidden = vec![0.0; self.hidden_size];
1568 }
1569 let dyn_prefill = router.is_some();
1574 let graph_prefill = self.graph_prefill_preferred();
1580 if task_mask.is_none()
1581 && !dyn_prefill
1582 && !graph_prefill
1583 && self.can_prefill_batched()
1584 && self.g3n.is_none()
1585 && input_ids.len() > 2
1586 {
1587 let chunk = prefill_chunk();
1593 let hs = self.hidden_size;
1594 while pos < input_ids.len() && !self.cancel.load(std::sync::atomic::Ordering::Relaxed) {
1595 let end = (pos + chunk).min(input_ids.len());
1596 let hb = self.prefill_batch(&input_ids[pos..end], pos);
1597 if let Some(m) = &mut mtp {
1598 for p in pos..end {
1599 if p + 1 < input_ids.len() {
1600 let _ = self.mtp_step(
1601 m,
1602 &hb[(p - pos) * hs..(p - pos + 1) * hs],
1603 input_ids[p + 1],
1604 p,
1605 );
1606 }
1607 }
1608 }
1609 hidden.copy_from_slice(&hb[(end - pos - 1) * hs..]);
1610 pos = end;
1611 }
1612 }
1613 let pair_off = std::env::var("CMF_PAIR").is_ok_and(|v| v == "0");
1614 if task_mask.is_none()
1615 && !dyn_prefill
1616 && !graph_prefill
1617 && !pair_off
1618 && self.pair_supported()
1619 {
1620 while pos + 1 < input_ids.len()
1621 && !self.cancel.load(std::sync::atomic::Ordering::Relaxed)
1622 {
1623 let e1 = self.embed_single(input_ids[pos]);
1624 let e2 = self.embed_single(input_ids[pos + 1]);
1625 let (h1, h2) = self.forward_pair(&e1, &e2, pos);
1626 self.commit_linear_scratch();
1628 if let Some(m) = &mut mtp {
1629 let _ = self.mtp_step(m, &h1, input_ids[pos + 1], pos);
1630 if pos + 2 < input_ids.len() {
1631 let _ = self.mtp_step(m, &h2, input_ids[pos + 2], pos + 1);
1632 }
1633 }
1634 hidden = h2;
1635 pos += 2;
1636 }
1637 }
1638 if batch_k > 0
1647 && graph_prefill
1648 && task_mask.is_none()
1649 && !self.o1_active()
1650 && mtp.is_none()
1651 && !dyn_prefill
1652 && pos + 1 < input_ids.len()
1653 {
1654 let hs = self.hidden_size;
1655 let chunk = batch_k;
1656 while pos < input_ids.len() {
1657 let end = (pos + chunk).min(input_ids.len());
1658 let bk = end - pos;
1659 let mut hiddens = vec![0f32; bk * hs];
1660 for (j, &id) in input_ids[pos..end].iter().enumerate() {
1661 hiddens[j * hs..(j + 1) * hs].copy_from_slice(&self.embed_single(id));
1662 }
1663 let positions: Vec<usize> = (pos..end).collect();
1664 let t_chunk = std::time::Instant::now();
1665 let ok_b = self.try_batch_graph_wgpu(&mut hiddens, &positions, bk);
1666 if std::env::var("CMF_GRAPH_PROF").is_ok() {
1667 let ms = t_chunk.elapsed().as_secs_f64() * 1000.0;
1668 eprintln!(
1669 "batch-chunk: k={bk} ok={ok_b} {ms:.1} ms ({:.1} tok/s)",
1670 bk as f64 / (ms / 1000.0)
1671 );
1672 }
1673 {
1674 use std::sync::atomic::{AtomicBool, Ordering};
1675 static SAID: AtomicBool = AtomicBool::new(false);
1676 if !SAID.swap(true, Ordering::Relaxed) {
1677 if ok_b {
1678 tracing::info!("batched prefill: ACTIVE (k={bk})");
1679 } else {
1680 tracing::warn!("batched prefill declined — per-position graph");
1681 }
1682 }
1683 }
1684 if ok_b {
1685 hidden.copy_from_slice(&hiddens[(bk - 1) * hs..]);
1686 pos = end;
1687 } else {
1688 break; }
1690 }
1691 }
1692 while pos < input_ids.len() && !self.cancel.load(std::sync::atomic::Ordering::Relaxed) {
1693 self.graph_want_logits = fuse_lm && pos + 1 == input_ids.len();
1694 hidden = self.forward_layers(&self.embed_single(input_ids[pos]), pos, task_mask);
1695 if let Some(m) = &mut mtp {
1696 if pos + 1 < input_ids.len() {
1697 let _ = self.mtp_step(m, &hidden, input_ids[pos + 1], pos);
1698 }
1699 }
1700 pos += 1;
1701 }
1702 if std::env::var("CMF_PREFILL_PROF").is_ok() {
1703 eprintln!(
1704 "prefill: {} tokens in {:.1} ms (batch_k={batch_k})",
1705 input_ids.len(),
1706 _tpf.elapsed().as_secs_f64() * 1000.0
1707 );
1708 }
1709 if self
1712 .cancel
1713 .swap(false, std::sync::atomic::Ordering::Relaxed)
1714 {
1715 self.kv_history.clear();
1716 if let Some(m) = mtp {
1717 self.mtp = Some(m);
1718 }
1719 return Ok(GenerateResult {
1720 text: String::new(),
1721 token_ids: Vec::new(),
1722 prompt_tokens: input_ids.len(),
1723 tokens_generated: 0,
1724 finish_reason: "cancelled".to_string(),
1725 mtp_drafted: 0,
1726 mtp_accepted: 0,
1727 token_confidence: Vec::new(),
1728 traces: Vec::new(),
1729 });
1730 }
1731
1732 self.o1_seal();
1735
1736 macro_rules! commit {
1738 ($id:expr) => {{
1739 all_ids.push($id);
1740 generated += 1;
1741 if self.tokenizer.is_eos($id) {
1742 finish_reason = "stop".to_string();
1743 false
1744 } else {
1745 let token_text = self.tokenizer.decode_token($id);
1746 let mut go = true;
1747 if let Some(ref mut cb) = on_token {
1748 if !cb(&token_text) {
1749 finish_reason = "cancelled".to_string();
1750 go = false;
1751 }
1752 }
1753 go
1754 }
1755 }};
1756 }
1757
1758 let mut next_pos = input_ids.len();
1760 'decode: while generated < max_tokens {
1761 if self
1762 .cancel
1763 .swap(false, std::sync::atomic::Ordering::Relaxed)
1764 {
1765 finish_reason = "cancelled".to_string();
1766 break 'decode;
1767 }
1768 let mut logits = match self.graph_logits.take() {
1769 Some(lg) => lg,
1770 None => {
1771 inference::rms_norm_into(
1772 &hidden,
1773 &self.weights.final_norm,
1774 self.rms_eps,
1775 self.norm_style,
1776 &mut self.ws.n1,
1777 );
1778 self.lm_head_forward(&self.ws.n1)
1779 }
1780 };
1781 let t_next = sampler::sample_with_scratch(
1782 &logits,
1783 &self.sampler_config,
1784 &all_ids,
1785 &mut self.rng,
1786 &mut self.sampler_scratch,
1787 );
1788 if self.confidence_on {
1789 confidence.push(top1_prob_t(&logits, t_next, calib_temp));
1790 }
1791 attention::recycle_buf(&mut logits);
1792 if trace_on {
1793 let skill = router.as_ref().and_then(|r| r.active_id());
1797 traces.push(TokenTrace {
1798 t: generated,
1799 token_id: t_next,
1800 confidence: confidence.last().copied().unwrap_or(0.0),
1801 active_skill: skill,
1802 recon: None,
1803 switched: false,
1804 });
1805 }
1806 if !commit!(t_next) {
1807 break 'decode;
1808 }
1809 if generated >= max_tokens {
1810 break 'decode;
1811 }
1812
1813 if self.kv_cache.needs_eviction() {
1814 let keep = (self.kv_cache.max_seq_len / 2).max(1);
1815 self.kv_cache.evict(keep);
1816 }
1817
1818 match &mut mtp {
1819 Some(m) if generated + 1 < max_tokens => {
1821 let draft = self.mtp_step(m, &hidden, t_next, next_pos - 1);
1822 drafted += 1;
1823 let emb1 = self.embed_single(t_next);
1824 let emb2 = self.embed_single(draft);
1825 let (h1, h2) = self.forward_pair(&emb1, &emb2, next_pos);
1826
1827 inference::rms_norm_into(
1828 &h1,
1829 &self.weights.final_norm,
1830 self.rms_eps,
1831 self.norm_style,
1832 &mut self.ws.n1,
1833 );
1834 let mut logits1 = self.lm_head_forward(&self.ws.n1);
1835 let t_after = sampler::sample_with_scratch(
1836 &logits1,
1837 &self.sampler_config,
1838 &all_ids,
1839 &mut self.rng,
1840 &mut self.sampler_scratch,
1841 );
1842 if self.confidence_on {
1843 confidence.push(top1_prob_t(&logits1, t_after, calib_temp));
1844 }
1845 attention::recycle_buf(&mut logits1);
1846 if trace_on {
1847 traces.push(TokenTrace {
1850 t: generated,
1851 token_id: t_after,
1852 confidence: confidence.last().copied().unwrap_or(0.0),
1853 active_skill: None,
1854 recon: None,
1855 switched: false,
1856 });
1857 }
1858 let stop = !commit!(t_after);
1859
1860 if t_after == draft {
1861 accepted += 1;
1862 self.commit_linear_scratch();
1863 let _ = self.mtp_step(m, &h1, t_after, next_pos);
1864 hidden = h2;
1865 next_pos += 2;
1866 } else {
1867 for layer in &mut self.kv_cache.layers {
1869 layer.truncate_last(1);
1870 }
1871 if !stop {
1872 let _ = self.mtp_step(m, &h1, t_after, next_pos);
1873 hidden = self.forward_layers(
1874 &self.embed_single(t_after),
1875 next_pos + 1,
1876 None,
1877 );
1878 }
1879 next_pos += 2;
1880 }
1881 if stop {
1882 break 'decode;
1883 }
1884 }
1885 _ => {
1887 #[cfg(feature = "gpu")]
1892 if Self::dsv4_spec_on() && self.dsv4.is_some() {
1893 static SAID: std::sync::Once = std::sync::Once::new();
1894 SAID.call_once(|| {
1895 eprintln!(
1896 "dsv4-spec гейт: mtp={} mask={} router={} trace={} temp={} rep={} ",
1897 !self.dsv4_mtp.is_empty(),
1898 task_mask.is_none(),
1899 router.is_none(),
1900 !trace_on,
1901 self.sampler_config.temperature < 1e-6,
1902 self.sampler_config.repetition_penalty == 1.0,
1903 );
1904 });
1905 }
1906 #[cfg(feature = "gpu")]
1907 if Self::dsv4_spec_on()
1908 && self.dsv4.is_some()
1909 && !self.dsv4_mtp.is_empty()
1910 && task_mask.is_none()
1911 && router.is_none()
1912 && !trace_on
1913 && self.sampler_config.temperature < 1e-6
1914 && self.sampler_config.repetition_penalty == 1.0
1915 && generated + 1 < max_tokens
1916 && all_ids.len() >= 2
1917 {
1918 let tip_token = all_ids[all_ids.len() - 2];
1919 if let Some((extra, n_pos)) = self.dsv4_spec_step(
1920 tip_token,
1921 t_next,
1922 next_pos,
1923 &mut drafted,
1924 &mut accepted,
1925 ) {
1926 next_pos = n_pos;
1927 let mut stopped = false;
1928 for &id in &extra {
1929 if self.confidence_on {
1930 confidence.push(0.0);
1931 }
1932 if !commit!(id) {
1933 stopped = true;
1934 break;
1935 }
1936 }
1937 if stopped {
1938 break 'decode;
1939 }
1940 continue 'decode;
1941 }
1942 }
1943 self.graph_want_logits = fuse_lm;
1944 let mut t_fwd = t_next;
1950 let pure_greedy = self.sampler_config.temperature < 1e-6
1951 && self.sampler_config.repetition_penalty == 1.0
1952 && self.sampler_config.suppress_tokens.is_empty();
1953 let burst_k = std::env::var("CMF_MULTISTEP")
1958 .ok()
1959 .and_then(|v| v.parse::<usize>().ok())
1960 .unwrap_or(0);
1961 if pure_greedy
1962 && burst_k >= 1
1963 && fuse_lm
1964 && task_mask.is_none()
1965 && router.is_none()
1966 && !trace_on
1967 && !self.confidence_on
1968 {
1969 let mut stopped = false;
1970 loop {
1971 let room = max_tokens.saturating_sub(generated);
1972 if room <= 2 {
1973 break;
1974 }
1975 let k = burst_k.min(room - 1);
1976 if k < 1 {
1977 break;
1978 }
1979 let Some(ids) = self.try_multi_burst(t_fwd, next_pos, k) else {
1980 break;
1981 };
1982 next_pos += k;
1983 for &id in &ids {
1984 if !commit!(id) {
1985 stopped = true;
1986 break;
1987 }
1988 }
1989 if stopped {
1990 break;
1991 }
1992 t_fwd = *ids.last().unwrap();
1993 }
1994 if stopped {
1995 break 'decode;
1996 }
1997 }
1998 hidden = self.forward_layers(&self.embed_single(t_fwd), next_pos, task_mask);
1999 next_pos += 1;
2000 if let Some(r) = &mut router {
2003 let phi = self.dyn_phi_ema.clone();
2004 let decision = r.step(&phi, generated);
2005 if let Some(new_active) = decision {
2006 let _ = self.set_active_skill(new_active);
2007 }
2008 if trace_on {
2011 if let Some(last) = traces.last_mut() {
2012 let e = r.last_best_e();
2013 last.recon = e.is_finite().then_some(e);
2014 last.switched = decision.is_some();
2015 }
2016 }
2017 }
2018 }
2019 }
2020 }
2021
2022 self.graph_want_logits = false;
2023 self.graph_logits = None;
2024 if router.is_some() {
2026 let _ = self.set_active_skill(None);
2027 }
2028 self.dyn_router = router.or(self.dyn_router.take());
2029 self.mtp = mtp.or(self.mtp.take());
2030
2031 let output_ids = &all_ids[input_ids.len()..];
2032 let forwarded = input_ids.len() + output_ids.len().saturating_sub(1);
2036 self.kv_history = all_ids[..forwarded.min(all_ids.len())].to_vec();
2037 confidence.truncate(output_ids.len()); traces.truncate(output_ids.len());
2039 Ok(GenerateResult {
2040 text: self.tokenizer.decode(output_ids),
2041 token_ids: output_ids.to_vec(),
2042 prompt_tokens: input_ids.len(),
2043 tokens_generated: generated,
2044 finish_reason,
2045 mtp_drafted: drafted,
2046 mtp_accepted: accepted,
2047 token_confidence: confidence,
2048 traces,
2049 })
2050 }
2051
2052 fn mtp_step(
2056 &mut self,
2057 m: &mut MtpModule,
2058 hidden: &[f32],
2059 next_token: u32,
2060 position: usize,
2061 ) -> u32 {
2062 let e = self.embed_single(next_token);
2066 let mut cat = vec![0.0f32; 2 * self.hidden_size];
2067 let (cat_e, cat_h) = cat.split_at_mut(self.hidden_size);
2068 inference::rms_norm_into(&e, &m.enorm, self.rms_eps, self.norm_style, cat_e);
2069 inference::rms_norm_into(hidden, &m.hnorm, self.rms_eps, self.norm_style, cat_h);
2070 let mut x = vec![0.0f32; self.hidden_size];
2071 m.eh_proj.matvec(&cat, &mut x, self.pool.as_deref());
2072
2073 let lw = &m.layer;
2075 inference::rms_norm_into(
2076 &x,
2077 &lw.input_norm,
2078 self.rms_eps,
2079 self.norm_style,
2080 &mut self.ws.n1,
2081 );
2082 let attn = match &lw.attn {
2083 AttnKind::Mla(_) => unreachable!("MLA has no MTP/pair path"),
2085 AttnKind::Kda(_) => unreachable!("KDA has no MTP/pair path"),
2086 AttnKind::Full {
2087 wq,
2088 wk,
2089 wv,
2090 wo,
2091 q_norm,
2092 k_norm,
2093 output_gate,
2094 softplus_gate,
2095 bias,
2096 } => {
2097 let mut cfg = self.attn_cfg(position);
2098 cfg.q_norm = q_norm.as_deref();
2099 cfg.k_norm = k_norm.as_deref();
2100 cfg.output_gate = *output_gate;
2101 cfg.softplus_gate = softplus_gate
2102 .as_ref()
2103 .map(|(gate, per_head)| (gate, *per_head));
2104 cfg.bias = bias
2105 .as_ref()
2106 .map(|(q, k, v)| (q.as_slice(), k.as_slice(), v.as_slice()));
2107 attention::qwen_attention(&self.ws.n1, wq, wk, wv, wo, &mut m.kv, &cfg)
2108 }
2109 AttnKind::Linear(_) | AttnKind::LinearGdn(_) | AttnKind::ShortConv(_) => {
2110 unreachable!("MTP block is full attention")
2111 }
2112 };
2113 for (i, &a) in attn.iter().enumerate() {
2114 x[i] += a;
2115 }
2116 inference::rms_norm_into(
2117 &x,
2118 &lw.post_norm,
2119 self.rms_eps,
2120 self.norm_style,
2121 &mut self.ws.p1,
2122 );
2123 let ffn = ffn_forward(&lw.ffn, &self.ws.p1, self.pool.as_deref(), None);
2124 for (i, &f) in ffn.iter().enumerate() {
2125 x[i] += f;
2126 }
2127
2128 inference::rms_norm_into(
2129 &x,
2130 &m.final_norm,
2131 self.rms_eps,
2132 self.norm_style,
2133 &mut self.ws.n1,
2134 );
2135 let mut lg = self.lm_head_forward(&self.ws.n1);
2136 let draft = sampler::argmax(&lg);
2137 attention::recycle_buf(&mut lg);
2138 draft
2139 }
2140
2141 pub fn measure_pair_fusion(&mut self, iters: usize) -> (f64, f64) {
2150 if !self.pair_supported() {
2151 return (0.0, 0.0);
2152 }
2153 let emb1 = self.embed_single(1);
2154 let emb2 = self.embed_single(2);
2155 let pos = self.kv_cache.seq_len();
2156
2157 let t0 = std::time::Instant::now();
2158 for _ in 0..iters {
2159 let _ = self.forward_layers(&emb1, pos, None);
2160 let _ = self.forward_layers(&emb2, pos + 1, None);
2161 for l in &mut self.kv_cache.layers {
2162 l.truncate_last(2);
2163 }
2164 }
2165 let singles_ms = t0.elapsed().as_secs_f64() * 1000.0 / iters as f64;
2166
2167 let t1 = std::time::Instant::now();
2168 for _ in 0..iters {
2169 let _ = self.forward_pair(&emb1, &emb2, pos);
2170 for l in &mut self.kv_cache.layers {
2171 l.truncate_last(2);
2172 }
2173 }
2174 let pair_ms = t1.elapsed().as_secs_f64() * 1000.0 / iters as f64;
2175 (singles_ms, pair_ms)
2176 }
2177
2178 fn pair_supported(&self) -> bool {
2186 !self.weights.layers.is_empty()
2193 && self.g3n.is_none()
2194 && !self
2195 .weights
2196 .layers
2197 .iter()
2198 .any(|lw| matches!(&lw.attn, AttnKind::Mla(_) | AttnKind::Kda(_)))
2199 }
2200
2201 fn forward_pair(
2202 &mut self,
2203 emb1: &[f32],
2204 emb2: &[f32],
2205 position: usize,
2206 ) -> (Vec<f32>, Vec<f32>) {
2207 let mut h1 = emb1.to_vec();
2208 let mut h2 = emb2.to_vec();
2209 let (_nkv, _hd, hs, _rd, eps) = (
2210 self.num_kv_heads,
2211 self.head_dim,
2212 self.hidden_size,
2213 self.rotary_dim,
2214 self.rms_eps,
2215 );
2216 let pool = self.pool.clone();
2217
2218 for li in 0..self.num_layers {
2219 let lw = &self.weights.layers[self.phys_layer(li)];
2220 inference::rms_norm_into(
2223 &h1,
2224 &lw.input_norm,
2225 self.rms_eps,
2226 self.norm_style,
2227 &mut self.ws.n1,
2228 );
2229 inference::rms_norm_into(
2230 &h2,
2231 &lw.input_norm,
2232 self.rms_eps,
2233 self.norm_style,
2234 &mut self.ws.n2,
2235 );
2236
2237 let (a1, a2) = match &lw.attn {
2238 AttnKind::Mla(_) => unreachable!("MLA has no MTP/pair path"),
2239 AttnKind::Kda(_) => unreachable!("KDA has no MTP/pair path"),
2240 AttnKind::Linear(w) => {
2241 let cfg = self.vmf_cfg.expect("linear layer without vmf_cfg");
2242 let layer = &mut self.kv_cache.layers[li];
2243 let (state, scratch) = (&mut layer.linear_state, &mut layer.linear_scratch);
2244 vmf_phase_pair(
2245 &self.ws.n1,
2246 &self.ws.n2,
2247 w,
2248 &cfg,
2249 state,
2250 scratch,
2251 self.pool.as_deref(),
2252 )
2253 }
2254 AttnKind::LinearGdn(w) => {
2255 let cfg = self.gdn_cfg.expect("gdn layer without gdn_cfg");
2256 let layer = &mut self.kv_cache.layers[li];
2257 let (state, scratch) = (&mut layer.linear_state, &mut layer.linear_scratch);
2258 gdn_pair(
2259 &self.ws.n1,
2260 &self.ws.n2,
2261 w,
2262 &cfg,
2263 state,
2264 scratch,
2265 self.pool.as_deref(),
2266 )
2267 }
2268 AttnKind::ShortConv(w) => {
2269 let cfg = self
2270 .short_conv_cfg
2271 .expect("short-conv layer without short_conv_cfg");
2272 let layer = &mut self.kv_cache.layers[li];
2273 let (state, scratch) = (&mut layer.linear_state, &mut layer.linear_scratch);
2274 short_conv_pair(
2275 &self.ws.n1,
2276 &self.ws.n2,
2277 w,
2278 &cfg,
2279 state,
2280 scratch,
2281 self.pool.as_deref(),
2282 )
2283 }
2284 AttnKind::Full {
2285 wq,
2286 wk,
2287 wv,
2288 wo,
2289 q_norm,
2290 k_norm,
2291 output_gate,
2292 softplus_gate,
2293 bias,
2294 } => {
2295 let inv_freq_l = self.layer_inv_freq(li);
2296 let (nkv_l, hd_l, rd_l) = self.layer_geom(li);
2297 let cfg = QwenAttnCfg {
2298 num_heads: self.layer_num_heads(li),
2299 num_kv_heads: nkv_l,
2300 head_dim: hd_l,
2301 hidden_size: hs,
2302 position,
2303 inv_freq: &inv_freq_l,
2304 rotary_dim: rd_l,
2305 scale: self.attn_scale,
2306 softcap: self.attn_softcap,
2307 window: self.layer_window(li),
2308 v_norm: self.attn_v_norm,
2309 q_norm: q_norm.as_deref(),
2310 k_norm: k_norm.as_deref(),
2311 output_gate: *output_gate,
2312 softplus_gate: softplus_gate
2313 .as_ref()
2314 .map(|(gate, per_head)| (gate, *per_head)),
2315 rope_scale: self.layer_rope_scale(li),
2316 bias: bias
2317 .as_ref()
2318 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
2319 rms_eps: eps,
2320 norm_style: self.norm_style,
2321 pool: pool.as_deref(),
2322 };
2323 attention::qwen_attention_pair(
2324 &self.ws.n1,
2325 &self.ws.n2,
2326 wq,
2327 wk,
2328 wv,
2329 wo,
2330 &mut self.kv_cache.layers[li],
2331 &cfg,
2332 )
2333 }
2334 };
2335 let (a1, a2) = match &self.weights.layers[self.phys_layer(li)].attn_out_norm {
2336 Some(w) => (
2337 inference::rms_norm(&a1, w, self.rms_eps, self.norm_style),
2338 inference::rms_norm(&a2, w, self.rms_eps, self.norm_style),
2339 ),
2340 None => (a1, a2),
2341 };
2342 for i in 0..self.hidden_size {
2343 h1[i] += a1[i];
2344 h2[i] += a2[i];
2345 }
2346 let (mut a1, mut a2) = (a1, a2);
2347 attention::recycle_buf(&mut a1);
2348 attention::recycle_buf(&mut a2);
2349
2350 let lw = &self.weights.layers[self.phys_layer(li)];
2351 inference::rms_norm_into(
2352 &h1,
2353 &lw.post_norm,
2354 self.rms_eps,
2355 self.norm_style,
2356 &mut self.ws.p1,
2357 );
2358 inference::rms_norm_into(
2359 &h2,
2360 &lw.post_norm,
2361 self.rms_eps,
2362 self.norm_style,
2363 &mut self.ws.p2,
2364 );
2365 let (f1, f2) = match &lw.ffn {
2366 FfnKind::DenseMoe(dm) => (
2369 dense_moe_ffn(
2370 dm,
2371 &self.ws.p1,
2372 &h1,
2373 self.rms_eps,
2374 self.norm_style,
2375 self.pool.as_deref(),
2376 ),
2377 dense_moe_ffn(
2378 dm,
2379 &self.ws.p2,
2380 &h2,
2381 self.rms_eps,
2382 self.norm_style,
2383 self.pool.as_deref(),
2384 ),
2385 ),
2386 _ => ffn_forward_pair(
2387 &lw.ffn,
2388 &self.ws.p1,
2389 &self.ws.p2,
2390 self.pool.as_deref(),
2391 None,
2392 ),
2393 };
2394 let (f1, f2) = match &self.weights.layers[self.phys_layer(li)].ffn_out_norm {
2395 Some(w) => (
2396 inference::rms_norm(&f1, w, self.rms_eps, self.norm_style),
2397 inference::rms_norm(&f2, w, self.rms_eps, self.norm_style),
2398 ),
2399 None => (f1, f2),
2400 };
2401 for i in 0..self.hidden_size {
2402 h1[i] += f1[i];
2403 h2[i] += f2[i];
2404 }
2405 let (mut f1, mut f2) = (f1, f2);
2406 attention::recycle_buf(&mut f1);
2407 attention::recycle_buf(&mut f2);
2408 if let Some(sc) = self.weights.layers[self.phys_layer(li)].layer_scale {
2409 for i in 0..self.hidden_size {
2410 h1[i] *= sc;
2411 h2[i] *= sc;
2412 }
2413 }
2414 if self.is_loop_end(li) && li + 1 < self.num_layers {
2416 h1 = inference::rms_norm(
2417 &h1,
2418 &self.weights.final_norm,
2419 self.rms_eps,
2420 self.norm_style,
2421 );
2422 h2 = inference::rms_norm(
2423 &h2,
2424 &self.weights.final_norm,
2425 self.rms_eps,
2426 self.norm_style,
2427 );
2428 }
2429 }
2430 (h1, h2)
2431 }
2432
2433 fn commit_linear_scratch(&mut self) {
2435 for layer in &mut self.kv_cache.layers {
2436 if !layer.linear_scratch.is_empty() {
2437 std::mem::swap(&mut layer.linear_state, &mut layer.linear_scratch);
2438 layer.linear_scratch.clear();
2439 }
2440 }
2441 }
2442
2443 pub fn forward_ids(
2446 &mut self,
2447 ids: &[u32],
2448 task_mask: Option<&TaskMask>,
2449 ) -> Result<Vec<f32>, String> {
2450 if ids.is_empty() {
2451 return Err("empty id sequence".to_string());
2452 }
2453 self.kv_cache.clear();
2454 self.kv_history.clear();
2455 self.o1_begin();
2456 let mut hidden = vec![0.0f32; self.hidden_size];
2457 let mut pos = 0usize;
2458 if task_mask.is_none() && self.can_prefill_batched() && ids.len() > 2 {
2459 let chunk = prefill_chunk();
2463 let hs = self.hidden_size;
2464 while pos < ids.len() {
2465 let end = (pos + chunk).min(ids.len());
2466 let hb = self.prefill_batch(&ids[pos..end], pos);
2467 hidden.copy_from_slice(&hb[(end - pos - 1) * hs..]);
2468 pos = end;
2469 }
2470 }
2471 if task_mask.is_none()
2475 && !std::env::var("CMF_PAIR").is_ok_and(|v| v == "0")
2476 && self.pair_supported()
2477 {
2478 while pos + 1 < ids.len() {
2479 let e1 = self.embed_single(ids[pos]);
2480 let e2 = self.embed_single(ids[pos + 1]);
2481 let (_, h2) = self.forward_pair(&e1, &e2, pos);
2482 self.commit_linear_scratch();
2483 hidden = h2;
2484 pos += 2;
2485 }
2486 }
2487 while pos < ids.len() {
2488 hidden = self.forward_layers(&self.embed_single(ids[pos]), pos, task_mask);
2489 pos += 1;
2490 }
2491 self.o1_seal();
2495 let normed = inference::rms_norm(
2496 &hidden,
2497 &self.weights.final_norm,
2498 self.rms_eps,
2499 self.norm_style,
2500 );
2501 Ok(self.lm_head_forward(&normed))
2502 }
2503
2504 pub fn ppl_ids(&mut self, ids: &[u32]) -> f64 {
2511 let (nll, cnt) = self.nll_ids_from(ids, 0);
2512 (nll / cnt.max(1) as f64).exp()
2513 }
2514
2515 pub fn probe_ffn_mass(&mut self, ids: &[u32]) -> Vec<Vec<f64>> {
2520 self.kv_cache.clear();
2521 self.kv_history.clear();
2522 FFN_PROBE.with(|p| {
2523 *p.borrow_mut() = Some(vec![vec![0f64; self.intermediate_size]; self.num_layers]);
2524 });
2525 crate::gpu::cpu_scope(|| {
2526 for (pos, &id) in ids.iter().enumerate() {
2527 let emb = self.embed_single(id);
2528 let _ = self.forward_layers(&emb, pos, None);
2529 }
2530 });
2531 self.kv_cache.clear();
2532 self.kv_history.clear();
2533 FFN_PROBE
2534 .with(|p| p.borrow_mut().take())
2535 .unwrap_or_default()
2536 }
2537
2538 pub fn ppl_ids_masked(&mut self, ids: &[u32], mask: &TaskMask) -> f64 {
2542 self.kv_cache.clear();
2543 self.kv_history.clear();
2544 let mut nll = 0f64;
2545 let mut cnt = 0usize;
2546 let mut hidden = vec![0f32; self.hidden_size];
2547 for (pos, &id) in ids.iter().enumerate() {
2548 if pos > 0 {
2549 inference::rms_norm_into(
2550 &hidden,
2551 &self.weights.final_norm,
2552 self.rms_eps,
2553 self.norm_style,
2554 &mut self.ws.n1,
2555 );
2556 let mut logits = self.lm_head_forward(&self.ws.n1);
2557 let max = logits.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
2558 let sum: f64 = logits.iter().map(|&v| ((v - max) as f64).exp()).sum();
2559 let p = ((logits[id as usize] - max) as f64).exp() / sum.max(1e-300);
2560 nll -= p.max(1e-300).ln();
2561 cnt += 1;
2562 attention::recycle_buf(&mut logits);
2563 }
2564 let emb = self.embed_single(id);
2565 hidden = self.forward_layers(&emb, pos, Some(mask));
2566 }
2567 self.kv_cache.clear();
2568 self.kv_history.clear();
2569 (nll / cnt.max(1) as f64).exp()
2570 }
2571
2572 pub fn nll_ids_from(&mut self, ids: &[u32], start: usize) -> (f64, usize) {
2581 self.kv_cache.clear();
2582 self.kv_history.clear();
2583 let mut nll = 0f64;
2584 let mut cnt = 0usize;
2585 if self.can_prefill_batched() {
2586 const CHUNK: usize = 128;
2592 const LM_SUB: usize = 32;
2593 let n = ids.len().saturating_sub(1);
2594 let hs = self.hidden_size;
2595 let rows = self.weights.lm_head.rows();
2596 let mut pos = 0usize;
2597 while pos < n {
2598 let end = (pos + CHUNK).min(n);
2599 let bsz = end - pos;
2600 let hb = self.prefill_batch(&ids[pos..end], pos);
2601 let mut k0 = 0usize;
2602 while k0 < bsz {
2603 let k1 = (k0 + LM_SUB).min(bsz);
2604 let sb = k1 - k0;
2605 if pos + k1 <= start {
2608 k0 = k1;
2609 continue;
2610 }
2611 let mut normed = vec![0.0f32; sb * hs];
2612 for k in 0..sb {
2613 let r = inference::rms_norm(
2614 &hb[(k0 + k) * hs..(k0 + k + 1) * hs],
2615 &self.weights.final_norm,
2616 self.rms_eps,
2617 self.norm_style,
2618 );
2619 normed[k * hs..(k + 1) * hs].copy_from_slice(&r);
2620 }
2621 let mut logits = vec![0.0f32; sb * rows];
2622 self.weights
2623 .lm_head
2624 .matmat(&normed, sb, &mut logits, self.pool.as_deref());
2625 for k in 0..sb {
2626 if pos + k0 + k < start {
2627 continue;
2628 }
2629 let lg = &mut logits[k * rows..k * rows + self.vocab_size.min(rows)];
2630 if let Some(mu) = self.logit_multiplier {
2631 for v in lg.iter_mut() {
2632 *v *= mu;
2633 }
2634 }
2635 if let Some(c) = self.final_softcap {
2639 for v in lg.iter_mut() {
2640 *v = c * (*v / c).tanh();
2641 }
2642 }
2643 let lg = &logits[k * rows..k * rows + self.vocab_size.min(rows)];
2644 let target = ids[pos + k0 + k + 1] as usize;
2645 let max = lg.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
2646 let lse: f64 = lg
2647 .iter()
2648 .map(|&v| ((v - max) as f64).exp())
2649 .sum::<f64>()
2650 .ln()
2651 + max as f64;
2652 nll += lse - lg[target] as f64;
2653 cnt += 1;
2654 if std::env::var("CMF_PPL_TRACE").is_ok() {
2655 let top = lg
2656 .iter()
2657 .enumerate()
2658 .max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
2659 .map(|(i, _)| i)
2660 .unwrap_or(0);
2661 eprintln!(
2662 "BTRACE pos {} target {} nll {:.4} top {} lg_t {:.3} lg_top {:.3}",
2663 pos + k0 + k,
2664 target,
2665 lse - lg[target] as f64,
2666 top,
2667 lg[target],
2668 lg[top]
2669 );
2670 }
2671 }
2672 k0 = k1;
2673 }
2674 pos = end;
2675 }
2676 self.kv_cache.clear();
2677 self.kv_history.clear();
2678 return (nll, cnt);
2679 }
2680 for pos in 0..ids.len().saturating_sub(1) {
2681 let hidden = self.forward_layers(&self.embed_single(ids[pos]), pos, None);
2682 let out_of_band = self.graph_logits.take();
2690 if pos < start {
2691 continue;
2692 }
2693 let logits = match out_of_band {
2694 Some(lg) => lg,
2695 None => {
2696 let normed = inference::rms_norm(
2697 &hidden,
2698 &self.weights.final_norm,
2699 self.rms_eps,
2700 self.norm_style,
2701 );
2702 self.lm_head_forward(&normed)
2706 }
2707 };
2708 let target = ids[pos + 1] as usize;
2709 let max = logits.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
2710 let lse: f64 = logits
2711 .iter()
2712 .map(|&v| ((v - max) as f64).exp())
2713 .sum::<f64>()
2714 .ln()
2715 + max as f64;
2716 let tok_nll = lse - logits[target] as f64;
2717 if std::env::var("CMF_PPL_TRACE").is_ok() && pos < 48 {
2718 let top = logits
2719 .iter()
2720 .enumerate()
2721 .max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
2722 .map(|(i, _)| i)
2723 .unwrap_or(0);
2724 eprintln!(
2725 "pos {pos:3} tgt {target:6} nll {tok_nll:7.3} | top1 {top:6} lg[t]={:.2} lg[top]={:.2}",
2726 logits[target], logits[top]
2727 );
2728 }
2729 nll += tok_nll;
2730 cnt += 1;
2731 }
2732 self.kv_cache.clear();
2733 self.kv_history.clear();
2734 (nll, cnt)
2735 }
2736
2737 pub fn nll_ids_o1(&mut self, ids: &[u32], prefill: usize) -> (f64, usize) {
2753 self.kv_cache.clear();
2754 self.kv_history.clear();
2755 self.o1_begin();
2756 let n = ids.len().saturating_sub(1);
2757 let p = prefill.min(n);
2758 let mut pos = 0usize;
2760 if self.can_prefill_batched() {
2761 const CHUNK: usize = 128;
2762 while pos < p {
2763 let end = (pos + CHUNK).min(p);
2764 let _ = self.prefill_batch(&ids[pos..end], pos);
2765 pos = end;
2766 }
2767 } else {
2768 while pos < p {
2769 let _ = self.forward_layers(&self.embed_single(ids[pos]), pos, None);
2770 pos += 1;
2771 }
2772 }
2773 self.o1_seal();
2774
2775 let mut nll = 0f64;
2776 let mut cnt = 0usize;
2777 for pos in p..n {
2778 let hidden = self.forward_layers(&self.embed_single(ids[pos]), pos, None);
2779 let normed = inference::rms_norm(
2780 &hidden,
2781 &self.weights.final_norm,
2782 self.rms_eps,
2783 self.norm_style,
2784 );
2785 let logits = self.lm_head_forward(&normed);
2789 let target = ids[pos + 1] as usize;
2790 let max = logits.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
2791 let lse: f64 = logits
2792 .iter()
2793 .map(|&v| ((v - max) as f64).exp())
2794 .sum::<f64>()
2795 .ln()
2796 + max as f64;
2797 let tok_nll = lse - logits[target] as f64;
2798 if std::env::var("CMF_PPL_TRACE").is_ok() && pos < 48 {
2799 let top = logits
2800 .iter()
2801 .enumerate()
2802 .max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
2803 .map(|(i, _)| i)
2804 .unwrap_or(0);
2805 eprintln!(
2806 "pos {pos:3} tgt {target:6} nll {tok_nll:7.3} | top1 {top:6} lg[t]={:.2} lg[top]={:.2}",
2807 logits[target], logits[top]
2808 );
2809 }
2810 nll += tok_nll;
2811 cnt += 1;
2812 }
2813 self.kv_cache.clear();
2814 self.kv_history.clear();
2815 (nll, cnt)
2816 }
2817
2818 pub fn calib_ids(&mut self, ids: &[u32], temps: &[f32]) -> (Vec<bool>, Vec<Vec<f32>>) {
2826 self.kv_cache.clear();
2827 self.kv_history.clear();
2828 let n = ids.len().saturating_sub(1);
2829 let mut correct = Vec::with_capacity(n);
2830 let mut pmax = Vec::with_capacity(n);
2831 for pos in 0..n {
2832 let emb = self.embed_single(ids[pos]);
2833 let hidden = self.forward_layers(&emb, pos, None);
2834 let normed = inference::rms_norm(
2835 &hidden,
2836 &self.weights.final_norm,
2837 self.rms_eps,
2838 self.norm_style,
2839 );
2840 let logits = self.lm_head_forward(&normed);
2844 let target = ids[pos + 1] as usize;
2845 let (mut amax, mut mval) = (0usize, f32::NEG_INFINITY);
2846 for (i, &v) in logits.iter().enumerate() {
2847 if v > mval {
2848 mval = v;
2849 amax = i;
2850 }
2851 }
2852 correct.push(amax == target);
2853 let row: Vec<f32> = temps
2854 .iter()
2855 .map(|&t| {
2856 let tt = t.max(1e-3);
2857 let s: f32 = logits.iter().map(|&v| ((v - mval) / tt).exp()).sum();
2858 1.0 / s.max(1e-12) })
2860 .collect();
2861 pmax.push(row);
2862 }
2863 self.kv_cache.clear();
2864 self.kv_history.clear();
2865 (correct, pmax)
2866 }
2867
2868 pub fn ppl_ids_dynamic(&mut self, ids: &[u32]) -> (f64, usize) {
2875 let mut router = match self.dyn_router.take() {
2876 Some(r) => r,
2877 None => return (self.ppl_ids(ids), 0),
2878 };
2879 router.reset();
2880 self.dyn_phi_seen = 0;
2881 let _ = self.set_active_skill(None);
2882
2883 self.kv_cache.clear();
2884
2885 self.kv_history.clear();
2886 let mut nll = 0f64;
2887 let mut cnt = 0usize;
2888 for pos in 0..ids.len().saturating_sub(1) {
2889 let hidden = self.forward_layers(&self.embed_single(ids[pos]), pos, None);
2890 let normed = inference::rms_norm(
2891 &hidden,
2892 &self.weights.final_norm,
2893 self.rms_eps,
2894 self.norm_style,
2895 );
2896 let logits = self.lm_head_forward(&normed);
2900 let target = ids[pos + 1] as usize;
2901 let max = logits.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
2902 let lse: f64 = logits
2903 .iter()
2904 .map(|&v| ((v - max) as f64).exp())
2905 .sum::<f64>()
2906 .ln()
2907 + max as f64;
2908 let tok_nll = lse - logits[target] as f64;
2909 if std::env::var("CMF_PPL_TRACE").is_ok() && pos < 48 {
2910 let top = logits
2911 .iter()
2912 .enumerate()
2913 .max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
2914 .map(|(i, _)| i)
2915 .unwrap_or(0);
2916 eprintln!(
2917 "pos {pos:3} tgt {target:6} nll {tok_nll:7.3} | top1 {top:6} lg[t]={:.2} lg[top]={:.2}",
2918 logits[target], logits[top]
2919 );
2920 }
2921 nll += tok_nll;
2922 cnt += 1;
2923 let phi = self.dyn_phi_ema.clone();
2925 if let Some(new_active) = router.step(&phi, pos) {
2926 let _ = self.set_active_skill(new_active);
2927 }
2928 }
2929 let switches = router.switches.len();
2930 let _ = self.set_active_skill(None);
2931 self.dyn_router = Some(router);
2932 self.kv_cache.clear();
2933 self.kv_history.clear();
2934 ((nll / cnt.max(1) as f64).exp(), switches)
2935 }
2936
2937 pub fn probe_phi(&mut self, ids: &[u32], layer: usize) -> Vec<f32> {
2939 self.kv_cache.clear();
2940 self.kv_history.clear();
2941 let mut acc = vec![0f32; self.hidden_size];
2942 for (pos, &id) in ids.iter().enumerate() {
2943 let h = self.forward_layers_upto(&self.embed_single(id), pos, None, Some(layer));
2944 for (a, v) in acc.iter_mut().zip(&h) {
2945 *a += v;
2946 }
2947 }
2948 let n = ids.len().max(1) as f32;
2949 for a in acc.iter_mut() {
2950 *a /= n;
2951 }
2952 self.kv_cache.clear();
2953 self.kv_history.clear();
2954 acc
2955 }
2956
2957 fn prefill_batch(&mut self, ids: &[u32], start_pos: usize) -> Vec<f32> {
2963 let b = ids.len();
2964 let hs = self.hidden_size;
2965 let mut h: Vec<f32> = vec![0.0; b * hs];
2968 let mut h_ready = false;
2969 let fill_h = |h: &mut Vec<f32>, me: &Self| {
2970 for (bi, &id) in ids.iter().enumerate() {
2971 let e = me.embed_single(id);
2972 h[bi * hs..(bi + 1) * hs].copy_from_slice(&e);
2973 }
2974 };
2975 let (_nkv, _hd, _rd, eps) = (
2976 self.num_kv_heads,
2977 self.head_dim,
2978 self.rotary_dim,
2979 self.rms_eps,
2980 );
2981 let pool = self.pool.clone();
2982 let norm_style = self.norm_style;
2983
2984 #[cfg(target_os = "macos")]
2985 let mut chunk_skip_until = 0usize;
2986 for li in 0..self.num_layers {
2987 crate::gpu::set_layer(li as i64); #[cfg(target_os = "macos")]
2994 {
2995 if li < chunk_skip_until {
2996 continue;
2997 }
2998 if !h_ready && li == 0 && self.weights.embed_tokens.q8_row_parts().is_none() {
3004 fill_h(&mut h, self);
3005 h_ready = true;
3006 }
3007 let ids_for_embed = (!h_ready && li == 0).then_some(ids);
3008 let end = self.chunk_run_gpu(li, &mut h, b, start_pos, ids_for_embed);
3009 if end > li {
3010 h_ready = true;
3011 chunk_skip_until = end;
3012 if self.is_loop_end(end - 1) && end < self.num_layers {
3015 for bi in 0..b {
3016 let normed = inference::rms_norm(
3017 &h[bi * hs..(bi + 1) * hs],
3018 &self.weights.final_norm,
3019 eps,
3020 norm_style,
3021 );
3022 h[bi * hs..(bi + 1) * hs].copy_from_slice(&normed);
3023 }
3024 }
3025 continue;
3026 }
3027 }
3028 if !h_ready {
3029 fill_h(&mut h, self);
3030 h_ready = true;
3031 }
3032 let lw = &self.weights.layers[self.phys_layer(li)];
3033 match &lw.attn {
3035 AttnKind::Kda(w) => {
3036 let cfg = self.kda_cfg.expect("kda layer without kda_cfg");
3038 let mut normed = vec![0.0f32; b * hs];
3039 for bi in 0..b {
3040 inference::rms_norm_into(
3041 &h[bi * hs..(bi + 1) * hs],
3042 &lw.input_norm,
3043 eps,
3044 norm_style,
3045 &mut normed[bi * hs..(bi + 1) * hs],
3046 );
3047 }
3048 let attn = crate::linear_core::kda_forward_batch(
3049 &normed,
3050 b,
3051 w,
3052 &cfg,
3053 &mut self.kv_cache.layers[li].linear_state,
3054 pool.as_deref(),
3055 );
3056 for (dst, &a) in h.iter_mut().zip(&attn) {
3057 *dst += a;
3058 }
3059 }
3060 AttnKind::LinearGdn(w) => {
3061 let cfg = self.gdn_cfg.expect("gdn layer without gdn_cfg");
3063 let mut normed = vec![0.0f32; b * hs];
3064 for bi in 0..b {
3065 let r = inference::rms_norm(
3066 &h[bi * hs..(bi + 1) * hs],
3067 &lw.input_norm,
3068 eps,
3069 norm_style,
3070 );
3071 normed[bi * hs..(bi + 1) * hs].copy_from_slice(&r);
3072 }
3073 let attn = crate::linear_core::gdn_forward_batch(
3074 &normed,
3075 b,
3076 w,
3077 &cfg,
3078 &mut self.kv_cache.layers[li].linear_state,
3079 pool.as_deref(),
3080 );
3081 for (dst, &a) in h.iter_mut().zip(&attn) {
3082 *dst += a;
3083 }
3084 }
3085 AttnKind::ShortConv(w) => {
3086 let cfg = self
3089 .short_conv_cfg
3090 .expect("short-conv layer without short_conv_cfg");
3091 let mut normed = vec![0.0f32; b * hs];
3092 for bi in 0..b {
3093 inference::rms_norm_into(
3094 &h[bi * hs..(bi + 1) * hs],
3095 &lw.input_norm,
3096 eps,
3097 norm_style,
3098 &mut normed[bi * hs..(bi + 1) * hs],
3099 );
3100 }
3101 let attn = short_conv_forward_batch(
3102 &normed,
3103 b,
3104 w,
3105 &cfg,
3106 &mut self.kv_cache.layers[li].linear_state,
3107 pool.as_deref(),
3108 );
3109 for (dst, &a) in h.iter_mut().zip(&attn) {
3110 *dst += a;
3111 }
3112 }
3113 AttnKind::Mla(w) => {
3114 let inv_freq_l = self.layer_inv_freq(li);
3117 let rs = self.layer_rope_scale(li);
3118 let mut normed = vec![0.0f32; hs];
3119 for bi in 0..b {
3120 inference::rms_norm_into(
3121 &h[bi * hs..(bi + 1) * hs],
3122 &lw.input_norm,
3123 eps,
3124 norm_style,
3125 &mut normed,
3126 );
3127 let ao = mla_attention(
3128 w,
3129 &normed,
3130 &mut self.kv_cache.layers[li],
3131 start_pos + bi,
3132 &inv_freq_l,
3133 rs,
3134 eps,
3135 pool.as_deref(),
3136 );
3137 for (dst, &a) in h[bi * hs..(bi + 1) * hs].iter_mut().zip(&ao) {
3138 *dst += a;
3139 }
3140 }
3141 }
3142 AttnKind::Full {
3143 wq,
3144 wk,
3145 wv,
3146 wo,
3147 q_norm,
3148 k_norm,
3149 output_gate,
3150 softplus_gate,
3151 bias,
3152 } => {
3153 let mut normed = vec![0.0f32; b * hs];
3157 for bi in 0..b {
3158 inference::rms_norm_into(
3159 &h[bi * hs..(bi + 1) * hs],
3160 &lw.input_norm,
3161 eps,
3162 norm_style,
3163 &mut normed[bi * hs..(bi + 1) * hs],
3164 );
3165 }
3166 let inv_freq_l = self.layer_inv_freq(li);
3167 let (nkv_l, hd_l, rd_l) = self.layer_geom(li);
3168 let cfg = QwenAttnCfg {
3169 num_heads: self.layer_num_heads(li),
3170 num_kv_heads: nkv_l,
3171 head_dim: hd_l,
3172 hidden_size: hs,
3173 position: start_pos,
3174 inv_freq: &inv_freq_l,
3175 rotary_dim: rd_l,
3176 scale: self.attn_scale,
3177 softcap: self.attn_softcap,
3178 window: self.layer_window(li),
3179 v_norm: self.attn_v_norm,
3180 q_norm: q_norm.as_deref(),
3181 k_norm: k_norm.as_deref(),
3182 output_gate: *output_gate,
3183 softplus_gate: softplus_gate
3184 .as_ref()
3185 .map(|(gate, per_head)| (gate, *per_head)),
3186 rope_scale: self.layer_rope_scale(li),
3187 bias: bias
3188 .as_ref()
3189 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
3190 rms_eps: eps,
3191 norm_style,
3192 pool: pool.as_deref(),
3193 };
3194 let mut attn = attention::qwen_attention_batch(
3195 &normed,
3196 b,
3197 wq,
3198 wk,
3199 wv,
3200 wo,
3201 &mut self.kv_cache.layers[li],
3202 &cfg,
3203 );
3204 if let Some(w) = &lw.attn_out_norm {
3205 for bi in 0..b {
3206 inference::rms_norm_into(
3207 &attn[bi * hs..(bi + 1) * hs],
3208 w,
3209 eps,
3210 norm_style,
3211 &mut normed[bi * hs..(bi + 1) * hs],
3212 );
3213 }
3214 attn.copy_from_slice(&normed);
3215 }
3216 for (dst, &a) in h.iter_mut().zip(&attn) {
3217 *dst += a;
3218 }
3219 }
3220 AttnKind::Linear(w) => {
3221 for bi in 0..b {
3222 let normed = inference::rms_norm(
3223 &h[bi * hs..(bi + 1) * hs],
3224 &lw.input_norm,
3225 eps,
3226 norm_style,
3227 );
3228 vmf_phase_forward(
3229 &normed,
3230 w,
3231 &self.vmf_cfg.expect("linear layer without vmf_cfg"),
3232 &mut self.kv_cache.layers[li].linear_state,
3233 pool.as_deref(),
3234 )
3235 .iter()
3236 .enumerate()
3237 .for_each(|(i, &a)| h[bi * hs + i] += a);
3238 }
3239 }
3240 }
3241
3242 let lw = &self.weights.layers[self.phys_layer(li)];
3244 let mut post = vec![0.0f32; b * hs];
3245 for bi in 0..b {
3246 let r =
3247 inference::rms_norm(&h[bi * hs..(bi + 1) * hs], &lw.post_norm, eps, norm_style);
3248 post[bi * hs..(bi + 1) * hs].copy_from_slice(&r);
3249 }
3250 let mut ffn = match &lw.ffn {
3251 FfnKind::Dense(d) => dense_ffn_batch(d, &post, b, pool.as_deref()),
3252 FfnKind::Moe(m) => moe_ffn_batch(m, &post, b, hs, pool.as_deref(), None),
3253 FfnKind::DenseMoe(dm) => {
3256 let mut out = vec![0.0f32; b * hs];
3257 for bi in 0..b {
3258 let r = dense_moe_ffn(
3259 dm,
3260 &post[bi * hs..(bi + 1) * hs],
3261 &h[bi * hs..(bi + 1) * hs],
3262 eps,
3263 norm_style,
3264 pool.as_deref(),
3265 );
3266 out[bi * hs..(bi + 1) * hs].copy_from_slice(&r);
3267 }
3268 out
3269 }
3270 };
3271 if let Some(w) = &lw.ffn_out_norm {
3272 for bi in 0..b {
3273 inference::rms_norm_into(
3274 &ffn[bi * hs..(bi + 1) * hs],
3275 w,
3276 eps,
3277 norm_style,
3278 &mut post[bi * hs..(bi + 1) * hs],
3279 );
3280 }
3281 ffn.copy_from_slice(&post);
3282 }
3283 for (dst, &f) in h.iter_mut().zip(&ffn) {
3284 *dst += f;
3285 }
3286 if let Some(sc) = lw.layer_scale {
3287 for v in h.iter_mut() {
3288 *v *= sc;
3289 }
3290 }
3291 if let Ok(tp) = std::env::var("CMF_TRACE_POS") {
3292 if let Some(t) = tp.parse::<usize>().ok() {
3293 if t >= start_pos && t < start_pos + b {
3294 let bi = t - start_pos;
3295 let row = &h[bi * hs..(bi + 1) * hs];
3296 let n: f32 = row.iter().map(|x| x * x).sum::<f32>().sqrt();
3297 eprintln!(
3298 "BATCH pos {t} after layer {li}: |h| = {n:.6} h0 {:.6} h1 {:.6}",
3299 row[0], row[1]
3300 );
3301 }
3302 }
3303 }
3304 if std::env::var("CMF_DEBUG_LAYERS").is_ok() {
3308 let row = &h[(b - 1) * hs..b * hs];
3309 let rms =
3310 (row.iter().map(|&v| (v as f64) * (v as f64)).sum::<f64>() / hs as f64).sqrt();
3311 let mx = row.iter().fold(0f32, |m, &v| m.max(v.abs()));
3312 eprintln!(
3313 "layer {li:>3} {:>10} ffn={:<5} rms={rms:>12.4} max={mx:>12.4}",
3314 match &self.weights.layers[self.phys_layer(li)].attn {
3315 AttnKind::LinearGdn(_) => "gdn",
3316 AttnKind::Linear(_) => "vmf",
3317 AttnKind::ShortConv(_) => "conv",
3318 _ => "attn",
3319 },
3320 match &lw.ffn {
3321 FfnKind::Moe(_) => "moe",
3322 FfnKind::Dense(_) => "dense",
3323 FfnKind::DenseMoe(_) => "dense+moe",
3324 },
3325 );
3326 }
3327 if self.is_loop_end(li) && li + 1 < self.num_layers {
3329 for bi in 0..b {
3330 let normed = inference::rms_norm(
3331 &h[bi * hs..(bi + 1) * hs],
3332 &self.weights.final_norm,
3333 eps,
3334 norm_style,
3335 );
3336 h[bi * hs..(bi + 1) * hs].copy_from_slice(&normed);
3337 }
3338 }
3339 if std::env::var("CMF_TRACE_H").is_ok() {
3340 let n = h[..hs].iter().map(|v| v.abs()).sum::<f32>() / hs as f32;
3341 let mx = h[..hs].iter().fold(0.0f32, |a, &v| a.max(v.abs()));
3342 eprintln!(
3343 "layer {li}: mean|h|={n:.4} max|h|={mx:.2} scale={:?}",
3344 lw.layer_scale
3345 );
3346 }
3347 }
3348 crate::gpu::set_layer(-1); h
3350 }
3351
3352 fn embed_single(&self, id: u32) -> Vec<f32> {
3354 let mut out = vec![0.0f32; self.hidden_size];
3355 if (id as usize) < self.weights.embed_tokens.rows() {
3356 self.weights.embed_tokens.row_f32(id as usize, &mut out);
3357 }
3358 if self.embed_multiplier != 1.0 {
3359 for v in out.iter_mut() {
3360 *v *= self.embed_multiplier;
3361 }
3362 }
3363 if self.dsv4.is_some() {
3367 let mut v = vec![0.0f32; self.hidden_size.max(1)];
3368 v[0] = id as f32;
3369 return v;
3370 }
3371 if let Some(b) = &self.g3n {
3374 return b.0.extend_embedding(id, &out, self.pool.as_deref());
3375 }
3376 out
3377 }
3378
3379 #[cfg(target_os = "macos")]
3385 fn chunk_run_gpu(
3386 &mut self,
3387 li0: usize,
3388 h: &mut [f32],
3389 b: usize,
3390 pos0: usize,
3391 embed_ids: Option<&[u32]>,
3392 ) -> usize {
3393 if !crate::gpu::enabled_here()
3397 || std::env::var("CMF_GPU_CHUNK")
3398 .map(|v| v == "0")
3399 .unwrap_or(false)
3400 || b < 32
3401 || self.swa.is_some()
3402 || self.global_attn.is_some()
3403 || self.attn_v_norm
3404 || (self.attn_scale - 1.0 / (self.head_dim as f32).sqrt()).abs() > 1e-9
3405 {
3406 return li0;
3407 }
3408 let Some(model) = self.model.clone() else {
3409 return li0;
3410 };
3411 let inv_freq = self.inv_freq.clone();
3412 let (nh, nkv, hd, hs) = (
3413 self.num_heads,
3414 self.num_kv_heads,
3415 self.head_dim,
3416 self.hidden_size,
3417 );
3418 let loop_end = if self.loop_final_norm {
3422 ((li0 / self.physical_layers) + 1) * self.physical_layers
3423 } else {
3424 self.num_layers
3425 };
3426 let mut layers: Vec<crate::gpu_metal::ChunkLayer> = Vec::new();
3427 let mut stored_at: Vec<usize> = Vec::new();
3428 for li in li0..self.num_layers.min(loop_end) {
3429 let lw = &self.weights.layers[self.phys_layer(li)];
3430 if lw.attn_out_norm.is_some() || lw.ffn_out_norm.is_some() || lw.layer_scale.is_some() {
3431 break;
3432 }
3433 let AttnKind::Full {
3434 wq,
3435 wk,
3436 wv,
3437 wo,
3438 q_norm,
3439 k_norm,
3440 output_gate: false,
3441 softplus_gate: None,
3442 bias,
3443 } = &lw.attn
3444 else {
3445 break;
3446 };
3447 let FfnKind::Dense(d) = &lw.ffn else { break };
3448 if d.act != Act::Silu {
3449 break;
3450 }
3451 fn cw(t: &QTensor) -> Option<(usize, usize, usize, &[f32])> {
3456 t.q8_row_parts()
3457 .or_else(|| t.q4t_parts().map(|(i, r, c)| (i, r, c, &[][..])))
3458 .or_else(|| t.q4tp_parts().map(|(i, r, c)| (i, r, c, &[][..])))
3459 }
3460 let parts = (
3461 cw(wq),
3462 cw(wk),
3463 cw(wv),
3464 cw(wo),
3465 cw(&d.gate_proj),
3466 cw(&d.up_proj),
3467 cw(&d.down_proj),
3468 );
3469 let (Some(pq), Some(pk), Some(pv), Some(po), Some(pg), Some(pu), Some(pd)) = parts
3470 else {
3471 break;
3472 };
3473 let layer = &self.kv_cache.layers[li];
3474 if layer.mode != crate::kv_cache::KvMode::F32 || layer.o1.is_some() {
3475 break;
3476 }
3477 stored_at.push(layer.head_len(0));
3478 layers.push(crate::gpu_metal::ChunkLayer {
3479 model: &model,
3480 kv_id: self.graph_kv_id,
3481 layer: li,
3482 wq: pq,
3483 wk: pk,
3484 wv: pv,
3485 wo: po,
3486 gate: pg,
3487 up: pu,
3488 down: pd,
3489 input_norm: &lw.input_norm,
3490 post_norm: &lw.post_norm,
3491 bias: bias
3492 .as_ref()
3493 .map(|(a, bb, cc)| (a.as_slice(), bb.as_slice(), cc.as_slice())),
3494 q_norm: q_norm.as_deref(),
3495 k_norm: k_norm.as_deref(),
3496 inv_freq: &inv_freq,
3497 rd: self.rotary_dim,
3498 nh,
3499 nkv,
3500 hd,
3501 hs,
3502 inter: d.gate_proj.rows(),
3503 gemma: matches!(self.norm_style, cortiq_core::NormStyle::Gemma),
3504 eps: self.rms_eps as f32,
3505 });
3506 }
3507 if layers.is_empty() {
3508 return li0;
3509 }
3510 let row = nkv * hd;
3511 let mut store: Vec<(Vec<f32>, Vec<f32>, Vec<f32>)> = stored_at
3512 .iter()
3513 .map(|&st| (vec![0f32; b * row], vec![0f32; b * row], vec![0f32; st + b]))
3514 .collect();
3515 let mut io: Vec<crate::gpu_metal::ChunkIo> = Vec::with_capacity(layers.len());
3516 for (i, (ok, ov, oi)) in store.iter_mut().enumerate() {
3517 let li = layers[i].layer;
3518 let layer = &self.kv_cache.layers[li];
3519 io.push(crate::gpu_metal::ChunkIo {
3520 cpu_stored: stored_at[i],
3521 cpu_k: (0..nkv).map(|g| layer.head_keys(g)).collect(),
3522 cpu_v: (0..nkv).map(|g| layer.head_values(g)).collect(),
3523 out_k: ok,
3524 out_v: ov,
3525 imp: oi,
3526 });
3527 }
3528 let n_run = layers.len();
3529 let last = layers.last().map(|l| l.layer + 1).unwrap_or(li0);
3530 let ep = embed_ids.and_then(|ids| {
3533 self.weights
3534 .embed_tokens
3535 .q8_row_parts()
3536 .map(|(idx, rows, _c, rs)| crate::gpu_metal::ChunkEmbed {
3537 idx,
3538 rows,
3539 row_scale: rs,
3540 ids,
3541 mult: self.embed_multiplier,
3542 })
3543 });
3544 if embed_ids.is_some() && ep.is_none() {
3545 return li0;
3546 }
3547 if !crate::gpu_metal::chunk_run_gpu(&layers, &mut io, h, b, pos0, ep.as_ref()) {
3548 return li0;
3549 }
3550 drop(io);
3551 drop(layers);
3552 for (i, (ok, ov, oi)) in store.iter().enumerate().take(n_run) {
3555 let li = li0 + i;
3556 let layer = &mut self.kv_cache.layers[li];
3557 for bi in 0..b {
3558 layer.append(
3559 &ok[bi * row..(bi + 1) * row],
3560 &ov[bi * row..(bi + 1) * row],
3561 &[],
3562 );
3563 }
3564 layer.accumulate_imp(oi);
3565 }
3566 last
3567 }
3568
3569 fn layer_is_local(&self, li: usize) -> bool {
3572 if let Some(layers) = &self.sliding_layers {
3573 return layers.get(li).copied().unwrap_or(false);
3574 }
3575 match self.swa {
3576 Some((_, pattern)) => (li + 1) % pattern.max(1) != 0,
3577 None => false,
3578 }
3579 }
3580
3581 fn layer_inv_freq(&self, li: usize) -> std::sync::Arc<Vec<f32>> {
3584 if self.layer_is_local(li) {
3585 if let Some(f) = &self.inv_freq_local {
3586 return f.clone();
3587 }
3588 } else if let Some(f) = &self.inv_freq_global {
3589 return f.clone();
3590 }
3591 self.inv_freq.clone()
3592 }
3593
3594 fn layer_window(&self, li: usize) -> Option<usize> {
3596 self.swa
3597 .and_then(|(w, _)| self.layer_is_local(li).then_some(w))
3598 }
3599
3600 fn layer_num_heads(&self, li: usize) -> usize {
3601 self.attention_heads_per_layer
3602 .as_ref()
3603 .and_then(|v| v.get(li).copied())
3604 .unwrap_or(self.num_heads)
3605 }
3606
3607 fn layer_rope_scale(&self, li: usize) -> f32 {
3608 if self.layer_is_local(li) {
3609 self.rope_scale_local
3610 } else {
3611 self.rope_scale
3612 }
3613 }
3614
3615 fn layer_geom(&self, li: usize) -> (usize, usize, usize) {
3618 if !self.layer_is_local(li) {
3619 if let Some((ghd, gkv)) = self.global_attn {
3620 return (gkv, ghd, ghd);
3621 }
3622 }
3623 (
3624 self.num_kv_heads,
3625 self.head_dim,
3626 if self.layer_is_local(li) {
3627 self.rotary_dim_local.unwrap_or(self.rotary_dim)
3628 } else {
3629 self.rotary_dim
3630 },
3631 )
3632 }
3633
3634 fn forward_layers(
3636 &mut self,
3637 hidden: &[f32],
3638 position: usize,
3639 task_mask: Option<&TaskMask>,
3640 ) -> Vec<f32> {
3641 self.forward_layers_upto(hidden, position, task_mask, None)
3642 }
3643
3644 fn try_token_graph_wgpu(
3648 &self,
3649 hidden: &[f32],
3650 position: usize,
3651 logits_out: &mut Vec<f32>,
3652 ) -> Option<Vec<f32>> {
3653 self.try_token_graph_wgpu_steps(hidden, position, logits_out, 1, None)
3654 }
3655
3656 fn try_multi_burst(&self, t_next: u32, position: usize, k: usize) -> Option<Vec<u32>> {
3660 if self.o1_active() || self.attn_softcap > 0.0 {
3661 return None;
3662 }
3663 let graph_on = match std::env::var("CMF_GPU_WGPU_GRAPH").ok().as_deref() {
3664 Some("0") => return None,
3665 Some(_) => true,
3666 None => crate::gpu::wgpu_graph_default(),
3667 };
3668 if !graph_on {
3669 return None;
3670 }
3671 let emb = self.embed_single(t_next);
3672 let mut lg = Vec::new();
3673 let mut ids = Vec::new();
3674 self.try_token_graph_wgpu_steps(&emb, position, &mut lg, k, Some(&mut ids))?;
3675 (ids.len() == k).then_some(ids)
3676 }
3677
3678 fn try_token_graph_wgpu_steps(
3682 &self,
3683 hidden: &[f32],
3684 position: usize,
3685 logits_out: &mut Vec<f32>,
3686 steps: usize,
3687 ids_out: Option<&mut Vec<u32>>,
3688 ) -> Option<Vec<f32>> {
3689 let o1_gpu = std::env::var("CMF_O1_GPU").as_deref() == Ok("1");
3692 if (self.o1_active() && !o1_gpu) || self.attn_softcap > 0.0 {
3693 return None;
3697 }
3698 let o1_views: Vec<Option<Vec<crate::nystrom::O1DeviceView<'_>>>> = (0..self.num_layers)
3703 .map(|li| {
3704 if !o1_gpu {
3705 return None;
3706 }
3707 self.kv_cache.layers[self.phys_layer(li)].o1_views()
3708 })
3709 .collect();
3710 if self.o1_active() && o1_gpu {
3711 let want: usize = (0..self.num_layers)
3714 .filter(|li| !matches!(self.kv_cache.layers[self.phys_layer(*li)].o1, None))
3715 .count();
3716 let have = o1_views.iter().filter(|v| v.is_some()).count();
3717 if want == 0 || have != want {
3718 return None;
3719 }
3720 }
3721 let nh = self.num_heads;
3722 let (nkv, hd, rd) = self.layer_geom(0);
3723 let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
3724 let mut layers = Vec::with_capacity(self.num_layers);
3725 let mut model = None;
3726 let dbg = std::env::var("CMF_GRAPH_DEBUG").is_ok();
3727 fn gw(t: &QTensor) -> Option<crate::gpu::GraphW<'_>> {
3728 if let Some((_, i, kind, rs)) = t.graph_weight() {
3729 return Some(crate::gpu::GraphW {
3730 idx: i,
3731 kind,
3732 row_scale: rs,
3733 data: &[],
3734 });
3735 }
3736 t.as_f32().map(|d| crate::gpu::GraphW {
3738 idx: 0,
3739 kind: 4,
3740 row_scale: &[],
3741 data: d,
3742 })
3743 }
3744 for li in 0..self.num_layers {
3745 let lw = &self.weights.layers[self.phys_layer(li)];
3746 if dbg {
3747 let ak = match &lw.attn {
3748 AttnKind::Mla(_) => "Mla".into(),
3749 AttnKind::Full {
3750 output_gate, bias, ..
3751 } => format!("Full gate={output_gate} bias={}", bias.is_some()),
3752 AttnKind::LinearGdn(_) => "LinearGdn".into(),
3753 AttnKind::Kda(_) => "Kda".into(),
3754 AttnKind::Linear(_) => "Linear".into(),
3755 AttnKind::ShortConv(_) => "ShortConv".into(),
3756 };
3757 let fk = match &lw.ffn {
3758 FfnKind::Dense(_) => "Dense",
3759 FfnKind::Moe(_) => "Moe",
3760 FfnKind::DenseMoe(_) => "DenseMoe",
3761 };
3762 eprintln!("graph L{li}: attn={ak} ffn={fk}");
3763 }
3764 let gffn = match &lw.ffn {
3765 FfnKind::DenseMoe(_) => return None, FfnKind::Dense(d) => crate::gpu::GraphFfn::Dense {
3767 gate: gw(&d.gate_proj)?,
3768 up: gw(&d.up_proj)?,
3769 down: gw(&d.down_proj)?,
3770 },
3771 FfnKind::Moe(m) => {
3772 if m.router_sigmoid
3777 || m.expert_bias.is_some()
3778 || m.route_tau.is_some()
3779 || m.mask.is_some()
3780 {
3781 return None;
3782 }
3783 let (se, sg) = m.shared.as_ref()?;
3784 let sgate = gw(sg.as_ref()?)?;
3785 let router = gw(&m.router)?;
3786 let inter = m.experts.first()?.gate_proj.rows();
3787 let mut experts = Vec::with_capacity(m.experts.len() + 1);
3788 let mut q4tp: Option<bool> = None;
3791 let mut gu_q2: Option<bool> = None;
3794 for e in m.experts.iter().chain(std::iter::once(se)) {
3795 if !matches!(e.act, Act::Silu)
3796 || e.gate_proj.rows() != inter
3797 || e.up_proj.rows() != inter
3798 {
3799 return None;
3800 }
3801 let (mm, gi, ui, di, is_p, is_q2) = match e.gate_proj.mapped_q4t() {
3802 Some((mm, gi)) => (
3803 mm,
3804 gi,
3805 e.up_proj.mapped_q4t()?.1,
3806 e.down_proj.mapped_q4t()?.1,
3807 false,
3808 false,
3809 ),
3810 None => match e.gate_proj.mapped_q2tp() {
3811 Some((mm, gi)) => (
3812 mm,
3813 gi,
3814 e.up_proj.mapped_q2tp()?.1,
3815 e.down_proj.mapped_q4tp()?.1,
3816 true,
3817 true,
3818 ),
3819 None => {
3820 let (mm, gi) = e.gate_proj.mapped_q4tp()?;
3821 (
3822 mm,
3823 gi,
3824 e.up_proj.mapped_q4tp()?.1,
3825 e.down_proj.mapped_q4tp()?.1,
3826 true,
3827 false,
3828 )
3829 }
3830 },
3831 };
3832 if *q4tp.get_or_insert(is_p) != is_p || *gu_q2.get_or_insert(is_q2) != is_q2
3833 {
3834 tracing::warn!(
3840 "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."
3841 );
3842 return None;
3843 }
3844 model.get_or_insert_with(|| mm.clone());
3845 experts.push((gi, ui, di));
3846 }
3847 crate::gpu::GraphFfn::Moe {
3848 router,
3849 shared_gate: sgate,
3850 experts,
3851 n_exp: m.experts.len(),
3852 top_k: std::env::var("CMF_TOPK_PROBE")
3858 .ok()
3859 .and_then(|v| v.parse::<usize>().ok())
3860 .filter(|k| *k > 0 && *k <= m.top_k)
3861 .unwrap_or(m.top_k),
3862 inter,
3863 norm_topk: m.norm_topk_prob,
3864 q4tp: q4tp?,
3865 gu_q2: gu_q2.unwrap_or(false),
3866 }
3867 }
3868 };
3869 let attn = match &lw.attn {
3870 AttnKind::Full {
3871 wq,
3872 wk,
3873 wv,
3874 wo,
3875 q_norm,
3876 k_norm,
3877 output_gate,
3878 softplus_gate,
3879 bias,
3880 } => {
3881 if softplus_gate.is_some() || self.attention_heads_per_layer.is_some() {
3882 return None;
3883 }
3884 let (m, _, _, _) = wq.graph_weight()?;
3885 model = Some(m.clone());
3886 crate::gpu::GraphAttn::Full {
3887 wq: gw(wq)?,
3888 wk: gw(wk)?,
3889 wv: gw(wv)?,
3890 wo: gw(wo)?,
3891 q_norm: q_norm.as_deref(),
3892 k_norm: k_norm.as_deref(),
3893 bias: bias
3894 .as_ref()
3895 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
3896 output_gate: *output_gate,
3897 cpu_k: self.kv_cache.layers[li].k_heads(),
3898 cpu_v: self.kv_cache.layers[li].v_heads(),
3899 }
3900 }
3901 AttnKind::LinearGdn(w) => {
3902 let cfg = self.gdn_cfg?;
3903 let (m, _, _, _) = w.in_proj_qkv.graph_weight()?;
3904 model = Some(m.clone());
3905 crate::gpu::GraphAttn::Gdn {
3906 qkv: gw(&w.in_proj_qkv)?,
3907 z: gw(&w.in_proj_z)?,
3908 a: gw(&w.in_proj_a)?,
3909 b: gw(&w.in_proj_b)?,
3910 out: gw(&w.out_proj)?,
3911 conv1d: &w.conv1d,
3912 a_log: &w.a_log,
3913 dt_bias: &w.dt_bias,
3914 norm: &w.norm,
3915 nv: cfg.num_v_heads,
3916 nk: cfg.num_k_heads,
3917 dk: cfg.key_head_dim,
3918 dv: cfg.value_head_dim,
3919 kk: cfg.conv_kernel,
3920 cpu_state: &self.kv_cache.layers[self.phys_layer(li)].linear_state,
3921 }
3922 }
3923 _ => return None,
3924 };
3925 layers.push(crate::gpu::GraphLayer {
3926 input_norm: &lw.input_norm,
3927 attn,
3928 post_norm: &lw.post_norm,
3929 ffn: gffn,
3930 });
3931 }
3932 let model = model?;
3933 let lm_gw = if self.graph_want_logits
3939 && std::env::var("CMF_GPU_LMHEAD")
3940 .map(|v| v != "0")
3941 .unwrap_or(true)
3942 {
3943 self.weights.lm_head.graph_weight().map(|(_, i, kind, rs)| {
3944 (
3945 crate::gpu::GraphW {
3946 idx: i,
3947 kind,
3948 row_scale: rs,
3949 data: &[],
3950 },
3951 self.weights.lm_head.rows(),
3952 )
3953 })
3954 } else {
3955 None
3956 };
3957 let lm = lm_gw.as_ref().map(|(gw, rows)| (gw, *rows));
3958 let emb_gw = if steps > 1 {
3960 self.weights
3961 .embed_tokens
3962 .graph_weight()
3963 .map(|(_, i, kind, rs)| {
3964 (
3965 crate::gpu::GraphW {
3966 idx: i,
3967 kind,
3968 row_scale: rs,
3969 data: &[],
3970 },
3971 self.weights.embed_tokens.rows(),
3972 self.embed_multiplier as f32,
3973 )
3974 })
3975 } else {
3976 None
3977 };
3978
3979 let loop_norm_at: Vec<usize> = if self.loop_final_norm {
3982 (0..self.num_layers - 1)
3983 .filter(|&li| (li + 1) % self.physical_layers == 0)
3984 .collect()
3985 } else {
3986 Vec::new()
3987 };
3988 let mut h = hidden.to_vec();
3989 crate::gpu::forward_token_graph(
3990 &model,
3991 self.graph_kv_id,
3992 &layers,
3993 &o1_views,
3994 self.o1_epoch,
3995 &self.inv_freq,
3996 &mut h,
3997 nh,
3998 nkv,
3999 hd,
4000 rd,
4001 self.hidden_size,
4002 self.intermediate_size,
4003 position,
4004 self.kv_cache.max_seq_len,
4005 gemma,
4006 self.rms_eps as f32,
4007 lm,
4008 &self.weights.final_norm,
4009 logits_out,
4010 &loop_norm_at,
4011 steps,
4012 emb_gw.as_ref().map(|(gw, rows, m)| (gw, *rows, *m)),
4013 ids_out,
4014 )
4015 .then_some(h)
4016 }
4017
4018 fn try_batch_graph_wgpu(&self, hiddens: &mut [f32], positions: &[usize], k: usize) -> bool {
4023 if self.attn_softcap > 0.0 {
4024 return false; }
4026 if self.o1_active() {
4027 return false;
4028 }
4029 let nh = self.num_heads;
4030 let (nkv, hd, rd) = self.layer_geom(0);
4031 let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
4032 fn gw(t: &QTensor) -> Option<crate::gpu::GraphW<'_>> {
4033 if let Some((_, i, kind, rs)) = t.graph_weight() {
4034 return Some(crate::gpu::GraphW {
4035 idx: i,
4036 kind,
4037 row_scale: rs,
4038 data: &[],
4039 });
4040 }
4041 t.as_f32().map(|d| crate::gpu::GraphW {
4042 idx: 0,
4043 kind: 4,
4044 row_scale: &[],
4045 data: d,
4046 })
4047 }
4048 let built: Option<(
4049 Vec<crate::gpu::GraphLayer<'_>>,
4050 std::sync::Arc<cortiq_core::CmfModel>,
4051 )> = (|| {
4052 let mut layers = Vec::with_capacity(self.num_layers);
4053 let mut model = None;
4054 for li in 0..self.num_layers {
4055 let lw = &self.weights.layers[self.phys_layer(li)];
4056 let gffn = match &lw.ffn {
4063 FfnKind::Dense(d) => crate::gpu::GraphFfn::Dense {
4064 gate: gw(&d.gate_proj)?,
4065 up: gw(&d.up_proj)?,
4066 down: gw(&d.down_proj)?,
4067 },
4068 FfnKind::Moe(m) => {
4069 if m.router_sigmoid
4070 || m.expert_bias.is_some()
4071 || m.route_tau.is_some()
4072 || m.mask.is_some()
4073 {
4074 return None;
4075 }
4076 let (se, sg) = m.shared.as_ref()?;
4077 let sgate = gw(sg.as_ref()?)?;
4078 let router = gw(&m.router)?;
4079 let inter = m.experts.first()?.gate_proj.rows();
4080 let mut experts = Vec::with_capacity(m.experts.len() + 1);
4081 let mut q4tp: Option<bool> = None;
4082 for e in m.experts.iter().chain(std::iter::once(se)) {
4083 if !matches!(e.act, Act::Silu)
4084 || e.gate_proj.rows() != inter
4085 || e.up_proj.rows() != inter
4086 {
4087 return None;
4088 }
4089 let (mm, gi, ui, di, is_p) = match e.gate_proj.mapped_q4t() {
4090 Some((mm, gi)) => (
4091 mm,
4092 gi,
4093 e.up_proj.mapped_q4t()?.1,
4094 e.down_proj.mapped_q4t()?.1,
4095 false,
4096 ),
4097 None => {
4098 let (mm, gi) = e.gate_proj.mapped_q4tp()?;
4099 (
4100 mm,
4101 gi,
4102 e.up_proj.mapped_q4tp()?.1,
4103 e.down_proj.mapped_q4tp()?.1,
4104 true,
4105 )
4106 }
4107 };
4108 if *q4tp.get_or_insert(is_p) != is_p {
4109 return None;
4110 }
4111 model.get_or_insert_with(|| mm.clone());
4112 experts.push((gi, ui, di));
4113 }
4114 crate::gpu::GraphFfn::Moe {
4115 router,
4116 shared_gate: sgate,
4117 experts,
4118 n_exp: m.experts.len(),
4119 top_k: m.top_k,
4120 inter,
4121 norm_topk: m.norm_topk_prob,
4122 q4tp: q4tp?,
4123 gu_q2: false,
4126 }
4127 }
4128 _ => return None,
4129 };
4130 let attn = match &lw.attn {
4131 AttnKind::Full {
4132 wq,
4133 wk,
4134 wv,
4135 wo,
4136 q_norm,
4137 k_norm,
4138 output_gate,
4139 softplus_gate,
4140 bias,
4141 } => {
4142 if softplus_gate.is_some() || self.attention_heads_per_layer.is_some() {
4143 return None;
4144 }
4145 let (m, _, _, _) = wq.graph_weight()?;
4146 model = Some(m.clone());
4147 crate::gpu::GraphAttn::Full {
4148 wq: gw(wq)?,
4149 wk: gw(wk)?,
4150 wv: gw(wv)?,
4151 wo: gw(wo)?,
4152 q_norm: q_norm.as_deref(),
4153 k_norm: k_norm.as_deref(),
4154 bias: bias
4155 .as_ref()
4156 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
4157 output_gate: *output_gate,
4158 cpu_k: self.kv_cache.layers[li].k_heads(),
4159 cpu_v: self.kv_cache.layers[li].v_heads(),
4160 }
4161 }
4162 AttnKind::LinearGdn(w) => {
4163 let cfg = self.gdn_cfg?;
4164 let (m, _, _, _) = w.in_proj_qkv.graph_weight()?;
4165 model = Some(m.clone());
4166 crate::gpu::GraphAttn::Gdn {
4167 qkv: gw(&w.in_proj_qkv)?,
4168 z: gw(&w.in_proj_z)?,
4169 a: gw(&w.in_proj_a)?,
4170 b: gw(&w.in_proj_b)?,
4171 out: gw(&w.out_proj)?,
4172 conv1d: &w.conv1d,
4173 a_log: &w.a_log,
4174 dt_bias: &w.dt_bias,
4175 norm: &w.norm,
4176 nv: cfg.num_v_heads,
4177 nk: cfg.num_k_heads,
4178 dk: cfg.key_head_dim,
4179 dv: cfg.value_head_dim,
4180 kk: cfg.conv_kernel,
4181 cpu_state: &self.kv_cache.layers[self.phys_layer(li)].linear_state,
4182 }
4183 }
4184 _ => return None,
4185 };
4186 layers.push(crate::gpu::GraphLayer {
4187 input_norm: &lw.input_norm,
4188 attn,
4189 post_norm: &lw.post_norm,
4190 ffn: gffn,
4191 });
4192 }
4193 Some((layers, model?))
4194 })();
4195 let Some((layers, model)) = built else {
4196 {
4197 use std::sync::atomic::{AtomicBool, Ordering};
4198 static SAID: AtomicBool = AtomicBool::new(false);
4199 if !SAID.swap(true, Ordering::Relaxed) {
4200 tracing::warn!("batch graph: BUILDER refused (layer weights/kinds)");
4201 }
4202 }
4203 return false;
4204 };
4205 crate::gpu::forward_batch_graph(
4206 &model,
4207 self.graph_kv_id,
4208 &layers,
4209 &self.inv_freq,
4210 hiddens,
4211 nh,
4212 nkv,
4213 hd,
4214 rd,
4215 self.hidden_size,
4216 self.intermediate_size,
4217 positions,
4218 self.kv_cache.max_seq_len,
4219 gemma,
4220 self.rms_eps as f32,
4221 k,
4222 )
4223 }
4224
4225 fn draft_probe() -> bool {
4229 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
4230 *ON.get_or_init(|| std::env::var("CMF_DSV4_DRAFT_PROBE").is_ok_and(|v| v != "0"))
4231}
4232
4233 #[cfg(feature = "gpu")]
4245 fn dsv4_spec_on() -> bool {
4246 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
4247 *ON.get_or_init(|| std::env::var("CMF_DSV4_SPEC").map(|v| v != "0").unwrap_or(true))
4248 }
4249
4250 #[cfg(feature = "gpu")]
4257 fn dsv4_spec_step(
4258 &mut self,
4259 tip_token: u32,
4260 t_next: u32,
4261 next_pos: usize,
4262 drafted: &mut usize,
4263 accepted_ctr: &mut usize,
4264 ) -> Option<(Vec<u32>, usize)> {
4265 let t_all = std::time::Instant::now();
4266 if std::env::var("CMF_DSV4_SPEC_TIME").is_ok() {
4267 thread_local! {
4268 static LAST: std::cell::Cell<Option<std::time::Instant>> =
4269 const { std::cell::Cell::new(None) };
4270 }
4271 LAST.with(|l| {
4272 if let Some(prev) = l.get() {
4273 eprintln!("между раундами {:.1} мс", prev.elapsed().as_secs_f64() * 1e3);
4274 }
4275 l.set(Some(std::time::Instant::now()));
4276 });
4277 }
4278 if std::env::var("CMF_DSV4_SPEC_DEBUG").is_ok() {
4279 eprintln!("spec_step: вход pos={next_pos}");
4280 }
4281 let n_layers = self.dsv4.as_ref().map(|b| b.1.len())?;
4282 let cfg = self.dsv4.as_ref().map(|b| b.2)?;
4283 if self.dspark.is_none() {
4285 let t = crate::dsv4::dspark_targets(&self.dsv4_mtp, &cfg, n_layers);
4286 if t.is_empty() {
4287 return None;
4288 }
4289 crate::dsv4::dspark_arm(&t, cfg.dim);
4290 self.dspark = Some(crate::dsv4::DsparkState::new(
4291 self.dsv4_mtp.len(),
4292 &cfg,
4293 t.len(),
4294 ));
4295 }
4296 let targets = crate::dsv4::dspark_targets(&self.dsv4_mtp, &cfg, n_layers);
4297 let pack = crate::dsv4::dspark_pack_get(&self.dsv4_mtp, &cfg);
4298 if pack.is_none() && std::env::var("CMF_DSV4_SPEC_DEBUG").is_ok() {
4299 eprintln!("spec_step: пак не построился (targets {targets:?})");
4300 }
4301 let pack = pack?;
4302 let block = crate::dsv4::dspark_block();
4303 let b_box = self.dsv4.as_mut()?;
4304 let (g, layers, st) = (&b_box.0, &b_box.1, &mut b_box.3);
4305 let ds = self.dspark.as_mut()?;
4306 let dbg = std::env::var("CMF_DSV4_SPEC_DEBUG").is_ok();
4309 if !crate::dsv4::dspark_take(&mut ds.main_hidden) && !ds.have_hidden {
4310 if dbg {
4311 eprintln!("spec_step: нет захвата");
4312 }
4313 return None;
4314 }
4315 ds.have_hidden = true;
4316 let tip_pos = next_pos.checked_sub(1)?;
4317 let draft_started = std::time::Instant::now();
4318 let mut conf = Vec::new();
4319 let props = crate::dsv4::dspark_draft_gpu(
4320 g,
4321 &self.dsv4_mtp,
4322 &cfg,
4323 ds,
4324 pack,
4325 st.kv_id,
4326 tip_token,
4327 tip_pos,
4328 self.pool.as_deref(),
4329 &mut conf,
4330 );
4331 self.dspark_draft_ns += draft_started.elapsed().as_nanos();
4332 *drafted += block;
4333 if props.is_empty() || props[0] != t_next {
4334 if dbg {
4335 eprintln!(
4336 "spec_step: черновик {} (props0={:?} t_next={t_next})",
4337 if props.is_empty() { "пуст" } else { "мимо" },
4338 props.first()
4339 );
4340 }
4341 return None;
4342 }
4343 let mut k_verify = crate::dsv4::dspark_verify_k().min(props.len());
4344 let conf_min = {
4350 static M: std::sync::OnceLock<f32> = std::sync::OnceLock::new();
4351 *M.get_or_init(|| {
4352 std::env::var("CMF_DSPARK_CONF_MIN")
4353 .ok()
4354 .and_then(|v| v.parse().ok())
4355 .unwrap_or(0.0)
4356 })
4357 };
4358 if conf_min > 0.0 && conf.len() >= props.len() {
4359 let mut keep = 1usize;
4360 while keep < k_verify && conf.get(keep).copied().unwrap_or(0.0) >= conf_min {
4361 keep += 1;
4362 }
4363 k_verify = k_verify.min(keep.max(2));
4364 }
4365 if k_verify < 2 {
4366 return None;
4367 }
4368 let mut fed = Vec::with_capacity(k_verify);
4369 fed.push(t_next);
4370 fed.extend_from_slice(&props[1..k_verify]);
4371 let mut argmax = Vec::new();
4372 let mut logits_all = Vec::new();
4373 let mut walked = Vec::new();
4374 let txn = crate::dsv4::dsv4_verify_chunk(
4375 g,
4376 layers,
4377 &cfg,
4378 st,
4379 &fed,
4380 next_pos,
4381 &self.inv_freq,
4382 self.pool.as_deref(),
4383 &targets,
4384 &mut argmax,
4385 &mut logits_all,
4386 &mut walked,
4387 );
4388 if txn.is_none() && dbg {
4389 eprintln!("spec_step: verify отказал");
4390 }
4391 let txn = txn?;
4392 let b = fed.len();
4393 let mut accepted = 1usize;
4394 while accepted < b && fed[accepted] == argmax[accepted - 1] {
4395 accepted += 1;
4396 }
4397 if std::env::var("CMF_DSV4_SPEC_FORCE_REJECT").is_ok_and(|v| v != "0") {
4402 accepted = 1;
4403 }
4404 if std::env::var("CMF_DSV4_SPEC_TRACE").is_ok() {
4405 eprintln!(
4406 "spec@{next_pos}: fed={fed:?} argmax={argmax:?} accepted={accepted}"
4407 );
4408 }
4409 let t_fin = std::time::Instant::now();
4410 if !crate::dsv4::dsv4_spec_finish(
4411 g,
4412 layers,
4413 &cfg,
4414 st,
4415 txn,
4416 accepted,
4417 &fed,
4418 &self.inv_freq,
4419 self.pool.as_deref(),
4420 ) {
4421 tracing::warn!("dsv4: спекулятивный откат не удался — состояние подозрительно");
4422 return None;
4423 }
4424 if std::env::var("CMF_DSV4_SPEC_TIME").is_ok() {
4425 eprintln!("finish(k={accepted}): {:.1} мс", t_fin.elapsed().as_secs_f64() * 1e3);
4426 }
4427 *accepted_ctr += accepted - 1;
4428 let (hc, dim) = (cfg.hc_mult, cfg.dim);
4433 let dev_caps: Vec<usize> = targets
4440 .iter()
4441 .copied()
4442 .filter(|&t| {
4443 st.dev_set.get(t).copied().unwrap_or(false)
4444 && !st.partial_set.get(t).copied().unwrap_or(false)
4445 })
4446 .collect();
4447 let mut caps_all = vec![0.0f32; dev_caps.len() * b * hc * dim];
4448 if !crate::gpu_wgpu::dsv4_spec_cap_read_all(b, dev_caps.len(), hc * dim, &mut caps_all) {
4449 return None;
4450 }
4451 for t in 0..accepted {
4452 let tip = t + 1 == accepted;
4453 for (slot, &tl) in targets.iter().enumerate() {
4454 if let Some(di) = dev_caps.iter().position(|&d| d == tl) {
4455 let lo = (di * b + t) * hc * dim;
4456 crate::dsv4::dspark_capture(
4457 &caps_all[lo..lo + hc * dim],
4458 &cfg,
4459 slot,
4460 &mut ds.main_hidden,
4461 );
4462 } else if tip
4463 && crate::dsv4::dspark_peek_slot(slot, dim, {
4464 let lo = slot * dim;
4465 &mut ds.main_hidden[lo..lo + dim]
4466 })
4467 {
4468 } else {
4473 crate::dsv4::dspark_capture(
4477 &walked[t * hc * dim..(t + 1) * hc * dim],
4478 &cfg,
4479 slot,
4480 &mut ds.main_hidden,
4481 );
4482 }
4483 }
4484 crate::dsv4::dspark_ring_append(g, &self.dsv4_mtp, &cfg, ds, next_pos + t, self.pool.as_deref());
4485 }
4486 let row = logits_all[(accepted - 1) * cfg.vocab..accepted * cfg.vocab].to_vec();
4487 self.graph_logits = Some(row);
4488 if std::env::var("CMF_DSV4_TRUNK_PICK_DUMP").is_ok() {
4493 crate::dsv4::trunk_freq_note(&crate::dsv4::pick_tally_take());
4494 crate::dsv4::pick_tally_arm();
4495 }
4496 if std::env::var("CMF_DSV4_SPEC_TIME").is_ok() {
4497 eprintln!("spec_step total {:.1} мс (k={accepted})", t_all.elapsed().as_secs_f64() * 1e3);
4498 }
4499 Some((fed[1..accepted].to_vec(), next_pos + accepted))
4500 }
4501
4502 fn dspark_probe(&mut self, position: usize, token_id: u32) {
4503 if self.dsv4_mtp.is_empty() || !Self::draft_probe() {
4504 return;
4505 }
4506 let trunk_now = crate::dsv4::pick_tally_take();
4508 crate::dsv4::trunk_freq_note(&trunk_now);
4509 if !trunk_now.is_empty() {
4510 self.dspark_trunk_picks.push(trunk_now);
4511 let keep = crate::dsv4::dspark_block();
4512 if self.dspark_trunk_picks.len() > keep {
4513 self.dspark_trunk_picks.remove(0);
4514 }
4515 }
4516 for p in std::mem::take(&mut self.dspark_pending) {
4519 let Some(i) = position.checked_sub(p.0 + 1) else {
4520 continue;
4521 };
4522 let mut p = p;
4523 if i < p.1.len() {
4524 if p.2 && p.1[i] == token_id {
4525 p.3 = i + 1;
4526 } else {
4527 p.2 = false;
4528 }
4529 if i + 1 < p.1.len() {
4530 self.dspark_pending.push(p);
4531 continue;
4532 }
4533 }
4534 self.dspark_hist.push(p.3);
4535 self.dspark_real.push(token_id);
4536 }
4537 let Some(b) = &mut self.dsv4 else { return };
4538 let (g, layers, cfg) = (&b.0, &b.1, b.2);
4539 let n_layers = layers.len();
4540 if self.dspark.is_none() {
4541 let t = crate::dsv4::dspark_targets(&self.dsv4_mtp, &cfg, n_layers);
4542 if t.is_empty() {
4543 return;
4544 }
4545 eprintln!("DSpark: захват со слоёв {t:?}, блок {}", crate::dsv4::dspark_block());
4546 crate::dsv4::dspark_arm(&t, cfg.dim);
4547 self.dspark = Some(crate::dsv4::DsparkState::new(
4548 self.dsv4_mtp.len(),
4549 &cfg,
4550 t.len(),
4551 ));
4552 }
4553 let ds = self.dspark.as_mut().unwrap();
4554 if !crate::dsv4::dspark_take(&mut ds.main_hidden) {
4555 return; }
4557 let mut conf = Vec::new();
4558 crate::dsv4::pick_tally_arm();
4559 let draft_started = std::time::Instant::now();
4564 #[cfg(feature = "gpu")]
4565 let gpu_draft = crate::dsv4::dspark_gpu_on();
4566 #[cfg(not(feature = "gpu"))]
4567 let gpu_draft = false;
4568 let props = if gpu_draft {
4569 #[cfg(feature = "gpu")]
4570 {
4571 let kv_id = b.3.kv_id;
4572 match crate::dsv4::dspark_pack_get(&self.dsv4_mtp, &cfg) {
4573 Some(pk) => crate::dsv4::dspark_draft_gpu(
4574 g,
4575 &self.dsv4_mtp,
4576 &cfg,
4577 ds,
4578 pk,
4579 kv_id,
4580 token_id,
4581 position,
4582 self.pool.as_deref(),
4583 &mut conf,
4584 ),
4585 None => Vec::new(),
4586 }
4587 }
4588 #[cfg(not(feature = "gpu"))]
4589 Vec::new()
4590 } else {
4591 crate::gpu::cpu_scope(|| {
4592 crate::dsv4::dspark_draft(
4593 g,
4594 &self.dsv4_mtp,
4595 &cfg,
4596 ds,
4597 token_id,
4598 position,
4599 self.pool.as_deref(),
4600 &mut conf,
4601 )
4602 })
4603 };
4604 self.dspark_draft_ns += draft_started.elapsed().as_nanos();
4605 let draft_picks = crate::dsv4::pick_tally_take();
4606 crate::dsv4::dspark_freq_note(&draft_picks);
4607 crate::dsv4::pick_tally_arm();
4610 if !props.is_empty() {
4611 let (tu, tt) = {
4615 let flat: Vec<(usize, Vec<usize>)> = self
4616 .dspark_trunk_picks
4617 .iter()
4618 .flat_map(|v| v.iter().cloned())
4619 .collect();
4620 let mut per: std::collections::HashMap<usize, Vec<usize>> =
4622 std::collections::HashMap::new();
4623 for (li, picks) in flat {
4624 per.entry(li).or_default().extend(picks);
4625 }
4626 let n = per.len().max(1);
4627 let mut u = 0usize;
4628 let mut t = 0usize;
4629 for (_, v) in per {
4630 t += v.len();
4631 u += v.iter().collect::<std::collections::HashSet<_>>().len();
4632 }
4633 (u / n, t / n)
4634 };
4635 let (du, dt) = crate::dsv4::tally_unique(&draft_picks);
4636 self.dspark_exp.push((tu, tt, du, dt));
4637 self.dspark_pending.push((position, props, true, 0));
4638 }
4639 if self.dspark_hist.len() >= 8 && self.dspark_hist.len() % 8 == 0 {
4640 let n = self.dspark_hist.len() as f32;
4641 let mean: f32 = self.dspark_hist.iter().sum::<usize>() as f32 / n;
4642 let block = crate::dsv4::dspark_block();
4643 let mut at = vec![0usize; block + 1];
4644 for &k in &self.dspark_hist {
4645 at[k] += 1;
4646 }
4647 let mut surv = Vec::with_capacity(block);
4649 for i in 1..=block {
4650 let k = at[i..].iter().sum::<usize>() as f32 / n;
4651 surv.push(format!("{k:.2}"));
4652 }
4653 let distinct = self
4654 .dspark_real
4655 .iter()
4656 .collect::<std::collections::HashSet<_>>()
4657 .len();
4658 let (tu, tt, du, dt) = self.dspark_exp.iter().fold((0, 0, 0, 0), |a, b| {
4659 (a.0 + b.0, a.1 + b.1, a.2 + b.2, a.3 + b.3)
4660 });
4661 let m = self.dspark_exp.len().max(1);
4662 eprintln!(
4663 "DSpark: черновиков {}, принято в среднем {mean:.2} из {block} \
4664 (токенов за проход {:.2}), распределение {at:?}, выживание [{}]",
4665 self.dspark_hist.len(),
4666 mean + 1.0,
4667 surv.join(" ")
4668 );
4669 eprintln!(
4670 "DSpark: разных токенов {distinct} из {} (вырожденность), \
4671 эксперты ствол {}/{} на слой за {block} токенов, \
4672 черновик {}/{} за блок, draft {:.2} мс/блок",
4673 self.dspark_real.len(),
4674 tu / m,
4675 tt / m,
4676 du / m,
4677 dt / m,
4678 self.dspark_draft_ns as f64 / self.dspark_exp.len().max(1) as f64 / 1e6
4679 );
4680 }
4681 }
4682
4683 fn forward_layers_upto(
4684 &mut self,
4685 hidden: &[f32],
4686 position: usize,
4687 task_mask: Option<&TaskMask>,
4688 upto: Option<usize>,
4689 ) -> Vec<f32> {
4690 if let Some(b) = &mut self.dsv4 {
4696 let _ = (task_mask, upto);
4697 let token_id = hidden.first().copied().unwrap_or(0.0) as u32;
4698 let (g, layers, cfg, st) = (&b.0, &b.1, b.2, &mut b.3);
4699 st.pos = position;
4700 let mut logits = Vec::new();
4701 crate::dsv4::forward_token(
4702 g,
4703 layers,
4704 &cfg,
4705 st,
4706 token_id,
4707 &self.inv_freq,
4708 self.pool.as_deref(),
4709 &mut logits,
4710 );
4711 self.graph_logits = Some(logits);
4712 self.dspark_probe(position, token_id);
4713 return vec![0.0; self.hidden_size];
4716 }
4717 if let Some(b) = &self.g3n {
4720 let _ = (task_mask, upto);
4721 return crate::g3n::g3n_forward(
4722 &b.0,
4723 &b.1,
4724 hidden,
4725 position,
4726 &mut self.kv_cache.layers,
4727 self.num_heads,
4728 self.num_kv_heads,
4729 self.head_dim,
4730 self.pool.as_deref(),
4731 );
4732 }
4733 let mut h = hidden.to_vec();
4734 let (nh, _nkv, _hd, hs, _rd, eps) = (
4737 self.num_heads,
4738 self.num_kv_heads,
4739 self.head_dim,
4740 self.hidden_size,
4741 self.rotary_dim,
4742 self.rms_eps,
4743 );
4744 let pool = self.pool.clone();
4745 let graph_env = std::env::var("CMF_GPU_WGPU_GRAPH").ok();
4757 let graph_on = match graph_env.as_deref() {
4758 Some("0") => false,
4759 Some(_) => true,
4760 None => crate::gpu::wgpu_graph_default(),
4766 };
4767 let graph_trusted =
4768 graph_env.is_some() || crate::gpu::wgpu_graph_default() || self.gdn_cfg.is_some();
4769 let race_eligible = graph_on && upto.is_none() && task_mask.is_none();
4770 if race_eligible && crate::gpu::graph_race_use_graph(graph_trusted) {
4771 let t_graph = std::time::Instant::now();
4772 let mut lg = Vec::new();
4773 let built = self.try_token_graph_wgpu(hidden, position, &mut lg);
4774 graph_note(built.is_some());
4775 if let Some(hh) = built {
4776 let dur = t_graph.elapsed();
4777 if std::env::var("CMF_GRAPH_PROF").is_ok() {
4778 eprintln!("graph-call: {:.2} ms total", dur.as_secs_f64() * 1000.0);
4779 }
4780 if graph_trusted || !crate::gpu::graph_race_first_token_hopeless(dur) {
4781 if !graph_trusted {
4782 crate::gpu::graph_race_record(true, dur);
4783 }
4784 if !lg.is_empty() {
4785 lg.resize(self.vocab_size, 0.0);
4788 if let Some(c) = self.final_softcap {
4789 for l in lg.iter_mut() {
4790 *l = c * (*l / c).tanh();
4791 }
4792 }
4793 self.graph_logits = Some(lg);
4794 }
4795 return hh;
4796 }
4797 }
4803 }
4804 let t_race_cpu = (race_eligible && !graph_trusted).then(std::time::Instant::now);
4805
4806 #[cfg(target_os = "macos")]
4807 let mut gpu_skip_until = 0usize;
4808 for li in 0..self.num_layers {
4809 crate::gpu::set_layer(li as i64); if let Some(u) = upto {
4811 if li > u {
4812 break;
4813 }
4814 }
4815 if let Some(mask) = task_mask {
4816 if !mask.layer_alive(li) {
4817 continue; }
4819 }
4820 #[cfg(target_os = "macos")]
4824 {
4825 if li < gpu_skip_until {
4826 continue;
4827 }
4828 if task_mask.is_none() {
4829 let end = self.q1_graph_gpu(li, upto, position, &mut h);
4830 if end > li {
4831 gpu_skip_until = end;
4832 if self.is_loop_end(end - 1) && end < self.num_layers {
4835 h = inference::rms_norm(
4836 &h,
4837 &self.weights.final_norm,
4838 self.rms_eps,
4839 self.norm_style,
4840 );
4841 }
4842 continue;
4843 }
4844 }
4845 }
4846
4847 let lw = &self.weights.layers[self.phys_layer(li)];
4848 if let Ok(tp) = std::env::var("CMF_TRACE_POS") {
4849 if tp.parse::<usize>().ok() == Some(position) {
4850 let n: f32 = h.iter().map(|x| x * x).sum::<f32>().sqrt();
4851 eprintln!(
4852 "TRACE pos {position} layer {li}: |h| = {n:.6} h0 {:.6} h1 {:.6}",
4853 h[0], h[1]
4854 );
4855 }
4856 }
4857 inference::rms_norm_into(
4860 &h,
4861 &lw.input_norm,
4862 self.rms_eps,
4863 self.norm_style,
4864 &mut self.ws.n1,
4865 );
4866
4867 let attn_out = match &lw.attn {
4868 AttnKind::Mla(w) => {
4869 let inv_freq_l = self.layer_inv_freq(li);
4870 let rs = self.layer_rope_scale(li);
4871 let eps = self.rms_eps;
4872 let pool = self.pool.clone();
4873 mla_attention(
4874 w,
4875 &self.ws.n1,
4876 &mut self.kv_cache.layers[li],
4877 position,
4878 &inv_freq_l,
4879 rs,
4880 eps,
4881 pool.as_deref(),
4882 )
4883 }
4884 AttnKind::Linear(w) => {
4885 let cfg = self.vmf_cfg.expect("linear layer without vmf_cfg");
4886 vmf_phase_forward(
4887 &self.ws.n1,
4888 w,
4889 &cfg,
4890 &mut self.kv_cache.layers[li].linear_state,
4891 self.pool.as_deref(),
4892 )
4893 }
4894 AttnKind::Kda(w) => {
4895 let cfg = self.kda_cfg.expect("kda layer without kda_cfg");
4896 crate::linear_core::kda_forward(
4897 &self.ws.n1,
4898 w,
4899 &cfg,
4900 &mut self.kv_cache.layers[li].linear_state,
4901 self.pool.as_deref(),
4902 )
4903 }
4904 AttnKind::LinearGdn(w) => {
4905 let cfg = self.gdn_cfg.expect("gdn layer without gdn_cfg");
4906 gdn_forward(
4907 &self.ws.n1,
4908 w,
4909 &cfg,
4910 &mut self.kv_cache.layers[li].linear_state,
4911 self.pool.as_deref(),
4912 )
4913 }
4914 AttnKind::ShortConv(w) => {
4915 let cfg = self
4916 .short_conv_cfg
4917 .expect("short-conv layer without short_conv_cfg");
4918 short_conv_forward(
4919 &self.ws.n1,
4920 w,
4921 &cfg,
4922 &mut self.kv_cache.layers[li].linear_state,
4923 self.pool.as_deref(),
4924 )
4925 }
4926 AttnKind::Full {
4927 wq,
4928 wk,
4929 wv,
4930 wo,
4931 q_norm,
4932 k_norm,
4933 output_gate,
4934 softplus_gate,
4935 bias,
4936 } if self.kv_cache.layers[li].o1_sealed() => {
4937 let inv_freq_l = self.layer_inv_freq(li);
4940 let (nkv_l, hd_l, rd_l) = self.layer_geom(li);
4941 let cfg = QwenAttnCfg {
4942 num_heads: self.layer_num_heads(li),
4943 num_kv_heads: nkv_l,
4944 head_dim: hd_l,
4945 hidden_size: hs,
4946 position,
4947 inv_freq: &inv_freq_l,
4948 rotary_dim: rd_l,
4949 scale: self.attn_scale,
4950 softcap: self.attn_softcap,
4951 window: None,
4952 v_norm: self.attn_v_norm,
4953 q_norm: q_norm.as_deref(),
4954 k_norm: k_norm.as_deref(),
4955 output_gate: *output_gate,
4956 softplus_gate: softplus_gate
4957 .as_ref()
4958 .map(|(gate, per_head)| (gate, *per_head)),
4959 rope_scale: self.layer_rope_scale(li),
4960 bias: bias
4961 .as_ref()
4962 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
4963 rms_eps: eps,
4964 norm_style: self.norm_style,
4965 pool: pool.as_deref(),
4966 };
4967 attention::qwen_attention_nystrom(
4968 &self.ws.n1,
4969 wq,
4970 wk,
4971 wv,
4972 wo,
4973 &mut self.kv_cache.layers[li],
4974 &cfg,
4975 )
4976 }
4977 AttnKind::Full {
4978 wq,
4979 wk,
4980 wv,
4981 wo,
4982 q_norm,
4983 k_norm,
4984 output_gate,
4985 softplus_gate,
4986 bias,
4987 } => 'attn: {
4988 if graph_on
4991 && !*output_gate
4992 && softplus_gate.is_none()
4993 && self.attention_heads_per_layer.is_none()
4994 && bias.is_none()
4995 && task_mask.is_none()
4996 {
4997 let inv_freq_l = self.layer_inv_freq(li);
4998 let (nkv_l, hd_l, rd_l) = self.layer_geom(li);
4999 let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
5000 if let (Some((gm, qi)), Some((_, ki)), Some((_, vi)), Some((_, oi))) = (
5001 wq.mapped_q1(),
5002 wk.mapped_q1(),
5003 wv.mapped_q1(),
5004 wo.mapped_q1(),
5005 ) {
5006 let gm = gm.clone();
5007 let mut out = vec![0f32; hs];
5008 let cache = &self.kv_cache.layers[li];
5009 if crate::gpu::attn_dropin(
5010 &gm,
5011 self.graph_kv_id,
5012 li,
5013 &self.ws.n1,
5014 qi,
5015 ki,
5016 vi,
5017 oi,
5018 q_norm.as_deref(),
5019 k_norm.as_deref(),
5020 &inv_freq_l,
5021 nh,
5022 nkv_l,
5023 hd_l,
5024 rd_l,
5025 hs,
5026 position,
5027 self.kv_cache.max_seq_len,
5028 gemma,
5029 eps as f32,
5030 cache.k_heads(),
5031 cache.v_heads(),
5032 &mut out,
5033 ) {
5034 break 'attn out;
5035 }
5036 }
5037 }
5038 let masked = task_mask
5039 .map(|m| m.head_flags(li, self.num_heads).iter().any(|&a| !a))
5040 .unwrap_or(false);
5041 let f32_view = (wq.as_f32(), wk.as_f32(), wv.as_f32(), wo.as_f32());
5042 match (masked, f32_view) {
5043 (true, (Some(q), Some(k), Some(v), Some(o))) => {
5046 let active_heads = task_mask.unwrap().head_flags(li, self.num_heads);
5047 attention::multi_head_attention(
5048 &self.ws.n1,
5049 q,
5050 k,
5051 v,
5052 o,
5053 &mut self.kv_cache.layers[li],
5054 self.num_heads,
5055 self.num_kv_heads,
5056 self.head_dim,
5057 self.hidden_size,
5058 position,
5059 &active_heads,
5060 &self.inv_freq,
5061 )
5062 }
5063 (masked, _) => {
5064 if masked {
5065 tracing::warn!(
5066 "layer {li}: head mask on quantized weights not \
5067 supported yet — executing dense"
5068 );
5069 }
5070 let inv_freq_l = self.layer_inv_freq(li);
5071 let (nkv_l, hd_l, rd_l) = self.layer_geom(li);
5072 let cfg = QwenAttnCfg {
5073 num_heads: self.layer_num_heads(li),
5074 num_kv_heads: nkv_l,
5075 head_dim: hd_l,
5076 hidden_size: hs,
5077 position,
5078 inv_freq: &inv_freq_l,
5079 rotary_dim: rd_l,
5080 scale: self.attn_scale,
5081 softcap: self.attn_softcap,
5082 window: self.layer_window(li),
5083 v_norm: self.attn_v_norm,
5084 q_norm: q_norm.as_deref(),
5085 k_norm: k_norm.as_deref(),
5086 output_gate: *output_gate,
5087 softplus_gate: softplus_gate
5088 .as_ref()
5089 .map(|(gate, per_head)| (gate, *per_head)),
5090 rope_scale: self.layer_rope_scale(li),
5091 bias: bias
5092 .as_ref()
5093 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
5094 rms_eps: eps,
5095 norm_style: self.norm_style,
5096 pool: pool.as_deref(),
5097 };
5098 attention::qwen_attention(
5099 &self.ws.n1,
5100 wq,
5101 wk,
5102 wv,
5103 wo,
5104 &mut self.kv_cache.layers[li],
5105 &cfg,
5106 )
5107 }
5108 }
5109 }
5110 };
5111 let attn_out = match &self.weights.layers[self.phys_layer(li)].attn_out_norm {
5114 Some(w) => inference::rms_norm(&attn_out, w, self.rms_eps, self.norm_style),
5115 None => attn_out,
5116 };
5117 let lw = &self.weights.layers[self.phys_layer(li)];
5118 inference::add_rmsnorm_fused_into(
5119 &mut h,
5120 &attn_out,
5121 &lw.post_norm,
5122 self.rms_eps,
5123 self.norm_style,
5124 &mut self.ws.p1,
5125 );
5126 let mut attn_out = attn_out;
5127 attention::recycle_buf(&mut attn_out);
5128 let post_normed = &self.ws.p1;
5129
5130 let ffn_masked = task_mask
5131 .map(|m| m.ffn_active_count(li) < self.intermediate_size)
5132 .unwrap_or(false);
5133 let f32_ffn = match &lw.ffn {
5136 FfnKind::Dense(d) => (
5137 d.gate_proj.as_f32(),
5138 d.up_proj.as_f32(),
5139 d.down_proj.as_f32(),
5140 ),
5141 FfnKind::Moe(_) | FfnKind::DenseMoe(_) => (None, None, None),
5142 };
5143 let ffn_out = match (ffn_masked, f32_ffn) {
5144 (true, (Some(g), Some(u), Some(d))) => {
5145 let active = task_mask.unwrap().ffn_active_indices(li);
5146 inference::sparse_ffn_forward(
5147 post_normed,
5148 g,
5149 u,
5150 d,
5151 self.hidden_size,
5152 self.intermediate_size,
5153 &active,
5154 self.pool.as_deref(),
5155 )
5156 }
5157 (true, _) => match &lw.ffn {
5161 FfnKind::Dense(d) if d.down_proj.sparse_col_ok() => {
5162 let active = task_mask.unwrap().ffn_active_indices(li);
5163 sparse_ffn_quant(
5164 d,
5165 post_normed,
5166 &active,
5167 self.hidden_size,
5168 self.pool.as_deref(),
5169 )
5170 }
5171 FfnKind::Dense(d) => {
5176 let active = task_mask.unwrap().ffn_active_indices(li);
5177 let (gf, uf, df) = dequant_dense_f32(d);
5178 inference::sparse_ffn_forward(
5179 post_normed,
5180 &gf,
5181 &uf,
5182 &df,
5183 self.hidden_size,
5184 self.intermediate_size,
5185 &active,
5186 self.pool.as_deref(),
5187 )
5188 }
5189 FfnKind::Moe(m) => {
5190 let allowed = task_mask.and_then(|tm| tm.expert_flags(li, m.experts.len()));
5194 ffn_forward(
5195 &lw.ffn,
5196 post_normed,
5197 self.pool.as_deref(),
5198 allowed.as_deref(),
5199 )
5200 }
5201 FfnKind::DenseMoe(dm) => dense_moe_ffn(
5202 dm,
5203 post_normed,
5204 &h,
5205 self.rms_eps,
5206 self.norm_style,
5207 self.pool.as_deref(),
5208 ),
5209 },
5210 (false, _) => match &lw.ffn {
5211 FfnKind::DenseMoe(dm) => dense_moe_ffn(
5212 dm,
5213 post_normed,
5214 &h,
5215 self.rms_eps,
5216 self.norm_style,
5217 self.pool.as_deref(),
5218 ),
5219 _ => {
5220 let allowed = match (&lw.ffn, task_mask) {
5221 (FfnKind::Moe(m), Some(tm)) => tm.expert_flags(li, m.experts.len()),
5222 _ => None,
5223 };
5224 ffn_forward(
5225 &lw.ffn,
5226 post_normed,
5227 self.pool.as_deref(),
5228 allowed.as_deref(),
5229 )
5230 }
5231 },
5232 };
5233 let ffn_out = match &self.weights.layers[self.phys_layer(li)].ffn_out_norm {
5234 Some(w) => inference::rms_norm(&ffn_out, w, self.rms_eps, self.norm_style),
5235 None => ffn_out,
5236 };
5237 for (i, &f) in ffn_out.iter().enumerate() {
5238 h[i] += f;
5239 }
5240 let mut ffn_out = ffn_out;
5241 attention::recycle_buf(&mut ffn_out);
5242
5243 if let Some(sc) = self.weights.layers[self.phys_layer(li)].layer_scale {
5245 for v in h.iter_mut() {
5246 *v *= sc;
5247 }
5248 }
5249
5250 if self.is_loop_end(li) && li + 1 < self.num_layers {
5253 h = inference::rms_norm(
5254 &h,
5255 &self.weights.final_norm,
5256 self.rms_eps,
5257 self.norm_style,
5258 );
5259 }
5260
5261 if self.dyn_phi_layer == Some(li) {
5265 self.update_dyn_phi(&h);
5266 }
5267 }
5268 crate::gpu::set_layer(-1); if let Some(t) = t_race_cpu {
5270 crate::gpu::graph_race_record(false, t.elapsed());
5271 }
5272
5273 h
5274 }
5275
5276 fn update_dyn_phi(&mut self, h: &[f32]) {
5279 const A: f32 = 0.2;
5280 if self.dyn_phi_ema.len() != h.len() {
5281 self.dyn_phi_ema = vec![0.0; h.len()];
5282 self.dyn_phi_seen = 0;
5283 }
5284 if self.dyn_phi_seen == 0 {
5285 self.dyn_phi_ema.copy_from_slice(h);
5286 } else {
5287 for (e, &v) in self.dyn_phi_ema.iter_mut().zip(h) {
5288 *e = (1.0 - A) * *e + A * v;
5289 }
5290 }
5291 self.dyn_phi_seen += 1;
5292 }
5293
5294 pub fn dyn_phi(&self) -> &[f32] {
5296 &self.dyn_phi_ema
5297 }
5298
5299 pub fn set_dyn_phi_layer(&mut self, layer: Option<usize>) {
5301 self.dyn_phi_layer = layer;
5302 self.dyn_phi_ema.clear();
5303 self.dyn_phi_seen = 0;
5304 }
5305
5306 pub fn dynamic_skills(&self) -> Vec<(usize, String, usize)> {
5308 let Some(model) = &self.model else {
5309 return Vec::new();
5310 };
5311 model
5312 .header
5313 .skills
5314 .iter()
5315 .enumerate()
5316 .filter_map(|(i, sk)| {
5317 let ok = matches!(self.dyn_skill_layers.get(i), Some(Some(_)));
5318 let sel = sk.selection.as_ref()?;
5319 (ok).then(|| (i, sk.id.clone(), sel.phi_layer))
5320 })
5321 .collect()
5322 }
5323
5324 pub fn active_skill(&self) -> Option<usize> {
5326 self.dyn_active
5327 }
5328
5329 pub fn enable_dynamic_routing(&mut self) -> usize {
5334 use crate::swarm::{DynRouter, RoutableSkill};
5335 let Some(model) = self.model.clone() else {
5336 return 0;
5337 };
5338 if self.dyn_blend_loaded {
5341 tracing::warn!("dynamic routing unavailable on a blend-loaded pipeline");
5342 return 0;
5343 }
5344 if let Some(a) = self.dyn_active {
5348 if !matches!(self.dyn_skill_layers.get(a), Some(Some(_))) {
5349 tracing::warn!("loaded skill is not FFN-eligible — dynamic routing unavailable");
5350 return 0;
5351 }
5352 }
5353 let hidden = self.hidden_size;
5354 let mut skills = Vec::new();
5355 for (idx, id, _phi) in self.dynamic_skills() {
5356 if let Some(sel) = model.header.skills[idx].selection.as_ref() {
5357 if let Some(rs) = RoutableSkill::from_descriptor(idx, id, sel, hidden) {
5358 skills.push(rs);
5359 }
5360 }
5361 }
5362 if skills.is_empty() {
5363 return 0;
5364 }
5365 let phi = skills[0].phi_layer;
5367 if skills.iter().any(|s| s.phi_layer != phi) {
5368 tracing::warn!("routable skills disagree on phi_layer; using {phi}");
5369 }
5370 let n = skills.len();
5371 self.set_dyn_phi_layer(Some(phi));
5372 self.dyn_router = Some(DynRouter::new(skills));
5373 n
5374 }
5375
5376 pub fn route_switches(&self) -> Vec<(usize, Option<String>, Option<String>)> {
5378 self.dyn_router
5379 .as_ref()
5380 .map(|r| r.switches.clone())
5381 .unwrap_or_default()
5382 }
5383
5384 fn lm_head_forward(&self, hidden: &[f32]) -> Vec<f32> {
5387 let rows = self.weights.lm_head.rows();
5388 let mut logits = attention::take_buf(rows.min(self.vocab_size));
5389 self.weights
5390 .lm_head
5391 .matvec(hidden, &mut logits, self.pool.as_deref());
5392 logits.resize(self.vocab_size, 0.0);
5393 if let Some(m) = self.logit_multiplier {
5394 for l in logits.iter_mut() {
5395 *l *= m;
5396 }
5397 }
5398 if let Some(c) = self.final_softcap {
5399 for l in logits.iter_mut() {
5400 *l = c * (*l / c).tanh();
5401 }
5402 }
5403 logits
5404 }
5405
5406 pub fn prefill_next_logits(&mut self, ids: &[u32], task_mask: Option<&TaskMask>) -> Vec<f32> {
5411 self.kv_cache.clear();
5412 self.kv_history.clear();
5413 let mut hidden = vec![0.0f32; self.hidden_size];
5414 for (pos, &id) in ids.iter().enumerate() {
5415 let emb = self.embed_single(id);
5416 hidden = self.forward_layers(&emb, pos, task_mask);
5417 }
5418 inference::rms_norm_into(
5419 &hidden,
5420 &self.weights.final_norm,
5421 self.rms_eps,
5422 self.norm_style,
5423 &mut self.ws.n1,
5424 );
5425 self.lm_head_forward(&self.ws.n1)
5426 }
5427}
5428
5429pub fn create_test_pipeline(
5431 hidden_size: usize,
5432 intermediate_size: usize,
5433 num_heads: usize,
5434 num_kv_heads: usize,
5435 head_dim: usize,
5436 num_layers: usize,
5437 vocab_size: usize,
5438) -> Pipeline {
5439 let synth = |n: usize, salt: usize| -> Vec<f32> {
5442 (0..n)
5443 .map(|i| (((i * 31 + salt * 17 + 7) % 97) as f32 / 97.0 - 0.5) * 0.2)
5444 .collect()
5445 };
5446 let qt = |rows: usize, cols: usize, salt: usize| -> QTensor {
5447 QTensor::from_f32(synth(rows * cols, salt), rows, cols)
5448 };
5449 let layer_weights: Vec<LayerWeights> = (0..num_layers)
5450 .map(|li| LayerWeights {
5451 input_norm: vec![1.0; hidden_size],
5452 post_norm: vec![1.0; hidden_size],
5453 attn_out_norm: None,
5454 ffn_out_norm: None,
5455 layer_scale: None,
5456 ffn: FfnKind::Dense(DenseFfn {
5457 gate_proj: qt(intermediate_size, hidden_size, li * 10 + 5),
5458 up_proj: qt(intermediate_size, hidden_size, li * 10 + 6),
5459 down_proj: qt(hidden_size, intermediate_size, li * 10 + 7),
5460 act: Act::Silu,
5461 }),
5462 attn: AttnKind::Full {
5463 bias: None,
5464 wq: qt(num_heads * head_dim, hidden_size, li * 10 + 1),
5465 wk: qt(num_kv_heads * head_dim, hidden_size, li * 10 + 2),
5466 wv: qt(num_kv_heads * head_dim, hidden_size, li * 10 + 3),
5467 wo: qt(hidden_size, num_heads * head_dim, li * 10 + 4),
5468 q_norm: None,
5469 k_norm: None,
5470 output_gate: false,
5471 softplus_gate: None,
5472 },
5473 })
5474 .collect();
5475
5476 Pipeline::new(
5477 Tokenizer::byte_level(),
5478 PipelineWeights {
5479 embed_tokens: qt(vocab_size, hidden_size, 100),
5480 layers: layer_weights,
5481 lm_head: qt(vocab_size, hidden_size, 200),
5482 final_norm: vec![1.0; hidden_size],
5483 },
5484 hidden_size,
5485 intermediate_size,
5486 num_heads,
5487 num_kv_heads,
5488 head_dim,
5489 num_layers,
5490 num_layers, false, vocab_size,
5493 1e-6,
5494 10_000.0,
5495 NormStyle::Qwen,
5496 4096,
5497 SamplerConfig {
5498 seed: Some(42),
5499 ..Default::default()
5500 },
5501 )
5502}
5503
5504fn dense_ffn_batch(d: &DenseFfn, xs: &[f32], b: usize, pool: Option<&Pool>) -> Vec<f32> {
5507 let inter = d.gate_proj.rows();
5508 let hidden = d.down_proj.rows();
5509 if d.act == Act::Silu && b >= 32 && crate::gpu::enabled_here() && !crate::gpu::mm_killed() {
5515 if let (Some((model, w1)), Some((_, w3)), Some((_, w2))) = (
5516 d.gate_proj.mapped_q4t(),
5517 d.up_proj.mapped_q4t(),
5518 d.down_proj.mapped_q4t(),
5519 ) {
5520 let mut out = vec![0.0f32; b * hidden];
5521 if crate::gpu::q4t_ffn(model, w1, w3, w2, xs, b, hidden, inter, &mut out) {
5522 return out;
5523 }
5524 }
5525 }
5526 let mut g = vec![0.0f32; b * inter];
5527 d.gate_proj.matmat(xs, b, &mut g, pool);
5528 let mut u = vec![0.0f32; b * inter];
5529 d.up_proj.matmat(xs, b, &mut u, pool);
5530 for i in 0..b * inter {
5531 g[i] = d.act.combine(g[i], u[i]);
5532 }
5533 let mut out = vec![0.0f32; b * hidden];
5534 d.down_proj.matmat(&g, b, &mut out, pool);
5535 out
5536}
5537
5538fn accumulate_act(m: &MoeFfn, xs: &[f32], b: usize) {
5543 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
5544 static DUMP: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
5545 let on = *ON.get_or_init(|| std::env::var("CMF_RMS_TRACE").is_ok());
5546 let dump = *DUMP.get_or_init(|| std::env::var("CMF_ACT_DUMP").is_ok());
5547 if (!on && !dump) || b == 0 {
5548 return;
5549 }
5550 let hidden = xs.len() / b;
5551 if on {
5552 let mut acc = m.act_sq.borrow_mut();
5553 if acc.len() < hidden {
5554 acc.resize(hidden, 0.0);
5555 }
5556 for t in 0..b {
5557 let row = &xs[t * hidden..(t + 1) * hidden];
5558 for (a, &v) in acc.iter_mut().zip(row) {
5559 *a += (v as f64) * (v as f64);
5560 }
5561 }
5562 }
5563 if dump {
5564 let cap: usize = std::env::var("CMF_ACT_DUMP_ROWS")
5567 .ok()
5568 .and_then(|v| v.parse().ok())
5569 .unwrap_or(4096);
5570 let mut rows = m.act_rows.borrow_mut();
5571 if rows.len() < cap * hidden {
5572 let take = b.min((cap * hidden - rows.len()) / hidden.max(1));
5573 rows.extend_from_slice(&xs[..take * hidden]);
5574 }
5575 }
5576}
5577
5578fn moe_ffn_batch(
5579 m: &MoeFfn,
5580 xs: &[f32],
5581 b: usize,
5582 hidden: usize,
5583 pool: Option<&Pool>,
5584 allowed: Option<&[bool]>,
5585) -> Vec<f32> {
5586 accumulate_act(m, xs, b);
5587 let ne = m.experts.len();
5588 let mut logits = vec![0.0f32; b * ne];
5589 m.router.matmat(xs, b, &mut logits, pool);
5590
5591 let mut assign: Vec<Vec<(usize, f32)>> = vec![Vec::new(); ne];
5594 {
5595 let mut st = m.stats.borrow_mut();
5596 if st.len() < ne {
5597 st.resize(ne, 0);
5598 }
5599 for bi in 0..b {
5600 let (idx, p, wsum) = moe_route(&logits[bi * ne..(bi + 1) * ne], m, allowed);
5601 for &e in &idx {
5602 st[e] += 1;
5603 assign[e].push((bi, p[e] / wsum));
5604 }
5605 }
5606 }
5607
5608 let mut out = vec![0.0f32; b * hidden];
5609 let cols = m.experts[0].gate_proj.cols();
5610 let mut run_expert = |d: &DenseFfn, list: &[(usize, f32)]| {
5611 let sb = list.len();
5612 let mut sub = vec![0.0f32; sb * cols];
5613 for (k, &(bi, _)) in list.iter().enumerate() {
5614 sub[k * cols..(k + 1) * cols].copy_from_slice(&xs[bi * cols..(bi + 1) * cols]);
5615 }
5616 let eo = dense_ffn_batch(d, &sub, sb, pool);
5617 for (k, &(bi, w)) in list.iter().enumerate() {
5618 for i in 0..hidden {
5619 out[bi * hidden + i] += w * eo[k * hidden + i];
5620 }
5621 }
5622 };
5623 for (e, a) in assign.iter().enumerate().take(ne) {
5624 if !a.is_empty() {
5625 run_expert(&m.experts[e], a);
5626 }
5627 }
5628 if let Some((se, gate)) = &m.shared {
5629 let all: Vec<(usize, f32)> = if let Some(gate) = gate {
5630 let mut gl = vec![0.0f32; b];
5631 gate.matmat(xs, b, &mut gl, pool);
5632 (0..b)
5633 .map(|bi| (bi, 1.0 / (1.0 + (-gl[bi]).exp())))
5634 .collect()
5635 } else {
5636 (0..b).map(|bi| (bi, 1.0)).collect()
5637 };
5638 run_expert(se, &all);
5639 }
5640 out
5641}
5642
5643thread_local! {
5644 static FFN_SCRATCH: std::cell::RefCell<[Vec<f32>; 4]> =
5648 const { std::cell::RefCell::new([Vec::new(), Vec::new(), Vec::new(), Vec::new()]) };
5649}
5650
5651fn dense_ffn(d: &DenseFfn, x: &[f32], pool: Option<&Pool>) -> Vec<f32> {
5653 if crate::gpu::enabled_here()
5664 && (d.gate_proj.rows() >= crate::gpu::min_rows() || d.gate_proj.is_q1())
5665 {
5666 let arm = if d.gate_proj.is_q1() && crate::gpu::q1_force() {
5667 crate::gpu::ProbeArm::Gpu
5668 } else {
5669 crate::gpu::probe_arm(crate::gpu::OpClass::Ffn)
5670 };
5671 match arm {
5672 crate::gpu::ProbeArm::Gpu => {
5673 let t0 = std::time::Instant::now();
5674 if let Some(out) = dense_ffn_gpu(d, x, pool) {
5675 crate::gpu::probe_record(crate::gpu::OpClass::Ffn, true, t0.elapsed());
5676 return out;
5677 }
5678 }
5679 crate::gpu::ProbeArm::CpuTimed => {
5680 let t0 = std::time::Instant::now();
5681 let out = crate::gpu::cpu_scope(|| dense_ffn_cpu(d, x, pool));
5682 crate::gpu::probe_record(crate::gpu::OpClass::Ffn, false, t0.elapsed());
5683 return out;
5684 }
5685 crate::gpu::ProbeArm::Cpu => {
5686 return crate::gpu::cpu_scope(|| dense_ffn_cpu(d, x, pool));
5687 }
5688 }
5689 }
5690 dense_ffn_cpu(d, x, pool)
5691}
5692
5693fn dense_ffn_cpu(d: &DenseFfn, x: &[f32], pool: Option<&Pool>) -> Vec<f32> {
5695 let inter = d.gate_proj.rows();
5696 FFN_SCRATCH.with(|s| {
5697 let mut s = s.borrow_mut();
5698 let [g, u, ..] = &mut *s;
5699 g.resize(inter, 0.0);
5700 if d.act == Act::Silu && QTensor::matvec_silu_mul(&d.gate_proj, &d.up_proj, x, g, pool) {
5703 } else {
5705 u.resize(inter, 0.0);
5706 QTensor::matvec_many([&d.gate_proj, &d.up_proj], x, [g, u], pool);
5708 for i in 0..inter {
5709 g[i] = d.act.combine(g[i], u[i]);
5710 }
5711 }
5712 FFN_PROBE.with(|pr| {
5715 if let Some(acc) = pr.borrow_mut().as_mut() {
5716 let li = crate::gpu::cur_layer();
5717 if li >= 0 {
5718 if let Some(row) = acc.get_mut(li as usize) {
5719 for (a, &v) in row.iter_mut().zip(g.iter()) {
5720 *a += (v as f64).abs();
5721 }
5722 }
5723 }
5724 }
5725 });
5726 let mut out = attention::take_buf(d.down_proj.rows());
5727 d.down_proj.matvec(g, &mut out, pool);
5728 out
5729 })
5730}
5731
5732thread_local! {
5733 static FFN_PROBE: std::cell::RefCell<Option<Vec<Vec<f64>>>> =
5736 const { std::cell::RefCell::new(None) };
5737}
5738
5739fn dense_ffn_gpu(d: &DenseFfn, x: &[f32], _pool: Option<&Pool>) -> Option<Vec<f32>> {
5745 if d.act != Act::Silu {
5747 return None;
5748 }
5749 if d.gate_proj.rows() < crate::gpu::min_rows() && !d.gate_proj.is_q1() {
5752 return None;
5753 }
5754 let mut jobs: Vec<crate::gpu::MoeJob> = Vec::with_capacity(1);
5755 let mut model_ref = None;
5756 moe_push_job(d, x, 1.0, &mut jobs, &mut model_ref)?;
5757 let model = model_ref?;
5758 let hidden = jobs[0].down.1;
5759 let mut out = attention::take_buf(hidden);
5760 if crate::gpu::moe_block(&model, &jobs, &mut out) {
5761 Some(out)
5762 } else {
5763 let mut out = out;
5764 attention::recycle_buf(&mut out);
5765 None
5766 }
5767}
5768
5769#[allow(clippy::type_complexity)]
5774#[allow(clippy::type_complexity)]
5775pub(crate) fn moe_parts(
5776 t: &QTensor,
5777) -> Option<(
5778 &std::sync::Arc<cortiq_core::CmfModel>,
5779 usize,
5780 usize,
5781 usize,
5782 &[f32],
5783 &[f32],
5784 bool,
5785 bool,
5786)> {
5787 match t {
5788 QTensor::Mapped {
5789 model,
5790 idx,
5791 dtype: dt @ (cortiq_core::TensorDtype::Q8_2f | cortiq_core::TensorDtype::Q8Row),
5792 rows,
5793 cols,
5794 row_scale,
5795 col_field,
5796 ..
5797 } if (*dt == cortiq_core::TensorDtype::Q8Row) || !col_field.is_empty() => Some((
5798 model, *idx, *rows, *cols, row_scale, col_field, false, false,
5799 )),
5800 QTensor::Mapped {
5802 model,
5803 idx,
5804 dtype: cortiq_core::TensorDtype::Q1,
5805 rows,
5806 cols,
5807 ..
5808 } => Some((model, *idx, *rows, *cols, &[][..], &[][..], true, false)),
5809 QTensor::Mapped {
5811 model,
5812 idx,
5813 dtype: cortiq_core::TensorDtype::Q4Tiled,
5814 rows,
5815 cols,
5816 ..
5817 } => Some((model, *idx, *rows, *cols, &[][..], &[][..], false, true)),
5818 QTensor::Mapped {
5820 model,
5821 idx,
5822 dtype: cortiq_core::TensorDtype::Q4TiledP,
5823 rows,
5824 cols,
5825 ..
5826 } => Some((model, *idx, *rows, *cols, &[][..], &[][..], false, true)),
5827 _ => None,
5828 }
5829}
5830
5831pub(crate) fn moe_push_job_parts<'a>(
5835 gate: &'a QTensor,
5836 up: &'a QTensor,
5837 down: &'a QTensor,
5838 x: &[f32],
5839 w: f32,
5840 swiglu_limit: f32,
5841 jobs: &mut Vec<crate::gpu::MoeJob<'a>>,
5842 model_ref: &mut Option<std::sync::Arc<cortiq_core::CmfModel>>,
5843) -> Option<()> {
5844 use crate::qtensor::prescale;
5845 let (gm, gi, gr, gc, grs, gcf, gq1, gq4) = moe_parts(gate)?;
5846 let (_, ui, ur, uc, urs, ucf, uq1, uq4) = moe_parts(up)?;
5847 let (_, di, dr, dc, drs, dcf, dq1, dq4) = moe_parts(down)?;
5848 if gq1 != uq1 || uq1 != dq1 || gq4 != uq4 || uq4 != dq4 {
5849 return None; }
5851 model_ref.get_or_insert_with(|| gm.clone());
5852 let dt = |cf: &[f32]| {
5853 if cf.is_empty() {
5854 cortiq_core::TensorDtype::Q8Row
5855 } else {
5856 cortiq_core::TensorDtype::Q8_2f
5857 }
5858 };
5859 jobs.push(crate::gpu::MoeJob {
5860 gate: (gi, gr, gc, grs),
5861 up: (ui, ur, uc, urs),
5862 down: (di, dr, dc, drs),
5863 xs_gate: prescale(x, gcf, dt(gcf)).into_owned(),
5864 xs_up: prescale(x, ucf, dt(ucf)).into_owned(),
5865 down_col: dcf,
5866 w,
5867 q1: gq1,
5868 q4t: gq4 && gate.mapped_q4tp().is_none(),
5869 q4tp: gq4 && gate.mapped_q4tp().is_some(),
5870 swiglu_limit,
5871 });
5872 Some(())
5873}
5874
5875fn moe_push_job<'a>(
5877 d: &'a DenseFfn,
5878 x: &[f32],
5879 w: f32,
5880 jobs: &mut Vec<crate::gpu::MoeJob<'a>>,
5881 model_ref: &mut Option<std::sync::Arc<cortiq_core::CmfModel>>,
5882) -> Option<()> {
5883 use crate::qtensor::prescale;
5884 if d.act != Act::Silu {
5885 return None; }
5887 let (gm, gi, gr, gc, grs, gcf, gq1, gq4) = moe_parts(&d.gate_proj)?;
5888 let (_, ui, ur, uc, urs, ucf, uq1, uq4) = moe_parts(&d.up_proj)?;
5889 let (_, di, dr, dc, drs, dcf, dq1, dq4) = moe_parts(&d.down_proj)?;
5890 if gq1 != uq1 || uq1 != dq1 || gq4 != uq4 || uq4 != dq4 {
5891 return None; }
5893 model_ref.get_or_insert_with(|| gm.clone());
5894 let gdt = if gcf.is_empty() {
5895 cortiq_core::TensorDtype::Q8Row
5896 } else {
5897 cortiq_core::TensorDtype::Q8_2f
5898 };
5899 let udt = if ucf.is_empty() {
5900 cortiq_core::TensorDtype::Q8Row
5901 } else {
5902 cortiq_core::TensorDtype::Q8_2f
5903 };
5904 jobs.push(crate::gpu::MoeJob {
5905 gate: (gi, gr, gc, grs),
5906 up: (ui, ur, uc, urs),
5907 down: (di, dr, dc, drs),
5908 xs_gate: prescale(x, gcf, gdt).into_owned(),
5909 xs_up: prescale(x, ucf, udt).into_owned(),
5910 down_col: dcf,
5911 w,
5912 q1: gq1,
5913 q4t: gq4 && d.gate_proj.mapped_q4tp().is_none(),
5914 q4tp: gq4 && d.gate_proj.mapped_q4tp().is_some(),
5915 swiglu_limit: 0.0,
5916 });
5917 Some(())
5918}
5919
5920fn sparse_ffn_quant(
5927 d: &DenseFfn,
5928 x: &[f32],
5929 active: &[u16],
5930 hidden: usize,
5931 pool: Option<&Pool>,
5932) -> Vec<f32> {
5933 let n = active.len();
5934 let inter = d.gate_proj.rows();
5935 let mut act = vec![0.0f32; n];
5936 let need_scratch = !(d.gate_proj.sparse_col_ok() && d.up_proj.sparse_col_ok());
5939 let compute = |ai: usize| -> f32 {
5940 let idx = active[ai] as usize;
5941 if idx >= inter {
5942 return 0.0; }
5944 let mut s = if need_scratch {
5945 vec![0.0f32; hidden]
5946 } else {
5947 Vec::new()
5948 };
5949 let gate = d.gate_proj.row_dot(idx, x, &mut s);
5950 let up = d.up_proj.row_dot(idx, x, &mut s);
5951 d.act.combine(gate, up)
5952 };
5953 match pool {
5954 Some(p) if n >= 256 => {
5955 let ptr = SendMut(act.as_mut_ptr());
5956 p.run(&|widx, nw| {
5957 let chunk = n.div_ceil(nw);
5958 let (s, e) = (widx * chunk, ((widx + 1) * chunk).min(n));
5959 for ai in s..e {
5960 unsafe { *ptr.at(ai) = compute(ai) };
5961 }
5962 });
5963 }
5964 _ => {
5965 for (ai, a) in act.iter_mut().enumerate() {
5966 *a = compute(ai);
5967 }
5968 }
5969 }
5970 let mut out = vec![0.0f32; hidden];
5972 for (ai, &idx) in active.iter().enumerate() {
5973 let w = act[ai];
5974 if w.abs() >= 1e-12 && (idx as usize) < inter {
5975 d.down_proj.add_col_scaled(idx as usize, w, &mut out);
5976 }
5977 }
5978 out
5979}
5980
5981#[doc(hidden)]
5983pub fn sparse_ffn_quant_for_test(
5984 d: &DenseFfn,
5985 x: &[f32],
5986 active: &[u16],
5987 hidden: usize,
5988) -> Vec<f32> {
5989 sparse_ffn_quant(d, x, active, hidden, None)
5990}
5991
5992fn dequant_dense_f32(d: &DenseFfn) -> (Vec<f32>, Vec<f32>, Vec<f32>) {
5996 let deq = |t: &QTensor| -> Vec<f32> {
5997 let (rows, cols) = (t.rows(), t.cols());
5998 let mut out = vec![0.0f32; rows * cols];
5999 for r in 0..rows {
6000 t.row_f32(r, &mut out[r * cols..(r + 1) * cols]);
6001 }
6002 out
6003 };
6004 (deq(&d.gate_proj), deq(&d.up_proj), deq(&d.down_proj))
6005}
6006
6007struct SendMut(*mut f32);
6009unsafe impl Send for SendMut {}
6010unsafe impl Sync for SendMut {}
6011impl SendMut {
6012 #[inline]
6013 #[allow(clippy::mut_from_ref)]
6016 unsafe fn at(&self, i: usize) -> &mut f32 {
6017 unsafe { &mut *self.0.add(i) }
6018 }
6019}
6020
6021fn moe_route(logits: &[f32], m: &MoeFfn, allowed: Option<&[bool]>) -> (Vec<usize>, Vec<f32>, f32) {
6031 let ne = logits.len();
6032 let p: Vec<f32> = if m.router_sigmoid {
6033 logits.iter().map(|&l| 1.0 / (1.0 + (-l).exp())).collect()
6034 } else {
6035 let mx = logits.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
6036 let mut e: Vec<f32> = logits.iter().map(|&l| (l - mx).exp()).collect();
6037 let s: f32 = e.iter().sum();
6038 for v in &mut e {
6039 *v /= s;
6040 }
6041 e
6042 };
6043 let admit = |e: usize| {
6049 m.mask.as_ref().is_none_or(|mk| mk[e])
6050 && allowed.is_none_or(|a| a.get(e).copied().unwrap_or(false))
6051 };
6052 let mut idx: Vec<usize> = (0..ne).filter(|&e| admit(e)).collect();
6053 match &m.expert_bias {
6055 Some(b) => idx.sort_unstable_by(|&x, &y| {
6056 (p[y] + b[y])
6057 .partial_cmp(&(p[x] + b[x]))
6058 .unwrap()
6059 .then(x.cmp(&y))
6060 }),
6061 None => idx.sort_unstable_by(|&x, &y| p[y].partial_cmp(&p[x]).unwrap().then(x.cmp(&y))),
6062 }
6063 idx.truncate(m.top_k);
6064 if let Some(tau) = m.route_tau {
6068 let total: f32 = idx.iter().map(|&e| p[e]).sum();
6069 if total > 0.0 {
6070 let mut acc = 0.0f32;
6071 let mut keep = idx.len();
6072 for (i, &e) in idx.iter().enumerate() {
6073 acc += p[e];
6074 if acc >= tau * total {
6075 keep = i + 1;
6076 break;
6077 }
6078 }
6079 idx.truncate(keep);
6080 }
6081 }
6082 let wsum: f32 = if m.norm_topk_prob {
6083 let s: f32 = idx.iter().map(|&e| p[e]).sum();
6084 (if m.router_sigmoid { s + 1e-6 } else { s }) / m.routed_scaling
6087 } else {
6088 1.0 / m.routed_scaling
6089 };
6090 (idx, p, wsum)
6091}
6092
6093fn moe_ffn(m: &MoeFfn, x: &[f32], pool: Option<&Pool>, allowed: Option<&[bool]>) -> Vec<f32> {
6096 accumulate_act(m, x, 1);
6097 let ne = m.experts.len();
6098 let mut logits = vec![0.0f32; ne];
6099 m.router.matvec(x, &mut logits, pool);
6100 let (idx, p, wsum) = moe_route(&logits, m, allowed);
6101 {
6102 let mut st = m.stats.borrow_mut();
6103 if st.len() < ne {
6104 st.resize(ne, 0);
6105 }
6106 for &e in &idx {
6107 st[e] += 1;
6108 }
6109 }
6110 if crate::gpu::enabled_here() {
6115 match crate::gpu::probe_arm(crate::gpu::OpClass::Ffn) {
6116 crate::gpu::ProbeArm::Gpu => {
6117 let t0 = std::time::Instant::now();
6118 if let Some(out) = moe_ffn_gpu(m, x, &idx, &p, wsum, pool) {
6119 crate::gpu::probe_record(crate::gpu::OpClass::Ffn, true, t0.elapsed());
6120 return out;
6121 }
6122 }
6123 crate::gpu::ProbeArm::CpuTimed => {
6124 let t0 = std::time::Instant::now();
6125 let out = crate::gpu::cpu_scope(|| moe_ffn_cpu(m, x, &idx, &p, wsum, pool));
6126 crate::gpu::probe_record(crate::gpu::OpClass::Ffn, false, t0.elapsed());
6127 return out;
6128 }
6129 crate::gpu::ProbeArm::Cpu => {
6130 return crate::gpu::cpu_scope(|| moe_ffn_cpu(m, x, &idx, &p, wsum, pool));
6131 }
6132 }
6133 }
6134 moe_ffn_cpu(m, x, &idx, &p, wsum, pool)
6135}
6136
6137fn graph_note(built: bool) {
6141 use std::sync::atomic::{AtomicBool, Ordering};
6142 static SAID: AtomicBool = AtomicBool::new(false);
6143 if !SAID.swap(true, Ordering::Relaxed) {
6144 if built {
6145 tracing::info!("wgpu whole-token graph: ACTIVE");
6146 } else {
6147 tracing::warn!("wgpu whole-token graph refused — per-op path");
6148 }
6149 }
6150}
6151
6152fn moe_batch_enabled() -> bool {
6155 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
6156 *ON.get_or_init(|| std::env::var("CMF_MOE_BATCH").as_deref() != Ok("0"))
6157}
6158
6159fn moe_ffn_cpu_batched(
6165 m: &MoeFfn,
6166 x: &[f32],
6167 idx: &[usize],
6168 p: &[f32],
6169 wsum: f32,
6170 pool: Option<&Pool>,
6171) -> Option<Vec<f32>> {
6172 if idx.is_empty() || !moe_batch_enabled() {
6173 return None;
6174 }
6175 if FFN_PROBE.with(|pr| pr.borrow().is_some()) {
6179 return None;
6180 }
6181 let n = idx.len() + usize::from(m.shared.is_some());
6182 let mut pairs = Vec::with_capacity(n);
6183 let mut downs = Vec::with_capacity(n);
6184 let mut ws = Vec::with_capacity(n);
6185 for &e in idx {
6186 let d = &m.experts[e];
6187 if d.act != Act::Silu {
6188 return None;
6189 }
6190 pairs.push((&d.gate_proj, &d.up_proj));
6191 downs.push(&d.down_proj);
6192 ws.push(p[e] / wsum * m.per_expert_scale.as_ref().map_or(1.0, |v| v[e]));
6193 }
6194 if let Some((se, gate)) = &m.shared {
6197 if se.act != Act::Silu {
6198 return None;
6199 }
6200 let g = gate.as_ref().map_or(1.0, |gate| {
6201 let mut gl = [0.0f32; 1];
6202 gate.matvec(x, &mut gl, pool);
6203 1.0 / (1.0 + (-gl[0]).exp())
6204 });
6205 pairs.push((&se.gate_proj, &se.up_proj));
6206 downs.push(&se.down_proj);
6207 ws.push(g);
6208 }
6209 let inter = pairs[0].0.rows();
6210 let mut gs: Vec<Vec<f32>> = (0..pairs.len()).map(|_| vec![0f32; inter]).collect();
6211 if !QTensor::moe_gate_up_many(&pairs, x, &mut gs, pool) {
6212 return None;
6213 }
6214 let mut out = attention::take_buf(x.len());
6215 if !QTensor::moe_down_many(&downs, &gs, &ws, &mut out, pool) {
6216 attention::recycle_buf(&mut out);
6217 return None;
6218 }
6219 Some(out)
6220}
6221
6222fn moe_ffn_cpu(
6224 m: &MoeFfn,
6225 x: &[f32],
6226 idx: &[usize],
6227 p: &[f32],
6228 wsum: f32,
6229 pool: Option<&Pool>,
6230) -> Vec<f32> {
6231 if let Some(out) = moe_ffn_cpu_batched(m, x, idx, p, wsum, pool) {
6232 return out;
6233 }
6234 let mut out = attention::take_buf(x.len());
6235 for &e in idx {
6236 let mut eo = dense_ffn(&m.experts[e], x, pool);
6237 let w = p[e] / wsum * m.per_expert_scale.as_ref().map_or(1.0, |v| v[e]);
6238 for i in 0..out.len() {
6239 out[i] += w * eo[i];
6240 }
6241 attention::recycle_buf(&mut eo);
6242 }
6243 if let Some((se, gate)) = &m.shared {
6244 let mut so = dense_ffn(se, x, pool);
6245 let g = gate.as_ref().map_or(1.0, |gate| {
6246 let mut gl = [0.0f32; 1];
6247 gate.matvec(x, &mut gl, pool);
6248 1.0 / (1.0 + (-gl[0]).exp())
6249 });
6250 for i in 0..out.len() {
6251 out[i] += g * so[i];
6252 }
6253 attention::recycle_buf(&mut so);
6254 }
6255 out
6256}
6257
6258#[allow(clippy::too_many_arguments)]
6266fn mla_attention(
6267 w: &MlaWeights,
6268 normed: &[f32],
6269 cache: &mut crate::kv_cache::LayerKvCache,
6270 position: usize,
6271 inv_freq: &[f32],
6272 rope_scale: f32,
6273 eps: f64,
6274 pool: Option<&Pool>,
6275) -> Vec<f32> {
6276 let (nh, dr, dn, dv, lora) = (w.nh, w.qk_rope, w.qk_nope, w.v_dim, w.lora);
6277 let hd = dr + dn;
6278 let mut q = vec![0.0f32; nh * hd];
6279 match (&w.q_a, &w.q_a_norm) {
6280 (Some(qa), Some(qn)) => {
6281 let mut t = vec![0.0f32; qa.rows()];
6282 qa.matvec(normed, &mut t, pool);
6283 let tn = inference::rms_norm(&t, qn, eps, NormStyle::Qwen);
6284 w.q_proj.matvec(&tn, &mut q, pool);
6285 }
6286 _ => w.q_proj.matvec(normed, &mut q, pool),
6287 }
6288 let mut ca = vec![0.0f32; lora + dr];
6289 w.kv_a.matvec(normed, &mut ca, pool);
6290 let (c_lat, k_rope) = ca.split_at_mut(lora);
6291 let latn = inference::rms_norm(c_lat, &w.kv_a_norm, eps, NormStyle::Qwen);
6292 let mut kvb = vec![0.0f32; nh * (dn + dv)];
6293 w.kv_b.matvec(&latn, &mut kvb, pool);
6294 if !w.nope {
6295 attention::rope_rotate_scaled(k_rope, position, inv_freq, rope_scale);
6296 }
6297 for h in 0..nh {
6298 if !w.nope {
6299 attention::rope_rotate_scaled(
6300 &mut q[h * hd..h * hd + dr],
6301 position,
6302 inv_freq,
6303 rope_scale,
6304 );
6305 }
6306 }
6307 let mut k = vec![0.0f32; nh * hd];
6308 let mut v = vec![0.0f32; nh * hd];
6309 for h in 0..nh {
6310 k[h * hd..h * hd + dr].copy_from_slice(k_rope);
6311 k[h * hd + dr..(h + 1) * hd].copy_from_slice(&kvb[h * (dn + dv)..h * (dn + dv) + dn]);
6312 v[h * hd..h * hd + dv].copy_from_slice(&kvb[h * (dn + dv) + dn..(h + 1) * (dn + dv)]);
6313 }
6314 cache.append(&k, &v, &vec![true; nh]);
6315 let (ao, mut imp) = attention::attend_all_heads(&q, cache, nh, 1, hd, w.scale, None, 0.0);
6316 attention::recycle_buf(&mut imp);
6317 let mut ov = vec![0.0f32; nh * dv];
6318 for h in 0..nh {
6319 ov[h * dv..(h + 1) * dv].copy_from_slice(&ao[h * hd..h * hd + dv]);
6320 }
6321 let mut out = vec![0.0f32; w.o_proj.rows()];
6322 w.o_proj.matvec(&ov, &mut out, pool);
6323 out
6324}
6325
6326fn dense_moe_ffn(
6333 dm: &DenseMoeFfn,
6334 x_normed: &[f32],
6335 h_raw: &[f32],
6336 eps: f64,
6337 norm_style: NormStyle,
6338 pool: Option<&Pool>,
6339) -> Vec<f32> {
6340 let mut d = dense_ffn(&dm.dense, x_normed, pool);
6341 d = inference::rms_norm(&d, &dm.post_norm_1, eps, norm_style);
6342 let m = &dm.moe;
6343 let ne = m.experts.len();
6344 let mut logits = vec![0.0f32; ne];
6345 if m.router_input_norm {
6346 let ss: f32 = h_raw.iter().map(|v| v * v).sum::<f32>() / h_raw.len() as f32;
6347 let inv = 1.0 / (ss + eps as f32).sqrt();
6348 let xr: Vec<f32> = h_raw.iter().map(|v| v * inv).collect();
6349 m.router.matvec(&xr, &mut logits, pool);
6350 } else {
6351 m.router.matvec(h_raw, &mut logits, pool);
6352 }
6353 let (idx, p, wsum) = moe_route(&logits, m, None);
6354 {
6355 let mut st = m.stats.borrow_mut();
6356 if st.len() < ne {
6357 st.resize(ne, 0);
6358 }
6359 for &e in &idx {
6360 st[e] += 1;
6361 }
6362 }
6363 let x2 = inference::rms_norm(h_raw, &dm.pre_norm_2, eps, norm_style);
6364 let mo = moe_ffn_cpu(m, &x2, &idx, &p, wsum, pool);
6365 let mo = inference::rms_norm(&mo, &dm.post_norm_2, eps, norm_style);
6366 for (di, mi) in d.iter_mut().zip(&mo) {
6367 *di += mi;
6368 }
6369 d
6370}
6371
6372fn moe_gpu_refused(why: &'static str) {
6379 use std::sync::atomic::{AtomicBool, Ordering};
6380 static SAID: AtomicBool = AtomicBool::new(false);
6381 if !SAID.swap(true, Ordering::Relaxed) {
6382 tracing::warn!("MoE GPU block refused ({why}) — experts run on the CPU");
6383 }
6384}
6385
6386fn moe_ffn_gpu(
6387 m: &MoeFfn,
6388 x: &[f32],
6389 idx: &[usize],
6390 p: &[f32],
6391 wsum: f32,
6392 pool: Option<&Pool>,
6393) -> Option<Vec<f32>> {
6394 use crate::gpu::MoeJob;
6395
6396 let mut jobs: Vec<MoeJob> = Vec::with_capacity(idx.len() + 1);
6397 let mut model_ref = None;
6398 for &e in idx {
6399 if moe_push_job(&m.experts[e], x, p[e] / wsum, &mut jobs, &mut model_ref).is_none() {
6400 moe_gpu_refused("push_job(expert)");
6401 return None;
6402 }
6403 }
6404 if let Some((se, gate)) = &m.shared {
6405 let g = gate.as_ref().map_or(1.0, |gate| {
6406 let mut gl = [0.0f32; 1];
6407 gate.matvec(x, &mut gl, pool);
6408 1.0 / (1.0 + (-gl[0]).exp())
6409 });
6410 if moe_push_job(se, x, g, &mut jobs, &mut model_ref).is_none() {
6411 moe_gpu_refused("push_job(shared)");
6412 return None;
6413 }
6414 }
6415 let Some(model) = model_ref else {
6416 moe_gpu_refused("no model_ref");
6417 return None;
6418 };
6419 let hidden = jobs[0].down.1;
6420 let mut out = vec![0.0f32; hidden];
6421 if crate::gpu::moe_block(&model, &jobs, &mut out) {
6422 Some(out)
6423 } else {
6424 moe_gpu_refused("gpu::moe_block");
6425 None
6426 }
6427}
6428
6429fn ffn_forward(
6431 ffn: &FfnKind,
6432 x: &[f32],
6433 pool: Option<&Pool>,
6434 experts_allowed: Option<&[bool]>,
6435) -> Vec<f32> {
6436 match ffn {
6437 FfnKind::Dense(d) => dense_ffn(d, x, pool),
6438 FfnKind::Moe(m) => moe_ffn(m, x, pool, experts_allowed),
6439 FfnKind::DenseMoe(_) => unreachable!("DenseMoe dispatches via dense_moe_ffn"),
6443 }
6444}
6445
6446fn ffn_forward_pair(
6450 ffn: &FfnKind,
6451 x1: &[f32],
6452 x2: &[f32],
6453 pool: Option<&Pool>,
6454 experts_allowed: Option<&[bool]>,
6455) -> (Vec<f32>, Vec<f32>) {
6456 let d = match ffn {
6457 FfnKind::Dense(d) => d,
6458 FfnKind::Moe(m) => {
6459 return (
6460 moe_ffn(m, x1, pool, experts_allowed),
6461 moe_ffn(m, x2, pool, experts_allowed),
6462 );
6463 }
6464 FfnKind::DenseMoe(_) => unreachable!("DenseMoe dispatches via dense_moe_ffn"),
6465 };
6466 let inter = d.gate_proj.rows();
6467 FFN_SCRATCH.with(|s| {
6468 let mut s = s.borrow_mut();
6469 let [g1, g2, u1, u2] = &mut *s;
6470 g1.resize(inter, 0.0);
6471 g2.resize(inter, 0.0);
6472 u1.resize(inter, 0.0);
6473 u2.resize(inter, 0.0);
6474 QTensor::matvec2_many(
6477 [&d.gate_proj, &d.up_proj],
6478 x1,
6479 x2,
6480 [g1.as_mut_slice(), u1.as_mut_slice()],
6481 [g2.as_mut_slice(), u2.as_mut_slice()],
6482 pool,
6483 );
6484 for i in 0..inter {
6485 g1[i] = d.act.combine(g1[i], u1[i]);
6486 g2[i] = d.act.combine(g2[i], u2[i]);
6487 }
6488 let mut o1 = attention::take_buf(d.down_proj.rows());
6489 let mut o2 = attention::take_buf(d.down_proj.rows());
6490 d.down_proj.matvec2(g1, g2, &mut o1, &mut o2, pool);
6491 (o1, o2)
6492 })
6493}
6494
6495#[cfg(test)]
6496mod tests {
6497
6498 #[test]
6499 fn cancel_flag_stops_generation() {
6500 let mut p = create_test_pipeline(16, 32, 2, 2, 8, 2, 32);
6501 p.cancel.store(true, std::sync::atomic::Ordering::Relaxed);
6504 let r = p.generate_from_ids(&[1, 2, 3], 8, None, None).unwrap();
6505 assert_eq!(r.finish_reason, "cancelled");
6506 assert!(
6507 r.token_ids.is_empty(),
6508 "no tokens after cancel: {:?}",
6509 r.token_ids
6510 );
6511 let r2 = p.generate_from_ids(&[1, 2, 3], 4, None, None).unwrap();
6513 assert_ne!(r2.finish_reason, "cancelled");
6514 }
6515 use super::*;
6516
6517 #[test]
6523 fn sparse_ffn_quant_equals_dense_with_inactive_zeroed() {
6524 let (hidden, inter) = (16usize, 40usize);
6525 let synth = |n: usize, salt: usize| -> Vec<f32> {
6526 (0..n)
6527 .map(|i| (((i * 37 + salt * 11 + 3) % 101) as f32 / 101.0 - 0.5) * 0.4)
6528 .collect()
6529 };
6530 let d = DenseFfn {
6531 gate_proj: QTensor::from_f32(synth(inter * hidden, 1), inter, hidden),
6532 up_proj: QTensor::from_f32(synth(inter * hidden, 2), inter, hidden),
6533 down_proj: QTensor::from_f32(synth(hidden * inter, 3), hidden, inter),
6534 act: Act::Silu,
6535 };
6536 let x = synth(hidden, 9);
6537 let active: Vec<u16> = (0..inter as u16).filter(|i| i % 3 == 0).collect();
6539
6540 let sparse = sparse_ffn_quant(&d, &x, &active, hidden, None);
6541
6542 let mut g = vec![0.0f32; inter];
6544 d.gate_proj.matvec(&x, &mut g, None);
6545 let mut u = vec![0.0f32; inter];
6546 d.up_proj.matvec(&x, &mut u, None);
6547 let act_set: std::collections::HashSet<u16> = active.iter().copied().collect();
6548 for i in 0..inter {
6549 g[i] = if act_set.contains(&(i as u16)) {
6550 inference::silu(g[i]) * u[i]
6551 } else {
6552 0.0
6553 };
6554 }
6555 let mut reference = vec![0.0f32; hidden];
6556 d.down_proj.matvec(&g, &mut reference, None);
6557
6558 let max_d = sparse
6559 .iter()
6560 .zip(&reference)
6561 .map(|(a, b)| (a - b).abs())
6562 .fold(0.0f32, f32::max);
6563 assert!(max_d < 1e-5, "sparse != dense-zeroed: max|Δ| = {max_d}");
6564 }
6565
6566 fn attach_test_mtp(p: &mut Pipeline) {
6568 let (h, inter, heads, kv, hd) = (
6569 p.hidden_size,
6570 p.intermediate_size,
6571 p.num_heads,
6572 p.num_kv_heads,
6573 p.head_dim,
6574 );
6575 let synth = |n: usize, salt: usize| -> Vec<f32> {
6576 (0..n)
6577 .map(|i| (((i * 29 + salt * 23 + 5) % 101) as f32 / 101.0 - 0.5) * 0.2)
6578 .collect()
6579 };
6580 let qt = |rows: usize, cols: usize, salt: usize| -> QTensor {
6581 QTensor::from_f32(synth(rows * cols, salt), rows, cols)
6582 };
6583 p.mtp = Some(MtpModule {
6584 enorm: vec![1.0; h],
6585 hnorm: vec![1.0; h],
6586 eh_proj: qt(h, 2 * h, 301),
6587 layer: LayerWeights {
6588 input_norm: vec![1.0; h],
6589 post_norm: vec![1.0; h],
6590 attn_out_norm: None,
6591 ffn_out_norm: None,
6592 layer_scale: None,
6593 ffn: FfnKind::Dense(DenseFfn {
6594 gate_proj: qt(inter, h, 315),
6595 up_proj: qt(inter, h, 316),
6596 down_proj: qt(h, inter, 317),
6597 act: Act::Silu,
6598 }),
6599 attn: AttnKind::Full {
6600 bias: None,
6601 wq: qt(heads * hd, h, 311),
6602 wk: qt(kv * hd, h, 312),
6603 wv: qt(kv * hd, h, 313),
6604 wo: qt(h, heads * hd, 314),
6605 q_norm: None,
6606 k_norm: None,
6607 output_gate: false,
6608 softplus_gate: None,
6609 },
6610 },
6611 final_norm: vec![1.0; h],
6612 kv: crate::kv_cache::LayerKvCache::new(kv, hd),
6613 });
6614 }
6615
6616 #[test]
6617 fn speculative_equals_vanilla_greedy() {
6618 unsafe { std::env::set_var("CMF_GPU_WGPU_GRAPH", "0") };
6622 let run = |spec: bool| {
6623 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
6624 p.sampler_config.temperature = 0.0;
6625 attach_test_mtp(&mut p);
6626 p.speculative = spec;
6627 let r = p.generate("abcdef", 12, None, None).unwrap();
6628 (r.token_ids, r.mtp_drafted, r.mtp_accepted)
6629 };
6630 let (vanilla, d0, _) = run(false);
6631 let (spec, d1, a1) = run(true);
6632 assert_eq!(d0, 0, "vanilla path must not draft");
6633 assert!(d1 > 0, "speculative path must draft");
6634 assert_eq!(
6635 vanilla, spec,
6636 "speculative must reproduce the exact greedy sequence (accepted {a1}/{d1})"
6637 );
6638 }
6639
6640 #[test]
6641 fn speculative_accepts_constant_oracle() {
6642 unsafe { std::env::set_var("CMF_GPU_WGPU_GRAPH", "0") };
6644 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
6645 p.sampler_config.temperature = 0.0;
6646 p.sampler_config.repetition_penalty = 1.0;
6647 p.weights.lm_head = QTensor::from_f32(vec![0.01; 64 * 8], 64, 8);
6650 attach_test_mtp(&mut p);
6651 p.speculative = true;
6652 let r = p.generate("abcd", 10, None, None).unwrap();
6653 assert!(r.mtp_drafted > 0);
6654 assert_eq!(
6655 r.mtp_accepted, r.mtp_drafted,
6656 "constant logits → every draft accepted"
6657 );
6658 assert!(r.token_ids.windows(2).all(|w| w[0] == w[1]));
6661 }
6662
6663 #[test]
6664 fn empty_prompt_is_an_error_not_a_panic() {
6665 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 260);
6666 let r = p.generate("", 4, None, None);
6667 assert!(r.is_err(), "empty prompt must be a clean error");
6668 }
6669
6670 #[test]
6671 fn every_token_enters_kv_exactly_once() {
6672 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
6673 p.sampler_config.temperature = 0.0;
6675 let r = p.generate("abc", 2, None, None).unwrap();
6676 assert_eq!(r.prompt_tokens, 3);
6677 assert_eq!(
6681 p.kv_cache.seq_len(),
6682 3 + r.tokens_generated - 1,
6683 "each token must be cached exactly once (v1 cached the last prompt token twice)"
6684 );
6685 }
6686
6687 #[test]
6688 fn generation_is_reproducible_with_seed() {
6689 let run = || {
6690 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
6691 p.generate("hello", 8, None, None).unwrap().token_ids
6692 };
6693 assert_eq!(run(), run());
6694 }
6695
6696 #[test]
6697 fn resetting_sampler_restarts_the_seeded_stream() {
6698 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
6699 let config = SamplerConfig {
6700 seed: Some(1234),
6701 ..SamplerConfig::default()
6702 };
6703 p.set_sampler_config(config.clone());
6704 let first = p.generate("hello", 8, None, None).unwrap().token_ids;
6705 p.set_sampler_config(config);
6706 let second = p.generate("hello", 8, None, None).unwrap().token_ids;
6707 assert_eq!(first, second);
6708 }
6709
6710 #[test]
6711 fn eviction_bounds_the_cache() {
6712 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 260);
6713 p.kv_cache.max_seq_len = 6;
6714 p.sampler_config.temperature = 0.0;
6715 let _ = p.generate("abcd", 12, None, None).unwrap();
6716 assert!(
6717 p.kv_cache.seq_len() <= 6 + 1,
6718 "cache must stay bounded by max_seq_len (got {})",
6719 p.kv_cache.seq_len()
6720 );
6721 }
6722
6723 #[test]
6724 fn confidence_matches_tokens_and_is_a_probability() {
6725 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
6726 p.sampler_config.temperature = 0.0;
6727 p.sampler_config.repetition_penalty = 1.0;
6728 let r = p.generate("abcd", 10, None, None).unwrap();
6729 assert_eq!(
6730 r.token_confidence.len(),
6731 r.token_ids.len(),
6732 "one confidence per emitted token"
6733 );
6734 for &c in &r.token_confidence {
6735 assert!((0.0..=1.0).contains(&c), "confidence out of [0,1]: {c}");
6736 }
6737 let logits = [1.0f32, 3.0, 0.5, 3.0];
6739 let p0 = top1_prob_t(&logits, 1, 1.0);
6740 let p1 = top1_prob_t(&logits, 3, 1.0);
6741 assert!((p0 - p1).abs() < 1e-6, "equal logits → equal prob");
6742 assert!(p0 > 0.0 && p0 < 1.0);
6743 let sharp = top1_prob_t(&logits, 1, 1.0);
6745 let soft = top1_prob_t(&logits, 1, 2.0);
6746 assert!(soft < sharp, "higher temperature lowers peak confidence");
6747 }
6748
6749 #[test]
6750 fn trace_is_opt_in_and_parallels_the_output() {
6751 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
6753 p.sampler_config.temperature = 0.0;
6754 p.sampler_config.repetition_penalty = 1.0;
6755 let r = p.generate("abcd", 10, None, None).unwrap();
6756 assert!(r.traces.is_empty(), "trace must be empty unless enabled");
6757
6758 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
6760 p.sampler_config.temperature = 0.0;
6761 p.sampler_config.repetition_penalty = 1.0;
6762 p.set_trace(true);
6763 let r = p.generate("abcd", 10, None, None).unwrap();
6764 assert_eq!(r.traces.len(), r.token_ids.len(), "one trace row per token");
6765 for (i, tr) in r.traces.iter().enumerate() {
6766 assert_eq!(tr.t, i, "trace index is sequential");
6767 assert_eq!(tr.token_id, r.token_ids[i], "trace token_id matches output");
6768 assert_eq!(
6769 tr.confidence, r.token_confidence[i],
6770 "trace confidence matches the confidence channel"
6771 );
6772 assert!(tr.active_skill.is_none() && tr.recon.is_none() && !tr.switched);
6774 }
6775 }
6776
6777 #[test]
6778 fn explain_prefill_logits_match_greedy_first_token() {
6779 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
6783 p.sampler_config.temperature = 0.0;
6784 p.sampler_config.repetition_penalty = 1.0;
6785 let ids = p.tokenizer.encode("abcd");
6786 let logits = p.prefill_next_logits(&ids, None);
6787 let argmax = logits
6788 .iter()
6789 .enumerate()
6790 .max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
6791 .unwrap()
6792 .0 as u32;
6793 let r = p.generate("abcd", 1, None, None).unwrap();
6794 assert_eq!(
6795 argmax, r.token_ids[0],
6796 "explain preview must match greedy emit"
6797 );
6798 }
6799
6800 #[test]
6801 fn laguna_shared_expert_is_unconditionally_added() {
6802 let matrix = |values: Vec<f32>| QTensor::from_f32(values, 2, 2);
6803 let identity = || matrix(vec![1.0, 0.0, 0.0, 1.0]);
6804 let zero_dense = || DenseFfn {
6805 gate_proj: matrix(vec![0.0; 4]),
6806 up_proj: matrix(vec![0.0; 4]),
6807 down_proj: matrix(vec![0.0; 4]),
6808 act: Act::Silu,
6809 };
6810 let shared = DenseFfn {
6811 gate_proj: identity(),
6812 up_proj: identity(),
6813 down_proj: identity(),
6814 act: Act::Silu,
6815 };
6816 let x = [1.0, 2.0];
6817 let expected = dense_ffn(&shared, &x, None);
6818 let moe = MoeFfn {
6819 router: QTensor::from_f32(vec![0.0, 0.0], 1, 2),
6820 experts: vec![zero_dense()],
6821 top_k: 1,
6822 norm_topk_prob: true,
6823 router_sigmoid: true,
6824 expert_bias: None,
6825 routed_scaling: 1.0,
6826 route_tau: None,
6827 shared: Some((shared, None)),
6828 stats: std::cell::RefCell::new(Vec::new()),
6829 act_sq: std::cell::RefCell::new(Vec::new()),
6830 act_rows: std::cell::RefCell::new(Vec::new()),
6831 mask: None,
6832 per_expert_scale: None,
6833 router_input_norm: false,
6834 };
6835 let actual = moe_ffn_cpu(&moe, &x, &[0], &[0.0], 1.0, None);
6836 for (actual, expected) in actual.iter().zip(expected) {
6837 assert!((actual - expected).abs() < 1e-6);
6838 }
6839 }
6840}