cluster_lighting.slang

cross platform rendering playground

src/shaders/lib/cluster_lighting.slang

3.21 KB
#pragma once

#include "shared.h"
#include "lib/brdf.slang"

u32 flatten_cluster_index(u32 tx, u32 ty, u32 tz, u32 countX, u32 countY) {
    return tz * countX * countY + ty * countX + tx;
}

uvec3 compute_cluster_counts(u32 screen_width, u32 screen_height) {
    u32 cluster_count_x = (screen_width + TILE_SIZE_X - 1) / TILE_SIZE_X;
    u32 cluster_count_y = (screen_height + TILE_SIZE_Y - 1) / TILE_SIZE_Y;
    u32 cluster_count_z = CLUSTER_COUNT_Z;
    return uvec3(cluster_count_x, cluster_count_x, cluster_count_z);
}

// eval the depth slice index for a given view-space depth
// logarithmic subdivision: slices are uniformly distributed in log space
// TODO: link the gpu gems
u32 compute_depth_slice(f32 viewDepth, f32 nearPlane, f32 farPlane, u32 numSlices) {
    f32 slice = log(viewDepth / nearPlane) / log(farPlane / nearPlane) * f32(numSlices);
    return min(u32(slice), numSlices - 1);
}

// evals all clustered point lights affecting the given fragment
// outputs the ClusterRecord via `cr` so callers (e.g. debug views) can
// inspect lightCount without re-loading from the buffer.
vec3 evaluate_cluster_lights(
    VtxOut input,
    vec3 worldPos,
    vec3 normalWS,
    vec3 V,
    f32 NdotV,
    f32 roughness,
    f32 metallic,
    vec3 F0,
    vec4 albedo,
    // FwdDrawPushConstants pc,
    FrameUBO frame,
    ClusterRecord cr,
) {
    vec2 pixel = input.clip_pos.xy;

    ClusterGrid data = frame.cluster;
    LightCull lc_data = frame.light_cull;

    uint tileX = uint(pixel.x) / TILE_SIZE_X;
    uint tileY = uint(pixel.y) / TILE_SIZE_Y;
    f32 vDepth = clamp(abs(input.view_pos.z), frame.near_plane, frame.far_plane);
    u32 tz = compute_depth_slice(vDepth, frame.near_plane, frame.far_plane, data.cluster_count_z);
    u32 flatIdx = flatten_cluster_index(tileX, tileY, tz, data.cluster_count_x, data.cluster_count_y);

    cr.lightOffset = lc_data.clusters_records[flatIdx].lightOffset;
    cr.lightCount = lc_data.clusters_records[flatIdx].lightCount;

    vec3 point_lighting = 0.0;

    for (u32 i = 0; i < cr.lightCount; ++i) {
        u32 li = lc_data.light_index_list[cr.lightOffset + i];
        PointLightGPU pl = frame.point_lights[li];
        vec3 pl_pos = pl.position;
        f32 pl_range = pl.range;
        vec3 pl_col = pl.color;
        f32 pl_int = pl.intensity;

        vec3 L_pl = pl_pos - worldPos;
        f32 dist = length(L_pl);
        L_pl /= max(dist, 1e-6);

        f32 atten = 1.0 / (dist * dist + 1.0);
        atten = saturate(1.0 - (dist * dist) / (pl_range * pl_range));
        atten *= atten;

        vec3 H_pl = normalize(V + L_pl);

        f32 NdotL_pl = saturate(dot(normalWS, L_pl));
        f32 NdotH_pl = saturate(dot(normalWS, H_pl));
        f32 VdotH_pl = saturate(dot(V, H_pl));

        f32 D_pl = D_GGX(NdotH_pl, roughness);
        f32 G_pl = G_Smith_Direct(NdotV, NdotL_pl, roughness);
        vec3 F_pl = F_Schlick(VdotH_pl, F0);

        vec3 specular_pl = (D_pl * G_pl * F_pl) / max(4.0 * NdotV * NdotL_pl, 1e-5);

        vec3 kS_pl = F_pl;
        vec3 kD_pl = (1.0 - kS_pl) * (1.0 - metallic);

        vec3 diffuse_pl = kD_pl * albedo.xyz * (1.0 / 3.14159265);

        point_lighting += (diffuse_pl + specular_pl) * pl_col * pl_int * NdotL_pl * atten;
    }

    return point_lighting;
}