1use crate::attention::{self, QwenAttnCfg};
10use crate::inference;
11use crate::kv_cache::KvCache;
12use crate::linear_core::{
13 gdn_forward, gdn_pair, vmf_phase_forward, vmf_phase_pair, GdnCfg, GdnWeights, VmfPhaseCfg,
14 VmfPhaseWeights,
15};
16use crate::pool::Pool;
17use crate::qtensor::QTensor;
18use crate::sampler::{self, SamplerConfig, SplitMix64};
19use crate::tokenizer::Tokenizer;
20use cortiq_core::mask::TaskMask;
21use cortiq_core::types::NormStyle;
22
23struct ForwardScratch {
27 n1: Vec<f32>,
28 n2: Vec<f32>,
29 p1: Vec<f32>,
30 p2: Vec<f32>,
31}
32
33impl ForwardScratch {
34 fn new(hidden: usize) -> Self {
35 Self {
36 n1: vec![0.0; hidden],
37 n2: vec![0.0; hidden],
38 p1: vec![0.0; hidden],
39 p2: vec![0.0; hidden],
40 }
41 }
42}
43
44pub struct Pipeline {
46 pub tokenizer: std::sync::Arc<Tokenizer>,
49 pub kv_cache: KvCache,
50 pub sampler_config: SamplerConfig,
51 pub weights: PipelineWeights,
52 pub hidden_size: usize,
53 pub intermediate_size: usize,
54 pub num_heads: usize,
55 pub num_kv_heads: usize,
56 pub head_dim: usize,
57 pub num_layers: usize,
58 pub vocab_size: usize,
59 pub rms_eps: f64,
60 pub rope_base: f32,
61 pub norm_style: NormStyle,
62 pub rotary_dim: usize,
64 pub vmf_cfg: Option<VmfPhaseCfg>,
66 pub gdn_cfg: Option<GdnCfg>,
68 pub mtp: Option<MtpModule>,
70 pub speculative: bool,
72 rng: SplitMix64,
73 inv_freq: std::sync::Arc<Vec<f32>>,
77 ws: ForwardScratch,
81 pool: Option<std::sync::Arc<Pool>>,
83 pub(crate) model: Option<std::sync::Arc<cortiq_core::CmfModel>>,
87 pub(crate) dyn_force_f32: bool,
89 pub(crate) dyn_skill_layers: Vec<Option<Vec<usize>>>,
94 pub(crate) dyn_active: Option<usize>,
100 pub(crate) dyn_blend_loaded: bool,
104 pub(crate) dyn_phi_layer: Option<usize>,
107 dyn_phi_ema: Vec<f32>,
109 dyn_phi_seen: usize,
110 pub dyn_router: Option<crate::swarm::DynRouter>,
113 o1_cfg: Option<crate::nystrom::O1Cfg>,
116 o1_flags: Vec<bool>,
118 trace: bool,
121 calib_temp: f32,
124 #[cfg_attr(not(target_os = "macos"), allow(dead_code))]
126 graph_kv_id: u64,
127 pub embed_multiplier: f32,
129 pub attn_scale: f32,
132 pub swa: Option<(usize, usize)>,
135 pub inv_freq_local: Option<std::sync::Arc<Vec<f32>>>,
138 pub global_attn: Option<(usize, usize)>,
141 pub inv_freq_global: Option<std::sync::Arc<Vec<f32>>>,
144 pub attn_v_norm: bool,
146 pub final_softcap: Option<f32>,
148 confidence_on: bool,
152}
153
154#[cfg(target_os = "macos")]
155impl Drop for Pipeline {
156 fn drop(&mut self) {
157 crate::gpu::kv_mirror_drop(self.graph_kv_id);
158 }
159}
160
161pub struct PipelineWeights {
166 pub embed_tokens: QTensor,
168 pub layers: Vec<LayerWeights>,
170 pub lm_head: QTensor,
172 pub final_norm: Vec<f32>,
174}
175
176pub struct LayerWeights {
178 pub input_norm: Vec<f32>,
179 pub post_norm: Vec<f32>,
182 pub attn_out_norm: Option<Vec<f32>>,
185 pub layer_scale: Option<f32>,
187 pub ffn_out_norm: Option<Vec<f32>>,
190 pub ffn: FfnKind,
191 pub attn: AttnKind,
192}
193
194#[derive(Clone, Copy, PartialEq, Eq, Debug, Default)]
197pub enum Act {
198 #[default]
199 Silu,
200 GeluTanh,
201}
202
203impl Act {
204 pub fn from_arch(name: &str) -> Self {
205 if name == "gelu_tanh" {
206 Self::GeluTanh
207 } else {
208 Self::Silu
209 }
210 }
211
212 #[inline]
213 pub fn apply(self, x: f32) -> f32 {
214 match self {
215 Self::Silu => inference::silu(x),
216 Self::GeluTanh => inference::gelu_tanh(x),
217 }
218 }
219}
220
221pub struct DenseFfn {
223 pub gate_proj: QTensor,
224 pub up_proj: QTensor,
225 pub down_proj: QTensor,
226 pub act: Act,
228}
229
230pub enum FfnKind {
233 Dense(DenseFfn),
234 Moe(MoeFfn),
238}
239
240pub struct MoeFfn {
241 pub router: QTensor,
243 pub experts: Vec<DenseFfn>,
244 pub top_k: usize,
245 pub norm_topk_prob: bool,
246 pub shared: Option<(DenseFfn, QTensor)>,
248 pub stats: std::cell::RefCell<Vec<u64>>,
252}
253
254pub enum AttnKind {
257 Full {
259 wq: QTensor,
260 wk: QTensor,
261 wv: QTensor,
262 wo: QTensor,
263 q_norm: Option<Vec<f32>>,
264 k_norm: Option<Vec<f32>>,
265 output_gate: bool,
266 bias: Option<(Vec<f32>, Vec<f32>, Vec<f32>)>,
268 },
269 Linear(VmfPhaseWeights),
271 LinearGdn(GdnWeights),
273}
274
275pub struct MtpModule {
280 pub enorm: Vec<f32>,
281 pub hnorm: Vec<f32>,
282 pub eh_proj: QTensor,
284 pub layer: LayerWeights,
285 pub final_norm: Vec<f32>,
286 pub kv: crate::kv_cache::LayerKvCache,
287}
288
289pub struct GenerateResult {
291 pub text: String,
292 pub token_ids: Vec<u32>,
293 pub prompt_tokens: usize,
294 pub tokens_generated: usize,
295 pub finish_reason: String,
296 pub mtp_drafted: usize,
298 pub mtp_accepted: usize,
299 pub token_confidence: Vec<f32>,
304 pub traces: Vec<TokenTrace>,
307}
308
309#[derive(Clone, Debug)]
314pub struct TokenTrace {
315 pub t: usize,
317 pub token_id: u32,
319 pub confidence: f32,
321 pub active_skill: Option<String>,
323 pub recon: Option<f32>,
327 pub switched: bool,
330}
331
332fn top1_prob_t(logits: &[f32], id: u32, temp: f32) -> f32 {
337 let t = if temp > 1e-3 { temp } else { 1.0 };
338 let max = logits.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
339 let sum: f32 = logits.iter().map(|&v| ((v - max) / t).exp()).sum();
340 if sum > 0.0 {
341 (((logits[id as usize] - max) / t).exp()) / sum
342 } else {
343 0.0
344 }
345}
346
347fn prefill_batched() -> bool {
350 std::env::var("CMF_PREFILL").map(|v| v != "seq").unwrap_or(true)
351}
352
353fn prefill_chunk() -> usize {
358 if let Some(n) =
359 std::env::var("CMF_PREFILL_CHUNK").ok().and_then(|v| v.parse::<usize>().ok())
360 {
361 return n.max(1);
362 }
363 if cfg!(target_os = "macos") {
364 512
365 } else {
366 48
367 }
368}
369
370pub type TokenCallback = Box<dyn FnMut(&str) -> bool + Send>;
372
373impl Pipeline {
374 #[allow(clippy::too_many_arguments)]
376
377 #[cfg(target_os = "macos")]
392 fn graph_prefill_preferred(&self) -> bool {
393 if !crate::gpu::enabled_here()
394 || !crate::gpu::q1_force()
395 || std::env::var("CMF_GPU_BLOCK").map(|v| v == "0").unwrap_or(false)
396 {
397 return false;
398 }
399 self.weights.layers.iter().any(
400 |lw| matches!(&lw.attn, AttnKind::LinearGdn(w) if w.in_proj_qkv.is_q1()),
401 )
402 }
403
404 #[cfg(not(target_os = "macos"))]
405 fn graph_prefill_preferred(&self) -> bool {
406 false
407 }
408
409 #[cfg(target_os = "macos")]
410 fn q1_graph_gpu(
411 &mut self,
412 start: usize,
413 upto: Option<usize>,
414 position: usize,
415 h: &mut [f32],
416 ) -> usize {
417 use crate::gpu::{AttnGpuLayer, GdnGpuCfg, GdnGpuLayer, GraphDims, TokenGraph};
418 if !crate::gpu::enabled_here()
419 || !crate::gpu::q1_force()
420 || std::env::var("CMF_GPU_BLOCK").map(|v| v == "0").unwrap_or(false)
421 {
422 return start;
423 }
424 if self.swa.is_some()
429 || self.global_attn.is_some()
430 || self.attn_v_norm
431 || (self.attn_scale - 1.0 / (self.head_dim as f32).sqrt()).abs() > 1e-9
432 || self.weights.layers.iter().any(|lw| {
433 lw.attn_out_norm.is_some()
434 || lw.ffn_out_norm.is_some()
435 || lw.layer_scale.is_some()
436 || matches!(&lw.ffn, FfnKind::Dense(d) if d.act != Act::Silu)
437 })
438 {
439 return start;
440 }
441 let limit = upto.map(|u| u + 1).unwrap_or(self.num_layers).min(self.num_layers);
442
443 enum Item<'a> {
444 Gdn {
445 run: Vec<GdnGpuLayer<'a>>,
446 first: usize,
447 },
448 Attn {
449 l: AttnGpuLayer<'a>,
450 li: usize,
451 q_norm: Option<&'a [f32]>,
452 k_norm: Option<&'a [f32]>,
453 output_gate: bool,
454 bias: Option<(&'a [f32], &'a [f32], &'a [f32])>,
455 full_gpu: bool,
458 },
459 }
460
461 let dev_attend = std::env::var("CMF_GPU_ATTEND").map(|v| v != "0").unwrap_or(true)
463 && self.head_dim % 4 == 0
464 && self.head_dim <= 128
465 && self.rotary_dim >= 2
466 && self.rotary_dim <= self.head_dim
467 && (self.rotary_dim / 2) % 32 == 0
468 && self.num_kv_heads > 0
469 && self.num_heads % self.num_kv_heads == 0;
470
471 let mut plan: Vec<Item> = Vec::new();
472 let mut model_ref: Option<std::sync::Arc<cortiq_core::CmfModel>> = None;
473 let mut scan = start;
474 while scan < limit {
475 let lw = &self.weights.layers[scan];
476 let FfnKind::Dense(d) = &lw.ffn else { break };
477 let (Some(g), Some(u), Some(dn)) =
478 (d.gate_proj.q1_parts(), d.up_proj.q1_parts(), d.down_proj.q1_parts())
479 else {
480 break;
481 };
482 match &lw.attn {
483 AttnKind::LinearGdn(w) if self.gdn_cfg.is_some() => {
484 let parts = (
485 w.in_proj_qkv.q1_parts(),
486 w.in_proj_z.q1_parts(),
487 w.in_proj_a.f32_parts(),
488 w.in_proj_b.f32_parts(),
489 w.out_proj.q1_parts(),
490 );
491 let (Some(qkv), Some(z), Some(a), Some(b), Some(out)) = parts else { break };
492 if let QTensor::Mapped { model, .. } = &w.in_proj_qkv {
493 model_ref.get_or_insert_with(|| model.clone());
494 }
495 let gl = GdnGpuLayer {
496 attn_norm: &lw.input_norm,
497 post_norm: &lw.post_norm,
498 qkv,
499 z,
500 a,
501 b,
502 out,
503 gate: g,
504 up: u,
505 down: dn,
506 conv1d: &w.conv1d,
507 a_log: &w.a_log,
508 dt_bias: &w.dt_bias,
509 gnorm: &w.norm,
510 };
511 match plan.last_mut() {
512 Some(Item::Gdn { run, .. }) => run.push(gl),
513 _ => plan.push(Item::Gdn { run: vec![gl], first: scan }),
514 }
515 }
516 AttnKind::Full { wq, wk, wv, wo, q_norm, k_norm, output_gate, bias }
517 if !self.kv_cache.layers[scan].o1_sealed() =>
518 {
519 let parts = (wq.q1_parts(), wk.q1_parts(), wv.q1_parts(), wo.q1_parts());
520 let (Some(pq), Some(pk), Some(pv), Some(po)) = parts else { break };
521 if let QTensor::Mapped { model, .. } = wq {
522 model_ref.get_or_insert_with(|| model.clone());
523 }
524 let cache = &self.kv_cache.layers[scan];
525 let full_gpu = dev_attend
526 && cache.mode == crate::kv_cache::KvMode::F32
527 && cache.o1.is_none()
528 && bias.is_none()
529 && pq.1 == self.num_heads * self.head_dim * (1 + *output_gate as usize)
530 && pk.1 == self.num_kv_heads * self.head_dim
531 && pv.1 == self.num_kv_heads * self.head_dim
532 && po.2 == self.num_heads * self.head_dim;
533 plan.push(Item::Attn {
534 l: AttnGpuLayer {
535 attn_norm: &lw.input_norm,
536 post_norm: &lw.post_norm,
537 wq: pq,
538 wk: pk,
539 wv: pv,
540 wo: po,
541 gate: g,
542 up: u,
543 down: dn,
544 },
545 li: scan,
546 q_norm: q_norm.as_deref(),
547 k_norm: k_norm.as_deref(),
548 output_gate: *output_gate,
549 bias: bias
550 .as_ref()
551 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
552 full_gpu,
553 });
554 }
555 _ => break,
556 }
557 scan += 1;
558 }
559 let Some(model) = model_ref else { return start };
560 if plan.is_empty() {
561 return start;
562 }
563 let dims = GraphDims {
564 hidden: self.hidden_size,
565 eps: self.rms_eps as f32,
566 gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
567 };
568 let Some(mut graph) = TokenGraph::new(&model, dims, h) else { return start };
569 let gcfg = self.gdn_cfg.map(|cfg| GdnGpuCfg {
570 nv: cfg.num_v_heads,
571 nk: cfg.num_k_heads,
572 dk: cfg.key_head_dim,
573 dv: cfg.value_head_dim,
574 kk: cfg.conv_kernel,
575 hidden: self.hidden_size,
576 inter: self.intermediate_size,
577 c_dim: cfg.conv_dim(),
578 eps: cfg.rms_eps as f32,
579 gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
580 });
581 let mut valid = 0usize;
585 let mut end = start;
586 for item in &plan {
587 let ok = match item {
588 Item::Gdn { run, .. } => gcfg
589 .as_ref()
590 .map(|gc| run.iter().all(|l| graph.gdn_ok(l, gc)))
591 .unwrap_or(false),
592 Item::Attn { l, .. } => graph.attn_ok(l),
593 };
594 if !ok {
595 break;
596 }
597 valid += 1;
598 end += match item {
599 Item::Gdn { run, .. } => run.len(),
600 Item::Attn { .. } => 1,
601 };
602 }
603 plan.truncate(valid);
604 if plan.is_empty() {
605 return start;
606 }
607
608 let inv_freq = self.inv_freq.clone();
609 let pool = self.pool.clone();
610 let (nh, nkv, hd, hs, rd, eps) = (
611 self.num_heads,
612 self.num_kv_heads,
613 self.head_dim,
614 self.hidden_size,
615 self.rotary_dim,
616 self.rms_eps,
617 );
618 let norm_style = self.norm_style;
619 let gemma = norm_style == cortiq_core::NormStyle::Gemma;
620 let want = self.gdn_cfg.map(|c| c.state_len()).unwrap_or(0);
621 let kv_id = self.graph_kv_id;
622 let mut pending: Vec<(usize, usize)> = Vec::new();
625 let mut dev_attn: Vec<usize> = Vec::new();
628 for item in &plan {
629 match item {
630 Item::Gdn { run, first } => {
631 for l in &mut self.kv_cache.layers[*first..*first + run.len()] {
632 if l.linear_state.len() != want {
633 l.linear_state = vec![0f32; want];
634 }
635 }
636 let ro: Vec<&[f32]> = self.kv_cache.layers[*first..*first + run.len()]
637 .iter()
638 .map(|l| l.linear_state.as_slice())
639 .collect();
640 if !graph.encode_gdn_run(run, &ro, gcfg.as_ref().unwrap()) {
641 tracing::error!("q1 graph: GDN run refused after validation");
643 return start;
644 }
645 graph.commit();
648 pending.push((*first, run.len()));
649 }
650 Item::Attn { l, li, q_norm, k_norm, output_gate, bias, full_gpu } => {
651 if *full_gpu {
653 let cache = &self.kv_cache.layers[*li];
654 let cpu_k: Vec<&[f32]> = (0..nkv).map(|g| cache.head_keys(g)).collect();
655 let cpu_v: Vec<&[f32]> = (0..nkv).map(|g| cache.head_values(g)).collect();
656 let cpu_stored = cpu_k[0].len() / hd;
657 let p = crate::gpu::AttnDeviceParams {
658 kv_id,
659 layer: *li,
660 nh,
661 nkv,
662 hd,
663 rd,
664 position,
665 eps: eps as f32,
666 gemma,
667 output_gate: *output_gate,
668 q_norm: *q_norm,
669 k_norm: *k_norm,
670 inv_freq: &inv_freq,
671 cpu_k,
672 cpu_v,
673 cpu_stored,
674 };
675 if graph.attn_device_ok(l, &p) && graph.encode_attn_device(l, &p) {
676 graph.commit();
677 dev_attn.push(*li);
678 continue;
679 }
680 }
682 graph.encode_attn_prefix(l);
683 graph.sync();
684 if !pending.is_empty() {
685 let idxs: Vec<usize> =
686 pending.drain(..).flat_map(|(f, n)| f..f + n).collect();
687 let mut outs: Vec<&mut [f32]> = self
688 .kv_cache
689 .layers
690 .iter_mut()
691 .enumerate()
692 .filter(|(i, _)| idxs.binary_search(i).is_ok())
693 .map(|(_, s)| s.linear_state.as_mut_slice())
694 .collect();
695 graph.read_states(&mut outs);
696 }
697 let mut q_raw = attention::take_buf(l.wq.1);
698 let mut k = attention::take_buf(l.wk.1);
699 let mut v = attention::take_buf(l.wv.1);
700 graph.read_qkv(&mut q_raw, &mut k, &mut v);
701 let cfg = QwenAttnCfg {
702 num_heads: nh,
703 num_kv_heads: nkv,
704 head_dim: hd,
705 hidden_size: hs,
706 position,
707 inv_freq: &inv_freq,
708 rotary_dim: rd,
709 scale: self.attn_scale,
710 window: None,
711 v_norm: false,
712 q_norm: *q_norm,
713 k_norm: *k_norm,
714 output_gate: *output_gate,
715 bias: *bias,
716 rms_eps: eps,
717 norm_style,
718 pool: pool.as_deref(),
719 };
720 let mut ao = attention::qwen_attention_core(
721 q_raw,
722 k,
723 v,
724 &mut self.kv_cache.layers[*li],
725 &cfg,
726 );
727 graph.encode_attn_suffix(l, &ao);
728 graph.commit();
731 attention::recycle_buf(&mut ao);
732 }
733 }
734 }
735 graph.sync();
736 if !pending.is_empty() {
737 let idxs: Vec<usize> = pending.drain(..).flat_map(|(f, n)| f..f + n).collect();
738 let mut outs: Vec<&mut [f32]> = self
739 .kv_cache
740 .layers
741 .iter_mut()
742 .enumerate()
743 .filter(|(i, _)| idxs.binary_search(i).is_ok())
744 .map(|(_, s)| s.linear_state.as_mut_slice())
745 .collect();
746 graph.read_states(&mut outs);
747 }
748 graph.finish(h);
749 for li in dev_attn {
753 let mut krow = attention::take_buf(nkv * hd);
754 let mut vrow = attention::take_buf(nkv * hd);
755 if crate::gpu::kv_mirror_read_last(kv_id, li, nkv, hd, &mut krow, &mut vrow) {
756 let cache = &mut self.kv_cache.layers[li];
757 cache.append(&krow, &vrow, &[]);
758 let n = cache.seq_len;
759 let mut imp = attention::take_buf(n);
760 crate::gpu::kv_mirror_take_imp(kv_id, li, &mut imp);
761 cache.accumulate_imp(&imp);
762 attention::recycle_buf(&mut imp);
763 }
764 attention::recycle_buf(&mut krow);
765 attention::recycle_buf(&mut vrow);
766 }
767 end
768 }
769
770 pub fn new(
771 tokenizer: Tokenizer,
772 weights: PipelineWeights,
773 hidden_size: usize,
774 intermediate_size: usize,
775 num_heads: usize,
776 num_kv_heads: usize,
777 head_dim: usize,
778 num_layers: usize,
779 vocab_size: usize,
780 rms_eps: f64,
781 rope_base: f32,
782 norm_style: NormStyle,
783 max_seq_len: usize,
784 sampler_config: SamplerConfig,
785 ) -> Self {
786 let rng = match sampler_config.seed {
787 Some(s) => SplitMix64::new(s),
788 None => SplitMix64::from_entropy(),
789 };
790 let inv_freq = std::sync::Arc::new(attention::rope_inv_freq(head_dim, rope_base));
791 let pool = Pool::from_env();
792 if let Some(p) = &pool {
793 tracing::info!("worker pool: {} threads", p.n_workers());
794 }
795 Self {
796 tokenizer: std::sync::Arc::new(tokenizer),
797 kv_cache: KvCache::new(num_layers, num_kv_heads, head_dim, max_seq_len),
798 sampler_config,
799 weights,
800 hidden_size,
801 intermediate_size,
802 num_heads,
803 num_kv_heads,
804 head_dim,
805 num_layers,
806 vocab_size,
807 rms_eps,
808 rope_base,
809 norm_style,
810 rotary_dim: head_dim,
811 vmf_cfg: None,
812 gdn_cfg: None,
813 mtp: None,
814 speculative: std::env::var("CMF_MTP").map(|v| v != "0").unwrap_or(true),
815 rng,
816 inv_freq,
817 ws: ForwardScratch::new(hidden_size),
818 pool,
819 model: None,
820 dyn_force_f32: false,
821 dyn_skill_layers: Vec::new(),
822 dyn_active: None,
823 dyn_blend_loaded: false,
824 dyn_phi_layer: None,
825 dyn_phi_ema: Vec::new(),
826 dyn_phi_seen: 0,
827 dyn_router: None,
828 o1_cfg: None,
829 o1_flags: Vec::new(),
830 trace: false,
831 calib_temp: 1.0,
832 confidence_on: true,
833 embed_multiplier: 1.0,
834 attn_scale: 1.0 / (head_dim as f32).sqrt(),
835 swa: None,
836 inv_freq_local: None,
837 global_attn: None,
838 inv_freq_global: None,
839 attn_v_norm: false,
840 final_softcap: None,
841 graph_kv_id: {
842 static NEXT: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(1);
843 NEXT.fetch_add(1, std::sync::atomic::Ordering::Relaxed)
844 },
845 }
846 }
847
848 pub fn set_o1(&mut self, cfg: Option<crate::nystrom::O1Cfg>) {
855 self.o1_flags = match &cfg {
856 Some(c) => {
857 let mut flags = c.layer_flags(self.num_layers);
858 for (li, f) in flags.iter_mut().enumerate() {
859 if *f && !matches!(self.weights.layers[li].attn, AttnKind::Full { .. }) {
860 *f = false;
861 }
862 }
863 flags
864 }
865 None => Vec::new(),
866 };
867 if let Some(c) = &cfg {
868 let n = self.o1_flags.iter().filter(|&&f| f).count();
869 tracing::info!(
870 "o1 nystrom attention: {n}/{} layer(s), m={} w={} sink={} rect={:?}",
871 self.num_layers, c.m, c.w, c.sink, c.rect
872 );
873 }
874 self.o1_cfg = cfg;
875 }
876
877 pub fn o1_active(&self) -> bool {
879 self.o1_cfg.is_some() && self.o1_flags.iter().any(|&f| f)
880 }
881
882 fn o1_begin(&mut self) {
884 if let Some(c) = &self.o1_cfg {
885 let (m, w, sink, rect) = (c.m, c.w, c.sink, c.rect);
886 for (li, &f) in self.o1_flags.iter().enumerate() {
887 if f {
888 self.kv_cache.layers[li].o1_begin(m, w, sink, rect);
889 }
890 }
891 }
892 }
893
894 fn o1_seal(&mut self) {
897 if self.o1_cfg.is_none() {
898 return;
899 }
900 for li in 0..self.num_layers {
901 if self.o1_flags.get(li).copied().unwrap_or(false) {
902 self.kv_cache.layers[li].o1_seal(self.num_heads);
903 }
904 }
905 }
906
907 pub fn set_trace(&mut self, on: bool) {
909 self.trace = on;
910 }
911
912 pub fn set_confidence(&mut self, on: bool) {
917 self.confidence_on = on;
918 }
919
920 pub fn set_calib_temp(&mut self, t: f32) {
923 self.calib_temp = if t > 1e-3 { t } else { 1.0 };
924 }
925
926 pub fn calib_temp(&self) -> f32 {
928 self.calib_temp
929 }
930
931 pub fn set_rotary(&mut self, rotary_dim: usize, base: f32) {
934 self.rotary_dim = rotary_dim.min(self.head_dim);
935 self.inv_freq = std::sync::Arc::new(attention::rope_inv_freq(self.rotary_dim, base));
936 }
937
938 fn attn_cfg(&self, position: usize) -> QwenAttnCfg<'_> {
939 QwenAttnCfg {
940 num_heads: self.num_heads,
941 num_kv_heads: self.num_kv_heads,
942 head_dim: self.head_dim,
943 hidden_size: self.hidden_size,
944 position,
945 inv_freq: &self.inv_freq,
946 rotary_dim: self.rotary_dim,
947 scale: self.attn_scale,
948 window: None,
949 v_norm: false,
950 q_norm: None,
951 k_norm: None,
952 output_gate: false,
953 bias: None,
954 rms_eps: self.rms_eps,
955 norm_style: self.norm_style,
956 pool: self.pool.as_deref(),
957 }
958 }
959
960 pub fn generate(
962 &mut self,
963 prompt: &str,
964 max_tokens: usize,
965 task_mask: Option<&TaskMask>,
966 on_token: Option<TokenCallback>,
967 ) -> Result<GenerateResult, String> {
968 let input_ids = self.tokenizer.with_bos(self.tokenizer.encode(prompt));
969 self.generate_from_ids(&input_ids, max_tokens, task_mask, on_token)
970 }
971
972 pub fn generate_from_ids(
980 &mut self,
981 input_ids: &[u32],
982 max_tokens: usize,
983 task_mask: Option<&TaskMask>,
984 mut on_token: Option<TokenCallback>,
985 ) -> Result<GenerateResult, String> {
986 if std::env::var("CMF_TRACE_H").is_ok() {
987 eprintln!("input_ids: {input_ids:?}");
988 }
989 if input_ids.is_empty() {
990 return Err("empty prompt: nothing to generate from".to_string());
991 }
992
993 self.kv_cache.clear();
995 self.o1_begin();
996
997 let spec_active = self.speculative
1001 && self.mtp.is_some()
1002 && task_mask.is_none()
1003 && !self.o1_active()
1004 && self.sampler_config.temperature < 1e-6;
1005 let mut mtp = if spec_active { self.mtp.take() } else { None };
1008 if let Some(m) = &mut mtp {
1009 m.kv.clear();
1010 }
1011 let mut router = if mtp.is_none() { self.dyn_router.take() } else { None };
1015 if let Some(r) = &mut router {
1016 r.reset(); self.dyn_phi_seen = 0; let _ = self.set_active_skill(None);
1019 }
1020
1021 let mut all_ids = input_ids.to_vec();
1022 let mut generated = 0usize;
1023 let mut finish_reason = "max_tokens".to_string();
1024 let mut drafted = 0usize;
1025 let mut accepted = 0usize;
1026 let mut confidence: Vec<f32> = Vec::new();
1027 let trace_on = self.trace;
1028 let calib_temp = self.calib_temp;
1029 let mut traces: Vec<TokenTrace> = Vec::new();
1030
1031 let mut hidden = vec![0.0f32; self.hidden_size];
1037 let mut pos = 0usize;
1038 let dyn_prefill = router.is_some();
1043 let graph_prefill = self.graph_prefill_preferred();
1049 if task_mask.is_none()
1050 && !dyn_prefill
1051 && !graph_prefill
1052 && prefill_batched()
1053 && input_ids.len() > 2
1054 {
1055 let chunk = prefill_chunk();
1061 let hs = self.hidden_size;
1062 while pos < input_ids.len() {
1063 let end = (pos + chunk).min(input_ids.len());
1064 let hb = self.prefill_batch(&input_ids[pos..end], pos);
1065 if let Some(m) = &mut mtp {
1066 for p in pos..end {
1067 if p + 1 < input_ids.len() {
1068 let _ = self.mtp_step(
1069 m,
1070 &hb[(p - pos) * hs..(p - pos + 1) * hs],
1071 input_ids[p + 1],
1072 p,
1073 );
1074 }
1075 }
1076 }
1077 hidden.copy_from_slice(&hb[(end - pos - 1) * hs..]);
1078 pos = end;
1079 }
1080 }
1081 if task_mask.is_none() && !dyn_prefill && !graph_prefill {
1082 while pos + 1 < input_ids.len() {
1083 let e1 = self.embed_single(input_ids[pos]);
1084 let e2 = self.embed_single(input_ids[pos + 1]);
1085 let (h1, h2) = self.forward_pair(&e1, &e2, pos);
1086 self.commit_linear_scratch();
1088 if let Some(m) = &mut mtp {
1089 let _ = self.mtp_step(m, &h1, input_ids[pos + 1], pos);
1090 if pos + 2 < input_ids.len() {
1091 let _ = self.mtp_step(m, &h2, input_ids[pos + 2], pos + 1);
1092 }
1093 }
1094 hidden = h2;
1095 pos += 2;
1096 }
1097 }
1098 while pos < input_ids.len() {
1099 hidden = self.forward_layers(&self.embed_single(input_ids[pos]), pos, task_mask);
1100 if let Some(m) = &mut mtp {
1101 if pos + 1 < input_ids.len() {
1102 let _ = self.mtp_step(m, &hidden, input_ids[pos + 1], pos);
1103 }
1104 }
1105 pos += 1;
1106 }
1107 self.o1_seal();
1110
1111 macro_rules! commit {
1113 ($id:expr) => {{
1114 all_ids.push($id);
1115 generated += 1;
1116 if self.tokenizer.is_eos($id) {
1117 finish_reason = "stop".to_string();
1118 false
1119 } else {
1120 let token_text = self.tokenizer.decode_token($id);
1121 let mut go = true;
1122 if let Some(ref mut cb) = on_token {
1123 if !cb(&token_text) {
1124 finish_reason = "cancelled".to_string();
1125 go = false;
1126 }
1127 }
1128 go
1129 }
1130 }};
1131 }
1132
1133 let mut next_pos = input_ids.len();
1135 'decode: while generated < max_tokens {
1136 inference::rms_norm_into(
1137 &hidden,
1138 &self.weights.final_norm,
1139 self.rms_eps,
1140 self.norm_style,
1141 &mut self.ws.n1,
1142 );
1143 let mut logits = self.lm_head_forward(&self.ws.n1);
1144 let t_next = sampler::sample(&logits, &self.sampler_config, &all_ids, &mut self.rng);
1145 if self.confidence_on {
1146 confidence.push(top1_prob_t(&logits, t_next, calib_temp));
1147 }
1148 attention::recycle_buf(&mut logits);
1149 if trace_on {
1150 let skill = router.as_ref().and_then(|r| r.active_id());
1154 traces.push(TokenTrace {
1155 t: generated,
1156 token_id: t_next,
1157 confidence: confidence.last().copied().unwrap_or(0.0),
1158 active_skill: skill,
1159 recon: None,
1160 switched: false,
1161 });
1162 }
1163 if !commit!(t_next) {
1164 break 'decode;
1165 }
1166 if generated >= max_tokens {
1167 break 'decode;
1168 }
1169
1170 if self.kv_cache.needs_eviction() {
1171 let keep = (self.kv_cache.max_seq_len / 2).max(1);
1172 self.kv_cache.evict(keep);
1173 }
1174
1175 match &mut mtp {
1176 Some(m) if generated + 1 < max_tokens => {
1178 let draft = self.mtp_step(m, &hidden, t_next, next_pos - 1);
1179 drafted += 1;
1180 let emb1 = self.embed_single(t_next);
1181 let emb2 = self.embed_single(draft);
1182 let (h1, h2) = self.forward_pair(&emb1, &emb2, next_pos);
1183
1184 inference::rms_norm_into(
1185 &h1,
1186 &self.weights.final_norm,
1187 self.rms_eps,
1188 self.norm_style,
1189 &mut self.ws.n1,
1190 );
1191 let mut logits1 = self.lm_head_forward(&self.ws.n1);
1192 let t_after =
1193 sampler::sample(&logits1, &self.sampler_config, &all_ids, &mut self.rng);
1194 if self.confidence_on {
1195 confidence.push(top1_prob_t(&logits1, t_after, calib_temp));
1196 }
1197 attention::recycle_buf(&mut logits1);
1198 if trace_on {
1199 traces.push(TokenTrace {
1202 t: generated,
1203 token_id: t_after,
1204 confidence: confidence.last().copied().unwrap_or(0.0),
1205 active_skill: None,
1206 recon: None,
1207 switched: false,
1208 });
1209 }
1210 let stop = !commit!(t_after);
1211
1212 if t_after == draft {
1213 accepted += 1;
1214 self.commit_linear_scratch();
1215 let _ = self.mtp_step(m, &h1, t_after, next_pos);
1216 hidden = h2;
1217 next_pos += 2;
1218 } else {
1219 for layer in &mut self.kv_cache.layers {
1221 layer.truncate_last(1);
1222 }
1223 if !stop {
1224 let _ = self.mtp_step(m, &h1, t_after, next_pos);
1225 hidden = self
1226 .forward_layers(&self.embed_single(t_after), next_pos + 1, None);
1227 }
1228 next_pos += 2;
1229 }
1230 if stop {
1231 break 'decode;
1232 }
1233 }
1234 _ => {
1236 hidden = self.forward_layers(&self.embed_single(t_next), next_pos, task_mask);
1237 next_pos += 1;
1238 if let Some(r) = &mut router {
1241 let phi = self.dyn_phi_ema.clone();
1242 let decision = r.step(&phi, generated);
1243 if let Some(new_active) = decision {
1244 let _ = self.set_active_skill(new_active);
1245 }
1246 if trace_on {
1249 if let Some(last) = traces.last_mut() {
1250 let e = r.last_best_e();
1251 last.recon = e.is_finite().then_some(e);
1252 last.switched = decision.is_some();
1253 }
1254 }
1255 }
1256 }
1257 }
1258 }
1259
1260 if router.is_some() {
1262 let _ = self.set_active_skill(None);
1263 }
1264 self.dyn_router = router.or(self.dyn_router.take());
1265 self.mtp = mtp.or(self.mtp.take());
1266
1267 let output_ids = &all_ids[input_ids.len()..];
1268 confidence.truncate(output_ids.len()); traces.truncate(output_ids.len());
1270 Ok(GenerateResult {
1271 text: self.tokenizer.decode(output_ids),
1272 token_ids: output_ids.to_vec(),
1273 prompt_tokens: input_ids.len(),
1274 tokens_generated: generated,
1275 finish_reason,
1276 mtp_drafted: drafted,
1277 mtp_accepted: accepted,
1278 token_confidence: confidence,
1279 traces,
1280 })
1281 }
1282
1283 fn mtp_step(&mut self, m: &mut MtpModule, hidden: &[f32], next_token: u32, position: usize) -> u32 {
1287 let e = self.embed_single(next_token);
1291 let mut cat = vec![0.0f32; 2 * self.hidden_size];
1292 let (cat_e, cat_h) = cat.split_at_mut(self.hidden_size);
1293 inference::rms_norm_into(&e, &m.enorm, self.rms_eps, self.norm_style, cat_e);
1294 inference::rms_norm_into(hidden, &m.hnorm, self.rms_eps, self.norm_style, cat_h);
1295 let mut x = vec![0.0f32; self.hidden_size];
1296 m.eh_proj.matvec(&cat, &mut x, self.pool.as_deref());
1297
1298 let lw = &m.layer;
1300 inference::rms_norm_into(&x, &lw.input_norm, self.rms_eps, self.norm_style, &mut self.ws.n1);
1301 let attn = match &lw.attn {
1302 AttnKind::Full {
1303 wq,
1304 wk,
1305 wv,
1306 wo,
1307 q_norm,
1308 k_norm,
1309 output_gate,
1310 bias,
1311 } => {
1312 let mut cfg = self.attn_cfg(position);
1313 cfg.q_norm = q_norm.as_deref();
1314 cfg.k_norm = k_norm.as_deref();
1315 cfg.output_gate = *output_gate;
1316 cfg.bias = bias
1317 .as_ref()
1318 .map(|(q, k, v)| (q.as_slice(), k.as_slice(), v.as_slice()));
1319 attention::qwen_attention(&self.ws.n1, wq, wk, wv, wo, &mut m.kv, &cfg)
1320 }
1321 AttnKind::Linear(_) | AttnKind::LinearGdn(_) => {
1322 unreachable!("MTP block is full attention")
1323 }
1324 };
1325 for (i, &a) in attn.iter().enumerate() {
1326 x[i] += a;
1327 }
1328 inference::rms_norm_into(&x, &lw.post_norm, self.rms_eps, self.norm_style, &mut self.ws.p1);
1329 let ffn = ffn_forward(&lw.ffn, &self.ws.p1, self.pool.as_deref());
1330 for (i, &f) in ffn.iter().enumerate() {
1331 x[i] += f;
1332 }
1333
1334 inference::rms_norm_into(&x, &m.final_norm, self.rms_eps, self.norm_style, &mut self.ws.n1);
1335 let mut lg = self.lm_head_forward(&self.ws.n1);
1336 let draft = sampler::argmax(&lg);
1337 attention::recycle_buf(&mut lg);
1338 draft
1339 }
1340
1341 pub fn measure_pair_fusion(&mut self, iters: usize) -> (f64, f64) {
1345 let emb1 = self.embed_single(1);
1346 let emb2 = self.embed_single(2);
1347 let pos = self.kv_cache.seq_len();
1348
1349 let t0 = std::time::Instant::now();
1350 for _ in 0..iters {
1351 let _ = self.forward_layers(&emb1, pos, None);
1352 let _ = self.forward_layers(&emb2, pos + 1, None);
1353 for l in &mut self.kv_cache.layers {
1354 l.truncate_last(2);
1355 }
1356 }
1357 let singles_ms = t0.elapsed().as_secs_f64() * 1000.0 / iters as f64;
1358
1359 let t1 = std::time::Instant::now();
1360 for _ in 0..iters {
1361 let _ = self.forward_pair(&emb1, &emb2, pos);
1362 for l in &mut self.kv_cache.layers {
1363 l.truncate_last(2);
1364 }
1365 }
1366 let pair_ms = t1.elapsed().as_secs_f64() * 1000.0 / iters as f64;
1367 (singles_ms, pair_ms)
1368 }
1369
1370 fn forward_pair(&mut self, emb1: &[f32], emb2: &[f32], position: usize) -> (Vec<f32>, Vec<f32>) {
1375 let mut h1 = emb1.to_vec();
1376 let mut h2 = emb2.to_vec();
1377 let (nh, _nkv, _hd, hs, _rd, eps) = (
1378 self.num_heads,
1379 self.num_kv_heads,
1380 self.head_dim,
1381 self.hidden_size,
1382 self.rotary_dim,
1383 self.rms_eps,
1384 );
1385 let pool = self.pool.clone();
1386
1387 for li in 0..self.num_layers {
1388 let lw = &self.weights.layers[li];
1389 inference::rms_norm_into(&h1, &lw.input_norm, self.rms_eps, self.norm_style, &mut self.ws.n1);
1392 inference::rms_norm_into(&h2, &lw.input_norm, self.rms_eps, self.norm_style, &mut self.ws.n2);
1393
1394 let (a1, a2) = match &lw.attn {
1395 AttnKind::Linear(w) => {
1396 let cfg = self.vmf_cfg.expect("linear layer without vmf_cfg");
1397 let layer = &mut self.kv_cache.layers[li];
1398 let (state, scratch) = (&mut layer.linear_state, &mut layer.linear_scratch);
1399 vmf_phase_pair(&self.ws.n1, &self.ws.n2, w, &cfg, state, scratch, self.pool.as_deref())
1400 }
1401 AttnKind::LinearGdn(w) => {
1402 let cfg = self.gdn_cfg.expect("gdn layer without gdn_cfg");
1403 let layer = &mut self.kv_cache.layers[li];
1404 let (state, scratch) = (&mut layer.linear_state, &mut layer.linear_scratch);
1405 gdn_pair(&self.ws.n1, &self.ws.n2, w, &cfg, state, scratch, self.pool.as_deref())
1406 }
1407 AttnKind::Full {
1408 wq,
1409 wk,
1410 wv,
1411 wo,
1412 q_norm,
1413 k_norm,
1414 output_gate,
1415 bias,
1416 } => {
1417 let inv_freq_l = self.layer_inv_freq(li);
1418 let (nkv_l, hd_l, rd_l) = self.layer_geom(li);
1419 let cfg = QwenAttnCfg {
1420 num_heads: nh,
1421 num_kv_heads: nkv_l,
1422 head_dim: hd_l,
1423 hidden_size: hs,
1424 position,
1425 inv_freq: &inv_freq_l,
1426 rotary_dim: rd_l,
1427 scale: self.attn_scale,
1428 window: self.layer_window(li),
1429 v_norm: self.attn_v_norm,
1430 q_norm: q_norm.as_deref(),
1431 k_norm: k_norm.as_deref(),
1432 output_gate: *output_gate,
1433 bias: bias
1434 .as_ref()
1435 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
1436 rms_eps: eps,
1437 norm_style: self.norm_style,
1438 pool: pool.as_deref(),
1439 };
1440 attention::qwen_attention_pair(
1441 &self.ws.n1,
1442 &self.ws.n2,
1443 wq,
1444 wk,
1445 wv,
1446 wo,
1447 &mut self.kv_cache.layers[li],
1448 &cfg,
1449 )
1450 }
1451 };
1452 let (a1, a2) = match &self.weights.layers[li].attn_out_norm {
1453 Some(w) => (
1454 inference::rms_norm(&a1, w, self.rms_eps, self.norm_style),
1455 inference::rms_norm(&a2, w, self.rms_eps, self.norm_style),
1456 ),
1457 None => (a1, a2),
1458 };
1459 for i in 0..self.hidden_size {
1460 h1[i] += a1[i];
1461 h2[i] += a2[i];
1462 }
1463 let (mut a1, mut a2) = (a1, a2);
1464 attention::recycle_buf(&mut a1);
1465 attention::recycle_buf(&mut a2);
1466
1467 let lw = &self.weights.layers[li];
1468 inference::rms_norm_into(&h1, &lw.post_norm, self.rms_eps, self.norm_style, &mut self.ws.p1);
1469 inference::rms_norm_into(&h2, &lw.post_norm, self.rms_eps, self.norm_style, &mut self.ws.p2);
1470 let (f1, f2) =
1471 ffn_forward_pair(&lw.ffn, &self.ws.p1, &self.ws.p2, self.pool.as_deref());
1472 let (f1, f2) = match &self.weights.layers[li].ffn_out_norm {
1473 Some(w) => (
1474 inference::rms_norm(&f1, w, self.rms_eps, self.norm_style),
1475 inference::rms_norm(&f2, w, self.rms_eps, self.norm_style),
1476 ),
1477 None => (f1, f2),
1478 };
1479 for i in 0..self.hidden_size {
1480 h1[i] += f1[i];
1481 h2[i] += f2[i];
1482 }
1483 let (mut f1, mut f2) = (f1, f2);
1484 attention::recycle_buf(&mut f1);
1485 attention::recycle_buf(&mut f2);
1486 if let Some(sc) = self.weights.layers[li].layer_scale {
1487 for i in 0..self.hidden_size {
1488 h1[i] *= sc;
1489 h2[i] *= sc;
1490 }
1491 }
1492 }
1493 (h1, h2)
1494 }
1495
1496 fn commit_linear_scratch(&mut self) {
1498 for layer in &mut self.kv_cache.layers {
1499 if !layer.linear_scratch.is_empty() {
1500 std::mem::swap(&mut layer.linear_state, &mut layer.linear_scratch);
1501 layer.linear_scratch.clear();
1502 }
1503 }
1504 }
1505
1506 pub fn forward_ids(
1509 &mut self,
1510 ids: &[u32],
1511 task_mask: Option<&TaskMask>,
1512 ) -> Result<Vec<f32>, String> {
1513 if ids.is_empty() {
1514 return Err("empty id sequence".to_string());
1515 }
1516 self.kv_cache.clear();
1517 self.o1_begin();
1518 let mut hidden = vec![0.0f32; self.hidden_size];
1519 let mut pos = 0usize;
1520 if task_mask.is_none() && prefill_batched() && ids.len() > 2 {
1521 let chunk = prefill_chunk();
1525 let hs = self.hidden_size;
1526 while pos < ids.len() {
1527 let end = (pos + chunk).min(ids.len());
1528 let hb = self.prefill_batch(&ids[pos..end], pos);
1529 hidden.copy_from_slice(&hb[(end - pos - 1) * hs..]);
1530 pos = end;
1531 }
1532 }
1533 if task_mask.is_none() {
1534 while pos + 1 < ids.len() {
1535 let e1 = self.embed_single(ids[pos]);
1536 let e2 = self.embed_single(ids[pos + 1]);
1537 let (_, h2) = self.forward_pair(&e1, &e2, pos);
1538 self.commit_linear_scratch();
1539 hidden = h2;
1540 pos += 2;
1541 }
1542 }
1543 while pos < ids.len() {
1544 hidden = self.forward_layers(&self.embed_single(ids[pos]), pos, task_mask);
1545 pos += 1;
1546 }
1547 self.o1_seal();
1551 let normed = inference::rms_norm(
1552 &hidden,
1553 &self.weights.final_norm,
1554 self.rms_eps,
1555 self.norm_style,
1556 );
1557 Ok(self.lm_head_forward(&normed))
1558 }
1559
1560 pub fn ppl_ids(&mut self, ids: &[u32]) -> f64 {
1567 let (nll, cnt) = self.nll_ids_from(ids, 0);
1568 (nll / cnt.max(1) as f64).exp()
1569 }
1570
1571 pub fn probe_ffn_mass(&mut self, ids: &[u32]) -> Vec<Vec<f64>> {
1576 self.kv_cache.clear();
1577 FFN_PROBE.with(|p| {
1578 *p.borrow_mut() =
1579 Some(vec![vec![0f64; self.intermediate_size]; self.num_layers]);
1580 });
1581 crate::gpu::cpu_scope(|| {
1582 for (pos, &id) in ids.iter().enumerate() {
1583 let emb = self.embed_single(id);
1584 let _ = self.forward_layers(&emb, pos, None);
1585 }
1586 });
1587 self.kv_cache.clear();
1588 FFN_PROBE.with(|p| p.borrow_mut().take()).unwrap_or_default()
1589 }
1590
1591 pub fn ppl_ids_masked(&mut self, ids: &[u32], mask: &TaskMask) -> f64 {
1595 self.kv_cache.clear();
1596 let mut nll = 0f64;
1597 let mut cnt = 0usize;
1598 let mut hidden = vec![0f32; self.hidden_size];
1599 for (pos, &id) in ids.iter().enumerate() {
1600 if pos > 0 {
1601 inference::rms_norm_into(
1602 &hidden,
1603 &self.weights.final_norm,
1604 self.rms_eps,
1605 self.norm_style,
1606 &mut self.ws.n1,
1607 );
1608 let mut logits = self.lm_head_forward(&self.ws.n1);
1609 let max = logits.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
1610 let sum: f64 = logits.iter().map(|&v| ((v - max) as f64).exp()).sum();
1611 let p = ((logits[id as usize] - max) as f64).exp() / sum.max(1e-300);
1612 nll -= p.max(1e-300).ln();
1613 cnt += 1;
1614 attention::recycle_buf(&mut logits);
1615 }
1616 let emb = self.embed_single(id);
1617 hidden = self.forward_layers(&emb, pos, Some(mask));
1618 }
1619 self.kv_cache.clear();
1620 (nll / cnt.max(1) as f64).exp()
1621 }
1622
1623 pub fn nll_ids_from(&mut self, ids: &[u32], start: usize) -> (f64, usize) {
1632 self.kv_cache.clear();
1633 let mut nll = 0f64;
1634 let mut cnt = 0usize;
1635 if prefill_batched() {
1636 const CHUNK: usize = 128;
1642 const LM_SUB: usize = 32;
1643 let n = ids.len().saturating_sub(1);
1644 let hs = self.hidden_size;
1645 let rows = self.weights.lm_head.rows();
1646 let mut pos = 0usize;
1647 while pos < n {
1648 let end = (pos + CHUNK).min(n);
1649 let bsz = end - pos;
1650 let hb = self.prefill_batch(&ids[pos..end], pos);
1651 let mut k0 = 0usize;
1652 while k0 < bsz {
1653 let k1 = (k0 + LM_SUB).min(bsz);
1654 let sb = k1 - k0;
1655 if pos + k1 <= start {
1658 k0 = k1;
1659 continue;
1660 }
1661 let mut normed = vec![0.0f32; sb * hs];
1662 for k in 0..sb {
1663 let r = inference::rms_norm(
1664 &hb[(k0 + k) * hs..(k0 + k + 1) * hs],
1665 &self.weights.final_norm,
1666 self.rms_eps,
1667 self.norm_style,
1668 );
1669 normed[k * hs..(k + 1) * hs].copy_from_slice(&r);
1670 }
1671 let mut logits = vec![0.0f32; sb * rows];
1672 self.weights
1673 .lm_head
1674 .matmat(&normed, sb, &mut logits, self.pool.as_deref());
1675 for k in 0..sb {
1676 if pos + k0 + k < start {
1677 continue;
1678 }
1679 let lg = &logits[k * rows..k * rows + self.vocab_size.min(rows)];
1680 let target = ids[pos + k0 + k + 1] as usize;
1681 let max = lg.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
1682 let lse: f64 = lg
1683 .iter()
1684 .map(|&v| ((v - max) as f64).exp())
1685 .sum::<f64>()
1686 .ln()
1687 + max as f64;
1688 nll += lse - lg[target] as f64;
1689 cnt += 1;
1690 }
1691 k0 = k1;
1692 }
1693 pos = end;
1694 }
1695 self.kv_cache.clear();
1696 return (nll, cnt);
1697 }
1698 for pos in 0..ids.len().saturating_sub(1) {
1699 let hidden = self.forward_layers(&self.embed_single(ids[pos]), pos, None);
1700 if pos < start {
1701 continue;
1702 }
1703 let normed = inference::rms_norm(
1704 &hidden,
1705 &self.weights.final_norm,
1706 self.rms_eps,
1707 self.norm_style,
1708 );
1709 let logits = self.lm_head_forward(&normed);
1710 let target = ids[pos + 1] as usize;
1711 let max = logits.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
1712 let lse: f64 = logits.iter().map(|&v| ((v - max) as f64).exp()).sum::<f64>().ln()
1713 + max as f64;
1714 nll += lse - logits[target] as f64;
1715 cnt += 1;
1716 }
1717 self.kv_cache.clear();
1718 (nll, cnt)
1719 }
1720
1721 pub fn nll_ids_o1(&mut self, ids: &[u32], prefill: usize) -> (f64, usize) {
1737 self.kv_cache.clear();
1738 self.o1_begin();
1739 let n = ids.len().saturating_sub(1);
1740 let p = prefill.min(n);
1741 let mut pos = 0usize;
1743 if prefill_batched() {
1744 const CHUNK: usize = 128;
1745 while pos < p {
1746 let end = (pos + CHUNK).min(p);
1747 let _ = self.prefill_batch(&ids[pos..end], pos);
1748 pos = end;
1749 }
1750 } else {
1751 while pos < p {
1752 let _ = self.forward_layers(&self.embed_single(ids[pos]), pos, None);
1753 pos += 1;
1754 }
1755 }
1756 self.o1_seal();
1757
1758 let mut nll = 0f64;
1759 let mut cnt = 0usize;
1760 for pos in p..n {
1761 let hidden = self.forward_layers(&self.embed_single(ids[pos]), pos, None);
1762 let normed = inference::rms_norm(
1763 &hidden,
1764 &self.weights.final_norm,
1765 self.rms_eps,
1766 self.norm_style,
1767 );
1768 let logits = self.lm_head_forward(&normed);
1769 let target = ids[pos + 1] as usize;
1770 let max = logits.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
1771 let lse: f64 = logits.iter().map(|&v| ((v - max) as f64).exp()).sum::<f64>().ln()
1772 + max as f64;
1773 nll += lse - logits[target] as f64;
1774 cnt += 1;
1775 }
1776 self.kv_cache.clear();
1777 (nll, cnt)
1778 }
1779
1780 pub fn calib_ids(&mut self, ids: &[u32], temps: &[f32]) -> (Vec<bool>, Vec<Vec<f32>>) {
1788 self.kv_cache.clear();
1789 let n = ids.len().saturating_sub(1);
1790 let mut correct = Vec::with_capacity(n);
1791 let mut pmax = Vec::with_capacity(n);
1792 for pos in 0..n {
1793 let emb = self.embed_single(ids[pos]);
1794 let hidden = self.forward_layers(&emb, pos, None);
1795 let normed = inference::rms_norm(
1796 &hidden,
1797 &self.weights.final_norm,
1798 self.rms_eps,
1799 self.norm_style,
1800 );
1801 let logits = self.lm_head_forward(&normed);
1802 let target = ids[pos + 1] as usize;
1803 let (mut amax, mut mval) = (0usize, f32::NEG_INFINITY);
1804 for (i, &v) in logits.iter().enumerate() {
1805 if v > mval {
1806 mval = v;
1807 amax = i;
1808 }
1809 }
1810 correct.push(amax == target);
1811 let row: Vec<f32> = temps
1812 .iter()
1813 .map(|&t| {
1814 let tt = t.max(1e-3);
1815 let s: f32 = logits.iter().map(|&v| ((v - mval) / tt).exp()).sum();
1816 1.0 / s.max(1e-12) })
1818 .collect();
1819 pmax.push(row);
1820 }
1821 self.kv_cache.clear();
1822 (correct, pmax)
1823 }
1824
1825 pub fn ppl_ids_dynamic(&mut self, ids: &[u32]) -> (f64, usize) {
1832 let mut router = match self.dyn_router.take() {
1833 Some(r) => r,
1834 None => return (self.ppl_ids(ids), 0),
1835 };
1836 router.reset();
1837 self.dyn_phi_seen = 0;
1838 let _ = self.set_active_skill(None);
1839
1840 self.kv_cache.clear();
1841 let mut nll = 0f64;
1842 let mut cnt = 0usize;
1843 for pos in 0..ids.len().saturating_sub(1) {
1844 let hidden = self.forward_layers(&self.embed_single(ids[pos]), pos, None);
1845 let normed = inference::rms_norm(
1846 &hidden,
1847 &self.weights.final_norm,
1848 self.rms_eps,
1849 self.norm_style,
1850 );
1851 let logits = self.lm_head_forward(&normed);
1852 let target = ids[pos + 1] as usize;
1853 let max = logits.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
1854 let lse: f64 = logits.iter().map(|&v| ((v - max) as f64).exp()).sum::<f64>().ln()
1855 + max as f64;
1856 nll += lse - logits[target] as f64;
1857 cnt += 1;
1858 let phi = self.dyn_phi_ema.clone();
1860 if let Some(new_active) = router.step(&phi, pos) {
1861 let _ = self.set_active_skill(new_active);
1862 }
1863 }
1864 let switches = router.switches.len();
1865 let _ = self.set_active_skill(None);
1866 self.dyn_router = Some(router);
1867 self.kv_cache.clear();
1868 ((nll / cnt.max(1) as f64).exp(), switches)
1869 }
1870
1871 pub fn probe_phi(&mut self, ids: &[u32], layer: usize) -> Vec<f32> {
1873 self.kv_cache.clear();
1874 let mut acc = vec![0f32; self.hidden_size];
1875 for (pos, &id) in ids.iter().enumerate() {
1876 let h = self.forward_layers_upto(&self.embed_single(id), pos, None, Some(layer));
1877 for (a, v) in acc.iter_mut().zip(&h) {
1878 *a += v;
1879 }
1880 }
1881 let n = ids.len().max(1) as f32;
1882 for a in acc.iter_mut() {
1883 *a /= n;
1884 }
1885 self.kv_cache.clear();
1886 acc
1887 }
1888
1889 fn prefill_batch(&mut self, ids: &[u32], start_pos: usize) -> Vec<f32> {
1895 let b = ids.len();
1896 let hs = self.hidden_size;
1897 let mut h: Vec<f32> = vec![0.0; b * hs];
1900 let mut h_ready = false;
1901 let mut fill_h = |h: &mut Vec<f32>, me: &Self| {
1902 for (bi, &id) in ids.iter().enumerate() {
1903 let e = me.embed_single(id);
1904 h[bi * hs..(bi + 1) * hs].copy_from_slice(&e);
1905 }
1906 };
1907 let (nh, _nkv, _hd, _rd, eps) = (
1908 self.num_heads,
1909 self.num_kv_heads,
1910 self.head_dim,
1911 self.rotary_dim,
1912 self.rms_eps,
1913 );
1914 let pool = self.pool.clone();
1915 let norm_style = self.norm_style;
1916
1917 #[cfg(target_os = "macos")]
1918 let mut chunk_skip_until = 0usize;
1919 for li in 0..self.num_layers {
1920 crate::gpu::set_layer(li as i64); #[cfg(target_os = "macos")]
1927 {
1928 if li < chunk_skip_until {
1929 continue;
1930 }
1931 let ids_for_embed = (!h_ready && li == 0).then_some(ids);
1932 let end = self.chunk_run_gpu(li, &mut h, b, start_pos, ids_for_embed);
1933 if end > li {
1934 h_ready = true;
1935 chunk_skip_until = end;
1936 continue;
1937 }
1938 }
1939 if !h_ready {
1940 fill_h(&mut h, self);
1941 h_ready = true;
1942 }
1943 let lw = &self.weights.layers[li];
1944 match &lw.attn {
1946 AttnKind::LinearGdn(w) => {
1947 let cfg = self.gdn_cfg.expect("gdn layer without gdn_cfg");
1949 let mut normed = vec![0.0f32; b * hs];
1950 for bi in 0..b {
1951 let r = inference::rms_norm(
1952 &h[bi * hs..(bi + 1) * hs], &lw.input_norm, eps, norm_style);
1953 normed[bi * hs..(bi + 1) * hs].copy_from_slice(&r);
1954 }
1955 let attn = crate::linear_core::gdn_forward_batch(
1956 &normed, b, w, &cfg,
1957 &mut self.kv_cache.layers[li].linear_state,
1958 pool.as_deref(),
1959 );
1960 for (dst, &a) in h.iter_mut().zip(&attn) {
1961 *dst += a;
1962 }
1963 }
1964 AttnKind::Full {
1965 wq, wk, wv, wo, q_norm, k_norm, output_gate, bias,
1966 } => {
1967 let mut normed = vec![0.0f32; b * hs];
1971 for bi in 0..b {
1972 inference::rms_norm_into(
1973 &h[bi * hs..(bi + 1) * hs],
1974 &lw.input_norm,
1975 eps,
1976 norm_style,
1977 &mut normed[bi * hs..(bi + 1) * hs],
1978 );
1979 }
1980 let inv_freq_l = self.layer_inv_freq(li);
1981 let (nkv_l, hd_l, rd_l) = self.layer_geom(li);
1982 let cfg = QwenAttnCfg {
1983 num_heads: nh,
1984 num_kv_heads: nkv_l,
1985 head_dim: hd_l,
1986 hidden_size: hs,
1987 position: start_pos,
1988 inv_freq: &inv_freq_l,
1989 rotary_dim: rd_l,
1990 scale: self.attn_scale,
1991 window: self.layer_window(li),
1992 v_norm: self.attn_v_norm,
1993 q_norm: q_norm.as_deref(),
1994 k_norm: k_norm.as_deref(),
1995 output_gate: *output_gate,
1996 bias: bias.as_ref().map(|(a, b, c)| {
1997 (a.as_slice(), b.as_slice(), c.as_slice())
1998 }),
1999 rms_eps: eps,
2000 norm_style,
2001 pool: pool.as_deref(),
2002 };
2003 let mut attn = attention::qwen_attention_batch(
2004 &normed, b, wq, wk, wv, wo,
2005 &mut self.kv_cache.layers[li], &cfg);
2006 if let Some(w) = &lw.attn_out_norm {
2007 for bi in 0..b {
2008 inference::rms_norm_into(
2009 &attn[bi * hs..(bi + 1) * hs], w, eps, norm_style,
2010 &mut normed[bi * hs..(bi + 1) * hs]);
2011 }
2012 attn.copy_from_slice(&normed);
2013 }
2014 for (dst, &a) in h.iter_mut().zip(&attn) {
2015 *dst += a;
2016 }
2017 }
2018 AttnKind::Linear(w) => {
2019 for bi in 0..b {
2020 let normed = inference::rms_norm(
2021 &h[bi * hs..(bi + 1) * hs], &lw.input_norm, eps, norm_style);
2022 vmf_phase_forward(
2023 &normed, w,
2024 &self.vmf_cfg.expect("linear layer without vmf_cfg"),
2025 &mut self.kv_cache.layers[li].linear_state,
2026 pool.as_deref(),
2027 )
2028 .iter()
2029 .enumerate()
2030 .for_each(|(i, &a)| h[bi * hs + i] += a);
2031 }
2032 }
2033 }
2034
2035 let lw = &self.weights.layers[li];
2037 let mut post = vec![0.0f32; b * hs];
2038 for bi in 0..b {
2039 let r = inference::rms_norm(
2040 &h[bi * hs..(bi + 1) * hs], &lw.post_norm, eps, norm_style);
2041 post[bi * hs..(bi + 1) * hs].copy_from_slice(&r);
2042 }
2043 let mut ffn = match &lw.ffn {
2044 FfnKind::Dense(d) => dense_ffn_batch(d, &post, b, pool.as_deref()),
2045 FfnKind::Moe(m) => moe_ffn_batch(m, &post, b, hs, pool.as_deref()),
2046 };
2047 if let Some(w) = &lw.ffn_out_norm {
2048 for bi in 0..b {
2049 inference::rms_norm_into(
2050 &ffn[bi * hs..(bi + 1) * hs], w, eps, norm_style,
2051 &mut post[bi * hs..(bi + 1) * hs]);
2052 }
2053 ffn.copy_from_slice(&post);
2054 }
2055 for (dst, &f) in h.iter_mut().zip(&ffn) {
2056 *dst += f;
2057 }
2058 if let Some(sc) = lw.layer_scale {
2059 for v in h.iter_mut() {
2060 *v *= sc;
2061 }
2062 }
2063 if std::env::var("CMF_TRACE_H").is_ok() {
2064 let n = h[..hs].iter().map(|v| v.abs()).sum::<f32>() / hs as f32;
2065 let mx = h[..hs].iter().fold(0.0f32, |a, &v| a.max(v.abs()));
2066 eprintln!("layer {li}: mean|h|={n:.4} max|h|={mx:.2} scale={:?}", lw.layer_scale);
2067 }
2068 }
2069 crate::gpu::set_layer(-1); h
2071 }
2072
2073 fn embed_single(&self, id: u32) -> Vec<f32> {
2075 let mut out = vec![0.0f32; self.hidden_size];
2076 if (id as usize) < self.weights.embed_tokens.rows() {
2077 self.weights.embed_tokens.row_f32(id as usize, &mut out);
2078 }
2079 if self.embed_multiplier != 1.0 {
2080 for v in out.iter_mut() {
2081 *v *= self.embed_multiplier;
2082 }
2083 }
2084 out
2085 }
2086
2087 #[cfg(target_os = "macos")]
2093 fn chunk_run_gpu(
2094 &mut self,
2095 li0: usize,
2096 h: &mut [f32],
2097 b: usize,
2098 pos0: usize,
2099 embed_ids: Option<&[u32]>,
2100 ) -> usize {
2101 if !crate::gpu::enabled_here()
2105 || std::env::var("CMF_GPU_CHUNK").map(|v| v == "0").unwrap_or(false)
2106 || b < 32
2107 || self.swa.is_some()
2108 || self.global_attn.is_some()
2109 || self.attn_v_norm
2110 || (self.attn_scale - 1.0 / (self.head_dim as f32).sqrt()).abs() > 1e-9
2111 {
2112 return li0;
2113 }
2114 let Some(model) = self.model.clone() else { return li0 };
2115 let inv_freq = self.inv_freq.clone();
2116 let (nh, nkv, hd, hs) = (self.num_heads, self.num_kv_heads, self.head_dim, self.hidden_size);
2117 let mut layers: Vec<crate::gpu_metal::ChunkLayer> = Vec::new();
2119 let mut stored_at: Vec<usize> = Vec::new();
2120 for li in li0..self.num_layers {
2121 let lw = &self.weights.layers[li];
2122 if lw.attn_out_norm.is_some() || lw.ffn_out_norm.is_some() || lw.layer_scale.is_some()
2123 {
2124 break;
2125 }
2126 let AttnKind::Full { wq, wk, wv, wo, q_norm, k_norm, output_gate: false, bias } =
2127 &lw.attn
2128 else {
2129 break;
2130 };
2131 let FfnKind::Dense(d) = &lw.ffn else { break };
2132 if d.act != Act::Silu {
2133 break;
2134 }
2135 let parts = (
2136 wq.q8_row_parts(),
2137 wk.q8_row_parts(),
2138 wv.q8_row_parts(),
2139 wo.q8_row_parts(),
2140 d.gate_proj.q8_row_parts(),
2141 d.up_proj.q8_row_parts(),
2142 d.down_proj.q8_row_parts(),
2143 );
2144 let (Some(pq), Some(pk), Some(pv), Some(po), Some(pg), Some(pu), Some(pd)) = parts
2145 else {
2146 break;
2147 };
2148 let layer = &self.kv_cache.layers[li];
2149 if layer.mode != crate::kv_cache::KvMode::F32 || layer.o1.is_some() {
2150 break;
2151 }
2152 stored_at.push(layer.head_len(0));
2153 layers.push(crate::gpu_metal::ChunkLayer {
2154 model: &model,
2155 kv_id: self.graph_kv_id,
2156 layer: li,
2157 wq: pq,
2158 wk: pk,
2159 wv: pv,
2160 wo: po,
2161 gate: pg,
2162 up: pu,
2163 down: pd,
2164 input_norm: &lw.input_norm,
2165 post_norm: &lw.post_norm,
2166 bias: bias
2167 .as_ref()
2168 .map(|(a, bb, cc)| (a.as_slice(), bb.as_slice(), cc.as_slice())),
2169 q_norm: q_norm.as_deref(),
2170 k_norm: k_norm.as_deref(),
2171 inv_freq: &inv_freq,
2172 rd: self.rotary_dim,
2173 nh,
2174 nkv,
2175 hd,
2176 hs,
2177 inter: d.gate_proj.rows(),
2178 gemma: matches!(self.norm_style, cortiq_core::NormStyle::Gemma),
2179 eps: self.rms_eps as f32,
2180 });
2181 }
2182 if layers.is_empty() {
2183 return li0;
2184 }
2185 let row = nkv * hd;
2186 let mut store: Vec<(Vec<f32>, Vec<f32>, Vec<f32>)> = stored_at
2187 .iter()
2188 .map(|&st| (vec![0f32; b * row], vec![0f32; b * row], vec![0f32; st + b]))
2189 .collect();
2190 let mut io: Vec<crate::gpu_metal::ChunkIo> = Vec::with_capacity(layers.len());
2191 for (i, (ok, ov, oi)) in store.iter_mut().enumerate() {
2192 let li = layers[i].layer;
2193 let layer = &self.kv_cache.layers[li];
2194 io.push(crate::gpu_metal::ChunkIo {
2195 cpu_stored: stored_at[i],
2196 cpu_k: (0..nkv).map(|g| layer.head_keys(g)).collect(),
2197 cpu_v: (0..nkv).map(|g| layer.head_values(g)).collect(),
2198 out_k: ok,
2199 out_v: ov,
2200 imp: oi,
2201 });
2202 }
2203 let n_run = layers.len();
2204 let last = layers.last().map(|l| l.layer + 1).unwrap_or(li0);
2205 let ep = embed_ids.and_then(|ids| {
2208 self.weights.embed_tokens.q8_row_parts().map(|(idx, rows, _c, rs)| {
2209 crate::gpu_metal::ChunkEmbed {
2210 idx,
2211 rows,
2212 row_scale: rs,
2213 ids,
2214 mult: self.embed_multiplier,
2215 }
2216 })
2217 });
2218 if embed_ids.is_some() && ep.is_none() {
2219 return li0;
2220 }
2221 if !crate::gpu_metal::chunk_run_gpu(&layers, &mut io, h, b, pos0, ep.as_ref()) {
2222 return li0;
2223 }
2224 drop(io);
2225 drop(layers);
2226 for (i, (ok, ov, oi)) in store.iter().enumerate().take(n_run) {
2229 let li = li0 + i;
2230 let layer = &mut self.kv_cache.layers[li];
2231 for bi in 0..b {
2232 layer.append(&ok[bi * row..(bi + 1) * row], &ov[bi * row..(bi + 1) * row], &[]);
2233 }
2234 layer.accumulate_imp(oi);
2235 }
2236 last
2237 }
2238
2239 fn layer_is_local(&self, li: usize) -> bool {
2242 match self.swa {
2243 Some((_, pattern)) => (li + 1) % pattern.max(1) != 0,
2244 None => false,
2245 }
2246 }
2247
2248 fn layer_inv_freq(&self, li: usize) -> std::sync::Arc<Vec<f32>> {
2251 if self.layer_is_local(li) {
2252 if let Some(f) = &self.inv_freq_local {
2253 return f.clone();
2254 }
2255 } else if let Some(f) = &self.inv_freq_global {
2256 return f.clone();
2257 }
2258 self.inv_freq.clone()
2259 }
2260
2261 fn layer_window(&self, li: usize) -> Option<usize> {
2263 match self.swa {
2264 Some((w, _)) if self.layer_is_local(li) => Some(w),
2265 _ => None,
2266 }
2267 }
2268
2269 fn layer_geom(&self, li: usize) -> (usize, usize, usize) {
2272 if !self.layer_is_local(li) {
2273 if let Some((ghd, gkv)) = self.global_attn {
2274 return (gkv, ghd, ghd);
2275 }
2276 }
2277 (self.num_kv_heads, self.head_dim, self.rotary_dim)
2278 }
2279
2280 fn forward_layers(
2282 &mut self,
2283 hidden: &[f32],
2284 position: usize,
2285 task_mask: Option<&TaskMask>,
2286 ) -> Vec<f32> {
2287 self.forward_layers_upto(hidden, position, task_mask, None)
2288 }
2289
2290 fn forward_layers_upto(
2292 &mut self,
2293 hidden: &[f32],
2294 position: usize,
2295 task_mask: Option<&TaskMask>,
2296 upto: Option<usize>,
2297 ) -> Vec<f32> {
2298 let mut h = hidden.to_vec();
2299 let (nh, _nkv, _hd, hs, _rd, eps) = (
2302 self.num_heads,
2303 self.num_kv_heads,
2304 self.head_dim,
2305 self.hidden_size,
2306 self.rotary_dim,
2307 self.rms_eps,
2308 );
2309 let pool = self.pool.clone();
2310
2311 #[cfg(target_os = "macos")]
2312 let mut gpu_skip_until = 0usize;
2313 for li in 0..self.num_layers {
2314 crate::gpu::set_layer(li as i64); if let Some(u) = upto {
2316 if li > u {
2317 break;
2318 }
2319 }
2320 if let Some(mask) = task_mask {
2321 if !mask.layer_alive(li) {
2322 continue; }
2324 }
2325 #[cfg(target_os = "macos")]
2329 {
2330 if li < gpu_skip_until {
2331 continue;
2332 }
2333 if task_mask.is_none() {
2334 let end = self.q1_graph_gpu(li, upto, position, &mut h);
2335 if end > li {
2336 gpu_skip_until = end;
2337 continue;
2338 }
2339 }
2340 }
2341
2342 let lw = &self.weights.layers[li];
2343 inference::rms_norm_into(&h, &lw.input_norm, self.rms_eps, self.norm_style, &mut self.ws.n1);
2346
2347 let attn_out = match &lw.attn {
2348 AttnKind::Linear(w) => {
2349 let cfg = self.vmf_cfg.expect("linear layer without vmf_cfg");
2350 vmf_phase_forward(
2351 &self.ws.n1,
2352 w,
2353 &cfg,
2354 &mut self.kv_cache.layers[li].linear_state,
2355 self.pool.as_deref(),
2356 )
2357 }
2358 AttnKind::LinearGdn(w) => {
2359 let cfg = self.gdn_cfg.expect("gdn layer without gdn_cfg");
2360 gdn_forward(
2361 &self.ws.n1,
2362 w,
2363 &cfg,
2364 &mut self.kv_cache.layers[li].linear_state,
2365 self.pool.as_deref(),
2366 )
2367 }
2368 AttnKind::Full {
2369 wq,
2370 wk,
2371 wv,
2372 wo,
2373 q_norm,
2374 k_norm,
2375 output_gate,
2376 bias,
2377 } if self.kv_cache.layers[li].o1_sealed() => {
2378 let inv_freq_l = self.layer_inv_freq(li);
2381 let (nkv_l, hd_l, rd_l) = self.layer_geom(li);
2382 let cfg = QwenAttnCfg {
2383 num_heads: nh,
2384 num_kv_heads: nkv_l,
2385 head_dim: hd_l,
2386 hidden_size: hs,
2387 position,
2388 inv_freq: &inv_freq_l,
2389 rotary_dim: rd_l,
2390 scale: self.attn_scale,
2391 window: None,
2392 v_norm: self.attn_v_norm,
2393 q_norm: q_norm.as_deref(),
2394 k_norm: k_norm.as_deref(),
2395 output_gate: *output_gate,
2396 bias: bias
2397 .as_ref()
2398 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
2399 rms_eps: eps,
2400 norm_style: self.norm_style,
2401 pool: pool.as_deref(),
2402 };
2403 attention::qwen_attention_nystrom(
2404 &self.ws.n1,
2405 wq,
2406 wk,
2407 wv,
2408 wo,
2409 &mut self.kv_cache.layers[li],
2410 &cfg,
2411 )
2412 }
2413 AttnKind::Full {
2414 wq,
2415 wk,
2416 wv,
2417 wo,
2418 q_norm,
2419 k_norm,
2420 output_gate,
2421 bias,
2422 } => {
2423 let masked = task_mask
2424 .map(|m| m.head_flags(li, self.num_heads).iter().any(|&a| !a))
2425 .unwrap_or(false);
2426 let f32_view = (wq.as_f32(), wk.as_f32(), wv.as_f32(), wo.as_f32());
2427 match (masked, f32_view) {
2428 (true, (Some(q), Some(k), Some(v), Some(o))) => {
2431 let active_heads = task_mask.unwrap().head_flags(li, self.num_heads);
2432 attention::multi_head_attention(
2433 &self.ws.n1,
2434 q,
2435 k,
2436 v,
2437 o,
2438 &mut self.kv_cache.layers[li],
2439 self.num_heads,
2440 self.num_kv_heads,
2441 self.head_dim,
2442 self.hidden_size,
2443 position,
2444 &active_heads,
2445 &self.inv_freq,
2446 )
2447 }
2448 (masked, _) => {
2449 if masked {
2450 tracing::warn!(
2451 "layer {li}: head mask on quantized weights not \
2452 supported yet — executing dense"
2453 );
2454 }
2455 let inv_freq_l = self.layer_inv_freq(li);
2456 let (nkv_l, hd_l, rd_l) = self.layer_geom(li);
2457 let cfg = QwenAttnCfg {
2458 num_heads: nh,
2459 num_kv_heads: nkv_l,
2460 head_dim: hd_l,
2461 hidden_size: hs,
2462 position,
2463 inv_freq: &inv_freq_l,
2464 rotary_dim: rd_l,
2465 scale: self.attn_scale,
2466 window: self.layer_window(li),
2467 v_norm: self.attn_v_norm,
2468 q_norm: q_norm.as_deref(),
2469 k_norm: k_norm.as_deref(),
2470 output_gate: *output_gate,
2471 bias: bias
2472 .as_ref()
2473 .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
2474 rms_eps: eps,
2475 norm_style: self.norm_style,
2476 pool: pool.as_deref(),
2477 };
2478 attention::qwen_attention(
2479 &self.ws.n1,
2480 wq,
2481 wk,
2482 wv,
2483 wo,
2484 &mut self.kv_cache.layers[li],
2485 &cfg,
2486 )
2487 }
2488 }
2489 }
2490 };
2491 let attn_out = match &self.weights.layers[li].attn_out_norm {
2494 Some(w) => inference::rms_norm(&attn_out, w, self.rms_eps, self.norm_style),
2495 None => attn_out,
2496 };
2497 for (i, &a) in attn_out.iter().enumerate() {
2498 h[i] += a;
2499 }
2500 let mut attn_out = attn_out;
2501 attention::recycle_buf(&mut attn_out);
2502
2503 let lw = &self.weights.layers[li];
2504 inference::rms_norm_into(&h, &lw.post_norm, self.rms_eps, self.norm_style, &mut self.ws.p1);
2505 let post_normed = &self.ws.p1;
2506
2507 let ffn_masked = task_mask
2508 .map(|m| m.ffn_active_count(li) < self.intermediate_size)
2509 .unwrap_or(false);
2510 let f32_ffn = match &lw.ffn {
2513 FfnKind::Dense(d) => {
2514 (d.gate_proj.as_f32(), d.up_proj.as_f32(), d.down_proj.as_f32())
2515 }
2516 FfnKind::Moe(_) => (None, None, None),
2517 };
2518 let ffn_out = match (ffn_masked, f32_ffn) {
2519 (true, (Some(g), Some(u), Some(d))) => {
2520 let active = task_mask.unwrap().ffn_active_indices(li);
2521 inference::sparse_ffn_forward(
2522 &post_normed,
2523 g,
2524 u,
2525 d,
2526 self.hidden_size,
2527 self.intermediate_size,
2528 &active,
2529 self.pool.as_deref(),
2530 )
2531 }
2532 (true, _) => match &lw.ffn {
2536 FfnKind::Dense(d) if d.down_proj.sparse_col_ok() => {
2537 let active = task_mask.unwrap().ffn_active_indices(li);
2538 sparse_ffn_quant(
2539 d,
2540 &post_normed,
2541 &active,
2542 self.hidden_size,
2543 self.pool.as_deref(),
2544 )
2545 }
2546 FfnKind::Dense(d) => {
2551 let active = task_mask.unwrap().ffn_active_indices(li);
2552 let (gf, uf, df) = dequant_dense_f32(d);
2553 inference::sparse_ffn_forward(
2554 &post_normed,
2555 &gf,
2556 &uf,
2557 &df,
2558 self.hidden_size,
2559 self.intermediate_size,
2560 &active,
2561 self.pool.as_deref(),
2562 )
2563 }
2564 FfnKind::Moe(_) => {
2565 ffn_forward(&lw.ffn, &post_normed, self.pool.as_deref())
2568 }
2569 },
2570 (false, _) => ffn_forward(&lw.ffn, &post_normed, self.pool.as_deref()),
2571 };
2572 let ffn_out = match &self.weights.layers[li].ffn_out_norm {
2573 Some(w) => inference::rms_norm(&ffn_out, w, self.rms_eps, self.norm_style),
2574 None => ffn_out,
2575 };
2576 for (i, &f) in ffn_out.iter().enumerate() {
2577 h[i] += f;
2578 }
2579 let mut ffn_out = ffn_out;
2580 attention::recycle_buf(&mut ffn_out);
2581
2582 if let Some(sc) = self.weights.layers[li].layer_scale {
2584 for v in h.iter_mut() {
2585 *v *= sc;
2586 }
2587 }
2588
2589 if self.dyn_phi_layer == Some(li) {
2593 self.update_dyn_phi(&h);
2594 }
2595 }
2596 crate::gpu::set_layer(-1); h
2599 }
2600
2601 fn update_dyn_phi(&mut self, h: &[f32]) {
2604 const A: f32 = 0.2;
2605 if self.dyn_phi_ema.len() != h.len() {
2606 self.dyn_phi_ema = vec![0.0; h.len()];
2607 self.dyn_phi_seen = 0;
2608 }
2609 if self.dyn_phi_seen == 0 {
2610 self.dyn_phi_ema.copy_from_slice(h);
2611 } else {
2612 for (e, &v) in self.dyn_phi_ema.iter_mut().zip(h) {
2613 *e = (1.0 - A) * *e + A * v;
2614 }
2615 }
2616 self.dyn_phi_seen += 1;
2617 }
2618
2619 pub fn dyn_phi(&self) -> &[f32] {
2621 &self.dyn_phi_ema
2622 }
2623
2624 pub fn set_dyn_phi_layer(&mut self, layer: Option<usize>) {
2626 self.dyn_phi_layer = layer;
2627 self.dyn_phi_ema.clear();
2628 self.dyn_phi_seen = 0;
2629 }
2630
2631 pub fn dynamic_skills(&self) -> Vec<(usize, String, usize)> {
2633 let Some(model) = &self.model else { return Vec::new() };
2634 model
2635 .header
2636 .skills
2637 .iter()
2638 .enumerate()
2639 .filter_map(|(i, sk)| {
2640 let ok = matches!(self.dyn_skill_layers.get(i), Some(Some(_)));
2641 let sel = sk.selection.as_ref()?;
2642 (ok).then(|| (i, sk.id.clone(), sel.phi_layer))
2643 })
2644 .collect()
2645 }
2646
2647 pub fn active_skill(&self) -> Option<usize> {
2649 self.dyn_active
2650 }
2651
2652 pub fn enable_dynamic_routing(&mut self) -> usize {
2657 use crate::swarm::{DynRouter, RoutableSkill};
2658 let Some(model) = self.model.clone() else { return 0 };
2659 if self.dyn_blend_loaded {
2662 tracing::warn!("dynamic routing unavailable on a blend-loaded pipeline");
2663 return 0;
2664 }
2665 if let Some(a) = self.dyn_active {
2669 if !matches!(self.dyn_skill_layers.get(a), Some(Some(_))) {
2670 tracing::warn!(
2671 "loaded skill is not FFN-eligible — dynamic routing unavailable"
2672 );
2673 return 0;
2674 }
2675 }
2676 let hidden = self.hidden_size;
2677 let mut skills = Vec::new();
2678 for (idx, id, _phi) in self.dynamic_skills() {
2679 if let Some(sel) = model.header.skills[idx].selection.as_ref() {
2680 if let Some(rs) = RoutableSkill::from_descriptor(idx, id, sel, hidden) {
2681 skills.push(rs);
2682 }
2683 }
2684 }
2685 if skills.is_empty() {
2686 return 0;
2687 }
2688 let phi = skills[0].phi_layer;
2690 if skills.iter().any(|s| s.phi_layer != phi) {
2691 tracing::warn!("routable skills disagree on phi_layer; using {phi}");
2692 }
2693 let n = skills.len();
2694 self.set_dyn_phi_layer(Some(phi));
2695 self.dyn_router = Some(DynRouter::new(skills));
2696 n
2697 }
2698
2699 pub fn route_switches(&self) -> Vec<(usize, Option<String>, Option<String>)> {
2701 self.dyn_router
2702 .as_ref()
2703 .map(|r| r.switches.clone())
2704 .unwrap_or_default()
2705 }
2706
2707 fn lm_head_forward(&self, hidden: &[f32]) -> Vec<f32> {
2710 let rows = self.weights.lm_head.rows();
2711 let mut logits = attention::take_buf(rows.min(self.vocab_size));
2712 self.weights
2713 .lm_head
2714 .matvec(hidden, &mut logits, self.pool.as_deref());
2715 logits.resize(self.vocab_size, 0.0);
2716 if let Some(c) = self.final_softcap {
2717 for l in logits.iter_mut() {
2718 *l = c * (*l / c).tanh();
2719 }
2720 }
2721 logits
2722 }
2723
2724 pub fn prefill_next_logits(&mut self, ids: &[u32], task_mask: Option<&TaskMask>) -> Vec<f32> {
2729 self.kv_cache.clear();
2730 let mut hidden = vec![0.0f32; self.hidden_size];
2731 for (pos, &id) in ids.iter().enumerate() {
2732 let emb = self.embed_single(id);
2733 hidden = self.forward_layers(&emb, pos, task_mask);
2734 }
2735 inference::rms_norm_into(
2736 &hidden,
2737 &self.weights.final_norm,
2738 self.rms_eps,
2739 self.norm_style,
2740 &mut self.ws.n1,
2741 );
2742 self.lm_head_forward(&self.ws.n1)
2743 }
2744}
2745
2746pub fn create_test_pipeline(
2748 hidden_size: usize,
2749 intermediate_size: usize,
2750 num_heads: usize,
2751 num_kv_heads: usize,
2752 head_dim: usize,
2753 num_layers: usize,
2754 vocab_size: usize,
2755) -> Pipeline {
2756 let synth = |n: usize, salt: usize| -> Vec<f32> {
2759 (0..n)
2760 .map(|i| (((i * 31 + salt * 17 + 7) % 97) as f32 / 97.0 - 0.5) * 0.2)
2761 .collect()
2762 };
2763 let qt = |rows: usize, cols: usize, salt: usize| -> QTensor {
2764 QTensor::from_f32(synth(rows * cols, salt), rows, cols)
2765 };
2766 let layer_weights: Vec<LayerWeights> = (0..num_layers)
2767 .map(|li| LayerWeights {
2768 input_norm: vec![1.0; hidden_size],
2769 post_norm: vec![1.0; hidden_size],
2770 attn_out_norm: None,
2771 ffn_out_norm: None,
2772 layer_scale: None,
2773 ffn: FfnKind::Dense(DenseFfn {
2774 gate_proj: qt(intermediate_size, hidden_size, li * 10 + 5),
2775 up_proj: qt(intermediate_size, hidden_size, li * 10 + 6),
2776 down_proj: qt(hidden_size, intermediate_size, li * 10 + 7),
2777 act: Act::Silu,
2778 }),
2779 attn: AttnKind::Full {
2780 bias: None,
2781 wq: qt(num_heads * head_dim, hidden_size, li * 10 + 1),
2782 wk: qt(num_kv_heads * head_dim, hidden_size, li * 10 + 2),
2783 wv: qt(num_kv_heads * head_dim, hidden_size, li * 10 + 3),
2784 wo: qt(hidden_size, num_heads * head_dim, li * 10 + 4),
2785 q_norm: None,
2786 k_norm: None,
2787 output_gate: false,
2788 },
2789 })
2790 .collect();
2791
2792 Pipeline::new(
2793 Tokenizer::byte_level(),
2794 PipelineWeights {
2795 embed_tokens: qt(vocab_size, hidden_size, 100),
2796 layers: layer_weights,
2797 lm_head: qt(vocab_size, hidden_size, 200),
2798 final_norm: vec![1.0; hidden_size],
2799 },
2800 hidden_size,
2801 intermediate_size,
2802 num_heads,
2803 num_kv_heads,
2804 head_dim,
2805 num_layers,
2806 vocab_size,
2807 1e-6,
2808 10_000.0,
2809 NormStyle::Qwen,
2810 4096,
2811 SamplerConfig {
2812 seed: Some(42),
2813 ..Default::default()
2814 },
2815 )
2816}
2817
2818fn dense_ffn_batch(d: &DenseFfn, xs: &[f32], b: usize, pool: Option<&Pool>) -> Vec<f32> {
2821 let inter = d.gate_proj.rows();
2822 let hidden = d.down_proj.rows();
2823 let mut g = vec![0.0f32; b * inter];
2824 d.gate_proj.matmat(xs, b, &mut g, pool);
2825 let mut u = vec![0.0f32; b * inter];
2826 d.up_proj.matmat(xs, b, &mut u, pool);
2827 for i in 0..b * inter {
2828 g[i] = d.act.apply(g[i]) * u[i];
2829 }
2830 let mut out = vec![0.0f32; b * hidden];
2831 d.down_proj.matmat(&g, b, &mut out, pool);
2832 out
2833}
2834
2835fn moe_ffn_batch(m: &MoeFfn, xs: &[f32], b: usize, hidden: usize, pool: Option<&Pool>) -> Vec<f32> {
2839 let ne = m.experts.len();
2840 let mut logits = vec![0.0f32; b * ne];
2841 m.router.matmat(xs, b, &mut logits, pool);
2842
2843 let mut assign: Vec<Vec<(usize, f32)>> = vec![Vec::new(); ne];
2846 {
2847 let mut st = m.stats.borrow_mut();
2848 if st.len() < ne {
2849 st.resize(ne, 0);
2850 }
2851 for bi in 0..b {
2852 let lg = &logits[bi * ne..(bi + 1) * ne];
2853 let mx = lg.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
2854 let mut p: Vec<f32> = lg.iter().map(|&l| (l - mx).exp()).collect();
2855 let sum: f32 = p.iter().sum();
2856 for v in &mut p {
2857 *v /= sum;
2858 }
2859 let mut order: Vec<usize> = (0..ne).collect();
2860 order.sort_unstable_by(|&x, &y| p[y].partial_cmp(&p[x]).unwrap().then(x.cmp(&y)));
2861 order.truncate(m.top_k);
2862 let wsum: f32 = if m.norm_topk_prob {
2863 order.iter().map(|&e| p[e]).sum()
2864 } else {
2865 1.0
2866 };
2867 for &e in &order {
2868 st[e] += 1;
2869 assign[e].push((bi, p[e] / wsum));
2870 }
2871 }
2872 }
2873
2874 let mut out = vec![0.0f32; b * hidden];
2875 let cols = m.experts[0].gate_proj.cols();
2876 let mut run_expert = |d: &DenseFfn, list: &[(usize, f32)]| {
2877 let sb = list.len();
2878 let mut sub = vec![0.0f32; sb * cols];
2879 for (k, &(bi, _)) in list.iter().enumerate() {
2880 sub[k * cols..(k + 1) * cols].copy_from_slice(&xs[bi * cols..(bi + 1) * cols]);
2881 }
2882 let eo = dense_ffn_batch(d, &sub, sb, pool);
2883 for (k, &(bi, w)) in list.iter().enumerate() {
2884 for i in 0..hidden {
2885 out[bi * hidden + i] += w * eo[k * hidden + i];
2886 }
2887 }
2888 };
2889 for e in 0..ne {
2890 if !assign[e].is_empty() {
2891 run_expert(&m.experts[e], &assign[e]);
2892 }
2893 }
2894 if let Some((se, gate)) = &m.shared {
2895 let mut gl = vec![0.0f32; b];
2896 gate.matmat(xs, b, &mut gl, pool);
2897 let all: Vec<(usize, f32)> = (0..b)
2898 .map(|bi| (bi, 1.0 / (1.0 + (-gl[bi]).exp())))
2899 .collect();
2900 run_expert(se, &all);
2901 }
2902 out
2903}
2904
2905thread_local! {
2906 static FFN_SCRATCH: std::cell::RefCell<[Vec<f32>; 4]> =
2910 const { std::cell::RefCell::new([Vec::new(), Vec::new(), Vec::new(), Vec::new()]) };
2911}
2912
2913fn dense_ffn(d: &DenseFfn, x: &[f32], pool: Option<&Pool>) -> Vec<f32> {
2915 if crate::gpu::enabled_here()
2926 && (d.gate_proj.rows() >= crate::gpu::min_rows() || d.gate_proj.is_q1())
2927 {
2928 let arm = if d.gate_proj.is_q1() && crate::gpu::q1_force() {
2929 crate::gpu::ProbeArm::Gpu
2930 } else {
2931 crate::gpu::probe_arm(crate::gpu::OpClass::Ffn)
2932 };
2933 match arm {
2934 crate::gpu::ProbeArm::Gpu => {
2935 let t0 = std::time::Instant::now();
2936 if let Some(out) = dense_ffn_gpu(d, x, pool) {
2937 crate::gpu::probe_record(crate::gpu::OpClass::Ffn, true, t0.elapsed());
2938 return out;
2939 }
2940 }
2941 crate::gpu::ProbeArm::CpuTimed => {
2942 let t0 = std::time::Instant::now();
2943 let out = crate::gpu::cpu_scope(|| dense_ffn_cpu(d, x, pool));
2944 crate::gpu::probe_record(crate::gpu::OpClass::Ffn, false, t0.elapsed());
2945 return out;
2946 }
2947 crate::gpu::ProbeArm::Cpu => {
2948 return crate::gpu::cpu_scope(|| dense_ffn_cpu(d, x, pool));
2949 }
2950 }
2951 }
2952 dense_ffn_cpu(d, x, pool)
2953}
2954
2955fn dense_ffn_cpu(d: &DenseFfn, x: &[f32], pool: Option<&Pool>) -> Vec<f32> {
2957 let inter = d.gate_proj.rows();
2958 FFN_SCRATCH.with(|s| {
2959 let mut s = s.borrow_mut();
2960 let [g, u, ..] = &mut *s;
2961 g.resize(inter, 0.0);
2962 u.resize(inter, 0.0);
2963 QTensor::matvec_many([&d.gate_proj, &d.up_proj], x, [g, u], pool);
2965 for i in 0..inter {
2966 g[i] = d.act.apply(g[i]) * u[i];
2967 }
2968 FFN_PROBE.with(|pr| {
2971 if let Some(acc) = pr.borrow_mut().as_mut() {
2972 let li = crate::gpu::cur_layer();
2973 if li >= 0 {
2974 if let Some(row) = acc.get_mut(li as usize) {
2975 for (a, &v) in row.iter_mut().zip(g.iter()) {
2976 *a += (v as f64).abs();
2977 }
2978 }
2979 }
2980 }
2981 });
2982 let mut out = attention::take_buf(d.down_proj.rows());
2983 d.down_proj.matvec(g, &mut out, pool);
2984 out
2985 })
2986}
2987
2988thread_local! {
2989 static FFN_PROBE: std::cell::RefCell<Option<Vec<Vec<f64>>>> =
2992 const { std::cell::RefCell::new(None) };
2993}
2994
2995fn dense_ffn_gpu(d: &DenseFfn, x: &[f32], _pool: Option<&Pool>) -> Option<Vec<f32>> {
3001 if d.act != Act::Silu {
3003 return None;
3004 }
3005 if d.gate_proj.rows() < crate::gpu::min_rows() && !d.gate_proj.is_q1() {
3008 return None;
3009 }
3010 let mut jobs: Vec<crate::gpu::MoeJob> = Vec::with_capacity(1);
3011 let mut model_ref = None;
3012 moe_push_job(d, x, 1.0, &mut jobs, &mut model_ref)?;
3013 let model = model_ref?;
3014 let hidden = jobs[0].down.1;
3015 let mut out = attention::take_buf(hidden);
3016 if crate::gpu::moe_block(&model, &jobs, &mut out) {
3017 Some(out)
3018 } else {
3019 let mut out = out;
3020 attention::recycle_buf(&mut out);
3021 None
3022 }
3023}
3024
3025#[allow(clippy::type_complexity)]
3030#[allow(clippy::type_complexity)]
3031fn moe_parts(
3032 t: &QTensor,
3033) -> Option<(&std::sync::Arc<cortiq_core::CmfModel>, usize, usize, usize, &[f32], &[f32], bool)> {
3034 match t {
3035 QTensor::Mapped {
3036 model,
3037 idx,
3038 dtype: dt @ (cortiq_core::TensorDtype::Q8_2f | cortiq_core::TensorDtype::Q8Row),
3039 rows,
3040 cols,
3041 row_scale,
3042 col_field,
3043 ..
3044 } if (*dt == cortiq_core::TensorDtype::Q8Row) || !col_field.is_empty() => {
3045 Some((model, *idx, *rows, *cols, row_scale, col_field, false))
3046 }
3047 QTensor::Mapped {
3049 model,
3050 idx,
3051 dtype: cortiq_core::TensorDtype::Q1,
3052 rows,
3053 cols,
3054 ..
3055 } => Some((model, *idx, *rows, *cols, &[][..], &[][..], true)),
3056 _ => None,
3057 }
3058}
3059
3060fn moe_push_job<'a>(
3062 d: &'a DenseFfn,
3063 x: &[f32],
3064 w: f32,
3065 jobs: &mut Vec<crate::gpu::MoeJob<'a>>,
3066 model_ref: &mut Option<std::sync::Arc<cortiq_core::CmfModel>>,
3067) -> Option<()> {
3068 use crate::qtensor::prescale;
3069 if d.act != Act::Silu {
3070 return None; }
3072 let (gm, gi, gr, gc, grs, gcf, gq1) = moe_parts(&d.gate_proj)?;
3073 let (_, ui, ur, uc, urs, ucf, uq1) = moe_parts(&d.up_proj)?;
3074 let (_, di, dr, dc, drs, dcf, dq1) = moe_parts(&d.down_proj)?;
3075 if gq1 != uq1 || uq1 != dq1 {
3076 return None; }
3078 model_ref.get_or_insert_with(|| gm.clone());
3079 let gdt = if gcf.is_empty() { cortiq_core::TensorDtype::Q8Row } else { cortiq_core::TensorDtype::Q8_2f };
3080 let udt = if ucf.is_empty() { cortiq_core::TensorDtype::Q8Row } else { cortiq_core::TensorDtype::Q8_2f };
3081 jobs.push(crate::gpu::MoeJob {
3082 gate: (gi, gr, gc, grs),
3083 up: (ui, ur, uc, urs),
3084 down: (di, dr, dc, drs),
3085 xs_gate: prescale(x, gcf, gdt).into_owned(),
3086 xs_up: prescale(x, ucf, udt).into_owned(),
3087 down_col: dcf,
3088 w,
3089 q1: gq1,
3090 });
3091 Some(())
3092}
3093
3094fn sparse_ffn_quant(
3101 d: &DenseFfn,
3102 x: &[f32],
3103 active: &[u16],
3104 hidden: usize,
3105 pool: Option<&Pool>,
3106) -> Vec<f32> {
3107 let n = active.len();
3108 let inter = d.gate_proj.rows();
3109 let mut act = vec![0.0f32; n];
3110 let need_scratch = !(d.gate_proj.sparse_col_ok() && d.up_proj.sparse_col_ok());
3113 let compute = |ai: usize| -> f32 {
3114 let idx = active[ai] as usize;
3115 if idx >= inter {
3116 return 0.0; }
3118 let mut s = if need_scratch { vec![0.0f32; hidden] } else { Vec::new() };
3119 let gate = d.gate_proj.row_dot(idx, x, &mut s);
3120 let up = d.up_proj.row_dot(idx, x, &mut s);
3121 d.act.apply(gate) * up
3122 };
3123 match pool {
3124 Some(p) if n >= 256 => {
3125 let ptr = SendMut(act.as_mut_ptr());
3126 p.run(&|widx, nw| {
3127 let chunk = n.div_ceil(nw);
3128 let (s, e) = (widx * chunk, ((widx + 1) * chunk).min(n));
3129 for ai in s..e {
3130 unsafe { *ptr.at(ai) = compute(ai) };
3131 }
3132 });
3133 }
3134 _ => {
3135 for (ai, a) in act.iter_mut().enumerate() {
3136 *a = compute(ai);
3137 }
3138 }
3139 }
3140 let mut out = vec![0.0f32; hidden];
3142 for (ai, &idx) in active.iter().enumerate() {
3143 let w = act[ai];
3144 if w.abs() >= 1e-12 && (idx as usize) < inter {
3145 d.down_proj.add_col_scaled(idx as usize, w, &mut out);
3146 }
3147 }
3148 out
3149}
3150
3151#[doc(hidden)]
3153pub fn sparse_ffn_quant_for_test(
3154 d: &DenseFfn,
3155 x: &[f32],
3156 active: &[u16],
3157 hidden: usize,
3158) -> Vec<f32> {
3159 sparse_ffn_quant(d, x, active, hidden, None)
3160}
3161
3162fn dequant_dense_f32(d: &DenseFfn) -> (Vec<f32>, Vec<f32>, Vec<f32>) {
3166 let deq = |t: &QTensor| -> Vec<f32> {
3167 let (rows, cols) = (t.rows(), t.cols());
3168 let mut out = vec![0.0f32; rows * cols];
3169 for r in 0..rows {
3170 t.row_f32(r, &mut out[r * cols..(r + 1) * cols]);
3171 }
3172 out
3173 };
3174 (deq(&d.gate_proj), deq(&d.up_proj), deq(&d.down_proj))
3175}
3176
3177struct SendMut(*mut f32);
3179unsafe impl Send for SendMut {}
3180unsafe impl Sync for SendMut {}
3181impl SendMut {
3182 #[inline]
3183 #[allow(clippy::mut_from_ref)]
3186 unsafe fn at(&self, i: usize) -> &mut f32 {
3187 unsafe { &mut *self.0.add(i) }
3188 }
3189}
3190
3191fn moe_ffn(m: &MoeFfn, x: &[f32], pool: Option<&Pool>) -> Vec<f32> {
3195 let ne = m.experts.len();
3196 let mut logits = vec![0.0f32; ne];
3197 m.router.matvec(x, &mut logits, pool);
3198 let mx = logits.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
3199 let mut p: Vec<f32> = logits.iter().map(|&l| (l - mx).exp()).collect();
3200 let s: f32 = p.iter().sum();
3201 for v in &mut p {
3202 *v /= s;
3203 }
3204 let mut idx: Vec<usize> = (0..ne).collect();
3205 idx.sort_unstable_by(|&a, &b| p[b].partial_cmp(&p[a]).unwrap().then(a.cmp(&b)));
3207 idx.truncate(m.top_k);
3208 let wsum: f32 = if m.norm_topk_prob {
3209 idx.iter().map(|&e| p[e]).sum()
3210 } else {
3211 1.0
3212 };
3213 {
3214 let mut st = m.stats.borrow_mut();
3215 if st.len() < ne {
3216 st.resize(ne, 0);
3217 }
3218 for &e in &idx {
3219 st[e] += 1;
3220 }
3221 }
3222 if crate::gpu::enabled_here() {
3227 match crate::gpu::probe_arm(crate::gpu::OpClass::Ffn) {
3228 crate::gpu::ProbeArm::Gpu => {
3229 let t0 = std::time::Instant::now();
3230 if let Some(out) = moe_ffn_gpu(m, x, &idx, &p, wsum, pool) {
3231 crate::gpu::probe_record(crate::gpu::OpClass::Ffn, true, t0.elapsed());
3232 return out;
3233 }
3234 }
3235 crate::gpu::ProbeArm::CpuTimed => {
3236 let t0 = std::time::Instant::now();
3237 let out = crate::gpu::cpu_scope(|| moe_ffn_cpu(m, x, &idx, &p, wsum, pool));
3238 crate::gpu::probe_record(crate::gpu::OpClass::Ffn, false, t0.elapsed());
3239 return out;
3240 }
3241 crate::gpu::ProbeArm::Cpu => {
3242 return crate::gpu::cpu_scope(|| moe_ffn_cpu(m, x, &idx, &p, wsum, pool));
3243 }
3244 }
3245 }
3246 moe_ffn_cpu(m, x, &idx, &p, wsum, pool)
3247}
3248
3249fn moe_ffn_cpu(
3251 m: &MoeFfn,
3252 x: &[f32],
3253 idx: &[usize],
3254 p: &[f32],
3255 wsum: f32,
3256 pool: Option<&Pool>,
3257) -> Vec<f32> {
3258 let mut out = attention::take_buf(x.len());
3259 for &e in idx {
3260 let mut eo = dense_ffn(&m.experts[e], x, pool);
3261 let w = p[e] / wsum;
3262 for i in 0..out.len() {
3263 out[i] += w * eo[i];
3264 }
3265 attention::recycle_buf(&mut eo);
3266 }
3267 if let Some((se, gate)) = &m.shared {
3268 let mut so = dense_ffn(se, x, pool);
3269 let mut gl = vec![0.0f32; 1];
3270 gate.matvec(x, &mut gl, pool);
3271 let g = 1.0 / (1.0 + (-gl[0]).exp());
3272 for i in 0..out.len() {
3273 out[i] += g * so[i];
3274 }
3275 attention::recycle_buf(&mut so);
3276 }
3277 out
3278}
3279
3280fn moe_ffn_gpu(
3283 m: &MoeFfn,
3284 x: &[f32],
3285 idx: &[usize],
3286 p: &[f32],
3287 wsum: f32,
3288 pool: Option<&Pool>,
3289) -> Option<Vec<f32>> {
3290 use crate::gpu::MoeJob;
3291
3292 let mut jobs: Vec<MoeJob> = Vec::with_capacity(idx.len() + 1);
3293 let mut model_ref = None;
3294 for &e in idx {
3295 moe_push_job(&m.experts[e], x, p[e] / wsum, &mut jobs, &mut model_ref)?;
3296 }
3297 if let Some((se, gate)) = &m.shared {
3298 let mut gl = vec![0.0f32; 1];
3299 gate.matvec(x, &mut gl, pool);
3300 let g = 1.0 / (1.0 + (-gl[0]).exp());
3301 moe_push_job(se, x, g, &mut jobs, &mut model_ref)?;
3302 }
3303 let model = model_ref?;
3304 let hidden = jobs[0].down.1;
3305 let mut out = vec![0.0f32; hidden];
3306 crate::gpu::moe_block(&model, &jobs, &mut out).then_some(out)
3307}
3308
3309fn ffn_forward(ffn: &FfnKind, x: &[f32], pool: Option<&Pool>) -> Vec<f32> {
3311 match ffn {
3312 FfnKind::Dense(d) => dense_ffn(d, x, pool),
3313 FfnKind::Moe(m) => moe_ffn(m, x, pool),
3314 }
3315}
3316
3317fn ffn_forward_pair(
3321 ffn: &FfnKind,
3322 x1: &[f32],
3323 x2: &[f32],
3324 pool: Option<&Pool>,
3325) -> (Vec<f32>, Vec<f32>) {
3326 let d = match ffn {
3327 FfnKind::Dense(d) => d,
3328 FfnKind::Moe(m) => return (moe_ffn(m, x1, pool), moe_ffn(m, x2, pool)),
3329 };
3330 let inter = d.gate_proj.rows();
3331 FFN_SCRATCH.with(|s| {
3332 let mut s = s.borrow_mut();
3333 let [g1, g2, u1, u2] = &mut *s;
3334 g1.resize(inter, 0.0);
3335 g2.resize(inter, 0.0);
3336 u1.resize(inter, 0.0);
3337 u2.resize(inter, 0.0);
3338 QTensor::matvec2_many(
3341 [&d.gate_proj, &d.up_proj],
3342 x1,
3343 x2,
3344 [g1.as_mut_slice(), u1.as_mut_slice()],
3345 [g2.as_mut_slice(), u2.as_mut_slice()],
3346 pool,
3347 );
3348 for i in 0..inter {
3349 g1[i] = d.act.apply(g1[i]) * u1[i];
3350 g2[i] = d.act.apply(g2[i]) * u2[i];
3351 }
3352 let mut o1 = attention::take_buf(d.down_proj.rows());
3353 let mut o2 = attention::take_buf(d.down_proj.rows());
3354 d.down_proj.matvec2(g1, g2, &mut o1, &mut o2, pool);
3355 (o1, o2)
3356 })
3357}
3358
3359#[cfg(test)]
3360mod tests {
3361 use super::*;
3362
3363 #[test]
3369 fn sparse_ffn_quant_equals_dense_with_inactive_zeroed() {
3370 let (hidden, inter) = (16usize, 40usize);
3371 let synth = |n: usize, salt: usize| -> Vec<f32> {
3372 (0..n)
3373 .map(|i| (((i * 37 + salt * 11 + 3) % 101) as f32 / 101.0 - 0.5) * 0.4)
3374 .collect()
3375 };
3376 let d = DenseFfn {
3377 gate_proj: QTensor::from_f32(synth(inter * hidden, 1), inter, hidden),
3378 up_proj: QTensor::from_f32(synth(inter * hidden, 2), inter, hidden),
3379 down_proj: QTensor::from_f32(synth(hidden * inter, 3), hidden, inter),
3380 act: Act::Silu,
3381 };
3382 let x = synth(hidden, 9);
3383 let active: Vec<u16> = (0..inter as u16).filter(|i| i % 3 == 0).collect();
3385
3386 let sparse = sparse_ffn_quant(&d, &x, &active, hidden, None);
3387
3388 let mut g = vec![0.0f32; inter];
3390 d.gate_proj.matvec(&x, &mut g, None);
3391 let mut u = vec![0.0f32; inter];
3392 d.up_proj.matvec(&x, &mut u, None);
3393 let act_set: std::collections::HashSet<u16> = active.iter().copied().collect();
3394 for i in 0..inter {
3395 g[i] = if act_set.contains(&(i as u16)) {
3396 inference::silu(g[i]) * u[i]
3397 } else {
3398 0.0
3399 };
3400 }
3401 let mut reference = vec![0.0f32; hidden];
3402 d.down_proj.matvec(&g, &mut reference, None);
3403
3404 let max_d = sparse
3405 .iter()
3406 .zip(&reference)
3407 .map(|(a, b)| (a - b).abs())
3408 .fold(0.0f32, f32::max);
3409 assert!(max_d < 1e-5, "sparse != dense-zeroed: max|Δ| = {max_d}");
3410 }
3411
3412 fn attach_test_mtp(p: &mut Pipeline) {
3414 let (h, inter, heads, kv, hd) = (
3415 p.hidden_size,
3416 p.intermediate_size,
3417 p.num_heads,
3418 p.num_kv_heads,
3419 p.head_dim,
3420 );
3421 let synth = |n: usize, salt: usize| -> Vec<f32> {
3422 (0..n)
3423 .map(|i| (((i * 29 + salt * 23 + 5) % 101) as f32 / 101.0 - 0.5) * 0.2)
3424 .collect()
3425 };
3426 let qt = |rows: usize, cols: usize, salt: usize| -> QTensor {
3427 QTensor::from_f32(synth(rows * cols, salt), rows, cols)
3428 };
3429 p.mtp = Some(MtpModule {
3430 enorm: vec![1.0; h],
3431 hnorm: vec![1.0; h],
3432 eh_proj: qt(h, 2 * h, 301),
3433 layer: LayerWeights {
3434 input_norm: vec![1.0; h],
3435 post_norm: vec![1.0; h],
3436 attn_out_norm: None,
3437 ffn_out_norm: None,
3438 layer_scale: None,
3439 ffn: FfnKind::Dense(DenseFfn {
3440 gate_proj: qt(inter, h, 315),
3441 up_proj: qt(inter, h, 316),
3442 down_proj: qt(h, inter, 317),
3443 act: Act::Silu,
3444 }),
3445 attn: AttnKind::Full {
3446 bias: None,
3447 wq: qt(heads * hd, h, 311),
3448 wk: qt(kv * hd, h, 312),
3449 wv: qt(kv * hd, h, 313),
3450 wo: qt(h, heads * hd, 314),
3451 q_norm: None,
3452 k_norm: None,
3453 output_gate: false,
3454 },
3455 },
3456 final_norm: vec![1.0; h],
3457 kv: crate::kv_cache::LayerKvCache::new(kv, hd),
3458 });
3459 }
3460
3461 #[test]
3462 fn speculative_equals_vanilla_greedy() {
3463 let run = |spec: bool| {
3464 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
3465 p.sampler_config.temperature = 0.0;
3466 attach_test_mtp(&mut p);
3467 p.speculative = spec;
3468 let r = p.generate("abcdef", 12, None, None).unwrap();
3469 (r.token_ids, r.mtp_drafted, r.mtp_accepted)
3470 };
3471 let (vanilla, d0, _) = run(false);
3472 let (spec, d1, a1) = run(true);
3473 assert_eq!(d0, 0, "vanilla path must not draft");
3474 assert!(d1 > 0, "speculative path must draft");
3475 assert_eq!(
3476 vanilla, spec,
3477 "speculative must reproduce the exact greedy sequence (accepted {a1}/{d1})"
3478 );
3479 }
3480
3481 #[test]
3482 fn speculative_accepts_constant_oracle() {
3483 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
3484 p.sampler_config.temperature = 0.0;
3485 p.sampler_config.repetition_penalty = 1.0;
3486 p.weights.lm_head = QTensor::from_f32(vec![0.01; 64 * 8], 64, 8);
3489 attach_test_mtp(&mut p);
3490 p.speculative = true;
3491 let r = p.generate("abcd", 10, None, None).unwrap();
3492 assert!(r.mtp_drafted > 0);
3493 assert_eq!(
3494 r.mtp_accepted, r.mtp_drafted,
3495 "constant logits → every draft accepted"
3496 );
3497 assert!(r.token_ids.windows(2).all(|w| w[0] == w[1]));
3500 }
3501
3502 #[test]
3503 fn empty_prompt_is_an_error_not_a_panic() {
3504 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 260);
3505 let r = p.generate("", 4, None, None);
3506 assert!(r.is_err(), "empty prompt must be a clean error");
3507 }
3508
3509 #[test]
3510 fn every_token_enters_kv_exactly_once() {
3511 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
3512 p.sampler_config.temperature = 0.0;
3514 let r = p.generate("abc", 2, None, None).unwrap();
3515 assert_eq!(r.prompt_tokens, 3);
3516 assert_eq!(
3520 p.kv_cache.seq_len(),
3521 3 + r.tokens_generated - 1,
3522 "each token must be cached exactly once (v1 cached the last prompt token twice)"
3523 );
3524 }
3525
3526 #[test]
3527 fn generation_is_reproducible_with_seed() {
3528 let run = || {
3529 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
3530 p.generate("hello", 8, None, None).unwrap().token_ids
3531 };
3532 assert_eq!(run(), run());
3533 }
3534
3535 #[test]
3536 fn eviction_bounds_the_cache() {
3537 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 260);
3538 p.kv_cache.max_seq_len = 6;
3539 p.sampler_config.temperature = 0.0;
3540 let _ = p.generate("abcd", 12, None, None).unwrap();
3541 assert!(
3542 p.kv_cache.seq_len() <= 6 + 1,
3543 "cache must stay bounded by max_seq_len (got {})",
3544 p.kv_cache.seq_len()
3545 );
3546 }
3547
3548 #[test]
3549 fn confidence_matches_tokens_and_is_a_probability() {
3550 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
3551 p.sampler_config.temperature = 0.0;
3552 p.sampler_config.repetition_penalty = 1.0;
3553 let r = p.generate("abcd", 10, None, None).unwrap();
3554 assert_eq!(
3555 r.token_confidence.len(),
3556 r.token_ids.len(),
3557 "one confidence per emitted token"
3558 );
3559 for &c in &r.token_confidence {
3560 assert!((0.0..=1.0).contains(&c), "confidence out of [0,1]: {c}");
3561 }
3562 let logits = [1.0f32, 3.0, 0.5, 3.0];
3564 let p0 = top1_prob_t(&logits, 1, 1.0);
3565 let p1 = top1_prob_t(&logits, 3, 1.0);
3566 assert!((p0 - p1).abs() < 1e-6, "equal logits → equal prob");
3567 assert!(p0 > 0.0 && p0 < 1.0);
3568 let sharp = top1_prob_t(&logits, 1, 1.0);
3570 let soft = top1_prob_t(&logits, 1, 2.0);
3571 assert!(soft < sharp, "higher temperature lowers peak confidence");
3572 }
3573
3574 #[test]
3575 fn trace_is_opt_in_and_parallels_the_output() {
3576 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
3578 p.sampler_config.temperature = 0.0;
3579 p.sampler_config.repetition_penalty = 1.0;
3580 let r = p.generate("abcd", 10, None, None).unwrap();
3581 assert!(r.traces.is_empty(), "trace must be empty unless enabled");
3582
3583 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
3585 p.sampler_config.temperature = 0.0;
3586 p.sampler_config.repetition_penalty = 1.0;
3587 p.set_trace(true);
3588 let r = p.generate("abcd", 10, None, None).unwrap();
3589 assert_eq!(r.traces.len(), r.token_ids.len(), "one trace row per token");
3590 for (i, tr) in r.traces.iter().enumerate() {
3591 assert_eq!(tr.t, i, "trace index is sequential");
3592 assert_eq!(tr.token_id, r.token_ids[i], "trace token_id matches output");
3593 assert_eq!(
3594 tr.confidence, r.token_confidence[i],
3595 "trace confidence matches the confidence channel"
3596 );
3597 assert!(tr.active_skill.is_none() && tr.recon.is_none() && !tr.switched);
3599 }
3600 }
3601
3602 #[test]
3603 fn explain_prefill_logits_match_greedy_first_token() {
3604 let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
3608 p.sampler_config.temperature = 0.0;
3609 p.sampler_config.repetition_penalty = 1.0;
3610 let ids = p.tokenizer.encode("abcd");
3611 let logits = p.prefill_next_logits(&ids, None);
3612 let argmax = logits
3613 .iter()
3614 .enumerate()
3615 .max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
3616 .unwrap()
3617 .0 as u32;
3618 let r = p.generate("abcd", 1, None, None).unwrap();
3619 assert_eq!(argmax, r.token_ids[0], "explain preview must match greedy emit");
3620 }
3621}