ddgi_trace.cpp

cross platform rendering playground

src/passes/ddgi/ddgi_trace.cpp

7.25 KB
#include "ddgi_trace.h"

#include <cmath>
#include <cstring>

#include "backend/rhi.h"
#include "core/random.h"
#include "passes/debug/debug_draw.h"
#include "passes/fg.h"
#include "pipelines.h"
#include "renderer.h"
#include "shaders/renderer_types.h"

namespace ddgi_trace {

static vec3 debug_probe_position(u32 probe_idx, vec3 origin, vec3 spacing, vec3 dims_f) {
    u32 dx = (u32)dims_f.x;
    u32 dy = (u32)dims_f.y;
    dx = dx == 0 ? 1 : dx;
    dy = dy == 0 ? 1 : dy;
    u32 cx = probe_idx % dx;
    u32 cy = (probe_idx / dx) % dy;
    u32 cz = probe_idx / (dx * dy);
    vec3 coord = vec3((f32)cx, (f32)cy, (f32)cz);
    vec3 center = (vec3(dims_f.x, dims_f.y, dims_f.z) - vec3(1.0f)) * 0.5f;
    return vec3(
        origin.x + (coord.x - center.x) * spacing.x,
        origin.y + (coord.y - center.y) * spacing.y,
        origin.z + (coord.z - center.z) * spacing.z
    );
}

void debug_draw(DebugDraw &draw, const RenderData &rd) {
    const DdgiGridParams &grid = rd.ddgi_grid;
    if (grid.probe_count == 0) {
        return;
    }
    vec3 origin = vec3(grid.origin.x, grid.origin.y, grid.origin.z);
    vec3 spacing = vec3(grid.spacing.x, grid.spacing.y, grid.spacing.z);
    vec3 dims_f = vec3(grid.dims.x, grid.dims.y, grid.dims.z);
    if (spacing.x <= 0.0f || spacing.y <= 0.0f || spacing.z <= 0.0f) {
        return;
    }
    vec3 half_extent = vec3(
        (dims_f.x - 1.0f) * 0.5f * spacing.x, (dims_f.y - 1.0f) * 0.5f * spacing.y, (dims_f.z - 1.0f) * 0.5f * spacing.z
    );

    vec4 grid_color = vec4(0.4f, 1.6f, 2.0f, 1.0f);
    draw.aabb(origin - half_extent, origin + half_extent, grid_color, DEBUG_CAT_DDGI, DEBUG_DEPTH_TEST, 2.0f);

    vec4 probe_color = vec4(3.0f, 2.7f, 0.6f, 1.0f);
    f32 cross_size = (spacing.x + spacing.y + spacing.z) / 3.0f * 0.25f;
    if (cross_size <= 0.0f) {
        cross_size = 0.1f;
    }
    for (u32 i = 0; i < grid.probe_count; ++i) {
        vec3 p = debug_probe_position(i, origin, spacing, dims_f);
        draw.cross(p, cross_size, probe_color, DEBUG_CAT_DDGI, DEBUG_DEPTH_TEST, 2.0f);
    }
}

Resources create(FrameGraph &fg) {
    Resources out;

    out.grid.origin = {0.0f, 0.0f, 6.0f, 0.0};
    out.grid.dims = {16, 24, 12, 0};
    out.grid.spacing = {2.0, 2.0, 2.0, 0.0};
    out.grid.spacing.w = fmax(fmax(out.grid.spacing.x, out.grid.spacing.y), out.grid.spacing.z);
    out.grid.probe_count = out.grid.dims.x * out.grid.dims.y * out.grid.dims.z;

    u64 buf_size = (u64)out.grid.probe_count * RAYS_PER_PROBE * sizeof(ProbeRaySample);
    out.ray_samples = fg.add_buffer({
        .desc =
            {
                .size = buf_size,
                .usage = rhi::BufferUsage::StorageBuffer | rhi::BufferUsage::ShaderAddr,
            },
        .name = "fg/ddgi/ray-samples",
    });
    out.probe_state = fg.add_buffer({
        .desc =
            {
                .size = (u64)out.grid.probe_count * sizeof(ProbeState),
                .usage = rhi::BufferUsage::StorageBuffer | rhi::BufferUsage::ShaderAddr,
                .memory = rhi::BufferMemType::Upload,
            },
        .name = "fg/ddgi/probe-state",
        .lifetime = Lifetime::Persistent,
    });

    u32 irr_tiles_per_row = (u32)ceilf(sqrtf((f32)out.grid.probe_count));
    u32 irr_row_count = (out.grid.probe_count + irr_tiles_per_row - 1) / irr_tiles_per_row;
    u32 irr_width = irr_tiles_per_row * IRRADIANCE_TILE_STRIDE;
    u32 irr_height = irr_row_count * IRRADIANCE_TILE_STRIDE;
    out.irradiance_img = fg.add_image({
        .desc = {.width = irr_width, .height = irr_height, .format = rhi::ImageFormat::RGBA16_FLOAT},
        .name = "fg/ddgi/irradiance",
        .lifetime = Lifetime::History,
        .size_class = SizeOp::Absolute,
        .view = {.type = rhi::ImageViewType::Sampled},
    });

    u32 dist_tiles_per_row = (u32)ceilf(sqrtf((f32)out.grid.probe_count));
    u32 dist_row_count = (out.grid.probe_count + dist_tiles_per_row - 1) / dist_tiles_per_row;
    u32 dist_width = dist_tiles_per_row * DISTANCE_TILE_STRIDE;
    u32 dist_height = dist_row_count * DISTANCE_TILE_STRIDE;
    out.distance_img = fg.add_image({
        .desc{
            .width = dist_width,
            .height = dist_height,
            .format = rhi::ImageFormat::RG16_FLOAT,
        },
        .name = "fg/ddgi/distance",
        .lifetime = Lifetime::History,
        .size_class = SizeOp::Absolute,
        .view = {.type = rhi::ImageViewType::Sampled},
    });

    return out;
}

void record(
    rhi::Device &device,
    Pipelines &pipelines,
    rhi::ShaderCompiler &sc,
    FrameGraph &fg,
    Resources &out,
    ImageHandle env,
    BufferHandle tlas_buf
) {
    rhi::ComputePipelineDesc pd;
    pd.device = &device;
    pd.set_shader(sc, "src/passes/ddgi/ddgi_trace.slang", "csMain");
    rhi::ComputePipeline &pipeline = pipelines.add_compute("DdgiTrace", pd);

    auto &builder = fg.add_pass("DDGI Trace");
    builder.read(tlas_buf, rhi::ResourceState::AccelTrace);
    builder.read(out.probe_state, rhi::ResourceState::StorageRead);
    builder.write(out.ray_samples, rhi::ResourceState::StorageReadWrite);
    builder.read(out.irradiance_img, rhi::ResourceState::TextureSampleNonFragment);
    builder.read(out.distance_img, rhi::ResourceState::TextureSampleNonFragment);
    builder.execute([=, &pipeline](PassContext &ctx) {
        u32 probe_count = out.grid.probe_count;
        u32 total_rays = probe_count * RAYS_PER_PROBE;
        u32 groups = (total_rays + 63) / 64;
        if (groups == 0) {
            return;
        }

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

        FrameData &frame = ctx.frame_data[ctx.frame_index];

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

        pc.tlas_idx = 0; // unused, tlas comes from frame tlas address
        pc.samples_buffer = rhi::buffer_device_address(ctx.device, ctx.get_buffer(out.ray_samples));
        pc.instance_data = rhi::buffer_device_address(ctx.device, ctx.render_data->rt_scene.instance_data_buffer);

        pc.cascade_count = Light::CASCADE_COUNT;
        pc.probe_state = rhi::buffer_device_address(ctx.device, ctx.get_buffer(out.probe_state));
        pc.irradiance = (u32)ctx.view(out.irradiance_img);
        pc.distance = (u32)ctx.view(out.distance_img);

        pc.vertex = (u64)ctx.render_data->vertex_view.slot;
        pc.index_bda = ctx.render_data->index_bda;

        static Random rng;
        vec3 axis = vec3(rng.Range(-1.0f, 1.0f), rng.Range(-1.0f, 1.0f), rng.Range(-1.0f, 1.0f)).normalized();
        f32 angle = rng.Range(0.0f, 2.0f * 3.14159265f);
        pc.random_rotation = vec4(axis.x, axis.y, axis.z, angle);

        {
            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);
        }

        static bool probe_state_initialized = false;
        if (!probe_state_initialized) {
            auto &ps = ctx.get_buffer(out.probe_state);
            if (ps.mapped) {
                memset(ps.mapped, 0, (u64)out.grid.probe_count * sizeof(ProbeState));
                probe_state_initialized = true;
            }
        }

        rhi::dispatch(ctx.cmd, groups, 1, 1);
    });
}

} // namespace ddgi_trace