use std::sync::atomic::{AtomicU32, AtomicU64, Ordering};
use std::sync::{Arc, Weak};
use parking_lot::Mutex;
use smallvec::SmallVec;
use crate::{BackendKind, Error};
#[derive(Debug, Default, Copy, Clone)]
pub struct BudgetCaps {
pub soft_cap_bytes: Option<u64>,
pub hard_cap_bytes: Option<u64>,
}
#[derive(Debug, Default, Clone)]
pub struct BudgetCapsSet {
pub cpu: BudgetCaps,
#[cfg(feature = "wgpu")]
pub wgpu: BudgetCaps,
#[cfg(feature = "opencl")]
pub opencl: BudgetCaps,
#[cfg(feature = "cuda")]
pub cuda: BudgetCaps,
}
impl BudgetCapsSet {
fn for_backend(&self, backend: BackendKind) -> BudgetCaps {
match backend {
BackendKind::Cpu => self.cpu,
#[cfg(feature = "wgpu")]
BackendKind::Wgpu => self.wgpu,
#[cfg(feature = "opencl")]
BackendKind::OpenCl => self.opencl,
#[cfg(feature = "cuda")]
BackendKind::Cuda => self.cuda,
#[cfg(feature = "wgpu")]
BackendKind::Vulkan
| BackendKind::D3D12
| BackendKind::D3D11
| BackendKind::Metal
| BackendKind::OpenGL => self.wgpu,
#[cfg(not(feature = "wgpu"))]
BackendKind::Vulkan
| BackendKind::D3D12
| BackendKind::D3D11
| BackendKind::Metal
| BackendKind::OpenGL => self.cpu,
#[cfg(all(target_os = "android", feature = "wgpu"))]
BackendKind::AHardwareBuffer
| BackendKind::AndroidPresentation
| BackendKind::AndroidSurfaceControl
| BackendKind::AImageWriter
| BackendKind::MediaCodec
| BackendKind::ForeignGl
| BackendKind::ForeignVulkan => self.wgpu,
#[cfg(all(target_os = "android", not(feature = "wgpu")))]
BackendKind::AHardwareBuffer
| BackendKind::AndroidPresentation
| BackendKind::AndroidSurfaceControl
| BackendKind::AImageWriter
| BackendKind::MediaCodec
| BackendKind::ForeignGl
| BackendKind::ForeignVulkan => self.cpu,
#[cfg(all(feature = "web-codecs", feature = "wgpu"))]
BackendKind::WebCodecs => self.wgpu,
#[cfg(all(feature = "web-codecs", not(feature = "wgpu")))]
BackendKind::WebCodecs => self.cpu,
}
}
}
#[derive(Default)]
struct BudgetCountersSet {
cpu: AtomicU64,
#[cfg(feature = "wgpu")]
wgpu: AtomicU64,
#[cfg(feature = "opencl")]
opencl: AtomicU64,
#[cfg(feature = "cuda")]
cuda: AtomicU64,
}
impl BudgetCountersSet {
fn get(&self, backend: BackendKind) -> &AtomicU64 {
match backend {
BackendKind::Cpu => &self.cpu,
#[cfg(feature = "wgpu")]
BackendKind::Wgpu => &self.wgpu,
#[cfg(feature = "opencl")]
BackendKind::OpenCl => &self.opencl,
#[cfg(feature = "cuda")]
BackendKind::Cuda => &self.cuda,
#[cfg(feature = "wgpu")]
BackendKind::Vulkan
| BackendKind::D3D12
| BackendKind::D3D11
| BackendKind::Metal
| BackendKind::OpenGL => &self.wgpu,
#[cfg(not(feature = "wgpu"))]
BackendKind::Vulkan
| BackendKind::D3D12
| BackendKind::D3D11
| BackendKind::Metal
| BackendKind::OpenGL => &self.cpu,
#[cfg(all(target_os = "android", feature = "wgpu"))]
BackendKind::AHardwareBuffer
| BackendKind::AndroidPresentation
| BackendKind::AndroidSurfaceControl
| BackendKind::AImageWriter
| BackendKind::MediaCodec
| BackendKind::ForeignGl
| BackendKind::ForeignVulkan => &self.wgpu,
#[cfg(all(target_os = "android", not(feature = "wgpu")))]
BackendKind::AHardwareBuffer
| BackendKind::AndroidPresentation
| BackendKind::AndroidSurfaceControl
| BackendKind::AImageWriter
| BackendKind::MediaCodec
| BackendKind::ForeignGl
| BackendKind::ForeignVulkan => &self.cpu,
#[cfg(all(feature = "web-codecs", feature = "wgpu"))]
BackendKind::WebCodecs => &self.wgpu,
#[cfg(all(feature = "web-codecs", not(feature = "wgpu")))]
BackendKind::WebCodecs => &self.cpu,
}
}
}
#[derive(Debug, Copy, Clone, PartialEq, Eq)]
pub enum BudgetPressure {
Ok,
Soft,
Hard,
}
#[derive(Debug, Copy, Clone)]
pub struct BudgetPressureEvent {
pub backend: BackendKind,
pub level: BudgetPressure,
pub current_usage_bytes: u64,
pub required_bytes: u64,
pub cap_bytes: u64,
}
#[cfg(not(target_family = "wasm"))]
pub type PressureCallback = Arc<dyn Fn(&BudgetPressureEvent) + Send + Sync>;
#[cfg(target_family = "wasm")]
pub type PressureCallback = Arc<dyn Fn(&BudgetPressureEvent)>;
pub trait EvictablePool: crate::MaybeSendSync {
fn evict_unused(&self) -> u64;
fn bytes_in_use(&self) -> u64;
}
#[derive(Debug, Copy, Clone, PartialEq, Eq, Hash)]
pub struct PoolHandle(pub u32);
#[derive(Debug, Copy, Clone, PartialEq, Eq, Hash)]
pub struct PressureCallbackHandle(pub u32);
type PoolEntry = (PoolHandle, BackendKind, Weak<dyn EvictablePool>);
type CallbackEntry = (PressureCallbackHandle, PressureCallback);
pub struct MemoryBudget {
caps: BudgetCapsSet,
usage: BudgetCountersSet,
pools: Mutex<SmallVec<[PoolEntry; 4]>>,
callbacks: Mutex<SmallVec<[CallbackEntry; 4]>>,
next_pool_handle: AtomicU32,
next_callback_handle: AtomicU32,
}
impl core::fmt::Debug for MemoryBudget {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
f.debug_struct("MemoryBudget").field("caps", &self.caps).finish_non_exhaustive()
}
}
pub struct BudgetReservation {
budget: Arc<MemoryBudget>,
backend: BackendKind,
bytes: u64,
}
impl BudgetReservation {
pub fn backend(&self) -> BackendKind {
self.backend
}
pub fn bytes(&self) -> u64 {
self.bytes
}
}
impl Drop for BudgetReservation {
fn drop(&mut self) {
self.budget.usage.get(self.backend).fetch_sub(self.bytes, Ordering::Release);
}
}
impl core::fmt::Debug for BudgetReservation {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
f.debug_struct("BudgetReservation").field("backend", &self.backend).field("bytes", &self.bytes).finish()
}
}
impl MemoryBudget {
pub fn new(caps: BudgetCapsSet) -> Arc<Self> {
Arc::new(Self {
caps,
usage: BudgetCountersSet::default(),
pools: Mutex::new(SmallVec::new()),
callbacks: Mutex::new(SmallVec::new()),
next_pool_handle: AtomicU32::new(1),
next_callback_handle: AtomicU32::new(1),
})
}
pub fn caps(&self, backend: BackendKind) -> BudgetCaps {
self.caps.for_backend(backend)
}
pub fn current_usage(&self, backend: BackendKind) -> u64 {
self.usage.get(backend).load(Ordering::Acquire)
}
pub fn available(&self, backend: BackendKind) -> u64 {
let cap = self.caps.for_backend(backend).hard_cap_bytes.unwrap_or(u64::MAX);
cap.saturating_sub(self.current_usage(backend))
}
pub fn pressure(&self, backend: BackendKind, required: u64) -> BudgetPressure {
let caps = self.caps.for_backend(backend);
let usage = self.current_usage(backend);
let projected = usage.saturating_add(required);
if let Some(hard) = caps.hard_cap_bytes
&& projected > hard
{
return BudgetPressure::Hard;
}
if let Some(soft) = caps.soft_cap_bytes
&& projected > soft
{
return BudgetPressure::Soft;
}
BudgetPressure::Ok
}
pub fn try_reserve(self: &Arc<Self>, backend: BackendKind, bytes: u64) -> Result<BudgetReservation, Error> {
let caps = self.caps.for_backend(backend);
if caps.soft_cap_bytes.is_none() && caps.hard_cap_bytes.is_none() {
self.usage.get(backend).fetch_add(bytes, Ordering::Release);
return Ok(BudgetReservation { budget: self.clone(), backend, bytes });
}
let level = self.pressure(backend, bytes);
if level != BudgetPressure::Ok {
let cap = caps.hard_cap_bytes.or(caps.soft_cap_bytes).unwrap_or(u64::MAX);
let event = BudgetPressureEvent {
backend,
level,
current_usage_bytes: self.current_usage(backend),
required_bytes: bytes,
cap_bytes: cap,
};
{
let cbs = self.callbacks.lock().clone();
for (_, cb) in cbs.iter() {
let cb = cb.clone();
let _ = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| cb(&event)));
}
}
{
let pools = self.pools.lock().clone();
for (_, bk, weak) in pools.iter() {
if *bk != backend {
continue;
}
if let Some(pool) = weak.upgrade() {
let _ = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| pool.evict_unused()));
}
}
}
}
if self.try_commit(backend, bytes) {
return Ok(BudgetReservation { budget: self.clone(), backend, bytes });
}
let available = self.available(backend);
Err(Error::OutOfGpuMemory { required_bytes: bytes, available_bytes: available, backend })
}
fn try_commit(&self, backend: BackendKind, bytes: u64) -> bool {
let counter = self.usage.get(backend);
let hard = self.caps.for_backend(backend).hard_cap_bytes;
loop {
let cur = counter.load(Ordering::Acquire);
let next = cur.saturating_add(bytes);
if let Some(h) = hard
&& next > h
{
return false;
}
if counter.compare_exchange(cur, next, Ordering::AcqRel, Ordering::Acquire).is_ok() {
return true;
}
}
}
pub fn register_pool(&self, backend: BackendKind, pool: &Arc<dyn EvictablePool>) -> PoolHandle {
let h = PoolHandle(self.next_pool_handle.fetch_add(1, Ordering::Relaxed));
self.pools.lock().push((h, backend, Arc::downgrade(pool)));
h
}
pub fn unregister_pool(&self, handle: PoolHandle) {
self.pools.lock().retain(|(h, _, _)| *h != handle);
}
pub fn register_pressure_callback(&self, cb: PressureCallback) -> PressureCallbackHandle {
let h = PressureCallbackHandle(self.next_callback_handle.fetch_add(1, Ordering::Relaxed));
self.callbacks.lock().push((h, cb));
h
}
pub fn unregister_pressure_callback(&self, handle: PressureCallbackHandle) {
self.callbacks.lock().retain(|(h, _)| *h != handle);
}
}