use std::sync::Arc;
use onnx_runtime_memory_governor::{
HolderId, MemoryError, MemoryGovernor, MemoryLease, MemoryRole, Tier,
};
use crate::VirtualMemoryError;
use crate::backing::{HostBacking, PhysicalMemoryAccounting, VirtualBacking};
pub struct VirtualBuffer<B: VirtualBacking = HostBacking> {
backing: B,
reservation: B::Reservation,
capacity_bytes: usize,
len: usize,
committed: usize,
governor: Arc<dyn MemoryGovernor + Send + Sync>,
tier: Tier,
role: MemoryRole,
holder: HolderId,
physical_memory_accounting: PhysicalMemoryAccounting,
lease: Option<MemoryLease>,
}
#[derive(Debug, thiserror::Error)]
pub enum VirtualBufferError {
#[error(transparent)]
Memory(#[from] VirtualMemoryError),
#[error(transparent)]
Budget(#[from] MemoryError),
#[error(
"virtual backing charges physical memory to {backing}, but the buffer governor uses \
{governor}; both must use the same memory authority"
)]
AuthorityMismatch {
backing: onnx_runtime_memory_governor::MemoryAuthorityId,
governor: onnx_runtime_memory_governor::MemoryAuthorityId,
},
#[error(
"cannot grow to {requested} bytes: this buffer reserved {capacity} bytes of address \
space and the reservation cannot be extended in place; construct it with a larger \
capacity"
)]
OverCapacity {
requested: usize,
capacity: usize,
},
}
impl VirtualBuffer<HostBacking> {
pub fn with_capacity(
capacity: usize,
governor: Arc<dyn MemoryGovernor + Send + Sync>,
tier: Tier,
role: MemoryRole,
holder: HolderId,
) -> Result<Self, VirtualBufferError> {
Self::with_backing(HostBacking, capacity, governor, tier, role, holder)
}
}
impl<B: VirtualBacking> VirtualBuffer<B> {
pub fn with_backing(
backing: B,
capacity: usize,
governor: Arc<dyn MemoryGovernor + Send + Sync>,
tier: Tier,
role: MemoryRole,
holder: HolderId,
) -> Result<Self, VirtualBufferError> {
let physical_memory_accounting = backing.physical_memory_accounting();
if let PhysicalMemoryAccounting::Backing { authority } = physical_memory_accounting {
let governor_authority = governor.authority_id();
if authority != governor_authority {
return Err(VirtualBufferError::AuthorityMismatch {
backing: authority,
governor: governor_authority,
});
}
}
let capacity = round_up(backing.granularity(), capacity.max(1));
let reservation = backing.reserve(capacity)?;
Ok(Self {
backing,
reservation,
capacity_bytes: capacity,
len: 0,
committed: 0,
governor,
tier,
role,
holder,
physical_memory_accounting,
lease: None,
})
}
pub fn len(&self) -> usize {
self.len
}
pub fn is_empty(&self) -> bool {
self.len == 0
}
pub fn committed(&self) -> usize {
self.committed
}
pub fn capacity(&self) -> usize {
self.capacity_bytes
}
pub fn as_ptr(&self) -> *const u8 {
B::base(&self.reservation) as *const u8
}
pub fn as_mut_ptr(&mut self) -> *mut u8 {
B::base(&self.reservation) as *mut u8
}
pub unsafe fn as_slice(&self) -> &[u8] {
unsafe { std::slice::from_raw_parts(self.as_ptr(), self.len) }
}
pub fn grow_to(&mut self, bytes: usize) -> Result<(), VirtualBufferError> {
if bytes <= self.len {
return Ok(());
}
if bytes > self.capacity() {
return Err(VirtualBufferError::OverCapacity {
requested: bytes,
capacity: self.capacity(),
});
}
let needed = round_up(self.backing.granularity(), bytes);
if needed > self.committed {
let extra = needed - self.committed;
let backing_accounts = matches!(
self.physical_memory_accounting,
PhysicalMemoryAccounting::Backing { .. }
);
if !backing_accounts {
match self.lease.as_mut() {
Some(lease) => lease.grow(extra as u64)?,
None => {
self.lease = Some(self.governor.reserve(
self.tier,
extra as u64,
self.role,
self.holder,
)?);
}
}
}
if let Err(error) = self
.backing
.commit(&mut self.reservation, self.committed, extra)
{
if !backing_accounts {
self.release(extra);
}
return Err(error.into());
}
self.committed = needed;
}
self.len = bytes;
Ok(())
}
pub fn shrink_to(&mut self, bytes: usize) -> Result<(), VirtualBufferError> {
if bytes >= self.len {
return Ok(());
}
let needed = round_up(self.backing.granularity(), bytes);
let mut offset = self.committed;
while offset > needed {
let granule = self.backing.granularity();
offset -= granule;
self.backing
.release(&mut self.reservation, offset, granule)?;
if matches!(
self.physical_memory_accounting,
PhysicalMemoryAccounting::Buffer
) {
self.release(granule);
}
self.committed = offset;
}
self.len = bytes;
Ok(())
}
fn release(&mut self, bytes: usize) {
let Some(lease) = self.lease.as_mut() else {
return;
};
lease.shrink(bytes as u64);
if lease.bytes() == 0 {
self.lease = None;
}
}
}
fn round_up(granule: usize, bytes: usize) -> usize {
bytes.div_ceil(granule) * granule
}
impl<B: VirtualBacking> std::fmt::Debug for VirtualBuffer<B> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("VirtualBuffer")
.field("base", &self.as_ptr())
.field("len", &self.len)
.field("committed", &self.committed)
.field("capacity", &self.capacity())
.field("tier", &self.tier)
.field("role", &self.role)
.finish()
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::granularity;
use onnx_runtime_memory_governor::{DeviceKey, LeaseLedger, LedgerGovernor, MemoryAuthorityId};
use std::sync::atomic::{AtomicUsize, Ordering};
const HOLDER: HolderId = HolderId::new(4);
fn buffer(capacity: usize, budget: u64) -> (VirtualBuffer, LedgerGovernor) {
let governor = LedgerGovernor::new(LeaseLedger::new(0, budget, 0));
let buffer = VirtualBuffer::with_capacity(
capacity,
Arc::new(governor.clone()),
Tier::Host,
MemoryRole::KvCache,
HOLDER,
)
.expect("address space");
(buffer, governor)
}
#[derive(Debug, Clone)]
struct AuthorityBacking {
authority: MemoryAuthorityId,
reserves: Arc<AtomicUsize>,
commits: Arc<AtomicUsize>,
}
unsafe impl VirtualBacking for AuthorityBacking {
type Reservation = usize;
fn granularity(&self) -> usize {
4096
}
fn physical_memory_accounting(&self) -> PhysicalMemoryAccounting {
PhysicalMemoryAccounting::Backing {
authority: self.authority,
}
}
fn reserve(&self, _len: usize) -> Result<Self::Reservation, VirtualMemoryError> {
self.reserves.fetch_add(1, Ordering::Relaxed);
Ok(0x1000)
}
fn base(reservation: &Self::Reservation) -> usize {
*reservation
}
fn commit(
&self,
_reservation: &mut Self::Reservation,
_offset: usize,
_len: usize,
) -> Result<(), VirtualMemoryError> {
self.commits.fetch_add(1, Ordering::Relaxed);
Ok(())
}
fn release(
&self,
_reservation: &mut Self::Reservation,
_offset: usize,
_len: usize,
) -> Result<(), VirtualMemoryError> {
Ok(())
}
}
#[test]
fn backing_accounting_accepts_the_same_authority_without_double_charge() {
let governor = LedgerGovernor::new(LeaseLedger::new_for_device(
DeviceKey::device(2),
8192,
0,
0,
));
let commits = Arc::new(AtomicUsize::new(0));
let backing = AuthorityBacking {
authority: governor.authority_id(),
reserves: Arc::new(AtomicUsize::new(0)),
commits: Arc::clone(&commits),
};
let mut buffer = VirtualBuffer::with_backing(
backing,
8192,
Arc::new(governor.clone()),
Tier::Device,
MemoryRole::KvCache,
HOLDER,
)
.expect("matching authority");
buffer.grow_to(4096).expect("backing owns the charge");
assert_eq!(commits.load(Ordering::Relaxed), 1);
assert_eq!(
governor.used(Tier::Device),
0,
"mapped attribution must not add a second physical charge"
);
}
#[test]
fn backing_accounting_rejects_a_different_authority_before_reservation() {
let backing_governor = LedgerGovernor::new(LeaseLedger::new_for_device(
DeviceKey::device(2),
8192,
0,
0,
));
let buffer_governor = LedgerGovernor::new(LeaseLedger::new_for_device(
DeviceKey::device(2),
8192,
0,
0,
));
let reserves = Arc::new(AtomicUsize::new(0));
let commits = Arc::new(AtomicUsize::new(0));
let backing = AuthorityBacking {
authority: backing_governor.authority_id(),
reserves: Arc::clone(&reserves),
commits: Arc::clone(&commits),
};
let error = VirtualBuffer::with_backing(
backing,
8192,
Arc::new(buffer_governor.clone()),
Tier::Device,
MemoryRole::KvCache,
HOLDER,
)
.expect_err("different accounting authorities must be rejected");
assert!(matches!(
error,
VirtualBufferError::AuthorityMismatch { backing, governor }
if backing == backing_governor.authority_id()
&& governor == buffer_governor.authority_id()
));
assert_eq!(reserves.load(Ordering::Relaxed), 0);
assert_eq!(commits.load(Ordering::Relaxed), 0);
assert_eq!(buffer_governor.used(Tier::Device), 0);
}
#[test]
fn the_address_does_not_change_as_the_buffer_grows() {
let (mut buffer, _) = buffer(64 << 20, 128 << 20);
let base = buffer.as_ptr();
for target in [1usize, 4096, 1 << 20, 8 << 20, 32 << 20] {
buffer.grow_to(target).expect("within capacity and budget");
assert_eq!(
buffer.as_ptr(),
base,
"growing to {target} moved the buffer, which is the one thing it must not do"
);
assert_eq!(buffer.len(), target);
}
}
#[test]
fn reserving_address_space_leases_nothing() {
let (buffer, governor) = buffer(1 << 30, 1 << 20);
assert_eq!(
governor.available(Tier::Host),
1 << 20,
"a 1 GiB reservation must not consume a 1 MiB budget"
);
assert_eq!(buffer.committed(), 0);
assert!(buffer.is_empty());
}
#[test]
fn growth_leases_exactly_what_it_commits() {
let (mut buffer, governor) = buffer(16 << 20, 32 << 20);
buffer.grow_to(1).expect("granted");
assert_eq!(
buffer.committed(),
granularity(),
"growth commits whole granules"
);
assert_eq!(
(32u64 << 20) - governor.available(Tier::Host),
buffer.committed() as u64,
"the governor must be charged the committed bytes, not the requested ones"
);
let before = governor.available(Tier::Host);
buffer.grow_to(2).expect("granted");
assert_eq!(
governor.available(Tier::Host),
before,
"growing within an already-committed granule must not lease again"
);
}
#[test]
fn contents_survive_growth() {
let (mut buffer, _) = buffer(8 << 20, 16 << 20);
buffer.grow_to(4096).expect("granted");
unsafe { std::ptr::write_bytes(buffer.as_mut_ptr(), 0xC7, 4096) };
buffer.grow_to(4 << 20).expect("granted");
let head = unsafe { std::slice::from_raw_parts(buffer.as_ptr(), 4096) };
assert!(
head.iter().all(|&byte| byte == 0xC7),
"growth lost the bytes that were already there"
);
}
#[test]
fn shrinking_returns_committed_pages() {
let (mut buffer, governor) = buffer(16 << 20, 32 << 20);
buffer.grow_to(4 << 20).expect("granted");
let held = buffer.committed();
assert!(held >= 4 << 20);
buffer.shrink_to(0).expect("shrunk");
assert_eq!(buffer.committed(), 0, "everything must come back");
assert_eq!(
governor.available(Tier::Host),
32 << 20,
"the governor must see the pages returned"
);
assert_eq!(buffer.len(), 0);
}
#[test]
fn a_shrink_within_a_granule_keeps_the_page() {
let (mut buffer, _) = buffer(16 << 20, 32 << 20);
buffer.grow_to(granularity()).expect("granted");
let committed = buffer.committed();
buffer.shrink_to(granularity() - 1).expect("shrunk");
assert_eq!(
buffer.committed(),
committed,
"the granule containing the new end must stay mapped"
);
assert_eq!(buffer.len(), granularity() - 1);
}
#[test]
fn growing_past_the_reservation_is_refused_and_names_the_capacity() {
let (mut buffer, _) = buffer(1 << 20, 64 << 20);
let capacity = buffer.capacity();
let error = buffer
.grow_to(capacity + 1)
.expect_err("the reservation cannot be extended");
let message = error.to_string();
assert!(
message.contains("larger capacity"),
"the error must say what to do, got: {message}"
);
assert_eq!(buffer.len(), 0, "a refused growth must change nothing");
}
#[test]
fn a_refused_lease_leaves_the_buffer_untouched() {
let (mut buffer, governor) = buffer(64 << 20, granularity() as u64);
buffer
.grow_to(1)
.expect("the first granule fits the budget");
let committed = buffer.committed();
let error = buffer.grow_to(32 << 20);
assert!(error.is_err(), "a 32 MiB growth cannot fit one granule");
assert_eq!(
buffer.committed(),
committed,
"a refused growth must not commit pages"
);
assert_eq!(governor.available(Tier::Host), 0);
unsafe { std::ptr::write_bytes(buffer.as_mut_ptr(), 0x11, committed) };
}
#[test]
fn dropping_the_buffer_returns_its_budget() {
let (mut buffer, governor) = buffer(16 << 20, 32 << 20);
buffer.grow_to(8 << 20).expect("granted");
assert!(governor.available(Tier::Host) < 32 << 20);
drop(buffer);
assert_eq!(
governor.available(Tier::Host),
32 << 20,
"the lease must be released when the buffer goes"
);
}
}