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    hal_label,
12    lock::RwLock,
13    lock::{rank, Mutex},
14    ray_tracing::{CreateBlasError, CreateTlasError},
15    resource,
16    resource::{BlasCompactState, TrackingData},
17    snatch::Snatchable,
18    LabelHelpers,
19};
20use hal::AccelerationStructureTriangleIndices;
21use wgt::{Features, AABB_GEOMETRY_MIN_STRIDE};
22
23impl Device {
24    pub fn create_blas(
25        self: &Arc<Self>,
26        blas_desc: &resource::BlasDescriptor,
27        sizes: wgt::BlasGeometrySizeDescriptors,
28    ) -> (Arc<resource::Blas>, Option<CreateBlasError>) {
29        #[cfg(feature = "trace")]
30        let trace_sizes = sizes.clone();
31
32        let (blas, error) = match self.create_blas_inner(blas_desc, sizes) {
33            Ok(blas) => (blas, None),
34            Err(err) => (resource::Blas::invalid(self.clone(), blas_desc), Some(err)),
35        };
36
37        #[cfg(feature = "trace")]
38        if let Some(trace) = self.trace.lock().as_mut() {
39            trace.add(Action::CreateBlas {
40                id: blas.to_trace(),
41                desc: blas_desc.clone(),
42                sizes: trace_sizes,
43            });
44        }
45
46        api_log!("Device::create_blas -> {:?}", Arc::as_ptr(&blas));
47        (blas, error)
48    }
49    pub(crate) fn create_blas_inner(
50        self: &Arc<Self>,
51        blas_desc: &resource::BlasDescriptor,
52        sizes: wgt::BlasGeometrySizeDescriptors,
53    ) -> Result<Arc<resource::Blas>, CreateBlasError> {
54        self.check_is_valid()?;
55        self.require_features(Features::EXPERIMENTAL_RAY_QUERY)?;
56
57        if blas_desc
58            .flags
59            .contains(wgt::AccelerationStructureFlags::ALLOW_RAY_HIT_VERTEX_RETURN)
60        {
61            self.require_features(Features::EXPERIMENTAL_RAY_HIT_VERTEX_RETURN)?;
62        }
63
64        let size_info = match &sizes {
65            wgt::BlasGeometrySizeDescriptors::Triangles { descriptors } => {
66                if descriptors.len() as u32 > self.limits.max_blas_geometry_count {
67                    return Err(CreateBlasError::TooManyGeometries(
68                        self.limits.max_blas_geometry_count,
69                        descriptors.len() as u32,
70                    ));
71                }
72
73                let mut entries =
74                    Vec::<hal::AccelerationStructureTriangles<dyn hal::DynBuffer>>::with_capacity(
75                        descriptors.len(),
76                    );
77                for desc in descriptors {
78                    if desc.index_count.is_some() != desc.index_format.is_some() {
79                        return Err(CreateBlasError::MissingIndexData);
80                    }
81                    let indices =
82                        desc.index_count
83                            .map(|count| AccelerationStructureTriangleIndices::<
84                                dyn hal::DynBuffer,
85                            > {
86                                format: desc.index_format.unwrap(),
87                                buffer: Some(self.zero_buffer.as_ref()),
88                                offset: 0,
89                                count,
90                            });
91                    if !self
92                        .features
93                        .allowed_vertex_formats_for_blas()
94                        .contains(&desc.vertex_format)
95                    {
96                        return Err(CreateBlasError::InvalidVertexFormat(
97                            desc.vertex_format,
98                            self.features.allowed_vertex_formats_for_blas(),
99                        ));
100                    }
101
102                    let mut transform = None;
103
104                    if blas_desc
105                        .flags
106                        .contains(wgt::AccelerationStructureFlags::USE_TRANSFORM)
107                    {
108                        transform = Some(wgpu_hal::AccelerationStructureTriangleTransform {
109                            buffer: self.zero_buffer.as_ref(),
110                            offset: 0,
111                        })
112                    }
113
114                    if desc.vertex_count > self.limits.max_blas_primitive_count {
115                        return Err(CreateBlasError::TooManyPrimitives(
116                            self.limits.max_blas_primitive_count,
117                            desc.vertex_count,
118                        ));
119                    }
120
121                    entries.push(hal::AccelerationStructureTriangles::<dyn hal::DynBuffer> {
122                        vertex_buffer: Some(self.zero_buffer.as_ref()),
123                        vertex_format: desc.vertex_format,
124                        first_vertex: 0,
125                        vertex_count: desc.vertex_count,
126                        vertex_stride: 0,
127                        indices,
128                        transform,
129                        flags: desc.flags,
130                    });
131                }
132                unsafe {
133                    self.raw().get_acceleration_structure_build_sizes(
134                        &hal::GetAccelerationStructureBuildSizesDescriptor {
135                            entries: &hal::AccelerationStructureEntries::Triangles(entries),
136                            flags: blas_desc.flags,
137                        },
138                    )
139                }
140            }
141            wgt::BlasGeometrySizeDescriptors::AABBs { descriptors } => {
142                if descriptors.len() as u32 > self.limits.max_blas_geometry_count {
143                    return Err(CreateBlasError::TooManyGeometries(
144                        self.limits.max_blas_geometry_count,
145                        descriptors.len() as u32,
146                    ));
147                }
148
149                let mut entries =
150                    Vec::<hal::AccelerationStructureAABBs<dyn hal::DynBuffer>>::with_capacity(
151                        descriptors.len(),
152                    );
153                for desc in descriptors {
154                    if desc.primitive_count > self.limits.max_blas_primitive_count {
155                        return Err(CreateBlasError::TooManyPrimitives(
156                            self.limits.max_blas_primitive_count,
157                            desc.primitive_count,
158                        ));
159                    }
160
161                    entries.push(hal::AccelerationStructureAABBs::<dyn hal::DynBuffer> {
162                        buffer: Some(self.zero_buffer.as_ref()),
163                        offset: 0,
164                        count: desc.primitive_count,
165                        stride: AABB_GEOMETRY_MIN_STRIDE,
166                        flags: desc.flags,
167                    });
168                }
169                unsafe {
170                    self.raw().get_acceleration_structure_build_sizes(
171                        &hal::GetAccelerationStructureBuildSizesDescriptor {
172                            entries: &hal::AccelerationStructureEntries::AABBs(entries),
173                            flags: blas_desc.flags,
174                        },
175                    )
176                }
177            }
178        };
179
180        let raw = unsafe {
181            self.raw()
182                .create_acceleration_structure(&hal::AccelerationStructureDescriptor {
183                    label: hal_label(blas_desc.label.as_deref(), self.instance_flags),
184                    size: size_info.acceleration_structure_size,
185                    format: hal::AccelerationStructureFormat::BottomLevel,
186                    allow_compaction: blas_desc
187                        .flags
188                        .contains(wgpu_types::AccelerationStructureFlags::ALLOW_COMPACTION),
189                })
190        }
191        .map_err(|e| self.handle_hal_error_with_nonfatal_oom(e))?;
192
193        let compaction_buffer = if blas_desc
194            .flags
195            .contains(wgpu_types::AccelerationStructureFlags::ALLOW_COMPACTION)
196        {
197            Some(ManuallyDrop::new(unsafe {
198                self.raw()
199                    .create_buffer(&hal::BufferDescriptor {
200                        label: hal_label(
201                            Some("(wgpu internal) compaction read-back buffer"),
202                            self.instance_flags,
203                        ),
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: hal_label(desc.label.as_deref(), self.instance_flags),
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}