Skip to main content

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                    .0
211            }))
212        } else {
213            None
214        };
215
216        let handle = unsafe {
217            self.raw()
218                .get_acceleration_structure_device_address(raw.as_ref())
219        };
220
221        Ok(Arc::new(resource::Blas {
222            state: ResourceState::Valid(resource::BlasState {
223                raw: Snatchable::new(raw),
224            }),
225            device: self.clone(),
226            size_info,
227            sizes,
228            flags: blas_desc.flags,
229            update_mode: blas_desc.update_mode,
230            handle,
231            label: blas_desc.label.to_string(),
232            built_index: RwLock::new(rank::BLAS_BUILT_INDEX, None),
233            tracking_data: TrackingData::new(self.tracker_indices.blas_s.clone()),
234            compaction_buffer,
235            compacted_state: Mutex::new(rank::BLAS_COMPACTION_STATE, BlasCompactState::Idle),
236        }))
237    }
238
239    pub fn create_tlas(
240        self: &Arc<Self>,
241        desc: &resource::TlasDescriptor,
242    ) -> (Arc<resource::Tlas>, Option<CreateTlasError>) {
243        let (tlas, error) = match self.create_tlas_inner(desc) {
244            Ok(tlas) => (tlas, None),
245            Err(e) => (resource::Tlas::invalid(Arc::clone(self), desc), Some(e)),
246        };
247        #[cfg(feature = "trace")]
248        if let Some(trace) = self.trace.lock().as_mut() {
249            trace.add(Action::CreateTlas {
250                id: tlas.to_trace(),
251                desc: desc.clone(),
252            });
253        }
254
255        api_log!("Device::create_tlas -> {:?}", Arc::as_ptr(&tlas));
256
257        (tlas, error)
258    }
259
260    pub(crate) fn create_tlas_inner(
261        self: &Arc<Self>,
262        desc: &resource::TlasDescriptor,
263    ) -> Result<Arc<resource::Tlas>, CreateTlasError> {
264        self.check_is_valid()?;
265        self.require_features(Features::EXPERIMENTAL_RAY_QUERY)?;
266
267        if desc.max_instances > self.limits.max_tlas_instance_count {
268            return Err(CreateTlasError::TooManyInstances(
269                self.limits.max_tlas_instance_count,
270                desc.max_instances,
271            ));
272        }
273
274        if desc
275            .flags
276            .contains(wgt::AccelerationStructureFlags::USE_TRANSFORM)
277        {
278            return Err(CreateTlasError::DisallowedFlag(
279                wgt::AccelerationStructureFlags::USE_TRANSFORM,
280            ));
281        }
282
283        if desc
284            .flags
285            .contains(wgt::AccelerationStructureFlags::ALLOW_RAY_HIT_VERTEX_RETURN)
286        {
287            self.require_features(Features::EXPERIMENTAL_RAY_HIT_VERTEX_RETURN)?;
288        }
289
290        let size_info = unsafe {
291            self.raw().get_acceleration_structure_build_sizes(
292                &hal::GetAccelerationStructureBuildSizesDescriptor {
293                    entries: &hal::AccelerationStructureEntries::Instances(
294                        hal::AccelerationStructureInstances {
295                            buffer: Some(self.zero_buffer.as_ref()),
296                            offset: 0,
297                            count: desc.max_instances,
298                        },
299                    ),
300                    flags: desc.flags,
301                },
302            )
303        };
304
305        let raw = unsafe {
306            self.raw()
307                .create_acceleration_structure(&hal::AccelerationStructureDescriptor {
308                    label: hal_label(desc.label.as_deref(), self.instance_flags),
309                    size: size_info.acceleration_structure_size,
310                    format: hal::AccelerationStructureFormat::TopLevel,
311                    allow_compaction: false,
312                })
313        }
314        .map_err(|e| self.handle_hal_error_with_nonfatal_oom(e))?;
315
316        let instance_buffer_size = self
317            .alignments
318            .raw_tlas_instance_size
319            .checked_mul(desc.max_instances.max(1))
320            .expect("max_tlas_instance_count should not allow excessive buffer size");
321        let (instance_buffer, _) = unsafe {
322            self.raw().create_buffer(&hal::BufferDescriptor {
323                label: hal_label(Some("(wgpu-core) instances_buffer"), self.instance_flags),
324                size: u64::from(instance_buffer_size),
325                usage: wgt::BufferUses::COPY_DST
326                    | wgt::BufferUses::TOP_LEVEL_ACCELERATION_STRUCTURE_INPUT,
327                memory_flags: hal::MemoryFlags::PREFER_COHERENT,
328            })
329        }
330        .map_err(|e| self.handle_hal_error_with_nonfatal_oom(e))?;
331
332        Ok(Arc::new(resource::Tlas {
333            state: ResourceState::Valid(resource::TlasState {
334                raw: Snatchable::new(raw),
335                instance_buffer,
336            }),
337            device: self.clone(),
338            size_info,
339            flags: desc.flags,
340            update_mode: desc.update_mode,
341            built_index: RwLock::new(rank::TLAS_BUILT_INDEX, None),
342            dependencies: RwLock::new(rank::TLAS_DEPENDENCIES, Vec::new()),
343            label: desc.label.to_string(),
344            max_instance_count: desc.max_instances,
345            tracking_data: TrackingData::new(self.tracker_indices.tlas_s.clone()),
346        }))
347    }
348}