468 lines
16 KiB
Rust
468 lines
16 KiB
Rust
//! 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();
|
|
}
|