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