Skip to content
Merged
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
140 changes: 122 additions & 18 deletions crates/onnx-runtime-ep-cpu/src/kernels/matmul_nbits.rs
Original file line number Diff line number Diff line change
Expand Up @@ -423,6 +423,15 @@ static N16_SDOT_M1_TEST_CALLS: AtomicUsize = AtomicUsize::new(0);
static KAI_SDOT_M1_TEST_CALLS: AtomicUsize = AtomicUsize::new(0);
#[cfg(all(test, feature = "mlas"))]
static MLAS_SQNBIT_TEST_CALLS: AtomicUsize = AtomicUsize::new(0);
/// Positive-proof counters for the zero-copy borrowed int4 decode path (#979).
/// Split by symmetry so a test can assert the *symmetric* branch is the one that
/// executed, not merely that some path avoided the resident `weight_nk` f32
/// cache. Incremented once per `execute` that routes into
/// [`borrowed_affine_int4_matmul`].
#[cfg(test)]
static BORROWED_INT4_SYMMETRIC_TEST_CALLS: AtomicUsize = AtomicUsize::new(0);
#[cfg(test)]
static BORROWED_INT4_ASYMMETRIC_TEST_CALLS: AtomicUsize = AtomicUsize::new(0);
/// Standard ONNX row-major packed NBits weight with affine block metadata.
///
/// This preserves the wire layout instead of expanding it to f32. Direct
Expand Down Expand Up @@ -780,17 +789,29 @@ impl Kernel for MatMulNBitsKernel {
&& self.accuracy_level == 0
&& !self.weight_prepacked
&& group_indices.is_none()
&& let Some(zero_points) = zero_points
&& let Some(packed) = contiguous_host_slice::<u8>(&inputs[1])
&& let Some(scales) = borrowed_scales(&inputs[2])
&& let Some(zero_points) = contiguous_host_slice::<u8>(zero_points)
&& let Some(borrowed_zero_points) = borrow_optional_int4_zero_points(zero_points)
{
// Gate on *symmetry* explicitly, not on "a zero_points input happens
// to exist". Symmetric int4 (no zero_points) has the implicit
// midpoint 8 and is mathematically simpler than the asymmetric case,
// yet it used to fall past this zero-copy path all the way to the
// resident f32 `weight_nk` cache (~8x the file size in RAM). It now
// borrows the packed int4 in place like the asymmetric case, using
// `None` zero points to mean the implicit midpoint (see #979).
#[cfg(test)]
if borrowed_zero_points.is_none() {
BORROWED_INT4_SYMMETRIC_TEST_CALLS.fetch_add(1, Ordering::Relaxed);
} else {
BORROWED_INT4_ASYMMETRIC_TEST_CALLS.fetch_add(1, Ordering::Relaxed);
}
with_decode_pool(|| {
borrowed_affine_int4_matmul(
&activations,
packed,
scales,
zero_points,
borrowed_zero_points,
bias.as_deref(),
result,
m,
Expand Down Expand Up @@ -5144,6 +5165,24 @@ fn packed_nbits_output_row(
}
}

/// Borrow the optional int4 zero-point tensor for the zero-copy decode path.
///
/// Returns `Some(None)` for a symmetric model (no zero_points input) — the
/// borrowed kernel then uses the implicit midpoint 8. Returns `Some(Some(zp))`
/// when an asymmetric uint8 zero_points input is present and host-contiguous.
/// Returns `None` only when a zero_points input exists but cannot be borrowed
/// in place, so the caller must fall through to another path. Gating on this
/// (rather than on "a zero_points input happens to exist") is what lets
/// symmetric int4 take the borrowed path instead of the resident f32 cache.
fn borrow_optional_int4_zero_points<'a>(
zero_points: Option<&TensorView<'a>>,
) -> Option<Option<&'a [u8]>> {
match zero_points {
None => Some(None),
Some(view) => contiguous_host_slice::<u8>(view).map(Some),
}
}

fn contiguous_host_slice<'a, T>(view: &TensorView<'a>) -> Option<&'a [T]> {
if !view.device.is_host_accessible() || !view.is_contiguous() {
return None;
Expand All @@ -5167,7 +5206,7 @@ fn borrowed_affine_int4_matmul(
activations: &[f32],
packed: &[u8],
scales: BorrowedScales<'_>,
zero_points: &[u8],
zero_points: Option<&[u8]>,
bias: Option<&[f32]>,
result: &mut [f32],
m: usize,
Expand Down Expand Up @@ -5214,16 +5253,18 @@ fn borrowed_affine_int4_matmul(
let output_index = output_start + offset;
let packed_row =
&packed[output_index * packed_row_size..(output_index + 1) * packed_row_size];
let zp_row = &zero_points
[output_index * zero_point_row_size..(output_index + 1) * zero_point_row_size];
let zp_row = zero_points.map(|zp| {
&zp[output_index * zero_point_row_size
..(output_index + 1) * zero_point_row_size]
});
let mut sum = bias.map_or(0.0, |values| values[output_index]);
for block in 0..block_count {
let depth_start = block * block_size;
let valid = k.saturating_sub(depth_start).min(block_size);
let block_values = &packed_row[block * layout.packed_block_size()
..(block + 1) * layout.packed_block_size()];
let scale = scales.get(output_index * block_count + block);
let zero_point = layout.zero_point(Some(zp_row), block) as f32;
let zero_point = layout.zero_point(zp_row, block) as f32;
let mut dot;
#[cfg(target_arch = "aarch64")]
if valid == 32 && block_size == 32 {
Expand Down Expand Up @@ -5262,7 +5303,7 @@ fn borrowed_affine_int4_matmul_m1_neon_dot(
activation: &[f32],
packed: &[u8],
scales: &BorrowedScales<'_>,
zero_points: &[u8],
zero_points: Option<&[u8]>,
bias: Option<&[f32]>,
result: &mut [f32],
k: usize,
Expand All @@ -5287,16 +5328,17 @@ fn borrowed_affine_int4_matmul_m1_neon_dot(
let output_index = output_start + offset;
let packed_row =
&packed[output_index * packed_row_size..(output_index + 1) * packed_row_size];
let zp_row = &zero_points
[output_index * zero_point_row_size..(output_index + 1) * zero_point_row_size];
let zp_row = zero_points.map(|zp| {
&zp[output_index * zero_point_row_size..(output_index + 1) * zero_point_row_size]
});
let mut sum = bias.map_or(0.0, |values| values[output_index]);
for block in 0..block_count {
let packed_block = &packed_row[block * 16..(block + 1) * 16];
let activation_block = &activation[block * block_size..(block + 1) * block_size];
// SAFETY: runtime dispatch selected FEAT_DotProd and both slices
// contain exactly one block32 payload.
let dot = unsafe { affine_int4_block32_sdot_neon(activation_block, packed_block) };
let zero_point = layout.zero_point(Some(zp_row), block) as i32;
let zero_point = layout.zero_point(zp_row, block) as i32;
let correction = (8 - zero_point) * activation_sums[block];
sum += (dot + correction) as f32
* (activation_scales[block] * scales.get(output_index * block_count + block));
Expand Down Expand Up @@ -11112,7 +11154,11 @@ mod tests {
}

#[test]
fn matmulnbits_prepacked_m1_block32_symmetric_reuses_weight_for_new_activations() {
fn matmulnbits_m1_block32_symmetric_borrows_weight_for_new_activations() {
// Post-#979: symmetric int4 (constant B) borrows the packed weight in
// place per call instead of building/reusing a resident cache. The
// per-activation correctness invariant still holds; the difference is
// that no prepack/f32 cache is ever populated.
let (k, n, block_size) = (35, 7, 32);
let a1_values: Vec<f32> = (0..k)
.map(|i| ((i * 11 % 41) as f32 - 20.0) / 13.0)
Expand Down Expand Up @@ -11141,18 +11187,17 @@ mod tests {

let cached_ptr = prepack_cache_ptr(&kernel);
assert!(
cached_ptr.is_some(),
"M=1 constant B must populate a prepacked weight cache"
cached_ptr.is_none(),
"symmetric M=1 constant B must borrow in place, not populate a weight cache (#979)"
);
let a2 = Owned::f32(&[1, k], &a2_values);
let mut y2 = Owned::zeros_f32(&[1, n]);
kernel
.execute(&[a2.view(), b.view(), scales.view()], &mut [y2.view_mut()])
.unwrap();
assert_eq!(
prepack_cache_ptr(&kernel),
cached_ptr,
"prepacked weight cache must be reused (stable) across activations"
assert!(
!prepack_cache_populated(&kernel),
"symmetric decode must stay on the borrowed path across activations (#979)"
);
assert_close(&y2.to_f32(), &reference(&a2_values, &dequantized, 1, k, n));
assert_ne!(y1.to_f32(), y2.to_f32());
Expand All @@ -11179,6 +11224,7 @@ mod tests {
let scales = Owned::f32(&[n, 2], &scales);
let zero_points = Owned::u8(&[n, 1], &zero_points);
let mut y = Owned::zeros_f32(&[1, n]);
let asym_before = BORROWED_INT4_ASYMMETRIC_TEST_CALLS.load(Ordering::Relaxed);
kernel
.execute(
&[a.view(), b.view(), scales.view(), zero_points.view()],
Expand All @@ -11187,12 +11233,70 @@ mod tests {
.unwrap();

assert_close(&y.to_f32(), &reference(&a_values, &dequantized, 1, k, n));
assert!(
BORROWED_INT4_ASYMMETRIC_TEST_CALLS.load(Ordering::Relaxed) > asym_before,
"asymmetric INT4 decode must route into the borrowed zero-copy path"
);
assert!(
!prepack_cache_populated(&kernel),
"asymmetric INT4 must borrow packed inputs instead of building a weight cache"
);
}

#[test]
fn matmulnbits_symmetric_m1_borrows_instead_of_building_f32_cache() {
// Regression for #979: symmetric int4 (no zero_points input) must take
// the zero-copy borrowed decode path using the implicit midpoint 8, and
// must NOT fall through to the resident f32 `weight_nk` cache (~8x the
// file size in RAM). With constant B/scales, the pre-#979 code populated
// `weight_nk`; the borrowed path returns before any cache is built, so
// `!prepack_cache_populated` is a *positive* proof the branch changed.
let (k, n, block_size) = (128, 7, 32);
let a_values: Vec<f32> = (0..k)
.map(|i| ((i * 11 % 41) as f32 - 20.0) / 13.0)
.collect();
let weights: Vec<f32> = (0..n * k)
.map(|i| ((i * 23 % 47) as f32 - 19.0) / 12.0)
.collect();
let (packed, scales, zero_points, _) = quantize(&weights, n, k, block_size, false);
assert!(
zero_points.is_none(),
"symmetric quantize must omit the zero_points tensor"
);
let dequantized = dequantize_reference(&packed, &scales, None, n, k, block_size);
let mut kernel = test_kernel(k, n, block_size);
kernel.set_constant_inputs(&[false, true, true]);

let a = Owned::f32(&[1, k], &a_values);
let b = Owned::u8(&[n, 4, 16], &packed);
let scales = Owned::f32(&[n, 4], &scales);
let mut y = Owned::zeros_f32(&[1, n]);
let sym_before = BORROWED_INT4_SYMMETRIC_TEST_CALLS.load(Ordering::Relaxed);
kernel
.execute(&[a.view(), b.view(), scales.view()], &mut [y.view_mut()])
.unwrap();

assert_close(&y.to_f32(), &reference(&a_values, &dequantized, 1, k, n));
// Positive proof the *symmetric borrowed* branch is the one that ran.
// The symmetric counter is bumped only by `borrowed_zero_points.is_none()`,
// so a strict increase across this symmetric-only `execute` attributes the
// route to that branch (`>` not `==` because the global counter is shared
// with other tests running in parallel).
assert!(
BORROWED_INT4_SYMMETRIC_TEST_CALLS.load(Ordering::Relaxed) > sym_before,
"symmetric INT4 decode must route into the borrowed zero-copy path (#979)"
);
// ...and negative proof no resident f32 / prepack cache was ever built.
assert!(
kernel.weight_nk.get().is_none(),
"symmetric INT4 must not expand the weight to a resident f32 cache (#979)"
);
assert!(
!prepack_cache_populated(&kernel),
"symmetric INT4 must borrow packed inputs instead of building any weight cache (#979)"
);
}

#[test]
fn matmulnbits_m1_dynamic_b_falls_back_without_populating_prepack_cache() {
let (k, n, block_size) = (35, 5, 32);
Expand Down
Loading