ssr.slang

cross platform rendering playground

src/passes/ssr/ssr.slang

2.53 KB
#include "shared.h"
#include "renderer_types.h"

[[vk::push_constant]]
SsrPushConstants pc;

[shader("compute")]
[numthreads(8, 8, 1)]
void csMain(uvec3 dtid: SV_DispatchThreadID) {
    if (dtid.x >= pc.width || dtid.y >= pc.height)
        return;
    int2 pix = int2(dtid.xy);

    f32 depth = pc.depth[pix].r;
    if (depth <= 0.0001f) {
        pc.output[pix] = vec4(0.0f, 0.0f, 0.0f, 0.0f);
        return;
    }

    vec3 N = pc.normal[pix].xyz;
    if (dot(N, N) < 1e-6f) {
        pc.output[pix] = vec4(0.0f, 0.0f, 0.0f, 0.0f);
        return;
    }
    N = normalize(N);

    FrameUBO frame = *pc.frame;

    vec2 uv = (vec2(dtid.xy) + 0.5f) / vec2(f32(pc.width), f32(pc.height));
    vec4 ndc = vec4(uv * 2.0f - 1.0f, depth, 1.0f);
    vec4 world4 = ndc * frame.inv_view_proj;
    vec3 world = world4.xyz / world4.w;

    vec3 V = normalize(frame.position_ws.xyz - world);
    vec3 R = reflect(-V, N);

    f32 NdotV = max(dot(N, V), 0.0f);
    // self-intersection, grazing
    f32 t = max(0.1f, 0.1f / max(dot(V, N), 0.1f));

    f32 t_hit = -1.0f;
    vec2 hit_uv = vec2(-1.0f);

    [loop]
    for (int i = 0; i < 128; i++) {
        if (i >= pc.max_steps)
            break;
        if (t > pc.max_dist)
            break;

        vec3 pos = world + R * t;
        vec4 clip = vec4(pos, 1.0f) * frame.view_proj;
        if (clip.w <= 0.0f) {
            t += pc.stride;
            continue;
        }

        vec2 suv = clip.xy / clip.w * 0.5f + 0.5f;
        if (any(suv < 0.0f) || any(suv > 1.0f)) {
            t += pc.stride;
            continue;
        }

        f32 ray_d = clip.z / clip.w;
        f32 scene_d = pc.depth.SampleLevel(frame.linear_clamp, suv, 0.0f).r;
        // rev-Z: closer surfaces hold LARGER depth, so a ray passing
        // behind a surface sees scene_d slightly above ray_d.
        f32 diff = scene_d - ray_d;
        if (diff >= 0.0f && diff < pc.thickness) {
            t_hit = t;
            hit_uv = suv;
            break;
        }
        t += pc.stride;
    }

    vec3 reflection = vec3(0.0f);
    if (t_hit > 0.0f) {
        vec3 sample = pc.color.SampleLevel(frame.linear_clamp, hit_uv, 0.0f).rgb;
        vec3 F0 = vec3(0.04f);
        vec3 fresnel = F0 + (vec3(1.0f) - F0) * pow(1.0f - NdotV, 5.0f);
        vec2 edge_dist = min(hit_uv, 1.0f - hit_uv);
        f32 edge_fade = saturate(min(edge_dist.x, edge_dist.y) * 8.0f);
        f32 dist_fade = 1.0f - t_hit / pc.max_dist;
        dist_fade *= dist_fade;
        reflection = sample * fresnel * edge_fade * dist_fade;
    }

    pc.output[pix] = vec4(reflection, 1.0f);
}