ddgi_probe.slang

cross platform rendering playground

src/passes/ddgi/ddgi_probe.slang

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

[[vk::push_constant]]
ProbeSpherePushConstants pc;

struct vtxout {
    vec4 position : SV_Position;
    vec3 world_pos : TEXCOORD0;
    uint probe_index : TEXCOORD1;
};

[shader("vertex")]
vtxout vsMain(uint vid: SV_VertexID, uint iid: SV_InstanceID) {
    vec4 pos_packed = pc.sphere_mesh[vid];

    FrameUBO frame = *pc.frame;
    DdgiGridParams grid = frame.ddgi_grid;

    uvec3 dims = uvec3(grid.dims.xyz);
    vec3 probe_pos = ProbePosition(iid, grid.origin.xyz, grid.spacing.xyz, dims) +
                     ProbeOffset(pc.probe_state, iid, grid.spacing.xyz);

    f32 scale = 0.1;
    vec3 world_pos = pos_packed.xyz * scale + probe_pos;

    vtxout o;
    o.position = vec4(world_pos, 1.0) * frame.view_proj;
    o.world_pos = world_pos;
    o.probe_index = iid;
    return o;
}

[shader("fragment")]
vec4 fsMain(vtxout in) : SV_Target {
    FrameUBO frame = *pc.frame;
    DdgiGridParams grid = frame.ddgi_grid;

    if (pc.debug_flag == DebugViewMode::PROBES_CLASSIFICATION) {
        // stored state, 0 active 1 inactive
        f32 state = LoadProbeState(pc.probe_state, in.probe_index);
        return state == DDGI_STATE_PROBE_INACTIVE ? vec4(1.0, 0.0, 0.0, 1.0) : vec4(0.0, 1.0, 0.0, 1.0);
    }

    uvec3 dims = uvec3(grid.dims.xyz);
    vec3 probe_pos = ProbePosition(in.probe_index, grid.origin.xyz, grid.spacing.xyz, dims) +
                     ProbeOffset(pc.probe_state, in.probe_index, grid.spacing.xyz);
    vec3 frag_dir = normalize(in.world_pos - probe_pos);

    uint base = in.probe_index * pc.rays_per_probe;
    uint closest = 0;
    f32 best_dot = -1.0;

    for (uint i = 0; i < pc.rays_per_probe; i++) {
        vec3 ray_dir = pc.samples_buffer[base + i].direction.xyz;
        f32 d = dot(frag_dir, ray_dir);
        if (d > best_dot) {
            best_dot = d;
            closest = i;
        }
    }

    vec4 sample = pc.samples_buffer[base + closest].radiance_distance;

    return vec4(sample.rgb, 1.0);
}

[shader("fragment")]
vec4 fsAtlasMain(vtxout in) : SV_Target {
    FrameUBO frame = *pc.frame;
    DdgiGridParams grid = frame.ddgi_grid;

    uvec3 dims = uvec3(grid.dims.xyz);
    vec3 probe_pos = ProbePosition(in.probe_index, grid.origin.xyz, grid.spacing.xyz, dims) +
                     ProbeOffset(pc.probe_state, in.probe_index, grid.spacing.xyz);

    vec3 dir = normalize(in.world_pos - probe_pos);
    vec2 uv = OctEncode(dir) * 0.5 + 0.5;

    u32 w, h;
    pc.atlas_sampled.GetDimensions(w, h);
    u32 tpr = w / IRRADIANCE_TILE_STRIDE;

    uvec2 tile_origin = ProbeTileOrigin(in.probe_index, tpr, IRRADIANCE_TILE_STRIDE);
    vec2 atlas_uv = OctUVtoAtlasUV(uv, tile_origin, vec2(w, h), IRRADIANCE_SIZE);

    vec4 irradiance = pc.atlas_sampled.SampleLevel(frame.linear_clamp, atlas_uv, 0);
    return vec4(irradiance.rgb, 1.0);
}

[shader("fragment")]
vec4 fsDistMain(vtxout in) : SV_Target {
    FrameUBO frame = *pc.frame;
    DdgiGridParams grid = frame.ddgi_grid;

    uvec3 dims = uvec3(grid.dims.xyz);
    vec3 probe_pos = ProbePosition(in.probe_index, grid.origin.xyz, grid.spacing.xyz, dims) +
                     ProbeOffset(pc.probe_state, in.probe_index, grid.spacing.xyz);
    vec3 dir = normalize(in.world_pos - probe_pos);
    vec2 uv = OctEncode(dir) * 0.5 + 0.5;

    u32 w, h;
    pc.atlas_sampled.GetDimensions(w, h);
    u32 tpr = w / DISTANCE_TILE_STRIDE;

    uvec2 tile_origin = ProbeTileOrigin(in.probe_index, tpr, DISTANCE_TILE_STRIDE);
    vec2 atlas_uv = OctUVtoAtlasUV(uv, tile_origin, vec2(w, h), ATLAS_SIZE);

    // dist2 lives in green
    f32 dist = pc.atlas_sampled.SampleLevel(frame.linear_clamp, atlas_uv, 0).r;

    f32 t = clamp(dist / 2.0, 0.0, 1.0);
    vec3 color = t < 0.5 ? vec3(0.0, t * 2.0, 1.0 - t * 2.0) : vec3((t - 0.5) * 2.0, 1.0 - (t - 0.5) * 2.0, 0.0);
    return vec4(color, 1.0);
}