1use alloc::{boxed::Box, sync::Arc, vec::Vec};
2use core::{
3 cmp::max,
4 num::NonZeroU64,
5 ops::{Deref, Range},
6};
7
8use wgt::{math::align_to, BufferUsages, BufferUses, Features};
9
10use crate::{
11 command::encoder::EncodingState,
12 ray_tracing::{
13 AsAction, AsBuild, BlasAabbGeometry, BlasTriangleGeometry, TlasBuild, TlasInstance,
14 ValidateAsActionsError,
15 },
16 resource::{Buffer, InvalidResourceError},
17};
18use crate::{command::EncoderStateError, device::resource::CommandIndices};
19use crate::{
20 command::{ArcCommand, ArcReferences, CommandBufferMutable},
21 device::queue::TempResource,
22 global::Global,
23 id::CommandEncoderId,
24 init_tracker::MemoryInitKind,
25 ray_tracing::{
26 ArcBlasAabbGeometry, ArcBlasBuildEntry, ArcBlasGeometries, ArcBlasTriangleGeometry,
27 ArcTlasInstance, ArcTlasPackage, BlasBuildEntry, BlasGeometries,
28 BuildAccelerationStructureError, OwnedBlasBuildEntry, OwnedTlasPackage, TlasPackage,
29 },
30 resource::{Blas, BlasCompactState, Labeled, StagingBuffer, Tlas, Trackable},
31 scratch::ScratchBuffer,
32 snatch::SnatchGuard,
33 track::TrackerIndex,
34 FastHashSet,
35};
36use crate::{lock::RwLockWriteGuard, resource::RawResourceAccess};
37
38use crate::id::{BlasId, TlasId};
39
40struct BlasStore<'a> {
41 blas: Arc<Blas>,
42 entries: hal::AccelerationStructureEntries<'a, dyn hal::DynBuffer>,
43 scratch_buffer_offset: u64,
44}
45
46struct UnsafeTlasStore<'a> {
47 tlas: Arc<Tlas>,
48 entries: hal::AccelerationStructureEntries<'a, dyn hal::DynBuffer>,
49 scratch_buffer_offset: u64,
50}
51
52struct TlasStore<'a> {
53 internal: UnsafeTlasStore<'a>,
54 range: Range<usize>,
55}
56
57impl super::CommandEncoder {
58 pub fn mark_acceleration_structures_built(
59 self: &Arc<Self>,
60 blases: &[Arc<Blas>],
61 tlases: &[Arc<Tlas>],
62 ) -> Result<(), EncoderStateError> {
63 profiling::scope!("CommandEncoder::mark_acceleration_structures_built");
64
65 let mut cmd_buf_data = self.data.lock();
66 cmd_buf_data.with_buffer(
67 crate::command::EncodingApi::Raw,
68 |cmd_buf_data| -> Result<(), BuildAccelerationStructureError> {
69 let device = &self.device;
70 device.check_is_valid()?;
71 device.require_features(Features::EXPERIMENTAL_RAY_QUERY)?;
72
73 let mut build_command = AsBuild::with_capacity(blases.len(), tlases.len());
74
75 for blas in blases {
76 let blas = blas.clone();
77 blas.check_is_valid()?;
78 build_command.blas_s_built.push(blas);
79 }
80
81 for tlas in tlases {
82 let tlas = tlas.clone();
83 tlas.check_is_valid()?;
84 build_command.tlas_s_built.push(TlasBuild {
85 tlas,
86 dependencies: Vec::new(),
87 });
88 }
89
90 cmd_buf_data.as_actions.push(AsAction::Build(build_command));
91 Ok(())
92 },
93 )
94 }
95
96 pub fn build_acceleration_structures<'a>(
97 self: &Arc<Self>,
98 blas_iter: impl Iterator<Item = BlasBuildEntry<'a, Arc<Blas>, Arc<Buffer>>>,
99 tlas_iter: impl Iterator<Item = TlasPackage<'a, Arc<Tlas>, Arc<Blas>>>,
100 ) -> Result<(), EncoderStateError> {
101 profiling::scope!("CommandEncoder::build_acceleration_structures");
102
103 let mut cmd_buf_data = self.data.lock();
104
105 cmd_buf_data.push_with(|| -> Result<_, BuildAccelerationStructureError> {
106 let blas = blas_iter
107 .map(|blas_entry| {
108 let geometries = match blas_entry.geometries {
109 BlasGeometries::TriangleGeometries(triangle_geometries) => {
110 let tri_geo = triangle_geometries
111 .map(|tg| {
112 let vertex_buffer = tg.vertex_buffer;
113 vertex_buffer.check_is_valid()?;
114 Ok(ArcBlasTriangleGeometry {
115 size: tg.size.clone(),
116 vertex_buffer,
117 index_buffer: tg
118 .index_buffer
119 .map(|index_buffer| -> Result<_, InvalidResourceError> {
120 index_buffer.check_is_valid()?;
121 Ok(index_buffer)
122 })
123 .transpose()?,
124 transform_buffer: tg
125 .transform_buffer
126 .map(|transform_buffer| -> Result<_, InvalidResourceError> {
127 transform_buffer.check_is_valid()?;
128 Ok(transform_buffer)
129 })
130 .transpose()?,
131 first_vertex: tg.first_vertex,
132 vertex_stride: tg.vertex_stride,
133 first_index: tg.first_index,
134 transform_buffer_offset: tg.transform_buffer_offset,
135 })
136 })
137 .collect::<Result<_, BuildAccelerationStructureError>>()?;
138 ArcBlasGeometries::TriangleGeometries(tri_geo)
139 }
140 BlasGeometries::AabbGeometries(aabb_geometries) => {
141 let aabb_geo = aabb_geometries
142 .map(|ag| {
143 ag.aabb_buffer.check_is_valid()?;
144 Ok(ArcBlasAabbGeometry {
145 size: ag.size.clone(),
146 stride: ag.stride,
147 aabb_buffer: ag.aabb_buffer,
148 primitive_offset: ag.primitive_offset,
149 })
150 })
151 .collect::<Result<_, BuildAccelerationStructureError>>()?;
152 ArcBlasGeometries::AabbGeometries(aabb_geo)
153 }
154 };
155 let blas = blas_entry.blas;
156 blas.check_is_valid()?;
157 Ok(ArcBlasBuildEntry { blas, geometries })
158 })
159 .collect::<Result<_, BuildAccelerationStructureError>>()?;
160
161 let tlas = tlas_iter
162 .map(|tlas_package| {
163 let instances = tlas_package
164 .instances
165 .map(|instance| {
166 instance
167 .as_ref()
168 .map(|instance| {
169 let blas = instance.blas.clone();
170 blas.check_is_valid()?;
171 Ok(ArcTlasInstance {
172 blas,
173 transform: *instance.transform,
174 custom_data: instance.custom_data,
175 mask: instance.mask,
176 })
177 })
178 .transpose()
179 })
180 .collect::<Result<_, BuildAccelerationStructureError>>()?;
181 let tlas = tlas_package.tlas;
182 tlas.check_is_valid()?;
183 Ok(ArcTlasPackage {
184 tlas,
185 instances,
186 lowest_unmodified: tlas_package.lowest_unmodified,
187 })
188 })
189 .collect::<Result<_, BuildAccelerationStructureError>>()?;
190
191 Ok(ArcCommand::BuildAccelerationStructures { blas, tlas })
192 })
193 }
194}
195
196impl Global {
197 pub fn command_encoder_mark_acceleration_structures_built(
198 &self,
199 command_encoder_id: CommandEncoderId,
200 blas_ids: &[BlasId],
201 tlas_ids: &[TlasId],
202 ) -> Result<(), EncoderStateError> {
203 let hub = &self.hub;
204
205 let cmd_enc = hub.command_encoders.get(command_encoder_id);
206
207 let blases = blas_ids
208 .iter()
209 .map(|&id| hub.blas_s.get(id))
210 .collect::<Vec<_>>();
211 let tlases = tlas_ids
212 .iter()
213 .map(|&id| hub.tlas_s.get(id))
214 .collect::<Vec<_>>();
215 cmd_enc.mark_acceleration_structures_built(&blases, &tlases)
216 }
217
218 pub fn command_encoder_build_acceleration_structures<'a>(
219 &self,
220 command_encoder_id: CommandEncoderId,
221 blas_iter: impl Iterator<Item = BlasBuildEntry<'a>>,
222 tlas_iter: impl Iterator<Item = TlasPackage<'a>>,
223 ) -> Result<(), EncoderStateError> {
224 let hub = &self.hub;
225
226 let cmd_enc = hub.command_encoders.get(command_encoder_id);
227 let blases = blas_iter.map(|e| BlasBuildEntry {
228 blas: hub.blas_s.get(e.blas),
229 geometries: match e.geometries {
230 BlasGeometries::TriangleGeometries(triangle_geometries) => {
231 let triangle_geometries = triangle_geometries.map(|tg| BlasTriangleGeometry {
232 size: tg.size,
233 vertex_buffer: hub.buffers.get(tg.vertex_buffer),
234 index_buffer: tg
235 .index_buffer
236 .map(|index_buffer| hub.buffers.get(index_buffer)),
237 transform_buffer: tg
238 .transform_buffer
239 .map(|transform_buffer| hub.buffers.get(transform_buffer)),
240 first_vertex: tg.first_vertex,
241 vertex_stride: tg.vertex_stride,
242 first_index: tg.first_index,
243 transform_buffer_offset: tg.transform_buffer_offset,
244 });
245 BlasGeometries::TriangleGeometries(Box::new(triangle_geometries.into_iter()))
246 }
247 BlasGeometries::AabbGeometries(aabb_geometries) => {
248 let aabb_geometries = aabb_geometries.map(|ag| BlasAabbGeometry {
249 size: ag.size,
250 stride: ag.stride,
251 aabb_buffer: hub.buffers.get(ag.aabb_buffer),
252 primitive_offset: ag.primitive_offset,
253 });
254 BlasGeometries::AabbGeometries(Box::new(aabb_geometries.into_iter()))
255 }
256 },
257 });
258 let tlases = tlas_iter.map(|e| TlasPackage {
259 tlas: hub.tlas_s.get(e.tlas),
260 instances: Box::new(e.instances.map(|instance| {
261 instance.as_ref().map(|instance| TlasInstance {
262 blas: hub.blas_s.get(instance.blas),
263 transform: instance.transform,
264 custom_data: instance.custom_data,
265 mask: instance.mask,
266 })
267 })),
268 lowest_unmodified: e.lowest_unmodified,
269 });
270 cmd_enc.build_acceleration_structures(blases.into_iter(), tlases.into_iter())
271 }
272}
273
274pub(crate) fn build_acceleration_structures(
275 state: &mut EncodingState,
276 blas: Vec<OwnedBlasBuildEntry<ArcReferences>>,
277 tlas: Vec<OwnedTlasPackage<ArcReferences>>,
278) -> Result<(), BuildAccelerationStructureError> {
279 state
280 .device
281 .require_features(Features::EXPERIMENTAL_RAY_QUERY)?;
282
283 let mut build_command = AsBuild::with_capacity(blas.len(), tlas.len());
284 let mut input_barriers = Vec::<hal::BufferBarrier<dyn hal::DynBuffer>>::new();
285 let mut scratch_buffer_blas_size = 0;
286 let mut blas_storage = Vec::with_capacity(blas.len());
287 iter_blas(
288 blas.iter(),
289 &mut build_command,
290 &mut input_barriers,
291 &mut scratch_buffer_blas_size,
292 &mut blas_storage,
293 state,
294 )?;
295
296 let mut scratch_buffer_tlas_size: u64 = 0;
297 let mut tlas_storage = Vec::<TlasStore>::with_capacity(tlas.len());
298 let mut instance_buffer_staging_source = Vec::<u8>::new();
299
300 for package in tlas.iter() {
301 let tlas = &package.tlas;
302 state.tracker.tlas_s.insert_single(tlas.clone());
303
304 let scratch_buffer_offset = scratch_buffer_tlas_size;
305 scratch_buffer_tlas_size = scratch_buffer_tlas_size.saturating_add(align_to(
306 tlas.size_info.build_scratch_size,
307 u64::from(state.device.alignments.ray_tracing_scratch_buffer_alignment),
308 ));
309
310 let first_byte_index = instance_buffer_staging_source.len();
311
312 let mut dependencies = Vec::new();
313 let mut seen_dependencies = FastHashSet::<TrackerIndex>::default();
314
315 let mut instance_count = 0;
316 for instance in package.instances.iter().flatten() {
317 if instance.custom_data >= (1u32 << 24u32) {
318 return Err(BuildAccelerationStructureError::TlasInvalidCustomIndex(
319 tlas.error_ident(),
320 ));
321 }
322 let blas = &instance.blas;
323 let is_new_dependency = seen_dependencies.insert(blas.tracker_index());
324
325 if is_new_dependency {
326 state.tracker.blas_s.insert_single(blas.clone());
327 }
328
329 instance_buffer_staging_source.extend(state.device.raw().tlas_instance_to_bytes(
330 hal::TlasInstance {
331 transform: instance.transform,
332 custom_data: instance.custom_data,
333 mask: instance.mask,
334 blas_address: blas.handle,
335 pipeline_intersection_data_offset: 0,
336 },
337 ));
338
339 if tlas
340 .flags
341 .contains(wgpu_types::AccelerationStructureFlags::ALLOW_RAY_HIT_VERTEX_RETURN)
342 && !blas
343 .flags
344 .contains(wgpu_types::AccelerationStructureFlags::ALLOW_RAY_HIT_VERTEX_RETURN)
345 {
346 return Err(
347 BuildAccelerationStructureError::TlasDependentMissingVertexReturn(
348 tlas.error_ident(),
349 blas.error_ident(),
350 ),
351 );
352 }
353
354 instance_count += 1;
355
356 if is_new_dependency {
357 dependencies.push(blas.clone());
358 }
359 }
360
361 build_command.tlas_s_built.push(TlasBuild {
362 tlas: tlas.clone(),
363 dependencies,
364 });
365
366 if instance_count > tlas.max_instance_count {
367 return Err(BuildAccelerationStructureError::TlasInstanceCountExceeded(
368 tlas.error_ident(),
369 instance_count,
370 tlas.max_instance_count,
371 ));
372 }
373
374 tlas_storage.push(TlasStore {
375 internal: UnsafeTlasStore {
376 tlas: tlas.clone(),
377 entries: hal::AccelerationStructureEntries::Instances(
378 hal::AccelerationStructureInstances {
379 buffer: Some(tlas.state()?.instance_buffer.as_ref()),
380 offset: 0,
381 count: instance_count,
382 },
383 ),
384 scratch_buffer_offset,
385 },
386 range: first_byte_index..instance_buffer_staging_source.len(),
387 });
388 }
389
390 let Some(scratch_size) =
391 wgt::BufferSize::new(max(scratch_buffer_blas_size, scratch_buffer_tlas_size))
392 else {
393 return Ok(());
395 };
396
397 if scratch_size.get() == u64::MAX {
398 return Err(crate::device::DeviceError::OutOfMemory.into());
399 }
400
401 let scratch_buffer = ScratchBuffer::new(state.device, scratch_size)?;
402
403 let scratch_buffer_barrier = hal::BufferBarrier::<dyn hal::DynBuffer> {
404 buffer: scratch_buffer.raw(),
405 usage: hal::StateTransition {
406 from: BufferUses::ACCELERATION_STRUCTURE_SCRATCH,
407 to: BufferUses::ACCELERATION_STRUCTURE_SCRATCH,
408 },
409 };
410
411 let mut tlas_descriptors = Vec::with_capacity(tlas_storage.len());
412
413 for &TlasStore {
414 internal:
415 UnsafeTlasStore {
416 ref tlas,
417 ref entries,
418 ref scratch_buffer_offset,
419 },
420 ..
421 } in &tlas_storage
422 {
423 if tlas.update_mode == wgt::AccelerationStructureUpdateMode::PreferUpdate {
424 log::warn!("build_acceleration_structures called with PreferUpdate, but only rebuild is implemented");
425 }
426 tlas_descriptors.push(hal::BuildAccelerationStructureDescriptor {
427 entries,
428 mode: hal::AccelerationStructureBuildMode::Build,
429 flags: tlas.flags,
430 source_acceleration_structure: None,
431 destination_acceleration_structure: tlas.try_raw(state.snatch_guard)?,
432 scratch_buffer: scratch_buffer.raw(),
433 scratch_buffer_offset: *scratch_buffer_offset,
434 })
435 }
436
437 let blas_present = !blas_storage.is_empty();
438 let tlas_present = !tlas_storage.is_empty();
439
440 let raw_encoder = &mut state.raw_encoder;
441
442 let mut blas_s_compactable = Vec::new();
443 let mut descriptors = Vec::with_capacity(blas.len());
444
445 for storage in &blas_storage {
446 descriptors.push(map_blas(
447 storage,
448 scratch_buffer.raw(),
449 state.snatch_guard,
450 &mut blas_s_compactable,
451 )?);
452 }
453
454 build_blas(
455 *raw_encoder,
456 blas_present,
457 tlas_present,
458 input_barriers,
459 &descriptors,
460 scratch_buffer_barrier,
461 blas_s_compactable,
462 );
463
464 if tlas_present {
465 let staging_buffer = if !instance_buffer_staging_source.is_empty() {
466 let mut staging_buffer = StagingBuffer::new(
467 state.device,
468 wgt::BufferSize::new(instance_buffer_staging_source.len() as u64).unwrap(),
469 )?;
470 staging_buffer.write(&instance_buffer_staging_source);
471 let flushed = staging_buffer.flush();
472 Some(flushed)
473 } else {
474 None
475 };
476
477 unsafe {
478 if let Some(ref staging_buffer) = staging_buffer {
479 raw_encoder.transition_buffers(&[hal::BufferBarrier::<dyn hal::DynBuffer> {
480 buffer: staging_buffer.raw(),
481 usage: hal::StateTransition {
482 from: BufferUses::MAP_WRITE,
483 to: BufferUses::COPY_SRC,
484 },
485 }]);
486 }
487 }
488
489 let mut instance_buffer_barriers = Vec::new();
490 for &TlasStore {
491 internal: UnsafeTlasStore { ref tlas, .. },
492 ref range,
493 } in &tlas_storage
494 {
495 let size = match wgt::BufferSize::new((range.end - range.start) as u64) {
496 None => continue,
497 Some(size) => size,
498 };
499 let tlas_state = tlas.state()?;
500 instance_buffer_barriers.push(hal::BufferBarrier::<dyn hal::DynBuffer> {
501 buffer: tlas_state.instance_buffer.as_ref(),
502 usage: hal::StateTransition {
503 from: BufferUses::COPY_DST,
504 to: BufferUses::TOP_LEVEL_ACCELERATION_STRUCTURE_INPUT,
505 },
506 });
507 unsafe {
508 raw_encoder.transition_buffers(&[hal::BufferBarrier::<dyn hal::DynBuffer> {
509 buffer: tlas_state.instance_buffer.as_ref(),
510 usage: hal::StateTransition {
511 from: BufferUses::TOP_LEVEL_ACCELERATION_STRUCTURE_INPUT,
512 to: BufferUses::COPY_DST,
513 },
514 }]);
515 let temp = hal::BufferCopy {
516 src_offset: range.start as u64,
517 dst_offset: 0,
518 size,
519 };
520 raw_encoder.copy_buffer_to_buffer(
521 staging_buffer.as_ref().unwrap().raw(),
522 tlas_state.instance_buffer.as_ref(),
523 &[temp],
524 );
525 }
526 }
527
528 unsafe {
529 raw_encoder.transition_buffers(&instance_buffer_barriers);
530
531 raw_encoder.build_acceleration_structures(&tlas_descriptors);
532
533 raw_encoder.place_acceleration_structure_barrier(hal::AccelerationStructureBarrier {
534 usage: hal::StateTransition {
535 from: hal::AccelerationStructureUses::BUILD_OUTPUT,
536 to: hal::AccelerationStructureUses::SHADER_INPUT,
537 },
538 });
539 }
540
541 if let Some(staging_buffer) = staging_buffer {
542 state
543 .temp_resources
544 .push(TempResource::StagingBuffer(staging_buffer));
545 }
546 }
547
548 state
549 .temp_resources
550 .push(TempResource::ScratchBuffer(scratch_buffer));
551
552 state.as_actions.push(AsAction::Build(build_command));
553
554 Ok(())
555}
556
557impl CommandBufferMutable {
558 pub(crate) fn validate_acceleration_structure_actions(
559 &self,
560 snatch_guard: &SnatchGuard,
561 command_index_guard: &mut RwLockWriteGuard<CommandIndices>,
562 ) -> Result<(), ValidateAsActionsError> {
563 profiling::scope!("CommandEncoder::[submission]::validate_as_actions");
564 for action in &self.as_actions {
565 match action {
566 AsAction::Build(build) => {
567 let build_command_index = NonZeroU64::new(
568 command_index_guard.next_acceleration_structure_build_command_index,
569 )
570 .unwrap();
571
572 command_index_guard.next_acceleration_structure_build_command_index += 1;
573 for blas in build.blas_s_built.iter() {
574 let mut state_lock = blas.compacted_state.lock();
575 *state_lock = match *state_lock {
576 BlasCompactState::Compacted => {
577 unreachable!("Should be validated out in build.")
578 }
579 _ => BlasCompactState::Idle,
582 };
583 *blas.built_index.write() = Some(build_command_index);
584 }
585
586 for tlas_build in build.tlas_s_built.iter() {
587 for blas in &tlas_build.dependencies {
588 if blas.built_index.read().is_none() {
589 return Err(ValidateAsActionsError::UsedUnbuiltBlas(
590 blas.error_ident(),
591 tlas_build.tlas.error_ident(),
592 ));
593 }
594 }
595 *tlas_build.tlas.built_index.write() = Some(build_command_index);
596 tlas_build
597 .tlas
598 .dependencies
599 .write()
600 .clone_from(&tlas_build.dependencies)
601 }
602 }
603 AsAction::UseTlas(tlas) => {
604 let tlas_build_index = tlas.built_index.read();
605 let dependencies = tlas.dependencies.read();
606
607 if (*tlas_build_index).is_none() {
608 return Err(ValidateAsActionsError::UsedUnbuiltTlas(tlas.error_ident()));
609 }
610 for blas in dependencies.deref() {
611 let blas_build_index = *blas.built_index.read();
612 if blas_build_index.is_none() {
613 return Err(ValidateAsActionsError::UsedUnbuiltBlas(
614 tlas.error_ident(),
615 blas.error_ident(),
616 ));
617 }
618 if blas_build_index.unwrap() > tlas_build_index.unwrap() {
619 return Err(ValidateAsActionsError::BlasNewerThenTlas(
620 blas.error_ident(),
621 tlas.error_ident(),
622 ));
623 }
624 blas.try_raw(snatch_guard)?;
625 }
626 }
627 }
628 }
629 Ok(())
630 }
631
632 pub(crate) fn set_acceleration_structure_dependencies(&self, snatch_guard: &SnatchGuard) {
633 profiling::scope!("CommandEncoder::[submission]::set_acceleration_structure_dependencies");
634 let tlas_dependencies_locks: Vec<_> = self
635 .as_actions
636 .iter()
637 .filter_map(|action| {
638 if let AsAction::UseTlas(tlas) = action {
639 Some(tlas.dependencies.read())
640 } else {
641 None
642 }
643 })
644 .collect();
645 let mut tlas_dependencies_lock_iter = tlas_dependencies_locks.iter();
646 let mut dependencies = Vec::new();
647 for action in &self.as_actions {
648 match action {
649 AsAction::Build(build) => {
650 for tlas_build in build.tlas_s_built.iter() {
651 for dependency in &tlas_build.dependencies {
652 if let Some(dependency) = dependency.raw(snatch_guard) {
653 dependencies.push(dependency);
654 }
655 }
656 }
657 }
658 AsAction::UseTlas(_tlas) => {
659 let tlas_dependencies = tlas_dependencies_lock_iter.next().unwrap(); for dependency in tlas_dependencies.iter() {
661 if let Some(dependency) = dependency.raw(snatch_guard) {
662 dependencies.push(dependency);
663 }
664 }
665 }
666 }
667 }
668 if !dependencies.is_empty() {
669 unsafe {
670 self.encoder
671 .raw
672 .set_acceleration_structure_dependencies(&self.encoder.list, &dependencies);
673 }
674 }
675 }
676}
677
678fn iter_blas<'snatch_guard: 'buffers, 'buffers>(
680 blas_iter: impl Iterator<Item = &'buffers OwnedBlasBuildEntry<ArcReferences>>,
681 build_command: &mut AsBuild,
682 input_barriers: &mut Vec<hal::BufferBarrier<'buffers, dyn hal::DynBuffer>>,
683 scratch_buffer_blas_size: &mut u64,
684 blas_storage: &mut Vec<BlasStore<'buffers>>,
685 state: &mut EncodingState<'snatch_guard, '_>,
686) -> Result<(), BuildAccelerationStructureError> {
687 for entry in blas_iter {
688 let blas = &entry.blas;
689 state.tracker.blas_s.insert_single(blas.clone());
690
691 build_command.blas_s_built.push(blas.clone());
692
693 match &entry.geometries {
694 ArcBlasGeometries::TriangleGeometries(triangle_geometries) => {
695 let mut triangle_entries =
696 Vec::<hal::AccelerationStructureTriangles<dyn hal::DynBuffer>>::new();
697
698 for (i, mesh) in triangle_geometries.iter().enumerate() {
699 let size_desc = match &blas.sizes {
700 wgt::BlasGeometrySizeDescriptors::Triangles { descriptors } => descriptors,
701 _ => {
702 return Err(BuildAccelerationStructureError::BlasGeometryKindMismatch(
703 blas.error_ident(),
704 ));
705 }
706 };
707 if i >= size_desc.len() {
708 return Err(BuildAccelerationStructureError::IncompatibleBlasBuildSizes(
709 blas.error_ident(),
710 ));
711 }
712 let size_desc = &size_desc[i];
713
714 if size_desc.flags != mesh.size.flags {
715 return Err(BuildAccelerationStructureError::IncompatibleBlasFlags(
716 blas.error_ident(),
717 size_desc.flags,
718 mesh.size.flags,
719 ));
720 }
721
722 if size_desc.vertex_count < mesh.size.vertex_count {
723 return Err(
724 BuildAccelerationStructureError::IncompatibleBlasVertexCount(
725 blas.error_ident(),
726 size_desc.vertex_count,
727 mesh.size.vertex_count,
728 ),
729 );
730 }
731
732 if size_desc.vertex_format != mesh.size.vertex_format {
733 return Err(BuildAccelerationStructureError::DifferentBlasVertexFormats(
734 blas.error_ident(),
735 size_desc.vertex_format,
736 mesh.size.vertex_format,
737 ));
738 }
739
740 if size_desc
741 .vertex_format
742 .min_acceleration_structure_vertex_stride()
743 > mesh.vertex_stride
744 {
745 return Err(BuildAccelerationStructureError::VertexStrideTooSmall(
746 blas.error_ident(),
747 size_desc
748 .vertex_format
749 .min_acceleration_structure_vertex_stride(),
750 mesh.vertex_stride,
751 ));
752 }
753
754 if mesh.vertex_stride
755 % size_desc
756 .vertex_format
757 .acceleration_structure_stride_alignment()
758 != 0
759 {
760 return Err(BuildAccelerationStructureError::VertexStrideUnaligned(
761 blas.error_ident(),
762 size_desc
763 .vertex_format
764 .acceleration_structure_stride_alignment(),
765 mesh.vertex_stride,
766 ));
767 }
768
769 match (size_desc.index_count, mesh.size.index_count) {
770 (Some(_), None) | (None, Some(_)) => {
771 return Err(
772 BuildAccelerationStructureError::BlasIndexCountProvidedMismatch(
773 blas.error_ident(),
774 ),
775 )
776 }
777 (Some(create), Some(build)) if create < build => {
778 return Err(
779 BuildAccelerationStructureError::IncompatibleBlasIndexCount(
780 blas.error_ident(),
781 create,
782 build,
783 ),
784 )
785 }
786 _ => {}
787 }
788
789 if size_desc.index_format != mesh.size.index_format {
790 return Err(BuildAccelerationStructureError::DifferentBlasIndexFormats(
791 blas.error_ident(),
792 size_desc.index_format,
793 mesh.size.index_format,
794 ));
795 }
796
797 if size_desc.index_count.is_some() && mesh.index_buffer.is_none() {
798 return Err(BuildAccelerationStructureError::MissingIndexBuffer(
799 blas.error_ident(),
800 ));
801 }
802 let vertex_buffer = mesh.vertex_buffer.clone();
803 let vertex_pending = state.tracker.buffers.set_single(
804 &vertex_buffer,
805 BufferUses::BOTTOM_LEVEL_ACCELERATION_STRUCTURE_INPUT,
806 );
807 let vertex_buffer = {
808 let vertex_raw = mesh.vertex_buffer.as_ref().try_raw(state.snatch_guard)?;
809 let vertex_buffer = &mesh.vertex_buffer;
810 vertex_buffer.check_usage(BufferUsages::BLAS_INPUT)?;
811
812 if let Some(barrier) = vertex_pending.map(|pending| {
813 pending.into_hal(vertex_buffer.as_ref(), state.snatch_guard)
814 }) {
815 input_barriers.push(barrier);
816 }
817 if u64::from(mesh.size.vertex_count)
818 .checked_add(u64::from(mesh.first_vertex))
819 .and_then(|end_vertex| end_vertex.checked_mul(mesh.vertex_stride))
820 .is_none_or(|end| vertex_buffer.size < end)
821 {
822 return Err(BuildAccelerationStructureError::InsufficientBufferSize {
823 buffer_ident: vertex_buffer.error_ident(),
824 offset: u64::from(mesh.first_vertex)
825 .saturating_mul(mesh.vertex_stride),
826 region_size: u64::from(mesh.size.vertex_count)
827 .saturating_mul(mesh.vertex_stride),
828 buffer_size: vertex_buffer.size,
829 });
830 }
831 let vertex_buffer_offset = mesh.first_vertex as u64 * mesh.vertex_stride;
832 state.buffer_memory_init_actions.extend(
833 vertex_buffer.initialization_status.read().create_action(
834 vertex_buffer,
835 vertex_buffer_offset
836 ..(vertex_buffer_offset
837 + mesh.size.vertex_count as u64 * mesh.vertex_stride),
838 MemoryInitKind::NeedsInitializedMemory,
839 ),
840 );
841 vertex_raw
842 };
843 let index_buffer = if let Some(ref index_buffer) = mesh.index_buffer {
844 if mesh.first_index.is_none()
845 || mesh.size.index_count.is_none()
846 || mesh.size.index_count.is_none()
847 {
848 return Err(BuildAccelerationStructureError::MissingAssociatedData(
849 index_buffer.error_ident(),
850 ));
851 }
852 let index_pending = state.tracker.buffers.set_single(
853 index_buffer,
854 BufferUses::BOTTOM_LEVEL_ACCELERATION_STRUCTURE_INPUT,
855 );
856 let index_raw = index_buffer.try_raw(state.snatch_guard)?;
857 index_buffer.check_usage(BufferUsages::BLAS_INPUT)?;
858
859 if let Some(barrier) = index_pending.map(|pending| {
860 pending.into_hal(index_buffer.as_ref(), state.snatch_guard)
861 }) {
862 input_barriers.push(barrier);
863 }
864 let index_stride = mesh.size.index_format.unwrap().byte_size();
865 let Some(vertex_offset) =
867 mesh.first_index.unwrap().checked_mul(index_stride)
868 else {
869 return Err(BuildAccelerationStructureError::OffsetLimitedTo4GB {
870 buffer_ident: index_buffer.error_ident(),
871 offset: u64::from(mesh.first_index.unwrap())
872 .saturating_mul(u64::from(index_stride)),
873 count: u64::from(mesh.first_index.unwrap()),
874 stride: u64::from(index_stride),
875 });
876 };
877 let Some(indexes_size) =
878 mesh.size.index_count.unwrap().checked_mul(index_stride)
879 else {
880 return Err(BuildAccelerationStructureError::OffsetLimitedTo4GB {
881 buffer_ident: index_buffer.error_ident(),
882 offset: u64::from(mesh.size.index_count.unwrap())
883 .saturating_mul(u64::from(index_stride)),
884 count: u64::from(mesh.size.index_count.unwrap()),
885 stride: u64::from(index_stride),
886 });
887 };
888
889 if mesh.size.index_count.unwrap() % 3 != 0 {
890 return Err(BuildAccelerationStructureError::InvalidIndexCount(
891 index_buffer.error_ident(),
892 mesh.size.index_count.unwrap(),
893 ));
894 }
895 if index_buffer.size < u64::from(vertex_offset)
896 || index_buffer.size - u64::from(vertex_offset)
897 < u64::from(indexes_size)
898 {
899 return Err(BuildAccelerationStructureError::InsufficientBufferSize {
900 buffer_ident: index_buffer.error_ident(),
901 offset: u64::from(vertex_offset),
902 region_size: u64::from(indexes_size),
903 buffer_size: index_buffer.size,
904 });
905 }
906
907 state.buffer_memory_init_actions.extend(
908 index_buffer.initialization_status.read().create_action(
909 index_buffer,
910 u64::from(vertex_offset)
911 ..(u64::from(vertex_offset) + u64::from(indexes_size)),
912 MemoryInitKind::NeedsInitializedMemory,
913 ),
914 );
915 Some((index_raw, vertex_offset))
916 } else {
917 None
918 };
919 let transform_buffer = if let Some(ref transform_buffer) = mesh.transform_buffer
920 {
921 if !blas
922 .flags
923 .contains(wgt::AccelerationStructureFlags::USE_TRANSFORM)
924 {
925 return Err(BuildAccelerationStructureError::UseTransformMissing(
926 blas.error_ident(),
927 ));
928 }
929 if mesh.transform_buffer_offset.is_none() {
930 return Err(BuildAccelerationStructureError::MissingAssociatedData(
931 transform_buffer.error_ident(),
932 ));
933 }
934 let transform_pending = state.tracker.buffers.set_single(
935 transform_buffer,
936 BufferUses::BOTTOM_LEVEL_ACCELERATION_STRUCTURE_INPUT,
937 );
938 if mesh.transform_buffer_offset.is_none() {
939 return Err(BuildAccelerationStructureError::MissingAssociatedData(
940 transform_buffer.error_ident(),
941 ));
942 }
943 let transform_raw = transform_buffer.try_raw(state.snatch_guard)?;
944 transform_buffer.check_usage(BufferUsages::BLAS_INPUT)?;
945
946 if let Some(barrier) = transform_pending.map(|pending| {
947 pending.into_hal(transform_buffer.as_ref(), state.snatch_guard)
948 }) {
949 input_barriers.push(barrier);
950 }
951
952 let Ok(offset) = u32::try_from(mesh.transform_buffer_offset.unwrap())
954 else {
955 return Err(BuildAccelerationStructureError::OffsetLimitedTo4GB {
956 buffer_ident: transform_buffer.error_ident(),
957 offset: mesh.transform_buffer_offset.unwrap(),
958 count: mesh.transform_buffer_offset.unwrap(),
959 stride: 1,
960 });
961 };
962
963 if offset % wgt::TRANSFORM_BUFFER_ALIGNMENT as u32 != 0 {
964 return Err(
965 BuildAccelerationStructureError::UnalignedTransformBufferOffset(
966 transform_buffer.error_ident(),
967 ),
968 );
969 }
970 if transform_buffer.size < 48
971 || transform_buffer.size - 48 < u64::from(offset)
972 {
973 return Err(BuildAccelerationStructureError::InsufficientBufferSize {
974 buffer_ident: transform_buffer.error_ident(),
975 region_size: 48,
976 offset: u64::from(offset),
977 buffer_size: transform_buffer.size,
978 });
979 }
980 state.buffer_memory_init_actions.extend(
981 transform_buffer.initialization_status.read().create_action(
982 transform_buffer,
983 u64::from(offset)..(u64::from(offset) + 48),
984 MemoryInitKind::NeedsInitializedMemory,
985 ),
986 );
987 Some((transform_raw, offset))
988 } else {
989 if blas
990 .flags
991 .contains(wgt::AccelerationStructureFlags::USE_TRANSFORM)
992 {
993 return Err(BuildAccelerationStructureError::TransformMissing(
994 blas.error_ident(),
995 ));
996 }
997 None
998 };
999
1000 let triangles = hal::AccelerationStructureTriangles {
1001 vertex_buffer: Some(vertex_buffer),
1002 vertex_format: mesh.size.vertex_format,
1003 first_vertex: mesh.first_vertex,
1004 vertex_count: mesh.size.vertex_count,
1005 vertex_stride: mesh.vertex_stride,
1006 indices: index_buffer.map(|(index_raw, vertex_offset)| {
1007 hal::AccelerationStructureTriangleIndices::<dyn hal::DynBuffer> {
1008 format: mesh.size.index_format.unwrap(),
1009 buffer: Some(index_raw),
1010 offset: vertex_offset,
1011 count: mesh.size.index_count.unwrap(),
1012 }
1013 }),
1014 transform: transform_buffer.map(|(transform_raw, offset)| {
1015 hal::AccelerationStructureTriangleTransform {
1016 buffer: transform_raw,
1017 offset,
1018 }
1019 }),
1020 flags: mesh.size.flags,
1021 };
1022 triangle_entries.push(triangles);
1023 }
1024
1025 {
1026 let scratch_buffer_offset = *scratch_buffer_blas_size;
1027 *scratch_buffer_blas_size = scratch_buffer_blas_size.saturating_add(align_to(
1028 blas.size_info.build_scratch_size,
1029 u64::from(state.device.alignments.ray_tracing_scratch_buffer_alignment),
1030 ));
1031
1032 blas_storage.push(BlasStore {
1033 blas: blas.clone(),
1034 entries: hal::AccelerationStructureEntries::Triangles(triangle_entries),
1035 scratch_buffer_offset,
1036 });
1037 }
1038 }
1039 ArcBlasGeometries::AabbGeometries(aabb_geometries) => {
1040 let mut aabb_entries =
1041 Vec::<hal::AccelerationStructureAABBs<dyn hal::DynBuffer>>::new();
1042
1043 for (i, aabb) in aabb_geometries.iter().enumerate() {
1044 let size_desc = match &blas.sizes {
1045 wgt::BlasGeometrySizeDescriptors::AABBs { descriptors } => descriptors,
1046 _ => {
1047 return Err(BuildAccelerationStructureError::BlasGeometryKindMismatch(
1048 blas.error_ident(),
1049 ));
1050 }
1051 };
1052 if i >= size_desc.len() {
1053 return Err(BuildAccelerationStructureError::IncompatibleBlasBuildSizes(
1054 blas.error_ident(),
1055 ));
1056 }
1057 let size_desc = &size_desc[i];
1058
1059 if size_desc.flags != aabb.size.flags {
1060 return Err(BuildAccelerationStructureError::IncompatibleBlasFlags(
1061 blas.error_ident(),
1062 size_desc.flags,
1063 aabb.size.flags,
1064 ));
1065 }
1066
1067 if size_desc.primitive_count < aabb.size.primitive_count {
1068 return Err(
1069 BuildAccelerationStructureError::IncompatibleBlasAabbPrimitiveCount(
1070 blas.error_ident(),
1071 size_desc.primitive_count,
1072 aabb.size.primitive_count,
1073 ),
1074 );
1075 }
1076
1077 if aabb.primitive_offset % 8 != 0 {
1078 return Err(
1079 BuildAccelerationStructureError::UnalignedAabbPrimitiveOffset(
1080 blas.error_ident(),
1081 ),
1082 );
1083 }
1084
1085 if aabb.stride < wgt::AABB_GEOMETRY_MIN_STRIDE || aabb.stride % 8 != 0 {
1086 return Err(BuildAccelerationStructureError::InvalidAabbStride(
1087 blas.error_ident(),
1088 aabb.stride,
1089 ));
1090 }
1091
1092 let aabb_buffer = aabb.aabb_buffer.clone();
1093 let aabb_pending = state.tracker.buffers.set_single(
1094 &aabb_buffer,
1095 BufferUses::BOTTOM_LEVEL_ACCELERATION_STRUCTURE_INPUT,
1096 );
1097 let aabb_raw = {
1098 let aabb_raw = aabb.aabb_buffer.as_ref().try_raw(state.snatch_guard)?;
1099 let aabb_buffer = &aabb.aabb_buffer;
1100 aabb_buffer.check_usage(BufferUsages::BLAS_INPUT)?;
1101
1102 if let Some(barrier) = aabb_pending.map(|pending| {
1103 pending.into_hal(aabb_buffer.as_ref(), state.snatch_guard)
1104 }) {
1105 input_barriers.push(barrier);
1106 }
1107
1108 let Some(aabb_size) = u32::try_from(aabb.stride)
1110 .ok()
1111 .and_then(|stride| aabb.size.primitive_count.checked_mul(stride))
1112 else {
1113 return Err(BuildAccelerationStructureError::OffsetLimitedTo4GB {
1114 buffer_ident: aabb_buffer.error_ident(),
1115 offset: u64::from(aabb.size.primitive_count)
1116 .saturating_mul(aabb.stride),
1117 count: u64::from(aabb.size.primitive_count),
1118 stride: aabb.stride,
1119 });
1120 };
1121
1122 if aabb_buffer.size < u64::from(aabb.primitive_offset)
1123 || aabb_buffer.size - u64::from(aabb.primitive_offset)
1124 < u64::from(aabb_size)
1125 {
1126 return Err(BuildAccelerationStructureError::InsufficientBufferSize {
1127 buffer_ident: aabb_buffer.error_ident(),
1128 region_size: u64::from(aabb_size),
1129 offset: u64::from(aabb.primitive_offset),
1130 buffer_size: aabb_buffer.size,
1131 });
1132 }
1133
1134 state.buffer_memory_init_actions.extend(
1135 aabb_buffer.initialization_status.read().create_action(
1136 aabb_buffer,
1137 u64::from(aabb.primitive_offset)
1138 ..u64::from(aabb.primitive_offset) + u64::from(aabb_size),
1139 MemoryInitKind::NeedsInitializedMemory,
1140 ),
1141 );
1142 aabb_raw
1143 };
1144
1145 aabb_entries.push(hal::AccelerationStructureAABBs {
1146 buffer: Some(aabb_raw),
1147 offset: aabb.primitive_offset,
1148 count: aabb.size.primitive_count,
1149 stride: aabb.stride,
1150 flags: aabb.size.flags,
1151 });
1152 }
1153
1154 {
1155 let scratch_buffer_offset = *scratch_buffer_blas_size;
1156 *scratch_buffer_blas_size = scratch_buffer_blas_size.saturating_add(align_to(
1157 blas.size_info.build_scratch_size,
1158 u64::from(state.device.alignments.ray_tracing_scratch_buffer_alignment),
1159 ));
1160
1161 blas_storage.push(BlasStore {
1162 blas: blas.clone(),
1163 entries: hal::AccelerationStructureEntries::AABBs(aabb_entries),
1164 scratch_buffer_offset,
1165 });
1166 }
1167 }
1168 }
1169 }
1170 Ok(())
1171}
1172
1173fn map_blas<'a>(
1174 storage: &'a BlasStore<'_>,
1175 scratch_buffer: &'a dyn hal::DynBuffer,
1176 snatch_guard: &'a SnatchGuard,
1177 blases_compactable: &mut Vec<(
1178 &'a dyn hal::DynBuffer,
1179 &'a dyn hal::DynAccelerationStructure,
1180 )>,
1181) -> Result<
1182 hal::BuildAccelerationStructureDescriptor<
1183 'a,
1184 dyn hal::DynBuffer,
1185 dyn hal::DynAccelerationStructure,
1186 >,
1187 BuildAccelerationStructureError,
1188> {
1189 let BlasStore {
1190 blas,
1191 entries,
1192 scratch_buffer_offset,
1193 } = storage;
1194 if blas.update_mode == wgt::AccelerationStructureUpdateMode::PreferUpdate {
1195 log::debug!("only rebuild implemented")
1196 }
1197 let raw = blas.try_raw(snatch_guard)?;
1198
1199 let state_lock = blas.compacted_state.lock();
1200 if let BlasCompactState::Compacted = *state_lock {
1201 return Err(BuildAccelerationStructureError::CompactedBlas(
1202 blas.error_ident(),
1203 ));
1204 }
1205
1206 if blas
1207 .flags
1208 .contains(wgpu_types::AccelerationStructureFlags::ALLOW_COMPACTION)
1209 {
1210 blases_compactable.push((blas.compaction_buffer.as_ref().unwrap().as_ref(), raw));
1211 }
1212 Ok(hal::BuildAccelerationStructureDescriptor {
1213 entries,
1214 mode: hal::AccelerationStructureBuildMode::Build,
1215 flags: blas.flags,
1216 source_acceleration_structure: None,
1217 destination_acceleration_structure: raw,
1218 scratch_buffer,
1219 scratch_buffer_offset: *scratch_buffer_offset,
1220 })
1221}
1222
1223fn build_blas<'a>(
1224 cmd_buf_raw: &mut dyn hal::DynCommandEncoder,
1225 blas_present: bool,
1226 tlas_present: bool,
1227 input_barriers: Vec<hal::BufferBarrier<dyn hal::DynBuffer>>,
1228 blas_descriptors: &[hal::BuildAccelerationStructureDescriptor<
1229 'a,
1230 dyn hal::DynBuffer,
1231 dyn hal::DynAccelerationStructure,
1232 >],
1233 scratch_buffer_barrier: hal::BufferBarrier<dyn hal::DynBuffer>,
1234 blas_s_for_compaction: Vec<(
1235 &'a dyn hal::DynBuffer,
1236 &'a dyn hal::DynAccelerationStructure,
1237 )>,
1238) {
1239 unsafe {
1240 cmd_buf_raw.transition_buffers(&input_barriers);
1241 }
1242
1243 if blas_present {
1244 unsafe {
1245 cmd_buf_raw.place_acceleration_structure_barrier(hal::AccelerationStructureBarrier {
1246 usage: hal::StateTransition {
1247 from: hal::AccelerationStructureUses::BUILD_INPUT,
1248 to: hal::AccelerationStructureUses::BUILD_OUTPUT,
1249 },
1250 });
1251
1252 cmd_buf_raw.build_acceleration_structures(blas_descriptors);
1253 }
1254 }
1255
1256 if blas_present && tlas_present {
1257 unsafe {
1258 cmd_buf_raw.transition_buffers(&[scratch_buffer_barrier]);
1259 }
1260 }
1261
1262 let mut source_usage = hal::AccelerationStructureUses::empty();
1263 let mut destination_usage = hal::AccelerationStructureUses::empty();
1264 for &(buf, blas) in blas_s_for_compaction.iter() {
1265 unsafe {
1266 cmd_buf_raw.transition_buffers(&[hal::BufferBarrier {
1267 buffer: buf,
1268 usage: hal::StateTransition {
1269 from: BufferUses::ACCELERATION_STRUCTURE_QUERY,
1270 to: BufferUses::ACCELERATION_STRUCTURE_QUERY,
1271 },
1272 }])
1273 }
1274 unsafe { cmd_buf_raw.read_acceleration_structure_compact_size(blas, buf) }
1275 destination_usage |= hal::AccelerationStructureUses::COPY_SRC;
1276 }
1277
1278 if blas_present {
1279 source_usage |= hal::AccelerationStructureUses::BUILD_OUTPUT;
1280 destination_usage |= hal::AccelerationStructureUses::BUILD_INPUT
1281 }
1282 if tlas_present {
1283 source_usage |= hal::AccelerationStructureUses::SHADER_INPUT;
1284 destination_usage |= hal::AccelerationStructureUses::BUILD_OUTPUT;
1285 }
1286 unsafe {
1287 cmd_buf_raw.place_acceleration_structure_barrier(hal::AccelerationStructureBarrier {
1288 usage: hal::StateTransition {
1289 from: source_usage,
1290 to: destination_usage,
1291 },
1292 });
1293 }
1294}