#![expect(
unsafe_code,
reason = "Building the TLAS without wgpu-core's per-instance cost requires wgpu_hal."
)]
#![cfg_attr(
not(any(
windows,
target_os = "linux",
target_os = "android",
target_os = "freebsd"
)),
expect(
dead_code,
unused_variables,
reason = "no backend with a raw TLAS build path is compiled in for this target"
)
)]
use bevy_render::{
render_resource::{Blas, Buffer, BufferDescriptor, BufferUsages, CommandEncoder, Tlas},
renderer::RenderDevice,
};
use core::marker::PhantomData;
use wgpu::{hal, BufferUses};
pub const INSTANCE_DESCRIPTOR_SIZE: u64 = 64;
const TLAS_BUILD_FLAGS: wgpu::AccelerationStructureFlags =
wgpu::AccelerationStructureFlags::PREFER_FAST_TRACE;
pub trait RawTlasBackend: Send + Sync {
fn scratch_size(
&self,
render_device: &RenderDevice,
instances: &Buffer,
instance_count: u32,
) -> Option<u64>;
fn create_scratch_buffer(&self, render_device: &RenderDevice, size: u64) -> Option<Buffer>;
fn build_tlas(
&self,
encoder: &mut CommandEncoder,
tlas: &mut Tlas,
instances: &Buffer,
instance_count: u32,
scratch: &Buffer,
) -> bool;
}
pub fn resolve(render_device: &RenderDevice) -> Option<&'static dyn RawTlasBackend> {
let device = render_device.wgpu_device();
#[cfg(any(
windows,
target_os = "linux",
target_os = "android",
target_os = "freebsd"
))]
if unsafe { device.as_hal::<hal::api::Vulkan>() }.is_some() {
return Some(&Hal::<hal::api::Vulkan>::NEW);
}
#[cfg(windows)]
if unsafe { device.as_hal::<hal::api::Dx12>() }.is_some() {
return Some(&Hal::<hal::api::Dx12>::NEW);
}
None
}
struct Hal<A>(PhantomData<fn() -> A>);
impl<A: hal::Api> Hal<A> {
const NEW: Self = Self(PhantomData);
fn record_build(
&self,
encoder: &mut CommandEncoder,
tlas: &mut Tlas,
instances: &Buffer,
instance_count: u32,
scratch: &Buffer,
) -> Option<()> {
use hal::CommandEncoder as _;
let buffer_address = |buffer: &Buffer| {
let guard = unsafe { buffer.as_hal::<A>() }?;
Some(core::ptr::from_ref::<A::Buffer>(&*guard))
};
let hal_instances = buffer_address(instances)?;
let hal_scratch = buffer_address(scratch)?;
let hal_tlas = unsafe { tlas.as_hal::<A>() }?;
let (hal_instances, hal_scratch) = unsafe { (&*hal_instances, &*hal_scratch) };
let entries =
hal::AccelerationStructureEntries::Instances(hal::AccelerationStructureInstances {
buffer: Some(hal_instances),
offset: 0,
count: instance_count,
});
let descriptor = hal::BuildAccelerationStructureDescriptor {
entries: &entries,
mode: hal::AccelerationStructureBuildMode::Build,
flags: TLAS_BUILD_FLAGS,
source_acceleration_structure: None,
destination_acceleration_structure: &*hal_tlas,
scratch_buffer: hal_scratch,
scratch_buffer_offset: 0,
};
unsafe {
encoder.as_hal_mut::<A, _, _>(|encoder| {
let encoder = encoder?;
encoder.place_acceleration_structure_barrier(hal::AccelerationStructureBarrier {
usage: hal::StateTransition {
from: hal::AccelerationStructureUses::SHADER_INPUT
| hal::AccelerationStructureUses::COPY_DST,
to: hal::AccelerationStructureUses::BUILD_OUTPUT
| hal::AccelerationStructureUses::BUILD_INPUT,
},
});
encoder.build_acceleration_structures(1, [descriptor]);
encoder.place_acceleration_structure_barrier(hal::AccelerationStructureBarrier {
usage: hal::StateTransition {
from: hal::AccelerationStructureUses::BUILD_OUTPUT,
to: hal::AccelerationStructureUses::SHADER_INPUT,
},
});
Some(())
})
}
}
}
impl<A: hal::Api> RawTlasBackend for Hal<A> {
fn scratch_size(
&self,
render_device: &RenderDevice,
instances: &Buffer,
instance_count: u32,
) -> Option<u64> {
use hal::Device as _;
let hal_device = unsafe { render_device.wgpu_device().as_hal::<A>() }?;
let hal_instances = unsafe { instances.as_hal::<A>() }?;
let entries =
hal::AccelerationStructureEntries::Instances(hal::AccelerationStructureInstances {
buffer: Some(&*hal_instances),
offset: 0,
count: instance_count,
});
let sizes = unsafe {
hal_device.get_acceleration_structure_build_sizes(
&hal::GetAccelerationStructureBuildSizesDescriptor {
entries: &entries,
flags: TLAS_BUILD_FLAGS,
},
)
};
Some(sizes.build_scratch_size)
}
fn create_scratch_buffer(&self, render_device: &RenderDevice, size: u64) -> Option<Buffer> {
use hal::Device as _;
let device = render_device.wgpu_device();
let hal_buffer = {
let hal_device = unsafe { device.as_hal::<A>() }?;
unsafe {
hal_device.create_buffer(&hal::BufferDescriptor {
label: Some("solari_tlas_scratch"),
size,
usage: BufferUses::ACCELERATION_STRUCTURE_SCRATCH,
memory_flags: hal::MemoryFlags::empty(),
})
}
.ok()?
};
let descriptor = BufferDescriptor {
label: Some("solari_tlas_scratch"),
size,
usage: BufferUsages::STORAGE,
mapped_at_creation: false,
};
Some(unsafe { device.create_buffer_from_hal::<A>(hal_buffer, &descriptor) }.into())
}
fn build_tlas(
&self,
encoder: &mut CommandEncoder,
tlas: &mut Tlas,
instances: &Buffer,
instance_count: u32,
scratch: &Buffer,
) -> bool {
if self
.record_build(encoder, tlas, instances, instance_count, scratch)
.is_none()
{
return false;
}
unsafe {
encoder.mark_acceleration_structures_built(core::iter::empty::<&Blas>(), [&*tlas]);
}
true
}
}