//! 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, 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, bytes: usize, } struct State { sets: Vec, placements: HashMap, 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, } impl ResidencySets { pub(super) fn new(recommended: usize) -> Result { 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::>(); 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::() ); } assert_eq!( state.wired, state.sets.iter().map(|set| set.bytes).sum::() ); }; 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::>(); 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::>(); 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::>(); 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::>(); 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::>(); 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(); }