1use alloc::{sync::Arc, vec::Vec};
2use core::mem::{size_of, ManuallyDrop};
3
4#[cfg(feature = "trace")]
5use crate::device::trace::{Action, IntoTrace};
6use crate::device::DeviceError;
7use crate::resource::ResourceState;
8use crate::{
9 api_log,
10 device::Device,
11 global::Global,
12 hal_label,
13 id::{self, BlasId, TlasId},
14 lock::RwLock,
15 lock::{rank, Mutex},
16 ray_tracing::BlasPrepareCompactError,
17 ray_tracing::{CreateBlasError, CreateTlasError},
18 resource,
19 resource::{BlasCompactCallback, BlasCompactState, InvalidResourceError, TrackingData},
20 snatch::Snatchable,
21 LabelHelpers,
22};
23use hal::AccelerationStructureTriangleIndices;
24use wgt::{Features, AABB_GEOMETRY_MIN_STRIDE};
25
26impl Device {
27 pub fn create_blas(
28 self: &Arc<Self>,
29 blas_desc: &resource::BlasDescriptor,
30 sizes: wgt::BlasGeometrySizeDescriptors,
31 ) -> (Arc<resource::Blas>, Option<CreateBlasError>) {
32 #[cfg(feature = "trace")]
33 let trace_sizes = sizes.clone();
34
35 let (blas, error) = match self.create_blas_inner(blas_desc, sizes) {
36 Ok(blas) => (blas, None),
37 Err(err) => (resource::Blas::invalid(self.clone(), blas_desc), Some(err)),
38 };
39
40 #[cfg(feature = "trace")]
41 if let Some(trace) = self.trace.lock().as_mut() {
42 trace.add(Action::CreateBlas {
43 id: blas.to_trace(),
44 desc: blas_desc.clone(),
45 sizes: trace_sizes,
46 });
47 }
48
49 api_log!("Device::create_blas -> {:?}", Arc::as_ptr(&blas));
50 (blas, error)
51 }
52 pub(crate) fn create_blas_inner(
53 self: &Arc<Self>,
54 blas_desc: &resource::BlasDescriptor,
55 sizes: wgt::BlasGeometrySizeDescriptors,
56 ) -> Result<Arc<resource::Blas>, CreateBlasError> {
57 self.check_is_valid()?;
58 self.require_features(Features::EXPERIMENTAL_RAY_QUERY)?;
59
60 if blas_desc
61 .flags
62 .contains(wgt::AccelerationStructureFlags::ALLOW_RAY_HIT_VERTEX_RETURN)
63 {
64 self.require_features(Features::EXPERIMENTAL_RAY_HIT_VERTEX_RETURN)?;
65 }
66
67 let size_info = match &sizes {
68 wgt::BlasGeometrySizeDescriptors::Triangles { descriptors } => {
69 if descriptors.len() as u32 > self.limits.max_blas_geometry_count {
70 return Err(CreateBlasError::TooManyGeometries(
71 self.limits.max_blas_geometry_count,
72 descriptors.len() as u32,
73 ));
74 }
75
76 let mut entries =
77 Vec::<hal::AccelerationStructureTriangles<dyn hal::DynBuffer>>::with_capacity(
78 descriptors.len(),
79 );
80 for desc in descriptors {
81 if desc.index_count.is_some() != desc.index_format.is_some() {
82 return Err(CreateBlasError::MissingIndexData);
83 }
84 let indices =
85 desc.index_count
86 .map(|count| AccelerationStructureTriangleIndices::<
87 dyn hal::DynBuffer,
88 > {
89 format: desc.index_format.unwrap(),
90 buffer: Some(self.zero_buffer.as_ref()),
91 offset: 0,
92 count,
93 });
94 if !self
95 .features
96 .allowed_vertex_formats_for_blas()
97 .contains(&desc.vertex_format)
98 {
99 return Err(CreateBlasError::InvalidVertexFormat(
100 desc.vertex_format,
101 self.features.allowed_vertex_formats_for_blas(),
102 ));
103 }
104
105 let mut transform = None;
106
107 if blas_desc
108 .flags
109 .contains(wgt::AccelerationStructureFlags::USE_TRANSFORM)
110 {
111 transform = Some(wgpu_hal::AccelerationStructureTriangleTransform {
112 buffer: self.zero_buffer.as_ref(),
113 offset: 0,
114 })
115 }
116
117 if desc.vertex_count > self.limits.max_blas_primitive_count {
118 return Err(CreateBlasError::TooManyPrimitives(
119 self.limits.max_blas_primitive_count,
120 desc.vertex_count,
121 ));
122 }
123
124 entries.push(hal::AccelerationStructureTriangles::<dyn hal::DynBuffer> {
125 vertex_buffer: Some(self.zero_buffer.as_ref()),
126 vertex_format: desc.vertex_format,
127 first_vertex: 0,
128 vertex_count: desc.vertex_count,
129 vertex_stride: 0,
130 indices,
131 transform,
132 flags: desc.flags,
133 });
134 }
135 unsafe {
136 self.raw().get_acceleration_structure_build_sizes(
137 &hal::GetAccelerationStructureBuildSizesDescriptor {
138 entries: &hal::AccelerationStructureEntries::Triangles(entries),
139 flags: blas_desc.flags,
140 },
141 )
142 }
143 }
144 wgt::BlasGeometrySizeDescriptors::AABBs { descriptors } => {
145 if descriptors.len() as u32 > self.limits.max_blas_geometry_count {
146 return Err(CreateBlasError::TooManyGeometries(
147 self.limits.max_blas_geometry_count,
148 descriptors.len() as u32,
149 ));
150 }
151
152 let mut entries =
153 Vec::<hal::AccelerationStructureAABBs<dyn hal::DynBuffer>>::with_capacity(
154 descriptors.len(),
155 );
156 for desc in descriptors {
157 if desc.primitive_count > self.limits.max_blas_primitive_count {
158 return Err(CreateBlasError::TooManyPrimitives(
159 self.limits.max_blas_primitive_count,
160 desc.primitive_count,
161 ));
162 }
163
164 entries.push(hal::AccelerationStructureAABBs::<dyn hal::DynBuffer> {
165 buffer: Some(self.zero_buffer.as_ref()),
166 offset: 0,
167 count: desc.primitive_count,
168 stride: AABB_GEOMETRY_MIN_STRIDE,
169 flags: desc.flags,
170 });
171 }
172 unsafe {
173 self.raw().get_acceleration_structure_build_sizes(
174 &hal::GetAccelerationStructureBuildSizesDescriptor {
175 entries: &hal::AccelerationStructureEntries::AABBs(entries),
176 flags: blas_desc.flags,
177 },
178 )
179 }
180 }
181 };
182
183 let raw = unsafe {
184 self.raw()
185 .create_acceleration_structure(&hal::AccelerationStructureDescriptor {
186 label: blas_desc.label.as_deref(),
187 size: size_info.acceleration_structure_size,
188 format: hal::AccelerationStructureFormat::BottomLevel,
189 allow_compaction: blas_desc
190 .flags
191 .contains(wgpu_types::AccelerationStructureFlags::ALLOW_COMPACTION),
192 })
193 }
194 .map_err(|e| self.handle_hal_error_with_nonfatal_oom(e))?;
195
196 let compaction_buffer = if blas_desc
197 .flags
198 .contains(wgpu_types::AccelerationStructureFlags::ALLOW_COMPACTION)
199 {
200 Some(ManuallyDrop::new(unsafe {
201 self.raw()
202 .create_buffer(&hal::BufferDescriptor {
203 label: Some("(wgpu internal) compaction read-back buffer"),
204 size: size_of::<wgpu_types::BufferAddress>() as wgpu_types::BufferAddress,
205 usage: wgpu_types::BufferUses::ACCELERATION_STRUCTURE_QUERY
206 | wgpu_types::BufferUses::MAP_READ,
207 memory_flags: hal::MemoryFlags::PREFER_COHERENT,
208 })
209 .map_err(DeviceError::from_hal)?
210 }))
211 } else {
212 None
213 };
214
215 let handle = unsafe {
216 self.raw()
217 .get_acceleration_structure_device_address(raw.as_ref())
218 };
219
220 Ok(Arc::new(resource::Blas {
221 state: ResourceState::Valid(resource::BlasState {
222 raw: Snatchable::new(raw),
223 }),
224 device: self.clone(),
225 size_info,
226 sizes,
227 flags: blas_desc.flags,
228 update_mode: blas_desc.update_mode,
229 handle,
230 label: blas_desc.label.to_string(),
231 built_index: RwLock::new(rank::BLAS_BUILT_INDEX, None),
232 tracking_data: TrackingData::new(self.tracker_indices.blas_s.clone()),
233 compaction_buffer,
234 compacted_state: Mutex::new(rank::BLAS_COMPACTION_STATE, BlasCompactState::Idle),
235 }))
236 }
237
238 pub fn create_tlas(
239 self: &Arc<Self>,
240 desc: &resource::TlasDescriptor,
241 ) -> (Arc<resource::Tlas>, Option<CreateTlasError>) {
242 let (tlas, error) = match self.create_tlas_inner(desc) {
243 Ok(tlas) => (tlas, None),
244 Err(e) => (resource::Tlas::invalid(Arc::clone(self), desc), Some(e)),
245 };
246 #[cfg(feature = "trace")]
247 if let Some(trace) = self.trace.lock().as_mut() {
248 trace.add(Action::CreateTlas {
249 id: tlas.to_trace(),
250 desc: desc.clone(),
251 });
252 }
253
254 api_log!("Device::create_tlas -> {:?}", Arc::as_ptr(&tlas));
255
256 (tlas, error)
257 }
258
259 pub(crate) fn create_tlas_inner(
260 self: &Arc<Self>,
261 desc: &resource::TlasDescriptor,
262 ) -> Result<Arc<resource::Tlas>, CreateTlasError> {
263 self.check_is_valid()?;
264 self.require_features(Features::EXPERIMENTAL_RAY_QUERY)?;
265
266 if desc.max_instances > self.limits.max_tlas_instance_count {
267 return Err(CreateTlasError::TooManyInstances(
268 self.limits.max_tlas_instance_count,
269 desc.max_instances,
270 ));
271 }
272
273 if desc
274 .flags
275 .contains(wgt::AccelerationStructureFlags::USE_TRANSFORM)
276 {
277 return Err(CreateTlasError::DisallowedFlag(
278 wgt::AccelerationStructureFlags::USE_TRANSFORM,
279 ));
280 }
281
282 if desc
283 .flags
284 .contains(wgt::AccelerationStructureFlags::ALLOW_RAY_HIT_VERTEX_RETURN)
285 {
286 self.require_features(Features::EXPERIMENTAL_RAY_HIT_VERTEX_RETURN)?;
287 }
288
289 let size_info = unsafe {
290 self.raw().get_acceleration_structure_build_sizes(
291 &hal::GetAccelerationStructureBuildSizesDescriptor {
292 entries: &hal::AccelerationStructureEntries::Instances(
293 hal::AccelerationStructureInstances {
294 buffer: Some(self.zero_buffer.as_ref()),
295 offset: 0,
296 count: desc.max_instances,
297 },
298 ),
299 flags: desc.flags,
300 },
301 )
302 };
303
304 let raw = unsafe {
305 self.raw()
306 .create_acceleration_structure(&hal::AccelerationStructureDescriptor {
307 label: desc.label.as_deref(),
308 size: size_info.acceleration_structure_size,
309 format: hal::AccelerationStructureFormat::TopLevel,
310 allow_compaction: false,
311 })
312 }
313 .map_err(|e| self.handle_hal_error_with_nonfatal_oom(e))?;
314
315 let instance_buffer_size = self
316 .alignments
317 .raw_tlas_instance_size
318 .checked_mul(desc.max_instances.max(1))
319 .expect("max_tlas_instance_count should not allow excessive buffer size");
320 let instance_buffer = unsafe {
321 self.raw().create_buffer(&hal::BufferDescriptor {
322 label: hal_label(Some("(wgpu-core) instances_buffer"), self.instance_flags),
323 size: u64::from(instance_buffer_size),
324 usage: wgt::BufferUses::COPY_DST
325 | wgt::BufferUses::TOP_LEVEL_ACCELERATION_STRUCTURE_INPUT,
326 memory_flags: hal::MemoryFlags::PREFER_COHERENT,
327 })
328 }
329 .map_err(|e| self.handle_hal_error_with_nonfatal_oom(e))?;
330
331 Ok(Arc::new(resource::Tlas {
332 state: ResourceState::Valid(resource::TlasState {
333 raw: Snatchable::new(raw),
334 instance_buffer,
335 }),
336 device: self.clone(),
337 size_info,
338 flags: desc.flags,
339 update_mode: desc.update_mode,
340 built_index: RwLock::new(rank::TLAS_BUILT_INDEX, None),
341 dependencies: RwLock::new(rank::TLAS_DEPENDENCIES, Vec::new()),
342 label: desc.label.to_string(),
343 max_instance_count: desc.max_instances,
344 tracking_data: TrackingData::new(self.tracker_indices.tlas_s.clone()),
345 }))
346 }
347}
348
349impl Global {
350 pub fn device_create_blas(
351 &self,
352 device_id: id::DeviceId,
353 desc: &resource::BlasDescriptor,
354 sizes: wgt::BlasGeometrySizeDescriptors,
355 id_in: Option<BlasId>,
356 ) -> (BlasId, Option<u64>, Option<CreateBlasError>) {
357 profiling::scope!("Device::create_blas");
358
359 let fid = self.hub.blas_s.prepare(id_in);
360
361 let device = self.hub.devices.get(device_id);
362
363 let (blas, error) = device.create_blas(desc, sizes);
364
365 let handle = blas.handle();
366
367 let id = fid.assign(blas);
368
369 (id, handle, error)
370 }
371
372 pub fn device_create_tlas(
373 &self,
374 device_id: id::DeviceId,
375 desc: &resource::TlasDescriptor,
376 id_in: Option<TlasId>,
377 ) -> (TlasId, Option<CreateTlasError>) {
378 profiling::scope!("Device::create_tlas");
379
380 let fid = self.hub.tlas_s.prepare(id_in);
381
382 let device = self.hub.devices.get(device_id);
383
384 let (tlas, error) = device.create_tlas(desc);
385
386 let id = fid.assign(tlas);
387
388 (id, error)
389 }
390
391 pub fn blas_drop(&self, blas_id: BlasId) {
392 let _blas = self.hub.blas_s.remove(blas_id);
393 }
394
395 pub fn tlas_drop(&self, tlas_id: TlasId) {
396 let _tlas = self.hub.tlas_s.remove(tlas_id);
397 }
398
399 pub fn blas_prepare_compact_async(
400 &self,
401 blas_id: BlasId,
402 callback: Option<BlasCompactCallback>,
403 ) -> Result<crate::SubmissionIndex, BlasPrepareCompactError> {
404 let hub = &self.hub;
405
406 let blas = hub.blas_s.get(blas_id);
407
408 blas.prepare_compact_async(callback)
409 }
410
411 pub fn ready_for_compaction(&self, blas_id: BlasId) -> Result<bool, InvalidResourceError> {
412 let hub = &self.hub;
413
414 let blas = hub.blas_s.get(blas_id);
415
416 blas.ready_for_compaction()
417 }
418}