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