#include <volk.h>
#include <cstring>
#include "backend/vulkan/debug.h"
#include "backend/vulkan/utils.h"
#include "backend/vulkan/vk_api.h"
#include "backend/vulkan/vk_conversion.h"
namespace rhi {
using namespace vk;
static_assert(
sizeof(AccelStructInstance) == sizeof(VkAccelerationStructureInstanceKHR), "AccelStructInstance size mismatch"
);
static_assert(
alignof(AccelStructInstance) == alignof(VkAccelerationStructureInstanceKHR),
"AccelStructInstance alignment mismatch"
);
static_assert(
static_cast<u32>(AccelStructInstanceFlags::TriangleFacingCullDisable) ==
VK_GEOMETRY_INSTANCE_TRIANGLE_FACING_CULL_DISABLE_BIT_KHR
);
static_assert(
static_cast<u32>(AccelStructInstanceFlags::TriangleFlipFacing) == VK_GEOMETRY_INSTANCE_TRIANGLE_FLIP_FACING_BIT_KHR
);
static_assert(static_cast<u32>(AccelStructInstanceFlags::ForceOpaque) == VK_GEOMETRY_INSTANCE_FORCE_OPAQUE_BIT_KHR);
static_assert(
static_cast<u32>(AccelStructInstanceFlags::ForceNoOpaque) == VK_GEOMETRY_INSTANCE_FORCE_NO_OPAQUE_BIT_KHR
);
static VkAccelerationStructureTypeKHR to_vk_accel_struct_type(AccelStructType type) {
switch (type) {
case AccelStructType::BottomLevel:
return VK_ACCELERATION_STRUCTURE_TYPE_BOTTOM_LEVEL_KHR;
case AccelStructType::TopLevel:
return VK_ACCELERATION_STRUCTURE_TYPE_TOP_LEVEL_KHR;
}
return VK_ACCELERATION_STRUCTURE_TYPE_BOTTOM_LEVEL_KHR;
}
static VkGeometryTypeKHR to_vk_geometry_type(AccelStructGeometryType type) {
switch (type) {
case AccelStructGeometryType::Triangles:
return VK_GEOMETRY_TYPE_TRIANGLES_KHR;
case AccelStructGeometryType::Instances:
return VK_GEOMETRY_TYPE_INSTANCES_KHR;
}
return VK_GEOMETRY_TYPE_TRIANGLES_KHR;
}
static VkBuildAccelerationStructureFlagsKHR to_vk_build_flags(AccelStructBuildFlags flags) {
VkBuildAccelerationStructureFlagsKHR result = 0;
if (Any(flags, AccelStructBuildFlags::AllowUpdate)) {
result |= VK_BUILD_ACCELERATION_STRUCTURE_ALLOW_UPDATE_BIT_KHR;
}
if (Any(flags, AccelStructBuildFlags::AllowCompaction)) {
result |= VK_BUILD_ACCELERATION_STRUCTURE_ALLOW_COMPACTION_BIT_KHR;
}
if (Any(flags, AccelStructBuildFlags::PreferFastTrace)) {
result |= VK_BUILD_ACCELERATION_STRUCTURE_PREFER_FAST_TRACE_BIT_KHR;
}
if (Any(flags, AccelStructBuildFlags::PreferFastBuild)) {
result |= VK_BUILD_ACCELERATION_STRUCTURE_PREFER_FAST_BUILD_BIT_KHR;
}
if (Any(flags, AccelStructBuildFlags::LowMemory)) {
result |= VK_BUILD_ACCELERATION_STRUCTURE_LOW_MEMORY_BIT_KHR;
}
return result;
}
static VkGeometryFlagsKHR to_vk_geometry_flags(AccelStructGeometryFlags flags) {
VkGeometryFlagsKHR result = 0;
if (Any(flags, AccelStructGeometryFlags::Opaque)) {
result |= VK_GEOMETRY_OPAQUE_BIT_KHR;
}
if (Any(flags, AccelStructGeometryFlags::NoDuplicateAnyHitInvocation)) {
result |= VK_GEOMETRY_NO_DUPLICATE_ANY_HIT_INVOCATION_BIT_KHR;
}
return result;
}
static VkGeometryInstanceFlagsKHR to_vk_instance_flag(AccelStructInstanceFlags flags) {
VkGeometryInstanceFlagsKHR result = 0;
if (Any(flags, AccelStructInstanceFlags::TriangleFacingCullDisable)) {
result |= VK_GEOMETRY_INSTANCE_TRIANGLE_FACING_CULL_DISABLE_BIT_KHR;
}
if (Any(flags, AccelStructInstanceFlags::TriangleFlipFacing)) {
result |= VK_GEOMETRY_INSTANCE_TRIANGLE_FLIP_FACING_BIT_KHR;
}
if (Any(flags, AccelStructInstanceFlags::ForceOpaque)) {
result |= VK_GEOMETRY_INSTANCE_FORCE_OPAQUE_BIT_KHR;
}
if (Any(flags, AccelStructInstanceFlags::ForceNoOpaque)) {
result |= VK_GEOMETRY_INSTANCE_FORCE_NO_OPAQUE_BIT_KHR;
}
return result;
}
static VkFormat to_vk_vertex_format(VertexFormat format) {
switch (format) {
case VertexFormat::R32G32B32_FLOAT:
return VK_FORMAT_R32G32B32_SFLOAT;
case VertexFormat::R32G32B32A32_FLOAT:
return VK_FORMAT_R32G32B32A32_SFLOAT;
default:
return VK_FORMAT_UNDEFINED;
}
}
static VkAccelerationStructureInstanceKHR to_vk_instance(const AccelStructInstance &in) {
VkAccelerationStructureInstanceKHR out{};
memcpy(&out.transform, &in.transform, sizeof(out.transform));
out.instanceCustomIndex = in.instance_custom_index;
out.mask = in.mask;
out.instanceShaderBindingTableRecordOffset = in.sbt_record_offset;
out.flags = to_vk_instance_flag(static_cast<AccelStructInstanceFlags>(in.flags));
out.accelerationStructureReference = in.acceleration_structure_reference;
return out;
}
static VkAccelerationStructureGeometryKHR to_vk_geometry(const AccelStructGeometry &in) {
VkAccelerationStructureGeometryKHR out{};
out.sType = VK_STRUCTURE_TYPE_ACCELERATION_STRUCTURE_GEOMETRY_KHR;
out.geometryType = to_vk_geometry_type(in.type);
out.flags = to_vk_geometry_flags(in.flags);
if (in.type == AccelStructGeometryType::Triangles) {
out.geometry.triangles.sType = VK_STRUCTURE_TYPE_ACCELERATION_STRUCTURE_GEOMETRY_TRIANGLES_DATA_KHR;
out.geometry.triangles.vertexFormat = to_vk_vertex_format(in.triangles.vertex_format);
out.geometry.triangles.vertexData = vk::device_address_const_to_vk(in.triangles.vertex_data);
out.geometry.triangles.vertexStride = in.triangles.vertex_stride;
out.geometry.triangles.maxVertex = in.triangles.max_vertex;
out.geometry.triangles.indexType =
in.triangles.index_type == IndexType::UINT32 ? VK_INDEX_TYPE_UINT32 : VK_INDEX_TYPE_UINT16;
out.geometry.triangles.indexData = vk::device_address_const_to_vk(in.triangles.index_data);
out.geometry.triangles.transformData = vk::device_address_const_to_vk(in.triangles.transform_data);
} else {
out.geometry.instances.sType = VK_STRUCTURE_TYPE_ACCELERATION_STRUCTURE_GEOMETRY_INSTANCES_DATA_KHR;
out.geometry.instances.arrayOfPointers = in.instances.array_of_pointers ? VK_TRUE : VK_FALSE;
out.geometry.instances.data = vk::device_address_const_to_vk(in.instances.data);
}
return out;
}
bool get_acceleration_structure_build_sizes(
Device &device,
const AccelStructBuildGeometryInfo &build_info,
const u32 *max_primitive_counts,
AccelStructBuildSizesInfo &sizes
) {
std::vector<VkAccelerationStructureGeometryKHR> geometries;
geometries.reserve(build_info.geometry_count);
for (u32 i = 0; i < build_info.geometry_count; ++i) {
geometries.push_back(to_vk_geometry(build_info.geometries[i]));
}
VkAccelerationStructureBuildGeometryInfoKHR vk_build_info{};
vk_build_info.sType = VK_STRUCTURE_TYPE_ACCELERATION_STRUCTURE_BUILD_GEOMETRY_INFO_KHR;
vk_build_info.type = to_vk_accel_struct_type(build_info.type);
vk_build_info.flags = to_vk_build_flags(build_info.flags);
vk_build_info.geometryCount = build_info.geometry_count;
vk_build_info.pGeometries = geometries.data();
VkAccelerationStructureBuildSizesInfoKHR vk_sizes{};
vk_sizes.sType = VK_STRUCTURE_TYPE_ACCELERATION_STRUCTURE_BUILD_SIZES_INFO_KHR;
vkGetAccelerationStructureBuildSizesKHR(
device.logical, VK_ACCELERATION_STRUCTURE_BUILD_TYPE_DEVICE_KHR, &vk_build_info, max_primitive_counts, &vk_sizes
);
sizes.acceleration_structure_size = vk_sizes.accelerationStructureSize;
sizes.update_scratch_size = vk_sizes.updateScratchSize;
sizes.build_scratch_size = vk_sizes.buildScratchSize;
return true;
}
bool create_acceleration_structure(Device &device, const AccelStructCreateInfo &info, AccelStruct &out) {
VkAccelerationStructureCreateInfoKHR create_info{};
create_info.sType = VK_STRUCTURE_TYPE_ACCELERATION_STRUCTURE_CREATE_INFO_KHR;
create_info.buffer = info.buffer->_handle;
create_info.offset = info.offset;
create_info.size = info.size;
create_info.type = to_vk_accel_struct_type(info.type);
VkAccelerationStructureKHR as = VK_NULL_HANDLE;
VK_CHECK(vkCreateAccelerationStructureKHR(device.logical, &create_info, nullptr, &as));
out.handle = as;
out.buffer = info.buffer;
if (!info.name.empty()) {
VK_NAME_EX(device, as, info.name.c_str());
}
return true;
}
void destroy_acceleration_structure(Device &device, AccelStruct &as) {
if (as.handle) {
vkDestroyAccelerationStructureKHR(device.logical, as.handle, nullptr);
as.handle = VK_NULL_HANDLE;
as.buffer = nullptr;
}
}
u64 get_acceleration_structure_device_address(Device &device, const AccelStruct &as) {
VkAccelerationStructureDeviceAddressInfoKHR addr_info{};
addr_info.sType = VK_STRUCTURE_TYPE_ACCELERATION_STRUCTURE_DEVICE_ADDRESS_INFO_KHR;
addr_info.accelerationStructure = as.handle;
return vkGetAccelerationStructureDeviceAddressKHR(device.logical, &addr_info);
}
void cmd_build_acceleration_structures(
CmdBuffer &cmd,
const AccelStructBuildGeometryInfo *build_infos,
const AccelStructBuildRangeInfo *const *ranges,
u32 count
) {
std::vector<VkAccelerationStructureBuildGeometryInfoKHR> vk_build_infos(count);
std::vector<VkAccelerationStructureGeometryKHR> vk_geometries;
std::vector<VkAccelerationStructureBuildRangeInfoKHR> vk_ranges(count);
std::vector<const VkAccelerationStructureBuildRangeInfoKHR *> vk_range_ptrs(count);
for (u32 i = 0; i < count; ++i) {
const AccelStructBuildGeometryInfo &in = build_infos[i];
for (u32 g = 0; g < in.geometry_count; ++g) {
vk_geometries.push_back(to_vk_geometry(in.geometries[g]));
}
}
u32 geometry_offset = 0;
for (u32 i = 0; i < count; ++i) {
const AccelStructBuildGeometryInfo &in = build_infos[i];
vk_build_infos[i].sType = VK_STRUCTURE_TYPE_ACCELERATION_STRUCTURE_BUILD_GEOMETRY_INFO_KHR;
vk_build_infos[i].type = to_vk_accel_struct_type(in.type);
vk_build_infos[i].flags = to_vk_build_flags(in.flags);
vk_build_infos[i].geometryCount = in.geometry_count;
vk_build_infos[i].pGeometries = vk_geometries.data() + geometry_offset;
if (in.dst_acceleration_structure) {
vk_build_infos[i].dstAccelerationStructure = in.dst_acceleration_structure->handle;
}
vk_build_infos[i].scratchData.deviceAddress = in.scratch_data;
geometry_offset += in.geometry_count;
vk_ranges[i].primitiveCount = ranges[i]->primitive_count;
vk_ranges[i].primitiveOffset = ranges[i]->primitive_offset;
vk_ranges[i].firstVertex = ranges[i]->first_vertex;
vk_ranges[i].transformOffset = ranges[i]->transform_offset;
vk_range_ptrs[i] = &vk_ranges[i];
}
vkCmdBuildAccelerationStructuresKHR(cmd.handle, count, vk_build_infos.data(), vk_range_ptrs.data());
}
} // namespace rhi