light_cull.slang

cross platform rendering playground

src/passes/cluster/light_cull.slang

2.45 KB
#include "shared.h"
#include "renderer_types.h"
#include "lib/frustum_culling.slang"

[[vk::push_constant]]
LightCullPushConstants pc;

// single-phase light culling: each thread counts, allocates, and fills
// for its cluster in one go

[shader("compute")]
[numthreads(64, 1, 1)]
void csMain(uvec3 dispatchThreadID: SV_DispatchThreadID) {
    u32 clusterIndex = dispatchThreadID.x;

    FrameUBO frame = *pc.frame;
    ClusterGrid data = frame.cluster;
    LightCull lc_data = frame.light_cull;

    if (clusterIndex >= data.cluster_total) {
        return;
    }

    // aabb of each cluster (cluster_bounds.sla)
    ClusterBounds bounds = data.clusters_bounds[clusterIndex];
    vec4 bmin = bounds.minPoint;
    vec4 bmax = bounds.maxPoint;

    mat4 view = frame.view;
    mat4 view_proj = frame.view_proj;
    Frustum frustum = extract_frustum(view_proj);
    u32 pl_count = frame.point_light_count;

    // lights intersection w the cluster
    u32 count = 0;
    for (u32 i = 0; i < pl_count; ++i) {
        PointLightGPU pl = frame.point_lights[i];
        vec3 posWS = pl.position;
        f32 range = pl.range;

        if (!sphere_visible_in_frustum(posWS, range, frustum))
            continue;

        vec4 posVS4 = vec4(posWS, 1.0) * view;
        vec3 posVS = posVS4.xyz / posVS4.w;

        if (sphere_intersects_aabb(posVS, range, bmin.xyz, bmax.xyz)) {
            count++;
        }
    }

    if (count == 0) {
        lc_data.clusters_records[clusterIndex].lightOffset = 0;
        lc_data.clusters_records[clusterIndex].lightCount = 0;
        return;
    }

    // alloc slot range from global atomic counter
    u32 offset;
    InterlockedAdd(pc.light_index_counter[0], count, offset);

    u32 writeCount = (offset < MAX_LIGHT_INDICES) ? min(count, MAX_LIGHT_INDICES - offset) : 0;

    lc_data.clusters_records[clusterIndex].lightOffset = offset;
    lc_data.clusters_records[clusterIndex].lightCount = writeCount;

    u32 localIdx = 0;
    for (u32 i = 0; i < pl_count && localIdx < writeCount; ++i) {
        PointLightGPU pl = frame.point_lights[i];
        vec3 posWS = pl.position;
        f32 range = pl.range;

        if (!sphere_visible_in_frustum(posWS, range, frustum)) {
            continue;
        }

        vec4 posVS4 = vec4(posWS, 1.0) * view;
        vec3 posVS = posVS4.xyz / posVS4.w;

        if (sphere_intersects_aabb(posVS, range, bmin.xyz, bmax.xyz)) {
            lc_data.light_index_list[offset + localIdx] = i;
            localIdx++;
        }
    }
}