light_cull.cpp

cross platform rendering playground

src/passes/cluster/light_cull.cpp

3.36 KB
#include "light_cull.h"

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

namespace light_cull {

Resources record(
    rhi::Device &device,
    Pipelines &pipelines,
    rhi::ShaderCompiler &sc,
    FrameGraph &fg,
    BufferHandle cluster_bounds,
    BufferHandle cluster_records
) {

    rhi::ComputePipelineDesc pd;
    pd.device = &device;
    pd.set_shader(sc, "src/passes/cluster/light_cull.slang", "csMain");
    rhi::ComputePipeline &pipeline = pipelines.add_compute("LightCull", pd);

    Resources out;
    out.light_index_list = fg.add_buffer({
        .desc =
            {
                .size = MAX_LIGHT_INDICES * sizeof(u32),
                .usage = rhi::BufferUsage::StorageBuffer | rhi::BufferUsage::ShaderAddr,
            },
        .name = "fg/light-cull/index-list",
    });
    out.light_index_counter = fg.add_buffer({
        .desc =
            {
                .size = sizeof(u32),
                .usage = rhi::BufferUsage::StorageBuffer | rhi::BufferUsage::ShaderAddr,

            },
        .name = "fg/light-cull/counter",
    });

    auto &builder = fg.add_pass("LightCull");
    builder.read(cluster_bounds, rhi::ResourceState::StorageRead);
    builder.write(cluster_records, rhi::ResourceState::StorageReadWrite);
    builder.write(out.light_index_list, rhi::ResourceState::StorageReadWrite);
    builder.write(out.light_index_counter, rhi::ResourceState::StorageReadWrite);

    builder.execute([=, &pipeline](PassContext &ctx) {
        auto &frame = ctx.frame_data[ctx.frame_index];

        u32 cx = (ctx.render_width + TILE_SIZE_X - 1) / TILE_SIZE_X;
        u32 cy = (ctx.render_height + TILE_SIZE_Y - 1) / TILE_SIZE_Y;
        u32 cz = CLUSTER_COUNT_Z;
        u32 total_clusters = cx * cy * cz;

        u32 group_count_x = (total_clusters + 63u) / 64u;
        if (group_count_x == 0) {
            return;
        }

        rhi::BufferView counter_view = ctx.buffer_view(out.light_index_counter, {.type = rhi::BufferViewType::Storage});
        rhi::Buffer &counter_buf = *counter_view.buffer;
        rhi::barrier(ctx.cmd, counter_buf, counter_buf.state, rhi::ResourceState::TransferTo);
        rhi::fill_buffer(ctx.cmd, counter_buf, 0, sizeof(u32), 0);
        rhi::barrier(ctx.cmd, counter_buf, rhi::ResourceState::TransferTo, rhi::ResourceState::StorageReadWrite);

        LightCull lc{};
        lc.clusters_records = rhi::buffer_device_address(ctx.device, ctx.get_buffer(cluster_records));
        lc.light_index_list = rhi::buffer_device_address(ctx.device, ctx.get_buffer(out.light_index_list));

        memcpy((u8 *)frame.frame_ubo.mapped + offsetof(FrameUBO, light_cull), &lc, sizeof(LightCull));

        LightCullPushConstants pc{};
        pc.frame = (u64)ctx.render_data->frame_ubo_view[ctx.frame_index].slot;
        pc.light_index_counter = rhi::buffer_device_address(ctx.device, ctx.get_buffer(out.light_index_counter));

        pc.light_count = 0;
        pc.phase = 0;

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

        {
            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);
        }
        rhi::dispatch(ctx.cmd, group_count_x);
    });
    return out;
}

} // namespace light_cull