Save inference parity implementation and evaluation harness
This commit is contained in:
@@ -0,0 +1,467 @@
|
||||
//! ResidencySets from the pinned MTPLX runtime. Metal calls stay in the bridge;
|
||||
//! set selection, budgets, membership and commit scheduling stay in Rust.
|
||||
use super::super::gpu::*;
|
||||
use std::collections::HashMap;
|
||||
use std::ffi::c_void;
|
||||
use std::ptr::NonNull;
|
||||
use std::sync::{
|
||||
Mutex, OnceLock,
|
||||
atomic::{AtomicUsize, Ordering},
|
||||
};
|
||||
|
||||
struct Set {
|
||||
raw: NonNull<c_void>,
|
||||
bytes: usize,
|
||||
}
|
||||
|
||||
// SAFETY: all operations on these owned Metal sets are serialized by State's
|
||||
// mutex. Queue attachment borrows handles while that same lock is held.
|
||||
unsafe impl Send for Set {}
|
||||
|
||||
impl Drop for Set {
|
||||
fn drop(&mut self) {
|
||||
unsafe { ds4_gpu_mtplx_residency_free(self.raw.as_ptr()) };
|
||||
}
|
||||
}
|
||||
|
||||
struct Placement {
|
||||
set: Option<usize>,
|
||||
bytes: usize,
|
||||
}
|
||||
|
||||
struct State {
|
||||
sets: Vec<Set>,
|
||||
placements: HashMap<usize, Placement>,
|
||||
capacity: usize,
|
||||
wired: usize,
|
||||
max_per_set: usize,
|
||||
debug: bool,
|
||||
}
|
||||
|
||||
impl State {
|
||||
fn add_set(&mut self, count: &AtomicUsize) -> bool {
|
||||
let Some(raw) = NonNull::new(unsafe { ds4_gpu_mtplx_residency_create() }) else {
|
||||
return false;
|
||||
};
|
||||
self.sets.push(Set { raw, bytes: 0 });
|
||||
count.store(self.sets.len(), Ordering::Release);
|
||||
if self.debug {
|
||||
eprintln!(
|
||||
"[residency] created residency set {} (max_bytes_per_set={} MB)",
|
||||
self.sets.len() - 1,
|
||||
self.max_per_set >> 20
|
||||
);
|
||||
}
|
||||
true
|
||||
}
|
||||
|
||||
fn choose(&mut self, bytes: usize, count: &AtomicUsize) -> usize {
|
||||
if self.max_per_set == 0 {
|
||||
return 0;
|
||||
}
|
||||
let mut empty = None;
|
||||
let mut emptiest = 0;
|
||||
for (i, set) in self.sets.iter().enumerate() {
|
||||
if set.bytes + bytes <= self.max_per_set {
|
||||
return i;
|
||||
}
|
||||
if i != 0 && set.bytes == 0 && empty.is_none() {
|
||||
empty = Some(i);
|
||||
}
|
||||
if set.bytes < self.sets[emptiest].bytes {
|
||||
emptiest = i;
|
||||
}
|
||||
}
|
||||
if let Some(i) = empty {
|
||||
i
|
||||
} else if self.sets.len() < 32 && self.add_set(count) {
|
||||
self.sets.len() - 1
|
||||
} else {
|
||||
emptiest
|
||||
}
|
||||
}
|
||||
|
||||
fn add(&mut self, key: usize, at: &mut Placement, count: &AtomicUsize) -> usize {
|
||||
let i = self.choose(at.bytes, count);
|
||||
unsafe {
|
||||
ds4_gpu_mtplx_residency_allocation(self.sets[i].raw.as_ptr(), key as *mut c_void, 1)
|
||||
};
|
||||
self.sets[i].bytes += at.bytes;
|
||||
self.wired += at.bytes;
|
||||
at.set = Some(i);
|
||||
i
|
||||
}
|
||||
|
||||
fn remove(&mut self, key: usize, at: &mut Placement) -> usize {
|
||||
let i = at.set.take().unwrap();
|
||||
unsafe {
|
||||
ds4_gpu_mtplx_residency_allocation(self.sets[i].raw.as_ptr(), key as *mut c_void, 0)
|
||||
};
|
||||
self.sets[i].bytes -= at.bytes;
|
||||
self.wired -= at.bytes;
|
||||
i
|
||||
}
|
||||
|
||||
fn commit(&self, i: usize) {
|
||||
unsafe { ds4_gpu_mtplx_residency_commit(self.sets[i].raw.as_ptr()) };
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) struct ResidencySets {
|
||||
enabled: bool,
|
||||
count: AtomicUsize,
|
||||
state: Mutex<State>,
|
||||
}
|
||||
|
||||
impl ResidencySets {
|
||||
pub(super) fn new(recommended: usize) -> Result<Self, String> {
|
||||
let enabled = unsafe { ds4_gpu_mtplx_residency_supported() } != 0;
|
||||
static ENV: OnceLock<(i32, bool)> = OnceLock::new();
|
||||
let (pct, debug) = if enabled {
|
||||
*ENV.get_or_init(|| {
|
||||
let integer = |name: &std::ffi::CStr, default| unsafe {
|
||||
let value = libc::getenv(name.as_ptr());
|
||||
if value.is_null() {
|
||||
default
|
||||
} else {
|
||||
libc::atoi(value)
|
||||
}
|
||||
};
|
||||
(
|
||||
integer(c"MLX_RESIDENCY_SET_MAX_PCT", 5),
|
||||
integer(c"MLX_RESIDENCY_DEBUG", 0) != 0,
|
||||
)
|
||||
})
|
||||
} else {
|
||||
(0, false)
|
||||
};
|
||||
let count = AtomicUsize::new(0);
|
||||
let mut state = State {
|
||||
sets: Vec::new(),
|
||||
placements: HashMap::new(),
|
||||
capacity: 0,
|
||||
wired: 0,
|
||||
debug,
|
||||
max_per_set: if pct <= 0 || pct >= 100 {
|
||||
0
|
||||
} else {
|
||||
(recommended / 100 * pct as usize).max(64 << 20)
|
||||
},
|
||||
};
|
||||
if enabled && !state.add_set(&count) {
|
||||
return Err("cannot construct initial Metal residency set".into());
|
||||
}
|
||||
Ok(Self {
|
||||
enabled,
|
||||
count,
|
||||
state: Mutex::new(state),
|
||||
})
|
||||
}
|
||||
|
||||
/// SAFETY: key is a live Metal allocation kept alive until erase or until
|
||||
/// this residency owner is destroyed. Heap-backed buffers use their heap.
|
||||
pub(super) unsafe fn insert(&self, key: usize) {
|
||||
if !self.enabled {
|
||||
return;
|
||||
}
|
||||
let bytes = unsafe { ds4_gpu_mtplx_allocated_size(key as *mut c_void) } as usize;
|
||||
let mut state = self.state.lock().unwrap();
|
||||
assert!(
|
||||
!state.placements.contains_key(&key),
|
||||
"duplicate residency allocation"
|
||||
);
|
||||
let mut at = Placement { set: None, bytes };
|
||||
if state.wired + bytes <= state.capacity {
|
||||
let i = state.add(key, &mut at, &self.count);
|
||||
state.commit(i);
|
||||
}
|
||||
state.placements.insert(key, at);
|
||||
}
|
||||
|
||||
pub(super) fn erase(&self, key: usize) {
|
||||
if !self.enabled {
|
||||
return;
|
||||
}
|
||||
let mut state = self.state.lock().unwrap();
|
||||
let mut at = state
|
||||
.placements
|
||||
.remove(&key)
|
||||
.expect("unknown residency allocation");
|
||||
if at.set.is_some() {
|
||||
let i = state.remove(key, &mut at);
|
||||
state.commit(i);
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg_attr(not(test), expect(dead_code, reason = "Reference diagnostic API"))]
|
||||
pub(super) fn resize(&self, capacity: usize) {
|
||||
if !self.enabled {
|
||||
return;
|
||||
}
|
||||
let mut state = self.state.lock().unwrap();
|
||||
if state.capacity == capacity {
|
||||
return;
|
||||
}
|
||||
state.capacity = capacity;
|
||||
let mut touched = vec![false; state.sets.len()];
|
||||
// Like the reference's unordered_map walk: membership is not sorted or
|
||||
// ranked. Moving the map preserves its iteration order and buckets;
|
||||
// only values change while set operations borrow the rest of State.
|
||||
let mut placements = std::mem::take(&mut state.placements);
|
||||
if state.wired < capacity {
|
||||
for (&key, at) in &mut placements {
|
||||
if at.set.is_some() || state.wired + at.bytes > capacity {
|
||||
continue;
|
||||
}
|
||||
let i = state.add(key, at, &self.count);
|
||||
touched.resize(state.sets.len(), false);
|
||||
touched[i] = true;
|
||||
}
|
||||
} else {
|
||||
for (&key, at) in &mut placements {
|
||||
if state.wired <= capacity {
|
||||
break;
|
||||
}
|
||||
if at.set.is_none() {
|
||||
continue;
|
||||
}
|
||||
let i = state.remove(key, at);
|
||||
touched[i] = true;
|
||||
}
|
||||
}
|
||||
state.placements = placements;
|
||||
for (i, changed) in touched.into_iter().enumerate() {
|
||||
if changed {
|
||||
state.commit(i);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// SAFETY: queue is live, and attached belongs to that queue and this owner.
|
||||
/// Call immediately before commit, including after new sets appeared.
|
||||
pub(super) unsafe fn attach_new_sets(&self, queue: *mut c_void, attached: &mut usize) {
|
||||
if self.count.load(Ordering::Acquire) == *attached {
|
||||
return;
|
||||
}
|
||||
let state = self.state.lock().unwrap();
|
||||
assert!(!queue.is_null() && *attached <= state.sets.len());
|
||||
let sets = state.sets[*attached..]
|
||||
.iter()
|
||||
.map(|s| s.raw.as_ptr())
|
||||
.collect::<Vec<_>>();
|
||||
unsafe { ds4_gpu_mtplx_residency_attach(queue, sets.as_ptr(), sets.len() as u64) };
|
||||
*attached = state.sets.len();
|
||||
}
|
||||
|
||||
#[cfg_attr(not(test), expect(dead_code, reason = "Reference diagnostic API"))]
|
||||
pub(super) fn wired_size(&self) -> usize {
|
||||
self.state.lock().unwrap().wired
|
||||
}
|
||||
#[cfg_attr(not(test), expect(dead_code, reason = "Reference diagnostic API"))]
|
||||
pub(super) fn num_sets(&self) -> usize {
|
||||
self.count.load(Ordering::Acquire)
|
||||
}
|
||||
|
||||
// Same test-only cap override as the reference's residency_tests.cpp.
|
||||
#[cfg_attr(not(test), expect(dead_code, reason = "Reference diagnostic API"))]
|
||||
pub(super) fn set_max_per_set(&self, bytes: usize) {
|
||||
self.state.lock().unwrap().max_per_set = bytes;
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[ignore = "requires Apple Metal 3 / macOS 15, no model weights"]
|
||||
fn mtplx_residency_matches_reference_lifecycle_cases() {
|
||||
use super::super::*;
|
||||
use super::allocator::Allocator;
|
||||
const MB: usize = 1 << 20;
|
||||
configure_sources().unwrap();
|
||||
let _context = Context::open_qwen(0).unwrap();
|
||||
let allocator = Allocator::new().unwrap();
|
||||
let residency = &allocator.residency;
|
||||
assert!(
|
||||
residency.enabled,
|
||||
"this test requires native Metal residency support"
|
||||
);
|
||||
allocator.set_cache_limit(0);
|
||||
assert_eq!(residency.wired_size(), 0);
|
||||
assert_eq!(residency.num_sets(), 1);
|
||||
let alloc = |bytes| allocator.allocate(bytes).unwrap().unwrap();
|
||||
let check_native = || {
|
||||
let state = residency.state.lock().unwrap();
|
||||
for (i, set) in state.sets.iter().enumerate() {
|
||||
let expected = state
|
||||
.placements
|
||||
.values()
|
||||
.filter(|at| at.set == Some(i))
|
||||
.count();
|
||||
assert_eq!(
|
||||
unsafe { ds4_gpu_mtplx_residency_count(set.raw.as_ptr()) } as usize,
|
||||
expected
|
||||
);
|
||||
for (&key, at) in &state.placements {
|
||||
assert_eq!(
|
||||
unsafe {
|
||||
ds4_gpu_mtplx_residency_contains(set.raw.as_ptr(), key as *mut c_void)
|
||||
} != 0,
|
||||
at.set == Some(i)
|
||||
);
|
||||
}
|
||||
assert_eq!(
|
||||
set.bytes,
|
||||
state
|
||||
.placements
|
||||
.values()
|
||||
.filter(|at| at.set == Some(i))
|
||||
.map(|at| at.bytes)
|
||||
.sum::<usize>()
|
||||
);
|
||||
}
|
||||
assert_eq!(
|
||||
state.wired,
|
||||
state.sets.iter().map(|set| set.bytes).sum::<usize>()
|
||||
);
|
||||
};
|
||||
let mut attached = 0;
|
||||
let check_gpu_work = |attached: &mut usize| {
|
||||
let input = alloc(16);
|
||||
let output = alloc(16);
|
||||
let expected = [0x80, 0x3f].repeat(8);
|
||||
input.write(0, &expected).unwrap();
|
||||
let commands = Commands::begin().unwrap();
|
||||
assert!(
|
||||
super::qsa_dynamic_copy_rows([&input, &output], [8, 8], [1, 8], [None, None], [8; 2])
|
||||
.unwrap()
|
||||
.is_empty()
|
||||
);
|
||||
// This is the reference commit boundary, not an early attach on malloc.
|
||||
unsafe { residency.attach_new_sets(ds4_gpu_mtplx_command_queue(), attached) };
|
||||
assert_eq!(*attached, residency.num_sets());
|
||||
commands.finish().unwrap();
|
||||
let mut actual = [0; 16];
|
||||
output.read(0, &mut actual).unwrap();
|
||||
assert_eq!(actual.as_slice(), expected);
|
||||
check_native();
|
||||
};
|
||||
check_gpu_work(&mut attached);
|
||||
|
||||
// These scenarios follow the pinned tests/residency_tests.cpp, including
|
||||
// its test-only set cap, rather than inventing a different budget policy.
|
||||
let buffers = (0..4).map(|_| alloc(4 * MB)).collect::<Vec<_>>();
|
||||
assert_eq!(residency.wired_size(), 0);
|
||||
drop(buffers);
|
||||
assert_eq!(residency.wired_size(), 0);
|
||||
assert_eq!(allocator.set_wired_limit(256 * MB).unwrap(), 0);
|
||||
let baseline = residency.wired_size();
|
||||
assert!(
|
||||
baseline >= MB,
|
||||
"the heap is wired once, not once per tiny buffer"
|
||||
);
|
||||
allocator.set_wired_limit(baseline + 8 * MB).unwrap();
|
||||
let mut buffers = Vec::new();
|
||||
for _ in 0..8 {
|
||||
buffers.push(alloc(4 * MB));
|
||||
assert!(residency.wired_size() <= baseline + 8 * MB);
|
||||
}
|
||||
assert!(residency.wired_size() >= baseline + 4 * MB);
|
||||
drop(buffers);
|
||||
assert_eq!(residency.wired_size(), baseline);
|
||||
|
||||
allocator.set_wired_limit(0).unwrap();
|
||||
let buffer = alloc(8 * MB);
|
||||
assert_eq!(residency.wired_size(), 0);
|
||||
allocator.set_wired_limit(64 * MB).unwrap();
|
||||
assert!(residency.wired_size() >= baseline + 8 * MB);
|
||||
allocator.set_wired_limit(0).unwrap();
|
||||
assert_eq!(residency.wired_size(), 0);
|
||||
allocator.set_wired_limit(64 * MB).unwrap();
|
||||
assert!(residency.wired_size() >= baseline + 8 * MB);
|
||||
drop(buffer);
|
||||
for _ in 0..32 {
|
||||
let buffer = alloc(8 * MB);
|
||||
assert!(residency.wired_size() >= baseline + 8 * MB);
|
||||
drop(buffer);
|
||||
assert_eq!(residency.wired_size(), baseline);
|
||||
}
|
||||
|
||||
residency.set_max_per_set(8 * MB);
|
||||
allocator.set_wired_limit(128 * MB).unwrap();
|
||||
let previous_sets = residency.num_sets();
|
||||
let buffers = (0..8).map(|_| alloc(5 * MB)).collect::<Vec<_>>();
|
||||
assert_eq!(residency.wired_size(), baseline + 40 * MB);
|
||||
assert!(residency.num_sets() >= 8 && residency.num_sets() > previous_sets);
|
||||
assert!(
|
||||
residency
|
||||
.state
|
||||
.lock()
|
||||
.unwrap()
|
||||
.sets
|
||||
.iter()
|
||||
.all(|s| s.bytes <= 8 * MB)
|
||||
);
|
||||
check_gpu_work(&mut attached);
|
||||
let count = residency.num_sets();
|
||||
drop(buffers);
|
||||
for _ in 0..4 {
|
||||
let buffers = (0..8).map(|_| alloc(5 * MB)).collect::<Vec<_>>();
|
||||
assert_eq!(residency.num_sets(), count);
|
||||
drop(buffers);
|
||||
}
|
||||
assert_eq!(residency.wired_size(), baseline);
|
||||
|
||||
residency.set_max_per_set(2 * MB);
|
||||
allocator.set_wired_limit(256 * MB).unwrap();
|
||||
let buffers = (0..64).map(|_| alloc(2 * MB)).collect::<Vec<_>>();
|
||||
assert_eq!(residency.num_sets(), 32);
|
||||
assert_eq!(residency.wired_size(), baseline + 128 * MB);
|
||||
check_gpu_work(&mut attached);
|
||||
drop(buffers);
|
||||
assert_eq!(residency.wired_size(), baseline);
|
||||
|
||||
residency.set_max_per_set(4 * MB);
|
||||
let buffer = alloc(32 * MB);
|
||||
assert_eq!(residency.wired_size(), baseline + 32 * MB);
|
||||
check_gpu_work(&mut attached);
|
||||
drop(buffer);
|
||||
assert_eq!(residency.wired_size(), baseline);
|
||||
residency.set_max_per_set(0);
|
||||
let count = residency.num_sets();
|
||||
let buffers = (0..16).map(|_| alloc(4 * MB)).collect::<Vec<_>>();
|
||||
assert_eq!(residency.num_sets(), count);
|
||||
assert_eq!(residency.wired_size(), baseline + 64 * MB);
|
||||
drop(buffers);
|
||||
assert_eq!(residency.wired_size(), baseline);
|
||||
|
||||
allocator.set_wired_limit(baseline + 4 * MB).unwrap();
|
||||
let first = alloc(4 * MB);
|
||||
let second = alloc(4 * MB);
|
||||
assert_eq!(residency.wired_size(), baseline + 4 * MB);
|
||||
drop(first);
|
||||
assert_eq!(
|
||||
residency.wired_size(),
|
||||
baseline,
|
||||
"free must not promote an unwired allocation"
|
||||
);
|
||||
// Reapplying the same limit also does not promote it.
|
||||
allocator.set_wired_limit(baseline + 4 * MB).unwrap();
|
||||
assert_eq!(residency.wired_size(), baseline);
|
||||
drop(second);
|
||||
assert!(allocator.set_wired_limit(usize::MAX).is_err());
|
||||
assert_eq!(residency.wired_size(), baseline);
|
||||
|
||||
allocator.set_wired_limit(64 * MB).unwrap();
|
||||
allocator.set_cache_limit(16 * MB);
|
||||
let buffer = alloc(4 * MB);
|
||||
drop(buffer);
|
||||
assert_eq!(
|
||||
residency.wired_size(),
|
||||
baseline + 4 * MB,
|
||||
"cached storage stays wired"
|
||||
);
|
||||
check_native();
|
||||
allocator.clear_cache();
|
||||
assert_eq!(residency.wired_size(), baseline);
|
||||
allocator.set_wired_limit(0).unwrap();
|
||||
assert_eq!(residency.wired_size(), 0);
|
||||
check_native();
|
||||
}
|
||||
Reference in New Issue
Block a user