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::sync::Arc;
21
22pub enum QTensor {
23    F32 {
24        data: Vec<f32>,
25        rows: usize,
26        cols: usize,
27    },
28    Mapped {
29        model: Arc<CmfModel>,
30        /// Index into the model's tensor directory.
31        idx: usize,
32        dtype: TensorDtype,
33        rows: usize,
34        cols: usize,
35        /// Per-row scales, dequantized to f32 up front (tiny).
36        row_scale: Vec<f32>,
37        /// q8_2f column field (θ), dequantized up front; empty for q8_row.
38        col_field: Vec<f32>,
39        /// Vbit only: byte offset of each row's packed data within the
40        /// tensor blob (`[rows + 1]`, computed once at load — the per-
41        /// matvec prefix scan over row bit-widths was O(rows) each call).
42        vbit_offsets: Vec<usize>,
43        /// q8-family decode repack (load-time, optional): rows in groups
44        /// of 4, interleaved in 16-byte units — one 64-byte line per
45        /// iteration feeds all 4 sdot lanes, ONE sequential weight
46        /// stream per worker instead of four (this is where llama.cpp's
47        /// repacked Q8 kernels get their bandwidth). Empty = off
48        /// (CMF_REPACK=0, non-SDOT arch, or an ineligible shape). Trades
49        /// an anonymous copy of the quants for mmap pages that go cold.
50        repack: Vec<u8>,
51    },
52}
53
54/// Load-time q8 repack gate (see `Mapped::repack`). OPT-IN
55/// (`CMF_REPACK=1`): the single-stream hypothesis LOST on Apple Silicon
56/// (M4, interleaved A/B: decode 101 vs 94 tok/s — four adjacent row
57/// streams per worker feed the prefetcher MORE memory-level parallelism
58/// than one); kept as an experiment flag for x86, where the tradeoff
59/// may land differently.
60fn repack_enabled() -> bool {
61    static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
62    *ON.get_or_init(|| {
63        std::env::var("CMF_REPACK")
64            .map(|v| v == "1")
65            .unwrap_or(cfg!(target_os = "android"))
66    })
67}
68
69/// Interleave q8 rows for the decode kernel: group g holds rows
70/// 4g..4g+4 as [r0[c], r1[c], r2[c], r3[c]] per 16-byte chunk c. Only
71/// full groups are packed — tail rows keep reading the mmap layout.
72fn q8_repack(bytes: &[u8], rows: usize, cols: usize) -> Vec<u8> {
73    #[cfg(target_arch = "aarch64")]
74    let arch_ok = sdot_enabled();
75    #[cfg(not(target_arch = "aarch64"))]
76    let arch_ok = false;
77    if !arch_ok || !repack_enabled() || rows < 256 || cols % 16 != 0 {
78        return Vec::new();
79    }
80    q8_repack_layout(bytes, rows, cols)
81}
82
83/// The pure layout transform behind `q8_repack` (tested directly —
84/// the gate depends on arch and env).
85fn q8_repack_layout(bytes: &[u8], rows: usize, cols: usize) -> Vec<u8> {
86    let groups = rows / 4;
87    let mut rep = vec![0u8; groups * 4 * cols];
88    for g in 0..groups {
89        let dst = &mut rep[g * 4 * cols..(g + 1) * 4 * cols];
90        for c in 0..cols / 16 {
91            for lane in 0..4 {
92                let src = (g * 4 + lane) * cols + c * 16;
93                dst[c * 64 + lane * 16..c * 64 + lane * 16 + 16]
94                    .copy_from_slice(&bytes[src..src + 16]);
95            }
96        }
97    }
98    rep
99}
100
101/// Prefix-sum of vbit row payload offsets (absolute within the tensor
102/// bytes). `offsets[r]..offsets[r+1]` is row r's packed data.
103fn vbit_row_offsets(bytes: &[u8], rows: usize, cols: usize) -> Vec<usize> {
104    let ng = cols / GROUP_SIZE;
105    let bits = &bytes[..rows];
106    let mut offsets = Vec::with_capacity(rows + 1);
107    let mut off = rows + rows * ng * 2;
108    for r in 0..rows {
109        offsets.push(off);
110        off += (cols * bits[r] as usize).div_ceil(8);
111    }
112    offsets.push(off);
113    offsets
114}
115
116/// `CMF_X86_BLOCKED` / `CMF_GPU_LMHEAD` / `CMF_GPU_SPLIT`, read once. They
117/// used to be read from the environment on every large matvec and on every
118/// matmat in six places — microseconds each, but also a knob that could
119/// change under a running process, which is not a thing a kernel choice
120/// should be able to do mid-sequence.
121fn blocked_enabled() -> bool {
122    use std::sync::atomic::Ordering::Relaxed;
123    match BLOCKED_OVERRIDE.load(Relaxed) {
124        1 => false,
125        2 => true,
126        _ => {
127            static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
128            *ON.get_or_init(|| {
129                std::env::var("CMF_X86_BLOCKED")
130                    .map(|v| v != "0")
131                    .unwrap_or(true)
132            })
133        }
134    }
135}
136
137static BLOCKED_OVERRIDE: std::sync::atomic::AtomicU8 = std::sync::atomic::AtomicU8::new(0);
138
139/// Force the blocked GEMM on or off, ignoring the environment; `None`
140/// restores it. For tests that need to run BOTH paths and compare them:
141/// `blocked_enabled` caches its answer for the life of the process, which
142/// is right when the environment is the only input, but leaves a test that
143/// flips `CMF_X86_BLOCKED` between two calls comparing a path against
144/// itself — or against whatever a test running in parallel latched first.
145pub fn set_blocked_override(on: Option<bool>) {
146    let v = match on {
147        None => 0,
148        Some(false) => 1,
149        Some(true) => 2,
150    };
151    BLOCKED_OVERRIDE.store(v, std::sync::atomic::Ordering::Relaxed);
152}
153
154fn gpu_lmhead_enabled() -> bool {
155    static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
156    *ON.get_or_init(|| {
157        std::env::var("CMF_GPU_LMHEAD")
158            .map(|v| v != "0")
159            .unwrap_or(true)
160    })
161}
162
163fn gpu_split_frac() -> f32 {
164    // MiMo's banked verification projects every row on the device. Plain
165    // decode must not quantize half of its O/head activation on the CPU.
166    if FULL_GPU_Q8.get() {
167        return 1.0;
168    }
169    static F: std::sync::OnceLock<f32> = std::sync::OnceLock::new();
170    *F.get_or_init(|| {
171        std::env::var("CMF_GPU_SPLIT")
172            .ok()
173            .and_then(|v| v.parse::<f32>().ok())
174            .unwrap_or(0.5)
175            .clamp(0.0, 1.0)
176    })
177}
178
179impl QTensor {
180    pub fn from_f32(data: Vec<f32>, rows: usize, cols: usize) -> Self {
181        debug_assert_eq!(data.len(), rows * cols);
182        Self::F32 { data, rows, cols }
183    }
184
185    /// Wrap a directory tensor without dequantizing the payload.
186    /// Falls back to dequantized f32 for dtypes without a fused kernel.
187    pub fn from_model(model: &Arc<CmfModel>, name: &str) -> Result<Self, String> {
188        // Indexed lookup: the linear directory scan made pipeline build
189        // O(N²) on MoE/skills files with thousands of tensors.
190        let idx = model
191            .tensor_index(name)
192            .ok_or_else(|| format!("tensor '{name}' not found in CMF directory"))?;
193        let entry = &model.tensors[idx];
194        if entry.shape.len() != 2 {
195            return Err(format!("QTensor::from_model needs 2-D, got '{name}'"));
196        }
197        let (rows, cols) = (entry.shape[0], entry.shape[1]);
198        let bytes = model.entry_bytes(entry);
199
200        match entry.dtype {
201            TensorDtype::Q8Row | TensorDtype::Q8_2f => {
202                let n = rows * cols;
203                let scales_off = n;
204                let row_scale: Vec<f32> = (0..rows)
205                    .map(|o| {
206                        f16_to_f32(u16::from_le_bytes([
207                            bytes[scales_off + o * 2],
208                            bytes[scales_off + o * 2 + 1],
209                        ]))
210                    })
211                    .collect();
212                let col_field: Vec<f32> = if entry.dtype == TensorDtype::Q8_2f {
213                    let col_off = n + rows * 2;
214                    (0..cols)
215                        .map(|i| {
216                            f16_to_f32(u16::from_le_bytes([
217                                bytes[col_off + i * 2],
218                                bytes[col_off + i * 2 + 1],
219                            ]))
220                        })
221                        .collect()
222                } else {
223                    Vec::new()
224                };
225                Ok(Self::Mapped {
226                    model: model.clone(),
227                    idx,
228                    dtype: entry.dtype,
229                    rows,
230                    cols,
231                    row_scale,
232                    col_field,
233                    vbit_offsets: Vec::new(),
234                    repack: q8_repack(bytes, rows, cols),
235                })
236            }
237            // vbit: fused kernel unpacks variable-bit rows from mmap.
238            TensorDtype::Vbit if cols % GROUP_SIZE == 0 => Ok(Self::Mapped {
239                model: model.clone(),
240                idx,
241                dtype: entry.dtype,
242                rows,
243                cols,
244                row_scale: Vec::new(),
245                col_field: Vec::new(),
246                vbit_offsets: vbit_row_offsets(bytes, rows, cols),
247                repack: Vec::new(),
248            }),
249            // vbit_ro (§4.2): the offset table comes straight from the
250            // file — no load-time prefix scan; kernels are shared with
251            // legacy vbit (they consume absolute offsets either way).
252            TensorDtype::VbitRo if cols % GROUP_SIZE == 0 => {
253                let (_, off_off, packed_off) = cortiq_core::quant::vbit_ro_sections(rows, cols);
254                let offsets: Vec<usize> = (0..=rows)
255                    .map(|r| packed_off + cortiq_core::quant::vbit_ro_offset(bytes, off_off, r))
256                    .collect();
257                Ok(Self::Mapped {
258                    model: model.clone(),
259                    idx,
260                    dtype: entry.dtype,
261                    rows,
262                    cols,
263                    row_scale: Vec::new(),
264                    col_field: Vec::new(),
265                    vbit_offsets: offsets,
266                    repack: Vec::new(),
267                })
268            }
269            // q4_block: fused kernel reads nibbles straight from mmap —
270            // a 14B q4 file no longer explodes into ×8 f32 RAM.
271            // q4_tiled (§4.3): interleaved [scale][nibbles] tiles — one
272            // sequential memory stream (measured ×1.66 ARM / ×1.13 AVX2
273            // at kernel level over the split layout).
274            TensorDtype::Q4Tiled if cols % GROUP_SIZE == 0 => Ok(Self::Mapped {
275                model: model.clone(),
276                idx,
277                dtype: entry.dtype,
278                rows,
279                cols,
280                row_scale: Vec::new(),
281                col_field: Vec::new(),
282                vbit_offsets: Vec::new(),
283                repack: Vec::new(),
284            }),
285            // q4tp (§4.10): nibbles from mmap, scale from the row ladder —
286            // 7.3% less file than q4t at the same 4-bit grid.
287            TensorDtype::Q4TiledP if cols % GROUP_SIZE == 0 => Ok(Self::Mapped {
288                model: model.clone(),
289                idx,
290                dtype: entry.dtype,
291                rows,
292                cols,
293                row_scale: Vec::new(),
294                col_field: Vec::new(),
295                vbit_offsets: Vec::new(),
296                repack: Vec::new(),
297            }),
298            // q2tp: 2-bit chunks from mmap, scale from the same row ladder.
299            TensorDtype::Q2TiledP if cols % GROUP_SIZE == 0 => Ok(Self::Mapped {
300                model: model.clone(),
301                idx,
302                dtype: entry.dtype,
303                rows,
304                cols,
305                row_scale: Vec::new(),
306                col_field: Vec::new(),
307                vbit_offsets: Vec::new(),
308                repack: Vec::new(),
309            }),
310            TensorDtype::Q4Block if cols % GROUP_SIZE == 0 => Ok(Self::Mapped {
311                model: model.clone(),
312                idx,
313                dtype: entry.dtype,
314                rows,
315                cols,
316                row_scale: Vec::new(),
317                col_field: Vec::new(),
318                vbit_offsets: Vec::new(),
319                repack: Vec::new(),
320            }),
321            // q1: binary sign-bit tiles from mmap (1-bit-trained models).
322            TensorDtype::Q1 if cols % GROUP_SIZE == 0 => Ok(Self::Mapped {
323                model: model.clone(),
324                idx,
325                dtype: entry.dtype,
326                rows,
327                cols,
328                row_scale: Vec::new(),
329                col_field: Vec::new(),
330                vbit_offsets: Vec::new(),
331                repack: Vec::new(),
332            }),
333            // q1t (ternary + outlier overlay): fused per-row dequant kernel
334            // reads straight from mmap — a 12B q1t stays ~its file size in
335            // RAM instead of dequantizing to ~48 GB of f32.
336            TensorDtype::Q1T if cols % GROUP_SIZE == 0 => Ok(Self::Mapped {
337                model: model.clone(),
338                idx,
339                dtype: entry.dtype,
340                rows,
341                cols,
342                row_scale: Vec::new(),
343                col_field: Vec::new(),
344                vbit_offsets: Vec::new(),
345                repack: Vec::new(),
346            }),
347            // No fused kernel yet → dequantize once (correct, more RAM).
348            _ => {
349                let mut data = vec![0.0f32; rows * cols];
350                cortiq_core::quant::dequant_tensor(entry, bytes, &mut data)?;
351                Ok(Self::from_f32(data, rows, cols))
352            }
353        }
354    }
355
356    /// q1-mapped tensor? (GPU gates: the q1 CPU kernel is
357    /// compute-bound, so offload pays at much smaller shapes than q8.)
358    pub(crate) fn is_q1(&self) -> bool {
359        matches!(
360            self,
361            Self::Mapped {
362                dtype: TensorDtype::Q1,
363                ..
364            }
365        )
366    }
367
368    /// Owned-f32 view (data, rows, cols) — the GDN a/b gate projections
369    /// arrive dequantized (force-f16 in the converter → F32 in RAM).
370    pub(crate) fn f32_parts(&self) -> Option<(&[f32], usize, usize)> {
371        match self {
372            Self::F32 { data, rows, cols } => Some((data, *rows, *cols)),
373            _ => None,
374        }
375    }
376
377    /// (directory idx, rows, cols) of a q1-mapped tensor — the
378    /// whole-block GPU path resolves offsets itself.
379    /// (idx, rows, cols) of a mapped tensor the whole-token GPU graph can drive
380    /// — Q1, Q1T or Q4-block (it resolves the offset and picks the kernel by
381    /// dtype). Q4-block lets a precise down_proj/lm_head stay on-device.
382    /// Named `q1_parts` for historical reasons.
383    pub(crate) fn q1_parts(&self) -> Option<(usize, usize, usize)> {
384        if self.has_prism_contract() {
385            return None;
386        }
387        match self {
388            #[cfg(target_os = "macos")]
389            Self::Mapped {
390                dtype: TensorDtype::Q1T,
391                ..
392            } if !crate::gpu::metal_q1t_enabled() => None,
393            Self::Mapped {
394                idx,
395                dtype:
396                    TensorDtype::Q1
397                    | TensorDtype::Q1T
398                    | TensorDtype::Q4Block
399                    | TensorDtype::Q4Tiled
400                    // Q2TiledP deliberately absent: the Metal graph has no
401                    // q2tp kernel, and advertising it here made the block
402                    // plan truncate mid-run at the first q2tp layer.
403                    | TensorDtype::Q4TiledP
404                    | TensorDtype::Q8Row
405                    | TensorDtype::Q8_2f,
406                rows,
407                cols,
408                ..
409            } => Some((*idx, *rows, *cols)),
410            _ => None,
411        }
412    }
413
414    /// `(directory idx, rows, cols)` for the native Metal token graph.  The
415    /// historical q1 graph gate intentionally refuses every Prism tensor so
416    /// an untransformed q2 payload cannot slip into the resident path.  The
417    /// Metal2 graph is descriptor-aware and admits only the production
418    /// q2tp-affine forward targets; ordinary q1/q4 callers retain the old
419    /// `q1_parts` behaviour.
420    #[cfg(target_os = "macos")]
421    pub(crate) fn metal_graph_parts(&self) -> Option<(usize, usize, usize)> {
422        if let Some((model, idx, kind, _)) = self.graph_weight_descriptor() {
423            let name = &model.tensors[idx].name;
424            let forward = kind == 9 && crate::prism::is_forward_weight(model, name);
425            let affine = kind == 9 && crate::prism::is_affine_target(model, name);
426            if forward && affine {
427                let e = model.tensors.get(idx)?;
428                return Some((idx, *e.shape.first()?, *e.shape.get(1)?));
429            }
430        }
431        self.q1_parts()
432    }
433
434    /// (directory idx, rows, cols) of a q4_tiled mapped tensor. The
435    /// chunk-prefill graph takes it in the same 4-tuple slot as
436    /// `q8_row_parts` with an EMPTY row_scale — q4t carries its scales
437    /// inside the 18-byte tiles, and the empty slice is what tells the
438    /// encoder to reach for the q4t kernels.
439    pub(crate) fn q4t_parts(&self) -> Option<(usize, usize, usize)> {
440        if self.has_prism_contract() {
441            return None;
442        }
443        match self {
444            Self::Mapped {
445                idx,
446                dtype: TensorDtype::Q4Tiled,
447                rows,
448                cols,
449                ..
450            } => Some((*idx, *rows, *cols)),
451            _ => None,
452        }
453    }
454
455    /// (directory idx, rows, cols) of a q4tp mapped tensor. Same empty-scale
456    /// slot as `q4t_parts` in the chunk graph — the encoder tells the two
457    /// apart by the tensor's dtype, not by the slot.
458    pub(crate) fn q4tp_parts(&self) -> Option<(usize, usize, usize)> {
459        if self.has_prism_contract() {
460            return None;
461        }
462        match self {
463            Self::Mapped {
464                idx,
465                dtype: TensorDtype::Q4TiledP,
466                rows,
467                cols,
468                ..
469            } => Some((*idx, *rows, *cols)),
470            _ => None,
471        }
472    }
473
474    /// (directory idx, rows, cols, row_scale) of a plain q8_row mapped
475    /// tensor — the chunk-prefill GPU graph resolves offsets itself.
476    /// q8_2f is excluded on purpose: its column field would need a
477    /// prescale stage on the device.
478    pub(crate) fn q8_row_parts(&self) -> Option<(usize, usize, usize, &[f32])> {
479        if self.has_prism_contract() {
480            return None;
481        }
482        match self {
483            Self::Mapped {
484                idx,
485                dtype: TensorDtype::Q8Row,
486                rows,
487                cols,
488                row_scale,
489                col_field,
490                ..
491            } if col_field.is_empty() => Some((*idx, *rows, *cols, row_scale)),
492            _ => None,
493        }
494    }
495
496    /// The layout this tensor is stored in, when it is mapped from a model.
497    /// The frames branch on it — a q2tp gate against a q4tp down is a real
498    /// combination in the 2-bit profile and needs a different kernel.
499    pub fn model_dtype(&self) -> Option<cortiq_core::TensorDtype> {
500        match self {
501            Self::Mapped { dtype, .. } => Some(*dtype),
502            _ => None,
503        }
504    }
505
506    /// The tensor's index in the model directory, when it is mapped from one.
507    /// The GPU frames bind by index rather than by name — a name lookup per
508    /// layer per token is not free, and the index is what the device cache is
509    /// keyed on anyway.
510    pub fn model_idx(&self) -> Option<usize> {
511        match self {
512            Self::Mapped { idx, .. } => Some(*idx),
513            _ => None,
514        }
515    }
516
517    /// The model this tensor is mapped from, when it is mapped at all. The
518    /// GPU frames need the container to reach the bytes; a QTensor already
519    /// holds it, and threading a second handle down every call site to say
520    /// the same thing invites the two to disagree.
521    pub fn model_arc(&self) -> Option<std::sync::Arc<cortiq_core::CmfModel>> {
522        match self {
523            Self::Mapped { model, .. } => Some(model.clone()),
524            _ => None,
525        }
526    }
527
528    /// Whether this mapped tensor belongs to the Prism/Bonsai transform
529    /// contract.  Device graphs do not carry the descriptor, so callers use
530    /// this conservative predicate to stay on the descriptor-aware CPU path
531    /// instead of silently executing an unrotated matrix.
532    pub(crate) fn has_prism_contract(&self) -> bool {
533        matches!(self, Self::Mapped { model, .. } if crate::prism::has_contract(model))
534    }
535
536    pub fn rows(&self) -> usize {
537        match self {
538            Self::F32 { rows, .. } | Self::Mapped { rows, .. } => *rows,
539        }
540    }
541
542    /// Mapped q4t handle (model + directory index) — the fused GPU FFN
543    /// needs the raw file coordinates of its three projections.
544    pub(crate) fn mapped_q4t(&self) -> Option<(&Arc<CmfModel>, usize)> {
545        if self.has_prism_contract() {
546            return None;
547        }
548        match self {
549            Self::Mapped {
550                model,
551                idx,
552                dtype: TensorDtype::Q4Tiled,
553                ..
554            } => Some((model, *idx)),
555            _ => None,
556        }
557    }
558
559    /// Same slot as `mapped_q4t` for a q4tp tensor — the fused DiT FFN picks
560    /// its kernels by which of the two answers.
561    pub fn mapped_q4tp(&self) -> Option<(&Arc<CmfModel>, usize)> {
562        if self.has_prism_contract() {
563            return None;
564        }
565        match self {
566            Self::Mapped {
567                model,
568                idx,
569                dtype: TensorDtype::Q4TiledP,
570                ..
571            } => Some((model, *idx)),
572            _ => None,
573        }
574    }
575
576    /// (model, tensor idx) for a mapped weight in ANY codec the fused device
577    /// paths can run — four-bit tiled or either int8 layout.
578    ///
579    /// The fused DiT chains asked for `mapped_q4tp` by name, so an eight-bit
580    /// container never reached them and rendered through per-op GEMMs even
581    /// after those kernels learned its codec. The gate is what the codec has
582    /// a device GEMM for, not which codec it is.
583    pub fn mapped_device_gemm(&self) -> Option<(&Arc<CmfModel>, usize)> {
584        if self.has_prism_contract() {
585            return None;
586        }
587        match self {
588            Self::Mapped {
589                model,
590                idx,
591                dtype: TensorDtype::Q4TiledP | TensorDtype::Q8Row | TensorDtype::Q8_2f,
592                ..
593            } => Some((model, *idx)),
594            _ => None,
595        }
596    }
597
598    /// (model, tensor idx) for a q2tp mapped weight — the 2-bit twin of
599    /// `mapped_q4tp`, used by the mixed MoE profile.
600    pub fn mapped_q2tp(&self) -> Option<(&Arc<CmfModel>, usize)> {
601        if self.has_prism_contract() {
602            return None;
603        }
604        match self {
605            Self::Mapped {
606                model,
607                idx,
608                dtype: TensorDtype::Q2TiledP,
609                ..
610            } => Some((model, *idx)),
611            _ => None,
612        }
613    }
614
615    pub fn cols(&self) -> usize {
616        match self {
617            Self::F32 { cols, .. } | Self::Mapped { cols, .. } => *cols,
618        }
619    }
620
621    /// (model, tensor idx) for a q1 mapped weight — the wgpu token graph
622    /// keys its resident VRAM cache by idx. None for any other dtype/kind.
623    pub fn mapped_q1(&self) -> Option<(&std::sync::Arc<CmfModel>, usize)> {
624        if self.has_prism_contract() {
625            return None;
626        }
627        match self {
628            Self::Mapped {
629                model,
630                idx,
631                dtype: TensorDtype::Q1,
632                ..
633            } => Some((model, *idx)),
634            _ => None,
635        }
636    }
637
638    /// (model, idx, kind, row_scale) for a graph-capable mapped weight.
639    /// kind: 0=q8_row (per-row scales), 1=q1, 2=q4_block, 3=q1t
640    /// (tile-embedded, no rs), 5=q4_tiled, 6=q4tp, 7=q8_2f (both scale
641    /// planes live inside the tensor). None only for `vbit`.
642    ///
643    /// The old comment here claimed q4_block was unhandled while the arm
644    /// right below mapped it, and it named q8_2f as unhandled after that
645    /// stopped being true — a stale comment on this function is how a
646    /// model silently loses the graph, so it is worth keeping honest.
647    pub fn graph_weight(&self) -> Option<(&std::sync::Arc<CmfModel>, usize, u8, &[f32])> {
648        if self.has_prism_contract() {
649            return None;
650        }
651        self.graph_weight_descriptor()
652    }
653
654    /// Descriptor-aware graph handle used only by the Prism token graph.
655    /// Ordinary graph callers continue to use [`graph_weight`] and therefore
656    /// remain fail-closed until they provide the same explicit transform
657    /// contract.
658    pub(crate) fn graph_weight_descriptor(
659        &self,
660    ) -> Option<(&std::sync::Arc<CmfModel>, usize, u8, &[f32])> {
661        match self {
662            Self::Mapped {
663                model,
664                idx,
665                dtype: TensorDtype::Q8Row,
666                row_scale,
667                ..
668            } => Some((model, *idx, 0, row_scale.as_slice())),
669            Self::Mapped {
670                model,
671                idx,
672                dtype: TensorDtype::Q1,
673                ..
674            } => Some((model, *idx, 1, &[])),
675            // Q4Tiled is kind 5, NOT 2: both carried 2 historically, and
676            // the wgpu token graph fed 18B interleaved tiles to the
677            // split-layout q4b kernel — garbage output on q4t models
678            // (caught by an end-to-end answer check on real Vulkan).
679            Self::Mapped {
680                model,
681                idx,
682                dtype: TensorDtype::Q4Tiled,
683                ..
684            } => Some((model, *idx, 5, &[])),
685            // Kind 6, not 5: q4tp's nibble stride and scale planes differ,
686            // and feeding them to the q4t kernel is exactly the mistake that
687            // produced garbage when Q4Tiled shared kind 2 with Q4Block.
688            Self::Mapped {
689                model,
690                idx,
691                dtype: TensorDtype::Q4TiledP,
692                ..
693            } => Some((model, *idx, 6, &[])),
694            Self::Mapped {
695                model,
696                idx,
697                dtype: TensorDtype::Q4Block,
698                ..
699            } => Some((model, *idx, 2, &[])),
700            // q8_2f carries BOTH scale planes after the int8 body (rows
701            // f16, then cols f16), so the graph takes the whole tensor
702            // and the kernel reads them where they lie — no host-side
703            // prescale, which is what the per-op path does instead.
704            Self::Mapped {
705                model,
706                idx,
707                dtype: TensorDtype::Q8_2f,
708                ..
709            } => Some((model, *idx, 7, &[])),
710            Self::Mapped {
711                model,
712                idx,
713                dtype: TensorDtype::Q1T,
714                ..
715            } => Some((model, *idx, 3, &[])),
716            // Kind 9: the 2-bit plane on the q4tp ladder (dense FFN gate/up
717            // of the q2tp profile). Its own kernel — 8 bytes a group where
718            // q4tp has 16, and rung 0 is the exact zero.
719            Self::Mapped {
720                model,
721                idx,
722                dtype: TensorDtype::Q2TiledP,
723                ..
724            } => Some((model, *idx, 9, &[])),
725            _ => None,
726        }
727    }
728
729    /// Dense f32 view — only for owned tensors. Masked/sparse execution
730    /// paths require it; quantized weights don't support masks yet.
731    pub fn as_f32(&self) -> Option<&[f32]> {
732        match self {
733            Self::F32 { data, .. } => Some(data),
734            Self::Mapped { .. } => None,
735        }
736    }
737
738    fn quant_bytes(&self) -> &[u8] {
739        match self {
740            Self::Mapped { model, idx, .. } => model.entry_bytes(&model.tensors[*idx]),
741            Self::F32 { .. } => unreachable!("quant_bytes on F32"),
742        }
743    }
744
745    /// Dequantize one row into `dst` (embedding lookup).
746    pub fn row_f32(&self, r: usize, dst: &mut [f32]) {
747        let cols = self.cols();
748        debug_assert_eq!(dst.len(), cols);
749        match self {
750            Self::F32 { data, .. } => dst.copy_from_slice(&data[r * cols..(r + 1) * cols]),
751            Self::Mapped {
752                model,
753                idx,
754                dtype,
755                row_scale,
756                col_field,
757                vbit_offsets,
758                ..
759            } => {
760                if *dtype == TensorDtype::Q4Tiled {
761                    let bytes = self.quant_bytes();
762                    let gpr = cols / GROUP_SIZE;
763                    for gi in 0..gpr {
764                        let tile = &bytes[(r * gpr + gi) * Q4_TILE..(r * gpr + gi + 1) * Q4_TILE];
765                        let s = f16_to_f32(u16::from_le_bytes([tile[0], tile[1]]));
766                        for (k, &b) in tile[2..].iter().enumerate() {
767                            dst[gi * GROUP_SIZE + k * 2] = ((b & 0x0F) as f32 - 8.0) * s;
768                            dst[gi * GROUP_SIZE + k * 2 + 1] = (((b >> 4) & 0x0F) as f32 - 8.0) * s;
769                        }
770                    }
771                    if crate::prism::is_inverse_embedding(model, &model.tensors[*idx].name) {
772                        crate::prism::inverse_embedding(model, dst);
773                    }
774                    return;
775                }
776                if *dtype == TensorDtype::Q4TiledP {
777                    let bytes = self.quant_bytes();
778                    let gpr = cols / GROUP_SIZE;
779                    let v = Q4tpView::new(bytes, self.rows(), cols);
780                    let mut sc = vec![0f32; gpr];
781                    v.scales_into(r, gpr, &mut sc);
782                    for gi in 0..gpr {
783                        let tile = &v.nib[(r * gpr + gi) * Q4TP_NIB..(r * gpr + gi + 1) * Q4TP_NIB];
784                        let s = sc[gi];
785                        for (k, &b) in tile.iter().enumerate() {
786                            dst[gi * GROUP_SIZE + k * 2] = ((b & 0x0F) as f32 - 8.0) * s;
787                            dst[gi * GROUP_SIZE + k * 2 + 1] = (((b >> 4) & 0x0F) as f32 - 8.0) * s;
788                        }
789                    }
790                    if crate::prism::is_inverse_embedding(model, &model.tensors[*idx].name) {
791                        crate::prism::inverse_embedding(model, dst);
792                    }
793                    return;
794                }
795                if *dtype == TensorDtype::Q2TiledP {
796                    let bytes = self.quant_bytes();
797                    let gpr = cols / GROUP_SIZE;
798                    let v = Q4tpView::new_q2(bytes, self.rows(), cols);
799                    let mut sc = vec![0f32; gpr];
800                    v.scales_into(r, gpr, &mut sc);
801                    for gi in 0..gpr {
802                        let ch =
803                            &v.nib[(r * gpr + gi) * Q2TP_CHUNK..(r * gpr + gi + 1) * Q2TP_CHUNK];
804                        let s = sc[gi];
805                        for (k, &b) in ch.iter().enumerate() {
806                            for j in 0..4 {
807                                let center = if crate::prism::is_affine_target(
808                                    model,
809                                    &model.tensors[*idx].name,
810                                ) {
811                                    1.0
812                                } else {
813                                    1.5
814                                };
815                                dst[gi * GROUP_SIZE + k * 4 + j] =
816                                    (((b >> (2 * j)) & 3) as f32 - center) * s;
817                            }
818                        }
819                    }
820                    if crate::prism::is_inverse_embedding(model, &model.tensors[*idx].name) {
821                        crate::prism::inverse_embedding(model, dst);
822                    }
823                    return;
824                }
825                if *dtype == TensorDtype::Q4Block {
826                    let (packed, scales) = q4_split(self.quant_bytes(), self.rows(), cols);
827                    let gpr = cols / GROUP_SIZE;
828                    for gi in 0..gpr {
829                        let g = r * gpr + gi;
830                        let s = f16_to_f32(u16::from_le_bytes([scales[g * 2], scales[g * 2 + 1]]));
831                        for (k, &b) in packed[g * 16..(g + 1) * 16].iter().enumerate() {
832                            dst[gi * GROUP_SIZE + k * 2] = ((b & 0x0F) as f32 - 8.0) * s;
833                            dst[gi * GROUP_SIZE + k * 2 + 1] = (((b >> 4) & 0x0F) as f32 - 8.0) * s;
834                        }
835                    }
836                    if crate::prism::is_inverse_embedding(model, &model.tensors[*idx].name) {
837                        crate::prism::inverse_embedding(model, dst);
838                    }
839                    return;
840                }
841                if *dtype == TensorDtype::Q1 {
842                    let bytes = self.quant_bytes();
843                    let gpr = cols / GROUP_SIZE;
844                    for gi in 0..gpr {
845                        let tile = &bytes[(r * gpr + gi) * Q1_TILE..(r * gpr + gi + 1) * Q1_TILE];
846                        let s = f16_to_f32(u16::from_le_bytes([tile[0], tile[1]]));
847                        for (j, &b) in tile[2..].iter().enumerate() {
848                            for k in 0..8 {
849                                dst[gi * GROUP_SIZE + j * 8 + k] =
850                                    (((b >> k) & 1) as f32 * 2.0 - 1.0) * s;
851                            }
852                        }
853                    }
854                    if crate::prism::is_inverse_embedding(model, &model.tensors[*idx].name) {
855                        crate::prism::inverse_embedding(model, dst);
856                    }
857                    return;
858                }
859                if *dtype == TensorDtype::Q1T {
860                    let bytes = self.quant_bytes();
861                    let gpr = cols / GROUP_SIZE;
862                    let base_len = self.rows() * gpr * cortiq_core::quant::Q1T_TILE;
863                    for gi in 0..gpr {
864                        let off = (r * gpr + gi) * cortiq_core::quant::Q1T_TILE;
865                        let s = cortiq_core::quant::f16_to_f32(u16::from_le_bytes([
866                            bytes[off],
867                            bytes[off + 1],
868                        ]));
869                        let codes = &bytes[off + 2..off + cortiq_core::quant::Q1T_TILE];
870                        for k in 0..GROUP_SIZE {
871                            dst[gi * GROUP_SIZE + k] = match cortiq_core::quant::q1t_code(codes, k)
872                            {
873                                1 => s,
874                                2 => -s,
875                                _ => 0.0,
876                            };
877                        }
878                    }
879                    // Overlay
880                    let rows = self.rows();
881                    let entries = base_len + (rows + 1) * 4;
882                    if entries <= bytes.len() {
883                        let ptrs = &bytes[base_len..base_len + (rows + 1) * 4];
884                        let r0 = u32::from_le_bytes([
885                            ptrs[r * 4],
886                            ptrs[r * 4 + 1],
887                            ptrs[r * 4 + 2],
888                            ptrs[r * 4 + 3],
889                        ]) as usize;
890                        let r1 = u32::from_le_bytes([
891                            ptrs[(r + 1) * 4],
892                            ptrs[(r + 1) * 4 + 1],
893                            ptrs[(r + 1) * 4 + 2],
894                            ptrs[(r + 1) * 4 + 3],
895                        ]) as usize;
896                        let off = entries + r0 * 4;
897                        for i in 0..r1 - r0 {
898                            let item = &bytes[off + i * 4..off + i * 4 + 4];
899                            let c = u16::from_le_bytes([item[0], item[1]]) as usize;
900                            let v = cortiq_core::quant::f16_to_f32(u16::from_le_bytes([
901                                item[2], item[3],
902                            ]));
903                            if c < cols {
904                                dst[c] = v;
905                            }
906                        }
907                    }
908                    if crate::prism::is_inverse_embedding(model, &model.tensors[*idx].name) {
909                        crate::prism::inverse_embedding(model, dst);
910                    }
911                    return;
912                }
913                if matches!(dtype, TensorDtype::Vbit | TensorDtype::VbitRo) {
914                    let bytes = self.quant_bytes();
915                    let rows = self.rows();
916                    let ng = cols / GROUP_SIZE;
917                    let bits = &bytes[..rows];
918                    let sc_off = rows;
919                    // Precomputed at load — embedding lookup used to scan
920                    // the bit-widths of every preceding row (O(token_id)).
921                    let off = vbit_offsets[r];
922                    let b = bits[r] as usize;
923                    let l = ((1usize << (b - 1)) - 1) as f32;
924                    let data = &bytes[off..];
925                    let (mut acc, mut nbits, mut byte_idx) = (0u64, 0usize, 0usize);
926                    for (i, d) in dst.iter_mut().enumerate() {
927                        while nbits < b {
928                            acc = (acc << 8) | data[byte_idx] as u64;
929                            byte_idx += 1;
930                            nbits += 8;
931                        }
932                        let u = ((acc >> (nbits - b)) & ((1u64 << b) - 1)) as f32;
933                        nbits -= b;
934                        let so = (r * ng + i / GROUP_SIZE) * 2;
935                        let sv = f16_to_f32(u16::from_le_bytes([
936                            bytes[sc_off + so],
937                            bytes[sc_off + so + 1],
938                        ]));
939                        *d = (u - l) * sv;
940                    }
941                    if crate::prism::is_inverse_embedding(model, &model.tensors[*idx].name) {
942                        crate::prism::inverse_embedding(model, dst);
943                    }
944                    return;
945                }
946                let q = &self.quant_bytes()[r * cols..(r + 1) * cols];
947                let s = row_scale[r];
948                match dtype {
949                    TensorDtype::Q8Row => {
950                        for (d, &b) in dst.iter_mut().zip(q) {
951                            *d = (b as i8) as f32 * s;
952                        }
953                    }
954                    TensorDtype::Q8_2f => {
955                        for (i, (d, &b)) in dst.iter_mut().zip(q).enumerate() {
956                            *d = (b as i8) as f32 * s * col_field[i];
957                        }
958                    }
959                    _ => unreachable!(),
960                }
961                if crate::prism::is_inverse_embedding(model, &model.tensors[*idx].name) {
962                    crate::prism::inverse_embedding(model, dst);
963                }
964            }
965        }
966    }
967
968    /// Can this tensor's columns be read cheaply (for sparse down_proj)?
969    /// True for F32/Q8Row/Q8_2f (per-row scale, direct strided access);
970    /// false for group-packed q4/vbit (column access would unpack whole
971    /// groups — sparse execution falls back to f32 for those).
972    pub fn sparse_col_ok(&self) -> bool {
973        match self {
974            Self::F32 { .. } => true,
975            Self::Mapped { dtype, .. } => {
976                matches!(dtype, TensorDtype::Q8Row | TensorDtype::Q8_2f)
977            }
978        }
979    }
980
981    /// down_proj [hidden, inter]: accumulate `w · col(c)` into `out`
982    /// [hidden] — reads ONLY column `c` (one neuron) from the mmap,
983    /// no full-matrix dequant. `out[k] += w · down[k, c]`.
984    pub fn add_col_scaled(&self, c: usize, w: f32, out: &mut [f32]) {
985        let inter = self.cols();
986        let hidden = self.rows();
987        debug_assert_eq!(out.len(), hidden);
988        match self {
989            Self::F32 { data, .. } => {
990                for (k, o) in out.iter_mut().enumerate() {
991                    *o += w * data[k * inter + c];
992                }
993            }
994            Self::Mapped {
995                dtype,
996                row_scale,
997                col_field,
998                ..
999            } => {
1000                let q = self.quant_bytes();
1001                let colf = if *dtype == TensorDtype::Q8_2f {
1002                    col_field[c]
1003                } else {
1004                    1.0
1005                };
1006                let wc = w * colf;
1007                for (k, o) in out.iter_mut().enumerate() {
1008                    let b = q[k * inter + c] as i8 as f32;
1009                    *o += wc * b * row_scale[k];
1010                }
1011            }
1012        }
1013    }
1014
1015    /// Touch the head of row `r` so the DRAM latency of the next
1016    /// neuron's weights overlaps the current one's arithmetic.
1017    ///
1018    /// Scattered rows are what per-token sparsity reads, and a 2 KB
1019    /// stride is past what the hardware prefetcher follows: without this
1020    /// every row starts with a cold miss that nothing hides. One touch
1021    /// per 512 bytes is enough — the rest of the row is a sequential run
1022    /// the prefetcher does pick up.
1023    #[inline]
1024    pub fn prefetch_row(&self, r: usize) {
1025        let Self::Mapped { dtype, .. } = self else {
1026            return;
1027        };
1028        if !matches!(dtype, TensorDtype::Q8Row | TensorDtype::Q8_2f) {
1029            return;
1030        }
1031        let cols = self.cols();
1032        let q = self.quant_bytes();
1033        let (a, b) = (r * cols, (r + 1) * cols);
1034        if b > q.len() {
1035            return;
1036        }
1037        let mut j = a;
1038        while j < b {
1039            unsafe { std::ptr::read_volatile(q.as_ptr().add(j)) };
1040            j += 512;
1041        }
1042    }
1043
1044    /// `out += w · row(r)` — the transposed twin of `add_col_scaled`.
1045    ///
1046    /// A neuron's `down` weights are a COLUMN of `[hidden, inter]`, and a
1047    /// column is strided: reading one costs a cache line per element, so
1048    /// per-neuron dynamic sparsity saves arithmetic and no bytes. Stored
1049    /// transposed (`down_proj.t.weight`, `[inter, hidden]`) the same
1050    /// weights are a contiguous ROW, and this accumulate reads exactly
1051    /// the neurons the token asked for.
1052    pub fn add_row_scaled(&self, r: usize, w: f32, out: &mut [f32], scratch: &mut [f32]) {
1053        let cols = self.cols();
1054        debug_assert_eq!(out.len(), cols);
1055        match self {
1056            Self::F32 { data, .. } => {
1057                let row = &data[r * cols..(r + 1) * cols];
1058                for (o, v) in out.iter_mut().zip(row) {
1059                    *o += w * v;
1060                }
1061            }
1062            Self::Mapped {
1063                dtype,
1064                row_scale,
1065                col_field,
1066                ..
1067            } => match dtype {
1068                TensorDtype::Q8Row => {
1069                    let q = &self.quant_bytes()[r * cols..(r + 1) * cols];
1070                    let ws = w * row_scale[r];
1071                    let row: &[i8] =
1072                        unsafe { std::slice::from_raw_parts(q.as_ptr() as *const i8, q.len()) };
1073                    axpy_i8_f32(out, row, ws);
1074                }
1075                TensorDtype::Q8_2f => {
1076                    let q = &self.quant_bytes()[r * cols..(r + 1) * cols];
1077                    let ws = w * row_scale[r];
1078                    for ((o, b), c) in out.iter_mut().zip(q).zip(col_field) {
1079                        *o += ws * c * (*b as i8 as f32);
1080                    }
1081                }
1082                _ => {
1083                    self.row_f32(r, scratch);
1084                    for (o, v) in out.iter_mut().zip(scratch.iter()) {
1085                        *o += w * v;
1086                    }
1087                }
1088            },
1089        }
1090    }
1091
1092    /// Dot of row `r` with `x` (gate/up active-neuron path). Reads only
1093    /// row `r` from the mmap — no full dequant. q4/vbit dequant the row
1094    /// into `scratch` first (rare for active-FFN weights).
1095    pub fn row_dot(&self, r: usize, x: &[f32], scratch: &mut [f32]) -> f32 {
1096        let cols = self.cols();
1097        match self {
1098            Self::F32 { data, .. } => {
1099                let row = &data[r * cols..(r + 1) * cols];
1100                row.iter().zip(x).map(|(w, v)| w * v).sum()
1101            }
1102            Self::Mapped {
1103                model,
1104                idx,
1105                dtype,
1106                row_scale,
1107                col_field,
1108                ..
1109            } => {
1110                let prism_forward =
1111                    crate::prism::is_forward_weight(model, &model.tensors[*idx].name);
1112                if prism_forward {
1113                    let transformed = crate::prism::forward(model, &x[..cols]);
1114                    let gpr = cols / GROUP_SIZE;
1115                    match dtype {
1116                        TensorDtype::Q2TiledP => {
1117                            let v = Q4tpView::new_q2(self.quant_bytes(), self.rows(), cols);
1118                            let mut sc = vec![0f32; gpr];
1119                            v.scales_into(r, gpr, &mut sc);
1120                            if crate::prism::is_affine_target(model, &model.tensors[*idx].name) {
1121                                return q2tp_affine_row_exact(v.nib, r, gpr, &transformed, &sc);
1122                            }
1123                            return q2tp_row_exact(v.nib, r, gpr, &transformed, &sc);
1124                        }
1125                        _ => {
1126                            self.row_f32(r, scratch);
1127                            return scratch.iter().zip(&transformed).map(|(w, v)| w * v).sum();
1128                        }
1129                    }
1130                }
1131                match dtype {
1132                    TensorDtype::Q8Row => {
1133                        let q = &self.quant_bytes()[r * cols..(r + 1) * cols];
1134                        dot_i8_f32(q, x) * row_scale[r]
1135                    }
1136                    TensorDtype::Q8_2f => {
1137                        let q = &self.quant_bytes()[r * cols..(r + 1) * cols];
1138                        dot_i8_col_f32(q, x, col_field) * row_scale[r]
1139                    }
1140                    _ => {
1141                        self.row_f32(r, scratch);
1142                        scratch.iter().zip(x).map(|(w, v)| w * v).sum()
1143                    }
1144                }
1145            }
1146        }
1147    }
1148
1149    /// `out = W · x` (row-major). F32 delegates to the historical
1150    /// bit-exact path; Mapped runs the fused int8 kernel.
1151    pub fn matvec(&self, x: &[f32], out: &mut [f32], pool: Option<&Pool>) {
1152        match self {
1153            // NOTE: `out.len()` DRIVES this arm — it computes that many rows,
1154            // and `x.len()` is the stride. A short `out` is legitimate here,
1155            // which is why the check below lives in the Mapped arm only.
1156            Self::F32 { data, .. } => matvec_rows(pool, data, x, out),
1157            Self::Mapped {
1158                model,
1159                idx,
1160                dtype,
1161                rows,
1162                cols,
1163                row_scale,
1164                col_field,
1165                vbit_offsets,
1166                repack,
1167            } => {
1168                let _ = (model, idx);
1169                // Every kernel below writes `rows` entries through a raw
1170                // pointer, so a short `out` is an out-of-bounds WRITE, not a
1171                // wrong answer: it scribbles on the allocator's metadata and
1172                // the process aborts much later, somewhere innocent
1173                // (`double free or corruption`, `corrupted double-linked
1174                // list`). The debug_assert two of the kernels carried is
1175                // compiled out of the release — exactly the build where it
1176                // matters. Fail here instead, while the caller is still on
1177                // the stack to be named.
1178                assert!(
1179                    out.len() >= *rows && x.len() >= *cols,
1180                    "matvec {rows}x{cols}: out {} (need {rows}), x {} (need {cols})",
1181                    out.len(),
1182                    x.len(),
1183                );
1184                let prism_forward =
1185                    crate::prism::is_forward_weight(model, &model.tensors[*idx].name);
1186                if *dtype == TensorDtype::Q2TiledP
1187                    && std::env::var("CMF_Q2TP_TRACE").as_deref() == Ok("1")
1188                {
1189                    use std::sync::atomic::{AtomicUsize, Ordering};
1190                    static N: AtomicUsize = AtomicUsize::new(0);
1191                    let n = N.fetch_add(1, Ordering::Relaxed);
1192                    if n < 128 {
1193                        eprintln!(
1194                            "q2tp-dispatch #{n} name={} prism={} rows={} cols={} gpu={} optin={} layer={}",
1195                            model.tensors[*idx].name,
1196                            prism_forward,
1197                            rows,
1198                            cols,
1199                            crate::gpu::enabled_here(),
1200                            crate::gpu::q2tp_gpu_opt_in(),
1201                            crate::gpu::cur_layer(),
1202                        );
1203                    }
1204                }
1205                // Prism stores every manifest-listed forward matrix in the
1206                // signed-Hadamard basis.  The q2tp WGSL path receives that
1207                // transformed vector and an explicit affine bit; codecs
1208                // without a descriptor-aware kernel remain on CPU below.
1209                if prism_forward {
1210                    let transformed = crate::prism::forward(model, &x[..*cols]);
1211                    match dtype {
1212                        TensorDtype::Q4Block => {
1213                            q4matvec(self.quant_bytes(), &transformed, *rows, *cols, out, pool)
1214                        }
1215                        TensorDtype::Q4Tiled => {
1216                            q4t_matvec(self.quant_bytes(), &transformed, *rows, *cols, out, pool)
1217                        }
1218                        TensorDtype::Q4TiledP => {
1219                            q4tp_matvec(self.quant_bytes(), &transformed, *rows, *cols, out, pool)
1220                        }
1221                        TensorDtype::Q2TiledP => {
1222                            let affine =
1223                                crate::prism::is_affine_target(model, &model.tensors[*idx].name);
1224                            if *rows * *cols >= 8_388_608
1225                                && crate::gpu::enabled_here()
1226                                && crate::gpu::q2tp_gpu_opt_in()
1227                            {
1228                                let gpu_ok = if affine {
1229                                    crate::gpu::q2tp_affine_matvec(
1230                                        model,
1231                                        *idx,
1232                                        &transformed,
1233                                        *rows,
1234                                        *cols,
1235                                        out,
1236                                    )
1237                                } else {
1238                                    crate::gpu::q2tp_matvec(
1239                                        model,
1240                                        *idx,
1241                                        &transformed,
1242                                        *rows,
1243                                        *cols,
1244                                        out,
1245                                    )
1246                                };
1247                                if gpu_ok {
1248                                    return;
1249                                }
1250                            }
1251                            if affine {
1252                                q2tp_affine_matvec(
1253                                    self.quant_bytes(),
1254                                    &transformed,
1255                                    *rows,
1256                                    *cols,
1257                                    out,
1258                                    pool,
1259                                )
1260                            } else {
1261                                q2tp_matvec(
1262                                    self.quant_bytes(),
1263                                    &transformed,
1264                                    *rows,
1265                                    *cols,
1266                                    out,
1267                                    pool,
1268                                )
1269                            }
1270                        }
1271                        TensorDtype::Q1 => {
1272                            q1_matvec(self.quant_bytes(), &transformed, *rows, *cols, out, pool)
1273                        }
1274                        TensorDtype::Q1T => {
1275                            q1t_matvec(self.quant_bytes(), &transformed, *rows, *cols, out, pool)
1276                        }
1277                        TensorDtype::Vbit | TensorDtype::VbitRo => vbitmatvec(
1278                            self.quant_bytes(),
1279                            vbit_offsets,
1280                            &transformed,
1281                            *rows,
1282                            *cols,
1283                            out,
1284                            pool,
1285                        ),
1286                        TensorDtype::Q8Row | TensorDtype::Q8_2f => qmatvec(
1287                            self.quant_bytes(),
1288                            repack,
1289                            row_scale,
1290                            &transformed,
1291                            col_field,
1292                            *dtype,
1293                            *rows,
1294                            *cols,
1295                            out,
1296                            pool,
1297                        ),
1298                        _ => unreachable!("unsupported mapped Prism dtype {dtype:?}"),
1299                    }
1300                    return;
1301                }
1302                if *dtype == TensorDtype::Q4Block {
1303                    // GPU route (wgpu q4b kernel) for large q4_block matvecs —
1304                    // gives NVIDIA/AMD/Intel q4 models a GPU path. Probe keeps
1305                    // the winner; Metal returns false → the CPU kernel below.
1306                    if *rows * *cols >= 8_388_608 && crate::gpu::enabled_here() {
1307                        let t0 = std::time::Instant::now();
1308                        match crate::gpu::probe_arm(crate::gpu::OpClass::Matvec) {
1309                            crate::gpu::ProbeArm::Gpu => {
1310                                if crate::gpu::q4b_matvec(model, *idx, x, *rows, *cols, out) {
1311                                    crate::gpu::probe_record(
1312                                        crate::gpu::OpClass::Matvec,
1313                                        true,
1314                                        t0.elapsed(),
1315                                    );
1316                                    return;
1317                                }
1318                            }
1319                            crate::gpu::ProbeArm::CpuTimed => {
1320                                q4matvec(self.quant_bytes(), x, *rows, *cols, out, pool);
1321                                crate::gpu::probe_record(
1322                                    crate::gpu::OpClass::Matvec,
1323                                    false,
1324                                    t0.elapsed(),
1325                                );
1326                                return;
1327                            }
1328                            crate::gpu::ProbeArm::Cpu => {}
1329                        }
1330                    }
1331                    q4matvec(self.quant_bytes(), x, *rows, *cols, out, pool);
1332                    return;
1333                }
1334                if *dtype == TensorDtype::Q4Tiled {
1335                    // GPU route for large q4t matvecs — the lm_head class,
1336                    // same shape as the q4tp arm below. The probe keeps the
1337                    // winner; a backend without the kernel refuses and the
1338                    // CPU path stays.
1339                    if *rows * *cols >= 8_388_608 && crate::gpu::enabled_here() {
1340                        let t0 = std::time::Instant::now();
1341                        let cls = crate::gpu::matvec_class(*rows, *cols);
1342                        match crate::gpu::probe_arm(cls) {
1343                            crate::gpu::ProbeArm::Gpu => {
1344                                if crate::gpu::q4t_matvec(model, *idx, x, *rows, *cols, out) {
1345                                    crate::gpu::probe_record(cls, true, t0.elapsed());
1346                                    return;
1347                                }
1348                            }
1349                            crate::gpu::ProbeArm::CpuTimed => {
1350                                q4t_matvec(self.quant_bytes(), x, *rows, *cols, out, pool);
1351                                crate::gpu::probe_record(cls, false, t0.elapsed());
1352                                return;
1353                            }
1354                            crate::gpu::ProbeArm::Cpu => {}
1355                        }
1356                    }
1357                    q4t_matvec(self.quant_bytes(), x, *rows, *cols, out, pool);
1358                    return;
1359                }
1360                if *dtype == TensorDtype::Q4TiledP {
1361                    // GPU route for large q4tp matvecs — the lm_head class.
1362                    // On a q4tp checkpoint the head is the biggest single
1363                    // host matvec left in the decode step, and the batched
1364                    // kernel at b=1 already exists on both backends. Probe
1365                    // keeps the winner, same as q4_block above.
1366                    if *rows * *cols >= 8_388_608 && crate::gpu::enabled_here() {
1367                        let t0 = std::time::Instant::now();
1368                        let cls = crate::gpu::matvec_class(*rows, *cols);
1369                        match crate::gpu::probe_arm(cls) {
1370                            crate::gpu::ProbeArm::Gpu => {
1371                                if crate::gpu::q4tp_matvec(model, *idx, x, *rows, *cols, out) {
1372                                    crate::gpu::probe_record(cls, true, t0.elapsed());
1373                                    return;
1374                                }
1375                            }
1376                            crate::gpu::ProbeArm::CpuTimed => {
1377                                q4tp_matvec(self.quant_bytes(), x, *rows, *cols, out, pool);
1378                                crate::gpu::probe_record(cls, false, t0.elapsed());
1379                                return;
1380                            }
1381                            crate::gpu::ProbeArm::Cpu => {}
1382                        }
1383                    }
1384                    q4tp_matvec(self.quant_bytes(), x, *rows, *cols, out, pool);
1385                    return;
1386                }
1387                if *dtype == TensorDtype::Q2TiledP {
1388                    q2tp_matvec(self.quant_bytes(), x, *rows, *cols, out, pool);
1389                    return;
1390                }
1391                if *dtype == TensorDtype::Q1 {
1392                    // GPU route for large q1 matvecs (out_proj / lm_head
1393                    // class): the CPU q1 kernel is load-port-bound at
1394                    // ~4 GB/s/core, the GPU one is bandwidth-bound — the
1395                    // probe measures both arms and keeps the winner.
1396                    if *rows * *cols >= 8_388_608 && crate::gpu::enabled_here() {
1397                        let t0 = std::time::Instant::now();
1398                        let arm = if crate::gpu::q1_force() {
1399                            crate::gpu::ProbeArm::Gpu
1400                        } else {
1401                            crate::gpu::probe_arm(crate::gpu::OpClass::Matvec)
1402                        };
1403                        match arm {
1404                            crate::gpu::ProbeArm::Gpu => {
1405                                if crate::gpu::q1_matvec(model, *idx, x, *rows, *cols, out) {
1406                                    crate::gpu::probe_record(
1407                                        crate::gpu::OpClass::Matvec,
1408                                        true,
1409                                        t0.elapsed(),
1410                                    );
1411                                    return;
1412                                }
1413                            }
1414                            crate::gpu::ProbeArm::CpuTimed => {
1415                                q1_matvec(self.quant_bytes(), x, *rows, *cols, out, pool);
1416                                crate::gpu::probe_record(
1417                                    crate::gpu::OpClass::Matvec,
1418                                    false,
1419                                    t0.elapsed(),
1420                                );
1421                                return;
1422                            }
1423                            crate::gpu::ProbeArm::Cpu => {}
1424                        }
1425                    }
1426                    q1_matvec(self.quant_bytes(), x, *rows, *cols, out, pool);
1427                    return;
1428                }
1429                if *dtype == TensorDtype::Q1T {
1430                    // GPU route for large q1t matvecs: the ternary BASE dot runs
1431                    // on the GPU (load-port-bound on CPU, like q1), then the
1432                    // sparse overlay is added on the CPU. Probe keeps the winner.
1433                    if *rows * *cols >= 8_388_608 && crate::gpu::enabled_here() {
1434                        let t0 = std::time::Instant::now();
1435                        match crate::gpu::probe_arm(crate::gpu::OpClass::Matvec) {
1436                            crate::gpu::ProbeArm::Gpu => {
1437                                if crate::gpu::q1t_matvec(model, *idx, x, *rows, *cols, out) {
1438                                    q1t_add_overlay(self.quant_bytes(), x, *rows, *cols, out, pool);
1439                                    crate::gpu::probe_record(
1440                                        crate::gpu::OpClass::Matvec,
1441                                        true,
1442                                        t0.elapsed(),
1443                                    );
1444                                    return;
1445                                }
1446                            }
1447                            crate::gpu::ProbeArm::CpuTimed => {
1448                                q1t_matvec(self.quant_bytes(), x, *rows, *cols, out, pool);
1449                                crate::gpu::probe_record(
1450                                    crate::gpu::OpClass::Matvec,
1451                                    false,
1452                                    t0.elapsed(),
1453                                );
1454                                return;
1455                            }
1456                            crate::gpu::ProbeArm::Cpu => {}
1457                        }
1458                    }
1459                    q1t_matvec(self.quant_bytes(), x, *rows, *cols, out, pool);
1460                    return;
1461                }
1462                if matches!(dtype, TensorDtype::Vbit | TensorDtype::VbitRo) {
1463                    vbitmatvec(self.quant_bytes(), vbit_offsets, x, *rows, *cols, out, pool);
1464                    return;
1465                }
1466                let xs = prescale(x, col_field, *dtype);
1467                // D5: large q8 matrices (lm_head-class) — hybrid
1468                // CPU∥GPU: split the rows, both sides compute
1469                // SIMULTANEOUSLY (same math, shared prescale).
1470                // GPU share: CMF_GPU_SPLIT (0..1, default 0.5).
1471                if *rows >= crate::gpu::min_rows()
1472                    && matches!(dtype, TensorDtype::Q8Row | TensorDtype::Q8_2f)
1473                    && gpu_lmhead_enabled()
1474                    && crate::gpu::enabled_here()
1475                {
1476                    // Runtime probe: alternate the hybrid against the
1477                    // pure-CPU matvec, keep whichever is faster HERE.
1478                    let t0 = std::time::Instant::now();
1479                    match crate::gpu::probe_arm(crate::gpu::OpClass::Matvec) {
1480                        crate::gpu::ProbeArm::Gpu => {}
1481                        crate::gpu::ProbeArm::CpuTimed => {
1482                            qmatvec(
1483                                self.quant_bytes(),
1484                                repack,
1485                                row_scale,
1486                                x,
1487                                col_field,
1488                                *dtype,
1489                                *rows,
1490                                *cols,
1491                                out,
1492                                pool,
1493                            );
1494                            crate::gpu::probe_record(
1495                                crate::gpu::OpClass::Matvec,
1496                                false,
1497                                t0.elapsed(),
1498                            );
1499                            return;
1500                        }
1501                        crate::gpu::ProbeArm::Cpu => {
1502                            qmatvec(
1503                                self.quant_bytes(),
1504                                repack,
1505                                row_scale,
1506                                x,
1507                                col_field,
1508                                *dtype,
1509                                *rows,
1510                                *cols,
1511                                out,
1512                                pool,
1513                            );
1514                            return;
1515                        }
1516                    }
1517                    let frac = gpu_split_frac();
1518                    let cpu_rows = ((*rows as f32) * (1.0 - frac)) as usize;
1519                    let (out_cpu, out_gpu) = out.split_at_mut(cpu_rows);
1520                    let bytes = self.quant_bytes();
1521                    let ok = std::thread::scope(|sc| {
1522                        let g = sc.spawn(|| {
1523                            crate::gpu::q8_matvec_range(
1524                                model,
1525                                *idx,
1526                                cpu_rows,
1527                                &row_scale[cpu_rows..],
1528                                &xs,
1529                                *rows - cpu_rows,
1530                                *cols,
1531                                out_gpu,
1532                            )
1533                        });
1534                        if cpu_rows > 0 {
1535                            // Repack prefix covers the full groups of the
1536                            // CPU half (the split starts at row 0).
1537                            let rep_cpu = if repack.is_empty() {
1538                                &[][..]
1539                            } else {
1540                                &repack[..(cpu_rows / 4) * 4 * *cols]
1541                            };
1542                            qmatvec(
1543                                &bytes[..cpu_rows * *cols],
1544                                rep_cpu,
1545                                &row_scale[..cpu_rows],
1546                                x,
1547                                col_field,
1548                                *dtype,
1549                                cpu_rows,
1550                                *cols,
1551                                out_cpu,
1552                                pool,
1553                            );
1554                        }
1555                        g.join().unwrap_or(false)
1556                    });
1557                    if ok {
1558                        crate::gpu::probe_record(crate::gpu::OpClass::Matvec, true, t0.elapsed());
1559                        return;
1560                    }
1561                    // GPU failed — CPU finishes its half (rows rebased —
1562                    // group offsets don't line up, mmap layout only).
1563                    qmatvec(
1564                        &bytes[cpu_rows * *cols..(*rows) * *cols],
1565                        &[],
1566                        &row_scale[cpu_rows..],
1567                        x,
1568                        col_field,
1569                        *dtype,
1570                        *rows - cpu_rows,
1571                        *cols,
1572                        out_gpu,
1573                        pool,
1574                    );
1575                    return;
1576                }
1577                qmatvec(
1578                    self.quant_bytes(),
1579                    repack,
1580                    row_scale,
1581                    x,
1582                    col_field,
1583                    *dtype,
1584                    *rows,
1585                    *cols,
1586                    out,
1587                    pool,
1588                );
1589            }
1590        }
1591    }
1592
1593    /// Fused two-input matvec (MTP verify pair): weights streamed once.
1594    pub fn matvec2(
1595        &self,
1596        x1: &[f32],
1597        x2: &[f32],
1598        o1: &mut [f32],
1599        o2: &mut [f32],
1600        pool: Option<&Pool>,
1601    ) {
1602        match self {
1603            Self::F32 { data, .. } => matvec_rows2(pool, data, x1, x2, o1, o2),
1604            Self::Mapped {
1605                model,
1606                idx,
1607                dtype,
1608                rows,
1609                cols,
1610                row_scale,
1611                col_field,
1612                vbit_offsets,
1613                ..
1614            } => {
1615                if crate::prism::is_forward_weight(model, &model.tensors[*idx].name) {
1616                    let tx1 = crate::prism::forward(model, &x1[..*cols]);
1617                    let tx2 = crate::prism::forward(model, &x2[..*cols]);
1618                    match dtype {
1619                        TensorDtype::Q4Block => {
1620                            q4matvec2(self.quant_bytes(), &tx1, &tx2, *rows, *cols, o1, o2, pool)
1621                        }
1622                        TensorDtype::Q4Tiled => {
1623                            q4t_matvec2(self.quant_bytes(), &tx1, &tx2, *rows, *cols, o1, o2, pool)
1624                        }
1625                        TensorDtype::Q4TiledP => {
1626                            q4tp_matvec2(self.quant_bytes(), &tx1, &tx2, *rows, *cols, o1, o2, pool)
1627                        }
1628                        TensorDtype::Q2TiledP => {
1629                            if crate::prism::is_affine_target(model, &model.tensors[*idx].name) {
1630                                q2tp_affine_matvec2(
1631                                    self.quant_bytes(),
1632                                    &tx1,
1633                                    &tx2,
1634                                    *rows,
1635                                    *cols,
1636                                    o1,
1637                                    o2,
1638                                    pool,
1639                                )
1640                            } else {
1641                                q2tp_matvec2(
1642                                    self.quant_bytes(),
1643                                    &tx1,
1644                                    &tx2,
1645                                    *rows,
1646                                    *cols,
1647                                    o1,
1648                                    o2,
1649                                    pool,
1650                                )
1651                            }
1652                        }
1653                        TensorDtype::Q1 => {
1654                            q1_matvec2(self.quant_bytes(), &tx1, &tx2, *rows, *cols, o1, o2, pool)
1655                        }
1656                        TensorDtype::Q1T => {
1657                            q1t_matvec2(self.quant_bytes(), &tx1, &tx2, *rows, *cols, o1, o2, pool)
1658                        }
1659                        TensorDtype::Vbit | TensorDtype::VbitRo => vbitmatvec2(
1660                            self.quant_bytes(),
1661                            vbit_offsets,
1662                            &tx1,
1663                            &tx2,
1664                            *rows,
1665                            *cols,
1666                            o1,
1667                            o2,
1668                            pool,
1669                        ),
1670                        TensorDtype::Q8Row | TensorDtype::Q8_2f => qmatvec2(
1671                            self.quant_bytes(),
1672                            row_scale,
1673                            &tx1,
1674                            &tx2,
1675                            col_field,
1676                            *dtype,
1677                            *rows,
1678                            *cols,
1679                            o1,
1680                            o2,
1681                            pool,
1682                        ),
1683                        _ => unreachable!("unsupported mapped Prism dtype {dtype:?}"),
1684                    }
1685                    return;
1686                }
1687                if *dtype == TensorDtype::Q4Block {
1688                    q4matvec2(self.quant_bytes(), x1, x2, *rows, *cols, o1, o2, pool);
1689                    return;
1690                }
1691                if *dtype == TensorDtype::Q4Tiled {
1692                    q4t_matvec2(self.quant_bytes(), x1, x2, *rows, *cols, o1, o2, pool);
1693                    return;
1694                }
1695                if *dtype == TensorDtype::Q4TiledP {
1696                    q4tp_matvec2(self.quant_bytes(), x1, x2, *rows, *cols, o1, o2, pool);
1697                    return;
1698                }
1699                if *dtype == TensorDtype::Q2TiledP {
1700                    q2tp_matvec2(self.quant_bytes(), x1, x2, *rows, *cols, o1, o2, pool);
1701                    return;
1702                }
1703                if *dtype == TensorDtype::Q1 {
1704                    q1_matvec2(self.quant_bytes(), x1, x2, *rows, *cols, o1, o2, pool);
1705                    return;
1706                }
1707                if *dtype == TensorDtype::Q1T {
1708                    // Fused ternary pair: one row pass, the register
1709                    // unpack shared across both streams on ARM. (Q1T
1710                    // lacks a row_scale array — scales live inline in
1711                    // the tiles — so it must not fall through to the
1712                    // q8 qmatvec2 below.)
1713                    q1t_matvec2(self.quant_bytes(), x1, x2, *rows, *cols, o1, o2, pool);
1714                    return;
1715                }
1716                if matches!(dtype, TensorDtype::Vbit | TensorDtype::VbitRo) {
1717                    vbitmatvec2(
1718                        self.quant_bytes(),
1719                        vbit_offsets,
1720                        x1,
1721                        x2,
1722                        *rows,
1723                        *cols,
1724                        o1,
1725                        o2,
1726                        pool,
1727                    );
1728                    return;
1729                }
1730                qmatvec2(
1731                    self.quant_bytes(),
1732                    row_scale,
1733                    x1,
1734                    x2,
1735                    col_field,
1736                    *dtype,
1737                    *rows,
1738                    *cols,
1739                    o1,
1740                    o2,
1741                    pool,
1742                );
1743            }
1744        }
1745    }
1746}
1747
1748impl QTensor {
1749    /// Batched matvec (prefill-GEMM): xs — row-major [b, cols],
1750    /// out — row-major [b, rows]. Element-wise semantics are IDENTICAL
1751    /// to b matvec calls (same dot kernels in the same order); the win —
1752    /// the weight row streams from DRAM once per batch, not b times.
1753    /// `(model, index)` when this is a memory-mapped q4tp tensor — the
1754    /// identity a device-resident chain needs to hand `tp_matmat` the
1755    /// weight without going through this struct's own dispatch.
1756    pub fn q4tp_mapped(&self) -> Option<(&std::sync::Arc<CmfModel>, usize)> {
1757        if self.has_prism_contract() {
1758            return None;
1759        }
1760        match self {
1761            Self::Mapped {
1762                model, idx, dtype, ..
1763            } if *dtype == TensorDtype::Q4TiledP => Some((model, *idx)),
1764            _ => None,
1765        }
1766    }
1767
1768    pub fn matmat(&self, xs_all: &[f32], b: usize, out: &mut [f32], pool: Option<&Pool>) {
1769        let cols = self.cols();
1770        let rows = self.rows();
1771        debug_assert_eq!(xs_all.len(), b * cols);
1772        debug_assert_eq!(out.len(), b * rows);
1773        let _prof = crate::cpuprof::time(crate::cpuprof::Slot::Matmat);
1774        // GPTQ calibration: fold this layer's inputs into its Hessian. Only
1775        // Mapped tensors carry a directory name; the check is a relaxed
1776        // atomic load, free when not calibrating.
1777        if crate::gptq_capture::capturing() {
1778            if let Self::Mapped { model, idx, .. } = self {
1779                crate::gptq_capture::accumulate(&model.tensors[*idx].name, xs_all, b, cols);
1780            }
1781        }
1782        match self {
1783            Self::F32 { data, .. } => {
1784                let out_addr = SendMut(out.as_mut_ptr());
1785                let run = |start: usize, end: usize| {
1786                    for o in start..end {
1787                        let row = &data[o * cols..(o + 1) * cols];
1788                        for bi in 0..b {
1789                            let x = &xs_all[bi * cols..(bi + 1) * cols];
1790                            let mut acc = 0f32;
1791                            for j in 0..cols {
1792                                acc += row[j] * x[j];
1793                            }
1794                            unsafe { *out_addr.at(bi * rows + o) = acc };
1795                        }
1796                    }
1797                };
1798                dispatch_rows(pool, rows, &run);
1799            }
1800            Self::Mapped {
1801                model,
1802                idx,
1803                dtype,
1804                row_scale,
1805                col_field,
1806                vbit_offsets,
1807                ..
1808            } => {
1809                if crate::prism::is_forward_weight(model, &model.tensors[*idx].name) {
1810                    let mut transformed = Vec::with_capacity(xs_all.len());
1811                    for bi in 0..b {
1812                        transformed.extend_from_slice(&crate::prism::forward(
1813                            model,
1814                            &xs_all[bi * cols..(bi + 1) * cols],
1815                        ));
1816                    }
1817                    match dtype {
1818                        TensorDtype::Q4Block => {
1819                            q4matmat(self.quant_bytes(), &transformed, b, rows, cols, out, pool)
1820                        }
1821                        TensorDtype::Q4Tiled => {
1822                            q4t_matmat(self.quant_bytes(), &transformed, b, rows, cols, out, pool)
1823                        }
1824                        TensorDtype::Q4TiledP => {
1825                            q4tp_matmat(self.quant_bytes(), &transformed, b, rows, cols, out, pool)
1826                        }
1827                        TensorDtype::Q2TiledP => {
1828                            let affine =
1829                                crate::prism::is_affine_target(model, &model.tensors[*idx].name);
1830                            // Affine Prism Q2TP has a descriptor-aware GPU
1831                            // kernel for short/tail batches too.  Unlike the
1832                            // ordinary Q2TP path, don't force b<32 back to a
1833                            // scalar CPU matmat: prefill chunks and the final
1834                            // tail both need to stay on the tested GPU arm.
1835                            let gpu_batch_ok = if affine {
1836                                b >= 2
1837                            } else {
1838                                b >= 32 && b * rows * cols >= 128_000_000
1839                            };
1840                            if gpu_batch_ok
1841                                && cols % 32 == 0
1842                                && crate::gpu::enabled_here()
1843                                && crate::gpu::q2tp_gpu_opt_in()
1844                            {
1845                                let gpu_ok = if affine {
1846                                    crate::gpu::q2tp_affine_matmat(
1847                                        model,
1848                                        *idx,
1849                                        &transformed,
1850                                        b,
1851                                        rows,
1852                                        cols,
1853                                        out,
1854                                    )
1855                                } else {
1856                                    crate::gpu::q2tp_matmat(
1857                                        model,
1858                                        *idx,
1859                                        &transformed,
1860                                        b,
1861                                        rows,
1862                                        cols,
1863                                        out,
1864                                    )
1865                                };
1866                                if gpu_ok {
1867                                    return;
1868                                }
1869                            }
1870                            // A one-token Prism decode is the other short
1871                            // case.  Use the descriptor-aware matvec kernel
1872                            // before falling back to the exact CPU path.
1873                            if affine
1874                                && b == 1
1875                                && cols % 32 == 0
1876                                && crate::gpu::enabled_here()
1877                                && crate::gpu::q2tp_gpu_opt_in()
1878                                && crate::gpu::q2tp_affine_matvec(
1879                                    model,
1880                                    *idx,
1881                                    &transformed[..cols],
1882                                    rows,
1883                                    cols,
1884                                    &mut out[..rows],
1885                                )
1886                            {
1887                                return;
1888                            }
1889                            if affine {
1890                                q2tp_affine_matmat(
1891                                    self.quant_bytes(),
1892                                    &transformed,
1893                                    b,
1894                                    rows,
1895                                    cols,
1896                                    out,
1897                                    pool,
1898                                )
1899                            } else {
1900                                q2tp_matmat(
1901                                    self.quant_bytes(),
1902                                    &transformed,
1903                                    b,
1904                                    rows,
1905                                    cols,
1906                                    out,
1907                                    pool,
1908                                )
1909                            }
1910                        }
1911                        TensorDtype::Q1 => {
1912                            q1_matmat(self.quant_bytes(), &transformed, b, rows, cols, out, pool)
1913                        }
1914                        TensorDtype::Q1T => {
1915                            q1t_matmat(self.quant_bytes(), &transformed, b, rows, cols, out, pool)
1916                        }
1917                        TensorDtype::Vbit | TensorDtype::VbitRo => vbitmatmat(
1918                            self.quant_bytes(),
1919                            vbit_offsets,
1920                            &transformed,
1921                            b,
1922                            rows,
1923                            cols,
1924                            out,
1925                            pool,
1926                        ),
1927                        TensorDtype::Q8Row | TensorDtype::Q8_2f => {
1928                            let pre: Vec<std::borrow::Cow<'_, [f32]>> = (0..b)
1929                                .map(|bi| {
1930                                    prescale(
1931                                        &transformed[bi * cols..(bi + 1) * cols],
1932                                        col_field,
1933                                        *dtype,
1934                                    )
1935                                })
1936                                .collect();
1937                            qmatmat(self.quant_bytes(), row_scale, &pre, rows, cols, out, pool)
1938                        }
1939                        _ => unreachable!("unsupported mapped Prism dtype {dtype:?}"),
1940                    }
1941                    return;
1942                }
1943                if *dtype == TensorDtype::Q4Block {
1944                    q4matmat(self.quant_bytes(), xs_all, b, rows, cols, out, pool);
1945                    return;
1946                }
1947                if *dtype == TensorDtype::Q4TiledP {
1948                    // GPU batched q4tp GEMM (dequant + f32nt mul_mm on the
1949                    // device); the probe keeps whichever beats the CPU arm.
1950                    // Narrow (prompt-encode) and wide (DiT) batches probe
1951                    // as separate classes — the regimes have opposite
1952                    // winners and one shared verdict locked the wrong arm.
1953                    // Kill switch (gpu::mm_kill): one grossly slow GPU op
1954                    // (a fair-condition op is ≤~100 ms even at 1024px)
1955                    // means the device is contended by another process
1956                    // (e.g. a simulator) — verdicts are per-process, so
1957                    // without the bail the whole render crawls behind
1958                    // someone else's queue.
1959                    // Row-exact batches stay on the host: the device GEMM
1960                    // is an f32 dequant-sgemm, not the host matvec's sum.
1961                    if b >= 32
1962                        && b * rows * cols >= 128_000_000
1963                        && cols % 32 == 0
1964                        && !row_exact()
1965                        && !crate::gpu::mm_killed()
1966                        && crate::gpu::enabled_here()
1967                    {
1968                        let class = if b >= 128 {
1969                            crate::gpu::OpClass::MatmatWide
1970                        } else {
1971                            crate::gpu::OpClass::Matmat
1972                        };
1973                        if let Self::Mapped { model, idx, .. } = self {
1974                            // In-process A/B (`CMF_MM_AB=1`). Three
1975                            // wall-clock A/Bs on a shared stand disagreed
1976                            // with each other by 25% on the same change,
1977                            // because the machine drifts between processes
1978                            // and interleaving whole renders does not fix
1979                            // that. Here both arms run back to back on the
1980                            // SAME data inside one call, so whatever the
1981                            // machine is doing, it does to both — and the
1982                            // disagreement between their outputs falls out
1983                            // for free. Doubles the work; a diagnostic,
1984                            // not a mode.
1985                            if crate::mm_ab::on() {
1986                                let mut g = vec![0f32; b * rows];
1987                                let t = std::time::Instant::now();
1988                                let took = crate::gpu::q4tp_matmat(
1989                                    model, *idx, xs_all, b, rows, cols, &mut g,
1990                                );
1991                                let dg = t.elapsed();
1992                                let t = std::time::Instant::now();
1993                                q4tp_matmat(self.quant_bytes(), xs_all, b, rows, cols, out, pool);
1994                                let dc = t.elapsed();
1995                                crate::mm_ab::record(b, rows, cols, took, dg, dc, &g, out);
1996                                return;
1997                            }
1998                            let t0 = std::time::Instant::now();
1999                            // A cold call takes the device arm: its sample
2000                            // is discarded either way, and the upload is
2001                            // what the next step needs.
2002                            let resident = crate::gpu::weight_is_resident(model, *idx);
2003                            match crate::gpu::probe_arm_cold_prefers_gpu(class, resident) {
2004                                crate::gpu::ProbeArm::Gpu => {
2005                                    if crate::gpu::q4tp_matmat(
2006                                        model, *idx, xs_all, b, rows, cols, out,
2007                                    ) {
2008                                        let el = t0.elapsed();
2009                                        // Work-proportional budget: ~8× the
2010                                        // fair-device estimate (+20 ms slack).
2011                                        // An absolute cap missed the worst
2012                                        // case — contended ops sit at
2013                                        // 100–240 ms each and still bury a
2014                                        // render whose fair op is 3–9 ms.
2015                                        // Cold ops (first PSO build, buffer
2016                                        // alloc) are exempt: a one-off
2017                                        // ~50 ms compile is not contention.
2018                                        let flops = 2.0 * b as f64 * rows as f64 * cols as f64;
2019                                        let budget = std::time::Duration::from_secs_f64(
2020                                            flops / 1.5e12 * 8.0 + 0.020,
2021                                        );
2022                                        crate::gpu::mm_budget_check(
2023                                            "q4tp matmat",
2024                                            el,
2025                                            budget,
2026                                            crate::gpu::probe_was_cold() || !resident,
2027                                        );
2028                                        crate::gpu::probe_record(class, true, el);
2029                                        return;
2030                                    }
2031                                }
2032                                crate::gpu::ProbeArm::CpuTimed => {
2033                                    q4tp_matmat(
2034                                        self.quant_bytes(),
2035                                        xs_all,
2036                                        b,
2037                                        rows,
2038                                        cols,
2039                                        out,
2040                                        pool,
2041                                    );
2042                                    crate::gpu::probe_record(class, false, t0.elapsed());
2043                                    return;
2044                                }
2045                                crate::gpu::ProbeArm::Cpu => {}
2046                            }
2047                        }
2048                    }
2049                    q4tp_matmat(self.quant_bytes(), xs_all, b, rows, cols, out, pool);
2050                    return;
2051                }
2052                if *dtype == TensorDtype::Q2TiledP {
2053                    // Same device arm as q4tp, behind the same probe:
2054                    // the planes differ, the dispatch does not. Without
2055                    // this a q2tp file ran its widest projections on the
2056                    // host while the 4-bit one had the card, which is a
2057                    // codec paying for its size twice.
2058                    if b >= 32
2059                        && b * rows * cols >= 128_000_000
2060                        && cols % 32 == 0
2061                        && !crate::gpu::mm_killed()
2062                        && crate::gpu::enabled_here()
2063                    {
2064                        let class = if b >= 128 {
2065                            crate::gpu::OpClass::MatmatWide
2066                        } else {
2067                            crate::gpu::OpClass::Matmat
2068                        };
2069                        if let Self::Mapped { model, idx, .. } = self {
2070                            let t0 = std::time::Instant::now();
2071                            match crate::gpu::probe_arm(class) {
2072                                crate::gpu::ProbeArm::Gpu => {
2073                                    if crate::gpu::q2tp_matmat(
2074                                        model, *idx, xs_all, b, rows, cols, out,
2075                                    ) {
2076                                        crate::gpu::probe_record(class, true, t0.elapsed());
2077                                        return;
2078                                    }
2079                                }
2080                                crate::gpu::ProbeArm::CpuTimed => {
2081                                    q2tp_matmat(
2082                                        self.quant_bytes(),
2083                                        xs_all,
2084                                        b,
2085                                        rows,
2086                                        cols,
2087                                        out,
2088                                        pool,
2089                                    );
2090                                    crate::gpu::probe_record(class, false, t0.elapsed());
2091                                    return;
2092                                }
2093                                crate::gpu::ProbeArm::Cpu => {}
2094                            }
2095                        }
2096                    }
2097                    // Without a host arm a q2tp tensor falls through to
2098                    // the q8 fallback, which reads it at one BYTE per
2099                    // weight — a 2x overrun that killed pool workers
2100                    // mid-prefill while the dispatcher waited forever.
2101                    q2tp_matmat(self.quant_bytes(), xs_all, b, rows, cols, out, pool);
2102                    return;
2103                }
2104                if *dtype == TensorDtype::Q4Tiled {
2105                    // GPU batched q4t GEMM (dequant + f32nt mul_mm on the
2106                    // device); the probe keeps whichever beats the CPU arm.
2107                    // Narrow (prompt-encode) and wide (DiT) batches probe
2108                    // as separate classes — the regimes have opposite
2109                    // winners and one shared verdict locked the wrong arm.
2110                    // Kill switch (gpu::mm_kill): one grossly slow GPU op
2111                    // (a fair-condition op is ≤~100 ms even at 1024px)
2112                    // means the device is contended by another process
2113                    // (e.g. a simulator) — verdicts are per-process, so
2114                    // without the bail the whole render crawls behind
2115                    // someone else's queue.
2116                    if b >= 32
2117                        && b * rows * cols >= 128_000_000
2118                        && cols % 32 == 0
2119                        && !crate::gpu::mm_killed()
2120                        && crate::gpu::enabled_here()
2121                    {
2122                        let class = if b >= 128 {
2123                            crate::gpu::OpClass::MatmatWide
2124                        } else {
2125                            crate::gpu::OpClass::Matmat
2126                        };
2127                        if let Self::Mapped { model, idx, .. } = self {
2128                            let t0 = std::time::Instant::now();
2129                            match crate::gpu::probe_arm(class) {
2130                                crate::gpu::ProbeArm::Gpu => {
2131                                    if crate::gpu::q4t_matmat(
2132                                        model, *idx, xs_all, b, rows, cols, out,
2133                                    ) {
2134                                        let el = t0.elapsed();
2135                                        // Work-proportional budget: ~8× the
2136                                        // fair-device estimate (+20 ms slack).
2137                                        // An absolute cap missed the worst
2138                                        // case — contended ops sit at
2139                                        // 100–240 ms each and still bury a
2140                                        // render whose fair op is 3–9 ms.
2141                                        // Cold ops (first PSO build, buffer
2142                                        // alloc) are exempt: a one-off
2143                                        // ~50 ms compile is not contention.
2144                                        let flops = 2.0 * b as f64 * rows as f64 * cols as f64;
2145                                        let budget = std::time::Duration::from_secs_f64(
2146                                            flops / 1.5e12 * 8.0 + 0.020,
2147                                        );
2148                                        crate::gpu::mm_budget_check(
2149                                            "q4t matmat",
2150                                            el,
2151                                            budget,
2152                                            crate::gpu::probe_was_cold(),
2153                                        );
2154                                        crate::gpu::probe_record(class, true, el);
2155                                        return;
2156                                    }
2157                                }
2158                                crate::gpu::ProbeArm::CpuTimed => {
2159                                    q4t_matmat(
2160                                        self.quant_bytes(),
2161                                        xs_all,
2162                                        b,
2163                                        rows,
2164                                        cols,
2165                                        out,
2166                                        pool,
2167                                    );
2168                                    crate::gpu::probe_record(class, false, t0.elapsed());
2169                                    return;
2170                                }
2171                                crate::gpu::ProbeArm::Cpu => {}
2172                            }
2173                        }
2174                    }
2175                    q4t_matmat(self.quant_bytes(), xs_all, b, rows, cols, out, pool);
2176                    return;
2177                }
2178                if *dtype == TensorDtype::Q1 {
2179                    // GPU batched q1 GEMM for wide prefill (q1_mul_mm on the
2180                    // device); the probe keeps whichever beats the CPU matmat.
2181                    if b >= 32
2182                        && b * rows * cols >= 128_000_000
2183                        && cols % 64 == 0
2184                        && crate::gpu::enabled_here()
2185                    {
2186                        if let Self::Mapped { model, idx, .. } = self {
2187                            let t0 = std::time::Instant::now();
2188                            match crate::gpu::probe_arm(crate::gpu::OpClass::Matmat) {
2189                                crate::gpu::ProbeArm::Gpu => {
2190                                    if crate::gpu::q1_matmat(
2191                                        model, *idx, xs_all, b, rows, cols, out,
2192                                    ) {
2193                                        crate::gpu::probe_record(
2194                                            crate::gpu::OpClass::Matmat,
2195                                            true,
2196                                            t0.elapsed(),
2197                                        );
2198                                        return;
2199                                    }
2200                                }
2201                                crate::gpu::ProbeArm::CpuTimed => {
2202                                    q1_matmat(self.quant_bytes(), xs_all, b, rows, cols, out, pool);
2203                                    crate::gpu::probe_record(
2204                                        crate::gpu::OpClass::Matmat,
2205                                        false,
2206                                        t0.elapsed(),
2207                                    );
2208                                    return;
2209                                }
2210                                crate::gpu::ProbeArm::Cpu => {}
2211                            }
2212                        }
2213                    }
2214                    q1_matmat(self.quant_bytes(), xs_all, b, rows, cols, out, pool);
2215                    return;
2216                }
2217                if *dtype == TensorDtype::Q1T {
2218                    // GPU batched GEMM for wide prefill (base + overlay on the
2219                    // device); probe keeps the winner vs the CPU matmat.
2220                    if b >= 32 && b * rows * cols >= 128_000_000 && crate::gpu::enabled_here() {
2221                        if let Self::Mapped { model, idx, .. } = self {
2222                            let t0 = std::time::Instant::now();
2223                            match crate::gpu::probe_arm(crate::gpu::OpClass::Matmat) {
2224                                crate::gpu::ProbeArm::Gpu => {
2225                                    if crate::gpu::q1t_matmat(
2226                                        model, *idx, xs_all, b, rows, cols, out,
2227                                    ) {
2228                                        crate::gpu::probe_record(
2229                                            crate::gpu::OpClass::Matmat,
2230                                            true,
2231                                            t0.elapsed(),
2232                                        );
2233                                        return;
2234                                    }
2235                                }
2236                                crate::gpu::ProbeArm::CpuTimed => {
2237                                    q1t_matmat(
2238                                        self.quant_bytes(),
2239                                        xs_all,
2240                                        b,
2241                                        rows,
2242                                        cols,
2243                                        out,
2244                                        pool,
2245                                    );
2246                                    crate::gpu::probe_record(
2247                                        crate::gpu::OpClass::Matmat,
2248                                        false,
2249                                        t0.elapsed(),
2250                                    );
2251                                    return;
2252                                }
2253                                crate::gpu::ProbeArm::Cpu => {}
2254                            }
2255                        }
2256                    }
2257                    q1t_matmat(self.quant_bytes(), xs_all, b, rows, cols, out, pool);
2258                    return;
2259                }
2260                if matches!(dtype, TensorDtype::Vbit | TensorDtype::VbitRo) {
2261                    vbitmatmat(
2262                        self.quant_bytes(),
2263                        vbit_offsets,
2264                        xs_all,
2265                        b,
2266                        rows,
2267                        cols,
2268                        out,
2269                        pool,
2270                    );
2271                    return;
2272                }
2273                let pre: Vec<std::borrow::Cow<'_, [f32]>> = (0..b)
2274                    .map(|bi| prescale(&xs_all[bi * cols..(bi + 1) * cols], col_field, *dtype))
2275                    .collect();
2276                // MiMo verification is a 2–4 row decode panel, not a wide
2277                // prompt GEMM. Keep q8 projections on the same device as
2278                // decode; the generic b>=8 gate otherwise silently moves
2279                // every projection back to CPU. The short wgpu matmat uses
2280                // the same 64-lane reduction as its single-token matvec.
2281                if row_exact()
2282                    && (1..=4).contains(&b)
2283                    && matches!(dtype, TensorDtype::Q8Row | TensorDtype::Q8_2f)
2284                    && crate::gpu::enabled_here()
2285                    && crate::gpu::wgpu_active()
2286                {
2287                    let flat: Vec<f32> = pre.iter().flat_map(|v| v.iter().copied()).collect();
2288                    if crate::gpu::q8_matmat(model, *idx, row_scale, &flat, b, rows, cols, out) {
2289                        return;
2290                    }
2291                }
2292                // D5: large prefill-batch GEMMs — on the GPU (threshold by
2293                // work volume: submission carries b×rows×cols MACs).
2294                // Runtime probe: the naive GEMM shader + sync readback
2295                // lose to the CPU GEMM on slow driver stacks — alternate
2296                // both arms and keep the winner.
2297                if b >= 8 && b * rows * cols >= 128_000_000 && crate::gpu::enabled_here() {
2298                    if let Self::Mapped { model, idx, .. } = self {
2299                        let t0 = std::time::Instant::now();
2300                        match crate::gpu::probe_arm(crate::gpu::OpClass::Matmat) {
2301                            crate::gpu::ProbeArm::Gpu
2302                                if crate::gpu::probe_deciding(crate::gpu::OpClass::Matmat)
2303                                    && !crate::gpu::q8_resident_or_upload(model, *idx) =>
2304                            {
2305                                // Cold weights during probing: the upload
2306                                // has started, the count runs on the CPU —
2307                                // the GPU arm samples on the next touch.
2308                                let q = self.quant_bytes();
2309                                qmatmat(q, row_scale, &pre, rows, cols, out, pool);
2310                                return;
2311                            }
2312                            crate::gpu::ProbeArm::Gpu => {
2313                                let flat: Vec<f32> =
2314                                    pre.iter().flat_map(|v| v.iter().copied()).collect();
2315                                if crate::gpu::q8_matmat(
2316                                    model, *idx, row_scale, &flat, b, rows, cols, out,
2317                                ) {
2318                                    crate::gpu::probe_record(
2319                                        crate::gpu::OpClass::Matmat,
2320                                        true,
2321                                        t0.elapsed(),
2322                                    );
2323                                    return;
2324                                }
2325                            }
2326                            crate::gpu::ProbeArm::CpuTimed => {
2327                                let q = self.quant_bytes();
2328                                qmatmat(q, row_scale, &pre, rows, cols, out, pool);
2329                                crate::gpu::probe_record(
2330                                    crate::gpu::OpClass::Matmat,
2331                                    false,
2332                                    t0.elapsed(),
2333                                );
2334                                return;
2335                            }
2336                            crate::gpu::ProbeArm::Cpu => {}
2337                        }
2338                    }
2339                }
2340                let q = self.quant_bytes();
2341                qmatmat(q, row_scale, &pre, rows, cols, out, pool);
2342            }
2343        }
2344    }
2345}
2346
2347impl QTensor {
2348    /// The device GEMM this tensor would take, run once on the caller's
2349    /// data — the startup parity probe's arm, and the one place that knows
2350    /// which entry point each codec has.
2351    ///
2352    /// It exists because the probe used to look for a `q4tp` weight by
2353    /// name AND dtype, and a container packed any other way was declared
2354    /// "host path" for the whole render even though its codec had a device
2355    /// GEMM of its own. A gate that only recognizes one codec is a gate
2356    /// that silently downgrades every other one.
2357    pub fn device_matmat(&self, xs: &[f32], b: usize, out: &mut [f32]) -> bool {
2358        let (rows, cols) = (self.rows(), self.cols());
2359        let Self::Mapped {
2360            model,
2361            idx,
2362            dtype,
2363            row_scale,
2364            col_field,
2365            ..
2366        } = self
2367        else {
2368            return false;
2369        };
2370        if crate::prism::has_contract(model) {
2371            return false;
2372        }
2373        match *dtype {
2374            TensorDtype::Q4TiledP => crate::gpu::q4tp_matmat(model, *idx, xs, b, rows, cols, out),
2375            // The two-field codec folds its column field into the
2376            // activation, which leaves a plain per-row int8 GEMM — the
2377            // same kernel `q8_row` uses, on both backends.
2378            TensorDtype::Q8Row | TensorDtype::Q8_2f => {
2379                // The field belongs to the weight; only a backend that cannot
2380                // apply it there makes a scaled copy of the activation.
2381                if *dtype == TensorDtype::Q8_2f
2382                    && std::env::var("CMF_Q8_2F_DEV").as_deref() != Ok("0")
2383                    && crate::gpu::q8_matmat_2f(
2384                        model, *idx, row_scale, col_field, xs, b, rows, cols, out,
2385                    )
2386                {
2387                    return true;
2388                }
2389                let flat: Vec<f32> = (0..b)
2390                    .flat_map(|bi| {
2391                        prescale(&xs[bi * cols..(bi + 1) * cols], col_field, *dtype).into_owned()
2392                    })
2393                    .collect();
2394                crate::gpu::q8_matmat(model, *idx, row_scale, &flat, b, rows, cols, out)
2395            }
2396            _ => false,
2397        }
2398    }
2399
2400    /// Multi-matrix job (roadmap §3 P0): N tensors sharing one input
2401    /// run under a SINGLE pool dispatch — QKV or gate+up cost one
2402    /// barrier instead of N. Per-row math is the exact same kernel as
2403    /// `matvec` (bit-identical outputs); only the dispatch is fused.
2404    /// Falls back to N sequential matvecs when the set is not a uniform
2405    /// q8-family/F32 group or there is no pool.
2406    pub fn matvec_many<const N: usize>(
2407        ts: [&QTensor; N],
2408        x: &[f32],
2409        mut outs: [&mut [f32]; N],
2410        pool: Option<&Pool>,
2411    ) {
2412        let total_rows: usize = ts.iter().map(|t| t.rows()).sum();
2413        if ts.iter().any(|t| t.has_prism_contract()) {
2414            // The fused range kernels have no transform descriptor.  Let
2415            // each tensor's ordinary matvec dispatch perform the explicit
2416            // signed FWHT (and retain CPU fallback for mixed q2tp/q4tp).
2417            for (t, o) in ts.iter().zip(outs.iter_mut()) {
2418                t.matvec(x, o, pool);
2419            }
2420            return;
2421        }
2422        let uniform_q8 = ts.iter().all(|t| {
2423            matches!(
2424                t,
2425                Self::Mapped {
2426                    dtype: TensorDtype::Q8Row | TensorDtype::Q8_2f,
2427                    ..
2428                }
2429            )
2430        });
2431        let uniform_f32 = ts.iter().all(|t| matches!(t, Self::F32 { .. }));
2432        let uniform_q4 = ts.iter().all(|t| {
2433            matches!(
2434                t,
2435                Self::Mapped {
2436                    dtype: TensorDtype::Q4Block,
2437                    ..
2438                }
2439            )
2440        });
2441        let uniform_vbit = ts.iter().all(|t| {
2442            matches!(
2443                t,
2444                Self::Mapped {
2445                    dtype: TensorDtype::Vbit | TensorDtype::VbitRo,
2446                    ..
2447                }
2448            )
2449        });
2450        let uniform_q1 = ts.iter().all(|t| {
2451            matches!(
2452                t,
2453                Self::Mapped {
2454                    dtype: TensorDtype::Q1,
2455                    ..
2456                }
2457            )
2458        });
2459        let uniform_q1t = ts.iter().all(|t| {
2460            matches!(
2461                t,
2462                Self::Mapped {
2463                    dtype: TensorDtype::Q1T,
2464                    ..
2465                }
2466            )
2467        });
2468        // q4tp is the skeleton dtype of the big MoE files, and without an arm
2469        // here every projection that shares an input paid its own pool
2470        // barrier: DeepSeek-V4's attention step alone hands this function
2471        // wq_a, wkv and both compressors' pairs off the same hidden state.
2472        let uniform_q4tp = ts.iter().all(|t| {
2473            matches!(
2474                t,
2475                Self::Mapped {
2476                    dtype: TensorDtype::Q4TiledP,
2477                    ..
2478                }
2479            )
2480        }) && ts
2481            .iter()
2482            .all(|t| t.cols() == ts[0].cols() && t.cols() % GROUP_SIZE == 0);
2483        let Some(pool) = pool else {
2484            for (t, o) in ts.iter().zip(outs.iter_mut()) {
2485                t.matvec(x, o, None);
2486            }
2487            return;
2488        };
2489        if total_rows < 256
2490            || !(uniform_q8
2491                || uniform_f32
2492                || uniform_q4
2493                || uniform_vbit
2494                || uniform_q1
2495                || uniform_q1t
2496                || uniform_q4tp)
2497        {
2498            for (t, o) in ts.iter().zip(outs.iter_mut()) {
2499                t.matvec(x, o, Some(pool));
2500            }
2501            return;
2502        }
2503
2504        if uniform_q4tp {
2505            // Every tensor's rows laid end to end in one virtual row space,
2506            // so the whole set is ONE dispatch. The per-row body is the
2507            // `q4tp_matvec` arm verbatim — same activation split, same
2508            // accumulation order — so the outputs are bit-identical to the
2509            // sequential calls this replaces.
2510            let cols = ts[0].cols();
2511            let gpr = cols / GROUP_SIZE;
2512            let views: Vec<Q4tpView> = ts
2513                .iter()
2514                .map(|t| Q4tpView::new(t.quant_bytes(), t.rows(), cols))
2515                .collect();
2516            let rows_of: Vec<usize> = ts.iter().map(|t| t.rows()).collect();
2517            let outs_addr: [SendMut; N] = std::array::from_fn(|i| SendMut(outs[i].as_mut_ptr()));
2518            // flat index -> (which tensor, which of its rows)
2519            let locate = |flat: usize| -> (usize, usize) {
2520                let mut acc = 0;
2521                for (i, &r) in rows_of.iter().enumerate() {
2522                    if flat < acc + r {
2523                        return (i, flat - acc);
2524                    }
2525                    acc += r;
2526                }
2527                (rows_of.len() - 1, 0)
2528            };
2529            let (views, outs_addr) = (&views, &outs_addr);
2530            if a8w8_enabled() {
2531                let act = split_act(x);
2532                let act = &act;
2533                let run = |start: usize, end: usize| {
2534                    let mut sc = vec![0f32; gpr];
2535                    for flat in start..end {
2536                        let (t, r) = locate(flat);
2537                        let v = &views[t];
2538                        v.scales_into(r, gpr, &mut sc);
2539                        let mut acc = dot_q4tp_row_i8(v.nib, r, gpr, &act.xq, &sc) * act.sx;
2540                        for &(j, xv) in &act.outliers {
2541                            let (w, s) = q4tp_outlier(v.nib, r, gpr, j, &sc);
2542                            acc += w * s * xv;
2543                        }
2544                        // SAFETY: one worker owns each (tensor, row) pair.
2545                        unsafe { *outs_addr[t].at(r) = acc };
2546                    }
2547                };
2548                pool.run_rows(total_rows, &run);
2549            } else {
2550                let run = |start: usize, end: usize| {
2551                    let mut sc = vec![0f32; gpr];
2552                    for flat in start..end {
2553                        let (t, r) = locate(flat);
2554                        let v = &views[t];
2555                        v.scales_into(r, gpr, &mut sc);
2556                        // SAFETY: one worker owns each (tensor, row) pair.
2557                        unsafe { *outs_addr[t].at(r) = q4tp_row_exact(v.nib, r, gpr, x, &sc) };
2558                    }
2559                };
2560                pool.run_rows(total_rows, &run);
2561            }
2562            return;
2563        }
2564
2565        if uniform_q1 {
2566            // One shared activation split + group sums (q1 has no col
2567            // field; the same input feeds every tensor).
2568            let outs_addr: [SendMut; N] = std::array::from_fn(|i| SendMut(outs[i].as_mut_ptr()));
2569            if a8w8_enabled() {
2570                let act = split_act(x);
2571                let gsum = q1_group_sums(&act.xq, ts[0].cols() / GROUP_SIZE);
2572                let (act, gsum) = (&act, &gsum);
2573                let closures: [_; N] = std::array::from_fn(|i| {
2574                    let (bytes, gpr, out) =
2575                        (ts[i].quant_bytes(), ts[i].cols() / GROUP_SIZE, outs_addr[i]);
2576                    move |s: usize, e: usize| q1_range_a8w8(bytes, gpr, act, gsum, out, s, e)
2577                });
2578                let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2579                    std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2580                pool.run_many(&parts);
2581            } else {
2582                let closures: [_; N] = std::array::from_fn(|i| {
2583                    let (bytes, gpr, out) =
2584                        (ts[i].quant_bytes(), ts[i].cols() / GROUP_SIZE, outs_addr[i]);
2585                    move |s: usize, e: usize| q1_range_f32(bytes, gpr, x, out, s, e)
2586                });
2587                let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2588                    std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2589                pool.run_many(&parts);
2590            }
2591            return;
2592        }
2593
2594        if uniform_q1t {
2595            // Q1T batched: one shared activation split + overlay decode,
2596            // all tensors' rows in ONE pool dispatch (saves N−1 dispatches
2597            // and N−1 redundant split_act calls per layer).
2598            let outs_addr: [SendMut; N] = std::array::from_fn(|i| SendMut(outs[i].as_mut_ptr()));
2599            const TILE: usize = cortiq_core::quant::Q1T_TILE;
2600            if a8w8_enabled() {
2601                let act = split_act(x);
2602                let act = &act;
2603                let x_ref = x;
2604                let closures: [_; N] = std::array::from_fn(|i| {
2605                    let bytes = ts[i].quant_bytes();
2606                    let (rows, cols) = (ts[i].rows(), ts[i].cols());
2607                    let gpr = cols / GROUP_SIZE;
2608                    let (rp_off, ent_off, has_ov) = q1t_overlay(bytes, rows * gpr * TILE, rows);
2609                    let out = outs_addr[i];
2610                    move |s: usize, e: usize| {
2611                        q1t_range_a8w8(bytes, gpr, rp_off, ent_off, has_ov, act, x_ref, out, s, e)
2612                    }
2613                });
2614                let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2615                    std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2616                pool.run_many(&parts);
2617            } else {
2618                let x_ref = x;
2619                let closures: [_; N] = std::array::from_fn(|i| {
2620                    let bytes = ts[i].quant_bytes();
2621                    let (rows, cols) = (ts[i].rows(), ts[i].cols());
2622                    let gpr = cols / GROUP_SIZE;
2623                    let (rp_off, ent_off, has_ov) = q1t_overlay(bytes, rows * gpr * TILE, rows);
2624                    let out = outs_addr[i];
2625                    move |s: usize, e: usize| {
2626                        q1t_range_f32_batch(bytes, gpr, rp_off, ent_off, has_ov, x_ref, out, s, e)
2627                    }
2628                });
2629                let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2630                    std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2631                pool.run_many(&parts);
2632            }
2633            return;
2634        }
2635
2636        if uniform_q4 || uniform_vbit {
2637            let outs_addr: [SendMut; N] = std::array::from_fn(|i| SendMut(outs[i].as_mut_ptr()));
2638            // q4/vbit share one activation split — no per-tensor col field.
2639            if a8w8_enabled() {
2640                let act = split_act(x);
2641                let act = &act;
2642                if uniform_q4 {
2643                    let closures: [_; N] = std::array::from_fn(|i| {
2644                        let (packed, scales) =
2645                            q4_split(ts[i].quant_bytes(), ts[i].rows(), ts[i].cols());
2646                        let (gpr, cols, out) =
2647                            (ts[i].cols() / GROUP_SIZE, ts[i].cols(), outs_addr[i]);
2648                        move |s: usize, e: usize| {
2649                            q4_range_a8w8(packed, scales, gpr, cols, act, out, s, e)
2650                        }
2651                    });
2652                    let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2653                        std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2654                    pool.run_many(&parts);
2655                } else {
2656                    let closures: [_; N] = std::array::from_fn(|i| {
2657                        let Self::Mapped { vbit_offsets, .. } = ts[i] else {
2658                            unreachable!()
2659                        };
2660                        let (bytes, rows, cols, out) = (
2661                            ts[i].quant_bytes(),
2662                            ts[i].rows(),
2663                            ts[i].cols(),
2664                            outs_addr[i],
2665                        );
2666                        move |s: usize, e: usize| {
2667                            vbit_range_a8w8(bytes, vbit_offsets, x, act, rows, cols, out, s, e)
2668                        }
2669                    });
2670                    let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2671                        std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2672                    pool.run_many(&parts);
2673                }
2674                return;
2675            }
2676            if uniform_q4 {
2677                let closures: [_; N] = std::array::from_fn(|i| {
2678                    let (packed, scales) =
2679                        q4_split(ts[i].quant_bytes(), ts[i].rows(), ts[i].cols());
2680                    let (gpr, out) = (ts[i].cols() / GROUP_SIZE, outs_addr[i]);
2681                    move |s: usize, e: usize| q4_range_f32(packed, scales, gpr, x, out, s, e)
2682                });
2683                let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2684                    std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2685                pool.run_many(&parts);
2686            } else {
2687                let closures: [_; N] = std::array::from_fn(|i| {
2688                    let Self::Mapped { vbit_offsets, .. } = ts[i] else {
2689                        unreachable!()
2690                    };
2691                    let (bytes, rows, cols, out) = (
2692                        ts[i].quant_bytes(),
2693                        ts[i].rows(),
2694                        ts[i].cols(),
2695                        outs_addr[i],
2696                    );
2697                    move |s: usize, e: usize| {
2698                        vbit_range_f32(bytes, vbit_offsets, x, rows, cols, out, s, e)
2699                    }
2700                });
2701                let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2702                    std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2703                pool.run_many(&parts);
2704            }
2705            return;
2706        }
2707
2708        if uniform_f32 {
2709            let outs_addr: [SendMut; N] = std::array::from_fn(|i| SendMut(outs[i].as_mut_ptr()));
2710            let closures: [_; N] = std::array::from_fn(|i| {
2711                let Self::F32 { data, cols, .. } = ts[i] else {
2712                    unreachable!()
2713                };
2714                let out = outs_addr[i];
2715                move |start: usize, end: usize| {
2716                    for o in start..end {
2717                        let row = &data[o * cols..(o + 1) * cols];
2718                        let mut sum = 0.0f32;
2719                        for j in 0..*cols {
2720                            sum += row[j] * x[j];
2721                        }
2722                        // SAFETY: disjoint (tensor, row) cells per worker.
2723                        unsafe { *out.at(o) = sum };
2724                    }
2725                }
2726            });
2727            let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2728                std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2729            pool.run_many(&parts);
2730            return;
2731        }
2732
2733        // Uniform q8-family: per-tensor prescale (q8_2f col fields
2734        // differ per tensor) + the shared range kernels.
2735        struct Ctx<'a> {
2736            bytes: &'a [u8],
2737            #[cfg_attr(not(target_arch = "aarch64"), allow(dead_code))]
2738            rep: &'a [u8],
2739            row_scale: &'a [f32],
2740            cols: usize,
2741            xs: std::borrow::Cow<'a, [f32]>,
2742        }
2743        let ctxs: [Ctx<'_>; N] = std::array::from_fn(|i| {
2744            let Self::Mapped {
2745                dtype,
2746                cols,
2747                row_scale,
2748                col_field,
2749                repack,
2750                ..
2751            } = ts[i]
2752            else {
2753                unreachable!()
2754            };
2755            Ctx {
2756                bytes: ts[i].quant_bytes(),
2757                rep: repack,
2758                row_scale,
2759                cols: *cols,
2760                xs: prescale(x, col_field, *dtype),
2761            }
2762        });
2763        let outs_addr: [SendMut; N] = std::array::from_fn(|i| SendMut(outs[i].as_mut_ptr()));
2764        #[cfg(target_arch = "aarch64")]
2765        if sdot_enabled() {
2766            let acts: [SplitAct; N] = std::array::from_fn(|i| split_act(&ctxs[i].xs));
2767            let closures: [_; N] = std::array::from_fn(|i| {
2768                let (c, act, out) = (&ctxs[i], &acts[i], outs_addr[i]);
2769                move |start: usize, end: usize| {
2770                    q8_range_sdot(c.bytes, c.rep, c.row_scale, act, c.cols, out, start, end)
2771                }
2772            });
2773            let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2774                std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2775            pool.run_many(&parts);
2776            return;
2777        }
2778        #[cfg(target_arch = "x86_64")]
2779        if avx2_a8w8_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_avx2(c.bytes, 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        let closures: [_; N] = std::array::from_fn(|i| {
2793            let (c, out) = (&ctxs[i], outs_addr[i]);
2794            move |start: usize, end: usize| {
2795                q8_range_f32(c.bytes, c.row_scale, &c.xs, c.cols, out, start, end)
2796            }
2797        });
2798        let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2799            std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2800        pool.run_many(&parts);
2801    }
2802}
2803
2804impl QTensor {
2805    /// Pair-input multi-matrix job: N tensors × 2 shared inputs under a
2806    /// single pool dispatch — the MTP/pair decode path publishes one job
2807    /// for Q/K/V (and one for gate+up) instead of one per tensor.
2808    /// Per-row math is exactly `matvec2`'s kernels; bit-identical.
2809    #[allow(clippy::needless_range_loop)]
2810    pub fn matvec2_many<const N: usize>(
2811        ts: [&QTensor; N],
2812        x1: &[f32],
2813        x2: &[f32],
2814        mut o1s: [&mut [f32]; N],
2815        mut o2s: [&mut [f32]; N],
2816        pool: Option<&Pool>,
2817    ) {
2818        let total_rows: usize = ts.iter().map(|t| t.rows()).sum();
2819        if ts.iter().any(|t| t.has_prism_contract()) {
2820            for i in 0..N {
2821                ts[i].matvec2(x1, x2, o1s[i], o2s[i], pool);
2822            }
2823            return;
2824        }
2825        let uniform_q8 = ts.iter().all(|t| {
2826            matches!(
2827                t,
2828                Self::Mapped {
2829                    dtype: TensorDtype::Q8Row | TensorDtype::Q8_2f,
2830                    ..
2831                }
2832            )
2833        });
2834        let uniform_f32 = ts.iter().all(|t| matches!(t, Self::F32 { .. }));
2835        let uniform_q4 = ts.iter().all(|t| {
2836            matches!(
2837                t,
2838                Self::Mapped {
2839                    dtype: TensorDtype::Q4Block,
2840                    ..
2841                }
2842            )
2843        });
2844        let uniform_vbit = ts.iter().all(|t| {
2845            matches!(
2846                t,
2847                Self::Mapped {
2848                    dtype: TensorDtype::Vbit | TensorDtype::VbitRo,
2849                    ..
2850                }
2851            )
2852        });
2853        let fusable = pool.is_some()
2854            && total_rows >= 256
2855            && (uniform_q8 || uniform_f32 || uniform_q4 || uniform_vbit);
2856        if !fusable {
2857            for i in 0..N {
2858                ts[i].matvec2(x1, x2, o1s[i], o2s[i], pool);
2859            }
2860            return;
2861        }
2862        let pool = pool.unwrap();
2863
2864        if uniform_q4 || uniform_vbit {
2865            let p1: [SendMut; N] = std::array::from_fn(|i| SendMut(o1s[i].as_mut_ptr()));
2866            let p2: [SendMut; N] = std::array::from_fn(|i| SendMut(o2s[i].as_mut_ptr()));
2867            // q4/vbit share activation splits — no per-tensor col field.
2868            if a8w8_enabled() {
2869                let a1 = split_act(x1);
2870                let a2 = split_act(x2);
2871                let (a1, a2) = (&a1, &a2);
2872                if uniform_q4 {
2873                    let closures: [_; N] = std::array::from_fn(|i| {
2874                        let (packed, scales) =
2875                            q4_split(ts[i].quant_bytes(), ts[i].rows(), ts[i].cols());
2876                        let (gpr, cols, o1, o2) =
2877                            (ts[i].cols() / GROUP_SIZE, ts[i].cols(), p1[i], p2[i]);
2878                        move |s: usize, e: usize| {
2879                            q4_range2_a8w8(packed, scales, gpr, cols, a1, a2, o1, o2, s, e)
2880                        }
2881                    });
2882                    let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2883                        std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2884                    pool.run_many(&parts);
2885                } else {
2886                    let closures: [_; N] = std::array::from_fn(|i| {
2887                        let Self::Mapped { vbit_offsets, .. } = ts[i] else {
2888                            unreachable!()
2889                        };
2890                        let (bytes, rows, cols, o1, o2) = (
2891                            ts[i].quant_bytes(),
2892                            ts[i].rows(),
2893                            ts[i].cols(),
2894                            p1[i],
2895                            p2[i],
2896                        );
2897                        move |s: usize, e: usize| {
2898                            vbit_range2_a8w8(
2899                                bytes,
2900                                vbit_offsets,
2901                                x1,
2902                                x2,
2903                                a1,
2904                                a2,
2905                                rows,
2906                                cols,
2907                                o1,
2908                                o2,
2909                                s,
2910                                e,
2911                            )
2912                        }
2913                    });
2914                    let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2915                        std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2916                    pool.run_many(&parts);
2917                }
2918                return;
2919            }
2920            if uniform_q4 {
2921                let closures: [_; N] = std::array::from_fn(|i| {
2922                    let (packed, scales) =
2923                        q4_split(ts[i].quant_bytes(), ts[i].rows(), ts[i].cols());
2924                    let (gpr, o1, o2) = (ts[i].cols() / GROUP_SIZE, p1[i], p2[i]);
2925                    move |s: usize, e: usize| {
2926                        q4_range2_f32(packed, scales, gpr, x1, x2, o1, o2, s, e)
2927                    }
2928                });
2929                let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2930                    std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2931                pool.run_many(&parts);
2932            } else {
2933                let closures: [_; N] = std::array::from_fn(|i| {
2934                    let Self::Mapped { vbit_offsets, .. } = ts[i] else {
2935                        unreachable!()
2936                    };
2937                    let (bytes, rows, cols, o1, o2) = (
2938                        ts[i].quant_bytes(),
2939                        ts[i].rows(),
2940                        ts[i].cols(),
2941                        p1[i],
2942                        p2[i],
2943                    );
2944                    move |s: usize, e: usize| {
2945                        vbit_range2_f32(bytes, vbit_offsets, x1, x2, rows, cols, o1, o2, s, e)
2946                    }
2947                });
2948                let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2949                    std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2950                pool.run_many(&parts);
2951            }
2952            return;
2953        }
2954
2955        if uniform_f32 {
2956            let p1: [SendMut; N] = std::array::from_fn(|i| SendMut(o1s[i].as_mut_ptr()));
2957            let p2: [SendMut; N] = std::array::from_fn(|i| SendMut(o2s[i].as_mut_ptr()));
2958            let closures: [_; N] = std::array::from_fn(|i| {
2959                let Self::F32 { data, cols, .. } = ts[i] else {
2960                    unreachable!()
2961                };
2962                let (o1, o2) = (p1[i], p2[i]);
2963                move |start: usize, end: usize| {
2964                    for o in start..end {
2965                        let row = &data[o * cols..(o + 1) * cols];
2966                        let (mut s1, mut s2) = (0.0f32, 0.0f32);
2967                        for j in 0..*cols {
2968                            s1 += row[j] * x1[j];
2969                            s2 += row[j] * x2[j];
2970                        }
2971                        // SAFETY: disjoint (tensor, row) cells per worker.
2972                        unsafe {
2973                            *o1.at(o) = s1;
2974                            *o2.at(o) = s2;
2975                        }
2976                    }
2977                }
2978            });
2979            let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2980                std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2981            pool.run_many(&parts);
2982            return;
2983        }
2984
2985        struct Ctx<'a> {
2986            bytes: &'a [u8],
2987            row_scale: &'a [f32],
2988            cols: usize,
2989            xs1: std::borrow::Cow<'a, [f32]>,
2990            xs2: std::borrow::Cow<'a, [f32]>,
2991        }
2992        let ctxs: [Ctx<'_>; N] = std::array::from_fn(|i| {
2993            let Self::Mapped {
2994                dtype,
2995                cols,
2996                row_scale,
2997                col_field,
2998                ..
2999            } = ts[i]
3000            else {
3001                unreachable!()
3002            };
3003            Ctx {
3004                bytes: ts[i].quant_bytes(),
3005                row_scale,
3006                cols: *cols,
3007                xs1: prescale(x1, col_field, *dtype),
3008                xs2: prescale(x2, col_field, *dtype),
3009            }
3010        });
3011        let p1: [SendMut; N] = std::array::from_fn(|i| SendMut(o1s[i].as_mut_ptr()));
3012        let p2: [SendMut; N] = std::array::from_fn(|i| SendMut(o2s[i].as_mut_ptr()));
3013        #[cfg(target_arch = "aarch64")]
3014        if sdot_enabled() {
3015            let acts: [(SplitAct, SplitAct); N] =
3016                std::array::from_fn(|i| (split_act(&ctxs[i].xs1), split_act(&ctxs[i].xs2)));
3017            let closures: [_; N] = std::array::from_fn(|i| {
3018                let (c, a, o1, o2) = (&ctxs[i], &acts[i], p1[i], p2[i]);
3019                move |start: usize, end: usize| {
3020                    q8_range2_sdot(c.bytes, c.row_scale, &a.0, &a.1, c.cols, o1, o2, start, end)
3021                }
3022            });
3023            let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
3024                std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
3025            pool.run_many(&parts);
3026            return;
3027        }
3028        #[cfg(target_arch = "x86_64")]
3029        if avx2_a8w8_enabled() {
3030            let acts: [(SplitAct, SplitAct); N] =
3031                std::array::from_fn(|i| (split_act(&ctxs[i].xs1), split_act(&ctxs[i].xs2)));
3032            let closures: [_; N] = std::array::from_fn(|i| {
3033                let (c, a, o1, o2) = (&ctxs[i], &acts[i], p1[i], p2[i]);
3034                move |start: usize, end: usize| {
3035                    q8_range2_avx2(c.bytes, c.row_scale, &a.0, &a.1, c.cols, o1, o2, start, end)
3036                }
3037            });
3038            let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
3039                std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
3040            pool.run_many(&parts);
3041            return;
3042        }
3043        let closures: [_; N] = std::array::from_fn(|i| {
3044            let (c, o1, o2) = (&ctxs[i], p1[i], p2[i]);
3045            move |start: usize, end: usize| {
3046                q8_range2_f32(
3047                    c.bytes,
3048                    c.row_scale,
3049                    &c.xs1,
3050                    &c.xs2,
3051                    c.cols,
3052                    o1,
3053                    o2,
3054                    start,
3055                    end,
3056                )
3057            }
3058        });
3059        let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
3060            std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
3061        pool.run_many(&parts);
3062    }
3063
3064    /// Fused gate+up matvec with SiLU·mul: for each row r, computes
3065    /// `silu(gate·x) * (up·x)` and writes to `out[r]`. ONE pool dispatch,
3066    /// no intermediate g/u buffers, no separate silu pass. Falls back
3067    /// (returns false) for unsupported dtype combos.
3068    pub fn matvec_silu_mul(
3069        gate: &QTensor,
3070        up: &QTensor,
3071        x: &[f32],
3072        out: &mut [f32],
3073        pool: Option<&Pool>,
3074    ) -> bool {
3075        Self::matvec_silu_mul_limited(gate, up, x, out, 0.0, pool)
3076    }
3077
3078    /// Fused gate+up+SiLU with the GLM asymmetrical clamp.  `limit == 0`
3079    /// preserves the historical unclamped helper; a positive limit clamps
3080    /// `up` to both sides and `gate` only from above, matching the GLM
3081    /// SwiGLU reference.  Keeping the limit in the row kernel avoids the two
3082    /// intermediate vectors and the extra combine pass on the Q2TP experts.
3083    pub fn matvec_silu_mul_limited(
3084        gate: &QTensor,
3085        up: &QTensor,
3086        x: &[f32],
3087        out: &mut [f32],
3088        limit: f32,
3089        pool: Option<&Pool>,
3090    ) -> bool {
3091        if gate.has_prism_contract() || up.has_prism_contract() {
3092            // The fused gate/up kernels consume x directly.  Prism requires
3093            // a per-matrix signed FWHT, so the caller must use two ordinary
3094            // descriptor-aware matvecs instead of an unrotated fast path.
3095            return false;
3096        }
3097        let inter = gate.rows();
3098        debug_assert_eq!(up.rows(), inter);
3099        debug_assert_eq!(out.len(), inter);
3100        debug_assert_eq!(gate.cols(), up.cols());
3101        if !a8w8_enabled() {
3102            return false;
3103        }
3104        let act = split_act(x);
3105        let act = &act;
3106        let x_ref = x;
3107        let out_addr = SendMut(out.as_mut_ptr());
3108
3109        match (gate, up) {
3110            // Q4Block gate + Q4Block up (most common mobile q4 models)
3111            (
3112                Self::Mapped {
3113                    dtype: TensorDtype::Q4Block,
3114                    ..
3115                },
3116                Self::Mapped {
3117                    dtype: TensorDtype::Q4Block,
3118                    ..
3119                },
3120            ) => {
3121                let (gp, gs) = q4_split(gate.quant_bytes(), gate.rows(), gate.cols());
3122                let (up_p, up_s) = q4_split(up.quant_bytes(), up.rows(), up.cols());
3123                let gpr = gate.cols() / GROUP_SIZE;
3124                let cols = gate.cols();
3125                let run = move |start: usize, end: usize| {
3126                    for r in start..end {
3127                        let mut gv = dot_q4_row_i8(gp, gs, r * gpr, gpr, &act.xq) * act.sx;
3128                        let mut uv = dot_q4_row_i8(up_p, up_s, r * gpr, gpr, &act.xq) * act.sx;
3129                        for &(j, xv) in &act.outliers {
3130                            let flat = r * cols + j;
3131                            let gb = gp[flat / 2];
3132                            let gn = if flat & 1 == 0 { gb & 0x0F } else { gb >> 4 };
3133                            let gsc = f16_to_f32(u16::from_le_bytes([
3134                                gs[(flat / GROUP_SIZE) * 2],
3135                                gs[(flat / GROUP_SIZE) * 2 + 1],
3136                            ]));
3137                            gv += ((gn as i32 - 8) as f32) * gsc * xv;
3138                            let ub = up_p[flat / 2];
3139                            let un = if flat & 1 == 0 { ub & 0x0F } else { ub >> 4 };
3140                            let usc = f16_to_f32(u16::from_le_bytes([
3141                                up_s[(flat / GROUP_SIZE) * 2],
3142                                up_s[(flat / GROUP_SIZE) * 2 + 1],
3143                            ]));
3144                            uv += ((un as i32 - 8) as f32) * usc * xv;
3145                        }
3146                        let silu_g = gv / (1.0 + (-gv).exp());
3147                        // SAFETY: disjoint row ranges per worker.
3148                        unsafe { *out_addr.at(r) = silu_g * uv };
3149                    }
3150                };
3151                dispatch_rows(pool, inter, &run);
3152                true
3153            }
3154            // Q4Tiled gate + Q4Tiled up — one row pass, both tile
3155            // streams sequential, silu·mul fused (same per-row math as
3156            // `q4t_matvec`).
3157            (
3158                Self::Mapped {
3159                    dtype: TensorDtype::Q4Tiled,
3160                    ..
3161                },
3162                Self::Mapped {
3163                    dtype: TensorDtype::Q4Tiled,
3164                    ..
3165                },
3166            ) => {
3167                let g_bytes = gate.quant_bytes();
3168                let u_bytes = up.quant_bytes();
3169                let gpr = gate.cols() / GROUP_SIZE;
3170                let run = move |start: usize, end: usize| {
3171                    for r in start..end {
3172                        let mut gv = dot_q4t_row_i8(g_bytes, r, gpr, &act.xq) * act.sx;
3173                        let mut uv = dot_q4t_row_i8(u_bytes, r, gpr, &act.xq) * act.sx;
3174                        for &(j, xv) in &act.outliers {
3175                            let (w, s) = q4t_outlier(g_bytes, r, gpr, j);
3176                            gv += w * s * xv;
3177                            let (w, s) = q4t_outlier(u_bytes, r, gpr, j);
3178                            uv += w * s * xv;
3179                        }
3180                        let silu_g = gv / (1.0 + (-gv).exp());
3181                        // SAFETY: disjoint row ranges per worker.
3182                        unsafe { *out_addr.at(r) = silu_g * uv };
3183                    }
3184                };
3185                dispatch_rows(pool, inter, &run);
3186                true
3187            }
3188            // Q4TiledP gate + Q4TiledP up — the same fused row pass, with
3189            // each row's two ladders built once and spent on both streams.
3190            (
3191                Self::Mapped {
3192                    dtype: TensorDtype::Q4TiledP,
3193                    ..
3194                },
3195                Self::Mapped {
3196                    dtype: TensorDtype::Q4TiledP,
3197                    ..
3198                },
3199            ) => {
3200                let cols = gate.cols();
3201                let gpr = cols / GROUP_SIZE;
3202                let gv_view = Q4tpView::new(gate.quant_bytes(), inter, cols);
3203                let uv_view = Q4tpView::new(up.quant_bytes(), inter, cols);
3204                let run = |start: usize, end: usize| {
3205                    let (mut gsc, mut usc) = (vec![0f32; gpr], vec![0f32; gpr]);
3206                    for r in start..end {
3207                        gv_view.scales_into(r, gpr, &mut gsc);
3208                        uv_view.scales_into(r, gpr, &mut usc);
3209                        let mut gv = dot_q4tp_row_i8(gv_view.nib, r, gpr, &act.xq, &gsc) * act.sx;
3210                        let mut uv = dot_q4tp_row_i8(uv_view.nib, r, gpr, &act.xq, &usc) * act.sx;
3211                        for &(j, xv) in &act.outliers {
3212                            let (w, s) = q4tp_outlier(gv_view.nib, r, gpr, j, &gsc);
3213                            gv += w * s * xv;
3214                            let (w, s) = q4tp_outlier(uv_view.nib, r, gpr, j, &usc);
3215                            uv += w * s * xv;
3216                        }
3217                        let silu_g = gv / (1.0 + (-gv).exp());
3218                        // SAFETY: disjoint row ranges per worker.
3219                        unsafe { *out_addr.at(r) = silu_g * uv };
3220                    }
3221                };
3222                dispatch_rows(pool, inter, &run);
3223                true
3224            }
3225            // Q1 gate + Q1 up — one row pass over both sign streams,
3226            // silu·mul fused (the per-row math of `q1_range_a8w8`); the
3227            // activation group sums are shared by both streams. Without
3228            // this arm a q1 dense FFN paid two dispatches + a combine
3229            // loop — the exact barrier this function exists to remove.
3230            (
3231                Self::Mapped {
3232                    dtype: TensorDtype::Q1,
3233                    ..
3234                },
3235                Self::Mapped {
3236                    dtype: TensorDtype::Q1,
3237                    ..
3238                },
3239            ) => {
3240                let g_bytes = gate.quant_bytes();
3241                let u_bytes = up.quant_bytes();
3242                let gpr = gate.cols() / GROUP_SIZE;
3243                let gsum = q1_group_sums(&act.xq, gpr);
3244                let gsum = &gsum;
3245                let run = move |start: usize, end: usize| {
3246                    for r in start..end {
3247                        let mut gv = dot_q1_row_i8(g_bytes, r, gpr, &act.xq, gsum) * act.sx;
3248                        let mut uv = dot_q1_row_i8(u_bytes, r, gpr, &act.xq, gsum) * act.sx;
3249                        for &(j, xv) in &act.outliers {
3250                            let (w, s) = q1_outlier(g_bytes, r, gpr, j);
3251                            gv += w * s * xv;
3252                            let (w, s) = q1_outlier(u_bytes, r, gpr, j);
3253                            uv += w * s * xv;
3254                        }
3255                        let silu_g = gv / (1.0 + (-gv).exp());
3256                        // SAFETY: disjoint row ranges per worker.
3257                        unsafe { *out_addr.at(r) = silu_g * uv };
3258                    }
3259                };
3260                dispatch_rows(pool, inter, &run);
3261                true
3262            }
3263            // Q2TiledP gate + Q2TiledP up — the 2-bit expert pair (MoE
3264            // FFNs of the W2 class): one row pass, both ladders built
3265            // once, integer code dots with shared group sums.
3266            (
3267                Self::Mapped {
3268                    dtype: TensorDtype::Q2TiledP,
3269                    ..
3270                },
3271                Self::Mapped {
3272                    dtype: TensorDtype::Q2TiledP,
3273                    ..
3274                },
3275            ) => {
3276                let cols = gate.cols();
3277                let gpr = cols / GROUP_SIZE;
3278                let gv_view = Q4tpView::new_q2(gate.quant_bytes(), inter, cols);
3279                let uv_view = Q4tpView::new_q2(up.quant_bytes(), inter, cols);
3280                let gsum = q1_group_sums(&act.xq, gpr);
3281                let gsum = &gsum;
3282                let run = move |start: usize, end: usize| {
3283                    let (mut gsc, mut usc) = (vec![0f32; gpr], vec![0f32; gpr]);
3284                    for r in start..end {
3285                        gv_view.scales_into(r, gpr, &mut gsc);
3286                        uv_view.scales_into(r, gpr, &mut usc);
3287                        let mut gv =
3288                            dot_q2tp_row_i8(gv_view.nib, r, gpr, &act.xq, gsum, &gsc) * act.sx;
3289                        let mut uv =
3290                            dot_q2tp_row_i8(uv_view.nib, r, gpr, &act.xq, gsum, &usc) * act.sx;
3291                        for &(j, xv) in &act.outliers {
3292                            let (w, s) = q2tp_outlier(gv_view.nib, r, gpr, j, &gsc);
3293                            gv += w * s * xv;
3294                            let (w, s) = q2tp_outlier(uv_view.nib, r, gpr, j, &usc);
3295                            uv += w * s * xv;
3296                        }
3297                        let silu_g = gv / (1.0 + (-gv).exp());
3298                        // SAFETY: disjoint row ranges per worker.
3299                        unsafe { *out_addr.at(r) = silu_g * uv };
3300                    }
3301                };
3302                dispatch_rows(pool, inter, &run);
3303                true
3304            }
3305            // Q8Row gate + Q8Row up — one row pass over both i8 streams.
3306            // Q8_2f stays out on purpose: its column field prescales the
3307            // activations PER TENSOR, which breaks this fn's shared
3308            // split_act contract — it keeps the two-dispatch path.
3309            (
3310                Self::Mapped {
3311                    dtype: TensorDtype::Q8Row,
3312                    row_scale: g_rs,
3313                    ..
3314                },
3315                Self::Mapped {
3316                    dtype: TensorDtype::Q8Row,
3317                    row_scale: u_rs,
3318                    ..
3319                },
3320            ) => {
3321                let g_bytes = gate.quant_bytes();
3322                let u_bytes = up.quant_bytes();
3323                let cols = gate.cols();
3324                let run = move |start: usize, end: usize| {
3325                    for r in start..end {
3326                        let gv = q8_row_dot(&g_bytes[r * cols..(r + 1) * cols], act) * g_rs[r];
3327                        let uv = q8_row_dot(&u_bytes[r * cols..(r + 1) * cols], act) * u_rs[r];
3328                        let silu_g = gv / (1.0 + (-gv).exp());
3329                        // SAFETY: disjoint row ranges per worker.
3330                        unsafe { *out_addr.at(r) = silu_g * uv };
3331                    }
3332                };
3333                dispatch_rows(pool, inter, &run);
3334                true
3335            }
3336            // Q1T gate + Q1T up
3337            (
3338                Self::Mapped {
3339                    dtype: TensorDtype::Q1T,
3340                    ..
3341                },
3342                Self::Mapped {
3343                    dtype: TensorDtype::Q1T,
3344                    ..
3345                },
3346            ) => {
3347                const TILE: usize = cortiq_core::quant::Q1T_TILE;
3348                let g_bytes = gate.quant_bytes();
3349                let u_bytes = up.quant_bytes();
3350                let gpr = gate.cols() / GROUP_SIZE;
3351                let (g_rp, g_ent, g_ov) = q1t_overlay(g_bytes, inter * gpr * TILE, inter);
3352                let (u_rp, u_ent, u_ov) = q1t_overlay(u_bytes, inter * gpr * TILE, inter);
3353                let run = move |start: usize, end: usize| {
3354                    for r in start..end {
3355                        let mut gv = q1t_dot_row_i8(g_bytes, r, gpr, &act.xq) * act.sx;
3356                        let mut uv = q1t_dot_row_i8(u_bytes, r, gpr, &act.xq) * act.sx;
3357                        for &(j, xv) in &act.outliers {
3358                            gv += q1t_base_weight(g_bytes, r, gpr, j) * xv;
3359                            uv += q1t_base_weight(u_bytes, r, gpr, j) * xv;
3360                        }
3361                        gv += q1t_row_outlier_correction(g_bytes, r, g_rp, g_ent, g_ov, x_ref);
3362                        uv += q1t_row_outlier_correction(u_bytes, r, u_rp, u_ent, u_ov, x_ref);
3363                        let silu_g = gv / (1.0 + (-gv).exp());
3364                        // SAFETY: disjoint row ranges per worker.
3365                        unsafe { *out_addr.at(r) = silu_g * uv };
3366                    }
3367                };
3368                dispatch_rows(pool, inter, &run);
3369                true
3370            }
3371            _ => false,
3372        }
3373    }
3374
3375    /// Every routed expert's fused gate/up/SiLU under ONE pool dispatch.
3376    ///
3377    /// The per-expert path pays a pool barrier per expert per stage: at 9
3378    /// experts over 40 layers that is ~720 barriers a token, and a decode
3379    /// profile of Qwen3.6-35B-A3B showed the pool parked in
3380    /// `psynch_cvwait` about twice as long as it spent computing. Laying
3381    /// every expert's rows end-to-end in one virtual row space collapses
3382    /// the stage to a single dispatch. The per-row body is the
3383    /// single-expert q4tp arm verbatim, so outputs are bit-identical.
3384    ///
3385    /// `false` = something is outside the fused q4tp kernel (dtype, shape,
3386    /// or a transformed tensor); the caller walks the ordinary per-expert
3387    /// path. Float activations use the same exact scalar rows, still fused
3388    /// under one pool dispatch.
3389    pub fn moe_gate_up_many(
3390        pairs: &[(&QTensor, &QTensor)],
3391        x: &[f32],
3392        outs: &mut [Vec<f32>],
3393        pool: Option<&Pool>,
3394    ) -> bool {
3395        if pairs.is_empty() || pairs.len() != outs.len() {
3396            return false;
3397        }
3398        if !a8w8_enabled() {
3399            let groups = vec![vec![0]; pairs.len()];
3400            return Self::moe_gate_up_rows(pairs, &groups, x, outs, pool);
3401        }
3402        let inter = pairs[0].0.rows();
3403        let cols = pairs[0].0.cols();
3404        if cols % GROUP_SIZE != 0 {
3405            return false;
3406        }
3407        let gpr = cols / GROUP_SIZE;
3408        // Uniform layout across every routed pair: q4tp, or the 2-bit
3409        // profile's q2tp gate/up (the W2 class). Mixed sets refuse.
3410        let q2 = matches!(
3411            pairs[0].0,
3412            Self::Mapped {
3413                dtype: TensorDtype::Q2TiledP,
3414                ..
3415            }
3416        );
3417        let want = if q2 {
3418            TensorDtype::Q2TiledP
3419        } else {
3420            TensorDtype::Q4TiledP
3421        };
3422        let mut views = Vec::with_capacity(pairs.len() * 2);
3423        for ((g, u), o) in pairs.iter().zip(outs.iter()) {
3424            let both = matches!(g, Self::Mapped { dtype, .. } if *dtype == want)
3425                && matches!(u, Self::Mapped { dtype, .. } if *dtype == want);
3426            if !both
3427                || g.rows() != inter
3428                || u.rows() != inter
3429                || g.cols() != cols
3430                || u.cols() != cols
3431                || o.len() != inter
3432            {
3433                return false;
3434            }
3435            let mk = if q2 { Q4tpView::new_q2 } else { Q4tpView::new };
3436            views.push(mk(g.quant_bytes(), inter, cols));
3437            views.push(mk(u.quant_bytes(), inter, cols));
3438        }
3439        let act = split_act(x);
3440        let gsum = if q2 {
3441            q1_group_sums(&act.xq, gpr)
3442        } else {
3443            Vec::new()
3444        };
3445        let (act, gsum) = (&act, &gsum);
3446        let ptrs: Vec<SendMut> = outs.iter_mut().map(|o| SendMut(o.as_mut_ptr())).collect();
3447        let (views, ptrs) = (&views, &ptrs);
3448        let run = |start: usize, end: usize| {
3449            let (mut gsc, mut usc) = (vec![0f32; gpr], vec![0f32; gpr]);
3450            for flat in start..end {
3451                let (e, r) = (flat / inter, flat % inter);
3452                let gv_view = &views[e * 2];
3453                let uv_view = &views[e * 2 + 1];
3454                gv_view.scales_into(r, gpr, &mut gsc);
3455                uv_view.scales_into(r, gpr, &mut usc);
3456                let (mut gv, mut uv) = if q2 {
3457                    (
3458                        dot_q2tp_row_i8(gv_view.nib, r, gpr, &act.xq, gsum, &gsc) * act.sx,
3459                        dot_q2tp_row_i8(uv_view.nib, r, gpr, &act.xq, gsum, &usc) * act.sx,
3460                    )
3461                } else {
3462                    (
3463                        dot_q4tp_row_i8(gv_view.nib, r, gpr, &act.xq, &gsc) * act.sx,
3464                        dot_q4tp_row_i8(uv_view.nib, r, gpr, &act.xq, &usc) * act.sx,
3465                    )
3466                };
3467                for &(j, xv) in &act.outliers {
3468                    let (og, ou) = if q2 {
3469                        (
3470                            q2tp_outlier(gv_view.nib, r, gpr, j, &gsc),
3471                            q2tp_outlier(uv_view.nib, r, gpr, j, &usc),
3472                        )
3473                    } else {
3474                        (
3475                            q4tp_outlier(gv_view.nib, r, gpr, j, &gsc),
3476                            q4tp_outlier(uv_view.nib, r, gpr, j, &usc),
3477                        )
3478                    };
3479                    gv += og.0 * og.1 * xv;
3480                    uv += ou.0 * ou.1 * xv;
3481                }
3482                let silu_g = gv / (1.0 + (-gv).exp());
3483                // SAFETY: one worker owns each (expert, row) pair.
3484                unsafe { *ptrs[e].at(r) = silu_g * uv };
3485            }
3486        };
3487        dispatch_rows(pool, pairs.len() * inter, &run);
3488        true
3489    }
3490
3491    /// Every routed expert's down projection, weighted and summed into
3492    /// `out`, under ONE pool dispatch.
3493    ///
3494    /// Partitioned by OUTPUT row rather than by expert: each row is owned
3495    /// by a single worker, so the experts are summed in the caller's order
3496    /// — the same sequence of f32 adds the serial `out[i] += w·eo[i]` loop
3497    /// performs, hence bit-identical. Partitioning by expert instead would
3498    /// race on the shared accumulator.
3499    pub fn moe_down_many(
3500        downs: &[&QTensor],
3501        gs: &[Vec<f32>],
3502        weights: &[f32],
3503        out: &mut [f32],
3504        pool: Option<&Pool>,
3505    ) -> bool {
3506        if downs.is_empty() || downs.len() != gs.len() || downs.len() != weights.len() {
3507            return false;
3508        }
3509        if !a8w8_enabled() {
3510            let mut terms = vec![vec![0.0; out.len()]; downs.len()];
3511            if !Self::moe_down_rows(downs, &vec![1; downs.len()], gs, &mut terms, pool) {
3512                return false;
3513            }
3514            out.fill(0.0);
3515            for (row, &w) in terms.iter().zip(weights) {
3516                for (o, &v) in out.iter_mut().zip(row) {
3517                    *o += w * v;
3518                }
3519            }
3520            return true;
3521        }
3522        let rows = out.len();
3523        let cols = downs[0].cols();
3524        if cols % GROUP_SIZE != 0 {
3525            return false;
3526        }
3527        let gpr = cols / GROUP_SIZE;
3528        let mut views = Vec::with_capacity(downs.len());
3529        for (d, g) in downs.iter().zip(gs.iter()) {
3530            if !matches!(
3531                d,
3532                Self::Mapped {
3533                    dtype: TensorDtype::Q4TiledP,
3534                    ..
3535                }
3536            ) || d.rows() != rows
3537                || d.cols() != cols
3538                || g.len() != cols
3539            {
3540                return false;
3541            }
3542            views.push(Q4tpView::new(d.quant_bytes(), rows, cols));
3543        }
3544        // One int8 split per expert — the activation vectors differ.
3545        let acts: Vec<SplitAct> = gs.iter().map(|g| split_act(g)).collect();
3546        // Partitioned by OUTPUT row, with the experts folded inside: each
3547        // row is owned by one worker, so they are summed in the caller's
3548        // order — the same f32 sequence the serial `out[i] += w·eo[i]`
3549        // loop produces. Partitioning by expert instead would either race
3550        // on the accumulator or need a scratch plane and a second pass;
3551        // measured, that variant was a wash, so this keeps the simpler
3552        // shape.
3553        let out_addr = SendMut(out.as_mut_ptr());
3554        let (views, acts, weights) = (&views, &acts, &weights);
3555        let run = |start: usize, end: usize| {
3556            let mut sc = vec![0f32; gpr];
3557            for r in start..end {
3558                let mut acc = 0f32;
3559                for (e, v) in views.iter().enumerate() {
3560                    v.scales_into(r, gpr, &mut sc);
3561                    let a = &acts[e];
3562                    let mut d = dot_q4tp_row_i8(v.nib, r, gpr, &a.xq, &sc) * a.sx;
3563                    for &(j, xv) in &a.outliers {
3564                        let (w, s) = q4tp_outlier(v.nib, r, gpr, j, &sc);
3565                        d += w * s * xv;
3566                    }
3567                    acc += weights[e] * d;
3568                }
3569                // SAFETY: disjoint row ranges per worker.
3570                unsafe { *out_addr.at(r) = acc };
3571            }
3572        };
3573        dispatch_rows(pool, rows, &run);
3574        true
3575    }
3576
3577    /// `moe_gate_up_many` for SEVERAL tokens at once, decode-exact: expert
3578    /// `e` (`pairs[e]`, q4tp) serves the tokens `groups[e]` (row indices
3579    /// into `xs`, each `cols` wide). Every (expert, token) output is
3580    /// bit-identical to `moe_gate_up_many` run on that token alone — the
3581    /// same int8 activation split, VNNI dots, outlier terms and inline
3582    /// SiLU — while each weight row is read once for all the tokens routed
3583    /// to its expert (the speculative verify's expert sharing). `outs` is
3584    /// flat in (expert, token-of-group) order. False = not covered (not
3585    /// q4tp): the caller takes the per-token path. With float activations,
3586    /// the exact scalar row kernel replaces the int8 dot without changing
3587    /// the shared dispatch or route-order reduction.
3588    pub fn moe_gate_up_rows(
3589        pairs: &[(&QTensor, &QTensor)],
3590        groups: &[Vec<usize>],
3591        xs: &[f32],
3592        outs: &mut [Vec<f32>],
3593        pool: Option<&Pool>,
3594    ) -> bool {
3595        if pairs.is_empty() || pairs.len() != groups.len() {
3596            return false;
3597        }
3598        let inter = pairs[0].0.rows();
3599        let cols = pairs[0].0.cols();
3600        let n_pairs: usize = groups.iter().map(|g| g.len()).sum();
3601        if cols == 0 || cols % GROUP_SIZE != 0 || outs.len() != n_pairs || xs.len() % cols != 0 {
3602            return false;
3603        }
3604        let b = xs.len() / cols;
3605        let gpr = cols / GROUP_SIZE;
3606        let mut views = Vec::with_capacity(pairs.len() * 2);
3607        for (g, u) in pairs {
3608            let q4tp = |t: &QTensor| {
3609                matches!(
3610                    t,
3611                    Self::Mapped {
3612                        dtype: TensorDtype::Q4TiledP,
3613                        ..
3614                    }
3615                )
3616            };
3617            if g.has_prism_contract()
3618                || u.has_prism_contract()
3619                || !q4tp(g)
3620                || !q4tp(u)
3621                || g.rows() != inter
3622                || u.rows() != inter
3623                || g.cols() != cols
3624                || u.cols() != cols
3625            {
3626                return false;
3627            }
3628            views.push(Q4tpView::new(g.quant_bytes(), inter, cols));
3629            views.push(Q4tpView::new(u.quant_bytes(), inter, cols));
3630        }
3631        if outs.iter().any(|o| o.len() != inter) || groups.iter().flatten().any(|&t| t >= b) {
3632            return false;
3633        }
3634        let quantized = a8w8_enabled();
3635        let acts: Vec<SplitAct> = if quantized {
3636            (0..b)
3637                .map(|t| split_act(&xs[t * cols..(t + 1) * cols]))
3638                .collect()
3639        } else {
3640            Vec::new()
3641        };
3642        let mut offs = Vec::with_capacity(groups.len());
3643        let mut o = 0usize;
3644        for g in groups {
3645            offs.push(o);
3646            o += g.len();
3647        }
3648        let ptrs: Vec<SendMut> = outs.iter_mut().map(|o| SendMut(o.as_mut_ptr())).collect();
3649        let (views, ptrs, acts, offs) = (&views, &ptrs, &acts, &offs);
3650        let run = |start: usize, end: usize| {
3651            let (mut gsc, mut usc) = (vec![0f32; gpr], vec![0f32; gpr]);
3652            for flat in start..end {
3653                let (e, r) = (flat / inter, flat % inter);
3654                let (gv_view, uv_view) = (&views[e * 2], &views[e * 2 + 1]);
3655                gv_view.scales_into(r, gpr, &mut gsc);
3656                uv_view.scales_into(r, gpr, &mut usc);
3657                for (k, &t) in groups[e].iter().enumerate() {
3658                    if !quantized {
3659                        let x = &xs[t * cols..(t + 1) * cols];
3660                        let gv = q4tp_row_exact(gv_view.nib, r, gpr, x, &gsc);
3661                        let uv = q4tp_row_exact(uv_view.nib, r, gpr, x, &usc);
3662                        unsafe { *ptrs[offs[e] + k].at(r) = (gv / (1.0 + (-gv).exp())) * uv };
3663                        continue;
3664                    }
3665                    let act = &acts[t];
3666                    let mut gv = dot_q4tp_row_i8(gv_view.nib, r, gpr, &act.xq, &gsc) * act.sx;
3667                    let mut uv = dot_q4tp_row_i8(uv_view.nib, r, gpr, &act.xq, &usc) * act.sx;
3668                    for &(j, xv) in &act.outliers {
3669                        let og = q4tp_outlier(gv_view.nib, r, gpr, j, &gsc);
3670                        let ou = q4tp_outlier(uv_view.nib, r, gpr, j, &usc);
3671                        gv += og.0 * og.1 * xv;
3672                        uv += ou.0 * ou.1 * xv;
3673                    }
3674                    let silu_g = gv / (1.0 + (-gv).exp());
3675                    // SAFETY: one worker owns each (expert, row) cell of
3676                    // every output of the expert's group.
3677                    unsafe { *ptrs[offs[e] + k].at(r) = silu_g * uv };
3678                }
3679            }
3680        };
3681        dispatch_rows(pool, pairs.len() * inter, &run);
3682        true
3683    }
3684
3685    /// The per-(expert, token) down terms `moe_down_many` weights and sums,
3686    /// for SEVERAL tokens: `outs[p][o] = down_e[o] · gs[p]` (int8 split of
3687    /// `gs[p]`, VNNI dot, outlier terms — bit-identical to that kernel's
3688    /// `d`), each down row read once for its expert's whole group. The
3689    /// caller sums `w·d` per token in its route order, which reproduces
3690    /// `moe_down_many`'s f32 sequence exactly. Layout as `moe_gate_up_rows`.
3691    pub fn moe_down_rows(
3692        downs: &[&QTensor],
3693        group_lens: &[usize],
3694        gs: &[Vec<f32>],
3695        outs: &mut [Vec<f32>],
3696        pool: Option<&Pool>,
3697    ) -> bool {
3698        if downs.is_empty() || downs.len() != group_lens.len() {
3699            return false;
3700        }
3701        let rows = downs[0].rows();
3702        let cols = downs[0].cols();
3703        let n_pairs: usize = group_lens.iter().sum();
3704        if cols == 0 || cols % GROUP_SIZE != 0 || gs.len() != n_pairs || outs.len() != n_pairs {
3705            return false;
3706        }
3707        let gpr = cols / GROUP_SIZE;
3708        let mut views = Vec::with_capacity(downs.len());
3709        for d in downs {
3710            if d.has_prism_contract()
3711                || !matches!(
3712                    d,
3713                    Self::Mapped {
3714                        dtype: TensorDtype::Q4TiledP,
3715                        ..
3716                    }
3717                )
3718                || d.rows() != rows
3719                || d.cols() != cols
3720            {
3721                return false;
3722            }
3723            views.push(Q4tpView::new(d.quant_bytes(), rows, cols));
3724        }
3725        if gs.iter().any(|g| g.len() != cols) || outs.iter().any(|o| o.len() != rows) {
3726            return false;
3727        }
3728        let quantized = a8w8_enabled();
3729        let acts: Vec<SplitAct> = if quantized {
3730            gs.iter().map(|g| split_act(g)).collect()
3731        } else {
3732            Vec::new()
3733        };
3734        let mut offs = Vec::with_capacity(group_lens.len());
3735        let mut o = 0usize;
3736        for &l in group_lens {
3737            offs.push(o);
3738            o += l;
3739        }
3740        let ptrs: Vec<SendMut> = outs.iter_mut().map(|o| SendMut(o.as_mut_ptr())).collect();
3741        let (views, ptrs, acts, offs) = (&views, &ptrs, &acts, &offs);
3742        let run = |start: usize, end: usize| {
3743            let mut sc = vec![0f32; gpr];
3744            for flat in start..end {
3745                let (e, r) = (flat / rows, flat % rows);
3746                let v = &views[e];
3747                v.scales_into(r, gpr, &mut sc);
3748                for k in 0..group_lens[e] {
3749                    if !quantized {
3750                        let d = q4tp_row_exact(v.nib, r, gpr, &gs[offs[e] + k], &sc);
3751                        unsafe { *ptrs[offs[e] + k].at(r) = d };
3752                        continue;
3753                    }
3754                    let a = &acts[offs[e] + k];
3755                    let mut d = dot_q4tp_row_i8(v.nib, r, gpr, &a.xq, &sc) * a.sx;
3756                    for &(j, xv) in &a.outliers {
3757                        let (w, s) = q4tp_outlier(v.nib, r, gpr, j, &sc);
3758                        d += w * s * xv;
3759                    }
3760                    // SAFETY: one worker owns each (expert, row) cell.
3761                    unsafe { *ptrs[offs[e] + k].at(r) = d };
3762                }
3763            }
3764        };
3765        dispatch_rows(pool, downs.len() * rows, &run);
3766        true
3767    }
3768}
3769
3770/// Batched q8 kernel: same math as qmatvec, the row makes a single
3771/// pass from memory for the whole batch.
3772/// Accelerate CBLAS — the Apple AMX matrix units, the same engine
3773/// llama.cpp's `-ngl 0` prefill rides via ggml-blas.
3774#[cfg(target_os = "macos")]
3775mod accel_blas {
3776    #[link(name = "Accelerate", kind = "framework")]
3777    unsafe extern "C" {
3778        pub fn cblas_sgemm(
3779            order: i32,
3780            trans_a: i32,
3781            trans_b: i32,
3782            m: i32,
3783            n: i32,
3784            k: i32,
3785            alpha: f32,
3786            a: *const f32,
3787            lda: i32,
3788            b: *const f32,
3789            ldb: i32,
3790            beta: f32,
3791            c: *mut f32,
3792            ldc: i32,
3793        );
3794    }
3795}
3796
3797#[cfg(target_os = "macos")]
3798pub(crate) fn accel_gemm_enabled() -> bool {
3799    static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
3800    *ON.get_or_init(|| std::env::var("CMF_ACCEL").map(|v| v != "0").unwrap_or(true))
3801}
3802
3803/// Off macOS the "accel" GEMM is the portable NEON micro-kernel below —
3804/// same entry point, so the batched-attention path opens on mobile.
3805#[cfg(all(target_arch = "aarch64", not(target_os = "macos")))]
3806pub(crate) fn accel_gemm_enabled() -> bool {
3807    true
3808}
3809
3810/// Portable NEON f32 GEMM (row-major, optional Bᵀ): a 4×8 fmla
3811/// micro-kernel with A broadcast against B panels — the mobile stand-in
3812/// for Accelerate in the batched causal attention (QKᵀ and P·V). Not a
3813/// BLAS: shapes here are the attention panels (m ≤ heads·chunk,
3814/// k = head_dim or context), and the goal is removing the per-position
3815/// quadratic wall, not peak GEMM.
3816#[cfg(target_arch = "aarch64")]
3817#[allow(clippy::too_many_arguments)]
3818pub(crate) fn neon_gemm_rm(
3819    m: usize,
3820    n: usize,
3821    k: usize,
3822    alpha: f32,
3823    a: &[f32],
3824    lda: usize,
3825    b_mat: &[f32],
3826    ldb: usize,
3827    b_rows_are_n: bool,
3828    c: &mut [f32],
3829    ldc: usize,
3830) {
3831    debug_assert!(a.len() >= (m - 1) * lda + k);
3832    debug_assert!(c.len() >= (m - 1) * ldc + n);
3833    // SAFETY: bounds asserted above; NEON is baseline on aarch64.
3834    unsafe {
3835        use core::arch::aarch64::*;
3836        let mut i = 0usize;
3837        while i < m {
3838            let mi = (m - i).min(4);
3839            let mut j = 0usize;
3840            while j < n {
3841                let nj = (n - j).min(8);
3842                if mi == 4 && nj == 8 {
3843                    let (mut c0a, mut c0b) = (vdupq_n_f32(0.0), vdupq_n_f32(0.0));
3844                    let (mut c1a, mut c1b) = (vdupq_n_f32(0.0), vdupq_n_f32(0.0));
3845                    let (mut c2a, mut c2b) = (vdupq_n_f32(0.0), vdupq_n_f32(0.0));
3846                    let (mut c3a, mut c3b) = (vdupq_n_f32(0.0), vdupq_n_f32(0.0));
3847                    for p in 0..k {
3848                        let (b0, b1) = if b_rows_are_n {
3849                            // B is [n, k]: column p of Bᵀ = element p of
3850                            // eight consecutive B rows — gathered.
3851                            let base = b_mat.as_ptr().add(j * ldb + p);
3852                            let g = |o: usize| *base.add(o * ldb);
3853                            ([g(0), g(1), g(2), g(3)], [g(4), g(5), g(6), g(7)])
3854                        } else {
3855                            let base = b_mat.as_ptr().add(p * ldb + j);
3856                            (
3857                                [*base, *base.add(1), *base.add(2), *base.add(3)],
3858                                [*base.add(4), *base.add(5), *base.add(6), *base.add(7)],
3859                            )
3860                        };
3861                        let bv0 = vld1q_f32(b0.as_ptr());
3862                        let bv1 = vld1q_f32(b1.as_ptr());
3863                        let a0 = vdupq_n_f32(*a.as_ptr().add(i * lda + p));
3864                        let a1 = vdupq_n_f32(*a.as_ptr().add((i + 1) * lda + p));
3865                        let a2 = vdupq_n_f32(*a.as_ptr().add((i + 2) * lda + p));
3866                        let a3 = vdupq_n_f32(*a.as_ptr().add((i + 3) * lda + p));
3867                        c0a = vfmaq_f32(c0a, a0, bv0);
3868                        c0b = vfmaq_f32(c0b, a0, bv1);
3869                        c1a = vfmaq_f32(c1a, a1, bv0);
3870                        c1b = vfmaq_f32(c1b, a1, bv1);
3871                        c2a = vfmaq_f32(c2a, a2, bv0);
3872                        c2b = vfmaq_f32(c2b, a2, bv1);
3873                        c3a = vfmaq_f32(c3a, a3, bv0);
3874                        c3b = vfmaq_f32(c3b, a3, bv1);
3875                    }
3876                    let al = vdupq_n_f32(alpha);
3877                    for (r, (ca, cb)) in [(c0a, c0b), (c1a, c1b), (c2a, c2b), (c3a, c3b)]
3878                        .iter()
3879                        .enumerate()
3880                    {
3881                        let dst = c.as_mut_ptr().add((i + r) * ldc + j);
3882                        vst1q_f32(dst, vmulq_f32(*ca, al));
3883                        vst1q_f32(dst.add(4), vmulq_f32(*cb, al));
3884                    }
3885                } else {
3886                    for r in 0..mi {
3887                        for q in 0..nj {
3888                            let mut acc = 0f32;
3889                            for p in 0..k {
3890                                let bv = if b_rows_are_n {
3891                                    b_mat[(j + q) * ldb + p]
3892                                } else {
3893                                    b_mat[p * ldb + j + q]
3894                                };
3895                                acc += a[(i + r) * lda + p] * bv;
3896                            }
3897                            c[(i + r) * ldc + j + q] = acc * alpha;
3898                        }
3899                    }
3900                }
3901                j += nj;
3902            }
3903            i += mi;
3904        }
3905    }
3906}
3907
3908/// Off-macOS aarch64: the batched attention rides the NEON micro-GEMM.
3909#[cfg(all(target_arch = "aarch64", not(target_os = "macos")))]
3910#[allow(clippy::too_many_arguments)]
3911pub(crate) fn sgemm_rm(
3912    m: usize,
3913    n: usize,
3914    k: usize,
3915    alpha: f32,
3916    a: &[f32],
3917    lda: usize,
3918    b_mat: &[f32],
3919    ldb: usize,
3920    b_rows_are_n: bool,
3921    c: &mut [f32],
3922    ldc: usize,
3923) {
3924    neon_gemm_rm(m, n, k, alpha, a, lda, b_mat, ldb, b_rows_are_n, c, ldc);
3925}
3926
3927/// Row-major f32 GEMM, exposed for offline tools (the AWNP pass builds a
3928/// per-layer projection and applies it to every expert; a naive triple loop
3929/// would turn a two-minute job into half an hour).
3930#[allow(clippy::too_many_arguments)]
3931pub fn sgemm_public(
3932    m: usize,
3933    n: usize,
3934    k: usize,
3935    alpha: f32,
3936    a: &[f32],
3937    lda: usize,
3938    b_mat: &[f32],
3939    ldb: usize,
3940    b_rows_are_n: bool,
3941    c: &mut [f32],
3942    ldc: usize,
3943) {
3944    #[cfg(any(target_os = "macos", target_arch = "aarch64"))]
3945    {
3946        sgemm_rm(m, n, k, alpha, a, lda, b_mat, ldb, b_rows_are_n, c, ldc);
3947    }
3948    // x86 without Accelerate has no sgemm_rm: the specialized paths there are
3949    // quantized kernels, not an f32 GEMM. Only the offline AWNP pass reaches
3950    // this, so correctness matters and throughput does not — a triple loop is
3951    // the honest fallback rather than a reason to make the tool macOS-only.
3952    #[cfg(not(any(target_os = "macos", target_arch = "aarch64")))]
3953    {
3954        for i in 0..m {
3955            for j in 0..n {
3956                let mut acc = 0f32;
3957                for p in 0..k {
3958                    let bv = if b_rows_are_n {
3959                        b_mat[j * ldb + p]
3960                    } else {
3961                        b_mat[p * ldb + j]
3962                    };
3963                    acc += a[i * lda + p] * bv;
3964                }
3965                c[i * ldc + j] = alpha * acc;
3966            }
3967        }
3968    }
3969}
3970
3971/// Row-major f32 GEMM on Accelerate: C[m,n] = alpha·A[m,k] × B(ᵀ).
3972/// `b_rows_are_n` = true multiplies by Bᵀ where B is stored [n, k].
3973#[cfg(target_os = "macos")]
3974#[allow(clippy::too_many_arguments)]
3975pub(crate) fn sgemm_rm(
3976    m: usize,
3977    n: usize,
3978    k: usize,
3979    alpha: f32,
3980    a: &[f32],
3981    lda: usize,
3982    b_mat: &[f32],
3983    ldb: usize,
3984    b_rows_are_n: bool,
3985    c: &mut [f32],
3986    ldc: usize,
3987) {
3988    debug_assert!(a.len() >= (m - 1) * lda + k);
3989    debug_assert!(c.len() >= (m - 1) * ldc + n);
3990    // Test hook: route the attention GEMMs through the portable NEON
3991    // micro-kernel ON APPLE SILICON — how the mobile batched attend is
3992    // measured without a phone in the loop. (Intel macOS has no NEON —
3993    // the hook is a no-op there, Accelerate continues below.)
3994    #[cfg(target_arch = "aarch64")]
3995    if std::env::var("CMF_FORCE_NEON_GEMM")
3996        .map(|v| v == "1")
3997        .unwrap_or(false)
3998    {
3999        return neon_gemm_rm(m, n, k, alpha, a, lda, b_mat, ldb, b_rows_are_n, c, ldc);
4000    }
4001    unsafe {
4002        accel_blas::cblas_sgemm(
4003            101, // RowMajor
4004            111, // NoTrans A
4005            if b_rows_are_n { 112 } else { 111 },
4006            m as i32,
4007            n as i32,
4008            k as i32,
4009            alpha,
4010            a.as_ptr(),
4011            lda as i32,
4012            b_mat.as_ptr(),
4013            ldb as i32,
4014            0.0,
4015            c.as_mut_ptr(),
4016            ldc as i32,
4017        );
4018    }
4019}
4020
4021/// Prefill GEMM through Accelerate (macOS): dequantize q8 rows into
4022/// f32 tiles (scale folded in, pool-parallel) and multiply each tile
4023/// on the AMX with one row-major sgemm. Tiles live in cache, weights
4024/// stream once. Numerics are f32-GEMM (not the int8 dot): prefill
4025/// logits shift within f32 rounding — tolerance-class, like every
4026/// reduction-order change; decode (M=1) never takes this path.
4027#[cfg(target_os = "macos")]
4028fn qmatmat_accel(
4029    q: &[u8],
4030    row_scale: &[f32],
4031    pre: &[std::borrow::Cow<'_, [f32]>],
4032    rows: usize,
4033    cols: usize,
4034    out: &mut [f32],
4035    pool: Option<&Pool>,
4036) {
4037    // NOTE: double-buffering the dequant against the sgemm (a scoped
4038    // thread driving the pool on tile k+1 while the caller multiplies
4039    // tile k) was tried and LOST ~6%: Accelerate's sgemm is itself
4040    // multithreaded, and the dequant workers just steal its cores.
4041    const TR: usize = 2048;
4042    let b = pre.len();
4043    thread_local! {
4044        static XPANEL: std::cell::RefCell<Vec<f32>> = const { std::cell::RefCell::new(Vec::new()) };
4045        static WTILE: std::cell::RefCell<Vec<f32>> = const { std::cell::RefCell::new(Vec::new()) };
4046    }
4047    XPANEL.with(|xp| {
4048        WTILE.with(|wt| {
4049            let mut xpanel = xp.borrow_mut();
4050            xpanel.clear();
4051            for x in pre {
4052                xpanel.extend_from_slice(x);
4053            }
4054            let mut wtile = wt.borrow_mut();
4055            wtile.resize(TR * cols, 0.0);
4056            let mut r0 = 0usize;
4057            while r0 < rows {
4058                let tr = TR.min(rows - r0);
4059                // Dequant the tile (scale folded) — pool-parallel.
4060                let wt_addr = SendMut(wtile.as_mut_ptr());
4061                let run = |start: usize, end: usize| {
4062                    for r in start..end {
4063                        let row = &q[(r0 + r) * cols..(r0 + r + 1) * cols];
4064                        let s = row_scale[r0 + r];
4065                        // SAFETY: workers cover disjoint r ranges.
4066                        let dst =
4067                            unsafe { std::slice::from_raw_parts_mut(wt_addr.at(r * cols), cols) };
4068                        for (d, &v) in dst.iter_mut().zip(row) {
4069                            *d = (v as i8) as f32 * s;
4070                        }
4071                    }
4072                };
4073                dispatch_rows(pool, tr, &run);
4074                // C[b, tr] (at column r0 of out[b, rows]) = X · Wtileᵀ
4075                unsafe {
4076                    accel_blas::cblas_sgemm(
4077                        101, // RowMajor
4078                        111, // NoTrans A
4079                        112, // Trans B
4080                        b as i32,
4081                        tr as i32,
4082                        cols as i32,
4083                        1.0,
4084                        xpanel.as_ptr(),
4085                        cols as i32,
4086                        wtile.as_ptr(),
4087                        cols as i32,
4088                        0.0,
4089                        out.as_mut_ptr().add(r0),
4090                        rows as i32,
4091                    );
4092                }
4093                r0 += tr;
4094            }
4095        })
4096    });
4097}
4098
4099fn qmatmat(
4100    q: &[u8],
4101    row_scale: &[f32],
4102    pre: &[std::borrow::Cow<'_, [f32]>],
4103    rows: usize,
4104    cols: usize,
4105    out: &mut [f32],
4106    pool: Option<&Pool>,
4107) {
4108    let b = pre.len();
4109    debug_assert_eq!(out.len(), b * rows);
4110    // Big prefill batches ride the AMX (roadmap PR3): the row×batch
4111    // SDOT loop below peaks near the CPU's dot throughput, an order
4112    // below the matrix units. Small tensors and tiny test models stay
4113    // on the exact integer path.
4114    #[cfg(target_os = "macos")]
4115    if b >= 8 && rows * cols >= 500_000 && accel_gemm_enabled() {
4116        qmatmat_accel(q, row_scale, pre, rows, cols, out, pool);
4117        return;
4118    }
4119    #[cfg(target_arch = "aarch64")]
4120    if sdot_enabled() {
4121        let acts: Vec<SplitAct> = pre.iter().map(|x| split_act(x)).collect();
4122        let out_addr = SendMut(out.as_mut_ptr());
4123        // Blocked 2×4 (mobile prefill: no AMX to fall back on — this
4124        // path IS the ARM prefill GEMM off Apple silicon).
4125        let blocked_ok = blocked_enabled();
4126        let use_i8mm = i8mm_enabled();
4127        if blocked_ok {
4128            let run = |start: usize, end: usize| {
4129                let mut o = start;
4130                while o < end {
4131                    if o + 2 <= end {
4132                        let r0 = &q[o * cols..(o + 1) * cols];
4133                        let r1 = &q[(o + 1) * cols..(o + 2) * cols];
4134                        let mut bi = 0usize;
4135                        while bi + 4 <= acts.len() {
4136                            let xs = [
4137                                acts[bi].xq.as_slice(),
4138                                acts[bi + 1].xq.as_slice(),
4139                                acts[bi + 2].xq.as_slice(),
4140                                acts[bi + 3].xq.as_slice(),
4141                            ];
4142                            let d = if use_i8mm {
4143                                unsafe { dot_i8_smmla_2x4(r0, r1, xs) }
4144                            } else {
4145                                unsafe { dot_i8_sdot_2x4(r0, r1, xs) }
4146                            };
4147                            for (r, row) in [r0, r1].into_iter().enumerate() {
4148                                for k in 0..4 {
4149                                    let act = &acts[bi + k];
4150                                    let mut v = d[r][k] as f32 * act.sx;
4151                                    for &(j, xv) in &act.outliers {
4152                                        v += (row[j] as i8) as f32 * xv;
4153                                    }
4154                                    unsafe {
4155                                        *out_addr.at((bi + k) * rows + o + r) = v * row_scale[o + r]
4156                                    };
4157                                }
4158                            }
4159                            bi += 4;
4160                        }
4161                        while bi < acts.len() {
4162                            for (r, row) in [r0, r1].into_iter().enumerate() {
4163                                let v = row_dot_sdot(row, &acts[bi]) * row_scale[o + r];
4164                                unsafe { *out_addr.at(bi * rows + o + r) = v };
4165                            }
4166                            bi += 1;
4167                        }
4168                        o += 2;
4169                    } else {
4170                        let row = &q[o * cols..(o + 1) * cols];
4171                        for (bi, act) in acts.iter().enumerate() {
4172                            let v = row_dot_sdot(row, act) * row_scale[o];
4173                            unsafe { *out_addr.at(bi * rows + o) = v };
4174                        }
4175                        o += 1;
4176                    }
4177                }
4178            };
4179            dispatch_rows(pool, rows, &run);
4180            return;
4181        }
4182        let run = |start: usize, end: usize| {
4183            for o in start..end {
4184                let row = &q[o * cols..(o + 1) * cols];
4185                for (bi, act) in acts.iter().enumerate() {
4186                    let v = row_dot_sdot(row, act) * row_scale[o];
4187                    unsafe { *out_addr.at(bi * rows + o) = v };
4188                }
4189            }
4190        };
4191        dispatch_rows(pool, rows, &run);
4192        return;
4193    }
4194    // x86 A8W8 batch. Non-VNNI parts take the BLOCKED 2×4 kernel
4195    // (roadmap P0: two weight rows' abs() stay in registers across four
4196    // activation streams); VNNI machines keep the per-row bias-trick
4197    // dot, which is already throughput-bound there.
4198    #[cfg(target_arch = "x86_64")]
4199    if avx2_a8w8_enabled() {
4200        let acts: Vec<SplitAct> = pre.iter().map(|x| split_act(x)).collect();
4201        let out_addr = SendMut(out.as_mut_ptr());
4202        // CMF_X86_BLOCKED=0 forces the per-row path (paired in-process
4203        // A/B on noisy shared-vCPU hosts).
4204        let blocked_ok = blocked_enabled();
4205        if !avx512vnni_enabled() && blocked_ok && !row_exact() {
4206            let run = |start: usize, end: usize| {
4207                let mut o = start;
4208                while o < end {
4209                    if o + 2 <= end {
4210                        let r0 = &q[o * cols..(o + 1) * cols];
4211                        let r1 = &q[(o + 1) * cols..(o + 2) * cols];
4212                        let mut bi = 0usize;
4213                        while bi + 4 <= acts.len() {
4214                            let xs = [
4215                                acts[bi].xq.as_slice(),
4216                                acts[bi + 1].xq.as_slice(),
4217                                acts[bi + 2].xq.as_slice(),
4218                                acts[bi + 3].xq.as_slice(),
4219                            ];
4220                            let d = unsafe { dot_i8_i8_avx2_2x4(r0, r1, xs) };
4221                            for (r, row) in [r0, r1].into_iter().enumerate() {
4222                                for k in 0..4 {
4223                                    let act = &acts[bi + k];
4224                                    let mut v = d[r][k] as f32 * act.sx;
4225                                    for &(j, xv) in &act.outliers {
4226                                        v += (row[j] as i8) as f32 * xv;
4227                                    }
4228                                    unsafe {
4229                                        *out_addr.at((bi + k) * rows + o + r) = v * row_scale[o + r]
4230                                    };
4231                                }
4232                            }
4233                            bi += 4;
4234                        }
4235                        while bi < acts.len() {
4236                            for (r, row) in [r0, r1].into_iter().enumerate() {
4237                                let v = row_dot_avx2(row, &acts[bi]) * row_scale[o + r];
4238                                unsafe { *out_addr.at(bi * rows + o + r) = v };
4239                            }
4240                            bi += 1;
4241                        }
4242                        o += 2;
4243                    } else {
4244                        let row = &q[o * cols..(o + 1) * cols];
4245                        for (bi, act) in acts.iter().enumerate() {
4246                            let v = row_dot_avx2(row, act) * row_scale[o];
4247                            unsafe { *out_addr.at(bi * rows + o) = v };
4248                        }
4249                        o += 1;
4250                    }
4251                }
4252            };
4253            dispatch_rows(pool, rows, &run);
4254            return;
4255        }
4256        let run = |start: usize, end: usize| {
4257            for o in start..end {
4258                let row = &q[o * cols..(o + 1) * cols];
4259                for (bi, act) in acts.iter().enumerate() {
4260                    let v = row_dot_avx2(row, act) * row_scale[o];
4261                    unsafe { *out_addr.at(bi * rows + o) = v };
4262                }
4263            }
4264        };
4265        dispatch_rows(pool, rows, &run);
4266        return;
4267    }
4268    let out_addr = SendMut(out.as_mut_ptr());
4269    let run = |start: usize, end: usize| {
4270        for o in start..end {
4271            let row = &q[o * cols..(o + 1) * cols];
4272            for (bi, x) in pre.iter().enumerate() {
4273                let mut acc = 0f32;
4274                for j in 0..cols {
4275                    acc += (row[j] as i8) as f32 * x[j];
4276                }
4277                unsafe { *out_addr.at(bi * rows + o) = acc * row_scale[o] };
4278            }
4279        }
4280    };
4281    dispatch_rows(pool, rows, &run);
4282}
4283
4284/// Split rows across pool workers (shared qmatvec pattern). Self-balancing
4285/// — see `Pool::run_rows` for why a static 1/n split is wrong here.
4286fn dispatch_rows(pool: Option<&Pool>, rows: usize, run: &(dyn Fn(usize, usize) + Sync)) {
4287    match pool {
4288        Some(pool) if rows >= 256 => pool.run_rows(rows, run),
4289        _ => run(0, rows),
4290    }
4291}
4292
4293/// Split a q4_block blob into (packed nibbles, f16 group scales).
4294fn q4_split(bytes: &[u8], rows: usize, cols: usize) -> (&[u8], &[u8]) {
4295    let groups = rows * cols / GROUP_SIZE;
4296    bytes.split_at(groups * 16)
4297}
4298
4299/// SIMD unpack for the dominant vbit width B=4 (94% of rows on the
4300/// log2-shape calibration): 16 packed bytes -> 32 centered i8 values.
4301/// vbit packs MSB-first, so the HIGH nibble is the even element
4302/// (opposite of q4_block's lo-first interleave). Centering is u-7.
4303#[inline]
4304fn vbit_fill4(data: &[u8], buf: &mut [u8]) {
4305    #[cfg(target_arch = "aarch64")]
4306    unsafe {
4307        return vbit_fill4_neon(data, buf);
4308    }
4309    #[cfg(target_arch = "x86_64")]
4310    if avx2_enabled() {
4311        return unsafe { vbit_fill4_avx2(data, buf) };
4312    }
4313    #[allow(unreachable_code)]
4314    for (blk, chunk) in buf.chunks_exact_mut(8).enumerate() {
4315        let u = unpack8::<4>(&data[blk * 4..]);
4316        for k in 0..8 {
4317            chunk[k] = (u[k] - 7) as i8 as u8;
4318        }
4319    }
4320}
4321
4322#[cfg(target_arch = "aarch64")]
4323#[target_feature(enable = "neon")]
4324unsafe fn vbit_fill4_neon(data: &[u8], buf: &mut [u8]) {
4325    // SAFETY: buf.len() is a multiple of GROUP_SIZE=32; data holds
4326    // buf.len()/2 packed bytes (validated at load).
4327    unsafe {
4328        use core::arch::aarch64::*;
4329        let n = buf.len();
4330        let mask = vdupq_n_u8(0x0F);
4331        let seven = vdupq_n_s8(7);
4332        let mut g = 0usize;
4333        while g * 32 + 32 <= n {
4334            let b = vld1q_u8(data.as_ptr().add(g * 16));
4335            let hi = vshrq_n_u8::<4>(b);
4336            let lo = vandq_u8(b, mask);
4337            let z0 = vsubq_s8(vreinterpretq_s8_u8(vzip1q_u8(hi, lo)), seven);
4338            let z1 = vsubq_s8(vreinterpretq_s8_u8(vzip2q_u8(hi, lo)), seven);
4339            vst1q_u8(buf.as_mut_ptr().add(g * 32), vreinterpretq_u8_s8(z0));
4340            vst1q_u8(buf.as_mut_ptr().add(g * 32 + 16), vreinterpretq_u8_s8(z1));
4341            g += 1;
4342        }
4343    }
4344}
4345
4346#[cfg(target_arch = "x86_64")]
4347#[target_feature(enable = "avx2")]
4348unsafe fn vbit_fill4_avx2(data: &[u8], buf: &mut [u8]) {
4349    // SAFETY: see vbit_fill4_neon.
4350    unsafe {
4351        use core::arch::x86_64::*;
4352        let n = buf.len();
4353        let mask = _mm_set1_epi8(0x0F);
4354        let seven = _mm256_set1_epi8(7);
4355        let mut g = 0usize;
4356        while g * 32 + 32 <= n {
4357            let b = _mm_loadu_si128(data.as_ptr().add(g * 16) as *const __m128i);
4358            let hi = _mm_and_si128(_mm_srli_epi16::<4>(b), mask);
4359            let lo = _mm_and_si128(b, mask);
4360            let z = _mm256_sub_epi8(
4361                _mm256_set_m128i(_mm_unpackhi_epi8(hi, lo), _mm_unpacklo_epi8(hi, lo)),
4362                seven,
4363            );
4364            _mm256_storeu_si256(buf.as_mut_ptr().add(g * 32) as *mut __m256i, z);
4365            g += 1;
4366        }
4367    }
4368}
4369
4370/// Unpack 8 MSB-first B-bit values from exactly B bytes (fixed shifts —
4371/// no serial bit-buffer, auto-vectorizable). Every 32-value group starts
4372/// byte-aligned (32·B/8 is integral for B∈3..8), so groups decompose
4373/// into 4 such blocks.
4374#[inline(always)]
4375fn unpack8<const B: usize>(data: &[u8]) -> [i32; 8] {
4376    let mut acc = 0u64;
4377    for i in 0..B {
4378        acc = (acc << 8) | data[i] as u64;
4379    }
4380    let mask = (1u64 << B) - 1;
4381    let mut out = [0i32; 8];
4382    for (k, o) in out.iter_mut().enumerate() {
4383        *o = ((acc >> ((7 - k) * B)) & mask) as i32;
4384    }
4385    out
4386}
4387
4388/// Fused vbit matvec straight from the mapped bytes (spec §3, P13
4389/// FIG.3): [u8 bits: rows][f16 scales: rows·cols/32][bit-packed rows,
4390/// MSB-first, byte-padded]. Row data offsets are precomputed at load
4391/// (`vbit_row_offsets`) — the per-call prefix scan was O(rows) pure
4392/// overhead on every matvec.
4393#[allow(clippy::too_many_arguments)]
4394fn vbitmatvec(
4395    bytes: &[u8],
4396    offsets: &[usize],
4397    x: &[f32],
4398    rows: usize,
4399    cols: usize,
4400    out: &mut [f32],
4401    pool: Option<&Pool>,
4402) {
4403    debug_assert_eq!(out.len(), rows);
4404    debug_assert_eq!(offsets.len(), rows + 1);
4405
4406    // SDOT path: unpack the row to centered i8 once, then per-group
4407    // int8 dot against the quantized activations — same A8W8 contract
4408    // as q8 (bounded noise; CMF_SDOT=0 keeps the exact scalar path).
4409    if a8w8_enabled() {
4410        let act = split_act(x);
4411        let out_addr = SendMut(out.as_mut_ptr());
4412        let run = move |start: usize, end: usize| {
4413            vbit_range_a8w8(bytes, offsets, x, &act, rows, cols, out_addr, start, end)
4414        };
4415        dispatch_rows(pool, rows, &run);
4416        return;
4417    }
4418
4419    let out_addr = SendMut(out.as_mut_ptr());
4420    let run = move |start: usize, end: usize| {
4421        vbit_range_f32(bytes, offsets, x, rows, cols, out_addr, start, end)
4422    };
4423    dispatch_rows(pool, rows, &run);
4424}
4425
4426/// One vbit row range via the A8W8 int8 path — kernel body of
4427/// `vbitmatvec`, extracted so multi-matrix jobs can drive it for
4428/// several tensors in one dispatch (b=8 rows go exact f32).
4429#[allow(clippy::too_many_arguments)]
4430fn vbit_range_a8w8(
4431    bytes: &[u8],
4432    offsets: &[usize],
4433    x: &[f32],
4434    act: &SplitAct,
4435    rows: usize,
4436    cols: usize,
4437    out: SendMut,
4438    start: usize,
4439    end: usize,
4440) {
4441    let ng = cols / GROUP_SIZE;
4442    let bits = &bytes[..rows];
4443    let sc_off = rows;
4444    let row_dot = |r: usize| -> f32 {
4445        let b = bits[r] as usize;
4446        let l = (1i32 << (b - 1)) - 1;
4447        let mask = (1u64 << b) - 1;
4448        let data = &bytes[offsets[r]..offsets[r + 1]];
4449        if b == 8 {
4450            // u−L reaches 128 → does not fit i8; exact f32 path.
4451            let (mut acc, mut nbits, mut idx) = (0u64, 0usize, 0usize);
4452            let mut dot = 0f32;
4453            for g in 0..ng {
4454                let so = (r * ng + g) * 2;
4455                let sgf = f16_to_f32(u16::from_le_bytes([
4456                    bytes[sc_off + so],
4457                    bytes[sc_off + so + 1],
4458                ]));
4459                let xg = &x[g * GROUP_SIZE..(g + 1) * GROUP_SIZE];
4460                let mut gd = 0f32;
4461                for &xv in xg.iter() {
4462                    if nbits < 8 {
4463                        acc = (acc << 8) | data[idx] as u64;
4464                        idx += 1;
4465                        nbits += 8;
4466                    }
4467                    let u = ((acc >> (nbits - 8)) & 0xFF) as i32;
4468                    nbits -= 8;
4469                    gd += (u - l) as f32 * xv;
4470                }
4471                dot += gd * sgf;
4472            }
4473            return dot;
4474        }
4475        // Per-worker scratch: this closure runs for every row of the
4476        // tensor (lm_head ≈ 150k rows/token) — a heap allocation per
4477        // row was measurable pure overhead.
4478        thread_local! {
4479            static VBIT_SCRATCH: std::cell::RefCell<Vec<u8>> =
4480                const { std::cell::RefCell::new(Vec::new()) };
4481        }
4482        #[inline(always)]
4483        fn fill<const B: usize>(data: &[u8], l: i32, buf: &mut [u8]) {
4484            for (blk, chunk) in buf.chunks_exact_mut(8).enumerate() {
4485                let u = unpack8::<B>(&data[blk * B..]);
4486                for k in 0..8 {
4487                    chunk[k] = (u[k] - l) as i8 as u8;
4488                }
4489            }
4490        }
4491        let _ = mask;
4492        VBIT_SCRATCH.with(|scratch| {
4493            let mut buf = scratch.borrow_mut();
4494            buf.resize(cols, 0);
4495            match b {
4496                3 => fill::<3>(data, l, &mut buf),
4497                4 => vbit_fill4(data, &mut buf),
4498                5 => fill::<5>(data, l, &mut buf),
4499                6 => fill::<6>(data, l, &mut buf),
4500                _ => unreachable!(),
4501            }
4502            let mut dot = 0f32;
4503            for g in 0..ng {
4504                let so = (r * ng + g) * 2;
4505                let s = f16_to_f32(u16::from_le_bytes([
4506                    bytes[sc_off + so],
4507                    bytes[sc_off + so + 1],
4508                ]));
4509                let d = dot_i8_i8(
4510                    &buf[g * GROUP_SIZE..(g + 1) * GROUP_SIZE],
4511                    &act.xq[g * GROUP_SIZE..(g + 1) * GROUP_SIZE],
4512                ) as f32
4513                    * act.sx;
4514                dot += d * s;
4515            }
4516            for &(j, xv) in &act.outliers {
4517                let so = (r * ng + j / GROUP_SIZE) * 2;
4518                let s = f16_to_f32(u16::from_le_bytes([
4519                    bytes[sc_off + so],
4520                    bytes[sc_off + so + 1],
4521                ]));
4522                // xq is zeroed at outlier slots — add the exact term.
4523                dot += (buf[j] as i8) as f32 * s * xv;
4524            }
4525            dot
4526        })
4527    };
4528    for r in start..end {
4529        // SAFETY: disjoint row ranges per worker.
4530        unsafe { *out.at(r) = row_dot(r) };
4531    }
4532}
4533
4534/// Exact scalar vbit row range (same extraction, non-SDOT path).
4535#[allow(clippy::too_many_arguments)]
4536fn vbit_range_f32(
4537    bytes: &[u8],
4538    offsets: &[usize],
4539    x: &[f32],
4540    rows: usize,
4541    cols: usize,
4542    out: SendMut,
4543    start: usize,
4544    end: usize,
4545) {
4546    let ng = cols / GROUP_SIZE;
4547    let bits = &bytes[..rows];
4548    let sc_off = rows;
4549    // Per-bit-width specialized inner loops: the compiler unrolls the
4550    // constant shifts (the generic bit-buffer loop was branch-bound —
4551    // 5.6 vs 13.2 tok/s q4 on the 0.8B).
4552    #[inline(always)]
4553    fn dot_row<const B: usize>(
4554        data: &[u8],
4555        bytes: &[u8],
4556        sc_off: usize,
4557        r: usize,
4558        ng: usize,
4559        x: &[f32],
4560    ) -> f32 {
4561        let l = ((1i32 << (B - 1)) - 1) as f32;
4562        let gbytes = GROUP_SIZE * B / 8;
4563        let mut dot = 0f32;
4564        for g in 0..ng {
4565            let so = (r * ng + g) * 2;
4566            let s = f16_to_f32(u16::from_le_bytes([
4567                bytes[sc_off + so],
4568                bytes[sc_off + so + 1],
4569            ]));
4570            let xg = &x[g * GROUP_SIZE..(g + 1) * GROUP_SIZE];
4571            let gd0 = &data[g * gbytes..(g + 1) * gbytes];
4572            let mut gd = 0f32;
4573            for blk in 0..GROUP_SIZE / 8 {
4574                let u = unpack8::<B>(&gd0[blk * B..]);
4575                let xb = &xg[blk * 8..blk * 8 + 8];
4576                for k in 0..8 {
4577                    gd += (u[k] as f32 - l) * xb[k];
4578                }
4579            }
4580            dot += gd * s;
4581        }
4582        dot
4583    }
4584    for r in start..end {
4585        let data = &bytes[offsets[r]..offsets[r + 1]];
4586        let v = match bits[r] {
4587            3 => dot_row::<3>(data, bytes, sc_off, r, ng, x),
4588            4 => dot_row::<4>(data, bytes, sc_off, r, ng, x),
4589            5 => dot_row::<5>(data, bytes, sc_off, r, ng, x),
4590            6 => dot_row::<6>(data, bytes, sc_off, r, ng, x),
4591            8 => dot_row::<8>(data, bytes, sc_off, r, ng, x),
4592            b => unreachable!("vbit bit-width {b} (validated at load)"),
4593        };
4594        // SAFETY: disjoint row ranges per worker.
4595        unsafe { *out.at(r) = v };
4596    }
4597}
4598
4599/// Fused two-input vbit matvec: each row is unpacked from the mmap ONCE
4600/// and dotted against BOTH activations (MTP verify / pair prefill used
4601/// to run two full matvecs — double weight traffic and double unpack).
4602/// Per-input math is identical to `vbitmatvec` → same accuracy contract.
4603#[allow(clippy::too_many_arguments)]
4604fn vbitmatvec2(
4605    bytes: &[u8],
4606    offsets: &[usize],
4607    x1: &[f32],
4608    x2: &[f32],
4609    rows: usize,
4610    cols: usize,
4611    o1: &mut [f32],
4612    o2: &mut [f32],
4613    pool: Option<&Pool>,
4614) {
4615    debug_assert_eq!(o1.len(), rows);
4616    debug_assert_eq!(o2.len(), rows);
4617
4618    if a8w8_enabled() {
4619        let a1 = split_act(x1);
4620        let a2 = split_act(x2);
4621        let p1 = SendMut(o1.as_mut_ptr());
4622        let p2 = SendMut(o2.as_mut_ptr());
4623        let run = move |start: usize, end: usize| {
4624            vbit_range2_a8w8(
4625                bytes, offsets, x1, x2, &a1, &a2, rows, cols, p1, p2, start, end,
4626            )
4627        };
4628        dispatch_rows(pool, rows, &run);
4629        return;
4630    }
4631
4632    let p1 = SendMut(o1.as_mut_ptr());
4633    let p2 = SendMut(o2.as_mut_ptr());
4634    let run = move |start: usize, end: usize| {
4635        vbit_range2_f32(bytes, offsets, x1, x2, rows, cols, p1, p2, start, end)
4636    };
4637    dispatch_rows(pool, rows, &run);
4638}
4639
4640/// Two-input vbit row range via the A8W8 int8 path — kernel body of
4641/// `vbitmatvec2`, extracted for pair multi-matrix jobs (b=8 rows go
4642/// exact f32 for both lanes, bits streamed once).
4643#[allow(clippy::too_many_arguments)]
4644fn vbit_range2_a8w8(
4645    bytes: &[u8],
4646    offsets: &[usize],
4647    x1: &[f32],
4648    x2: &[f32],
4649    a1: &SplitAct,
4650    a2: &SplitAct,
4651    rows: usize,
4652    cols: usize,
4653    p1: SendMut,
4654    p2: SendMut,
4655    start: usize,
4656    end: usize,
4657) {
4658    let ng = cols / GROUP_SIZE;
4659    let bits = &bytes[..rows];
4660    let sc_off = rows;
4661    let row_dots = |r: usize| -> (f32, f32) {
4662        let b = bits[r] as usize;
4663        let l = (1i32 << (b - 1)) - 1;
4664        let data = &bytes[offsets[r]..offsets[r + 1]];
4665        if b == 8 {
4666            // u−L reaches 128 → does not fit i8; exact f32 path,
4667            // bits still streamed once for both lanes.
4668            let (mut acc, mut nbits, mut idx) = (0u64, 0usize, 0usize);
4669            let (mut d1, mut d2) = (0f32, 0f32);
4670            for g in 0..ng {
4671                let so = (r * ng + g) * 2;
4672                let sgf = f16_to_f32(u16::from_le_bytes([
4673                    bytes[sc_off + so],
4674                    bytes[sc_off + so + 1],
4675                ]));
4676                let (mut g1, mut g2) = (0f32, 0f32);
4677                for k in 0..GROUP_SIZE {
4678                    if nbits < 8 {
4679                        acc = (acc << 8) | data[idx] as u64;
4680                        idx += 1;
4681                        nbits += 8;
4682                    }
4683                    let u = ((acc >> (nbits - 8)) & 0xFF) as i32;
4684                    nbits -= 8;
4685                    let w = (u - l) as f32;
4686                    g1 += w * x1[g * GROUP_SIZE + k];
4687                    g2 += w * x2[g * GROUP_SIZE + k];
4688                }
4689                d1 += g1 * sgf;
4690                d2 += g2 * sgf;
4691            }
4692            return (d1, d2);
4693        }
4694        thread_local! {
4695            static VBIT_SCRATCH2: std::cell::RefCell<Vec<u8>> =
4696                const { std::cell::RefCell::new(Vec::new()) };
4697        }
4698        #[inline(always)]
4699        fn fill<const B: usize>(data: &[u8], l: i32, buf: &mut [u8]) {
4700            for (blk, chunk) in buf.chunks_exact_mut(8).enumerate() {
4701                let u = unpack8::<B>(&data[blk * B..]);
4702                for k in 0..8 {
4703                    chunk[k] = (u[k] - l) as i8 as u8;
4704                }
4705            }
4706        }
4707        VBIT_SCRATCH2.with(|scratch| {
4708            let mut buf = scratch.borrow_mut();
4709            buf.resize(cols, 0);
4710            match b {
4711                3 => fill::<3>(data, l, &mut buf),
4712                4 => vbit_fill4(data, &mut buf),
4713                5 => fill::<5>(data, l, &mut buf),
4714                6 => fill::<6>(data, l, &mut buf),
4715                _ => unreachable!(),
4716            }
4717            let (mut d1, mut d2) = (0f32, 0f32);
4718            for g in 0..ng {
4719                let so = (r * ng + g) * 2;
4720                let s = f16_to_f32(u16::from_le_bytes([
4721                    bytes[sc_off + so],
4722                    bytes[sc_off + so + 1],
4723                ]));
4724                let wg = &buf[g * GROUP_SIZE..(g + 1) * GROUP_SIZE];
4725                let v1 = dot_i8_i8(wg, &a1.xq[g * GROUP_SIZE..(g + 1) * GROUP_SIZE]) as f32 * a1.sx;
4726                let v2 = dot_i8_i8(wg, &a2.xq[g * GROUP_SIZE..(g + 1) * GROUP_SIZE]) as f32 * a2.sx;
4727                d1 += v1 * s;
4728                d2 += v2 * s;
4729            }
4730            for &(j, xv) in &a1.outliers {
4731                let so = (r * ng + j / GROUP_SIZE) * 2;
4732                let s = f16_to_f32(u16::from_le_bytes([
4733                    bytes[sc_off + so],
4734                    bytes[sc_off + so + 1],
4735                ]));
4736                d1 += (buf[j] as i8) as f32 * s * xv;
4737            }
4738            for &(j, xv) in &a2.outliers {
4739                let so = (r * ng + j / GROUP_SIZE) * 2;
4740                let s = f16_to_f32(u16::from_le_bytes([
4741                    bytes[sc_off + so],
4742                    bytes[sc_off + so + 1],
4743                ]));
4744                d2 += (buf[j] as i8) as f32 * s * xv;
4745            }
4746            (d1, d2)
4747        })
4748    };
4749    for r in start..end {
4750        let (v1, v2) = row_dots(r);
4751        // SAFETY: disjoint row ranges per worker.
4752        unsafe {
4753            *p1.at(r) = v1;
4754            *p2.at(r) = v2;
4755        }
4756    }
4757}
4758
4759/// Two-input exact scalar vbit row range (same extraction) —
4760/// per-bit-width specialized, two accumulators per row; per-lane
4761/// accumulation order matches `vbitmatvec` exactly.
4762#[allow(clippy::too_many_arguments)]
4763fn vbit_range2_f32(
4764    bytes: &[u8],
4765    offsets: &[usize],
4766    x1: &[f32],
4767    x2: &[f32],
4768    rows: usize,
4769    cols: usize,
4770    p1: SendMut,
4771    p2: SendMut,
4772    start: usize,
4773    end: usize,
4774) {
4775    let ng = cols / GROUP_SIZE;
4776    let bits = &bytes[..rows];
4777    let sc_off = rows;
4778    #[inline(always)]
4779    #[allow(clippy::too_many_arguments)]
4780    fn dot_row2<const B: usize>(
4781        data: &[u8],
4782        bytes: &[u8],
4783        sc_off: usize,
4784        r: usize,
4785        ng: usize,
4786        x1: &[f32],
4787        x2: &[f32],
4788    ) -> (f32, f32) {
4789        let l = ((1i32 << (B - 1)) - 1) as f32;
4790        let gbytes = GROUP_SIZE * B / 8;
4791        let (mut d1, mut d2) = (0f32, 0f32);
4792        for g in 0..ng {
4793            let so = (r * ng + g) * 2;
4794            let s = f16_to_f32(u16::from_le_bytes([
4795                bytes[sc_off + so],
4796                bytes[sc_off + so + 1],
4797            ]));
4798            let x1g = &x1[g * GROUP_SIZE..(g + 1) * GROUP_SIZE];
4799            let x2g = &x2[g * GROUP_SIZE..(g + 1) * GROUP_SIZE];
4800            let gd0 = &data[g * gbytes..(g + 1) * gbytes];
4801            let (mut g1, mut g2) = (0f32, 0f32);
4802            for blk in 0..GROUP_SIZE / 8 {
4803                let u = unpack8::<B>(&gd0[blk * B..]);
4804                for k in 0..8 {
4805                    let w = u[k] as f32 - l;
4806                    g1 += w * x1g[blk * 8 + k];
4807                    g2 += w * x2g[blk * 8 + k];
4808                }
4809            }
4810            d1 += g1 * s;
4811            d2 += g2 * s;
4812        }
4813        (d1, d2)
4814    }
4815    for r in start..end {
4816        let data = &bytes[offsets[r]..offsets[r + 1]];
4817        let (v1, v2) = match bits[r] {
4818            3 => dot_row2::<3>(data, bytes, sc_off, r, ng, x1, x2),
4819            4 => dot_row2::<4>(data, bytes, sc_off, r, ng, x1, x2),
4820            5 => dot_row2::<5>(data, bytes, sc_off, r, ng, x1, x2),
4821            6 => dot_row2::<6>(data, bytes, sc_off, r, ng, x1, x2),
4822            8 => dot_row2::<8>(data, bytes, sc_off, r, ng, x1, x2),
4823            b => unreachable!("vbit bit-width {b} (validated at load)"),
4824        };
4825        // SAFETY: disjoint row ranges per worker.
4826        unsafe {
4827            *p1.at(r) = v1;
4828            *p2.at(r) = v2;
4829        }
4830    }
4831}
4832
4833// ───────────────────── q4_tiled kernels (§4.3) ─────────────────────
4834
4835/// One q4_tiled row dot on the A8W8 int8 path: per 32-group the tile
4836/// is ONE sequential read — [f16 scale][16B nibbles] — versus the two
4837/// distant streams of the split layout. Values/order identical to the
4838/// split kernels.
4839#[inline]
4840#[allow(unreachable_code)]
4841fn dot_q4t_row_i8(bytes: &[u8], r: usize, gpr: usize, xq: &[i8]) -> f32 {
4842    #[cfg(target_arch = "aarch64")]
4843    unsafe {
4844        return dot_q4t_row_sdot(bytes, r, gpr, xq);
4845    }
4846    #[cfg(target_arch = "x86_64")]
4847    unsafe {
4848        if vnni_tiles_enabled() {
4849            return dot_q4t_row_vnni(bytes, r, gpr, xq);
4850        }
4851        return dot_q4t_row_avx2(bytes, r, gpr, xq);
4852    }
4853    let mut acc = 0f32;
4854    for gi in 0..gpr {
4855        let tile = &bytes[(r * gpr + gi) * Q4_TILE..(r * gpr + gi + 1) * Q4_TILE];
4856        let s = f16_to_f32(u16::from_le_bytes([tile[0], tile[1]]));
4857        let mut d = 0i32;
4858        for (k, &b) in tile[2..].iter().enumerate() {
4859            d += ((b & 0x0F) as i32 - 8) * xq[gi * GROUP_SIZE + k * 2] as i32
4860                + (((b >> 4) & 0x0F) as i32 - 8) * xq[gi * GROUP_SIZE + k * 2 + 1] as i32;
4861        }
4862        acc += d as f32 * s;
4863    }
4864    acc
4865}
4866
4867#[cfg(target_arch = "aarch64")]
4868#[target_feature(enable = "neon,dotprod")]
4869unsafe fn dot_q4t_row_sdot(bytes: &[u8], r: usize, gpr: usize, xq: &[i8]) -> f32 {
4870    // SAFETY: callers uphold slice-length contracts (18B tile per group,
4871    // xq.len() == gpr·GROUP_SIZE).
4872    unsafe {
4873        use core::arch::aarch64::*;
4874        use core::arch::asm;
4875        let lomask = vdupq_n_u8(0x0F);
4876        let eight = vdupq_n_s8(8);
4877        let mut acc = 0f32;
4878        for gi in 0..gpr {
4879            let t = bytes.as_ptr().add((r * gpr + gi) * Q4_TILE);
4880            let s = f16_to_f32(u16::from_le_bytes([*t, *t.add(1)]));
4881            let b = vld1q_u8(t.add(2));
4882            let lo = vandq_u8(b, lomask);
4883            let hi = vshrq_n_u8::<4>(b);
4884            let e0 = vsubq_s8(vreinterpretq_s8_u8(vzip1q_u8(lo, hi)), eight);
4885            let e1 = vsubq_s8(vreinterpretq_s8_u8(vzip2q_u8(lo, hi)), eight);
4886            let x0 = vld1q_s8(xq.as_ptr().add(gi * GROUP_SIZE));
4887            let x1 = vld1q_s8(xq.as_ptr().add(gi * GROUP_SIZE + 16));
4888            let (mut a0, mut a1) = (vdupq_n_s32(0), vdupq_n_s32(0));
4889            asm!(
4890                "sdot {a0:v}.4s, {e0:v}.16b, {x0:v}.16b",
4891                "sdot {a1:v}.4s, {e1:v}.16b, {x1:v}.16b",
4892                a0 = inout(vreg) a0, a1 = inout(vreg) a1,
4893                e0 = in(vreg) e0, x0 = in(vreg) x0, e1 = in(vreg) e1, x1 = in(vreg) x1,
4894                options(pure, nomem, nostack),
4895            );
4896            acc += vaddvq_s32(vaddq_s32(a0, a1)) as f32 * s;
4897        }
4898        acc
4899    }
4900}
4901
4902#[cfg(target_arch = "x86_64")]
4903#[target_feature(enable = "avx2")]
4904unsafe fn dot_q4t_row_avx2(bytes: &[u8], r: usize, gpr: usize, xq: &[i8]) -> f32 {
4905    // SAFETY: see dot_q4t_row_sdot.
4906    unsafe {
4907        use core::arch::x86_64::*;
4908        let lomask = _mm_set1_epi8(0x0F);
4909        let eight = _mm256_set1_epi8(8);
4910        let ones = _mm256_set1_epi16(1);
4911        let mut acc = 0f32;
4912        for gi in 0..gpr {
4913            let t = bytes.as_ptr().add((r * gpr + gi) * Q4_TILE);
4914            let s = f16_to_f32(u16::from_le_bytes([*t, *t.add(1)]));
4915            let b = _mm_loadu_si128(t.add(2) as *const __m128i);
4916            let lo = _mm_and_si128(b, lomask);
4917            let hi = _mm_and_si128(_mm_srli_epi16::<4>(b), lomask);
4918            let w = _mm256_sub_epi8(
4919                _mm256_set_m128i(_mm_unpackhi_epi8(lo, hi), _mm_unpacklo_epi8(lo, hi)),
4920                eight,
4921            );
4922            let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
4923            let p16 = _mm256_maddubs_epi16(_mm256_abs_epi8(w), _mm256_sign_epi8(x, w));
4924            let d = _mm256_madd_epi16(p16, ones);
4925            let hi128 = _mm256_extracti128_si256::<1>(d);
4926            let s128 = _mm_add_epi32(_mm256_castsi256_si128(d), hi128);
4927            let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
4928            let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
4929            acc += _mm_cvtsi128_si32(s32) as f32 * s;
4930        }
4931        acc
4932    }
4933}
4934
4935/// VNNI twin of `dot_q4t_row_avx2`: same unpack, `vpdpbusd` replaces
4936/// the maddubs+madd pair (see `dpbusd_hsum` — sums are bit-identical).
4937/// 256-bit VL encoding, so the VEX `vpsignb` stays usable.
4938#[cfg(target_arch = "x86_64")]
4939#[target_feature(enable = "avx2,avx512f,avx512bw,avx512vl,avx512vnni")]
4940unsafe fn dot_q4t_row_vnni(bytes: &[u8], r: usize, gpr: usize, xq: &[i8]) -> f32 {
4941    // SAFETY: see dot_q4t_row_sdot.
4942    unsafe {
4943        use core::arch::x86_64::*;
4944        let lomask = _mm_set1_epi8(0x0F);
4945        let eight = _mm256_set1_epi8(8);
4946        let mut acc = 0f32;
4947        for gi in 0..gpr {
4948            let t = bytes.as_ptr().add((r * gpr + gi) * Q4_TILE);
4949            let s = f16_to_f32(u16::from_le_bytes([*t, *t.add(1)]));
4950            let b = _mm_loadu_si128(t.add(2) as *const __m128i);
4951            let lo = _mm_and_si128(b, lomask);
4952            let hi = _mm_and_si128(_mm_srli_epi16::<4>(b), lomask);
4953            let w = _mm256_sub_epi8(
4954                _mm256_set_m128i(_mm_unpackhi_epi8(lo, hi), _mm_unpacklo_epi8(lo, hi)),
4955                eight,
4956            );
4957            let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
4958            let d = dpbusd_hsum(_mm256_abs_epi8(w), _mm256_sign_epi8(x, w));
4959            acc += d as f32 * s;
4960        }
4961        acc
4962    }
4963}
4964
4965/// One q4_tiled row against FOUR activation streams: the nibble unpack
4966/// and abs() happen once per group instead of once per (group,
4967/// activation) — the unpack is the dominant per-element cost of the
4968/// tiled format (roadmap P0 portable blocking, q4t leg).
4969#[cfg(target_arch = "x86_64")]
4970// `fma` is NOT implied by `avx2`: without it LLVM lowers _mm256_fmadd_ps
4971// to a libm call per lane — measured 2x slower than the reduction this
4972// kernel replaces. The runtime gate (`avx2_enabled`) already requires
4973// both features, so declaring it here is safe.
4974#[target_feature(enable = "avx2,fma")]
4975unsafe fn dot_q4t_row_1x4_avx2(bytes: &[u8], r: usize, gpr: usize, xs: [&[i8]; 4]) -> [f32; 4] {
4976    // SAFETY: callers uphold the 18B-tile and xq-length contracts.
4977    unsafe {
4978        use core::arch::x86_64::*;
4979        let lomask = _mm_set1_epi8(0x0F);
4980        let eight = _mm256_set1_epi8(8);
4981        let ones = _mm256_set1_epi16(1);
4982        // One f32 accumulator VECTOR per activation, reduced once at the
4983        // end. Folding each group's i32 lanes to a scalar inside the loop
4984        // costs an extracti128 + three shift/add + a movd — a cross-lane
4985        // dependency chain per (group, activation), 288 of them per row at
4986        // cols=2304. The per-group scale is what forces a float
4987        // accumulator; it does not force a horizontal sum.
4988        //
4989        // The four accumulators are NAMED, not an array: as `[__m256; 4]`
4990        // indexed by a loop variable LLVM keeps them in memory and every
4991        // group pays four 32-byte loads and stores. That alone made this
4992        // kernel 2x SLOWER than the per-group reduction it replaces
4993        // (measured on the EPYC box: 150 s vs 71 s for two 256² steps).
4994        let mut f0 = _mm256_setzero_ps();
4995        let mut f1 = _mm256_setzero_ps();
4996        let mut f2 = _mm256_setzero_ps();
4997        let mut f3 = _mm256_setzero_ps();
4998        for gi in 0..gpr {
4999            let t = bytes.as_ptr().add((r * gpr + gi) * Q4_TILE);
5000            let s = f16_to_f32(u16::from_le_bytes([*t, *t.add(1)]));
5001            let sv = _mm256_set1_ps(s);
5002            let bb = _mm_loadu_si128(t.add(2) as *const __m128i);
5003            let lo = _mm_and_si128(bb, lomask);
5004            let hi = _mm_and_si128(_mm_srli_epi16::<4>(bb), lomask);
5005            let w = _mm256_sub_epi8(
5006                _mm256_set_m128i(_mm_unpackhi_epi8(lo, hi), _mm_unpacklo_epi8(lo, hi)),
5007                eight,
5008            );
5009            let aw = _mm256_abs_epi8(w);
5010            let off = gi * GROUP_SIZE;
5011            let dot = |xq: &[i8]| {
5012                let x = _mm256_loadu_si256(xq.as_ptr().add(off) as *const __m256i);
5013                let p16 = _mm256_maddubs_epi16(aw, _mm256_sign_epi8(x, w));
5014                _mm256_cvtepi32_ps(_mm256_madd_epi16(p16, ones))
5015            };
5016            f0 = _mm256_fmadd_ps(dot(xs[0]), sv, f0);
5017            f1 = _mm256_fmadd_ps(dot(xs[1]), sv, f1);
5018            f2 = _mm256_fmadd_ps(dot(xs[2]), sv, f2);
5019            f3 = _mm256_fmadd_ps(dot(xs[3]), sv, f3);
5020        }
5021        [
5022            hsum256_ps(f0),
5023            hsum256_ps(f1),
5024            hsum256_ps(f2),
5025            hsum256_ps(f3),
5026        ]
5027    }
5028}
5029
5030/// Horizontal sum of eight f32 lanes — the one cross-lane reduction the
5031/// blocked kernels pay, once per row instead of once per group.
5032#[cfg(target_arch = "x86_64")]
5033#[target_feature(enable = "avx2")]
5034#[inline]
5035unsafe fn hsum256_ps(v: core::arch::x86_64::__m256) -> f32 {
5036    // SAFETY: pure register arithmetic on the caller's vector.
5037    unsafe {
5038        use core::arch::x86_64::*;
5039        let hi = _mm256_extractf128_ps::<1>(v);
5040        let s = _mm_add_ps(_mm256_castps256_ps128(v), hi);
5041        let s = _mm_add_ps(s, _mm_movehl_ps(s, s));
5042        let s = _mm_add_ss(s, _mm_shuffle_ps::<0x55>(s, s));
5043        _mm_cvtss_f32(s)
5044    }
5045}
5046
5047/// VNNI twin of `dot_q4t_row_1x4_avx2` (see `dpbusd_hsum`).
5048#[cfg(target_arch = "x86_64")]
5049#[target_feature(enable = "avx2,fma,avx512f,avx512bw,avx512vl,avx512vnni")]
5050unsafe fn dot_q4t_row_1x4_vnni(bytes: &[u8], r: usize, gpr: usize, xs: [&[i8]; 4]) -> [f32; 4] {
5051    // SAFETY: callers uphold the 18B-tile and xq-length contracts.
5052    unsafe {
5053        use core::arch::x86_64::*;
5054        let lomask = _mm_set1_epi8(0x0F);
5055        let eight = _mm256_set1_epi8(8);
5056        // Same shape as the AVX2 twin: accumulate in f32 vectors and pay
5057        // one cross-lane reduction per row, not per (group, activation).
5058        let mut f0 = _mm256_setzero_ps();
5059        let mut f1 = _mm256_setzero_ps();
5060        let mut f2 = _mm256_setzero_ps();
5061        let mut f3 = _mm256_setzero_ps();
5062        for gi in 0..gpr {
5063            let t = bytes.as_ptr().add((r * gpr + gi) * Q4_TILE);
5064            let s = f16_to_f32(u16::from_le_bytes([*t, *t.add(1)]));
5065            let sv = _mm256_set1_ps(s);
5066            let bb = _mm_loadu_si128(t.add(2) as *const __m128i);
5067            let lo = _mm_and_si128(bb, lomask);
5068            let hi = _mm_and_si128(_mm_srli_epi16::<4>(bb), lomask);
5069            let w = _mm256_sub_epi8(
5070                _mm256_set_m128i(_mm_unpackhi_epi8(lo, hi), _mm_unpacklo_epi8(lo, hi)),
5071                eight,
5072            );
5073            let aw = _mm256_abs_epi8(w);
5074            let off = gi * GROUP_SIZE;
5075            let dot = |xq: &[i8]| {
5076                let x = _mm256_loadu_si256(xq.as_ptr().add(off) as *const __m256i);
5077                _mm256_cvtepi32_ps(_mm256_dpbusd_epi32(
5078                    _mm256_setzero_si256(),
5079                    aw,
5080                    _mm256_sign_epi8(x, w),
5081                ))
5082            };
5083            f0 = _mm256_fmadd_ps(dot(xs[0]), sv, f0);
5084            f1 = _mm256_fmadd_ps(dot(xs[1]), sv, f1);
5085            f2 = _mm256_fmadd_ps(dot(xs[2]), sv, f2);
5086            f3 = _mm256_fmadd_ps(dot(xs[3]), sv, f3);
5087        }
5088        let acc = [
5089            hsum256_ps(f0),
5090            hsum256_ps(f1),
5091            hsum256_ps(f2),
5092            hsum256_ps(f3),
5093        ];
5094        acc
5095    }
5096}
5097
5098/// ARM twin of `dot_q4t_row_1x4_avx2`: one nibble unpack per group
5099/// serves FOUR activation streams. Per stream the group order and f32
5100/// accumulation match `dot_q4t_row_sdot` exactly — batch == matvec
5101/// bit-for-bit.
5102#[cfg(target_arch = "aarch64")]
5103#[target_feature(enable = "neon,dotprod")]
5104unsafe fn dot_q4t_row_1x4_sdot(bytes: &[u8], r: usize, gpr: usize, xs: [&[i8]; 4]) -> [f32; 4] {
5105    // SAFETY: callers uphold the 18B-tile and xq-length contracts.
5106    unsafe {
5107        use core::arch::aarch64::*;
5108        use core::arch::asm;
5109        let lomask = vdupq_n_u8(0x0F);
5110        let eight = vdupq_n_s8(8);
5111        let mut acc = [0f32; 4];
5112        for gi in 0..gpr {
5113            let t = bytes.as_ptr().add((r * gpr + gi) * Q4_TILE);
5114            let s = f16_to_f32(u16::from_le_bytes([*t, *t.add(1)]));
5115            let b = vld1q_u8(t.add(2));
5116            let lo = vandq_u8(b, lomask);
5117            let hi = vshrq_n_u8::<4>(b);
5118            let e0 = vsubq_s8(vreinterpretq_s8_u8(vzip1q_u8(lo, hi)), eight);
5119            let e1 = vsubq_s8(vreinterpretq_s8_u8(vzip2q_u8(lo, hi)), eight);
5120            for (k, xq) in xs.iter().enumerate() {
5121                let x0 = vld1q_s8(xq.as_ptr().add(gi * GROUP_SIZE));
5122                let x1 = vld1q_s8(xq.as_ptr().add(gi * GROUP_SIZE + 16));
5123                let (mut a0, mut a1) = (vdupq_n_s32(0), vdupq_n_s32(0));
5124                asm!(
5125                    "sdot {a0:v}.4s, {e0:v}.16b, {x0:v}.16b",
5126                    "sdot {a1:v}.4s, {e1:v}.16b, {x1:v}.16b",
5127                    a0 = inout(vreg) a0, a1 = inout(vreg) a1,
5128                    e0 = in(vreg) e0, x0 = in(vreg) x0, e1 = in(vreg) e1, x1 = in(vreg) x1,
5129                    options(pure, nomem, nostack),
5130                );
5131                acc[k] += vaddvq_s32(vaddq_s32(a0, a1)) as f32 * s;
5132            }
5133        }
5134        acc
5135    }
5136}
5137
5138/// Exact-term correction for A8W8 outliers on a tiled row.
5139#[inline]
5140fn q4t_outlier(bytes: &[u8], r: usize, gpr: usize, j: usize) -> (f32, f32) {
5141    let gi = j / GROUP_SIZE;
5142    let k = j % GROUP_SIZE;
5143    let tile = &bytes[(r * gpr + gi) * Q4_TILE..(r * gpr + gi + 1) * Q4_TILE];
5144    let s = f16_to_f32(u16::from_le_bytes([tile[0], tile[1]]));
5145    let byte = tile[2 + k / 2];
5146    let nib = if k & 1 == 0 { byte & 0x0F } else { byte >> 4 };
5147    ((nib as i32 - 8) as f32, s)
5148}
5149
5150/// Exact scalar q4_tiled row (CMF_SDOT=0 contract) — same pairwise
5151/// accumulation shape as `q4_range_f32`.
5152#[inline]
5153fn q4t_row_exact(bytes: &[u8], r: usize, gpr: usize, x: &[f32]) -> f32 {
5154    let mut acc = 0f32;
5155    for gi in 0..gpr {
5156        let tile = &bytes[(r * gpr + gi) * Q4_TILE..(r * gpr + gi + 1) * Q4_TILE];
5157        let s = f16_to_f32(u16::from_le_bytes([tile[0], tile[1]]));
5158        let xg = &x[gi * GROUP_SIZE..(gi + 1) * GROUP_SIZE];
5159        let mut ga = 0f32;
5160        for (k, &b) in tile[2..].iter().enumerate() {
5161            ga += ((b & 0x0F) as f32 - 8.0) * xg[k * 2]
5162                + (((b >> 4) & 0x0F) as f32 - 8.0) * xg[k * 2 + 1];
5163        }
5164        acc += ga * s;
5165    }
5166    acc
5167}
5168
5169/// Split view of a `q4tp` payload. The three planes are resolved once per
5170/// matvec instead of per row — `q4tp_sections` is cheap, but doing it inside
5171/// the row loop would put a division on the hot path for nothing.
5172struct Q4tpView<'a> {
5173    nib: &'a [u8],
5174    params: &'a [u8],
5175    codes: &'a [u8],
5176    stride: usize,
5177    /// q2tp reads the ladder with rung 0 = exact zero.
5178    zero_rung: bool,
5179}
5180
5181impl<'a> Q4tpView<'a> {
5182    fn new(bytes: &'a [u8], rows: usize, cols: usize) -> Self {
5183        let (params_off, codes_off, stride) = q4tp_sections(rows, cols);
5184        Self {
5185            nib: &bytes[..params_off],
5186            params: &bytes[params_off..codes_off],
5187            codes: &bytes[codes_off..],
5188            stride,
5189            zero_rung: false,
5190        }
5191    }
5192
5193    /// The q2tp view: identical params/codes planes, 8 B weight chunks.
5194    fn new_q2(bytes: &'a [u8], rows: usize, cols: usize) -> Self {
5195        let (params_off, codes_off, stride) = q2tp_sections(rows, cols);
5196        Self {
5197            nib: &bytes[..params_off],
5198            params: &bytes[params_off..codes_off],
5199            codes: &bytes[codes_off..],
5200            stride,
5201            zero_rung: true,
5202        }
5203    }
5204
5205    /// Expand row `r`'s per-tile scales into `out` (length `gpr`).
5206    ///
5207    /// Doing this once per row — rather than decoding a 5-bit code inside the
5208    /// tile loop — is what makes the format free at runtime. Random access to
5209    /// a packed 5-bit field costs a division, two bounds checks and a branch;
5210    /// the tile's actual work is two `sdot`s, so per-tile decoding dominated
5211    /// the kernel and cost 5x (measured: 1.4 vs 6.9 tok/s on Nanbeige-3B).
5212    /// Walking the plane sequentially with a bit accumulator is ~3 ops.
5213    /// Eight 5-bit codes are exactly five bytes, so a whole group of
5214    /// eight decodes from one little-endian word at fixed shifts. The
5215    /// bit-accumulator this replaces carried a data-dependent `while
5216    /// have < 5` refill whose branch sat in the innermost loop of every
5217    /// q4tp row; a decode profile put this function above the dot
5218    /// products it feeds. Same bitstream, same codes — just no branch
5219    /// and eight independent extractions.
5220    #[inline]
5221    fn scales_into(&self, r: usize, gpr: usize, out: &mut [f32]) {
5222        let tab = if self.zero_rung {
5223            q2tp_ladder(self.params, r)
5224        } else {
5225            q4tp_ladder(self.params, r)
5226        };
5227        let codes = &self.codes[r * self.stride..(r + 1) * self.stride];
5228        let out = &mut out[..gpr];
5229        let mut chunks = out.chunks_exact_mut(8);
5230        let mut ci = 0usize;
5231        for c in &mut chunks {
5232            let w = u64::from(codes[ci])
5233                | u64::from(codes[ci + 1]) << 8
5234                | u64::from(codes[ci + 2]) << 16
5235                | u64::from(codes[ci + 3]) << 24
5236                | u64::from(codes[ci + 4]) << 32;
5237            for (k, o) in c.iter_mut().enumerate() {
5238                *o = tab[((w >> (5 * k)) & 31) as usize];
5239            }
5240            ci += 5;
5241        }
5242        // Fewer than eight codes left: the shared total accessor, which
5243        // tolerates a 5-bit field whose spill byte is past the stride.
5244        let tail = &codes[ci..];
5245        for (k, o) in chunks.into_remainder().iter_mut().enumerate() {
5246            *o = tab[q4tp_code(tail, k)];
5247        }
5248    }
5249}
5250
5251#[inline]
5252fn dot_q4tp_row_i8(nib: &[u8], r: usize, gpr: usize, xq: &[i8], scales: &[f32]) -> f32 {
5253    #[cfg(target_arch = "aarch64")]
5254    unsafe {
5255        return dot_q4tp_row_sdot(nib, r, gpr, xq, scales);
5256    }
5257    #[cfg(target_arch = "x86_64")]
5258    unsafe {
5259        if vnni_tiles_enabled() {
5260            return dot_q4tp_row_vnni(nib, r, gpr, xq, scales);
5261        }
5262        return dot_q4tp_row_avx2(nib, r, gpr, xq, scales);
5263    }
5264    #[allow(unreachable_code)]
5265    {
5266        let mut acc = 0f32;
5267        for gi in 0..gpr {
5268            let tile = &nib[(r * gpr + gi) * Q4TP_NIB..(r * gpr + gi + 1) * Q4TP_NIB];
5269            let s = scales[gi];
5270            let mut d = 0i32;
5271            for (k, &b) in tile.iter().enumerate() {
5272                d += ((b & 0x0F) as i32 - 8) * xq[gi * GROUP_SIZE + k * 2] as i32
5273                    + (((b >> 4) & 0x0F) as i32 - 8) * xq[gi * GROUP_SIZE + k * 2 + 1] as i32;
5274            }
5275            acc += d as f32 * s;
5276        }
5277        acc
5278    }
5279}
5280
5281/// q4tp twin of `dot_q4t_row_sdot`: identical nibble math, but the tile
5282/// stride is 16 B (no inline scale) and the scale is a ladder lookup.
5283#[cfg(target_arch = "aarch64")]
5284#[target_feature(enable = "neon,dotprod")]
5285unsafe fn dot_q4tp_row_sdot(nib: &[u8], r: usize, gpr: usize, xq: &[i8], scales: &[f32]) -> f32 {
5286    // SAFETY: callers uphold slice-length contracts (16B tile per group,
5287    // xq.len() == gpr·GROUP_SIZE, codes covering gpr 5-bit fields).
5288    unsafe {
5289        use core::arch::aarch64::*;
5290        use core::arch::asm;
5291        let lomask = vdupq_n_u8(0x0F);
5292        let eight = vdupq_n_s8(8);
5293        let mut acc = 0f32;
5294        for gi in 0..gpr {
5295            let t = nib.as_ptr().add((r * gpr + gi) * Q4TP_NIB);
5296            let s = *scales.get_unchecked(gi);
5297            let b = vld1q_u8(t);
5298            let lo = vandq_u8(b, lomask);
5299            let hi = vshrq_n_u8::<4>(b);
5300            let e0 = vsubq_s8(vreinterpretq_s8_u8(vzip1q_u8(lo, hi)), eight);
5301            let e1 = vsubq_s8(vreinterpretq_s8_u8(vzip2q_u8(lo, hi)), eight);
5302            let x0 = vld1q_s8(xq.as_ptr().add(gi * GROUP_SIZE));
5303            let x1 = vld1q_s8(xq.as_ptr().add(gi * GROUP_SIZE + 16));
5304            let (mut a0, mut a1) = (vdupq_n_s32(0), vdupq_n_s32(0));
5305            asm!(
5306                "sdot {a0:v}.4s, {e0:v}.16b, {x0:v}.16b",
5307                "sdot {a1:v}.4s, {e1:v}.16b, {x1:v}.16b",
5308                a0 = inout(vreg) a0, a1 = inout(vreg) a1,
5309                e0 = in(vreg) e0, x0 = in(vreg) x0, e1 = in(vreg) e1, x1 = in(vreg) x1,
5310                options(pure, nomem, nostack),
5311            );
5312            acc += vaddvq_s32(vaddq_s32(a0, a1)) as f32 * s;
5313        }
5314        acc
5315    }
5316}
5317
5318#[cfg(target_arch = "x86_64")]
5319#[target_feature(enable = "avx2")]
5320unsafe fn dot_q4tp_row_avx2(nib: &[u8], r: usize, gpr: usize, xq: &[i8], scales: &[f32]) -> f32 {
5321    // SAFETY: see dot_q4tp_row_sdot.
5322    unsafe {
5323        use core::arch::x86_64::*;
5324        let lomask = _mm_set1_epi8(0x0F);
5325        let eight = _mm256_set1_epi8(8);
5326        let ones = _mm256_set1_epi16(1);
5327        let mut acc = 0f32;
5328        for gi in 0..gpr {
5329            let t = nib.as_ptr().add((r * gpr + gi) * Q4TP_NIB);
5330            let s = *scales.get_unchecked(gi);
5331            let b = _mm_loadu_si128(t as *const __m128i);
5332            let lo = _mm_and_si128(b, lomask);
5333            let hi = _mm_and_si128(_mm_srli_epi16::<4>(b), lomask);
5334            let w = _mm256_sub_epi8(
5335                _mm256_set_m128i(_mm_unpackhi_epi8(lo, hi), _mm_unpacklo_epi8(lo, hi)),
5336                eight,
5337            );
5338            let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
5339            let p16 = _mm256_maddubs_epi16(_mm256_abs_epi8(w), _mm256_sign_epi8(x, w));
5340            let d = _mm256_madd_epi16(p16, ones);
5341            let hi128 = _mm256_extracti128_si256::<1>(d);
5342            let s128 = _mm_add_epi32(_mm256_castsi256_si128(d), hi128);
5343            let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
5344            let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
5345            acc += _mm_cvtsi128_si32(s32) as f32 * s;
5346        }
5347        acc
5348    }
5349}
5350
5351/// VNNI twin of `dot_q4tp_row_avx2` (see `dot_q4t_row_vnni` for why the
5352/// 256-bit VL encoding is the one to use here).
5353#[cfg(target_arch = "x86_64")]
5354#[target_feature(enable = "avx2,avx512f,avx512bw,avx512vl,avx512vnni")]
5355unsafe fn dot_q4tp_row_vnni(nib: &[u8], r: usize, gpr: usize, xq: &[i8], scales: &[f32]) -> f32 {
5356    // SAFETY: see dot_q4tp_row_sdot.
5357    unsafe {
5358        use core::arch::x86_64::*;
5359        let lomask = _mm_set1_epi8(0x0F);
5360        let eight = _mm256_set1_epi8(8);
5361        let mut acc = 0f32;
5362        for gi in 0..gpr {
5363            let t = nib.as_ptr().add((r * gpr + gi) * Q4TP_NIB);
5364            let s = *scales.get_unchecked(gi);
5365            let b = _mm_loadu_si128(t as *const __m128i);
5366            let lo = _mm_and_si128(b, lomask);
5367            let hi = _mm_and_si128(_mm_srli_epi16::<4>(b), lomask);
5368            let w = _mm256_sub_epi8(
5369                _mm256_set_m128i(_mm_unpackhi_epi8(lo, hi), _mm_unpacklo_epi8(lo, hi)),
5370                eight,
5371            );
5372            let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
5373            acc += dpbusd_hsum(_mm256_abs_epi8(w), _mm256_sign_epi8(x, w)) as f32 * s;
5374        }
5375        acc
5376    }
5377}
5378
5379/// Exact scalar q4tp row — the `CMF_SDOT=0` contract, same pairwise
5380/// accumulation shape as `q4t_row_exact`.
5381#[inline]
5382fn q4tp_row_exact(nib: &[u8], r: usize, gpr: usize, x: &[f32], scales: &[f32]) -> f32 {
5383    #[cfg(target_arch = "x86_64")]
5384    if avx2_enabled() {
5385        // Keep the scalar pair/group reduction order, not a vector sum.
5386        return unsafe { q4tp_row_float_avx2(nib, r, gpr, x, scales) };
5387    }
5388    q4tp_row_float_scalar(nib, r, gpr, x, scales)
5389}
5390
5391#[inline]
5392fn q4tp_row_float_scalar(nib: &[u8], r: usize, gpr: usize, x: &[f32], scales: &[f32]) -> f32 {
5393    let mut acc = 0f32;
5394    for gi in 0..gpr {
5395        let tile = &nib[(r * gpr + gi) * Q4TP_NIB..(r * gpr + gi + 1) * Q4TP_NIB];
5396        let s = scales[gi];
5397        let xg = &x[gi * GROUP_SIZE..(gi + 1) * GROUP_SIZE];
5398        let mut ga = 0f32;
5399        for (k, &b) in tile.iter().enumerate() {
5400            ga += ((b & 0x0F) as f32 - 8.0) * xg[k * 2]
5401                + (((b >> 4) & 0x0F) as f32 - 8.0) * xg[k * 2 + 1];
5402        }
5403        acc += ga * s;
5404    }
5405    acc
5406}
5407
5408/// Vectorize unpack, conversion and multiplication, but preserve every
5409/// pair addition and the scalar accumulation order. No activation rounding
5410/// or FMA: bit-identical to the float scalar row, including its group scale.
5411#[cfg(target_arch = "x86_64")]
5412#[target_feature(enable = "avx2")]
5413unsafe fn q4tp_row_float_avx2(nib: &[u8], r: usize, gpr: usize, x: &[f32], scales: &[f32]) -> f32 {
5414    // SAFETY: caller checks AVX2 and provides the same complete 32-element
5415    // groups as the scalar row. Loads/stores are explicitly unaligned.
5416    unsafe {
5417        use core::arch::x86_64::*;
5418        let mask = _mm_set1_epi8(15);
5419        let eight = _mm_set1_epi8(8);
5420        let order = _mm256_setr_epi32(0, 1, 4, 5, 2, 3, 6, 7);
5421        let mut acc = 0.0f32;
5422        for gi in 0..gpr {
5423            let packed = _mm_loadu_si128(nib.as_ptr().add((r * gpr + gi) * Q4TP_NIB).cast());
5424            let lo = _mm_and_si128(packed, mask);
5425            let hi = _mm_and_si128(_mm_srli_epi16::<4>(packed), mask);
5426            let w0 = _mm_sub_epi8(_mm_unpacklo_epi8(lo, hi), eight);
5427            let w1 = _mm_sub_epi8(_mm_unpackhi_epi8(lo, hi), eight);
5428            let xp = x.as_ptr().add(gi * GROUP_SIZE);
5429            let a = _mm256_mul_ps(
5430                _mm256_cvtepi32_ps(_mm256_cvtepi8_epi32(w0)),
5431                _mm256_loadu_ps(xp),
5432            );
5433            let b = _mm256_mul_ps(
5434                _mm256_cvtepi32_ps(_mm256_cvtepi8_epi32(_mm_srli_si128::<8>(w0))),
5435                _mm256_loadu_ps(xp.add(8)),
5436            );
5437            let c = _mm256_mul_ps(
5438                _mm256_cvtepi32_ps(_mm256_cvtepi8_epi32(w1)),
5439                _mm256_loadu_ps(xp.add(16)),
5440            );
5441            let d = _mm256_mul_ps(
5442                _mm256_cvtepi32_ps(_mm256_cvtepi8_epi32(_mm_srli_si128::<8>(w1))),
5443                _mm256_loadu_ps(xp.add(24)),
5444            );
5445            let mut pairs = [0.0f32; 16];
5446            _mm256_storeu_ps(
5447                pairs.as_mut_ptr(),
5448                _mm256_permutevar8x32_ps(_mm256_hadd_ps(a, b), order),
5449            );
5450            _mm256_storeu_ps(
5451                pairs.as_mut_ptr().add(8),
5452                _mm256_permutevar8x32_ps(_mm256_hadd_ps(c, d), order),
5453            );
5454            let mut ga = 0.0f32;
5455            for v in pairs {
5456                ga += v;
5457            }
5458            acc += ga * scales[gi];
5459        }
5460        acc
5461    }
5462}
5463
5464/// Single weight of a q4tp tensor — the a8w8 outlier path, which restores
5465/// activation outliers at full precision after the int8 pass.
5466#[inline]
5467fn q4tp_outlier(nib: &[u8], r: usize, gpr: usize, j: usize, scales: &[f32]) -> (f32, f32) {
5468    let (gi, k) = (j / GROUP_SIZE, j % GROUP_SIZE);
5469    let byte = nib[(r * gpr + gi) * Q4TP_NIB + k / 2];
5470    let n = if k & 1 == 0 { byte & 0x0F } else { byte >> 4 };
5471    ((n as i32 - 8) as f32, scales[gi])
5472}
5473
5474/// Fused q4tp matvec (dispatch mirrors `q4t_matvec`).
5475fn q4tp_matvec(
5476    bytes: &[u8],
5477    x: &[f32],
5478    rows: usize,
5479    cols: usize,
5480    out: &mut [f32],
5481    pool: Option<&Pool>,
5482) {
5483    debug_assert_eq!(out.len(), rows);
5484    let gpr = cols / GROUP_SIZE;
5485    let v = Q4tpView::new(bytes, rows, cols);
5486    let out_addr = SendMut(out.as_mut_ptr());
5487    if a8w8_enabled() {
5488        let act = split_act(x);
5489        let run = |start: usize, end: usize| {
5490            // One scratch row of scales per worker — borrowed, not minted.
5491            with_krow(gpr, |sc| {
5492                for r in start..end {
5493                    v.scales_into(r, gpr, sc);
5494                    let mut acc = dot_q4tp_row_i8(v.nib, r, gpr, &act.xq, sc) * act.sx;
5495                    for &(j, xv) in &act.outliers {
5496                        let (w, s) = q4tp_outlier(v.nib, r, gpr, j, sc);
5497                        acc += w * s * xv;
5498                    }
5499                    // SAFETY: disjoint row ranges per worker.
5500                    unsafe { *out_addr.at(r) = acc };
5501                }
5502            })
5503        };
5504        dispatch_rows(pool, rows, &run);
5505        return;
5506    }
5507    let run = |start: usize, end: usize| {
5508        with_krow(gpr, |sc| {
5509            for r in start..end {
5510                v.scales_into(r, gpr, sc);
5511                // SAFETY: disjoint row ranges per worker.
5512                unsafe { *out_addr.at(r) = q4tp_row_exact(v.nib, r, gpr, x, sc) };
5513            }
5514        })
5515    };
5516    dispatch_rows(pool, rows, &run);
5517}
5518
5519/// Fused two-input q4tp matvec — the SwiGLU gate/up pair. Weights and the
5520/// row ladder are read once and spent on both activation streams.
5521#[allow(clippy::too_many_arguments)]
5522fn q4tp_matvec2(
5523    bytes: &[u8],
5524    x1: &[f32],
5525    x2: &[f32],
5526    rows: usize,
5527    cols: usize,
5528    o1: &mut [f32],
5529    o2: &mut [f32],
5530    pool: Option<&Pool>,
5531) {
5532    let gpr = cols / GROUP_SIZE;
5533    let v = Q4tpView::new(bytes, rows, cols);
5534    let (p1, p2) = (SendMut(o1.as_mut_ptr()), SendMut(o2.as_mut_ptr()));
5535    let run = |start: usize, end: usize| {
5536        let mut sc = vec![0f32; gpr];
5537        for r in start..end {
5538            v.scales_into(r, gpr, &mut sc);
5539            // SAFETY: disjoint row ranges per worker.
5540            unsafe {
5541                *p1.at(r) = q4tp_row_exact(v.nib, r, gpr, x1, &sc);
5542                *p2.at(r) = q4tp_row_exact(v.nib, r, gpr, x2, &sc);
5543            }
5544        }
5545    };
5546    dispatch_rows(pool, rows, &run);
5547}
5548
5549/// One q2tp outlier weight at column `j` of row `r`: the 2-bit code and
5550/// its group scale, mirrored on `q4tp_outlier`.
5551#[inline]
5552fn q2tp_outlier(chunks: &[u8], r: usize, gpr: usize, j: usize, scales: &[f32]) -> (f32, f32) {
5553    let (gi, k) = (j / GROUP_SIZE, j % GROUP_SIZE);
5554    let byte = chunks[(r * gpr + gi) * Q2TP_CHUNK + k / 4];
5555    let c = (byte >> (2 * (k % 4))) & 3;
5556    (c as f32 - 1.5, scales[gi])
5557}
5558
5559#[cfg(target_arch = "x86_64")]
5560const Q2TP_DECODE_U32: [u32; 256] = {
5561    let mut tab = [0u32; 256];
5562    let mut b = 0usize;
5563    while b < 256 {
5564        tab[b] = ((b as u32) & 3)
5565            | ((((b as u32) >> 2) & 3) << 8)
5566            | ((((b as u32) >> 4) & 3) << 16)
5567            | ((((b as u32) >> 6) & 3) << 24);
5568        b += 1;
5569    }
5570    tab
5571};
5572
5573/// Eight packed q2tp bytes against 32 signed activation bytes. `maddubs`
5574/// exactly computes unsigned 2-bit code × signed i8; its pair sums cannot
5575/// saturate (2 × 3 × 127 < i16::MAX), and the second madd widens to i32.
5576#[cfg(target_arch = "x86_64")]
5577#[target_feature(enable = "avx2")]
5578unsafe fn q2tp_code_dot_avx2(ch: &[u8], x: &[i8]) -> i32 {
5579    use core::arch::x86_64::*;
5580    debug_assert!(ch.len() >= Q2TP_CHUNK && x.len() >= GROUP_SIZE);
5581    let codes = _mm256_setr_epi32(
5582        Q2TP_DECODE_U32[ch[0] as usize] as i32,
5583        Q2TP_DECODE_U32[ch[1] as usize] as i32,
5584        Q2TP_DECODE_U32[ch[2] as usize] as i32,
5585        Q2TP_DECODE_U32[ch[3] as usize] as i32,
5586        Q2TP_DECODE_U32[ch[4] as usize] as i32,
5587        Q2TP_DECODE_U32[ch[5] as usize] as i32,
5588        Q2TP_DECODE_U32[ch[6] as usize] as i32,
5589        Q2TP_DECODE_U32[ch[7] as usize] as i32,
5590    );
5591    let xv = unsafe { _mm256_loadu_si256(x.as_ptr().cast()) };
5592    let pair = _mm256_maddubs_epi16(codes, xv);
5593    let quad = _mm256_madd_epi16(pair, _mm256_set1_epi16(1));
5594    let sum128 = _mm_add_epi32(
5595        _mm256_castsi256_si128(quad),
5596        _mm256_extracti128_si256(quad, 1),
5597    );
5598    let sum64 = _mm_hadd_epi32(sum128, sum128);
5599    _mm_cvtsi128_si32(_mm_hadd_epi32(sum64, sum64))
5600}
5601
5602/// Integer dot of one q2tp row against pre-quantized activations:
5603/// Σ_g s_g · (Σ c·xq − 1.5·Σ xq). The half-integer grid (c − 1.5)
5604/// becomes exact integer math through the group sums — the same trick
5605/// every a8w8 kernel in this file rides. The codes decode into a
5606/// 32-byte scratch in natural order and the dot itself is the shared
5607/// SDOT primitive; elsewhere a scalar integer loop.
5608#[inline]
5609fn dot_q2tp_row_i8(
5610    chunks: &[u8],
5611    r: usize,
5612    gpr: usize,
5613    xq: &[i8],
5614    gsum: &[i32],
5615    scales: &[f32],
5616) -> f32 {
5617    let mut acc = 0f32;
5618    let base = r * gpr * Q2TP_CHUNK;
5619    #[cfg(not(any(target_arch = "aarch64", target_arch = "x86_64")))]
5620    let mut codes = [0i8; GROUP_SIZE];
5621    #[cfg(target_arch = "x86_64")]
5622    let avx2 = std::arch::is_x86_feature_detected!("avx2");
5623    for gi in 0..gpr {
5624        let ch = &chunks[base + gi * Q2TP_CHUNK..base + (gi + 1) * Q2TP_CHUNK];
5625        let xg = &xq[gi * GROUP_SIZE..(gi + 1) * GROUP_SIZE];
5626        #[cfg(target_arch = "aarch64")]
5627        // NEON: the byte's four 2-bit fields land in four lane vectors
5628        // (shift+mask), vld4 de-interleaves xq to match (xj[k] =
5629        // xq[4k+j]), widening MACs accumulate exactly in i32. A scalar
5630        // decode here cost as much as the dot it fed — the profile put
5631        // it at the top of the whole W2 decode.
5632        let dot = unsafe {
5633            use core::arch::aarch64::*;
5634            let b = vld1_u8(ch.as_ptr());
5635            let three = vdup_n_u8(3);
5636            let c0 = vreinterpret_s8_u8(vand_u8(b, three));
5637            let c1 = vreinterpret_s8_u8(vand_u8(vshr_n_u8(b, 2), three));
5638            let c2 = vreinterpret_s8_u8(vand_u8(vshr_n_u8(b, 4), three));
5639            let c3 = vreinterpret_s8_u8(vand_u8(vshr_n_u8(b, 6), three));
5640            let x4 = vld4_s8(xg.as_ptr());
5641            let mut acc4 = vdupq_n_s32(0);
5642            acc4 = vpadalq_s16(acc4, vmull_s8(c0, x4.0));
5643            acc4 = vpadalq_s16(acc4, vmull_s8(c1, x4.1));
5644            acc4 = vpadalq_s16(acc4, vmull_s8(c2, x4.2));
5645            acc4 = vpadalq_s16(acc4, vmull_s8(c3, x4.3));
5646            vaddvq_s32(acc4)
5647        };
5648        #[cfg(target_arch = "x86_64")]
5649        let dot: i32 = if avx2 {
5650            // SAFETY: the runtime feature check gates the target-feature body;
5651            // the group slices above are exactly 8 and 32 bytes long.
5652            unsafe { q2tp_code_dot_avx2(ch, xg) }
5653        } else {
5654            ch.iter()
5655                .enumerate()
5656                .map(|(k, &b)| {
5657                    ((b & 3) as i32) * xg[k * 4] as i32
5658                        + (((b >> 2) & 3) as i32) * xg[k * 4 + 1] as i32
5659                        + (((b >> 4) & 3) as i32) * xg[k * 4 + 2] as i32
5660                        + (((b >> 6) & 3) as i32) * xg[k * 4 + 3] as i32
5661                })
5662                .sum()
5663        };
5664        #[cfg(not(any(target_arch = "aarch64", target_arch = "x86_64")))]
5665        let dot: i32 = {
5666            for (k, &b) in ch.iter().enumerate() {
5667                codes[k * 4] = (b & 3) as i8;
5668                codes[k * 4 + 1] = ((b >> 2) & 3) as i8;
5669                codes[k * 4 + 2] = ((b >> 4) & 3) as i8;
5670                codes[k * 4 + 3] = ((b >> 6) & 3) as i8;
5671            }
5672            codes
5673                .iter()
5674                .zip(xg)
5675                .map(|(&c, &x)| c as i32 * x as i32)
5676                .sum()
5677        };
5678        acc += scales[gi] * (dot as f32 - 1.5 * gsum[gi] as f32);
5679    }
5680    acc
5681}
5682
5683/// Exact f32 dot of one q2tp row: 2-bit fields LSB-first, (c − 1.5)·s.
5684/// Scalar on purpose — the 2-bit class targets the GPU graph; the CPU
5685/// path exists for parity gates and small-machine fallback.
5686fn q2tp_row_exact(chunks: &[u8], r: usize, gpr: usize, x: &[f32], scales: &[f32]) -> f32 {
5687    q2tp_row_exact_center(chunks, r, gpr, x, scales, 1.5)
5688}
5689
5690/// Fused Prism affine row: the derived correction is applied inside the
5691/// decoded code, avoiding a second accumulated dot and avoiding cancellation
5692/// between `B=(c-1.5)s` and `+.5s` for long 5120/17408 rows.
5693#[inline]
5694fn q2tp_affine_row_exact(chunks: &[u8], r: usize, gpr: usize, x: &[f32], scales: &[f32]) -> f32 {
5695    q2tp_row_exact_center(chunks, r, gpr, x, scales, 1.0)
5696}
5697
5698#[inline]
5699fn q2tp_row_exact_center(
5700    chunks: &[u8],
5701    r: usize,
5702    gpr: usize,
5703    x: &[f32],
5704    scales: &[f32],
5705    center: f32,
5706) -> f32 {
5707    let mut acc = 0f32;
5708    for gi in 0..gpr {
5709        let ch = &chunks[(r * gpr + gi) * Q2TP_CHUNK..(r * gpr + gi + 1) * Q2TP_CHUNK];
5710        let s = scales[gi];
5711        let xb = &x[gi * GROUP_SIZE..(gi + 1) * GROUP_SIZE];
5712        let mut g = 0f32;
5713        for (k, &b) in ch.iter().enumerate() {
5714            g += ((b & 3) as f32 - center) * xb[k * 4]
5715                + (((b >> 2) & 3) as f32 - center) * xb[k * 4 + 1]
5716                + (((b >> 4) & 3) as f32 - center) * xb[k * 4 + 2]
5717                + (((b >> 6) & 3) as f32 - center) * xb[k * 4 + 3];
5718        }
5719        acc += s * g;
5720    }
5721    acc
5722}
5723
5724fn q2tp_matvec(
5725    bytes: &[u8],
5726    x: &[f32],
5727    rows: usize,
5728    cols: usize,
5729    out: &mut [f32],
5730    pool: Option<&Pool>,
5731) {
5732    q2tp_matvec_mode(bytes, x, rows, cols, out, pool, false);
5733}
5734
5735fn q2tp_affine_matvec(
5736    bytes: &[u8],
5737    x: &[f32],
5738    rows: usize,
5739    cols: usize,
5740    out: &mut [f32],
5741    pool: Option<&Pool>,
5742) {
5743    q2tp_matvec_mode(bytes, x, rows, cols, out, pool, true);
5744}
5745
5746fn q2tp_matvec_mode(
5747    bytes: &[u8],
5748    x: &[f32],
5749    rows: usize,
5750    cols: usize,
5751    out: &mut [f32],
5752    pool: Option<&Pool>,
5753    affine: bool,
5754) {
5755    debug_assert_eq!(out.len(), rows);
5756    let gpr = cols / GROUP_SIZE;
5757    let v = Q4tpView::new_q2(bytes, rows, cols);
5758    let out_addr = SendMut(out.as_mut_ptr());
5759    // a8w8 fast path (CMF_SDOT=0 keeps the exact scalar walk): integer
5760    // code dots + group sums, exact outlier correction — the same
5761    // contract as every sibling kernel; measured 2-bit rows were the
5762    // only scalar holdout in the family.
5763    if !affine && a8w8_enabled() {
5764        let act = split_act(x);
5765        let gsum = q1_group_sums(&act.xq, gpr);
5766        let (act, gsum) = (&act, &gsum);
5767        let run = move |start: usize, end: usize| {
5768            with_krow(gpr, |sc| {
5769                for r in start..end {
5770                    v.scales_into(r, gpr, sc);
5771                    let mut acc = dot_q2tp_row_i8(v.nib, r, gpr, &act.xq, gsum, sc) * act.sx;
5772                    for &(j, xv) in &act.outliers {
5773                        let (w, s) = q2tp_outlier(v.nib, r, gpr, j, sc);
5774                        acc += w * s * xv;
5775                    }
5776                    // SAFETY: disjoint row ranges per worker.
5777                    unsafe { *out_addr.at(r) = acc };
5778                }
5779            })
5780        };
5781        dispatch_rows(pool, rows, &run);
5782        return;
5783    }
5784    let run = |start: usize, end: usize| {
5785        with_krow(gpr, |sc| {
5786            for r in start..end {
5787                v.scales_into(r, gpr, sc);
5788                // SAFETY: disjoint row ranges per worker.
5789                unsafe {
5790                    *out_addr.at(r) = if affine {
5791                        q2tp_affine_row_exact(v.nib, r, gpr, x, sc)
5792                    } else {
5793                        q2tp_row_exact(v.nib, r, gpr, x, sc)
5794                    }
5795                };
5796            }
5797        })
5798    };
5799    dispatch_rows(pool, rows, &run);
5800}
5801
5802/// Fused two-input q2tp matvec — the SwiGLU gate/up pair.
5803#[allow(clippy::too_many_arguments)]
5804fn q2tp_matvec2(
5805    bytes: &[u8],
5806    x1: &[f32],
5807    x2: &[f32],
5808    rows: usize,
5809    cols: usize,
5810    o1: &mut [f32],
5811    o2: &mut [f32],
5812    pool: Option<&Pool>,
5813) {
5814    q2tp_matvec2_mode(bytes, x1, x2, rows, cols, o1, o2, pool, false);
5815}
5816
5817#[allow(clippy::too_many_arguments)]
5818fn q2tp_affine_matvec2(
5819    bytes: &[u8],
5820    x1: &[f32],
5821    x2: &[f32],
5822    rows: usize,
5823    cols: usize,
5824    o1: &mut [f32],
5825    o2: &mut [f32],
5826    pool: Option<&Pool>,
5827) {
5828    q2tp_matvec2_mode(bytes, x1, x2, rows, cols, o1, o2, pool, true);
5829}
5830
5831#[allow(clippy::too_many_arguments)]
5832fn q2tp_matvec2_mode(
5833    bytes: &[u8],
5834    x1: &[f32],
5835    x2: &[f32],
5836    rows: usize,
5837    cols: usize,
5838    o1: &mut [f32],
5839    o2: &mut [f32],
5840    pool: Option<&Pool>,
5841    affine: bool,
5842) {
5843    let gpr = cols / GROUP_SIZE;
5844    let v = Q4tpView::new_q2(bytes, rows, cols);
5845    let (p1, p2) = (SendMut(o1.as_mut_ptr()), SendMut(o2.as_mut_ptr()));
5846    let run = |start: usize, end: usize| {
5847        let mut sc = vec![0f32; gpr];
5848        for r in start..end {
5849            v.scales_into(r, gpr, &mut sc);
5850            // SAFETY: disjoint row ranges per worker.
5851            unsafe {
5852                *p1.at(r) = if affine {
5853                    q2tp_affine_row_exact(v.nib, r, gpr, x1, &sc)
5854                } else {
5855                    q2tp_row_exact(v.nib, r, gpr, x1, &sc)
5856                };
5857                *p2.at(r) = if affine {
5858                    q2tp_affine_row_exact(v.nib, r, gpr, x2, &sc)
5859                } else {
5860                    q2tp_row_exact(v.nib, r, gpr, x2, &sc)
5861                };
5862            }
5863        }
5864    };
5865    dispatch_rows(pool, rows, &run);
5866}
5867
5868/// Batched q2tp matmat: scalar row kernel over every batch column. CPU
5869/// prefill only — decode rides the graph, so plain and correct beats
5870/// clever here.
5871/// Test doors into the host 2-bit kernels: the stand's heap corruption
5872/// pointed at down-shaped tensors, and the private fns need a way to be
5873/// held to a reference without a model file around them.
5874pub fn q2tp_matvec_for_test(bytes: &[u8], x: &[f32], rows: usize, cols: usize, out: &mut [f32]) {
5875    // The facade IS the reference: encoder oracles hold requant output
5876    // to the exact scalar walk. The production dispatch may take the i8
5877    // fast path, whose error scale is the ACTIVATIONS' — a different
5878    // claim than the encoder correctness these tests pin.
5879    let gpr = cols / GROUP_SIZE;
5880    let v = Q4tpView::new_q2(bytes, rows, cols);
5881    with_krow(gpr, |sc| {
5882        for r in 0..rows {
5883            v.scales_into(r, gpr, sc);
5884            out[r] = q2tp_row_exact(v.nib, r, gpr, x, sc);
5885        }
5886    });
5887}
5888
5889/// Test door for the descriptor-specific fused affine decode.  Production
5890/// callers select this through a validated Prism header, never by dtype alone.
5891pub fn q2tp_affine_matvec_for_test(
5892    bytes: &[u8],
5893    x: &[f32],
5894    rows: usize,
5895    cols: usize,
5896    out: &mut [f32],
5897) {
5898    q2tp_affine_matvec(bytes, x, rows, cols, out, None);
5899}
5900
5901pub fn q2tp_matmat_for_test(
5902    bytes: &[u8],
5903    xs_all: &[f32],
5904    b: usize,
5905    rows: usize,
5906    cols: usize,
5907    out: &mut [f32],
5908) {
5909    q2tp_matmat(bytes, xs_all, b, rows, cols, out, None);
5910}
5911
5912fn q2tp_matmat(
5913    bytes: &[u8],
5914    xs_all: &[f32],
5915    b: usize,
5916    rows: usize,
5917    cols: usize,
5918    out: &mut [f32],
5919    pool: Option<&Pool>,
5920) {
5921    q2tp_matmat_mode(bytes, xs_all, b, rows, cols, out, pool, false);
5922}
5923
5924fn q2tp_affine_matmat(
5925    bytes: &[u8],
5926    xs_all: &[f32],
5927    b: usize,
5928    rows: usize,
5929    cols: usize,
5930    out: &mut [f32],
5931    pool: Option<&Pool>,
5932) {
5933    q2tp_matmat_mode(bytes, xs_all, b, rows, cols, out, pool, true);
5934}
5935
5936fn q2tp_matmat_mode(
5937    bytes: &[u8],
5938    xs_all: &[f32],
5939    b: usize,
5940    rows: usize,
5941    cols: usize,
5942    out: &mut [f32],
5943    pool: Option<&Pool>,
5944    affine: bool,
5945) {
5946    debug_assert_eq!(out.len(), b * rows);
5947    let gpr = cols / GROUP_SIZE;
5948    let v = Q4tpView::new_q2(bytes, rows, cols);
5949    let out_addr = SendMut(out.as_mut_ptr());
5950    let run = |start: usize, end: usize| {
5951        let mut sc = vec![0f32; gpr];
5952        for r in start..end {
5953            v.scales_into(r, gpr, &mut sc);
5954            for bi in 0..b {
5955                let x = &xs_all[bi * cols..(bi + 1) * cols];
5956                // SAFETY: disjoint row ranges per worker.
5957                unsafe {
5958                    *out_addr.at(bi * rows + r) = if affine {
5959                        q2tp_affine_row_exact(v.nib, r, gpr, x, &sc)
5960                    } else {
5961                        q2tp_row_exact(v.nib, r, gpr, x, &sc)
5962                    }
5963                };
5964            }
5965        }
5966    };
5967    dispatch_rows(pool, rows, &run);
5968}
5969
5970/// The pre-vectorised shape, kept for A/B (`CMF_Q4TP_V1=1`): the
5971/// horizontal add lands once per group per column instead of once per
5972/// row. Same weights, same activations — only the reduction differs.
5973/// It is also the row-exact batch kernel: per group and column it forms
5974/// `int dot as f32 * scale` and adds it to a scalar running sum, which is
5975/// `dot_q4tp_row_sdot` step for step (Rust never contracts to an fma),
5976/// so each column is bit-identical to that column's matvec.
5977#[cfg(target_arch = "aarch64")]
5978#[target_feature(enable = "neon,dotprod")]
5979unsafe fn dot_q4tp_row_1x4_sdot_v1(
5980    nib: &[u8],
5981    r: usize,
5982    gpr: usize,
5983    xs: [&[i8]; 4],
5984    scales: &[f32],
5985) -> [f32; 4] {
5986    unsafe {
5987        use core::arch::aarch64::*;
5988        use core::arch::asm;
5989        let lomask = vdupq_n_u8(0x0F);
5990        let eight = vdupq_n_s8(8);
5991        let (mut f0, mut f1, mut f2, mut f3) = (0f32, 0f32, 0f32, 0f32);
5992        for gi in 0..gpr {
5993            let t = nib.as_ptr().add((r * gpr + gi) * Q4TP_NIB);
5994            let s = *scales.get_unchecked(gi);
5995            let bb = vld1q_u8(t);
5996            let lo = vandq_u8(bb, lomask);
5997            let hi = vshrq_n_u8::<4>(bb);
5998            let e0 = vsubq_s8(vreinterpretq_s8_u8(vzip1q_u8(lo, hi)), eight);
5999            let e1 = vsubq_s8(vreinterpretq_s8_u8(vzip2q_u8(lo, hi)), eight);
6000            let mut d = [0f32; 4];
6001            for (k, dk) in d.iter_mut().enumerate() {
6002                let x0 = vld1q_s8(xs[k].as_ptr().add(gi * GROUP_SIZE));
6003                let x1 = vld1q_s8(xs[k].as_ptr().add(gi * GROUP_SIZE + 16));
6004                let (mut a0, mut a1) = (vdupq_n_s32(0), vdupq_n_s32(0));
6005                asm!(
6006                    "sdot {a0:v}.4s, {e0:v}.16b, {x0:v}.16b",
6007                    "sdot {a1:v}.4s, {e1:v}.16b, {x1:v}.16b",
6008                    a0 = inout(vreg) a0, a1 = inout(vreg) a1,
6009                    e0 = in(vreg) e0, x0 = in(vreg) x0, e1 = in(vreg) e1, x1 = in(vreg) x1,
6010                    options(pure, nomem, nostack),
6011                );
6012                *dk = vaddvq_s32(vaddq_s32(a0, a1)) as f32 * s;
6013            }
6014            f0 += d[0];
6015            f1 += d[1];
6016            f2 += d[2];
6017            f3 += d[3];
6018        }
6019        [f0, f1, f2, f3]
6020    }
6021}
6022
6023/// Which q4tp batch kernel to run: 1 = the previous one, 2 = the tuned
6024/// one, 0 = decide from the CPU. An atomic rather than a `OnceLock` so a
6025/// benchmark can alternate the two inside one process, where the machine's
6026/// mood — a shared box drifts ±25% between runs — is the same for both.
6027/// What the two mean is per-architecture: on x86 the blocked AVX-512 path
6028/// against the per-column one, on ARM the two reduction shapes.
6029#[allow(dead_code)]
6030static Q4TP_ALT: std::sync::atomic::AtomicU8 = std::sync::atomic::AtomicU8::new(0);
6031
6032/// Tests that store `Q4TP_ALT` hold this, so one test's kernel pick does
6033/// not leak into another's bit-exact comparison running in parallel.
6034#[cfg(test)]
6035static Q4TP_ALT_TEST_LOCK: std::sync::Mutex<()> = std::sync::Mutex::new(());
6036
6037/// Blocking pays on x86 only with 512-bit VNNI. With AVX2 alone, four
6038/// columns sharing an unpack still measured slower than the per-column
6039/// path (23.2 ms against 19.4 on a 48-thread EPYC), because that path
6040/// already dequantizes the row once — so the blocked kernel bought a
6041/// second unpack-free pass at the price of half the vector width.
6042#[cfg(target_arch = "x86_64")]
6043fn q4tp_blocked_x86() -> bool {
6044    match Q4TP_ALT.load(std::sync::atomic::Ordering::Relaxed) {
6045        1 => false,
6046        // A forced ON still asks the CPU. The switch exists so a bench can
6047        // pick a kernel, not so it can promise instructions the machine
6048        // does not have — CI caught that as a SIGILL on a runner without
6049        // AVX-512, where the parity test had turned the path on by hand.
6050        2 => avx512vnni_enabled(),
6051        // Deliberately not cached back into the switch: both gates below
6052        // hold their own `OnceLock`, and latching their answer here would
6053        // make a test's override outlive the test that set it.
6054        _ => blocked_enabled() && avx512vnni_enabled(),
6055    }
6056}
6057
6058/// `CMF_Q4TP_V1=1` picks the old reduction shape (A/B only).
6059#[cfg(target_arch = "aarch64")]
6060#[allow(dead_code)]
6061fn q4tp_v1() -> bool {
6062    match Q4TP_ALT.load(std::sync::atomic::Ordering::Relaxed) {
6063        1 => true,
6064        2 => false,
6065        _ => {
6066            static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
6067            *ON.get_or_init(|| std::env::var("CMF_Q4TP_V1").is_ok_and(|v| v != "0"))
6068        }
6069    }
6070}
6071
6072/// Two weight rows against eight columns. The activation load is the
6073/// same for both rows, so it is paid once for twice the arithmetic, and
6074/// sixteen accumulator chains run where eight did — which is what a kernel
6075/// retiring 0.29 instructions a cycle is short of. Register pressure is
6076/// the limit: sixteen `zmm` accumulators, two weight tiles, one
6077/// activation, of thirty-two.
6078///
6079/// Four rows by four columns spends the same sixteen accumulators the
6080/// other way and measured worse — 1488 GFLOP/s against 1644 — so the
6081/// unpack, which four rows pay twice as often, costs more than the extra
6082/// sharing of one activation load buys.
6083#[cfg(target_arch = "x86_64")]
6084#[target_feature(enable = "avx512f,avx512bw,avx512vnni")]
6085unsafe fn dot_q4tp_2x8_avx512(
6086    nib: &[u8],
6087    r0: usize,
6088    gpr: usize,
6089    xs: [&[i8]; 8],
6090    sc0: &[f32],
6091    sc1: &[f32],
6092) -> [[f32; 8]; 2] {
6093    // SAFETY: as dot_q4tp_row_1x8_avx512, two adjacent rows at once; the
6094    // caller guarantees r0 + 1 < rows and the ISA.
6095    unsafe {
6096        use core::arch::x86_64::*;
6097        let lomask = _mm256_set1_epi8(0x0F);
6098        let eight = _mm256_set1_epi8(8);
6099        let zero = _mm512_setzero_si512();
6100        let mut v0 = [_mm512_setzero_ps(); 8];
6101        let mut v1 = [_mm512_setzero_ps(); 8];
6102        let pairs = gpr / 2;
6103        let unpack = |r: usize, gi: usize| -> (__m512i, __mmask64) {
6104            let t = nib.as_ptr().add((r * gpr + gi) * Q4TP_NIB);
6105            let bb = _mm256_loadu_si256(t as *const __m256i);
6106            let lo = _mm256_and_si256(bb, lomask);
6107            let hi = _mm256_and_si256(_mm256_srli_epi16::<4>(bb), lomask);
6108            let ul = _mm256_sub_epi8(_mm256_unpacklo_epi8(lo, hi), eight);
6109            let uh = _mm256_sub_epi8(_mm256_unpackhi_epi8(lo, hi), eight);
6110            let cat = _mm512_inserti64x4::<1>(_mm512_castsi256_si512(ul), uh);
6111            let w = _mm512_shuffle_i64x2::<0b11_01_10_00>(cat, cat);
6112            (_mm512_abs_epi8(w), _mm512_movepi8_mask(w))
6113        };
6114        for gp in 0..pairs {
6115            let gi = gp * 2;
6116            let (wa0, neg0) = unpack(r0, gi);
6117            let (wa1, neg1) = unpack(r0 + 1, gi);
6118            let off = gi * GROUP_SIZE;
6119            let sv = |sc: &[f32]| {
6120                _mm512_insertf32x8::<1>(
6121                    _mm512_castps256_ps512(_mm256_set1_ps(*sc.get_unchecked(gi))),
6122                    _mm256_set1_ps(*sc.get_unchecked(gi + 1)),
6123                )
6124            };
6125            let s0 = sv(sc0);
6126            let s1 = sv(sc1);
6127            for k in 0..8 {
6128                let xv = _mm512_loadu_si512(xs[k].as_ptr().add(off) as *const __m512i);
6129                let d0 = _mm512_cvtepi32_ps(_mm512_dpbusd_epi32(
6130                    zero,
6131                    wa0,
6132                    _mm512_mask_sub_epi8(xv, neg0, zero, xv),
6133                ));
6134                let d1 = _mm512_cvtepi32_ps(_mm512_dpbusd_epi32(
6135                    zero,
6136                    wa1,
6137                    _mm512_mask_sub_epi8(xv, neg1, zero, xv),
6138                ));
6139                v0[k] = _mm512_fmadd_ps(d0, s0, v0[k]);
6140                v1[k] = _mm512_fmadd_ps(d1, s1, v1[k]);
6141            }
6142        }
6143        let mut acc = [[0f32; 8]; 2];
6144        for k in 0..8 {
6145            acc[0][k] = _mm512_reduce_add_ps(v0[k]);
6146            acc[1][k] = _mm512_reduce_add_ps(v1[k]);
6147        }
6148        if gpr % 2 == 1 {
6149            let off = (gpr - 1) * GROUP_SIZE;
6150            for j in off..off + GROUP_SIZE {
6151                let (w0, sa) = q4tp_outlier(nib, r0, gpr, j, sc0);
6152                let (w1, sb) = q4tp_outlier(nib, r0 + 1, gpr, j, sc1);
6153                for k in 0..8 {
6154                    let x = *xs[k].get_unchecked(j) as f32;
6155                    acc[0][k] += w0 * sa * x;
6156                    acc[1][k] += w1 * sb * x;
6157                }
6158            }
6159        }
6160        acc
6161    }
6162}
6163
6164/// The same, eight columns at a time. One unpack then feeds twice as many
6165/// activation streams, so a wide batch reads the weight tile half as
6166/// often; the price is eight accumulators live at once. Measured 9.0 ->
6167/// 8.3 ms at 9216x2304, b=296 on a 48-thread EPYC 9B45.
6168#[cfg(target_arch = "x86_64")]
6169#[target_feature(enable = "avx512f,avx512bw,avx512vnni")]
6170unsafe fn dot_q4tp_row_1x8_avx512(
6171    nib: &[u8],
6172    r: usize,
6173    gpr: usize,
6174    xs: [&[i8]; 8],
6175    scales: &[f32],
6176) -> [f32; 8] {
6177    // SAFETY: as dot_q4tp_row_1x4_avx2; caller guarantees the ISA.
6178    unsafe {
6179        use core::arch::x86_64::*;
6180        let lomask = _mm256_set1_epi8(0x0F);
6181        let eight = _mm256_set1_epi8(8);
6182        let zero = _mm512_setzero_si512();
6183        let (mut v0, mut v1, mut v2, mut v3) = (
6184            _mm512_setzero_ps(),
6185            _mm512_setzero_ps(),
6186            _mm512_setzero_ps(),
6187            _mm512_setzero_ps(),
6188        );
6189        let (mut v4, mut v5, mut v6, mut v7) = (
6190            _mm512_setzero_ps(),
6191            _mm512_setzero_ps(),
6192            _mm512_setzero_ps(),
6193            _mm512_setzero_ps(),
6194        );
6195        let pairs = gpr / 2;
6196        for gp in 0..pairs {
6197            let gi = gp * 2;
6198            let t = nib.as_ptr().add((r * gpr + gi) * Q4TP_NIB);
6199            let bb = _mm256_loadu_si256(t as *const __m256i);
6200            let lo = _mm256_and_si256(bb, lomask);
6201            let hi = _mm256_and_si256(_mm256_srli_epi16::<4>(bb), lomask);
6202            // `unpack` works per 128-bit lane, so the halves come out as
6203            // [A.lo, B.lo] and [A.hi, B.hi]; the shuffle reorders the four
6204            // 128-bit lanes into the weights' natural order, which is what
6205            // the straight activation load expects.
6206            let ul = _mm256_sub_epi8(_mm256_unpacklo_epi8(lo, hi), eight);
6207            let uh = _mm256_sub_epi8(_mm256_unpackhi_epi8(lo, hi), eight);
6208            let cat = _mm512_inserti64x4::<1>(_mm512_castsi256_si512(ul), uh);
6209            let w = _mm512_shuffle_i64x2::<0b11_01_10_00>(cat, cat);
6210            let wabs = _mm512_abs_epi8(w);
6211            let neg = _mm512_movepi8_mask(w);
6212            let off = gi * GROUP_SIZE;
6213            let sv = _mm512_insertf32x8::<1>(
6214                _mm512_castps256_ps512(_mm256_set1_ps(*scales.get_unchecked(gi))),
6215                _mm256_set1_ps(*scales.get_unchecked(gi + 1)),
6216            );
6217            let dot = |x: &[i8]| -> __m512 {
6218                let xv = _mm512_loadu_si512(x.as_ptr().add(off) as *const __m512i);
6219                let sx = _mm512_mask_sub_epi8(xv, neg, zero, xv);
6220                _mm512_cvtepi32_ps(_mm512_dpbusd_epi32(zero, wabs, sx))
6221            };
6222            v0 = _mm512_fmadd_ps(dot(xs[0]), sv, v0);
6223            v1 = _mm512_fmadd_ps(dot(xs[1]), sv, v1);
6224            v2 = _mm512_fmadd_ps(dot(xs[2]), sv, v2);
6225            v3 = _mm512_fmadd_ps(dot(xs[3]), sv, v3);
6226            v4 = _mm512_fmadd_ps(dot(xs[4]), sv, v4);
6227            v5 = _mm512_fmadd_ps(dot(xs[5]), sv, v5);
6228            v6 = _mm512_fmadd_ps(dot(xs[6]), sv, v6);
6229            v7 = _mm512_fmadd_ps(dot(xs[7]), sv, v7);
6230        }
6231        let mut acc = [
6232            _mm512_reduce_add_ps(v0),
6233            _mm512_reduce_add_ps(v1),
6234            _mm512_reduce_add_ps(v2),
6235            _mm512_reduce_add_ps(v3),
6236            _mm512_reduce_add_ps(v4),
6237            _mm512_reduce_add_ps(v5),
6238            _mm512_reduce_add_ps(v6),
6239            _mm512_reduce_add_ps(v7),
6240        ];
6241        // An odd group count leaves one group over; the narrow kernel
6242        // finishes it rather than the tail being a special case here.
6243        if gpr % 2 == 1 {
6244            let off = (gpr - 1) * GROUP_SIZE;
6245            for j in off..off + GROUP_SIZE {
6246                let (w, s) = q4tp_outlier(nib, r, gpr, j, scales);
6247                let ws = w * s;
6248                for k in 0..8 {
6249                    acc[k] += ws * *xs[k].get_unchecked(j) as f32;
6250                }
6251            }
6252        }
6253        acc
6254    }
6255}
6256
6257/// The same four columns, 512 bits wide. Two groups (64 weights) ride one
6258/// unpack and one `vpdpbusd`, where AVX2 needs two unpacks and four
6259/// `maddubs`/`madd` pairs — about 2.3x fewer instructions for the same
6260/// arithmetic. The two groups carry different scales, so the fma takes a
6261/// vector whose halves hold each group's scale rather than a broadcast.
6262///
6263/// There is no 512-bit `vpsignb`, so the activation's sign is applied by
6264/// negating under a mask taken from the weight's sign bits. That mask is
6265/// per-tile, so it is hoisted out of the column loop and the per-column
6266/// cost stays exactly one instruction, as with `sign_epi8`. Weights of
6267/// zero are not zeroed by the mask trick and do not need to be: their
6268/// magnitude is zero, so the product is.
6269#[cfg(target_arch = "x86_64")]
6270#[target_feature(enable = "avx512f,avx512bw,avx512vnni")]
6271unsafe fn dot_q4tp_row_1x4_avx512(
6272    nib: &[u8],
6273    r: usize,
6274    gpr: usize,
6275    xs: [&[i8]; 4],
6276    scales: &[f32],
6277) -> [f32; 4] {
6278    // SAFETY: as dot_q4tp_row_1x4_avx2; caller guarantees the ISA.
6279    unsafe {
6280        use core::arch::x86_64::*;
6281        let lomask = _mm256_set1_epi8(0x0F);
6282        let eight = _mm256_set1_epi8(8);
6283        let zero = _mm512_setzero_si512();
6284        let (mut v0, mut v1, mut v2, mut v3) = (
6285            _mm512_setzero_ps(),
6286            _mm512_setzero_ps(),
6287            _mm512_setzero_ps(),
6288            _mm512_setzero_ps(),
6289        );
6290        let pairs = gpr / 2;
6291        for gp in 0..pairs {
6292            let gi = gp * 2;
6293            let t = nib.as_ptr().add((r * gpr + gi) * Q4TP_NIB);
6294            let bb = _mm256_loadu_si256(t as *const __m256i);
6295            let lo = _mm256_and_si256(bb, lomask);
6296            let hi = _mm256_and_si256(_mm256_srli_epi16::<4>(bb), lomask);
6297            // `unpack` works per 128-bit lane, so the halves come out as
6298            // [A.lo, B.lo] and [A.hi, B.hi]; the shuffle reorders the four
6299            // 128-bit lanes into the weights' natural order, which is what
6300            // the straight activation load expects.
6301            let ul = _mm256_sub_epi8(_mm256_unpacklo_epi8(lo, hi), eight);
6302            let uh = _mm256_sub_epi8(_mm256_unpackhi_epi8(lo, hi), eight);
6303            let cat = _mm512_inserti64x4::<1>(_mm512_castsi256_si512(ul), uh);
6304            let w = _mm512_shuffle_i64x2::<0b11_01_10_00>(cat, cat);
6305            let wabs = _mm512_abs_epi8(w);
6306            let neg = _mm512_movepi8_mask(w);
6307            let off = gi * GROUP_SIZE;
6308            let sv = _mm512_insertf32x8::<1>(
6309                _mm512_castps256_ps512(_mm256_set1_ps(*scales.get_unchecked(gi))),
6310                _mm256_set1_ps(*scales.get_unchecked(gi + 1)),
6311            );
6312            let dot = |x: &[i8]| -> __m512 {
6313                let xv = _mm512_loadu_si512(x.as_ptr().add(off) as *const __m512i);
6314                let sx = _mm512_mask_sub_epi8(xv, neg, zero, xv);
6315                _mm512_cvtepi32_ps(_mm512_dpbusd_epi32(zero, wabs, sx))
6316            };
6317            v0 = _mm512_fmadd_ps(dot(xs[0]), sv, v0);
6318            v1 = _mm512_fmadd_ps(dot(xs[1]), sv, v1);
6319            v2 = _mm512_fmadd_ps(dot(xs[2]), sv, v2);
6320            v3 = _mm512_fmadd_ps(dot(xs[3]), sv, v3);
6321        }
6322        let mut acc = [
6323            _mm512_reduce_add_ps(v0),
6324            _mm512_reduce_add_ps(v1),
6325            _mm512_reduce_add_ps(v2),
6326            _mm512_reduce_add_ps(v3),
6327        ];
6328        // An odd group count leaves one group over; the narrow kernel
6329        // finishes it rather than the tail being a special case here.
6330        if gpr % 2 == 1 {
6331            let off = (gpr - 1) * GROUP_SIZE;
6332            for j in off..off + GROUP_SIZE {
6333                let (w, s) = q4tp_outlier(nib, r, gpr, j, scales);
6334                let ws = w * s;
6335                for k in 0..4 {
6336                    acc[k] += ws * *xs[k].get_unchecked(j) as f32;
6337                }
6338            }
6339        }
6340        acc
6341    }
6342}
6343
6344/// Four batch columns against one q4tp row: the tile is unpacked ONCE and
6345/// spent on four activation streams, which is where a prefill batch stops
6346/// being weight-bandwidth-bound. Twin of `dot_q4t_row_1x4_sdot`.
6347#[cfg(target_arch = "aarch64")]
6348#[target_feature(enable = "neon,dotprod")]
6349unsafe fn dot_q4tp_row_1x4_sdot(
6350    nib: &[u8],
6351    r: usize,
6352    gpr: usize,
6353    xs: [&[i8]; 4],
6354    scales: &[f32],
6355) -> [f32; 4] {
6356    // SAFETY: see dot_q4tp_row_sdot; every xs[k] is gpr·GROUP_SIZE long.
6357    unsafe {
6358        use core::arch::aarch64::*;
6359        use core::arch::asm;
6360        let lomask = vdupq_n_u8(0x0F);
6361        let eight = vdupq_n_s8(8);
6362        // Named accumulators, NOT an array indexed by a loop variable: the
6363        // latter does not stay in registers (the same defect cost 2x in the
6364        // AVX2 q4t kernel and again in WGSL).
6365        //
6366        // They are VECTORS, and the horizontal add happens once at the end
6367        // instead of once per group per column. `vaddvq` is a cross-lane
6368        // reduction — with 72 groups and four columns the old shape paid
6369        // 288 of them per row, each one a dependency stall the pipeline
6370        // cannot hide, to save four float adds. The group's scale now
6371        // rides an fma into the lane accumulators, so the arithmetic per
6372        // group is one convert and one fma. Summation order changes (the
6373        // lanes carry independent partial sums), which is the same
6374        // round-off class the SDOT path already lives in — the strict
6375        // kernel (`CMF_SDOT=0`, what `cortiq ppl` runs) is unchanged and
6376        // stays the reference.
6377        let (mut v0, mut v1, mut v2, mut v3) = (
6378            vdupq_n_f32(0.0),
6379            vdupq_n_f32(0.0),
6380            vdupq_n_f32(0.0),
6381            vdupq_n_f32(0.0),
6382        );
6383        for gi in 0..gpr {
6384            let t = nib.as_ptr().add((r * gpr + gi) * Q4TP_NIB);
6385            let s = *scales.get_unchecked(gi);
6386            let bb = vld1q_u8(t);
6387            let lo = vandq_u8(bb, lomask);
6388            let hi = vshrq_n_u8::<4>(bb);
6389            let e0 = vsubq_s8(vreinterpretq_s8_u8(vzip1q_u8(lo, hi)), eight);
6390            let e1 = vsubq_s8(vreinterpretq_s8_u8(vzip2q_u8(lo, hi)), eight);
6391            let off = gi * GROUP_SIZE;
6392            let dot4 = |x: &[i8]| -> int32x4_t {
6393                let x0 = vld1q_s8(x.as_ptr().add(off));
6394                let x1 = vld1q_s8(x.as_ptr().add(off + 16));
6395                let (mut a0, mut a1) = (vdupq_n_s32(0), vdupq_n_s32(0));
6396                asm!(
6397                    "sdot {a0:v}.4s, {e0:v}.16b, {x0:v}.16b",
6398                    "sdot {a1:v}.4s, {e1:v}.16b, {x1:v}.16b",
6399                    a0 = inout(vreg) a0, a1 = inout(vreg) a1,
6400                    e0 = in(vreg) e0, x0 = in(vreg) x0, e1 = in(vreg) e1, x1 = in(vreg) x1,
6401                    options(pure, nomem, nostack),
6402                );
6403                vaddq_s32(a0, a1)
6404            };
6405            v0 = vfmaq_n_f32(v0, vcvtq_f32_s32(dot4(xs[0])), s);
6406            v1 = vfmaq_n_f32(v1, vcvtq_f32_s32(dot4(xs[1])), s);
6407            v2 = vfmaq_n_f32(v2, vcvtq_f32_s32(dot4(xs[2])), s);
6408            v3 = vfmaq_n_f32(v3, vcvtq_f32_s32(dot4(xs[3])), s);
6409        }
6410        [
6411            vaddvq_f32(v0),
6412            vaddvq_f32(v1),
6413            vaddvq_f32(v2),
6414            vaddvq_f32(v3),
6415        ]
6416    }
6417}
6418
6419/// Fused q4tp matmat — the same three arms `q4t_matmat` has. Shipping only
6420/// the scalar one made Nanbeige-3B decode at 1.2 tok/s against q4t's 5.9:
6421/// the format was fine, the missing arms were the whole regression.
6422///
6423/// Under `row_exact()` every cell is summed in `q4tp_matvec`'s order on
6424/// every architecture; the mode is read once, so one call never mixes
6425/// the two contracts when a concurrent scope opens or closes mid-call.
6426fn q4tp_matmat(
6427    bytes: &[u8],
6428    xs_all: &[f32],
6429    b: usize,
6430    rows: usize,
6431    cols: usize,
6432    out: &mut [f32],
6433    pool: Option<&Pool>,
6434) {
6435    q4tp_matmat_with(bytes, xs_all, b, rows, cols, out, pool, row_exact())
6436}
6437
6438/// `q4tp_matmat` with the row-exact mode passed in rather than read from
6439/// the shared counter: `exact` = each cell equals its token's matvec bit
6440/// for bit, otherwise the fast blocked / AMX arms are free to reorder.
6441#[allow(clippy::too_many_arguments)]
6442fn q4tp_matmat_with(
6443    bytes: &[u8],
6444    xs_all: &[f32],
6445    b: usize,
6446    rows: usize,
6447    cols: usize,
6448    out: &mut [f32],
6449    pool: Option<&Pool>,
6450    exact: bool,
6451) {
6452    debug_assert_eq!(out.len(), b * rows);
6453    let gpr = cols / GROUP_SIZE;
6454    let v = Q4tpView::new(bytes, rows, cols);
6455
6456    // Wide batches ride the AMX through a dequant-tile sgemm, as in q4t.
6457    // An f32 GEMM over dequantized weights is not the matvec's int8 sum,
6458    // so a row-exact batch never takes it.
6459    #[cfg(target_os = "macos")]
6460    if !exact && b >= 8 && rows * cols >= 500_000 && accel_gemm_enabled() {
6461        dequant_matmat_accel(
6462            &|r, dst| {
6463                let mut sc = [0f32; 32];
6464                let mut scv;
6465                let s: &[f32] = if gpr <= 32 {
6466                    v.scales_into(r, gpr, &mut sc);
6467                    &sc[..gpr]
6468                } else {
6469                    scv = vec![0f32; gpr];
6470                    v.scales_into(r, gpr, &mut scv);
6471                    &scv
6472                };
6473                for gi in 0..gpr {
6474                    let tile = &v.nib[(r * gpr + gi) * Q4TP_NIB..(r * gpr + gi + 1) * Q4TP_NIB];
6475                    for (k, &bb) in tile.iter().enumerate() {
6476                        dst[gi * GROUP_SIZE + k * 2] = ((bb & 0x0F) as f32 - 8.0) * s[gi];
6477                        dst[gi * GROUP_SIZE + k * 2 + 1] =
6478                            (((bb >> 4) & 0x0F) as f32 - 8.0) * s[gi];
6479                    }
6480                }
6481            },
6482            xs_all,
6483            b,
6484            rows,
6485            cols,
6486            out,
6487            pool,
6488        );
6489        return;
6490    }
6491
6492    let out_addr = SendMut(out.as_mut_ptr());
6493    if a8w8_enabled() {
6494        let acts: Vec<SplitAct> = (0..b)
6495            .map(|bi| split_act(&xs_all[bi * cols..(bi + 1) * cols]))
6496            .collect();
6497        let acts = &acts;
6498        // ARM stays blocked under `exact` too: the 1x4 kernel then runs in
6499        // its v1 shape, whose per-group `int dot as f32 * scale` and scalar
6500        // running sum are exactly `dot_q4tp_row_sdot`'s, so the tile is
6501        // still unpacked once for four columns and each column equals its
6502        // matvec. The tuned shape (fma into lane partials, one horizontal
6503        // add per row) is 1-16 ulp off the matvec and is kept for `!exact`.
6504        #[cfg(target_arch = "aarch64")]
6505        let blocked_ok = sdot_enabled() && blocked_enabled();
6506        // x86 gets the same blocking: one tile unpack spent on four
6507        // columns. Without it every column re-decoded the row, which is
6508        // why a 48-core EPYC measured a sixth of an M4's per-core rate.
6509        // The gate is `avx2_enabled`, as in q4t — `sdot_enabled` answers
6510        // for ARM's dotprod and is hard-wired false everywhere else, so
6511        // asking it here left the whole blocked path unreachable on x86.
6512        #[cfg(target_arch = "x86_64")]
6513        let blocked_ok = q4tp_blocked_x86() && !exact;
6514        #[cfg(not(any(target_arch = "aarch64", target_arch = "x86_64")))]
6515        let blocked_ok = {
6516            let _ = exact;
6517            false
6518        };
6519        // Columns are swept in panels that fit L2. Without this a
6520        // row-pair walks every activation in the batch — 4.8 MB at
6521        // 512x512 — and does it again for the next pair, so the whole
6522        // batch streams out of the shared cache once per row. Measured
6523        // 800 GB/s of it, flat across batch sizes, which is the signature
6524        // of a loop bound by traffic rather than by arithmetic. A panel of
6525        // 256 columns is 590 KB beside 221 KB of this worker's weights:
6526        // both stay resident and the batch crosses L3 once instead of
6527        // once per row.
6528        let panel_cols: usize = std::env::var("CMF_Q4TP_PANEL")
6529            .ok()
6530            .and_then(|v| v.parse().ok())
6531            .filter(|v| *v > 0)
6532            .unwrap_or(256);
6533        let run = |start: usize, end: usize| {
6534            for abase in (0..acts.len()).step_by(panel_cols) {
6535                let alen = (acts.len() - abase).min(panel_cols);
6536                let mut sc = vec![0f32; gpr];
6537                #[cfg(target_arch = "x86_64")]
6538                let mut r_lo = start;
6539                #[cfg(target_arch = "x86_64")]
6540                if blocked_ok && alen >= 8 {
6541                    let mut sc1 = vec![0f32; gpr];
6542                    while r_lo + 2 <= end {
6543                        v.scales_into(r_lo, gpr, &mut sc);
6544                        v.scales_into(r_lo + 1, gpr, &mut sc1);
6545                        let mut bi = 0usize;
6546                        while bi + 8 <= alen {
6547                            let xs = [
6548                                acts[abase + bi].xq.as_slice(),
6549                                acts[abase + bi + 1].xq.as_slice(),
6550                                acts[abase + bi + 2].xq.as_slice(),
6551                                acts[abase + bi + 3].xq.as_slice(),
6552                                acts[abase + bi + 4].xq.as_slice(),
6553                                acts[abase + bi + 5].xq.as_slice(),
6554                                acts[abase + bi + 6].xq.as_slice(),
6555                                acts[abase + bi + 7].xq.as_slice(),
6556                            ];
6557                            let d = unsafe { dot_q4tp_2x8_avx512(v.nib, r_lo, gpr, xs, &sc, &sc1) };
6558                            for (row, dr, scr) in [(r_lo, &d[0], &sc), (r_lo + 1, &d[1], &sc1)] {
6559                                for k in 0..8 {
6560                                    let act = &acts[abase + bi + k];
6561                                    let mut acc = dr[k] * act.sx;
6562                                    for &(j, xv) in &act.outliers {
6563                                        let (w, s) = q4tp_outlier(v.nib, row, gpr, j, scr);
6564                                        acc += w * s * xv;
6565                                    }
6566                                    // SAFETY: disjoint (bi, r) cells per worker.
6567                                    unsafe { *out_addr.at((abase + bi + k) * rows + row) = acc };
6568                                }
6569                            }
6570                            bi += 8;
6571                        }
6572                        // Columns past the last group of eight, both rows —
6573                        // the same single-row kernel the tail below uses.
6574                        for row in [r_lo, r_lo + 1] {
6575                            let scr: &[f32] = if row == r_lo { &sc } else { &sc1 };
6576                            for b2 in bi..alen {
6577                                let act = &acts[abase + b2];
6578                                let xs4 = [
6579                                    act.xq.as_slice(),
6580                                    act.xq.as_slice(),
6581                                    act.xq.as_slice(),
6582                                    act.xq.as_slice(),
6583                                ];
6584                                let d =
6585                                    unsafe { dot_q4tp_row_1x4_avx512(v.nib, row, gpr, xs4, scr) };
6586                                let mut acc = d[0] * act.sx;
6587                                for &(j, xv) in &act.outliers {
6588                                    let (w, s) = q4tp_outlier(v.nib, row, gpr, j, scr);
6589                                    acc += w * s * xv;
6590                                }
6591                                // SAFETY: disjoint (bi, r) cells per worker.
6592                                unsafe { *out_addr.at((abase + b2) * rows + row) = acc };
6593                            }
6594                        }
6595                        r_lo += 2;
6596                    }
6597                }
6598                #[cfg(target_arch = "x86_64")]
6599                let row_start = r_lo;
6600                #[cfg(not(target_arch = "x86_64"))]
6601                let row_start = start;
6602                for r in row_start..end {
6603                    v.scales_into(r, gpr, &mut sc);
6604                    let mut bi = 0usize;
6605                    #[cfg(target_arch = "x86_64")]
6606                    if blocked_ok {
6607                        while bi + 8 <= alen {
6608                            let xs = [
6609                                acts[abase + bi].xq.as_slice(),
6610                                acts[abase + bi + 1].xq.as_slice(),
6611                                acts[abase + bi + 2].xq.as_slice(),
6612                                acts[abase + bi + 3].xq.as_slice(),
6613                                acts[abase + bi + 4].xq.as_slice(),
6614                                acts[abase + bi + 5].xq.as_slice(),
6615                                acts[abase + bi + 6].xq.as_slice(),
6616                                acts[abase + bi + 7].xq.as_slice(),
6617                            ];
6618                            let d = unsafe { dot_q4tp_row_1x8_avx512(v.nib, r, gpr, xs, &sc) };
6619                            for k in 0..8 {
6620                                let act = &acts[abase + bi + k];
6621                                let mut acc = d[k] * act.sx;
6622                                for &(j, xv) in &act.outliers {
6623                                    let (w, s) = q4tp_outlier(v.nib, r, gpr, j, &sc);
6624                                    acc += w * s * xv;
6625                                }
6626                                // SAFETY: disjoint (bi, r) cells per worker.
6627                                unsafe { *out_addr.at((abase + bi + k) * rows + r) = acc };
6628                            }
6629                            bi += 8;
6630                        }
6631                        while bi + 4 <= alen {
6632                            let xs = [
6633                                acts[abase + bi].xq.as_slice(),
6634                                acts[abase + bi + 1].xq.as_slice(),
6635                                acts[abase + bi + 2].xq.as_slice(),
6636                                acts[abase + bi + 3].xq.as_slice(),
6637                            ];
6638                            let d = unsafe { dot_q4tp_row_1x4_avx512(v.nib, r, gpr, xs, &sc) };
6639                            for k in 0..4 {
6640                                let act = &acts[abase + bi + k];
6641                                let mut acc = d[k] * act.sx;
6642                                for &(j, xv) in &act.outliers {
6643                                    let (w, s) = q4tp_outlier(v.nib, r, gpr, j, &sc);
6644                                    acc += w * s * xv;
6645                                }
6646                                // SAFETY: disjoint (bi, r) cells per worker.
6647                                unsafe { *out_addr.at((abase + bi + k) * rows + r) = acc };
6648                            }
6649                            bi += 4;
6650                        }
6651                    }
6652                    #[cfg(target_arch = "aarch64")]
6653                    if blocked_ok {
6654                        while bi + 4 <= alen {
6655                            let xs = [
6656                                acts[abase + bi].xq.as_slice(),
6657                                acts[abase + bi + 1].xq.as_slice(),
6658                                acts[abase + bi + 2].xq.as_slice(),
6659                                acts[abase + bi + 3].xq.as_slice(),
6660                            ];
6661                            let d = unsafe {
6662                                if exact || q4tp_v1() {
6663                                    dot_q4tp_row_1x4_sdot_v1(v.nib, r, gpr, xs, &sc)
6664                                } else {
6665                                    dot_q4tp_row_1x4_sdot(v.nib, r, gpr, xs, &sc)
6666                                }
6667                            };
6668                            for k in 0..4 {
6669                                let act = &acts[abase + bi + k];
6670                                let mut acc = d[k] * act.sx;
6671                                for &(j, xv) in &act.outliers {
6672                                    let (w, s) = q4tp_outlier(v.nib, r, gpr, j, &sc);
6673                                    acc += w * s * xv;
6674                                }
6675                                // SAFETY: disjoint (bi, r) cells per worker.
6676                                unsafe { *out_addr.at((abase + bi + k) * rows + r) = acc };
6677                            }
6678                            bi += 4;
6679                        }
6680                    }
6681                    let _ = blocked_ok;
6682                    while bi < alen {
6683                        let act = &acts[abase + bi];
6684                        let mut acc = dot_q4tp_row_i8(v.nib, r, gpr, &act.xq, &sc) * act.sx;
6685                        for &(j, xv) in &act.outliers {
6686                            let (w, s) = q4tp_outlier(v.nib, r, gpr, j, &sc);
6687                            acc += w * s * xv;
6688                        }
6689                        // SAFETY: disjoint (bi, r) cells per worker range.
6690                        unsafe { *out_addr.at((abase + bi) * rows + r) = acc };
6691                        bi += 1;
6692                    }
6693                }
6694            }
6695        };
6696        dispatch_rows(pool, rows, &run);
6697        return;
6698    }
6699
6700    let run = |start: usize, end: usize| {
6701        let mut sc = vec![0f32; gpr];
6702        for r in start..end {
6703            v.scales_into(r, gpr, &mut sc);
6704            for bi in 0..b {
6705                let x = &xs_all[bi * cols..(bi + 1) * cols];
6706                // SAFETY: disjoint (bi, r) cells per worker range.
6707                unsafe { *out_addr.at(bi * rows + r) = q4tp_row_exact(v.nib, r, gpr, x, &sc) };
6708            }
6709        }
6710    };
6711    dispatch_rows(pool, rows, &run);
6712}
6713
6714/// Fused q4_tiled matvec (dispatch mirrors `q4matvec`).
6715fn q4t_matvec(
6716    bytes: &[u8],
6717    x: &[f32],
6718    rows: usize,
6719    cols: usize,
6720    out: &mut [f32],
6721    pool: Option<&Pool>,
6722) {
6723    debug_assert_eq!(out.len(), rows);
6724    let gpr = cols / GROUP_SIZE;
6725    let out_addr = SendMut(out.as_mut_ptr());
6726    if a8w8_enabled() {
6727        let act = split_act(x);
6728        let run = move |start: usize, end: usize| {
6729            for r in start..end {
6730                let mut acc = dot_q4t_row_i8(bytes, r, gpr, &act.xq) * act.sx;
6731                for &(j, xv) in &act.outliers {
6732                    let (w, s) = q4t_outlier(bytes, r, gpr, j);
6733                    acc += w * s * xv;
6734                }
6735                // SAFETY: disjoint row ranges per worker.
6736                unsafe { *out_addr.at(r) = acc };
6737            }
6738        };
6739        dispatch_rows(pool, rows, &run);
6740        return;
6741    }
6742    let run = move |start: usize, end: usize| {
6743        for r in start..end {
6744            // SAFETY: disjoint row ranges per worker.
6745            unsafe { *out_addr.at(r) = q4t_row_exact(bytes, r, gpr, x) };
6746        }
6747    };
6748    dispatch_rows(pool, rows, &run);
6749}
6750
6751/// Fused two-input q4_tiled matvec (weights read once per pair).
6752#[allow(clippy::too_many_arguments)]
6753fn q4t_matvec2(
6754    bytes: &[u8],
6755    x1: &[f32],
6756    x2: &[f32],
6757    rows: usize,
6758    cols: usize,
6759    o1: &mut [f32],
6760    o2: &mut [f32],
6761    pool: Option<&Pool>,
6762) {
6763    let gpr = cols / GROUP_SIZE;
6764    let p1 = SendMut(o1.as_mut_ptr());
6765    let p2 = SendMut(o2.as_mut_ptr());
6766    if a8w8_enabled() {
6767        let a1 = split_act(x1);
6768        let a2 = split_act(x2);
6769        let run = move |start: usize, end: usize| {
6770            for r in start..end {
6771                let mut v1 = dot_q4t_row_i8(bytes, r, gpr, &a1.xq) * a1.sx;
6772                let mut v2 = dot_q4t_row_i8(bytes, r, gpr, &a2.xq) * a2.sx;
6773                for &(j, xv) in &a1.outliers {
6774                    let (w, s) = q4t_outlier(bytes, r, gpr, j);
6775                    v1 += w * s * xv;
6776                }
6777                for &(j, xv) in &a2.outliers {
6778                    let (w, s) = q4t_outlier(bytes, r, gpr, j);
6779                    v2 += w * s * xv;
6780                }
6781                // SAFETY: disjoint row ranges per worker.
6782                unsafe {
6783                    *p1.at(r) = v1;
6784                    *p2.at(r) = v2;
6785                }
6786            }
6787        };
6788        dispatch_rows(pool, rows, &run);
6789        return;
6790    }
6791    let run = move |start: usize, end: usize| {
6792        for r in start..end {
6793            // SAFETY: disjoint row ranges per worker.
6794            unsafe {
6795                *p1.at(r) = q4t_row_exact(bytes, r, gpr, x1);
6796                *p2.at(r) = q4t_row_exact(bytes, r, gpr, x2);
6797            }
6798        }
6799    };
6800    dispatch_rows(pool, rows, &run);
6801}
6802
6803/// Batched q4_tiled matmat: each row's tiles stream once per microbatch.
6804#[allow(clippy::too_many_arguments)]
6805/// Prefill GEMM through Accelerate for group-quantized codecs: a
6806/// caller-supplied row dequantizer fills f32 tiles (pool-parallel) and
6807/// each tile rides the AMX with one sgemm — the generic sibling of
6808/// `qmatmat_accel` (q8). Numerics are f32-GEMM (tolerance class);
6809/// decode (b=1) never takes this path.
6810#[cfg(target_os = "macos")]
6811fn dequant_matmat_accel(
6812    dequant_row: &(dyn Fn(usize, &mut [f32]) + Sync),
6813    xs_all: &[f32],
6814    b: usize,
6815    rows: usize,
6816    cols: usize,
6817    out: &mut [f32],
6818    pool: Option<&Pool>,
6819) {
6820    const TR: usize = 2048;
6821    thread_local! {
6822        static WTILE: std::cell::RefCell<Vec<f32>> = const { std::cell::RefCell::new(Vec::new()) };
6823    }
6824    WTILE.with(|wt| {
6825        let mut wtile = wt.borrow_mut();
6826        wtile.resize(TR * cols, 0.0);
6827        let mut r0 = 0usize;
6828        while r0 < rows {
6829            let tr = TR.min(rows - r0);
6830            let wt_addr = SendMut(wtile.as_mut_ptr());
6831            let run = |start: usize, end: usize| {
6832                for r in start..end {
6833                    // SAFETY: workers cover disjoint r ranges.
6834                    let dst = unsafe { std::slice::from_raw_parts_mut(wt_addr.at(r * cols), cols) };
6835                    dequant_row(r0 + r, dst);
6836                }
6837            };
6838            dispatch_rows(pool, tr, &run);
6839            unsafe {
6840                accel_blas::cblas_sgemm(
6841                    101, // RowMajor
6842                    111, // NoTrans A
6843                    112, // Trans B
6844                    b as i32,
6845                    tr as i32,
6846                    cols as i32,
6847                    1.0,
6848                    xs_all.as_ptr(),
6849                    cols as i32,
6850                    wtile.as_ptr(),
6851                    cols as i32,
6852                    0.0,
6853                    out.as_mut_ptr().add(r0),
6854                    rows as i32,
6855                );
6856            }
6857            r0 += tr;
6858        }
6859    });
6860}
6861
6862fn q4t_matmat(
6863    bytes: &[u8],
6864    xs_all: &[f32],
6865    b: usize,
6866    rows: usize,
6867    cols: usize,
6868    out: &mut [f32],
6869    pool: Option<&Pool>,
6870) {
6871    debug_assert_eq!(out.len(), b * rows);
6872    let gpr = cols / GROUP_SIZE;
6873    // Wide batches ride the AMX like q8's qmatmat: on Apple silicon
6874    // the dequant-tile sgemm is an order above the SDOT row loop for
6875    // prefill shapes (imagegen DiT forwards are exactly this).
6876    #[cfg(target_os = "macos")]
6877    if b >= 8 && rows * cols >= 500_000 && accel_gemm_enabled() {
6878        dequant_matmat_accel(
6879            &|r, dst| {
6880                for gi in 0..gpr {
6881                    let tile = &bytes[(r * gpr + gi) * Q4_TILE..(r * gpr + gi + 1) * Q4_TILE];
6882                    let s = f16_to_f32(u16::from_le_bytes([tile[0], tile[1]]));
6883                    for (k, &bb) in tile[2..].iter().enumerate() {
6884                        dst[gi * GROUP_SIZE + k * 2] = ((bb & 0x0F) as f32 - 8.0) * s;
6885                        dst[gi * GROUP_SIZE + k * 2 + 1] = (((bb >> 4) & 0x0F) as f32 - 8.0) * s;
6886                    }
6887                }
6888            },
6889            xs_all,
6890            b,
6891            rows,
6892            cols,
6893            out,
6894            pool,
6895        );
6896        return;
6897    }
6898    let out_addr = SendMut(out.as_mut_ptr());
6899    if a8w8_enabled() {
6900        let acts: Vec<SplitAct> = (0..b)
6901            .map(|bi| split_act(&xs_all[bi * cols..(bi + 1) * cols]))
6902            .collect();
6903        let acts = &acts;
6904        #[cfg(target_arch = "x86_64")]
6905        let blocked_ok = avx2_enabled() && blocked_enabled();
6906        #[cfg(target_arch = "aarch64")]
6907        let blocked_ok = sdot_enabled() && blocked_enabled();
6908        #[cfg(not(any(target_arch = "x86_64", target_arch = "aarch64")))]
6909        let blocked_ok = false;
6910        let run = move |start: usize, end: usize| {
6911            for r in start..end {
6912                let mut bi = 0usize;
6913                #[cfg(target_arch = "aarch64")]
6914                if blocked_ok {
6915                    while bi + 4 <= acts.len() {
6916                        let xs = [
6917                            acts[bi].xq.as_slice(),
6918                            acts[bi + 1].xq.as_slice(),
6919                            acts[bi + 2].xq.as_slice(),
6920                            acts[bi + 3].xq.as_slice(),
6921                        ];
6922                        let d = unsafe { dot_q4t_row_1x4_sdot(bytes, r, gpr, xs) };
6923                        for k in 0..4 {
6924                            let act = &acts[bi + k];
6925                            let mut acc = d[k] * act.sx;
6926                            for &(j, xv) in &act.outliers {
6927                                let (w, sc) = q4t_outlier(bytes, r, gpr, j);
6928                                acc += w * sc * xv;
6929                            }
6930                            // SAFETY: disjoint (bi, r) cells per worker.
6931                            unsafe { *out_addr.at((bi + k) * rows + r) = acc };
6932                        }
6933                        bi += 4;
6934                    }
6935                }
6936                #[cfg(target_arch = "x86_64")]
6937                if blocked_ok {
6938                    while bi + 4 <= acts.len() {
6939                        let xs = [
6940                            acts[bi].xq.as_slice(),
6941                            acts[bi + 1].xq.as_slice(),
6942                            acts[bi + 2].xq.as_slice(),
6943                            acts[bi + 3].xq.as_slice(),
6944                        ];
6945                        let d = unsafe {
6946                            if vnni_tiles_enabled() {
6947                                dot_q4t_row_1x4_vnni(bytes, r, gpr, xs)
6948                            } else {
6949                                dot_q4t_row_1x4_avx2(bytes, r, gpr, xs)
6950                            }
6951                        };
6952                        for k in 0..4 {
6953                            let act = &acts[bi + k];
6954                            let mut acc = d[k] * act.sx;
6955                            for &(j, xv) in &act.outliers {
6956                                let (w, sc) = q4t_outlier(bytes, r, gpr, j);
6957                                acc += w * sc * xv;
6958                            }
6959                            // SAFETY: disjoint (bi, r) cells per worker.
6960                            unsafe { *out_addr.at((bi + k) * rows + r) = acc };
6961                        }
6962                        bi += 4;
6963                    }
6964                }
6965                let _ = blocked_ok;
6966                while bi < acts.len() {
6967                    let act = &acts[bi];
6968                    let mut acc = dot_q4t_row_i8(bytes, r, gpr, &act.xq) * act.sx;
6969                    for &(j, xv) in &act.outliers {
6970                        let (w, s) = q4t_outlier(bytes, r, gpr, j);
6971                        acc += w * s * xv;
6972                    }
6973                    // SAFETY: disjoint (bi, r) cells per worker range.
6974                    unsafe { *out_addr.at(bi * rows + r) = acc };
6975                    bi += 1;
6976                }
6977            }
6978        };
6979        dispatch_rows(pool, rows, &run);
6980        return;
6981    }
6982    let run = move |start: usize, end: usize| {
6983        for r in start..end {
6984            for bi in 0..b {
6985                let x = &xs_all[bi * cols..(bi + 1) * cols];
6986                // SAFETY: disjoint (bi, r) cells per worker range.
6987                unsafe { *out_addr.at(bi * rows + r) = q4t_row_exact(bytes, r, gpr, x) };
6988            }
6989        }
6990    };
6991    dispatch_rows(pool, rows, &run);
6992}
6993
6994// ── q1 (dtype 12): binary weights, [f16 scale][4B sign bits] per
6995// 32-group tile. The kernel family mirrors q4_tiled: one sequential
6996// stream of 6-byte tiles, per-tile integer dot × scale, exact outlier
6997// correction (A8W8 contract), exact scalar path under CMF_SDOT=0. ──
6998
6999/// Per-32-group sums of the quantized activation — the ±1 identity's
7000/// shared half: `dot = −2·sdot(mask, x) − gsum[g]`, computed ONCE per
7001/// matvec and reused by every row.
7002fn q1_group_sums(xq: &[i8], gpr: usize) -> Vec<i32> {
7003    (0..gpr)
7004        .map(|gi| {
7005            xq[gi * GROUP_SIZE..(gi + 1) * GROUP_SIZE]
7006                .iter()
7007                .map(|&v| v as i32)
7008                .sum()
7009        })
7010        .collect()
7011}
7012
7013/// One q1 row via the A8W8 int8 path — mask-SDOT on ARM (no ±1
7014/// expansion at all), scalar bit loop elsewhere (AVX2 queued with the
7015/// x86 pass).
7016#[inline]
7017#[allow(unreachable_code)]
7018/// AVX2 q1 row via the same ±1 identity as the ARM sdot kernel: the
7019/// sign bits expand to a {0, −1} byte mask through shuffle+cmpeq, the
7020/// masked activation sums through maddubs(1, x&mask), and
7021/// `dot = −(2·masked_sum + Σx_group)` — bit-identical integer math.
7022#[cfg(target_arch = "x86_64")]
7023#[target_feature(enable = "avx2")]
7024unsafe fn dot_q1_row_avx2(bytes: &[u8], r: usize, gpr: usize, xq: &[i8], gsum: &[i32]) -> f32 {
7025    // SAFETY: callers uphold the 6B-tile and xq/gsum length contracts.
7026    unsafe {
7027        use core::arch::x86_64::*;
7028        // Byte j of the mask must replicate bits-byte j/8.
7029        let expand = _mm256_setr_epi8(
7030            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,
7031            3, 3, 3,
7032        );
7033        let bitsel = _mm256_setr_epi8(
7034            1, 2, 4, 8, 16, 32, 64, -128, 1, 2, 4, 8, 16, 32, 64, -128, 1, 2, 4, 8, 16, 32, 64,
7035            -128, 1, 2, 4, 8, 16, 32, 64, -128,
7036        );
7037        let ones8 = _mm256_set1_epi8(1);
7038        let ones16 = _mm256_set1_epi16(1);
7039        let mut acc = 0f32;
7040        for gi in 0..gpr {
7041            let t = bytes.as_ptr().add((r * gpr + gi) * Q1_TILE);
7042            let s = f16_to_f32(u16::from_le_bytes([*t, *t.add(1)]));
7043            let bits = u32::from_le_bytes([*t.add(2), *t.add(3), *t.add(4), *t.add(5)]);
7044            let bc = _mm256_shuffle_epi8(_mm256_set1_epi32(bits as i32), expand);
7045            let mask = _mm256_cmpeq_epi8(_mm256_and_si256(bc, bitsel), bitsel);
7046            let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
7047            let sel = _mm256_and_si256(x, mask);
7048            // Σ of selected i8 lanes: maddubs(1u8, sel_i8) pairs → madd.
7049            let p16 = _mm256_maddubs_epi16(ones8, sel);
7050            let d32 = _mm256_madd_epi16(p16, ones16);
7051            let hi128 = _mm256_extracti128_si256::<1>(d32);
7052            let s128 = _mm_add_epi32(_mm256_castsi256_si128(d32), hi128);
7053            let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
7054            let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
7055            let msum = _mm_cvtsi128_si32(s32);
7056            // The and-select keeps x UN-negated (unlike ARM's −1-mask
7057            // sdot): d = Σ_set − Σ_unset = 2·Σ_set − Σ_all.
7058            let d = 2 * msum - gsum[gi];
7059            acc += d as f32 * s;
7060        }
7061        acc
7062    }
7063}
7064
7065/// VNNI twin of `dot_q1_row_avx2`: the masked-select sum goes through
7066/// one `vpdpbusd(1u8, sel)` (see `dpbusd_hsum` — bit-identical).
7067#[cfg(target_arch = "x86_64")]
7068#[target_feature(enable = "avx2,avx512f,avx512bw,avx512vl,avx512vnni")]
7069unsafe fn dot_q1_row_vnni(bytes: &[u8], r: usize, gpr: usize, xq: &[i8], gsum: &[i32]) -> f32 {
7070    // SAFETY: callers uphold the 6B-tile and xq/gsum length contracts.
7071    unsafe {
7072        use core::arch::x86_64::*;
7073        let expand = _mm256_setr_epi8(
7074            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,
7075            3, 3, 3,
7076        );
7077        let bitsel = _mm256_setr_epi8(
7078            1, 2, 4, 8, 16, 32, 64, -128, 1, 2, 4, 8, 16, 32, 64, -128, 1, 2, 4, 8, 16, 32, 64,
7079            -128, 1, 2, 4, 8, 16, 32, 64, -128,
7080        );
7081        let ones8 = _mm256_set1_epi8(1);
7082        let mut acc = 0f32;
7083        for gi in 0..gpr {
7084            let t = bytes.as_ptr().add((r * gpr + gi) * Q1_TILE);
7085            let s = f16_to_f32(u16::from_le_bytes([*t, *t.add(1)]));
7086            let bits = u32::from_le_bytes([*t.add(2), *t.add(3), *t.add(4), *t.add(5)]);
7087            let bc = _mm256_shuffle_epi8(_mm256_set1_epi32(bits as i32), expand);
7088            let mask = _mm256_cmpeq_epi8(_mm256_and_si256(bc, bitsel), bitsel);
7089            let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
7090            let msum = dpbusd_hsum(ones8, _mm256_and_si256(x, mask));
7091            let d = 2 * msum - gsum[gi];
7092            acc += d as f32 * s;
7093        }
7094        acc
7095    }
7096}
7097
7098/// VNNI twin of `dot_q1_row_1x4_avx2` (see `dpbusd_hsum`).
7099#[cfg(target_arch = "x86_64")]
7100#[target_feature(enable = "avx2,avx512f,avx512bw,avx512vl,avx512vnni")]
7101unsafe fn dot_q1_row_1x4_vnni(
7102    bytes: &[u8],
7103    r: usize,
7104    gpr: usize,
7105    xs: [&[i8]; 4],
7106    gsums: [&[i32]; 4],
7107) -> [f32; 4] {
7108    // SAFETY: callers uphold the 6B-tile and xq/gsum length contracts.
7109    unsafe {
7110        use core::arch::x86_64::*;
7111        let expand = _mm256_setr_epi8(
7112            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,
7113            3, 3, 3,
7114        );
7115        let bitsel = _mm256_setr_epi8(
7116            1, 2, 4, 8, 16, 32, 64, -128, 1, 2, 4, 8, 16, 32, 64, -128, 1, 2, 4, 8, 16, 32, 64,
7117            -128, 1, 2, 4, 8, 16, 32, 64, -128,
7118        );
7119        let ones8 = _mm256_set1_epi8(1);
7120        let mut acc = [0f32; 4];
7121        for gi in 0..gpr {
7122            let t = bytes.as_ptr().add((r * gpr + gi) * Q1_TILE);
7123            let s = f16_to_f32(u16::from_le_bytes([*t, *t.add(1)]));
7124            let bits = u32::from_le_bytes([*t.add(2), *t.add(3), *t.add(4), *t.add(5)]);
7125            let bc = _mm256_shuffle_epi8(_mm256_set1_epi32(bits as i32), expand);
7126            let mask = _mm256_cmpeq_epi8(_mm256_and_si256(bc, bitsel), bitsel);
7127            for (k, xq) in xs.iter().enumerate() {
7128                let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
7129                let msum = dpbusd_hsum(ones8, _mm256_and_si256(x, mask));
7130                let d = 2 * msum - gsums[k][gi];
7131                acc[k] += d as f32 * s;
7132            }
7133        }
7134        acc
7135    }
7136}
7137
7138/// The blocked 1×4 flavor: the expanded bit mask serves four activation
7139/// streams per group (mask build once, four select+reduce chains).
7140#[cfg(target_arch = "x86_64")]
7141#[target_feature(enable = "avx2")]
7142unsafe fn dot_q1_row_1x4_avx2(
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 ones16 = _mm256_set1_epi16(1);
7162        let mut acc = [0f32; 4];
7163        for gi in 0..gpr {
7164            let t = bytes.as_ptr().add((r * gpr + gi) * Q1_TILE);
7165            let s = f16_to_f32(u16::from_le_bytes([*t, *t.add(1)]));
7166            let bits = u32::from_le_bytes([*t.add(2), *t.add(3), *t.add(4), *t.add(5)]);
7167            let bc = _mm256_shuffle_epi8(_mm256_set1_epi32(bits as i32), expand);
7168            let mask = _mm256_cmpeq_epi8(_mm256_and_si256(bc, bitsel), bitsel);
7169            for (k, xq) in xs.iter().enumerate() {
7170                let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
7171                let sel = _mm256_and_si256(x, mask);
7172                let p16 = _mm256_maddubs_epi16(ones8, sel);
7173                let d32 = _mm256_madd_epi16(p16, ones16);
7174                let hi128 = _mm256_extracti128_si256::<1>(d32);
7175                let s128 = _mm_add_epi32(_mm256_castsi256_si128(d32), hi128);
7176                let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
7177                let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
7178                let msum = _mm_cvtsi128_si32(s32);
7179                let d = 2 * msum - gsums[k][gi];
7180                acc[k] += d as f32 * s;
7181            }
7182        }
7183        acc
7184    }
7185}
7186
7187#[allow(unreachable_code)]
7188fn dot_q1_row_i8(bytes: &[u8], r: usize, gpr: usize, xq: &[i8], gsum: &[i32]) -> f32 {
7189    #[cfg(target_arch = "aarch64")]
7190    unsafe {
7191        return dot_q1_row_sdot(bytes, r, gpr, xq, gsum);
7192    }
7193    #[cfg(target_arch = "x86_64")]
7194    if avx2_enabled() {
7195        unsafe {
7196            if vnni_tiles_enabled() {
7197                return dot_q1_row_vnni(bytes, r, gpr, xq, gsum);
7198            }
7199            return dot_q1_row_avx2(bytes, r, gpr, xq, gsum);
7200        }
7201    }
7202    let _ = gsum;
7203    let mut acc = 0f32;
7204    for gi in 0..gpr {
7205        let tile = &bytes[(r * gpr + gi) * Q1_TILE..(r * gpr + gi + 1) * Q1_TILE];
7206        let s = f16_to_f32(u16::from_le_bytes([tile[0], tile[1]]));
7207        let mut d = 0i32;
7208        for (j, &b) in tile[2..].iter().enumerate() {
7209            for k in 0..8 {
7210                let w = ((b >> k) & 1) as i32 * 2 - 1;
7211                d += w * xq[gi * GROUP_SIZE + j * 8 + k] as i32;
7212            }
7213        }
7214        acc += d as f32 * s;
7215    }
7216    acc
7217}
7218
7219/// SDOT q1 row via the ±1 identity: the vtst mask (0xFF where the bit
7220/// is set, i.e. −1 as i8) feeds `sdot` DIRECTLY — no expansion to ±1
7221/// lanes at all — and `dot = −(2·sdot(mask, x) + Σx_group)`, with the
7222/// per-group activation sums shared across every row of the matvec.
7223/// Four tiles (128 weights) per iteration: integer dots reduce through
7224/// a vpaddq tree into ONE i32x4 that meets its four scales in a single
7225/// fused f32 multiply-add. Integer math throughout — bit-identical to
7226/// the scalar ±1 reference.
7227#[cfg(target_arch = "aarch64")]
7228#[target_feature(enable = "neon,dotprod")]
7229unsafe fn dot_q1_row_sdot(bytes: &[u8], r: usize, gpr: usize, xq: &[i8], gsum: &[i32]) -> f32 {
7230    // SAFETY: callers uphold slice-length contracts (6B tile per group,
7231    // xq.len() == gpr·GROUP_SIZE, gsum.len() == gpr).
7232    unsafe {
7233        use core::arch::aarch64::*;
7234        use core::arch::asm;
7235        const MASKS: [u8; 16] = [1, 2, 4, 8, 16, 32, 64, 128, 1, 2, 4, 8, 16, 32, 64, 128];
7236        let m = vld1q_u8(MASKS.as_ptr());
7237        // One tile's −Σ_set(x) as an UNREDUCED i32x4 (two mask-sdots).
7238        macro_rules! tile_dot {
7239            ($t:expr, $x:expr) => {{
7240                let v0 = vcombine_u8(vdup_n_u8(*$t.add(2)), vdup_n_u8(*$t.add(3)));
7241                let v1 = vcombine_u8(vdup_n_u8(*$t.add(4)), vdup_n_u8(*$t.add(5)));
7242                let w0 = vreinterpretq_s8_u8(vtstq_u8(v0, m));
7243                let w1 = vreinterpretq_s8_u8(vtstq_u8(v1, m));
7244                let x0 = vld1q_s8($x);
7245                let x1 = vld1q_s8($x.add(16));
7246                let (mut a0, mut a1) = (vdupq_n_s32(0), vdupq_n_s32(0));
7247                asm!(
7248                    "sdot {a0:v}.4s, {w0:v}.16b, {x0:v}.16b",
7249                    "sdot {a1:v}.4s, {w1:v}.16b, {x1:v}.16b",
7250                    a0 = inout(vreg) a0, a1 = inout(vreg) a1,
7251                    w0 = in(vreg) w0, x0 = in(vreg) x0, w1 = in(vreg) w1, x1 = in(vreg) x1,
7252                    options(pure, nomem, nostack),
7253                );
7254                vaddq_s32(a0, a1)
7255            }};
7256        }
7257        // TBL unpack over PAIR loads: one vld1q covers two 6B tiles
7258        // ([s s b b b b][s s b b b b] + 4B slack), TBL replicates each
7259        // bit-byte across 8 lanes for vtst, and the four scales gather
7260        // through tbl2 into one fcvtl — the 16 ld1r broadcast loads and
7261        // 4 branchy software f16 conversions per 128 weights (the
7262        // measured load-port wall of this kernel) become 2 vector
7263        // loads + 9 table lookups. Integer math order is unchanged —
7264        // bit-identical results (FCVTL is exact on every f16).
7265        const IW00: [u8; 16] = [2, 2, 2, 2, 2, 2, 2, 2, 3, 3, 3, 3, 3, 3, 3, 3];
7266        const IW01: [u8; 16] = [4, 4, 4, 4, 4, 4, 4, 4, 5, 5, 5, 5, 5, 5, 5, 5];
7267        const IW10: [u8; 16] = [8, 8, 8, 8, 8, 8, 8, 8, 9, 9, 9, 9, 9, 9, 9, 9];
7268        const IW11: [u8; 16] = [
7269            10, 10, 10, 10, 10, 10, 10, 10, 11, 11, 11, 11, 11, 11, 11, 11,
7270        ];
7271        const ISC: [u8; 8] = [0, 1, 6, 7, 16, 17, 22, 23];
7272        let (iw00, iw01) = (vld1q_u8(IW00.as_ptr()), vld1q_u8(IW01.as_ptr()));
7273        let (iw10, iw11) = (vld1q_u8(IW10.as_ptr()), vld1q_u8(IW11.as_ptr()));
7274        let isc = vld1_u8(ISC.as_ptr());
7275        // One tile's −Σ_set(x) from a TBL-unpacked pair load.
7276        macro_rules! tile_dot_tbl {
7277            ($ld:expr, $i0:expr, $i1:expr, $x:expr) => {{
7278                let w0 = vreinterpretq_s8_u8(vtstq_u8(vqtbl1q_u8($ld, $i0), m));
7279                let w1 = vreinterpretq_s8_u8(vtstq_u8(vqtbl1q_u8($ld, $i1), m));
7280                let x0 = vld1q_s8($x);
7281                let x1 = vld1q_s8($x.add(16));
7282                let (mut a0, mut a1) = (vdupq_n_s32(0), vdupq_n_s32(0));
7283                asm!(
7284                    "sdot {a0:v}.4s, {w0:v}.16b, {x0:v}.16b",
7285                    "sdot {a1:v}.4s, {w1:v}.16b, {x1:v}.16b",
7286                    a0 = inout(vreg) a0, a1 = inout(vreg) a1,
7287                    w0 = in(vreg) w0, x0 = in(vreg) x0, w1 = in(vreg) w1, x1 = in(vreg) x1,
7288                    options(pure, nomem, nostack),
7289                );
7290                vaddq_s32(a0, a1)
7291            }};
7292        }
7293        let base = bytes.as_ptr().add(r * gpr * Q1_TILE);
7294        let row_base = r * gpr * Q1_TILE;
7295        let abs_end = bytes.len();
7296        let xp = xq.as_ptr();
7297        let gp = gsum.as_ptr();
7298        let mut accv = vdupq_n_f32(0.0);
7299        let mut gi = 0;
7300        // The second pair load reads 4B past tile gi+3 — stay inside
7301        // the payload slice (only the file's final tiles fall back).
7302        while gi + 4 <= gpr && row_base + (gi + 4) * Q1_TILE + 4 <= abs_end {
7303            let t0 = base.add(gi * Q1_TILE);
7304            let ld_a = vld1q_u8(t0);
7305            let ld_b = vld1q_u8(t0.add(2 * Q1_TILE));
7306            let d0 = tile_dot_tbl!(ld_a, iw00, iw01, xp.add(gi * GROUP_SIZE));
7307            let d1 = tile_dot_tbl!(ld_a, iw10, iw11, xp.add((gi + 1) * GROUP_SIZE));
7308            let d2 = tile_dot_tbl!(ld_b, iw00, iw01, xp.add((gi + 2) * GROUP_SIZE));
7309            let d3 = tile_dot_tbl!(ld_b, iw10, iw11, xp.add((gi + 3) * GROUP_SIZE));
7310            // [−Σ0, −Σ1, −Σ2, −Σ3] → dots = −(2·Σset_neg + gsum)
7311            let neg = vpaddq_s32(vpaddq_s32(d0, d1), vpaddq_s32(d2, d3));
7312            let g = vld1q_s32(gp.add(gi));
7313            let dots = vnegq_s32(vaddq_s32(vshlq_n_s32::<1>(neg), g));
7314            let sc16 = vqtbl2_u8(uint8x16x2_t(ld_a, ld_b), isc);
7315            let scf: float32x4_t;
7316            asm!(
7317                "fcvtl {o:v}.4s, {i:v}.4h",
7318                o = out(vreg) scf, i = in(vreg) sc16,
7319                options(pure, nomem, nostack),
7320            );
7321            accv = vfmaq_f32(accv, vcvtq_f32_s32(dots), scf);
7322            gi += 4;
7323        }
7324        let mut acc = vaddvq_f32(accv);
7325        while gi < gpr {
7326            let t = base.add(gi * Q1_TILE);
7327            let s = f16_to_f32(u16::from_le_bytes([*t, *t.add(1)]));
7328            let d = vaddvq_s32(tile_dot!(t, xp.add(gi * GROUP_SIZE)));
7329            acc += (-(2 * d + *gp.add(gi))) as f32 * s;
7330            gi += 1;
7331        }
7332        acc
7333    }
7334}
7335
7336/// Blocked q1 1×4: one TBL unpack of the tile pair serves FOUR
7337/// activation streams (prefill amortization — the same idea as the
7338/// AVX2 twin; per stream the group order, fma order and tail match the
7339/// single-row kernel exactly, so batch == matvec bit-for-bit).
7340#[cfg(target_arch = "aarch64")]
7341#[target_feature(enable = "neon,dotprod")]
7342unsafe fn dot_q1_row_1x4_sdot(
7343    bytes: &[u8],
7344    r: usize,
7345    gpr: usize,
7346    xs: [&[i8]; 4],
7347    gs: [&[i32]; 4],
7348) -> [f32; 4] {
7349    // SAFETY: same slice-length contracts as `dot_q1_row_sdot`, ×4.
7350    unsafe {
7351        use core::arch::aarch64::*;
7352        use core::arch::asm;
7353        const MASKS: [u8; 16] = [1, 2, 4, 8, 16, 32, 64, 128, 1, 2, 4, 8, 16, 32, 64, 128];
7354        const IW00: [u8; 16] = [2, 2, 2, 2, 2, 2, 2, 2, 3, 3, 3, 3, 3, 3, 3, 3];
7355        const IW01: [u8; 16] = [4, 4, 4, 4, 4, 4, 4, 4, 5, 5, 5, 5, 5, 5, 5, 5];
7356        const IW10: [u8; 16] = [8, 8, 8, 8, 8, 8, 8, 8, 9, 9, 9, 9, 9, 9, 9, 9];
7357        const IW11: [u8; 16] = [
7358            10, 10, 10, 10, 10, 10, 10, 10, 11, 11, 11, 11, 11, 11, 11, 11,
7359        ];
7360        const ISC: [u8; 8] = [0, 1, 6, 7, 16, 17, 22, 23];
7361        let m = vld1q_u8(MASKS.as_ptr());
7362        let (iw00, iw01) = (vld1q_u8(IW00.as_ptr()), vld1q_u8(IW01.as_ptr()));
7363        let (iw10, iw11) = (vld1q_u8(IW10.as_ptr()), vld1q_u8(IW11.as_ptr()));
7364        let isc = vld1_u8(ISC.as_ptr());
7365        macro_rules! sdot2 {
7366            ($w0:expr, $w1:expr, $x:expr) => {{
7367                let x0 = vld1q_s8($x);
7368                let x1 = vld1q_s8($x.add(16));
7369                let (mut a0, mut a1) = (vdupq_n_s32(0), vdupq_n_s32(0));
7370                asm!(
7371                    "sdot {a0:v}.4s, {w0:v}.16b, {x0:v}.16b",
7372                    "sdot {a1:v}.4s, {w1:v}.16b, {x1:v}.16b",
7373                    a0 = inout(vreg) a0, a1 = inout(vreg) a1,
7374                    w0 = in(vreg) $w0, x0 = in(vreg) x0, w1 = in(vreg) $w1, x1 = in(vreg) x1,
7375                    options(pure, nomem, nostack),
7376                );
7377                vaddq_s32(a0, a1)
7378            }};
7379        }
7380        let base = bytes.as_ptr().add(r * gpr * Q1_TILE);
7381        let row_base = r * gpr * Q1_TILE;
7382        let abs_end = bytes.len();
7383        let mut accv = [vdupq_n_f32(0.0); 4];
7384        let mut gi = 0;
7385        while gi + 4 <= gpr && row_base + (gi + 4) * Q1_TILE + 4 <= abs_end {
7386            let t0 = base.add(gi * Q1_TILE);
7387            let ld_a = vld1q_u8(t0);
7388            let ld_b = vld1q_u8(t0.add(2 * Q1_TILE));
7389            // Unpack ONCE — eight ±mask vectors serve all four streams.
7390            let w00 = vreinterpretq_s8_u8(vtstq_u8(vqtbl1q_u8(ld_a, iw00), m));
7391            let w01 = vreinterpretq_s8_u8(vtstq_u8(vqtbl1q_u8(ld_a, iw01), m));
7392            let w10 = vreinterpretq_s8_u8(vtstq_u8(vqtbl1q_u8(ld_a, iw10), m));
7393            let w11 = vreinterpretq_s8_u8(vtstq_u8(vqtbl1q_u8(ld_a, iw11), m));
7394            let w20 = vreinterpretq_s8_u8(vtstq_u8(vqtbl1q_u8(ld_b, iw00), m));
7395            let w21 = vreinterpretq_s8_u8(vtstq_u8(vqtbl1q_u8(ld_b, iw01), m));
7396            let w30 = vreinterpretq_s8_u8(vtstq_u8(vqtbl1q_u8(ld_b, iw10), m));
7397            let w31 = vreinterpretq_s8_u8(vtstq_u8(vqtbl1q_u8(ld_b, iw11), m));
7398            let sc16 = vqtbl2_u8(uint8x16x2_t(ld_a, ld_b), isc);
7399            let scf: float32x4_t;
7400            asm!(
7401                "fcvtl {o:v}.4s, {i:v}.4h",
7402                o = out(vreg) scf, i = in(vreg) sc16,
7403                options(pure, nomem, nostack),
7404            );
7405            for k in 0..4 {
7406                let xp = xs[k].as_ptr();
7407                let d0 = sdot2!(w00, w01, xp.add(gi * GROUP_SIZE));
7408                let d1 = sdot2!(w10, w11, xp.add((gi + 1) * GROUP_SIZE));
7409                let d2 = sdot2!(w20, w21, xp.add((gi + 2) * GROUP_SIZE));
7410                let d3 = sdot2!(w30, w31, xp.add((gi + 3) * GROUP_SIZE));
7411                let neg = vpaddq_s32(vpaddq_s32(d0, d1), vpaddq_s32(d2, d3));
7412                let g = vld1q_s32(gs[k].as_ptr().add(gi));
7413                let dots = vnegq_s32(vaddq_s32(vshlq_n_s32::<1>(neg), g));
7414                accv[k] = vfmaq_f32(accv[k], vcvtq_f32_s32(dots), scf);
7415            }
7416            gi += 4;
7417        }
7418        let mut acc = [
7419            vaddvq_f32(accv[0]),
7420            vaddvq_f32(accv[1]),
7421            vaddvq_f32(accv[2]),
7422            vaddvq_f32(accv[3]),
7423        ];
7424        while gi < gpr {
7425            let t = base.add(gi * Q1_TILE);
7426            let sc = f16_to_f32(u16::from_le_bytes([*t, *t.add(1)]));
7427            let v0 = vcombine_u8(vdup_n_u8(*t.add(2)), vdup_n_u8(*t.add(3)));
7428            let v1 = vcombine_u8(vdup_n_u8(*t.add(4)), vdup_n_u8(*t.add(5)));
7429            let w0 = vreinterpretq_s8_u8(vtstq_u8(v0, m));
7430            let w1 = vreinterpretq_s8_u8(vtstq_u8(v1, m));
7431            for k in 0..4 {
7432                let d = vaddvq_s32(sdot2!(w0, w1, xs[k].as_ptr().add(gi * GROUP_SIZE)));
7433                acc[k] += (-(2 * d + *gs[k].as_ptr().add(gi))) as f32 * sc;
7434            }
7435            gi += 1;
7436        }
7437        acc
7438    }
7439}
7440
7441/// (weight ±1, scale) of one q1 element — the exact outlier term.
7442#[inline]
7443fn q1_outlier(bytes: &[u8], r: usize, gpr: usize, j: usize) -> (f32, f32) {
7444    let gi = j / GROUP_SIZE;
7445    let k = j % GROUP_SIZE;
7446    let tile = &bytes[(r * gpr + gi) * Q1_TILE..(r * gpr + gi + 1) * Q1_TILE];
7447    let s = f16_to_f32(u16::from_le_bytes([tile[0], tile[1]]));
7448    let bit = (tile[2 + k / 8] >> (k % 8)) & 1;
7449    ((bit as i32 * 2 - 1) as f32, s)
7450}
7451
7452/// Exact scalar q1 row (CMF_SDOT=0 contract).
7453#[inline]
7454fn q1_row_exact(bytes: &[u8], r: usize, gpr: usize, x: &[f32]) -> f32 {
7455    let mut acc = 0f32;
7456    for gi in 0..gpr {
7457        let tile = &bytes[(r * gpr + gi) * Q1_TILE..(r * gpr + gi + 1) * Q1_TILE];
7458        let s = f16_to_f32(u16::from_le_bytes([tile[0], tile[1]]));
7459        let xg = &x[gi * GROUP_SIZE..(gi + 1) * GROUP_SIZE];
7460        let mut ga = 0f32;
7461        for (j, &b) in tile[2..].iter().enumerate() {
7462            for k in 0..8 {
7463                ga += (((b >> k) & 1) as f32 * 2.0 - 1.0) * xg[j * 8 + k];
7464            }
7465        }
7466        acc += ga * s;
7467    }
7468    acc
7469}
7470
7471/// One q1 row range via A8W8 (the body of `q1_matvec`'s hot loop,
7472/// extracted so multi-matrix jobs drive the same kernel).
7473#[allow(clippy::too_many_arguments)]
7474fn q1_range_a8w8(
7475    bytes: &[u8],
7476    gpr: usize,
7477    act: &SplitAct,
7478    gsum: &[i32],
7479    out: SendMut,
7480    start: usize,
7481    end: usize,
7482) {
7483    for r in start..end {
7484        let mut acc = dot_q1_row_i8(bytes, r, gpr, &act.xq, gsum) * act.sx;
7485        for &(j, xv) in &act.outliers {
7486            let (w, s) = q1_outlier(bytes, r, gpr, j);
7487            acc += w * s * xv;
7488        }
7489        // SAFETY: disjoint row ranges per worker.
7490        unsafe { *out.at(r) = acc };
7491    }
7492}
7493
7494/// Exact-scalar q1 row range (CMF_SDOT=0 contract).
7495fn q1_range_f32(bytes: &[u8], gpr: usize, x: &[f32], out: SendMut, start: usize, end: usize) {
7496    for r in start..end {
7497        // SAFETY: disjoint row ranges per worker.
7498        unsafe { *out.at(r) = q1_row_exact(bytes, r, gpr, x) };
7499    }
7500}
7501
7502/// q1t per-row overlay locator. After the base (`base_len`) come
7503/// `[u32 row_ptr[rows+1]]` then `[(u16 col, f16 val)]` grouped by row (row
7504/// `r`'s entries are `[row_ptr[r], row_ptr[r+1])`). Returns
7505/// `(row_ptr offset, entries offset, present)`.
7506fn q1t_overlay(bytes: &[u8], base_len: usize, rows: usize) -> (usize, usize, bool) {
7507    let entries = base_len + (rows + 1) * 4;
7508    (base_len, entries, entries <= bytes.len())
7509}
7510
7511/// Read `row_ptr[r]` from the overlay's prefix-sum table.
7512#[inline]
7513fn q1t_rowptr(bytes: &[u8], rp_off: usize, r: usize) -> usize {
7514    let o = rp_off + r * 4;
7515    u32::from_le_bytes([bytes[o], bytes[o + 1], bytes[o + 2], bytes[o + 3]]) as usize
7516}
7517
7518/// Byte → the 5 ternary signs it packs `{−1,0,+1}` as f32, precomputed so
7519/// decoding a q1t code is a table load, not the base-3 divide/modulo per
7520/// weight (division is ~20–40× the cost of a load). Built at compile time.
7521const SIGN5: [[f32; 5]; 256] = {
7522    let mut lut = [[0.0f32; 5]; 256];
7523    let pow3 = [1u16, 3, 9, 27, 81];
7524    let mut byte = 0usize;
7525    while byte < 256 {
7526        let mut i = 0usize;
7527        while i < 5 {
7528            let code = (byte as u16 / pow3[i]) % 3;
7529            lut[byte][i] = if code == 1 {
7530                1.0
7531            } else if code == 2 {
7532                -1.0
7533            } else {
7534                0.0
7535            };
7536            i += 1;
7537        }
7538        byte += 1;
7539    }
7540    lut
7541};
7542
7543/// Same table, as i8 signs — the operand for the int8 SDOT base kernel.
7544const SIGN5_I8: [[i8; 5]; 256] = {
7545    let mut lut = [[0i8; 5]; 256];
7546    let pow3 = [1u16, 3, 9, 27, 81];
7547    let mut byte = 0usize;
7548    while byte < 256 {
7549        let mut i = 0usize;
7550        while i < 5 {
7551            let code = (byte as u16 / pow3[i]) % 3;
7552            lut[byte][i] = if code == 1 {
7553                1
7554            } else if code == 2 {
7555                -1
7556            } else {
7557                0
7558            };
7559            i += 1;
7560        }
7561        byte += 1;
7562    }
7563    lut
7564};
7565
7566/// The same 5 i8 signs packed into a u64 (`[s0 s1 s2 s3 s4 0 0 0]`, LE) so the
7567/// group unpack is 7 unaligned u64 stores at offsets 0,5,10,…,30 instead of
7568/// six 5-byte copies + LUT indexing — each store's trailing zeros are fixed by
7569/// the next store, and the last one runs 6 B past the 32nd weight (the unpack
7570/// buffer is padded to 40). This is the decode/prefill hot inner op.
7571const SIGN5_U64: [u64; 256] = {
7572    let mut lut = [0u64; 256];
7573    let pow3 = [1u16, 3, 9, 27, 81];
7574    let mut byte = 0usize;
7575    while byte < 256 {
7576        let mut v = 0u64;
7577        let mut i = 0usize;
7578        while i < 5 {
7579            let code = (byte as u16 / pow3[i]) % 3;
7580            let s: u8 = if code == 1 {
7581                1
7582            } else if code == 2 {
7583                0xFF
7584            } else {
7585                0
7586            };
7587            v |= (s as u64) << (i * 8);
7588            i += 1;
7589        }
7590        lut[byte] = v;
7591        byte += 1;
7592    }
7593    lut
7594};
7595
7596/// Ternary base weight at `(row r, col j)` = `sign(code)·s_group`. Used to add
7597/// back activation-outlier columns, whose `x` was zeroed for the int8 bulk dot
7598/// (`split_act`). At a weight-outlier position the code is 0, so this is 0 and
7599/// the overlay correction owns that column — no double counting.
7600#[inline]
7601fn q1t_base_weight(bytes: &[u8], r: usize, gpr: usize, j: usize) -> f32 {
7602    const TILE: usize = cortiq_core::quant::Q1T_TILE;
7603    let off = (r * gpr + j / GROUP_SIZE) * TILE;
7604    let s = f16_to_f32(u16::from_le_bytes([bytes[off], bytes[off + 1]]));
7605    let within = j % GROUP_SIZE;
7606    SIGN5[bytes[off + 2 + within / 5] as usize][within % 5] * s
7607}
7608
7609/// One 32-group int8 dot via two SDOTs. Bit-exact vs the scalar i8 sum
7610/// (integer accumulation is order-independent).
7611#[cfg(target_arch = "aarch64")]
7612#[target_feature(enable = "neon,dotprod")]
7613#[inline]
7614unsafe fn sdot32_i8(w: *const i8, x: *const i8) -> i32 {
7615    // SAFETY: caller guarantees 32 readable i8 at each pointer.
7616    unsafe {
7617        use core::arch::aarch64::*;
7618        use core::arch::asm;
7619        let w0 = vld1q_s8(w);
7620        let w1 = vld1q_s8(w.add(16));
7621        let x0 = vld1q_s8(x);
7622        let x1 = vld1q_s8(x.add(16));
7623        let (mut a0, mut a1) = (vdupq_n_s32(0), vdupq_n_s32(0));
7624        asm!(
7625            "sdot {a0:v}.4s, {w0:v}.16b, {x0:v}.16b",
7626            "sdot {a1:v}.4s, {w1:v}.16b, {x1:v}.16b",
7627            a0 = inout(vreg) a0, a1 = inout(vreg) a1,
7628            w0 = in(vreg) w0, x0 = in(vreg) x0, w1 = in(vreg) w1, x1 = in(vreg) x1,
7629            options(pure, nomem, nostack),
7630        );
7631        vaddvq_s32(vaddq_s32(a0, a1))
7632    }
7633}
7634
7635/// One 32-group int8 dot via AVX2: signed·signed as `maddubs(|w|, sign(x,w))`
7636/// then `madd` and a horizontal reduce (the same idiom as `dot_q4t_row_avx2`).
7637#[cfg(target_arch = "x86_64")]
7638#[target_feature(enable = "avx2")]
7639#[inline]
7640unsafe fn i8dot32_avx2(w: *const i8, x: *const i8) -> i32 {
7641    // SAFETY: caller guarantees 32 readable i8 at each pointer.
7642    unsafe {
7643        use core::arch::x86_64::*;
7644        let wv = _mm256_loadu_si256(w as *const __m256i);
7645        let xv = _mm256_loadu_si256(x as *const __m256i);
7646        let p16 = _mm256_maddubs_epi16(_mm256_abs_epi8(wv), _mm256_sign_epi8(xv, wv));
7647        let d = _mm256_madd_epi16(p16, _mm256_set1_epi16(1));
7648        let hi128 = _mm256_extracti128_si256::<1>(d);
7649        let s128 = _mm_add_epi32(_mm256_castsi256_si128(d), hi128);
7650        let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
7651        let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
7652        _mm_cvtsi128_si32(s32)
7653    }
7654}
7655
7656/// Unpack one q1t group's base-3 codes into 32 i8 signs via 7 unaligned u64
7657/// stores (see `SIGN5_U64`). `dst` MUST have ≥ 40 bytes: the 7th store writes
7658/// `dst[30..38]`. Stores go in order so each one's trailing zeros are
7659/// overwritten by the next; the final 6 padding bytes are unused by the dot.
7660#[inline]
7661fn q1t_unpack_group_i8(codes: *const u8, dst: &mut [i8]) {
7662    debug_assert!(dst.len() >= 40);
7663    // SAFETY: codes points at 7 readable bytes; dst has ≥ 40 bytes so every
7664    // 8-byte store at offset bi*5 (bi ≤ 6 → ≤ 30) stays in bounds.
7665    unsafe {
7666        let p = dst.as_mut_ptr();
7667        for bi in 0..7 {
7668            core::ptr::write_unaligned(
7669                p.add(bi * 5) as *mut u64,
7670                SIGN5_U64[*codes.add(bi) as usize],
7671            );
7672        }
7673    }
7674}
7675
7676/// One 32-group int8 dot, arch-dispatched (the matmat inner loop, where the
7677/// row's signs are unpacked once and dotted against every batch input).
7678/// Callers are gated by `a8w8_enabled()`, so the target-feature arms are
7679/// reachable; the scalar arm is a non-SIMD-arch fallback.
7680#[inline]
7681fn q1t_i8dot32(w: *const i8, x: *const i8) -> i32 {
7682    #[cfg(target_arch = "aarch64")]
7683    unsafe {
7684        return sdot32_i8(w, x);
7685    }
7686    #[cfg(target_arch = "x86_64")]
7687    unsafe {
7688        return i8dot32_avx2(w, x);
7689    }
7690    #[allow(unreachable_code)]
7691    unsafe {
7692        let mut s = 0i32;
7693        for k in 0..GROUP_SIZE {
7694            s += *w.add(k) as i32 * *x.add(k) as i32;
7695        }
7696        s
7697    }
7698}
7699
7700#[inline]
7701unsafe fn q1t_unpack_reg_u64s(codes: *const u8) -> (u64, u64, u64, u64) {
7702    let (s0, s1, s2, s3, s4, s5, s6) = unsafe {
7703        (
7704            SIGN5_U64[*codes as usize],
7705            SIGN5_U64[*codes.add(1) as usize],
7706            SIGN5_U64[*codes.add(2) as usize],
7707            SIGN5_U64[*codes.add(3) as usize],
7708            SIGN5_U64[*codes.add(4) as usize],
7709            SIGN5_U64[*codes.add(5) as usize],
7710            SIGN5_U64[*codes.add(6) as usize],
7711        )
7712    };
7713
7714    let u0 = s0 | (s1 << 40);
7715    let u1 = (s1 >> 24) | (s2 << 16) | (s3 << 56);
7716    let u2 = (s3 >> 8) | (s4 << 32);
7717    let u3 = (s4 >> 32) | (s5 << 8) | (s6 << 48);
7718
7719    (u0, u1, u2, u3)
7720}
7721
7722/// One q1t row's int8 base dot: `Σ_group s·dot(signs, xq)` (before the shared
7723/// `sx`). Direct register unpacking (zero stack stores/loads, no STLF stalls).
7724/// ARM SDOT.
7725#[cfg(target_arch = "aarch64")]
7726#[target_feature(enable = "neon,dotprod")]
7727unsafe fn q1t_dot_row_sdot(bytes: &[u8], r: usize, gpr: usize, xq: &[i8]) -> f32 {
7728    use core::arch::aarch64::*;
7729    use core::arch::asm;
7730    unsafe {
7731        const TILE: usize = cortiq_core::quant::Q1T_TILE;
7732        let mut acc = 0f32;
7733        let bytes_ptr = bytes.as_ptr();
7734        let xq_ptr = xq.as_ptr();
7735        let row_off = r * gpr * TILE;
7736
7737        let gpr2 = gpr & !1;
7738        let mut gi = 0;
7739        while gi < gpr2 {
7740            let off0 = row_off + gi * TILE;
7741            let off1 = off0 + TILE;
7742            let s0 = f16_to_f32(u16::from_le_bytes([
7743                *bytes_ptr.add(off0),
7744                *bytes_ptr.add(off0 + 1),
7745            ]));
7746            let s1 = f16_to_f32(u16::from_le_bytes([
7747                *bytes_ptr.add(off1),
7748                *bytes_ptr.add(off1 + 1),
7749            ]));
7750
7751            let (u0_0, u1_0, u2_0, u3_0) = q1t_unpack_reg_u64s(bytes_ptr.add(off0 + 2));
7752            let (u0_1, u1_1, u2_1, u3_1) = q1t_unpack_reg_u64s(bytes_ptr.add(off1 + 2));
7753
7754            let w0_0 = vreinterpretq_s8_u64(vcombine_u64(vcreate_u64(u0_0), vcreate_u64(u1_0)));
7755            let w1_0 = vreinterpretq_s8_u64(vcombine_u64(vcreate_u64(u2_0), vcreate_u64(u3_0)));
7756            let w0_1 = vreinterpretq_s8_u64(vcombine_u64(vcreate_u64(u0_1), vcreate_u64(u1_1)));
7757            let w1_1 = vreinterpretq_s8_u64(vcombine_u64(vcreate_u64(u2_1), vcreate_u64(u3_1)));
7758
7759            let x0_0 = vld1q_s8(xq_ptr.add(gi * GROUP_SIZE));
7760            let x1_0 = vld1q_s8(xq_ptr.add(gi * GROUP_SIZE + 16));
7761            let x0_1 = vld1q_s8(xq_ptr.add((gi + 1) * GROUP_SIZE));
7762            let x1_1 = vld1q_s8(xq_ptr.add((gi + 1) * GROUP_SIZE + 16));
7763
7764            let (mut a0_0, mut a1_0) = (vdupq_n_s32(0), vdupq_n_s32(0));
7765            let (mut a0_1, mut a1_1) = (vdupq_n_s32(0), vdupq_n_s32(0));
7766            asm!(
7767                "sdot {a0_0:v}.4s, {w0_0:v}.16b, {x0_0:v}.16b",
7768                "sdot {a1_0:v}.4s, {w1_0:v}.16b, {x1_0:v}.16b",
7769                "sdot {a0_1:v}.4s, {w0_1:v}.16b, {x0_1:v}.16b",
7770                "sdot {a1_1:v}.4s, {w1_1:v}.16b, {x1_1:v}.16b",
7771                a0_0 = inout(vreg) a0_0, a1_0 = inout(vreg) a1_0,
7772                a0_1 = inout(vreg) a0_1, a1_1 = inout(vreg) a1_1,
7773                w0_0 = in(vreg) w0_0, x0_0 = in(vreg) x0_0, w1_0 = in(vreg) w1_0, x1_0 = in(vreg) x1_0,
7774                w0_1 = in(vreg) w0_1, x0_1 = in(vreg) x0_1, w1_1 = in(vreg) w1_1, x1_1 = in(vreg) x1_1,
7775                options(pure, nomem, nostack),
7776            );
7777            let d0 = vaddvq_s32(vaddq_s32(a0_0, a1_0));
7778            let d1 = vaddvq_s32(vaddq_s32(a0_1, a1_1));
7779            acc += d0 as f32 * s0 + d1 as f32 * s1;
7780            gi += 2;
7781        }
7782
7783        if gi < gpr {
7784            let off = row_off + gi * TILE;
7785            let s = f16_to_f32(u16::from_le_bytes([
7786                *bytes_ptr.add(off),
7787                *bytes_ptr.add(off + 1),
7788            ]));
7789            let (u0, u1, u2, u3) = q1t_unpack_reg_u64s(bytes_ptr.add(off + 2));
7790            let w0 = vreinterpretq_s8_u64(vcombine_u64(vcreate_u64(u0), vcreate_u64(u1)));
7791            let w1 = vreinterpretq_s8_u64(vcombine_u64(vcreate_u64(u2), vcreate_u64(u3)));
7792            let x0 = vld1q_s8(xq_ptr.add(gi * GROUP_SIZE));
7793            let x1 = vld1q_s8(xq_ptr.add(gi * GROUP_SIZE + 16));
7794            let (mut a0, mut a1) = (vdupq_n_s32(0), vdupq_n_s32(0));
7795            asm!(
7796                "sdot {a0:v}.4s, {w0:v}.16b, {x0:v}.16b",
7797                "sdot {a1:v}.4s, {w1:v}.16b, {x1:v}.16b",
7798                a0 = inout(vreg) a0, a1 = inout(vreg) a1,
7799                w0 = in(vreg) w0, x0 = in(vreg) x0, w1 = in(vreg) w1, x1 = in(vreg) x1,
7800                options(pure, nomem, nostack),
7801            );
7802            let d = vaddvq_s32(vaddq_s32(a0, a1));
7803            acc += d as f32 * s;
7804        }
7805        acc
7806    }
7807}
7808
7809/// x86 AVX2 mirror of `q1t_dot_row_sdot` (maddubs int8 dot per group).
7810#[cfg(target_arch = "x86_64")]
7811#[target_feature(enable = "avx2")]
7812unsafe fn q1t_dot_row_avx2(bytes: &[u8], r: usize, gpr: usize, xq: &[i8]) -> f32 {
7813    use core::arch::x86_64::*;
7814    unsafe {
7815        const TILE: usize = cortiq_core::quant::Q1T_TILE;
7816        let mut acc = 0f32;
7817        let bytes_ptr = bytes.as_ptr();
7818        let xq_ptr = xq.as_ptr();
7819        let row_off = r * gpr * TILE;
7820
7821        let ones = _mm256_set1_epi16(1);
7822        for gi in 0..gpr {
7823            let off = row_off + gi * TILE;
7824            let s = f16_to_f32(u16::from_le_bytes([
7825                *bytes_ptr.add(off),
7826                *bytes_ptr.add(off + 1),
7827            ]));
7828            let (u0, u1, u2, u3) = q1t_unpack_reg_u64s(bytes_ptr.add(off + 2));
7829            let wv = _mm256_set_epi64x(u3 as i64, u2 as i64, u1 as i64, u0 as i64);
7830            let xv = _mm256_loadu_si256(xq_ptr.add(gi * GROUP_SIZE) as *const __m256i);
7831            let p16 = _mm256_maddubs_epi16(_mm256_abs_epi8(wv), _mm256_sign_epi8(xv, wv));
7832            let d256 = _mm256_madd_epi16(p16, ones);
7833            let d128 = _mm_add_epi32(
7834                _mm256_castsi256_si128(d256),
7835                _mm256_extracti128_si256(d256, 1),
7836            );
7837            let d64 = _mm_add_epi32(d128, _mm_shuffle_epi32(d128, 0xee));
7838            let d32 = _mm_cvtsi128_si32(_mm_add_epi32(d64, _mm_shuffle_epi32(d64, 0x55)));
7839            acc += d32 as f32 * s;
7840        }
7841        acc
7842    }
7843}
7844
7845/// VNNI twin of `q1t_dot_row_avx2` (see `dpbusd_hsum`).
7846#[cfg(target_arch = "x86_64")]
7847#[target_feature(enable = "avx2,avx512f,avx512bw,avx512vl,avx512vnni")]
7848unsafe fn q1t_dot_row_vnni(bytes: &[u8], r: usize, gpr: usize, xq: &[i8]) -> f32 {
7849    use core::arch::x86_64::*;
7850    // SAFETY: same tile/xq contracts as `q1t_dot_row_avx2`.
7851    unsafe {
7852        const TILE: usize = cortiq_core::quant::Q1T_TILE;
7853        let mut acc = 0f32;
7854        let bytes_ptr = bytes.as_ptr();
7855        let xq_ptr = xq.as_ptr();
7856        let row_off = r * gpr * TILE;
7857        for gi in 0..gpr {
7858            let off = row_off + gi * TILE;
7859            let s = f16_to_f32(u16::from_le_bytes([
7860                *bytes_ptr.add(off),
7861                *bytes_ptr.add(off + 1),
7862            ]));
7863            let (u0, u1, u2, u3) = q1t_unpack_reg_u64s(bytes_ptr.add(off + 2));
7864            let wv = _mm256_set_epi64x(u3 as i64, u2 as i64, u1 as i64, u0 as i64);
7865            let xv = _mm256_loadu_si256(xq_ptr.add(gi * GROUP_SIZE) as *const __m256i);
7866            let d = dpbusd_hsum(_mm256_abs_epi8(wv), _mm256_sign_epi8(xv, wv));
7867            acc += d as f32 * s;
7868        }
7869        acc
7870    }
7871}
7872
7873/// Per-row int8 base dot, dispatched once per row (matvec decode hot path).
7874/// Callers are gated by `a8w8_enabled()`, so the target-feature kernels are
7875/// reachable.
7876#[inline]
7877fn q1t_dot_row_i8(bytes: &[u8], r: usize, gpr: usize, xq: &[i8]) -> f32 {
7878    #[cfg(target_arch = "aarch64")]
7879    unsafe {
7880        return q1t_dot_row_sdot(bytes, r, gpr, xq);
7881    }
7882    #[cfg(target_arch = "x86_64")]
7883    unsafe {
7884        if vnni_tiles_enabled() {
7885            return q1t_dot_row_vnni(bytes, r, gpr, xq);
7886        }
7887        return q1t_dot_row_avx2(bytes, r, gpr, xq);
7888    }
7889    #[allow(unreachable_code)]
7890    {
7891        const TILE: usize = cortiq_core::quant::Q1T_TILE;
7892        let mut acc = 0f32;
7893        let mut sg = [0i8; GROUP_SIZE + 8]; // +8 slack for the u64-store unpack
7894        for gi in 0..gpr {
7895            let off = (r * gpr + gi) * TILE;
7896            let s = f16_to_f32(u16::from_le_bytes([bytes[off], bytes[off + 1]]));
7897            q1t_unpack_group_i8(bytes.as_ptr().wrapping_add(off + 2), &mut sg);
7898            let mut d = 0i32;
7899            for k in 0..GROUP_SIZE {
7900                d += sg[k] as i32 * xq[gi * GROUP_SIZE + k] as i32;
7901            }
7902            acc += d as f32 * s;
7903        }
7904        acc
7905    }
7906}
7907
7908/// Σ over a row's outliers of `value·x[col]` — the correction that adds the
7909/// overlay's exact weights on top of the base dot. INVARIANT: the encoder
7910/// writes ternary code 0 at every outlier position (`quantize_q1t`), so the
7911/// base contributes nothing there and this is a plain `value·x`, not
7912/// `(value − base)·x` — no scattered per-outlier scale read. Row `r`'s entries
7913/// are the contiguous slice `[row_ptr[r], row_ptr[r+1])`, so no binary search.
7914fn q1t_row_outlier_correction(
7915    bytes: &[u8],
7916    r: usize,
7917    rp_off: usize,
7918    entries_off: usize,
7919    has_ov: bool,
7920    x: &[f32],
7921) -> f32 {
7922    if !has_ov {
7923        return 0.0;
7924    }
7925    let (c0, c1) = (
7926        q1t_rowptr(bytes, rp_off, r),
7927        q1t_rowptr(bytes, rp_off, r + 1),
7928    );
7929    let mut corr = 0f32;
7930    for p in c0..c1 {
7931        let e = entries_off + p * 4;
7932        let col = u16::from_le_bytes([bytes[e], bytes[e + 1]]) as usize;
7933        let val = f16_to_f32(u16::from_le_bytes([bytes[e + 2], bytes[e + 3]]));
7934        corr += val * x[col];
7935    }
7936    corr
7937}
7938
7939/// Dequantize one q1t row into `buf[..cols]` via the sign LUT (no division),
7940/// then apply the row's outliers (its `[row_ptr[r], row_ptr[r+1])` slice).
7941/// Used by the batched (prefill) path where the decode amortizes over the batch.
7942fn q1t_dequant_row(
7943    bytes: &[u8],
7944    r: usize,
7945    gpr: usize,
7946    rp_off: usize,
7947    entries_off: usize,
7948    has_ov: bool,
7949    buf: &mut [f32],
7950) {
7951    const TILE: usize = cortiq_core::quant::Q1T_TILE;
7952    for g in 0..gpr {
7953        let off = (r * gpr + g) * TILE;
7954        let s = f16_to_f32(u16::from_le_bytes([bytes[off], bytes[off + 1]]));
7955        let codes = &bytes[off + 2..off + TILE];
7956        let bc = g * GROUP_SIZE;
7957        // 6 full bytes (30 codes) + a 7th byte holding the last 2.
7958        for bi in 0..6 {
7959            let lut = &SIGN5[codes[bi] as usize];
7960            let d = &mut buf[bc + bi * 5..bc + bi * 5 + 5];
7961            for i in 0..5 {
7962                d[i] = lut[i] * s;
7963            }
7964        }
7965        let lut = &SIGN5[codes[6] as usize];
7966        buf[bc + 30] = lut[0] * s;
7967        buf[bc + 31] = lut[1] * s;
7968    }
7969    if !has_ov {
7970        return;
7971    }
7972    let (c0, c1) = (
7973        q1t_rowptr(bytes, rp_off, r),
7974        q1t_rowptr(bytes, rp_off, r + 1),
7975    );
7976    for p in c0..c1 {
7977        let e = entries_off + p * 4;
7978        let col = u16::from_le_bytes([bytes[e], bytes[e + 1]]) as usize;
7979        buf[col] = f16_to_f32(u16::from_le_bytes([bytes[e + 2], bytes[e + 3]]));
7980    }
7981}
7982
7983/// Add the sparse outlier overlay onto a base dot already in `out` (the GPU
7984/// computes the ternary base; the overlay stays on the CPU — its entries are
7985/// few and its per-row gather doesn't vectorize on the GPU). Row-parallel.
7986fn q1t_add_overlay(
7987    bytes: &[u8],
7988    x: &[f32],
7989    rows: usize,
7990    cols: usize,
7991    out: &mut [f32],
7992    pool: Option<&Pool>,
7993) {
7994    const TILE: usize = cortiq_core::quant::Q1T_TILE;
7995    let gpr = cols / GROUP_SIZE;
7996    let (rp_off, ent_off, has_ov) = q1t_overlay(bytes, rows * gpr * TILE, rows);
7997    if !has_ov {
7998        return;
7999    }
8000    let out_addr = SendMut(out.as_mut_ptr());
8001    let run = move |start: usize, end: usize| {
8002        for r in start..end {
8003            let corr = q1t_row_outlier_correction(bytes, r, rp_off, ent_off, has_ov, x);
8004            // SAFETY: disjoint rows; add onto the base the GPU already wrote.
8005            unsafe { *out_addr.at(r) += corr };
8006        }
8007    };
8008    dispatch_rows(pool, rows, &run);
8009}
8010
8011/// Q1T row range via the A8W8 int8 path — shared activation split,
8012/// per-row: base SDOT dot + outlier correction + overlay.
8013#[allow(clippy::too_many_arguments)]
8014fn q1t_range_a8w8(
8015    bytes: &[u8],
8016    gpr: usize,
8017    rp_off: usize,
8018    ent_off: usize,
8019    has_ov: bool,
8020    act: &SplitAct,
8021    x: &[f32],
8022    out: SendMut,
8023    start: usize,
8024    end: usize,
8025) {
8026    for r in start..end {
8027        let mut acc = q1t_dot_row_i8(bytes, r, gpr, &act.xq) * act.sx;
8028        for &(j, xv) in &act.outliers {
8029            acc += q1t_base_weight(bytes, r, gpr, j) * xv;
8030        }
8031        acc += q1t_row_outlier_correction(bytes, r, rp_off, ent_off, has_ov, x);
8032        // SAFETY: disjoint row ranges per worker.
8033        unsafe { *out.at(r) = acc };
8034    }
8035}
8036
8037/// Q1T row range via the f32 path (no SDOT) — for matvec_many batched
8038/// dispatch when a8w8 is unavailable.
8039#[allow(clippy::too_many_arguments)]
8040fn q1t_range_f32_batch(
8041    bytes: &[u8],
8042    gpr: usize,
8043    rp_off: usize,
8044    ent_off: usize,
8045    has_ov: bool,
8046    x: &[f32],
8047    out: SendMut,
8048    start: usize,
8049    end: usize,
8050) {
8051    const TILE: usize = cortiq_core::quant::Q1T_TILE;
8052    let mut sg = [0f32; GROUP_SIZE];
8053    for r in start..end {
8054        let mut acc = 0f32;
8055        for g in 0..gpr {
8056            let off = (r * gpr + g) * TILE;
8057            let s = f16_to_f32(u16::from_le_bytes([bytes[off], bytes[off + 1]]));
8058            let codes = &bytes[off + 2..off + TILE];
8059            let xg = &x[g * GROUP_SIZE..g * GROUP_SIZE + GROUP_SIZE];
8060            for bi in 0..6 {
8061                sg[bi * 5..bi * 5 + 5].copy_from_slice(&SIGN5[codes[bi] as usize]);
8062            }
8063            let lut = &SIGN5[codes[6] as usize];
8064            sg[30] = lut[0];
8065            sg[31] = lut[1];
8066            let mut gsum = 0f32;
8067            for k in 0..GROUP_SIZE {
8068                gsum += sg[k] * xg[k];
8069            }
8070            acc += s * gsum;
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/// Ternary (q1t) matvec — decode+dot straight from mmap, one group at a time:
8079/// no per-ROW buffer, no division (the sign LUT), and a tiny per-group sign
8080/// buffer so the 32-wide dot vectorizes. This is the decode hot path.
8081fn q1t_matvec(
8082    bytes: &[u8],
8083    x: &[f32],
8084    rows: usize,
8085    cols: usize,
8086    out: &mut [f32],
8087    pool: Option<&Pool>,
8088) {
8089    debug_assert_eq!(out.len(), rows);
8090    const TILE: usize = cortiq_core::quant::Q1T_TILE;
8091    let gpr = cols / GROUP_SIZE;
8092    let (rp_off, ent_off, has_ov) = q1t_overlay(bytes, rows * gpr * TILE, rows);
8093    let out_addr = SendMut(out.as_mut_ptr());
8094    // int8 SDOT base dot (ARM dotprod): ~4× the f32 arithmetic. x → i8 once
8095    // (`split_act`), activation outliers added back exactly in f32, weight
8096    // overlay on top. ARM SDOT / x86 AVX2; CMF_SDOT=0 keeps the exact f32 path.
8097    if a8w8_enabled() {
8098        let act = split_act(x);
8099        let act = &act;
8100        let run = move |start: usize, end: usize| {
8101            for r in start..end {
8102                let mut acc = q1t_dot_row_i8(bytes, r, gpr, &act.xq) * act.sx;
8103                for &(j, xv) in &act.outliers {
8104                    acc += q1t_base_weight(bytes, r, gpr, j) * xv;
8105                }
8106                acc += q1t_row_outlier_correction(bytes, r, rp_off, ent_off, has_ov, x);
8107                // SAFETY: disjoint row ranges per worker.
8108                unsafe { *out_addr.at(r) = acc };
8109            }
8110        };
8111        dispatch_rows(pool, rows, &run);
8112        return;
8113    }
8114    let run = move |start: usize, end: usize| {
8115        // Per-group signs, unpacked contiguously so the dot below is a clean
8116        // 32-wide reduction the autovectorizer turns into f32x4 FMAs — the
8117        // 5-values-per-byte base-3 layout won't SIMD in place.
8118        let mut sg = [0f32; GROUP_SIZE];
8119        for r in start..end {
8120            let mut acc = 0f32;
8121            for g in 0..gpr {
8122                let off = (r * gpr + g) * TILE;
8123                let s = f16_to_f32(u16::from_le_bytes([bytes[off], bytes[off + 1]]));
8124                let codes = &bytes[off + 2..off + TILE];
8125                let xg = &x[g * GROUP_SIZE..g * GROUP_SIZE + GROUP_SIZE];
8126                for bi in 0..6 {
8127                    sg[bi * 5..bi * 5 + 5].copy_from_slice(&SIGN5[codes[bi] as usize]);
8128                }
8129                let lut = &SIGN5[codes[6] as usize];
8130                sg[30] = lut[0];
8131                sg[31] = lut[1];
8132                let mut gsum = 0f32;
8133                for k in 0..GROUP_SIZE {
8134                    gsum += sg[k] * xg[k];
8135                }
8136                acc += s * gsum;
8137            }
8138            acc += q1t_row_outlier_correction(bytes, r, rp_off, ent_off, has_ov, x);
8139            unsafe { *out_addr.at(r) = acc };
8140        }
8141    };
8142    dispatch_rows(pool, rows, &run);
8143}
8144
8145/// Fused-pair twin of `q1t_dot_row_sdot`: ONE register unpack of the
8146/// ternary codes serves BOTH activation streams (the unpack chain is
8147/// the dominant per-row cost — MTP verify pairs paid it twice). Per
8148/// stream the group order and f32 accumulation match the single-row
8149/// kernel exactly, so pair == 2×matvec bit-for-bit.
8150#[cfg(target_arch = "aarch64")]
8151#[target_feature(enable = "neon,dotprod")]
8152unsafe fn q1t_dot_row_sdot2(bytes: &[u8], r: usize, gpr: usize, xa: &[i8], xb: &[i8]) -> [f32; 2] {
8153    use core::arch::aarch64::*;
8154    use core::arch::asm;
8155    // SAFETY: same slice-length contracts as `q1t_dot_row_sdot`, ×2.
8156    unsafe {
8157        const TILE: usize = cortiq_core::quant::Q1T_TILE;
8158        let bytes_ptr = bytes.as_ptr();
8159        let row_off = r * gpr * TILE;
8160        let xp = [xa.as_ptr(), xb.as_ptr()];
8161        let mut acc = [0f32; 2];
8162        macro_rules! sdot2 {
8163            ($w0:expr, $w1:expr, $x:expr) => {{
8164                let x0 = vld1q_s8($x);
8165                let x1 = vld1q_s8($x.add(16));
8166                let (mut a0, mut a1) = (vdupq_n_s32(0), vdupq_n_s32(0));
8167                asm!(
8168                    "sdot {a0:v}.4s, {w0:v}.16b, {x0:v}.16b",
8169                    "sdot {a1:v}.4s, {w1:v}.16b, {x1:v}.16b",
8170                    a0 = inout(vreg) a0, a1 = inout(vreg) a1,
8171                    w0 = in(vreg) $w0, x0 = in(vreg) x0, w1 = in(vreg) $w1, x1 = in(vreg) x1,
8172                    options(pure, nomem, nostack),
8173                );
8174                vaddvq_s32(vaddq_s32(a0, a1))
8175            }};
8176        }
8177        let gpr2 = gpr & !1;
8178        let mut gi = 0;
8179        while gi < gpr2 {
8180            let off0 = row_off + gi * TILE;
8181            let off1 = off0 + TILE;
8182            let s0 = f16_to_f32(u16::from_le_bytes([
8183                *bytes_ptr.add(off0),
8184                *bytes_ptr.add(off0 + 1),
8185            ]));
8186            let s1 = f16_to_f32(u16::from_le_bytes([
8187                *bytes_ptr.add(off1),
8188                *bytes_ptr.add(off1 + 1),
8189            ]));
8190            let (u0_0, u1_0, u2_0, u3_0) = q1t_unpack_reg_u64s(bytes_ptr.add(off0 + 2));
8191            let (u0_1, u1_1, u2_1, u3_1) = q1t_unpack_reg_u64s(bytes_ptr.add(off1 + 2));
8192            let w0_0 = vreinterpretq_s8_u64(vcombine_u64(vcreate_u64(u0_0), vcreate_u64(u1_0)));
8193            let w1_0 = vreinterpretq_s8_u64(vcombine_u64(vcreate_u64(u2_0), vcreate_u64(u3_0)));
8194            let w0_1 = vreinterpretq_s8_u64(vcombine_u64(vcreate_u64(u0_1), vcreate_u64(u1_1)));
8195            let w1_1 = vreinterpretq_s8_u64(vcombine_u64(vcreate_u64(u2_1), vcreate_u64(u3_1)));
8196            for k in 0..2 {
8197                let d0 = sdot2!(w0_0, w1_0, xp[k].add(gi * GROUP_SIZE));
8198                let d1 = sdot2!(w0_1, w1_1, xp[k].add((gi + 1) * GROUP_SIZE));
8199                acc[k] += d0 as f32 * s0 + d1 as f32 * s1;
8200            }
8201            gi += 2;
8202        }
8203        if gi < gpr {
8204            let off = row_off + gi * TILE;
8205            let s = f16_to_f32(u16::from_le_bytes([
8206                *bytes_ptr.add(off),
8207                *bytes_ptr.add(off + 1),
8208            ]));
8209            let (u0, u1, u2, u3) = q1t_unpack_reg_u64s(bytes_ptr.add(off + 2));
8210            let w0 = vreinterpretq_s8_u64(vcombine_u64(vcreate_u64(u0), vcreate_u64(u1)));
8211            let w1 = vreinterpretq_s8_u64(vcombine_u64(vcreate_u64(u2), vcreate_u64(u3)));
8212            for k in 0..2 {
8213                let d = sdot2!(w0, w1, xp[k].add(gi * GROUP_SIZE));
8214                acc[k] += d as f32 * s;
8215            }
8216        }
8217        acc
8218    }
8219}
8220
8221/// Fused Q1T pair matvec: ONE pass over the rows serves both
8222/// activation streams — on ARM the ternary register unpack happens
8223/// once per tile pair (`q1t_dot_row_sdot2`); elsewhere the second dot
8224/// rides the row's L1-warm tile bytes. Per stream the math matches
8225/// `q1t_matvec` exactly.
8226fn q1t_matvec2(
8227    bytes: &[u8],
8228    x1: &[f32],
8229    x2: &[f32],
8230    rows: usize,
8231    cols: usize,
8232    o1: &mut [f32],
8233    o2: &mut [f32],
8234    pool: Option<&Pool>,
8235) {
8236    debug_assert_eq!(o1.len(), rows);
8237    debug_assert_eq!(o2.len(), rows);
8238    const TILE: usize = cortiq_core::quant::Q1T_TILE;
8239    let gpr = cols / GROUP_SIZE;
8240    let (rp_off, ent_off, has_ov) = q1t_overlay(bytes, rows * gpr * TILE, rows);
8241    let out1 = SendMut(o1.as_mut_ptr());
8242    let out2 = SendMut(o2.as_mut_ptr());
8243    if a8w8_enabled() {
8244        let a1 = split_act(x1);
8245        let a2 = split_act(x2);
8246        let (a1, a2) = (&a1, &a2);
8247        let run = move |start: usize, end: usize| {
8248            for r in start..end {
8249                #[cfg(target_arch = "aarch64")]
8250                // a8w8 on aarch64 ⇔ sdot_enabled(), so the kernel's
8251                // target features are present.
8252                let ds = unsafe { q1t_dot_row_sdot2(bytes, r, gpr, &a1.xq, &a2.xq) };
8253                #[cfg(not(target_arch = "aarch64"))]
8254                let ds = [
8255                    q1t_dot_row_i8(bytes, r, gpr, &a1.xq),
8256                    q1t_dot_row_i8(bytes, r, gpr, &a2.xq),
8257                ];
8258                let mut acc1 = ds[0] * a1.sx;
8259                for &(j, xv) in &a1.outliers {
8260                    acc1 += q1t_base_weight(bytes, r, gpr, j) * xv;
8261                }
8262                acc1 += q1t_row_outlier_correction(bytes, r, rp_off, ent_off, has_ov, x1);
8263                let mut acc2 = ds[1] * a2.sx;
8264                for &(j, xv) in &a2.outliers {
8265                    acc2 += q1t_base_weight(bytes, r, gpr, j) * xv;
8266                }
8267                acc2 += q1t_row_outlier_correction(bytes, r, rp_off, ent_off, has_ov, x2);
8268                // SAFETY: disjoint row ranges per worker.
8269                unsafe {
8270                    *out1.at(r) = acc1;
8271                    *out2.at(r) = acc2;
8272                }
8273            }
8274        };
8275        dispatch_rows(pool, rows, &run);
8276        return;
8277    }
8278    let run = move |start: usize, end: usize| {
8279        // Exact path (CMF_SDOT=0): unpack the sign LUT once per group,
8280        // dot both streams — same op order per stream as `q1t_matvec`.
8281        let mut sg = [0f32; GROUP_SIZE];
8282        for r in start..end {
8283            let mut acc1 = 0f32;
8284            let mut acc2 = 0f32;
8285            for g in 0..gpr {
8286                let off = (r * gpr + g) * TILE;
8287                let s = f16_to_f32(u16::from_le_bytes([bytes[off], bytes[off + 1]]));
8288                let codes = &bytes[off + 2..off + TILE];
8289                for bi in 0..6 {
8290                    sg[bi * 5..bi * 5 + 5].copy_from_slice(&SIGN5[codes[bi] as usize]);
8291                }
8292                let lut = &SIGN5[codes[6] as usize];
8293                sg[30] = lut[0];
8294                sg[31] = lut[1];
8295                let xg1 = &x1[g * GROUP_SIZE..g * GROUP_SIZE + GROUP_SIZE];
8296                let xg2 = &x2[g * GROUP_SIZE..g * GROUP_SIZE + GROUP_SIZE];
8297                let mut gsum1 = 0f32;
8298                for k in 0..GROUP_SIZE {
8299                    gsum1 += sg[k] * xg1[k];
8300                }
8301                acc1 += s * gsum1;
8302                let mut gsum2 = 0f32;
8303                for k in 0..GROUP_SIZE {
8304                    gsum2 += sg[k] * xg2[k];
8305                }
8306                acc2 += s * gsum2;
8307            }
8308            acc1 += q1t_row_outlier_correction(bytes, r, rp_off, ent_off, has_ov, x1);
8309            acc2 += q1t_row_outlier_correction(bytes, r, rp_off, ent_off, has_ov, x2);
8310            // SAFETY: disjoint row ranges per worker.
8311            unsafe {
8312                *out1.at(r) = acc1;
8313                *out2.at(r) = acc2;
8314            }
8315        }
8316    };
8317    dispatch_rows(pool, rows, &run);
8318}
8319
8320/// Ternary (q1t) matmat (prefill) — dequant each row once, dot the whole
8321/// batch against it (amortizes the per-row decode).
8322fn q1t_matmat(
8323    bytes: &[u8],
8324    xs: &[f32],
8325    b: usize,
8326    rows: usize,
8327    cols: usize,
8328    out: &mut [f32],
8329    pool: Option<&Pool>,
8330) {
8331    debug_assert_eq!(out.len(), b * rows);
8332    const TILE: usize = cortiq_core::quant::Q1T_TILE;
8333    let gpr = cols / GROUP_SIZE;
8334    let (rp_off, ent_off, has_ov) = q1t_overlay(bytes, rows * gpr * TILE, rows);
8335    let out_addr = SendMut(out.as_mut_ptr());
8336    // int8 prefill (ARM SDOT / x86 AVX2): quantize the B inputs once, unpack
8337    // each weight row's signs to i8 ONCE, then int8-dot against every input —
8338    // the row sign-decode amortizes over the whole batch. CMF_SDOT=0 → f32.
8339    if a8w8_enabled() {
8340        let acts: Vec<SplitAct> = (0..b)
8341            .map(|bi| split_act(&xs[bi * cols..(bi + 1) * cols]))
8342            .collect();
8343        let acts = &acts;
8344        let run = move |start: usize, end: usize| {
8345            let mut sg = vec![0i8; cols + 8]; // row signs, i8 (+8 unpack slack)
8346            let mut sc = vec![0f32; gpr]; // per-group scales
8347            let mut accs = vec![0f32; b]; // per-batch accumulators, reused per row
8348            for r in start..end {
8349                for g in 0..gpr {
8350                    let off = (r * gpr + g) * TILE;
8351                    sc[g] = f16_to_f32(u16::from_le_bytes([bytes[off], bytes[off + 1]]));
8352                    q1t_unpack_group_i8(
8353                        bytes.as_ptr().wrapping_add(off + 2),
8354                        &mut sg[g * GROUP_SIZE..],
8355                    );
8356                }
8357                for bi in 0..b {
8358                    let act = &acts[bi];
8359                    let mut isum = 0f32;
8360                    for g in 0..gpr {
8361                        let d = q1t_i8dot32(
8362                            sg.as_ptr().wrapping_add(g * GROUP_SIZE),
8363                            act.xq.as_ptr().wrapping_add(g * GROUP_SIZE),
8364                        );
8365                        isum += d as f32 * sc[g];
8366                    }
8367                    let mut acc = isum * act.sx;
8368                    for &(j, xv) in &act.outliers {
8369                        acc += q1t_base_weight(bytes, r, gpr, j) * xv;
8370                    }
8371                    accs[bi] = acc;
8372                }
8373                // Overlay ONCE per row for the whole batch: read each (col, val)
8374                // from mmap a single time (was b× — the re-read dominated prefill)
8375                // and fan it out over the batch via the cached inputs.
8376                if has_ov {
8377                    let (c0, c1) = (
8378                        q1t_rowptr(bytes, rp_off, r),
8379                        q1t_rowptr(bytes, rp_off, r + 1),
8380                    );
8381                    for p in c0..c1 {
8382                        let e = ent_off + p * 4;
8383                        let col = u16::from_le_bytes([bytes[e], bytes[e + 1]]) as usize;
8384                        let val = f16_to_f32(u16::from_le_bytes([bytes[e + 2], bytes[e + 3]]));
8385                        for bi in 0..b {
8386                            accs[bi] += val * xs[bi * cols + col];
8387                        }
8388                    }
8389                }
8390                for bi in 0..b {
8391                    unsafe { *out_addr.at(bi * rows + r) = accs[bi] };
8392                }
8393            }
8394        };
8395        dispatch_rows(pool, rows, &run);
8396        return;
8397    }
8398    let run = move |start: usize, end: usize| {
8399        let mut buf = vec![0f32; cols];
8400        for r in start..end {
8401            q1t_dequant_row(bytes, r, gpr, rp_off, ent_off, has_ov, &mut buf);
8402            for bi in 0..b {
8403                let xr = &xs[bi * cols..(bi + 1) * cols];
8404                let mut acc = 0f32;
8405                for j in 0..cols {
8406                    acc += buf[j] * xr[j];
8407                }
8408                unsafe { *out_addr.at(bi * rows + r) = acc };
8409            }
8410        }
8411    };
8412    dispatch_rows(pool, rows, &run);
8413}
8414
8415fn q1_matvec(
8416    bytes: &[u8],
8417    x: &[f32],
8418    rows: usize,
8419    cols: usize,
8420    out: &mut [f32],
8421    pool: Option<&Pool>,
8422) {
8423    debug_assert_eq!(out.len(), rows);
8424    let gpr = cols / GROUP_SIZE;
8425    let out_addr = SendMut(out.as_mut_ptr());
8426    if a8w8_enabled() {
8427        let act = split_act(x);
8428        let gsum = q1_group_sums(&act.xq, gpr);
8429        let (act, gsum) = (&act, &gsum);
8430        let run = move |start: usize, end: usize| {
8431            q1_range_a8w8(bytes, gpr, act, gsum, out_addr, start, end)
8432        };
8433        dispatch_rows(pool, rows, &run);
8434        return;
8435    }
8436    let run = move |start: usize, end: usize| q1_range_f32(bytes, gpr, x, out_addr, start, end);
8437    dispatch_rows(pool, rows, &run);
8438}
8439
8440/// Fused two-input q1 matvec (weights read once per pair).
8441#[allow(clippy::too_many_arguments)]
8442fn q1_matvec2(
8443    bytes: &[u8],
8444    x1: &[f32],
8445    x2: &[f32],
8446    rows: usize,
8447    cols: usize,
8448    o1: &mut [f32],
8449    o2: &mut [f32],
8450    pool: Option<&Pool>,
8451) {
8452    let gpr = cols / GROUP_SIZE;
8453    let p1 = SendMut(o1.as_mut_ptr());
8454    let p2 = SendMut(o2.as_mut_ptr());
8455    if a8w8_enabled() {
8456        let a1 = split_act(x1);
8457        let a2 = split_act(x2);
8458        let g1 = q1_group_sums(&a1.xq, gpr);
8459        let g2 = q1_group_sums(&a2.xq, gpr);
8460        let (a1, a2, g1, g2) = (&a1, &a2, &g1, &g2);
8461        let run = move |start: usize, end: usize| {
8462            for r in start..end {
8463                let mut v1 = dot_q1_row_i8(bytes, r, gpr, &a1.xq, g1) * a1.sx;
8464                let mut v2 = dot_q1_row_i8(bytes, r, gpr, &a2.xq, g2) * a2.sx;
8465                for &(j, xv) in &a1.outliers {
8466                    let (w, s) = q1_outlier(bytes, r, gpr, j);
8467                    v1 += w * s * xv;
8468                }
8469                for &(j, xv) in &a2.outliers {
8470                    let (w, s) = q1_outlier(bytes, r, gpr, j);
8471                    v2 += w * s * xv;
8472                }
8473                // SAFETY: disjoint row ranges per worker.
8474                unsafe {
8475                    *p1.at(r) = v1;
8476                    *p2.at(r) = v2;
8477                }
8478            }
8479        };
8480        dispatch_rows(pool, rows, &run);
8481        return;
8482    }
8483    let run = move |start: usize, end: usize| {
8484        for r in start..end {
8485            // SAFETY: disjoint row ranges per worker.
8486            unsafe {
8487                *p1.at(r) = q1_row_exact(bytes, r, gpr, x1);
8488                *p2.at(r) = q1_row_exact(bytes, r, gpr, x2);
8489            }
8490        }
8491    };
8492    dispatch_rows(pool, rows, &run);
8493}
8494
8495/// Batched q1 matmat: each row's tiles stream once per microbatch.
8496#[allow(clippy::too_many_arguments)]
8497fn q1_matmat(
8498    bytes: &[u8],
8499    xs_all: &[f32],
8500    b: usize,
8501    rows: usize,
8502    cols: usize,
8503    out: &mut [f32],
8504    pool: Option<&Pool>,
8505) {
8506    debug_assert_eq!(out.len(), b * rows);
8507    let gpr = cols / GROUP_SIZE;
8508    let out_addr = SendMut(out.as_mut_ptr());
8509    if a8w8_enabled() {
8510        let acts: Vec<(SplitAct, Vec<i32>)> = (0..b)
8511            .map(|bi| {
8512                let act = split_act(&xs_all[bi * cols..(bi + 1) * cols]);
8513                let gsum = q1_group_sums(&act.xq, gpr);
8514                (act, gsum)
8515            })
8516            .collect();
8517        let acts = &acts;
8518        #[cfg(target_arch = "x86_64")]
8519        let blocked_ok = avx2_enabled() && blocked_enabled();
8520        #[cfg(target_arch = "aarch64")]
8521        let blocked_ok = sdot_enabled() && blocked_enabled();
8522        let run = move |start: usize, end: usize| {
8523            for r in start..end {
8524                let mut bi = 0usize;
8525                // Blocked 1×4: the unpacked bit mask serves four
8526                // activation streams per group.
8527                #[cfg(target_arch = "aarch64")]
8528                if blocked_ok {
8529                    while bi + 4 <= acts.len() {
8530                        let xs = [
8531                            acts[bi].0.xq.as_slice(),
8532                            acts[bi + 1].0.xq.as_slice(),
8533                            acts[bi + 2].0.xq.as_slice(),
8534                            acts[bi + 3].0.xq.as_slice(),
8535                        ];
8536                        let gs = [
8537                            acts[bi].1.as_slice(),
8538                            acts[bi + 1].1.as_slice(),
8539                            acts[bi + 2].1.as_slice(),
8540                            acts[bi + 3].1.as_slice(),
8541                        ];
8542                        let d = unsafe { dot_q1_row_1x4_sdot(bytes, r, gpr, xs, gs) };
8543                        for k in 0..4 {
8544                            let (act, _) = &acts[bi + k];
8545                            let mut acc = d[k] * act.sx;
8546                            for &(j, xv) in &act.outliers {
8547                                let (w, sc) = q1_outlier(bytes, r, gpr, j);
8548                                acc += w * sc * xv;
8549                            }
8550                            // SAFETY: disjoint (bi, r) cells per worker.
8551                            unsafe { *out_addr.at((bi + k) * rows + r) = acc };
8552                        }
8553                        bi += 4;
8554                    }
8555                }
8556                #[cfg(target_arch = "x86_64")]
8557                if blocked_ok {
8558                    while bi + 4 <= acts.len() {
8559                        let xs = [
8560                            acts[bi].0.xq.as_slice(),
8561                            acts[bi + 1].0.xq.as_slice(),
8562                            acts[bi + 2].0.xq.as_slice(),
8563                            acts[bi + 3].0.xq.as_slice(),
8564                        ];
8565                        let gs = [
8566                            acts[bi].1.as_slice(),
8567                            acts[bi + 1].1.as_slice(),
8568                            acts[bi + 2].1.as_slice(),
8569                            acts[bi + 3].1.as_slice(),
8570                        ];
8571                        let d = unsafe {
8572                            if vnni_tiles_enabled() {
8573                                dot_q1_row_1x4_vnni(bytes, r, gpr, xs, gs)
8574                            } else {
8575                                dot_q1_row_1x4_avx2(bytes, r, gpr, xs, gs)
8576                            }
8577                        };
8578                        for k in 0..4 {
8579                            let (act, _) = &acts[bi + k];
8580                            let mut acc = d[k] * act.sx;
8581                            for &(j, xv) in &act.outliers {
8582                                let (w, sc) = q1_outlier(bytes, r, gpr, j);
8583                                acc += w * sc * xv;
8584                            }
8585                            // SAFETY: disjoint (bi, r) cells per worker.
8586                            unsafe { *out_addr.at((bi + k) * rows + r) = acc };
8587                        }
8588                        bi += 4;
8589                    }
8590                }
8591                while bi < acts.len() {
8592                    let (act, gsum) = &acts[bi];
8593                    let mut acc = dot_q1_row_i8(bytes, r, gpr, &act.xq, gsum) * act.sx;
8594                    for &(j, xv) in &act.outliers {
8595                        let (w, s) = q1_outlier(bytes, r, gpr, j);
8596                        acc += w * s * xv;
8597                    }
8598                    // SAFETY: disjoint (bi, r) cells per worker range.
8599                    unsafe { *out_addr.at(bi * rows + r) = acc };
8600                    bi += 1;
8601                }
8602            }
8603        };
8604        dispatch_rows(pool, rows, &run);
8605        return;
8606    }
8607    let run = move |start: usize, end: usize| {
8608        for r in start..end {
8609            for bi in 0..b {
8610                let x = &xs_all[bi * cols..(bi + 1) * cols];
8611                // SAFETY: disjoint (bi, r) cells per worker range.
8612                unsafe { *out_addr.at(bi * rows + r) = q1_row_exact(bytes, r, gpr, x) };
8613            }
8614        }
8615    };
8616    dispatch_rows(pool, rows, &run);
8617}
8618
8619/// Fused q4_block matvec straight from the mapped bytes. SDOT path when
8620/// dotprod is available (port of vmfcore `dot_q4_block_sdot`, measured
8621/// +23% on q4 decode): nibbles → centered i8, int8×int8 `sdot` per
8622/// 32-group, exact outlier correction — the same A8W8 contract as q8.
8623/// `CMF_SDOT=0` keeps the exact scalar path.
8624fn q4matvec(
8625    bytes: &[u8],
8626    x: &[f32],
8627    rows: usize,
8628    cols: usize,
8629    out: &mut [f32],
8630    pool: Option<&Pool>,
8631) {
8632    debug_assert_eq!(out.len(), rows);
8633    let (packed, scales) = q4_split(bytes, rows, cols);
8634    let gpr = cols / GROUP_SIZE;
8635    let out_addr = SendMut(out.as_mut_ptr());
8636
8637    if a8w8_enabled() {
8638        let act = split_act(x);
8639        let run = move |start: usize, end: usize| {
8640            q4_range_a8w8(packed, scales, gpr, cols, &act, out_addr, start, end)
8641        };
8642        dispatch_rows(pool, rows, &run);
8643        return;
8644    }
8645
8646    let run =
8647        move |start: usize, end: usize| q4_range_f32(packed, scales, gpr, x, out_addr, start, end);
8648    dispatch_rows(pool, rows, &run);
8649}
8650
8651/// One q4 row via the A8W8 int8 path — SDOT on ARM, AVX2 maddubs on
8652/// x86 (scalar fallback is unreachable: callers gate on a8w8_enabled).
8653#[inline]
8654#[allow(unreachable_code)]
8655/// One UNPACKED q4 row (centered i8 in `buf`) against four activation
8656/// streams: the 32-byte weight chunk and its abs() load once per group,
8657/// the per-group f16 scale decodes once — four maddubs+reduce chains
8658/// instead of four full (load, abs, dot) rounds.
8659#[cfg(target_arch = "x86_64")]
8660#[target_feature(enable = "avx2")]
8661unsafe fn dot_q4b_row_1x4_avx2(
8662    buf: &[u8],
8663    scales: &[u8],
8664    g0: usize,
8665    gpr: usize,
8666    xs: [&[i8]; 4],
8667) -> [f32; 4] {
8668    // SAFETY: callers uphold buffer contracts (buf.len() == gpr·32).
8669    unsafe {
8670        use core::arch::x86_64::*;
8671        let ones = _mm256_set1_epi16(1);
8672        let mut acc = [0f32; 4];
8673        for gi in 0..gpr {
8674            let s = f16_to_f32(u16::from_le_bytes([
8675                scales[(g0 + gi) * 2],
8676                scales[(g0 + gi) * 2 + 1],
8677            ]));
8678            let w = _mm256_loadu_si256(buf.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
8679            let aw = _mm256_abs_epi8(w);
8680            for (k, xq) in xs.iter().enumerate() {
8681                let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
8682                let p16 = _mm256_maddubs_epi16(aw, _mm256_sign_epi8(x, w));
8683                let d = _mm256_madd_epi16(p16, ones);
8684                let hi128 = _mm256_extracti128_si256::<1>(d);
8685                let s128 = _mm_add_epi32(_mm256_castsi256_si128(d), hi128);
8686                let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
8687                let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
8688                acc[k] += _mm_cvtsi128_si32(s32) as f32 * s;
8689            }
8690        }
8691        acc
8692    }
8693}
8694
8695/// VNNI twin of `dot_q4b_row_1x4_avx2` (see `dpbusd_hsum`).
8696#[cfg(target_arch = "x86_64")]
8697#[target_feature(enable = "avx2,avx512f,avx512bw,avx512vl,avx512vnni")]
8698unsafe fn dot_q4b_row_1x4_vnni(
8699    buf: &[u8],
8700    scales: &[u8],
8701    g0: usize,
8702    gpr: usize,
8703    xs: [&[i8]; 4],
8704) -> [f32; 4] {
8705    // SAFETY: callers uphold buffer contracts (buf.len() == gpr·32).
8706    unsafe {
8707        use core::arch::x86_64::*;
8708        let mut acc = [0f32; 4];
8709        for gi in 0..gpr {
8710            let s = f16_to_f32(u16::from_le_bytes([
8711                scales[(g0 + gi) * 2],
8712                scales[(g0 + gi) * 2 + 1],
8713            ]));
8714            let w = _mm256_loadu_si256(buf.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
8715            let aw = _mm256_abs_epi8(w);
8716            for (k, xq) in xs.iter().enumerate() {
8717                let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
8718                let d = dpbusd_hsum(aw, _mm256_sign_epi8(x, w));
8719                acc[k] += d as f32 * s;
8720            }
8721        }
8722        acc
8723    }
8724}
8725
8726/// The vbit flavor of the blocked 1×4: the per-activation A8W8 scale
8727/// folds in PER GROUP as `(d·sx)·s` — bit-matching the single-matvec
8728/// accumulation order (the q4_block flavor applies sx once at the end,
8729/// matching ITS single path; the two conventions are historical and
8730/// each blocked leg must mirror its own).
8731#[cfg(target_arch = "x86_64")]
8732#[target_feature(enable = "avx2")]
8733unsafe fn dot_q4b_row_1x4_sx_avx2(
8734    buf: &[u8],
8735    scales: &[u8],
8736    g0: usize,
8737    gpr: usize,
8738    xs: [&[i8]; 4],
8739    sxs: [f32; 4],
8740) -> [f32; 4] {
8741    // SAFETY: callers uphold buffer contracts (buf.len() == gpr·32).
8742    unsafe {
8743        use core::arch::x86_64::*;
8744        let ones = _mm256_set1_epi16(1);
8745        let mut acc = [0f32; 4];
8746        for gi in 0..gpr {
8747            let s = f16_to_f32(u16::from_le_bytes([
8748                scales[(g0 + gi) * 2],
8749                scales[(g0 + gi) * 2 + 1],
8750            ]));
8751            let w = _mm256_loadu_si256(buf.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
8752            let aw = _mm256_abs_epi8(w);
8753            for (k, xq) in xs.iter().enumerate() {
8754                let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
8755                let p16 = _mm256_maddubs_epi16(aw, _mm256_sign_epi8(x, w));
8756                let d = _mm256_madd_epi16(p16, ones);
8757                let hi128 = _mm256_extracti128_si256::<1>(d);
8758                let s128 = _mm_add_epi32(_mm256_castsi256_si128(d), hi128);
8759                let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
8760                let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
8761                acc[k] += (_mm_cvtsi128_si32(s32) as f32 * sxs[k]) * s;
8762            }
8763        }
8764        acc
8765    }
8766}
8767
8768/// VNNI twin of `dot_q4b_row_1x4_sx_avx2` (see `dpbusd_hsum`; the
8769/// per-group `(d·sx)·s` fold mirrors the vbit single path).
8770#[cfg(target_arch = "x86_64")]
8771#[target_feature(enable = "avx2,avx512f,avx512bw,avx512vl,avx512vnni")]
8772unsafe fn dot_q4b_row_1x4_sx_vnni(
8773    buf: &[u8],
8774    scales: &[u8],
8775    g0: usize,
8776    gpr: usize,
8777    xs: [&[i8]; 4],
8778    sxs: [f32; 4],
8779) -> [f32; 4] {
8780    // SAFETY: callers uphold buffer contracts (buf.len() == gpr·32).
8781    unsafe {
8782        use core::arch::x86_64::*;
8783        let mut acc = [0f32; 4];
8784        for gi in 0..gpr {
8785            let s = f16_to_f32(u16::from_le_bytes([
8786                scales[(g0 + gi) * 2],
8787                scales[(g0 + gi) * 2 + 1],
8788            ]));
8789            let w = _mm256_loadu_si256(buf.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
8790            let aw = _mm256_abs_epi8(w);
8791            for (k, xq) in xs.iter().enumerate() {
8792                let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
8793                let d = dpbusd_hsum(aw, _mm256_sign_epi8(x, w));
8794                acc[k] += (d as f32 * sxs[k]) * s;
8795            }
8796        }
8797        acc
8798    }
8799}
8800
8801#[allow(unreachable_code)]
8802fn dot_q4_row_i8(packed: &[u8], scales: &[u8], g0: usize, gpr: usize, xq: &[i8]) -> f32 {
8803    #[cfg(target_arch = "aarch64")]
8804    unsafe {
8805        return dot_q4_row_sdot(packed, scales, g0, gpr, xq);
8806    }
8807    #[cfg(target_arch = "x86_64")]
8808    unsafe {
8809        return dot_q4_row_avx2(packed, scales, g0, gpr, xq);
8810    }
8811    let mut acc = 0f32;
8812    for gi in 0..gpr {
8813        let g = g0 + gi;
8814        let s = f16_to_f32(u16::from_le_bytes([scales[g * 2], scales[g * 2 + 1]]));
8815        let mut d = 0i32;
8816        for (k, &b) in packed[g * 16..(g + 1) * 16].iter().enumerate() {
8817            d += ((b & 0x0F) as i32 - 8) * xq[gi * GROUP_SIZE + k * 2] as i32
8818                + (((b >> 4) & 0x0F) as i32 - 8) * xq[gi * GROUP_SIZE + k * 2 + 1] as i32;
8819        }
8820        acc += d as f32 * s;
8821    }
8822    acc
8823}
8824
8825/// Two-activation q4 row via the A8W8 int8 path (see `dot_q4_row_i8`).
8826#[inline]
8827#[allow(unreachable_code)]
8828fn dot_q4_row_i8_2(
8829    packed: &[u8],
8830    scales: &[u8],
8831    g0: usize,
8832    gpr: usize,
8833    xq1: &[i8],
8834    xq2: &[i8],
8835) -> (f32, f32) {
8836    #[cfg(target_arch = "aarch64")]
8837    unsafe {
8838        return dot_q4_row_sdot2(packed, scales, g0, gpr, xq1, xq2);
8839    }
8840    #[cfg(target_arch = "x86_64")]
8841    unsafe {
8842        return dot_q4_row_avx2_2(packed, scales, g0, gpr, xq1, xq2);
8843    }
8844    (
8845        dot_q4_row_i8(packed, scales, g0, gpr, xq1),
8846        dot_q4_row_i8(packed, scales, g0, gpr, xq2),
8847    )
8848}
8849
8850/// One q4 row range via SDOT (kernel body of `q4matvec`, extracted so
8851/// multi-matrix jobs can drive it for several tensors in one dispatch).
8852#[allow(clippy::too_many_arguments)]
8853fn q4_range_a8w8(
8854    packed: &[u8],
8855    scales: &[u8],
8856    gpr: usize,
8857    cols: usize,
8858    act: &SplitAct,
8859    out: SendMut,
8860    start: usize,
8861    end: usize,
8862) {
8863    for r in start..end {
8864        let mut acc = dot_q4_row_i8(packed, scales, r * gpr, gpr, &act.xq) * act.sx;
8865        // xq is zeroed at outlier slots — add the exact terms.
8866        for &(j, xv) in &act.outliers {
8867            let flat = r * cols + j;
8868            let byte = packed[flat / 2];
8869            let nib = if flat & 1 == 0 {
8870                byte & 0x0F
8871            } else {
8872                byte >> 4
8873            };
8874            let s = f16_to_f32(u16::from_le_bytes([
8875                scales[(flat / GROUP_SIZE) * 2],
8876                scales[(flat / GROUP_SIZE) * 2 + 1],
8877            ]));
8878            acc += ((nib as i32 - 8) as f32) * s * xv;
8879        }
8880        // SAFETY: disjoint row ranges per worker.
8881        unsafe { *out.at(r) = acc };
8882    }
8883}
8884
8885/// Two-input q4 row range via the A8W8 int8 path — kernel body of
8886/// `q4matvec2`, extracted for pair multi-matrix jobs.
8887#[allow(clippy::too_many_arguments)]
8888fn q4_range2_a8w8(
8889    packed: &[u8],
8890    scales: &[u8],
8891    gpr: usize,
8892    cols: usize,
8893    a1: &SplitAct,
8894    a2: &SplitAct,
8895    p1: SendMut,
8896    p2: SendMut,
8897    start: usize,
8898    end: usize,
8899) {
8900    for r in start..end {
8901        let (s1, s2) = dot_q4_row_i8_2(packed, scales, r * gpr, gpr, &a1.xq, &a2.xq);
8902        let mut acc1 = s1 * a1.sx;
8903        let mut acc2 = s2 * a2.sx;
8904        // xq is zeroed at outlier slots — add the exact terms.
8905        let fix = |outliers: &[(usize, f32)], acc: &mut f32| {
8906            for &(j, xv) in outliers {
8907                let flat = r * cols + j;
8908                let byte = packed[flat / 2];
8909                let nib = if flat & 1 == 0 {
8910                    byte & 0x0F
8911                } else {
8912                    byte >> 4
8913                };
8914                let s = f16_to_f32(u16::from_le_bytes([
8915                    scales[(flat / GROUP_SIZE) * 2],
8916                    scales[(flat / GROUP_SIZE) * 2 + 1],
8917                ]));
8918                *acc += ((nib as i32 - 8) as f32) * s * xv;
8919            }
8920        };
8921        fix(&a1.outliers, &mut acc1);
8922        fix(&a2.outliers, &mut acc2);
8923        // SAFETY: disjoint row ranges per worker.
8924        unsafe {
8925            *p1.at(r) = acc1;
8926            *p2.at(r) = acc2;
8927        }
8928    }
8929}
8930
8931/// Exact scalar q4 row range (same extraction, non-SDOT path).
8932fn q4_range_f32(
8933    packed: &[u8],
8934    scales: &[u8],
8935    gpr: usize,
8936    x: &[f32],
8937    out: SendMut,
8938    start: usize,
8939    end: usize,
8940) {
8941    for r in start..end {
8942        let mut acc = 0f32;
8943        for gi in 0..gpr {
8944            let g = r * gpr + gi;
8945            let s = f16_to_f32(u16::from_le_bytes([scales[g * 2], scales[g * 2 + 1]]));
8946            let pk = &packed[g * 16..(g + 1) * 16];
8947            let xg = &x[gi * GROUP_SIZE..(gi + 1) * GROUP_SIZE];
8948            let mut ga = 0f32;
8949            for (k, &b) in pk.iter().enumerate() {
8950                ga += ((b & 0x0F) as f32 - 8.0) * xg[k * 2]
8951                    + (((b >> 4) & 0x0F) as f32 - 8.0) * xg[k * 2 + 1];
8952            }
8953            acc += ga * s;
8954        }
8955        // SAFETY: disjoint row ranges per worker.
8956        unsafe { *out.at(r) = acc };
8957    }
8958}
8959
8960/// Fused two-input q4 matvec: nibbles are unpacked ONCE per group and
8961/// dotted against both activations (was: two full matvecs — double
8962/// weight traffic). Per-lane math matches `q4matvec` exactly.
8963#[allow(clippy::too_many_arguments)]
8964fn q4matvec2(
8965    bytes: &[u8],
8966    x1: &[f32],
8967    x2: &[f32],
8968    rows: usize,
8969    cols: usize,
8970    o1: &mut [f32],
8971    o2: &mut [f32],
8972    pool: Option<&Pool>,
8973) {
8974    debug_assert_eq!(o1.len(), rows);
8975    debug_assert_eq!(o2.len(), rows);
8976    let (packed, scales) = q4_split(bytes, rows, cols);
8977    let gpr = cols / GROUP_SIZE;
8978
8979    if a8w8_enabled() {
8980        let a1 = split_act(x1);
8981        let a2 = split_act(x2);
8982        let p1 = SendMut(o1.as_mut_ptr());
8983        let p2 = SendMut(o2.as_mut_ptr());
8984        let run = move |start: usize, end: usize| {
8985            q4_range2_a8w8(packed, scales, gpr, cols, &a1, &a2, p1, p2, start, end)
8986        };
8987        dispatch_rows(pool, rows, &run);
8988        return;
8989    }
8990
8991    let p1 = SendMut(o1.as_mut_ptr());
8992    let p2 = SendMut(o2.as_mut_ptr());
8993    let run = move |start: usize, end: usize| {
8994        q4_range2_f32(packed, scales, gpr, x1, x2, p1, p2, start, end)
8995    };
8996    dispatch_rows(pool, rows, &run);
8997}
8998
8999/// Two-input exact scalar q4 row range (same extraction).
9000#[allow(clippy::too_many_arguments)]
9001fn q4_range2_f32(
9002    packed: &[u8],
9003    scales: &[u8],
9004    gpr: usize,
9005    x1: &[f32],
9006    x2: &[f32],
9007    p1: SendMut,
9008    p2: SendMut,
9009    start: usize,
9010    end: usize,
9011) {
9012    for r in start..end {
9013        let (mut acc1, mut acc2) = (0f32, 0f32);
9014        for gi in 0..gpr {
9015            let g = r * gpr + gi;
9016            let s = f16_to_f32(u16::from_le_bytes([scales[g * 2], scales[g * 2 + 1]]));
9017            let pk = &packed[g * 16..(g + 1) * 16];
9018            let x1g = &x1[gi * GROUP_SIZE..(gi + 1) * GROUP_SIZE];
9019            let x2g = &x2[gi * GROUP_SIZE..(gi + 1) * GROUP_SIZE];
9020            let (mut g1, mut g2) = (0f32, 0f32);
9021            for (k, &b) in pk.iter().enumerate() {
9022                let wl = (b & 0x0F) as f32 - 8.0;
9023                let wh = ((b >> 4) & 0x0F) as f32 - 8.0;
9024                g1 += wl * x1g[k * 2] + wh * x1g[k * 2 + 1];
9025                g2 += wl * x2g[k * 2] + wh * x2g[k * 2 + 1];
9026            }
9027            acc1 += g1 * s;
9028            acc2 += g2 * s;
9029        }
9030        // SAFETY: disjoint row ranges per worker.
9031        unsafe {
9032            *p1.at(r) = acc1;
9033            *p2.at(r) = acc2;
9034        }
9035    }
9036}
9037
9038thread_local! {
9039    /// Per-worker decoded-row scratch for the batched q4/vbit kernels
9040    /// (centered i8 for SDOT, f32 for the exact/scalar paths).
9041    static ROW_I8: std::cell::RefCell<Vec<u8>> = const { std::cell::RefCell::new(Vec::new()) };
9042    static ROW_F32: std::cell::RefCell<Vec<f32>> = const { std::cell::RefCell::new(Vec::new()) };
9043}
9044
9045/// Batched q4 matmat: each weight row is unpacked from the mmap ONCE
9046/// and dotted against ALL b activations (prefill used to fall back to b
9047/// full matvecs — b× weight traffic and b× nibble decode). Per-position
9048/// math matches `q4matvec` exactly: same group order, same accumulation.
9049/// `out` is row-major [b, rows] like `qmatmat`.
9050#[allow(clippy::too_many_arguments)]
9051fn q4matmat(
9052    bytes: &[u8],
9053    xs_all: &[f32],
9054    b: usize,
9055    rows: usize,
9056    cols: usize,
9057    out: &mut [f32],
9058    pool: Option<&Pool>,
9059) {
9060    debug_assert_eq!(xs_all.len(), b * cols);
9061    debug_assert_eq!(out.len(), b * rows);
9062    let (packed, scales) = q4_split(bytes, rows, cols);
9063    let gpr = cols / GROUP_SIZE;
9064    let gscale = |g: usize| f16_to_f32(u16::from_le_bytes([scales[g * 2], scales[g * 2 + 1]]));
9065
9066    if a8w8_enabled() {
9067        let acts: Vec<SplitAct> = (0..b)
9068            .map(|bi| split_act(&xs_all[bi * cols..(bi + 1) * cols]))
9069            .collect();
9070        let acts = &acts;
9071        let out_addr = SendMut(out.as_mut_ptr());
9072        let run = move |start: usize, end: usize| {
9073            ROW_I8.with(|rb| {
9074                let mut buf = rb.borrow_mut();
9075                buf.resize(cols, 0);
9076                for r in start..end {
9077                    // Unpack the row's nibbles to centered i8 once
9078                    // (element 2k = low nibble, 2k+1 = high — flat order,
9079                    // same as dot_q4_row_sdot's zip).
9080                    for gi in 0..gpr {
9081                        let g = r * gpr + gi;
9082                        for (k, &bt) in packed[g * 16..(g + 1) * 16].iter().enumerate() {
9083                            buf[gi * GROUP_SIZE + k * 2] = ((bt & 0x0F) as i32 - 8) as i8 as u8;
9084                            buf[gi * GROUP_SIZE + k * 2 + 1] =
9085                                (((bt >> 4) & 0x0F) as i32 - 8) as i8 as u8;
9086                        }
9087                    }
9088                    let mut bi = 0usize;
9089                    #[cfg(target_arch = "x86_64")]
9090                    if avx2_enabled() && blocked_enabled() {
9091                        while bi + 4 <= acts.len() {
9092                            let xs = [
9093                                acts[bi].xq.as_slice(),
9094                                acts[bi + 1].xq.as_slice(),
9095                                acts[bi + 2].xq.as_slice(),
9096                                acts[bi + 3].xq.as_slice(),
9097                            ];
9098                            let d = unsafe {
9099                                if vnni_tiles_enabled() {
9100                                    dot_q4b_row_1x4_vnni(&buf, scales, r * gpr, gpr, xs)
9101                                } else {
9102                                    dot_q4b_row_1x4_avx2(&buf, scales, r * gpr, gpr, xs)
9103                                }
9104                            };
9105                            for k in 0..4 {
9106                                let act = &acts[bi + k];
9107                                let mut acc = d[k] * act.sx;
9108                                for &(j, xv) in &act.outliers {
9109                                    acc += (buf[j] as i8) as f32
9110                                        * gscale((r * cols + j) / GROUP_SIZE)
9111                                        * xv;
9112                                }
9113                                // SAFETY: disjoint (bi, r) cells per worker.
9114                                unsafe { *out_addr.at((bi + k) * rows + r) = acc };
9115                            }
9116                            bi += 4;
9117                        }
9118                    }
9119                    while bi < acts.len() {
9120                        let act = &acts[bi];
9121                        let mut acc = 0f32;
9122                        for gi in 0..gpr {
9123                            let d = dot_i8_i8(
9124                                &buf[gi * GROUP_SIZE..(gi + 1) * GROUP_SIZE],
9125                                &act.xq[gi * GROUP_SIZE..(gi + 1) * GROUP_SIZE],
9126                            );
9127                            acc += d as f32 * gscale(r * gpr + gi);
9128                        }
9129                        acc *= act.sx;
9130                        // xq is zeroed at outlier slots — exact terms.
9131                        for &(j, xv) in &act.outliers {
9132                            acc += (buf[j] as i8) as f32 * gscale((r * cols + j) / GROUP_SIZE) * xv;
9133                        }
9134                        // SAFETY: disjoint (bi, r) cells per worker row range.
9135                        unsafe { *out_addr.at(bi * rows + r) = acc };
9136                        bi += 1;
9137                    }
9138                }
9139            })
9140        };
9141        dispatch_rows(pool, rows, &run);
9142        return;
9143    }
9144
9145    let out_addr = SendMut(out.as_mut_ptr());
9146    let run = move |start: usize, end: usize| {
9147        ROW_F32.with(|rb| {
9148            let mut buf = rb.borrow_mut();
9149            buf.resize(cols, 0.0);
9150            for r in start..end {
9151                // Decode raw (nib − 8) values once; scales stay per-group
9152                // so the accumulation order matches q4matvec bit-for-bit.
9153                for gi in 0..gpr {
9154                    let g = r * gpr + gi;
9155                    for (k, &bt) in packed[g * 16..(g + 1) * 16].iter().enumerate() {
9156                        buf[gi * GROUP_SIZE + k * 2] = (bt & 0x0F) as f32 - 8.0;
9157                        buf[gi * GROUP_SIZE + k * 2 + 1] = ((bt >> 4) & 0x0F) as f32 - 8.0;
9158                    }
9159                }
9160                for bi in 0..b {
9161                    let x = &xs_all[bi * cols..(bi + 1) * cols];
9162                    let mut acc = 0f32;
9163                    for gi in 0..gpr {
9164                        let mut ga = 0f32;
9165                        // Pairwise (lo + hi) addition, matching
9166                        // q4matvec's `ga += lo·x + hi·x` shape exactly —
9167                        // a flat one-per-element loop rounds differently
9168                        // and broke bit-parity on the scalar (x86) path.
9169                        for k in 0..GROUP_SIZE / 2 {
9170                            let e = gi * GROUP_SIZE + k * 2;
9171                            ga += buf[e] * x[e] + buf[e + 1] * x[e + 1];
9172                        }
9173                        acc += ga * gscale(r * gpr + gi);
9174                    }
9175                    // SAFETY: disjoint (bi, r) cells per worker row range.
9176                    unsafe { *out_addr.at(bi * rows + r) = acc };
9177                }
9178            }
9179        })
9180    };
9181    dispatch_rows(pool, rows, &run);
9182}
9183
9184/// Batched vbit matmat: each variable-bit row is decoded from the mmap
9185/// ONCE for the whole microbatch. Same per-position math as
9186/// `vbitmatvec` (SDOT A8W8 with exact outliers / exact f32 for b=8 rows
9187/// and the scalar path).
9188#[allow(clippy::too_many_arguments)]
9189fn vbitmatmat(
9190    bytes: &[u8],
9191    offsets: &[usize],
9192    xs_all: &[f32],
9193    b: usize,
9194    rows: usize,
9195    cols: usize,
9196    out: &mut [f32],
9197    pool: Option<&Pool>,
9198) {
9199    debug_assert_eq!(xs_all.len(), b * cols);
9200    debug_assert_eq!(out.len(), b * rows);
9201    debug_assert_eq!(offsets.len(), rows + 1);
9202    let ng = cols / GROUP_SIZE;
9203    let bits = &bytes[..rows];
9204    let sc_off = rows;
9205    let gscale = |r: usize, g: usize| {
9206        let so = (r * ng + g) * 2;
9207        f16_to_f32(u16::from_le_bytes([
9208            bytes[sc_off + so],
9209            bytes[sc_off + so + 1],
9210        ]))
9211    };
9212
9213    // Decode row r's raw (u − L) values into `dst` (f32, unscaled).
9214    let decode_f32 = |r: usize, dst: &mut [f32]| {
9215        let bw = bits[r] as usize;
9216        let l = ((1i32 << (bw - 1)) - 1) as f32;
9217        let data = &bytes[offsets[r]..offsets[r + 1]];
9218        let (mut acc, mut nbits, mut idx) = (0u64, 0usize, 0usize);
9219        for d in dst.iter_mut() {
9220            while nbits < bw {
9221                acc = (acc << 8) | data[idx] as u64;
9222                idx += 1;
9223                nbits += 8;
9224            }
9225            let u = ((acc >> (nbits - bw)) & ((1u64 << bw) - 1)) as f32;
9226            nbits -= bw;
9227            *d = u - l;
9228        }
9229    };
9230
9231    if a8w8_enabled() {
9232        let acts: Vec<SplitAct> = (0..b)
9233            .map(|bi| split_act(&xs_all[bi * cols..(bi + 1) * cols]))
9234            .collect();
9235        let acts = &acts;
9236        let out_addr = SendMut(out.as_mut_ptr());
9237        let run = move |start: usize, end: usize| {
9238            for r in start..end {
9239                let bw = bits[r] as usize;
9240                if bw == 8 {
9241                    // u−L reaches 128 → no i8 path; decode once, exact
9242                    // f32 dots for every position (same as vbitmatvec).
9243                    ROW_F32.with(|rb| {
9244                        let mut buf = rb.borrow_mut();
9245                        buf.resize(cols, 0.0);
9246                        decode_f32(r, &mut buf);
9247                        for bi in 0..b {
9248                            let x = &xs_all[bi * cols..(bi + 1) * cols];
9249                            let mut dot = 0f32;
9250                            for g in 0..ng {
9251                                let mut gd = 0f32;
9252                                for k in 0..GROUP_SIZE {
9253                                    gd += buf[g * GROUP_SIZE + k] * x[g * GROUP_SIZE + k];
9254                                }
9255                                dot += gd * gscale(r, g);
9256                            }
9257                            // SAFETY: disjoint (bi, r) cells per worker range.
9258                            unsafe { *out_addr.at(bi * rows + r) = dot };
9259                        }
9260                    });
9261                    continue;
9262                }
9263                let l = (1i32 << (bw - 1)) - 1;
9264                let data = &bytes[offsets[r]..offsets[r + 1]];
9265                ROW_I8.with(|rb| {
9266                    let mut buf = rb.borrow_mut();
9267                    buf.resize(cols, 0);
9268                    #[inline(always)]
9269                    fn fill<const B: usize>(data: &[u8], l: i32, buf: &mut [u8]) {
9270                        for (blk, chunk) in buf.chunks_exact_mut(8).enumerate() {
9271                            let u = unpack8::<B>(&data[blk * B..]);
9272                            for k in 0..8 {
9273                                chunk[k] = (u[k] - l) as i8 as u8;
9274                            }
9275                        }
9276                    }
9277                    match bw {
9278                        3 => fill::<3>(data, l, &mut buf),
9279                        4 => vbit_fill4(data, &mut buf),
9280                        5 => fill::<5>(data, l, &mut buf),
9281                        6 => fill::<6>(data, l, &mut buf),
9282                        _ => unreachable!("vbit bit-width {bw} (validated at load)"),
9283                    }
9284                    let mut bi = 0usize;
9285                    // The vbit scale table shares q4_block's layout
9286                    // (contiguous f16 per (row·ng + g)), so the same
9287                    // blocked 1×4 kernel serves the decoded row.
9288                    #[cfg(target_arch = "x86_64")]
9289                    if avx2_enabled() && blocked_enabled() {
9290                        while bi + 4 <= acts.len() {
9291                            let xs = [
9292                                acts[bi].xq.as_slice(),
9293                                acts[bi + 1].xq.as_slice(),
9294                                acts[bi + 2].xq.as_slice(),
9295                                acts[bi + 3].xq.as_slice(),
9296                            ];
9297                            let sxs = [
9298                                acts[bi].sx,
9299                                acts[bi + 1].sx,
9300                                acts[bi + 2].sx,
9301                                acts[bi + 3].sx,
9302                            ];
9303                            let d = unsafe {
9304                                if vnni_tiles_enabled() {
9305                                    dot_q4b_row_1x4_sx_vnni(
9306                                        &buf,
9307                                        &bytes[sc_off..],
9308                                        r * ng,
9309                                        ng,
9310                                        xs,
9311                                        sxs,
9312                                    )
9313                                } else {
9314                                    dot_q4b_row_1x4_sx_avx2(
9315                                        &buf,
9316                                        &bytes[sc_off..],
9317                                        r * ng,
9318                                        ng,
9319                                        xs,
9320                                        sxs,
9321                                    )
9322                                }
9323                            };
9324                            for k in 0..4 {
9325                                let act = &acts[bi + k];
9326                                let mut dot = d[k];
9327                                for &(j, xv) in &act.outliers {
9328                                    dot += (buf[j] as i8) as f32 * gscale(r, j / GROUP_SIZE) * xv;
9329                                }
9330                                // SAFETY: disjoint (bi, r) cells per worker.
9331                                unsafe { *out_addr.at((bi + k) * rows + r) = dot };
9332                            }
9333                            bi += 4;
9334                        }
9335                    }
9336                    while bi < acts.len() {
9337                        let act = &acts[bi];
9338                        let mut dot = 0f32;
9339                        for g in 0..ng {
9340                            let d = dot_i8_i8(
9341                                &buf[g * GROUP_SIZE..(g + 1) * GROUP_SIZE],
9342                                &act.xq[g * GROUP_SIZE..(g + 1) * GROUP_SIZE],
9343                            ) as f32
9344                                * act.sx;
9345                            dot += d * gscale(r, g);
9346                        }
9347                        for &(j, xv) in &act.outliers {
9348                            dot += (buf[j] as i8) as f32 * gscale(r, j / GROUP_SIZE) * xv;
9349                        }
9350                        // SAFETY: disjoint (bi, r) cells per worker range.
9351                        unsafe { *out_addr.at(bi * rows + r) = dot };
9352                        bi += 1;
9353                    }
9354                });
9355            }
9356        };
9357        dispatch_rows(pool, rows, &run);
9358        return;
9359    }
9360
9361    let out_addr = SendMut(out.as_mut_ptr());
9362    let run = move |start: usize, end: usize| {
9363        ROW_F32.with(|rb| {
9364            let mut buf = rb.borrow_mut();
9365            buf.resize(cols, 0.0);
9366            for r in start..end {
9367                decode_f32(r, &mut buf);
9368                for bi in 0..b {
9369                    let x = &xs_all[bi * cols..(bi + 1) * cols];
9370                    let mut dot = 0f32;
9371                    for g in 0..ng {
9372                        let mut gd = 0f32;
9373                        for k in 0..GROUP_SIZE {
9374                            gd += buf[g * GROUP_SIZE + k] * x[g * GROUP_SIZE + k];
9375                        }
9376                        dot += gd * gscale(r, g);
9377                    }
9378                    // SAFETY: disjoint (bi, r) cells per worker range.
9379                    unsafe { *out_addr.at(bi * rows + r) = dot };
9380                }
9381            }
9382        })
9383    };
9384    dispatch_rows(pool, rows, &run);
9385}
9386
9387/// Build a GPU batch job for a q8-family mapped tensor (primary
9388/// shard): prescaled input + directory coordinates. None → not
9389/// GPU-eligible, caller stays on the CPU.
9390pub(crate) fn gpu_batch_job<'a>(
9391    t: &'a QTensor,
9392    x: &[f32],
9393) -> Option<(std::sync::Arc<CmfModel>, crate::gpu::BatchJob<'a>)> {
9394    match t {
9395        QTensor::Mapped {
9396            model,
9397            idx,
9398            dtype: dt @ (TensorDtype::Q8Row | TensorDtype::Q8_2f),
9399            rows,
9400            cols,
9401            row_scale,
9402            col_field,
9403            ..
9404        } => Some((
9405            model.clone(),
9406            crate::gpu::BatchJob {
9407                idx: *idx,
9408                rows: *rows,
9409                cols: *cols,
9410                row_scale,
9411                xs: prescale(x, col_field, *dt).into_owned(),
9412                layout: crate::gpu::BatchLayout::Q8,
9413            },
9414        )),
9415        // q1: raw f32 activations, tile-embedded scales.
9416        QTensor::Mapped {
9417            model,
9418            idx,
9419            dtype: TensorDtype::Q1,
9420            rows,
9421            cols,
9422            ..
9423        } => Some((
9424            model.clone(),
9425            crate::gpu::BatchJob {
9426                idx: *idx,
9427                rows: *rows,
9428                cols: *cols,
9429                row_scale: &[],
9430                xs: x.to_vec(),
9431                layout: crate::gpu::BatchLayout::Q1,
9432            },
9433        )),
9434        // q4_tiled / q4tp: raw f32 activations; the scales live in the
9435        // payload (inline tiles / row ladder), so row_scale stays empty.
9436        // The GDN projection batch already runs these layouts on Metal —
9437        // this arm lets the attention QKV batch reach the same kernels.
9438        QTensor::Mapped {
9439            model,
9440            idx,
9441            dtype: dt @ (TensorDtype::Q4Tiled | TensorDtype::Q4TiledP),
9442            rows,
9443            cols,
9444            ..
9445        } => Some((
9446            model.clone(),
9447            crate::gpu::BatchJob {
9448                idx: *idx,
9449                rows: *rows,
9450                cols: *cols,
9451                row_scale: &[],
9452                xs: x.to_vec(),
9453                layout: if *dt == TensorDtype::Q4Tiled {
9454                    crate::gpu::BatchLayout::Q4t
9455                } else {
9456                    crate::gpu::BatchLayout::Q4tp
9457                },
9458            },
9459        )),
9460        _ => None,
9461    }
9462}
9463
9464thread_local! {
9465    static PRESCALE_BUF1: std::cell::RefCell<Vec<f32>> = const { std::cell::RefCell::new(Vec::new()) };
9466    static PRESCALE_BUF2: std::cell::RefCell<Vec<f32>> = const { std::cell::RefCell::new(Vec::new()) };
9467}
9468
9469pub(crate) fn prescale<'a>(
9470    x: &'a [f32],
9471    col_field: &[f32],
9472    dtype: TensorDtype,
9473) -> std::borrow::Cow<'a, [f32]> {
9474    if dtype == TensorDtype::Q8_2f {
9475        x.iter().zip(col_field).map(|(a, c)| a * c).collect()
9476    } else {
9477        std::borrow::Cow::Borrowed(x)
9478    }
9479}
9480
9481/// θ col-field fold for q8_2f activations. Borrowed pass-through for
9482/// every other dtype, using thread-local buffers to eliminate per-matvec allocations.
9483pub(crate) fn prescale_with<R, F: FnOnce(&[f32]) -> R>(
9484    x: &[f32],
9485    col_field: &[f32],
9486    dtype: TensorDtype,
9487    buf_id: u8,
9488    f: F,
9489) -> R {
9490    if dtype == TensorDtype::Q8_2f {
9491        if buf_id == 1 {
9492            PRESCALE_BUF1.with(|b| {
9493                let mut buf = b.borrow_mut();
9494                buf.clear();
9495                buf.extend(x.iter().zip(col_field).map(|(a, c)| a * c));
9496                f(&buf)
9497            })
9498        } else {
9499            PRESCALE_BUF2.with(|b| {
9500                let mut buf = b.borrow_mut();
9501                buf.clear();
9502                buf.extend(x.iter().zip(col_field).map(|(a, c)| a * c));
9503                f(&buf)
9504            })
9505        }
9506    } else {
9507        f(x)
9508    }
9509}
9510
9511// ───────────────────── x86-64 AVX2 kernels (roadmap этап 2) ─────────────────────
9512
9513/// AVX2+FMA available? Default ON when the CPU supports both;
9514/// `CMF_AVX2=0` disables (falls back to the autovectorized loops).
9515#[cfg(target_arch = "x86_64")]
9516pub(crate) fn avx2_enabled() -> bool {
9517    use std::sync::OnceLock;
9518    static ON: OnceLock<bool> = OnceLock::new();
9519    *ON.get_or_init(|| {
9520        std::env::var("CMF_AVX2").map(|v| v != "0").unwrap_or(true)
9521            && std::arch::is_x86_feature_detected!("avx2")
9522            && std::arch::is_x86_feature_detected!("fma")
9523    })
9524}
9525
9526/// AVX2 A8W8 allowed? The quantized-activation contract is switched by
9527/// the SAME env as the ARM SDOT path: `CMF_SDOT=0` keeps exact kernels
9528/// (the golden-parity exact gate relies on it) — AVX2 f32 kernels stay
9529/// active either way, they are exact (regrouped sums only).
9530#[cfg(target_arch = "x86_64")]
9531fn avx2_a8w8_enabled() -> bool {
9532    if FLOAT_ACTIVATIONS.get() {
9533        return false;
9534    }
9535    use std::sync::OnceLock;
9536    static ON: OnceLock<bool> = OnceLock::new();
9537    *ON.get_or_init(|| {
9538        avx2_enabled() && std::env::var("CMF_SDOT").map(|v| v != "0").unwrap_or(true)
9539    })
9540}
9541
9542thread_local! {
9543    static FULL_GPU_Q8: std::cell::Cell<bool> = const { std::cell::Cell::new(false) };
9544}
9545
9546/// Match the graph's full-device q8 projection precision on MiMo's host
9547/// tail. Does not enable the GPU or bypass a CPU-only/device-refusal gate.
9548pub(crate) fn enter_full_gpu_q8_scope() -> impl Drop {
9549    struct Restore(bool, std::marker::PhantomData<std::rc::Rc<()>>);
9550    impl Drop for Restore {
9551        fn drop(&mut self) {
9552            FULL_GPU_Q8.set(self.0);
9553        }
9554    }
9555    Restore(FULL_GPU_Q8.replace(true), std::marker::PhantomData)
9556}
9557
9558// Dynamic MiMo experts must not change activation precision when a cache
9559// fill moves them from CPU to GPU. Thread-local: only the cold-expert
9560// dispatch selects float kernels; concurrent pipelines keep their policy.
9561thread_local! {
9562    static FLOAT_ACTIVATIONS: std::cell::Cell<bool> = const { std::cell::Cell::new(false) };
9563}
9564
9565pub(crate) fn float_activations_scope<R>(f: impl FnOnce() -> R) -> R {
9566    struct Restore(bool);
9567    impl Drop for Restore {
9568        fn drop(&mut self) {
9569            FLOAT_ACTIVATIONS.set(self.0);
9570        }
9571    }
9572    let _restore = Restore(FLOAT_ACTIVATIONS.replace(true));
9573    f()
9574}
9575
9576/// Row-exact batching: while set, the x86 batched kernels (`qmatmat`,
9577/// `q4tp_matmat`) compute every (weight row, token) cell with the
9578/// single-token kernel instead of the blocked 2×4 / 1×4 / 1×8 tiles, so a
9579/// token's result does not depend on the batch it rides in and equals its
9580/// matvec. On ARM `q4tp_matmat` keeps its 1×4 tile but in the matvec's
9581/// reduction order, and no q4tp batch takes the AMX or device GEMM. The
9582/// MiMo speculative verify holds it (`row_exact_scope`) — its accepted
9583/// rows must be the rows plain decode would have produced.
9584// Shared with pool workers, so overlapping requests must keep the mode
9585// enabled until the LAST scope leaves. Saving/restoring a global bool is
9586// incorrect when two threads enter and leave in a non-LIFO order.
9587static ROW_EXACT: std::sync::atomic::AtomicUsize = std::sync::atomic::AtomicUsize::new(0);
9588
9589pub(crate) fn row_exact() -> bool {
9590    ROW_EXACT.load(std::sync::atomic::Ordering::Acquire) != 0
9591}
9592
9593fn counted_row_exact_scope<R>(active: &std::sync::atomic::AtomicUsize, f: impl FnOnce() -> R) -> R {
9594    struct Restore<'a>(&'a std::sync::atomic::AtomicUsize);
9595    impl Drop for Restore<'_> {
9596        fn drop(&mut self) {
9597            self.0.fetch_sub(1, std::sync::atomic::Ordering::AcqRel);
9598        }
9599    }
9600    active.fetch_add(1, std::sync::atomic::Ordering::AcqRel);
9601    let _restore = Restore(active);
9602    f()
9603}
9604
9605/// Run `f` with row-exact batching on (also released on unwind).
9606pub(crate) fn row_exact_scope<R>(f: impl FnOnce() -> R) -> R {
9607    counted_row_exact_scope(&ROW_EXACT, f)
9608}
9609
9610/// A8W8 quantized-activation path available on THIS machine? One
9611/// switch across architectures: ARM dotprod (CMF_SDOT) or x86 AVX2
9612/// (CMF_AVX2 + the same CMF_SDOT exact-contract override).
9613#[inline]
9614pub(crate) fn a8w8_enabled() -> bool {
9615    #[cfg(target_arch = "aarch64")]
9616    {
9617        sdot_enabled()
9618    }
9619    #[cfg(target_arch = "x86_64")]
9620    {
9621        avx2_a8w8_enabled()
9622    }
9623    #[cfg(not(any(target_arch = "aarch64", target_arch = "x86_64")))]
9624    {
9625        false
9626    }
9627}
9628
9629/// int8·int8 dot dispatch: SDOT on ARM; AVX-512 VNNI (vpdpbusd) or AVX2
9630/// maddubs on x86. Callers are gated by `a8w8_enabled()`.
9631#[inline]
9632#[allow(unreachable_code)]
9633fn dot_i8_i8(w: &[u8], xq: &[i8]) -> i32 {
9634    #[cfg(target_arch = "aarch64")]
9635    unsafe {
9636        return dot_i8_sdot(w, xq);
9637    }
9638    #[cfg(target_arch = "x86_64")]
9639    unsafe {
9640        if avx512vnni_enabled() {
9641            return dot_i8_i8_vnni(w, xq);
9642        }
9643        return dot_i8_i8_avx2(w, xq);
9644    }
9645    w.iter()
9646        .zip(xq)
9647        .map(|(&a, &b)| (a as i8) as i32 * b as i32)
9648        .sum()
9649}
9650
9651/// AVX-512 VNNI available? (F+BW+VL+VNNI; `CMF_AVX512=0` falls back to
9652/// AVX2.) VL matters: short 32-byte groups (q4/vbit) ride the 256-bit
9653/// `vpdpbusd` encoding.
9654#[cfg(target_arch = "x86_64")]
9655fn avx512vnni_enabled() -> bool {
9656    use std::sync::OnceLock;
9657    static ON: OnceLock<bool> = OnceLock::new();
9658    *ON.get_or_init(|| {
9659        std::env::var("CMF_AVX512")
9660            .map(|v| v != "0")
9661            .unwrap_or(true)
9662            && std::arch::is_x86_feature_detected!("avx512f")
9663            && std::arch::is_x86_feature_detected!("avx512bw")
9664            && std::arch::is_x86_feature_detected!("avx512vl")
9665            && std::arch::is_x86_feature_detected!("avx512vnni")
9666    })
9667}
9668
9669/// Grouped-codec VNNI arms (the q4t/q4b/q1/q1t tile kernels): default
9670/// ON where AVX-512 VNNI exists (`CMF_VNNI_TILES=0` opt-out). Measured
9671/// on Ryzen 7950X (Zen4, 3 alternating process pairs, blocked GEMM
9672/// 4864×896 b=256): q4t 63→68 GF/s (+8%), q1 53→56 (+6%), q4b 72→75
9673/// (+4%) — consistent, no leg regressed. The tile kernels keep a
9674/// horizontal reduce per 32-weight group, so the `vpdpbusd` saving is
9675/// smaller than the long-dot q8 win (+13%), but it is real and free.
9676#[cfg(target_arch = "x86_64")]
9677fn vnni_tiles_enabled() -> bool {
9678    use std::sync::OnceLock;
9679    static ON: OnceLock<bool> = OnceLock::new();
9680    *ON.get_or_init(|| {
9681        std::env::var("CMF_VNNI_TILES")
9682            .map(|v| v != "0")
9683            .unwrap_or(true)
9684            && avx512vnni_enabled()
9685    })
9686}
9687
9688/// One 256-bit u8×i8 dot → i32 via `vpdpbusd` into a fresh accumulator
9689/// plus the same horizontal reduce the AVX2 kernels use. Products are
9690/// bounded (|w| ≤ 8 or ≤ 1), so maddubs never saturated — the i32 sum
9691/// is bit-identical to the maddubs+madd pair it replaces.
9692#[cfg(target_arch = "x86_64")]
9693#[target_feature(enable = "avx2,avx512f,avx512bw,avx512vl,avx512vnni")]
9694#[inline]
9695unsafe fn dpbusd_hsum(aw: core::arch::x86_64::__m256i, xs: core::arch::x86_64::__m256i) -> i32 {
9696    // SAFETY: pure register math.
9697    unsafe {
9698        use core::arch::x86_64::*;
9699        let d = _mm256_dpbusd_epi32(_mm256_setzero_si256(), aw, xs);
9700        let hi128 = _mm256_extracti128_si256::<1>(d);
9701        let s128 = _mm_add_epi32(_mm256_castsi256_si128(d), hi128);
9702        let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
9703        let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
9704        _mm_cvtsi128_si32(s32)
9705    }
9706}
9707
9708/// int8·int8 via AVX-512 VNNI: `vpdpbusd` fuses the maddubs+madd+add
9709/// triple into one u8×i8 dot-accumulate. AVX-512 has no vpsignb, so the
9710/// |w|·sign(x,w) trick becomes |w| × (x negated where w<0) via a mask
9711/// subtract — w==0 lanes contribute 0 through |w|=0 either way.
9712#[cfg(target_arch = "x86_64")]
9713#[target_feature(enable = "avx2,avx512f,avx512bw,avx512vl,avx512vnni")]
9714unsafe fn dot_i8_i8_vnni(w: &[u8], xq: &[i8]) -> i32 {
9715    // SAFETY: callers uphold slice-length contracts (see call sites).
9716    unsafe {
9717        use core::arch::x86_64::*;
9718        let n = w.len();
9719        let mut j = 0usize;
9720        let mut total: i32;
9721        // 4 independent accumulators: vpdpbusd is its own loop-carried
9722        // dependency (~5-cycle latency) — a single-acc loop runs
9723        // latency-bound and LOSES to the AVX2 maddubs kernel, measured
9724        // on Granite Rapids.
9725        {
9726            #[inline(always)]
9727            unsafe fn step(
9728                w: *const u8,
9729                x: *const i8,
9730                acc: core::arch::x86_64::__m512i,
9731            ) -> core::arch::x86_64::__m512i {
9732                unsafe {
9733                    use core::arch::x86_64::*;
9734                    let wv = _mm512_loadu_si512(w as *const _);
9735                    let xv = _mm512_loadu_si512(x as *const _);
9736                    let aw = _mm512_abs_epi8(wv);
9737                    let neg = _mm512_movepi8_mask(wv);
9738                    let sx = _mm512_mask_sub_epi8(xv, neg, _mm512_setzero_si512(), xv);
9739                    _mm512_dpbusd_epi32(acc, aw, sx)
9740                }
9741            }
9742            let (mut a0, mut a1, mut a2, mut a3) = (
9743                _mm512_setzero_si512(),
9744                _mm512_setzero_si512(),
9745                _mm512_setzero_si512(),
9746                _mm512_setzero_si512(),
9747            );
9748            while j + 256 <= n {
9749                a0 = step(w.as_ptr().add(j), xq.as_ptr().add(j), a0);
9750                a1 = step(w.as_ptr().add(j + 64), xq.as_ptr().add(j + 64), a1);
9751                a2 = step(w.as_ptr().add(j + 128), xq.as_ptr().add(j + 128), a2);
9752                a3 = step(w.as_ptr().add(j + 192), xq.as_ptr().add(j + 192), a3);
9753                j += 256;
9754            }
9755            while j + 64 <= n {
9756                a0 = step(w.as_ptr().add(j), xq.as_ptr().add(j), a0);
9757                j += 64;
9758            }
9759            let s01 = _mm512_add_epi32(a0, a1);
9760            let s23 = _mm512_add_epi32(a2, a3);
9761            total = _mm512_reduce_add_epi32(_mm512_add_epi32(s01, s23));
9762        }
9763        // 32-wide (q4/vbit groups are exactly 32 bytes).
9764        if j + 32 <= n {
9765            let wv = _mm256_loadu_si256(w.as_ptr().add(j) as *const __m256i);
9766            let xv = _mm256_loadu_si256(xq.as_ptr().add(j) as *const __m256i);
9767            let d = _mm256_dpbusd_epi32(
9768                _mm256_setzero_si256(),
9769                _mm256_abs_epi8(wv),
9770                _mm256_sign_epi8(xv, wv),
9771            );
9772            let hi128 = _mm256_extracti128_si256::<1>(d);
9773            let s128 = _mm_add_epi32(_mm256_castsi256_si128(d), hi128);
9774            let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
9775            let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
9776            total += _mm_cvtsi128_si32(s32);
9777            j += 32;
9778        }
9779        while j < n {
9780            total += (w[j] as i8) as i32 * xq[j] as i32;
9781            j += 1;
9782        }
9783        total
9784    }
9785}
9786
9787/// i8 row · f32 x via AVX2/FMA (x86 mirror of `dot_i8_f32_neon`).
9788#[cfg(target_arch = "x86_64")]
9789#[target_feature(enable = "avx2,fma")]
9790unsafe fn dot_i8_f32_avx2(w: &[u8], x: &[f32]) -> f32 {
9791    // SAFETY: callers uphold slice-length contracts (see call sites).
9792    unsafe {
9793        use core::arch::x86_64::*;
9794        let n = x.len();
9795        let wp = w.as_ptr();
9796        let xp = x.as_ptr();
9797        let (mut a0, mut a1) = (_mm256_setzero_ps(), _mm256_setzero_ps());
9798        let mut j = 0usize;
9799        while j + 16 <= n {
9800            let wb = _mm_loadu_si128(wp.add(j) as *const __m128i);
9801            let lo = _mm256_cvtepi8_epi32(wb);
9802            let hi = _mm256_cvtepi8_epi32(_mm_srli_si128::<8>(wb));
9803            a0 = _mm256_fmadd_ps(_mm256_cvtepi32_ps(lo), _mm256_loadu_ps(xp.add(j)), a0);
9804            a1 = _mm256_fmadd_ps(_mm256_cvtepi32_ps(hi), _mm256_loadu_ps(xp.add(j + 8)), a1);
9805            j += 16;
9806        }
9807        let acc = _mm256_add_ps(a0, a1);
9808        let hi128 = _mm256_extractf128_ps::<1>(acc);
9809        let s128 = _mm_add_ps(_mm256_castps256_ps128(acc), hi128);
9810        let s64 = _mm_add_ps(s128, _mm_movehl_ps(s128, s128));
9811        let s32 = _mm_add_ss(s64, _mm_shuffle_ps::<1>(s64, s64));
9812        let mut sum = _mm_cvtss_f32(s32);
9813        while j < n {
9814            sum += (*wp.add(j) as i8) as f32 * *xp.add(j);
9815            j += 1;
9816        }
9817        sum
9818    }
9819}
9820
9821/// int8(weight)·int8(activation) → i32 via AVX2 maddubs — the x86
9822/// analogue of the SDOT A8W8 path. `maddubs` takes u8×i8, so the
9823/// standard sign trick applies: |w| × sign(x, w) ≡ w × x per lane.
9824/// Pair saturation is safe: |w|≤128, |x|≤127 → 2·128·127 < 32767.
9825#[cfg(target_arch = "x86_64")]
9826#[target_feature(enable = "avx2")]
9827unsafe fn dot_i8_i8_avx2(w: &[u8], xq: &[i8]) -> i32 {
9828    // SAFETY: callers uphold slice-length contracts (see call sites).
9829    unsafe {
9830        use core::arch::x86_64::*;
9831        let n = w.len();
9832        let ones = _mm256_set1_epi16(1);
9833        let mut acc = _mm256_setzero_si256();
9834        let mut j = 0usize;
9835        while j + 32 <= n {
9836            let wv = _mm256_loadu_si256(w.as_ptr().add(j) as *const __m256i);
9837            let xv = _mm256_loadu_si256(xq.as_ptr().add(j) as *const __m256i);
9838            let p16 = _mm256_maddubs_epi16(_mm256_abs_epi8(wv), _mm256_sign_epi8(xv, wv));
9839            acc = _mm256_add_epi32(acc, _mm256_madd_epi16(p16, ones));
9840            j += 32;
9841        }
9842        let hi128 = _mm256_extracti128_si256::<1>(acc);
9843        let s128 = _mm_add_epi32(_mm256_castsi256_si128(acc), hi128);
9844        let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
9845        let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
9846        let mut s = _mm_cvtsi128_si32(s32);
9847        while j < n {
9848            s += (w[j] as i8) as i32 * xq[j] as i32;
9849            j += 1;
9850        }
9851        s
9852    }
9853}
9854
9855/// smmla 2×4: one instruction covers a 2-row × 2-activation × 8-deep
9856/// tile (32 MACs vs sdot's 16) — the weight pair loads once per 8-k
9857/// slice as a combined 2×8 register and meets two activation pairs.
9858#[cfg(target_arch = "aarch64")]
9859#[target_feature(enable = "neon,i8mm")]
9860unsafe fn dot_i8_smmla_2x4(w0: &[u8], w1: &[u8], xs: [&[i8]; 4]) -> [[i32; 4]; 2] {
9861    // SAFETY: callers uphold slice-length contracts.
9862    unsafe {
9863        use core::arch::aarch64::*;
9864        use core::arch::asm;
9865        let n = w0.len();
9866        let w0p = w0.as_ptr() as *const i8;
9867        let w1p = w1.as_ptr() as *const i8;
9868        // acc01 holds [c(r0,x0) c(r0,x1) c(r1,x0) c(r1,x1)]; acc23 the
9869        // same for x2/x3.
9870        let mut acc01 = vdupq_n_s32(0);
9871        let mut acc23 = vdupq_n_s32(0);
9872        let mut i = 0usize;
9873        while i + 8 <= n {
9874            let wa = vcombine_s8(vld1_s8(w0p.add(i)), vld1_s8(w1p.add(i)));
9875            let xb01 = vcombine_s8(
9876                vld1_s8(xs[0].as_ptr().add(i)),
9877                vld1_s8(xs[1].as_ptr().add(i)),
9878            );
9879            let xb23 = vcombine_s8(
9880                vld1_s8(xs[2].as_ptr().add(i)),
9881                vld1_s8(xs[3].as_ptr().add(i)),
9882            );
9883            asm!(
9884                "smmla {a01:v}.4s, {w:v}.16b, {x01:v}.16b",
9885                "smmla {a23:v}.4s, {w:v}.16b, {x23:v}.16b",
9886                a01 = inout(vreg) acc01, a23 = inout(vreg) acc23,
9887                w = in(vreg) wa, x01 = in(vreg) xb01, x23 = in(vreg) xb23,
9888                options(pure, nomem, nostack),
9889            );
9890            i += 8;
9891        }
9892        let mut out = [[0i32; 4]; 2];
9893        let a01: [i32; 4] = core::mem::transmute(acc01);
9894        let a23: [i32; 4] = core::mem::transmute(acc23);
9895        out[0][0] = a01[0];
9896        out[0][1] = a01[1];
9897        out[1][0] = a01[2];
9898        out[1][1] = a01[3];
9899        out[0][2] = a23[0];
9900        out[0][3] = a23[1];
9901        out[1][2] = a23[2];
9902        out[1][3] = a23[3];
9903        if i < n {
9904            for (k, x) in xs.iter().enumerate() {
9905                for j in i..n {
9906                    out[0][k] += (w0[j] as i8) as i32 * x[j] as i32;
9907                    out[1][k] += (w1[j] as i8) as i32 * x[j] as i32;
9908                }
9909            }
9910        }
9911        out
9912    }
9913}
9914
9915/// ARM twin of the x86 blocked prefill GEMM: two weight rows stay in
9916/// registers across four activation streams, eight sdot accumulators.
9917/// (The per-row form re-read each W row once per activation.)
9918#[cfg(target_arch = "aarch64")]
9919#[target_feature(enable = "neon,dotprod")]
9920unsafe fn dot_i8_sdot_2x4(w0: &[u8], w1: &[u8], xs: [&[i8]; 4]) -> [[i32; 4]; 2] {
9921    // SAFETY: callers uphold slice-length contracts.
9922    unsafe {
9923        use core::arch::aarch64::*;
9924        use core::arch::asm;
9925        let n = w0.len();
9926        let w0p = w0.as_ptr() as *const i8;
9927        let w1p = w1.as_ptr() as *const i8;
9928        let mut acc = [[vdupq_n_s32(0); 4]; 2];
9929        let mut i = 0usize;
9930        while i + 16 <= n {
9931            let wv0 = vld1q_s8(w0p.add(i));
9932            let wv1 = vld1q_s8(w1p.add(i));
9933            for (k, x) in xs.iter().enumerate() {
9934                let xv = vld1q_s8(x.as_ptr().add(i));
9935                let (mut a0, mut a1) = (acc[0][k], acc[1][k]);
9936                asm!(
9937                    "sdot {a0:v}.4s, {w0:v}.16b, {x:v}.16b",
9938                    "sdot {a1:v}.4s, {w1:v}.16b, {x:v}.16b",
9939                    a0 = inout(vreg) a0, a1 = inout(vreg) a1,
9940                    w0 = in(vreg) wv0, w1 = in(vreg) wv1, x = in(vreg) xv,
9941                    options(pure, nomem, nostack),
9942                );
9943                acc[0][k] = a0;
9944                acc[1][k] = a1;
9945            }
9946            i += 16;
9947        }
9948        let mut out = [[0i32; 4]; 2];
9949        for r in 0..2 {
9950            for k in 0..4 {
9951                out[r][k] = vaddvq_s32(acc[r][k]);
9952            }
9953        }
9954        if i < n {
9955            for (k, x) in xs.iter().enumerate() {
9956                for j in i..n {
9957                    out[0][k] += (w0[j] as i8) as i32 * x[j] as i32;
9958                    out[1][k] += (w1[j] as i8) as i32 * x[j] as i32;
9959                }
9960            }
9961        }
9962        out
9963    }
9964}
9965
9966/// Blocked 2 weight rows × 4 activations for the prefill GEMM
9967/// (roadmap P0: packed panels + multi-row accumulators). The two rows'
9968/// abs() live in registers across all four activation streams; the
9969/// sign-fixup is recomputed per pair (the price of the maddubs trick).
9970/// Returns raw i8·i8 dots; the caller applies scales and outliers.
9971#[cfg(target_arch = "x86_64")]
9972#[target_feature(enable = "avx2")]
9973unsafe fn dot_i8_i8_avx2_2x4(w0: &[u8], w1: &[u8], xs: [&[i8]; 4]) -> [[i32; 4]; 2] {
9974    // SAFETY: callers uphold slice-length contracts.
9975    unsafe {
9976        use core::arch::x86_64::*;
9977        let n = w0.len();
9978        let ones = _mm256_set1_epi16(1);
9979        let mut acc = [[_mm256_setzero_si256(); 4]; 2];
9980        let mut j = 0usize;
9981        while j + 32 <= n {
9982            let wv0 = _mm256_loadu_si256(w0.as_ptr().add(j) as *const __m256i);
9983            let wv1 = _mm256_loadu_si256(w1.as_ptr().add(j) as *const __m256i);
9984            let aw0 = _mm256_abs_epi8(wv0);
9985            let aw1 = _mm256_abs_epi8(wv1);
9986            for (k, x) in xs.iter().enumerate() {
9987                let xv = _mm256_loadu_si256(x.as_ptr().add(j) as *const __m256i);
9988                let p0 = _mm256_maddubs_epi16(aw0, _mm256_sign_epi8(xv, wv0));
9989                acc[0][k] = _mm256_add_epi32(acc[0][k], _mm256_madd_epi16(p0, ones));
9990                let p1 = _mm256_maddubs_epi16(aw1, _mm256_sign_epi8(xv, wv1));
9991                acc[1][k] = _mm256_add_epi32(acc[1][k], _mm256_madd_epi16(p1, ones));
9992            }
9993            j += 32;
9994        }
9995        let mut out = [[0i32; 4]; 2];
9996        for r in 0..2 {
9997            for k in 0..4 {
9998                let a = acc[r][k];
9999                let hi128 = _mm256_extracti128_si256::<1>(a);
10000                let s128 = _mm_add_epi32(_mm256_castsi256_si128(a), hi128);
10001                let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
10002                let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
10003                out[r][k] = _mm_cvtsi128_si32(s32);
10004            }
10005        }
10006        if j < n {
10007            for (k, x) in xs.iter().enumerate() {
10008                for i in j..n {
10009                    out[0][k] += (w0[i] as i8) as i32 * x[i] as i32;
10010                    out[1][k] += (w1[i] as i8) as i32 * x[i] as i32;
10011                }
10012            }
10013        }
10014        out
10015    }
10016}
10017
10018/// AVX2/VNNI q8 row dot with exact outlier correction (x86 mirror of
10019/// `row_dot_sdot` — same A8W8 contract). With AVX-512 VNNI the row goes
10020/// through the bias trick: Σ(w+128)·x via pure `vpdpbusd` (no per-lane
10021/// sign fixups), corrected by −128·Σx with Σx precomputed per split.
10022#[cfg(target_arch = "x86_64")]
10023#[inline]
10024fn row_dot_avx2(row: &[u8], act: &SplitAct) -> f32 {
10025    let dot = if avx512vnni_enabled() && row.len() >= 64 {
10026        (unsafe { dot_u8p128_i8_vnni(row, &act.xq) }) - 128 * act.xsum
10027    } else {
10028        unsafe { dot_i8_i8_avx2(row, &act.xq) }
10029    };
10030    let mut acc = dot as f32 * act.sx;
10031    for &(j, xv) in &act.outliers {
10032        acc += (row[j] as i8) as f32 * xv;
10033    }
10034    acc
10035}
10036
10037/// Σ (w[i]+128)·x[i] via pure `vpdpbusd` — the caller subtracts
10038/// 128·Σx. Four independent accumulators (dpbusd is ~5-cycle latency;
10039/// a single-acc loop runs latency-bound, measured on Granite Rapids).
10040#[cfg(target_arch = "x86_64")]
10041#[target_feature(enable = "avx2,avx512f,avx512bw,avx512vl,avx512vnni")]
10042unsafe fn dot_u8p128_i8_vnni(w: &[u8], xq: &[i8]) -> i32 {
10043    // SAFETY: callers uphold slice-length contracts (see call sites).
10044    unsafe {
10045        use core::arch::x86_64::*;
10046        let n = w.len();
10047        let flip = _mm512_set1_epi8(-128); // XOR 0x80: i8 w → u8 (w+128)
10048        #[inline(always)]
10049        unsafe fn step(
10050            w: *const u8,
10051            x: *const i8,
10052            flip: core::arch::x86_64::__m512i,
10053            acc: core::arch::x86_64::__m512i,
10054        ) -> core::arch::x86_64::__m512i {
10055            unsafe {
10056                use core::arch::x86_64::*;
10057                let wv = _mm512_xor_si512(_mm512_loadu_si512(w as *const _), flip);
10058                _mm512_dpbusd_epi32(acc, wv, _mm512_loadu_si512(x as *const _))
10059            }
10060        }
10061        let (mut a0, mut a1, mut a2, mut a3) = (
10062            _mm512_setzero_si512(),
10063            _mm512_setzero_si512(),
10064            _mm512_setzero_si512(),
10065            _mm512_setzero_si512(),
10066        );
10067        let mut j = 0usize;
10068        while j + 256 <= n {
10069            a0 = step(w.as_ptr().add(j), xq.as_ptr().add(j), flip, a0);
10070            a1 = step(w.as_ptr().add(j + 64), xq.as_ptr().add(j + 64), flip, a1);
10071            a2 = step(w.as_ptr().add(j + 128), xq.as_ptr().add(j + 128), flip, a2);
10072            a3 = step(w.as_ptr().add(j + 192), xq.as_ptr().add(j + 192), flip, a3);
10073            j += 256;
10074        }
10075        while j + 64 <= n {
10076            a0 = step(w.as_ptr().add(j), xq.as_ptr().add(j), flip, a0);
10077            j += 64;
10078        }
10079        let mut total = _mm512_reduce_add_epi32(_mm512_add_epi32(
10080            _mm512_add_epi32(a0, a1),
10081            _mm512_add_epi32(a2, a3),
10082        ));
10083        // Scalar tail: (w as i8) + 128 ≡ (w as u8) ^ 0x80.
10084        while j < n {
10085            total += ((w[j] ^ 0x80) as i32) * xq[j] as i32;
10086            j += 1;
10087        }
10088        total
10089    }
10090}
10091
10092/// One q4 row via AVX2: nibbles → centered i8 (unpacklo/hi restores the
10093/// writer's flat order, same as the NEON vzip pair), maddubs against
10094/// the pre-quantized activation group, × the group's f16 scale. Pair
10095/// saturation safe: |w|≤8, |x|≤127 → 2·8·127 ≪ 32767. Mirror of
10096/// `dot_q4_row_sdot`.
10097#[cfg(target_arch = "x86_64")]
10098#[target_feature(enable = "avx2")]
10099unsafe fn dot_q4_row_avx2(packed: &[u8], scales: &[u8], g0: usize, gpr: usize, xq: &[i8]) -> f32 {
10100    // SAFETY: callers uphold slice-length contracts (16 packed bytes and
10101    // 2 scale bytes per group; xq.len() == gpr·GROUP_SIZE).
10102    unsafe {
10103        use core::arch::x86_64::*;
10104        let lomask = _mm_set1_epi8(0x0F);
10105        let eight = _mm256_set1_epi8(8);
10106        let ones = _mm256_set1_epi16(1);
10107        let mut acc = 0f32;
10108        for gi in 0..gpr {
10109            let g = g0 + gi;
10110            let s = f16_to_f32(u16::from_le_bytes([scales[g * 2], scales[g * 2 + 1]]));
10111            let b = _mm_loadu_si128(packed.as_ptr().add(g * 16) as *const __m128i);
10112            let lo = _mm_and_si128(b, lomask);
10113            let hi = _mm_and_si128(_mm_srli_epi16::<4>(b), lomask);
10114            let w = _mm256_sub_epi8(
10115                _mm256_set_m128i(_mm_unpackhi_epi8(lo, hi), _mm_unpacklo_epi8(lo, hi)),
10116                eight,
10117            );
10118            let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
10119            let p16 = _mm256_maddubs_epi16(_mm256_abs_epi8(w), _mm256_sign_epi8(x, w));
10120            let d = _mm256_madd_epi16(p16, ones);
10121            let hi128 = _mm256_extracti128_si256::<1>(d);
10122            let s128 = _mm_add_epi32(_mm256_castsi256_si128(d), hi128);
10123            let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
10124            let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
10125            acc += _mm_cvtsi128_si32(s32) as f32 * s;
10126        }
10127        acc
10128    }
10129}
10130
10131/// Two-activation q4 row via AVX2: nibbles unpacked ONCE per group,
10132/// both activations dotted against the same centered i8 register.
10133#[cfg(target_arch = "x86_64")]
10134#[target_feature(enable = "avx2")]
10135unsafe fn dot_q4_row_avx2_2(
10136    packed: &[u8],
10137    scales: &[u8],
10138    g0: usize,
10139    gpr: usize,
10140    xq1: &[i8],
10141    xq2: &[i8],
10142) -> (f32, f32) {
10143    // SAFETY: callers uphold slice-length contracts (see dot_q4_row_avx2).
10144    unsafe {
10145        use core::arch::x86_64::*;
10146        let lomask = _mm_set1_epi8(0x0F);
10147        let eight = _mm256_set1_epi8(8);
10148        let ones = _mm256_set1_epi16(1);
10149        let (mut acc1, mut acc2) = (0f32, 0f32);
10150        #[inline(always)]
10151        unsafe fn hsum(d: core::arch::x86_64::__m256i) -> i32 {
10152            unsafe {
10153                use core::arch::x86_64::*;
10154                let hi128 = _mm256_extracti128_si256::<1>(d);
10155                let s128 = _mm_add_epi32(_mm256_castsi256_si128(d), hi128);
10156                let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
10157                let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
10158                _mm_cvtsi128_si32(s32)
10159            }
10160        }
10161        for gi in 0..gpr {
10162            let g = g0 + gi;
10163            let s = f16_to_f32(u16::from_le_bytes([scales[g * 2], scales[g * 2 + 1]]));
10164            let b = _mm_loadu_si128(packed.as_ptr().add(g * 16) as *const __m128i);
10165            let lo = _mm_and_si128(b, lomask);
10166            let hi = _mm_and_si128(_mm_srli_epi16::<4>(b), lomask);
10167            let w = _mm256_sub_epi8(
10168                _mm256_set_m128i(_mm_unpackhi_epi8(lo, hi), _mm_unpacklo_epi8(lo, hi)),
10169                eight,
10170            );
10171            let aw = _mm256_abs_epi8(w);
10172            let x1 = _mm256_loadu_si256(xq1.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
10173            let x2 = _mm256_loadu_si256(xq2.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
10174            let d1 = _mm256_madd_epi16(_mm256_maddubs_epi16(aw, _mm256_sign_epi8(x1, w)), ones);
10175            let d2 = _mm256_madd_epi16(_mm256_maddubs_epi16(aw, _mm256_sign_epi8(x2, w)), ones);
10176            acc1 += hsum(d1) as f32 * s;
10177            acc2 += hsum(d2) as f32 * s;
10178        }
10179        (acc1, acc2)
10180    }
10181}
10182
10183/// One q8 row range via AVX2 (x86 mirror of `q8_range_sdot`).
10184#[cfg(target_arch = "x86_64")]
10185fn q8_range_avx2(
10186    q: &[u8],
10187    row_scale: &[f32],
10188    act: &SplitAct,
10189    cols: usize,
10190    out_addr: SendMut,
10191    start: usize,
10192    end: usize,
10193) {
10194    for o in start..end {
10195        let v = row_dot_avx2(&q[o * cols..(o + 1) * cols], act) * row_scale[o];
10196        // SAFETY: disjoint row ranges per worker.
10197        unsafe { *out_addr.at(o) = v };
10198    }
10199}
10200
10201/// Two-input q8 row range via AVX2 (x86 mirror of `q8_range2_sdot`).
10202#[cfg(target_arch = "x86_64")]
10203#[allow(clippy::too_many_arguments)]
10204fn q8_range2_avx2(
10205    q: &[u8],
10206    row_scale: &[f32],
10207    a1: &SplitAct,
10208    a2: &SplitAct,
10209    cols: usize,
10210    p1: SendMut,
10211    p2: SendMut,
10212    start: usize,
10213    end: usize,
10214) {
10215    for o in start..end {
10216        let row = &q[o * cols..(o + 1) * cols];
10217        // SAFETY: disjoint row ranges per worker.
10218        unsafe {
10219            *p1.at(o) = row_dot_avx2(row, a1) * row_scale[o];
10220            *p2.at(o) = row_dot_avx2(row, a2) * row_scale[o];
10221        }
10222    }
10223}
10224
10225// ───────────────────── A8W8 SDOT path (port of vmfcore, ×1.78 decode) ─────────────────────
10226
10227/// ARMv8.6 i8mm (smmla): 32 int8 MACs per instruction vs sdot's 16 —
10228/// yet MEASURED 2.4× SLOWER than the blocked sdot on Apple silicon
10229/// (108 vs 264 GF/s): the on-the-fly vcombine packing and the two-
10230/// accumulator dependency chain swamp the MAC advantage, and Apple's
10231/// four SIMD pipes already keep sdot fed. OPT-IN (CMF_I8MM=1) for
10232/// field trials on Cortex-A710/X-class parts with two pipes, where the
10233/// balance may differ; a pre-interleaved weight layout (repack infra)
10234/// is the known path if it ever earns its keep.
10235#[cfg(target_arch = "aarch64")]
10236fn i8mm_enabled() -> bool {
10237    static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
10238    *ON.get_or_init(|| {
10239        std::env::var("CMF_I8MM").map(|v| v == "1").unwrap_or(false)
10240            && std::arch::is_aarch64_feature_detected!("i8mm")
10241    })
10242}
10243
10244/// SDOT enabled? Default ON when the CPU has ARMv8.2 dotprod;
10245/// `CMF_SDOT=0` disables (falls back to i8×f32 NEON).
10246/// (On non-ARM release builds only the test tolerance switch calls it.)
10247#[cfg_attr(not(target_arch = "aarch64"), allow(dead_code))]
10248fn sdot_enabled() -> bool {
10249    if FLOAT_ACTIVATIONS.get() {
10250        return false;
10251    }
10252    use std::sync::OnceLock;
10253    static ON: OnceLock<bool> = OnceLock::new();
10254    *ON.get_or_init(|| {
10255        let want = std::env::var("CMF_SDOT").map(|v| v != "0").unwrap_or(true);
10256        if !want {
10257            return false;
10258        }
10259
10260        #[cfg(target_arch = "aarch64")]
10261        {
10262            if std::arch::is_aarch64_feature_detected!("dotprod") {
10263                return true;
10264            }
10265            #[cfg(target_os = "android")]
10266            {
10267                if let Ok(cpuinfo) = std::fs::read_to_string("/proc/cpuinfo") {
10268                    if cpuinfo.lines().any(|l| {
10269                        (l.starts_with("Features") || l.starts_with("features"))
10270                            && l.contains("asimddp")
10271                    }) {
10272                        return true;
10273                    }
10274                }
10275            }
10276            false
10277        }
10278        #[cfg(not(target_arch = "aarch64"))]
10279        {
10280            false
10281        }
10282    })
10283}
10284
10285/// Two-field activation split (≡ vmfcore `q8_split_prep`): outlier
10286/// channels (>8·rms) are computed exactly in f32; the bulk (outliers
10287/// zeroed → clean absmax) goes through int8 SDOT. Computed ONCE per
10288/// matvec, shared by all rows/workers.
10289struct SplitAct {
10290    xq: Vec<i8>,
10291    sx: f32,
10292    outliers: Vec<(usize, f32)>,
10293    /// Σ xq — the VNNI bias-trick correction (`(w+128)·x` sums need
10294    /// `−128·Σx`); one i32 per split, computed once per matvec.
10295    #[cfg_attr(not(target_arch = "x86_64"), allow(dead_code))]
10296    xsum: i32,
10297}
10298
10299thread_local! {
10300    /// Recycled xq buffers: split_act runs for every matvec (~200/token)
10301    /// and its hidden-size allocation was steady-state heap churn.
10302    static XQ_FREE: std::cell::RefCell<Vec<Vec<i8>>> =
10303        const { std::cell::RefCell::new(Vec::new()) };
10304}
10305
10306impl Drop for SplitAct {
10307    fn drop(&mut self) {
10308        let buf = std::mem::take(&mut self.xq);
10309        if buf.capacity() > 0 {
10310            XQ_FREE.with(|f| {
10311                let mut f = f.borrow_mut();
10312                if f.len() < 16 {
10313                    f.push(buf);
10314                }
10315            });
10316        }
10317    }
10318}
10319
10320thread_local! {
10321    /// One scratch row per WORKER, kept for the life of the thread.
10322    ///
10323    /// The kernels take a row of group scales per dispatch, and a fresh
10324    /// `vec![0f32; gpr]` inside the closure is one allocation per worker per
10325    /// dispatch — on the release checkpoint about six thousand a token, a
10326    /// quarter of everything the benchmark counts.
10327    static KROW: std::cell::RefCell<Vec<f32>> = const { std::cell::RefCell::new(Vec::new()) };
10328}
10329
10330/// Borrow `n` floats of the calling worker's scratch. Nothing inside a
10331/// kernel body borrows it again, which is what keeps the RefCell honest.
10332#[inline]
10333fn with_krow<R>(n: usize, f: impl FnOnce(&mut [f32]) -> R) -> R {
10334    KROW.with(|s| {
10335        let mut b = s.borrow_mut();
10336        if b.len() < n {
10337            b.resize(n, 0.0);
10338        }
10339        f(&mut b[..n])
10340    })
10341}
10342
10343/// `t.round().clamp(-127.0, 127.0) as i8`, bit for bit, without the libm
10344/// call. On baseline x86-64 (no SSE4.1 `roundps`) `f32::round` is a
10345/// function call per element, and split_act runs it over every hidden
10346/// state before every matvec: measured 27 us a call on a 2048-wide
10347/// activation on an EPYC 7763 — 5.4 ms of a 55 ms decode token, all of
10348/// it on the caller's thread while thirty workers wait. Clamping first is
10349/// equivalent (round is monotonic and ±127 are integers), and after the
10350/// clamp `t - trunc(t)` is exact, so the half-away-from-zero decision is
10351/// the one `round` makes. NaN clamps to NaN and converts to 0, as before.
10352/// The loop vectorizes (cvttps2dq + compare/select).
10353#[inline(always)]
10354fn q8_round(t: f32) -> i8 {
10355    let t = t.clamp(-127.0, 127.0);
10356    let i = t as i32;
10357    let f = t - i as f32;
10358    let r = if f >= 0.5 {
10359        i + 1
10360    } else if f <= -0.5 {
10361        i - 1
10362    } else {
10363        i
10364    };
10365    r as i8
10366}
10367
10368fn split_act(x: &[f32]) -> SplitAct {
10369    let _prof = crate::cpuprof::time(crate::cpuprof::Slot::SplitAct);
10370    let n = x.len();
10371    let rms = (x.iter().map(|&v| (v * v) as f64).sum::<f64>() / n.max(1) as f64).sqrt() as f32;
10372    let thr = 8.0 * rms;
10373    // One pass: collect outliers and the bulk absmax (outliers excluded —
10374    // identical to the old zero-then-fold over a copied buffer, minus the
10375    // full-vector copy).
10376    let mut outliers: Vec<(usize, f32)> = Vec::new();
10377    let mut amax = 0f32;
10378    for (j, &v) in x.iter().enumerate() {
10379        let a = v.abs();
10380        if a > thr {
10381            outliers.push((j, v));
10382        } else if a > amax {
10383            amax = a;
10384        }
10385    }
10386    let sx = if amax > 0.0 { amax / 127.0 } else { 1.0 };
10387    let inv = 1.0 / sx;
10388    let mut xq = XQ_FREE.with(|f| f.borrow_mut().pop()).unwrap_or_default();
10389    xq.clear();
10390    xq.reserve(n);
10391    if outliers.is_empty() {
10392        xq.extend(
10393            x.iter()
10394                .map(|&v| q8_round(v * inv)),
10395        );
10396    } else {
10397        // Outlier slots quantize to 0 (their exact term is added later).
10398        xq.extend(x.iter().map(|&v| {
10399            if v.abs() > thr {
10400                0
10401            } else {
10402                q8_round(v * inv)
10403            }
10404        }));
10405    }
10406    let xsum = xq.iter().map(|&v| v as i32).sum();
10407    SplitAct {
10408        xq,
10409        sx,
10410        outliers,
10411        xsum,
10412    }
10413}
10414
10415fn split_act_q8_2f(x: &[f32], col: &[f32]) -> SplitAct {
10416    let _prof = crate::cpuprof::time(crate::cpuprof::Slot::SplitAct);
10417    let n = x.len();
10418    let rms = (x
10419        .iter()
10420        .zip(col)
10421        .map(|(&a, &c)| {
10422            let v = a * c;
10423            (v * v) as f64
10424        })
10425        .sum::<f64>()
10426        / n.max(1) as f64)
10427        .sqrt() as f32;
10428    let thr = 8.0 * rms;
10429
10430    let mut outliers = Vec::new();
10431    let mut amax = 0f32;
10432    for (j, (&a, &c)) in x.iter().zip(col).enumerate() {
10433        let v = a * c;
10434        let s = v.abs();
10435        if s > thr {
10436            outliers.push((j, v));
10437        } else if s > amax {
10438            amax = s;
10439        }
10440    }
10441
10442    let sx = if amax > 0.0 { amax / 127.0 } else { 1.0 };
10443    let inv = 1.0 / sx;
10444    let mut xq = XQ_FREE.with(|f| f.borrow_mut().pop()).unwrap_or_default();
10445    xq.clear();
10446    xq.reserve(n);
10447    if outliers.is_empty() {
10448        xq.extend(
10449            x.iter()
10450                .zip(col)
10451                .map(|(&a, &c)| q8_round((a * c) * inv)),
10452        );
10453    } else {
10454        xq.extend(x.iter().zip(col).map(|(&a, &c)| {
10455            let v = a * c;
10456            if v.abs() > thr {
10457                0
10458            } else {
10459                q8_round(v * inv)
10460            }
10461        }));
10462    }
10463    let xsum = xq.iter().map(|&v| v as i32).sum();
10464    SplitAct {
10465        xq,
10466        sx,
10467        outliers,
10468        xsum,
10469    }
10470}
10471
10472/// int8(weight)·int8(activation) → i32 via `sdot` (inline asm — the
10473/// vdotq intrinsic is unstable; port of vmfcore `dot_i8_sdot`).
10474#[cfg(target_arch = "aarch64")]
10475#[target_feature(enable = "neon,dotprod")]
10476unsafe fn dot_i8_sdot(w: &[u8], xq: &[i8]) -> i32 {
10477    // SAFETY: callers uphold slice-length contracts (see call sites).
10478    unsafe {
10479        use core::arch::aarch64::*;
10480        use core::arch::asm;
10481        let wp = w.as_ptr() as *const i8;
10482        let n = w.len();
10483        let (mut a0, mut a1, mut a2, mut a3) = (
10484            vdupq_n_s32(0),
10485            vdupq_n_s32(0),
10486            vdupq_n_s32(0),
10487            vdupq_n_s32(0),
10488        );
10489        let mut i = 0;
10490        while i + 64 <= n {
10491            let (w0, x0) = (vld1q_s8(wp.add(i)), vld1q_s8(xq.as_ptr().add(i)));
10492            let (w1, x1) = (vld1q_s8(wp.add(i + 16)), vld1q_s8(xq.as_ptr().add(i + 16)));
10493            let (w2, x2) = (vld1q_s8(wp.add(i + 32)), vld1q_s8(xq.as_ptr().add(i + 32)));
10494            let (w3, x3) = (vld1q_s8(wp.add(i + 48)), vld1q_s8(xq.as_ptr().add(i + 48)));
10495            asm!(
10496                "sdot {a0:v}.4s, {w0:v}.16b, {x0:v}.16b",
10497                "sdot {a1:v}.4s, {w1:v}.16b, {x1:v}.16b",
10498                "sdot {a2:v}.4s, {w2:v}.16b, {x2:v}.16b",
10499                "sdot {a3:v}.4s, {w3:v}.16b, {x3:v}.16b",
10500                a0 = inout(vreg) a0, a1 = inout(vreg) a1, a2 = inout(vreg) a2, a3 = inout(vreg) a3,
10501                w0 = in(vreg) w0, x0 = in(vreg) x0, w1 = in(vreg) w1, x1 = in(vreg) x1,
10502                w2 = in(vreg) w2, x2 = in(vreg) x2, w3 = in(vreg) w3, x3 = in(vreg) x3,
10503                options(pure, nomem, nostack),
10504            );
10505            i += 64;
10506        }
10507        while i + 16 <= n {
10508            let (wv, xv) = (vld1q_s8(wp.add(i)), vld1q_s8(xq.as_ptr().add(i)));
10509            asm!("sdot {a:v}.4s, {w:v}.16b, {x:v}.16b",
10510                 a = inout(vreg) a0, w = in(vreg) wv, x = in(vreg) xv, options(pure, nomem, nostack));
10511            i += 16;
10512        }
10513        let mut s = vaddvq_s32(vaddq_s32(vaddq_s32(a0, a1), vaddq_s32(a2, a3)));
10514        while i < n {
10515            s += (*wp.add(i)) as i32 * xq[i] as i32;
10516            i += 1;
10517        }
10518        s
10519    }
10520}
10521
10522/// Row-blocked SDOT: 4 output rows per pass — the activation chunk is
10523/// loaded once and reused, 4 independent accumulators hide sdot latency
10524/// (port of vmfcore `dot_i8_sdot_4rows`).
10525#[cfg(target_arch = "aarch64")]
10526#[target_feature(enable = "neon,dotprod")]
10527unsafe fn dot_i8_sdot_4rows(w0: &[u8], w1: &[u8], w2: &[u8], w3: &[u8], xq: &[i8]) -> [i32; 4] {
10528    // SAFETY: callers uphold slice-length contracts (see call sites).
10529    unsafe {
10530        use core::arch::aarch64::*;
10531        use core::arch::asm;
10532        let n = xq.len();
10533        let px = xq.as_ptr();
10534        let (p0, p1, p2, p3) = (
10535            w0.as_ptr() as *const i8,
10536            w1.as_ptr() as *const i8,
10537            w2.as_ptr() as *const i8,
10538            w3.as_ptr() as *const i8,
10539        );
10540        let (mut a0, mut a1, mut a2, mut a3) = (
10541            vdupq_n_s32(0),
10542            vdupq_n_s32(0),
10543            vdupq_n_s32(0),
10544            vdupq_n_s32(0),
10545        );
10546        let mut i = 0;
10547        while i + 16 <= n {
10548            let x = vld1q_s8(px.add(i));
10549            let v0 = vld1q_s8(p0.add(i));
10550            let v1 = vld1q_s8(p1.add(i));
10551            let v2 = vld1q_s8(p2.add(i));
10552            let v3 = vld1q_s8(p3.add(i));
10553            asm!(
10554                "sdot {a0:v}.4s, {v0:v}.16b, {x:v}.16b",
10555                "sdot {a1:v}.4s, {v1:v}.16b, {x:v}.16b",
10556                "sdot {a2:v}.4s, {v2:v}.16b, {x:v}.16b",
10557                "sdot {a3:v}.4s, {v3:v}.16b, {x:v}.16b",
10558                a0 = inout(vreg) a0, a1 = inout(vreg) a1, a2 = inout(vreg) a2, a3 = inout(vreg) a3,
10559                v0 = in(vreg) v0, v1 = in(vreg) v1, v2 = in(vreg) v2, v3 = in(vreg) v3, x = in(vreg) x,
10560                options(pure, nomem, nostack),
10561            );
10562            i += 16;
10563        }
10564        let mut r = [
10565            vaddvq_s32(a0),
10566            vaddvq_s32(a1),
10567            vaddvq_s32(a2),
10568            vaddvq_s32(a3),
10569        ];
10570        while i < n {
10571            let xi = *px.add(i) as i32;
10572            r[0] += (*p0.add(i)) as i32 * xi;
10573            r[1] += (*p1.add(i)) as i32 * xi;
10574            r[2] += (*p2.add(i)) as i32 * xi;
10575            r[3] += (*p3.add(i)) as i32 * xi;
10576            i += 1;
10577        }
10578        r
10579    }
10580}
10581
10582/// 4 interleaved rows in one pass: the repacked group is [r0[c], r1[c],
10583/// r2[c], r3[c]] per 16-byte chunk, so each iteration reads ONE 64-byte
10584/// line plus the shared activation chunk — a single sequential weight
10585/// stream per worker. Per-row accumulation is the same one-accumulator
10586/// scheme as `dot_i8_sdot_4rows`; integer sums are exact, so outputs
10587/// are bit-identical to the mmap-layout kernel.
10588#[cfg(target_arch = "aarch64")]
10589#[target_feature(enable = "neon,dotprod")]
10590unsafe fn dot_i8_sdot_4rows_il(g: &[u8], xq: &[i8]) -> [i32; 4] {
10591    // SAFETY: callers uphold slice-length contracts (g.len() == 4·n,
10592    // n % 16 == 0 — guaranteed by the repack gate).
10593    unsafe {
10594        use core::arch::aarch64::*;
10595        use core::arch::asm;
10596        let n = xq.len();
10597        let px = xq.as_ptr();
10598        let pg = g.as_ptr() as *const i8;
10599        let (mut a0, mut a1, mut a2, mut a3) = (
10600            vdupq_n_s32(0),
10601            vdupq_n_s32(0),
10602            vdupq_n_s32(0),
10603            vdupq_n_s32(0),
10604        );
10605        let mut i = 0;
10606        while i + 16 <= n {
10607            let x = vld1q_s8(px.add(i));
10608            let base = pg.add(4 * i);
10609            let v0 = vld1q_s8(base);
10610            let v1 = vld1q_s8(base.add(16));
10611            let v2 = vld1q_s8(base.add(32));
10612            let v3 = vld1q_s8(base.add(48));
10613            asm!(
10614                "sdot {a0:v}.4s, {v0:v}.16b, {x:v}.16b",
10615                "sdot {a1:v}.4s, {v1:v}.16b, {x:v}.16b",
10616                "sdot {a2:v}.4s, {v2:v}.16b, {x:v}.16b",
10617                "sdot {a3:v}.4s, {v3:v}.16b, {x:v}.16b",
10618                a0 = inout(vreg) a0, a1 = inout(vreg) a1, a2 = inout(vreg) a2, a3 = inout(vreg) a3,
10619                v0 = in(vreg) v0, v1 = in(vreg) v1, v2 = in(vreg) v2, v3 = in(vreg) v3, x = in(vreg) x,
10620                options(pure, nomem, nostack),
10621            );
10622            i += 16;
10623        }
10624        [
10625            vaddvq_s32(a0),
10626            vaddvq_s32(a1),
10627            vaddvq_s32(a2),
10628            vaddvq_s32(a3),
10629        ]
10630    }
10631}
10632
10633/// One q8 row range via SDOT (4-row blocks + tail) — the body of
10634/// `qmatvec`'s hot loop, extracted so multi-matrix jobs can drive the
10635/// SAME kernel for several tensors under one pool dispatch. `rep` — the
10636/// load-time interleaved repack (empty = mmap layout only); rows outside
10637/// full 4-row groups always come from the mmap layout.
10638#[cfg(target_arch = "aarch64")]
10639fn q8_range_sdot(
10640    q: &[u8],
10641    rep: &[u8],
10642    row_scale: &[f32],
10643    act: &SplitAct,
10644    cols: usize,
10645    out_addr: SendMut,
10646    start: usize,
10647    end: usize,
10648) {
10649    let mut o = start;
10650    // Leading rows to the group boundary (repack path only): the pool
10651    // splits row ranges arbitrarily, groups are absolute.
10652    if !rep.is_empty() {
10653        while o < end && o % 4 != 0 {
10654            let v = row_dot_sdot(&q[o * cols..(o + 1) * cols], act) * row_scale[o];
10655            unsafe { *out_addr.at(o) = v };
10656            o += 1;
10657        }
10658    }
10659    while o + 4 <= end {
10660        let r = if rep.is_empty() {
10661            unsafe {
10662                dot_i8_sdot_4rows(
10663                    &q[o * cols..(o + 1) * cols],
10664                    &q[(o + 1) * cols..(o + 2) * cols],
10665                    &q[(o + 2) * cols..(o + 3) * cols],
10666                    &q[(o + 3) * cols..(o + 4) * cols],
10667                    &act.xq,
10668                )
10669            }
10670        } else {
10671            unsafe { dot_i8_sdot_4rows_il(&rep[o * cols..(o + 4) * cols], &act.xq) }
10672        };
10673        for k in 0..4 {
10674            let mut acc = r[k] as f32 * act.sx;
10675            for &(j, xv) in &act.outliers {
10676                acc += (q[(o + k) * cols + j] as i8) as f32 * xv;
10677            }
10678            // SAFETY: disjoint row ranges per worker.
10679            unsafe { *out_addr.at(o + k) = acc * row_scale[o + k] };
10680        }
10681        o += 4;
10682    }
10683    while o < end {
10684        let v = row_dot_sdot(&q[o * cols..(o + 1) * cols], act) * row_scale[o];
10685        unsafe { *out_addr.at(o) = v };
10686        o += 1;
10687    }
10688}
10689
10690/// Two-input q8 row range via SDOT — `qmatvec2`'s hot loop, extracted
10691/// for the fused pair multi-matrix job (`matvec2_many`).
10692#[cfg(target_arch = "aarch64")]
10693#[allow(clippy::too_many_arguments)]
10694fn q8_range2_sdot(
10695    q: &[u8],
10696    row_scale: &[f32],
10697    a1: &SplitAct,
10698    a2: &SplitAct,
10699    cols: usize,
10700    p1: SendMut,
10701    p2: SendMut,
10702    start: usize,
10703    end: usize,
10704) {
10705    for o in start..end {
10706        let row = &q[o * cols..(o + 1) * cols];
10707        // SAFETY: disjoint row ranges per worker.
10708        unsafe {
10709            *p1.at(o) = row_dot_sdot(row, a1) * row_scale[o];
10710            *p2.at(o) = row_dot_sdot(row, a2) * row_scale[o];
10711        }
10712    }
10713}
10714
10715/// Two-input q8 row range, f32 kernel (non-SDOT) — same extraction.
10716#[allow(clippy::too_many_arguments)]
10717fn q8_range2_f32(
10718    q: &[u8],
10719    row_scale: &[f32],
10720    x1: &[f32],
10721    x2: &[f32],
10722    cols: usize,
10723    p1: SendMut,
10724    p2: SendMut,
10725    start: usize,
10726    end: usize,
10727) {
10728    for o in start..end {
10729        let row = &q[o * cols..(o + 1) * cols];
10730        // SAFETY: disjoint row ranges per worker.
10731        unsafe {
10732            *p1.at(o) = dot_i8_f32(row, x1) * row_scale[o];
10733            *p2.at(o) = dot_i8_f32(row, x2) * row_scale[o];
10734        }
10735    }
10736}
10737
10738/// Scalar/NEON-f32 q8 row range (non-SDOT platforms) — same extraction.
10739fn q8_range_f32(
10740    q: &[u8],
10741    row_scale: &[f32],
10742    xs: &[f32],
10743    cols: usize,
10744    out_addr: SendMut,
10745    start: usize,
10746    end: usize,
10747) {
10748    for o in start..end {
10749        let v = dot_i8_f32(&q[o * cols..(o + 1) * cols], xs) * row_scale[o];
10750        // SAFETY: disjoint row ranges per worker.
10751        unsafe { *out_addr.at(o) = v };
10752    }
10753}
10754
10755/// One q8 row against a split activation, portable: the per-arch fast
10756/// dots where they exist, the exact scalar loop elsewhere. The scalar
10757/// arm is also the test oracle for both fast arms.
10758#[inline]
10759fn q8_row_dot(row: &[u8], act: &SplitAct) -> f32 {
10760    #[cfg(target_arch = "aarch64")]
10761    return row_dot_sdot(row, act);
10762    #[cfg(target_arch = "x86_64")]
10763    return row_dot_avx2(row, act);
10764    #[allow(unreachable_code)]
10765    q8_row_dot_scalar(row, act)
10766}
10767
10768#[allow(dead_code)]
10769fn q8_row_dot_scalar(row: &[u8], act: &SplitAct) -> f32 {
10770    let mut acc = 0i32;
10771    for (k, &b) in row.iter().enumerate() {
10772        acc += (b as i8) as i32 * act.xq[k] as i32;
10773    }
10774    let mut acc = acc as f32 * act.sx;
10775    for &(j, xv) in &act.outliers {
10776        acc += (row[j] as i8) as f32 * xv;
10777    }
10778    acc
10779}
10780
10781/// SDOT row dot with exact outlier correction:
10782/// `dot = sdot(w, xq)·sx + Σ_outl w[j]·x[j]` (then × row_scale by caller).
10783#[cfg(target_arch = "aarch64")]
10784#[inline]
10785fn row_dot_sdot(row: &[u8], act: &SplitAct) -> f32 {
10786    let mut acc = unsafe { dot_i8_sdot(row, &act.xq) } as f32 * act.sx;
10787    for &(j, xv) in &act.outliers {
10788        acc += (row[j] as i8) as f32 * xv;
10789    }
10790    acc
10791}
10792
10793/// One q4 row via SDOT: each 32-group's nibbles unpack to centered i8
10794/// (nib−8 ∈ [−8,7]), int8×int8 `sdot` against the pre-quantized
10795/// activation group, × the group's f16 scale. Returns Σ_g dot_g·s_g;
10796/// the caller multiplies by the activation scale and adds the exact
10797/// outlier terms (port of vmfcore `dot_q4_block_sdot`, +23% measured).
10798/// Nibble order matches the writer: element 2k = low nibble, 2k+1 = high
10799/// → zip(lo,hi) restores flat order.
10800#[cfg(target_arch = "aarch64")]
10801#[target_feature(enable = "neon,dotprod")]
10802unsafe fn dot_q4_row_sdot(packed: &[u8], scales: &[u8], g0: usize, gpr: usize, xq: &[i8]) -> f32 {
10803    // SAFETY: callers uphold slice-length contracts (16 packed bytes and
10804    // 2 scale bytes per group; xq.len() == gpr·GROUP_SIZE).
10805    unsafe {
10806        use core::arch::aarch64::*;
10807        use core::arch::asm;
10808        let lomask = vdupq_n_u8(0x0F);
10809        let eight = vdupq_n_s8(8);
10810        let mut acc = 0f32;
10811        for gi in 0..gpr {
10812            let g = g0 + gi;
10813            let s = f16_to_f32(u16::from_le_bytes([scales[g * 2], scales[g * 2 + 1]]));
10814            let b = vld1q_u8(packed.as_ptr().add(g * 16));
10815            let lo = vandq_u8(b, lomask);
10816            let hi = vshrq_n_u8::<4>(b);
10817            let e0 = vsubq_s8(vreinterpretq_s8_u8(vzip1q_u8(lo, hi)), eight);
10818            let e1 = vsubq_s8(vreinterpretq_s8_u8(vzip2q_u8(lo, hi)), eight);
10819            let x0 = vld1q_s8(xq.as_ptr().add(gi * GROUP_SIZE));
10820            let x1 = vld1q_s8(xq.as_ptr().add(gi * GROUP_SIZE + 16));
10821            let (mut a0, mut a1) = (vdupq_n_s32(0), vdupq_n_s32(0));
10822            asm!(
10823                "sdot {a0:v}.4s, {e0:v}.16b, {x0:v}.16b",
10824                "sdot {a1:v}.4s, {e1:v}.16b, {x1:v}.16b",
10825                a0 = inout(vreg) a0, a1 = inout(vreg) a1,
10826                e0 = in(vreg) e0, x0 = in(vreg) x0, e1 = in(vreg) e1, x1 = in(vreg) x1,
10827                options(pure, nomem, nostack),
10828            );
10829            acc += vaddvq_s32(vaddq_s32(a0, a1)) as f32 * s;
10830        }
10831        acc
10832    }
10833}
10834
10835/// Two-activation q4 row via SDOT: the nibble unpack (the expensive
10836/// part) happens ONCE per group; both pre-quantized activations are
10837/// dotted against the same centered i8 registers. Per-lane math matches
10838/// `dot_q4_row_sdot` exactly.
10839#[cfg(target_arch = "aarch64")]
10840#[target_feature(enable = "neon,dotprod")]
10841unsafe fn dot_q4_row_sdot2(
10842    packed: &[u8],
10843    scales: &[u8],
10844    g0: usize,
10845    gpr: usize,
10846    xq1: &[i8],
10847    xq2: &[i8],
10848) -> (f32, f32) {
10849    // SAFETY: callers uphold slice-length contracts (16 packed bytes and
10850    // 2 scale bytes per group; xq*.len() == gpr·GROUP_SIZE).
10851    unsafe {
10852        use core::arch::aarch64::*;
10853        use core::arch::asm;
10854        let lomask = vdupq_n_u8(0x0F);
10855        let eight = vdupq_n_s8(8);
10856        let (mut acc1, mut acc2) = (0f32, 0f32);
10857        for gi in 0..gpr {
10858            let g = g0 + gi;
10859            let s = f16_to_f32(u16::from_le_bytes([scales[g * 2], scales[g * 2 + 1]]));
10860            let b = vld1q_u8(packed.as_ptr().add(g * 16));
10861            let lo = vandq_u8(b, lomask);
10862            let hi = vshrq_n_u8::<4>(b);
10863            let e0 = vsubq_s8(vreinterpretq_s8_u8(vzip1q_u8(lo, hi)), eight);
10864            let e1 = vsubq_s8(vreinterpretq_s8_u8(vzip2q_u8(lo, hi)), eight);
10865            let x10 = vld1q_s8(xq1.as_ptr().add(gi * GROUP_SIZE));
10866            let x11 = vld1q_s8(xq1.as_ptr().add(gi * GROUP_SIZE + 16));
10867            let x20 = vld1q_s8(xq2.as_ptr().add(gi * GROUP_SIZE));
10868            let x21 = vld1q_s8(xq2.as_ptr().add(gi * GROUP_SIZE + 16));
10869            let (mut a0, mut a1, mut b0, mut b1) = (
10870                vdupq_n_s32(0),
10871                vdupq_n_s32(0),
10872                vdupq_n_s32(0),
10873                vdupq_n_s32(0),
10874            );
10875            asm!(
10876                "sdot {a0:v}.4s, {e0:v}.16b, {x10:v}.16b",
10877                "sdot {a1:v}.4s, {e1:v}.16b, {x11:v}.16b",
10878                "sdot {b0:v}.4s, {e0:v}.16b, {x20:v}.16b",
10879                "sdot {b1:v}.4s, {e1:v}.16b, {x21:v}.16b",
10880                a0 = inout(vreg) a0, a1 = inout(vreg) a1,
10881                b0 = inout(vreg) b0, b1 = inout(vreg) b1,
10882                e0 = in(vreg) e0, e1 = in(vreg) e1,
10883                x10 = in(vreg) x10, x11 = in(vreg) x11,
10884                x20 = in(vreg) x20, x21 = in(vreg) x21,
10885                options(pure, nomem, nostack),
10886            );
10887            acc1 += vaddvq_s32(vaddq_s32(a0, a1)) as f32 * s;
10888            acc2 += vaddvq_s32(vaddq_s32(b0, b1)) as f32 * s;
10889        }
10890        (acc1, acc2)
10891    }
10892}
10893
10894// ───────────────────── fused int8 kernels ─────────────────────
10895
10896/// `acc += w · row` where the row is centered i8 — NEON widen+fma on
10897/// aarch64, scalar elsewhere. The KV-cache q8 value path rides on this.
10898#[inline]
10899pub(crate) fn axpy_i8_f32(acc: &mut [f32], row: &[i8], w: f32) {
10900    #[cfg(target_arch = "aarch64")]
10901    unsafe {
10902        return axpy_i8_f32_neon(acc, row, w);
10903    }
10904    #[cfg(target_arch = "x86_64")]
10905    if avx2_enabled() {
10906        return unsafe { axpy_i8_f32_avx2(acc, row, w) };
10907    }
10908    #[allow(unreachable_code)]
10909    {
10910        for (a, &b) in acc.iter_mut().zip(row) {
10911            *a += w * b as f32;
10912        }
10913    }
10914}
10915
10916/// i8→f32 axpy via AVX2/FMA (x86 mirror of `axpy_i8_f32_neon`).
10917#[cfg(target_arch = "x86_64")]
10918#[target_feature(enable = "avx2,fma")]
10919unsafe fn axpy_i8_f32_avx2(acc: &mut [f32], row: &[i8], w: f32) {
10920    // SAFETY: callers uphold slice-length contracts (see call sites).
10921    unsafe {
10922        use core::arch::x86_64::*;
10923        let n = acc.len().min(row.len());
10924        let ap = acc.as_mut_ptr();
10925        let rp = row.as_ptr();
10926        let wv = _mm256_set1_ps(w);
10927        let mut j = 0usize;
10928        while j + 16 <= n {
10929            let rb = _mm_loadu_si128(rp.add(j) as *const __m128i);
10930            let lo = _mm256_cvtepi32_ps(_mm256_cvtepi8_epi32(rb));
10931            let hi = _mm256_cvtepi32_ps(_mm256_cvtepi8_epi32(_mm_srli_si128::<8>(rb)));
10932            let v0 = _mm256_fmadd_ps(wv, lo, _mm256_loadu_ps(ap.add(j)));
10933            let v1 = _mm256_fmadd_ps(wv, hi, _mm256_loadu_ps(ap.add(j + 8)));
10934            _mm256_storeu_ps(ap.add(j), v0);
10935            _mm256_storeu_ps(ap.add(j + 8), v1);
10936            j += 16;
10937        }
10938        while j < n {
10939            *ap.add(j) += w * (*rp.add(j)) as f32;
10940            j += 1;
10941        }
10942    }
10943}
10944
10945#[cfg(target_arch = "aarch64")]
10946#[target_feature(enable = "neon")]
10947unsafe fn axpy_i8_f32_neon(acc: &mut [f32], row: &[i8], w: f32) {
10948    // SAFETY: callers uphold slice-length contracts (see call sites).
10949    unsafe {
10950        use core::arch::aarch64::*;
10951        let n = acc.len().min(row.len());
10952        let ap = acc.as_mut_ptr();
10953        let rp = row.as_ptr();
10954        let wv = vdupq_n_f32(w);
10955        let mut j = 0usize;
10956        while j + 16 <= n {
10957            let rb = vld1q_s8(rp.add(j));
10958            let lo = vmovl_s8(vget_low_s8(rb));
10959            let hi = vmovl_s8(vget_high_s8(rb));
10960            for (off, half) in [(0, lo), (8, hi)] {
10961                let f0 = vcvtq_f32_s32(vmovl_s16(vget_low_s16(half)));
10962                let f1 = vcvtq_f32_s32(vmovl_s16(vget_high_s16(half)));
10963                let o = j + off;
10964                vst1q_f32(ap.add(o), vfmaq_f32(vld1q_f32(ap.add(o)), wv, f0));
10965                vst1q_f32(ap.add(o + 4), vfmaq_f32(vld1q_f32(ap.add(o + 4)), wv, f1));
10966            }
10967            j += 16;
10968        }
10969        while j < n {
10970            *ap.add(j) += w * (*rp.add(j)) as f32;
10971            j += 1;
10972        }
10973    }
10974}
10975
10976/// i8 row · f32 x. NEON on aarch64 (ported from vmfcore `dot_i8_f32_neon`,
10977/// ≈9× scalar), scalar elsewhere.
10978#[inline]
10979pub(crate) fn dot_i8_f32(w: &[u8], x: &[f32]) -> f32 {
10980    #[cfg(target_arch = "aarch64")]
10981    unsafe {
10982        return dot_i8_f32_neon(w, x);
10983    }
10984    #[cfg(target_arch = "x86_64")]
10985    if avx2_enabled() {
10986        return unsafe { dot_i8_f32_avx2(w, x) };
10987    }
10988    #[allow(unreachable_code)]
10989    {
10990        let mut sum = 0.0f32;
10991        for (j, &b) in w.iter().enumerate() {
10992            sum += (b as i8) as f32 * x[j];
10993        }
10994        sum
10995    }
10996}
10997
10998/// i8 row · (x ⊙ col_field) — the q8_2f row dot with the θ col-field
10999/// folded into the product (no prescaled copy of x). NEON on aarch64,
11000/// scalar elsewhere. Used by the active-neuron path `row_dot`.
11001#[inline]
11002fn dot_i8_col_f32(w: &[u8], x: &[f32], col: &[f32]) -> f32 {
11003    #[cfg(target_arch = "aarch64")]
11004    unsafe {
11005        return dot_i8_col_f32_neon(w, x, col);
11006    }
11007    #[allow(unreachable_code)]
11008    {
11009        let mut sum = 0.0f32;
11010        for (j, &b) in w.iter().enumerate() {
11011            sum += (b as i8) as f32 * x[j] * col[j];
11012        }
11013        sum
11014    }
11015}
11016
11017#[cfg(target_arch = "aarch64")]
11018#[target_feature(enable = "neon")]
11019unsafe fn dot_i8_col_f32_neon(w: &[u8], x: &[f32], col: &[f32]) -> f32 {
11020    // SAFETY: callers uphold slice-length contracts (see call sites).
11021    unsafe {
11022        use core::arch::aarch64::*;
11023        let n = x.len();
11024        let wp = w.as_ptr() as *const i8;
11025        let xp = x.as_ptr();
11026        let cp = col.as_ptr();
11027        let (mut a0, mut a1, mut a2, mut a3) = (
11028            vdupq_n_f32(0.0),
11029            vdupq_n_f32(0.0),
11030            vdupq_n_f32(0.0),
11031            vdupq_n_f32(0.0),
11032        );
11033        let mut j = 0usize;
11034        while j + 16 <= n {
11035            let wb = vld1q_s8(wp.add(j));
11036            let lo = vmovl_s8(vget_low_s8(wb));
11037            let hi = vmovl_s8(vget_high_s8(wb));
11038            let w0 = vcvtq_f32_s32(vmovl_s16(vget_low_s16(lo)));
11039            let w1 = vcvtq_f32_s32(vmovl_s16(vget_high_s16(lo)));
11040            let w2 = vcvtq_f32_s32(vmovl_s16(vget_low_s16(hi)));
11041            let w3 = vcvtq_f32_s32(vmovl_s16(vget_high_s16(hi)));
11042            a0 = vfmaq_f32(
11043                a0,
11044                w0,
11045                vmulq_f32(vld1q_f32(xp.add(j)), vld1q_f32(cp.add(j))),
11046            );
11047            a1 = vfmaq_f32(
11048                a1,
11049                w1,
11050                vmulq_f32(vld1q_f32(xp.add(j + 4)), vld1q_f32(cp.add(j + 4))),
11051            );
11052            a2 = vfmaq_f32(
11053                a2,
11054                w2,
11055                vmulq_f32(vld1q_f32(xp.add(j + 8)), vld1q_f32(cp.add(j + 8))),
11056            );
11057            a3 = vfmaq_f32(
11058                a3,
11059                w3,
11060                vmulq_f32(vld1q_f32(xp.add(j + 12)), vld1q_f32(cp.add(j + 12))),
11061            );
11062            j += 16;
11063        }
11064        let mut sum = vaddvq_f32(vaddq_f32(vaddq_f32(a0, a1), vaddq_f32(a2, a3)));
11065        while j < n {
11066            sum += (*wp.add(j)) as f32 * *xp.add(j) * *cp.add(j);
11067            j += 1;
11068        }
11069        sum
11070    }
11071}
11072
11073#[cfg(target_arch = "aarch64")]
11074#[target_feature(enable = "neon")]
11075unsafe fn dot_i8_f32_neon(w: &[u8], x: &[f32]) -> f32 {
11076    // SAFETY: callers uphold slice-length contracts (see call sites).
11077    unsafe {
11078        use core::arch::aarch64::*;
11079        let n = x.len();
11080        let wp = w.as_ptr() as *const i8;
11081        let xp = x.as_ptr();
11082        let (mut a0, mut a1, mut a2, mut a3) = (
11083            vdupq_n_f32(0.0),
11084            vdupq_n_f32(0.0),
11085            vdupq_n_f32(0.0),
11086            vdupq_n_f32(0.0),
11087        );
11088        let mut j = 0usize;
11089        while j + 16 <= n {
11090            let wb = vld1q_s8(wp.add(j));
11091            let lo = vmovl_s8(vget_low_s8(wb));
11092            let hi = vmovl_s8(vget_high_s8(wb));
11093            let w0 = vcvtq_f32_s32(vmovl_s16(vget_low_s16(lo)));
11094            let w1 = vcvtq_f32_s32(vmovl_s16(vget_high_s16(lo)));
11095            let w2 = vcvtq_f32_s32(vmovl_s16(vget_low_s16(hi)));
11096            let w3 = vcvtq_f32_s32(vmovl_s16(vget_high_s16(hi)));
11097            a0 = vfmaq_f32(a0, w0, vld1q_f32(xp.add(j)));
11098            a1 = vfmaq_f32(a1, w1, vld1q_f32(xp.add(j + 4)));
11099            a2 = vfmaq_f32(a2, w2, vld1q_f32(xp.add(j + 8)));
11100            a3 = vfmaq_f32(a3, w3, vld1q_f32(xp.add(j + 12)));
11101            j += 16;
11102        }
11103        let mut sum = vaddvq_f32(vaddq_f32(vaddq_f32(a0, a1), vaddq_f32(a2, a3)));
11104        while j < n {
11105            sum += (*wp.add(j)) as f32 * *xp.add(j);
11106            j += 1;
11107        }
11108        sum
11109    }
11110}
11111
11112#[allow(clippy::too_many_arguments)]
11113fn qmatvec(
11114    q: &[u8],
11115    rep: &[u8],
11116    row_scale: &[f32],
11117    x: &[f32],
11118    col_field: &[f32],
11119    dtype: TensorDtype,
11120    rows: usize,
11121    cols: usize,
11122    out: &mut [f32],
11123    pool: Option<&Pool>,
11124) {
11125    debug_assert_eq!(out.len(), rows);
11126    #[cfg(not(target_arch = "aarch64"))]
11127    let _ = rep;
11128
11129    #[cfg(target_arch = "aarch64")]
11130    if sdot_enabled() {
11131        let act = if dtype == TensorDtype::Q8_2f {
11132            split_act_q8_2f(x, col_field)
11133        } else {
11134            split_act(x)
11135        };
11136        let out_addr = SendMut(out.as_mut_ptr());
11137        let run_range = |start: usize, end: usize| {
11138            q8_range_sdot(q, rep, row_scale, &act, cols, out_addr, start, end)
11139        };
11140        match pool {
11141            Some(pool) if rows >= 256 => pool.run_rows(rows, &run_range),
11142            _ => run_range(0, rows),
11143        }
11144        return;
11145    }
11146    // x86 A8W8 via AVX2 maddubs — same quantized-activation contract as
11147    // the SDOT path (CMF_AVX2=0 keeps the exact i8×f32 loop).
11148    #[cfg(target_arch = "x86_64")]
11149    if avx2_a8w8_enabled() {
11150        let act = if dtype == TensorDtype::Q8_2f {
11151            split_act_q8_2f(x, col_field)
11152        } else {
11153            split_act(x)
11154        };
11155        let out_addr = SendMut(out.as_mut_ptr());
11156        let run_range = |start: usize, end: usize| {
11157            q8_range_avx2(q, row_scale, &act, cols, out_addr, start, end)
11158        };
11159        match pool {
11160            Some(pool) if rows >= 256 => pool.run_rows(rows, &run_range),
11161            _ => run_range(0, rows),
11162        }
11163        return;
11164    }
11165
11166    prescale_with(x, col_field, dtype, 1, |xs| {
11167        let out_addr = SendMut(out.as_mut_ptr());
11168        let run_range = move |start: usize, end: usize| {
11169            for o in start..end {
11170                let v = dot_i8_f32(&q[o * cols..(o + 1) * cols], xs) * row_scale[o];
11171                // SAFETY: disjoint row ranges per worker.
11172                unsafe { *out_addr.at(o) = v };
11173            }
11174        };
11175        match pool {
11176            Some(pool) if rows >= 256 => pool.run_rows(rows, &run_range),
11177            _ => run_range(0, rows),
11178        }
11179    });
11180}
11181
11182#[allow(clippy::too_many_arguments)]
11183fn qmatvec2(
11184    q: &[u8],
11185    row_scale: &[f32],
11186    x1: &[f32],
11187    x2: &[f32],
11188    col_field: &[f32],
11189    dtype: TensorDtype,
11190    rows: usize,
11191    cols: usize,
11192    o1: &mut [f32],
11193    o2: &mut [f32],
11194    pool: Option<&Pool>,
11195) {
11196    #[cfg(target_arch = "aarch64")]
11197    if sdot_enabled() {
11198        let a1s = if dtype == TensorDtype::Q8_2f {
11199            split_act_q8_2f(x1, col_field)
11200        } else {
11201            split_act(x1)
11202        };
11203        let a2s = if dtype == TensorDtype::Q8_2f {
11204            split_act_q8_2f(x2, col_field)
11205        } else {
11206            split_act(x2)
11207        };
11208        let p1 = SendMut(o1.as_mut_ptr());
11209        let p2 = SendMut(o2.as_mut_ptr());
11210        let run_range = |start: usize, end: usize| {
11211            q8_range2_sdot(q, row_scale, &a1s, &a2s, cols, p1, p2, start, end)
11212        };
11213        match pool {
11214            Some(pool) if rows >= 256 => pool.run_rows(rows, &run_range),
11215            _ => run_range(0, rows),
11216        }
11217        return;
11218    }
11219    #[cfg(target_arch = "x86_64")]
11220    if avx2_a8w8_enabled() {
11221        let a1s = if dtype == TensorDtype::Q8_2f {
11222            split_act_q8_2f(x1, col_field)
11223        } else {
11224            split_act(x1)
11225        };
11226        let a2s = if dtype == TensorDtype::Q8_2f {
11227            split_act_q8_2f(x2, col_field)
11228        } else {
11229            split_act(x2)
11230        };
11231        let p1 = SendMut(o1.as_mut_ptr());
11232        let p2 = SendMut(o2.as_mut_ptr());
11233        let run_range = |start: usize, end: usize| {
11234            q8_range2_avx2(q, row_scale, &a1s, &a2s, cols, p1, p2, start, end)
11235        };
11236        match pool {
11237            Some(pool) if rows >= 256 => pool.run_rows(rows, &run_range),
11238            _ => run_range(0, rows),
11239        }
11240        return;
11241    }
11242
11243    prescale_with(x1, col_field, dtype, 1, |x1s| {
11244        prescale_with(x2, col_field, dtype, 2, |x2s| {
11245            let p1 = SendMut(o1.as_mut_ptr());
11246            let p2 = SendMut(o2.as_mut_ptr());
11247            let run_range = move |start: usize, end: usize| {
11248                for o in start..end {
11249                    let row = &q[o * cols..(o + 1) * cols];
11250                    let s1 = dot_i8_f32(row, x1s) * row_scale[o];
11251                    let s2 = dot_i8_f32(row, x2s) * row_scale[o];
11252                    // SAFETY: disjoint row ranges per worker.
11253                    unsafe {
11254                        *p1.at(o) = s1;
11255                        *p2.at(o) = s2;
11256                    }
11257                }
11258            };
11259            match pool {
11260                Some(pool) if rows >= 256 => pool.run_rows(rows, &run_range),
11261                _ => run_range(0, rows),
11262            }
11263        });
11264    });
11265}
11266
11267#[derive(Clone, Copy)]
11268struct SendMut(*mut f32);
11269unsafe impl Send for SendMut {}
11270unsafe impl Sync for SendMut {}
11271
11272impl SendMut {
11273    #[inline]
11274    fn at(self, i: usize) -> *mut f32 {
11275        unsafe { self.0.add(i) }
11276    }
11277}
11278
11279#[cfg(test)]
11280mod tests {
11281    /// `q8_round` must be `round().clamp(±127) as i8` bit for bit: every
11282    /// half-integer, their neighbours one ulp either side, the clamp
11283    /// boundary, huge values, infinities and NaN, plus a dense sweep.
11284    #[test]
11285    fn q8_round_is_round_clamp() {
11286        let reference = |t: f32| t.round().clamp(-127.0, 127.0) as i8;
11287        let mut probe = vec![
11288            0.0f32,
11289            -0.0,
11290            f32::NAN,
11291            f32::INFINITY,
11292            f32::NEG_INFINITY,
11293            f32::MAX,
11294            f32::MIN,
11295            1e30,
11296            -1e30,
11297            f32::MIN_POSITIVE,
11298            -f32::MIN_POSITIVE,
11299        ];
11300        for k in -300i32..=300 {
11301            let h = k as f32 * 0.5;
11302            let up = f32::from_bits(h.to_bits() + 1);
11303            let down = f32::from_bits(h.to_bits().wrapping_sub(1));
11304            for t in [h, up, down] {
11305                probe.push(t);
11306                probe.push(-t);
11307            }
11308        }
11309        let mut t = -140.0f32;
11310        while t < 140.0 {
11311            probe.push(t);
11312            t += 0.000_731;
11313        }
11314        for t in probe {
11315            assert_eq!(q8_round(t), reference(t), "t = {t:e} ({:#x})", t.to_bits());
11316        }
11317    }
11318
11319    use super::*;
11320
11321    #[test]
11322    fn q2tp_i8_dot_matches_exact_on_grid() {
11323        // On-grid activations (±1 → sx=1/127, xq=±127 dequantizes
11324        // exactly, no outliers) must make the integer path agree with
11325        // the exact scalar walk to f32 rounding.
11326        let (rows, cols) = (5, 64);
11327        let gpr = cols / GROUP_SIZE;
11328        // Synthetic codes plane + a flat ladder: scales_into is not under
11329        // test here, so drive dot_q2tp_row_i8 / q2tp_row_exact directly
11330        // with hand-made scales.
11331        let chunks: Vec<u8> = (0..rows * gpr * Q2TP_CHUNK)
11332            .map(|i| (i as u32).wrapping_mul(2654435761) as u8)
11333            .collect();
11334        let scales: Vec<f32> = (0..gpr).map(|g| 0.5 + g as f32 * 0.25).collect();
11335        let x: Vec<f32> = (0..cols)
11336            .map(|i| if i % 3 == 0 { -1.0 } else { 1.0 })
11337            .collect();
11338        let act = split_act(&x);
11339        assert!(
11340            act.outliers.is_empty(),
11341            "on-grid input must have no outliers"
11342        );
11343        let gsum = q1_group_sums(&act.xq, gpr);
11344        for r in 0..rows {
11345            let exact = q2tp_row_exact(&chunks, r, gpr, &x, &scales);
11346            let fast = dot_q2tp_row_i8(&chunks, r, gpr, &act.xq, &gsum, &scales) * act.sx;
11347            assert!(
11348                (exact - fast).abs() <= exact.abs() * 1e-5 + 1e-5,
11349                "row {r}: exact {exact} vs i8 {fast}"
11350            );
11351        }
11352    }
11353
11354    #[test]
11355    fn q2tp_affine_fuses_half_scale_correction_without_changing_raw_decode() {
11356        let (rows, cols) = (1usize, GROUP_SIZE);
11357        let mut bytes = vec![0u8; Q2TP_CHUNK + 4 + 1];
11358        // Repeating symbols 0,1,2,0 at unit scale.  q2tp's raw B is
11359        // (c-1.5), while the affine Prism operator is (c-1.0).
11360        bytes[..Q2TP_CHUNK].fill(0x24); // codes 0,1,2,0 in LSB-first order
11361        bytes[Q2TP_CHUNK..Q2TP_CHUNK + 2].copy_from_slice(&0u16.to_le_bytes());
11362        bytes[Q2TP_CHUNK + 2..Q2TP_CHUNK + 4].copy_from_slice(&0u16.to_le_bytes());
11363        bytes[Q2TP_CHUNK + 4] = 1; // dtype16 rung 1 = 1.0
11364        let x = vec![1.0f32; cols];
11365        let mut raw = vec![0.0f32; rows];
11366        let mut affine = vec![0.0f32; rows];
11367        q2tp_matvec_for_test(&bytes, &x, rows, cols, &mut raw);
11368        q2tp_affine_matvec_for_test(&bytes, &x, rows, cols, &mut affine);
11369        assert_eq!(raw, vec![-24.0]);
11370        assert_eq!(affine, vec![-8.0]);
11371        assert!((affine[0] - (raw[0] + 0.5 * cols as f32)).abs() < 1e-6);
11372    }
11373
11374    #[cfg(target_arch = "x86_64")]
11375    #[test]
11376    fn q2tp_avx2_dot_matches_scalar_for_random_patterns() {
11377        // Compare the release AVX2 integer dot against the scalar oracle over
11378        // arbitrary packed bytes/activation signs.  This guards the exact
11379        // table-load path used after rejecting a faster-looking decoder whose
11380        // full-checkpoint greedy output drifted.
11381        if !std::arch::is_x86_feature_detected!("avx2") {
11382            return;
11383        }
11384        let mut seed = 0x9e3779b9u32;
11385        let mut next = || {
11386            seed = seed.wrapping_mul(1664525).wrapping_add(1013904223);
11387            seed
11388        };
11389        for _ in 0..20_000 {
11390            let mut ch = [0u8; Q2TP_CHUNK];
11391            let mut x = [0i8; GROUP_SIZE];
11392            for b in &mut ch {
11393                *b = next() as u8;
11394            }
11395            for v in &mut x {
11396                *v = (next() >> 24) as i8;
11397            }
11398            let mut reference = 0i32;
11399            for (k, &b) in ch.iter().enumerate() {
11400                reference += (b & 3) as i32 * x[k * 4] as i32;
11401                reference += ((b >> 2) & 3) as i32 * x[k * 4 + 1] as i32;
11402                reference += ((b >> 4) & 3) as i32 * x[k * 4 + 2] as i32;
11403                reference += ((b >> 6) & 3) as i32 * x[k * 4 + 3] as i32;
11404            }
11405            // SAFETY: guarded by the runtime AVX2 feature check and fixed
11406            // 8-byte/32-byte slice lengths above.
11407            let got = unsafe { q2tp_code_dot_avx2(&ch, &x) };
11408            assert_eq!(got, reference, "packed q2 lane mismatch");
11409        }
11410    }
11411
11412    #[test]
11413    fn q8_row_dot_fast_matches_scalar() {
11414        // The per-arch fast dot must agree with the exact scalar oracle
11415        // (same contract the fused q8 FFN arm rides on).
11416        let cols = 96;
11417        let row: Vec<u8> = (0..cols)
11418            .map(|i| ((i * 37 % 251) - 125) as i8 as u8)
11419            .collect();
11420        let x: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.13).sin()).collect();
11421        let act = split_act(&x);
11422        let fast = q8_row_dot(&row, &act);
11423        let scalar = q8_row_dot_scalar(&row, &act);
11424        assert!(
11425            (fast - scalar).abs() <= scalar.abs() * 1e-5 + 1e-5,
11426            "fast {fast} vs scalar {scalar}"
11427        );
11428    }
11429
11430    #[test]
11431    fn f32_matvec_matches_matvec_rows_bitexact() {
11432        let (rows, cols) = (300, 40);
11433        let w: Vec<f32> = (0..rows * cols).map(|i| (i as f32 * 0.017).sin()).collect();
11434        let x: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.05).cos()).collect();
11435        let qt = QTensor::from_f32(w.clone(), rows, cols);
11436
11437        let mut a = vec![0.0f32; rows];
11438        matvec_rows(None, &w, &x, &mut a);
11439        let mut b = vec![0.0f32; rows];
11440        qt.matvec(&x, &mut b, None);
11441        assert_eq!(a, b);
11442    }
11443
11444    #[test]
11445    fn sdot_kernel_exact_on_grid() {
11446        // Activations already on the i8 grid (±1 with amax=1 → sx=1/127,
11447        // xq=±127 dequantizes EXACTLY) → the SDOT path must match the
11448        // exact f32 dot to float rounding. This isolates kernel
11449        // correctness from quantization noise.
11450        eprintln!("sdot_enabled = {}", sdot_enabled());
11451        let (rows, cols) = (9, 80); // odd rows → exercises 4-row + tail
11452        let w: Vec<u8> = (0..rows * cols)
11453            .map(|i| (((i * 37) % 251) as i32 - 125) as i8 as u8)
11454            .collect();
11455        let scales: Vec<f32> = (0..rows).map(|o| 0.005 + o as f32 * 0.001).collect();
11456        let x: Vec<f32> = (0..cols)
11457            .map(|i| match i % 3 {
11458                0 => 1.0,
11459                1 => -1.0,
11460                _ => 0.0,
11461            })
11462            .collect();
11463        let mut a = vec![0.0f32; rows];
11464        qmatvec(
11465            &w,
11466            &[],
11467            &scales,
11468            &x,
11469            &[],
11470            TensorDtype::Q8Row,
11471            rows,
11472            cols,
11473            &mut a,
11474            None,
11475        );
11476        for o in 0..rows {
11477            let mut acc = 0.0f32;
11478            for j in 0..cols {
11479                acc += (w[o * cols + j] as i8) as f32 * x[j];
11480            }
11481            let expect = acc * scales[o];
11482            assert!(
11483                (a[o] - expect).abs() < 1e-3 * expect.abs().max(1e-3),
11484                "row {o}: {} vs {expect}",
11485                a[o]
11486            );
11487        }
11488    }
11489
11490    #[test]
11491    fn q1_tbl_fast_path_matches_reference() {
11492        // gpr = 8 exercises the TBL pair-load fast loop, and the LAST
11493        // row's final 4-tile window trips the 4B-overread guard (the
11494        // payload ends exactly at the last tile) — both paths must
11495        // agree with the dequant reference.
11496        let (rows, cols) = (5, 256);
11497        let gpr = cols / GROUP_SIZE;
11498        let mut bytes = Vec::new();
11499        for t in 0..rows * gpr {
11500            let s = 0.007 + (t % 11) as f32 * 0.004;
11501            bytes.extend_from_slice(&cortiq_core::quant::f32_to_f16(s).to_le_bytes());
11502            for j in 0..4 {
11503                bytes.push(((t * 53 + j * 89 + 7) % 249) as u8);
11504            }
11505        }
11506        let x: Vec<f32> = (0..cols)
11507            .map(|i| if (i * 5) % 7 < 3 { 1.0 } else { -1.0 })
11508            .collect();
11509        let mut w = vec![0.0f32; rows * cols];
11510        cortiq_core::quant::dequant_q1(&bytes, &mut w);
11511        let mut got = vec![0.0f32; rows];
11512        q1_matvec(&bytes, &x, rows, cols, &mut got, None);
11513        for o in 0..rows {
11514            let expect: f32 = (0..cols).map(|j| w[o * cols + j] * x[j]).sum();
11515            assert!(
11516                (got[o] - expect).abs() < 1e-3 * expect.abs().max(1e-3),
11517                "row {o}: {} vs {expect}",
11518                got[o]
11519            );
11520        }
11521        // Blocked 1×4 batch (b=5: one quad + remainder) must equal the
11522        // single-matvec path bit-for-bit.
11523        let b = 5usize;
11524        let mut xs_all = Vec::new();
11525        for bi in 0..b {
11526            xs_all.extend(x.iter().map(|v| if bi % 2 == 0 { *v } else { -*v }));
11527        }
11528        let mut mm = vec![0.0f32; b * rows];
11529        q1_matmat(&bytes, &xs_all, b, rows, cols, &mut mm, None);
11530        for bi in 0..b {
11531            let mut single = vec![0.0f32; rows];
11532            q1_matvec(
11533                &bytes,
11534                &xs_all[bi * cols..(bi + 1) * cols],
11535                rows,
11536                cols,
11537                &mut single,
11538                None,
11539            );
11540            assert_eq!(&mm[bi * rows..(bi + 1) * rows], &single[..], "stream {bi}");
11541        }
11542    }
11543
11544    #[test]
11545    fn q1_kernels_match_exact_reference() {
11546        // Synthetic q1 payload: 6-byte tiles [f16 scale][4B bits].
11547        let (rows, cols) = (7, 96);
11548        let gpr = cols / GROUP_SIZE;
11549        let mut bytes = Vec::new();
11550        for t in 0..rows * gpr {
11551            let s = 0.01 + (t % 13) as f32 * 0.003;
11552            bytes.extend_from_slice(&cortiq_core::quant::f32_to_f16(s).to_le_bytes());
11553            for j in 0..4 {
11554                bytes.push(((t * 31 + j * 97) % 251) as u8);
11555            }
11556        }
11557        // On-grid activations (±1, amax 1) → the SDOT path is exact.
11558        let x: Vec<f32> = (0..cols)
11559            .map(|i| if i % 3 == 0 { 1.0 } else { -1.0 })
11560            .collect();
11561        // Reference through the core dequant.
11562        let mut w = vec![0.0f32; rows * cols];
11563        cortiq_core::quant::dequant_q1(&bytes, &mut w);
11564        let mut expect = vec![0.0f32; rows];
11565        for o in 0..rows {
11566            expect[o] = (0..cols).map(|j| w[o * cols + j] * x[j]).sum();
11567        }
11568        let mut got = vec![0.0f32; rows];
11569        q1_matvec(&bytes, &x, rows, cols, &mut got, None);
11570        for o in 0..rows {
11571            assert!(
11572                (got[o] - expect[o]).abs() < 1e-3 * expect[o].abs().max(1e-3),
11573                "row {o}: {} vs {}",
11574                got[o],
11575                expect[o]
11576            );
11577        }
11578        // Pair and batch paths agree with the single path.
11579        let x2: Vec<f32> = x.iter().map(|v| -v).collect();
11580        let (mut a1, mut a2) = (vec![0.0f32; rows], vec![0.0f32; rows]);
11581        q1_matvec2(&bytes, &x, &x2, rows, cols, &mut a1, &mut a2, None);
11582        assert_eq!(a1, got);
11583        let mut xs = x.clone();
11584        xs.extend_from_slice(&x2);
11585        let mut mm = vec![0.0f32; 2 * rows];
11586        q1_matmat(&bytes, &xs, 2, rows, cols, &mut mm, None);
11587        assert_eq!(&mm[..rows], got.as_slice());
11588        assert_eq!(&mm[rows..], a2.as_slice());
11589    }
11590
11591    #[test]
11592    fn repack_is_bit_identical() {
11593        // The interleaved-repack kernel must produce EXACTLY the same
11594        // bits as the mmap-layout kernel: integer accumulation is order-
11595        // exact, the f32 epilogue is identical. Odd rows exercise the
11596        // tail; direct range calls exercise unaligned pool splits.
11597        let (rows, cols) = (267, 96); // 66 groups + 3 tail rows, cols % 16 == 0
11598        let w: Vec<u8> = (0..rows * cols)
11599            .map(|i| (((i * 89) % 253) as i32 - 126) as i8 as u8)
11600            .collect();
11601        let scales: Vec<f32> = (0..rows).map(|o| 0.003 + o as f32 * 0.0007).collect();
11602        let x: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.37).sin() * 2.0).collect();
11603        let rep = q8_repack_layout(&w, rows, cols);
11604        // Group interleave round-trips.
11605        for g in 0..rows / 4 {
11606            for c in 0..cols / 16 {
11607                for lane in 0..4 {
11608                    assert_eq!(
11609                        &rep[g * 4 * cols + c * 64 + lane * 16
11610                            ..g * 4 * cols + c * 64 + lane * 16 + 16],
11611                        &w[(g * 4 + lane) * cols + c * 16..(g * 4 + lane) * cols + c * 16 + 16],
11612                    );
11613                }
11614            }
11615        }
11616        let mut a = vec![0.0f32; rows];
11617        qmatvec(
11618            &w,
11619            &[],
11620            &scales,
11621            &x,
11622            &[],
11623            TensorDtype::Q8Row,
11624            rows,
11625            cols,
11626            &mut a,
11627            None,
11628        );
11629        let mut b = vec![0.0f32; rows];
11630        qmatvec(
11631            &w,
11632            &rep,
11633            &scales,
11634            &x,
11635            &[],
11636            TensorDtype::Q8Row,
11637            rows,
11638            cols,
11639            &mut b,
11640            None,
11641        );
11642        assert_eq!(a, b, "full-range repack output diverged");
11643
11644        #[cfg(target_arch = "aarch64")]
11645        if sdot_enabled() {
11646            // Unaligned range split (pool workers get arbitrary bounds).
11647            let act = split_act(&x);
11648            let mut c1 = vec![0.0f32; rows];
11649            let mut c2 = vec![0.0f32; rows];
11650            q8_range_sdot(
11651                &w,
11652                &[],
11653                &scales,
11654                &act,
11655                cols,
11656                SendMut(c1.as_mut_ptr()),
11657                3,
11658                rows - 2,
11659            );
11660            q8_range_sdot(
11661                &w,
11662                &rep,
11663                &scales,
11664                &act,
11665                cols,
11666                SendMut(c2.as_mut_ptr()),
11667                3,
11668                rows - 2,
11669            );
11670            assert_eq!(c1, c2, "unaligned-range repack output diverged");
11671        }
11672    }
11673
11674    #[test]
11675    fn sdot_a8w8_noise_is_bounded() {
11676        // Off-grid activations: A8 quantization noise must stay small in
11677        // relative L2 over the whole output (realistic accuracy contract;
11678        // vmfcore measured argmax-identical decode on real models).
11679        let (rows, cols) = (16, 512);
11680        let w: Vec<u8> = (0..rows * cols)
11681            .map(|i| (((i * 37) % 251) as i32 - 125) as i8 as u8)
11682            .collect();
11683        let scales = vec![0.01f32; rows];
11684        let x: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.21).sin()).collect();
11685        let mut a = vec![0.0f32; rows];
11686        qmatvec(
11687            &w,
11688            &[],
11689            &scales,
11690            &x,
11691            &[],
11692            TensorDtype::Q8Row,
11693            rows,
11694            cols,
11695            &mut a,
11696            None,
11697        );
11698        let (mut num, mut den) = (0f64, 0f64);
11699        for o in 0..rows {
11700            let mut acc = 0.0f32;
11701            for j in 0..cols {
11702                acc += (w[o * cols + j] as i8) as f32 * x[j];
11703            }
11704            let expect = acc * scales[o];
11705            num += ((a[o] - expect) as f64).powi(2);
11706            den += (expect as f64).powi(2);
11707        }
11708        let rel = (num / den.max(1e-12)).sqrt();
11709        assert!(rel < 0.05, "A8W8 relative L2 error too high: {rel}");
11710    }
11711
11712    #[test]
11713    fn i8_dot_neon_matches_scalar() {
11714        let n = 100;
11715        let w: Vec<u8> = (0..n).map(|i| ((i * 37 + 11) % 251) as u8).collect();
11716        let x: Vec<f32> = (0..n).map(|i| (i as f32 * 0.13).sin()).collect();
11717        let mut scalar = 0.0f32;
11718        for j in 0..n {
11719            scalar += (w[j] as i8) as f32 * x[j];
11720        }
11721        let fast = dot_i8_f32(&w, &x);
11722        assert!((scalar - fast).abs() < 1e-3 * scalar.abs().max(1.0));
11723    }
11724
11725    /// Fused vbit matvec must match full dequant_vbit + dense matvec.
11726    #[test]
11727    fn vbitmatvec_matches_full_dequant() {
11728        let (rows, cols) = (6, 64);
11729        let ng = cols / GROUP_SIZE;
11730        // Hand-craft: bits per row, f16 scales, packed rows.
11731        let bits: Vec<u8> = vec![3, 4, 5, 6, 8, 4];
11732        let mut bytes = bits.clone();
11733        for g in 0..rows * ng {
11734            let s = 0.02 + 0.001 * g as f32;
11735            bytes.extend_from_slice(&cortiq_core::quant::f32_to_f16(s).to_le_bytes());
11736        }
11737        for r in 0..rows {
11738            let b = bits[r] as usize;
11739            let (mut acc, mut nb) = (0u64, 0usize);
11740            let mut rowbytes = Vec::new();
11741            for i in 0..cols {
11742                let v = ((i * 7 + r * 13) % (1 << b)) as u64;
11743                acc = (acc << b) | v;
11744                nb += b;
11745                while nb >= 8 {
11746                    nb -= 8;
11747                    rowbytes.push(((acc >> nb) & 0xFF) as u8);
11748                }
11749            }
11750            if nb > 0 {
11751                rowbytes.push(((acc << (8 - nb)) & 0xFF) as u8);
11752            }
11753            bytes.extend_from_slice(&rowbytes);
11754        }
11755        let x: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.19).sin()).collect();
11756
11757        let mut reference = vec![0f32; rows * cols];
11758        cortiq_core::quant::dequant_vbit(&bytes, rows, cols, &mut reference).unwrap();
11759        let mut expect = vec![0f32; rows];
11760        for r in 0..rows {
11761            expect[r] = reference[r * cols..(r + 1) * cols]
11762                .iter()
11763                .zip(&x)
11764                .map(|(w, xv)| w * xv)
11765                .sum();
11766        }
11767        let mut got = vec![0f32; rows];
11768        let offsets = vbit_row_offsets(&bytes, rows, cols);
11769        vbitmatvec(&bytes, &offsets, &x, rows, cols, &mut got, None);
11770        // SDOT path quantizes activations to i8 (A8W8): bounded noise,
11771        // same contract as q8 (exact path is pinned by CMF_SDOT=0 in
11772        // the golden-parity gate).
11773        let tol = if a8w8_enabled() { 6e-2 } else { 1e-4 };
11774        let scale = expect.iter().fold(0f32, |m, v| m.max(v.abs())).max(1e-6);
11775        for r in 0..rows {
11776            assert!(
11777                (got[r] - expect[r]).abs() < tol * scale,
11778                "row {r}: {} vs {}",
11779                got[r],
11780                expect[r]
11781            );
11782        }
11783    }
11784
11785    /// Fused q4 matvec must match the reference full-dequant + dense
11786    /// matvec bit-for-bit in structure (same f32 math, group order).
11787    /// vbit matmat: the blocked 1×4 leg must match the per-row path
11788    /// (paired env toggle; larger shape so both code paths engage).
11789    #[test]
11790    #[cfg(target_arch = "x86_64")]
11791    fn vbit_matmat_blocked_matches_per_row() {
11792        let (rows, cols, b) = (64usize, 128usize, 9usize);
11793        let ng = cols / GROUP_SIZE;
11794        let bits: Vec<u8> = (0..rows).map(|r| [3u8, 4, 5, 6][r % 4]).collect();
11795        let mut bytes = bits.clone();
11796        for g in 0..rows * ng {
11797            let sc = 0.02 + 0.0005 * g as f32;
11798            bytes.extend_from_slice(&cortiq_core::quant::f32_to_f16(sc).to_le_bytes());
11799        }
11800        for r in 0..rows {
11801            let bw = bits[r] as usize;
11802            let (mut acc, mut nb) = (0u64, 0usize);
11803            let mut rowbytes = Vec::new();
11804            for i in 0..cols {
11805                let v = ((i * 7 + r * 13) % (1 << bw)) as u64;
11806                acc = (acc << bw) | v;
11807                nb += bw;
11808                while nb >= 8 {
11809                    nb -= 8;
11810                    rowbytes.push(((acc >> nb) & 0xFF) as u8);
11811                }
11812            }
11813            if nb > 0 {
11814                rowbytes.push(((acc << (8 - nb)) & 0xFF) as u8);
11815            }
11816            bytes.extend_from_slice(&rowbytes);
11817        }
11818        let x: Vec<f32> = (0..b * cols)
11819            .map(|i| ((i * 13 + 7) % 97) as f32 / 97.0 - 0.5)
11820            .collect();
11821        let offsets = vbit_row_offsets(&bytes, rows, cols);
11822        let mut y_a = vec![0f32; b * rows];
11823        let mut y_b = vec![0f32; b * rows];
11824        unsafe { std::env::set_var("CMF_X86_BLOCKED", "1") };
11825        vbitmatmat(&bytes, &offsets, &x, b, rows, cols, &mut y_a, None);
11826        unsafe { std::env::set_var("CMF_X86_BLOCKED", "0") };
11827        vbitmatmat(&bytes, &offsets, &x, b, rows, cols, &mut y_b, None);
11828        unsafe { std::env::remove_var("CMF_X86_BLOCKED") };
11829        let max_d = y_a
11830            .iter()
11831            .zip(&y_b)
11832            .map(|(p, q)| (p - q).abs())
11833            .fold(0.0f32, f32::max);
11834        assert!(max_d < 1e-4, "vbit blocked ≠ per-row: max|Δ| = {max_d}");
11835    }
11836
11837    /// q4t blocked 1×4 (SDOT on ARM, AVX2 on x86) must equal the
11838    /// per-row path exactly: same nibble unpack, same group order,
11839    /// same f32 accumulation — batch == matvec bit-for-bit. b=9 covers
11840    /// two full 1×4 blocks plus a remainder through the single-row
11841    /// kernel. (Both paths produce identical output, so the shared
11842    /// CMF_X86_BLOCKED env var racing with other tests cannot flip
11843    /// the verdict — worst case both sides take the same path.)
11844    #[test]
11845    fn q4t_matmat_blocked_matches_per_row() {
11846        let (rows, cols, b) = (16usize, 64usize, 9usize);
11847        let gpr = cols / GROUP_SIZE;
11848        let mut bytes = vec![0u8; rows * gpr * Q4_TILE];
11849        for r in 0..rows {
11850            for g in 0..gpr {
11851                let t = (r * gpr + g) * Q4_TILE;
11852                let sc = 0.02 + 0.001 * (r * gpr + g) as f32;
11853                bytes[t..t + 2].copy_from_slice(&cortiq_core::quant::f32_to_f16(sc).to_le_bytes());
11854                for k in 0..16 {
11855                    bytes[t + 2 + k] = ((r * 31 + g * 7 + k * 13) % 251) as u8;
11856                }
11857            }
11858        }
11859        let x: Vec<f32> = (0..b * cols)
11860            .map(|i| ((i * 13 + 7) % 97) as f32 / 97.0 - 0.5)
11861            .collect();
11862        let mut y_blk = vec![0f32; b * rows];
11863        let mut y_row = vec![0f32; b * rows];
11864        unsafe { std::env::set_var("CMF_X86_BLOCKED", "1") };
11865        q4t_matmat(&bytes, &x, b, rows, cols, &mut y_blk, None);
11866        unsafe { std::env::set_var("CMF_X86_BLOCKED", "0") };
11867        q4t_matmat(&bytes, &x, b, rows, cols, &mut y_row, None);
11868        unsafe { std::env::remove_var("CMF_X86_BLOCKED") };
11869        assert_eq!(y_blk, y_row, "q4t blocked 1x4 ≠ per-row");
11870    }
11871
11872    /// The wide-batch Accelerate arm of q4t_matmat vs a brute-force
11873    /// f32 dequant matmul: both are f32 GEMMs, so only reduction
11874    /// order differs — tight tolerance.
11875    /// A synthetic q4tp payload: random nibbles plus a per-row ladder whose
11876    /// span varies row to row, so the codes actually exercise the full 0..31
11877    /// range rather than clustering on one rung.
11878    fn synth_q4tp(rows: usize, cols: usize) -> Vec<u8> {
11879        use cortiq_core::quant::{f32_to_f16, q4tp_code_stride, q4tp_put_code};
11880        let gpr = cols / GROUP_SIZE;
11881        let stride = q4tp_code_stride(gpr);
11882        let (params_off, codes_off, _) = q4tp_sections(rows, cols);
11883        let mut b = vec![0u8; codes_off + rows * stride];
11884        for r in 0..rows {
11885            for g in 0..gpr {
11886                let t = (r * gpr + g) * Q4TP_NIB;
11887                for k in 0..16 {
11888                    b[t + k] = ((r * 31 + g * 7 + k * 13) % 251) as u8;
11889                }
11890            }
11891            let lo = -6.0 - 0.03 * (r % 17) as f32;
11892            let step = 0.01 + 0.004 * (r % 11) as f32;
11893            let p = params_off + r * 4;
11894            b[p..p + 2].copy_from_slice(&f32_to_f16(lo).to_le_bytes());
11895            b[p + 2..p + 4].copy_from_slice(&f32_to_f16(step).to_le_bytes());
11896            let crow = &mut b[codes_off + r * stride..codes_off + (r + 1) * stride];
11897            for g in 0..gpr {
11898                q4tp_put_code(crow, g, (r * 5 + g * 3) % 32);
11899            }
11900        }
11901        b
11902    }
11903
11904    /// The same weights re-expressed as q4_tiled, so the proven kernel can
11905    /// be the reference: each tile stores the ladder scale its code selects.
11906    /// Only the f16 rounding of that scale separates the two payloads.
11907    fn q4tp_as_q4t(bytes: &[u8], rows: usize, cols: usize) -> Vec<u8> {
11908        let gpr = cols / GROUP_SIZE;
11909        let v = Q4tpView::new(bytes, rows, cols);
11910        let mut out = vec![0u8; rows * gpr * Q4_TILE];
11911        let mut sc = vec![0f32; gpr];
11912        for r in 0..rows {
11913            v.scales_into(r, gpr, &mut sc);
11914            for g in 0..gpr {
11915                let t = (r * gpr + g) * Q4_TILE;
11916                let s = sc[g];
11917                out[t..t + 2].copy_from_slice(&cortiq_core::quant::f32_to_f16(s).to_le_bytes());
11918                let src = (r * gpr + g) * Q4TP_NIB;
11919                out[t + 2..t + Q4_TILE].copy_from_slice(&v.nib[src..src + Q4TP_NIB]);
11920            }
11921        }
11922        out
11923    }
11924
11925    /// The exact (`CMF_SDOT=0`) path must reproduce `dequant_q4tp` to f32
11926    /// rounding — that scalar routine is the format's definition, and the
11927    /// kernels re-derive the scale from the ladder independently. Call the
11928    /// row kernel directly: `matmat` picks the int8 arm when a8w8 is on,
11929    /// so routing through it would test the other path by accident.
11930    #[test]
11931    fn q4tp_exact_path_matches_dequant_reference() {
11932        let (rows, cols) = (256usize, 512usize);
11933        let gpr = cols / GROUP_SIZE;
11934        let bytes = synth_q4tp(rows, cols);
11935        let mut w = vec![0f32; rows * cols];
11936        cortiq_core::quant::dequant_q4tp(&bytes, rows, cols, &mut w);
11937
11938        let x: Vec<f32> = (0..cols)
11939            .map(|i| ((i * 13 + 7) % 97) as f32 / 97.0 - 0.5)
11940            .collect();
11941        let v = Q4tpView::new(&bytes, rows, cols);
11942        let mut sc = vec![0f32; gpr];
11943        for r in 0..rows {
11944            v.scales_into(r, gpr, &mut sc);
11945            let got = q4tp_row_exact(v.nib, r, gpr, &x, &sc);
11946            let want: f32 = (0..cols).map(|c| w[r * cols + c] * x[c]).sum();
11947            // These dot products cancel down to ~1e-3 from terms of ~5e-2, so
11948            // the meaningful yardstick is the summed magnitude, not the result:
11949            // against the result any reordering of a 512-term f32 sum "fails".
11950            let mag: f32 = (0..cols).map(|c| (w[r * cols + c] * x[c]).abs()).sum();
11951            assert!(
11952                (got - want).abs() <= 1e-5 * mag,
11953                "row {r}: kernel {got} vs dequant {want}"
11954            );
11955        }
11956    }
11957
11958    /// The int8 (a8w8) path can't be checked against an f32 reference — the
11959    /// activation quantization dominates. Check it against the q4t kernel it
11960    /// was ported from instead, on payloads holding the same weights: that
11961    /// isolates exactly what the port could break (16 B stride, ladder
11962    /// lookup, nibble unpack) from what it deliberately shares.
11963    #[test]
11964    fn q4tp_matvec_matches_the_q4t_kernel_it_was_ported_from() {
11965        let (rows, cols) = (256usize, 512usize);
11966        let bytes = synth_q4tp(rows, cols);
11967        let twin = q4tp_as_q4t(&bytes, rows, cols);
11968        let x: Vec<f32> = (0..cols)
11969            .map(|i| ((i * 13 + 7) % 97) as f32 / 97.0 - 0.5)
11970            .collect();
11971
11972        let mut got = vec![0f32; rows];
11973        q4tp_matvec(&bytes, &x, rows, cols, &mut got, None);
11974        let mut want = vec![0f32; rows];
11975        q4t_matvec(&twin, &x, rows, cols, &mut want, None);
11976
11977        // Scale is f16 in the twin and f32 here, so allow that rounding on
11978        // top of the summed magnitude (same cancellation argument as above).
11979        let mut w = vec![0f32; rows * cols];
11980        cortiq_core::quant::dequant_q4tp(&bytes, rows, cols, &mut w);
11981        for r in 0..rows {
11982            let mag: f32 = (0..cols).map(|c| (w[r * cols + c] * x[c]).abs()).sum();
11983            assert!(
11984                (got[r] - want[r]).abs() <= 1e-3 * mag,
11985                "row {r}: q4tp {} vs q4t {}",
11986                got[r],
11987                want[r]
11988            );
11989        }
11990    }
11991
11992    /// `matmat` carries three arms (Accelerate, blocked int8 1x4, scalar).
11993    /// Batch 5 crosses the blocked kernel's stride, so this exercises the
11994    /// 1x4 path AND its scalar tail in one run — the blocked kernel is new
11995    /// code and its four accumulators are exactly what tends to go wrong.
11996    #[test]
11997    fn q4tp_matmat_matches_the_q4t_kernel_it_was_ported_from() {
11998        let (rows, cols, b) = (256usize, 512usize, 5usize);
11999        let bytes = synth_q4tp(rows, cols);
12000        let twin = q4tp_as_q4t(&bytes, rows, cols);
12001        let xs: Vec<f32> = (0..b * cols)
12002            .map(|i| ((i * 29 + 11) % 89) as f32 / 89.0 - 0.5)
12003            .collect();
12004
12005        let mut got = vec![0f32; b * rows];
12006        q4tp_matmat(&bytes, &xs, b, rows, cols, &mut got, None);
12007        let mut want = vec![0f32; b * rows];
12008        q4t_matmat(&twin, &xs, b, rows, cols, &mut want, None);
12009
12010        let mut w = vec![0f32; rows * cols];
12011        cortiq_core::quant::dequant_q4tp(&bytes, rows, cols, &mut w);
12012        for t in 0..b {
12013            for r in 0..rows {
12014                let mag: f32 = (0..cols)
12015                    .map(|c| (w[r * cols + c] * xs[t * cols + c]).abs())
12016                    .sum();
12017                let (g, wa) = (got[t * rows + r], want[t * rows + r]);
12018                assert!(
12019                    (g - wa).abs() <= 1e-3 * mag,
12020                    "batch {t} row {r}: q4tp {g} vs q4t {wa}"
12021                );
12022            }
12023        }
12024    }
12025
12026    #[test]
12027    fn q4tp_matvec2_matches_the_single_stream_kernel() {
12028        let (rows, cols) = (128usize, 256usize);
12029        let gpr = cols / GROUP_SIZE;
12030        let bytes = synth_q4tp(rows, cols);
12031        let xs: Vec<f32> = (0..2 * cols)
12032            .map(|i| ((i * 29 + 11) % 89) as f32 / 89.0 - 0.5)
12033            .collect();
12034
12035        let (mut o1, mut o2) = (vec![0f32; rows], vec![0f32; rows]);
12036        q4tp_matvec2(
12037            &bytes,
12038            &xs[..cols],
12039            &xs[cols..],
12040            rows,
12041            cols,
12042            &mut o1,
12043            &mut o2,
12044            None,
12045        );
12046
12047        // matvec2 takes the exact path for both streams, so the single-row
12048        // kernel is an exact reference — no tolerance for path differences.
12049        let v = Q4tpView::new(&bytes, rows, cols);
12050        let mut sc = vec![0f32; gpr];
12051        for r in 0..rows {
12052            v.scales_into(r, gpr, &mut sc);
12053            assert_eq!(o1[r], q4tp_row_exact(v.nib, r, gpr, &xs[..cols], &sc));
12054            assert_eq!(o2[r], q4tp_row_exact(v.nib, r, gpr, &xs[cols..], &sc));
12055        }
12056    }
12057
12058    /// q4tp must not COST speed — it exists to save bytes, and a format that
12059    /// trades 7% of a file for a slower model is a bad trade. This guard is
12060    /// here because correctness tests happily passed while `q4tp_matmat` was
12061    /// missing its int8 and Accelerate arms and the model ran 5x slower.
12062    /// Measured on M-series: 0.97-1.04x, i.e. parity (16 B tiles are better
12063    /// aligned than q4t's 18 B, which pays for the scale indirection).
12064    #[test]
12065    fn q4tp_matvec_keeps_pace_with_q4t() {
12066        let (rows, cols) = (4096usize, 3072usize);
12067        let bytes = synth_q4tp(rows, cols);
12068        let twin = q4tp_as_q4t(&bytes, rows, cols);
12069        let x: Vec<f32> = (0..cols).map(|i| (i % 97) as f32 / 97.0 - 0.5).collect();
12070        let mut o = vec![0f32; rows];
12071        let n = 12;
12072        let mut best = (f64::MAX, f64::MAX);
12073        // Interleaved A/B, minimum statistic: this machine throttles, and a
12074        // mean over a thermal ramp reliably indicts whichever ran second.
12075        for _ in 0..3 {
12076            let t0 = std::time::Instant::now();
12077            for _ in 0..n {
12078                q4t_matvec(&twin, &x, rows, cols, &mut o, None);
12079            }
12080            best.0 = best.0.min(t0.elapsed().as_secs_f64());
12081            let t0 = std::time::Instant::now();
12082            for _ in 0..n {
12083                q4tp_matvec(&bytes, &x, rows, cols, &mut o, None);
12084            }
12085            best.1 = best.1.min(t0.elapsed().as_secs_f64());
12086        }
12087        let ratio = best.1 / best.0;
12088        println!(
12089            "q4t {:.3} ms | q4tp {:.3} ms | {ratio:.2}x",
12090            best.0 * 1e3 / n as f64,
12091            best.1 * 1e3 / n as f64
12092        );
12093        assert!(ratio < 2.0, "q4tp matvec {ratio:.2}x slower than q4t");
12094    }
12095
12096    #[cfg(target_os = "macos")]
12097    #[test]
12098    fn q4t_matmat_accel_matches_dequant_reference() {
12099        if !accel_gemm_enabled() {
12100            return; // CMF_ACCEL=0
12101        }
12102        let (rows, cols, b) = (512usize, 1024usize, 8usize); // ≥500K → accel arm
12103        let gpr = cols / GROUP_SIZE;
12104        let mut bytes = vec![0u8; rows * gpr * Q4_TILE];
12105        for r in 0..rows {
12106            for g in 0..gpr {
12107                let t = (r * gpr + g) * Q4_TILE;
12108                let sc = 0.02 + 0.0005 * ((r * gpr + g) % 64) as f32;
12109                bytes[t..t + 2].copy_from_slice(&cortiq_core::quant::f32_to_f16(sc).to_le_bytes());
12110                for k in 0..16 {
12111                    bytes[t + 2 + k] = ((r * 31 + g * 7 + k * 13) % 251) as u8;
12112                }
12113            }
12114        }
12115        let x: Vec<f32> = (0..b * cols)
12116            .map(|i| ((i * 13 + 7) % 97) as f32 / 97.0 - 0.5)
12117            .collect();
12118        let mut got = vec![0f32; b * rows];
12119        q4t_matmat(&bytes, &x, b, rows, cols, &mut got, None);
12120        // Brute-force reference off the same tiles.
12121        let mut w = vec![0f32; rows * cols];
12122        for r in 0..rows {
12123            for g in 0..gpr {
12124                let t = (r * gpr + g) * Q4_TILE;
12125                let s = f16_to_f32(u16::from_le_bytes([bytes[t], bytes[t + 1]]));
12126                for (k, &bb) in bytes[t + 2..t + Q4_TILE].iter().enumerate() {
12127                    w[r * cols + g * GROUP_SIZE + k * 2] = ((bb & 0x0F) as f32 - 8.0) * s;
12128                    w[r * cols + g * GROUP_SIZE + k * 2 + 1] =
12129                        (((bb >> 4) & 0x0F) as f32 - 8.0) * s;
12130                }
12131            }
12132        }
12133        for bi in 0..b {
12134            for r in 0..rows {
12135                let want: f32 = (0..cols).map(|j| x[bi * cols + j] * w[r * cols + j]).sum();
12136                let d = (got[bi * rows + r] - want).abs();
12137                assert!(
12138                    d <= want.abs().max(1.0) * 1e-4,
12139                    "accel q4t GEMM diverged at ({bi},{r}): {} vs {want}",
12140                    got[bi * rows + r]
12141                );
12142            }
12143        }
12144    }
12145
12146    #[test]
12147    fn q4matvec_matches_full_dequant() {
12148        let (rows, cols) = (8, 64);
12149        let groups = rows * cols / GROUP_SIZE;
12150        // Hand-craft a q4_block blob: nibbles then f16 scales.
12151        let mut bytes = Vec::with_capacity(groups * 16 + groups * 2);
12152        for i in 0..groups * 16 {
12153            bytes.push((((i * 7 + 3) % 256) & 0xFF) as u8);
12154        }
12155        for g in 0..groups {
12156            let s = 0.01 + 0.003 * g as f32;
12157            bytes.extend_from_slice(&cortiq_core::quant::f32_to_f16(s).to_le_bytes());
12158        }
12159        let x: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.17).sin()).collect();
12160
12161        let mut reference = vec![0.0f32; rows * cols];
12162        cortiq_core::quant::dequant_q4_block(&bytes, &mut reference);
12163        let mut expect = vec![0.0f32; rows];
12164        for r in 0..rows {
12165            expect[r] = reference[r * cols..(r + 1) * cols]
12166                .iter()
12167                .zip(&x)
12168                .map(|(w, xv)| w * xv)
12169                .sum();
12170        }
12171
12172        let mut got = vec![0.0f32; rows];
12173        q4matvec(&bytes, &x, rows, cols, &mut got, None);
12174        // SDOT path quantizes activations to i8 (A8W8): bounded noise,
12175        // same contract as q8/vbit (exact path is pinned by CMF_SDOT=0
12176        // in the golden-parity gate).
12177        let tol = if a8w8_enabled() { 6e-2 } else { 1e-4 };
12178        let scale = expect.iter().fold(0f32, |m, v| m.max(v.abs())).max(1.0);
12179        for r in 0..rows {
12180            assert!(
12181                (got[r] - expect[r]).abs() < tol * scale,
12182                "row {r}: {} vs {}",
12183                got[r],
12184                expect[r]
12185            );
12186        }
12187    }
12188
12189    /// Fused two-input vbit matvec must equal two single matvecs exactly
12190    /// (same per-lane accumulation order on both scalar and SDOT paths).
12191    #[test]
12192    fn vbitmatvec2_equals_two_singles() {
12193        let (rows, cols) = (6, 64);
12194        let ng = cols / GROUP_SIZE;
12195        let bits: Vec<u8> = vec![3, 4, 5, 6, 8, 4];
12196        let mut bytes = bits.clone();
12197        for g in 0..rows * ng {
12198            let s = 0.02 + 0.001 * g as f32;
12199            bytes.extend_from_slice(&cortiq_core::quant::f32_to_f16(s).to_le_bytes());
12200        }
12201        for r in 0..rows {
12202            let b = bits[r] as usize;
12203            let (mut acc, mut nb) = (0u64, 0usize);
12204            let mut rowbytes = Vec::new();
12205            for i in 0..cols {
12206                let v = ((i * 7 + r * 13) % (1 << b)) as u64;
12207                acc = (acc << b) | v;
12208                nb += b;
12209                while nb >= 8 {
12210                    nb -= 8;
12211                    rowbytes.push(((acc >> nb) & 0xFF) as u8);
12212                }
12213            }
12214            if nb > 0 {
12215                rowbytes.push(((acc << (8 - nb)) & 0xFF) as u8);
12216            }
12217            bytes.extend_from_slice(&rowbytes);
12218        }
12219        let x1: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.19).sin()).collect();
12220        let x2: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.11).cos()).collect();
12221        let offsets = vbit_row_offsets(&bytes, rows, cols);
12222
12223        let (mut a1, mut a2) = (vec![0f32; rows], vec![0f32; rows]);
12224        vbitmatvec(&bytes, &offsets, &x1, rows, cols, &mut a1, None);
12225        vbitmatvec(&bytes, &offsets, &x2, rows, cols, &mut a2, None);
12226        let (mut b1, mut b2) = (vec![0f32; rows], vec![0f32; rows]);
12227        vbitmatvec2(
12228            &bytes, &offsets, &x1, &x2, rows, cols, &mut b1, &mut b2, None,
12229        );
12230        assert_eq!(a1, b1, "fused vbit lane 1 must be bit-identical");
12231        assert_eq!(a2, b2, "fused vbit lane 2 must be bit-identical");
12232    }
12233
12234    /// Fused two-input q4 matvec must equal two single matvecs exactly.
12235    #[test]
12236    fn q4matvec2_equals_two_singles() {
12237        let (rows, cols) = (8, 128);
12238        let groups = rows * cols / GROUP_SIZE;
12239        let mut bytes = Vec::with_capacity(groups * 16 + groups * 2);
12240        for i in 0..groups * 16 {
12241            bytes.push((((i * 7 + 3) % 256) & 0xFF) as u8);
12242        }
12243        for g in 0..groups {
12244            let s = 0.01 + 0.003 * g as f32;
12245            bytes.extend_from_slice(&cortiq_core::quant::f32_to_f16(s).to_le_bytes());
12246        }
12247        // Include an outlier channel so the SDOT correction path is
12248        // exercised in the pair kernel too.
12249        let mut x1: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.17).sin()).collect();
12250        x1[9] = 250.0;
12251        let x2: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.23).cos()).collect();
12252
12253        let (mut a1, mut a2) = (vec![0f32; rows], vec![0f32; rows]);
12254        q4matvec(&bytes, &x1, rows, cols, &mut a1, None);
12255        q4matvec(&bytes, &x2, rows, cols, &mut a2, None);
12256        let (mut b1, mut b2) = (vec![0f32; rows], vec![0f32; rows]);
12257        q4matvec2(&bytes, &x1, &x2, rows, cols, &mut b1, &mut b2, None);
12258        assert_eq!(a1, b1, "fused q4 lane 1 must be bit-identical");
12259        assert_eq!(a2, b2, "fused q4 lane 2 must be bit-identical");
12260    }
12261
12262    /// Multi-matrix job must equal separate matvecs exactly — same
12263    /// kernels, only the dispatch is fused.
12264    #[test]
12265    fn matvec_many_equals_separate_matvecs() {
12266        use crate::pool::Pool;
12267        let (r1, r2, cols) = (300, 200, 64);
12268        let mk = |salt: usize, rows: usize| {
12269            QTensor::from_f32(
12270                (0..rows * cols)
12271                    .map(|i| ((i * 7 + salt) % 97) as f32 / 97.0 - 0.5)
12272                    .collect(),
12273                rows,
12274                cols,
12275            )
12276        };
12277        let (a, b) = (mk(1, r1), mk(5, r2));
12278        let x: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.11).sin()).collect();
12279        let pool = Pool::new(3);
12280
12281        let (mut ea, mut eb) = (vec![0f32; r1], vec![0f32; r2]);
12282        a.matvec(&x, &mut ea, Some(&pool));
12283        b.matvec(&x, &mut eb, Some(&pool));
12284        let (mut ga, mut gb) = (vec![0f32; r1], vec![0f32; r2]);
12285        QTensor::matvec_many([&a, &b], &x, [&mut ga, &mut gb], Some(&pool));
12286        assert_eq!(ea, ga, "fused multi-matrix lane 1 must be bit-identical");
12287        assert_eq!(eb, gb, "fused multi-matrix lane 2 must be bit-identical");
12288    }
12289
12290    /// The public Q4TP operator must take the real mapped matvec_many arm,
12291    /// rather than the F32 fallback above.  Build a tiny valid CMF so both
12292    /// handles retain their mmap payloads, then compare the fused dispatch
12293    /// with two ordinary mapped matvec calls bit-for-bit.
12294    #[test]
12295    fn q4tp_matvec_many_equals_separate_matvecs() {
12296        use crate::pool::Pool;
12297        use cortiq_core::{CMF_VERSION, CmfHeader, CmfModel, QuantType, TensorSpec};
12298
12299        let (r1, r2, cols) = (300usize, 200usize, 64usize);
12300        let arch: cortiq_core::ModelArch = serde_json::from_value(serde_json::json!({
12301            "arch_name": "tiny-q4tp",
12302            "hidden_size": cols,
12303            "intermediate_size": cols * 2,
12304            "num_layers": 1,
12305            "num_attention_heads": 2,
12306            "num_kv_heads": 1,
12307            "head_dim": 32,
12308            "vocab_size": r1,
12309            "layer_types": ["FullAttention"],
12310            "rms_norm_eps": 1e-6,
12311            "max_position_embeddings": 8,
12312            "linear_conv_kernel_dim": 0,
12313            "linear_num_key_heads": 0,
12314            "linear_num_value_heads": 0
12315        }))
12316        .unwrap();
12317        let header = CmfHeader {
12318            format: "cmf".into(),
12319            version: CMF_VERSION,
12320            arch,
12321            quant_type: QuantType::Q4Block,
12322            provenance: None,
12323            tokenizer_config: None,
12324            section_hashes: None,
12325            skills: Vec::new(),
12326            shard: None,
12327            calibration: None,
12328            routing: None,
12329        };
12330        let specs = [
12331            TensorSpec {
12332                name: "q".into(),
12333                dtype: TensorDtype::Q4TiledP,
12334                shape: vec![r1, cols],
12335                data: synth_q4tp(r1, cols),
12336            },
12337            TensorSpec {
12338                name: "kv".into(),
12339                dtype: TensorDtype::Q4TiledP,
12340                shape: vec![r2, cols],
12341                data: synth_q4tp(r2, cols),
12342            },
12343        ];
12344        let dir = std::env::temp_dir().join(format!("cmf-q4tp-many-{}", std::process::id()));
12345        std::fs::create_dir_all(&dir).unwrap();
12346        let path = dir.join("m.cmf");
12347        CmfModel::write(&path, &header, &specs, None, None).unwrap();
12348        let model = Arc::new(CmfModel::open(&path).unwrap());
12349        let (a, b) = (
12350            QTensor::from_model(&model, "q").unwrap(),
12351            QTensor::from_model(&model, "kv").unwrap(),
12352        );
12353        assert_eq!(a.model_dtype(), Some(TensorDtype::Q4TiledP));
12354        assert_eq!(b.model_dtype(), Some(TensorDtype::Q4TiledP));
12355        let x: Vec<f32> = (0..cols)
12356            .map(|i| ((i * 17 + 3) % 97) as f32 / 97.0 - 0.5)
12357            .collect();
12358        let pool = Pool::new(3);
12359        let (mut ea, mut eb) = (vec![0.0f32; r1], vec![0.0f32; r2]);
12360        a.matvec(&x, &mut ea, Some(&pool));
12361        b.matvec(&x, &mut eb, Some(&pool));
12362        let (mut ga, mut gb) = (vec![0.0f32; r1], vec![0.0f32; r2]);
12363        QTensor::matvec_many([&a, &b], &x, [&mut ga, &mut gb], Some(&pool));
12364        assert_eq!(ea, ga, "Q4TP fused lane 1 must be bit-identical");
12365        assert_eq!(eb, gb, "Q4TP fused lane 2 must be bit-identical");
12366        let _ = std::fs::remove_dir_all(&dir);
12367    }
12368
12369    /// The MiMo speculative verify's kernels: several tokens' MoE through
12370    /// `moe_gate_up_rows` / `moe_down_rows` (+ the caller's route-order sum)
12371    /// is bit-identical to each token alone through `moe_gate_up_many` /
12372    /// `moe_down_many` (decode), and a row-exact `q4tp_matmat` of five
12373    /// tokens (wide enough for the blocked tiles) equals five matvecs.
12374    #[test]
12375    fn multi_token_moe_rows_equal_single_token_decode() {
12376        use crate::pool::Pool;
12377        use cortiq_core::{CMF_VERSION, CmfHeader, CmfModel, QuantType, TensorSpec};
12378
12379        let (h, inter, ne) = (64usize, 128usize, 3usize);
12380        let arch: cortiq_core::ModelArch = serde_json::from_value(serde_json::json!({
12381            "arch_name": "tiny-q4tp-moe",
12382            "hidden_size": h,
12383            "intermediate_size": inter,
12384            "num_layers": 1,
12385            "num_attention_heads": 2,
12386            "num_kv_heads": 1,
12387            "head_dim": 32,
12388            "vocab_size": 8,
12389            "layer_types": ["FullAttention"],
12390            "rms_norm_eps": 1e-6,
12391            "max_position_embeddings": 8,
12392            "linear_conv_kernel_dim": 0,
12393            "linear_num_key_heads": 0,
12394            "linear_num_value_heads": 0
12395        }))
12396        .unwrap();
12397        let header = CmfHeader {
12398            format: "cmf".into(),
12399            version: CMF_VERSION,
12400            arch,
12401            quant_type: QuantType::Q4Block,
12402            provenance: None,
12403            tokenizer_config: None,
12404            section_hashes: None,
12405            skills: Vec::new(),
12406            shard: None,
12407            calibration: None,
12408            routing: None,
12409        };
12410        let mut specs = Vec::new();
12411        for e in 0..ne {
12412            for (k, (n, r, c)) in [("g", inter, h), ("u", inter, h), ("d", h, inter)]
12413                .into_iter()
12414                .enumerate()
12415            {
12416                // Distinct experts: perturb only the nibble plane (any byte
12417                // is a valid pair of codes; the ladder stays intact).
12418                let mut data = synth_q4tp(r, c);
12419                for (i, byte) in data[..r * (c / GROUP_SIZE) * Q4TP_NIB]
12420                    .iter_mut()
12421                    .enumerate()
12422                {
12423                    *byte ^= ((i * (e * 3 + k + 1)) % 251) as u8;
12424                }
12425                specs.push(TensorSpec {
12426                    name: format!("{n}{e}"),
12427                    dtype: TensorDtype::Q4TiledP,
12428                    shape: vec![r, c],
12429                    data,
12430                });
12431            }
12432        }
12433        let dir = std::env::temp_dir().join(format!(
12434            "cmf-moe-rows-{}-{}",
12435            std::process::id(),
12436            FLOAT_ACTIVATIONS.get()
12437        ));
12438        std::fs::create_dir_all(&dir).unwrap();
12439        let path = dir.join("m.cmf");
12440        CmfModel::write(&path, &header, &specs, None, None).unwrap();
12441        let model = Arc::new(CmfModel::open(&path).unwrap());
12442        let t = |n: String| QTensor::from_model(&model, &n).unwrap();
12443        let g: Vec<QTensor> = (0..ne).map(|e| t(format!("g{e}"))).collect();
12444        let u: Vec<QTensor> = (0..ne).map(|e| t(format!("u{e}"))).collect();
12445        let d: Vec<QTensor> = (0..ne).map(|e| t(format!("d{e}"))).collect();
12446        let b = 4usize;
12447        let mut xs: Vec<f32> = (0..b * h)
12448            .map(|i| ((i * 31 + 7) % 89) as f32 / 89.0 - 0.5)
12449            .collect();
12450        xs[5] = 9.0; // an activation outlier on token 0
12451        // Token -> (experts in route order, weights).
12452        let routes: Vec<(Vec<usize>, Vec<f32>)> = vec![
12453            (vec![2, 0], vec![0.6, 0.4]),
12454            (vec![0, 1, 2], vec![0.2, 0.5, 0.3]),
12455            (vec![1], vec![1.0]),
12456            (vec![2, 1, 0], vec![0.25, 0.25, 0.5]),
12457        ];
12458        let pool = Pool::new(3);
12459        // Decode reference, token by token.
12460        let mut want = vec![0f32; b * h];
12461        for (tk, (idx, w)) in routes.iter().enumerate() {
12462            let x = &xs[tk * h..(tk + 1) * h];
12463            let pairs: Vec<(&QTensor, &QTensor)> = idx.iter().map(|&e| (&g[e], &u[e])).collect();
12464            let mut gs: Vec<Vec<f32>> = idx.iter().map(|_| vec![0f32; inter]).collect();
12465            assert!(QTensor::moe_gate_up_many(&pairs, x, &mut gs, Some(&pool)));
12466            if FLOAT_ACTIVATIONS.get() {
12467                for (slot, &e) in idx.iter().enumerate() {
12468                    let (mut gate, mut up) = (vec![0.0; inter], vec![0.0; inter]);
12469                    g[e].matvec(x, &mut gate, Some(&pool));
12470                    u[e].matvec(x, &mut up, Some(&pool));
12471                    for (v, u) in gate.iter_mut().zip(up) {
12472                        *v = (*v / (1.0 + (-*v).exp())) * u;
12473                    }
12474                    assert_eq!(gs[slot], gate, "float gate/up must equal ordinary matvecs");
12475                }
12476            }
12477            let downs: Vec<&QTensor> = idx.iter().map(|&e| &d[e]).collect();
12478            assert!(QTensor::moe_down_many(
12479                &downs,
12480                &gs,
12481                w,
12482                &mut want[tk * h..(tk + 1) * h],
12483                Some(&pool)
12484            ));
12485        }
12486        if FLOAT_ACTIVATIONS.get() {
12487            for (tk, (idx, w)) in routes.iter().enumerate() {
12488                let mut scalar = vec![0.0; h];
12489                for (&e, &weight) in idx.iter().zip(w) {
12490                    let (mut gate, mut up, mut down) =
12491                        (vec![0.0; inter], vec![0.0; inter], vec![0.0; h]);
12492                    g[e].matvec(&xs[tk * h..(tk + 1) * h], &mut gate, Some(&pool));
12493                    u[e].matvec(&xs[tk * h..(tk + 1) * h], &mut up, Some(&pool));
12494                    for (v, u) in gate.iter_mut().zip(up) {
12495                        *v = (*v / (1.0 + (-*v).exp())) * u;
12496                    }
12497                    d[e].matvec(&gate, &mut down, Some(&pool));
12498                    for (v, d) in scalar.iter_mut().zip(down) {
12499                        *v += weight * d;
12500                    }
12501                }
12502                assert_eq!(
12503                    &want[tk * h..(tk + 1) * h],
12504                    scalar,
12505                    "float many equals scalar experts"
12506                );
12507            }
12508        }
12509        // All four tokens at once, grouped by expert.
12510        let mut experts: Vec<usize> = Vec::new();
12511        let mut groups: Vec<Vec<usize>> = Vec::new();
12512        for (tk, (idx, _)) in routes.iter().enumerate() {
12513            for &e in idx {
12514                match experts.iter().position(|&x| x == e) {
12515                    Some(k) => groups[k].push(tk),
12516                    None => {
12517                        experts.push(e);
12518                        groups.push(vec![tk]);
12519                    }
12520                }
12521            }
12522        }
12523        let n_pairs: usize = groups.iter().map(|g| g.len()).sum();
12524        let pairs: Vec<(&QTensor, &QTensor)> = experts.iter().map(|&e| (&g[e], &u[e])).collect();
12525        let mut gs: Vec<Vec<f32>> = (0..n_pairs).map(|_| vec![0f32; inter]).collect();
12526        assert!(QTensor::moe_gate_up_rows(
12527            &pairs,
12528            &groups,
12529            &xs,
12530            &mut gs,
12531            Some(&pool)
12532        ));
12533        let downs: Vec<&QTensor> = experts.iter().map(|&e| &d[e]).collect();
12534        let lens: Vec<usize> = groups.iter().map(|g| g.len()).collect();
12535        let mut ds: Vec<Vec<f32>> = (0..n_pairs).map(|_| vec![0f32; h]).collect();
12536        assert!(QTensor::moe_down_rows(
12537            &downs,
12538            &lens,
12539            &gs,
12540            &mut ds,
12541            Some(&pool)
12542        ));
12543        let slot = |tk: usize, e: usize| {
12544            let k = experts.iter().position(|&x| x == e).unwrap();
12545            groups[..k].iter().map(|g| g.len()).sum::<usize>()
12546                + groups[k].iter().position(|&x| x == tk).unwrap()
12547        };
12548        let mut got = vec![0f32; b * h];
12549        for (tk, (idx, w)) in routes.iter().enumerate() {
12550            for i in 0..h {
12551                let mut acc = 0f32;
12552                for (&e, &we) in idx.iter().zip(w) {
12553                    acc += we * ds[slot(tk, e)][i];
12554                }
12555                got[tk * h + i] = acc;
12556            }
12557        }
12558        assert!(want.iter().any(|v| *v != 0.0));
12559        assert_eq!(
12560            want.iter().map(|v| v.to_bits()).collect::<Vec<_>>(),
12561            got.iter().map(|v| v.to_bits()).collect::<Vec<_>>(),
12562            "multi-token MoE must equal decode bit for bit"
12563        );
12564
12565        // Row-exact q4tp_matmat: five tokens (a blocked 1x4 tile + tail
12566        // otherwise) equal five matvecs.
12567        let b5 = 5usize;
12568        let x5: Vec<f32> = (0..b5 * h)
12569            .map(|i| ((i * 13 + 5) % 71) as f32 / 71.0 - 0.5)
12570            .collect();
12571        let mut mm = vec![0f32; b5 * inter];
12572        row_exact_scope(|| g[1].matmat(&x5, b5, &mut mm, Some(&pool)));
12573        for tk in 0..b5 {
12574            let mut mv = vec![0f32; inter];
12575            g[1].matvec(&x5[tk * h..(tk + 1) * h], &mut mv, Some(&pool));
12576            assert_eq!(
12577                mv.iter().map(|v| v.to_bits()).collect::<Vec<_>>(),
12578                mm[tk * inter..(tk + 1) * inter]
12579                    .iter()
12580                    .map(|v| v.to_bits())
12581                    .collect::<Vec<_>>(),
12582                "row-exact matmat token {tk}"
12583            );
12584        }
12585        // Other concurrent tests/requests may still hold the shared mode.
12586        // Nested, overlapping and unwind restoration is checked separately.
12587        let _ = std::fs::remove_dir_all(&dir);
12588    }
12589
12590    /// The row-exact fix must not touch the fast path. Outside the scope
12591    /// `q4tp_matmat` has to produce exactly what it did before, and on ARM
12592    /// "before" is spelled out below: the tuned 1x4 SDOT tile for every
12593    /// four columns and the single-row kernel for the tail. Inside the
12594    /// scope every column equals its token's matvec. `q4tp_matmat_with`
12595    /// takes the mode as an argument, so a concurrent test holding the
12596    /// shared scope cannot flip it under this one.
12597    #[test]
12598    fn q4tp_matmat_fast_path_unchanged_outside_row_exact() {
12599        use crate::pool::Pool;
12600        use std::sync::atomic::Ordering::Relaxed;
12601        let _alt = Q4TP_ALT_TEST_LOCK.lock().unwrap_or_else(|e| e.into_inner());
12602        // The tuned ARM shape, which is also what an unset switch picks
12603        // unless CMF_Q4TP_V1 is exported.
12604        Q4TP_ALT.store(2, Relaxed);
12605        let pool = Pool::new(3);
12606        // Under 500k cells, so macOS keeps the matmat off the AMX; the
12607        // second shape runs across pool workers (rows >= 256) with 32
12608        // groups of accumulation and a tail after two 1x4 tiles.
12609        for &(rows, cols, b) in &[(64usize, 256usize, 7usize), (320, 1024, 9)] {
12610            let bytes = synth_q4tp(rows, cols);
12611            let mut xs: Vec<f32> = (0..b * cols)
12612                .map(|i| ((i * 29 + 11) % 83) as f32 / 83.0 - 0.5)
12613                .collect();
12614            xs[3] = 7.5; // an activation outlier on token 0
12615            let run = |exact: bool| {
12616                let mut out = vec![0f32; b * rows];
12617                q4tp_matmat_with(&bytes, &xs, b, rows, cols, &mut out, Some(&pool), exact);
12618                out
12619            };
12620            let (fast, exact) = (run(false), run(true));
12621            #[cfg(not(target_arch = "aarch64"))]
12622            let _ = fast;
12623            let bits = |v: &[f32]| v.iter().map(|x| x.to_bits()).collect::<Vec<_>>();
12624            let mut matvecs = vec![0f32; b * rows];
12625            for (bi, o) in matvecs.chunks_mut(rows).enumerate() {
12626                q4tp_matvec(
12627                    &bytes,
12628                    &xs[bi * cols..(bi + 1) * cols],
12629                    rows,
12630                    cols,
12631                    o,
12632                    Some(&pool),
12633                );
12634            }
12635            assert!(matvecs.iter().any(|v| *v != 0.0));
12636            assert_eq!(
12637                bits(&exact),
12638                bits(&matvecs),
12639                "{rows}x{cols} b={b}: row-exact matmat must equal per-token matvecs"
12640            );
12641            #[cfg(target_arch = "aarch64")]
12642            {
12643                // The pre-fix ARM loop, cell for cell.
12644                let gpr = cols / GROUP_SIZE;
12645                let v = Q4tpView::new(&bytes, rows, cols);
12646                let mut old = vec![0f32; b * rows];
12647                let mut sc = vec![0f32; gpr];
12648                let a8w8 = a8w8_enabled();
12649                let blocked = sdot_enabled() && blocked_enabled();
12650                let acts: Vec<SplitAct> = (0..b)
12651                    .map(|bi| split_act(&xs[bi * cols..(bi + 1) * cols]))
12652                    .collect();
12653                for r in 0..rows {
12654                    v.scales_into(r, gpr, &mut sc);
12655                    if !a8w8 {
12656                        for bi in 0..b {
12657                            let x = &xs[bi * cols..(bi + 1) * cols];
12658                            old[bi * rows + r] = q4tp_row_exact(v.nib, r, gpr, x, &sc);
12659                        }
12660                        continue;
12661                    }
12662                    let finish = |d: f32, act: &SplitAct| {
12663                        let mut acc = d * act.sx;
12664                        for &(j, xv) in &act.outliers {
12665                            let (w, s) = q4tp_outlier(v.nib, r, gpr, j, &sc);
12666                            acc += w * s * xv;
12667                        }
12668                        acc
12669                    };
12670                    let mut bi = 0usize;
12671                    while blocked && bi + 4 <= b {
12672                        let xs4 = [
12673                            acts[bi].xq.as_slice(),
12674                            acts[bi + 1].xq.as_slice(),
12675                            acts[bi + 2].xq.as_slice(),
12676                            acts[bi + 3].xq.as_slice(),
12677                        ];
12678                        let d = unsafe { dot_q4tp_row_1x4_sdot(v.nib, r, gpr, xs4, &sc) };
12679                        for k in 0..4 {
12680                            old[(bi + k) * rows + r] = finish(d[k], &acts[bi + k]);
12681                        }
12682                        bi += 4;
12683                    }
12684                    for (bi, act) in acts.iter().enumerate().skip(bi) {
12685                        let d = dot_q4tp_row_i8(v.nib, r, gpr, &act.xq, &sc);
12686                        old[bi * rows + r] = finish(d, act);
12687                    }
12688                }
12689                assert_eq!(
12690                    bits(&fast),
12691                    bits(&old),
12692                    "{rows}x{cols} b={b}: the fast path changed outside row_exact"
12693                );
12694                // And it is still the fast tile that runs: its lane-parallel
12695                // fma sum rounds differently from the matvec somewhere.
12696                if blocked {
12697                    assert_ne!(
12698                        bits(&fast),
12699                        bits(&matvecs),
12700                        "{rows}x{cols} b={b}: the tuned tile no longer runs outside row_exact"
12701                    );
12702                }
12703            }
12704        }
12705        Q4TP_ALT.store(0, Relaxed);
12706    }
12707
12708    #[test]
12709    fn row_exact_scopes_survive_overlap_nesting_and_unwind() {
12710        use std::sync::{Barrier, atomic::{AtomicUsize, Ordering}};
12711        // A private counter makes this restoration test independent of
12712        // numerical tests concurrently using the production counter.
12713        let active = AtomicUsize::new(0);
12714        counted_row_exact_scope(&active, || {
12715            assert_eq!(active.load(Ordering::Acquire), 1);
12716            counted_row_exact_scope(&active, || {
12717                assert_eq!(active.load(Ordering::Acquire), 2);
12718            });
12719            assert_eq!(active.load(Ordering::Acquire), 1);
12720        });
12721        assert_eq!(active.load(Ordering::Acquire), 0);
12722
12723        let both_entered = Barrier::new(2);
12724        let release_last = Barrier::new(2);
12725        std::thread::scope(|s| {
12726            let first = s.spawn(|| counted_row_exact_scope(&active, || {
12727                both_entered.wait();
12728            }));
12729            let last = s.spawn(|| counted_row_exact_scope(&active, || {
12730                both_entered.wait();
12731                release_last.wait();
12732            }));
12733            first.join().unwrap();
12734            let after_first = active.load(Ordering::Acquire);
12735            release_last.wait();
12736            last.join().unwrap();
12737            assert_eq!(after_first, 1, "second request must remain exact");
12738        });
12739        assert_eq!(active.load(Ordering::Acquire), 0);
12740        let panic = std::panic::catch_unwind(|| {
12741            counted_row_exact_scope(&active, || panic!("scope unwind"));
12742        });
12743        assert!(panic.is_err());
12744        assert_eq!(active.load(Ordering::Acquire), 0);
12745    }
12746
12747    #[test]
12748    #[cfg(target_arch = "x86_64")]
12749    fn q4tp_float_avx2_is_bitwise_scalar() {
12750        if !avx2_enabled() {
12751            return;
12752        }
12753        for cols in [32, 64, 96, 2048, 4096] {
12754            let rows = 9;
12755            let bytes = synth_q4tp(rows, cols);
12756            let v = Q4tpView::new(&bytes, rows, cols);
12757            let gpr = cols / GROUP_SIZE;
12758            let mut sc = vec![0.0; gpr];
12759            for seed in 1..=5 {
12760                let xs: Vec<f32> = (0..cols)
12761                    .map(|i| (((i * 104729 + seed * 8191) % 100003) as f32 - 50001.0) / 7919.0)
12762                    .collect();
12763                for r in 0..rows {
12764                    v.scales_into(r, gpr, &mut sc);
12765                    let scalar = q4tp_row_float_scalar(v.nib, r, gpr, &xs, &sc);
12766                    let vector = unsafe { q4tp_row_float_avx2(v.nib, r, gpr, &xs, &sc) };
12767                    assert_eq!(
12768                        scalar.to_bits(),
12769                        vector.to_bits(),
12770                        "cols={cols} row={r} seed={seed}"
12771                    );
12772                }
12773            }
12774        }
12775    }
12776
12777    #[test]
12778    fn multi_token_moe_rows_float_equal_single_token_decode() {
12779        float_activations_scope(multi_token_moe_rows_equal_single_token_decode);
12780    }
12781
12782    #[test]
12783    fn full_gpu_q8_scope_is_nested_and_thread_local() {
12784        assert!(!FULL_GPU_Q8.get());
12785        let before = gpu_split_frac();
12786        {
12787            let _guard = enter_full_gpu_q8_scope();
12788            assert_eq!(gpu_split_frac(), 1.0);
12789            {
12790                let _nested = enter_full_gpu_q8_scope();
12791            }
12792            assert_eq!(gpu_split_frac(), 1.0);
12793            std::thread::spawn(|| assert!(!FULL_GPU_Q8.get())).join().unwrap();
12794        }
12795        assert!(!FULL_GPU_Q8.get());
12796        assert_eq!(gpu_split_frac(), before);
12797    }
12798
12799    #[test]
12800    fn float_activation_scope_is_nested_thread_local_and_unwind_safe() {
12801        assert!(!FLOAT_ACTIVATIONS.get());
12802        let before = a8w8_enabled();
12803        float_activations_scope(|| {
12804            assert!(!a8w8_enabled());
12805            float_activations_scope(|| assert!(!a8w8_enabled()));
12806            assert!(FLOAT_ACTIVATIONS.get());
12807            std::thread::spawn(|| assert!(!FLOAT_ACTIVATIONS.get()))
12808                .join()
12809                .unwrap();
12810        });
12811        assert!(!FLOAT_ACTIVATIONS.get());
12812        assert_eq!(a8w8_enabled(), before);
12813        let _ = std::panic::catch_unwind(|| float_activations_scope(|| panic!("test unwind")));
12814        assert!(!FLOAT_ACTIVATIONS.get());
12815    }
12816
12817    /// Batched q4/vbit matmat must equal per-position matvec calls
12818    /// exactly (the fallback it replaced) — same kernels, same order.
12819    #[test]
12820    fn batched_matmat_equals_per_position_matvec() {
12821        let (rows, cols, b) = (8, 64, 5);
12822        // q4 blob.
12823        let groups = rows * cols / GROUP_SIZE;
12824        let mut q4 = Vec::new();
12825        for i in 0..groups * 16 {
12826            q4.push((((i * 7 + 3) % 256) & 0xFF) as u8);
12827        }
12828        for g in 0..groups {
12829            q4.extend_from_slice(
12830                &cortiq_core::quant::f32_to_f16(0.01 + 0.003 * g as f32).to_le_bytes(),
12831            );
12832        }
12833        // vbit blob (mixed widths incl. 8).
12834        let ng = cols / GROUP_SIZE;
12835        let bits: Vec<u8> = vec![3, 4, 5, 6, 8, 4, 5, 3];
12836        let mut vb = bits.clone();
12837        for g in 0..rows * ng {
12838            vb.extend_from_slice(
12839                &cortiq_core::quant::f32_to_f16(0.02 + 0.001 * g as f32).to_le_bytes(),
12840            );
12841        }
12842        for r in 0..rows {
12843            let bw = bits[r] as usize;
12844            let (mut acc, mut nb) = (0u64, 0usize);
12845            let mut rowbytes = Vec::new();
12846            for i in 0..cols {
12847                let v = ((i * 7 + r * 13) % (1 << bw)) as u64;
12848                acc = (acc << bw) | v;
12849                nb += bw;
12850                while nb >= 8 {
12851                    nb -= 8;
12852                    rowbytes.push(((acc >> nb) & 0xFF) as u8);
12853                }
12854            }
12855            if nb > 0 {
12856                rowbytes.push(((acc << (8 - nb)) & 0xFF) as u8);
12857            }
12858            vb.extend_from_slice(&rowbytes);
12859        }
12860        let offsets = vbit_row_offsets(&vb, rows, cols);
12861
12862        let xs: Vec<f32> = (0..b * cols).map(|i| (i as f32 * 0.13).sin()).collect();
12863
12864        // q4: batch vs singles.
12865        let mut got = vec![0f32; b * rows];
12866        q4matmat(&q4, &xs, b, rows, cols, &mut got, None);
12867        for bi in 0..b {
12868            let mut expect = vec![0f32; rows];
12869            q4matvec(
12870                &q4,
12871                &xs[bi * cols..(bi + 1) * cols],
12872                rows,
12873                cols,
12874                &mut expect,
12875                None,
12876            );
12877            assert_eq!(
12878                &got[bi * rows..(bi + 1) * rows],
12879                &expect[..],
12880                "q4 batch pos {bi}"
12881            );
12882        }
12883
12884        // vbit: batch vs singles.
12885        let mut got = vec![0f32; b * rows];
12886        vbitmatmat(&vb, &offsets, &xs, b, rows, cols, &mut got, None);
12887        for bi in 0..b {
12888            let mut expect = vec![0f32; rows];
12889            vbitmatvec(
12890                &vb,
12891                &offsets,
12892                &xs[bi * cols..(bi + 1) * cols],
12893                rows,
12894                cols,
12895                &mut expect,
12896                None,
12897            );
12898            assert_eq!(
12899                &got[bi * rows..(bi + 1) * rows],
12900                &expect[..],
12901                "vbit batch pos {bi}"
12902            );
12903        }
12904    }
12905
12906    /// q4_tiled kernels must produce BIT-identical outputs to the q4
12907    /// split kernels on the same values (same ints, same order — only
12908    /// the byte placement differs).
12909    #[test]
12910    fn q4_tiled_matches_q4_block_bitexact() {
12911        let (rows, cols, b) = (8usize, 128usize, 3usize);
12912        let groups = rows * cols / GROUP_SIZE;
12913        let mut split = Vec::with_capacity(groups * 18);
12914        for i in 0..groups * 16 {
12915            split.push((((i * 7 + 3) % 256) & 0xFF) as u8);
12916        }
12917        for g in 0..groups {
12918            split.extend_from_slice(
12919                &cortiq_core::quant::f32_to_f16(0.01 + 0.003 * g as f32).to_le_bytes(),
12920            );
12921        }
12922        // Re-tile: [scale][nibbles] per group.
12923        let (packed, scales) = split.split_at(groups * 16);
12924        let mut tiled = Vec::with_capacity(groups * Q4_TILE);
12925        for g in 0..groups {
12926            tiled.extend_from_slice(&scales[g * 2..g * 2 + 2]);
12927            tiled.extend_from_slice(&packed[g * 16..(g + 1) * 16]);
12928        }
12929
12930        let mut x1: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.17).sin()).collect();
12931        x1[9] = 250.0; // exercise the outlier path
12932        let x2: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.23).cos()).collect();
12933
12934        let (mut a, mut t) = (vec![0f32; rows], vec![0f32; rows]);
12935        q4matvec(&split, &x1, rows, cols, &mut a, None);
12936        q4t_matvec(&tiled, &x1, rows, cols, &mut t, None);
12937        assert_eq!(a, t, "q4t matvec must match q4 bit-for-bit");
12938
12939        let (mut a1, mut a2) = (vec![0f32; rows], vec![0f32; rows]);
12940        let (mut t1, mut t2) = (vec![0f32; rows], vec![0f32; rows]);
12941        q4matvec2(&split, &x1, &x2, rows, cols, &mut a1, &mut a2, None);
12942        q4t_matvec2(&tiled, &x1, &x2, rows, cols, &mut t1, &mut t2, None);
12943        assert_eq!(a1, t1);
12944        assert_eq!(a2, t2);
12945
12946        let xs: Vec<f32> = (0..b * cols).map(|i| (i as f32 * 0.13).sin()).collect();
12947        let (mut am, mut tm) = (vec![0f32; b * rows], vec![0f32; b * rows]);
12948        q4matmat(&split, &xs, b, rows, cols, &mut am, None);
12949        q4t_matmat(&tiled, &xs, b, rows, cols, &mut tm, None);
12950        assert_eq!(am, tm, "q4t matmat must match q4 bit-for-bit");
12951    }
12952
12953    /// q4 SDOT outlier correction: a single huge activation channel
12954    /// (>8·rms → outlier, zeroed in xq) must still contribute its EXACT
12955    /// term. On-grid bulk (±1/0 → xq dequantizes exactly) isolates the
12956    /// correction from A8W8 noise. cols must exceed 64: at n=64 the
12957    /// 8·rms threshold equals sqrt(v²+rest) ≥ v, so a single outlier
12958    /// can never qualify (8² = n).
12959    #[test]
12960    fn q4matvec_sdot_outlier_exact() {
12961        let (rows, cols) = (4, 128);
12962        let groups = rows * cols / GROUP_SIZE;
12963        let mut bytes = Vec::with_capacity(groups * 16 + groups * 2);
12964        for i in 0..groups * 16 {
12965            bytes.push(((i * 11 + 5) % 256) as u8);
12966        }
12967        for g in 0..groups {
12968            let s = 0.02 + 0.002 * g as f32;
12969            bytes.extend_from_slice(&cortiq_core::quant::f32_to_f16(s).to_le_bytes());
12970        }
12971        let mut x: Vec<f32> = (0..cols)
12972            .map(|i| match i % 3 {
12973                0 => 1.0,
12974                1 => -1.0,
12975                _ => 0.0,
12976            })
12977            .collect();
12978        x[17] = 300.0; // ≫ 8·rms → outlier channel
12979
12980        let mut reference = vec![0.0f32; rows * cols];
12981        cortiq_core::quant::dequant_q4_block(&bytes, &mut reference);
12982        let mut expect = vec![0.0f32; rows];
12983        for r in 0..rows {
12984            expect[r] = reference[r * cols..(r + 1) * cols]
12985                .iter()
12986                .zip(&x)
12987                .map(|(w, xv)| w * xv)
12988                .sum();
12989        }
12990        let mut got = vec![0.0f32; rows];
12991        q4matvec(&bytes, &x, rows, cols, &mut got, None);
12992        let scale = expect.iter().fold(0f32, |m, v| m.max(v.abs())).max(1.0);
12993        for r in 0..rows {
12994            assert!(
12995                (got[r] - expect[r]).abs() < 2e-3 * scale,
12996                "row {r}: {} vs {} (outlier term must be exact)",
12997                got[r],
12998                expect[r]
12999            );
13000        }
13001    }
13002
13003    /// The fused q1t matvec must equal the reference (dequant_q1t → dot),
13004    /// including the ternary zero level and the binary-searched outlier
13005    /// overlay. Guards the mmap kernel that makes a 12B q1t runnable.
13006    #[test]
13007    fn q1t_matvec_matches_reference() {
13008        use cortiq_core::quant::{dequant_q1t, f32_to_f16};
13009        let (rows, cols) = (3usize, 64usize); // gpr = 2
13010        let gpr = cols / GROUP_SIZE;
13011        let scales = [0.5f32, 0.3, 0.7, 0.2, 0.6, 0.15];
13012        // Overlay (must be sorted by flat index): a few spikes across rows.
13013        let outliers: [(u32, f32); 3] = [(5, 9.0), (70, -4.5), (150, 3.25)];
13014        let is_out = |flat: usize| outliers.iter().any(|&(i, _)| i as usize == flat);
13015        let mut bytes = Vec::new();
13016        for r in 0..rows {
13017            for g in 0..gpr {
13018                bytes.extend_from_slice(&f32_to_f16(scales[r * gpr + g]).to_le_bytes());
13019                let mut c = [0u8; 7];
13020                for k in 0..GROUP_SIZE {
13021                    // Encoder invariant: code 0 at outlier positions.
13022                    let code = if is_out(r * cols + g * GROUP_SIZE + k) {
13023                        0
13024                    } else {
13025                        ((k + r * 3 + g) % 3) as u8 // 0,1,2
13026                    };
13027                    cortiq_core::quant::q1t_pack(&mut c, k, code);
13028                }
13029                bytes.extend_from_slice(&c);
13030            }
13031        }
13032        // Per-row overlay: [u32 row_ptr[rows+1]] then [(u16 col, f16 val)] by
13033        // row (outliers are sorted by flat index → already grouped by row).
13034        let mut row_ptr = vec![0u32; rows + 1];
13035        for &(idx, _) in &outliers {
13036            row_ptr[idx as usize / cols + 1] += 1;
13037        }
13038        for r in 0..rows {
13039            row_ptr[r + 1] += row_ptr[r];
13040        }
13041        for &p in &row_ptr {
13042            bytes.extend_from_slice(&p.to_le_bytes());
13043        }
13044        for &(idx, v) in &outliers {
13045            bytes.extend_from_slice(&((idx as usize % cols) as u16).to_le_bytes());
13046            bytes.extend_from_slice(&f32_to_f16(v).to_le_bytes());
13047        }
13048
13049        let mut refw = vec![0f32; rows * cols];
13050        dequant_q1t(&bytes, rows, cols, &mut refw);
13051        // On-grid activations (±1, amax 1) so the int8 SDOT path reconstructs
13052        // x exactly and matches the f32 reference (same trick as the q1 test).
13053        let x: Vec<f32> = (0..cols)
13054            .map(|j| if j % 3 == 0 { 1.0 } else { -1.0 })
13055            .collect();
13056        let mut expect = vec![0f32; rows];
13057        for r in 0..rows {
13058            let mut a = 0.0f32;
13059            for j in 0..cols {
13060                a += refw[r * cols + j] * x[j];
13061            }
13062            expect[r] = a;
13063        }
13064        let tol = |e: f32| 1e-3 * e.abs().max(1e-3);
13065        let mut got = vec![0f32; rows];
13066        q1t_matvec(&bytes, &x, rows, cols, &mut got, None);
13067        for r in 0..rows {
13068            assert!(
13069                (got[r] - expect[r]).abs() < tol(expect[r]),
13070                "row {r}: {} vs {}",
13071                got[r],
13072                expect[r]
13073            );
13074        }
13075        // matmat (b=2, f32 decode path) must agree too.
13076        let x2: Vec<f32> = x.iter().chain(x.iter()).copied().collect();
13077        let mut gm = vec![0f32; 2 * rows];
13078        q1t_matmat(&bytes, &x2, 2, rows, cols, &mut gm, None);
13079        for r in 0..rows {
13080            assert!((gm[r] - expect[r]).abs() < tol(expect[r]));
13081            assert!((gm[rows + r] - expect[r]).abs() < tol(expect[r]));
13082        }
13083        // Fused pair (q1t_matvec2) must equal two single matvecs
13084        // bit-for-bit: same unpack, same group order, same f32
13085        // accumulation per stream. Distinct x2 exercises both lanes.
13086        let xb: Vec<f32> = (0..cols)
13087            .map(|j| if j % 5 == 0 { -1.0 } else { 1.0 })
13088            .collect();
13089        let (mut s1, mut s2) = (vec![0f32; rows], vec![0f32; rows]);
13090        q1t_matvec(&bytes, &x, rows, cols, &mut s1, None);
13091        q1t_matvec(&bytes, &xb, rows, cols, &mut s2, None);
13092        let (mut p1, mut p2) = (vec![0f32; rows], vec![0f32; rows]);
13093        q1t_matvec2(&bytes, &x, &xb, rows, cols, &mut p1, &mut p2, None);
13094        assert_eq!(p1, s1, "q1t pair lane 1 ≠ single matvec");
13095        assert_eq!(p2, s2, "q1t pair lane 2 ≠ single matvec");
13096    }
13097
13098    /// Pair == 2×matvec with an ODD group count (the kernel's tail
13099    /// group) and no overlay section.
13100    #[test]
13101    fn q1t_matvec2_odd_gpr_matches_singles() {
13102        use cortiq_core::quant::{Q1T_TILE, f32_to_f16, q1t_pack};
13103        let (rows, cols) = (5usize, 96usize); // gpr = 3 → paired + tail
13104        let gpr = cols / GROUP_SIZE;
13105        let mut bytes = Vec::with_capacity(rows * gpr * Q1T_TILE);
13106        for r in 0..rows {
13107            for g in 0..gpr {
13108                bytes.extend_from_slice(&f32_to_f16(0.1 + 0.05 * (r + g) as f32).to_le_bytes());
13109                let mut c = [0u8; 7];
13110                for k in 0..GROUP_SIZE {
13111                    q1t_pack(&mut c, k, ((k * 7 + r * 5 + g * 3) % 3) as u8);
13112                }
13113                bytes.extend_from_slice(&c);
13114            }
13115        }
13116        let x1: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.31).sin()).collect();
13117        let x2: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.17).cos()).collect();
13118        let (mut s1, mut s2) = (vec![0f32; rows], vec![0f32; rows]);
13119        q1t_matvec(&bytes, &x1, rows, cols, &mut s1, None);
13120        q1t_matvec(&bytes, &x2, rows, cols, &mut s2, None);
13121        let (mut p1, mut p2) = (vec![0f32; rows], vec![0f32; rows]);
13122        q1t_matvec2(&bytes, &x1, &x2, rows, cols, &mut p1, &mut p2, None);
13123        assert_eq!(p1, s1, "odd-gpr pair lane 1 ≠ single");
13124        assert_eq!(p2, s2, "odd-gpr pair lane 2 ≠ single");
13125    }
13126
13127    // Speed A/B: fused pair (one unpack, two streams) vs two single
13128    // matvecs. Single-threaded, FFN-sized, min-of paired in-process.
13129    //   cargo test -p cortiq-engine --release q1t_matvec2_speed -- --ignored --nocapture
13130    #[test]
13131    #[ignore]
13132    fn q1t_matvec2_speed() {
13133        use cortiq_core::quant::{Q1T_TILE, f32_to_f16, q1t_pack};
13134        use std::time::Instant;
13135        let (rows, cols) = (8192usize, 4096usize);
13136        let gpr = cols / GROUP_SIZE;
13137        let mut bytes = Vec::with_capacity(rows * gpr * Q1T_TILE);
13138        for r in 0..rows {
13139            for g in 0..gpr {
13140                let s = 0.1 + ((r + g) % 7) as f32 * 0.01;
13141                bytes.extend_from_slice(&f32_to_f16(s).to_le_bytes());
13142                let mut c = [0u8; 7];
13143                for k in 0..GROUP_SIZE {
13144                    q1t_pack(&mut c, k, ((k * 7 + r + g) % 3) as u8);
13145                }
13146                bytes.extend_from_slice(&c);
13147            }
13148        }
13149        let x1: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.31).sin()).collect();
13150        let x2: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.17).cos()).collect();
13151        let (mut s1, mut s2) = (vec![0f32; rows], vec![0f32; rows]);
13152        let (mut p1, mut p2) = (vec![0f32; rows], vec![0f32; rows]);
13153        // Warm both paths once.
13154        q1t_matvec(&bytes, &x1, rows, cols, &mut s1, None);
13155        q1t_matvec2(&bytes, &x1, &x2, rows, cols, &mut p1, &mut p2, None);
13156        let (mut t_pair, mut t_two) = (f64::MAX, f64::MAX);
13157        for _ in 0..8 {
13158            let t0 = Instant::now();
13159            q1t_matvec2(&bytes, &x1, &x2, rows, cols, &mut p1, &mut p2, None);
13160            t_pair = t_pair.min(t0.elapsed().as_secs_f64() * 1000.0);
13161            let t1 = Instant::now();
13162            q1t_matvec(&bytes, &x1, rows, cols, &mut s1, None);
13163            q1t_matvec(&bytes, &x2, rows, cols, &mut s2, None);
13164            t_two = t_two.min(t1.elapsed().as_secs_f64() * 1000.0);
13165        }
13166        assert_eq!(p1, s1);
13167        assert_eq!(p2, s2);
13168        println!("q1t pair {rows}x{cols}: fused {t_pair:.2} ms | two singles {t_two:.2} ms");
13169    }
13170
13171    // Speed A/B: the base-3-division decode (what the packing commit left in
13172    // place) vs the fused sign-LUT matvec. Both single-threaded, same bytes.
13173    //   cargo test -p cortiq-engine q1t_matvec_speed -- --ignored --nocapture
13174    #[test]
13175    #[ignore]
13176    fn q1t_matvec_speed() {
13177        use cortiq_core::quant::{Q1T_TILE, f32_to_f16, q1t_code, q1t_pack};
13178        use std::time::Instant;
13179        let (rows, cols) = (8192usize, 4096usize); // FFN-sized
13180        let gpr = cols / GROUP_SIZE;
13181        let mut bytes = Vec::with_capacity(rows * gpr * Q1T_TILE + 16);
13182        for r in 0..rows {
13183            for g in 0..gpr {
13184                let s = 0.1 + ((r + g) % 7) as f32 * 0.01;
13185                bytes.extend_from_slice(&f32_to_f16(s).to_le_bytes());
13186                let mut c = [0u8; 7];
13187                for k in 0..GROUP_SIZE {
13188                    q1t_pack(&mut c, k, ((k * 7 + r + g) % 3) as u8);
13189                }
13190                bytes.extend_from_slice(&c);
13191            }
13192        }
13193        let (n, stride) = (rows * cols, 40usize); // ~2.5% outliers, per-row overlay
13194        let mut row_ptr = vec![0u32; rows + 1];
13195        let mut idx = 0usize;
13196        while idx < n {
13197            row_ptr[idx / cols + 1] += 1;
13198            idx += stride;
13199        }
13200        for r in 0..rows {
13201            row_ptr[r + 1] += row_ptr[r];
13202        }
13203        for &p in &row_ptr {
13204            bytes.extend_from_slice(&p.to_le_bytes());
13205        }
13206        let mut idx = 0usize;
13207        while idx < n {
13208            bytes.extend_from_slice(&((idx % cols) as u16).to_le_bytes());
13209            bytes.extend_from_slice(&f32_to_f16((idx % 13) as f32 * 0.1 - 0.6).to_le_bytes());
13210            idx += stride;
13211        }
13212        // On-grid ±1 so the fast path's int8 SDOT is exact vs the f32 "slow"
13213        // reference (the A/B is a timing check; values must still agree).
13214        let x: Vec<f32> = (0..cols)
13215            .map(|j| if j % 3 == 0 { 1.0 } else { -1.0 })
13216            .collect();
13217        let (rp_off, ent_off, has_ov) = q1t_overlay(&bytes, rows * gpr * Q1T_TILE, rows);
13218
13219        // "before": base-3 division decode into a buffer, then dot.
13220        let slow = |out: &mut [f32]| {
13221            let mut buf = vec![0f32; cols];
13222            for r in 0..rows {
13223                for g in 0..gpr {
13224                    let off = (r * gpr + g) * Q1T_TILE;
13225                    let s = f16_to_f32(u16::from_le_bytes([bytes[off], bytes[off + 1]]));
13226                    let codes = &bytes[off + 2..off + Q1T_TILE];
13227                    for k in 0..GROUP_SIZE {
13228                        buf[g * GROUP_SIZE + k] = match q1t_code(codes, k) {
13229                            1 => s,
13230                            2 => -s,
13231                            _ => 0.0,
13232                        };
13233                    }
13234                }
13235                out[r] = q1t_row_outlier_correction(&bytes, r, rp_off, ent_off, has_ov, &x)
13236                    + (0..cols).map(|j| buf[j] * x[j]).sum::<f32>();
13237            }
13238        };
13239        let iters = 5;
13240        let mut a = vec![0f32; rows];
13241        slow(&mut a); // warm
13242        let t = Instant::now();
13243        for _ in 0..iters {
13244            slow(&mut a);
13245        }
13246        let slow_ms = t.elapsed().as_secs_f64() * 1e3 / iters as f64;
13247
13248        let mut b = vec![0f32; rows];
13249        q1t_matvec(&bytes, &x, rows, cols, &mut b, None); // warm
13250        let t = Instant::now();
13251        for _ in 0..iters {
13252            q1t_matvec(&bytes, &x, rows, cols, &mut b, None);
13253        }
13254        let fast_ms = t.elapsed().as_secs_f64() * 1e3 / iters as f64;
13255
13256        for r in 0..rows {
13257            assert!((a[r] - b[r]).abs() < 1e-2, "mismatch row {r}");
13258        }
13259        println!(
13260            "q1t matvec {rows}x{cols} (1 thread): div-decode {slow_ms:.2} ms  fused-LUT {fast_ms:.2} ms  => {:.2}x",
13261            slow_ms / fast_ms
13262        );
13263    }
13264}
13265
13266#[cfg(test)]
13267mod gemm_bench {
13268    /// `cargo test -p cortiq-engine --release q4tp_matmat_throughput -- --ignored --nocapture`
13269    /// Times the batched q4tp GEMM at the shapes the image DiT runs
13270    /// (b=296 tokens, 2304 -> 9216), on synthetic bytes: no model, no
13271    /// mmap, no thermal drift over minutes — a kernel change shows up
13272    /// here in seconds where a full render hides it in noise.
13273    ///
13274    /// On macOS add `CMF_ACCEL=0`: this shape is over the 500k-cell mark
13275    /// where the matmat hands off to Accelerate's dequant sgemm, and
13276    /// without the opt-out both rows below measure the AMX, not the
13277    /// kernel under test.
13278    #[test]
13279    #[ignore]
13280    fn q4tp_matmat_throughput() {
13281        let _alt = super::Q4TP_ALT_TEST_LOCK
13282            .lock()
13283            .unwrap_or_else(|e| e.into_inner());
13284        // 296 is a prompt-encode batch; the image DiT runs 2085 at
13285        // 512x512, where the activation panel stops fitting L2 and the
13286        // loop's shape starts to matter more than its instructions.
13287        let b: usize = std::env::var("CMF_BENCH_B")
13288            .ok()
13289            .and_then(|v| v.parse().ok())
13290            .unwrap_or(296);
13291        let (rows, cols) = (9216usize, 2304usize);
13292        let (_, _, _) = (rows, cols, b);
13293        let total =
13294            cortiq_core::quant::expected_nbytes(cortiq_core::TensorDtype::Q4TiledP, &[rows, cols])
13295                .unwrap();
13296        // Random nibbles are fine, but the row params are f16 (lo, step)
13297        // of a geometric ladder: garbage there gives exp2 of a huge
13298        // exponent, the scales come back inf, and the whole bench times
13299        // NaN arithmetic instead of the kernel.
13300        let (params_off, codes_off, _) = cortiq_core::quant::q4tp_sections(rows, cols);
13301        let mut bytes: Vec<u8> = (0..total).map(|i| (i * 37 % 251) as u8).collect();
13302        let lo = cortiq_core::quant::f32_to_f16(-4.0);
13303        let step = cortiq_core::quant::f32_to_f16(0.1);
13304        for r in 0..rows {
13305            let o = params_off + r * 4;
13306            bytes[o..o + 2].copy_from_slice(&lo.to_le_bytes());
13307            bytes[o + 2..o + 4].copy_from_slice(&step.to_le_bytes());
13308        }
13309        let _ = codes_off;
13310        let xs: Vec<f32> = (0..b * cols)
13311            .map(|i| ((i % 97) as f32 - 48.0) / 48.0)
13312            .collect();
13313        let mut out = vec![0f32; b * rows];
13314        let pool = crate::pool::Pool::from_env();
13315        // A shared 48-core stand drifts ±25% run to run, which is wider
13316        // than any kernel change worth making. So: alternate the two
13317        // kernels inside one process and keep the BEST time for
13318        // each. Interleaving makes both see the same interference, and a
13319        // minimum is the one statistic another tenant cannot inflate.
13320        super::q4tp_matmat(&bytes, &xs, b, rows, cols, &mut out, pool.as_deref());
13321        let reps: usize = std::env::var("CMF_BENCH_REPS")
13322            .ok()
13323            .and_then(|v| v.parse().ok())
13324            .unwrap_or(10);
13325        let mut best = [f64::MAX; 2];
13326        let mut sums = [0f32; 2];
13327        for _ in 0..reps {
13328            for (k, w) in [(0usize, 1u8), (1usize, 2u8)] {
13329                super::Q4TP_ALT.store(w, std::sync::atomic::Ordering::Relaxed);
13330                let t = std::time::Instant::now();
13331                super::q4tp_matmat(&bytes, &xs, b, rows, cols, &mut out, pool.as_deref());
13332                best[k] = best[k].min(t.elapsed().as_secs_f64());
13333                sums[k] = out.iter().take(64).sum::<f32>();
13334            }
13335        }
13336        let flops = 2.0 * b as f64 * rows as f64 * cols as f64;
13337        for (k, name) in ["previous", "tuned   "].iter().enumerate() {
13338            println!(
13339                "q4tp matmat {rows}x{cols} b={b} {name}: {:.1} ms  {:.1} GFLOP/s  (checksum {:.3})",
13340                best[k] * 1e3,
13341                flops / best[k] / 1e9,
13342                sums[k]
13343            );
13344        }
13345        assert!(
13346            (sums[0] - sums[1]).abs() < 1e-2,
13347            "the tuned kernel changed the result: {} vs {}",
13348            sums[0],
13349            sums[1]
13350        );
13351    }
13352
13353    /// The blocked kernel must agree with the per-column path exactly —
13354    /// same weights, same activation split, only a different instruction
13355    /// mix. Shapes are chosen to hit the awkward cases: a column count
13356    /// that leaves an odd group (the 512-bit kernel does two at a time),
13357    /// and a batch that does not divide by four.
13358    #[test]
13359    fn q4tp_matmat_blocked_matches_scalar() {
13360        use std::sync::atomic::Ordering::Relaxed;
13361        let _alt = super::Q4TP_ALT_TEST_LOCK
13362            .lock()
13363            .unwrap_or_else(|e| e.into_inner());
13364        // The last shape carries the image DiT's column count — 2304, so
13365        // 72 groups of accumulation, which is where a reordered sum can
13366        // actually drift — and runs through the thread pool, since the
13367        // blocked path splits rows across workers. Its row count stays
13368        // under 500k cells on purpose: above that, macOS diverts the whole
13369        // matmat to the Accelerate/AMX dequant sgemm and neither kernel
13370        // here would run.
13371        for &(rows, cols, b) in &[
13372            (64usize, 128usize, 7usize),
13373            (33, 96, 4),
13374            (16, 256, 9),
13375            (192, 2304, 37),
13376        ] {
13377            let total = cortiq_core::quant::expected_nbytes(
13378                cortiq_core::TensorDtype::Q4TiledP,
13379                &[rows, cols],
13380            )
13381            .unwrap();
13382            let (params_off, _, _) = cortiq_core::quant::q4tp_sections(rows, cols);
13383            let mut bytes: Vec<u8> = (0..total).map(|i| (i * 61 % 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 xs: Vec<f32> = (0..b * cols)
13392                .map(|i| ((i % 89) as f32 - 44.0) / 44.0)
13393                .collect();
13394            let mut got = vec![0f32; b * rows];
13395            let mut want = vec![0f32; b * rows];
13396            let gpr = cols / 32;
13397            let view = super::Q4tpView::new(&bytes, rows, cols);
13398            let pool = crate::pool::Pool::from_env();
13399            super::Q4TP_ALT.store(2, Relaxed);
13400            super::q4tp_matmat(&bytes, &xs, b, rows, cols, &mut got, pool.as_deref());
13401            super::Q4TP_ALT.store(1, Relaxed);
13402            super::q4tp_matmat(&bytes, &xs, b, rows, cols, &mut want, pool.as_deref());
13403            super::Q4TP_ALT.store(0, Relaxed);
13404            // Measured against the output's scale, not cell by cell: a
13405            // dot product of 2304 terms lands near zero wherever the row
13406            // and the activation nearly cancel, and there a per-cell
13407            // ratio reports 1e-3 for an absolute error of 5e-6 — f32's
13408            // own rounding, reordered. What must stay small is the error
13409            // relative to what the layer actually outputs.
13410            let scale = want.iter().fold(0f32, |m, v| m.max(v.abs())).max(1e-6);
13411            let (mut worst, mut at) = (0f32, 0usize);
13412            for (i, (g, w)) in got.iter().zip(&want).enumerate() {
13413                if (g - w).abs() > worst {
13414                    worst = (g - w).abs();
13415                    at = i;
13416                }
13417            }
13418            assert!(
13419                worst <= 1e-4 * scale,
13420                "{rows}x{cols} b={b}: blocked and scalar disagree by {worst:.3e} \
13421                 (scale {scale:.3e}) at cell {at}: {} vs {}",
13422                got[at],
13423                want[at]
13424            );
13425
13426            // "Same speed, no quality loss" is a claim about which answer
13427            // is RIGHT, not about which two agree. Both paths sum the same
13428            // 2304 products in different orders, so f64 decides: the
13429            // blocked kernel keeps sixteen partial sums and folds them at
13430            // the end, which is a shallower addition tree than the
13431            // per-column path's running scalar, and it must not be worse.
13432            let (mut e_blocked, mut e_scalar) = (0f64, 0f64);
13433            for bi in 0..b {
13434                let act = super::split_act(&xs[bi * cols..(bi + 1) * cols]);
13435                for r in 0..rows {
13436                    let mut sc = vec![0f32; gpr];
13437                    view.scales_into(r, gpr, &mut sc);
13438                    let mut exact = 0f64;
13439                    for j in 0..cols {
13440                        let (w, sq) = super::q4tp_outlier(view.nib, r, gpr, j, &sc);
13441                        exact += w as f64 * sq as f64 * act.xq[j] as f64;
13442                    }
13443                    exact *= act.sx as f64;
13444                    for &(j, xv) in &act.outliers {
13445                        let (w, sq) = super::q4tp_outlier(view.nib, r, gpr, j, &sc);
13446                        exact += w as f64 * sq as f64 * xv as f64;
13447                    }
13448                    let i = bi * rows + r;
13449                    e_blocked = e_blocked.max((got[i] as f64 - exact).abs());
13450                    e_scalar = e_scalar.max((want[i] as f64 - exact).abs());
13451                }
13452            }
13453            println!(
13454                "{rows}x{cols} b={b}: worst error vs f64 — blocked {e_blocked:.3e}, \
13455                 per-column {e_scalar:.3e}"
13456            );
13457            // An absolute bar, not a race between the two: at these
13458            // magnitudes both sit in f32's last bits, and on a small shape
13459            // whichever one happens to round the unluckiest cell "wins" by
13460            // a factor the next seed reverses.
13461            assert!(
13462                e_blocked <= 1e-5 * scale as f64 && e_scalar <= 1e-5 * scale as f64,
13463                "{rows}x{cols} b={b}: error against f64 too large — blocked \
13464                 {e_blocked:.3e}, per-column {e_scalar:.3e}, scale {scale:.3e}"
13465            );
13466        }
13467    }
13468
13469    /// What the row-exact contract costs a speculative-verify panel on the
13470    /// host: the same batch through the fast arms, through the row-exact
13471    /// arms, and as one matvec per token (the other way to be exact).
13472    /// Arms alternate inside one process and keep their best time.
13473    /// `CMF_GPU=0 cargo test -p cortiq-engine --release q4tp_matmat_row_exact_cost -- --ignored --nocapture`
13474    #[test]
13475    #[ignore]
13476    fn q4tp_matmat_row_exact_cost() {
13477        let _alt = super::Q4TP_ALT_TEST_LOCK
13478            .lock()
13479            .unwrap_or_else(|e| e.into_inner());
13480        let pool = crate::pool::Pool::from_env();
13481        let reps: usize = std::env::var("CMF_BENCH_REPS")
13482            .ok()
13483            .and_then(|v| v.parse().ok())
13484            .unwrap_or(30);
13485        for &(rows, cols) in &[(2048usize, 4096usize), (4096, 2048), (4096, 4096)] {
13486            let total = cortiq_core::quant::expected_nbytes(
13487                cortiq_core::TensorDtype::Q4TiledP,
13488                &[rows, cols],
13489            )
13490            .unwrap();
13491            let (params_off, _, _) = cortiq_core::quant::q4tp_sections(rows, cols);
13492            let mut bytes: Vec<u8> = (0..total).map(|i| (i * 37 % 251) as u8).collect();
13493            let lo = cortiq_core::quant::f32_to_f16(-4.0);
13494            let step = cortiq_core::quant::f32_to_f16(0.1);
13495            for r in 0..rows {
13496                let o = params_off + r * 4;
13497                bytes[o..o + 2].copy_from_slice(&lo.to_le_bytes());
13498                bytes[o + 2..o + 4].copy_from_slice(&step.to_le_bytes());
13499            }
13500            for &b in &[2usize, 4, 5, 8] {
13501                let xs: Vec<f32> = (0..b * cols)
13502                    .map(|i| ((i % 97) as f32 - 48.0) / 48.0)
13503                    .collect();
13504                let mut out = vec![0f32; b * rows];
13505                let mut best = [f64::MAX; 3];
13506                for _ in 0..reps {
13507                    for (k, best_k) in best.iter_mut().enumerate() {
13508                        let t = std::time::Instant::now();
13509                        match k {
13510                            0 | 1 => super::q4tp_matmat_with(
13511                                &bytes,
13512                                &xs,
13513                                b,
13514                                rows,
13515                                cols,
13516                                &mut out,
13517                                pool.as_deref(),
13518                                k == 1,
13519                            ),
13520                            _ => {
13521                                for (bi, o) in out.chunks_mut(rows).enumerate() {
13522                                    super::q4tp_matvec(
13523                                        &bytes,
13524                                        &xs[bi * cols..(bi + 1) * cols],
13525                                        rows,
13526                                        cols,
13527                                        o,
13528                                        pool.as_deref(),
13529                                    );
13530                                }
13531                            }
13532                        }
13533                        *best_k = best_k.min(t.elapsed().as_secs_f64());
13534                    }
13535                }
13536                println!(
13537                    "q4tp {rows}x{cols} b={b}: fast {:.3} ms, row-exact {:.3} ms, \
13538                     {b} matvecs {:.3} ms",
13539                    best[0] * 1e3,
13540                    best[1] * 1e3,
13541                    best[2] * 1e3
13542                );
13543            }
13544        }
13545    }
13546
13547    /// The q4t twin of the throughput bench, same shape and rules, so the
13548    /// two quantisations' batch kernels can be read against each other.
13549    /// `cargo test -p cortiq-engine --release q4t_matmat_throughput -- --ignored --nocapture`
13550    #[test]
13551    #[ignore]
13552    fn q4t_matmat_throughput() {
13553        let (rows, cols, b) = (9216usize, 2304usize, 296usize);
13554        let total =
13555            cortiq_core::quant::expected_nbytes(cortiq_core::TensorDtype::Q4Tiled, &[rows, cols])
13556                .unwrap();
13557        // q4t carries a per-group f16 scale in the tile's first two bytes;
13558        // random bytes there decode to inf and the bench would time NaNs.
13559        let mut bytes: Vec<u8> = (0..total).map(|i| (i * 37 % 251) as u8).collect();
13560        let sc = cortiq_core::quant::f32_to_f16(0.02);
13561        for t in bytes.chunks_mut(super::Q4_TILE) {
13562            t[..2].copy_from_slice(&sc.to_le_bytes());
13563        }
13564        let xs: Vec<f32> = (0..b * cols)
13565            .map(|i| ((i % 97) as f32 - 48.0) / 48.0)
13566            .collect();
13567        let mut out = vec![0f32; b * rows];
13568        let pool = crate::pool::Pool::from_env();
13569        super::q4t_matmat(&bytes, &xs, b, rows, cols, &mut out, pool.as_deref());
13570        let reps: usize = std::env::var("CMF_BENCH_REPS")
13571            .ok()
13572            .and_then(|v| v.parse().ok())
13573            .unwrap_or(10);
13574        let mut best = f64::MAX;
13575        for _ in 0..reps {
13576            let t = std::time::Instant::now();
13577            super::q4t_matmat(&bytes, &xs, b, rows, cols, &mut out, pool.as_deref());
13578            best = best.min(t.elapsed().as_secs_f64());
13579        }
13580        let flops = 2.0 * b as f64 * rows as f64 * cols as f64;
13581        println!(
13582            "q4t matmat {rows}x{cols} b={b}: {:.1} ms  {:.1} GFLOP/s  (checksum {:.3})",
13583            best * 1e3,
13584            flops / best / 1e9,
13585            out.iter().take(64).sum::<f32>()
13586        );
13587    }
13588}