use std::collections::HashMap;
use std::sync::Mutex;
use cudarc::driver::sys::CUdeviceptr;
use onnx_runtime_ep_api::{EpError, Result};
fn error(message: impl Into<String>) -> EpError {
EpError::KernelFailed(format!("cuda_ep int4 interleave: {}", message.into()))
}
pub(crate) trait InterleaveDevice {
fn interleave_device_id(&self) -> u64;
fn interleave_alloc(&self, bytes: usize) -> Result<CUdeviceptr>;
unsafe fn interleave_free(&self, ptr: CUdeviceptr);
fn interleave_build(&self, src: CUdeviceptr, dst: CUdeviceptr, bytes: usize) -> Result<()>;
fn interleave_is_capturing(&self) -> Result<bool>;
fn interleave_frees_are_observed(&self) -> bool;
fn interleave_drain_before_free(&self) -> Result<()>;
}
#[derive(Debug, Default)]
pub(crate) struct InterleaveCache {
entries: Mutex<HashMap<(CUdeviceptr, usize), CUdeviceptr>>,
live: std::sync::atomic::AtomicUsize,
retired: Mutex<Vec<CUdeviceptr>>,
device: Mutex<Option<u64>>,
}
impl InterleaveCache {
pub(crate) fn ensure<D: InterleaveDevice>(
&self,
device: &D,
packed: CUdeviceptr,
bytes: usize,
) -> Result<(CUdeviceptr, bool)> {
self.bind(device.interleave_device_id())?;
if !device.interleave_frees_are_observed() {
return Err(error(
"int4 interleave is unavailable while device weight offload is on: a paged \
weight's pages are retired without passing through this cache, so its address \
can be recycled to a different weight behind the cache's back (#1726)",
));
}
let key = (packed, bytes);
if let Some(hit) = self.lookup(&key) {
return Ok((hit, true));
}
if device.interleave_is_capturing()? {
return Err(error(
"int4 interleave cannot allocate during CUDA-graph capture; the weight must be \
interleaved during warmup before capture",
));
}
let built = device.interleave_alloc(bytes)?;
if let Err(error) = device.interleave_build(packed, built, bytes) {
unsafe { device.interleave_free(built) };
return Err(error);
}
let mut entries = self.lock();
if let Some(&winner) = entries.get(&key) {
drop(entries);
unsafe { device.interleave_free(built) };
return Ok((winner, true));
}
entries.insert(key, built);
self.live
.store(entries.len(), std::sync::atomic::Ordering::Release);
Ok((built, false))
}
pub(crate) fn invalidate<D: InterleaveDevice>(
&self,
device: &D,
base: CUdeviceptr,
len: usize,
) {
if self.live.load(std::sync::atomic::Ordering::Acquire) == 0 {
return;
}
let end = base.saturating_add(len as CUdeviceptr);
let mut entries = self.lock();
let doomed: Vec<(CUdeviceptr, usize)> = entries
.keys()
.filter(|(packed, _)| *packed >= base && *packed < end)
.copied()
.collect();
let freed: Vec<CUdeviceptr> = doomed
.iter()
.filter_map(|key| entries.remove(key))
.collect();
self.live
.store(entries.len(), std::sync::atomic::Ordering::Release);
drop(entries);
if freed.is_empty() {
return;
}
if device.interleave_is_capturing().unwrap_or(true) {
self.retire(freed);
return;
}
if device.interleave_drain_before_free().is_err() {
self.retire(freed);
return;
}
for ptr in freed {
unsafe { device.interleave_free(ptr) };
}
}
pub(crate) fn release_all<D: InterleaveDevice>(&self, device: &D) {
let mut entries = self.lock();
let mut drained: Vec<CUdeviceptr> = entries.drain().map(|(_, ptr)| ptr).collect();
self.live.store(0, std::sync::atomic::Ordering::Release);
drop(entries);
drained.append(&mut self.retired.lock().unwrap_or_else(|e| e.into_inner()));
for ptr in drained {
unsafe { device.interleave_free(ptr) };
}
}
fn retire(&self, buffers: Vec<CUdeviceptr>) {
self.retired
.lock()
.unwrap_or_else(|e| e.into_inner())
.extend(buffers);
}
#[cfg(test)]
pub(crate) fn len(&self) -> usize {
self.lock().len()
}
#[cfg(test)]
pub(crate) fn retired_len(&self) -> usize {
self.retired.lock().unwrap_or_else(|e| e.into_inner()).len()
}
fn bind(&self, device: u64) -> Result<()> {
let mut bound = self.device.lock().unwrap_or_else(|e| e.into_inner());
match *bound {
None => {
*bound = Some(device);
Ok(())
}
Some(owner) if owner == device => Ok(()),
Some(owner) => Err(error(format!(
"cache built for device {owner} was asked to serve device {device}; an entry is \
keyed by a device address, which only names a weight for the device that minted \
it (#1726)"
))),
}
}
fn lookup(&self, key: &(CUdeviceptr, usize)) -> Option<CUdeviceptr> {
self.lock().get(key).copied()
}
fn lock(&self) -> std::sync::MutexGuard<'_, HashMap<(CUdeviceptr, usize), CUdeviceptr>> {
self.entries.lock().unwrap_or_else(|e| e.into_inner())
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::atomic::{AtomicUsize, Ordering};
#[derive(Default)]
struct RecyclingDevice {
blocks: Mutex<HashMap<CUdeviceptr, Vec<u8>>>,
free_lists: Mutex<HashMap<usize, Vec<CUdeviceptr>>>,
next_address: AtomicUsize,
allocations: AtomicUsize,
frees: AtomicUsize,
builds: AtomicUsize,
drains: AtomicUsize,
capturing: std::sync::atomic::AtomicBool,
readers: Mutex<Vec<CUdeviceptr>>,
frees_observed: std::sync::atomic::AtomicBool,
id: u64,
}
impl RecyclingDevice {
fn new() -> Self {
static NEXT_ID: AtomicUsize = AtomicUsize::new(1);
let device = Self {
next_address: AtomicUsize::new(0x1000),
id: NEXT_ID.fetch_add(1, Ordering::Relaxed) as u64,
..Default::default()
};
device.frees_observed.store(true, Ordering::Relaxed);
device
}
fn launch_reading(&self, ptr: CUdeviceptr) {
self.readers.lock().unwrap().push(ptr);
}
fn put(&self, contents: &[u8]) -> CUdeviceptr {
let ptr = self.raw_alloc(contents.len());
self.blocks.lock().unwrap().insert(ptr, contents.to_owned());
ptr
}
fn free(&self, ptr: CUdeviceptr) {
unsafe { self.interleave_free(ptr) };
}
fn contents(&self, ptr: CUdeviceptr) -> Vec<u8> {
self.blocks.lock().unwrap().get(&ptr).cloned().unwrap()
}
fn read(&self, ptr: CUdeviceptr, bytes: usize) -> Vec<u8> {
let blocks = self.blocks.lock().unwrap();
let (base, block) = blocks
.iter()
.find(|(base, block)| **base <= ptr && ptr < **base + block.len() as CUdeviceptr)
.unwrap_or_else(|| panic!("read from {ptr:#x}, which is in no live block"));
let offset = (ptr - base) as usize;
assert!(
offset + bytes <= block.len(),
"read {bytes} bytes at offset {offset} past the end of a {}-byte block",
block.len()
);
block[offset..offset + bytes].to_vec()
}
fn raw_alloc(&self, bytes: usize) -> CUdeviceptr {
self.allocations.fetch_add(1, Ordering::Relaxed);
if let Some(recycled) = self
.free_lists
.lock()
.unwrap()
.get_mut(&bytes)
.and_then(Vec::pop)
{
return recycled;
}
self.next_address.fetch_add(0x1_0000, Ordering::Relaxed) as CUdeviceptr
}
fn live_blocks(&self) -> usize {
self.blocks.lock().unwrap().len()
}
fn block_len(&self, ptr: CUdeviceptr) -> usize {
self.blocks.lock().unwrap().get(&ptr).map_or(0, Vec::len)
}
}
impl InterleaveDevice for RecyclingDevice {
fn interleave_device_id(&self) -> u64 {
self.id
}
fn interleave_alloc(&self, bytes: usize) -> Result<CUdeviceptr> {
let ptr = self.raw_alloc(bytes);
self.blocks.lock().unwrap().insert(ptr, vec![0; bytes]);
Ok(ptr)
}
unsafe fn interleave_free(&self, ptr: CUdeviceptr) {
self.frees.fetch_add(1, Ordering::Relaxed);
assert!(
!self.readers.lock().unwrap().contains(&ptr),
"freed the block at {ptr:#x} while a launch was still reading it; it goes onto \
the reuse free list immediately, so the next allocation can overwrite memory a \
live kernel is reading"
);
let bytes = self
.blocks
.lock()
.unwrap()
.remove(&ptr)
.expect("freed a live block")
.len();
self.free_lists
.lock()
.unwrap()
.entry(bytes)
.or_default()
.push(ptr);
}
fn interleave_build(&self, src: CUdeviceptr, dst: CUdeviceptr, bytes: usize) -> Result<()> {
self.builds.fetch_add(1, Ordering::Relaxed);
let source = self.read(src, bytes);
let built: Vec<u8> = source.iter().map(|b| b.rotate_left(4)).collect();
self.blocks.lock().unwrap().insert(dst, built);
Ok(())
}
fn interleave_is_capturing(&self) -> Result<bool> {
Ok(self.capturing.load(Ordering::Relaxed))
}
fn interleave_frees_are_observed(&self) -> bool {
self.frees_observed.load(Ordering::Relaxed)
}
fn interleave_drain_before_free(&self) -> Result<()> {
self.drains.fetch_add(1, Ordering::Relaxed);
self.readers.lock().unwrap().clear();
Ok(())
}
}
struct FakeRuntime<'a> {
device: &'a RecyclingDevice,
interleave: InterleaveCache,
}
impl<'a> FakeRuntime<'a> {
fn new(device: &'a RecyclingDevice) -> Self {
Self {
device,
interleave: InterleaveCache::default(),
}
}
fn ensure_interleaved_int4(
&self,
packed: CUdeviceptr,
bytes: usize,
) -> Result<(CUdeviceptr, bool)> {
self.interleave.ensure(self.device, packed, bytes)
}
fn interleaved_weight_count(&self) -> usize {
self.interleave.len()
}
fn deallocate(&self, base: CUdeviceptr) {
let len = self.device.block_len(base);
self.interleave.invalidate(self.device, base, len);
self.device.free(base);
}
}
impl Drop for FakeRuntime<'_> {
fn drop(&mut self) {
self.device.interleave_drain_before_free().unwrap();
self.interleave.release_all(self.device);
}
}
fn weight(seed: u8, bytes: usize) -> Vec<u8> {
(0..bytes).map(|i| seed.wrapping_add(i as u8)).collect()
}
fn interleaved(source: &[u8]) -> Vec<u8> {
source.iter().map(|b| b.rotate_left(4)).collect()
}
fn digest(bytes: &[u8]) -> String {
let sum = bytes
.iter()
.fold(0u64, |acc, &b| acc.wrapping_mul(31).wrapping_add(b as u64));
format!("{sum:016x}/{}b", bytes.len())
}
#[test]
fn a_recycled_weight_address_must_not_serve_the_previous_weights_interleave() {
const BYTES: usize = 4096;
let device = RecyclingDevice::new();
let first = weight(0x11, BYTES);
let first_ptr = device.put(&first);
{
let provider = FakeRuntime::new(&device);
let (built, warm) = provider.ensure_interleaved_int4(first_ptr, BYTES).unwrap();
assert!(!warm, "the first sight of a weight must build it");
assert_eq!(provider.interleaved_weight_count(), 1);
assert_eq!(
digest(&device.contents(built)),
digest(&interleaved(&first))
);
}
device.free(first_ptr);
let second = weight(0x77, BYTES);
let second_ptr = device.put(&second);
assert_eq!(
second_ptr, first_ptr,
"the falsifier is vacuous unless the address is actually recycled"
);
assert_ne!(first, second, "the two weights must be distinguishable");
let provider = FakeRuntime::new(&device);
let (built, warm) = provider.ensure_interleaved_int4(second_ptr, BYTES).unwrap();
assert!(
!warm,
"a weight this provider has never seen must be built, not served warm"
);
assert_eq!(provider.interleaved_weight_count(), 1);
assert_eq!(
digest(&device.contents(built)),
digest(&interleaved(&second)),
"served the previous weight's interleave at recycled address {second_ptr:#x}"
);
}
#[test]
fn a_weight_freed_by_one_executor_must_not_lend_its_interleave_to_the_next() {
const BYTES: usize = 4096;
let device = RecyclingDevice::new();
let runtime = FakeRuntime::new(&device);
let first = weight(0x11, BYTES);
let first_ptr = device.put(&first);
let (built, warm) = runtime.ensure_interleaved_int4(first_ptr, BYTES).unwrap();
assert!(!warm, "the first sight of a weight must build it");
assert_eq!(
digest(&device.contents(built)),
digest(&interleaved(&first)),
"executor 1 was served an interleave that is not its own weight's"
);
runtime.deallocate(first_ptr);
let second = weight(0xa5, BYTES);
let second_ptr = device.put(&second);
assert_eq!(
second_ptr, first_ptr,
"the arena must recycle the address or this test proves nothing"
);
assert_ne!(first, second, "the two weights must be distinguishable");
let (rebuilt, warm) = runtime.ensure_interleaved_int4(second_ptr, BYTES).unwrap();
assert_eq!(
digest(&device.contents(rebuilt)),
digest(&interleaved(&second)),
"executor 2 was served the interleave of the weight executor 1 freed, at recycled \
address {second_ptr:#x}"
);
assert!(
!warm,
"a recycled address must be a cold miss; a warm hit here is the freed weight's entry"
);
assert_eq!(
runtime.interleaved_weight_count(),
1,
"executor 2's weight, and only executor 2's, should be cached"
);
}
#[test]
fn repeated_executor_teardowns_never_serve_a_previous_plans_interleave() {
const BYTES: usize = 2048;
let device = RecyclingDevice::new();
let runtime = FakeRuntime::new(&device);
let mut addresses = Vec::new();
for round in 0..6u8 {
let w = weight(0x20u8.wrapping_add(round.wrapping_mul(0x31)), BYTES);
let ptr = device.put(&w);
addresses.push(ptr);
let (built, warm) = runtime.ensure_interleaved_int4(ptr, BYTES).unwrap();
assert_eq!(
digest(&device.contents(built)),
digest(&interleaved(&w)),
"round {round}: served an interleave built for a different weight at {ptr:#x}"
);
assert!(
!warm,
"round {round}: every plan builds its own weight; a warm hit means the entry \
outlived the plan that installed it"
);
runtime.deallocate(ptr);
}
assert!(
addresses.windows(2).any(|w| w[0] == w[1]),
"the arena never recycled an address across rounds, so nothing was falsified: {addresses:#x?}"
);
assert_eq!(
runtime.interleaved_weight_count(),
0,
"every round's entry must have been released with its plan"
);
}
#[test]
fn freeing_one_weight_leaves_every_other_weights_interleave_alone() {
const BYTES: usize = 512;
let device = RecyclingDevice::new();
let runtime = FakeRuntime::new(&device);
let mine = weight(0x77, BYTES);
let theirs = weight(0x0e, BYTES);
let mine_ptr = device.put(&mine);
let theirs_ptr = device.put(&theirs);
let (mine_built, _) = runtime.ensure_interleaved_int4(mine_ptr, BYTES).unwrap();
let (theirs_built, _) = runtime.ensure_interleaved_int4(theirs_ptr, BYTES).unwrap();
assert_ne!(mine_built, theirs_built);
assert_eq!(runtime.interleaved_weight_count(), 2);
runtime.deallocate(mine_ptr);
assert_eq!(
runtime.interleaved_weight_count(),
1,
"freeing one weight must evict exactly its own entry"
);
let (still, warm) = runtime.ensure_interleaved_int4(theirs_ptr, BYTES).unwrap();
assert!(warm, "the surviving executor's weight must still be cached");
assert_eq!(
still, theirs_built,
"the surviving executor's interleaved buffer moved; a captured graph holding the old \
pointer would now replay into freed memory"
);
assert_eq!(
digest(&device.contents(still)),
digest(&interleaved(&theirs)),
"the surviving buffer no longer holds its own weight's interleave"
);
runtime.deallocate(theirs_ptr);
assert_eq!(runtime.interleaved_weight_count(), 0);
}
#[test]
fn invalidating_an_address_the_cache_never_saw_is_a_no_op() {
const BYTES: usize = 512;
let device = RecyclingDevice::new();
let runtime = FakeRuntime::new(&device);
let scratch = device.put(&weight(0x01, BYTES));
runtime.deallocate(scratch);
assert_eq!(runtime.interleaved_weight_count(), 0);
let w = weight(0x77, BYTES);
let ptr = device.put(&w);
runtime.ensure_interleaved_int4(ptr, BYTES).unwrap();
assert_eq!(runtime.interleaved_weight_count(), 1);
runtime.interleave.invalidate(&device, ptr, BYTES);
runtime.interleave.invalidate(&device, ptr, BYTES);
assert_eq!(runtime.interleaved_weight_count(), 0);
assert_eq!(
device.live_blocks(),
1,
"only the caller-owned weight should remain live"
);
device.free(ptr);
}
#[test]
fn a_weight_at_an_offset_inside_the_freed_buffer_loses_its_interleave() {
const BYTES: usize = 512;
const OFFSET: CUdeviceptr = BYTES as CUdeviceptr;
let device = RecyclingDevice::new();
let runtime = FakeRuntime::new(&device);
let first = weight(0x11, BYTES);
let second = weight(0x99, BYTES);
let mut whole = first.clone();
whole.extend_from_slice(&second);
let base = device.put(&whole);
let view = base + OFFSET;
let (built, _) = runtime.ensure_interleaved_int4(view, BYTES).unwrap();
assert_eq!(
digest(&device.contents(built)),
digest(&interleaved(&second)),
"the offset view's interleave should be built from the bytes at that offset"
);
runtime.deallocate(base);
assert_eq!(
runtime.interleaved_weight_count(),
0,
"the entry keyed at base+{OFFSET:#x} must die with the buffer that contained it; \
matching only the base leaves it alive to serve the next weight at that offset"
);
}
#[test]
fn an_interleave_buffer_is_not_freed_under_a_launch_still_reading_it() {
const BYTES: usize = 1024;
let device = RecyclingDevice::new();
let runtime = FakeRuntime::new(&device);
let ptr = device.put(&weight(0x5a, BYTES));
let (built, _) = runtime.ensure_interleaved_int4(ptr, BYTES).unwrap();
device.launch_reading(ptr);
device.launch_reading(built);
runtime.deallocate(ptr);
assert_eq!(
device.drains.load(Ordering::Relaxed),
1,
"evicting an entry must drain in-flight work before returning the buffer to the \
allocator's reuse pool"
);
assert_eq!(runtime.interleaved_weight_count(), 0);
}
#[test]
fn a_device_with_unobserved_frees_is_refused_rather_than_cached_for() {
const BYTES: usize = 512;
let device = RecyclingDevice::new();
device.frees_observed.store(false, Ordering::Relaxed);
let runtime = FakeRuntime::new(&device);
let ptr = device.put(&weight(0x21, BYTES));
assert!(
runtime.ensure_interleaved_int4(ptr, BYTES).is_err(),
"a device that does not report every weight free must be refused, not served"
);
assert_eq!(
runtime.interleaved_weight_count(),
0,
"a refused call must install nothing; an entry here would outlive the page it was \
keyed on"
);
assert_eq!(
device.builds.load(Ordering::Relaxed),
0,
"a refused call must not build either"
);
device.free(ptr);
let recycled = device.put(&weight(0xc4, BYTES));
assert_eq!(
recycled, ptr,
"the harness must actually recycle the address for this test to mean anything"
);
assert!(
runtime.ensure_interleaved_int4(recycled, BYTES).is_err(),
"still refused, so the second weight cannot be served the first one's bytes"
);
device.free(recycled);
}
#[test]
fn alternating_weights_on_one_recycled_address_each_get_their_own_interleave() {
const BYTES: usize = 2048;
let device = RecyclingDevice::new();
let weights = [weight(0x05, BYTES), weight(0xa0, BYTES)];
let mut seen: Vec<(CUdeviceptr, usize)> = Vec::new();
let mut recycled_across_weights = false;
for round in 0..8 {
let which = round % 2;
let source = &weights[which];
let ptr = device.put(source);
let recycled = seen
.iter()
.any(|&(seen_ptr, seen_which)| seen_ptr == ptr && seen_which != which);
recycled_across_weights |= recycled;
seen.push((ptr, which));
let provider = FakeRuntime::new(&device);
let (built, _) = provider.ensure_interleaved_int4(ptr, BYTES).unwrap();
assert_eq!(
digest(&device.contents(built)),
digest(&interleaved(source)),
"round {round} (address {ptr:#x} recycled={recycled}): served an interleave \
built for a different weight at this address"
);
drop(provider);
device.free(ptr);
}
assert!(
recycled_across_weights,
"no address was reused across two different weights, so this proved nothing"
);
}
#[test]
fn releasing_the_cache_frees_every_buffer_it_built() {
const BYTES: usize = 512;
let device = RecyclingDevice::new();
let sources: Vec<CUdeviceptr> = (0..4).map(|i| device.put(&weight(i, BYTES))).collect();
let cache = InterleaveCache::default();
for &ptr in &sources {
cache.ensure(&device, ptr, BYTES).unwrap();
}
assert_eq!(cache.len(), 4);
assert_eq!(device.live_blocks(), 8, "four weights and four interleaves");
cache.release_all(&device);
assert_eq!(cache.len(), 0);
assert_eq!(
device.live_blocks(),
4,
"release_all must free the interleaves and nothing else"
);
assert_eq!(device.frees.load(Ordering::Relaxed), 4);
}
#[test]
fn a_second_ask_for_a_live_weight_is_served_without_building() {
const BYTES: usize = 256;
let device = RecyclingDevice::new();
let ptr = device.put(&weight(3, BYTES));
let cache = InterleaveCache::default();
let (first, warm) = cache.ensure(&device, ptr, BYTES).unwrap();
assert!(!warm);
let builds = device.builds.load(Ordering::Relaxed);
let allocations = device.allocations.load(Ordering::Relaxed);
let (second, warm) = cache.ensure(&device, ptr, BYTES).unwrap();
assert!(warm, "a cached weight must report warm");
assert_eq!(first, second);
assert_eq!(device.builds.load(Ordering::Relaxed), builds);
assert_eq!(device.allocations.load(Ordering::Relaxed), allocations);
cache.release_all(&device);
device.free(ptr);
}
#[test]
fn a_cold_miss_during_capture_is_refused() {
const BYTES: usize = 256;
let device = RecyclingDevice::new();
let ptr = device.put(&weight(9, BYTES));
let cache = InterleaveCache::default();
device.capturing.store(true, Ordering::Relaxed);
assert!(cache.ensure(&device, ptr, BYTES).is_err());
assert_eq!(device.builds.load(Ordering::Relaxed), 0);
assert_eq!(cache.len(), 0);
device.capturing.store(false, Ordering::Relaxed);
cache.ensure(&device, ptr, BYTES).unwrap();
device.capturing.store(true, Ordering::Relaxed);
let (_, warm) = cache.ensure(&device, ptr, BYTES).unwrap();
assert!(warm);
device.capturing.store(false, Ordering::Relaxed);
cache.release_all(&device);
device.free(ptr);
}
#[test]
fn byte_length_is_part_of_the_identity() {
const LONG: usize = 1024;
const SHORT: usize = 512;
let device = RecyclingDevice::new();
let long = weight(0x21, LONG);
let ptr = device.put(&long);
let cache = InterleaveCache::default();
let (built_long, warm) = cache.ensure(&device, ptr, LONG).unwrap();
assert!(!warm);
assert_eq!(
digest(&device.contents(built_long)),
digest(&interleaved(&long))
);
let (built_short, warm) = cache.ensure(&device, ptr, SHORT).unwrap();
assert!(
!warm,
"a different byte length at the same address hit the longer entry, \
so the length is not part of the identity"
);
assert_ne!(built_short, built_long);
assert_eq!(cache.len(), 2, "the two lengths must be two entries");
assert_eq!(
digest(&device.contents(built_short)),
digest(&interleaved(&long[..SHORT]))
);
cache.release_all(&device);
}
#[test]
fn a_cache_that_has_served_one_device_refuses_another() {
const BYTES: usize = 256;
let first = RecyclingDevice::new();
let second = RecyclingDevice::new();
assert_ne!(first.interleave_device_id(), second.interleave_device_id());
let shared = InterleaveCache::default();
let ptr = first.put(&weight(1, BYTES));
shared.ensure(&first, ptr, BYTES).unwrap();
let error = shared
.ensure(&second, second.put(&weight(2, BYTES)), BYTES)
.expect_err("a cache bound to one device must refuse another");
let message = format!("{error:?}");
assert!(
message.contains("1726"),
"the refusal must name the defect it prevents: {message}"
);
assert_eq!(
second.builds.load(Ordering::Relaxed),
0,
"the refused device must not have built anything"
);
shared.release_all(&first);
}
#[test]
fn concurrent_first_sights_keep_one_buffer_and_leak_nothing() {
const BYTES: usize = 1024;
let device = RecyclingDevice::new();
let ptr = device.put(&weight(0x5a, BYTES));
let cache = InterleaveCache::default();
let served: Vec<CUdeviceptr> = std::thread::scope(|scope| {
let handles: Vec<_> = (0..8)
.map(|_| scope.spawn(|| cache.ensure(&device, ptr, BYTES).unwrap().0))
.collect();
handles.into_iter().map(|h| h.join().unwrap()).collect()
});
let winner = served[0];
assert!(
served.iter().all(|&p| p == winner),
"racers were served different buffers: {served:?}"
);
assert_eq!(cache.len(), 1);
assert_eq!(
digest(&device.contents(winner)),
digest(&interleaved(&device.contents(ptr)))
);
assert_eq!(
device.allocations.load(Ordering::Relaxed) - device.frees.load(Ordering::Relaxed),
2,
"exactly the weight and the surviving interleave stay live"
);
cache.release_all(&device);
device.free(ptr);
}
#[test]
fn an_eviction_during_capture_drops_the_entry_without_synchronizing() {
const BYTES: usize = 512;
let device = RecyclingDevice::new();
let runtime = FakeRuntime::new(&device);
let ptr = device.put(&weight(0x3c, BYTES));
let (built, _) = runtime.ensure_interleaved_int4(ptr, BYTES).unwrap();
assert_eq!(runtime.interleaved_weight_count(), 1);
device.capturing.store(true, Ordering::Relaxed);
device.launch_reading(built);
runtime.interleave.invalidate(&device, ptr, BYTES);
assert_eq!(
runtime.interleaved_weight_count(),
0,
"the entry must go even during capture; it is the entry that would serve the next \
weight at this address the wrong bytes"
);
assert_eq!(
device.drains.load(Ordering::Relaxed),
0,
"synchronizing during a capture is illegal and would invalidate it"
);
assert_eq!(
device.frees.load(Ordering::Relaxed),
0,
"the buffer must not be handed back under a graph that references it"
);
assert_eq!(
runtime.interleave.retired_len(),
1,
"a buffer that could not be handed back must be parked, not dropped on the floor"
);
device.capturing.store(false, Ordering::Relaxed);
drop(runtime);
assert_eq!(
device.frees.load(Ordering::Relaxed),
1,
"teardown must reclaim the parked buffer"
);
device.free(ptr);
}
#[test]
#[cfg_attr(miri, ignore)]
fn the_pager_constructors_mark_the_runtime_as_paging() {
let source = include_str!("weight_paging.rs");
let code: String = source
.lines()
.map(|line| match line.find("//") {
Some(at) => &line[..at],
None => line,
})
.collect::<Vec<_>>()
.join("\n");
let marker = "set_weights_may_be_paged()";
for owner in ["CudaWeightPager", "CudaWeightResidency"] {
let impl_at = code
.find(&format!(
"impl<'a, S: MmapRegionSource + ?Sized> {owner}<'a, S> {{"
))
.or_else(|| code.find(&format!("impl {owner} {{")))
.unwrap_or_else(|| {
panic!(
"could not find the `impl {owner}` block in weight_paging.rs; if it was \
renamed, update this test rather than deleting it -- otherwise it \
passes by finding nothing to check"
)
});
let ctor_at = code[impl_at..]
.find("pub fn new(")
.map(|at| impl_at + at)
.unwrap_or_else(|| panic!("`{owner}` has no `new` constructor any more"));
let sig_end = code[ctor_at..]
.find('\n')
.map(|at| ctor_at + at)
.expect("a constructor signature is followed by a body");
let body_end = code[sig_end..]
.find("Self {")
.map(|at| sig_end + at)
.unwrap_or_else(|| panic!("`{owner}::new` no longer returns a `Self` literal"));
assert!(
code[sig_end..body_end].contains(marker),
"`{owner}::new` does not call {marker}. A paged weight's pages are retired by \
weight_paging rather than by the provider's deallocate, so the interleave cache \
is never told the address died -- it has to refuse to cache on this runtime at \
all, and it only knows to when this is called. Marking inside the constructor \
is what makes that unmissable; moving it to the call sites means the next call \
site added silently reopens #1726."
);
}
}
}