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::{GROUP_SIZE, Q1_TILE, Q4_TILE, f16_to_f32};
16use cortiq_core::{CmfModel, TensorDtype};
17use std::sync::Arc;
18
19pub enum QTensor {
20    F32 {
21        data: Vec<f32>,
22        rows: usize,
23        cols: usize,
24    },
25    Mapped {
26        model: Arc<CmfModel>,
27        /// Index into the model's tensor directory.
28        idx: usize,
29        dtype: TensorDtype,
30        rows: usize,
31        cols: usize,
32        /// Per-row scales, dequantized to f32 up front (tiny).
33        row_scale: Vec<f32>,
34        /// q8_2f column field (θ), dequantized up front; empty for q8_row.
35        col_field: Vec<f32>,
36        /// Vbit only: byte offset of each row's packed data within the
37        /// tensor blob (`[rows + 1]`, computed once at load — the per-
38        /// matvec prefix scan over row bit-widths was O(rows) each call).
39        vbit_offsets: Vec<usize>,
40        /// q8-family decode repack (load-time, optional): rows in groups
41        /// of 4, interleaved in 16-byte units — one 64-byte line per
42        /// iteration feeds all 4 sdot lanes, ONE sequential weight
43        /// stream per worker instead of four (this is where llama.cpp's
44        /// repacked Q8 kernels get their bandwidth). Empty = off
45        /// (CMF_REPACK=0, non-SDOT arch, or an ineligible shape). Trades
46        /// an anonymous copy of the quants for mmap pages that go cold.
47        repack: Vec<u8>,
48    },
49}
50
51/// Load-time q8 repack gate (see `Mapped::repack`). OPT-IN
52/// (`CMF_REPACK=1`): the single-stream hypothesis LOST on Apple Silicon
53/// (M4, interleaved A/B: decode 101 vs 94 tok/s — four adjacent row
54/// streams per worker feed the prefetcher MORE memory-level parallelism
55/// than one); kept as an experiment flag for x86, where the tradeoff
56/// may land differently.
57fn repack_enabled() -> bool {
58    static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
59    *ON.get_or_init(|| {
60        std::env::var("CMF_REPACK")
61            .map(|v| v == "1")
62            .unwrap_or(cfg!(target_os = "android"))
63    })
64}
65
66/// Interleave q8 rows for the decode kernel: group g holds rows
67/// 4g..4g+4 as [r0[c], r1[c], r2[c], r3[c]] per 16-byte chunk c. Only
68/// full groups are packed — tail rows keep reading the mmap layout.
69fn q8_repack(bytes: &[u8], rows: usize, cols: usize) -> Vec<u8> {
70    #[cfg(target_arch = "aarch64")]
71    let arch_ok = sdot_enabled();
72    #[cfg(not(target_arch = "aarch64"))]
73    let arch_ok = false;
74    if !arch_ok || !repack_enabled() || rows < 256 || cols % 16 != 0 {
75        return Vec::new();
76    }
77    q8_repack_layout(bytes, rows, cols)
78}
79
80/// The pure layout transform behind `q8_repack` (tested directly —
81/// the gate depends on arch and env).
82fn q8_repack_layout(bytes: &[u8], rows: usize, cols: usize) -> Vec<u8> {
83    let groups = rows / 4;
84    let mut rep = vec![0u8; groups * 4 * cols];
85    for g in 0..groups {
86        let dst = &mut rep[g * 4 * cols..(g + 1) * 4 * cols];
87        for c in 0..cols / 16 {
88            for lane in 0..4 {
89                let src = (g * 4 + lane) * cols + c * 16;
90                dst[c * 64 + lane * 16..c * 64 + lane * 16 + 16]
91                    .copy_from_slice(&bytes[src..src + 16]);
92            }
93        }
94    }
95    rep
96}
97
98/// Prefix-sum of vbit row payload offsets (absolute within the tensor
99/// bytes). `offsets[r]..offsets[r+1]` is row r's packed data.
100fn vbit_row_offsets(bytes: &[u8], rows: usize, cols: usize) -> Vec<usize> {
101    let ng = cols / GROUP_SIZE;
102    let bits = &bytes[..rows];
103    let mut offsets = Vec::with_capacity(rows + 1);
104    let mut off = rows + rows * ng * 2;
105    for r in 0..rows {
106        offsets.push(off);
107        off += (cols * bits[r] as usize).div_ceil(8);
108    }
109    offsets.push(off);
110    offsets
111}
112
113impl QTensor {
114    pub fn from_f32(data: Vec<f32>, rows: usize, cols: usize) -> Self {
115        debug_assert_eq!(data.len(), rows * cols);
116        Self::F32 { data, rows, cols }
117    }
118
119    /// Wrap a directory tensor without dequantizing the payload.
120    /// Falls back to dequantized f32 for dtypes without a fused kernel.
121    pub fn from_model(model: &Arc<CmfModel>, name: &str) -> Result<Self, String> {
122        // Indexed lookup: the linear directory scan made pipeline build
123        // O(N²) on MoE/skills files with thousands of tensors.
124        let idx = model
125            .tensor_index(name)
126            .ok_or_else(|| format!("tensor '{name}' not found in CMF directory"))?;
127        let entry = &model.tensors[idx];
128        if entry.shape.len() != 2 {
129            return Err(format!("QTensor::from_model needs 2-D, got '{name}'"));
130        }
131        let (rows, cols) = (entry.shape[0], entry.shape[1]);
132        let bytes = model.entry_bytes(entry);
133
134        match entry.dtype {
135            TensorDtype::Q8Row | TensorDtype::Q8_2f => {
136                let n = rows * cols;
137                let scales_off = n;
138                let row_scale: Vec<f32> = (0..rows)
139                    .map(|o| {
140                        f16_to_f32(u16::from_le_bytes([
141                            bytes[scales_off + o * 2],
142                            bytes[scales_off + o * 2 + 1],
143                        ]))
144                    })
145                    .collect();
146                let col_field: Vec<f32> = if entry.dtype == TensorDtype::Q8_2f {
147                    let col_off = n + rows * 2;
148                    (0..cols)
149                        .map(|i| {
150                            f16_to_f32(u16::from_le_bytes([
151                                bytes[col_off + i * 2],
152                                bytes[col_off + i * 2 + 1],
153                            ]))
154                        })
155                        .collect()
156                } else {
157                    Vec::new()
158                };
159                Ok(Self::Mapped {
160                    model: model.clone(),
161                    idx,
162                    dtype: entry.dtype,
163                    rows,
164                    cols,
165                    row_scale,
166                    col_field,
167                    vbit_offsets: Vec::new(),
168                    repack: q8_repack(bytes, rows, cols),
169                })
170            }
171            // vbit: fused kernel unpacks variable-bit rows from mmap.
172            TensorDtype::Vbit if cols % GROUP_SIZE == 0 => Ok(Self::Mapped {
173                model: model.clone(),
174                idx,
175                dtype: entry.dtype,
176                rows,
177                cols,
178                row_scale: Vec::new(),
179                col_field: Vec::new(),
180                vbit_offsets: vbit_row_offsets(bytes, rows, cols),
181                repack: Vec::new(),
182            }),
183            // vbit_ro (§4.2): the offset table comes straight from the
184            // file — no load-time prefix scan; kernels are shared with
185            // legacy vbit (they consume absolute offsets either way).
186            TensorDtype::VbitRo if cols % GROUP_SIZE == 0 => {
187                let (_, off_off, packed_off) = cortiq_core::quant::vbit_ro_sections(rows, cols);
188                let offsets: Vec<usize> = (0..=rows)
189                    .map(|r| packed_off + cortiq_core::quant::vbit_ro_offset(bytes, off_off, r))
190                    .collect();
191                Ok(Self::Mapped {
192                    model: model.clone(),
193                    idx,
194                    dtype: entry.dtype,
195                    rows,
196                    cols,
197                    row_scale: Vec::new(),
198                    col_field: Vec::new(),
199                    vbit_offsets: offsets,
200                    repack: Vec::new(),
201                })
202            }
203            // q4_block: fused kernel reads nibbles straight from mmap —
204            // a 14B q4 file no longer explodes into ×8 f32 RAM.
205            // q4_tiled (§4.3): interleaved [scale][nibbles] tiles — one
206            // sequential memory stream (measured ×1.66 ARM / ×1.13 AVX2
207            // at kernel level over the split layout).
208            TensorDtype::Q4Tiled if cols % GROUP_SIZE == 0 => Ok(Self::Mapped {
209                model: model.clone(),
210                idx,
211                dtype: entry.dtype,
212                rows,
213                cols,
214                row_scale: Vec::new(),
215                col_field: Vec::new(),
216                vbit_offsets: Vec::new(),
217                repack: Vec::new(),
218            }),
219            TensorDtype::Q4Block if cols % GROUP_SIZE == 0 => Ok(Self::Mapped {
220                model: model.clone(),
221                idx,
222                dtype: entry.dtype,
223                rows,
224                cols,
225                row_scale: Vec::new(),
226                col_field: Vec::new(),
227                vbit_offsets: Vec::new(),
228                repack: Vec::new(),
229            }),
230            // q1: binary sign-bit tiles from mmap (1-bit-trained models).
231            TensorDtype::Q1 if cols % GROUP_SIZE == 0 => Ok(Self::Mapped {
232                model: model.clone(),
233                idx,
234                dtype: entry.dtype,
235                rows,
236                cols,
237                row_scale: Vec::new(),
238                col_field: Vec::new(),
239                vbit_offsets: Vec::new(),
240                repack: Vec::new(),
241            }),
242            // q1t (ternary + outlier overlay): fused per-row dequant kernel
243            // reads straight from mmap — a 12B q1t stays ~its file size in
244            // RAM instead of dequantizing to ~48 GB of f32.
245            TensorDtype::Q1T if cols % GROUP_SIZE == 0 => Ok(Self::Mapped {
246                model: model.clone(),
247                idx,
248                dtype: entry.dtype,
249                rows,
250                cols,
251                row_scale: Vec::new(),
252                col_field: Vec::new(),
253                vbit_offsets: Vec::new(),
254                repack: Vec::new(),
255            }),
256            // No fused kernel yet → dequantize once (correct, more RAM).
257            _ => {
258                let mut data = vec![0.0f32; rows * cols];
259                cortiq_core::quant::dequant_tensor(entry, bytes, &mut data)?;
260                Ok(Self::from_f32(data, rows, cols))
261            }
262        }
263    }
264
265    /// q1-mapped tensor? (GPU gates: the q1 CPU kernel is
266    /// compute-bound, so offload pays at much smaller shapes than q8.)
267    pub(crate) fn is_q1(&self) -> bool {
268        matches!(
269            self,
270            Self::Mapped {
271                dtype: TensorDtype::Q1,
272                ..
273            }
274        )
275    }
276
277    /// Owned-f32 view (data, rows, cols) — the GDN a/b gate projections
278    /// arrive dequantized (force-f16 in the converter → F32 in RAM).
279    pub(crate) fn f32_parts(&self) -> Option<(&[f32], usize, usize)> {
280        match self {
281            Self::F32 { data, rows, cols } => Some((data, *rows, *cols)),
282            _ => None,
283        }
284    }
285
286    /// (directory idx, rows, cols) of a q1-mapped tensor — the
287    /// whole-block GPU path resolves offsets itself.
288    /// (idx, rows, cols) of a mapped tensor the whole-token GPU graph can drive
289    /// — Q1, Q1T or Q4-block (it resolves the offset and picks the kernel by
290    /// dtype). Q4-block lets a precise down_proj/lm_head stay on-device.
291    /// Named `q1_parts` for historical reasons.
292    pub(crate) fn q1_parts(&self) -> Option<(usize, usize, usize)> {
293        match self {
294            #[cfg(target_os = "macos")]
295            Self::Mapped {
296                dtype: TensorDtype::Q1T,
297                ..
298            } if !crate::gpu::metal_q1t_enabled() => None,
299            Self::Mapped {
300                idx,
301                dtype:
302                    TensorDtype::Q1
303                    | TensorDtype::Q1T
304                    | TensorDtype::Q4Block
305                    | TensorDtype::Q4Tiled
306                    | TensorDtype::Q8Row
307                    | TensorDtype::Q8_2f,
308                rows,
309                cols,
310                ..
311            } => Some((*idx, *rows, *cols)),
312            _ => None,
313        }
314    }
315
316    /// (directory idx, rows, cols, row_scale) of a plain q8_row mapped
317    /// tensor — the chunk-prefill GPU graph resolves offsets itself.
318    /// q8_2f is excluded on purpose: its column field would need a
319    /// prescale stage on the device.
320    pub(crate) fn q8_row_parts(&self) -> Option<(usize, usize, usize, &[f32])> {
321        match self {
322            Self::Mapped {
323                idx,
324                dtype: TensorDtype::Q8Row,
325                rows,
326                cols,
327                row_scale,
328                col_field,
329                ..
330            } if col_field.is_empty() => Some((*idx, *rows, *cols, row_scale)),
331            _ => None,
332        }
333    }
334
335    pub fn rows(&self) -> usize {
336        match self {
337            Self::F32 { rows, .. } | Self::Mapped { rows, .. } => *rows,
338        }
339    }
340
341    pub fn cols(&self) -> usize {
342        match self {
343            Self::F32 { cols, .. } | Self::Mapped { cols, .. } => *cols,
344        }
345    }
346
347    /// (model, tensor idx) for a q1 mapped weight — the wgpu token graph
348    /// keys its resident VRAM cache by idx. None for any other dtype/kind.
349    pub fn mapped_q1(&self) -> Option<(&std::sync::Arc<CmfModel>, usize)> {
350        match self {
351            Self::Mapped {
352                model,
353                idx,
354                dtype: TensorDtype::Q1,
355                ..
356            } => Some((model, *idx)),
357            _ => None,
358        }
359    }
360
361    /// (model, idx, kind, row_scale) for a graph-capable mapped weight. kind:
362    /// 0=q8_row (per-row scales), 1=q1, 2=q4_tiled, 3=q1t (tile-embedded, no
363    /// rs). None for dtypes the token graph does not handle (q8_2f/q4_block/vbit).
364    pub fn graph_weight(&self) -> Option<(&std::sync::Arc<CmfModel>, usize, u8, &[f32])> {
365        match self {
366            Self::Mapped {
367                model,
368                idx,
369                dtype: TensorDtype::Q8Row,
370                row_scale,
371                ..
372            } => Some((model, *idx, 0, row_scale.as_slice())),
373            Self::Mapped {
374                model,
375                idx,
376                dtype: TensorDtype::Q1,
377                ..
378            } => Some((model, *idx, 1, &[])),
379            // Q4Tiled is kind 5, NOT 2: both carried 2 historically, and
380            // the wgpu token graph fed 18B interleaved tiles to the
381            // split-layout q4b kernel — garbage output on q4t models
382            // (caught by an end-to-end answer check on real Vulkan).
383            Self::Mapped {
384                model,
385                idx,
386                dtype: TensorDtype::Q4Tiled,
387                ..
388            } => Some((model, *idx, 5, &[])),
389            Self::Mapped {
390                model,
391                idx,
392                dtype: TensorDtype::Q4Block,
393                ..
394            } => Some((model, *idx, 2, &[])),
395            Self::Mapped {
396                model,
397                idx,
398                dtype: TensorDtype::Q1T,
399                ..
400            } => Some((model, *idx, 3, &[])),
401            _ => None,
402        }
403    }
404
405    /// Dense f32 view — only for owned tensors. Masked/sparse execution
406    /// paths require it; quantized weights don't support masks yet.
407    pub fn as_f32(&self) -> Option<&[f32]> {
408        match self {
409            Self::F32 { data, .. } => Some(data),
410            Self::Mapped { .. } => None,
411        }
412    }
413
414    fn quant_bytes(&self) -> &[u8] {
415        match self {
416            Self::Mapped { model, idx, .. } => model.entry_bytes(&model.tensors[*idx]),
417            Self::F32 { .. } => unreachable!("quant_bytes on F32"),
418        }
419    }
420
421    /// Dequantize one row into `dst` (embedding lookup).
422    pub fn row_f32(&self, r: usize, dst: &mut [f32]) {
423        let cols = self.cols();
424        debug_assert_eq!(dst.len(), cols);
425        match self {
426            Self::F32 { data, .. } => dst.copy_from_slice(&data[r * cols..(r + 1) * cols]),
427            Self::Mapped {
428                dtype,
429                row_scale,
430                col_field,
431                vbit_offsets,
432                ..
433            } => {
434                if *dtype == TensorDtype::Q4Tiled {
435                    let bytes = self.quant_bytes();
436                    let gpr = cols / GROUP_SIZE;
437                    for gi in 0..gpr {
438                        let tile = &bytes[(r * gpr + gi) * Q4_TILE..(r * gpr + gi + 1) * Q4_TILE];
439                        let s = f16_to_f32(u16::from_le_bytes([tile[0], tile[1]]));
440                        for (k, &b) in tile[2..].iter().enumerate() {
441                            dst[gi * GROUP_SIZE + k * 2] = ((b & 0x0F) as f32 - 8.0) * s;
442                            dst[gi * GROUP_SIZE + k * 2 + 1] = (((b >> 4) & 0x0F) as f32 - 8.0) * s;
443                        }
444                    }
445                    return;
446                }
447                if *dtype == TensorDtype::Q4Block {
448                    let (packed, scales) = q4_split(self.quant_bytes(), self.rows(), cols);
449                    let gpr = cols / GROUP_SIZE;
450                    for gi in 0..gpr {
451                        let g = r * gpr + gi;
452                        let s = f16_to_f32(u16::from_le_bytes([scales[g * 2], scales[g * 2 + 1]]));
453                        for (k, &b) in packed[g * 16..(g + 1) * 16].iter().enumerate() {
454                            dst[gi * GROUP_SIZE + k * 2] = ((b & 0x0F) as f32 - 8.0) * s;
455                            dst[gi * GROUP_SIZE + k * 2 + 1] = (((b >> 4) & 0x0F) as f32 - 8.0) * s;
456                        }
457                    }
458                    return;
459                }
460                if *dtype == TensorDtype::Q1 {
461                    let bytes = self.quant_bytes();
462                    let gpr = cols / GROUP_SIZE;
463                    for gi in 0..gpr {
464                        let tile = &bytes[(r * gpr + gi) * Q1_TILE..(r * gpr + gi + 1) * Q1_TILE];
465                        let s = f16_to_f32(u16::from_le_bytes([tile[0], tile[1]]));
466                        for (j, &b) in tile[2..].iter().enumerate() {
467                            for k in 0..8 {
468                                dst[gi * GROUP_SIZE + j * 8 + k] =
469                                    (((b >> k) & 1) as f32 * 2.0 - 1.0) * s;
470                            }
471                        }
472                    }
473                    return;
474                }
475                if *dtype == TensorDtype::Q1T {
476                    let bytes = self.quant_bytes();
477                    let gpr = cols / GROUP_SIZE;
478                    let base_len = self.rows() * gpr * cortiq_core::quant::Q1T_TILE;
479                    for gi in 0..gpr {
480                        let off = (r * gpr + gi) * cortiq_core::quant::Q1T_TILE;
481                        let s = cortiq_core::quant::f16_to_f32(u16::from_le_bytes([
482                            bytes[off],
483                            bytes[off + 1],
484                        ]));
485                        let codes = &bytes[off + 2..off + cortiq_core::quant::Q1T_TILE];
486                        for k in 0..GROUP_SIZE {
487                            dst[gi * GROUP_SIZE + k] = match cortiq_core::quant::q1t_code(codes, k)
488                            {
489                                1 => s,
490                                2 => -s,
491                                _ => 0.0,
492                            };
493                        }
494                    }
495                    // Overlay
496                    let rows = self.rows();
497                    let entries = base_len + (rows + 1) * 4;
498                    if entries <= bytes.len() {
499                        let ptrs = &bytes[base_len..base_len + (rows + 1) * 4];
500                        let r0 = u32::from_le_bytes([
501                            ptrs[r * 4],
502                            ptrs[r * 4 + 1],
503                            ptrs[r * 4 + 2],
504                            ptrs[r * 4 + 3],
505                        ]) as usize;
506                        let r1 = u32::from_le_bytes([
507                            ptrs[(r + 1) * 4],
508                            ptrs[(r + 1) * 4 + 1],
509                            ptrs[(r + 1) * 4 + 2],
510                            ptrs[(r + 1) * 4 + 3],
511                        ]) as usize;
512                        let off = entries + r0 * 4;
513                        for i in 0..r1 - r0 {
514                            let item = &bytes[off + i * 4..off + i * 4 + 4];
515                            let c = u16::from_le_bytes([item[0], item[1]]) as usize;
516                            let v = cortiq_core::quant::f16_to_f32(u16::from_le_bytes([
517                                item[2], item[3],
518                            ]));
519                            if c < cols {
520                                dst[c] = v;
521                            }
522                        }
523                    }
524                    return;
525                }
526                if matches!(dtype, TensorDtype::Vbit | TensorDtype::VbitRo) {
527                    let bytes = self.quant_bytes();
528                    let rows = self.rows();
529                    let ng = cols / GROUP_SIZE;
530                    let bits = &bytes[..rows];
531                    let sc_off = rows;
532                    // Precomputed at load — embedding lookup used to scan
533                    // the bit-widths of every preceding row (O(token_id)).
534                    let off = vbit_offsets[r];
535                    let b = bits[r] as usize;
536                    let l = ((1usize << (b - 1)) - 1) as f32;
537                    let data = &bytes[off..];
538                    let (mut acc, mut nbits, mut idx) = (0u64, 0usize, 0usize);
539                    for (i, d) in dst.iter_mut().enumerate() {
540                        while nbits < b {
541                            acc = (acc << 8) | data[idx] as u64;
542                            idx += 1;
543                            nbits += 8;
544                        }
545                        let u = ((acc >> (nbits - b)) & ((1u64 << b) - 1)) as f32;
546                        nbits -= b;
547                        let so = (r * ng + i / GROUP_SIZE) * 2;
548                        let sv = f16_to_f32(u16::from_le_bytes([
549                            bytes[sc_off + so],
550                            bytes[sc_off + so + 1],
551                        ]));
552                        *d = (u - l) * sv;
553                    }
554                    return;
555                }
556                let q = &self.quant_bytes()[r * cols..(r + 1) * cols];
557                let s = row_scale[r];
558                match dtype {
559                    TensorDtype::Q8Row => {
560                        for (d, &b) in dst.iter_mut().zip(q) {
561                            *d = (b as i8) as f32 * s;
562                        }
563                    }
564                    TensorDtype::Q8_2f => {
565                        for (i, (d, &b)) in dst.iter_mut().zip(q).enumerate() {
566                            *d = (b as i8) as f32 * s * col_field[i];
567                        }
568                    }
569                    _ => unreachable!(),
570                }
571            }
572        }
573    }
574
575    /// Can this tensor's columns be read cheaply (for sparse down_proj)?
576    /// True for F32/Q8Row/Q8_2f (per-row scale, direct strided access);
577    /// false for group-packed q4/vbit (column access would unpack whole
578    /// groups — sparse execution falls back to f32 for those).
579    pub fn sparse_col_ok(&self) -> bool {
580        match self {
581            Self::F32 { .. } => true,
582            Self::Mapped { dtype, .. } => {
583                matches!(dtype, TensorDtype::Q8Row | TensorDtype::Q8_2f)
584            }
585        }
586    }
587
588    /// down_proj [hidden, inter]: accumulate `w · col(c)` into `out`
589    /// [hidden] — reads ONLY column `c` (one neuron) from the mmap,
590    /// no full-matrix dequant. `out[k] += w · down[k, c]`.
591    pub fn add_col_scaled(&self, c: usize, w: f32, out: &mut [f32]) {
592        let inter = self.cols();
593        let hidden = self.rows();
594        debug_assert_eq!(out.len(), hidden);
595        match self {
596            Self::F32 { data, .. } => {
597                for (k, o) in out.iter_mut().enumerate() {
598                    *o += w * data[k * inter + c];
599                }
600            }
601            Self::Mapped {
602                dtype,
603                row_scale,
604                col_field,
605                ..
606            } => {
607                let q = self.quant_bytes();
608                let colf = if *dtype == TensorDtype::Q8_2f {
609                    col_field[c]
610                } else {
611                    1.0
612                };
613                let wc = w * colf;
614                for (k, o) in out.iter_mut().enumerate() {
615                    let b = q[k * inter + c] as i8 as f32;
616                    *o += wc * b * row_scale[k];
617                }
618            }
619        }
620    }
621
622    /// Dot of row `r` with `x` (gate/up active-neuron path). Reads only
623    /// row `r` from the mmap — no full dequant. q4/vbit dequant the row
624    /// into `scratch` first (rare for active-FFN weights).
625    pub fn row_dot(&self, r: usize, x: &[f32], scratch: &mut [f32]) -> f32 {
626        let cols = self.cols();
627        match self {
628            Self::F32 { data, .. } => {
629                let row = &data[r * cols..(r + 1) * cols];
630                row.iter().zip(x).map(|(w, v)| w * v).sum()
631            }
632            Self::Mapped {
633                dtype,
634                row_scale,
635                col_field,
636                ..
637            } => match dtype {
638                TensorDtype::Q8Row => {
639                    let q = &self.quant_bytes()[r * cols..(r + 1) * cols];
640                    dot_i8_f32(q, x) * row_scale[r]
641                }
642                TensorDtype::Q8_2f => {
643                    let q = &self.quant_bytes()[r * cols..(r + 1) * cols];
644                    dot_i8_col_f32(q, x, col_field) * row_scale[r]
645                }
646                _ => {
647                    self.row_f32(r, scratch);
648                    scratch.iter().zip(x).map(|(w, v)| w * v).sum()
649                }
650            },
651        }
652    }
653
654    /// `out = W · x` (row-major). F32 delegates to the historical
655    /// bit-exact path; Mapped runs the fused int8 kernel.
656    pub fn matvec(&self, x: &[f32], out: &mut [f32], pool: Option<&Pool>) {
657        match self {
658            Self::F32 { data, .. } => matvec_rows(pool, data, x, out),
659            Self::Mapped {
660                model,
661                idx,
662                dtype,
663                rows,
664                cols,
665                row_scale,
666                col_field,
667                vbit_offsets,
668                repack,
669            } => {
670                let _ = (model, idx);
671                if *dtype == TensorDtype::Q4Block {
672                    // GPU route (wgpu q4b kernel) for large q4_block matvecs —
673                    // gives NVIDIA/AMD/Intel q4 models a GPU path. Probe keeps
674                    // the winner; Metal returns false → the CPU kernel below.
675                    if *rows * *cols >= 8_388_608 && crate::gpu::enabled_here() {
676                        let t0 = std::time::Instant::now();
677                        match crate::gpu::probe_arm(crate::gpu::OpClass::Matvec) {
678                            crate::gpu::ProbeArm::Gpu => {
679                                if crate::gpu::q4b_matvec(model, *idx, x, *rows, *cols, out) {
680                                    crate::gpu::probe_record(
681                                        crate::gpu::OpClass::Matvec,
682                                        true,
683                                        t0.elapsed(),
684                                    );
685                                    return;
686                                }
687                            }
688                            crate::gpu::ProbeArm::CpuTimed => {
689                                q4matvec(self.quant_bytes(), x, *rows, *cols, out, pool);
690                                crate::gpu::probe_record(
691                                    crate::gpu::OpClass::Matvec,
692                                    false,
693                                    t0.elapsed(),
694                                );
695                                return;
696                            }
697                            crate::gpu::ProbeArm::Cpu => {}
698                        }
699                    }
700                    q4matvec(self.quant_bytes(), x, *rows, *cols, out, pool);
701                    return;
702                }
703                if *dtype == TensorDtype::Q4Tiled {
704                    q4t_matvec(self.quant_bytes(), x, *rows, *cols, out, pool);
705                    return;
706                }
707                if *dtype == TensorDtype::Q1 {
708                    // GPU route for large q1 matvecs (out_proj / lm_head
709                    // class): the CPU q1 kernel is load-port-bound at
710                    // ~4 GB/s/core, the GPU one is bandwidth-bound — the
711                    // probe measures both arms and keeps the winner.
712                    if *rows * *cols >= 8_388_608 && crate::gpu::enabled_here() {
713                        let t0 = std::time::Instant::now();
714                        let arm = if crate::gpu::q1_force() {
715                            crate::gpu::ProbeArm::Gpu
716                        } else {
717                            crate::gpu::probe_arm(crate::gpu::OpClass::Matvec)
718                        };
719                        match arm {
720                            crate::gpu::ProbeArm::Gpu => {
721                                if crate::gpu::q1_matvec(model, *idx, x, *rows, *cols, out) {
722                                    crate::gpu::probe_record(
723                                        crate::gpu::OpClass::Matvec,
724                                        true,
725                                        t0.elapsed(),
726                                    );
727                                    return;
728                                }
729                            }
730                            crate::gpu::ProbeArm::CpuTimed => {
731                                q1_matvec(self.quant_bytes(), x, *rows, *cols, out, pool);
732                                crate::gpu::probe_record(
733                                    crate::gpu::OpClass::Matvec,
734                                    false,
735                                    t0.elapsed(),
736                                );
737                                return;
738                            }
739                            crate::gpu::ProbeArm::Cpu => {}
740                        }
741                    }
742                    q1_matvec(self.quant_bytes(), x, *rows, *cols, out, pool);
743                    return;
744                }
745                if *dtype == TensorDtype::Q1T {
746                    // GPU route for large q1t matvecs: the ternary BASE dot runs
747                    // on the GPU (load-port-bound on CPU, like q1), then the
748                    // sparse overlay is added on the CPU. Probe keeps the winner.
749                    if *rows * *cols >= 8_388_608 && crate::gpu::enabled_here() {
750                        let t0 = std::time::Instant::now();
751                        match crate::gpu::probe_arm(crate::gpu::OpClass::Matvec) {
752                            crate::gpu::ProbeArm::Gpu => {
753                                if crate::gpu::q1t_matvec(model, *idx, x, *rows, *cols, out) {
754                                    q1t_add_overlay(self.quant_bytes(), x, *rows, *cols, out, pool);
755                                    crate::gpu::probe_record(
756                                        crate::gpu::OpClass::Matvec,
757                                        true,
758                                        t0.elapsed(),
759                                    );
760                                    return;
761                                }
762                            }
763                            crate::gpu::ProbeArm::CpuTimed => {
764                                q1t_matvec(self.quant_bytes(), x, *rows, *cols, out, pool);
765                                crate::gpu::probe_record(
766                                    crate::gpu::OpClass::Matvec,
767                                    false,
768                                    t0.elapsed(),
769                                );
770                                return;
771                            }
772                            crate::gpu::ProbeArm::Cpu => {}
773                        }
774                    }
775                    q1t_matvec(self.quant_bytes(), x, *rows, *cols, out, pool);
776                    return;
777                }
778                if matches!(dtype, TensorDtype::Vbit | TensorDtype::VbitRo) {
779                    vbitmatvec(self.quant_bytes(), vbit_offsets, x, *rows, *cols, out, pool);
780                    return;
781                }
782                let xs = prescale(x, col_field, *dtype);
783                // D5: large q8 matrices (lm_head-class) — hybrid
784                // CPU∥GPU: split the rows, both sides compute
785                // SIMULTANEOUSLY (same math, shared prescale).
786                // GPU share: CMF_GPU_SPLIT (0..1, default 0.5).
787                if *rows >= crate::gpu::min_rows()
788                    && matches!(dtype, TensorDtype::Q8Row | TensorDtype::Q8_2f)
789                    && std::env::var("CMF_GPU_LMHEAD")
790                        .map(|v| v != "0")
791                        .unwrap_or(true)
792                    && crate::gpu::enabled_here()
793                {
794                    // Runtime probe: alternate the hybrid against the
795                    // pure-CPU matvec, keep whichever is faster HERE.
796                    let t0 = std::time::Instant::now();
797                    match crate::gpu::probe_arm(crate::gpu::OpClass::Matvec) {
798                        crate::gpu::ProbeArm::Gpu => {}
799                        crate::gpu::ProbeArm::CpuTimed => {
800                            qmatvec(
801                                self.quant_bytes(),
802                                repack,
803                                row_scale,
804                                x,
805                                col_field,
806                                *dtype,
807                                *rows,
808                                *cols,
809                                out,
810                                pool,
811                            );
812                            crate::gpu::probe_record(
813                                crate::gpu::OpClass::Matvec,
814                                false,
815                                t0.elapsed(),
816                            );
817                            return;
818                        }
819                        crate::gpu::ProbeArm::Cpu => {
820                            qmatvec(
821                                self.quant_bytes(),
822                                repack,
823                                row_scale,
824                                x,
825                                col_field,
826                                *dtype,
827                                *rows,
828                                *cols,
829                                out,
830                                pool,
831                            );
832                            return;
833                        }
834                    }
835                    let frac = std::env::var("CMF_GPU_SPLIT")
836                        .ok()
837                        .and_then(|v| v.parse::<f32>().ok())
838                        .unwrap_or(0.5)
839                        .clamp(0.0, 1.0);
840                    let cpu_rows = ((*rows as f32) * (1.0 - frac)) as usize;
841                    let (out_cpu, out_gpu) = out.split_at_mut(cpu_rows);
842                    let bytes = self.quant_bytes();
843                    let ok = std::thread::scope(|sc| {
844                        let g = sc.spawn(|| {
845                            crate::gpu::q8_matvec_range(
846                                model,
847                                *idx,
848                                cpu_rows,
849                                &row_scale[cpu_rows..],
850                                &xs,
851                                *rows - cpu_rows,
852                                *cols,
853                                out_gpu,
854                            )
855                        });
856                        if cpu_rows > 0 {
857                            // Repack prefix covers the full groups of the
858                            // CPU half (the split starts at row 0).
859                            let rep_cpu = if repack.is_empty() {
860                                &[][..]
861                            } else {
862                                &repack[..(cpu_rows / 4) * 4 * *cols]
863                            };
864                            qmatvec(
865                                &bytes[..cpu_rows * *cols],
866                                rep_cpu,
867                                &row_scale[..cpu_rows],
868                                x,
869                                col_field,
870                                *dtype,
871                                cpu_rows,
872                                *cols,
873                                out_cpu,
874                                pool,
875                            );
876                        }
877                        g.join().unwrap_or(false)
878                    });
879                    if ok {
880                        crate::gpu::probe_record(crate::gpu::OpClass::Matvec, true, t0.elapsed());
881                        return;
882                    }
883                    // GPU failed — CPU finishes its half (rows rebased —
884                    // group offsets don't line up, mmap layout only).
885                    qmatvec(
886                        &bytes[cpu_rows * *cols..(*rows) * *cols],
887                        &[],
888                        &row_scale[cpu_rows..],
889                        x,
890                        col_field,
891                        *dtype,
892                        *rows - cpu_rows,
893                        *cols,
894                        out_gpu,
895                        pool,
896                    );
897                    return;
898                }
899                qmatvec(
900                    self.quant_bytes(),
901                    repack,
902                    row_scale,
903                    x,
904                    col_field,
905                    *dtype,
906                    *rows,
907                    *cols,
908                    out,
909                    pool,
910                );
911            }
912        }
913    }
914
915    /// Fused two-input matvec (MTP verify pair): weights streamed once.
916    pub fn matvec2(
917        &self,
918        x1: &[f32],
919        x2: &[f32],
920        o1: &mut [f32],
921        o2: &mut [f32],
922        pool: Option<&Pool>,
923    ) {
924        match self {
925            Self::F32 { data, .. } => matvec_rows2(pool, data, x1, x2, o1, o2),
926            Self::Mapped {
927                dtype,
928                rows,
929                cols,
930                row_scale,
931                col_field,
932                vbit_offsets,
933                ..
934            } => {
935                if *dtype == TensorDtype::Q4Block {
936                    q4matvec2(self.quant_bytes(), x1, x2, *rows, *cols, o1, o2, pool);
937                    return;
938                }
939                if *dtype == TensorDtype::Q4Tiled {
940                    q4t_matvec2(self.quant_bytes(), x1, x2, *rows, *cols, o1, o2, pool);
941                    return;
942                }
943                if *dtype == TensorDtype::Q1 {
944                    q1_matvec2(self.quant_bytes(), x1, x2, *rows, *cols, o1, o2, pool);
945                    return;
946                }
947                if *dtype == TensorDtype::Q1T {
948                    // Fused ternary pair: one row pass, the register
949                    // unpack shared across both streams on ARM. (Q1T
950                    // lacks a row_scale array — scales live inline in
951                    // the tiles — so it must not fall through to the
952                    // q8 qmatvec2 below.)
953                    q1t_matvec2(self.quant_bytes(), x1, x2, *rows, *cols, o1, o2, pool);
954                    return;
955                }
956                if matches!(dtype, TensorDtype::Vbit | TensorDtype::VbitRo) {
957                    vbitmatvec2(
958                        self.quant_bytes(),
959                        vbit_offsets,
960                        x1,
961                        x2,
962                        *rows,
963                        *cols,
964                        o1,
965                        o2,
966                        pool,
967                    );
968                    return;
969                }
970                qmatvec2(
971                    self.quant_bytes(),
972                    row_scale,
973                    x1,
974                    x2,
975                    col_field,
976                    *dtype,
977                    *rows,
978                    *cols,
979                    o1,
980                    o2,
981                    pool,
982                );
983            }
984        }
985    }
986}
987
988impl QTensor {
989    /// Batched matvec (prefill-GEMM): xs — row-major [b, cols],
990    /// out — row-major [b, rows]. Element-wise semantics are IDENTICAL
991    /// to b matvec calls (same dot kernels in the same order); the win —
992    /// the weight row streams from DRAM once per batch, not b times.
993    pub fn matmat(&self, xs_all: &[f32], b: usize, out: &mut [f32], pool: Option<&Pool>) {
994        let cols = self.cols();
995        let rows = self.rows();
996        debug_assert_eq!(xs_all.len(), b * cols);
997        debug_assert_eq!(out.len(), b * rows);
998        // GPTQ calibration: fold this layer's inputs into its Hessian. Only
999        // Mapped tensors carry a directory name; the check is a relaxed
1000        // atomic load, free when not calibrating.
1001        if crate::gptq_capture::capturing() {
1002            if let Self::Mapped { model, idx, .. } = self {
1003                crate::gptq_capture::accumulate(&model.tensors[*idx].name, xs_all, b, cols);
1004            }
1005        }
1006        match self {
1007            Self::F32 { data, .. } => {
1008                let out_addr = SendMut(out.as_mut_ptr());
1009                let run = |start: usize, end: usize| {
1010                    for o in start..end {
1011                        let row = &data[o * cols..(o + 1) * cols];
1012                        for bi in 0..b {
1013                            let x = &xs_all[bi * cols..(bi + 1) * cols];
1014                            let mut acc = 0f32;
1015                            for j in 0..cols {
1016                                acc += row[j] * x[j];
1017                            }
1018                            unsafe { *out_addr.at(bi * rows + o) = acc };
1019                        }
1020                    }
1021                };
1022                dispatch_rows(pool, rows, &run);
1023            }
1024            Self::Mapped {
1025                dtype,
1026                row_scale,
1027                col_field,
1028                vbit_offsets,
1029                ..
1030            } => {
1031                if *dtype == TensorDtype::Q4Block {
1032                    q4matmat(self.quant_bytes(), xs_all, b, rows, cols, out, pool);
1033                    return;
1034                }
1035                if *dtype == TensorDtype::Q4Tiled {
1036                    q4t_matmat(self.quant_bytes(), xs_all, b, rows, cols, out, pool);
1037                    return;
1038                }
1039                if *dtype == TensorDtype::Q1 {
1040                    // GPU batched q1 GEMM for wide prefill (q1_mul_mm on the
1041                    // device); the probe keeps whichever beats the CPU matmat.
1042                    if b >= 32
1043                        && b * rows * cols >= 128_000_000
1044                        && cols % 64 == 0
1045                        && crate::gpu::enabled_here()
1046                    {
1047                        if let Self::Mapped { model, idx, .. } = self {
1048                            let t0 = std::time::Instant::now();
1049                            match crate::gpu::probe_arm(crate::gpu::OpClass::Matmat) {
1050                                crate::gpu::ProbeArm::Gpu => {
1051                                    if crate::gpu::q1_matmat(
1052                                        model, *idx, xs_all, b, rows, cols, out,
1053                                    ) {
1054                                        crate::gpu::probe_record(
1055                                            crate::gpu::OpClass::Matmat,
1056                                            true,
1057                                            t0.elapsed(),
1058                                        );
1059                                        return;
1060                                    }
1061                                }
1062                                crate::gpu::ProbeArm::CpuTimed => {
1063                                    q1_matmat(self.quant_bytes(), xs_all, b, rows, cols, out, pool);
1064                                    crate::gpu::probe_record(
1065                                        crate::gpu::OpClass::Matmat,
1066                                        false,
1067                                        t0.elapsed(),
1068                                    );
1069                                    return;
1070                                }
1071                                crate::gpu::ProbeArm::Cpu => {}
1072                            }
1073                        }
1074                    }
1075                    q1_matmat(self.quant_bytes(), xs_all, b, rows, cols, out, pool);
1076                    return;
1077                }
1078                if *dtype == TensorDtype::Q1T {
1079                    // GPU batched GEMM for wide prefill (base + overlay on the
1080                    // device); probe keeps the winner vs the CPU matmat.
1081                    if b >= 32 && b * rows * cols >= 128_000_000 && crate::gpu::enabled_here() {
1082                        if let Self::Mapped { model, idx, .. } = self {
1083                            let t0 = std::time::Instant::now();
1084                            match crate::gpu::probe_arm(crate::gpu::OpClass::Matmat) {
1085                                crate::gpu::ProbeArm::Gpu => {
1086                                    if crate::gpu::q1t_matmat(
1087                                        model, *idx, xs_all, b, rows, cols, out,
1088                                    ) {
1089                                        crate::gpu::probe_record(
1090                                            crate::gpu::OpClass::Matmat,
1091                                            true,
1092                                            t0.elapsed(),
1093                                        );
1094                                        return;
1095                                    }
1096                                }
1097                                crate::gpu::ProbeArm::CpuTimed => {
1098                                    q1t_matmat(
1099                                        self.quant_bytes(),
1100                                        xs_all,
1101                                        b,
1102                                        rows,
1103                                        cols,
1104                                        out,
1105                                        pool,
1106                                    );
1107                                    crate::gpu::probe_record(
1108                                        crate::gpu::OpClass::Matmat,
1109                                        false,
1110                                        t0.elapsed(),
1111                                    );
1112                                    return;
1113                                }
1114                                crate::gpu::ProbeArm::Cpu => {}
1115                            }
1116                        }
1117                    }
1118                    q1t_matmat(self.quant_bytes(), xs_all, b, rows, cols, out, pool);
1119                    return;
1120                }
1121                if matches!(dtype, TensorDtype::Vbit | TensorDtype::VbitRo) {
1122                    vbitmatmat(
1123                        self.quant_bytes(),
1124                        vbit_offsets,
1125                        xs_all,
1126                        b,
1127                        rows,
1128                        cols,
1129                        out,
1130                        pool,
1131                    );
1132                    return;
1133                }
1134                let pre: Vec<std::borrow::Cow<'_, [f32]>> = (0..b)
1135                    .map(|bi| prescale(&xs_all[bi * cols..(bi + 1) * cols], col_field, *dtype))
1136                    .collect();
1137                // D5: large prefill-batch GEMMs — on the GPU (threshold by
1138                // work volume: submission carries b×rows×cols MACs).
1139                // Runtime probe: the naive GEMM shader + sync readback
1140                // lose to the CPU GEMM on slow driver stacks — alternate
1141                // both arms and keep the winner.
1142                if b >= 8 && b * rows * cols >= 128_000_000 && crate::gpu::enabled_here() {
1143                    if let Self::Mapped { model, idx, .. } = self {
1144                        let t0 = std::time::Instant::now();
1145                        match crate::gpu::probe_arm(crate::gpu::OpClass::Matmat) {
1146                            crate::gpu::ProbeArm::Gpu
1147                                if crate::gpu::probe_deciding(crate::gpu::OpClass::Matmat)
1148                                    && !crate::gpu::q8_resident_or_upload(model, *idx) =>
1149                            {
1150                                // Cold weights during probing: the upload
1151                                // has started, the count runs on the CPU —
1152                                // the GPU arm samples on the next touch.
1153                                let q = self.quant_bytes();
1154                                qmatmat(q, row_scale, &pre, rows, cols, out, pool);
1155                                return;
1156                            }
1157                            crate::gpu::ProbeArm::Gpu => {
1158                                let flat: Vec<f32> =
1159                                    pre.iter().flat_map(|v| v.iter().copied()).collect();
1160                                if crate::gpu::q8_matmat(
1161                                    model, *idx, row_scale, &flat, b, rows, cols, out,
1162                                ) {
1163                                    crate::gpu::probe_record(
1164                                        crate::gpu::OpClass::Matmat,
1165                                        true,
1166                                        t0.elapsed(),
1167                                    );
1168                                    return;
1169                                }
1170                            }
1171                            crate::gpu::ProbeArm::CpuTimed => {
1172                                let q = self.quant_bytes();
1173                                qmatmat(q, row_scale, &pre, rows, cols, out, pool);
1174                                crate::gpu::probe_record(
1175                                    crate::gpu::OpClass::Matmat,
1176                                    false,
1177                                    t0.elapsed(),
1178                                );
1179                                return;
1180                            }
1181                            crate::gpu::ProbeArm::Cpu => {}
1182                        }
1183                    }
1184                }
1185                let q = self.quant_bytes();
1186                qmatmat(q, row_scale, &pre, rows, cols, out, pool);
1187            }
1188        }
1189    }
1190}
1191
1192impl QTensor {
1193    /// Multi-matrix job (roadmap §3 P0): N tensors sharing one input
1194    /// run under a SINGLE pool dispatch — QKV or gate+up cost one
1195    /// barrier instead of N. Per-row math is the exact same kernel as
1196    /// `matvec` (bit-identical outputs); only the dispatch is fused.
1197    /// Falls back to N sequential matvecs when the set is not a uniform
1198    /// q8-family/F32 group or there is no pool.
1199    pub fn matvec_many<const N: usize>(
1200        ts: [&QTensor; N],
1201        x: &[f32],
1202        mut outs: [&mut [f32]; N],
1203        pool: Option<&Pool>,
1204    ) {
1205        let total_rows: usize = ts.iter().map(|t| t.rows()).sum();
1206        let uniform_q8 = ts.iter().all(|t| {
1207            matches!(
1208                t,
1209                Self::Mapped {
1210                    dtype: TensorDtype::Q8Row | TensorDtype::Q8_2f,
1211                    ..
1212                }
1213            )
1214        });
1215        let uniform_f32 = ts.iter().all(|t| matches!(t, Self::F32 { .. }));
1216        let uniform_q4 = ts.iter().all(|t| {
1217            matches!(
1218                t,
1219                Self::Mapped {
1220                    dtype: TensorDtype::Q4Block,
1221                    ..
1222                }
1223            )
1224        });
1225        let uniform_vbit = ts.iter().all(|t| {
1226            matches!(
1227                t,
1228                Self::Mapped {
1229                    dtype: TensorDtype::Vbit | TensorDtype::VbitRo,
1230                    ..
1231                }
1232            )
1233        });
1234        let uniform_q1 = ts.iter().all(|t| {
1235            matches!(
1236                t,
1237                Self::Mapped {
1238                    dtype: TensorDtype::Q1,
1239                    ..
1240                }
1241            )
1242        });
1243        let uniform_q1t = ts.iter().all(|t| {
1244            matches!(
1245                t,
1246                Self::Mapped {
1247                    dtype: TensorDtype::Q1T,
1248                    ..
1249                }
1250            )
1251        });
1252        let Some(pool) = pool else {
1253            for (t, o) in ts.iter().zip(outs.iter_mut()) {
1254                t.matvec(x, o, None);
1255            }
1256            return;
1257        };
1258        if total_rows < 256
1259            || !(uniform_q8
1260                || uniform_f32
1261                || uniform_q4
1262                || uniform_vbit
1263                || uniform_q1
1264                || uniform_q1t)
1265        {
1266            for (t, o) in ts.iter().zip(outs.iter_mut()) {
1267                t.matvec(x, o, Some(pool));
1268            }
1269            return;
1270        }
1271
1272        if uniform_q1 {
1273            // One shared activation split + group sums (q1 has no col
1274            // field; the same input feeds every tensor).
1275            let outs_addr: [SendMut; N] = std::array::from_fn(|i| SendMut(outs[i].as_mut_ptr()));
1276            if a8w8_enabled() {
1277                let act = split_act(x);
1278                let gsum = q1_group_sums(&act.xq, ts[0].cols() / GROUP_SIZE);
1279                let (act, gsum) = (&act, &gsum);
1280                let closures: [_; N] = std::array::from_fn(|i| {
1281                    let (bytes, gpr, out) =
1282                        (ts[i].quant_bytes(), ts[i].cols() / GROUP_SIZE, outs_addr[i]);
1283                    move |s: usize, e: usize| q1_range_a8w8(bytes, gpr, act, gsum, out, s, e)
1284                });
1285                let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
1286                    std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
1287                pool.run_many(&parts);
1288            } else {
1289                let closures: [_; N] = std::array::from_fn(|i| {
1290                    let (bytes, gpr, out) =
1291                        (ts[i].quant_bytes(), ts[i].cols() / GROUP_SIZE, outs_addr[i]);
1292                    move |s: usize, e: usize| q1_range_f32(bytes, gpr, x, out, s, e)
1293                });
1294                let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
1295                    std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
1296                pool.run_many(&parts);
1297            }
1298            return;
1299        }
1300
1301        if uniform_q1t {
1302            // Q1T batched: one shared activation split + overlay decode,
1303            // all tensors' rows in ONE pool dispatch (saves N−1 dispatches
1304            // and N−1 redundant split_act calls per layer).
1305            let outs_addr: [SendMut; N] = std::array::from_fn(|i| SendMut(outs[i].as_mut_ptr()));
1306            const TILE: usize = cortiq_core::quant::Q1T_TILE;
1307            if a8w8_enabled() {
1308                let act = split_act(x);
1309                let act = &act;
1310                let x_ref = x;
1311                let closures: [_; N] = std::array::from_fn(|i| {
1312                    let bytes = ts[i].quant_bytes();
1313                    let (rows, cols) = (ts[i].rows(), ts[i].cols());
1314                    let gpr = cols / GROUP_SIZE;
1315                    let (rp_off, ent_off, has_ov) = q1t_overlay(bytes, rows * gpr * TILE, rows);
1316                    let out = outs_addr[i];
1317                    move |s: usize, e: usize| {
1318                        q1t_range_a8w8(bytes, gpr, rp_off, ent_off, has_ov, act, x_ref, out, s, e)
1319                    }
1320                });
1321                let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
1322                    std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
1323                pool.run_many(&parts);
1324            } else {
1325                let x_ref = x;
1326                let closures: [_; N] = std::array::from_fn(|i| {
1327                    let bytes = ts[i].quant_bytes();
1328                    let (rows, cols) = (ts[i].rows(), ts[i].cols());
1329                    let gpr = cols / GROUP_SIZE;
1330                    let (rp_off, ent_off, has_ov) = q1t_overlay(bytes, rows * gpr * TILE, rows);
1331                    let out = outs_addr[i];
1332                    move |s: usize, e: usize| {
1333                        q1t_range_f32_batch(bytes, gpr, rp_off, ent_off, has_ov, x_ref, out, s, e)
1334                    }
1335                });
1336                let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
1337                    std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
1338                pool.run_many(&parts);
1339            }
1340            return;
1341        }
1342
1343        if uniform_q4 || uniform_vbit {
1344            let outs_addr: [SendMut; N] = std::array::from_fn(|i| SendMut(outs[i].as_mut_ptr()));
1345            // q4/vbit share one activation split — no per-tensor col field.
1346            if a8w8_enabled() {
1347                let act = split_act(x);
1348                let act = &act;
1349                if uniform_q4 {
1350                    let closures: [_; N] = std::array::from_fn(|i| {
1351                        let (packed, scales) =
1352                            q4_split(ts[i].quant_bytes(), ts[i].rows(), ts[i].cols());
1353                        let (gpr, cols, out) =
1354                            (ts[i].cols() / GROUP_SIZE, ts[i].cols(), outs_addr[i]);
1355                        move |s: usize, e: usize| {
1356                            q4_range_a8w8(packed, scales, gpr, cols, act, out, s, e)
1357                        }
1358                    });
1359                    let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
1360                        std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
1361                    pool.run_many(&parts);
1362                } else {
1363                    let closures: [_; N] = std::array::from_fn(|i| {
1364                        let Self::Mapped { vbit_offsets, .. } = ts[i] else {
1365                            unreachable!()
1366                        };
1367                        let (bytes, rows, cols, out) = (
1368                            ts[i].quant_bytes(),
1369                            ts[i].rows(),
1370                            ts[i].cols(),
1371                            outs_addr[i],
1372                        );
1373                        move |s: usize, e: usize| {
1374                            vbit_range_a8w8(bytes, vbit_offsets, x, act, rows, cols, out, s, e)
1375                        }
1376                    });
1377                    let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
1378                        std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
1379                    pool.run_many(&parts);
1380                }
1381                return;
1382            }
1383            if uniform_q4 {
1384                let closures: [_; N] = std::array::from_fn(|i| {
1385                    let (packed, scales) =
1386                        q4_split(ts[i].quant_bytes(), ts[i].rows(), ts[i].cols());
1387                    let (gpr, out) = (ts[i].cols() / GROUP_SIZE, outs_addr[i]);
1388                    move |s: usize, e: usize| q4_range_f32(packed, scales, gpr, x, out, s, e)
1389                });
1390                let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
1391                    std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
1392                pool.run_many(&parts);
1393            } else {
1394                let closures: [_; N] = std::array::from_fn(|i| {
1395                    let Self::Mapped { vbit_offsets, .. } = ts[i] else {
1396                        unreachable!()
1397                    };
1398                    let (bytes, rows, cols, out) = (
1399                        ts[i].quant_bytes(),
1400                        ts[i].rows(),
1401                        ts[i].cols(),
1402                        outs_addr[i],
1403                    );
1404                    move |s: usize, e: usize| {
1405                        vbit_range_f32(bytes, vbit_offsets, x, rows, cols, out, s, e)
1406                    }
1407                });
1408                let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
1409                    std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
1410                pool.run_many(&parts);
1411            }
1412            return;
1413        }
1414
1415        if uniform_f32 {
1416            let outs_addr: [SendMut; N] = std::array::from_fn(|i| SendMut(outs[i].as_mut_ptr()));
1417            let closures: [_; N] = std::array::from_fn(|i| {
1418                let Self::F32 { data, cols, .. } = ts[i] else {
1419                    unreachable!()
1420                };
1421                let out = outs_addr[i];
1422                move |start: usize, end: usize| {
1423                    for o in start..end {
1424                        let row = &data[o * cols..(o + 1) * cols];
1425                        let mut sum = 0.0f32;
1426                        for j in 0..*cols {
1427                            sum += row[j] * x[j];
1428                        }
1429                        // SAFETY: disjoint (tensor, row) cells per worker.
1430                        unsafe { *out.at(o) = sum };
1431                    }
1432                }
1433            });
1434            let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
1435                std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
1436            pool.run_many(&parts);
1437            return;
1438        }
1439
1440        // Uniform q8-family: per-tensor prescale (q8_2f col fields
1441        // differ per tensor) + the shared range kernels.
1442        struct Ctx<'a> {
1443            bytes: &'a [u8],
1444            #[cfg_attr(not(target_arch = "aarch64"), allow(dead_code))]
1445            rep: &'a [u8],
1446            row_scale: &'a [f32],
1447            cols: usize,
1448            xs: std::borrow::Cow<'a, [f32]>,
1449        }
1450        let ctxs: [Ctx<'_>; N] = std::array::from_fn(|i| {
1451            let Self::Mapped {
1452                dtype,
1453                cols,
1454                row_scale,
1455                col_field,
1456                repack,
1457                ..
1458            } = ts[i]
1459            else {
1460                unreachable!()
1461            };
1462            Ctx {
1463                bytes: ts[i].quant_bytes(),
1464                rep: repack,
1465                row_scale,
1466                cols: *cols,
1467                xs: prescale(x, col_field, *dtype),
1468            }
1469        });
1470        let outs_addr: [SendMut; N] = std::array::from_fn(|i| SendMut(outs[i].as_mut_ptr()));
1471        #[cfg(target_arch = "aarch64")]
1472        if sdot_enabled() {
1473            let acts: [SplitAct; N] = std::array::from_fn(|i| split_act(&ctxs[i].xs));
1474            let closures: [_; N] = std::array::from_fn(|i| {
1475                let (c, act, out) = (&ctxs[i], &acts[i], outs_addr[i]);
1476                move |start: usize, end: usize| {
1477                    q8_range_sdot(c.bytes, c.rep, c.row_scale, act, c.cols, out, start, end)
1478                }
1479            });
1480            let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
1481                std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
1482            pool.run_many(&parts);
1483            return;
1484        }
1485        #[cfg(target_arch = "x86_64")]
1486        if avx2_a8w8_enabled() {
1487            let acts: [SplitAct; N] = std::array::from_fn(|i| split_act(&ctxs[i].xs));
1488            let closures: [_; N] = std::array::from_fn(|i| {
1489                let (c, act, out) = (&ctxs[i], &acts[i], outs_addr[i]);
1490                move |start: usize, end: usize| {
1491                    q8_range_avx2(c.bytes, c.row_scale, act, c.cols, out, start, end)
1492                }
1493            });
1494            let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
1495                std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
1496            pool.run_many(&parts);
1497            return;
1498        }
1499        let closures: [_; N] = std::array::from_fn(|i| {
1500            let (c, out) = (&ctxs[i], outs_addr[i]);
1501            move |start: usize, end: usize| {
1502                q8_range_f32(c.bytes, c.row_scale, &c.xs, c.cols, out, start, end)
1503            }
1504        });
1505        let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
1506            std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
1507        pool.run_many(&parts);
1508    }
1509}
1510
1511impl QTensor {
1512    /// Pair-input multi-matrix job: N tensors × 2 shared inputs under a
1513    /// single pool dispatch — the MTP/pair decode path publishes one job
1514    /// for Q/K/V (and one for gate+up) instead of one per tensor.
1515    /// Per-row math is exactly `matvec2`'s kernels; bit-identical.
1516    #[allow(clippy::needless_range_loop)]
1517    pub fn matvec2_many<const N: usize>(
1518        ts: [&QTensor; N],
1519        x1: &[f32],
1520        x2: &[f32],
1521        mut o1s: [&mut [f32]; N],
1522        mut o2s: [&mut [f32]; N],
1523        pool: Option<&Pool>,
1524    ) {
1525        let total_rows: usize = ts.iter().map(|t| t.rows()).sum();
1526        let uniform_q8 = ts.iter().all(|t| {
1527            matches!(
1528                t,
1529                Self::Mapped {
1530                    dtype: TensorDtype::Q8Row | TensorDtype::Q8_2f,
1531                    ..
1532                }
1533            )
1534        });
1535        let uniform_f32 = ts.iter().all(|t| matches!(t, Self::F32 { .. }));
1536        let uniform_q4 = ts.iter().all(|t| {
1537            matches!(
1538                t,
1539                Self::Mapped {
1540                    dtype: TensorDtype::Q4Block,
1541                    ..
1542                }
1543            )
1544        });
1545        let uniform_vbit = ts.iter().all(|t| {
1546            matches!(
1547                t,
1548                Self::Mapped {
1549                    dtype: TensorDtype::Vbit | TensorDtype::VbitRo,
1550                    ..
1551                }
1552            )
1553        });
1554        let fusable = pool.is_some()
1555            && total_rows >= 256
1556            && (uniform_q8 || uniform_f32 || uniform_q4 || uniform_vbit);
1557        if !fusable {
1558            for i in 0..N {
1559                ts[i].matvec2(x1, x2, o1s[i], o2s[i], pool);
1560            }
1561            return;
1562        }
1563        let pool = pool.unwrap();
1564
1565        if uniform_q4 || uniform_vbit {
1566            let p1: [SendMut; N] = std::array::from_fn(|i| SendMut(o1s[i].as_mut_ptr()));
1567            let p2: [SendMut; N] = std::array::from_fn(|i| SendMut(o2s[i].as_mut_ptr()));
1568            // q4/vbit share activation splits — no per-tensor col field.
1569            if a8w8_enabled() {
1570                let a1 = split_act(x1);
1571                let a2 = split_act(x2);
1572                let (a1, a2) = (&a1, &a2);
1573                if uniform_q4 {
1574                    let closures: [_; N] = std::array::from_fn(|i| {
1575                        let (packed, scales) =
1576                            q4_split(ts[i].quant_bytes(), ts[i].rows(), ts[i].cols());
1577                        let (gpr, cols, o1, o2) =
1578                            (ts[i].cols() / GROUP_SIZE, ts[i].cols(), p1[i], p2[i]);
1579                        move |s: usize, e: usize| {
1580                            q4_range2_a8w8(packed, scales, gpr, cols, a1, a2, o1, o2, s, e)
1581                        }
1582                    });
1583                    let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
1584                        std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
1585                    pool.run_many(&parts);
1586                } else {
1587                    let closures: [_; N] = std::array::from_fn(|i| {
1588                        let Self::Mapped { vbit_offsets, .. } = ts[i] else {
1589                            unreachable!()
1590                        };
1591                        let (bytes, rows, cols, o1, o2) = (
1592                            ts[i].quant_bytes(),
1593                            ts[i].rows(),
1594                            ts[i].cols(),
1595                            p1[i],
1596                            p2[i],
1597                        );
1598                        move |s: usize, e: usize| {
1599                            vbit_range2_a8w8(
1600                                bytes,
1601                                vbit_offsets,
1602                                x1,
1603                                x2,
1604                                a1,
1605                                a2,
1606                                rows,
1607                                cols,
1608                                o1,
1609                                o2,
1610                                s,
1611                                e,
1612                            )
1613                        }
1614                    });
1615                    let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
1616                        std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
1617                    pool.run_many(&parts);
1618                }
1619                return;
1620            }
1621            if uniform_q4 {
1622                let closures: [_; N] = std::array::from_fn(|i| {
1623                    let (packed, scales) =
1624                        q4_split(ts[i].quant_bytes(), ts[i].rows(), ts[i].cols());
1625                    let (gpr, o1, o2) = (ts[i].cols() / GROUP_SIZE, p1[i], p2[i]);
1626                    move |s: usize, e: usize| {
1627                        q4_range2_f32(packed, scales, gpr, x1, x2, o1, o2, s, e)
1628                    }
1629                });
1630                let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
1631                    std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
1632                pool.run_many(&parts);
1633            } else {
1634                let closures: [_; N] = std::array::from_fn(|i| {
1635                    let Self::Mapped { vbit_offsets, .. } = ts[i] else {
1636                        unreachable!()
1637                    };
1638                    let (bytes, rows, cols, o1, o2) = (
1639                        ts[i].quant_bytes(),
1640                        ts[i].rows(),
1641                        ts[i].cols(),
1642                        p1[i],
1643                        p2[i],
1644                    );
1645                    move |s: usize, e: usize| {
1646                        vbit_range2_f32(bytes, vbit_offsets, x1, x2, rows, cols, o1, o2, s, e)
1647                    }
1648                });
1649                let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
1650                    std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
1651                pool.run_many(&parts);
1652            }
1653            return;
1654        }
1655
1656        if uniform_f32 {
1657            let p1: [SendMut; N] = std::array::from_fn(|i| SendMut(o1s[i].as_mut_ptr()));
1658            let p2: [SendMut; N] = std::array::from_fn(|i| SendMut(o2s[i].as_mut_ptr()));
1659            let closures: [_; N] = std::array::from_fn(|i| {
1660                let Self::F32 { data, cols, .. } = ts[i] else {
1661                    unreachable!()
1662                };
1663                let (o1, o2) = (p1[i], p2[i]);
1664                move |start: usize, end: usize| {
1665                    for o in start..end {
1666                        let row = &data[o * cols..(o + 1) * cols];
1667                        let (mut s1, mut s2) = (0.0f32, 0.0f32);
1668                        for j in 0..*cols {
1669                            s1 += row[j] * x1[j];
1670                            s2 += row[j] * x2[j];
1671                        }
1672                        // SAFETY: disjoint (tensor, row) cells per worker.
1673                        unsafe {
1674                            *o1.at(o) = s1;
1675                            *o2.at(o) = s2;
1676                        }
1677                    }
1678                }
1679            });
1680            let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
1681                std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
1682            pool.run_many(&parts);
1683            return;
1684        }
1685
1686        struct Ctx<'a> {
1687            bytes: &'a [u8],
1688            row_scale: &'a [f32],
1689            cols: usize,
1690            xs1: std::borrow::Cow<'a, [f32]>,
1691            xs2: std::borrow::Cow<'a, [f32]>,
1692        }
1693        let ctxs: [Ctx<'_>; N] = std::array::from_fn(|i| {
1694            let Self::Mapped {
1695                dtype,
1696                cols,
1697                row_scale,
1698                col_field,
1699                ..
1700            } = ts[i]
1701            else {
1702                unreachable!()
1703            };
1704            Ctx {
1705                bytes: ts[i].quant_bytes(),
1706                row_scale,
1707                cols: *cols,
1708                xs1: prescale(x1, col_field, *dtype),
1709                xs2: prescale(x2, col_field, *dtype),
1710            }
1711        });
1712        let p1: [SendMut; N] = std::array::from_fn(|i| SendMut(o1s[i].as_mut_ptr()));
1713        let p2: [SendMut; N] = std::array::from_fn(|i| SendMut(o2s[i].as_mut_ptr()));
1714        #[cfg(target_arch = "aarch64")]
1715        if sdot_enabled() {
1716            let acts: [(SplitAct, SplitAct); N] =
1717                std::array::from_fn(|i| (split_act(&ctxs[i].xs1), split_act(&ctxs[i].xs2)));
1718            let closures: [_; N] = std::array::from_fn(|i| {
1719                let (c, a, o1, o2) = (&ctxs[i], &acts[i], p1[i], p2[i]);
1720                move |start: usize, end: usize| {
1721                    q8_range2_sdot(c.bytes, c.row_scale, &a.0, &a.1, c.cols, o1, o2, start, end)
1722                }
1723            });
1724            let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
1725                std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
1726            pool.run_many(&parts);
1727            return;
1728        }
1729        #[cfg(target_arch = "x86_64")]
1730        if avx2_a8w8_enabled() {
1731            let acts: [(SplitAct, SplitAct); N] =
1732                std::array::from_fn(|i| (split_act(&ctxs[i].xs1), split_act(&ctxs[i].xs2)));
1733            let closures: [_; N] = std::array::from_fn(|i| {
1734                let (c, a, o1, o2) = (&ctxs[i], &acts[i], p1[i], p2[i]);
1735                move |start: usize, end: usize| {
1736                    q8_range2_avx2(c.bytes, c.row_scale, &a.0, &a.1, c.cols, o1, o2, start, end)
1737                }
1738            });
1739            let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
1740                std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
1741            pool.run_many(&parts);
1742            return;
1743        }
1744        let closures: [_; N] = std::array::from_fn(|i| {
1745            let (c, o1, o2) = (&ctxs[i], p1[i], p2[i]);
1746            move |start: usize, end: usize| {
1747                q8_range2_f32(
1748                    c.bytes,
1749                    c.row_scale,
1750                    &c.xs1,
1751                    &c.xs2,
1752                    c.cols,
1753                    o1,
1754                    o2,
1755                    start,
1756                    end,
1757                )
1758            }
1759        });
1760        let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
1761            std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
1762        pool.run_many(&parts);
1763    }
1764
1765    /// Fused gate+up matvec with SiLU·mul: for each row r, computes
1766    /// `silu(gate·x) * (up·x)` and writes to `out[r]`. ONE pool dispatch,
1767    /// no intermediate g/u buffers, no separate silu pass. Falls back
1768    /// (returns false) for unsupported dtype combos.
1769    pub fn matvec_silu_mul(
1770        gate: &QTensor,
1771        up: &QTensor,
1772        x: &[f32],
1773        out: &mut [f32],
1774        pool: Option<&Pool>,
1775    ) -> bool {
1776        let inter = gate.rows();
1777        debug_assert_eq!(up.rows(), inter);
1778        debug_assert_eq!(out.len(), inter);
1779        debug_assert_eq!(gate.cols(), up.cols());
1780        if !a8w8_enabled() {
1781            return false;
1782        }
1783        let act = split_act(x);
1784        let act = &act;
1785        let x_ref = x;
1786        let out_addr = SendMut(out.as_mut_ptr());
1787
1788        match (gate, up) {
1789            // Q4Block gate + Q4Block up (most common mobile q4 models)
1790            (
1791                Self::Mapped {
1792                    dtype: TensorDtype::Q4Block,
1793                    ..
1794                },
1795                Self::Mapped {
1796                    dtype: TensorDtype::Q4Block,
1797                    ..
1798                },
1799            ) => {
1800                let (gp, gs) = q4_split(gate.quant_bytes(), gate.rows(), gate.cols());
1801                let (up_p, up_s) = q4_split(up.quant_bytes(), up.rows(), up.cols());
1802                let gpr = gate.cols() / GROUP_SIZE;
1803                let cols = gate.cols();
1804                let run = move |start: usize, end: usize| {
1805                    for r in start..end {
1806                        let mut gv = dot_q4_row_i8(gp, gs, r * gpr, gpr, &act.xq) * act.sx;
1807                        let mut uv = dot_q4_row_i8(up_p, up_s, r * gpr, gpr, &act.xq) * act.sx;
1808                        for &(j, xv) in &act.outliers {
1809                            let flat = r * cols + j;
1810                            let gb = gp[flat / 2];
1811                            let gn = if flat & 1 == 0 { gb & 0x0F } else { gb >> 4 };
1812                            let gsc = f16_to_f32(u16::from_le_bytes([
1813                                gs[(flat / GROUP_SIZE) * 2],
1814                                gs[(flat / GROUP_SIZE) * 2 + 1],
1815                            ]));
1816                            gv += ((gn as i32 - 8) as f32) * gsc * xv;
1817                            let ub = up_p[flat / 2];
1818                            let un = if flat & 1 == 0 { ub & 0x0F } else { ub >> 4 };
1819                            let usc = f16_to_f32(u16::from_le_bytes([
1820                                up_s[(flat / GROUP_SIZE) * 2],
1821                                up_s[(flat / GROUP_SIZE) * 2 + 1],
1822                            ]));
1823                            uv += ((un as i32 - 8) as f32) * usc * xv;
1824                        }
1825                        let silu_g = gv / (1.0 + (-gv).exp());
1826                        // SAFETY: disjoint row ranges per worker.
1827                        unsafe { *out_addr.at(r) = silu_g * uv };
1828                    }
1829                };
1830                dispatch_rows(pool, inter, &run);
1831                true
1832            }
1833            // Q4Tiled gate + Q4Tiled up — one row pass, both tile
1834            // streams sequential, silu·mul fused (same per-row math as
1835            // `q4t_matvec`).
1836            (
1837                Self::Mapped {
1838                    dtype: TensorDtype::Q4Tiled,
1839                    ..
1840                },
1841                Self::Mapped {
1842                    dtype: TensorDtype::Q4Tiled,
1843                    ..
1844                },
1845            ) => {
1846                let g_bytes = gate.quant_bytes();
1847                let u_bytes = up.quant_bytes();
1848                let gpr = gate.cols() / GROUP_SIZE;
1849                let run = move |start: usize, end: usize| {
1850                    for r in start..end {
1851                        let mut gv = dot_q4t_row_i8(g_bytes, r, gpr, &act.xq) * act.sx;
1852                        let mut uv = dot_q4t_row_i8(u_bytes, r, gpr, &act.xq) * act.sx;
1853                        for &(j, xv) in &act.outliers {
1854                            let (w, s) = q4t_outlier(g_bytes, r, gpr, j);
1855                            gv += w * s * xv;
1856                            let (w, s) = q4t_outlier(u_bytes, r, gpr, j);
1857                            uv += w * s * xv;
1858                        }
1859                        let silu_g = gv / (1.0 + (-gv).exp());
1860                        // SAFETY: disjoint row ranges per worker.
1861                        unsafe { *out_addr.at(r) = silu_g * uv };
1862                    }
1863                };
1864                dispatch_rows(pool, inter, &run);
1865                true
1866            }
1867            // Q1T gate + Q1T up
1868            (
1869                Self::Mapped {
1870                    dtype: TensorDtype::Q1T,
1871                    ..
1872                },
1873                Self::Mapped {
1874                    dtype: TensorDtype::Q1T,
1875                    ..
1876                },
1877            ) => {
1878                const TILE: usize = cortiq_core::quant::Q1T_TILE;
1879                let g_bytes = gate.quant_bytes();
1880                let u_bytes = up.quant_bytes();
1881                let gpr = gate.cols() / GROUP_SIZE;
1882                let (g_rp, g_ent, g_ov) = q1t_overlay(g_bytes, inter * gpr * TILE, inter);
1883                let (u_rp, u_ent, u_ov) = q1t_overlay(u_bytes, inter * gpr * TILE, inter);
1884                let run = move |start: usize, end: usize| {
1885                    for r in start..end {
1886                        let mut gv = q1t_dot_row_i8(g_bytes, r, gpr, &act.xq) * act.sx;
1887                        let mut uv = q1t_dot_row_i8(u_bytes, r, gpr, &act.xq) * act.sx;
1888                        for &(j, xv) in &act.outliers {
1889                            gv += q1t_base_weight(g_bytes, r, gpr, j) * xv;
1890                            uv += q1t_base_weight(u_bytes, r, gpr, j) * xv;
1891                        }
1892                        gv += q1t_row_outlier_correction(g_bytes, r, g_rp, g_ent, g_ov, x_ref);
1893                        uv += q1t_row_outlier_correction(u_bytes, r, u_rp, u_ent, u_ov, x_ref);
1894                        let silu_g = gv / (1.0 + (-gv).exp());
1895                        // SAFETY: disjoint row ranges per worker.
1896                        unsafe { *out_addr.at(r) = silu_g * uv };
1897                    }
1898                };
1899                dispatch_rows(pool, inter, &run);
1900                true
1901            }
1902            _ => false,
1903        }
1904    }
1905}
1906
1907/// Batched q8 kernel: same math as qmatvec, the row makes a single
1908/// pass from memory for the whole batch.
1909/// Accelerate CBLAS — the Apple AMX matrix units, the same engine
1910/// llama.cpp's `-ngl 0` prefill rides via ggml-blas.
1911#[cfg(target_os = "macos")]
1912mod accel_blas {
1913    #[link(name = "Accelerate", kind = "framework")]
1914    unsafe extern "C" {
1915        pub fn cblas_sgemm(
1916            order: i32,
1917            trans_a: i32,
1918            trans_b: i32,
1919            m: i32,
1920            n: i32,
1921            k: i32,
1922            alpha: f32,
1923            a: *const f32,
1924            lda: i32,
1925            b: *const f32,
1926            ldb: i32,
1927            beta: f32,
1928            c: *mut f32,
1929            ldc: i32,
1930        );
1931    }
1932}
1933
1934#[cfg(target_os = "macos")]
1935pub(crate) fn accel_gemm_enabled() -> bool {
1936    static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
1937    *ON.get_or_init(|| std::env::var("CMF_ACCEL").map(|v| v != "0").unwrap_or(true))
1938}
1939
1940/// Off macOS the "accel" GEMM is the portable NEON micro-kernel below —
1941/// same entry point, so the batched-attention path opens on mobile.
1942#[cfg(all(target_arch = "aarch64", not(target_os = "macos")))]
1943pub(crate) fn accel_gemm_enabled() -> bool {
1944    true
1945}
1946
1947/// Portable NEON f32 GEMM (row-major, optional Bᵀ): a 4×8 fmla
1948/// micro-kernel with A broadcast against B panels — the mobile stand-in
1949/// for Accelerate in the batched causal attention (QKᵀ and P·V). Not a
1950/// BLAS: shapes here are the attention panels (m ≤ heads·chunk,
1951/// k = head_dim or context), and the goal is removing the per-position
1952/// quadratic wall, not peak GEMM.
1953#[cfg(target_arch = "aarch64")]
1954#[allow(clippy::too_many_arguments)]
1955pub(crate) fn neon_gemm_rm(
1956    m: usize,
1957    n: usize,
1958    k: usize,
1959    alpha: f32,
1960    a: &[f32],
1961    lda: usize,
1962    b_mat: &[f32],
1963    ldb: usize,
1964    b_rows_are_n: bool,
1965    c: &mut [f32],
1966    ldc: usize,
1967) {
1968    debug_assert!(a.len() >= (m - 1) * lda + k);
1969    debug_assert!(c.len() >= (m - 1) * ldc + n);
1970    // SAFETY: bounds asserted above; NEON is baseline on aarch64.
1971    unsafe {
1972        use core::arch::aarch64::*;
1973        let mut i = 0usize;
1974        while i < m {
1975            let mi = (m - i).min(4);
1976            let mut j = 0usize;
1977            while j < n {
1978                let nj = (n - j).min(8);
1979                if mi == 4 && nj == 8 {
1980                    let (mut c0a, mut c0b) = (vdupq_n_f32(0.0), vdupq_n_f32(0.0));
1981                    let (mut c1a, mut c1b) = (vdupq_n_f32(0.0), vdupq_n_f32(0.0));
1982                    let (mut c2a, mut c2b) = (vdupq_n_f32(0.0), vdupq_n_f32(0.0));
1983                    let (mut c3a, mut c3b) = (vdupq_n_f32(0.0), vdupq_n_f32(0.0));
1984                    for p in 0..k {
1985                        let (b0, b1) = if b_rows_are_n {
1986                            // B is [n, k]: column p of Bᵀ = element p of
1987                            // eight consecutive B rows — gathered.
1988                            let base = b_mat.as_ptr().add(j * ldb + p);
1989                            let g = |o: usize| *base.add(o * ldb);
1990                            ([g(0), g(1), g(2), g(3)], [g(4), g(5), g(6), g(7)])
1991                        } else {
1992                            let base = b_mat.as_ptr().add(p * ldb + j);
1993                            (
1994                                [*base, *base.add(1), *base.add(2), *base.add(3)],
1995                                [*base.add(4), *base.add(5), *base.add(6), *base.add(7)],
1996                            )
1997                        };
1998                        let bv0 = vld1q_f32(b0.as_ptr());
1999                        let bv1 = vld1q_f32(b1.as_ptr());
2000                        let a0 = vdupq_n_f32(*a.as_ptr().add(i * lda + p));
2001                        let a1 = vdupq_n_f32(*a.as_ptr().add((i + 1) * lda + p));
2002                        let a2 = vdupq_n_f32(*a.as_ptr().add((i + 2) * lda + p));
2003                        let a3 = vdupq_n_f32(*a.as_ptr().add((i + 3) * lda + p));
2004                        c0a = vfmaq_f32(c0a, a0, bv0);
2005                        c0b = vfmaq_f32(c0b, a0, bv1);
2006                        c1a = vfmaq_f32(c1a, a1, bv0);
2007                        c1b = vfmaq_f32(c1b, a1, bv1);
2008                        c2a = vfmaq_f32(c2a, a2, bv0);
2009                        c2b = vfmaq_f32(c2b, a2, bv1);
2010                        c3a = vfmaq_f32(c3a, a3, bv0);
2011                        c3b = vfmaq_f32(c3b, a3, bv1);
2012                    }
2013                    let al = vdupq_n_f32(alpha);
2014                    for (r, (ca, cb)) in [(c0a, c0b), (c1a, c1b), (c2a, c2b), (c3a, c3b)]
2015                        .iter()
2016                        .enumerate()
2017                    {
2018                        let dst = c.as_mut_ptr().add((i + r) * ldc + j);
2019                        vst1q_f32(dst, vmulq_f32(*ca, al));
2020                        vst1q_f32(dst.add(4), vmulq_f32(*cb, al));
2021                    }
2022                } else {
2023                    for r in 0..mi {
2024                        for q in 0..nj {
2025                            let mut acc = 0f32;
2026                            for p in 0..k {
2027                                let bv = if b_rows_are_n {
2028                                    b_mat[(j + q) * ldb + p]
2029                                } else {
2030                                    b_mat[p * ldb + j + q]
2031                                };
2032                                acc += a[(i + r) * lda + p] * bv;
2033                            }
2034                            c[(i + r) * ldc + j + q] = acc * alpha;
2035                        }
2036                    }
2037                }
2038                j += nj;
2039            }
2040            i += mi;
2041        }
2042    }
2043}
2044
2045/// Off-macOS aarch64: the batched attention rides the NEON micro-GEMM.
2046#[cfg(all(target_arch = "aarch64", not(target_os = "macos")))]
2047#[allow(clippy::too_many_arguments)]
2048pub(crate) fn sgemm_rm(
2049    m: usize,
2050    n: usize,
2051    k: usize,
2052    alpha: f32,
2053    a: &[f32],
2054    lda: usize,
2055    b_mat: &[f32],
2056    ldb: usize,
2057    b_rows_are_n: bool,
2058    c: &mut [f32],
2059    ldc: usize,
2060) {
2061    neon_gemm_rm(m, n, k, alpha, a, lda, b_mat, ldb, b_rows_are_n, c, ldc);
2062}
2063
2064/// Row-major f32 GEMM on Accelerate: C[m,n] = alpha·A[m,k] × B(ᵀ).
2065/// `b_rows_are_n` = true multiplies by Bᵀ where B is stored [n, k].
2066#[cfg(target_os = "macos")]
2067#[allow(clippy::too_many_arguments)]
2068pub(crate) fn sgemm_rm(
2069    m: usize,
2070    n: usize,
2071    k: usize,
2072    alpha: f32,
2073    a: &[f32],
2074    lda: usize,
2075    b_mat: &[f32],
2076    ldb: usize,
2077    b_rows_are_n: bool,
2078    c: &mut [f32],
2079    ldc: usize,
2080) {
2081    debug_assert!(a.len() >= (m - 1) * lda + k);
2082    debug_assert!(c.len() >= (m - 1) * ldc + n);
2083    // Test hook: route the attention GEMMs through the portable NEON
2084    // micro-kernel ON APPLE SILICON — how the mobile batched attend is
2085    // measured without a phone in the loop. (Intel macOS has no NEON —
2086    // the hook is a no-op there, Accelerate continues below.)
2087    #[cfg(target_arch = "aarch64")]
2088    if std::env::var("CMF_FORCE_NEON_GEMM")
2089        .map(|v| v == "1")
2090        .unwrap_or(false)
2091    {
2092        return neon_gemm_rm(m, n, k, alpha, a, lda, b_mat, ldb, b_rows_are_n, c, ldc);
2093    }
2094    unsafe {
2095        accel_blas::cblas_sgemm(
2096            101, // RowMajor
2097            111, // NoTrans A
2098            if b_rows_are_n { 112 } else { 111 },
2099            m as i32,
2100            n as i32,
2101            k as i32,
2102            alpha,
2103            a.as_ptr(),
2104            lda as i32,
2105            b_mat.as_ptr(),
2106            ldb as i32,
2107            0.0,
2108            c.as_mut_ptr(),
2109            ldc as i32,
2110        );
2111    }
2112}
2113
2114/// Prefill GEMM through Accelerate (macOS): dequantize q8 rows into
2115/// f32 tiles (scale folded in, pool-parallel) and multiply each tile
2116/// on the AMX with one row-major sgemm. Tiles live in cache, weights
2117/// stream once. Numerics are f32-GEMM (not the int8 dot): prefill
2118/// logits shift within f32 rounding — tolerance-class, like every
2119/// reduction-order change; decode (M=1) never takes this path.
2120#[cfg(target_os = "macos")]
2121fn qmatmat_accel(
2122    q: &[u8],
2123    row_scale: &[f32],
2124    pre: &[std::borrow::Cow<'_, [f32]>],
2125    rows: usize,
2126    cols: usize,
2127    out: &mut [f32],
2128    pool: Option<&Pool>,
2129) {
2130    // NOTE: double-buffering the dequant against the sgemm (a scoped
2131    // thread driving the pool on tile k+1 while the caller multiplies
2132    // tile k) was tried and LOST ~6%: Accelerate's sgemm is itself
2133    // multithreaded, and the dequant workers just steal its cores.
2134    const TR: usize = 2048;
2135    let b = pre.len();
2136    thread_local! {
2137        static XPANEL: std::cell::RefCell<Vec<f32>> = const { std::cell::RefCell::new(Vec::new()) };
2138        static WTILE: std::cell::RefCell<Vec<f32>> = const { std::cell::RefCell::new(Vec::new()) };
2139    }
2140    XPANEL.with(|xp| {
2141        WTILE.with(|wt| {
2142            let mut xpanel = xp.borrow_mut();
2143            xpanel.clear();
2144            for x in pre {
2145                xpanel.extend_from_slice(x);
2146            }
2147            let mut wtile = wt.borrow_mut();
2148            wtile.resize(TR * cols, 0.0);
2149            let mut r0 = 0usize;
2150            while r0 < rows {
2151                let tr = TR.min(rows - r0);
2152                // Dequant the tile (scale folded) — pool-parallel.
2153                let wt_addr = SendMut(wtile.as_mut_ptr());
2154                let run = |start: usize, end: usize| {
2155                    for r in start..end {
2156                        let row = &q[(r0 + r) * cols..(r0 + r + 1) * cols];
2157                        let s = row_scale[r0 + r];
2158                        // SAFETY: workers cover disjoint r ranges.
2159                        let dst =
2160                            unsafe { std::slice::from_raw_parts_mut(wt_addr.at(r * cols), cols) };
2161                        for (d, &v) in dst.iter_mut().zip(row) {
2162                            *d = (v as i8) as f32 * s;
2163                        }
2164                    }
2165                };
2166                dispatch_rows(pool, tr, &run);
2167                // C[b, tr] (at column r0 of out[b, rows]) = X · Wtileᵀ
2168                unsafe {
2169                    accel_blas::cblas_sgemm(
2170                        101, // RowMajor
2171                        111, // NoTrans A
2172                        112, // Trans B
2173                        b as i32,
2174                        tr as i32,
2175                        cols as i32,
2176                        1.0,
2177                        xpanel.as_ptr(),
2178                        cols as i32,
2179                        wtile.as_ptr(),
2180                        cols as i32,
2181                        0.0,
2182                        out.as_mut_ptr().add(r0),
2183                        rows as i32,
2184                    );
2185                }
2186                r0 += tr;
2187            }
2188        })
2189    });
2190}
2191
2192fn qmatmat(
2193    q: &[u8],
2194    row_scale: &[f32],
2195    pre: &[std::borrow::Cow<'_, [f32]>],
2196    rows: usize,
2197    cols: usize,
2198    out: &mut [f32],
2199    pool: Option<&Pool>,
2200) {
2201    let b = pre.len();
2202    debug_assert_eq!(out.len(), b * rows);
2203    // Big prefill batches ride the AMX (roadmap PR3): the row×batch
2204    // SDOT loop below peaks near the CPU's dot throughput, an order
2205    // below the matrix units. Small tensors and tiny test models stay
2206    // on the exact integer path.
2207    #[cfg(target_os = "macos")]
2208    if b >= 8 && rows * cols >= 500_000 && accel_gemm_enabled() {
2209        qmatmat_accel(q, row_scale, pre, rows, cols, out, pool);
2210        return;
2211    }
2212    #[cfg(target_arch = "aarch64")]
2213    if sdot_enabled() {
2214        let acts: Vec<SplitAct> = pre.iter().map(|x| split_act(x)).collect();
2215        let out_addr = SendMut(out.as_mut_ptr());
2216        // Blocked 2×4 (mobile prefill: no AMX to fall back on — this
2217        // path IS the ARM prefill GEMM off Apple silicon).
2218        let blocked_ok = std::env::var("CMF_X86_BLOCKED")
2219            .map(|v| v != "0")
2220            .unwrap_or(true);
2221        let use_i8mm = i8mm_enabled();
2222        if blocked_ok {
2223            let run = |start: usize, end: usize| {
2224                let mut o = start;
2225                while o < end {
2226                    if o + 2 <= end {
2227                        let r0 = &q[o * cols..(o + 1) * cols];
2228                        let r1 = &q[(o + 1) * cols..(o + 2) * cols];
2229                        let mut bi = 0usize;
2230                        while bi + 4 <= acts.len() {
2231                            let xs = [
2232                                acts[bi].xq.as_slice(),
2233                                acts[bi + 1].xq.as_slice(),
2234                                acts[bi + 2].xq.as_slice(),
2235                                acts[bi + 3].xq.as_slice(),
2236                            ];
2237                            let d = if use_i8mm {
2238                                unsafe { dot_i8_smmla_2x4(r0, r1, xs) }
2239                            } else {
2240                                unsafe { dot_i8_sdot_2x4(r0, r1, xs) }
2241                            };
2242                            for (r, row) in [r0, r1].into_iter().enumerate() {
2243                                for k in 0..4 {
2244                                    let act = &acts[bi + k];
2245                                    let mut v = d[r][k] as f32 * act.sx;
2246                                    for &(j, xv) in &act.outliers {
2247                                        v += (row[j] as i8) as f32 * xv;
2248                                    }
2249                                    unsafe {
2250                                        *out_addr.at((bi + k) * rows + o + r) = v * row_scale[o + r]
2251                                    };
2252                                }
2253                            }
2254                            bi += 4;
2255                        }
2256                        while bi < acts.len() {
2257                            for (r, row) in [r0, r1].into_iter().enumerate() {
2258                                let v = row_dot_sdot(row, &acts[bi]) * row_scale[o + r];
2259                                unsafe { *out_addr.at(bi * rows + o + r) = v };
2260                            }
2261                            bi += 1;
2262                        }
2263                        o += 2;
2264                    } else {
2265                        let row = &q[o * cols..(o + 1) * cols];
2266                        for (bi, act) in acts.iter().enumerate() {
2267                            let v = row_dot_sdot(row, act) * row_scale[o];
2268                            unsafe { *out_addr.at(bi * rows + o) = v };
2269                        }
2270                        o += 1;
2271                    }
2272                }
2273            };
2274            dispatch_rows(pool, rows, &run);
2275            return;
2276        }
2277        let run = |start: usize, end: usize| {
2278            for o in start..end {
2279                let row = &q[o * cols..(o + 1) * cols];
2280                for (bi, act) in acts.iter().enumerate() {
2281                    let v = row_dot_sdot(row, act) * row_scale[o];
2282                    unsafe { *out_addr.at(bi * rows + o) = v };
2283                }
2284            }
2285        };
2286        dispatch_rows(pool, rows, &run);
2287        return;
2288    }
2289    // x86 A8W8 batch. Non-VNNI parts take the BLOCKED 2×4 kernel
2290    // (roadmap P0: two weight rows' abs() stay in registers across four
2291    // activation streams); VNNI machines keep the per-row bias-trick
2292    // dot, which is already throughput-bound there.
2293    #[cfg(target_arch = "x86_64")]
2294    if avx2_a8w8_enabled() {
2295        let acts: Vec<SplitAct> = pre.iter().map(|x| split_act(x)).collect();
2296        let out_addr = SendMut(out.as_mut_ptr());
2297        // CMF_X86_BLOCKED=0 forces the per-row path (paired in-process
2298        // A/B on noisy shared-vCPU hosts).
2299        let blocked_ok = std::env::var("CMF_X86_BLOCKED")
2300            .map(|v| v != "0")
2301            .unwrap_or(true);
2302        if !avx512vnni_enabled() && blocked_ok {
2303            let run = |start: usize, end: usize| {
2304                let mut o = start;
2305                while o < end {
2306                    if o + 2 <= end {
2307                        let r0 = &q[o * cols..(o + 1) * cols];
2308                        let r1 = &q[(o + 1) * cols..(o + 2) * cols];
2309                        let mut bi = 0usize;
2310                        while bi + 4 <= acts.len() {
2311                            let xs = [
2312                                acts[bi].xq.as_slice(),
2313                                acts[bi + 1].xq.as_slice(),
2314                                acts[bi + 2].xq.as_slice(),
2315                                acts[bi + 3].xq.as_slice(),
2316                            ];
2317                            let d = unsafe { dot_i8_i8_avx2_2x4(r0, r1, xs) };
2318                            for (r, row) in [r0, r1].into_iter().enumerate() {
2319                                for k in 0..4 {
2320                                    let act = &acts[bi + k];
2321                                    let mut v = d[r][k] as f32 * act.sx;
2322                                    for &(j, xv) in &act.outliers {
2323                                        v += (row[j] as i8) as f32 * xv;
2324                                    }
2325                                    unsafe {
2326                                        *out_addr.at((bi + k) * rows + o + r) = v * row_scale[o + r]
2327                                    };
2328                                }
2329                            }
2330                            bi += 4;
2331                        }
2332                        while bi < acts.len() {
2333                            for (r, row) in [r0, r1].into_iter().enumerate() {
2334                                let v = row_dot_avx2(row, &acts[bi]) * row_scale[o + r];
2335                                unsafe { *out_addr.at(bi * rows + o + r) = v };
2336                            }
2337                            bi += 1;
2338                        }
2339                        o += 2;
2340                    } else {
2341                        let row = &q[o * cols..(o + 1) * cols];
2342                        for (bi, act) in acts.iter().enumerate() {
2343                            let v = row_dot_avx2(row, act) * row_scale[o];
2344                            unsafe { *out_addr.at(bi * rows + o) = v };
2345                        }
2346                        o += 1;
2347                    }
2348                }
2349            };
2350            dispatch_rows(pool, rows, &run);
2351            return;
2352        }
2353        let run = |start: usize, end: usize| {
2354            for o in start..end {
2355                let row = &q[o * cols..(o + 1) * cols];
2356                for (bi, act) in acts.iter().enumerate() {
2357                    let v = row_dot_avx2(row, act) * row_scale[o];
2358                    unsafe { *out_addr.at(bi * rows + o) = v };
2359                }
2360            }
2361        };
2362        dispatch_rows(pool, rows, &run);
2363        return;
2364    }
2365    let out_addr = SendMut(out.as_mut_ptr());
2366    let run = |start: usize, end: usize| {
2367        for o in start..end {
2368            let row = &q[o * cols..(o + 1) * cols];
2369            for (bi, x) in pre.iter().enumerate() {
2370                let mut acc = 0f32;
2371                for j in 0..cols {
2372                    acc += (row[j] as i8) as f32 * x[j];
2373                }
2374                unsafe { *out_addr.at(bi * rows + o) = acc * row_scale[o] };
2375            }
2376        }
2377    };
2378    dispatch_rows(pool, rows, &run);
2379}
2380
2381/// Split rows across pool workers (shared qmatvec pattern). Self-balancing
2382/// — see `Pool::run_rows` for why a static 1/n split is wrong here.
2383fn dispatch_rows(pool: Option<&Pool>, rows: usize, run: &(dyn Fn(usize, usize) + Sync)) {
2384    match pool {
2385        Some(pool) if rows >= 256 => pool.run_rows(rows, run),
2386        _ => run(0, rows),
2387    }
2388}
2389
2390/// Split a q4_block blob into (packed nibbles, f16 group scales).
2391fn q4_split(bytes: &[u8], rows: usize, cols: usize) -> (&[u8], &[u8]) {
2392    let groups = rows * cols / GROUP_SIZE;
2393    bytes.split_at(groups * 16)
2394}
2395
2396/// SIMD unpack for the dominant vbit width B=4 (94% of rows on the
2397/// log2-shape calibration): 16 packed bytes -> 32 centered i8 values.
2398/// vbit packs MSB-first, so the HIGH nibble is the even element
2399/// (opposite of q4_block's lo-first interleave). Centering is u-7.
2400#[inline]
2401fn vbit_fill4(data: &[u8], buf: &mut [u8]) {
2402    #[cfg(target_arch = "aarch64")]
2403    unsafe {
2404        return vbit_fill4_neon(data, buf);
2405    }
2406    #[cfg(target_arch = "x86_64")]
2407    if avx2_enabled() {
2408        return unsafe { vbit_fill4_avx2(data, buf) };
2409    }
2410    #[allow(unreachable_code)]
2411    for (blk, chunk) in buf.chunks_exact_mut(8).enumerate() {
2412        let u = unpack8::<4>(&data[blk * 4..]);
2413        for k in 0..8 {
2414            chunk[k] = (u[k] - 7) as i8 as u8;
2415        }
2416    }
2417}
2418
2419#[cfg(target_arch = "aarch64")]
2420#[target_feature(enable = "neon")]
2421unsafe fn vbit_fill4_neon(data: &[u8], buf: &mut [u8]) {
2422    // SAFETY: buf.len() is a multiple of GROUP_SIZE=32; data holds
2423    // buf.len()/2 packed bytes (validated at load).
2424    unsafe {
2425        use core::arch::aarch64::*;
2426        let n = buf.len();
2427        let mask = vdupq_n_u8(0x0F);
2428        let seven = vdupq_n_s8(7);
2429        let mut g = 0usize;
2430        while g * 32 + 32 <= n {
2431            let b = vld1q_u8(data.as_ptr().add(g * 16));
2432            let hi = vshrq_n_u8::<4>(b);
2433            let lo = vandq_u8(b, mask);
2434            let z0 = vsubq_s8(vreinterpretq_s8_u8(vzip1q_u8(hi, lo)), seven);
2435            let z1 = vsubq_s8(vreinterpretq_s8_u8(vzip2q_u8(hi, lo)), seven);
2436            vst1q_u8(buf.as_mut_ptr().add(g * 32), vreinterpretq_u8_s8(z0));
2437            vst1q_u8(buf.as_mut_ptr().add(g * 32 + 16), vreinterpretq_u8_s8(z1));
2438            g += 1;
2439        }
2440    }
2441}
2442
2443#[cfg(target_arch = "x86_64")]
2444#[target_feature(enable = "avx2")]
2445unsafe fn vbit_fill4_avx2(data: &[u8], buf: &mut [u8]) {
2446    // SAFETY: see vbit_fill4_neon.
2447    unsafe {
2448        use core::arch::x86_64::*;
2449        let n = buf.len();
2450        let mask = _mm_set1_epi8(0x0F);
2451        let seven = _mm256_set1_epi8(7);
2452        let mut g = 0usize;
2453        while g * 32 + 32 <= n {
2454            let b = _mm_loadu_si128(data.as_ptr().add(g * 16) as *const __m128i);
2455            let hi = _mm_and_si128(_mm_srli_epi16::<4>(b), mask);
2456            let lo = _mm_and_si128(b, mask);
2457            let z = _mm256_sub_epi8(
2458                _mm256_set_m128i(_mm_unpackhi_epi8(hi, lo), _mm_unpacklo_epi8(hi, lo)),
2459                seven,
2460            );
2461            _mm256_storeu_si256(buf.as_mut_ptr().add(g * 32) as *mut __m256i, z);
2462            g += 1;
2463        }
2464    }
2465}
2466
2467/// Unpack 8 MSB-first B-bit values from exactly B bytes (fixed shifts —
2468/// no serial bit-buffer, auto-vectorizable). Every 32-value group starts
2469/// byte-aligned (32·B/8 is integral for B∈3..8), so groups decompose
2470/// into 4 such blocks.
2471#[inline(always)]
2472fn unpack8<const B: usize>(data: &[u8]) -> [i32; 8] {
2473    let mut acc = 0u64;
2474    for i in 0..B {
2475        acc = (acc << 8) | data[i] as u64;
2476    }
2477    let mask = (1u64 << B) - 1;
2478    let mut out = [0i32; 8];
2479    for (k, o) in out.iter_mut().enumerate() {
2480        *o = ((acc >> ((7 - k) * B)) & mask) as i32;
2481    }
2482    out
2483}
2484
2485/// Fused vbit matvec straight from the mapped bytes (spec §3, P13
2486/// FIG.3): [u8 bits: rows][f16 scales: rows·cols/32][bit-packed rows,
2487/// MSB-first, byte-padded]. Row data offsets are precomputed at load
2488/// (`vbit_row_offsets`) — the per-call prefix scan was O(rows) pure
2489/// overhead on every matvec.
2490#[allow(clippy::too_many_arguments)]
2491fn vbitmatvec(
2492    bytes: &[u8],
2493    offsets: &[usize],
2494    x: &[f32],
2495    rows: usize,
2496    cols: usize,
2497    out: &mut [f32],
2498    pool: Option<&Pool>,
2499) {
2500    debug_assert_eq!(out.len(), rows);
2501    debug_assert_eq!(offsets.len(), rows + 1);
2502
2503    // SDOT path: unpack the row to centered i8 once, then per-group
2504    // int8 dot against the quantized activations — same A8W8 contract
2505    // as q8 (bounded noise; CMF_SDOT=0 keeps the exact scalar path).
2506    if a8w8_enabled() {
2507        let act = split_act(x);
2508        let out_addr = SendMut(out.as_mut_ptr());
2509        let run = move |start: usize, end: usize| {
2510            vbit_range_a8w8(bytes, offsets, x, &act, rows, cols, out_addr, start, end)
2511        };
2512        dispatch_rows(pool, rows, &run);
2513        return;
2514    }
2515
2516    let out_addr = SendMut(out.as_mut_ptr());
2517    let run = move |start: usize, end: usize| {
2518        vbit_range_f32(bytes, offsets, x, rows, cols, out_addr, start, end)
2519    };
2520    dispatch_rows(pool, rows, &run);
2521}
2522
2523/// One vbit row range via the A8W8 int8 path — kernel body of
2524/// `vbitmatvec`, extracted so multi-matrix jobs can drive it for
2525/// several tensors in one dispatch (b=8 rows go exact f32).
2526#[allow(clippy::too_many_arguments)]
2527fn vbit_range_a8w8(
2528    bytes: &[u8],
2529    offsets: &[usize],
2530    x: &[f32],
2531    act: &SplitAct,
2532    rows: usize,
2533    cols: usize,
2534    out: SendMut,
2535    start: usize,
2536    end: usize,
2537) {
2538    let ng = cols / GROUP_SIZE;
2539    let bits = &bytes[..rows];
2540    let sc_off = rows;
2541    let row_dot = |r: usize| -> f32 {
2542        let b = bits[r] as usize;
2543        let l = (1i32 << (b - 1)) - 1;
2544        let mask = (1u64 << b) - 1;
2545        let data = &bytes[offsets[r]..offsets[r + 1]];
2546        if b == 8 {
2547            // u−L reaches 128 → does not fit i8; exact f32 path.
2548            let (mut acc, mut nbits, mut idx) = (0u64, 0usize, 0usize);
2549            let mut dot = 0f32;
2550            for g in 0..ng {
2551                let so = (r * ng + g) * 2;
2552                let sgf = f16_to_f32(u16::from_le_bytes([
2553                    bytes[sc_off + so],
2554                    bytes[sc_off + so + 1],
2555                ]));
2556                let xg = &x[g * GROUP_SIZE..(g + 1) * GROUP_SIZE];
2557                let mut gd = 0f32;
2558                for &xv in xg.iter() {
2559                    if nbits < 8 {
2560                        acc = (acc << 8) | data[idx] as u64;
2561                        idx += 1;
2562                        nbits += 8;
2563                    }
2564                    let u = ((acc >> (nbits - 8)) & 0xFF) as i32;
2565                    nbits -= 8;
2566                    gd += (u - l) as f32 * xv;
2567                }
2568                dot += gd * sgf;
2569            }
2570            return dot;
2571        }
2572        // Per-worker scratch: this closure runs for every row of the
2573        // tensor (lm_head ≈ 150k rows/token) — a heap allocation per
2574        // row was measurable pure overhead.
2575        thread_local! {
2576            static VBIT_SCRATCH: std::cell::RefCell<Vec<u8>> =
2577                const { std::cell::RefCell::new(Vec::new()) };
2578        }
2579        #[inline(always)]
2580        fn fill<const B: usize>(data: &[u8], l: i32, buf: &mut [u8]) {
2581            for (blk, chunk) in buf.chunks_exact_mut(8).enumerate() {
2582                let u = unpack8::<B>(&data[blk * B..]);
2583                for k in 0..8 {
2584                    chunk[k] = (u[k] - l) as i8 as u8;
2585                }
2586            }
2587        }
2588        let _ = mask;
2589        VBIT_SCRATCH.with(|scratch| {
2590            let mut buf = scratch.borrow_mut();
2591            buf.resize(cols, 0);
2592            match b {
2593                3 => fill::<3>(data, l, &mut buf),
2594                4 => vbit_fill4(data, &mut buf),
2595                5 => fill::<5>(data, l, &mut buf),
2596                6 => fill::<6>(data, l, &mut buf),
2597                _ => unreachable!(),
2598            }
2599            let mut dot = 0f32;
2600            for g in 0..ng {
2601                let so = (r * ng + g) * 2;
2602                let s = f16_to_f32(u16::from_le_bytes([
2603                    bytes[sc_off + so],
2604                    bytes[sc_off + so + 1],
2605                ]));
2606                let d = dot_i8_i8(
2607                    &buf[g * GROUP_SIZE..(g + 1) * GROUP_SIZE],
2608                    &act.xq[g * GROUP_SIZE..(g + 1) * GROUP_SIZE],
2609                ) as f32
2610                    * act.sx;
2611                dot += d * s;
2612            }
2613            for &(j, xv) in &act.outliers {
2614                let so = (r * ng + j / GROUP_SIZE) * 2;
2615                let s = f16_to_f32(u16::from_le_bytes([
2616                    bytes[sc_off + so],
2617                    bytes[sc_off + so + 1],
2618                ]));
2619                // xq is zeroed at outlier slots — add the exact term.
2620                dot += (buf[j] as i8) as f32 * s * xv;
2621            }
2622            dot
2623        })
2624    };
2625    for r in start..end {
2626        // SAFETY: disjoint row ranges per worker.
2627        unsafe { *out.at(r) = row_dot(r) };
2628    }
2629}
2630
2631/// Exact scalar vbit row range (same extraction, non-SDOT path).
2632#[allow(clippy::too_many_arguments)]
2633fn vbit_range_f32(
2634    bytes: &[u8],
2635    offsets: &[usize],
2636    x: &[f32],
2637    rows: usize,
2638    cols: usize,
2639    out: SendMut,
2640    start: usize,
2641    end: usize,
2642) {
2643    let ng = cols / GROUP_SIZE;
2644    let bits = &bytes[..rows];
2645    let sc_off = rows;
2646    // Per-bit-width specialized inner loops: the compiler unrolls the
2647    // constant shifts (the generic bit-buffer loop was branch-bound —
2648    // 5.6 vs 13.2 tok/s q4 on the 0.8B).
2649    #[inline(always)]
2650    fn dot_row<const B: usize>(
2651        data: &[u8],
2652        bytes: &[u8],
2653        sc_off: usize,
2654        r: usize,
2655        ng: usize,
2656        x: &[f32],
2657    ) -> f32 {
2658        let l = ((1i32 << (B - 1)) - 1) as f32;
2659        let gbytes = GROUP_SIZE * B / 8;
2660        let mut dot = 0f32;
2661        for g in 0..ng {
2662            let so = (r * ng + g) * 2;
2663            let s = f16_to_f32(u16::from_le_bytes([
2664                bytes[sc_off + so],
2665                bytes[sc_off + so + 1],
2666            ]));
2667            let xg = &x[g * GROUP_SIZE..(g + 1) * GROUP_SIZE];
2668            let gd0 = &data[g * gbytes..(g + 1) * gbytes];
2669            let mut gd = 0f32;
2670            for blk in 0..GROUP_SIZE / 8 {
2671                let u = unpack8::<B>(&gd0[blk * B..]);
2672                let xb = &xg[blk * 8..blk * 8 + 8];
2673                for k in 0..8 {
2674                    gd += (u[k] as f32 - l) * xb[k];
2675                }
2676            }
2677            dot += gd * s;
2678        }
2679        dot
2680    }
2681    for r in start..end {
2682        let data = &bytes[offsets[r]..offsets[r + 1]];
2683        let v = match bits[r] {
2684            3 => dot_row::<3>(data, bytes, sc_off, r, ng, x),
2685            4 => dot_row::<4>(data, bytes, sc_off, r, ng, x),
2686            5 => dot_row::<5>(data, bytes, sc_off, r, ng, x),
2687            6 => dot_row::<6>(data, bytes, sc_off, r, ng, x),
2688            8 => dot_row::<8>(data, bytes, sc_off, r, ng, x),
2689            b => unreachable!("vbit bit-width {b} (validated at load)"),
2690        };
2691        // SAFETY: disjoint row ranges per worker.
2692        unsafe { *out.at(r) = v };
2693    }
2694}
2695
2696/// Fused two-input vbit matvec: each row is unpacked from the mmap ONCE
2697/// and dotted against BOTH activations (MTP verify / pair prefill used
2698/// to run two full matvecs — double weight traffic and double unpack).
2699/// Per-input math is identical to `vbitmatvec` → same accuracy contract.
2700#[allow(clippy::too_many_arguments)]
2701fn vbitmatvec2(
2702    bytes: &[u8],
2703    offsets: &[usize],
2704    x1: &[f32],
2705    x2: &[f32],
2706    rows: usize,
2707    cols: usize,
2708    o1: &mut [f32],
2709    o2: &mut [f32],
2710    pool: Option<&Pool>,
2711) {
2712    debug_assert_eq!(o1.len(), rows);
2713    debug_assert_eq!(o2.len(), rows);
2714
2715    if a8w8_enabled() {
2716        let a1 = split_act(x1);
2717        let a2 = split_act(x2);
2718        let p1 = SendMut(o1.as_mut_ptr());
2719        let p2 = SendMut(o2.as_mut_ptr());
2720        let run = move |start: usize, end: usize| {
2721            vbit_range2_a8w8(
2722                bytes, offsets, x1, x2, &a1, &a2, rows, cols, p1, p2, start, end,
2723            )
2724        };
2725        dispatch_rows(pool, rows, &run);
2726        return;
2727    }
2728
2729    let p1 = SendMut(o1.as_mut_ptr());
2730    let p2 = SendMut(o2.as_mut_ptr());
2731    let run = move |start: usize, end: usize| {
2732        vbit_range2_f32(bytes, offsets, x1, x2, rows, cols, p1, p2, start, end)
2733    };
2734    dispatch_rows(pool, rows, &run);
2735}
2736
2737/// Two-input vbit row range via the A8W8 int8 path — kernel body of
2738/// `vbitmatvec2`, extracted for pair multi-matrix jobs (b=8 rows go
2739/// exact f32 for both lanes, bits streamed once).
2740#[allow(clippy::too_many_arguments)]
2741fn vbit_range2_a8w8(
2742    bytes: &[u8],
2743    offsets: &[usize],
2744    x1: &[f32],
2745    x2: &[f32],
2746    a1: &SplitAct,
2747    a2: &SplitAct,
2748    rows: usize,
2749    cols: usize,
2750    p1: SendMut,
2751    p2: SendMut,
2752    start: usize,
2753    end: usize,
2754) {
2755    let ng = cols / GROUP_SIZE;
2756    let bits = &bytes[..rows];
2757    let sc_off = rows;
2758    let row_dots = |r: usize| -> (f32, f32) {
2759        let b = bits[r] as usize;
2760        let l = (1i32 << (b - 1)) - 1;
2761        let data = &bytes[offsets[r]..offsets[r + 1]];
2762        if b == 8 {
2763            // u−L reaches 128 → does not fit i8; exact f32 path,
2764            // bits still streamed once for both lanes.
2765            let (mut acc, mut nbits, mut idx) = (0u64, 0usize, 0usize);
2766            let (mut d1, mut d2) = (0f32, 0f32);
2767            for g in 0..ng {
2768                let so = (r * ng + g) * 2;
2769                let sgf = f16_to_f32(u16::from_le_bytes([
2770                    bytes[sc_off + so],
2771                    bytes[sc_off + so + 1],
2772                ]));
2773                let (mut g1, mut g2) = (0f32, 0f32);
2774                for k in 0..GROUP_SIZE {
2775                    if nbits < 8 {
2776                        acc = (acc << 8) | data[idx] as u64;
2777                        idx += 1;
2778                        nbits += 8;
2779                    }
2780                    let u = ((acc >> (nbits - 8)) & 0xFF) as i32;
2781                    nbits -= 8;
2782                    let w = (u - l) as f32;
2783                    g1 += w * x1[g * GROUP_SIZE + k];
2784                    g2 += w * x2[g * GROUP_SIZE + k];
2785                }
2786                d1 += g1 * sgf;
2787                d2 += g2 * sgf;
2788            }
2789            return (d1, d2);
2790        }
2791        thread_local! {
2792            static VBIT_SCRATCH2: std::cell::RefCell<Vec<u8>> =
2793                const { std::cell::RefCell::new(Vec::new()) };
2794        }
2795        #[inline(always)]
2796        fn fill<const B: usize>(data: &[u8], l: i32, buf: &mut [u8]) {
2797            for (blk, chunk) in buf.chunks_exact_mut(8).enumerate() {
2798                let u = unpack8::<B>(&data[blk * B..]);
2799                for k in 0..8 {
2800                    chunk[k] = (u[k] - l) as i8 as u8;
2801                }
2802            }
2803        }
2804        VBIT_SCRATCH2.with(|scratch| {
2805            let mut buf = scratch.borrow_mut();
2806            buf.resize(cols, 0);
2807            match b {
2808                3 => fill::<3>(data, l, &mut buf),
2809                4 => vbit_fill4(data, &mut buf),
2810                5 => fill::<5>(data, l, &mut buf),
2811                6 => fill::<6>(data, l, &mut buf),
2812                _ => unreachable!(),
2813            }
2814            let (mut d1, mut d2) = (0f32, 0f32);
2815            for g in 0..ng {
2816                let so = (r * ng + g) * 2;
2817                let s = f16_to_f32(u16::from_le_bytes([
2818                    bytes[sc_off + so],
2819                    bytes[sc_off + so + 1],
2820                ]));
2821                let wg = &buf[g * GROUP_SIZE..(g + 1) * GROUP_SIZE];
2822                let v1 = dot_i8_i8(wg, &a1.xq[g * GROUP_SIZE..(g + 1) * GROUP_SIZE]) as f32 * a1.sx;
2823                let v2 = dot_i8_i8(wg, &a2.xq[g * GROUP_SIZE..(g + 1) * GROUP_SIZE]) as f32 * a2.sx;
2824                d1 += v1 * s;
2825                d2 += v2 * s;
2826            }
2827            for &(j, xv) in &a1.outliers {
2828                let so = (r * ng + j / GROUP_SIZE) * 2;
2829                let s = f16_to_f32(u16::from_le_bytes([
2830                    bytes[sc_off + so],
2831                    bytes[sc_off + so + 1],
2832                ]));
2833                d1 += (buf[j] as i8) as f32 * s * xv;
2834            }
2835            for &(j, xv) in &a2.outliers {
2836                let so = (r * ng + j / GROUP_SIZE) * 2;
2837                let s = f16_to_f32(u16::from_le_bytes([
2838                    bytes[sc_off + so],
2839                    bytes[sc_off + so + 1],
2840                ]));
2841                d2 += (buf[j] as i8) as f32 * s * xv;
2842            }
2843            (d1, d2)
2844        })
2845    };
2846    for r in start..end {
2847        let (v1, v2) = row_dots(r);
2848        // SAFETY: disjoint row ranges per worker.
2849        unsafe {
2850            *p1.at(r) = v1;
2851            *p2.at(r) = v2;
2852        }
2853    }
2854}
2855
2856/// Two-input exact scalar vbit row range (same extraction) —
2857/// per-bit-width specialized, two accumulators per row; per-lane
2858/// accumulation order matches `vbitmatvec` exactly.
2859#[allow(clippy::too_many_arguments)]
2860fn vbit_range2_f32(
2861    bytes: &[u8],
2862    offsets: &[usize],
2863    x1: &[f32],
2864    x2: &[f32],
2865    rows: usize,
2866    cols: usize,
2867    p1: SendMut,
2868    p2: SendMut,
2869    start: usize,
2870    end: usize,
2871) {
2872    let ng = cols / GROUP_SIZE;
2873    let bits = &bytes[..rows];
2874    let sc_off = rows;
2875    #[inline(always)]
2876    #[allow(clippy::too_many_arguments)]
2877    fn dot_row2<const B: usize>(
2878        data: &[u8],
2879        bytes: &[u8],
2880        sc_off: usize,
2881        r: usize,
2882        ng: usize,
2883        x1: &[f32],
2884        x2: &[f32],
2885    ) -> (f32, f32) {
2886        let l = ((1i32 << (B - 1)) - 1) as f32;
2887        let gbytes = GROUP_SIZE * B / 8;
2888        let (mut d1, mut d2) = (0f32, 0f32);
2889        for g in 0..ng {
2890            let so = (r * ng + g) * 2;
2891            let s = f16_to_f32(u16::from_le_bytes([
2892                bytes[sc_off + so],
2893                bytes[sc_off + so + 1],
2894            ]));
2895            let x1g = &x1[g * GROUP_SIZE..(g + 1) * GROUP_SIZE];
2896            let x2g = &x2[g * GROUP_SIZE..(g + 1) * GROUP_SIZE];
2897            let gd0 = &data[g * gbytes..(g + 1) * gbytes];
2898            let (mut g1, mut g2) = (0f32, 0f32);
2899            for blk in 0..GROUP_SIZE / 8 {
2900                let u = unpack8::<B>(&gd0[blk * B..]);
2901                for k in 0..8 {
2902                    let w = u[k] as f32 - l;
2903                    g1 += w * x1g[blk * 8 + k];
2904                    g2 += w * x2g[blk * 8 + k];
2905                }
2906            }
2907            d1 += g1 * s;
2908            d2 += g2 * s;
2909        }
2910        (d1, d2)
2911    }
2912    for r in start..end {
2913        let data = &bytes[offsets[r]..offsets[r + 1]];
2914        let (v1, v2) = match bits[r] {
2915            3 => dot_row2::<3>(data, bytes, sc_off, r, ng, x1, x2),
2916            4 => dot_row2::<4>(data, bytes, sc_off, r, ng, x1, x2),
2917            5 => dot_row2::<5>(data, bytes, sc_off, r, ng, x1, x2),
2918            6 => dot_row2::<6>(data, bytes, sc_off, r, ng, x1, x2),
2919            8 => dot_row2::<8>(data, bytes, sc_off, r, ng, x1, x2),
2920            b => unreachable!("vbit bit-width {b} (validated at load)"),
2921        };
2922        // SAFETY: disjoint row ranges per worker.
2923        unsafe {
2924            *p1.at(r) = v1;
2925            *p2.at(r) = v2;
2926        }
2927    }
2928}
2929
2930// ───────────────────── q4_tiled kernels (§4.3) ─────────────────────
2931
2932/// One q4_tiled row dot on the A8W8 int8 path: per 32-group the tile
2933/// is ONE sequential read — [f16 scale][16B nibbles] — versus the two
2934/// distant streams of the split layout. Values/order identical to the
2935/// split kernels.
2936#[inline]
2937#[allow(unreachable_code)]
2938fn dot_q4t_row_i8(bytes: &[u8], r: usize, gpr: usize, xq: &[i8]) -> f32 {
2939    #[cfg(target_arch = "aarch64")]
2940    unsafe {
2941        return dot_q4t_row_sdot(bytes, r, gpr, xq);
2942    }
2943    #[cfg(target_arch = "x86_64")]
2944    unsafe {
2945        if vnni_tiles_enabled() {
2946            return dot_q4t_row_vnni(bytes, r, gpr, xq);
2947        }
2948        return dot_q4t_row_avx2(bytes, r, gpr, xq);
2949    }
2950    let mut acc = 0f32;
2951    for gi in 0..gpr {
2952        let tile = &bytes[(r * gpr + gi) * Q4_TILE..(r * gpr + gi + 1) * Q4_TILE];
2953        let s = f16_to_f32(u16::from_le_bytes([tile[0], tile[1]]));
2954        let mut d = 0i32;
2955        for (k, &b) in tile[2..].iter().enumerate() {
2956            d += ((b & 0x0F) as i32 - 8) * xq[gi * GROUP_SIZE + k * 2] as i32
2957                + (((b >> 4) & 0x0F) as i32 - 8) * xq[gi * GROUP_SIZE + k * 2 + 1] as i32;
2958        }
2959        acc += d as f32 * s;
2960    }
2961    acc
2962}
2963
2964#[cfg(target_arch = "aarch64")]
2965#[target_feature(enable = "neon,dotprod")]
2966unsafe fn dot_q4t_row_sdot(bytes: &[u8], r: usize, gpr: usize, xq: &[i8]) -> f32 {
2967    // SAFETY: callers uphold slice-length contracts (18B tile per group,
2968    // xq.len() == gpr·GROUP_SIZE).
2969    unsafe {
2970        use core::arch::aarch64::*;
2971        use core::arch::asm;
2972        let lomask = vdupq_n_u8(0x0F);
2973        let eight = vdupq_n_s8(8);
2974        let mut acc = 0f32;
2975        for gi in 0..gpr {
2976            let t = bytes.as_ptr().add((r * gpr + gi) * Q4_TILE);
2977            let s = f16_to_f32(u16::from_le_bytes([*t, *t.add(1)]));
2978            let b = vld1q_u8(t.add(2));
2979            let lo = vandq_u8(b, lomask);
2980            let hi = vshrq_n_u8::<4>(b);
2981            let e0 = vsubq_s8(vreinterpretq_s8_u8(vzip1q_u8(lo, hi)), eight);
2982            let e1 = vsubq_s8(vreinterpretq_s8_u8(vzip2q_u8(lo, hi)), eight);
2983            let x0 = vld1q_s8(xq.as_ptr().add(gi * GROUP_SIZE));
2984            let x1 = vld1q_s8(xq.as_ptr().add(gi * GROUP_SIZE + 16));
2985            let (mut a0, mut a1) = (vdupq_n_s32(0), vdupq_n_s32(0));
2986            asm!(
2987                "sdot {a0:v}.4s, {e0:v}.16b, {x0:v}.16b",
2988                "sdot {a1:v}.4s, {e1:v}.16b, {x1:v}.16b",
2989                a0 = inout(vreg) a0, a1 = inout(vreg) a1,
2990                e0 = in(vreg) e0, x0 = in(vreg) x0, e1 = in(vreg) e1, x1 = in(vreg) x1,
2991                options(pure, nomem, nostack),
2992            );
2993            acc += vaddvq_s32(vaddq_s32(a0, a1)) as f32 * s;
2994        }
2995        acc
2996    }
2997}
2998
2999#[cfg(target_arch = "x86_64")]
3000#[target_feature(enable = "avx2")]
3001unsafe fn dot_q4t_row_avx2(bytes: &[u8], r: usize, gpr: usize, xq: &[i8]) -> f32 {
3002    // SAFETY: see dot_q4t_row_sdot.
3003    unsafe {
3004        use core::arch::x86_64::*;
3005        let lomask = _mm_set1_epi8(0x0F);
3006        let eight = _mm256_set1_epi8(8);
3007        let ones = _mm256_set1_epi16(1);
3008        let mut acc = 0f32;
3009        for gi in 0..gpr {
3010            let t = bytes.as_ptr().add((r * gpr + gi) * Q4_TILE);
3011            let s = f16_to_f32(u16::from_le_bytes([*t, *t.add(1)]));
3012            let b = _mm_loadu_si128(t.add(2) as *const __m128i);
3013            let lo = _mm_and_si128(b, lomask);
3014            let hi = _mm_and_si128(_mm_srli_epi16::<4>(b), lomask);
3015            let w = _mm256_sub_epi8(
3016                _mm256_set_m128i(_mm_unpackhi_epi8(lo, hi), _mm_unpacklo_epi8(lo, hi)),
3017                eight,
3018            );
3019            let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
3020            let p16 = _mm256_maddubs_epi16(_mm256_abs_epi8(w), _mm256_sign_epi8(x, w));
3021            let d = _mm256_madd_epi16(p16, ones);
3022            let hi128 = _mm256_extracti128_si256::<1>(d);
3023            let s128 = _mm_add_epi32(_mm256_castsi256_si128(d), hi128);
3024            let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
3025            let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
3026            acc += _mm_cvtsi128_si32(s32) as f32 * s;
3027        }
3028        acc
3029    }
3030}
3031
3032/// VNNI twin of `dot_q4t_row_avx2`: same unpack, `vpdpbusd` replaces
3033/// the maddubs+madd pair (see `dpbusd_hsum` — sums are bit-identical).
3034/// 256-bit VL encoding, so the VEX `vpsignb` stays usable.
3035#[cfg(target_arch = "x86_64")]
3036#[target_feature(enable = "avx2,avx512f,avx512bw,avx512vl,avx512vnni")]
3037unsafe fn dot_q4t_row_vnni(bytes: &[u8], r: usize, gpr: usize, xq: &[i8]) -> f32 {
3038    // SAFETY: see dot_q4t_row_sdot.
3039    unsafe {
3040        use core::arch::x86_64::*;
3041        let lomask = _mm_set1_epi8(0x0F);
3042        let eight = _mm256_set1_epi8(8);
3043        let mut acc = 0f32;
3044        for gi in 0..gpr {
3045            let t = bytes.as_ptr().add((r * gpr + gi) * Q4_TILE);
3046            let s = f16_to_f32(u16::from_le_bytes([*t, *t.add(1)]));
3047            let b = _mm_loadu_si128(t.add(2) as *const __m128i);
3048            let lo = _mm_and_si128(b, lomask);
3049            let hi = _mm_and_si128(_mm_srli_epi16::<4>(b), lomask);
3050            let w = _mm256_sub_epi8(
3051                _mm256_set_m128i(_mm_unpackhi_epi8(lo, hi), _mm_unpacklo_epi8(lo, hi)),
3052                eight,
3053            );
3054            let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
3055            let d = dpbusd_hsum(_mm256_abs_epi8(w), _mm256_sign_epi8(x, w));
3056            acc += d as f32 * s;
3057        }
3058        acc
3059    }
3060}
3061
3062/// One q4_tiled row against FOUR activation streams: the nibble unpack
3063/// and abs() happen once per group instead of once per (group,
3064/// activation) — the unpack is the dominant per-element cost of the
3065/// tiled format (roadmap P0 portable blocking, q4t leg).
3066#[cfg(target_arch = "x86_64")]
3067#[target_feature(enable = "avx2")]
3068unsafe fn dot_q4t_row_1x4_avx2(bytes: &[u8], r: usize, gpr: usize, xs: [&[i8]; 4]) -> [f32; 4] {
3069    // SAFETY: callers uphold the 18B-tile and xq-length contracts.
3070    unsafe {
3071        use core::arch::x86_64::*;
3072        let lomask = _mm_set1_epi8(0x0F);
3073        let eight = _mm256_set1_epi8(8);
3074        let ones = _mm256_set1_epi16(1);
3075        let mut acc = [0f32; 4];
3076        for gi in 0..gpr {
3077            let t = bytes.as_ptr().add((r * gpr + gi) * Q4_TILE);
3078            let s = f16_to_f32(u16::from_le_bytes([*t, *t.add(1)]));
3079            let bb = _mm_loadu_si128(t.add(2) as *const __m128i);
3080            let lo = _mm_and_si128(bb, lomask);
3081            let hi = _mm_and_si128(_mm_srli_epi16::<4>(bb), lomask);
3082            let w = _mm256_sub_epi8(
3083                _mm256_set_m128i(_mm_unpackhi_epi8(lo, hi), _mm_unpacklo_epi8(lo, hi)),
3084                eight,
3085            );
3086            let aw = _mm256_abs_epi8(w);
3087            for (k, xq) in xs.iter().enumerate() {
3088                let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
3089                let p16 = _mm256_maddubs_epi16(aw, _mm256_sign_epi8(x, w));
3090                let d = _mm256_madd_epi16(p16, ones);
3091                let hi128 = _mm256_extracti128_si256::<1>(d);
3092                let s128 = _mm_add_epi32(_mm256_castsi256_si128(d), hi128);
3093                let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
3094                let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
3095                acc[k] += _mm_cvtsi128_si32(s32) as f32 * s;
3096            }
3097        }
3098        acc
3099    }
3100}
3101
3102/// VNNI twin of `dot_q4t_row_1x4_avx2` (see `dpbusd_hsum`).
3103#[cfg(target_arch = "x86_64")]
3104#[target_feature(enable = "avx2,avx512f,avx512bw,avx512vl,avx512vnni")]
3105unsafe fn dot_q4t_row_1x4_vnni(bytes: &[u8], r: usize, gpr: usize, xs: [&[i8]; 4]) -> [f32; 4] {
3106    // SAFETY: callers uphold the 18B-tile and xq-length contracts.
3107    unsafe {
3108        use core::arch::x86_64::*;
3109        let lomask = _mm_set1_epi8(0x0F);
3110        let eight = _mm256_set1_epi8(8);
3111        let mut acc = [0f32; 4];
3112        for gi in 0..gpr {
3113            let t = bytes.as_ptr().add((r * gpr + gi) * Q4_TILE);
3114            let s = f16_to_f32(u16::from_le_bytes([*t, *t.add(1)]));
3115            let bb = _mm_loadu_si128(t.add(2) as *const __m128i);
3116            let lo = _mm_and_si128(bb, lomask);
3117            let hi = _mm_and_si128(_mm_srli_epi16::<4>(bb), lomask);
3118            let w = _mm256_sub_epi8(
3119                _mm256_set_m128i(_mm_unpackhi_epi8(lo, hi), _mm_unpacklo_epi8(lo, hi)),
3120                eight,
3121            );
3122            let aw = _mm256_abs_epi8(w);
3123            for (k, xq) in xs.iter().enumerate() {
3124                let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
3125                let d = dpbusd_hsum(aw, _mm256_sign_epi8(x, w));
3126                acc[k] += d as f32 * s;
3127            }
3128        }
3129        acc
3130    }
3131}
3132
3133/// ARM twin of `dot_q4t_row_1x4_avx2`: one nibble unpack per group
3134/// serves FOUR activation streams. Per stream the group order and f32
3135/// accumulation match `dot_q4t_row_sdot` exactly — batch == matvec
3136/// bit-for-bit.
3137#[cfg(target_arch = "aarch64")]
3138#[target_feature(enable = "neon,dotprod")]
3139unsafe fn dot_q4t_row_1x4_sdot(bytes: &[u8], r: usize, gpr: usize, xs: [&[i8]; 4]) -> [f32; 4] {
3140    // SAFETY: callers uphold the 18B-tile and xq-length contracts.
3141    unsafe {
3142        use core::arch::aarch64::*;
3143        use core::arch::asm;
3144        let lomask = vdupq_n_u8(0x0F);
3145        let eight = vdupq_n_s8(8);
3146        let mut acc = [0f32; 4];
3147        for gi in 0..gpr {
3148            let t = bytes.as_ptr().add((r * gpr + gi) * Q4_TILE);
3149            let s = f16_to_f32(u16::from_le_bytes([*t, *t.add(1)]));
3150            let b = vld1q_u8(t.add(2));
3151            let lo = vandq_u8(b, lomask);
3152            let hi = vshrq_n_u8::<4>(b);
3153            let e0 = vsubq_s8(vreinterpretq_s8_u8(vzip1q_u8(lo, hi)), eight);
3154            let e1 = vsubq_s8(vreinterpretq_s8_u8(vzip2q_u8(lo, hi)), eight);
3155            for (k, xq) in xs.iter().enumerate() {
3156                let x0 = vld1q_s8(xq.as_ptr().add(gi * GROUP_SIZE));
3157                let x1 = vld1q_s8(xq.as_ptr().add(gi * GROUP_SIZE + 16));
3158                let (mut a0, mut a1) = (vdupq_n_s32(0), vdupq_n_s32(0));
3159                asm!(
3160                    "sdot {a0:v}.4s, {e0:v}.16b, {x0:v}.16b",
3161                    "sdot {a1:v}.4s, {e1:v}.16b, {x1:v}.16b",
3162                    a0 = inout(vreg) a0, a1 = inout(vreg) a1,
3163                    e0 = in(vreg) e0, x0 = in(vreg) x0, e1 = in(vreg) e1, x1 = in(vreg) x1,
3164                    options(pure, nomem, nostack),
3165                );
3166                acc[k] += vaddvq_s32(vaddq_s32(a0, a1)) as f32 * s;
3167            }
3168        }
3169        acc
3170    }
3171}
3172
3173/// Exact-term correction for A8W8 outliers on a tiled row.
3174#[inline]
3175fn q4t_outlier(bytes: &[u8], r: usize, gpr: usize, j: usize) -> (f32, f32) {
3176    let gi = j / GROUP_SIZE;
3177    let k = j % GROUP_SIZE;
3178    let tile = &bytes[(r * gpr + gi) * Q4_TILE..(r * gpr + gi + 1) * Q4_TILE];
3179    let s = f16_to_f32(u16::from_le_bytes([tile[0], tile[1]]));
3180    let byte = tile[2 + k / 2];
3181    let nib = if k & 1 == 0 { byte & 0x0F } else { byte >> 4 };
3182    ((nib as i32 - 8) as f32, s)
3183}
3184
3185/// Exact scalar q4_tiled row (CMF_SDOT=0 contract) — same pairwise
3186/// accumulation shape as `q4_range_f32`.
3187#[inline]
3188fn q4t_row_exact(bytes: &[u8], r: usize, gpr: usize, x: &[f32]) -> f32 {
3189    let mut acc = 0f32;
3190    for gi in 0..gpr {
3191        let tile = &bytes[(r * gpr + gi) * Q4_TILE..(r * gpr + gi + 1) * Q4_TILE];
3192        let s = f16_to_f32(u16::from_le_bytes([tile[0], tile[1]]));
3193        let xg = &x[gi * GROUP_SIZE..(gi + 1) * GROUP_SIZE];
3194        let mut ga = 0f32;
3195        for (k, &b) in tile[2..].iter().enumerate() {
3196            ga += ((b & 0x0F) as f32 - 8.0) * xg[k * 2]
3197                + (((b >> 4) & 0x0F) as f32 - 8.0) * xg[k * 2 + 1];
3198        }
3199        acc += ga * s;
3200    }
3201    acc
3202}
3203
3204/// Fused q4_tiled matvec (dispatch mirrors `q4matvec`).
3205fn q4t_matvec(
3206    bytes: &[u8],
3207    x: &[f32],
3208    rows: usize,
3209    cols: usize,
3210    out: &mut [f32],
3211    pool: Option<&Pool>,
3212) {
3213    debug_assert_eq!(out.len(), rows);
3214    let gpr = cols / GROUP_SIZE;
3215    let out_addr = SendMut(out.as_mut_ptr());
3216    if a8w8_enabled() {
3217        let act = split_act(x);
3218        let run = move |start: usize, end: usize| {
3219            for r in start..end {
3220                let mut acc = dot_q4t_row_i8(bytes, r, gpr, &act.xq) * act.sx;
3221                for &(j, xv) in &act.outliers {
3222                    let (w, s) = q4t_outlier(bytes, r, gpr, j);
3223                    acc += w * s * xv;
3224                }
3225                // SAFETY: disjoint row ranges per worker.
3226                unsafe { *out_addr.at(r) = acc };
3227            }
3228        };
3229        dispatch_rows(pool, rows, &run);
3230        return;
3231    }
3232    let run = move |start: usize, end: usize| {
3233        for r in start..end {
3234            // SAFETY: disjoint row ranges per worker.
3235            unsafe { *out_addr.at(r) = q4t_row_exact(bytes, r, gpr, x) };
3236        }
3237    };
3238    dispatch_rows(pool, rows, &run);
3239}
3240
3241/// Fused two-input q4_tiled matvec (weights read once per pair).
3242#[allow(clippy::too_many_arguments)]
3243fn q4t_matvec2(
3244    bytes: &[u8],
3245    x1: &[f32],
3246    x2: &[f32],
3247    rows: usize,
3248    cols: usize,
3249    o1: &mut [f32],
3250    o2: &mut [f32],
3251    pool: Option<&Pool>,
3252) {
3253    let gpr = cols / GROUP_SIZE;
3254    let p1 = SendMut(o1.as_mut_ptr());
3255    let p2 = SendMut(o2.as_mut_ptr());
3256    if a8w8_enabled() {
3257        let a1 = split_act(x1);
3258        let a2 = split_act(x2);
3259        let run = move |start: usize, end: usize| {
3260            for r in start..end {
3261                let mut v1 = dot_q4t_row_i8(bytes, r, gpr, &a1.xq) * a1.sx;
3262                let mut v2 = dot_q4t_row_i8(bytes, r, gpr, &a2.xq) * a2.sx;
3263                for &(j, xv) in &a1.outliers {
3264                    let (w, s) = q4t_outlier(bytes, r, gpr, j);
3265                    v1 += w * s * xv;
3266                }
3267                for &(j, xv) in &a2.outliers {
3268                    let (w, s) = q4t_outlier(bytes, r, gpr, j);
3269                    v2 += w * s * xv;
3270                }
3271                // SAFETY: disjoint row ranges per worker.
3272                unsafe {
3273                    *p1.at(r) = v1;
3274                    *p2.at(r) = v2;
3275                }
3276            }
3277        };
3278        dispatch_rows(pool, rows, &run);
3279        return;
3280    }
3281    let run = move |start: usize, end: usize| {
3282        for r in start..end {
3283            // SAFETY: disjoint row ranges per worker.
3284            unsafe {
3285                *p1.at(r) = q4t_row_exact(bytes, r, gpr, x1);
3286                *p2.at(r) = q4t_row_exact(bytes, r, gpr, x2);
3287            }
3288        }
3289    };
3290    dispatch_rows(pool, rows, &run);
3291}
3292
3293/// Batched q4_tiled matmat: each row's tiles stream once per microbatch.
3294#[allow(clippy::too_many_arguments)]
3295fn q4t_matmat(
3296    bytes: &[u8],
3297    xs_all: &[f32],
3298    b: usize,
3299    rows: usize,
3300    cols: usize,
3301    out: &mut [f32],
3302    pool: Option<&Pool>,
3303) {
3304    debug_assert_eq!(out.len(), b * rows);
3305    let gpr = cols / GROUP_SIZE;
3306    let out_addr = SendMut(out.as_mut_ptr());
3307    if a8w8_enabled() {
3308        let acts: Vec<SplitAct> = (0..b)
3309            .map(|bi| split_act(&xs_all[bi * cols..(bi + 1) * cols]))
3310            .collect();
3311        let acts = &acts;
3312        #[cfg(target_arch = "x86_64")]
3313        let blocked_ok = avx2_enabled()
3314            && std::env::var("CMF_X86_BLOCKED")
3315                .map(|v| v != "0")
3316                .unwrap_or(true);
3317        #[cfg(target_arch = "aarch64")]
3318        let blocked_ok = sdot_enabled()
3319            && std::env::var("CMF_X86_BLOCKED")
3320                .map(|v| v != "0")
3321                .unwrap_or(true);
3322        #[cfg(not(any(target_arch = "x86_64", target_arch = "aarch64")))]
3323        let blocked_ok = false;
3324        let run = move |start: usize, end: usize| {
3325            for r in start..end {
3326                let mut bi = 0usize;
3327                #[cfg(target_arch = "aarch64")]
3328                if blocked_ok {
3329                    while bi + 4 <= acts.len() {
3330                        let xs = [
3331                            acts[bi].xq.as_slice(),
3332                            acts[bi + 1].xq.as_slice(),
3333                            acts[bi + 2].xq.as_slice(),
3334                            acts[bi + 3].xq.as_slice(),
3335                        ];
3336                        let d = unsafe { dot_q4t_row_1x4_sdot(bytes, r, gpr, xs) };
3337                        for k in 0..4 {
3338                            let act = &acts[bi + k];
3339                            let mut acc = d[k] * act.sx;
3340                            for &(j, xv) in &act.outliers {
3341                                let (w, sc) = q4t_outlier(bytes, r, gpr, j);
3342                                acc += w * sc * xv;
3343                            }
3344                            // SAFETY: disjoint (bi, r) cells per worker.
3345                            unsafe { *out_addr.at((bi + k) * rows + r) = acc };
3346                        }
3347                        bi += 4;
3348                    }
3349                }
3350                #[cfg(target_arch = "x86_64")]
3351                if blocked_ok {
3352                    while bi + 4 <= acts.len() {
3353                        let xs = [
3354                            acts[bi].xq.as_slice(),
3355                            acts[bi + 1].xq.as_slice(),
3356                            acts[bi + 2].xq.as_slice(),
3357                            acts[bi + 3].xq.as_slice(),
3358                        ];
3359                        let d = unsafe {
3360                            if vnni_tiles_enabled() {
3361                                dot_q4t_row_1x4_vnni(bytes, r, gpr, xs)
3362                            } else {
3363                                dot_q4t_row_1x4_avx2(bytes, r, gpr, xs)
3364                            }
3365                        };
3366                        for k in 0..4 {
3367                            let act = &acts[bi + k];
3368                            let mut acc = d[k] * act.sx;
3369                            for &(j, xv) in &act.outliers {
3370                                let (w, sc) = q4t_outlier(bytes, r, gpr, j);
3371                                acc += w * sc * xv;
3372                            }
3373                            // SAFETY: disjoint (bi, r) cells per worker.
3374                            unsafe { *out_addr.at((bi + k) * rows + r) = acc };
3375                        }
3376                        bi += 4;
3377                    }
3378                }
3379                let _ = blocked_ok;
3380                while bi < acts.len() {
3381                    let act = &acts[bi];
3382                    let mut acc = dot_q4t_row_i8(bytes, r, gpr, &act.xq) * act.sx;
3383                    for &(j, xv) in &act.outliers {
3384                        let (w, s) = q4t_outlier(bytes, r, gpr, j);
3385                        acc += w * s * xv;
3386                    }
3387                    // SAFETY: disjoint (bi, r) cells per worker range.
3388                    unsafe { *out_addr.at(bi * rows + r) = acc };
3389                    bi += 1;
3390                }
3391            }
3392        };
3393        dispatch_rows(pool, rows, &run);
3394        return;
3395    }
3396    let run = move |start: usize, end: usize| {
3397        for r in start..end {
3398            for bi in 0..b {
3399                let x = &xs_all[bi * cols..(bi + 1) * cols];
3400                // SAFETY: disjoint (bi, r) cells per worker range.
3401                unsafe { *out_addr.at(bi * rows + r) = q4t_row_exact(bytes, r, gpr, x) };
3402            }
3403        }
3404    };
3405    dispatch_rows(pool, rows, &run);
3406}
3407
3408// ── q1 (dtype 12): binary weights, [f16 scale][4B sign bits] per
3409// 32-group tile. The kernel family mirrors q4_tiled: one sequential
3410// stream of 6-byte tiles, per-tile integer dot × scale, exact outlier
3411// correction (A8W8 contract), exact scalar path under CMF_SDOT=0. ──
3412
3413/// Per-32-group sums of the quantized activation — the ±1 identity's
3414/// shared half: `dot = −2·sdot(mask, x) − gsum[g]`, computed ONCE per
3415/// matvec and reused by every row.
3416fn q1_group_sums(xq: &[i8], gpr: usize) -> Vec<i32> {
3417    (0..gpr)
3418        .map(|gi| {
3419            xq[gi * GROUP_SIZE..(gi + 1) * GROUP_SIZE]
3420                .iter()
3421                .map(|&v| v as i32)
3422                .sum()
3423        })
3424        .collect()
3425}
3426
3427/// One q1 row via the A8W8 int8 path — mask-SDOT on ARM (no ±1
3428/// expansion at all), scalar bit loop elsewhere (AVX2 queued with the
3429/// x86 pass).
3430#[inline]
3431#[allow(unreachable_code)]
3432/// AVX2 q1 row via the same ±1 identity as the ARM sdot kernel: the
3433/// sign bits expand to a {0, −1} byte mask through shuffle+cmpeq, the
3434/// masked activation sums through maddubs(1, x&mask), and
3435/// `dot = −(2·masked_sum + Σx_group)` — bit-identical integer math.
3436#[cfg(target_arch = "x86_64")]
3437#[target_feature(enable = "avx2")]
3438unsafe fn dot_q1_row_avx2(bytes: &[u8], r: usize, gpr: usize, xq: &[i8], gsum: &[i32]) -> f32 {
3439    // SAFETY: callers uphold the 6B-tile and xq/gsum length contracts.
3440    unsafe {
3441        use core::arch::x86_64::*;
3442        // Byte j of the mask must replicate bits-byte j/8.
3443        let expand = _mm256_setr_epi8(
3444            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,
3445            3, 3, 3,
3446        );
3447        let bitsel = _mm256_setr_epi8(
3448            1, 2, 4, 8, 16, 32, 64, -128, 1, 2, 4, 8, 16, 32, 64, -128, 1, 2, 4, 8, 16, 32, 64,
3449            -128, 1, 2, 4, 8, 16, 32, 64, -128,
3450        );
3451        let ones8 = _mm256_set1_epi8(1);
3452        let ones16 = _mm256_set1_epi16(1);
3453        let mut acc = 0f32;
3454        for gi in 0..gpr {
3455            let t = bytes.as_ptr().add((r * gpr + gi) * Q1_TILE);
3456            let s = f16_to_f32(u16::from_le_bytes([*t, *t.add(1)]));
3457            let bits = u32::from_le_bytes([*t.add(2), *t.add(3), *t.add(4), *t.add(5)]);
3458            let bc = _mm256_shuffle_epi8(_mm256_set1_epi32(bits as i32), expand);
3459            let mask = _mm256_cmpeq_epi8(_mm256_and_si256(bc, bitsel), bitsel);
3460            let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
3461            let sel = _mm256_and_si256(x, mask);
3462            // Σ of selected i8 lanes: maddubs(1u8, sel_i8) pairs → madd.
3463            let p16 = _mm256_maddubs_epi16(ones8, sel);
3464            let d32 = _mm256_madd_epi16(p16, ones16);
3465            let hi128 = _mm256_extracti128_si256::<1>(d32);
3466            let s128 = _mm_add_epi32(_mm256_castsi256_si128(d32), hi128);
3467            let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
3468            let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
3469            let msum = _mm_cvtsi128_si32(s32);
3470            // The and-select keeps x UN-negated (unlike ARM's −1-mask
3471            // sdot): d = Σ_set − Σ_unset = 2·Σ_set − Σ_all.
3472            let d = 2 * msum - gsum[gi];
3473            acc += d as f32 * s;
3474        }
3475        acc
3476    }
3477}
3478
3479/// VNNI twin of `dot_q1_row_avx2`: the masked-select sum goes through
3480/// one `vpdpbusd(1u8, sel)` (see `dpbusd_hsum` — bit-identical).
3481#[cfg(target_arch = "x86_64")]
3482#[target_feature(enable = "avx2,avx512f,avx512bw,avx512vl,avx512vnni")]
3483unsafe fn dot_q1_row_vnni(bytes: &[u8], r: usize, gpr: usize, xq: &[i8], gsum: &[i32]) -> f32 {
3484    // SAFETY: callers uphold the 6B-tile and xq/gsum length contracts.
3485    unsafe {
3486        use core::arch::x86_64::*;
3487        let expand = _mm256_setr_epi8(
3488            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,
3489            3, 3, 3,
3490        );
3491        let bitsel = _mm256_setr_epi8(
3492            1, 2, 4, 8, 16, 32, 64, -128, 1, 2, 4, 8, 16, 32, 64, -128, 1, 2, 4, 8, 16, 32, 64,
3493            -128, 1, 2, 4, 8, 16, 32, 64, -128,
3494        );
3495        let ones8 = _mm256_set1_epi8(1);
3496        let mut acc = 0f32;
3497        for gi in 0..gpr {
3498            let t = bytes.as_ptr().add((r * gpr + gi) * Q1_TILE);
3499            let s = f16_to_f32(u16::from_le_bytes([*t, *t.add(1)]));
3500            let bits = u32::from_le_bytes([*t.add(2), *t.add(3), *t.add(4), *t.add(5)]);
3501            let bc = _mm256_shuffle_epi8(_mm256_set1_epi32(bits as i32), expand);
3502            let mask = _mm256_cmpeq_epi8(_mm256_and_si256(bc, bitsel), bitsel);
3503            let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
3504            let msum = dpbusd_hsum(ones8, _mm256_and_si256(x, mask));
3505            let d = 2 * msum - gsum[gi];
3506            acc += d as f32 * s;
3507        }
3508        acc
3509    }
3510}
3511
3512/// VNNI twin of `dot_q1_row_1x4_avx2` (see `dpbusd_hsum`).
3513#[cfg(target_arch = "x86_64")]
3514#[target_feature(enable = "avx2,avx512f,avx512bw,avx512vl,avx512vnni")]
3515unsafe fn dot_q1_row_1x4_vnni(
3516    bytes: &[u8],
3517    r: usize,
3518    gpr: usize,
3519    xs: [&[i8]; 4],
3520    gsums: [&[i32]; 4],
3521) -> [f32; 4] {
3522    // SAFETY: callers uphold the 6B-tile and xq/gsum length contracts.
3523    unsafe {
3524        use core::arch::x86_64::*;
3525        let expand = _mm256_setr_epi8(
3526            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,
3527            3, 3, 3,
3528        );
3529        let bitsel = _mm256_setr_epi8(
3530            1, 2, 4, 8, 16, 32, 64, -128, 1, 2, 4, 8, 16, 32, 64, -128, 1, 2, 4, 8, 16, 32, 64,
3531            -128, 1, 2, 4, 8, 16, 32, 64, -128,
3532        );
3533        let ones8 = _mm256_set1_epi8(1);
3534        let mut acc = [0f32; 4];
3535        for gi in 0..gpr {
3536            let t = bytes.as_ptr().add((r * gpr + gi) * Q1_TILE);
3537            let s = f16_to_f32(u16::from_le_bytes([*t, *t.add(1)]));
3538            let bits = u32::from_le_bytes([*t.add(2), *t.add(3), *t.add(4), *t.add(5)]);
3539            let bc = _mm256_shuffle_epi8(_mm256_set1_epi32(bits as i32), expand);
3540            let mask = _mm256_cmpeq_epi8(_mm256_and_si256(bc, bitsel), bitsel);
3541            for (k, xq) in xs.iter().enumerate() {
3542                let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
3543                let msum = dpbusd_hsum(ones8, _mm256_and_si256(x, mask));
3544                let d = 2 * msum - gsums[k][gi];
3545                acc[k] += d as f32 * s;
3546            }
3547        }
3548        acc
3549    }
3550}
3551
3552/// The blocked 1×4 flavor: the expanded bit mask serves four activation
3553/// streams per group (mask build once, four select+reduce chains).
3554#[cfg(target_arch = "x86_64")]
3555#[target_feature(enable = "avx2")]
3556unsafe fn dot_q1_row_1x4_avx2(
3557    bytes: &[u8],
3558    r: usize,
3559    gpr: usize,
3560    xs: [&[i8]; 4],
3561    gsums: [&[i32]; 4],
3562) -> [f32; 4] {
3563    // SAFETY: callers uphold the 6B-tile and xq/gsum length contracts.
3564    unsafe {
3565        use core::arch::x86_64::*;
3566        let expand = _mm256_setr_epi8(
3567            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,
3568            3, 3, 3,
3569        );
3570        let bitsel = _mm256_setr_epi8(
3571            1, 2, 4, 8, 16, 32, 64, -128, 1, 2, 4, 8, 16, 32, 64, -128, 1, 2, 4, 8, 16, 32, 64,
3572            -128, 1, 2, 4, 8, 16, 32, 64, -128,
3573        );
3574        let ones8 = _mm256_set1_epi8(1);
3575        let ones16 = _mm256_set1_epi16(1);
3576        let mut acc = [0f32; 4];
3577        for gi in 0..gpr {
3578            let t = bytes.as_ptr().add((r * gpr + gi) * Q1_TILE);
3579            let s = f16_to_f32(u16::from_le_bytes([*t, *t.add(1)]));
3580            let bits = u32::from_le_bytes([*t.add(2), *t.add(3), *t.add(4), *t.add(5)]);
3581            let bc = _mm256_shuffle_epi8(_mm256_set1_epi32(bits as i32), expand);
3582            let mask = _mm256_cmpeq_epi8(_mm256_and_si256(bc, bitsel), bitsel);
3583            for (k, xq) in xs.iter().enumerate() {
3584                let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
3585                let sel = _mm256_and_si256(x, mask);
3586                let p16 = _mm256_maddubs_epi16(ones8, sel);
3587                let d32 = _mm256_madd_epi16(p16, ones16);
3588                let hi128 = _mm256_extracti128_si256::<1>(d32);
3589                let s128 = _mm_add_epi32(_mm256_castsi256_si128(d32), hi128);
3590                let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
3591                let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
3592                let msum = _mm_cvtsi128_si32(s32);
3593                let d = 2 * msum - gsums[k][gi];
3594                acc[k] += d as f32 * s;
3595            }
3596        }
3597        acc
3598    }
3599}
3600
3601#[allow(unreachable_code)]
3602fn dot_q1_row_i8(bytes: &[u8], r: usize, gpr: usize, xq: &[i8], gsum: &[i32]) -> f32 {
3603    #[cfg(target_arch = "aarch64")]
3604    unsafe {
3605        return dot_q1_row_sdot(bytes, r, gpr, xq, gsum);
3606    }
3607    #[cfg(target_arch = "x86_64")]
3608    if avx2_enabled() {
3609        unsafe {
3610            if vnni_tiles_enabled() {
3611                return dot_q1_row_vnni(bytes, r, gpr, xq, gsum);
3612            }
3613            return dot_q1_row_avx2(bytes, r, gpr, xq, gsum);
3614        }
3615    }
3616    let _ = gsum;
3617    let mut acc = 0f32;
3618    for gi in 0..gpr {
3619        let tile = &bytes[(r * gpr + gi) * Q1_TILE..(r * gpr + gi + 1) * Q1_TILE];
3620        let s = f16_to_f32(u16::from_le_bytes([tile[0], tile[1]]));
3621        let mut d = 0i32;
3622        for (j, &b) in tile[2..].iter().enumerate() {
3623            for k in 0..8 {
3624                let w = ((b >> k) & 1) as i32 * 2 - 1;
3625                d += w * xq[gi * GROUP_SIZE + j * 8 + k] as i32;
3626            }
3627        }
3628        acc += d as f32 * s;
3629    }
3630    acc
3631}
3632
3633/// SDOT q1 row via the ±1 identity: the vtst mask (0xFF where the bit
3634/// is set, i.e. −1 as i8) feeds `sdot` DIRECTLY — no expansion to ±1
3635/// lanes at all — and `dot = −(2·sdot(mask, x) + Σx_group)`, with the
3636/// per-group activation sums shared across every row of the matvec.
3637/// Four tiles (128 weights) per iteration: integer dots reduce through
3638/// a vpaddq tree into ONE i32x4 that meets its four scales in a single
3639/// fused f32 multiply-add. Integer math throughout — bit-identical to
3640/// the scalar ±1 reference.
3641#[cfg(target_arch = "aarch64")]
3642#[target_feature(enable = "neon,dotprod")]
3643unsafe fn dot_q1_row_sdot(bytes: &[u8], r: usize, gpr: usize, xq: &[i8], gsum: &[i32]) -> f32 {
3644    // SAFETY: callers uphold slice-length contracts (6B tile per group,
3645    // xq.len() == gpr·GROUP_SIZE, gsum.len() == gpr).
3646    unsafe {
3647        use core::arch::aarch64::*;
3648        use core::arch::asm;
3649        const MASKS: [u8; 16] = [1, 2, 4, 8, 16, 32, 64, 128, 1, 2, 4, 8, 16, 32, 64, 128];
3650        let m = vld1q_u8(MASKS.as_ptr());
3651        // One tile's −Σ_set(x) as an UNREDUCED i32x4 (two mask-sdots).
3652        macro_rules! tile_dot {
3653            ($t:expr, $x:expr) => {{
3654                let v0 = vcombine_u8(vdup_n_u8(*$t.add(2)), vdup_n_u8(*$t.add(3)));
3655                let v1 = vcombine_u8(vdup_n_u8(*$t.add(4)), vdup_n_u8(*$t.add(5)));
3656                let w0 = vreinterpretq_s8_u8(vtstq_u8(v0, m));
3657                let w1 = vreinterpretq_s8_u8(vtstq_u8(v1, m));
3658                let x0 = vld1q_s8($x);
3659                let x1 = vld1q_s8($x.add(16));
3660                let (mut a0, mut a1) = (vdupq_n_s32(0), vdupq_n_s32(0));
3661                asm!(
3662                    "sdot {a0:v}.4s, {w0:v}.16b, {x0:v}.16b",
3663                    "sdot {a1:v}.4s, {w1:v}.16b, {x1:v}.16b",
3664                    a0 = inout(vreg) a0, a1 = inout(vreg) a1,
3665                    w0 = in(vreg) w0, x0 = in(vreg) x0, w1 = in(vreg) w1, x1 = in(vreg) x1,
3666                    options(pure, nomem, nostack),
3667                );
3668                vaddq_s32(a0, a1)
3669            }};
3670        }
3671        // TBL unpack over PAIR loads: one vld1q covers two 6B tiles
3672        // ([s s b b b b][s s b b b b] + 4B slack), TBL replicates each
3673        // bit-byte across 8 lanes for vtst, and the four scales gather
3674        // through tbl2 into one fcvtl — the 16 ld1r broadcast loads and
3675        // 4 branchy software f16 conversions per 128 weights (the
3676        // measured load-port wall of this kernel) become 2 vector
3677        // loads + 9 table lookups. Integer math order is unchanged —
3678        // bit-identical results (FCVTL is exact on every f16).
3679        const IW00: [u8; 16] = [2, 2, 2, 2, 2, 2, 2, 2, 3, 3, 3, 3, 3, 3, 3, 3];
3680        const IW01: [u8; 16] = [4, 4, 4, 4, 4, 4, 4, 4, 5, 5, 5, 5, 5, 5, 5, 5];
3681        const IW10: [u8; 16] = [8, 8, 8, 8, 8, 8, 8, 8, 9, 9, 9, 9, 9, 9, 9, 9];
3682        const IW11: [u8; 16] = [
3683            10, 10, 10, 10, 10, 10, 10, 10, 11, 11, 11, 11, 11, 11, 11, 11,
3684        ];
3685        const ISC: [u8; 8] = [0, 1, 6, 7, 16, 17, 22, 23];
3686        let (iw00, iw01) = (vld1q_u8(IW00.as_ptr()), vld1q_u8(IW01.as_ptr()));
3687        let (iw10, iw11) = (vld1q_u8(IW10.as_ptr()), vld1q_u8(IW11.as_ptr()));
3688        let isc = vld1_u8(ISC.as_ptr());
3689        // One tile's −Σ_set(x) from a TBL-unpacked pair load.
3690        macro_rules! tile_dot_tbl {
3691            ($ld:expr, $i0:expr, $i1:expr, $x:expr) => {{
3692                let w0 = vreinterpretq_s8_u8(vtstq_u8(vqtbl1q_u8($ld, $i0), m));
3693                let w1 = vreinterpretq_s8_u8(vtstq_u8(vqtbl1q_u8($ld, $i1), m));
3694                let x0 = vld1q_s8($x);
3695                let x1 = vld1q_s8($x.add(16));
3696                let (mut a0, mut a1) = (vdupq_n_s32(0), vdupq_n_s32(0));
3697                asm!(
3698                    "sdot {a0:v}.4s, {w0:v}.16b, {x0:v}.16b",
3699                    "sdot {a1:v}.4s, {w1:v}.16b, {x1:v}.16b",
3700                    a0 = inout(vreg) a0, a1 = inout(vreg) a1,
3701                    w0 = in(vreg) w0, x0 = in(vreg) x0, w1 = in(vreg) w1, x1 = in(vreg) x1,
3702                    options(pure, nomem, nostack),
3703                );
3704                vaddq_s32(a0, a1)
3705            }};
3706        }
3707        let base = bytes.as_ptr().add(r * gpr * Q1_TILE);
3708        let row_base = r * gpr * Q1_TILE;
3709        let abs_end = bytes.len();
3710        let xp = xq.as_ptr();
3711        let gp = gsum.as_ptr();
3712        let mut accv = vdupq_n_f32(0.0);
3713        let mut gi = 0;
3714        // The second pair load reads 4B past tile gi+3 — stay inside
3715        // the payload slice (only the file's final tiles fall back).
3716        while gi + 4 <= gpr && row_base + (gi + 4) * Q1_TILE + 4 <= abs_end {
3717            let t0 = base.add(gi * Q1_TILE);
3718            let ld_a = vld1q_u8(t0);
3719            let ld_b = vld1q_u8(t0.add(2 * Q1_TILE));
3720            let d0 = tile_dot_tbl!(ld_a, iw00, iw01, xp.add(gi * GROUP_SIZE));
3721            let d1 = tile_dot_tbl!(ld_a, iw10, iw11, xp.add((gi + 1) * GROUP_SIZE));
3722            let d2 = tile_dot_tbl!(ld_b, iw00, iw01, xp.add((gi + 2) * GROUP_SIZE));
3723            let d3 = tile_dot_tbl!(ld_b, iw10, iw11, xp.add((gi + 3) * GROUP_SIZE));
3724            // [−Σ0, −Σ1, −Σ2, −Σ3] → dots = −(2·Σset_neg + gsum)
3725            let neg = vpaddq_s32(vpaddq_s32(d0, d1), vpaddq_s32(d2, d3));
3726            let g = vld1q_s32(gp.add(gi));
3727            let dots = vnegq_s32(vaddq_s32(vshlq_n_s32::<1>(neg), g));
3728            let sc16 = vqtbl2_u8(uint8x16x2_t(ld_a, ld_b), isc);
3729            let scf: float32x4_t;
3730            asm!(
3731                "fcvtl {o:v}.4s, {i:v}.4h",
3732                o = out(vreg) scf, i = in(vreg) sc16,
3733                options(pure, nomem, nostack),
3734            );
3735            accv = vfmaq_f32(accv, vcvtq_f32_s32(dots), scf);
3736            gi += 4;
3737        }
3738        let mut acc = vaddvq_f32(accv);
3739        while gi < gpr {
3740            let t = base.add(gi * Q1_TILE);
3741            let s = f16_to_f32(u16::from_le_bytes([*t, *t.add(1)]));
3742            let d = vaddvq_s32(tile_dot!(t, xp.add(gi * GROUP_SIZE)));
3743            acc += (-(2 * d + *gp.add(gi))) as f32 * s;
3744            gi += 1;
3745        }
3746        acc
3747    }
3748}
3749
3750/// Blocked q1 1×4: one TBL unpack of the tile pair serves FOUR
3751/// activation streams (prefill amortization — the same idea as the
3752/// AVX2 twin; per stream the group order, fma order and tail match the
3753/// single-row kernel exactly, so batch == matvec bit-for-bit).
3754#[cfg(target_arch = "aarch64")]
3755#[target_feature(enable = "neon,dotprod")]
3756unsafe fn dot_q1_row_1x4_sdot(
3757    bytes: &[u8],
3758    r: usize,
3759    gpr: usize,
3760    xs: [&[i8]; 4],
3761    gs: [&[i32]; 4],
3762) -> [f32; 4] {
3763    // SAFETY: same slice-length contracts as `dot_q1_row_sdot`, ×4.
3764    unsafe {
3765        use core::arch::aarch64::*;
3766        use core::arch::asm;
3767        const MASKS: [u8; 16] = [1, 2, 4, 8, 16, 32, 64, 128, 1, 2, 4, 8, 16, 32, 64, 128];
3768        const IW00: [u8; 16] = [2, 2, 2, 2, 2, 2, 2, 2, 3, 3, 3, 3, 3, 3, 3, 3];
3769        const IW01: [u8; 16] = [4, 4, 4, 4, 4, 4, 4, 4, 5, 5, 5, 5, 5, 5, 5, 5];
3770        const IW10: [u8; 16] = [8, 8, 8, 8, 8, 8, 8, 8, 9, 9, 9, 9, 9, 9, 9, 9];
3771        const IW11: [u8; 16] = [
3772            10, 10, 10, 10, 10, 10, 10, 10, 11, 11, 11, 11, 11, 11, 11, 11,
3773        ];
3774        const ISC: [u8; 8] = [0, 1, 6, 7, 16, 17, 22, 23];
3775        let m = vld1q_u8(MASKS.as_ptr());
3776        let (iw00, iw01) = (vld1q_u8(IW00.as_ptr()), vld1q_u8(IW01.as_ptr()));
3777        let (iw10, iw11) = (vld1q_u8(IW10.as_ptr()), vld1q_u8(IW11.as_ptr()));
3778        let isc = vld1_u8(ISC.as_ptr());
3779        macro_rules! sdot2 {
3780            ($w0:expr, $w1:expr, $x:expr) => {{
3781                let x0 = vld1q_s8($x);
3782                let x1 = vld1q_s8($x.add(16));
3783                let (mut a0, mut a1) = (vdupq_n_s32(0), vdupq_n_s32(0));
3784                asm!(
3785                    "sdot {a0:v}.4s, {w0:v}.16b, {x0:v}.16b",
3786                    "sdot {a1:v}.4s, {w1:v}.16b, {x1:v}.16b",
3787                    a0 = inout(vreg) a0, a1 = inout(vreg) a1,
3788                    w0 = in(vreg) $w0, x0 = in(vreg) x0, w1 = in(vreg) $w1, x1 = in(vreg) x1,
3789                    options(pure, nomem, nostack),
3790                );
3791                vaddq_s32(a0, a1)
3792            }};
3793        }
3794        let base = bytes.as_ptr().add(r * gpr * Q1_TILE);
3795        let row_base = r * gpr * Q1_TILE;
3796        let abs_end = bytes.len();
3797        let mut accv = [vdupq_n_f32(0.0); 4];
3798        let mut gi = 0;
3799        while gi + 4 <= gpr && row_base + (gi + 4) * Q1_TILE + 4 <= abs_end {
3800            let t0 = base.add(gi * Q1_TILE);
3801            let ld_a = vld1q_u8(t0);
3802            let ld_b = vld1q_u8(t0.add(2 * Q1_TILE));
3803            // Unpack ONCE — eight ±mask vectors serve all four streams.
3804            let w00 = vreinterpretq_s8_u8(vtstq_u8(vqtbl1q_u8(ld_a, iw00), m));
3805            let w01 = vreinterpretq_s8_u8(vtstq_u8(vqtbl1q_u8(ld_a, iw01), m));
3806            let w10 = vreinterpretq_s8_u8(vtstq_u8(vqtbl1q_u8(ld_a, iw10), m));
3807            let w11 = vreinterpretq_s8_u8(vtstq_u8(vqtbl1q_u8(ld_a, iw11), m));
3808            let w20 = vreinterpretq_s8_u8(vtstq_u8(vqtbl1q_u8(ld_b, iw00), m));
3809            let w21 = vreinterpretq_s8_u8(vtstq_u8(vqtbl1q_u8(ld_b, iw01), m));
3810            let w30 = vreinterpretq_s8_u8(vtstq_u8(vqtbl1q_u8(ld_b, iw10), m));
3811            let w31 = vreinterpretq_s8_u8(vtstq_u8(vqtbl1q_u8(ld_b, iw11), m));
3812            let sc16 = vqtbl2_u8(uint8x16x2_t(ld_a, ld_b), isc);
3813            let scf: float32x4_t;
3814            asm!(
3815                "fcvtl {o:v}.4s, {i:v}.4h",
3816                o = out(vreg) scf, i = in(vreg) sc16,
3817                options(pure, nomem, nostack),
3818            );
3819            for k in 0..4 {
3820                let xp = xs[k].as_ptr();
3821                let d0 = sdot2!(w00, w01, xp.add(gi * GROUP_SIZE));
3822                let d1 = sdot2!(w10, w11, xp.add((gi + 1) * GROUP_SIZE));
3823                let d2 = sdot2!(w20, w21, xp.add((gi + 2) * GROUP_SIZE));
3824                let d3 = sdot2!(w30, w31, xp.add((gi + 3) * GROUP_SIZE));
3825                let neg = vpaddq_s32(vpaddq_s32(d0, d1), vpaddq_s32(d2, d3));
3826                let g = vld1q_s32(gs[k].as_ptr().add(gi));
3827                let dots = vnegq_s32(vaddq_s32(vshlq_n_s32::<1>(neg), g));
3828                accv[k] = vfmaq_f32(accv[k], vcvtq_f32_s32(dots), scf);
3829            }
3830            gi += 4;
3831        }
3832        let mut acc = [
3833            vaddvq_f32(accv[0]),
3834            vaddvq_f32(accv[1]),
3835            vaddvq_f32(accv[2]),
3836            vaddvq_f32(accv[3]),
3837        ];
3838        while gi < gpr {
3839            let t = base.add(gi * Q1_TILE);
3840            let sc = f16_to_f32(u16::from_le_bytes([*t, *t.add(1)]));
3841            let v0 = vcombine_u8(vdup_n_u8(*t.add(2)), vdup_n_u8(*t.add(3)));
3842            let v1 = vcombine_u8(vdup_n_u8(*t.add(4)), vdup_n_u8(*t.add(5)));
3843            let w0 = vreinterpretq_s8_u8(vtstq_u8(v0, m));
3844            let w1 = vreinterpretq_s8_u8(vtstq_u8(v1, m));
3845            for k in 0..4 {
3846                let d = vaddvq_s32(sdot2!(w0, w1, xs[k].as_ptr().add(gi * GROUP_SIZE)));
3847                acc[k] += (-(2 * d + *gs[k].as_ptr().add(gi))) as f32 * sc;
3848            }
3849            gi += 1;
3850        }
3851        acc
3852    }
3853}
3854
3855/// (weight ±1, scale) of one q1 element — the exact outlier term.
3856#[inline]
3857fn q1_outlier(bytes: &[u8], r: usize, gpr: usize, j: usize) -> (f32, f32) {
3858    let gi = j / GROUP_SIZE;
3859    let k = j % GROUP_SIZE;
3860    let tile = &bytes[(r * gpr + gi) * Q1_TILE..(r * gpr + gi + 1) * Q1_TILE];
3861    let s = f16_to_f32(u16::from_le_bytes([tile[0], tile[1]]));
3862    let bit = (tile[2 + k / 8] >> (k % 8)) & 1;
3863    ((bit as i32 * 2 - 1) as f32, s)
3864}
3865
3866/// Exact scalar q1 row (CMF_SDOT=0 contract).
3867#[inline]
3868fn q1_row_exact(bytes: &[u8], r: usize, gpr: usize, x: &[f32]) -> f32 {
3869    let mut acc = 0f32;
3870    for gi in 0..gpr {
3871        let tile = &bytes[(r * gpr + gi) * Q1_TILE..(r * gpr + gi + 1) * Q1_TILE];
3872        let s = f16_to_f32(u16::from_le_bytes([tile[0], tile[1]]));
3873        let xg = &x[gi * GROUP_SIZE..(gi + 1) * GROUP_SIZE];
3874        let mut ga = 0f32;
3875        for (j, &b) in tile[2..].iter().enumerate() {
3876            for k in 0..8 {
3877                ga += (((b >> k) & 1) as f32 * 2.0 - 1.0) * xg[j * 8 + k];
3878            }
3879        }
3880        acc += ga * s;
3881    }
3882    acc
3883}
3884
3885/// One q1 row range via A8W8 (the body of `q1_matvec`'s hot loop,
3886/// extracted so multi-matrix jobs drive the same kernel).
3887#[allow(clippy::too_many_arguments)]
3888fn q1_range_a8w8(
3889    bytes: &[u8],
3890    gpr: usize,
3891    act: &SplitAct,
3892    gsum: &[i32],
3893    out: SendMut,
3894    start: usize,
3895    end: usize,
3896) {
3897    for r in start..end {
3898        let mut acc = dot_q1_row_i8(bytes, r, gpr, &act.xq, gsum) * act.sx;
3899        for &(j, xv) in &act.outliers {
3900            let (w, s) = q1_outlier(bytes, r, gpr, j);
3901            acc += w * s * xv;
3902        }
3903        // SAFETY: disjoint row ranges per worker.
3904        unsafe { *out.at(r) = acc };
3905    }
3906}
3907
3908/// Exact-scalar q1 row range (CMF_SDOT=0 contract).
3909fn q1_range_f32(bytes: &[u8], gpr: usize, x: &[f32], out: SendMut, start: usize, end: usize) {
3910    for r in start..end {
3911        // SAFETY: disjoint row ranges per worker.
3912        unsafe { *out.at(r) = q1_row_exact(bytes, r, gpr, x) };
3913    }
3914}
3915
3916/// q1t per-row overlay locator. After the base (`base_len`) come
3917/// `[u32 row_ptr[rows+1]]` then `[(u16 col, f16 val)]` grouped by row (row
3918/// `r`'s entries are `[row_ptr[r], row_ptr[r+1])`). Returns
3919/// `(row_ptr offset, entries offset, present)`.
3920fn q1t_overlay(bytes: &[u8], base_len: usize, rows: usize) -> (usize, usize, bool) {
3921    let entries = base_len + (rows + 1) * 4;
3922    (base_len, entries, entries <= bytes.len())
3923}
3924
3925/// Read `row_ptr[r]` from the overlay's prefix-sum table.
3926#[inline]
3927fn q1t_rowptr(bytes: &[u8], rp_off: usize, r: usize) -> usize {
3928    let o = rp_off + r * 4;
3929    u32::from_le_bytes([bytes[o], bytes[o + 1], bytes[o + 2], bytes[o + 3]]) as usize
3930}
3931
3932/// Byte → the 5 ternary signs it packs `{−1,0,+1}` as f32, precomputed so
3933/// decoding a q1t code is a table load, not the base-3 divide/modulo per
3934/// weight (division is ~20–40× the cost of a load). Built at compile time.
3935const SIGN5: [[f32; 5]; 256] = {
3936    let mut lut = [[0.0f32; 5]; 256];
3937    let pow3 = [1u16, 3, 9, 27, 81];
3938    let mut byte = 0usize;
3939    while byte < 256 {
3940        let mut i = 0usize;
3941        while i < 5 {
3942            let code = (byte as u16 / pow3[i]) % 3;
3943            lut[byte][i] = if code == 1 {
3944                1.0
3945            } else if code == 2 {
3946                -1.0
3947            } else {
3948                0.0
3949            };
3950            i += 1;
3951        }
3952        byte += 1;
3953    }
3954    lut
3955};
3956
3957/// Same table, as i8 signs — the operand for the int8 SDOT base kernel.
3958const SIGN5_I8: [[i8; 5]; 256] = {
3959    let mut lut = [[0i8; 5]; 256];
3960    let pow3 = [1u16, 3, 9, 27, 81];
3961    let mut byte = 0usize;
3962    while byte < 256 {
3963        let mut i = 0usize;
3964        while i < 5 {
3965            let code = (byte as u16 / pow3[i]) % 3;
3966            lut[byte][i] = if code == 1 {
3967                1
3968            } else if code == 2 {
3969                -1
3970            } else {
3971                0
3972            };
3973            i += 1;
3974        }
3975        byte += 1;
3976    }
3977    lut
3978};
3979
3980/// The same 5 i8 signs packed into a u64 (`[s0 s1 s2 s3 s4 0 0 0]`, LE) so the
3981/// group unpack is 7 unaligned u64 stores at offsets 0,5,10,…,30 instead of
3982/// six 5-byte copies + LUT indexing — each store's trailing zeros are fixed by
3983/// the next store, and the last one runs 6 B past the 32nd weight (the unpack
3984/// buffer is padded to 40). This is the decode/prefill hot inner op.
3985const SIGN5_U64: [u64; 256] = {
3986    let mut lut = [0u64; 256];
3987    let pow3 = [1u16, 3, 9, 27, 81];
3988    let mut byte = 0usize;
3989    while byte < 256 {
3990        let mut v = 0u64;
3991        let mut i = 0usize;
3992        while i < 5 {
3993            let code = (byte as u16 / pow3[i]) % 3;
3994            let s: u8 = if code == 1 {
3995                1
3996            } else if code == 2 {
3997                0xFF
3998            } else {
3999                0
4000            };
4001            v |= (s as u64) << (i * 8);
4002            i += 1;
4003        }
4004        lut[byte] = v;
4005        byte += 1;
4006    }
4007    lut
4008};
4009
4010/// Ternary base weight at `(row r, col j)` = `sign(code)·s_group`. Used to add
4011/// back activation-outlier columns, whose `x` was zeroed for the int8 bulk dot
4012/// (`split_act`). At a weight-outlier position the code is 0, so this is 0 and
4013/// the overlay correction owns that column — no double counting.
4014#[inline]
4015fn q1t_base_weight(bytes: &[u8], r: usize, gpr: usize, j: usize) -> f32 {
4016    const TILE: usize = cortiq_core::quant::Q1T_TILE;
4017    let off = (r * gpr + j / GROUP_SIZE) * TILE;
4018    let s = f16_to_f32(u16::from_le_bytes([bytes[off], bytes[off + 1]]));
4019    let within = j % GROUP_SIZE;
4020    SIGN5[bytes[off + 2 + within / 5] as usize][within % 5] * s
4021}
4022
4023/// One 32-group int8 dot via two SDOTs. Bit-exact vs the scalar i8 sum
4024/// (integer accumulation is order-independent).
4025#[cfg(target_arch = "aarch64")]
4026#[target_feature(enable = "neon,dotprod")]
4027#[inline]
4028unsafe fn sdot32_i8(w: *const i8, x: *const i8) -> i32 {
4029    // SAFETY: caller guarantees 32 readable i8 at each pointer.
4030    unsafe {
4031        use core::arch::aarch64::*;
4032        use core::arch::asm;
4033        let w0 = vld1q_s8(w);
4034        let w1 = vld1q_s8(w.add(16));
4035        let x0 = vld1q_s8(x);
4036        let x1 = vld1q_s8(x.add(16));
4037        let (mut a0, mut a1) = (vdupq_n_s32(0), vdupq_n_s32(0));
4038        asm!(
4039            "sdot {a0:v}.4s, {w0:v}.16b, {x0:v}.16b",
4040            "sdot {a1:v}.4s, {w1:v}.16b, {x1:v}.16b",
4041            a0 = inout(vreg) a0, a1 = inout(vreg) a1,
4042            w0 = in(vreg) w0, x0 = in(vreg) x0, w1 = in(vreg) w1, x1 = in(vreg) x1,
4043            options(pure, nomem, nostack),
4044        );
4045        vaddvq_s32(vaddq_s32(a0, a1))
4046    }
4047}
4048
4049/// One 32-group int8 dot via AVX2: signed·signed as `maddubs(|w|, sign(x,w))`
4050/// then `madd` and a horizontal reduce (the same idiom as `dot_q4t_row_avx2`).
4051#[cfg(target_arch = "x86_64")]
4052#[target_feature(enable = "avx2")]
4053#[inline]
4054unsafe fn i8dot32_avx2(w: *const i8, x: *const i8) -> i32 {
4055    // SAFETY: caller guarantees 32 readable i8 at each pointer.
4056    unsafe {
4057        use core::arch::x86_64::*;
4058        let wv = _mm256_loadu_si256(w as *const __m256i);
4059        let xv = _mm256_loadu_si256(x as *const __m256i);
4060        let p16 = _mm256_maddubs_epi16(_mm256_abs_epi8(wv), _mm256_sign_epi8(xv, wv));
4061        let d = _mm256_madd_epi16(p16, _mm256_set1_epi16(1));
4062        let hi128 = _mm256_extracti128_si256::<1>(d);
4063        let s128 = _mm_add_epi32(_mm256_castsi256_si128(d), hi128);
4064        let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
4065        let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
4066        _mm_cvtsi128_si32(s32)
4067    }
4068}
4069
4070/// Unpack one q1t group's base-3 codes into 32 i8 signs via 7 unaligned u64
4071/// stores (see `SIGN5_U64`). `dst` MUST have ≥ 40 bytes: the 7th store writes
4072/// `dst[30..38]`. Stores go in order so each one's trailing zeros are
4073/// overwritten by the next; the final 6 padding bytes are unused by the dot.
4074#[inline]
4075fn q1t_unpack_group_i8(codes: *const u8, dst: &mut [i8]) {
4076    debug_assert!(dst.len() >= 40);
4077    // SAFETY: codes points at 7 readable bytes; dst has ≥ 40 bytes so every
4078    // 8-byte store at offset bi*5 (bi ≤ 6 → ≤ 30) stays in bounds.
4079    unsafe {
4080        let p = dst.as_mut_ptr();
4081        for bi in 0..7 {
4082            core::ptr::write_unaligned(
4083                p.add(bi * 5) as *mut u64,
4084                SIGN5_U64[*codes.add(bi) as usize],
4085            );
4086        }
4087    }
4088}
4089
4090/// One 32-group int8 dot, arch-dispatched (the matmat inner loop, where the
4091/// row's signs are unpacked once and dotted against every batch input).
4092/// Callers are gated by `a8w8_enabled()`, so the target-feature arms are
4093/// reachable; the scalar arm is a non-SIMD-arch fallback.
4094#[inline]
4095fn q1t_i8dot32(w: *const i8, x: *const i8) -> i32 {
4096    #[cfg(target_arch = "aarch64")]
4097    unsafe {
4098        return sdot32_i8(w, x);
4099    }
4100    #[cfg(target_arch = "x86_64")]
4101    unsafe {
4102        return i8dot32_avx2(w, x);
4103    }
4104    #[allow(unreachable_code)]
4105    unsafe {
4106        let mut s = 0i32;
4107        for k in 0..GROUP_SIZE {
4108            s += *w.add(k) as i32 * *x.add(k) as i32;
4109        }
4110        s
4111    }
4112}
4113
4114#[inline]
4115unsafe fn q1t_unpack_reg_u64s(codes: *const u8) -> (u64, u64, u64, u64) {
4116    let (s0, s1, s2, s3, s4, s5, s6) = unsafe {
4117        (
4118            SIGN5_U64[*codes as usize],
4119            SIGN5_U64[*codes.add(1) as usize],
4120            SIGN5_U64[*codes.add(2) as usize],
4121            SIGN5_U64[*codes.add(3) as usize],
4122            SIGN5_U64[*codes.add(4) as usize],
4123            SIGN5_U64[*codes.add(5) as usize],
4124            SIGN5_U64[*codes.add(6) as usize],
4125        )
4126    };
4127
4128    let u0 = s0 | (s1 << 40);
4129    let u1 = (s1 >> 24) | (s2 << 16) | (s3 << 56);
4130    let u2 = (s3 >> 8) | (s4 << 32);
4131    let u3 = (s4 >> 32) | (s5 << 8) | (s6 << 48);
4132
4133    (u0, u1, u2, u3)
4134}
4135
4136/// One q1t row's int8 base dot: `Σ_group s·dot(signs, xq)` (before the shared
4137/// `sx`). Direct register unpacking (zero stack stores/loads, no STLF stalls).
4138/// ARM SDOT.
4139#[cfg(target_arch = "aarch64")]
4140#[target_feature(enable = "neon,dotprod")]
4141unsafe fn q1t_dot_row_sdot(bytes: &[u8], r: usize, gpr: usize, xq: &[i8]) -> f32 {
4142    use core::arch::aarch64::*;
4143    use core::arch::asm;
4144    unsafe {
4145        const TILE: usize = cortiq_core::quant::Q1T_TILE;
4146        let mut acc = 0f32;
4147        let bytes_ptr = bytes.as_ptr();
4148        let xq_ptr = xq.as_ptr();
4149        let row_off = r * gpr * TILE;
4150
4151        let gpr2 = gpr & !1;
4152        let mut gi = 0;
4153        while gi < gpr2 {
4154            let off0 = row_off + gi * TILE;
4155            let off1 = off0 + TILE;
4156            let s0 = f16_to_f32(u16::from_le_bytes([
4157                *bytes_ptr.add(off0),
4158                *bytes_ptr.add(off0 + 1),
4159            ]));
4160            let s1 = f16_to_f32(u16::from_le_bytes([
4161                *bytes_ptr.add(off1),
4162                *bytes_ptr.add(off1 + 1),
4163            ]));
4164
4165            let (u0_0, u1_0, u2_0, u3_0) = q1t_unpack_reg_u64s(bytes_ptr.add(off0 + 2));
4166            let (u0_1, u1_1, u2_1, u3_1) = q1t_unpack_reg_u64s(bytes_ptr.add(off1 + 2));
4167
4168            let w0_0 = vreinterpretq_s8_u64(vcombine_u64(vcreate_u64(u0_0), vcreate_u64(u1_0)));
4169            let w1_0 = vreinterpretq_s8_u64(vcombine_u64(vcreate_u64(u2_0), vcreate_u64(u3_0)));
4170            let w0_1 = vreinterpretq_s8_u64(vcombine_u64(vcreate_u64(u0_1), vcreate_u64(u1_1)));
4171            let w1_1 = vreinterpretq_s8_u64(vcombine_u64(vcreate_u64(u2_1), vcreate_u64(u3_1)));
4172
4173            let x0_0 = vld1q_s8(xq_ptr.add(gi * GROUP_SIZE));
4174            let x1_0 = vld1q_s8(xq_ptr.add(gi * GROUP_SIZE + 16));
4175            let x0_1 = vld1q_s8(xq_ptr.add((gi + 1) * GROUP_SIZE));
4176            let x1_1 = vld1q_s8(xq_ptr.add((gi + 1) * GROUP_SIZE + 16));
4177
4178            let (mut a0_0, mut a1_0) = (vdupq_n_s32(0), vdupq_n_s32(0));
4179            let (mut a0_1, mut a1_1) = (vdupq_n_s32(0), vdupq_n_s32(0));
4180            asm!(
4181                "sdot {a0_0:v}.4s, {w0_0:v}.16b, {x0_0:v}.16b",
4182                "sdot {a1_0:v}.4s, {w1_0:v}.16b, {x1_0:v}.16b",
4183                "sdot {a0_1:v}.4s, {w0_1:v}.16b, {x0_1:v}.16b",
4184                "sdot {a1_1:v}.4s, {w1_1:v}.16b, {x1_1:v}.16b",
4185                a0_0 = inout(vreg) a0_0, a1_0 = inout(vreg) a1_0,
4186                a0_1 = inout(vreg) a0_1, a1_1 = inout(vreg) a1_1,
4187                w0_0 = in(vreg) w0_0, x0_0 = in(vreg) x0_0, w1_0 = in(vreg) w1_0, x1_0 = in(vreg) x1_0,
4188                w0_1 = in(vreg) w0_1, x0_1 = in(vreg) x0_1, w1_1 = in(vreg) w1_1, x1_1 = in(vreg) x1_1,
4189                options(pure, nomem, nostack),
4190            );
4191            let d0 = vaddvq_s32(vaddq_s32(a0_0, a1_0));
4192            let d1 = vaddvq_s32(vaddq_s32(a0_1, a1_1));
4193            acc += d0 as f32 * s0 + d1 as f32 * s1;
4194            gi += 2;
4195        }
4196
4197        if gi < gpr {
4198            let off = row_off + gi * TILE;
4199            let s = f16_to_f32(u16::from_le_bytes([
4200                *bytes_ptr.add(off),
4201                *bytes_ptr.add(off + 1),
4202            ]));
4203            let (u0, u1, u2, u3) = q1t_unpack_reg_u64s(bytes_ptr.add(off + 2));
4204            let w0 = vreinterpretq_s8_u64(vcombine_u64(vcreate_u64(u0), vcreate_u64(u1)));
4205            let w1 = vreinterpretq_s8_u64(vcombine_u64(vcreate_u64(u2), vcreate_u64(u3)));
4206            let x0 = vld1q_s8(xq_ptr.add(gi * GROUP_SIZE));
4207            let x1 = vld1q_s8(xq_ptr.add(gi * GROUP_SIZE + 16));
4208            let (mut a0, mut a1) = (vdupq_n_s32(0), vdupq_n_s32(0));
4209            asm!(
4210                "sdot {a0:v}.4s, {w0:v}.16b, {x0:v}.16b",
4211                "sdot {a1:v}.4s, {w1:v}.16b, {x1:v}.16b",
4212                a0 = inout(vreg) a0, a1 = inout(vreg) a1,
4213                w0 = in(vreg) w0, x0 = in(vreg) x0, w1 = in(vreg) w1, x1 = in(vreg) x1,
4214                options(pure, nomem, nostack),
4215            );
4216            let d = vaddvq_s32(vaddq_s32(a0, a1));
4217            acc += d as f32 * s;
4218        }
4219        acc
4220    }
4221}
4222
4223/// x86 AVX2 mirror of `q1t_dot_row_sdot` (maddubs int8 dot per group).
4224#[cfg(target_arch = "x86_64")]
4225#[target_feature(enable = "avx2")]
4226unsafe fn q1t_dot_row_avx2(bytes: &[u8], r: usize, gpr: usize, xq: &[i8]) -> f32 {
4227    use core::arch::x86_64::*;
4228    unsafe {
4229        const TILE: usize = cortiq_core::quant::Q1T_TILE;
4230        let mut acc = 0f32;
4231        let bytes_ptr = bytes.as_ptr();
4232        let xq_ptr = xq.as_ptr();
4233        let row_off = r * gpr * TILE;
4234
4235        let ones = _mm256_set1_epi16(1);
4236        for gi in 0..gpr {
4237            let off = row_off + gi * TILE;
4238            let s = f16_to_f32(u16::from_le_bytes([
4239                *bytes_ptr.add(off),
4240                *bytes_ptr.add(off + 1),
4241            ]));
4242            let (u0, u1, u2, u3) = q1t_unpack_reg_u64s(bytes_ptr.add(off + 2));
4243            let wv = _mm256_set_epi64x(u3 as i64, u2 as i64, u1 as i64, u0 as i64);
4244            let xv = _mm256_loadu_si256(xq_ptr.add(gi * GROUP_SIZE) as *const __m256i);
4245            let p16 = _mm256_maddubs_epi16(_mm256_abs_epi8(wv), _mm256_sign_epi8(xv, wv));
4246            let d256 = _mm256_madd_epi16(p16, ones);
4247            let d128 = _mm_add_epi32(
4248                _mm256_castsi256_si128(d256),
4249                _mm256_extracti128_si256(d256, 1),
4250            );
4251            let d64 = _mm_add_epi32(d128, _mm_shuffle_epi32(d128, 0xee));
4252            let d32 = _mm_cvtsi128_si32(_mm_add_epi32(d64, _mm_shuffle_epi32(d64, 0x55)));
4253            acc += d32 as f32 * s;
4254        }
4255        acc
4256    }
4257}
4258
4259/// VNNI twin of `q1t_dot_row_avx2` (see `dpbusd_hsum`).
4260#[cfg(target_arch = "x86_64")]
4261#[target_feature(enable = "avx2,avx512f,avx512bw,avx512vl,avx512vnni")]
4262unsafe fn q1t_dot_row_vnni(bytes: &[u8], r: usize, gpr: usize, xq: &[i8]) -> f32 {
4263    use core::arch::x86_64::*;
4264    // SAFETY: same tile/xq contracts as `q1t_dot_row_avx2`.
4265    unsafe {
4266        const TILE: usize = cortiq_core::quant::Q1T_TILE;
4267        let mut acc = 0f32;
4268        let bytes_ptr = bytes.as_ptr();
4269        let xq_ptr = xq.as_ptr();
4270        let row_off = r * gpr * TILE;
4271        for gi in 0..gpr {
4272            let off = row_off + gi * TILE;
4273            let s = f16_to_f32(u16::from_le_bytes([
4274                *bytes_ptr.add(off),
4275                *bytes_ptr.add(off + 1),
4276            ]));
4277            let (u0, u1, u2, u3) = q1t_unpack_reg_u64s(bytes_ptr.add(off + 2));
4278            let wv = _mm256_set_epi64x(u3 as i64, u2 as i64, u1 as i64, u0 as i64);
4279            let xv = _mm256_loadu_si256(xq_ptr.add(gi * GROUP_SIZE) as *const __m256i);
4280            let d = dpbusd_hsum(_mm256_abs_epi8(wv), _mm256_sign_epi8(xv, wv));
4281            acc += d as f32 * s;
4282        }
4283        acc
4284    }
4285}
4286
4287/// Per-row int8 base dot, dispatched once per row (matvec decode hot path).
4288/// Callers are gated by `a8w8_enabled()`, so the target-feature kernels are
4289/// reachable.
4290#[inline]
4291fn q1t_dot_row_i8(bytes: &[u8], r: usize, gpr: usize, xq: &[i8]) -> f32 {
4292    #[cfg(target_arch = "aarch64")]
4293    unsafe {
4294        return q1t_dot_row_sdot(bytes, r, gpr, xq);
4295    }
4296    #[cfg(target_arch = "x86_64")]
4297    unsafe {
4298        if vnni_tiles_enabled() {
4299            return q1t_dot_row_vnni(bytes, r, gpr, xq);
4300        }
4301        return q1t_dot_row_avx2(bytes, r, gpr, xq);
4302    }
4303    #[allow(unreachable_code)]
4304    {
4305        const TILE: usize = cortiq_core::quant::Q1T_TILE;
4306        let mut acc = 0f32;
4307        let mut sg = [0i8; GROUP_SIZE + 8]; // +8 slack for the u64-store unpack
4308        for gi in 0..gpr {
4309            let off = (r * gpr + gi) * TILE;
4310            let s = f16_to_f32(u16::from_le_bytes([bytes[off], bytes[off + 1]]));
4311            q1t_unpack_group_i8(bytes.as_ptr().wrapping_add(off + 2), &mut sg);
4312            let mut d = 0i32;
4313            for k in 0..GROUP_SIZE {
4314                d += sg[k] as i32 * xq[gi * GROUP_SIZE + k] as i32;
4315            }
4316            acc += d as f32 * s;
4317        }
4318        acc
4319    }
4320}
4321
4322/// Σ over a row's outliers of `value·x[col]` — the correction that adds the
4323/// overlay's exact weights on top of the base dot. INVARIANT: the encoder
4324/// writes ternary code 0 at every outlier position (`quantize_q1t`), so the
4325/// base contributes nothing there and this is a plain `value·x`, not
4326/// `(value − base)·x` — no scattered per-outlier scale read. Row `r`'s entries
4327/// are the contiguous slice `[row_ptr[r], row_ptr[r+1])`, so no binary search.
4328fn q1t_row_outlier_correction(
4329    bytes: &[u8],
4330    r: usize,
4331    rp_off: usize,
4332    entries_off: usize,
4333    has_ov: bool,
4334    x: &[f32],
4335) -> f32 {
4336    if !has_ov {
4337        return 0.0;
4338    }
4339    let (c0, c1) = (
4340        q1t_rowptr(bytes, rp_off, r),
4341        q1t_rowptr(bytes, rp_off, r + 1),
4342    );
4343    let mut corr = 0f32;
4344    for p in c0..c1 {
4345        let e = entries_off + p * 4;
4346        let col = u16::from_le_bytes([bytes[e], bytes[e + 1]]) as usize;
4347        let val = f16_to_f32(u16::from_le_bytes([bytes[e + 2], bytes[e + 3]]));
4348        corr += val * x[col];
4349    }
4350    corr
4351}
4352
4353/// Dequantize one q1t row into `buf[..cols]` via the sign LUT (no division),
4354/// then apply the row's outliers (its `[row_ptr[r], row_ptr[r+1])` slice).
4355/// Used by the batched (prefill) path where the decode amortizes over the batch.
4356fn q1t_dequant_row(
4357    bytes: &[u8],
4358    r: usize,
4359    gpr: usize,
4360    rp_off: usize,
4361    entries_off: usize,
4362    has_ov: bool,
4363    buf: &mut [f32],
4364) {
4365    const TILE: usize = cortiq_core::quant::Q1T_TILE;
4366    for g in 0..gpr {
4367        let off = (r * gpr + g) * TILE;
4368        let s = f16_to_f32(u16::from_le_bytes([bytes[off], bytes[off + 1]]));
4369        let codes = &bytes[off + 2..off + TILE];
4370        let bc = g * GROUP_SIZE;
4371        // 6 full bytes (30 codes) + a 7th byte holding the last 2.
4372        for bi in 0..6 {
4373            let lut = &SIGN5[codes[bi] as usize];
4374            let d = &mut buf[bc + bi * 5..bc + bi * 5 + 5];
4375            for i in 0..5 {
4376                d[i] = lut[i] * s;
4377            }
4378        }
4379        let lut = &SIGN5[codes[6] as usize];
4380        buf[bc + 30] = lut[0] * s;
4381        buf[bc + 31] = lut[1] * s;
4382    }
4383    if !has_ov {
4384        return;
4385    }
4386    let (c0, c1) = (
4387        q1t_rowptr(bytes, rp_off, r),
4388        q1t_rowptr(bytes, rp_off, r + 1),
4389    );
4390    for p in c0..c1 {
4391        let e = entries_off + p * 4;
4392        let col = u16::from_le_bytes([bytes[e], bytes[e + 1]]) as usize;
4393        buf[col] = f16_to_f32(u16::from_le_bytes([bytes[e + 2], bytes[e + 3]]));
4394    }
4395}
4396
4397/// Add the sparse outlier overlay onto a base dot already in `out` (the GPU
4398/// computes the ternary base; the overlay stays on the CPU — its entries are
4399/// few and its per-row gather doesn't vectorize on the GPU). Row-parallel.
4400fn q1t_add_overlay(
4401    bytes: &[u8],
4402    x: &[f32],
4403    rows: usize,
4404    cols: usize,
4405    out: &mut [f32],
4406    pool: Option<&Pool>,
4407) {
4408    const TILE: usize = cortiq_core::quant::Q1T_TILE;
4409    let gpr = cols / GROUP_SIZE;
4410    let (rp_off, ent_off, has_ov) = q1t_overlay(bytes, rows * gpr * TILE, rows);
4411    if !has_ov {
4412        return;
4413    }
4414    let out_addr = SendMut(out.as_mut_ptr());
4415    let run = move |start: usize, end: usize| {
4416        for r in start..end {
4417            let corr = q1t_row_outlier_correction(bytes, r, rp_off, ent_off, has_ov, x);
4418            // SAFETY: disjoint rows; add onto the base the GPU already wrote.
4419            unsafe { *out_addr.at(r) += corr };
4420        }
4421    };
4422    dispatch_rows(pool, rows, &run);
4423}
4424
4425/// Q1T row range via the A8W8 int8 path — shared activation split,
4426/// per-row: base SDOT dot + outlier correction + overlay.
4427#[allow(clippy::too_many_arguments)]
4428fn q1t_range_a8w8(
4429    bytes: &[u8],
4430    gpr: usize,
4431    rp_off: usize,
4432    ent_off: usize,
4433    has_ov: bool,
4434    act: &SplitAct,
4435    x: &[f32],
4436    out: SendMut,
4437    start: usize,
4438    end: usize,
4439) {
4440    for r in start..end {
4441        let mut acc = q1t_dot_row_i8(bytes, r, gpr, &act.xq) * act.sx;
4442        for &(j, xv) in &act.outliers {
4443            acc += q1t_base_weight(bytes, r, gpr, j) * xv;
4444        }
4445        acc += q1t_row_outlier_correction(bytes, r, rp_off, ent_off, has_ov, x);
4446        // SAFETY: disjoint row ranges per worker.
4447        unsafe { *out.at(r) = acc };
4448    }
4449}
4450
4451/// Q1T row range via the f32 path (no SDOT) — for matvec_many batched
4452/// dispatch when a8w8 is unavailable.
4453#[allow(clippy::too_many_arguments)]
4454fn q1t_range_f32_batch(
4455    bytes: &[u8],
4456    gpr: usize,
4457    rp_off: usize,
4458    ent_off: usize,
4459    has_ov: bool,
4460    x: &[f32],
4461    out: SendMut,
4462    start: usize,
4463    end: usize,
4464) {
4465    const TILE: usize = cortiq_core::quant::Q1T_TILE;
4466    let mut sg = [0f32; GROUP_SIZE];
4467    for r in start..end {
4468        let mut acc = 0f32;
4469        for g in 0..gpr {
4470            let off = (r * gpr + g) * TILE;
4471            let s = f16_to_f32(u16::from_le_bytes([bytes[off], bytes[off + 1]]));
4472            let codes = &bytes[off + 2..off + TILE];
4473            let xg = &x[g * GROUP_SIZE..g * GROUP_SIZE + GROUP_SIZE];
4474            for bi in 0..6 {
4475                sg[bi * 5..bi * 5 + 5].copy_from_slice(&SIGN5[codes[bi] as usize]);
4476            }
4477            let lut = &SIGN5[codes[6] as usize];
4478            sg[30] = lut[0];
4479            sg[31] = lut[1];
4480            let mut gsum = 0f32;
4481            for k in 0..GROUP_SIZE {
4482                gsum += sg[k] * xg[k];
4483            }
4484            acc += s * gsum;
4485        }
4486        acc += q1t_row_outlier_correction(bytes, r, rp_off, ent_off, has_ov, x);
4487        // SAFETY: disjoint row ranges per worker.
4488        unsafe { *out.at(r) = acc };
4489    }
4490}
4491
4492/// Ternary (q1t) matvec — decode+dot straight from mmap, one group at a time:
4493/// no per-ROW buffer, no division (the sign LUT), and a tiny per-group sign
4494/// buffer so the 32-wide dot vectorizes. This is the decode hot path.
4495fn q1t_matvec(
4496    bytes: &[u8],
4497    x: &[f32],
4498    rows: usize,
4499    cols: usize,
4500    out: &mut [f32],
4501    pool: Option<&Pool>,
4502) {
4503    debug_assert_eq!(out.len(), rows);
4504    const TILE: usize = cortiq_core::quant::Q1T_TILE;
4505    let gpr = cols / GROUP_SIZE;
4506    let (rp_off, ent_off, has_ov) = q1t_overlay(bytes, rows * gpr * TILE, rows);
4507    let out_addr = SendMut(out.as_mut_ptr());
4508    // int8 SDOT base dot (ARM dotprod): ~4× the f32 arithmetic. x → i8 once
4509    // (`split_act`), activation outliers added back exactly in f32, weight
4510    // overlay on top. ARM SDOT / x86 AVX2; CMF_SDOT=0 keeps the exact f32 path.
4511    if a8w8_enabled() {
4512        let act = split_act(x);
4513        let act = &act;
4514        let run = move |start: usize, end: usize| {
4515            for r in start..end {
4516                let mut acc = q1t_dot_row_i8(bytes, r, gpr, &act.xq) * act.sx;
4517                for &(j, xv) in &act.outliers {
4518                    acc += q1t_base_weight(bytes, r, gpr, j) * xv;
4519                }
4520                acc += q1t_row_outlier_correction(bytes, r, rp_off, ent_off, has_ov, x);
4521                // SAFETY: disjoint row ranges per worker.
4522                unsafe { *out_addr.at(r) = acc };
4523            }
4524        };
4525        dispatch_rows(pool, rows, &run);
4526        return;
4527    }
4528    let run = move |start: usize, end: usize| {
4529        // Per-group signs, unpacked contiguously so the dot below is a clean
4530        // 32-wide reduction the autovectorizer turns into f32x4 FMAs — the
4531        // 5-values-per-byte base-3 layout won't SIMD in place.
4532        let mut sg = [0f32; GROUP_SIZE];
4533        for r in start..end {
4534            let mut acc = 0f32;
4535            for g in 0..gpr {
4536                let off = (r * gpr + g) * TILE;
4537                let s = f16_to_f32(u16::from_le_bytes([bytes[off], bytes[off + 1]]));
4538                let codes = &bytes[off + 2..off + TILE];
4539                let xg = &x[g * GROUP_SIZE..g * GROUP_SIZE + GROUP_SIZE];
4540                for bi in 0..6 {
4541                    sg[bi * 5..bi * 5 + 5].copy_from_slice(&SIGN5[codes[bi] as usize]);
4542                }
4543                let lut = &SIGN5[codes[6] as usize];
4544                sg[30] = lut[0];
4545                sg[31] = lut[1];
4546                let mut gsum = 0f32;
4547                for k in 0..GROUP_SIZE {
4548                    gsum += sg[k] * xg[k];
4549                }
4550                acc += s * gsum;
4551            }
4552            acc += q1t_row_outlier_correction(bytes, r, rp_off, ent_off, has_ov, x);
4553            unsafe { *out_addr.at(r) = acc };
4554        }
4555    };
4556    dispatch_rows(pool, rows, &run);
4557}
4558
4559/// Fused-pair twin of `q1t_dot_row_sdot`: ONE register unpack of the
4560/// ternary codes serves BOTH activation streams (the unpack chain is
4561/// the dominant per-row cost — MTP verify pairs paid it twice). Per
4562/// stream the group order and f32 accumulation match the single-row
4563/// kernel exactly, so pair == 2×matvec bit-for-bit.
4564#[cfg(target_arch = "aarch64")]
4565#[target_feature(enable = "neon,dotprod")]
4566unsafe fn q1t_dot_row_sdot2(bytes: &[u8], r: usize, gpr: usize, xa: &[i8], xb: &[i8]) -> [f32; 2] {
4567    use core::arch::aarch64::*;
4568    use core::arch::asm;
4569    // SAFETY: same slice-length contracts as `q1t_dot_row_sdot`, ×2.
4570    unsafe {
4571        const TILE: usize = cortiq_core::quant::Q1T_TILE;
4572        let bytes_ptr = bytes.as_ptr();
4573        let row_off = r * gpr * TILE;
4574        let xp = [xa.as_ptr(), xb.as_ptr()];
4575        let mut acc = [0f32; 2];
4576        macro_rules! sdot2 {
4577            ($w0:expr, $w1:expr, $x:expr) => {{
4578                let x0 = vld1q_s8($x);
4579                let x1 = vld1q_s8($x.add(16));
4580                let (mut a0, mut a1) = (vdupq_n_s32(0), vdupq_n_s32(0));
4581                asm!(
4582                    "sdot {a0:v}.4s, {w0:v}.16b, {x0:v}.16b",
4583                    "sdot {a1:v}.4s, {w1:v}.16b, {x1:v}.16b",
4584                    a0 = inout(vreg) a0, a1 = inout(vreg) a1,
4585                    w0 = in(vreg) $w0, x0 = in(vreg) x0, w1 = in(vreg) $w1, x1 = in(vreg) x1,
4586                    options(pure, nomem, nostack),
4587                );
4588                vaddvq_s32(vaddq_s32(a0, a1))
4589            }};
4590        }
4591        let gpr2 = gpr & !1;
4592        let mut gi = 0;
4593        while gi < gpr2 {
4594            let off0 = row_off + gi * TILE;
4595            let off1 = off0 + TILE;
4596            let s0 = f16_to_f32(u16::from_le_bytes([
4597                *bytes_ptr.add(off0),
4598                *bytes_ptr.add(off0 + 1),
4599            ]));
4600            let s1 = f16_to_f32(u16::from_le_bytes([
4601                *bytes_ptr.add(off1),
4602                *bytes_ptr.add(off1 + 1),
4603            ]));
4604            let (u0_0, u1_0, u2_0, u3_0) = q1t_unpack_reg_u64s(bytes_ptr.add(off0 + 2));
4605            let (u0_1, u1_1, u2_1, u3_1) = q1t_unpack_reg_u64s(bytes_ptr.add(off1 + 2));
4606            let w0_0 = vreinterpretq_s8_u64(vcombine_u64(vcreate_u64(u0_0), vcreate_u64(u1_0)));
4607            let w1_0 = vreinterpretq_s8_u64(vcombine_u64(vcreate_u64(u2_0), vcreate_u64(u3_0)));
4608            let w0_1 = vreinterpretq_s8_u64(vcombine_u64(vcreate_u64(u0_1), vcreate_u64(u1_1)));
4609            let w1_1 = vreinterpretq_s8_u64(vcombine_u64(vcreate_u64(u2_1), vcreate_u64(u3_1)));
4610            for k in 0..2 {
4611                let d0 = sdot2!(w0_0, w1_0, xp[k].add(gi * GROUP_SIZE));
4612                let d1 = sdot2!(w0_1, w1_1, xp[k].add((gi + 1) * GROUP_SIZE));
4613                acc[k] += d0 as f32 * s0 + d1 as f32 * s1;
4614            }
4615            gi += 2;
4616        }
4617        if gi < gpr {
4618            let off = row_off + gi * TILE;
4619            let s = f16_to_f32(u16::from_le_bytes([
4620                *bytes_ptr.add(off),
4621                *bytes_ptr.add(off + 1),
4622            ]));
4623            let (u0, u1, u2, u3) = q1t_unpack_reg_u64s(bytes_ptr.add(off + 2));
4624            let w0 = vreinterpretq_s8_u64(vcombine_u64(vcreate_u64(u0), vcreate_u64(u1)));
4625            let w1 = vreinterpretq_s8_u64(vcombine_u64(vcreate_u64(u2), vcreate_u64(u3)));
4626            for k in 0..2 {
4627                let d = sdot2!(w0, w1, xp[k].add(gi * GROUP_SIZE));
4628                acc[k] += d as f32 * s;
4629            }
4630        }
4631        acc
4632    }
4633}
4634
4635/// Fused Q1T pair matvec: ONE pass over the rows serves both
4636/// activation streams — on ARM the ternary register unpack happens
4637/// once per tile pair (`q1t_dot_row_sdot2`); elsewhere the second dot
4638/// rides the row's L1-warm tile bytes. Per stream the math matches
4639/// `q1t_matvec` exactly.
4640fn q1t_matvec2(
4641    bytes: &[u8],
4642    x1: &[f32],
4643    x2: &[f32],
4644    rows: usize,
4645    cols: usize,
4646    o1: &mut [f32],
4647    o2: &mut [f32],
4648    pool: Option<&Pool>,
4649) {
4650    debug_assert_eq!(o1.len(), rows);
4651    debug_assert_eq!(o2.len(), rows);
4652    const TILE: usize = cortiq_core::quant::Q1T_TILE;
4653    let gpr = cols / GROUP_SIZE;
4654    let (rp_off, ent_off, has_ov) = q1t_overlay(bytes, rows * gpr * TILE, rows);
4655    let out1 = SendMut(o1.as_mut_ptr());
4656    let out2 = SendMut(o2.as_mut_ptr());
4657    if a8w8_enabled() {
4658        let a1 = split_act(x1);
4659        let a2 = split_act(x2);
4660        let (a1, a2) = (&a1, &a2);
4661        let run = move |start: usize, end: usize| {
4662            for r in start..end {
4663                #[cfg(target_arch = "aarch64")]
4664                // a8w8 on aarch64 ⇔ sdot_enabled(), so the kernel's
4665                // target features are present.
4666                let ds = unsafe { q1t_dot_row_sdot2(bytes, r, gpr, &a1.xq, &a2.xq) };
4667                #[cfg(not(target_arch = "aarch64"))]
4668                let ds = [
4669                    q1t_dot_row_i8(bytes, r, gpr, &a1.xq),
4670                    q1t_dot_row_i8(bytes, r, gpr, &a2.xq),
4671                ];
4672                let mut acc1 = ds[0] * a1.sx;
4673                for &(j, xv) in &a1.outliers {
4674                    acc1 += q1t_base_weight(bytes, r, gpr, j) * xv;
4675                }
4676                acc1 += q1t_row_outlier_correction(bytes, r, rp_off, ent_off, has_ov, x1);
4677                let mut acc2 = ds[1] * a2.sx;
4678                for &(j, xv) in &a2.outliers {
4679                    acc2 += q1t_base_weight(bytes, r, gpr, j) * xv;
4680                }
4681                acc2 += q1t_row_outlier_correction(bytes, r, rp_off, ent_off, has_ov, x2);
4682                // SAFETY: disjoint row ranges per worker.
4683                unsafe {
4684                    *out1.at(r) = acc1;
4685                    *out2.at(r) = acc2;
4686                }
4687            }
4688        };
4689        dispatch_rows(pool, rows, &run);
4690        return;
4691    }
4692    let run = move |start: usize, end: usize| {
4693        // Exact path (CMF_SDOT=0): unpack the sign LUT once per group,
4694        // dot both streams — same op order per stream as `q1t_matvec`.
4695        let mut sg = [0f32; GROUP_SIZE];
4696        for r in start..end {
4697            let mut acc1 = 0f32;
4698            let mut acc2 = 0f32;
4699            for g in 0..gpr {
4700                let off = (r * gpr + g) * TILE;
4701                let s = f16_to_f32(u16::from_le_bytes([bytes[off], bytes[off + 1]]));
4702                let codes = &bytes[off + 2..off + TILE];
4703                for bi in 0..6 {
4704                    sg[bi * 5..bi * 5 + 5].copy_from_slice(&SIGN5[codes[bi] as usize]);
4705                }
4706                let lut = &SIGN5[codes[6] as usize];
4707                sg[30] = lut[0];
4708                sg[31] = lut[1];
4709                let xg1 = &x1[g * GROUP_SIZE..g * GROUP_SIZE + GROUP_SIZE];
4710                let xg2 = &x2[g * GROUP_SIZE..g * GROUP_SIZE + GROUP_SIZE];
4711                let mut gsum1 = 0f32;
4712                for k in 0..GROUP_SIZE {
4713                    gsum1 += sg[k] * xg1[k];
4714                }
4715                acc1 += s * gsum1;
4716                let mut gsum2 = 0f32;
4717                for k in 0..GROUP_SIZE {
4718                    gsum2 += sg[k] * xg2[k];
4719                }
4720                acc2 += s * gsum2;
4721            }
4722            acc1 += q1t_row_outlier_correction(bytes, r, rp_off, ent_off, has_ov, x1);
4723            acc2 += q1t_row_outlier_correction(bytes, r, rp_off, ent_off, has_ov, x2);
4724            // SAFETY: disjoint row ranges per worker.
4725            unsafe {
4726                *out1.at(r) = acc1;
4727                *out2.at(r) = acc2;
4728            }
4729        }
4730    };
4731    dispatch_rows(pool, rows, &run);
4732}
4733
4734/// Ternary (q1t) matmat (prefill) — dequant each row once, dot the whole
4735/// batch against it (amortizes the per-row decode).
4736fn q1t_matmat(
4737    bytes: &[u8],
4738    xs: &[f32],
4739    b: usize,
4740    rows: usize,
4741    cols: usize,
4742    out: &mut [f32],
4743    pool: Option<&Pool>,
4744) {
4745    debug_assert_eq!(out.len(), b * rows);
4746    const TILE: usize = cortiq_core::quant::Q1T_TILE;
4747    let gpr = cols / GROUP_SIZE;
4748    let (rp_off, ent_off, has_ov) = q1t_overlay(bytes, rows * gpr * TILE, rows);
4749    let out_addr = SendMut(out.as_mut_ptr());
4750    // int8 prefill (ARM SDOT / x86 AVX2): quantize the B inputs once, unpack
4751    // each weight row's signs to i8 ONCE, then int8-dot against every input —
4752    // the row sign-decode amortizes over the whole batch. CMF_SDOT=0 → f32.
4753    if a8w8_enabled() {
4754        let acts: Vec<SplitAct> = (0..b)
4755            .map(|bi| split_act(&xs[bi * cols..(bi + 1) * cols]))
4756            .collect();
4757        let acts = &acts;
4758        let run = move |start: usize, end: usize| {
4759            let mut sg = vec![0i8; cols + 8]; // row signs, i8 (+8 unpack slack)
4760            let mut sc = vec![0f32; gpr]; // per-group scales
4761            let mut accs = vec![0f32; b]; // per-batch accumulators, reused per row
4762            for r in start..end {
4763                for g in 0..gpr {
4764                    let off = (r * gpr + g) * TILE;
4765                    sc[g] = f16_to_f32(u16::from_le_bytes([bytes[off], bytes[off + 1]]));
4766                    q1t_unpack_group_i8(
4767                        bytes.as_ptr().wrapping_add(off + 2),
4768                        &mut sg[g * GROUP_SIZE..],
4769                    );
4770                }
4771                for bi in 0..b {
4772                    let act = &acts[bi];
4773                    let mut isum = 0f32;
4774                    for g in 0..gpr {
4775                        let d = q1t_i8dot32(
4776                            sg.as_ptr().wrapping_add(g * GROUP_SIZE),
4777                            act.xq.as_ptr().wrapping_add(g * GROUP_SIZE),
4778                        );
4779                        isum += d as f32 * sc[g];
4780                    }
4781                    let mut acc = isum * act.sx;
4782                    for &(j, xv) in &act.outliers {
4783                        acc += q1t_base_weight(bytes, r, gpr, j) * xv;
4784                    }
4785                    accs[bi] = acc;
4786                }
4787                // Overlay ONCE per row for the whole batch: read each (col, val)
4788                // from mmap a single time (was b× — the re-read dominated prefill)
4789                // and fan it out over the batch via the cached inputs.
4790                if has_ov {
4791                    let (c0, c1) = (
4792                        q1t_rowptr(bytes, rp_off, r),
4793                        q1t_rowptr(bytes, rp_off, r + 1),
4794                    );
4795                    for p in c0..c1 {
4796                        let e = ent_off + p * 4;
4797                        let col = u16::from_le_bytes([bytes[e], bytes[e + 1]]) as usize;
4798                        let val = f16_to_f32(u16::from_le_bytes([bytes[e + 2], bytes[e + 3]]));
4799                        for bi in 0..b {
4800                            accs[bi] += val * xs[bi * cols + col];
4801                        }
4802                    }
4803                }
4804                for bi in 0..b {
4805                    unsafe { *out_addr.at(bi * rows + r) = accs[bi] };
4806                }
4807            }
4808        };
4809        dispatch_rows(pool, rows, &run);
4810        return;
4811    }
4812    let run = move |start: usize, end: usize| {
4813        let mut buf = vec![0f32; cols];
4814        for r in start..end {
4815            q1t_dequant_row(bytes, r, gpr, rp_off, ent_off, has_ov, &mut buf);
4816            for bi in 0..b {
4817                let xr = &xs[bi * cols..(bi + 1) * cols];
4818                let mut acc = 0f32;
4819                for j in 0..cols {
4820                    acc += buf[j] * xr[j];
4821                }
4822                unsafe { *out_addr.at(bi * rows + r) = acc };
4823            }
4824        }
4825    };
4826    dispatch_rows(pool, rows, &run);
4827}
4828
4829fn q1_matvec(
4830    bytes: &[u8],
4831    x: &[f32],
4832    rows: usize,
4833    cols: usize,
4834    out: &mut [f32],
4835    pool: Option<&Pool>,
4836) {
4837    debug_assert_eq!(out.len(), rows);
4838    let gpr = cols / GROUP_SIZE;
4839    let out_addr = SendMut(out.as_mut_ptr());
4840    if a8w8_enabled() {
4841        let act = split_act(x);
4842        let gsum = q1_group_sums(&act.xq, gpr);
4843        let (act, gsum) = (&act, &gsum);
4844        let run = move |start: usize, end: usize| {
4845            q1_range_a8w8(bytes, gpr, act, gsum, out_addr, start, end)
4846        };
4847        dispatch_rows(pool, rows, &run);
4848        return;
4849    }
4850    let run = move |start: usize, end: usize| q1_range_f32(bytes, gpr, x, out_addr, start, end);
4851    dispatch_rows(pool, rows, &run);
4852}
4853
4854/// Fused two-input q1 matvec (weights read once per pair).
4855#[allow(clippy::too_many_arguments)]
4856fn q1_matvec2(
4857    bytes: &[u8],
4858    x1: &[f32],
4859    x2: &[f32],
4860    rows: usize,
4861    cols: usize,
4862    o1: &mut [f32],
4863    o2: &mut [f32],
4864    pool: Option<&Pool>,
4865) {
4866    let gpr = cols / GROUP_SIZE;
4867    let p1 = SendMut(o1.as_mut_ptr());
4868    let p2 = SendMut(o2.as_mut_ptr());
4869    if a8w8_enabled() {
4870        let a1 = split_act(x1);
4871        let a2 = split_act(x2);
4872        let g1 = q1_group_sums(&a1.xq, gpr);
4873        let g2 = q1_group_sums(&a2.xq, gpr);
4874        let (a1, a2, g1, g2) = (&a1, &a2, &g1, &g2);
4875        let run = move |start: usize, end: usize| {
4876            for r in start..end {
4877                let mut v1 = dot_q1_row_i8(bytes, r, gpr, &a1.xq, g1) * a1.sx;
4878                let mut v2 = dot_q1_row_i8(bytes, r, gpr, &a2.xq, g2) * a2.sx;
4879                for &(j, xv) in &a1.outliers {
4880                    let (w, s) = q1_outlier(bytes, r, gpr, j);
4881                    v1 += w * s * xv;
4882                }
4883                for &(j, xv) in &a2.outliers {
4884                    let (w, s) = q1_outlier(bytes, r, gpr, j);
4885                    v2 += w * s * xv;
4886                }
4887                // SAFETY: disjoint row ranges per worker.
4888                unsafe {
4889                    *p1.at(r) = v1;
4890                    *p2.at(r) = v2;
4891                }
4892            }
4893        };
4894        dispatch_rows(pool, rows, &run);
4895        return;
4896    }
4897    let run = move |start: usize, end: usize| {
4898        for r in start..end {
4899            // SAFETY: disjoint row ranges per worker.
4900            unsafe {
4901                *p1.at(r) = q1_row_exact(bytes, r, gpr, x1);
4902                *p2.at(r) = q1_row_exact(bytes, r, gpr, x2);
4903            }
4904        }
4905    };
4906    dispatch_rows(pool, rows, &run);
4907}
4908
4909/// Batched q1 matmat: each row's tiles stream once per microbatch.
4910#[allow(clippy::too_many_arguments)]
4911fn q1_matmat(
4912    bytes: &[u8],
4913    xs_all: &[f32],
4914    b: usize,
4915    rows: usize,
4916    cols: usize,
4917    out: &mut [f32],
4918    pool: Option<&Pool>,
4919) {
4920    debug_assert_eq!(out.len(), b * rows);
4921    let gpr = cols / GROUP_SIZE;
4922    let out_addr = SendMut(out.as_mut_ptr());
4923    if a8w8_enabled() {
4924        let acts: Vec<(SplitAct, Vec<i32>)> = (0..b)
4925            .map(|bi| {
4926                let act = split_act(&xs_all[bi * cols..(bi + 1) * cols]);
4927                let gsum = q1_group_sums(&act.xq, gpr);
4928                (act, gsum)
4929            })
4930            .collect();
4931        let acts = &acts;
4932        #[cfg(target_arch = "x86_64")]
4933        let blocked_ok = avx2_enabled()
4934            && std::env::var("CMF_X86_BLOCKED")
4935                .map(|v| v != "0")
4936                .unwrap_or(true);
4937        #[cfg(target_arch = "aarch64")]
4938        let blocked_ok = sdot_enabled()
4939            && std::env::var("CMF_X86_BLOCKED")
4940                .map(|v| v != "0")
4941                .unwrap_or(true);
4942        let run = move |start: usize, end: usize| {
4943            for r in start..end {
4944                let mut bi = 0usize;
4945                // Blocked 1×4: the unpacked bit mask serves four
4946                // activation streams per group.
4947                #[cfg(target_arch = "aarch64")]
4948                if blocked_ok {
4949                    while bi + 4 <= acts.len() {
4950                        let xs = [
4951                            acts[bi].0.xq.as_slice(),
4952                            acts[bi + 1].0.xq.as_slice(),
4953                            acts[bi + 2].0.xq.as_slice(),
4954                            acts[bi + 3].0.xq.as_slice(),
4955                        ];
4956                        let gs = [
4957                            acts[bi].1.as_slice(),
4958                            acts[bi + 1].1.as_slice(),
4959                            acts[bi + 2].1.as_slice(),
4960                            acts[bi + 3].1.as_slice(),
4961                        ];
4962                        let d = unsafe { dot_q1_row_1x4_sdot(bytes, r, gpr, xs, gs) };
4963                        for k in 0..4 {
4964                            let (act, _) = &acts[bi + k];
4965                            let mut acc = d[k] * act.sx;
4966                            for &(j, xv) in &act.outliers {
4967                                let (w, sc) = q1_outlier(bytes, r, gpr, j);
4968                                acc += w * sc * xv;
4969                            }
4970                            // SAFETY: disjoint (bi, r) cells per worker.
4971                            unsafe { *out_addr.at((bi + k) * rows + r) = acc };
4972                        }
4973                        bi += 4;
4974                    }
4975                }
4976                #[cfg(target_arch = "x86_64")]
4977                if blocked_ok {
4978                    while bi + 4 <= acts.len() {
4979                        let xs = [
4980                            acts[bi].0.xq.as_slice(),
4981                            acts[bi + 1].0.xq.as_slice(),
4982                            acts[bi + 2].0.xq.as_slice(),
4983                            acts[bi + 3].0.xq.as_slice(),
4984                        ];
4985                        let gs = [
4986                            acts[bi].1.as_slice(),
4987                            acts[bi + 1].1.as_slice(),
4988                            acts[bi + 2].1.as_slice(),
4989                            acts[bi + 3].1.as_slice(),
4990                        ];
4991                        let d = unsafe {
4992                            if vnni_tiles_enabled() {
4993                                dot_q1_row_1x4_vnni(bytes, r, gpr, xs, gs)
4994                            } else {
4995                                dot_q1_row_1x4_avx2(bytes, r, gpr, xs, gs)
4996                            }
4997                        };
4998                        for k in 0..4 {
4999                            let (act, _) = &acts[bi + k];
5000                            let mut acc = d[k] * act.sx;
5001                            for &(j, xv) in &act.outliers {
5002                                let (w, sc) = q1_outlier(bytes, r, gpr, j);
5003                                acc += w * sc * xv;
5004                            }
5005                            // SAFETY: disjoint (bi, r) cells per worker.
5006                            unsafe { *out_addr.at((bi + k) * rows + r) = acc };
5007                        }
5008                        bi += 4;
5009                    }
5010                }
5011                while bi < acts.len() {
5012                    let (act, gsum) = &acts[bi];
5013                    let mut acc = dot_q1_row_i8(bytes, r, gpr, &act.xq, gsum) * act.sx;
5014                    for &(j, xv) in &act.outliers {
5015                        let (w, s) = q1_outlier(bytes, r, gpr, j);
5016                        acc += w * s * xv;
5017                    }
5018                    // SAFETY: disjoint (bi, r) cells per worker range.
5019                    unsafe { *out_addr.at(bi * rows + r) = acc };
5020                    bi += 1;
5021                }
5022            }
5023        };
5024        dispatch_rows(pool, rows, &run);
5025        return;
5026    }
5027    let run = move |start: usize, end: usize| {
5028        for r in start..end {
5029            for bi in 0..b {
5030                let x = &xs_all[bi * cols..(bi + 1) * cols];
5031                // SAFETY: disjoint (bi, r) cells per worker range.
5032                unsafe { *out_addr.at(bi * rows + r) = q1_row_exact(bytes, r, gpr, x) };
5033            }
5034        }
5035    };
5036    dispatch_rows(pool, rows, &run);
5037}
5038
5039/// Fused q4_block matvec straight from the mapped bytes. SDOT path when
5040/// dotprod is available (port of vmfcore `dot_q4_block_sdot`, measured
5041/// +23% on q4 decode): nibbles → centered i8, int8×int8 `sdot` per
5042/// 32-group, exact outlier correction — the same A8W8 contract as q8.
5043/// `CMF_SDOT=0` keeps the exact scalar path.
5044fn q4matvec(
5045    bytes: &[u8],
5046    x: &[f32],
5047    rows: usize,
5048    cols: usize,
5049    out: &mut [f32],
5050    pool: Option<&Pool>,
5051) {
5052    debug_assert_eq!(out.len(), rows);
5053    let (packed, scales) = q4_split(bytes, rows, cols);
5054    let gpr = cols / GROUP_SIZE;
5055    let out_addr = SendMut(out.as_mut_ptr());
5056
5057    if a8w8_enabled() {
5058        let act = split_act(x);
5059        let run = move |start: usize, end: usize| {
5060            q4_range_a8w8(packed, scales, gpr, cols, &act, out_addr, start, end)
5061        };
5062        dispatch_rows(pool, rows, &run);
5063        return;
5064    }
5065
5066    let run =
5067        move |start: usize, end: usize| q4_range_f32(packed, scales, gpr, x, out_addr, start, end);
5068    dispatch_rows(pool, rows, &run);
5069}
5070
5071/// One q4 row via the A8W8 int8 path — SDOT on ARM, AVX2 maddubs on
5072/// x86 (scalar fallback is unreachable: callers gate on a8w8_enabled).
5073#[inline]
5074#[allow(unreachable_code)]
5075/// One UNPACKED q4 row (centered i8 in `buf`) against four activation
5076/// streams: the 32-byte weight chunk and its abs() load once per group,
5077/// the per-group f16 scale decodes once — four maddubs+reduce chains
5078/// instead of four full (load, abs, dot) rounds.
5079#[cfg(target_arch = "x86_64")]
5080#[target_feature(enable = "avx2")]
5081unsafe fn dot_q4b_row_1x4_avx2(
5082    buf: &[u8],
5083    scales: &[u8],
5084    g0: usize,
5085    gpr: usize,
5086    xs: [&[i8]; 4],
5087) -> [f32; 4] {
5088    // SAFETY: callers uphold buffer contracts (buf.len() == gpr·32).
5089    unsafe {
5090        use core::arch::x86_64::*;
5091        let ones = _mm256_set1_epi16(1);
5092        let mut acc = [0f32; 4];
5093        for gi in 0..gpr {
5094            let s = f16_to_f32(u16::from_le_bytes([
5095                scales[(g0 + gi) * 2],
5096                scales[(g0 + gi) * 2 + 1],
5097            ]));
5098            let w = _mm256_loadu_si256(buf.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
5099            let aw = _mm256_abs_epi8(w);
5100            for (k, xq) in xs.iter().enumerate() {
5101                let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
5102                let p16 = _mm256_maddubs_epi16(aw, _mm256_sign_epi8(x, w));
5103                let d = _mm256_madd_epi16(p16, ones);
5104                let hi128 = _mm256_extracti128_si256::<1>(d);
5105                let s128 = _mm_add_epi32(_mm256_castsi256_si128(d), hi128);
5106                let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
5107                let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
5108                acc[k] += _mm_cvtsi128_si32(s32) as f32 * s;
5109            }
5110        }
5111        acc
5112    }
5113}
5114
5115/// VNNI twin of `dot_q4b_row_1x4_avx2` (see `dpbusd_hsum`).
5116#[cfg(target_arch = "x86_64")]
5117#[target_feature(enable = "avx2,avx512f,avx512bw,avx512vl,avx512vnni")]
5118unsafe fn dot_q4b_row_1x4_vnni(
5119    buf: &[u8],
5120    scales: &[u8],
5121    g0: usize,
5122    gpr: usize,
5123    xs: [&[i8]; 4],
5124) -> [f32; 4] {
5125    // SAFETY: callers uphold buffer contracts (buf.len() == gpr·32).
5126    unsafe {
5127        use core::arch::x86_64::*;
5128        let mut acc = [0f32; 4];
5129        for gi in 0..gpr {
5130            let s = f16_to_f32(u16::from_le_bytes([
5131                scales[(g0 + gi) * 2],
5132                scales[(g0 + gi) * 2 + 1],
5133            ]));
5134            let w = _mm256_loadu_si256(buf.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
5135            let aw = _mm256_abs_epi8(w);
5136            for (k, xq) in xs.iter().enumerate() {
5137                let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
5138                let d = dpbusd_hsum(aw, _mm256_sign_epi8(x, w));
5139                acc[k] += d as f32 * s;
5140            }
5141        }
5142        acc
5143    }
5144}
5145
5146/// The vbit flavor of the blocked 1×4: the per-activation A8W8 scale
5147/// folds in PER GROUP as `(d·sx)·s` — bit-matching the single-matvec
5148/// accumulation order (the q4_block flavor applies sx once at the end,
5149/// matching ITS single path; the two conventions are historical and
5150/// each blocked leg must mirror its own).
5151#[cfg(target_arch = "x86_64")]
5152#[target_feature(enable = "avx2")]
5153unsafe fn dot_q4b_row_1x4_sx_avx2(
5154    buf: &[u8],
5155    scales: &[u8],
5156    g0: usize,
5157    gpr: usize,
5158    xs: [&[i8]; 4],
5159    sxs: [f32; 4],
5160) -> [f32; 4] {
5161    // SAFETY: callers uphold buffer contracts (buf.len() == gpr·32).
5162    unsafe {
5163        use core::arch::x86_64::*;
5164        let ones = _mm256_set1_epi16(1);
5165        let mut acc = [0f32; 4];
5166        for gi in 0..gpr {
5167            let s = f16_to_f32(u16::from_le_bytes([
5168                scales[(g0 + gi) * 2],
5169                scales[(g0 + gi) * 2 + 1],
5170            ]));
5171            let w = _mm256_loadu_si256(buf.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
5172            let aw = _mm256_abs_epi8(w);
5173            for (k, xq) in xs.iter().enumerate() {
5174                let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
5175                let p16 = _mm256_maddubs_epi16(aw, _mm256_sign_epi8(x, w));
5176                let d = _mm256_madd_epi16(p16, ones);
5177                let hi128 = _mm256_extracti128_si256::<1>(d);
5178                let s128 = _mm_add_epi32(_mm256_castsi256_si128(d), hi128);
5179                let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
5180                let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
5181                acc[k] += (_mm_cvtsi128_si32(s32) as f32 * sxs[k]) * s;
5182            }
5183        }
5184        acc
5185    }
5186}
5187
5188/// VNNI twin of `dot_q4b_row_1x4_sx_avx2` (see `dpbusd_hsum`; the
5189/// per-group `(d·sx)·s` fold mirrors the vbit single path).
5190#[cfg(target_arch = "x86_64")]
5191#[target_feature(enable = "avx2,avx512f,avx512bw,avx512vl,avx512vnni")]
5192unsafe fn dot_q4b_row_1x4_sx_vnni(
5193    buf: &[u8],
5194    scales: &[u8],
5195    g0: usize,
5196    gpr: usize,
5197    xs: [&[i8]; 4],
5198    sxs: [f32; 4],
5199) -> [f32; 4] {
5200    // SAFETY: callers uphold buffer contracts (buf.len() == gpr·32).
5201    unsafe {
5202        use core::arch::x86_64::*;
5203        let mut acc = [0f32; 4];
5204        for gi in 0..gpr {
5205            let s = f16_to_f32(u16::from_le_bytes([
5206                scales[(g0 + gi) * 2],
5207                scales[(g0 + gi) * 2 + 1],
5208            ]));
5209            let w = _mm256_loadu_si256(buf.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
5210            let aw = _mm256_abs_epi8(w);
5211            for (k, xq) in xs.iter().enumerate() {
5212                let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
5213                let d = dpbusd_hsum(aw, _mm256_sign_epi8(x, w));
5214                acc[k] += (d as f32 * sxs[k]) * s;
5215            }
5216        }
5217        acc
5218    }
5219}
5220
5221#[allow(unreachable_code)]
5222fn dot_q4_row_i8(packed: &[u8], scales: &[u8], g0: usize, gpr: usize, xq: &[i8]) -> f32 {
5223    #[cfg(target_arch = "aarch64")]
5224    unsafe {
5225        return dot_q4_row_sdot(packed, scales, g0, gpr, xq);
5226    }
5227    #[cfg(target_arch = "x86_64")]
5228    unsafe {
5229        return dot_q4_row_avx2(packed, scales, g0, gpr, xq);
5230    }
5231    let mut acc = 0f32;
5232    for gi in 0..gpr {
5233        let g = g0 + gi;
5234        let s = f16_to_f32(u16::from_le_bytes([scales[g * 2], scales[g * 2 + 1]]));
5235        let mut d = 0i32;
5236        for (k, &b) in packed[g * 16..(g + 1) * 16].iter().enumerate() {
5237            d += ((b & 0x0F) as i32 - 8) * xq[gi * GROUP_SIZE + k * 2] as i32
5238                + (((b >> 4) & 0x0F) as i32 - 8) * xq[gi * GROUP_SIZE + k * 2 + 1] as i32;
5239        }
5240        acc += d as f32 * s;
5241    }
5242    acc
5243}
5244
5245/// Two-activation q4 row via the A8W8 int8 path (see `dot_q4_row_i8`).
5246#[inline]
5247#[allow(unreachable_code)]
5248fn dot_q4_row_i8_2(
5249    packed: &[u8],
5250    scales: &[u8],
5251    g0: usize,
5252    gpr: usize,
5253    xq1: &[i8],
5254    xq2: &[i8],
5255) -> (f32, f32) {
5256    #[cfg(target_arch = "aarch64")]
5257    unsafe {
5258        return dot_q4_row_sdot2(packed, scales, g0, gpr, xq1, xq2);
5259    }
5260    #[cfg(target_arch = "x86_64")]
5261    unsafe {
5262        return dot_q4_row_avx2_2(packed, scales, g0, gpr, xq1, xq2);
5263    }
5264    (
5265        dot_q4_row_i8(packed, scales, g0, gpr, xq1),
5266        dot_q4_row_i8(packed, scales, g0, gpr, xq2),
5267    )
5268}
5269
5270/// One q4 row range via SDOT (kernel body of `q4matvec`, extracted so
5271/// multi-matrix jobs can drive it for several tensors in one dispatch).
5272#[allow(clippy::too_many_arguments)]
5273fn q4_range_a8w8(
5274    packed: &[u8],
5275    scales: &[u8],
5276    gpr: usize,
5277    cols: usize,
5278    act: &SplitAct,
5279    out: SendMut,
5280    start: usize,
5281    end: usize,
5282) {
5283    for r in start..end {
5284        let mut acc = dot_q4_row_i8(packed, scales, r * gpr, gpr, &act.xq) * act.sx;
5285        // xq is zeroed at outlier slots — add the exact terms.
5286        for &(j, xv) in &act.outliers {
5287            let flat = r * cols + j;
5288            let byte = packed[flat / 2];
5289            let nib = if flat & 1 == 0 {
5290                byte & 0x0F
5291            } else {
5292                byte >> 4
5293            };
5294            let s = f16_to_f32(u16::from_le_bytes([
5295                scales[(flat / GROUP_SIZE) * 2],
5296                scales[(flat / GROUP_SIZE) * 2 + 1],
5297            ]));
5298            acc += ((nib as i32 - 8) as f32) * s * xv;
5299        }
5300        // SAFETY: disjoint row ranges per worker.
5301        unsafe { *out.at(r) = acc };
5302    }
5303}
5304
5305/// Two-input q4 row range via the A8W8 int8 path — kernel body of
5306/// `q4matvec2`, extracted for pair multi-matrix jobs.
5307#[allow(clippy::too_many_arguments)]
5308fn q4_range2_a8w8(
5309    packed: &[u8],
5310    scales: &[u8],
5311    gpr: usize,
5312    cols: usize,
5313    a1: &SplitAct,
5314    a2: &SplitAct,
5315    p1: SendMut,
5316    p2: SendMut,
5317    start: usize,
5318    end: usize,
5319) {
5320    for r in start..end {
5321        let (s1, s2) = dot_q4_row_i8_2(packed, scales, r * gpr, gpr, &a1.xq, &a2.xq);
5322        let mut acc1 = s1 * a1.sx;
5323        let mut acc2 = s2 * a2.sx;
5324        // xq is zeroed at outlier slots — add the exact terms.
5325        let fix = |outliers: &[(usize, f32)], acc: &mut f32| {
5326            for &(j, xv) in outliers {
5327                let flat = r * cols + j;
5328                let byte = packed[flat / 2];
5329                let nib = if flat & 1 == 0 {
5330                    byte & 0x0F
5331                } else {
5332                    byte >> 4
5333                };
5334                let s = f16_to_f32(u16::from_le_bytes([
5335                    scales[(flat / GROUP_SIZE) * 2],
5336                    scales[(flat / GROUP_SIZE) * 2 + 1],
5337                ]));
5338                *acc += ((nib as i32 - 8) as f32) * s * xv;
5339            }
5340        };
5341        fix(&a1.outliers, &mut acc1);
5342        fix(&a2.outliers, &mut acc2);
5343        // SAFETY: disjoint row ranges per worker.
5344        unsafe {
5345            *p1.at(r) = acc1;
5346            *p2.at(r) = acc2;
5347        }
5348    }
5349}
5350
5351/// Exact scalar q4 row range (same extraction, non-SDOT path).
5352fn q4_range_f32(
5353    packed: &[u8],
5354    scales: &[u8],
5355    gpr: usize,
5356    x: &[f32],
5357    out: SendMut,
5358    start: usize,
5359    end: usize,
5360) {
5361    for r in start..end {
5362        let mut acc = 0f32;
5363        for gi in 0..gpr {
5364            let g = r * gpr + gi;
5365            let s = f16_to_f32(u16::from_le_bytes([scales[g * 2], scales[g * 2 + 1]]));
5366            let pk = &packed[g * 16..(g + 1) * 16];
5367            let xg = &x[gi * GROUP_SIZE..(gi + 1) * GROUP_SIZE];
5368            let mut ga = 0f32;
5369            for (k, &b) in pk.iter().enumerate() {
5370                ga += ((b & 0x0F) as f32 - 8.0) * xg[k * 2]
5371                    + (((b >> 4) & 0x0F) as f32 - 8.0) * xg[k * 2 + 1];
5372            }
5373            acc += ga * s;
5374        }
5375        // SAFETY: disjoint row ranges per worker.
5376        unsafe { *out.at(r) = acc };
5377    }
5378}
5379
5380/// Fused two-input q4 matvec: nibbles are unpacked ONCE per group and
5381/// dotted against both activations (was: two full matvecs — double
5382/// weight traffic). Per-lane math matches `q4matvec` exactly.
5383#[allow(clippy::too_many_arguments)]
5384fn q4matvec2(
5385    bytes: &[u8],
5386    x1: &[f32],
5387    x2: &[f32],
5388    rows: usize,
5389    cols: usize,
5390    o1: &mut [f32],
5391    o2: &mut [f32],
5392    pool: Option<&Pool>,
5393) {
5394    debug_assert_eq!(o1.len(), rows);
5395    debug_assert_eq!(o2.len(), rows);
5396    let (packed, scales) = q4_split(bytes, rows, cols);
5397    let gpr = cols / GROUP_SIZE;
5398
5399    if a8w8_enabled() {
5400        let a1 = split_act(x1);
5401        let a2 = split_act(x2);
5402        let p1 = SendMut(o1.as_mut_ptr());
5403        let p2 = SendMut(o2.as_mut_ptr());
5404        let run = move |start: usize, end: usize| {
5405            q4_range2_a8w8(packed, scales, gpr, cols, &a1, &a2, p1, p2, start, end)
5406        };
5407        dispatch_rows(pool, rows, &run);
5408        return;
5409    }
5410
5411    let p1 = SendMut(o1.as_mut_ptr());
5412    let p2 = SendMut(o2.as_mut_ptr());
5413    let run = move |start: usize, end: usize| {
5414        q4_range2_f32(packed, scales, gpr, x1, x2, p1, p2, start, end)
5415    };
5416    dispatch_rows(pool, rows, &run);
5417}
5418
5419/// Two-input exact scalar q4 row range (same extraction).
5420#[allow(clippy::too_many_arguments)]
5421fn q4_range2_f32(
5422    packed: &[u8],
5423    scales: &[u8],
5424    gpr: usize,
5425    x1: &[f32],
5426    x2: &[f32],
5427    p1: SendMut,
5428    p2: SendMut,
5429    start: usize,
5430    end: usize,
5431) {
5432    for r in start..end {
5433        let (mut acc1, mut acc2) = (0f32, 0f32);
5434        for gi in 0..gpr {
5435            let g = r * gpr + gi;
5436            let s = f16_to_f32(u16::from_le_bytes([scales[g * 2], scales[g * 2 + 1]]));
5437            let pk = &packed[g * 16..(g + 1) * 16];
5438            let x1g = &x1[gi * GROUP_SIZE..(gi + 1) * GROUP_SIZE];
5439            let x2g = &x2[gi * GROUP_SIZE..(gi + 1) * GROUP_SIZE];
5440            let (mut g1, mut g2) = (0f32, 0f32);
5441            for (k, &b) in pk.iter().enumerate() {
5442                let wl = (b & 0x0F) as f32 - 8.0;
5443                let wh = ((b >> 4) & 0x0F) as f32 - 8.0;
5444                g1 += wl * x1g[k * 2] + wh * x1g[k * 2 + 1];
5445                g2 += wl * x2g[k * 2] + wh * x2g[k * 2 + 1];
5446            }
5447            acc1 += g1 * s;
5448            acc2 += g2 * s;
5449        }
5450        // SAFETY: disjoint row ranges per worker.
5451        unsafe {
5452            *p1.at(r) = acc1;
5453            *p2.at(r) = acc2;
5454        }
5455    }
5456}
5457
5458thread_local! {
5459    /// Per-worker decoded-row scratch for the batched q4/vbit kernels
5460    /// (centered i8 for SDOT, f32 for the exact/scalar paths).
5461    static ROW_I8: std::cell::RefCell<Vec<u8>> = const { std::cell::RefCell::new(Vec::new()) };
5462    static ROW_F32: std::cell::RefCell<Vec<f32>> = const { std::cell::RefCell::new(Vec::new()) };
5463}
5464
5465/// Batched q4 matmat: each weight row is unpacked from the mmap ONCE
5466/// and dotted against ALL b activations (prefill used to fall back to b
5467/// full matvecs — b× weight traffic and b× nibble decode). Per-position
5468/// math matches `q4matvec` exactly: same group order, same accumulation.
5469/// `out` is row-major [b, rows] like `qmatmat`.
5470#[allow(clippy::too_many_arguments)]
5471fn q4matmat(
5472    bytes: &[u8],
5473    xs_all: &[f32],
5474    b: usize,
5475    rows: usize,
5476    cols: usize,
5477    out: &mut [f32],
5478    pool: Option<&Pool>,
5479) {
5480    debug_assert_eq!(xs_all.len(), b * cols);
5481    debug_assert_eq!(out.len(), b * rows);
5482    let (packed, scales) = q4_split(bytes, rows, cols);
5483    let gpr = cols / GROUP_SIZE;
5484    let gscale = |g: usize| f16_to_f32(u16::from_le_bytes([scales[g * 2], scales[g * 2 + 1]]));
5485
5486    if a8w8_enabled() {
5487        let acts: Vec<SplitAct> = (0..b)
5488            .map(|bi| split_act(&xs_all[bi * cols..(bi + 1) * cols]))
5489            .collect();
5490        let acts = &acts;
5491        let out_addr = SendMut(out.as_mut_ptr());
5492        let run = move |start: usize, end: usize| {
5493            ROW_I8.with(|rb| {
5494                let mut buf = rb.borrow_mut();
5495                buf.resize(cols, 0);
5496                for r in start..end {
5497                    // Unpack the row's nibbles to centered i8 once
5498                    // (element 2k = low nibble, 2k+1 = high — flat order,
5499                    // same as dot_q4_row_sdot's zip).
5500                    for gi in 0..gpr {
5501                        let g = r * gpr + gi;
5502                        for (k, &bt) in packed[g * 16..(g + 1) * 16].iter().enumerate() {
5503                            buf[gi * GROUP_SIZE + k * 2] = ((bt & 0x0F) as i32 - 8) as i8 as u8;
5504                            buf[gi * GROUP_SIZE + k * 2 + 1] =
5505                                (((bt >> 4) & 0x0F) as i32 - 8) as i8 as u8;
5506                        }
5507                    }
5508                    let mut bi = 0usize;
5509                    #[cfg(target_arch = "x86_64")]
5510                    if avx2_enabled()
5511                        && std::env::var("CMF_X86_BLOCKED")
5512                            .map(|v| v != "0")
5513                            .unwrap_or(true)
5514                    {
5515                        while bi + 4 <= acts.len() {
5516                            let xs = [
5517                                acts[bi].xq.as_slice(),
5518                                acts[bi + 1].xq.as_slice(),
5519                                acts[bi + 2].xq.as_slice(),
5520                                acts[bi + 3].xq.as_slice(),
5521                            ];
5522                            let d = unsafe {
5523                                if vnni_tiles_enabled() {
5524                                    dot_q4b_row_1x4_vnni(&buf, scales, r * gpr, gpr, xs)
5525                                } else {
5526                                    dot_q4b_row_1x4_avx2(&buf, scales, r * gpr, gpr, xs)
5527                                }
5528                            };
5529                            for k in 0..4 {
5530                                let act = &acts[bi + k];
5531                                let mut acc = d[k] * act.sx;
5532                                for &(j, xv) in &act.outliers {
5533                                    acc += (buf[j] as i8) as f32
5534                                        * gscale((r * cols + j) / GROUP_SIZE)
5535                                        * xv;
5536                                }
5537                                // SAFETY: disjoint (bi, r) cells per worker.
5538                                unsafe { *out_addr.at((bi + k) * rows + r) = acc };
5539                            }
5540                            bi += 4;
5541                        }
5542                    }
5543                    while bi < acts.len() {
5544                        let act = &acts[bi];
5545                        let mut acc = 0f32;
5546                        for gi in 0..gpr {
5547                            let d = dot_i8_i8(
5548                                &buf[gi * GROUP_SIZE..(gi + 1) * GROUP_SIZE],
5549                                &act.xq[gi * GROUP_SIZE..(gi + 1) * GROUP_SIZE],
5550                            );
5551                            acc += d as f32 * gscale(r * gpr + gi);
5552                        }
5553                        acc *= act.sx;
5554                        // xq is zeroed at outlier slots — exact terms.
5555                        for &(j, xv) in &act.outliers {
5556                            acc += (buf[j] as i8) as f32 * gscale((r * cols + j) / GROUP_SIZE) * xv;
5557                        }
5558                        // SAFETY: disjoint (bi, r) cells per worker row range.
5559                        unsafe { *out_addr.at(bi * rows + r) = acc };
5560                        bi += 1;
5561                    }
5562                }
5563            })
5564        };
5565        dispatch_rows(pool, rows, &run);
5566        return;
5567    }
5568
5569    let out_addr = SendMut(out.as_mut_ptr());
5570    let run = move |start: usize, end: usize| {
5571        ROW_F32.with(|rb| {
5572            let mut buf = rb.borrow_mut();
5573            buf.resize(cols, 0.0);
5574            for r in start..end {
5575                // Decode raw (nib − 8) values once; scales stay per-group
5576                // so the accumulation order matches q4matvec bit-for-bit.
5577                for gi in 0..gpr {
5578                    let g = r * gpr + gi;
5579                    for (k, &bt) in packed[g * 16..(g + 1) * 16].iter().enumerate() {
5580                        buf[gi * GROUP_SIZE + k * 2] = (bt & 0x0F) as f32 - 8.0;
5581                        buf[gi * GROUP_SIZE + k * 2 + 1] = ((bt >> 4) & 0x0F) as f32 - 8.0;
5582                    }
5583                }
5584                for bi in 0..b {
5585                    let x = &xs_all[bi * cols..(bi + 1) * cols];
5586                    let mut acc = 0f32;
5587                    for gi in 0..gpr {
5588                        let mut ga = 0f32;
5589                        // Pairwise (lo + hi) addition, matching
5590                        // q4matvec's `ga += lo·x + hi·x` shape exactly —
5591                        // a flat one-per-element loop rounds differently
5592                        // and broke bit-parity on the scalar (x86) path.
5593                        for k in 0..GROUP_SIZE / 2 {
5594                            let e = gi * GROUP_SIZE + k * 2;
5595                            ga += buf[e] * x[e] + buf[e + 1] * x[e + 1];
5596                        }
5597                        acc += ga * gscale(r * gpr + gi);
5598                    }
5599                    // SAFETY: disjoint (bi, r) cells per worker row range.
5600                    unsafe { *out_addr.at(bi * rows + r) = acc };
5601                }
5602            }
5603        })
5604    };
5605    dispatch_rows(pool, rows, &run);
5606}
5607
5608/// Batched vbit matmat: each variable-bit row is decoded from the mmap
5609/// ONCE for the whole microbatch. Same per-position math as
5610/// `vbitmatvec` (SDOT A8W8 with exact outliers / exact f32 for b=8 rows
5611/// and the scalar path).
5612#[allow(clippy::too_many_arguments)]
5613fn vbitmatmat(
5614    bytes: &[u8],
5615    offsets: &[usize],
5616    xs_all: &[f32],
5617    b: usize,
5618    rows: usize,
5619    cols: usize,
5620    out: &mut [f32],
5621    pool: Option<&Pool>,
5622) {
5623    debug_assert_eq!(xs_all.len(), b * cols);
5624    debug_assert_eq!(out.len(), b * rows);
5625    debug_assert_eq!(offsets.len(), rows + 1);
5626    let ng = cols / GROUP_SIZE;
5627    let bits = &bytes[..rows];
5628    let sc_off = rows;
5629    let gscale = |r: usize, g: usize| {
5630        let so = (r * ng + g) * 2;
5631        f16_to_f32(u16::from_le_bytes([
5632            bytes[sc_off + so],
5633            bytes[sc_off + so + 1],
5634        ]))
5635    };
5636
5637    // Decode row r's raw (u − L) values into `dst` (f32, unscaled).
5638    let decode_f32 = |r: usize, dst: &mut [f32]| {
5639        let bw = bits[r] as usize;
5640        let l = ((1i32 << (bw - 1)) - 1) as f32;
5641        let data = &bytes[offsets[r]..offsets[r + 1]];
5642        let (mut acc, mut nbits, mut idx) = (0u64, 0usize, 0usize);
5643        for d in dst.iter_mut() {
5644            while nbits < bw {
5645                acc = (acc << 8) | data[idx] as u64;
5646                idx += 1;
5647                nbits += 8;
5648            }
5649            let u = ((acc >> (nbits - bw)) & ((1u64 << bw) - 1)) as f32;
5650            nbits -= bw;
5651            *d = u - l;
5652        }
5653    };
5654
5655    if a8w8_enabled() {
5656        let acts: Vec<SplitAct> = (0..b)
5657            .map(|bi| split_act(&xs_all[bi * cols..(bi + 1) * cols]))
5658            .collect();
5659        let acts = &acts;
5660        let out_addr = SendMut(out.as_mut_ptr());
5661        let run = move |start: usize, end: usize| {
5662            for r in start..end {
5663                let bw = bits[r] as usize;
5664                if bw == 8 {
5665                    // u−L reaches 128 → no i8 path; decode once, exact
5666                    // f32 dots for every position (same as vbitmatvec).
5667                    ROW_F32.with(|rb| {
5668                        let mut buf = rb.borrow_mut();
5669                        buf.resize(cols, 0.0);
5670                        decode_f32(r, &mut buf);
5671                        for bi in 0..b {
5672                            let x = &xs_all[bi * cols..(bi + 1) * cols];
5673                            let mut dot = 0f32;
5674                            for g in 0..ng {
5675                                let mut gd = 0f32;
5676                                for k in 0..GROUP_SIZE {
5677                                    gd += buf[g * GROUP_SIZE + k] * x[g * GROUP_SIZE + k];
5678                                }
5679                                dot += gd * gscale(r, g);
5680                            }
5681                            // SAFETY: disjoint (bi, r) cells per worker range.
5682                            unsafe { *out_addr.at(bi * rows + r) = dot };
5683                        }
5684                    });
5685                    continue;
5686                }
5687                let l = (1i32 << (bw - 1)) - 1;
5688                let data = &bytes[offsets[r]..offsets[r + 1]];
5689                ROW_I8.with(|rb| {
5690                    let mut buf = rb.borrow_mut();
5691                    buf.resize(cols, 0);
5692                    #[inline(always)]
5693                    fn fill<const B: usize>(data: &[u8], l: i32, buf: &mut [u8]) {
5694                        for (blk, chunk) in buf.chunks_exact_mut(8).enumerate() {
5695                            let u = unpack8::<B>(&data[blk * B..]);
5696                            for k in 0..8 {
5697                                chunk[k] = (u[k] - l) as i8 as u8;
5698                            }
5699                        }
5700                    }
5701                    match bw {
5702                        3 => fill::<3>(data, l, &mut buf),
5703                        4 => vbit_fill4(data, &mut buf),
5704                        5 => fill::<5>(data, l, &mut buf),
5705                        6 => fill::<6>(data, l, &mut buf),
5706                        _ => unreachable!("vbit bit-width {bw} (validated at load)"),
5707                    }
5708                    let mut bi = 0usize;
5709                    // The vbit scale table shares q4_block's layout
5710                    // (contiguous f16 per (row·ng + g)), so the same
5711                    // blocked 1×4 kernel serves the decoded row.
5712                    #[cfg(target_arch = "x86_64")]
5713                    if avx2_enabled()
5714                        && std::env::var("CMF_X86_BLOCKED")
5715                            .map(|v| v != "0")
5716                            .unwrap_or(true)
5717                    {
5718                        while bi + 4 <= acts.len() {
5719                            let xs = [
5720                                acts[bi].xq.as_slice(),
5721                                acts[bi + 1].xq.as_slice(),
5722                                acts[bi + 2].xq.as_slice(),
5723                                acts[bi + 3].xq.as_slice(),
5724                            ];
5725                            let sxs = [
5726                                acts[bi].sx,
5727                                acts[bi + 1].sx,
5728                                acts[bi + 2].sx,
5729                                acts[bi + 3].sx,
5730                            ];
5731                            let d = unsafe {
5732                                if vnni_tiles_enabled() {
5733                                    dot_q4b_row_1x4_sx_vnni(
5734                                        &buf,
5735                                        &bytes[sc_off..],
5736                                        r * ng,
5737                                        ng,
5738                                        xs,
5739                                        sxs,
5740                                    )
5741                                } else {
5742                                    dot_q4b_row_1x4_sx_avx2(
5743                                        &buf,
5744                                        &bytes[sc_off..],
5745                                        r * ng,
5746                                        ng,
5747                                        xs,
5748                                        sxs,
5749                                    )
5750                                }
5751                            };
5752                            for k in 0..4 {
5753                                let act = &acts[bi + k];
5754                                let mut dot = d[k];
5755                                for &(j, xv) in &act.outliers {
5756                                    dot += (buf[j] as i8) as f32 * gscale(r, j / GROUP_SIZE) * xv;
5757                                }
5758                                // SAFETY: disjoint (bi, r) cells per worker.
5759                                unsafe { *out_addr.at((bi + k) * rows + r) = dot };
5760                            }
5761                            bi += 4;
5762                        }
5763                    }
5764                    while bi < acts.len() {
5765                        let act = &acts[bi];
5766                        let mut dot = 0f32;
5767                        for g in 0..ng {
5768                            let d = dot_i8_i8(
5769                                &buf[g * GROUP_SIZE..(g + 1) * GROUP_SIZE],
5770                                &act.xq[g * GROUP_SIZE..(g + 1) * GROUP_SIZE],
5771                            ) as f32
5772                                * act.sx;
5773                            dot += d * gscale(r, g);
5774                        }
5775                        for &(j, xv) in &act.outliers {
5776                            dot += (buf[j] as i8) as f32 * gscale(r, j / GROUP_SIZE) * xv;
5777                        }
5778                        // SAFETY: disjoint (bi, r) cells per worker range.
5779                        unsafe { *out_addr.at(bi * rows + r) = dot };
5780                        bi += 1;
5781                    }
5782                });
5783            }
5784        };
5785        dispatch_rows(pool, rows, &run);
5786        return;
5787    }
5788
5789    let out_addr = SendMut(out.as_mut_ptr());
5790    let run = move |start: usize, end: usize| {
5791        ROW_F32.with(|rb| {
5792            let mut buf = rb.borrow_mut();
5793            buf.resize(cols, 0.0);
5794            for r in start..end {
5795                decode_f32(r, &mut buf);
5796                for bi in 0..b {
5797                    let x = &xs_all[bi * cols..(bi + 1) * cols];
5798                    let mut dot = 0f32;
5799                    for g in 0..ng {
5800                        let mut gd = 0f32;
5801                        for k in 0..GROUP_SIZE {
5802                            gd += buf[g * GROUP_SIZE + k] * x[g * GROUP_SIZE + k];
5803                        }
5804                        dot += gd * gscale(r, g);
5805                    }
5806                    // SAFETY: disjoint (bi, r) cells per worker range.
5807                    unsafe { *out_addr.at(bi * rows + r) = dot };
5808                }
5809            }
5810        })
5811    };
5812    dispatch_rows(pool, rows, &run);
5813}
5814
5815/// Build a GPU batch job for a q8-family mapped tensor (primary
5816/// shard): prescaled input + directory coordinates. None → not
5817/// GPU-eligible, caller stays on the CPU.
5818pub(crate) fn gpu_batch_job<'a>(
5819    t: &'a QTensor,
5820    x: &[f32],
5821) -> Option<(std::sync::Arc<CmfModel>, crate::gpu::BatchJob<'a>)> {
5822    match t {
5823        QTensor::Mapped {
5824            model,
5825            idx,
5826            dtype: dt @ (TensorDtype::Q8Row | TensorDtype::Q8_2f),
5827            rows,
5828            cols,
5829            row_scale,
5830            col_field,
5831            ..
5832        } => Some((
5833            model.clone(),
5834            crate::gpu::BatchJob {
5835                idx: *idx,
5836                rows: *rows,
5837                cols: *cols,
5838                row_scale,
5839                xs: prescale(x, col_field, *dt).into_owned(),
5840                q1: false,
5841            },
5842        )),
5843        // q1: raw f32 activations, tile-embedded scales.
5844        QTensor::Mapped {
5845            model,
5846            idx,
5847            dtype: TensorDtype::Q1,
5848            rows,
5849            cols,
5850            ..
5851        } => Some((
5852            model.clone(),
5853            crate::gpu::BatchJob {
5854                idx: *idx,
5855                rows: *rows,
5856                cols: *cols,
5857                row_scale: &[],
5858                xs: x.to_vec(),
5859                q1: true,
5860            },
5861        )),
5862        _ => None,
5863    }
5864}
5865
5866thread_local! {
5867    static PRESCALE_BUF1: std::cell::RefCell<Vec<f32>> = const { std::cell::RefCell::new(Vec::new()) };
5868    static PRESCALE_BUF2: std::cell::RefCell<Vec<f32>> = const { std::cell::RefCell::new(Vec::new()) };
5869}
5870
5871pub(crate) fn prescale<'a>(
5872    x: &'a [f32],
5873    col_field: &[f32],
5874    dtype: TensorDtype,
5875) -> std::borrow::Cow<'a, [f32]> {
5876    if dtype == TensorDtype::Q8_2f {
5877        x.iter().zip(col_field).map(|(a, c)| a * c).collect()
5878    } else {
5879        std::borrow::Cow::Borrowed(x)
5880    }
5881}
5882
5883/// θ col-field fold for q8_2f activations. Borrowed pass-through for
5884/// every other dtype, using thread-local buffers to eliminate per-matvec allocations.
5885pub(crate) fn prescale_with<R, F: FnOnce(&[f32]) -> R>(
5886    x: &[f32],
5887    col_field: &[f32],
5888    dtype: TensorDtype,
5889    buf_id: u8,
5890    f: F,
5891) -> R {
5892    if dtype == TensorDtype::Q8_2f {
5893        if buf_id == 1 {
5894            PRESCALE_BUF1.with(|b| {
5895                let mut buf = b.borrow_mut();
5896                buf.clear();
5897                buf.extend(x.iter().zip(col_field).map(|(a, c)| a * c));
5898                f(&buf)
5899            })
5900        } else {
5901            PRESCALE_BUF2.with(|b| {
5902                let mut buf = b.borrow_mut();
5903                buf.clear();
5904                buf.extend(x.iter().zip(col_field).map(|(a, c)| a * c));
5905                f(&buf)
5906            })
5907        }
5908    } else {
5909        f(x)
5910    }
5911}
5912
5913// ───────────────────── x86-64 AVX2 kernels (roadmap этап 2) ─────────────────────
5914
5915/// AVX2+FMA available? Default ON when the CPU supports both;
5916/// `CMF_AVX2=0` disables (falls back to the autovectorized loops).
5917#[cfg(target_arch = "x86_64")]
5918pub(crate) fn avx2_enabled() -> bool {
5919    use std::sync::OnceLock;
5920    static ON: OnceLock<bool> = OnceLock::new();
5921    *ON.get_or_init(|| {
5922        std::env::var("CMF_AVX2").map(|v| v != "0").unwrap_or(true)
5923            && std::arch::is_x86_feature_detected!("avx2")
5924            && std::arch::is_x86_feature_detected!("fma")
5925    })
5926}
5927
5928/// AVX2 A8W8 allowed? The quantized-activation contract is switched by
5929/// the SAME env as the ARM SDOT path: `CMF_SDOT=0` keeps exact kernels
5930/// (the golden-parity exact gate relies on it) — AVX2 f32 kernels stay
5931/// active either way, they are exact (regrouped sums only).
5932#[cfg(target_arch = "x86_64")]
5933fn avx2_a8w8_enabled() -> bool {
5934    use std::sync::OnceLock;
5935    static ON: OnceLock<bool> = OnceLock::new();
5936    *ON.get_or_init(|| {
5937        avx2_enabled() && std::env::var("CMF_SDOT").map(|v| v != "0").unwrap_or(true)
5938    })
5939}
5940
5941/// A8W8 quantized-activation path available on THIS machine? One
5942/// switch across architectures: ARM dotprod (CMF_SDOT) or x86 AVX2
5943/// (CMF_AVX2 + the same CMF_SDOT exact-contract override).
5944#[inline]
5945pub(crate) fn a8w8_enabled() -> bool {
5946    #[cfg(target_arch = "aarch64")]
5947    {
5948        sdot_enabled()
5949    }
5950    #[cfg(target_arch = "x86_64")]
5951    {
5952        avx2_a8w8_enabled()
5953    }
5954    #[cfg(not(any(target_arch = "aarch64", target_arch = "x86_64")))]
5955    {
5956        false
5957    }
5958}
5959
5960/// int8·int8 dot dispatch: SDOT on ARM; AVX-512 VNNI (vpdpbusd) or AVX2
5961/// maddubs on x86. Callers are gated by `a8w8_enabled()`.
5962#[inline]
5963#[allow(unreachable_code)]
5964fn dot_i8_i8(w: &[u8], xq: &[i8]) -> i32 {
5965    #[cfg(target_arch = "aarch64")]
5966    unsafe {
5967        return dot_i8_sdot(w, xq);
5968    }
5969    #[cfg(target_arch = "x86_64")]
5970    unsafe {
5971        if avx512vnni_enabled() {
5972            return dot_i8_i8_vnni(w, xq);
5973        }
5974        return dot_i8_i8_avx2(w, xq);
5975    }
5976    w.iter()
5977        .zip(xq)
5978        .map(|(&a, &b)| (a as i8) as i32 * b as i32)
5979        .sum()
5980}
5981
5982/// AVX-512 VNNI available? (F+BW+VL+VNNI; `CMF_AVX512=0` falls back to
5983/// AVX2.) VL matters: short 32-byte groups (q4/vbit) ride the 256-bit
5984/// `vpdpbusd` encoding.
5985#[cfg(target_arch = "x86_64")]
5986fn avx512vnni_enabled() -> bool {
5987    use std::sync::OnceLock;
5988    static ON: OnceLock<bool> = OnceLock::new();
5989    *ON.get_or_init(|| {
5990        std::env::var("CMF_AVX512")
5991            .map(|v| v != "0")
5992            .unwrap_or(true)
5993            && std::arch::is_x86_feature_detected!("avx512f")
5994            && std::arch::is_x86_feature_detected!("avx512bw")
5995            && std::arch::is_x86_feature_detected!("avx512vl")
5996            && std::arch::is_x86_feature_detected!("avx512vnni")
5997    })
5998}
5999
6000/// Grouped-codec VNNI arms (the q4t/q4b/q1/q1t tile kernels): default
6001/// ON where AVX-512 VNNI exists (`CMF_VNNI_TILES=0` opt-out). Measured
6002/// on Ryzen 7950X (Zen4, 3 alternating process pairs, blocked GEMM
6003/// 4864×896 b=256): q4t 63→68 GF/s (+8%), q1 53→56 (+6%), q4b 72→75
6004/// (+4%) — consistent, no leg regressed. The tile kernels keep a
6005/// horizontal reduce per 32-weight group, so the `vpdpbusd` saving is
6006/// smaller than the long-dot q8 win (+13%), but it is real and free.
6007#[cfg(target_arch = "x86_64")]
6008fn vnni_tiles_enabled() -> bool {
6009    use std::sync::OnceLock;
6010    static ON: OnceLock<bool> = OnceLock::new();
6011    *ON.get_or_init(|| {
6012        std::env::var("CMF_VNNI_TILES")
6013            .map(|v| v != "0")
6014            .unwrap_or(true)
6015            && avx512vnni_enabled()
6016    })
6017}
6018
6019/// One 256-bit u8×i8 dot → i32 via `vpdpbusd` into a fresh accumulator
6020/// plus the same horizontal reduce the AVX2 kernels use. Products are
6021/// bounded (|w| ≤ 8 or ≤ 1), so maddubs never saturated — the i32 sum
6022/// is bit-identical to the maddubs+madd pair it replaces.
6023#[cfg(target_arch = "x86_64")]
6024#[target_feature(enable = "avx2,avx512f,avx512bw,avx512vl,avx512vnni")]
6025#[inline]
6026unsafe fn dpbusd_hsum(aw: core::arch::x86_64::__m256i, xs: core::arch::x86_64::__m256i) -> i32 {
6027    // SAFETY: pure register math.
6028    unsafe {
6029        use core::arch::x86_64::*;
6030        let d = _mm256_dpbusd_epi32(_mm256_setzero_si256(), aw, xs);
6031        let hi128 = _mm256_extracti128_si256::<1>(d);
6032        let s128 = _mm_add_epi32(_mm256_castsi256_si128(d), hi128);
6033        let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
6034        let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
6035        _mm_cvtsi128_si32(s32)
6036    }
6037}
6038
6039/// int8·int8 via AVX-512 VNNI: `vpdpbusd` fuses the maddubs+madd+add
6040/// triple into one u8×i8 dot-accumulate. AVX-512 has no vpsignb, so the
6041/// |w|·sign(x,w) trick becomes |w| × (x negated where w<0) via a mask
6042/// subtract — w==0 lanes contribute 0 through |w|=0 either way.
6043#[cfg(target_arch = "x86_64")]
6044#[target_feature(enable = "avx2,avx512f,avx512bw,avx512vl,avx512vnni")]
6045unsafe fn dot_i8_i8_vnni(w: &[u8], xq: &[i8]) -> i32 {
6046    // SAFETY: callers uphold slice-length contracts (see call sites).
6047    unsafe {
6048        use core::arch::x86_64::*;
6049        let n = w.len();
6050        let mut j = 0usize;
6051        let mut total: i32;
6052        // 4 independent accumulators: vpdpbusd is its own loop-carried
6053        // dependency (~5-cycle latency) — a single-acc loop runs
6054        // latency-bound and LOSES to the AVX2 maddubs kernel, measured
6055        // on Granite Rapids.
6056        {
6057            #[inline(always)]
6058            unsafe fn step(
6059                w: *const u8,
6060                x: *const i8,
6061                acc: core::arch::x86_64::__m512i,
6062            ) -> core::arch::x86_64::__m512i {
6063                unsafe {
6064                    use core::arch::x86_64::*;
6065                    let wv = _mm512_loadu_si512(w as *const _);
6066                    let xv = _mm512_loadu_si512(x as *const _);
6067                    let aw = _mm512_abs_epi8(wv);
6068                    let neg = _mm512_movepi8_mask(wv);
6069                    let sx = _mm512_mask_sub_epi8(xv, neg, _mm512_setzero_si512(), xv);
6070                    _mm512_dpbusd_epi32(acc, aw, sx)
6071                }
6072            }
6073            let (mut a0, mut a1, mut a2, mut a3) = (
6074                _mm512_setzero_si512(),
6075                _mm512_setzero_si512(),
6076                _mm512_setzero_si512(),
6077                _mm512_setzero_si512(),
6078            );
6079            while j + 256 <= n {
6080                a0 = step(w.as_ptr().add(j), xq.as_ptr().add(j), a0);
6081                a1 = step(w.as_ptr().add(j + 64), xq.as_ptr().add(j + 64), a1);
6082                a2 = step(w.as_ptr().add(j + 128), xq.as_ptr().add(j + 128), a2);
6083                a3 = step(w.as_ptr().add(j + 192), xq.as_ptr().add(j + 192), a3);
6084                j += 256;
6085            }
6086            while j + 64 <= n {
6087                a0 = step(w.as_ptr().add(j), xq.as_ptr().add(j), a0);
6088                j += 64;
6089            }
6090            let s01 = _mm512_add_epi32(a0, a1);
6091            let s23 = _mm512_add_epi32(a2, a3);
6092            total = _mm512_reduce_add_epi32(_mm512_add_epi32(s01, s23));
6093        }
6094        // 32-wide (q4/vbit groups are exactly 32 bytes).
6095        if j + 32 <= n {
6096            let wv = _mm256_loadu_si256(w.as_ptr().add(j) as *const __m256i);
6097            let xv = _mm256_loadu_si256(xq.as_ptr().add(j) as *const __m256i);
6098            let d = _mm256_dpbusd_epi32(
6099                _mm256_setzero_si256(),
6100                _mm256_abs_epi8(wv),
6101                _mm256_sign_epi8(xv, wv),
6102            );
6103            let hi128 = _mm256_extracti128_si256::<1>(d);
6104            let s128 = _mm_add_epi32(_mm256_castsi256_si128(d), hi128);
6105            let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
6106            let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
6107            total += _mm_cvtsi128_si32(s32);
6108            j += 32;
6109        }
6110        while j < n {
6111            total += (w[j] as i8) as i32 * xq[j] as i32;
6112            j += 1;
6113        }
6114        total
6115    }
6116}
6117
6118/// i8 row · f32 x via AVX2/FMA (x86 mirror of `dot_i8_f32_neon`).
6119#[cfg(target_arch = "x86_64")]
6120#[target_feature(enable = "avx2,fma")]
6121unsafe fn dot_i8_f32_avx2(w: &[u8], x: &[f32]) -> f32 {
6122    // SAFETY: callers uphold slice-length contracts (see call sites).
6123    unsafe {
6124        use core::arch::x86_64::*;
6125        let n = x.len();
6126        let wp = w.as_ptr();
6127        let xp = x.as_ptr();
6128        let (mut a0, mut a1) = (_mm256_setzero_ps(), _mm256_setzero_ps());
6129        let mut j = 0usize;
6130        while j + 16 <= n {
6131            let wb = _mm_loadu_si128(wp.add(j) as *const __m128i);
6132            let lo = _mm256_cvtepi8_epi32(wb);
6133            let hi = _mm256_cvtepi8_epi32(_mm_srli_si128::<8>(wb));
6134            a0 = _mm256_fmadd_ps(_mm256_cvtepi32_ps(lo), _mm256_loadu_ps(xp.add(j)), a0);
6135            a1 = _mm256_fmadd_ps(_mm256_cvtepi32_ps(hi), _mm256_loadu_ps(xp.add(j + 8)), a1);
6136            j += 16;
6137        }
6138        let acc = _mm256_add_ps(a0, a1);
6139        let hi128 = _mm256_extractf128_ps::<1>(acc);
6140        let s128 = _mm_add_ps(_mm256_castps256_ps128(acc), hi128);
6141        let s64 = _mm_add_ps(s128, _mm_movehl_ps(s128, s128));
6142        let s32 = _mm_add_ss(s64, _mm_shuffle_ps::<1>(s64, s64));
6143        let mut sum = _mm_cvtss_f32(s32);
6144        while j < n {
6145            sum += (*wp.add(j) as i8) as f32 * *xp.add(j);
6146            j += 1;
6147        }
6148        sum
6149    }
6150}
6151
6152/// int8(weight)·int8(activation) → i32 via AVX2 maddubs — the x86
6153/// analogue of the SDOT A8W8 path. `maddubs` takes u8×i8, so the
6154/// standard sign trick applies: |w| × sign(x, w) ≡ w × x per lane.
6155/// Pair saturation is safe: |w|≤128, |x|≤127 → 2·128·127 < 32767.
6156#[cfg(target_arch = "x86_64")]
6157#[target_feature(enable = "avx2")]
6158unsafe fn dot_i8_i8_avx2(w: &[u8], xq: &[i8]) -> i32 {
6159    // SAFETY: callers uphold slice-length contracts (see call sites).
6160    unsafe {
6161        use core::arch::x86_64::*;
6162        let n = w.len();
6163        let ones = _mm256_set1_epi16(1);
6164        let mut acc = _mm256_setzero_si256();
6165        let mut j = 0usize;
6166        while j + 32 <= n {
6167            let wv = _mm256_loadu_si256(w.as_ptr().add(j) as *const __m256i);
6168            let xv = _mm256_loadu_si256(xq.as_ptr().add(j) as *const __m256i);
6169            let p16 = _mm256_maddubs_epi16(_mm256_abs_epi8(wv), _mm256_sign_epi8(xv, wv));
6170            acc = _mm256_add_epi32(acc, _mm256_madd_epi16(p16, ones));
6171            j += 32;
6172        }
6173        let hi128 = _mm256_extracti128_si256::<1>(acc);
6174        let s128 = _mm_add_epi32(_mm256_castsi256_si128(acc), hi128);
6175        let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
6176        let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
6177        let mut s = _mm_cvtsi128_si32(s32);
6178        while j < n {
6179            s += (w[j] as i8) as i32 * xq[j] as i32;
6180            j += 1;
6181        }
6182        s
6183    }
6184}
6185
6186/// smmla 2×4: one instruction covers a 2-row × 2-activation × 8-deep
6187/// tile (32 MACs vs sdot's 16) — the weight pair loads once per 8-k
6188/// slice as a combined 2×8 register and meets two activation pairs.
6189#[cfg(target_arch = "aarch64")]
6190#[target_feature(enable = "neon,i8mm")]
6191unsafe fn dot_i8_smmla_2x4(w0: &[u8], w1: &[u8], xs: [&[i8]; 4]) -> [[i32; 4]; 2] {
6192    // SAFETY: callers uphold slice-length contracts.
6193    unsafe {
6194        use core::arch::aarch64::*;
6195        use core::arch::asm;
6196        let n = w0.len();
6197        let w0p = w0.as_ptr() as *const i8;
6198        let w1p = w1.as_ptr() as *const i8;
6199        // acc01 holds [c(r0,x0) c(r0,x1) c(r1,x0) c(r1,x1)]; acc23 the
6200        // same for x2/x3.
6201        let mut acc01 = vdupq_n_s32(0);
6202        let mut acc23 = vdupq_n_s32(0);
6203        let mut i = 0usize;
6204        while i + 8 <= n {
6205            let wa = vcombine_s8(vld1_s8(w0p.add(i)), vld1_s8(w1p.add(i)));
6206            let xb01 = vcombine_s8(
6207                vld1_s8(xs[0].as_ptr().add(i)),
6208                vld1_s8(xs[1].as_ptr().add(i)),
6209            );
6210            let xb23 = vcombine_s8(
6211                vld1_s8(xs[2].as_ptr().add(i)),
6212                vld1_s8(xs[3].as_ptr().add(i)),
6213            );
6214            asm!(
6215                "smmla {a01:v}.4s, {w:v}.16b, {x01:v}.16b",
6216                "smmla {a23:v}.4s, {w:v}.16b, {x23:v}.16b",
6217                a01 = inout(vreg) acc01, a23 = inout(vreg) acc23,
6218                w = in(vreg) wa, x01 = in(vreg) xb01, x23 = in(vreg) xb23,
6219                options(pure, nomem, nostack),
6220            );
6221            i += 8;
6222        }
6223        let mut out = [[0i32; 4]; 2];
6224        let a01: [i32; 4] = core::mem::transmute(acc01);
6225        let a23: [i32; 4] = core::mem::transmute(acc23);
6226        out[0][0] = a01[0];
6227        out[0][1] = a01[1];
6228        out[1][0] = a01[2];
6229        out[1][1] = a01[3];
6230        out[0][2] = a23[0];
6231        out[0][3] = a23[1];
6232        out[1][2] = a23[2];
6233        out[1][3] = a23[3];
6234        if i < n {
6235            for (k, x) in xs.iter().enumerate() {
6236                for j in i..n {
6237                    out[0][k] += (w0[j] as i8) as i32 * x[j] as i32;
6238                    out[1][k] += (w1[j] as i8) as i32 * x[j] as i32;
6239                }
6240            }
6241        }
6242        out
6243    }
6244}
6245
6246/// ARM twin of the x86 blocked prefill GEMM: two weight rows stay in
6247/// registers across four activation streams, eight sdot accumulators.
6248/// (The per-row form re-read each W row once per activation.)
6249#[cfg(target_arch = "aarch64")]
6250#[target_feature(enable = "neon,dotprod")]
6251unsafe fn dot_i8_sdot_2x4(w0: &[u8], w1: &[u8], xs: [&[i8]; 4]) -> [[i32; 4]; 2] {
6252    // SAFETY: callers uphold slice-length contracts.
6253    unsafe {
6254        use core::arch::aarch64::*;
6255        use core::arch::asm;
6256        let n = w0.len();
6257        let w0p = w0.as_ptr() as *const i8;
6258        let w1p = w1.as_ptr() as *const i8;
6259        let mut acc = [[vdupq_n_s32(0); 4]; 2];
6260        let mut i = 0usize;
6261        while i + 16 <= n {
6262            let wv0 = vld1q_s8(w0p.add(i));
6263            let wv1 = vld1q_s8(w1p.add(i));
6264            for (k, x) in xs.iter().enumerate() {
6265                let xv = vld1q_s8(x.as_ptr().add(i));
6266                let (mut a0, mut a1) = (acc[0][k], acc[1][k]);
6267                asm!(
6268                    "sdot {a0:v}.4s, {w0:v}.16b, {x:v}.16b",
6269                    "sdot {a1:v}.4s, {w1:v}.16b, {x:v}.16b",
6270                    a0 = inout(vreg) a0, a1 = inout(vreg) a1,
6271                    w0 = in(vreg) wv0, w1 = in(vreg) wv1, x = in(vreg) xv,
6272                    options(pure, nomem, nostack),
6273                );
6274                acc[0][k] = a0;
6275                acc[1][k] = a1;
6276            }
6277            i += 16;
6278        }
6279        let mut out = [[0i32; 4]; 2];
6280        for r in 0..2 {
6281            for k in 0..4 {
6282                out[r][k] = vaddvq_s32(acc[r][k]);
6283            }
6284        }
6285        if i < n {
6286            for (k, x) in xs.iter().enumerate() {
6287                for j in i..n {
6288                    out[0][k] += (w0[j] as i8) as i32 * x[j] as i32;
6289                    out[1][k] += (w1[j] as i8) as i32 * x[j] as i32;
6290                }
6291            }
6292        }
6293        out
6294    }
6295}
6296
6297/// Blocked 2 weight rows × 4 activations for the prefill GEMM
6298/// (roadmap P0: packed panels + multi-row accumulators). The two rows'
6299/// abs() live in registers across all four activation streams; the
6300/// sign-fixup is recomputed per pair (the price of the maddubs trick).
6301/// Returns raw i8·i8 dots; the caller applies scales and outliers.
6302#[cfg(target_arch = "x86_64")]
6303#[target_feature(enable = "avx2")]
6304unsafe fn dot_i8_i8_avx2_2x4(w0: &[u8], w1: &[u8], xs: [&[i8]; 4]) -> [[i32; 4]; 2] {
6305    // SAFETY: callers uphold slice-length contracts.
6306    unsafe {
6307        use core::arch::x86_64::*;
6308        let n = w0.len();
6309        let ones = _mm256_set1_epi16(1);
6310        let mut acc = [[_mm256_setzero_si256(); 4]; 2];
6311        let mut j = 0usize;
6312        while j + 32 <= n {
6313            let wv0 = _mm256_loadu_si256(w0.as_ptr().add(j) as *const __m256i);
6314            let wv1 = _mm256_loadu_si256(w1.as_ptr().add(j) as *const __m256i);
6315            let aw0 = _mm256_abs_epi8(wv0);
6316            let aw1 = _mm256_abs_epi8(wv1);
6317            for (k, x) in xs.iter().enumerate() {
6318                let xv = _mm256_loadu_si256(x.as_ptr().add(j) as *const __m256i);
6319                let p0 = _mm256_maddubs_epi16(aw0, _mm256_sign_epi8(xv, wv0));
6320                acc[0][k] = _mm256_add_epi32(acc[0][k], _mm256_madd_epi16(p0, ones));
6321                let p1 = _mm256_maddubs_epi16(aw1, _mm256_sign_epi8(xv, wv1));
6322                acc[1][k] = _mm256_add_epi32(acc[1][k], _mm256_madd_epi16(p1, ones));
6323            }
6324            j += 32;
6325        }
6326        let mut out = [[0i32; 4]; 2];
6327        for r in 0..2 {
6328            for k in 0..4 {
6329                let a = acc[r][k];
6330                let hi128 = _mm256_extracti128_si256::<1>(a);
6331                let s128 = _mm_add_epi32(_mm256_castsi256_si128(a), hi128);
6332                let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
6333                let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
6334                out[r][k] = _mm_cvtsi128_si32(s32);
6335            }
6336        }
6337        if j < n {
6338            for (k, x) in xs.iter().enumerate() {
6339                for i in j..n {
6340                    out[0][k] += (w0[i] as i8) as i32 * x[i] as i32;
6341                    out[1][k] += (w1[i] as i8) as i32 * x[i] as i32;
6342                }
6343            }
6344        }
6345        out
6346    }
6347}
6348
6349/// AVX2/VNNI q8 row dot with exact outlier correction (x86 mirror of
6350/// `row_dot_sdot` — same A8W8 contract). With AVX-512 VNNI the row goes
6351/// through the bias trick: Σ(w+128)·x via pure `vpdpbusd` (no per-lane
6352/// sign fixups), corrected by −128·Σx with Σx precomputed per split.
6353#[cfg(target_arch = "x86_64")]
6354#[inline]
6355fn row_dot_avx2(row: &[u8], act: &SplitAct) -> f32 {
6356    let dot = if avx512vnni_enabled() && row.len() >= 64 {
6357        (unsafe { dot_u8p128_i8_vnni(row, &act.xq) }) - 128 * act.xsum
6358    } else {
6359        unsafe { dot_i8_i8_avx2(row, &act.xq) }
6360    };
6361    let mut acc = dot as f32 * act.sx;
6362    for &(j, xv) in &act.outliers {
6363        acc += (row[j] as i8) as f32 * xv;
6364    }
6365    acc
6366}
6367
6368/// Σ (w[i]+128)·x[i] via pure `vpdpbusd` — the caller subtracts
6369/// 128·Σx. Four independent accumulators (dpbusd is ~5-cycle latency;
6370/// a single-acc loop runs latency-bound, measured on Granite Rapids).
6371#[cfg(target_arch = "x86_64")]
6372#[target_feature(enable = "avx2,avx512f,avx512bw,avx512vl,avx512vnni")]
6373unsafe fn dot_u8p128_i8_vnni(w: &[u8], xq: &[i8]) -> i32 {
6374    // SAFETY: callers uphold slice-length contracts (see call sites).
6375    unsafe {
6376        use core::arch::x86_64::*;
6377        let n = w.len();
6378        let flip = _mm512_set1_epi8(-128); // XOR 0x80: i8 w → u8 (w+128)
6379        #[inline(always)]
6380        unsafe fn step(
6381            w: *const u8,
6382            x: *const i8,
6383            flip: core::arch::x86_64::__m512i,
6384            acc: core::arch::x86_64::__m512i,
6385        ) -> core::arch::x86_64::__m512i {
6386            unsafe {
6387                use core::arch::x86_64::*;
6388                let wv = _mm512_xor_si512(_mm512_loadu_si512(w as *const _), flip);
6389                _mm512_dpbusd_epi32(acc, wv, _mm512_loadu_si512(x as *const _))
6390            }
6391        }
6392        let (mut a0, mut a1, mut a2, mut a3) = (
6393            _mm512_setzero_si512(),
6394            _mm512_setzero_si512(),
6395            _mm512_setzero_si512(),
6396            _mm512_setzero_si512(),
6397        );
6398        let mut j = 0usize;
6399        while j + 256 <= n {
6400            a0 = step(w.as_ptr().add(j), xq.as_ptr().add(j), flip, a0);
6401            a1 = step(w.as_ptr().add(j + 64), xq.as_ptr().add(j + 64), flip, a1);
6402            a2 = step(w.as_ptr().add(j + 128), xq.as_ptr().add(j + 128), flip, a2);
6403            a3 = step(w.as_ptr().add(j + 192), xq.as_ptr().add(j + 192), flip, a3);
6404            j += 256;
6405        }
6406        while j + 64 <= n {
6407            a0 = step(w.as_ptr().add(j), xq.as_ptr().add(j), flip, a0);
6408            j += 64;
6409        }
6410        let mut total = _mm512_reduce_add_epi32(_mm512_add_epi32(
6411            _mm512_add_epi32(a0, a1),
6412            _mm512_add_epi32(a2, a3),
6413        ));
6414        // Scalar tail: (w as i8) + 128 ≡ (w as u8) ^ 0x80.
6415        while j < n {
6416            total += ((w[j] ^ 0x80) as i32) * xq[j] as i32;
6417            j += 1;
6418        }
6419        total
6420    }
6421}
6422
6423/// One q4 row via AVX2: nibbles → centered i8 (unpacklo/hi restores the
6424/// writer's flat order, same as the NEON vzip pair), maddubs against
6425/// the pre-quantized activation group, × the group's f16 scale. Pair
6426/// saturation safe: |w|≤8, |x|≤127 → 2·8·127 ≪ 32767. Mirror of
6427/// `dot_q4_row_sdot`.
6428#[cfg(target_arch = "x86_64")]
6429#[target_feature(enable = "avx2")]
6430unsafe fn dot_q4_row_avx2(packed: &[u8], scales: &[u8], g0: usize, gpr: usize, xq: &[i8]) -> f32 {
6431    // SAFETY: callers uphold slice-length contracts (16 packed bytes and
6432    // 2 scale bytes per group; xq.len() == gpr·GROUP_SIZE).
6433    unsafe {
6434        use core::arch::x86_64::*;
6435        let lomask = _mm_set1_epi8(0x0F);
6436        let eight = _mm256_set1_epi8(8);
6437        let ones = _mm256_set1_epi16(1);
6438        let mut acc = 0f32;
6439        for gi in 0..gpr {
6440            let g = g0 + gi;
6441            let s = f16_to_f32(u16::from_le_bytes([scales[g * 2], scales[g * 2 + 1]]));
6442            let b = _mm_loadu_si128(packed.as_ptr().add(g * 16) as *const __m128i);
6443            let lo = _mm_and_si128(b, lomask);
6444            let hi = _mm_and_si128(_mm_srli_epi16::<4>(b), lomask);
6445            let w = _mm256_sub_epi8(
6446                _mm256_set_m128i(_mm_unpackhi_epi8(lo, hi), _mm_unpacklo_epi8(lo, hi)),
6447                eight,
6448            );
6449            let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
6450            let p16 = _mm256_maddubs_epi16(_mm256_abs_epi8(w), _mm256_sign_epi8(x, w));
6451            let d = _mm256_madd_epi16(p16, ones);
6452            let hi128 = _mm256_extracti128_si256::<1>(d);
6453            let s128 = _mm_add_epi32(_mm256_castsi256_si128(d), hi128);
6454            let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
6455            let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
6456            acc += _mm_cvtsi128_si32(s32) as f32 * s;
6457        }
6458        acc
6459    }
6460}
6461
6462/// Two-activation q4 row via AVX2: nibbles unpacked ONCE per group,
6463/// both activations dotted against the same centered i8 register.
6464#[cfg(target_arch = "x86_64")]
6465#[target_feature(enable = "avx2")]
6466unsafe fn dot_q4_row_avx2_2(
6467    packed: &[u8],
6468    scales: &[u8],
6469    g0: usize,
6470    gpr: usize,
6471    xq1: &[i8],
6472    xq2: &[i8],
6473) -> (f32, f32) {
6474    // SAFETY: callers uphold slice-length contracts (see dot_q4_row_avx2).
6475    unsafe {
6476        use core::arch::x86_64::*;
6477        let lomask = _mm_set1_epi8(0x0F);
6478        let eight = _mm256_set1_epi8(8);
6479        let ones = _mm256_set1_epi16(1);
6480        let (mut acc1, mut acc2) = (0f32, 0f32);
6481        #[inline(always)]
6482        unsafe fn hsum(d: core::arch::x86_64::__m256i) -> i32 {
6483            unsafe {
6484                use core::arch::x86_64::*;
6485                let hi128 = _mm256_extracti128_si256::<1>(d);
6486                let s128 = _mm_add_epi32(_mm256_castsi256_si128(d), hi128);
6487                let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
6488                let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
6489                _mm_cvtsi128_si32(s32)
6490            }
6491        }
6492        for gi in 0..gpr {
6493            let g = g0 + gi;
6494            let s = f16_to_f32(u16::from_le_bytes([scales[g * 2], scales[g * 2 + 1]]));
6495            let b = _mm_loadu_si128(packed.as_ptr().add(g * 16) as *const __m128i);
6496            let lo = _mm_and_si128(b, lomask);
6497            let hi = _mm_and_si128(_mm_srli_epi16::<4>(b), lomask);
6498            let w = _mm256_sub_epi8(
6499                _mm256_set_m128i(_mm_unpackhi_epi8(lo, hi), _mm_unpacklo_epi8(lo, hi)),
6500                eight,
6501            );
6502            let aw = _mm256_abs_epi8(w);
6503            let x1 = _mm256_loadu_si256(xq1.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
6504            let x2 = _mm256_loadu_si256(xq2.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
6505            let d1 = _mm256_madd_epi16(_mm256_maddubs_epi16(aw, _mm256_sign_epi8(x1, w)), ones);
6506            let d2 = _mm256_madd_epi16(_mm256_maddubs_epi16(aw, _mm256_sign_epi8(x2, w)), ones);
6507            acc1 += hsum(d1) as f32 * s;
6508            acc2 += hsum(d2) as f32 * s;
6509        }
6510        (acc1, acc2)
6511    }
6512}
6513
6514/// One q8 row range via AVX2 (x86 mirror of `q8_range_sdot`).
6515#[cfg(target_arch = "x86_64")]
6516fn q8_range_avx2(
6517    q: &[u8],
6518    row_scale: &[f32],
6519    act: &SplitAct,
6520    cols: usize,
6521    out_addr: SendMut,
6522    start: usize,
6523    end: usize,
6524) {
6525    for o in start..end {
6526        let v = row_dot_avx2(&q[o * cols..(o + 1) * cols], act) * row_scale[o];
6527        // SAFETY: disjoint row ranges per worker.
6528        unsafe { *out_addr.at(o) = v };
6529    }
6530}
6531
6532/// Two-input q8 row range via AVX2 (x86 mirror of `q8_range2_sdot`).
6533#[cfg(target_arch = "x86_64")]
6534#[allow(clippy::too_many_arguments)]
6535fn q8_range2_avx2(
6536    q: &[u8],
6537    row_scale: &[f32],
6538    a1: &SplitAct,
6539    a2: &SplitAct,
6540    cols: usize,
6541    p1: SendMut,
6542    p2: SendMut,
6543    start: usize,
6544    end: usize,
6545) {
6546    for o in start..end {
6547        let row = &q[o * cols..(o + 1) * cols];
6548        // SAFETY: disjoint row ranges per worker.
6549        unsafe {
6550            *p1.at(o) = row_dot_avx2(row, a1) * row_scale[o];
6551            *p2.at(o) = row_dot_avx2(row, a2) * row_scale[o];
6552        }
6553    }
6554}
6555
6556// ───────────────────── A8W8 SDOT path (port of vmfcore, ×1.78 decode) ─────────────────────
6557
6558/// ARMv8.6 i8mm (smmla): 32 int8 MACs per instruction vs sdot's 16 —
6559/// yet MEASURED 2.4× SLOWER than the blocked sdot on Apple silicon
6560/// (108 vs 264 GF/s): the on-the-fly vcombine packing and the two-
6561/// accumulator dependency chain swamp the MAC advantage, and Apple's
6562/// four SIMD pipes already keep sdot fed. OPT-IN (CMF_I8MM=1) for
6563/// field trials on Cortex-A710/X-class parts with two pipes, where the
6564/// balance may differ; a pre-interleaved weight layout (repack infra)
6565/// is the known path if it ever earns its keep.
6566#[cfg(target_arch = "aarch64")]
6567fn i8mm_enabled() -> bool {
6568    static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
6569    *ON.get_or_init(|| {
6570        std::env::var("CMF_I8MM").map(|v| v == "1").unwrap_or(false)
6571            && std::arch::is_aarch64_feature_detected!("i8mm")
6572    })
6573}
6574
6575/// SDOT enabled? Default ON when the CPU has ARMv8.2 dotprod;
6576/// `CMF_SDOT=0` disables (falls back to i8×f32 NEON).
6577/// (On non-ARM release builds only the test tolerance switch calls it.)
6578#[cfg_attr(not(target_arch = "aarch64"), allow(dead_code))]
6579fn sdot_enabled() -> bool {
6580    use std::sync::OnceLock;
6581    static ON: OnceLock<bool> = OnceLock::new();
6582    *ON.get_or_init(|| {
6583        let want = std::env::var("CMF_SDOT").map(|v| v != "0").unwrap_or(true);
6584        if !want {
6585            return false;
6586        }
6587
6588        #[cfg(target_arch = "aarch64")]
6589        {
6590            if std::arch::is_aarch64_feature_detected!("dotprod") {
6591                return true;
6592            }
6593            #[cfg(target_os = "android")]
6594            {
6595                if let Ok(cpuinfo) = std::fs::read_to_string("/proc/cpuinfo") {
6596                    if cpuinfo.lines().any(|l| {
6597                        (l.starts_with("Features") || l.starts_with("features"))
6598                            && l.contains("asimddp")
6599                    }) {
6600                        return true;
6601                    }
6602                }
6603            }
6604            false
6605        }
6606        #[cfg(not(target_arch = "aarch64"))]
6607        {
6608            false
6609        }
6610    })
6611}
6612
6613/// Two-field activation split (≡ vmfcore `q8_split_prep`): outlier
6614/// channels (>8·rms) are computed exactly in f32; the bulk (outliers
6615/// zeroed → clean absmax) goes through int8 SDOT. Computed ONCE per
6616/// matvec, shared by all rows/workers.
6617struct SplitAct {
6618    xq: Vec<i8>,
6619    sx: f32,
6620    outliers: Vec<(usize, f32)>,
6621    /// Σ xq — the VNNI bias-trick correction (`(w+128)·x` sums need
6622    /// `−128·Σx`); one i32 per split, computed once per matvec.
6623    #[cfg_attr(not(target_arch = "x86_64"), allow(dead_code))]
6624    xsum: i32,
6625}
6626
6627thread_local! {
6628    /// Recycled xq buffers: split_act runs for every matvec (~200/token)
6629    /// and its hidden-size allocation was steady-state heap churn.
6630    static XQ_FREE: std::cell::RefCell<Vec<Vec<i8>>> =
6631        const { std::cell::RefCell::new(Vec::new()) };
6632}
6633
6634impl Drop for SplitAct {
6635    fn drop(&mut self) {
6636        let buf = std::mem::take(&mut self.xq);
6637        if buf.capacity() > 0 {
6638            XQ_FREE.with(|f| {
6639                let mut f = f.borrow_mut();
6640                if f.len() < 16 {
6641                    f.push(buf);
6642                }
6643            });
6644        }
6645    }
6646}
6647
6648fn split_act(x: &[f32]) -> SplitAct {
6649    let n = x.len();
6650    let rms = (x.iter().map(|&v| (v * v) as f64).sum::<f64>() / n.max(1) as f64).sqrt() as f32;
6651    let thr = 8.0 * rms;
6652    // One pass: collect outliers and the bulk absmax (outliers excluded —
6653    // identical to the old zero-then-fold over a copied buffer, minus the
6654    // full-vector copy).
6655    let mut outliers: Vec<(usize, f32)> = Vec::new();
6656    let mut amax = 0f32;
6657    for (j, &v) in x.iter().enumerate() {
6658        let a = v.abs();
6659        if a > thr {
6660            outliers.push((j, v));
6661        } else if a > amax {
6662            amax = a;
6663        }
6664    }
6665    let sx = if amax > 0.0 { amax / 127.0 } else { 1.0 };
6666    let inv = 1.0 / sx;
6667    let mut xq = XQ_FREE.with(|f| f.borrow_mut().pop()).unwrap_or_default();
6668    xq.clear();
6669    xq.reserve(n);
6670    if outliers.is_empty() {
6671        xq.extend(
6672            x.iter()
6673                .map(|&v| (v * inv).round().clamp(-127.0, 127.0) as i8),
6674        );
6675    } else {
6676        // Outlier slots quantize to 0 (their exact term is added later).
6677        xq.extend(x.iter().map(|&v| {
6678            if v.abs() > thr {
6679                0
6680            } else {
6681                (v * inv).round().clamp(-127.0, 127.0) as i8
6682            }
6683        }));
6684    }
6685    let xsum = xq.iter().map(|&v| v as i32).sum();
6686    SplitAct {
6687        xq,
6688        sx,
6689        outliers,
6690        xsum,
6691    }
6692}
6693
6694fn split_act_q8_2f(x: &[f32], col: &[f32]) -> SplitAct {
6695    let n = x.len();
6696    let rms = (x
6697        .iter()
6698        .zip(col)
6699        .map(|(&a, &c)| {
6700            let v = a * c;
6701            (v * v) as f64
6702        })
6703        .sum::<f64>()
6704        / n.max(1) as f64)
6705        .sqrt() as f32;
6706    let thr = 8.0 * rms;
6707
6708    let mut outliers = Vec::new();
6709    let mut amax = 0f32;
6710    for (j, (&a, &c)) in x.iter().zip(col).enumerate() {
6711        let v = a * c;
6712        let s = v.abs();
6713        if s > thr {
6714            outliers.push((j, v));
6715        } else if s > amax {
6716            amax = s;
6717        }
6718    }
6719
6720    let sx = if amax > 0.0 { amax / 127.0 } else { 1.0 };
6721    let inv = 1.0 / sx;
6722    let mut xq = XQ_FREE.with(|f| f.borrow_mut().pop()).unwrap_or_default();
6723    xq.clear();
6724    xq.reserve(n);
6725    if outliers.is_empty() {
6726        xq.extend(
6727            x.iter()
6728                .zip(col)
6729                .map(|(&a, &c)| ((a * c) * inv).round().clamp(-127.0, 127.0) as i8),
6730        );
6731    } else {
6732        xq.extend(x.iter().zip(col).map(|(&a, &c)| {
6733            let v = a * c;
6734            if v.abs() > thr {
6735                0
6736            } else {
6737                (v * inv).round().clamp(-127.0, 127.0) as i8
6738            }
6739        }));
6740    }
6741    let xsum = xq.iter().map(|&v| v as i32).sum();
6742    SplitAct {
6743        xq,
6744        sx,
6745        outliers,
6746        xsum,
6747    }
6748}
6749
6750/// int8(weight)·int8(activation) → i32 via `sdot` (inline asm — the
6751/// vdotq intrinsic is unstable; port of vmfcore `dot_i8_sdot`).
6752#[cfg(target_arch = "aarch64")]
6753#[target_feature(enable = "neon,dotprod")]
6754unsafe fn dot_i8_sdot(w: &[u8], xq: &[i8]) -> i32 {
6755    // SAFETY: callers uphold slice-length contracts (see call sites).
6756    unsafe {
6757        use core::arch::aarch64::*;
6758        use core::arch::asm;
6759        let wp = w.as_ptr() as *const i8;
6760        let n = w.len();
6761        let (mut a0, mut a1, mut a2, mut a3) = (
6762            vdupq_n_s32(0),
6763            vdupq_n_s32(0),
6764            vdupq_n_s32(0),
6765            vdupq_n_s32(0),
6766        );
6767        let mut i = 0;
6768        while i + 64 <= n {
6769            let (w0, x0) = (vld1q_s8(wp.add(i)), vld1q_s8(xq.as_ptr().add(i)));
6770            let (w1, x1) = (vld1q_s8(wp.add(i + 16)), vld1q_s8(xq.as_ptr().add(i + 16)));
6771            let (w2, x2) = (vld1q_s8(wp.add(i + 32)), vld1q_s8(xq.as_ptr().add(i + 32)));
6772            let (w3, x3) = (vld1q_s8(wp.add(i + 48)), vld1q_s8(xq.as_ptr().add(i + 48)));
6773            asm!(
6774                "sdot {a0:v}.4s, {w0:v}.16b, {x0:v}.16b",
6775                "sdot {a1:v}.4s, {w1:v}.16b, {x1:v}.16b",
6776                "sdot {a2:v}.4s, {w2:v}.16b, {x2:v}.16b",
6777                "sdot {a3:v}.4s, {w3:v}.16b, {x3:v}.16b",
6778                a0 = inout(vreg) a0, a1 = inout(vreg) a1, a2 = inout(vreg) a2, a3 = inout(vreg) a3,
6779                w0 = in(vreg) w0, x0 = in(vreg) x0, w1 = in(vreg) w1, x1 = in(vreg) x1,
6780                w2 = in(vreg) w2, x2 = in(vreg) x2, w3 = in(vreg) w3, x3 = in(vreg) x3,
6781                options(pure, nomem, nostack),
6782            );
6783            i += 64;
6784        }
6785        while i + 16 <= n {
6786            let (wv, xv) = (vld1q_s8(wp.add(i)), vld1q_s8(xq.as_ptr().add(i)));
6787            asm!("sdot {a:v}.4s, {w:v}.16b, {x:v}.16b",
6788                 a = inout(vreg) a0, w = in(vreg) wv, x = in(vreg) xv, options(pure, nomem, nostack));
6789            i += 16;
6790        }
6791        let mut s = vaddvq_s32(vaddq_s32(vaddq_s32(a0, a1), vaddq_s32(a2, a3)));
6792        while i < n {
6793            s += (*wp.add(i)) as i32 * xq[i] as i32;
6794            i += 1;
6795        }
6796        s
6797    }
6798}
6799
6800/// Row-blocked SDOT: 4 output rows per pass — the activation chunk is
6801/// loaded once and reused, 4 independent accumulators hide sdot latency
6802/// (port of vmfcore `dot_i8_sdot_4rows`).
6803#[cfg(target_arch = "aarch64")]
6804#[target_feature(enable = "neon,dotprod")]
6805unsafe fn dot_i8_sdot_4rows(w0: &[u8], w1: &[u8], w2: &[u8], w3: &[u8], xq: &[i8]) -> [i32; 4] {
6806    // SAFETY: callers uphold slice-length contracts (see call sites).
6807    unsafe {
6808        use core::arch::aarch64::*;
6809        use core::arch::asm;
6810        let n = xq.len();
6811        let px = xq.as_ptr();
6812        let (p0, p1, p2, p3) = (
6813            w0.as_ptr() as *const i8,
6814            w1.as_ptr() as *const i8,
6815            w2.as_ptr() as *const i8,
6816            w3.as_ptr() as *const i8,
6817        );
6818        let (mut a0, mut a1, mut a2, mut a3) = (
6819            vdupq_n_s32(0),
6820            vdupq_n_s32(0),
6821            vdupq_n_s32(0),
6822            vdupq_n_s32(0),
6823        );
6824        let mut i = 0;
6825        while i + 16 <= n {
6826            let x = vld1q_s8(px.add(i));
6827            let v0 = vld1q_s8(p0.add(i));
6828            let v1 = vld1q_s8(p1.add(i));
6829            let v2 = vld1q_s8(p2.add(i));
6830            let v3 = vld1q_s8(p3.add(i));
6831            asm!(
6832                "sdot {a0:v}.4s, {v0:v}.16b, {x:v}.16b",
6833                "sdot {a1:v}.4s, {v1:v}.16b, {x:v}.16b",
6834                "sdot {a2:v}.4s, {v2:v}.16b, {x:v}.16b",
6835                "sdot {a3:v}.4s, {v3:v}.16b, {x:v}.16b",
6836                a0 = inout(vreg) a0, a1 = inout(vreg) a1, a2 = inout(vreg) a2, a3 = inout(vreg) a3,
6837                v0 = in(vreg) v0, v1 = in(vreg) v1, v2 = in(vreg) v2, v3 = in(vreg) v3, x = in(vreg) x,
6838                options(pure, nomem, nostack),
6839            );
6840            i += 16;
6841        }
6842        let mut r = [
6843            vaddvq_s32(a0),
6844            vaddvq_s32(a1),
6845            vaddvq_s32(a2),
6846            vaddvq_s32(a3),
6847        ];
6848        while i < n {
6849            let xi = *px.add(i) as i32;
6850            r[0] += (*p0.add(i)) as i32 * xi;
6851            r[1] += (*p1.add(i)) as i32 * xi;
6852            r[2] += (*p2.add(i)) as i32 * xi;
6853            r[3] += (*p3.add(i)) as i32 * xi;
6854            i += 1;
6855        }
6856        r
6857    }
6858}
6859
6860/// 4 interleaved rows in one pass: the repacked group is [r0[c], r1[c],
6861/// r2[c], r3[c]] per 16-byte chunk, so each iteration reads ONE 64-byte
6862/// line plus the shared activation chunk — a single sequential weight
6863/// stream per worker. Per-row accumulation is the same one-accumulator
6864/// scheme as `dot_i8_sdot_4rows`; integer sums are exact, so outputs
6865/// are bit-identical to the mmap-layout kernel.
6866#[cfg(target_arch = "aarch64")]
6867#[target_feature(enable = "neon,dotprod")]
6868unsafe fn dot_i8_sdot_4rows_il(g: &[u8], xq: &[i8]) -> [i32; 4] {
6869    // SAFETY: callers uphold slice-length contracts (g.len() == 4·n,
6870    // n % 16 == 0 — guaranteed by the repack gate).
6871    unsafe {
6872        use core::arch::aarch64::*;
6873        use core::arch::asm;
6874        let n = xq.len();
6875        let px = xq.as_ptr();
6876        let pg = g.as_ptr() as *const i8;
6877        let (mut a0, mut a1, mut a2, mut a3) = (
6878            vdupq_n_s32(0),
6879            vdupq_n_s32(0),
6880            vdupq_n_s32(0),
6881            vdupq_n_s32(0),
6882        );
6883        let mut i = 0;
6884        while i + 16 <= n {
6885            let x = vld1q_s8(px.add(i));
6886            let base = pg.add(4 * i);
6887            let v0 = vld1q_s8(base);
6888            let v1 = vld1q_s8(base.add(16));
6889            let v2 = vld1q_s8(base.add(32));
6890            let v3 = vld1q_s8(base.add(48));
6891            asm!(
6892                "sdot {a0:v}.4s, {v0:v}.16b, {x:v}.16b",
6893                "sdot {a1:v}.4s, {v1:v}.16b, {x:v}.16b",
6894                "sdot {a2:v}.4s, {v2:v}.16b, {x:v}.16b",
6895                "sdot {a3:v}.4s, {v3:v}.16b, {x:v}.16b",
6896                a0 = inout(vreg) a0, a1 = inout(vreg) a1, a2 = inout(vreg) a2, a3 = inout(vreg) a3,
6897                v0 = in(vreg) v0, v1 = in(vreg) v1, v2 = in(vreg) v2, v3 = in(vreg) v3, x = in(vreg) x,
6898                options(pure, nomem, nostack),
6899            );
6900            i += 16;
6901        }
6902        [
6903            vaddvq_s32(a0),
6904            vaddvq_s32(a1),
6905            vaddvq_s32(a2),
6906            vaddvq_s32(a3),
6907        ]
6908    }
6909}
6910
6911/// One q8 row range via SDOT (4-row blocks + tail) — the body of
6912/// `qmatvec`'s hot loop, extracted so multi-matrix jobs can drive the
6913/// SAME kernel for several tensors under one pool dispatch. `rep` — the
6914/// load-time interleaved repack (empty = mmap layout only); rows outside
6915/// full 4-row groups always come from the mmap layout.
6916#[cfg(target_arch = "aarch64")]
6917fn q8_range_sdot(
6918    q: &[u8],
6919    rep: &[u8],
6920    row_scale: &[f32],
6921    act: &SplitAct,
6922    cols: usize,
6923    out_addr: SendMut,
6924    start: usize,
6925    end: usize,
6926) {
6927    let mut o = start;
6928    // Leading rows to the group boundary (repack path only): the pool
6929    // splits row ranges arbitrarily, groups are absolute.
6930    if !rep.is_empty() {
6931        while o < end && o % 4 != 0 {
6932            let v = row_dot_sdot(&q[o * cols..(o + 1) * cols], act) * row_scale[o];
6933            unsafe { *out_addr.at(o) = v };
6934            o += 1;
6935        }
6936    }
6937    while o + 4 <= end {
6938        let r = if rep.is_empty() {
6939            unsafe {
6940                dot_i8_sdot_4rows(
6941                    &q[o * cols..(o + 1) * cols],
6942                    &q[(o + 1) * cols..(o + 2) * cols],
6943                    &q[(o + 2) * cols..(o + 3) * cols],
6944                    &q[(o + 3) * cols..(o + 4) * cols],
6945                    &act.xq,
6946                )
6947            }
6948        } else {
6949            unsafe { dot_i8_sdot_4rows_il(&rep[o * cols..(o + 4) * cols], &act.xq) }
6950        };
6951        for k in 0..4 {
6952            let mut acc = r[k] as f32 * act.sx;
6953            for &(j, xv) in &act.outliers {
6954                acc += (q[(o + k) * cols + j] as i8) as f32 * xv;
6955            }
6956            // SAFETY: disjoint row ranges per worker.
6957            unsafe { *out_addr.at(o + k) = acc * row_scale[o + k] };
6958        }
6959        o += 4;
6960    }
6961    while o < end {
6962        let v = row_dot_sdot(&q[o * cols..(o + 1) * cols], act) * row_scale[o];
6963        unsafe { *out_addr.at(o) = v };
6964        o += 1;
6965    }
6966}
6967
6968/// Two-input q8 row range via SDOT — `qmatvec2`'s hot loop, extracted
6969/// for the fused pair multi-matrix job (`matvec2_many`).
6970#[cfg(target_arch = "aarch64")]
6971#[allow(clippy::too_many_arguments)]
6972fn q8_range2_sdot(
6973    q: &[u8],
6974    row_scale: &[f32],
6975    a1: &SplitAct,
6976    a2: &SplitAct,
6977    cols: usize,
6978    p1: SendMut,
6979    p2: SendMut,
6980    start: usize,
6981    end: usize,
6982) {
6983    for o in start..end {
6984        let row = &q[o * cols..(o + 1) * cols];
6985        // SAFETY: disjoint row ranges per worker.
6986        unsafe {
6987            *p1.at(o) = row_dot_sdot(row, a1) * row_scale[o];
6988            *p2.at(o) = row_dot_sdot(row, a2) * row_scale[o];
6989        }
6990    }
6991}
6992
6993/// Two-input q8 row range, f32 kernel (non-SDOT) — same extraction.
6994#[allow(clippy::too_many_arguments)]
6995fn q8_range2_f32(
6996    q: &[u8],
6997    row_scale: &[f32],
6998    x1: &[f32],
6999    x2: &[f32],
7000    cols: usize,
7001    p1: SendMut,
7002    p2: SendMut,
7003    start: usize,
7004    end: usize,
7005) {
7006    for o in start..end {
7007        let row = &q[o * cols..(o + 1) * cols];
7008        // SAFETY: disjoint row ranges per worker.
7009        unsafe {
7010            *p1.at(o) = dot_i8_f32(row, x1) * row_scale[o];
7011            *p2.at(o) = dot_i8_f32(row, x2) * row_scale[o];
7012        }
7013    }
7014}
7015
7016/// Scalar/NEON-f32 q8 row range (non-SDOT platforms) — same extraction.
7017fn q8_range_f32(
7018    q: &[u8],
7019    row_scale: &[f32],
7020    xs: &[f32],
7021    cols: usize,
7022    out_addr: SendMut,
7023    start: usize,
7024    end: usize,
7025) {
7026    for o in start..end {
7027        let v = dot_i8_f32(&q[o * cols..(o + 1) * cols], xs) * row_scale[o];
7028        // SAFETY: disjoint row ranges per worker.
7029        unsafe { *out_addr.at(o) = v };
7030    }
7031}
7032
7033/// SDOT row dot with exact outlier correction:
7034/// `dot = sdot(w, xq)·sx + Σ_outl w[j]·x[j]` (then × row_scale by caller).
7035#[cfg(target_arch = "aarch64")]
7036#[inline]
7037fn row_dot_sdot(row: &[u8], act: &SplitAct) -> f32 {
7038    let mut acc = unsafe { dot_i8_sdot(row, &act.xq) } as f32 * act.sx;
7039    for &(j, xv) in &act.outliers {
7040        acc += (row[j] as i8) as f32 * xv;
7041    }
7042    acc
7043}
7044
7045/// One q4 row via SDOT: each 32-group's nibbles unpack to centered i8
7046/// (nib−8 ∈ [−8,7]), int8×int8 `sdot` against the pre-quantized
7047/// activation group, × the group's f16 scale. Returns Σ_g dot_g·s_g;
7048/// the caller multiplies by the activation scale and adds the exact
7049/// outlier terms (port of vmfcore `dot_q4_block_sdot`, +23% measured).
7050/// Nibble order matches the writer: element 2k = low nibble, 2k+1 = high
7051/// → zip(lo,hi) restores flat order.
7052#[cfg(target_arch = "aarch64")]
7053#[target_feature(enable = "neon,dotprod")]
7054unsafe fn dot_q4_row_sdot(packed: &[u8], scales: &[u8], g0: usize, gpr: usize, xq: &[i8]) -> f32 {
7055    // SAFETY: callers uphold slice-length contracts (16 packed bytes and
7056    // 2 scale bytes per group; xq.len() == gpr·GROUP_SIZE).
7057    unsafe {
7058        use core::arch::aarch64::*;
7059        use core::arch::asm;
7060        let lomask = vdupq_n_u8(0x0F);
7061        let eight = vdupq_n_s8(8);
7062        let mut acc = 0f32;
7063        for gi in 0..gpr {
7064            let g = g0 + gi;
7065            let s = f16_to_f32(u16::from_le_bytes([scales[g * 2], scales[g * 2 + 1]]));
7066            let b = vld1q_u8(packed.as_ptr().add(g * 16));
7067            let lo = vandq_u8(b, lomask);
7068            let hi = vshrq_n_u8::<4>(b);
7069            let e0 = vsubq_s8(vreinterpretq_s8_u8(vzip1q_u8(lo, hi)), eight);
7070            let e1 = vsubq_s8(vreinterpretq_s8_u8(vzip2q_u8(lo, hi)), eight);
7071            let x0 = vld1q_s8(xq.as_ptr().add(gi * GROUP_SIZE));
7072            let x1 = vld1q_s8(xq.as_ptr().add(gi * GROUP_SIZE + 16));
7073            let (mut a0, mut a1) = (vdupq_n_s32(0), vdupq_n_s32(0));
7074            asm!(
7075                "sdot {a0:v}.4s, {e0:v}.16b, {x0:v}.16b",
7076                "sdot {a1:v}.4s, {e1:v}.16b, {x1:v}.16b",
7077                a0 = inout(vreg) a0, a1 = inout(vreg) a1,
7078                e0 = in(vreg) e0, x0 = in(vreg) x0, e1 = in(vreg) e1, x1 = in(vreg) x1,
7079                options(pure, nomem, nostack),
7080            );
7081            acc += vaddvq_s32(vaddq_s32(a0, a1)) as f32 * s;
7082        }
7083        acc
7084    }
7085}
7086
7087/// Two-activation q4 row via SDOT: the nibble unpack (the expensive
7088/// part) happens ONCE per group; both pre-quantized activations are
7089/// dotted against the same centered i8 registers. Per-lane math matches
7090/// `dot_q4_row_sdot` exactly.
7091#[cfg(target_arch = "aarch64")]
7092#[target_feature(enable = "neon,dotprod")]
7093unsafe fn dot_q4_row_sdot2(
7094    packed: &[u8],
7095    scales: &[u8],
7096    g0: usize,
7097    gpr: usize,
7098    xq1: &[i8],
7099    xq2: &[i8],
7100) -> (f32, f32) {
7101    // SAFETY: callers uphold slice-length contracts (16 packed bytes and
7102    // 2 scale bytes per group; xq*.len() == gpr·GROUP_SIZE).
7103    unsafe {
7104        use core::arch::aarch64::*;
7105        use core::arch::asm;
7106        let lomask = vdupq_n_u8(0x0F);
7107        let eight = vdupq_n_s8(8);
7108        let (mut acc1, mut acc2) = (0f32, 0f32);
7109        for gi in 0..gpr {
7110            let g = g0 + gi;
7111            let s = f16_to_f32(u16::from_le_bytes([scales[g * 2], scales[g * 2 + 1]]));
7112            let b = vld1q_u8(packed.as_ptr().add(g * 16));
7113            let lo = vandq_u8(b, lomask);
7114            let hi = vshrq_n_u8::<4>(b);
7115            let e0 = vsubq_s8(vreinterpretq_s8_u8(vzip1q_u8(lo, hi)), eight);
7116            let e1 = vsubq_s8(vreinterpretq_s8_u8(vzip2q_u8(lo, hi)), eight);
7117            let x10 = vld1q_s8(xq1.as_ptr().add(gi * GROUP_SIZE));
7118            let x11 = vld1q_s8(xq1.as_ptr().add(gi * GROUP_SIZE + 16));
7119            let x20 = vld1q_s8(xq2.as_ptr().add(gi * GROUP_SIZE));
7120            let x21 = vld1q_s8(xq2.as_ptr().add(gi * GROUP_SIZE + 16));
7121            let (mut a0, mut a1, mut b0, mut b1) = (
7122                vdupq_n_s32(0),
7123                vdupq_n_s32(0),
7124                vdupq_n_s32(0),
7125                vdupq_n_s32(0),
7126            );
7127            asm!(
7128                "sdot {a0:v}.4s, {e0:v}.16b, {x10:v}.16b",
7129                "sdot {a1:v}.4s, {e1:v}.16b, {x11:v}.16b",
7130                "sdot {b0:v}.4s, {e0:v}.16b, {x20:v}.16b",
7131                "sdot {b1:v}.4s, {e1:v}.16b, {x21:v}.16b",
7132                a0 = inout(vreg) a0, a1 = inout(vreg) a1,
7133                b0 = inout(vreg) b0, b1 = inout(vreg) b1,
7134                e0 = in(vreg) e0, e1 = in(vreg) e1,
7135                x10 = in(vreg) x10, x11 = in(vreg) x11,
7136                x20 = in(vreg) x20, x21 = in(vreg) x21,
7137                options(pure, nomem, nostack),
7138            );
7139            acc1 += vaddvq_s32(vaddq_s32(a0, a1)) as f32 * s;
7140            acc2 += vaddvq_s32(vaddq_s32(b0, b1)) as f32 * s;
7141        }
7142        (acc1, acc2)
7143    }
7144}
7145
7146// ───────────────────── fused int8 kernels ─────────────────────
7147
7148/// `acc += w · row` where the row is centered i8 — NEON widen+fma on
7149/// aarch64, scalar elsewhere. The KV-cache q8 value path rides on this.
7150#[inline]
7151pub(crate) fn axpy_i8_f32(acc: &mut [f32], row: &[i8], w: f32) {
7152    #[cfg(target_arch = "aarch64")]
7153    unsafe {
7154        return axpy_i8_f32_neon(acc, row, w);
7155    }
7156    #[cfg(target_arch = "x86_64")]
7157    if avx2_enabled() {
7158        return unsafe { axpy_i8_f32_avx2(acc, row, w) };
7159    }
7160    #[allow(unreachable_code)]
7161    {
7162        for (a, &b) in acc.iter_mut().zip(row) {
7163            *a += w * b as f32;
7164        }
7165    }
7166}
7167
7168/// i8→f32 axpy via AVX2/FMA (x86 mirror of `axpy_i8_f32_neon`).
7169#[cfg(target_arch = "x86_64")]
7170#[target_feature(enable = "avx2,fma")]
7171unsafe fn axpy_i8_f32_avx2(acc: &mut [f32], row: &[i8], w: f32) {
7172    // SAFETY: callers uphold slice-length contracts (see call sites).
7173    unsafe {
7174        use core::arch::x86_64::*;
7175        let n = acc.len().min(row.len());
7176        let ap = acc.as_mut_ptr();
7177        let rp = row.as_ptr();
7178        let wv = _mm256_set1_ps(w);
7179        let mut j = 0usize;
7180        while j + 16 <= n {
7181            let rb = _mm_loadu_si128(rp.add(j) as *const __m128i);
7182            let lo = _mm256_cvtepi32_ps(_mm256_cvtepi8_epi32(rb));
7183            let hi = _mm256_cvtepi32_ps(_mm256_cvtepi8_epi32(_mm_srli_si128::<8>(rb)));
7184            let v0 = _mm256_fmadd_ps(wv, lo, _mm256_loadu_ps(ap.add(j)));
7185            let v1 = _mm256_fmadd_ps(wv, hi, _mm256_loadu_ps(ap.add(j + 8)));
7186            _mm256_storeu_ps(ap.add(j), v0);
7187            _mm256_storeu_ps(ap.add(j + 8), v1);
7188            j += 16;
7189        }
7190        while j < n {
7191            *ap.add(j) += w * (*rp.add(j)) as f32;
7192            j += 1;
7193        }
7194    }
7195}
7196
7197#[cfg(target_arch = "aarch64")]
7198#[target_feature(enable = "neon")]
7199unsafe fn axpy_i8_f32_neon(acc: &mut [f32], row: &[i8], w: f32) {
7200    // SAFETY: callers uphold slice-length contracts (see call sites).
7201    unsafe {
7202        use core::arch::aarch64::*;
7203        let n = acc.len().min(row.len());
7204        let ap = acc.as_mut_ptr();
7205        let rp = row.as_ptr();
7206        let wv = vdupq_n_f32(w);
7207        let mut j = 0usize;
7208        while j + 16 <= n {
7209            let rb = vld1q_s8(rp.add(j));
7210            let lo = vmovl_s8(vget_low_s8(rb));
7211            let hi = vmovl_s8(vget_high_s8(rb));
7212            for (off, half) in [(0, lo), (8, hi)] {
7213                let f0 = vcvtq_f32_s32(vmovl_s16(vget_low_s16(half)));
7214                let f1 = vcvtq_f32_s32(vmovl_s16(vget_high_s16(half)));
7215                let o = j + off;
7216                vst1q_f32(ap.add(o), vfmaq_f32(vld1q_f32(ap.add(o)), wv, f0));
7217                vst1q_f32(ap.add(o + 4), vfmaq_f32(vld1q_f32(ap.add(o + 4)), wv, f1));
7218            }
7219            j += 16;
7220        }
7221        while j < n {
7222            *ap.add(j) += w * (*rp.add(j)) as f32;
7223            j += 1;
7224        }
7225    }
7226}
7227
7228/// i8 row · f32 x. NEON on aarch64 (ported from vmfcore `dot_i8_f32_neon`,
7229/// ≈9× scalar), scalar elsewhere.
7230#[inline]
7231pub(crate) fn dot_i8_f32(w: &[u8], x: &[f32]) -> f32 {
7232    #[cfg(target_arch = "aarch64")]
7233    unsafe {
7234        return dot_i8_f32_neon(w, x);
7235    }
7236    #[cfg(target_arch = "x86_64")]
7237    if avx2_enabled() {
7238        return unsafe { dot_i8_f32_avx2(w, x) };
7239    }
7240    #[allow(unreachable_code)]
7241    {
7242        let mut sum = 0.0f32;
7243        for (j, &b) in w.iter().enumerate() {
7244            sum += (b as i8) as f32 * x[j];
7245        }
7246        sum
7247    }
7248}
7249
7250/// i8 row · (x ⊙ col_field) — the q8_2f row dot with the θ col-field
7251/// folded into the product (no prescaled copy of x). NEON on aarch64,
7252/// scalar elsewhere. Used by the active-neuron path `row_dot`.
7253#[inline]
7254fn dot_i8_col_f32(w: &[u8], x: &[f32], col: &[f32]) -> f32 {
7255    #[cfg(target_arch = "aarch64")]
7256    unsafe {
7257        return dot_i8_col_f32_neon(w, x, col);
7258    }
7259    #[allow(unreachable_code)]
7260    {
7261        let mut sum = 0.0f32;
7262        for (j, &b) in w.iter().enumerate() {
7263            sum += (b as i8) as f32 * x[j] * col[j];
7264        }
7265        sum
7266    }
7267}
7268
7269#[cfg(target_arch = "aarch64")]
7270#[target_feature(enable = "neon")]
7271unsafe fn dot_i8_col_f32_neon(w: &[u8], x: &[f32], col: &[f32]) -> f32 {
7272    // SAFETY: callers uphold slice-length contracts (see call sites).
7273    unsafe {
7274        use core::arch::aarch64::*;
7275        let n = x.len();
7276        let wp = w.as_ptr() as *const i8;
7277        let xp = x.as_ptr();
7278        let cp = col.as_ptr();
7279        let (mut a0, mut a1, mut a2, mut a3) = (
7280            vdupq_n_f32(0.0),
7281            vdupq_n_f32(0.0),
7282            vdupq_n_f32(0.0),
7283            vdupq_n_f32(0.0),
7284        );
7285        let mut j = 0usize;
7286        while j + 16 <= n {
7287            let wb = vld1q_s8(wp.add(j));
7288            let lo = vmovl_s8(vget_low_s8(wb));
7289            let hi = vmovl_s8(vget_high_s8(wb));
7290            let w0 = vcvtq_f32_s32(vmovl_s16(vget_low_s16(lo)));
7291            let w1 = vcvtq_f32_s32(vmovl_s16(vget_high_s16(lo)));
7292            let w2 = vcvtq_f32_s32(vmovl_s16(vget_low_s16(hi)));
7293            let w3 = vcvtq_f32_s32(vmovl_s16(vget_high_s16(hi)));
7294            a0 = vfmaq_f32(
7295                a0,
7296                w0,
7297                vmulq_f32(vld1q_f32(xp.add(j)), vld1q_f32(cp.add(j))),
7298            );
7299            a1 = vfmaq_f32(
7300                a1,
7301                w1,
7302                vmulq_f32(vld1q_f32(xp.add(j + 4)), vld1q_f32(cp.add(j + 4))),
7303            );
7304            a2 = vfmaq_f32(
7305                a2,
7306                w2,
7307                vmulq_f32(vld1q_f32(xp.add(j + 8)), vld1q_f32(cp.add(j + 8))),
7308            );
7309            a3 = vfmaq_f32(
7310                a3,
7311                w3,
7312                vmulq_f32(vld1q_f32(xp.add(j + 12)), vld1q_f32(cp.add(j + 12))),
7313            );
7314            j += 16;
7315        }
7316        let mut sum = vaddvq_f32(vaddq_f32(vaddq_f32(a0, a1), vaddq_f32(a2, a3)));
7317        while j < n {
7318            sum += (*wp.add(j)) as f32 * *xp.add(j) * *cp.add(j);
7319            j += 1;
7320        }
7321        sum
7322    }
7323}
7324
7325#[cfg(target_arch = "aarch64")]
7326#[target_feature(enable = "neon")]
7327unsafe fn dot_i8_f32_neon(w: &[u8], x: &[f32]) -> f32 {
7328    // SAFETY: callers uphold slice-length contracts (see call sites).
7329    unsafe {
7330        use core::arch::aarch64::*;
7331        let n = x.len();
7332        let wp = w.as_ptr() as *const i8;
7333        let xp = x.as_ptr();
7334        let (mut a0, mut a1, mut a2, mut a3) = (
7335            vdupq_n_f32(0.0),
7336            vdupq_n_f32(0.0),
7337            vdupq_n_f32(0.0),
7338            vdupq_n_f32(0.0),
7339        );
7340        let mut j = 0usize;
7341        while j + 16 <= n {
7342            let wb = vld1q_s8(wp.add(j));
7343            let lo = vmovl_s8(vget_low_s8(wb));
7344            let hi = vmovl_s8(vget_high_s8(wb));
7345            let w0 = vcvtq_f32_s32(vmovl_s16(vget_low_s16(lo)));
7346            let w1 = vcvtq_f32_s32(vmovl_s16(vget_high_s16(lo)));
7347            let w2 = vcvtq_f32_s32(vmovl_s16(vget_low_s16(hi)));
7348            let w3 = vcvtq_f32_s32(vmovl_s16(vget_high_s16(hi)));
7349            a0 = vfmaq_f32(a0, w0, vld1q_f32(xp.add(j)));
7350            a1 = vfmaq_f32(a1, w1, vld1q_f32(xp.add(j + 4)));
7351            a2 = vfmaq_f32(a2, w2, vld1q_f32(xp.add(j + 8)));
7352            a3 = vfmaq_f32(a3, w3, vld1q_f32(xp.add(j + 12)));
7353            j += 16;
7354        }
7355        let mut sum = vaddvq_f32(vaddq_f32(vaddq_f32(a0, a1), vaddq_f32(a2, a3)));
7356        while j < n {
7357            sum += (*wp.add(j)) as f32 * *xp.add(j);
7358            j += 1;
7359        }
7360        sum
7361    }
7362}
7363
7364#[allow(clippy::too_many_arguments)]
7365fn qmatvec(
7366    q: &[u8],
7367    rep: &[u8],
7368    row_scale: &[f32],
7369    x: &[f32],
7370    col_field: &[f32],
7371    dtype: TensorDtype,
7372    rows: usize,
7373    cols: usize,
7374    out: &mut [f32],
7375    pool: Option<&Pool>,
7376) {
7377    debug_assert_eq!(out.len(), rows);
7378    #[cfg(not(target_arch = "aarch64"))]
7379    let _ = rep;
7380
7381    #[cfg(target_arch = "aarch64")]
7382    if sdot_enabled() {
7383        let act = if dtype == TensorDtype::Q8_2f {
7384            split_act_q8_2f(x, col_field)
7385        } else {
7386            split_act(x)
7387        };
7388        let out_addr = SendMut(out.as_mut_ptr());
7389        let run_range = |start: usize, end: usize| {
7390            q8_range_sdot(q, rep, row_scale, &act, cols, out_addr, start, end)
7391        };
7392        match pool {
7393            Some(pool) if rows >= 256 => pool.run_rows(rows, &run_range),
7394            _ => run_range(0, rows),
7395        }
7396        return;
7397    }
7398    // x86 A8W8 via AVX2 maddubs — same quantized-activation contract as
7399    // the SDOT path (CMF_AVX2=0 keeps the exact i8×f32 loop).
7400    #[cfg(target_arch = "x86_64")]
7401    if avx2_a8w8_enabled() {
7402        let act = if dtype == TensorDtype::Q8_2f {
7403            split_act_q8_2f(x, col_field)
7404        } else {
7405            split_act(x)
7406        };
7407        let out_addr = SendMut(out.as_mut_ptr());
7408        let run_range = |start: usize, end: usize| {
7409            q8_range_avx2(q, row_scale, &act, cols, out_addr, start, end)
7410        };
7411        match pool {
7412            Some(pool) if rows >= 256 => pool.run_rows(rows, &run_range),
7413            _ => run_range(0, rows),
7414        }
7415        return;
7416    }
7417
7418    prescale_with(x, col_field, dtype, 1, |xs| {
7419        let out_addr = SendMut(out.as_mut_ptr());
7420        let run_range = move |start: usize, end: usize| {
7421            for o in start..end {
7422                let v = dot_i8_f32(&q[o * cols..(o + 1) * cols], xs) * row_scale[o];
7423                // SAFETY: disjoint row ranges per worker.
7424                unsafe { *out_addr.at(o) = v };
7425            }
7426        };
7427        match pool {
7428            Some(pool) if rows >= 256 => pool.run_rows(rows, &run_range),
7429            _ => run_range(0, rows),
7430        }
7431    });
7432}
7433
7434#[allow(clippy::too_many_arguments)]
7435fn qmatvec2(
7436    q: &[u8],
7437    row_scale: &[f32],
7438    x1: &[f32],
7439    x2: &[f32],
7440    col_field: &[f32],
7441    dtype: TensorDtype,
7442    rows: usize,
7443    cols: usize,
7444    o1: &mut [f32],
7445    o2: &mut [f32],
7446    pool: Option<&Pool>,
7447) {
7448    #[cfg(target_arch = "aarch64")]
7449    if sdot_enabled() {
7450        let a1s = if dtype == TensorDtype::Q8_2f {
7451            split_act_q8_2f(x1, col_field)
7452        } else {
7453            split_act(x1)
7454        };
7455        let a2s = if dtype == TensorDtype::Q8_2f {
7456            split_act_q8_2f(x2, col_field)
7457        } else {
7458            split_act(x2)
7459        };
7460        let p1 = SendMut(o1.as_mut_ptr());
7461        let p2 = SendMut(o2.as_mut_ptr());
7462        let run_range = |start: usize, end: usize| {
7463            q8_range2_sdot(q, row_scale, &a1s, &a2s, cols, p1, p2, start, end)
7464        };
7465        match pool {
7466            Some(pool) if rows >= 256 => pool.run_rows(rows, &run_range),
7467            _ => run_range(0, rows),
7468        }
7469        return;
7470    }
7471    #[cfg(target_arch = "x86_64")]
7472    if avx2_a8w8_enabled() {
7473        let a1s = if dtype == TensorDtype::Q8_2f {
7474            split_act_q8_2f(x1, col_field)
7475        } else {
7476            split_act(x1)
7477        };
7478        let a2s = if dtype == TensorDtype::Q8_2f {
7479            split_act_q8_2f(x2, col_field)
7480        } else {
7481            split_act(x2)
7482        };
7483        let p1 = SendMut(o1.as_mut_ptr());
7484        let p2 = SendMut(o2.as_mut_ptr());
7485        let run_range = |start: usize, end: usize| {
7486            q8_range2_avx2(q, row_scale, &a1s, &a2s, cols, p1, p2, start, end)
7487        };
7488        match pool {
7489            Some(pool) if rows >= 256 => pool.run_rows(rows, &run_range),
7490            _ => run_range(0, rows),
7491        }
7492        return;
7493    }
7494
7495    prescale_with(x1, col_field, dtype, 1, |x1s| {
7496        prescale_with(x2, col_field, dtype, 2, |x2s| {
7497            let p1 = SendMut(o1.as_mut_ptr());
7498            let p2 = SendMut(o2.as_mut_ptr());
7499            let run_range = move |start: usize, end: usize| {
7500                for o in start..end {
7501                    let row = &q[o * cols..(o + 1) * cols];
7502                    let s1 = dot_i8_f32(row, x1s) * row_scale[o];
7503                    let s2 = dot_i8_f32(row, x2s) * row_scale[o];
7504                    // SAFETY: disjoint row ranges per worker.
7505                    unsafe {
7506                        *p1.at(o) = s1;
7507                        *p2.at(o) = s2;
7508                    }
7509                }
7510            };
7511            match pool {
7512                Some(pool) if rows >= 256 => pool.run_rows(rows, &run_range),
7513                _ => run_range(0, rows),
7514            }
7515        });
7516    });
7517}
7518
7519#[derive(Clone, Copy)]
7520struct SendMut(*mut f32);
7521unsafe impl Send for SendMut {}
7522unsafe impl Sync for SendMut {}
7523
7524impl SendMut {
7525    #[inline]
7526    fn at(self, i: usize) -> *mut f32 {
7527        unsafe { self.0.add(i) }
7528    }
7529}
7530
7531#[cfg(test)]
7532mod tests {
7533    use super::*;
7534
7535    #[test]
7536    fn f32_matvec_matches_matvec_rows_bitexact() {
7537        let (rows, cols) = (300, 40);
7538        let w: Vec<f32> = (0..rows * cols).map(|i| (i as f32 * 0.017).sin()).collect();
7539        let x: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.05).cos()).collect();
7540        let qt = QTensor::from_f32(w.clone(), rows, cols);
7541
7542        let mut a = vec![0.0f32; rows];
7543        matvec_rows(None, &w, &x, &mut a);
7544        let mut b = vec![0.0f32; rows];
7545        qt.matvec(&x, &mut b, None);
7546        assert_eq!(a, b);
7547    }
7548
7549    #[test]
7550    fn sdot_kernel_exact_on_grid() {
7551        // Activations already on the i8 grid (±1 with amax=1 → sx=1/127,
7552        // xq=±127 dequantizes EXACTLY) → the SDOT path must match the
7553        // exact f32 dot to float rounding. This isolates kernel
7554        // correctness from quantization noise.
7555        eprintln!("sdot_enabled = {}", sdot_enabled());
7556        let (rows, cols) = (9, 80); // odd rows → exercises 4-row + tail
7557        let w: Vec<u8> = (0..rows * cols)
7558            .map(|i| (((i * 37) % 251) as i32 - 125) as i8 as u8)
7559            .collect();
7560        let scales: Vec<f32> = (0..rows).map(|o| 0.005 + o as f32 * 0.001).collect();
7561        let x: Vec<f32> = (0..cols)
7562            .map(|i| match i % 3 {
7563                0 => 1.0,
7564                1 => -1.0,
7565                _ => 0.0,
7566            })
7567            .collect();
7568        let mut a = vec![0.0f32; rows];
7569        qmatvec(
7570            &w,
7571            &[],
7572            &scales,
7573            &x,
7574            &[],
7575            TensorDtype::Q8Row,
7576            rows,
7577            cols,
7578            &mut a,
7579            None,
7580        );
7581        for o in 0..rows {
7582            let mut acc = 0.0f32;
7583            for j in 0..cols {
7584                acc += (w[o * cols + j] as i8) as f32 * x[j];
7585            }
7586            let expect = acc * scales[o];
7587            assert!(
7588                (a[o] - expect).abs() < 1e-3 * expect.abs().max(1e-3),
7589                "row {o}: {} vs {expect}",
7590                a[o]
7591            );
7592        }
7593    }
7594
7595    #[test]
7596    fn q1_tbl_fast_path_matches_reference() {
7597        // gpr = 8 exercises the TBL pair-load fast loop, and the LAST
7598        // row's final 4-tile window trips the 4B-overread guard (the
7599        // payload ends exactly at the last tile) — both paths must
7600        // agree with the dequant reference.
7601        let (rows, cols) = (5, 256);
7602        let gpr = cols / GROUP_SIZE;
7603        let mut bytes = Vec::new();
7604        for t in 0..rows * gpr {
7605            let s = 0.007 + (t % 11) as f32 * 0.004;
7606            bytes.extend_from_slice(&cortiq_core::quant::f32_to_f16(s).to_le_bytes());
7607            for j in 0..4 {
7608                bytes.push(((t * 53 + j * 89 + 7) % 249) as u8);
7609            }
7610        }
7611        let x: Vec<f32> = (0..cols)
7612            .map(|i| if (i * 5) % 7 < 3 { 1.0 } else { -1.0 })
7613            .collect();
7614        let mut w = vec![0.0f32; rows * cols];
7615        cortiq_core::quant::dequant_q1(&bytes, &mut w);
7616        let mut got = vec![0.0f32; rows];
7617        q1_matvec(&bytes, &x, rows, cols, &mut got, None);
7618        for o in 0..rows {
7619            let expect: f32 = (0..cols).map(|j| w[o * cols + j] * x[j]).sum();
7620            assert!(
7621                (got[o] - expect).abs() < 1e-3 * expect.abs().max(1e-3),
7622                "row {o}: {} vs {expect}",
7623                got[o]
7624            );
7625        }
7626        // Blocked 1×4 batch (b=5: one quad + remainder) must equal the
7627        // single-matvec path bit-for-bit.
7628        let b = 5usize;
7629        let mut xs_all = Vec::new();
7630        for bi in 0..b {
7631            xs_all.extend(x.iter().map(|v| if bi % 2 == 0 { *v } else { -*v }));
7632        }
7633        let mut mm = vec![0.0f32; b * rows];
7634        q1_matmat(&bytes, &xs_all, b, rows, cols, &mut mm, None);
7635        for bi in 0..b {
7636            let mut single = vec![0.0f32; rows];
7637            q1_matvec(
7638                &bytes,
7639                &xs_all[bi * cols..(bi + 1) * cols],
7640                rows,
7641                cols,
7642                &mut single,
7643                None,
7644            );
7645            assert_eq!(&mm[bi * rows..(bi + 1) * rows], &single[..], "stream {bi}");
7646        }
7647    }
7648
7649    #[test]
7650    fn q1_kernels_match_exact_reference() {
7651        // Synthetic q1 payload: 6-byte tiles [f16 scale][4B bits].
7652        let (rows, cols) = (7, 96);
7653        let gpr = cols / GROUP_SIZE;
7654        let mut bytes = Vec::new();
7655        for t in 0..rows * gpr {
7656            let s = 0.01 + (t % 13) as f32 * 0.003;
7657            bytes.extend_from_slice(&cortiq_core::quant::f32_to_f16(s).to_le_bytes());
7658            for j in 0..4 {
7659                bytes.push(((t * 31 + j * 97) % 251) as u8);
7660            }
7661        }
7662        // On-grid activations (±1, amax 1) → the SDOT path is exact.
7663        let x: Vec<f32> = (0..cols)
7664            .map(|i| if i % 3 == 0 { 1.0 } else { -1.0 })
7665            .collect();
7666        // Reference through the core dequant.
7667        let mut w = vec![0.0f32; rows * cols];
7668        cortiq_core::quant::dequant_q1(&bytes, &mut w);
7669        let mut expect = vec![0.0f32; rows];
7670        for o in 0..rows {
7671            expect[o] = (0..cols).map(|j| w[o * cols + j] * x[j]).sum();
7672        }
7673        let mut got = vec![0.0f32; rows];
7674        q1_matvec(&bytes, &x, rows, cols, &mut got, None);
7675        for o in 0..rows {
7676            assert!(
7677                (got[o] - expect[o]).abs() < 1e-3 * expect[o].abs().max(1e-3),
7678                "row {o}: {} vs {}",
7679                got[o],
7680                expect[o]
7681            );
7682        }
7683        // Pair and batch paths agree with the single path.
7684        let x2: Vec<f32> = x.iter().map(|v| -v).collect();
7685        let (mut a1, mut a2) = (vec![0.0f32; rows], vec![0.0f32; rows]);
7686        q1_matvec2(&bytes, &x, &x2, rows, cols, &mut a1, &mut a2, None);
7687        assert_eq!(a1, got);
7688        let mut xs = x.clone();
7689        xs.extend_from_slice(&x2);
7690        let mut mm = vec![0.0f32; 2 * rows];
7691        q1_matmat(&bytes, &xs, 2, rows, cols, &mut mm, None);
7692        assert_eq!(&mm[..rows], got.as_slice());
7693        assert_eq!(&mm[rows..], a2.as_slice());
7694    }
7695
7696    #[test]
7697    fn repack_is_bit_identical() {
7698        // The interleaved-repack kernel must produce EXACTLY the same
7699        // bits as the mmap-layout kernel: integer accumulation is order-
7700        // exact, the f32 epilogue is identical. Odd rows exercise the
7701        // tail; direct range calls exercise unaligned pool splits.
7702        let (rows, cols) = (267, 96); // 66 groups + 3 tail rows, cols % 16 == 0
7703        let w: Vec<u8> = (0..rows * cols)
7704            .map(|i| (((i * 89) % 253) as i32 - 126) as i8 as u8)
7705            .collect();
7706        let scales: Vec<f32> = (0..rows).map(|o| 0.003 + o as f32 * 0.0007).collect();
7707        let x: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.37).sin() * 2.0).collect();
7708        let rep = q8_repack_layout(&w, rows, cols);
7709        // Group interleave round-trips.
7710        for g in 0..rows / 4 {
7711            for c in 0..cols / 16 {
7712                for lane in 0..4 {
7713                    assert_eq!(
7714                        &rep[g * 4 * cols + c * 64 + lane * 16
7715                            ..g * 4 * cols + c * 64 + lane * 16 + 16],
7716                        &w[(g * 4 + lane) * cols + c * 16..(g * 4 + lane) * cols + c * 16 + 16],
7717                    );
7718                }
7719            }
7720        }
7721        let mut a = vec![0.0f32; rows];
7722        qmatvec(
7723            &w,
7724            &[],
7725            &scales,
7726            &x,
7727            &[],
7728            TensorDtype::Q8Row,
7729            rows,
7730            cols,
7731            &mut a,
7732            None,
7733        );
7734        let mut b = vec![0.0f32; rows];
7735        qmatvec(
7736            &w,
7737            &rep,
7738            &scales,
7739            &x,
7740            &[],
7741            TensorDtype::Q8Row,
7742            rows,
7743            cols,
7744            &mut b,
7745            None,
7746        );
7747        assert_eq!(a, b, "full-range repack output diverged");
7748
7749        #[cfg(target_arch = "aarch64")]
7750        if sdot_enabled() {
7751            // Unaligned range split (pool workers get arbitrary bounds).
7752            let act = split_act(&x);
7753            let mut c1 = vec![0.0f32; rows];
7754            let mut c2 = vec![0.0f32; rows];
7755            q8_range_sdot(
7756                &w,
7757                &[],
7758                &scales,
7759                &act,
7760                cols,
7761                SendMut(c1.as_mut_ptr()),
7762                3,
7763                rows - 2,
7764            );
7765            q8_range_sdot(
7766                &w,
7767                &rep,
7768                &scales,
7769                &act,
7770                cols,
7771                SendMut(c2.as_mut_ptr()),
7772                3,
7773                rows - 2,
7774            );
7775            assert_eq!(c1, c2, "unaligned-range repack output diverged");
7776        }
7777    }
7778
7779    #[test]
7780    fn sdot_a8w8_noise_is_bounded() {
7781        // Off-grid activations: A8 quantization noise must stay small in
7782        // relative L2 over the whole output (realistic accuracy contract;
7783        // vmfcore measured argmax-identical decode on real models).
7784        let (rows, cols) = (16, 512);
7785        let w: Vec<u8> = (0..rows * cols)
7786            .map(|i| (((i * 37) % 251) as i32 - 125) as i8 as u8)
7787            .collect();
7788        let scales = vec![0.01f32; rows];
7789        let x: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.21).sin()).collect();
7790        let mut a = vec![0.0f32; rows];
7791        qmatvec(
7792            &w,
7793            &[],
7794            &scales,
7795            &x,
7796            &[],
7797            TensorDtype::Q8Row,
7798            rows,
7799            cols,
7800            &mut a,
7801            None,
7802        );
7803        let (mut num, mut den) = (0f64, 0f64);
7804        for o in 0..rows {
7805            let mut acc = 0.0f32;
7806            for j in 0..cols {
7807                acc += (w[o * cols + j] as i8) as f32 * x[j];
7808            }
7809            let expect = acc * scales[o];
7810            num += ((a[o] - expect) as f64).powi(2);
7811            den += (expect as f64).powi(2);
7812        }
7813        let rel = (num / den.max(1e-12)).sqrt();
7814        assert!(rel < 0.05, "A8W8 relative L2 error too high: {rel}");
7815    }
7816
7817    #[test]
7818    fn i8_dot_neon_matches_scalar() {
7819        let n = 100;
7820        let w: Vec<u8> = (0..n).map(|i| ((i * 37 + 11) % 251) as u8).collect();
7821        let x: Vec<f32> = (0..n).map(|i| (i as f32 * 0.13).sin()).collect();
7822        let mut scalar = 0.0f32;
7823        for j in 0..n {
7824            scalar += (w[j] as i8) as f32 * x[j];
7825        }
7826        let fast = dot_i8_f32(&w, &x);
7827        assert!((scalar - fast).abs() < 1e-3 * scalar.abs().max(1.0));
7828    }
7829
7830    /// Fused vbit matvec must match full dequant_vbit + dense matvec.
7831    #[test]
7832    fn vbitmatvec_matches_full_dequant() {
7833        let (rows, cols) = (6, 64);
7834        let ng = cols / GROUP_SIZE;
7835        // Hand-craft: bits per row, f16 scales, packed rows.
7836        let bits: Vec<u8> = vec![3, 4, 5, 6, 8, 4];
7837        let mut bytes = bits.clone();
7838        for g in 0..rows * ng {
7839            let s = 0.02 + 0.001 * g as f32;
7840            bytes.extend_from_slice(&cortiq_core::quant::f32_to_f16(s).to_le_bytes());
7841        }
7842        for r in 0..rows {
7843            let b = bits[r] as usize;
7844            let (mut acc, mut nb) = (0u64, 0usize);
7845            let mut rowbytes = Vec::new();
7846            for i in 0..cols {
7847                let v = ((i * 7 + r * 13) % (1 << b)) as u64;
7848                acc = (acc << b) | v;
7849                nb += b;
7850                while nb >= 8 {
7851                    nb -= 8;
7852                    rowbytes.push(((acc >> nb) & 0xFF) as u8);
7853                }
7854            }
7855            if nb > 0 {
7856                rowbytes.push(((acc << (8 - nb)) & 0xFF) as u8);
7857            }
7858            bytes.extend_from_slice(&rowbytes);
7859        }
7860        let x: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.19).sin()).collect();
7861
7862        let mut reference = vec![0f32; rows * cols];
7863        cortiq_core::quant::dequant_vbit(&bytes, rows, cols, &mut reference).unwrap();
7864        let mut expect = vec![0f32; rows];
7865        for r in 0..rows {
7866            expect[r] = reference[r * cols..(r + 1) * cols]
7867                .iter()
7868                .zip(&x)
7869                .map(|(w, xv)| w * xv)
7870                .sum();
7871        }
7872        let mut got = vec![0f32; rows];
7873        let offsets = vbit_row_offsets(&bytes, rows, cols);
7874        vbitmatvec(&bytes, &offsets, &x, rows, cols, &mut got, None);
7875        // SDOT path quantizes activations to i8 (A8W8): bounded noise,
7876        // same contract as q8 (exact path is pinned by CMF_SDOT=0 in
7877        // the golden-parity gate).
7878        let tol = if a8w8_enabled() { 6e-2 } else { 1e-4 };
7879        let scale = expect.iter().fold(0f32, |m, v| m.max(v.abs())).max(1e-6);
7880        for r in 0..rows {
7881            assert!(
7882                (got[r] - expect[r]).abs() < tol * scale,
7883                "row {r}: {} vs {}",
7884                got[r],
7885                expect[r]
7886            );
7887        }
7888    }
7889
7890    /// Fused q4 matvec must match the reference full-dequant + dense
7891    /// matvec bit-for-bit in structure (same f32 math, group order).
7892    /// vbit matmat: the blocked 1×4 leg must match the per-row path
7893    /// (paired env toggle; larger shape so both code paths engage).
7894    #[test]
7895    #[cfg(target_arch = "x86_64")]
7896    fn vbit_matmat_blocked_matches_per_row() {
7897        let (rows, cols, b) = (64usize, 128usize, 9usize);
7898        let ng = cols / GROUP_SIZE;
7899        let bits: Vec<u8> = (0..rows).map(|r| [3u8, 4, 5, 6][r % 4]).collect();
7900        let mut bytes = bits.clone();
7901        for g in 0..rows * ng {
7902            let sc = 0.02 + 0.0005 * g as f32;
7903            bytes.extend_from_slice(&cortiq_core::quant::f32_to_f16(sc).to_le_bytes());
7904        }
7905        for r in 0..rows {
7906            let bw = bits[r] as usize;
7907            let (mut acc, mut nb) = (0u64, 0usize);
7908            let mut rowbytes = Vec::new();
7909            for i in 0..cols {
7910                let v = ((i * 7 + r * 13) % (1 << bw)) as u64;
7911                acc = (acc << bw) | v;
7912                nb += bw;
7913                while nb >= 8 {
7914                    nb -= 8;
7915                    rowbytes.push(((acc >> nb) & 0xFF) as u8);
7916                }
7917            }
7918            if nb > 0 {
7919                rowbytes.push(((acc << (8 - nb)) & 0xFF) as u8);
7920            }
7921            bytes.extend_from_slice(&rowbytes);
7922        }
7923        let x: Vec<f32> = (0..b * cols)
7924            .map(|i| ((i * 13 + 7) % 97) as f32 / 97.0 - 0.5)
7925            .collect();
7926        let offsets = vbit_row_offsets(&bytes, rows, cols);
7927        let mut y_a = vec![0f32; b * rows];
7928        let mut y_b = vec![0f32; b * rows];
7929        unsafe { std::env::set_var("CMF_X86_BLOCKED", "1") };
7930        vbitmatmat(&bytes, &offsets, &x, b, rows, cols, &mut y_a, None);
7931        unsafe { std::env::set_var("CMF_X86_BLOCKED", "0") };
7932        vbitmatmat(&bytes, &offsets, &x, b, rows, cols, &mut y_b, None);
7933        unsafe { std::env::remove_var("CMF_X86_BLOCKED") };
7934        let max_d = y_a
7935            .iter()
7936            .zip(&y_b)
7937            .map(|(p, q)| (p - q).abs())
7938            .fold(0.0f32, f32::max);
7939        assert!(max_d < 1e-4, "vbit blocked ≠ per-row: max|Δ| = {max_d}");
7940    }
7941
7942    /// q4t blocked 1×4 (SDOT on ARM, AVX2 on x86) must equal the
7943    /// per-row path exactly: same nibble unpack, same group order,
7944    /// same f32 accumulation — batch == matvec bit-for-bit. b=9 covers
7945    /// two full 1×4 blocks plus a remainder through the single-row
7946    /// kernel. (Both paths produce identical output, so the shared
7947    /// CMF_X86_BLOCKED env var racing with other tests cannot flip
7948    /// the verdict — worst case both sides take the same path.)
7949    #[test]
7950    fn q4t_matmat_blocked_matches_per_row() {
7951        let (rows, cols, b) = (16usize, 64usize, 9usize);
7952        let gpr = cols / GROUP_SIZE;
7953        let mut bytes = vec![0u8; rows * gpr * Q4_TILE];
7954        for r in 0..rows {
7955            for g in 0..gpr {
7956                let t = (r * gpr + g) * Q4_TILE;
7957                let sc = 0.02 + 0.001 * (r * gpr + g) as f32;
7958                bytes[t..t + 2].copy_from_slice(&cortiq_core::quant::f32_to_f16(sc).to_le_bytes());
7959                for k in 0..16 {
7960                    bytes[t + 2 + k] = ((r * 31 + g * 7 + k * 13) % 251) as u8;
7961                }
7962            }
7963        }
7964        let x: Vec<f32> = (0..b * cols)
7965            .map(|i| ((i * 13 + 7) % 97) as f32 / 97.0 - 0.5)
7966            .collect();
7967        let mut y_blk = vec![0f32; b * rows];
7968        let mut y_row = vec![0f32; b * rows];
7969        unsafe { std::env::set_var("CMF_X86_BLOCKED", "1") };
7970        q4t_matmat(&bytes, &x, b, rows, cols, &mut y_blk, None);
7971        unsafe { std::env::set_var("CMF_X86_BLOCKED", "0") };
7972        q4t_matmat(&bytes, &x, b, rows, cols, &mut y_row, None);
7973        unsafe { std::env::remove_var("CMF_X86_BLOCKED") };
7974        assert_eq!(y_blk, y_row, "q4t blocked 1x4 ≠ per-row");
7975    }
7976
7977    #[test]
7978    fn q4matvec_matches_full_dequant() {
7979        let (rows, cols) = (8, 64);
7980        let groups = rows * cols / GROUP_SIZE;
7981        // Hand-craft a q4_block blob: nibbles then f16 scales.
7982        let mut bytes = Vec::with_capacity(groups * 16 + groups * 2);
7983        for i in 0..groups * 16 {
7984            bytes.push((((i * 7 + 3) % 256) & 0xFF) as u8);
7985        }
7986        for g in 0..groups {
7987            let s = 0.01 + 0.003 * g as f32;
7988            bytes.extend_from_slice(&cortiq_core::quant::f32_to_f16(s).to_le_bytes());
7989        }
7990        let x: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.17).sin()).collect();
7991
7992        let mut reference = vec![0.0f32; rows * cols];
7993        cortiq_core::quant::dequant_q4_block(&bytes, &mut reference);
7994        let mut expect = vec![0.0f32; rows];
7995        for r in 0..rows {
7996            expect[r] = reference[r * cols..(r + 1) * cols]
7997                .iter()
7998                .zip(&x)
7999                .map(|(w, xv)| w * xv)
8000                .sum();
8001        }
8002
8003        let mut got = vec![0.0f32; rows];
8004        q4matvec(&bytes, &x, rows, cols, &mut got, None);
8005        // SDOT path quantizes activations to i8 (A8W8): bounded noise,
8006        // same contract as q8/vbit (exact path is pinned by CMF_SDOT=0
8007        // in the golden-parity gate).
8008        let tol = if a8w8_enabled() { 6e-2 } else { 1e-4 };
8009        let scale = expect.iter().fold(0f32, |m, v| m.max(v.abs())).max(1.0);
8010        for r in 0..rows {
8011            assert!(
8012                (got[r] - expect[r]).abs() < tol * scale,
8013                "row {r}: {} vs {}",
8014                got[r],
8015                expect[r]
8016            );
8017        }
8018    }
8019
8020    /// Fused two-input vbit matvec must equal two single matvecs exactly
8021    /// (same per-lane accumulation order on both scalar and SDOT paths).
8022    #[test]
8023    fn vbitmatvec2_equals_two_singles() {
8024        let (rows, cols) = (6, 64);
8025        let ng = cols / GROUP_SIZE;
8026        let bits: Vec<u8> = vec![3, 4, 5, 6, 8, 4];
8027        let mut bytes = bits.clone();
8028        for g in 0..rows * ng {
8029            let s = 0.02 + 0.001 * g as f32;
8030            bytes.extend_from_slice(&cortiq_core::quant::f32_to_f16(s).to_le_bytes());
8031        }
8032        for r in 0..rows {
8033            let b = bits[r] as usize;
8034            let (mut acc, mut nb) = (0u64, 0usize);
8035            let mut rowbytes = Vec::new();
8036            for i in 0..cols {
8037                let v = ((i * 7 + r * 13) % (1 << b)) as u64;
8038                acc = (acc << b) | v;
8039                nb += b;
8040                while nb >= 8 {
8041                    nb -= 8;
8042                    rowbytes.push(((acc >> nb) & 0xFF) as u8);
8043                }
8044            }
8045            if nb > 0 {
8046                rowbytes.push(((acc << (8 - nb)) & 0xFF) as u8);
8047            }
8048            bytes.extend_from_slice(&rowbytes);
8049        }
8050        let x1: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.19).sin()).collect();
8051        let x2: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.11).cos()).collect();
8052        let offsets = vbit_row_offsets(&bytes, rows, cols);
8053
8054        let (mut a1, mut a2) = (vec![0f32; rows], vec![0f32; rows]);
8055        vbitmatvec(&bytes, &offsets, &x1, rows, cols, &mut a1, None);
8056        vbitmatvec(&bytes, &offsets, &x2, rows, cols, &mut a2, None);
8057        let (mut b1, mut b2) = (vec![0f32; rows], vec![0f32; rows]);
8058        vbitmatvec2(
8059            &bytes, &offsets, &x1, &x2, rows, cols, &mut b1, &mut b2, None,
8060        );
8061        assert_eq!(a1, b1, "fused vbit lane 1 must be bit-identical");
8062        assert_eq!(a2, b2, "fused vbit lane 2 must be bit-identical");
8063    }
8064
8065    /// Fused two-input q4 matvec must equal two single matvecs exactly.
8066    #[test]
8067    fn q4matvec2_equals_two_singles() {
8068        let (rows, cols) = (8, 128);
8069        let groups = rows * cols / GROUP_SIZE;
8070        let mut bytes = Vec::with_capacity(groups * 16 + groups * 2);
8071        for i in 0..groups * 16 {
8072            bytes.push((((i * 7 + 3) % 256) & 0xFF) as u8);
8073        }
8074        for g in 0..groups {
8075            let s = 0.01 + 0.003 * g as f32;
8076            bytes.extend_from_slice(&cortiq_core::quant::f32_to_f16(s).to_le_bytes());
8077        }
8078        // Include an outlier channel so the SDOT correction path is
8079        // exercised in the pair kernel too.
8080        let mut x1: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.17).sin()).collect();
8081        x1[9] = 250.0;
8082        let x2: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.23).cos()).collect();
8083
8084        let (mut a1, mut a2) = (vec![0f32; rows], vec![0f32; rows]);
8085        q4matvec(&bytes, &x1, rows, cols, &mut a1, None);
8086        q4matvec(&bytes, &x2, rows, cols, &mut a2, None);
8087        let (mut b1, mut b2) = (vec![0f32; rows], vec![0f32; rows]);
8088        q4matvec2(&bytes, &x1, &x2, rows, cols, &mut b1, &mut b2, None);
8089        assert_eq!(a1, b1, "fused q4 lane 1 must be bit-identical");
8090        assert_eq!(a2, b2, "fused q4 lane 2 must be bit-identical");
8091    }
8092
8093    /// Multi-matrix job must equal separate matvecs exactly — same
8094    /// kernels, only the dispatch is fused.
8095    #[test]
8096    fn matvec_many_equals_separate_matvecs() {
8097        use crate::pool::Pool;
8098        let (r1, r2, cols) = (300, 200, 64);
8099        let mk = |salt: usize, rows: usize| {
8100            QTensor::from_f32(
8101                (0..rows * cols)
8102                    .map(|i| ((i * 7 + salt) % 97) as f32 / 97.0 - 0.5)
8103                    .collect(),
8104                rows,
8105                cols,
8106            )
8107        };
8108        let (a, b) = (mk(1, r1), mk(5, r2));
8109        let x: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.11).sin()).collect();
8110        let pool = Pool::new(3);
8111
8112        let (mut ea, mut eb) = (vec![0f32; r1], vec![0f32; r2]);
8113        a.matvec(&x, &mut ea, Some(&pool));
8114        b.matvec(&x, &mut eb, Some(&pool));
8115        let (mut ga, mut gb) = (vec![0f32; r1], vec![0f32; r2]);
8116        QTensor::matvec_many([&a, &b], &x, [&mut ga, &mut gb], Some(&pool));
8117        assert_eq!(ea, ga, "fused multi-matrix lane 1 must be bit-identical");
8118        assert_eq!(eb, gb, "fused multi-matrix lane 2 must be bit-identical");
8119    }
8120
8121    /// Batched q4/vbit matmat must equal per-position matvec calls
8122    /// exactly (the fallback it replaced) — same kernels, same order.
8123    #[test]
8124    fn batched_matmat_equals_per_position_matvec() {
8125        let (rows, cols, b) = (8, 64, 5);
8126        // q4 blob.
8127        let groups = rows * cols / GROUP_SIZE;
8128        let mut q4 = Vec::new();
8129        for i in 0..groups * 16 {
8130            q4.push((((i * 7 + 3) % 256) & 0xFF) as u8);
8131        }
8132        for g in 0..groups {
8133            q4.extend_from_slice(
8134                &cortiq_core::quant::f32_to_f16(0.01 + 0.003 * g as f32).to_le_bytes(),
8135            );
8136        }
8137        // vbit blob (mixed widths incl. 8).
8138        let ng = cols / GROUP_SIZE;
8139        let bits: Vec<u8> = vec![3, 4, 5, 6, 8, 4, 5, 3];
8140        let mut vb = bits.clone();
8141        for g in 0..rows * ng {
8142            vb.extend_from_slice(
8143                &cortiq_core::quant::f32_to_f16(0.02 + 0.001 * g as f32).to_le_bytes(),
8144            );
8145        }
8146        for r in 0..rows {
8147            let bw = bits[r] as usize;
8148            let (mut acc, mut nb) = (0u64, 0usize);
8149            let mut rowbytes = Vec::new();
8150            for i in 0..cols {
8151                let v = ((i * 7 + r * 13) % (1 << bw)) as u64;
8152                acc = (acc << bw) | v;
8153                nb += bw;
8154                while nb >= 8 {
8155                    nb -= 8;
8156                    rowbytes.push(((acc >> nb) & 0xFF) as u8);
8157                }
8158            }
8159            if nb > 0 {
8160                rowbytes.push(((acc << (8 - nb)) & 0xFF) as u8);
8161            }
8162            vb.extend_from_slice(&rowbytes);
8163        }
8164        let offsets = vbit_row_offsets(&vb, rows, cols);
8165
8166        let xs: Vec<f32> = (0..b * cols).map(|i| (i as f32 * 0.13).sin()).collect();
8167
8168        // q4: batch vs singles.
8169        let mut got = vec![0f32; b * rows];
8170        q4matmat(&q4, &xs, b, rows, cols, &mut got, None);
8171        for bi in 0..b {
8172            let mut expect = vec![0f32; rows];
8173            q4matvec(
8174                &q4,
8175                &xs[bi * cols..(bi + 1) * cols],
8176                rows,
8177                cols,
8178                &mut expect,
8179                None,
8180            );
8181            assert_eq!(
8182                &got[bi * rows..(bi + 1) * rows],
8183                &expect[..],
8184                "q4 batch pos {bi}"
8185            );
8186        }
8187
8188        // vbit: batch vs singles.
8189        let mut got = vec![0f32; b * rows];
8190        vbitmatmat(&vb, &offsets, &xs, b, rows, cols, &mut got, None);
8191        for bi in 0..b {
8192            let mut expect = vec![0f32; rows];
8193            vbitmatvec(
8194                &vb,
8195                &offsets,
8196                &xs[bi * cols..(bi + 1) * cols],
8197                rows,
8198                cols,
8199                &mut expect,
8200                None,
8201            );
8202            assert_eq!(
8203                &got[bi * rows..(bi + 1) * rows],
8204                &expect[..],
8205                "vbit batch pos {bi}"
8206            );
8207        }
8208    }
8209
8210    /// q4_tiled kernels must produce BIT-identical outputs to the q4
8211    /// split kernels on the same values (same ints, same order — only
8212    /// the byte placement differs).
8213    #[test]
8214    fn q4_tiled_matches_q4_block_bitexact() {
8215        let (rows, cols, b) = (8usize, 128usize, 3usize);
8216        let groups = rows * cols / GROUP_SIZE;
8217        let mut split = Vec::with_capacity(groups * 18);
8218        for i in 0..groups * 16 {
8219            split.push((((i * 7 + 3) % 256) & 0xFF) as u8);
8220        }
8221        for g in 0..groups {
8222            split.extend_from_slice(
8223                &cortiq_core::quant::f32_to_f16(0.01 + 0.003 * g as f32).to_le_bytes(),
8224            );
8225        }
8226        // Re-tile: [scale][nibbles] per group.
8227        let (packed, scales) = split.split_at(groups * 16);
8228        let mut tiled = Vec::with_capacity(groups * Q4_TILE);
8229        for g in 0..groups {
8230            tiled.extend_from_slice(&scales[g * 2..g * 2 + 2]);
8231            tiled.extend_from_slice(&packed[g * 16..(g + 1) * 16]);
8232        }
8233
8234        let mut x1: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.17).sin()).collect();
8235        x1[9] = 250.0; // exercise the outlier path
8236        let x2: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.23).cos()).collect();
8237
8238        let (mut a, mut t) = (vec![0f32; rows], vec![0f32; rows]);
8239        q4matvec(&split, &x1, rows, cols, &mut a, None);
8240        q4t_matvec(&tiled, &x1, rows, cols, &mut t, None);
8241        assert_eq!(a, t, "q4t matvec must match q4 bit-for-bit");
8242
8243        let (mut a1, mut a2) = (vec![0f32; rows], vec![0f32; rows]);
8244        let (mut t1, mut t2) = (vec![0f32; rows], vec![0f32; rows]);
8245        q4matvec2(&split, &x1, &x2, rows, cols, &mut a1, &mut a2, None);
8246        q4t_matvec2(&tiled, &x1, &x2, rows, cols, &mut t1, &mut t2, None);
8247        assert_eq!(a1, t1);
8248        assert_eq!(a2, t2);
8249
8250        let xs: Vec<f32> = (0..b * cols).map(|i| (i as f32 * 0.13).sin()).collect();
8251        let (mut am, mut tm) = (vec![0f32; b * rows], vec![0f32; b * rows]);
8252        q4matmat(&split, &xs, b, rows, cols, &mut am, None);
8253        q4t_matmat(&tiled, &xs, b, rows, cols, &mut tm, None);
8254        assert_eq!(am, tm, "q4t matmat must match q4 bit-for-bit");
8255    }
8256
8257    /// q4 SDOT outlier correction: a single huge activation channel
8258    /// (>8·rms → outlier, zeroed in xq) must still contribute its EXACT
8259    /// term. On-grid bulk (±1/0 → xq dequantizes exactly) isolates the
8260    /// correction from A8W8 noise. cols must exceed 64: at n=64 the
8261    /// 8·rms threshold equals sqrt(v²+rest) ≥ v, so a single outlier
8262    /// can never qualify (8² = n).
8263    #[test]
8264    fn q4matvec_sdot_outlier_exact() {
8265        let (rows, cols) = (4, 128);
8266        let groups = rows * cols / GROUP_SIZE;
8267        let mut bytes = Vec::with_capacity(groups * 16 + groups * 2);
8268        for i in 0..groups * 16 {
8269            bytes.push(((i * 11 + 5) % 256) as u8);
8270        }
8271        for g in 0..groups {
8272            let s = 0.02 + 0.002 * g as f32;
8273            bytes.extend_from_slice(&cortiq_core::quant::f32_to_f16(s).to_le_bytes());
8274        }
8275        let mut x: Vec<f32> = (0..cols)
8276            .map(|i| match i % 3 {
8277                0 => 1.0,
8278                1 => -1.0,
8279                _ => 0.0,
8280            })
8281            .collect();
8282        x[17] = 300.0; // ≫ 8·rms → outlier channel
8283
8284        let mut reference = vec![0.0f32; rows * cols];
8285        cortiq_core::quant::dequant_q4_block(&bytes, &mut reference);
8286        let mut expect = vec![0.0f32; rows];
8287        for r in 0..rows {
8288            expect[r] = reference[r * cols..(r + 1) * cols]
8289                .iter()
8290                .zip(&x)
8291                .map(|(w, xv)| w * xv)
8292                .sum();
8293        }
8294        let mut got = vec![0.0f32; rows];
8295        q4matvec(&bytes, &x, rows, cols, &mut got, None);
8296        let scale = expect.iter().fold(0f32, |m, v| m.max(v.abs())).max(1.0);
8297        for r in 0..rows {
8298            assert!(
8299                (got[r] - expect[r]).abs() < 2e-3 * scale,
8300                "row {r}: {} vs {} (outlier term must be exact)",
8301                got[r],
8302                expect[r]
8303            );
8304        }
8305    }
8306
8307    /// The fused q1t matvec must equal the reference (dequant_q1t → dot),
8308    /// including the ternary zero level and the binary-searched outlier
8309    /// overlay. Guards the mmap kernel that makes a 12B q1t runnable.
8310    #[test]
8311    fn q1t_matvec_matches_reference() {
8312        use cortiq_core::quant::{dequant_q1t, f32_to_f16};
8313        let (rows, cols) = (3usize, 64usize); // gpr = 2
8314        let gpr = cols / GROUP_SIZE;
8315        let scales = [0.5f32, 0.3, 0.7, 0.2, 0.6, 0.15];
8316        // Overlay (must be sorted by flat index): a few spikes across rows.
8317        let outliers: [(u32, f32); 3] = [(5, 9.0), (70, -4.5), (150, 3.25)];
8318        let is_out = |flat: usize| outliers.iter().any(|&(i, _)| i as usize == flat);
8319        let mut bytes = Vec::new();
8320        for r in 0..rows {
8321            for g in 0..gpr {
8322                bytes.extend_from_slice(&f32_to_f16(scales[r * gpr + g]).to_le_bytes());
8323                let mut c = [0u8; 7];
8324                for k in 0..GROUP_SIZE {
8325                    // Encoder invariant: code 0 at outlier positions.
8326                    let code = if is_out(r * cols + g * GROUP_SIZE + k) {
8327                        0
8328                    } else {
8329                        ((k + r * 3 + g) % 3) as u8 // 0,1,2
8330                    };
8331                    cortiq_core::quant::q1t_pack(&mut c, k, code);
8332                }
8333                bytes.extend_from_slice(&c);
8334            }
8335        }
8336        // Per-row overlay: [u32 row_ptr[rows+1]] then [(u16 col, f16 val)] by
8337        // row (outliers are sorted by flat index → already grouped by row).
8338        let mut row_ptr = vec![0u32; rows + 1];
8339        for &(idx, _) in &outliers {
8340            row_ptr[idx as usize / cols + 1] += 1;
8341        }
8342        for r in 0..rows {
8343            row_ptr[r + 1] += row_ptr[r];
8344        }
8345        for &p in &row_ptr {
8346            bytes.extend_from_slice(&p.to_le_bytes());
8347        }
8348        for &(idx, v) in &outliers {
8349            bytes.extend_from_slice(&((idx as usize % cols) as u16).to_le_bytes());
8350            bytes.extend_from_slice(&f32_to_f16(v).to_le_bytes());
8351        }
8352
8353        let mut refw = vec![0f32; rows * cols];
8354        dequant_q1t(&bytes, rows, cols, &mut refw);
8355        // On-grid activations (±1, amax 1) so the int8 SDOT path reconstructs
8356        // x exactly and matches the f32 reference (same trick as the q1 test).
8357        let x: Vec<f32> = (0..cols)
8358            .map(|j| if j % 3 == 0 { 1.0 } else { -1.0 })
8359            .collect();
8360        let mut expect = vec![0f32; rows];
8361        for r in 0..rows {
8362            let mut a = 0.0f32;
8363            for j in 0..cols {
8364                a += refw[r * cols + j] * x[j];
8365            }
8366            expect[r] = a;
8367        }
8368        let tol = |e: f32| 1e-3 * e.abs().max(1e-3);
8369        let mut got = vec![0f32; rows];
8370        q1t_matvec(&bytes, &x, rows, cols, &mut got, None);
8371        for r in 0..rows {
8372            assert!(
8373                (got[r] - expect[r]).abs() < tol(expect[r]),
8374                "row {r}: {} vs {}",
8375                got[r],
8376                expect[r]
8377            );
8378        }
8379        // matmat (b=2, f32 decode path) must agree too.
8380        let x2: Vec<f32> = x.iter().chain(x.iter().map(|v| v)).copied().collect();
8381        let mut gm = vec![0f32; 2 * rows];
8382        q1t_matmat(&bytes, &x2, 2, rows, cols, &mut gm, None);
8383        for r in 0..rows {
8384            assert!((gm[r] - expect[r]).abs() < tol(expect[r]));
8385            assert!((gm[rows + r] - expect[r]).abs() < tol(expect[r]));
8386        }
8387        // Fused pair (q1t_matvec2) must equal two single matvecs
8388        // bit-for-bit: same unpack, same group order, same f32
8389        // accumulation per stream. Distinct x2 exercises both lanes.
8390        let xb: Vec<f32> = (0..cols)
8391            .map(|j| if j % 5 == 0 { -1.0 } else { 1.0 })
8392            .collect();
8393        let (mut s1, mut s2) = (vec![0f32; rows], vec![0f32; rows]);
8394        q1t_matvec(&bytes, &x, rows, cols, &mut s1, None);
8395        q1t_matvec(&bytes, &xb, rows, cols, &mut s2, None);
8396        let (mut p1, mut p2) = (vec![0f32; rows], vec![0f32; rows]);
8397        q1t_matvec2(&bytes, &x, &xb, rows, cols, &mut p1, &mut p2, None);
8398        assert_eq!(p1, s1, "q1t pair lane 1 ≠ single matvec");
8399        assert_eq!(p2, s2, "q1t pair lane 2 ≠ single matvec");
8400    }
8401
8402    /// Pair == 2×matvec with an ODD group count (the kernel's tail
8403    /// group) and no overlay section.
8404    #[test]
8405    fn q1t_matvec2_odd_gpr_matches_singles() {
8406        use cortiq_core::quant::{Q1T_TILE, f32_to_f16, q1t_pack};
8407        let (rows, cols) = (5usize, 96usize); // gpr = 3 → paired + tail
8408        let gpr = cols / GROUP_SIZE;
8409        let mut bytes = Vec::with_capacity(rows * gpr * Q1T_TILE);
8410        for r in 0..rows {
8411            for g in 0..gpr {
8412                bytes.extend_from_slice(&f32_to_f16(0.1 + 0.05 * (r + g) as f32).to_le_bytes());
8413                let mut c = [0u8; 7];
8414                for k in 0..GROUP_SIZE {
8415                    q1t_pack(&mut c, k, ((k * 7 + r * 5 + g * 3) % 3) as u8);
8416                }
8417                bytes.extend_from_slice(&c);
8418            }
8419        }
8420        let x1: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.31).sin()).collect();
8421        let x2: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.17).cos()).collect();
8422        let (mut s1, mut s2) = (vec![0f32; rows], vec![0f32; rows]);
8423        q1t_matvec(&bytes, &x1, rows, cols, &mut s1, None);
8424        q1t_matvec(&bytes, &x2, rows, cols, &mut s2, None);
8425        let (mut p1, mut p2) = (vec![0f32; rows], vec![0f32; rows]);
8426        q1t_matvec2(&bytes, &x1, &x2, rows, cols, &mut p1, &mut p2, None);
8427        assert_eq!(p1, s1, "odd-gpr pair lane 1 ≠ single");
8428        assert_eq!(p2, s2, "odd-gpr pair lane 2 ≠ single");
8429    }
8430
8431    // Speed A/B: fused pair (one unpack, two streams) vs two single
8432    // matvecs. Single-threaded, FFN-sized, min-of paired in-process.
8433    //   cargo test -p cortiq-engine --release q1t_matvec2_speed -- --ignored --nocapture
8434    #[test]
8435    #[ignore]
8436    fn q1t_matvec2_speed() {
8437        use cortiq_core::quant::{Q1T_TILE, f32_to_f16, q1t_pack};
8438        use std::time::Instant;
8439        let (rows, cols) = (8192usize, 4096usize);
8440        let gpr = cols / GROUP_SIZE;
8441        let mut bytes = Vec::with_capacity(rows * gpr * Q1T_TILE);
8442        for r in 0..rows {
8443            for g in 0..gpr {
8444                let s = 0.1 + ((r + g) % 7) as f32 * 0.01;
8445                bytes.extend_from_slice(&f32_to_f16(s).to_le_bytes());
8446                let mut c = [0u8; 7];
8447                for k in 0..GROUP_SIZE {
8448                    q1t_pack(&mut c, k, ((k * 7 + r + g) % 3) as u8);
8449                }
8450                bytes.extend_from_slice(&c);
8451            }
8452        }
8453        let x1: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.31).sin()).collect();
8454        let x2: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.17).cos()).collect();
8455        let (mut s1, mut s2) = (vec![0f32; rows], vec![0f32; rows]);
8456        let (mut p1, mut p2) = (vec![0f32; rows], vec![0f32; rows]);
8457        // Warm both paths once.
8458        q1t_matvec(&bytes, &x1, rows, cols, &mut s1, None);
8459        q1t_matvec2(&bytes, &x1, &x2, rows, cols, &mut p1, &mut p2, None);
8460        let (mut t_pair, mut t_two) = (f64::MAX, f64::MAX);
8461        for _ in 0..8 {
8462            let t0 = Instant::now();
8463            q1t_matvec2(&bytes, &x1, &x2, rows, cols, &mut p1, &mut p2, None);
8464            t_pair = t_pair.min(t0.elapsed().as_secs_f64() * 1000.0);
8465            let t1 = Instant::now();
8466            q1t_matvec(&bytes, &x1, rows, cols, &mut s1, None);
8467            q1t_matvec(&bytes, &x2, rows, cols, &mut s2, None);
8468            t_two = t_two.min(t1.elapsed().as_secs_f64() * 1000.0);
8469        }
8470        assert_eq!(p1, s1);
8471        assert_eq!(p2, s2);
8472        println!("q1t pair {rows}x{cols}: fused {t_pair:.2} ms | two singles {t_two:.2} ms");
8473    }
8474
8475    // Speed A/B: the base-3-division decode (what the packing commit left in
8476    // place) vs the fused sign-LUT matvec. Both single-threaded, same bytes.
8477    //   cargo test -p cortiq-engine q1t_matvec_speed -- --ignored --nocapture
8478    #[test]
8479    #[ignore]
8480    fn q1t_matvec_speed() {
8481        use cortiq_core::quant::{Q1T_TILE, f32_to_f16, q1t_code, q1t_pack};
8482        use std::time::Instant;
8483        let (rows, cols) = (8192usize, 4096usize); // FFN-sized
8484        let gpr = cols / GROUP_SIZE;
8485        let mut bytes = Vec::with_capacity(rows * gpr * Q1T_TILE + 16);
8486        for r in 0..rows {
8487            for g in 0..gpr {
8488                let s = 0.1 + ((r + g) % 7) as f32 * 0.01;
8489                bytes.extend_from_slice(&f32_to_f16(s).to_le_bytes());
8490                let mut c = [0u8; 7];
8491                for k in 0..GROUP_SIZE {
8492                    q1t_pack(&mut c, k, ((k * 7 + r + g) % 3) as u8);
8493                }
8494                bytes.extend_from_slice(&c);
8495            }
8496        }
8497        let (n, stride) = (rows * cols, 40usize); // ~2.5% outliers, per-row overlay
8498        let mut row_ptr = vec![0u32; rows + 1];
8499        let mut idx = 0usize;
8500        while idx < n {
8501            row_ptr[idx / cols + 1] += 1;
8502            idx += stride;
8503        }
8504        for r in 0..rows {
8505            row_ptr[r + 1] += row_ptr[r];
8506        }
8507        for &p in &row_ptr {
8508            bytes.extend_from_slice(&p.to_le_bytes());
8509        }
8510        let mut idx = 0usize;
8511        while idx < n {
8512            bytes.extend_from_slice(&((idx % cols) as u16).to_le_bytes());
8513            bytes.extend_from_slice(&f32_to_f16((idx % 13) as f32 * 0.1 - 0.6).to_le_bytes());
8514            idx += stride;
8515        }
8516        // On-grid ±1 so the fast path's int8 SDOT is exact vs the f32 "slow"
8517        // reference (the A/B is a timing check; values must still agree).
8518        let x: Vec<f32> = (0..cols)
8519            .map(|j| if j % 3 == 0 { 1.0 } else { -1.0 })
8520            .collect();
8521        let (rp_off, ent_off, has_ov) = q1t_overlay(&bytes, rows * gpr * Q1T_TILE, rows);
8522
8523        // "before": base-3 division decode into a buffer, then dot.
8524        let slow = |out: &mut [f32]| {
8525            let mut buf = vec![0f32; cols];
8526            for r in 0..rows {
8527                for g in 0..gpr {
8528                    let off = (r * gpr + g) * Q1T_TILE;
8529                    let s = f16_to_f32(u16::from_le_bytes([bytes[off], bytes[off + 1]]));
8530                    let codes = &bytes[off + 2..off + Q1T_TILE];
8531                    for k in 0..GROUP_SIZE {
8532                        buf[g * GROUP_SIZE + k] = match q1t_code(codes, k) {
8533                            1 => s,
8534                            2 => -s,
8535                            _ => 0.0,
8536                        };
8537                    }
8538                }
8539                out[r] = q1t_row_outlier_correction(&bytes, r, rp_off, ent_off, has_ov, &x)
8540                    + (0..cols).map(|j| buf[j] * x[j]).sum::<f32>();
8541            }
8542        };
8543        let iters = 5;
8544        let mut a = vec![0f32; rows];
8545        slow(&mut a); // warm
8546        let t = Instant::now();
8547        for _ in 0..iters {
8548            slow(&mut a);
8549        }
8550        let slow_ms = t.elapsed().as_secs_f64() * 1e3 / iters as f64;
8551
8552        let mut b = vec![0f32; rows];
8553        q1t_matvec(&bytes, &x, rows, cols, &mut b, None); // warm
8554        let t = Instant::now();
8555        for _ in 0..iters {
8556            q1t_matvec(&bytes, &x, rows, cols, &mut b, None);
8557        }
8558        let fast_ms = t.elapsed().as_secs_f64() * 1e3 / iters as f64;
8559
8560        for r in 0..rows {
8561            assert!((a[r] - b[r]).abs() < 1e-2, "mismatch row {r}");
8562        }
8563        println!(
8564            "q1t matvec {rows}x{cols} (1 thread): div-decode {slow_ms:.2} ms  fused-LUT {fast_ms:.2} ms  => {:.2}x",
8565            slow_ms / fast_ms
8566        );
8567    }
8568}