cluster_bounds.slang

cross platform rendering playground

src/passes/cluster/cluster_bounds.slang

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

[[vk::push_constant]]
ClusterBoundsPushConstants pc;

f32 slice_far_depth(u32 slice, f32 near, f32 far, u32 num_slices) {
    f32 t = (f32(slice) + 1.0) / f32(num_slices);
    return near * pow(far / near, t);
}

vec3 make_view_ray(f32 sx, f32 sy, mat4 invProj, f32 screenW, f32 screenH) {
    f32 ndcX = sx / screenW * 2.0 - 1.0;
    f32 ndcY = 1.0 - sy / screenH * 2.0;
    // f32 ndcY = sy / screenH * 2.0 - 1.0;

    vec4 p = vec4(ndcX, ndcY, 1.0, 1.0) * invProj;

    vec3 ray = p.xyz / p.w;

    // Make z = 1
    return ray / ray.z;
}

[shader("compute")]
[numthreads(64, 1, 1)]
void csMain(uvec3 dispatchThreadID: SV_DispatchThreadID) {
    u32 clusterIndex = dispatchThreadID.x;
    FrameUBO frame = *pc.frame;
    ClusterGrid data = frame.cluster;

    u32 totalClusters = data.cluster_count_x * data.cluster_count_y * data.cluster_count_z;
    if (clusterIndex >= totalClusters) {
        return;
    }

    // create a flat index to (tx, ty, tz)
    u32 tz = clusterIndex / (data.cluster_count_x * data.cluster_count_y);
    u32 xy = clusterIndex % (data.cluster_count_x * data.cluster_count_y);
    u32 ty = xy / data.cluster_count_x;
    u32 tx = xy % data.cluster_count_x;

    // screen-space tile bounds
    f32 left = f32(tx) * f32(TILE_SIZE_X);
    f32 top = f32(ty) * f32(TILE_SIZE_Y);
    f32 right = f32(tx + 1) * f32(TILE_SIZE_X);
    f32 bottom = f32(ty + 1) * f32(TILE_SIZE_Y);

    // depth bounds for this slice (view-space Z, positive away from camera)
    f32 nearZ =
        tz == 0 ? frame.near_plane : slice_far_depth(tz - 1, frame.near_plane, frame.far_plane, data.cluster_count_z);
    f32 farZ = slice_far_depth(tz, frame.near_plane, frame.far_plane, data.cluster_count_z);

    // sample 4 tile corners at both near and far depth → 8 points
    vec3 corners[8];
    mat4 invProj = frame.inv_proj;

    vec3 rayTL = make_view_ray(left, top, invProj, f32(frame.screen_w), f32(frame.screen_h));
    vec3 rayTR = make_view_ray(right, top, invProj, f32(frame.screen_w), f32(frame.screen_h));
    vec3 rayBL = make_view_ray(left, bottom, invProj, f32(frame.screen_w), f32(frame.screen_h));
    vec3 rayBR = make_view_ray(right, bottom, invProj, f32(frame.screen_w), f32(frame.screen_h));

    corners[0] = rayTL * nearZ;
    corners[1] = rayTL * farZ;

    corners[2] = rayTR * nearZ;
    corners[3] = rayTR * farZ;

    corners[4] = rayBL * nearZ;
    corners[5] = rayBL * farZ;

    corners[6] = rayBR * nearZ;
    corners[7] = rayBR * farZ;

    // view space bounding box
    vec3 minP = corners[0];
    vec3 maxP = corners[0];
    [unroll]
    for (u32 i = 1; i < 8; ++i) {
        minP = min(minP, corners[i]);
        maxP = max(maxP, corners[i]);
    }

    data.clusters_bounds[clusterIndex].minPoint = vec4(minP, 0.0);
    data.clusters_bounds[clusterIndex].maxPoint = vec4(maxP, 0.0);
}