Skip to main content

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_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                            // 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 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
587///iterates over the blas iterator, and it's geometry, pushing the buffers into a storage vector (and also some validation).
588fn 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                        // `hal::AccelerationStructureTriangleIndices` accepts only `u32` offset
775                        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                        // `hal::AccelerationStructureTriangleTransform` accepts only `u32` offset
862                        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                        // `hal::AccelerationStructureAABBs` accepts only `u32` offset
1018                        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}