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                let run = |start: usize, end: usize| {
1794                    for o in start..end {
1795                        let row = &data[o * cols..(o + 1) * cols];
1796                        for bi in 0..b {
1797                            let x = &xs_all[bi * cols..(bi + 1) * cols];
1798                            let mut acc = 0f32;
1799                            for j in 0..cols {
1800                                acc += row[j] * x[j];
1801                            }
1802                            unsafe { *out_addr.at(bi * rows + o) = acc };
1803                        }
1804                    }
1805                };
1806                dispatch_rows(pool, rows, &run);
1807            }
1808            Self::Mapped {
1809                model,
1810                idx,
1811                dtype,
1812                row_scale,
1813                col_field,
1814                vbit_offsets,
1815                ..
1816            } => {
1817                if crate::prism::is_forward_weight(model, &model.tensors[*idx].name) {
1818                    let mut transformed = Vec::with_capacity(xs_all.len());
1819                    for bi in 0..b {
1820                        transformed.extend_from_slice(&crate::prism::forward(
1821                            model,
1822                            &xs_all[bi * cols..(bi + 1) * cols],
1823                        ));
1824                    }
1825                    match dtype {
1826                        TensorDtype::Q4Block => {
1827                            q4matmat(self.quant_bytes(), &transformed, b, rows, cols, out, pool)
1828                        }
1829                        TensorDtype::Q4Tiled => {
1830                            q4t_matmat(self.quant_bytes(), &transformed, b, rows, cols, out, pool)
1831                        }
1832                        TensorDtype::Q4TiledP => {
1833                            q4tp_matmat(self.quant_bytes(), &transformed, b, rows, cols, out, pool)
1834                        }
1835                        TensorDtype::Q2TiledP => {
1836                            let affine =
1837                                crate::prism::is_affine_target(model, &model.tensors[*idx].name);
1838                            // Affine Prism Q2TP has a descriptor-aware GPU
1839                            // kernel for short/tail batches too.  Unlike the
1840                            // ordinary Q2TP path, don't force b<32 back to a
1841                            // scalar CPU matmat: prefill chunks and the final
1842                            // tail both need to stay on the tested GPU arm.
1843                            let gpu_batch_ok = if affine {
1844                                b >= 2
1845                            } else {
1846                                b >= 32 && b * rows * cols >= 128_000_000
1847                            };
1848                            if gpu_batch_ok
1849                                && cols % 32 == 0
1850                                && crate::gpu::enabled_here()
1851                                && crate::gpu::q2tp_gpu_opt_in()
1852                            {
1853                                let gpu_ok = if affine {
1854                                    crate::gpu::q2tp_affine_matmat(
1855                                        model,
1856                                        *idx,
1857                                        &transformed,
1858                                        b,
1859                                        rows,
1860                                        cols,
1861                                        out,
1862                                    )
1863                                } else {
1864                                    crate::gpu::q2tp_matmat(
1865                                        model,
1866                                        *idx,
1867                                        &transformed,
1868                                        b,
1869                                        rows,
1870                                        cols,
1871                                        out,
1872                                    )
1873                                };
1874                                if gpu_ok {
1875                                    return;
1876                                }
1877                            }
1878                            // A one-token Prism decode is the other short
1879                            // case.  Use the descriptor-aware matvec kernel
1880                            // before falling back to the exact CPU path.
1881                            if affine
1882                                && b == 1
1883                                && cols % 32 == 0
1884                                && crate::gpu::enabled_here()
1885                                && crate::gpu::q2tp_gpu_opt_in()
1886                                && crate::gpu::q2tp_affine_matvec(
1887                                    model,
1888                                    *idx,
1889                                    &transformed[..cols],
1890                                    rows,
1891                                    cols,
1892                                    &mut out[..rows],
1893                                )
1894                            {
1895                                return;
1896                            }
1897                            if affine {
1898                                q2tp_affine_matmat(
1899                                    self.quant_bytes(),
1900                                    &transformed,
1901                                    b,
1902                                    rows,
1903                                    cols,
1904                                    out,
1905                                    pool,
1906                                )
1907                            } else {
1908                                q2tp_matmat(
1909                                    self.quant_bytes(),
1910                                    &transformed,
1911                                    b,
1912                                    rows,
1913                                    cols,
1914                                    out,
1915                                    pool,
1916                                )
1917                            }
1918                        }
1919                        TensorDtype::Q1 => {
1920                            q1_matmat(self.quant_bytes(), &transformed, b, rows, cols, out, pool)
1921                        }
1922                        TensorDtype::Q1T => {
1923                            q1t_matmat(self.quant_bytes(), &transformed, b, rows, cols, out, pool)
1924                        }
1925                        TensorDtype::Vbit | TensorDtype::VbitRo => vbitmatmat(
1926                            self.quant_bytes(),
1927                            vbit_offsets,
1928                            &transformed,
1929                            b,
1930                            rows,
1931                            cols,
1932                            out,
1933                            pool,
1934                        ),
1935                        TensorDtype::Q8Row | TensorDtype::Q8_2f => {
1936                            let pre: Vec<std::borrow::Cow<'_, [f32]>> = (0..b)
1937                                .map(|bi| {
1938                                    prescale(
1939                                        &transformed[bi * cols..(bi + 1) * cols],
1940                                        col_field,
1941                                        *dtype,
1942                                    )
1943                                })
1944                                .collect();
1945                            qmatmat(self.quant_bytes(), row_scale, &pre, rows, cols, out, pool)
1946                        }
1947                        _ => unreachable!("unsupported mapped Prism dtype {dtype:?}"),
1948                    }
1949                    return;
1950                }
1951                if *dtype == TensorDtype::Q4Block {
1952                    q4matmat(self.quant_bytes(), xs_all, b, rows, cols, out, pool);
1953                    return;
1954                }
1955                if *dtype == TensorDtype::Q4TiledP {
1956                    // GPU batched q4tp GEMM (dequant + f32nt mul_mm on the
1957                    // device); the probe keeps whichever beats the CPU arm.
1958                    // Narrow (prompt-encode) and wide (DiT) batches probe
1959                    // as separate classes — the regimes have opposite
1960                    // winners and one shared verdict locked the wrong arm.
1961                    // Kill switch (gpu::mm_kill): one grossly slow GPU op
1962                    // (a fair-condition op is ≤~100 ms even at 1024px)
1963                    // means the device is contended by another process
1964                    // (e.g. a simulator) — verdicts are per-process, so
1965                    // without the bail the whole render crawls behind
1966                    // someone else's queue.
1967                    // Row-exact batches stay on the host: the device GEMM
1968                    // is an f32 dequant-sgemm, not the host matvec's sum.
1969                    if b >= 32
1970                        && b * rows * cols >= 128_000_000
1971                        && cols % 32 == 0
1972                        && !row_exact()
1973                        && !crate::gpu::mm_killed()
1974                        && crate::gpu::enabled_here()
1975                    {
1976                        let class = if b >= 128 {
1977                            crate::gpu::OpClass::MatmatWide
1978                        } else {
1979                            crate::gpu::OpClass::Matmat
1980                        };
1981                        if let Self::Mapped { model, idx, .. } = self {
1982                            // In-process A/B (`CMF_MM_AB=1`). Three
1983                            // wall-clock A/Bs on a shared stand disagreed
1984                            // with each other by 25% on the same change,
1985                            // because the machine drifts between processes
1986                            // and interleaving whole renders does not fix
1987                            // that. Here both arms run back to back on the
1988                            // SAME data inside one call, so whatever the
1989                            // machine is doing, it does to both — and the
1990                            // disagreement between their outputs falls out
1991                            // for free. Doubles the work; a diagnostic,
1992                            // not a mode.
1993                            if crate::mm_ab::on() {
1994                                let mut g = vec![0f32; b * rows];
1995                                let t = std::time::Instant::now();
1996                                let took = crate::gpu::q4tp_matmat(
1997                                    model, *idx, xs_all, b, rows, cols, &mut g,
1998                                );
1999                                let dg = t.elapsed();
2000                                let t = std::time::Instant::now();
2001                                q4tp_matmat(self.quant_bytes(), xs_all, b, rows, cols, out, pool);
2002                                let dc = t.elapsed();
2003                                crate::mm_ab::record(b, rows, cols, took, dg, dc, &g, out);
2004                                return;
2005                            }
2006                            let t0 = std::time::Instant::now();
2007                            // A cold call takes the device arm: its sample
2008                            // is discarded either way, and the upload is
2009                            // what the next step needs.
2010                            let resident = crate::gpu::weight_is_resident(model, *idx);
2011                            match crate::gpu::probe_arm_cold_prefers_gpu(class, resident) {
2012                                crate::gpu::ProbeArm::Gpu => {
2013                                    if crate::gpu::q4tp_matmat(
2014                                        model, *idx, xs_all, b, rows, cols, out,
2015                                    ) {
2016                                        let el = t0.elapsed();
2017                                        // Work-proportional budget: ~8× the
2018                                        // fair-device estimate (+20 ms slack).
2019                                        // An absolute cap missed the worst
2020                                        // case — contended ops sit at
2021                                        // 100–240 ms each and still bury a
2022                                        // render whose fair op is 3–9 ms.
2023                                        // Cold ops (first PSO build, buffer
2024                                        // alloc) are exempt: a one-off
2025                                        // ~50 ms compile is not contention.
2026                                        let flops = 2.0 * b as f64 * rows as f64 * cols as f64;
2027                                        let budget = std::time::Duration::from_secs_f64(
2028                                            flops / 1.5e12 * 8.0 + 0.020,
2029                                        );
2030                                        crate::gpu::mm_budget_check(
2031                                            "q4tp matmat",
2032                                            el,
2033                                            budget,
2034                                            crate::gpu::probe_was_cold() || !resident,
2035                                        );
2036                                        crate::gpu::probe_record(class, true, el);
2037                                        return;
2038                                    }
2039                                }
2040                                crate::gpu::ProbeArm::CpuTimed => {
2041                                    q4tp_matmat(
2042                                        self.quant_bytes(),
2043                                        xs_all,
2044                                        b,
2045                                        rows,
2046                                        cols,
2047                                        out,
2048                                        pool,
2049                                    );
2050                                    crate::gpu::probe_record(class, false, t0.elapsed());
2051                                    return;
2052                                }
2053                                crate::gpu::ProbeArm::Cpu => {}
2054                            }
2055                        }
2056                    }
2057                    q4tp_matmat(self.quant_bytes(), xs_all, b, rows, cols, out, pool);
2058                    return;
2059                }
2060                if *dtype == TensorDtype::Q2TiledP {
2061                    // Same device arm as q4tp, behind the same probe:
2062                    // the planes differ, the dispatch does not. Without
2063                    // this a q2tp file ran its widest projections on the
2064                    // host while the 4-bit one had the card, which is a
2065                    // codec paying for its size twice.
2066                    if b >= 32
2067                        && b * rows * cols >= 128_000_000
2068                        && cols % 32 == 0
2069                        && !crate::gpu::mm_killed()
2070                        && crate::gpu::enabled_here()
2071                    {
2072                        let class = if b >= 128 {
2073                            crate::gpu::OpClass::MatmatWide
2074                        } else {
2075                            crate::gpu::OpClass::Matmat
2076                        };
2077                        if let Self::Mapped { model, idx, .. } = self {
2078                            let t0 = std::time::Instant::now();
2079                            match crate::gpu::probe_arm(class) {
2080                                crate::gpu::ProbeArm::Gpu => {
2081                                    if crate::gpu::q2tp_matmat(
2082                                        model, *idx, xs_all, b, rows, cols, out,
2083                                    ) {
2084                                        crate::gpu::probe_record(class, true, t0.elapsed());
2085                                        return;
2086                                    }
2087                                }
2088                                crate::gpu::ProbeArm::CpuTimed => {
2089                                    q2tp_matmat(
2090                                        self.quant_bytes(),
2091                                        xs_all,
2092                                        b,
2093                                        rows,
2094                                        cols,
2095                                        out,
2096                                        pool,
2097                                    );
2098                                    crate::gpu::probe_record(class, false, t0.elapsed());
2099                                    return;
2100                                }
2101                                crate::gpu::ProbeArm::Cpu => {}
2102                            }
2103                        }
2104                    }
2105                    // Without a host arm a q2tp tensor falls through to
2106                    // the q8 fallback, which reads it at one BYTE per
2107                    // weight — a 2x overrun that killed pool workers
2108                    // mid-prefill while the dispatcher waited forever.
2109                    q2tp_matmat(self.quant_bytes(), xs_all, b, rows, cols, out, pool);
2110                    return;
2111                }
2112                if *dtype == TensorDtype::Q4Tiled {
2113                    // GPU batched q4t GEMM (dequant + f32nt mul_mm on the
2114                    // device); the probe keeps whichever beats the CPU arm.
2115                    // Narrow (prompt-encode) and wide (DiT) batches probe
2116                    // as separate classes — the regimes have opposite
2117                    // winners and one shared verdict locked the wrong arm.
2118                    // Kill switch (gpu::mm_kill): one grossly slow GPU op
2119                    // (a fair-condition op is ≤~100 ms even at 1024px)
2120                    // means the device is contended by another process
2121                    // (e.g. a simulator) — verdicts are per-process, so
2122                    // without the bail the whole render crawls behind
2123                    // someone else's queue.
2124                    if b >= 32
2125                        && b * rows * cols >= 128_000_000
2126                        && cols % 32 == 0
2127                        && !crate::gpu::mm_killed()
2128                        && crate::gpu::enabled_here()
2129                    {
2130                        let class = if b >= 128 {
2131                            crate::gpu::OpClass::MatmatWide
2132                        } else {
2133                            crate::gpu::OpClass::Matmat
2134                        };
2135                        if let Self::Mapped { model, idx, .. } = self {
2136                            let t0 = std::time::Instant::now();
2137                            match crate::gpu::probe_arm(class) {
2138                                crate::gpu::ProbeArm::Gpu => {
2139                                    if crate::gpu::q4t_matmat(
2140                                        model, *idx, xs_all, b, rows, cols, out,
2141                                    ) {
2142                                        let el = t0.elapsed();
2143                                        // Work-proportional budget: ~8× the
2144                                        // fair-device estimate (+20 ms slack).
2145                                        // An absolute cap missed the worst
2146                                        // case — contended ops sit at
2147                                        // 100–240 ms each and still bury a
2148                                        // render whose fair op is 3–9 ms.
2149                                        // Cold ops (first PSO build, buffer
2150                                        // alloc) are exempt: a one-off
2151                                        // ~50 ms compile is not contention.
2152                                        let flops = 2.0 * b as f64 * rows as f64 * cols as f64;
2153                                        let budget = std::time::Duration::from_secs_f64(
2154                                            flops / 1.5e12 * 8.0 + 0.020,
2155                                        );
2156                                        crate::gpu::mm_budget_check(
2157                                            "q4t matmat",
2158                                            el,
2159                                            budget,
2160                                            crate::gpu::probe_was_cold(),
2161                                        );
2162                                        crate::gpu::probe_record(class, true, el);
2163                                        return;
2164                                    }
2165                                }
2166                                crate::gpu::ProbeArm::CpuTimed => {
2167                                    q4t_matmat(
2168                                        self.quant_bytes(),
2169                                        xs_all,
2170                                        b,
2171                                        rows,
2172                                        cols,
2173                                        out,
2174                                        pool,
2175                                    );
2176                                    crate::gpu::probe_record(class, false, t0.elapsed());
2177                                    return;
2178                                }
2179                                crate::gpu::ProbeArm::Cpu => {}
2180                            }
2181                        }
2182                    }
2183                    q4t_matmat(self.quant_bytes(), xs_all, b, rows, cols, out, pool);
2184                    return;
2185                }
2186                if *dtype == TensorDtype::Q1 {
2187                    // GPU batched q1 GEMM for wide prefill (q1_mul_mm on the
2188                    // device); the probe keeps whichever beats the CPU matmat.
2189                    if b >= 32
2190                        && b * rows * cols >= 128_000_000
2191                        && cols % 64 == 0
2192                        && crate::gpu::enabled_here()
2193                    {
2194                        if let Self::Mapped { model, idx, .. } = self {
2195                            let t0 = std::time::Instant::now();
2196                            match crate::gpu::probe_arm(crate::gpu::OpClass::Matmat) {
2197                                crate::gpu::ProbeArm::Gpu => {
2198                                    if crate::gpu::q1_matmat(
2199                                        model, *idx, xs_all, b, rows, cols, out,
2200                                    ) {
2201                                        crate::gpu::probe_record(
2202                                            crate::gpu::OpClass::Matmat,
2203                                            true,
2204                                            t0.elapsed(),
2205                                        );
2206                                        return;
2207                                    }
2208                                }
2209                                crate::gpu::ProbeArm::CpuTimed => {
2210                                    q1_matmat(self.quant_bytes(), xs_all, b, rows, cols, out, pool);
2211                                    crate::gpu::probe_record(
2212                                        crate::gpu::OpClass::Matmat,
2213                                        false,
2214                                        t0.elapsed(),
2215                                    );
2216                                    return;
2217                                }
2218                                crate::gpu::ProbeArm::Cpu => {}
2219                            }
2220                        }
2221                    }
2222                    q1_matmat(self.quant_bytes(), xs_all, b, rows, cols, out, pool);
2223                    return;
2224                }
2225                if *dtype == TensorDtype::Q1T {
2226                    // GPU batched GEMM for wide prefill (base + overlay on the
2227                    // device); probe keeps the winner vs the CPU matmat.
2228                    if b >= 32 && b * rows * cols >= 128_000_000 && crate::gpu::enabled_here() {
2229                        if let Self::Mapped { model, idx, .. } = self {
2230                            let t0 = std::time::Instant::now();
2231                            match crate::gpu::probe_arm(crate::gpu::OpClass::Matmat) {
2232                                crate::gpu::ProbeArm::Gpu => {
2233                                    if crate::gpu::q1t_matmat(
2234                                        model, *idx, xs_all, b, rows, cols, out,
2235                                    ) {
2236                                        crate::gpu::probe_record(
2237                                            crate::gpu::OpClass::Matmat,
2238                                            true,
2239                                            t0.elapsed(),
2240                                        );
2241                                        return;
2242                                    }
2243                                }
2244                                crate::gpu::ProbeArm::CpuTimed => {
2245                                    q1t_matmat(
2246                                        self.quant_bytes(),
2247                                        xs_all,
2248                                        b,
2249                                        rows,
2250                                        cols,
2251                                        out,
2252                                        pool,
2253                                    );
2254                                    crate::gpu::probe_record(
2255                                        crate::gpu::OpClass::Matmat,
2256                                        false,
2257                                        t0.elapsed(),
2258                                    );
2259                                    return;
2260                                }
2261                                crate::gpu::ProbeArm::Cpu => {}
2262                            }
2263                        }
2264                    }
2265                    q1t_matmat(self.quant_bytes(), xs_all, b, rows, cols, out, pool);
2266                    return;
2267                }
2268                if matches!(dtype, TensorDtype::Vbit | TensorDtype::VbitRo) {
2269                    vbitmatmat(
2270                        self.quant_bytes(),
2271                        vbit_offsets,
2272                        xs_all,
2273                        b,
2274                        rows,
2275                        cols,
2276                        out,
2277                        pool,
2278                    );
2279                    return;
2280                }
2281                let pre: Vec<std::borrow::Cow<'_, [f32]>> = (0..b)
2282                    .map(|bi| prescale(&xs_all[bi * cols..(bi + 1) * cols], col_field, *dtype))
2283                    .collect();
2284                // MiMo verification is a 2–4 row decode panel, not a wide
2285                // prompt GEMM. Keep q8 projections on the same device as
2286                // decode; the generic b>=8 gate otherwise silently moves
2287                // every projection back to CPU. The short wgpu matmat uses
2288                // the same 64-lane reduction as its single-token matvec.
2289                if row_exact()
2290                    && (1..=4).contains(&b)
2291                    && matches!(dtype, TensorDtype::Q8Row | TensorDtype::Q8_2f)
2292                    && crate::gpu::enabled_here()
2293                    && crate::gpu::wgpu_active()
2294                {
2295                    let flat: Vec<f32> = pre.iter().flat_map(|v| v.iter().copied()).collect();
2296                    if crate::gpu::q8_matmat(model, *idx, row_scale, &flat, b, rows, cols, out) {
2297                        return;
2298                    }
2299                }
2300                // D5: large prefill-batch GEMMs — on the GPU (threshold by
2301                // work volume: submission carries b×rows×cols MACs).
2302                // Runtime probe: the naive GEMM shader + sync readback
2303                // lose to the CPU GEMM on slow driver stacks — alternate
2304                // both arms and keep the winner.
2305                if b >= 8 && b * rows * cols >= 128_000_000 && crate::gpu::enabled_here() {
2306                    if let Self::Mapped { model, idx, .. } = self {
2307                        let t0 = std::time::Instant::now();
2308                        match crate::gpu::probe_arm(crate::gpu::OpClass::Matmat) {
2309                            crate::gpu::ProbeArm::Gpu
2310                                if crate::gpu::probe_deciding(crate::gpu::OpClass::Matmat)
2311                                    && !crate::gpu::q8_resident_or_upload(model, *idx) =>
2312                            {
2313                                // Cold weights during probing: the upload
2314                                // has started, the count runs on the CPU —
2315                                // the GPU arm samples on the next touch.
2316                                let q = self.quant_bytes();
2317                                qmatmat(q, row_scale, &pre, rows, cols, out, pool);
2318                                return;
2319                            }
2320                            crate::gpu::ProbeArm::Gpu => {
2321                                let flat: Vec<f32> =
2322                                    pre.iter().flat_map(|v| v.iter().copied()).collect();
2323                                if crate::gpu::q8_matmat(
2324                                    model, *idx, row_scale, &flat, b, rows, cols, out,
2325                                ) {
2326                                    crate::gpu::probe_record(
2327                                        crate::gpu::OpClass::Matmat,
2328                                        true,
2329                                        t0.elapsed(),
2330                                    );
2331                                    return;
2332                                }
2333                            }
2334                            crate::gpu::ProbeArm::CpuTimed => {
2335                                let q = self.quant_bytes();
2336                                qmatmat(q, row_scale, &pre, rows, cols, out, pool);
2337                                crate::gpu::probe_record(
2338                                    crate::gpu::OpClass::Matmat,
2339                                    false,
2340                                    t0.elapsed(),
2341                                );
2342                                return;
2343                            }
2344                            crate::gpu::ProbeArm::Cpu => {}
2345                        }
2346                    }
2347                }
2348                let q = self.quant_bytes();
2349                qmatmat(q, row_scale, &pre, rows, cols, out, pool);
2350            }
2351        }
2352    }
2353}
2354
2355impl QTensor {
2356    /// The device GEMM this tensor would take, run once on the caller's
2357    /// data — the startup parity probe's arm, and the one place that knows
2358    /// which entry point each codec has.
2359    ///
2360    /// It exists because the probe used to look for a `q4tp` weight by
2361    /// name AND dtype, and a container packed any other way was declared
2362    /// "host path" for the whole render even though its codec had a device
2363    /// GEMM of its own. A gate that only recognizes one codec is a gate
2364    /// that silently downgrades every other one.
2365    pub fn device_matmat(&self, xs: &[f32], b: usize, out: &mut [f32]) -> bool {
2366        let (rows, cols) = (self.rows(), self.cols());
2367        let Self::Mapped {
2368            model,
2369            idx,
2370            dtype,
2371            row_scale,
2372            col_field,
2373            ..
2374        } = self
2375        else {
2376            return false;
2377        };
2378        if crate::prism::has_contract(model) {
2379            return false;
2380        }
2381        match *dtype {
2382            TensorDtype::Q4TiledP => crate::gpu::q4tp_matmat(model, *idx, xs, b, rows, cols, out),
2383            // The two-field codec folds its column field into the
2384            // activation, which leaves a plain per-row int8 GEMM — the
2385            // same kernel `q8_row` uses, on both backends.
2386            TensorDtype::Q8Row | TensorDtype::Q8_2f => {
2387                // The field belongs to the weight; only a backend that cannot
2388                // apply it there makes a scaled copy of the activation.
2389                if *dtype == TensorDtype::Q8_2f
2390                    && std::env::var("CMF_Q8_2F_DEV").as_deref() != Ok("0")
2391                    && crate::gpu::q8_matmat_2f(
2392                        model, *idx, row_scale, col_field, xs, b, rows, cols, out,
2393                    )
2394                {
2395                    return true;
2396                }
2397                let flat: Vec<f32> = (0..b)
2398                    .flat_map(|bi| {
2399                        prescale(&xs[bi * cols..(bi + 1) * cols], col_field, *dtype).into_owned()
2400                    })
2401                    .collect();
2402                crate::gpu::q8_matmat(model, *idx, row_scale, &flat, b, rows, cols, out)
2403            }
2404            _ => false,
2405        }
2406    }
2407
2408    /// Multi-matrix job (roadmap §3 P0): N tensors sharing one input
2409    /// run under a SINGLE pool dispatch — QKV or gate+up cost one
2410    /// barrier instead of N. Per-row math is the exact same kernel as
2411    /// `matvec` (bit-identical outputs); only the dispatch is fused.
2412    /// Falls back to N sequential matvecs when the set is not a uniform
2413    /// q8-family/F32 group or there is no pool.
2414    pub fn matvec_many<const N: usize>(
2415        ts: [&QTensor; N],
2416        x: &[f32],
2417        mut outs: [&mut [f32]; N],
2418        pool: Option<&Pool>,
2419    ) {
2420        let total_rows: usize = ts.iter().map(|t| t.rows()).sum();
2421        if ts.iter().any(|t| t.has_prism_contract()) {
2422            // The fused range kernels have no transform descriptor.  Let
2423            // each tensor's ordinary matvec dispatch perform the explicit
2424            // signed FWHT (and retain CPU fallback for mixed q2tp/q4tp).
2425            for (t, o) in ts.iter().zip(outs.iter_mut()) {
2426                t.matvec(x, o, pool);
2427            }
2428            return;
2429        }
2430        let uniform_q8 = ts.iter().all(|t| {
2431            matches!(
2432                t,
2433                Self::Mapped {
2434                    dtype: TensorDtype::Q8Row | TensorDtype::Q8_2f,
2435                    ..
2436                }
2437            )
2438        });
2439        let uniform_f32 = ts.iter().all(|t| matches!(t, Self::F32 { .. }));
2440        if uniform_f32 && crate::f32_backend::active() {
2441            for (t, o) in ts.iter().zip(outs.iter_mut()) {
2442                t.matvec(x, o, pool);
2443            }
2444            return;
2445        }
2446        let uniform_q4 = ts.iter().all(|t| {
2447            matches!(
2448                t,
2449                Self::Mapped {
2450                    dtype: TensorDtype::Q4Block,
2451                    ..
2452                }
2453            )
2454        });
2455        let uniform_vbit = ts.iter().all(|t| {
2456            matches!(
2457                t,
2458                Self::Mapped {
2459                    dtype: TensorDtype::Vbit | TensorDtype::VbitRo,
2460                    ..
2461                }
2462            )
2463        });
2464        let uniform_q1 = ts.iter().all(|t| {
2465            matches!(
2466                t,
2467                Self::Mapped {
2468                    dtype: TensorDtype::Q1,
2469                    ..
2470                }
2471            )
2472        });
2473        let uniform_q1t = ts.iter().all(|t| {
2474            matches!(
2475                t,
2476                Self::Mapped {
2477                    dtype: TensorDtype::Q1T,
2478                    ..
2479                }
2480            )
2481        });
2482        // q4tp is the skeleton dtype of the big MoE files, and without an arm
2483        // here every projection that shares an input paid its own pool
2484        // barrier: DeepSeek-V4's attention step alone hands this function
2485        // wq_a, wkv and both compressors' pairs off the same hidden state.
2486        let uniform_q4tp = ts.iter().all(|t| {
2487            matches!(
2488                t,
2489                Self::Mapped {
2490                    dtype: TensorDtype::Q4TiledP,
2491                    ..
2492                }
2493            )
2494        }) && ts
2495            .iter()
2496            .all(|t| t.cols() == ts[0].cols() && t.cols() % GROUP_SIZE == 0);
2497        let Some(pool) = pool else {
2498            for (t, o) in ts.iter().zip(outs.iter_mut()) {
2499                t.matvec(x, o, None);
2500            }
2501            return;
2502        };
2503        if total_rows < 256
2504            || !(uniform_q8
2505                || uniform_f32
2506                || uniform_q4
2507                || uniform_vbit
2508                || uniform_q1
2509                || uniform_q1t
2510                || uniform_q4tp)
2511        {
2512            for (t, o) in ts.iter().zip(outs.iter_mut()) {
2513                t.matvec(x, o, Some(pool));
2514            }
2515            return;
2516        }
2517
2518        if uniform_q4tp {
2519            // Every tensor's rows laid end to end in one virtual row space,
2520            // so the whole set is ONE dispatch. The per-row body is the
2521            // `q4tp_matvec` arm verbatim — same activation split, same
2522            // accumulation order — so the outputs are bit-identical to the
2523            // sequential calls this replaces.
2524            let cols = ts[0].cols();
2525            let gpr = cols / GROUP_SIZE;
2526            let views: [Q4tpView; N] =
2527                std::array::from_fn(|i| Q4tpView::new(ts[i].quant_bytes(), ts[i].rows(), cols));
2528            let rows_of: [usize; N] = std::array::from_fn(|i| ts[i].rows());
2529            let outs_addr: [SendMut; N] = std::array::from_fn(|i| SendMut(outs[i].as_mut_ptr()));
2530            // flat index -> (which tensor, which of its rows)
2531            let locate = |flat: usize| -> (usize, usize) {
2532                let mut acc = 0;
2533                for (i, &r) in rows_of.iter().enumerate() {
2534                    if flat < acc + r {
2535                        return (i, flat - acc);
2536                    }
2537                    acc += r;
2538                }
2539                (rows_of.len() - 1, 0)
2540            };
2541            let (views, outs_addr) = (&views, &outs_addr);
2542            if a8w8_enabled() {
2543                let act = split_act(x);
2544                let act = &act;
2545                let run = |start: usize, end: usize| {
2546                    with_krow(gpr, |sc| {
2547                        for flat in start..end {
2548                            let (t, r) = locate(flat);
2549                            let v = &views[t];
2550                            v.scales_into(r, gpr, sc);
2551                            let mut acc = dot_q4tp_row_i8(v.nib, r, gpr, &act.xq, sc) * act.sx;
2552                            for &(j, xv) in &act.outliers {
2553                                let (w, s) = q4tp_outlier(v.nib, r, gpr, j, sc);
2554                                acc += w * s * xv;
2555                            }
2556                            // SAFETY: one worker owns each (tensor, row) pair.
2557                            unsafe { *outs_addr[t].at(r) = acc };
2558                        }
2559                    });
2560                };
2561                pool.run_rows(total_rows, &run);
2562            } else {
2563                let run = |start: usize, end: usize| {
2564                    with_krow(gpr, |sc| {
2565                        for flat in start..end {
2566                            let (t, r) = locate(flat);
2567                            let v = &views[t];
2568                            v.scales_into(r, gpr, sc);
2569                            // SAFETY: one worker owns each (tensor, row) pair.
2570                            unsafe { *outs_addr[t].at(r) = q4tp_row_exact(v.nib, r, gpr, x, sc) };
2571                        }
2572                    });
2573                };
2574                pool.run_rows(total_rows, &run);
2575            }
2576            return;
2577        }
2578
2579        if uniform_q1 {
2580            // One shared activation split + group sums (q1 has no col
2581            // field; the same input feeds every tensor).
2582            let outs_addr: [SendMut; N] = std::array::from_fn(|i| SendMut(outs[i].as_mut_ptr()));
2583            if a8w8_enabled() {
2584                let act = split_act(x);
2585                let gsum = q1_group_sums(&act.xq, ts[0].cols() / GROUP_SIZE);
2586                let (act, gsum) = (&act, &gsum);
2587                let closures: [_; N] = std::array::from_fn(|i| {
2588                    let (bytes, gpr, out) =
2589                        (ts[i].quant_bytes(), ts[i].cols() / GROUP_SIZE, outs_addr[i]);
2590                    move |s: usize, e: usize| q1_range_a8w8(bytes, gpr, act, gsum, out, s, e)
2591                });
2592                let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2593                    std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2594                pool.run_many(&parts);
2595            } else {
2596                let closures: [_; N] = std::array::from_fn(|i| {
2597                    let (bytes, gpr, out) =
2598                        (ts[i].quant_bytes(), ts[i].cols() / GROUP_SIZE, outs_addr[i]);
2599                    move |s: usize, e: usize| q1_range_f32(bytes, gpr, x, out, s, e)
2600                });
2601                let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2602                    std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2603                pool.run_many(&parts);
2604            }
2605            return;
2606        }
2607
2608        if uniform_q1t {
2609            // Q1T batched: one shared activation split + overlay decode,
2610            // all tensors' rows in ONE pool dispatch (saves N−1 dispatches
2611            // and N−1 redundant split_act calls per layer).
2612            let outs_addr: [SendMut; N] = std::array::from_fn(|i| SendMut(outs[i].as_mut_ptr()));
2613            const TILE: usize = cortiq_core::quant::Q1T_TILE;
2614            if a8w8_enabled() {
2615                let act = split_act(x);
2616                let act = &act;
2617                let x_ref = x;
2618                let closures: [_; N] = std::array::from_fn(|i| {
2619                    let bytes = ts[i].quant_bytes();
2620                    let (rows, cols) = (ts[i].rows(), ts[i].cols());
2621                    let gpr = cols / GROUP_SIZE;
2622                    let (rp_off, ent_off, has_ov) = q1t_overlay(bytes, rows * gpr * TILE, rows);
2623                    let out = outs_addr[i];
2624                    move |s: usize, e: usize| {
2625                        q1t_range_a8w8(bytes, gpr, rp_off, ent_off, has_ov, act, x_ref, out, s, e)
2626                    }
2627                });
2628                let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2629                    std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2630                pool.run_many(&parts);
2631            } else {
2632                let x_ref = x;
2633                let closures: [_; N] = std::array::from_fn(|i| {
2634                    let bytes = ts[i].quant_bytes();
2635                    let (rows, cols) = (ts[i].rows(), ts[i].cols());
2636                    let gpr = cols / GROUP_SIZE;
2637                    let (rp_off, ent_off, has_ov) = q1t_overlay(bytes, rows * gpr * TILE, rows);
2638                    let out = outs_addr[i];
2639                    move |s: usize, e: usize| {
2640                        q1t_range_f32_batch(bytes, gpr, rp_off, ent_off, has_ov, x_ref, out, s, e)
2641                    }
2642                });
2643                let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2644                    std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2645                pool.run_many(&parts);
2646            }
2647            return;
2648        }
2649
2650        if uniform_q4 || uniform_vbit {
2651            let outs_addr: [SendMut; N] = std::array::from_fn(|i| SendMut(outs[i].as_mut_ptr()));
2652            // q4/vbit share one activation split — no per-tensor col field.
2653            if a8w8_enabled() {
2654                let act = split_act(x);
2655                let act = &act;
2656                if uniform_q4 {
2657                    let closures: [_; N] = std::array::from_fn(|i| {
2658                        let (packed, scales) =
2659                            q4_split(ts[i].quant_bytes(), ts[i].rows(), ts[i].cols());
2660                        let (gpr, cols, out) =
2661                            (ts[i].cols() / GROUP_SIZE, ts[i].cols(), outs_addr[i]);
2662                        move |s: usize, e: usize| {
2663                            q4_range_a8w8(packed, scales, gpr, cols, act, out, s, e)
2664                        }
2665                    });
2666                    let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2667                        std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2668                    pool.run_many(&parts);
2669                } else {
2670                    let closures: [_; N] = std::array::from_fn(|i| {
2671                        let Self::Mapped { vbit_offsets, .. } = ts[i] else {
2672                            unreachable!()
2673                        };
2674                        let (bytes, rows, cols, out) = (
2675                            ts[i].quant_bytes(),
2676                            ts[i].rows(),
2677                            ts[i].cols(),
2678                            outs_addr[i],
2679                        );
2680                        move |s: usize, e: usize| {
2681                            vbit_range_a8w8(bytes, vbit_offsets, x, act, rows, cols, out, s, e)
2682                        }
2683                    });
2684                    let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2685                        std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2686                    pool.run_many(&parts);
2687                }
2688                return;
2689            }
2690            if uniform_q4 {
2691                let closures: [_; N] = std::array::from_fn(|i| {
2692                    let (packed, scales) =
2693                        q4_split(ts[i].quant_bytes(), ts[i].rows(), ts[i].cols());
2694                    let (gpr, out) = (ts[i].cols() / GROUP_SIZE, outs_addr[i]);
2695                    move |s: usize, e: usize| q4_range_f32(packed, scales, gpr, x, out, s, e)
2696                });
2697                let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2698                    std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2699                pool.run_many(&parts);
2700            } else {
2701                let closures: [_; N] = std::array::from_fn(|i| {
2702                    let Self::Mapped { vbit_offsets, .. } = ts[i] else {
2703                        unreachable!()
2704                    };
2705                    let (bytes, rows, cols, out) = (
2706                        ts[i].quant_bytes(),
2707                        ts[i].rows(),
2708                        ts[i].cols(),
2709                        outs_addr[i],
2710                    );
2711                    move |s: usize, e: usize| {
2712                        vbit_range_f32(bytes, vbit_offsets, x, rows, cols, out, s, e)
2713                    }
2714                });
2715                let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2716                    std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2717                pool.run_many(&parts);
2718            }
2719            return;
2720        }
2721
2722        if uniform_f32 {
2723            let outs_addr: [SendMut; N] = std::array::from_fn(|i| SendMut(outs[i].as_mut_ptr()));
2724            let closures: [_; N] = std::array::from_fn(|i| {
2725                let Self::F32 { data, cols, .. } = ts[i] else {
2726                    unreachable!()
2727                };
2728                let out = outs_addr[i];
2729                move |start: usize, end: usize| {
2730                    for o in start..end {
2731                        let row = &data[o * cols..(o + 1) * cols];
2732                        let mut sum = 0.0f32;
2733                        for j in 0..*cols {
2734                            sum += row[j] * x[j];
2735                        }
2736                        // SAFETY: disjoint (tensor, row) cells per worker.
2737                        unsafe { *out.at(o) = sum };
2738                    }
2739                }
2740            });
2741            let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2742                std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2743            pool.run_many(&parts);
2744            return;
2745        }
2746
2747        // Uniform q8-family: per-tensor prescale (q8_2f col fields
2748        // differ per tensor) + the shared range kernels.
2749        struct Ctx<'a> {
2750            bytes: &'a [u8],
2751            #[cfg_attr(not(target_arch = "aarch64"), allow(dead_code))]
2752            rep: &'a [u8],
2753            row_scale: &'a [f32],
2754            cols: usize,
2755            xs: std::borrow::Cow<'a, [f32]>,
2756        }
2757        let ctxs: [Ctx<'_>; N] = std::array::from_fn(|i| {
2758            let Self::Mapped {
2759                dtype,
2760                cols,
2761                row_scale,
2762                col_field,
2763                repack,
2764                ..
2765            } = ts[i]
2766            else {
2767                unreachable!()
2768            };
2769            Ctx {
2770                bytes: ts[i].quant_bytes(),
2771                rep: repack,
2772                row_scale,
2773                cols: *cols,
2774                xs: prescale(x, col_field, *dtype),
2775            }
2776        });
2777        let outs_addr: [SendMut; N] = std::array::from_fn(|i| SendMut(outs[i].as_mut_ptr()));
2778        #[cfg(target_arch = "aarch64")]
2779        if sdot_enabled() {
2780            let acts: [SplitAct; N] = std::array::from_fn(|i| split_act(&ctxs[i].xs));
2781            let closures: [_; N] = std::array::from_fn(|i| {
2782                let (c, act, out) = (&ctxs[i], &acts[i], outs_addr[i]);
2783                move |start: usize, end: usize| {
2784                    q8_range_sdot(c.bytes, c.rep, c.row_scale, act, c.cols, out, start, end)
2785                }
2786            });
2787            let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2788                std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2789            pool.run_many(&parts);
2790            return;
2791        }
2792        #[cfg(target_arch = "x86_64")]
2793        if avx2_a8w8_enabled() {
2794            let acts: [SplitAct; N] = std::array::from_fn(|i| split_act(&ctxs[i].xs));
2795            let closures: [_; N] = std::array::from_fn(|i| {
2796                let (c, act, out) = (&ctxs[i], &acts[i], outs_addr[i]);
2797                move |start: usize, end: usize| {
2798                    q8_range_avx2(c.bytes, c.row_scale, act, c.cols, out, start, end)
2799                }
2800            });
2801            let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2802                std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2803            pool.run_many(&parts);
2804            return;
2805        }
2806        let closures: [_; N] = std::array::from_fn(|i| {
2807            let (c, out) = (&ctxs[i], outs_addr[i]);
2808            move |start: usize, end: usize| {
2809                q8_range_f32(c.bytes, c.row_scale, &c.xs, c.cols, out, start, end)
2810            }
2811        });
2812        let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2813            std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2814        pool.run_many(&parts);
2815    }
2816}
2817
2818impl QTensor {
2819    /// Pair-input multi-matrix job: N tensors × 2 shared inputs under a
2820    /// single pool dispatch — the MTP/pair decode path publishes one job
2821    /// for Q/K/V (and one for gate+up) instead of one per tensor.
2822    /// Per-row math is exactly `matvec2`'s kernels; bit-identical.
2823    #[allow(clippy::needless_range_loop)]
2824    pub fn matvec2_many<const N: usize>(
2825        ts: [&QTensor; N],
2826        x1: &[f32],
2827        x2: &[f32],
2828        mut o1s: [&mut [f32]; N],
2829        mut o2s: [&mut [f32]; N],
2830        pool: Option<&Pool>,
2831    ) {
2832        let total_rows: usize = ts.iter().map(|t| t.rows()).sum();
2833        if ts.iter().any(|t| t.has_prism_contract()) {
2834            for i in 0..N {
2835                ts[i].matvec2(x1, x2, o1s[i], o2s[i], pool);
2836            }
2837            return;
2838        }
2839        let uniform_q8 = ts.iter().all(|t| {
2840            matches!(
2841                t,
2842                Self::Mapped {
2843                    dtype: TensorDtype::Q8Row | TensorDtype::Q8_2f,
2844                    ..
2845                }
2846            )
2847        });
2848        let uniform_f32 = ts.iter().all(|t| matches!(t, Self::F32 { .. }));
2849        let uniform_q4 = ts.iter().all(|t| {
2850            matches!(
2851                t,
2852                Self::Mapped {
2853                    dtype: TensorDtype::Q4Block,
2854                    ..
2855                }
2856            )
2857        });
2858        let uniform_vbit = ts.iter().all(|t| {
2859            matches!(
2860                t,
2861                Self::Mapped {
2862                    dtype: TensorDtype::Vbit | TensorDtype::VbitRo,
2863                    ..
2864                }
2865            )
2866        });
2867        let fusable = pool.is_some()
2868            && total_rows >= 256
2869            && (uniform_q8 || uniform_f32 || uniform_q4 || uniform_vbit);
2870        if !fusable {
2871            for i in 0..N {
2872                ts[i].matvec2(x1, x2, o1s[i], o2s[i], pool);
2873            }
2874            return;
2875        }
2876        let pool = pool.unwrap();
2877
2878        if uniform_q4 || uniform_vbit {
2879            let p1: [SendMut; N] = std::array::from_fn(|i| SendMut(o1s[i].as_mut_ptr()));
2880            let p2: [SendMut; N] = std::array::from_fn(|i| SendMut(o2s[i].as_mut_ptr()));
2881            // q4/vbit share activation splits — no per-tensor col field.
2882            if a8w8_enabled() {
2883                let a1 = split_act(x1);
2884                let a2 = split_act(x2);
2885                let (a1, a2) = (&a1, &a2);
2886                if uniform_q4 {
2887                    let closures: [_; N] = std::array::from_fn(|i| {
2888                        let (packed, scales) =
2889                            q4_split(ts[i].quant_bytes(), ts[i].rows(), ts[i].cols());
2890                        let (gpr, cols, o1, o2) =
2891                            (ts[i].cols() / GROUP_SIZE, ts[i].cols(), p1[i], p2[i]);
2892                        move |s: usize, e: usize| {
2893                            q4_range2_a8w8(packed, scales, gpr, cols, a1, a2, o1, o2, s, e)
2894                        }
2895                    });
2896                    let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2897                        std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2898                    pool.run_many(&parts);
2899                } else {
2900                    let closures: [_; N] = std::array::from_fn(|i| {
2901                        let Self::Mapped { vbit_offsets, .. } = ts[i] else {
2902                            unreachable!()
2903                        };
2904                        let (bytes, rows, cols, o1, o2) = (
2905                            ts[i].quant_bytes(),
2906                            ts[i].rows(),
2907                            ts[i].cols(),
2908                            p1[i],
2909                            p2[i],
2910                        );
2911                        move |s: usize, e: usize| {
2912                            vbit_range2_a8w8(
2913                                bytes,
2914                                vbit_offsets,
2915                                x1,
2916                                x2,
2917                                a1,
2918                                a2,
2919                                rows,
2920                                cols,
2921                                o1,
2922                                o2,
2923                                s,
2924                                e,
2925                            )
2926                        }
2927                    });
2928                    let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2929                        std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2930                    pool.run_many(&parts);
2931                }
2932                return;
2933            }
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, o1, o2) = (ts[i].cols() / GROUP_SIZE, p1[i], p2[i]);
2939                    move |s: usize, e: usize| {
2940                        q4_range2_f32(packed, scales, gpr, x1, x2, o1, o2, s, e)
2941                    }
2942                });
2943                let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2944                    std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2945                pool.run_many(&parts);
2946            } else {
2947                let closures: [_; N] = std::array::from_fn(|i| {
2948                    let Self::Mapped { vbit_offsets, .. } = ts[i] else {
2949                        unreachable!()
2950                    };
2951                    let (bytes, rows, cols, o1, o2) = (
2952                        ts[i].quant_bytes(),
2953                        ts[i].rows(),
2954                        ts[i].cols(),
2955                        p1[i],
2956                        p2[i],
2957                    );
2958                    move |s: usize, e: usize| {
2959                        vbit_range2_f32(bytes, vbit_offsets, x1, x2, rows, cols, o1, o2, s, e)
2960                    }
2961                });
2962                let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2963                    std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2964                pool.run_many(&parts);
2965            }
2966            return;
2967        }
2968
2969        if uniform_f32 {
2970            let p1: [SendMut; N] = std::array::from_fn(|i| SendMut(o1s[i].as_mut_ptr()));
2971            let p2: [SendMut; N] = std::array::from_fn(|i| SendMut(o2s[i].as_mut_ptr()));
2972            let closures: [_; N] = std::array::from_fn(|i| {
2973                let Self::F32 { data, cols, .. } = ts[i] else {
2974                    unreachable!()
2975                };
2976                let (o1, o2) = (p1[i], p2[i]);
2977                move |start: usize, end: usize| {
2978                    for o in start..end {
2979                        let row = &data[o * cols..(o + 1) * cols];
2980                        let (mut s1, mut s2) = (0.0f32, 0.0f32);
2981                        for j in 0..*cols {
2982                            s1 += row[j] * x1[j];
2983                            s2 += row[j] * x2[j];
2984                        }
2985                        // SAFETY: disjoint (tensor, row) cells per worker.
2986                        unsafe {
2987                            *o1.at(o) = s1;
2988                            *o2.at(o) = s2;
2989                        }
2990                    }
2991                }
2992            });
2993            let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2994                std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2995            pool.run_many(&parts);
2996            return;
2997        }
2998
2999        struct Ctx<'a> {
3000            bytes: &'a [u8],
3001            row_scale: &'a [f32],
3002            cols: usize,
3003            xs1: std::borrow::Cow<'a, [f32]>,
3004            xs2: std::borrow::Cow<'a, [f32]>,
3005        }
3006        let ctxs: [Ctx<'_>; N] = std::array::from_fn(|i| {
3007            let Self::Mapped {
3008                dtype,
3009                cols,
3010                row_scale,
3011                col_field,
3012                ..
3013            } = ts[i]
3014            else {
3015                unreachable!()
3016            };
3017            Ctx {
3018                bytes: ts[i].quant_bytes(),
3019                row_scale,
3020                cols: *cols,
3021                xs1: prescale(x1, col_field, *dtype),
3022                xs2: prescale(x2, col_field, *dtype),
3023            }
3024        });
3025        let p1: [SendMut; N] = std::array::from_fn(|i| SendMut(o1s[i].as_mut_ptr()));
3026        let p2: [SendMut; N] = std::array::from_fn(|i| SendMut(o2s[i].as_mut_ptr()));
3027        #[cfg(target_arch = "aarch64")]
3028        if sdot_enabled() {
3029            let acts: [(SplitAct, SplitAct); N] =
3030                std::array::from_fn(|i| (split_act(&ctxs[i].xs1), split_act(&ctxs[i].xs2)));
3031            let closures: [_; N] = std::array::from_fn(|i| {
3032                let (c, a, o1, o2) = (&ctxs[i], &acts[i], p1[i], p2[i]);
3033                move |start: usize, end: usize| {
3034                    q8_range2_sdot(c.bytes, c.row_scale, &a.0, &a.1, c.cols, o1, o2, start, end)
3035                }
3036            });
3037            let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
3038                std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
3039            pool.run_many(&parts);
3040            return;
3041        }
3042        #[cfg(target_arch = "x86_64")]
3043        if avx2_a8w8_enabled() {
3044            let acts: [(SplitAct, SplitAct); N] =
3045                std::array::from_fn(|i| (split_act(&ctxs[i].xs1), split_act(&ctxs[i].xs2)));
3046            let closures: [_; N] = std::array::from_fn(|i| {
3047                let (c, a, o1, o2) = (&ctxs[i], &acts[i], p1[i], p2[i]);
3048                move |start: usize, end: usize| {
3049                    q8_range2_avx2(c.bytes, c.row_scale, &a.0, &a.1, c.cols, o1, o2, start, end)
3050                }
3051            });
3052            let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
3053                std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
3054            pool.run_many(&parts);
3055            return;
3056        }
3057        let closures: [_; N] = std::array::from_fn(|i| {
3058            let (c, o1, o2) = (&ctxs[i], p1[i], p2[i]);
3059            move |start: usize, end: usize| {
3060                q8_range2_f32(
3061                    c.bytes,
3062                    c.row_scale,
3063                    &c.xs1,
3064                    &c.xs2,
3065                    c.cols,
3066                    o1,
3067                    o2,
3068                    start,
3069                    end,
3070                )
3071            }
3072        });
3073        let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
3074            std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
3075        pool.run_many(&parts);
3076    }
3077
3078    /// Fused gate+up matvec with SiLU·mul: for each row r, computes
3079    /// `silu(gate·x) * (up·x)` and writes to `out[r]`. ONE pool dispatch,
3080    /// no intermediate g/u buffers, no separate silu pass. Falls back
3081    /// (returns false) for unsupported dtype combos.
3082    pub fn matvec_silu_mul(
3083        gate: &QTensor,
3084        up: &QTensor,
3085        x: &[f32],
3086        out: &mut [f32],
3087        pool: Option<&Pool>,
3088    ) -> bool {
3089        Self::matvec_silu_mul_limited(gate, up, x, out, 0.0, pool)
3090    }
3091
3092    /// Fused gate+up+SiLU with the GLM asymmetrical clamp.  `limit == 0`
3093    /// preserves the historical unclamped helper; a positive limit clamps
3094    /// `up` to both sides and `gate` only from above, matching the GLM
3095    /// SwiGLU reference.  Keeping the limit in the row kernel avoids the two
3096    /// intermediate vectors and the extra combine pass on the Q2TP experts.
3097    pub fn matvec_silu_mul_limited(
3098        gate: &QTensor,
3099        up: &QTensor,
3100        x: &[f32],
3101        out: &mut [f32],
3102        limit: f32,
3103        pool: Option<&Pool>,
3104    ) -> bool {
3105        if gate.has_prism_contract() || up.has_prism_contract() {
3106            // The fused gate/up kernels consume x directly.  Prism requires
3107            // a per-matrix signed FWHT, so the caller must use two ordinary
3108            // descriptor-aware matvecs instead of an unrotated fast path.
3109            return false;
3110        }
3111        let inter = gate.rows();
3112        debug_assert_eq!(up.rows(), inter);
3113        debug_assert_eq!(out.len(), inter);
3114        debug_assert_eq!(gate.cols(), up.cols());
3115        if !a8w8_enabled() {
3116            return false;
3117        }
3118        let act = split_act(x);
3119        let act = &act;
3120        let x_ref = x;
3121        let out_addr = SendMut(out.as_mut_ptr());
3122
3123        match (gate, up) {
3124            // Q4Block gate + Q4Block up (most common mobile q4 models)
3125            (
3126                Self::Mapped {
3127                    dtype: TensorDtype::Q4Block,
3128                    ..
3129                },
3130                Self::Mapped {
3131                    dtype: TensorDtype::Q4Block,
3132                    ..
3133                },
3134            ) => {
3135                let (gp, gs) = q4_split(gate.quant_bytes(), gate.rows(), gate.cols());
3136                let (up_p, up_s) = q4_split(up.quant_bytes(), up.rows(), up.cols());
3137                let gpr = gate.cols() / GROUP_SIZE;
3138                let cols = gate.cols();
3139                let run = move |start: usize, end: usize| {
3140                    for r in start..end {
3141                        let mut gv = dot_q4_row_i8(gp, gs, r * gpr, gpr, &act.xq) * act.sx;
3142                        let mut uv = dot_q4_row_i8(up_p, up_s, r * gpr, gpr, &act.xq) * act.sx;
3143                        for &(j, xv) in &act.outliers {
3144                            let flat = r * cols + j;
3145                            let gb = gp[flat / 2];
3146                            let gn = if flat & 1 == 0 { gb & 0x0F } else { gb >> 4 };
3147                            let gsc = f16_to_f32(u16::from_le_bytes([
3148                                gs[(flat / GROUP_SIZE) * 2],
3149                                gs[(flat / GROUP_SIZE) * 2 + 1],
3150                            ]));
3151                            gv += ((gn as i32 - 8) as f32) * gsc * xv;
3152                            let ub = up_p[flat / 2];
3153                            let un = if flat & 1 == 0 { ub & 0x0F } else { ub >> 4 };
3154                            let usc = f16_to_f32(u16::from_le_bytes([
3155                                up_s[(flat / GROUP_SIZE) * 2],
3156                                up_s[(flat / GROUP_SIZE) * 2 + 1],
3157                            ]));
3158                            uv += ((un as i32 - 8) as f32) * usc * xv;
3159                        }
3160                        // SAFETY: disjoint row ranges per worker.
3161                        unsafe { *out_addr.at(r) = silu_mul_limited(gv, uv, limit) };
3162                    }
3163                };
3164                dispatch_rows(pool, inter, &run);
3165                true
3166            }
3167            // Q4Tiled gate + Q4Tiled up — one row pass, both tile
3168            // streams sequential, silu·mul fused (same per-row math as
3169            // `q4t_matvec`).
3170            (
3171                Self::Mapped {
3172                    dtype: TensorDtype::Q4Tiled,
3173                    ..
3174                },
3175                Self::Mapped {
3176                    dtype: TensorDtype::Q4Tiled,
3177                    ..
3178                },
3179            ) => {
3180                let g_bytes = gate.quant_bytes();
3181                let u_bytes = up.quant_bytes();
3182                let gpr = gate.cols() / GROUP_SIZE;
3183                let run = move |start: usize, end: usize| {
3184                    for r in start..end {
3185                        let mut gv = dot_q4t_row_i8(g_bytes, r, gpr, &act.xq) * act.sx;
3186                        let mut uv = dot_q4t_row_i8(u_bytes, r, gpr, &act.xq) * act.sx;
3187                        for &(j, xv) in &act.outliers {
3188                            let (w, s) = q4t_outlier(g_bytes, r, gpr, j);
3189                            gv += w * s * xv;
3190                            let (w, s) = q4t_outlier(u_bytes, r, gpr, j);
3191                            uv += w * s * xv;
3192                        }
3193                        // SAFETY: disjoint row ranges per worker.
3194                        unsafe { *out_addr.at(r) = silu_mul_limited(gv, uv, limit) };
3195                    }
3196                };
3197                dispatch_rows(pool, inter, &run);
3198                true
3199            }
3200            // Q4TiledP gate + Q4TiledP up — the same fused row pass, with
3201            // each row's two ladders built once and spent on both streams.
3202            (
3203                Self::Mapped {
3204                    dtype: TensorDtype::Q4TiledP,
3205                    ..
3206                },
3207                Self::Mapped {
3208                    dtype: TensorDtype::Q4TiledP,
3209                    ..
3210                },
3211            ) => {
3212                let cols = gate.cols();
3213                let gpr = cols / GROUP_SIZE;
3214                let gv_view = Q4tpView::new(gate.quant_bytes(), inter, cols);
3215                let uv_view = Q4tpView::new(up.quant_bytes(), inter, cols);
3216                let run = |start: usize, end: usize| {
3217                    with_krows(gpr, |gsc, usc| {
3218                        for r in start..end {
3219                            gv_view.scales_into(r, gpr, gsc);
3220                            uv_view.scales_into(r, gpr, usc);
3221                            let mut gv =
3222                                dot_q4tp_row_i8(gv_view.nib, r, gpr, &act.xq, gsc) * act.sx;
3223                            let mut uv =
3224                                dot_q4tp_row_i8(uv_view.nib, r, gpr, &act.xq, usc) * act.sx;
3225                            for &(j, xv) in &act.outliers {
3226                                let (w, s) = q4tp_outlier(gv_view.nib, r, gpr, j, gsc);
3227                                gv += w * s * xv;
3228                                let (w, s) = q4tp_outlier(uv_view.nib, r, gpr, j, usc);
3229                                uv += w * s * xv;
3230                            }
3231                            // SAFETY: disjoint row ranges per worker.
3232                            unsafe { *out_addr.at(r) = silu_mul_limited(gv, uv, limit) };
3233                        }
3234                    });
3235                };
3236                dispatch_rows(pool, inter, &run);
3237                true
3238            }
3239            // Q1 gate + Q1 up — one row pass over both sign streams,
3240            // silu·mul fused (the per-row math of `q1_range_a8w8`); the
3241            // activation group sums are shared by both streams. Without
3242            // this arm a q1 dense FFN paid two dispatches + a combine
3243            // loop — the exact barrier this function exists to remove.
3244            (
3245                Self::Mapped {
3246                    dtype: TensorDtype::Q1,
3247                    ..
3248                },
3249                Self::Mapped {
3250                    dtype: TensorDtype::Q1,
3251                    ..
3252                },
3253            ) => {
3254                let g_bytes = gate.quant_bytes();
3255                let u_bytes = up.quant_bytes();
3256                let gpr = gate.cols() / GROUP_SIZE;
3257                let gsum = q1_group_sums(&act.xq, gpr);
3258                let gsum = &gsum;
3259                let run = move |start: usize, end: usize| {
3260                    for r in start..end {
3261                        let mut gv = dot_q1_row_i8(g_bytes, r, gpr, &act.xq, gsum) * act.sx;
3262                        let mut uv = dot_q1_row_i8(u_bytes, r, gpr, &act.xq, gsum) * act.sx;
3263                        for &(j, xv) in &act.outliers {
3264                            let (w, s) = q1_outlier(g_bytes, r, gpr, j);
3265                            gv += w * s * xv;
3266                            let (w, s) = q1_outlier(u_bytes, r, gpr, j);
3267                            uv += w * s * xv;
3268                        }
3269                        // SAFETY: disjoint row ranges per worker.
3270                        unsafe { *out_addr.at(r) = silu_mul_limited(gv, uv, limit) };
3271                    }
3272                };
3273                dispatch_rows(pool, inter, &run);
3274                true
3275            }
3276            // Q2TiledP gate + Q2TiledP up — the 2-bit expert pair (MoE
3277            // FFNs of the W2 class): one row pass, both ladders built
3278            // once, integer code dots with shared group sums.
3279            (
3280                Self::Mapped {
3281                    dtype: TensorDtype::Q2TiledP,
3282                    ..
3283                },
3284                Self::Mapped {
3285                    dtype: TensorDtype::Q2TiledP,
3286                    ..
3287                },
3288            ) => {
3289                let cols = gate.cols();
3290                let gpr = cols / GROUP_SIZE;
3291                let gv_view = Q4tpView::new_q2(gate.quant_bytes(), inter, cols);
3292                let uv_view = Q4tpView::new_q2(up.quant_bytes(), inter, cols);
3293                let gsum = q1_group_sums(&act.xq, gpr);
3294                let gsum = &gsum;
3295                let run = move |start: usize, end: usize| {
3296                    with_krows(gpr, |gsc, usc| {
3297                        for r in start..end {
3298                            gv_view.scales_into(r, gpr, gsc);
3299                            uv_view.scales_into(r, gpr, usc);
3300                            let mut gv =
3301                                dot_q2tp_row_i8(gv_view.nib, r, gpr, &act.xq, gsum, gsc) * act.sx;
3302                            let mut uv =
3303                                dot_q2tp_row_i8(uv_view.nib, r, gpr, &act.xq, gsum, usc) * act.sx;
3304                            for &(j, xv) in &act.outliers {
3305                                let (w, s) = q2tp_outlier(gv_view.nib, r, gpr, j, gsc);
3306                                gv += w * s * xv;
3307                                let (w, s) = q2tp_outlier(uv_view.nib, r, gpr, j, usc);
3308                                uv += w * s * xv;
3309                            }
3310                            // SAFETY: disjoint row ranges per worker.
3311                            unsafe { *out_addr.at(r) = silu_mul_limited(gv, uv, limit) };
3312                        }
3313                    });
3314                };
3315                dispatch_rows(pool, inter, &run);
3316                true
3317            }
3318            // Q8Row gate + Q8Row up — one row pass over both i8 streams.
3319            // Q8_2f stays out on purpose: its column field prescales the
3320            // activations PER TENSOR, which breaks this fn's shared
3321            // split_act contract — it keeps the two-dispatch path.
3322            (
3323                Self::Mapped {
3324                    dtype: TensorDtype::Q8Row,
3325                    row_scale: g_rs,
3326                    ..
3327                },
3328                Self::Mapped {
3329                    dtype: TensorDtype::Q8Row,
3330                    row_scale: u_rs,
3331                    ..
3332                },
3333            ) => {
3334                let g_bytes = gate.quant_bytes();
3335                let u_bytes = up.quant_bytes();
3336                let cols = gate.cols();
3337                let run = move |start: usize, end: usize| {
3338                    for r in start..end {
3339                        let gv = q8_row_dot(&g_bytes[r * cols..(r + 1) * cols], act) * g_rs[r];
3340                        let uv = q8_row_dot(&u_bytes[r * cols..(r + 1) * cols], act) * u_rs[r];
3341                        // SAFETY: disjoint row ranges per worker.
3342                        unsafe { *out_addr.at(r) = silu_mul_limited(gv, uv, limit) };
3343                    }
3344                };
3345                dispatch_rows(pool, inter, &run);
3346                true
3347            }
3348            // Q1T gate + Q1T up
3349            (
3350                Self::Mapped {
3351                    dtype: TensorDtype::Q1T,
3352                    ..
3353                },
3354                Self::Mapped {
3355                    dtype: TensorDtype::Q1T,
3356                    ..
3357                },
3358            ) => {
3359                const TILE: usize = cortiq_core::quant::Q1T_TILE;
3360                let g_bytes = gate.quant_bytes();
3361                let u_bytes = up.quant_bytes();
3362                let gpr = gate.cols() / GROUP_SIZE;
3363                let (g_rp, g_ent, g_ov) = q1t_overlay(g_bytes, inter * gpr * TILE, inter);
3364                let (u_rp, u_ent, u_ov) = q1t_overlay(u_bytes, inter * gpr * TILE, inter);
3365                let run = move |start: usize, end: usize| {
3366                    for r in start..end {
3367                        let mut gv = q1t_dot_row_i8(g_bytes, r, gpr, &act.xq) * act.sx;
3368                        let mut uv = q1t_dot_row_i8(u_bytes, r, gpr, &act.xq) * act.sx;
3369                        for &(j, xv) in &act.outliers {
3370                            gv += q1t_base_weight(g_bytes, r, gpr, j) * xv;
3371                            uv += q1t_base_weight(u_bytes, r, gpr, j) * xv;
3372                        }
3373                        gv += q1t_row_outlier_correction(g_bytes, r, g_rp, g_ent, g_ov, x_ref);
3374                        uv += q1t_row_outlier_correction(u_bytes, r, u_rp, u_ent, u_ov, x_ref);
3375                        // SAFETY: disjoint row ranges per worker.
3376                        unsafe { *out_addr.at(r) = silu_mul_limited(gv, uv, limit) };
3377                    }
3378                };
3379                dispatch_rows(pool, inter, &run);
3380                true
3381            }
3382            _ => false,
3383        }
3384    }
3385
3386    /// Every routed expert's fused gate/up/SiLU under ONE pool dispatch.
3387    ///
3388    /// The per-expert path pays a pool barrier per expert per stage: at 9
3389    /// experts over 40 layers that is ~720 barriers a token, and a decode
3390    /// profile of Qwen3.6-35B-A3B showed the pool parked in
3391    /// `psynch_cvwait` about twice as long as it spent computing. Laying
3392    /// every expert's rows end-to-end in one virtual row space collapses
3393    /// the stage to a single dispatch. The per-row body is the
3394    /// single-expert q4tp arm verbatim, so outputs are bit-identical.
3395    ///
3396    /// `false` = something is outside the fused q4tp kernel (dtype, shape,
3397    /// or a transformed tensor); the caller walks the ordinary per-expert
3398    /// path. Float activations use the same exact scalar rows, still fused
3399    /// under one pool dispatch.
3400    pub fn moe_gate_up_many(
3401        pairs: &[(&QTensor, &QTensor)],
3402        x: &[f32],
3403        outs: &mut [Vec<f32>],
3404        pool: Option<&Pool>,
3405    ) -> bool {
3406        Self::moe_gate_up_many_limited(pairs, x, outs, 0.0, pool)
3407    }
3408
3409    /// Batched gate/up/SiLU with the optional GLM clamp.  The public legacy
3410    /// helper above keeps its historical unclamped semantics; callers that
3411    /// implement a reference with a positive SwiGLU limit use this variant.
3412    pub fn moe_gate_up_many_limited(
3413        pairs: &[(&QTensor, &QTensor)],
3414        x: &[f32],
3415        outs: &mut [Vec<f32>],
3416        limit: f32,
3417        pool: Option<&Pool>,
3418    ) -> bool {
3419        if pairs.is_empty() || pairs.len() != outs.len() {
3420            return false;
3421        }
3422        if !a8w8_enabled() {
3423            if limit > 0.0 {
3424                // The exact-row fallback applies no SwiGLU clamp; the
3425                // caller's per-expert path carries it instead.
3426                return false;
3427            }
3428            let groups = vec![vec![0]; pairs.len()];
3429            return Self::moe_gate_up_rows(pairs, &groups, x, outs, pool);
3430        }
3431        let inter = pairs[0].0.rows();
3432        let cols = pairs[0].0.cols();
3433        if cols % GROUP_SIZE != 0 {
3434            return false;
3435        }
3436        let gpr = cols / GROUP_SIZE;
3437        // Uniform layout across every routed pair: q4tp, or the 2-bit
3438        // profile's q2tp gate/up (the W2 class). Mixed sets refuse.
3439        let q2 = matches!(
3440            pairs[0].0,
3441            Self::Mapped {
3442                dtype: TensorDtype::Q2TiledP,
3443                ..
3444            }
3445        );
3446        let want = if q2 {
3447            TensorDtype::Q2TiledP
3448        } else {
3449            TensorDtype::Q4TiledP
3450        };
3451        let mut views = Vec::with_capacity(pairs.len() * 2);
3452        for ((g, u), o) in pairs.iter().zip(outs.iter()) {
3453            let both = matches!(g, Self::Mapped { dtype, .. } if *dtype == want)
3454                && matches!(u, Self::Mapped { dtype, .. } if *dtype == want);
3455            if !both
3456                || g.rows() != inter
3457                || u.rows() != inter
3458                || g.cols() != cols
3459                || u.cols() != cols
3460                || o.len() != inter
3461            {
3462                return false;
3463            }
3464            let mk = if q2 { Q4tpView::new_q2 } else { Q4tpView::new };
3465            views.push(mk(g.quant_bytes(), inter, cols));
3466            views.push(mk(u.quant_bytes(), inter, cols));
3467        }
3468        let act = split_act(x);
3469        let gsum = if q2 {
3470            q1_group_sums(&act.xq, gpr)
3471        } else {
3472            Vec::new()
3473        };
3474        let (act, gsum) = (&act, &gsum);
3475        let ptrs: Vec<SendMut> = outs.iter_mut().map(|o| SendMut(o.as_mut_ptr())).collect();
3476        let (views, ptrs) = (&views, &ptrs);
3477        let run = |start: usize, end: usize| {
3478            with_krows(gpr, |gsc, usc| {
3479                for flat in start..end {
3480                    let (e, r) = (flat / inter, flat % inter);
3481                    let gv_view = &views[e * 2];
3482                    let uv_view = &views[e * 2 + 1];
3483                    gv_view.scales_into(r, gpr, gsc);
3484                    uv_view.scales_into(r, gpr, usc);
3485                    let (mut gv, mut uv) = if q2 {
3486                        (
3487                            dot_q2tp_row_i8(gv_view.nib, r, gpr, &act.xq, gsum, gsc) * act.sx,
3488                            dot_q2tp_row_i8(uv_view.nib, r, gpr, &act.xq, gsum, usc) * act.sx,
3489                        )
3490                    } else {
3491                        (
3492                            dot_q4tp_row_i8(gv_view.nib, r, gpr, &act.xq, gsc) * act.sx,
3493                            dot_q4tp_row_i8(uv_view.nib, r, gpr, &act.xq, usc) * act.sx,
3494                        )
3495                    };
3496                    for &(j, xv) in &act.outliers {
3497                        let (og, ou) = if q2 {
3498                            (
3499                                q2tp_outlier(gv_view.nib, r, gpr, j, gsc),
3500                                q2tp_outlier(uv_view.nib, r, gpr, j, usc),
3501                            )
3502                        } else {
3503                            (
3504                                q4tp_outlier(gv_view.nib, r, gpr, j, gsc),
3505                                q4tp_outlier(uv_view.nib, r, gpr, j, usc),
3506                            )
3507                        };
3508                        gv += og.0 * og.1 * xv;
3509                        uv += ou.0 * ou.1 * xv;
3510                    }
3511                    // SAFETY: one worker owns each (expert, row) pair.
3512                    unsafe { *ptrs[e].at(r) = silu_mul_limited(gv, uv, limit) };
3513                }
3514            });
3515        };
3516        dispatch_rows(pool, pairs.len() * inter, &run);
3517        true
3518    }
3519
3520    /// Every routed expert's down projection, weighted and summed into
3521    /// `out`, under ONE pool dispatch.
3522    ///
3523    /// Partitioned by OUTPUT row rather than by expert: each row is owned
3524    /// by a single worker, so the experts are summed in the caller's order
3525    /// — the same sequence of f32 adds the serial `out[i] += w·eo[i]` loop
3526    /// performs, hence bit-identical. Partitioning by expert instead would
3527    /// race on the shared accumulator.
3528    pub fn moe_down_many(
3529        downs: &[&QTensor],
3530        gs: &[Vec<f32>],
3531        weights: &[f32],
3532        out: &mut [f32],
3533        pool: Option<&Pool>,
3534    ) -> bool {
3535        if downs.is_empty() || downs.len() != gs.len() || downs.len() != weights.len() {
3536            return false;
3537        }
3538        if !a8w8_enabled() {
3539            let mut terms = vec![vec![0.0; out.len()]; downs.len()];
3540            if !Self::moe_down_rows(downs, &vec![1; downs.len()], gs, &mut terms, pool) {
3541                return false;
3542            }
3543            out.fill(0.0);
3544            for (row, &w) in terms.iter().zip(weights) {
3545                for (o, &v) in out.iter_mut().zip(row) {
3546                    *o += w * v;
3547                }
3548            }
3549            return true;
3550        }
3551        let rows = out.len();
3552        let cols = downs[0].cols();
3553        if cols % GROUP_SIZE != 0 {
3554            return false;
3555        }
3556        let gpr = cols / GROUP_SIZE;
3557        let mut views = Vec::with_capacity(downs.len());
3558        for (d, g) in downs.iter().zip(gs.iter()) {
3559            if !matches!(
3560                d,
3561                Self::Mapped {
3562                    dtype: TensorDtype::Q4TiledP,
3563                    ..
3564                }
3565            ) || d.rows() != rows
3566                || d.cols() != cols
3567                || g.len() != cols
3568            {
3569                return false;
3570            }
3571            views.push(Q4tpView::new(d.quant_bytes(), rows, cols));
3572        }
3573        // One int8 split per expert — the activation vectors differ.
3574        let acts: Vec<SplitAct> = gs.iter().map(|g| split_act(g)).collect();
3575        // Partitioned by OUTPUT row, with the experts folded inside: each
3576        // row is owned by one worker, so they are summed in the caller's
3577        // order — the same f32 sequence the serial `out[i] += w·eo[i]`
3578        // loop produces. Partitioning by expert instead would either race
3579        // on the accumulator or need a scratch plane and a second pass;
3580        // measured, that variant was a wash, so this keeps the simpler
3581        // shape.
3582        let out_addr = SendMut(out.as_mut_ptr());
3583        let (views, acts, weights) = (&views, &acts, &weights);
3584        let run = |start: usize, end: usize| {
3585            with_krow(gpr, |sc| {
3586                for r in start..end {
3587                    let mut acc = 0f32;
3588                    for (e, v) in views.iter().enumerate() {
3589                        v.scales_into(r, gpr, sc);
3590                        let a = &acts[e];
3591                        let mut d = dot_q4tp_row_i8(v.nib, r, gpr, &a.xq, sc) * a.sx;
3592                        for &(j, xv) in &a.outliers {
3593                            let (w, s) = q4tp_outlier(v.nib, r, gpr, j, sc);
3594                            d += w * s * xv;
3595                        }
3596                        acc += weights[e] * d;
3597                    }
3598                    // SAFETY: disjoint row ranges per worker.
3599                    unsafe { *out_addr.at(r) = acc };
3600                }
3601            });
3602        };
3603        dispatch_rows(pool, rows, &run);
3604        true
3605    }
3606
3607    /// `moe_gate_up_many` for SEVERAL tokens at once, decode-exact: expert
3608    /// `e` (`pairs[e]`, q4tp) serves the tokens `groups[e]` (row indices
3609    /// into `xs`, each `cols` wide). Every (expert, token) output is
3610    /// bit-identical to `moe_gate_up_many` run on that token alone — the
3611    /// same int8 activation split, VNNI dots, outlier terms and inline
3612    /// SiLU — while each weight row is read once for all the tokens routed
3613    /// to its expert (the speculative verify's expert sharing). `outs` is
3614    /// flat in (expert, token-of-group) order. False = not covered (not
3615    /// q4tp): the caller takes the per-token path. With float activations,
3616    /// the exact scalar row kernel replaces the int8 dot without changing
3617    /// the shared dispatch or route-order reduction.
3618    pub fn moe_gate_up_rows(
3619        pairs: &[(&QTensor, &QTensor)],
3620        groups: &[Vec<usize>],
3621        xs: &[f32],
3622        outs: &mut [Vec<f32>],
3623        pool: Option<&Pool>,
3624    ) -> bool {
3625        if pairs.is_empty() || pairs.len() != groups.len() {
3626            return false;
3627        }
3628        let inter = pairs[0].0.rows();
3629        let cols = pairs[0].0.cols();
3630        let n_pairs: usize = groups.iter().map(|g| g.len()).sum();
3631        if cols == 0 || cols % GROUP_SIZE != 0 || outs.len() != n_pairs || xs.len() % cols != 0 {
3632            return false;
3633        }
3634        let b = xs.len() / cols;
3635        let gpr = cols / GROUP_SIZE;
3636        let mut views = Vec::with_capacity(pairs.len() * 2);
3637        for (g, u) in pairs {
3638            let q4tp = |t: &QTensor| {
3639                matches!(
3640                    t,
3641                    Self::Mapped {
3642                        dtype: TensorDtype::Q4TiledP,
3643                        ..
3644                    }
3645                )
3646            };
3647            if g.has_prism_contract()
3648                || u.has_prism_contract()
3649                || !q4tp(g)
3650                || !q4tp(u)
3651                || g.rows() != inter
3652                || u.rows() != inter
3653                || g.cols() != cols
3654                || u.cols() != cols
3655            {
3656                return false;
3657            }
3658            views.push(Q4tpView::new(g.quant_bytes(), inter, cols));
3659            views.push(Q4tpView::new(u.quant_bytes(), inter, cols));
3660        }
3661        if outs.iter().any(|o| o.len() != inter) || groups.iter().flatten().any(|&t| t >= b) {
3662            return false;
3663        }
3664        let quantized = a8w8_enabled();
3665        let acts: Vec<SplitAct> = if quantized {
3666            (0..b)
3667                .map(|t| split_act(&xs[t * cols..(t + 1) * cols]))
3668                .collect()
3669        } else {
3670            Vec::new()
3671        };
3672        let mut offs = Vec::with_capacity(groups.len());
3673        let mut o = 0usize;
3674        for g in groups {
3675            offs.push(o);
3676            o += g.len();
3677        }
3678        let ptrs: Vec<SendMut> = outs.iter_mut().map(|o| SendMut(o.as_mut_ptr())).collect();
3679        let (views, ptrs, acts, offs) = (&views, &ptrs, &acts, &offs);
3680        let run = |start: usize, end: usize| {
3681            let (mut gsc, mut usc) = (vec![0f32; gpr], vec![0f32; gpr]);
3682            for flat in start..end {
3683                let (e, r) = (flat / inter, flat % inter);
3684                let (gv_view, uv_view) = (&views[e * 2], &views[e * 2 + 1]);
3685                gv_view.scales_into(r, gpr, &mut gsc);
3686                uv_view.scales_into(r, gpr, &mut usc);
3687                for (k, &t) in groups[e].iter().enumerate() {
3688                    if !quantized {
3689                        let x = &xs[t * cols..(t + 1) * cols];
3690                        let gv = q4tp_row_exact(gv_view.nib, r, gpr, x, &gsc);
3691                        let uv = q4tp_row_exact(uv_view.nib, r, gpr, x, &usc);
3692                        unsafe { *ptrs[offs[e] + k].at(r) = (gv / (1.0 + (-gv).exp())) * uv };
3693                        continue;
3694                    }
3695                    let act = &acts[t];
3696                    let mut gv = dot_q4tp_row_i8(gv_view.nib, r, gpr, &act.xq, &gsc) * act.sx;
3697                    let mut uv = dot_q4tp_row_i8(uv_view.nib, r, gpr, &act.xq, &usc) * act.sx;
3698                    for &(j, xv) in &act.outliers {
3699                        let og = q4tp_outlier(gv_view.nib, r, gpr, j, &gsc);
3700                        let ou = q4tp_outlier(uv_view.nib, r, gpr, j, &usc);
3701                        gv += og.0 * og.1 * xv;
3702                        uv += ou.0 * ou.1 * xv;
3703                    }
3704                    let silu_g = gv / (1.0 + (-gv).exp());
3705                    // SAFETY: one worker owns each (expert, row) cell of
3706                    // every output of the expert's group.
3707                    unsafe { *ptrs[offs[e] + k].at(r) = silu_g * uv };
3708                }
3709            }
3710        };
3711        dispatch_rows(pool, pairs.len() * inter, &run);
3712        true
3713    }
3714
3715    /// The per-(expert, token) down terms `moe_down_many` weights and sums,
3716    /// for SEVERAL tokens: `outs[p][o] = down_e[o] · gs[p]` (int8 split of
3717    /// `gs[p]`, VNNI dot, outlier terms — bit-identical to that kernel's
3718    /// `d`), each down row read once for its expert's whole group. The
3719    /// caller sums `w·d` per token in its route order, which reproduces
3720    /// `moe_down_many`'s f32 sequence exactly. Layout as `moe_gate_up_rows`.
3721    pub fn moe_down_rows(
3722        downs: &[&QTensor],
3723        group_lens: &[usize],
3724        gs: &[Vec<f32>],
3725        outs: &mut [Vec<f32>],
3726        pool: Option<&Pool>,
3727    ) -> bool {
3728        if downs.is_empty() || downs.len() != group_lens.len() {
3729            return false;
3730        }
3731        let rows = downs[0].rows();
3732        let cols = downs[0].cols();
3733        let n_pairs: usize = group_lens.iter().sum();
3734        if cols == 0 || cols % GROUP_SIZE != 0 || gs.len() != n_pairs || outs.len() != n_pairs {
3735            return false;
3736        }
3737        let gpr = cols / GROUP_SIZE;
3738        let mut views = Vec::with_capacity(downs.len());
3739        for d in downs {
3740            if d.has_prism_contract()
3741                || !matches!(
3742                    d,
3743                    Self::Mapped {
3744                        dtype: TensorDtype::Q4TiledP,
3745                        ..
3746                    }
3747                )
3748                || d.rows() != rows
3749                || d.cols() != cols
3750            {
3751                return false;
3752            }
3753            views.push(Q4tpView::new(d.quant_bytes(), rows, cols));
3754        }
3755        if gs.iter().any(|g| g.len() != cols) || outs.iter().any(|o| o.len() != rows) {
3756            return false;
3757        }
3758        let quantized = a8w8_enabled();
3759        let acts: Vec<SplitAct> = if quantized {
3760            gs.iter().map(|g| split_act(g)).collect()
3761        } else {
3762            Vec::new()
3763        };
3764        let mut offs = Vec::with_capacity(group_lens.len());
3765        let mut o = 0usize;
3766        for &l in group_lens {
3767            offs.push(o);
3768            o += l;
3769        }
3770        let ptrs: Vec<SendMut> = outs.iter_mut().map(|o| SendMut(o.as_mut_ptr())).collect();
3771        let (views, ptrs, acts, offs) = (&views, &ptrs, &acts, &offs);
3772        let run = |start: usize, end: usize| {
3773            let mut sc = vec![0f32; gpr];
3774            for flat in start..end {
3775                let (e, r) = (flat / rows, flat % rows);
3776                let v = &views[e];
3777                v.scales_into(r, gpr, &mut sc);
3778                for k in 0..group_lens[e] {
3779                    if !quantized {
3780                        let d = q4tp_row_exact(v.nib, r, gpr, &gs[offs[e] + k], &sc);
3781                        unsafe { *ptrs[offs[e] + k].at(r) = d };
3782                        continue;
3783                    }
3784                    let a = &acts[offs[e] + k];
3785                    let mut d = dot_q4tp_row_i8(v.nib, r, gpr, &a.xq, &sc) * a.sx;
3786                    for &(j, xv) in &a.outliers {
3787                        let (w, s) = q4tp_outlier(v.nib, r, gpr, j, &sc);
3788                        d += w * s * xv;
3789                    }
3790                    // SAFETY: one worker owns each (expert, row) cell.
3791                    unsafe { *ptrs[offs[e] + k].at(r) = d };
3792                }
3793            }
3794        };
3795        dispatch_rows(pool, downs.len() * rows, &run);
3796        true
3797    }
3798}
3799
3800/// Batched q8 kernel: same math as qmatvec, the row makes a single
3801/// pass from memory for the whole batch.
3802/// Accelerate CBLAS — the Apple AMX matrix units, the same engine
3803/// llama.cpp's `-ngl 0` prefill rides via ggml-blas.
3804#[cfg(target_os = "macos")]
3805mod accel_blas {
3806    #[link(name = "Accelerate", kind = "framework")]
3807    unsafe extern "C" {
3808        pub fn cblas_sgemm(
3809            order: i32,
3810            trans_a: i32,
3811            trans_b: i32,
3812            m: i32,
3813            n: i32,
3814            k: i32,
3815            alpha: f32,
3816            a: *const f32,
3817            lda: i32,
3818            b: *const f32,
3819            ldb: i32,
3820            beta: f32,
3821            c: *mut f32,
3822            ldc: i32,
3823        );
3824    }
3825}
3826
3827#[cfg(target_os = "macos")]
3828pub(crate) fn accel_gemm_enabled() -> bool {
3829    static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
3830    *ON.get_or_init(|| std::env::var("CMF_ACCEL").map(|v| v != "0").unwrap_or(true))
3831}
3832
3833/// Off macOS the "accel" GEMM is the portable NEON micro-kernel below —
3834/// same entry point, so the batched-attention path opens on mobile.
3835#[cfg(all(target_arch = "aarch64", not(target_os = "macos")))]
3836pub(crate) fn accel_gemm_enabled() -> bool {
3837    true
3838}
3839
3840/// Portable NEON f32 GEMM (row-major, optional Bᵀ): a 4×8 fmla
3841/// micro-kernel with A broadcast against B panels — the mobile stand-in
3842/// for Accelerate in the batched causal attention (QKᵀ and P·V). Not a
3843/// BLAS: shapes here are the attention panels (m ≤ heads·chunk,
3844/// k = head_dim or context), and the goal is removing the per-position
3845/// quadratic wall, not peak GEMM.
3846#[cfg(target_arch = "aarch64")]
3847#[allow(clippy::too_many_arguments)]
3848pub(crate) fn neon_gemm_rm(
3849    m: usize,
3850    n: usize,
3851    k: usize,
3852    alpha: f32,
3853    a: &[f32],
3854    lda: usize,
3855    b_mat: &[f32],
3856    ldb: usize,
3857    b_rows_are_n: bool,
3858    c: &mut [f32],
3859    ldc: usize,
3860) {
3861    debug_assert!(a.len() >= (m - 1) * lda + k);
3862    debug_assert!(c.len() >= (m - 1) * ldc + n);
3863    // SAFETY: bounds asserted above; NEON is baseline on aarch64.
3864    unsafe {
3865        use core::arch::aarch64::*;
3866        let mut i = 0usize;
3867        while i < m {
3868            let mi = (m - i).min(4);
3869            let mut j = 0usize;
3870            while j < n {
3871                let nj = (n - j).min(8);
3872                if mi == 4 && nj == 8 {
3873                    let (mut c0a, mut c0b) = (vdupq_n_f32(0.0), vdupq_n_f32(0.0));
3874                    let (mut c1a, mut c1b) = (vdupq_n_f32(0.0), vdupq_n_f32(0.0));
3875                    let (mut c2a, mut c2b) = (vdupq_n_f32(0.0), vdupq_n_f32(0.0));
3876                    let (mut c3a, mut c3b) = (vdupq_n_f32(0.0), vdupq_n_f32(0.0));
3877                    for p in 0..k {
3878                        let (b0, b1) = if b_rows_are_n {
3879                            // B is [n, k]: column p of Bᵀ = element p of
3880                            // eight consecutive B rows — gathered.
3881                            let base = b_mat.as_ptr().add(j * ldb + p);
3882                            let g = |o: usize| *base.add(o * ldb);
3883                            ([g(0), g(1), g(2), g(3)], [g(4), g(5), g(6), g(7)])
3884                        } else {
3885                            let base = b_mat.as_ptr().add(p * ldb + j);
3886                            (
3887                                [*base, *base.add(1), *base.add(2), *base.add(3)],
3888                                [*base.add(4), *base.add(5), *base.add(6), *base.add(7)],
3889                            )
3890                        };
3891                        let bv0 = vld1q_f32(b0.as_ptr());
3892                        let bv1 = vld1q_f32(b1.as_ptr());
3893                        let a0 = vdupq_n_f32(*a.as_ptr().add(i * lda + p));
3894                        let a1 = vdupq_n_f32(*a.as_ptr().add((i + 1) * lda + p));
3895                        let a2 = vdupq_n_f32(*a.as_ptr().add((i + 2) * lda + p));
3896                        let a3 = vdupq_n_f32(*a.as_ptr().add((i + 3) * lda + p));
3897                        c0a = vfmaq_f32(c0a, a0, bv0);
3898                        c0b = vfmaq_f32(c0b, a0, bv1);
3899                        c1a = vfmaq_f32(c1a, a1, bv0);
3900                        c1b = vfmaq_f32(c1b, a1, bv1);
3901                        c2a = vfmaq_f32(c2a, a2, bv0);
3902                        c2b = vfmaq_f32(c2b, a2, bv1);
3903                        c3a = vfmaq_f32(c3a, a3, bv0);
3904                        c3b = vfmaq_f32(c3b, a3, bv1);
3905                    }
3906                    let al = vdupq_n_f32(alpha);
3907                    for (r, (ca, cb)) in [(c0a, c0b), (c1a, c1b), (c2a, c2b), (c3a, c3b)]
3908                        .iter()
3909                        .enumerate()
3910                    {
3911                        let dst = c.as_mut_ptr().add((i + r) * ldc + j);
3912                        vst1q_f32(dst, vmulq_f32(*ca, al));
3913                        vst1q_f32(dst.add(4), vmulq_f32(*cb, al));
3914                    }
3915                } else {
3916                    for r in 0..mi {
3917                        for q in 0..nj {
3918                            let mut acc = 0f32;
3919                            for p in 0..k {
3920                                let bv = if b_rows_are_n {
3921                                    b_mat[(j + q) * ldb + p]
3922                                } else {
3923                                    b_mat[p * ldb + j + q]
3924                                };
3925                                acc += a[(i + r) * lda + p] * bv;
3926                            }
3927                            c[(i + r) * ldc + j + q] = acc * alpha;
3928                        }
3929                    }
3930                }
3931                j += nj;
3932            }
3933            i += mi;
3934        }
3935    }
3936}
3937
3938/// Off-macOS aarch64: the batched attention rides the NEON micro-GEMM.
3939#[cfg(all(target_arch = "aarch64", not(target_os = "macos")))]
3940#[allow(clippy::too_many_arguments)]
3941pub(crate) fn sgemm_rm(
3942    m: usize,
3943    n: usize,
3944    k: usize,
3945    alpha: f32,
3946    a: &[f32],
3947    lda: usize,
3948    b_mat: &[f32],
3949    ldb: usize,
3950    b_rows_are_n: bool,
3951    c: &mut [f32],
3952    ldc: usize,
3953) {
3954    neon_gemm_rm(m, n, k, alpha, a, lda, b_mat, ldb, b_rows_are_n, c, ldc);
3955}
3956
3957/// Row-major f32 GEMM, exposed for offline tools (the AWNP pass builds a
3958/// per-layer projection and applies it to every expert; a naive triple loop
3959/// would turn a two-minute job into half an hour).
3960#[allow(clippy::too_many_arguments)]
3961pub fn sgemm_public(
3962    m: usize,
3963    n: usize,
3964    k: usize,
3965    alpha: f32,
3966    a: &[f32],
3967    lda: usize,
3968    b_mat: &[f32],
3969    ldb: usize,
3970    b_rows_are_n: bool,
3971    c: &mut [f32],
3972    ldc: usize,
3973) {
3974    #[cfg(any(target_os = "macos", target_arch = "aarch64"))]
3975    {
3976        sgemm_rm(m, n, k, alpha, a, lda, b_mat, ldb, b_rows_are_n, c, ldc);
3977    }
3978    // x86 without Accelerate has no sgemm_rm: the specialized paths there are
3979    // quantized kernels, not an f32 GEMM. Only the offline AWNP pass reaches
3980    // this, so correctness matters and throughput does not — a triple loop is
3981    // the honest fallback rather than a reason to make the tool macOS-only.
3982    #[cfg(not(any(target_os = "macos", target_arch = "aarch64")))]
3983    {
3984        for i in 0..m {
3985            for j in 0..n {
3986                let mut acc = 0f32;
3987                for p in 0..k {
3988                    let bv = if b_rows_are_n {
3989                        b_mat[j * ldb + p]
3990                    } else {
3991                        b_mat[p * ldb + j]
3992                    };
3993                    acc += a[i * lda + p] * bv;
3994                }
3995                c[i * ldc + j] = alpha * acc;
3996            }
3997        }
3998    }
3999}
4000
4001/// Row-major f32 GEMM on Accelerate: C[m,n] = alpha·A[m,k] × B(ᵀ).
4002/// `b_rows_are_n` = true multiplies by Bᵀ where B is stored [n, k].
4003#[cfg(target_os = "macos")]
4004#[allow(clippy::too_many_arguments)]
4005pub(crate) fn sgemm_rm(
4006    m: usize,
4007    n: usize,
4008    k: usize,
4009    alpha: f32,
4010    a: &[f32],
4011    lda: usize,
4012    b_mat: &[f32],
4013    ldb: usize,
4014    b_rows_are_n: bool,
4015    c: &mut [f32],
4016    ldc: usize,
4017) {
4018    debug_assert!(a.len() >= (m - 1) * lda + k);
4019    debug_assert!(c.len() >= (m - 1) * ldc + n);
4020    // Test hook: route the attention GEMMs through the portable NEON
4021    // micro-kernel ON APPLE SILICON — how the mobile batched attend is
4022    // measured without a phone in the loop. (Intel macOS has no NEON —
4023    // the hook is a no-op there, Accelerate continues below.)
4024    #[cfg(target_arch = "aarch64")]
4025    if std::env::var("CMF_FORCE_NEON_GEMM")
4026        .map(|v| v == "1")
4027        .unwrap_or(false)
4028    {
4029        return neon_gemm_rm(m, n, k, alpha, a, lda, b_mat, ldb, b_rows_are_n, c, ldc);
4030    }
4031    unsafe {
4032        accel_blas::cblas_sgemm(
4033            101, // RowMajor
4034            111, // NoTrans A
4035            if b_rows_are_n { 112 } else { 111 },
4036            m as i32,
4037            n as i32,
4038            k as i32,
4039            alpha,
4040            a.as_ptr(),
4041            lda as i32,
4042            b_mat.as_ptr(),
4043            ldb as i32,
4044            0.0,
4045            c.as_mut_ptr(),
4046            ldc as i32,
4047        );
4048    }
4049}
4050
4051/// Prefill GEMM through Accelerate (macOS): dequantize q8 rows into
4052/// f32 tiles (scale folded in, pool-parallel) and multiply each tile
4053/// on the AMX with one row-major sgemm. Tiles live in cache, weights
4054/// stream once. Numerics are f32-GEMM (not the int8 dot): prefill
4055/// logits shift within f32 rounding — tolerance-class, like every
4056/// reduction-order change; decode (M=1) never takes this path.
4057#[cfg(target_os = "macos")]
4058fn qmatmat_accel(
4059    q: &[u8],
4060    row_scale: &[f32],
4061    pre: &[std::borrow::Cow<'_, [f32]>],
4062    rows: usize,
4063    cols: usize,
4064    out: &mut [f32],
4065    pool: Option<&Pool>,
4066) {
4067    // NOTE: double-buffering the dequant against the sgemm (a scoped
4068    // thread driving the pool on tile k+1 while the caller multiplies
4069    // tile k) was tried and LOST ~6%: Accelerate's sgemm is itself
4070    // multithreaded, and the dequant workers just steal its cores.
4071    const TR: usize = 2048;
4072    let b = pre.len();
4073    thread_local! {
4074        static XPANEL: std::cell::RefCell<Vec<f32>> = const { std::cell::RefCell::new(Vec::new()) };
4075        static WTILE: std::cell::RefCell<Vec<f32>> = const { std::cell::RefCell::new(Vec::new()) };
4076    }
4077    XPANEL.with(|xp| {
4078        WTILE.with(|wt| {
4079            let mut xpanel = xp.borrow_mut();
4080            xpanel.clear();
4081            for x in pre {
4082                xpanel.extend_from_slice(x);
4083            }
4084            let mut wtile = wt.borrow_mut();
4085            wtile.resize(TR * cols, 0.0);
4086            let mut r0 = 0usize;
4087            while r0 < rows {
4088                let tr = TR.min(rows - r0);
4089                // Dequant the tile (scale folded) — pool-parallel.
4090                let wt_addr = SendMut(wtile.as_mut_ptr());
4091                let run = |start: usize, end: usize| {
4092                    for r in start..end {
4093                        let row = &q[(r0 + r) * cols..(r0 + r + 1) * cols];
4094                        let s = row_scale[r0 + r];
4095                        // SAFETY: workers cover disjoint r ranges.
4096                        let dst =
4097                            unsafe { std::slice::from_raw_parts_mut(wt_addr.at(r * cols), cols) };
4098                        for (d, &v) in dst.iter_mut().zip(row) {
4099                            *d = (v as i8) as f32 * s;
4100                        }
4101                    }
4102                };
4103                dispatch_rows(pool, tr, &run);
4104                // C[b, tr] (at column r0 of out[b, rows]) = X · Wtileᵀ
4105                unsafe {
4106                    accel_blas::cblas_sgemm(
4107                        101, // RowMajor
4108                        111, // NoTrans A
4109                        112, // Trans B
4110                        b as i32,
4111                        tr as i32,
4112                        cols as i32,
4113                        1.0,
4114                        xpanel.as_ptr(),
4115                        cols as i32,
4116                        wtile.as_ptr(),
4117                        cols as i32,
4118                        0.0,
4119                        out.as_mut_ptr().add(r0),
4120                        rows as i32,
4121                    );
4122                }
4123                r0 += tr;
4124            }
4125        })
4126    });
4127}
4128
4129fn qmatmat(
4130    q: &[u8],
4131    row_scale: &[f32],
4132    pre: &[std::borrow::Cow<'_, [f32]>],
4133    rows: usize,
4134    cols: usize,
4135    out: &mut [f32],
4136    pool: Option<&Pool>,
4137) {
4138    let b = pre.len();
4139    debug_assert_eq!(out.len(), b * rows);
4140    // Big prefill batches ride the AMX (roadmap PR3): the row×batch
4141    // SDOT loop below peaks near the CPU's dot throughput, an order
4142    // below the matrix units. Small tensors and tiny test models stay
4143    // on the exact integer path.
4144    #[cfg(target_os = "macos")]
4145    if b >= 8 && rows * cols >= 500_000 && accel_gemm_enabled() {
4146        qmatmat_accel(q, row_scale, pre, rows, cols, out, pool);
4147        return;
4148    }
4149    #[cfg(target_arch = "aarch64")]
4150    if sdot_enabled() {
4151        let acts: Vec<SplitAct> = pre.iter().map(|x| split_act(x)).collect();
4152        let out_addr = SendMut(out.as_mut_ptr());
4153        // Blocked 2×4 (mobile prefill: no AMX to fall back on — this
4154        // path IS the ARM prefill GEMM off Apple silicon).
4155        let blocked_ok = blocked_enabled();
4156        let use_i8mm = i8mm_enabled();
4157        if blocked_ok {
4158            let run = |start: usize, end: usize| {
4159                let mut o = start;
4160                while o < end {
4161                    if o + 2 <= end {
4162                        let r0 = &q[o * cols..(o + 1) * cols];
4163                        let r1 = &q[(o + 1) * cols..(o + 2) * cols];
4164                        let mut bi = 0usize;
4165                        while bi + 4 <= acts.len() {
4166                            let xs = [
4167                                acts[bi].xq.as_slice(),
4168                                acts[bi + 1].xq.as_slice(),
4169                                acts[bi + 2].xq.as_slice(),
4170                                acts[bi + 3].xq.as_slice(),
4171                            ];
4172                            let d = if use_i8mm {
4173                                unsafe { dot_i8_smmla_2x4(r0, r1, xs) }
4174                            } else {
4175                                unsafe { dot_i8_sdot_2x4(r0, r1, xs) }
4176                            };
4177                            for (r, row) in [r0, r1].into_iter().enumerate() {
4178                                for k in 0..4 {
4179                                    let act = &acts[bi + k];
4180                                    let mut v = d[r][k] as f32 * act.sx;
4181                                    for &(j, xv) in &act.outliers {
4182                                        v += (row[j] as i8) as f32 * xv;
4183                                    }
4184                                    unsafe {
4185                                        *out_addr.at((bi + k) * rows + o + r) = v * row_scale[o + r]
4186                                    };
4187                                }
4188                            }
4189                            bi += 4;
4190                        }
4191                        while bi < acts.len() {
4192                            for (r, row) in [r0, r1].into_iter().enumerate() {
4193                                let v = row_dot_sdot(row, &acts[bi]) * row_scale[o + r];
4194                                unsafe { *out_addr.at(bi * rows + o + r) = v };
4195                            }
4196                            bi += 1;
4197                        }
4198                        o += 2;
4199                    } else {
4200                        let row = &q[o * cols..(o + 1) * cols];
4201                        for (bi, act) in acts.iter().enumerate() {
4202                            let v = row_dot_sdot(row, act) * row_scale[o];
4203                            unsafe { *out_addr.at(bi * rows + o) = v };
4204                        }
4205                        o += 1;
4206                    }
4207                }
4208            };
4209            dispatch_rows(pool, rows, &run);
4210            return;
4211        }
4212        let run = |start: usize, end: usize| {
4213            for o in start..end {
4214                let row = &q[o * cols..(o + 1) * cols];
4215                for (bi, act) in acts.iter().enumerate() {
4216                    let v = row_dot_sdot(row, act) * row_scale[o];
4217                    unsafe { *out_addr.at(bi * rows + o) = v };
4218                }
4219            }
4220        };
4221        dispatch_rows(pool, rows, &run);
4222        return;
4223    }
4224    // x86 A8W8 batch. Non-VNNI parts take the BLOCKED 2×4 kernel
4225    // (roadmap P0: two weight rows' abs() stay in registers across four
4226    // activation streams); VNNI machines keep the per-row bias-trick
4227    // dot, which is already throughput-bound there.
4228    #[cfg(target_arch = "x86_64")]
4229    if avx2_a8w8_enabled() {
4230        let acts: Vec<SplitAct> = pre.iter().map(|x| split_act(x)).collect();
4231        let out_addr = SendMut(out.as_mut_ptr());
4232        // CMF_X86_BLOCKED=0 forces the per-row path (paired in-process
4233        // A/B on noisy shared-vCPU hosts).
4234        let blocked_ok = blocked_enabled();
4235        if !avx512vnni_enabled() && blocked_ok && !row_exact() {
4236            let run = |start: usize, end: usize| {
4237                let mut o = start;
4238                while o < end {
4239                    if o + 2 <= end {
4240                        let r0 = &q[o * cols..(o + 1) * cols];
4241                        let r1 = &q[(o + 1) * cols..(o + 2) * cols];
4242                        let mut bi = 0usize;
4243                        while bi + 4 <= acts.len() {
4244                            let xs = [
4245                                acts[bi].xq.as_slice(),
4246                                acts[bi + 1].xq.as_slice(),
4247                                acts[bi + 2].xq.as_slice(),
4248                                acts[bi + 3].xq.as_slice(),
4249                            ];
4250                            let d = unsafe { dot_i8_i8_avx2_2x4(r0, r1, xs) };
4251                            for (r, row) in [r0, r1].into_iter().enumerate() {
4252                                for k in 0..4 {
4253                                    let act = &acts[bi + k];
4254                                    let mut v = d[r][k] as f32 * act.sx;
4255                                    for &(j, xv) in &act.outliers {
4256                                        v += (row[j] as i8) as f32 * xv;
4257                                    }
4258                                    unsafe {
4259                                        *out_addr.at((bi + k) * rows + o + r) = v * row_scale[o + r]
4260                                    };
4261                                }
4262                            }
4263                            bi += 4;
4264                        }
4265                        while bi < acts.len() {
4266                            for (r, row) in [r0, r1].into_iter().enumerate() {
4267                                let v = row_dot_avx2(row, &acts[bi]) * row_scale[o + r];
4268                                unsafe { *out_addr.at(bi * rows + o + r) = v };
4269                            }
4270                            bi += 1;
4271                        }
4272                        o += 2;
4273                    } else {
4274                        let row = &q[o * cols..(o + 1) * cols];
4275                        for (bi, act) in acts.iter().enumerate() {
4276                            let v = row_dot_avx2(row, act) * row_scale[o];
4277                            unsafe { *out_addr.at(bi * rows + o) = v };
4278                        }
4279                        o += 1;
4280                    }
4281                }
4282            };
4283            dispatch_rows(pool, rows, &run);
4284            return;
4285        }
4286        let run = |start: usize, end: usize| {
4287            for o in start..end {
4288                let row = &q[o * cols..(o + 1) * cols];
4289                for (bi, act) in acts.iter().enumerate() {
4290                    let v = row_dot_avx2(row, act) * row_scale[o];
4291                    unsafe { *out_addr.at(bi * rows + o) = v };
4292                }
4293            }
4294        };
4295        dispatch_rows(pool, rows, &run);
4296        return;
4297    }
4298    let out_addr = SendMut(out.as_mut_ptr());
4299    let run = |start: usize, end: usize| {
4300        for o in start..end {
4301            let row = &q[o * cols..(o + 1) * cols];
4302            for (bi, x) in pre.iter().enumerate() {
4303                let mut acc = 0f32;
4304                for j in 0..cols {
4305                    acc += (row[j] as i8) as f32 * x[j];
4306                }
4307                unsafe { *out_addr.at(bi * rows + o) = acc * row_scale[o] };
4308            }
4309        }
4310    };
4311    dispatch_rows(pool, rows, &run);
4312}
4313
4314/// Split rows across pool workers (shared qmatvec pattern). Self-balancing
4315/// — see `Pool::run_rows` for why a static 1/n split is wrong here.
4316fn dispatch_rows(pool: Option<&Pool>, rows: usize, run: &(dyn Fn(usize, usize) + Sync)) {
4317    match pool {
4318        Some(pool) if rows >= 256 => pool.run_rows(rows, run),
4319        _ => run(0, rows),
4320    }
4321}
4322
4323/// Split a q4_block blob into (packed nibbles, f16 group scales).
4324fn q4_split(bytes: &[u8], rows: usize, cols: usize) -> (&[u8], &[u8]) {
4325    let groups = rows * cols / GROUP_SIZE;
4326    bytes.split_at(groups * 16)
4327}
4328
4329/// SIMD unpack for the dominant vbit width B=4 (94% of rows on the
4330/// log2-shape calibration): 16 packed bytes -> 32 centered i8 values.
4331/// vbit packs MSB-first, so the HIGH nibble is the even element
4332/// (opposite of q4_block's lo-first interleave). Centering is u-7.
4333#[inline]
4334fn vbit_fill4(data: &[u8], buf: &mut [u8]) {
4335    #[cfg(target_arch = "aarch64")]
4336    unsafe {
4337        return vbit_fill4_neon(data, buf);
4338    }
4339    #[cfg(target_arch = "x86_64")]
4340    if avx2_enabled() {
4341        return unsafe { vbit_fill4_avx2(data, buf) };
4342    }
4343    #[allow(unreachable_code)]
4344    for (blk, chunk) in buf.chunks_exact_mut(8).enumerate() {
4345        let u = unpack8::<4>(&data[blk * 4..]);
4346        for k in 0..8 {
4347            chunk[k] = (u[k] - 7) as i8 as u8;
4348        }
4349    }
4350}
4351
4352#[cfg(target_arch = "aarch64")]
4353#[target_feature(enable = "neon")]
4354unsafe fn vbit_fill4_neon(data: &[u8], buf: &mut [u8]) {
4355    // SAFETY: buf.len() is a multiple of GROUP_SIZE=32; data holds
4356    // buf.len()/2 packed bytes (validated at load).
4357    unsafe {
4358        use core::arch::aarch64::*;
4359        let n = buf.len();
4360        let mask = vdupq_n_u8(0x0F);
4361        let seven = vdupq_n_s8(7);
4362        let mut g = 0usize;
4363        while g * 32 + 32 <= n {
4364            let b = vld1q_u8(data.as_ptr().add(g * 16));
4365            let hi = vshrq_n_u8::<4>(b);
4366            let lo = vandq_u8(b, mask);
4367            let z0 = vsubq_s8(vreinterpretq_s8_u8(vzip1q_u8(hi, lo)), seven);
4368            let z1 = vsubq_s8(vreinterpretq_s8_u8(vzip2q_u8(hi, lo)), seven);
4369            vst1q_u8(buf.as_mut_ptr().add(g * 32), vreinterpretq_u8_s8(z0));
4370            vst1q_u8(buf.as_mut_ptr().add(g * 32 + 16), vreinterpretq_u8_s8(z1));
4371            g += 1;
4372        }
4373    }
4374}
4375
4376#[cfg(target_arch = "x86_64")]
4377#[target_feature(enable = "avx2")]
4378unsafe fn vbit_fill4_avx2(data: &[u8], buf: &mut [u8]) {
4379    // SAFETY: see vbit_fill4_neon.
4380    unsafe {
4381        use core::arch::x86_64::*;
4382        let n = buf.len();
4383        let mask = _mm_set1_epi8(0x0F);
4384        let seven = _mm256_set1_epi8(7);
4385        let mut g = 0usize;
4386        while g * 32 + 32 <= n {
4387            let b = _mm_loadu_si128(data.as_ptr().add(g * 16) as *const __m128i);
4388            let hi = _mm_and_si128(_mm_srli_epi16::<4>(b), mask);
4389            let lo = _mm_and_si128(b, mask);
4390            let z = _mm256_sub_epi8(
4391                _mm256_set_m128i(_mm_unpackhi_epi8(hi, lo), _mm_unpacklo_epi8(hi, lo)),
4392                seven,
4393            );
4394            _mm256_storeu_si256(buf.as_mut_ptr().add(g * 32) as *mut __m256i, z);
4395            g += 1;
4396        }
4397    }
4398}
4399
4400/// Unpack 8 MSB-first B-bit values from exactly B bytes (fixed shifts —
4401/// no serial bit-buffer, auto-vectorizable). Every 32-value group starts
4402/// byte-aligned (32·B/8 is integral for B∈3..8), so groups decompose
4403/// into 4 such blocks.
4404#[inline(always)]
4405fn unpack8<const B: usize>(data: &[u8]) -> [i32; 8] {
4406    let mut acc = 0u64;
4407    for i in 0..B {
4408        acc = (acc << 8) | data[i] as u64;
4409    }
4410    let mask = (1u64 << B) - 1;
4411    let mut out = [0i32; 8];
4412    for (k, o) in out.iter_mut().enumerate() {
4413        *o = ((acc >> ((7 - k) * B)) & mask) as i32;
4414    }
4415    out
4416}
4417
4418/// Fused vbit matvec straight from the mapped bytes (spec §3, P13
4419/// FIG.3): [u8 bits: rows][f16 scales: rows·cols/32][bit-packed rows,
4420/// MSB-first, byte-padded]. Row data offsets are precomputed at load
4421/// (`vbit_row_offsets`) — the per-call prefix scan was O(rows) pure
4422/// overhead on every matvec.
4423#[allow(clippy::too_many_arguments)]
4424fn vbitmatvec(
4425    bytes: &[u8],
4426    offsets: &[usize],
4427    x: &[f32],
4428    rows: usize,
4429    cols: usize,
4430    out: &mut [f32],
4431    pool: Option<&Pool>,
4432) {
4433    debug_assert_eq!(out.len(), rows);
4434    debug_assert_eq!(offsets.len(), rows + 1);
4435
4436    // SDOT path: unpack the row to centered i8 once, then per-group
4437    // int8 dot against the quantized activations — same A8W8 contract
4438    // as q8 (bounded noise; CMF_SDOT=0 keeps the exact scalar path).
4439    if a8w8_enabled() {
4440        let act = split_act(x);
4441        let out_addr = SendMut(out.as_mut_ptr());
4442        let run = move |start: usize, end: usize| {
4443            vbit_range_a8w8(bytes, offsets, x, &act, rows, cols, out_addr, start, end)
4444        };
4445        dispatch_rows(pool, rows, &run);
4446        return;
4447    }
4448
4449    let out_addr = SendMut(out.as_mut_ptr());
4450    let run = move |start: usize, end: usize| {
4451        vbit_range_f32(bytes, offsets, x, rows, cols, out_addr, start, end)
4452    };
4453    dispatch_rows(pool, rows, &run);
4454}
4455
4456/// One vbit row range via the A8W8 int8 path — kernel body of
4457/// `vbitmatvec`, extracted so multi-matrix jobs can drive it for
4458/// several tensors in one dispatch (b=8 rows go exact f32).
4459#[allow(clippy::too_many_arguments)]
4460fn vbit_range_a8w8(
4461    bytes: &[u8],
4462    offsets: &[usize],
4463    x: &[f32],
4464    act: &SplitAct,
4465    rows: usize,
4466    cols: usize,
4467    out: SendMut,
4468    start: usize,
4469    end: usize,
4470) {
4471    let ng = cols / GROUP_SIZE;
4472    let bits = &bytes[..rows];
4473    let sc_off = rows;
4474    let row_dot = |r: usize| -> f32 {
4475        let b = bits[r] as usize;
4476        let l = (1i32 << (b - 1)) - 1;
4477        let mask = (1u64 << b) - 1;
4478        let data = &bytes[offsets[r]..offsets[r + 1]];
4479        if b == 8 {
4480            // u−L reaches 128 → does not fit i8; exact f32 path.
4481            let (mut acc, mut nbits, mut idx) = (0u64, 0usize, 0usize);
4482            let mut dot = 0f32;
4483            for g in 0..ng {
4484                let so = (r * ng + g) * 2;
4485                let sgf = f16_to_f32(u16::from_le_bytes([
4486                    bytes[sc_off + so],
4487                    bytes[sc_off + so + 1],
4488                ]));
4489                let xg = &x[g * GROUP_SIZE..(g + 1) * GROUP_SIZE];
4490                let mut gd = 0f32;
4491                for &xv in xg.iter() {
4492                    if nbits < 8 {
4493                        acc = (acc << 8) | data[idx] as u64;
4494                        idx += 1;
4495                        nbits += 8;
4496                    }
4497                    let u = ((acc >> (nbits - 8)) & 0xFF) as i32;
4498                    nbits -= 8;
4499                    gd += (u - l) as f32 * xv;
4500                }
4501                dot += gd * sgf;
4502            }
4503            return dot;
4504        }
4505        // Per-worker scratch: this closure runs for every row of the
4506        // tensor (lm_head ≈ 150k rows/token) — a heap allocation per
4507        // row was measurable pure overhead.
4508        thread_local! {
4509            static VBIT_SCRATCH: std::cell::RefCell<Vec<u8>> =
4510                const { std::cell::RefCell::new(Vec::new()) };
4511        }
4512        #[inline(always)]
4513        fn fill<const B: usize>(data: &[u8], l: i32, buf: &mut [u8]) {
4514            for (blk, chunk) in buf.chunks_exact_mut(8).enumerate() {
4515                let u = unpack8::<B>(&data[blk * B..]);
4516                for k in 0..8 {
4517                    chunk[k] = (u[k] - l) as i8 as u8;
4518                }
4519            }
4520        }
4521        let _ = mask;
4522        VBIT_SCRATCH.with(|scratch| {
4523            let mut buf = scratch.borrow_mut();
4524            buf.resize(cols, 0);
4525            match b {
4526                3 => fill::<3>(data, l, &mut buf),
4527                4 => vbit_fill4(data, &mut buf),
4528                5 => fill::<5>(data, l, &mut buf),
4529                6 => fill::<6>(data, l, &mut buf),
4530                _ => unreachable!(),
4531            }
4532            let mut dot = 0f32;
4533            for g in 0..ng {
4534                let so = (r * ng + g) * 2;
4535                let s = f16_to_f32(u16::from_le_bytes([
4536                    bytes[sc_off + so],
4537                    bytes[sc_off + so + 1],
4538                ]));
4539                let d = dot_i8_i8(
4540                    &buf[g * GROUP_SIZE..(g + 1) * GROUP_SIZE],
4541                    &act.xq[g * GROUP_SIZE..(g + 1) * GROUP_SIZE],
4542                ) as f32
4543                    * act.sx;
4544                dot += d * s;
4545            }
4546            for &(j, xv) in &act.outliers {
4547                let so = (r * ng + j / GROUP_SIZE) * 2;
4548                let s = f16_to_f32(u16::from_le_bytes([
4549                    bytes[sc_off + so],
4550                    bytes[sc_off + so + 1],
4551                ]));
4552                // xq is zeroed at outlier slots — add the exact term.
4553                dot += (buf[j] as i8) as f32 * s * xv;
4554            }
4555            dot
4556        })
4557    };
4558    for r in start..end {
4559        // SAFETY: disjoint row ranges per worker.
4560        unsafe { *out.at(r) = row_dot(r) };
4561    }
4562}
4563
4564/// Exact scalar vbit row range (same extraction, non-SDOT path).
4565#[allow(clippy::too_many_arguments)]
4566fn vbit_range_f32(
4567    bytes: &[u8],
4568    offsets: &[usize],
4569    x: &[f32],
4570    rows: usize,
4571    cols: usize,
4572    out: SendMut,
4573    start: usize,
4574    end: usize,
4575) {
4576    let ng = cols / GROUP_SIZE;
4577    let bits = &bytes[..rows];
4578    let sc_off = rows;
4579    // Per-bit-width specialized inner loops: the compiler unrolls the
4580    // constant shifts (the generic bit-buffer loop was branch-bound —
4581    // 5.6 vs 13.2 tok/s q4 on the 0.8B).
4582    #[inline(always)]
4583    fn dot_row<const B: usize>(
4584        data: &[u8],
4585        bytes: &[u8],
4586        sc_off: usize,
4587        r: usize,
4588        ng: usize,
4589        x: &[f32],
4590    ) -> f32 {
4591        let l = ((1i32 << (B - 1)) - 1) as f32;
4592        let gbytes = GROUP_SIZE * B / 8;
4593        let mut dot = 0f32;
4594        for g in 0..ng {
4595            let so = (r * ng + g) * 2;
4596            let s = f16_to_f32(u16::from_le_bytes([
4597                bytes[sc_off + so],
4598                bytes[sc_off + so + 1],
4599            ]));
4600            let xg = &x[g * GROUP_SIZE..(g + 1) * GROUP_SIZE];
4601            let gd0 = &data[g * gbytes..(g + 1) * gbytes];
4602            let mut gd = 0f32;
4603            for blk in 0..GROUP_SIZE / 8 {
4604                let u = unpack8::<B>(&gd0[blk * B..]);
4605                let xb = &xg[blk * 8..blk * 8 + 8];
4606                for k in 0..8 {
4607                    gd += (u[k] as f32 - l) * xb[k];
4608                }
4609            }
4610            dot += gd * s;
4611        }
4612        dot
4613    }
4614    for r in start..end {
4615        let data = &bytes[offsets[r]..offsets[r + 1]];
4616        let v = match bits[r] {
4617            3 => dot_row::<3>(data, bytes, sc_off, r, ng, x),
4618            4 => dot_row::<4>(data, bytes, sc_off, r, ng, x),
4619            5 => dot_row::<5>(data, bytes, sc_off, r, ng, x),
4620            6 => dot_row::<6>(data, bytes, sc_off, r, ng, x),
4621            8 => dot_row::<8>(data, bytes, sc_off, r, ng, x),
4622            b => unreachable!("vbit bit-width {b} (validated at load)"),
4623        };
4624        // SAFETY: disjoint row ranges per worker.
4625        unsafe { *out.at(r) = v };
4626    }
4627}
4628
4629/// Fused two-input vbit matvec: each row is unpacked from the mmap ONCE
4630/// and dotted against BOTH activations (MTP verify / pair prefill used
4631/// to run two full matvecs — double weight traffic and double unpack).
4632/// Per-input math is identical to `vbitmatvec` → same accuracy contract.
4633#[allow(clippy::too_many_arguments)]
4634fn vbitmatvec2(
4635    bytes: &[u8],
4636    offsets: &[usize],
4637    x1: &[f32],
4638    x2: &[f32],
4639    rows: usize,
4640    cols: usize,
4641    o1: &mut [f32],
4642    o2: &mut [f32],
4643    pool: Option<&Pool>,
4644) {
4645    debug_assert_eq!(o1.len(), rows);
4646    debug_assert_eq!(o2.len(), rows);
4647
4648    if a8w8_enabled() {
4649        let a1 = split_act(x1);
4650        let a2 = split_act(x2);
4651        let p1 = SendMut(o1.as_mut_ptr());
4652        let p2 = SendMut(o2.as_mut_ptr());
4653        let run = move |start: usize, end: usize| {
4654            vbit_range2_a8w8(
4655                bytes, offsets, x1, x2, &a1, &a2, rows, cols, p1, p2, start, end,
4656            )
4657        };
4658        dispatch_rows(pool, rows, &run);
4659        return;
4660    }
4661
4662    let p1 = SendMut(o1.as_mut_ptr());
4663    let p2 = SendMut(o2.as_mut_ptr());
4664    let run = move |start: usize, end: usize| {
4665        vbit_range2_f32(bytes, offsets, x1, x2, rows, cols, p1, p2, start, end)
4666    };
4667    dispatch_rows(pool, rows, &run);
4668}
4669
4670/// Two-input vbit row range via the A8W8 int8 path — kernel body of
4671/// `vbitmatvec2`, extracted for pair multi-matrix jobs (b=8 rows go
4672/// exact f32 for both lanes, bits streamed once).
4673#[allow(clippy::too_many_arguments)]
4674fn vbit_range2_a8w8(
4675    bytes: &[u8],
4676    offsets: &[usize],
4677    x1: &[f32],
4678    x2: &[f32],
4679    a1: &SplitAct,
4680    a2: &SplitAct,
4681    rows: usize,
4682    cols: usize,
4683    p1: SendMut,
4684    p2: SendMut,
4685    start: usize,
4686    end: usize,
4687) {
4688    let ng = cols / GROUP_SIZE;
4689    let bits = &bytes[..rows];
4690    let sc_off = rows;
4691    let row_dots = |r: usize| -> (f32, f32) {
4692        let b = bits[r] as usize;
4693        let l = (1i32 << (b - 1)) - 1;
4694        let data = &bytes[offsets[r]..offsets[r + 1]];
4695        if b == 8 {
4696            // u−L reaches 128 → does not fit i8; exact f32 path,
4697            // bits still streamed once for both lanes.
4698            let (mut acc, mut nbits, mut idx) = (0u64, 0usize, 0usize);
4699            let (mut d1, mut d2) = (0f32, 0f32);
4700            for g in 0..ng {
4701                let so = (r * ng + g) * 2;
4702                let sgf = f16_to_f32(u16::from_le_bytes([
4703                    bytes[sc_off + so],
4704                    bytes[sc_off + so + 1],
4705                ]));
4706                let (mut g1, mut g2) = (0f32, 0f32);
4707                for k in 0..GROUP_SIZE {
4708                    if nbits < 8 {
4709                        acc = (acc << 8) | data[idx] as u64;
4710                        idx += 1;
4711                        nbits += 8;
4712                    }
4713                    let u = ((acc >> (nbits - 8)) & 0xFF) as i32;
4714                    nbits -= 8;
4715                    let w = (u - l) as f32;
4716                    g1 += w * x1[g * GROUP_SIZE + k];
4717                    g2 += w * x2[g * GROUP_SIZE + k];
4718                }
4719                d1 += g1 * sgf;
4720                d2 += g2 * sgf;
4721            }
4722            return (d1, d2);
4723        }
4724        thread_local! {
4725            static VBIT_SCRATCH2: std::cell::RefCell<Vec<u8>> =
4726                const { std::cell::RefCell::new(Vec::new()) };
4727        }
4728        #[inline(always)]
4729        fn fill<const B: usize>(data: &[u8], l: i32, buf: &mut [u8]) {
4730            for (blk, chunk) in buf.chunks_exact_mut(8).enumerate() {
4731                let u = unpack8::<B>(&data[blk * B..]);
4732                for k in 0..8 {
4733                    chunk[k] = (u[k] - l) as i8 as u8;
4734                }
4735            }
4736        }
4737        VBIT_SCRATCH2.with(|scratch| {
4738            let mut buf = scratch.borrow_mut();
4739            buf.resize(cols, 0);
4740            match b {
4741                3 => fill::<3>(data, l, &mut buf),
4742                4 => vbit_fill4(data, &mut buf),
4743                5 => fill::<5>(data, l, &mut buf),
4744                6 => fill::<6>(data, l, &mut buf),
4745                _ => unreachable!(),
4746            }
4747            let (mut d1, mut d2) = (0f32, 0f32);
4748            for g in 0..ng {
4749                let so = (r * ng + g) * 2;
4750                let s = f16_to_f32(u16::from_le_bytes([
4751                    bytes[sc_off + so],
4752                    bytes[sc_off + so + 1],
4753                ]));
4754                let wg = &buf[g * GROUP_SIZE..(g + 1) * GROUP_SIZE];
4755                let v1 = dot_i8_i8(wg, &a1.xq[g * GROUP_SIZE..(g + 1) * GROUP_SIZE]) as f32 * a1.sx;
4756                let v2 = dot_i8_i8(wg, &a2.xq[g * GROUP_SIZE..(g + 1) * GROUP_SIZE]) as f32 * a2.sx;
4757                d1 += v1 * s;
4758                d2 += v2 * s;
4759            }
4760            for &(j, xv) in &a1.outliers {
4761                let so = (r * ng + j / GROUP_SIZE) * 2;
4762                let s = f16_to_f32(u16::from_le_bytes([
4763                    bytes[sc_off + so],
4764                    bytes[sc_off + so + 1],
4765                ]));
4766                d1 += (buf[j] as i8) as f32 * s * xv;
4767            }
4768            for &(j, xv) in &a2.outliers {
4769                let so = (r * ng + j / GROUP_SIZE) * 2;
4770                let s = f16_to_f32(u16::from_le_bytes([
4771                    bytes[sc_off + so],
4772                    bytes[sc_off + so + 1],
4773                ]));
4774                d2 += (buf[j] as i8) as f32 * s * xv;
4775            }
4776            (d1, d2)
4777        })
4778    };
4779    for r in start..end {
4780        let (v1, v2) = row_dots(r);
4781        // SAFETY: disjoint row ranges per worker.
4782        unsafe {
4783            *p1.at(r) = v1;
4784            *p2.at(r) = v2;
4785        }
4786    }
4787}
4788
4789/// Two-input exact scalar vbit row range (same extraction) —
4790/// per-bit-width specialized, two accumulators per row; per-lane
4791/// accumulation order matches `vbitmatvec` exactly.
4792#[allow(clippy::too_many_arguments)]
4793fn vbit_range2_f32(
4794    bytes: &[u8],
4795    offsets: &[usize],
4796    x1: &[f32],
4797    x2: &[f32],
4798    rows: usize,
4799    cols: usize,
4800    p1: SendMut,
4801    p2: SendMut,
4802    start: usize,
4803    end: usize,
4804) {
4805    let ng = cols / GROUP_SIZE;
4806    let bits = &bytes[..rows];
4807    let sc_off = rows;
4808    #[inline(always)]
4809    #[allow(clippy::too_many_arguments)]
4810    fn dot_row2<const B: usize>(
4811        data: &[u8],
4812        bytes: &[u8],
4813        sc_off: usize,
4814        r: usize,
4815        ng: usize,
4816        x1: &[f32],
4817        x2: &[f32],
4818    ) -> (f32, f32) {
4819        let l = ((1i32 << (B - 1)) - 1) as f32;
4820        let gbytes = GROUP_SIZE * B / 8;
4821        let (mut d1, mut d2) = (0f32, 0f32);
4822        for g in 0..ng {
4823            let so = (r * ng + g) * 2;
4824            let s = f16_to_f32(u16::from_le_bytes([
4825                bytes[sc_off + so],
4826                bytes[sc_off + so + 1],
4827            ]));
4828            let x1g = &x1[g * GROUP_SIZE..(g + 1) * GROUP_SIZE];
4829            let x2g = &x2[g * GROUP_SIZE..(g + 1) * GROUP_SIZE];
4830            let gd0 = &data[g * gbytes..(g + 1) * gbytes];
4831            let (mut g1, mut g2) = (0f32, 0f32);
4832            for blk in 0..GROUP_SIZE / 8 {
4833                let u = unpack8::<B>(&gd0[blk * B..]);
4834                for k in 0..8 {
4835                    let w = u[k] as f32 - l;
4836                    g1 += w * x1g[blk * 8 + k];
4837                    g2 += w * x2g[blk * 8 + k];
4838                }
4839            }
4840            d1 += g1 * s;
4841            d2 += g2 * s;
4842        }
4843        (d1, d2)
4844    }
4845    for r in start..end {
4846        let data = &bytes[offsets[r]..offsets[r + 1]];
4847        let (v1, v2) = match bits[r] {
4848            3 => dot_row2::<3>(data, bytes, sc_off, r, ng, x1, x2),
4849            4 => dot_row2::<4>(data, bytes, sc_off, r, ng, x1, x2),
4850            5 => dot_row2::<5>(data, bytes, sc_off, r, ng, x1, x2),
4851            6 => dot_row2::<6>(data, bytes, sc_off, r, ng, x1, x2),
4852            8 => dot_row2::<8>(data, bytes, sc_off, r, ng, x1, x2),
4853            b => unreachable!("vbit bit-width {b} (validated at load)"),
4854        };
4855        // SAFETY: disjoint row ranges per worker.
4856        unsafe {
4857            *p1.at(r) = v1;
4858            *p2.at(r) = v2;
4859        }
4860    }
4861}
4862
4863// ───────────────────── q4_tiled kernels (§4.3) ─────────────────────
4864
4865/// One q4_tiled row dot on the A8W8 int8 path: per 32-group the tile
4866/// is ONE sequential read — [f16 scale][16B nibbles] — versus the two
4867/// distant streams of the split layout. Values/order identical to the
4868/// split kernels.
4869#[inline]
4870#[allow(unreachable_code)]
4871fn dot_q4t_row_i8(bytes: &[u8], r: usize, gpr: usize, xq: &[i8]) -> f32 {
4872    #[cfg(target_arch = "aarch64")]
4873    unsafe {
4874        return dot_q4t_row_sdot(bytes, r, gpr, xq);
4875    }
4876    #[cfg(target_arch = "x86_64")]
4877    unsafe {
4878        if vnni_tiles_enabled() {
4879            return dot_q4t_row_vnni(bytes, r, gpr, xq);
4880        }
4881        return dot_q4t_row_avx2(bytes, r, gpr, xq);
4882    }
4883    let mut acc = 0f32;
4884    for gi in 0..gpr {
4885        let tile = &bytes[(r * gpr + gi) * Q4_TILE..(r * gpr + gi + 1) * Q4_TILE];
4886        let s = f16_to_f32(u16::from_le_bytes([tile[0], tile[1]]));
4887        let mut d = 0i32;
4888        for (k, &b) in tile[2..].iter().enumerate() {
4889            d += ((b & 0x0F) as i32 - 8) * xq[gi * GROUP_SIZE + k * 2] as i32
4890                + (((b >> 4) & 0x0F) as i32 - 8) * xq[gi * GROUP_SIZE + k * 2 + 1] as i32;
4891        }
4892        acc += d as f32 * s;
4893    }
4894    acc
4895}
4896
4897#[cfg(target_arch = "aarch64")]
4898#[target_feature(enable = "neon,dotprod")]
4899unsafe fn dot_q4t_row_sdot(bytes: &[u8], r: usize, gpr: usize, xq: &[i8]) -> f32 {
4900    // SAFETY: callers uphold slice-length contracts (18B tile per group,
4901    // xq.len() == gpr·GROUP_SIZE).
4902    unsafe {
4903        use core::arch::aarch64::*;
4904        use core::arch::asm;
4905        let lomask = vdupq_n_u8(0x0F);
4906        let eight = vdupq_n_s8(8);
4907        let mut acc = 0f32;
4908        for gi in 0..gpr {
4909            let t = bytes.as_ptr().add((r * gpr + gi) * Q4_TILE);
4910            let s = f16_to_f32(u16::from_le_bytes([*t, *t.add(1)]));
4911            let b = vld1q_u8(t.add(2));
4912            let lo = vandq_u8(b, lomask);
4913            let hi = vshrq_n_u8::<4>(b);
4914            let e0 = vsubq_s8(vreinterpretq_s8_u8(vzip1q_u8(lo, hi)), eight);
4915            let e1 = vsubq_s8(vreinterpretq_s8_u8(vzip2q_u8(lo, hi)), eight);
4916            let x0 = vld1q_s8(xq.as_ptr().add(gi * GROUP_SIZE));
4917            let x1 = vld1q_s8(xq.as_ptr().add(gi * GROUP_SIZE + 16));
4918            let (mut a0, mut a1) = (vdupq_n_s32(0), vdupq_n_s32(0));
4919            asm!(
4920                "sdot {a0:v}.4s, {e0:v}.16b, {x0:v}.16b",
4921                "sdot {a1:v}.4s, {e1:v}.16b, {x1:v}.16b",
4922                a0 = inout(vreg) a0, a1 = inout(vreg) a1,
4923                e0 = in(vreg) e0, x0 = in(vreg) x0, e1 = in(vreg) e1, x1 = in(vreg) x1,
4924                options(pure, nomem, nostack),
4925            );
4926            acc += vaddvq_s32(vaddq_s32(a0, a1)) as f32 * s;
4927        }
4928        acc
4929    }
4930}
4931
4932#[cfg(target_arch = "x86_64")]
4933#[target_feature(enable = "avx2")]
4934unsafe fn dot_q4t_row_avx2(bytes: &[u8], r: usize, gpr: usize, xq: &[i8]) -> f32 {
4935    // SAFETY: see dot_q4t_row_sdot.
4936    unsafe {
4937        use core::arch::x86_64::*;
4938        let lomask = _mm_set1_epi8(0x0F);
4939        let eight = _mm256_set1_epi8(8);
4940        let ones = _mm256_set1_epi16(1);
4941        let mut acc = 0f32;
4942        for gi in 0..gpr {
4943            let t = bytes.as_ptr().add((r * gpr + gi) * Q4_TILE);
4944            let s = f16_to_f32(u16::from_le_bytes([*t, *t.add(1)]));
4945            let b = _mm_loadu_si128(t.add(2) as *const __m128i);
4946            let lo = _mm_and_si128(b, lomask);
4947            let hi = _mm_and_si128(_mm_srli_epi16::<4>(b), lomask);
4948            let w = _mm256_sub_epi8(
4949                _mm256_set_m128i(_mm_unpackhi_epi8(lo, hi), _mm_unpacklo_epi8(lo, hi)),
4950                eight,
4951            );
4952            let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
4953            let p16 = _mm256_maddubs_epi16(_mm256_abs_epi8(w), _mm256_sign_epi8(x, w));
4954            let d = _mm256_madd_epi16(p16, ones);
4955            let hi128 = _mm256_extracti128_si256::<1>(d);
4956            let s128 = _mm_add_epi32(_mm256_castsi256_si128(d), hi128);
4957            let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
4958            let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
4959            acc += _mm_cvtsi128_si32(s32) as f32 * s;
4960        }
4961        acc
4962    }
4963}
4964
4965/// VNNI twin of `dot_q4t_row_avx2`: same unpack, `vpdpbusd` replaces
4966/// the maddubs+madd pair (see `dpbusd_hsum` — sums are bit-identical).
4967/// 256-bit VL encoding, so the VEX `vpsignb` stays usable.
4968#[cfg(target_arch = "x86_64")]
4969#[target_feature(enable = "avx2,avx512f,avx512bw,avx512vl,avx512vnni")]
4970unsafe fn dot_q4t_row_vnni(bytes: &[u8], r: usize, gpr: usize, xq: &[i8]) -> f32 {
4971    // SAFETY: see dot_q4t_row_sdot.
4972    unsafe {
4973        use core::arch::x86_64::*;
4974        let lomask = _mm_set1_epi8(0x0F);
4975        let eight = _mm256_set1_epi8(8);
4976        let mut acc = 0f32;
4977        for gi in 0..gpr {
4978            let t = bytes.as_ptr().add((r * gpr + gi) * Q4_TILE);
4979            let s = f16_to_f32(u16::from_le_bytes([*t, *t.add(1)]));
4980            let b = _mm_loadu_si128(t.add(2) as *const __m128i);
4981            let lo = _mm_and_si128(b, lomask);
4982            let hi = _mm_and_si128(_mm_srli_epi16::<4>(b), lomask);
4983            let w = _mm256_sub_epi8(
4984                _mm256_set_m128i(_mm_unpackhi_epi8(lo, hi), _mm_unpacklo_epi8(lo, hi)),
4985                eight,
4986            );
4987            let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
4988            let d = dpbusd_hsum(_mm256_abs_epi8(w), _mm256_sign_epi8(x, w));
4989            acc += d as f32 * s;
4990        }
4991        acc
4992    }
4993}
4994
4995/// One q4_tiled row against FOUR activation streams: the nibble unpack
4996/// and abs() happen once per group instead of once per (group,
4997/// activation) — the unpack is the dominant per-element cost of the
4998/// tiled format (roadmap P0 portable blocking, q4t leg).
4999#[cfg(target_arch = "x86_64")]
5000// `fma` is NOT implied by `avx2`: without it LLVM lowers _mm256_fmadd_ps
5001// to a libm call per lane — measured 2x slower than the reduction this
5002// kernel replaces. The runtime gate (`avx2_enabled`) already requires
5003// both features, so declaring it here is safe.
5004#[target_feature(enable = "avx2,fma")]
5005unsafe fn dot_q4t_row_1x4_avx2(bytes: &[u8], r: usize, gpr: usize, xs: [&[i8]; 4]) -> [f32; 4] {
5006    // SAFETY: callers uphold the 18B-tile and xq-length contracts.
5007    unsafe {
5008        use core::arch::x86_64::*;
5009        let lomask = _mm_set1_epi8(0x0F);
5010        let eight = _mm256_set1_epi8(8);
5011        let ones = _mm256_set1_epi16(1);
5012        // One f32 accumulator VECTOR per activation, reduced once at the
5013        // end. Folding each group's i32 lanes to a scalar inside the loop
5014        // costs an extracti128 + three shift/add + a movd — a cross-lane
5015        // dependency chain per (group, activation), 288 of them per row at
5016        // cols=2304. The per-group scale is what forces a float
5017        // accumulator; it does not force a horizontal sum.
5018        //
5019        // The four accumulators are NAMED, not an array: as `[__m256; 4]`
5020        // indexed by a loop variable LLVM keeps them in memory and every
5021        // group pays four 32-byte loads and stores. That alone made this
5022        // kernel 2x SLOWER than the per-group reduction it replaces
5023        // (measured on the EPYC box: 150 s vs 71 s for two 256² steps).
5024        let mut f0 = _mm256_setzero_ps();
5025        let mut f1 = _mm256_setzero_ps();
5026        let mut f2 = _mm256_setzero_ps();
5027        let mut f3 = _mm256_setzero_ps();
5028        for gi in 0..gpr {
5029            let t = bytes.as_ptr().add((r * gpr + gi) * Q4_TILE);
5030            let s = f16_to_f32(u16::from_le_bytes([*t, *t.add(1)]));
5031            let sv = _mm256_set1_ps(s);
5032            let bb = _mm_loadu_si128(t.add(2) as *const __m128i);
5033            let lo = _mm_and_si128(bb, lomask);
5034            let hi = _mm_and_si128(_mm_srli_epi16::<4>(bb), lomask);
5035            let w = _mm256_sub_epi8(
5036                _mm256_set_m128i(_mm_unpackhi_epi8(lo, hi), _mm_unpacklo_epi8(lo, hi)),
5037                eight,
5038            );
5039            let aw = _mm256_abs_epi8(w);
5040            let off = gi * GROUP_SIZE;
5041            let dot = |xq: &[i8]| {
5042                let x = _mm256_loadu_si256(xq.as_ptr().add(off) as *const __m256i);
5043                let p16 = _mm256_maddubs_epi16(aw, _mm256_sign_epi8(x, w));
5044                _mm256_cvtepi32_ps(_mm256_madd_epi16(p16, ones))
5045            };
5046            f0 = _mm256_fmadd_ps(dot(xs[0]), sv, f0);
5047            f1 = _mm256_fmadd_ps(dot(xs[1]), sv, f1);
5048            f2 = _mm256_fmadd_ps(dot(xs[2]), sv, f2);
5049            f3 = _mm256_fmadd_ps(dot(xs[3]), sv, f3);
5050        }
5051        [
5052            hsum256_ps(f0),
5053            hsum256_ps(f1),
5054            hsum256_ps(f2),
5055            hsum256_ps(f3),
5056        ]
5057    }
5058}
5059
5060/// Horizontal sum of eight f32 lanes — the one cross-lane reduction the
5061/// blocked kernels pay, once per row instead of once per group.
5062#[cfg(target_arch = "x86_64")]
5063#[target_feature(enable = "avx2")]
5064#[inline]
5065unsafe fn hsum256_ps(v: core::arch::x86_64::__m256) -> f32 {
5066    // SAFETY: pure register arithmetic on the caller's vector.
5067    unsafe {
5068        use core::arch::x86_64::*;
5069        let hi = _mm256_extractf128_ps::<1>(v);
5070        let s = _mm_add_ps(_mm256_castps256_ps128(v), hi);
5071        let s = _mm_add_ps(s, _mm_movehl_ps(s, s));
5072        let s = _mm_add_ss(s, _mm_shuffle_ps::<0x55>(s, s));
5073        _mm_cvtss_f32(s)
5074    }
5075}
5076
5077/// VNNI twin of `dot_q4t_row_1x4_avx2` (see `dpbusd_hsum`).
5078#[cfg(target_arch = "x86_64")]
5079#[target_feature(enable = "avx2,fma,avx512f,avx512bw,avx512vl,avx512vnni")]
5080unsafe fn dot_q4t_row_1x4_vnni(bytes: &[u8], r: usize, gpr: usize, xs: [&[i8]; 4]) -> [f32; 4] {
5081    // SAFETY: callers uphold the 18B-tile and xq-length contracts.
5082    unsafe {
5083        use core::arch::x86_64::*;
5084        let lomask = _mm_set1_epi8(0x0F);
5085        let eight = _mm256_set1_epi8(8);
5086        // Same shape as the AVX2 twin: accumulate in f32 vectors and pay
5087        // one cross-lane reduction per row, not per (group, activation).
5088        let mut f0 = _mm256_setzero_ps();
5089        let mut f1 = _mm256_setzero_ps();
5090        let mut f2 = _mm256_setzero_ps();
5091        let mut f3 = _mm256_setzero_ps();
5092        for gi in 0..gpr {
5093            let t = bytes.as_ptr().add((r * gpr + gi) * Q4_TILE);
5094            let s = f16_to_f32(u16::from_le_bytes([*t, *t.add(1)]));
5095            let sv = _mm256_set1_ps(s);
5096            let bb = _mm_loadu_si128(t.add(2) as *const __m128i);
5097            let lo = _mm_and_si128(bb, lomask);
5098            let hi = _mm_and_si128(_mm_srli_epi16::<4>(bb), lomask);
5099            let w = _mm256_sub_epi8(
5100                _mm256_set_m128i(_mm_unpackhi_epi8(lo, hi), _mm_unpacklo_epi8(lo, hi)),
5101                eight,
5102            );
5103            let aw = _mm256_abs_epi8(w);
5104            let off = gi * GROUP_SIZE;
5105            let dot = |xq: &[i8]| {
5106                let x = _mm256_loadu_si256(xq.as_ptr().add(off) as *const __m256i);
5107                _mm256_cvtepi32_ps(_mm256_dpbusd_epi32(
5108                    _mm256_setzero_si256(),
5109                    aw,
5110                    _mm256_sign_epi8(x, w),
5111                ))
5112            };
5113            f0 = _mm256_fmadd_ps(dot(xs[0]), sv, f0);
5114            f1 = _mm256_fmadd_ps(dot(xs[1]), sv, f1);
5115            f2 = _mm256_fmadd_ps(dot(xs[2]), sv, f2);
5116            f3 = _mm256_fmadd_ps(dot(xs[3]), sv, f3);
5117        }
5118        let acc = [
5119            hsum256_ps(f0),
5120            hsum256_ps(f1),
5121            hsum256_ps(f2),
5122            hsum256_ps(f3),
5123        ];
5124        acc
5125    }
5126}
5127
5128/// ARM twin of `dot_q4t_row_1x4_avx2`: one nibble unpack per group
5129/// serves FOUR activation streams. Per stream the group order and f32
5130/// accumulation match `dot_q4t_row_sdot` exactly — batch == matvec
5131/// bit-for-bit.
5132#[cfg(target_arch = "aarch64")]
5133#[target_feature(enable = "neon,dotprod")]
5134unsafe fn dot_q4t_row_1x4_sdot(bytes: &[u8], r: usize, gpr: usize, xs: [&[i8]; 4]) -> [f32; 4] {
5135    // SAFETY: callers uphold the 18B-tile and xq-length contracts.
5136    unsafe {
5137        use core::arch::aarch64::*;
5138        use core::arch::asm;
5139        let lomask = vdupq_n_u8(0x0F);
5140        let eight = vdupq_n_s8(8);
5141        let mut acc = [0f32; 4];
5142        for gi in 0..gpr {
5143            let t = bytes.as_ptr().add((r * gpr + gi) * Q4_TILE);
5144            let s = f16_to_f32(u16::from_le_bytes([*t, *t.add(1)]));
5145            let b = vld1q_u8(t.add(2));
5146            let lo = vandq_u8(b, lomask);
5147            let hi = vshrq_n_u8::<4>(b);
5148            let e0 = vsubq_s8(vreinterpretq_s8_u8(vzip1q_u8(lo, hi)), eight);
5149            let e1 = vsubq_s8(vreinterpretq_s8_u8(vzip2q_u8(lo, hi)), eight);
5150            for (k, xq) in xs.iter().enumerate() {
5151                let x0 = vld1q_s8(xq.as_ptr().add(gi * GROUP_SIZE));
5152                let x1 = vld1q_s8(xq.as_ptr().add(gi * GROUP_SIZE + 16));
5153                let (mut a0, mut a1) = (vdupq_n_s32(0), vdupq_n_s32(0));
5154                asm!(
5155                    "sdot {a0:v}.4s, {e0:v}.16b, {x0:v}.16b",
5156                    "sdot {a1:v}.4s, {e1:v}.16b, {x1:v}.16b",
5157                    a0 = inout(vreg) a0, a1 = inout(vreg) a1,
5158                    e0 = in(vreg) e0, x0 = in(vreg) x0, e1 = in(vreg) e1, x1 = in(vreg) x1,
5159                    options(pure, nomem, nostack),
5160                );
5161                acc[k] += vaddvq_s32(vaddq_s32(a0, a1)) as f32 * s;
5162            }
5163        }
5164        acc
5165    }
5166}
5167
5168/// Exact-term correction for A8W8 outliers on a tiled row.
5169#[inline]
5170fn q4t_outlier(bytes: &[u8], r: usize, gpr: usize, j: usize) -> (f32, f32) {
5171    let gi = j / GROUP_SIZE;
5172    let k = j % GROUP_SIZE;
5173    let tile = &bytes[(r * gpr + gi) * Q4_TILE..(r * gpr + gi + 1) * Q4_TILE];
5174    let s = f16_to_f32(u16::from_le_bytes([tile[0], tile[1]]));
5175    let byte = tile[2 + k / 2];
5176    let nib = if k & 1 == 0 { byte & 0x0F } else { byte >> 4 };
5177    ((nib as i32 - 8) as f32, s)
5178}
5179
5180/// Exact scalar q4_tiled row (CMF_SDOT=0 contract) — same pairwise
5181/// accumulation shape as `q4_range_f32`.
5182#[inline]
5183fn q4t_row_exact(bytes: &[u8], r: usize, gpr: usize, x: &[f32]) -> f32 {
5184    let mut acc = 0f32;
5185    for gi in 0..gpr {
5186        let tile = &bytes[(r * gpr + gi) * Q4_TILE..(r * gpr + gi + 1) * Q4_TILE];
5187        let s = f16_to_f32(u16::from_le_bytes([tile[0], tile[1]]));
5188        let xg = &x[gi * GROUP_SIZE..(gi + 1) * GROUP_SIZE];
5189        let mut ga = 0f32;
5190        for (k, &b) in tile[2..].iter().enumerate() {
5191            ga += ((b & 0x0F) as f32 - 8.0) * xg[k * 2]
5192                + (((b >> 4) & 0x0F) as f32 - 8.0) * xg[k * 2 + 1];
5193        }
5194        acc += ga * s;
5195    }
5196    acc
5197}
5198
5199/// Split view of a `q4tp` payload. The three planes are resolved once per
5200/// matvec instead of per row — `q4tp_sections` is cheap, but doing it inside
5201/// the row loop would put a division on the hot path for nothing.
5202struct Q4tpView<'a> {
5203    nib: &'a [u8],
5204    params: &'a [u8],
5205    codes: &'a [u8],
5206    stride: usize,
5207    /// q2tp reads the ladder with rung 0 = exact zero.
5208    zero_rung: bool,
5209}
5210
5211impl<'a> Q4tpView<'a> {
5212    fn new(bytes: &'a [u8], rows: usize, cols: usize) -> Self {
5213        let (params_off, codes_off, stride) = q4tp_sections(rows, cols);
5214        Self {
5215            nib: &bytes[..params_off],
5216            params: &bytes[params_off..codes_off],
5217            codes: &bytes[codes_off..],
5218            stride,
5219            zero_rung: false,
5220        }
5221    }
5222
5223    /// The q2tp view: identical params/codes planes, 8 B weight chunks.
5224    fn new_q2(bytes: &'a [u8], rows: usize, cols: usize) -> Self {
5225        let (params_off, codes_off, stride) = q2tp_sections(rows, cols);
5226        Self {
5227            nib: &bytes[..params_off],
5228            params: &bytes[params_off..codes_off],
5229            codes: &bytes[codes_off..],
5230            stride,
5231            zero_rung: true,
5232        }
5233    }
5234
5235    /// Expand row `r`'s per-tile scales into `out` (length `gpr`).
5236    ///
5237    /// Doing this once per row — rather than decoding a 5-bit code inside the
5238    /// tile loop — is what makes the format free at runtime. Random access to
5239    /// a packed 5-bit field costs a division, two bounds checks and a branch;
5240    /// the tile's actual work is two `sdot`s, so per-tile decoding dominated
5241    /// the kernel and cost 5x (measured: 1.4 vs 6.9 tok/s on Nanbeige-3B).
5242    /// Walking the plane sequentially with a bit accumulator is ~3 ops.
5243    /// Eight 5-bit codes are exactly five bytes, so a whole group of
5244    /// eight decodes from one little-endian word at fixed shifts. The
5245    /// bit-accumulator this replaces carried a data-dependent `while
5246    /// have < 5` refill whose branch sat in the innermost loop of every
5247    /// q4tp row; a decode profile put this function above the dot
5248    /// products it feeds. Same bitstream, same codes — just no branch
5249    /// and eight independent extractions.
5250    #[inline]
5251    fn scales_into(&self, r: usize, gpr: usize, out: &mut [f32]) {
5252        let tab = if self.zero_rung {
5253            q2tp_ladder(self.params, r)
5254        } else {
5255            q4tp_ladder(self.params, r)
5256        };
5257        let codes = &self.codes[r * self.stride..(r + 1) * self.stride];
5258        let out = &mut out[..gpr];
5259        let mut chunks = out.chunks_exact_mut(8);
5260        let mut ci = 0usize;
5261        for c in &mut chunks {
5262            let w = u64::from(codes[ci])
5263                | u64::from(codes[ci + 1]) << 8
5264                | u64::from(codes[ci + 2]) << 16
5265                | u64::from(codes[ci + 3]) << 24
5266                | u64::from(codes[ci + 4]) << 32;
5267            for (k, o) in c.iter_mut().enumerate() {
5268                *o = tab[((w >> (5 * k)) & 31) as usize];
5269            }
5270            ci += 5;
5271        }
5272        // Fewer than eight codes left: the shared total accessor, which
5273        // tolerates a 5-bit field whose spill byte is past the stride.
5274        let tail = &codes[ci..];
5275        for (k, o) in chunks.into_remainder().iter_mut().enumerate() {
5276            *o = tab[q4tp_code(tail, k)];
5277        }
5278    }
5279}
5280
5281#[inline]
5282fn dot_q4tp_row_i8(nib: &[u8], r: usize, gpr: usize, xq: &[i8], scales: &[f32]) -> f32 {
5283    #[cfg(target_arch = "aarch64")]
5284    unsafe {
5285        return dot_q4tp_row_sdot(nib, r, gpr, xq, scales);
5286    }
5287    #[cfg(target_arch = "x86_64")]
5288    unsafe {
5289        if vnni_tiles_enabled() {
5290            return dot_q4tp_row_vnni(nib, r, gpr, xq, scales);
5291        }
5292        return dot_q4tp_row_avx2(nib, r, gpr, xq, scales);
5293    }
5294    #[allow(unreachable_code)]
5295    {
5296        let mut acc = 0f32;
5297        for gi in 0..gpr {
5298            let tile = &nib[(r * gpr + gi) * Q4TP_NIB..(r * gpr + gi + 1) * Q4TP_NIB];
5299            let s = scales[gi];
5300            let mut d = 0i32;
5301            for (k, &b) in tile.iter().enumerate() {
5302                d += ((b & 0x0F) as i32 - 8) * xq[gi * GROUP_SIZE + k * 2] as i32
5303                    + (((b >> 4) & 0x0F) as i32 - 8) * xq[gi * GROUP_SIZE + k * 2 + 1] as i32;
5304            }
5305            acc += d as f32 * s;
5306        }
5307        acc
5308    }
5309}
5310
5311/// q4tp twin of `dot_q4t_row_sdot`: identical nibble math, but the tile
5312/// stride is 16 B (no inline scale) and the scale is a ladder lookup.
5313#[cfg(target_arch = "aarch64")]
5314#[target_feature(enable = "neon,dotprod")]
5315unsafe fn dot_q4tp_row_sdot(nib: &[u8], r: usize, gpr: usize, xq: &[i8], scales: &[f32]) -> f32 {
5316    // SAFETY: callers uphold slice-length contracts (16B tile per group,
5317    // xq.len() == gpr·GROUP_SIZE, codes covering gpr 5-bit fields).
5318    unsafe {
5319        use core::arch::aarch64::*;
5320        use core::arch::asm;
5321        let lomask = vdupq_n_u8(0x0F);
5322        let eight = vdupq_n_s8(8);
5323        let mut acc = 0f32;
5324        for gi in 0..gpr {
5325            let t = nib.as_ptr().add((r * gpr + gi) * Q4TP_NIB);
5326            let s = *scales.get_unchecked(gi);
5327            let b = vld1q_u8(t);
5328            let lo = vandq_u8(b, lomask);
5329            let hi = vshrq_n_u8::<4>(b);
5330            let e0 = vsubq_s8(vreinterpretq_s8_u8(vzip1q_u8(lo, hi)), eight);
5331            let e1 = vsubq_s8(vreinterpretq_s8_u8(vzip2q_u8(lo, hi)), eight);
5332            let x0 = vld1q_s8(xq.as_ptr().add(gi * GROUP_SIZE));
5333            let x1 = vld1q_s8(xq.as_ptr().add(gi * GROUP_SIZE + 16));
5334            let (mut a0, mut a1) = (vdupq_n_s32(0), vdupq_n_s32(0));
5335            asm!(
5336                "sdot {a0:v}.4s, {e0:v}.16b, {x0:v}.16b",
5337                "sdot {a1:v}.4s, {e1:v}.16b, {x1:v}.16b",
5338                a0 = inout(vreg) a0, a1 = inout(vreg) a1,
5339                e0 = in(vreg) e0, x0 = in(vreg) x0, e1 = in(vreg) e1, x1 = in(vreg) x1,
5340                options(pure, nomem, nostack),
5341            );
5342            acc += vaddvq_s32(vaddq_s32(a0, a1)) as f32 * s;
5343        }
5344        acc
5345    }
5346}
5347
5348#[cfg(target_arch = "x86_64")]
5349#[target_feature(enable = "avx2")]
5350unsafe fn dot_q4tp_row_avx2(nib: &[u8], r: usize, gpr: usize, xq: &[i8], scales: &[f32]) -> f32 {
5351    // SAFETY: see dot_q4tp_row_sdot.
5352    unsafe {
5353        use core::arch::x86_64::*;
5354        let lomask = _mm_set1_epi8(0x0F);
5355        let eight = _mm256_set1_epi8(8);
5356        let ones = _mm256_set1_epi16(1);
5357        let mut acc = 0f32;
5358        for gi in 0..gpr {
5359            let t = nib.as_ptr().add((r * gpr + gi) * Q4TP_NIB);
5360            let s = *scales.get_unchecked(gi);
5361            let b = _mm_loadu_si128(t as *const __m128i);
5362            let lo = _mm_and_si128(b, lomask);
5363            let hi = _mm_and_si128(_mm_srli_epi16::<4>(b), lomask);
5364            let w = _mm256_sub_epi8(
5365                _mm256_set_m128i(_mm_unpackhi_epi8(lo, hi), _mm_unpacklo_epi8(lo, hi)),
5366                eight,
5367            );
5368            let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
5369            let p16 = _mm256_maddubs_epi16(_mm256_abs_epi8(w), _mm256_sign_epi8(x, w));
5370            let d = _mm256_madd_epi16(p16, ones);
5371            let hi128 = _mm256_extracti128_si256::<1>(d);
5372            let s128 = _mm_add_epi32(_mm256_castsi256_si128(d), hi128);
5373            let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
5374            let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
5375            acc += _mm_cvtsi128_si32(s32) as f32 * s;
5376        }
5377        acc
5378    }
5379}
5380
5381/// VNNI twin of `dot_q4tp_row_avx2` (see `dot_q4t_row_vnni` for why the
5382/// 256-bit VL encoding is the one to use here).
5383#[cfg(target_arch = "x86_64")]
5384#[target_feature(enable = "avx2,avx512f,avx512bw,avx512vl,avx512vnni")]
5385unsafe fn dot_q4tp_row_vnni(nib: &[u8], r: usize, gpr: usize, xq: &[i8], scales: &[f32]) -> f32 {
5386    // SAFETY: see dot_q4tp_row_sdot.
5387    unsafe {
5388        use core::arch::x86_64::*;
5389        let lomask = _mm_set1_epi8(0x0F);
5390        let eight = _mm256_set1_epi8(8);
5391        let mut acc = 0f32;
5392        for gi in 0..gpr {
5393            let t = nib.as_ptr().add((r * gpr + gi) * Q4TP_NIB);
5394            let s = *scales.get_unchecked(gi);
5395            let b = _mm_loadu_si128(t as *const __m128i);
5396            let lo = _mm_and_si128(b, lomask);
5397            let hi = _mm_and_si128(_mm_srli_epi16::<4>(b), lomask);
5398            let w = _mm256_sub_epi8(
5399                _mm256_set_m128i(_mm_unpackhi_epi8(lo, hi), _mm_unpacklo_epi8(lo, hi)),
5400                eight,
5401            );
5402            let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
5403            acc += dpbusd_hsum(_mm256_abs_epi8(w), _mm256_sign_epi8(x, w)) as f32 * s;
5404        }
5405        acc
5406    }
5407}
5408
5409/// Exact scalar q4tp row — the `CMF_SDOT=0` contract, same pairwise
5410/// accumulation shape as `q4t_row_exact`.
5411#[inline]
5412fn q4tp_row_exact(nib: &[u8], r: usize, gpr: usize, x: &[f32], scales: &[f32]) -> f32 {
5413    #[cfg(target_arch = "x86_64")]
5414    if avx2_enabled() {
5415        // Keep the scalar pair/group reduction order, not a vector sum.
5416        return unsafe { q4tp_row_float_avx2(nib, r, gpr, x, scales) };
5417    }
5418    q4tp_row_float_scalar(nib, r, gpr, x, scales)
5419}
5420
5421#[inline]
5422fn q4tp_row_float_scalar(nib: &[u8], r: usize, gpr: usize, x: &[f32], scales: &[f32]) -> f32 {
5423    let mut acc = 0f32;
5424    for gi in 0..gpr {
5425        let tile = &nib[(r * gpr + gi) * Q4TP_NIB..(r * gpr + gi + 1) * Q4TP_NIB];
5426        let s = scales[gi];
5427        let xg = &x[gi * GROUP_SIZE..(gi + 1) * GROUP_SIZE];
5428        let mut ga = 0f32;
5429        for (k, &b) in tile.iter().enumerate() {
5430            ga += ((b & 0x0F) as f32 - 8.0) * xg[k * 2]
5431                + (((b >> 4) & 0x0F) as f32 - 8.0) * xg[k * 2 + 1];
5432        }
5433        acc += ga * s;
5434    }
5435    acc
5436}
5437
5438/// Vectorize unpack, conversion and multiplication, but preserve every
5439/// pair addition and the scalar accumulation order. No activation rounding
5440/// or FMA: bit-identical to the float scalar row, including its group scale.
5441#[cfg(target_arch = "x86_64")]
5442#[target_feature(enable = "avx2")]
5443unsafe fn q4tp_row_float_avx2(nib: &[u8], r: usize, gpr: usize, x: &[f32], scales: &[f32]) -> f32 {
5444    // SAFETY: caller checks AVX2 and provides the same complete 32-element
5445    // groups as the scalar row. Loads/stores are explicitly unaligned.
5446    unsafe {
5447        use core::arch::x86_64::*;
5448        let mask = _mm_set1_epi8(15);
5449        let eight = _mm_set1_epi8(8);
5450        let order = _mm256_setr_epi32(0, 1, 4, 5, 2, 3, 6, 7);
5451        let mut acc = 0.0f32;
5452        for gi in 0..gpr {
5453            let packed = _mm_loadu_si128(nib.as_ptr().add((r * gpr + gi) * Q4TP_NIB).cast());
5454            let lo = _mm_and_si128(packed, mask);
5455            let hi = _mm_and_si128(_mm_srli_epi16::<4>(packed), mask);
5456            let w0 = _mm_sub_epi8(_mm_unpacklo_epi8(lo, hi), eight);
5457            let w1 = _mm_sub_epi8(_mm_unpackhi_epi8(lo, hi), eight);
5458            let xp = x.as_ptr().add(gi * GROUP_SIZE);
5459            let a = _mm256_mul_ps(
5460                _mm256_cvtepi32_ps(_mm256_cvtepi8_epi32(w0)),
5461                _mm256_loadu_ps(xp),
5462            );
5463            let b = _mm256_mul_ps(
5464                _mm256_cvtepi32_ps(_mm256_cvtepi8_epi32(_mm_srli_si128::<8>(w0))),
5465                _mm256_loadu_ps(xp.add(8)),
5466            );
5467            let c = _mm256_mul_ps(
5468                _mm256_cvtepi32_ps(_mm256_cvtepi8_epi32(w1)),
5469                _mm256_loadu_ps(xp.add(16)),
5470            );
5471            let d = _mm256_mul_ps(
5472                _mm256_cvtepi32_ps(_mm256_cvtepi8_epi32(_mm_srli_si128::<8>(w1))),
5473                _mm256_loadu_ps(xp.add(24)),
5474            );
5475            let mut pairs = [0.0f32; 16];
5476            _mm256_storeu_ps(
5477                pairs.as_mut_ptr(),
5478                _mm256_permutevar8x32_ps(_mm256_hadd_ps(a, b), order),
5479            );
5480            _mm256_storeu_ps(
5481                pairs.as_mut_ptr().add(8),
5482                _mm256_permutevar8x32_ps(_mm256_hadd_ps(c, d), order),
5483            );
5484            let mut ga = 0.0f32;
5485            for v in pairs {
5486                ga += v;
5487            }
5488            acc += ga * scales[gi];
5489        }
5490        acc
5491    }
5492}
5493
5494/// Single weight of a q4tp tensor — the a8w8 outlier path, which restores
5495/// activation outliers at full precision after the int8 pass.
5496#[inline]
5497fn q4tp_outlier(nib: &[u8], r: usize, gpr: usize, j: usize, scales: &[f32]) -> (f32, f32) {
5498    let (gi, k) = (j / GROUP_SIZE, j % GROUP_SIZE);
5499    let byte = nib[(r * gpr + gi) * Q4TP_NIB + k / 2];
5500    let n = if k & 1 == 0 { byte & 0x0F } else { byte >> 4 };
5501    ((n as i32 - 8) as f32, scales[gi])
5502}
5503
5504/// Fused q4tp matvec (dispatch mirrors `q4t_matvec`).
5505fn q4tp_matvec(
5506    bytes: &[u8],
5507    x: &[f32],
5508    rows: usize,
5509    cols: usize,
5510    out: &mut [f32],
5511    pool: Option<&Pool>,
5512) {
5513    debug_assert_eq!(out.len(), rows);
5514    let gpr = cols / GROUP_SIZE;
5515    let v = Q4tpView::new(bytes, rows, cols);
5516    let out_addr = SendMut(out.as_mut_ptr());
5517    if a8w8_enabled() {
5518        let act = split_act(x);
5519        let run = |start: usize, end: usize| {
5520            // One scratch row of scales per worker — borrowed, not minted.
5521            with_krow(gpr, |sc| {
5522                for r in start..end {
5523                    v.scales_into(r, gpr, sc);
5524                    let mut acc = dot_q4tp_row_i8(v.nib, r, gpr, &act.xq, sc) * act.sx;
5525                    for &(j, xv) in &act.outliers {
5526                        let (w, s) = q4tp_outlier(v.nib, r, gpr, j, sc);
5527                        acc += w * s * xv;
5528                    }
5529                    // SAFETY: disjoint row ranges per worker.
5530                    unsafe { *out_addr.at(r) = acc };
5531                }
5532            })
5533        };
5534        dispatch_rows(pool, rows, &run);
5535        return;
5536    }
5537    let run = |start: usize, end: usize| {
5538        with_krow(gpr, |sc| {
5539            for r in start..end {
5540                v.scales_into(r, gpr, sc);
5541                // SAFETY: disjoint row ranges per worker.
5542                unsafe { *out_addr.at(r) = q4tp_row_exact(v.nib, r, gpr, x, sc) };
5543            }
5544        })
5545    };
5546    dispatch_rows(pool, rows, &run);
5547}
5548
5549/// Fused two-input q4tp matvec — the SwiGLU gate/up pair. Weights and the
5550/// row ladder are read once and spent on both activation streams.
5551#[allow(clippy::too_many_arguments)]
5552fn q4tp_matvec2(
5553    bytes: &[u8],
5554    x1: &[f32],
5555    x2: &[f32],
5556    rows: usize,
5557    cols: usize,
5558    o1: &mut [f32],
5559    o2: &mut [f32],
5560    pool: Option<&Pool>,
5561) {
5562    let gpr = cols / GROUP_SIZE;
5563    let v = Q4tpView::new(bytes, rows, cols);
5564    let (p1, p2) = (SendMut(o1.as_mut_ptr()), SendMut(o2.as_mut_ptr()));
5565    let run = |start: usize, end: usize| {
5566        let mut sc = vec![0f32; gpr];
5567        for r in start..end {
5568            v.scales_into(r, gpr, &mut sc);
5569            // SAFETY: disjoint row ranges per worker.
5570            unsafe {
5571                *p1.at(r) = q4tp_row_exact(v.nib, r, gpr, x1, &sc);
5572                *p2.at(r) = q4tp_row_exact(v.nib, r, gpr, x2, &sc);
5573            }
5574        }
5575    };
5576    dispatch_rows(pool, rows, &run);
5577}
5578
5579/// One q2tp outlier weight at column `j` of row `r`: the 2-bit code and
5580/// its group scale, mirrored on `q4tp_outlier`.
5581#[inline]
5582fn q2tp_outlier(chunks: &[u8], r: usize, gpr: usize, j: usize, scales: &[f32]) -> (f32, f32) {
5583    let (gi, k) = (j / GROUP_SIZE, j % GROUP_SIZE);
5584    let byte = chunks[(r * gpr + gi) * Q2TP_CHUNK + k / 4];
5585    let c = (byte >> (2 * (k % 4))) & 3;
5586    (c as f32 - 1.5, scales[gi])
5587}
5588
5589#[cfg(target_arch = "x86_64")]
5590const Q2TP_DECODE_U32: [u32; 256] = {
5591    let mut tab = [0u32; 256];
5592    let mut b = 0usize;
5593    while b < 256 {
5594        tab[b] = ((b as u32) & 3)
5595            | ((((b as u32) >> 2) & 3) << 8)
5596            | ((((b as u32) >> 4) & 3) << 16)
5597            | ((((b as u32) >> 6) & 3) << 24);
5598        b += 1;
5599    }
5600    tab
5601};
5602
5603/// Eight packed q2tp bytes against 32 signed activation bytes. `maddubs`
5604/// exactly computes unsigned 2-bit code × signed i8; its pair sums cannot
5605/// saturate (2 × 3 × 127 < i16::MAX), and the second madd widens to i32.
5606#[cfg(target_arch = "x86_64")]
5607#[target_feature(enable = "avx2")]
5608unsafe fn q2tp_code_dot_avx2(ch: &[u8], x: &[i8]) -> i32 {
5609    use core::arch::x86_64::*;
5610    debug_assert!(ch.len() >= Q2TP_CHUNK && x.len() >= GROUP_SIZE);
5611    let codes = _mm256_setr_epi32(
5612        Q2TP_DECODE_U32[ch[0] as usize] as i32,
5613        Q2TP_DECODE_U32[ch[1] as usize] as i32,
5614        Q2TP_DECODE_U32[ch[2] as usize] as i32,
5615        Q2TP_DECODE_U32[ch[3] as usize] as i32,
5616        Q2TP_DECODE_U32[ch[4] as usize] as i32,
5617        Q2TP_DECODE_U32[ch[5] as usize] as i32,
5618        Q2TP_DECODE_U32[ch[6] as usize] as i32,
5619        Q2TP_DECODE_U32[ch[7] as usize] as i32,
5620    );
5621    let xv = unsafe { _mm256_loadu_si256(x.as_ptr().cast()) };
5622    let pair = _mm256_maddubs_epi16(codes, xv);
5623    let quad = _mm256_madd_epi16(pair, _mm256_set1_epi16(1));
5624    let sum128 = _mm_add_epi32(
5625        _mm256_castsi256_si128(quad),
5626        _mm256_extracti128_si256(quad, 1),
5627    );
5628    let sum64 = _mm_hadd_epi32(sum128, sum128);
5629    _mm_cvtsi128_si32(_mm_hadd_epi32(sum64, sum64))
5630}
5631
5632#[cfg(target_arch = "x86_64")]
5633#[inline]
5634fn q2tp_avx2_enabled() -> bool {
5635    static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
5636    *ON.get_or_init(|| std::arch::is_x86_feature_detected!("avx2"))
5637}
5638
5639/// Integer dot of one q2tp row against pre-quantized activations:
5640/// Σ_g s_g · (Σ c·xq − 1.5·Σ xq). The half-integer grid (c − 1.5)
5641/// becomes exact integer math through the group sums — the same trick
5642/// every a8w8 kernel in this file rides. The codes decode into a
5643/// 32-byte scratch in natural order and the dot itself is the shared
5644/// SDOT primitive; elsewhere a scalar integer loop.
5645#[inline]
5646fn dot_q2tp_row_i8(
5647    chunks: &[u8],
5648    r: usize,
5649    gpr: usize,
5650    xq: &[i8],
5651    gsum: &[i32],
5652    scales: &[f32],
5653) -> f32 {
5654    let mut acc = 0f32;
5655    let base = r * gpr * Q2TP_CHUNK;
5656    #[cfg(not(any(target_arch = "aarch64", target_arch = "x86_64")))]
5657    let mut codes = [0i8; GROUP_SIZE];
5658    #[cfg(target_arch = "x86_64")]
5659    // Cache the process-wide AVX2 decision; probing it for every 32-weight
5660    // group adds a branch to the hottest Q2TP row loop while retaining the
5661    // established table-load AVX2 arithmetic (and without coupling this path
5662    // to the separate FMA-gated A8W8 switch).
5663    let avx2 = q2tp_avx2_enabled();
5664    for gi in 0..gpr {
5665        let ch = &chunks[base + gi * Q2TP_CHUNK..base + (gi + 1) * Q2TP_CHUNK];
5666        let xg = &xq[gi * GROUP_SIZE..(gi + 1) * GROUP_SIZE];
5667        #[cfg(target_arch = "aarch64")]
5668        // NEON: the byte's four 2-bit fields land in four lane vectors
5669        // (shift+mask), vld4 de-interleaves xq to match (xj[k] =
5670        // xq[4k+j]), widening MACs accumulate exactly in i32. A scalar
5671        // decode here cost as much as the dot it fed — the profile put
5672        // it at the top of the whole W2 decode.
5673        let dot = unsafe {
5674            use core::arch::aarch64::*;
5675            let b = vld1_u8(ch.as_ptr());
5676            let three = vdup_n_u8(3);
5677            let c0 = vreinterpret_s8_u8(vand_u8(b, three));
5678            let c1 = vreinterpret_s8_u8(vand_u8(vshr_n_u8(b, 2), three));
5679            let c2 = vreinterpret_s8_u8(vand_u8(vshr_n_u8(b, 4), three));
5680            let c3 = vreinterpret_s8_u8(vand_u8(vshr_n_u8(b, 6), three));
5681            let x4 = vld4_s8(xg.as_ptr());
5682            let mut acc4 = vdupq_n_s32(0);
5683            acc4 = vpadalq_s16(acc4, vmull_s8(c0, x4.0));
5684            acc4 = vpadalq_s16(acc4, vmull_s8(c1, x4.1));
5685            acc4 = vpadalq_s16(acc4, vmull_s8(c2, x4.2));
5686            acc4 = vpadalq_s16(acc4, vmull_s8(c3, x4.3));
5687            vaddvq_s32(acc4)
5688        };
5689        #[cfg(target_arch = "x86_64")]
5690        let dot: i32 = if avx2 {
5691            // SAFETY: the runtime feature check gates the target-feature body;
5692            // the group slices above are exactly 8 and 32 bytes long.
5693            unsafe { q2tp_code_dot_avx2(ch, xg) }
5694        } else {
5695            ch.iter()
5696                .enumerate()
5697                .map(|(k, &b)| {
5698                    ((b & 3) as i32) * xg[k * 4] as i32
5699                        + (((b >> 2) & 3) as i32) * xg[k * 4 + 1] as i32
5700                        + (((b >> 4) & 3) as i32) * xg[k * 4 + 2] as i32
5701                        + (((b >> 6) & 3) as i32) * xg[k * 4 + 3] as i32
5702                })
5703                .sum()
5704        };
5705        #[cfg(not(any(target_arch = "aarch64", target_arch = "x86_64")))]
5706        let dot: i32 = {
5707            for (k, &b) in ch.iter().enumerate() {
5708                codes[k * 4] = (b & 3) as i8;
5709                codes[k * 4 + 1] = ((b >> 2) & 3) as i8;
5710                codes[k * 4 + 2] = ((b >> 4) & 3) as i8;
5711                codes[k * 4 + 3] = ((b >> 6) & 3) as i8;
5712            }
5713            codes
5714                .iter()
5715                .zip(xg)
5716                .map(|(&c, &x)| c as i32 * x as i32)
5717                .sum()
5718        };
5719        acc += scales[gi] * (dot as f32 - 1.5 * gsum[gi] as f32);
5720    }
5721    acc
5722}
5723
5724/// Exact f32 dot of one q2tp row: 2-bit fields LSB-first, (c − 1.5)·s.
5725/// Scalar on purpose — the 2-bit class targets the GPU graph; the CPU
5726/// path exists for parity gates and small-machine fallback.
5727fn q2tp_row_exact(chunks: &[u8], r: usize, gpr: usize, x: &[f32], scales: &[f32]) -> f32 {
5728    q2tp_row_exact_center(chunks, r, gpr, x, scales, 1.5)
5729}
5730
5731/// Fused Prism affine row: the derived correction is applied inside the
5732/// decoded code, avoiding a second accumulated dot and avoiding cancellation
5733/// between `B=(c-1.5)s` and `+.5s` for long 5120/17408 rows.
5734#[inline]
5735fn q2tp_affine_row_exact(chunks: &[u8], r: usize, gpr: usize, x: &[f32], scales: &[f32]) -> f32 {
5736    q2tp_row_exact_center(chunks, r, gpr, x, scales, 1.0)
5737}
5738
5739#[inline]
5740fn q2tp_row_exact_center(
5741    chunks: &[u8],
5742    r: usize,
5743    gpr: usize,
5744    x: &[f32],
5745    scales: &[f32],
5746    center: f32,
5747) -> f32 {
5748    let mut acc = 0f32;
5749    for gi in 0..gpr {
5750        let ch = &chunks[(r * gpr + gi) * Q2TP_CHUNK..(r * gpr + gi + 1) * Q2TP_CHUNK];
5751        let s = scales[gi];
5752        let xb = &x[gi * GROUP_SIZE..(gi + 1) * GROUP_SIZE];
5753        let mut g = 0f32;
5754        for (k, &b) in ch.iter().enumerate() {
5755            g += ((b & 3) as f32 - center) * xb[k * 4]
5756                + (((b >> 2) & 3) as f32 - center) * xb[k * 4 + 1]
5757                + (((b >> 4) & 3) as f32 - center) * xb[k * 4 + 2]
5758                + (((b >> 6) & 3) as f32 - center) * xb[k * 4 + 3];
5759        }
5760        acc += s * g;
5761    }
5762    acc
5763}
5764
5765fn q2tp_matvec(
5766    bytes: &[u8],
5767    x: &[f32],
5768    rows: usize,
5769    cols: usize,
5770    out: &mut [f32],
5771    pool: Option<&Pool>,
5772) {
5773    q2tp_matvec_mode(bytes, x, rows, cols, out, pool, false);
5774}
5775
5776fn q2tp_affine_matvec(
5777    bytes: &[u8],
5778    x: &[f32],
5779    rows: usize,
5780    cols: usize,
5781    out: &mut [f32],
5782    pool: Option<&Pool>,
5783) {
5784    q2tp_matvec_mode(bytes, x, rows, cols, out, pool, true);
5785}
5786
5787fn q2tp_matvec_mode(
5788    bytes: &[u8],
5789    x: &[f32],
5790    rows: usize,
5791    cols: usize,
5792    out: &mut [f32],
5793    pool: Option<&Pool>,
5794    affine: bool,
5795) {
5796    debug_assert_eq!(out.len(), rows);
5797    let gpr = cols / GROUP_SIZE;
5798    let v = Q4tpView::new_q2(bytes, rows, cols);
5799    let out_addr = SendMut(out.as_mut_ptr());
5800    // a8w8 fast path (CMF_SDOT=0 keeps the exact scalar walk): integer
5801    // code dots + group sums, exact outlier correction — the same
5802    // contract as every sibling kernel; measured 2-bit rows were the
5803    // only scalar holdout in the family.
5804    if !affine && a8w8_enabled() {
5805        let act = split_act(x);
5806        let gsum = q1_group_sums(&act.xq, gpr);
5807        let (act, gsum) = (&act, &gsum);
5808        let run = move |start: usize, end: usize| {
5809            with_krow(gpr, |sc| {
5810                for r in start..end {
5811                    v.scales_into(r, gpr, sc);
5812                    let mut acc = dot_q2tp_row_i8(v.nib, r, gpr, &act.xq, gsum, sc) * act.sx;
5813                    for &(j, xv) in &act.outliers {
5814                        let (w, s) = q2tp_outlier(v.nib, r, gpr, j, sc);
5815                        acc += w * s * xv;
5816                    }
5817                    // SAFETY: disjoint row ranges per worker.
5818                    unsafe { *out_addr.at(r) = acc };
5819                }
5820            })
5821        };
5822        dispatch_rows(pool, rows, &run);
5823        return;
5824    }
5825    let run = |start: usize, end: usize| {
5826        with_krow(gpr, |sc| {
5827            for r in start..end {
5828                v.scales_into(r, gpr, sc);
5829                // SAFETY: disjoint row ranges per worker.
5830                unsafe {
5831                    *out_addr.at(r) = if affine {
5832                        q2tp_affine_row_exact(v.nib, r, gpr, x, sc)
5833                    } else {
5834                        q2tp_row_exact(v.nib, r, gpr, x, sc)
5835                    }
5836                };
5837            }
5838        })
5839    };
5840    dispatch_rows(pool, rows, &run);
5841}
5842
5843/// Fused two-input q2tp matvec — the SwiGLU gate/up pair.
5844#[allow(clippy::too_many_arguments)]
5845fn q2tp_matvec2(
5846    bytes: &[u8],
5847    x1: &[f32],
5848    x2: &[f32],
5849    rows: usize,
5850    cols: usize,
5851    o1: &mut [f32],
5852    o2: &mut [f32],
5853    pool: Option<&Pool>,
5854) {
5855    q2tp_matvec2_mode(bytes, x1, x2, rows, cols, o1, o2, pool, false);
5856}
5857
5858#[allow(clippy::too_many_arguments)]
5859fn q2tp_affine_matvec2(
5860    bytes: &[u8],
5861    x1: &[f32],
5862    x2: &[f32],
5863    rows: usize,
5864    cols: usize,
5865    o1: &mut [f32],
5866    o2: &mut [f32],
5867    pool: Option<&Pool>,
5868) {
5869    q2tp_matvec2_mode(bytes, x1, x2, rows, cols, o1, o2, pool, true);
5870}
5871
5872#[allow(clippy::too_many_arguments)]
5873fn q2tp_matvec2_mode(
5874    bytes: &[u8],
5875    x1: &[f32],
5876    x2: &[f32],
5877    rows: usize,
5878    cols: usize,
5879    o1: &mut [f32],
5880    o2: &mut [f32],
5881    pool: Option<&Pool>,
5882    affine: bool,
5883) {
5884    let gpr = cols / GROUP_SIZE;
5885    let v = Q4tpView::new_q2(bytes, rows, cols);
5886    let (p1, p2) = (SendMut(o1.as_mut_ptr()), SendMut(o2.as_mut_ptr()));
5887    let run = |start: usize, end: usize| {
5888        let mut sc = vec![0f32; gpr];
5889        for r in start..end {
5890            v.scales_into(r, gpr, &mut sc);
5891            // SAFETY: disjoint row ranges per worker.
5892            unsafe {
5893                *p1.at(r) = if affine {
5894                    q2tp_affine_row_exact(v.nib, r, gpr, x1, &sc)
5895                } else {
5896                    q2tp_row_exact(v.nib, r, gpr, x1, &sc)
5897                };
5898                *p2.at(r) = if affine {
5899                    q2tp_affine_row_exact(v.nib, r, gpr, x2, &sc)
5900                } else {
5901                    q2tp_row_exact(v.nib, r, gpr, x2, &sc)
5902                };
5903            }
5904        }
5905    };
5906    dispatch_rows(pool, rows, &run);
5907}
5908
5909/// Batched q2tp matmat: scalar row kernel over every batch column. CPU
5910/// prefill only — decode rides the graph, so plain and correct beats
5911/// clever here.
5912/// Test doors into the host 2-bit kernels: the stand's heap corruption
5913/// pointed at down-shaped tensors, and the private fns need a way to be
5914/// held to a reference without a model file around them.
5915pub fn q2tp_matvec_for_test(bytes: &[u8], x: &[f32], rows: usize, cols: usize, out: &mut [f32]) {
5916    // The facade IS the reference: encoder oracles hold requant output
5917    // to the exact scalar walk. The production dispatch may take the i8
5918    // fast path, whose error scale is the ACTIVATIONS' — a different
5919    // claim than the encoder correctness these tests pin.
5920    let gpr = cols / GROUP_SIZE;
5921    let v = Q4tpView::new_q2(bytes, rows, cols);
5922    with_krow(gpr, |sc| {
5923        for r in 0..rows {
5924            v.scales_into(r, gpr, sc);
5925            out[r] = q2tp_row_exact(v.nib, r, gpr, x, sc);
5926        }
5927    });
5928}
5929
5930/// Test door for the descriptor-specific fused affine decode.  Production
5931/// callers select this through a validated Prism header, never by dtype alone.
5932pub fn q2tp_affine_matvec_for_test(
5933    bytes: &[u8],
5934    x: &[f32],
5935    rows: usize,
5936    cols: usize,
5937    out: &mut [f32],
5938) {
5939    q2tp_affine_matvec(bytes, x, rows, cols, out, None);
5940}
5941
5942pub fn q2tp_matmat_for_test(
5943    bytes: &[u8],
5944    xs_all: &[f32],
5945    b: usize,
5946    rows: usize,
5947    cols: usize,
5948    out: &mut [f32],
5949) {
5950    q2tp_matmat(bytes, xs_all, b, rows, cols, out, None);
5951}
5952
5953fn q2tp_matmat(
5954    bytes: &[u8],
5955    xs_all: &[f32],
5956    b: usize,
5957    rows: usize,
5958    cols: usize,
5959    out: &mut [f32],
5960    pool: Option<&Pool>,
5961) {
5962    q2tp_matmat_mode(bytes, xs_all, b, rows, cols, out, pool, false);
5963}
5964
5965fn q2tp_affine_matmat(
5966    bytes: &[u8],
5967    xs_all: &[f32],
5968    b: usize,
5969    rows: usize,
5970    cols: usize,
5971    out: &mut [f32],
5972    pool: Option<&Pool>,
5973) {
5974    q2tp_matmat_mode(bytes, xs_all, b, rows, cols, out, pool, true);
5975}
5976
5977fn q2tp_matmat_mode(
5978    bytes: &[u8],
5979    xs_all: &[f32],
5980    b: usize,
5981    rows: usize,
5982    cols: usize,
5983    out: &mut [f32],
5984    pool: Option<&Pool>,
5985    affine: bool,
5986) {
5987    debug_assert_eq!(out.len(), b * rows);
5988    let gpr = cols / GROUP_SIZE;
5989    let v = Q4tpView::new_q2(bytes, rows, cols);
5990    let out_addr = SendMut(out.as_mut_ptr());
5991    let run = |start: usize, end: usize| {
5992        let mut sc = vec![0f32; gpr];
5993        for r in start..end {
5994            v.scales_into(r, gpr, &mut sc);
5995            for bi in 0..b {
5996                let x = &xs_all[bi * cols..(bi + 1) * cols];
5997                // SAFETY: disjoint row ranges per worker.
5998                unsafe {
5999                    *out_addr.at(bi * rows + r) = if affine {
6000                        q2tp_affine_row_exact(v.nib, r, gpr, x, &sc)
6001                    } else {
6002                        q2tp_row_exact(v.nib, r, gpr, x, &sc)
6003                    }
6004                };
6005            }
6006        }
6007    };
6008    dispatch_rows(pool, rows, &run);
6009}
6010
6011/// The pre-vectorised shape, kept for A/B (`CMF_Q4TP_V1=1`): the
6012/// horizontal add lands once per group per column instead of once per
6013/// row. Same weights, same activations — only the reduction differs.
6014/// It is also the row-exact batch kernel: per group and column it forms
6015/// `int dot as f32 * scale` and adds it to a scalar running sum, which is
6016/// `dot_q4tp_row_sdot` step for step (Rust never contracts to an fma),
6017/// so each column is bit-identical to that column's matvec.
6018#[cfg(target_arch = "aarch64")]
6019#[target_feature(enable = "neon,dotprod")]
6020unsafe fn dot_q4tp_row_1x4_sdot_v1(
6021    nib: &[u8],
6022    r: usize,
6023    gpr: usize,
6024    xs: [&[i8]; 4],
6025    scales: &[f32],
6026) -> [f32; 4] {
6027    unsafe {
6028        use core::arch::aarch64::*;
6029        use core::arch::asm;
6030        let lomask = vdupq_n_u8(0x0F);
6031        let eight = vdupq_n_s8(8);
6032        let (mut f0, mut f1, mut f2, mut f3) = (0f32, 0f32, 0f32, 0f32);
6033        for gi in 0..gpr {
6034            let t = nib.as_ptr().add((r * gpr + gi) * Q4TP_NIB);
6035            let s = *scales.get_unchecked(gi);
6036            let bb = vld1q_u8(t);
6037            let lo = vandq_u8(bb, lomask);
6038            let hi = vshrq_n_u8::<4>(bb);
6039            let e0 = vsubq_s8(vreinterpretq_s8_u8(vzip1q_u8(lo, hi)), eight);
6040            let e1 = vsubq_s8(vreinterpretq_s8_u8(vzip2q_u8(lo, hi)), eight);
6041            let mut d = [0f32; 4];
6042            for (k, dk) in d.iter_mut().enumerate() {
6043                let x0 = vld1q_s8(xs[k].as_ptr().add(gi * GROUP_SIZE));
6044                let x1 = vld1q_s8(xs[k].as_ptr().add(gi * GROUP_SIZE + 16));
6045                let (mut a0, mut a1) = (vdupq_n_s32(0), vdupq_n_s32(0));
6046                asm!(
6047                    "sdot {a0:v}.4s, {e0:v}.16b, {x0:v}.16b",
6048                    "sdot {a1:v}.4s, {e1:v}.16b, {x1:v}.16b",
6049                    a0 = inout(vreg) a0, a1 = inout(vreg) a1,
6050                    e0 = in(vreg) e0, x0 = in(vreg) x0, e1 = in(vreg) e1, x1 = in(vreg) x1,
6051                    options(pure, nomem, nostack),
6052                );
6053                *dk = vaddvq_s32(vaddq_s32(a0, a1)) as f32 * s;
6054            }
6055            f0 += d[0];
6056            f1 += d[1];
6057            f2 += d[2];
6058            f3 += d[3];
6059        }
6060        [f0, f1, f2, f3]
6061    }
6062}
6063
6064/// Which q4tp batch kernel to run: 1 = the previous one, 2 = the tuned
6065/// one, 0 = decide from the CPU. An atomic rather than a `OnceLock` so a
6066/// benchmark can alternate the two inside one process, where the machine's
6067/// mood — a shared box drifts ±25% between runs — is the same for both.
6068/// What the two mean is per-architecture: on x86 the blocked AVX-512 path
6069/// against the per-column one, on ARM the two reduction shapes.
6070#[allow(dead_code)]
6071static Q4TP_ALT: std::sync::atomic::AtomicU8 = std::sync::atomic::AtomicU8::new(0);
6072
6073/// Tests that store `Q4TP_ALT` hold this, so one test's kernel pick does
6074/// not leak into another's bit-exact comparison running in parallel.
6075#[cfg(test)]
6076static Q4TP_ALT_TEST_LOCK: std::sync::Mutex<()> = std::sync::Mutex::new(());
6077
6078/// Blocking pays on x86 only with 512-bit VNNI. With AVX2 alone, four
6079/// columns sharing an unpack still measured slower than the per-column
6080/// path (23.2 ms against 19.4 on a 48-thread EPYC), because that path
6081/// already dequantizes the row once — so the blocked kernel bought a
6082/// second unpack-free pass at the price of half the vector width.
6083#[cfg(target_arch = "x86_64")]
6084fn q4tp_blocked_x86() -> bool {
6085    match Q4TP_ALT.load(std::sync::atomic::Ordering::Relaxed) {
6086        1 => false,
6087        // A forced ON still asks the CPU. The switch exists so a bench can
6088        // pick a kernel, not so it can promise instructions the machine
6089        // does not have — CI caught that as a SIGILL on a runner without
6090        // AVX-512, where the parity test had turned the path on by hand.
6091        2 => avx512vnni_enabled(),
6092        // Deliberately not cached back into the switch: both gates below
6093        // hold their own `OnceLock`, and latching their answer here would
6094        // make a test's override outlive the test that set it.
6095        _ => blocked_enabled() && avx512vnni_enabled(),
6096    }
6097}
6098
6099/// `CMF_Q4TP_V1=1` picks the old reduction shape (A/B only).
6100#[cfg(target_arch = "aarch64")]
6101#[allow(dead_code)]
6102fn q4tp_v1() -> bool {
6103    match Q4TP_ALT.load(std::sync::atomic::Ordering::Relaxed) {
6104        1 => true,
6105        2 => false,
6106        _ => {
6107            static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
6108            *ON.get_or_init(|| std::env::var("CMF_Q4TP_V1").is_ok_and(|v| v != "0"))
6109        }
6110    }
6111}
6112
6113/// Two weight rows against eight columns. The activation load is the
6114/// same for both rows, so it is paid once for twice the arithmetic, and
6115/// sixteen accumulator chains run where eight did — which is what a kernel
6116/// retiring 0.29 instructions a cycle is short of. Register pressure is
6117/// the limit: sixteen `zmm` accumulators, two weight tiles, one
6118/// activation, of thirty-two.
6119///
6120/// Four rows by four columns spends the same sixteen accumulators the
6121/// other way and measured worse — 1488 GFLOP/s against 1644 — so the
6122/// unpack, which four rows pay twice as often, costs more than the extra
6123/// sharing of one activation load buys.
6124#[cfg(target_arch = "x86_64")]
6125#[target_feature(enable = "avx512f,avx512bw,avx512vnni")]
6126unsafe fn dot_q4tp_2x8_avx512(
6127    nib: &[u8],
6128    r0: usize,
6129    gpr: usize,
6130    xs: [&[i8]; 8],
6131    sc0: &[f32],
6132    sc1: &[f32],
6133) -> [[f32; 8]; 2] {
6134    // SAFETY: as dot_q4tp_row_1x8_avx512, two adjacent rows at once; the
6135    // caller guarantees r0 + 1 < rows and the ISA.
6136    unsafe {
6137        use core::arch::x86_64::*;
6138        let lomask = _mm256_set1_epi8(0x0F);
6139        let eight = _mm256_set1_epi8(8);
6140        let zero = _mm512_setzero_si512();
6141        let mut v0 = [_mm512_setzero_ps(); 8];
6142        let mut v1 = [_mm512_setzero_ps(); 8];
6143        let pairs = gpr / 2;
6144        let unpack = |r: usize, gi: usize| -> (__m512i, __mmask64) {
6145            let t = nib.as_ptr().add((r * gpr + gi) * Q4TP_NIB);
6146            let bb = _mm256_loadu_si256(t as *const __m256i);
6147            let lo = _mm256_and_si256(bb, lomask);
6148            let hi = _mm256_and_si256(_mm256_srli_epi16::<4>(bb), lomask);
6149            let ul = _mm256_sub_epi8(_mm256_unpacklo_epi8(lo, hi), eight);
6150            let uh = _mm256_sub_epi8(_mm256_unpackhi_epi8(lo, hi), eight);
6151            let cat = _mm512_inserti64x4::<1>(_mm512_castsi256_si512(ul), uh);
6152            let w = _mm512_shuffle_i64x2::<0b11_01_10_00>(cat, cat);
6153            (_mm512_abs_epi8(w), _mm512_movepi8_mask(w))
6154        };
6155        for gp in 0..pairs {
6156            let gi = gp * 2;
6157            let (wa0, neg0) = unpack(r0, gi);
6158            let (wa1, neg1) = unpack(r0 + 1, gi);
6159            let off = gi * GROUP_SIZE;
6160            let sv = |sc: &[f32]| {
6161                _mm512_insertf32x8::<1>(
6162                    _mm512_castps256_ps512(_mm256_set1_ps(*sc.get_unchecked(gi))),
6163                    _mm256_set1_ps(*sc.get_unchecked(gi + 1)),
6164                )
6165            };
6166            let s0 = sv(sc0);
6167            let s1 = sv(sc1);
6168            for k in 0..8 {
6169                let xv = _mm512_loadu_si512(xs[k].as_ptr().add(off) as *const __m512i);
6170                let d0 = _mm512_cvtepi32_ps(_mm512_dpbusd_epi32(
6171                    zero,
6172                    wa0,
6173                    _mm512_mask_sub_epi8(xv, neg0, zero, xv),
6174                ));
6175                let d1 = _mm512_cvtepi32_ps(_mm512_dpbusd_epi32(
6176                    zero,
6177                    wa1,
6178                    _mm512_mask_sub_epi8(xv, neg1, zero, xv),
6179                ));
6180                v0[k] = _mm512_fmadd_ps(d0, s0, v0[k]);
6181                v1[k] = _mm512_fmadd_ps(d1, s1, v1[k]);
6182            }
6183        }
6184        let mut acc = [[0f32; 8]; 2];
6185        for k in 0..8 {
6186            acc[0][k] = _mm512_reduce_add_ps(v0[k]);
6187            acc[1][k] = _mm512_reduce_add_ps(v1[k]);
6188        }
6189        if gpr % 2 == 1 {
6190            let off = (gpr - 1) * GROUP_SIZE;
6191            for j in off..off + GROUP_SIZE {
6192                let (w0, sa) = q4tp_outlier(nib, r0, gpr, j, sc0);
6193                let (w1, sb) = q4tp_outlier(nib, r0 + 1, gpr, j, sc1);
6194                for k in 0..8 {
6195                    let x = *xs[k].get_unchecked(j) as f32;
6196                    acc[0][k] += w0 * sa * x;
6197                    acc[1][k] += w1 * sb * x;
6198                }
6199            }
6200        }
6201        acc
6202    }
6203}
6204
6205/// The same, eight columns at a time. One unpack then feeds twice as many
6206/// activation streams, so a wide batch reads the weight tile half as
6207/// often; the price is eight accumulators live at once. Measured 9.0 ->
6208/// 8.3 ms at 9216x2304, b=296 on a 48-thread EPYC 9B45.
6209#[cfg(target_arch = "x86_64")]
6210#[target_feature(enable = "avx512f,avx512bw,avx512vnni")]
6211unsafe fn dot_q4tp_row_1x8_avx512(
6212    nib: &[u8],
6213    r: usize,
6214    gpr: usize,
6215    xs: [&[i8]; 8],
6216    scales: &[f32],
6217) -> [f32; 8] {
6218    // SAFETY: as dot_q4tp_row_1x4_avx2; caller guarantees the ISA.
6219    unsafe {
6220        use core::arch::x86_64::*;
6221        let lomask = _mm256_set1_epi8(0x0F);
6222        let eight = _mm256_set1_epi8(8);
6223        let zero = _mm512_setzero_si512();
6224        let (mut v0, mut v1, mut v2, mut v3) = (
6225            _mm512_setzero_ps(),
6226            _mm512_setzero_ps(),
6227            _mm512_setzero_ps(),
6228            _mm512_setzero_ps(),
6229        );
6230        let (mut v4, mut v5, mut v6, mut v7) = (
6231            _mm512_setzero_ps(),
6232            _mm512_setzero_ps(),
6233            _mm512_setzero_ps(),
6234            _mm512_setzero_ps(),
6235        );
6236        let pairs = gpr / 2;
6237        for gp in 0..pairs {
6238            let gi = gp * 2;
6239            let t = nib.as_ptr().add((r * gpr + gi) * Q4TP_NIB);
6240            let bb = _mm256_loadu_si256(t as *const __m256i);
6241            let lo = _mm256_and_si256(bb, lomask);
6242            let hi = _mm256_and_si256(_mm256_srli_epi16::<4>(bb), lomask);
6243            // `unpack` works per 128-bit lane, so the halves come out as
6244            // [A.lo, B.lo] and [A.hi, B.hi]; the shuffle reorders the four
6245            // 128-bit lanes into the weights' natural order, which is what
6246            // the straight activation load expects.
6247            let ul = _mm256_sub_epi8(_mm256_unpacklo_epi8(lo, hi), eight);
6248            let uh = _mm256_sub_epi8(_mm256_unpackhi_epi8(lo, hi), eight);
6249            let cat = _mm512_inserti64x4::<1>(_mm512_castsi256_si512(ul), uh);
6250            let w = _mm512_shuffle_i64x2::<0b11_01_10_00>(cat, cat);
6251            let wabs = _mm512_abs_epi8(w);
6252            let neg = _mm512_movepi8_mask(w);
6253            let off = gi * GROUP_SIZE;
6254            let sv = _mm512_insertf32x8::<1>(
6255                _mm512_castps256_ps512(_mm256_set1_ps(*scales.get_unchecked(gi))),
6256                _mm256_set1_ps(*scales.get_unchecked(gi + 1)),
6257            );
6258            let dot = |x: &[i8]| -> __m512 {
6259                let xv = _mm512_loadu_si512(x.as_ptr().add(off) as *const __m512i);
6260                let sx = _mm512_mask_sub_epi8(xv, neg, zero, xv);
6261                _mm512_cvtepi32_ps(_mm512_dpbusd_epi32(zero, wabs, sx))
6262            };
6263            v0 = _mm512_fmadd_ps(dot(xs[0]), sv, v0);
6264            v1 = _mm512_fmadd_ps(dot(xs[1]), sv, v1);
6265            v2 = _mm512_fmadd_ps(dot(xs[2]), sv, v2);
6266            v3 = _mm512_fmadd_ps(dot(xs[3]), sv, v3);
6267            v4 = _mm512_fmadd_ps(dot(xs[4]), sv, v4);
6268            v5 = _mm512_fmadd_ps(dot(xs[5]), sv, v5);
6269            v6 = _mm512_fmadd_ps(dot(xs[6]), sv, v6);
6270            v7 = _mm512_fmadd_ps(dot(xs[7]), sv, v7);
6271        }
6272        let mut acc = [
6273            _mm512_reduce_add_ps(v0),
6274            _mm512_reduce_add_ps(v1),
6275            _mm512_reduce_add_ps(v2),
6276            _mm512_reduce_add_ps(v3),
6277            _mm512_reduce_add_ps(v4),
6278            _mm512_reduce_add_ps(v5),
6279            _mm512_reduce_add_ps(v6),
6280            _mm512_reduce_add_ps(v7),
6281        ];
6282        // An odd group count leaves one group over; the narrow kernel
6283        // finishes it rather than the tail being a special case here.
6284        if gpr % 2 == 1 {
6285            let off = (gpr - 1) * GROUP_SIZE;
6286            for j in off..off + GROUP_SIZE {
6287                let (w, s) = q4tp_outlier(nib, r, gpr, j, scales);
6288                let ws = w * s;
6289                for k in 0..8 {
6290                    acc[k] += ws * *xs[k].get_unchecked(j) as f32;
6291                }
6292            }
6293        }
6294        acc
6295    }
6296}
6297
6298/// The same four columns, 512 bits wide. Two groups (64 weights) ride one
6299/// unpack and one `vpdpbusd`, where AVX2 needs two unpacks and four
6300/// `maddubs`/`madd` pairs — about 2.3x fewer instructions for the same
6301/// arithmetic. The two groups carry different scales, so the fma takes a
6302/// vector whose halves hold each group's scale rather than a broadcast.
6303///
6304/// There is no 512-bit `vpsignb`, so the activation's sign is applied by
6305/// negating under a mask taken from the weight's sign bits. That mask is
6306/// per-tile, so it is hoisted out of the column loop and the per-column
6307/// cost stays exactly one instruction, as with `sign_epi8`. Weights of
6308/// zero are not zeroed by the mask trick and do not need to be: their
6309/// magnitude is zero, so the product is.
6310#[cfg(target_arch = "x86_64")]
6311#[target_feature(enable = "avx512f,avx512bw,avx512vnni")]
6312unsafe fn dot_q4tp_row_1x4_avx512(
6313    nib: &[u8],
6314    r: usize,
6315    gpr: usize,
6316    xs: [&[i8]; 4],
6317    scales: &[f32],
6318) -> [f32; 4] {
6319    // SAFETY: as dot_q4tp_row_1x4_avx2; caller guarantees the ISA.
6320    unsafe {
6321        use core::arch::x86_64::*;
6322        let lomask = _mm256_set1_epi8(0x0F);
6323        let eight = _mm256_set1_epi8(8);
6324        let zero = _mm512_setzero_si512();
6325        let (mut v0, mut v1, mut v2, mut v3) = (
6326            _mm512_setzero_ps(),
6327            _mm512_setzero_ps(),
6328            _mm512_setzero_ps(),
6329            _mm512_setzero_ps(),
6330        );
6331        let pairs = gpr / 2;
6332        for gp in 0..pairs {
6333            let gi = gp * 2;
6334            let t = nib.as_ptr().add((r * gpr + gi) * Q4TP_NIB);
6335            let bb = _mm256_loadu_si256(t as *const __m256i);
6336            let lo = _mm256_and_si256(bb, lomask);
6337            let hi = _mm256_and_si256(_mm256_srli_epi16::<4>(bb), lomask);
6338            // `unpack` works per 128-bit lane, so the halves come out as
6339            // [A.lo, B.lo] and [A.hi, B.hi]; the shuffle reorders the four
6340            // 128-bit lanes into the weights' natural order, which is what
6341            // the straight activation load expects.
6342            let ul = _mm256_sub_epi8(_mm256_unpacklo_epi8(lo, hi), eight);
6343            let uh = _mm256_sub_epi8(_mm256_unpackhi_epi8(lo, hi), eight);
6344            let cat = _mm512_inserti64x4::<1>(_mm512_castsi256_si512(ul), uh);
6345            let w = _mm512_shuffle_i64x2::<0b11_01_10_00>(cat, cat);
6346            let wabs = _mm512_abs_epi8(w);
6347            let neg = _mm512_movepi8_mask(w);
6348            let off = gi * GROUP_SIZE;
6349            let sv = _mm512_insertf32x8::<1>(
6350                _mm512_castps256_ps512(_mm256_set1_ps(*scales.get_unchecked(gi))),
6351                _mm256_set1_ps(*scales.get_unchecked(gi + 1)),
6352            );
6353            let dot = |x: &[i8]| -> __m512 {
6354                let xv = _mm512_loadu_si512(x.as_ptr().add(off) as *const __m512i);
6355                let sx = _mm512_mask_sub_epi8(xv, neg, zero, xv);
6356                _mm512_cvtepi32_ps(_mm512_dpbusd_epi32(zero, wabs, sx))
6357            };
6358            v0 = _mm512_fmadd_ps(dot(xs[0]), sv, v0);
6359            v1 = _mm512_fmadd_ps(dot(xs[1]), sv, v1);
6360            v2 = _mm512_fmadd_ps(dot(xs[2]), sv, v2);
6361            v3 = _mm512_fmadd_ps(dot(xs[3]), sv, v3);
6362        }
6363        let mut acc = [
6364            _mm512_reduce_add_ps(v0),
6365            _mm512_reduce_add_ps(v1),
6366            _mm512_reduce_add_ps(v2),
6367            _mm512_reduce_add_ps(v3),
6368        ];
6369        // An odd group count leaves one group over; the narrow kernel
6370        // finishes it rather than the tail being a special case here.
6371        if gpr % 2 == 1 {
6372            let off = (gpr - 1) * GROUP_SIZE;
6373            for j in off..off + GROUP_SIZE {
6374                let (w, s) = q4tp_outlier(nib, r, gpr, j, scales);
6375                let ws = w * s;
6376                for k in 0..4 {
6377                    acc[k] += ws * *xs[k].get_unchecked(j) as f32;
6378                }
6379            }
6380        }
6381        acc
6382    }
6383}
6384
6385/// Four batch columns against one q4tp row: the tile is unpacked ONCE and
6386/// spent on four activation streams, which is where a prefill batch stops
6387/// being weight-bandwidth-bound. Twin of `dot_q4t_row_1x4_sdot`.
6388#[cfg(target_arch = "aarch64")]
6389#[target_feature(enable = "neon,dotprod")]
6390unsafe fn dot_q4tp_row_1x4_sdot(
6391    nib: &[u8],
6392    r: usize,
6393    gpr: usize,
6394    xs: [&[i8]; 4],
6395    scales: &[f32],
6396) -> [f32; 4] {
6397    // SAFETY: see dot_q4tp_row_sdot; every xs[k] is gpr·GROUP_SIZE long.
6398    unsafe {
6399        use core::arch::aarch64::*;
6400        use core::arch::asm;
6401        let lomask = vdupq_n_u8(0x0F);
6402        let eight = vdupq_n_s8(8);
6403        // Named accumulators, NOT an array indexed by a loop variable: the
6404        // latter does not stay in registers (the same defect cost 2x in the
6405        // AVX2 q4t kernel and again in WGSL).
6406        //
6407        // They are VECTORS, and the horizontal add happens once at the end
6408        // instead of once per group per column. `vaddvq` is a cross-lane
6409        // reduction — with 72 groups and four columns the old shape paid
6410        // 288 of them per row, each one a dependency stall the pipeline
6411        // cannot hide, to save four float adds. The group's scale now
6412        // rides an fma into the lane accumulators, so the arithmetic per
6413        // group is one convert and one fma. Summation order changes (the
6414        // lanes carry independent partial sums), which is the same
6415        // round-off class the SDOT path already lives in — the strict
6416        // kernel (`CMF_SDOT=0`, what `cortiq ppl` runs) is unchanged and
6417        // stays the reference.
6418        let (mut v0, mut v1, mut v2, mut v3) = (
6419            vdupq_n_f32(0.0),
6420            vdupq_n_f32(0.0),
6421            vdupq_n_f32(0.0),
6422            vdupq_n_f32(0.0),
6423        );
6424        for gi in 0..gpr {
6425            let t = nib.as_ptr().add((r * gpr + gi) * Q4TP_NIB);
6426            let s = *scales.get_unchecked(gi);
6427            let bb = vld1q_u8(t);
6428            let lo = vandq_u8(bb, lomask);
6429            let hi = vshrq_n_u8::<4>(bb);
6430            let e0 = vsubq_s8(vreinterpretq_s8_u8(vzip1q_u8(lo, hi)), eight);
6431            let e1 = vsubq_s8(vreinterpretq_s8_u8(vzip2q_u8(lo, hi)), eight);
6432            let off = gi * GROUP_SIZE;
6433            let dot4 = |x: &[i8]| -> int32x4_t {
6434                let x0 = vld1q_s8(x.as_ptr().add(off));
6435                let x1 = vld1q_s8(x.as_ptr().add(off + 16));
6436                let (mut a0, mut a1) = (vdupq_n_s32(0), vdupq_n_s32(0));
6437                asm!(
6438                    "sdot {a0:v}.4s, {e0:v}.16b, {x0:v}.16b",
6439                    "sdot {a1:v}.4s, {e1:v}.16b, {x1:v}.16b",
6440                    a0 = inout(vreg) a0, a1 = inout(vreg) a1,
6441                    e0 = in(vreg) e0, x0 = in(vreg) x0, e1 = in(vreg) e1, x1 = in(vreg) x1,
6442                    options(pure, nomem, nostack),
6443                );
6444                vaddq_s32(a0, a1)
6445            };
6446            v0 = vfmaq_n_f32(v0, vcvtq_f32_s32(dot4(xs[0])), s);
6447            v1 = vfmaq_n_f32(v1, vcvtq_f32_s32(dot4(xs[1])), s);
6448            v2 = vfmaq_n_f32(v2, vcvtq_f32_s32(dot4(xs[2])), s);
6449            v3 = vfmaq_n_f32(v3, vcvtq_f32_s32(dot4(xs[3])), s);
6450        }
6451        [
6452            vaddvq_f32(v0),
6453            vaddvq_f32(v1),
6454            vaddvq_f32(v2),
6455            vaddvq_f32(v3),
6456        ]
6457    }
6458}
6459
6460/// Fused q4tp matmat — the same three arms `q4t_matmat` has. Shipping only
6461/// the scalar one made Nanbeige-3B decode at 1.2 tok/s against q4t's 5.9:
6462/// the format was fine, the missing arms were the whole regression.
6463///
6464/// Under `row_exact()` every cell is summed in `q4tp_matvec`'s order on
6465/// every architecture; the mode is read once, so one call never mixes
6466/// the two contracts when a concurrent scope opens or closes mid-call.
6467fn q4tp_matmat(
6468    bytes: &[u8],
6469    xs_all: &[f32],
6470    b: usize,
6471    rows: usize,
6472    cols: usize,
6473    out: &mut [f32],
6474    pool: Option<&Pool>,
6475) {
6476    q4tp_matmat_with(bytes, xs_all, b, rows, cols, out, pool, row_exact())
6477}
6478
6479/// `q4tp_matmat` with the row-exact mode passed in rather than read from
6480/// the shared counter: `exact` = each cell equals its token's matvec bit
6481/// for bit, otherwise the fast blocked / AMX arms are free to reorder.
6482#[allow(clippy::too_many_arguments)]
6483fn q4tp_matmat_with(
6484    bytes: &[u8],
6485    xs_all: &[f32],
6486    b: usize,
6487    rows: usize,
6488    cols: usize,
6489    out: &mut [f32],
6490    pool: Option<&Pool>,
6491    exact: bool,
6492) {
6493    debug_assert_eq!(out.len(), b * rows);
6494    let gpr = cols / GROUP_SIZE;
6495    let v = Q4tpView::new(bytes, rows, cols);
6496
6497    // Wide batches ride the AMX through a dequant-tile sgemm, as in q4t.
6498    // An f32 GEMM over dequantized weights is not the matvec's int8 sum,
6499    // so a row-exact batch never takes it.
6500    #[cfg(target_os = "macos")]
6501    if !exact && b >= 8 && rows * cols >= 500_000 && accel_gemm_enabled() {
6502        dequant_matmat_accel(
6503            &|r, dst| {
6504                let mut sc = [0f32; 32];
6505                let mut scv;
6506                let s: &[f32] = if gpr <= 32 {
6507                    v.scales_into(r, gpr, &mut sc);
6508                    &sc[..gpr]
6509                } else {
6510                    scv = vec![0f32; gpr];
6511                    v.scales_into(r, gpr, &mut scv);
6512                    &scv
6513                };
6514                for gi in 0..gpr {
6515                    let tile = &v.nib[(r * gpr + gi) * Q4TP_NIB..(r * gpr + gi + 1) * Q4TP_NIB];
6516                    for (k, &bb) in tile.iter().enumerate() {
6517                        dst[gi * GROUP_SIZE + k * 2] = ((bb & 0x0F) as f32 - 8.0) * s[gi];
6518                        dst[gi * GROUP_SIZE + k * 2 + 1] =
6519                            (((bb >> 4) & 0x0F) as f32 - 8.0) * s[gi];
6520                    }
6521                }
6522            },
6523            xs_all,
6524            b,
6525            rows,
6526            cols,
6527            out,
6528            pool,
6529        );
6530        return;
6531    }
6532
6533    let out_addr = SendMut(out.as_mut_ptr());
6534    if a8w8_enabled() {
6535        let acts: Vec<SplitAct> = (0..b)
6536            .map(|bi| split_act(&xs_all[bi * cols..(bi + 1) * cols]))
6537            .collect();
6538        let acts = &acts;
6539        // ARM stays blocked under `exact` too: the 1x4 kernel then runs in
6540        // its v1 shape, whose per-group `int dot as f32 * scale` and scalar
6541        // running sum are exactly `dot_q4tp_row_sdot`'s, so the tile is
6542        // still unpacked once for four columns and each column equals its
6543        // matvec. The tuned shape (fma into lane partials, one horizontal
6544        // add per row) is 1-16 ulp off the matvec and is kept for `!exact`.
6545        #[cfg(target_arch = "aarch64")]
6546        let blocked_ok = sdot_enabled() && blocked_enabled();
6547        // x86 gets the same blocking: one tile unpack spent on four
6548        // columns. Without it every column re-decoded the row, which is
6549        // why a 48-core EPYC measured a sixth of an M4's per-core rate.
6550        // The gate is `avx2_enabled`, as in q4t — `sdot_enabled` answers
6551        // for ARM's dotprod and is hard-wired false everywhere else, so
6552        // asking it here left the whole blocked path unreachable on x86.
6553        #[cfg(target_arch = "x86_64")]
6554        let blocked_ok = q4tp_blocked_x86() && !exact;
6555        #[cfg(not(any(target_arch = "aarch64", target_arch = "x86_64")))]
6556        let blocked_ok = {
6557            let _ = exact;
6558            false
6559        };
6560        // Columns are swept in panels that fit L2. Without this a
6561        // row-pair walks every activation in the batch — 4.8 MB at
6562        // 512x512 — and does it again for the next pair, so the whole
6563        // batch streams out of the shared cache once per row. Measured
6564        // 800 GB/s of it, flat across batch sizes, which is the signature
6565        // of a loop bound by traffic rather than by arithmetic. A panel of
6566        // 256 columns is 590 KB beside 221 KB of this worker's weights:
6567        // both stay resident and the batch crosses L3 once instead of
6568        // once per row.
6569        let panel_cols: usize = std::env::var("CMF_Q4TP_PANEL")
6570            .ok()
6571            .and_then(|v| v.parse().ok())
6572            .filter(|v| *v > 0)
6573            .unwrap_or(256);
6574        let run = |start: usize, end: usize| {
6575            for abase in (0..acts.len()).step_by(panel_cols) {
6576                let alen = (acts.len() - abase).min(panel_cols);
6577                let mut sc = vec![0f32; gpr];
6578                #[cfg(target_arch = "x86_64")]
6579                let mut r_lo = start;
6580                #[cfg(target_arch = "x86_64")]
6581                if blocked_ok && alen >= 8 {
6582                    let mut sc1 = vec![0f32; gpr];
6583                    while r_lo + 2 <= end {
6584                        v.scales_into(r_lo, gpr, &mut sc);
6585                        v.scales_into(r_lo + 1, gpr, &mut sc1);
6586                        let mut bi = 0usize;
6587                        while bi + 8 <= alen {
6588                            let xs = [
6589                                acts[abase + bi].xq.as_slice(),
6590                                acts[abase + bi + 1].xq.as_slice(),
6591                                acts[abase + bi + 2].xq.as_slice(),
6592                                acts[abase + bi + 3].xq.as_slice(),
6593                                acts[abase + bi + 4].xq.as_slice(),
6594                                acts[abase + bi + 5].xq.as_slice(),
6595                                acts[abase + bi + 6].xq.as_slice(),
6596                                acts[abase + bi + 7].xq.as_slice(),
6597                            ];
6598                            let d = unsafe { dot_q4tp_2x8_avx512(v.nib, r_lo, gpr, xs, &sc, &sc1) };
6599                            for (row, dr, scr) in [(r_lo, &d[0], &sc), (r_lo + 1, &d[1], &sc1)] {
6600                                for k in 0..8 {
6601                                    let act = &acts[abase + bi + k];
6602                                    let mut acc = dr[k] * act.sx;
6603                                    for &(j, xv) in &act.outliers {
6604                                        let (w, s) = q4tp_outlier(v.nib, row, gpr, j, scr);
6605                                        acc += w * s * xv;
6606                                    }
6607                                    // SAFETY: disjoint (bi, r) cells per worker.
6608                                    unsafe { *out_addr.at((abase + bi + k) * rows + row) = acc };
6609                                }
6610                            }
6611                            bi += 8;
6612                        }
6613                        // Columns past the last group of eight, both rows —
6614                        // the same single-row kernel the tail below uses.
6615                        for row in [r_lo, r_lo + 1] {
6616                            let scr: &[f32] = if row == r_lo { &sc } else { &sc1 };
6617                            for b2 in bi..alen {
6618                                let act = &acts[abase + b2];
6619                                let xs4 = [
6620                                    act.xq.as_slice(),
6621                                    act.xq.as_slice(),
6622                                    act.xq.as_slice(),
6623                                    act.xq.as_slice(),
6624                                ];
6625                                let d =
6626                                    unsafe { dot_q4tp_row_1x4_avx512(v.nib, row, gpr, xs4, scr) };
6627                                let mut acc = d[0] * act.sx;
6628                                for &(j, xv) in &act.outliers {
6629                                    let (w, s) = q4tp_outlier(v.nib, row, gpr, j, scr);
6630                                    acc += w * s * xv;
6631                                }
6632                                // SAFETY: disjoint (bi, r) cells per worker.
6633                                unsafe { *out_addr.at((abase + b2) * rows + row) = acc };
6634                            }
6635                        }
6636                        r_lo += 2;
6637                    }
6638                }
6639                #[cfg(target_arch = "x86_64")]
6640                let row_start = r_lo;
6641                #[cfg(not(target_arch = "x86_64"))]
6642                let row_start = start;
6643                for r in row_start..end {
6644                    v.scales_into(r, gpr, &mut sc);
6645                    let mut bi = 0usize;
6646                    #[cfg(target_arch = "x86_64")]
6647                    if blocked_ok {
6648                        while bi + 8 <= alen {
6649                            let xs = [
6650                                acts[abase + bi].xq.as_slice(),
6651                                acts[abase + bi + 1].xq.as_slice(),
6652                                acts[abase + bi + 2].xq.as_slice(),
6653                                acts[abase + bi + 3].xq.as_slice(),
6654                                acts[abase + bi + 4].xq.as_slice(),
6655                                acts[abase + bi + 5].xq.as_slice(),
6656                                acts[abase + bi + 6].xq.as_slice(),
6657                                acts[abase + bi + 7].xq.as_slice(),
6658                            ];
6659                            let d = unsafe { dot_q4tp_row_1x8_avx512(v.nib, r, gpr, xs, &sc) };
6660                            for k in 0..8 {
6661                                let act = &acts[abase + bi + k];
6662                                let mut acc = d[k] * act.sx;
6663                                for &(j, xv) in &act.outliers {
6664                                    let (w, s) = q4tp_outlier(v.nib, r, gpr, j, &sc);
6665                                    acc += w * s * xv;
6666                                }
6667                                // SAFETY: disjoint (bi, r) cells per worker.
6668                                unsafe { *out_addr.at((abase + bi + k) * rows + r) = acc };
6669                            }
6670                            bi += 8;
6671                        }
6672                        while bi + 4 <= alen {
6673                            let xs = [
6674                                acts[abase + bi].xq.as_slice(),
6675                                acts[abase + bi + 1].xq.as_slice(),
6676                                acts[abase + bi + 2].xq.as_slice(),
6677                                acts[abase + bi + 3].xq.as_slice(),
6678                            ];
6679                            let d = unsafe { dot_q4tp_row_1x4_avx512(v.nib, r, gpr, xs, &sc) };
6680                            for k in 0..4 {
6681                                let act = &acts[abase + bi + k];
6682                                let mut acc = d[k] * act.sx;
6683                                for &(j, xv) in &act.outliers {
6684                                    let (w, s) = q4tp_outlier(v.nib, r, gpr, j, &sc);
6685                                    acc += w * s * xv;
6686                                }
6687                                // SAFETY: disjoint (bi, r) cells per worker.
6688                                unsafe { *out_addr.at((abase + bi + k) * rows + r) = acc };
6689                            }
6690                            bi += 4;
6691                        }
6692                    }
6693                    #[cfg(target_arch = "aarch64")]
6694                    if blocked_ok {
6695                        while bi + 4 <= alen {
6696                            let xs = [
6697                                acts[abase + bi].xq.as_slice(),
6698                                acts[abase + bi + 1].xq.as_slice(),
6699                                acts[abase + bi + 2].xq.as_slice(),
6700                                acts[abase + bi + 3].xq.as_slice(),
6701                            ];
6702                            let d = unsafe {
6703                                if exact || q4tp_v1() {
6704                                    dot_q4tp_row_1x4_sdot_v1(v.nib, r, gpr, xs, &sc)
6705                                } else {
6706                                    dot_q4tp_row_1x4_sdot(v.nib, r, gpr, xs, &sc)
6707                                }
6708                            };
6709                            for k in 0..4 {
6710                                let act = &acts[abase + bi + k];
6711                                let mut acc = d[k] * act.sx;
6712                                for &(j, xv) in &act.outliers {
6713                                    let (w, s) = q4tp_outlier(v.nib, r, gpr, j, &sc);
6714                                    acc += w * s * xv;
6715                                }
6716                                // SAFETY: disjoint (bi, r) cells per worker.
6717                                unsafe { *out_addr.at((abase + bi + k) * rows + r) = acc };
6718                            }
6719                            bi += 4;
6720                        }
6721                    }
6722                    let _ = blocked_ok;
6723                    while bi < alen {
6724                        let act = &acts[abase + bi];
6725                        let mut acc = dot_q4tp_row_i8(v.nib, r, gpr, &act.xq, &sc) * act.sx;
6726                        for &(j, xv) in &act.outliers {
6727                            let (w, s) = q4tp_outlier(v.nib, r, gpr, j, &sc);
6728                            acc += w * s * xv;
6729                        }
6730                        // SAFETY: disjoint (bi, r) cells per worker range.
6731                        unsafe { *out_addr.at((abase + bi) * rows + r) = acc };
6732                        bi += 1;
6733                    }
6734                }
6735            }
6736        };
6737        dispatch_rows(pool, rows, &run);
6738        return;
6739    }
6740
6741    let run = |start: usize, end: usize| {
6742        let mut sc = vec![0f32; gpr];
6743        for r in start..end {
6744            v.scales_into(r, gpr, &mut sc);
6745            for bi in 0..b {
6746                let x = &xs_all[bi * cols..(bi + 1) * cols];
6747                // SAFETY: disjoint (bi, r) cells per worker range.
6748                unsafe { *out_addr.at(bi * rows + r) = q4tp_row_exact(v.nib, r, gpr, x, &sc) };
6749            }
6750        }
6751    };
6752    dispatch_rows(pool, rows, &run);
6753}
6754
6755/// Fused q4_tiled matvec (dispatch mirrors `q4matvec`).
6756fn q4t_matvec(
6757    bytes: &[u8],
6758    x: &[f32],
6759    rows: usize,
6760    cols: usize,
6761    out: &mut [f32],
6762    pool: Option<&Pool>,
6763) {
6764    debug_assert_eq!(out.len(), rows);
6765    let gpr = cols / GROUP_SIZE;
6766    let out_addr = SendMut(out.as_mut_ptr());
6767    if a8w8_enabled() {
6768        let act = split_act(x);
6769        let run = move |start: usize, end: usize| {
6770            for r in start..end {
6771                let mut acc = dot_q4t_row_i8(bytes, r, gpr, &act.xq) * act.sx;
6772                for &(j, xv) in &act.outliers {
6773                    let (w, s) = q4t_outlier(bytes, r, gpr, j);
6774                    acc += w * s * xv;
6775                }
6776                // SAFETY: disjoint row ranges per worker.
6777                unsafe { *out_addr.at(r) = acc };
6778            }
6779        };
6780        dispatch_rows(pool, rows, &run);
6781        return;
6782    }
6783    let run = move |start: usize, end: usize| {
6784        for r in start..end {
6785            // SAFETY: disjoint row ranges per worker.
6786            unsafe { *out_addr.at(r) = q4t_row_exact(bytes, r, gpr, x) };
6787        }
6788    };
6789    dispatch_rows(pool, rows, &run);
6790}
6791
6792/// Fused two-input q4_tiled matvec (weights read once per pair).
6793#[allow(clippy::too_many_arguments)]
6794fn q4t_matvec2(
6795    bytes: &[u8],
6796    x1: &[f32],
6797    x2: &[f32],
6798    rows: usize,
6799    cols: usize,
6800    o1: &mut [f32],
6801    o2: &mut [f32],
6802    pool: Option<&Pool>,
6803) {
6804    let gpr = cols / GROUP_SIZE;
6805    let p1 = SendMut(o1.as_mut_ptr());
6806    let p2 = SendMut(o2.as_mut_ptr());
6807    if a8w8_enabled() {
6808        let a1 = split_act(x1);
6809        let a2 = split_act(x2);
6810        let run = move |start: usize, end: usize| {
6811            for r in start..end {
6812                let mut v1 = dot_q4t_row_i8(bytes, r, gpr, &a1.xq) * a1.sx;
6813                let mut v2 = dot_q4t_row_i8(bytes, r, gpr, &a2.xq) * a2.sx;
6814                for &(j, xv) in &a1.outliers {
6815                    let (w, s) = q4t_outlier(bytes, r, gpr, j);
6816                    v1 += w * s * xv;
6817                }
6818                for &(j, xv) in &a2.outliers {
6819                    let (w, s) = q4t_outlier(bytes, r, gpr, j);
6820                    v2 += w * s * xv;
6821                }
6822                // SAFETY: disjoint row ranges per worker.
6823                unsafe {
6824                    *p1.at(r) = v1;
6825                    *p2.at(r) = v2;
6826                }
6827            }
6828        };
6829        dispatch_rows(pool, rows, &run);
6830        return;
6831    }
6832    let run = move |start: usize, end: usize| {
6833        for r in start..end {
6834            // SAFETY: disjoint row ranges per worker.
6835            unsafe {
6836                *p1.at(r) = q4t_row_exact(bytes, r, gpr, x1);
6837                *p2.at(r) = q4t_row_exact(bytes, r, gpr, x2);
6838            }
6839        }
6840    };
6841    dispatch_rows(pool, rows, &run);
6842}
6843
6844/// Batched q4_tiled matmat: each row's tiles stream once per microbatch.
6845#[allow(clippy::too_many_arguments)]
6846/// Prefill GEMM through Accelerate for group-quantized codecs: a
6847/// caller-supplied row dequantizer fills f32 tiles (pool-parallel) and
6848/// each tile rides the AMX with one sgemm — the generic sibling of
6849/// `qmatmat_accel` (q8). Numerics are f32-GEMM (tolerance class);
6850/// decode (b=1) never takes this path.
6851#[cfg(target_os = "macos")]
6852fn dequant_matmat_accel(
6853    dequant_row: &(dyn Fn(usize, &mut [f32]) + Sync),
6854    xs_all: &[f32],
6855    b: usize,
6856    rows: usize,
6857    cols: usize,
6858    out: &mut [f32],
6859    pool: Option<&Pool>,
6860) {
6861    const TR: usize = 2048;
6862    thread_local! {
6863        static WTILE: std::cell::RefCell<Vec<f32>> = const { std::cell::RefCell::new(Vec::new()) };
6864    }
6865    WTILE.with(|wt| {
6866        let mut wtile = wt.borrow_mut();
6867        wtile.resize(TR * cols, 0.0);
6868        let mut r0 = 0usize;
6869        while r0 < rows {
6870            let tr = TR.min(rows - r0);
6871            let wt_addr = SendMut(wtile.as_mut_ptr());
6872            let run = |start: usize, end: usize| {
6873                for r in start..end {
6874                    // SAFETY: workers cover disjoint r ranges.
6875                    let dst = unsafe { std::slice::from_raw_parts_mut(wt_addr.at(r * cols), cols) };
6876                    dequant_row(r0 + r, dst);
6877                }
6878            };
6879            dispatch_rows(pool, tr, &run);
6880            unsafe {
6881                accel_blas::cblas_sgemm(
6882                    101, // RowMajor
6883                    111, // NoTrans A
6884                    112, // Trans B
6885                    b as i32,
6886                    tr as i32,
6887                    cols as i32,
6888                    1.0,
6889                    xs_all.as_ptr(),
6890                    cols as i32,
6891                    wtile.as_ptr(),
6892                    cols as i32,
6893                    0.0,
6894                    out.as_mut_ptr().add(r0),
6895                    rows as i32,
6896                );
6897            }
6898            r0 += tr;
6899        }
6900    });
6901}
6902
6903fn q4t_matmat(
6904    bytes: &[u8],
6905    xs_all: &[f32],
6906    b: usize,
6907    rows: usize,
6908    cols: usize,
6909    out: &mut [f32],
6910    pool: Option<&Pool>,
6911) {
6912    debug_assert_eq!(out.len(), b * rows);
6913    let gpr = cols / GROUP_SIZE;
6914    // Wide batches ride the AMX like q8's qmatmat: on Apple silicon
6915    // the dequant-tile sgemm is an order above the SDOT row loop for
6916    // prefill shapes (imagegen DiT forwards are exactly this).
6917    #[cfg(target_os = "macos")]
6918    if b >= 8 && rows * cols >= 500_000 && accel_gemm_enabled() {
6919        dequant_matmat_accel(
6920            &|r, dst| {
6921                for gi in 0..gpr {
6922                    let tile = &bytes[(r * gpr + gi) * Q4_TILE..(r * gpr + gi + 1) * Q4_TILE];
6923                    let s = f16_to_f32(u16::from_le_bytes([tile[0], tile[1]]));
6924                    for (k, &bb) in tile[2..].iter().enumerate() {
6925                        dst[gi * GROUP_SIZE + k * 2] = ((bb & 0x0F) as f32 - 8.0) * s;
6926                        dst[gi * GROUP_SIZE + k * 2 + 1] = (((bb >> 4) & 0x0F) as f32 - 8.0) * s;
6927                    }
6928                }
6929            },
6930            xs_all,
6931            b,
6932            rows,
6933            cols,
6934            out,
6935            pool,
6936        );
6937        return;
6938    }
6939    let out_addr = SendMut(out.as_mut_ptr());
6940    if a8w8_enabled() {
6941        let acts: Vec<SplitAct> = (0..b)
6942            .map(|bi| split_act(&xs_all[bi * cols..(bi + 1) * cols]))
6943            .collect();
6944        let acts = &acts;
6945        #[cfg(target_arch = "x86_64")]
6946        let blocked_ok = avx2_enabled() && blocked_enabled();
6947        #[cfg(target_arch = "aarch64")]
6948        let blocked_ok = sdot_enabled() && blocked_enabled();
6949        #[cfg(not(any(target_arch = "x86_64", target_arch = "aarch64")))]
6950        let blocked_ok = false;
6951        let run = move |start: usize, end: usize| {
6952            for r in start..end {
6953                let mut bi = 0usize;
6954                #[cfg(target_arch = "aarch64")]
6955                if blocked_ok {
6956                    while bi + 4 <= acts.len() {
6957                        let xs = [
6958                            acts[bi].xq.as_slice(),
6959                            acts[bi + 1].xq.as_slice(),
6960                            acts[bi + 2].xq.as_slice(),
6961                            acts[bi + 3].xq.as_slice(),
6962                        ];
6963                        let d = unsafe { dot_q4t_row_1x4_sdot(bytes, r, gpr, xs) };
6964                        for k in 0..4 {
6965                            let act = &acts[bi + k];
6966                            let mut acc = d[k] * act.sx;
6967                            for &(j, xv) in &act.outliers {
6968                                let (w, sc) = q4t_outlier(bytes, r, gpr, j);
6969                                acc += w * sc * xv;
6970                            }
6971                            // SAFETY: disjoint (bi, r) cells per worker.
6972                            unsafe { *out_addr.at((bi + k) * rows + r) = acc };
6973                        }
6974                        bi += 4;
6975                    }
6976                }
6977                #[cfg(target_arch = "x86_64")]
6978                if blocked_ok {
6979                    while bi + 4 <= acts.len() {
6980                        let xs = [
6981                            acts[bi].xq.as_slice(),
6982                            acts[bi + 1].xq.as_slice(),
6983                            acts[bi + 2].xq.as_slice(),
6984                            acts[bi + 3].xq.as_slice(),
6985                        ];
6986                        let d = unsafe {
6987                            if vnni_tiles_enabled() {
6988                                dot_q4t_row_1x4_vnni(bytes, r, gpr, xs)
6989                            } else {
6990                                dot_q4t_row_1x4_avx2(bytes, r, gpr, xs)
6991                            }
6992                        };
6993                        for k in 0..4 {
6994                            let act = &acts[bi + k];
6995                            let mut acc = d[k] * act.sx;
6996                            for &(j, xv) in &act.outliers {
6997                                let (w, sc) = q4t_outlier(bytes, r, gpr, j);
6998                                acc += w * sc * xv;
6999                            }
7000                            // SAFETY: disjoint (bi, r) cells per worker.
7001                            unsafe { *out_addr.at((bi + k) * rows + r) = acc };
7002                        }
7003                        bi += 4;
7004                    }
7005                }
7006                let _ = blocked_ok;
7007                while bi < acts.len() {
7008                    let act = &acts[bi];
7009                    let mut acc = dot_q4t_row_i8(bytes, r, gpr, &act.xq) * act.sx;
7010                    for &(j, xv) in &act.outliers {
7011                        let (w, s) = q4t_outlier(bytes, r, gpr, j);
7012                        acc += w * s * xv;
7013                    }
7014                    // SAFETY: disjoint (bi, r) cells per worker range.
7015                    unsafe { *out_addr.at(bi * rows + r) = acc };
7016                    bi += 1;
7017                }
7018            }
7019        };
7020        dispatch_rows(pool, rows, &run);
7021        return;
7022    }
7023    let run = move |start: usize, end: usize| {
7024        for r in start..end {
7025            for bi in 0..b {
7026                let x = &xs_all[bi * cols..(bi + 1) * cols];
7027                // SAFETY: disjoint (bi, r) cells per worker range.
7028                unsafe { *out_addr.at(bi * rows + r) = q4t_row_exact(bytes, r, gpr, x) };
7029            }
7030        }
7031    };
7032    dispatch_rows(pool, rows, &run);
7033}
7034
7035// ── q1 (dtype 12): binary weights, [f16 scale][4B sign bits] per
7036// 32-group tile. The kernel family mirrors q4_tiled: one sequential
7037// stream of 6-byte tiles, per-tile integer dot × scale, exact outlier
7038// correction (A8W8 contract), exact scalar path under CMF_SDOT=0. ──
7039
7040/// Per-32-group sums of the quantized activation — the ±1 identity's
7041/// shared half: `dot = −2·sdot(mask, x) − gsum[g]`, computed ONCE per
7042/// matvec and reused by every row.
7043fn q1_group_sums(xq: &[i8], gpr: usize) -> Vec<i32> {
7044    (0..gpr)
7045        .map(|gi| {
7046            xq[gi * GROUP_SIZE..(gi + 1) * GROUP_SIZE]
7047                .iter()
7048                .map(|&v| v as i32)
7049                .sum()
7050        })
7051        .collect()
7052}
7053
7054/// One q1 row via the A8W8 int8 path — mask-SDOT on ARM (no ±1
7055/// expansion at all), scalar bit loop elsewhere (AVX2 queued with the
7056/// x86 pass).
7057#[inline]
7058#[allow(unreachable_code)]
7059/// AVX2 q1 row via the same ±1 identity as the ARM sdot kernel: the
7060/// sign bits expand to a {0, −1} byte mask through shuffle+cmpeq, the
7061/// masked activation sums through maddubs(1, x&mask), and
7062/// `dot = −(2·masked_sum + Σx_group)` — bit-identical integer math.
7063#[cfg(target_arch = "x86_64")]
7064#[target_feature(enable = "avx2")]
7065unsafe fn dot_q1_row_avx2(bytes: &[u8], r: usize, gpr: usize, xq: &[i8], gsum: &[i32]) -> f32 {
7066    // SAFETY: callers uphold the 6B-tile and xq/gsum length contracts.
7067    unsafe {
7068        use core::arch::x86_64::*;
7069        // Byte j of the mask must replicate bits-byte j/8.
7070        let expand = _mm256_setr_epi8(
7071            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,
7072            3, 3, 3,
7073        );
7074        let bitsel = _mm256_setr_epi8(
7075            1, 2, 4, 8, 16, 32, 64, -128, 1, 2, 4, 8, 16, 32, 64, -128, 1, 2, 4, 8, 16, 32, 64,
7076            -128, 1, 2, 4, 8, 16, 32, 64, -128,
7077        );
7078        let ones8 = _mm256_set1_epi8(1);
7079        let ones16 = _mm256_set1_epi16(1);
7080        let mut acc = 0f32;
7081        for gi in 0..gpr {
7082            let t = bytes.as_ptr().add((r * gpr + gi) * Q1_TILE);
7083            let s = f16_to_f32(u16::from_le_bytes([*t, *t.add(1)]));
7084            let bits = u32::from_le_bytes([*t.add(2), *t.add(3), *t.add(4), *t.add(5)]);
7085            let bc = _mm256_shuffle_epi8(_mm256_set1_epi32(bits as i32), expand);
7086            let mask = _mm256_cmpeq_epi8(_mm256_and_si256(bc, bitsel), bitsel);
7087            let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
7088            let sel = _mm256_and_si256(x, mask);
7089            // Σ of selected i8 lanes: maddubs(1u8, sel_i8) pairs → madd.
7090            let p16 = _mm256_maddubs_epi16(ones8, sel);
7091            let d32 = _mm256_madd_epi16(p16, ones16);
7092            let hi128 = _mm256_extracti128_si256::<1>(d32);
7093            let s128 = _mm_add_epi32(_mm256_castsi256_si128(d32), hi128);
7094            let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
7095            let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
7096            let msum = _mm_cvtsi128_si32(s32);
7097            // The and-select keeps x UN-negated (unlike ARM's −1-mask
7098            // sdot): d = Σ_set − Σ_unset = 2·Σ_set − Σ_all.
7099            let d = 2 * msum - gsum[gi];
7100            acc += d as f32 * s;
7101        }
7102        acc
7103    }
7104}
7105
7106/// VNNI twin of `dot_q1_row_avx2`: the masked-select sum goes through
7107/// one `vpdpbusd(1u8, sel)` (see `dpbusd_hsum` — bit-identical).
7108#[cfg(target_arch = "x86_64")]
7109#[target_feature(enable = "avx2,avx512f,avx512bw,avx512vl,avx512vnni")]
7110unsafe fn dot_q1_row_vnni(bytes: &[u8], r: usize, gpr: usize, xq: &[i8], gsum: &[i32]) -> f32 {
7111    // SAFETY: callers uphold the 6B-tile and xq/gsum length contracts.
7112    unsafe {
7113        use core::arch::x86_64::*;
7114        let expand = _mm256_setr_epi8(
7115            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,
7116            3, 3, 3,
7117        );
7118        let bitsel = _mm256_setr_epi8(
7119            1, 2, 4, 8, 16, 32, 64, -128, 1, 2, 4, 8, 16, 32, 64, -128, 1, 2, 4, 8, 16, 32, 64,
7120            -128, 1, 2, 4, 8, 16, 32, 64, -128,
7121        );
7122        let ones8 = _mm256_set1_epi8(1);
7123        let mut acc = 0f32;
7124        for gi in 0..gpr {
7125            let t = bytes.as_ptr().add((r * gpr + gi) * Q1_TILE);
7126            let s = f16_to_f32(u16::from_le_bytes([*t, *t.add(1)]));
7127            let bits = u32::from_le_bytes([*t.add(2), *t.add(3), *t.add(4), *t.add(5)]);
7128            let bc = _mm256_shuffle_epi8(_mm256_set1_epi32(bits as i32), expand);
7129            let mask = _mm256_cmpeq_epi8(_mm256_and_si256(bc, bitsel), bitsel);
7130            let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
7131            let msum = dpbusd_hsum(ones8, _mm256_and_si256(x, mask));
7132            let d = 2 * msum - gsum[gi];
7133            acc += d as f32 * s;
7134        }
7135        acc
7136    }
7137}
7138
7139/// VNNI twin of `dot_q1_row_1x4_avx2` (see `dpbusd_hsum`).
7140#[cfg(target_arch = "x86_64")]
7141#[target_feature(enable = "avx2,avx512f,avx512bw,avx512vl,avx512vnni")]
7142unsafe fn dot_q1_row_1x4_vnni(
7143    bytes: &[u8],
7144    r: usize,
7145    gpr: usize,
7146    xs: [&[i8]; 4],
7147    gsums: [&[i32]; 4],
7148) -> [f32; 4] {
7149    // SAFETY: callers uphold the 6B-tile and xq/gsum length contracts.
7150    unsafe {
7151        use core::arch::x86_64::*;
7152        let expand = _mm256_setr_epi8(
7153            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,
7154            3, 3, 3,
7155        );
7156        let bitsel = _mm256_setr_epi8(
7157            1, 2, 4, 8, 16, 32, 64, -128, 1, 2, 4, 8, 16, 32, 64, -128, 1, 2, 4, 8, 16, 32, 64,
7158            -128, 1, 2, 4, 8, 16, 32, 64, -128,
7159        );
7160        let ones8 = _mm256_set1_epi8(1);
7161        let mut acc = [0f32; 4];
7162        for gi in 0..gpr {
7163            let t = bytes.as_ptr().add((r * gpr + gi) * Q1_TILE);
7164            let s = f16_to_f32(u16::from_le_bytes([*t, *t.add(1)]));
7165            let bits = u32::from_le_bytes([*t.add(2), *t.add(3), *t.add(4), *t.add(5)]);
7166            let bc = _mm256_shuffle_epi8(_mm256_set1_epi32(bits as i32), expand);
7167            let mask = _mm256_cmpeq_epi8(_mm256_and_si256(bc, bitsel), bitsel);
7168            for (k, xq) in xs.iter().enumerate() {
7169                let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
7170                let msum = dpbusd_hsum(ones8, _mm256_and_si256(x, mask));
7171                let d = 2 * msum - gsums[k][gi];
7172                acc[k] += d as f32 * s;
7173            }
7174        }
7175        acc
7176    }
7177}
7178
7179/// The blocked 1×4 flavor: the expanded bit mask serves four activation
7180/// streams per group (mask build once, four select+reduce chains).
7181#[cfg(target_arch = "x86_64")]
7182#[target_feature(enable = "avx2")]
7183unsafe fn dot_q1_row_1x4_avx2(
7184    bytes: &[u8],
7185    r: usize,
7186    gpr: usize,
7187    xs: [&[i8]; 4],
7188    gsums: [&[i32]; 4],
7189) -> [f32; 4] {
7190    // SAFETY: callers uphold the 6B-tile and xq/gsum length contracts.
7191    unsafe {
7192        use core::arch::x86_64::*;
7193        let expand = _mm256_setr_epi8(
7194            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,
7195            3, 3, 3,
7196        );
7197        let bitsel = _mm256_setr_epi8(
7198            1, 2, 4, 8, 16, 32, 64, -128, 1, 2, 4, 8, 16, 32, 64, -128, 1, 2, 4, 8, 16, 32, 64,
7199            -128, 1, 2, 4, 8, 16, 32, 64, -128,
7200        );
7201        let ones8 = _mm256_set1_epi8(1);
7202        let ones16 = _mm256_set1_epi16(1);
7203        let mut acc = [0f32; 4];
7204        for gi in 0..gpr {
7205            let t = bytes.as_ptr().add((r * gpr + gi) * Q1_TILE);
7206            let s = f16_to_f32(u16::from_le_bytes([*t, *t.add(1)]));
7207            let bits = u32::from_le_bytes([*t.add(2), *t.add(3), *t.add(4), *t.add(5)]);
7208            let bc = _mm256_shuffle_epi8(_mm256_set1_epi32(bits as i32), expand);
7209            let mask = _mm256_cmpeq_epi8(_mm256_and_si256(bc, bitsel), bitsel);
7210            for (k, xq) in xs.iter().enumerate() {
7211                let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
7212                let sel = _mm256_and_si256(x, mask);
7213                let p16 = _mm256_maddubs_epi16(ones8, sel);
7214                let d32 = _mm256_madd_epi16(p16, ones16);
7215                let hi128 = _mm256_extracti128_si256::<1>(d32);
7216                let s128 = _mm_add_epi32(_mm256_castsi256_si128(d32), hi128);
7217                let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
7218                let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
7219                let msum = _mm_cvtsi128_si32(s32);
7220                let d = 2 * msum - gsums[k][gi];
7221                acc[k] += d as f32 * s;
7222            }
7223        }
7224        acc
7225    }
7226}
7227
7228#[allow(unreachable_code)]
7229fn dot_q1_row_i8(bytes: &[u8], r: usize, gpr: usize, xq: &[i8], gsum: &[i32]) -> f32 {
7230    #[cfg(target_arch = "aarch64")]
7231    unsafe {
7232        return dot_q1_row_sdot(bytes, r, gpr, xq, gsum);
7233    }
7234    #[cfg(target_arch = "x86_64")]
7235    if avx2_enabled() {
7236        unsafe {
7237            if vnni_tiles_enabled() {
7238                return dot_q1_row_vnni(bytes, r, gpr, xq, gsum);
7239            }
7240            return dot_q1_row_avx2(bytes, r, gpr, xq, gsum);
7241        }
7242    }
7243    let _ = gsum;
7244    let mut acc = 0f32;
7245    for gi in 0..gpr {
7246        let tile = &bytes[(r * gpr + gi) * Q1_TILE..(r * gpr + gi + 1) * Q1_TILE];
7247        let s = f16_to_f32(u16::from_le_bytes([tile[0], tile[1]]));
7248        let mut d = 0i32;
7249        for (j, &b) in tile[2..].iter().enumerate() {
7250            for k in 0..8 {
7251                let w = ((b >> k) & 1) as i32 * 2 - 1;
7252                d += w * xq[gi * GROUP_SIZE + j * 8 + k] as i32;
7253            }
7254        }
7255        acc += d as f32 * s;
7256    }
7257    acc
7258}
7259
7260/// SDOT q1 row via the ±1 identity: the vtst mask (0xFF where the bit
7261/// is set, i.e. −1 as i8) feeds `sdot` DIRECTLY — no expansion to ±1
7262/// lanes at all — and `dot = −(2·sdot(mask, x) + Σx_group)`, with the
7263/// per-group activation sums shared across every row of the matvec.
7264/// Four tiles (128 weights) per iteration: integer dots reduce through
7265/// a vpaddq tree into ONE i32x4 that meets its four scales in a single
7266/// fused f32 multiply-add. Integer math throughout — bit-identical to
7267/// the scalar ±1 reference.
7268#[cfg(target_arch = "aarch64")]
7269#[target_feature(enable = "neon,dotprod")]
7270unsafe fn dot_q1_row_sdot(bytes: &[u8], r: usize, gpr: usize, xq: &[i8], gsum: &[i32]) -> f32 {
7271    // SAFETY: callers uphold slice-length contracts (6B tile per group,
7272    // xq.len() == gpr·GROUP_SIZE, gsum.len() == gpr).
7273    unsafe {
7274        use core::arch::aarch64::*;
7275        use core::arch::asm;
7276        const MASKS: [u8; 16] = [1, 2, 4, 8, 16, 32, 64, 128, 1, 2, 4, 8, 16, 32, 64, 128];
7277        let m = vld1q_u8(MASKS.as_ptr());
7278        // One tile's −Σ_set(x) as an UNREDUCED i32x4 (two mask-sdots).
7279        macro_rules! tile_dot {
7280            ($t:expr, $x:expr) => {{
7281                let v0 = vcombine_u8(vdup_n_u8(*$t.add(2)), vdup_n_u8(*$t.add(3)));
7282                let v1 = vcombine_u8(vdup_n_u8(*$t.add(4)), vdup_n_u8(*$t.add(5)));
7283                let w0 = vreinterpretq_s8_u8(vtstq_u8(v0, m));
7284                let w1 = vreinterpretq_s8_u8(vtstq_u8(v1, m));
7285                let x0 = vld1q_s8($x);
7286                let x1 = vld1q_s8($x.add(16));
7287                let (mut a0, mut a1) = (vdupq_n_s32(0), vdupq_n_s32(0));
7288                asm!(
7289                    "sdot {a0:v}.4s, {w0:v}.16b, {x0:v}.16b",
7290                    "sdot {a1:v}.4s, {w1:v}.16b, {x1:v}.16b",
7291                    a0 = inout(vreg) a0, a1 = inout(vreg) a1,
7292                    w0 = in(vreg) w0, x0 = in(vreg) x0, w1 = in(vreg) w1, x1 = in(vreg) x1,
7293                    options(pure, nomem, nostack),
7294                );
7295                vaddq_s32(a0, a1)
7296            }};
7297        }
7298        // TBL unpack over PAIR loads: one vld1q covers two 6B tiles
7299        // ([s s b b b b][s s b b b b] + 4B slack), TBL replicates each
7300        // bit-byte across 8 lanes for vtst, and the four scales gather
7301        // through tbl2 into one fcvtl — the 16 ld1r broadcast loads and
7302        // 4 branchy software f16 conversions per 128 weights (the
7303        // measured load-port wall of this kernel) become 2 vector
7304        // loads + 9 table lookups. Integer math order is unchanged —
7305        // bit-identical results (FCVTL is exact on every f16).
7306        const IW00: [u8; 16] = [2, 2, 2, 2, 2, 2, 2, 2, 3, 3, 3, 3, 3, 3, 3, 3];
7307        const IW01: [u8; 16] = [4, 4, 4, 4, 4, 4, 4, 4, 5, 5, 5, 5, 5, 5, 5, 5];
7308        const IW10: [u8; 16] = [8, 8, 8, 8, 8, 8, 8, 8, 9, 9, 9, 9, 9, 9, 9, 9];
7309        const IW11: [u8; 16] = [
7310            10, 10, 10, 10, 10, 10, 10, 10, 11, 11, 11, 11, 11, 11, 11, 11,
7311        ];
7312        const ISC: [u8; 8] = [0, 1, 6, 7, 16, 17, 22, 23];
7313        let (iw00, iw01) = (vld1q_u8(IW00.as_ptr()), vld1q_u8(IW01.as_ptr()));
7314        let (iw10, iw11) = (vld1q_u8(IW10.as_ptr()), vld1q_u8(IW11.as_ptr()));
7315        let isc = vld1_u8(ISC.as_ptr());
7316        // One tile's −Σ_set(x) from a TBL-unpacked pair load.
7317        macro_rules! tile_dot_tbl {
7318            ($ld:expr, $i0:expr, $i1:expr, $x:expr) => {{
7319                let w0 = vreinterpretq_s8_u8(vtstq_u8(vqtbl1q_u8($ld, $i0), m));
7320                let w1 = vreinterpretq_s8_u8(vtstq_u8(vqtbl1q_u8($ld, $i1), m));
7321                let x0 = vld1q_s8($x);
7322                let x1 = vld1q_s8($x.add(16));
7323                let (mut a0, mut a1) = (vdupq_n_s32(0), vdupq_n_s32(0));
7324                asm!(
7325                    "sdot {a0:v}.4s, {w0:v}.16b, {x0:v}.16b",
7326                    "sdot {a1:v}.4s, {w1:v}.16b, {x1:v}.16b",
7327                    a0 = inout(vreg) a0, a1 = inout(vreg) a1,
7328                    w0 = in(vreg) w0, x0 = in(vreg) x0, w1 = in(vreg) w1, x1 = in(vreg) x1,
7329                    options(pure, nomem, nostack),
7330                );
7331                vaddq_s32(a0, a1)
7332            }};
7333        }
7334        let base = bytes.as_ptr().add(r * gpr * Q1_TILE);
7335        let row_base = r * gpr * Q1_TILE;
7336        let abs_end = bytes.len();
7337        let xp = xq.as_ptr();
7338        let gp = gsum.as_ptr();
7339        let mut accv = vdupq_n_f32(0.0);
7340        let mut gi = 0;
7341        // The second pair load reads 4B past tile gi+3 — stay inside
7342        // the payload slice (only the file's final tiles fall back).
7343        while gi + 4 <= gpr && row_base + (gi + 4) * Q1_TILE + 4 <= abs_end {
7344            let t0 = base.add(gi * Q1_TILE);
7345            let ld_a = vld1q_u8(t0);
7346            let ld_b = vld1q_u8(t0.add(2 * Q1_TILE));
7347            let d0 = tile_dot_tbl!(ld_a, iw00, iw01, xp.add(gi * GROUP_SIZE));
7348            let d1 = tile_dot_tbl!(ld_a, iw10, iw11, xp.add((gi + 1) * GROUP_SIZE));
7349            let d2 = tile_dot_tbl!(ld_b, iw00, iw01, xp.add((gi + 2) * GROUP_SIZE));
7350            let d3 = tile_dot_tbl!(ld_b, iw10, iw11, xp.add((gi + 3) * GROUP_SIZE));
7351            // [−Σ0, −Σ1, −Σ2, −Σ3] → dots = −(2·Σset_neg + gsum)
7352            let neg = vpaddq_s32(vpaddq_s32(d0, d1), vpaddq_s32(d2, d3));
7353            let g = vld1q_s32(gp.add(gi));
7354            let dots = vnegq_s32(vaddq_s32(vshlq_n_s32::<1>(neg), g));
7355            let sc16 = vqtbl2_u8(uint8x16x2_t(ld_a, ld_b), isc);
7356            let scf: float32x4_t;
7357            asm!(
7358                "fcvtl {o:v}.4s, {i:v}.4h",
7359                o = out(vreg) scf, i = in(vreg) sc16,
7360                options(pure, nomem, nostack),
7361            );
7362            accv = vfmaq_f32(accv, vcvtq_f32_s32(dots), scf);
7363            gi += 4;
7364        }
7365        let mut acc = vaddvq_f32(accv);
7366        while gi < gpr {
7367            let t = base.add(gi * Q1_TILE);
7368            let s = f16_to_f32(u16::from_le_bytes([*t, *t.add(1)]));
7369            let d = vaddvq_s32(tile_dot!(t, xp.add(gi * GROUP_SIZE)));
7370            acc += (-(2 * d + *gp.add(gi))) as f32 * s;
7371            gi += 1;
7372        }
7373        acc
7374    }
7375}
7376
7377/// Blocked q1 1×4: one TBL unpack of the tile pair serves FOUR
7378/// activation streams (prefill amortization — the same idea as the
7379/// AVX2 twin; per stream the group order, fma order and tail match the
7380/// single-row kernel exactly, so batch == matvec bit-for-bit).
7381#[cfg(target_arch = "aarch64")]
7382#[target_feature(enable = "neon,dotprod")]
7383unsafe fn dot_q1_row_1x4_sdot(
7384    bytes: &[u8],
7385    r: usize,
7386    gpr: usize,
7387    xs: [&[i8]; 4],
7388    gs: [&[i32]; 4],
7389) -> [f32; 4] {
7390    // SAFETY: same slice-length contracts as `dot_q1_row_sdot`, ×4.
7391    unsafe {
7392        use core::arch::aarch64::*;
7393        use core::arch::asm;
7394        const MASKS: [u8; 16] = [1, 2, 4, 8, 16, 32, 64, 128, 1, 2, 4, 8, 16, 32, 64, 128];
7395        const IW00: [u8; 16] = [2, 2, 2, 2, 2, 2, 2, 2, 3, 3, 3, 3, 3, 3, 3, 3];
7396        const IW01: [u8; 16] = [4, 4, 4, 4, 4, 4, 4, 4, 5, 5, 5, 5, 5, 5, 5, 5];
7397        const IW10: [u8; 16] = [8, 8, 8, 8, 8, 8, 8, 8, 9, 9, 9, 9, 9, 9, 9, 9];
7398        const IW11: [u8; 16] = [
7399            10, 10, 10, 10, 10, 10, 10, 10, 11, 11, 11, 11, 11, 11, 11, 11,
7400        ];
7401        const ISC: [u8; 8] = [0, 1, 6, 7, 16, 17, 22, 23];
7402        let m = vld1q_u8(MASKS.as_ptr());
7403        let (iw00, iw01) = (vld1q_u8(IW00.as_ptr()), vld1q_u8(IW01.as_ptr()));
7404        let (iw10, iw11) = (vld1q_u8(IW10.as_ptr()), vld1q_u8(IW11.as_ptr()));
7405        let isc = vld1_u8(ISC.as_ptr());
7406        macro_rules! sdot2 {
7407            ($w0:expr, $w1:expr, $x:expr) => {{
7408                let x0 = vld1q_s8($x);
7409                let x1 = vld1q_s8($x.add(16));
7410                let (mut a0, mut a1) = (vdupq_n_s32(0), vdupq_n_s32(0));
7411                asm!(
7412                    "sdot {a0:v}.4s, {w0:v}.16b, {x0:v}.16b",
7413                    "sdot {a1:v}.4s, {w1:v}.16b, {x1:v}.16b",
7414                    a0 = inout(vreg) a0, a1 = inout(vreg) a1,
7415                    w0 = in(vreg) $w0, x0 = in(vreg) x0, w1 = in(vreg) $w1, x1 = in(vreg) x1,
7416                    options(pure, nomem, nostack),
7417                );
7418                vaddq_s32(a0, a1)
7419            }};
7420        }
7421        let base = bytes.as_ptr().add(r * gpr * Q1_TILE);
7422        let row_base = r * gpr * Q1_TILE;
7423        let abs_end = bytes.len();
7424        let mut accv = [vdupq_n_f32(0.0); 4];
7425        let mut gi = 0;
7426        while gi + 4 <= gpr && row_base + (gi + 4) * Q1_TILE + 4 <= abs_end {
7427            let t0 = base.add(gi * Q1_TILE);
7428            let ld_a = vld1q_u8(t0);
7429            let ld_b = vld1q_u8(t0.add(2 * Q1_TILE));
7430            // Unpack ONCE — eight ±mask vectors serve all four streams.
7431            let w00 = vreinterpretq_s8_u8(vtstq_u8(vqtbl1q_u8(ld_a, iw00), m));
7432            let w01 = vreinterpretq_s8_u8(vtstq_u8(vqtbl1q_u8(ld_a, iw01), m));
7433            let w10 = vreinterpretq_s8_u8(vtstq_u8(vqtbl1q_u8(ld_a, iw10), m));
7434            let w11 = vreinterpretq_s8_u8(vtstq_u8(vqtbl1q_u8(ld_a, iw11), m));
7435            let w20 = vreinterpretq_s8_u8(vtstq_u8(vqtbl1q_u8(ld_b, iw00), m));
7436            let w21 = vreinterpretq_s8_u8(vtstq_u8(vqtbl1q_u8(ld_b, iw01), m));
7437            let w30 = vreinterpretq_s8_u8(vtstq_u8(vqtbl1q_u8(ld_b, iw10), m));
7438            let w31 = vreinterpretq_s8_u8(vtstq_u8(vqtbl1q_u8(ld_b, iw11), m));
7439            let sc16 = vqtbl2_u8(uint8x16x2_t(ld_a, ld_b), isc);
7440            let scf: float32x4_t;
7441            asm!(
7442                "fcvtl {o:v}.4s, {i:v}.4h",
7443                o = out(vreg) scf, i = in(vreg) sc16,
7444                options(pure, nomem, nostack),
7445            );
7446            for k in 0..4 {
7447                let xp = xs[k].as_ptr();
7448                let d0 = sdot2!(w00, w01, xp.add(gi * GROUP_SIZE));
7449                let d1 = sdot2!(w10, w11, xp.add((gi + 1) * GROUP_SIZE));
7450                let d2 = sdot2!(w20, w21, xp.add((gi + 2) * GROUP_SIZE));
7451                let d3 = sdot2!(w30, w31, xp.add((gi + 3) * GROUP_SIZE));
7452                let neg = vpaddq_s32(vpaddq_s32(d0, d1), vpaddq_s32(d2, d3));
7453                let g = vld1q_s32(gs[k].as_ptr().add(gi));
7454                let dots = vnegq_s32(vaddq_s32(vshlq_n_s32::<1>(neg), g));
7455                accv[k] = vfmaq_f32(accv[k], vcvtq_f32_s32(dots), scf);
7456            }
7457            gi += 4;
7458        }
7459        let mut acc = [
7460            vaddvq_f32(accv[0]),
7461            vaddvq_f32(accv[1]),
7462            vaddvq_f32(accv[2]),
7463            vaddvq_f32(accv[3]),
7464        ];
7465        while gi < gpr {
7466            let t = base.add(gi * Q1_TILE);
7467            let sc = f16_to_f32(u16::from_le_bytes([*t, *t.add(1)]));
7468            let v0 = vcombine_u8(vdup_n_u8(*t.add(2)), vdup_n_u8(*t.add(3)));
7469            let v1 = vcombine_u8(vdup_n_u8(*t.add(4)), vdup_n_u8(*t.add(5)));
7470            let w0 = vreinterpretq_s8_u8(vtstq_u8(v0, m));
7471            let w1 = vreinterpretq_s8_u8(vtstq_u8(v1, m));
7472            for k in 0..4 {
7473                let d = vaddvq_s32(sdot2!(w0, w1, xs[k].as_ptr().add(gi * GROUP_SIZE)));
7474                acc[k] += (-(2 * d + *gs[k].as_ptr().add(gi))) as f32 * sc;
7475            }
7476            gi += 1;
7477        }
7478        acc
7479    }
7480}
7481
7482/// (weight ±1, scale) of one q1 element — the exact outlier term.
7483#[inline]
7484fn q1_outlier(bytes: &[u8], r: usize, gpr: usize, j: usize) -> (f32, f32) {
7485    let gi = j / GROUP_SIZE;
7486    let k = j % GROUP_SIZE;
7487    let tile = &bytes[(r * gpr + gi) * Q1_TILE..(r * gpr + gi + 1) * Q1_TILE];
7488    let s = f16_to_f32(u16::from_le_bytes([tile[0], tile[1]]));
7489    let bit = (tile[2 + k / 8] >> (k % 8)) & 1;
7490    ((bit as i32 * 2 - 1) as f32, s)
7491}
7492
7493/// Exact scalar q1 row (CMF_SDOT=0 contract).
7494#[inline]
7495fn q1_row_exact(bytes: &[u8], r: usize, gpr: usize, x: &[f32]) -> f32 {
7496    let mut acc = 0f32;
7497    for gi in 0..gpr {
7498        let tile = &bytes[(r * gpr + gi) * Q1_TILE..(r * gpr + gi + 1) * Q1_TILE];
7499        let s = f16_to_f32(u16::from_le_bytes([tile[0], tile[1]]));
7500        let xg = &x[gi * GROUP_SIZE..(gi + 1) * GROUP_SIZE];
7501        let mut ga = 0f32;
7502        for (j, &b) in tile[2..].iter().enumerate() {
7503            for k in 0..8 {
7504                ga += (((b >> k) & 1) as f32 * 2.0 - 1.0) * xg[j * 8 + k];
7505            }
7506        }
7507        acc += ga * s;
7508    }
7509    acc
7510}
7511
7512/// One q1 row range via A8W8 (the body of `q1_matvec`'s hot loop,
7513/// extracted so multi-matrix jobs drive the same kernel).
7514#[allow(clippy::too_many_arguments)]
7515fn q1_range_a8w8(
7516    bytes: &[u8],
7517    gpr: usize,
7518    act: &SplitAct,
7519    gsum: &[i32],
7520    out: SendMut,
7521    start: usize,
7522    end: usize,
7523) {
7524    for r in start..end {
7525        let mut acc = dot_q1_row_i8(bytes, r, gpr, &act.xq, gsum) * act.sx;
7526        for &(j, xv) in &act.outliers {
7527            let (w, s) = q1_outlier(bytes, r, gpr, j);
7528            acc += w * s * xv;
7529        }
7530        // SAFETY: disjoint row ranges per worker.
7531        unsafe { *out.at(r) = acc };
7532    }
7533}
7534
7535/// Exact-scalar q1 row range (CMF_SDOT=0 contract).
7536fn q1_range_f32(bytes: &[u8], gpr: usize, x: &[f32], out: SendMut, start: usize, end: usize) {
7537    for r in start..end {
7538        // SAFETY: disjoint row ranges per worker.
7539        unsafe { *out.at(r) = q1_row_exact(bytes, r, gpr, x) };
7540    }
7541}
7542
7543/// q1t per-row overlay locator. After the base (`base_len`) come
7544/// `[u32 row_ptr[rows+1]]` then `[(u16 col, f16 val)]` grouped by row (row
7545/// `r`'s entries are `[row_ptr[r], row_ptr[r+1])`). Returns
7546/// `(row_ptr offset, entries offset, present)`.
7547fn q1t_overlay(bytes: &[u8], base_len: usize, rows: usize) -> (usize, usize, bool) {
7548    let entries = base_len + (rows + 1) * 4;
7549    (base_len, entries, entries <= bytes.len())
7550}
7551
7552/// Read `row_ptr[r]` from the overlay's prefix-sum table.
7553#[inline]
7554fn q1t_rowptr(bytes: &[u8], rp_off: usize, r: usize) -> usize {
7555    let o = rp_off + r * 4;
7556    u32::from_le_bytes([bytes[o], bytes[o + 1], bytes[o + 2], bytes[o + 3]]) as usize
7557}
7558
7559/// Byte → the 5 ternary signs it packs `{−1,0,+1}` as f32, precomputed so
7560/// decoding a q1t code is a table load, not the base-3 divide/modulo per
7561/// weight (division is ~20–40× the cost of a load). Built at compile time.
7562const SIGN5: [[f32; 5]; 256] = {
7563    let mut lut = [[0.0f32; 5]; 256];
7564    let pow3 = [1u16, 3, 9, 27, 81];
7565    let mut byte = 0usize;
7566    while byte < 256 {
7567        let mut i = 0usize;
7568        while i < 5 {
7569            let code = (byte as u16 / pow3[i]) % 3;
7570            lut[byte][i] = if code == 1 {
7571                1.0
7572            } else if code == 2 {
7573                -1.0
7574            } else {
7575                0.0
7576            };
7577            i += 1;
7578        }
7579        byte += 1;
7580    }
7581    lut
7582};
7583
7584/// Same table, as i8 signs — the operand for the int8 SDOT base kernel.
7585const SIGN5_I8: [[i8; 5]; 256] = {
7586    let mut lut = [[0i8; 5]; 256];
7587    let pow3 = [1u16, 3, 9, 27, 81];
7588    let mut byte = 0usize;
7589    while byte < 256 {
7590        let mut i = 0usize;
7591        while i < 5 {
7592            let code = (byte as u16 / pow3[i]) % 3;
7593            lut[byte][i] = if code == 1 {
7594                1
7595            } else if code == 2 {
7596                -1
7597            } else {
7598                0
7599            };
7600            i += 1;
7601        }
7602        byte += 1;
7603    }
7604    lut
7605};
7606
7607/// The same 5 i8 signs packed into a u64 (`[s0 s1 s2 s3 s4 0 0 0]`, LE) so the
7608/// group unpack is 7 unaligned u64 stores at offsets 0,5,10,…,30 instead of
7609/// six 5-byte copies + LUT indexing — each store's trailing zeros are fixed by
7610/// the next store, and the last one runs 6 B past the 32nd weight (the unpack
7611/// buffer is padded to 40). This is the decode/prefill hot inner op.
7612const SIGN5_U64: [u64; 256] = {
7613    let mut lut = [0u64; 256];
7614    let pow3 = [1u16, 3, 9, 27, 81];
7615    let mut byte = 0usize;
7616    while byte < 256 {
7617        let mut v = 0u64;
7618        let mut i = 0usize;
7619        while i < 5 {
7620            let code = (byte as u16 / pow3[i]) % 3;
7621            let s: u8 = if code == 1 {
7622                1
7623            } else if code == 2 {
7624                0xFF
7625            } else {
7626                0
7627            };
7628            v |= (s as u64) << (i * 8);
7629            i += 1;
7630        }
7631        lut[byte] = v;
7632        byte += 1;
7633    }
7634    lut
7635};
7636
7637/// Ternary base weight at `(row r, col j)` = `sign(code)·s_group`. Used to add
7638/// back activation-outlier columns, whose `x` was zeroed for the int8 bulk dot
7639/// (`split_act`). At a weight-outlier position the code is 0, so this is 0 and
7640/// the overlay correction owns that column — no double counting.
7641#[inline]
7642fn q1t_base_weight(bytes: &[u8], r: usize, gpr: usize, j: usize) -> f32 {
7643    const TILE: usize = cortiq_core::quant::Q1T_TILE;
7644    let off = (r * gpr + j / GROUP_SIZE) * TILE;
7645    let s = f16_to_f32(u16::from_le_bytes([bytes[off], bytes[off + 1]]));
7646    let within = j % GROUP_SIZE;
7647    SIGN5[bytes[off + 2 + within / 5] as usize][within % 5] * s
7648}
7649
7650/// One 32-group int8 dot via two SDOTs. Bit-exact vs the scalar i8 sum
7651/// (integer accumulation is order-independent).
7652#[cfg(target_arch = "aarch64")]
7653#[target_feature(enable = "neon,dotprod")]
7654#[inline]
7655unsafe fn sdot32_i8(w: *const i8, x: *const i8) -> i32 {
7656    // SAFETY: caller guarantees 32 readable i8 at each pointer.
7657    unsafe {
7658        use core::arch::aarch64::*;
7659        use core::arch::asm;
7660        let w0 = vld1q_s8(w);
7661        let w1 = vld1q_s8(w.add(16));
7662        let x0 = vld1q_s8(x);
7663        let x1 = vld1q_s8(x.add(16));
7664        let (mut a0, mut a1) = (vdupq_n_s32(0), vdupq_n_s32(0));
7665        asm!(
7666            "sdot {a0:v}.4s, {w0:v}.16b, {x0:v}.16b",
7667            "sdot {a1:v}.4s, {w1:v}.16b, {x1:v}.16b",
7668            a0 = inout(vreg) a0, a1 = inout(vreg) a1,
7669            w0 = in(vreg) w0, x0 = in(vreg) x0, w1 = in(vreg) w1, x1 = in(vreg) x1,
7670            options(pure, nomem, nostack),
7671        );
7672        vaddvq_s32(vaddq_s32(a0, a1))
7673    }
7674}
7675
7676/// One 32-group int8 dot via AVX2: signed·signed as `maddubs(|w|, sign(x,w))`
7677/// then `madd` and a horizontal reduce (the same idiom as `dot_q4t_row_avx2`).
7678#[cfg(target_arch = "x86_64")]
7679#[target_feature(enable = "avx2")]
7680#[inline]
7681unsafe fn i8dot32_avx2(w: *const i8, x: *const i8) -> i32 {
7682    // SAFETY: caller guarantees 32 readable i8 at each pointer.
7683    unsafe {
7684        use core::arch::x86_64::*;
7685        let wv = _mm256_loadu_si256(w as *const __m256i);
7686        let xv = _mm256_loadu_si256(x as *const __m256i);
7687        let p16 = _mm256_maddubs_epi16(_mm256_abs_epi8(wv), _mm256_sign_epi8(xv, wv));
7688        let d = _mm256_madd_epi16(p16, _mm256_set1_epi16(1));
7689        let hi128 = _mm256_extracti128_si256::<1>(d);
7690        let s128 = _mm_add_epi32(_mm256_castsi256_si128(d), hi128);
7691        let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
7692        let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
7693        _mm_cvtsi128_si32(s32)
7694    }
7695}
7696
7697/// Unpack one q1t group's base-3 codes into 32 i8 signs via 7 unaligned u64
7698/// stores (see `SIGN5_U64`). `dst` MUST have ≥ 40 bytes: the 7th store writes
7699/// `dst[30..38]`. Stores go in order so each one's trailing zeros are
7700/// overwritten by the next; the final 6 padding bytes are unused by the dot.
7701#[inline]
7702fn q1t_unpack_group_i8(codes: *const u8, dst: &mut [i8]) {
7703    debug_assert!(dst.len() >= 40);
7704    // SAFETY: codes points at 7 readable bytes; dst has ≥ 40 bytes so every
7705    // 8-byte store at offset bi*5 (bi ≤ 6 → ≤ 30) stays in bounds.
7706    unsafe {
7707        let p = dst.as_mut_ptr();
7708        for bi in 0..7 {
7709            core::ptr::write_unaligned(
7710                p.add(bi * 5) as *mut u64,
7711                SIGN5_U64[*codes.add(bi) as usize],
7712            );
7713        }
7714    }
7715}
7716
7717/// One 32-group int8 dot, arch-dispatched (the matmat inner loop, where the
7718/// row's signs are unpacked once and dotted against every batch input).
7719/// Callers are gated by `a8w8_enabled()`, so the target-feature arms are
7720/// reachable; the scalar arm is a non-SIMD-arch fallback.
7721#[inline]
7722fn q1t_i8dot32(w: *const i8, x: *const i8) -> i32 {
7723    #[cfg(target_arch = "aarch64")]
7724    unsafe {
7725        return sdot32_i8(w, x);
7726    }
7727    #[cfg(target_arch = "x86_64")]
7728    unsafe {
7729        return i8dot32_avx2(w, x);
7730    }
7731    #[allow(unreachable_code)]
7732    unsafe {
7733        let mut s = 0i32;
7734        for k in 0..GROUP_SIZE {
7735            s += *w.add(k) as i32 * *x.add(k) as i32;
7736        }
7737        s
7738    }
7739}
7740
7741#[inline]
7742unsafe fn q1t_unpack_reg_u64s(codes: *const u8) -> (u64, u64, u64, u64) {
7743    let (s0, s1, s2, s3, s4, s5, s6) = unsafe {
7744        (
7745            SIGN5_U64[*codes as usize],
7746            SIGN5_U64[*codes.add(1) as usize],
7747            SIGN5_U64[*codes.add(2) as usize],
7748            SIGN5_U64[*codes.add(3) as usize],
7749            SIGN5_U64[*codes.add(4) as usize],
7750            SIGN5_U64[*codes.add(5) as usize],
7751            SIGN5_U64[*codes.add(6) as usize],
7752        )
7753    };
7754
7755    let u0 = s0 | (s1 << 40);
7756    let u1 = (s1 >> 24) | (s2 << 16) | (s3 << 56);
7757    let u2 = (s3 >> 8) | (s4 << 32);
7758    let u3 = (s4 >> 32) | (s5 << 8) | (s6 << 48);
7759
7760    (u0, u1, u2, u3)
7761}
7762
7763/// One q1t row's int8 base dot: `Σ_group s·dot(signs, xq)` (before the shared
7764/// `sx`). Direct register unpacking (zero stack stores/loads, no STLF stalls).
7765/// ARM SDOT.
7766#[cfg(target_arch = "aarch64")]
7767#[target_feature(enable = "neon,dotprod")]
7768unsafe fn q1t_dot_row_sdot(bytes: &[u8], r: usize, gpr: usize, xq: &[i8]) -> f32 {
7769    use core::arch::aarch64::*;
7770    use core::arch::asm;
7771    unsafe {
7772        const TILE: usize = cortiq_core::quant::Q1T_TILE;
7773        let mut acc = 0f32;
7774        let bytes_ptr = bytes.as_ptr();
7775        let xq_ptr = xq.as_ptr();
7776        let row_off = r * gpr * TILE;
7777
7778        let gpr2 = gpr & !1;
7779        let mut gi = 0;
7780        while gi < gpr2 {
7781            let off0 = row_off + gi * TILE;
7782            let off1 = off0 + TILE;
7783            let s0 = f16_to_f32(u16::from_le_bytes([
7784                *bytes_ptr.add(off0),
7785                *bytes_ptr.add(off0 + 1),
7786            ]));
7787            let s1 = f16_to_f32(u16::from_le_bytes([
7788                *bytes_ptr.add(off1),
7789                *bytes_ptr.add(off1 + 1),
7790            ]));
7791
7792            let (u0_0, u1_0, u2_0, u3_0) = q1t_unpack_reg_u64s(bytes_ptr.add(off0 + 2));
7793            let (u0_1, u1_1, u2_1, u3_1) = q1t_unpack_reg_u64s(bytes_ptr.add(off1 + 2));
7794
7795            let w0_0 = vreinterpretq_s8_u64(vcombine_u64(vcreate_u64(u0_0), vcreate_u64(u1_0)));
7796            let w1_0 = vreinterpretq_s8_u64(vcombine_u64(vcreate_u64(u2_0), vcreate_u64(u3_0)));
7797            let w0_1 = vreinterpretq_s8_u64(vcombine_u64(vcreate_u64(u0_1), vcreate_u64(u1_1)));
7798            let w1_1 = vreinterpretq_s8_u64(vcombine_u64(vcreate_u64(u2_1), vcreate_u64(u3_1)));
7799
7800            let x0_0 = vld1q_s8(xq_ptr.add(gi * GROUP_SIZE));
7801            let x1_0 = vld1q_s8(xq_ptr.add(gi * GROUP_SIZE + 16));
7802            let x0_1 = vld1q_s8(xq_ptr.add((gi + 1) * GROUP_SIZE));
7803            let x1_1 = vld1q_s8(xq_ptr.add((gi + 1) * GROUP_SIZE + 16));
7804
7805            let (mut a0_0, mut a1_0) = (vdupq_n_s32(0), vdupq_n_s32(0));
7806            let (mut a0_1, mut a1_1) = (vdupq_n_s32(0), vdupq_n_s32(0));
7807            asm!(
7808                "sdot {a0_0:v}.4s, {w0_0:v}.16b, {x0_0:v}.16b",
7809                "sdot {a1_0:v}.4s, {w1_0:v}.16b, {x1_0:v}.16b",
7810                "sdot {a0_1:v}.4s, {w0_1:v}.16b, {x0_1:v}.16b",
7811                "sdot {a1_1:v}.4s, {w1_1:v}.16b, {x1_1:v}.16b",
7812                a0_0 = inout(vreg) a0_0, a1_0 = inout(vreg) a1_0,
7813                a0_1 = inout(vreg) a0_1, a1_1 = inout(vreg) a1_1,
7814                w0_0 = in(vreg) w0_0, x0_0 = in(vreg) x0_0, w1_0 = in(vreg) w1_0, x1_0 = in(vreg) x1_0,
7815                w0_1 = in(vreg) w0_1, x0_1 = in(vreg) x0_1, w1_1 = in(vreg) w1_1, x1_1 = in(vreg) x1_1,
7816                options(pure, nomem, nostack),
7817            );
7818            let d0 = vaddvq_s32(vaddq_s32(a0_0, a1_0));
7819            let d1 = vaddvq_s32(vaddq_s32(a0_1, a1_1));
7820            acc += d0 as f32 * s0 + d1 as f32 * s1;
7821            gi += 2;
7822        }
7823
7824        if gi < gpr {
7825            let off = row_off + gi * TILE;
7826            let s = f16_to_f32(u16::from_le_bytes([
7827                *bytes_ptr.add(off),
7828                *bytes_ptr.add(off + 1),
7829            ]));
7830            let (u0, u1, u2, u3) = q1t_unpack_reg_u64s(bytes_ptr.add(off + 2));
7831            let w0 = vreinterpretq_s8_u64(vcombine_u64(vcreate_u64(u0), vcreate_u64(u1)));
7832            let w1 = vreinterpretq_s8_u64(vcombine_u64(vcreate_u64(u2), vcreate_u64(u3)));
7833            let x0 = vld1q_s8(xq_ptr.add(gi * GROUP_SIZE));
7834            let x1 = vld1q_s8(xq_ptr.add(gi * GROUP_SIZE + 16));
7835            let (mut a0, mut a1) = (vdupq_n_s32(0), vdupq_n_s32(0));
7836            asm!(
7837                "sdot {a0:v}.4s, {w0:v}.16b, {x0:v}.16b",
7838                "sdot {a1:v}.4s, {w1:v}.16b, {x1:v}.16b",
7839                a0 = inout(vreg) a0, a1 = inout(vreg) a1,
7840                w0 = in(vreg) w0, x0 = in(vreg) x0, w1 = in(vreg) w1, x1 = in(vreg) x1,
7841                options(pure, nomem, nostack),
7842            );
7843            let d = vaddvq_s32(vaddq_s32(a0, a1));
7844            acc += d as f32 * s;
7845        }
7846        acc
7847    }
7848}
7849
7850/// x86 AVX2 mirror of `q1t_dot_row_sdot` (maddubs int8 dot per group).
7851#[cfg(target_arch = "x86_64")]
7852#[target_feature(enable = "avx2")]
7853unsafe fn q1t_dot_row_avx2(bytes: &[u8], r: usize, gpr: usize, xq: &[i8]) -> f32 {
7854    use core::arch::x86_64::*;
7855    unsafe {
7856        const TILE: usize = cortiq_core::quant::Q1T_TILE;
7857        let mut acc = 0f32;
7858        let bytes_ptr = bytes.as_ptr();
7859        let xq_ptr = xq.as_ptr();
7860        let row_off = r * gpr * TILE;
7861
7862        let ones = _mm256_set1_epi16(1);
7863        for gi in 0..gpr {
7864            let off = row_off + gi * TILE;
7865            let s = f16_to_f32(u16::from_le_bytes([
7866                *bytes_ptr.add(off),
7867                *bytes_ptr.add(off + 1),
7868            ]));
7869            let (u0, u1, u2, u3) = q1t_unpack_reg_u64s(bytes_ptr.add(off + 2));
7870            let wv = _mm256_set_epi64x(u3 as i64, u2 as i64, u1 as i64, u0 as i64);
7871            let xv = _mm256_loadu_si256(xq_ptr.add(gi * GROUP_SIZE) as *const __m256i);
7872            let p16 = _mm256_maddubs_epi16(_mm256_abs_epi8(wv), _mm256_sign_epi8(xv, wv));
7873            let d256 = _mm256_madd_epi16(p16, ones);
7874            let d128 = _mm_add_epi32(
7875                _mm256_castsi256_si128(d256),
7876                _mm256_extracti128_si256(d256, 1),
7877            );
7878            let d64 = _mm_add_epi32(d128, _mm_shuffle_epi32(d128, 0xee));
7879            let d32 = _mm_cvtsi128_si32(_mm_add_epi32(d64, _mm_shuffle_epi32(d64, 0x55)));
7880            acc += d32 as f32 * s;
7881        }
7882        acc
7883    }
7884}
7885
7886/// VNNI twin of `q1t_dot_row_avx2` (see `dpbusd_hsum`).
7887#[cfg(target_arch = "x86_64")]
7888#[target_feature(enable = "avx2,avx512f,avx512bw,avx512vl,avx512vnni")]
7889unsafe fn q1t_dot_row_vnni(bytes: &[u8], r: usize, gpr: usize, xq: &[i8]) -> f32 {
7890    use core::arch::x86_64::*;
7891    // SAFETY: same tile/xq contracts as `q1t_dot_row_avx2`.
7892    unsafe {
7893        const TILE: usize = cortiq_core::quant::Q1T_TILE;
7894        let mut acc = 0f32;
7895        let bytes_ptr = bytes.as_ptr();
7896        let xq_ptr = xq.as_ptr();
7897        let row_off = r * gpr * TILE;
7898        for gi in 0..gpr {
7899            let off = row_off + gi * TILE;
7900            let s = f16_to_f32(u16::from_le_bytes([
7901                *bytes_ptr.add(off),
7902                *bytes_ptr.add(off + 1),
7903            ]));
7904            let (u0, u1, u2, u3) = q1t_unpack_reg_u64s(bytes_ptr.add(off + 2));
7905            let wv = _mm256_set_epi64x(u3 as i64, u2 as i64, u1 as i64, u0 as i64);
7906            let xv = _mm256_loadu_si256(xq_ptr.add(gi * GROUP_SIZE) as *const __m256i);
7907            let d = dpbusd_hsum(_mm256_abs_epi8(wv), _mm256_sign_epi8(xv, wv));
7908            acc += d as f32 * s;
7909        }
7910        acc
7911    }
7912}
7913
7914/// Per-row int8 base dot, dispatched once per row (matvec decode hot path).
7915/// Callers are gated by `a8w8_enabled()`, so the target-feature kernels are
7916/// reachable.
7917#[inline]
7918fn q1t_dot_row_i8(bytes: &[u8], r: usize, gpr: usize, xq: &[i8]) -> f32 {
7919    #[cfg(target_arch = "aarch64")]
7920    unsafe {
7921        return q1t_dot_row_sdot(bytes, r, gpr, xq);
7922    }
7923    #[cfg(target_arch = "x86_64")]
7924    unsafe {
7925        if vnni_tiles_enabled() {
7926            return q1t_dot_row_vnni(bytes, r, gpr, xq);
7927        }
7928        return q1t_dot_row_avx2(bytes, r, gpr, xq);
7929    }
7930    #[allow(unreachable_code)]
7931    {
7932        const TILE: usize = cortiq_core::quant::Q1T_TILE;
7933        let mut acc = 0f32;
7934        let mut sg = [0i8; GROUP_SIZE + 8]; // +8 slack for the u64-store unpack
7935        for gi in 0..gpr {
7936            let off = (r * gpr + gi) * TILE;
7937            let s = f16_to_f32(u16::from_le_bytes([bytes[off], bytes[off + 1]]));
7938            q1t_unpack_group_i8(bytes.as_ptr().wrapping_add(off + 2), &mut sg);
7939            let mut d = 0i32;
7940            for k in 0..GROUP_SIZE {
7941                d += sg[k] as i32 * xq[gi * GROUP_SIZE + k] as i32;
7942            }
7943            acc += d as f32 * s;
7944        }
7945        acc
7946    }
7947}
7948
7949/// Σ over a row's outliers of `value·x[col]` — the correction that adds the
7950/// overlay's exact weights on top of the base dot. INVARIANT: the encoder
7951/// writes ternary code 0 at every outlier position (`quantize_q1t`), so the
7952/// base contributes nothing there and this is a plain `value·x`, not
7953/// `(value − base)·x` — no scattered per-outlier scale read. Row `r`'s entries
7954/// are the contiguous slice `[row_ptr[r], row_ptr[r+1])`, so no binary search.
7955fn q1t_row_outlier_correction(
7956    bytes: &[u8],
7957    r: usize,
7958    rp_off: usize,
7959    entries_off: usize,
7960    has_ov: bool,
7961    x: &[f32],
7962) -> f32 {
7963    if !has_ov {
7964        return 0.0;
7965    }
7966    let (c0, c1) = (
7967        q1t_rowptr(bytes, rp_off, r),
7968        q1t_rowptr(bytes, rp_off, r + 1),
7969    );
7970    let mut corr = 0f32;
7971    for p in c0..c1 {
7972        let e = entries_off + p * 4;
7973        let col = u16::from_le_bytes([bytes[e], bytes[e + 1]]) as usize;
7974        let val = f16_to_f32(u16::from_le_bytes([bytes[e + 2], bytes[e + 3]]));
7975        corr += val * x[col];
7976    }
7977    corr
7978}
7979
7980/// Dequantize one q1t row into `buf[..cols]` via the sign LUT (no division),
7981/// then apply the row's outliers (its `[row_ptr[r], row_ptr[r+1])` slice).
7982/// Used by the batched (prefill) path where the decode amortizes over the batch.
7983fn q1t_dequant_row(
7984    bytes: &[u8],
7985    r: usize,
7986    gpr: usize,
7987    rp_off: usize,
7988    entries_off: usize,
7989    has_ov: bool,
7990    buf: &mut [f32],
7991) {
7992    const TILE: usize = cortiq_core::quant::Q1T_TILE;
7993    for g in 0..gpr {
7994        let off = (r * gpr + g) * TILE;
7995        let s = f16_to_f32(u16::from_le_bytes([bytes[off], bytes[off + 1]]));
7996        let codes = &bytes[off + 2..off + TILE];
7997        let bc = g * GROUP_SIZE;
7998        // 6 full bytes (30 codes) + a 7th byte holding the last 2.
7999        for bi in 0..6 {
8000            let lut = &SIGN5[codes[bi] as usize];
8001            let d = &mut buf[bc + bi * 5..bc + bi * 5 + 5];
8002            for i in 0..5 {
8003                d[i] = lut[i] * s;
8004            }
8005        }
8006        let lut = &SIGN5[codes[6] as usize];
8007        buf[bc + 30] = lut[0] * s;
8008        buf[bc + 31] = lut[1] * s;
8009    }
8010    if !has_ov {
8011        return;
8012    }
8013    let (c0, c1) = (
8014        q1t_rowptr(bytes, rp_off, r),
8015        q1t_rowptr(bytes, rp_off, r + 1),
8016    );
8017    for p in c0..c1 {
8018        let e = entries_off + p * 4;
8019        let col = u16::from_le_bytes([bytes[e], bytes[e + 1]]) as usize;
8020        buf[col] = f16_to_f32(u16::from_le_bytes([bytes[e + 2], bytes[e + 3]]));
8021    }
8022}
8023
8024/// Add the sparse outlier overlay onto a base dot already in `out` (the GPU
8025/// computes the ternary base; the overlay stays on the CPU — its entries are
8026/// few and its per-row gather doesn't vectorize on the GPU). Row-parallel.
8027fn q1t_add_overlay(
8028    bytes: &[u8],
8029    x: &[f32],
8030    rows: usize,
8031    cols: usize,
8032    out: &mut [f32],
8033    pool: Option<&Pool>,
8034) {
8035    const TILE: usize = cortiq_core::quant::Q1T_TILE;
8036    let gpr = cols / GROUP_SIZE;
8037    let (rp_off, ent_off, has_ov) = q1t_overlay(bytes, rows * gpr * TILE, rows);
8038    if !has_ov {
8039        return;
8040    }
8041    let out_addr = SendMut(out.as_mut_ptr());
8042    let run = move |start: usize, end: usize| {
8043        for r in start..end {
8044            let corr = q1t_row_outlier_correction(bytes, r, rp_off, ent_off, has_ov, x);
8045            // SAFETY: disjoint rows; add onto the base the GPU already wrote.
8046            unsafe { *out_addr.at(r) += corr };
8047        }
8048    };
8049    dispatch_rows(pool, rows, &run);
8050}
8051
8052/// Q1T row range via the A8W8 int8 path — shared activation split,
8053/// per-row: base SDOT dot + outlier correction + overlay.
8054#[allow(clippy::too_many_arguments)]
8055fn q1t_range_a8w8(
8056    bytes: &[u8],
8057    gpr: usize,
8058    rp_off: usize,
8059    ent_off: usize,
8060    has_ov: bool,
8061    act: &SplitAct,
8062    x: &[f32],
8063    out: SendMut,
8064    start: usize,
8065    end: usize,
8066) {
8067    for r in start..end {
8068        let mut acc = q1t_dot_row_i8(bytes, r, gpr, &act.xq) * act.sx;
8069        for &(j, xv) in &act.outliers {
8070            acc += q1t_base_weight(bytes, r, gpr, j) * xv;
8071        }
8072        acc += q1t_row_outlier_correction(bytes, r, rp_off, ent_off, has_ov, x);
8073        // SAFETY: disjoint row ranges per worker.
8074        unsafe { *out.at(r) = acc };
8075    }
8076}
8077
8078/// Q1T row range via the f32 path (no SDOT) — for matvec_many batched
8079/// dispatch when a8w8 is unavailable.
8080#[allow(clippy::too_many_arguments)]
8081fn q1t_range_f32_batch(
8082    bytes: &[u8],
8083    gpr: usize,
8084    rp_off: usize,
8085    ent_off: usize,
8086    has_ov: bool,
8087    x: &[f32],
8088    out: SendMut,
8089    start: usize,
8090    end: usize,
8091) {
8092    const TILE: usize = cortiq_core::quant::Q1T_TILE;
8093    let mut sg = [0f32; GROUP_SIZE];
8094    for r in start..end {
8095        let mut acc = 0f32;
8096        for g in 0..gpr {
8097            let off = (r * gpr + g) * TILE;
8098            let s = f16_to_f32(u16::from_le_bytes([bytes[off], bytes[off + 1]]));
8099            let codes = &bytes[off + 2..off + TILE];
8100            let xg = &x[g * GROUP_SIZE..g * GROUP_SIZE + GROUP_SIZE];
8101            for bi in 0..6 {
8102                sg[bi * 5..bi * 5 + 5].copy_from_slice(&SIGN5[codes[bi] as usize]);
8103            }
8104            let lut = &SIGN5[codes[6] as usize];
8105            sg[30] = lut[0];
8106            sg[31] = lut[1];
8107            let mut gsum = 0f32;
8108            for k in 0..GROUP_SIZE {
8109                gsum += sg[k] * xg[k];
8110            }
8111            acc += s * gsum;
8112        }
8113        acc += q1t_row_outlier_correction(bytes, r, rp_off, ent_off, has_ov, x);
8114        // SAFETY: disjoint row ranges per worker.
8115        unsafe { *out.at(r) = acc };
8116    }
8117}
8118
8119/// Ternary (q1t) matvec — decode+dot straight from mmap, one group at a time:
8120/// no per-ROW buffer, no division (the sign LUT), and a tiny per-group sign
8121/// buffer so the 32-wide dot vectorizes. This is the decode hot path.
8122fn q1t_matvec(
8123    bytes: &[u8],
8124    x: &[f32],
8125    rows: usize,
8126    cols: usize,
8127    out: &mut [f32],
8128    pool: Option<&Pool>,
8129) {
8130    debug_assert_eq!(out.len(), rows);
8131    const TILE: usize = cortiq_core::quant::Q1T_TILE;
8132    let gpr = cols / GROUP_SIZE;
8133    let (rp_off, ent_off, has_ov) = q1t_overlay(bytes, rows * gpr * TILE, rows);
8134    let out_addr = SendMut(out.as_mut_ptr());
8135    // int8 SDOT base dot (ARM dotprod): ~4× the f32 arithmetic. x → i8 once
8136    // (`split_act`), activation outliers added back exactly in f32, weight
8137    // overlay on top. ARM SDOT / x86 AVX2; CMF_SDOT=0 keeps the exact f32 path.
8138    if a8w8_enabled() {
8139        let act = split_act(x);
8140        let act = &act;
8141        let run = move |start: usize, end: usize| {
8142            for r in start..end {
8143                let mut acc = q1t_dot_row_i8(bytes, r, gpr, &act.xq) * act.sx;
8144                for &(j, xv) in &act.outliers {
8145                    acc += q1t_base_weight(bytes, r, gpr, j) * xv;
8146                }
8147                acc += q1t_row_outlier_correction(bytes, r, rp_off, ent_off, has_ov, x);
8148                // SAFETY: disjoint row ranges per worker.
8149                unsafe { *out_addr.at(r) = acc };
8150            }
8151        };
8152        dispatch_rows(pool, rows, &run);
8153        return;
8154    }
8155    let run = move |start: usize, end: usize| {
8156        // Per-group signs, unpacked contiguously so the dot below is a clean
8157        // 32-wide reduction the autovectorizer turns into f32x4 FMAs — the
8158        // 5-values-per-byte base-3 layout won't SIMD in place.
8159        let mut sg = [0f32; GROUP_SIZE];
8160        for r in start..end {
8161            let mut acc = 0f32;
8162            for g in 0..gpr {
8163                let off = (r * gpr + g) * TILE;
8164                let s = f16_to_f32(u16::from_le_bytes([bytes[off], bytes[off + 1]]));
8165                let codes = &bytes[off + 2..off + TILE];
8166                let xg = &x[g * GROUP_SIZE..g * GROUP_SIZE + GROUP_SIZE];
8167                for bi in 0..6 {
8168                    sg[bi * 5..bi * 5 + 5].copy_from_slice(&SIGN5[codes[bi] as usize]);
8169                }
8170                let lut = &SIGN5[codes[6] as usize];
8171                sg[30] = lut[0];
8172                sg[31] = lut[1];
8173                let mut gsum = 0f32;
8174                for k in 0..GROUP_SIZE {
8175                    gsum += sg[k] * xg[k];
8176                }
8177                acc += s * gsum;
8178            }
8179            acc += q1t_row_outlier_correction(bytes, r, rp_off, ent_off, has_ov, x);
8180            unsafe { *out_addr.at(r) = acc };
8181        }
8182    };
8183    dispatch_rows(pool, rows, &run);
8184}
8185
8186/// Fused-pair twin of `q1t_dot_row_sdot`: ONE register unpack of the
8187/// ternary codes serves BOTH activation streams (the unpack chain is
8188/// the dominant per-row cost — MTP verify pairs paid it twice). Per
8189/// stream the group order and f32 accumulation match the single-row
8190/// kernel exactly, so pair == 2×matvec bit-for-bit.
8191#[cfg(target_arch = "aarch64")]
8192#[target_feature(enable = "neon,dotprod")]
8193unsafe fn q1t_dot_row_sdot2(bytes: &[u8], r: usize, gpr: usize, xa: &[i8], xb: &[i8]) -> [f32; 2] {
8194    use core::arch::aarch64::*;
8195    use core::arch::asm;
8196    // SAFETY: same slice-length contracts as `q1t_dot_row_sdot`, ×2.
8197    unsafe {
8198        const TILE: usize = cortiq_core::quant::Q1T_TILE;
8199        let bytes_ptr = bytes.as_ptr();
8200        let row_off = r * gpr * TILE;
8201        let xp = [xa.as_ptr(), xb.as_ptr()];
8202        let mut acc = [0f32; 2];
8203        macro_rules! sdot2 {
8204            ($w0:expr, $w1:expr, $x:expr) => {{
8205                let x0 = vld1q_s8($x);
8206                let x1 = vld1q_s8($x.add(16));
8207                let (mut a0, mut a1) = (vdupq_n_s32(0), vdupq_n_s32(0));
8208                asm!(
8209                    "sdot {a0:v}.4s, {w0:v}.16b, {x0:v}.16b",
8210                    "sdot {a1:v}.4s, {w1:v}.16b, {x1:v}.16b",
8211                    a0 = inout(vreg) a0, a1 = inout(vreg) a1,
8212                    w0 = in(vreg) $w0, x0 = in(vreg) x0, w1 = in(vreg) $w1, x1 = in(vreg) x1,
8213                    options(pure, nomem, nostack),
8214                );
8215                vaddvq_s32(vaddq_s32(a0, a1))
8216            }};
8217        }
8218        let gpr2 = gpr & !1;
8219        let mut gi = 0;
8220        while gi < gpr2 {
8221            let off0 = row_off + gi * TILE;
8222            let off1 = off0 + TILE;
8223            let s0 = f16_to_f32(u16::from_le_bytes([
8224                *bytes_ptr.add(off0),
8225                *bytes_ptr.add(off0 + 1),
8226            ]));
8227            let s1 = f16_to_f32(u16::from_le_bytes([
8228                *bytes_ptr.add(off1),
8229                *bytes_ptr.add(off1 + 1),
8230            ]));
8231            let (u0_0, u1_0, u2_0, u3_0) = q1t_unpack_reg_u64s(bytes_ptr.add(off0 + 2));
8232            let (u0_1, u1_1, u2_1, u3_1) = q1t_unpack_reg_u64s(bytes_ptr.add(off1 + 2));
8233            let w0_0 = vreinterpretq_s8_u64(vcombine_u64(vcreate_u64(u0_0), vcreate_u64(u1_0)));
8234            let w1_0 = vreinterpretq_s8_u64(vcombine_u64(vcreate_u64(u2_0), vcreate_u64(u3_0)));
8235            let w0_1 = vreinterpretq_s8_u64(vcombine_u64(vcreate_u64(u0_1), vcreate_u64(u1_1)));
8236            let w1_1 = vreinterpretq_s8_u64(vcombine_u64(vcreate_u64(u2_1), vcreate_u64(u3_1)));
8237            for k in 0..2 {
8238                let d0 = sdot2!(w0_0, w1_0, xp[k].add(gi * GROUP_SIZE));
8239                let d1 = sdot2!(w0_1, w1_1, xp[k].add((gi + 1) * GROUP_SIZE));
8240                acc[k] += d0 as f32 * s0 + d1 as f32 * s1;
8241            }
8242            gi += 2;
8243        }
8244        if gi < gpr {
8245            let off = row_off + gi * TILE;
8246            let s = f16_to_f32(u16::from_le_bytes([
8247                *bytes_ptr.add(off),
8248                *bytes_ptr.add(off + 1),
8249            ]));
8250            let (u0, u1, u2, u3) = q1t_unpack_reg_u64s(bytes_ptr.add(off + 2));
8251            let w0 = vreinterpretq_s8_u64(vcombine_u64(vcreate_u64(u0), vcreate_u64(u1)));
8252            let w1 = vreinterpretq_s8_u64(vcombine_u64(vcreate_u64(u2), vcreate_u64(u3)));
8253            for k in 0..2 {
8254                let d = sdot2!(w0, w1, xp[k].add(gi * GROUP_SIZE));
8255                acc[k] += d as f32 * s;
8256            }
8257        }
8258        acc
8259    }
8260}
8261
8262/// Fused Q1T pair matvec: ONE pass over the rows serves both
8263/// activation streams — on ARM the ternary register unpack happens
8264/// once per tile pair (`q1t_dot_row_sdot2`); elsewhere the second dot
8265/// rides the row's L1-warm tile bytes. Per stream the math matches
8266/// `q1t_matvec` exactly.
8267fn q1t_matvec2(
8268    bytes: &[u8],
8269    x1: &[f32],
8270    x2: &[f32],
8271    rows: usize,
8272    cols: usize,
8273    o1: &mut [f32],
8274    o2: &mut [f32],
8275    pool: Option<&Pool>,
8276) {
8277    debug_assert_eq!(o1.len(), rows);
8278    debug_assert_eq!(o2.len(), rows);
8279    const TILE: usize = cortiq_core::quant::Q1T_TILE;
8280    let gpr = cols / GROUP_SIZE;
8281    let (rp_off, ent_off, has_ov) = q1t_overlay(bytes, rows * gpr * TILE, rows);
8282    let out1 = SendMut(o1.as_mut_ptr());
8283    let out2 = SendMut(o2.as_mut_ptr());
8284    if a8w8_enabled() {
8285        let a1 = split_act(x1);
8286        let a2 = split_act(x2);
8287        let (a1, a2) = (&a1, &a2);
8288        let run = move |start: usize, end: usize| {
8289            for r in start..end {
8290                #[cfg(target_arch = "aarch64")]
8291                // a8w8 on aarch64 ⇔ sdot_enabled(), so the kernel's
8292                // target features are present.
8293                let ds = unsafe { q1t_dot_row_sdot2(bytes, r, gpr, &a1.xq, &a2.xq) };
8294                #[cfg(not(target_arch = "aarch64"))]
8295                let ds = [
8296                    q1t_dot_row_i8(bytes, r, gpr, &a1.xq),
8297                    q1t_dot_row_i8(bytes, r, gpr, &a2.xq),
8298                ];
8299                let mut acc1 = ds[0] * a1.sx;
8300                for &(j, xv) in &a1.outliers {
8301                    acc1 += q1t_base_weight(bytes, r, gpr, j) * xv;
8302                }
8303                acc1 += q1t_row_outlier_correction(bytes, r, rp_off, ent_off, has_ov, x1);
8304                let mut acc2 = ds[1] * a2.sx;
8305                for &(j, xv) in &a2.outliers {
8306                    acc2 += q1t_base_weight(bytes, r, gpr, j) * xv;
8307                }
8308                acc2 += q1t_row_outlier_correction(bytes, r, rp_off, ent_off, has_ov, x2);
8309                // SAFETY: disjoint row ranges per worker.
8310                unsafe {
8311                    *out1.at(r) = acc1;
8312                    *out2.at(r) = acc2;
8313                }
8314            }
8315        };
8316        dispatch_rows(pool, rows, &run);
8317        return;
8318    }
8319    let run = move |start: usize, end: usize| {
8320        // Exact path (CMF_SDOT=0): unpack the sign LUT once per group,
8321        // dot both streams — same op order per stream as `q1t_matvec`.
8322        let mut sg = [0f32; GROUP_SIZE];
8323        for r in start..end {
8324            let mut acc1 = 0f32;
8325            let mut acc2 = 0f32;
8326            for g in 0..gpr {
8327                let off = (r * gpr + g) * TILE;
8328                let s = f16_to_f32(u16::from_le_bytes([bytes[off], bytes[off + 1]]));
8329                let codes = &bytes[off + 2..off + TILE];
8330                for bi in 0..6 {
8331                    sg[bi * 5..bi * 5 + 5].copy_from_slice(&SIGN5[codes[bi] as usize]);
8332                }
8333                let lut = &SIGN5[codes[6] as usize];
8334                sg[30] = lut[0];
8335                sg[31] = lut[1];
8336                let xg1 = &x1[g * GROUP_SIZE..g * GROUP_SIZE + GROUP_SIZE];
8337                let xg2 = &x2[g * GROUP_SIZE..g * GROUP_SIZE + GROUP_SIZE];
8338                let mut gsum1 = 0f32;
8339                for k in 0..GROUP_SIZE {
8340                    gsum1 += sg[k] * xg1[k];
8341                }
8342                acc1 += s * gsum1;
8343                let mut gsum2 = 0f32;
8344                for k in 0..GROUP_SIZE {
8345                    gsum2 += sg[k] * xg2[k];
8346                }
8347                acc2 += s * gsum2;
8348            }
8349            acc1 += q1t_row_outlier_correction(bytes, r, rp_off, ent_off, has_ov, x1);
8350            acc2 += q1t_row_outlier_correction(bytes, r, rp_off, ent_off, has_ov, x2);
8351            // SAFETY: disjoint row ranges per worker.
8352            unsafe {
8353                *out1.at(r) = acc1;
8354                *out2.at(r) = acc2;
8355            }
8356        }
8357    };
8358    dispatch_rows(pool, rows, &run);
8359}
8360
8361/// Ternary (q1t) matmat (prefill) — dequant each row once, dot the whole
8362/// batch against it (amortizes the per-row decode).
8363fn q1t_matmat(
8364    bytes: &[u8],
8365    xs: &[f32],
8366    b: usize,
8367    rows: usize,
8368    cols: usize,
8369    out: &mut [f32],
8370    pool: Option<&Pool>,
8371) {
8372    debug_assert_eq!(out.len(), b * rows);
8373    const TILE: usize = cortiq_core::quant::Q1T_TILE;
8374    let gpr = cols / GROUP_SIZE;
8375    let (rp_off, ent_off, has_ov) = q1t_overlay(bytes, rows * gpr * TILE, rows);
8376    let out_addr = SendMut(out.as_mut_ptr());
8377    // int8 prefill (ARM SDOT / x86 AVX2): quantize the B inputs once, unpack
8378    // each weight row's signs to i8 ONCE, then int8-dot against every input —
8379    // the row sign-decode amortizes over the whole batch. CMF_SDOT=0 → f32.
8380    if a8w8_enabled() {
8381        let acts: Vec<SplitAct> = (0..b)
8382            .map(|bi| split_act(&xs[bi * cols..(bi + 1) * cols]))
8383            .collect();
8384        let acts = &acts;
8385        let run = move |start: usize, end: usize| {
8386            let mut sg = vec![0i8; cols + 8]; // row signs, i8 (+8 unpack slack)
8387            let mut sc = vec![0f32; gpr]; // per-group scales
8388            let mut accs = vec![0f32; b]; // per-batch accumulators, reused per row
8389            for r in start..end {
8390                for g in 0..gpr {
8391                    let off = (r * gpr + g) * TILE;
8392                    sc[g] = f16_to_f32(u16::from_le_bytes([bytes[off], bytes[off + 1]]));
8393                    q1t_unpack_group_i8(
8394                        bytes.as_ptr().wrapping_add(off + 2),
8395                        &mut sg[g * GROUP_SIZE..],
8396                    );
8397                }
8398                for bi in 0..b {
8399                    let act = &acts[bi];
8400                    let mut isum = 0f32;
8401                    for g in 0..gpr {
8402                        let d = q1t_i8dot32(
8403                            sg.as_ptr().wrapping_add(g * GROUP_SIZE),
8404                            act.xq.as_ptr().wrapping_add(g * GROUP_SIZE),
8405                        );
8406                        isum += d as f32 * sc[g];
8407                    }
8408                    let mut acc = isum * act.sx;
8409                    for &(j, xv) in &act.outliers {
8410                        acc += q1t_base_weight(bytes, r, gpr, j) * xv;
8411                    }
8412                    accs[bi] = acc;
8413                }
8414                // Overlay ONCE per row for the whole batch: read each (col, val)
8415                // from mmap a single time (was b× — the re-read dominated prefill)
8416                // and fan it out over the batch via the cached inputs.
8417                if has_ov {
8418                    let (c0, c1) = (
8419                        q1t_rowptr(bytes, rp_off, r),
8420                        q1t_rowptr(bytes, rp_off, r + 1),
8421                    );
8422                    for p in c0..c1 {
8423                        let e = ent_off + p * 4;
8424                        let col = u16::from_le_bytes([bytes[e], bytes[e + 1]]) as usize;
8425                        let val = f16_to_f32(u16::from_le_bytes([bytes[e + 2], bytes[e + 3]]));
8426                        for bi in 0..b {
8427                            accs[bi] += val * xs[bi * cols + col];
8428                        }
8429                    }
8430                }
8431                for bi in 0..b {
8432                    unsafe { *out_addr.at(bi * rows + r) = accs[bi] };
8433                }
8434            }
8435        };
8436        dispatch_rows(pool, rows, &run);
8437        return;
8438    }
8439    let run = move |start: usize, end: usize| {
8440        let mut buf = vec![0f32; cols];
8441        for r in start..end {
8442            q1t_dequant_row(bytes, r, gpr, rp_off, ent_off, has_ov, &mut buf);
8443            for bi in 0..b {
8444                let xr = &xs[bi * cols..(bi + 1) * cols];
8445                let mut acc = 0f32;
8446                for j in 0..cols {
8447                    acc += buf[j] * xr[j];
8448                }
8449                unsafe { *out_addr.at(bi * rows + r) = acc };
8450            }
8451        }
8452    };
8453    dispatch_rows(pool, rows, &run);
8454}
8455
8456fn q1_matvec(
8457    bytes: &[u8],
8458    x: &[f32],
8459    rows: usize,
8460    cols: usize,
8461    out: &mut [f32],
8462    pool: Option<&Pool>,
8463) {
8464    debug_assert_eq!(out.len(), rows);
8465    let gpr = cols / GROUP_SIZE;
8466    let out_addr = SendMut(out.as_mut_ptr());
8467    if a8w8_enabled() {
8468        let act = split_act(x);
8469        let gsum = q1_group_sums(&act.xq, gpr);
8470        let (act, gsum) = (&act, &gsum);
8471        let run = move |start: usize, end: usize| {
8472            q1_range_a8w8(bytes, gpr, act, gsum, out_addr, start, end)
8473        };
8474        dispatch_rows(pool, rows, &run);
8475        return;
8476    }
8477    let run = move |start: usize, end: usize| q1_range_f32(bytes, gpr, x, out_addr, start, end);
8478    dispatch_rows(pool, rows, &run);
8479}
8480
8481/// Fused two-input q1 matvec (weights read once per pair).
8482#[allow(clippy::too_many_arguments)]
8483fn q1_matvec2(
8484    bytes: &[u8],
8485    x1: &[f32],
8486    x2: &[f32],
8487    rows: usize,
8488    cols: usize,
8489    o1: &mut [f32],
8490    o2: &mut [f32],
8491    pool: Option<&Pool>,
8492) {
8493    let gpr = cols / GROUP_SIZE;
8494    let p1 = SendMut(o1.as_mut_ptr());
8495    let p2 = SendMut(o2.as_mut_ptr());
8496    if a8w8_enabled() {
8497        let a1 = split_act(x1);
8498        let a2 = split_act(x2);
8499        let g1 = q1_group_sums(&a1.xq, gpr);
8500        let g2 = q1_group_sums(&a2.xq, gpr);
8501        let (a1, a2, g1, g2) = (&a1, &a2, &g1, &g2);
8502        let run = move |start: usize, end: usize| {
8503            for r in start..end {
8504                let mut v1 = dot_q1_row_i8(bytes, r, gpr, &a1.xq, g1) * a1.sx;
8505                let mut v2 = dot_q1_row_i8(bytes, r, gpr, &a2.xq, g2) * a2.sx;
8506                for &(j, xv) in &a1.outliers {
8507                    let (w, s) = q1_outlier(bytes, r, gpr, j);
8508                    v1 += w * s * xv;
8509                }
8510                for &(j, xv) in &a2.outliers {
8511                    let (w, s) = q1_outlier(bytes, r, gpr, j);
8512                    v2 += w * s * xv;
8513                }
8514                // SAFETY: disjoint row ranges per worker.
8515                unsafe {
8516                    *p1.at(r) = v1;
8517                    *p2.at(r) = v2;
8518                }
8519            }
8520        };
8521        dispatch_rows(pool, rows, &run);
8522        return;
8523    }
8524    let run = move |start: usize, end: usize| {
8525        for r in start..end {
8526            // SAFETY: disjoint row ranges per worker.
8527            unsafe {
8528                *p1.at(r) = q1_row_exact(bytes, r, gpr, x1);
8529                *p2.at(r) = q1_row_exact(bytes, r, gpr, x2);
8530            }
8531        }
8532    };
8533    dispatch_rows(pool, rows, &run);
8534}
8535
8536/// Batched q1 matmat: each row's tiles stream once per microbatch.
8537#[allow(clippy::too_many_arguments)]
8538fn q1_matmat(
8539    bytes: &[u8],
8540    xs_all: &[f32],
8541    b: usize,
8542    rows: usize,
8543    cols: usize,
8544    out: &mut [f32],
8545    pool: Option<&Pool>,
8546) {
8547    debug_assert_eq!(out.len(), b * rows);
8548    let gpr = cols / GROUP_SIZE;
8549    let out_addr = SendMut(out.as_mut_ptr());
8550    if a8w8_enabled() {
8551        let acts: Vec<(SplitAct, Vec<i32>)> = (0..b)
8552            .map(|bi| {
8553                let act = split_act(&xs_all[bi * cols..(bi + 1) * cols]);
8554                let gsum = q1_group_sums(&act.xq, gpr);
8555                (act, gsum)
8556            })
8557            .collect();
8558        let acts = &acts;
8559        #[cfg(target_arch = "x86_64")]
8560        let blocked_ok = avx2_enabled() && blocked_enabled();
8561        #[cfg(target_arch = "aarch64")]
8562        let blocked_ok = sdot_enabled() && blocked_enabled();
8563        let run = move |start: usize, end: usize| {
8564            for r in start..end {
8565                let mut bi = 0usize;
8566                // Blocked 1×4: the unpacked bit mask serves four
8567                // activation streams per group.
8568                #[cfg(target_arch = "aarch64")]
8569                if blocked_ok {
8570                    while bi + 4 <= acts.len() {
8571                        let xs = [
8572                            acts[bi].0.xq.as_slice(),
8573                            acts[bi + 1].0.xq.as_slice(),
8574                            acts[bi + 2].0.xq.as_slice(),
8575                            acts[bi + 3].0.xq.as_slice(),
8576                        ];
8577                        let gs = [
8578                            acts[bi].1.as_slice(),
8579                            acts[bi + 1].1.as_slice(),
8580                            acts[bi + 2].1.as_slice(),
8581                            acts[bi + 3].1.as_slice(),
8582                        ];
8583                        let d = unsafe { dot_q1_row_1x4_sdot(bytes, r, gpr, xs, gs) };
8584                        for k in 0..4 {
8585                            let (act, _) = &acts[bi + k];
8586                            let mut acc = d[k] * act.sx;
8587                            for &(j, xv) in &act.outliers {
8588                                let (w, sc) = q1_outlier(bytes, r, gpr, j);
8589                                acc += w * sc * xv;
8590                            }
8591                            // SAFETY: disjoint (bi, r) cells per worker.
8592                            unsafe { *out_addr.at((bi + k) * rows + r) = acc };
8593                        }
8594                        bi += 4;
8595                    }
8596                }
8597                #[cfg(target_arch = "x86_64")]
8598                if blocked_ok {
8599                    while bi + 4 <= acts.len() {
8600                        let xs = [
8601                            acts[bi].0.xq.as_slice(),
8602                            acts[bi + 1].0.xq.as_slice(),
8603                            acts[bi + 2].0.xq.as_slice(),
8604                            acts[bi + 3].0.xq.as_slice(),
8605                        ];
8606                        let gs = [
8607                            acts[bi].1.as_slice(),
8608                            acts[bi + 1].1.as_slice(),
8609                            acts[bi + 2].1.as_slice(),
8610                            acts[bi + 3].1.as_slice(),
8611                        ];
8612                        let d = unsafe {
8613                            if vnni_tiles_enabled() {
8614                                dot_q1_row_1x4_vnni(bytes, r, gpr, xs, gs)
8615                            } else {
8616                                dot_q1_row_1x4_avx2(bytes, r, gpr, xs, gs)
8617                            }
8618                        };
8619                        for k in 0..4 {
8620                            let (act, _) = &acts[bi + k];
8621                            let mut acc = d[k] * act.sx;
8622                            for &(j, xv) in &act.outliers {
8623                                let (w, sc) = q1_outlier(bytes, r, gpr, j);
8624                                acc += w * sc * xv;
8625                            }
8626                            // SAFETY: disjoint (bi, r) cells per worker.
8627                            unsafe { *out_addr.at((bi + k) * rows + r) = acc };
8628                        }
8629                        bi += 4;
8630                    }
8631                }
8632                while bi < acts.len() {
8633                    let (act, gsum) = &acts[bi];
8634                    let mut acc = dot_q1_row_i8(bytes, r, gpr, &act.xq, gsum) * act.sx;
8635                    for &(j, xv) in &act.outliers {
8636                        let (w, s) = q1_outlier(bytes, r, gpr, j);
8637                        acc += w * s * xv;
8638                    }
8639                    // SAFETY: disjoint (bi, r) cells per worker range.
8640                    unsafe { *out_addr.at(bi * rows + r) = acc };
8641                    bi += 1;
8642                }
8643            }
8644        };
8645        dispatch_rows(pool, rows, &run);
8646        return;
8647    }
8648    let run = move |start: usize, end: usize| {
8649        for r in start..end {
8650            for bi in 0..b {
8651                let x = &xs_all[bi * cols..(bi + 1) * cols];
8652                // SAFETY: disjoint (bi, r) cells per worker range.
8653                unsafe { *out_addr.at(bi * rows + r) = q1_row_exact(bytes, r, gpr, x) };
8654            }
8655        }
8656    };
8657    dispatch_rows(pool, rows, &run);
8658}
8659
8660/// Fused q4_block matvec straight from the mapped bytes. SDOT path when
8661/// dotprod is available (port of vmfcore `dot_q4_block_sdot`, measured
8662/// +23% on q4 decode): nibbles → centered i8, int8×int8 `sdot` per
8663/// 32-group, exact outlier correction — the same A8W8 contract as q8.
8664/// `CMF_SDOT=0` keeps the exact scalar path.
8665fn q4matvec(
8666    bytes: &[u8],
8667    x: &[f32],
8668    rows: usize,
8669    cols: usize,
8670    out: &mut [f32],
8671    pool: Option<&Pool>,
8672) {
8673    debug_assert_eq!(out.len(), rows);
8674    let (packed, scales) = q4_split(bytes, rows, cols);
8675    let gpr = cols / GROUP_SIZE;
8676    let out_addr = SendMut(out.as_mut_ptr());
8677
8678    if a8w8_enabled() {
8679        let act = split_act(x);
8680        let run = move |start: usize, end: usize| {
8681            q4_range_a8w8(packed, scales, gpr, cols, &act, out_addr, start, end)
8682        };
8683        dispatch_rows(pool, rows, &run);
8684        return;
8685    }
8686
8687    let run =
8688        move |start: usize, end: usize| q4_range_f32(packed, scales, gpr, x, out_addr, start, end);
8689    dispatch_rows(pool, rows, &run);
8690}
8691
8692/// One q4 row via the A8W8 int8 path — SDOT on ARM, AVX2 maddubs on
8693/// x86 (scalar fallback is unreachable: callers gate on a8w8_enabled).
8694#[inline]
8695#[allow(unreachable_code)]
8696/// One UNPACKED q4 row (centered i8 in `buf`) against four activation
8697/// streams: the 32-byte weight chunk and its abs() load once per group,
8698/// the per-group f16 scale decodes once — four maddubs+reduce chains
8699/// instead of four full (load, abs, dot) rounds.
8700#[cfg(target_arch = "x86_64")]
8701#[target_feature(enable = "avx2")]
8702unsafe fn dot_q4b_row_1x4_avx2(
8703    buf: &[u8],
8704    scales: &[u8],
8705    g0: usize,
8706    gpr: usize,
8707    xs: [&[i8]; 4],
8708) -> [f32; 4] {
8709    // SAFETY: callers uphold buffer contracts (buf.len() == gpr·32).
8710    unsafe {
8711        use core::arch::x86_64::*;
8712        let ones = _mm256_set1_epi16(1);
8713        let mut acc = [0f32; 4];
8714        for gi in 0..gpr {
8715            let s = f16_to_f32(u16::from_le_bytes([
8716                scales[(g0 + gi) * 2],
8717                scales[(g0 + gi) * 2 + 1],
8718            ]));
8719            let w = _mm256_loadu_si256(buf.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
8720            let aw = _mm256_abs_epi8(w);
8721            for (k, xq) in xs.iter().enumerate() {
8722                let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
8723                let p16 = _mm256_maddubs_epi16(aw, _mm256_sign_epi8(x, w));
8724                let d = _mm256_madd_epi16(p16, ones);
8725                let hi128 = _mm256_extracti128_si256::<1>(d);
8726                let s128 = _mm_add_epi32(_mm256_castsi256_si128(d), hi128);
8727                let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
8728                let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
8729                acc[k] += _mm_cvtsi128_si32(s32) as f32 * s;
8730            }
8731        }
8732        acc
8733    }
8734}
8735
8736/// VNNI twin of `dot_q4b_row_1x4_avx2` (see `dpbusd_hsum`).
8737#[cfg(target_arch = "x86_64")]
8738#[target_feature(enable = "avx2,avx512f,avx512bw,avx512vl,avx512vnni")]
8739unsafe fn dot_q4b_row_1x4_vnni(
8740    buf: &[u8],
8741    scales: &[u8],
8742    g0: usize,
8743    gpr: usize,
8744    xs: [&[i8]; 4],
8745) -> [f32; 4] {
8746    // SAFETY: callers uphold buffer contracts (buf.len() == gpr·32).
8747    unsafe {
8748        use core::arch::x86_64::*;
8749        let mut acc = [0f32; 4];
8750        for gi in 0..gpr {
8751            let s = f16_to_f32(u16::from_le_bytes([
8752                scales[(g0 + gi) * 2],
8753                scales[(g0 + gi) * 2 + 1],
8754            ]));
8755            let w = _mm256_loadu_si256(buf.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
8756            let aw = _mm256_abs_epi8(w);
8757            for (k, xq) in xs.iter().enumerate() {
8758                let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
8759                let d = dpbusd_hsum(aw, _mm256_sign_epi8(x, w));
8760                acc[k] += d as f32 * s;
8761            }
8762        }
8763        acc
8764    }
8765}
8766
8767/// The vbit flavor of the blocked 1×4: the per-activation A8W8 scale
8768/// folds in PER GROUP as `(d·sx)·s` — bit-matching the single-matvec
8769/// accumulation order (the q4_block flavor applies sx once at the end,
8770/// matching ITS single path; the two conventions are historical and
8771/// each blocked leg must mirror its own).
8772#[cfg(target_arch = "x86_64")]
8773#[target_feature(enable = "avx2")]
8774unsafe fn dot_q4b_row_1x4_sx_avx2(
8775    buf: &[u8],
8776    scales: &[u8],
8777    g0: usize,
8778    gpr: usize,
8779    xs: [&[i8]; 4],
8780    sxs: [f32; 4],
8781) -> [f32; 4] {
8782    // SAFETY: callers uphold buffer contracts (buf.len() == gpr·32).
8783    unsafe {
8784        use core::arch::x86_64::*;
8785        let ones = _mm256_set1_epi16(1);
8786        let mut acc = [0f32; 4];
8787        for gi in 0..gpr {
8788            let s = f16_to_f32(u16::from_le_bytes([
8789                scales[(g0 + gi) * 2],
8790                scales[(g0 + gi) * 2 + 1],
8791            ]));
8792            let w = _mm256_loadu_si256(buf.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
8793            let aw = _mm256_abs_epi8(w);
8794            for (k, xq) in xs.iter().enumerate() {
8795                let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
8796                let p16 = _mm256_maddubs_epi16(aw, _mm256_sign_epi8(x, w));
8797                let d = _mm256_madd_epi16(p16, ones);
8798                let hi128 = _mm256_extracti128_si256::<1>(d);
8799                let s128 = _mm_add_epi32(_mm256_castsi256_si128(d), hi128);
8800                let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
8801                let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
8802                acc[k] += (_mm_cvtsi128_si32(s32) as f32 * sxs[k]) * s;
8803            }
8804        }
8805        acc
8806    }
8807}
8808
8809/// VNNI twin of `dot_q4b_row_1x4_sx_avx2` (see `dpbusd_hsum`; the
8810/// per-group `(d·sx)·s` fold mirrors the vbit single path).
8811#[cfg(target_arch = "x86_64")]
8812#[target_feature(enable = "avx2,avx512f,avx512bw,avx512vl,avx512vnni")]
8813unsafe fn dot_q4b_row_1x4_sx_vnni(
8814    buf: &[u8],
8815    scales: &[u8],
8816    g0: usize,
8817    gpr: usize,
8818    xs: [&[i8]; 4],
8819    sxs: [f32; 4],
8820) -> [f32; 4] {
8821    // SAFETY: callers uphold buffer contracts (buf.len() == gpr·32).
8822    unsafe {
8823        use core::arch::x86_64::*;
8824        let mut acc = [0f32; 4];
8825        for gi in 0..gpr {
8826            let s = f16_to_f32(u16::from_le_bytes([
8827                scales[(g0 + gi) * 2],
8828                scales[(g0 + gi) * 2 + 1],
8829            ]));
8830            let w = _mm256_loadu_si256(buf.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
8831            let aw = _mm256_abs_epi8(w);
8832            for (k, xq) in xs.iter().enumerate() {
8833                let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
8834                let d = dpbusd_hsum(aw, _mm256_sign_epi8(x, w));
8835                acc[k] += (d as f32 * sxs[k]) * s;
8836            }
8837        }
8838        acc
8839    }
8840}
8841
8842#[allow(unreachable_code)]
8843fn dot_q4_row_i8(packed: &[u8], scales: &[u8], g0: usize, gpr: usize, xq: &[i8]) -> f32 {
8844    #[cfg(target_arch = "aarch64")]
8845    unsafe {
8846        return dot_q4_row_sdot(packed, scales, g0, gpr, xq);
8847    }
8848    #[cfg(target_arch = "x86_64")]
8849    unsafe {
8850        return dot_q4_row_avx2(packed, scales, g0, gpr, xq);
8851    }
8852    let mut acc = 0f32;
8853    for gi in 0..gpr {
8854        let g = g0 + gi;
8855        let s = f16_to_f32(u16::from_le_bytes([scales[g * 2], scales[g * 2 + 1]]));
8856        let mut d = 0i32;
8857        for (k, &b) in packed[g * 16..(g + 1) * 16].iter().enumerate() {
8858            d += ((b & 0x0F) as i32 - 8) * xq[gi * GROUP_SIZE + k * 2] as i32
8859                + (((b >> 4) & 0x0F) as i32 - 8) * xq[gi * GROUP_SIZE + k * 2 + 1] as i32;
8860        }
8861        acc += d as f32 * s;
8862    }
8863    acc
8864}
8865
8866/// Two-activation q4 row via the A8W8 int8 path (see `dot_q4_row_i8`).
8867#[inline]
8868#[allow(unreachable_code)]
8869fn dot_q4_row_i8_2(
8870    packed: &[u8],
8871    scales: &[u8],
8872    g0: usize,
8873    gpr: usize,
8874    xq1: &[i8],
8875    xq2: &[i8],
8876) -> (f32, f32) {
8877    #[cfg(target_arch = "aarch64")]
8878    unsafe {
8879        return dot_q4_row_sdot2(packed, scales, g0, gpr, xq1, xq2);
8880    }
8881    #[cfg(target_arch = "x86_64")]
8882    unsafe {
8883        return dot_q4_row_avx2_2(packed, scales, g0, gpr, xq1, xq2);
8884    }
8885    (
8886        dot_q4_row_i8(packed, scales, g0, gpr, xq1),
8887        dot_q4_row_i8(packed, scales, g0, gpr, xq2),
8888    )
8889}
8890
8891/// One q4 row range via SDOT (kernel body of `q4matvec`, extracted so
8892/// multi-matrix jobs can drive it for several tensors in one dispatch).
8893#[allow(clippy::too_many_arguments)]
8894fn q4_range_a8w8(
8895    packed: &[u8],
8896    scales: &[u8],
8897    gpr: usize,
8898    cols: usize,
8899    act: &SplitAct,
8900    out: SendMut,
8901    start: usize,
8902    end: usize,
8903) {
8904    for r in start..end {
8905        let mut acc = dot_q4_row_i8(packed, scales, r * gpr, gpr, &act.xq) * act.sx;
8906        // xq is zeroed at outlier slots — add the exact terms.
8907        for &(j, xv) in &act.outliers {
8908            let flat = r * cols + j;
8909            let byte = packed[flat / 2];
8910            let nib = if flat & 1 == 0 {
8911                byte & 0x0F
8912            } else {
8913                byte >> 4
8914            };
8915            let s = f16_to_f32(u16::from_le_bytes([
8916                scales[(flat / GROUP_SIZE) * 2],
8917                scales[(flat / GROUP_SIZE) * 2 + 1],
8918            ]));
8919            acc += ((nib as i32 - 8) as f32) * s * xv;
8920        }
8921        // SAFETY: disjoint row ranges per worker.
8922        unsafe { *out.at(r) = acc };
8923    }
8924}
8925
8926/// Two-input q4 row range via the A8W8 int8 path — kernel body of
8927/// `q4matvec2`, extracted for pair multi-matrix jobs.
8928#[allow(clippy::too_many_arguments)]
8929fn q4_range2_a8w8(
8930    packed: &[u8],
8931    scales: &[u8],
8932    gpr: usize,
8933    cols: usize,
8934    a1: &SplitAct,
8935    a2: &SplitAct,
8936    p1: SendMut,
8937    p2: SendMut,
8938    start: usize,
8939    end: usize,
8940) {
8941    for r in start..end {
8942        let (s1, s2) = dot_q4_row_i8_2(packed, scales, r * gpr, gpr, &a1.xq, &a2.xq);
8943        let mut acc1 = s1 * a1.sx;
8944        let mut acc2 = s2 * a2.sx;
8945        // xq is zeroed at outlier slots — add the exact terms.
8946        let fix = |outliers: &[(usize, f32)], acc: &mut f32| {
8947            for &(j, xv) in outliers {
8948                let flat = r * cols + j;
8949                let byte = packed[flat / 2];
8950                let nib = if flat & 1 == 0 {
8951                    byte & 0x0F
8952                } else {
8953                    byte >> 4
8954                };
8955                let s = f16_to_f32(u16::from_le_bytes([
8956                    scales[(flat / GROUP_SIZE) * 2],
8957                    scales[(flat / GROUP_SIZE) * 2 + 1],
8958                ]));
8959                *acc += ((nib as i32 - 8) as f32) * s * xv;
8960            }
8961        };
8962        fix(&a1.outliers, &mut acc1);
8963        fix(&a2.outliers, &mut acc2);
8964        // SAFETY: disjoint row ranges per worker.
8965        unsafe {
8966            *p1.at(r) = acc1;
8967            *p2.at(r) = acc2;
8968        }
8969    }
8970}
8971
8972/// Exact scalar q4 row range (same extraction, non-SDOT path).
8973fn q4_range_f32(
8974    packed: &[u8],
8975    scales: &[u8],
8976    gpr: usize,
8977    x: &[f32],
8978    out: SendMut,
8979    start: usize,
8980    end: usize,
8981) {
8982    for r in start..end {
8983        let mut acc = 0f32;
8984        for gi in 0..gpr {
8985            let g = r * gpr + gi;
8986            let s = f16_to_f32(u16::from_le_bytes([scales[g * 2], scales[g * 2 + 1]]));
8987            let pk = &packed[g * 16..(g + 1) * 16];
8988            let xg = &x[gi * GROUP_SIZE..(gi + 1) * GROUP_SIZE];
8989            let mut ga = 0f32;
8990            for (k, &b) in pk.iter().enumerate() {
8991                ga += ((b & 0x0F) as f32 - 8.0) * xg[k * 2]
8992                    + (((b >> 4) & 0x0F) as f32 - 8.0) * xg[k * 2 + 1];
8993            }
8994            acc += ga * s;
8995        }
8996        // SAFETY: disjoint row ranges per worker.
8997        unsafe { *out.at(r) = acc };
8998    }
8999}
9000
9001/// Fused two-input q4 matvec: nibbles are unpacked ONCE per group and
9002/// dotted against both activations (was: two full matvecs — double
9003/// weight traffic). Per-lane math matches `q4matvec` exactly.
9004#[allow(clippy::too_many_arguments)]
9005fn q4matvec2(
9006    bytes: &[u8],
9007    x1: &[f32],
9008    x2: &[f32],
9009    rows: usize,
9010    cols: usize,
9011    o1: &mut [f32],
9012    o2: &mut [f32],
9013    pool: Option<&Pool>,
9014) {
9015    debug_assert_eq!(o1.len(), rows);
9016    debug_assert_eq!(o2.len(), rows);
9017    let (packed, scales) = q4_split(bytes, rows, cols);
9018    let gpr = cols / GROUP_SIZE;
9019
9020    if a8w8_enabled() {
9021        let a1 = split_act(x1);
9022        let a2 = split_act(x2);
9023        let p1 = SendMut(o1.as_mut_ptr());
9024        let p2 = SendMut(o2.as_mut_ptr());
9025        let run = move |start: usize, end: usize| {
9026            q4_range2_a8w8(packed, scales, gpr, cols, &a1, &a2, p1, p2, start, end)
9027        };
9028        dispatch_rows(pool, rows, &run);
9029        return;
9030    }
9031
9032    let p1 = SendMut(o1.as_mut_ptr());
9033    let p2 = SendMut(o2.as_mut_ptr());
9034    let run = move |start: usize, end: usize| {
9035        q4_range2_f32(packed, scales, gpr, x1, x2, p1, p2, start, end)
9036    };
9037    dispatch_rows(pool, rows, &run);
9038}
9039
9040/// Two-input exact scalar q4 row range (same extraction).
9041#[allow(clippy::too_many_arguments)]
9042fn q4_range2_f32(
9043    packed: &[u8],
9044    scales: &[u8],
9045    gpr: usize,
9046    x1: &[f32],
9047    x2: &[f32],
9048    p1: SendMut,
9049    p2: SendMut,
9050    start: usize,
9051    end: usize,
9052) {
9053    for r in start..end {
9054        let (mut acc1, mut acc2) = (0f32, 0f32);
9055        for gi in 0..gpr {
9056            let g = r * gpr + gi;
9057            let s = f16_to_f32(u16::from_le_bytes([scales[g * 2], scales[g * 2 + 1]]));
9058            let pk = &packed[g * 16..(g + 1) * 16];
9059            let x1g = &x1[gi * GROUP_SIZE..(gi + 1) * GROUP_SIZE];
9060            let x2g = &x2[gi * GROUP_SIZE..(gi + 1) * GROUP_SIZE];
9061            let (mut g1, mut g2) = (0f32, 0f32);
9062            for (k, &b) in pk.iter().enumerate() {
9063                let wl = (b & 0x0F) as f32 - 8.0;
9064                let wh = ((b >> 4) & 0x0F) as f32 - 8.0;
9065                g1 += wl * x1g[k * 2] + wh * x1g[k * 2 + 1];
9066                g2 += wl * x2g[k * 2] + wh * x2g[k * 2 + 1];
9067            }
9068            acc1 += g1 * s;
9069            acc2 += g2 * s;
9070        }
9071        // SAFETY: disjoint row ranges per worker.
9072        unsafe {
9073            *p1.at(r) = acc1;
9074            *p2.at(r) = acc2;
9075        }
9076    }
9077}
9078
9079thread_local! {
9080    /// Per-worker decoded-row scratch for the batched q4/vbit kernels
9081    /// (centered i8 for SDOT, f32 for the exact/scalar paths).
9082    static ROW_I8: std::cell::RefCell<Vec<u8>> = const { std::cell::RefCell::new(Vec::new()) };
9083    static ROW_F32: std::cell::RefCell<Vec<f32>> = const { std::cell::RefCell::new(Vec::new()) };
9084}
9085
9086/// Batched q4 matmat: each weight row is unpacked from the mmap ONCE
9087/// and dotted against ALL b activations (prefill used to fall back to b
9088/// full matvecs — b× weight traffic and b× nibble decode). Per-position
9089/// math matches `q4matvec` exactly: same group order, same accumulation.
9090/// `out` is row-major [b, rows] like `qmatmat`.
9091#[allow(clippy::too_many_arguments)]
9092fn q4matmat(
9093    bytes: &[u8],
9094    xs_all: &[f32],
9095    b: usize,
9096    rows: usize,
9097    cols: usize,
9098    out: &mut [f32],
9099    pool: Option<&Pool>,
9100) {
9101    debug_assert_eq!(xs_all.len(), b * cols);
9102    debug_assert_eq!(out.len(), b * rows);
9103    let (packed, scales) = q4_split(bytes, rows, cols);
9104    let gpr = cols / GROUP_SIZE;
9105    let gscale = |g: usize| f16_to_f32(u16::from_le_bytes([scales[g * 2], scales[g * 2 + 1]]));
9106
9107    if a8w8_enabled() {
9108        let acts: Vec<SplitAct> = (0..b)
9109            .map(|bi| split_act(&xs_all[bi * cols..(bi + 1) * cols]))
9110            .collect();
9111        let acts = &acts;
9112        let out_addr = SendMut(out.as_mut_ptr());
9113        let run = move |start: usize, end: usize| {
9114            ROW_I8.with(|rb| {
9115                let mut buf = rb.borrow_mut();
9116                buf.resize(cols, 0);
9117                for r in start..end {
9118                    // Unpack the row's nibbles to centered i8 once
9119                    // (element 2k = low nibble, 2k+1 = high — flat order,
9120                    // same as dot_q4_row_sdot's zip).
9121                    for gi in 0..gpr {
9122                        let g = r * gpr + gi;
9123                        for (k, &bt) in packed[g * 16..(g + 1) * 16].iter().enumerate() {
9124                            buf[gi * GROUP_SIZE + k * 2] = ((bt & 0x0F) as i32 - 8) as i8 as u8;
9125                            buf[gi * GROUP_SIZE + k * 2 + 1] =
9126                                (((bt >> 4) & 0x0F) as i32 - 8) as i8 as u8;
9127                        }
9128                    }
9129                    let mut bi = 0usize;
9130                    #[cfg(target_arch = "x86_64")]
9131                    if avx2_enabled() && blocked_enabled() {
9132                        while bi + 4 <= acts.len() {
9133                            let xs = [
9134                                acts[bi].xq.as_slice(),
9135                                acts[bi + 1].xq.as_slice(),
9136                                acts[bi + 2].xq.as_slice(),
9137                                acts[bi + 3].xq.as_slice(),
9138                            ];
9139                            let d = unsafe {
9140                                if vnni_tiles_enabled() {
9141                                    dot_q4b_row_1x4_vnni(&buf, scales, r * gpr, gpr, xs)
9142                                } else {
9143                                    dot_q4b_row_1x4_avx2(&buf, scales, r * gpr, gpr, xs)
9144                                }
9145                            };
9146                            for k in 0..4 {
9147                                let act = &acts[bi + k];
9148                                let mut acc = d[k] * act.sx;
9149                                for &(j, xv) in &act.outliers {
9150                                    acc += (buf[j] as i8) as f32
9151                                        * gscale((r * cols + j) / GROUP_SIZE)
9152                                        * xv;
9153                                }
9154                                // SAFETY: disjoint (bi, r) cells per worker.
9155                                unsafe { *out_addr.at((bi + k) * rows + r) = acc };
9156                            }
9157                            bi += 4;
9158                        }
9159                    }
9160                    while bi < acts.len() {
9161                        let act = &acts[bi];
9162                        let mut acc = 0f32;
9163                        for gi in 0..gpr {
9164                            let d = dot_i8_i8(
9165                                &buf[gi * GROUP_SIZE..(gi + 1) * GROUP_SIZE],
9166                                &act.xq[gi * GROUP_SIZE..(gi + 1) * GROUP_SIZE],
9167                            );
9168                            acc += d as f32 * gscale(r * gpr + gi);
9169                        }
9170                        acc *= act.sx;
9171                        // xq is zeroed at outlier slots — exact terms.
9172                        for &(j, xv) in &act.outliers {
9173                            acc += (buf[j] as i8) as f32 * gscale((r * cols + j) / GROUP_SIZE) * xv;
9174                        }
9175                        // SAFETY: disjoint (bi, r) cells per worker row range.
9176                        unsafe { *out_addr.at(bi * rows + r) = acc };
9177                        bi += 1;
9178                    }
9179                }
9180            })
9181        };
9182        dispatch_rows(pool, rows, &run);
9183        return;
9184    }
9185
9186    let out_addr = SendMut(out.as_mut_ptr());
9187    let run = move |start: usize, end: usize| {
9188        ROW_F32.with(|rb| {
9189            let mut buf = rb.borrow_mut();
9190            buf.resize(cols, 0.0);
9191            for r in start..end {
9192                // Decode raw (nib − 8) values once; scales stay per-group
9193                // so the accumulation order matches q4matvec bit-for-bit.
9194                for gi in 0..gpr {
9195                    let g = r * gpr + gi;
9196                    for (k, &bt) in packed[g * 16..(g + 1) * 16].iter().enumerate() {
9197                        buf[gi * GROUP_SIZE + k * 2] = (bt & 0x0F) as f32 - 8.0;
9198                        buf[gi * GROUP_SIZE + k * 2 + 1] = ((bt >> 4) & 0x0F) as f32 - 8.0;
9199                    }
9200                }
9201                for bi in 0..b {
9202                    let x = &xs_all[bi * cols..(bi + 1) * cols];
9203                    let mut acc = 0f32;
9204                    for gi in 0..gpr {
9205                        let mut ga = 0f32;
9206                        // Pairwise (lo + hi) addition, matching
9207                        // q4matvec's `ga += lo·x + hi·x` shape exactly —
9208                        // a flat one-per-element loop rounds differently
9209                        // and broke bit-parity on the scalar (x86) path.
9210                        for k in 0..GROUP_SIZE / 2 {
9211                            let e = gi * GROUP_SIZE + k * 2;
9212                            ga += buf[e] * x[e] + buf[e + 1] * x[e + 1];
9213                        }
9214                        acc += ga * gscale(r * gpr + gi);
9215                    }
9216                    // SAFETY: disjoint (bi, r) cells per worker row range.
9217                    unsafe { *out_addr.at(bi * rows + r) = acc };
9218                }
9219            }
9220        })
9221    };
9222    dispatch_rows(pool, rows, &run);
9223}
9224
9225/// Batched vbit matmat: each variable-bit row is decoded from the mmap
9226/// ONCE for the whole microbatch. Same per-position math as
9227/// `vbitmatvec` (SDOT A8W8 with exact outliers / exact f32 for b=8 rows
9228/// and the scalar path).
9229#[allow(clippy::too_many_arguments)]
9230fn vbitmatmat(
9231    bytes: &[u8],
9232    offsets: &[usize],
9233    xs_all: &[f32],
9234    b: usize,
9235    rows: usize,
9236    cols: usize,
9237    out: &mut [f32],
9238    pool: Option<&Pool>,
9239) {
9240    debug_assert_eq!(xs_all.len(), b * cols);
9241    debug_assert_eq!(out.len(), b * rows);
9242    debug_assert_eq!(offsets.len(), rows + 1);
9243    let ng = cols / GROUP_SIZE;
9244    let bits = &bytes[..rows];
9245    let sc_off = rows;
9246    let gscale = |r: usize, g: usize| {
9247        let so = (r * ng + g) * 2;
9248        f16_to_f32(u16::from_le_bytes([
9249            bytes[sc_off + so],
9250            bytes[sc_off + so + 1],
9251        ]))
9252    };
9253
9254    // Decode row r's raw (u − L) values into `dst` (f32, unscaled).
9255    let decode_f32 = |r: usize, dst: &mut [f32]| {
9256        let bw = bits[r] as usize;
9257        let l = ((1i32 << (bw - 1)) - 1) as f32;
9258        let data = &bytes[offsets[r]..offsets[r + 1]];
9259        let (mut acc, mut nbits, mut idx) = (0u64, 0usize, 0usize);
9260        for d in dst.iter_mut() {
9261            while nbits < bw {
9262                acc = (acc << 8) | data[idx] as u64;
9263                idx += 1;
9264                nbits += 8;
9265            }
9266            let u = ((acc >> (nbits - bw)) & ((1u64 << bw) - 1)) as f32;
9267            nbits -= bw;
9268            *d = u - l;
9269        }
9270    };
9271
9272    if a8w8_enabled() {
9273        let acts: Vec<SplitAct> = (0..b)
9274            .map(|bi| split_act(&xs_all[bi * cols..(bi + 1) * cols]))
9275            .collect();
9276        let acts = &acts;
9277        let out_addr = SendMut(out.as_mut_ptr());
9278        let run = move |start: usize, end: usize| {
9279            for r in start..end {
9280                let bw = bits[r] as usize;
9281                if bw == 8 {
9282                    // u−L reaches 128 → no i8 path; decode once, exact
9283                    // f32 dots for every position (same as vbitmatvec).
9284                    ROW_F32.with(|rb| {
9285                        let mut buf = rb.borrow_mut();
9286                        buf.resize(cols, 0.0);
9287                        decode_f32(r, &mut buf);
9288                        for bi in 0..b {
9289                            let x = &xs_all[bi * cols..(bi + 1) * cols];
9290                            let mut dot = 0f32;
9291                            for g in 0..ng {
9292                                let mut gd = 0f32;
9293                                for k in 0..GROUP_SIZE {
9294                                    gd += buf[g * GROUP_SIZE + k] * x[g * GROUP_SIZE + k];
9295                                }
9296                                dot += gd * gscale(r, g);
9297                            }
9298                            // SAFETY: disjoint (bi, r) cells per worker range.
9299                            unsafe { *out_addr.at(bi * rows + r) = dot };
9300                        }
9301                    });
9302                    continue;
9303                }
9304                let l = (1i32 << (bw - 1)) - 1;
9305                let data = &bytes[offsets[r]..offsets[r + 1]];
9306                ROW_I8.with(|rb| {
9307                    let mut buf = rb.borrow_mut();
9308                    buf.resize(cols, 0);
9309                    #[inline(always)]
9310                    fn fill<const B: usize>(data: &[u8], l: i32, buf: &mut [u8]) {
9311                        for (blk, chunk) in buf.chunks_exact_mut(8).enumerate() {
9312                            let u = unpack8::<B>(&data[blk * B..]);
9313                            for k in 0..8 {
9314                                chunk[k] = (u[k] - l) as i8 as u8;
9315                            }
9316                        }
9317                    }
9318                    match bw {
9319                        3 => fill::<3>(data, l, &mut buf),
9320                        4 => vbit_fill4(data, &mut buf),
9321                        5 => fill::<5>(data, l, &mut buf),
9322                        6 => fill::<6>(data, l, &mut buf),
9323                        _ => unreachable!("vbit bit-width {bw} (validated at load)"),
9324                    }
9325                    let mut bi = 0usize;
9326                    // The vbit scale table shares q4_block's layout
9327                    // (contiguous f16 per (row·ng + g)), so the same
9328                    // blocked 1×4 kernel serves the decoded row.
9329                    #[cfg(target_arch = "x86_64")]
9330                    if avx2_enabled() && blocked_enabled() {
9331                        while bi + 4 <= acts.len() {
9332                            let xs = [
9333                                acts[bi].xq.as_slice(),
9334                                acts[bi + 1].xq.as_slice(),
9335                                acts[bi + 2].xq.as_slice(),
9336                                acts[bi + 3].xq.as_slice(),
9337                            ];
9338                            let sxs = [
9339                                acts[bi].sx,
9340                                acts[bi + 1].sx,
9341                                acts[bi + 2].sx,
9342                                acts[bi + 3].sx,
9343                            ];
9344                            let d = unsafe {
9345                                if vnni_tiles_enabled() {
9346                                    dot_q4b_row_1x4_sx_vnni(
9347                                        &buf,
9348                                        &bytes[sc_off..],
9349                                        r * ng,
9350                                        ng,
9351                                        xs,
9352                                        sxs,
9353                                    )
9354                                } else {
9355                                    dot_q4b_row_1x4_sx_avx2(
9356                                        &buf,
9357                                        &bytes[sc_off..],
9358                                        r * ng,
9359                                        ng,
9360                                        xs,
9361                                        sxs,
9362                                    )
9363                                }
9364                            };
9365                            for k in 0..4 {
9366                                let act = &acts[bi + k];
9367                                let mut dot = d[k];
9368                                for &(j, xv) in &act.outliers {
9369                                    dot += (buf[j] as i8) as f32 * gscale(r, j / GROUP_SIZE) * xv;
9370                                }
9371                                // SAFETY: disjoint (bi, r) cells per worker.
9372                                unsafe { *out_addr.at((bi + k) * rows + r) = dot };
9373                            }
9374                            bi += 4;
9375                        }
9376                    }
9377                    while bi < acts.len() {
9378                        let act = &acts[bi];
9379                        let mut dot = 0f32;
9380                        for g in 0..ng {
9381                            let d = dot_i8_i8(
9382                                &buf[g * GROUP_SIZE..(g + 1) * GROUP_SIZE],
9383                                &act.xq[g * GROUP_SIZE..(g + 1) * GROUP_SIZE],
9384                            ) as f32
9385                                * act.sx;
9386                            dot += d * gscale(r, g);
9387                        }
9388                        for &(j, xv) in &act.outliers {
9389                            dot += (buf[j] as i8) as f32 * gscale(r, j / GROUP_SIZE) * xv;
9390                        }
9391                        // SAFETY: disjoint (bi, r) cells per worker range.
9392                        unsafe { *out_addr.at(bi * rows + r) = dot };
9393                        bi += 1;
9394                    }
9395                });
9396            }
9397        };
9398        dispatch_rows(pool, rows, &run);
9399        return;
9400    }
9401
9402    let out_addr = SendMut(out.as_mut_ptr());
9403    let run = move |start: usize, end: usize| {
9404        ROW_F32.with(|rb| {
9405            let mut buf = rb.borrow_mut();
9406            buf.resize(cols, 0.0);
9407            for r in start..end {
9408                decode_f32(r, &mut buf);
9409                for bi in 0..b {
9410                    let x = &xs_all[bi * cols..(bi + 1) * cols];
9411                    let mut dot = 0f32;
9412                    for g in 0..ng {
9413                        let mut gd = 0f32;
9414                        for k in 0..GROUP_SIZE {
9415                            gd += buf[g * GROUP_SIZE + k] * x[g * GROUP_SIZE + k];
9416                        }
9417                        dot += gd * gscale(r, g);
9418                    }
9419                    // SAFETY: disjoint (bi, r) cells per worker range.
9420                    unsafe { *out_addr.at(bi * rows + r) = dot };
9421                }
9422            }
9423        })
9424    };
9425    dispatch_rows(pool, rows, &run);
9426}
9427
9428/// Build a GPU batch job for a q8-family mapped tensor (primary
9429/// shard): prescaled input + directory coordinates. None → not
9430/// GPU-eligible, caller stays on the CPU.
9431pub(crate) fn gpu_batch_job<'a>(
9432    t: &'a QTensor,
9433    x: &[f32],
9434) -> Option<(std::sync::Arc<CmfModel>, crate::gpu::BatchJob<'a>)> {
9435    match t {
9436        QTensor::Mapped {
9437            model,
9438            idx,
9439            dtype: dt @ (TensorDtype::Q8Row | TensorDtype::Q8_2f),
9440            rows,
9441            cols,
9442            row_scale,
9443            col_field,
9444            ..
9445        } => Some((
9446            model.clone(),
9447            crate::gpu::BatchJob {
9448                idx: *idx,
9449                rows: *rows,
9450                cols: *cols,
9451                row_scale,
9452                xs: prescale(x, col_field, *dt).into_owned(),
9453                layout: crate::gpu::BatchLayout::Q8,
9454            },
9455        )),
9456        // q1: raw f32 activations, tile-embedded scales.
9457        QTensor::Mapped {
9458            model,
9459            idx,
9460            dtype: TensorDtype::Q1,
9461            rows,
9462            cols,
9463            ..
9464        } => Some((
9465            model.clone(),
9466            crate::gpu::BatchJob {
9467                idx: *idx,
9468                rows: *rows,
9469                cols: *cols,
9470                row_scale: &[],
9471                xs: x.to_vec(),
9472                layout: crate::gpu::BatchLayout::Q1,
9473            },
9474        )),
9475        // q4_tiled / q4tp: raw f32 activations; the scales live in the
9476        // payload (inline tiles / row ladder), so row_scale stays empty.
9477        // The GDN projection batch already runs these layouts on Metal —
9478        // this arm lets the attention QKV batch reach the same kernels.
9479        QTensor::Mapped {
9480            model,
9481            idx,
9482            dtype: dt @ (TensorDtype::Q4Tiled | TensorDtype::Q4TiledP),
9483            rows,
9484            cols,
9485            ..
9486        } => Some((
9487            model.clone(),
9488            crate::gpu::BatchJob {
9489                idx: *idx,
9490                rows: *rows,
9491                cols: *cols,
9492                row_scale: &[],
9493                xs: x.to_vec(),
9494                layout: if *dt == TensorDtype::Q4Tiled {
9495                    crate::gpu::BatchLayout::Q4t
9496                } else {
9497                    crate::gpu::BatchLayout::Q4tp
9498                },
9499            },
9500        )),
9501        _ => None,
9502    }
9503}
9504
9505thread_local! {
9506    static PRESCALE_BUF1: std::cell::RefCell<Vec<f32>> = const { std::cell::RefCell::new(Vec::new()) };
9507    static PRESCALE_BUF2: std::cell::RefCell<Vec<f32>> = const { std::cell::RefCell::new(Vec::new()) };
9508}
9509
9510pub(crate) fn prescale<'a>(
9511    x: &'a [f32],
9512    col_field: &[f32],
9513    dtype: TensorDtype,
9514) -> std::borrow::Cow<'a, [f32]> {
9515    if dtype == TensorDtype::Q8_2f {
9516        x.iter().zip(col_field).map(|(a, c)| a * c).collect()
9517    } else {
9518        std::borrow::Cow::Borrowed(x)
9519    }
9520}
9521
9522/// θ col-field fold for q8_2f activations. Borrowed pass-through for
9523/// every other dtype, using thread-local buffers to eliminate per-matvec allocations.
9524pub(crate) fn prescale_with<R, F: FnOnce(&[f32]) -> R>(
9525    x: &[f32],
9526    col_field: &[f32],
9527    dtype: TensorDtype,
9528    buf_id: u8,
9529    f: F,
9530) -> R {
9531    if dtype == TensorDtype::Q8_2f {
9532        if buf_id == 1 {
9533            PRESCALE_BUF1.with(|b| {
9534                let mut buf = b.borrow_mut();
9535                buf.clear();
9536                buf.extend(x.iter().zip(col_field).map(|(a, c)| a * c));
9537                f(&buf)
9538            })
9539        } else {
9540            PRESCALE_BUF2.with(|b| {
9541                let mut buf = b.borrow_mut();
9542                buf.clear();
9543                buf.extend(x.iter().zip(col_field).map(|(a, c)| a * c));
9544                f(&buf)
9545            })
9546        }
9547    } else {
9548        f(x)
9549    }
9550}
9551
9552// ───────────────────── x86-64 AVX2 kernels (roadmap этап 2) ─────────────────────
9553
9554/// AVX2+FMA available? Default ON when the CPU supports both;
9555/// `CMF_AVX2=0` disables (falls back to the autovectorized loops).
9556#[cfg(target_arch = "x86_64")]
9557pub(crate) fn avx2_enabled() -> bool {
9558    use std::sync::OnceLock;
9559    static ON: OnceLock<bool> = OnceLock::new();
9560    *ON.get_or_init(|| {
9561        std::env::var("CMF_AVX2").map(|v| v != "0").unwrap_or(true)
9562            && std::arch::is_x86_feature_detected!("avx2")
9563            && std::arch::is_x86_feature_detected!("fma")
9564    })
9565}
9566
9567/// AVX2 A8W8 allowed? The quantized-activation contract is switched by
9568/// the SAME env as the ARM SDOT path: `CMF_SDOT=0` keeps exact kernels
9569/// (the golden-parity exact gate relies on it) — AVX2 f32 kernels stay
9570/// active either way, they are exact (regrouped sums only).
9571#[cfg(target_arch = "x86_64")]
9572fn avx2_a8w8_enabled() -> bool {
9573    if FLOAT_ACTIVATIONS.get() {
9574        return false;
9575    }
9576    use std::sync::OnceLock;
9577    static ON: OnceLock<bool> = OnceLock::new();
9578    *ON.get_or_init(|| {
9579        avx2_enabled() && std::env::var("CMF_SDOT").map(|v| v != "0").unwrap_or(true)
9580    })
9581}
9582
9583thread_local! {
9584    static FULL_GPU_Q8: std::cell::Cell<bool> = const { std::cell::Cell::new(false) };
9585}
9586
9587/// Match the graph's full-device q8 projection precision on MiMo's host
9588/// tail. Does not enable the GPU or bypass a CPU-only/device-refusal gate.
9589pub(crate) fn enter_full_gpu_q8_scope() -> impl Drop {
9590    struct Restore(bool, std::marker::PhantomData<std::rc::Rc<()>>);
9591    impl Drop for Restore {
9592        fn drop(&mut self) {
9593            FULL_GPU_Q8.set(self.0);
9594        }
9595    }
9596    Restore(FULL_GPU_Q8.replace(true), std::marker::PhantomData)
9597}
9598
9599// Dynamic MiMo experts must not change activation precision when a cache
9600// fill moves them from CPU to GPU. Thread-local: only the cold-expert
9601// dispatch selects float kernels; concurrent pipelines keep their policy.
9602thread_local! {
9603    static FLOAT_ACTIVATIONS: std::cell::Cell<bool> = const { std::cell::Cell::new(false) };
9604}
9605
9606pub(crate) fn float_activations_scope<R>(f: impl FnOnce() -> R) -> R {
9607    struct Restore(bool);
9608    impl Drop for Restore {
9609        fn drop(&mut self) {
9610            FLOAT_ACTIVATIONS.set(self.0);
9611        }
9612    }
9613    let _restore = Restore(FLOAT_ACTIVATIONS.replace(true));
9614    f()
9615}
9616
9617/// Row-exact batching: while set, the x86 batched kernels (`qmatmat`,
9618/// `q4tp_matmat`) compute every (weight row, token) cell with the
9619/// single-token kernel instead of the blocked 2×4 / 1×4 / 1×8 tiles, so a
9620/// token's result does not depend on the batch it rides in and equals its
9621/// matvec. On ARM `q4tp_matmat` keeps its 1×4 tile but in the matvec's
9622/// reduction order, and no q4tp batch takes the AMX or device GEMM. The
9623/// MiMo speculative verify holds it (`row_exact_scope`) — its accepted
9624/// rows must be the rows plain decode would have produced.
9625// Shared with pool workers, so overlapping requests must keep the mode
9626// enabled until the LAST scope leaves. Saving/restoring a global bool is
9627// incorrect when two threads enter and leave in a non-LIFO order.
9628static ROW_EXACT: std::sync::atomic::AtomicUsize = std::sync::atomic::AtomicUsize::new(0);
9629
9630pub(crate) fn row_exact() -> bool {
9631    ROW_EXACT.load(std::sync::atomic::Ordering::Acquire) != 0
9632}
9633
9634fn counted_row_exact_scope<R>(active: &std::sync::atomic::AtomicUsize, f: impl FnOnce() -> R) -> R {
9635    struct Restore<'a>(&'a std::sync::atomic::AtomicUsize);
9636    impl Drop for Restore<'_> {
9637        fn drop(&mut self) {
9638            self.0.fetch_sub(1, std::sync::atomic::Ordering::AcqRel);
9639        }
9640    }
9641    active.fetch_add(1, std::sync::atomic::Ordering::AcqRel);
9642    let _restore = Restore(active);
9643    f()
9644}
9645
9646/// Run `f` with row-exact batching on (also released on unwind).
9647pub(crate) fn row_exact_scope<R>(f: impl FnOnce() -> R) -> R {
9648    counted_row_exact_scope(&ROW_EXACT, f)
9649}
9650
9651/// A8W8 quantized-activation path available on THIS machine? One
9652/// switch across architectures: ARM dotprod (CMF_SDOT) or x86 AVX2
9653/// (CMF_AVX2 + the same CMF_SDOT exact-contract override).
9654#[inline]
9655pub(crate) fn a8w8_enabled() -> bool {
9656    #[cfg(target_arch = "aarch64")]
9657    {
9658        sdot_enabled()
9659    }
9660    #[cfg(target_arch = "x86_64")]
9661    {
9662        avx2_a8w8_enabled()
9663    }
9664    #[cfg(not(any(target_arch = "aarch64", target_arch = "x86_64")))]
9665    {
9666        false
9667    }
9668}
9669
9670/// int8·int8 dot dispatch: SDOT on ARM; AVX-512 VNNI (vpdpbusd) or AVX2
9671/// maddubs on x86. Callers are gated by `a8w8_enabled()`.
9672#[inline]
9673#[allow(unreachable_code)]
9674fn dot_i8_i8(w: &[u8], xq: &[i8]) -> i32 {
9675    #[cfg(target_arch = "aarch64")]
9676    unsafe {
9677        return dot_i8_sdot(w, xq);
9678    }
9679    #[cfg(target_arch = "x86_64")]
9680    unsafe {
9681        if avx512vnni_enabled() {
9682            return dot_i8_i8_vnni(w, xq);
9683        }
9684        return dot_i8_i8_avx2(w, xq);
9685    }
9686    w.iter()
9687        .zip(xq)
9688        .map(|(&a, &b)| (a as i8) as i32 * b as i32)
9689        .sum()
9690}
9691
9692/// AVX-512 VNNI available? (F+BW+VL+VNNI; `CMF_AVX512=0` falls back to
9693/// AVX2.) VL matters: short 32-byte groups (q4/vbit) ride the 256-bit
9694/// `vpdpbusd` encoding.
9695#[cfg(target_arch = "x86_64")]
9696fn avx512vnni_enabled() -> bool {
9697    use std::sync::OnceLock;
9698    static ON: OnceLock<bool> = OnceLock::new();
9699    *ON.get_or_init(|| {
9700        std::env::var("CMF_AVX512")
9701            .map(|v| v != "0")
9702            .unwrap_or(true)
9703            && std::arch::is_x86_feature_detected!("avx512f")
9704            && std::arch::is_x86_feature_detected!("avx512bw")
9705            && std::arch::is_x86_feature_detected!("avx512vl")
9706            && std::arch::is_x86_feature_detected!("avx512vnni")
9707    })
9708}
9709
9710/// Grouped-codec VNNI arms (the q4t/q4b/q1/q1t tile kernels): default
9711/// ON where AVX-512 VNNI exists (`CMF_VNNI_TILES=0` opt-out). Measured
9712/// on Ryzen 7950X (Zen4, 3 alternating process pairs, blocked GEMM
9713/// 4864×896 b=256): q4t 63→68 GF/s (+8%), q1 53→56 (+6%), q4b 72→75
9714/// (+4%) — consistent, no leg regressed. The tile kernels keep a
9715/// horizontal reduce per 32-weight group, so the `vpdpbusd` saving is
9716/// smaller than the long-dot q8 win (+13%), but it is real and free.
9717#[cfg(target_arch = "x86_64")]
9718fn vnni_tiles_enabled() -> bool {
9719    use std::sync::OnceLock;
9720    static ON: OnceLock<bool> = OnceLock::new();
9721    *ON.get_or_init(|| {
9722        std::env::var("CMF_VNNI_TILES")
9723            .map(|v| v != "0")
9724            .unwrap_or(true)
9725            && avx512vnni_enabled()
9726    })
9727}
9728
9729/// One 256-bit u8×i8 dot → i32 via `vpdpbusd` into a fresh accumulator
9730/// plus the same horizontal reduce the AVX2 kernels use. Products are
9731/// bounded (|w| ≤ 8 or ≤ 1), so maddubs never saturated — the i32 sum
9732/// is bit-identical to the maddubs+madd pair it replaces.
9733#[cfg(target_arch = "x86_64")]
9734#[target_feature(enable = "avx2,avx512f,avx512bw,avx512vl,avx512vnni")]
9735#[inline]
9736unsafe fn dpbusd_hsum(aw: core::arch::x86_64::__m256i, xs: core::arch::x86_64::__m256i) -> i32 {
9737    // SAFETY: pure register math.
9738    unsafe {
9739        use core::arch::x86_64::*;
9740        let d = _mm256_dpbusd_epi32(_mm256_setzero_si256(), aw, xs);
9741        let hi128 = _mm256_extracti128_si256::<1>(d);
9742        let s128 = _mm_add_epi32(_mm256_castsi256_si128(d), hi128);
9743        let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
9744        let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
9745        _mm_cvtsi128_si32(s32)
9746    }
9747}
9748
9749/// int8·int8 via AVX-512 VNNI: `vpdpbusd` fuses the maddubs+madd+add
9750/// triple into one u8×i8 dot-accumulate. AVX-512 has no vpsignb, so the
9751/// |w|·sign(x,w) trick becomes |w| × (x negated where w<0) via a mask
9752/// subtract — w==0 lanes contribute 0 through |w|=0 either way.
9753#[cfg(target_arch = "x86_64")]
9754#[target_feature(enable = "avx2,avx512f,avx512bw,avx512vl,avx512vnni")]
9755unsafe fn dot_i8_i8_vnni(w: &[u8], xq: &[i8]) -> i32 {
9756    // SAFETY: callers uphold slice-length contracts (see call sites).
9757    unsafe {
9758        use core::arch::x86_64::*;
9759        let n = w.len();
9760        let mut j = 0usize;
9761        let mut total: i32;
9762        // 4 independent accumulators: vpdpbusd is its own loop-carried
9763        // dependency (~5-cycle latency) — a single-acc loop runs
9764        // latency-bound and LOSES to the AVX2 maddubs kernel, measured
9765        // on Granite Rapids.
9766        {
9767            #[inline(always)]
9768            unsafe fn step(
9769                w: *const u8,
9770                x: *const i8,
9771                acc: core::arch::x86_64::__m512i,
9772            ) -> core::arch::x86_64::__m512i {
9773                unsafe {
9774                    use core::arch::x86_64::*;
9775                    let wv = _mm512_loadu_si512(w as *const _);
9776                    let xv = _mm512_loadu_si512(x as *const _);
9777                    let aw = _mm512_abs_epi8(wv);
9778                    let neg = _mm512_movepi8_mask(wv);
9779                    let sx = _mm512_mask_sub_epi8(xv, neg, _mm512_setzero_si512(), xv);
9780                    _mm512_dpbusd_epi32(acc, aw, sx)
9781                }
9782            }
9783            let (mut a0, mut a1, mut a2, mut a3) = (
9784                _mm512_setzero_si512(),
9785                _mm512_setzero_si512(),
9786                _mm512_setzero_si512(),
9787                _mm512_setzero_si512(),
9788            );
9789            while j + 256 <= n {
9790                a0 = step(w.as_ptr().add(j), xq.as_ptr().add(j), a0);
9791                a1 = step(w.as_ptr().add(j + 64), xq.as_ptr().add(j + 64), a1);
9792                a2 = step(w.as_ptr().add(j + 128), xq.as_ptr().add(j + 128), a2);
9793                a3 = step(w.as_ptr().add(j + 192), xq.as_ptr().add(j + 192), a3);
9794                j += 256;
9795            }
9796            while j + 64 <= n {
9797                a0 = step(w.as_ptr().add(j), xq.as_ptr().add(j), a0);
9798                j += 64;
9799            }
9800            let s01 = _mm512_add_epi32(a0, a1);
9801            let s23 = _mm512_add_epi32(a2, a3);
9802            total = _mm512_reduce_add_epi32(_mm512_add_epi32(s01, s23));
9803        }
9804        // 32-wide (q4/vbit groups are exactly 32 bytes).
9805        if j + 32 <= n {
9806            let wv = _mm256_loadu_si256(w.as_ptr().add(j) as *const __m256i);
9807            let xv = _mm256_loadu_si256(xq.as_ptr().add(j) as *const __m256i);
9808            let d = _mm256_dpbusd_epi32(
9809                _mm256_setzero_si256(),
9810                _mm256_abs_epi8(wv),
9811                _mm256_sign_epi8(xv, wv),
9812            );
9813            let hi128 = _mm256_extracti128_si256::<1>(d);
9814            let s128 = _mm_add_epi32(_mm256_castsi256_si128(d), hi128);
9815            let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
9816            let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
9817            total += _mm_cvtsi128_si32(s32);
9818            j += 32;
9819        }
9820        while j < n {
9821            total += (w[j] as i8) as i32 * xq[j] as i32;
9822            j += 1;
9823        }
9824        total
9825    }
9826}
9827
9828/// i8 row · f32 x via AVX2/FMA (x86 mirror of `dot_i8_f32_neon`).
9829#[cfg(target_arch = "x86_64")]
9830#[target_feature(enable = "avx2,fma")]
9831unsafe fn dot_i8_f32_avx2(w: &[u8], x: &[f32]) -> f32 {
9832    // SAFETY: callers uphold slice-length contracts (see call sites).
9833    unsafe {
9834        use core::arch::x86_64::*;
9835        let n = x.len();
9836        let wp = w.as_ptr();
9837        let xp = x.as_ptr();
9838        let (mut a0, mut a1) = (_mm256_setzero_ps(), _mm256_setzero_ps());
9839        let mut j = 0usize;
9840        while j + 16 <= n {
9841            let wb = _mm_loadu_si128(wp.add(j) as *const __m128i);
9842            let lo = _mm256_cvtepi8_epi32(wb);
9843            let hi = _mm256_cvtepi8_epi32(_mm_srli_si128::<8>(wb));
9844            a0 = _mm256_fmadd_ps(_mm256_cvtepi32_ps(lo), _mm256_loadu_ps(xp.add(j)), a0);
9845            a1 = _mm256_fmadd_ps(_mm256_cvtepi32_ps(hi), _mm256_loadu_ps(xp.add(j + 8)), a1);
9846            j += 16;
9847        }
9848        let acc = _mm256_add_ps(a0, a1);
9849        let hi128 = _mm256_extractf128_ps::<1>(acc);
9850        let s128 = _mm_add_ps(_mm256_castps256_ps128(acc), hi128);
9851        let s64 = _mm_add_ps(s128, _mm_movehl_ps(s128, s128));
9852        let s32 = _mm_add_ss(s64, _mm_shuffle_ps::<1>(s64, s64));
9853        let mut sum = _mm_cvtss_f32(s32);
9854        while j < n {
9855            sum += (*wp.add(j) as i8) as f32 * *xp.add(j);
9856            j += 1;
9857        }
9858        sum
9859    }
9860}
9861
9862/// int8(weight)·int8(activation) → i32 via AVX2 maddubs — the x86
9863/// analogue of the SDOT A8W8 path. `maddubs` takes u8×i8, so the
9864/// standard sign trick applies: |w| × sign(x, w) ≡ w × x per lane.
9865/// Pair saturation is safe: |w|≤128, |x|≤127 → 2·128·127 < 32767.
9866#[cfg(target_arch = "x86_64")]
9867#[target_feature(enable = "avx2")]
9868unsafe fn dot_i8_i8_avx2(w: &[u8], xq: &[i8]) -> i32 {
9869    // SAFETY: callers uphold slice-length contracts (see call sites).
9870    unsafe {
9871        use core::arch::x86_64::*;
9872        let n = w.len();
9873        let ones = _mm256_set1_epi16(1);
9874        let mut acc = _mm256_setzero_si256();
9875        let mut j = 0usize;
9876        while j + 32 <= n {
9877            let wv = _mm256_loadu_si256(w.as_ptr().add(j) as *const __m256i);
9878            let xv = _mm256_loadu_si256(xq.as_ptr().add(j) as *const __m256i);
9879            let p16 = _mm256_maddubs_epi16(_mm256_abs_epi8(wv), _mm256_sign_epi8(xv, wv));
9880            acc = _mm256_add_epi32(acc, _mm256_madd_epi16(p16, ones));
9881            j += 32;
9882        }
9883        let hi128 = _mm256_extracti128_si256::<1>(acc);
9884        let s128 = _mm_add_epi32(_mm256_castsi256_si128(acc), hi128);
9885        let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
9886        let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
9887        let mut s = _mm_cvtsi128_si32(s32);
9888        while j < n {
9889            s += (w[j] as i8) as i32 * xq[j] as i32;
9890            j += 1;
9891        }
9892        s
9893    }
9894}
9895
9896/// smmla 2×4: one instruction covers a 2-row × 2-activation × 8-deep
9897/// tile (32 MACs vs sdot's 16) — the weight pair loads once per 8-k
9898/// slice as a combined 2×8 register and meets two activation pairs.
9899#[cfg(target_arch = "aarch64")]
9900#[target_feature(enable = "neon,i8mm")]
9901unsafe fn dot_i8_smmla_2x4(w0: &[u8], w1: &[u8], xs: [&[i8]; 4]) -> [[i32; 4]; 2] {
9902    // SAFETY: callers uphold slice-length contracts.
9903    unsafe {
9904        use core::arch::aarch64::*;
9905        use core::arch::asm;
9906        let n = w0.len();
9907        let w0p = w0.as_ptr() as *const i8;
9908        let w1p = w1.as_ptr() as *const i8;
9909        // acc01 holds [c(r0,x0) c(r0,x1) c(r1,x0) c(r1,x1)]; acc23 the
9910        // same for x2/x3.
9911        let mut acc01 = vdupq_n_s32(0);
9912        let mut acc23 = vdupq_n_s32(0);
9913        let mut i = 0usize;
9914        while i + 8 <= n {
9915            let wa = vcombine_s8(vld1_s8(w0p.add(i)), vld1_s8(w1p.add(i)));
9916            let xb01 = vcombine_s8(
9917                vld1_s8(xs[0].as_ptr().add(i)),
9918                vld1_s8(xs[1].as_ptr().add(i)),
9919            );
9920            let xb23 = vcombine_s8(
9921                vld1_s8(xs[2].as_ptr().add(i)),
9922                vld1_s8(xs[3].as_ptr().add(i)),
9923            );
9924            asm!(
9925                "smmla {a01:v}.4s, {w:v}.16b, {x01:v}.16b",
9926                "smmla {a23:v}.4s, {w:v}.16b, {x23:v}.16b",
9927                a01 = inout(vreg) acc01, a23 = inout(vreg) acc23,
9928                w = in(vreg) wa, x01 = in(vreg) xb01, x23 = in(vreg) xb23,
9929                options(pure, nomem, nostack),
9930            );
9931            i += 8;
9932        }
9933        let mut out = [[0i32; 4]; 2];
9934        let a01: [i32; 4] = core::mem::transmute(acc01);
9935        let a23: [i32; 4] = core::mem::transmute(acc23);
9936        out[0][0] = a01[0];
9937        out[0][1] = a01[1];
9938        out[1][0] = a01[2];
9939        out[1][1] = a01[3];
9940        out[0][2] = a23[0];
9941        out[0][3] = a23[1];
9942        out[1][2] = a23[2];
9943        out[1][3] = a23[3];
9944        if i < n {
9945            for (k, x) in xs.iter().enumerate() {
9946                for j in i..n {
9947                    out[0][k] += (w0[j] as i8) as i32 * x[j] as i32;
9948                    out[1][k] += (w1[j] as i8) as i32 * x[j] as i32;
9949                }
9950            }
9951        }
9952        out
9953    }
9954}
9955
9956/// ARM twin of the x86 blocked prefill GEMM: two weight rows stay in
9957/// registers across four activation streams, eight sdot accumulators.
9958/// (The per-row form re-read each W row once per activation.)
9959#[cfg(target_arch = "aarch64")]
9960#[target_feature(enable = "neon,dotprod")]
9961unsafe fn dot_i8_sdot_2x4(w0: &[u8], w1: &[u8], xs: [&[i8]; 4]) -> [[i32; 4]; 2] {
9962    // SAFETY: callers uphold slice-length contracts.
9963    unsafe {
9964        use core::arch::aarch64::*;
9965        use core::arch::asm;
9966        let n = w0.len();
9967        let w0p = w0.as_ptr() as *const i8;
9968        let w1p = w1.as_ptr() as *const i8;
9969        let mut acc = [[vdupq_n_s32(0); 4]; 2];
9970        let mut i = 0usize;
9971        while i + 16 <= n {
9972            let wv0 = vld1q_s8(w0p.add(i));
9973            let wv1 = vld1q_s8(w1p.add(i));
9974            for (k, x) in xs.iter().enumerate() {
9975                let xv = vld1q_s8(x.as_ptr().add(i));
9976                let (mut a0, mut a1) = (acc[0][k], acc[1][k]);
9977                asm!(
9978                    "sdot {a0:v}.4s, {w0:v}.16b, {x:v}.16b",
9979                    "sdot {a1:v}.4s, {w1:v}.16b, {x:v}.16b",
9980                    a0 = inout(vreg) a0, a1 = inout(vreg) a1,
9981                    w0 = in(vreg) wv0, w1 = in(vreg) wv1, x = in(vreg) xv,
9982                    options(pure, nomem, nostack),
9983                );
9984                acc[0][k] = a0;
9985                acc[1][k] = a1;
9986            }
9987            i += 16;
9988        }
9989        let mut out = [[0i32; 4]; 2];
9990        for r in 0..2 {
9991            for k in 0..4 {
9992                out[r][k] = vaddvq_s32(acc[r][k]);
9993            }
9994        }
9995        if i < n {
9996            for (k, x) in xs.iter().enumerate() {
9997                for j in i..n {
9998                    out[0][k] += (w0[j] as i8) as i32 * x[j] as i32;
9999                    out[1][k] += (w1[j] as i8) as i32 * x[j] as i32;
10000                }
10001            }
10002        }
10003        out
10004    }
10005}
10006
10007/// Blocked 2 weight rows × 4 activations for the prefill GEMM
10008/// (roadmap P0: packed panels + multi-row accumulators). The two rows'
10009/// abs() live in registers across all four activation streams; the
10010/// sign-fixup is recomputed per pair (the price of the maddubs trick).
10011/// Returns raw i8·i8 dots; the caller applies scales and outliers.
10012#[cfg(target_arch = "x86_64")]
10013#[target_feature(enable = "avx2")]
10014unsafe fn dot_i8_i8_avx2_2x4(w0: &[u8], w1: &[u8], xs: [&[i8]; 4]) -> [[i32; 4]; 2] {
10015    // SAFETY: callers uphold slice-length contracts.
10016    unsafe {
10017        use core::arch::x86_64::*;
10018        let n = w0.len();
10019        let ones = _mm256_set1_epi16(1);
10020        let mut acc = [[_mm256_setzero_si256(); 4]; 2];
10021        let mut j = 0usize;
10022        while j + 32 <= n {
10023            let wv0 = _mm256_loadu_si256(w0.as_ptr().add(j) as *const __m256i);
10024            let wv1 = _mm256_loadu_si256(w1.as_ptr().add(j) as *const __m256i);
10025            let aw0 = _mm256_abs_epi8(wv0);
10026            let aw1 = _mm256_abs_epi8(wv1);
10027            for (k, x) in xs.iter().enumerate() {
10028                let xv = _mm256_loadu_si256(x.as_ptr().add(j) as *const __m256i);
10029                let p0 = _mm256_maddubs_epi16(aw0, _mm256_sign_epi8(xv, wv0));
10030                acc[0][k] = _mm256_add_epi32(acc[0][k], _mm256_madd_epi16(p0, ones));
10031                let p1 = _mm256_maddubs_epi16(aw1, _mm256_sign_epi8(xv, wv1));
10032                acc[1][k] = _mm256_add_epi32(acc[1][k], _mm256_madd_epi16(p1, ones));
10033            }
10034            j += 32;
10035        }
10036        let mut out = [[0i32; 4]; 2];
10037        for r in 0..2 {
10038            for k in 0..4 {
10039                let a = acc[r][k];
10040                let hi128 = _mm256_extracti128_si256::<1>(a);
10041                let s128 = _mm_add_epi32(_mm256_castsi256_si128(a), hi128);
10042                let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
10043                let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
10044                out[r][k] = _mm_cvtsi128_si32(s32);
10045            }
10046        }
10047        if j < n {
10048            for (k, x) in xs.iter().enumerate() {
10049                for i in j..n {
10050                    out[0][k] += (w0[i] as i8) as i32 * x[i] as i32;
10051                    out[1][k] += (w1[i] as i8) as i32 * x[i] as i32;
10052                }
10053            }
10054        }
10055        out
10056    }
10057}
10058
10059/// AVX2/VNNI q8 row dot with exact outlier correction (x86 mirror of
10060/// `row_dot_sdot` — same A8W8 contract). With AVX-512 VNNI the row goes
10061/// through the bias trick: Σ(w+128)·x via pure `vpdpbusd` (no per-lane
10062/// sign fixups), corrected by −128·Σx with Σx precomputed per split.
10063#[cfg(target_arch = "x86_64")]
10064#[inline]
10065fn row_dot_avx2(row: &[u8], act: &SplitAct) -> f32 {
10066    let dot = if avx512vnni_enabled() && row.len() >= 64 {
10067        (unsafe { dot_u8p128_i8_vnni(row, &act.xq) }) - 128 * act.xsum
10068    } else {
10069        unsafe { dot_i8_i8_avx2(row, &act.xq) }
10070    };
10071    let mut acc = dot as f32 * act.sx;
10072    for &(j, xv) in &act.outliers {
10073        acc += (row[j] as i8) as f32 * xv;
10074    }
10075    acc
10076}
10077
10078/// Σ (w[i]+128)·x[i] via pure `vpdpbusd` — the caller subtracts
10079/// 128·Σx. Four independent accumulators (dpbusd is ~5-cycle latency;
10080/// a single-acc loop runs latency-bound, measured on Granite Rapids).
10081#[cfg(target_arch = "x86_64")]
10082#[target_feature(enable = "avx2,avx512f,avx512bw,avx512vl,avx512vnni")]
10083unsafe fn dot_u8p128_i8_vnni(w: &[u8], xq: &[i8]) -> i32 {
10084    // SAFETY: callers uphold slice-length contracts (see call sites).
10085    unsafe {
10086        use core::arch::x86_64::*;
10087        let n = w.len();
10088        let flip = _mm512_set1_epi8(-128); // XOR 0x80: i8 w → u8 (w+128)
10089        #[inline(always)]
10090        unsafe fn step(
10091            w: *const u8,
10092            x: *const i8,
10093            flip: core::arch::x86_64::__m512i,
10094            acc: core::arch::x86_64::__m512i,
10095        ) -> core::arch::x86_64::__m512i {
10096            unsafe {
10097                use core::arch::x86_64::*;
10098                let wv = _mm512_xor_si512(_mm512_loadu_si512(w as *const _), flip);
10099                _mm512_dpbusd_epi32(acc, wv, _mm512_loadu_si512(x as *const _))
10100            }
10101        }
10102        let (mut a0, mut a1, mut a2, mut a3) = (
10103            _mm512_setzero_si512(),
10104            _mm512_setzero_si512(),
10105            _mm512_setzero_si512(),
10106            _mm512_setzero_si512(),
10107        );
10108        let mut j = 0usize;
10109        while j + 256 <= n {
10110            a0 = step(w.as_ptr().add(j), xq.as_ptr().add(j), flip, a0);
10111            a1 = step(w.as_ptr().add(j + 64), xq.as_ptr().add(j + 64), flip, a1);
10112            a2 = step(w.as_ptr().add(j + 128), xq.as_ptr().add(j + 128), flip, a2);
10113            a3 = step(w.as_ptr().add(j + 192), xq.as_ptr().add(j + 192), flip, a3);
10114            j += 256;
10115        }
10116        while j + 64 <= n {
10117            a0 = step(w.as_ptr().add(j), xq.as_ptr().add(j), flip, a0);
10118            j += 64;
10119        }
10120        let mut total = _mm512_reduce_add_epi32(_mm512_add_epi32(
10121            _mm512_add_epi32(a0, a1),
10122            _mm512_add_epi32(a2, a3),
10123        ));
10124        // Scalar tail: (w as i8) + 128 ≡ (w as u8) ^ 0x80.
10125        while j < n {
10126            total += ((w[j] ^ 0x80) as i32) * xq[j] as i32;
10127            j += 1;
10128        }
10129        total
10130    }
10131}
10132
10133/// One q4 row via AVX2: nibbles → centered i8 (unpacklo/hi restores the
10134/// writer's flat order, same as the NEON vzip pair), maddubs against
10135/// the pre-quantized activation group, × the group's f16 scale. Pair
10136/// saturation safe: |w|≤8, |x|≤127 → 2·8·127 ≪ 32767. Mirror of
10137/// `dot_q4_row_sdot`.
10138#[cfg(target_arch = "x86_64")]
10139#[target_feature(enable = "avx2")]
10140unsafe fn dot_q4_row_avx2(packed: &[u8], scales: &[u8], g0: usize, gpr: usize, xq: &[i8]) -> f32 {
10141    // SAFETY: callers uphold slice-length contracts (16 packed bytes and
10142    // 2 scale bytes per group; xq.len() == gpr·GROUP_SIZE).
10143    unsafe {
10144        use core::arch::x86_64::*;
10145        let lomask = _mm_set1_epi8(0x0F);
10146        let eight = _mm256_set1_epi8(8);
10147        let ones = _mm256_set1_epi16(1);
10148        let mut acc = 0f32;
10149        for gi in 0..gpr {
10150            let g = g0 + gi;
10151            let s = f16_to_f32(u16::from_le_bytes([scales[g * 2], scales[g * 2 + 1]]));
10152            let b = _mm_loadu_si128(packed.as_ptr().add(g * 16) as *const __m128i);
10153            let lo = _mm_and_si128(b, lomask);
10154            let hi = _mm_and_si128(_mm_srli_epi16::<4>(b), lomask);
10155            let w = _mm256_sub_epi8(
10156                _mm256_set_m128i(_mm_unpackhi_epi8(lo, hi), _mm_unpacklo_epi8(lo, hi)),
10157                eight,
10158            );
10159            let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
10160            let p16 = _mm256_maddubs_epi16(_mm256_abs_epi8(w), _mm256_sign_epi8(x, w));
10161            let d = _mm256_madd_epi16(p16, ones);
10162            let hi128 = _mm256_extracti128_si256::<1>(d);
10163            let s128 = _mm_add_epi32(_mm256_castsi256_si128(d), hi128);
10164            let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
10165            let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
10166            acc += _mm_cvtsi128_si32(s32) as f32 * s;
10167        }
10168        acc
10169    }
10170}
10171
10172/// Two-activation q4 row via AVX2: nibbles unpacked ONCE per group,
10173/// both activations dotted against the same centered i8 register.
10174#[cfg(target_arch = "x86_64")]
10175#[target_feature(enable = "avx2")]
10176unsafe fn dot_q4_row_avx2_2(
10177    packed: &[u8],
10178    scales: &[u8],
10179    g0: usize,
10180    gpr: usize,
10181    xq1: &[i8],
10182    xq2: &[i8],
10183) -> (f32, f32) {
10184    // SAFETY: callers uphold slice-length contracts (see dot_q4_row_avx2).
10185    unsafe {
10186        use core::arch::x86_64::*;
10187        let lomask = _mm_set1_epi8(0x0F);
10188        let eight = _mm256_set1_epi8(8);
10189        let ones = _mm256_set1_epi16(1);
10190        let (mut acc1, mut acc2) = (0f32, 0f32);
10191        #[inline(always)]
10192        unsafe fn hsum(d: core::arch::x86_64::__m256i) -> i32 {
10193            unsafe {
10194                use core::arch::x86_64::*;
10195                let hi128 = _mm256_extracti128_si256::<1>(d);
10196                let s128 = _mm_add_epi32(_mm256_castsi256_si128(d), hi128);
10197                let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
10198                let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
10199                _mm_cvtsi128_si32(s32)
10200            }
10201        }
10202        for gi in 0..gpr {
10203            let g = g0 + gi;
10204            let s = f16_to_f32(u16::from_le_bytes([scales[g * 2], scales[g * 2 + 1]]));
10205            let b = _mm_loadu_si128(packed.as_ptr().add(g * 16) as *const __m128i);
10206            let lo = _mm_and_si128(b, lomask);
10207            let hi = _mm_and_si128(_mm_srli_epi16::<4>(b), lomask);
10208            let w = _mm256_sub_epi8(
10209                _mm256_set_m128i(_mm_unpackhi_epi8(lo, hi), _mm_unpacklo_epi8(lo, hi)),
10210                eight,
10211            );
10212            let aw = _mm256_abs_epi8(w);
10213            let x1 = _mm256_loadu_si256(xq1.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
10214            let x2 = _mm256_loadu_si256(xq2.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
10215            let d1 = _mm256_madd_epi16(_mm256_maddubs_epi16(aw, _mm256_sign_epi8(x1, w)), ones);
10216            let d2 = _mm256_madd_epi16(_mm256_maddubs_epi16(aw, _mm256_sign_epi8(x2, w)), ones);
10217            acc1 += hsum(d1) as f32 * s;
10218            acc2 += hsum(d2) as f32 * s;
10219        }
10220        (acc1, acc2)
10221    }
10222}
10223
10224/// One q8 row range via AVX2 (x86 mirror of `q8_range_sdot`).
10225#[cfg(target_arch = "x86_64")]
10226fn q8_range_avx2(
10227    q: &[u8],
10228    row_scale: &[f32],
10229    act: &SplitAct,
10230    cols: usize,
10231    out_addr: SendMut,
10232    start: usize,
10233    end: usize,
10234) {
10235    for o in start..end {
10236        let v = row_dot_avx2(&q[o * cols..(o + 1) * cols], act) * row_scale[o];
10237        // SAFETY: disjoint row ranges per worker.
10238        unsafe { *out_addr.at(o) = v };
10239    }
10240}
10241
10242/// Two-input q8 row range via AVX2 (x86 mirror of `q8_range2_sdot`).
10243#[cfg(target_arch = "x86_64")]
10244#[allow(clippy::too_many_arguments)]
10245fn q8_range2_avx2(
10246    q: &[u8],
10247    row_scale: &[f32],
10248    a1: &SplitAct,
10249    a2: &SplitAct,
10250    cols: usize,
10251    p1: SendMut,
10252    p2: SendMut,
10253    start: usize,
10254    end: usize,
10255) {
10256    for o in start..end {
10257        let row = &q[o * cols..(o + 1) * cols];
10258        // SAFETY: disjoint row ranges per worker.
10259        unsafe {
10260            *p1.at(o) = row_dot_avx2(row, a1) * row_scale[o];
10261            *p2.at(o) = row_dot_avx2(row, a2) * row_scale[o];
10262        }
10263    }
10264}
10265
10266// ───────────────────── A8W8 SDOT path (port of vmfcore, ×1.78 decode) ─────────────────────
10267
10268/// ARMv8.6 i8mm (smmla): 32 int8 MACs per instruction vs sdot's 16 —
10269/// yet MEASURED 2.4× SLOWER than the blocked sdot on Apple silicon
10270/// (108 vs 264 GF/s): the on-the-fly vcombine packing and the two-
10271/// accumulator dependency chain swamp the MAC advantage, and Apple's
10272/// four SIMD pipes already keep sdot fed. OPT-IN (CMF_I8MM=1) for
10273/// field trials on Cortex-A710/X-class parts with two pipes, where the
10274/// balance may differ; a pre-interleaved weight layout (repack infra)
10275/// is the known path if it ever earns its keep.
10276#[cfg(target_arch = "aarch64")]
10277fn i8mm_enabled() -> bool {
10278    static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
10279    *ON.get_or_init(|| {
10280        std::env::var("CMF_I8MM").map(|v| v == "1").unwrap_or(false)
10281            && std::arch::is_aarch64_feature_detected!("i8mm")
10282    })
10283}
10284
10285/// SDOT enabled? Default ON when the CPU has ARMv8.2 dotprod;
10286/// `CMF_SDOT=0` disables (falls back to i8×f32 NEON).
10287/// (On non-ARM release builds only the test tolerance switch calls it.)
10288#[cfg_attr(not(target_arch = "aarch64"), allow(dead_code))]
10289fn sdot_enabled() -> bool {
10290    if FLOAT_ACTIVATIONS.get() {
10291        return false;
10292    }
10293    use std::sync::OnceLock;
10294    static ON: OnceLock<bool> = OnceLock::new();
10295    *ON.get_or_init(|| {
10296        let want = std::env::var("CMF_SDOT").map(|v| v != "0").unwrap_or(true);
10297        if !want {
10298            return false;
10299        }
10300
10301        #[cfg(target_arch = "aarch64")]
10302        {
10303            if std::arch::is_aarch64_feature_detected!("dotprod") {
10304                return true;
10305            }
10306            #[cfg(target_os = "android")]
10307            {
10308                if let Ok(cpuinfo) = std::fs::read_to_string("/proc/cpuinfo") {
10309                    if cpuinfo.lines().any(|l| {
10310                        (l.starts_with("Features") || l.starts_with("features"))
10311                            && l.contains("asimddp")
10312                    }) {
10313                        return true;
10314                    }
10315                }
10316            }
10317            false
10318        }
10319        #[cfg(not(target_arch = "aarch64"))]
10320        {
10321            false
10322        }
10323    })
10324}
10325
10326/// Two-field activation split (≡ vmfcore `q8_split_prep`): outlier
10327/// channels (>8·rms) are computed exactly in f32; the bulk (outliers
10328/// zeroed → clean absmax) goes through int8 SDOT. Computed ONCE per
10329/// matvec, shared by all rows/workers.
10330struct SplitAct {
10331    xq: Vec<i8>,
10332    sx: f32,
10333    outliers: Vec<(usize, f32)>,
10334    /// Σ xq — the VNNI bias-trick correction (`(w+128)·x` sums need
10335    /// `−128·Σx`); one i32 per split, computed once per matvec.
10336    #[cfg_attr(not(target_arch = "x86_64"), allow(dead_code))]
10337    xsum: i32,
10338}
10339
10340thread_local! {
10341    /// Recycled xq buffers: split_act runs for every matvec (~200/token)
10342    /// and its hidden-size allocation was steady-state heap churn.
10343    static XQ_FREE: std::cell::RefCell<Vec<Vec<i8>>> =
10344        const { std::cell::RefCell::new(Vec::new()) };
10345}
10346
10347impl Drop for SplitAct {
10348    fn drop(&mut self) {
10349        let buf = std::mem::take(&mut self.xq);
10350        if buf.capacity() > 0 {
10351            XQ_FREE.with(|f| {
10352                let mut f = f.borrow_mut();
10353                if f.len() < 16 {
10354                    f.push(buf);
10355                }
10356            });
10357        }
10358    }
10359}
10360
10361thread_local! {
10362    /// One scratch row per WORKER, kept for the life of the thread.
10363    ///
10364    /// The kernels take a row of group scales per dispatch, and a fresh
10365    /// `vec![0f32; gpr]` inside the closure is one allocation per worker per
10366    /// dispatch — on the release checkpoint about six thousand a token, a
10367    /// quarter of everything the benchmark counts.
10368    static KROW: UnsafeCell<Vec<f32>> = const { UnsafeCell::new(Vec::new()) };
10369    /// Two scratch rows for a gate/up pair.  Q2TP and Q4TP fused SwiGLU
10370    /// decode each projection's ladder independently, but both ladders have
10371    /// the same lifetime; retaining them per worker removes two heap trips
10372    /// from every dispatched expert pair.
10373    static KROWS: UnsafeCell<[Vec<f32>; 2]> =
10374        const { UnsafeCell::new([Vec::new(), Vec::new()]) };
10375}
10376
10377/// Borrow `n` floats of the calling worker's scratch. Nothing inside a
10378/// kernel body borrows it again, which is what keeps the RefCell honest.
10379#[inline]
10380fn with_krow<R>(n: usize, f: impl FnOnce(&mut [f32]) -> R) -> R {
10381    KROW.with(|s| {
10382        // SAFETY: `KROW` is thread-local and this function does not recurse;
10383        // each worker has exclusive access to its own scratch row.
10384        let b = unsafe { &mut *s.get() };
10385        if b.len() < n {
10386            b.resize(n, 0.0);
10387        }
10388        f(&mut b[..n])
10389    })
10390}
10391
10392/// `t.round().clamp(-127.0, 127.0) as i8`, bit for bit, without the libm
10393/// call. On baseline x86-64 (no SSE4.1 `roundps`) `f32::round` is a
10394/// function call per element, and split_act runs it over every hidden
10395/// state before every matvec: measured 27 us a call on a 2048-wide
10396/// activation on an EPYC 7763 — 5.4 ms of a 55 ms decode token, all of
10397/// it on the caller's thread while thirty workers wait. Clamping first is
10398/// equivalent (round is monotonic and ±127 are integers), and after the
10399/// clamp `t - trunc(t)` is exact, so the half-away-from-zero decision is
10400/// the one `round` makes. NaN clamps to NaN and converts to 0, as before.
10401/// The loop vectorizes (cvttps2dq + compare/select).
10402#[inline(always)]
10403fn q8_round(t: f32) -> i8 {
10404    let t = t.clamp(-127.0, 127.0);
10405    let i = t as i32;
10406    let f = t - i as f32;
10407    let r = if f >= 0.5 {
10408        i + 1
10409    } else if f <= -0.5 {
10410        i - 1
10411    } else {
10412        i
10413    };
10414    r as i8
10415}
10416
10417#[inline]
10418fn with_krows<R>(n: usize, f: impl FnOnce(&mut [f32], &mut [f32]) -> R) -> R {
10419    KROWS.with(|s| {
10420        // SAFETY: KROWS is thread-local and no kernel body recursively calls
10421        // this helper; each participant owns its two vectors exclusively.
10422        let b = unsafe { &mut *s.get() };
10423        for row in b.iter_mut() {
10424            if row.len() < n {
10425                row.resize(n, 0.0);
10426            }
10427        }
10428        let (a, b) = b.split_at_mut(1);
10429        f(&mut a[0][..n], &mut b[0][..n])
10430    })
10431}
10432
10433#[inline]
10434fn silu_mul_limited(mut gate: f32, mut up: f32, limit: f32) -> f32 {
10435    if limit > 0.0 {
10436        up = up.clamp(-limit, limit);
10437        gate = gate.min(limit);
10438    }
10439    gate / (1.0 + (-gate).exp()) * up
10440}
10441
10442fn split_act(x: &[f32]) -> SplitAct {
10443    let _prof = crate::cpuprof::time(crate::cpuprof::Slot::SplitAct);
10444    let n = x.len();
10445    let rms = (x.iter().map(|&v| (v * v) as f64).sum::<f64>() / n.max(1) as f64).sqrt() as f32;
10446    let thr = 8.0 * rms;
10447    // One pass: collect outliers and the bulk absmax (outliers excluded —
10448    // identical to the old zero-then-fold over a copied buffer, minus the
10449    // full-vector copy).
10450    let mut outliers: Vec<(usize, f32)> = Vec::new();
10451    let mut amax = 0f32;
10452    for (j, &v) in x.iter().enumerate() {
10453        let a = v.abs();
10454        if a > thr {
10455            outliers.push((j, v));
10456        } else if a > amax {
10457            amax = a;
10458        }
10459    }
10460    let sx = if amax > 0.0 { amax / 127.0 } else { 1.0 };
10461    let inv = 1.0 / sx;
10462    let mut xq = XQ_FREE.with(|f| f.borrow_mut().pop()).unwrap_or_default();
10463    xq.clear();
10464    xq.reserve(n);
10465    if outliers.is_empty() {
10466        xq.extend(
10467            x.iter()
10468                .map(|&v| q8_round(v * inv)),
10469        );
10470    } else {
10471        // Outlier slots quantize to 0 (their exact term is added later).
10472        xq.extend(x.iter().map(|&v| {
10473            if v.abs() > thr {
10474                0
10475            } else {
10476                q8_round(v * inv)
10477            }
10478        }));
10479    }
10480    let xsum = xq.iter().map(|&v| v as i32).sum();
10481    SplitAct {
10482        xq,
10483        sx,
10484        outliers,
10485        xsum,
10486    }
10487}
10488
10489fn split_act_q8_2f(x: &[f32], col: &[f32]) -> SplitAct {
10490    let _prof = crate::cpuprof::time(crate::cpuprof::Slot::SplitAct);
10491    let n = x.len();
10492    let rms = (x
10493        .iter()
10494        .zip(col)
10495        .map(|(&a, &c)| {
10496            let v = a * c;
10497            (v * v) as f64
10498        })
10499        .sum::<f64>()
10500        / n.max(1) as f64)
10501        .sqrt() as f32;
10502    let thr = 8.0 * rms;
10503
10504    let mut outliers = Vec::new();
10505    let mut amax = 0f32;
10506    for (j, (&a, &c)) in x.iter().zip(col).enumerate() {
10507        let v = a * c;
10508        let s = v.abs();
10509        if s > thr {
10510            outliers.push((j, v));
10511        } else if s > amax {
10512            amax = s;
10513        }
10514    }
10515
10516    let sx = if amax > 0.0 { amax / 127.0 } else { 1.0 };
10517    let inv = 1.0 / sx;
10518    let mut xq = XQ_FREE.with(|f| f.borrow_mut().pop()).unwrap_or_default();
10519    xq.clear();
10520    xq.reserve(n);
10521    if outliers.is_empty() {
10522        xq.extend(
10523            x.iter()
10524                .zip(col)
10525                .map(|(&a, &c)| q8_round((a * c) * inv)),
10526        );
10527    } else {
10528        xq.extend(x.iter().zip(col).map(|(&a, &c)| {
10529            let v = a * c;
10530            if v.abs() > thr {
10531                0
10532            } else {
10533                q8_round(v * inv)
10534            }
10535        }));
10536    }
10537    let xsum = xq.iter().map(|&v| v as i32).sum();
10538    SplitAct {
10539        xq,
10540        sx,
10541        outliers,
10542        xsum,
10543    }
10544}
10545
10546/// int8(weight)·int8(activation) → i32 via `sdot` (inline asm — the
10547/// vdotq intrinsic is unstable; port of vmfcore `dot_i8_sdot`).
10548#[cfg(target_arch = "aarch64")]
10549#[target_feature(enable = "neon,dotprod")]
10550unsafe fn dot_i8_sdot(w: &[u8], xq: &[i8]) -> i32 {
10551    // SAFETY: callers uphold slice-length contracts (see call sites).
10552    unsafe {
10553        use core::arch::aarch64::*;
10554        use core::arch::asm;
10555        let wp = w.as_ptr() as *const i8;
10556        let n = w.len();
10557        let (mut a0, mut a1, mut a2, mut a3) = (
10558            vdupq_n_s32(0),
10559            vdupq_n_s32(0),
10560            vdupq_n_s32(0),
10561            vdupq_n_s32(0),
10562        );
10563        let mut i = 0;
10564        while i + 64 <= n {
10565            let (w0, x0) = (vld1q_s8(wp.add(i)), vld1q_s8(xq.as_ptr().add(i)));
10566            let (w1, x1) = (vld1q_s8(wp.add(i + 16)), vld1q_s8(xq.as_ptr().add(i + 16)));
10567            let (w2, x2) = (vld1q_s8(wp.add(i + 32)), vld1q_s8(xq.as_ptr().add(i + 32)));
10568            let (w3, x3) = (vld1q_s8(wp.add(i + 48)), vld1q_s8(xq.as_ptr().add(i + 48)));
10569            asm!(
10570                "sdot {a0:v}.4s, {w0:v}.16b, {x0:v}.16b",
10571                "sdot {a1:v}.4s, {w1:v}.16b, {x1:v}.16b",
10572                "sdot {a2:v}.4s, {w2:v}.16b, {x2:v}.16b",
10573                "sdot {a3:v}.4s, {w3:v}.16b, {x3:v}.16b",
10574                a0 = inout(vreg) a0, a1 = inout(vreg) a1, a2 = inout(vreg) a2, a3 = inout(vreg) a3,
10575                w0 = in(vreg) w0, x0 = in(vreg) x0, w1 = in(vreg) w1, x1 = in(vreg) x1,
10576                w2 = in(vreg) w2, x2 = in(vreg) x2, w3 = in(vreg) w3, x3 = in(vreg) x3,
10577                options(pure, nomem, nostack),
10578            );
10579            i += 64;
10580        }
10581        while i + 16 <= n {
10582            let (wv, xv) = (vld1q_s8(wp.add(i)), vld1q_s8(xq.as_ptr().add(i)));
10583            asm!("sdot {a:v}.4s, {w:v}.16b, {x:v}.16b",
10584                 a = inout(vreg) a0, w = in(vreg) wv, x = in(vreg) xv, options(pure, nomem, nostack));
10585            i += 16;
10586        }
10587        let mut s = vaddvq_s32(vaddq_s32(vaddq_s32(a0, a1), vaddq_s32(a2, a3)));
10588        while i < n {
10589            s += (*wp.add(i)) as i32 * xq[i] as i32;
10590            i += 1;
10591        }
10592        s
10593    }
10594}
10595
10596/// Row-blocked SDOT: 4 output rows per pass — the activation chunk is
10597/// loaded once and reused, 4 independent accumulators hide sdot latency
10598/// (port of vmfcore `dot_i8_sdot_4rows`).
10599#[cfg(target_arch = "aarch64")]
10600#[target_feature(enable = "neon,dotprod")]
10601unsafe fn dot_i8_sdot_4rows(w0: &[u8], w1: &[u8], w2: &[u8], w3: &[u8], xq: &[i8]) -> [i32; 4] {
10602    // SAFETY: callers uphold slice-length contracts (see call sites).
10603    unsafe {
10604        use core::arch::aarch64::*;
10605        use core::arch::asm;
10606        let n = xq.len();
10607        let px = xq.as_ptr();
10608        let (p0, p1, p2, p3) = (
10609            w0.as_ptr() as *const i8,
10610            w1.as_ptr() as *const i8,
10611            w2.as_ptr() as *const i8,
10612            w3.as_ptr() as *const i8,
10613        );
10614        let (mut a0, mut a1, mut a2, mut a3) = (
10615            vdupq_n_s32(0),
10616            vdupq_n_s32(0),
10617            vdupq_n_s32(0),
10618            vdupq_n_s32(0),
10619        );
10620        let mut i = 0;
10621        while i + 16 <= n {
10622            let x = vld1q_s8(px.add(i));
10623            let v0 = vld1q_s8(p0.add(i));
10624            let v1 = vld1q_s8(p1.add(i));
10625            let v2 = vld1q_s8(p2.add(i));
10626            let v3 = vld1q_s8(p3.add(i));
10627            asm!(
10628                "sdot {a0:v}.4s, {v0:v}.16b, {x:v}.16b",
10629                "sdot {a1:v}.4s, {v1:v}.16b, {x:v}.16b",
10630                "sdot {a2:v}.4s, {v2:v}.16b, {x:v}.16b",
10631                "sdot {a3:v}.4s, {v3:v}.16b, {x:v}.16b",
10632                a0 = inout(vreg) a0, a1 = inout(vreg) a1, a2 = inout(vreg) a2, a3 = inout(vreg) a3,
10633                v0 = in(vreg) v0, v1 = in(vreg) v1, v2 = in(vreg) v2, v3 = in(vreg) v3, x = in(vreg) x,
10634                options(pure, nomem, nostack),
10635            );
10636            i += 16;
10637        }
10638        let mut r = [
10639            vaddvq_s32(a0),
10640            vaddvq_s32(a1),
10641            vaddvq_s32(a2),
10642            vaddvq_s32(a3),
10643        ];
10644        while i < n {
10645            let xi = *px.add(i) as i32;
10646            r[0] += (*p0.add(i)) as i32 * xi;
10647            r[1] += (*p1.add(i)) as i32 * xi;
10648            r[2] += (*p2.add(i)) as i32 * xi;
10649            r[3] += (*p3.add(i)) as i32 * xi;
10650            i += 1;
10651        }
10652        r
10653    }
10654}
10655
10656/// 4 interleaved rows in one pass: the repacked group is [r0[c], r1[c],
10657/// r2[c], r3[c]] per 16-byte chunk, so each iteration reads ONE 64-byte
10658/// line plus the shared activation chunk — a single sequential weight
10659/// stream per worker. Per-row accumulation is the same one-accumulator
10660/// scheme as `dot_i8_sdot_4rows`; integer sums are exact, so outputs
10661/// are bit-identical to the mmap-layout kernel.
10662#[cfg(target_arch = "aarch64")]
10663#[target_feature(enable = "neon,dotprod")]
10664unsafe fn dot_i8_sdot_4rows_il(g: &[u8], xq: &[i8]) -> [i32; 4] {
10665    // SAFETY: callers uphold slice-length contracts (g.len() == 4·n,
10666    // n % 16 == 0 — guaranteed by the repack gate).
10667    unsafe {
10668        use core::arch::aarch64::*;
10669        use core::arch::asm;
10670        let n = xq.len();
10671        let px = xq.as_ptr();
10672        let pg = g.as_ptr() as *const i8;
10673        let (mut a0, mut a1, mut a2, mut a3) = (
10674            vdupq_n_s32(0),
10675            vdupq_n_s32(0),
10676            vdupq_n_s32(0),
10677            vdupq_n_s32(0),
10678        );
10679        let mut i = 0;
10680        while i + 16 <= n {
10681            let x = vld1q_s8(px.add(i));
10682            let base = pg.add(4 * i);
10683            let v0 = vld1q_s8(base);
10684            let v1 = vld1q_s8(base.add(16));
10685            let v2 = vld1q_s8(base.add(32));
10686            let v3 = vld1q_s8(base.add(48));
10687            asm!(
10688                "sdot {a0:v}.4s, {v0:v}.16b, {x:v}.16b",
10689                "sdot {a1:v}.4s, {v1:v}.16b, {x:v}.16b",
10690                "sdot {a2:v}.4s, {v2:v}.16b, {x:v}.16b",
10691                "sdot {a3:v}.4s, {v3:v}.16b, {x:v}.16b",
10692                a0 = inout(vreg) a0, a1 = inout(vreg) a1, a2 = inout(vreg) a2, a3 = inout(vreg) a3,
10693                v0 = in(vreg) v0, v1 = in(vreg) v1, v2 = in(vreg) v2, v3 = in(vreg) v3, x = in(vreg) x,
10694                options(pure, nomem, nostack),
10695            );
10696            i += 16;
10697        }
10698        [
10699            vaddvq_s32(a0),
10700            vaddvq_s32(a1),
10701            vaddvq_s32(a2),
10702            vaddvq_s32(a3),
10703        ]
10704    }
10705}
10706
10707/// One q8 row range via SDOT (4-row blocks + tail) — the body of
10708/// `qmatvec`'s hot loop, extracted so multi-matrix jobs can drive the
10709/// SAME kernel for several tensors under one pool dispatch. `rep` — the
10710/// load-time interleaved repack (empty = mmap layout only); rows outside
10711/// full 4-row groups always come from the mmap layout.
10712#[cfg(target_arch = "aarch64")]
10713fn q8_range_sdot(
10714    q: &[u8],
10715    rep: &[u8],
10716    row_scale: &[f32],
10717    act: &SplitAct,
10718    cols: usize,
10719    out_addr: SendMut,
10720    start: usize,
10721    end: usize,
10722) {
10723    let mut o = start;
10724    // Leading rows to the group boundary (repack path only): the pool
10725    // splits row ranges arbitrarily, groups are absolute.
10726    if !rep.is_empty() {
10727        while o < end && o % 4 != 0 {
10728            let v = row_dot_sdot(&q[o * cols..(o + 1) * cols], act) * row_scale[o];
10729            unsafe { *out_addr.at(o) = v };
10730            o += 1;
10731        }
10732    }
10733    while o + 4 <= end {
10734        let r = if rep.is_empty() {
10735            unsafe {
10736                dot_i8_sdot_4rows(
10737                    &q[o * cols..(o + 1) * cols],
10738                    &q[(o + 1) * cols..(o + 2) * cols],
10739                    &q[(o + 2) * cols..(o + 3) * cols],
10740                    &q[(o + 3) * cols..(o + 4) * cols],
10741                    &act.xq,
10742                )
10743            }
10744        } else {
10745            unsafe { dot_i8_sdot_4rows_il(&rep[o * cols..(o + 4) * cols], &act.xq) }
10746        };
10747        for k in 0..4 {
10748            let mut acc = r[k] as f32 * act.sx;
10749            for &(j, xv) in &act.outliers {
10750                acc += (q[(o + k) * cols + j] as i8) as f32 * xv;
10751            }
10752            // SAFETY: disjoint row ranges per worker.
10753            unsafe { *out_addr.at(o + k) = acc * row_scale[o + k] };
10754        }
10755        o += 4;
10756    }
10757    while o < end {
10758        let v = row_dot_sdot(&q[o * cols..(o + 1) * cols], act) * row_scale[o];
10759        unsafe { *out_addr.at(o) = v };
10760        o += 1;
10761    }
10762}
10763
10764/// Two-input q8 row range via SDOT — `qmatvec2`'s hot loop, extracted
10765/// for the fused pair multi-matrix job (`matvec2_many`).
10766#[cfg(target_arch = "aarch64")]
10767#[allow(clippy::too_many_arguments)]
10768fn q8_range2_sdot(
10769    q: &[u8],
10770    row_scale: &[f32],
10771    a1: &SplitAct,
10772    a2: &SplitAct,
10773    cols: usize,
10774    p1: SendMut,
10775    p2: SendMut,
10776    start: usize,
10777    end: usize,
10778) {
10779    for o in start..end {
10780        let row = &q[o * cols..(o + 1) * cols];
10781        // SAFETY: disjoint row ranges per worker.
10782        unsafe {
10783            *p1.at(o) = row_dot_sdot(row, a1) * row_scale[o];
10784            *p2.at(o) = row_dot_sdot(row, a2) * row_scale[o];
10785        }
10786    }
10787}
10788
10789/// Two-input q8 row range, f32 kernel (non-SDOT) — same extraction.
10790#[allow(clippy::too_many_arguments)]
10791fn q8_range2_f32(
10792    q: &[u8],
10793    row_scale: &[f32],
10794    x1: &[f32],
10795    x2: &[f32],
10796    cols: usize,
10797    p1: SendMut,
10798    p2: SendMut,
10799    start: usize,
10800    end: usize,
10801) {
10802    for o in start..end {
10803        let row = &q[o * cols..(o + 1) * cols];
10804        // SAFETY: disjoint row ranges per worker.
10805        unsafe {
10806            *p1.at(o) = dot_i8_f32(row, x1) * row_scale[o];
10807            *p2.at(o) = dot_i8_f32(row, x2) * row_scale[o];
10808        }
10809    }
10810}
10811
10812/// Scalar/NEON-f32 q8 row range (non-SDOT platforms) — same extraction.
10813fn q8_range_f32(
10814    q: &[u8],
10815    row_scale: &[f32],
10816    xs: &[f32],
10817    cols: usize,
10818    out_addr: SendMut,
10819    start: usize,
10820    end: usize,
10821) {
10822    for o in start..end {
10823        let v = dot_i8_f32(&q[o * cols..(o + 1) * cols], xs) * row_scale[o];
10824        // SAFETY: disjoint row ranges per worker.
10825        unsafe { *out_addr.at(o) = v };
10826    }
10827}
10828
10829/// One q8 row against a split activation, portable: the per-arch fast
10830/// dots where they exist, the exact scalar loop elsewhere. The scalar
10831/// arm is also the test oracle for both fast arms.
10832#[inline]
10833fn q8_row_dot(row: &[u8], act: &SplitAct) -> f32 {
10834    #[cfg(target_arch = "aarch64")]
10835    return row_dot_sdot(row, act);
10836    #[cfg(target_arch = "x86_64")]
10837    return row_dot_avx2(row, act);
10838    #[allow(unreachable_code)]
10839    q8_row_dot_scalar(row, act)
10840}
10841
10842#[allow(dead_code)]
10843fn q8_row_dot_scalar(row: &[u8], act: &SplitAct) -> f32 {
10844    let mut acc = 0i32;
10845    for (k, &b) in row.iter().enumerate() {
10846        acc += (b as i8) as i32 * act.xq[k] as i32;
10847    }
10848    let mut acc = acc as f32 * act.sx;
10849    for &(j, xv) in &act.outliers {
10850        acc += (row[j] as i8) as f32 * xv;
10851    }
10852    acc
10853}
10854
10855/// SDOT row dot with exact outlier correction:
10856/// `dot = sdot(w, xq)·sx + Σ_outl w[j]·x[j]` (then × row_scale by caller).
10857#[cfg(target_arch = "aarch64")]
10858#[inline]
10859fn row_dot_sdot(row: &[u8], act: &SplitAct) -> f32 {
10860    let mut acc = unsafe { dot_i8_sdot(row, &act.xq) } as f32 * act.sx;
10861    for &(j, xv) in &act.outliers {
10862        acc += (row[j] as i8) as f32 * xv;
10863    }
10864    acc
10865}
10866
10867/// One q4 row via SDOT: each 32-group's nibbles unpack to centered i8
10868/// (nib−8 ∈ [−8,7]), int8×int8 `sdot` against the pre-quantized
10869/// activation group, × the group's f16 scale. Returns Σ_g dot_g·s_g;
10870/// the caller multiplies by the activation scale and adds the exact
10871/// outlier terms (port of vmfcore `dot_q4_block_sdot`, +23% measured).
10872/// Nibble order matches the writer: element 2k = low nibble, 2k+1 = high
10873/// → zip(lo,hi) restores flat order.
10874#[cfg(target_arch = "aarch64")]
10875#[target_feature(enable = "neon,dotprod")]
10876unsafe fn dot_q4_row_sdot(packed: &[u8], scales: &[u8], g0: usize, gpr: usize, xq: &[i8]) -> f32 {
10877    // SAFETY: callers uphold slice-length contracts (16 packed bytes and
10878    // 2 scale bytes per group; xq.len() == gpr·GROUP_SIZE).
10879    unsafe {
10880        use core::arch::aarch64::*;
10881        use core::arch::asm;
10882        let lomask = vdupq_n_u8(0x0F);
10883        let eight = vdupq_n_s8(8);
10884        let mut acc = 0f32;
10885        for gi in 0..gpr {
10886            let g = g0 + gi;
10887            let s = f16_to_f32(u16::from_le_bytes([scales[g * 2], scales[g * 2 + 1]]));
10888            let b = vld1q_u8(packed.as_ptr().add(g * 16));
10889            let lo = vandq_u8(b, lomask);
10890            let hi = vshrq_n_u8::<4>(b);
10891            let e0 = vsubq_s8(vreinterpretq_s8_u8(vzip1q_u8(lo, hi)), eight);
10892            let e1 = vsubq_s8(vreinterpretq_s8_u8(vzip2q_u8(lo, hi)), eight);
10893            let x0 = vld1q_s8(xq.as_ptr().add(gi * GROUP_SIZE));
10894            let x1 = vld1q_s8(xq.as_ptr().add(gi * GROUP_SIZE + 16));
10895            let (mut a0, mut a1) = (vdupq_n_s32(0), vdupq_n_s32(0));
10896            asm!(
10897                "sdot {a0:v}.4s, {e0:v}.16b, {x0:v}.16b",
10898                "sdot {a1:v}.4s, {e1:v}.16b, {x1:v}.16b",
10899                a0 = inout(vreg) a0, a1 = inout(vreg) a1,
10900                e0 = in(vreg) e0, x0 = in(vreg) x0, e1 = in(vreg) e1, x1 = in(vreg) x1,
10901                options(pure, nomem, nostack),
10902            );
10903            acc += vaddvq_s32(vaddq_s32(a0, a1)) as f32 * s;
10904        }
10905        acc
10906    }
10907}
10908
10909/// Two-activation q4 row via SDOT: the nibble unpack (the expensive
10910/// part) happens ONCE per group; both pre-quantized activations are
10911/// dotted against the same centered i8 registers. Per-lane math matches
10912/// `dot_q4_row_sdot` exactly.
10913#[cfg(target_arch = "aarch64")]
10914#[target_feature(enable = "neon,dotprod")]
10915unsafe fn dot_q4_row_sdot2(
10916    packed: &[u8],
10917    scales: &[u8],
10918    g0: usize,
10919    gpr: usize,
10920    xq1: &[i8],
10921    xq2: &[i8],
10922) -> (f32, f32) {
10923    // SAFETY: callers uphold slice-length contracts (16 packed bytes and
10924    // 2 scale bytes per group; xq*.len() == gpr·GROUP_SIZE).
10925    unsafe {
10926        use core::arch::aarch64::*;
10927        use core::arch::asm;
10928        let lomask = vdupq_n_u8(0x0F);
10929        let eight = vdupq_n_s8(8);
10930        let (mut acc1, mut acc2) = (0f32, 0f32);
10931        for gi in 0..gpr {
10932            let g = g0 + gi;
10933            let s = f16_to_f32(u16::from_le_bytes([scales[g * 2], scales[g * 2 + 1]]));
10934            let b = vld1q_u8(packed.as_ptr().add(g * 16));
10935            let lo = vandq_u8(b, lomask);
10936            let hi = vshrq_n_u8::<4>(b);
10937            let e0 = vsubq_s8(vreinterpretq_s8_u8(vzip1q_u8(lo, hi)), eight);
10938            let e1 = vsubq_s8(vreinterpretq_s8_u8(vzip2q_u8(lo, hi)), eight);
10939            let x10 = vld1q_s8(xq1.as_ptr().add(gi * GROUP_SIZE));
10940            let x11 = vld1q_s8(xq1.as_ptr().add(gi * GROUP_SIZE + 16));
10941            let x20 = vld1q_s8(xq2.as_ptr().add(gi * GROUP_SIZE));
10942            let x21 = vld1q_s8(xq2.as_ptr().add(gi * GROUP_SIZE + 16));
10943            let (mut a0, mut a1, mut b0, mut b1) = (
10944                vdupq_n_s32(0),
10945                vdupq_n_s32(0),
10946                vdupq_n_s32(0),
10947                vdupq_n_s32(0),
10948            );
10949            asm!(
10950                "sdot {a0:v}.4s, {e0:v}.16b, {x10:v}.16b",
10951                "sdot {a1:v}.4s, {e1:v}.16b, {x11:v}.16b",
10952                "sdot {b0:v}.4s, {e0:v}.16b, {x20:v}.16b",
10953                "sdot {b1:v}.4s, {e1:v}.16b, {x21:v}.16b",
10954                a0 = inout(vreg) a0, a1 = inout(vreg) a1,
10955                b0 = inout(vreg) b0, b1 = inout(vreg) b1,
10956                e0 = in(vreg) e0, e1 = in(vreg) e1,
10957                x10 = in(vreg) x10, x11 = in(vreg) x11,
10958                x20 = in(vreg) x20, x21 = in(vreg) x21,
10959                options(pure, nomem, nostack),
10960            );
10961            acc1 += vaddvq_s32(vaddq_s32(a0, a1)) as f32 * s;
10962            acc2 += vaddvq_s32(vaddq_s32(b0, b1)) as f32 * s;
10963        }
10964        (acc1, acc2)
10965    }
10966}
10967
10968// ───────────────────── fused int8 kernels ─────────────────────
10969
10970/// `acc += w · row` where the row is centered i8 — NEON widen+fma on
10971/// aarch64, scalar elsewhere. The KV-cache q8 value path rides on this.
10972#[inline]
10973pub(crate) fn axpy_i8_f32(acc: &mut [f32], row: &[i8], w: f32) {
10974    #[cfg(target_arch = "aarch64")]
10975    unsafe {
10976        return axpy_i8_f32_neon(acc, row, w);
10977    }
10978    #[cfg(target_arch = "x86_64")]
10979    if avx2_enabled() {
10980        return unsafe { axpy_i8_f32_avx2(acc, row, w) };
10981    }
10982    #[allow(unreachable_code)]
10983    {
10984        for (a, &b) in acc.iter_mut().zip(row) {
10985            *a += w * b as f32;
10986        }
10987    }
10988}
10989
10990/// i8→f32 axpy via AVX2/FMA (x86 mirror of `axpy_i8_f32_neon`).
10991#[cfg(target_arch = "x86_64")]
10992#[target_feature(enable = "avx2,fma")]
10993unsafe fn axpy_i8_f32_avx2(acc: &mut [f32], row: &[i8], w: f32) {
10994    // SAFETY: callers uphold slice-length contracts (see call sites).
10995    unsafe {
10996        use core::arch::x86_64::*;
10997        let n = acc.len().min(row.len());
10998        let ap = acc.as_mut_ptr();
10999        let rp = row.as_ptr();
11000        let wv = _mm256_set1_ps(w);
11001        let mut j = 0usize;
11002        while j + 16 <= n {
11003            let rb = _mm_loadu_si128(rp.add(j) as *const __m128i);
11004            let lo = _mm256_cvtepi32_ps(_mm256_cvtepi8_epi32(rb));
11005            let hi = _mm256_cvtepi32_ps(_mm256_cvtepi8_epi32(_mm_srli_si128::<8>(rb)));
11006            let v0 = _mm256_fmadd_ps(wv, lo, _mm256_loadu_ps(ap.add(j)));
11007            let v1 = _mm256_fmadd_ps(wv, hi, _mm256_loadu_ps(ap.add(j + 8)));
11008            _mm256_storeu_ps(ap.add(j), v0);
11009            _mm256_storeu_ps(ap.add(j + 8), v1);
11010            j += 16;
11011        }
11012        while j < n {
11013            *ap.add(j) += w * (*rp.add(j)) as f32;
11014            j += 1;
11015        }
11016    }
11017}
11018
11019#[cfg(target_arch = "aarch64")]
11020#[target_feature(enable = "neon")]
11021unsafe fn axpy_i8_f32_neon(acc: &mut [f32], row: &[i8], w: f32) {
11022    // SAFETY: callers uphold slice-length contracts (see call sites).
11023    unsafe {
11024        use core::arch::aarch64::*;
11025        let n = acc.len().min(row.len());
11026        let ap = acc.as_mut_ptr();
11027        let rp = row.as_ptr();
11028        let wv = vdupq_n_f32(w);
11029        let mut j = 0usize;
11030        while j + 16 <= n {
11031            let rb = vld1q_s8(rp.add(j));
11032            let lo = vmovl_s8(vget_low_s8(rb));
11033            let hi = vmovl_s8(vget_high_s8(rb));
11034            for (off, half) in [(0, lo), (8, hi)] {
11035                let f0 = vcvtq_f32_s32(vmovl_s16(vget_low_s16(half)));
11036                let f1 = vcvtq_f32_s32(vmovl_s16(vget_high_s16(half)));
11037                let o = j + off;
11038                vst1q_f32(ap.add(o), vfmaq_f32(vld1q_f32(ap.add(o)), wv, f0));
11039                vst1q_f32(ap.add(o + 4), vfmaq_f32(vld1q_f32(ap.add(o + 4)), wv, f1));
11040            }
11041            j += 16;
11042        }
11043        while j < n {
11044            *ap.add(j) += w * (*rp.add(j)) as f32;
11045            j += 1;
11046        }
11047    }
11048}
11049
11050/// i8 row · f32 x. NEON on aarch64 (ported from vmfcore `dot_i8_f32_neon`,
11051/// ≈9× scalar), scalar elsewhere.
11052#[inline]
11053pub(crate) fn dot_i8_f32(w: &[u8], x: &[f32]) -> f32 {
11054    #[cfg(target_arch = "aarch64")]
11055    unsafe {
11056        return dot_i8_f32_neon(w, x);
11057    }
11058    #[cfg(target_arch = "x86_64")]
11059    if avx2_enabled() {
11060        return unsafe { dot_i8_f32_avx2(w, x) };
11061    }
11062    #[allow(unreachable_code)]
11063    {
11064        let mut sum = 0.0f32;
11065        for (j, &b) in w.iter().enumerate() {
11066            sum += (b as i8) as f32 * x[j];
11067        }
11068        sum
11069    }
11070}
11071
11072/// i8 row · (x ⊙ col_field) — the q8_2f row dot with the θ col-field
11073/// folded into the product (no prescaled copy of x). NEON on aarch64,
11074/// scalar elsewhere. Used by the active-neuron path `row_dot`.
11075#[inline]
11076fn dot_i8_col_f32(w: &[u8], x: &[f32], col: &[f32]) -> f32 {
11077    #[cfg(target_arch = "aarch64")]
11078    unsafe {
11079        return dot_i8_col_f32_neon(w, x, col);
11080    }
11081    #[allow(unreachable_code)]
11082    {
11083        let mut sum = 0.0f32;
11084        for (j, &b) in w.iter().enumerate() {
11085            sum += (b as i8) as f32 * x[j] * col[j];
11086        }
11087        sum
11088    }
11089}
11090
11091#[cfg(target_arch = "aarch64")]
11092#[target_feature(enable = "neon")]
11093unsafe fn dot_i8_col_f32_neon(w: &[u8], x: &[f32], col: &[f32]) -> f32 {
11094    // SAFETY: callers uphold slice-length contracts (see call sites).
11095    unsafe {
11096        use core::arch::aarch64::*;
11097        let n = x.len();
11098        let wp = w.as_ptr() as *const i8;
11099        let xp = x.as_ptr();
11100        let cp = col.as_ptr();
11101        let (mut a0, mut a1, mut a2, mut a3) = (
11102            vdupq_n_f32(0.0),
11103            vdupq_n_f32(0.0),
11104            vdupq_n_f32(0.0),
11105            vdupq_n_f32(0.0),
11106        );
11107        let mut j = 0usize;
11108        while j + 16 <= n {
11109            let wb = vld1q_s8(wp.add(j));
11110            let lo = vmovl_s8(vget_low_s8(wb));
11111            let hi = vmovl_s8(vget_high_s8(wb));
11112            let w0 = vcvtq_f32_s32(vmovl_s16(vget_low_s16(lo)));
11113            let w1 = vcvtq_f32_s32(vmovl_s16(vget_high_s16(lo)));
11114            let w2 = vcvtq_f32_s32(vmovl_s16(vget_low_s16(hi)));
11115            let w3 = vcvtq_f32_s32(vmovl_s16(vget_high_s16(hi)));
11116            a0 = vfmaq_f32(
11117                a0,
11118                w0,
11119                vmulq_f32(vld1q_f32(xp.add(j)), vld1q_f32(cp.add(j))),
11120            );
11121            a1 = vfmaq_f32(
11122                a1,
11123                w1,
11124                vmulq_f32(vld1q_f32(xp.add(j + 4)), vld1q_f32(cp.add(j + 4))),
11125            );
11126            a2 = vfmaq_f32(
11127                a2,
11128                w2,
11129                vmulq_f32(vld1q_f32(xp.add(j + 8)), vld1q_f32(cp.add(j + 8))),
11130            );
11131            a3 = vfmaq_f32(
11132                a3,
11133                w3,
11134                vmulq_f32(vld1q_f32(xp.add(j + 12)), vld1q_f32(cp.add(j + 12))),
11135            );
11136            j += 16;
11137        }
11138        let mut sum = vaddvq_f32(vaddq_f32(vaddq_f32(a0, a1), vaddq_f32(a2, a3)));
11139        while j < n {
11140            sum += (*wp.add(j)) as f32 * *xp.add(j) * *cp.add(j);
11141            j += 1;
11142        }
11143        sum
11144    }
11145}
11146
11147#[cfg(target_arch = "aarch64")]
11148#[target_feature(enable = "neon")]
11149unsafe fn dot_i8_f32_neon(w: &[u8], x: &[f32]) -> f32 {
11150    // SAFETY: callers uphold slice-length contracts (see call sites).
11151    unsafe {
11152        use core::arch::aarch64::*;
11153        let n = x.len();
11154        let wp = w.as_ptr() as *const i8;
11155        let xp = x.as_ptr();
11156        let (mut a0, mut a1, mut a2, mut a3) = (
11157            vdupq_n_f32(0.0),
11158            vdupq_n_f32(0.0),
11159            vdupq_n_f32(0.0),
11160            vdupq_n_f32(0.0),
11161        );
11162        let mut j = 0usize;
11163        while j + 16 <= n {
11164            let wb = vld1q_s8(wp.add(j));
11165            let lo = vmovl_s8(vget_low_s8(wb));
11166            let hi = vmovl_s8(vget_high_s8(wb));
11167            let w0 = vcvtq_f32_s32(vmovl_s16(vget_low_s16(lo)));
11168            let w1 = vcvtq_f32_s32(vmovl_s16(vget_high_s16(lo)));
11169            let w2 = vcvtq_f32_s32(vmovl_s16(vget_low_s16(hi)));
11170            let w3 = vcvtq_f32_s32(vmovl_s16(vget_high_s16(hi)));
11171            a0 = vfmaq_f32(a0, w0, vld1q_f32(xp.add(j)));
11172            a1 = vfmaq_f32(a1, w1, vld1q_f32(xp.add(j + 4)));
11173            a2 = vfmaq_f32(a2, w2, vld1q_f32(xp.add(j + 8)));
11174            a3 = vfmaq_f32(a3, w3, vld1q_f32(xp.add(j + 12)));
11175            j += 16;
11176        }
11177        let mut sum = vaddvq_f32(vaddq_f32(vaddq_f32(a0, a1), vaddq_f32(a2, a3)));
11178        while j < n {
11179            sum += (*wp.add(j)) as f32 * *xp.add(j);
11180            j += 1;
11181        }
11182        sum
11183    }
11184}
11185
11186#[allow(clippy::too_many_arguments)]
11187fn qmatvec(
11188    q: &[u8],
11189    rep: &[u8],
11190    row_scale: &[f32],
11191    x: &[f32],
11192    col_field: &[f32],
11193    dtype: TensorDtype,
11194    rows: usize,
11195    cols: usize,
11196    out: &mut [f32],
11197    pool: Option<&Pool>,
11198) {
11199    debug_assert_eq!(out.len(), rows);
11200    #[cfg(not(target_arch = "aarch64"))]
11201    let _ = rep;
11202
11203    #[cfg(target_arch = "aarch64")]
11204    if sdot_enabled() {
11205        let act = if dtype == TensorDtype::Q8_2f {
11206            split_act_q8_2f(x, col_field)
11207        } else {
11208            split_act(x)
11209        };
11210        let out_addr = SendMut(out.as_mut_ptr());
11211        let run_range = |start: usize, end: usize| {
11212            q8_range_sdot(q, rep, row_scale, &act, cols, out_addr, start, end)
11213        };
11214        match pool {
11215            Some(pool) if rows >= 256 => pool.run_rows(rows, &run_range),
11216            _ => run_range(0, rows),
11217        }
11218        return;
11219    }
11220    // x86 A8W8 via AVX2 maddubs — same quantized-activation contract as
11221    // the SDOT path (CMF_AVX2=0 keeps the exact i8×f32 loop).
11222    #[cfg(target_arch = "x86_64")]
11223    if avx2_a8w8_enabled() {
11224        let act = if dtype == TensorDtype::Q8_2f {
11225            split_act_q8_2f(x, col_field)
11226        } else {
11227            split_act(x)
11228        };
11229        let out_addr = SendMut(out.as_mut_ptr());
11230        let run_range = |start: usize, end: usize| {
11231            q8_range_avx2(q, row_scale, &act, cols, out_addr, start, end)
11232        };
11233        match pool {
11234            Some(pool) if rows >= 256 => pool.run_rows(rows, &run_range),
11235            _ => run_range(0, rows),
11236        }
11237        return;
11238    }
11239
11240    prescale_with(x, col_field, dtype, 1, |xs| {
11241        let out_addr = SendMut(out.as_mut_ptr());
11242        let run_range = move |start: usize, end: usize| {
11243            for o in start..end {
11244                let v = dot_i8_f32(&q[o * cols..(o + 1) * cols], xs) * row_scale[o];
11245                // SAFETY: disjoint row ranges per worker.
11246                unsafe { *out_addr.at(o) = v };
11247            }
11248        };
11249        match pool {
11250            Some(pool) if rows >= 256 => pool.run_rows(rows, &run_range),
11251            _ => run_range(0, rows),
11252        }
11253    });
11254}
11255
11256#[allow(clippy::too_many_arguments)]
11257fn qmatvec2(
11258    q: &[u8],
11259    row_scale: &[f32],
11260    x1: &[f32],
11261    x2: &[f32],
11262    col_field: &[f32],
11263    dtype: TensorDtype,
11264    rows: usize,
11265    cols: usize,
11266    o1: &mut [f32],
11267    o2: &mut [f32],
11268    pool: Option<&Pool>,
11269) {
11270    #[cfg(target_arch = "aarch64")]
11271    if sdot_enabled() {
11272        let a1s = if dtype == TensorDtype::Q8_2f {
11273            split_act_q8_2f(x1, col_field)
11274        } else {
11275            split_act(x1)
11276        };
11277        let a2s = if dtype == TensorDtype::Q8_2f {
11278            split_act_q8_2f(x2, col_field)
11279        } else {
11280            split_act(x2)
11281        };
11282        let p1 = SendMut(o1.as_mut_ptr());
11283        let p2 = SendMut(o2.as_mut_ptr());
11284        let run_range = |start: usize, end: usize| {
11285            q8_range2_sdot(q, row_scale, &a1s, &a2s, cols, p1, p2, start, end)
11286        };
11287        match pool {
11288            Some(pool) if rows >= 256 => pool.run_rows(rows, &run_range),
11289            _ => run_range(0, rows),
11290        }
11291        return;
11292    }
11293    #[cfg(target_arch = "x86_64")]
11294    if avx2_a8w8_enabled() {
11295        let a1s = if dtype == TensorDtype::Q8_2f {
11296            split_act_q8_2f(x1, col_field)
11297        } else {
11298            split_act(x1)
11299        };
11300        let a2s = if dtype == TensorDtype::Q8_2f {
11301            split_act_q8_2f(x2, col_field)
11302        } else {
11303            split_act(x2)
11304        };
11305        let p1 = SendMut(o1.as_mut_ptr());
11306        let p2 = SendMut(o2.as_mut_ptr());
11307        let run_range = |start: usize, end: usize| {
11308            q8_range2_avx2(q, row_scale, &a1s, &a2s, cols, p1, p2, start, end)
11309        };
11310        match pool {
11311            Some(pool) if rows >= 256 => pool.run_rows(rows, &run_range),
11312            _ => run_range(0, rows),
11313        }
11314        return;
11315    }
11316
11317    prescale_with(x1, col_field, dtype, 1, |x1s| {
11318        prescale_with(x2, col_field, dtype, 2, |x2s| {
11319            let p1 = SendMut(o1.as_mut_ptr());
11320            let p2 = SendMut(o2.as_mut_ptr());
11321            let run_range = move |start: usize, end: usize| {
11322                for o in start..end {
11323                    let row = &q[o * cols..(o + 1) * cols];
11324                    let s1 = dot_i8_f32(row, x1s) * row_scale[o];
11325                    let s2 = dot_i8_f32(row, x2s) * row_scale[o];
11326                    // SAFETY: disjoint row ranges per worker.
11327                    unsafe {
11328                        *p1.at(o) = s1;
11329                        *p2.at(o) = s2;
11330                    }
11331                }
11332            };
11333            match pool {
11334                Some(pool) if rows >= 256 => pool.run_rows(rows, &run_range),
11335                _ => run_range(0, rows),
11336            }
11337        });
11338    });
11339}
11340
11341#[derive(Clone, Copy)]
11342struct SendMut(*mut f32);
11343unsafe impl Send for SendMut {}
11344unsafe impl Sync for SendMut {}
11345
11346impl SendMut {
11347    #[inline]
11348    fn at(self, i: usize) -> *mut f32 {
11349        unsafe { self.0.add(i) }
11350    }
11351}
11352
11353#[cfg(test)]
11354mod tests {
11355    /// `q8_round` must be `round().clamp(±127) as i8` bit for bit: every
11356    /// half-integer, their neighbours one ulp either side, the clamp
11357    /// boundary, huge values, infinities and NaN, plus a dense sweep.
11358    #[test]
11359    fn q8_round_is_round_clamp() {
11360        let reference = |t: f32| t.round().clamp(-127.0, 127.0) as i8;
11361        let mut probe = vec![
11362            0.0f32,
11363            -0.0,
11364            f32::NAN,
11365            f32::INFINITY,
11366            f32::NEG_INFINITY,
11367            f32::MAX,
11368            f32::MIN,
11369            1e30,
11370            -1e30,
11371            f32::MIN_POSITIVE,
11372            -f32::MIN_POSITIVE,
11373        ];
11374        for k in -300i32..=300 {
11375            let h = k as f32 * 0.5;
11376            let up = f32::from_bits(h.to_bits() + 1);
11377            let down = f32::from_bits(h.to_bits().wrapping_sub(1));
11378            for t in [h, up, down] {
11379                probe.push(t);
11380                probe.push(-t);
11381            }
11382        }
11383        let mut t = -140.0f32;
11384        while t < 140.0 {
11385            probe.push(t);
11386            t += 0.000_731;
11387        }
11388        for t in probe {
11389            assert_eq!(q8_round(t), reference(t), "t = {t:e} ({:#x})", t.to_bits());
11390        }
11391    }
11392
11393    use super::*;
11394
11395    #[test]
11396    fn q2tp_i8_dot_matches_exact_on_grid() {
11397        // On-grid activations (±1 → sx=1/127, xq=±127 dequantizes
11398        // exactly, no outliers) must make the integer path agree with
11399        // the exact scalar walk to f32 rounding.
11400        let (rows, cols) = (5, 64);
11401        let gpr = cols / GROUP_SIZE;
11402        // Synthetic codes plane + a flat ladder: scales_into is not under
11403        // test here, so drive dot_q2tp_row_i8 / q2tp_row_exact directly
11404        // with hand-made scales.
11405        let chunks: Vec<u8> = (0..rows * gpr * Q2TP_CHUNK)
11406            .map(|i| (i as u32).wrapping_mul(2654435761) as u8)
11407            .collect();
11408        let scales: Vec<f32> = (0..gpr).map(|g| 0.5 + g as f32 * 0.25).collect();
11409        let x: Vec<f32> = (0..cols)
11410            .map(|i| if i % 3 == 0 { -1.0 } else { 1.0 })
11411            .collect();
11412        let act = split_act(&x);
11413        assert!(
11414            act.outliers.is_empty(),
11415            "on-grid input must have no outliers"
11416        );
11417        let gsum = q1_group_sums(&act.xq, gpr);
11418        for r in 0..rows {
11419            let exact = q2tp_row_exact(&chunks, r, gpr, &x, &scales);
11420            let fast = dot_q2tp_row_i8(&chunks, r, gpr, &act.xq, &gsum, &scales) * act.sx;
11421            assert!(
11422                (exact - fast).abs() <= exact.abs() * 1e-5 + 1e-5,
11423                "row {r}: exact {exact} vs i8 {fast}"
11424            );
11425        }
11426    }
11427
11428    #[cfg(target_arch = "x86_64")]
11429    #[test]
11430    fn q2tp_avx2_dot_matches_scalar_for_random_patterns() {
11431        // Compare the release AVX2 integer dot against the scalar oracle over
11432        // arbitrary packed bytes/activation signs.  This guards the exact
11433        // table-load path used after rejecting a faster-looking decoder whose
11434        // full-checkpoint greedy output drifted.
11435        if !std::arch::is_x86_feature_detected!("avx2") {
11436            return;
11437        }
11438        let mut seed = 0x9e3779b9u32;
11439        let mut next = || {
11440            seed = seed.wrapping_mul(1664525).wrapping_add(1013904223);
11441            seed
11442        };
11443        for _ in 0..20_000 {
11444            let mut ch = [0u8; Q2TP_CHUNK];
11445            let mut x = [0i8; GROUP_SIZE];
11446            for b in &mut ch {
11447                *b = next() as u8;
11448            }
11449            for v in &mut x {
11450                *v = (next() >> 24) as i8;
11451            }
11452            let mut reference = 0i32;
11453            for (k, &b) in ch.iter().enumerate() {
11454                reference += (b & 3) as i32 * x[k * 4] as i32;
11455                reference += ((b >> 2) & 3) as i32 * x[k * 4 + 1] as i32;
11456                reference += ((b >> 4) & 3) as i32 * x[k * 4 + 2] as i32;
11457                reference += ((b >> 6) & 3) as i32 * x[k * 4 + 3] as i32;
11458            }
11459            // SAFETY: guarded by the runtime AVX2 feature check and fixed
11460            // 8-byte/32-byte slice lengths above.
11461            let got = unsafe { q2tp_code_dot_avx2(&ch, &x) };
11462            assert_eq!(got, reference, "packed q2 lane mismatch");
11463        }
11464    }
11465
11466    #[test]
11467    fn q2tp_affine_fuses_half_scale_correction_without_changing_raw_decode() {
11468        let (rows, cols) = (1usize, GROUP_SIZE);
11469        let mut bytes = vec![0u8; Q2TP_CHUNK + 4 + 1];
11470        // Repeating symbols 0,1,2,0 at unit scale.  q2tp's raw B is
11471        // (c-1.5), while the affine Prism operator is (c-1.0).
11472        bytes[..Q2TP_CHUNK].fill(0x24); // codes 0,1,2,0 in LSB-first order
11473        bytes[Q2TP_CHUNK..Q2TP_CHUNK + 2].copy_from_slice(&0u16.to_le_bytes());
11474        bytes[Q2TP_CHUNK + 2..Q2TP_CHUNK + 4].copy_from_slice(&0u16.to_le_bytes());
11475        bytes[Q2TP_CHUNK + 4] = 1; // dtype16 rung 1 = 1.0
11476        let x = vec![1.0f32; cols];
11477        let mut raw = vec![0.0f32; rows];
11478        let mut affine = vec![0.0f32; rows];
11479        q2tp_matvec_for_test(&bytes, &x, rows, cols, &mut raw);
11480        q2tp_affine_matvec_for_test(&bytes, &x, rows, cols, &mut affine);
11481        assert_eq!(raw, vec![-24.0]);
11482        assert_eq!(affine, vec![-8.0]);
11483        assert!((affine[0] - (raw[0] + 0.5 * cols as f32)).abs() < 1e-6);
11484    }
11485
11486    #[test]
11487    fn q8_row_dot_fast_matches_scalar() {
11488        // The per-arch fast dot must agree with the exact scalar oracle
11489        // (same contract the fused q8 FFN arm rides on).
11490        let cols = 96;
11491        let row: Vec<u8> = (0..cols)
11492            .map(|i| ((i * 37 % 251) - 125) as i8 as u8)
11493            .collect();
11494        let x: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.13).sin()).collect();
11495        let act = split_act(&x);
11496        let fast = q8_row_dot(&row, &act);
11497        let scalar = q8_row_dot_scalar(&row, &act);
11498        assert!(
11499            (fast - scalar).abs() <= scalar.abs() * 1e-5 + 1e-5,
11500            "fast {fast} vs scalar {scalar}"
11501        );
11502    }
11503
11504    #[test]
11505    fn f32_matvec_matches_matvec_rows_bitexact() {
11506        let (rows, cols) = (300, 40);
11507        let w: Vec<f32> = (0..rows * cols).map(|i| (i as f32 * 0.017).sin()).collect();
11508        let x: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.05).cos()).collect();
11509        let qt = QTensor::from_f32(w.clone(), rows, cols);
11510
11511        let mut a = vec![0.0f32; rows];
11512        matvec_rows(None, &w, &x, &mut a);
11513        let mut b = vec![0.0f32; rows];
11514        qt.matvec(&x, &mut b, None);
11515        assert_eq!(a, b);
11516    }
11517
11518    #[test]
11519    fn sdot_kernel_exact_on_grid() {
11520        // Activations already on the i8 grid (±1 with amax=1 → sx=1/127,
11521        // xq=±127 dequantizes EXACTLY) → the SDOT path must match the
11522        // exact f32 dot to float rounding. This isolates kernel
11523        // correctness from quantization noise.
11524        eprintln!("sdot_enabled = {}", sdot_enabled());
11525        let (rows, cols) = (9, 80); // odd rows → exercises 4-row + tail
11526        let w: Vec<u8> = (0..rows * cols)
11527            .map(|i| (((i * 37) % 251) as i32 - 125) as i8 as u8)
11528            .collect();
11529        let scales: Vec<f32> = (0..rows).map(|o| 0.005 + o as f32 * 0.001).collect();
11530        let x: Vec<f32> = (0..cols)
11531            .map(|i| match i % 3 {
11532                0 => 1.0,
11533                1 => -1.0,
11534                _ => 0.0,
11535            })
11536            .collect();
11537        let mut a = vec![0.0f32; rows];
11538        qmatvec(
11539            &w,
11540            &[],
11541            &scales,
11542            &x,
11543            &[],
11544            TensorDtype::Q8Row,
11545            rows,
11546            cols,
11547            &mut a,
11548            None,
11549        );
11550        for o in 0..rows {
11551            let mut acc = 0.0f32;
11552            for j in 0..cols {
11553                acc += (w[o * cols + j] as i8) as f32 * x[j];
11554            }
11555            let expect = acc * scales[o];
11556            assert!(
11557                (a[o] - expect).abs() < 1e-3 * expect.abs().max(1e-3),
11558                "row {o}: {} vs {expect}",
11559                a[o]
11560            );
11561        }
11562    }
11563
11564    #[test]
11565    fn q1_tbl_fast_path_matches_reference() {
11566        // gpr = 8 exercises the TBL pair-load fast loop, and the LAST
11567        // row's final 4-tile window trips the 4B-overread guard (the
11568        // payload ends exactly at the last tile) — both paths must
11569        // agree with the dequant reference.
11570        let (rows, cols) = (5, 256);
11571        let gpr = cols / GROUP_SIZE;
11572        let mut bytes = Vec::new();
11573        for t in 0..rows * gpr {
11574            let s = 0.007 + (t % 11) as f32 * 0.004;
11575            bytes.extend_from_slice(&cortiq_core::quant::f32_to_f16(s).to_le_bytes());
11576            for j in 0..4 {
11577                bytes.push(((t * 53 + j * 89 + 7) % 249) as u8);
11578            }
11579        }
11580        let x: Vec<f32> = (0..cols)
11581            .map(|i| if (i * 5) % 7 < 3 { 1.0 } else { -1.0 })
11582            .collect();
11583        let mut w = vec![0.0f32; rows * cols];
11584        cortiq_core::quant::dequant_q1(&bytes, &mut w);
11585        let mut got = vec![0.0f32; rows];
11586        q1_matvec(&bytes, &x, rows, cols, &mut got, None);
11587        for o in 0..rows {
11588            let expect: f32 = (0..cols).map(|j| w[o * cols + j] * x[j]).sum();
11589            assert!(
11590                (got[o] - expect).abs() < 1e-3 * expect.abs().max(1e-3),
11591                "row {o}: {} vs {expect}",
11592                got[o]
11593            );
11594        }
11595        // Blocked 1×4 batch (b=5: one quad + remainder) must equal the
11596        // single-matvec path bit-for-bit.
11597        let b = 5usize;
11598        let mut xs_all = Vec::new();
11599        for bi in 0..b {
11600            xs_all.extend(x.iter().map(|v| if bi % 2 == 0 { *v } else { -*v }));
11601        }
11602        let mut mm = vec![0.0f32; b * rows];
11603        q1_matmat(&bytes, &xs_all, b, rows, cols, &mut mm, None);
11604        for bi in 0..b {
11605            let mut single = vec![0.0f32; rows];
11606            q1_matvec(
11607                &bytes,
11608                &xs_all[bi * cols..(bi + 1) * cols],
11609                rows,
11610                cols,
11611                &mut single,
11612                None,
11613            );
11614            assert_eq!(&mm[bi * rows..(bi + 1) * rows], &single[..], "stream {bi}");
11615        }
11616    }
11617
11618    #[test]
11619    fn q1_kernels_match_exact_reference() {
11620        // Synthetic q1 payload: 6-byte tiles [f16 scale][4B bits].
11621        let (rows, cols) = (7, 96);
11622        let gpr = cols / GROUP_SIZE;
11623        let mut bytes = Vec::new();
11624        for t in 0..rows * gpr {
11625            let s = 0.01 + (t % 13) as f32 * 0.003;
11626            bytes.extend_from_slice(&cortiq_core::quant::f32_to_f16(s).to_le_bytes());
11627            for j in 0..4 {
11628                bytes.push(((t * 31 + j * 97) % 251) as u8);
11629            }
11630        }
11631        // On-grid activations (±1, amax 1) → the SDOT path is exact.
11632        let x: Vec<f32> = (0..cols)
11633            .map(|i| if i % 3 == 0 { 1.0 } else { -1.0 })
11634            .collect();
11635        // Reference through the core dequant.
11636        let mut w = vec![0.0f32; rows * cols];
11637        cortiq_core::quant::dequant_q1(&bytes, &mut w);
11638        let mut expect = vec![0.0f32; rows];
11639        for o in 0..rows {
11640            expect[o] = (0..cols).map(|j| w[o * cols + j] * x[j]).sum();
11641        }
11642        let mut got = vec![0.0f32; rows];
11643        q1_matvec(&bytes, &x, rows, cols, &mut got, None);
11644        for o in 0..rows {
11645            assert!(
11646                (got[o] - expect[o]).abs() < 1e-3 * expect[o].abs().max(1e-3),
11647                "row {o}: {} vs {}",
11648                got[o],
11649                expect[o]
11650            );
11651        }
11652        // Pair and batch paths agree with the single path.
11653        let x2: Vec<f32> = x.iter().map(|v| -v).collect();
11654        let (mut a1, mut a2) = (vec![0.0f32; rows], vec![0.0f32; rows]);
11655        q1_matvec2(&bytes, &x, &x2, rows, cols, &mut a1, &mut a2, None);
11656        assert_eq!(a1, got);
11657        let mut xs = x.clone();
11658        xs.extend_from_slice(&x2);
11659        let mut mm = vec![0.0f32; 2 * rows];
11660        q1_matmat(&bytes, &xs, 2, rows, cols, &mut mm, None);
11661        assert_eq!(&mm[..rows], got.as_slice());
11662        assert_eq!(&mm[rows..], a2.as_slice());
11663    }
11664
11665    #[test]
11666    fn repack_is_bit_identical() {
11667        // The interleaved-repack kernel must produce EXACTLY the same
11668        // bits as the mmap-layout kernel: integer accumulation is order-
11669        // exact, the f32 epilogue is identical. Odd rows exercise the
11670        // tail; direct range calls exercise unaligned pool splits.
11671        let (rows, cols) = (267, 96); // 66 groups + 3 tail rows, cols % 16 == 0
11672        let w: Vec<u8> = (0..rows * cols)
11673            .map(|i| (((i * 89) % 253) as i32 - 126) as i8 as u8)
11674            .collect();
11675        let scales: Vec<f32> = (0..rows).map(|o| 0.003 + o as f32 * 0.0007).collect();
11676        let x: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.37).sin() * 2.0).collect();
11677        let rep = q8_repack_layout(&w, rows, cols);
11678        // Group interleave round-trips.
11679        for g in 0..rows / 4 {
11680            for c in 0..cols / 16 {
11681                for lane in 0..4 {
11682                    assert_eq!(
11683                        &rep[g * 4 * cols + c * 64 + lane * 16
11684                            ..g * 4 * cols + c * 64 + lane * 16 + 16],
11685                        &w[(g * 4 + lane) * cols + c * 16..(g * 4 + lane) * cols + c * 16 + 16],
11686                    );
11687                }
11688            }
11689        }
11690        let mut a = vec![0.0f32; rows];
11691        qmatvec(
11692            &w,
11693            &[],
11694            &scales,
11695            &x,
11696            &[],
11697            TensorDtype::Q8Row,
11698            rows,
11699            cols,
11700            &mut a,
11701            None,
11702        );
11703        let mut b = vec![0.0f32; rows];
11704        qmatvec(
11705            &w,
11706            &rep,
11707            &scales,
11708            &x,
11709            &[],
11710            TensorDtype::Q8Row,
11711            rows,
11712            cols,
11713            &mut b,
11714            None,
11715        );
11716        assert_eq!(a, b, "full-range repack output diverged");
11717
11718        #[cfg(target_arch = "aarch64")]
11719        if sdot_enabled() {
11720            // Unaligned range split (pool workers get arbitrary bounds).
11721            let act = split_act(&x);
11722            let mut c1 = vec![0.0f32; rows];
11723            let mut c2 = vec![0.0f32; rows];
11724            q8_range_sdot(
11725                &w,
11726                &[],
11727                &scales,
11728                &act,
11729                cols,
11730                SendMut(c1.as_mut_ptr()),
11731                3,
11732                rows - 2,
11733            );
11734            q8_range_sdot(
11735                &w,
11736                &rep,
11737                &scales,
11738                &act,
11739                cols,
11740                SendMut(c2.as_mut_ptr()),
11741                3,
11742                rows - 2,
11743            );
11744            assert_eq!(c1, c2, "unaligned-range repack output diverged");
11745        }
11746    }
11747
11748    #[test]
11749    fn sdot_a8w8_noise_is_bounded() {
11750        // Off-grid activations: A8 quantization noise must stay small in
11751        // relative L2 over the whole output (realistic accuracy contract;
11752        // vmfcore measured argmax-identical decode on real models).
11753        let (rows, cols) = (16, 512);
11754        let w: Vec<u8> = (0..rows * cols)
11755            .map(|i| (((i * 37) % 251) as i32 - 125) as i8 as u8)
11756            .collect();
11757        let scales = vec![0.01f32; rows];
11758        let x: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.21).sin()).collect();
11759        let mut a = vec![0.0f32; rows];
11760        qmatvec(
11761            &w,
11762            &[],
11763            &scales,
11764            &x,
11765            &[],
11766            TensorDtype::Q8Row,
11767            rows,
11768            cols,
11769            &mut a,
11770            None,
11771        );
11772        let (mut num, mut den) = (0f64, 0f64);
11773        for o in 0..rows {
11774            let mut acc = 0.0f32;
11775            for j in 0..cols {
11776                acc += (w[o * cols + j] as i8) as f32 * x[j];
11777            }
11778            let expect = acc * scales[o];
11779            num += ((a[o] - expect) as f64).powi(2);
11780            den += (expect as f64).powi(2);
11781        }
11782        let rel = (num / den.max(1e-12)).sqrt();
11783        assert!(rel < 0.05, "A8W8 relative L2 error too high: {rel}");
11784    }
11785
11786    #[test]
11787    fn i8_dot_neon_matches_scalar() {
11788        let n = 100;
11789        let w: Vec<u8> = (0..n).map(|i| ((i * 37 + 11) % 251) as u8).collect();
11790        let x: Vec<f32> = (0..n).map(|i| (i as f32 * 0.13).sin()).collect();
11791        let mut scalar = 0.0f32;
11792        for j in 0..n {
11793            scalar += (w[j] as i8) as f32 * x[j];
11794        }
11795        let fast = dot_i8_f32(&w, &x);
11796        assert!((scalar - fast).abs() < 1e-3 * scalar.abs().max(1.0));
11797    }
11798
11799    /// Fused vbit matvec must match full dequant_vbit + dense matvec.
11800    #[test]
11801    fn vbitmatvec_matches_full_dequant() {
11802        let (rows, cols) = (6, 64);
11803        let ng = cols / GROUP_SIZE;
11804        // Hand-craft: bits per row, f16 scales, packed rows.
11805        let bits: Vec<u8> = vec![3, 4, 5, 6, 8, 4];
11806        let mut bytes = bits.clone();
11807        for g in 0..rows * ng {
11808            let s = 0.02 + 0.001 * g as f32;
11809            bytes.extend_from_slice(&cortiq_core::quant::f32_to_f16(s).to_le_bytes());
11810        }
11811        for r in 0..rows {
11812            let b = bits[r] as usize;
11813            let (mut acc, mut nb) = (0u64, 0usize);
11814            let mut rowbytes = Vec::new();
11815            for i in 0..cols {
11816                let v = ((i * 7 + r * 13) % (1 << b)) as u64;
11817                acc = (acc << b) | v;
11818                nb += b;
11819                while nb >= 8 {
11820                    nb -= 8;
11821                    rowbytes.push(((acc >> nb) & 0xFF) as u8);
11822                }
11823            }
11824            if nb > 0 {
11825                rowbytes.push(((acc << (8 - nb)) & 0xFF) as u8);
11826            }
11827            bytes.extend_from_slice(&rowbytes);
11828        }
11829        let x: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.19).sin()).collect();
11830
11831        let mut reference = vec![0f32; rows * cols];
11832        cortiq_core::quant::dequant_vbit(&bytes, rows, cols, &mut reference).unwrap();
11833        let mut expect = vec![0f32; rows];
11834        for r in 0..rows {
11835            expect[r] = reference[r * cols..(r + 1) * cols]
11836                .iter()
11837                .zip(&x)
11838                .map(|(w, xv)| w * xv)
11839                .sum();
11840        }
11841        let mut got = vec![0f32; rows];
11842        let offsets = vbit_row_offsets(&bytes, rows, cols);
11843        vbitmatvec(&bytes, &offsets, &x, rows, cols, &mut got, None);
11844        // SDOT path quantizes activations to i8 (A8W8): bounded noise,
11845        // same contract as q8 (exact path is pinned by CMF_SDOT=0 in
11846        // the golden-parity gate).
11847        let tol = if a8w8_enabled() { 6e-2 } else { 1e-4 };
11848        let scale = expect.iter().fold(0f32, |m, v| m.max(v.abs())).max(1e-6);
11849        for r in 0..rows {
11850            assert!(
11851                (got[r] - expect[r]).abs() < tol * scale,
11852                "row {r}: {} vs {}",
11853                got[r],
11854                expect[r]
11855            );
11856        }
11857    }
11858
11859    /// Fused q4 matvec must match the reference full-dequant + dense
11860    /// matvec bit-for-bit in structure (same f32 math, group order).
11861    /// vbit matmat: the blocked 1×4 leg must match the per-row path
11862    /// (paired env toggle; larger shape so both code paths engage).
11863    #[test]
11864    #[cfg(target_arch = "x86_64")]
11865    fn vbit_matmat_blocked_matches_per_row() {
11866        let (rows, cols, b) = (64usize, 128usize, 9usize);
11867        let ng = cols / GROUP_SIZE;
11868        let bits: Vec<u8> = (0..rows).map(|r| [3u8, 4, 5, 6][r % 4]).collect();
11869        let mut bytes = bits.clone();
11870        for g in 0..rows * ng {
11871            let sc = 0.02 + 0.0005 * g as f32;
11872            bytes.extend_from_slice(&cortiq_core::quant::f32_to_f16(sc).to_le_bytes());
11873        }
11874        for r in 0..rows {
11875            let bw = bits[r] as usize;
11876            let (mut acc, mut nb) = (0u64, 0usize);
11877            let mut rowbytes = Vec::new();
11878            for i in 0..cols {
11879                let v = ((i * 7 + r * 13) % (1 << bw)) as u64;
11880                acc = (acc << bw) | v;
11881                nb += bw;
11882                while nb >= 8 {
11883                    nb -= 8;
11884                    rowbytes.push(((acc >> nb) & 0xFF) as u8);
11885                }
11886            }
11887            if nb > 0 {
11888                rowbytes.push(((acc << (8 - nb)) & 0xFF) as u8);
11889            }
11890            bytes.extend_from_slice(&rowbytes);
11891        }
11892        let x: Vec<f32> = (0..b * cols)
11893            .map(|i| ((i * 13 + 7) % 97) as f32 / 97.0 - 0.5)
11894            .collect();
11895        let offsets = vbit_row_offsets(&bytes, rows, cols);
11896        let mut y_a = vec![0f32; b * rows];
11897        let mut y_b = vec![0f32; b * rows];
11898        unsafe { std::env::set_var("CMF_X86_BLOCKED", "1") };
11899        vbitmatmat(&bytes, &offsets, &x, b, rows, cols, &mut y_a, None);
11900        unsafe { std::env::set_var("CMF_X86_BLOCKED", "0") };
11901        vbitmatmat(&bytes, &offsets, &x, b, rows, cols, &mut y_b, None);
11902        unsafe { std::env::remove_var("CMF_X86_BLOCKED") };
11903        let max_d = y_a
11904            .iter()
11905            .zip(&y_b)
11906            .map(|(p, q)| (p - q).abs())
11907            .fold(0.0f32, f32::max);
11908        assert!(max_d < 1e-4, "vbit blocked ≠ per-row: max|Δ| = {max_d}");
11909    }
11910
11911    /// q4t blocked 1×4 (SDOT on ARM, AVX2 on x86) must equal the
11912    /// per-row path exactly: same nibble unpack, same group order,
11913    /// same f32 accumulation — batch == matvec bit-for-bit. b=9 covers
11914    /// two full 1×4 blocks plus a remainder through the single-row
11915    /// kernel. (Both paths produce identical output, so the shared
11916    /// CMF_X86_BLOCKED env var racing with other tests cannot flip
11917    /// the verdict — worst case both sides take the same path.)
11918    #[test]
11919    fn q4t_matmat_blocked_matches_per_row() {
11920        let (rows, cols, b) = (16usize, 64usize, 9usize);
11921        let gpr = cols / GROUP_SIZE;
11922        let mut bytes = vec![0u8; rows * gpr * Q4_TILE];
11923        for r in 0..rows {
11924            for g in 0..gpr {
11925                let t = (r * gpr + g) * Q4_TILE;
11926                let sc = 0.02 + 0.001 * (r * gpr + g) as f32;
11927                bytes[t..t + 2].copy_from_slice(&cortiq_core::quant::f32_to_f16(sc).to_le_bytes());
11928                for k in 0..16 {
11929                    bytes[t + 2 + k] = ((r * 31 + g * 7 + k * 13) % 251) as u8;
11930                }
11931            }
11932        }
11933        let x: Vec<f32> = (0..b * cols)
11934            .map(|i| ((i * 13 + 7) % 97) as f32 / 97.0 - 0.5)
11935            .collect();
11936        let mut y_blk = vec![0f32; b * rows];
11937        let mut y_row = vec![0f32; b * rows];
11938        unsafe { std::env::set_var("CMF_X86_BLOCKED", "1") };
11939        q4t_matmat(&bytes, &x, b, rows, cols, &mut y_blk, None);
11940        unsafe { std::env::set_var("CMF_X86_BLOCKED", "0") };
11941        q4t_matmat(&bytes, &x, b, rows, cols, &mut y_row, None);
11942        unsafe { std::env::remove_var("CMF_X86_BLOCKED") };
11943        assert_eq!(y_blk, y_row, "q4t blocked 1x4 ≠ per-row");
11944    }
11945
11946    /// The wide-batch Accelerate arm of q4t_matmat vs a brute-force
11947    /// f32 dequant matmul: both are f32 GEMMs, so only reduction
11948    /// order differs — tight tolerance.
11949    /// A synthetic q4tp payload: random nibbles plus a per-row ladder whose
11950    /// span varies row to row, so the codes actually exercise the full 0..31
11951    /// range rather than clustering on one rung.
11952    fn synth_q4tp(rows: usize, cols: usize) -> Vec<u8> {
11953        use cortiq_core::quant::{f32_to_f16, q4tp_code_stride, q4tp_put_code};
11954        let gpr = cols / GROUP_SIZE;
11955        let stride = q4tp_code_stride(gpr);
11956        let (params_off, codes_off, _) = q4tp_sections(rows, cols);
11957        let mut b = vec![0u8; codes_off + rows * stride];
11958        for r in 0..rows {
11959            for g in 0..gpr {
11960                let t = (r * gpr + g) * Q4TP_NIB;
11961                for k in 0..16 {
11962                    b[t + k] = ((r * 31 + g * 7 + k * 13) % 251) as u8;
11963                }
11964            }
11965            let lo = -6.0 - 0.03 * (r % 17) as f32;
11966            let step = 0.01 + 0.004 * (r % 11) as f32;
11967            let p = params_off + r * 4;
11968            b[p..p + 2].copy_from_slice(&f32_to_f16(lo).to_le_bytes());
11969            b[p + 2..p + 4].copy_from_slice(&f32_to_f16(step).to_le_bytes());
11970            let crow = &mut b[codes_off + r * stride..codes_off + (r + 1) * stride];
11971            for g in 0..gpr {
11972                q4tp_put_code(crow, g, (r * 5 + g * 3) % 32);
11973            }
11974        }
11975        b
11976    }
11977
11978    /// The same weights re-expressed as q4_tiled, so the proven kernel can
11979    /// be the reference: each tile stores the ladder scale its code selects.
11980    /// Only the f16 rounding of that scale separates the two payloads.
11981    fn q4tp_as_q4t(bytes: &[u8], rows: usize, cols: usize) -> Vec<u8> {
11982        let gpr = cols / GROUP_SIZE;
11983        let v = Q4tpView::new(bytes, rows, cols);
11984        let mut out = vec![0u8; rows * gpr * Q4_TILE];
11985        let mut sc = vec![0f32; gpr];
11986        for r in 0..rows {
11987            v.scales_into(r, gpr, &mut sc);
11988            for g in 0..gpr {
11989                let t = (r * gpr + g) * Q4_TILE;
11990                let s = sc[g];
11991                out[t..t + 2].copy_from_slice(&cortiq_core::quant::f32_to_f16(s).to_le_bytes());
11992                let src = (r * gpr + g) * Q4TP_NIB;
11993                out[t + 2..t + Q4_TILE].copy_from_slice(&v.nib[src..src + Q4TP_NIB]);
11994            }
11995        }
11996        out
11997    }
11998
11999    /// The exact (`CMF_SDOT=0`) path must reproduce `dequant_q4tp` to f32
12000    /// rounding — that scalar routine is the format's definition, and the
12001    /// kernels re-derive the scale from the ladder independently. Call the
12002    /// row kernel directly: `matmat` picks the int8 arm when a8w8 is on,
12003    /// so routing through it would test the other path by accident.
12004    #[test]
12005    fn q4tp_exact_path_matches_dequant_reference() {
12006        let (rows, cols) = (256usize, 512usize);
12007        let gpr = cols / GROUP_SIZE;
12008        let bytes = synth_q4tp(rows, cols);
12009        let mut w = vec![0f32; rows * cols];
12010        cortiq_core::quant::dequant_q4tp(&bytes, rows, cols, &mut w);
12011
12012        let x: Vec<f32> = (0..cols)
12013            .map(|i| ((i * 13 + 7) % 97) as f32 / 97.0 - 0.5)
12014            .collect();
12015        let v = Q4tpView::new(&bytes, rows, cols);
12016        let mut sc = vec![0f32; gpr];
12017        for r in 0..rows {
12018            v.scales_into(r, gpr, &mut sc);
12019            let got = q4tp_row_exact(v.nib, r, gpr, &x, &sc);
12020            let want: f32 = (0..cols).map(|c| w[r * cols + c] * x[c]).sum();
12021            // These dot products cancel down to ~1e-3 from terms of ~5e-2, so
12022            // the meaningful yardstick is the summed magnitude, not the result:
12023            // against the result any reordering of a 512-term f32 sum "fails".
12024            let mag: f32 = (0..cols).map(|c| (w[r * cols + c] * x[c]).abs()).sum();
12025            assert!(
12026                (got - want).abs() <= 1e-5 * mag,
12027                "row {r}: kernel {got} vs dequant {want}"
12028            );
12029        }
12030    }
12031
12032    /// The int8 (a8w8) path can't be checked against an f32 reference — the
12033    /// activation quantization dominates. Check it against the q4t kernel it
12034    /// was ported from instead, on payloads holding the same weights: that
12035    /// isolates exactly what the port could break (16 B stride, ladder
12036    /// lookup, nibble unpack) from what it deliberately shares.
12037    #[test]
12038    fn q4tp_matvec_matches_the_q4t_kernel_it_was_ported_from() {
12039        let (rows, cols) = (256usize, 512usize);
12040        let bytes = synth_q4tp(rows, cols);
12041        let twin = q4tp_as_q4t(&bytes, rows, cols);
12042        let x: Vec<f32> = (0..cols)
12043            .map(|i| ((i * 13 + 7) % 97) as f32 / 97.0 - 0.5)
12044            .collect();
12045
12046        let mut got = vec![0f32; rows];
12047        q4tp_matvec(&bytes, &x, rows, cols, &mut got, None);
12048        let mut want = vec![0f32; rows];
12049        q4t_matvec(&twin, &x, rows, cols, &mut want, None);
12050
12051        // Scale is f16 in the twin and f32 here, so allow that rounding on
12052        // top of the summed magnitude (same cancellation argument as above).
12053        let mut w = vec![0f32; rows * cols];
12054        cortiq_core::quant::dequant_q4tp(&bytes, rows, cols, &mut w);
12055        for r in 0..rows {
12056            let mag: f32 = (0..cols).map(|c| (w[r * cols + c] * x[c]).abs()).sum();
12057            assert!(
12058                (got[r] - want[r]).abs() <= 1e-3 * mag,
12059                "row {r}: q4tp {} vs q4t {}",
12060                got[r],
12061                want[r]
12062            );
12063        }
12064    }
12065
12066    /// `matmat` carries three arms (Accelerate, blocked int8 1x4, scalar).
12067    /// Batch 5 crosses the blocked kernel's stride, so this exercises the
12068    /// 1x4 path AND its scalar tail in one run — the blocked kernel is new
12069    /// code and its four accumulators are exactly what tends to go wrong.
12070    #[test]
12071    fn q4tp_matmat_matches_the_q4t_kernel_it_was_ported_from() {
12072        let (rows, cols, b) = (256usize, 512usize, 5usize);
12073        let bytes = synth_q4tp(rows, cols);
12074        let twin = q4tp_as_q4t(&bytes, rows, cols);
12075        let xs: Vec<f32> = (0..b * cols)
12076            .map(|i| ((i * 29 + 11) % 89) as f32 / 89.0 - 0.5)
12077            .collect();
12078
12079        let mut got = vec![0f32; b * rows];
12080        q4tp_matmat(&bytes, &xs, b, rows, cols, &mut got, None);
12081        let mut want = vec![0f32; b * rows];
12082        q4t_matmat(&twin, &xs, b, rows, cols, &mut want, None);
12083
12084        let mut w = vec![0f32; rows * cols];
12085        cortiq_core::quant::dequant_q4tp(&bytes, rows, cols, &mut w);
12086        for t in 0..b {
12087            for r in 0..rows {
12088                let mag: f32 = (0..cols)
12089                    .map(|c| (w[r * cols + c] * xs[t * cols + c]).abs())
12090                    .sum();
12091                let (g, wa) = (got[t * rows + r], want[t * rows + r]);
12092                assert!(
12093                    (g - wa).abs() <= 1e-3 * mag,
12094                    "batch {t} row {r}: q4tp {g} vs q4t {wa}"
12095                );
12096            }
12097        }
12098    }
12099
12100    #[test]
12101    fn q4tp_matvec2_matches_the_single_stream_kernel() {
12102        let (rows, cols) = (128usize, 256usize);
12103        let gpr = cols / GROUP_SIZE;
12104        let bytes = synth_q4tp(rows, cols);
12105        let xs: Vec<f32> = (0..2 * cols)
12106            .map(|i| ((i * 29 + 11) % 89) as f32 / 89.0 - 0.5)
12107            .collect();
12108
12109        let (mut o1, mut o2) = (vec![0f32; rows], vec![0f32; rows]);
12110        q4tp_matvec2(
12111            &bytes,
12112            &xs[..cols],
12113            &xs[cols..],
12114            rows,
12115            cols,
12116            &mut o1,
12117            &mut o2,
12118            None,
12119        );
12120
12121        // matvec2 takes the exact path for both streams, so the single-row
12122        // kernel is an exact reference — no tolerance for path differences.
12123        let v = Q4tpView::new(&bytes, rows, cols);
12124        let mut sc = vec![0f32; gpr];
12125        for r in 0..rows {
12126            v.scales_into(r, gpr, &mut sc);
12127            assert_eq!(o1[r], q4tp_row_exact(v.nib, r, gpr, &xs[..cols], &sc));
12128            assert_eq!(o2[r], q4tp_row_exact(v.nib, r, gpr, &xs[cols..], &sc));
12129        }
12130    }
12131
12132    /// q4tp must not COST speed — it exists to save bytes, and a format that
12133    /// trades 7% of a file for a slower model is a bad trade. This guard is
12134    /// here because correctness tests happily passed while `q4tp_matmat` was
12135    /// missing its int8 and Accelerate arms and the model ran 5x slower.
12136    /// Measured on M-series: 0.97-1.04x, i.e. parity (16 B tiles are better
12137    /// aligned than q4t's 18 B, which pays for the scale indirection).
12138    #[test]
12139    fn q4tp_matvec_keeps_pace_with_q4t() {
12140        let (rows, cols) = (4096usize, 3072usize);
12141        let bytes = synth_q4tp(rows, cols);
12142        let twin = q4tp_as_q4t(&bytes, rows, cols);
12143        let x: Vec<f32> = (0..cols).map(|i| (i % 97) as f32 / 97.0 - 0.5).collect();
12144        let mut o = vec![0f32; rows];
12145        let n = 12;
12146        let mut best = (f64::MAX, f64::MAX);
12147        // Interleaved A/B, minimum statistic: this machine throttles, and a
12148        // mean over a thermal ramp reliably indicts whichever ran second.
12149        for _ in 0..3 {
12150            let t0 = std::time::Instant::now();
12151            for _ in 0..n {
12152                q4t_matvec(&twin, &x, rows, cols, &mut o, None);
12153            }
12154            best.0 = best.0.min(t0.elapsed().as_secs_f64());
12155            let t0 = std::time::Instant::now();
12156            for _ in 0..n {
12157                q4tp_matvec(&bytes, &x, rows, cols, &mut o, None);
12158            }
12159            best.1 = best.1.min(t0.elapsed().as_secs_f64());
12160        }
12161        let ratio = best.1 / best.0;
12162        println!(
12163            "q4t {:.3} ms | q4tp {:.3} ms | {ratio:.2}x",
12164            best.0 * 1e3 / n as f64,
12165            best.1 * 1e3 / n as f64
12166        );
12167        assert!(ratio < 2.0, "q4tp matvec {ratio:.2}x slower than q4t");
12168    }
12169
12170    #[cfg(target_os = "macos")]
12171    #[test]
12172    fn q4t_matmat_accel_matches_dequant_reference() {
12173        if !accel_gemm_enabled() {
12174            return; // CMF_ACCEL=0
12175        }
12176        let (rows, cols, b) = (512usize, 1024usize, 8usize); // ≥500K → accel arm
12177        let gpr = cols / GROUP_SIZE;
12178        let mut bytes = vec![0u8; rows * gpr * Q4_TILE];
12179        for r in 0..rows {
12180            for g in 0..gpr {
12181                let t = (r * gpr + g) * Q4_TILE;
12182                let sc = 0.02 + 0.0005 * ((r * gpr + g) % 64) as f32;
12183                bytes[t..t + 2].copy_from_slice(&cortiq_core::quant::f32_to_f16(sc).to_le_bytes());
12184                for k in 0..16 {
12185                    bytes[t + 2 + k] = ((r * 31 + g * 7 + k * 13) % 251) as u8;
12186                }
12187            }
12188        }
12189        let x: Vec<f32> = (0..b * cols)
12190            .map(|i| ((i * 13 + 7) % 97) as f32 / 97.0 - 0.5)
12191            .collect();
12192        let mut got = vec![0f32; b * rows];
12193        q4t_matmat(&bytes, &x, b, rows, cols, &mut got, None);
12194        // Brute-force reference off the same tiles.
12195        let mut w = vec![0f32; rows * cols];
12196        for r in 0..rows {
12197            for g in 0..gpr {
12198                let t = (r * gpr + g) * Q4_TILE;
12199                let s = f16_to_f32(u16::from_le_bytes([bytes[t], bytes[t + 1]]));
12200                for (k, &bb) in bytes[t + 2..t + Q4_TILE].iter().enumerate() {
12201                    w[r * cols + g * GROUP_SIZE + k * 2] = ((bb & 0x0F) as f32 - 8.0) * s;
12202                    w[r * cols + g * GROUP_SIZE + k * 2 + 1] =
12203                        (((bb >> 4) & 0x0F) as f32 - 8.0) * s;
12204                }
12205            }
12206        }
12207        for bi in 0..b {
12208            for r in 0..rows {
12209                let want: f32 = (0..cols).map(|j| x[bi * cols + j] * w[r * cols + j]).sum();
12210                let d = (got[bi * rows + r] - want).abs();
12211                assert!(
12212                    d <= want.abs().max(1.0) * 1e-4,
12213                    "accel q4t GEMM diverged at ({bi},{r}): {} vs {want}",
12214                    got[bi * rows + r]
12215                );
12216            }
12217        }
12218    }
12219
12220    #[test]
12221    fn q4matvec_matches_full_dequant() {
12222        let (rows, cols) = (8, 64);
12223        let groups = rows * cols / GROUP_SIZE;
12224        // Hand-craft a q4_block blob: nibbles then f16 scales.
12225        let mut bytes = Vec::with_capacity(groups * 16 + groups * 2);
12226        for i in 0..groups * 16 {
12227            bytes.push((((i * 7 + 3) % 256) & 0xFF) as u8);
12228        }
12229        for g in 0..groups {
12230            let s = 0.01 + 0.003 * g as f32;
12231            bytes.extend_from_slice(&cortiq_core::quant::f32_to_f16(s).to_le_bytes());
12232        }
12233        let x: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.17).sin()).collect();
12234
12235        let mut reference = vec![0.0f32; rows * cols];
12236        cortiq_core::quant::dequant_q4_block(&bytes, &mut reference);
12237        let mut expect = vec![0.0f32; rows];
12238        for r in 0..rows {
12239            expect[r] = reference[r * cols..(r + 1) * cols]
12240                .iter()
12241                .zip(&x)
12242                .map(|(w, xv)| w * xv)
12243                .sum();
12244        }
12245
12246        let mut got = vec![0.0f32; rows];
12247        q4matvec(&bytes, &x, rows, cols, &mut got, None);
12248        // SDOT path quantizes activations to i8 (A8W8): bounded noise,
12249        // same contract as q8/vbit (exact path is pinned by CMF_SDOT=0
12250        // in the golden-parity gate).
12251        let tol = if a8w8_enabled() { 6e-2 } else { 1e-4 };
12252        let scale = expect.iter().fold(0f32, |m, v| m.max(v.abs())).max(1.0);
12253        for r in 0..rows {
12254            assert!(
12255                (got[r] - expect[r]).abs() < tol * scale,
12256                "row {r}: {} vs {}",
12257                got[r],
12258                expect[r]
12259            );
12260        }
12261    }
12262
12263    /// Fused two-input vbit matvec must equal two single matvecs exactly
12264    /// (same per-lane accumulation order on both scalar and SDOT paths).
12265    #[test]
12266    fn vbitmatvec2_equals_two_singles() {
12267        let (rows, cols) = (6, 64);
12268        let ng = cols / GROUP_SIZE;
12269        let bits: Vec<u8> = vec![3, 4, 5, 6, 8, 4];
12270        let mut bytes = bits.clone();
12271        for g in 0..rows * ng {
12272            let s = 0.02 + 0.001 * g as f32;
12273            bytes.extend_from_slice(&cortiq_core::quant::f32_to_f16(s).to_le_bytes());
12274        }
12275        for r in 0..rows {
12276            let b = bits[r] as usize;
12277            let (mut acc, mut nb) = (0u64, 0usize);
12278            let mut rowbytes = Vec::new();
12279            for i in 0..cols {
12280                let v = ((i * 7 + r * 13) % (1 << b)) as u64;
12281                acc = (acc << b) | v;
12282                nb += b;
12283                while nb >= 8 {
12284                    nb -= 8;
12285                    rowbytes.push(((acc >> nb) & 0xFF) as u8);
12286                }
12287            }
12288            if nb > 0 {
12289                rowbytes.push(((acc << (8 - nb)) & 0xFF) as u8);
12290            }
12291            bytes.extend_from_slice(&rowbytes);
12292        }
12293        let x1: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.19).sin()).collect();
12294        let x2: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.11).cos()).collect();
12295        let offsets = vbit_row_offsets(&bytes, rows, cols);
12296
12297        let (mut a1, mut a2) = (vec![0f32; rows], vec![0f32; rows]);
12298        vbitmatvec(&bytes, &offsets, &x1, rows, cols, &mut a1, None);
12299        vbitmatvec(&bytes, &offsets, &x2, rows, cols, &mut a2, None);
12300        let (mut b1, mut b2) = (vec![0f32; rows], vec![0f32; rows]);
12301        vbitmatvec2(
12302            &bytes, &offsets, &x1, &x2, rows, cols, &mut b1, &mut b2, None,
12303        );
12304        assert_eq!(a1, b1, "fused vbit lane 1 must be bit-identical");
12305        assert_eq!(a2, b2, "fused vbit lane 2 must be bit-identical");
12306    }
12307
12308    /// Fused two-input q4 matvec must equal two single matvecs exactly.
12309    #[test]
12310    fn q4matvec2_equals_two_singles() {
12311        let (rows, cols) = (8, 128);
12312        let groups = rows * cols / GROUP_SIZE;
12313        let mut bytes = Vec::with_capacity(groups * 16 + groups * 2);
12314        for i in 0..groups * 16 {
12315            bytes.push((((i * 7 + 3) % 256) & 0xFF) as u8);
12316        }
12317        for g in 0..groups {
12318            let s = 0.01 + 0.003 * g as f32;
12319            bytes.extend_from_slice(&cortiq_core::quant::f32_to_f16(s).to_le_bytes());
12320        }
12321        // Include an outlier channel so the SDOT correction path is
12322        // exercised in the pair kernel too.
12323        let mut x1: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.17).sin()).collect();
12324        x1[9] = 250.0;
12325        let x2: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.23).cos()).collect();
12326
12327        let (mut a1, mut a2) = (vec![0f32; rows], vec![0f32; rows]);
12328        q4matvec(&bytes, &x1, rows, cols, &mut a1, None);
12329        q4matvec(&bytes, &x2, rows, cols, &mut a2, None);
12330        let (mut b1, mut b2) = (vec![0f32; rows], vec![0f32; rows]);
12331        q4matvec2(&bytes, &x1, &x2, rows, cols, &mut b1, &mut b2, None);
12332        assert_eq!(a1, b1, "fused q4 lane 1 must be bit-identical");
12333        assert_eq!(a2, b2, "fused q4 lane 2 must be bit-identical");
12334    }
12335
12336    /// Multi-matrix job must equal separate matvecs exactly — same
12337    /// kernels, only the dispatch is fused.
12338    #[test]
12339    fn matvec_many_equals_separate_matvecs() {
12340        use crate::pool::Pool;
12341        let (r1, r2, cols) = (300, 200, 64);
12342        let mk = |salt: usize, rows: usize| {
12343            QTensor::from_f32(
12344                (0..rows * cols)
12345                    .map(|i| ((i * 7 + salt) % 97) as f32 / 97.0 - 0.5)
12346                    .collect(),
12347                rows,
12348                cols,
12349            )
12350        };
12351        let (a, b) = (mk(1, r1), mk(5, r2));
12352        let x: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.11).sin()).collect();
12353        let pool = Pool::new(3);
12354
12355        let (mut ea, mut eb) = (vec![0f32; r1], vec![0f32; r2]);
12356        a.matvec(&x, &mut ea, Some(&pool));
12357        b.matvec(&x, &mut eb, Some(&pool));
12358        let (mut ga, mut gb) = (vec![0f32; r1], vec![0f32; r2]);
12359        QTensor::matvec_many([&a, &b], &x, [&mut ga, &mut gb], Some(&pool));
12360        assert_eq!(ea, ga, "fused multi-matrix lane 1 must be bit-identical");
12361        assert_eq!(eb, gb, "fused multi-matrix lane 2 must be bit-identical");
12362    }
12363
12364    /// The public Q4TP operator must take the real mapped matvec_many arm,
12365    /// rather than the F32 fallback above.  Build a tiny valid CMF so both
12366    /// handles retain their mmap payloads, then compare the fused dispatch
12367    /// with two ordinary mapped matvec calls bit-for-bit.
12368    #[test]
12369    fn q4tp_matvec_many_equals_separate_matvecs() {
12370        use crate::pool::Pool;
12371        use cortiq_core::{CMF_VERSION, CmfHeader, CmfModel, QuantType, TensorSpec};
12372
12373        let (r1, r2, cols) = (300usize, 200usize, 64usize);
12374        let arch: cortiq_core::ModelArch = serde_json::from_value(serde_json::json!({
12375            "arch_name": "tiny-q4tp",
12376            "hidden_size": cols,
12377            "intermediate_size": cols * 2,
12378            "num_layers": 1,
12379            "num_attention_heads": 2,
12380            "num_kv_heads": 1,
12381            "head_dim": 32,
12382            "vocab_size": r1,
12383            "layer_types": ["FullAttention"],
12384            "rms_norm_eps": 1e-6,
12385            "max_position_embeddings": 8,
12386            "linear_conv_kernel_dim": 0,
12387            "linear_num_key_heads": 0,
12388            "linear_num_value_heads": 0
12389        }))
12390        .unwrap();
12391        let header = CmfHeader {
12392            format: "cmf".into(),
12393            version: CMF_VERSION,
12394            arch,
12395            quant_type: QuantType::Q4Block,
12396            provenance: None,
12397            tokenizer_config: None,
12398            section_hashes: None,
12399            skills: Vec::new(),
12400            shard: None,
12401            calibration: None,
12402            routing: None,
12403            genome: None,
12404            lineage: Vec::new(),
12405            router: None,
12406            segments: Vec::new(),
12407        };
12408        let specs = [
12409            TensorSpec {
12410                name: "q".into(),
12411                dtype: TensorDtype::Q4TiledP,
12412                shape: vec![r1, cols],
12413                data: synth_q4tp(r1, cols),
12414            },
12415            TensorSpec {
12416                name: "kv".into(),
12417                dtype: TensorDtype::Q4TiledP,
12418                shape: vec![r2, cols],
12419                data: synth_q4tp(r2, cols),
12420            },
12421        ];
12422        let dir = std::env::temp_dir().join(format!("cmf-q4tp-many-{}", std::process::id()));
12423        std::fs::create_dir_all(&dir).unwrap();
12424        let path = dir.join("m.cmf");
12425        CmfModel::write(&path, &header, &specs, None, None).unwrap();
12426        let model = Arc::new(CmfModel::open(&path).unwrap());
12427        let (a, b) = (
12428            QTensor::from_model(&model, "q").unwrap(),
12429            QTensor::from_model(&model, "kv").unwrap(),
12430        );
12431        assert_eq!(a.model_dtype(), Some(TensorDtype::Q4TiledP));
12432        assert_eq!(b.model_dtype(), Some(TensorDtype::Q4TiledP));
12433        let x: Vec<f32> = (0..cols)
12434            .map(|i| ((i * 17 + 3) % 97) as f32 / 97.0 - 0.5)
12435            .collect();
12436        let pool = Pool::new(3);
12437        let (mut ea, mut eb) = (vec![0.0f32; r1], vec![0.0f32; r2]);
12438        a.matvec(&x, &mut ea, Some(&pool));
12439        b.matvec(&x, &mut eb, Some(&pool));
12440        let (mut ga, mut gb) = (vec![0.0f32; r1], vec![0.0f32; r2]);
12441        QTensor::matvec_many([&a, &b], &x, [&mut ga, &mut gb], Some(&pool));
12442        assert_eq!(ea, ga, "Q4TP fused lane 1 must be bit-identical");
12443        assert_eq!(eb, gb, "Q4TP fused lane 2 must be bit-identical");
12444        let _ = std::fs::remove_dir_all(&dir);
12445    }
12446
12447    /// The MiMo speculative verify's kernels: several tokens' MoE through
12448    /// `moe_gate_up_rows` / `moe_down_rows` (+ the caller's route-order sum)
12449    /// is bit-identical to each token alone through `moe_gate_up_many` /
12450    /// `moe_down_many` (decode), and a row-exact `q4tp_matmat` of five
12451    /// tokens (wide enough for the blocked tiles) equals five matvecs.
12452    #[test]
12453    fn multi_token_moe_rows_equal_single_token_decode() {
12454        use crate::pool::Pool;
12455        use cortiq_core::{CMF_VERSION, CmfHeader, CmfModel, QuantType, TensorSpec};
12456
12457        let (h, inter, ne) = (64usize, 128usize, 3usize);
12458        let arch: cortiq_core::ModelArch = serde_json::from_value(serde_json::json!({
12459            "arch_name": "tiny-q4tp-moe",
12460            "hidden_size": h,
12461            "intermediate_size": inter,
12462            "num_layers": 1,
12463            "num_attention_heads": 2,
12464            "num_kv_heads": 1,
12465            "head_dim": 32,
12466            "vocab_size": 8,
12467            "layer_types": ["FullAttention"],
12468            "rms_norm_eps": 1e-6,
12469            "max_position_embeddings": 8,
12470            "linear_conv_kernel_dim": 0,
12471            "linear_num_key_heads": 0,
12472            "linear_num_value_heads": 0
12473        }))
12474        .unwrap();
12475        let header = CmfHeader {
12476            format: "cmf".into(),
12477            version: CMF_VERSION,
12478            arch,
12479            quant_type: QuantType::Q4Block,
12480            provenance: None,
12481            tokenizer_config: None,
12482            section_hashes: None,
12483            skills: Vec::new(),
12484            shard: None,
12485            calibration: None,
12486            routing: None,
12487            genome: None,
12488            lineage: Vec::new(),
12489            router: None,
12490            segments: Vec::new(),
12491        };
12492        let mut specs = Vec::new();
12493        for e in 0..ne {
12494            for (k, (n, r, c)) in [("g", inter, h), ("u", inter, h), ("d", h, inter)]
12495                .into_iter()
12496                .enumerate()
12497            {
12498                // Distinct experts: perturb only the nibble plane (any byte
12499                // is a valid pair of codes; the ladder stays intact).
12500                let mut data = synth_q4tp(r, c);
12501                for (i, byte) in data[..r * (c / GROUP_SIZE) * Q4TP_NIB]
12502                    .iter_mut()
12503                    .enumerate()
12504                {
12505                    *byte ^= ((i * (e * 3 + k + 1)) % 251) as u8;
12506                }
12507                specs.push(TensorSpec {
12508                    name: format!("{n}{e}"),
12509                    dtype: TensorDtype::Q4TiledP,
12510                    shape: vec![r, c],
12511                    data,
12512                });
12513            }
12514        }
12515        let dir = std::env::temp_dir().join(format!(
12516            "cmf-moe-rows-{}-{}",
12517            std::process::id(),
12518            FLOAT_ACTIVATIONS.get()
12519        ));
12520        std::fs::create_dir_all(&dir).unwrap();
12521        let path = dir.join("m.cmf");
12522        CmfModel::write(&path, &header, &specs, None, None).unwrap();
12523        let model = Arc::new(CmfModel::open(&path).unwrap());
12524        let t = |n: String| QTensor::from_model(&model, &n).unwrap();
12525        let g: Vec<QTensor> = (0..ne).map(|e| t(format!("g{e}"))).collect();
12526        let u: Vec<QTensor> = (0..ne).map(|e| t(format!("u{e}"))).collect();
12527        let d: Vec<QTensor> = (0..ne).map(|e| t(format!("d{e}"))).collect();
12528        let b = 4usize;
12529        let mut xs: Vec<f32> = (0..b * h)
12530            .map(|i| ((i * 31 + 7) % 89) as f32 / 89.0 - 0.5)
12531            .collect();
12532        xs[5] = 9.0; // an activation outlier on token 0
12533        // Token -> (experts in route order, weights).
12534        let routes: Vec<(Vec<usize>, Vec<f32>)> = vec![
12535            (vec![2, 0], vec![0.6, 0.4]),
12536            (vec![0, 1, 2], vec![0.2, 0.5, 0.3]),
12537            (vec![1], vec![1.0]),
12538            (vec![2, 1, 0], vec![0.25, 0.25, 0.5]),
12539        ];
12540        let pool = Pool::new(3);
12541        // Decode reference, token by token.
12542        let mut want = vec![0f32; b * h];
12543        for (tk, (idx, w)) in routes.iter().enumerate() {
12544            let x = &xs[tk * h..(tk + 1) * h];
12545            let pairs: Vec<(&QTensor, &QTensor)> = idx.iter().map(|&e| (&g[e], &u[e])).collect();
12546            let mut gs: Vec<Vec<f32>> = idx.iter().map(|_| vec![0f32; inter]).collect();
12547            assert!(QTensor::moe_gate_up_many(&pairs, x, &mut gs, Some(&pool)));
12548            if FLOAT_ACTIVATIONS.get() {
12549                for (slot, &e) in idx.iter().enumerate() {
12550                    let (mut gate, mut up) = (vec![0.0; inter], vec![0.0; inter]);
12551                    g[e].matvec(x, &mut gate, Some(&pool));
12552                    u[e].matvec(x, &mut up, Some(&pool));
12553                    for (v, u) in gate.iter_mut().zip(up) {
12554                        *v = (*v / (1.0 + (-*v).exp())) * u;
12555                    }
12556                    assert_eq!(gs[slot], gate, "float gate/up must equal ordinary matvecs");
12557                }
12558            }
12559            let downs: Vec<&QTensor> = idx.iter().map(|&e| &d[e]).collect();
12560            assert!(QTensor::moe_down_many(
12561                &downs,
12562                &gs,
12563                w,
12564                &mut want[tk * h..(tk + 1) * h],
12565                Some(&pool)
12566            ));
12567        }
12568        if FLOAT_ACTIVATIONS.get() {
12569            for (tk, (idx, w)) in routes.iter().enumerate() {
12570                let mut scalar = vec![0.0; h];
12571                for (&e, &weight) in idx.iter().zip(w) {
12572                    let (mut gate, mut up, mut down) =
12573                        (vec![0.0; inter], vec![0.0; inter], vec![0.0; h]);
12574                    g[e].matvec(&xs[tk * h..(tk + 1) * h], &mut gate, Some(&pool));
12575                    u[e].matvec(&xs[tk * h..(tk + 1) * h], &mut up, Some(&pool));
12576                    for (v, u) in gate.iter_mut().zip(up) {
12577                        *v = (*v / (1.0 + (-*v).exp())) * u;
12578                    }
12579                    d[e].matvec(&gate, &mut down, Some(&pool));
12580                    for (v, d) in scalar.iter_mut().zip(down) {
12581                        *v += weight * d;
12582                    }
12583                }
12584                assert_eq!(
12585                    &want[tk * h..(tk + 1) * h],
12586                    scalar,
12587                    "float many equals scalar experts"
12588                );
12589            }
12590        }
12591        // All four tokens at once, grouped by expert.
12592        let mut experts: Vec<usize> = Vec::new();
12593        let mut groups: Vec<Vec<usize>> = Vec::new();
12594        for (tk, (idx, _)) in routes.iter().enumerate() {
12595            for &e in idx {
12596                match experts.iter().position(|&x| x == e) {
12597                    Some(k) => groups[k].push(tk),
12598                    None => {
12599                        experts.push(e);
12600                        groups.push(vec![tk]);
12601                    }
12602                }
12603            }
12604        }
12605        let n_pairs: usize = groups.iter().map(|g| g.len()).sum();
12606        let pairs: Vec<(&QTensor, &QTensor)> = experts.iter().map(|&e| (&g[e], &u[e])).collect();
12607        let mut gs: Vec<Vec<f32>> = (0..n_pairs).map(|_| vec![0f32; inter]).collect();
12608        assert!(QTensor::moe_gate_up_rows(
12609            &pairs,
12610            &groups,
12611            &xs,
12612            &mut gs,
12613            Some(&pool)
12614        ));
12615        let downs: Vec<&QTensor> = experts.iter().map(|&e| &d[e]).collect();
12616        let lens: Vec<usize> = groups.iter().map(|g| g.len()).collect();
12617        let mut ds: Vec<Vec<f32>> = (0..n_pairs).map(|_| vec![0f32; h]).collect();
12618        assert!(QTensor::moe_down_rows(
12619            &downs,
12620            &lens,
12621            &gs,
12622            &mut ds,
12623            Some(&pool)
12624        ));
12625        let slot = |tk: usize, e: usize| {
12626            let k = experts.iter().position(|&x| x == e).unwrap();
12627            groups[..k].iter().map(|g| g.len()).sum::<usize>()
12628                + groups[k].iter().position(|&x| x == tk).unwrap()
12629        };
12630        let mut got = vec![0f32; b * h];
12631        for (tk, (idx, w)) in routes.iter().enumerate() {
12632            for i in 0..h {
12633                let mut acc = 0f32;
12634                for (&e, &we) in idx.iter().zip(w) {
12635                    acc += we * ds[slot(tk, e)][i];
12636                }
12637                got[tk * h + i] = acc;
12638            }
12639        }
12640        assert!(want.iter().any(|v| *v != 0.0));
12641        assert_eq!(
12642            want.iter().map(|v| v.to_bits()).collect::<Vec<_>>(),
12643            got.iter().map(|v| v.to_bits()).collect::<Vec<_>>(),
12644            "multi-token MoE must equal decode bit for bit"
12645        );
12646
12647        // Row-exact q4tp_matmat: five tokens (a blocked 1x4 tile + tail
12648        // otherwise) equal five matvecs.
12649        let b5 = 5usize;
12650        let x5: Vec<f32> = (0..b5 * h)
12651            .map(|i| ((i * 13 + 5) % 71) as f32 / 71.0 - 0.5)
12652            .collect();
12653        let mut mm = vec![0f32; b5 * inter];
12654        row_exact_scope(|| g[1].matmat(&x5, b5, &mut mm, Some(&pool)));
12655        for tk in 0..b5 {
12656            let mut mv = vec![0f32; inter];
12657            g[1].matvec(&x5[tk * h..(tk + 1) * h], &mut mv, Some(&pool));
12658            assert_eq!(
12659                mv.iter().map(|v| v.to_bits()).collect::<Vec<_>>(),
12660                mm[tk * inter..(tk + 1) * inter]
12661                    .iter()
12662                    .map(|v| v.to_bits())
12663                    .collect::<Vec<_>>(),
12664                "row-exact matmat token {tk}"
12665            );
12666        }
12667        // Other concurrent tests/requests may still hold the shared mode.
12668        // Nested, overlapping and unwind restoration is checked separately.
12669        let _ = std::fs::remove_dir_all(&dir);
12670    }
12671
12672    /// The row-exact fix must not touch the fast path. Outside the scope
12673    /// `q4tp_matmat` has to produce exactly what it did before, and on ARM
12674    /// "before" is spelled out below: the tuned 1x4 SDOT tile for every
12675    /// four columns and the single-row kernel for the tail. Inside the
12676    /// scope every column equals its token's matvec. `q4tp_matmat_with`
12677    /// takes the mode as an argument, so a concurrent test holding the
12678    /// shared scope cannot flip it under this one.
12679    #[test]
12680    fn q4tp_matmat_fast_path_unchanged_outside_row_exact() {
12681        use crate::pool::Pool;
12682        use std::sync::atomic::Ordering::Relaxed;
12683        let _alt = Q4TP_ALT_TEST_LOCK.lock().unwrap_or_else(|e| e.into_inner());
12684        // The tuned ARM shape, which is also what an unset switch picks
12685        // unless CMF_Q4TP_V1 is exported.
12686        Q4TP_ALT.store(2, Relaxed);
12687        let pool = Pool::new(3);
12688        // Under 500k cells, so macOS keeps the matmat off the AMX; the
12689        // second shape runs across pool workers (rows >= 256) with 32
12690        // groups of accumulation and a tail after two 1x4 tiles.
12691        for &(rows, cols, b) in &[(64usize, 256usize, 7usize), (320, 1024, 9)] {
12692            let bytes = synth_q4tp(rows, cols);
12693            let mut xs: Vec<f32> = (0..b * cols)
12694                .map(|i| ((i * 29 + 11) % 83) as f32 / 83.0 - 0.5)
12695                .collect();
12696            xs[3] = 7.5; // an activation outlier on token 0
12697            let run = |exact: bool| {
12698                let mut out = vec![0f32; b * rows];
12699                q4tp_matmat_with(&bytes, &xs, b, rows, cols, &mut out, Some(&pool), exact);
12700                out
12701            };
12702            let (fast, exact) = (run(false), run(true));
12703            #[cfg(not(target_arch = "aarch64"))]
12704            let _ = fast;
12705            let bits = |v: &[f32]| v.iter().map(|x| x.to_bits()).collect::<Vec<_>>();
12706            let mut matvecs = vec![0f32; b * rows];
12707            for (bi, o) in matvecs.chunks_mut(rows).enumerate() {
12708                q4tp_matvec(
12709                    &bytes,
12710                    &xs[bi * cols..(bi + 1) * cols],
12711                    rows,
12712                    cols,
12713                    o,
12714                    Some(&pool),
12715                );
12716            }
12717            assert!(matvecs.iter().any(|v| *v != 0.0));
12718            assert_eq!(
12719                bits(&exact),
12720                bits(&matvecs),
12721                "{rows}x{cols} b={b}: row-exact matmat must equal per-token matvecs"
12722            );
12723            #[cfg(target_arch = "aarch64")]
12724            {
12725                // The pre-fix ARM loop, cell for cell.
12726                let gpr = cols / GROUP_SIZE;
12727                let v = Q4tpView::new(&bytes, rows, cols);
12728                let mut old = vec![0f32; b * rows];
12729                let mut sc = vec![0f32; gpr];
12730                let a8w8 = a8w8_enabled();
12731                let blocked = sdot_enabled() && blocked_enabled();
12732                let acts: Vec<SplitAct> = (0..b)
12733                    .map(|bi| split_act(&xs[bi * cols..(bi + 1) * cols]))
12734                    .collect();
12735                for r in 0..rows {
12736                    v.scales_into(r, gpr, &mut sc);
12737                    if !a8w8 {
12738                        for bi in 0..b {
12739                            let x = &xs[bi * cols..(bi + 1) * cols];
12740                            old[bi * rows + r] = q4tp_row_exact(v.nib, r, gpr, x, &sc);
12741                        }
12742                        continue;
12743                    }
12744                    let finish = |d: f32, act: &SplitAct| {
12745                        let mut acc = d * act.sx;
12746                        for &(j, xv) in &act.outliers {
12747                            let (w, s) = q4tp_outlier(v.nib, r, gpr, j, &sc);
12748                            acc += w * s * xv;
12749                        }
12750                        acc
12751                    };
12752                    let mut bi = 0usize;
12753                    while blocked && bi + 4 <= b {
12754                        let xs4 = [
12755                            acts[bi].xq.as_slice(),
12756                            acts[bi + 1].xq.as_slice(),
12757                            acts[bi + 2].xq.as_slice(),
12758                            acts[bi + 3].xq.as_slice(),
12759                        ];
12760                        let d = unsafe { dot_q4tp_row_1x4_sdot(v.nib, r, gpr, xs4, &sc) };
12761                        for k in 0..4 {
12762                            old[(bi + k) * rows + r] = finish(d[k], &acts[bi + k]);
12763                        }
12764                        bi += 4;
12765                    }
12766                    for (bi, act) in acts.iter().enumerate().skip(bi) {
12767                        let d = dot_q4tp_row_i8(v.nib, r, gpr, &act.xq, &sc);
12768                        old[bi * rows + r] = finish(d, act);
12769                    }
12770                }
12771                assert_eq!(
12772                    bits(&fast),
12773                    bits(&old),
12774                    "{rows}x{cols} b={b}: the fast path changed outside row_exact"
12775                );
12776                // And it is still the fast tile that runs: its lane-parallel
12777                // fma sum rounds differently from the matvec somewhere.
12778                if blocked {
12779                    assert_ne!(
12780                        bits(&fast),
12781                        bits(&matvecs),
12782                        "{rows}x{cols} b={b}: the tuned tile no longer runs outside row_exact"
12783                    );
12784                }
12785            }
12786        }
12787        Q4TP_ALT.store(0, Relaxed);
12788    }
12789
12790    #[test]
12791    fn row_exact_scopes_survive_overlap_nesting_and_unwind() {
12792        use std::sync::{Barrier, atomic::{AtomicUsize, Ordering}};
12793        // A private counter makes this restoration test independent of
12794        // numerical tests concurrently using the production counter.
12795        let active = AtomicUsize::new(0);
12796        counted_row_exact_scope(&active, || {
12797            assert_eq!(active.load(Ordering::Acquire), 1);
12798            counted_row_exact_scope(&active, || {
12799                assert_eq!(active.load(Ordering::Acquire), 2);
12800            });
12801            assert_eq!(active.load(Ordering::Acquire), 1);
12802        });
12803        assert_eq!(active.load(Ordering::Acquire), 0);
12804
12805        let both_entered = Barrier::new(2);
12806        let release_last = Barrier::new(2);
12807        std::thread::scope(|s| {
12808            let first = s.spawn(|| counted_row_exact_scope(&active, || {
12809                both_entered.wait();
12810            }));
12811            let last = s.spawn(|| counted_row_exact_scope(&active, || {
12812                both_entered.wait();
12813                release_last.wait();
12814            }));
12815            first.join().unwrap();
12816            let after_first = active.load(Ordering::Acquire);
12817            release_last.wait();
12818            last.join().unwrap();
12819            assert_eq!(after_first, 1, "second request must remain exact");
12820        });
12821        assert_eq!(active.load(Ordering::Acquire), 0);
12822        let panic = std::panic::catch_unwind(|| {
12823            counted_row_exact_scope(&active, || panic!("scope unwind"));
12824        });
12825        assert!(panic.is_err());
12826        assert_eq!(active.load(Ordering::Acquire), 0);
12827    }
12828
12829    #[test]
12830    #[cfg(target_arch = "x86_64")]
12831    fn q4tp_float_avx2_is_bitwise_scalar() {
12832        if !avx2_enabled() {
12833            return;
12834        }
12835        for cols in [32, 64, 96, 2048, 4096] {
12836            let rows = 9;
12837            let bytes = synth_q4tp(rows, cols);
12838            let v = Q4tpView::new(&bytes, rows, cols);
12839            let gpr = cols / GROUP_SIZE;
12840            let mut sc = vec![0.0; gpr];
12841            for seed in 1..=5 {
12842                let xs: Vec<f32> = (0..cols)
12843                    .map(|i| (((i * 104729 + seed * 8191) % 100003) as f32 - 50001.0) / 7919.0)
12844                    .collect();
12845                for r in 0..rows {
12846                    v.scales_into(r, gpr, &mut sc);
12847                    let scalar = q4tp_row_float_scalar(v.nib, r, gpr, &xs, &sc);
12848                    let vector = unsafe { q4tp_row_float_avx2(v.nib, r, gpr, &xs, &sc) };
12849                    assert_eq!(
12850                        scalar.to_bits(),
12851                        vector.to_bits(),
12852                        "cols={cols} row={r} seed={seed}"
12853                    );
12854                }
12855            }
12856        }
12857    }
12858
12859    #[test]
12860    fn multi_token_moe_rows_float_equal_single_token_decode() {
12861        float_activations_scope(multi_token_moe_rows_equal_single_token_decode);
12862    }
12863
12864    #[test]
12865    fn full_gpu_q8_scope_is_nested_and_thread_local() {
12866        assert!(!FULL_GPU_Q8.get());
12867        let before = gpu_split_frac();
12868        {
12869            let _guard = enter_full_gpu_q8_scope();
12870            assert_eq!(gpu_split_frac(), 1.0);
12871            {
12872                let _nested = enter_full_gpu_q8_scope();
12873            }
12874            assert_eq!(gpu_split_frac(), 1.0);
12875            std::thread::spawn(|| assert!(!FULL_GPU_Q8.get())).join().unwrap();
12876        }
12877        assert!(!FULL_GPU_Q8.get());
12878        assert_eq!(gpu_split_frac(), before);
12879    }
12880
12881    #[test]
12882    fn float_activation_scope_is_nested_thread_local_and_unwind_safe() {
12883        assert!(!FLOAT_ACTIVATIONS.get());
12884        let before = a8w8_enabled();
12885        float_activations_scope(|| {
12886            assert!(!a8w8_enabled());
12887            float_activations_scope(|| assert!(!a8w8_enabled()));
12888            assert!(FLOAT_ACTIVATIONS.get());
12889            std::thread::spawn(|| assert!(!FLOAT_ACTIVATIONS.get()))
12890                .join()
12891                .unwrap();
12892        });
12893        assert!(!FLOAT_ACTIVATIONS.get());
12894        assert_eq!(a8w8_enabled(), before);
12895        let _ = std::panic::catch_unwind(|| float_activations_scope(|| panic!("test unwind")));
12896        assert!(!FLOAT_ACTIVATIONS.get());
12897    }
12898
12899    /// Batched q4/vbit matmat must equal per-position matvec calls
12900    /// exactly (the fallback it replaced) — same kernels, same order.
12901    #[test]
12902    fn batched_matmat_equals_per_position_matvec() {
12903        let (rows, cols, b) = (8, 64, 5);
12904        // q4 blob.
12905        let groups = rows * cols / GROUP_SIZE;
12906        let mut q4 = Vec::new();
12907        for i in 0..groups * 16 {
12908            q4.push((((i * 7 + 3) % 256) & 0xFF) as u8);
12909        }
12910        for g in 0..groups {
12911            q4.extend_from_slice(
12912                &cortiq_core::quant::f32_to_f16(0.01 + 0.003 * g as f32).to_le_bytes(),
12913            );
12914        }
12915        // vbit blob (mixed widths incl. 8).
12916        let ng = cols / GROUP_SIZE;
12917        let bits: Vec<u8> = vec![3, 4, 5, 6, 8, 4, 5, 3];
12918        let mut vb = bits.clone();
12919        for g in 0..rows * ng {
12920            vb.extend_from_slice(
12921                &cortiq_core::quant::f32_to_f16(0.02 + 0.001 * g as f32).to_le_bytes(),
12922            );
12923        }
12924        for r in 0..rows {
12925            let bw = bits[r] as usize;
12926            let (mut acc, mut nb) = (0u64, 0usize);
12927            let mut rowbytes = Vec::new();
12928            for i in 0..cols {
12929                let v = ((i * 7 + r * 13) % (1 << bw)) as u64;
12930                acc = (acc << bw) | v;
12931                nb += bw;
12932                while nb >= 8 {
12933                    nb -= 8;
12934                    rowbytes.push(((acc >> nb) & 0xFF) as u8);
12935                }
12936            }
12937            if nb > 0 {
12938                rowbytes.push(((acc << (8 - nb)) & 0xFF) as u8);
12939            }
12940            vb.extend_from_slice(&rowbytes);
12941        }
12942        let offsets = vbit_row_offsets(&vb, rows, cols);
12943
12944        let xs: Vec<f32> = (0..b * cols).map(|i| (i as f32 * 0.13).sin()).collect();
12945
12946        // q4: batch vs singles.
12947        let mut got = vec![0f32; b * rows];
12948        q4matmat(&q4, &xs, b, rows, cols, &mut got, None);
12949        for bi in 0..b {
12950            let mut expect = vec![0f32; rows];
12951            q4matvec(
12952                &q4,
12953                &xs[bi * cols..(bi + 1) * cols],
12954                rows,
12955                cols,
12956                &mut expect,
12957                None,
12958            );
12959            assert_eq!(
12960                &got[bi * rows..(bi + 1) * rows],
12961                &expect[..],
12962                "q4 batch pos {bi}"
12963            );
12964        }
12965
12966        // vbit: batch vs singles.
12967        let mut got = vec![0f32; b * rows];
12968        vbitmatmat(&vb, &offsets, &xs, b, rows, cols, &mut got, None);
12969        for bi in 0..b {
12970            let mut expect = vec![0f32; rows];
12971            vbitmatvec(
12972                &vb,
12973                &offsets,
12974                &xs[bi * cols..(bi + 1) * cols],
12975                rows,
12976                cols,
12977                &mut expect,
12978                None,
12979            );
12980            assert_eq!(
12981                &got[bi * rows..(bi + 1) * rows],
12982                &expect[..],
12983                "vbit batch pos {bi}"
12984            );
12985        }
12986    }
12987
12988    /// q4_tiled kernels must produce BIT-identical outputs to the q4
12989    /// split kernels on the same values (same ints, same order — only
12990    /// the byte placement differs).
12991    #[test]
12992    fn q4_tiled_matches_q4_block_bitexact() {
12993        let (rows, cols, b) = (8usize, 128usize, 3usize);
12994        let groups = rows * cols / GROUP_SIZE;
12995        let mut split = Vec::with_capacity(groups * 18);
12996        for i in 0..groups * 16 {
12997            split.push((((i * 7 + 3) % 256) & 0xFF) as u8);
12998        }
12999        for g in 0..groups {
13000            split.extend_from_slice(
13001                &cortiq_core::quant::f32_to_f16(0.01 + 0.003 * g as f32).to_le_bytes(),
13002            );
13003        }
13004        // Re-tile: [scale][nibbles] per group.
13005        let (packed, scales) = split.split_at(groups * 16);
13006        let mut tiled = Vec::with_capacity(groups * Q4_TILE);
13007        for g in 0..groups {
13008            tiled.extend_from_slice(&scales[g * 2..g * 2 + 2]);
13009            tiled.extend_from_slice(&packed[g * 16..(g + 1) * 16]);
13010        }
13011
13012        let mut x1: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.17).sin()).collect();
13013        x1[9] = 250.0; // exercise the outlier path
13014        let x2: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.23).cos()).collect();
13015
13016        let (mut a, mut t) = (vec![0f32; rows], vec![0f32; rows]);
13017        q4matvec(&split, &x1, rows, cols, &mut a, None);
13018        q4t_matvec(&tiled, &x1, rows, cols, &mut t, None);
13019        assert_eq!(a, t, "q4t matvec must match q4 bit-for-bit");
13020
13021        let (mut a1, mut a2) = (vec![0f32; rows], vec![0f32; rows]);
13022        let (mut t1, mut t2) = (vec![0f32; rows], vec![0f32; rows]);
13023        q4matvec2(&split, &x1, &x2, rows, cols, &mut a1, &mut a2, None);
13024        q4t_matvec2(&tiled, &x1, &x2, rows, cols, &mut t1, &mut t2, None);
13025        assert_eq!(a1, t1);
13026        assert_eq!(a2, t2);
13027
13028        let xs: Vec<f32> = (0..b * cols).map(|i| (i as f32 * 0.13).sin()).collect();
13029        let (mut am, mut tm) = (vec![0f32; b * rows], vec![0f32; b * rows]);
13030        q4matmat(&split, &xs, b, rows, cols, &mut am, None);
13031        q4t_matmat(&tiled, &xs, b, rows, cols, &mut tm, None);
13032        assert_eq!(am, tm, "q4t matmat must match q4 bit-for-bit");
13033    }
13034
13035    /// q4 SDOT outlier correction: a single huge activation channel
13036    /// (>8·rms → outlier, zeroed in xq) must still contribute its EXACT
13037    /// term. On-grid bulk (±1/0 → xq dequantizes exactly) isolates the
13038    /// correction from A8W8 noise. cols must exceed 64: at n=64 the
13039    /// 8·rms threshold equals sqrt(v²+rest) ≥ v, so a single outlier
13040    /// can never qualify (8² = n).
13041    #[test]
13042    fn q4matvec_sdot_outlier_exact() {
13043        let (rows, cols) = (4, 128);
13044        let groups = rows * cols / GROUP_SIZE;
13045        let mut bytes = Vec::with_capacity(groups * 16 + groups * 2);
13046        for i in 0..groups * 16 {
13047            bytes.push(((i * 11 + 5) % 256) as u8);
13048        }
13049        for g in 0..groups {
13050            let s = 0.02 + 0.002 * g as f32;
13051            bytes.extend_from_slice(&cortiq_core::quant::f32_to_f16(s).to_le_bytes());
13052        }
13053        let mut x: Vec<f32> = (0..cols)
13054            .map(|i| match i % 3 {
13055                0 => 1.0,
13056                1 => -1.0,
13057                _ => 0.0,
13058            })
13059            .collect();
13060        x[17] = 300.0; // ≫ 8·rms → outlier channel
13061
13062        let mut reference = vec![0.0f32; rows * cols];
13063        cortiq_core::quant::dequant_q4_block(&bytes, &mut reference);
13064        let mut expect = vec![0.0f32; rows];
13065        for r in 0..rows {
13066            expect[r] = reference[r * cols..(r + 1) * cols]
13067                .iter()
13068                .zip(&x)
13069                .map(|(w, xv)| w * xv)
13070                .sum();
13071        }
13072        let mut got = vec![0.0f32; rows];
13073        q4matvec(&bytes, &x, rows, cols, &mut got, None);
13074        let scale = expect.iter().fold(0f32, |m, v| m.max(v.abs())).max(1.0);
13075        for r in 0..rows {
13076            assert!(
13077                (got[r] - expect[r]).abs() < 2e-3 * scale,
13078                "row {r}: {} vs {} (outlier term must be exact)",
13079                got[r],
13080                expect[r]
13081            );
13082        }
13083    }
13084
13085    /// The fused q1t matvec must equal the reference (dequant_q1t → dot),
13086    /// including the ternary zero level and the binary-searched outlier
13087    /// overlay. Guards the mmap kernel that makes a 12B q1t runnable.
13088    #[test]
13089    fn q1t_matvec_matches_reference() {
13090        use cortiq_core::quant::{dequant_q1t, f32_to_f16};
13091        let (rows, cols) = (3usize, 64usize); // gpr = 2
13092        let gpr = cols / GROUP_SIZE;
13093        let scales = [0.5f32, 0.3, 0.7, 0.2, 0.6, 0.15];
13094        // Overlay (must be sorted by flat index): a few spikes across rows.
13095        let outliers: [(u32, f32); 3] = [(5, 9.0), (70, -4.5), (150, 3.25)];
13096        let is_out = |flat: usize| outliers.iter().any(|&(i, _)| i as usize == flat);
13097        let mut bytes = Vec::new();
13098        for r in 0..rows {
13099            for g in 0..gpr {
13100                bytes.extend_from_slice(&f32_to_f16(scales[r * gpr + g]).to_le_bytes());
13101                let mut c = [0u8; 7];
13102                for k in 0..GROUP_SIZE {
13103                    // Encoder invariant: code 0 at outlier positions.
13104                    let code = if is_out(r * cols + g * GROUP_SIZE + k) {
13105                        0
13106                    } else {
13107                        ((k + r * 3 + g) % 3) as u8 // 0,1,2
13108                    };
13109                    cortiq_core::quant::q1t_pack(&mut c, k, code);
13110                }
13111                bytes.extend_from_slice(&c);
13112            }
13113        }
13114        // Per-row overlay: [u32 row_ptr[rows+1]] then [(u16 col, f16 val)] by
13115        // row (outliers are sorted by flat index → already grouped by row).
13116        let mut row_ptr = vec![0u32; rows + 1];
13117        for &(idx, _) in &outliers {
13118            row_ptr[idx as usize / cols + 1] += 1;
13119        }
13120        for r in 0..rows {
13121            row_ptr[r + 1] += row_ptr[r];
13122        }
13123        for &p in &row_ptr {
13124            bytes.extend_from_slice(&p.to_le_bytes());
13125        }
13126        for &(idx, v) in &outliers {
13127            bytes.extend_from_slice(&((idx as usize % cols) as u16).to_le_bytes());
13128            bytes.extend_from_slice(&f32_to_f16(v).to_le_bytes());
13129        }
13130
13131        let mut refw = vec![0f32; rows * cols];
13132        dequant_q1t(&bytes, rows, cols, &mut refw);
13133        // On-grid activations (±1, amax 1) so the int8 SDOT path reconstructs
13134        // x exactly and matches the f32 reference (same trick as the q1 test).
13135        let x: Vec<f32> = (0..cols)
13136            .map(|j| if j % 3 == 0 { 1.0 } else { -1.0 })
13137            .collect();
13138        let mut expect = vec![0f32; rows];
13139        for r in 0..rows {
13140            let mut a = 0.0f32;
13141            for j in 0..cols {
13142                a += refw[r * cols + j] * x[j];
13143            }
13144            expect[r] = a;
13145        }
13146        let tol = |e: f32| 1e-3 * e.abs().max(1e-3);
13147        let mut got = vec![0f32; rows];
13148        q1t_matvec(&bytes, &x, rows, cols, &mut got, None);
13149        for r in 0..rows {
13150            assert!(
13151                (got[r] - expect[r]).abs() < tol(expect[r]),
13152                "row {r}: {} vs {}",
13153                got[r],
13154                expect[r]
13155            );
13156        }
13157        // matmat (b=2, f32 decode path) must agree too.
13158        let x2: Vec<f32> = x.iter().chain(x.iter()).copied().collect();
13159        let mut gm = vec![0f32; 2 * rows];
13160        q1t_matmat(&bytes, &x2, 2, rows, cols, &mut gm, None);
13161        for r in 0..rows {
13162            assert!((gm[r] - expect[r]).abs() < tol(expect[r]));
13163            assert!((gm[rows + r] - expect[r]).abs() < tol(expect[r]));
13164        }
13165        // Fused pair (q1t_matvec2) must equal two single matvecs
13166        // bit-for-bit: same unpack, same group order, same f32
13167        // accumulation per stream. Distinct x2 exercises both lanes.
13168        let xb: Vec<f32> = (0..cols)
13169            .map(|j| if j % 5 == 0 { -1.0 } else { 1.0 })
13170            .collect();
13171        let (mut s1, mut s2) = (vec![0f32; rows], vec![0f32; rows]);
13172        q1t_matvec(&bytes, &x, rows, cols, &mut s1, None);
13173        q1t_matvec(&bytes, &xb, rows, cols, &mut s2, None);
13174        let (mut p1, mut p2) = (vec![0f32; rows], vec![0f32; rows]);
13175        q1t_matvec2(&bytes, &x, &xb, rows, cols, &mut p1, &mut p2, None);
13176        assert_eq!(p1, s1, "q1t pair lane 1 ≠ single matvec");
13177        assert_eq!(p2, s2, "q1t pair lane 2 ≠ single matvec");
13178    }
13179
13180    /// Pair == 2×matvec with an ODD group count (the kernel's tail
13181    /// group) and no overlay section.
13182    #[test]
13183    fn q1t_matvec2_odd_gpr_matches_singles() {
13184        use cortiq_core::quant::{Q1T_TILE, f32_to_f16, q1t_pack};
13185        let (rows, cols) = (5usize, 96usize); // gpr = 3 → paired + tail
13186        let gpr = cols / GROUP_SIZE;
13187        let mut bytes = Vec::with_capacity(rows * gpr * Q1T_TILE);
13188        for r in 0..rows {
13189            for g in 0..gpr {
13190                bytes.extend_from_slice(&f32_to_f16(0.1 + 0.05 * (r + g) as f32).to_le_bytes());
13191                let mut c = [0u8; 7];
13192                for k in 0..GROUP_SIZE {
13193                    q1t_pack(&mut c, k, ((k * 7 + r * 5 + g * 3) % 3) as u8);
13194                }
13195                bytes.extend_from_slice(&c);
13196            }
13197        }
13198        let x1: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.31).sin()).collect();
13199        let x2: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.17).cos()).collect();
13200        let (mut s1, mut s2) = (vec![0f32; rows], vec![0f32; rows]);
13201        q1t_matvec(&bytes, &x1, rows, cols, &mut s1, None);
13202        q1t_matvec(&bytes, &x2, rows, cols, &mut s2, None);
13203        let (mut p1, mut p2) = (vec![0f32; rows], vec![0f32; rows]);
13204        q1t_matvec2(&bytes, &x1, &x2, rows, cols, &mut p1, &mut p2, None);
13205        assert_eq!(p1, s1, "odd-gpr pair lane 1 ≠ single");
13206        assert_eq!(p2, s2, "odd-gpr pair lane 2 ≠ single");
13207    }
13208
13209    // Speed A/B: fused pair (one unpack, two streams) vs two single
13210    // matvecs. Single-threaded, FFN-sized, min-of paired in-process.
13211    //   cargo test -p cortiq-engine --release q1t_matvec2_speed -- --ignored --nocapture
13212    #[test]
13213    #[ignore]
13214    fn q1t_matvec2_speed() {
13215        use cortiq_core::quant::{Q1T_TILE, f32_to_f16, q1t_pack};
13216        use std::time::Instant;
13217        let (rows, cols) = (8192usize, 4096usize);
13218        let gpr = cols / GROUP_SIZE;
13219        let mut bytes = Vec::with_capacity(rows * gpr * Q1T_TILE);
13220        for r in 0..rows {
13221            for g in 0..gpr {
13222                let s = 0.1 + ((r + g) % 7) as f32 * 0.01;
13223                bytes.extend_from_slice(&f32_to_f16(s).to_le_bytes());
13224                let mut c = [0u8; 7];
13225                for k in 0..GROUP_SIZE {
13226                    q1t_pack(&mut c, k, ((k * 7 + r + g) % 3) as u8);
13227                }
13228                bytes.extend_from_slice(&c);
13229            }
13230        }
13231        let x1: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.31).sin()).collect();
13232        let x2: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.17).cos()).collect();
13233        let (mut s1, mut s2) = (vec![0f32; rows], vec![0f32; rows]);
13234        let (mut p1, mut p2) = (vec![0f32; rows], vec![0f32; rows]);
13235        // Warm both paths once.
13236        q1t_matvec(&bytes, &x1, rows, cols, &mut s1, None);
13237        q1t_matvec2(&bytes, &x1, &x2, rows, cols, &mut p1, &mut p2, None);
13238        let (mut t_pair, mut t_two) = (f64::MAX, f64::MAX);
13239        for _ in 0..8 {
13240            let t0 = Instant::now();
13241            q1t_matvec2(&bytes, &x1, &x2, rows, cols, &mut p1, &mut p2, None);
13242            t_pair = t_pair.min(t0.elapsed().as_secs_f64() * 1000.0);
13243            let t1 = Instant::now();
13244            q1t_matvec(&bytes, &x1, rows, cols, &mut s1, None);
13245            q1t_matvec(&bytes, &x2, rows, cols, &mut s2, None);
13246            t_two = t_two.min(t1.elapsed().as_secs_f64() * 1000.0);
13247        }
13248        assert_eq!(p1, s1);
13249        assert_eq!(p2, s2);
13250        println!("q1t pair {rows}x{cols}: fused {t_pair:.2} ms | two singles {t_two:.2} ms");
13251    }
13252
13253    // Speed A/B: the base-3-division decode (what the packing commit left in
13254    // place) vs the fused sign-LUT matvec. Both single-threaded, same bytes.
13255    //   cargo test -p cortiq-engine q1t_matvec_speed -- --ignored --nocapture
13256    #[test]
13257    #[ignore]
13258    fn q1t_matvec_speed() {
13259        use cortiq_core::quant::{Q1T_TILE, f32_to_f16, q1t_code, q1t_pack};
13260        use std::time::Instant;
13261        let (rows, cols) = (8192usize, 4096usize); // FFN-sized
13262        let gpr = cols / GROUP_SIZE;
13263        let mut bytes = Vec::with_capacity(rows * gpr * Q1T_TILE + 16);
13264        for r in 0..rows {
13265            for g in 0..gpr {
13266                let s = 0.1 + ((r + g) % 7) as f32 * 0.01;
13267                bytes.extend_from_slice(&f32_to_f16(s).to_le_bytes());
13268                let mut c = [0u8; 7];
13269                for k in 0..GROUP_SIZE {
13270                    q1t_pack(&mut c, k, ((k * 7 + r + g) % 3) as u8);
13271                }
13272                bytes.extend_from_slice(&c);
13273            }
13274        }
13275        let (n, stride) = (rows * cols, 40usize); // ~2.5% outliers, per-row overlay
13276        let mut row_ptr = vec![0u32; rows + 1];
13277        let mut idx = 0usize;
13278        while idx < n {
13279            row_ptr[idx / cols + 1] += 1;
13280            idx += stride;
13281        }
13282        for r in 0..rows {
13283            row_ptr[r + 1] += row_ptr[r];
13284        }
13285        for &p in &row_ptr {
13286            bytes.extend_from_slice(&p.to_le_bytes());
13287        }
13288        let mut idx = 0usize;
13289        while idx < n {
13290            bytes.extend_from_slice(&((idx % cols) as u16).to_le_bytes());
13291            bytes.extend_from_slice(&f32_to_f16((idx % 13) as f32 * 0.1 - 0.6).to_le_bytes());
13292            idx += stride;
13293        }
13294        // On-grid ±1 so the fast path's int8 SDOT is exact vs the f32 "slow"
13295        // reference (the A/B is a timing check; values must still agree).
13296        let x: Vec<f32> = (0..cols)
13297            .map(|j| if j % 3 == 0 { 1.0 } else { -1.0 })
13298            .collect();
13299        let (rp_off, ent_off, has_ov) = q1t_overlay(&bytes, rows * gpr * Q1T_TILE, rows);
13300
13301        // "before": base-3 division decode into a buffer, then dot.
13302        let slow = |out: &mut [f32]| {
13303            let mut buf = vec![0f32; cols];
13304            for r in 0..rows {
13305                for g in 0..gpr {
13306                    let off = (r * gpr + g) * Q1T_TILE;
13307                    let s = f16_to_f32(u16::from_le_bytes([bytes[off], bytes[off + 1]]));
13308                    let codes = &bytes[off + 2..off + Q1T_TILE];
13309                    for k in 0..GROUP_SIZE {
13310                        buf[g * GROUP_SIZE + k] = match q1t_code(codes, k) {
13311                            1 => s,
13312                            2 => -s,
13313                            _ => 0.0,
13314                        };
13315                    }
13316                }
13317                out[r] = q1t_row_outlier_correction(&bytes, r, rp_off, ent_off, has_ov, &x)
13318                    + (0..cols).map(|j| buf[j] * x[j]).sum::<f32>();
13319            }
13320        };
13321        let iters = 5;
13322        let mut a = vec![0f32; rows];
13323        slow(&mut a); // warm
13324        let t = Instant::now();
13325        for _ in 0..iters {
13326            slow(&mut a);
13327        }
13328        let slow_ms = t.elapsed().as_secs_f64() * 1e3 / iters as f64;
13329
13330        let mut b = vec![0f32; rows];
13331        q1t_matvec(&bytes, &x, rows, cols, &mut b, None); // warm
13332        let t = Instant::now();
13333        for _ in 0..iters {
13334            q1t_matvec(&bytes, &x, rows, cols, &mut b, None);
13335        }
13336        let fast_ms = t.elapsed().as_secs_f64() * 1e3 / iters as f64;
13337
13338        for r in 0..rows {
13339            assert!((a[r] - b[r]).abs() < 1e-2, "mismatch row {r}");
13340        }
13341        println!(
13342            "q1t matvec {rows}x{cols} (1 thread): div-decode {slow_ms:.2} ms  fused-LUT {fast_ms:.2} ms  => {:.2}x",
13343            slow_ms / fast_ms
13344        );
13345    }
13346}
13347
13348#[cfg(test)]
13349mod gemm_bench {
13350    /// `cargo test -p cortiq-engine --release q4tp_matmat_throughput -- --ignored --nocapture`
13351    /// Times the batched q4tp GEMM at the shapes the image DiT runs
13352    /// (b=296 tokens, 2304 -> 9216), on synthetic bytes: no model, no
13353    /// mmap, no thermal drift over minutes — a kernel change shows up
13354    /// here in seconds where a full render hides it in noise.
13355    ///
13356    /// On macOS add `CMF_ACCEL=0`: this shape is over the 500k-cell mark
13357    /// where the matmat hands off to Accelerate's dequant sgemm, and
13358    /// without the opt-out both rows below measure the AMX, not the
13359    /// kernel under test.
13360    #[test]
13361    #[ignore]
13362    fn q4tp_matmat_throughput() {
13363        let _alt = super::Q4TP_ALT_TEST_LOCK
13364            .lock()
13365            .unwrap_or_else(|e| e.into_inner());
13366        // 296 is a prompt-encode batch; the image DiT runs 2085 at
13367        // 512x512, where the activation panel stops fitting L2 and the
13368        // loop's shape starts to matter more than its instructions.
13369        let b: usize = std::env::var("CMF_BENCH_B")
13370            .ok()
13371            .and_then(|v| v.parse().ok())
13372            .unwrap_or(296);
13373        let (rows, cols) = (9216usize, 2304usize);
13374        let (_, _, _) = (rows, cols, b);
13375        let total =
13376            cortiq_core::quant::expected_nbytes(cortiq_core::TensorDtype::Q4TiledP, &[rows, cols])
13377                .unwrap();
13378        // Random nibbles are fine, but the row params are f16 (lo, step)
13379        // of a geometric ladder: garbage there gives exp2 of a huge
13380        // exponent, the scales come back inf, and the whole bench times
13381        // NaN arithmetic instead of the kernel.
13382        let (params_off, codes_off, _) = cortiq_core::quant::q4tp_sections(rows, cols);
13383        let mut bytes: Vec<u8> = (0..total).map(|i| (i * 37 % 251) as u8).collect();
13384        let lo = cortiq_core::quant::f32_to_f16(-4.0);
13385        let step = cortiq_core::quant::f32_to_f16(0.1);
13386        for r in 0..rows {
13387            let o = params_off + r * 4;
13388            bytes[o..o + 2].copy_from_slice(&lo.to_le_bytes());
13389            bytes[o + 2..o + 4].copy_from_slice(&step.to_le_bytes());
13390        }
13391        let _ = codes_off;
13392        let xs: Vec<f32> = (0..b * cols)
13393            .map(|i| ((i % 97) as f32 - 48.0) / 48.0)
13394            .collect();
13395        let mut out = vec![0f32; b * rows];
13396        let pool = crate::pool::Pool::from_env();
13397        // A shared 48-core stand drifts ±25% run to run, which is wider
13398        // than any kernel change worth making. So: alternate the two
13399        // kernels inside one process and keep the BEST time for
13400        // each. Interleaving makes both see the same interference, and a
13401        // minimum is the one statistic another tenant cannot inflate.
13402        super::q4tp_matmat(&bytes, &xs, b, rows, cols, &mut out, pool.as_deref());
13403        let reps: usize = std::env::var("CMF_BENCH_REPS")
13404            .ok()
13405            .and_then(|v| v.parse().ok())
13406            .unwrap_or(10);
13407        let mut best = [f64::MAX; 2];
13408        let mut sums = [0f32; 2];
13409        for _ in 0..reps {
13410            for (k, w) in [(0usize, 1u8), (1usize, 2u8)] {
13411                super::Q4TP_ALT.store(w, std::sync::atomic::Ordering::Relaxed);
13412                let t = std::time::Instant::now();
13413                super::q4tp_matmat(&bytes, &xs, b, rows, cols, &mut out, pool.as_deref());
13414                best[k] = best[k].min(t.elapsed().as_secs_f64());
13415                sums[k] = out.iter().take(64).sum::<f32>();
13416            }
13417        }
13418        let flops = 2.0 * b as f64 * rows as f64 * cols as f64;
13419        for (k, name) in ["previous", "tuned   "].iter().enumerate() {
13420            println!(
13421                "q4tp matmat {rows}x{cols} b={b} {name}: {:.1} ms  {:.1} GFLOP/s  (checksum {:.3})",
13422                best[k] * 1e3,
13423                flops / best[k] / 1e9,
13424                sums[k]
13425            );
13426        }
13427        assert!(
13428            (sums[0] - sums[1]).abs() < 1e-2,
13429            "the tuned kernel changed the result: {} vs {}",
13430            sums[0],
13431            sums[1]
13432        );
13433    }
13434
13435    /// The blocked kernel must agree with the per-column path exactly —
13436    /// same weights, same activation split, only a different instruction
13437    /// mix. Shapes are chosen to hit the awkward cases: a column count
13438    /// that leaves an odd group (the 512-bit kernel does two at a time),
13439    /// and a batch that does not divide by four.
13440    #[test]
13441    fn q4tp_matmat_blocked_matches_scalar() {
13442        use std::sync::atomic::Ordering::Relaxed;
13443        let _alt = super::Q4TP_ALT_TEST_LOCK
13444            .lock()
13445            .unwrap_or_else(|e| e.into_inner());
13446        // The last shape carries the image DiT's column count — 2304, so
13447        // 72 groups of accumulation, which is where a reordered sum can
13448        // actually drift — and runs through the thread pool, since the
13449        // blocked path splits rows across workers. Its row count stays
13450        // under 500k cells on purpose: above that, macOS diverts the whole
13451        // matmat to the Accelerate/AMX dequant sgemm and neither kernel
13452        // here would run.
13453        for &(rows, cols, b) in &[
13454            (64usize, 128usize, 7usize),
13455            (33, 96, 4),
13456            (16, 256, 9),
13457            (192, 2304, 37),
13458        ] {
13459            let total = cortiq_core::quant::expected_nbytes(
13460                cortiq_core::TensorDtype::Q4TiledP,
13461                &[rows, cols],
13462            )
13463            .unwrap();
13464            let (params_off, _, _) = cortiq_core::quant::q4tp_sections(rows, cols);
13465            let mut bytes: Vec<u8> = (0..total).map(|i| (i * 61 % 251) as u8).collect();
13466            let lo = cortiq_core::quant::f32_to_f16(-4.0);
13467            let step = cortiq_core::quant::f32_to_f16(0.1);
13468            for r in 0..rows {
13469                let o = params_off + r * 4;
13470                bytes[o..o + 2].copy_from_slice(&lo.to_le_bytes());
13471                bytes[o + 2..o + 4].copy_from_slice(&step.to_le_bytes());
13472            }
13473            let xs: Vec<f32> = (0..b * cols)
13474                .map(|i| ((i % 89) as f32 - 44.0) / 44.0)
13475                .collect();
13476            let mut got = vec![0f32; b * rows];
13477            let mut want = vec![0f32; b * rows];
13478            let gpr = cols / 32;
13479            let view = super::Q4tpView::new(&bytes, rows, cols);
13480            let pool = crate::pool::Pool::from_env();
13481            super::Q4TP_ALT.store(2, Relaxed);
13482            super::q4tp_matmat(&bytes, &xs, b, rows, cols, &mut got, pool.as_deref());
13483            super::Q4TP_ALT.store(1, Relaxed);
13484            super::q4tp_matmat(&bytes, &xs, b, rows, cols, &mut want, pool.as_deref());
13485            super::Q4TP_ALT.store(0, Relaxed);
13486            // Measured against the output's scale, not cell by cell: a
13487            // dot product of 2304 terms lands near zero wherever the row
13488            // and the activation nearly cancel, and there a per-cell
13489            // ratio reports 1e-3 for an absolute error of 5e-6 — f32's
13490            // own rounding, reordered. What must stay small is the error
13491            // relative to what the layer actually outputs.
13492            let scale = want.iter().fold(0f32, |m, v| m.max(v.abs())).max(1e-6);
13493            let (mut worst, mut at) = (0f32, 0usize);
13494            for (i, (g, w)) in got.iter().zip(&want).enumerate() {
13495                if (g - w).abs() > worst {
13496                    worst = (g - w).abs();
13497                    at = i;
13498                }
13499            }
13500            assert!(
13501                worst <= 1e-4 * scale,
13502                "{rows}x{cols} b={b}: blocked and scalar disagree by {worst:.3e} \
13503                 (scale {scale:.3e}) at cell {at}: {} vs {}",
13504                got[at],
13505                want[at]
13506            );
13507
13508            // "Same speed, no quality loss" is a claim about which answer
13509            // is RIGHT, not about which two agree. Both paths sum the same
13510            // 2304 products in different orders, so f64 decides: the
13511            // blocked kernel keeps sixteen partial sums and folds them at
13512            // the end, which is a shallower addition tree than the
13513            // per-column path's running scalar, and it must not be worse.
13514            let (mut e_blocked, mut e_scalar) = (0f64, 0f64);
13515            for bi in 0..b {
13516                let act = super::split_act(&xs[bi * cols..(bi + 1) * cols]);
13517                for r in 0..rows {
13518                    let mut sc = vec![0f32; gpr];
13519                    view.scales_into(r, gpr, &mut sc);
13520                    let mut exact = 0f64;
13521                    for j in 0..cols {
13522                        let (w, sq) = super::q4tp_outlier(view.nib, r, gpr, j, &sc);
13523                        exact += w as f64 * sq as f64 * act.xq[j] as f64;
13524                    }
13525                    exact *= act.sx as f64;
13526                    for &(j, xv) in &act.outliers {
13527                        let (w, sq) = super::q4tp_outlier(view.nib, r, gpr, j, &sc);
13528                        exact += w as f64 * sq as f64 * xv as f64;
13529                    }
13530                    let i = bi * rows + r;
13531                    e_blocked = e_blocked.max((got[i] as f64 - exact).abs());
13532                    e_scalar = e_scalar.max((want[i] as f64 - exact).abs());
13533                }
13534            }
13535            println!(
13536                "{rows}x{cols} b={b}: worst error vs f64 — blocked {e_blocked:.3e}, \
13537                 per-column {e_scalar:.3e}"
13538            );
13539            // An absolute bar, not a race between the two: at these
13540            // magnitudes both sit in f32's last bits, and on a small shape
13541            // whichever one happens to round the unluckiest cell "wins" by
13542            // a factor the next seed reverses.
13543            assert!(
13544                e_blocked <= 1e-5 * scale as f64 && e_scalar <= 1e-5 * scale as f64,
13545                "{rows}x{cols} b={b}: error against f64 too large — blocked \
13546                 {e_blocked:.3e}, per-column {e_scalar:.3e}, scale {scale:.3e}"
13547            );
13548        }
13549    }
13550
13551    /// What the row-exact contract costs a speculative-verify panel on the
13552    /// host: the same batch through the fast arms, through the row-exact
13553    /// arms, and as one matvec per token (the other way to be exact).
13554    /// Arms alternate inside one process and keep their best time.
13555    /// `CMF_GPU=0 cargo test -p cortiq-engine --release q4tp_matmat_row_exact_cost -- --ignored --nocapture`
13556    #[test]
13557    #[ignore]
13558    fn q4tp_matmat_row_exact_cost() {
13559        let _alt = super::Q4TP_ALT_TEST_LOCK
13560            .lock()
13561            .unwrap_or_else(|e| e.into_inner());
13562        let pool = crate::pool::Pool::from_env();
13563        let reps: usize = std::env::var("CMF_BENCH_REPS")
13564            .ok()
13565            .and_then(|v| v.parse().ok())
13566            .unwrap_or(30);
13567        for &(rows, cols) in &[(2048usize, 4096usize), (4096, 2048), (4096, 4096)] {
13568            let total = cortiq_core::quant::expected_nbytes(
13569                cortiq_core::TensorDtype::Q4TiledP,
13570                &[rows, cols],
13571            )
13572            .unwrap();
13573            let (params_off, _, _) = cortiq_core::quant::q4tp_sections(rows, cols);
13574            let mut bytes: Vec<u8> = (0..total).map(|i| (i * 37 % 251) as u8).collect();
13575            let lo = cortiq_core::quant::f32_to_f16(-4.0);
13576            let step = cortiq_core::quant::f32_to_f16(0.1);
13577            for r in 0..rows {
13578                let o = params_off + r * 4;
13579                bytes[o..o + 2].copy_from_slice(&lo.to_le_bytes());
13580                bytes[o + 2..o + 4].copy_from_slice(&step.to_le_bytes());
13581            }
13582            for &b in &[2usize, 4, 5, 8] {
13583                let xs: Vec<f32> = (0..b * cols)
13584                    .map(|i| ((i % 97) as f32 - 48.0) / 48.0)
13585                    .collect();
13586                let mut out = vec![0f32; b * rows];
13587                let mut best = [f64::MAX; 3];
13588                for _ in 0..reps {
13589                    for (k, best_k) in best.iter_mut().enumerate() {
13590                        let t = std::time::Instant::now();
13591                        match k {
13592                            0 | 1 => super::q4tp_matmat_with(
13593                                &bytes,
13594                                &xs,
13595                                b,
13596                                rows,
13597                                cols,
13598                                &mut out,
13599                                pool.as_deref(),
13600                                k == 1,
13601                            ),
13602                            _ => {
13603                                for (bi, o) in out.chunks_mut(rows).enumerate() {
13604                                    super::q4tp_matvec(
13605                                        &bytes,
13606                                        &xs[bi * cols..(bi + 1) * cols],
13607                                        rows,
13608                                        cols,
13609                                        o,
13610                                        pool.as_deref(),
13611                                    );
13612                                }
13613                            }
13614                        }
13615                        *best_k = best_k.min(t.elapsed().as_secs_f64());
13616                    }
13617                }
13618                println!(
13619                    "q4tp {rows}x{cols} b={b}: fast {:.3} ms, row-exact {:.3} ms, \
13620                     {b} matvecs {:.3} ms",
13621                    best[0] * 1e3,
13622                    best[1] * 1e3,
13623                    best[2] * 1e3
13624                );
13625            }
13626        }
13627    }
13628
13629    /// The q4t twin of the throughput bench, same shape and rules, so the
13630    /// two quantisations' batch kernels can be read against each other.
13631    /// `cargo test -p cortiq-engine --release q4t_matmat_throughput -- --ignored --nocapture`
13632    #[test]
13633    #[ignore]
13634    fn q4t_matmat_throughput() {
13635        let (rows, cols, b) = (9216usize, 2304usize, 296usize);
13636        let total =
13637            cortiq_core::quant::expected_nbytes(cortiq_core::TensorDtype::Q4Tiled, &[rows, cols])
13638                .unwrap();
13639        // q4t carries a per-group f16 scale in the tile's first two bytes;
13640        // random bytes there decode to inf and the bench would time NaNs.
13641        let mut bytes: Vec<u8> = (0..total).map(|i| (i * 37 % 251) as u8).collect();
13642        let sc = cortiq_core::quant::f32_to_f16(0.02);
13643        for t in bytes.chunks_mut(super::Q4_TILE) {
13644            t[..2].copy_from_slice(&sc.to_le_bytes());
13645        }
13646        let xs: Vec<f32> = (0..b * cols)
13647            .map(|i| ((i % 97) as f32 - 48.0) / 48.0)
13648            .collect();
13649        let mut out = vec![0f32; b * rows];
13650        let pool = crate::pool::Pool::from_env();
13651        super::q4t_matmat(&bytes, &xs, b, rows, cols, &mut out, pool.as_deref());
13652        let reps: usize = std::env::var("CMF_BENCH_REPS")
13653            .ok()
13654            .and_then(|v| v.parse().ok())
13655            .unwrap_or(10);
13656        let mut best = f64::MAX;
13657        for _ in 0..reps {
13658            let t = std::time::Instant::now();
13659            super::q4t_matmat(&bytes, &xs, b, rows, cols, &mut out, pool.as_deref());
13660            best = best.min(t.elapsed().as_secs_f64());
13661        }
13662        let flops = 2.0 * b as f64 * rows as f64 * cols as f64;
13663        println!(
13664            "q4t matmat {rows}x{cols} b={b}: {:.1} ms  {:.1} GFLOP/s  (checksum {:.3})",
13665            best * 1e3,
13666            flops / best / 1e9,
13667            out.iter().take(64).sum::<f32>()
13668        );
13669    }
13670}