#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