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_exact(&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 mut dependencies = Vec::new();
561 for action in &self.as_actions {
562 match action {
563 AsAction::Build(build) => {
564 dependencies.extend(build.blas_s_built.iter().map(Arc::clone));
565 }
566 AsAction::UseTlas(tlas) => {
567 dependencies.extend(tlas.dependencies.read().iter().map(Arc::clone));
568 }
569 }
570 }
571
572 let raw_dependencies = dependencies
573 .iter()
574 .flat_map(|blas| blas.raw(snatch_guard))
575 .collect::<Vec<_>>();
576
577 if !dependencies.is_empty() {
578 unsafe {
579 self.encoder
580 .raw
581 .set_acceleration_structure_dependencies(&self.encoder.list, &raw_dependencies);
582 }
583 }
584 }
585}
586
587fn iter_blas<'snatch_guard: 'buffers, 'buffers>(
589 blas_iter: impl Iterator<Item = &'buffers OwnedBlasBuildEntry<ArcReferences>>,
590 build_command: &mut AsBuild,
591 input_barriers: &mut Vec<hal::BufferBarrier<'buffers, dyn hal::DynBuffer>>,
592 scratch_buffer_blas_size: &mut u64,
593 blas_storage: &mut Vec<BlasStore<'buffers>>,
594 state: &mut EncodingState<'snatch_guard, '_>,
595) -> Result<(), BuildAccelerationStructureError> {
596 for entry in blas_iter {
597 let blas = &entry.blas;
598 state.tracker.blas_s.insert_single(blas.clone());
599
600 build_command.blas_s_built.push(blas.clone());
601
602 match &entry.geometries {
603 ArcBlasGeometries::TriangleGeometries(triangle_geometries) => {
604 let mut triangle_entries =
605 Vec::<hal::AccelerationStructureTriangles<dyn hal::DynBuffer>>::new();
606
607 for (i, mesh) in triangle_geometries.iter().enumerate() {
608 let size_desc = match &blas.sizes {
609 wgt::BlasGeometrySizeDescriptors::Triangles { descriptors } => descriptors,
610 _ => {
611 return Err(BuildAccelerationStructureError::BlasGeometryKindMismatch(
612 blas.error_ident(),
613 ));
614 }
615 };
616 if i >= size_desc.len() {
617 return Err(BuildAccelerationStructureError::IncompatibleBlasBuildSizes(
618 blas.error_ident(),
619 ));
620 }
621 let size_desc = &size_desc[i];
622
623 if size_desc.flags != mesh.size.flags {
624 return Err(BuildAccelerationStructureError::IncompatibleBlasFlags(
625 blas.error_ident(),
626 size_desc.flags,
627 mesh.size.flags,
628 ));
629 }
630
631 if size_desc.vertex_count < mesh.size.vertex_count {
632 return Err(
633 BuildAccelerationStructureError::IncompatibleBlasVertexCount(
634 blas.error_ident(),
635 size_desc.vertex_count,
636 mesh.size.vertex_count,
637 ),
638 );
639 }
640
641 if size_desc.vertex_format != mesh.size.vertex_format {
642 return Err(BuildAccelerationStructureError::DifferentBlasVertexFormats(
643 blas.error_ident(),
644 size_desc.vertex_format,
645 mesh.size.vertex_format,
646 ));
647 }
648
649 if size_desc
650 .vertex_format
651 .min_acceleration_structure_vertex_stride()
652 > mesh.vertex_stride
653 {
654 return Err(BuildAccelerationStructureError::VertexStrideTooSmall(
655 blas.error_ident(),
656 size_desc
657 .vertex_format
658 .min_acceleration_structure_vertex_stride(),
659 mesh.vertex_stride,
660 ));
661 }
662
663 if mesh.vertex_stride
664 % size_desc
665 .vertex_format
666 .acceleration_structure_stride_alignment()
667 != 0
668 {
669 return Err(BuildAccelerationStructureError::VertexStrideUnaligned(
670 blas.error_ident(),
671 size_desc
672 .vertex_format
673 .acceleration_structure_stride_alignment(),
674 mesh.vertex_stride,
675 ));
676 }
677
678 match (size_desc.index_count, mesh.size.index_count) {
679 (Some(_), None) | (None, Some(_)) => {
680 return Err(
681 BuildAccelerationStructureError::BlasIndexCountProvidedMismatch(
682 blas.error_ident(),
683 ),
684 )
685 }
686 (Some(create), Some(build)) if create < build => {
687 return Err(
688 BuildAccelerationStructureError::IncompatibleBlasIndexCount(
689 blas.error_ident(),
690 create,
691 build,
692 ),
693 )
694 }
695 _ => {}
696 }
697
698 if size_desc.index_format != mesh.size.index_format {
699 return Err(BuildAccelerationStructureError::DifferentBlasIndexFormats(
700 blas.error_ident(),
701 size_desc.index_format,
702 mesh.size.index_format,
703 ));
704 }
705
706 if size_desc.index_count.is_some() && mesh.index_buffer.is_none() {
707 return Err(BuildAccelerationStructureError::MissingIndexBuffer(
708 blas.error_ident(),
709 ));
710 }
711 let vertex_buffer = mesh.vertex_buffer.clone();
712 let vertex_pending = state.tracker.buffers.set_single(
713 &vertex_buffer,
714 BufferUses::BOTTOM_LEVEL_ACCELERATION_STRUCTURE_INPUT,
715 );
716 let vertex_buffer = {
717 let vertex_raw = mesh.vertex_buffer.as_ref().try_raw(state.snatch_guard)?;
718 let vertex_buffer = &mesh.vertex_buffer;
719 vertex_buffer.check_usage(BufferUsages::BLAS_INPUT)?;
720
721 if let Some(barrier) = vertex_pending.map(|pending| {
722 pending.into_hal(vertex_buffer.as_ref(), state.snatch_guard)
723 }) {
724 input_barriers.push(barrier);
725 }
726 if u64::from(mesh.size.vertex_count)
727 .checked_add(u64::from(mesh.first_vertex))
728 .and_then(|end_vertex| end_vertex.checked_mul(mesh.vertex_stride))
729 .is_none_or(|end| vertex_buffer.size < end)
730 {
731 return Err(BuildAccelerationStructureError::InsufficientBufferSize {
732 buffer_ident: vertex_buffer.error_ident(),
733 offset: u64::from(mesh.first_vertex)
734 .saturating_mul(mesh.vertex_stride),
735 region_size: u64::from(mesh.size.vertex_count)
736 .saturating_mul(mesh.vertex_stride),
737 buffer_size: vertex_buffer.size,
738 });
739 }
740 let vertex_buffer_offset = mesh.first_vertex as u64 * mesh.vertex_stride;
741 state.buffer_memory_init_actions.extend(
742 vertex_buffer.initialization_status.read().create_action(
743 vertex_buffer,
744 vertex_buffer_offset
745 ..(vertex_buffer_offset
746 + mesh.size.vertex_count as u64 * mesh.vertex_stride),
747 MemoryInitKind::NeedsInitializedMemory,
748 ),
749 );
750 vertex_raw
751 };
752 let index_buffer = if let Some(ref index_buffer) = mesh.index_buffer {
753 if mesh.first_index.is_none()
754 || mesh.size.index_count.is_none()
755 || mesh.size.index_count.is_none()
756 {
757 return Err(BuildAccelerationStructureError::MissingAssociatedData(
758 index_buffer.error_ident(),
759 ));
760 }
761 let index_pending = state.tracker.buffers.set_single(
762 index_buffer,
763 BufferUses::BOTTOM_LEVEL_ACCELERATION_STRUCTURE_INPUT,
764 );
765 let index_raw = index_buffer.try_raw(state.snatch_guard)?;
766 index_buffer.check_usage(BufferUsages::BLAS_INPUT)?;
767
768 if let Some(barrier) = index_pending.map(|pending| {
769 pending.into_hal(index_buffer.as_ref(), state.snatch_guard)
770 }) {
771 input_barriers.push(barrier);
772 }
773 let index_stride = mesh.size.index_format.unwrap().byte_size();
774 let Some(vertex_offset) =
776 mesh.first_index.unwrap().checked_mul(index_stride)
777 else {
778 return Err(BuildAccelerationStructureError::OffsetLimitedTo4GB {
779 buffer_ident: index_buffer.error_ident(),
780 offset: u64::from(mesh.first_index.unwrap())
781 .saturating_mul(u64::from(index_stride)),
782 count: u64::from(mesh.first_index.unwrap()),
783 stride: u64::from(index_stride),
784 });
785 };
786 let Some(indexes_size) =
787 mesh.size.index_count.unwrap().checked_mul(index_stride)
788 else {
789 return Err(BuildAccelerationStructureError::OffsetLimitedTo4GB {
790 buffer_ident: index_buffer.error_ident(),
791 offset: u64::from(mesh.size.index_count.unwrap())
792 .saturating_mul(u64::from(index_stride)),
793 count: u64::from(mesh.size.index_count.unwrap()),
794 stride: u64::from(index_stride),
795 });
796 };
797
798 if mesh.size.index_count.unwrap() % 3 != 0 {
799 return Err(BuildAccelerationStructureError::InvalidIndexCount(
800 index_buffer.error_ident(),
801 mesh.size.index_count.unwrap(),
802 ));
803 }
804 if index_buffer.size < u64::from(vertex_offset)
805 || index_buffer.size - u64::from(vertex_offset)
806 < u64::from(indexes_size)
807 {
808 return Err(BuildAccelerationStructureError::InsufficientBufferSize {
809 buffer_ident: index_buffer.error_ident(),
810 offset: u64::from(vertex_offset),
811 region_size: u64::from(indexes_size),
812 buffer_size: index_buffer.size,
813 });
814 }
815
816 state.buffer_memory_init_actions.extend(
817 index_buffer.initialization_status.read().create_action(
818 index_buffer,
819 u64::from(vertex_offset)
820 ..(u64::from(vertex_offset) + u64::from(indexes_size)),
821 MemoryInitKind::NeedsInitializedMemory,
822 ),
823 );
824 Some((index_raw, vertex_offset))
825 } else {
826 None
827 };
828 let transform_buffer = if let Some(ref transform_buffer) = mesh.transform_buffer
829 {
830 if !blas
831 .flags
832 .contains(wgt::AccelerationStructureFlags::USE_TRANSFORM)
833 {
834 return Err(BuildAccelerationStructureError::UseTransformMissing(
835 blas.error_ident(),
836 ));
837 }
838 if mesh.transform_buffer_offset.is_none() {
839 return Err(BuildAccelerationStructureError::MissingAssociatedData(
840 transform_buffer.error_ident(),
841 ));
842 }
843 let transform_pending = state.tracker.buffers.set_single(
844 transform_buffer,
845 BufferUses::BOTTOM_LEVEL_ACCELERATION_STRUCTURE_INPUT,
846 );
847 if mesh.transform_buffer_offset.is_none() {
848 return Err(BuildAccelerationStructureError::MissingAssociatedData(
849 transform_buffer.error_ident(),
850 ));
851 }
852 let transform_raw = transform_buffer.try_raw(state.snatch_guard)?;
853 transform_buffer.check_usage(BufferUsages::BLAS_INPUT)?;
854
855 if let Some(barrier) = transform_pending.map(|pending| {
856 pending.into_hal(transform_buffer.as_ref(), state.snatch_guard)
857 }) {
858 input_barriers.push(barrier);
859 }
860
861 let Ok(offset) = u32::try_from(mesh.transform_buffer_offset.unwrap())
863 else {
864 return Err(BuildAccelerationStructureError::OffsetLimitedTo4GB {
865 buffer_ident: transform_buffer.error_ident(),
866 offset: mesh.transform_buffer_offset.unwrap(),
867 count: mesh.transform_buffer_offset.unwrap(),
868 stride: 1,
869 });
870 };
871
872 if offset % wgt::TRANSFORM_BUFFER_ALIGNMENT as u32 != 0 {
873 return Err(
874 BuildAccelerationStructureError::UnalignedTransformBufferOffset(
875 transform_buffer.error_ident(),
876 ),
877 );
878 }
879 if transform_buffer.size < 48
880 || transform_buffer.size - 48 < u64::from(offset)
881 {
882 return Err(BuildAccelerationStructureError::InsufficientBufferSize {
883 buffer_ident: transform_buffer.error_ident(),
884 region_size: 48,
885 offset: u64::from(offset),
886 buffer_size: transform_buffer.size,
887 });
888 }
889 state.buffer_memory_init_actions.extend(
890 transform_buffer.initialization_status.read().create_action(
891 transform_buffer,
892 u64::from(offset)..(u64::from(offset) + 48),
893 MemoryInitKind::NeedsInitializedMemory,
894 ),
895 );
896 Some((transform_raw, offset))
897 } else {
898 if blas
899 .flags
900 .contains(wgt::AccelerationStructureFlags::USE_TRANSFORM)
901 {
902 return Err(BuildAccelerationStructureError::TransformMissing(
903 blas.error_ident(),
904 ));
905 }
906 None
907 };
908
909 let triangles = hal::AccelerationStructureTriangles {
910 vertex_buffer: Some(vertex_buffer),
911 vertex_format: mesh.size.vertex_format,
912 first_vertex: mesh.first_vertex,
913 vertex_count: mesh.size.vertex_count,
914 vertex_stride: mesh.vertex_stride,
915 indices: index_buffer.map(|(index_raw, vertex_offset)| {
916 hal::AccelerationStructureTriangleIndices::<dyn hal::DynBuffer> {
917 format: mesh.size.index_format.unwrap(),
918 buffer: Some(index_raw),
919 offset: vertex_offset,
920 count: mesh.size.index_count.unwrap(),
921 }
922 }),
923 transform: transform_buffer.map(|(transform_raw, offset)| {
924 hal::AccelerationStructureTriangleTransform {
925 buffer: transform_raw,
926 offset,
927 }
928 }),
929 flags: mesh.size.flags,
930 };
931 triangle_entries.push(triangles);
932 }
933
934 {
935 let scratch_buffer_offset = *scratch_buffer_blas_size;
936 *scratch_buffer_blas_size = scratch_buffer_blas_size.saturating_add(align_to(
937 blas.size_info.build_scratch_size,
938 u64::from(state.device.alignments.ray_tracing_scratch_buffer_alignment),
939 ));
940
941 blas_storage.push(BlasStore {
942 blas: blas.clone(),
943 entries: hal::AccelerationStructureEntries::Triangles(triangle_entries),
944 scratch_buffer_offset,
945 });
946 }
947 }
948 ArcBlasGeometries::AabbGeometries(aabb_geometries) => {
949 let mut aabb_entries =
950 Vec::<hal::AccelerationStructureAABBs<dyn hal::DynBuffer>>::new();
951
952 for (i, aabb) in aabb_geometries.iter().enumerate() {
953 let size_desc = match &blas.sizes {
954 wgt::BlasGeometrySizeDescriptors::AABBs { descriptors } => descriptors,
955 _ => {
956 return Err(BuildAccelerationStructureError::BlasGeometryKindMismatch(
957 blas.error_ident(),
958 ));
959 }
960 };
961 if i >= size_desc.len() {
962 return Err(BuildAccelerationStructureError::IncompatibleBlasBuildSizes(
963 blas.error_ident(),
964 ));
965 }
966 let size_desc = &size_desc[i];
967
968 if size_desc.flags != aabb.size.flags {
969 return Err(BuildAccelerationStructureError::IncompatibleBlasFlags(
970 blas.error_ident(),
971 size_desc.flags,
972 aabb.size.flags,
973 ));
974 }
975
976 if size_desc.primitive_count < aabb.size.primitive_count {
977 return Err(
978 BuildAccelerationStructureError::IncompatibleBlasAabbPrimitiveCount(
979 blas.error_ident(),
980 size_desc.primitive_count,
981 aabb.size.primitive_count,
982 ),
983 );
984 }
985
986 if aabb.primitive_offset % 8 != 0 {
987 return Err(
988 BuildAccelerationStructureError::UnalignedAabbPrimitiveOffset(
989 blas.error_ident(),
990 ),
991 );
992 }
993
994 if aabb.stride < wgt::AABB_GEOMETRY_MIN_STRIDE || aabb.stride % 8 != 0 {
995 return Err(BuildAccelerationStructureError::InvalidAabbStride(
996 blas.error_ident(),
997 aabb.stride,
998 ));
999 }
1000
1001 let aabb_buffer = aabb.aabb_buffer.clone();
1002 let aabb_pending = state.tracker.buffers.set_single(
1003 &aabb_buffer,
1004 BufferUses::BOTTOM_LEVEL_ACCELERATION_STRUCTURE_INPUT,
1005 );
1006 let aabb_raw = {
1007 let aabb_raw = aabb.aabb_buffer.as_ref().try_raw(state.snatch_guard)?;
1008 let aabb_buffer = &aabb.aabb_buffer;
1009 aabb_buffer.check_usage(BufferUsages::BLAS_INPUT)?;
1010
1011 if let Some(barrier) = aabb_pending.map(|pending| {
1012 pending.into_hal(aabb_buffer.as_ref(), state.snatch_guard)
1013 }) {
1014 input_barriers.push(barrier);
1015 }
1016
1017 let Some(aabb_size) = u32::try_from(aabb.stride)
1019 .ok()
1020 .and_then(|stride| aabb.size.primitive_count.checked_mul(stride))
1021 else {
1022 return Err(BuildAccelerationStructureError::OffsetLimitedTo4GB {
1023 buffer_ident: aabb_buffer.error_ident(),
1024 offset: u64::from(aabb.size.primitive_count)
1025 .saturating_mul(aabb.stride),
1026 count: u64::from(aabb.size.primitive_count),
1027 stride: aabb.stride,
1028 });
1029 };
1030
1031 if aabb_buffer.size < u64::from(aabb.primitive_offset)
1032 || aabb_buffer.size - u64::from(aabb.primitive_offset)
1033 < u64::from(aabb_size)
1034 {
1035 return Err(BuildAccelerationStructureError::InsufficientBufferSize {
1036 buffer_ident: aabb_buffer.error_ident(),
1037 region_size: u64::from(aabb_size),
1038 offset: u64::from(aabb.primitive_offset),
1039 buffer_size: aabb_buffer.size,
1040 });
1041 }
1042
1043 state.buffer_memory_init_actions.extend(
1044 aabb_buffer.initialization_status.read().create_action(
1045 aabb_buffer,
1046 u64::from(aabb.primitive_offset)
1047 ..u64::from(aabb.primitive_offset) + u64::from(aabb_size),
1048 MemoryInitKind::NeedsInitializedMemory,
1049 ),
1050 );
1051 aabb_raw
1052 };
1053
1054 aabb_entries.push(hal::AccelerationStructureAABBs {
1055 buffer: Some(aabb_raw),
1056 offset: aabb.primitive_offset,
1057 count: aabb.size.primitive_count,
1058 stride: aabb.stride,
1059 flags: aabb.size.flags,
1060 });
1061 }
1062
1063 {
1064 let scratch_buffer_offset = *scratch_buffer_blas_size;
1065 *scratch_buffer_blas_size = scratch_buffer_blas_size.saturating_add(align_to(
1066 blas.size_info.build_scratch_size,
1067 u64::from(state.device.alignments.ray_tracing_scratch_buffer_alignment),
1068 ));
1069
1070 blas_storage.push(BlasStore {
1071 blas: blas.clone(),
1072 entries: hal::AccelerationStructureEntries::AABBs(aabb_entries),
1073 scratch_buffer_offset,
1074 });
1075 }
1076 }
1077 }
1078 }
1079 Ok(())
1080}
1081
1082fn map_blas<'a>(
1083 storage: &'a BlasStore<'_>,
1084 scratch_buffer: &'a dyn hal::DynBuffer,
1085 snatch_guard: &'a SnatchGuard,
1086 blases_compactable: &mut Vec<(
1087 &'a dyn hal::DynBuffer,
1088 &'a dyn hal::DynAccelerationStructure,
1089 )>,
1090) -> Result<
1091 hal::BuildAccelerationStructureDescriptor<
1092 'a,
1093 dyn hal::DynBuffer,
1094 dyn hal::DynAccelerationStructure,
1095 >,
1096 BuildAccelerationStructureError,
1097> {
1098 let BlasStore {
1099 blas,
1100 entries,
1101 scratch_buffer_offset,
1102 } = storage;
1103 if blas.update_mode == wgt::AccelerationStructureUpdateMode::PreferUpdate {
1104 log::debug!("only rebuild implemented")
1105 }
1106 let raw = blas.try_raw(snatch_guard)?;
1107
1108 let state_lock = blas.compacted_state.lock();
1109 if let BlasCompactState::Compacted = *state_lock {
1110 return Err(BuildAccelerationStructureError::CompactedBlas(
1111 blas.error_ident(),
1112 ));
1113 }
1114
1115 if blas
1116 .flags
1117 .contains(wgpu_types::AccelerationStructureFlags::ALLOW_COMPACTION)
1118 {
1119 blases_compactable.push((blas.compaction_buffer.as_ref().unwrap().as_ref(), raw));
1120 }
1121 Ok(hal::BuildAccelerationStructureDescriptor {
1122 entries,
1123 mode: hal::AccelerationStructureBuildMode::Build,
1124 flags: blas.flags,
1125 source_acceleration_structure: None,
1126 destination_acceleration_structure: raw,
1127 scratch_buffer,
1128 scratch_buffer_offset: *scratch_buffer_offset,
1129 })
1130}
1131
1132fn build_blas<'a>(
1133 cmd_buf_raw: &mut dyn hal::DynCommandEncoder,
1134 blas_present: bool,
1135 tlas_present: bool,
1136 input_barriers: Vec<hal::BufferBarrier<dyn hal::DynBuffer>>,
1137 blas_descriptors: &[hal::BuildAccelerationStructureDescriptor<
1138 'a,
1139 dyn hal::DynBuffer,
1140 dyn hal::DynAccelerationStructure,
1141 >],
1142 scratch_buffer_barrier: hal::BufferBarrier<dyn hal::DynBuffer>,
1143 blas_s_for_compaction: Vec<(
1144 &'a dyn hal::DynBuffer,
1145 &'a dyn hal::DynAccelerationStructure,
1146 )>,
1147) {
1148 unsafe {
1149 cmd_buf_raw.transition_buffers(&input_barriers);
1150 }
1151
1152 if blas_present {
1153 unsafe {
1154 cmd_buf_raw.place_acceleration_structure_barrier(hal::AccelerationStructureBarrier {
1155 usage: hal::StateTransition {
1156 from: hal::AccelerationStructureUses::BUILD_INPUT,
1157 to: hal::AccelerationStructureUses::BUILD_OUTPUT,
1158 },
1159 });
1160
1161 cmd_buf_raw.build_acceleration_structures(blas_descriptors);
1162 }
1163 }
1164
1165 if blas_present && tlas_present {
1166 unsafe {
1167 cmd_buf_raw.transition_buffers(&[scratch_buffer_barrier]);
1168 }
1169 }
1170
1171 let mut source_usage = hal::AccelerationStructureUses::empty();
1172 let mut destination_usage = hal::AccelerationStructureUses::empty();
1173 for &(buf, blas) in blas_s_for_compaction.iter() {
1174 unsafe {
1175 cmd_buf_raw.transition_buffers(&[hal::BufferBarrier {
1176 buffer: buf,
1177 usage: hal::StateTransition {
1178 from: BufferUses::ACCELERATION_STRUCTURE_QUERY,
1179 to: BufferUses::ACCELERATION_STRUCTURE_QUERY,
1180 },
1181 }])
1182 }
1183 unsafe { cmd_buf_raw.read_acceleration_structure_compact_size(blas, buf) }
1184 destination_usage |= hal::AccelerationStructureUses::COPY_SRC;
1185 }
1186
1187 if blas_present {
1188 source_usage |= hal::AccelerationStructureUses::BUILD_OUTPUT;
1189 destination_usage |= hal::AccelerationStructureUses::BUILD_INPUT
1190 }
1191 if tlas_present {
1192 source_usage |= hal::AccelerationStructureUses::SHADER_INPUT;
1193 destination_usage |= hal::AccelerationStructureUses::BUILD_OUTPUT;
1194 }
1195 unsafe {
1196 cmd_buf_raw.place_acceleration_structure_barrier(hal::AccelerationStructureBarrier {
1197 usage: hal::StateTransition {
1198 from: source_usage,
1199 to: destination_usage,
1200 },
1201 });
1202 }
1203}