#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);
}