ddgi_common.slang

cross platform rendering playground

src/shaders/lib/ddgi_common.slang

8.62 KB
#pragma once

#include "shared.h"
#include "octahedral.slang"

// linear probe index to 3d grid coord
uvec3 ProbeCoord(u32 probe_idx, uvec3 dims) {
    return uvec3(probe_idx % dims.x, (probe_idx / dims.x) % dims.y, probe_idx / (dims.x * dims.y));
}

u32 ProbeIndex(uvec3 coord, uvec3 dims) {
    return coord.x + coord.y * dims.x + coord.z * dims.x * dims.y;
}

// world pos of probe from linear index
vec3 ProbePosition(u32 probe_idx, vec3 origin, vec3 spacing, uvec3 dims) {
    uvec3 coord = ProbeCoord(probe_idx, dims);
    return origin + (vec3(coord) - (vec3(dims) - 1.0f) * 0.5f) * spacing;
}

// world pos of probe from grid coord
vec3 ProbePositionFromCoord(uvec3 coord, vec3 origin, vec3 spacing, uvec3 dims) {
    return origin + (vec3(coord) - (vec3(dims) - 1.0f) * 0.5f) * spacing;
}

// relocation offset, stored normalized (offset / spacing)
vec3 ProbeOffset(ProbeState *probe_state, u32 probe_idx, vec3 spacing) {
    return probe_state[probe_idx].offset_active.xyz * spacing;
}

// probe state, 0 = active, 1 = inactive
f32 LoadProbeState(ProbeState *probe_state, u32 probe_idx) {
    return probe_state[probe_idx].offset_active.w;
}

// spherical fibonacci dir for even ray spread
vec3 fibonacci_ray_direction(uint ray_slot, uint rays_per_probe) {
    const f32 PI = 3.14159265f;
    const f32 GOLDEN_ANGLE = 2.39996323f;

    f32 t = (f32(ray_slot) + 0.5f) / f32(rays_per_probe);
    f32 y = 1.0f - 2.0f * t;
    f32 r = sqrt(max(0.0f, 1.0f - y * y));
    f32 phi = GOLDEN_ANGLE * f32(ray_slot);

    return vec3(cos(phi) * r, y, sin(phi) * r);
}

// probe index to tile pixel origin in atlas
uvec2 ProbeTileOrigin(uint probe_idx, uint tiles_per_row, uint tile_stride) {
    return uvec2((probe_idx % tiles_per_row) * tile_stride, (probe_idx / tiles_per_row) * tile_stride);
}

// octahedral uv [0,1] to atlas uv for samplelevel
vec2 OctUVtoAtlasUV(vec2 oct_uv, uvec2 tile_origin, vec2 atlas_size, uint atlas_tile_size) {
    // interior texel i (local 1+i) holds oct_uv = (i + 0.5) / tile_size,
    // so the inverse is local = oct_uv * tile_size + 1 (matches DDGIGetProbeUV)
    vec2 local_pixel = oct_uv * f32(atlas_tile_size) + 1.0f;
    return (vec2(tile_origin) + local_pixel) / atlas_size;
}

// low discrepancy sphere dir for sample index, same set every time
static const f32 DDGI_PI = 3.1415926535897932f;
static const f32 DDGI_2PI = 6.2831853071795864f;
vec3 SphericalFibonacci(f32 sampleIndex, f32 numSamples) {
    const f32 b = (sqrt(5.f) * 0.5f + 0.5f) - 1.f;
    f32 phi = DDGI_2PI * frac(sampleIndex * b);
    f32 cosTheta = 1.f - (2.f * sampleIndex + 1.f) * (1.f / numSamples);
    f32 sinTheta = sqrt(saturate(1.f - (cosTheta * cosTheta)));

    return vec3((cos(phi) * sinTheta), (sin(phi) * sinTheta), cosTheta);
}

f32 chebyshev_visibility(f32 surface_distance, f32 mean_depth, f32 mean_depth_squared) {
    f32 variance = abs(mean_depth * mean_depth - mean_depth_squared);

    f32 visibility = 1.0f;
    if (surface_distance > mean_depth) {
        f32 delta = surface_distance - mean_depth;
        visibility = variance / (variance + delta * delta);
        visibility = max(pow(visibility, 3.0f), 0.0f);
    }

    return max(visibility, 0.05f);
}

vec3 eval_sampling_bias(f32 probe_max_spacing, vec3 camera_position, vec3 position, vec3 geo_normal) {
    const f32 BIAS_FACTOR = 0.25f;
    const f32 NORMAL_TO_VIEW_WEIGHT = 0.3f;
    const f32 origin_to_sample_dst = length(camera_position - position);
    const f32 sample_offset = min(probe_max_spacing * BIAS_FACTOR, origin_to_sample_dst * 0.5f);
    const vec3 sample_to_origin = normalize(camera_position - position);
    return position + lerp(sample_to_origin, geo_normal, NORMAL_TO_VIEW_WEIGHT) * sample_offset;
}

// fade weight [0,1] near grid edges
f32 ddgi_volume_blend_weight(vec3 world_position, DdgiGridParams grid) {
    vec3 probe_coords = (world_position - grid.origin.xyz) / grid.spacing.xyz + (vec3(grid.dims.xyz) - 1.0f) * 0.5f;
    vec3 over_max = vec3(grid.dims.xyz) - 1.0f - probe_coords;

    f32 w = 1.0f;
    w *= clamp(probe_coords.x, 0.0f, 1.0f) * clamp(over_max.x, 0.0f, 1.0f);
    w *= clamp(probe_coords.y, 0.0f, 1.0f) * clamp(over_max.y, 0.0f, 1.0f);
    w *= clamp(probe_coords.z, 0.0f, 1.0f) * clamp(over_max.z, 0.0f, 1.0f);
    return w;
}

vec3 sample_ddgi_irradiance(
    vec3 world_position,
    vec3 world_normal,
    DdgiGridParams grid,
    FrameUBO frame,
    Texture2D irradiance_atlas,
    Texture2D distance_atlas,
    ProbeState *probe_state,
    bool is_camera_pass
) {
    vec4 probe_origin_packed = grid.origin;
    vec3 probe_origin = probe_origin_packed.xyz;

    vec4 probe_spacing_packed = grid.spacing;
    vec3 probe_spacing = probe_spacing_packed.xyz;
    f32 max_spacing = probe_spacing_packed.w;

    vec4 probe_dims_packed = grid.dims;
    uvec3 probe_dims = uvec3(probe_dims_packed.xyz);

    vec4 cam_pos = frame.position_ws;
    world_position = eval_sampling_bias(max_spacing, cam_pos.xyz, world_position, world_normal);

    // grid coord, [0, dims-1] spans first to last probe
    vec3 grid_fraction = (world_position - probe_origin) / probe_spacing + (vec3(probe_dims) - 1.0f) * 0.5f;

    // clamp to grid range
    vec3 clamped_fraction = clamp(grid_fraction, vec3(0.0f), vec3(probe_dims) - 1.0f - 1e-6f);
    ivec3 cell_lower = ivec3(floor(clamped_fraction));
    vec3 trilinear_alpha = clamped_fraction - vec3(cell_lower);

    u32 irr_w;
    u32 irr_h;
    irradiance_atlas.GetDimensions(irr_w, irr_h);
    u32 irr_tiles_per_row = irr_w / IRRADIANCE_TILE_STRIDE;

    u32 dist_w;
    u32 dist_h;
    distance_atlas.GetDimensions(dist_w, dist_h);
    u32 dist_tiles_per_row = dist_w / DISTANCE_TILE_STRIDE;

    vec2 octahedral_uv = OctEncode(world_normal) * 0.5f + 0.5f;

    vec3 accumulated_irradiance = vec3(0.0f);
    f32 accumulated_weight = 0.0f;

    for (int corner_index = 0; corner_index < 8; corner_index++) {
        ivec3 corner_offset = ivec3(corner_index & 1, (corner_index >> 1) & 1, (corner_index >> 2) & 1);
        ivec3 corner_coord = clamp(cell_lower + corner_offset, ivec3(0), ivec3(probe_dims) - 1);
        u32 probe_index = corner_coord.z * probe_dims.y * probe_dims.x + corner_coord.y * probe_dims.x + corner_coord.x;

        f32 weight = 1.0f;

        f32 probe_state_value = LoadProbeState(probe_state, probe_index);
        if (probe_state_value >= DDGI_STATE_PROBE_INACTIVE) {
            continue;
        }

        vec3 trilinear = max(vec3(0.001f), lerp(1.0f - trilinear_alpha, trilinear_alpha, vec3(corner_offset)));
        f32 trilinear_weight = trilinear.x * trilinear.y * trilinear.z;

        uvec2 tile_origin = ProbeTileOrigin(probe_index, irr_tiles_per_row, IRRADIANCE_TILE_STRIDE);
        vec2 atlas_uv = OctUVtoAtlasUV(octahedral_uv, tile_origin, vec2(irr_w, irr_h), IRRADIANCE_SIZE);
        vec4 irradiance = irradiance_atlas.SampleLevel(frame.linear_clamp, atlas_uv, 0);

        // dir from probe to surface, atlas is keyed by probe dir
        vec3 grid_pos = probe_origin + (vec3(corner_coord) - (vec3(probe_dims) - 1.0f) * 0.5f) * probe_spacing;
        vec3 probe_off = ProbeOffset(probe_state, probe_index, probe_spacing);
        vec3 actual_probe_position = grid_pos + probe_off;
        vec3 probe_to_surface = normalize(world_position - actual_probe_position);
        f32 distance_to_surface = length(world_position - actual_probe_position);

        vec3 unbiased_dir = normalize(actual_probe_position - world_position);
        f32 wrap_shading = (dot(unbiased_dir, world_normal) + 1.0f) * 0.5f;
        weight *= (wrap_shading * wrap_shading) + 0.2f;

        // sample distance atlas
        vec2 distance_uv = OctEncode(probe_to_surface) * 0.5f + 0.5f;
        uvec2 dist_tile = ProbeTileOrigin(probe_index, dist_tiles_per_row, DISTANCE_TILE_STRIDE);
        vec2 dist_atlas_uv = OctUVtoAtlasUV(distance_uv, dist_tile, vec2(dist_w, dist_h), ATLAS_SIZE);
        vec4 distance_tex = distance_atlas.SampleLevel(frame.linear_clamp, dist_atlas_uv, 0);

        f32 mean_depth = distance_tex.r * 2.0f;
        f32 mean_depth2 = distance_tex.g * 2.0f;

        weight *= chebyshev_visibility(distance_to_surface, mean_depth, mean_depth2);

        // keep weight off zero
        weight = max(0.00001f, weight);

        const f32 crush_threshold = 0.2f;
        if (weight < crush_threshold) {
            weight *= (weight * weight) * (1.0f / (crush_threshold * crush_threshold));
        }

        weight *= trilinear_weight;
        accumulated_irradiance += weight * pow(irradiance.rgb, frame.ddgi_cfg.irradiance_encoding_gamma * 0.5f);
        accumulated_weight += weight;
    }

    // no live probes around, return black
    vec3 result = accumulated_weight > 0.0f ? accumulated_irradiance / accumulated_weight : vec3(0.0f);
    result = result * result * DDGI_2PI;

    return result;
}