Files
DS4Server/src/engine/metal/qwen_mtplx/residency.rs
T

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();
}