Skip to content
Merged
Show file tree
Hide file tree
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
21 changes: 20 additions & 1 deletion crates/onnx-genai-engine/src/engine/load.rs
Original file line number Diff line number Diff line change
Expand Up @@ -477,11 +477,21 @@ impl Engine {
// (which owns the kernel dispatch) how many extra bytes the cache costs,
// rather than re-deriving the rule here where it would drift from the
// kernel (#947). CUDA decode uses different kernels, so it never applies.
//
// #1056: the weight-transpose cache is the third such buffer. It holds
// one full `K x N` f32/f16 copy per transposed constant weight for the
// session (populated on all platforms by `Gemm` with `transB`, and on
// Apple also by `MatMul`). We ask the CPU EP to predict its bytes from
// the same graph and fold them into the resident total, so the single
// admission verdict below governs it alongside the other two.
let resident_f32_cache_bytes = if native_cuda_load {
0
} else {
match onnx_runtime_loader::load_model(&model_directory.model_path) {
Ok(graph) => onnx_runtime_ep_cpu::resident_dequant_f32_cache_bytes(&graph),
Ok(graph) => onnx_runtime_ep_cpu::resident_dequant_f32_cache_bytes(&graph)
.saturating_add(onnx_runtime_ep_cpu::weight_transpose_cache_predicted_bytes(
&graph,
)),
Err(_) => 0,
}
};
Expand Down Expand Up @@ -546,6 +556,15 @@ impl Engine {
onnx_runtime_ep_cpu::set_mlas_sqnbit_packing_enabled(
memory_strategy_plan.f32_weight_cache_admitted,
);
// #1056: the weight-transpose cache (one resident `K x N` f32/f16 copy
// per transposed constant weight) is the third buffer folded into
// `resident_f32_cache_bytes`, so the same verdict governs it. When
// declined, the `MatMul`/`Gemm` kernels transpose per call and retain
// nothing (byte-identical output, only slower decode) instead of holding
// the session-lifetime copies resident over budget.
onnx_runtime_ep_cpu::set_weight_transpose_cache_enabled(
memory_strategy_plan.f32_weight_cache_admitted,
);
#[cfg(feature = "cuda")]
let cuda_offload_resolution =
cuda_offload_resolution_from_plan(&native_device, &memory_strategy_plan);
Expand Down
206 changes: 206 additions & 0 deletions crates/onnx-runtime-ep-cpu/src/kernels/gemm.rs
Original file line number Diff line number Diff line change
Expand Up @@ -208,6 +208,7 @@ fn try_half_gemm(
mod tests {
use super::*;
use crate::kernels::testutil::Owned;
use onnx_runtime_ir::{Attribute, Graph, NodeId, TensorData, WeightRef};

#[allow(clippy::too_many_arguments)]
fn gemm(
Expand Down Expand Up @@ -669,4 +670,209 @@ mod tests {
it is no longer using the shared GEMM backend"
);
}

/// Serializes the tests that read the process-global weight-transpose cache
/// bytes or toggle its admission flag, so they never observe each other's
/// entries or setting under Rust's parallel test harness (#1056). Analogous
/// to `matmul_nbits`'s `CACHE_FLAG_TEST_LOCK`.
static TRANSPOSE_TEST_LOCK: std::sync::Mutex<()> = std::sync::Mutex::new(());

/// Build a single-node `Gemm` graph with a constant `[n, k]` `B` initializer
/// and `transB = 1`, for driving [`weight_transpose_cache_predicted_bytes`].
/// The initializer's *bytes* are irrelevant (the predictor reads only its
/// dims), so they are left zeroed; the executed kernel below uses its own
/// `Owned` `B` carrying real data.
fn trans_b_gemm_graph(n: usize, k: usize) -> Graph {
use onnx_runtime_ir::static_shape;
let mut graph = Graph::new();
let a = graph.create_named_value("A", DataType::Float32, static_shape([1, k]));
let b = graph.create_named_value("B", DataType::Float32, static_shape([n, k]));
let y = graph.create_named_value("Y", DataType::Float32, static_shape([1, n]));
graph.add_input(a);
graph.add_input(b);
let mut node = Node::new(NodeId(0), "Gemm", vec![Some(a), Some(b)], vec![y]);
node.attributes.insert("transB".into(), Attribute::Int(1));
graph.insert_node(node);
graph.add_output(y);
graph.set_initializer(
b,
WeightRef::Inline(TensorData::from_raw(
DataType::Float32,
vec![n, k],
vec![0u8; n * k * 4],
)),
);
graph
}

/// #1056 acceptance criterion: the predicted transpose-cache bytes must
/// equal the bytes actually held after a real run, ratio 1.00.
///
/// This runs the constant-weight `transB` `Gemm` kernel through the `Kernel`
/// trait at *two* activation shapes (prefill `m = 4`, decode `m = 1`) — the
/// exact duplication the shape-keyed kernel cache produces in an
/// autoregressive session — and asserts that the process-global cache grows
/// by the predicted amount *once*, proving effect (b) from #1051 (per-
/// instantiation copies) does **not** apply here: the global cache keys on
/// the weight address, so the decode instance hits the prefill entry.
#[test]
fn predicted_transpose_bytes_equal_actual_after_gemm_execution() {
let _guard = TRANSPOSE_TEST_LOCK.lock().unwrap_or_else(|e| e.into_inner());
// Thread-local admit: never mutates the process-global the concurrent
// `transposed_b` tests read (#1056 isolation).
let _admit = crate::kernels::weight_transpose::CacheEnabledScope::new(true);

// Distinctive geometry so a stray same-shaped entry cannot alias it.
let (n, k) = (37usize, 91usize);
let graph = trans_b_gemm_graph(n, k);
let predicted = matmul::weight_transpose_cache_predicted_bytes(&graph);
assert_eq!(
predicted,
(n as u64) * (k as u64) * 4,
"one f32 [k,n] transpose per constant transB Gemm weight"
);

// One `B` buffer, shared by both kernel instances, so both key the
// global cache identically.
let b_data: Vec<f32> = (0..n * k).map(|i| (i as f32 * 0.013).sin()).collect();
let b = Owned::f32(&[n, k], &b_data);
let b_ptr = b.view().data_ptr::<f32>();

// Evict any stale entry a since-freed weight of these dims may have left
// at this recycled address, so the first `run` below is a genuine miss
// and the byte total grows by `predicted` (rather than silently hitting a
// pre-existing entry and asserting against no growth). Safe: no live
// concurrent allocation can share `b_ptr`.
crate::kernels::weight_transpose::f32_cache_evict(b_ptr, n, k);

let run = |m: usize| -> Vec<f32> {
let a_data: Vec<f32> = (0..m * k).map(|i| (i as f32 * 0.007).cos()).collect();
let a = Owned::f32(&[m, k], &a_data);
let mut y = Owned::zeros_f32(&[m, n]);
let mut kernel = GemmKernel {
alpha: 1.0,
beta: 1.0,
trans_a: false,
trans_b: true,
prepack: MatMulPrepack::default(),
};
kernel.set_constant_inputs(&[false, true]);
kernel
.execute(&[a.view(), b.view()], &mut [y.view_mut()])
.unwrap();
y.to_f32()
};

let before = matmul::weight_transpose_cache_bytes();
let _prefill = run(4);
let _decode = run(1);
let after = matmul::weight_transpose_cache_bytes();

// The process-global byte total is shared across the parallel test
// harness, so a *concurrent* test caching an unrelated weight can also
// grow it. Measure the bytes held for THIS weight directly by re-keying
// the global cache with the exact `B` pointer the kernel used (a
// constant, contiguous f32 input densifies to a borrow of its own
// buffer, so `dense(1)` and this query share one address): a hit returns
// the entry the two runs installed, never a fresh allocation.
//
// SAFETY: `b` is a live, contiguous `[n, k]` f32 tensor for the whole
// test, so its data pointer addresses exactly `n * k` f32 values.
let b_slice = unsafe { std::slice::from_raw_parts(b.view().data_ptr::<f32>(), n * k) };
let entry = crate::kernels::weight_transpose::cached_transpose_f32(b_slice, n, k)
.expect("the executed kernel cached this constant weight's transpose");
let actual = (entry.len() * std::mem::size_of::<f32>()) as u64;
assert_eq!(
actual, predicted,
"predicted transpose-cache bytes must equal the bytes actually held \
for this weight (ratio 1.00); two shape instantiations retain one \
shared copy"
);
// And that copy is genuinely part of the global total the plan budgets.
assert!(
after >= before + predicted as usize,
"the cached transpose must be reflected in the global byte total"
);
}

/// #1056 decline contract: when the cache is declined, a constant `transB`
/// `Gemm` weight retains **nothing** (the transpose is recomputed per call
/// and freed), and the result is byte-identical to the admitted run — a pure
/// performance tradeoff, never a numerical one.
#[test]
fn declined_transpose_cache_retains_nothing_and_is_byte_identical() {
let _guard = TRANSPOSE_TEST_LOCK.lock().unwrap_or_else(|e| e.into_inner());

// A fresh `B` never seen by the global cache. Built once so its data
// pointer is stable across both runs and can be probed by exact key.
let (n, k, m) = (29usize, 53usize, 3usize);
let b_data: Vec<f32> = (0..n * k).map(|i| (i as f32 * 0.019).sin()).collect();
let a_data: Vec<f32> = (0..m * k).map(|i| (i as f32 * 0.011).cos()).collect();
let b = Owned::f32(&[n, k], &b_data);
// `b` outlives every use of this pointer below and is a contiguous
// `[n, k]` f32 tensor, so the key names exactly this weight.
let b_ptr = b.view().data_ptr::<f32>();

// Address reuse hygiene: the cache keys on `(addr, K, N)`, and an earlier
// (now-freed) weight of the same dims could have been recycled onto this
// exact address, leaving a stale entry that would make the probe below a
// false positive. Evict that key first so the "before" state is
// deterministically empty. Safe under the parallel harness: no *live*
// concurrent allocation can share `b_ptr`, so the only entry this can
// touch is a stale one nobody is using.
crate::kernels::weight_transpose::f32_cache_evict(b_ptr, n, k);
assert!(
!crate::kernels::weight_transpose::f32_cache_contains(b_ptr, n, k),
"precondition: this weight's key must start absent"
);

let run = || -> Vec<f32> {
let a = Owned::f32(&[m, k], &a_data);
let mut y = Owned::zeros_f32(&[m, n]);
let mut kernel = GemmKernel {
alpha: 1.0,
beta: 1.0,
trans_a: false,
trans_b: true,
prepack: MatMulPrepack::default(),
};
kernel.set_constant_inputs(&[false, true]);
kernel
.execute(&[a.view(), b.view()], &mut [y.view_mut()])
.unwrap();
y.to_f32()
};

// Declined: this weight's transpose must not be resident afterward.
// Probe the exact `(addr, n, k)` key the kernel installs (it calls
// `transposed_b(&b, n, k)`) rather than a global byte total, so a
// concurrent test caching an unrelated weight cannot mask a leak.
let declined = {
let _decline = crate::kernels::weight_transpose::CacheEnabledScope::new(false);
let out = run();
assert!(
!crate::kernels::weight_transpose::f32_cache_contains(b_ptr, n, k),
"a declined transpose cache must retain nothing for this weight"
);
out
};

// Admitted: byte-identical output (transpose is the same math either
// way) and now this weight *is* resident, confirming the two paths
// differ only in retention, not in numerics.
let admitted = {
let _admit = crate::kernels::weight_transpose::CacheEnabledScope::new(true);
let out = run();
assert!(
crate::kernels::weight_transpose::f32_cache_contains(b_ptr, n, k),
"an admitted transpose cache must retain this weight"
);
out
};
assert_eq!(
declined, admitted,
"declining the transpose cache must not change results"
);
}
}

Loading
Loading