context.cpp

cross platform rendering playground

src/backend/vulkan/context.cpp

14.13 KB
#define VOLK_IMPLEMENTATION
#include <volk.h>

#include <cstring>

#include "backend/rhi_api.h"
#include "backend/vulkan/debug.h"
#include "backend/vulkan/vk_api.h"
#include "backend/vulkan/vk_conversion.h"

namespace rhi {
using namespace vk;

static bool ext_available(const char *name, const std::vector<VkExtensionProperties> &list) {
    for (auto &e : list) {
        if (strcmp(name, e.extensionName) == 0) {
            return true;
        }
    }
    return false;
}

static void create_instance(Device &device, Config &cfg, VkInstance &vk_instance) {
    VEL_INFO("Initializing Vulkan Context...");

    u32 extension_count = 0;
    VK_CHECK(vkEnumerateInstanceExtensionProperties(nullptr, &extension_count, nullptr));

    std::vector<VkExtensionProperties> available(extension_count);
    VK_CHECK(vkEnumerateInstanceExtensionProperties(nullptr, &extension_count, available.data()));

    cfg.instance_extensions.push_back(VK_KHR_GET_SURFACE_CAPABILITIES_2_EXTENSION_NAME);

    VkDebugUtilsMessengerCreateInfoEXT debug_ci{
        .sType = VK_STRUCTURE_TYPE_DEBUG_UTILS_MESSENGER_CREATE_INFO_EXT,
        .pNext = nullptr,
        .messageSeverity =
            VK_DEBUG_UTILS_MESSAGE_SEVERITY_WARNING_BIT_EXT | VK_DEBUG_UTILS_MESSAGE_SEVERITY_ERROR_BIT_EXT,
        .messageType = VK_DEBUG_UTILS_MESSAGE_TYPE_GENERAL_BIT_EXT | VK_DEBUG_UTILS_MESSAGE_TYPE_VALIDATION_BIT_EXT |
                       VK_DEBUG_UTILS_MESSAGE_TYPE_PERFORMANCE_BIT_EXT,
        .pfnUserCallback = debug_callback,
    };

    static constexpr VkValidationFeatureEnableEXT val_enable[] = {
        VK_VALIDATION_FEATURE_ENABLE_SYNCHRONIZATION_VALIDATION_EXT,
        VK_VALIDATION_FEATURE_ENABLE_BEST_PRACTICES_EXT,
        // VK_VALIDATION_FEATURE_ENABLE_GPU_ASSISTED_EXT,
        // VK_VALIDATION_FEATURE_ENABLE_GPU_ASSISTED_RESERVE_BINDING_SLOT_EXT,
        // VK_VALIDATION_FEATURE_ENABLE_DEBUG_PRINTF_EXT, // Keep disabled if GPU_ASSISTED is active
    };

    VkValidationFeaturesEXT val_features{
        .sType = VK_STRUCTURE_TYPE_VALIDATION_FEATURES_EXT,
        .pNext = &debug_ci, // Chain debug messenger into validation setup
        .enabledValidationFeatureCount = static_cast<u32>(std::size(val_enable)),
        .pEnabledValidationFeatures = val_enable,
    };

    void *pnext_chain = nullptr;

    if (cfg.enable_validation) {
        cfg.instance_layers.push_back("VK_LAYER_KHRONOS_validation");

        if (ext_available(VK_EXT_DEBUG_UTILS_EXTENSION_NAME, available)) {
            cfg.instance_extensions.push_back(VK_EXT_DEBUG_UTILS_EXTENSION_NAME);
        }
        if (ext_available(VK_EXT_VALIDATION_FEATURES_EXTENSION_NAME, available)) {
            cfg.instance_extensions.push_back(VK_EXT_VALIDATION_FEATURES_EXTENSION_NAME);
        }

        pnext_chain = &val_features;
    }

    const VkApplicationInfo app_info{
        .sType = VK_STRUCTURE_TYPE_APPLICATION_INFO,
        .pApplicationName = "vel-renderer",
        .applicationVersion = VK_MAKE_VERSION(1, 0, 0),
        .pEngineName = "vel-renderer",
        .engineVersion = VK_MAKE_VERSION(1, 0, 0),
        .apiVersion = VK_API_VERSION_1_3,
    };

    const VkInstanceCreateInfo ci{
        .sType = VK_STRUCTURE_TYPE_INSTANCE_CREATE_INFO,
        .pNext = pnext_chain,
        .pApplicationInfo = &app_info,
        .enabledLayerCount = static_cast<u32>(cfg.instance_layers.size()),
        .ppEnabledLayerNames = cfg.instance_layers.data(),
        .enabledExtensionCount = static_cast<u32>(cfg.instance_extensions.size()),
        .ppEnabledExtensionNames = cfg.instance_extensions.data(),
    };

    VK_CHECK(vkCreateInstance(&ci, nullptr, &vk_instance));

    device.instance = vk_instance;
    volkLoadInstance(vk_instance);

    if (cfg.enable_validation) {
        create_debug_messenger(device);
    }
}

static void create_phys_device(Device &device, Config &cfg, VkInstance &vk_instance, VkPhysicalDevice &vk_physical) {
    u32 n = 0;
    VK_CHECK(vkEnumeratePhysicalDevices(vk_instance, &n, nullptr));
    if (n == 0) {
        VEL_CRITICAL("No Vulkan physical devices found");
        std::abort();
    }

    std::vector<VkPhysicalDevice> devices(n);
    VK_CHECK(vkEnumeratePhysicalDevices(vk_instance, &n, devices.data()));

    for (auto pd : devices) {
        VkPhysicalDeviceProperties props;
        vkGetPhysicalDeviceProperties(pd, &props);
        if (props.deviceType == VK_PHYSICAL_DEVICE_TYPE_DISCRETE_GPU) {
            vk_physical = pd;
            break;
        }
    }
    if (!vk_physical) {
        vk_physical = devices[0];
    }

    VkPhysicalDeviceProperties props;
    vkGetPhysicalDeviceProperties(vk_physical, &props);
    const char *type_str = (props.deviceType == VK_PHYSICAL_DEVICE_TYPE_DISCRETE_GPU) ? "Discrete" : "Integrated/Other";
    VEL_INFO("Selected GPU: {} ({})", props.deviceName, type_str);
    VEL_INFO(
        "Driver: {}.{}.{}",
        VK_API_VERSION_MAJOR(props.driverVersion),
        VK_API_VERSION_MINOR(props.driverVersion),
        VK_API_VERSION_PATCH(props.driverVersion)
    );

    device.physical = vk_physical;
    device.min_ubo_alignment = props.limits.minUniformBufferOffsetAlignment;
    device.timestamp_period = props.limits.timestampPeriod;

    VkPhysicalDeviceDescriptorIndexingProperties indexing{
        .sType = VK_STRUCTURE_TYPE_PHYSICAL_DEVICE_DESCRIPTOR_INDEXING_PROPERTIES
    };
    VkPhysicalDeviceAccelerationStructurePropertiesKHR accel_structure_props{
        .sType = VK_STRUCTURE_TYPE_PHYSICAL_DEVICE_ACCELERATION_STRUCTURE_PROPERTIES_KHR,
        .pNext = &indexing,
    };
    VkPhysicalDeviceVulkan12Properties props12{
        .sType = VK_STRUCTURE_TYPE_PHYSICAL_DEVICE_VULKAN_1_2_PROPERTIES,
        .pNext = &accel_structure_props,
    };
    VkPhysicalDeviceProperties2 props2{
        .sType = VK_STRUCTURE_TYPE_PHYSICAL_DEVICE_PROPERTIES_2,
        .pNext = &props12,
    };
    vkGetPhysicalDeviceProperties2(vk_physical, &props2);

    device.max_sampled_textures = props12.maxDescriptorSetUpdateAfterBindSampledImages;
    device.max_storage_buffers = props12.maxDescriptorSetUpdateAfterBindStorageBuffers;
    device.max_storage_images = props12.maxDescriptorSetUpdateAfterBindStorageImages;
    device.max_uniform_buffers = props12.maxDescriptorSetUpdateAfterBindUniformBuffers;
    device.max_samplers = props12.maxDescriptorSetUpdateAfterBindSamplers;
    device.max_acceleration_structures = accel_structure_props.maxDescriptorSetUpdateAfterBindAccelerationStructures;
};

static void create_logi_device(
    Device &device, Config &cfg, VkInstance &vk_instance, VkPhysicalDevice &vk_physical, VkDevice &vk_device
) {
    device.graphics = find_queue(device, VK_QUEUE_GRAPHICS_BIT, 0);

    f32 priority = 1.0f;
    std::vector<VkDeviceQueueCreateInfo> queue_infos;
    if (device.graphics.family_index != UINT32_MAX) {
        queue_infos.push_back(
            VkDeviceQueueCreateInfo{
                .sType = VK_STRUCTURE_TYPE_DEVICE_QUEUE_CREATE_INFO,
                .queueFamilyIndex = device.graphics.family_index,
                .queueCount = 1,
                .pQueuePriorities = &priority,
            }
        );
    }

    std::vector<VkExtensionProperties> available_exts;
    {
        u32 n = 0;
        VK_CHECK(vkEnumerateDeviceExtensionProperties(vk_physical, nullptr, &n, nullptr));
        available_exts.resize(n);
        VK_CHECK(vkEnumerateDeviceExtensionProperties(vk_physical, nullptr, &n, available_exts.data()));
    }

    std::vector<const char *> exts = cfg.device_extensions;
    auto push_ext = [&](const char *name) {
        if (ext_available(name, available_exts)) {
            exts.push_back(name);
        }
    };
    push_ext(VK_KHR_SWAPCHAIN_EXTENSION_NAME);
    push_ext(VK_EXT_EXTENDED_DYNAMIC_STATE_2_EXTENSION_NAME);
    push_ext(VK_KHR_DEFERRED_HOST_OPERATIONS_EXTENSION_NAME);
    push_ext(VK_KHR_ACCELERATION_STRUCTURE_EXTENSION_NAME);
    push_ext(VK_KHR_RAY_QUERY_EXTENSION_NAME);
    push_ext(VK_KHR_RAY_TRACING_PIPELINE_EXTENSION_NAME);
    push_ext(VK_EXT_MUTABLE_DESCRIPTOR_TYPE_EXTENSION_NAME);
    push_ext(VK_KHR_PUSH_DESCRIPTOR_EXTENSION_NAME);

    // gate behind validation?
    push_ext(VK_EXT_DEBUG_UTILS_EXTENSION_NAME);

    VkPhysicalDeviceFeatures2 features{VK_STRUCTURE_TYPE_PHYSICAL_DEVICE_FEATURES_2};
    PNEXTCHAIN_DECLARE(features.pNext);

    VkPhysicalDeviceVulkan11Features f11{VK_STRUCTURE_TYPE_PHYSICAL_DEVICE_VULKAN_1_1_FEATURES};
    PNEXTCHAIN_APPEND_STRUCT(f11);

    VkPhysicalDeviceVulkan12Features f12{VK_STRUCTURE_TYPE_PHYSICAL_DEVICE_VULKAN_1_2_FEATURES};
    PNEXTCHAIN_APPEND_STRUCT(f12);

    VkPhysicalDeviceVulkan13Features f13{VK_STRUCTURE_TYPE_PHYSICAL_DEVICE_VULKAN_1_3_FEATURES};
    PNEXTCHAIN_APPEND_STRUCT(f13);

    VkPhysicalDeviceExtendedDynamicState2FeaturesEXT eds2{
        VK_STRUCTURE_TYPE_PHYSICAL_DEVICE_EXTENDED_DYNAMIC_STATE_2_FEATURES_EXT
    };
    PNEXTCHAIN_APPEND_STRUCT(eds2);

    VkPhysicalDeviceAccelerationStructureFeaturesKHR as_f{
        VK_STRUCTURE_TYPE_PHYSICAL_DEVICE_ACCELERATION_STRUCTURE_FEATURES_KHR
    };
    PNEXTCHAIN_APPEND_STRUCT(as_f);

    VkPhysicalDeviceRayQueryFeaturesKHR rq_f{VK_STRUCTURE_TYPE_PHYSICAL_DEVICE_RAY_QUERY_FEATURES_KHR};
    PNEXTCHAIN_APPEND_STRUCT(rq_f);

    VkPhysicalDeviceMutableDescriptorTypeFeaturesEXT mutable_f{
        VK_STRUCTURE_TYPE_PHYSICAL_DEVICE_MUTABLE_DESCRIPTOR_TYPE_FEATURES_EXT
    };
    PNEXTCHAIN_APPEND_STRUCT(mutable_f);

    vkGetPhysicalDeviceFeatures2(vk_physical, &features);

    // VK 1.0 features
    features.features.samplerAnisotropy = VK_TRUE;
    features.features.multiDrawIndirect = VK_TRUE;
    features.features.drawIndirectFirstInstance = VK_TRUE;
    features.features.fillModeNonSolid = VK_TRUE;
    features.features.depthClamp = VK_TRUE;

    // needed for vk shader validation as well
    features.features.vertexPipelineStoresAndAtomics = VK_TRUE;
    features.features.fragmentStoresAndAtomics = VK_TRUE;

    features.features.shaderInt64 = VK_TRUE;
    features.features.shaderStorageImageReadWithoutFormat = VK_TRUE;
    features.features.shaderStorageImageWriteWithoutFormat = VK_TRUE;

    // VK 1.1 features
    f11.shaderDrawParameters = VK_TRUE;

    // VK 1.2 features (bindless architecture base)
    f12.descriptorIndexing = VK_TRUE;
    f12.runtimeDescriptorArray = VK_TRUE;
    f12.descriptorBindingPartiallyBound = VK_TRUE;
    f12.descriptorBindingVariableDescriptorCount = VK_TRUE;
    f12.descriptorBindingUpdateUnusedWhilePending = VK_TRUE;
    f12.descriptorBindingUniformBufferUpdateAfterBind = VK_TRUE;
    f12.descriptorBindingStorageBufferUpdateAfterBind = VK_TRUE;
    f12.descriptorBindingStorageImageUpdateAfterBind = VK_TRUE;
    f12.descriptorBindingSampledImageUpdateAfterBind = VK_TRUE;
    f12.shaderSampledImageArrayNonUniformIndexing = VK_TRUE;
    f12.shaderStorageBufferArrayNonUniformIndexing = VK_TRUE;
    f12.shaderUniformBufferArrayNonUniformIndexing = VK_TRUE;
    f12.shaderStorageImageArrayNonUniformIndexing = VK_TRUE;
    f12.shaderUniformTexelBufferArrayNonUniformIndexing = VK_TRUE;
    f12.shaderStorageTexelBufferArrayNonUniformIndexing = VK_TRUE;
    f12.timelineSemaphore = VK_TRUE;
    f12.drawIndirectCount = VK_TRUE;

    // I just end up manually aligning stuff anyway, but a nice to have
    f12.scalarBlockLayout = VK_TRUE;
    f12.bufferDeviceAddress = VK_TRUE;

    // VK 1.3 features
    f13.synchronization2 = VK_TRUE;
    f13.dynamicRendering = VK_TRUE;
    f13.pipelineCreationCacheControl = VK_TRUE;
    f13.shaderDemoteToHelperInvocation = VK_TRUE;

    eds2.extendedDynamicState2 = VK_TRUE;
    as_f.accelerationStructure = VK_TRUE;
    as_f.descriptorBindingAccelerationStructureUpdateAfterBind = VK_TRUE;
    rq_f.rayQuery = VK_TRUE;
    mutable_f.mutableDescriptorType = VK_TRUE;

    // requirements check
    if (!f13.synchronization2) {
        VEL_CRITICAL("synchronization2 not supported (Vulkan 1.3 required)");
        std::abort();
    }
    if (!f13.dynamicRendering) {
        VEL_CRITICAL("dynamicRendering not supported (Vulkan 1.3 required)");
        std::abort();
    }
    if (!f12.bufferDeviceAddress) {
        VEL_CRITICAL("bufferDeviceAddress not supported (Vulkan 1.2 required)");
        std::abort();
    }
    if (!f12.descriptorIndexing) {
        VEL_CRITICAL("descriptorIndexing not supported (Vulkan 1.2 required)");
        std::abort();
    }
    if (!mutable_f.mutableDescriptorType) {
        VEL_CRITICAL("VK_EXT_mutable_descriptor_type not supported (required)");
        std::abort();
    }

    VkDeviceCreateInfo ci{VK_STRUCTURE_TYPE_DEVICE_CREATE_INFO};
    ci.pNext = &features;
    ci.queueCreateInfoCount = (u32)queue_infos.size();
    ci.pQueueCreateInfos = queue_infos.data();
    ci.enabledExtensionCount = (u32)exts.size();
    ci.ppEnabledExtensionNames = exts.data();

    VK_CHECK(vkCreateDevice(vk_physical, &ci, nullptr, &vk_device));
    device.logical = vk_device;

    if (device.graphics.family_index != UINT32_MAX) {
        VkQueue q;
        vkGetDeviceQueue(vk_device, device.graphics.family_index, 0, &q);
        device.graphics.handle = q;
    }
}

void create_context(Device &device, Config &cfg) {
    VEL_INFO("Initializing Vulkan Context...");

    volkInitialize();

    VkInstance vk_inst{};
    create_instance(device, cfg, vk_inst);
    VkPhysicalDevice vk_phys{};
    create_phys_device(device, cfg, vk_inst, vk_phys);
    VkDevice vk_dev{};
    create_logi_device(device, cfg, vk_inst, vk_phys, vk_dev);

    volkLoadDevice(vk_dev);

    VK_NAME(device, device.instance);
    VK_NAME(device, device.logical);
    VK_NAME(device, device.graphics.handle);

    create_global_descriptors(device);
}

void destroy_context(Device &device) {
    VEL_INFO("Shutting down Vulkan Context...");
    vkDeviceWaitIdle(device.logical);

    destroy_global_descriptors(device);

    if (device.debug_messenger && device.instance) {
        destory_debug_messenger(device);
        device.debug_messenger = VK_NULL_HANDLE;
    }

    vkDestroyDevice(device.logical, nullptr);
    device.logical = VK_NULL_HANDLE;

    vkDestroyInstance(device.instance, nullptr);
    device.instance = VK_NULL_HANDLE;

    device.physical = VK_NULL_HANDLE;
    device.graphics = {};
    device.pipeline_cache = VK_NULL_HANDLE;
}

void device_wait_idle(Device &device) {
    VK_CHECK(vkDeviceWaitIdle(device.logical));
}

void queue_wait_idle(Device &device) {
    VK_CHECK(vkQueueWaitIdle(device.graphics.handle));
}

} // namespace rhi