wgpu_core/device/
ray_tracing.rs

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}