build_indirect.cpp

cross platform rendering playground

src/passes/dispatch_mdi/build_indirect.cpp

4.39 KB
#include "build_indirect.h"

#include <cstdio>

#include "pipelines.h"
#include "renderer.h"
#include "shaders/renderer_types.h"

namespace build_indirect {

Resources create(FrameGraph &fg, u32 render_item_count, const char *tag) {
    char indirect_name[128];
    char visible_name[128];
    char count_name[128];
    (void)snprintf(indirect_name, sizeof(indirect_name), "fg/build-indirect/%s/indirect", tag);
    (void)snprintf(visible_name, sizeof(visible_name), "fg/build-indirect/%s/visible", tag);
    (void)snprintf(count_name, sizeof(count_name), "fg/build-indirect/%s/draw-count", tag);

    Resources res;
    const u64 indirect_size = render_item_count * rhi::DRAW_INDEXED_INDIRECT_STRIDE;
    const u64 visible_size = render_item_count * sizeof(u32) * 2;

    res.indirect = fg.add_buffer({
        .desc = {.size = indirect_size, .usage = rhi::BufferUsage::StorageBuffer | rhi::BufferUsage::ShaderAddr},
        .name = indirect_name,
    });
    res.visible = fg.add_buffer({
        .desc = {.size = visible_size, .usage = rhi::BufferUsage::StorageBuffer | rhi::BufferUsage::ShaderAddr},
        .name = visible_name,
    });
    res.draw_count = fg.add_buffer({
        .desc = {.size = sizeof(u32), .usage = rhi::BufferUsage::StorageBuffer | rhi::BufferUsage::ShaderAddr},
        .name = count_name,
    });
    return res;
};

void record(
    rhi::Device &device,
    Pipelines &pipelines,
    rhi::ShaderCompiler &sc,
    FrameGraph &fg,
    BufferHandle render_items,
    BufferHandle submeshes,
    Resources &res,
    ImageHandle hiz_image,
    bool enable_culling,
    const char *pass_name
) {

    rhi::ComputePipelineDesc pd;
    pd.device = &device;
    pd.set_shader(sc, "src/passes/dispatch_mdi/build_indirect.slang", "csMain");
    pipelines.add_compute("BuildIndirect", pd);

    rhi::ComputePipeline &pipeline = pipelines.get_compute("BuildIndirect");

    auto &clear = fg.add_pass("Clear Draw Count");
    clear.write(res.draw_count, rhi::ResourceState::TransferTo);
    clear.execute([=](PassContext &ctx) {
        auto &count_buf = ctx.get_buffer(res.draw_count);
        rhi::fill_buffer(ctx.cmd, count_buf, 0, sizeof(u32), 0);
    });

    auto &builder = fg.add_pass(pass_name);
    builder.read(render_items, rhi::ResourceState::StorageRead);
    builder.read(submeshes, rhi::ResourceState::StorageRead);
    if (enable_culling && hiz_image.id != UINT32_MAX) {
        builder.read(hiz_image, rhi::ResourceState::TextureSampleNonFragment);
    }
    builder.write(res.indirect, rhi::ResourceState::StorageReadWrite);
    builder.write(res.visible, rhi::ResourceState::StorageReadWrite);
    builder.write(res.draw_count, rhi::ResourceState::StorageReadWrite);

    builder.execute([=, &pipeline](PassContext &ctx) {
        // render-item buffer is the candidate list: one thread per item.
        u32 item_count = (u32)ctx.render_data->render_items.size();

        rhi::set_pipeline(ctx.cmd, pipeline);

        BuildIndirectPushConstants pc{};
        pc.frame = (u64)ctx.render_data->frame_ubo_view[ctx.frame_index].slot;

        // (BDA path: render-item/submesh buffers travel by device address.)

        pc.indirect_buffer = rhi::buffer_device_address(ctx.device, ctx.get_buffer(res.indirect));
        pc.visible_items = rhi::buffer_device_address(ctx.device, ctx.get_buffer(res.visible));
        pc.draw_count_buffer = rhi::buffer_device_address(ctx.device, ctx.get_buffer(res.draw_count));

        pc.item_count = item_count;

        if (enable_culling) {
            auto hiz_img = ctx.get_image(hiz_image);
            pc.hiz = (u32)ctx.view(hiz_image);
            pc.hiz_width = hiz_img.desc.width;
            pc.hiz_height = hiz_img.desc.height;
            pc.hiz_mip_count = hiz_img.desc.mip_levels;
        } else {
            pc.hiz = UINT64_MAX; // canonical null handle: skips occlusion test
            pc.hiz_width = 0;
            pc.hiz_height = 0;
            pc.hiz_mip_count = 0;
        }

        pc.enable_culling = enable_culling ? 1 : 0;
        {
            static_assert(sizeof(pc) <= 128, "push constants exceed global 128B range");
            char pc128[128] = {};
            std::memcpy(pc128, &pc, sizeof(pc));
            rhi::set_constants(ctx.cmd, pipeline, rhi::ShaderStage::ALL, sizeof(pc128), pc128);
        }

        u32 group_count_x = (item_count + 63u) / 64u;
        if (group_count_x > 0) {
            rhi::dispatch(ctx.cmd, group_count_x);
        }
    });
}

} // namespace build_indirect