aprender-serve 0.70.2

Pure Rust ML inference engine built from scratch - model serving for GGUF and safetensors
//! #3953: IQ2_S block geometry and oracle strength, asserted as PRECONDITIONS of
//! the GPU kernel rather than recorded in a comment.
//!
//! Two kinds of precondition, and the second is the one a port usually skips.
//!
//! GEOMETRY. `iq_dispatch`'s header records the IQ4_NL failure: a type wired
//! against the wrong element count "does not produce an error — it produces
//! `blocks_per_row = in_dim/256`, so a row that occupies 8 blocks is read as 1".
//! With codebooks in the decode path a layout error and a decode error are
//! indistinguishable by eye, so the numbers are pinned before any value is
//! compared.
//!
//! ORACLE STRENGTH. The device A/B will plant faults — sign bits ignored, the
//! `qh` high bits dropped, the wrong scale nibble — and require each to go RED.
//! A planted fault proves nothing if the oracle's data never exercises the
//! mechanism it breaks. If the generated bytes never produced a grid index
//! >= 256, a kernel that drops `qh` entirely would agree with the CPU exactly,
//! and the "qh dropped" fault would pass for the wrong reason. So the data is
//! shown to exercise each mechanism HERE, before the A/B relies on it.

use super::iq2_s::{GGML_TYPE_IQ2_S, IQ2_S_BLOCK_BYTES, IQ2_S_BLOCK_ELEMS};
use super::iq_dispatch::{iq_block_bytes, iq_block_elems};

/// Byte offsets of the five fields in one 82-byte super-block. The kernel
/// hardcodes these; the decoder slices at the same numbers.
pub(crate) const IQ2_S_OFF_QS: usize = 2;
pub(crate) const IQ2_S_OFF_SIGNS: usize = 34;
pub(crate) const IQ2_S_OFF_QH: usize = 66;
pub(crate) const IQ2_S_OFF_SCALES: usize = 74;

/// 2.5 bits/weight: 82 bytes per 256 values.
///   2  f16 super-block scale `d`
///  32  `qs`     low 8 bits of each of 32 grid indices
///  32  `signs`  one sign byte per 8 outputs
///   8  `qh`     high 2 bits of four indices, per sub-block
///   8  `scales` two 4-bit scales per sub-block
#[test]
fn iq2_s_is_256_elements_in_82_bytes_3953() {
    assert_eq!(IQ2_S_BLOCK_ELEMS, 256);
    assert_eq!(IQ2_S_BLOCK_BYTES, 82);
    assert_eq!(2 + 32 + 32 + 8 + 8, IQ2_S_BLOCK_BYTES);
    let bpw = (IQ2_S_BLOCK_BYTES * 8) as f64 / IQ2_S_BLOCK_ELEMS as f64;
    assert!(
        (bpw - 2.5625).abs() < 1e-9,
        "82*8/256 = 2.5625 bpw; got {bpw}"
    );
}

/// The field offsets tile the block exactly, in order, with no gap or overlap.
#[test]
fn the_field_offsets_tile_the_block_3953() {
    assert_eq!(IQ2_S_OFF_QS, 2, "qs follows the 2-byte f16 d");
    assert_eq!(IQ2_S_OFF_SIGNS, IQ2_S_OFF_QS + 32);
    assert_eq!(IQ2_S_OFF_QH, IQ2_S_OFF_SIGNS + 32);
    assert_eq!(IQ2_S_OFF_SCALES, IQ2_S_OFF_QH + 8);
    assert_eq!(
        IQ2_S_OFF_SCALES + 8,
        IQ2_S_BLOCK_BYTES,
        "scales end the block"
    );
}

/// The dispatch must agree with the module's own constants. The kernel reads the
/// dispatch, not the constants.
#[test]
fn the_dispatch_agrees_with_the_block_constants_3953() {
    assert_eq!(iq_block_elems(GGML_TYPE_IQ2_S), Some(IQ2_S_BLOCK_ELEMS));
    assert_eq!(iq_block_bytes(GGML_TYPE_IQ2_S), Some(IQ2_S_BLOCK_BYTES));
}

/// The row stride the kernel indexes with: a row of `k` values is `ceil(k/256)`
/// blocks of 82 bytes.
#[test]
fn the_row_stride_is_blocks_per_row_times_82_3953() {
    for k in [256usize, 512, 2048, 5120] {
        assert_eq!(
            k.div_ceil(IQ2_S_BLOCK_ELEMS) * IQ2_S_BLOCK_BYTES,
            (k / 256) * 82
        );
    }
}

/// Deterministic IQ2_S weights for `n` rows of `k` values, every block's f16
/// scale pinned to an exactly representable 1.0 so any GPU/CPU disagreement is
/// about INDEXING rather than rounding. Shared with the device A/B so the oracle
/// and the A/B cannot drift apart.
pub(crate) fn iq2_s_weights(n: usize, k: usize) -> Vec<u8> {
    let blocks = n * k.div_ceil(IQ2_S_BLOCK_ELEMS);
    let mut data = Vec::with_capacity(blocks * IQ2_S_BLOCK_BYTES);
    let mut state: u32 = 0x9E37_79B9;
    for _ in 0..blocks {
        data.push(0x00);
        data.push(0x3c); // f16 1.0
        for _ in 2..IQ2_S_BLOCK_BYTES {
            state = state.wrapping_mul(1_664_525).wrapping_add(1_013_904_223);
            data.push((state >> 24) as u8);
        }
    }
    data
}

/// ORACLE STRENGTH, per mechanism. Each assertion is the precondition for one
/// planted fault in the device A/B to be able to go red.
#[test]
fn the_oracle_data_exercises_every_mechanism_a_fault_targets_3953() {
    let (k, n) = (512usize, 48usize);
    let w = iq2_s_weights(n, k);
    let blocks = w.chunks_exact(IQ2_S_BLOCK_BYTES).count();

    // qh DROPPED: needs grid indices >= 256, i.e. a non-zero 2-bit high part.
    let mut high_indices = 0usize;
    // WRONG SCALE NIBBLE (l&1 vs l>>1): needs sub-blocks whose two nibbles differ.
    let mut unequal_nibbles = 0usize;
    // SIGNS IGNORED: needs sign bytes with bits set.
    let mut set_sign_bits = 0usize;
    for b in w.chunks_exact(IQ2_S_BLOCK_BYTES) {
        for ib in 0..8 {
            let qh = b[IQ2_S_OFF_QH + ib];
            for l in 0..4 {
                if (qh >> (2 * l)) & 3 != 0 {
                    high_indices += 1;
                }
                set_sign_bits += b[IQ2_S_OFF_SIGNS + 4 * ib + l].count_ones() as usize;
            }
            let sc = b[IQ2_S_OFF_SCALES + ib];
            if sc & 0xf != sc >> 4 {
                unequal_nibbles += 1;
            }
        }
    }
    let lanes = blocks * 32;
    assert!(
        high_indices >= lanes / 2,
        "only {high_indices}/{lanes} grid indices use the qh high bits; a kernel that \
         drops qh would agree with the oracle and that fault could not go red"
    );
    assert!(
        unequal_nibbles >= (blocks * 8) / 2,
        "only {unequal_nibbles}/{} sub-blocks have two different scale nibbles; picking \
         the wrong nibble would be invisible",
        blocks * 8
    );
    assert!(
        set_sign_bits >= (lanes * 8) / 4,
        "too few sign bits set ({set_sign_bits}) for an ignored-signs fault to show"
    );
}

/// The decoder is the A/B's oracle, so it must produce values at all — an oracle
/// returning zeros would agree with an all-zero kernel perfectly.
#[test]
fn the_cpu_decoder_produces_a_nondegenerate_block_3953() {
    let block = iq2_s_weights(1, 256);
    let mut out = vec![0.0f32; IQ2_S_BLOCK_ELEMS];
    super::iq2_s::dequantize_iq2_s_block(&block, &mut out);
    let nonzero = out.iter().filter(|v| v.abs() > 1e-9).count();
    assert!(
        nonzero >= IQ2_S_BLOCK_ELEMS / 2,
        "only {nonzero} non-zero decoded values"
    );
    assert!(
        out.iter().any(|v| *v < 0.0) && out.iter().any(|v| *v > 0.0),
        "the sign bytes must produce both signs, or the sign path is untested"
    );
    assert!(out.iter().all(|v| v.is_finite()));
}

/// The oracle the device A/B compares against, end to end on the CPU, at a shape
/// where the row stride spans two super-blocks.
#[test]
fn iq_parallel_matvec_decodes_iq2_s_end_to_end_3953() {
    let (k, n) = (512usize, 48usize);
    let weights = iq2_s_weights(n, k);
    assert_eq!(
        weights.len(),
        n * (k / IQ2_S_BLOCK_ELEMS) * IQ2_S_BLOCK_BYTES
    );
    let input: Vec<f32> = (0..k).map(|i| ((i % 13) as f32) - 6.0).collect();
    let got = super::iq_dispatch::iq_parallel_matvec(GGML_TYPE_IQ2_S, &weights, &input, k, n)
        .expect("IQ2_S is dispatched by iq_parallel_matvec — this is the A/B's oracle");
    assert_eq!(got.len(), n);
    assert!(got.iter().all(|v| v.is_finite()));
    assert!(
        got.iter().filter(|v| v.abs() > 1e-6).count() >= n / 2,
        "fewer than half the oracle rows are non-zero"
    );
    let distinct = got
        .iter()
        .map(|v| v.to_bits())
        .collect::<std::collections::BTreeSet<_>>();
    assert!(
        distinct.len() >= n / 2,
        "only {} of {n} oracle rows are distinct; a kernel ignoring the row index would \
         still match",
        distinct.len()
    );
}

/// A row whose last block is padding: `k` not a multiple of 256.
#[test]
fn the_oracle_handles_a_row_whose_last_block_is_padding_3953() {
    let (k, n) = (300usize, 8usize);
    assert_eq!(k.div_ceil(IQ2_S_BLOCK_ELEMS), 2);
    let weights = iq2_s_weights(n, k);
    let input: Vec<f32> = (0..k).map(|i| ((i % 7) as f32) - 3.0).collect();
    let got = super::iq_dispatch::iq_parallel_matvec(GGML_TYPE_IQ2_S, &weights, &input, k, n)
        .expect("a padded tail must decode, not error");
    assert!(got.iter().all(|v| v.is_finite()));
    assert!(got.iter().filter(|v| v.abs() > 1e-6).count() >= n / 2);
}

/// The kernel reads a sign as "bit j of the sign byte". The reference reads it as
/// `sign_byte & KMASK_IQ2XS[j]`. Those are the same only if `KMASK_IQ2XS[j] == 1 << j`;
/// #3884's receipt verified that in Python, and here it is a test, so a change to the
/// table cannot silently desynchronise the kernel from its oracle.
#[test]
fn kmask_is_bit_j_so_the_kernel_may_read_the_sign_as_bit_j_3953() {
    for (j, m) in super::iq_grids::KMASK_IQ2XS.iter().enumerate() {
        assert_eq!(u32::from(*m), 1u32 << j, "KMASK_IQ2XS[{j}] must be bit {j}");
    }
}

/// ONE-SHOT PROOF of the IQ2_S decoder against gguf-py on every real type-22 tensor, NOT a
/// regular test (hence `#[ignore]`). Same method as #3960's Q2_K proof: `<name>.bin` holds
/// raw bytes written by gguf-py's own reader, so no offset convention can enter; this writes
/// `<name>.apr.f32` for an element-wise comparison against gguf-py's `<name>.ref.f32`.
/// The IQ2_S decoder's own comment cites 50 random blocks (PMAT-3477); this covers the
/// actual tensors the kernel's oracle is used on.
#[test]
#[ignore = "one-shot probe: needs IQ2S_PROBE_DIR from the gguf-py dump"]
fn probe_dump_iq2_s_decodes_for_gguf_py_comparison_3953() {
    let Ok(dir) = std::env::var("IQ2S_PROBE_DIR") else {
        panic!("PROBE: IQ2S_PROBE_DIR is unset -- this probe did NOT run");
    };
    let mut done = 0usize;
    for entry in std::fs::read_dir(&dir).expect("probe dir") {
        let path = entry.expect("dir entry").path();
        if path.extension().and_then(|e| e.to_str()) != Some("bin") {
            continue;
        }
        let bytes = std::fs::read(&path).expect("read bin");
        let out = super::iq2_s::dequantize_iq2_s(&bytes).expect("dequantize_iq2_s");
        let mut buf = Vec::with_capacity(out.len() * 4);
        for v in &out {
            buf.extend_from_slice(&v.to_le_bytes());
        }
        std::fs::write(path.with_extension("apr.f32"), &buf).expect("write apr.f32");
        eprintln!(
            "PROBE {}: {} bytes -> {} values",
            path.display(),
            bytes.len(),
            out.len()
        );
        done += 1;
    }
    assert!(
        done > 0,
        "PROBE: no .bin files in {dir} -- nothing was compared"
    );
}