accel_struct.cpp

cross platform rendering playground

src/backend/vulkan/accel_struct.cpp

10.51 KB
#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