use ruda_kernel::dsl::server::IoError;
use ruda::runtime::storage::{ComputeStorage, StorageHandle, StorageId, StorageUtilization};
use hashbrown::HashMap;
use std::num::NonZeroU64;
use wgpu::BufferUsages;
const MIN_BUFFER_SIZE: u64 = 32;
pub struct WgpuStorage {
memory: HashMap<StorageId, wgpu::Buffer>,
device: wgpu::Device,
buffer_usages: BufferUsages,
mem_alignment: usize,
relocation_queue: Option<wgpu::Queue>,
}
impl core::fmt::Debug for WgpuStorage {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str(format!("WgpuStorage {{ device: {:?} }}", self.device).as_str())
}
}
#[derive(new, Debug)]
pub struct WgpuResource {
#[new(default)]
pin: Option<ruda::runtime::memory_management::MemoryResourcePin>,
pub buffer: wgpu::Buffer,
pub offset: u64,
pub size: u64,
}
impl WgpuResource {
pub(crate) fn address_pin(&self) -> Option<ruda::runtime::memory_management::MemoryResourcePin> {
self.pin.clone()
}
pub fn as_wgpu_bind_resource(&self) -> wgpu::BindingResource<'_> {
let size = NonZeroU64::new(self.size.next_multiple_of(4));
let binding = wgpu::BufferBinding {
buffer: &self.buffer,
offset: self.offset,
size,
};
wgpu::BindingResource::Buffer(binding)
}
}
impl WgpuStorage {
pub fn new(mem_alignment: usize, device: wgpu::Device, usages: BufferUsages) -> Self {
Self {
memory: HashMap::new(),
device,
buffer_usages: usages,
mem_alignment,
relocation_queue: None,
}
}
pub(crate) fn with_relocation_queue(mut self, queue: wgpu::Queue) -> Self {
self.relocation_queue = Some(queue);
self
}
fn encode_relocation_copy(&self, encoder: &mut wgpu::CommandEncoder, source: &StorageHandle, target: &StorageHandle) -> Result<(), IoError> {
let src = self.memory.get(&source.id).expect("relocation source");
let dst = self.memory.get(&target.id).expect("relocation target");
let size = source.size().next_multiple_of(wgpu::COPY_BUFFER_ALIGNMENT);
if source.offset() % wgpu::COPY_BUFFER_ALIGNMENT != 0 || target.offset() % wgpu::COPY_BUFFER_ALIGNMENT != 0
|| source.offset().checked_add(size).is_none_or(|end| end > src.size())
|| target.offset().checked_add(size).is_none_or(|end| end > dst.size()) {
return Err(IoError::UnsupportedIoOperation { backtrace: ruda_core::backtrace::BackTrace::capture() });
}
encoder.copy_buffer_to_buffer(src, source.offset(), dst, target.offset(), size);
Ok(())
}
fn wait_relocation(&self) -> Result<(), IoError> {
#[cfg(not(target_family = "wasm"))]
{
self.device.poll(wgpu::PollType::Wait { submission_index: None, timeout: None })
.map(|_| ()).map_err(|error| IoError::Unknown {
description: format!("WGPU relocation wait: {error}"),
backtrace: ruda_core::backtrace::BackTrace::capture(),
})
}
#[cfg(target_family = "wasm")]
{
Err(IoError::UnsupportedIoOperation { backtrace: ruda_core::backtrace::BackTrace::capture() })
}
}
}
impl ComputeStorage for WgpuStorage {
type Resource = WgpuResource;
fn alignment(&self) -> usize {
self.mem_alignment
}
fn get_pinned(&mut self, handle: &StorageHandle, binding: ruda::runtime::memory_management::ManagedMemoryBinding) -> Self::Resource {
let mut resource = self.get(handle);
resource.pin = Some(binding.pin());
resource
}
fn supports_relocation(&self) -> bool {
self.relocation_queue.is_some() && cfg!(not(target_family = "wasm"))
}
fn relocation_barrier(&mut self) -> Result<(), IoError> {
let queue = self.relocation_queue.as_ref().expect("relocation queue");
queue.submit([]);
self.wait_relocation()
}
fn relocation_copy(&mut self, source: &StorageHandle, target: &StorageHandle) -> Result<(), IoError> {
self.relocation_copy_batch(core::iter::once((source, target)))
}
fn relocation_copy_batch<'a>(
&mut self,
copies: impl IntoIterator<Item = (&'a StorageHandle, &'a StorageHandle)>,
) -> Result<(), IoError> {
let mut copies = copies.into_iter().peekable();
if copies.peek().is_none() { return Ok(()); }
let validation = self.device.push_error_scope(wgpu::ErrorFilter::Validation);
let allocation = self.device.push_error_scope(wgpu::ErrorFilter::OutOfMemory);
let internal = self.device.push_error_scope(wgpu::ErrorFilter::Internal);
let mut encoder = self.device.create_command_encoder(&wgpu::CommandEncoderDescriptor { label: Some("RUDA adaptive memory copies") });
let copied = copies.try_for_each(|(source, target)| self.encode_relocation_copy(&mut encoder, source, target));
self.relocation_queue.as_ref().expect("relocation queue").submit([encoder.finish()]);
let errors = [internal.pop(), allocation.pop(), validation.pop()];
#[cfg(not(target_family = "wasm"))]
{
let mut first = None;
for error in errors {
if let Some(error) = ruda_core::future::block_on(error) {
if first.is_none() { first = Some(error); }
}
}
copied?;
if let Some(error) = first {
return Err(IoError::Unknown { description: format!("WGPU relocation copy: {error}"),
backtrace: ruda_core::backtrace::BackTrace::capture() });
}
}
#[cfg(target_family = "wasm")]
{ let _ = errors; copied?; }
Ok(())
}
fn relocation_complete(&mut self) -> Result<(), IoError> { self.wait_relocation() }
fn get(&mut self, handle: &StorageHandle) -> Self::Resource {
let buffer = self.memory.get(&handle.id).unwrap();
WgpuResource::new(buffer.clone(), handle.offset(), handle.size())
}
#[cfg_attr(
feature = "tracing",
tracing::instrument(level = "trace", skip(self, size))
)]
fn alloc(&mut self, size: u64) -> Result<StorageHandle, IoError> {
let id = StorageId::new();
let alloc_size = size.max(MIN_BUFFER_SIZE);
let scopes = self.relocation_queue.as_ref().map(|_| [
self.device.push_error_scope(wgpu::ErrorFilter::Validation),
self.device.push_error_scope(wgpu::ErrorFilter::OutOfMemory),
self.device.push_error_scope(wgpu::ErrorFilter::Internal),
]);
let buffer = self.device.create_buffer(&wgpu::BufferDescriptor {
label: None,
size: alloc_size,
usage: self.buffer_usages,
mapped_at_creation: false,
});
if let Some(scopes) = scopes {
let errors = scopes.into_iter().rev().map(|scope| scope.pop()).collect::<Vec<_>>();
#[cfg(not(target_family = "wasm"))]
{
let mut first = None;
for error in errors {
if let Some(error) = ruda_core::future::block_on(error) {
if first.is_none() { first = Some(error); }
}
}
if let Some(error) = first {
return Err(IoError::Unknown { description: format!("WGPU adaptive allocation: {error}"),
backtrace: ruda_core::backtrace::BackTrace::capture() });
}
}
#[cfg(target_family = "wasm")]
let _ = errors;
}
self.memory.insert(id, buffer);
Ok(StorageHandle::new(
id,
StorageUtilization { offset: 0, size },
))
}
#[cfg_attr(feature = "tracing", tracing::instrument(level = "trace", skip(self)))]
fn dealloc(&mut self, id: StorageId) {
self.memory.remove(&id);
}
fn flush(&mut self) {
}
}