ddgi_trace.slang

cross platform rendering playground

src/passes/ddgi/ddgi_trace.slang

6.22 KB
#include "shared.h"
#include "renderer_types.h"
#include "lib/ddgi_common.slang"
#include "lib/octahedral.slang"
#include "lib/shadow_sample.slang"

[vk::push_constant]
DdgiTracePushConstants pc;

static const f32 DDGI_BOUNCE_CONTRIBUTION_MULTIPLIER = 0.9f;
static const f32 DDGI_MISS_DISTANCE = 1e27f;

static const f32 PI = 3.1415926535f;
static const f32 M_PI_INV = 0.318309886f;

[shader("compute")]
[numthreads(64, 1, 1)]
void csMain(uvec3 dtid: SV_DispatchThreadID) {
    uint ray_index = dtid.x;
    FrameUBO frame = *pc.frame;
    DdgiGridParams grid = frame.ddgi_grid;
    uvec3 dims = uvec3(grid.dims.xyz);
    uint total_rays = dims.x * dims.y * dims.z * RAYS_PER_PROBE;
    if (ray_index >= total_rays)
        return;

    uint probe_idx = ray_index / RAYS_PER_PROBE;
    uint ray_slot = ray_index % RAYS_PER_PROBE;
    // probe world position includes relocation offset
    vec3 probe_pos = ProbePosition(probe_idx, grid.origin.xyz, grid.spacing.xyz, dims) +
                     ProbeOffset(pc.probe_state, probe_idx, grid.spacing.xyz);

    vec3 dir;
    if (ray_slot < FIXED_RAY_COUNT) {
        // fixed rays, same dirs every frame for stable relocate/classify
        dir = SphericalFibonacci(ray_slot, FIXED_RAY_COUNT);
    } else {
        // stochastic rays, rotated each frame
        uint sampleIndex = ray_slot - FIXED_RAY_COUNT;
        uint numSamples = RAYS_PER_PROBE - FIXED_RAY_COUNT;
        vec3 raw_dir = SphericalFibonacci(sampleIndex, numSamples);
        vec3 axis = normalize(pc.random_rotation.xyz);
        f32 angle = pc.random_rotation.w;
        f32 s, c;
        sincos(angle, s, c);
        dir = raw_dir * c + cross(axis, raw_dir) * s + axis * dot(axis, raw_dir) * (1.0f - c);
        dir = normalize(dir);
    }

    // skip stochastic rays on dead probes, fixed rays still trace for classify.
    // zero the slot so readers never see stale data
    f32 state = LoadProbeState(pc.probe_state, probe_idx);
    if (state >= DDGI_STATE_PROBE_INACTIVE && ray_slot >= FIXED_RAY_COUNT) {
        pc.samples_buffer[ray_index].radiance_distance = vec4(0.0f);
        pc.samples_buffer[ray_index].direction = vec4(dir, 0.0f);
        return;
    }

    RayQuery<0> q;
    RayDesc ray;
    ray.Origin = probe_pos;
    ray.Direction = dir;
    ray.TMin = 0.0f;
    f32 grid_extent = max(max(grid.dims.x, grid.dims.y), grid.dims.z) * grid.spacing.w;
    ray.TMax = grid_extent * 2.0;
    q.TraceRayInline(RaytracingAccelerationStructure(frame.tlas_address), 0, 0xFF, ray);
    q.Proceed();

    ProbeRaySample sample;
    sample.direction = vec4(dir, 0.0);
    sample.radiance_distance = vec4(0.0);

    if (q.CommittedStatus() == COMMITTED_TRIANGLE_HIT) {
        f32 dist = q.CommittedRayT();

        if (!q.CommittedTriangleFrontFace()) {
            // backface, probe is inside geometry. no light here, flag with negative dist
            sample.radiance_distance = vec4(vec3(0.0f, 0.0f, 0.0f), -dist * 0.2f);
        } else {
            // fixed rays store dist only, no lighting, must not bias irradiance
            if (ray_slot < FIXED_RAY_COUNT) {
                sample.radiance_distance = vec4(vec3(0.0f, 0.0f, 0.0f), dist);
            } else {
                u32 inst_idx = q.CommittedInstanceID();
                u32 item_idx = pc.instance_data[inst_idx].render_item_idx;
                u32 prim_idx = q.CommittedPrimitiveIndex();
                vec2 bc = q.CommittedTriangleBarycentrics();

                f32 bary_u = 1.0 - bc.x - bc.y;
                f32 bary_v = bc.x;
                f32 bary_w = bc.y;

                RenderItemGPU item = frame.render_items[item_idx];
                SubmeshGPU submesh = frame.submeshes[item.submesh_index];

                u32 i0 = pc.index_bda[submesh.first_index + prim_idx * 3 + 0];
                u32 i1 = pc.index_bda[submesh.first_index + prim_idx * 3 + 1];
                u32 i2 = pc.index_bda[submesh.first_index + prim_idx * 3 + 2];

                StructuredBuffer<Vertex> verts = pc.vertex;
                Vertex v0 = verts[i0 + submesh.base_vertex];
                Vertex v1 = verts[i1 + submesh.base_vertex];
                Vertex v2 = verts[i2 + submesh.base_vertex];

                vec2 uv = v0.uv * bary_u + v1.uv * bary_v + v2.uv * bary_w;

                MaterialGPU mat = frame.materials[item.material_index];
                vec4 albedo = mat.albedo.SampleLevel(frame.aniso_wrap_mips, uv, 0) * mat.base_color;

                vec3 objN = normalize(v0.normal * bary_u + v1.normal * bary_v + v2.normal * bary_w);
                TransformGPU xform = frame.transforms[item.instance_id];
                vec3 N = normalize(objN * (mat3)transpose(xform.normal_world));

                vec3 hit_ws = ray.Origin + ray.Direction * dist;

                vec3 L = normalize(frame.light_dir);
                f32 NdotL = saturate(dot(N, L));
                f32 intensity = f32(frame.light_intensity);
                vec3 light_color = frame.light_color;

                f32 vis = 1.0;
                {
                    RayDesc shadow_ray;
                    shadow_ray.Origin = hit_ws + N * 0.01f;
                    shadow_ray.Direction = L;
                    shadow_ray.TMin = 0.0f;
                    shadow_ray.TMax = 1e6f;
                    RayQuery<0> sq;
                    sq.TraceRayInline(
                        RaytracingAccelerationStructure(frame.tlas_address),
                        RAY_FLAG_ACCEPT_FIRST_HIT_AND_END_SEARCH | RAY_FLAG_SKIP_CLOSEST_HIT_SHADER,
                        0xFF,
                        shadow_ray
                    );
                    sq.Proceed();
                    if (sq.CommittedStatus() == COMMITTED_TRIANGLE_HIT)
                        vis = 0.0f;
                }
                vec3 direct = light_color * intensity * NdotL * vis;
                vec3 indirect =
                    sample_ddgi_irradiance(hit_ws, N, grid, frame, pc.irradiance, pc.distance, pc.probe_state, false);
                vec3 radiance = (direct + indirect) * (min(albedo.rgb, 0.9) / PI);
                sample.radiance_distance = vec4(radiance, dist);
            }
        }
    } else {
        vec3 sky = frame.env.SampleLevel(frame.linear_clamp, dir, 0.0f).rgb;
        sample.radiance_distance = vec4(sky, DDGI_MISS_DISTANCE);
    }

    pc.samples_buffer[ray_index] = sample;
}