Skip to main content

cortiq_engine/
qtensor.rs

1//! QTensor — weight tensor with pluggable storage.
2//!
3//! Two backings, one interface:
4//! - `F32`   — owned dense floats (small models, tests). Every operation
5//!   is bit-identical to the historical `&[f32]` code paths.
6//! - `Mapped` — quantized bytes zero-copy from the CMF mmap (`q8_row` /
7//!   `q8_2f`). The matvec is fused: int8 rows × f32 activations, the
8//!   q8_2f column field folds into a pre-scale of the input
9//!   (`x'[i] = col[i]·x[i]`), so the inner loop is the same i8 dot as
10//!   q8_row. This is what lets a 15B file run in a few GB of RSS.
11//!
12//! Extension point: new dtypes = new match arm here, nothing else moves.
13
14use crate::pool::{Pool, matvec_rows, matvec_rows2};
15use cortiq_core::quant::{
16    GROUP_SIZE, Q1_TILE, Q2TP_CHUNK, Q4_TILE, Q4TP_NIB, f16_to_f32, q2tp_ladder, q2tp_sections,
17    q4tp_code, q4tp_ladder, q4tp_sections,
18};
19use cortiq_core::{CmfModel, TensorDtype};
20use std::cell::UnsafeCell;
21use std::sync::Arc;
22
23pub enum QTensor {
24    F32 {
25        data: Vec<f32>,
26        rows: usize,
27        cols: usize,
28    },
29    Mapped {
30        model: Arc<CmfModel>,
31        /// Index into the model's tensor directory.
32        idx: usize,
33        dtype: TensorDtype,
34        rows: usize,
35        cols: usize,
36        /// Per-row scales, dequantized to f32 up front (tiny).
37        row_scale: Vec<f32>,
38        /// q8_2f column field (θ), dequantized up front; empty for q8_row.
39        col_field: Vec<f32>,
40        /// Vbit only: byte offset of each row's packed data within the
41        /// tensor blob (`[rows + 1]`, computed once at load — the per-
42        /// matvec prefix scan over row bit-widths was O(rows) each call).
43        vbit_offsets: Vec<usize>,
44        /// q8-family decode repack (load-time, optional): rows in groups
45        /// of 4, interleaved in 16-byte units — one 64-byte line per
46        /// iteration feeds all 4 sdot lanes, ONE sequential weight
47        /// stream per worker instead of four (this is where llama.cpp's
48        /// repacked Q8 kernels get their bandwidth). Empty = off
49        /// (CMF_REPACK=0, non-SDOT arch, or an ineligible shape). Trades
50        /// an anonymous copy of the quants for mmap pages that go cold.
51        repack: Vec<u8>,
52    },
53}
54
55/// Load-time q8 repack gate (see `Mapped::repack`). OPT-IN
56/// (`CMF_REPACK=1`): the single-stream hypothesis LOST on Apple Silicon
57/// (M4, interleaved A/B: decode 101 vs 94 tok/s — four adjacent row
58/// streams per worker feed the prefetcher MORE memory-level parallelism
59/// than one); kept as an experiment flag for x86, where the tradeoff
60/// may land differently.
61fn repack_enabled() -> bool {
62    static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
63    *ON.get_or_init(|| {
64        std::env::var("CMF_REPACK")
65            .map(|v| v == "1")
66            .unwrap_or(cfg!(target_os = "android"))
67    })
68}
69
70/// Interleave q8 rows for the decode kernel: group g holds rows
71/// 4g..4g+4 as [r0[c], r1[c], r2[c], r3[c]] per 16-byte chunk c. Only
72/// full groups are packed — tail rows keep reading the mmap layout.
73fn q8_repack(bytes: &[u8], rows: usize, cols: usize) -> Vec<u8> {
74    #[cfg(target_arch = "aarch64")]
75    let arch_ok = sdot_enabled();
76    #[cfg(not(target_arch = "aarch64"))]
77    let arch_ok = false;
78    if !arch_ok || !repack_enabled() || rows < 256 || cols % 16 != 0 {
79        return Vec::new();
80    }
81    q8_repack_layout(bytes, rows, cols)
82}
83
84/// The pure layout transform behind `q8_repack` (tested directly —
85/// the gate depends on arch and env).
86fn q8_repack_layout(bytes: &[u8], rows: usize, cols: usize) -> Vec<u8> {
87    let groups = rows / 4;
88    let mut rep = vec![0u8; groups * 4 * cols];
89    for g in 0..groups {
90        let dst = &mut rep[g * 4 * cols..(g + 1) * 4 * cols];
91        for c in 0..cols / 16 {
92            for lane in 0..4 {
93                let src = (g * 4 + lane) * cols + c * 16;
94                dst[c * 64 + lane * 16..c * 64 + lane * 16 + 16]
95                    .copy_from_slice(&bytes[src..src + 16]);
96            }
97        }
98    }
99    rep
100}
101
102/// Prefix-sum of vbit row payload offsets (absolute within the tensor
103/// bytes). `offsets[r]..offsets[r+1]` is row r's packed data.
104fn vbit_row_offsets(bytes: &[u8], rows: usize, cols: usize) -> Vec<usize> {
105    let ng = cols / GROUP_SIZE;
106    let bits = &bytes[..rows];
107    let mut offsets = Vec::with_capacity(rows + 1);
108    let mut off = rows + rows * ng * 2;
109    for r in 0..rows {
110        offsets.push(off);
111        off += (cols * bits[r] as usize).div_ceil(8);
112    }
113    offsets.push(off);
114    offsets
115}
116
117/// `CMF_X86_BLOCKED` / `CMF_GPU_LMHEAD` / `CMF_GPU_SPLIT`, read once. They
118/// used to be read from the environment on every large matvec and on every
119/// matmat in six places — microseconds each, but also a knob that could
120/// change under a running process, which is not a thing a kernel choice
121/// should be able to do mid-sequence.
122fn blocked_enabled() -> bool {
123    use std::sync::atomic::Ordering::Relaxed;
124    match BLOCKED_OVERRIDE.load(Relaxed) {
125        1 => false,
126        2 => true,
127        _ => {
128            static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
129            *ON.get_or_init(|| {
130                std::env::var("CMF_X86_BLOCKED")
131                    .map(|v| v != "0")
132                    .unwrap_or(true)
133            })
134        }
135    }
136}
137
138static BLOCKED_OVERRIDE: std::sync::atomic::AtomicU8 = std::sync::atomic::AtomicU8::new(0);
139
140/// Force the blocked GEMM on or off, ignoring the environment; `None`
141/// restores it. For tests that need to run BOTH paths and compare them:
142/// `blocked_enabled` caches its answer for the life of the process, which
143/// is right when the environment is the only input, but leaves a test that
144/// flips `CMF_X86_BLOCKED` between two calls comparing a path against
145/// itself — or against whatever a test running in parallel latched first.
146pub fn set_blocked_override(on: Option<bool>) {
147    let v = match on {
148        None => 0,
149        Some(false) => 1,
150        Some(true) => 2,
151    };
152    BLOCKED_OVERRIDE.store(v, std::sync::atomic::Ordering::Relaxed);
153}
154
155fn gpu_lmhead_enabled() -> bool {
156    static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
157    *ON.get_or_init(|| {
158        std::env::var("CMF_GPU_LMHEAD")
159            .map(|v| v != "0")
160            .unwrap_or(true)
161    })
162}
163
164fn gpu_split_frac() -> f32 {
165    // MiMo's banked verification projects every row on the device. Plain
166    // decode must not quantize half of its O/head activation on the CPU.
167    if FULL_GPU_Q8.get() {
168        return 1.0;
169    }
170    static F: std::sync::OnceLock<f32> = std::sync::OnceLock::new();
171    *F.get_or_init(|| {
172        std::env::var("CMF_GPU_SPLIT")
173            .ok()
174            .and_then(|v| v.parse::<f32>().ok())
175            .unwrap_or(0.5)
176            .clamp(0.0, 1.0)
177    })
178}
179
180impl QTensor {
181    pub fn from_f32(data: Vec<f32>, rows: usize, cols: usize) -> Self {
182        debug_assert_eq!(data.len(), rows * cols);
183        Self::F32 { data, rows, cols }
184    }
185
186    /// Wrap a directory tensor without dequantizing the payload.
187    /// Falls back to dequantized f32 for dtypes without a fused kernel.
188    pub fn from_model(model: &Arc<CmfModel>, name: &str) -> Result<Self, String> {
189        // Indexed lookup: the linear directory scan made pipeline build
190        // O(N²) on MoE/skills files with thousands of tensors.
191        let idx = model
192            .tensor_index(name)
193            .ok_or_else(|| format!("tensor '{name}' not found in CMF directory"))?;
194        let entry = &model.tensors[idx];
195        if entry.shape.len() != 2 {
196            return Err(format!("QTensor::from_model needs 2-D, got '{name}'"));
197        }
198        let (rows, cols) = (entry.shape[0], entry.shape[1]);
199        let bytes = model.entry_bytes(entry);
200
201        match entry.dtype {
202            TensorDtype::Q8Row | TensorDtype::Q8_2f => {
203                let n = rows * cols;
204                let scales_off = n;
205                let row_scale: Vec<f32> = (0..rows)
206                    .map(|o| {
207                        f16_to_f32(u16::from_le_bytes([
208                            bytes[scales_off + o * 2],
209                            bytes[scales_off + o * 2 + 1],
210                        ]))
211                    })
212                    .collect();
213                let col_field: Vec<f32> = if entry.dtype == TensorDtype::Q8_2f {
214                    let col_off = n + rows * 2;
215                    (0..cols)
216                        .map(|i| {
217                            f16_to_f32(u16::from_le_bytes([
218                                bytes[col_off + i * 2],
219                                bytes[col_off + i * 2 + 1],
220                            ]))
221                        })
222                        .collect()
223                } else {
224                    Vec::new()
225                };
226                Ok(Self::Mapped {
227                    model: model.clone(),
228                    idx,
229                    dtype: entry.dtype,
230                    rows,
231                    cols,
232                    row_scale,
233                    col_field,
234                    vbit_offsets: Vec::new(),
235                    repack: q8_repack(bytes, rows, cols),
236                })
237            }
238            // vbit: fused kernel unpacks variable-bit rows from mmap.
239            TensorDtype::Vbit if cols % GROUP_SIZE == 0 => Ok(Self::Mapped {
240                model: model.clone(),
241                idx,
242                dtype: entry.dtype,
243                rows,
244                cols,
245                row_scale: Vec::new(),
246                col_field: Vec::new(),
247                vbit_offsets: vbit_row_offsets(bytes, rows, cols),
248                repack: Vec::new(),
249            }),
250            // vbit_ro (§4.2): the offset table comes straight from the
251            // file — no load-time prefix scan; kernels are shared with
252            // legacy vbit (they consume absolute offsets either way).
253            TensorDtype::VbitRo if cols % GROUP_SIZE == 0 => {
254                let (_, off_off, packed_off) = cortiq_core::quant::vbit_ro_sections(rows, cols);
255                let offsets: Vec<usize> = (0..=rows)
256                    .map(|r| packed_off + cortiq_core::quant::vbit_ro_offset(bytes, off_off, r))
257                    .collect();
258                Ok(Self::Mapped {
259                    model: model.clone(),
260                    idx,
261                    dtype: entry.dtype,
262                    rows,
263                    cols,
264                    row_scale: Vec::new(),
265                    col_field: Vec::new(),
266                    vbit_offsets: offsets,
267                    repack: Vec::new(),
268                })
269            }
270            // q4_block: fused kernel reads nibbles straight from mmap —
271            // a 14B q4 file no longer explodes into ×8 f32 RAM.
272            // q4_tiled (§4.3): interleaved [scale][nibbles] tiles — one
273            // sequential memory stream (measured ×1.66 ARM / ×1.13 AVX2
274            // at kernel level over the split layout).
275            TensorDtype::Q4Tiled if cols % GROUP_SIZE == 0 => Ok(Self::Mapped {
276                model: model.clone(),
277                idx,
278                dtype: entry.dtype,
279                rows,
280                cols,
281                row_scale: Vec::new(),
282                col_field: Vec::new(),
283                vbit_offsets: Vec::new(),
284                repack: Vec::new(),
285            }),
286            // q4tp (§4.10): nibbles from mmap, scale from the row ladder —
287            // 7.3% less file than q4t at the same 4-bit grid.
288            TensorDtype::Q4TiledP if cols % GROUP_SIZE == 0 => Ok(Self::Mapped {
289                model: model.clone(),
290                idx,
291                dtype: entry.dtype,
292                rows,
293                cols,
294                row_scale: Vec::new(),
295                col_field: Vec::new(),
296                vbit_offsets: Vec::new(),
297                repack: Vec::new(),
298            }),
299            // q2tp: 2-bit chunks from mmap, scale from the same row ladder.
300            TensorDtype::Q2TiledP if cols % GROUP_SIZE == 0 => Ok(Self::Mapped {
301                model: model.clone(),
302                idx,
303                dtype: entry.dtype,
304                rows,
305                cols,
306                row_scale: Vec::new(),
307                col_field: Vec::new(),
308                vbit_offsets: Vec::new(),
309                repack: Vec::new(),
310            }),
311            TensorDtype::Q4Block if cols % GROUP_SIZE == 0 => Ok(Self::Mapped {
312                model: model.clone(),
313                idx,
314                dtype: entry.dtype,
315                rows,
316                cols,
317                row_scale: Vec::new(),
318                col_field: Vec::new(),
319                vbit_offsets: Vec::new(),
320                repack: Vec::new(),
321            }),
322            // q1: binary sign-bit tiles from mmap (1-bit-trained models).
323            TensorDtype::Q1 if cols % GROUP_SIZE == 0 => Ok(Self::Mapped {
324                model: model.clone(),
325                idx,
326                dtype: entry.dtype,
327                rows,
328                cols,
329                row_scale: Vec::new(),
330                col_field: Vec::new(),
331                vbit_offsets: Vec::new(),
332                repack: Vec::new(),
333            }),
334            // q1t (ternary + outlier overlay): fused per-row dequant kernel
335            // reads straight from mmap — a 12B q1t stays ~its file size in
336            // RAM instead of dequantizing to ~48 GB of f32.
337            TensorDtype::Q1T if cols % GROUP_SIZE == 0 => Ok(Self::Mapped {
338                model: model.clone(),
339                idx,
340                dtype: entry.dtype,
341                rows,
342                cols,
343                row_scale: Vec::new(),
344                col_field: Vec::new(),
345                vbit_offsets: Vec::new(),
346                repack: Vec::new(),
347            }),
348            // No fused kernel yet → dequantize once (correct, more RAM).
349            _ => {
350                let mut data = vec![0.0f32; rows * cols];
351                cortiq_core::quant::dequant_tensor(entry, bytes, &mut data)?;
352                Ok(Self::from_f32(data, rows, cols))
353            }
354        }
355    }
356
357    /// q1-mapped tensor? (GPU gates: the q1 CPU kernel is
358    /// compute-bound, so offload pays at much smaller shapes than q8.)
359    pub(crate) fn is_q1(&self) -> bool {
360        matches!(
361            self,
362            Self::Mapped {
363                dtype: TensorDtype::Q1,
364                ..
365            }
366        )
367    }
368
369    /// Owned-f32 view (data, rows, cols) — the GDN a/b gate projections
370    /// arrive dequantized (force-f16 in the converter → F32 in RAM).
371    pub(crate) fn f32_parts(&self) -> Option<(&[f32], usize, usize)> {
372        match self {
373            Self::F32 { data, rows, cols } => Some((data, *rows, *cols)),
374            _ => None,
375        }
376    }
377
378    /// (directory idx, rows, cols) of a q1-mapped tensor — the
379    /// whole-block GPU path resolves offsets itself.
380    /// (idx, rows, cols) of a mapped tensor the whole-token GPU graph can drive
381    /// — Q1, Q1T or Q4-block (it resolves the offset and picks the kernel by
382    /// dtype). Q4-block lets a precise down_proj/lm_head stay on-device.
383    /// Named `q1_parts` for historical reasons.
384    pub(crate) fn q1_parts(&self) -> Option<(usize, usize, usize)> {
385        if self.has_prism_contract() {
386            return None;
387        }
388        match self {
389            #[cfg(target_os = "macos")]
390            Self::Mapped {
391                dtype: TensorDtype::Q1T,
392                ..
393            } if !crate::gpu::metal_q1t_enabled() => None,
394            Self::Mapped {
395                idx,
396                dtype:
397                    TensorDtype::Q1
398                    | TensorDtype::Q1T
399                    | TensorDtype::Q4Block
400                    | TensorDtype::Q4Tiled
401                    // Q2TiledP deliberately absent: the Metal graph has no
402                    // q2tp kernel, and advertising it here made the block
403                    // plan truncate mid-run at the first q2tp layer.
404                    | TensorDtype::Q4TiledP
405                    | TensorDtype::Q8Row
406                    | TensorDtype::Q8_2f,
407                rows,
408                cols,
409                ..
410            } => Some((*idx, *rows, *cols)),
411            _ => None,
412        }
413    }
414
415    /// `(directory idx, rows, cols)` for the native Metal token graph.  The
416    /// historical q1 graph gate intentionally refuses every Prism tensor so
417    /// an untransformed q2 payload cannot slip into the resident path.  The
418    /// Metal2 graph is descriptor-aware and admits only the production
419    /// q2tp-affine forward targets; ordinary q1/q4 callers retain the old
420    /// `q1_parts` behaviour.
421    #[cfg(target_os = "macos")]
422    pub(crate) fn metal_graph_parts(&self) -> Option<(usize, usize, usize)> {
423        if let Some((model, idx, kind, _)) = self.graph_weight_descriptor() {
424            let name = &model.tensors[idx].name;
425            let forward = kind == 9 && crate::prism::is_forward_weight(model, name);
426            let affine = kind == 9 && crate::prism::is_affine_target(model, name);
427            if forward && affine {
428                let e = model.tensors.get(idx)?;
429                return Some((idx, *e.shape.first()?, *e.shape.get(1)?));
430            }
431        }
432        self.q1_parts()
433    }
434
435    /// (directory idx, rows, cols) of a q4_tiled mapped tensor. The
436    /// chunk-prefill graph takes it in the same 4-tuple slot as
437    /// `q8_row_parts` with an EMPTY row_scale — q4t carries its scales
438    /// inside the 18-byte tiles, and the empty slice is what tells the
439    /// encoder to reach for the q4t kernels.
440    pub(crate) fn q4t_parts(&self) -> Option<(usize, usize, usize)> {
441        if self.has_prism_contract() {
442            return None;
443        }
444        match self {
445            Self::Mapped {
446                idx,
447                dtype: TensorDtype::Q4Tiled,
448                rows,
449                cols,
450                ..
451            } => Some((*idx, *rows, *cols)),
452            _ => None,
453        }
454    }
455
456    /// (directory idx, rows, cols) of a q4tp mapped tensor. Same empty-scale
457    /// slot as `q4t_parts` in the chunk graph — the encoder tells the two
458    /// apart by the tensor's dtype, not by the slot.
459    pub(crate) fn q4tp_parts(&self) -> Option<(usize, usize, usize)> {
460        if self.has_prism_contract() {
461            return None;
462        }
463        match self {
464            Self::Mapped {
465                idx,
466                dtype: TensorDtype::Q4TiledP,
467                rows,
468                cols,
469                ..
470            } => Some((*idx, *rows, *cols)),
471            _ => None,
472        }
473    }
474
475    /// (directory idx, rows, cols, row_scale) of a plain q8_row mapped
476    /// tensor — the chunk-prefill GPU graph resolves offsets itself.
477    /// q8_2f is excluded on purpose: its column field would need a
478    /// prescale stage on the device.
479    pub(crate) fn q8_row_parts(&self) -> Option<(usize, usize, usize, &[f32])> {
480        if self.has_prism_contract() {
481            return None;
482        }
483        match self {
484            Self::Mapped {
485                idx,
486                dtype: TensorDtype::Q8Row,
487                rows,
488                cols,
489                row_scale,
490                col_field,
491                ..
492            } if col_field.is_empty() => Some((*idx, *rows, *cols, row_scale)),
493            _ => None,
494        }
495    }
496
497    /// The layout this tensor is stored in, when it is mapped from a model.
498    /// The frames branch on it — a q2tp gate against a q4tp down is a real
499    /// combination in the 2-bit profile and needs a different kernel.
500    pub fn model_dtype(&self) -> Option<cortiq_core::TensorDtype> {
501        match self {
502            Self::Mapped { dtype, .. } => Some(*dtype),
503            _ => None,
504        }
505    }
506
507    /// The tensor's index in the model directory, when it is mapped from one.
508    /// The GPU frames bind by index rather than by name — a name lookup per
509    /// layer per token is not free, and the index is what the device cache is
510    /// keyed on anyway.
511    pub fn model_idx(&self) -> Option<usize> {
512        match self {
513            Self::Mapped { idx, .. } => Some(*idx),
514            _ => None,
515        }
516    }
517
518    /// The model this tensor is mapped from, when it is mapped at all. The
519    /// GPU frames need the container to reach the bytes; a QTensor already
520    /// holds it, and threading a second handle down every call site to say
521    /// the same thing invites the two to disagree.
522    pub fn model_arc(&self) -> Option<std::sync::Arc<cortiq_core::CmfModel>> {
523        match self {
524            Self::Mapped { model, .. } => Some(model.clone()),
525            _ => None,
526        }
527    }
528
529    /// Whether this mapped tensor belongs to the Prism/Bonsai transform
530    /// contract.  Device graphs do not carry the descriptor, so callers use
531    /// this conservative predicate to stay on the descriptor-aware CPU path
532    /// instead of silently executing an unrotated matrix.
533    pub(crate) fn has_prism_contract(&self) -> bool {
534        matches!(self, Self::Mapped { model, .. } if crate::prism::has_contract(model))
535    }
536
537    pub fn rows(&self) -> usize {
538        match self {
539            Self::F32 { rows, .. } | Self::Mapped { rows, .. } => *rows,
540        }
541    }
542
543    /// Mapped q4t handle (model + directory index) — the fused GPU FFN
544    /// needs the raw file coordinates of its three projections.
545    pub(crate) fn mapped_q4t(&self) -> Option<(&Arc<CmfModel>, usize)> {
546        if self.has_prism_contract() {
547            return None;
548        }
549        match self {
550            Self::Mapped {
551                model,
552                idx,
553                dtype: TensorDtype::Q4Tiled,
554                ..
555            } => Some((model, *idx)),
556            _ => None,
557        }
558    }
559
560    /// Same slot as `mapped_q4t` for a q4tp tensor — the fused DiT FFN picks
561    /// its kernels by which of the two answers.
562    pub fn mapped_q4tp(&self) -> Option<(&Arc<CmfModel>, usize)> {
563        if self.has_prism_contract() {
564            return None;
565        }
566        match self {
567            Self::Mapped {
568                model,
569                idx,
570                dtype: TensorDtype::Q4TiledP,
571                ..
572            } => Some((model, *idx)),
573            _ => None,
574        }
575    }
576
577    /// (model, tensor idx) for a mapped weight in ANY codec the fused device
578    /// paths can run — four-bit tiled or either int8 layout.
579    ///
580    /// The fused DiT chains asked for `mapped_q4tp` by name, so an eight-bit
581    /// container never reached them and rendered through per-op GEMMs even
582    /// after those kernels learned its codec. The gate is what the codec has
583    /// a device GEMM for, not which codec it is.
584    pub fn mapped_device_gemm(&self) -> Option<(&Arc<CmfModel>, usize)> {
585        if self.has_prism_contract() {
586            return None;
587        }
588        match self {
589            Self::Mapped {
590                model,
591                idx,
592                dtype: TensorDtype::Q4TiledP | TensorDtype::Q8Row | TensorDtype::Q8_2f,
593                ..
594            } => Some((model, *idx)),
595            _ => None,
596        }
597    }
598
599    /// (model, tensor idx) for a q2tp mapped weight — the 2-bit twin of
600    /// `mapped_q4tp`, used by the mixed MoE profile.
601    pub fn mapped_q2tp(&self) -> Option<(&Arc<CmfModel>, usize)> {
602        if self.has_prism_contract() {
603            return None;
604        }
605        match self {
606            Self::Mapped {
607                model,
608                idx,
609                dtype: TensorDtype::Q2TiledP,
610                ..
611            } => Some((model, *idx)),
612            _ => None,
613        }
614    }
615
616    pub fn cols(&self) -> usize {
617        match self {
618            Self::F32 { cols, .. } | Self::Mapped { cols, .. } => *cols,
619        }
620    }
621
622    /// (model, tensor idx) for a q1 mapped weight — the wgpu token graph
623    /// keys its resident VRAM cache by idx. None for any other dtype/kind.
624    pub fn mapped_q1(&self) -> Option<(&std::sync::Arc<CmfModel>, usize)> {
625        if self.has_prism_contract() {
626            return None;
627        }
628        match self {
629            Self::Mapped {
630                model,
631                idx,
632                dtype: TensorDtype::Q1,
633                ..
634            } => Some((model, *idx)),
635            _ => None,
636        }
637    }
638
639    /// (model, idx, kind, row_scale) for a graph-capable mapped weight.
640    /// kind: 0=q8_row (per-row scales), 1=q1, 2=q4_block, 3=q1t
641    /// (tile-embedded, no rs), 5=q4_tiled, 6=q4tp, 7=q8_2f (both scale
642    /// planes live inside the tensor). None only for `vbit`.
643    ///
644    /// The old comment here claimed q4_block was unhandled while the arm
645    /// right below mapped it, and it named q8_2f as unhandled after that
646    /// stopped being true — a stale comment on this function is how a
647    /// model silently loses the graph, so it is worth keeping honest.
648    pub fn graph_weight(&self) -> Option<(&std::sync::Arc<CmfModel>, usize, u8, &[f32])> {
649        if self.has_prism_contract() {
650            return None;
651        }
652        self.graph_weight_descriptor()
653    }
654
655    /// Descriptor-aware graph handle used only by the Prism token graph.
656    /// Ordinary graph callers continue to use [`graph_weight`] and therefore
657    /// remain fail-closed until they provide the same explicit transform
658    /// contract.
659    pub(crate) fn graph_weight_descriptor(
660        &self,
661    ) -> Option<(&std::sync::Arc<CmfModel>, usize, u8, &[f32])> {
662        match self {
663            Self::Mapped {
664                model,
665                idx,
666                dtype: TensorDtype::Q8Row,
667                row_scale,
668                ..
669            } => Some((model, *idx, 0, row_scale.as_slice())),
670            Self::Mapped {
671                model,
672                idx,
673                dtype: TensorDtype::Q1,
674                ..
675            } => Some((model, *idx, 1, &[])),
676            // Q4Tiled is kind 5, NOT 2: both carried 2 historically, and
677            // the wgpu token graph fed 18B interleaved tiles to the
678            // split-layout q4b kernel — garbage output on q4t models
679            // (caught by an end-to-end answer check on real Vulkan).
680            Self::Mapped {
681                model,
682                idx,
683                dtype: TensorDtype::Q4Tiled,
684                ..
685            } => Some((model, *idx, 5, &[])),
686            // Kind 6, not 5: q4tp's nibble stride and scale planes differ,
687            // and feeding them to the q4t kernel is exactly the mistake that
688            // produced garbage when Q4Tiled shared kind 2 with Q4Block.
689            Self::Mapped {
690                model,
691                idx,
692                dtype: TensorDtype::Q4TiledP,
693                ..
694            } => Some((model, *idx, 6, &[])),
695            Self::Mapped {
696                model,
697                idx,
698                dtype: TensorDtype::Q4Block,
699                ..
700            } => Some((model, *idx, 2, &[])),
701            // q8_2f carries BOTH scale planes after the int8 body (rows
702            // f16, then cols f16), so the graph takes the whole tensor
703            // and the kernel reads them where they lie — no host-side
704            // prescale, which is what the per-op path does instead.
705            Self::Mapped {
706                model,
707                idx,
708                dtype: TensorDtype::Q8_2f,
709                ..
710            } => Some((model, *idx, 7, &[])),
711            Self::Mapped {
712                model,
713                idx,
714                dtype: TensorDtype::Q1T,
715                ..
716            } => Some((model, *idx, 3, &[])),
717            // Kind 9: the 2-bit plane on the q4tp ladder (dense FFN gate/up
718            // of the q2tp profile). Its own kernel — 8 bytes a group where
719            // q4tp has 16, and rung 0 is the exact zero.
720            Self::Mapped {
721                model,
722                idx,
723                dtype: TensorDtype::Q2TiledP,
724                ..
725            } => Some((model, *idx, 9, &[])),
726            _ => None,
727        }
728    }
729
730    /// Dense f32 view — only for owned tensors. Masked/sparse execution
731    /// paths require it; quantized weights don't support masks yet.
732    pub fn as_f32(&self) -> Option<&[f32]> {
733        match self {
734            Self::F32 { data, .. } => Some(data),
735            Self::Mapped { .. } => None,
736        }
737    }
738
739    fn quant_bytes(&self) -> &[u8] {
740        match self {
741            Self::Mapped { model, idx, .. } => model.entry_bytes(&model.tensors[*idx]),
742            Self::F32 { .. } => unreachable!("quant_bytes on F32"),
743        }
744    }
745
746    /// Dequantize one row into `dst` (embedding lookup).
747    pub fn row_f32(&self, r: usize, dst: &mut [f32]) {
748        let cols = self.cols();
749        debug_assert_eq!(dst.len(), cols);
750        match self {
751            Self::F32 { data, .. } => dst.copy_from_slice(&data[r * cols..(r + 1) * cols]),
752            Self::Mapped {
753                model,
754                idx,
755                dtype,
756                row_scale,
757                col_field,
758                vbit_offsets,
759                ..
760            } => {
761                if *dtype == TensorDtype::Q4Tiled {
762                    let bytes = self.quant_bytes();
763                    let gpr = cols / GROUP_SIZE;
764                    for gi in 0..gpr {
765                        let tile = &bytes[(r * gpr + gi) * Q4_TILE..(r * gpr + gi + 1) * Q4_TILE];
766                        let s = f16_to_f32(u16::from_le_bytes([tile[0], tile[1]]));
767                        for (k, &b) in tile[2..].iter().enumerate() {
768                            dst[gi * GROUP_SIZE + k * 2] = ((b & 0x0F) as f32 - 8.0) * s;
769                            dst[gi * GROUP_SIZE + k * 2 + 1] = (((b >> 4) & 0x0F) as f32 - 8.0) * s;
770                        }
771                    }
772                    if crate::prism::is_inverse_embedding(model, &model.tensors[*idx].name) {
773                        crate::prism::inverse_embedding(model, dst);
774                    }
775                    return;
776                }
777                if *dtype == TensorDtype::Q4TiledP {
778                    let bytes = self.quant_bytes();
779                    let gpr = cols / GROUP_SIZE;
780                    let v = Q4tpView::new(bytes, self.rows(), cols);
781                    let mut sc = vec![0f32; gpr];
782                    v.scales_into(r, gpr, &mut sc);
783                    for gi in 0..gpr {
784                        let tile = &v.nib[(r * gpr + gi) * Q4TP_NIB..(r * gpr + gi + 1) * Q4TP_NIB];
785                        let s = sc[gi];
786                        for (k, &b) in tile.iter().enumerate() {
787                            dst[gi * GROUP_SIZE + k * 2] = ((b & 0x0F) as f32 - 8.0) * s;
788                            dst[gi * GROUP_SIZE + k * 2 + 1] = (((b >> 4) & 0x0F) as f32 - 8.0) * s;
789                        }
790                    }
791                    if crate::prism::is_inverse_embedding(model, &model.tensors[*idx].name) {
792                        crate::prism::inverse_embedding(model, dst);
793                    }
794                    return;
795                }
796                if *dtype == TensorDtype::Q2TiledP {
797                    let bytes = self.quant_bytes();
798                    let gpr = cols / GROUP_SIZE;
799                    let v = Q4tpView::new_q2(bytes, self.rows(), cols);
800                    let mut sc = vec![0f32; gpr];
801                    v.scales_into(r, gpr, &mut sc);
802                    for gi in 0..gpr {
803                        let ch =
804                            &v.nib[(r * gpr + gi) * Q2TP_CHUNK..(r * gpr + gi + 1) * Q2TP_CHUNK];
805                        let s = sc[gi];
806                        for (k, &b) in ch.iter().enumerate() {
807                            for j in 0..4 {
808                                let center = if crate::prism::is_affine_target(
809                                    model,
810                                    &model.tensors[*idx].name,
811                                ) {
812                                    1.0
813                                } else {
814                                    1.5
815                                };
816                                dst[gi * GROUP_SIZE + k * 4 + j] =
817                                    (((b >> (2 * j)) & 3) as f32 - center) * s;
818                            }
819                        }
820                    }
821                    if crate::prism::is_inverse_embedding(model, &model.tensors[*idx].name) {
822                        crate::prism::inverse_embedding(model, dst);
823                    }
824                    return;
825                }
826                if *dtype == TensorDtype::Q4Block {
827                    let (packed, scales) = q4_split(self.quant_bytes(), self.rows(), cols);
828                    let gpr = cols / GROUP_SIZE;
829                    for gi in 0..gpr {
830                        let g = r * gpr + gi;
831                        let s = f16_to_f32(u16::from_le_bytes([scales[g * 2], scales[g * 2 + 1]]));
832                        for (k, &b) in packed[g * 16..(g + 1) * 16].iter().enumerate() {
833                            dst[gi * GROUP_SIZE + k * 2] = ((b & 0x0F) as f32 - 8.0) * s;
834                            dst[gi * GROUP_SIZE + k * 2 + 1] = (((b >> 4) & 0x0F) as f32 - 8.0) * s;
835                        }
836                    }
837                    if crate::prism::is_inverse_embedding(model, &model.tensors[*idx].name) {
838                        crate::prism::inverse_embedding(model, dst);
839                    }
840                    return;
841                }
842                if *dtype == TensorDtype::Q1 {
843                    let bytes = self.quant_bytes();
844                    let gpr = cols / GROUP_SIZE;
845                    for gi in 0..gpr {
846                        let tile = &bytes[(r * gpr + gi) * Q1_TILE..(r * gpr + gi + 1) * Q1_TILE];
847                        let s = f16_to_f32(u16::from_le_bytes([tile[0], tile[1]]));
848                        for (j, &b) in tile[2..].iter().enumerate() {
849                            for k in 0..8 {
850                                dst[gi * GROUP_SIZE + j * 8 + k] =
851                                    (((b >> k) & 1) as f32 * 2.0 - 1.0) * s;
852                            }
853                        }
854                    }
855                    if crate::prism::is_inverse_embedding(model, &model.tensors[*idx].name) {
856                        crate::prism::inverse_embedding(model, dst);
857                    }
858                    return;
859                }
860                if *dtype == TensorDtype::Q1T {
861                    let bytes = self.quant_bytes();
862                    let gpr = cols / GROUP_SIZE;
863                    let base_len = self.rows() * gpr * cortiq_core::quant::Q1T_TILE;
864                    for gi in 0..gpr {
865                        let off = (r * gpr + gi) * cortiq_core::quant::Q1T_TILE;
866                        let s = cortiq_core::quant::f16_to_f32(u16::from_le_bytes([
867                            bytes[off],
868                            bytes[off + 1],
869                        ]));
870                        let codes = &bytes[off + 2..off + cortiq_core::quant::Q1T_TILE];
871                        for k in 0..GROUP_SIZE {
872                            dst[gi * GROUP_SIZE + k] = match cortiq_core::quant::q1t_code(codes, k)
873                            {
874                                1 => s,
875                                2 => -s,
876                                _ => 0.0,
877                            };
878                        }
879                    }
880                    // Overlay
881                    let rows = self.rows();
882                    let entries = base_len + (rows + 1) * 4;
883                    if entries <= bytes.len() {
884                        let ptrs = &bytes[base_len..base_len + (rows + 1) * 4];
885                        let r0 = u32::from_le_bytes([
886                            ptrs[r * 4],
887                            ptrs[r * 4 + 1],
888                            ptrs[r * 4 + 2],
889                            ptrs[r * 4 + 3],
890                        ]) as usize;
891                        let r1 = u32::from_le_bytes([
892                            ptrs[(r + 1) * 4],
893                            ptrs[(r + 1) * 4 + 1],
894                            ptrs[(r + 1) * 4 + 2],
895                            ptrs[(r + 1) * 4 + 3],
896                        ]) as usize;
897                        let off = entries + r0 * 4;
898                        for i in 0..r1 - r0 {
899                            let item = &bytes[off + i * 4..off + i * 4 + 4];
900                            let c = u16::from_le_bytes([item[0], item[1]]) as usize;
901                            let v = cortiq_core::quant::f16_to_f32(u16::from_le_bytes([
902                                item[2], item[3],
903                            ]));
904                            if c < cols {
905                                dst[c] = v;
906                            }
907                        }
908                    }
909                    if crate::prism::is_inverse_embedding(model, &model.tensors[*idx].name) {
910                        crate::prism::inverse_embedding(model, dst);
911                    }
912                    return;
913                }
914                if matches!(dtype, TensorDtype::Vbit | TensorDtype::VbitRo) {
915                    let bytes = self.quant_bytes();
916                    let rows = self.rows();
917                    let ng = cols / GROUP_SIZE;
918                    let bits = &bytes[..rows];
919                    let sc_off = rows;
920                    // Precomputed at load — embedding lookup used to scan
921                    // the bit-widths of every preceding row (O(token_id)).
922                    let off = vbit_offsets[r];
923                    let b = bits[r] as usize;
924                    let l = ((1usize << (b - 1)) - 1) as f32;
925                    let data = &bytes[off..];
926                    let (mut acc, mut nbits, mut byte_idx) = (0u64, 0usize, 0usize);
927                    for (i, d) in dst.iter_mut().enumerate() {
928                        while nbits < b {
929                            acc = (acc << 8) | data[byte_idx] as u64;
930                            byte_idx += 1;
931                            nbits += 8;
932                        }
933                        let u = ((acc >> (nbits - b)) & ((1u64 << b) - 1)) as f32;
934                        nbits -= b;
935                        let so = (r * ng + i / GROUP_SIZE) * 2;
936                        let sv = f16_to_f32(u16::from_le_bytes([
937                            bytes[sc_off + so],
938                            bytes[sc_off + so + 1],
939                        ]));
940                        *d = (u - l) * sv;
941                    }
942                    if crate::prism::is_inverse_embedding(model, &model.tensors[*idx].name) {
943                        crate::prism::inverse_embedding(model, dst);
944                    }
945                    return;
946                }
947                let q = &self.quant_bytes()[r * cols..(r + 1) * cols];
948                let s = row_scale[r];
949                match dtype {
950                    TensorDtype::Q8Row => {
951                        for (d, &b) in dst.iter_mut().zip(q) {
952                            *d = (b as i8) as f32 * s;
953                        }
954                    }
955                    TensorDtype::Q8_2f => {
956                        for (i, (d, &b)) in dst.iter_mut().zip(q).enumerate() {
957                            *d = (b as i8) as f32 * s * col_field[i];
958                        }
959                    }
960                    _ => unreachable!(),
961                }
962                if crate::prism::is_inverse_embedding(model, &model.tensors[*idx].name) {
963                    crate::prism::inverse_embedding(model, dst);
964                }
965            }
966        }
967    }
968
969    /// Can this tensor's columns be read cheaply (for sparse down_proj)?
970    /// True for F32/Q8Row/Q8_2f (per-row scale, direct strided access);
971    /// false for group-packed q4/vbit (column access would unpack whole
972    /// groups — sparse execution falls back to f32 for those).
973    pub fn sparse_col_ok(&self) -> bool {
974        match self {
975            Self::F32 { .. } => true,
976            Self::Mapped { dtype, .. } => {
977                matches!(dtype, TensorDtype::Q8Row | TensorDtype::Q8_2f)
978            }
979        }
980    }
981
982    /// down_proj [hidden, inter]: accumulate `w · col(c)` into `out`
983    /// [hidden] — reads ONLY column `c` (one neuron) from the mmap,
984    /// no full-matrix dequant. `out[k] += w · down[k, c]`.
985    pub fn add_col_scaled(&self, c: usize, w: f32, out: &mut [f32]) {
986        let inter = self.cols();
987        let hidden = self.rows();
988        debug_assert_eq!(out.len(), hidden);
989        match self {
990            Self::F32 { data, .. } => {
991                for (k, o) in out.iter_mut().enumerate() {
992                    *o += w * data[k * inter + c];
993                }
994            }
995            Self::Mapped {
996                dtype,
997                row_scale,
998                col_field,
999                ..
1000            } => {
1001                let q = self.quant_bytes();
1002                let colf = if *dtype == TensorDtype::Q8_2f {
1003                    col_field[c]
1004                } else {
1005                    1.0
1006                };
1007                let wc = w * colf;
1008                for (k, o) in out.iter_mut().enumerate() {
1009                    let b = q[k * inter + c] as i8 as f32;
1010                    *o += wc * b * row_scale[k];
1011                }
1012            }
1013        }
1014    }
1015
1016    /// Touch the head of row `r` so the DRAM latency of the next
1017    /// neuron's weights overlaps the current one's arithmetic.
1018    ///
1019    /// Scattered rows are what per-token sparsity reads, and a 2 KB
1020    /// stride is past what the hardware prefetcher follows: without this
1021    /// every row starts with a cold miss that nothing hides. One touch
1022    /// per 512 bytes is enough — the rest of the row is a sequential run
1023    /// the prefetcher does pick up.
1024    #[inline]
1025    pub fn prefetch_row(&self, r: usize) {
1026        let Self::Mapped { dtype, .. } = self else {
1027            return;
1028        };
1029        if !matches!(dtype, TensorDtype::Q8Row | TensorDtype::Q8_2f) {
1030            return;
1031        }
1032        let cols = self.cols();
1033        let q = self.quant_bytes();
1034        let (a, b) = (r * cols, (r + 1) * cols);
1035        if b > q.len() {
1036            return;
1037        }
1038        let mut j = a;
1039        while j < b {
1040            unsafe { std::ptr::read_volatile(q.as_ptr().add(j)) };
1041            j += 512;
1042        }
1043    }
1044
1045    /// `out += w · row(r)` — the transposed twin of `add_col_scaled`.
1046    ///
1047    /// A neuron's `down` weights are a COLUMN of `[hidden, inter]`, and a
1048    /// column is strided: reading one costs a cache line per element, so
1049    /// per-neuron dynamic sparsity saves arithmetic and no bytes. Stored
1050    /// transposed (`down_proj.t.weight`, `[inter, hidden]`) the same
1051    /// weights are a contiguous ROW, and this accumulate reads exactly
1052    /// the neurons the token asked for.
1053    pub fn add_row_scaled(&self, r: usize, w: f32, out: &mut [f32], scratch: &mut [f32]) {
1054        let cols = self.cols();
1055        debug_assert_eq!(out.len(), cols);
1056        match self {
1057            Self::F32 { data, .. } => {
1058                let row = &data[r * cols..(r + 1) * cols];
1059                for (o, v) in out.iter_mut().zip(row) {
1060                    *o += w * v;
1061                }
1062            }
1063            Self::Mapped {
1064                dtype,
1065                row_scale,
1066                col_field,
1067                ..
1068            } => match dtype {
1069                TensorDtype::Q8Row => {
1070                    let q = &self.quant_bytes()[r * cols..(r + 1) * cols];
1071                    let ws = w * row_scale[r];
1072                    let row: &[i8] =
1073                        unsafe { std::slice::from_raw_parts(q.as_ptr() as *const i8, q.len()) };
1074                    axpy_i8_f32(out, row, ws);
1075                }
1076                TensorDtype::Q8_2f => {
1077                    let q = &self.quant_bytes()[r * cols..(r + 1) * cols];
1078                    let ws = w * row_scale[r];
1079                    for ((o, b), c) in out.iter_mut().zip(q).zip(col_field) {
1080                        *o += ws * c * (*b as i8 as f32);
1081                    }
1082                }
1083                _ => {
1084                    self.row_f32(r, scratch);
1085                    for (o, v) in out.iter_mut().zip(scratch.iter()) {
1086                        *o += w * v;
1087                    }
1088                }
1089            },
1090        }
1091    }
1092
1093    /// Dot of row `r` with `x` (gate/up active-neuron path). Reads only
1094    /// row `r` from the mmap — no full dequant. q4/vbit dequant the row
1095    /// into `scratch` first (rare for active-FFN weights).
1096    pub fn row_dot(&self, r: usize, x: &[f32], scratch: &mut [f32]) -> f32 {
1097        let cols = self.cols();
1098        match self {
1099            Self::F32 { data, .. } => {
1100                let row = &data[r * cols..(r + 1) * cols];
1101                row.iter().zip(x).map(|(w, v)| w * v).sum()
1102            }
1103            Self::Mapped {
1104                model,
1105                idx,
1106                dtype,
1107                row_scale,
1108                col_field,
1109                ..
1110            } => {
1111                let prism_forward =
1112                    crate::prism::is_forward_weight(model, &model.tensors[*idx].name);
1113                if prism_forward {
1114                    let transformed = crate::prism::forward(model, &x[..cols]);
1115                    let gpr = cols / GROUP_SIZE;
1116                    match dtype {
1117                        TensorDtype::Q2TiledP => {
1118                            let v = Q4tpView::new_q2(self.quant_bytes(), self.rows(), cols);
1119                            let mut sc = vec![0f32; gpr];
1120                            v.scales_into(r, gpr, &mut sc);
1121                            if crate::prism::is_affine_target(model, &model.tensors[*idx].name) {
1122                                return q2tp_affine_row_exact(v.nib, r, gpr, &transformed, &sc);
1123                            }
1124                            return q2tp_row_exact(v.nib, r, gpr, &transformed, &sc);
1125                        }
1126                        _ => {
1127                            self.row_f32(r, scratch);
1128                            return scratch.iter().zip(&transformed).map(|(w, v)| w * v).sum();
1129                        }
1130                    }
1131                }
1132                match dtype {
1133                    TensorDtype::Q8Row => {
1134                        let q = &self.quant_bytes()[r * cols..(r + 1) * cols];
1135                        dot_i8_f32(q, x) * row_scale[r]
1136                    }
1137                    TensorDtype::Q8_2f => {
1138                        let q = &self.quant_bytes()[r * cols..(r + 1) * cols];
1139                        dot_i8_col_f32(q, x, col_field) * row_scale[r]
1140                    }
1141                    _ => {
1142                        self.row_f32(r, scratch);
1143                        scratch.iter().zip(x).map(|(w, v)| w * v).sum()
1144                    }
1145                }
1146            }
1147        }
1148    }
1149
1150    /// `out = W · x` (row-major). F32 delegates to the historical
1151    /// bit-exact path; Mapped runs the fused int8 kernel.
1152    pub fn matvec(&self, x: &[f32], out: &mut [f32], pool: Option<&Pool>) {
1153        match self {
1154            // NOTE: `out.len()` DRIVES this arm — it computes that many rows,
1155            // and `x.len()` is the stride. A short `out` is legitimate here,
1156            // which is why the check below lives in the Mapped arm only.
1157            Self::F32 { data, .. } => {
1158                if !crate::f32_backend::matvec(data, x, out) {
1159                    matvec_rows(pool, data, x, out);
1160                }
1161            }
1162            Self::Mapped {
1163                model,
1164                idx,
1165                dtype,
1166                rows,
1167                cols,
1168                row_scale,
1169                col_field,
1170                vbit_offsets,
1171                repack,
1172            } => {
1173                let _ = (model, idx);
1174                // Every kernel below writes `rows` entries through a raw
1175                // pointer, so a short `out` is an out-of-bounds WRITE, not a
1176                // wrong answer: it scribbles on the allocator's metadata and
1177                // the process aborts much later, somewhere innocent
1178                // (`double free or corruption`, `corrupted double-linked
1179                // list`). The debug_assert two of the kernels carried is
1180                // compiled out of the release — exactly the build where it
1181                // matters. Fail here instead, while the caller is still on
1182                // the stack to be named.
1183                assert!(
1184                    out.len() >= *rows && x.len() >= *cols,
1185                    "matvec {rows}x{cols}: out {} (need {rows}), x {} (need {cols})",
1186                    out.len(),
1187                    x.len(),
1188                );
1189                let prism_forward =
1190                    crate::prism::is_forward_weight(model, &model.tensors[*idx].name);
1191                if *dtype == TensorDtype::Q2TiledP
1192                    && std::env::var("CMF_Q2TP_TRACE").as_deref() == Ok("1")
1193                {
1194                    use std::sync::atomic::{AtomicUsize, Ordering};
1195                    static N: AtomicUsize = AtomicUsize::new(0);
1196                    let n = N.fetch_add(1, Ordering::Relaxed);
1197                    if n < 128 {
1198                        eprintln!(
1199                            "q2tp-dispatch #{n} name={} prism={} rows={} cols={} gpu={} optin={} layer={}",
1200                            model.tensors[*idx].name,
1201                            prism_forward,
1202                            rows,
1203                            cols,
1204                            crate::gpu::enabled_here(),
1205                            crate::gpu::q2tp_gpu_opt_in(),
1206                            crate::gpu::cur_layer(),
1207                        );
1208                    }
1209                }
1210                // Prism stores every manifest-listed forward matrix in the
1211                // signed-Hadamard basis.  The q2tp WGSL path receives that
1212                // transformed vector and an explicit affine bit; codecs
1213                // without a descriptor-aware kernel remain on CPU below.
1214                if prism_forward {
1215                    let transformed = crate::prism::forward(model, &x[..*cols]);
1216                    match dtype {
1217                        TensorDtype::Q4Block => {
1218                            q4matvec(self.quant_bytes(), &transformed, *rows, *cols, out, pool)
1219                        }
1220                        TensorDtype::Q4Tiled => {
1221                            q4t_matvec(self.quant_bytes(), &transformed, *rows, *cols, out, pool)
1222                        }
1223                        TensorDtype::Q4TiledP => {
1224                            q4tp_matvec(self.quant_bytes(), &transformed, *rows, *cols, out, pool)
1225                        }
1226                        TensorDtype::Q2TiledP => {
1227                            let affine =
1228                                crate::prism::is_affine_target(model, &model.tensors[*idx].name);
1229                            if *rows * *cols >= 8_388_608
1230                                && crate::gpu::enabled_here()
1231                                && crate::gpu::q2tp_gpu_opt_in()
1232                            {
1233                                let gpu_ok = if affine {
1234                                    crate::gpu::q2tp_affine_matvec(
1235                                        model,
1236                                        *idx,
1237                                        &transformed,
1238                                        *rows,
1239                                        *cols,
1240                                        out,
1241                                    )
1242                                } else {
1243                                    crate::gpu::q2tp_matvec(
1244                                        model,
1245                                        *idx,
1246                                        &transformed,
1247                                        *rows,
1248                                        *cols,
1249                                        out,
1250                                    )
1251                                };
1252                                if gpu_ok {
1253                                    return;
1254                                }
1255                            }
1256                            if affine {
1257                                q2tp_affine_matvec(
1258                                    self.quant_bytes(),
1259                                    &transformed,
1260                                    *rows,
1261                                    *cols,
1262                                    out,
1263                                    pool,
1264                                )
1265                            } else {
1266                                q2tp_matvec(
1267                                    self.quant_bytes(),
1268                                    &transformed,
1269                                    *rows,
1270                                    *cols,
1271                                    out,
1272                                    pool,
1273                                )
1274                            }
1275                        }
1276                        TensorDtype::Q1 => {
1277                            q1_matvec(self.quant_bytes(), &transformed, *rows, *cols, out, pool)
1278                        }
1279                        TensorDtype::Q1T => {
1280                            q1t_matvec(self.quant_bytes(), &transformed, *rows, *cols, out, pool)
1281                        }
1282                        TensorDtype::Vbit | TensorDtype::VbitRo => vbitmatvec(
1283                            self.quant_bytes(),
1284                            vbit_offsets,
1285                            &transformed,
1286                            *rows,
1287                            *cols,
1288                            out,
1289                            pool,
1290                        ),
1291                        TensorDtype::Q8Row | TensorDtype::Q8_2f => qmatvec(
1292                            self.quant_bytes(),
1293                            repack,
1294                            row_scale,
1295                            &transformed,
1296                            col_field,
1297                            *dtype,
1298                            *rows,
1299                            *cols,
1300                            out,
1301                            pool,
1302                        ),
1303                        _ => unreachable!("unsupported mapped Prism dtype {dtype:?}"),
1304                    }
1305                    return;
1306                }
1307                if *dtype == TensorDtype::Q4Block {
1308                    // GPU route (wgpu q4b kernel) for large q4_block matvecs —
1309                    // gives NVIDIA/AMD/Intel q4 models a GPU path. Probe keeps
1310                    // the winner; Metal returns false → the CPU kernel below.
1311                    if *rows * *cols >= 8_388_608 && crate::gpu::enabled_here() {
1312                        let t0 = std::time::Instant::now();
1313                        match crate::gpu::probe_arm(crate::gpu::OpClass::Matvec) {
1314                            crate::gpu::ProbeArm::Gpu => {
1315                                if crate::gpu::q4b_matvec(model, *idx, x, *rows, *cols, out) {
1316                                    crate::gpu::probe_record(
1317                                        crate::gpu::OpClass::Matvec,
1318                                        true,
1319                                        t0.elapsed(),
1320                                    );
1321                                    return;
1322                                }
1323                            }
1324                            crate::gpu::ProbeArm::CpuTimed => {
1325                                q4matvec(self.quant_bytes(), x, *rows, *cols, out, pool);
1326                                crate::gpu::probe_record(
1327                                    crate::gpu::OpClass::Matvec,
1328                                    false,
1329                                    t0.elapsed(),
1330                                );
1331                                return;
1332                            }
1333                            crate::gpu::ProbeArm::Cpu => {}
1334                        }
1335                    }
1336                    q4matvec(self.quant_bytes(), x, *rows, *cols, out, pool);
1337                    return;
1338                }
1339                if *dtype == TensorDtype::Q4Tiled {
1340                    // GPU route for large q4t matvecs — the lm_head class,
1341                    // same shape as the q4tp arm below. The probe keeps the
1342                    // winner; a backend without the kernel refuses and the
1343                    // CPU path stays.
1344                    if *rows * *cols >= 8_388_608 && crate::gpu::enabled_here() {
1345                        let t0 = std::time::Instant::now();
1346                        let cls = crate::gpu::matvec_class(*rows, *cols);
1347                        match crate::gpu::probe_arm(cls) {
1348                            crate::gpu::ProbeArm::Gpu => {
1349                                if crate::gpu::q4t_matvec(model, *idx, x, *rows, *cols, out) {
1350                                    crate::gpu::probe_record(cls, true, t0.elapsed());
1351                                    return;
1352                                }
1353                            }
1354                            crate::gpu::ProbeArm::CpuTimed => {
1355                                q4t_matvec(self.quant_bytes(), x, *rows, *cols, out, pool);
1356                                crate::gpu::probe_record(cls, false, t0.elapsed());
1357                                return;
1358                            }
1359                            crate::gpu::ProbeArm::Cpu => {}
1360                        }
1361                    }
1362                    q4t_matvec(self.quant_bytes(), x, *rows, *cols, out, pool);
1363                    return;
1364                }
1365                if *dtype == TensorDtype::Q4TiledP {
1366                    // GPU route for large q4tp matvecs — the lm_head class.
1367                    // On a q4tp checkpoint the head is the biggest single
1368                    // host matvec left in the decode step, and the batched
1369                    // kernel at b=1 already exists on both backends. Probe
1370                    // keeps the winner, same as q4_block above.
1371                    if *rows * *cols >= 8_388_608 && crate::gpu::enabled_here() {
1372                        let t0 = std::time::Instant::now();
1373                        let cls = crate::gpu::matvec_class(*rows, *cols);
1374                        match crate::gpu::probe_arm(cls) {
1375                            crate::gpu::ProbeArm::Gpu => {
1376                                if crate::gpu::q4tp_matvec(model, *idx, x, *rows, *cols, out) {
1377                                    crate::gpu::probe_record(cls, true, t0.elapsed());
1378                                    return;
1379                                }
1380                            }
1381                            crate::gpu::ProbeArm::CpuTimed => {
1382                                q4tp_matvec(self.quant_bytes(), x, *rows, *cols, out, pool);
1383                                crate::gpu::probe_record(cls, false, t0.elapsed());
1384                                return;
1385                            }
1386                            crate::gpu::ProbeArm::Cpu => {}
1387                        }
1388                    }
1389                    q4tp_matvec(self.quant_bytes(), x, *rows, *cols, out, pool);
1390                    return;
1391                }
1392                if *dtype == TensorDtype::Q2TiledP {
1393                    q2tp_matvec(self.quant_bytes(), x, *rows, *cols, out, pool);
1394                    return;
1395                }
1396                if *dtype == TensorDtype::Q1 {
1397                    // GPU route for large q1 matvecs (out_proj / lm_head
1398                    // class): the CPU q1 kernel is load-port-bound at
1399                    // ~4 GB/s/core, the GPU one is bandwidth-bound — the
1400                    // probe measures both arms and keeps the winner.
1401                    if *rows * *cols >= 8_388_608 && crate::gpu::enabled_here() {
1402                        let t0 = std::time::Instant::now();
1403                        let arm = if crate::gpu::q1_force() {
1404                            crate::gpu::ProbeArm::Gpu
1405                        } else {
1406                            crate::gpu::probe_arm(crate::gpu::OpClass::Matvec)
1407                        };
1408                        match arm {
1409                            crate::gpu::ProbeArm::Gpu => {
1410                                if crate::gpu::q1_matvec(model, *idx, x, *rows, *cols, out) {
1411                                    crate::gpu::probe_record(
1412                                        crate::gpu::OpClass::Matvec,
1413                                        true,
1414                                        t0.elapsed(),
1415                                    );
1416                                    return;
1417                                }
1418                            }
1419                            crate::gpu::ProbeArm::CpuTimed => {
1420                                q1_matvec(self.quant_bytes(), x, *rows, *cols, out, pool);
1421                                crate::gpu::probe_record(
1422                                    crate::gpu::OpClass::Matvec,
1423                                    false,
1424                                    t0.elapsed(),
1425                                );
1426                                return;
1427                            }
1428                            crate::gpu::ProbeArm::Cpu => {}
1429                        }
1430                    }
1431                    q1_matvec(self.quant_bytes(), x, *rows, *cols, out, pool);
1432                    return;
1433                }
1434                if *dtype == TensorDtype::Q1T {
1435                    // GPU route for large q1t matvecs: the ternary BASE dot runs
1436                    // on the GPU (load-port-bound on CPU, like q1), then the
1437                    // sparse overlay is added on the CPU. Probe keeps the winner.
1438                    if *rows * *cols >= 8_388_608 && crate::gpu::enabled_here() {
1439                        let t0 = std::time::Instant::now();
1440                        match crate::gpu::probe_arm(crate::gpu::OpClass::Matvec) {
1441                            crate::gpu::ProbeArm::Gpu => {
1442                                if crate::gpu::q1t_matvec(model, *idx, x, *rows, *cols, out) {
1443                                    q1t_add_overlay(self.quant_bytes(), x, *rows, *cols, out, pool);
1444                                    crate::gpu::probe_record(
1445                                        crate::gpu::OpClass::Matvec,
1446                                        true,
1447                                        t0.elapsed(),
1448                                    );
1449                                    return;
1450                                }
1451                            }
1452                            crate::gpu::ProbeArm::CpuTimed => {
1453                                q1t_matvec(self.quant_bytes(), x, *rows, *cols, out, pool);
1454                                crate::gpu::probe_record(
1455                                    crate::gpu::OpClass::Matvec,
1456                                    false,
1457                                    t0.elapsed(),
1458                                );
1459                                return;
1460                            }
1461                            crate::gpu::ProbeArm::Cpu => {}
1462                        }
1463                    }
1464                    q1t_matvec(self.quant_bytes(), x, *rows, *cols, out, pool);
1465                    return;
1466                }
1467                if matches!(dtype, TensorDtype::Vbit | TensorDtype::VbitRo) {
1468                    vbitmatvec(self.quant_bytes(), vbit_offsets, x, *rows, *cols, out, pool);
1469                    return;
1470                }
1471                let xs = prescale(x, col_field, *dtype);
1472                // D5: large q8 matrices (lm_head-class) — hybrid
1473                // CPU∥GPU: split the rows, both sides compute
1474                // SIMULTANEOUSLY (same math, shared prescale).
1475                // GPU share: CMF_GPU_SPLIT (0..1, default 0.5).
1476                if *rows >= crate::gpu::min_rows()
1477                    && matches!(dtype, TensorDtype::Q8Row | TensorDtype::Q8_2f)
1478                    && gpu_lmhead_enabled()
1479                    && crate::gpu::enabled_here()
1480                {
1481                    // Runtime probe: alternate the hybrid against the
1482                    // pure-CPU matvec, keep whichever is faster HERE.
1483                    let t0 = std::time::Instant::now();
1484                    match crate::gpu::probe_arm(crate::gpu::OpClass::Matvec) {
1485                        crate::gpu::ProbeArm::Gpu => {}
1486                        crate::gpu::ProbeArm::CpuTimed => {
1487                            qmatvec(
1488                                self.quant_bytes(),
1489                                repack,
1490                                row_scale,
1491                                x,
1492                                col_field,
1493                                *dtype,
1494                                *rows,
1495                                *cols,
1496                                out,
1497                                pool,
1498                            );
1499                            crate::gpu::probe_record(
1500                                crate::gpu::OpClass::Matvec,
1501                                false,
1502                                t0.elapsed(),
1503                            );
1504                            return;
1505                        }
1506                        crate::gpu::ProbeArm::Cpu => {
1507                            qmatvec(
1508                                self.quant_bytes(),
1509                                repack,
1510                                row_scale,
1511                                x,
1512                                col_field,
1513                                *dtype,
1514                                *rows,
1515                                *cols,
1516                                out,
1517                                pool,
1518                            );
1519                            return;
1520                        }
1521                    }
1522                    let frac = gpu_split_frac();
1523                    let cpu_rows = ((*rows as f32) * (1.0 - frac)) as usize;
1524                    let (out_cpu, out_gpu) = out.split_at_mut(cpu_rows);
1525                    let bytes = self.quant_bytes();
1526                    let ok = std::thread::scope(|sc| {
1527                        let g = sc.spawn(|| {
1528                            crate::gpu::q8_matvec_range(
1529                                model,
1530                                *idx,
1531                                cpu_rows,
1532                                &row_scale[cpu_rows..],
1533                                &xs,
1534                                *rows - cpu_rows,
1535                                *cols,
1536                                out_gpu,
1537                            )
1538                        });
1539                        if cpu_rows > 0 {
1540                            // Repack prefix covers the full groups of the
1541                            // CPU half (the split starts at row 0).
1542                            let rep_cpu = if repack.is_empty() {
1543                                &[][..]
1544                            } else {
1545                                &repack[..(cpu_rows / 4) * 4 * *cols]
1546                            };
1547                            qmatvec(
1548                                &bytes[..cpu_rows * *cols],
1549                                rep_cpu,
1550                                &row_scale[..cpu_rows],
1551                                x,
1552                                col_field,
1553                                *dtype,
1554                                cpu_rows,
1555                                *cols,
1556                                out_cpu,
1557                                pool,
1558                            );
1559                        }
1560                        g.join().unwrap_or(false)
1561                    });
1562                    if ok {
1563                        crate::gpu::probe_record(crate::gpu::OpClass::Matvec, true, t0.elapsed());
1564                        return;
1565                    }
1566                    // GPU failed — CPU finishes its half (rows rebased —
1567                    // group offsets don't line up, mmap layout only).
1568                    qmatvec(
1569                        &bytes[cpu_rows * *cols..(*rows) * *cols],
1570                        &[],
1571                        &row_scale[cpu_rows..],
1572                        x,
1573                        col_field,
1574                        *dtype,
1575                        *rows - cpu_rows,
1576                        *cols,
1577                        out_gpu,
1578                        pool,
1579                    );
1580                    return;
1581                }
1582                qmatvec(
1583                    self.quant_bytes(),
1584                    repack,
1585                    row_scale,
1586                    x,
1587                    col_field,
1588                    *dtype,
1589                    *rows,
1590                    *cols,
1591                    out,
1592                    pool,
1593                );
1594            }
1595        }
1596    }
1597
1598    /// Fused two-input matvec (MTP verify pair): weights streamed once.
1599    pub fn matvec2(
1600        &self,
1601        x1: &[f32],
1602        x2: &[f32],
1603        o1: &mut [f32],
1604        o2: &mut [f32],
1605        pool: Option<&Pool>,
1606    ) {
1607        match self {
1608            Self::F32 { data, .. } => matvec_rows2(pool, data, x1, x2, o1, o2),
1609            Self::Mapped {
1610                model,
1611                idx,
1612                dtype,
1613                rows,
1614                cols,
1615                row_scale,
1616                col_field,
1617                vbit_offsets,
1618                ..
1619            } => {
1620                if crate::prism::is_forward_weight(model, &model.tensors[*idx].name) {
1621                    let tx1 = crate::prism::forward(model, &x1[..*cols]);
1622                    let tx2 = crate::prism::forward(model, &x2[..*cols]);
1623                    match dtype {
1624                        TensorDtype::Q4Block => {
1625                            q4matvec2(self.quant_bytes(), &tx1, &tx2, *rows, *cols, o1, o2, pool)
1626                        }
1627                        TensorDtype::Q4Tiled => {
1628                            q4t_matvec2(self.quant_bytes(), &tx1, &tx2, *rows, *cols, o1, o2, pool)
1629                        }
1630                        TensorDtype::Q4TiledP => {
1631                            q4tp_matvec2(self.quant_bytes(), &tx1, &tx2, *rows, *cols, o1, o2, pool)
1632                        }
1633                        TensorDtype::Q2TiledP => {
1634                            if crate::prism::is_affine_target(model, &model.tensors[*idx].name) {
1635                                q2tp_affine_matvec2(
1636                                    self.quant_bytes(),
1637                                    &tx1,
1638                                    &tx2,
1639                                    *rows,
1640                                    *cols,
1641                                    o1,
1642                                    o2,
1643                                    pool,
1644                                )
1645                            } else {
1646                                q2tp_matvec2(
1647                                    self.quant_bytes(),
1648                                    &tx1,
1649                                    &tx2,
1650                                    *rows,
1651                                    *cols,
1652                                    o1,
1653                                    o2,
1654                                    pool,
1655                                )
1656                            }
1657                        }
1658                        TensorDtype::Q1 => {
1659                            q1_matvec2(self.quant_bytes(), &tx1, &tx2, *rows, *cols, o1, o2, pool)
1660                        }
1661                        TensorDtype::Q1T => {
1662                            q1t_matvec2(self.quant_bytes(), &tx1, &tx2, *rows, *cols, o1, o2, pool)
1663                        }
1664                        TensorDtype::Vbit | TensorDtype::VbitRo => vbitmatvec2(
1665                            self.quant_bytes(),
1666                            vbit_offsets,
1667                            &tx1,
1668                            &tx2,
1669                            *rows,
1670                            *cols,
1671                            o1,
1672                            o2,
1673                            pool,
1674                        ),
1675                        TensorDtype::Q8Row | TensorDtype::Q8_2f => qmatvec2(
1676                            self.quant_bytes(),
1677                            row_scale,
1678                            &tx1,
1679                            &tx2,
1680                            col_field,
1681                            *dtype,
1682                            *rows,
1683                            *cols,
1684                            o1,
1685                            o2,
1686                            pool,
1687                        ),
1688                        _ => unreachable!("unsupported mapped Prism dtype {dtype:?}"),
1689                    }
1690                    return;
1691                }
1692                if *dtype == TensorDtype::Q4Block {
1693                    q4matvec2(self.quant_bytes(), x1, x2, *rows, *cols, o1, o2, pool);
1694                    return;
1695                }
1696                if *dtype == TensorDtype::Q4Tiled {
1697                    q4t_matvec2(self.quant_bytes(), x1, x2, *rows, *cols, o1, o2, pool);
1698                    return;
1699                }
1700                if *dtype == TensorDtype::Q4TiledP {
1701                    q4tp_matvec2(self.quant_bytes(), x1, x2, *rows, *cols, o1, o2, pool);
1702                    return;
1703                }
1704                if *dtype == TensorDtype::Q2TiledP {
1705                    q2tp_matvec2(self.quant_bytes(), x1, x2, *rows, *cols, o1, o2, pool);
1706                    return;
1707                }
1708                if *dtype == TensorDtype::Q1 {
1709                    q1_matvec2(self.quant_bytes(), x1, x2, *rows, *cols, o1, o2, pool);
1710                    return;
1711                }
1712                if *dtype == TensorDtype::Q1T {
1713                    // Fused ternary pair: one row pass, the register
1714                    // unpack shared across both streams on ARM. (Q1T
1715                    // lacks a row_scale array — scales live inline in
1716                    // the tiles — so it must not fall through to the
1717                    // q8 qmatvec2 below.)
1718                    q1t_matvec2(self.quant_bytes(), x1, x2, *rows, *cols, o1, o2, pool);
1719                    return;
1720                }
1721                if matches!(dtype, TensorDtype::Vbit | TensorDtype::VbitRo) {
1722                    vbitmatvec2(
1723                        self.quant_bytes(),
1724                        vbit_offsets,
1725                        x1,
1726                        x2,
1727                        *rows,
1728                        *cols,
1729                        o1,
1730                        o2,
1731                        pool,
1732                    );
1733                    return;
1734                }
1735                qmatvec2(
1736                    self.quant_bytes(),
1737                    row_scale,
1738                    x1,
1739                    x2,
1740                    col_field,
1741                    *dtype,
1742                    *rows,
1743                    *cols,
1744                    o1,
1745                    o2,
1746                    pool,
1747                );
1748            }
1749        }
1750    }
1751}
1752
1753impl QTensor {
1754    /// Batched matvec (prefill-GEMM): xs — row-major [b, cols],
1755    /// out — row-major [b, rows]. Element-wise semantics are IDENTICAL
1756    /// to b matvec calls (same dot kernels in the same order); the win —
1757    /// the weight row streams from DRAM once per batch, not b times.
1758    /// `(model, index)` when this is a memory-mapped q4tp tensor — the
1759    /// identity a device-resident chain needs to hand `tp_matmat` the
1760    /// weight without going through this struct's own dispatch.
1761    pub fn q4tp_mapped(&self) -> Option<(&std::sync::Arc<CmfModel>, usize)> {
1762        if self.has_prism_contract() {
1763            return None;
1764        }
1765        match self {
1766            Self::Mapped {
1767                model, idx, dtype, ..
1768            } if *dtype == TensorDtype::Q4TiledP => Some((model, *idx)),
1769            _ => None,
1770        }
1771    }
1772
1773    pub fn matmat(&self, xs_all: &[f32], b: usize, out: &mut [f32], pool: Option<&Pool>) {
1774        let cols = self.cols();
1775        let rows = self.rows();
1776        debug_assert_eq!(xs_all.len(), b * cols);
1777        debug_assert_eq!(out.len(), b * rows);
1778        let _prof = crate::cpuprof::time(crate::cpuprof::Slot::Matmat);
1779        // GPTQ calibration: fold this layer's inputs into its Hessian. Only
1780        // Mapped tensors carry a directory name; the check is a relaxed
1781        // atomic load, free when not calibrating.
1782        if crate::gptq_capture::capturing() {
1783            if let Self::Mapped { model, idx, .. } = self {
1784                crate::gptq_capture::accumulate(&model.tensors[*idx].name, xs_all, b, cols);
1785            }
1786        }
1787        match self {
1788            Self::F32 { data, .. } => {
1789                if crate::f32_backend::matmat(data, xs_all, b, rows, cols, out) {
1790                    return;
1791                }
1792                let out_addr = SendMut(out.as_mut_ptr());
1793                // Over the batch, four positions a row at a time: four
1794                // independent chains, each still summing j = 0..cols in the
1795                // scalar order (bit-identical to the one-at-a-time loop and
1796                // to `matvec_rows`). The batch is the wide axis of a prefill;
1797                // the rows may be few (Spark-X2.5's 16-row g_proj spent
1798                // 17 ms a layer here split over its rows).
1799                let run = |start: usize, end: usize| {
1800                    let mut bi = start;
1801                    while bi < end {
1802                        let n = (end - bi).min(4);
1803                        for o in 0..rows {
1804                            let row = &data[o * cols..(o + 1) * cols];
1805                            if n == 4 {
1806                                let x = |k: usize| &xs_all[(bi + k) * cols..(bi + k + 1) * cols];
1807                                let (x0, x1, x2, x3) = (x(0), x(1), x(2), x(3));
1808                                let (mut a0, mut a1, mut a2, mut a3) = (0f32, 0f32, 0f32, 0f32);
1809                                for j in 0..cols {
1810                                    let w = row[j];
1811                                    a0 += w * x0[j];
1812                                    a1 += w * x1[j];
1813                                    a2 += w * x2[j];
1814                                    a3 += w * x3[j];
1815                                }
1816                                for (k, acc) in [a0, a1, a2, a3].into_iter().enumerate() {
1817                                    unsafe { *out_addr.at((bi + k) * rows + o) = acc };
1818                                }
1819                            } else {
1820                                for k in 0..n {
1821                                    let x = &xs_all[(bi + k) * cols..(bi + k + 1) * cols];
1822                                    let mut acc = 0f32;
1823                                    for j in 0..cols {
1824                                        acc += row[j] * x[j];
1825                                    }
1826                                    unsafe { *out_addr.at((bi + k) * rows + o) = acc };
1827                                }
1828                            }
1829                        }
1830                        bi += n;
1831                    }
1832                };
1833                dispatch_rows(pool, b, &run);
1834            }
1835            Self::Mapped {
1836                model,
1837                idx,
1838                dtype,
1839                row_scale,
1840                col_field,
1841                vbit_offsets,
1842                ..
1843            } => {
1844                if crate::prism::is_forward_weight(model, &model.tensors[*idx].name) {
1845                    let mut transformed = Vec::with_capacity(xs_all.len());
1846                    for bi in 0..b {
1847                        transformed.extend_from_slice(&crate::prism::forward(
1848                            model,
1849                            &xs_all[bi * cols..(bi + 1) * cols],
1850                        ));
1851                    }
1852                    match dtype {
1853                        TensorDtype::Q4Block => {
1854                            q4matmat(self.quant_bytes(), &transformed, b, rows, cols, out, pool)
1855                        }
1856                        TensorDtype::Q4Tiled => {
1857                            q4t_matmat(self.quant_bytes(), &transformed, b, rows, cols, out, pool)
1858                        }
1859                        TensorDtype::Q4TiledP => {
1860                            q4tp_matmat(self.quant_bytes(), &transformed, b, rows, cols, out, pool)
1861                        }
1862                        TensorDtype::Q2TiledP => {
1863                            let affine =
1864                                crate::prism::is_affine_target(model, &model.tensors[*idx].name);
1865                            // Affine Prism Q2TP has a descriptor-aware GPU
1866                            // kernel for short/tail batches too.  Unlike the
1867                            // ordinary Q2TP path, don't force b<32 back to a
1868                            // scalar CPU matmat: prefill chunks and the final
1869                            // tail both need to stay on the tested GPU arm.
1870                            let gpu_batch_ok = if affine {
1871                                b >= 2
1872                            } else {
1873                                b >= 32 && b * rows * cols >= 128_000_000
1874                            };
1875                            if gpu_batch_ok
1876                                && cols % 32 == 0
1877                                && crate::gpu::enabled_here()
1878                                && crate::gpu::q2tp_gpu_opt_in()
1879                            {
1880                                let gpu_ok = if affine {
1881                                    crate::gpu::q2tp_affine_matmat(
1882                                        model,
1883                                        *idx,
1884                                        &transformed,
1885                                        b,
1886                                        rows,
1887                                        cols,
1888                                        out,
1889                                    )
1890                                } else {
1891                                    crate::gpu::q2tp_matmat(
1892                                        model,
1893                                        *idx,
1894                                        &transformed,
1895                                        b,
1896                                        rows,
1897                                        cols,
1898                                        out,
1899                                    )
1900                                };
1901                                if gpu_ok {
1902                                    return;
1903                                }
1904                            }
1905                            // A one-token Prism decode is the other short
1906                            // case.  Use the descriptor-aware matvec kernel
1907                            // before falling back to the exact CPU path.
1908                            if affine
1909                                && b == 1
1910                                && cols % 32 == 0
1911                                && crate::gpu::enabled_here()
1912                                && crate::gpu::q2tp_gpu_opt_in()
1913                                && crate::gpu::q2tp_affine_matvec(
1914                                    model,
1915                                    *idx,
1916                                    &transformed[..cols],
1917                                    rows,
1918                                    cols,
1919                                    &mut out[..rows],
1920                                )
1921                            {
1922                                return;
1923                            }
1924                            if affine {
1925                                q2tp_affine_matmat(
1926                                    self.quant_bytes(),
1927                                    &transformed,
1928                                    b,
1929                                    rows,
1930                                    cols,
1931                                    out,
1932                                    pool,
1933                                )
1934                            } else {
1935                                q2tp_matmat(
1936                                    self.quant_bytes(),
1937                                    &transformed,
1938                                    b,
1939                                    rows,
1940                                    cols,
1941                                    out,
1942                                    pool,
1943                                )
1944                            }
1945                        }
1946                        TensorDtype::Q1 => {
1947                            q1_matmat(self.quant_bytes(), &transformed, b, rows, cols, out, pool)
1948                        }
1949                        TensorDtype::Q1T => {
1950                            q1t_matmat(self.quant_bytes(), &transformed, b, rows, cols, out, pool)
1951                        }
1952                        TensorDtype::Vbit | TensorDtype::VbitRo => vbitmatmat(
1953                            self.quant_bytes(),
1954                            vbit_offsets,
1955                            &transformed,
1956                            b,
1957                            rows,
1958                            cols,
1959                            out,
1960                            pool,
1961                        ),
1962                        TensorDtype::Q8Row | TensorDtype::Q8_2f => {
1963                            let pre: Vec<std::borrow::Cow<'_, [f32]>> = (0..b)
1964                                .map(|bi| {
1965                                    prescale(
1966                                        &transformed[bi * cols..(bi + 1) * cols],
1967                                        col_field,
1968                                        *dtype,
1969                                    )
1970                                })
1971                                .collect();
1972                            qmatmat(self.quant_bytes(), row_scale, &pre, rows, cols, out, pool)
1973                        }
1974                        _ => unreachable!("unsupported mapped Prism dtype {dtype:?}"),
1975                    }
1976                    return;
1977                }
1978                if *dtype == TensorDtype::Q4Block {
1979                    q4matmat(self.quant_bytes(), xs_all, b, rows, cols, out, pool);
1980                    return;
1981                }
1982                if *dtype == TensorDtype::Q4TiledP {
1983                    // GPU batched q4tp GEMM (dequant + f32nt mul_mm on the
1984                    // device); the probe keeps whichever beats the CPU arm.
1985                    // Narrow (prompt-encode) and wide (DiT) batches probe
1986                    // as separate classes — the regimes have opposite
1987                    // winners and one shared verdict locked the wrong arm.
1988                    // Kill switch (gpu::mm_kill): one grossly slow GPU op
1989                    // (a fair-condition op is ≤~100 ms even at 1024px)
1990                    // means the device is contended by another process
1991                    // (e.g. a simulator) — verdicts are per-process, so
1992                    // without the bail the whole render crawls behind
1993                    // someone else's queue.
1994                    // Row-exact batches stay on the host: the device GEMM
1995                    // is an f32 dequant-sgemm, not the host matvec's sum.
1996                    if b >= 32
1997                        && b * rows * cols >= 128_000_000
1998                        && cols % 32 == 0
1999                        && !row_exact()
2000                        && !crate::gpu::mm_killed()
2001                        && crate::gpu::enabled_here()
2002                    {
2003                        let class = if b >= 128 {
2004                            crate::gpu::OpClass::MatmatWide
2005                        } else {
2006                            crate::gpu::OpClass::Matmat
2007                        };
2008                        if let Self::Mapped { model, idx, .. } = self {
2009                            // In-process A/B (`CMF_MM_AB=1`). Three
2010                            // wall-clock A/Bs on a shared stand disagreed
2011                            // with each other by 25% on the same change,
2012                            // because the machine drifts between processes
2013                            // and interleaving whole renders does not fix
2014                            // that. Here both arms run back to back on the
2015                            // SAME data inside one call, so whatever the
2016                            // machine is doing, it does to both — and the
2017                            // disagreement between their outputs falls out
2018                            // for free. Doubles the work; a diagnostic,
2019                            // not a mode.
2020                            if crate::mm_ab::on() {
2021                                let mut g = vec![0f32; b * rows];
2022                                let t = std::time::Instant::now();
2023                                let took = crate::gpu::q4tp_matmat(
2024                                    model, *idx, xs_all, b, rows, cols, &mut g,
2025                                );
2026                                let dg = t.elapsed();
2027                                let t = std::time::Instant::now();
2028                                q4tp_matmat(self.quant_bytes(), xs_all, b, rows, cols, out, pool);
2029                                let dc = t.elapsed();
2030                                crate::mm_ab::record(b, rows, cols, took, dg, dc, &g, out);
2031                                return;
2032                            }
2033                            let t0 = std::time::Instant::now();
2034                            // A cold call takes the device arm: its sample
2035                            // is discarded either way, and the upload is
2036                            // what the next step needs.
2037                            let resident = crate::gpu::weight_is_resident(model, *idx);
2038                            match crate::gpu::probe_arm_cold_prefers_gpu(class, resident) {
2039                                crate::gpu::ProbeArm::Gpu => {
2040                                    if crate::gpu::q4tp_matmat(
2041                                        model, *idx, xs_all, b, rows, cols, out,
2042                                    ) {
2043                                        let el = t0.elapsed();
2044                                        // Work-proportional budget: ~8× the
2045                                        // fair-device estimate (+20 ms slack).
2046                                        // An absolute cap missed the worst
2047                                        // case — contended ops sit at
2048                                        // 100–240 ms each and still bury a
2049                                        // render whose fair op is 3–9 ms.
2050                                        // Cold ops (first PSO build, buffer
2051                                        // alloc) are exempt: a one-off
2052                                        // ~50 ms compile is not contention.
2053                                        let flops = 2.0 * b as f64 * rows as f64 * cols as f64;
2054                                        let budget = std::time::Duration::from_secs_f64(
2055                                            flops / 1.5e12 * 8.0 + 0.020,
2056                                        );
2057                                        crate::gpu::mm_budget_check(
2058                                            "q4tp matmat",
2059                                            el,
2060                                            budget,
2061                                            crate::gpu::probe_was_cold() || !resident,
2062                                        );
2063                                        crate::gpu::probe_record(class, true, el);
2064                                        return;
2065                                    }
2066                                }
2067                                crate::gpu::ProbeArm::CpuTimed => {
2068                                    q4tp_matmat(
2069                                        self.quant_bytes(),
2070                                        xs_all,
2071                                        b,
2072                                        rows,
2073                                        cols,
2074                                        out,
2075                                        pool,
2076                                    );
2077                                    crate::gpu::probe_record(class, false, t0.elapsed());
2078                                    return;
2079                                }
2080                                crate::gpu::ProbeArm::Cpu => {}
2081                            }
2082                        }
2083                    }
2084                    q4tp_matmat(self.quant_bytes(), xs_all, b, rows, cols, out, pool);
2085                    return;
2086                }
2087                if *dtype == TensorDtype::Q2TiledP {
2088                    // Same device arm as q4tp, behind the same probe:
2089                    // the planes differ, the dispatch does not. Without
2090                    // this a q2tp file ran its widest projections on the
2091                    // host while the 4-bit one had the card, which is a
2092                    // codec paying for its size twice.
2093                    if b >= 32
2094                        && b * rows * cols >= 128_000_000
2095                        && cols % 32 == 0
2096                        && !crate::gpu::mm_killed()
2097                        && crate::gpu::enabled_here()
2098                    {
2099                        let class = if b >= 128 {
2100                            crate::gpu::OpClass::MatmatWide
2101                        } else {
2102                            crate::gpu::OpClass::Matmat
2103                        };
2104                        if let Self::Mapped { model, idx, .. } = self {
2105                            let t0 = std::time::Instant::now();
2106                            match crate::gpu::probe_arm(class) {
2107                                crate::gpu::ProbeArm::Gpu => {
2108                                    if crate::gpu::q2tp_matmat(
2109                                        model, *idx, xs_all, b, rows, cols, out,
2110                                    ) {
2111                                        crate::gpu::probe_record(class, true, t0.elapsed());
2112                                        return;
2113                                    }
2114                                }
2115                                crate::gpu::ProbeArm::CpuTimed => {
2116                                    q2tp_matmat(
2117                                        self.quant_bytes(),
2118                                        xs_all,
2119                                        b,
2120                                        rows,
2121                                        cols,
2122                                        out,
2123                                        pool,
2124                                    );
2125                                    crate::gpu::probe_record(class, false, t0.elapsed());
2126                                    return;
2127                                }
2128                                crate::gpu::ProbeArm::Cpu => {}
2129                            }
2130                        }
2131                    }
2132                    // Without a host arm a q2tp tensor falls through to
2133                    // the q8 fallback, which reads it at one BYTE per
2134                    // weight — a 2x overrun that killed pool workers
2135                    // mid-prefill while the dispatcher waited forever.
2136                    q2tp_matmat(self.quant_bytes(), xs_all, b, rows, cols, out, pool);
2137                    return;
2138                }
2139                if *dtype == TensorDtype::Q4Tiled {
2140                    // GPU batched q4t GEMM (dequant + f32nt mul_mm on the
2141                    // device); the probe keeps whichever beats the CPU arm.
2142                    // Narrow (prompt-encode) and wide (DiT) batches probe
2143                    // as separate classes — the regimes have opposite
2144                    // winners and one shared verdict locked the wrong arm.
2145                    // Kill switch (gpu::mm_kill): one grossly slow GPU op
2146                    // (a fair-condition op is ≤~100 ms even at 1024px)
2147                    // means the device is contended by another process
2148                    // (e.g. a simulator) — verdicts are per-process, so
2149                    // without the bail the whole render crawls behind
2150                    // someone else's queue.
2151                    if b >= 32
2152                        && b * rows * cols >= 128_000_000
2153                        && cols % 32 == 0
2154                        && !crate::gpu::mm_killed()
2155                        && crate::gpu::enabled_here()
2156                    {
2157                        let class = if b >= 128 {
2158                            crate::gpu::OpClass::MatmatWide
2159                        } else {
2160                            crate::gpu::OpClass::Matmat
2161                        };
2162                        if let Self::Mapped { model, idx, .. } = self {
2163                            let t0 = std::time::Instant::now();
2164                            match crate::gpu::probe_arm(class) {
2165                                crate::gpu::ProbeArm::Gpu => {
2166                                    if crate::gpu::q4t_matmat(
2167                                        model, *idx, xs_all, b, rows, cols, out,
2168                                    ) {
2169                                        let el = t0.elapsed();
2170                                        // Work-proportional budget: ~8× the
2171                                        // fair-device estimate (+20 ms slack).
2172                                        // An absolute cap missed the worst
2173                                        // case — contended ops sit at
2174                                        // 100–240 ms each and still bury a
2175                                        // render whose fair op is 3–9 ms.
2176                                        // Cold ops (first PSO build, buffer
2177                                        // alloc) are exempt: a one-off
2178                                        // ~50 ms compile is not contention.
2179                                        let flops = 2.0 * b as f64 * rows as f64 * cols as f64;
2180                                        let budget = std::time::Duration::from_secs_f64(
2181                                            flops / 1.5e12 * 8.0 + 0.020,
2182                                        );
2183                                        crate::gpu::mm_budget_check(
2184                                            "q4t matmat",
2185                                            el,
2186                                            budget,
2187                                            crate::gpu::probe_was_cold(),
2188                                        );
2189                                        crate::gpu::probe_record(class, true, el);
2190                                        return;
2191                                    }
2192                                }
2193                                crate::gpu::ProbeArm::CpuTimed => {
2194                                    q4t_matmat(
2195                                        self.quant_bytes(),
2196                                        xs_all,
2197                                        b,
2198                                        rows,
2199                                        cols,
2200                                        out,
2201                                        pool,
2202                                    );
2203                                    crate::gpu::probe_record(class, false, t0.elapsed());
2204                                    return;
2205                                }
2206                                crate::gpu::ProbeArm::Cpu => {}
2207                            }
2208                        }
2209                    }
2210                    q4t_matmat(self.quant_bytes(), xs_all, b, rows, cols, out, pool);
2211                    return;
2212                }
2213                if *dtype == TensorDtype::Q1 {
2214                    // GPU batched q1 GEMM for wide prefill (q1_mul_mm on the
2215                    // device); the probe keeps whichever beats the CPU matmat.
2216                    if b >= 32
2217                        && b * rows * cols >= 128_000_000
2218                        && cols % 64 == 0
2219                        && crate::gpu::enabled_here()
2220                    {
2221                        if let Self::Mapped { model, idx, .. } = self {
2222                            let t0 = std::time::Instant::now();
2223                            match crate::gpu::probe_arm(crate::gpu::OpClass::Matmat) {
2224                                crate::gpu::ProbeArm::Gpu => {
2225                                    if crate::gpu::q1_matmat(
2226                                        model, *idx, xs_all, b, rows, cols, out,
2227                                    ) {
2228                                        crate::gpu::probe_record(
2229                                            crate::gpu::OpClass::Matmat,
2230                                            true,
2231                                            t0.elapsed(),
2232                                        );
2233                                        return;
2234                                    }
2235                                }
2236                                crate::gpu::ProbeArm::CpuTimed => {
2237                                    q1_matmat(self.quant_bytes(), xs_all, b, rows, cols, out, pool);
2238                                    crate::gpu::probe_record(
2239                                        crate::gpu::OpClass::Matmat,
2240                                        false,
2241                                        t0.elapsed(),
2242                                    );
2243                                    return;
2244                                }
2245                                crate::gpu::ProbeArm::Cpu => {}
2246                            }
2247                        }
2248                    }
2249                    q1_matmat(self.quant_bytes(), xs_all, b, rows, cols, out, pool);
2250                    return;
2251                }
2252                if *dtype == TensorDtype::Q1T {
2253                    // GPU batched GEMM for wide prefill (base + overlay on the
2254                    // device); probe keeps the winner vs the CPU matmat.
2255                    if b >= 32 && b * rows * cols >= 128_000_000 && crate::gpu::enabled_here() {
2256                        if let Self::Mapped { model, idx, .. } = self {
2257                            let t0 = std::time::Instant::now();
2258                            match crate::gpu::probe_arm(crate::gpu::OpClass::Matmat) {
2259                                crate::gpu::ProbeArm::Gpu => {
2260                                    if crate::gpu::q1t_matmat(
2261                                        model, *idx, xs_all, b, rows, cols, out,
2262                                    ) {
2263                                        crate::gpu::probe_record(
2264                                            crate::gpu::OpClass::Matmat,
2265                                            true,
2266                                            t0.elapsed(),
2267                                        );
2268                                        return;
2269                                    }
2270                                }
2271                                crate::gpu::ProbeArm::CpuTimed => {
2272                                    q1t_matmat(
2273                                        self.quant_bytes(),
2274                                        xs_all,
2275                                        b,
2276                                        rows,
2277                                        cols,
2278                                        out,
2279                                        pool,
2280                                    );
2281                                    crate::gpu::probe_record(
2282                                        crate::gpu::OpClass::Matmat,
2283                                        false,
2284                                        t0.elapsed(),
2285                                    );
2286                                    return;
2287                                }
2288                                crate::gpu::ProbeArm::Cpu => {}
2289                            }
2290                        }
2291                    }
2292                    q1t_matmat(self.quant_bytes(), xs_all, b, rows, cols, out, pool);
2293                    return;
2294                }
2295                if matches!(dtype, TensorDtype::Vbit | TensorDtype::VbitRo) {
2296                    vbitmatmat(
2297                        self.quant_bytes(),
2298                        vbit_offsets,
2299                        xs_all,
2300                        b,
2301                        rows,
2302                        cols,
2303                        out,
2304                        pool,
2305                    );
2306                    return;
2307                }
2308                let pre: Vec<std::borrow::Cow<'_, [f32]>> = (0..b)
2309                    .map(|bi| prescale(&xs_all[bi * cols..(bi + 1) * cols], col_field, *dtype))
2310                    .collect();
2311                // MiMo verification is a 2–4 row decode panel, not a wide
2312                // prompt GEMM. Keep q8 projections on the same device as
2313                // decode; the generic b>=8 gate otherwise silently moves
2314                // every projection back to CPU. The short wgpu matmat uses
2315                // the same 64-lane reduction as its single-token matvec.
2316                if row_exact()
2317                    && (1..=4).contains(&b)
2318                    && matches!(dtype, TensorDtype::Q8Row | TensorDtype::Q8_2f)
2319                    && crate::gpu::enabled_here()
2320                    && crate::gpu::wgpu_active()
2321                {
2322                    let flat: Vec<f32> = pre.iter().flat_map(|v| v.iter().copied()).collect();
2323                    if crate::gpu::q8_matmat(model, *idx, row_scale, &flat, b, rows, cols, out) {
2324                        return;
2325                    }
2326                }
2327                // D5: large prefill-batch GEMMs — on the GPU (threshold by
2328                // work volume: submission carries b×rows×cols MACs).
2329                // Runtime probe: the naive GEMM shader + sync readback
2330                // lose to the CPU GEMM on slow driver stacks — alternate
2331                // both arms and keep the winner.
2332                if b >= 8 && b * rows * cols >= 128_000_000 && crate::gpu::enabled_here() {
2333                    if let Self::Mapped { model, idx, .. } = self {
2334                        let t0 = std::time::Instant::now();
2335                        match crate::gpu::probe_arm(crate::gpu::OpClass::Matmat) {
2336                            crate::gpu::ProbeArm::Gpu
2337                                if crate::gpu::probe_deciding(crate::gpu::OpClass::Matmat)
2338                                    && !crate::gpu::q8_resident_or_upload(model, *idx) =>
2339                            {
2340                                // Cold weights during probing: the upload
2341                                // has started, the count runs on the CPU —
2342                                // the GPU arm samples on the next touch.
2343                                let q = self.quant_bytes();
2344                                qmatmat(q, row_scale, &pre, rows, cols, out, pool);
2345                                return;
2346                            }
2347                            crate::gpu::ProbeArm::Gpu => {
2348                                // The two-field codec's own entry first, as
2349                                // `device_matmat` takes it: the column field
2350                                // folds into the weight plane and the GEMM
2351                                // runs on the matrix units. The plain int8
2352                                // entry below is the scalar f32 GEMM — on a
2353                                // Spark-X2.5 4B prefill it ran q|k|v|o at
2354                                // host speed.
2355                                if *dtype == TensorDtype::Q8_2f
2356                                    && std::env::var("CMF_Q8_2F_DEV").as_deref() != Ok("0")
2357                                    && crate::gpu::q8_matmat_2f(
2358                                        model, *idx, row_scale, col_field, xs_all, b, rows,
2359                                        cols, out,
2360                                    )
2361                                {
2362                                    crate::gpu::probe_record(
2363                                        crate::gpu::OpClass::Matmat,
2364                                        true,
2365                                        t0.elapsed(),
2366                                    );
2367                                    return;
2368                                }
2369                                let flat: Vec<f32> =
2370                                    pre.iter().flat_map(|v| v.iter().copied()).collect();
2371                                if crate::gpu::q8_matmat(
2372                                    model, *idx, row_scale, &flat, b, rows, cols, out,
2373                                ) {
2374                                    crate::gpu::probe_record(
2375                                        crate::gpu::OpClass::Matmat,
2376                                        true,
2377                                        t0.elapsed(),
2378                                    );
2379                                    return;
2380                                }
2381                            }
2382                            crate::gpu::ProbeArm::CpuTimed => {
2383                                let q = self.quant_bytes();
2384                                qmatmat(q, row_scale, &pre, rows, cols, out, pool);
2385                                crate::gpu::probe_record(
2386                                    crate::gpu::OpClass::Matmat,
2387                                    false,
2388                                    t0.elapsed(),
2389                                );
2390                                return;
2391                            }
2392                            crate::gpu::ProbeArm::Cpu => {}
2393                        }
2394                    }
2395                }
2396                let q = self.quant_bytes();
2397                qmatmat(q, row_scale, &pre, rows, cols, out, pool);
2398            }
2399        }
2400    }
2401}
2402
2403impl QTensor {
2404    /// The device GEMM this tensor would take, run once on the caller's
2405    /// data — the startup parity probe's arm, and the one place that knows
2406    /// which entry point each codec has.
2407    ///
2408    /// It exists because the probe used to look for a `q4tp` weight by
2409    /// name AND dtype, and a container packed any other way was declared
2410    /// "host path" for the whole render even though its codec had a device
2411    /// GEMM of its own. A gate that only recognizes one codec is a gate
2412    /// that silently downgrades every other one.
2413    pub fn device_matmat(&self, xs: &[f32], b: usize, out: &mut [f32]) -> bool {
2414        let (rows, cols) = (self.rows(), self.cols());
2415        let Self::Mapped {
2416            model,
2417            idx,
2418            dtype,
2419            row_scale,
2420            col_field,
2421            ..
2422        } = self
2423        else {
2424            return false;
2425        };
2426        if crate::prism::has_contract(model) {
2427            return false;
2428        }
2429        match *dtype {
2430            TensorDtype::Q4TiledP => crate::gpu::q4tp_matmat(model, *idx, xs, b, rows, cols, out),
2431            // The two-field codec folds its column field into the
2432            // activation, which leaves a plain per-row int8 GEMM — the
2433            // same kernel `q8_row` uses, on both backends.
2434            TensorDtype::Q8Row | TensorDtype::Q8_2f => {
2435                // The field belongs to the weight; only a backend that cannot
2436                // apply it there makes a scaled copy of the activation.
2437                if *dtype == TensorDtype::Q8_2f
2438                    && std::env::var("CMF_Q8_2F_DEV").as_deref() != Ok("0")
2439                    && crate::gpu::q8_matmat_2f(
2440                        model, *idx, row_scale, col_field, xs, b, rows, cols, out,
2441                    )
2442                {
2443                    return true;
2444                }
2445                let flat: Vec<f32> = (0..b)
2446                    .flat_map(|bi| {
2447                        prescale(&xs[bi * cols..(bi + 1) * cols], col_field, *dtype).into_owned()
2448                    })
2449                    .collect();
2450                crate::gpu::q8_matmat(model, *idx, row_scale, &flat, b, rows, cols, out)
2451            }
2452            _ => false,
2453        }
2454    }
2455
2456    /// Multi-matrix job (roadmap §3 P0): N tensors sharing one input
2457    /// run under a SINGLE pool dispatch — QKV or gate+up cost one
2458    /// barrier instead of N. Per-row math is the exact same kernel as
2459    /// `matvec` (bit-identical outputs); only the dispatch is fused.
2460    /// Falls back to N sequential matvecs when the set is not a uniform
2461    /// q8-family/F32 group or there is no pool.
2462    pub fn matvec_many<const N: usize>(
2463        ts: [&QTensor; N],
2464        x: &[f32],
2465        mut outs: [&mut [f32]; N],
2466        pool: Option<&Pool>,
2467    ) {
2468        let total_rows: usize = ts.iter().map(|t| t.rows()).sum();
2469        if ts.iter().any(|t| t.has_prism_contract()) {
2470            // The fused range kernels have no transform descriptor.  Let
2471            // each tensor's ordinary matvec dispatch perform the explicit
2472            // signed FWHT (and retain CPU fallback for mixed q2tp/q4tp).
2473            for (t, o) in ts.iter().zip(outs.iter_mut()) {
2474                t.matvec(x, o, pool);
2475            }
2476            return;
2477        }
2478        let uniform_q8 = ts.iter().all(|t| {
2479            matches!(
2480                t,
2481                Self::Mapped {
2482                    dtype: TensorDtype::Q8Row | TensorDtype::Q8_2f,
2483                    ..
2484                }
2485            )
2486        });
2487        let uniform_f32 = ts.iter().all(|t| matches!(t, Self::F32 { .. }));
2488        if uniform_f32 && crate::f32_backend::active() {
2489            for (t, o) in ts.iter().zip(outs.iter_mut()) {
2490                t.matvec(x, o, pool);
2491            }
2492            return;
2493        }
2494        let uniform_q4 = ts.iter().all(|t| {
2495            matches!(
2496                t,
2497                Self::Mapped {
2498                    dtype: TensorDtype::Q4Block,
2499                    ..
2500                }
2501            )
2502        });
2503        let uniform_vbit = ts.iter().all(|t| {
2504            matches!(
2505                t,
2506                Self::Mapped {
2507                    dtype: TensorDtype::Vbit | TensorDtype::VbitRo,
2508                    ..
2509                }
2510            )
2511        });
2512        let uniform_q1 = ts.iter().all(|t| {
2513            matches!(
2514                t,
2515                Self::Mapped {
2516                    dtype: TensorDtype::Q1,
2517                    ..
2518                }
2519            )
2520        });
2521        let uniform_q1t = ts.iter().all(|t| {
2522            matches!(
2523                t,
2524                Self::Mapped {
2525                    dtype: TensorDtype::Q1T,
2526                    ..
2527                }
2528            )
2529        });
2530        // q4tp is the skeleton dtype of the big MoE files, and without an arm
2531        // here every projection that shares an input paid its own pool
2532        // barrier: DeepSeek-V4's attention step alone hands this function
2533        // wq_a, wkv and both compressors' pairs off the same hidden state.
2534        let uniform_q4tp = ts.iter().all(|t| {
2535            matches!(
2536                t,
2537                Self::Mapped {
2538                    dtype: TensorDtype::Q4TiledP,
2539                    ..
2540                }
2541            )
2542        }) && ts
2543            .iter()
2544            .all(|t| t.cols() == ts[0].cols() && t.cols() % GROUP_SIZE == 0);
2545        let Some(pool) = pool else {
2546            for (t, o) in ts.iter().zip(outs.iter_mut()) {
2547                t.matvec(x, o, None);
2548            }
2549            return;
2550        };
2551        if total_rows < 256
2552            || !(uniform_q8
2553                || uniform_f32
2554                || uniform_q4
2555                || uniform_vbit
2556                || uniform_q1
2557                || uniform_q1t
2558                || uniform_q4tp)
2559        {
2560            for (t, o) in ts.iter().zip(outs.iter_mut()) {
2561                t.matvec(x, o, Some(pool));
2562            }
2563            return;
2564        }
2565
2566        if uniform_q4tp {
2567            // Every tensor's rows laid end to end in one virtual row space,
2568            // so the whole set is ONE dispatch. The per-row body is the
2569            // `q4tp_matvec` arm verbatim — same activation split, same
2570            // accumulation order — so the outputs are bit-identical to the
2571            // sequential calls this replaces.
2572            let cols = ts[0].cols();
2573            let gpr = cols / GROUP_SIZE;
2574            let views: [Q4tpView; N] =
2575                std::array::from_fn(|i| Q4tpView::new(ts[i].quant_bytes(), ts[i].rows(), cols));
2576            let rows_of: [usize; N] = std::array::from_fn(|i| ts[i].rows());
2577            let outs_addr: [SendMut; N] = std::array::from_fn(|i| SendMut(outs[i].as_mut_ptr()));
2578            // flat index -> (which tensor, which of its rows)
2579            let locate = |flat: usize| -> (usize, usize) {
2580                let mut acc = 0;
2581                for (i, &r) in rows_of.iter().enumerate() {
2582                    if flat < acc + r {
2583                        return (i, flat - acc);
2584                    }
2585                    acc += r;
2586                }
2587                (rows_of.len() - 1, 0)
2588            };
2589            let (views, outs_addr) = (&views, &outs_addr);
2590            if a8w8_enabled() {
2591                let act = split_act(x);
2592                let act = &act;
2593                let run = |start: usize, end: usize| {
2594                    with_krow(gpr, |sc| {
2595                        for flat in start..end {
2596                            let (t, r) = locate(flat);
2597                            let v = &views[t];
2598                            v.scales_into(r, gpr, sc);
2599                            let mut acc = dot_q4tp_row_i8(v.nib, r, gpr, &act.xq, sc) * act.sx;
2600                            for &(j, xv) in &act.outliers {
2601                                let (w, s) = q4tp_outlier(v.nib, r, gpr, j, sc);
2602                                acc += w * s * xv;
2603                            }
2604                            // SAFETY: one worker owns each (tensor, row) pair.
2605                            unsafe { *outs_addr[t].at(r) = acc };
2606                        }
2607                    });
2608                };
2609                pool.run_rows(total_rows, &run);
2610            } else {
2611                let run = |start: usize, end: usize| {
2612                    with_krow(gpr, |sc| {
2613                        for flat in start..end {
2614                            let (t, r) = locate(flat);
2615                            let v = &views[t];
2616                            v.scales_into(r, gpr, sc);
2617                            // SAFETY: one worker owns each (tensor, row) pair.
2618                            unsafe { *outs_addr[t].at(r) = q4tp_row_exact(v.nib, r, gpr, x, sc) };
2619                        }
2620                    });
2621                };
2622                pool.run_rows(total_rows, &run);
2623            }
2624            return;
2625        }
2626
2627        if uniform_q1 {
2628            // One shared activation split + group sums (q1 has no col
2629            // field; the same input feeds every tensor).
2630            let outs_addr: [SendMut; N] = std::array::from_fn(|i| SendMut(outs[i].as_mut_ptr()));
2631            if a8w8_enabled() {
2632                let act = split_act(x);
2633                let gsum = q1_group_sums(&act.xq, ts[0].cols() / GROUP_SIZE);
2634                let (act, gsum) = (&act, &gsum);
2635                let closures: [_; N] = std::array::from_fn(|i| {
2636                    let (bytes, gpr, out) =
2637                        (ts[i].quant_bytes(), ts[i].cols() / GROUP_SIZE, outs_addr[i]);
2638                    move |s: usize, e: usize| q1_range_a8w8(bytes, gpr, act, gsum, out, s, e)
2639                });
2640                let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2641                    std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2642                pool.run_many(&parts);
2643            } else {
2644                let closures: [_; N] = std::array::from_fn(|i| {
2645                    let (bytes, gpr, out) =
2646                        (ts[i].quant_bytes(), ts[i].cols() / GROUP_SIZE, outs_addr[i]);
2647                    move |s: usize, e: usize| q1_range_f32(bytes, gpr, x, out, s, e)
2648                });
2649                let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2650                    std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2651                pool.run_many(&parts);
2652            }
2653            return;
2654        }
2655
2656        if uniform_q1t {
2657            // Q1T batched: one shared activation split + overlay decode,
2658            // all tensors' rows in ONE pool dispatch (saves N−1 dispatches
2659            // and N−1 redundant split_act calls per layer).
2660            let outs_addr: [SendMut; N] = std::array::from_fn(|i| SendMut(outs[i].as_mut_ptr()));
2661            const TILE: usize = cortiq_core::quant::Q1T_TILE;
2662            if a8w8_enabled() {
2663                let act = split_act(x);
2664                let act = &act;
2665                let x_ref = x;
2666                let closures: [_; N] = std::array::from_fn(|i| {
2667                    let bytes = ts[i].quant_bytes();
2668                    let (rows, cols) = (ts[i].rows(), ts[i].cols());
2669                    let gpr = cols / GROUP_SIZE;
2670                    let (rp_off, ent_off, has_ov) = q1t_overlay(bytes, rows * gpr * TILE, rows);
2671                    let out = outs_addr[i];
2672                    move |s: usize, e: usize| {
2673                        q1t_range_a8w8(bytes, gpr, rp_off, ent_off, has_ov, act, x_ref, out, s, e)
2674                    }
2675                });
2676                let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2677                    std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2678                pool.run_many(&parts);
2679            } else {
2680                let x_ref = x;
2681                let closures: [_; N] = std::array::from_fn(|i| {
2682                    let bytes = ts[i].quant_bytes();
2683                    let (rows, cols) = (ts[i].rows(), ts[i].cols());
2684                    let gpr = cols / GROUP_SIZE;
2685                    let (rp_off, ent_off, has_ov) = q1t_overlay(bytes, rows * gpr * TILE, rows);
2686                    let out = outs_addr[i];
2687                    move |s: usize, e: usize| {
2688                        q1t_range_f32_batch(bytes, gpr, rp_off, ent_off, has_ov, x_ref, out, s, e)
2689                    }
2690                });
2691                let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2692                    std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2693                pool.run_many(&parts);
2694            }
2695            return;
2696        }
2697
2698        if uniform_q4 || uniform_vbit {
2699            let outs_addr: [SendMut; N] = std::array::from_fn(|i| SendMut(outs[i].as_mut_ptr()));
2700            // q4/vbit share one activation split — no per-tensor col field.
2701            if a8w8_enabled() {
2702                let act = split_act(x);
2703                let act = &act;
2704                if uniform_q4 {
2705                    let closures: [_; N] = std::array::from_fn(|i| {
2706                        let (packed, scales) =
2707                            q4_split(ts[i].quant_bytes(), ts[i].rows(), ts[i].cols());
2708                        let (gpr, cols, out) =
2709                            (ts[i].cols() / GROUP_SIZE, ts[i].cols(), outs_addr[i]);
2710                        move |s: usize, e: usize| {
2711                            q4_range_a8w8(packed, scales, gpr, cols, act, out, s, e)
2712                        }
2713                    });
2714                    let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2715                        std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2716                    pool.run_many(&parts);
2717                } else {
2718                    let closures: [_; N] = std::array::from_fn(|i| {
2719                        let Self::Mapped { vbit_offsets, .. } = ts[i] else {
2720                            unreachable!()
2721                        };
2722                        let (bytes, rows, cols, out) = (
2723                            ts[i].quant_bytes(),
2724                            ts[i].rows(),
2725                            ts[i].cols(),
2726                            outs_addr[i],
2727                        );
2728                        move |s: usize, e: usize| {
2729                            vbit_range_a8w8(bytes, vbit_offsets, x, act, rows, cols, out, s, e)
2730                        }
2731                    });
2732                    let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2733                        std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2734                    pool.run_many(&parts);
2735                }
2736                return;
2737            }
2738            if uniform_q4 {
2739                let closures: [_; N] = std::array::from_fn(|i| {
2740                    let (packed, scales) =
2741                        q4_split(ts[i].quant_bytes(), ts[i].rows(), ts[i].cols());
2742                    let (gpr, out) = (ts[i].cols() / GROUP_SIZE, outs_addr[i]);
2743                    move |s: usize, e: usize| q4_range_f32(packed, scales, gpr, x, out, s, e)
2744                });
2745                let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2746                    std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2747                pool.run_many(&parts);
2748            } else {
2749                let closures: [_; N] = std::array::from_fn(|i| {
2750                    let Self::Mapped { vbit_offsets, .. } = ts[i] else {
2751                        unreachable!()
2752                    };
2753                    let (bytes, rows, cols, out) = (
2754                        ts[i].quant_bytes(),
2755                        ts[i].rows(),
2756                        ts[i].cols(),
2757                        outs_addr[i],
2758                    );
2759                    move |s: usize, e: usize| {
2760                        vbit_range_f32(bytes, vbit_offsets, x, rows, cols, out, s, e)
2761                    }
2762                });
2763                let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2764                    std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2765                pool.run_many(&parts);
2766            }
2767            return;
2768        }
2769
2770        if uniform_f32 {
2771            let outs_addr: [SendMut; N] = std::array::from_fn(|i| SendMut(outs[i].as_mut_ptr()));
2772            let closures: [_; N] = std::array::from_fn(|i| {
2773                let Self::F32 { data, cols, .. } = ts[i] else {
2774                    unreachable!()
2775                };
2776                let out = outs_addr[i];
2777                move |start: usize, end: usize| {
2778                    for o in start..end {
2779                        let row = &data[o * cols..(o + 1) * cols];
2780                        let mut sum = 0.0f32;
2781                        for j in 0..*cols {
2782                            sum += row[j] * x[j];
2783                        }
2784                        // SAFETY: disjoint (tensor, row) cells per worker.
2785                        unsafe { *out.at(o) = sum };
2786                    }
2787                }
2788            });
2789            let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2790                std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2791            pool.run_many(&parts);
2792            return;
2793        }
2794
2795        // Uniform q8-family: per-tensor prescale (q8_2f col fields
2796        // differ per tensor) + the shared range kernels.
2797        struct Ctx<'a> {
2798            bytes: &'a [u8],
2799            #[cfg_attr(not(target_arch = "aarch64"), allow(dead_code))]
2800            rep: &'a [u8],
2801            row_scale: &'a [f32],
2802            cols: usize,
2803            xs: std::borrow::Cow<'a, [f32]>,
2804        }
2805        let ctxs: [Ctx<'_>; N] = std::array::from_fn(|i| {
2806            let Self::Mapped {
2807                dtype,
2808                cols,
2809                row_scale,
2810                col_field,
2811                repack,
2812                ..
2813            } = ts[i]
2814            else {
2815                unreachable!()
2816            };
2817            Ctx {
2818                bytes: ts[i].quant_bytes(),
2819                rep: repack,
2820                row_scale,
2821                cols: *cols,
2822                xs: prescale(x, col_field, *dtype),
2823            }
2824        });
2825        let outs_addr: [SendMut; N] = std::array::from_fn(|i| SendMut(outs[i].as_mut_ptr()));
2826        #[cfg(target_arch = "aarch64")]
2827        if sdot_enabled() {
2828            let acts: [SplitAct; N] = std::array::from_fn(|i| split_act(&ctxs[i].xs));
2829            let closures: [_; N] = std::array::from_fn(|i| {
2830                let (c, act, out) = (&ctxs[i], &acts[i], outs_addr[i]);
2831                move |start: usize, end: usize| {
2832                    q8_range_sdot(c.bytes, c.rep, c.row_scale, act, c.cols, out, start, end)
2833                }
2834            });
2835            let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2836                std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2837            pool.run_many(&parts);
2838            return;
2839        }
2840        #[cfg(target_arch = "x86_64")]
2841        if avx2_a8w8_enabled() {
2842            let acts: [SplitAct; N] = std::array::from_fn(|i| split_act(&ctxs[i].xs));
2843            let closures: [_; N] = std::array::from_fn(|i| {
2844                let (c, act, out) = (&ctxs[i], &acts[i], outs_addr[i]);
2845                move |start: usize, end: usize| {
2846                    q8_range_avx2(c.bytes, c.row_scale, act, c.cols, out, start, end)
2847                }
2848            });
2849            let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2850                std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2851            pool.run_many(&parts);
2852            return;
2853        }
2854        let closures: [_; N] = std::array::from_fn(|i| {
2855            let (c, out) = (&ctxs[i], outs_addr[i]);
2856            move |start: usize, end: usize| {
2857                q8_range_f32(c.bytes, c.row_scale, &c.xs, c.cols, out, start, end)
2858            }
2859        });
2860        let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2861            std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2862        pool.run_many(&parts);
2863    }
2864}
2865
2866impl QTensor {
2867    /// Pair-input multi-matrix job: N tensors × 2 shared inputs under a
2868    /// single pool dispatch — the MTP/pair decode path publishes one job
2869    /// for Q/K/V (and one for gate+up) instead of one per tensor.
2870    /// Per-row math is exactly `matvec2`'s kernels; bit-identical.
2871    #[allow(clippy::needless_range_loop)]
2872    pub fn matvec2_many<const N: usize>(
2873        ts: [&QTensor; N],
2874        x1: &[f32],
2875        x2: &[f32],
2876        mut o1s: [&mut [f32]; N],
2877        mut o2s: [&mut [f32]; N],
2878        pool: Option<&Pool>,
2879    ) {
2880        let total_rows: usize = ts.iter().map(|t| t.rows()).sum();
2881        if ts.iter().any(|t| t.has_prism_contract()) {
2882            for i in 0..N {
2883                ts[i].matvec2(x1, x2, o1s[i], o2s[i], pool);
2884            }
2885            return;
2886        }
2887        let uniform_q8 = ts.iter().all(|t| {
2888            matches!(
2889                t,
2890                Self::Mapped {
2891                    dtype: TensorDtype::Q8Row | TensorDtype::Q8_2f,
2892                    ..
2893                }
2894            )
2895        });
2896        let uniform_f32 = ts.iter().all(|t| matches!(t, Self::F32 { .. }));
2897        let uniform_q4 = ts.iter().all(|t| {
2898            matches!(
2899                t,
2900                Self::Mapped {
2901                    dtype: TensorDtype::Q4Block,
2902                    ..
2903                }
2904            )
2905        });
2906        let uniform_vbit = ts.iter().all(|t| {
2907            matches!(
2908                t,
2909                Self::Mapped {
2910                    dtype: TensorDtype::Vbit | TensorDtype::VbitRo,
2911                    ..
2912                }
2913            )
2914        });
2915        let fusable = pool.is_some()
2916            && total_rows >= 256
2917            && (uniform_q8 || uniform_f32 || uniform_q4 || uniform_vbit);
2918        if !fusable {
2919            for i in 0..N {
2920                ts[i].matvec2(x1, x2, o1s[i], o2s[i], pool);
2921            }
2922            return;
2923        }
2924        let pool = pool.unwrap();
2925
2926        if uniform_q4 || uniform_vbit {
2927            let p1: [SendMut; N] = std::array::from_fn(|i| SendMut(o1s[i].as_mut_ptr()));
2928            let p2: [SendMut; N] = std::array::from_fn(|i| SendMut(o2s[i].as_mut_ptr()));
2929            // q4/vbit share activation splits — no per-tensor col field.
2930            if a8w8_enabled() {
2931                let a1 = split_act(x1);
2932                let a2 = split_act(x2);
2933                let (a1, a2) = (&a1, &a2);
2934                if uniform_q4 {
2935                    let closures: [_; N] = std::array::from_fn(|i| {
2936                        let (packed, scales) =
2937                            q4_split(ts[i].quant_bytes(), ts[i].rows(), ts[i].cols());
2938                        let (gpr, cols, o1, o2) =
2939                            (ts[i].cols() / GROUP_SIZE, ts[i].cols(), p1[i], p2[i]);
2940                        move |s: usize, e: usize| {
2941                            q4_range2_a8w8(packed, scales, gpr, cols, a1, a2, o1, o2, s, e)
2942                        }
2943                    });
2944                    let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2945                        std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2946                    pool.run_many(&parts);
2947                } else {
2948                    let closures: [_; N] = std::array::from_fn(|i| {
2949                        let Self::Mapped { vbit_offsets, .. } = ts[i] else {
2950                            unreachable!()
2951                        };
2952                        let (bytes, rows, cols, o1, o2) = (
2953                            ts[i].quant_bytes(),
2954                            ts[i].rows(),
2955                            ts[i].cols(),
2956                            p1[i],
2957                            p2[i],
2958                        );
2959                        move |s: usize, e: usize| {
2960                            vbit_range2_a8w8(
2961                                bytes,
2962                                vbit_offsets,
2963                                x1,
2964                                x2,
2965                                a1,
2966                                a2,
2967                                rows,
2968                                cols,
2969                                o1,
2970                                o2,
2971                                s,
2972                                e,
2973                            )
2974                        }
2975                    });
2976                    let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2977                        std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2978                    pool.run_many(&parts);
2979                }
2980                return;
2981            }
2982            if uniform_q4 {
2983                let closures: [_; N] = std::array::from_fn(|i| {
2984                    let (packed, scales) =
2985                        q4_split(ts[i].quant_bytes(), ts[i].rows(), ts[i].cols());
2986                    let (gpr, o1, o2) = (ts[i].cols() / GROUP_SIZE, p1[i], p2[i]);
2987                    move |s: usize, e: usize| {
2988                        q4_range2_f32(packed, scales, gpr, x1, x2, o1, o2, s, e)
2989                    }
2990                });
2991                let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2992                    std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2993                pool.run_many(&parts);
2994            } else {
2995                let closures: [_; N] = std::array::from_fn(|i| {
2996                    let Self::Mapped { vbit_offsets, .. } = ts[i] else {
2997                        unreachable!()
2998                    };
2999                    let (bytes, rows, cols, o1, o2) = (
3000                        ts[i].quant_bytes(),
3001                        ts[i].rows(),
3002                        ts[i].cols(),
3003                        p1[i],
3004                        p2[i],
3005                    );
3006                    move |s: usize, e: usize| {
3007                        vbit_range2_f32(bytes, vbit_offsets, x1, x2, rows, cols, o1, o2, s, e)
3008                    }
3009                });
3010                let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
3011                    std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
3012                pool.run_many(&parts);
3013            }
3014            return;
3015        }
3016
3017        if uniform_f32 {
3018            let p1: [SendMut; N] = std::array::from_fn(|i| SendMut(o1s[i].as_mut_ptr()));
3019            let p2: [SendMut; N] = std::array::from_fn(|i| SendMut(o2s[i].as_mut_ptr()));
3020            let closures: [_; N] = std::array::from_fn(|i| {
3021                let Self::F32 { data, cols, .. } = ts[i] else {
3022                    unreachable!()
3023                };
3024                let (o1, o2) = (p1[i], p2[i]);
3025                move |start: usize, end: usize| {
3026                    for o in start..end {
3027                        let row = &data[o * cols..(o + 1) * cols];
3028                        let (mut s1, mut s2) = (0.0f32, 0.0f32);
3029                        for j in 0..*cols {
3030                            s1 += row[j] * x1[j];
3031                            s2 += row[j] * x2[j];
3032                        }
3033                        // SAFETY: disjoint (tensor, row) cells per worker.
3034                        unsafe {
3035                            *o1.at(o) = s1;
3036                            *o2.at(o) = s2;
3037                        }
3038                    }
3039                }
3040            });
3041            let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
3042                std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
3043            pool.run_many(&parts);
3044            return;
3045        }
3046
3047        struct Ctx<'a> {
3048            bytes: &'a [u8],
3049            row_scale: &'a [f32],
3050            cols: usize,
3051            xs1: std::borrow::Cow<'a, [f32]>,
3052            xs2: std::borrow::Cow<'a, [f32]>,
3053        }
3054        let ctxs: [Ctx<'_>; N] = std::array::from_fn(|i| {
3055            let Self::Mapped {
3056                dtype,
3057                cols,
3058                row_scale,
3059                col_field,
3060                ..
3061            } = ts[i]
3062            else {
3063                unreachable!()
3064            };
3065            Ctx {
3066                bytes: ts[i].quant_bytes(),
3067                row_scale,
3068                cols: *cols,
3069                xs1: prescale(x1, col_field, *dtype),
3070                xs2: prescale(x2, col_field, *dtype),
3071            }
3072        });
3073        let p1: [SendMut; N] = std::array::from_fn(|i| SendMut(o1s[i].as_mut_ptr()));
3074        let p2: [SendMut; N] = std::array::from_fn(|i| SendMut(o2s[i].as_mut_ptr()));
3075        #[cfg(target_arch = "aarch64")]
3076        if sdot_enabled() {
3077            let acts: [(SplitAct, SplitAct); N] =
3078                std::array::from_fn(|i| (split_act(&ctxs[i].xs1), split_act(&ctxs[i].xs2)));
3079            let closures: [_; N] = std::array::from_fn(|i| {
3080                let (c, a, o1, o2) = (&ctxs[i], &acts[i], p1[i], p2[i]);
3081                move |start: usize, end: usize| {
3082                    q8_range2_sdot(c.bytes, c.row_scale, &a.0, &a.1, c.cols, o1, o2, start, end)
3083                }
3084            });
3085            let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
3086                std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
3087            pool.run_many(&parts);
3088            return;
3089        }
3090        #[cfg(target_arch = "x86_64")]
3091        if avx2_a8w8_enabled() {
3092            let acts: [(SplitAct, SplitAct); N] =
3093                std::array::from_fn(|i| (split_act(&ctxs[i].xs1), split_act(&ctxs[i].xs2)));
3094            let closures: [_; N] = std::array::from_fn(|i| {
3095                let (c, a, o1, o2) = (&ctxs[i], &acts[i], p1[i], p2[i]);
3096                move |start: usize, end: usize| {
3097                    q8_range2_avx2(c.bytes, c.row_scale, &a.0, &a.1, c.cols, o1, o2, start, end)
3098                }
3099            });
3100            let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
3101                std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
3102            pool.run_many(&parts);
3103            return;
3104        }
3105        let closures: [_; N] = std::array::from_fn(|i| {
3106            let (c, o1, o2) = (&ctxs[i], p1[i], p2[i]);
3107            move |start: usize, end: usize| {
3108                q8_range2_f32(
3109                    c.bytes,
3110                    c.row_scale,
3111                    &c.xs1,
3112                    &c.xs2,
3113                    c.cols,
3114                    o1,
3115                    o2,
3116                    start,
3117                    end,
3118                )
3119            }
3120        });
3121        let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
3122            std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
3123        pool.run_many(&parts);
3124    }
3125
3126    /// Fused gate+up matvec with SiLU·mul: for each row r, computes
3127    /// `silu(gate·x) * (up·x)` and writes to `out[r]`. ONE pool dispatch,
3128    /// no intermediate g/u buffers, no separate silu pass. Falls back
3129    /// (returns false) for unsupported dtype combos.
3130    pub fn matvec_silu_mul(
3131        gate: &QTensor,
3132        up: &QTensor,
3133        x: &[f32],
3134        out: &mut [f32],
3135        pool: Option<&Pool>,
3136    ) -> bool {
3137        Self::matvec_silu_mul_limited(gate, up, x, out, 0.0, pool)
3138    }
3139
3140    /// Fused gate+up+SiLU with the GLM asymmetrical clamp.  `limit == 0`
3141    /// preserves the historical unclamped helper; a positive limit clamps
3142    /// `up` to both sides and `gate` only from above, matching the GLM
3143    /// SwiGLU reference.  Keeping the limit in the row kernel avoids the two
3144    /// intermediate vectors and the extra combine pass on the Q2TP experts.
3145    pub fn matvec_silu_mul_limited(
3146        gate: &QTensor,
3147        up: &QTensor,
3148        x: &[f32],
3149        out: &mut [f32],
3150        limit: f32,
3151        pool: Option<&Pool>,
3152    ) -> bool {
3153        if gate.has_prism_contract() || up.has_prism_contract() {
3154            // The fused gate/up kernels consume x directly.  Prism requires
3155            // a per-matrix signed FWHT, so the caller must use two ordinary
3156            // descriptor-aware matvecs instead of an unrotated fast path.
3157            return false;
3158        }
3159        let inter = gate.rows();
3160        debug_assert_eq!(up.rows(), inter);
3161        debug_assert_eq!(out.len(), inter);
3162        debug_assert_eq!(gate.cols(), up.cols());
3163        if !a8w8_enabled() {
3164            return false;
3165        }
3166        let act = split_act(x);
3167        let act = &act;
3168        let x_ref = x;
3169        let out_addr = SendMut(out.as_mut_ptr());
3170
3171        match (gate, up) {
3172            // Q4Block gate + Q4Block up (most common mobile q4 models)
3173            (
3174                Self::Mapped {
3175                    dtype: TensorDtype::Q4Block,
3176                    ..
3177                },
3178                Self::Mapped {
3179                    dtype: TensorDtype::Q4Block,
3180                    ..
3181                },
3182            ) => {
3183                let (gp, gs) = q4_split(gate.quant_bytes(), gate.rows(), gate.cols());
3184                let (up_p, up_s) = q4_split(up.quant_bytes(), up.rows(), up.cols());
3185                let gpr = gate.cols() / GROUP_SIZE;
3186                let cols = gate.cols();
3187                let run = move |start: usize, end: usize| {
3188                    for r in start..end {
3189                        let mut gv = dot_q4_row_i8(gp, gs, r * gpr, gpr, &act.xq) * act.sx;
3190                        let mut uv = dot_q4_row_i8(up_p, up_s, r * gpr, gpr, &act.xq) * act.sx;
3191                        for &(j, xv) in &act.outliers {
3192                            let flat = r * cols + j;
3193                            let gb = gp[flat / 2];
3194                            let gn = if flat & 1 == 0 { gb & 0x0F } else { gb >> 4 };
3195                            let gsc = f16_to_f32(u16::from_le_bytes([
3196                                gs[(flat / GROUP_SIZE) * 2],
3197                                gs[(flat / GROUP_SIZE) * 2 + 1],
3198                            ]));
3199                            gv += ((gn as i32 - 8) as f32) * gsc * xv;
3200                            let ub = up_p[flat / 2];
3201                            let un = if flat & 1 == 0 { ub & 0x0F } else { ub >> 4 };
3202                            let usc = f16_to_f32(u16::from_le_bytes([
3203                                up_s[(flat / GROUP_SIZE) * 2],
3204                                up_s[(flat / GROUP_SIZE) * 2 + 1],
3205                            ]));
3206                            uv += ((un as i32 - 8) as f32) * usc * xv;
3207                        }
3208                        // SAFETY: disjoint row ranges per worker.
3209                        unsafe { *out_addr.at(r) = silu_mul_limited(gv, uv, limit) };
3210                    }
3211                };
3212                dispatch_rows(pool, inter, &run);
3213                true
3214            }
3215            // Q4Tiled gate + Q4Tiled up — one row pass, both tile
3216            // streams sequential, silu·mul fused (same per-row math as
3217            // `q4t_matvec`).
3218            (
3219                Self::Mapped {
3220                    dtype: TensorDtype::Q4Tiled,
3221                    ..
3222                },
3223                Self::Mapped {
3224                    dtype: TensorDtype::Q4Tiled,
3225                    ..
3226                },
3227            ) => {
3228                let g_bytes = gate.quant_bytes();
3229                let u_bytes = up.quant_bytes();
3230                let gpr = gate.cols() / GROUP_SIZE;
3231                let run = move |start: usize, end: usize| {
3232                    for r in start..end {
3233                        let mut gv = dot_q4t_row_i8(g_bytes, r, gpr, &act.xq) * act.sx;
3234                        let mut uv = dot_q4t_row_i8(u_bytes, r, gpr, &act.xq) * act.sx;
3235                        for &(j, xv) in &act.outliers {
3236                            let (w, s) = q4t_outlier(g_bytes, r, gpr, j);
3237                            gv += w * s * xv;
3238                            let (w, s) = q4t_outlier(u_bytes, r, gpr, j);
3239                            uv += w * s * xv;
3240                        }
3241                        // SAFETY: disjoint row ranges per worker.
3242                        unsafe { *out_addr.at(r) = silu_mul_limited(gv, uv, limit) };
3243                    }
3244                };
3245                dispatch_rows(pool, inter, &run);
3246                true
3247            }
3248            // Q4TiledP gate + Q4TiledP up — the same fused row pass, with
3249            // each row's two ladders built once and spent on both streams.
3250            (
3251                Self::Mapped {
3252                    dtype: TensorDtype::Q4TiledP,
3253                    ..
3254                },
3255                Self::Mapped {
3256                    dtype: TensorDtype::Q4TiledP,
3257                    ..
3258                },
3259            ) => {
3260                let cols = gate.cols();
3261                let gpr = cols / GROUP_SIZE;
3262                let gv_view = Q4tpView::new(gate.quant_bytes(), inter, cols);
3263                let uv_view = Q4tpView::new(up.quant_bytes(), inter, cols);
3264                let run = |start: usize, end: usize| {
3265                    with_krows(gpr, |gsc, usc| {
3266                        for r in start..end {
3267                            gv_view.scales_into(r, gpr, gsc);
3268                            uv_view.scales_into(r, gpr, usc);
3269                            let mut gv =
3270                                dot_q4tp_row_i8(gv_view.nib, r, gpr, &act.xq, gsc) * act.sx;
3271                            let mut uv =
3272                                dot_q4tp_row_i8(uv_view.nib, r, gpr, &act.xq, usc) * act.sx;
3273                            for &(j, xv) in &act.outliers {
3274                                let (w, s) = q4tp_outlier(gv_view.nib, r, gpr, j, gsc);
3275                                gv += w * s * xv;
3276                                let (w, s) = q4tp_outlier(uv_view.nib, r, gpr, j, usc);
3277                                uv += w * s * xv;
3278                            }
3279                            // SAFETY: disjoint row ranges per worker.
3280                            unsafe { *out_addr.at(r) = silu_mul_limited(gv, uv, limit) };
3281                        }
3282                    });
3283                };
3284                dispatch_rows(pool, inter, &run);
3285                true
3286            }
3287            // Q1 gate + Q1 up — one row pass over both sign streams,
3288            // silu·mul fused (the per-row math of `q1_range_a8w8`); the
3289            // activation group sums are shared by both streams. Without
3290            // this arm a q1 dense FFN paid two dispatches + a combine
3291            // loop — the exact barrier this function exists to remove.
3292            (
3293                Self::Mapped {
3294                    dtype: TensorDtype::Q1,
3295                    ..
3296                },
3297                Self::Mapped {
3298                    dtype: TensorDtype::Q1,
3299                    ..
3300                },
3301            ) => {
3302                let g_bytes = gate.quant_bytes();
3303                let u_bytes = up.quant_bytes();
3304                let gpr = gate.cols() / GROUP_SIZE;
3305                let gsum = q1_group_sums(&act.xq, gpr);
3306                let gsum = &gsum;
3307                let run = move |start: usize, end: usize| {
3308                    for r in start..end {
3309                        let mut gv = dot_q1_row_i8(g_bytes, r, gpr, &act.xq, gsum) * act.sx;
3310                        let mut uv = dot_q1_row_i8(u_bytes, r, gpr, &act.xq, gsum) * act.sx;
3311                        for &(j, xv) in &act.outliers {
3312                            let (w, s) = q1_outlier(g_bytes, r, gpr, j);
3313                            gv += w * s * xv;
3314                            let (w, s) = q1_outlier(u_bytes, r, gpr, j);
3315                            uv += w * s * xv;
3316                        }
3317                        // SAFETY: disjoint row ranges per worker.
3318                        unsafe { *out_addr.at(r) = silu_mul_limited(gv, uv, limit) };
3319                    }
3320                };
3321                dispatch_rows(pool, inter, &run);
3322                true
3323            }
3324            // Q2TiledP gate + Q2TiledP up — the 2-bit expert pair (MoE
3325            // FFNs of the W2 class): one row pass, both ladders built
3326            // once, integer code dots with shared group sums.
3327            (
3328                Self::Mapped {
3329                    dtype: TensorDtype::Q2TiledP,
3330                    ..
3331                },
3332                Self::Mapped {
3333                    dtype: TensorDtype::Q2TiledP,
3334                    ..
3335                },
3336            ) => {
3337                let cols = gate.cols();
3338                let gpr = cols / GROUP_SIZE;
3339                let gv_view = Q4tpView::new_q2(gate.quant_bytes(), inter, cols);
3340                let uv_view = Q4tpView::new_q2(up.quant_bytes(), inter, cols);
3341                let gsum = q1_group_sums(&act.xq, gpr);
3342                let gsum = &gsum;
3343                let run = move |start: usize, end: usize| {
3344                    with_krows(gpr, |gsc, usc| {
3345                        for r in start..end {
3346                            gv_view.scales_into(r, gpr, gsc);
3347                            uv_view.scales_into(r, gpr, usc);
3348                            let mut gv =
3349                                dot_q2tp_row_i8(gv_view.nib, r, gpr, &act.xq, gsum, gsc) * act.sx;
3350                            let mut uv =
3351                                dot_q2tp_row_i8(uv_view.nib, r, gpr, &act.xq, gsum, usc) * act.sx;
3352                            for &(j, xv) in &act.outliers {
3353                                let (w, s) = q2tp_outlier(gv_view.nib, r, gpr, j, gsc);
3354                                gv += w * s * xv;
3355                                let (w, s) = q2tp_outlier(uv_view.nib, r, gpr, j, usc);
3356                                uv += w * s * xv;
3357                            }
3358                            // SAFETY: disjoint row ranges per worker.
3359                            unsafe { *out_addr.at(r) = silu_mul_limited(gv, uv, limit) };
3360                        }
3361                    });
3362                };
3363                dispatch_rows(pool, inter, &run);
3364                true
3365            }
3366            // Q8Row gate + Q8Row up — one row pass over both i8 streams.
3367            // Q8_2f stays out on purpose: its column field prescales the
3368            // activations PER TENSOR, which breaks this fn's shared
3369            // split_act contract — it keeps the two-dispatch path.
3370            (
3371                Self::Mapped {
3372                    dtype: TensorDtype::Q8Row,
3373                    row_scale: g_rs,
3374                    ..
3375                },
3376                Self::Mapped {
3377                    dtype: TensorDtype::Q8Row,
3378                    row_scale: u_rs,
3379                    ..
3380                },
3381            ) => {
3382                let g_bytes = gate.quant_bytes();
3383                let u_bytes = up.quant_bytes();
3384                let cols = gate.cols();
3385                let run = move |start: usize, end: usize| {
3386                    for r in start..end {
3387                        let gv = q8_row_dot(&g_bytes[r * cols..(r + 1) * cols], act) * g_rs[r];
3388                        let uv = q8_row_dot(&u_bytes[r * cols..(r + 1) * cols], act) * u_rs[r];
3389                        // SAFETY: disjoint row ranges per worker.
3390                        unsafe { *out_addr.at(r) = silu_mul_limited(gv, uv, limit) };
3391                    }
3392                };
3393                dispatch_rows(pool, inter, &run);
3394                true
3395            }
3396            // Q1T gate + Q1T up
3397            (
3398                Self::Mapped {
3399                    dtype: TensorDtype::Q1T,
3400                    ..
3401                },
3402                Self::Mapped {
3403                    dtype: TensorDtype::Q1T,
3404                    ..
3405                },
3406            ) => {
3407                const TILE: usize = cortiq_core::quant::Q1T_TILE;
3408                let g_bytes = gate.quant_bytes();
3409                let u_bytes = up.quant_bytes();
3410                let gpr = gate.cols() / GROUP_SIZE;
3411                let (g_rp, g_ent, g_ov) = q1t_overlay(g_bytes, inter * gpr * TILE, inter);
3412                let (u_rp, u_ent, u_ov) = q1t_overlay(u_bytes, inter * gpr * TILE, inter);
3413                let run = move |start: usize, end: usize| {
3414                    for r in start..end {
3415                        let mut gv = q1t_dot_row_i8(g_bytes, r, gpr, &act.xq) * act.sx;
3416                        let mut uv = q1t_dot_row_i8(u_bytes, r, gpr, &act.xq) * act.sx;
3417                        for &(j, xv) in &act.outliers {
3418                            gv += q1t_base_weight(g_bytes, r, gpr, j) * xv;
3419                            uv += q1t_base_weight(u_bytes, r, gpr, j) * xv;
3420                        }
3421                        gv += q1t_row_outlier_correction(g_bytes, r, g_rp, g_ent, g_ov, x_ref);
3422                        uv += q1t_row_outlier_correction(u_bytes, r, u_rp, u_ent, u_ov, x_ref);
3423                        // SAFETY: disjoint row ranges per worker.
3424                        unsafe { *out_addr.at(r) = silu_mul_limited(gv, uv, limit) };
3425                    }
3426                };
3427                dispatch_rows(pool, inter, &run);
3428                true
3429            }
3430            _ => false,
3431        }
3432    }
3433
3434    /// Every routed expert's fused gate/up/SiLU under ONE pool dispatch.
3435    ///
3436    /// The per-expert path pays a pool barrier per expert per stage: at 9
3437    /// experts over 40 layers that is ~720 barriers a token, and a decode
3438    /// profile of Qwen3.6-35B-A3B showed the pool parked in
3439    /// `psynch_cvwait` about twice as long as it spent computing. Laying
3440    /// every expert's rows end-to-end in one virtual row space collapses
3441    /// the stage to a single dispatch. The per-row body is the
3442    /// single-expert q4tp arm verbatim, so outputs are bit-identical.
3443    ///
3444    /// `false` = something is outside the fused q4tp kernel (dtype, shape,
3445    /// or a transformed tensor); the caller walks the ordinary per-expert
3446    /// path. Float activations use the same exact scalar rows, still fused
3447    /// under one pool dispatch.
3448    pub fn moe_gate_up_many(
3449        pairs: &[(&QTensor, &QTensor)],
3450        x: &[f32],
3451        outs: &mut [Vec<f32>],
3452        pool: Option<&Pool>,
3453    ) -> bool {
3454        Self::moe_gate_up_many_limited(pairs, x, outs, 0.0, pool)
3455    }
3456
3457    /// Batched gate/up/SiLU with the optional GLM clamp.  The public legacy
3458    /// helper above keeps its historical unclamped semantics; callers that
3459    /// implement a reference with a positive SwiGLU limit use this variant.
3460    pub fn moe_gate_up_many_limited(
3461        pairs: &[(&QTensor, &QTensor)],
3462        x: &[f32],
3463        outs: &mut [Vec<f32>],
3464        limit: f32,
3465        pool: Option<&Pool>,
3466    ) -> bool {
3467        if pairs.is_empty() || pairs.len() != outs.len() {
3468            return false;
3469        }
3470        if !a8w8_enabled() {
3471            if limit > 0.0 {
3472                // The exact-row fallback applies no SwiGLU clamp; the
3473                // caller's per-expert path carries it instead.
3474                return false;
3475            }
3476            let groups = vec![vec![0]; pairs.len()];
3477            return Self::moe_gate_up_rows(pairs, &groups, x, outs, pool);
3478        }
3479        let inter = pairs[0].0.rows();
3480        let cols = pairs[0].0.cols();
3481        if cols % GROUP_SIZE != 0 {
3482            return false;
3483        }
3484        let gpr = cols / GROUP_SIZE;
3485        // Uniform layout across every routed pair: q4tp, or the 2-bit
3486        // profile's q2tp gate/up (the W2 class). Mixed sets refuse.
3487        let q2 = matches!(
3488            pairs[0].0,
3489            Self::Mapped {
3490                dtype: TensorDtype::Q2TiledP,
3491                ..
3492            }
3493        );
3494        let want = if q2 {
3495            TensorDtype::Q2TiledP
3496        } else {
3497            TensorDtype::Q4TiledP
3498        };
3499        let mut views = Vec::with_capacity(pairs.len() * 2);
3500        for ((g, u), o) in pairs.iter().zip(outs.iter()) {
3501            let both = matches!(g, Self::Mapped { dtype, .. } if *dtype == want)
3502                && matches!(u, Self::Mapped { dtype, .. } if *dtype == want);
3503            if !both
3504                || g.rows() != inter
3505                || u.rows() != inter
3506                || g.cols() != cols
3507                || u.cols() != cols
3508                || o.len() != inter
3509            {
3510                return false;
3511            }
3512            let mk = if q2 { Q4tpView::new_q2 } else { Q4tpView::new };
3513            views.push(mk(g.quant_bytes(), inter, cols));
3514            views.push(mk(u.quant_bytes(), inter, cols));
3515        }
3516        let act = split_act(x);
3517        let gsum = if q2 {
3518            q1_group_sums(&act.xq, gpr)
3519        } else {
3520            Vec::new()
3521        };
3522        let (act, gsum) = (&act, &gsum);
3523        let ptrs: Vec<SendMut> = outs.iter_mut().map(|o| SendMut(o.as_mut_ptr())).collect();
3524        let (views, ptrs) = (&views, &ptrs);
3525        let run = |start: usize, end: usize| {
3526            with_krows(gpr, |gsc, usc| {
3527                for flat in start..end {
3528                    let (e, r) = (flat / inter, flat % inter);
3529                    let gv_view = &views[e * 2];
3530                    let uv_view = &views[e * 2 + 1];
3531                    gv_view.scales_into(r, gpr, gsc);
3532                    uv_view.scales_into(r, gpr, usc);
3533                    let (mut gv, mut uv) = if q2 {
3534                        (
3535                            dot_q2tp_row_i8(gv_view.nib, r, gpr, &act.xq, gsum, gsc) * act.sx,
3536                            dot_q2tp_row_i8(uv_view.nib, r, gpr, &act.xq, gsum, usc) * act.sx,
3537                        )
3538                    } else {
3539                        (
3540                            dot_q4tp_row_i8(gv_view.nib, r, gpr, &act.xq, gsc) * act.sx,
3541                            dot_q4tp_row_i8(uv_view.nib, r, gpr, &act.xq, usc) * act.sx,
3542                        )
3543                    };
3544                    for &(j, xv) in &act.outliers {
3545                        let (og, ou) = if q2 {
3546                            (
3547                                q2tp_outlier(gv_view.nib, r, gpr, j, gsc),
3548                                q2tp_outlier(uv_view.nib, r, gpr, j, usc),
3549                            )
3550                        } else {
3551                            (
3552                                q4tp_outlier(gv_view.nib, r, gpr, j, gsc),
3553                                q4tp_outlier(uv_view.nib, r, gpr, j, usc),
3554                            )
3555                        };
3556                        gv += og.0 * og.1 * xv;
3557                        uv += ou.0 * ou.1 * xv;
3558                    }
3559                    // SAFETY: one worker owns each (expert, row) pair.
3560                    unsafe { *ptrs[e].at(r) = silu_mul_limited(gv, uv, limit) };
3561                }
3562            });
3563        };
3564        dispatch_rows(pool, pairs.len() * inter, &run);
3565        true
3566    }
3567
3568    /// Every routed expert's down projection, weighted and summed into
3569    /// `out`, under ONE pool dispatch.
3570    ///
3571    /// Partitioned by OUTPUT row rather than by expert: each row is owned
3572    /// by a single worker, so the experts are summed in the caller's order
3573    /// — the same sequence of f32 adds the serial `out[i] += w·eo[i]` loop
3574    /// performs, hence bit-identical. Partitioning by expert instead would
3575    /// race on the shared accumulator.
3576    pub fn moe_down_many(
3577        downs: &[&QTensor],
3578        gs: &[Vec<f32>],
3579        weights: &[f32],
3580        out: &mut [f32],
3581        pool: Option<&Pool>,
3582    ) -> bool {
3583        if downs.is_empty() || downs.len() != gs.len() || downs.len() != weights.len() {
3584            return false;
3585        }
3586        if !a8w8_enabled() {
3587            let mut terms = vec![vec![0.0; out.len()]; downs.len()];
3588            if !Self::moe_down_rows(downs, &vec![1; downs.len()], gs, &mut terms, pool) {
3589                return false;
3590            }
3591            out.fill(0.0);
3592            for (row, &w) in terms.iter().zip(weights) {
3593                for (o, &v) in out.iter_mut().zip(row) {
3594                    *o += w * v;
3595                }
3596            }
3597            return true;
3598        }
3599        let rows = out.len();
3600        let cols = downs[0].cols();
3601        if cols % GROUP_SIZE != 0 {
3602            return false;
3603        }
3604        let gpr = cols / GROUP_SIZE;
3605        let mut views = Vec::with_capacity(downs.len());
3606        for (d, g) in downs.iter().zip(gs.iter()) {
3607            if !matches!(
3608                d,
3609                Self::Mapped {
3610                    dtype: TensorDtype::Q4TiledP,
3611                    ..
3612                }
3613            ) || d.rows() != rows
3614                || d.cols() != cols
3615                || g.len() != cols
3616            {
3617                return false;
3618            }
3619            views.push(Q4tpView::new(d.quant_bytes(), rows, cols));
3620        }
3621        // One int8 split per expert — the activation vectors differ.
3622        let acts: Vec<SplitAct> = gs.iter().map(|g| split_act(g)).collect();
3623        // Partitioned by OUTPUT row, with the experts folded inside: each
3624        // row is owned by one worker, so they are summed in the caller's
3625        // order — the same f32 sequence the serial `out[i] += w·eo[i]`
3626        // loop produces. Partitioning by expert instead would either race
3627        // on the accumulator or need a scratch plane and a second pass;
3628        // measured, that variant was a wash, so this keeps the simpler
3629        // shape.
3630        let out_addr = SendMut(out.as_mut_ptr());
3631        let (views, acts, weights) = (&views, &acts, &weights);
3632        let run = |start: usize, end: usize| {
3633            with_krow(gpr, |sc| {
3634                for r in start..end {
3635                    let mut acc = 0f32;
3636                    for (e, v) in views.iter().enumerate() {
3637                        v.scales_into(r, gpr, sc);
3638                        let a = &acts[e];
3639                        let mut d = dot_q4tp_row_i8(v.nib, r, gpr, &a.xq, sc) * a.sx;
3640                        for &(j, xv) in &a.outliers {
3641                            let (w, s) = q4tp_outlier(v.nib, r, gpr, j, sc);
3642                            d += w * s * xv;
3643                        }
3644                        acc += weights[e] * d;
3645                    }
3646                    // SAFETY: disjoint row ranges per worker.
3647                    unsafe { *out_addr.at(r) = acc };
3648                }
3649            });
3650        };
3651        dispatch_rows(pool, rows, &run);
3652        true
3653    }
3654
3655    /// `moe_gate_up_many` for SEVERAL tokens at once, decode-exact: expert
3656    /// `e` (`pairs[e]`, q4tp) serves the tokens `groups[e]` (row indices
3657    /// into `xs`, each `cols` wide). Every (expert, token) output is
3658    /// bit-identical to `moe_gate_up_many` run on that token alone — the
3659    /// same int8 activation split, VNNI dots, outlier terms and inline
3660    /// SiLU — while each weight row is read once for all the tokens routed
3661    /// to its expert (the speculative verify's expert sharing). `outs` is
3662    /// flat in (expert, token-of-group) order. False = not covered (not
3663    /// q4tp): the caller takes the per-token path. With float activations,
3664    /// the exact scalar row kernel replaces the int8 dot without changing
3665    /// the shared dispatch or route-order reduction.
3666    pub fn moe_gate_up_rows(
3667        pairs: &[(&QTensor, &QTensor)],
3668        groups: &[Vec<usize>],
3669        xs: &[f32],
3670        outs: &mut [Vec<f32>],
3671        pool: Option<&Pool>,
3672    ) -> bool {
3673        if pairs.is_empty() || pairs.len() != groups.len() {
3674            return false;
3675        }
3676        let inter = pairs[0].0.rows();
3677        let cols = pairs[0].0.cols();
3678        let n_pairs: usize = groups.iter().map(|g| g.len()).sum();
3679        if cols == 0 || cols % GROUP_SIZE != 0 || outs.len() != n_pairs || xs.len() % cols != 0 {
3680            return false;
3681        }
3682        let b = xs.len() / cols;
3683        let gpr = cols / GROUP_SIZE;
3684        let mut views = Vec::with_capacity(pairs.len() * 2);
3685        for (g, u) in pairs {
3686            let q4tp = |t: &QTensor| {
3687                matches!(
3688                    t,
3689                    Self::Mapped {
3690                        dtype: TensorDtype::Q4TiledP,
3691                        ..
3692                    }
3693                )
3694            };
3695            if g.has_prism_contract()
3696                || u.has_prism_contract()
3697                || !q4tp(g)
3698                || !q4tp(u)
3699                || g.rows() != inter
3700                || u.rows() != inter
3701                || g.cols() != cols
3702                || u.cols() != cols
3703            {
3704                return false;
3705            }
3706            views.push(Q4tpView::new(g.quant_bytes(), inter, cols));
3707            views.push(Q4tpView::new(u.quant_bytes(), inter, cols));
3708        }
3709        if outs.iter().any(|o| o.len() != inter) || groups.iter().flatten().any(|&t| t >= b) {
3710            return false;
3711        }
3712        let quantized = a8w8_enabled();
3713        let acts: Vec<SplitAct> = if quantized {
3714            (0..b)
3715                .map(|t| split_act(&xs[t * cols..(t + 1) * cols]))
3716                .collect()
3717        } else {
3718            Vec::new()
3719        };
3720        let mut offs = Vec::with_capacity(groups.len());
3721        let mut o = 0usize;
3722        for g in groups {
3723            offs.push(o);
3724            o += g.len();
3725        }
3726        let ptrs: Vec<SendMut> = outs.iter_mut().map(|o| SendMut(o.as_mut_ptr())).collect();
3727        let (views, ptrs, acts, offs) = (&views, &ptrs, &acts, &offs);
3728        let run = |start: usize, end: usize| {
3729            let (mut gsc, mut usc) = (vec![0f32; gpr], vec![0f32; gpr]);
3730            for flat in start..end {
3731                let (e, r) = (flat / inter, flat % inter);
3732                let (gv_view, uv_view) = (&views[e * 2], &views[e * 2 + 1]);
3733                gv_view.scales_into(r, gpr, &mut gsc);
3734                uv_view.scales_into(r, gpr, &mut usc);
3735                for (k, &t) in groups[e].iter().enumerate() {
3736                    if !quantized {
3737                        let x = &xs[t * cols..(t + 1) * cols];
3738                        let gv = q4tp_row_exact(gv_view.nib, r, gpr, x, &gsc);
3739                        let uv = q4tp_row_exact(uv_view.nib, r, gpr, x, &usc);
3740                        unsafe { *ptrs[offs[e] + k].at(r) = (gv / (1.0 + (-gv).exp())) * uv };
3741                        continue;
3742                    }
3743                    let act = &acts[t];
3744                    let mut gv = dot_q4tp_row_i8(gv_view.nib, r, gpr, &act.xq, &gsc) * act.sx;
3745                    let mut uv = dot_q4tp_row_i8(uv_view.nib, r, gpr, &act.xq, &usc) * act.sx;
3746                    for &(j, xv) in &act.outliers {
3747                        let og = q4tp_outlier(gv_view.nib, r, gpr, j, &gsc);
3748                        let ou = q4tp_outlier(uv_view.nib, r, gpr, j, &usc);
3749                        gv += og.0 * og.1 * xv;
3750                        uv += ou.0 * ou.1 * xv;
3751                    }
3752                    let silu_g = gv / (1.0 + (-gv).exp());
3753                    // SAFETY: one worker owns each (expert, row) cell of
3754                    // every output of the expert's group.
3755                    unsafe { *ptrs[offs[e] + k].at(r) = silu_g * uv };
3756                }
3757            }
3758        };
3759        dispatch_rows(pool, pairs.len() * inter, &run);
3760        true
3761    }
3762
3763    /// The per-(expert, token) down terms `moe_down_many` weights and sums,
3764    /// for SEVERAL tokens: `outs[p][o] = down_e[o] · gs[p]` (int8 split of
3765    /// `gs[p]`, VNNI dot, outlier terms — bit-identical to that kernel's
3766    /// `d`), each down row read once for its expert's whole group. The
3767    /// caller sums `w·d` per token in its route order, which reproduces
3768    /// `moe_down_many`'s f32 sequence exactly. Layout as `moe_gate_up_rows`.
3769    pub fn moe_down_rows(
3770        downs: &[&QTensor],
3771        group_lens: &[usize],
3772        gs: &[Vec<f32>],
3773        outs: &mut [Vec<f32>],
3774        pool: Option<&Pool>,
3775    ) -> bool {
3776        if downs.is_empty() || downs.len() != group_lens.len() {
3777            return false;
3778        }
3779        let rows = downs[0].rows();
3780        let cols = downs[0].cols();
3781        let n_pairs: usize = group_lens.iter().sum();
3782        if cols == 0 || cols % GROUP_SIZE != 0 || gs.len() != n_pairs || outs.len() != n_pairs {
3783            return false;
3784        }
3785        let gpr = cols / GROUP_SIZE;
3786        let mut views = Vec::with_capacity(downs.len());
3787        for d in downs {
3788            if d.has_prism_contract()
3789                || !matches!(
3790                    d,
3791                    Self::Mapped {
3792                        dtype: TensorDtype::Q4TiledP,
3793                        ..
3794                    }
3795                )
3796                || d.rows() != rows
3797                || d.cols() != cols
3798            {
3799                return false;
3800            }
3801            views.push(Q4tpView::new(d.quant_bytes(), rows, cols));
3802        }
3803        if gs.iter().any(|g| g.len() != cols) || outs.iter().any(|o| o.len() != rows) {
3804            return false;
3805        }
3806        let quantized = a8w8_enabled();
3807        let acts: Vec<SplitAct> = if quantized {
3808            gs.iter().map(|g| split_act(g)).collect()
3809        } else {
3810            Vec::new()
3811        };
3812        let mut offs = Vec::with_capacity(group_lens.len());
3813        let mut o = 0usize;
3814        for &l in group_lens {
3815            offs.push(o);
3816            o += l;
3817        }
3818        let ptrs: Vec<SendMut> = outs.iter_mut().map(|o| SendMut(o.as_mut_ptr())).collect();
3819        let (views, ptrs, acts, offs) = (&views, &ptrs, &acts, &offs);
3820        let run = |start: usize, end: usize| {
3821            let mut sc = vec![0f32; gpr];
3822            for flat in start..end {
3823                let (e, r) = (flat / rows, flat % rows);
3824                let v = &views[e];
3825                v.scales_into(r, gpr, &mut sc);
3826                for k in 0..group_lens[e] {
3827                    if !quantized {
3828                        let d = q4tp_row_exact(v.nib, r, gpr, &gs[offs[e] + k], &sc);
3829                        unsafe { *ptrs[offs[e] + k].at(r) = d };
3830                        continue;
3831                    }
3832                    let a = &acts[offs[e] + k];
3833                    let mut d = dot_q4tp_row_i8(v.nib, r, gpr, &a.xq, &sc) * a.sx;
3834                    for &(j, xv) in &a.outliers {
3835                        let (w, s) = q4tp_outlier(v.nib, r, gpr, j, &sc);
3836                        d += w * s * xv;
3837                    }
3838                    // SAFETY: one worker owns each (expert, row) cell.
3839                    unsafe { *ptrs[offs[e] + k].at(r) = d };
3840                }
3841            }
3842        };
3843        dispatch_rows(pool, downs.len() * rows, &run);
3844        true
3845    }
3846}
3847
3848/// Batched q8 kernel: same math as qmatvec, the row makes a single
3849/// pass from memory for the whole batch.
3850/// Accelerate CBLAS — the Apple AMX matrix units, the same engine
3851/// llama.cpp's `-ngl 0` prefill rides via ggml-blas.
3852#[cfg(target_os = "macos")]
3853mod accel_blas {
3854    #[link(name = "Accelerate", kind = "framework")]
3855    unsafe extern "C" {
3856        pub fn cblas_sgemm(
3857            order: i32,
3858            trans_a: i32,
3859            trans_b: i32,
3860            m: i32,
3861            n: i32,
3862            k: i32,
3863            alpha: f32,
3864            a: *const f32,
3865            lda: i32,
3866            b: *const f32,
3867            ldb: i32,
3868            beta: f32,
3869            c: *mut f32,
3870            ldc: i32,
3871        );
3872    }
3873}
3874
3875#[cfg(target_os = "macos")]
3876pub(crate) fn accel_gemm_enabled() -> bool {
3877    static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
3878    *ON.get_or_init(|| std::env::var("CMF_ACCEL").map(|v| v != "0").unwrap_or(true))
3879}
3880
3881/// Off macOS the "accel" GEMM is the portable NEON micro-kernel below —
3882/// same entry point, so the batched-attention path opens on mobile.
3883#[cfg(all(target_arch = "aarch64", not(target_os = "macos")))]
3884pub(crate) fn accel_gemm_enabled() -> bool {
3885    true
3886}
3887
3888/// Portable NEON f32 GEMM (row-major, optional Bᵀ): a 4×8 fmla
3889/// micro-kernel with A broadcast against B panels — the mobile stand-in
3890/// for Accelerate in the batched causal attention (QKᵀ and P·V). Not a
3891/// BLAS: shapes here are the attention panels (m ≤ heads·chunk,
3892/// k = head_dim or context), and the goal is removing the per-position
3893/// quadratic wall, not peak GEMM.
3894#[cfg(target_arch = "aarch64")]
3895#[allow(clippy::too_many_arguments)]
3896pub(crate) fn neon_gemm_rm(
3897    m: usize,
3898    n: usize,
3899    k: usize,
3900    alpha: f32,
3901    a: &[f32],
3902    lda: usize,
3903    b_mat: &[f32],
3904    ldb: usize,
3905    b_rows_are_n: bool,
3906    c: &mut [f32],
3907    ldc: usize,
3908) {
3909    debug_assert!(a.len() >= (m - 1) * lda + k);
3910    debug_assert!(c.len() >= (m - 1) * ldc + n);
3911    // SAFETY: bounds asserted above; NEON is baseline on aarch64.
3912    unsafe {
3913        use core::arch::aarch64::*;
3914        let mut i = 0usize;
3915        while i < m {
3916            let mi = (m - i).min(4);
3917            let mut j = 0usize;
3918            while j < n {
3919                let nj = (n - j).min(8);
3920                if mi == 4 && nj == 8 {
3921                    let (mut c0a, mut c0b) = (vdupq_n_f32(0.0), vdupq_n_f32(0.0));
3922                    let (mut c1a, mut c1b) = (vdupq_n_f32(0.0), vdupq_n_f32(0.0));
3923                    let (mut c2a, mut c2b) = (vdupq_n_f32(0.0), vdupq_n_f32(0.0));
3924                    let (mut c3a, mut c3b) = (vdupq_n_f32(0.0), vdupq_n_f32(0.0));
3925                    for p in 0..k {
3926                        let (b0, b1) = if b_rows_are_n {
3927                            // B is [n, k]: column p of Bᵀ = element p of
3928                            // eight consecutive B rows — gathered.
3929                            let base = b_mat.as_ptr().add(j * ldb + p);
3930                            let g = |o: usize| *base.add(o * ldb);
3931                            ([g(0), g(1), g(2), g(3)], [g(4), g(5), g(6), g(7)])
3932                        } else {
3933                            let base = b_mat.as_ptr().add(p * ldb + j);
3934                            (
3935                                [*base, *base.add(1), *base.add(2), *base.add(3)],
3936                                [*base.add(4), *base.add(5), *base.add(6), *base.add(7)],
3937                            )
3938                        };
3939                        let bv0 = vld1q_f32(b0.as_ptr());
3940                        let bv1 = vld1q_f32(b1.as_ptr());
3941                        let a0 = vdupq_n_f32(*a.as_ptr().add(i * lda + p));
3942                        let a1 = vdupq_n_f32(*a.as_ptr().add((i + 1) * lda + p));
3943                        let a2 = vdupq_n_f32(*a.as_ptr().add((i + 2) * lda + p));
3944                        let a3 = vdupq_n_f32(*a.as_ptr().add((i + 3) * lda + p));
3945                        c0a = vfmaq_f32(c0a, a0, bv0);
3946                        c0b = vfmaq_f32(c0b, a0, bv1);
3947                        c1a = vfmaq_f32(c1a, a1, bv0);
3948                        c1b = vfmaq_f32(c1b, a1, bv1);
3949                        c2a = vfmaq_f32(c2a, a2, bv0);
3950                        c2b = vfmaq_f32(c2b, a2, bv1);
3951                        c3a = vfmaq_f32(c3a, a3, bv0);
3952                        c3b = vfmaq_f32(c3b, a3, bv1);
3953                    }
3954                    let al = vdupq_n_f32(alpha);
3955                    for (r, (ca, cb)) in [(c0a, c0b), (c1a, c1b), (c2a, c2b), (c3a, c3b)]
3956                        .iter()
3957                        .enumerate()
3958                    {
3959                        let dst = c.as_mut_ptr().add((i + r) * ldc + j);
3960                        vst1q_f32(dst, vmulq_f32(*ca, al));
3961                        vst1q_f32(dst.add(4), vmulq_f32(*cb, al));
3962                    }
3963                } else {
3964                    for r in 0..mi {
3965                        for q in 0..nj {
3966                            let mut acc = 0f32;
3967                            for p in 0..k {
3968                                let bv = if b_rows_are_n {
3969                                    b_mat[(j + q) * ldb + p]
3970                                } else {
3971                                    b_mat[p * ldb + j + q]
3972                                };
3973                                acc += a[(i + r) * lda + p] * bv;
3974                            }
3975                            c[(i + r) * ldc + j + q] = acc * alpha;
3976                        }
3977                    }
3978                }
3979                j += nj;
3980            }
3981            i += mi;
3982        }
3983    }
3984}
3985
3986/// Off-macOS aarch64: the batched attention rides the NEON micro-GEMM.
3987#[cfg(all(target_arch = "aarch64", not(target_os = "macos")))]
3988#[allow(clippy::too_many_arguments)]
3989pub(crate) fn sgemm_rm(
3990    m: usize,
3991    n: usize,
3992    k: usize,
3993    alpha: f32,
3994    a: &[f32],
3995    lda: usize,
3996    b_mat: &[f32],
3997    ldb: usize,
3998    b_rows_are_n: bool,
3999    c: &mut [f32],
4000    ldc: usize,
4001) {
4002    neon_gemm_rm(m, n, k, alpha, a, lda, b_mat, ldb, b_rows_are_n, c, ldc);
4003}
4004
4005/// Row-major f32 GEMM, exposed for offline tools (the AWNP pass builds a
4006/// per-layer projection and applies it to every expert; a naive triple loop
4007/// would turn a two-minute job into half an hour).
4008#[allow(clippy::too_many_arguments)]
4009pub fn sgemm_public(
4010    m: usize,
4011    n: usize,
4012    k: usize,
4013    alpha: f32,
4014    a: &[f32],
4015    lda: usize,
4016    b_mat: &[f32],
4017    ldb: usize,
4018    b_rows_are_n: bool,
4019    c: &mut [f32],
4020    ldc: usize,
4021) {
4022    #[cfg(any(target_os = "macos", target_arch = "aarch64"))]
4023    {
4024        sgemm_rm(m, n, k, alpha, a, lda, b_mat, ldb, b_rows_are_n, c, ldc);
4025    }
4026    // x86 without Accelerate has no sgemm_rm: the specialized paths there are
4027    // quantized kernels, not an f32 GEMM. Only the offline AWNP pass reaches
4028    // this, so correctness matters and throughput does not — a triple loop is
4029    // the honest fallback rather than a reason to make the tool macOS-only.
4030    #[cfg(not(any(target_os = "macos", target_arch = "aarch64")))]
4031    {
4032        for i in 0..m {
4033            for j in 0..n {
4034                let mut acc = 0f32;
4035                for p in 0..k {
4036                    let bv = if b_rows_are_n {
4037                        b_mat[j * ldb + p]
4038                    } else {
4039                        b_mat[p * ldb + j]
4040                    };
4041                    acc += a[i * lda + p] * bv;
4042                }
4043                c[i * ldc + j] = alpha * acc;
4044            }
4045        }
4046    }
4047}
4048
4049/// Row-major f32 GEMM on Accelerate: C[m,n] = alpha·A[m,k] × B(ᵀ).
4050/// `b_rows_are_n` = true multiplies by Bᵀ where B is stored [n, k].
4051#[cfg(target_os = "macos")]
4052#[allow(clippy::too_many_arguments)]
4053pub(crate) fn sgemm_rm(
4054    m: usize,
4055    n: usize,
4056    k: usize,
4057    alpha: f32,
4058    a: &[f32],
4059    lda: usize,
4060    b_mat: &[f32],
4061    ldb: usize,
4062    b_rows_are_n: bool,
4063    c: &mut [f32],
4064    ldc: usize,
4065) {
4066    debug_assert!(a.len() >= (m - 1) * lda + k);
4067    debug_assert!(c.len() >= (m - 1) * ldc + n);
4068    // Test hook: route the attention GEMMs through the portable NEON
4069    // micro-kernel ON APPLE SILICON — how the mobile batched attend is
4070    // measured without a phone in the loop. (Intel macOS has no NEON —
4071    // the hook is a no-op there, Accelerate continues below.)
4072    #[cfg(target_arch = "aarch64")]
4073    if std::env::var("CMF_FORCE_NEON_GEMM")
4074        .map(|v| v == "1")
4075        .unwrap_or(false)
4076    {
4077        return neon_gemm_rm(m, n, k, alpha, a, lda, b_mat, ldb, b_rows_are_n, c, ldc);
4078    }
4079    unsafe {
4080        accel_blas::cblas_sgemm(
4081            101, // RowMajor
4082            111, // NoTrans A
4083            if b_rows_are_n { 112 } else { 111 },
4084            m as i32,
4085            n as i32,
4086            k as i32,
4087            alpha,
4088            a.as_ptr(),
4089            lda as i32,
4090            b_mat.as_ptr(),
4091            ldb as i32,
4092            0.0,
4093            c.as_mut_ptr(),
4094            ldc as i32,
4095        );
4096    }
4097}
4098
4099/// Prefill GEMM through Accelerate (macOS): dequantize q8 rows into
4100/// f32 tiles (scale folded in, pool-parallel) and multiply each tile
4101/// on the AMX with one row-major sgemm. Tiles live in cache, weights
4102/// stream once. Numerics are f32-GEMM (not the int8 dot): prefill
4103/// logits shift within f32 rounding — tolerance-class, like every
4104/// reduction-order change; decode (M=1) never takes this path.
4105#[cfg(target_os = "macos")]
4106fn qmatmat_accel(
4107    q: &[u8],
4108    row_scale: &[f32],
4109    pre: &[std::borrow::Cow<'_, [f32]>],
4110    rows: usize,
4111    cols: usize,
4112    out: &mut [f32],
4113    pool: Option<&Pool>,
4114) {
4115    // NOTE: double-buffering the dequant against the sgemm (a scoped
4116    // thread driving the pool on tile k+1 while the caller multiplies
4117    // tile k) was tried and LOST ~6%: Accelerate's sgemm is itself
4118    // multithreaded, and the dequant workers just steal its cores.
4119    const TR: usize = 2048;
4120    let b = pre.len();
4121    thread_local! {
4122        static XPANEL: std::cell::RefCell<Vec<f32>> = const { std::cell::RefCell::new(Vec::new()) };
4123        static WTILE: std::cell::RefCell<Vec<f32>> = const { std::cell::RefCell::new(Vec::new()) };
4124    }
4125    XPANEL.with(|xp| {
4126        WTILE.with(|wt| {
4127            let mut xpanel = xp.borrow_mut();
4128            xpanel.clear();
4129            for x in pre {
4130                xpanel.extend_from_slice(x);
4131            }
4132            let mut wtile = wt.borrow_mut();
4133            wtile.resize(TR * cols, 0.0);
4134            let mut r0 = 0usize;
4135            while r0 < rows {
4136                let tr = TR.min(rows - r0);
4137                // Dequant the tile (scale folded) — pool-parallel.
4138                let wt_addr = SendMut(wtile.as_mut_ptr());
4139                let run = |start: usize, end: usize| {
4140                    for r in start..end {
4141                        let row = &q[(r0 + r) * cols..(r0 + r + 1) * cols];
4142                        let s = row_scale[r0 + r];
4143                        // SAFETY: workers cover disjoint r ranges.
4144                        let dst =
4145                            unsafe { std::slice::from_raw_parts_mut(wt_addr.at(r * cols), cols) };
4146                        for (d, &v) in dst.iter_mut().zip(row) {
4147                            *d = (v as i8) as f32 * s;
4148                        }
4149                    }
4150                };
4151                dispatch_rows(pool, tr, &run);
4152                // C[b, tr] (at column r0 of out[b, rows]) = X · Wtileᵀ
4153                unsafe {
4154                    accel_blas::cblas_sgemm(
4155                        101, // RowMajor
4156                        111, // NoTrans A
4157                        112, // Trans B
4158                        b as i32,
4159                        tr as i32,
4160                        cols as i32,
4161                        1.0,
4162                        xpanel.as_ptr(),
4163                        cols as i32,
4164                        wtile.as_ptr(),
4165                        cols as i32,
4166                        0.0,
4167                        out.as_mut_ptr().add(r0),
4168                        rows as i32,
4169                    );
4170                }
4171                r0 += tr;
4172            }
4173        })
4174    });
4175}
4176
4177fn qmatmat(
4178    q: &[u8],
4179    row_scale: &[f32],
4180    pre: &[std::borrow::Cow<'_, [f32]>],
4181    rows: usize,
4182    cols: usize,
4183    out: &mut [f32],
4184    pool: Option<&Pool>,
4185) {
4186    let b = pre.len();
4187    debug_assert_eq!(out.len(), b * rows);
4188    // Big prefill batches ride the AMX (roadmap PR3): the row×batch
4189    // SDOT loop below peaks near the CPU's dot throughput, an order
4190    // below the matrix units. Small tensors and tiny test models stay
4191    // on the exact integer path.
4192    #[cfg(target_os = "macos")]
4193    if b >= 8 && rows * cols >= 500_000 && accel_gemm_enabled() {
4194        qmatmat_accel(q, row_scale, pre, rows, cols, out, pool);
4195        return;
4196    }
4197    #[cfg(target_arch = "aarch64")]
4198    if sdot_enabled() {
4199        let acts: Vec<SplitAct> = pre.iter().map(|x| split_act(x)).collect();
4200        let out_addr = SendMut(out.as_mut_ptr());
4201        // Blocked 2×4 (mobile prefill: no AMX to fall back on — this
4202        // path IS the ARM prefill GEMM off Apple silicon).
4203        let blocked_ok = blocked_enabled();
4204        let use_i8mm = i8mm_enabled();
4205        if blocked_ok {
4206            let run = |start: usize, end: usize| {
4207                let mut o = start;
4208                while o < end {
4209                    if o + 2 <= end {
4210                        let r0 = &q[o * cols..(o + 1) * cols];
4211                        let r1 = &q[(o + 1) * cols..(o + 2) * cols];
4212                        let mut bi = 0usize;
4213                        while bi + 4 <= acts.len() {
4214                            let xs = [
4215                                acts[bi].xq.as_slice(),
4216                                acts[bi + 1].xq.as_slice(),
4217                                acts[bi + 2].xq.as_slice(),
4218                                acts[bi + 3].xq.as_slice(),
4219                            ];
4220                            let d = if use_i8mm {
4221                                unsafe { dot_i8_smmla_2x4(r0, r1, xs) }
4222                            } else {
4223                                unsafe { dot_i8_sdot_2x4(r0, r1, xs) }
4224                            };
4225                            for (r, row) in [r0, r1].into_iter().enumerate() {
4226                                for k in 0..4 {
4227                                    let act = &acts[bi + k];
4228                                    let mut v = d[r][k] as f32 * act.sx;
4229                                    for &(j, xv) in &act.outliers {
4230                                        v += (row[j] as i8) as f32 * xv;
4231                                    }
4232                                    unsafe {
4233                                        *out_addr.at((bi + k) * rows + o + r) = v * row_scale[o + r]
4234                                    };
4235                                }
4236                            }
4237                            bi += 4;
4238                        }
4239                        while bi < acts.len() {
4240                            for (r, row) in [r0, r1].into_iter().enumerate() {
4241                                let v = row_dot_sdot(row, &acts[bi]) * row_scale[o + r];
4242                                unsafe { *out_addr.at(bi * rows + o + r) = v };
4243                            }
4244                            bi += 1;
4245                        }
4246                        o += 2;
4247                    } else {
4248                        let row = &q[o * cols..(o + 1) * cols];
4249                        for (bi, act) in acts.iter().enumerate() {
4250                            let v = row_dot_sdot(row, act) * row_scale[o];
4251                            unsafe { *out_addr.at(bi * rows + o) = v };
4252                        }
4253                        o += 1;
4254                    }
4255                }
4256            };
4257            dispatch_rows(pool, rows, &run);
4258            return;
4259        }
4260        let run = |start: usize, end: usize| {
4261            for o in start..end {
4262                let row = &q[o * cols..(o + 1) * cols];
4263                for (bi, act) in acts.iter().enumerate() {
4264                    let v = row_dot_sdot(row, act) * row_scale[o];
4265                    unsafe { *out_addr.at(bi * rows + o) = v };
4266                }
4267            }
4268        };
4269        dispatch_rows(pool, rows, &run);
4270        return;
4271    }
4272    // x86 A8W8 batch. Non-VNNI parts take the BLOCKED 2×4 kernel
4273    // (roadmap P0: two weight rows' abs() stay in registers across four
4274    // activation streams); VNNI machines keep the per-row bias-trick
4275    // dot, which is already throughput-bound there.
4276    #[cfg(target_arch = "x86_64")]
4277    if avx2_a8w8_enabled() {
4278        let acts: Vec<SplitAct> = pre.iter().map(|x| split_act(x)).collect();
4279        let out_addr = SendMut(out.as_mut_ptr());
4280        // CMF_X86_BLOCKED=0 forces the per-row path (paired in-process
4281        // A/B on noisy shared-vCPU hosts).
4282        let blocked_ok = blocked_enabled();
4283        if !avx512vnni_enabled() && blocked_ok && !row_exact() {
4284            let run = |start: usize, end: usize| {
4285                let mut o = start;
4286                while o < end {
4287                    if o + 2 <= end {
4288                        let r0 = &q[o * cols..(o + 1) * cols];
4289                        let r1 = &q[(o + 1) * cols..(o + 2) * cols];
4290                        let mut bi = 0usize;
4291                        while bi + 4 <= acts.len() {
4292                            let xs = [
4293                                acts[bi].xq.as_slice(),
4294                                acts[bi + 1].xq.as_slice(),
4295                                acts[bi + 2].xq.as_slice(),
4296                                acts[bi + 3].xq.as_slice(),
4297                            ];
4298                            let d = unsafe { dot_i8_i8_avx2_2x4(r0, r1, xs) };
4299                            for (r, row) in [r0, r1].into_iter().enumerate() {
4300                                for k in 0..4 {
4301                                    let act = &acts[bi + k];
4302                                    let mut v = d[r][k] as f32 * act.sx;
4303                                    for &(j, xv) in &act.outliers {
4304                                        v += (row[j] as i8) as f32 * xv;
4305                                    }
4306                                    unsafe {
4307                                        *out_addr.at((bi + k) * rows + o + r) = v * row_scale[o + r]
4308                                    };
4309                                }
4310                            }
4311                            bi += 4;
4312                        }
4313                        while bi < acts.len() {
4314                            for (r, row) in [r0, r1].into_iter().enumerate() {
4315                                let v = row_dot_avx2(row, &acts[bi]) * row_scale[o + r];
4316                                unsafe { *out_addr.at(bi * rows + o + r) = v };
4317                            }
4318                            bi += 1;
4319                        }
4320                        o += 2;
4321                    } else {
4322                        let row = &q[o * cols..(o + 1) * cols];
4323                        for (bi, act) in acts.iter().enumerate() {
4324                            let v = row_dot_avx2(row, act) * row_scale[o];
4325                            unsafe { *out_addr.at(bi * rows + o) = v };
4326                        }
4327                        o += 1;
4328                    }
4329                }
4330            };
4331            dispatch_rows(pool, rows, &run);
4332            return;
4333        }
4334        let run = |start: usize, end: usize| {
4335            for o in start..end {
4336                let row = &q[o * cols..(o + 1) * cols];
4337                for (bi, act) in acts.iter().enumerate() {
4338                    let v = row_dot_avx2(row, act) * row_scale[o];
4339                    unsafe { *out_addr.at(bi * rows + o) = v };
4340                }
4341            }
4342        };
4343        dispatch_rows(pool, rows, &run);
4344        return;
4345    }
4346    let out_addr = SendMut(out.as_mut_ptr());
4347    let run = |start: usize, end: usize| {
4348        for o in start..end {
4349            let row = &q[o * cols..(o + 1) * cols];
4350            for (bi, x) in pre.iter().enumerate() {
4351                let mut acc = 0f32;
4352                for j in 0..cols {
4353                    acc += (row[j] as i8) as f32 * x[j];
4354                }
4355                unsafe { *out_addr.at(bi * rows + o) = acc * row_scale[o] };
4356            }
4357        }
4358    };
4359    dispatch_rows(pool, rows, &run);
4360}
4361
4362/// Split rows across pool workers (shared qmatvec pattern). Self-balancing
4363/// — see `Pool::run_rows` for why a static 1/n split is wrong here.
4364fn dispatch_rows(pool: Option<&Pool>, rows: usize, run: &(dyn Fn(usize, usize) + Sync)) {
4365    match pool {
4366        Some(pool) if rows >= 256 => pool.run_rows(rows, run),
4367        _ => run(0, rows),
4368    }
4369}
4370
4371/// Split a q4_block blob into (packed nibbles, f16 group scales).
4372fn q4_split(bytes: &[u8], rows: usize, cols: usize) -> (&[u8], &[u8]) {
4373    let groups = rows * cols / GROUP_SIZE;
4374    bytes.split_at(groups * 16)
4375}
4376
4377/// SIMD unpack for the dominant vbit width B=4 (94% of rows on the
4378/// log2-shape calibration): 16 packed bytes -> 32 centered i8 values.
4379/// vbit packs MSB-first, so the HIGH nibble is the even element
4380/// (opposite of q4_block's lo-first interleave). Centering is u-7.
4381#[inline]
4382fn vbit_fill4(data: &[u8], buf: &mut [u8]) {
4383    #[cfg(target_arch = "aarch64")]
4384    unsafe {
4385        return vbit_fill4_neon(data, buf);
4386    }
4387    #[cfg(target_arch = "x86_64")]
4388    if avx2_enabled() {
4389        return unsafe { vbit_fill4_avx2(data, buf) };
4390    }
4391    #[allow(unreachable_code)]
4392    for (blk, chunk) in buf.chunks_exact_mut(8).enumerate() {
4393        let u = unpack8::<4>(&data[blk * 4..]);
4394        for k in 0..8 {
4395            chunk[k] = (u[k] - 7) as i8 as u8;
4396        }
4397    }
4398}
4399
4400#[cfg(target_arch = "aarch64")]
4401#[target_feature(enable = "neon")]
4402unsafe fn vbit_fill4_neon(data: &[u8], buf: &mut [u8]) {
4403    // SAFETY: buf.len() is a multiple of GROUP_SIZE=32; data holds
4404    // buf.len()/2 packed bytes (validated at load).
4405    unsafe {
4406        use core::arch::aarch64::*;
4407        let n = buf.len();
4408        let mask = vdupq_n_u8(0x0F);
4409        let seven = vdupq_n_s8(7);
4410        let mut g = 0usize;
4411        while g * 32 + 32 <= n {
4412            let b = vld1q_u8(data.as_ptr().add(g * 16));
4413            let hi = vshrq_n_u8::<4>(b);
4414            let lo = vandq_u8(b, mask);
4415            let z0 = vsubq_s8(vreinterpretq_s8_u8(vzip1q_u8(hi, lo)), seven);
4416            let z1 = vsubq_s8(vreinterpretq_s8_u8(vzip2q_u8(hi, lo)), seven);
4417            vst1q_u8(buf.as_mut_ptr().add(g * 32), vreinterpretq_u8_s8(z0));
4418            vst1q_u8(buf.as_mut_ptr().add(g * 32 + 16), vreinterpretq_u8_s8(z1));
4419            g += 1;
4420        }
4421    }
4422}
4423
4424#[cfg(target_arch = "x86_64")]
4425#[target_feature(enable = "avx2")]
4426unsafe fn vbit_fill4_avx2(data: &[u8], buf: &mut [u8]) {
4427    // SAFETY: see vbit_fill4_neon.
4428    unsafe {
4429        use core::arch::x86_64::*;
4430        let n = buf.len();
4431        let mask = _mm_set1_epi8(0x0F);
4432        let seven = _mm256_set1_epi8(7);
4433        let mut g = 0usize;
4434        while g * 32 + 32 <= n {
4435            let b = _mm_loadu_si128(data.as_ptr().add(g * 16) as *const __m128i);
4436            let hi = _mm_and_si128(_mm_srli_epi16::<4>(b), mask);
4437            let lo = _mm_and_si128(b, mask);
4438            let z = _mm256_sub_epi8(
4439                _mm256_set_m128i(_mm_unpackhi_epi8(hi, lo), _mm_unpacklo_epi8(hi, lo)),
4440                seven,
4441            );
4442            _mm256_storeu_si256(buf.as_mut_ptr().add(g * 32) as *mut __m256i, z);
4443            g += 1;
4444        }
4445    }
4446}
4447
4448/// Unpack 8 MSB-first B-bit values from exactly B bytes (fixed shifts —
4449/// no serial bit-buffer, auto-vectorizable). Every 32-value group starts
4450/// byte-aligned (32·B/8 is integral for B∈3..8), so groups decompose
4451/// into 4 such blocks.
4452#[inline(always)]
4453fn unpack8<const B: usize>(data: &[u8]) -> [i32; 8] {
4454    let mut acc = 0u64;
4455    for i in 0..B {
4456        acc = (acc << 8) | data[i] as u64;
4457    }
4458    let mask = (1u64 << B) - 1;
4459    let mut out = [0i32; 8];
4460    for (k, o) in out.iter_mut().enumerate() {
4461        *o = ((acc >> ((7 - k) * B)) & mask) as i32;
4462    }
4463    out
4464}
4465
4466/// Fused vbit matvec straight from the mapped bytes (spec §3, P13
4467/// FIG.3): [u8 bits: rows][f16 scales: rows·cols/32][bit-packed rows,
4468/// MSB-first, byte-padded]. Row data offsets are precomputed at load
4469/// (`vbit_row_offsets`) — the per-call prefix scan was O(rows) pure
4470/// overhead on every matvec.
4471#[allow(clippy::too_many_arguments)]
4472fn vbitmatvec(
4473    bytes: &[u8],
4474    offsets: &[usize],
4475    x: &[f32],
4476    rows: usize,
4477    cols: usize,
4478    out: &mut [f32],
4479    pool: Option<&Pool>,
4480) {
4481    debug_assert_eq!(out.len(), rows);
4482    debug_assert_eq!(offsets.len(), rows + 1);
4483
4484    // SDOT path: unpack the row to centered i8 once, then per-group
4485    // int8 dot against the quantized activations — same A8W8 contract
4486    // as q8 (bounded noise; CMF_SDOT=0 keeps the exact scalar path).
4487    if a8w8_enabled() {
4488        let act = split_act(x);
4489        let out_addr = SendMut(out.as_mut_ptr());
4490        let run = move |start: usize, end: usize| {
4491            vbit_range_a8w8(bytes, offsets, x, &act, rows, cols, out_addr, start, end)
4492        };
4493        dispatch_rows(pool, rows, &run);
4494        return;
4495    }
4496
4497    let out_addr = SendMut(out.as_mut_ptr());
4498    let run = move |start: usize, end: usize| {
4499        vbit_range_f32(bytes, offsets, x, rows, cols, out_addr, start, end)
4500    };
4501    dispatch_rows(pool, rows, &run);
4502}
4503
4504/// One vbit row range via the A8W8 int8 path — kernel body of
4505/// `vbitmatvec`, extracted so multi-matrix jobs can drive it for
4506/// several tensors in one dispatch (b=8 rows go exact f32).
4507#[allow(clippy::too_many_arguments)]
4508fn vbit_range_a8w8(
4509    bytes: &[u8],
4510    offsets: &[usize],
4511    x: &[f32],
4512    act: &SplitAct,
4513    rows: usize,
4514    cols: usize,
4515    out: SendMut,
4516    start: usize,
4517    end: usize,
4518) {
4519    let ng = cols / GROUP_SIZE;
4520    let bits = &bytes[..rows];
4521    let sc_off = rows;
4522    let row_dot = |r: usize| -> f32 {
4523        let b = bits[r] as usize;
4524        let l = (1i32 << (b - 1)) - 1;
4525        let mask = (1u64 << b) - 1;
4526        let data = &bytes[offsets[r]..offsets[r + 1]];
4527        if b == 8 {
4528            // u−L reaches 128 → does not fit i8; exact f32 path.
4529            let (mut acc, mut nbits, mut idx) = (0u64, 0usize, 0usize);
4530            let mut dot = 0f32;
4531            for g in 0..ng {
4532                let so = (r * ng + g) * 2;
4533                let sgf = f16_to_f32(u16::from_le_bytes([
4534                    bytes[sc_off + so],
4535                    bytes[sc_off + so + 1],
4536                ]));
4537                let xg = &x[g * GROUP_SIZE..(g + 1) * GROUP_SIZE];
4538                let mut gd = 0f32;
4539                for &xv in xg.iter() {
4540                    if nbits < 8 {
4541                        acc = (acc << 8) | data[idx] as u64;
4542                        idx += 1;
4543                        nbits += 8;
4544                    }
4545                    let u = ((acc >> (nbits - 8)) & 0xFF) as i32;
4546                    nbits -= 8;
4547                    gd += (u - l) as f32 * xv;
4548                }
4549                dot += gd * sgf;
4550            }
4551            return dot;
4552        }
4553        // Per-worker scratch: this closure runs for every row of the
4554        // tensor (lm_head ≈ 150k rows/token) — a heap allocation per
4555        // row was measurable pure overhead.
4556        thread_local! {
4557            static VBIT_SCRATCH: std::cell::RefCell<Vec<u8>> =
4558                const { std::cell::RefCell::new(Vec::new()) };
4559        }
4560        #[inline(always)]
4561        fn fill<const B: usize>(data: &[u8], l: i32, buf: &mut [u8]) {
4562            for (blk, chunk) in buf.chunks_exact_mut(8).enumerate() {
4563                let u = unpack8::<B>(&data[blk * B..]);
4564                for k in 0..8 {
4565                    chunk[k] = (u[k] - l) as i8 as u8;
4566                }
4567            }
4568        }
4569        let _ = mask;
4570        VBIT_SCRATCH.with(|scratch| {
4571            let mut buf = scratch.borrow_mut();
4572            buf.resize(cols, 0);
4573            match b {
4574                3 => fill::<3>(data, l, &mut buf),
4575                4 => vbit_fill4(data, &mut buf),
4576                5 => fill::<5>(data, l, &mut buf),
4577                6 => fill::<6>(data, l, &mut buf),
4578                _ => unreachable!(),
4579            }
4580            let mut dot = 0f32;
4581            for g in 0..ng {
4582                let so = (r * ng + g) * 2;
4583                let s = f16_to_f32(u16::from_le_bytes([
4584                    bytes[sc_off + so],
4585                    bytes[sc_off + so + 1],
4586                ]));
4587                let d = dot_i8_i8(
4588                    &buf[g * GROUP_SIZE..(g + 1) * GROUP_SIZE],
4589                    &act.xq[g * GROUP_SIZE..(g + 1) * GROUP_SIZE],
4590                ) as f32
4591                    * act.sx;
4592                dot += d * s;
4593            }
4594            for &(j, xv) in &act.outliers {
4595                let so = (r * ng + j / GROUP_SIZE) * 2;
4596                let s = f16_to_f32(u16::from_le_bytes([
4597                    bytes[sc_off + so],
4598                    bytes[sc_off + so + 1],
4599                ]));
4600                // xq is zeroed at outlier slots — add the exact term.
4601                dot += (buf[j] as i8) as f32 * s * xv;
4602            }
4603            dot
4604        })
4605    };
4606    for r in start..end {
4607        // SAFETY: disjoint row ranges per worker.
4608        unsafe { *out.at(r) = row_dot(r) };
4609    }
4610}
4611
4612/// Exact scalar vbit row range (same extraction, non-SDOT path).
4613#[allow(clippy::too_many_arguments)]
4614fn vbit_range_f32(
4615    bytes: &[u8],
4616    offsets: &[usize],
4617    x: &[f32],
4618    rows: usize,
4619    cols: usize,
4620    out: SendMut,
4621    start: usize,
4622    end: usize,
4623) {
4624    let ng = cols / GROUP_SIZE;
4625    let bits = &bytes[..rows];
4626    let sc_off = rows;
4627    // Per-bit-width specialized inner loops: the compiler unrolls the
4628    // constant shifts (the generic bit-buffer loop was branch-bound —
4629    // 5.6 vs 13.2 tok/s q4 on the 0.8B).
4630    #[inline(always)]
4631    fn dot_row<const B: usize>(
4632        data: &[u8],
4633        bytes: &[u8],
4634        sc_off: usize,
4635        r: usize,
4636        ng: usize,
4637        x: &[f32],
4638    ) -> f32 {
4639        let l = ((1i32 << (B - 1)) - 1) as f32;
4640        let gbytes = GROUP_SIZE * B / 8;
4641        let mut dot = 0f32;
4642        for g in 0..ng {
4643            let so = (r * ng + g) * 2;
4644            let s = f16_to_f32(u16::from_le_bytes([
4645                bytes[sc_off + so],
4646                bytes[sc_off + so + 1],
4647            ]));
4648            let xg = &x[g * GROUP_SIZE..(g + 1) * GROUP_SIZE];
4649            let gd0 = &data[g * gbytes..(g + 1) * gbytes];
4650            let mut gd = 0f32;
4651            for blk in 0..GROUP_SIZE / 8 {
4652                let u = unpack8::<B>(&gd0[blk * B..]);
4653                let xb = &xg[blk * 8..blk * 8 + 8];
4654                for k in 0..8 {
4655                    gd += (u[k] as f32 - l) * xb[k];
4656                }
4657            }
4658            dot += gd * s;
4659        }
4660        dot
4661    }
4662    for r in start..end {
4663        let data = &bytes[offsets[r]..offsets[r + 1]];
4664        let v = match bits[r] {
4665            3 => dot_row::<3>(data, bytes, sc_off, r, ng, x),
4666            4 => dot_row::<4>(data, bytes, sc_off, r, ng, x),
4667            5 => dot_row::<5>(data, bytes, sc_off, r, ng, x),
4668            6 => dot_row::<6>(data, bytes, sc_off, r, ng, x),
4669            8 => dot_row::<8>(data, bytes, sc_off, r, ng, x),
4670            b => unreachable!("vbit bit-width {b} (validated at load)"),
4671        };
4672        // SAFETY: disjoint row ranges per worker.
4673        unsafe { *out.at(r) = v };
4674    }
4675}
4676
4677/// Fused two-input vbit matvec: each row is unpacked from the mmap ONCE
4678/// and dotted against BOTH activations (MTP verify / pair prefill used
4679/// to run two full matvecs — double weight traffic and double unpack).
4680/// Per-input math is identical to `vbitmatvec` → same accuracy contract.
4681#[allow(clippy::too_many_arguments)]
4682fn vbitmatvec2(
4683    bytes: &[u8],
4684    offsets: &[usize],
4685    x1: &[f32],
4686    x2: &[f32],
4687    rows: usize,
4688    cols: usize,
4689    o1: &mut [f32],
4690    o2: &mut [f32],
4691    pool: Option<&Pool>,
4692) {
4693    debug_assert_eq!(o1.len(), rows);
4694    debug_assert_eq!(o2.len(), rows);
4695
4696    if a8w8_enabled() {
4697        let a1 = split_act(x1);
4698        let a2 = split_act(x2);
4699        let p1 = SendMut(o1.as_mut_ptr());
4700        let p2 = SendMut(o2.as_mut_ptr());
4701        let run = move |start: usize, end: usize| {
4702            vbit_range2_a8w8(
4703                bytes, offsets, x1, x2, &a1, &a2, rows, cols, p1, p2, start, end,
4704            )
4705        };
4706        dispatch_rows(pool, rows, &run);
4707        return;
4708    }
4709
4710    let p1 = SendMut(o1.as_mut_ptr());
4711    let p2 = SendMut(o2.as_mut_ptr());
4712    let run = move |start: usize, end: usize| {
4713        vbit_range2_f32(bytes, offsets, x1, x2, rows, cols, p1, p2, start, end)
4714    };
4715    dispatch_rows(pool, rows, &run);
4716}
4717
4718/// Two-input vbit row range via the A8W8 int8 path — kernel body of
4719/// `vbitmatvec2`, extracted for pair multi-matrix jobs (b=8 rows go
4720/// exact f32 for both lanes, bits streamed once).
4721#[allow(clippy::too_many_arguments)]
4722fn vbit_range2_a8w8(
4723    bytes: &[u8],
4724    offsets: &[usize],
4725    x1: &[f32],
4726    x2: &[f32],
4727    a1: &SplitAct,
4728    a2: &SplitAct,
4729    rows: usize,
4730    cols: usize,
4731    p1: SendMut,
4732    p2: SendMut,
4733    start: usize,
4734    end: usize,
4735) {
4736    let ng = cols / GROUP_SIZE;
4737    let bits = &bytes[..rows];
4738    let sc_off = rows;
4739    let row_dots = |r: usize| -> (f32, f32) {
4740        let b = bits[r] as usize;
4741        let l = (1i32 << (b - 1)) - 1;
4742        let data = &bytes[offsets[r]..offsets[r + 1]];
4743        if b == 8 {
4744            // u−L reaches 128 → does not fit i8; exact f32 path,
4745            // bits still streamed once for both lanes.
4746            let (mut acc, mut nbits, mut idx) = (0u64, 0usize, 0usize);
4747            let (mut d1, mut d2) = (0f32, 0f32);
4748            for g in 0..ng {
4749                let so = (r * ng + g) * 2;
4750                let sgf = f16_to_f32(u16::from_le_bytes([
4751                    bytes[sc_off + so],
4752                    bytes[sc_off + so + 1],
4753                ]));
4754                let (mut g1, mut g2) = (0f32, 0f32);
4755                for k in 0..GROUP_SIZE {
4756                    if nbits < 8 {
4757                        acc = (acc << 8) | data[idx] as u64;
4758                        idx += 1;
4759                        nbits += 8;
4760                    }
4761                    let u = ((acc >> (nbits - 8)) & 0xFF) as i32;
4762                    nbits -= 8;
4763                    let w = (u - l) as f32;
4764                    g1 += w * x1[g * GROUP_SIZE + k];
4765                    g2 += w * x2[g * GROUP_SIZE + k];
4766                }
4767                d1 += g1 * sgf;
4768                d2 += g2 * sgf;
4769            }
4770            return (d1, d2);
4771        }
4772        thread_local! {
4773            static VBIT_SCRATCH2: std::cell::RefCell<Vec<u8>> =
4774                const { std::cell::RefCell::new(Vec::new()) };
4775        }
4776        #[inline(always)]
4777        fn fill<const B: usize>(data: &[u8], l: i32, buf: &mut [u8]) {
4778            for (blk, chunk) in buf.chunks_exact_mut(8).enumerate() {
4779                let u = unpack8::<B>(&data[blk * B..]);
4780                for k in 0..8 {
4781                    chunk[k] = (u[k] - l) as i8 as u8;
4782                }
4783            }
4784        }
4785        VBIT_SCRATCH2.with(|scratch| {
4786            let mut buf = scratch.borrow_mut();
4787            buf.resize(cols, 0);
4788            match b {
4789                3 => fill::<3>(data, l, &mut buf),
4790                4 => vbit_fill4(data, &mut buf),
4791                5 => fill::<5>(data, l, &mut buf),
4792                6 => fill::<6>(data, l, &mut buf),
4793                _ => unreachable!(),
4794            }
4795            let (mut d1, mut d2) = (0f32, 0f32);
4796            for g in 0..ng {
4797                let so = (r * ng + g) * 2;
4798                let s = f16_to_f32(u16::from_le_bytes([
4799                    bytes[sc_off + so],
4800                    bytes[sc_off + so + 1],
4801                ]));
4802                let wg = &buf[g * GROUP_SIZE..(g + 1) * GROUP_SIZE];
4803                let v1 = dot_i8_i8(wg, &a1.xq[g * GROUP_SIZE..(g + 1) * GROUP_SIZE]) as f32 * a1.sx;
4804                let v2 = dot_i8_i8(wg, &a2.xq[g * GROUP_SIZE..(g + 1) * GROUP_SIZE]) as f32 * a2.sx;
4805                d1 += v1 * s;
4806                d2 += v2 * s;
4807            }
4808            for &(j, xv) in &a1.outliers {
4809                let so = (r * ng + j / GROUP_SIZE) * 2;
4810                let s = f16_to_f32(u16::from_le_bytes([
4811                    bytes[sc_off + so],
4812                    bytes[sc_off + so + 1],
4813                ]));
4814                d1 += (buf[j] as i8) as f32 * s * xv;
4815            }
4816            for &(j, xv) in &a2.outliers {
4817                let so = (r * ng + j / GROUP_SIZE) * 2;
4818                let s = f16_to_f32(u16::from_le_bytes([
4819                    bytes[sc_off + so],
4820                    bytes[sc_off + so + 1],
4821                ]));
4822                d2 += (buf[j] as i8) as f32 * s * xv;
4823            }
4824            (d1, d2)
4825        })
4826    };
4827    for r in start..end {
4828        let (v1, v2) = row_dots(r);
4829        // SAFETY: disjoint row ranges per worker.
4830        unsafe {
4831            *p1.at(r) = v1;
4832            *p2.at(r) = v2;
4833        }
4834    }
4835}
4836
4837/// Two-input exact scalar vbit row range (same extraction) —
4838/// per-bit-width specialized, two accumulators per row; per-lane
4839/// accumulation order matches `vbitmatvec` exactly.
4840#[allow(clippy::too_many_arguments)]
4841fn vbit_range2_f32(
4842    bytes: &[u8],
4843    offsets: &[usize],
4844    x1: &[f32],
4845    x2: &[f32],
4846    rows: usize,
4847    cols: usize,
4848    p1: SendMut,
4849    p2: SendMut,
4850    start: usize,
4851    end: usize,
4852) {
4853    let ng = cols / GROUP_SIZE;
4854    let bits = &bytes[..rows];
4855    let sc_off = rows;
4856    #[inline(always)]
4857    #[allow(clippy::too_many_arguments)]
4858    fn dot_row2<const B: usize>(
4859        data: &[u8],
4860        bytes: &[u8],
4861        sc_off: usize,
4862        r: usize,
4863        ng: usize,
4864        x1: &[f32],
4865        x2: &[f32],
4866    ) -> (f32, f32) {
4867        let l = ((1i32 << (B - 1)) - 1) as f32;
4868        let gbytes = GROUP_SIZE * B / 8;
4869        let (mut d1, mut d2) = (0f32, 0f32);
4870        for g in 0..ng {
4871            let so = (r * ng + g) * 2;
4872            let s = f16_to_f32(u16::from_le_bytes([
4873                bytes[sc_off + so],
4874                bytes[sc_off + so + 1],
4875            ]));
4876            let x1g = &x1[g * GROUP_SIZE..(g + 1) * GROUP_SIZE];
4877            let x2g = &x2[g * GROUP_SIZE..(g + 1) * GROUP_SIZE];
4878            let gd0 = &data[g * gbytes..(g + 1) * gbytes];
4879            let (mut g1, mut g2) = (0f32, 0f32);
4880            for blk in 0..GROUP_SIZE / 8 {
4881                let u = unpack8::<B>(&gd0[blk * B..]);
4882                for k in 0..8 {
4883                    let w = u[k] as f32 - l;
4884                    g1 += w * x1g[blk * 8 + k];
4885                    g2 += w * x2g[blk * 8 + k];
4886                }
4887            }
4888            d1 += g1 * s;
4889            d2 += g2 * s;
4890        }
4891        (d1, d2)
4892    }
4893    for r in start..end {
4894        let data = &bytes[offsets[r]..offsets[r + 1]];
4895        let (v1, v2) = match bits[r] {
4896            3 => dot_row2::<3>(data, bytes, sc_off, r, ng, x1, x2),
4897            4 => dot_row2::<4>(data, bytes, sc_off, r, ng, x1, x2),
4898            5 => dot_row2::<5>(data, bytes, sc_off, r, ng, x1, x2),
4899            6 => dot_row2::<6>(data, bytes, sc_off, r, ng, x1, x2),
4900            8 => dot_row2::<8>(data, bytes, sc_off, r, ng, x1, x2),
4901            b => unreachable!("vbit bit-width {b} (validated at load)"),
4902        };
4903        // SAFETY: disjoint row ranges per worker.
4904        unsafe {
4905            *p1.at(r) = v1;
4906            *p2.at(r) = v2;
4907        }
4908    }
4909}
4910
4911// ───────────────────── q4_tiled kernels (§4.3) ─────────────────────
4912
4913/// One q4_tiled row dot on the A8W8 int8 path: per 32-group the tile
4914/// is ONE sequential read — [f16 scale][16B nibbles] — versus the two
4915/// distant streams of the split layout. Values/order identical to the
4916/// split kernels.
4917#[inline]
4918#[allow(unreachable_code)]
4919fn dot_q4t_row_i8(bytes: &[u8], r: usize, gpr: usize, xq: &[i8]) -> f32 {
4920    #[cfg(target_arch = "aarch64")]
4921    unsafe {
4922        return dot_q4t_row_sdot(bytes, r, gpr, xq);
4923    }
4924    #[cfg(target_arch = "x86_64")]
4925    unsafe {
4926        if vnni_tiles_enabled() {
4927            return dot_q4t_row_vnni(bytes, r, gpr, xq);
4928        }
4929        return dot_q4t_row_avx2(bytes, r, gpr, xq);
4930    }
4931    let mut acc = 0f32;
4932    for gi in 0..gpr {
4933        let tile = &bytes[(r * gpr + gi) * Q4_TILE..(r * gpr + gi + 1) * Q4_TILE];
4934        let s = f16_to_f32(u16::from_le_bytes([tile[0], tile[1]]));
4935        let mut d = 0i32;
4936        for (k, &b) in tile[2..].iter().enumerate() {
4937            d += ((b & 0x0F) as i32 - 8) * xq[gi * GROUP_SIZE + k * 2] as i32
4938                + (((b >> 4) & 0x0F) as i32 - 8) * xq[gi * GROUP_SIZE + k * 2 + 1] as i32;
4939        }
4940        acc += d as f32 * s;
4941    }
4942    acc
4943}
4944
4945#[cfg(target_arch = "aarch64")]
4946#[target_feature(enable = "neon,dotprod")]
4947unsafe fn dot_q4t_row_sdot(bytes: &[u8], r: usize, gpr: usize, xq: &[i8]) -> f32 {
4948    // SAFETY: callers uphold slice-length contracts (18B tile per group,
4949    // xq.len() == gpr·GROUP_SIZE).
4950    unsafe {
4951        use core::arch::aarch64::*;
4952        use core::arch::asm;
4953        let lomask = vdupq_n_u8(0x0F);
4954        let eight = vdupq_n_s8(8);
4955        let mut acc = 0f32;
4956        for gi in 0..gpr {
4957            let t = bytes.as_ptr().add((r * gpr + gi) * Q4_TILE);
4958            let s = f16_to_f32(u16::from_le_bytes([*t, *t.add(1)]));
4959            let b = vld1q_u8(t.add(2));
4960            let lo = vandq_u8(b, lomask);
4961            let hi = vshrq_n_u8::<4>(b);
4962            let e0 = vsubq_s8(vreinterpretq_s8_u8(vzip1q_u8(lo, hi)), eight);
4963            let e1 = vsubq_s8(vreinterpretq_s8_u8(vzip2q_u8(lo, hi)), eight);
4964            let x0 = vld1q_s8(xq.as_ptr().add(gi * GROUP_SIZE));
4965            let x1 = vld1q_s8(xq.as_ptr().add(gi * GROUP_SIZE + 16));
4966            let (mut a0, mut a1) = (vdupq_n_s32(0), vdupq_n_s32(0));
4967            asm!(
4968                "sdot {a0:v}.4s, {e0:v}.16b, {x0:v}.16b",
4969                "sdot {a1:v}.4s, {e1:v}.16b, {x1:v}.16b",
4970                a0 = inout(vreg) a0, a1 = inout(vreg) a1,
4971                e0 = in(vreg) e0, x0 = in(vreg) x0, e1 = in(vreg) e1, x1 = in(vreg) x1,
4972                options(pure, nomem, nostack),
4973            );
4974            acc += vaddvq_s32(vaddq_s32(a0, a1)) as f32 * s;
4975        }
4976        acc
4977    }
4978}
4979
4980#[cfg(target_arch = "x86_64")]
4981#[target_feature(enable = "avx2")]
4982unsafe fn dot_q4t_row_avx2(bytes: &[u8], r: usize, gpr: usize, xq: &[i8]) -> f32 {
4983    // SAFETY: see dot_q4t_row_sdot.
4984    unsafe {
4985        use core::arch::x86_64::*;
4986        let lomask = _mm_set1_epi8(0x0F);
4987        let eight = _mm256_set1_epi8(8);
4988        let ones = _mm256_set1_epi16(1);
4989        let mut acc = 0f32;
4990        for gi in 0..gpr {
4991            let t = bytes.as_ptr().add((r * gpr + gi) * Q4_TILE);
4992            let s = f16_to_f32(u16::from_le_bytes([*t, *t.add(1)]));
4993            let b = _mm_loadu_si128(t.add(2) as *const __m128i);
4994            let lo = _mm_and_si128(b, lomask);
4995            let hi = _mm_and_si128(_mm_srli_epi16::<4>(b), lomask);
4996            let w = _mm256_sub_epi8(
4997                _mm256_set_m128i(_mm_unpackhi_epi8(lo, hi), _mm_unpacklo_epi8(lo, hi)),
4998                eight,
4999            );
5000            let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
5001            let p16 = _mm256_maddubs_epi16(_mm256_abs_epi8(w), _mm256_sign_epi8(x, w));
5002            let d = _mm256_madd_epi16(p16, ones);
5003            let hi128 = _mm256_extracti128_si256::<1>(d);
5004            let s128 = _mm_add_epi32(_mm256_castsi256_si128(d), hi128);
5005            let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
5006            let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
5007            acc += _mm_cvtsi128_si32(s32) as f32 * s;
5008        }
5009        acc
5010    }
5011}
5012
5013/// VNNI twin of `dot_q4t_row_avx2`: same unpack, `vpdpbusd` replaces
5014/// the maddubs+madd pair (see `dpbusd_hsum` — sums are bit-identical).
5015/// 256-bit VL encoding, so the VEX `vpsignb` stays usable.
5016#[cfg(target_arch = "x86_64")]
5017#[target_feature(enable = "avx2,avx512f,avx512bw,avx512vl,avx512vnni")]
5018unsafe fn dot_q4t_row_vnni(bytes: &[u8], r: usize, gpr: usize, xq: &[i8]) -> f32 {
5019    // SAFETY: see dot_q4t_row_sdot.
5020    unsafe {
5021        use core::arch::x86_64::*;
5022        let lomask = _mm_set1_epi8(0x0F);
5023        let eight = _mm256_set1_epi8(8);
5024        let mut acc = 0f32;
5025        for gi in 0..gpr {
5026            let t = bytes.as_ptr().add((r * gpr + gi) * Q4_TILE);
5027            let s = f16_to_f32(u16::from_le_bytes([*t, *t.add(1)]));
5028            let b = _mm_loadu_si128(t.add(2) as *const __m128i);
5029            let lo = _mm_and_si128(b, lomask);
5030            let hi = _mm_and_si128(_mm_srli_epi16::<4>(b), lomask);
5031            let w = _mm256_sub_epi8(
5032                _mm256_set_m128i(_mm_unpackhi_epi8(lo, hi), _mm_unpacklo_epi8(lo, hi)),
5033                eight,
5034            );
5035            let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
5036            let d = dpbusd_hsum(_mm256_abs_epi8(w), _mm256_sign_epi8(x, w));
5037            acc += d as f32 * s;
5038        }
5039        acc
5040    }
5041}
5042
5043/// One q4_tiled row against FOUR activation streams: the nibble unpack
5044/// and abs() happen once per group instead of once per (group,
5045/// activation) — the unpack is the dominant per-element cost of the
5046/// tiled format (roadmap P0 portable blocking, q4t leg).
5047#[cfg(target_arch = "x86_64")]
5048// `fma` is NOT implied by `avx2`: without it LLVM lowers _mm256_fmadd_ps
5049// to a libm call per lane — measured 2x slower than the reduction this
5050// kernel replaces. The runtime gate (`avx2_enabled`) already requires
5051// both features, so declaring it here is safe.
5052#[target_feature(enable = "avx2,fma")]
5053unsafe fn dot_q4t_row_1x4_avx2(bytes: &[u8], r: usize, gpr: usize, xs: [&[i8]; 4]) -> [f32; 4] {
5054    // SAFETY: callers uphold the 18B-tile and xq-length contracts.
5055    unsafe {
5056        use core::arch::x86_64::*;
5057        let lomask = _mm_set1_epi8(0x0F);
5058        let eight = _mm256_set1_epi8(8);
5059        let ones = _mm256_set1_epi16(1);
5060        // One f32 accumulator VECTOR per activation, reduced once at the
5061        // end. Folding each group's i32 lanes to a scalar inside the loop
5062        // costs an extracti128 + three shift/add + a movd — a cross-lane
5063        // dependency chain per (group, activation), 288 of them per row at
5064        // cols=2304. The per-group scale is what forces a float
5065        // accumulator; it does not force a horizontal sum.
5066        //
5067        // The four accumulators are NAMED, not an array: as `[__m256; 4]`
5068        // indexed by a loop variable LLVM keeps them in memory and every
5069        // group pays four 32-byte loads and stores. That alone made this
5070        // kernel 2x SLOWER than the per-group reduction it replaces
5071        // (measured on the EPYC box: 150 s vs 71 s for two 256² steps).
5072        let mut f0 = _mm256_setzero_ps();
5073        let mut f1 = _mm256_setzero_ps();
5074        let mut f2 = _mm256_setzero_ps();
5075        let mut f3 = _mm256_setzero_ps();
5076        for gi in 0..gpr {
5077            let t = bytes.as_ptr().add((r * gpr + gi) * Q4_TILE);
5078            let s = f16_to_f32(u16::from_le_bytes([*t, *t.add(1)]));
5079            let sv = _mm256_set1_ps(s);
5080            let bb = _mm_loadu_si128(t.add(2) as *const __m128i);
5081            let lo = _mm_and_si128(bb, lomask);
5082            let hi = _mm_and_si128(_mm_srli_epi16::<4>(bb), lomask);
5083            let w = _mm256_sub_epi8(
5084                _mm256_set_m128i(_mm_unpackhi_epi8(lo, hi), _mm_unpacklo_epi8(lo, hi)),
5085                eight,
5086            );
5087            let aw = _mm256_abs_epi8(w);
5088            let off = gi * GROUP_SIZE;
5089            let dot = |xq: &[i8]| {
5090                let x = _mm256_loadu_si256(xq.as_ptr().add(off) as *const __m256i);
5091                let p16 = _mm256_maddubs_epi16(aw, _mm256_sign_epi8(x, w));
5092                _mm256_cvtepi32_ps(_mm256_madd_epi16(p16, ones))
5093            };
5094            f0 = _mm256_fmadd_ps(dot(xs[0]), sv, f0);
5095            f1 = _mm256_fmadd_ps(dot(xs[1]), sv, f1);
5096            f2 = _mm256_fmadd_ps(dot(xs[2]), sv, f2);
5097            f3 = _mm256_fmadd_ps(dot(xs[3]), sv, f3);
5098        }
5099        [
5100            hsum256_ps(f0),
5101            hsum256_ps(f1),
5102            hsum256_ps(f2),
5103            hsum256_ps(f3),
5104        ]
5105    }
5106}
5107
5108/// Horizontal sum of eight f32 lanes — the one cross-lane reduction the
5109/// blocked kernels pay, once per row instead of once per group.
5110#[cfg(target_arch = "x86_64")]
5111#[target_feature(enable = "avx2")]
5112#[inline]
5113unsafe fn hsum256_ps(v: core::arch::x86_64::__m256) -> f32 {
5114    // SAFETY: pure register arithmetic on the caller's vector.
5115    unsafe {
5116        use core::arch::x86_64::*;
5117        let hi = _mm256_extractf128_ps::<1>(v);
5118        let s = _mm_add_ps(_mm256_castps256_ps128(v), hi);
5119        let s = _mm_add_ps(s, _mm_movehl_ps(s, s));
5120        let s = _mm_add_ss(s, _mm_shuffle_ps::<0x55>(s, s));
5121        _mm_cvtss_f32(s)
5122    }
5123}
5124
5125/// VNNI twin of `dot_q4t_row_1x4_avx2` (see `dpbusd_hsum`).
5126#[cfg(target_arch = "x86_64")]
5127#[target_feature(enable = "avx2,fma,avx512f,avx512bw,avx512vl,avx512vnni")]
5128unsafe fn dot_q4t_row_1x4_vnni(bytes: &[u8], r: usize, gpr: usize, xs: [&[i8]; 4]) -> [f32; 4] {
5129    // SAFETY: callers uphold the 18B-tile and xq-length contracts.
5130    unsafe {
5131        use core::arch::x86_64::*;
5132        let lomask = _mm_set1_epi8(0x0F);
5133        let eight = _mm256_set1_epi8(8);
5134        // Same shape as the AVX2 twin: accumulate in f32 vectors and pay
5135        // one cross-lane reduction per row, not per (group, activation).
5136        let mut f0 = _mm256_setzero_ps();
5137        let mut f1 = _mm256_setzero_ps();
5138        let mut f2 = _mm256_setzero_ps();
5139        let mut f3 = _mm256_setzero_ps();
5140        for gi in 0..gpr {
5141            let t = bytes.as_ptr().add((r * gpr + gi) * Q4_TILE);
5142            let s = f16_to_f32(u16::from_le_bytes([*t, *t.add(1)]));
5143            let sv = _mm256_set1_ps(s);
5144            let bb = _mm_loadu_si128(t.add(2) as *const __m128i);
5145            let lo = _mm_and_si128(bb, lomask);
5146            let hi = _mm_and_si128(_mm_srli_epi16::<4>(bb), lomask);
5147            let w = _mm256_sub_epi8(
5148                _mm256_set_m128i(_mm_unpackhi_epi8(lo, hi), _mm_unpacklo_epi8(lo, hi)),
5149                eight,
5150            );
5151            let aw = _mm256_abs_epi8(w);
5152            let off = gi * GROUP_SIZE;
5153            let dot = |xq: &[i8]| {
5154                let x = _mm256_loadu_si256(xq.as_ptr().add(off) as *const __m256i);
5155                _mm256_cvtepi32_ps(_mm256_dpbusd_epi32(
5156                    _mm256_setzero_si256(),
5157                    aw,
5158                    _mm256_sign_epi8(x, w),
5159                ))
5160            };
5161            f0 = _mm256_fmadd_ps(dot(xs[0]), sv, f0);
5162            f1 = _mm256_fmadd_ps(dot(xs[1]), sv, f1);
5163            f2 = _mm256_fmadd_ps(dot(xs[2]), sv, f2);
5164            f3 = _mm256_fmadd_ps(dot(xs[3]), sv, f3);
5165        }
5166        let acc = [
5167            hsum256_ps(f0),
5168            hsum256_ps(f1),
5169            hsum256_ps(f2),
5170            hsum256_ps(f3),
5171        ];
5172        acc
5173    }
5174}
5175
5176/// ARM twin of `dot_q4t_row_1x4_avx2`: one nibble unpack per group
5177/// serves FOUR activation streams. Per stream the group order and f32
5178/// accumulation match `dot_q4t_row_sdot` exactly — batch == matvec
5179/// bit-for-bit.
5180#[cfg(target_arch = "aarch64")]
5181#[target_feature(enable = "neon,dotprod")]
5182unsafe fn dot_q4t_row_1x4_sdot(bytes: &[u8], r: usize, gpr: usize, xs: [&[i8]; 4]) -> [f32; 4] {
5183    // SAFETY: callers uphold the 18B-tile and xq-length contracts.
5184    unsafe {
5185        use core::arch::aarch64::*;
5186        use core::arch::asm;
5187        let lomask = vdupq_n_u8(0x0F);
5188        let eight = vdupq_n_s8(8);
5189        let mut acc = [0f32; 4];
5190        for gi in 0..gpr {
5191            let t = bytes.as_ptr().add((r * gpr + gi) * Q4_TILE);
5192            let s = f16_to_f32(u16::from_le_bytes([*t, *t.add(1)]));
5193            let b = vld1q_u8(t.add(2));
5194            let lo = vandq_u8(b, lomask);
5195            let hi = vshrq_n_u8::<4>(b);
5196            let e0 = vsubq_s8(vreinterpretq_s8_u8(vzip1q_u8(lo, hi)), eight);
5197            let e1 = vsubq_s8(vreinterpretq_s8_u8(vzip2q_u8(lo, hi)), eight);
5198            for (k, xq) in xs.iter().enumerate() {
5199                let x0 = vld1q_s8(xq.as_ptr().add(gi * GROUP_SIZE));
5200                let x1 = vld1q_s8(xq.as_ptr().add(gi * GROUP_SIZE + 16));
5201                let (mut a0, mut a1) = (vdupq_n_s32(0), vdupq_n_s32(0));
5202                asm!(
5203                    "sdot {a0:v}.4s, {e0:v}.16b, {x0:v}.16b",
5204                    "sdot {a1:v}.4s, {e1:v}.16b, {x1:v}.16b",
5205                    a0 = inout(vreg) a0, a1 = inout(vreg) a1,
5206                    e0 = in(vreg) e0, x0 = in(vreg) x0, e1 = in(vreg) e1, x1 = in(vreg) x1,
5207                    options(pure, nomem, nostack),
5208                );
5209                acc[k] += vaddvq_s32(vaddq_s32(a0, a1)) as f32 * s;
5210            }
5211        }
5212        acc
5213    }
5214}
5215
5216/// Exact-term correction for A8W8 outliers on a tiled row.
5217#[inline]
5218fn q4t_outlier(bytes: &[u8], r: usize, gpr: usize, j: usize) -> (f32, f32) {
5219    let gi = j / GROUP_SIZE;
5220    let k = j % GROUP_SIZE;
5221    let tile = &bytes[(r * gpr + gi) * Q4_TILE..(r * gpr + gi + 1) * Q4_TILE];
5222    let s = f16_to_f32(u16::from_le_bytes([tile[0], tile[1]]));
5223    let byte = tile[2 + k / 2];
5224    let nib = if k & 1 == 0 { byte & 0x0F } else { byte >> 4 };
5225    ((nib as i32 - 8) as f32, s)
5226}
5227
5228/// Exact scalar q4_tiled row (CMF_SDOT=0 contract) — same pairwise
5229/// accumulation shape as `q4_range_f32`.
5230#[inline]
5231fn q4t_row_exact(bytes: &[u8], r: usize, gpr: usize, x: &[f32]) -> f32 {
5232    let mut acc = 0f32;
5233    for gi in 0..gpr {
5234        let tile = &bytes[(r * gpr + gi) * Q4_TILE..(r * gpr + gi + 1) * Q4_TILE];
5235        let s = f16_to_f32(u16::from_le_bytes([tile[0], tile[1]]));
5236        let xg = &x[gi * GROUP_SIZE..(gi + 1) * GROUP_SIZE];
5237        let mut ga = 0f32;
5238        for (k, &b) in tile[2..].iter().enumerate() {
5239            ga += ((b & 0x0F) as f32 - 8.0) * xg[k * 2]
5240                + (((b >> 4) & 0x0F) as f32 - 8.0) * xg[k * 2 + 1];
5241        }
5242        acc += ga * s;
5243    }
5244    acc
5245}
5246
5247/// Split view of a `q4tp` payload. The three planes are resolved once per
5248/// matvec instead of per row — `q4tp_sections` is cheap, but doing it inside
5249/// the row loop would put a division on the hot path for nothing.
5250struct Q4tpView<'a> {
5251    nib: &'a [u8],
5252    params: &'a [u8],
5253    codes: &'a [u8],
5254    stride: usize,
5255    /// q2tp reads the ladder with rung 0 = exact zero.
5256    zero_rung: bool,
5257}
5258
5259impl<'a> Q4tpView<'a> {
5260    fn new(bytes: &'a [u8], rows: usize, cols: usize) -> Self {
5261        let (params_off, codes_off, stride) = q4tp_sections(rows, cols);
5262        Self {
5263            nib: &bytes[..params_off],
5264            params: &bytes[params_off..codes_off],
5265            codes: &bytes[codes_off..],
5266            stride,
5267            zero_rung: false,
5268        }
5269    }
5270
5271    /// The q2tp view: identical params/codes planes, 8 B weight chunks.
5272    fn new_q2(bytes: &'a [u8], rows: usize, cols: usize) -> Self {
5273        let (params_off, codes_off, stride) = q2tp_sections(rows, cols);
5274        Self {
5275            nib: &bytes[..params_off],
5276            params: &bytes[params_off..codes_off],
5277            codes: &bytes[codes_off..],
5278            stride,
5279            zero_rung: true,
5280        }
5281    }
5282
5283    /// Expand row `r`'s per-tile scales into `out` (length `gpr`).
5284    ///
5285    /// Doing this once per row — rather than decoding a 5-bit code inside the
5286    /// tile loop — is what makes the format free at runtime. Random access to
5287    /// a packed 5-bit field costs a division, two bounds checks and a branch;
5288    /// the tile's actual work is two `sdot`s, so per-tile decoding dominated
5289    /// the kernel and cost 5x (measured: 1.4 vs 6.9 tok/s on Nanbeige-3B).
5290    /// Walking the plane sequentially with a bit accumulator is ~3 ops.
5291    /// Eight 5-bit codes are exactly five bytes, so a whole group of
5292    /// eight decodes from one little-endian word at fixed shifts. The
5293    /// bit-accumulator this replaces carried a data-dependent `while
5294    /// have < 5` refill whose branch sat in the innermost loop of every
5295    /// q4tp row; a decode profile put this function above the dot
5296    /// products it feeds. Same bitstream, same codes — just no branch
5297    /// and eight independent extractions.
5298    #[inline]
5299    fn scales_into(&self, r: usize, gpr: usize, out: &mut [f32]) {
5300        let tab = if self.zero_rung {
5301            q2tp_ladder(self.params, r)
5302        } else {
5303            q4tp_ladder(self.params, r)
5304        };
5305        let codes = &self.codes[r * self.stride..(r + 1) * self.stride];
5306        let out = &mut out[..gpr];
5307        let mut chunks = out.chunks_exact_mut(8);
5308        let mut ci = 0usize;
5309        for c in &mut chunks {
5310            let w = u64::from(codes[ci])
5311                | u64::from(codes[ci + 1]) << 8
5312                | u64::from(codes[ci + 2]) << 16
5313                | u64::from(codes[ci + 3]) << 24
5314                | u64::from(codes[ci + 4]) << 32;
5315            for (k, o) in c.iter_mut().enumerate() {
5316                *o = tab[((w >> (5 * k)) & 31) as usize];
5317            }
5318            ci += 5;
5319        }
5320        // Fewer than eight codes left: the shared total accessor, which
5321        // tolerates a 5-bit field whose spill byte is past the stride.
5322        let tail = &codes[ci..];
5323        for (k, o) in chunks.into_remainder().iter_mut().enumerate() {
5324            *o = tab[q4tp_code(tail, k)];
5325        }
5326    }
5327}
5328
5329#[inline]
5330fn dot_q4tp_row_i8(nib: &[u8], r: usize, gpr: usize, xq: &[i8], scales: &[f32]) -> f32 {
5331    #[cfg(target_arch = "aarch64")]
5332    unsafe {
5333        return dot_q4tp_row_sdot(nib, r, gpr, xq, scales);
5334    }
5335    #[cfg(target_arch = "x86_64")]
5336    unsafe {
5337        if vnni_tiles_enabled() {
5338            return dot_q4tp_row_vnni(nib, r, gpr, xq, scales);
5339        }
5340        return dot_q4tp_row_avx2(nib, r, gpr, xq, scales);
5341    }
5342    #[allow(unreachable_code)]
5343    {
5344        let mut acc = 0f32;
5345        for gi in 0..gpr {
5346            let tile = &nib[(r * gpr + gi) * Q4TP_NIB..(r * gpr + gi + 1) * Q4TP_NIB];
5347            let s = scales[gi];
5348            let mut d = 0i32;
5349            for (k, &b) in tile.iter().enumerate() {
5350                d += ((b & 0x0F) as i32 - 8) * xq[gi * GROUP_SIZE + k * 2] as i32
5351                    + (((b >> 4) & 0x0F) as i32 - 8) * xq[gi * GROUP_SIZE + k * 2 + 1] as i32;
5352            }
5353            acc += d as f32 * s;
5354        }
5355        acc
5356    }
5357}
5358
5359/// q4tp twin of `dot_q4t_row_sdot`: identical nibble math, but the tile
5360/// stride is 16 B (no inline scale) and the scale is a ladder lookup.
5361#[cfg(target_arch = "aarch64")]
5362#[target_feature(enable = "neon,dotprod")]
5363unsafe fn dot_q4tp_row_sdot(nib: &[u8], r: usize, gpr: usize, xq: &[i8], scales: &[f32]) -> f32 {
5364    // SAFETY: callers uphold slice-length contracts (16B tile per group,
5365    // xq.len() == gpr·GROUP_SIZE, codes covering gpr 5-bit fields).
5366    unsafe {
5367        use core::arch::aarch64::*;
5368        use core::arch::asm;
5369        let lomask = vdupq_n_u8(0x0F);
5370        let eight = vdupq_n_s8(8);
5371        let mut acc = 0f32;
5372        for gi in 0..gpr {
5373            let t = nib.as_ptr().add((r * gpr + gi) * Q4TP_NIB);
5374            let s = *scales.get_unchecked(gi);
5375            let b = vld1q_u8(t);
5376            let lo = vandq_u8(b, lomask);
5377            let hi = vshrq_n_u8::<4>(b);
5378            let e0 = vsubq_s8(vreinterpretq_s8_u8(vzip1q_u8(lo, hi)), eight);
5379            let e1 = vsubq_s8(vreinterpretq_s8_u8(vzip2q_u8(lo, hi)), eight);
5380            let x0 = vld1q_s8(xq.as_ptr().add(gi * GROUP_SIZE));
5381            let x1 = vld1q_s8(xq.as_ptr().add(gi * GROUP_SIZE + 16));
5382            let (mut a0, mut a1) = (vdupq_n_s32(0), vdupq_n_s32(0));
5383            asm!(
5384                "sdot {a0:v}.4s, {e0:v}.16b, {x0:v}.16b",
5385                "sdot {a1:v}.4s, {e1:v}.16b, {x1:v}.16b",
5386                a0 = inout(vreg) a0, a1 = inout(vreg) a1,
5387                e0 = in(vreg) e0, x0 = in(vreg) x0, e1 = in(vreg) e1, x1 = in(vreg) x1,
5388                options(pure, nomem, nostack),
5389            );
5390            acc += vaddvq_s32(vaddq_s32(a0, a1)) as f32 * s;
5391        }
5392        acc
5393    }
5394}
5395
5396#[cfg(target_arch = "x86_64")]
5397#[target_feature(enable = "avx2")]
5398unsafe fn dot_q4tp_row_avx2(nib: &[u8], r: usize, gpr: usize, xq: &[i8], scales: &[f32]) -> f32 {
5399    // SAFETY: see dot_q4tp_row_sdot.
5400    unsafe {
5401        use core::arch::x86_64::*;
5402        let lomask = _mm_set1_epi8(0x0F);
5403        let eight = _mm256_set1_epi8(8);
5404        let ones = _mm256_set1_epi16(1);
5405        let mut acc = 0f32;
5406        for gi in 0..gpr {
5407            let t = nib.as_ptr().add((r * gpr + gi) * Q4TP_NIB);
5408            let s = *scales.get_unchecked(gi);
5409            let b = _mm_loadu_si128(t as *const __m128i);
5410            let lo = _mm_and_si128(b, lomask);
5411            let hi = _mm_and_si128(_mm_srli_epi16::<4>(b), lomask);
5412            let w = _mm256_sub_epi8(
5413                _mm256_set_m128i(_mm_unpackhi_epi8(lo, hi), _mm_unpacklo_epi8(lo, hi)),
5414                eight,
5415            );
5416            let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
5417            let p16 = _mm256_maddubs_epi16(_mm256_abs_epi8(w), _mm256_sign_epi8(x, w));
5418            let d = _mm256_madd_epi16(p16, ones);
5419            let hi128 = _mm256_extracti128_si256::<1>(d);
5420            let s128 = _mm_add_epi32(_mm256_castsi256_si128(d), hi128);
5421            let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
5422            let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
5423            acc += _mm_cvtsi128_si32(s32) as f32 * s;
5424        }
5425        acc
5426    }
5427}
5428
5429/// VNNI twin of `dot_q4tp_row_avx2` (see `dot_q4t_row_vnni` for why the
5430/// 256-bit VL encoding is the one to use here).
5431#[cfg(target_arch = "x86_64")]
5432#[target_feature(enable = "avx2,avx512f,avx512bw,avx512vl,avx512vnni")]
5433unsafe fn dot_q4tp_row_vnni(nib: &[u8], r: usize, gpr: usize, xq: &[i8], scales: &[f32]) -> f32 {
5434    // SAFETY: see dot_q4tp_row_sdot.
5435    unsafe {
5436        use core::arch::x86_64::*;
5437        let lomask = _mm_set1_epi8(0x0F);
5438        let eight = _mm256_set1_epi8(8);
5439        let mut acc = 0f32;
5440        for gi in 0..gpr {
5441            let t = nib.as_ptr().add((r * gpr + gi) * Q4TP_NIB);
5442            let s = *scales.get_unchecked(gi);
5443            let b = _mm_loadu_si128(t as *const __m128i);
5444            let lo = _mm_and_si128(b, lomask);
5445            let hi = _mm_and_si128(_mm_srli_epi16::<4>(b), lomask);
5446            let w = _mm256_sub_epi8(
5447                _mm256_set_m128i(_mm_unpackhi_epi8(lo, hi), _mm_unpacklo_epi8(lo, hi)),
5448                eight,
5449            );
5450            let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
5451            acc += dpbusd_hsum(_mm256_abs_epi8(w), _mm256_sign_epi8(x, w)) as f32 * s;
5452        }
5453        acc
5454    }
5455}
5456
5457/// Exact scalar q4tp row — the `CMF_SDOT=0` contract, same pairwise
5458/// accumulation shape as `q4t_row_exact`.
5459#[inline]
5460fn q4tp_row_exact(nib: &[u8], r: usize, gpr: usize, x: &[f32], scales: &[f32]) -> f32 {
5461    #[cfg(target_arch = "x86_64")]
5462    if avx2_enabled() {
5463        // Keep the scalar pair/group reduction order, not a vector sum.
5464        return unsafe { q4tp_row_float_avx2(nib, r, gpr, x, scales) };
5465    }
5466    q4tp_row_float_scalar(nib, r, gpr, x, scales)
5467}
5468
5469#[inline]
5470fn q4tp_row_float_scalar(nib: &[u8], r: usize, gpr: usize, x: &[f32], scales: &[f32]) -> f32 {
5471    let mut acc = 0f32;
5472    for gi in 0..gpr {
5473        let tile = &nib[(r * gpr + gi) * Q4TP_NIB..(r * gpr + gi + 1) * Q4TP_NIB];
5474        let s = scales[gi];
5475        let xg = &x[gi * GROUP_SIZE..(gi + 1) * GROUP_SIZE];
5476        let mut ga = 0f32;
5477        for (k, &b) in tile.iter().enumerate() {
5478            ga += ((b & 0x0F) as f32 - 8.0) * xg[k * 2]
5479                + (((b >> 4) & 0x0F) as f32 - 8.0) * xg[k * 2 + 1];
5480        }
5481        acc += ga * s;
5482    }
5483    acc
5484}
5485
5486/// Vectorize unpack, conversion and multiplication, but preserve every
5487/// pair addition and the scalar accumulation order. No activation rounding
5488/// or FMA: bit-identical to the float scalar row, including its group scale.
5489#[cfg(target_arch = "x86_64")]
5490#[target_feature(enable = "avx2")]
5491unsafe fn q4tp_row_float_avx2(nib: &[u8], r: usize, gpr: usize, x: &[f32], scales: &[f32]) -> f32 {
5492    // SAFETY: caller checks AVX2 and provides the same complete 32-element
5493    // groups as the scalar row. Loads/stores are explicitly unaligned.
5494    unsafe {
5495        use core::arch::x86_64::*;
5496        let mask = _mm_set1_epi8(15);
5497        let eight = _mm_set1_epi8(8);
5498        let order = _mm256_setr_epi32(0, 1, 4, 5, 2, 3, 6, 7);
5499        let mut acc = 0.0f32;
5500        for gi in 0..gpr {
5501            let packed = _mm_loadu_si128(nib.as_ptr().add((r * gpr + gi) * Q4TP_NIB).cast());
5502            let lo = _mm_and_si128(packed, mask);
5503            let hi = _mm_and_si128(_mm_srli_epi16::<4>(packed), mask);
5504            let w0 = _mm_sub_epi8(_mm_unpacklo_epi8(lo, hi), eight);
5505            let w1 = _mm_sub_epi8(_mm_unpackhi_epi8(lo, hi), eight);
5506            let xp = x.as_ptr().add(gi * GROUP_SIZE);
5507            let a = _mm256_mul_ps(
5508                _mm256_cvtepi32_ps(_mm256_cvtepi8_epi32(w0)),
5509                _mm256_loadu_ps(xp),
5510            );
5511            let b = _mm256_mul_ps(
5512                _mm256_cvtepi32_ps(_mm256_cvtepi8_epi32(_mm_srli_si128::<8>(w0))),
5513                _mm256_loadu_ps(xp.add(8)),
5514            );
5515            let c = _mm256_mul_ps(
5516                _mm256_cvtepi32_ps(_mm256_cvtepi8_epi32(w1)),
5517                _mm256_loadu_ps(xp.add(16)),
5518            );
5519            let d = _mm256_mul_ps(
5520                _mm256_cvtepi32_ps(_mm256_cvtepi8_epi32(_mm_srli_si128::<8>(w1))),
5521                _mm256_loadu_ps(xp.add(24)),
5522            );
5523            let mut pairs = [0.0f32; 16];
5524            _mm256_storeu_ps(
5525                pairs.as_mut_ptr(),
5526                _mm256_permutevar8x32_ps(_mm256_hadd_ps(a, b), order),
5527            );
5528            _mm256_storeu_ps(
5529                pairs.as_mut_ptr().add(8),
5530                _mm256_permutevar8x32_ps(_mm256_hadd_ps(c, d), order),
5531            );
5532            let mut ga = 0.0f32;
5533            for v in pairs {
5534                ga += v;
5535            }
5536            acc += ga * scales[gi];
5537        }
5538        acc
5539    }
5540}
5541
5542/// Single weight of a q4tp tensor — the a8w8 outlier path, which restores
5543/// activation outliers at full precision after the int8 pass.
5544#[inline]
5545fn q4tp_outlier(nib: &[u8], r: usize, gpr: usize, j: usize, scales: &[f32]) -> (f32, f32) {
5546    let (gi, k) = (j / GROUP_SIZE, j % GROUP_SIZE);
5547    let byte = nib[(r * gpr + gi) * Q4TP_NIB + k / 2];
5548    let n = if k & 1 == 0 { byte & 0x0F } else { byte >> 4 };
5549    ((n as i32 - 8) as f32, scales[gi])
5550}
5551
5552/// Fused q4tp matvec (dispatch mirrors `q4t_matvec`).
5553fn q4tp_matvec(
5554    bytes: &[u8],
5555    x: &[f32],
5556    rows: usize,
5557    cols: usize,
5558    out: &mut [f32],
5559    pool: Option<&Pool>,
5560) {
5561    debug_assert_eq!(out.len(), rows);
5562    let gpr = cols / GROUP_SIZE;
5563    let v = Q4tpView::new(bytes, rows, cols);
5564    let out_addr = SendMut(out.as_mut_ptr());
5565    if a8w8_enabled() {
5566        let act = split_act(x);
5567        let run = |start: usize, end: usize| {
5568            // One scratch row of scales per worker — borrowed, not minted.
5569            with_krow(gpr, |sc| {
5570                for r in start..end {
5571                    v.scales_into(r, gpr, sc);
5572                    let mut acc = dot_q4tp_row_i8(v.nib, r, gpr, &act.xq, sc) * act.sx;
5573                    for &(j, xv) in &act.outliers {
5574                        let (w, s) = q4tp_outlier(v.nib, r, gpr, j, sc);
5575                        acc += w * s * xv;
5576                    }
5577                    // SAFETY: disjoint row ranges per worker.
5578                    unsafe { *out_addr.at(r) = acc };
5579                }
5580            })
5581        };
5582        dispatch_rows(pool, rows, &run);
5583        return;
5584    }
5585    let run = |start: usize, end: usize| {
5586        with_krow(gpr, |sc| {
5587            for r in start..end {
5588                v.scales_into(r, gpr, sc);
5589                // SAFETY: disjoint row ranges per worker.
5590                unsafe { *out_addr.at(r) = q4tp_row_exact(v.nib, r, gpr, x, sc) };
5591            }
5592        })
5593    };
5594    dispatch_rows(pool, rows, &run);
5595}
5596
5597/// Fused two-input q4tp matvec — the SwiGLU gate/up pair. Weights and the
5598/// row ladder are read once and spent on both activation streams.
5599#[allow(clippy::too_many_arguments)]
5600fn q4tp_matvec2(
5601    bytes: &[u8],
5602    x1: &[f32],
5603    x2: &[f32],
5604    rows: usize,
5605    cols: usize,
5606    o1: &mut [f32],
5607    o2: &mut [f32],
5608    pool: Option<&Pool>,
5609) {
5610    let gpr = cols / GROUP_SIZE;
5611    let v = Q4tpView::new(bytes, rows, cols);
5612    let (p1, p2) = (SendMut(o1.as_mut_ptr()), SendMut(o2.as_mut_ptr()));
5613    let run = |start: usize, end: usize| {
5614        let mut sc = vec![0f32; gpr];
5615        for r in start..end {
5616            v.scales_into(r, gpr, &mut sc);
5617            // SAFETY: disjoint row ranges per worker.
5618            unsafe {
5619                *p1.at(r) = q4tp_row_exact(v.nib, r, gpr, x1, &sc);
5620                *p2.at(r) = q4tp_row_exact(v.nib, r, gpr, x2, &sc);
5621            }
5622        }
5623    };
5624    dispatch_rows(pool, rows, &run);
5625}
5626
5627/// One q2tp outlier weight at column `j` of row `r`: the 2-bit code and
5628/// its group scale, mirrored on `q4tp_outlier`.
5629#[inline]
5630fn q2tp_outlier(chunks: &[u8], r: usize, gpr: usize, j: usize, scales: &[f32]) -> (f32, f32) {
5631    let (gi, k) = (j / GROUP_SIZE, j % GROUP_SIZE);
5632    let byte = chunks[(r * gpr + gi) * Q2TP_CHUNK + k / 4];
5633    let c = (byte >> (2 * (k % 4))) & 3;
5634    (c as f32 - 1.5, scales[gi])
5635}
5636
5637#[cfg(target_arch = "x86_64")]
5638const Q2TP_DECODE_U32: [u32; 256] = {
5639    let mut tab = [0u32; 256];
5640    let mut b = 0usize;
5641    while b < 256 {
5642        tab[b] = ((b as u32) & 3)
5643            | ((((b as u32) >> 2) & 3) << 8)
5644            | ((((b as u32) >> 4) & 3) << 16)
5645            | ((((b as u32) >> 6) & 3) << 24);
5646        b += 1;
5647    }
5648    tab
5649};
5650
5651/// Eight packed q2tp bytes against 32 signed activation bytes. `maddubs`
5652/// exactly computes unsigned 2-bit code × signed i8; its pair sums cannot
5653/// saturate (2 × 3 × 127 < i16::MAX), and the second madd widens to i32.
5654#[cfg(target_arch = "x86_64")]
5655#[target_feature(enable = "avx2")]
5656unsafe fn q2tp_code_dot_avx2(ch: &[u8], x: &[i8]) -> i32 {
5657    use core::arch::x86_64::*;
5658    debug_assert!(ch.len() >= Q2TP_CHUNK && x.len() >= GROUP_SIZE);
5659    let codes = _mm256_setr_epi32(
5660        Q2TP_DECODE_U32[ch[0] as usize] as i32,
5661        Q2TP_DECODE_U32[ch[1] as usize] as i32,
5662        Q2TP_DECODE_U32[ch[2] as usize] as i32,
5663        Q2TP_DECODE_U32[ch[3] as usize] as i32,
5664        Q2TP_DECODE_U32[ch[4] as usize] as i32,
5665        Q2TP_DECODE_U32[ch[5] as usize] as i32,
5666        Q2TP_DECODE_U32[ch[6] as usize] as i32,
5667        Q2TP_DECODE_U32[ch[7] as usize] as i32,
5668    );
5669    let xv = unsafe { _mm256_loadu_si256(x.as_ptr().cast()) };
5670    let pair = _mm256_maddubs_epi16(codes, xv);
5671    let quad = _mm256_madd_epi16(pair, _mm256_set1_epi16(1));
5672    let sum128 = _mm_add_epi32(
5673        _mm256_castsi256_si128(quad),
5674        _mm256_extracti128_si256(quad, 1),
5675    );
5676    let sum64 = _mm_hadd_epi32(sum128, sum128);
5677    _mm_cvtsi128_si32(_mm_hadd_epi32(sum64, sum64))
5678}
5679
5680#[cfg(target_arch = "x86_64")]
5681#[inline]
5682fn q2tp_avx2_enabled() -> bool {
5683    static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
5684    *ON.get_or_init(|| std::arch::is_x86_feature_detected!("avx2"))
5685}
5686
5687/// Integer dot of one q2tp row against pre-quantized activations:
5688/// Σ_g s_g · (Σ c·xq − 1.5·Σ xq). The half-integer grid (c − 1.5)
5689/// becomes exact integer math through the group sums — the same trick
5690/// every a8w8 kernel in this file rides. The codes decode into a
5691/// 32-byte scratch in natural order and the dot itself is the shared
5692/// SDOT primitive; elsewhere a scalar integer loop.
5693#[inline]
5694fn dot_q2tp_row_i8(
5695    chunks: &[u8],
5696    r: usize,
5697    gpr: usize,
5698    xq: &[i8],
5699    gsum: &[i32],
5700    scales: &[f32],
5701) -> f32 {
5702    let mut acc = 0f32;
5703    let base = r * gpr * Q2TP_CHUNK;
5704    #[cfg(not(any(target_arch = "aarch64", target_arch = "x86_64")))]
5705    let mut codes = [0i8; GROUP_SIZE];
5706    #[cfg(target_arch = "x86_64")]
5707    // Cache the process-wide AVX2 decision; probing it for every 32-weight
5708    // group adds a branch to the hottest Q2TP row loop while retaining the
5709    // established table-load AVX2 arithmetic (and without coupling this path
5710    // to the separate FMA-gated A8W8 switch).
5711    let avx2 = q2tp_avx2_enabled();
5712    for gi in 0..gpr {
5713        let ch = &chunks[base + gi * Q2TP_CHUNK..base + (gi + 1) * Q2TP_CHUNK];
5714        let xg = &xq[gi * GROUP_SIZE..(gi + 1) * GROUP_SIZE];
5715        #[cfg(target_arch = "aarch64")]
5716        // NEON: the byte's four 2-bit fields land in four lane vectors
5717        // (shift+mask), vld4 de-interleaves xq to match (xj[k] =
5718        // xq[4k+j]), widening MACs accumulate exactly in i32. A scalar
5719        // decode here cost as much as the dot it fed — the profile put
5720        // it at the top of the whole W2 decode.
5721        let dot = unsafe {
5722            use core::arch::aarch64::*;
5723            let b = vld1_u8(ch.as_ptr());
5724            let three = vdup_n_u8(3);
5725            let c0 = vreinterpret_s8_u8(vand_u8(b, three));
5726            let c1 = vreinterpret_s8_u8(vand_u8(vshr_n_u8(b, 2), three));
5727            let c2 = vreinterpret_s8_u8(vand_u8(vshr_n_u8(b, 4), three));
5728            let c3 = vreinterpret_s8_u8(vand_u8(vshr_n_u8(b, 6), three));
5729            let x4 = vld4_s8(xg.as_ptr());
5730            let mut acc4 = vdupq_n_s32(0);
5731            acc4 = vpadalq_s16(acc4, vmull_s8(c0, x4.0));
5732            acc4 = vpadalq_s16(acc4, vmull_s8(c1, x4.1));
5733            acc4 = vpadalq_s16(acc4, vmull_s8(c2, x4.2));
5734            acc4 = vpadalq_s16(acc4, vmull_s8(c3, x4.3));
5735            vaddvq_s32(acc4)
5736        };
5737        #[cfg(target_arch = "x86_64")]
5738        let dot: i32 = if avx2 {
5739            // SAFETY: the runtime feature check gates the target-feature body;
5740            // the group slices above are exactly 8 and 32 bytes long.
5741            unsafe { q2tp_code_dot_avx2(ch, xg) }
5742        } else {
5743            ch.iter()
5744                .enumerate()
5745                .map(|(k, &b)| {
5746                    ((b & 3) as i32) * xg[k * 4] as i32
5747                        + (((b >> 2) & 3) as i32) * xg[k * 4 + 1] as i32
5748                        + (((b >> 4) & 3) as i32) * xg[k * 4 + 2] as i32
5749                        + (((b >> 6) & 3) as i32) * xg[k * 4 + 3] as i32
5750                })
5751                .sum()
5752        };
5753        #[cfg(not(any(target_arch = "aarch64", target_arch = "x86_64")))]
5754        let dot: i32 = {
5755            for (k, &b) in ch.iter().enumerate() {
5756                codes[k * 4] = (b & 3) as i8;
5757                codes[k * 4 + 1] = ((b >> 2) & 3) as i8;
5758                codes[k * 4 + 2] = ((b >> 4) & 3) as i8;
5759                codes[k * 4 + 3] = ((b >> 6) & 3) as i8;
5760            }
5761            codes
5762                .iter()
5763                .zip(xg)
5764                .map(|(&c, &x)| c as i32 * x as i32)
5765                .sum()
5766        };
5767        acc += scales[gi] * (dot as f32 - 1.5 * gsum[gi] as f32);
5768    }
5769    acc
5770}
5771
5772/// Exact f32 dot of one q2tp row: 2-bit fields LSB-first, (c − 1.5)·s.
5773/// Scalar on purpose — the 2-bit class targets the GPU graph; the CPU
5774/// path exists for parity gates and small-machine fallback.
5775fn q2tp_row_exact(chunks: &[u8], r: usize, gpr: usize, x: &[f32], scales: &[f32]) -> f32 {
5776    q2tp_row_exact_center(chunks, r, gpr, x, scales, 1.5)
5777}
5778
5779/// Fused Prism affine row: the derived correction is applied inside the
5780/// decoded code, avoiding a second accumulated dot and avoiding cancellation
5781/// between `B=(c-1.5)s` and `+.5s` for long 5120/17408 rows.
5782#[inline]
5783fn q2tp_affine_row_exact(chunks: &[u8], r: usize, gpr: usize, x: &[f32], scales: &[f32]) -> f32 {
5784    q2tp_row_exact_center(chunks, r, gpr, x, scales, 1.0)
5785}
5786
5787#[inline]
5788fn q2tp_row_exact_center(
5789    chunks: &[u8],
5790    r: usize,
5791    gpr: usize,
5792    x: &[f32],
5793    scales: &[f32],
5794    center: f32,
5795) -> f32 {
5796    let mut acc = 0f32;
5797    for gi in 0..gpr {
5798        let ch = &chunks[(r * gpr + gi) * Q2TP_CHUNK..(r * gpr + gi + 1) * Q2TP_CHUNK];
5799        let s = scales[gi];
5800        let xb = &x[gi * GROUP_SIZE..(gi + 1) * GROUP_SIZE];
5801        let mut g = 0f32;
5802        for (k, &b) in ch.iter().enumerate() {
5803            g += ((b & 3) as f32 - center) * xb[k * 4]
5804                + (((b >> 2) & 3) as f32 - center) * xb[k * 4 + 1]
5805                + (((b >> 4) & 3) as f32 - center) * xb[k * 4 + 2]
5806                + (((b >> 6) & 3) as f32 - center) * xb[k * 4 + 3];
5807        }
5808        acc += s * g;
5809    }
5810    acc
5811}
5812
5813fn q2tp_matvec(
5814    bytes: &[u8],
5815    x: &[f32],
5816    rows: usize,
5817    cols: usize,
5818    out: &mut [f32],
5819    pool: Option<&Pool>,
5820) {
5821    q2tp_matvec_mode(bytes, x, rows, cols, out, pool, false);
5822}
5823
5824fn q2tp_affine_matvec(
5825    bytes: &[u8],
5826    x: &[f32],
5827    rows: usize,
5828    cols: usize,
5829    out: &mut [f32],
5830    pool: Option<&Pool>,
5831) {
5832    q2tp_matvec_mode(bytes, x, rows, cols, out, pool, true);
5833}
5834
5835fn q2tp_matvec_mode(
5836    bytes: &[u8],
5837    x: &[f32],
5838    rows: usize,
5839    cols: usize,
5840    out: &mut [f32],
5841    pool: Option<&Pool>,
5842    affine: bool,
5843) {
5844    debug_assert_eq!(out.len(), rows);
5845    let gpr = cols / GROUP_SIZE;
5846    let v = Q4tpView::new_q2(bytes, rows, cols);
5847    let out_addr = SendMut(out.as_mut_ptr());
5848    // a8w8 fast path (CMF_SDOT=0 keeps the exact scalar walk): integer
5849    // code dots + group sums, exact outlier correction — the same
5850    // contract as every sibling kernel; measured 2-bit rows were the
5851    // only scalar holdout in the family.
5852    if !affine && a8w8_enabled() {
5853        let act = split_act(x);
5854        let gsum = q1_group_sums(&act.xq, gpr);
5855        let (act, gsum) = (&act, &gsum);
5856        let run = move |start: usize, end: usize| {
5857            with_krow(gpr, |sc| {
5858                for r in start..end {
5859                    v.scales_into(r, gpr, sc);
5860                    let mut acc = dot_q2tp_row_i8(v.nib, r, gpr, &act.xq, gsum, sc) * act.sx;
5861                    for &(j, xv) in &act.outliers {
5862                        let (w, s) = q2tp_outlier(v.nib, r, gpr, j, sc);
5863                        acc += w * s * xv;
5864                    }
5865                    // SAFETY: disjoint row ranges per worker.
5866                    unsafe { *out_addr.at(r) = acc };
5867                }
5868            })
5869        };
5870        dispatch_rows(pool, rows, &run);
5871        return;
5872    }
5873    let run = |start: usize, end: usize| {
5874        with_krow(gpr, |sc| {
5875            for r in start..end {
5876                v.scales_into(r, gpr, sc);
5877                // SAFETY: disjoint row ranges per worker.
5878                unsafe {
5879                    *out_addr.at(r) = if affine {
5880                        q2tp_affine_row_exact(v.nib, r, gpr, x, sc)
5881                    } else {
5882                        q2tp_row_exact(v.nib, r, gpr, x, sc)
5883                    }
5884                };
5885            }
5886        })
5887    };
5888    dispatch_rows(pool, rows, &run);
5889}
5890
5891/// Fused two-input q2tp matvec — the SwiGLU gate/up pair.
5892#[allow(clippy::too_many_arguments)]
5893fn q2tp_matvec2(
5894    bytes: &[u8],
5895    x1: &[f32],
5896    x2: &[f32],
5897    rows: usize,
5898    cols: usize,
5899    o1: &mut [f32],
5900    o2: &mut [f32],
5901    pool: Option<&Pool>,
5902) {
5903    q2tp_matvec2_mode(bytes, x1, x2, rows, cols, o1, o2, pool, false);
5904}
5905
5906#[allow(clippy::too_many_arguments)]
5907fn q2tp_affine_matvec2(
5908    bytes: &[u8],
5909    x1: &[f32],
5910    x2: &[f32],
5911    rows: usize,
5912    cols: usize,
5913    o1: &mut [f32],
5914    o2: &mut [f32],
5915    pool: Option<&Pool>,
5916) {
5917    q2tp_matvec2_mode(bytes, x1, x2, rows, cols, o1, o2, pool, true);
5918}
5919
5920#[allow(clippy::too_many_arguments)]
5921fn q2tp_matvec2_mode(
5922    bytes: &[u8],
5923    x1: &[f32],
5924    x2: &[f32],
5925    rows: usize,
5926    cols: usize,
5927    o1: &mut [f32],
5928    o2: &mut [f32],
5929    pool: Option<&Pool>,
5930    affine: bool,
5931) {
5932    let gpr = cols / GROUP_SIZE;
5933    let v = Q4tpView::new_q2(bytes, rows, cols);
5934    let (p1, p2) = (SendMut(o1.as_mut_ptr()), SendMut(o2.as_mut_ptr()));
5935    let run = |start: usize, end: usize| {
5936        let mut sc = vec![0f32; gpr];
5937        for r in start..end {
5938            v.scales_into(r, gpr, &mut sc);
5939            // SAFETY: disjoint row ranges per worker.
5940            unsafe {
5941                *p1.at(r) = if affine {
5942                    q2tp_affine_row_exact(v.nib, r, gpr, x1, &sc)
5943                } else {
5944                    q2tp_row_exact(v.nib, r, gpr, x1, &sc)
5945                };
5946                *p2.at(r) = if affine {
5947                    q2tp_affine_row_exact(v.nib, r, gpr, x2, &sc)
5948                } else {
5949                    q2tp_row_exact(v.nib, r, gpr, x2, &sc)
5950                };
5951            }
5952        }
5953    };
5954    dispatch_rows(pool, rows, &run);
5955}
5956
5957/// Batched q2tp matmat: scalar row kernel over every batch column. CPU
5958/// prefill only — decode rides the graph, so plain and correct beats
5959/// clever here.
5960/// Test doors into the host 2-bit kernels: the stand's heap corruption
5961/// pointed at down-shaped tensors, and the private fns need a way to be
5962/// held to a reference without a model file around them.
5963pub fn q2tp_matvec_for_test(bytes: &[u8], x: &[f32], rows: usize, cols: usize, out: &mut [f32]) {
5964    // The facade IS the reference: encoder oracles hold requant output
5965    // to the exact scalar walk. The production dispatch may take the i8
5966    // fast path, whose error scale is the ACTIVATIONS' — a different
5967    // claim than the encoder correctness these tests pin.
5968    let gpr = cols / GROUP_SIZE;
5969    let v = Q4tpView::new_q2(bytes, rows, cols);
5970    with_krow(gpr, |sc| {
5971        for r in 0..rows {
5972            v.scales_into(r, gpr, sc);
5973            out[r] = q2tp_row_exact(v.nib, r, gpr, x, sc);
5974        }
5975    });
5976}
5977
5978/// Test door for the descriptor-specific fused affine decode.  Production
5979/// callers select this through a validated Prism header, never by dtype alone.
5980pub fn q2tp_affine_matvec_for_test(
5981    bytes: &[u8],
5982    x: &[f32],
5983    rows: usize,
5984    cols: usize,
5985    out: &mut [f32],
5986) {
5987    q2tp_affine_matvec(bytes, x, rows, cols, out, None);
5988}
5989
5990pub fn q2tp_matmat_for_test(
5991    bytes: &[u8],
5992    xs_all: &[f32],
5993    b: usize,
5994    rows: usize,
5995    cols: usize,
5996    out: &mut [f32],
5997) {
5998    q2tp_matmat(bytes, xs_all, b, rows, cols, out, None);
5999}
6000
6001fn q2tp_matmat(
6002    bytes: &[u8],
6003    xs_all: &[f32],
6004    b: usize,
6005    rows: usize,
6006    cols: usize,
6007    out: &mut [f32],
6008    pool: Option<&Pool>,
6009) {
6010    q2tp_matmat_mode(bytes, xs_all, b, rows, cols, out, pool, false);
6011}
6012
6013fn q2tp_affine_matmat(
6014    bytes: &[u8],
6015    xs_all: &[f32],
6016    b: usize,
6017    rows: usize,
6018    cols: usize,
6019    out: &mut [f32],
6020    pool: Option<&Pool>,
6021) {
6022    q2tp_matmat_mode(bytes, xs_all, b, rows, cols, out, pool, true);
6023}
6024
6025fn q2tp_matmat_mode(
6026    bytes: &[u8],
6027    xs_all: &[f32],
6028    b: usize,
6029    rows: usize,
6030    cols: usize,
6031    out: &mut [f32],
6032    pool: Option<&Pool>,
6033    affine: bool,
6034) {
6035    debug_assert_eq!(out.len(), b * rows);
6036    let gpr = cols / GROUP_SIZE;
6037    let v = Q4tpView::new_q2(bytes, rows, cols);
6038    let out_addr = SendMut(out.as_mut_ptr());
6039    let run = |start: usize, end: usize| {
6040        let mut sc = vec![0f32; gpr];
6041        for r in start..end {
6042            v.scales_into(r, gpr, &mut sc);
6043            for bi in 0..b {
6044                let x = &xs_all[bi * cols..(bi + 1) * cols];
6045                // SAFETY: disjoint row ranges per worker.
6046                unsafe {
6047                    *out_addr.at(bi * rows + r) = if affine {
6048                        q2tp_affine_row_exact(v.nib, r, gpr, x, &sc)
6049                    } else {
6050                        q2tp_row_exact(v.nib, r, gpr, x, &sc)
6051                    }
6052                };
6053            }
6054        }
6055    };
6056    dispatch_rows(pool, rows, &run);
6057}
6058
6059/// The pre-vectorised shape, kept for A/B (`CMF_Q4TP_V1=1`): the
6060/// horizontal add lands once per group per column instead of once per
6061/// row. Same weights, same activations — only the reduction differs.
6062/// It is also the row-exact batch kernel: per group and column it forms
6063/// `int dot as f32 * scale` and adds it to a scalar running sum, which is
6064/// `dot_q4tp_row_sdot` step for step (Rust never contracts to an fma),
6065/// so each column is bit-identical to that column's matvec.
6066#[cfg(target_arch = "aarch64")]
6067#[target_feature(enable = "neon,dotprod")]
6068unsafe fn dot_q4tp_row_1x4_sdot_v1(
6069    nib: &[u8],
6070    r: usize,
6071    gpr: usize,
6072    xs: [&[i8]; 4],
6073    scales: &[f32],
6074) -> [f32; 4] {
6075    unsafe {
6076        use core::arch::aarch64::*;
6077        use core::arch::asm;
6078        let lomask = vdupq_n_u8(0x0F);
6079        let eight = vdupq_n_s8(8);
6080        let (mut f0, mut f1, mut f2, mut f3) = (0f32, 0f32, 0f32, 0f32);
6081        for gi in 0..gpr {
6082            let t = nib.as_ptr().add((r * gpr + gi) * Q4TP_NIB);
6083            let s = *scales.get_unchecked(gi);
6084            let bb = vld1q_u8(t);
6085            let lo = vandq_u8(bb, lomask);
6086            let hi = vshrq_n_u8::<4>(bb);
6087            let e0 = vsubq_s8(vreinterpretq_s8_u8(vzip1q_u8(lo, hi)), eight);
6088            let e1 = vsubq_s8(vreinterpretq_s8_u8(vzip2q_u8(lo, hi)), eight);
6089            let mut d = [0f32; 4];
6090            for (k, dk) in d.iter_mut().enumerate() {
6091                let x0 = vld1q_s8(xs[k].as_ptr().add(gi * GROUP_SIZE));
6092                let x1 = vld1q_s8(xs[k].as_ptr().add(gi * GROUP_SIZE + 16));
6093                let (mut a0, mut a1) = (vdupq_n_s32(0), vdupq_n_s32(0));
6094                asm!(
6095                    "sdot {a0:v}.4s, {e0:v}.16b, {x0:v}.16b",
6096                    "sdot {a1:v}.4s, {e1:v}.16b, {x1:v}.16b",
6097                    a0 = inout(vreg) a0, a1 = inout(vreg) a1,
6098                    e0 = in(vreg) e0, x0 = in(vreg) x0, e1 = in(vreg) e1, x1 = in(vreg) x1,
6099                    options(pure, nomem, nostack),
6100                );
6101                *dk = vaddvq_s32(vaddq_s32(a0, a1)) as f32 * s;
6102            }
6103            f0 += d[0];
6104            f1 += d[1];
6105            f2 += d[2];
6106            f3 += d[3];
6107        }
6108        [f0, f1, f2, f3]
6109    }
6110}
6111
6112/// Which q4tp batch kernel to run: 1 = the previous one, 2 = the tuned
6113/// one, 0 = decide from the CPU. An atomic rather than a `OnceLock` so a
6114/// benchmark can alternate the two inside one process, where the machine's
6115/// mood — a shared box drifts ±25% between runs — is the same for both.
6116/// What the two mean is per-architecture: on x86 the blocked AVX-512 path
6117/// against the per-column one, on ARM the two reduction shapes.
6118#[allow(dead_code)]
6119static Q4TP_ALT: std::sync::atomic::AtomicU8 = std::sync::atomic::AtomicU8::new(0);
6120
6121/// Tests that store `Q4TP_ALT` hold this, so one test's kernel pick does
6122/// not leak into another's bit-exact comparison running in parallel.
6123#[cfg(test)]
6124static Q4TP_ALT_TEST_LOCK: std::sync::Mutex<()> = std::sync::Mutex::new(());
6125
6126/// Blocking pays on x86 only with 512-bit VNNI. With AVX2 alone, four
6127/// columns sharing an unpack still measured slower than the per-column
6128/// path (23.2 ms against 19.4 on a 48-thread EPYC), because that path
6129/// already dequantizes the row once — so the blocked kernel bought a
6130/// second unpack-free pass at the price of half the vector width.
6131#[cfg(target_arch = "x86_64")]
6132fn q4tp_blocked_x86() -> bool {
6133    match Q4TP_ALT.load(std::sync::atomic::Ordering::Relaxed) {
6134        1 => false,
6135        // A forced ON still asks the CPU. The switch exists so a bench can
6136        // pick a kernel, not so it can promise instructions the machine
6137        // does not have — CI caught that as a SIGILL on a runner without
6138        // AVX-512, where the parity test had turned the path on by hand.
6139        2 => avx512vnni_enabled(),
6140        // Deliberately not cached back into the switch: both gates below
6141        // hold their own `OnceLock`, and latching their answer here would
6142        // make a test's override outlive the test that set it.
6143        _ => blocked_enabled() && avx512vnni_enabled(),
6144    }
6145}
6146
6147/// `CMF_Q4TP_V1=1` picks the old reduction shape (A/B only).
6148#[cfg(target_arch = "aarch64")]
6149#[allow(dead_code)]
6150fn q4tp_v1() -> bool {
6151    match Q4TP_ALT.load(std::sync::atomic::Ordering::Relaxed) {
6152        1 => true,
6153        2 => false,
6154        _ => {
6155            static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
6156            *ON.get_or_init(|| std::env::var("CMF_Q4TP_V1").is_ok_and(|v| v != "0"))
6157        }
6158    }
6159}
6160
6161/// Two weight rows against eight columns. The activation load is the
6162/// same for both rows, so it is paid once for twice the arithmetic, and
6163/// sixteen accumulator chains run where eight did — which is what a kernel
6164/// retiring 0.29 instructions a cycle is short of. Register pressure is
6165/// the limit: sixteen `zmm` accumulators, two weight tiles, one
6166/// activation, of thirty-two.
6167///
6168/// Four rows by four columns spends the same sixteen accumulators the
6169/// other way and measured worse — 1488 GFLOP/s against 1644 — so the
6170/// unpack, which four rows pay twice as often, costs more than the extra
6171/// sharing of one activation load buys.
6172#[cfg(target_arch = "x86_64")]
6173#[target_feature(enable = "avx512f,avx512bw,avx512vnni")]
6174unsafe fn dot_q4tp_2x8_avx512(
6175    nib: &[u8],
6176    r0: usize,
6177    gpr: usize,
6178    xs: [&[i8]; 8],
6179    sc0: &[f32],
6180    sc1: &[f32],
6181) -> [[f32; 8]; 2] {
6182    // SAFETY: as dot_q4tp_row_1x8_avx512, two adjacent rows at once; the
6183    // caller guarantees r0 + 1 < rows and the ISA.
6184    unsafe {
6185        use core::arch::x86_64::*;
6186        let lomask = _mm256_set1_epi8(0x0F);
6187        let eight = _mm256_set1_epi8(8);
6188        let zero = _mm512_setzero_si512();
6189        let mut v0 = [_mm512_setzero_ps(); 8];
6190        let mut v1 = [_mm512_setzero_ps(); 8];
6191        let pairs = gpr / 2;
6192        let unpack = |r: usize, gi: usize| -> (__m512i, __mmask64) {
6193            let t = nib.as_ptr().add((r * gpr + gi) * Q4TP_NIB);
6194            let bb = _mm256_loadu_si256(t as *const __m256i);
6195            let lo = _mm256_and_si256(bb, lomask);
6196            let hi = _mm256_and_si256(_mm256_srli_epi16::<4>(bb), lomask);
6197            let ul = _mm256_sub_epi8(_mm256_unpacklo_epi8(lo, hi), eight);
6198            let uh = _mm256_sub_epi8(_mm256_unpackhi_epi8(lo, hi), eight);
6199            let cat = _mm512_inserti64x4::<1>(_mm512_castsi256_si512(ul), uh);
6200            let w = _mm512_shuffle_i64x2::<0b11_01_10_00>(cat, cat);
6201            (_mm512_abs_epi8(w), _mm512_movepi8_mask(w))
6202        };
6203        for gp in 0..pairs {
6204            let gi = gp * 2;
6205            let (wa0, neg0) = unpack(r0, gi);
6206            let (wa1, neg1) = unpack(r0 + 1, gi);
6207            let off = gi * GROUP_SIZE;
6208            let sv = |sc: &[f32]| {
6209                _mm512_insertf32x8::<1>(
6210                    _mm512_castps256_ps512(_mm256_set1_ps(*sc.get_unchecked(gi))),
6211                    _mm256_set1_ps(*sc.get_unchecked(gi + 1)),
6212                )
6213            };
6214            let s0 = sv(sc0);
6215            let s1 = sv(sc1);
6216            for k in 0..8 {
6217                let xv = _mm512_loadu_si512(xs[k].as_ptr().add(off) as *const __m512i);
6218                let d0 = _mm512_cvtepi32_ps(_mm512_dpbusd_epi32(
6219                    zero,
6220                    wa0,
6221                    _mm512_mask_sub_epi8(xv, neg0, zero, xv),
6222                ));
6223                let d1 = _mm512_cvtepi32_ps(_mm512_dpbusd_epi32(
6224                    zero,
6225                    wa1,
6226                    _mm512_mask_sub_epi8(xv, neg1, zero, xv),
6227                ));
6228                v0[k] = _mm512_fmadd_ps(d0, s0, v0[k]);
6229                v1[k] = _mm512_fmadd_ps(d1, s1, v1[k]);
6230            }
6231        }
6232        let mut acc = [[0f32; 8]; 2];
6233        for k in 0..8 {
6234            acc[0][k] = _mm512_reduce_add_ps(v0[k]);
6235            acc[1][k] = _mm512_reduce_add_ps(v1[k]);
6236        }
6237        if gpr % 2 == 1 {
6238            let off = (gpr - 1) * GROUP_SIZE;
6239            for j in off..off + GROUP_SIZE {
6240                let (w0, sa) = q4tp_outlier(nib, r0, gpr, j, sc0);
6241                let (w1, sb) = q4tp_outlier(nib, r0 + 1, gpr, j, sc1);
6242                for k in 0..8 {
6243                    let x = *xs[k].get_unchecked(j) as f32;
6244                    acc[0][k] += w0 * sa * x;
6245                    acc[1][k] += w1 * sb * x;
6246                }
6247            }
6248        }
6249        acc
6250    }
6251}
6252
6253/// The same, eight columns at a time. One unpack then feeds twice as many
6254/// activation streams, so a wide batch reads the weight tile half as
6255/// often; the price is eight accumulators live at once. Measured 9.0 ->
6256/// 8.3 ms at 9216x2304, b=296 on a 48-thread EPYC 9B45.
6257#[cfg(target_arch = "x86_64")]
6258#[target_feature(enable = "avx512f,avx512bw,avx512vnni")]
6259unsafe fn dot_q4tp_row_1x8_avx512(
6260    nib: &[u8],
6261    r: usize,
6262    gpr: usize,
6263    xs: [&[i8]; 8],
6264    scales: &[f32],
6265) -> [f32; 8] {
6266    // SAFETY: as dot_q4tp_row_1x4_avx2; caller guarantees the ISA.
6267    unsafe {
6268        use core::arch::x86_64::*;
6269        let lomask = _mm256_set1_epi8(0x0F);
6270        let eight = _mm256_set1_epi8(8);
6271        let zero = _mm512_setzero_si512();
6272        let (mut v0, mut v1, mut v2, mut v3) = (
6273            _mm512_setzero_ps(),
6274            _mm512_setzero_ps(),
6275            _mm512_setzero_ps(),
6276            _mm512_setzero_ps(),
6277        );
6278        let (mut v4, mut v5, mut v6, mut v7) = (
6279            _mm512_setzero_ps(),
6280            _mm512_setzero_ps(),
6281            _mm512_setzero_ps(),
6282            _mm512_setzero_ps(),
6283        );
6284        let pairs = gpr / 2;
6285        for gp in 0..pairs {
6286            let gi = gp * 2;
6287            let t = nib.as_ptr().add((r * gpr + gi) * Q4TP_NIB);
6288            let bb = _mm256_loadu_si256(t as *const __m256i);
6289            let lo = _mm256_and_si256(bb, lomask);
6290            let hi = _mm256_and_si256(_mm256_srli_epi16::<4>(bb), lomask);
6291            // `unpack` works per 128-bit lane, so the halves come out as
6292            // [A.lo, B.lo] and [A.hi, B.hi]; the shuffle reorders the four
6293            // 128-bit lanes into the weights' natural order, which is what
6294            // the straight activation load expects.
6295            let ul = _mm256_sub_epi8(_mm256_unpacklo_epi8(lo, hi), eight);
6296            let uh = _mm256_sub_epi8(_mm256_unpackhi_epi8(lo, hi), eight);
6297            let cat = _mm512_inserti64x4::<1>(_mm512_castsi256_si512(ul), uh);
6298            let w = _mm512_shuffle_i64x2::<0b11_01_10_00>(cat, cat);
6299            let wabs = _mm512_abs_epi8(w);
6300            let neg = _mm512_movepi8_mask(w);
6301            let off = gi * GROUP_SIZE;
6302            let sv = _mm512_insertf32x8::<1>(
6303                _mm512_castps256_ps512(_mm256_set1_ps(*scales.get_unchecked(gi))),
6304                _mm256_set1_ps(*scales.get_unchecked(gi + 1)),
6305            );
6306            let dot = |x: &[i8]| -> __m512 {
6307                let xv = _mm512_loadu_si512(x.as_ptr().add(off) as *const __m512i);
6308                let sx = _mm512_mask_sub_epi8(xv, neg, zero, xv);
6309                _mm512_cvtepi32_ps(_mm512_dpbusd_epi32(zero, wabs, sx))
6310            };
6311            v0 = _mm512_fmadd_ps(dot(xs[0]), sv, v0);
6312            v1 = _mm512_fmadd_ps(dot(xs[1]), sv, v1);
6313            v2 = _mm512_fmadd_ps(dot(xs[2]), sv, v2);
6314            v3 = _mm512_fmadd_ps(dot(xs[3]), sv, v3);
6315            v4 = _mm512_fmadd_ps(dot(xs[4]), sv, v4);
6316            v5 = _mm512_fmadd_ps(dot(xs[5]), sv, v5);
6317            v6 = _mm512_fmadd_ps(dot(xs[6]), sv, v6);
6318            v7 = _mm512_fmadd_ps(dot(xs[7]), sv, v7);
6319        }
6320        let mut acc = [
6321            _mm512_reduce_add_ps(v0),
6322            _mm512_reduce_add_ps(v1),
6323            _mm512_reduce_add_ps(v2),
6324            _mm512_reduce_add_ps(v3),
6325            _mm512_reduce_add_ps(v4),
6326            _mm512_reduce_add_ps(v5),
6327            _mm512_reduce_add_ps(v6),
6328            _mm512_reduce_add_ps(v7),
6329        ];
6330        // An odd group count leaves one group over; the narrow kernel
6331        // finishes it rather than the tail being a special case here.
6332        if gpr % 2 == 1 {
6333            let off = (gpr - 1) * GROUP_SIZE;
6334            for j in off..off + GROUP_SIZE {
6335                let (w, s) = q4tp_outlier(nib, r, gpr, j, scales);
6336                let ws = w * s;
6337                for k in 0..8 {
6338                    acc[k] += ws * *xs[k].get_unchecked(j) as f32;
6339                }
6340            }
6341        }
6342        acc
6343    }
6344}
6345
6346/// The same four columns, 512 bits wide. Two groups (64 weights) ride one
6347/// unpack and one `vpdpbusd`, where AVX2 needs two unpacks and four
6348/// `maddubs`/`madd` pairs — about 2.3x fewer instructions for the same
6349/// arithmetic. The two groups carry different scales, so the fma takes a
6350/// vector whose halves hold each group's scale rather than a broadcast.
6351///
6352/// There is no 512-bit `vpsignb`, so the activation's sign is applied by
6353/// negating under a mask taken from the weight's sign bits. That mask is
6354/// per-tile, so it is hoisted out of the column loop and the per-column
6355/// cost stays exactly one instruction, as with `sign_epi8`. Weights of
6356/// zero are not zeroed by the mask trick and do not need to be: their
6357/// magnitude is zero, so the product is.
6358#[cfg(target_arch = "x86_64")]
6359#[target_feature(enable = "avx512f,avx512bw,avx512vnni")]
6360unsafe fn dot_q4tp_row_1x4_avx512(
6361    nib: &[u8],
6362    r: usize,
6363    gpr: usize,
6364    xs: [&[i8]; 4],
6365    scales: &[f32],
6366) -> [f32; 4] {
6367    // SAFETY: as dot_q4tp_row_1x4_avx2; caller guarantees the ISA.
6368    unsafe {
6369        use core::arch::x86_64::*;
6370        let lomask = _mm256_set1_epi8(0x0F);
6371        let eight = _mm256_set1_epi8(8);
6372        let zero = _mm512_setzero_si512();
6373        let (mut v0, mut v1, mut v2, mut v3) = (
6374            _mm512_setzero_ps(),
6375            _mm512_setzero_ps(),
6376            _mm512_setzero_ps(),
6377            _mm512_setzero_ps(),
6378        );
6379        let pairs = gpr / 2;
6380        for gp in 0..pairs {
6381            let gi = gp * 2;
6382            let t = nib.as_ptr().add((r * gpr + gi) * Q4TP_NIB);
6383            let bb = _mm256_loadu_si256(t as *const __m256i);
6384            let lo = _mm256_and_si256(bb, lomask);
6385            let hi = _mm256_and_si256(_mm256_srli_epi16::<4>(bb), lomask);
6386            // `unpack` works per 128-bit lane, so the halves come out as
6387            // [A.lo, B.lo] and [A.hi, B.hi]; the shuffle reorders the four
6388            // 128-bit lanes into the weights' natural order, which is what
6389            // the straight activation load expects.
6390            let ul = _mm256_sub_epi8(_mm256_unpacklo_epi8(lo, hi), eight);
6391            let uh = _mm256_sub_epi8(_mm256_unpackhi_epi8(lo, hi), eight);
6392            let cat = _mm512_inserti64x4::<1>(_mm512_castsi256_si512(ul), uh);
6393            let w = _mm512_shuffle_i64x2::<0b11_01_10_00>(cat, cat);
6394            let wabs = _mm512_abs_epi8(w);
6395            let neg = _mm512_movepi8_mask(w);
6396            let off = gi * GROUP_SIZE;
6397            let sv = _mm512_insertf32x8::<1>(
6398                _mm512_castps256_ps512(_mm256_set1_ps(*scales.get_unchecked(gi))),
6399                _mm256_set1_ps(*scales.get_unchecked(gi + 1)),
6400            );
6401            let dot = |x: &[i8]| -> __m512 {
6402                let xv = _mm512_loadu_si512(x.as_ptr().add(off) as *const __m512i);
6403                let sx = _mm512_mask_sub_epi8(xv, neg, zero, xv);
6404                _mm512_cvtepi32_ps(_mm512_dpbusd_epi32(zero, wabs, sx))
6405            };
6406            v0 = _mm512_fmadd_ps(dot(xs[0]), sv, v0);
6407            v1 = _mm512_fmadd_ps(dot(xs[1]), sv, v1);
6408            v2 = _mm512_fmadd_ps(dot(xs[2]), sv, v2);
6409            v3 = _mm512_fmadd_ps(dot(xs[3]), sv, v3);
6410        }
6411        let mut acc = [
6412            _mm512_reduce_add_ps(v0),
6413            _mm512_reduce_add_ps(v1),
6414            _mm512_reduce_add_ps(v2),
6415            _mm512_reduce_add_ps(v3),
6416        ];
6417        // An odd group count leaves one group over; the narrow kernel
6418        // finishes it rather than the tail being a special case here.
6419        if gpr % 2 == 1 {
6420            let off = (gpr - 1) * GROUP_SIZE;
6421            for j in off..off + GROUP_SIZE {
6422                let (w, s) = q4tp_outlier(nib, r, gpr, j, scales);
6423                let ws = w * s;
6424                for k in 0..4 {
6425                    acc[k] += ws * *xs[k].get_unchecked(j) as f32;
6426                }
6427            }
6428        }
6429        acc
6430    }
6431}
6432
6433/// Four batch columns against one q4tp row: the tile is unpacked ONCE and
6434/// spent on four activation streams, which is where a prefill batch stops
6435/// being weight-bandwidth-bound. Twin of `dot_q4t_row_1x4_sdot`.
6436#[cfg(target_arch = "aarch64")]
6437#[target_feature(enable = "neon,dotprod")]
6438unsafe fn dot_q4tp_row_1x4_sdot(
6439    nib: &[u8],
6440    r: usize,
6441    gpr: usize,
6442    xs: [&[i8]; 4],
6443    scales: &[f32],
6444) -> [f32; 4] {
6445    // SAFETY: see dot_q4tp_row_sdot; every xs[k] is gpr·GROUP_SIZE long.
6446    unsafe {
6447        use core::arch::aarch64::*;
6448        use core::arch::asm;
6449        let lomask = vdupq_n_u8(0x0F);
6450        let eight = vdupq_n_s8(8);
6451        // Named accumulators, NOT an array indexed by a loop variable: the
6452        // latter does not stay in registers (the same defect cost 2x in the
6453        // AVX2 q4t kernel and again in WGSL).
6454        //
6455        // They are VECTORS, and the horizontal add happens once at the end
6456        // instead of once per group per column. `vaddvq` is a cross-lane
6457        // reduction — with 72 groups and four columns the old shape paid
6458        // 288 of them per row, each one a dependency stall the pipeline
6459        // cannot hide, to save four float adds. The group's scale now
6460        // rides an fma into the lane accumulators, so the arithmetic per
6461        // group is one convert and one fma. Summation order changes (the
6462        // lanes carry independent partial sums), which is the same
6463        // round-off class the SDOT path already lives in — the strict
6464        // kernel (`CMF_SDOT=0`, what `cortiq ppl` runs) is unchanged and
6465        // stays the reference.
6466        let (mut v0, mut v1, mut v2, mut v3) = (
6467            vdupq_n_f32(0.0),
6468            vdupq_n_f32(0.0),
6469            vdupq_n_f32(0.0),
6470            vdupq_n_f32(0.0),
6471        );
6472        for gi in 0..gpr {
6473            let t = nib.as_ptr().add((r * gpr + gi) * Q4TP_NIB);
6474            let s = *scales.get_unchecked(gi);
6475            let bb = vld1q_u8(t);
6476            let lo = vandq_u8(bb, lomask);
6477            let hi = vshrq_n_u8::<4>(bb);
6478            let e0 = vsubq_s8(vreinterpretq_s8_u8(vzip1q_u8(lo, hi)), eight);
6479            let e1 = vsubq_s8(vreinterpretq_s8_u8(vzip2q_u8(lo, hi)), eight);
6480            let off = gi * GROUP_SIZE;
6481            let dot4 = |x: &[i8]| -> int32x4_t {
6482                let x0 = vld1q_s8(x.as_ptr().add(off));
6483                let x1 = vld1q_s8(x.as_ptr().add(off + 16));
6484                let (mut a0, mut a1) = (vdupq_n_s32(0), vdupq_n_s32(0));
6485                asm!(
6486                    "sdot {a0:v}.4s, {e0:v}.16b, {x0:v}.16b",
6487                    "sdot {a1:v}.4s, {e1:v}.16b, {x1:v}.16b",
6488                    a0 = inout(vreg) a0, a1 = inout(vreg) a1,
6489                    e0 = in(vreg) e0, x0 = in(vreg) x0, e1 = in(vreg) e1, x1 = in(vreg) x1,
6490                    options(pure, nomem, nostack),
6491                );
6492                vaddq_s32(a0, a1)
6493            };
6494            v0 = vfmaq_n_f32(v0, vcvtq_f32_s32(dot4(xs[0])), s);
6495            v1 = vfmaq_n_f32(v1, vcvtq_f32_s32(dot4(xs[1])), s);
6496            v2 = vfmaq_n_f32(v2, vcvtq_f32_s32(dot4(xs[2])), s);
6497            v3 = vfmaq_n_f32(v3, vcvtq_f32_s32(dot4(xs[3])), s);
6498        }
6499        [
6500            vaddvq_f32(v0),
6501            vaddvq_f32(v1),
6502            vaddvq_f32(v2),
6503            vaddvq_f32(v3),
6504        ]
6505    }
6506}
6507
6508/// Fused q4tp matmat — the same three arms `q4t_matmat` has. Shipping only
6509/// the scalar one made Nanbeige-3B decode at 1.2 tok/s against q4t's 5.9:
6510/// the format was fine, the missing arms were the whole regression.
6511///
6512/// Under `row_exact()` every cell is summed in `q4tp_matvec`'s order on
6513/// every architecture; the mode is read once, so one call never mixes
6514/// the two contracts when a concurrent scope opens or closes mid-call.
6515fn q4tp_matmat(
6516    bytes: &[u8],
6517    xs_all: &[f32],
6518    b: usize,
6519    rows: usize,
6520    cols: usize,
6521    out: &mut [f32],
6522    pool: Option<&Pool>,
6523) {
6524    q4tp_matmat_with(bytes, xs_all, b, rows, cols, out, pool, row_exact())
6525}
6526
6527/// `q4tp_matmat` with the row-exact mode passed in rather than read from
6528/// the shared counter: `exact` = each cell equals its token's matvec bit
6529/// for bit, otherwise the fast blocked / AMX arms are free to reorder.
6530#[allow(clippy::too_many_arguments)]
6531fn q4tp_matmat_with(
6532    bytes: &[u8],
6533    xs_all: &[f32],
6534    b: usize,
6535    rows: usize,
6536    cols: usize,
6537    out: &mut [f32],
6538    pool: Option<&Pool>,
6539    exact: bool,
6540) {
6541    debug_assert_eq!(out.len(), b * rows);
6542    let gpr = cols / GROUP_SIZE;
6543    let v = Q4tpView::new(bytes, rows, cols);
6544
6545    // Wide batches ride the AMX through a dequant-tile sgemm, as in q4t.
6546    // An f32 GEMM over dequantized weights is not the matvec's int8 sum,
6547    // so a row-exact batch never takes it.
6548    #[cfg(target_os = "macos")]
6549    if !exact && b >= 8 && rows * cols >= 500_000 && accel_gemm_enabled() {
6550        dequant_matmat_accel(
6551            &|r, dst| {
6552                let mut sc = [0f32; 32];
6553                let mut scv;
6554                let s: &[f32] = if gpr <= 32 {
6555                    v.scales_into(r, gpr, &mut sc);
6556                    &sc[..gpr]
6557                } else {
6558                    scv = vec![0f32; gpr];
6559                    v.scales_into(r, gpr, &mut scv);
6560                    &scv
6561                };
6562                for gi in 0..gpr {
6563                    let tile = &v.nib[(r * gpr + gi) * Q4TP_NIB..(r * gpr + gi + 1) * Q4TP_NIB];
6564                    for (k, &bb) in tile.iter().enumerate() {
6565                        dst[gi * GROUP_SIZE + k * 2] = ((bb & 0x0F) as f32 - 8.0) * s[gi];
6566                        dst[gi * GROUP_SIZE + k * 2 + 1] =
6567                            (((bb >> 4) & 0x0F) as f32 - 8.0) * s[gi];
6568                    }
6569                }
6570            },
6571            xs_all,
6572            b,
6573            rows,
6574            cols,
6575            out,
6576            pool,
6577        );
6578        return;
6579    }
6580
6581    let out_addr = SendMut(out.as_mut_ptr());
6582    if a8w8_enabled() {
6583        let acts: Vec<SplitAct> = (0..b)
6584            .map(|bi| split_act(&xs_all[bi * cols..(bi + 1) * cols]))
6585            .collect();
6586        let acts = &acts;
6587        // ARM stays blocked under `exact` too: the 1x4 kernel then runs in
6588        // its v1 shape, whose per-group `int dot as f32 * scale` and scalar
6589        // running sum are exactly `dot_q4tp_row_sdot`'s, so the tile is
6590        // still unpacked once for four columns and each column equals its
6591        // matvec. The tuned shape (fma into lane partials, one horizontal
6592        // add per row) is 1-16 ulp off the matvec and is kept for `!exact`.
6593        #[cfg(target_arch = "aarch64")]
6594        let blocked_ok = sdot_enabled() && blocked_enabled();
6595        // x86 gets the same blocking: one tile unpack spent on four
6596        // columns. Without it every column re-decoded the row, which is
6597        // why a 48-core EPYC measured a sixth of an M4's per-core rate.
6598        // The gate is `avx2_enabled`, as in q4t — `sdot_enabled` answers
6599        // for ARM's dotprod and is hard-wired false everywhere else, so
6600        // asking it here left the whole blocked path unreachable on x86.
6601        #[cfg(target_arch = "x86_64")]
6602        let blocked_ok = q4tp_blocked_x86() && !exact;
6603        #[cfg(not(any(target_arch = "aarch64", target_arch = "x86_64")))]
6604        let blocked_ok = {
6605            let _ = exact;
6606            false
6607        };
6608        // Columns are swept in panels that fit L2. Without this a
6609        // row-pair walks every activation in the batch — 4.8 MB at
6610        // 512x512 — and does it again for the next pair, so the whole
6611        // batch streams out of the shared cache once per row. Measured
6612        // 800 GB/s of it, flat across batch sizes, which is the signature
6613        // of a loop bound by traffic rather than by arithmetic. A panel of
6614        // 256 columns is 590 KB beside 221 KB of this worker's weights:
6615        // both stay resident and the batch crosses L3 once instead of
6616        // once per row.
6617        let panel_cols: usize = std::env::var("CMF_Q4TP_PANEL")
6618            .ok()
6619            .and_then(|v| v.parse().ok())
6620            .filter(|v| *v > 0)
6621            .unwrap_or(256);
6622        let run = |start: usize, end: usize| {
6623            for abase in (0..acts.len()).step_by(panel_cols) {
6624                let alen = (acts.len() - abase).min(panel_cols);
6625                let mut sc = vec![0f32; gpr];
6626                #[cfg(target_arch = "x86_64")]
6627                let mut r_lo = start;
6628                #[cfg(target_arch = "x86_64")]
6629                if blocked_ok && alen >= 8 {
6630                    let mut sc1 = vec![0f32; gpr];
6631                    while r_lo + 2 <= end {
6632                        v.scales_into(r_lo, gpr, &mut sc);
6633                        v.scales_into(r_lo + 1, gpr, &mut sc1);
6634                        let mut bi = 0usize;
6635                        while bi + 8 <= alen {
6636                            let xs = [
6637                                acts[abase + bi].xq.as_slice(),
6638                                acts[abase + bi + 1].xq.as_slice(),
6639                                acts[abase + bi + 2].xq.as_slice(),
6640                                acts[abase + bi + 3].xq.as_slice(),
6641                                acts[abase + bi + 4].xq.as_slice(),
6642                                acts[abase + bi + 5].xq.as_slice(),
6643                                acts[abase + bi + 6].xq.as_slice(),
6644                                acts[abase + bi + 7].xq.as_slice(),
6645                            ];
6646                            let d = unsafe { dot_q4tp_2x8_avx512(v.nib, r_lo, gpr, xs, &sc, &sc1) };
6647                            for (row, dr, scr) in [(r_lo, &d[0], &sc), (r_lo + 1, &d[1], &sc1)] {
6648                                for k in 0..8 {
6649                                    let act = &acts[abase + bi + k];
6650                                    let mut acc = dr[k] * act.sx;
6651                                    for &(j, xv) in &act.outliers {
6652                                        let (w, s) = q4tp_outlier(v.nib, row, gpr, j, scr);
6653                                        acc += w * s * xv;
6654                                    }
6655                                    // SAFETY: disjoint (bi, r) cells per worker.
6656                                    unsafe { *out_addr.at((abase + bi + k) * rows + row) = acc };
6657                                }
6658                            }
6659                            bi += 8;
6660                        }
6661                        // Columns past the last group of eight, both rows —
6662                        // the same single-row kernel the tail below uses.
6663                        for row in [r_lo, r_lo + 1] {
6664                            let scr: &[f32] = if row == r_lo { &sc } else { &sc1 };
6665                            for b2 in bi..alen {
6666                                let act = &acts[abase + b2];
6667                                let xs4 = [
6668                                    act.xq.as_slice(),
6669                                    act.xq.as_slice(),
6670                                    act.xq.as_slice(),
6671                                    act.xq.as_slice(),
6672                                ];
6673                                let d =
6674                                    unsafe { dot_q4tp_row_1x4_avx512(v.nib, row, gpr, xs4, scr) };
6675                                let mut acc = d[0] * act.sx;
6676                                for &(j, xv) in &act.outliers {
6677                                    let (w, s) = q4tp_outlier(v.nib, row, gpr, j, scr);
6678                                    acc += w * s * xv;
6679                                }
6680                                // SAFETY: disjoint (bi, r) cells per worker.
6681                                unsafe { *out_addr.at((abase + b2) * rows + row) = acc };
6682                            }
6683                        }
6684                        r_lo += 2;
6685                    }
6686                }
6687                #[cfg(target_arch = "x86_64")]
6688                let row_start = r_lo;
6689                #[cfg(not(target_arch = "x86_64"))]
6690                let row_start = start;
6691                for r in row_start..end {
6692                    v.scales_into(r, gpr, &mut sc);
6693                    let mut bi = 0usize;
6694                    #[cfg(target_arch = "x86_64")]
6695                    if blocked_ok {
6696                        while bi + 8 <= alen {
6697                            let xs = [
6698                                acts[abase + bi].xq.as_slice(),
6699                                acts[abase + bi + 1].xq.as_slice(),
6700                                acts[abase + bi + 2].xq.as_slice(),
6701                                acts[abase + bi + 3].xq.as_slice(),
6702                                acts[abase + bi + 4].xq.as_slice(),
6703                                acts[abase + bi + 5].xq.as_slice(),
6704                                acts[abase + bi + 6].xq.as_slice(),
6705                                acts[abase + bi + 7].xq.as_slice(),
6706                            ];
6707                            let d = unsafe { dot_q4tp_row_1x8_avx512(v.nib, r, gpr, xs, &sc) };
6708                            for k in 0..8 {
6709                                let act = &acts[abase + bi + k];
6710                                let mut acc = d[k] * act.sx;
6711                                for &(j, xv) in &act.outliers {
6712                                    let (w, s) = q4tp_outlier(v.nib, r, gpr, j, &sc);
6713                                    acc += w * s * xv;
6714                                }
6715                                // SAFETY: disjoint (bi, r) cells per worker.
6716                                unsafe { *out_addr.at((abase + bi + k) * rows + r) = acc };
6717                            }
6718                            bi += 8;
6719                        }
6720                        while bi + 4 <= alen {
6721                            let xs = [
6722                                acts[abase + bi].xq.as_slice(),
6723                                acts[abase + bi + 1].xq.as_slice(),
6724                                acts[abase + bi + 2].xq.as_slice(),
6725                                acts[abase + bi + 3].xq.as_slice(),
6726                            ];
6727                            let d = unsafe { dot_q4tp_row_1x4_avx512(v.nib, r, gpr, xs, &sc) };
6728                            for k in 0..4 {
6729                                let act = &acts[abase + bi + k];
6730                                let mut acc = d[k] * act.sx;
6731                                for &(j, xv) in &act.outliers {
6732                                    let (w, s) = q4tp_outlier(v.nib, r, gpr, j, &sc);
6733                                    acc += w * s * xv;
6734                                }
6735                                // SAFETY: disjoint (bi, r) cells per worker.
6736                                unsafe { *out_addr.at((abase + bi + k) * rows + r) = acc };
6737                            }
6738                            bi += 4;
6739                        }
6740                    }
6741                    #[cfg(target_arch = "aarch64")]
6742                    if blocked_ok {
6743                        while bi + 4 <= alen {
6744                            let xs = [
6745                                acts[abase + bi].xq.as_slice(),
6746                                acts[abase + bi + 1].xq.as_slice(),
6747                                acts[abase + bi + 2].xq.as_slice(),
6748                                acts[abase + bi + 3].xq.as_slice(),
6749                            ];
6750                            let d = unsafe {
6751                                if exact || q4tp_v1() {
6752                                    dot_q4tp_row_1x4_sdot_v1(v.nib, r, gpr, xs, &sc)
6753                                } else {
6754                                    dot_q4tp_row_1x4_sdot(v.nib, r, gpr, xs, &sc)
6755                                }
6756                            };
6757                            for k in 0..4 {
6758                                let act = &acts[abase + bi + k];
6759                                let mut acc = d[k] * act.sx;
6760                                for &(j, xv) in &act.outliers {
6761                                    let (w, s) = q4tp_outlier(v.nib, r, gpr, j, &sc);
6762                                    acc += w * s * xv;
6763                                }
6764                                // SAFETY: disjoint (bi, r) cells per worker.
6765                                unsafe { *out_addr.at((abase + bi + k) * rows + r) = acc };
6766                            }
6767                            bi += 4;
6768                        }
6769                    }
6770                    let _ = blocked_ok;
6771                    while bi < alen {
6772                        let act = &acts[abase + bi];
6773                        let mut acc = dot_q4tp_row_i8(v.nib, r, gpr, &act.xq, &sc) * act.sx;
6774                        for &(j, xv) in &act.outliers {
6775                            let (w, s) = q4tp_outlier(v.nib, r, gpr, j, &sc);
6776                            acc += w * s * xv;
6777                        }
6778                        // SAFETY: disjoint (bi, r) cells per worker range.
6779                        unsafe { *out_addr.at((abase + bi) * rows + r) = acc };
6780                        bi += 1;
6781                    }
6782                }
6783            }
6784        };
6785        dispatch_rows(pool, rows, &run);
6786        return;
6787    }
6788
6789    let run = |start: usize, end: usize| {
6790        let mut sc = vec![0f32; gpr];
6791        for r in start..end {
6792            v.scales_into(r, gpr, &mut sc);
6793            for bi in 0..b {
6794                let x = &xs_all[bi * cols..(bi + 1) * cols];
6795                // SAFETY: disjoint (bi, r) cells per worker range.
6796                unsafe { *out_addr.at(bi * rows + r) = q4tp_row_exact(v.nib, r, gpr, x, &sc) };
6797            }
6798        }
6799    };
6800    dispatch_rows(pool, rows, &run);
6801}
6802
6803/// Fused q4_tiled matvec (dispatch mirrors `q4matvec`).
6804fn q4t_matvec(
6805    bytes: &[u8],
6806    x: &[f32],
6807    rows: usize,
6808    cols: usize,
6809    out: &mut [f32],
6810    pool: Option<&Pool>,
6811) {
6812    debug_assert_eq!(out.len(), rows);
6813    let gpr = cols / GROUP_SIZE;
6814    let out_addr = SendMut(out.as_mut_ptr());
6815    if a8w8_enabled() {
6816        let act = split_act(x);
6817        let run = move |start: usize, end: usize| {
6818            for r in start..end {
6819                let mut acc = dot_q4t_row_i8(bytes, r, gpr, &act.xq) * act.sx;
6820                for &(j, xv) in &act.outliers {
6821                    let (w, s) = q4t_outlier(bytes, r, gpr, j);
6822                    acc += w * s * xv;
6823                }
6824                // SAFETY: disjoint row ranges per worker.
6825                unsafe { *out_addr.at(r) = acc };
6826            }
6827        };
6828        dispatch_rows(pool, rows, &run);
6829        return;
6830    }
6831    let run = move |start: usize, end: usize| {
6832        for r in start..end {
6833            // SAFETY: disjoint row ranges per worker.
6834            unsafe { *out_addr.at(r) = q4t_row_exact(bytes, r, gpr, x) };
6835        }
6836    };
6837    dispatch_rows(pool, rows, &run);
6838}
6839
6840/// Fused two-input q4_tiled matvec (weights read once per pair).
6841#[allow(clippy::too_many_arguments)]
6842fn q4t_matvec2(
6843    bytes: &[u8],
6844    x1: &[f32],
6845    x2: &[f32],
6846    rows: usize,
6847    cols: usize,
6848    o1: &mut [f32],
6849    o2: &mut [f32],
6850    pool: Option<&Pool>,
6851) {
6852    let gpr = cols / GROUP_SIZE;
6853    let p1 = SendMut(o1.as_mut_ptr());
6854    let p2 = SendMut(o2.as_mut_ptr());
6855    if a8w8_enabled() {
6856        let a1 = split_act(x1);
6857        let a2 = split_act(x2);
6858        let run = move |start: usize, end: usize| {
6859            for r in start..end {
6860                let mut v1 = dot_q4t_row_i8(bytes, r, gpr, &a1.xq) * a1.sx;
6861                let mut v2 = dot_q4t_row_i8(bytes, r, gpr, &a2.xq) * a2.sx;
6862                for &(j, xv) in &a1.outliers {
6863                    let (w, s) = q4t_outlier(bytes, r, gpr, j);
6864                    v1 += w * s * xv;
6865                }
6866                for &(j, xv) in &a2.outliers {
6867                    let (w, s) = q4t_outlier(bytes, r, gpr, j);
6868                    v2 += w * s * xv;
6869                }
6870                // SAFETY: disjoint row ranges per worker.
6871                unsafe {
6872                    *p1.at(r) = v1;
6873                    *p2.at(r) = v2;
6874                }
6875            }
6876        };
6877        dispatch_rows(pool, rows, &run);
6878        return;
6879    }
6880    let run = move |start: usize, end: usize| {
6881        for r in start..end {
6882            // SAFETY: disjoint row ranges per worker.
6883            unsafe {
6884                *p1.at(r) = q4t_row_exact(bytes, r, gpr, x1);
6885                *p2.at(r) = q4t_row_exact(bytes, r, gpr, x2);
6886            }
6887        }
6888    };
6889    dispatch_rows(pool, rows, &run);
6890}
6891
6892/// Batched q4_tiled matmat: each row's tiles stream once per microbatch.
6893#[allow(clippy::too_many_arguments)]
6894/// Prefill GEMM through Accelerate for group-quantized codecs: a
6895/// caller-supplied row dequantizer fills f32 tiles (pool-parallel) and
6896/// each tile rides the AMX with one sgemm — the generic sibling of
6897/// `qmatmat_accel` (q8). Numerics are f32-GEMM (tolerance class);
6898/// decode (b=1) never takes this path.
6899#[cfg(target_os = "macos")]
6900fn dequant_matmat_accel(
6901    dequant_row: &(dyn Fn(usize, &mut [f32]) + Sync),
6902    xs_all: &[f32],
6903    b: usize,
6904    rows: usize,
6905    cols: usize,
6906    out: &mut [f32],
6907    pool: Option<&Pool>,
6908) {
6909    const TR: usize = 2048;
6910    thread_local! {
6911        static WTILE: std::cell::RefCell<Vec<f32>> = const { std::cell::RefCell::new(Vec::new()) };
6912    }
6913    WTILE.with(|wt| {
6914        let mut wtile = wt.borrow_mut();
6915        wtile.resize(TR * cols, 0.0);
6916        let mut r0 = 0usize;
6917        while r0 < rows {
6918            let tr = TR.min(rows - r0);
6919            let wt_addr = SendMut(wtile.as_mut_ptr());
6920            let run = |start: usize, end: usize| {
6921                for r in start..end {
6922                    // SAFETY: workers cover disjoint r ranges.
6923                    let dst = unsafe { std::slice::from_raw_parts_mut(wt_addr.at(r * cols), cols) };
6924                    dequant_row(r0 + r, dst);
6925                }
6926            };
6927            dispatch_rows(pool, tr, &run);
6928            unsafe {
6929                accel_blas::cblas_sgemm(
6930                    101, // RowMajor
6931                    111, // NoTrans A
6932                    112, // Trans B
6933                    b as i32,
6934                    tr as i32,
6935                    cols as i32,
6936                    1.0,
6937                    xs_all.as_ptr(),
6938                    cols as i32,
6939                    wtile.as_ptr(),
6940                    cols as i32,
6941                    0.0,
6942                    out.as_mut_ptr().add(r0),
6943                    rows as i32,
6944                );
6945            }
6946            r0 += tr;
6947        }
6948    });
6949}
6950
6951fn q4t_matmat(
6952    bytes: &[u8],
6953    xs_all: &[f32],
6954    b: usize,
6955    rows: usize,
6956    cols: usize,
6957    out: &mut [f32],
6958    pool: Option<&Pool>,
6959) {
6960    debug_assert_eq!(out.len(), b * rows);
6961    let gpr = cols / GROUP_SIZE;
6962    // Wide batches ride the AMX like q8's qmatmat: on Apple silicon
6963    // the dequant-tile sgemm is an order above the SDOT row loop for
6964    // prefill shapes (imagegen DiT forwards are exactly this).
6965    #[cfg(target_os = "macos")]
6966    if b >= 8 && rows * cols >= 500_000 && accel_gemm_enabled() {
6967        dequant_matmat_accel(
6968            &|r, dst| {
6969                for gi in 0..gpr {
6970                    let tile = &bytes[(r * gpr + gi) * Q4_TILE..(r * gpr + gi + 1) * Q4_TILE];
6971                    let s = f16_to_f32(u16::from_le_bytes([tile[0], tile[1]]));
6972                    for (k, &bb) in tile[2..].iter().enumerate() {
6973                        dst[gi * GROUP_SIZE + k * 2] = ((bb & 0x0F) as f32 - 8.0) * s;
6974                        dst[gi * GROUP_SIZE + k * 2 + 1] = (((bb >> 4) & 0x0F) as f32 - 8.0) * s;
6975                    }
6976                }
6977            },
6978            xs_all,
6979            b,
6980            rows,
6981            cols,
6982            out,
6983            pool,
6984        );
6985        return;
6986    }
6987    let out_addr = SendMut(out.as_mut_ptr());
6988    if a8w8_enabled() {
6989        let acts: Vec<SplitAct> = (0..b)
6990            .map(|bi| split_act(&xs_all[bi * cols..(bi + 1) * cols]))
6991            .collect();
6992        let acts = &acts;
6993        #[cfg(target_arch = "x86_64")]
6994        let blocked_ok = avx2_enabled() && blocked_enabled();
6995        #[cfg(target_arch = "aarch64")]
6996        let blocked_ok = sdot_enabled() && blocked_enabled();
6997        #[cfg(not(any(target_arch = "x86_64", target_arch = "aarch64")))]
6998        let blocked_ok = false;
6999        let run = move |start: usize, end: usize| {
7000            for r in start..end {
7001                let mut bi = 0usize;
7002                #[cfg(target_arch = "aarch64")]
7003                if blocked_ok {
7004                    while bi + 4 <= acts.len() {
7005                        let xs = [
7006                            acts[bi].xq.as_slice(),
7007                            acts[bi + 1].xq.as_slice(),
7008                            acts[bi + 2].xq.as_slice(),
7009                            acts[bi + 3].xq.as_slice(),
7010                        ];
7011                        let d = unsafe { dot_q4t_row_1x4_sdot(bytes, r, gpr, xs) };
7012                        for k in 0..4 {
7013                            let act = &acts[bi + k];
7014                            let mut acc = d[k] * act.sx;
7015                            for &(j, xv) in &act.outliers {
7016                                let (w, sc) = q4t_outlier(bytes, r, gpr, j);
7017                                acc += w * sc * xv;
7018                            }
7019                            // SAFETY: disjoint (bi, r) cells per worker.
7020                            unsafe { *out_addr.at((bi + k) * rows + r) = acc };
7021                        }
7022                        bi += 4;
7023                    }
7024                }
7025                #[cfg(target_arch = "x86_64")]
7026                if blocked_ok {
7027                    while bi + 4 <= acts.len() {
7028                        let xs = [
7029                            acts[bi].xq.as_slice(),
7030                            acts[bi + 1].xq.as_slice(),
7031                            acts[bi + 2].xq.as_slice(),
7032                            acts[bi + 3].xq.as_slice(),
7033                        ];
7034                        let d = unsafe {
7035                            if vnni_tiles_enabled() {
7036                                dot_q4t_row_1x4_vnni(bytes, r, gpr, xs)
7037                            } else {
7038                                dot_q4t_row_1x4_avx2(bytes, r, gpr, xs)
7039                            }
7040                        };
7041                        for k in 0..4 {
7042                            let act = &acts[bi + k];
7043                            let mut acc = d[k] * act.sx;
7044                            for &(j, xv) in &act.outliers {
7045                                let (w, sc) = q4t_outlier(bytes, r, gpr, j);
7046                                acc += w * sc * xv;
7047                            }
7048                            // SAFETY: disjoint (bi, r) cells per worker.
7049                            unsafe { *out_addr.at((bi + k) * rows + r) = acc };
7050                        }
7051                        bi += 4;
7052                    }
7053                }
7054                let _ = blocked_ok;
7055                while bi < acts.len() {
7056                    let act = &acts[bi];
7057                    let mut acc = dot_q4t_row_i8(bytes, r, gpr, &act.xq) * act.sx;
7058                    for &(j, xv) in &act.outliers {
7059                        let (w, s) = q4t_outlier(bytes, r, gpr, j);
7060                        acc += w * s * xv;
7061                    }
7062                    // SAFETY: disjoint (bi, r) cells per worker range.
7063                    unsafe { *out_addr.at(bi * rows + r) = acc };
7064                    bi += 1;
7065                }
7066            }
7067        };
7068        dispatch_rows(pool, rows, &run);
7069        return;
7070    }
7071    let run = move |start: usize, end: usize| {
7072        for r in start..end {
7073            for bi in 0..b {
7074                let x = &xs_all[bi * cols..(bi + 1) * cols];
7075                // SAFETY: disjoint (bi, r) cells per worker range.
7076                unsafe { *out_addr.at(bi * rows + r) = q4t_row_exact(bytes, r, gpr, x) };
7077            }
7078        }
7079    };
7080    dispatch_rows(pool, rows, &run);
7081}
7082
7083// ── q1 (dtype 12): binary weights, [f16 scale][4B sign bits] per
7084// 32-group tile. The kernel family mirrors q4_tiled: one sequential
7085// stream of 6-byte tiles, per-tile integer dot × scale, exact outlier
7086// correction (A8W8 contract), exact scalar path under CMF_SDOT=0. ──
7087
7088/// Per-32-group sums of the quantized activation — the ±1 identity's
7089/// shared half: `dot = −2·sdot(mask, x) − gsum[g]`, computed ONCE per
7090/// matvec and reused by every row.
7091fn q1_group_sums(xq: &[i8], gpr: usize) -> Vec<i32> {
7092    (0..gpr)
7093        .map(|gi| {
7094            xq[gi * GROUP_SIZE..(gi + 1) * GROUP_SIZE]
7095                .iter()
7096                .map(|&v| v as i32)
7097                .sum()
7098        })
7099        .collect()
7100}
7101
7102/// One q1 row via the A8W8 int8 path — mask-SDOT on ARM (no ±1
7103/// expansion at all), scalar bit loop elsewhere (AVX2 queued with the
7104/// x86 pass).
7105#[inline]
7106#[allow(unreachable_code)]
7107/// AVX2 q1 row via the same ±1 identity as the ARM sdot kernel: the
7108/// sign bits expand to a {0, −1} byte mask through shuffle+cmpeq, the
7109/// masked activation sums through maddubs(1, x&mask), and
7110/// `dot = −(2·masked_sum + Σx_group)` — bit-identical integer math.
7111#[cfg(target_arch = "x86_64")]
7112#[target_feature(enable = "avx2")]
7113unsafe fn dot_q1_row_avx2(bytes: &[u8], r: usize, gpr: usize, xq: &[i8], gsum: &[i32]) -> f32 {
7114    // SAFETY: callers uphold the 6B-tile and xq/gsum length contracts.
7115    unsafe {
7116        use core::arch::x86_64::*;
7117        // Byte j of the mask must replicate bits-byte j/8.
7118        let expand = _mm256_setr_epi8(
7119            0, 0, 0, 0, 0, 0, 0, 0, 1, 1, 1, 1, 1, 1, 1, 1, 2, 2, 2, 2, 2, 2, 2, 2, 3, 3, 3, 3, 3,
7120            3, 3, 3,
7121        );
7122        let bitsel = _mm256_setr_epi8(
7123            1, 2, 4, 8, 16, 32, 64, -128, 1, 2, 4, 8, 16, 32, 64, -128, 1, 2, 4, 8, 16, 32, 64,
7124            -128, 1, 2, 4, 8, 16, 32, 64, -128,
7125        );
7126        let ones8 = _mm256_set1_epi8(1);
7127        let ones16 = _mm256_set1_epi16(1);
7128        let mut acc = 0f32;
7129        for gi in 0..gpr {
7130            let t = bytes.as_ptr().add((r * gpr + gi) * Q1_TILE);
7131            let s = f16_to_f32(u16::from_le_bytes([*t, *t.add(1)]));
7132            let bits = u32::from_le_bytes([*t.add(2), *t.add(3), *t.add(4), *t.add(5)]);
7133            let bc = _mm256_shuffle_epi8(_mm256_set1_epi32(bits as i32), expand);
7134            let mask = _mm256_cmpeq_epi8(_mm256_and_si256(bc, bitsel), bitsel);
7135            let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
7136            let sel = _mm256_and_si256(x, mask);
7137            // Σ of selected i8 lanes: maddubs(1u8, sel_i8) pairs → madd.
7138            let p16 = _mm256_maddubs_epi16(ones8, sel);
7139            let d32 = _mm256_madd_epi16(p16, ones16);
7140            let hi128 = _mm256_extracti128_si256::<1>(d32);
7141            let s128 = _mm_add_epi32(_mm256_castsi256_si128(d32), hi128);
7142            let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
7143            let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
7144            let msum = _mm_cvtsi128_si32(s32);
7145            // The and-select keeps x UN-negated (unlike ARM's −1-mask
7146            // sdot): d = Σ_set − Σ_unset = 2·Σ_set − Σ_all.
7147            let d = 2 * msum - gsum[gi];
7148            acc += d as f32 * s;
7149        }
7150        acc
7151    }
7152}
7153
7154/// VNNI twin of `dot_q1_row_avx2`: the masked-select sum goes through
7155/// one `vpdpbusd(1u8, sel)` (see `dpbusd_hsum` — bit-identical).
7156#[cfg(target_arch = "x86_64")]
7157#[target_feature(enable = "avx2,avx512f,avx512bw,avx512vl,avx512vnni")]
7158unsafe fn dot_q1_row_vnni(bytes: &[u8], r: usize, gpr: usize, xq: &[i8], gsum: &[i32]) -> f32 {
7159    // SAFETY: callers uphold the 6B-tile and xq/gsum length contracts.
7160    unsafe {
7161        use core::arch::x86_64::*;
7162        let expand = _mm256_setr_epi8(
7163            0, 0, 0, 0, 0, 0, 0, 0, 1, 1, 1, 1, 1, 1, 1, 1, 2, 2, 2, 2, 2, 2, 2, 2, 3, 3, 3, 3, 3,
7164            3, 3, 3,
7165        );
7166        let bitsel = _mm256_setr_epi8(
7167            1, 2, 4, 8, 16, 32, 64, -128, 1, 2, 4, 8, 16, 32, 64, -128, 1, 2, 4, 8, 16, 32, 64,
7168            -128, 1, 2, 4, 8, 16, 32, 64, -128,
7169        );
7170        let ones8 = _mm256_set1_epi8(1);
7171        let mut acc = 0f32;
7172        for gi in 0..gpr {
7173            let t = bytes.as_ptr().add((r * gpr + gi) * Q1_TILE);
7174            let s = f16_to_f32(u16::from_le_bytes([*t, *t.add(1)]));
7175            let bits = u32::from_le_bytes([*t.add(2), *t.add(3), *t.add(4), *t.add(5)]);
7176            let bc = _mm256_shuffle_epi8(_mm256_set1_epi32(bits as i32), expand);
7177            let mask = _mm256_cmpeq_epi8(_mm256_and_si256(bc, bitsel), bitsel);
7178            let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
7179            let msum = dpbusd_hsum(ones8, _mm256_and_si256(x, mask));
7180            let d = 2 * msum - gsum[gi];
7181            acc += d as f32 * s;
7182        }
7183        acc
7184    }
7185}
7186
7187/// VNNI twin of `dot_q1_row_1x4_avx2` (see `dpbusd_hsum`).
7188#[cfg(target_arch = "x86_64")]
7189#[target_feature(enable = "avx2,avx512f,avx512bw,avx512vl,avx512vnni")]
7190unsafe fn dot_q1_row_1x4_vnni(
7191    bytes: &[u8],
7192    r: usize,
7193    gpr: usize,
7194    xs: [&[i8]; 4],
7195    gsums: [&[i32]; 4],
7196) -> [f32; 4] {
7197    // SAFETY: callers uphold the 6B-tile and xq/gsum length contracts.
7198    unsafe {
7199        use core::arch::x86_64::*;
7200        let expand = _mm256_setr_epi8(
7201            0, 0, 0, 0, 0, 0, 0, 0, 1, 1, 1, 1, 1, 1, 1, 1, 2, 2, 2, 2, 2, 2, 2, 2, 3, 3, 3, 3, 3,
7202            3, 3, 3,
7203        );
7204        let bitsel = _mm256_setr_epi8(
7205            1, 2, 4, 8, 16, 32, 64, -128, 1, 2, 4, 8, 16, 32, 64, -128, 1, 2, 4, 8, 16, 32, 64,
7206            -128, 1, 2, 4, 8, 16, 32, 64, -128,
7207        );
7208        let ones8 = _mm256_set1_epi8(1);
7209        let mut acc = [0f32; 4];
7210        for gi in 0..gpr {
7211            let t = bytes.as_ptr().add((r * gpr + gi) * Q1_TILE);
7212            let s = f16_to_f32(u16::from_le_bytes([*t, *t.add(1)]));
7213            let bits = u32::from_le_bytes([*t.add(2), *t.add(3), *t.add(4), *t.add(5)]);
7214            let bc = _mm256_shuffle_epi8(_mm256_set1_epi32(bits as i32), expand);
7215            let mask = _mm256_cmpeq_epi8(_mm256_and_si256(bc, bitsel), bitsel);
7216            for (k, xq) in xs.iter().enumerate() {
7217                let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
7218                let msum = dpbusd_hsum(ones8, _mm256_and_si256(x, mask));
7219                let d = 2 * msum - gsums[k][gi];
7220                acc[k] += d as f32 * s;
7221            }
7222        }
7223        acc
7224    }
7225}
7226
7227/// The blocked 1×4 flavor: the expanded bit mask serves four activation
7228/// streams per group (mask build once, four select+reduce chains).
7229#[cfg(target_arch = "x86_64")]
7230#[target_feature(enable = "avx2")]
7231unsafe fn dot_q1_row_1x4_avx2(
7232    bytes: &[u8],
7233    r: usize,
7234    gpr: usize,
7235    xs: [&[i8]; 4],
7236    gsums: [&[i32]; 4],
7237) -> [f32; 4] {
7238    // SAFETY: callers uphold the 6B-tile and xq/gsum length contracts.
7239    unsafe {
7240        use core::arch::x86_64::*;
7241        let expand = _mm256_setr_epi8(
7242            0, 0, 0, 0, 0, 0, 0, 0, 1, 1, 1, 1, 1, 1, 1, 1, 2, 2, 2, 2, 2, 2, 2, 2, 3, 3, 3, 3, 3,
7243            3, 3, 3,
7244        );
7245        let bitsel = _mm256_setr_epi8(
7246            1, 2, 4, 8, 16, 32, 64, -128, 1, 2, 4, 8, 16, 32, 64, -128, 1, 2, 4, 8, 16, 32, 64,
7247            -128, 1, 2, 4, 8, 16, 32, 64, -128,
7248        );
7249        let ones8 = _mm256_set1_epi8(1);
7250        let ones16 = _mm256_set1_epi16(1);
7251        let mut acc = [0f32; 4];
7252        for gi in 0..gpr {
7253            let t = bytes.as_ptr().add((r * gpr + gi) * Q1_TILE);
7254            let s = f16_to_f32(u16::from_le_bytes([*t, *t.add(1)]));
7255            let bits = u32::from_le_bytes([*t.add(2), *t.add(3), *t.add(4), *t.add(5)]);
7256            let bc = _mm256_shuffle_epi8(_mm256_set1_epi32(bits as i32), expand);
7257            let mask = _mm256_cmpeq_epi8(_mm256_and_si256(bc, bitsel), bitsel);
7258            for (k, xq) in xs.iter().enumerate() {
7259                let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
7260                let sel = _mm256_and_si256(x, mask);
7261                let p16 = _mm256_maddubs_epi16(ones8, sel);
7262                let d32 = _mm256_madd_epi16(p16, ones16);
7263                let hi128 = _mm256_extracti128_si256::<1>(d32);
7264                let s128 = _mm_add_epi32(_mm256_castsi256_si128(d32), hi128);
7265                let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
7266                let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
7267                let msum = _mm_cvtsi128_si32(s32);
7268                let d = 2 * msum - gsums[k][gi];
7269                acc[k] += d as f32 * s;
7270            }
7271        }
7272        acc
7273    }
7274}
7275
7276#[allow(unreachable_code)]
7277fn dot_q1_row_i8(bytes: &[u8], r: usize, gpr: usize, xq: &[i8], gsum: &[i32]) -> f32 {
7278    #[cfg(target_arch = "aarch64")]
7279    unsafe {
7280        return dot_q1_row_sdot(bytes, r, gpr, xq, gsum);
7281    }
7282    #[cfg(target_arch = "x86_64")]
7283    if avx2_enabled() {
7284        unsafe {
7285            if vnni_tiles_enabled() {
7286                return dot_q1_row_vnni(bytes, r, gpr, xq, gsum);
7287            }
7288            return dot_q1_row_avx2(bytes, r, gpr, xq, gsum);
7289        }
7290    }
7291    let _ = gsum;
7292    let mut acc = 0f32;
7293    for gi in 0..gpr {
7294        let tile = &bytes[(r * gpr + gi) * Q1_TILE..(r * gpr + gi + 1) * Q1_TILE];
7295        let s = f16_to_f32(u16::from_le_bytes([tile[0], tile[1]]));
7296        let mut d = 0i32;
7297        for (j, &b) in tile[2..].iter().enumerate() {
7298            for k in 0..8 {
7299                let w = ((b >> k) & 1) as i32 * 2 - 1;
7300                d += w * xq[gi * GROUP_SIZE + j * 8 + k] as i32;
7301            }
7302        }
7303        acc += d as f32 * s;
7304    }
7305    acc
7306}
7307
7308/// SDOT q1 row via the ±1 identity: the vtst mask (0xFF where the bit
7309/// is set, i.e. −1 as i8) feeds `sdot` DIRECTLY — no expansion to ±1
7310/// lanes at all — and `dot = −(2·sdot(mask, x) + Σx_group)`, with the
7311/// per-group activation sums shared across every row of the matvec.
7312/// Four tiles (128 weights) per iteration: integer dots reduce through
7313/// a vpaddq tree into ONE i32x4 that meets its four scales in a single
7314/// fused f32 multiply-add. Integer math throughout — bit-identical to
7315/// the scalar ±1 reference.
7316#[cfg(target_arch = "aarch64")]
7317#[target_feature(enable = "neon,dotprod")]
7318unsafe fn dot_q1_row_sdot(bytes: &[u8], r: usize, gpr: usize, xq: &[i8], gsum: &[i32]) -> f32 {
7319    // SAFETY: callers uphold slice-length contracts (6B tile per group,
7320    // xq.len() == gpr·GROUP_SIZE, gsum.len() == gpr).
7321    unsafe {
7322        use core::arch::aarch64::*;
7323        use core::arch::asm;
7324        const MASKS: [u8; 16] = [1, 2, 4, 8, 16, 32, 64, 128, 1, 2, 4, 8, 16, 32, 64, 128];
7325        let m = vld1q_u8(MASKS.as_ptr());
7326        // One tile's −Σ_set(x) as an UNREDUCED i32x4 (two mask-sdots).
7327        macro_rules! tile_dot {
7328            ($t:expr, $x:expr) => {{
7329                let v0 = vcombine_u8(vdup_n_u8(*$t.add(2)), vdup_n_u8(*$t.add(3)));
7330                let v1 = vcombine_u8(vdup_n_u8(*$t.add(4)), vdup_n_u8(*$t.add(5)));
7331                let w0 = vreinterpretq_s8_u8(vtstq_u8(v0, m));
7332                let w1 = vreinterpretq_s8_u8(vtstq_u8(v1, m));
7333                let x0 = vld1q_s8($x);
7334                let x1 = vld1q_s8($x.add(16));
7335                let (mut a0, mut a1) = (vdupq_n_s32(0), vdupq_n_s32(0));
7336                asm!(
7337                    "sdot {a0:v}.4s, {w0:v}.16b, {x0:v}.16b",
7338                    "sdot {a1:v}.4s, {w1:v}.16b, {x1:v}.16b",
7339                    a0 = inout(vreg) a0, a1 = inout(vreg) a1,
7340                    w0 = in(vreg) w0, x0 = in(vreg) x0, w1 = in(vreg) w1, x1 = in(vreg) x1,
7341                    options(pure, nomem, nostack),
7342                );
7343                vaddq_s32(a0, a1)
7344            }};
7345        }
7346        // TBL unpack over PAIR loads: one vld1q covers two 6B tiles
7347        // ([s s b b b b][s s b b b b] + 4B slack), TBL replicates each
7348        // bit-byte across 8 lanes for vtst, and the four scales gather
7349        // through tbl2 into one fcvtl — the 16 ld1r broadcast loads and
7350        // 4 branchy software f16 conversions per 128 weights (the
7351        // measured load-port wall of this kernel) become 2 vector
7352        // loads + 9 table lookups. Integer math order is unchanged —
7353        // bit-identical results (FCVTL is exact on every f16).
7354        const IW00: [u8; 16] = [2, 2, 2, 2, 2, 2, 2, 2, 3, 3, 3, 3, 3, 3, 3, 3];
7355        const IW01: [u8; 16] = [4, 4, 4, 4, 4, 4, 4, 4, 5, 5, 5, 5, 5, 5, 5, 5];
7356        const IW10: [u8; 16] = [8, 8, 8, 8, 8, 8, 8, 8, 9, 9, 9, 9, 9, 9, 9, 9];
7357        const IW11: [u8; 16] = [
7358            10, 10, 10, 10, 10, 10, 10, 10, 11, 11, 11, 11, 11, 11, 11, 11,
7359        ];
7360        const ISC: [u8; 8] = [0, 1, 6, 7, 16, 17, 22, 23];
7361        let (iw00, iw01) = (vld1q_u8(IW00.as_ptr()), vld1q_u8(IW01.as_ptr()));
7362        let (iw10, iw11) = (vld1q_u8(IW10.as_ptr()), vld1q_u8(IW11.as_ptr()));
7363        let isc = vld1_u8(ISC.as_ptr());
7364        // One tile's −Σ_set(x) from a TBL-unpacked pair load.
7365        macro_rules! tile_dot_tbl {
7366            ($ld:expr, $i0:expr, $i1:expr, $x:expr) => {{
7367                let w0 = vreinterpretq_s8_u8(vtstq_u8(vqtbl1q_u8($ld, $i0), m));
7368                let w1 = vreinterpretq_s8_u8(vtstq_u8(vqtbl1q_u8($ld, $i1), m));
7369                let x0 = vld1q_s8($x);
7370                let x1 = vld1q_s8($x.add(16));
7371                let (mut a0, mut a1) = (vdupq_n_s32(0), vdupq_n_s32(0));
7372                asm!(
7373                    "sdot {a0:v}.4s, {w0:v}.16b, {x0:v}.16b",
7374                    "sdot {a1:v}.4s, {w1:v}.16b, {x1:v}.16b",
7375                    a0 = inout(vreg) a0, a1 = inout(vreg) a1,
7376                    w0 = in(vreg) w0, x0 = in(vreg) x0, w1 = in(vreg) w1, x1 = in(vreg) x1,
7377                    options(pure, nomem, nostack),
7378                );
7379                vaddq_s32(a0, a1)
7380            }};
7381        }
7382        let base = bytes.as_ptr().add(r * gpr * Q1_TILE);
7383        let row_base = r * gpr * Q1_TILE;
7384        let abs_end = bytes.len();
7385        let xp = xq.as_ptr();
7386        let gp = gsum.as_ptr();
7387        let mut accv = vdupq_n_f32(0.0);
7388        let mut gi = 0;
7389        // The second pair load reads 4B past tile gi+3 — stay inside
7390        // the payload slice (only the file's final tiles fall back).
7391        while gi + 4 <= gpr && row_base + (gi + 4) * Q1_TILE + 4 <= abs_end {
7392            let t0 = base.add(gi * Q1_TILE);
7393            let ld_a = vld1q_u8(t0);
7394            let ld_b = vld1q_u8(t0.add(2 * Q1_TILE));
7395            let d0 = tile_dot_tbl!(ld_a, iw00, iw01, xp.add(gi * GROUP_SIZE));
7396            let d1 = tile_dot_tbl!(ld_a, iw10, iw11, xp.add((gi + 1) * GROUP_SIZE));
7397            let d2 = tile_dot_tbl!(ld_b, iw00, iw01, xp.add((gi + 2) * GROUP_SIZE));
7398            let d3 = tile_dot_tbl!(ld_b, iw10, iw11, xp.add((gi + 3) * GROUP_SIZE));
7399            // [−Σ0, −Σ1, −Σ2, −Σ3] → dots = −(2·Σset_neg + gsum)
7400            let neg = vpaddq_s32(vpaddq_s32(d0, d1), vpaddq_s32(d2, d3));
7401            let g = vld1q_s32(gp.add(gi));
7402            let dots = vnegq_s32(vaddq_s32(vshlq_n_s32::<1>(neg), g));
7403            let sc16 = vqtbl2_u8(uint8x16x2_t(ld_a, ld_b), isc);
7404            let scf: float32x4_t;
7405            asm!(
7406                "fcvtl {o:v}.4s, {i:v}.4h",
7407                o = out(vreg) scf, i = in(vreg) sc16,
7408                options(pure, nomem, nostack),
7409            );
7410            accv = vfmaq_f32(accv, vcvtq_f32_s32(dots), scf);
7411            gi += 4;
7412        }
7413        let mut acc = vaddvq_f32(accv);
7414        while gi < gpr {
7415            let t = base.add(gi * Q1_TILE);
7416            let s = f16_to_f32(u16::from_le_bytes([*t, *t.add(1)]));
7417            let d = vaddvq_s32(tile_dot!(t, xp.add(gi * GROUP_SIZE)));
7418            acc += (-(2 * d + *gp.add(gi))) as f32 * s;
7419            gi += 1;
7420        }
7421        acc
7422    }
7423}
7424
7425/// Blocked q1 1×4: one TBL unpack of the tile pair serves FOUR
7426/// activation streams (prefill amortization — the same idea as the
7427/// AVX2 twin; per stream the group order, fma order and tail match the
7428/// single-row kernel exactly, so batch == matvec bit-for-bit).
7429#[cfg(target_arch = "aarch64")]
7430#[target_feature(enable = "neon,dotprod")]
7431unsafe fn dot_q1_row_1x4_sdot(
7432    bytes: &[u8],
7433    r: usize,
7434    gpr: usize,
7435    xs: [&[i8]; 4],
7436    gs: [&[i32]; 4],
7437) -> [f32; 4] {
7438    // SAFETY: same slice-length contracts as `dot_q1_row_sdot`, ×4.
7439    unsafe {
7440        use core::arch::aarch64::*;
7441        use core::arch::asm;
7442        const MASKS: [u8; 16] = [1, 2, 4, 8, 16, 32, 64, 128, 1, 2, 4, 8, 16, 32, 64, 128];
7443        const IW00: [u8; 16] = [2, 2, 2, 2, 2, 2, 2, 2, 3, 3, 3, 3, 3, 3, 3, 3];
7444        const IW01: [u8; 16] = [4, 4, 4, 4, 4, 4, 4, 4, 5, 5, 5, 5, 5, 5, 5, 5];
7445        const IW10: [u8; 16] = [8, 8, 8, 8, 8, 8, 8, 8, 9, 9, 9, 9, 9, 9, 9, 9];
7446        const IW11: [u8; 16] = [
7447            10, 10, 10, 10, 10, 10, 10, 10, 11, 11, 11, 11, 11, 11, 11, 11,
7448        ];
7449        const ISC: [u8; 8] = [0, 1, 6, 7, 16, 17, 22, 23];
7450        let m = vld1q_u8(MASKS.as_ptr());
7451        let (iw00, iw01) = (vld1q_u8(IW00.as_ptr()), vld1q_u8(IW01.as_ptr()));
7452        let (iw10, iw11) = (vld1q_u8(IW10.as_ptr()), vld1q_u8(IW11.as_ptr()));
7453        let isc = vld1_u8(ISC.as_ptr());
7454        macro_rules! sdot2 {
7455            ($w0:expr, $w1:expr, $x:expr) => {{
7456                let x0 = vld1q_s8($x);
7457                let x1 = vld1q_s8($x.add(16));
7458                let (mut a0, mut a1) = (vdupq_n_s32(0), vdupq_n_s32(0));
7459                asm!(
7460                    "sdot {a0:v}.4s, {w0:v}.16b, {x0:v}.16b",
7461                    "sdot {a1:v}.4s, {w1:v}.16b, {x1:v}.16b",
7462                    a0 = inout(vreg) a0, a1 = inout(vreg) a1,
7463                    w0 = in(vreg) $w0, x0 = in(vreg) x0, w1 = in(vreg) $w1, x1 = in(vreg) x1,
7464                    options(pure, nomem, nostack),
7465                );
7466                vaddq_s32(a0, a1)
7467            }};
7468        }
7469        let base = bytes.as_ptr().add(r * gpr * Q1_TILE);
7470        let row_base = r * gpr * Q1_TILE;
7471        let abs_end = bytes.len();
7472        let mut accv = [vdupq_n_f32(0.0); 4];
7473        let mut gi = 0;
7474        while gi + 4 <= gpr && row_base + (gi + 4) * Q1_TILE + 4 <= abs_end {
7475            let t0 = base.add(gi * Q1_TILE);
7476            let ld_a = vld1q_u8(t0);
7477            let ld_b = vld1q_u8(t0.add(2 * Q1_TILE));
7478            // Unpack ONCE — eight ±mask vectors serve all four streams.
7479            let w00 = vreinterpretq_s8_u8(vtstq_u8(vqtbl1q_u8(ld_a, iw00), m));
7480            let w01 = vreinterpretq_s8_u8(vtstq_u8(vqtbl1q_u8(ld_a, iw01), m));
7481            let w10 = vreinterpretq_s8_u8(vtstq_u8(vqtbl1q_u8(ld_a, iw10), m));
7482            let w11 = vreinterpretq_s8_u8(vtstq_u8(vqtbl1q_u8(ld_a, iw11), m));
7483            let w20 = vreinterpretq_s8_u8(vtstq_u8(vqtbl1q_u8(ld_b, iw00), m));
7484            let w21 = vreinterpretq_s8_u8(vtstq_u8(vqtbl1q_u8(ld_b, iw01), m));
7485            let w30 = vreinterpretq_s8_u8(vtstq_u8(vqtbl1q_u8(ld_b, iw10), m));
7486            let w31 = vreinterpretq_s8_u8(vtstq_u8(vqtbl1q_u8(ld_b, iw11), m));
7487            let sc16 = vqtbl2_u8(uint8x16x2_t(ld_a, ld_b), isc);
7488            let scf: float32x4_t;
7489            asm!(
7490                "fcvtl {o:v}.4s, {i:v}.4h",
7491                o = out(vreg) scf, i = in(vreg) sc16,
7492                options(pure, nomem, nostack),
7493            );
7494            for k in 0..4 {
7495                let xp = xs[k].as_ptr();
7496                let d0 = sdot2!(w00, w01, xp.add(gi * GROUP_SIZE));
7497                let d1 = sdot2!(w10, w11, xp.add((gi + 1) * GROUP_SIZE));
7498                let d2 = sdot2!(w20, w21, xp.add((gi + 2) * GROUP_SIZE));
7499                let d3 = sdot2!(w30, w31, xp.add((gi + 3) * GROUP_SIZE));
7500                let neg = vpaddq_s32(vpaddq_s32(d0, d1), vpaddq_s32(d2, d3));
7501                let g = vld1q_s32(gs[k].as_ptr().add(gi));
7502                let dots = vnegq_s32(vaddq_s32(vshlq_n_s32::<1>(neg), g));
7503                accv[k] = vfmaq_f32(accv[k], vcvtq_f32_s32(dots), scf);
7504            }
7505            gi += 4;
7506        }
7507        let mut acc = [
7508            vaddvq_f32(accv[0]),
7509            vaddvq_f32(accv[1]),
7510            vaddvq_f32(accv[2]),
7511            vaddvq_f32(accv[3]),
7512        ];
7513        while gi < gpr {
7514            let t = base.add(gi * Q1_TILE);
7515            let sc = f16_to_f32(u16::from_le_bytes([*t, *t.add(1)]));
7516            let v0 = vcombine_u8(vdup_n_u8(*t.add(2)), vdup_n_u8(*t.add(3)));
7517            let v1 = vcombine_u8(vdup_n_u8(*t.add(4)), vdup_n_u8(*t.add(5)));
7518            let w0 = vreinterpretq_s8_u8(vtstq_u8(v0, m));
7519            let w1 = vreinterpretq_s8_u8(vtstq_u8(v1, m));
7520            for k in 0..4 {
7521                let d = vaddvq_s32(sdot2!(w0, w1, xs[k].as_ptr().add(gi * GROUP_SIZE)));
7522                acc[k] += (-(2 * d + *gs[k].as_ptr().add(gi))) as f32 * sc;
7523            }
7524            gi += 1;
7525        }
7526        acc
7527    }
7528}
7529
7530/// (weight ±1, scale) of one q1 element — the exact outlier term.
7531#[inline]
7532fn q1_outlier(bytes: &[u8], r: usize, gpr: usize, j: usize) -> (f32, f32) {
7533    let gi = j / GROUP_SIZE;
7534    let k = j % GROUP_SIZE;
7535    let tile = &bytes[(r * gpr + gi) * Q1_TILE..(r * gpr + gi + 1) * Q1_TILE];
7536    let s = f16_to_f32(u16::from_le_bytes([tile[0], tile[1]]));
7537    let bit = (tile[2 + k / 8] >> (k % 8)) & 1;
7538    ((bit as i32 * 2 - 1) as f32, s)
7539}
7540
7541/// Exact scalar q1 row (CMF_SDOT=0 contract).
7542#[inline]
7543fn q1_row_exact(bytes: &[u8], r: usize, gpr: usize, x: &[f32]) -> f32 {
7544    let mut acc = 0f32;
7545    for gi in 0..gpr {
7546        let tile = &bytes[(r * gpr + gi) * Q1_TILE..(r * gpr + gi + 1) * Q1_TILE];
7547        let s = f16_to_f32(u16::from_le_bytes([tile[0], tile[1]]));
7548        let xg = &x[gi * GROUP_SIZE..(gi + 1) * GROUP_SIZE];
7549        let mut ga = 0f32;
7550        for (j, &b) in tile[2..].iter().enumerate() {
7551            for k in 0..8 {
7552                ga += (((b >> k) & 1) as f32 * 2.0 - 1.0) * xg[j * 8 + k];
7553            }
7554        }
7555        acc += ga * s;
7556    }
7557    acc
7558}
7559
7560/// One q1 row range via A8W8 (the body of `q1_matvec`'s hot loop,
7561/// extracted so multi-matrix jobs drive the same kernel).
7562#[allow(clippy::too_many_arguments)]
7563fn q1_range_a8w8(
7564    bytes: &[u8],
7565    gpr: usize,
7566    act: &SplitAct,
7567    gsum: &[i32],
7568    out: SendMut,
7569    start: usize,
7570    end: usize,
7571) {
7572    for r in start..end {
7573        let mut acc = dot_q1_row_i8(bytes, r, gpr, &act.xq, gsum) * act.sx;
7574        for &(j, xv) in &act.outliers {
7575            let (w, s) = q1_outlier(bytes, r, gpr, j);
7576            acc += w * s * xv;
7577        }
7578        // SAFETY: disjoint row ranges per worker.
7579        unsafe { *out.at(r) = acc };
7580    }
7581}
7582
7583/// Exact-scalar q1 row range (CMF_SDOT=0 contract).
7584fn q1_range_f32(bytes: &[u8], gpr: usize, x: &[f32], out: SendMut, start: usize, end: usize) {
7585    for r in start..end {
7586        // SAFETY: disjoint row ranges per worker.
7587        unsafe { *out.at(r) = q1_row_exact(bytes, r, gpr, x) };
7588    }
7589}
7590
7591/// q1t per-row overlay locator. After the base (`base_len`) come
7592/// `[u32 row_ptr[rows+1]]` then `[(u16 col, f16 val)]` grouped by row (row
7593/// `r`'s entries are `[row_ptr[r], row_ptr[r+1])`). Returns
7594/// `(row_ptr offset, entries offset, present)`.
7595fn q1t_overlay(bytes: &[u8], base_len: usize, rows: usize) -> (usize, usize, bool) {
7596    let entries = base_len + (rows + 1) * 4;
7597    (base_len, entries, entries <= bytes.len())
7598}
7599
7600/// Read `row_ptr[r]` from the overlay's prefix-sum table.
7601#[inline]
7602fn q1t_rowptr(bytes: &[u8], rp_off: usize, r: usize) -> usize {
7603    let o = rp_off + r * 4;
7604    u32::from_le_bytes([bytes[o], bytes[o + 1], bytes[o + 2], bytes[o + 3]]) as usize
7605}
7606
7607/// Byte → the 5 ternary signs it packs `{−1,0,+1}` as f32, precomputed so
7608/// decoding a q1t code is a table load, not the base-3 divide/modulo per
7609/// weight (division is ~20–40× the cost of a load). Built at compile time.
7610const SIGN5: [[f32; 5]; 256] = {
7611    let mut lut = [[0.0f32; 5]; 256];
7612    let pow3 = [1u16, 3, 9, 27, 81];
7613    let mut byte = 0usize;
7614    while byte < 256 {
7615        let mut i = 0usize;
7616        while i < 5 {
7617            let code = (byte as u16 / pow3[i]) % 3;
7618            lut[byte][i] = if code == 1 {
7619                1.0
7620            } else if code == 2 {
7621                -1.0
7622            } else {
7623                0.0
7624            };
7625            i += 1;
7626        }
7627        byte += 1;
7628    }
7629    lut
7630};
7631
7632/// Same table, as i8 signs — the operand for the int8 SDOT base kernel.
7633const SIGN5_I8: [[i8; 5]; 256] = {
7634    let mut lut = [[0i8; 5]; 256];
7635    let pow3 = [1u16, 3, 9, 27, 81];
7636    let mut byte = 0usize;
7637    while byte < 256 {
7638        let mut i = 0usize;
7639        while i < 5 {
7640            let code = (byte as u16 / pow3[i]) % 3;
7641            lut[byte][i] = if code == 1 {
7642                1
7643            } else if code == 2 {
7644                -1
7645            } else {
7646                0
7647            };
7648            i += 1;
7649        }
7650        byte += 1;
7651    }
7652    lut
7653};
7654
7655/// The same 5 i8 signs packed into a u64 (`[s0 s1 s2 s3 s4 0 0 0]`, LE) so the
7656/// group unpack is 7 unaligned u64 stores at offsets 0,5,10,…,30 instead of
7657/// six 5-byte copies + LUT indexing — each store's trailing zeros are fixed by
7658/// the next store, and the last one runs 6 B past the 32nd weight (the unpack
7659/// buffer is padded to 40). This is the decode/prefill hot inner op.
7660const SIGN5_U64: [u64; 256] = {
7661    let mut lut = [0u64; 256];
7662    let pow3 = [1u16, 3, 9, 27, 81];
7663    let mut byte = 0usize;
7664    while byte < 256 {
7665        let mut v = 0u64;
7666        let mut i = 0usize;
7667        while i < 5 {
7668            let code = (byte as u16 / pow3[i]) % 3;
7669            let s: u8 = if code == 1 {
7670                1
7671            } else if code == 2 {
7672                0xFF
7673            } else {
7674                0
7675            };
7676            v |= (s as u64) << (i * 8);
7677            i += 1;
7678        }
7679        lut[byte] = v;
7680        byte += 1;
7681    }
7682    lut
7683};
7684
7685/// Ternary base weight at `(row r, col j)` = `sign(code)·s_group`. Used to add
7686/// back activation-outlier columns, whose `x` was zeroed for the int8 bulk dot
7687/// (`split_act`). At a weight-outlier position the code is 0, so this is 0 and
7688/// the overlay correction owns that column — no double counting.
7689#[inline]
7690fn q1t_base_weight(bytes: &[u8], r: usize, gpr: usize, j: usize) -> f32 {
7691    const TILE: usize = cortiq_core::quant::Q1T_TILE;
7692    let off = (r * gpr + j / GROUP_SIZE) * TILE;
7693    let s = f16_to_f32(u16::from_le_bytes([bytes[off], bytes[off + 1]]));
7694    let within = j % GROUP_SIZE;
7695    SIGN5[bytes[off + 2 + within / 5] as usize][within % 5] * s
7696}
7697
7698/// One 32-group int8 dot via two SDOTs. Bit-exact vs the scalar i8 sum
7699/// (integer accumulation is order-independent).
7700#[cfg(target_arch = "aarch64")]
7701#[target_feature(enable = "neon,dotprod")]
7702#[inline]
7703unsafe fn sdot32_i8(w: *const i8, x: *const i8) -> i32 {
7704    // SAFETY: caller guarantees 32 readable i8 at each pointer.
7705    unsafe {
7706        use core::arch::aarch64::*;
7707        use core::arch::asm;
7708        let w0 = vld1q_s8(w);
7709        let w1 = vld1q_s8(w.add(16));
7710        let x0 = vld1q_s8(x);
7711        let x1 = vld1q_s8(x.add(16));
7712        let (mut a0, mut a1) = (vdupq_n_s32(0), vdupq_n_s32(0));
7713        asm!(
7714            "sdot {a0:v}.4s, {w0:v}.16b, {x0:v}.16b",
7715            "sdot {a1:v}.4s, {w1:v}.16b, {x1:v}.16b",
7716            a0 = inout(vreg) a0, a1 = inout(vreg) a1,
7717            w0 = in(vreg) w0, x0 = in(vreg) x0, w1 = in(vreg) w1, x1 = in(vreg) x1,
7718            options(pure, nomem, nostack),
7719        );
7720        vaddvq_s32(vaddq_s32(a0, a1))
7721    }
7722}
7723
7724/// One 32-group int8 dot via AVX2: signed·signed as `maddubs(|w|, sign(x,w))`
7725/// then `madd` and a horizontal reduce (the same idiom as `dot_q4t_row_avx2`).
7726#[cfg(target_arch = "x86_64")]
7727#[target_feature(enable = "avx2")]
7728#[inline]
7729unsafe fn i8dot32_avx2(w: *const i8, x: *const i8) -> i32 {
7730    // SAFETY: caller guarantees 32 readable i8 at each pointer.
7731    unsafe {
7732        use core::arch::x86_64::*;
7733        let wv = _mm256_loadu_si256(w as *const __m256i);
7734        let xv = _mm256_loadu_si256(x as *const __m256i);
7735        let p16 = _mm256_maddubs_epi16(_mm256_abs_epi8(wv), _mm256_sign_epi8(xv, wv));
7736        let d = _mm256_madd_epi16(p16, _mm256_set1_epi16(1));
7737        let hi128 = _mm256_extracti128_si256::<1>(d);
7738        let s128 = _mm_add_epi32(_mm256_castsi256_si128(d), hi128);
7739        let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
7740        let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
7741        _mm_cvtsi128_si32(s32)
7742    }
7743}
7744
7745/// Unpack one q1t group's base-3 codes into 32 i8 signs via 7 unaligned u64
7746/// stores (see `SIGN5_U64`). `dst` MUST have ≥ 40 bytes: the 7th store writes
7747/// `dst[30..38]`. Stores go in order so each one's trailing zeros are
7748/// overwritten by the next; the final 6 padding bytes are unused by the dot.
7749#[inline]
7750fn q1t_unpack_group_i8(codes: *const u8, dst: &mut [i8]) {
7751    debug_assert!(dst.len() >= 40);
7752    // SAFETY: codes points at 7 readable bytes; dst has ≥ 40 bytes so every
7753    // 8-byte store at offset bi*5 (bi ≤ 6 → ≤ 30) stays in bounds.
7754    unsafe {
7755        let p = dst.as_mut_ptr();
7756        for bi in 0..7 {
7757            core::ptr::write_unaligned(
7758                p.add(bi * 5) as *mut u64,
7759                SIGN5_U64[*codes.add(bi) as usize],
7760            );
7761        }
7762    }
7763}
7764
7765/// One 32-group int8 dot, arch-dispatched (the matmat inner loop, where the
7766/// row's signs are unpacked once and dotted against every batch input).
7767/// Callers are gated by `a8w8_enabled()`, so the target-feature arms are
7768/// reachable; the scalar arm is a non-SIMD-arch fallback.
7769#[inline]
7770fn q1t_i8dot32(w: *const i8, x: *const i8) -> i32 {
7771    #[cfg(target_arch = "aarch64")]
7772    unsafe {
7773        return sdot32_i8(w, x);
7774    }
7775    #[cfg(target_arch = "x86_64")]
7776    unsafe {
7777        return i8dot32_avx2(w, x);
7778    }
7779    #[allow(unreachable_code)]
7780    unsafe {
7781        let mut s = 0i32;
7782        for k in 0..GROUP_SIZE {
7783            s += *w.add(k) as i32 * *x.add(k) as i32;
7784        }
7785        s
7786    }
7787}
7788
7789#[inline]
7790unsafe fn q1t_unpack_reg_u64s(codes: *const u8) -> (u64, u64, u64, u64) {
7791    let (s0, s1, s2, s3, s4, s5, s6) = unsafe {
7792        (
7793            SIGN5_U64[*codes as usize],
7794            SIGN5_U64[*codes.add(1) as usize],
7795            SIGN5_U64[*codes.add(2) as usize],
7796            SIGN5_U64[*codes.add(3) as usize],
7797            SIGN5_U64[*codes.add(4) as usize],
7798            SIGN5_U64[*codes.add(5) as usize],
7799            SIGN5_U64[*codes.add(6) as usize],
7800        )
7801    };
7802
7803    let u0 = s0 | (s1 << 40);
7804    let u1 = (s1 >> 24) | (s2 << 16) | (s3 << 56);
7805    let u2 = (s3 >> 8) | (s4 << 32);
7806    let u3 = (s4 >> 32) | (s5 << 8) | (s6 << 48);
7807
7808    (u0, u1, u2, u3)
7809}
7810
7811/// One q1t row's int8 base dot: `Σ_group s·dot(signs, xq)` (before the shared
7812/// `sx`). Direct register unpacking (zero stack stores/loads, no STLF stalls).
7813/// ARM SDOT.
7814#[cfg(target_arch = "aarch64")]
7815#[target_feature(enable = "neon,dotprod")]
7816unsafe fn q1t_dot_row_sdot(bytes: &[u8], r: usize, gpr: usize, xq: &[i8]) -> f32 {
7817    use core::arch::aarch64::*;
7818    use core::arch::asm;
7819    unsafe {
7820        const TILE: usize = cortiq_core::quant::Q1T_TILE;
7821        let mut acc = 0f32;
7822        let bytes_ptr = bytes.as_ptr();
7823        let xq_ptr = xq.as_ptr();
7824        let row_off = r * gpr * TILE;
7825
7826        let gpr2 = gpr & !1;
7827        let mut gi = 0;
7828        while gi < gpr2 {
7829            let off0 = row_off + gi * TILE;
7830            let off1 = off0 + TILE;
7831            let s0 = f16_to_f32(u16::from_le_bytes([
7832                *bytes_ptr.add(off0),
7833                *bytes_ptr.add(off0 + 1),
7834            ]));
7835            let s1 = f16_to_f32(u16::from_le_bytes([
7836                *bytes_ptr.add(off1),
7837                *bytes_ptr.add(off1 + 1),
7838            ]));
7839
7840            let (u0_0, u1_0, u2_0, u3_0) = q1t_unpack_reg_u64s(bytes_ptr.add(off0 + 2));
7841            let (u0_1, u1_1, u2_1, u3_1) = q1t_unpack_reg_u64s(bytes_ptr.add(off1 + 2));
7842
7843            let w0_0 = vreinterpretq_s8_u64(vcombine_u64(vcreate_u64(u0_0), vcreate_u64(u1_0)));
7844            let w1_0 = vreinterpretq_s8_u64(vcombine_u64(vcreate_u64(u2_0), vcreate_u64(u3_0)));
7845            let w0_1 = vreinterpretq_s8_u64(vcombine_u64(vcreate_u64(u0_1), vcreate_u64(u1_1)));
7846            let w1_1 = vreinterpretq_s8_u64(vcombine_u64(vcreate_u64(u2_1), vcreate_u64(u3_1)));
7847
7848            let x0_0 = vld1q_s8(xq_ptr.add(gi * GROUP_SIZE));
7849            let x1_0 = vld1q_s8(xq_ptr.add(gi * GROUP_SIZE + 16));
7850            let x0_1 = vld1q_s8(xq_ptr.add((gi + 1) * GROUP_SIZE));
7851            let x1_1 = vld1q_s8(xq_ptr.add((gi + 1) * GROUP_SIZE + 16));
7852
7853            let (mut a0_0, mut a1_0) = (vdupq_n_s32(0), vdupq_n_s32(0));
7854            let (mut a0_1, mut a1_1) = (vdupq_n_s32(0), vdupq_n_s32(0));
7855            asm!(
7856                "sdot {a0_0:v}.4s, {w0_0:v}.16b, {x0_0:v}.16b",
7857                "sdot {a1_0:v}.4s, {w1_0:v}.16b, {x1_0:v}.16b",
7858                "sdot {a0_1:v}.4s, {w0_1:v}.16b, {x0_1:v}.16b",
7859                "sdot {a1_1:v}.4s, {w1_1:v}.16b, {x1_1:v}.16b",
7860                a0_0 = inout(vreg) a0_0, a1_0 = inout(vreg) a1_0,
7861                a0_1 = inout(vreg) a0_1, a1_1 = inout(vreg) a1_1,
7862                w0_0 = in(vreg) w0_0, x0_0 = in(vreg) x0_0, w1_0 = in(vreg) w1_0, x1_0 = in(vreg) x1_0,
7863                w0_1 = in(vreg) w0_1, x0_1 = in(vreg) x0_1, w1_1 = in(vreg) w1_1, x1_1 = in(vreg) x1_1,
7864                options(pure, nomem, nostack),
7865            );
7866            let d0 = vaddvq_s32(vaddq_s32(a0_0, a1_0));
7867            let d1 = vaddvq_s32(vaddq_s32(a0_1, a1_1));
7868            acc += d0 as f32 * s0 + d1 as f32 * s1;
7869            gi += 2;
7870        }
7871
7872        if gi < gpr {
7873            let off = row_off + gi * TILE;
7874            let s = f16_to_f32(u16::from_le_bytes([
7875                *bytes_ptr.add(off),
7876                *bytes_ptr.add(off + 1),
7877            ]));
7878            let (u0, u1, u2, u3) = q1t_unpack_reg_u64s(bytes_ptr.add(off + 2));
7879            let w0 = vreinterpretq_s8_u64(vcombine_u64(vcreate_u64(u0), vcreate_u64(u1)));
7880            let w1 = vreinterpretq_s8_u64(vcombine_u64(vcreate_u64(u2), vcreate_u64(u3)));
7881            let x0 = vld1q_s8(xq_ptr.add(gi * GROUP_SIZE));
7882            let x1 = vld1q_s8(xq_ptr.add(gi * GROUP_SIZE + 16));
7883            let (mut a0, mut a1) = (vdupq_n_s32(0), vdupq_n_s32(0));
7884            asm!(
7885                "sdot {a0:v}.4s, {w0:v}.16b, {x0:v}.16b",
7886                "sdot {a1:v}.4s, {w1:v}.16b, {x1:v}.16b",
7887                a0 = inout(vreg) a0, a1 = inout(vreg) a1,
7888                w0 = in(vreg) w0, x0 = in(vreg) x0, w1 = in(vreg) w1, x1 = in(vreg) x1,
7889                options(pure, nomem, nostack),
7890            );
7891            let d = vaddvq_s32(vaddq_s32(a0, a1));
7892            acc += d as f32 * s;
7893        }
7894        acc
7895    }
7896}
7897
7898/// x86 AVX2 mirror of `q1t_dot_row_sdot` (maddubs int8 dot per group).
7899#[cfg(target_arch = "x86_64")]
7900#[target_feature(enable = "avx2")]
7901unsafe fn q1t_dot_row_avx2(bytes: &[u8], r: usize, gpr: usize, xq: &[i8]) -> f32 {
7902    use core::arch::x86_64::*;
7903    unsafe {
7904        const TILE: usize = cortiq_core::quant::Q1T_TILE;
7905        let mut acc = 0f32;
7906        let bytes_ptr = bytes.as_ptr();
7907        let xq_ptr = xq.as_ptr();
7908        let row_off = r * gpr * TILE;
7909
7910        let ones = _mm256_set1_epi16(1);
7911        for gi in 0..gpr {
7912            let off = row_off + gi * TILE;
7913            let s = f16_to_f32(u16::from_le_bytes([
7914                *bytes_ptr.add(off),
7915                *bytes_ptr.add(off + 1),
7916            ]));
7917            let (u0, u1, u2, u3) = q1t_unpack_reg_u64s(bytes_ptr.add(off + 2));
7918            let wv = _mm256_set_epi64x(u3 as i64, u2 as i64, u1 as i64, u0 as i64);
7919            let xv = _mm256_loadu_si256(xq_ptr.add(gi * GROUP_SIZE) as *const __m256i);
7920            let p16 = _mm256_maddubs_epi16(_mm256_abs_epi8(wv), _mm256_sign_epi8(xv, wv));
7921            let d256 = _mm256_madd_epi16(p16, ones);
7922            let d128 = _mm_add_epi32(
7923                _mm256_castsi256_si128(d256),
7924                _mm256_extracti128_si256(d256, 1),
7925            );
7926            let d64 = _mm_add_epi32(d128, _mm_shuffle_epi32(d128, 0xee));
7927            let d32 = _mm_cvtsi128_si32(_mm_add_epi32(d64, _mm_shuffle_epi32(d64, 0x55)));
7928            acc += d32 as f32 * s;
7929        }
7930        acc
7931    }
7932}
7933
7934/// VNNI twin of `q1t_dot_row_avx2` (see `dpbusd_hsum`).
7935#[cfg(target_arch = "x86_64")]
7936#[target_feature(enable = "avx2,avx512f,avx512bw,avx512vl,avx512vnni")]
7937unsafe fn q1t_dot_row_vnni(bytes: &[u8], r: usize, gpr: usize, xq: &[i8]) -> f32 {
7938    use core::arch::x86_64::*;
7939    // SAFETY: same tile/xq contracts as `q1t_dot_row_avx2`.
7940    unsafe {
7941        const TILE: usize = cortiq_core::quant::Q1T_TILE;
7942        let mut acc = 0f32;
7943        let bytes_ptr = bytes.as_ptr();
7944        let xq_ptr = xq.as_ptr();
7945        let row_off = r * gpr * TILE;
7946        for gi in 0..gpr {
7947            let off = row_off + gi * TILE;
7948            let s = f16_to_f32(u16::from_le_bytes([
7949                *bytes_ptr.add(off),
7950                *bytes_ptr.add(off + 1),
7951            ]));
7952            let (u0, u1, u2, u3) = q1t_unpack_reg_u64s(bytes_ptr.add(off + 2));
7953            let wv = _mm256_set_epi64x(u3 as i64, u2 as i64, u1 as i64, u0 as i64);
7954            let xv = _mm256_loadu_si256(xq_ptr.add(gi * GROUP_SIZE) as *const __m256i);
7955            let d = dpbusd_hsum(_mm256_abs_epi8(wv), _mm256_sign_epi8(xv, wv));
7956            acc += d as f32 * s;
7957        }
7958        acc
7959    }
7960}
7961
7962/// Per-row int8 base dot, dispatched once per row (matvec decode hot path).
7963/// Callers are gated by `a8w8_enabled()`, so the target-feature kernels are
7964/// reachable.
7965#[inline]
7966fn q1t_dot_row_i8(bytes: &[u8], r: usize, gpr: usize, xq: &[i8]) -> f32 {
7967    #[cfg(target_arch = "aarch64")]
7968    unsafe {
7969        return q1t_dot_row_sdot(bytes, r, gpr, xq);
7970    }
7971    #[cfg(target_arch = "x86_64")]
7972    unsafe {
7973        if vnni_tiles_enabled() {
7974            return q1t_dot_row_vnni(bytes, r, gpr, xq);
7975        }
7976        return q1t_dot_row_avx2(bytes, r, gpr, xq);
7977    }
7978    #[allow(unreachable_code)]
7979    {
7980        const TILE: usize = cortiq_core::quant::Q1T_TILE;
7981        let mut acc = 0f32;
7982        let mut sg = [0i8; GROUP_SIZE + 8]; // +8 slack for the u64-store unpack
7983        for gi in 0..gpr {
7984            let off = (r * gpr + gi) * TILE;
7985            let s = f16_to_f32(u16::from_le_bytes([bytes[off], bytes[off + 1]]));
7986            q1t_unpack_group_i8(bytes.as_ptr().wrapping_add(off + 2), &mut sg);
7987            let mut d = 0i32;
7988            for k in 0..GROUP_SIZE {
7989                d += sg[k] as i32 * xq[gi * GROUP_SIZE + k] as i32;
7990            }
7991            acc += d as f32 * s;
7992        }
7993        acc
7994    }
7995}
7996
7997/// Σ over a row's outliers of `value·x[col]` — the correction that adds the
7998/// overlay's exact weights on top of the base dot. INVARIANT: the encoder
7999/// writes ternary code 0 at every outlier position (`quantize_q1t`), so the
8000/// base contributes nothing there and this is a plain `value·x`, not
8001/// `(value − base)·x` — no scattered per-outlier scale read. Row `r`'s entries
8002/// are the contiguous slice `[row_ptr[r], row_ptr[r+1])`, so no binary search.
8003fn q1t_row_outlier_correction(
8004    bytes: &[u8],
8005    r: usize,
8006    rp_off: usize,
8007    entries_off: usize,
8008    has_ov: bool,
8009    x: &[f32],
8010) -> f32 {
8011    if !has_ov {
8012        return 0.0;
8013    }
8014    let (c0, c1) = (
8015        q1t_rowptr(bytes, rp_off, r),
8016        q1t_rowptr(bytes, rp_off, r + 1),
8017    );
8018    let mut corr = 0f32;
8019    for p in c0..c1 {
8020        let e = entries_off + p * 4;
8021        let col = u16::from_le_bytes([bytes[e], bytes[e + 1]]) as usize;
8022        let val = f16_to_f32(u16::from_le_bytes([bytes[e + 2], bytes[e + 3]]));
8023        corr += val * x[col];
8024    }
8025    corr
8026}
8027
8028/// Dequantize one q1t row into `buf[..cols]` via the sign LUT (no division),
8029/// then apply the row's outliers (its `[row_ptr[r], row_ptr[r+1])` slice).
8030/// Used by the batched (prefill) path where the decode amortizes over the batch.
8031fn q1t_dequant_row(
8032    bytes: &[u8],
8033    r: usize,
8034    gpr: usize,
8035    rp_off: usize,
8036    entries_off: usize,
8037    has_ov: bool,
8038    buf: &mut [f32],
8039) {
8040    const TILE: usize = cortiq_core::quant::Q1T_TILE;
8041    for g in 0..gpr {
8042        let off = (r * gpr + g) * TILE;
8043        let s = f16_to_f32(u16::from_le_bytes([bytes[off], bytes[off + 1]]));
8044        let codes = &bytes[off + 2..off + TILE];
8045        let bc = g * GROUP_SIZE;
8046        // 6 full bytes (30 codes) + a 7th byte holding the last 2.
8047        for bi in 0..6 {
8048            let lut = &SIGN5[codes[bi] as usize];
8049            let d = &mut buf[bc + bi * 5..bc + bi * 5 + 5];
8050            for i in 0..5 {
8051                d[i] = lut[i] * s;
8052            }
8053        }
8054        let lut = &SIGN5[codes[6] as usize];
8055        buf[bc + 30] = lut[0] * s;
8056        buf[bc + 31] = lut[1] * s;
8057    }
8058    if !has_ov {
8059        return;
8060    }
8061    let (c0, c1) = (
8062        q1t_rowptr(bytes, rp_off, r),
8063        q1t_rowptr(bytes, rp_off, r + 1),
8064    );
8065    for p in c0..c1 {
8066        let e = entries_off + p * 4;
8067        let col = u16::from_le_bytes([bytes[e], bytes[e + 1]]) as usize;
8068        buf[col] = f16_to_f32(u16::from_le_bytes([bytes[e + 2], bytes[e + 3]]));
8069    }
8070}
8071
8072/// Add the sparse outlier overlay onto a base dot already in `out` (the GPU
8073/// computes the ternary base; the overlay stays on the CPU — its entries are
8074/// few and its per-row gather doesn't vectorize on the GPU). Row-parallel.
8075fn q1t_add_overlay(
8076    bytes: &[u8],
8077    x: &[f32],
8078    rows: usize,
8079    cols: usize,
8080    out: &mut [f32],
8081    pool: Option<&Pool>,
8082) {
8083    const TILE: usize = cortiq_core::quant::Q1T_TILE;
8084    let gpr = cols / GROUP_SIZE;
8085    let (rp_off, ent_off, has_ov) = q1t_overlay(bytes, rows * gpr * TILE, rows);
8086    if !has_ov {
8087        return;
8088    }
8089    let out_addr = SendMut(out.as_mut_ptr());
8090    let run = move |start: usize, end: usize| {
8091        for r in start..end {
8092            let corr = q1t_row_outlier_correction(bytes, r, rp_off, ent_off, has_ov, x);
8093            // SAFETY: disjoint rows; add onto the base the GPU already wrote.
8094            unsafe { *out_addr.at(r) += corr };
8095        }
8096    };
8097    dispatch_rows(pool, rows, &run);
8098}
8099
8100/// Q1T row range via the A8W8 int8 path — shared activation split,
8101/// per-row: base SDOT dot + outlier correction + overlay.
8102#[allow(clippy::too_many_arguments)]
8103fn q1t_range_a8w8(
8104    bytes: &[u8],
8105    gpr: usize,
8106    rp_off: usize,
8107    ent_off: usize,
8108    has_ov: bool,
8109    act: &SplitAct,
8110    x: &[f32],
8111    out: SendMut,
8112    start: usize,
8113    end: usize,
8114) {
8115    for r in start..end {
8116        let mut acc = q1t_dot_row_i8(bytes, r, gpr, &act.xq) * act.sx;
8117        for &(j, xv) in &act.outliers {
8118            acc += q1t_base_weight(bytes, r, gpr, j) * xv;
8119        }
8120        acc += q1t_row_outlier_correction(bytes, r, rp_off, ent_off, has_ov, x);
8121        // SAFETY: disjoint row ranges per worker.
8122        unsafe { *out.at(r) = acc };
8123    }
8124}
8125
8126/// Q1T row range via the f32 path (no SDOT) — for matvec_many batched
8127/// dispatch when a8w8 is unavailable.
8128#[allow(clippy::too_many_arguments)]
8129fn q1t_range_f32_batch(
8130    bytes: &[u8],
8131    gpr: usize,
8132    rp_off: usize,
8133    ent_off: usize,
8134    has_ov: bool,
8135    x: &[f32],
8136    out: SendMut,
8137    start: usize,
8138    end: usize,
8139) {
8140    const TILE: usize = cortiq_core::quant::Q1T_TILE;
8141    let mut sg = [0f32; GROUP_SIZE];
8142    for r in start..end {
8143        let mut acc = 0f32;
8144        for g in 0..gpr {
8145            let off = (r * gpr + g) * TILE;
8146            let s = f16_to_f32(u16::from_le_bytes([bytes[off], bytes[off + 1]]));
8147            let codes = &bytes[off + 2..off + TILE];
8148            let xg = &x[g * GROUP_SIZE..g * GROUP_SIZE + GROUP_SIZE];
8149            for bi in 0..6 {
8150                sg[bi * 5..bi * 5 + 5].copy_from_slice(&SIGN5[codes[bi] as usize]);
8151            }
8152            let lut = &SIGN5[codes[6] as usize];
8153            sg[30] = lut[0];
8154            sg[31] = lut[1];
8155            let mut gsum = 0f32;
8156            for k in 0..GROUP_SIZE {
8157                gsum += sg[k] * xg[k];
8158            }
8159            acc += s * gsum;
8160        }
8161        acc += q1t_row_outlier_correction(bytes, r, rp_off, ent_off, has_ov, x);
8162        // SAFETY: disjoint row ranges per worker.
8163        unsafe { *out.at(r) = acc };
8164    }
8165}
8166
8167/// Ternary (q1t) matvec — decode+dot straight from mmap, one group at a time:
8168/// no per-ROW buffer, no division (the sign LUT), and a tiny per-group sign
8169/// buffer so the 32-wide dot vectorizes. This is the decode hot path.
8170fn q1t_matvec(
8171    bytes: &[u8],
8172    x: &[f32],
8173    rows: usize,
8174    cols: usize,
8175    out: &mut [f32],
8176    pool: Option<&Pool>,
8177) {
8178    debug_assert_eq!(out.len(), rows);
8179    const TILE: usize = cortiq_core::quant::Q1T_TILE;
8180    let gpr = cols / GROUP_SIZE;
8181    let (rp_off, ent_off, has_ov) = q1t_overlay(bytes, rows * gpr * TILE, rows);
8182    let out_addr = SendMut(out.as_mut_ptr());
8183    // int8 SDOT base dot (ARM dotprod): ~4× the f32 arithmetic. x → i8 once
8184    // (`split_act`), activation outliers added back exactly in f32, weight
8185    // overlay on top. ARM SDOT / x86 AVX2; CMF_SDOT=0 keeps the exact f32 path.
8186    if a8w8_enabled() {
8187        let act = split_act(x);
8188        let act = &act;
8189        let run = move |start: usize, end: usize| {
8190            for r in start..end {
8191                let mut acc = q1t_dot_row_i8(bytes, r, gpr, &act.xq) * act.sx;
8192                for &(j, xv) in &act.outliers {
8193                    acc += q1t_base_weight(bytes, r, gpr, j) * xv;
8194                }
8195                acc += q1t_row_outlier_correction(bytes, r, rp_off, ent_off, has_ov, x);
8196                // SAFETY: disjoint row ranges per worker.
8197                unsafe { *out_addr.at(r) = acc };
8198            }
8199        };
8200        dispatch_rows(pool, rows, &run);
8201        return;
8202    }
8203    let run = move |start: usize, end: usize| {
8204        // Per-group signs, unpacked contiguously so the dot below is a clean
8205        // 32-wide reduction the autovectorizer turns into f32x4 FMAs — the
8206        // 5-values-per-byte base-3 layout won't SIMD in place.
8207        let mut sg = [0f32; GROUP_SIZE];
8208        for r in start..end {
8209            let mut acc = 0f32;
8210            for g in 0..gpr {
8211                let off = (r * gpr + g) * TILE;
8212                let s = f16_to_f32(u16::from_le_bytes([bytes[off], bytes[off + 1]]));
8213                let codes = &bytes[off + 2..off + TILE];
8214                let xg = &x[g * GROUP_SIZE..g * GROUP_SIZE + GROUP_SIZE];
8215                for bi in 0..6 {
8216                    sg[bi * 5..bi * 5 + 5].copy_from_slice(&SIGN5[codes[bi] as usize]);
8217                }
8218                let lut = &SIGN5[codes[6] as usize];
8219                sg[30] = lut[0];
8220                sg[31] = lut[1];
8221                let mut gsum = 0f32;
8222                for k in 0..GROUP_SIZE {
8223                    gsum += sg[k] * xg[k];
8224                }
8225                acc += s * gsum;
8226            }
8227            acc += q1t_row_outlier_correction(bytes, r, rp_off, ent_off, has_ov, x);
8228            unsafe { *out_addr.at(r) = acc };
8229        }
8230    };
8231    dispatch_rows(pool, rows, &run);
8232}
8233
8234/// Fused-pair twin of `q1t_dot_row_sdot`: ONE register unpack of the
8235/// ternary codes serves BOTH activation streams (the unpack chain is
8236/// the dominant per-row cost — MTP verify pairs paid it twice). Per
8237/// stream the group order and f32 accumulation match the single-row
8238/// kernel exactly, so pair == 2×matvec bit-for-bit.
8239#[cfg(target_arch = "aarch64")]
8240#[target_feature(enable = "neon,dotprod")]
8241unsafe fn q1t_dot_row_sdot2(bytes: &[u8], r: usize, gpr: usize, xa: &[i8], xb: &[i8]) -> [f32; 2] {
8242    use core::arch::aarch64::*;
8243    use core::arch::asm;
8244    // SAFETY: same slice-length contracts as `q1t_dot_row_sdot`, ×2.
8245    unsafe {
8246        const TILE: usize = cortiq_core::quant::Q1T_TILE;
8247        let bytes_ptr = bytes.as_ptr();
8248        let row_off = r * gpr * TILE;
8249        let xp = [xa.as_ptr(), xb.as_ptr()];
8250        let mut acc = [0f32; 2];
8251        macro_rules! sdot2 {
8252            ($w0:expr, $w1:expr, $x:expr) => {{
8253                let x0 = vld1q_s8($x);
8254                let x1 = vld1q_s8($x.add(16));
8255                let (mut a0, mut a1) = (vdupq_n_s32(0), vdupq_n_s32(0));
8256                asm!(
8257                    "sdot {a0:v}.4s, {w0:v}.16b, {x0:v}.16b",
8258                    "sdot {a1:v}.4s, {w1:v}.16b, {x1:v}.16b",
8259                    a0 = inout(vreg) a0, a1 = inout(vreg) a1,
8260                    w0 = in(vreg) $w0, x0 = in(vreg) x0, w1 = in(vreg) $w1, x1 = in(vreg) x1,
8261                    options(pure, nomem, nostack),
8262                );
8263                vaddvq_s32(vaddq_s32(a0, a1))
8264            }};
8265        }
8266        let gpr2 = gpr & !1;
8267        let mut gi = 0;
8268        while gi < gpr2 {
8269            let off0 = row_off + gi * TILE;
8270            let off1 = off0 + TILE;
8271            let s0 = f16_to_f32(u16::from_le_bytes([
8272                *bytes_ptr.add(off0),
8273                *bytes_ptr.add(off0 + 1),
8274            ]));
8275            let s1 = f16_to_f32(u16::from_le_bytes([
8276                *bytes_ptr.add(off1),
8277                *bytes_ptr.add(off1 + 1),
8278            ]));
8279            let (u0_0, u1_0, u2_0, u3_0) = q1t_unpack_reg_u64s(bytes_ptr.add(off0 + 2));
8280            let (u0_1, u1_1, u2_1, u3_1) = q1t_unpack_reg_u64s(bytes_ptr.add(off1 + 2));
8281            let w0_0 = vreinterpretq_s8_u64(vcombine_u64(vcreate_u64(u0_0), vcreate_u64(u1_0)));
8282            let w1_0 = vreinterpretq_s8_u64(vcombine_u64(vcreate_u64(u2_0), vcreate_u64(u3_0)));
8283            let w0_1 = vreinterpretq_s8_u64(vcombine_u64(vcreate_u64(u0_1), vcreate_u64(u1_1)));
8284            let w1_1 = vreinterpretq_s8_u64(vcombine_u64(vcreate_u64(u2_1), vcreate_u64(u3_1)));
8285            for k in 0..2 {
8286                let d0 = sdot2!(w0_0, w1_0, xp[k].add(gi * GROUP_SIZE));
8287                let d1 = sdot2!(w0_1, w1_1, xp[k].add((gi + 1) * GROUP_SIZE));
8288                acc[k] += d0 as f32 * s0 + d1 as f32 * s1;
8289            }
8290            gi += 2;
8291        }
8292        if gi < gpr {
8293            let off = row_off + gi * TILE;
8294            let s = f16_to_f32(u16::from_le_bytes([
8295                *bytes_ptr.add(off),
8296                *bytes_ptr.add(off + 1),
8297            ]));
8298            let (u0, u1, u2, u3) = q1t_unpack_reg_u64s(bytes_ptr.add(off + 2));
8299            let w0 = vreinterpretq_s8_u64(vcombine_u64(vcreate_u64(u0), vcreate_u64(u1)));
8300            let w1 = vreinterpretq_s8_u64(vcombine_u64(vcreate_u64(u2), vcreate_u64(u3)));
8301            for k in 0..2 {
8302                let d = sdot2!(w0, w1, xp[k].add(gi * GROUP_SIZE));
8303                acc[k] += d as f32 * s;
8304            }
8305        }
8306        acc
8307    }
8308}
8309
8310/// Fused Q1T pair matvec: ONE pass over the rows serves both
8311/// activation streams — on ARM the ternary register unpack happens
8312/// once per tile pair (`q1t_dot_row_sdot2`); elsewhere the second dot
8313/// rides the row's L1-warm tile bytes. Per stream the math matches
8314/// `q1t_matvec` exactly.
8315fn q1t_matvec2(
8316    bytes: &[u8],
8317    x1: &[f32],
8318    x2: &[f32],
8319    rows: usize,
8320    cols: usize,
8321    o1: &mut [f32],
8322    o2: &mut [f32],
8323    pool: Option<&Pool>,
8324) {
8325    debug_assert_eq!(o1.len(), rows);
8326    debug_assert_eq!(o2.len(), rows);
8327    const TILE: usize = cortiq_core::quant::Q1T_TILE;
8328    let gpr = cols / GROUP_SIZE;
8329    let (rp_off, ent_off, has_ov) = q1t_overlay(bytes, rows * gpr * TILE, rows);
8330    let out1 = SendMut(o1.as_mut_ptr());
8331    let out2 = SendMut(o2.as_mut_ptr());
8332    if a8w8_enabled() {
8333        let a1 = split_act(x1);
8334        let a2 = split_act(x2);
8335        let (a1, a2) = (&a1, &a2);
8336        let run = move |start: usize, end: usize| {
8337            for r in start..end {
8338                #[cfg(target_arch = "aarch64")]
8339                // a8w8 on aarch64 ⇔ sdot_enabled(), so the kernel's
8340                // target features are present.
8341                let ds = unsafe { q1t_dot_row_sdot2(bytes, r, gpr, &a1.xq, &a2.xq) };
8342                #[cfg(not(target_arch = "aarch64"))]
8343                let ds = [
8344                    q1t_dot_row_i8(bytes, r, gpr, &a1.xq),
8345                    q1t_dot_row_i8(bytes, r, gpr, &a2.xq),
8346                ];
8347                let mut acc1 = ds[0] * a1.sx;
8348                for &(j, xv) in &a1.outliers {
8349                    acc1 += q1t_base_weight(bytes, r, gpr, j) * xv;
8350                }
8351                acc1 += q1t_row_outlier_correction(bytes, r, rp_off, ent_off, has_ov, x1);
8352                let mut acc2 = ds[1] * a2.sx;
8353                for &(j, xv) in &a2.outliers {
8354                    acc2 += q1t_base_weight(bytes, r, gpr, j) * xv;
8355                }
8356                acc2 += q1t_row_outlier_correction(bytes, r, rp_off, ent_off, has_ov, x2);
8357                // SAFETY: disjoint row ranges per worker.
8358                unsafe {
8359                    *out1.at(r) = acc1;
8360                    *out2.at(r) = acc2;
8361                }
8362            }
8363        };
8364        dispatch_rows(pool, rows, &run);
8365        return;
8366    }
8367    let run = move |start: usize, end: usize| {
8368        // Exact path (CMF_SDOT=0): unpack the sign LUT once per group,
8369        // dot both streams — same op order per stream as `q1t_matvec`.
8370        let mut sg = [0f32; GROUP_SIZE];
8371        for r in start..end {
8372            let mut acc1 = 0f32;
8373            let mut acc2 = 0f32;
8374            for g in 0..gpr {
8375                let off = (r * gpr + g) * TILE;
8376                let s = f16_to_f32(u16::from_le_bytes([bytes[off], bytes[off + 1]]));
8377                let codes = &bytes[off + 2..off + TILE];
8378                for bi in 0..6 {
8379                    sg[bi * 5..bi * 5 + 5].copy_from_slice(&SIGN5[codes[bi] as usize]);
8380                }
8381                let lut = &SIGN5[codes[6] as usize];
8382                sg[30] = lut[0];
8383                sg[31] = lut[1];
8384                let xg1 = &x1[g * GROUP_SIZE..g * GROUP_SIZE + GROUP_SIZE];
8385                let xg2 = &x2[g * GROUP_SIZE..g * GROUP_SIZE + GROUP_SIZE];
8386                let mut gsum1 = 0f32;
8387                for k in 0..GROUP_SIZE {
8388                    gsum1 += sg[k] * xg1[k];
8389                }
8390                acc1 += s * gsum1;
8391                let mut gsum2 = 0f32;
8392                for k in 0..GROUP_SIZE {
8393                    gsum2 += sg[k] * xg2[k];
8394                }
8395                acc2 += s * gsum2;
8396            }
8397            acc1 += q1t_row_outlier_correction(bytes, r, rp_off, ent_off, has_ov, x1);
8398            acc2 += q1t_row_outlier_correction(bytes, r, rp_off, ent_off, has_ov, x2);
8399            // SAFETY: disjoint row ranges per worker.
8400            unsafe {
8401                *out1.at(r) = acc1;
8402                *out2.at(r) = acc2;
8403            }
8404        }
8405    };
8406    dispatch_rows(pool, rows, &run);
8407}
8408
8409/// Ternary (q1t) matmat (prefill) — dequant each row once, dot the whole
8410/// batch against it (amortizes the per-row decode).
8411fn q1t_matmat(
8412    bytes: &[u8],
8413    xs: &[f32],
8414    b: usize,
8415    rows: usize,
8416    cols: usize,
8417    out: &mut [f32],
8418    pool: Option<&Pool>,
8419) {
8420    debug_assert_eq!(out.len(), b * rows);
8421    const TILE: usize = cortiq_core::quant::Q1T_TILE;
8422    let gpr = cols / GROUP_SIZE;
8423    let (rp_off, ent_off, has_ov) = q1t_overlay(bytes, rows * gpr * TILE, rows);
8424    let out_addr = SendMut(out.as_mut_ptr());
8425    // int8 prefill (ARM SDOT / x86 AVX2): quantize the B inputs once, unpack
8426    // each weight row's signs to i8 ONCE, then int8-dot against every input —
8427    // the row sign-decode amortizes over the whole batch. CMF_SDOT=0 → f32.
8428    if a8w8_enabled() {
8429        let acts: Vec<SplitAct> = (0..b)
8430            .map(|bi| split_act(&xs[bi * cols..(bi + 1) * cols]))
8431            .collect();
8432        let acts = &acts;
8433        let run = move |start: usize, end: usize| {
8434            let mut sg = vec![0i8; cols + 8]; // row signs, i8 (+8 unpack slack)
8435            let mut sc = vec![0f32; gpr]; // per-group scales
8436            let mut accs = vec![0f32; b]; // per-batch accumulators, reused per row
8437            for r in start..end {
8438                for g in 0..gpr {
8439                    let off = (r * gpr + g) * TILE;
8440                    sc[g] = f16_to_f32(u16::from_le_bytes([bytes[off], bytes[off + 1]]));
8441                    q1t_unpack_group_i8(
8442                        bytes.as_ptr().wrapping_add(off + 2),
8443                        &mut sg[g * GROUP_SIZE..],
8444                    );
8445                }
8446                for bi in 0..b {
8447                    let act = &acts[bi];
8448                    let mut isum = 0f32;
8449                    for g in 0..gpr {
8450                        let d = q1t_i8dot32(
8451                            sg.as_ptr().wrapping_add(g * GROUP_SIZE),
8452                            act.xq.as_ptr().wrapping_add(g * GROUP_SIZE),
8453                        );
8454                        isum += d as f32 * sc[g];
8455                    }
8456                    let mut acc = isum * act.sx;
8457                    for &(j, xv) in &act.outliers {
8458                        acc += q1t_base_weight(bytes, r, gpr, j) * xv;
8459                    }
8460                    accs[bi] = acc;
8461                }
8462                // Overlay ONCE per row for the whole batch: read each (col, val)
8463                // from mmap a single time (was b× — the re-read dominated prefill)
8464                // and fan it out over the batch via the cached inputs.
8465                if has_ov {
8466                    let (c0, c1) = (
8467                        q1t_rowptr(bytes, rp_off, r),
8468                        q1t_rowptr(bytes, rp_off, r + 1),
8469                    );
8470                    for p in c0..c1 {
8471                        let e = ent_off + p * 4;
8472                        let col = u16::from_le_bytes([bytes[e], bytes[e + 1]]) as usize;
8473                        let val = f16_to_f32(u16::from_le_bytes([bytes[e + 2], bytes[e + 3]]));
8474                        for bi in 0..b {
8475                            accs[bi] += val * xs[bi * cols + col];
8476                        }
8477                    }
8478                }
8479                for bi in 0..b {
8480                    unsafe { *out_addr.at(bi * rows + r) = accs[bi] };
8481                }
8482            }
8483        };
8484        dispatch_rows(pool, rows, &run);
8485        return;
8486    }
8487    let run = move |start: usize, end: usize| {
8488        let mut buf = vec![0f32; cols];
8489        for r in start..end {
8490            q1t_dequant_row(bytes, r, gpr, rp_off, ent_off, has_ov, &mut buf);
8491            for bi in 0..b {
8492                let xr = &xs[bi * cols..(bi + 1) * cols];
8493                let mut acc = 0f32;
8494                for j in 0..cols {
8495                    acc += buf[j] * xr[j];
8496                }
8497                unsafe { *out_addr.at(bi * rows + r) = acc };
8498            }
8499        }
8500    };
8501    dispatch_rows(pool, rows, &run);
8502}
8503
8504fn q1_matvec(
8505    bytes: &[u8],
8506    x: &[f32],
8507    rows: usize,
8508    cols: usize,
8509    out: &mut [f32],
8510    pool: Option<&Pool>,
8511) {
8512    debug_assert_eq!(out.len(), rows);
8513    let gpr = cols / GROUP_SIZE;
8514    let out_addr = SendMut(out.as_mut_ptr());
8515    if a8w8_enabled() {
8516        let act = split_act(x);
8517        let gsum = q1_group_sums(&act.xq, gpr);
8518        let (act, gsum) = (&act, &gsum);
8519        let run = move |start: usize, end: usize| {
8520            q1_range_a8w8(bytes, gpr, act, gsum, out_addr, start, end)
8521        };
8522        dispatch_rows(pool, rows, &run);
8523        return;
8524    }
8525    let run = move |start: usize, end: usize| q1_range_f32(bytes, gpr, x, out_addr, start, end);
8526    dispatch_rows(pool, rows, &run);
8527}
8528
8529/// Fused two-input q1 matvec (weights read once per pair).
8530#[allow(clippy::too_many_arguments)]
8531fn q1_matvec2(
8532    bytes: &[u8],
8533    x1: &[f32],
8534    x2: &[f32],
8535    rows: usize,
8536    cols: usize,
8537    o1: &mut [f32],
8538    o2: &mut [f32],
8539    pool: Option<&Pool>,
8540) {
8541    let gpr = cols / GROUP_SIZE;
8542    let p1 = SendMut(o1.as_mut_ptr());
8543    let p2 = SendMut(o2.as_mut_ptr());
8544    if a8w8_enabled() {
8545        let a1 = split_act(x1);
8546        let a2 = split_act(x2);
8547        let g1 = q1_group_sums(&a1.xq, gpr);
8548        let g2 = q1_group_sums(&a2.xq, gpr);
8549        let (a1, a2, g1, g2) = (&a1, &a2, &g1, &g2);
8550        let run = move |start: usize, end: usize| {
8551            for r in start..end {
8552                let mut v1 = dot_q1_row_i8(bytes, r, gpr, &a1.xq, g1) * a1.sx;
8553                let mut v2 = dot_q1_row_i8(bytes, r, gpr, &a2.xq, g2) * a2.sx;
8554                for &(j, xv) in &a1.outliers {
8555                    let (w, s) = q1_outlier(bytes, r, gpr, j);
8556                    v1 += w * s * xv;
8557                }
8558                for &(j, xv) in &a2.outliers {
8559                    let (w, s) = q1_outlier(bytes, r, gpr, j);
8560                    v2 += w * s * xv;
8561                }
8562                // SAFETY: disjoint row ranges per worker.
8563                unsafe {
8564                    *p1.at(r) = v1;
8565                    *p2.at(r) = v2;
8566                }
8567            }
8568        };
8569        dispatch_rows(pool, rows, &run);
8570        return;
8571    }
8572    let run = move |start: usize, end: usize| {
8573        for r in start..end {
8574            // SAFETY: disjoint row ranges per worker.
8575            unsafe {
8576                *p1.at(r) = q1_row_exact(bytes, r, gpr, x1);
8577                *p2.at(r) = q1_row_exact(bytes, r, gpr, x2);
8578            }
8579        }
8580    };
8581    dispatch_rows(pool, rows, &run);
8582}
8583
8584/// Batched q1 matmat: each row's tiles stream once per microbatch.
8585#[allow(clippy::too_many_arguments)]
8586fn q1_matmat(
8587    bytes: &[u8],
8588    xs_all: &[f32],
8589    b: usize,
8590    rows: usize,
8591    cols: usize,
8592    out: &mut [f32],
8593    pool: Option<&Pool>,
8594) {
8595    debug_assert_eq!(out.len(), b * rows);
8596    let gpr = cols / GROUP_SIZE;
8597    let out_addr = SendMut(out.as_mut_ptr());
8598    if a8w8_enabled() {
8599        let acts: Vec<(SplitAct, Vec<i32>)> = (0..b)
8600            .map(|bi| {
8601                let act = split_act(&xs_all[bi * cols..(bi + 1) * cols]);
8602                let gsum = q1_group_sums(&act.xq, gpr);
8603                (act, gsum)
8604            })
8605            .collect();
8606        let acts = &acts;
8607        #[cfg(target_arch = "x86_64")]
8608        let blocked_ok = avx2_enabled() && blocked_enabled();
8609        #[cfg(target_arch = "aarch64")]
8610        let blocked_ok = sdot_enabled() && blocked_enabled();
8611        let run = move |start: usize, end: usize| {
8612            for r in start..end {
8613                let mut bi = 0usize;
8614                // Blocked 1×4: the unpacked bit mask serves four
8615                // activation streams per group.
8616                #[cfg(target_arch = "aarch64")]
8617                if blocked_ok {
8618                    while bi + 4 <= acts.len() {
8619                        let xs = [
8620                            acts[bi].0.xq.as_slice(),
8621                            acts[bi + 1].0.xq.as_slice(),
8622                            acts[bi + 2].0.xq.as_slice(),
8623                            acts[bi + 3].0.xq.as_slice(),
8624                        ];
8625                        let gs = [
8626                            acts[bi].1.as_slice(),
8627                            acts[bi + 1].1.as_slice(),
8628                            acts[bi + 2].1.as_slice(),
8629                            acts[bi + 3].1.as_slice(),
8630                        ];
8631                        let d = unsafe { dot_q1_row_1x4_sdot(bytes, r, gpr, xs, gs) };
8632                        for k in 0..4 {
8633                            let (act, _) = &acts[bi + k];
8634                            let mut acc = d[k] * act.sx;
8635                            for &(j, xv) in &act.outliers {
8636                                let (w, sc) = q1_outlier(bytes, r, gpr, j);
8637                                acc += w * sc * xv;
8638                            }
8639                            // SAFETY: disjoint (bi, r) cells per worker.
8640                            unsafe { *out_addr.at((bi + k) * rows + r) = acc };
8641                        }
8642                        bi += 4;
8643                    }
8644                }
8645                #[cfg(target_arch = "x86_64")]
8646                if blocked_ok {
8647                    while bi + 4 <= acts.len() {
8648                        let xs = [
8649                            acts[bi].0.xq.as_slice(),
8650                            acts[bi + 1].0.xq.as_slice(),
8651                            acts[bi + 2].0.xq.as_slice(),
8652                            acts[bi + 3].0.xq.as_slice(),
8653                        ];
8654                        let gs = [
8655                            acts[bi].1.as_slice(),
8656                            acts[bi + 1].1.as_slice(),
8657                            acts[bi + 2].1.as_slice(),
8658                            acts[bi + 3].1.as_slice(),
8659                        ];
8660                        let d = unsafe {
8661                            if vnni_tiles_enabled() {
8662                                dot_q1_row_1x4_vnni(bytes, r, gpr, xs, gs)
8663                            } else {
8664                                dot_q1_row_1x4_avx2(bytes, r, gpr, xs, gs)
8665                            }
8666                        };
8667                        for k in 0..4 {
8668                            let (act, _) = &acts[bi + k];
8669                            let mut acc = d[k] * act.sx;
8670                            for &(j, xv) in &act.outliers {
8671                                let (w, sc) = q1_outlier(bytes, r, gpr, j);
8672                                acc += w * sc * xv;
8673                            }
8674                            // SAFETY: disjoint (bi, r) cells per worker.
8675                            unsafe { *out_addr.at((bi + k) * rows + r) = acc };
8676                        }
8677                        bi += 4;
8678                    }
8679                }
8680                while bi < acts.len() {
8681                    let (act, gsum) = &acts[bi];
8682                    let mut acc = dot_q1_row_i8(bytes, r, gpr, &act.xq, gsum) * act.sx;
8683                    for &(j, xv) in &act.outliers {
8684                        let (w, s) = q1_outlier(bytes, r, gpr, j);
8685                        acc += w * s * xv;
8686                    }
8687                    // SAFETY: disjoint (bi, r) cells per worker range.
8688                    unsafe { *out_addr.at(bi * rows + r) = acc };
8689                    bi += 1;
8690                }
8691            }
8692        };
8693        dispatch_rows(pool, rows, &run);
8694        return;
8695    }
8696    let run = move |start: usize, end: usize| {
8697        for r in start..end {
8698            for bi in 0..b {
8699                let x = &xs_all[bi * cols..(bi + 1) * cols];
8700                // SAFETY: disjoint (bi, r) cells per worker range.
8701                unsafe { *out_addr.at(bi * rows + r) = q1_row_exact(bytes, r, gpr, x) };
8702            }
8703        }
8704    };
8705    dispatch_rows(pool, rows, &run);
8706}
8707
8708/// Fused q4_block matvec straight from the mapped bytes. SDOT path when
8709/// dotprod is available (port of vmfcore `dot_q4_block_sdot`, measured
8710/// +23% on q4 decode): nibbles → centered i8, int8×int8 `sdot` per
8711/// 32-group, exact outlier correction — the same A8W8 contract as q8.
8712/// `CMF_SDOT=0` keeps the exact scalar path.
8713fn q4matvec(
8714    bytes: &[u8],
8715    x: &[f32],
8716    rows: usize,
8717    cols: usize,
8718    out: &mut [f32],
8719    pool: Option<&Pool>,
8720) {
8721    debug_assert_eq!(out.len(), rows);
8722    let (packed, scales) = q4_split(bytes, rows, cols);
8723    let gpr = cols / GROUP_SIZE;
8724    let out_addr = SendMut(out.as_mut_ptr());
8725
8726    if a8w8_enabled() {
8727        let act = split_act(x);
8728        let run = move |start: usize, end: usize| {
8729            q4_range_a8w8(packed, scales, gpr, cols, &act, out_addr, start, end)
8730        };
8731        dispatch_rows(pool, rows, &run);
8732        return;
8733    }
8734
8735    let run =
8736        move |start: usize, end: usize| q4_range_f32(packed, scales, gpr, x, out_addr, start, end);
8737    dispatch_rows(pool, rows, &run);
8738}
8739
8740/// One q4 row via the A8W8 int8 path — SDOT on ARM, AVX2 maddubs on
8741/// x86 (scalar fallback is unreachable: callers gate on a8w8_enabled).
8742#[inline]
8743#[allow(unreachable_code)]
8744/// One UNPACKED q4 row (centered i8 in `buf`) against four activation
8745/// streams: the 32-byte weight chunk and its abs() load once per group,
8746/// the per-group f16 scale decodes once — four maddubs+reduce chains
8747/// instead of four full (load, abs, dot) rounds.
8748#[cfg(target_arch = "x86_64")]
8749#[target_feature(enable = "avx2")]
8750unsafe fn dot_q4b_row_1x4_avx2(
8751    buf: &[u8],
8752    scales: &[u8],
8753    g0: usize,
8754    gpr: usize,
8755    xs: [&[i8]; 4],
8756) -> [f32; 4] {
8757    // SAFETY: callers uphold buffer contracts (buf.len() == gpr·32).
8758    unsafe {
8759        use core::arch::x86_64::*;
8760        let ones = _mm256_set1_epi16(1);
8761        let mut acc = [0f32; 4];
8762        for gi in 0..gpr {
8763            let s = f16_to_f32(u16::from_le_bytes([
8764                scales[(g0 + gi) * 2],
8765                scales[(g0 + gi) * 2 + 1],
8766            ]));
8767            let w = _mm256_loadu_si256(buf.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
8768            let aw = _mm256_abs_epi8(w);
8769            for (k, xq) in xs.iter().enumerate() {
8770                let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
8771                let p16 = _mm256_maddubs_epi16(aw, _mm256_sign_epi8(x, w));
8772                let d = _mm256_madd_epi16(p16, ones);
8773                let hi128 = _mm256_extracti128_si256::<1>(d);
8774                let s128 = _mm_add_epi32(_mm256_castsi256_si128(d), hi128);
8775                let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
8776                let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
8777                acc[k] += _mm_cvtsi128_si32(s32) as f32 * s;
8778            }
8779        }
8780        acc
8781    }
8782}
8783
8784/// VNNI twin of `dot_q4b_row_1x4_avx2` (see `dpbusd_hsum`).
8785#[cfg(target_arch = "x86_64")]
8786#[target_feature(enable = "avx2,avx512f,avx512bw,avx512vl,avx512vnni")]
8787unsafe fn dot_q4b_row_1x4_vnni(
8788    buf: &[u8],
8789    scales: &[u8],
8790    g0: usize,
8791    gpr: usize,
8792    xs: [&[i8]; 4],
8793) -> [f32; 4] {
8794    // SAFETY: callers uphold buffer contracts (buf.len() == gpr·32).
8795    unsafe {
8796        use core::arch::x86_64::*;
8797        let mut acc = [0f32; 4];
8798        for gi in 0..gpr {
8799            let s = f16_to_f32(u16::from_le_bytes([
8800                scales[(g0 + gi) * 2],
8801                scales[(g0 + gi) * 2 + 1],
8802            ]));
8803            let w = _mm256_loadu_si256(buf.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
8804            let aw = _mm256_abs_epi8(w);
8805            for (k, xq) in xs.iter().enumerate() {
8806                let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
8807                let d = dpbusd_hsum(aw, _mm256_sign_epi8(x, w));
8808                acc[k] += d as f32 * s;
8809            }
8810        }
8811        acc
8812    }
8813}
8814
8815/// The vbit flavor of the blocked 1×4: the per-activation A8W8 scale
8816/// folds in PER GROUP as `(d·sx)·s` — bit-matching the single-matvec
8817/// accumulation order (the q4_block flavor applies sx once at the end,
8818/// matching ITS single path; the two conventions are historical and
8819/// each blocked leg must mirror its own).
8820#[cfg(target_arch = "x86_64")]
8821#[target_feature(enable = "avx2")]
8822unsafe fn dot_q4b_row_1x4_sx_avx2(
8823    buf: &[u8],
8824    scales: &[u8],
8825    g0: usize,
8826    gpr: usize,
8827    xs: [&[i8]; 4],
8828    sxs: [f32; 4],
8829) -> [f32; 4] {
8830    // SAFETY: callers uphold buffer contracts (buf.len() == gpr·32).
8831    unsafe {
8832        use core::arch::x86_64::*;
8833        let ones = _mm256_set1_epi16(1);
8834        let mut acc = [0f32; 4];
8835        for gi in 0..gpr {
8836            let s = f16_to_f32(u16::from_le_bytes([
8837                scales[(g0 + gi) * 2],
8838                scales[(g0 + gi) * 2 + 1],
8839            ]));
8840            let w = _mm256_loadu_si256(buf.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
8841            let aw = _mm256_abs_epi8(w);
8842            for (k, xq) in xs.iter().enumerate() {
8843                let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
8844                let p16 = _mm256_maddubs_epi16(aw, _mm256_sign_epi8(x, w));
8845                let d = _mm256_madd_epi16(p16, ones);
8846                let hi128 = _mm256_extracti128_si256::<1>(d);
8847                let s128 = _mm_add_epi32(_mm256_castsi256_si128(d), hi128);
8848                let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
8849                let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
8850                acc[k] += (_mm_cvtsi128_si32(s32) as f32 * sxs[k]) * s;
8851            }
8852        }
8853        acc
8854    }
8855}
8856
8857/// VNNI twin of `dot_q4b_row_1x4_sx_avx2` (see `dpbusd_hsum`; the
8858/// per-group `(d·sx)·s` fold mirrors the vbit single path).
8859#[cfg(target_arch = "x86_64")]
8860#[target_feature(enable = "avx2,avx512f,avx512bw,avx512vl,avx512vnni")]
8861unsafe fn dot_q4b_row_1x4_sx_vnni(
8862    buf: &[u8],
8863    scales: &[u8],
8864    g0: usize,
8865    gpr: usize,
8866    xs: [&[i8]; 4],
8867    sxs: [f32; 4],
8868) -> [f32; 4] {
8869    // SAFETY: callers uphold buffer contracts (buf.len() == gpr·32).
8870    unsafe {
8871        use core::arch::x86_64::*;
8872        let mut acc = [0f32; 4];
8873        for gi in 0..gpr {
8874            let s = f16_to_f32(u16::from_le_bytes([
8875                scales[(g0 + gi) * 2],
8876                scales[(g0 + gi) * 2 + 1],
8877            ]));
8878            let w = _mm256_loadu_si256(buf.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
8879            let aw = _mm256_abs_epi8(w);
8880            for (k, xq) in xs.iter().enumerate() {
8881                let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
8882                let d = dpbusd_hsum(aw, _mm256_sign_epi8(x, w));
8883                acc[k] += (d as f32 * sxs[k]) * s;
8884            }
8885        }
8886        acc
8887    }
8888}
8889
8890#[allow(unreachable_code)]
8891fn dot_q4_row_i8(packed: &[u8], scales: &[u8], g0: usize, gpr: usize, xq: &[i8]) -> f32 {
8892    #[cfg(target_arch = "aarch64")]
8893    unsafe {
8894        return dot_q4_row_sdot(packed, scales, g0, gpr, xq);
8895    }
8896    #[cfg(target_arch = "x86_64")]
8897    unsafe {
8898        return dot_q4_row_avx2(packed, scales, g0, gpr, xq);
8899    }
8900    let mut acc = 0f32;
8901    for gi in 0..gpr {
8902        let g = g0 + gi;
8903        let s = f16_to_f32(u16::from_le_bytes([scales[g * 2], scales[g * 2 + 1]]));
8904        let mut d = 0i32;
8905        for (k, &b) in packed[g * 16..(g + 1) * 16].iter().enumerate() {
8906            d += ((b & 0x0F) as i32 - 8) * xq[gi * GROUP_SIZE + k * 2] as i32
8907                + (((b >> 4) & 0x0F) as i32 - 8) * xq[gi * GROUP_SIZE + k * 2 + 1] as i32;
8908        }
8909        acc += d as f32 * s;
8910    }
8911    acc
8912}
8913
8914/// Two-activation q4 row via the A8W8 int8 path (see `dot_q4_row_i8`).
8915#[inline]
8916#[allow(unreachable_code)]
8917fn dot_q4_row_i8_2(
8918    packed: &[u8],
8919    scales: &[u8],
8920    g0: usize,
8921    gpr: usize,
8922    xq1: &[i8],
8923    xq2: &[i8],
8924) -> (f32, f32) {
8925    #[cfg(target_arch = "aarch64")]
8926    unsafe {
8927        return dot_q4_row_sdot2(packed, scales, g0, gpr, xq1, xq2);
8928    }
8929    #[cfg(target_arch = "x86_64")]
8930    unsafe {
8931        return dot_q4_row_avx2_2(packed, scales, g0, gpr, xq1, xq2);
8932    }
8933    (
8934        dot_q4_row_i8(packed, scales, g0, gpr, xq1),
8935        dot_q4_row_i8(packed, scales, g0, gpr, xq2),
8936    )
8937}
8938
8939/// One q4 row range via SDOT (kernel body of `q4matvec`, extracted so
8940/// multi-matrix jobs can drive it for several tensors in one dispatch).
8941#[allow(clippy::too_many_arguments)]
8942fn q4_range_a8w8(
8943    packed: &[u8],
8944    scales: &[u8],
8945    gpr: usize,
8946    cols: usize,
8947    act: &SplitAct,
8948    out: SendMut,
8949    start: usize,
8950    end: usize,
8951) {
8952    for r in start..end {
8953        let mut acc = dot_q4_row_i8(packed, scales, r * gpr, gpr, &act.xq) * act.sx;
8954        // xq is zeroed at outlier slots — add the exact terms.
8955        for &(j, xv) in &act.outliers {
8956            let flat = r * cols + j;
8957            let byte = packed[flat / 2];
8958            let nib = if flat & 1 == 0 {
8959                byte & 0x0F
8960            } else {
8961                byte >> 4
8962            };
8963            let s = f16_to_f32(u16::from_le_bytes([
8964                scales[(flat / GROUP_SIZE) * 2],
8965                scales[(flat / GROUP_SIZE) * 2 + 1],
8966            ]));
8967            acc += ((nib as i32 - 8) as f32) * s * xv;
8968        }
8969        // SAFETY: disjoint row ranges per worker.
8970        unsafe { *out.at(r) = acc };
8971    }
8972}
8973
8974/// Two-input q4 row range via the A8W8 int8 path — kernel body of
8975/// `q4matvec2`, extracted for pair multi-matrix jobs.
8976#[allow(clippy::too_many_arguments)]
8977fn q4_range2_a8w8(
8978    packed: &[u8],
8979    scales: &[u8],
8980    gpr: usize,
8981    cols: usize,
8982    a1: &SplitAct,
8983    a2: &SplitAct,
8984    p1: SendMut,
8985    p2: SendMut,
8986    start: usize,
8987    end: usize,
8988) {
8989    for r in start..end {
8990        let (s1, s2) = dot_q4_row_i8_2(packed, scales, r * gpr, gpr, &a1.xq, &a2.xq);
8991        let mut acc1 = s1 * a1.sx;
8992        let mut acc2 = s2 * a2.sx;
8993        // xq is zeroed at outlier slots — add the exact terms.
8994        let fix = |outliers: &[(usize, f32)], acc: &mut f32| {
8995            for &(j, xv) in outliers {
8996                let flat = r * cols + j;
8997                let byte = packed[flat / 2];
8998                let nib = if flat & 1 == 0 {
8999                    byte & 0x0F
9000                } else {
9001                    byte >> 4
9002                };
9003                let s = f16_to_f32(u16::from_le_bytes([
9004                    scales[(flat / GROUP_SIZE) * 2],
9005                    scales[(flat / GROUP_SIZE) * 2 + 1],
9006                ]));
9007                *acc += ((nib as i32 - 8) as f32) * s * xv;
9008            }
9009        };
9010        fix(&a1.outliers, &mut acc1);
9011        fix(&a2.outliers, &mut acc2);
9012        // SAFETY: disjoint row ranges per worker.
9013        unsafe {
9014            *p1.at(r) = acc1;
9015            *p2.at(r) = acc2;
9016        }
9017    }
9018}
9019
9020/// Exact scalar q4 row range (same extraction, non-SDOT path).
9021fn q4_range_f32(
9022    packed: &[u8],
9023    scales: &[u8],
9024    gpr: usize,
9025    x: &[f32],
9026    out: SendMut,
9027    start: usize,
9028    end: usize,
9029) {
9030    for r in start..end {
9031        let mut acc = 0f32;
9032        for gi in 0..gpr {
9033            let g = r * gpr + gi;
9034            let s = f16_to_f32(u16::from_le_bytes([scales[g * 2], scales[g * 2 + 1]]));
9035            let pk = &packed[g * 16..(g + 1) * 16];
9036            let xg = &x[gi * GROUP_SIZE..(gi + 1) * GROUP_SIZE];
9037            let mut ga = 0f32;
9038            for (k, &b) in pk.iter().enumerate() {
9039                ga += ((b & 0x0F) as f32 - 8.0) * xg[k * 2]
9040                    + (((b >> 4) & 0x0F) as f32 - 8.0) * xg[k * 2 + 1];
9041            }
9042            acc += ga * s;
9043        }
9044        // SAFETY: disjoint row ranges per worker.
9045        unsafe { *out.at(r) = acc };
9046    }
9047}
9048
9049/// Fused two-input q4 matvec: nibbles are unpacked ONCE per group and
9050/// dotted against both activations (was: two full matvecs — double
9051/// weight traffic). Per-lane math matches `q4matvec` exactly.
9052#[allow(clippy::too_many_arguments)]
9053fn q4matvec2(
9054    bytes: &[u8],
9055    x1: &[f32],
9056    x2: &[f32],
9057    rows: usize,
9058    cols: usize,
9059    o1: &mut [f32],
9060    o2: &mut [f32],
9061    pool: Option<&Pool>,
9062) {
9063    debug_assert_eq!(o1.len(), rows);
9064    debug_assert_eq!(o2.len(), rows);
9065    let (packed, scales) = q4_split(bytes, rows, cols);
9066    let gpr = cols / GROUP_SIZE;
9067
9068    if a8w8_enabled() {
9069        let a1 = split_act(x1);
9070        let a2 = split_act(x2);
9071        let p1 = SendMut(o1.as_mut_ptr());
9072        let p2 = SendMut(o2.as_mut_ptr());
9073        let run = move |start: usize, end: usize| {
9074            q4_range2_a8w8(packed, scales, gpr, cols, &a1, &a2, p1, p2, start, end)
9075        };
9076        dispatch_rows(pool, rows, &run);
9077        return;
9078    }
9079
9080    let p1 = SendMut(o1.as_mut_ptr());
9081    let p2 = SendMut(o2.as_mut_ptr());
9082    let run = move |start: usize, end: usize| {
9083        q4_range2_f32(packed, scales, gpr, x1, x2, p1, p2, start, end)
9084    };
9085    dispatch_rows(pool, rows, &run);
9086}
9087
9088/// Two-input exact scalar q4 row range (same extraction).
9089#[allow(clippy::too_many_arguments)]
9090fn q4_range2_f32(
9091    packed: &[u8],
9092    scales: &[u8],
9093    gpr: usize,
9094    x1: &[f32],
9095    x2: &[f32],
9096    p1: SendMut,
9097    p2: SendMut,
9098    start: usize,
9099    end: usize,
9100) {
9101    for r in start..end {
9102        let (mut acc1, mut acc2) = (0f32, 0f32);
9103        for gi in 0..gpr {
9104            let g = r * gpr + gi;
9105            let s = f16_to_f32(u16::from_le_bytes([scales[g * 2], scales[g * 2 + 1]]));
9106            let pk = &packed[g * 16..(g + 1) * 16];
9107            let x1g = &x1[gi * GROUP_SIZE..(gi + 1) * GROUP_SIZE];
9108            let x2g = &x2[gi * GROUP_SIZE..(gi + 1) * GROUP_SIZE];
9109            let (mut g1, mut g2) = (0f32, 0f32);
9110            for (k, &b) in pk.iter().enumerate() {
9111                let wl = (b & 0x0F) as f32 - 8.0;
9112                let wh = ((b >> 4) & 0x0F) as f32 - 8.0;
9113                g1 += wl * x1g[k * 2] + wh * x1g[k * 2 + 1];
9114                g2 += wl * x2g[k * 2] + wh * x2g[k * 2 + 1];
9115            }
9116            acc1 += g1 * s;
9117            acc2 += g2 * s;
9118        }
9119        // SAFETY: disjoint row ranges per worker.
9120        unsafe {
9121            *p1.at(r) = acc1;
9122            *p2.at(r) = acc2;
9123        }
9124    }
9125}
9126
9127thread_local! {
9128    /// Per-worker decoded-row scratch for the batched q4/vbit kernels
9129    /// (centered i8 for SDOT, f32 for the exact/scalar paths).
9130    static ROW_I8: std::cell::RefCell<Vec<u8>> = const { std::cell::RefCell::new(Vec::new()) };
9131    static ROW_F32: std::cell::RefCell<Vec<f32>> = const { std::cell::RefCell::new(Vec::new()) };
9132}
9133
9134/// Batched q4 matmat: each weight row is unpacked from the mmap ONCE
9135/// and dotted against ALL b activations (prefill used to fall back to b
9136/// full matvecs — b× weight traffic and b× nibble decode). Per-position
9137/// math matches `q4matvec` exactly: same group order, same accumulation.
9138/// `out` is row-major [b, rows] like `qmatmat`.
9139#[allow(clippy::too_many_arguments)]
9140fn q4matmat(
9141    bytes: &[u8],
9142    xs_all: &[f32],
9143    b: usize,
9144    rows: usize,
9145    cols: usize,
9146    out: &mut [f32],
9147    pool: Option<&Pool>,
9148) {
9149    debug_assert_eq!(xs_all.len(), b * cols);
9150    debug_assert_eq!(out.len(), b * rows);
9151    let (packed, scales) = q4_split(bytes, rows, cols);
9152    let gpr = cols / GROUP_SIZE;
9153    let gscale = |g: usize| f16_to_f32(u16::from_le_bytes([scales[g * 2], scales[g * 2 + 1]]));
9154
9155    if a8w8_enabled() {
9156        let acts: Vec<SplitAct> = (0..b)
9157            .map(|bi| split_act(&xs_all[bi * cols..(bi + 1) * cols]))
9158            .collect();
9159        let acts = &acts;
9160        let out_addr = SendMut(out.as_mut_ptr());
9161        let run = move |start: usize, end: usize| {
9162            ROW_I8.with(|rb| {
9163                let mut buf = rb.borrow_mut();
9164                buf.resize(cols, 0);
9165                for r in start..end {
9166                    // Unpack the row's nibbles to centered i8 once
9167                    // (element 2k = low nibble, 2k+1 = high — flat order,
9168                    // same as dot_q4_row_sdot's zip).
9169                    for gi in 0..gpr {
9170                        let g = r * gpr + gi;
9171                        for (k, &bt) in packed[g * 16..(g + 1) * 16].iter().enumerate() {
9172                            buf[gi * GROUP_SIZE + k * 2] = ((bt & 0x0F) as i32 - 8) as i8 as u8;
9173                            buf[gi * GROUP_SIZE + k * 2 + 1] =
9174                                (((bt >> 4) & 0x0F) as i32 - 8) as i8 as u8;
9175                        }
9176                    }
9177                    let mut bi = 0usize;
9178                    #[cfg(target_arch = "x86_64")]
9179                    if avx2_enabled() && blocked_enabled() {
9180                        while bi + 4 <= acts.len() {
9181                            let xs = [
9182                                acts[bi].xq.as_slice(),
9183                                acts[bi + 1].xq.as_slice(),
9184                                acts[bi + 2].xq.as_slice(),
9185                                acts[bi + 3].xq.as_slice(),
9186                            ];
9187                            let d = unsafe {
9188                                if vnni_tiles_enabled() {
9189                                    dot_q4b_row_1x4_vnni(&buf, scales, r * gpr, gpr, xs)
9190                                } else {
9191                                    dot_q4b_row_1x4_avx2(&buf, scales, r * gpr, gpr, xs)
9192                                }
9193                            };
9194                            for k in 0..4 {
9195                                let act = &acts[bi + k];
9196                                let mut acc = d[k] * act.sx;
9197                                for &(j, xv) in &act.outliers {
9198                                    acc += (buf[j] as i8) as f32
9199                                        * gscale((r * cols + j) / GROUP_SIZE)
9200                                        * xv;
9201                                }
9202                                // SAFETY: disjoint (bi, r) cells per worker.
9203                                unsafe { *out_addr.at((bi + k) * rows + r) = acc };
9204                            }
9205                            bi += 4;
9206                        }
9207                    }
9208                    while bi < acts.len() {
9209                        let act = &acts[bi];
9210                        let mut acc = 0f32;
9211                        for gi in 0..gpr {
9212                            let d = dot_i8_i8(
9213                                &buf[gi * GROUP_SIZE..(gi + 1) * GROUP_SIZE],
9214                                &act.xq[gi * GROUP_SIZE..(gi + 1) * GROUP_SIZE],
9215                            );
9216                            acc += d as f32 * gscale(r * gpr + gi);
9217                        }
9218                        acc *= act.sx;
9219                        // xq is zeroed at outlier slots — exact terms.
9220                        for &(j, xv) in &act.outliers {
9221                            acc += (buf[j] as i8) as f32 * gscale((r * cols + j) / GROUP_SIZE) * xv;
9222                        }
9223                        // SAFETY: disjoint (bi, r) cells per worker row range.
9224                        unsafe { *out_addr.at(bi * rows + r) = acc };
9225                        bi += 1;
9226                    }
9227                }
9228            })
9229        };
9230        dispatch_rows(pool, rows, &run);
9231        return;
9232    }
9233
9234    let out_addr = SendMut(out.as_mut_ptr());
9235    let run = move |start: usize, end: usize| {
9236        ROW_F32.with(|rb| {
9237            let mut buf = rb.borrow_mut();
9238            buf.resize(cols, 0.0);
9239            for r in start..end {
9240                // Decode raw (nib − 8) values once; scales stay per-group
9241                // so the accumulation order matches q4matvec bit-for-bit.
9242                for gi in 0..gpr {
9243                    let g = r * gpr + gi;
9244                    for (k, &bt) in packed[g * 16..(g + 1) * 16].iter().enumerate() {
9245                        buf[gi * GROUP_SIZE + k * 2] = (bt & 0x0F) as f32 - 8.0;
9246                        buf[gi * GROUP_SIZE + k * 2 + 1] = ((bt >> 4) & 0x0F) as f32 - 8.0;
9247                    }
9248                }
9249                for bi in 0..b {
9250                    let x = &xs_all[bi * cols..(bi + 1) * cols];
9251                    let mut acc = 0f32;
9252                    for gi in 0..gpr {
9253                        let mut ga = 0f32;
9254                        // Pairwise (lo + hi) addition, matching
9255                        // q4matvec's `ga += lo·x + hi·x` shape exactly —
9256                        // a flat one-per-element loop rounds differently
9257                        // and broke bit-parity on the scalar (x86) path.
9258                        for k in 0..GROUP_SIZE / 2 {
9259                            let e = gi * GROUP_SIZE + k * 2;
9260                            ga += buf[e] * x[e] + buf[e + 1] * x[e + 1];
9261                        }
9262                        acc += ga * gscale(r * gpr + gi);
9263                    }
9264                    // SAFETY: disjoint (bi, r) cells per worker row range.
9265                    unsafe { *out_addr.at(bi * rows + r) = acc };
9266                }
9267            }
9268        })
9269    };
9270    dispatch_rows(pool, rows, &run);
9271}
9272
9273/// Batched vbit matmat: each variable-bit row is decoded from the mmap
9274/// ONCE for the whole microbatch. Same per-position math as
9275/// `vbitmatvec` (SDOT A8W8 with exact outliers / exact f32 for b=8 rows
9276/// and the scalar path).
9277#[allow(clippy::too_many_arguments)]
9278fn vbitmatmat(
9279    bytes: &[u8],
9280    offsets: &[usize],
9281    xs_all: &[f32],
9282    b: usize,
9283    rows: usize,
9284    cols: usize,
9285    out: &mut [f32],
9286    pool: Option<&Pool>,
9287) {
9288    debug_assert_eq!(xs_all.len(), b * cols);
9289    debug_assert_eq!(out.len(), b * rows);
9290    debug_assert_eq!(offsets.len(), rows + 1);
9291    let ng = cols / GROUP_SIZE;
9292    let bits = &bytes[..rows];
9293    let sc_off = rows;
9294    let gscale = |r: usize, g: usize| {
9295        let so = (r * ng + g) * 2;
9296        f16_to_f32(u16::from_le_bytes([
9297            bytes[sc_off + so],
9298            bytes[sc_off + so + 1],
9299        ]))
9300    };
9301
9302    // Decode row r's raw (u − L) values into `dst` (f32, unscaled).
9303    let decode_f32 = |r: usize, dst: &mut [f32]| {
9304        let bw = bits[r] as usize;
9305        let l = ((1i32 << (bw - 1)) - 1) as f32;
9306        let data = &bytes[offsets[r]..offsets[r + 1]];
9307        let (mut acc, mut nbits, mut idx) = (0u64, 0usize, 0usize);
9308        for d in dst.iter_mut() {
9309            while nbits < bw {
9310                acc = (acc << 8) | data[idx] as u64;
9311                idx += 1;
9312                nbits += 8;
9313            }
9314            let u = ((acc >> (nbits - bw)) & ((1u64 << bw) - 1)) as f32;
9315            nbits -= bw;
9316            *d = u - l;
9317        }
9318    };
9319
9320    if a8w8_enabled() {
9321        let acts: Vec<SplitAct> = (0..b)
9322            .map(|bi| split_act(&xs_all[bi * cols..(bi + 1) * cols]))
9323            .collect();
9324        let acts = &acts;
9325        let out_addr = SendMut(out.as_mut_ptr());
9326        let run = move |start: usize, end: usize| {
9327            for r in start..end {
9328                let bw = bits[r] as usize;
9329                if bw == 8 {
9330                    // u−L reaches 128 → no i8 path; decode once, exact
9331                    // f32 dots for every position (same as vbitmatvec).
9332                    ROW_F32.with(|rb| {
9333                        let mut buf = rb.borrow_mut();
9334                        buf.resize(cols, 0.0);
9335                        decode_f32(r, &mut buf);
9336                        for bi in 0..b {
9337                            let x = &xs_all[bi * cols..(bi + 1) * cols];
9338                            let mut dot = 0f32;
9339                            for g in 0..ng {
9340                                let mut gd = 0f32;
9341                                for k in 0..GROUP_SIZE {
9342                                    gd += buf[g * GROUP_SIZE + k] * x[g * GROUP_SIZE + k];
9343                                }
9344                                dot += gd * gscale(r, g);
9345                            }
9346                            // SAFETY: disjoint (bi, r) cells per worker range.
9347                            unsafe { *out_addr.at(bi * rows + r) = dot };
9348                        }
9349                    });
9350                    continue;
9351                }
9352                let l = (1i32 << (bw - 1)) - 1;
9353                let data = &bytes[offsets[r]..offsets[r + 1]];
9354                ROW_I8.with(|rb| {
9355                    let mut buf = rb.borrow_mut();
9356                    buf.resize(cols, 0);
9357                    #[inline(always)]
9358                    fn fill<const B: usize>(data: &[u8], l: i32, buf: &mut [u8]) {
9359                        for (blk, chunk) in buf.chunks_exact_mut(8).enumerate() {
9360                            let u = unpack8::<B>(&data[blk * B..]);
9361                            for k in 0..8 {
9362                                chunk[k] = (u[k] - l) as i8 as u8;
9363                            }
9364                        }
9365                    }
9366                    match bw {
9367                        3 => fill::<3>(data, l, &mut buf),
9368                        4 => vbit_fill4(data, &mut buf),
9369                        5 => fill::<5>(data, l, &mut buf),
9370                        6 => fill::<6>(data, l, &mut buf),
9371                        _ => unreachable!("vbit bit-width {bw} (validated at load)"),
9372                    }
9373                    let mut bi = 0usize;
9374                    // The vbit scale table shares q4_block's layout
9375                    // (contiguous f16 per (row·ng + g)), so the same
9376                    // blocked 1×4 kernel serves the decoded row.
9377                    #[cfg(target_arch = "x86_64")]
9378                    if avx2_enabled() && blocked_enabled() {
9379                        while bi + 4 <= acts.len() {
9380                            let xs = [
9381                                acts[bi].xq.as_slice(),
9382                                acts[bi + 1].xq.as_slice(),
9383                                acts[bi + 2].xq.as_slice(),
9384                                acts[bi + 3].xq.as_slice(),
9385                            ];
9386                            let sxs = [
9387                                acts[bi].sx,
9388                                acts[bi + 1].sx,
9389                                acts[bi + 2].sx,
9390                                acts[bi + 3].sx,
9391                            ];
9392                            let d = unsafe {
9393                                if vnni_tiles_enabled() {
9394                                    dot_q4b_row_1x4_sx_vnni(
9395                                        &buf,
9396                                        &bytes[sc_off..],
9397                                        r * ng,
9398                                        ng,
9399                                        xs,
9400                                        sxs,
9401                                    )
9402                                } else {
9403                                    dot_q4b_row_1x4_sx_avx2(
9404                                        &buf,
9405                                        &bytes[sc_off..],
9406                                        r * ng,
9407                                        ng,
9408                                        xs,
9409                                        sxs,
9410                                    )
9411                                }
9412                            };
9413                            for k in 0..4 {
9414                                let act = &acts[bi + k];
9415                                let mut dot = d[k];
9416                                for &(j, xv) in &act.outliers {
9417                                    dot += (buf[j] as i8) as f32 * gscale(r, j / GROUP_SIZE) * xv;
9418                                }
9419                                // SAFETY: disjoint (bi, r) cells per worker.
9420                                unsafe { *out_addr.at((bi + k) * rows + r) = dot };
9421                            }
9422                            bi += 4;
9423                        }
9424                    }
9425                    while bi < acts.len() {
9426                        let act = &acts[bi];
9427                        let mut dot = 0f32;
9428                        for g in 0..ng {
9429                            let d = dot_i8_i8(
9430                                &buf[g * GROUP_SIZE..(g + 1) * GROUP_SIZE],
9431                                &act.xq[g * GROUP_SIZE..(g + 1) * GROUP_SIZE],
9432                            ) as f32
9433                                * act.sx;
9434                            dot += d * gscale(r, g);
9435                        }
9436                        for &(j, xv) in &act.outliers {
9437                            dot += (buf[j] as i8) as f32 * gscale(r, j / GROUP_SIZE) * xv;
9438                        }
9439                        // SAFETY: disjoint (bi, r) cells per worker range.
9440                        unsafe { *out_addr.at(bi * rows + r) = dot };
9441                        bi += 1;
9442                    }
9443                });
9444            }
9445        };
9446        dispatch_rows(pool, rows, &run);
9447        return;
9448    }
9449
9450    let out_addr = SendMut(out.as_mut_ptr());
9451    let run = move |start: usize, end: usize| {
9452        ROW_F32.with(|rb| {
9453            let mut buf = rb.borrow_mut();
9454            buf.resize(cols, 0.0);
9455            for r in start..end {
9456                decode_f32(r, &mut buf);
9457                for bi in 0..b {
9458                    let x = &xs_all[bi * cols..(bi + 1) * cols];
9459                    let mut dot = 0f32;
9460                    for g in 0..ng {
9461                        let mut gd = 0f32;
9462                        for k in 0..GROUP_SIZE {
9463                            gd += buf[g * GROUP_SIZE + k] * x[g * GROUP_SIZE + k];
9464                        }
9465                        dot += gd * gscale(r, g);
9466                    }
9467                    // SAFETY: disjoint (bi, r) cells per worker range.
9468                    unsafe { *out_addr.at(bi * rows + r) = dot };
9469                }
9470            }
9471        })
9472    };
9473    dispatch_rows(pool, rows, &run);
9474}
9475
9476/// Build a GPU batch job for a q8-family mapped tensor (primary
9477/// shard): prescaled input + directory coordinates. None → not
9478/// GPU-eligible, caller stays on the CPU.
9479pub(crate) fn gpu_batch_job<'a>(
9480    t: &'a QTensor,
9481    x: &[f32],
9482) -> Option<(std::sync::Arc<CmfModel>, crate::gpu::BatchJob<'a>)> {
9483    match t {
9484        QTensor::Mapped {
9485            model,
9486            idx,
9487            dtype: dt @ (TensorDtype::Q8Row | TensorDtype::Q8_2f),
9488            rows,
9489            cols,
9490            row_scale,
9491            col_field,
9492            ..
9493        } => Some((
9494            model.clone(),
9495            crate::gpu::BatchJob {
9496                idx: *idx,
9497                rows: *rows,
9498                cols: *cols,
9499                row_scale,
9500                xs: prescale(x, col_field, *dt).into_owned(),
9501                layout: crate::gpu::BatchLayout::Q8,
9502            },
9503        )),
9504        // q1: raw f32 activations, tile-embedded scales.
9505        QTensor::Mapped {
9506            model,
9507            idx,
9508            dtype: TensorDtype::Q1,
9509            rows,
9510            cols,
9511            ..
9512        } => Some((
9513            model.clone(),
9514            crate::gpu::BatchJob {
9515                idx: *idx,
9516                rows: *rows,
9517                cols: *cols,
9518                row_scale: &[],
9519                xs: x.to_vec(),
9520                layout: crate::gpu::BatchLayout::Q1,
9521            },
9522        )),
9523        // q4_tiled / q4tp: raw f32 activations; the scales live in the
9524        // payload (inline tiles / row ladder), so row_scale stays empty.
9525        // The GDN projection batch already runs these layouts on Metal —
9526        // this arm lets the attention QKV batch reach the same kernels.
9527        QTensor::Mapped {
9528            model,
9529            idx,
9530            dtype: dt @ (TensorDtype::Q4Tiled | TensorDtype::Q4TiledP),
9531            rows,
9532            cols,
9533            ..
9534        } => Some((
9535            model.clone(),
9536            crate::gpu::BatchJob {
9537                idx: *idx,
9538                rows: *rows,
9539                cols: *cols,
9540                row_scale: &[],
9541                xs: x.to_vec(),
9542                layout: if *dt == TensorDtype::Q4Tiled {
9543                    crate::gpu::BatchLayout::Q4t
9544                } else {
9545                    crate::gpu::BatchLayout::Q4tp
9546                },
9547            },
9548        )),
9549        _ => None,
9550    }
9551}
9552
9553thread_local! {
9554    static PRESCALE_BUF1: std::cell::RefCell<Vec<f32>> = const { std::cell::RefCell::new(Vec::new()) };
9555    static PRESCALE_BUF2: std::cell::RefCell<Vec<f32>> = const { std::cell::RefCell::new(Vec::new()) };
9556}
9557
9558pub(crate) fn prescale<'a>(
9559    x: &'a [f32],
9560    col_field: &[f32],
9561    dtype: TensorDtype,
9562) -> std::borrow::Cow<'a, [f32]> {
9563    if dtype == TensorDtype::Q8_2f {
9564        x.iter().zip(col_field).map(|(a, c)| a * c).collect()
9565    } else {
9566        std::borrow::Cow::Borrowed(x)
9567    }
9568}
9569
9570/// θ col-field fold for q8_2f activations. Borrowed pass-through for
9571/// every other dtype, using thread-local buffers to eliminate per-matvec allocations.
9572pub(crate) fn prescale_with<R, F: FnOnce(&[f32]) -> R>(
9573    x: &[f32],
9574    col_field: &[f32],
9575    dtype: TensorDtype,
9576    buf_id: u8,
9577    f: F,
9578) -> R {
9579    if dtype == TensorDtype::Q8_2f {
9580        if buf_id == 1 {
9581            PRESCALE_BUF1.with(|b| {
9582                let mut buf = b.borrow_mut();
9583                buf.clear();
9584                buf.extend(x.iter().zip(col_field).map(|(a, c)| a * c));
9585                f(&buf)
9586            })
9587        } else {
9588            PRESCALE_BUF2.with(|b| {
9589                let mut buf = b.borrow_mut();
9590                buf.clear();
9591                buf.extend(x.iter().zip(col_field).map(|(a, c)| a * c));
9592                f(&buf)
9593            })
9594        }
9595    } else {
9596        f(x)
9597    }
9598}
9599
9600// ───────────────────── x86-64 AVX2 kernels (roadmap этап 2) ─────────────────────
9601
9602/// AVX2+FMA available? Default ON when the CPU supports both;
9603/// `CMF_AVX2=0` disables (falls back to the autovectorized loops).
9604#[cfg(target_arch = "x86_64")]
9605pub(crate) fn avx2_enabled() -> bool {
9606    use std::sync::OnceLock;
9607    static ON: OnceLock<bool> = OnceLock::new();
9608    *ON.get_or_init(|| {
9609        std::env::var("CMF_AVX2").map(|v| v != "0").unwrap_or(true)
9610            && std::arch::is_x86_feature_detected!("avx2")
9611            && std::arch::is_x86_feature_detected!("fma")
9612    })
9613}
9614
9615/// AVX2 A8W8 allowed? The quantized-activation contract is switched by
9616/// the SAME env as the ARM SDOT path: `CMF_SDOT=0` keeps exact kernels
9617/// (the golden-parity exact gate relies on it) — AVX2 f32 kernels stay
9618/// active either way, they are exact (regrouped sums only).
9619#[cfg(target_arch = "x86_64")]
9620fn avx2_a8w8_enabled() -> bool {
9621    if FLOAT_ACTIVATIONS.get() {
9622        return false;
9623    }
9624    use std::sync::OnceLock;
9625    static ON: OnceLock<bool> = OnceLock::new();
9626    *ON.get_or_init(|| {
9627        avx2_enabled() && std::env::var("CMF_SDOT").map(|v| v != "0").unwrap_or(true)
9628    })
9629}
9630
9631thread_local! {
9632    static FULL_GPU_Q8: std::cell::Cell<bool> = const { std::cell::Cell::new(false) };
9633}
9634
9635/// Match the graph's full-device q8 projection precision on MiMo's host
9636/// tail. Does not enable the GPU or bypass a CPU-only/device-refusal gate.
9637pub(crate) fn enter_full_gpu_q8_scope() -> impl Drop {
9638    struct Restore(bool, std::marker::PhantomData<std::rc::Rc<()>>);
9639    impl Drop for Restore {
9640        fn drop(&mut self) {
9641            FULL_GPU_Q8.set(self.0);
9642        }
9643    }
9644    Restore(FULL_GPU_Q8.replace(true), std::marker::PhantomData)
9645}
9646
9647// Dynamic MiMo experts must not change activation precision when a cache
9648// fill moves them from CPU to GPU. Thread-local: only the cold-expert
9649// dispatch selects float kernels; concurrent pipelines keep their policy.
9650thread_local! {
9651    static FLOAT_ACTIVATIONS: std::cell::Cell<bool> = const { std::cell::Cell::new(false) };
9652}
9653
9654pub(crate) fn float_activations_scope<R>(f: impl FnOnce() -> R) -> R {
9655    struct Restore(bool);
9656    impl Drop for Restore {
9657        fn drop(&mut self) {
9658            FLOAT_ACTIVATIONS.set(self.0);
9659        }
9660    }
9661    let _restore = Restore(FLOAT_ACTIVATIONS.replace(true));
9662    f()
9663}
9664
9665/// Row-exact batching: while set, the x86 batched kernels (`qmatmat`,
9666/// `q4tp_matmat`) compute every (weight row, token) cell with the
9667/// single-token kernel instead of the blocked 2×4 / 1×4 / 1×8 tiles, so a
9668/// token's result does not depend on the batch it rides in and equals its
9669/// matvec. On ARM `q4tp_matmat` keeps its 1×4 tile but in the matvec's
9670/// reduction order, and no q4tp batch takes the AMX or device GEMM. The
9671/// MiMo speculative verify holds it (`row_exact_scope`) — its accepted
9672/// rows must be the rows plain decode would have produced.
9673// Shared with pool workers, so overlapping requests must keep the mode
9674// enabled until the LAST scope leaves. Saving/restoring a global bool is
9675// incorrect when two threads enter and leave in a non-LIFO order.
9676static ROW_EXACT: std::sync::atomic::AtomicUsize = std::sync::atomic::AtomicUsize::new(0);
9677
9678pub(crate) fn row_exact() -> bool {
9679    ROW_EXACT.load(std::sync::atomic::Ordering::Acquire) != 0
9680}
9681
9682fn counted_row_exact_scope<R>(active: &std::sync::atomic::AtomicUsize, f: impl FnOnce() -> R) -> R {
9683    struct Restore<'a>(&'a std::sync::atomic::AtomicUsize);
9684    impl Drop for Restore<'_> {
9685        fn drop(&mut self) {
9686            self.0.fetch_sub(1, std::sync::atomic::Ordering::AcqRel);
9687        }
9688    }
9689    active.fetch_add(1, std::sync::atomic::Ordering::AcqRel);
9690    let _restore = Restore(active);
9691    f()
9692}
9693
9694/// Run `f` with row-exact batching on (also released on unwind).
9695pub(crate) fn row_exact_scope<R>(f: impl FnOnce() -> R) -> R {
9696    counted_row_exact_scope(&ROW_EXACT, f)
9697}
9698
9699/// A8W8 quantized-activation path available on THIS machine? One
9700/// switch across architectures: ARM dotprod (CMF_SDOT) or x86 AVX2
9701/// (CMF_AVX2 + the same CMF_SDOT exact-contract override).
9702#[inline]
9703pub(crate) fn a8w8_enabled() -> bool {
9704    #[cfg(target_arch = "aarch64")]
9705    {
9706        sdot_enabled()
9707    }
9708    #[cfg(target_arch = "x86_64")]
9709    {
9710        avx2_a8w8_enabled()
9711    }
9712    #[cfg(not(any(target_arch = "aarch64", target_arch = "x86_64")))]
9713    {
9714        false
9715    }
9716}
9717
9718/// int8·int8 dot dispatch: SDOT on ARM; AVX-512 VNNI (vpdpbusd) or AVX2
9719/// maddubs on x86. Callers are gated by `a8w8_enabled()`.
9720#[inline]
9721#[allow(unreachable_code)]
9722fn dot_i8_i8(w: &[u8], xq: &[i8]) -> i32 {
9723    #[cfg(target_arch = "aarch64")]
9724    unsafe {
9725        return dot_i8_sdot(w, xq);
9726    }
9727    #[cfg(target_arch = "x86_64")]
9728    unsafe {
9729        if avx512vnni_enabled() {
9730            return dot_i8_i8_vnni(w, xq);
9731        }
9732        return dot_i8_i8_avx2(w, xq);
9733    }
9734    w.iter()
9735        .zip(xq)
9736        .map(|(&a, &b)| (a as i8) as i32 * b as i32)
9737        .sum()
9738}
9739
9740/// AVX-512 VNNI available? (F+BW+VL+VNNI; `CMF_AVX512=0` falls back to
9741/// AVX2.) VL matters: short 32-byte groups (q4/vbit) ride the 256-bit
9742/// `vpdpbusd` encoding.
9743#[cfg(target_arch = "x86_64")]
9744fn avx512vnni_enabled() -> bool {
9745    use std::sync::OnceLock;
9746    static ON: OnceLock<bool> = OnceLock::new();
9747    *ON.get_or_init(|| {
9748        std::env::var("CMF_AVX512")
9749            .map(|v| v != "0")
9750            .unwrap_or(true)
9751            && std::arch::is_x86_feature_detected!("avx512f")
9752            && std::arch::is_x86_feature_detected!("avx512bw")
9753            && std::arch::is_x86_feature_detected!("avx512vl")
9754            && std::arch::is_x86_feature_detected!("avx512vnni")
9755    })
9756}
9757
9758/// Grouped-codec VNNI arms (the q4t/q4b/q1/q1t tile kernels): default
9759/// ON where AVX-512 VNNI exists (`CMF_VNNI_TILES=0` opt-out). Measured
9760/// on Ryzen 7950X (Zen4, 3 alternating process pairs, blocked GEMM
9761/// 4864×896 b=256): q4t 63→68 GF/s (+8%), q1 53→56 (+6%), q4b 72→75
9762/// (+4%) — consistent, no leg regressed. The tile kernels keep a
9763/// horizontal reduce per 32-weight group, so the `vpdpbusd` saving is
9764/// smaller than the long-dot q8 win (+13%), but it is real and free.
9765#[cfg(target_arch = "x86_64")]
9766fn vnni_tiles_enabled() -> bool {
9767    use std::sync::OnceLock;
9768    static ON: OnceLock<bool> = OnceLock::new();
9769    *ON.get_or_init(|| {
9770        std::env::var("CMF_VNNI_TILES")
9771            .map(|v| v != "0")
9772            .unwrap_or(true)
9773            && avx512vnni_enabled()
9774    })
9775}
9776
9777/// One 256-bit u8×i8 dot → i32 via `vpdpbusd` into a fresh accumulator
9778/// plus the same horizontal reduce the AVX2 kernels use. Products are
9779/// bounded (|w| ≤ 8 or ≤ 1), so maddubs never saturated — the i32 sum
9780/// is bit-identical to the maddubs+madd pair it replaces.
9781#[cfg(target_arch = "x86_64")]
9782#[target_feature(enable = "avx2,avx512f,avx512bw,avx512vl,avx512vnni")]
9783#[inline]
9784unsafe fn dpbusd_hsum(aw: core::arch::x86_64::__m256i, xs: core::arch::x86_64::__m256i) -> i32 {
9785    // SAFETY: pure register math.
9786    unsafe {
9787        use core::arch::x86_64::*;
9788        let d = _mm256_dpbusd_epi32(_mm256_setzero_si256(), aw, xs);
9789        let hi128 = _mm256_extracti128_si256::<1>(d);
9790        let s128 = _mm_add_epi32(_mm256_castsi256_si128(d), hi128);
9791        let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
9792        let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
9793        _mm_cvtsi128_si32(s32)
9794    }
9795}
9796
9797/// int8·int8 via AVX-512 VNNI: `vpdpbusd` fuses the maddubs+madd+add
9798/// triple into one u8×i8 dot-accumulate. AVX-512 has no vpsignb, so the
9799/// |w|·sign(x,w) trick becomes |w| × (x negated where w<0) via a mask
9800/// subtract — w==0 lanes contribute 0 through |w|=0 either way.
9801#[cfg(target_arch = "x86_64")]
9802#[target_feature(enable = "avx2,avx512f,avx512bw,avx512vl,avx512vnni")]
9803unsafe fn dot_i8_i8_vnni(w: &[u8], xq: &[i8]) -> i32 {
9804    // SAFETY: callers uphold slice-length contracts (see call sites).
9805    unsafe {
9806        use core::arch::x86_64::*;
9807        let n = w.len();
9808        let mut j = 0usize;
9809        let mut total: i32;
9810        // 4 independent accumulators: vpdpbusd is its own loop-carried
9811        // dependency (~5-cycle latency) — a single-acc loop runs
9812        // latency-bound and LOSES to the AVX2 maddubs kernel, measured
9813        // on Granite Rapids.
9814        {
9815            #[inline(always)]
9816            unsafe fn step(
9817                w: *const u8,
9818                x: *const i8,
9819                acc: core::arch::x86_64::__m512i,
9820            ) -> core::arch::x86_64::__m512i {
9821                unsafe {
9822                    use core::arch::x86_64::*;
9823                    let wv = _mm512_loadu_si512(w as *const _);
9824                    let xv = _mm512_loadu_si512(x as *const _);
9825                    let aw = _mm512_abs_epi8(wv);
9826                    let neg = _mm512_movepi8_mask(wv);
9827                    let sx = _mm512_mask_sub_epi8(xv, neg, _mm512_setzero_si512(), xv);
9828                    _mm512_dpbusd_epi32(acc, aw, sx)
9829                }
9830            }
9831            let (mut a0, mut a1, mut a2, mut a3) = (
9832                _mm512_setzero_si512(),
9833                _mm512_setzero_si512(),
9834                _mm512_setzero_si512(),
9835                _mm512_setzero_si512(),
9836            );
9837            while j + 256 <= n {
9838                a0 = step(w.as_ptr().add(j), xq.as_ptr().add(j), a0);
9839                a1 = step(w.as_ptr().add(j + 64), xq.as_ptr().add(j + 64), a1);
9840                a2 = step(w.as_ptr().add(j + 128), xq.as_ptr().add(j + 128), a2);
9841                a3 = step(w.as_ptr().add(j + 192), xq.as_ptr().add(j + 192), a3);
9842                j += 256;
9843            }
9844            while j + 64 <= n {
9845                a0 = step(w.as_ptr().add(j), xq.as_ptr().add(j), a0);
9846                j += 64;
9847            }
9848            let s01 = _mm512_add_epi32(a0, a1);
9849            let s23 = _mm512_add_epi32(a2, a3);
9850            total = _mm512_reduce_add_epi32(_mm512_add_epi32(s01, s23));
9851        }
9852        // 32-wide (q4/vbit groups are exactly 32 bytes).
9853        if j + 32 <= n {
9854            let wv = _mm256_loadu_si256(w.as_ptr().add(j) as *const __m256i);
9855            let xv = _mm256_loadu_si256(xq.as_ptr().add(j) as *const __m256i);
9856            let d = _mm256_dpbusd_epi32(
9857                _mm256_setzero_si256(),
9858                _mm256_abs_epi8(wv),
9859                _mm256_sign_epi8(xv, wv),
9860            );
9861            let hi128 = _mm256_extracti128_si256::<1>(d);
9862            let s128 = _mm_add_epi32(_mm256_castsi256_si128(d), hi128);
9863            let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
9864            let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
9865            total += _mm_cvtsi128_si32(s32);
9866            j += 32;
9867        }
9868        while j < n {
9869            total += (w[j] as i8) as i32 * xq[j] as i32;
9870            j += 1;
9871        }
9872        total
9873    }
9874}
9875
9876/// i8 row · f32 x via AVX2/FMA (x86 mirror of `dot_i8_f32_neon`).
9877#[cfg(target_arch = "x86_64")]
9878#[target_feature(enable = "avx2,fma")]
9879unsafe fn dot_i8_f32_avx2(w: &[u8], x: &[f32]) -> f32 {
9880    // SAFETY: callers uphold slice-length contracts (see call sites).
9881    unsafe {
9882        use core::arch::x86_64::*;
9883        let n = x.len();
9884        let wp = w.as_ptr();
9885        let xp = x.as_ptr();
9886        let (mut a0, mut a1) = (_mm256_setzero_ps(), _mm256_setzero_ps());
9887        let mut j = 0usize;
9888        while j + 16 <= n {
9889            let wb = _mm_loadu_si128(wp.add(j) as *const __m128i);
9890            let lo = _mm256_cvtepi8_epi32(wb);
9891            let hi = _mm256_cvtepi8_epi32(_mm_srli_si128::<8>(wb));
9892            a0 = _mm256_fmadd_ps(_mm256_cvtepi32_ps(lo), _mm256_loadu_ps(xp.add(j)), a0);
9893            a1 = _mm256_fmadd_ps(_mm256_cvtepi32_ps(hi), _mm256_loadu_ps(xp.add(j + 8)), a1);
9894            j += 16;
9895        }
9896        let acc = _mm256_add_ps(a0, a1);
9897        let hi128 = _mm256_extractf128_ps::<1>(acc);
9898        let s128 = _mm_add_ps(_mm256_castps256_ps128(acc), hi128);
9899        let s64 = _mm_add_ps(s128, _mm_movehl_ps(s128, s128));
9900        let s32 = _mm_add_ss(s64, _mm_shuffle_ps::<1>(s64, s64));
9901        let mut sum = _mm_cvtss_f32(s32);
9902        while j < n {
9903            sum += (*wp.add(j) as i8) as f32 * *xp.add(j);
9904            j += 1;
9905        }
9906        sum
9907    }
9908}
9909
9910/// int8(weight)·int8(activation) → i32 via AVX2 maddubs — the x86
9911/// analogue of the SDOT A8W8 path. `maddubs` takes u8×i8, so the
9912/// standard sign trick applies: |w| × sign(x, w) ≡ w × x per lane.
9913/// Pair saturation is safe: |w|≤128, |x|≤127 → 2·128·127 < 32767.
9914#[cfg(target_arch = "x86_64")]
9915#[target_feature(enable = "avx2")]
9916unsafe fn dot_i8_i8_avx2(w: &[u8], xq: &[i8]) -> i32 {
9917    // SAFETY: callers uphold slice-length contracts (see call sites).
9918    unsafe {
9919        use core::arch::x86_64::*;
9920        let n = w.len();
9921        let ones = _mm256_set1_epi16(1);
9922        let mut acc = _mm256_setzero_si256();
9923        let mut j = 0usize;
9924        while j + 32 <= n {
9925            let wv = _mm256_loadu_si256(w.as_ptr().add(j) as *const __m256i);
9926            let xv = _mm256_loadu_si256(xq.as_ptr().add(j) as *const __m256i);
9927            let p16 = _mm256_maddubs_epi16(_mm256_abs_epi8(wv), _mm256_sign_epi8(xv, wv));
9928            acc = _mm256_add_epi32(acc, _mm256_madd_epi16(p16, ones));
9929            j += 32;
9930        }
9931        let hi128 = _mm256_extracti128_si256::<1>(acc);
9932        let s128 = _mm_add_epi32(_mm256_castsi256_si128(acc), hi128);
9933        let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
9934        let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
9935        let mut s = _mm_cvtsi128_si32(s32);
9936        while j < n {
9937            s += (w[j] as i8) as i32 * xq[j] as i32;
9938            j += 1;
9939        }
9940        s
9941    }
9942}
9943
9944/// smmla 2×4: one instruction covers a 2-row × 2-activation × 8-deep
9945/// tile (32 MACs vs sdot's 16) — the weight pair loads once per 8-k
9946/// slice as a combined 2×8 register and meets two activation pairs.
9947#[cfg(target_arch = "aarch64")]
9948#[target_feature(enable = "neon,i8mm")]
9949unsafe fn dot_i8_smmla_2x4(w0: &[u8], w1: &[u8], xs: [&[i8]; 4]) -> [[i32; 4]; 2] {
9950    // SAFETY: callers uphold slice-length contracts.
9951    unsafe {
9952        use core::arch::aarch64::*;
9953        use core::arch::asm;
9954        let n = w0.len();
9955        let w0p = w0.as_ptr() as *const i8;
9956        let w1p = w1.as_ptr() as *const i8;
9957        // acc01 holds [c(r0,x0) c(r0,x1) c(r1,x0) c(r1,x1)]; acc23 the
9958        // same for x2/x3.
9959        let mut acc01 = vdupq_n_s32(0);
9960        let mut acc23 = vdupq_n_s32(0);
9961        let mut i = 0usize;
9962        while i + 8 <= n {
9963            let wa = vcombine_s8(vld1_s8(w0p.add(i)), vld1_s8(w1p.add(i)));
9964            let xb01 = vcombine_s8(
9965                vld1_s8(xs[0].as_ptr().add(i)),
9966                vld1_s8(xs[1].as_ptr().add(i)),
9967            );
9968            let xb23 = vcombine_s8(
9969                vld1_s8(xs[2].as_ptr().add(i)),
9970                vld1_s8(xs[3].as_ptr().add(i)),
9971            );
9972            asm!(
9973                "smmla {a01:v}.4s, {w:v}.16b, {x01:v}.16b",
9974                "smmla {a23:v}.4s, {w:v}.16b, {x23:v}.16b",
9975                a01 = inout(vreg) acc01, a23 = inout(vreg) acc23,
9976                w = in(vreg) wa, x01 = in(vreg) xb01, x23 = in(vreg) xb23,
9977                options(pure, nomem, nostack),
9978            );
9979            i += 8;
9980        }
9981        let mut out = [[0i32; 4]; 2];
9982        let a01: [i32; 4] = core::mem::transmute(acc01);
9983        let a23: [i32; 4] = core::mem::transmute(acc23);
9984        out[0][0] = a01[0];
9985        out[0][1] = a01[1];
9986        out[1][0] = a01[2];
9987        out[1][1] = a01[3];
9988        out[0][2] = a23[0];
9989        out[0][3] = a23[1];
9990        out[1][2] = a23[2];
9991        out[1][3] = a23[3];
9992        if i < n {
9993            for (k, x) in xs.iter().enumerate() {
9994                for j in i..n {
9995                    out[0][k] += (w0[j] as i8) as i32 * x[j] as i32;
9996                    out[1][k] += (w1[j] as i8) as i32 * x[j] as i32;
9997                }
9998            }
9999        }
10000        out
10001    }
10002}
10003
10004/// ARM twin of the x86 blocked prefill GEMM: two weight rows stay in
10005/// registers across four activation streams, eight sdot accumulators.
10006/// (The per-row form re-read each W row once per activation.)
10007#[cfg(target_arch = "aarch64")]
10008#[target_feature(enable = "neon,dotprod")]
10009unsafe fn dot_i8_sdot_2x4(w0: &[u8], w1: &[u8], xs: [&[i8]; 4]) -> [[i32; 4]; 2] {
10010    // SAFETY: callers uphold slice-length contracts.
10011    unsafe {
10012        use core::arch::aarch64::*;
10013        use core::arch::asm;
10014        let n = w0.len();
10015        let w0p = w0.as_ptr() as *const i8;
10016        let w1p = w1.as_ptr() as *const i8;
10017        let mut acc = [[vdupq_n_s32(0); 4]; 2];
10018        let mut i = 0usize;
10019        while i + 16 <= n {
10020            let wv0 = vld1q_s8(w0p.add(i));
10021            let wv1 = vld1q_s8(w1p.add(i));
10022            for (k, x) in xs.iter().enumerate() {
10023                let xv = vld1q_s8(x.as_ptr().add(i));
10024                let (mut a0, mut a1) = (acc[0][k], acc[1][k]);
10025                asm!(
10026                    "sdot {a0:v}.4s, {w0:v}.16b, {x:v}.16b",
10027                    "sdot {a1:v}.4s, {w1:v}.16b, {x:v}.16b",
10028                    a0 = inout(vreg) a0, a1 = inout(vreg) a1,
10029                    w0 = in(vreg) wv0, w1 = in(vreg) wv1, x = in(vreg) xv,
10030                    options(pure, nomem, nostack),
10031                );
10032                acc[0][k] = a0;
10033                acc[1][k] = a1;
10034            }
10035            i += 16;
10036        }
10037        let mut out = [[0i32; 4]; 2];
10038        for r in 0..2 {
10039            for k in 0..4 {
10040                out[r][k] = vaddvq_s32(acc[r][k]);
10041            }
10042        }
10043        if i < n {
10044            for (k, x) in xs.iter().enumerate() {
10045                for j in i..n {
10046                    out[0][k] += (w0[j] as i8) as i32 * x[j] as i32;
10047                    out[1][k] += (w1[j] as i8) as i32 * x[j] as i32;
10048                }
10049            }
10050        }
10051        out
10052    }
10053}
10054
10055/// Blocked 2 weight rows × 4 activations for the prefill GEMM
10056/// (roadmap P0: packed panels + multi-row accumulators). The two rows'
10057/// abs() live in registers across all four activation streams; the
10058/// sign-fixup is recomputed per pair (the price of the maddubs trick).
10059/// Returns raw i8·i8 dots; the caller applies scales and outliers.
10060#[cfg(target_arch = "x86_64")]
10061#[target_feature(enable = "avx2")]
10062unsafe fn dot_i8_i8_avx2_2x4(w0: &[u8], w1: &[u8], xs: [&[i8]; 4]) -> [[i32; 4]; 2] {
10063    // SAFETY: callers uphold slice-length contracts.
10064    unsafe {
10065        use core::arch::x86_64::*;
10066        let n = w0.len();
10067        let ones = _mm256_set1_epi16(1);
10068        let mut acc = [[_mm256_setzero_si256(); 4]; 2];
10069        let mut j = 0usize;
10070        while j + 32 <= n {
10071            let wv0 = _mm256_loadu_si256(w0.as_ptr().add(j) as *const __m256i);
10072            let wv1 = _mm256_loadu_si256(w1.as_ptr().add(j) as *const __m256i);
10073            let aw0 = _mm256_abs_epi8(wv0);
10074            let aw1 = _mm256_abs_epi8(wv1);
10075            for (k, x) in xs.iter().enumerate() {
10076                let xv = _mm256_loadu_si256(x.as_ptr().add(j) as *const __m256i);
10077                let p0 = _mm256_maddubs_epi16(aw0, _mm256_sign_epi8(xv, wv0));
10078                acc[0][k] = _mm256_add_epi32(acc[0][k], _mm256_madd_epi16(p0, ones));
10079                let p1 = _mm256_maddubs_epi16(aw1, _mm256_sign_epi8(xv, wv1));
10080                acc[1][k] = _mm256_add_epi32(acc[1][k], _mm256_madd_epi16(p1, ones));
10081            }
10082            j += 32;
10083        }
10084        let mut out = [[0i32; 4]; 2];
10085        for r in 0..2 {
10086            for k in 0..4 {
10087                let a = acc[r][k];
10088                let hi128 = _mm256_extracti128_si256::<1>(a);
10089                let s128 = _mm_add_epi32(_mm256_castsi256_si128(a), hi128);
10090                let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
10091                let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
10092                out[r][k] = _mm_cvtsi128_si32(s32);
10093            }
10094        }
10095        if j < n {
10096            for (k, x) in xs.iter().enumerate() {
10097                for i in j..n {
10098                    out[0][k] += (w0[i] as i8) as i32 * x[i] as i32;
10099                    out[1][k] += (w1[i] as i8) as i32 * x[i] as i32;
10100                }
10101            }
10102        }
10103        out
10104    }
10105}
10106
10107/// AVX2/VNNI q8 row dot with exact outlier correction (x86 mirror of
10108/// `row_dot_sdot` — same A8W8 contract). With AVX-512 VNNI the row goes
10109/// through the bias trick: Σ(w+128)·x via pure `vpdpbusd` (no per-lane
10110/// sign fixups), corrected by −128·Σx with Σx precomputed per split.
10111#[cfg(target_arch = "x86_64")]
10112#[inline]
10113fn row_dot_avx2(row: &[u8], act: &SplitAct) -> f32 {
10114    let dot = if avx512vnni_enabled() && row.len() >= 64 {
10115        (unsafe { dot_u8p128_i8_vnni(row, &act.xq) }) - 128 * act.xsum
10116    } else {
10117        unsafe { dot_i8_i8_avx2(row, &act.xq) }
10118    };
10119    let mut acc = dot as f32 * act.sx;
10120    for &(j, xv) in &act.outliers {
10121        acc += (row[j] as i8) as f32 * xv;
10122    }
10123    acc
10124}
10125
10126/// Σ (w[i]+128)·x[i] via pure `vpdpbusd` — the caller subtracts
10127/// 128·Σx. Four independent accumulators (dpbusd is ~5-cycle latency;
10128/// a single-acc loop runs latency-bound, measured on Granite Rapids).
10129#[cfg(target_arch = "x86_64")]
10130#[target_feature(enable = "avx2,avx512f,avx512bw,avx512vl,avx512vnni")]
10131unsafe fn dot_u8p128_i8_vnni(w: &[u8], xq: &[i8]) -> i32 {
10132    // SAFETY: callers uphold slice-length contracts (see call sites).
10133    unsafe {
10134        use core::arch::x86_64::*;
10135        let n = w.len();
10136        let flip = _mm512_set1_epi8(-128); // XOR 0x80: i8 w → u8 (w+128)
10137        #[inline(always)]
10138        unsafe fn step(
10139            w: *const u8,
10140            x: *const i8,
10141            flip: core::arch::x86_64::__m512i,
10142            acc: core::arch::x86_64::__m512i,
10143        ) -> core::arch::x86_64::__m512i {
10144            unsafe {
10145                use core::arch::x86_64::*;
10146                let wv = _mm512_xor_si512(_mm512_loadu_si512(w as *const _), flip);
10147                _mm512_dpbusd_epi32(acc, wv, _mm512_loadu_si512(x as *const _))
10148            }
10149        }
10150        let (mut a0, mut a1, mut a2, mut a3) = (
10151            _mm512_setzero_si512(),
10152            _mm512_setzero_si512(),
10153            _mm512_setzero_si512(),
10154            _mm512_setzero_si512(),
10155        );
10156        let mut j = 0usize;
10157        while j + 256 <= n {
10158            a0 = step(w.as_ptr().add(j), xq.as_ptr().add(j), flip, a0);
10159            a1 = step(w.as_ptr().add(j + 64), xq.as_ptr().add(j + 64), flip, a1);
10160            a2 = step(w.as_ptr().add(j + 128), xq.as_ptr().add(j + 128), flip, a2);
10161            a3 = step(w.as_ptr().add(j + 192), xq.as_ptr().add(j + 192), flip, a3);
10162            j += 256;
10163        }
10164        while j + 64 <= n {
10165            a0 = step(w.as_ptr().add(j), xq.as_ptr().add(j), flip, a0);
10166            j += 64;
10167        }
10168        let mut total = _mm512_reduce_add_epi32(_mm512_add_epi32(
10169            _mm512_add_epi32(a0, a1),
10170            _mm512_add_epi32(a2, a3),
10171        ));
10172        // Scalar tail: (w as i8) + 128 ≡ (w as u8) ^ 0x80.
10173        while j < n {
10174            total += ((w[j] ^ 0x80) as i32) * xq[j] as i32;
10175            j += 1;
10176        }
10177        total
10178    }
10179}
10180
10181/// One q4 row via AVX2: nibbles → centered i8 (unpacklo/hi restores the
10182/// writer's flat order, same as the NEON vzip pair), maddubs against
10183/// the pre-quantized activation group, × the group's f16 scale. Pair
10184/// saturation safe: |w|≤8, |x|≤127 → 2·8·127 ≪ 32767. Mirror of
10185/// `dot_q4_row_sdot`.
10186#[cfg(target_arch = "x86_64")]
10187#[target_feature(enable = "avx2")]
10188unsafe fn dot_q4_row_avx2(packed: &[u8], scales: &[u8], g0: usize, gpr: usize, xq: &[i8]) -> f32 {
10189    // SAFETY: callers uphold slice-length contracts (16 packed bytes and
10190    // 2 scale bytes per group; xq.len() == gpr·GROUP_SIZE).
10191    unsafe {
10192        use core::arch::x86_64::*;
10193        let lomask = _mm_set1_epi8(0x0F);
10194        let eight = _mm256_set1_epi8(8);
10195        let ones = _mm256_set1_epi16(1);
10196        let mut acc = 0f32;
10197        for gi in 0..gpr {
10198            let g = g0 + gi;
10199            let s = f16_to_f32(u16::from_le_bytes([scales[g * 2], scales[g * 2 + 1]]));
10200            let b = _mm_loadu_si128(packed.as_ptr().add(g * 16) as *const __m128i);
10201            let lo = _mm_and_si128(b, lomask);
10202            let hi = _mm_and_si128(_mm_srli_epi16::<4>(b), lomask);
10203            let w = _mm256_sub_epi8(
10204                _mm256_set_m128i(_mm_unpackhi_epi8(lo, hi), _mm_unpacklo_epi8(lo, hi)),
10205                eight,
10206            );
10207            let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
10208            let p16 = _mm256_maddubs_epi16(_mm256_abs_epi8(w), _mm256_sign_epi8(x, w));
10209            let d = _mm256_madd_epi16(p16, ones);
10210            let hi128 = _mm256_extracti128_si256::<1>(d);
10211            let s128 = _mm_add_epi32(_mm256_castsi256_si128(d), hi128);
10212            let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
10213            let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
10214            acc += _mm_cvtsi128_si32(s32) as f32 * s;
10215        }
10216        acc
10217    }
10218}
10219
10220/// Two-activation q4 row via AVX2: nibbles unpacked ONCE per group,
10221/// both activations dotted against the same centered i8 register.
10222#[cfg(target_arch = "x86_64")]
10223#[target_feature(enable = "avx2")]
10224unsafe fn dot_q4_row_avx2_2(
10225    packed: &[u8],
10226    scales: &[u8],
10227    g0: usize,
10228    gpr: usize,
10229    xq1: &[i8],
10230    xq2: &[i8],
10231) -> (f32, f32) {
10232    // SAFETY: callers uphold slice-length contracts (see dot_q4_row_avx2).
10233    unsafe {
10234        use core::arch::x86_64::*;
10235        let lomask = _mm_set1_epi8(0x0F);
10236        let eight = _mm256_set1_epi8(8);
10237        let ones = _mm256_set1_epi16(1);
10238        let (mut acc1, mut acc2) = (0f32, 0f32);
10239        #[inline(always)]
10240        unsafe fn hsum(d: core::arch::x86_64::__m256i) -> i32 {
10241            unsafe {
10242                use core::arch::x86_64::*;
10243                let hi128 = _mm256_extracti128_si256::<1>(d);
10244                let s128 = _mm_add_epi32(_mm256_castsi256_si128(d), hi128);
10245                let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
10246                let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
10247                _mm_cvtsi128_si32(s32)
10248            }
10249        }
10250        for gi in 0..gpr {
10251            let g = g0 + gi;
10252            let s = f16_to_f32(u16::from_le_bytes([scales[g * 2], scales[g * 2 + 1]]));
10253            let b = _mm_loadu_si128(packed.as_ptr().add(g * 16) as *const __m128i);
10254            let lo = _mm_and_si128(b, lomask);
10255            let hi = _mm_and_si128(_mm_srli_epi16::<4>(b), lomask);
10256            let w = _mm256_sub_epi8(
10257                _mm256_set_m128i(_mm_unpackhi_epi8(lo, hi), _mm_unpacklo_epi8(lo, hi)),
10258                eight,
10259            );
10260            let aw = _mm256_abs_epi8(w);
10261            let x1 = _mm256_loadu_si256(xq1.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
10262            let x2 = _mm256_loadu_si256(xq2.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
10263            let d1 = _mm256_madd_epi16(_mm256_maddubs_epi16(aw, _mm256_sign_epi8(x1, w)), ones);
10264            let d2 = _mm256_madd_epi16(_mm256_maddubs_epi16(aw, _mm256_sign_epi8(x2, w)), ones);
10265            acc1 += hsum(d1) as f32 * s;
10266            acc2 += hsum(d2) as f32 * s;
10267        }
10268        (acc1, acc2)
10269    }
10270}
10271
10272/// One q8 row range via AVX2 (x86 mirror of `q8_range_sdot`).
10273#[cfg(target_arch = "x86_64")]
10274fn q8_range_avx2(
10275    q: &[u8],
10276    row_scale: &[f32],
10277    act: &SplitAct,
10278    cols: usize,
10279    out_addr: SendMut,
10280    start: usize,
10281    end: usize,
10282) {
10283    for o in start..end {
10284        let v = row_dot_avx2(&q[o * cols..(o + 1) * cols], act) * row_scale[o];
10285        // SAFETY: disjoint row ranges per worker.
10286        unsafe { *out_addr.at(o) = v };
10287    }
10288}
10289
10290/// Two-input q8 row range via AVX2 (x86 mirror of `q8_range2_sdot`).
10291#[cfg(target_arch = "x86_64")]
10292#[allow(clippy::too_many_arguments)]
10293fn q8_range2_avx2(
10294    q: &[u8],
10295    row_scale: &[f32],
10296    a1: &SplitAct,
10297    a2: &SplitAct,
10298    cols: usize,
10299    p1: SendMut,
10300    p2: SendMut,
10301    start: usize,
10302    end: usize,
10303) {
10304    for o in start..end {
10305        let row = &q[o * cols..(o + 1) * cols];
10306        // SAFETY: disjoint row ranges per worker.
10307        unsafe {
10308            *p1.at(o) = row_dot_avx2(row, a1) * row_scale[o];
10309            *p2.at(o) = row_dot_avx2(row, a2) * row_scale[o];
10310        }
10311    }
10312}
10313
10314// ───────────────────── A8W8 SDOT path (port of vmfcore, ×1.78 decode) ─────────────────────
10315
10316/// ARMv8.6 i8mm (smmla): 32 int8 MACs per instruction vs sdot's 16 —
10317/// yet MEASURED 2.4× SLOWER than the blocked sdot on Apple silicon
10318/// (108 vs 264 GF/s): the on-the-fly vcombine packing and the two-
10319/// accumulator dependency chain swamp the MAC advantage, and Apple's
10320/// four SIMD pipes already keep sdot fed. OPT-IN (CMF_I8MM=1) for
10321/// field trials on Cortex-A710/X-class parts with two pipes, where the
10322/// balance may differ; a pre-interleaved weight layout (repack infra)
10323/// is the known path if it ever earns its keep.
10324#[cfg(target_arch = "aarch64")]
10325fn i8mm_enabled() -> bool {
10326    static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
10327    *ON.get_or_init(|| {
10328        std::env::var("CMF_I8MM").map(|v| v == "1").unwrap_or(false)
10329            && std::arch::is_aarch64_feature_detected!("i8mm")
10330    })
10331}
10332
10333/// SDOT enabled? Default ON when the CPU has ARMv8.2 dotprod;
10334/// `CMF_SDOT=0` disables (falls back to i8×f32 NEON).
10335/// (On non-ARM release builds only the test tolerance switch calls it.)
10336#[cfg_attr(not(target_arch = "aarch64"), allow(dead_code))]
10337fn sdot_enabled() -> bool {
10338    if FLOAT_ACTIVATIONS.get() {
10339        return false;
10340    }
10341    use std::sync::OnceLock;
10342    static ON: OnceLock<bool> = OnceLock::new();
10343    *ON.get_or_init(|| {
10344        let want = std::env::var("CMF_SDOT").map(|v| v != "0").unwrap_or(true);
10345        if !want {
10346            return false;
10347        }
10348
10349        #[cfg(target_arch = "aarch64")]
10350        {
10351            if std::arch::is_aarch64_feature_detected!("dotprod") {
10352                return true;
10353            }
10354            #[cfg(target_os = "android")]
10355            {
10356                if let Ok(cpuinfo) = std::fs::read_to_string("/proc/cpuinfo") {
10357                    if cpuinfo.lines().any(|l| {
10358                        (l.starts_with("Features") || l.starts_with("features"))
10359                            && l.contains("asimddp")
10360                    }) {
10361                        return true;
10362                    }
10363                }
10364            }
10365            false
10366        }
10367        #[cfg(not(target_arch = "aarch64"))]
10368        {
10369            false
10370        }
10371    })
10372}
10373
10374/// Two-field activation split (≡ vmfcore `q8_split_prep`): outlier
10375/// channels (>8·rms) are computed exactly in f32; the bulk (outliers
10376/// zeroed → clean absmax) goes through int8 SDOT. Computed ONCE per
10377/// matvec, shared by all rows/workers.
10378struct SplitAct {
10379    xq: Vec<i8>,
10380    sx: f32,
10381    outliers: Vec<(usize, f32)>,
10382    /// Σ xq — the VNNI bias-trick correction (`(w+128)·x` sums need
10383    /// `−128·Σx`); one i32 per split, computed once per matvec.
10384    #[cfg_attr(not(target_arch = "x86_64"), allow(dead_code))]
10385    xsum: i32,
10386}
10387
10388thread_local! {
10389    /// Recycled xq buffers: split_act runs for every matvec (~200/token)
10390    /// and its hidden-size allocation was steady-state heap churn.
10391    static XQ_FREE: std::cell::RefCell<Vec<Vec<i8>>> =
10392        const { std::cell::RefCell::new(Vec::new()) };
10393}
10394
10395impl Drop for SplitAct {
10396    fn drop(&mut self) {
10397        let buf = std::mem::take(&mut self.xq);
10398        if buf.capacity() > 0 {
10399            XQ_FREE.with(|f| {
10400                let mut f = f.borrow_mut();
10401                if f.len() < 16 {
10402                    f.push(buf);
10403                }
10404            });
10405        }
10406    }
10407}
10408
10409thread_local! {
10410    /// One scratch row per WORKER, kept for the life of the thread.
10411    ///
10412    /// The kernels take a row of group scales per dispatch, and a fresh
10413    /// `vec![0f32; gpr]` inside the closure is one allocation per worker per
10414    /// dispatch — on the release checkpoint about six thousand a token, a
10415    /// quarter of everything the benchmark counts.
10416    static KROW: UnsafeCell<Vec<f32>> = const { UnsafeCell::new(Vec::new()) };
10417    /// Two scratch rows for a gate/up pair.  Q2TP and Q4TP fused SwiGLU
10418    /// decode each projection's ladder independently, but both ladders have
10419    /// the same lifetime; retaining them per worker removes two heap trips
10420    /// from every dispatched expert pair.
10421    static KROWS: UnsafeCell<[Vec<f32>; 2]> =
10422        const { UnsafeCell::new([Vec::new(), Vec::new()]) };
10423}
10424
10425/// Borrow `n` floats of the calling worker's scratch. Nothing inside a
10426/// kernel body borrows it again, which is what keeps the RefCell honest.
10427#[inline]
10428fn with_krow<R>(n: usize, f: impl FnOnce(&mut [f32]) -> R) -> R {
10429    KROW.with(|s| {
10430        // SAFETY: `KROW` is thread-local and this function does not recurse;
10431        // each worker has exclusive access to its own scratch row.
10432        let b = unsafe { &mut *s.get() };
10433        if b.len() < n {
10434            b.resize(n, 0.0);
10435        }
10436        f(&mut b[..n])
10437    })
10438}
10439
10440/// `t.round().clamp(-127.0, 127.0) as i8`, bit for bit, without the libm
10441/// call. On baseline x86-64 (no SSE4.1 `roundps`) `f32::round` is a
10442/// function call per element, and split_act runs it over every hidden
10443/// state before every matvec: measured 27 us a call on a 2048-wide
10444/// activation on an EPYC 7763 — 5.4 ms of a 55 ms decode token, all of
10445/// it on the caller's thread while thirty workers wait. Clamping first is
10446/// equivalent (round is monotonic and ±127 are integers), and after the
10447/// clamp `t - trunc(t)` is exact, so the half-away-from-zero decision is
10448/// the one `round` makes. NaN clamps to NaN and converts to 0, as before.
10449/// The loop vectorizes (cvttps2dq + compare/select).
10450#[inline(always)]
10451fn q8_round(t: f32) -> i8 {
10452    let t = t.clamp(-127.0, 127.0);
10453    let i = t as i32;
10454    let f = t - i as f32;
10455    let r = if f >= 0.5 {
10456        i + 1
10457    } else if f <= -0.5 {
10458        i - 1
10459    } else {
10460        i
10461    };
10462    r as i8
10463}
10464
10465#[inline]
10466fn with_krows<R>(n: usize, f: impl FnOnce(&mut [f32], &mut [f32]) -> R) -> R {
10467    KROWS.with(|s| {
10468        // SAFETY: KROWS is thread-local and no kernel body recursively calls
10469        // this helper; each participant owns its two vectors exclusively.
10470        let b = unsafe { &mut *s.get() };
10471        for row in b.iter_mut() {
10472            if row.len() < n {
10473                row.resize(n, 0.0);
10474            }
10475        }
10476        let (a, b) = b.split_at_mut(1);
10477        f(&mut a[0][..n], &mut b[0][..n])
10478    })
10479}
10480
10481#[inline]
10482fn silu_mul_limited(mut gate: f32, mut up: f32, limit: f32) -> f32 {
10483    if limit > 0.0 {
10484        up = up.clamp(-limit, limit);
10485        gate = gate.min(limit);
10486    }
10487    gate / (1.0 + (-gate).exp()) * up
10488}
10489
10490fn split_act(x: &[f32]) -> SplitAct {
10491    let _prof = crate::cpuprof::time(crate::cpuprof::Slot::SplitAct);
10492    let n = x.len();
10493    let rms = (x.iter().map(|&v| (v * v) as f64).sum::<f64>() / n.max(1) as f64).sqrt() as f32;
10494    let thr = 8.0 * rms;
10495    // One pass: collect outliers and the bulk absmax (outliers excluded —
10496    // identical to the old zero-then-fold over a copied buffer, minus the
10497    // full-vector copy).
10498    let mut outliers: Vec<(usize, f32)> = Vec::new();
10499    let mut amax = 0f32;
10500    for (j, &v) in x.iter().enumerate() {
10501        let a = v.abs();
10502        if a > thr {
10503            outliers.push((j, v));
10504        } else if a > amax {
10505            amax = a;
10506        }
10507    }
10508    let sx = if amax > 0.0 { amax / 127.0 } else { 1.0 };
10509    let inv = 1.0 / sx;
10510    let mut xq = XQ_FREE.with(|f| f.borrow_mut().pop()).unwrap_or_default();
10511    xq.clear();
10512    xq.reserve(n);
10513    if outliers.is_empty() {
10514        xq.extend(
10515            x.iter()
10516                .map(|&v| q8_round(v * inv)),
10517        );
10518    } else {
10519        // Outlier slots quantize to 0 (their exact term is added later).
10520        xq.extend(x.iter().map(|&v| {
10521            if v.abs() > thr {
10522                0
10523            } else {
10524                q8_round(v * inv)
10525            }
10526        }));
10527    }
10528    let xsum = xq.iter().map(|&v| v as i32).sum();
10529    SplitAct {
10530        xq,
10531        sx,
10532        outliers,
10533        xsum,
10534    }
10535}
10536
10537fn split_act_q8_2f(x: &[f32], col: &[f32]) -> SplitAct {
10538    let _prof = crate::cpuprof::time(crate::cpuprof::Slot::SplitAct);
10539    let n = x.len();
10540    let rms = (x
10541        .iter()
10542        .zip(col)
10543        .map(|(&a, &c)| {
10544            let v = a * c;
10545            (v * v) as f64
10546        })
10547        .sum::<f64>()
10548        / n.max(1) as f64)
10549        .sqrt() as f32;
10550    let thr = 8.0 * rms;
10551
10552    let mut outliers = Vec::new();
10553    let mut amax = 0f32;
10554    for (j, (&a, &c)) in x.iter().zip(col).enumerate() {
10555        let v = a * c;
10556        let s = v.abs();
10557        if s > thr {
10558            outliers.push((j, v));
10559        } else if s > amax {
10560            amax = s;
10561        }
10562    }
10563
10564    let sx = if amax > 0.0 { amax / 127.0 } else { 1.0 };
10565    let inv = 1.0 / sx;
10566    let mut xq = XQ_FREE.with(|f| f.borrow_mut().pop()).unwrap_or_default();
10567    xq.clear();
10568    xq.reserve(n);
10569    if outliers.is_empty() {
10570        xq.extend(
10571            x.iter()
10572                .zip(col)
10573                .map(|(&a, &c)| q8_round((a * c) * inv)),
10574        );
10575    } else {
10576        xq.extend(x.iter().zip(col).map(|(&a, &c)| {
10577            let v = a * c;
10578            if v.abs() > thr {
10579                0
10580            } else {
10581                q8_round(v * inv)
10582            }
10583        }));
10584    }
10585    let xsum = xq.iter().map(|&v| v as i32).sum();
10586    SplitAct {
10587        xq,
10588        sx,
10589        outliers,
10590        xsum,
10591    }
10592}
10593
10594/// int8(weight)·int8(activation) → i32 via `sdot` (inline asm — the
10595/// vdotq intrinsic is unstable; port of vmfcore `dot_i8_sdot`).
10596#[cfg(target_arch = "aarch64")]
10597#[target_feature(enable = "neon,dotprod")]
10598unsafe fn dot_i8_sdot(w: &[u8], xq: &[i8]) -> i32 {
10599    // SAFETY: callers uphold slice-length contracts (see call sites).
10600    unsafe {
10601        use core::arch::aarch64::*;
10602        use core::arch::asm;
10603        let wp = w.as_ptr() as *const i8;
10604        let n = w.len();
10605        let (mut a0, mut a1, mut a2, mut a3) = (
10606            vdupq_n_s32(0),
10607            vdupq_n_s32(0),
10608            vdupq_n_s32(0),
10609            vdupq_n_s32(0),
10610        );
10611        let mut i = 0;
10612        while i + 64 <= n {
10613            let (w0, x0) = (vld1q_s8(wp.add(i)), vld1q_s8(xq.as_ptr().add(i)));
10614            let (w1, x1) = (vld1q_s8(wp.add(i + 16)), vld1q_s8(xq.as_ptr().add(i + 16)));
10615            let (w2, x2) = (vld1q_s8(wp.add(i + 32)), vld1q_s8(xq.as_ptr().add(i + 32)));
10616            let (w3, x3) = (vld1q_s8(wp.add(i + 48)), vld1q_s8(xq.as_ptr().add(i + 48)));
10617            asm!(
10618                "sdot {a0:v}.4s, {w0:v}.16b, {x0:v}.16b",
10619                "sdot {a1:v}.4s, {w1:v}.16b, {x1:v}.16b",
10620                "sdot {a2:v}.4s, {w2:v}.16b, {x2:v}.16b",
10621                "sdot {a3:v}.4s, {w3:v}.16b, {x3:v}.16b",
10622                a0 = inout(vreg) a0, a1 = inout(vreg) a1, a2 = inout(vreg) a2, a3 = inout(vreg) a3,
10623                w0 = in(vreg) w0, x0 = in(vreg) x0, w1 = in(vreg) w1, x1 = in(vreg) x1,
10624                w2 = in(vreg) w2, x2 = in(vreg) x2, w3 = in(vreg) w3, x3 = in(vreg) x3,
10625                options(pure, nomem, nostack),
10626            );
10627            i += 64;
10628        }
10629        while i + 16 <= n {
10630            let (wv, xv) = (vld1q_s8(wp.add(i)), vld1q_s8(xq.as_ptr().add(i)));
10631            asm!("sdot {a:v}.4s, {w:v}.16b, {x:v}.16b",
10632                 a = inout(vreg) a0, w = in(vreg) wv, x = in(vreg) xv, options(pure, nomem, nostack));
10633            i += 16;
10634        }
10635        let mut s = vaddvq_s32(vaddq_s32(vaddq_s32(a0, a1), vaddq_s32(a2, a3)));
10636        while i < n {
10637            s += (*wp.add(i)) as i32 * xq[i] as i32;
10638            i += 1;
10639        }
10640        s
10641    }
10642}
10643
10644/// Row-blocked SDOT: 4 output rows per pass — the activation chunk is
10645/// loaded once and reused, 4 independent accumulators hide sdot latency
10646/// (port of vmfcore `dot_i8_sdot_4rows`).
10647#[cfg(target_arch = "aarch64")]
10648#[target_feature(enable = "neon,dotprod")]
10649unsafe fn dot_i8_sdot_4rows(w0: &[u8], w1: &[u8], w2: &[u8], w3: &[u8], xq: &[i8]) -> [i32; 4] {
10650    // SAFETY: callers uphold slice-length contracts (see call sites).
10651    unsafe {
10652        use core::arch::aarch64::*;
10653        use core::arch::asm;
10654        let n = xq.len();
10655        let px = xq.as_ptr();
10656        let (p0, p1, p2, p3) = (
10657            w0.as_ptr() as *const i8,
10658            w1.as_ptr() as *const i8,
10659            w2.as_ptr() as *const i8,
10660            w3.as_ptr() as *const i8,
10661        );
10662        let (mut a0, mut a1, mut a2, mut a3) = (
10663            vdupq_n_s32(0),
10664            vdupq_n_s32(0),
10665            vdupq_n_s32(0),
10666            vdupq_n_s32(0),
10667        );
10668        let mut i = 0;
10669        while i + 16 <= n {
10670            let x = vld1q_s8(px.add(i));
10671            let v0 = vld1q_s8(p0.add(i));
10672            let v1 = vld1q_s8(p1.add(i));
10673            let v2 = vld1q_s8(p2.add(i));
10674            let v3 = vld1q_s8(p3.add(i));
10675            asm!(
10676                "sdot {a0:v}.4s, {v0:v}.16b, {x:v}.16b",
10677                "sdot {a1:v}.4s, {v1:v}.16b, {x:v}.16b",
10678                "sdot {a2:v}.4s, {v2:v}.16b, {x:v}.16b",
10679                "sdot {a3:v}.4s, {v3:v}.16b, {x:v}.16b",
10680                a0 = inout(vreg) a0, a1 = inout(vreg) a1, a2 = inout(vreg) a2, a3 = inout(vreg) a3,
10681                v0 = in(vreg) v0, v1 = in(vreg) v1, v2 = in(vreg) v2, v3 = in(vreg) v3, x = in(vreg) x,
10682                options(pure, nomem, nostack),
10683            );
10684            i += 16;
10685        }
10686        let mut r = [
10687            vaddvq_s32(a0),
10688            vaddvq_s32(a1),
10689            vaddvq_s32(a2),
10690            vaddvq_s32(a3),
10691        ];
10692        while i < n {
10693            let xi = *px.add(i) as i32;
10694            r[0] += (*p0.add(i)) as i32 * xi;
10695            r[1] += (*p1.add(i)) as i32 * xi;
10696            r[2] += (*p2.add(i)) as i32 * xi;
10697            r[3] += (*p3.add(i)) as i32 * xi;
10698            i += 1;
10699        }
10700        r
10701    }
10702}
10703
10704/// 4 interleaved rows in one pass: the repacked group is [r0[c], r1[c],
10705/// r2[c], r3[c]] per 16-byte chunk, so each iteration reads ONE 64-byte
10706/// line plus the shared activation chunk — a single sequential weight
10707/// stream per worker. Per-row accumulation is the same one-accumulator
10708/// scheme as `dot_i8_sdot_4rows`; integer sums are exact, so outputs
10709/// are bit-identical to the mmap-layout kernel.
10710#[cfg(target_arch = "aarch64")]
10711#[target_feature(enable = "neon,dotprod")]
10712unsafe fn dot_i8_sdot_4rows_il(g: &[u8], xq: &[i8]) -> [i32; 4] {
10713    // SAFETY: callers uphold slice-length contracts (g.len() == 4·n,
10714    // n % 16 == 0 — guaranteed by the repack gate).
10715    unsafe {
10716        use core::arch::aarch64::*;
10717        use core::arch::asm;
10718        let n = xq.len();
10719        let px = xq.as_ptr();
10720        let pg = g.as_ptr() as *const i8;
10721        let (mut a0, mut a1, mut a2, mut a3) = (
10722            vdupq_n_s32(0),
10723            vdupq_n_s32(0),
10724            vdupq_n_s32(0),
10725            vdupq_n_s32(0),
10726        );
10727        let mut i = 0;
10728        while i + 16 <= n {
10729            let x = vld1q_s8(px.add(i));
10730            let base = pg.add(4 * i);
10731            let v0 = vld1q_s8(base);
10732            let v1 = vld1q_s8(base.add(16));
10733            let v2 = vld1q_s8(base.add(32));
10734            let v3 = vld1q_s8(base.add(48));
10735            asm!(
10736                "sdot {a0:v}.4s, {v0:v}.16b, {x:v}.16b",
10737                "sdot {a1:v}.4s, {v1:v}.16b, {x:v}.16b",
10738                "sdot {a2:v}.4s, {v2:v}.16b, {x:v}.16b",
10739                "sdot {a3:v}.4s, {v3:v}.16b, {x:v}.16b",
10740                a0 = inout(vreg) a0, a1 = inout(vreg) a1, a2 = inout(vreg) a2, a3 = inout(vreg) a3,
10741                v0 = in(vreg) v0, v1 = in(vreg) v1, v2 = in(vreg) v2, v3 = in(vreg) v3, x = in(vreg) x,
10742                options(pure, nomem, nostack),
10743            );
10744            i += 16;
10745        }
10746        [
10747            vaddvq_s32(a0),
10748            vaddvq_s32(a1),
10749            vaddvq_s32(a2),
10750            vaddvq_s32(a3),
10751        ]
10752    }
10753}
10754
10755/// One q8 row range via SDOT (4-row blocks + tail) — the body of
10756/// `qmatvec`'s hot loop, extracted so multi-matrix jobs can drive the
10757/// SAME kernel for several tensors under one pool dispatch. `rep` — the
10758/// load-time interleaved repack (empty = mmap layout only); rows outside
10759/// full 4-row groups always come from the mmap layout.
10760#[cfg(target_arch = "aarch64")]
10761fn q8_range_sdot(
10762    q: &[u8],
10763    rep: &[u8],
10764    row_scale: &[f32],
10765    act: &SplitAct,
10766    cols: usize,
10767    out_addr: SendMut,
10768    start: usize,
10769    end: usize,
10770) {
10771    let mut o = start;
10772    // Leading rows to the group boundary (repack path only): the pool
10773    // splits row ranges arbitrarily, groups are absolute.
10774    if !rep.is_empty() {
10775        while o < end && o % 4 != 0 {
10776            let v = row_dot_sdot(&q[o * cols..(o + 1) * cols], act) * row_scale[o];
10777            unsafe { *out_addr.at(o) = v };
10778            o += 1;
10779        }
10780    }
10781    while o + 4 <= end {
10782        let r = if rep.is_empty() {
10783            unsafe {
10784                dot_i8_sdot_4rows(
10785                    &q[o * cols..(o + 1) * cols],
10786                    &q[(o + 1) * cols..(o + 2) * cols],
10787                    &q[(o + 2) * cols..(o + 3) * cols],
10788                    &q[(o + 3) * cols..(o + 4) * cols],
10789                    &act.xq,
10790                )
10791            }
10792        } else {
10793            unsafe { dot_i8_sdot_4rows_il(&rep[o * cols..(o + 4) * cols], &act.xq) }
10794        };
10795        for k in 0..4 {
10796            let mut acc = r[k] as f32 * act.sx;
10797            for &(j, xv) in &act.outliers {
10798                acc += (q[(o + k) * cols + j] as i8) as f32 * xv;
10799            }
10800            // SAFETY: disjoint row ranges per worker.
10801            unsafe { *out_addr.at(o + k) = acc * row_scale[o + k] };
10802        }
10803        o += 4;
10804    }
10805    while o < end {
10806        let v = row_dot_sdot(&q[o * cols..(o + 1) * cols], act) * row_scale[o];
10807        unsafe { *out_addr.at(o) = v };
10808        o += 1;
10809    }
10810}
10811
10812/// Two-input q8 row range via SDOT — `qmatvec2`'s hot loop, extracted
10813/// for the fused pair multi-matrix job (`matvec2_many`).
10814#[cfg(target_arch = "aarch64")]
10815#[allow(clippy::too_many_arguments)]
10816fn q8_range2_sdot(
10817    q: &[u8],
10818    row_scale: &[f32],
10819    a1: &SplitAct,
10820    a2: &SplitAct,
10821    cols: usize,
10822    p1: SendMut,
10823    p2: SendMut,
10824    start: usize,
10825    end: usize,
10826) {
10827    for o in start..end {
10828        let row = &q[o * cols..(o + 1) * cols];
10829        // SAFETY: disjoint row ranges per worker.
10830        unsafe {
10831            *p1.at(o) = row_dot_sdot(row, a1) * row_scale[o];
10832            *p2.at(o) = row_dot_sdot(row, a2) * row_scale[o];
10833        }
10834    }
10835}
10836
10837/// Two-input q8 row range, f32 kernel (non-SDOT) — same extraction.
10838#[allow(clippy::too_many_arguments)]
10839fn q8_range2_f32(
10840    q: &[u8],
10841    row_scale: &[f32],
10842    x1: &[f32],
10843    x2: &[f32],
10844    cols: usize,
10845    p1: SendMut,
10846    p2: SendMut,
10847    start: usize,
10848    end: usize,
10849) {
10850    for o in start..end {
10851        let row = &q[o * cols..(o + 1) * cols];
10852        // SAFETY: disjoint row ranges per worker.
10853        unsafe {
10854            *p1.at(o) = dot_i8_f32(row, x1) * row_scale[o];
10855            *p2.at(o) = dot_i8_f32(row, x2) * row_scale[o];
10856        }
10857    }
10858}
10859
10860/// Scalar/NEON-f32 q8 row range (non-SDOT platforms) — same extraction.
10861fn q8_range_f32(
10862    q: &[u8],
10863    row_scale: &[f32],
10864    xs: &[f32],
10865    cols: usize,
10866    out_addr: SendMut,
10867    start: usize,
10868    end: usize,
10869) {
10870    for o in start..end {
10871        let v = dot_i8_f32(&q[o * cols..(o + 1) * cols], xs) * row_scale[o];
10872        // SAFETY: disjoint row ranges per worker.
10873        unsafe { *out_addr.at(o) = v };
10874    }
10875}
10876
10877/// One q8 row against a split activation, portable: the per-arch fast
10878/// dots where they exist, the exact scalar loop elsewhere. The scalar
10879/// arm is also the test oracle for both fast arms.
10880#[inline]
10881fn q8_row_dot(row: &[u8], act: &SplitAct) -> f32 {
10882    #[cfg(target_arch = "aarch64")]
10883    return row_dot_sdot(row, act);
10884    #[cfg(target_arch = "x86_64")]
10885    return row_dot_avx2(row, act);
10886    #[allow(unreachable_code)]
10887    q8_row_dot_scalar(row, act)
10888}
10889
10890#[allow(dead_code)]
10891fn q8_row_dot_scalar(row: &[u8], act: &SplitAct) -> f32 {
10892    let mut acc = 0i32;
10893    for (k, &b) in row.iter().enumerate() {
10894        acc += (b as i8) as i32 * act.xq[k] as i32;
10895    }
10896    let mut acc = acc as f32 * act.sx;
10897    for &(j, xv) in &act.outliers {
10898        acc += (row[j] as i8) as f32 * xv;
10899    }
10900    acc
10901}
10902
10903/// SDOT row dot with exact outlier correction:
10904/// `dot = sdot(w, xq)·sx + Σ_outl w[j]·x[j]` (then × row_scale by caller).
10905#[cfg(target_arch = "aarch64")]
10906#[inline]
10907fn row_dot_sdot(row: &[u8], act: &SplitAct) -> f32 {
10908    let mut acc = unsafe { dot_i8_sdot(row, &act.xq) } as f32 * act.sx;
10909    for &(j, xv) in &act.outliers {
10910        acc += (row[j] as i8) as f32 * xv;
10911    }
10912    acc
10913}
10914
10915/// One q4 row via SDOT: each 32-group's nibbles unpack to centered i8
10916/// (nib−8 ∈ [−8,7]), int8×int8 `sdot` against the pre-quantized
10917/// activation group, × the group's f16 scale. Returns Σ_g dot_g·s_g;
10918/// the caller multiplies by the activation scale and adds the exact
10919/// outlier terms (port of vmfcore `dot_q4_block_sdot`, +23% measured).
10920/// Nibble order matches the writer: element 2k = low nibble, 2k+1 = high
10921/// → zip(lo,hi) restores flat order.
10922#[cfg(target_arch = "aarch64")]
10923#[target_feature(enable = "neon,dotprod")]
10924unsafe fn dot_q4_row_sdot(packed: &[u8], scales: &[u8], g0: usize, gpr: usize, xq: &[i8]) -> f32 {
10925    // SAFETY: callers uphold slice-length contracts (16 packed bytes and
10926    // 2 scale bytes per group; xq.len() == gpr·GROUP_SIZE).
10927    unsafe {
10928        use core::arch::aarch64::*;
10929        use core::arch::asm;
10930        let lomask = vdupq_n_u8(0x0F);
10931        let eight = vdupq_n_s8(8);
10932        let mut acc = 0f32;
10933        for gi in 0..gpr {
10934            let g = g0 + gi;
10935            let s = f16_to_f32(u16::from_le_bytes([scales[g * 2], scales[g * 2 + 1]]));
10936            let b = vld1q_u8(packed.as_ptr().add(g * 16));
10937            let lo = vandq_u8(b, lomask);
10938            let hi = vshrq_n_u8::<4>(b);
10939            let e0 = vsubq_s8(vreinterpretq_s8_u8(vzip1q_u8(lo, hi)), eight);
10940            let e1 = vsubq_s8(vreinterpretq_s8_u8(vzip2q_u8(lo, hi)), eight);
10941            let x0 = vld1q_s8(xq.as_ptr().add(gi * GROUP_SIZE));
10942            let x1 = vld1q_s8(xq.as_ptr().add(gi * GROUP_SIZE + 16));
10943            let (mut a0, mut a1) = (vdupq_n_s32(0), vdupq_n_s32(0));
10944            asm!(
10945                "sdot {a0:v}.4s, {e0:v}.16b, {x0:v}.16b",
10946                "sdot {a1:v}.4s, {e1:v}.16b, {x1:v}.16b",
10947                a0 = inout(vreg) a0, a1 = inout(vreg) a1,
10948                e0 = in(vreg) e0, x0 = in(vreg) x0, e1 = in(vreg) e1, x1 = in(vreg) x1,
10949                options(pure, nomem, nostack),
10950            );
10951            acc += vaddvq_s32(vaddq_s32(a0, a1)) as f32 * s;
10952        }
10953        acc
10954    }
10955}
10956
10957/// Two-activation q4 row via SDOT: the nibble unpack (the expensive
10958/// part) happens ONCE per group; both pre-quantized activations are
10959/// dotted against the same centered i8 registers. Per-lane math matches
10960/// `dot_q4_row_sdot` exactly.
10961#[cfg(target_arch = "aarch64")]
10962#[target_feature(enable = "neon,dotprod")]
10963unsafe fn dot_q4_row_sdot2(
10964    packed: &[u8],
10965    scales: &[u8],
10966    g0: usize,
10967    gpr: usize,
10968    xq1: &[i8],
10969    xq2: &[i8],
10970) -> (f32, f32) {
10971    // SAFETY: callers uphold slice-length contracts (16 packed bytes and
10972    // 2 scale bytes per group; xq*.len() == gpr·GROUP_SIZE).
10973    unsafe {
10974        use core::arch::aarch64::*;
10975        use core::arch::asm;
10976        let lomask = vdupq_n_u8(0x0F);
10977        let eight = vdupq_n_s8(8);
10978        let (mut acc1, mut acc2) = (0f32, 0f32);
10979        for gi in 0..gpr {
10980            let g = g0 + gi;
10981            let s = f16_to_f32(u16::from_le_bytes([scales[g * 2], scales[g * 2 + 1]]));
10982            let b = vld1q_u8(packed.as_ptr().add(g * 16));
10983            let lo = vandq_u8(b, lomask);
10984            let hi = vshrq_n_u8::<4>(b);
10985            let e0 = vsubq_s8(vreinterpretq_s8_u8(vzip1q_u8(lo, hi)), eight);
10986            let e1 = vsubq_s8(vreinterpretq_s8_u8(vzip2q_u8(lo, hi)), eight);
10987            let x10 = vld1q_s8(xq1.as_ptr().add(gi * GROUP_SIZE));
10988            let x11 = vld1q_s8(xq1.as_ptr().add(gi * GROUP_SIZE + 16));
10989            let x20 = vld1q_s8(xq2.as_ptr().add(gi * GROUP_SIZE));
10990            let x21 = vld1q_s8(xq2.as_ptr().add(gi * GROUP_SIZE + 16));
10991            let (mut a0, mut a1, mut b0, mut b1) = (
10992                vdupq_n_s32(0),
10993                vdupq_n_s32(0),
10994                vdupq_n_s32(0),
10995                vdupq_n_s32(0),
10996            );
10997            asm!(
10998                "sdot {a0:v}.4s, {e0:v}.16b, {x10:v}.16b",
10999                "sdot {a1:v}.4s, {e1:v}.16b, {x11:v}.16b",
11000                "sdot {b0:v}.4s, {e0:v}.16b, {x20:v}.16b",
11001                "sdot {b1:v}.4s, {e1:v}.16b, {x21:v}.16b",
11002                a0 = inout(vreg) a0, a1 = inout(vreg) a1,
11003                b0 = inout(vreg) b0, b1 = inout(vreg) b1,
11004                e0 = in(vreg) e0, e1 = in(vreg) e1,
11005                x10 = in(vreg) x10, x11 = in(vreg) x11,
11006                x20 = in(vreg) x20, x21 = in(vreg) x21,
11007                options(pure, nomem, nostack),
11008            );
11009            acc1 += vaddvq_s32(vaddq_s32(a0, a1)) as f32 * s;
11010            acc2 += vaddvq_s32(vaddq_s32(b0, b1)) as f32 * s;
11011        }
11012        (acc1, acc2)
11013    }
11014}
11015
11016// ───────────────────── fused int8 kernels ─────────────────────
11017
11018/// `acc += w · row` where the row is centered i8 — NEON widen+fma on
11019/// aarch64, scalar elsewhere. The KV-cache q8 value path rides on this.
11020#[inline]
11021pub(crate) fn axpy_i8_f32(acc: &mut [f32], row: &[i8], w: f32) {
11022    #[cfg(target_arch = "aarch64")]
11023    unsafe {
11024        return axpy_i8_f32_neon(acc, row, w);
11025    }
11026    #[cfg(target_arch = "x86_64")]
11027    if avx2_enabled() {
11028        return unsafe { axpy_i8_f32_avx2(acc, row, w) };
11029    }
11030    #[allow(unreachable_code)]
11031    {
11032        for (a, &b) in acc.iter_mut().zip(row) {
11033            *a += w * b as f32;
11034        }
11035    }
11036}
11037
11038/// i8→f32 axpy via AVX2/FMA (x86 mirror of `axpy_i8_f32_neon`).
11039#[cfg(target_arch = "x86_64")]
11040#[target_feature(enable = "avx2,fma")]
11041unsafe fn axpy_i8_f32_avx2(acc: &mut [f32], row: &[i8], w: f32) {
11042    // SAFETY: callers uphold slice-length contracts (see call sites).
11043    unsafe {
11044        use core::arch::x86_64::*;
11045        let n = acc.len().min(row.len());
11046        let ap = acc.as_mut_ptr();
11047        let rp = row.as_ptr();
11048        let wv = _mm256_set1_ps(w);
11049        let mut j = 0usize;
11050        while j + 16 <= n {
11051            let rb = _mm_loadu_si128(rp.add(j) as *const __m128i);
11052            let lo = _mm256_cvtepi32_ps(_mm256_cvtepi8_epi32(rb));
11053            let hi = _mm256_cvtepi32_ps(_mm256_cvtepi8_epi32(_mm_srli_si128::<8>(rb)));
11054            let v0 = _mm256_fmadd_ps(wv, lo, _mm256_loadu_ps(ap.add(j)));
11055            let v1 = _mm256_fmadd_ps(wv, hi, _mm256_loadu_ps(ap.add(j + 8)));
11056            _mm256_storeu_ps(ap.add(j), v0);
11057            _mm256_storeu_ps(ap.add(j + 8), v1);
11058            j += 16;
11059        }
11060        while j < n {
11061            *ap.add(j) += w * (*rp.add(j)) as f32;
11062            j += 1;
11063        }
11064    }
11065}
11066
11067#[cfg(target_arch = "aarch64")]
11068#[target_feature(enable = "neon")]
11069unsafe fn axpy_i8_f32_neon(acc: &mut [f32], row: &[i8], w: f32) {
11070    // SAFETY: callers uphold slice-length contracts (see call sites).
11071    unsafe {
11072        use core::arch::aarch64::*;
11073        let n = acc.len().min(row.len());
11074        let ap = acc.as_mut_ptr();
11075        let rp = row.as_ptr();
11076        let wv = vdupq_n_f32(w);
11077        let mut j = 0usize;
11078        while j + 16 <= n {
11079            let rb = vld1q_s8(rp.add(j));
11080            let lo = vmovl_s8(vget_low_s8(rb));
11081            let hi = vmovl_s8(vget_high_s8(rb));
11082            for (off, half) in [(0, lo), (8, hi)] {
11083                let f0 = vcvtq_f32_s32(vmovl_s16(vget_low_s16(half)));
11084                let f1 = vcvtq_f32_s32(vmovl_s16(vget_high_s16(half)));
11085                let o = j + off;
11086                vst1q_f32(ap.add(o), vfmaq_f32(vld1q_f32(ap.add(o)), wv, f0));
11087                vst1q_f32(ap.add(o + 4), vfmaq_f32(vld1q_f32(ap.add(o + 4)), wv, f1));
11088            }
11089            j += 16;
11090        }
11091        while j < n {
11092            *ap.add(j) += w * (*rp.add(j)) as f32;
11093            j += 1;
11094        }
11095    }
11096}
11097
11098/// i8 row · f32 x. NEON on aarch64 (ported from vmfcore `dot_i8_f32_neon`,
11099/// ≈9× scalar), scalar elsewhere.
11100#[inline]
11101pub(crate) fn dot_i8_f32(w: &[u8], x: &[f32]) -> f32 {
11102    #[cfg(target_arch = "aarch64")]
11103    unsafe {
11104        return dot_i8_f32_neon(w, x);
11105    }
11106    #[cfg(target_arch = "x86_64")]
11107    if avx2_enabled() {
11108        return unsafe { dot_i8_f32_avx2(w, x) };
11109    }
11110    #[allow(unreachable_code)]
11111    {
11112        let mut sum = 0.0f32;
11113        for (j, &b) in w.iter().enumerate() {
11114            sum += (b as i8) as f32 * x[j];
11115        }
11116        sum
11117    }
11118}
11119
11120/// i8 row · (x ⊙ col_field) — the q8_2f row dot with the θ col-field
11121/// folded into the product (no prescaled copy of x). NEON on aarch64,
11122/// scalar elsewhere. Used by the active-neuron path `row_dot`.
11123#[inline]
11124fn dot_i8_col_f32(w: &[u8], x: &[f32], col: &[f32]) -> f32 {
11125    #[cfg(target_arch = "aarch64")]
11126    unsafe {
11127        return dot_i8_col_f32_neon(w, x, col);
11128    }
11129    #[allow(unreachable_code)]
11130    {
11131        let mut sum = 0.0f32;
11132        for (j, &b) in w.iter().enumerate() {
11133            sum += (b as i8) as f32 * x[j] * col[j];
11134        }
11135        sum
11136    }
11137}
11138
11139#[cfg(target_arch = "aarch64")]
11140#[target_feature(enable = "neon")]
11141unsafe fn dot_i8_col_f32_neon(w: &[u8], x: &[f32], col: &[f32]) -> f32 {
11142    // SAFETY: callers uphold slice-length contracts (see call sites).
11143    unsafe {
11144        use core::arch::aarch64::*;
11145        let n = x.len();
11146        let wp = w.as_ptr() as *const i8;
11147        let xp = x.as_ptr();
11148        let cp = col.as_ptr();
11149        let (mut a0, mut a1, mut a2, mut a3) = (
11150            vdupq_n_f32(0.0),
11151            vdupq_n_f32(0.0),
11152            vdupq_n_f32(0.0),
11153            vdupq_n_f32(0.0),
11154        );
11155        let mut j = 0usize;
11156        while j + 16 <= n {
11157            let wb = vld1q_s8(wp.add(j));
11158            let lo = vmovl_s8(vget_low_s8(wb));
11159            let hi = vmovl_s8(vget_high_s8(wb));
11160            let w0 = vcvtq_f32_s32(vmovl_s16(vget_low_s16(lo)));
11161            let w1 = vcvtq_f32_s32(vmovl_s16(vget_high_s16(lo)));
11162            let w2 = vcvtq_f32_s32(vmovl_s16(vget_low_s16(hi)));
11163            let w3 = vcvtq_f32_s32(vmovl_s16(vget_high_s16(hi)));
11164            a0 = vfmaq_f32(
11165                a0,
11166                w0,
11167                vmulq_f32(vld1q_f32(xp.add(j)), vld1q_f32(cp.add(j))),
11168            );
11169            a1 = vfmaq_f32(
11170                a1,
11171                w1,
11172                vmulq_f32(vld1q_f32(xp.add(j + 4)), vld1q_f32(cp.add(j + 4))),
11173            );
11174            a2 = vfmaq_f32(
11175                a2,
11176                w2,
11177                vmulq_f32(vld1q_f32(xp.add(j + 8)), vld1q_f32(cp.add(j + 8))),
11178            );
11179            a3 = vfmaq_f32(
11180                a3,
11181                w3,
11182                vmulq_f32(vld1q_f32(xp.add(j + 12)), vld1q_f32(cp.add(j + 12))),
11183            );
11184            j += 16;
11185        }
11186        let mut sum = vaddvq_f32(vaddq_f32(vaddq_f32(a0, a1), vaddq_f32(a2, a3)));
11187        while j < n {
11188            sum += (*wp.add(j)) as f32 * *xp.add(j) * *cp.add(j);
11189            j += 1;
11190        }
11191        sum
11192    }
11193}
11194
11195#[cfg(target_arch = "aarch64")]
11196#[target_feature(enable = "neon")]
11197unsafe fn dot_i8_f32_neon(w: &[u8], x: &[f32]) -> f32 {
11198    // SAFETY: callers uphold slice-length contracts (see call sites).
11199    unsafe {
11200        use core::arch::aarch64::*;
11201        let n = x.len();
11202        let wp = w.as_ptr() as *const i8;
11203        let xp = x.as_ptr();
11204        let (mut a0, mut a1, mut a2, mut a3) = (
11205            vdupq_n_f32(0.0),
11206            vdupq_n_f32(0.0),
11207            vdupq_n_f32(0.0),
11208            vdupq_n_f32(0.0),
11209        );
11210        let mut j = 0usize;
11211        while j + 16 <= n {
11212            let wb = vld1q_s8(wp.add(j));
11213            let lo = vmovl_s8(vget_low_s8(wb));
11214            let hi = vmovl_s8(vget_high_s8(wb));
11215            let w0 = vcvtq_f32_s32(vmovl_s16(vget_low_s16(lo)));
11216            let w1 = vcvtq_f32_s32(vmovl_s16(vget_high_s16(lo)));
11217            let w2 = vcvtq_f32_s32(vmovl_s16(vget_low_s16(hi)));
11218            let w3 = vcvtq_f32_s32(vmovl_s16(vget_high_s16(hi)));
11219            a0 = vfmaq_f32(a0, w0, vld1q_f32(xp.add(j)));
11220            a1 = vfmaq_f32(a1, w1, vld1q_f32(xp.add(j + 4)));
11221            a2 = vfmaq_f32(a2, w2, vld1q_f32(xp.add(j + 8)));
11222            a3 = vfmaq_f32(a3, w3, vld1q_f32(xp.add(j + 12)));
11223            j += 16;
11224        }
11225        let mut sum = vaddvq_f32(vaddq_f32(vaddq_f32(a0, a1), vaddq_f32(a2, a3)));
11226        while j < n {
11227            sum += (*wp.add(j)) as f32 * *xp.add(j);
11228            j += 1;
11229        }
11230        sum
11231    }
11232}
11233
11234#[allow(clippy::too_many_arguments)]
11235fn qmatvec(
11236    q: &[u8],
11237    rep: &[u8],
11238    row_scale: &[f32],
11239    x: &[f32],
11240    col_field: &[f32],
11241    dtype: TensorDtype,
11242    rows: usize,
11243    cols: usize,
11244    out: &mut [f32],
11245    pool: Option<&Pool>,
11246) {
11247    debug_assert_eq!(out.len(), rows);
11248    #[cfg(not(target_arch = "aarch64"))]
11249    let _ = rep;
11250
11251    #[cfg(target_arch = "aarch64")]
11252    if sdot_enabled() {
11253        let act = if dtype == TensorDtype::Q8_2f {
11254            split_act_q8_2f(x, col_field)
11255        } else {
11256            split_act(x)
11257        };
11258        let out_addr = SendMut(out.as_mut_ptr());
11259        let run_range = |start: usize, end: usize| {
11260            q8_range_sdot(q, rep, row_scale, &act, cols, out_addr, start, end)
11261        };
11262        match pool {
11263            Some(pool) if rows >= 256 => pool.run_rows(rows, &run_range),
11264            _ => run_range(0, rows),
11265        }
11266        return;
11267    }
11268    // x86 A8W8 via AVX2 maddubs — same quantized-activation contract as
11269    // the SDOT path (CMF_AVX2=0 keeps the exact i8×f32 loop).
11270    #[cfg(target_arch = "x86_64")]
11271    if avx2_a8w8_enabled() {
11272        let act = if dtype == TensorDtype::Q8_2f {
11273            split_act_q8_2f(x, col_field)
11274        } else {
11275            split_act(x)
11276        };
11277        let out_addr = SendMut(out.as_mut_ptr());
11278        let run_range = |start: usize, end: usize| {
11279            q8_range_avx2(q, row_scale, &act, cols, out_addr, start, end)
11280        };
11281        match pool {
11282            Some(pool) if rows >= 256 => pool.run_rows(rows, &run_range),
11283            _ => run_range(0, rows),
11284        }
11285        return;
11286    }
11287
11288    prescale_with(x, col_field, dtype, 1, |xs| {
11289        let out_addr = SendMut(out.as_mut_ptr());
11290        let run_range = move |start: usize, end: usize| {
11291            for o in start..end {
11292                let v = dot_i8_f32(&q[o * cols..(o + 1) * cols], xs) * row_scale[o];
11293                // SAFETY: disjoint row ranges per worker.
11294                unsafe { *out_addr.at(o) = v };
11295            }
11296        };
11297        match pool {
11298            Some(pool) if rows >= 256 => pool.run_rows(rows, &run_range),
11299            _ => run_range(0, rows),
11300        }
11301    });
11302}
11303
11304#[allow(clippy::too_many_arguments)]
11305fn qmatvec2(
11306    q: &[u8],
11307    row_scale: &[f32],
11308    x1: &[f32],
11309    x2: &[f32],
11310    col_field: &[f32],
11311    dtype: TensorDtype,
11312    rows: usize,
11313    cols: usize,
11314    o1: &mut [f32],
11315    o2: &mut [f32],
11316    pool: Option<&Pool>,
11317) {
11318    #[cfg(target_arch = "aarch64")]
11319    if sdot_enabled() {
11320        let a1s = if dtype == TensorDtype::Q8_2f {
11321            split_act_q8_2f(x1, col_field)
11322        } else {
11323            split_act(x1)
11324        };
11325        let a2s = if dtype == TensorDtype::Q8_2f {
11326            split_act_q8_2f(x2, col_field)
11327        } else {
11328            split_act(x2)
11329        };
11330        let p1 = SendMut(o1.as_mut_ptr());
11331        let p2 = SendMut(o2.as_mut_ptr());
11332        let run_range = |start: usize, end: usize| {
11333            q8_range2_sdot(q, row_scale, &a1s, &a2s, cols, p1, p2, start, end)
11334        };
11335        match pool {
11336            Some(pool) if rows >= 256 => pool.run_rows(rows, &run_range),
11337            _ => run_range(0, rows),
11338        }
11339        return;
11340    }
11341    #[cfg(target_arch = "x86_64")]
11342    if avx2_a8w8_enabled() {
11343        let a1s = if dtype == TensorDtype::Q8_2f {
11344            split_act_q8_2f(x1, col_field)
11345        } else {
11346            split_act(x1)
11347        };
11348        let a2s = if dtype == TensorDtype::Q8_2f {
11349            split_act_q8_2f(x2, col_field)
11350        } else {
11351            split_act(x2)
11352        };
11353        let p1 = SendMut(o1.as_mut_ptr());
11354        let p2 = SendMut(o2.as_mut_ptr());
11355        let run_range = |start: usize, end: usize| {
11356            q8_range2_avx2(q, row_scale, &a1s, &a2s, cols, p1, p2, start, end)
11357        };
11358        match pool {
11359            Some(pool) if rows >= 256 => pool.run_rows(rows, &run_range),
11360            _ => run_range(0, rows),
11361        }
11362        return;
11363    }
11364
11365    prescale_with(x1, col_field, dtype, 1, |x1s| {
11366        prescale_with(x2, col_field, dtype, 2, |x2s| {
11367            let p1 = SendMut(o1.as_mut_ptr());
11368            let p2 = SendMut(o2.as_mut_ptr());
11369            let run_range = move |start: usize, end: usize| {
11370                for o in start..end {
11371                    let row = &q[o * cols..(o + 1) * cols];
11372                    let s1 = dot_i8_f32(row, x1s) * row_scale[o];
11373                    let s2 = dot_i8_f32(row, x2s) * row_scale[o];
11374                    // SAFETY: disjoint row ranges per worker.
11375                    unsafe {
11376                        *p1.at(o) = s1;
11377                        *p2.at(o) = s2;
11378                    }
11379                }
11380            };
11381            match pool {
11382                Some(pool) if rows >= 256 => pool.run_rows(rows, &run_range),
11383                _ => run_range(0, rows),
11384            }
11385        });
11386    });
11387}
11388
11389#[derive(Clone, Copy)]
11390struct SendMut(*mut f32);
11391unsafe impl Send for SendMut {}
11392unsafe impl Sync for SendMut {}
11393
11394impl SendMut {
11395    #[inline]
11396    fn at(self, i: usize) -> *mut f32 {
11397        unsafe { self.0.add(i) }
11398    }
11399}
11400
11401#[cfg(test)]
11402mod tests {
11403    /// `q8_round` must be `round().clamp(±127) as i8` bit for bit: every
11404    /// half-integer, their neighbours one ulp either side, the clamp
11405    /// boundary, huge values, infinities and NaN, plus a dense sweep.
11406    #[test]
11407    fn q8_round_is_round_clamp() {
11408        let reference = |t: f32| t.round().clamp(-127.0, 127.0) as i8;
11409        let mut probe = vec![
11410            0.0f32,
11411            -0.0,
11412            f32::NAN,
11413            f32::INFINITY,
11414            f32::NEG_INFINITY,
11415            f32::MAX,
11416            f32::MIN,
11417            1e30,
11418            -1e30,
11419            f32::MIN_POSITIVE,
11420            -f32::MIN_POSITIVE,
11421        ];
11422        for k in -300i32..=300 {
11423            let h = k as f32 * 0.5;
11424            let up = f32::from_bits(h.to_bits() + 1);
11425            let down = f32::from_bits(h.to_bits().wrapping_sub(1));
11426            for t in [h, up, down] {
11427                probe.push(t);
11428                probe.push(-t);
11429            }
11430        }
11431        let mut t = -140.0f32;
11432        while t < 140.0 {
11433            probe.push(t);
11434            t += 0.000_731;
11435        }
11436        for t in probe {
11437            assert_eq!(q8_round(t), reference(t), "t = {t:e} ({:#x})", t.to_bits());
11438        }
11439    }
11440
11441    use super::*;
11442
11443    /// The f32 `matmat` walks the batch four positions a row at a time; every
11444    /// output must still be the scalar-order sum `matvec` gives, bit for bit
11445    /// (a 7-row batch covers the four-wide block and the remainder).
11446    #[test]
11447    fn f32_matmat_equals_matvec_bitwise() {
11448        let (rows, cols, b) = (5usize, 37usize, 7usize);
11449        let w: Vec<f32> = (0..rows * cols).map(|i| ((i * 7919 % 113) as f32 - 56.0) / 37.0).collect();
11450        let xs: Vec<f32> = (0..b * cols).map(|i| ((i * 104_729 % 97) as f32 - 48.0) / 29.0).collect();
11451        let t = QTensor::from_f32(w, rows, cols);
11452        let mut out = vec![0f32; b * rows];
11453        t.matmat(&xs, b, &mut out, None);
11454        for bi in 0..b {
11455            let mut one = vec![0f32; rows];
11456            t.matvec(&xs[bi * cols..(bi + 1) * cols], &mut one, None);
11457            for o in 0..rows {
11458                assert_eq!(out[bi * rows + o].to_bits(), one[o].to_bits(), "bi {bi} row {o}");
11459            }
11460        }
11461    }
11462
11463    #[test]
11464    fn q2tp_i8_dot_matches_exact_on_grid() {
11465        // On-grid activations (±1 → sx=1/127, xq=±127 dequantizes
11466        // exactly, no outliers) must make the integer path agree with
11467        // the exact scalar walk to f32 rounding.
11468        let (rows, cols) = (5, 64);
11469        let gpr = cols / GROUP_SIZE;
11470        // Synthetic codes plane + a flat ladder: scales_into is not under
11471        // test here, so drive dot_q2tp_row_i8 / q2tp_row_exact directly
11472        // with hand-made scales.
11473        let chunks: Vec<u8> = (0..rows * gpr * Q2TP_CHUNK)
11474            .map(|i| (i as u32).wrapping_mul(2654435761) as u8)
11475            .collect();
11476        let scales: Vec<f32> = (0..gpr).map(|g| 0.5 + g as f32 * 0.25).collect();
11477        let x: Vec<f32> = (0..cols)
11478            .map(|i| if i % 3 == 0 { -1.0 } else { 1.0 })
11479            .collect();
11480        let act = split_act(&x);
11481        assert!(
11482            act.outliers.is_empty(),
11483            "on-grid input must have no outliers"
11484        );
11485        let gsum = q1_group_sums(&act.xq, gpr);
11486        for r in 0..rows {
11487            let exact = q2tp_row_exact(&chunks, r, gpr, &x, &scales);
11488            let fast = dot_q2tp_row_i8(&chunks, r, gpr, &act.xq, &gsum, &scales) * act.sx;
11489            assert!(
11490                (exact - fast).abs() <= exact.abs() * 1e-5 + 1e-5,
11491                "row {r}: exact {exact} vs i8 {fast}"
11492            );
11493        }
11494    }
11495
11496    #[cfg(target_arch = "x86_64")]
11497    #[test]
11498    fn q2tp_avx2_dot_matches_scalar_for_random_patterns() {
11499        // Compare the release AVX2 integer dot against the scalar oracle over
11500        // arbitrary packed bytes/activation signs.  This guards the exact
11501        // table-load path used after rejecting a faster-looking decoder whose
11502        // full-checkpoint greedy output drifted.
11503        if !std::arch::is_x86_feature_detected!("avx2") {
11504            return;
11505        }
11506        let mut seed = 0x9e3779b9u32;
11507        let mut next = || {
11508            seed = seed.wrapping_mul(1664525).wrapping_add(1013904223);
11509            seed
11510        };
11511        for _ in 0..20_000 {
11512            let mut ch = [0u8; Q2TP_CHUNK];
11513            let mut x = [0i8; GROUP_SIZE];
11514            for b in &mut ch {
11515                *b = next() as u8;
11516            }
11517            for v in &mut x {
11518                *v = (next() >> 24) as i8;
11519            }
11520            let mut reference = 0i32;
11521            for (k, &b) in ch.iter().enumerate() {
11522                reference += (b & 3) as i32 * x[k * 4] as i32;
11523                reference += ((b >> 2) & 3) as i32 * x[k * 4 + 1] as i32;
11524                reference += ((b >> 4) & 3) as i32 * x[k * 4 + 2] as i32;
11525                reference += ((b >> 6) & 3) as i32 * x[k * 4 + 3] as i32;
11526            }
11527            // SAFETY: guarded by the runtime AVX2 feature check and fixed
11528            // 8-byte/32-byte slice lengths above.
11529            let got = unsafe { q2tp_code_dot_avx2(&ch, &x) };
11530            assert_eq!(got, reference, "packed q2 lane mismatch");
11531        }
11532    }
11533
11534    #[test]
11535    fn q2tp_affine_fuses_half_scale_correction_without_changing_raw_decode() {
11536        let (rows, cols) = (1usize, GROUP_SIZE);
11537        let mut bytes = vec![0u8; Q2TP_CHUNK + 4 + 1];
11538        // Repeating symbols 0,1,2,0 at unit scale.  q2tp's raw B is
11539        // (c-1.5), while the affine Prism operator is (c-1.0).
11540        bytes[..Q2TP_CHUNK].fill(0x24); // codes 0,1,2,0 in LSB-first order
11541        bytes[Q2TP_CHUNK..Q2TP_CHUNK + 2].copy_from_slice(&0u16.to_le_bytes());
11542        bytes[Q2TP_CHUNK + 2..Q2TP_CHUNK + 4].copy_from_slice(&0u16.to_le_bytes());
11543        bytes[Q2TP_CHUNK + 4] = 1; // dtype16 rung 1 = 1.0
11544        let x = vec![1.0f32; cols];
11545        let mut raw = vec![0.0f32; rows];
11546        let mut affine = vec![0.0f32; rows];
11547        q2tp_matvec_for_test(&bytes, &x, rows, cols, &mut raw);
11548        q2tp_affine_matvec_for_test(&bytes, &x, rows, cols, &mut affine);
11549        assert_eq!(raw, vec![-24.0]);
11550        assert_eq!(affine, vec![-8.0]);
11551        assert!((affine[0] - (raw[0] + 0.5 * cols as f32)).abs() < 1e-6);
11552    }
11553
11554    #[test]
11555    fn q8_row_dot_fast_matches_scalar() {
11556        // The per-arch fast dot must agree with the exact scalar oracle
11557        // (same contract the fused q8 FFN arm rides on).
11558        let cols = 96;
11559        let row: Vec<u8> = (0..cols)
11560            .map(|i| ((i * 37 % 251) - 125) as i8 as u8)
11561            .collect();
11562        let x: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.13).sin()).collect();
11563        let act = split_act(&x);
11564        let fast = q8_row_dot(&row, &act);
11565        let scalar = q8_row_dot_scalar(&row, &act);
11566        assert!(
11567            (fast - scalar).abs() <= scalar.abs() * 1e-5 + 1e-5,
11568            "fast {fast} vs scalar {scalar}"
11569        );
11570    }
11571
11572    #[test]
11573    fn f32_matvec_matches_matvec_rows_bitexact() {
11574        let (rows, cols) = (300, 40);
11575        let w: Vec<f32> = (0..rows * cols).map(|i| (i as f32 * 0.017).sin()).collect();
11576        let x: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.05).cos()).collect();
11577        let qt = QTensor::from_f32(w.clone(), rows, cols);
11578
11579        let mut a = vec![0.0f32; rows];
11580        matvec_rows(None, &w, &x, &mut a);
11581        let mut b = vec![0.0f32; rows];
11582        qt.matvec(&x, &mut b, None);
11583        assert_eq!(a, b);
11584    }
11585
11586    #[test]
11587    fn sdot_kernel_exact_on_grid() {
11588        // Activations already on the i8 grid (±1 with amax=1 → sx=1/127,
11589        // xq=±127 dequantizes EXACTLY) → the SDOT path must match the
11590        // exact f32 dot to float rounding. This isolates kernel
11591        // correctness from quantization noise.
11592        eprintln!("sdot_enabled = {}", sdot_enabled());
11593        let (rows, cols) = (9, 80); // odd rows → exercises 4-row + tail
11594        let w: Vec<u8> = (0..rows * cols)
11595            .map(|i| (((i * 37) % 251) as i32 - 125) as i8 as u8)
11596            .collect();
11597        let scales: Vec<f32> = (0..rows).map(|o| 0.005 + o as f32 * 0.001).collect();
11598        let x: Vec<f32> = (0..cols)
11599            .map(|i| match i % 3 {
11600                0 => 1.0,
11601                1 => -1.0,
11602                _ => 0.0,
11603            })
11604            .collect();
11605        let mut a = vec![0.0f32; rows];
11606        qmatvec(
11607            &w,
11608            &[],
11609            &scales,
11610            &x,
11611            &[],
11612            TensorDtype::Q8Row,
11613            rows,
11614            cols,
11615            &mut a,
11616            None,
11617        );
11618        for o in 0..rows {
11619            let mut acc = 0.0f32;
11620            for j in 0..cols {
11621                acc += (w[o * cols + j] as i8) as f32 * x[j];
11622            }
11623            let expect = acc * scales[o];
11624            assert!(
11625                (a[o] - expect).abs() < 1e-3 * expect.abs().max(1e-3),
11626                "row {o}: {} vs {expect}",
11627                a[o]
11628            );
11629        }
11630    }
11631
11632    #[test]
11633    fn q1_tbl_fast_path_matches_reference() {
11634        // gpr = 8 exercises the TBL pair-load fast loop, and the LAST
11635        // row's final 4-tile window trips the 4B-overread guard (the
11636        // payload ends exactly at the last tile) — both paths must
11637        // agree with the dequant reference.
11638        let (rows, cols) = (5, 256);
11639        let gpr = cols / GROUP_SIZE;
11640        let mut bytes = Vec::new();
11641        for t in 0..rows * gpr {
11642            let s = 0.007 + (t % 11) as f32 * 0.004;
11643            bytes.extend_from_slice(&cortiq_core::quant::f32_to_f16(s).to_le_bytes());
11644            for j in 0..4 {
11645                bytes.push(((t * 53 + j * 89 + 7) % 249) as u8);
11646            }
11647        }
11648        let x: Vec<f32> = (0..cols)
11649            .map(|i| if (i * 5) % 7 < 3 { 1.0 } else { -1.0 })
11650            .collect();
11651        let mut w = vec![0.0f32; rows * cols];
11652        cortiq_core::quant::dequant_q1(&bytes, &mut w);
11653        let mut got = vec![0.0f32; rows];
11654        q1_matvec(&bytes, &x, rows, cols, &mut got, None);
11655        for o in 0..rows {
11656            let expect: f32 = (0..cols).map(|j| w[o * cols + j] * x[j]).sum();
11657            assert!(
11658                (got[o] - expect).abs() < 1e-3 * expect.abs().max(1e-3),
11659                "row {o}: {} vs {expect}",
11660                got[o]
11661            );
11662        }
11663        // Blocked 1×4 batch (b=5: one quad + remainder) must equal the
11664        // single-matvec path bit-for-bit.
11665        let b = 5usize;
11666        let mut xs_all = Vec::new();
11667        for bi in 0..b {
11668            xs_all.extend(x.iter().map(|v| if bi % 2 == 0 { *v } else { -*v }));
11669        }
11670        let mut mm = vec![0.0f32; b * rows];
11671        q1_matmat(&bytes, &xs_all, b, rows, cols, &mut mm, None);
11672        for bi in 0..b {
11673            let mut single = vec![0.0f32; rows];
11674            q1_matvec(
11675                &bytes,
11676                &xs_all[bi * cols..(bi + 1) * cols],
11677                rows,
11678                cols,
11679                &mut single,
11680                None,
11681            );
11682            assert_eq!(&mm[bi * rows..(bi + 1) * rows], &single[..], "stream {bi}");
11683        }
11684    }
11685
11686    #[test]
11687    fn q1_kernels_match_exact_reference() {
11688        // Synthetic q1 payload: 6-byte tiles [f16 scale][4B bits].
11689        let (rows, cols) = (7, 96);
11690        let gpr = cols / GROUP_SIZE;
11691        let mut bytes = Vec::new();
11692        for t in 0..rows * gpr {
11693            let s = 0.01 + (t % 13) as f32 * 0.003;
11694            bytes.extend_from_slice(&cortiq_core::quant::f32_to_f16(s).to_le_bytes());
11695            for j in 0..4 {
11696                bytes.push(((t * 31 + j * 97) % 251) as u8);
11697            }
11698        }
11699        // On-grid activations (±1, amax 1) → the SDOT path is exact.
11700        let x: Vec<f32> = (0..cols)
11701            .map(|i| if i % 3 == 0 { 1.0 } else { -1.0 })
11702            .collect();
11703        // Reference through the core dequant.
11704        let mut w = vec![0.0f32; rows * cols];
11705        cortiq_core::quant::dequant_q1(&bytes, &mut w);
11706        let mut expect = vec![0.0f32; rows];
11707        for o in 0..rows {
11708            expect[o] = (0..cols).map(|j| w[o * cols + j] * x[j]).sum();
11709        }
11710        let mut got = vec![0.0f32; rows];
11711        q1_matvec(&bytes, &x, rows, cols, &mut got, None);
11712        for o in 0..rows {
11713            assert!(
11714                (got[o] - expect[o]).abs() < 1e-3 * expect[o].abs().max(1e-3),
11715                "row {o}: {} vs {}",
11716                got[o],
11717                expect[o]
11718            );
11719        }
11720        // Pair and batch paths agree with the single path.
11721        let x2: Vec<f32> = x.iter().map(|v| -v).collect();
11722        let (mut a1, mut a2) = (vec![0.0f32; rows], vec![0.0f32; rows]);
11723        q1_matvec2(&bytes, &x, &x2, rows, cols, &mut a1, &mut a2, None);
11724        assert_eq!(a1, got);
11725        let mut xs = x.clone();
11726        xs.extend_from_slice(&x2);
11727        let mut mm = vec![0.0f32; 2 * rows];
11728        q1_matmat(&bytes, &xs, 2, rows, cols, &mut mm, None);
11729        assert_eq!(&mm[..rows], got.as_slice());
11730        assert_eq!(&mm[rows..], a2.as_slice());
11731    }
11732
11733    #[test]
11734    fn repack_is_bit_identical() {
11735        // The interleaved-repack kernel must produce EXACTLY the same
11736        // bits as the mmap-layout kernel: integer accumulation is order-
11737        // exact, the f32 epilogue is identical. Odd rows exercise the
11738        // tail; direct range calls exercise unaligned pool splits.
11739        let (rows, cols) = (267, 96); // 66 groups + 3 tail rows, cols % 16 == 0
11740        let w: Vec<u8> = (0..rows * cols)
11741            .map(|i| (((i * 89) % 253) as i32 - 126) as i8 as u8)
11742            .collect();
11743        let scales: Vec<f32> = (0..rows).map(|o| 0.003 + o as f32 * 0.0007).collect();
11744        let x: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.37).sin() * 2.0).collect();
11745        let rep = q8_repack_layout(&w, rows, cols);
11746        // Group interleave round-trips.
11747        for g in 0..rows / 4 {
11748            for c in 0..cols / 16 {
11749                for lane in 0..4 {
11750                    assert_eq!(
11751                        &rep[g * 4 * cols + c * 64 + lane * 16
11752                            ..g * 4 * cols + c * 64 + lane * 16 + 16],
11753                        &w[(g * 4 + lane) * cols + c * 16..(g * 4 + lane) * cols + c * 16 + 16],
11754                    );
11755                }
11756            }
11757        }
11758        let mut a = vec![0.0f32; rows];
11759        qmatvec(
11760            &w,
11761            &[],
11762            &scales,
11763            &x,
11764            &[],
11765            TensorDtype::Q8Row,
11766            rows,
11767            cols,
11768            &mut a,
11769            None,
11770        );
11771        let mut b = vec![0.0f32; rows];
11772        qmatvec(
11773            &w,
11774            &rep,
11775            &scales,
11776            &x,
11777            &[],
11778            TensorDtype::Q8Row,
11779            rows,
11780            cols,
11781            &mut b,
11782            None,
11783        );
11784        assert_eq!(a, b, "full-range repack output diverged");
11785
11786        #[cfg(target_arch = "aarch64")]
11787        if sdot_enabled() {
11788            // Unaligned range split (pool workers get arbitrary bounds).
11789            let act = split_act(&x);
11790            let mut c1 = vec![0.0f32; rows];
11791            let mut c2 = vec![0.0f32; rows];
11792            q8_range_sdot(
11793                &w,
11794                &[],
11795                &scales,
11796                &act,
11797                cols,
11798                SendMut(c1.as_mut_ptr()),
11799                3,
11800                rows - 2,
11801            );
11802            q8_range_sdot(
11803                &w,
11804                &rep,
11805                &scales,
11806                &act,
11807                cols,
11808                SendMut(c2.as_mut_ptr()),
11809                3,
11810                rows - 2,
11811            );
11812            assert_eq!(c1, c2, "unaligned-range repack output diverged");
11813        }
11814    }
11815
11816    #[test]
11817    fn sdot_a8w8_noise_is_bounded() {
11818        // Off-grid activations: A8 quantization noise must stay small in
11819        // relative L2 over the whole output (realistic accuracy contract;
11820        // vmfcore measured argmax-identical decode on real models).
11821        let (rows, cols) = (16, 512);
11822        let w: Vec<u8> = (0..rows * cols)
11823            .map(|i| (((i * 37) % 251) as i32 - 125) as i8 as u8)
11824            .collect();
11825        let scales = vec![0.01f32; rows];
11826        let x: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.21).sin()).collect();
11827        let mut a = vec![0.0f32; rows];
11828        qmatvec(
11829            &w,
11830            &[],
11831            &scales,
11832            &x,
11833            &[],
11834            TensorDtype::Q8Row,
11835            rows,
11836            cols,
11837            &mut a,
11838            None,
11839        );
11840        let (mut num, mut den) = (0f64, 0f64);
11841        for o in 0..rows {
11842            let mut acc = 0.0f32;
11843            for j in 0..cols {
11844                acc += (w[o * cols + j] as i8) as f32 * x[j];
11845            }
11846            let expect = acc * scales[o];
11847            num += ((a[o] - expect) as f64).powi(2);
11848            den += (expect as f64).powi(2);
11849        }
11850        let rel = (num / den.max(1e-12)).sqrt();
11851        assert!(rel < 0.05, "A8W8 relative L2 error too high: {rel}");
11852    }
11853
11854    #[test]
11855    fn i8_dot_neon_matches_scalar() {
11856        let n = 100;
11857        let w: Vec<u8> = (0..n).map(|i| ((i * 37 + 11) % 251) as u8).collect();
11858        let x: Vec<f32> = (0..n).map(|i| (i as f32 * 0.13).sin()).collect();
11859        let mut scalar = 0.0f32;
11860        for j in 0..n {
11861            scalar += (w[j] as i8) as f32 * x[j];
11862        }
11863        let fast = dot_i8_f32(&w, &x);
11864        assert!((scalar - fast).abs() < 1e-3 * scalar.abs().max(1.0));
11865    }
11866
11867    /// Fused vbit matvec must match full dequant_vbit + dense matvec.
11868    #[test]
11869    fn vbitmatvec_matches_full_dequant() {
11870        let (rows, cols) = (6, 64);
11871        let ng = cols / GROUP_SIZE;
11872        // Hand-craft: bits per row, f16 scales, packed rows.
11873        let bits: Vec<u8> = vec![3, 4, 5, 6, 8, 4];
11874        let mut bytes = bits.clone();
11875        for g in 0..rows * ng {
11876            let s = 0.02 + 0.001 * g as f32;
11877            bytes.extend_from_slice(&cortiq_core::quant::f32_to_f16(s).to_le_bytes());
11878        }
11879        for r in 0..rows {
11880            let b = bits[r] as usize;
11881            let (mut acc, mut nb) = (0u64, 0usize);
11882            let mut rowbytes = Vec::new();
11883            for i in 0..cols {
11884                let v = ((i * 7 + r * 13) % (1 << b)) as u64;
11885                acc = (acc << b) | v;
11886                nb += b;
11887                while nb >= 8 {
11888                    nb -= 8;
11889                    rowbytes.push(((acc >> nb) & 0xFF) as u8);
11890                }
11891            }
11892            if nb > 0 {
11893                rowbytes.push(((acc << (8 - nb)) & 0xFF) as u8);
11894            }
11895            bytes.extend_from_slice(&rowbytes);
11896        }
11897        let x: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.19).sin()).collect();
11898
11899        let mut reference = vec![0f32; rows * cols];
11900        cortiq_core::quant::dequant_vbit(&bytes, rows, cols, &mut reference).unwrap();
11901        let mut expect = vec![0f32; rows];
11902        for r in 0..rows {
11903            expect[r] = reference[r * cols..(r + 1) * cols]
11904                .iter()
11905                .zip(&x)
11906                .map(|(w, xv)| w * xv)
11907                .sum();
11908        }
11909        let mut got = vec![0f32; rows];
11910        let offsets = vbit_row_offsets(&bytes, rows, cols);
11911        vbitmatvec(&bytes, &offsets, &x, rows, cols, &mut got, None);
11912        // SDOT path quantizes activations to i8 (A8W8): bounded noise,
11913        // same contract as q8 (exact path is pinned by CMF_SDOT=0 in
11914        // the golden-parity gate).
11915        let tol = if a8w8_enabled() { 6e-2 } else { 1e-4 };
11916        let scale = expect.iter().fold(0f32, |m, v| m.max(v.abs())).max(1e-6);
11917        for r in 0..rows {
11918            assert!(
11919                (got[r] - expect[r]).abs() < tol * scale,
11920                "row {r}: {} vs {}",
11921                got[r],
11922                expect[r]
11923            );
11924        }
11925    }
11926
11927    /// Fused q4 matvec must match the reference full-dequant + dense
11928    /// matvec bit-for-bit in structure (same f32 math, group order).
11929    /// vbit matmat: the blocked 1×4 leg must match the per-row path
11930    /// (paired env toggle; larger shape so both code paths engage).
11931    #[test]
11932    #[cfg(target_arch = "x86_64")]
11933    fn vbit_matmat_blocked_matches_per_row() {
11934        let (rows, cols, b) = (64usize, 128usize, 9usize);
11935        let ng = cols / GROUP_SIZE;
11936        let bits: Vec<u8> = (0..rows).map(|r| [3u8, 4, 5, 6][r % 4]).collect();
11937        let mut bytes = bits.clone();
11938        for g in 0..rows * ng {
11939            let sc = 0.02 + 0.0005 * g as f32;
11940            bytes.extend_from_slice(&cortiq_core::quant::f32_to_f16(sc).to_le_bytes());
11941        }
11942        for r in 0..rows {
11943            let bw = bits[r] as usize;
11944            let (mut acc, mut nb) = (0u64, 0usize);
11945            let mut rowbytes = Vec::new();
11946            for i in 0..cols {
11947                let v = ((i * 7 + r * 13) % (1 << bw)) as u64;
11948                acc = (acc << bw) | v;
11949                nb += bw;
11950                while nb >= 8 {
11951                    nb -= 8;
11952                    rowbytes.push(((acc >> nb) & 0xFF) as u8);
11953                }
11954            }
11955            if nb > 0 {
11956                rowbytes.push(((acc << (8 - nb)) & 0xFF) as u8);
11957            }
11958            bytes.extend_from_slice(&rowbytes);
11959        }
11960        let x: Vec<f32> = (0..b * cols)
11961            .map(|i| ((i * 13 + 7) % 97) as f32 / 97.0 - 0.5)
11962            .collect();
11963        let offsets = vbit_row_offsets(&bytes, rows, cols);
11964        let mut y_a = vec![0f32; b * rows];
11965        let mut y_b = vec![0f32; b * rows];
11966        unsafe { std::env::set_var("CMF_X86_BLOCKED", "1") };
11967        vbitmatmat(&bytes, &offsets, &x, b, rows, cols, &mut y_a, None);
11968        unsafe { std::env::set_var("CMF_X86_BLOCKED", "0") };
11969        vbitmatmat(&bytes, &offsets, &x, b, rows, cols, &mut y_b, None);
11970        unsafe { std::env::remove_var("CMF_X86_BLOCKED") };
11971        let max_d = y_a
11972            .iter()
11973            .zip(&y_b)
11974            .map(|(p, q)| (p - q).abs())
11975            .fold(0.0f32, f32::max);
11976        assert!(max_d < 1e-4, "vbit blocked ≠ per-row: max|Δ| = {max_d}");
11977    }
11978
11979    /// q4t blocked 1×4 (SDOT on ARM, AVX2 on x86) must equal the
11980    /// per-row path exactly: same nibble unpack, same group order,
11981    /// same f32 accumulation — batch == matvec bit-for-bit. b=9 covers
11982    /// two full 1×4 blocks plus a remainder through the single-row
11983    /// kernel. (Both paths produce identical output, so the shared
11984    /// CMF_X86_BLOCKED env var racing with other tests cannot flip
11985    /// the verdict — worst case both sides take the same path.)
11986    #[test]
11987    fn q4t_matmat_blocked_matches_per_row() {
11988        let (rows, cols, b) = (16usize, 64usize, 9usize);
11989        let gpr = cols / GROUP_SIZE;
11990        let mut bytes = vec![0u8; rows * gpr * Q4_TILE];
11991        for r in 0..rows {
11992            for g in 0..gpr {
11993                let t = (r * gpr + g) * Q4_TILE;
11994                let sc = 0.02 + 0.001 * (r * gpr + g) as f32;
11995                bytes[t..t + 2].copy_from_slice(&cortiq_core::quant::f32_to_f16(sc).to_le_bytes());
11996                for k in 0..16 {
11997                    bytes[t + 2 + k] = ((r * 31 + g * 7 + k * 13) % 251) as u8;
11998                }
11999            }
12000        }
12001        let x: Vec<f32> = (0..b * cols)
12002            .map(|i| ((i * 13 + 7) % 97) as f32 / 97.0 - 0.5)
12003            .collect();
12004        let mut y_blk = vec![0f32; b * rows];
12005        let mut y_row = vec![0f32; b * rows];
12006        unsafe { std::env::set_var("CMF_X86_BLOCKED", "1") };
12007        q4t_matmat(&bytes, &x, b, rows, cols, &mut y_blk, None);
12008        unsafe { std::env::set_var("CMF_X86_BLOCKED", "0") };
12009        q4t_matmat(&bytes, &x, b, rows, cols, &mut y_row, None);
12010        unsafe { std::env::remove_var("CMF_X86_BLOCKED") };
12011        assert_eq!(y_blk, y_row, "q4t blocked 1x4 ≠ per-row");
12012    }
12013
12014    /// The wide-batch Accelerate arm of q4t_matmat vs a brute-force
12015    /// f32 dequant matmul: both are f32 GEMMs, so only reduction
12016    /// order differs — tight tolerance.
12017    /// A synthetic q4tp payload: random nibbles plus a per-row ladder whose
12018    /// span varies row to row, so the codes actually exercise the full 0..31
12019    /// range rather than clustering on one rung.
12020    fn synth_q4tp(rows: usize, cols: usize) -> Vec<u8> {
12021        use cortiq_core::quant::{f32_to_f16, q4tp_code_stride, q4tp_put_code};
12022        let gpr = cols / GROUP_SIZE;
12023        let stride = q4tp_code_stride(gpr);
12024        let (params_off, codes_off, _) = q4tp_sections(rows, cols);
12025        let mut b = vec![0u8; codes_off + rows * stride];
12026        for r in 0..rows {
12027            for g in 0..gpr {
12028                let t = (r * gpr + g) * Q4TP_NIB;
12029                for k in 0..16 {
12030                    b[t + k] = ((r * 31 + g * 7 + k * 13) % 251) as u8;
12031                }
12032            }
12033            let lo = -6.0 - 0.03 * (r % 17) as f32;
12034            let step = 0.01 + 0.004 * (r % 11) as f32;
12035            let p = params_off + r * 4;
12036            b[p..p + 2].copy_from_slice(&f32_to_f16(lo).to_le_bytes());
12037            b[p + 2..p + 4].copy_from_slice(&f32_to_f16(step).to_le_bytes());
12038            let crow = &mut b[codes_off + r * stride..codes_off + (r + 1) * stride];
12039            for g in 0..gpr {
12040                q4tp_put_code(crow, g, (r * 5 + g * 3) % 32);
12041            }
12042        }
12043        b
12044    }
12045
12046    /// The same weights re-expressed as q4_tiled, so the proven kernel can
12047    /// be the reference: each tile stores the ladder scale its code selects.
12048    /// Only the f16 rounding of that scale separates the two payloads.
12049    fn q4tp_as_q4t(bytes: &[u8], rows: usize, cols: usize) -> Vec<u8> {
12050        let gpr = cols / GROUP_SIZE;
12051        let v = Q4tpView::new(bytes, rows, cols);
12052        let mut out = vec![0u8; rows * gpr * Q4_TILE];
12053        let mut sc = vec![0f32; gpr];
12054        for r in 0..rows {
12055            v.scales_into(r, gpr, &mut sc);
12056            for g in 0..gpr {
12057                let t = (r * gpr + g) * Q4_TILE;
12058                let s = sc[g];
12059                out[t..t + 2].copy_from_slice(&cortiq_core::quant::f32_to_f16(s).to_le_bytes());
12060                let src = (r * gpr + g) * Q4TP_NIB;
12061                out[t + 2..t + Q4_TILE].copy_from_slice(&v.nib[src..src + Q4TP_NIB]);
12062            }
12063        }
12064        out
12065    }
12066
12067    /// The exact (`CMF_SDOT=0`) path must reproduce `dequant_q4tp` to f32
12068    /// rounding — that scalar routine is the format's definition, and the
12069    /// kernels re-derive the scale from the ladder independently. Call the
12070    /// row kernel directly: `matmat` picks the int8 arm when a8w8 is on,
12071    /// so routing through it would test the other path by accident.
12072    #[test]
12073    fn q4tp_exact_path_matches_dequant_reference() {
12074        let (rows, cols) = (256usize, 512usize);
12075        let gpr = cols / GROUP_SIZE;
12076        let bytes = synth_q4tp(rows, cols);
12077        let mut w = vec![0f32; rows * cols];
12078        cortiq_core::quant::dequant_q4tp(&bytes, rows, cols, &mut w);
12079
12080        let x: Vec<f32> = (0..cols)
12081            .map(|i| ((i * 13 + 7) % 97) as f32 / 97.0 - 0.5)
12082            .collect();
12083        let v = Q4tpView::new(&bytes, rows, cols);
12084        let mut sc = vec![0f32; gpr];
12085        for r in 0..rows {
12086            v.scales_into(r, gpr, &mut sc);
12087            let got = q4tp_row_exact(v.nib, r, gpr, &x, &sc);
12088            let want: f32 = (0..cols).map(|c| w[r * cols + c] * x[c]).sum();
12089            // These dot products cancel down to ~1e-3 from terms of ~5e-2, so
12090            // the meaningful yardstick is the summed magnitude, not the result:
12091            // against the result any reordering of a 512-term f32 sum "fails".
12092            let mag: f32 = (0..cols).map(|c| (w[r * cols + c] * x[c]).abs()).sum();
12093            assert!(
12094                (got - want).abs() <= 1e-5 * mag,
12095                "row {r}: kernel {got} vs dequant {want}"
12096            );
12097        }
12098    }
12099
12100    /// The int8 (a8w8) path can't be checked against an f32 reference — the
12101    /// activation quantization dominates. Check it against the q4t kernel it
12102    /// was ported from instead, on payloads holding the same weights: that
12103    /// isolates exactly what the port could break (16 B stride, ladder
12104    /// lookup, nibble unpack) from what it deliberately shares.
12105    #[test]
12106    fn q4tp_matvec_matches_the_q4t_kernel_it_was_ported_from() {
12107        let (rows, cols) = (256usize, 512usize);
12108        let bytes = synth_q4tp(rows, cols);
12109        let twin = q4tp_as_q4t(&bytes, rows, cols);
12110        let x: Vec<f32> = (0..cols)
12111            .map(|i| ((i * 13 + 7) % 97) as f32 / 97.0 - 0.5)
12112            .collect();
12113
12114        let mut got = vec![0f32; rows];
12115        q4tp_matvec(&bytes, &x, rows, cols, &mut got, None);
12116        let mut want = vec![0f32; rows];
12117        q4t_matvec(&twin, &x, rows, cols, &mut want, None);
12118
12119        // Scale is f16 in the twin and f32 here, so allow that rounding on
12120        // top of the summed magnitude (same cancellation argument as above).
12121        let mut w = vec![0f32; rows * cols];
12122        cortiq_core::quant::dequant_q4tp(&bytes, rows, cols, &mut w);
12123        for r in 0..rows {
12124            let mag: f32 = (0..cols).map(|c| (w[r * cols + c] * x[c]).abs()).sum();
12125            assert!(
12126                (got[r] - want[r]).abs() <= 1e-3 * mag,
12127                "row {r}: q4tp {} vs q4t {}",
12128                got[r],
12129                want[r]
12130            );
12131        }
12132    }
12133
12134    /// `matmat` carries three arms (Accelerate, blocked int8 1x4, scalar).
12135    /// Batch 5 crosses the blocked kernel's stride, so this exercises the
12136    /// 1x4 path AND its scalar tail in one run — the blocked kernel is new
12137    /// code and its four accumulators are exactly what tends to go wrong.
12138    #[test]
12139    fn q4tp_matmat_matches_the_q4t_kernel_it_was_ported_from() {
12140        let (rows, cols, b) = (256usize, 512usize, 5usize);
12141        let bytes = synth_q4tp(rows, cols);
12142        let twin = q4tp_as_q4t(&bytes, rows, cols);
12143        let xs: Vec<f32> = (0..b * cols)
12144            .map(|i| ((i * 29 + 11) % 89) as f32 / 89.0 - 0.5)
12145            .collect();
12146
12147        let mut got = vec![0f32; b * rows];
12148        q4tp_matmat(&bytes, &xs, b, rows, cols, &mut got, None);
12149        let mut want = vec![0f32; b * rows];
12150        q4t_matmat(&twin, &xs, b, rows, cols, &mut want, None);
12151
12152        let mut w = vec![0f32; rows * cols];
12153        cortiq_core::quant::dequant_q4tp(&bytes, rows, cols, &mut w);
12154        for t in 0..b {
12155            for r in 0..rows {
12156                let mag: f32 = (0..cols)
12157                    .map(|c| (w[r * cols + c] * xs[t * cols + c]).abs())
12158                    .sum();
12159                let (g, wa) = (got[t * rows + r], want[t * rows + r]);
12160                assert!(
12161                    (g - wa).abs() <= 1e-3 * mag,
12162                    "batch {t} row {r}: q4tp {g} vs q4t {wa}"
12163                );
12164            }
12165        }
12166    }
12167
12168    #[test]
12169    fn q4tp_matvec2_matches_the_single_stream_kernel() {
12170        let (rows, cols) = (128usize, 256usize);
12171        let gpr = cols / GROUP_SIZE;
12172        let bytes = synth_q4tp(rows, cols);
12173        let xs: Vec<f32> = (0..2 * cols)
12174            .map(|i| ((i * 29 + 11) % 89) as f32 / 89.0 - 0.5)
12175            .collect();
12176
12177        let (mut o1, mut o2) = (vec![0f32; rows], vec![0f32; rows]);
12178        q4tp_matvec2(
12179            &bytes,
12180            &xs[..cols],
12181            &xs[cols..],
12182            rows,
12183            cols,
12184            &mut o1,
12185            &mut o2,
12186            None,
12187        );
12188
12189        // matvec2 takes the exact path for both streams, so the single-row
12190        // kernel is an exact reference — no tolerance for path differences.
12191        let v = Q4tpView::new(&bytes, rows, cols);
12192        let mut sc = vec![0f32; gpr];
12193        for r in 0..rows {
12194            v.scales_into(r, gpr, &mut sc);
12195            assert_eq!(o1[r], q4tp_row_exact(v.nib, r, gpr, &xs[..cols], &sc));
12196            assert_eq!(o2[r], q4tp_row_exact(v.nib, r, gpr, &xs[cols..], &sc));
12197        }
12198    }
12199
12200    /// q4tp must not COST speed — it exists to save bytes, and a format that
12201    /// trades 7% of a file for a slower model is a bad trade. This guard is
12202    /// here because correctness tests happily passed while `q4tp_matmat` was
12203    /// missing its int8 and Accelerate arms and the model ran 5x slower.
12204    /// Measured on M-series: 0.97-1.04x, i.e. parity (16 B tiles are better
12205    /// aligned than q4t's 18 B, which pays for the scale indirection).
12206    #[test]
12207    fn q4tp_matvec_keeps_pace_with_q4t() {
12208        let (rows, cols) = (4096usize, 3072usize);
12209        let bytes = synth_q4tp(rows, cols);
12210        let twin = q4tp_as_q4t(&bytes, rows, cols);
12211        let x: Vec<f32> = (0..cols).map(|i| (i % 97) as f32 / 97.0 - 0.5).collect();
12212        let mut o = vec![0f32; rows];
12213        let n = 12;
12214        let mut best = (f64::MAX, f64::MAX);
12215        // Interleaved A/B, minimum statistic: this machine throttles, and a
12216        // mean over a thermal ramp reliably indicts whichever ran second.
12217        for _ in 0..3 {
12218            let t0 = std::time::Instant::now();
12219            for _ in 0..n {
12220                q4t_matvec(&twin, &x, rows, cols, &mut o, None);
12221            }
12222            best.0 = best.0.min(t0.elapsed().as_secs_f64());
12223            let t0 = std::time::Instant::now();
12224            for _ in 0..n {
12225                q4tp_matvec(&bytes, &x, rows, cols, &mut o, None);
12226            }
12227            best.1 = best.1.min(t0.elapsed().as_secs_f64());
12228        }
12229        let ratio = best.1 / best.0;
12230        println!(
12231            "q4t {:.3} ms | q4tp {:.3} ms | {ratio:.2}x",
12232            best.0 * 1e3 / n as f64,
12233            best.1 * 1e3 / n as f64
12234        );
12235        assert!(ratio < 2.0, "q4tp matvec {ratio:.2}x slower than q4t");
12236    }
12237
12238    #[cfg(target_os = "macos")]
12239    #[test]
12240    fn q4t_matmat_accel_matches_dequant_reference() {
12241        if !accel_gemm_enabled() {
12242            return; // CMF_ACCEL=0
12243        }
12244        let (rows, cols, b) = (512usize, 1024usize, 8usize); // ≥500K → accel arm
12245        let gpr = cols / GROUP_SIZE;
12246        let mut bytes = vec![0u8; rows * gpr * Q4_TILE];
12247        for r in 0..rows {
12248            for g in 0..gpr {
12249                let t = (r * gpr + g) * Q4_TILE;
12250                let sc = 0.02 + 0.0005 * ((r * gpr + g) % 64) as f32;
12251                bytes[t..t + 2].copy_from_slice(&cortiq_core::quant::f32_to_f16(sc).to_le_bytes());
12252                for k in 0..16 {
12253                    bytes[t + 2 + k] = ((r * 31 + g * 7 + k * 13) % 251) as u8;
12254                }
12255            }
12256        }
12257        let x: Vec<f32> = (0..b * cols)
12258            .map(|i| ((i * 13 + 7) % 97) as f32 / 97.0 - 0.5)
12259            .collect();
12260        let mut got = vec![0f32; b * rows];
12261        q4t_matmat(&bytes, &x, b, rows, cols, &mut got, None);
12262        // Brute-force reference off the same tiles.
12263        let mut w = vec![0f32; rows * cols];
12264        for r in 0..rows {
12265            for g in 0..gpr {
12266                let t = (r * gpr + g) * Q4_TILE;
12267                let s = f16_to_f32(u16::from_le_bytes([bytes[t], bytes[t + 1]]));
12268                for (k, &bb) in bytes[t + 2..t + Q4_TILE].iter().enumerate() {
12269                    w[r * cols + g * GROUP_SIZE + k * 2] = ((bb & 0x0F) as f32 - 8.0) * s;
12270                    w[r * cols + g * GROUP_SIZE + k * 2 + 1] =
12271                        (((bb >> 4) & 0x0F) as f32 - 8.0) * s;
12272                }
12273            }
12274        }
12275        for bi in 0..b {
12276            for r in 0..rows {
12277                let want: f32 = (0..cols).map(|j| x[bi * cols + j] * w[r * cols + j]).sum();
12278                let d = (got[bi * rows + r] - want).abs();
12279                assert!(
12280                    d <= want.abs().max(1.0) * 1e-4,
12281                    "accel q4t GEMM diverged at ({bi},{r}): {} vs {want}",
12282                    got[bi * rows + r]
12283                );
12284            }
12285        }
12286    }
12287
12288    #[test]
12289    fn q4matvec_matches_full_dequant() {
12290        let (rows, cols) = (8, 64);
12291        let groups = rows * cols / GROUP_SIZE;
12292        // Hand-craft a q4_block blob: nibbles then f16 scales.
12293        let mut bytes = Vec::with_capacity(groups * 16 + groups * 2);
12294        for i in 0..groups * 16 {
12295            bytes.push((((i * 7 + 3) % 256) & 0xFF) as u8);
12296        }
12297        for g in 0..groups {
12298            let s = 0.01 + 0.003 * g as f32;
12299            bytes.extend_from_slice(&cortiq_core::quant::f32_to_f16(s).to_le_bytes());
12300        }
12301        let x: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.17).sin()).collect();
12302
12303        let mut reference = vec![0.0f32; rows * cols];
12304        cortiq_core::quant::dequant_q4_block(&bytes, &mut reference);
12305        let mut expect = vec![0.0f32; rows];
12306        for r in 0..rows {
12307            expect[r] = reference[r * cols..(r + 1) * cols]
12308                .iter()
12309                .zip(&x)
12310                .map(|(w, xv)| w * xv)
12311                .sum();
12312        }
12313
12314        let mut got = vec![0.0f32; rows];
12315        q4matvec(&bytes, &x, rows, cols, &mut got, None);
12316        // SDOT path quantizes activations to i8 (A8W8): bounded noise,
12317        // same contract as q8/vbit (exact path is pinned by CMF_SDOT=0
12318        // in the golden-parity gate).
12319        let tol = if a8w8_enabled() { 6e-2 } else { 1e-4 };
12320        let scale = expect.iter().fold(0f32, |m, v| m.max(v.abs())).max(1.0);
12321        for r in 0..rows {
12322            assert!(
12323                (got[r] - expect[r]).abs() < tol * scale,
12324                "row {r}: {} vs {}",
12325                got[r],
12326                expect[r]
12327            );
12328        }
12329    }
12330
12331    /// Fused two-input vbit matvec must equal two single matvecs exactly
12332    /// (same per-lane accumulation order on both scalar and SDOT paths).
12333    #[test]
12334    fn vbitmatvec2_equals_two_singles() {
12335        let (rows, cols) = (6, 64);
12336        let ng = cols / GROUP_SIZE;
12337        let bits: Vec<u8> = vec![3, 4, 5, 6, 8, 4];
12338        let mut bytes = bits.clone();
12339        for g in 0..rows * ng {
12340            let s = 0.02 + 0.001 * g as f32;
12341            bytes.extend_from_slice(&cortiq_core::quant::f32_to_f16(s).to_le_bytes());
12342        }
12343        for r in 0..rows {
12344            let b = bits[r] as usize;
12345            let (mut acc, mut nb) = (0u64, 0usize);
12346            let mut rowbytes = Vec::new();
12347            for i in 0..cols {
12348                let v = ((i * 7 + r * 13) % (1 << b)) as u64;
12349                acc = (acc << b) | v;
12350                nb += b;
12351                while nb >= 8 {
12352                    nb -= 8;
12353                    rowbytes.push(((acc >> nb) & 0xFF) as u8);
12354                }
12355            }
12356            if nb > 0 {
12357                rowbytes.push(((acc << (8 - nb)) & 0xFF) as u8);
12358            }
12359            bytes.extend_from_slice(&rowbytes);
12360        }
12361        let x1: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.19).sin()).collect();
12362        let x2: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.11).cos()).collect();
12363        let offsets = vbit_row_offsets(&bytes, rows, cols);
12364
12365        let (mut a1, mut a2) = (vec![0f32; rows], vec![0f32; rows]);
12366        vbitmatvec(&bytes, &offsets, &x1, rows, cols, &mut a1, None);
12367        vbitmatvec(&bytes, &offsets, &x2, rows, cols, &mut a2, None);
12368        let (mut b1, mut b2) = (vec![0f32; rows], vec![0f32; rows]);
12369        vbitmatvec2(
12370            &bytes, &offsets, &x1, &x2, rows, cols, &mut b1, &mut b2, None,
12371        );
12372        assert_eq!(a1, b1, "fused vbit lane 1 must be bit-identical");
12373        assert_eq!(a2, b2, "fused vbit lane 2 must be bit-identical");
12374    }
12375
12376    /// Fused two-input q4 matvec must equal two single matvecs exactly.
12377    #[test]
12378    fn q4matvec2_equals_two_singles() {
12379        let (rows, cols) = (8, 128);
12380        let groups = rows * cols / GROUP_SIZE;
12381        let mut bytes = Vec::with_capacity(groups * 16 + groups * 2);
12382        for i in 0..groups * 16 {
12383            bytes.push((((i * 7 + 3) % 256) & 0xFF) as u8);
12384        }
12385        for g in 0..groups {
12386            let s = 0.01 + 0.003 * g as f32;
12387            bytes.extend_from_slice(&cortiq_core::quant::f32_to_f16(s).to_le_bytes());
12388        }
12389        // Include an outlier channel so the SDOT correction path is
12390        // exercised in the pair kernel too.
12391        let mut x1: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.17).sin()).collect();
12392        x1[9] = 250.0;
12393        let x2: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.23).cos()).collect();
12394
12395        let (mut a1, mut a2) = (vec![0f32; rows], vec![0f32; rows]);
12396        q4matvec(&bytes, &x1, rows, cols, &mut a1, None);
12397        q4matvec(&bytes, &x2, rows, cols, &mut a2, None);
12398        let (mut b1, mut b2) = (vec![0f32; rows], vec![0f32; rows]);
12399        q4matvec2(&bytes, &x1, &x2, rows, cols, &mut b1, &mut b2, None);
12400        assert_eq!(a1, b1, "fused q4 lane 1 must be bit-identical");
12401        assert_eq!(a2, b2, "fused q4 lane 2 must be bit-identical");
12402    }
12403
12404    /// Multi-matrix job must equal separate matvecs exactly — same
12405    /// kernels, only the dispatch is fused.
12406    #[test]
12407    fn matvec_many_equals_separate_matvecs() {
12408        use crate::pool::Pool;
12409        let (r1, r2, cols) = (300, 200, 64);
12410        let mk = |salt: usize, rows: usize| {
12411            QTensor::from_f32(
12412                (0..rows * cols)
12413                    .map(|i| ((i * 7 + salt) % 97) as f32 / 97.0 - 0.5)
12414                    .collect(),
12415                rows,
12416                cols,
12417            )
12418        };
12419        let (a, b) = (mk(1, r1), mk(5, r2));
12420        let x: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.11).sin()).collect();
12421        let pool = Pool::new(3);
12422
12423        let (mut ea, mut eb) = (vec![0f32; r1], vec![0f32; r2]);
12424        a.matvec(&x, &mut ea, Some(&pool));
12425        b.matvec(&x, &mut eb, Some(&pool));
12426        let (mut ga, mut gb) = (vec![0f32; r1], vec![0f32; r2]);
12427        QTensor::matvec_many([&a, &b], &x, [&mut ga, &mut gb], Some(&pool));
12428        assert_eq!(ea, ga, "fused multi-matrix lane 1 must be bit-identical");
12429        assert_eq!(eb, gb, "fused multi-matrix lane 2 must be bit-identical");
12430    }
12431
12432    /// The public Q4TP operator must take the real mapped matvec_many arm,
12433    /// rather than the F32 fallback above.  Build a tiny valid CMF so both
12434    /// handles retain their mmap payloads, then compare the fused dispatch
12435    /// with two ordinary mapped matvec calls bit-for-bit.
12436    #[test]
12437    fn q4tp_matvec_many_equals_separate_matvecs() {
12438        use crate::pool::Pool;
12439        use cortiq_core::{CMF_VERSION, CmfHeader, CmfModel, QuantType, TensorSpec};
12440
12441        let (r1, r2, cols) = (300usize, 200usize, 64usize);
12442        let arch: cortiq_core::ModelArch = serde_json::from_value(serde_json::json!({
12443            "arch_name": "tiny-q4tp",
12444            "hidden_size": cols,
12445            "intermediate_size": cols * 2,
12446            "num_layers": 1,
12447            "num_attention_heads": 2,
12448            "num_kv_heads": 1,
12449            "head_dim": 32,
12450            "vocab_size": r1,
12451            "layer_types": ["FullAttention"],
12452            "rms_norm_eps": 1e-6,
12453            "max_position_embeddings": 8,
12454            "linear_conv_kernel_dim": 0,
12455            "linear_num_key_heads": 0,
12456            "linear_num_value_heads": 0
12457        }))
12458        .unwrap();
12459        let header = CmfHeader {
12460            format: "cmf".into(),
12461            version: CMF_VERSION,
12462            arch,
12463            quant_type: QuantType::Q4Block,
12464            provenance: None,
12465            tokenizer_config: None,
12466            section_hashes: None,
12467            skills: Vec::new(),
12468            shard: None,
12469            calibration: None,
12470            routing: None,
12471            genome: None,
12472            lineage: Vec::new(),
12473            router: None,
12474            segments: Vec::new(),
12475        };
12476        let specs = [
12477            TensorSpec {
12478                name: "q".into(),
12479                dtype: TensorDtype::Q4TiledP,
12480                shape: vec![r1, cols],
12481                data: synth_q4tp(r1, cols),
12482            },
12483            TensorSpec {
12484                name: "kv".into(),
12485                dtype: TensorDtype::Q4TiledP,
12486                shape: vec![r2, cols],
12487                data: synth_q4tp(r2, cols),
12488            },
12489        ];
12490        let dir = std::env::temp_dir().join(format!("cmf-q4tp-many-{}", std::process::id()));
12491        std::fs::create_dir_all(&dir).unwrap();
12492        let path = dir.join("m.cmf");
12493        CmfModel::write(&path, &header, &specs, None, None).unwrap();
12494        let model = Arc::new(CmfModel::open(&path).unwrap());
12495        let (a, b) = (
12496            QTensor::from_model(&model, "q").unwrap(),
12497            QTensor::from_model(&model, "kv").unwrap(),
12498        );
12499        assert_eq!(a.model_dtype(), Some(TensorDtype::Q4TiledP));
12500        assert_eq!(b.model_dtype(), Some(TensorDtype::Q4TiledP));
12501        let x: Vec<f32> = (0..cols)
12502            .map(|i| ((i * 17 + 3) % 97) as f32 / 97.0 - 0.5)
12503            .collect();
12504        let pool = Pool::new(3);
12505        let (mut ea, mut eb) = (vec![0.0f32; r1], vec![0.0f32; r2]);
12506        a.matvec(&x, &mut ea, Some(&pool));
12507        b.matvec(&x, &mut eb, Some(&pool));
12508        let (mut ga, mut gb) = (vec![0.0f32; r1], vec![0.0f32; r2]);
12509        QTensor::matvec_many([&a, &b], &x, [&mut ga, &mut gb], Some(&pool));
12510        assert_eq!(ea, ga, "Q4TP fused lane 1 must be bit-identical");
12511        assert_eq!(eb, gb, "Q4TP fused lane 2 must be bit-identical");
12512        let _ = std::fs::remove_dir_all(&dir);
12513    }
12514
12515    /// The MiMo speculative verify's kernels: several tokens' MoE through
12516    /// `moe_gate_up_rows` / `moe_down_rows` (+ the caller's route-order sum)
12517    /// is bit-identical to each token alone through `moe_gate_up_many` /
12518    /// `moe_down_many` (decode), and a row-exact `q4tp_matmat` of five
12519    /// tokens (wide enough for the blocked tiles) equals five matvecs.
12520    #[test]
12521    fn multi_token_moe_rows_equal_single_token_decode() {
12522        use crate::pool::Pool;
12523        use cortiq_core::{CMF_VERSION, CmfHeader, CmfModel, QuantType, TensorSpec};
12524
12525        let (h, inter, ne) = (64usize, 128usize, 3usize);
12526        let arch: cortiq_core::ModelArch = serde_json::from_value(serde_json::json!({
12527            "arch_name": "tiny-q4tp-moe",
12528            "hidden_size": h,
12529            "intermediate_size": inter,
12530            "num_layers": 1,
12531            "num_attention_heads": 2,
12532            "num_kv_heads": 1,
12533            "head_dim": 32,
12534            "vocab_size": 8,
12535            "layer_types": ["FullAttention"],
12536            "rms_norm_eps": 1e-6,
12537            "max_position_embeddings": 8,
12538            "linear_conv_kernel_dim": 0,
12539            "linear_num_key_heads": 0,
12540            "linear_num_value_heads": 0
12541        }))
12542        .unwrap();
12543        let header = CmfHeader {
12544            format: "cmf".into(),
12545            version: CMF_VERSION,
12546            arch,
12547            quant_type: QuantType::Q4Block,
12548            provenance: None,
12549            tokenizer_config: None,
12550            section_hashes: None,
12551            skills: Vec::new(),
12552            shard: None,
12553            calibration: None,
12554            routing: None,
12555            genome: None,
12556            lineage: Vec::new(),
12557            router: None,
12558            segments: Vec::new(),
12559        };
12560        let mut specs = Vec::new();
12561        for e in 0..ne {
12562            for (k, (n, r, c)) in [("g", inter, h), ("u", inter, h), ("d", h, inter)]
12563                .into_iter()
12564                .enumerate()
12565            {
12566                // Distinct experts: perturb only the nibble plane (any byte
12567                // is a valid pair of codes; the ladder stays intact).
12568                let mut data = synth_q4tp(r, c);
12569                for (i, byte) in data[..r * (c / GROUP_SIZE) * Q4TP_NIB]
12570                    .iter_mut()
12571                    .enumerate()
12572                {
12573                    *byte ^= ((i * (e * 3 + k + 1)) % 251) as u8;
12574                }
12575                specs.push(TensorSpec {
12576                    name: format!("{n}{e}"),
12577                    dtype: TensorDtype::Q4TiledP,
12578                    shape: vec![r, c],
12579                    data,
12580                });
12581            }
12582        }
12583        let dir = std::env::temp_dir().join(format!(
12584            "cmf-moe-rows-{}-{}",
12585            std::process::id(),
12586            FLOAT_ACTIVATIONS.get()
12587        ));
12588        std::fs::create_dir_all(&dir).unwrap();
12589        let path = dir.join("m.cmf");
12590        CmfModel::write(&path, &header, &specs, None, None).unwrap();
12591        let model = Arc::new(CmfModel::open(&path).unwrap());
12592        let t = |n: String| QTensor::from_model(&model, &n).unwrap();
12593        let g: Vec<QTensor> = (0..ne).map(|e| t(format!("g{e}"))).collect();
12594        let u: Vec<QTensor> = (0..ne).map(|e| t(format!("u{e}"))).collect();
12595        let d: Vec<QTensor> = (0..ne).map(|e| t(format!("d{e}"))).collect();
12596        let b = 4usize;
12597        let mut xs: Vec<f32> = (0..b * h)
12598            .map(|i| ((i * 31 + 7) % 89) as f32 / 89.0 - 0.5)
12599            .collect();
12600        xs[5] = 9.0; // an activation outlier on token 0
12601        // Token -> (experts in route order, weights).
12602        let routes: Vec<(Vec<usize>, Vec<f32>)> = vec![
12603            (vec![2, 0], vec![0.6, 0.4]),
12604            (vec![0, 1, 2], vec![0.2, 0.5, 0.3]),
12605            (vec![1], vec![1.0]),
12606            (vec![2, 1, 0], vec![0.25, 0.25, 0.5]),
12607        ];
12608        let pool = Pool::new(3);
12609        // Decode reference, token by token.
12610        let mut want = vec![0f32; b * h];
12611        for (tk, (idx, w)) in routes.iter().enumerate() {
12612            let x = &xs[tk * h..(tk + 1) * h];
12613            let pairs: Vec<(&QTensor, &QTensor)> = idx.iter().map(|&e| (&g[e], &u[e])).collect();
12614            let mut gs: Vec<Vec<f32>> = idx.iter().map(|_| vec![0f32; inter]).collect();
12615            assert!(QTensor::moe_gate_up_many(&pairs, x, &mut gs, Some(&pool)));
12616            if FLOAT_ACTIVATIONS.get() {
12617                for (slot, &e) in idx.iter().enumerate() {
12618                    let (mut gate, mut up) = (vec![0.0; inter], vec![0.0; inter]);
12619                    g[e].matvec(x, &mut gate, Some(&pool));
12620                    u[e].matvec(x, &mut up, Some(&pool));
12621                    for (v, u) in gate.iter_mut().zip(up) {
12622                        *v = (*v / (1.0 + (-*v).exp())) * u;
12623                    }
12624                    assert_eq!(gs[slot], gate, "float gate/up must equal ordinary matvecs");
12625                }
12626            }
12627            let downs: Vec<&QTensor> = idx.iter().map(|&e| &d[e]).collect();
12628            assert!(QTensor::moe_down_many(
12629                &downs,
12630                &gs,
12631                w,
12632                &mut want[tk * h..(tk + 1) * h],
12633                Some(&pool)
12634            ));
12635        }
12636        if FLOAT_ACTIVATIONS.get() {
12637            for (tk, (idx, w)) in routes.iter().enumerate() {
12638                let mut scalar = vec![0.0; h];
12639                for (&e, &weight) in idx.iter().zip(w) {
12640                    let (mut gate, mut up, mut down) =
12641                        (vec![0.0; inter], vec![0.0; inter], vec![0.0; h]);
12642                    g[e].matvec(&xs[tk * h..(tk + 1) * h], &mut gate, Some(&pool));
12643                    u[e].matvec(&xs[tk * h..(tk + 1) * h], &mut up, Some(&pool));
12644                    for (v, u) in gate.iter_mut().zip(up) {
12645                        *v = (*v / (1.0 + (-*v).exp())) * u;
12646                    }
12647                    d[e].matvec(&gate, &mut down, Some(&pool));
12648                    for (v, d) in scalar.iter_mut().zip(down) {
12649                        *v += weight * d;
12650                    }
12651                }
12652                assert_eq!(
12653                    &want[tk * h..(tk + 1) * h],
12654                    scalar,
12655                    "float many equals scalar experts"
12656                );
12657            }
12658        }
12659        // All four tokens at once, grouped by expert.
12660        let mut experts: Vec<usize> = Vec::new();
12661        let mut groups: Vec<Vec<usize>> = Vec::new();
12662        for (tk, (idx, _)) in routes.iter().enumerate() {
12663            for &e in idx {
12664                match experts.iter().position(|&x| x == e) {
12665                    Some(k) => groups[k].push(tk),
12666                    None => {
12667                        experts.push(e);
12668                        groups.push(vec![tk]);
12669                    }
12670                }
12671            }
12672        }
12673        let n_pairs: usize = groups.iter().map(|g| g.len()).sum();
12674        let pairs: Vec<(&QTensor, &QTensor)> = experts.iter().map(|&e| (&g[e], &u[e])).collect();
12675        let mut gs: Vec<Vec<f32>> = (0..n_pairs).map(|_| vec![0f32; inter]).collect();
12676        assert!(QTensor::moe_gate_up_rows(
12677            &pairs,
12678            &groups,
12679            &xs,
12680            &mut gs,
12681            Some(&pool)
12682        ));
12683        let downs: Vec<&QTensor> = experts.iter().map(|&e| &d[e]).collect();
12684        let lens: Vec<usize> = groups.iter().map(|g| g.len()).collect();
12685        let mut ds: Vec<Vec<f32>> = (0..n_pairs).map(|_| vec![0f32; h]).collect();
12686        assert!(QTensor::moe_down_rows(
12687            &downs,
12688            &lens,
12689            &gs,
12690            &mut ds,
12691            Some(&pool)
12692        ));
12693        let slot = |tk: usize, e: usize| {
12694            let k = experts.iter().position(|&x| x == e).unwrap();
12695            groups[..k].iter().map(|g| g.len()).sum::<usize>()
12696                + groups[k].iter().position(|&x| x == tk).unwrap()
12697        };
12698        let mut got = vec![0f32; b * h];
12699        for (tk, (idx, w)) in routes.iter().enumerate() {
12700            for i in 0..h {
12701                let mut acc = 0f32;
12702                for (&e, &we) in idx.iter().zip(w) {
12703                    acc += we * ds[slot(tk, e)][i];
12704                }
12705                got[tk * h + i] = acc;
12706            }
12707        }
12708        assert!(want.iter().any(|v| *v != 0.0));
12709        assert_eq!(
12710            want.iter().map(|v| v.to_bits()).collect::<Vec<_>>(),
12711            got.iter().map(|v| v.to_bits()).collect::<Vec<_>>(),
12712            "multi-token MoE must equal decode bit for bit"
12713        );
12714
12715        // Row-exact q4tp_matmat: five tokens (a blocked 1x4 tile + tail
12716        // otherwise) equal five matvecs.
12717        let b5 = 5usize;
12718        let x5: Vec<f32> = (0..b5 * h)
12719            .map(|i| ((i * 13 + 5) % 71) as f32 / 71.0 - 0.5)
12720            .collect();
12721        let mut mm = vec![0f32; b5 * inter];
12722        row_exact_scope(|| g[1].matmat(&x5, b5, &mut mm, Some(&pool)));
12723        for tk in 0..b5 {
12724            let mut mv = vec![0f32; inter];
12725            g[1].matvec(&x5[tk * h..(tk + 1) * h], &mut mv, Some(&pool));
12726            assert_eq!(
12727                mv.iter().map(|v| v.to_bits()).collect::<Vec<_>>(),
12728                mm[tk * inter..(tk + 1) * inter]
12729                    .iter()
12730                    .map(|v| v.to_bits())
12731                    .collect::<Vec<_>>(),
12732                "row-exact matmat token {tk}"
12733            );
12734        }
12735        // Other concurrent tests/requests may still hold the shared mode.
12736        // Nested, overlapping and unwind restoration is checked separately.
12737        let _ = std::fs::remove_dir_all(&dir);
12738    }
12739
12740    /// The row-exact fix must not touch the fast path. Outside the scope
12741    /// `q4tp_matmat` has to produce exactly what it did before, and on ARM
12742    /// "before" is spelled out below: the tuned 1x4 SDOT tile for every
12743    /// four columns and the single-row kernel for the tail. Inside the
12744    /// scope every column equals its token's matvec. `q4tp_matmat_with`
12745    /// takes the mode as an argument, so a concurrent test holding the
12746    /// shared scope cannot flip it under this one.
12747    #[test]
12748    fn q4tp_matmat_fast_path_unchanged_outside_row_exact() {
12749        use crate::pool::Pool;
12750        use std::sync::atomic::Ordering::Relaxed;
12751        let _alt = Q4TP_ALT_TEST_LOCK.lock().unwrap_or_else(|e| e.into_inner());
12752        // The tuned ARM shape, which is also what an unset switch picks
12753        // unless CMF_Q4TP_V1 is exported.
12754        Q4TP_ALT.store(2, Relaxed);
12755        let pool = Pool::new(3);
12756        // Under 500k cells, so macOS keeps the matmat off the AMX; the
12757        // second shape runs across pool workers (rows >= 256) with 32
12758        // groups of accumulation and a tail after two 1x4 tiles.
12759        for &(rows, cols, b) in &[(64usize, 256usize, 7usize), (320, 1024, 9)] {
12760            let bytes = synth_q4tp(rows, cols);
12761            let mut xs: Vec<f32> = (0..b * cols)
12762                .map(|i| ((i * 29 + 11) % 83) as f32 / 83.0 - 0.5)
12763                .collect();
12764            xs[3] = 7.5; // an activation outlier on token 0
12765            let run = |exact: bool| {
12766                let mut out = vec![0f32; b * rows];
12767                q4tp_matmat_with(&bytes, &xs, b, rows, cols, &mut out, Some(&pool), exact);
12768                out
12769            };
12770            let (fast, exact) = (run(false), run(true));
12771            #[cfg(not(target_arch = "aarch64"))]
12772            let _ = fast;
12773            let bits = |v: &[f32]| v.iter().map(|x| x.to_bits()).collect::<Vec<_>>();
12774            let mut matvecs = vec![0f32; b * rows];
12775            for (bi, o) in matvecs.chunks_mut(rows).enumerate() {
12776                q4tp_matvec(
12777                    &bytes,
12778                    &xs[bi * cols..(bi + 1) * cols],
12779                    rows,
12780                    cols,
12781                    o,
12782                    Some(&pool),
12783                );
12784            }
12785            assert!(matvecs.iter().any(|v| *v != 0.0));
12786            assert_eq!(
12787                bits(&exact),
12788                bits(&matvecs),
12789                "{rows}x{cols} b={b}: row-exact matmat must equal per-token matvecs"
12790            );
12791            #[cfg(target_arch = "aarch64")]
12792            {
12793                // The pre-fix ARM loop, cell for cell.
12794                let gpr = cols / GROUP_SIZE;
12795                let v = Q4tpView::new(&bytes, rows, cols);
12796                let mut old = vec![0f32; b * rows];
12797                let mut sc = vec![0f32; gpr];
12798                let a8w8 = a8w8_enabled();
12799                let blocked = sdot_enabled() && blocked_enabled();
12800                let acts: Vec<SplitAct> = (0..b)
12801                    .map(|bi| split_act(&xs[bi * cols..(bi + 1) * cols]))
12802                    .collect();
12803                for r in 0..rows {
12804                    v.scales_into(r, gpr, &mut sc);
12805                    if !a8w8 {
12806                        for bi in 0..b {
12807                            let x = &xs[bi * cols..(bi + 1) * cols];
12808                            old[bi * rows + r] = q4tp_row_exact(v.nib, r, gpr, x, &sc);
12809                        }
12810                        continue;
12811                    }
12812                    let finish = |d: f32, act: &SplitAct| {
12813                        let mut acc = d * act.sx;
12814                        for &(j, xv) in &act.outliers {
12815                            let (w, s) = q4tp_outlier(v.nib, r, gpr, j, &sc);
12816                            acc += w * s * xv;
12817                        }
12818                        acc
12819                    };
12820                    let mut bi = 0usize;
12821                    while blocked && bi + 4 <= b {
12822                        let xs4 = [
12823                            acts[bi].xq.as_slice(),
12824                            acts[bi + 1].xq.as_slice(),
12825                            acts[bi + 2].xq.as_slice(),
12826                            acts[bi + 3].xq.as_slice(),
12827                        ];
12828                        let d = unsafe { dot_q4tp_row_1x4_sdot(v.nib, r, gpr, xs4, &sc) };
12829                        for k in 0..4 {
12830                            old[(bi + k) * rows + r] = finish(d[k], &acts[bi + k]);
12831                        }
12832                        bi += 4;
12833                    }
12834                    for (bi, act) in acts.iter().enumerate().skip(bi) {
12835                        let d = dot_q4tp_row_i8(v.nib, r, gpr, &act.xq, &sc);
12836                        old[bi * rows + r] = finish(d, act);
12837                    }
12838                }
12839                assert_eq!(
12840                    bits(&fast),
12841                    bits(&old),
12842                    "{rows}x{cols} b={b}: the fast path changed outside row_exact"
12843                );
12844                // And it is still the fast tile that runs: its lane-parallel
12845                // fma sum rounds differently from the matvec somewhere.
12846                if blocked {
12847                    assert_ne!(
12848                        bits(&fast),
12849                        bits(&matvecs),
12850                        "{rows}x{cols} b={b}: the tuned tile no longer runs outside row_exact"
12851                    );
12852                }
12853            }
12854        }
12855        Q4TP_ALT.store(0, Relaxed);
12856    }
12857
12858    #[test]
12859    fn row_exact_scopes_survive_overlap_nesting_and_unwind() {
12860        use std::sync::{Barrier, atomic::{AtomicUsize, Ordering}};
12861        // A private counter makes this restoration test independent of
12862        // numerical tests concurrently using the production counter.
12863        let active = AtomicUsize::new(0);
12864        counted_row_exact_scope(&active, || {
12865            assert_eq!(active.load(Ordering::Acquire), 1);
12866            counted_row_exact_scope(&active, || {
12867                assert_eq!(active.load(Ordering::Acquire), 2);
12868            });
12869            assert_eq!(active.load(Ordering::Acquire), 1);
12870        });
12871        assert_eq!(active.load(Ordering::Acquire), 0);
12872
12873        let both_entered = Barrier::new(2);
12874        let release_last = Barrier::new(2);
12875        std::thread::scope(|s| {
12876            let first = s.spawn(|| counted_row_exact_scope(&active, || {
12877                both_entered.wait();
12878            }));
12879            let last = s.spawn(|| counted_row_exact_scope(&active, || {
12880                both_entered.wait();
12881                release_last.wait();
12882            }));
12883            first.join().unwrap();
12884            let after_first = active.load(Ordering::Acquire);
12885            release_last.wait();
12886            last.join().unwrap();
12887            assert_eq!(after_first, 1, "second request must remain exact");
12888        });
12889        assert_eq!(active.load(Ordering::Acquire), 0);
12890        let panic = std::panic::catch_unwind(|| {
12891            counted_row_exact_scope(&active, || panic!("scope unwind"));
12892        });
12893        assert!(panic.is_err());
12894        assert_eq!(active.load(Ordering::Acquire), 0);
12895    }
12896
12897    #[test]
12898    #[cfg(target_arch = "x86_64")]
12899    fn q4tp_float_avx2_is_bitwise_scalar() {
12900        if !avx2_enabled() {
12901            return;
12902        }
12903        for cols in [32, 64, 96, 2048, 4096] {
12904            let rows = 9;
12905            let bytes = synth_q4tp(rows, cols);
12906            let v = Q4tpView::new(&bytes, rows, cols);
12907            let gpr = cols / GROUP_SIZE;
12908            let mut sc = vec![0.0; gpr];
12909            for seed in 1..=5 {
12910                let xs: Vec<f32> = (0..cols)
12911                    .map(|i| (((i * 104729 + seed * 8191) % 100003) as f32 - 50001.0) / 7919.0)
12912                    .collect();
12913                for r in 0..rows {
12914                    v.scales_into(r, gpr, &mut sc);
12915                    let scalar = q4tp_row_float_scalar(v.nib, r, gpr, &xs, &sc);
12916                    let vector = unsafe { q4tp_row_float_avx2(v.nib, r, gpr, &xs, &sc) };
12917                    assert_eq!(
12918                        scalar.to_bits(),
12919                        vector.to_bits(),
12920                        "cols={cols} row={r} seed={seed}"
12921                    );
12922                }
12923            }
12924        }
12925    }
12926
12927    #[test]
12928    fn multi_token_moe_rows_float_equal_single_token_decode() {
12929        float_activations_scope(multi_token_moe_rows_equal_single_token_decode);
12930    }
12931
12932    #[test]
12933    fn full_gpu_q8_scope_is_nested_and_thread_local() {
12934        assert!(!FULL_GPU_Q8.get());
12935        let before = gpu_split_frac();
12936        {
12937            let _guard = enter_full_gpu_q8_scope();
12938            assert_eq!(gpu_split_frac(), 1.0);
12939            {
12940                let _nested = enter_full_gpu_q8_scope();
12941            }
12942            assert_eq!(gpu_split_frac(), 1.0);
12943            std::thread::spawn(|| assert!(!FULL_GPU_Q8.get())).join().unwrap();
12944        }
12945        assert!(!FULL_GPU_Q8.get());
12946        assert_eq!(gpu_split_frac(), before);
12947    }
12948
12949    #[test]
12950    fn float_activation_scope_is_nested_thread_local_and_unwind_safe() {
12951        assert!(!FLOAT_ACTIVATIONS.get());
12952        let before = a8w8_enabled();
12953        float_activations_scope(|| {
12954            assert!(!a8w8_enabled());
12955            float_activations_scope(|| assert!(!a8w8_enabled()));
12956            assert!(FLOAT_ACTIVATIONS.get());
12957            std::thread::spawn(|| assert!(!FLOAT_ACTIVATIONS.get()))
12958                .join()
12959                .unwrap();
12960        });
12961        assert!(!FLOAT_ACTIVATIONS.get());
12962        assert_eq!(a8w8_enabled(), before);
12963        let _ = std::panic::catch_unwind(|| float_activations_scope(|| panic!("test unwind")));
12964        assert!(!FLOAT_ACTIVATIONS.get());
12965    }
12966
12967    /// Batched q4/vbit matmat must equal per-position matvec calls
12968    /// exactly (the fallback it replaced) — same kernels, same order.
12969    #[test]
12970    fn batched_matmat_equals_per_position_matvec() {
12971        let (rows, cols, b) = (8, 64, 5);
12972        // q4 blob.
12973        let groups = rows * cols / GROUP_SIZE;
12974        let mut q4 = Vec::new();
12975        for i in 0..groups * 16 {
12976            q4.push((((i * 7 + 3) % 256) & 0xFF) as u8);
12977        }
12978        for g in 0..groups {
12979            q4.extend_from_slice(
12980                &cortiq_core::quant::f32_to_f16(0.01 + 0.003 * g as f32).to_le_bytes(),
12981            );
12982        }
12983        // vbit blob (mixed widths incl. 8).
12984        let ng = cols / GROUP_SIZE;
12985        let bits: Vec<u8> = vec![3, 4, 5, 6, 8, 4, 5, 3];
12986        let mut vb = bits.clone();
12987        for g in 0..rows * ng {
12988            vb.extend_from_slice(
12989                &cortiq_core::quant::f32_to_f16(0.02 + 0.001 * g as f32).to_le_bytes(),
12990            );
12991        }
12992        for r in 0..rows {
12993            let bw = bits[r] as usize;
12994            let (mut acc, mut nb) = (0u64, 0usize);
12995            let mut rowbytes = Vec::new();
12996            for i in 0..cols {
12997                let v = ((i * 7 + r * 13) % (1 << bw)) as u64;
12998                acc = (acc << bw) | v;
12999                nb += bw;
13000                while nb >= 8 {
13001                    nb -= 8;
13002                    rowbytes.push(((acc >> nb) & 0xFF) as u8);
13003                }
13004            }
13005            if nb > 0 {
13006                rowbytes.push(((acc << (8 - nb)) & 0xFF) as u8);
13007            }
13008            vb.extend_from_slice(&rowbytes);
13009        }
13010        let offsets = vbit_row_offsets(&vb, rows, cols);
13011
13012        let xs: Vec<f32> = (0..b * cols).map(|i| (i as f32 * 0.13).sin()).collect();
13013
13014        // q4: batch vs singles.
13015        let mut got = vec![0f32; b * rows];
13016        q4matmat(&q4, &xs, b, rows, cols, &mut got, None);
13017        for bi in 0..b {
13018            let mut expect = vec![0f32; rows];
13019            q4matvec(
13020                &q4,
13021                &xs[bi * cols..(bi + 1) * cols],
13022                rows,
13023                cols,
13024                &mut expect,
13025                None,
13026            );
13027            assert_eq!(
13028                &got[bi * rows..(bi + 1) * rows],
13029                &expect[..],
13030                "q4 batch pos {bi}"
13031            );
13032        }
13033
13034        // vbit: batch vs singles.
13035        let mut got = vec![0f32; b * rows];
13036        vbitmatmat(&vb, &offsets, &xs, b, rows, cols, &mut got, None);
13037        for bi in 0..b {
13038            let mut expect = vec![0f32; rows];
13039            vbitmatvec(
13040                &vb,
13041                &offsets,
13042                &xs[bi * cols..(bi + 1) * cols],
13043                rows,
13044                cols,
13045                &mut expect,
13046                None,
13047            );
13048            assert_eq!(
13049                &got[bi * rows..(bi + 1) * rows],
13050                &expect[..],
13051                "vbit batch pos {bi}"
13052            );
13053        }
13054    }
13055
13056    /// q4_tiled kernels must produce BIT-identical outputs to the q4
13057    /// split kernels on the same values (same ints, same order — only
13058    /// the byte placement differs).
13059    #[test]
13060    fn q4_tiled_matches_q4_block_bitexact() {
13061        let (rows, cols, b) = (8usize, 128usize, 3usize);
13062        let groups = rows * cols / GROUP_SIZE;
13063        let mut split = Vec::with_capacity(groups * 18);
13064        for i in 0..groups * 16 {
13065            split.push((((i * 7 + 3) % 256) & 0xFF) as u8);
13066        }
13067        for g in 0..groups {
13068            split.extend_from_slice(
13069                &cortiq_core::quant::f32_to_f16(0.01 + 0.003 * g as f32).to_le_bytes(),
13070            );
13071        }
13072        // Re-tile: [scale][nibbles] per group.
13073        let (packed, scales) = split.split_at(groups * 16);
13074        let mut tiled = Vec::with_capacity(groups * Q4_TILE);
13075        for g in 0..groups {
13076            tiled.extend_from_slice(&scales[g * 2..g * 2 + 2]);
13077            tiled.extend_from_slice(&packed[g * 16..(g + 1) * 16]);
13078        }
13079
13080        let mut x1: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.17).sin()).collect();
13081        x1[9] = 250.0; // exercise the outlier path
13082        let x2: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.23).cos()).collect();
13083
13084        let (mut a, mut t) = (vec![0f32; rows], vec![0f32; rows]);
13085        q4matvec(&split, &x1, rows, cols, &mut a, None);
13086        q4t_matvec(&tiled, &x1, rows, cols, &mut t, None);
13087        assert_eq!(a, t, "q4t matvec must match q4 bit-for-bit");
13088
13089        let (mut a1, mut a2) = (vec![0f32; rows], vec![0f32; rows]);
13090        let (mut t1, mut t2) = (vec![0f32; rows], vec![0f32; rows]);
13091        q4matvec2(&split, &x1, &x2, rows, cols, &mut a1, &mut a2, None);
13092        q4t_matvec2(&tiled, &x1, &x2, rows, cols, &mut t1, &mut t2, None);
13093        assert_eq!(a1, t1);
13094        assert_eq!(a2, t2);
13095
13096        let xs: Vec<f32> = (0..b * cols).map(|i| (i as f32 * 0.13).sin()).collect();
13097        let (mut am, mut tm) = (vec![0f32; b * rows], vec![0f32; b * rows]);
13098        q4matmat(&split, &xs, b, rows, cols, &mut am, None);
13099        q4t_matmat(&tiled, &xs, b, rows, cols, &mut tm, None);
13100        assert_eq!(am, tm, "q4t matmat must match q4 bit-for-bit");
13101    }
13102
13103    /// q4 SDOT outlier correction: a single huge activation channel
13104    /// (>8·rms → outlier, zeroed in xq) must still contribute its EXACT
13105    /// term. On-grid bulk (±1/0 → xq dequantizes exactly) isolates the
13106    /// correction from A8W8 noise. cols must exceed 64: at n=64 the
13107    /// 8·rms threshold equals sqrt(v²+rest) ≥ v, so a single outlier
13108    /// can never qualify (8² = n).
13109    #[test]
13110    fn q4matvec_sdot_outlier_exact() {
13111        let (rows, cols) = (4, 128);
13112        let groups = rows * cols / GROUP_SIZE;
13113        let mut bytes = Vec::with_capacity(groups * 16 + groups * 2);
13114        for i in 0..groups * 16 {
13115            bytes.push(((i * 11 + 5) % 256) as u8);
13116        }
13117        for g in 0..groups {
13118            let s = 0.02 + 0.002 * g as f32;
13119            bytes.extend_from_slice(&cortiq_core::quant::f32_to_f16(s).to_le_bytes());
13120        }
13121        let mut x: Vec<f32> = (0..cols)
13122            .map(|i| match i % 3 {
13123                0 => 1.0,
13124                1 => -1.0,
13125                _ => 0.0,
13126            })
13127            .collect();
13128        x[17] = 300.0; // ≫ 8·rms → outlier channel
13129
13130        let mut reference = vec![0.0f32; rows * cols];
13131        cortiq_core::quant::dequant_q4_block(&bytes, &mut reference);
13132        let mut expect = vec![0.0f32; rows];
13133        for r in 0..rows {
13134            expect[r] = reference[r * cols..(r + 1) * cols]
13135                .iter()
13136                .zip(&x)
13137                .map(|(w, xv)| w * xv)
13138                .sum();
13139        }
13140        let mut got = vec![0.0f32; rows];
13141        q4matvec(&bytes, &x, rows, cols, &mut got, None);
13142        let scale = expect.iter().fold(0f32, |m, v| m.max(v.abs())).max(1.0);
13143        for r in 0..rows {
13144            assert!(
13145                (got[r] - expect[r]).abs() < 2e-3 * scale,
13146                "row {r}: {} vs {} (outlier term must be exact)",
13147                got[r],
13148                expect[r]
13149            );
13150        }
13151    }
13152
13153    /// The fused q1t matvec must equal the reference (dequant_q1t → dot),
13154    /// including the ternary zero level and the binary-searched outlier
13155    /// overlay. Guards the mmap kernel that makes a 12B q1t runnable.
13156    #[test]
13157    fn q1t_matvec_matches_reference() {
13158        use cortiq_core::quant::{dequant_q1t, f32_to_f16};
13159        let (rows, cols) = (3usize, 64usize); // gpr = 2
13160        let gpr = cols / GROUP_SIZE;
13161        let scales = [0.5f32, 0.3, 0.7, 0.2, 0.6, 0.15];
13162        // Overlay (must be sorted by flat index): a few spikes across rows.
13163        let outliers: [(u32, f32); 3] = [(5, 9.0), (70, -4.5), (150, 3.25)];
13164        let is_out = |flat: usize| outliers.iter().any(|&(i, _)| i as usize == flat);
13165        let mut bytes = Vec::new();
13166        for r in 0..rows {
13167            for g in 0..gpr {
13168                bytes.extend_from_slice(&f32_to_f16(scales[r * gpr + g]).to_le_bytes());
13169                let mut c = [0u8; 7];
13170                for k in 0..GROUP_SIZE {
13171                    // Encoder invariant: code 0 at outlier positions.
13172                    let code = if is_out(r * cols + g * GROUP_SIZE + k) {
13173                        0
13174                    } else {
13175                        ((k + r * 3 + g) % 3) as u8 // 0,1,2
13176                    };
13177                    cortiq_core::quant::q1t_pack(&mut c, k, code);
13178                }
13179                bytes.extend_from_slice(&c);
13180            }
13181        }
13182        // Per-row overlay: [u32 row_ptr[rows+1]] then [(u16 col, f16 val)] by
13183        // row (outliers are sorted by flat index → already grouped by row).
13184        let mut row_ptr = vec![0u32; rows + 1];
13185        for &(idx, _) in &outliers {
13186            row_ptr[idx as usize / cols + 1] += 1;
13187        }
13188        for r in 0..rows {
13189            row_ptr[r + 1] += row_ptr[r];
13190        }
13191        for &p in &row_ptr {
13192            bytes.extend_from_slice(&p.to_le_bytes());
13193        }
13194        for &(idx, v) in &outliers {
13195            bytes.extend_from_slice(&((idx as usize % cols) as u16).to_le_bytes());
13196            bytes.extend_from_slice(&f32_to_f16(v).to_le_bytes());
13197        }
13198
13199        let mut refw = vec![0f32; rows * cols];
13200        dequant_q1t(&bytes, rows, cols, &mut refw);
13201        // On-grid activations (±1, amax 1) so the int8 SDOT path reconstructs
13202        // x exactly and matches the f32 reference (same trick as the q1 test).
13203        let x: Vec<f32> = (0..cols)
13204            .map(|j| if j % 3 == 0 { 1.0 } else { -1.0 })
13205            .collect();
13206        let mut expect = vec![0f32; rows];
13207        for r in 0..rows {
13208            let mut a = 0.0f32;
13209            for j in 0..cols {
13210                a += refw[r * cols + j] * x[j];
13211            }
13212            expect[r] = a;
13213        }
13214        let tol = |e: f32| 1e-3 * e.abs().max(1e-3);
13215        let mut got = vec![0f32; rows];
13216        q1t_matvec(&bytes, &x, rows, cols, &mut got, None);
13217        for r in 0..rows {
13218            assert!(
13219                (got[r] - expect[r]).abs() < tol(expect[r]),
13220                "row {r}: {} vs {}",
13221                got[r],
13222                expect[r]
13223            );
13224        }
13225        // matmat (b=2, f32 decode path) must agree too.
13226        let x2: Vec<f32> = x.iter().chain(x.iter()).copied().collect();
13227        let mut gm = vec![0f32; 2 * rows];
13228        q1t_matmat(&bytes, &x2, 2, rows, cols, &mut gm, None);
13229        for r in 0..rows {
13230            assert!((gm[r] - expect[r]).abs() < tol(expect[r]));
13231            assert!((gm[rows + r] - expect[r]).abs() < tol(expect[r]));
13232        }
13233        // Fused pair (q1t_matvec2) must equal two single matvecs
13234        // bit-for-bit: same unpack, same group order, same f32
13235        // accumulation per stream. Distinct x2 exercises both lanes.
13236        let xb: Vec<f32> = (0..cols)
13237            .map(|j| if j % 5 == 0 { -1.0 } else { 1.0 })
13238            .collect();
13239        let (mut s1, mut s2) = (vec![0f32; rows], vec![0f32; rows]);
13240        q1t_matvec(&bytes, &x, rows, cols, &mut s1, None);
13241        q1t_matvec(&bytes, &xb, rows, cols, &mut s2, None);
13242        let (mut p1, mut p2) = (vec![0f32; rows], vec![0f32; rows]);
13243        q1t_matvec2(&bytes, &x, &xb, rows, cols, &mut p1, &mut p2, None);
13244        assert_eq!(p1, s1, "q1t pair lane 1 ≠ single matvec");
13245        assert_eq!(p2, s2, "q1t pair lane 2 ≠ single matvec");
13246    }
13247
13248    /// Pair == 2×matvec with an ODD group count (the kernel's tail
13249    /// group) and no overlay section.
13250    #[test]
13251    fn q1t_matvec2_odd_gpr_matches_singles() {
13252        use cortiq_core::quant::{Q1T_TILE, f32_to_f16, q1t_pack};
13253        let (rows, cols) = (5usize, 96usize); // gpr = 3 → paired + tail
13254        let gpr = cols / GROUP_SIZE;
13255        let mut bytes = Vec::with_capacity(rows * gpr * Q1T_TILE);
13256        for r in 0..rows {
13257            for g in 0..gpr {
13258                bytes.extend_from_slice(&f32_to_f16(0.1 + 0.05 * (r + g) as f32).to_le_bytes());
13259                let mut c = [0u8; 7];
13260                for k in 0..GROUP_SIZE {
13261                    q1t_pack(&mut c, k, ((k * 7 + r * 5 + g * 3) % 3) as u8);
13262                }
13263                bytes.extend_from_slice(&c);
13264            }
13265        }
13266        let x1: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.31).sin()).collect();
13267        let x2: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.17).cos()).collect();
13268        let (mut s1, mut s2) = (vec![0f32; rows], vec![0f32; rows]);
13269        q1t_matvec(&bytes, &x1, rows, cols, &mut s1, None);
13270        q1t_matvec(&bytes, &x2, rows, cols, &mut s2, None);
13271        let (mut p1, mut p2) = (vec![0f32; rows], vec![0f32; rows]);
13272        q1t_matvec2(&bytes, &x1, &x2, rows, cols, &mut p1, &mut p2, None);
13273        assert_eq!(p1, s1, "odd-gpr pair lane 1 ≠ single");
13274        assert_eq!(p2, s2, "odd-gpr pair lane 2 ≠ single");
13275    }
13276
13277    // Speed A/B: fused pair (one unpack, two streams) vs two single
13278    // matvecs. Single-threaded, FFN-sized, min-of paired in-process.
13279    //   cargo test -p cortiq-engine --release q1t_matvec2_speed -- --ignored --nocapture
13280    #[test]
13281    #[ignore]
13282    fn q1t_matvec2_speed() {
13283        use cortiq_core::quant::{Q1T_TILE, f32_to_f16, q1t_pack};
13284        use std::time::Instant;
13285        let (rows, cols) = (8192usize, 4096usize);
13286        let gpr = cols / GROUP_SIZE;
13287        let mut bytes = Vec::with_capacity(rows * gpr * Q1T_TILE);
13288        for r in 0..rows {
13289            for g in 0..gpr {
13290                let s = 0.1 + ((r + g) % 7) as f32 * 0.01;
13291                bytes.extend_from_slice(&f32_to_f16(s).to_le_bytes());
13292                let mut c = [0u8; 7];
13293                for k in 0..GROUP_SIZE {
13294                    q1t_pack(&mut c, k, ((k * 7 + r + g) % 3) as u8);
13295                }
13296                bytes.extend_from_slice(&c);
13297            }
13298        }
13299        let x1: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.31).sin()).collect();
13300        let x2: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.17).cos()).collect();
13301        let (mut s1, mut s2) = (vec![0f32; rows], vec![0f32; rows]);
13302        let (mut p1, mut p2) = (vec![0f32; rows], vec![0f32; rows]);
13303        // Warm both paths once.
13304        q1t_matvec(&bytes, &x1, rows, cols, &mut s1, None);
13305        q1t_matvec2(&bytes, &x1, &x2, rows, cols, &mut p1, &mut p2, None);
13306        let (mut t_pair, mut t_two) = (f64::MAX, f64::MAX);
13307        for _ in 0..8 {
13308            let t0 = Instant::now();
13309            q1t_matvec2(&bytes, &x1, &x2, rows, cols, &mut p1, &mut p2, None);
13310            t_pair = t_pair.min(t0.elapsed().as_secs_f64() * 1000.0);
13311            let t1 = Instant::now();
13312            q1t_matvec(&bytes, &x1, rows, cols, &mut s1, None);
13313            q1t_matvec(&bytes, &x2, rows, cols, &mut s2, None);
13314            t_two = t_two.min(t1.elapsed().as_secs_f64() * 1000.0);
13315        }
13316        assert_eq!(p1, s1);
13317        assert_eq!(p2, s2);
13318        println!("q1t pair {rows}x{cols}: fused {t_pair:.2} ms | two singles {t_two:.2} ms");
13319    }
13320
13321    // Speed A/B: the base-3-division decode (what the packing commit left in
13322    // place) vs the fused sign-LUT matvec. Both single-threaded, same bytes.
13323    //   cargo test -p cortiq-engine q1t_matvec_speed -- --ignored --nocapture
13324    #[test]
13325    #[ignore]
13326    fn q1t_matvec_speed() {
13327        use cortiq_core::quant::{Q1T_TILE, f32_to_f16, q1t_code, q1t_pack};
13328        use std::time::Instant;
13329        let (rows, cols) = (8192usize, 4096usize); // FFN-sized
13330        let gpr = cols / GROUP_SIZE;
13331        let mut bytes = Vec::with_capacity(rows * gpr * Q1T_TILE + 16);
13332        for r in 0..rows {
13333            for g in 0..gpr {
13334                let s = 0.1 + ((r + g) % 7) as f32 * 0.01;
13335                bytes.extend_from_slice(&f32_to_f16(s).to_le_bytes());
13336                let mut c = [0u8; 7];
13337                for k in 0..GROUP_SIZE {
13338                    q1t_pack(&mut c, k, ((k * 7 + r + g) % 3) as u8);
13339                }
13340                bytes.extend_from_slice(&c);
13341            }
13342        }
13343        let (n, stride) = (rows * cols, 40usize); // ~2.5% outliers, per-row overlay
13344        let mut row_ptr = vec![0u32; rows + 1];
13345        let mut idx = 0usize;
13346        while idx < n {
13347            row_ptr[idx / cols + 1] += 1;
13348            idx += stride;
13349        }
13350        for r in 0..rows {
13351            row_ptr[r + 1] += row_ptr[r];
13352        }
13353        for &p in &row_ptr {
13354            bytes.extend_from_slice(&p.to_le_bytes());
13355        }
13356        let mut idx = 0usize;
13357        while idx < n {
13358            bytes.extend_from_slice(&((idx % cols) as u16).to_le_bytes());
13359            bytes.extend_from_slice(&f32_to_f16((idx % 13) as f32 * 0.1 - 0.6).to_le_bytes());
13360            idx += stride;
13361        }
13362        // On-grid ±1 so the fast path's int8 SDOT is exact vs the f32 "slow"
13363        // reference (the A/B is a timing check; values must still agree).
13364        let x: Vec<f32> = (0..cols)
13365            .map(|j| if j % 3 == 0 { 1.0 } else { -1.0 })
13366            .collect();
13367        let (rp_off, ent_off, has_ov) = q1t_overlay(&bytes, rows * gpr * Q1T_TILE, rows);
13368
13369        // "before": base-3 division decode into a buffer, then dot.
13370        let slow = |out: &mut [f32]| {
13371            let mut buf = vec![0f32; cols];
13372            for r in 0..rows {
13373                for g in 0..gpr {
13374                    let off = (r * gpr + g) * Q1T_TILE;
13375                    let s = f16_to_f32(u16::from_le_bytes([bytes[off], bytes[off + 1]]));
13376                    let codes = &bytes[off + 2..off + Q1T_TILE];
13377                    for k in 0..GROUP_SIZE {
13378                        buf[g * GROUP_SIZE + k] = match q1t_code(codes, k) {
13379                            1 => s,
13380                            2 => -s,
13381                            _ => 0.0,
13382                        };
13383                    }
13384                }
13385                out[r] = q1t_row_outlier_correction(&bytes, r, rp_off, ent_off, has_ov, &x)
13386                    + (0..cols).map(|j| buf[j] * x[j]).sum::<f32>();
13387            }
13388        };
13389        let iters = 5;
13390        let mut a = vec![0f32; rows];
13391        slow(&mut a); // warm
13392        let t = Instant::now();
13393        for _ in 0..iters {
13394            slow(&mut a);
13395        }
13396        let slow_ms = t.elapsed().as_secs_f64() * 1e3 / iters as f64;
13397
13398        let mut b = vec![0f32; rows];
13399        q1t_matvec(&bytes, &x, rows, cols, &mut b, None); // warm
13400        let t = Instant::now();
13401        for _ in 0..iters {
13402            q1t_matvec(&bytes, &x, rows, cols, &mut b, None);
13403        }
13404        let fast_ms = t.elapsed().as_secs_f64() * 1e3 / iters as f64;
13405
13406        for r in 0..rows {
13407            assert!((a[r] - b[r]).abs() < 1e-2, "mismatch row {r}");
13408        }
13409        println!(
13410            "q1t matvec {rows}x{cols} (1 thread): div-decode {slow_ms:.2} ms  fused-LUT {fast_ms:.2} ms  => {:.2}x",
13411            slow_ms / fast_ms
13412        );
13413    }
13414}
13415
13416#[cfg(test)]
13417mod gemm_bench {
13418    /// `cargo test -p cortiq-engine --release q4tp_matmat_throughput -- --ignored --nocapture`
13419    /// Times the batched q4tp GEMM at the shapes the image DiT runs
13420    /// (b=296 tokens, 2304 -> 9216), on synthetic bytes: no model, no
13421    /// mmap, no thermal drift over minutes — a kernel change shows up
13422    /// here in seconds where a full render hides it in noise.
13423    ///
13424    /// On macOS add `CMF_ACCEL=0`: this shape is over the 500k-cell mark
13425    /// where the matmat hands off to Accelerate's dequant sgemm, and
13426    /// without the opt-out both rows below measure the AMX, not the
13427    /// kernel under test.
13428    #[test]
13429    #[ignore]
13430    fn q4tp_matmat_throughput() {
13431        let _alt = super::Q4TP_ALT_TEST_LOCK
13432            .lock()
13433            .unwrap_or_else(|e| e.into_inner());
13434        // 296 is a prompt-encode batch; the image DiT runs 2085 at
13435        // 512x512, where the activation panel stops fitting L2 and the
13436        // loop's shape starts to matter more than its instructions.
13437        let b: usize = std::env::var("CMF_BENCH_B")
13438            .ok()
13439            .and_then(|v| v.parse().ok())
13440            .unwrap_or(296);
13441        let (rows, cols) = (9216usize, 2304usize);
13442        let (_, _, _) = (rows, cols, b);
13443        let total =
13444            cortiq_core::quant::expected_nbytes(cortiq_core::TensorDtype::Q4TiledP, &[rows, cols])
13445                .unwrap();
13446        // Random nibbles are fine, but the row params are f16 (lo, step)
13447        // of a geometric ladder: garbage there gives exp2 of a huge
13448        // exponent, the scales come back inf, and the whole bench times
13449        // NaN arithmetic instead of the kernel.
13450        let (params_off, codes_off, _) = cortiq_core::quant::q4tp_sections(rows, cols);
13451        let mut bytes: Vec<u8> = (0..total).map(|i| (i * 37 % 251) as u8).collect();
13452        let lo = cortiq_core::quant::f32_to_f16(-4.0);
13453        let step = cortiq_core::quant::f32_to_f16(0.1);
13454        for r in 0..rows {
13455            let o = params_off + r * 4;
13456            bytes[o..o + 2].copy_from_slice(&lo.to_le_bytes());
13457            bytes[o + 2..o + 4].copy_from_slice(&step.to_le_bytes());
13458        }
13459        let _ = codes_off;
13460        let xs: Vec<f32> = (0..b * cols)
13461            .map(|i| ((i % 97) as f32 - 48.0) / 48.0)
13462            .collect();
13463        let mut out = vec![0f32; b * rows];
13464        let pool = crate::pool::Pool::from_env();
13465        // A shared 48-core stand drifts ±25% run to run, which is wider
13466        // than any kernel change worth making. So: alternate the two
13467        // kernels inside one process and keep the BEST time for
13468        // each. Interleaving makes both see the same interference, and a
13469        // minimum is the one statistic another tenant cannot inflate.
13470        super::q4tp_matmat(&bytes, &xs, b, rows, cols, &mut out, pool.as_deref());
13471        let reps: usize = std::env::var("CMF_BENCH_REPS")
13472            .ok()
13473            .and_then(|v| v.parse().ok())
13474            .unwrap_or(10);
13475        let mut best = [f64::MAX; 2];
13476        let mut sums = [0f32; 2];
13477        for _ in 0..reps {
13478            for (k, w) in [(0usize, 1u8), (1usize, 2u8)] {
13479                super::Q4TP_ALT.store(w, std::sync::atomic::Ordering::Relaxed);
13480                let t = std::time::Instant::now();
13481                super::q4tp_matmat(&bytes, &xs, b, rows, cols, &mut out, pool.as_deref());
13482                best[k] = best[k].min(t.elapsed().as_secs_f64());
13483                sums[k] = out.iter().take(64).sum::<f32>();
13484            }
13485        }
13486        let flops = 2.0 * b as f64 * rows as f64 * cols as f64;
13487        for (k, name) in ["previous", "tuned   "].iter().enumerate() {
13488            println!(
13489                "q4tp matmat {rows}x{cols} b={b} {name}: {:.1} ms  {:.1} GFLOP/s  (checksum {:.3})",
13490                best[k] * 1e3,
13491                flops / best[k] / 1e9,
13492                sums[k]
13493            );
13494        }
13495        assert!(
13496            (sums[0] - sums[1]).abs() < 1e-2,
13497            "the tuned kernel changed the result: {} vs {}",
13498            sums[0],
13499            sums[1]
13500        );
13501    }
13502
13503    /// The blocked kernel must agree with the per-column path exactly —
13504    /// same weights, same activation split, only a different instruction
13505    /// mix. Shapes are chosen to hit the awkward cases: a column count
13506    /// that leaves an odd group (the 512-bit kernel does two at a time),
13507    /// and a batch that does not divide by four.
13508    #[test]
13509    fn q4tp_matmat_blocked_matches_scalar() {
13510        use std::sync::atomic::Ordering::Relaxed;
13511        let _alt = super::Q4TP_ALT_TEST_LOCK
13512            .lock()
13513            .unwrap_or_else(|e| e.into_inner());
13514        // The last shape carries the image DiT's column count — 2304, so
13515        // 72 groups of accumulation, which is where a reordered sum can
13516        // actually drift — and runs through the thread pool, since the
13517        // blocked path splits rows across workers. Its row count stays
13518        // under 500k cells on purpose: above that, macOS diverts the whole
13519        // matmat to the Accelerate/AMX dequant sgemm and neither kernel
13520        // here would run.
13521        for &(rows, cols, b) in &[
13522            (64usize, 128usize, 7usize),
13523            (33, 96, 4),
13524            (16, 256, 9),
13525            (192, 2304, 37),
13526        ] {
13527            let total = cortiq_core::quant::expected_nbytes(
13528                cortiq_core::TensorDtype::Q4TiledP,
13529                &[rows, cols],
13530            )
13531            .unwrap();
13532            let (params_off, _, _) = cortiq_core::quant::q4tp_sections(rows, cols);
13533            let mut bytes: Vec<u8> = (0..total).map(|i| (i * 61 % 251) as u8).collect();
13534            let lo = cortiq_core::quant::f32_to_f16(-4.0);
13535            let step = cortiq_core::quant::f32_to_f16(0.1);
13536            for r in 0..rows {
13537                let o = params_off + r * 4;
13538                bytes[o..o + 2].copy_from_slice(&lo.to_le_bytes());
13539                bytes[o + 2..o + 4].copy_from_slice(&step.to_le_bytes());
13540            }
13541            let xs: Vec<f32> = (0..b * cols)
13542                .map(|i| ((i % 89) as f32 - 44.0) / 44.0)
13543                .collect();
13544            let mut got = vec![0f32; b * rows];
13545            let mut want = vec![0f32; b * rows];
13546            let gpr = cols / 32;
13547            let view = super::Q4tpView::new(&bytes, rows, cols);
13548            let pool = crate::pool::Pool::from_env();
13549            super::Q4TP_ALT.store(2, Relaxed);
13550            super::q4tp_matmat(&bytes, &xs, b, rows, cols, &mut got, pool.as_deref());
13551            super::Q4TP_ALT.store(1, Relaxed);
13552            super::q4tp_matmat(&bytes, &xs, b, rows, cols, &mut want, pool.as_deref());
13553            super::Q4TP_ALT.store(0, Relaxed);
13554            // Measured against the output's scale, not cell by cell: a
13555            // dot product of 2304 terms lands near zero wherever the row
13556            // and the activation nearly cancel, and there a per-cell
13557            // ratio reports 1e-3 for an absolute error of 5e-6 — f32's
13558            // own rounding, reordered. What must stay small is the error
13559            // relative to what the layer actually outputs.
13560            let scale = want.iter().fold(0f32, |m, v| m.max(v.abs())).max(1e-6);
13561            let (mut worst, mut at) = (0f32, 0usize);
13562            for (i, (g, w)) in got.iter().zip(&want).enumerate() {
13563                if (g - w).abs() > worst {
13564                    worst = (g - w).abs();
13565                    at = i;
13566                }
13567            }
13568            assert!(
13569                worst <= 1e-4 * scale,
13570                "{rows}x{cols} b={b}: blocked and scalar disagree by {worst:.3e} \
13571                 (scale {scale:.3e}) at cell {at}: {} vs {}",
13572                got[at],
13573                want[at]
13574            );
13575
13576            // "Same speed, no quality loss" is a claim about which answer
13577            // is RIGHT, not about which two agree. Both paths sum the same
13578            // 2304 products in different orders, so f64 decides: the
13579            // blocked kernel keeps sixteen partial sums and folds them at
13580            // the end, which is a shallower addition tree than the
13581            // per-column path's running scalar, and it must not be worse.
13582            let (mut e_blocked, mut e_scalar) = (0f64, 0f64);
13583            for bi in 0..b {
13584                let act = super::split_act(&xs[bi * cols..(bi + 1) * cols]);
13585                for r in 0..rows {
13586                    let mut sc = vec![0f32; gpr];
13587                    view.scales_into(r, gpr, &mut sc);
13588                    let mut exact = 0f64;
13589                    for j in 0..cols {
13590                        let (w, sq) = super::q4tp_outlier(view.nib, r, gpr, j, &sc);
13591                        exact += w as f64 * sq as f64 * act.xq[j] as f64;
13592                    }
13593                    exact *= act.sx as f64;
13594                    for &(j, xv) in &act.outliers {
13595                        let (w, sq) = super::q4tp_outlier(view.nib, r, gpr, j, &sc);
13596                        exact += w as f64 * sq as f64 * xv as f64;
13597                    }
13598                    let i = bi * rows + r;
13599                    e_blocked = e_blocked.max((got[i] as f64 - exact).abs());
13600                    e_scalar = e_scalar.max((want[i] as f64 - exact).abs());
13601                }
13602            }
13603            println!(
13604                "{rows}x{cols} b={b}: worst error vs f64 — blocked {e_blocked:.3e}, \
13605                 per-column {e_scalar:.3e}"
13606            );
13607            // An absolute bar, not a race between the two: at these
13608            // magnitudes both sit in f32's last bits, and on a small shape
13609            // whichever one happens to round the unluckiest cell "wins" by
13610            // a factor the next seed reverses.
13611            assert!(
13612                e_blocked <= 1e-5 * scale as f64 && e_scalar <= 1e-5 * scale as f64,
13613                "{rows}x{cols} b={b}: error against f64 too large — blocked \
13614                 {e_blocked:.3e}, per-column {e_scalar:.3e}, scale {scale:.3e}"
13615            );
13616        }
13617    }
13618
13619    /// What the row-exact contract costs a speculative-verify panel on the
13620    /// host: the same batch through the fast arms, through the row-exact
13621    /// arms, and as one matvec per token (the other way to be exact).
13622    /// Arms alternate inside one process and keep their best time.
13623    /// `CMF_GPU=0 cargo test -p cortiq-engine --release q4tp_matmat_row_exact_cost -- --ignored --nocapture`
13624    #[test]
13625    #[ignore]
13626    fn q4tp_matmat_row_exact_cost() {
13627        let _alt = super::Q4TP_ALT_TEST_LOCK
13628            .lock()
13629            .unwrap_or_else(|e| e.into_inner());
13630        let pool = crate::pool::Pool::from_env();
13631        let reps: usize = std::env::var("CMF_BENCH_REPS")
13632            .ok()
13633            .and_then(|v| v.parse().ok())
13634            .unwrap_or(30);
13635        for &(rows, cols) in &[(2048usize, 4096usize), (4096, 2048), (4096, 4096)] {
13636            let total = cortiq_core::quant::expected_nbytes(
13637                cortiq_core::TensorDtype::Q4TiledP,
13638                &[rows, cols],
13639            )
13640            .unwrap();
13641            let (params_off, _, _) = cortiq_core::quant::q4tp_sections(rows, cols);
13642            let mut bytes: Vec<u8> = (0..total).map(|i| (i * 37 % 251) as u8).collect();
13643            let lo = cortiq_core::quant::f32_to_f16(-4.0);
13644            let step = cortiq_core::quant::f32_to_f16(0.1);
13645            for r in 0..rows {
13646                let o = params_off + r * 4;
13647                bytes[o..o + 2].copy_from_slice(&lo.to_le_bytes());
13648                bytes[o + 2..o + 4].copy_from_slice(&step.to_le_bytes());
13649            }
13650            for &b in &[2usize, 4, 5, 8] {
13651                let xs: Vec<f32> = (0..b * cols)
13652                    .map(|i| ((i % 97) as f32 - 48.0) / 48.0)
13653                    .collect();
13654                let mut out = vec![0f32; b * rows];
13655                let mut best = [f64::MAX; 3];
13656                for _ in 0..reps {
13657                    for (k, best_k) in best.iter_mut().enumerate() {
13658                        let t = std::time::Instant::now();
13659                        match k {
13660                            0 | 1 => super::q4tp_matmat_with(
13661                                &bytes,
13662                                &xs,
13663                                b,
13664                                rows,
13665                                cols,
13666                                &mut out,
13667                                pool.as_deref(),
13668                                k == 1,
13669                            ),
13670                            _ => {
13671                                for (bi, o) in out.chunks_mut(rows).enumerate() {
13672                                    super::q4tp_matvec(
13673                                        &bytes,
13674                                        &xs[bi * cols..(bi + 1) * cols],
13675                                        rows,
13676                                        cols,
13677                                        o,
13678                                        pool.as_deref(),
13679                                    );
13680                                }
13681                            }
13682                        }
13683                        *best_k = best_k.min(t.elapsed().as_secs_f64());
13684                    }
13685                }
13686                println!(
13687                    "q4tp {rows}x{cols} b={b}: fast {:.3} ms, row-exact {:.3} ms, \
13688                     {b} matvecs {:.3} ms",
13689                    best[0] * 1e3,
13690                    best[1] * 1e3,
13691                    best[2] * 1e3
13692                );
13693            }
13694        }
13695    }
13696
13697    /// The q4t twin of the throughput bench, same shape and rules, so the
13698    /// two quantisations' batch kernels can be read against each other.
13699    /// `cargo test -p cortiq-engine --release q4t_matmat_throughput -- --ignored --nocapture`
13700    #[test]
13701    #[ignore]
13702    fn q4t_matmat_throughput() {
13703        let (rows, cols, b) = (9216usize, 2304usize, 296usize);
13704        let total =
13705            cortiq_core::quant::expected_nbytes(cortiq_core::TensorDtype::Q4Tiled, &[rows, cols])
13706                .unwrap();
13707        // q4t carries a per-group f16 scale in the tile's first two bytes;
13708        // random bytes there decode to inf and the bench would time NaNs.
13709        let mut bytes: Vec<u8> = (0..total).map(|i| (i * 37 % 251) as u8).collect();
13710        let sc = cortiq_core::quant::f32_to_f16(0.02);
13711        for t in bytes.chunks_mut(super::Q4_TILE) {
13712            t[..2].copy_from_slice(&sc.to_le_bytes());
13713        }
13714        let xs: Vec<f32> = (0..b * cols)
13715            .map(|i| ((i % 97) as f32 - 48.0) / 48.0)
13716            .collect();
13717        let mut out = vec![0f32; b * rows];
13718        let pool = crate::pool::Pool::from_env();
13719        super::q4t_matmat(&bytes, &xs, b, rows, cols, &mut out, pool.as_deref());
13720        let reps: usize = std::env::var("CMF_BENCH_REPS")
13721            .ok()
13722            .and_then(|v| v.parse().ok())
13723            .unwrap_or(10);
13724        let mut best = f64::MAX;
13725        for _ in 0..reps {
13726            let t = std::time::Instant::now();
13727            super::q4t_matmat(&bytes, &xs, b, rows, cols, &mut out, pool.as_deref());
13728            best = best.min(t.elapsed().as_secs_f64());
13729        }
13730        let flops = 2.0 * b as f64 * rows as f64 * cols as f64;
13731        println!(
13732            "q4t matmat {rows}x{cols} b={b}: {:.1} ms  {:.1} GFLOP/s  (checksum {:.3})",
13733            best * 1e3,
13734            flops / best / 1e9,
13735            out.iter().take(64).sum::<f32>()
13736        );
13737    }
13738}