#include "ddgi_classify.h"
#include "pipelines.h"
#include "renderer.h"
#include "shaders/renderer_types.h"
namespace ddgi_classify {
void record(
rhi::Device &device, Pipelines &pipelines, rhi::ShaderCompiler &sc, FrameGraph &fg, ddgi_trace::Resources &res
) {
rhi::ComputePipelineDesc pd;
pd.device = &device;
pd.set_shader(sc, "src/passes/ddgi/ddgi_classify.slang", "csMain");
rhi::ComputePipeline &pipeline = pipelines.add_compute("DdgiClassify", pd);
auto &builder = fg.add_pass("DDGI Classify");
builder.read(res.ray_samples, rhi::ResourceState::StorageRead);
builder.read_write(res.probe_state, rhi::ResourceState::StorageReadWrite);
builder.execute([=, &pipeline](PassContext &ctx) {
FrameData &frame = ctx.frame_data[ctx.frame_index];
u32 groups = (res.grid.probe_count + 63) / 64;
if (groups == 0) {
return;
}
rhi::set_pipeline(ctx.cmd, pipeline);
DdgiClassifyPushConstants pc{};
pc.frame = (u64)ctx.render_data->frame_ubo_view[ctx.frame_index].slot;
pc.frame_counter = ctx.frame_count;
pc.samples_buffer = rhi::buffer_device_address(ctx.device, ctx.get_buffer(res.ray_samples));
pc.probe_state = rhi::buffer_device_address(ctx.device, ctx.get_buffer(res.probe_state));
{
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, groups, 1, 1);
});
}
} // namespace ddgi_classify