ddgi_update.slang

cross platform rendering playground

src/passes/ddgi/ddgi_update.slang

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

[[vk::push_constant]]
DdgiUpdatePushConstants pc;

#ifndef TILE_SIZE
#error "TILE_SIZE must be defined before including ddgi_update_common.slang"
#endif

static const f32 SHARPNESS = 50.0f;

groupshared vec4 gs_tile[TILE_STRIDE][TILE_STRIDE];

[shader("compute")]
[numthreads(TILE_STRIDE, TILE_STRIDE, 1)]
void csMain(
    uvec3 dtid: SV_DispatchThreadID, uvec3 gtid: SV_GroupThreadID, uvec3 gid: SV_GroupID, uint gidx: SV_GroupIndex
) {
    FrameUBO frame = *pc.frame;
    DdgiConfig cfg = frame.ddgi_cfg;
    DdgiGridParams grid = frame.ddgi_grid;
    uvec4 dims = uvec4(grid.dims);

    uint probe_idx = gid.z * dims.y * dims.x + gid.y * dims.x + gid.x;
    if (probe_idx >= dims.x * dims.y * dims.z)
        return;

    uint texel_x = gtid.x;
    uint texel_y = gtid.y;

    RWTexture2D<float4> atlas = pc.atlas_storage;
    u32 atlas_w;
    u32 atlas_h;
    atlas.GetDimensions(atlas_w, atlas_h);

    uint tiles_per_row = atlas_w / TILE_STRIDE;
    uint tile_grid_x = probe_idx % tiles_per_row;
    uint tile_grid_y = probe_idx / tiles_per_row;

    uvec2 atlas_pixel;
    atlas_pixel.x = tile_grid_x * TILE_STRIDE + texel_x;
    atlas_pixel.y = tile_grid_y * TILE_STRIDE + texel_y;

    f32 state = pc.probe_state[probe_idx].offset_active.w;

    f32 probeMaxDistance = length(grid.spacing.w) * 1.5f;

    vec4 prev = pc.atlas_sample.Load(ivec3(atlas_pixel, 0));
    gs_tile[texel_x][texel_y] = prev;

    // dead probes keep old data, fading ones still blend so they stay fresh
    if (state >= DDGI_STATE_PROBE_INACTIVE) {
        atlas[atlas_pixel] = prev;
        return;
    }

    // step 1: load, prev frame value already in groupshared for border mirror
    GroupMemoryBarrierWithGroupSync();

    // step 2: compute, interior texels only
    bool is_interior = texel_x >= 1 && texel_x <= TILE_SIZE && texel_y >= 1 && texel_y <= TILE_SIZE;

    if (is_interior) {
        vec2 uv = (vec2(texel_x - 1, texel_y - 1) + 0.5) / f32(TILE_SIZE);
        vec3 texel_dir = OctDecode(uv * 2.0 - 1.0);

        vec4 accumulated_result = 0.0;
        f32 accumulated_weight = 0.0;

        uint startRay = RAYS_PER_PROBE > FIXED_RAY_COUNT ? FIXED_RAY_COUNT : 0;
        uint backfaces = 0;
        bool skip_texel = false;

        for (uint i = startRay; i < RAYS_PER_PROBE; i++) {
            uint base = probe_idx * RAYS_PER_PROBE;
            vec4 rad_dist = pc.samples_buffer[base + i].radiance_distance;
            vec4 ray_dir = pc.samples_buffer[base + i].direction;

            vec3 ray_radiance = rad_dist.xyz;
            f32 ray_distance = rad_dist.w;
            bool is_backface = (rad_dist.w < 0.0f);

            f32 w = dot(texel_dir, ray_dir.xyz);
            if (w <= 0.0f)
                continue;

#if defined(MODE_IRRADIANCE)
            if (is_backface) {
                backfaces++;
                if (backfaces > uint((RAYS_PER_PROBE - FIXED_RAY_COUNT) * cfg.random_ray_backface_threshold)) {
                    // too many backfaces, skip texel but keep threads converged
                    skip_texel = true;
                    break;
                }
                continue;
            }
            accumulated_result.rgb += ray_radiance * w;
#elif defined(MODE_DISTANCE)
            if (is_backface) {
                ray_distance = abs(ray_distance);
            }

            w += pow(w, SHARPNESS);

            ray_distance = min(ray_distance, probeMaxDistance);

            accumulated_result.x += abs(ray_distance) * w;
            accumulated_result.y += abs(ray_distance) * abs(ray_distance) * w;
#endif
            accumulated_weight += w;
        }

        f32 hw = pc.frame_counter > 0 ? (1.0 - cfg.history_alpha) : 1.0;

#if defined(MODE_IRRADIANCE)
        vec3 irr = accumulated_weight > 0.0 ? accumulated_result.rgb / (2.0 * accumulated_weight) : 0.0;
        irr = pow(irr, 1.0f / cfg.irradiance_encoding_gamma);
        irr = lerp(prev.rgb, irr, hw);
        gs_tile[texel_x][texel_y] = skip_texel ? prev : vec4(irr, 1.0);
#elif defined(MODE_DISTANCE)
        f32 mean_dist = accumulated_weight > 0.0 ? accumulated_result.x / (2.0 * accumulated_weight) : 0.0;
        f32 mean_dist2 = accumulated_weight > 0.0 ? accumulated_result.y / (2.0 * accumulated_weight) : 0.0;
        mean_dist = lerp(prev.x, mean_dist, hw);
        mean_dist2 = lerp(prev.y, mean_dist2, hw);
        gs_tile[texel_x][texel_y] = vec4(mean_dist, mean_dist2, 1.0, 1.0);
#endif
    }

    GroupMemoryBarrierWithGroupSync();

    // step 3: mirror, border texels copy interior
    if (!is_interior) {
        uint src_x = texel_x;
        uint src_y = texel_y;

        if (texel_x == 0)
            src_x = TILE_SIZE;
        else if (texel_x == TILE_STRIDE - 1)
            src_x = 1;

        if (texel_y == 0)
            src_y = TILE_SIZE;
        else if (texel_y == TILE_STRIDE - 1)
            src_y = 1;

        src_x = TILE_SIZE + 1 - src_x;
        src_y = TILE_SIZE + 1 - src_y;

        gs_tile[texel_x][texel_y] = gs_tile[src_x][src_y];
    }

    GroupMemoryBarrierWithGroupSync();

    // step 4: store, all threads write atlas
    atlas[atlas_pixel] = gs_tile[texel_x][texel_y];
    return;
}