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}