use cubecl_core::server::IoError;
use cubecl_environment::collections::HashMap;
use cubecl_runtime::storage::{ComputeStorage, StorageHandle, StorageId, StorageUtilization};
use std::num::NonZeroU64;
use wgpu::BufferUsages;
const MIN_BUFFER_SIZE: u64 = 32;
pub struct WgpuStorage {
memory: HashMap<StorageId, WgpuMemory>,
device: wgpu::Device,
buffer_usages: BufferUsages,
mem_alignment: usize,
#[allow(unused, reason = "keep it simple")]
vk_storage: bool,
}
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 {
pub buffer: wgpu::Buffer,
pub address: Option<NonZeroU64>,
pub offset: u64,
pub size: u64,
}
#[derive(new, Debug)]
pub struct WgpuMemory {
pub buffer: wgpu::Buffer,
pub address: Option<NonZeroU64>,
}
impl WgpuResource {
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,
vk_storage: bool,
) -> Self {
Self {
memory: HashMap::new(),
device,
buffer_usages: usages,
mem_alignment,
vk_storage,
}
}
}
impl ComputeStorage for WgpuStorage {
type Resource = WgpuResource;
fn alignment(&self) -> usize {
self.mem_alignment
}
fn get(&mut self, handle: &StorageHandle) -> Self::Resource {
let memory = self.memory.get(&handle.id).unwrap();
WgpuResource::new(
memory.buffer.clone(),
memory.address,
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 memory = self.create_buffer(&wgpu::BufferDescriptor {
label: None,
size: alloc_size,
usage: self.buffer_usages,
mapped_at_creation: false,
})?;
self.memory.insert(id, memory);
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) {
}
}
impl WgpuStorage {
#[cfg(feature = "spirv")]
fn create_buffer(&self, desc: &wgpu::BufferDescriptor<'_>) -> Result<WgpuMemory, IoError> {
if self.vk_storage {
let (buffer, addr) = crate::backend::vulkan::create_storage_buffer(&self.device, desc)?;
Ok(WgpuMemory::new(buffer, NonZeroU64::new(addr)))
} else {
Ok(WgpuMemory::new(self.device.create_buffer(desc), None))
}
}
#[cfg(not(feature = "spirv"))]
fn create_buffer(&self, desc: &wgpu::BufferDescriptor<'_>) -> Result<WgpuMemory, IoError> {
Ok(WgpuMemory::new(self.device.create_buffer(desc), None))
}
}