cubecl_wgpu/compute/
storage.rs1use cubecl_core::server::IoError;
2use cubecl_environment::backtrace::BackTrace;
3use cubecl_environment::collections::HashMap;
4use cubecl_runtime::storage::{ComputeStorage, StorageHandle, StorageId, StorageUtilization};
5use std::num::NonZeroU64;
6use wgpu::BufferUsages;
7
8const MIN_BUFFER_SIZE: u64 = 32;
12
13pub struct WgpuStorage {
15 memory: HashMap<StorageId, WgpuMemory>,
16 device: wgpu::Device,
17 buffer_usages: BufferUsages,
18 mem_alignment: usize,
19 #[allow(unused, reason = "keep it simple")]
20 vk_storage: bool,
21}
22
23impl core::fmt::Debug for WgpuStorage {
24 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
25 f.write_str(format!("WgpuStorage {{ device: {:?} }}", self.device).as_str())
26 }
27}
28
29#[derive(new, Debug, Clone)]
31pub struct WgpuResource {
32 pub buffer: wgpu::Buffer,
34 pub address: Option<NonZeroU64>,
36 pub offset: u64,
38 pub size: u64,
44}
45
46#[derive(new, Debug)]
48pub struct WgpuMemory {
49 pub buffer: wgpu::Buffer,
51 pub address: Option<NonZeroU64>,
53}
54
55impl WgpuResource {
56 pub fn as_wgpu_bind_resource(&self) -> wgpu::BindingResource<'_> {
58 let size = NonZeroU64::new(self.size.next_multiple_of(4));
68
69 let binding = wgpu::BufferBinding {
70 buffer: &self.buffer,
71 offset: self.offset,
72 size,
73 };
74 wgpu::BindingResource::Buffer(binding)
75 }
76}
77
78impl WgpuStorage {
80 pub fn new(
82 mem_alignment: usize,
83 device: wgpu::Device,
84 usages: BufferUsages,
85 vk_storage: bool,
86 ) -> Self {
87 Self {
88 memory: HashMap::new(),
89 device,
90 buffer_usages: usages,
91 mem_alignment,
92 vk_storage,
93 }
94 }
95}
96
97impl ComputeStorage for WgpuStorage {
98 type Resource = WgpuResource;
99
100 fn alignment(&self) -> usize {
101 self.mem_alignment
102 }
103
104 fn get(&mut self, handle: &StorageHandle) -> Result<Self::Resource, IoError> {
105 let memory = self
106 .memory
107 .get(&handle.id)
108 .ok_or_else(|| IoError::StorageHandleNotFound {
109 reason: format!("{} in the wgpu buffer storage", handle.id).into(),
110 backtrace: BackTrace::capture(),
111 })?;
112 Ok(WgpuResource::new(
113 memory.buffer.clone(),
114 memory.address,
115 handle.offset(),
116 handle.size(),
117 ))
118 }
119
120 #[cfg_attr(
121 feature = "tracing",
122 tracing::instrument(level = "trace", skip(self, size))
123 )]
124 fn alloc(&mut self, size: u64) -> Result<StorageHandle, IoError> {
125 let id = StorageId::new();
126
127 let alloc_size = size.max(MIN_BUFFER_SIZE);
128
129 let memory = self.create_buffer(&wgpu::BufferDescriptor {
130 label: None,
131 size: alloc_size,
132 usage: self.buffer_usages,
133 mapped_at_creation: false,
134 })?;
135
136 self.memory.insert(id, memory);
137 Ok(StorageHandle::new(
138 id,
139 StorageUtilization { offset: 0, size },
140 ))
141 }
142
143 #[cfg_attr(feature = "tracing", tracing::instrument(level = "trace", skip(self)))]
144 fn dealloc(&mut self, id: StorageId) {
145 self.memory.remove(&id);
146 }
147
148 fn flush(&mut self) {
149 }
151}
152
153impl WgpuStorage {
154 #[cfg(feature = "spirv")]
155 fn create_buffer(&self, desc: &wgpu::BufferDescriptor<'_>) -> Result<WgpuMemory, IoError> {
156 if self.vk_storage {
157 let (buffer, addr) = crate::backend::vulkan::create_storage_buffer(&self.device, desc)?;
163 Ok(WgpuMemory::new(buffer, NonZeroU64::new(addr)))
164 } else {
165 Ok(WgpuMemory::new(self.device.create_buffer(desc), None))
166 }
167 }
168
169 #[cfg(not(feature = "spirv"))]
170 fn create_buffer(&self, desc: &wgpu::BufferDescriptor<'_>) -> Result<WgpuMemory, IoError> {
171 Ok(WgpuMemory::new(self.device.create_buffer(desc), None))
172 }
173}