wgpu_core/command/
ray_tracing.rs

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    // We cannot move out of `tlases` because it will be borrowed by `tlas_storage`,
223    // but each package’s `instances` is consumed by the `mem::take` below.
224    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        // if the size is zero there is nothing to build
320        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                            // Reset the compacted state to idle. This means any prepares, before mapping their
506                            // internal buffer, will terminate.
507                            _ => 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(); // _tlas.dependencies.read();
586                    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
604///iterates over the blas iterator, and it's geometry, pushing the buffers into a storage vector (and also some validation).
605fn 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                        // `hal::AccelerationStructureTriangleIndices` accepts only `u32` offset
792                        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                        // `hal::AccelerationStructureTriangleTransform` accepts only `u32` offset
879                        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                        // `hal::AccelerationStructureAABBs` accepts only `u32` offset
1035                        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}