Skip to main content

qualia_core_db/inference/
ggml_quants.rs

1//! GGML quantization block layout and zero-heap row dequantization.
2//!
3//! Byte strides match `ggml_row_size()` in llama.cpp / ggml. Embedding lookup slices
4//! raw mmap bytes via `fetch_token_embedding`; this module dequantizes into
5//! caller-supplied `&mut [f32]` buffers (no `Vec` in the hot path).
6
7use crate::gguf_sharder::GgufTensorInfo;
8
9/// GGML element-type identifiers used in GGUF tensor-info headers.
10pub const GGML_TYPE_F32: u32 = 0;
11pub const GGML_TYPE_F16: u32 = 1;
12pub const GGML_TYPE_Q4_0: u32 = 2;
13pub const GGML_TYPE_Q5_0: u32 = 6;
14pub const GGML_TYPE_Q8_0: u32 = 8;
15pub const GGML_TYPE_Q4_K: u32 = 12;
16pub const GGML_TYPE_Q6_K: u32 = 14;
17/// Brain float16 (1 sign / 8 exp / 7 mantissa) — used by Gemma-4 and other modern GGUFs
18/// for norms / residual scales alongside Q4_K weights (`ggml_type` enum value 30).
19pub const GGML_TYPE_BF16: u32 = 30;
20/// Qualia conversion-time **SoA Q4_K** (not a stock GGML type).
21///
22/// Per 256-weight superblock (160 bytes, vs 144 AoS):
23/// - `[0..128)`: qs nibbles (same layout as Q4_K)
24/// - `[128..144)`: 8× f16 `d * sub_scale` (pre-expanded)
25/// - `[144..160)`: 8× f16 `dmin * sub_min` (pre-expanded)
26///
27/// Decode GEMV loads scales directly (no 6-bit scale unpack, no shared-header
28/// barriers). Type id 112 is outside the stock ggml enum range.
29pub const GGML_TYPE_Q4_K_SOA: u32 = 112;
30/// Bytes per SoA superblock (256 weights).
31pub const BLOCK_Q4K_SOA_BYTES: usize = 160;
32pub const BLOCK_Q4K_SOA_ELEMS: usize = 256;
33
34/// GGML `block_q6_K` — 210 bytes, 256 weights. Mirrors WGSL `BlockQ6K` layout.
35#[repr(C)]
36#[derive(Copy, Clone, bytemuck::Pod, bytemuck::Zeroable)]
37pub struct BlockQ6K {
38    pub ql: [u8; 128],
39    pub qh: [u8; 64],
40    pub scales: [i8; 16],
41    pub d: u16,
42}
43
44pub const BLOCK_Q6K_BYTES: usize = 210;
45pub const BLOCK_Q6K_ELEMS: usize = 256;
46
47/// Elements per quantization block and packed byte size (from ggml).
48#[derive(Debug, Clone, Copy, PartialEq, Eq)]
49pub struct GgmlBlockLayout {
50    pub block_elems: usize,
51    pub block_bytes: usize,
52}
53
54/// Return block layout for a GGML type, or `None` if unsupported.
55pub fn ggml_block_layout(ggml_type: u32) -> Option<GgmlBlockLayout> {
56    match ggml_type {
57        GGML_TYPE_Q4_0 => Some(GgmlBlockLayout {
58            block_elems: 32,
59            block_bytes: 18,
60        }),
61        GGML_TYPE_Q5_0 => Some(GgmlBlockLayout {
62            block_elems: 32,
63            block_bytes: 22,
64        }),
65        GGML_TYPE_Q8_0 => Some(GgmlBlockLayout {
66            block_elems: 32,
67            block_bytes: 34,
68        }),
69        GGML_TYPE_Q4_K => Some(GgmlBlockLayout {
70            block_elems: 256,
71            block_bytes: 144,
72        }),
73        GGML_TYPE_Q4_K_SOA => Some(GgmlBlockLayout {
74            block_elems: BLOCK_Q4K_SOA_ELEMS,
75            block_bytes: BLOCK_Q4K_SOA_BYTES,
76        }),
77        GGML_TYPE_Q6_K => Some(GgmlBlockLayout {
78            block_elems: 256,
79            block_bytes: 210,
80        }),
81        _ => None,
82    }
83}
84
85/// Packed byte length of one logical row (`n_elems` weights) for the given GGML type.
86pub fn ggml_row_bytes(ggml_type: u32, n_elems: usize) -> Option<usize> {
87    match ggml_type {
88        GGML_TYPE_F32 => Some(n_elems.checked_mul(4)?),
89        GGML_TYPE_F16 | GGML_TYPE_BF16 => Some(n_elems.checked_mul(2)?),
90        _ => {
91            let layout = ggml_block_layout(ggml_type)?;
92            if n_elems == 0 {
93                return Some(0);
94            }
95            Some(n_elems.div_ceil(layout.block_elems) * layout.block_bytes)
96        }
97    }
98}
99
100#[derive(Debug, Clone, Copy, PartialEq, Eq)]
101pub enum GgmlDequantError {
102    UnsupportedType,
103    BufferTooSmall,
104    TruncatedInput,
105}
106
107/// Errors from zero-copy mmap tensor slicing.
108#[derive(Debug, Clone, Copy, PartialEq, Eq)]
109pub enum ExecutionError {
110    TensorNotFound,
111    TokenOutOfRange,
112    UnsupportedType,
113    MmapBounds,
114}
115
116/// Return a zero-copy `&[u8]` slice of the packed embedding row for `token_id`.
117pub fn fetch_token_embedding<'a>(
118    mmap: &'a [u8],
119    tensor_data_start: u64,
120    tensor: &GgufTensorInfo,
121    token_id: u32,
122) -> Result<&'a [u8], ExecutionError> {
123    let n_embd = tensor.dims[0] as usize;
124    let n_vocab = tensor.dims[1] as usize;
125    if n_embd == 0 {
126        return Err(ExecutionError::TensorNotFound);
127    }
128    if token_id as usize >= n_vocab {
129        return Err(ExecutionError::TokenOutOfRange);
130    }
131    let bytes_per_token =
132        ggml_row_bytes(tensor.ggml_type, n_embd).ok_or(ExecutionError::UnsupportedType)?;
133    let start =
134        (tensor_data_start + tensor.byte_offset) as usize + token_id as usize * bytes_per_token;
135    let end = start + bytes_per_token;
136    if end > mmap.len() {
137        return Err(ExecutionError::MmapBounds);
138    }
139    Ok(&mmap[start..end])
140}
141
142/// Total packed byte length of a GGUF tensor from its shape and `ggml_type`.
143pub fn tensor_byte_len(tensor: &GgufTensorInfo) -> Option<usize> {
144    let n0 = tensor.dims[0] as usize;
145    if n0 == 0 {
146        return None;
147    }
148    // BitNet-1.58b ternary blob (STELLAR §A, P64 FFN): a SINGLE per-tensor `[scale f32][packed
149    // trits]` payload over ALL elements — NOT a per-row block format, so `ggml_row_bytes * dims[1]`
150    // does not apply. Compute the whole-tensor packed length directly from the element count.
151    if tensor.ggml_type == crate::ternary::GGML_TYPE_TERNARY_158 {
152        let n_elems = if tensor.n_dims > 1 && tensor.dims[1] > 0 {
153            n0.checked_mul(tensor.dims[1] as usize)?
154        } else {
155            n0
156        };
157        return Some(crate::ternary::ternary_blob_len(n_elems));
158    }
159    let row = ggml_row_bytes(tensor.ggml_type, n0)?;
160    if tensor.n_dims <= 1 || tensor.dims[1] == 0 {
161        Some(row)
162    } else {
163        Some(row.checked_mul(tensor.dims[1] as usize)?)
164    }
165}
166
167/// Zero-copy slice of an entire tensor payload from the mmap.
168pub fn fetch_tensor_bytes<'a>(
169    mmap: &'a [u8],
170    tensor_data_start: u64,
171    tensor: &GgufTensorInfo,
172) -> Result<&'a [u8], ExecutionError> {
173    let len = tensor_byte_len(tensor).ok_or(ExecutionError::UnsupportedType)?;
174    let start = (tensor_data_start + tensor.byte_offset) as usize;
175    let end = start + len;
176    if end > mmap.len() {
177        return Err(ExecutionError::MmapBounds);
178    }
179    Ok(&mmap[start..end])
180}
181
182/// Packed byte width of one logical matrix row (`dims[0]` elements).
183pub fn tensor_row_byte_len(tensor: &GgufTensorInfo) -> Result<usize, ExecutionError> {
184    let n0 = tensor.dims[0] as usize;
185    if n0 == 0 {
186        return Err(ExecutionError::TensorNotFound);
187    }
188    ggml_row_bytes(tensor.ggml_type, n0).ok_or(ExecutionError::UnsupportedType)
189}
190
191/// Zero-copy slice covering vocabulary rows `[row_start, row_start + row_count)`.
192pub fn fetch_tensor_row_range_bytes<'a>(
193    mmap: &'a [u8],
194    tensor_data_start: u64,
195    tensor: &GgufTensorInfo,
196    row_start: usize,
197    row_count: usize,
198) -> Result<&'a [u8], ExecutionError> {
199    let row_bytes = tensor_row_byte_len(tensor)?;
200    let n_rows = if tensor.n_dims > 1 && tensor.dims[1] > 0 {
201        tensor.dims[1] as usize
202    } else {
203        1
204    };
205    if row_start >= n_rows || row_count == 0 {
206        return Err(ExecutionError::TokenOutOfRange);
207    }
208    let rows = row_count.min(n_rows - row_start);
209    let start = (tensor_data_start + tensor.byte_offset) as usize + row_start * row_bytes;
210    let end = start + rows * row_bytes;
211    if end > mmap.len() {
212        return Err(ExecutionError::MmapBounds);
213    }
214    Ok(&mmap[start..end])
215}
216
217/// Dequantize one matrix row (`row` index along `dims[1]`) into `out`.
218pub fn dequant_matrix_row_into(
219    raw: &[u8],
220    info: &GgufTensorInfo,
221    row: usize,
222    out: &mut [f32],
223) -> Result<usize, GgmlDequantError> {
224    let n0 = info.dims[0] as usize;
225    let row_bytes = ggml_row_bytes(info.ggml_type, n0).ok_or(GgmlDequantError::UnsupportedType)?;
226    let start = row
227        .checked_mul(row_bytes)
228        .ok_or(GgmlDequantError::TruncatedInput)?;
229    if start + row_bytes > raw.len() {
230        return Err(GgmlDequantError::TruncatedInput);
231    }
232    dequantize_row_into(&raw[start..start + row_bytes], info.ggml_type, n0, out)
233}
234
235/// Dequantize one embedding row from raw mmap bytes into `out`.
236/// Returns the number of `f32` elements written (≤ `out.len()`).
237pub fn dequantize_row_into(
238    raw: &[u8],
239    ggml_type: u32,
240    n_elems: usize,
241    out: &mut [f32],
242) -> Result<usize, GgmlDequantError> {
243    if out.len() < n_elems {
244        return Err(GgmlDequantError::BufferTooSmall);
245    }
246    match ggml_type {
247        GGML_TYPE_F32 => dequant_f32(raw, n_elems, out),
248        GGML_TYPE_F16 => dequant_f16(raw, n_elems, out),
249        GGML_TYPE_BF16 => dequant_bf16(raw, n_elems, out),
250        GGML_TYPE_Q4_0 => dequant_q4_0(raw, n_elems, out),
251        GGML_TYPE_Q5_0 => dequant_q5_0(raw, n_elems, out),
252        GGML_TYPE_Q8_0 => dequant_q8_0(raw, n_elems, out),
253        GGML_TYPE_Q4_K => dequant_q4_k(raw, n_elems, out),
254        GGML_TYPE_Q4_K_SOA => dequant_q4_k_soa(raw, n_elems, out),
255        GGML_TYPE_Q6_K => dequant_q6_k(raw, n_elems, out),
256        _ => Err(GgmlDequantError::UnsupportedType),
257    }
258}
259
260/// Convert one stock Q4_K superblock (144 B) → SoA superblock (160 B).
261///
262/// Pre-expands the 8 sub-block `(d·scale, dmin·min)` pairs to f16 so the GPU
263/// GEMV never runs `get_scale_min_k4` / shared-header decode.
264#[inline]
265pub fn q4k_block_to_soa(src: &[u8], dst: &mut [u8]) -> Result<(), GgmlDequantError> {
266    if src.len() < 144 || dst.len() < BLOCK_Q4K_SOA_BYTES {
267        return Err(GgmlDequantError::TruncatedInput);
268    }
269    // qs first (same nibble layout as stock Q4_K).
270    dst[..128].copy_from_slice(&src[16..144]);
271    let d = half::f16::from_le_bytes([src[0], src[1]]).to_f32();
272    let dmin = half::f16::from_le_bytes([src[2], src[3]]).to_f32();
273    let scales: [u8; 12] = src[4..16].try_into().unwrap_or([0; 12]);
274    for j in 0..8 {
275        let mut sc = 0u8;
276        let mut m = 0u8;
277        get_scale_min_k4(j, &scales, &mut sc, &mut m);
278        let d_bits = half::f16::from_f32(d * sc as f32).to_le_bytes();
279        let m_bits = half::f16::from_f32(dmin * m as f32).to_le_bytes();
280        let o = 128 + j * 2;
281        dst[o] = d_bits[0];
282        dst[o + 1] = d_bits[1];
283        let om = 144 + j * 2;
284        dst[om] = m_bits[0];
285        dst[om + 1] = m_bits[1];
286    }
287    Ok(())
288}
289
290/// Expand a full Q4_K tensor blob (row-major superblocks) into SoA layout.
291/// `n_row_elems` = dims[0] (weights per row). `n_rows` = dims[1].
292pub fn expand_q4k_tensor_to_soa(
293    raw: &[u8],
294    n_row_elems: usize,
295    n_rows: usize,
296    out: &mut [u8],
297) -> Result<(), GgmlDequantError> {
298    let src_row =
299        ggml_row_bytes(GGML_TYPE_Q4_K, n_row_elems).ok_or(GgmlDequantError::UnsupportedType)?;
300    let dst_row =
301        ggml_row_bytes(GGML_TYPE_Q4_K_SOA, n_row_elems).ok_or(GgmlDequantError::UnsupportedType)?;
302    let need = dst_row
303        .checked_mul(n_rows)
304        .ok_or(GgmlDequantError::TruncatedInput)?;
305    if out.len() < need || raw.len() < src_row.saturating_mul(n_rows) {
306        return Err(GgmlDequantError::TruncatedInput);
307    }
308    let n_blocks = n_row_elems.div_ceil(256);
309    for r in 0..n_rows {
310        let src_base = r * src_row;
311        let dst_base = r * dst_row;
312        for b in 0..n_blocks {
313            let s = src_base + b * 144;
314            let d = dst_base + b * BLOCK_Q4K_SOA_BYTES;
315            if s + 144 > raw.len() || d + BLOCK_Q4K_SOA_BYTES > out.len() {
316                return Err(GgmlDequantError::TruncatedInput);
317            }
318            q4k_block_to_soa(&raw[s..s + 144], &mut out[d..d + BLOCK_Q4K_SOA_BYTES])?;
319        }
320    }
321    Ok(())
322}
323
324fn dequant_q4_k_soa(
325    raw: &[u8],
326    n_elems: usize,
327    out: &mut [f32],
328) -> Result<usize, GgmlDequantError> {
329    let n_blocks = n_elems.div_ceil(BLOCK_Q4K_SOA_ELEMS);
330    if raw.len() < n_blocks * BLOCK_Q4K_SOA_BYTES {
331        return Err(GgmlDequantError::TruncatedInput);
332    }
333    let mut out_idx = 0usize;
334    for b in 0..n_blocks {
335        let block = &raw[b * BLOCK_Q4K_SOA_BYTES..(b + 1) * BLOCK_Q4K_SOA_BYTES];
336        let qs = &block[0..128];
337        // Pre-expanded scales.
338        let mut d_sub = [0f32; 8];
339        let mut m_sub = [0f32; 8];
340        for j in 0..8 {
341            let o = 128 + j * 2;
342            d_sub[j] = half::f16::from_le_bytes([block[o], block[o + 1]]).to_f32();
343            let om = 144 + j * 2;
344            m_sub[j] = half::f16::from_le_bytes([block[om], block[om + 1]]).to_f32();
345        }
346        // Same element order as stock Q4_K: for each of 4 groups of 64:
347        //   32 low nibbles (sub 2g), then 32 high nibbles (sub 2g+1).
348        let mut q_off = 0usize;
349        for g in 0..4 {
350            let sub0 = g * 2;
351            let sub1 = g * 2 + 1;
352            for l in 0..32 {
353                if out_idx >= n_elems {
354                    return Ok(out_idx);
355                }
356                let nib = (qs[q_off + l] & 0xF) as f32;
357                out[out_idx] = d_sub[sub0] * nib - m_sub[sub0];
358                out_idx += 1;
359            }
360            for l in 0..32 {
361                if out_idx >= n_elems {
362                    return Ok(out_idx);
363                }
364                let nib = (qs[q_off + l] >> 4) as f32;
365                out[out_idx] = d_sub[sub1] * nib - m_sub[sub1];
366                out_idx += 1;
367            }
368            q_off += 32;
369        }
370    }
371    Ok(out_idx.min(n_elems))
372}
373
374fn dequant_f32(raw: &[u8], n_elems: usize, out: &mut [f32]) -> Result<usize, GgmlDequantError> {
375    let need = n_elems * 4;
376    if raw.len() < need {
377        return Err(GgmlDequantError::TruncatedInput);
378    }
379    for i in 0..n_elems {
380        out[i] = f32::from_le_bytes(raw[i * 4..i * 4 + 4].try_into().unwrap_or([0; 4]));
381    }
382    Ok(n_elems)
383}
384
385fn dequant_f16(raw: &[u8], n_elems: usize, out: &mut [f32]) -> Result<usize, GgmlDequantError> {
386    let need = n_elems * 2;
387    if raw.len() < need {
388        return Err(GgmlDequantError::TruncatedInput);
389    }
390    for i in 0..n_elems {
391        out[i] =
392            half::f16::from_le_bytes(raw[i * 2..i * 2 + 2].try_into().unwrap_or([0; 2])).to_f32();
393    }
394    Ok(n_elems)
395}
396
397/// BF16 → f32: shift 16-bit code into the high half of an f32 bit pattern.
398fn dequant_bf16(raw: &[u8], n_elems: usize, out: &mut [f32]) -> Result<usize, GgmlDequantError> {
399    let need = n_elems * 2;
400    if raw.len() < need {
401        return Err(GgmlDequantError::TruncatedInput);
402    }
403    for i in 0..n_elems {
404        let bits = u16::from_le_bytes(raw[i * 2..i * 2 + 2].try_into().unwrap_or([0; 2]));
405        out[i] = f32::from_bits((bits as u32) << 16);
406    }
407    Ok(n_elems)
408}
409
410fn dequant_q4_0(raw: &[u8], n_elems: usize, out: &mut [f32]) -> Result<usize, GgmlDequantError> {
411    const BLOCK_ELEMS: usize = 32;
412    const BLOCK_BYTES: usize = 18;
413    let n_blocks = n_elems.div_ceil(BLOCK_ELEMS);
414    if raw.len() < n_blocks * BLOCK_BYTES {
415        return Err(GgmlDequantError::TruncatedInput);
416    }
417    for b in 0..n_blocks {
418        let bs = b * BLOCK_BYTES;
419        let scale = half::f16::from_le_bytes([raw[bs], raw[bs + 1]]).to_f32();
420        let half = BLOCK_ELEMS / 2;
421        for j in 0..half {
422            if b * BLOCK_ELEMS + j >= n_elems {
423                break;
424            }
425            let byte = raw[bs + 2 + j];
426            let x0 = (byte & 0x0F) as i32 - 8;
427            let x1 = ((byte >> 4) & 0x0F) as i32 - 8;
428            out[b * BLOCK_ELEMS + j] = x0 as f32 * scale;
429            let hi = b * BLOCK_ELEMS + j + half;
430            if hi < n_elems {
431                out[hi] = x1 as f32 * scale;
432            }
433        }
434    }
435    Ok(n_elems)
436}
437
438/// `dequantize_row_q5_0` from ggml-quants.c — 5-bit weights, 32 elems per 22-byte block.
439fn dequant_q5_0(raw: &[u8], n_elems: usize, out: &mut [f32]) -> Result<usize, GgmlDequantError> {
440    const BLOCK_ELEMS: usize = 32;
441    const BLOCK_BYTES: usize = 22;
442    let n_blocks = n_elems.div_ceil(BLOCK_ELEMS);
443    if raw.len() < n_blocks * BLOCK_BYTES {
444        return Err(GgmlDequantError::TruncatedInput);
445    }
446    for b in 0..n_blocks {
447        let bs = b * BLOCK_BYTES;
448        let d = half::f16::from_le_bytes([raw[bs], raw[bs + 1]]).to_f32();
449        let qh = u32::from_le_bytes([raw[bs + 2], raw[bs + 3], raw[bs + 4], raw[bs + 5]]);
450        let qs = &raw[bs + 6..bs + 22];
451        let half = BLOCK_ELEMS / 2;
452        for j in 0..half {
453            let xh_0 = ((qh >> j) << 4) & 0x10;
454            let xh_1 = (qh >> (j + 12)) & 0x10;
455            let x0 = ((qs[j] & 0x0F) as u32 | xh_0) as i32 - 16;
456            let x1 = ((qs[j] >> 4) as u32 | xh_1) as i32 - 16;
457            let lo = b * BLOCK_ELEMS + j;
458            if lo < n_elems {
459                out[lo] = x0 as f32 * d;
460            }
461            let hi = lo + half;
462            if hi < n_elems {
463                out[hi] = x1 as f32 * d;
464            }
465        }
466    }
467    Ok(n_elems)
468}
469
470fn dequant_q8_0(raw: &[u8], n_elems: usize, out: &mut [f32]) -> Result<usize, GgmlDequantError> {
471    const BLOCK_ELEMS: usize = 32;
472    const BLOCK_BYTES: usize = 34;
473    let n_blocks = n_elems.div_ceil(BLOCK_ELEMS);
474    if raw.len() < n_blocks * BLOCK_BYTES {
475        return Err(GgmlDequantError::TruncatedInput);
476    }
477    for b in 0..n_blocks {
478        let bs = b * BLOCK_BYTES;
479        let scale = half::f16::from_le_bytes([raw[bs], raw[bs + 1]]).to_f32();
480        let elems = BLOCK_ELEMS.min(n_elems - b * BLOCK_ELEMS);
481        for j in 0..elems {
482            out[b * BLOCK_ELEMS + j] = raw[bs + 2 + j] as i8 as f32 * scale;
483        }
484    }
485    Ok(n_elems)
486}
487
488/// `get_scale_min_k4` from ggml-quants.c — unpack 6-bit scale/min pairs.
489#[inline]
490fn get_scale_min_k4(j: usize, scales: &[u8; 12], sc: &mut u8, m: &mut u8) {
491    if j < 4 {
492        *sc = scales[j] & 63;
493        *m = scales[j + 4] & 63;
494    } else {
495        *sc = (scales[j + 4] & 0xF) | ((scales[j - 4] >> 6) << 4);
496        *m = (scales[j + 4] >> 4) | ((scales[j] >> 6) << 4);
497    }
498}
499
500fn dequant_q4_k(raw: &[u8], n_elems: usize, out: &mut [f32]) -> Result<usize, GgmlDequantError> {
501    const BLOCK_ELEMS: usize = 256;
502    const BLOCK_BYTES: usize = 144;
503    let n_blocks = n_elems.div_ceil(BLOCK_ELEMS);
504    if raw.len() < n_blocks * BLOCK_BYTES {
505        return Err(GgmlDequantError::TruncatedInput);
506    }
507
508    let mut out_idx = 0usize;
509    for b in 0..n_blocks {
510        let block = &raw[b * BLOCK_BYTES..b * BLOCK_BYTES + BLOCK_BYTES];
511        let d = half::f16::from_le_bytes([block[0], block[1]]).to_f32();
512        let dmin = half::f16::from_le_bytes([block[2], block[3]]).to_f32();
513        let scales: [u8; 12] = block[4..16].try_into().unwrap_or([0; 12]);
514        let qs = &block[16..144];
515
516        let block_elems = BLOCK_ELEMS.min(n_elems - b * BLOCK_ELEMS);
517        let mut q_off = 0usize;
518        let mut is = 0usize;
519        let mut j = 0usize;
520        while j < block_elems && out_idx < n_elems {
521            let mut sc = 0u8;
522            let mut m = 0u8;
523            get_scale_min_k4(is, &scales, &mut sc, &mut m);
524            let d1 = d * sc as f32;
525            let m1 = dmin * m as f32;
526            get_scale_min_k4(is + 1, &scales, &mut sc, &mut m);
527            let d2 = d * sc as f32;
528            let m2 = dmin * m as f32;
529
530            for l in 0..32 {
531                if out_idx >= n_elems || j >= block_elems {
532                    break;
533                }
534                out[out_idx] = d1 * (qs[q_off + l] & 0xF) as f32 - m1;
535                out_idx += 1;
536                j += 1;
537            }
538            for l in 0..32 {
539                if out_idx >= n_elems || j >= block_elems {
540                    break;
541                }
542                out[out_idx] = d2 * (qs[q_off + l] >> 4) as f32 - m2;
543                out_idx += 1;
544                j += 1;
545            }
546            q_off += 32;
547            is += 2;
548        }
549    }
550    Ok(out_idx.min(n_elems))
551}
552
553fn dequant_q6_k_block(block: &[u8; 210], out: &mut [f32]) {
554    let blk = bytemuck::from_bytes::<BlockQ6K>(block);
555    let d = half::f16::from_bits(blk.d).to_f32();
556    let mut ql_off = 0usize;
557    let mut qh_off = 0usize;
558    let mut sc_off = 0usize;
559    let mut y_off = 0usize;
560
561    for _ in 0..2 {
562        for l in 0..32 {
563            let is = l / 16;
564            let q1 =
565                ((blk.ql[ql_off + l] & 0xF) | (((blk.qh[qh_off + l] >> 0) & 3) << 4)) as i8 - 32;
566            let q2 = ((blk.ql[ql_off + l + 32] & 0xF) | (((blk.qh[qh_off + l] >> 2) & 3) << 4))
567                as i8
568                - 32;
569            let q3 =
570                ((blk.ql[ql_off + l] >> 4) | (((blk.qh[qh_off + l] >> 4) & 3) << 4)) as i8 - 32;
571            let q4 = ((blk.ql[ql_off + l + 32] >> 4) | (((blk.qh[qh_off + l] >> 6) & 3) << 4))
572                as i8
573                - 32;
574            let sc = &blk.scales[sc_off..sc_off + 8];
575            out[y_off + l] = d * sc[is] as f32 * q1 as f32;
576            out[y_off + l + 32] = d * sc[is + 2] as f32 * q2 as f32;
577            out[y_off + l + 64] = d * sc[is + 4] as f32 * q3 as f32;
578            out[y_off + l + 96] = d * sc[is + 6] as f32 * q4 as f32;
579        }
580        y_off += 128;
581        ql_off += 64;
582        qh_off += 32;
583        sc_off += 8;
584    }
585}
586
587fn dequant_q6_k(raw: &[u8], n_elems: usize, out: &mut [f32]) -> Result<usize, GgmlDequantError> {
588    const BLOCK_ELEMS: usize = 256;
589    const BLOCK_BYTES: usize = 210;
590    let n_blocks = n_elems.div_ceil(BLOCK_ELEMS);
591    if raw.len() < n_blocks * BLOCK_BYTES {
592        return Err(GgmlDequantError::TruncatedInput);
593    }
594
595    let mut written = 0usize;
596    for b in 0..n_blocks {
597        let block: &[u8; 210] = raw[b * BLOCK_BYTES..b * BLOCK_BYTES + BLOCK_BYTES]
598            .try_into()
599            .map_err(|_| GgmlDequantError::TruncatedInput)?;
600        let elems = BLOCK_ELEMS.min(n_elems - written);
601        let mut block_out = [0f32; BLOCK_ELEMS];
602        dequant_q6_k_block(block, &mut block_out);
603        out[written..written + elems].copy_from_slice(&block_out[..elems]);
604        written += elems;
605    }
606    Ok(written)
607}
608
609/// Quantize 256 f32 weights into one Q4_K_SOA superblock (160 bytes).
610///
611/// Layout produced:
612/// - `[0..128)`: nibbles (same as Q4_K: 4 groups × 32 bytes, low=even sub, high=odd sub)
613/// - `[128..144)`: 8 × f16 effective scale (d_sub[j])
614/// - `[144..160)`: 8 × f16 effective min (m_sub[j])
615///
616/// Dequant formula: `w = d_sub[j] * nibble - m_sub[j]`
617/// So: `d_sub[j] = (max_j - min_j) / 15`, `m_sub[j] = -min_j`
618/// And: `nibble = round((w + m_sub[j]) / d_sub[j])` clamped to [0, 15].
619fn quantize_block_f32_to_q4_k_soa(src: &[f32], out: &mut [u8]) {
620    debug_assert!(out.len() >= BLOCK_Q4K_SOA_BYTES);
621    // Zero qs region so unused nibbles are 0.
622    out[..128].fill(0);
623
624    for j in 0..8 {
625        let sub_start = j * 32;
626        let sub_end = (sub_start + 32).min(src.len());
627
628        // Find min/max for this sub-block.
629        let mut min_val = 0.0f32;
630        let mut max_val = 0.0f32;
631        if sub_start < sub_end {
632            min_val = src[sub_start];
633            max_val = src[sub_start];
634            for i in sub_start + 1..sub_end {
635                let w = src[i];
636                if w < min_val {
637                    min_val = w;
638                }
639                if w > max_val {
640                    max_val = w;
641                }
642            }
643        }
644
645        let scale = if max_val > min_val {
646            (max_val - min_val) / 15.0
647        } else {
648            1.0
649        };
650        // Dequant: w = d_sub * nibble - m_sub
651        //   nibble=0 → w = -m_sub = min_val  →  m_sub = -min_val
652        //   nibble=15 → w = d_sub*15 - m_sub = max_val  →  d_sub = (max_val - min_val)/15 = scale
653        let d_sub = half::f16::from_f32(scale);
654        let m_sub = half::f16::from_f32(-min_val);
655        let d_actual = d_sub.to_f32();
656        let m_actual = m_sub.to_f32();
657
658        // Store f16 scale/min in SoA region.
659        let d_bytes = d_sub.to_le_bytes();
660        let m_bytes = m_sub.to_le_bytes();
661        out[128 + j * 2] = d_bytes[0];
662        out[128 + j * 2 + 1] = d_bytes[1];
663        out[144 + j * 2] = m_bytes[0];
664        out[144 + j * 2 + 1] = m_bytes[1];
665
666        // Quantize and pack nibbles.
667        // Group g = j / 2; low nibble if j even, high if j odd.
668        let g = j / 2;
669        let q_off = g * 32;
670        let is_high = j % 2 == 1;
671        for l in 0..32 {
672            let nibble: u8 = if sub_start + l < sub_end {
673                let w = src[sub_start + l];
674                let q = ((w + m_actual) / d_actual).round();
675                q.clamp(0.0, 15.0) as u8
676            } else {
677                0
678            };
679            let bi = q_off + l;
680            if is_high {
681                out[bi] = (out[bi] & 0x0F) | (nibble << 4);
682            } else {
683                out[bi] = (out[bi] & 0xF0) | (nibble & 0x0F);
684            }
685        }
686    }
687}
688
689/// Quantize a full f32 weight matrix to Q4_K_SOA layout.
690///
691/// `src`: row-major f32 weights, `n_row_elems × n_rows` elements.
692/// `out`: destination buffer, must be at least `ggml_row_bytes(Q4_K_SOA, n_row_elems) * n_rows` bytes.
693pub fn quantize_f32_to_q4_k_soa_tensor(
694    src: &[f32],
695    n_row_elems: usize,
696    n_rows: usize,
697    out: &mut [u8],
698) -> Result<(), GgmlDequantError> {
699    let dst_row =
700        ggml_row_bytes(GGML_TYPE_Q4_K_SOA, n_row_elems).ok_or(GgmlDequantError::UnsupportedType)?;
701    let need = dst_row
702        .checked_mul(n_rows)
703        .ok_or(GgmlDequantError::TruncatedInput)?;
704    if out.len() < need {
705        return Err(GgmlDequantError::BufferTooSmall);
706    }
707    if src.len() < n_row_elems * n_rows {
708        return Err(GgmlDequantError::TruncatedInput);
709    }
710
711    let n_blocks = n_row_elems.div_ceil(BLOCK_Q4K_SOA_ELEMS);
712    for r in 0..n_rows {
713        let src_base = r * n_row_elems;
714        let dst_base = r * dst_row;
715        for b in 0..n_blocks {
716            let block_start = src_base + b * BLOCK_Q4K_SOA_ELEMS;
717            let block_end = (block_start + BLOCK_Q4K_SOA_ELEMS).min(src_base + n_row_elems);
718            let block_src = &src[block_start..block_end];
719            let dst_off = dst_base + b * BLOCK_Q4K_SOA_BYTES;
720            quantize_block_f32_to_q4_k_soa(
721                block_src,
722                &mut out[dst_off..dst_off + BLOCK_Q4K_SOA_BYTES],
723            );
724        }
725    }
726    Ok(())
727}
728
729#[cfg(test)]
730mod tests {
731    use super::*;
732
733    #[test]
734    fn q4k_soa_roundtrip_matches_aos_dequant() {
735        // Build one synthetic Q4_K superblock and check SoA dequant ≈ stock dequant.
736        let mut aos = [0u8; 144];
737        // d = 1.0 f16, dmin = 0.1 f16
738        aos[0..2].copy_from_slice(&half::f16::from_f32(1.0).to_le_bytes());
739        aos[2..4].copy_from_slice(&half::f16::from_f32(0.1).to_le_bytes());
740        // scales: all 1 for sc, 0 for m (low 6 bits)
741        for i in 0..4 {
742            aos[4 + i] = 1;
743            aos[8 + i] = 0;
744        }
745        for i in 0..4 {
746            aos[12 + i] = 0;
747        }
748        // qs: low nibble = 3, high nibble = 5
749        for i in 16..144 {
750            aos[i] = 0x53;
751        }
752        let mut soa = [0u8; BLOCK_Q4K_SOA_BYTES];
753        q4k_block_to_soa(&aos, &mut soa).unwrap();
754        let mut out_aos = [0f32; 256];
755        let mut out_soa = [0f32; 256];
756        dequant_q4_k(&aos, 256, &mut out_aos).unwrap();
757        dequant_q4_k_soa(&soa, 256, &mut out_soa).unwrap();
758        for i in 0..256 {
759            let d = (out_aos[i] - out_soa[i]).abs();
760            assert!(
761                d < 1e-3,
762                "elem {i}: aos={} soa={} δ={d}",
763                out_aos[i],
764                out_soa[i]
765            );
766        }
767        assert_eq!(ggml_row_bytes(GGML_TYPE_Q4_K_SOA, 256), Some(160));
768        assert_eq!(ggml_row_bytes(GGML_TYPE_Q4_K_SOA, 512), Some(320));
769    }
770
771    #[test]
772    fn quantize_f32_to_soa_roundtrip() {
773        // Synthetic weights: 256 values with varied ranges per sub-block.
774        let mut src = [0f32; 256];
775        for j in 0..8 {
776            let base = j as f32 * 0.1 - 0.4;
777            let amp = (j as f32 + 1.0) * 0.05;
778            for l in 0..32 {
779                src[j * 32 + l] = base + amp * (l as f32 / 31.0);
780            }
781        }
782        let mut soa = [0u8; BLOCK_Q4K_SOA_BYTES];
783        quantize_block_f32_to_q4_k_soa(&src, &mut soa);
784        let mut deq = [0f32; 256];
785        dequant_q4_k_soa(&soa, 256, &mut deq).unwrap();
786        // Q4_K has 4-bit precision (~1/15 of the sub-block range). Check that
787        // the round-trip error is within the quantization step size.
788        for j in 0..8 {
789            let sub_min = src[j * 32..j * 32 + 32]
790                .iter()
791                .cloned()
792                .fold(f32::MAX, f32::min);
793            let sub_max = src[j * 32..j * 32 + 32]
794                .iter()
795                .cloned()
796                .fold(f32::MIN, f32::max);
797            let step = (sub_max - sub_min) / 15.0;
798            for l in 0..32 {
799                let i = j * 32 + l;
800                let err = (src[i] - deq[i]).abs();
801                assert!(
802                    err <= step + 1e-4,
803                    "elem {i}: src={} deq={} err={} step={}",
804                    src[i],
805                    deq[i],
806                    err,
807                    step
808                );
809            }
810        }
811    }
812
813    #[test]
814    fn quantize_f32_to_soa_tensor_roundtrip() {
815        // 2 rows × 512 elements (2 blocks per row).
816        let n_row_elems = 512;
817        let n_rows = 2;
818        let mut src = vec![0f32; n_row_elems * n_rows];
819        for i in 0..src.len() {
820            src[i] = ((i as f32) * 0.01 - 5.0).sin() * 0.5;
821        }
822        let dst_row = ggml_row_bytes(GGML_TYPE_Q4_K_SOA, n_row_elems).unwrap();
823        let mut out = vec![0u8; dst_row * n_rows];
824        quantize_f32_to_q4_k_soa_tensor(&src, n_row_elems, n_rows, &mut out).unwrap();
825        let mut deq = vec![0f32; n_row_elems * n_rows];
826        for r in 0..n_rows {
827            let off = r * dst_row;
828            dequant_q4_k_soa(
829                &out[off..off + dst_row],
830                n_row_elems,
831                &mut deq[r * n_row_elems..(r + 1) * n_row_elems],
832            )
833            .unwrap();
834        }
835        // Check average error is reasonable for 4-bit quantization.
836        let mut total_err = 0.0;
837        for i in 0..src.len() {
838            total_err += (src[i] - deq[i]).abs();
839        }
840        let avg_err = total_err / src.len() as f32;
841        assert!(avg_err < 0.05, "avg quantization error too high: {avg_err}");
842    }
843
844    #[test]
845    fn bf16_row_bytes_and_dequant() {
846        // Gemma-4 norms/scales use ggml_type 30 (BF16): 2 bytes/elem.
847        assert_eq!(ggml_row_bytes(GGML_TYPE_BF16, 1024), Some(2048));
848        // 1.0_bf16 = 0x3F80, -2.0_bf16 = 0xC000
849        let raw: [u8; 4] = [0x80, 0x3F, 0x00, 0xC0];
850        let mut out = [0.0f32; 2];
851        assert_eq!(
852            dequantize_row_into(&raw, GGML_TYPE_BF16, 2, &mut out),
853            Ok(2)
854        );
855        assert!((out[0] - 1.0).abs() < 1e-6, "got {}", out[0]);
856        assert!((out[1] + 2.0).abs() < 1e-6, "got {}", out[1]);
857        let info = GgufTensorInfo {
858            dims: [128, 1, 0, 0],
859            n_dims: 2,
860            ggml_type: GGML_TYPE_BF16,
861            byte_offset: 0,
862        };
863        assert_eq!(tensor_byte_len(&info), Some(256));
864    }
865
866    #[test]
867    fn q4_0_row_bytes_stride() {
868        // hidden_dim=4096 → (4096/32)*18 = 2304
869        assert_eq!(ggml_row_bytes(GGML_TYPE_Q4_0, 4096), Some(2304));
870    }
871
872    #[test]
873    fn q5_0_row_bytes_stride() {
874        // SmolLM2 attn_k row: hidden_dim=960 → (960/32)*22 = 660
875        assert_eq!(ggml_row_bytes(GGML_TYPE_Q5_0, 960), Some(660));
876    }
877
878    #[test]
879    fn ternary_tensor_byte_len_is_whole_blob_not_row_strided() {
880        // A1b inc 2a: a ternary FFN tensor is ONE `[scale f32][5-trits/byte]` blob over all
881        // dims[0]*dims[1] elements (per-tensor scale), so `tensor_byte_len`/`fetch_tensor_bytes`
882        // must return the whole-blob length — NOT the (None) row-based path that choked before.
883        let info = GgufTensorInfo {
884            dims: [960, 2560, 0, 0], // SmolLM2 ffn_gate-shape
885            n_dims: 2,
886            ggml_type: crate::ternary::GGML_TYPE_TERNARY_158,
887            byte_offset: 0,
888        };
889        let n = 960 * 2560;
890        assert_eq!(
891            tensor_byte_len(&info),
892            Some(crate::ternary::ternary_blob_len(n))
893        );
894
895        // and fetch returns exactly that slice from a buffer holding a real ternary blob.
896        let weights: Vec<f32> = (0..n).map(|i| (i as f32 * 0.001).sin()).collect();
897        let blob = crate::ternary::ternary_blob(&weights);
898        assert_eq!(blob.len(), crate::ternary::ternary_blob_len(n));
899        let got = fetch_tensor_bytes(&blob, 0, &info).expect("ternary fetch must slice");
900        assert_eq!(got.len(), blob.len());
901        assert_eq!(&got[..4], &blob[..4]); // scale preserved at the front
902    }
903
904    #[test]
905    fn q5_0_dequant_matches_gguf_smollm2_row0() {
906        let path = std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
907            .join("../../docs/models/SmolLM2-360M-Instruct-Q4_K_M.gguf");
908        if !path.exists() {
909            return;
910        }
911        let mmap = std::fs::read(&path).expect("read gguf");
912        let index = crate::gguf_sharder::GgufTensorIndex::from_gguf(&mmap);
913        let info = index.get_layer_tensors(0).attn_k.expect("blk.0.attn_k");
914        let raw = fetch_tensor_bytes(&mmap, index.tensor_data_start, &info).expect("fetch attn_k");
915        let row_bytes = ggml_row_bytes(GGML_TYPE_Q5_0, info.dims[0] as usize).unwrap();
916        assert_eq!(row_bytes, 660);
917        let mut out = [0f32; 960];
918        dequantize_row_into(&raw[..row_bytes], GGML_TYPE_Q5_0, 960, &mut out).unwrap();
919        // Reference: gguf-py dequantize_row_q5_0 on blk.0.attn_k.weight row 0
920        let expected = [
921            -0.0f32,
922            -0.02502441,
923            -0.10009766,
924            -0.07507324,
925            -0.05004883,
926            -0.3503418,
927            -0.0,
928            0.05004883,
929            -0.07507324,
930            0.1751709,
931        ];
932        for (i, &exp) in expected.iter().enumerate() {
933            assert!(
934                (out[i] - exp).abs() < 1e-5,
935                "elem {i}: got {} expected {}",
936                out[i],
937                exp
938            );
939        }
940    }
941
942    #[test]
943    fn smollm_q4km_tensor_types_and_q4k_row0() {
944        let path = std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
945            .join("../../docs/models/SmolLM2-360M-Instruct-Q4_K_M.gguf");
946        if !path.exists() {
947            return;
948        }
949        let mmap = std::fs::read(&path).expect("read gguf");
950        let index = crate::gguf_sharder::GgufTensorIndex::from_gguf(&mmap);
951        let emb = index.token_embd_info().expect("token_embd");
952        println!("token_embd type={} dims={:?}", emb.ggml_type, emb.dims);
953        let lt = index.get_layer_tensors(0);
954        for (name, info) in [
955            ("attn_q", lt.attn_q),
956            ("attn_k", lt.attn_k),
957            ("attn_v", lt.attn_v),
958            ("ffn_gate", lt.ffn_gate),
959            ("ffn_down", lt.ffn_down),
960        ] {
961            if let Some(i) = info {
962                println!("{name} type={} dims={:?}", i.ggml_type, i.dims);
963            }
964        }
965        // token 504 = "The" in naked prompt
966        let mut out = [0f32; 960];
967        let n = index.dequantize_token_embedding_into(&mmap, 504, &mut out);
968        assert_eq!(n, 960);
969        println!(
970            "token504 first10: {:?}",
971            &out[..10]
972                .iter()
973                .map(|v| format!("{v:.6}"))
974                .collect::<Vec<_>>()
975        );
976        assert!(out.iter().any(|&v| v.is_finite() && v != 0.0));
977    }
978
979    #[test]
980    fn q4_k_row_bytes_stride() {
981        // hidden_dim=2560 → (2560/256)*144 = 1440
982        assert_eq!(ggml_row_bytes(GGML_TYPE_Q4_K, 2560), Some(1440));
983    }
984
985    #[test]
986    fn q6_k_row_bytes_stride() {
987        // Gemma 4B token_embd: hidden_dim=2560 → (2560/256)*210 = 2100
988        assert_eq!(ggml_row_bytes(GGML_TYPE_Q6_K, 2560), Some(2100));
989    }
990
991    #[test]
992    fn q6_k_dequant_matches_gguf_smollm2_ffn_down_row0() {
993        let path = std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
994            .join("../../docs/models/SmolLM2-360M-Instruct-Q4_K_M.gguf");
995        if !path.exists() {
996            return;
997        }
998        let mmap = std::fs::read(&path).expect("read gguf");
999        let index = crate::gguf_sharder::GgufTensorIndex::from_gguf(&mmap);
1000        let info = index.get_layer_tensors(0).ffn_down.expect("blk.0.ffn_down");
1001        assert_eq!(info.ggml_type, GGML_TYPE_Q6_K);
1002        let raw =
1003            fetch_tensor_bytes(&mmap, index.tensor_data_start, &info).expect("fetch ffn_down");
1004        let row_bytes = ggml_row_bytes(GGML_TYPE_Q6_K, info.dims[0] as usize).unwrap();
1005        assert_eq!(row_bytes, 2100);
1006        let mut out = [0f32; 2560];
1007        dequantize_row_into(&raw[..row_bytes], GGML_TYPE_Q6_K, 2560, &mut out).unwrap();
1008        // Reference: llama.cpp dequantize_row_q6_K (signed int8 scales) on row 0
1009        let expected = [
1010            -0.11712998f32,
1011            -0.16535997,
1012            -0.0,
1013            0.15157998,
1014            -0.08956999,
1015            -0.02067,
1016            0.04133999,
1017            0.03445,
1018            0.17913999,
1019            0.22047997,
1020        ];
1021        for (i, &exp) in expected.iter().enumerate() {
1022            assert!(
1023                (out[i] - exp).abs() < 1e-5,
1024                "elem {i}: got {} expected {}",
1025                out[i],
1026                exp
1027            );
1028        }
1029        assert!(out.iter().any(|&v| v.is_finite() && v != 0.0));
1030    }
1031
1032    #[test]
1033    fn block_q6k_layout_matches_ggml() {
1034        assert_eq!(std::mem::size_of::<BlockQ6K>(), BLOCK_Q6K_BYTES);
1035        assert_eq!(std::mem::align_of::<BlockQ6K>(), 2);
1036    }
1037
1038    #[test]
1039    fn fetch_row_range_bounds() {
1040        use crate::gguf_sharder::GgufTensorInfo;
1041        let info = GgufTensorInfo {
1042            dims: [2560, 100, 0, 0],
1043            n_dims: 2,
1044            ggml_type: GGML_TYPE_Q6_K,
1045            byte_offset: 0,
1046        };
1047        let row = tensor_row_byte_len(&info).unwrap();
1048        let total = tensor_byte_len(&info).unwrap();
1049        assert_eq!(total, row * 100);
1050        let fake = vec![0u8; total];
1051        let chunk = fetch_tensor_row_range_bytes(&fake, 0, &info, 10, 8).unwrap();
1052        assert_eq!(chunk.len(), row * 8);
1053    }
1054
1055    #[test]
1056    fn q8_0_block_roundtrip() {
1057        let mut block = [0u8; 34];
1058        block[0] = 0x00;
1059        block[1] = 0x3C; // f16 1.0
1060        for i in 0..32 {
1061            block[2 + i] = (i + 1) as u8;
1062        }
1063        let mut out = [0f32; 32];
1064        dequant_q8_0(&block, 32, &mut out).unwrap();
1065        assert!((out[0] - 1.0).abs() < 0.01);
1066        assert!((out[31] - 32.0).abs() < 0.01);
1067    }
1068}