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