#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