ddgi_classify.cpp

cross platform rendering playground

src/passes/ddgi/ddgi_classify.cpp

1.65 KB
#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