use windows::Win32::Graphics::Direct3D12::*;
use super::allocator::{DeviceAllocator, PooledBuffer};
use super::com;
use crate::directx::context::DxContext;
use crate::directx::pipeline::serialize_desc_and_create;
use crate::directx::slang_builtins;
use crate::directx::slang_builtins::SlangCompile;
use crate::directx::texture::{create_buffer, create_uav_buffer};
use crate::gfx::render_types::{CLUSTER_COUNT, CLUSTER_LIGHT_LIST_STRIDE, ClusterParams};
const CLUSTER_PARAMS_SLOT_STRIDE: u64 = 256;
const CLUSTER_SLOT_CLUSTERED: u64 = 0;
const CLUSTER_SLOT_UNCLUSTERED: u64 = 1;
pub(in crate::directx) struct LightCullState {
pub root_sig: Option<ID3D12RootSignature>,
pub pso: Option<ID3D12PipelineState>,
pub cluster_buffer: ID3D12Resource,
pub params_resources: Vec<PooledBuffer>,
pub params_ptrs: Vec<*mut u8>,
}
impl LightCullState {
pub(in crate::directx) fn unmap(&self) {
for res in &self.params_resources {
unsafe { res.Unmap(0, None) };
}
}
}
pub(in crate::directx) fn compile_light_cull_shader(hot_reload: bool) -> Result<Vec<u8>, String> {
slang_builtins::LIGHT_CULL.compile(hot_reload)
}
pub(in crate::directx) fn create_light_cull_root_signature(
device: &ID3D12Device,
) -> Result<ID3D12RootSignature, String> {
let params = [
D3D12_ROOT_PARAMETER {
ParameterType: D3D12_ROOT_PARAMETER_TYPE_CBV,
Anonymous: D3D12_ROOT_PARAMETER_0 {
Descriptor: D3D12_ROOT_DESCRIPTOR {
ShaderRegister: 0,
RegisterSpace: 0,
},
},
ShaderVisibility: D3D12_SHADER_VISIBILITY_ALL,
},
D3D12_ROOT_PARAMETER {
ParameterType: D3D12_ROOT_PARAMETER_TYPE_SRV,
Anonymous: D3D12_ROOT_PARAMETER_0 {
Descriptor: D3D12_ROOT_DESCRIPTOR {
ShaderRegister: 0,
RegisterSpace: 0,
},
},
ShaderVisibility: D3D12_SHADER_VISIBILITY_ALL,
},
D3D12_ROOT_PARAMETER {
ParameterType: D3D12_ROOT_PARAMETER_TYPE_UAV,
Anonymous: D3D12_ROOT_PARAMETER_0 {
Descriptor: D3D12_ROOT_DESCRIPTOR {
ShaderRegister: 0,
RegisterSpace: 0,
},
},
ShaderVisibility: D3D12_SHADER_VISIBILITY_ALL,
},
];
let desc = D3D12_ROOT_SIGNATURE_DESC {
NumParameters: params.len() as u32,
pParameters: params.as_ptr(),
Flags: D3D12_ROOT_SIGNATURE_FLAG_NONE,
..Default::default()
};
serialize_desc_and_create(device, &desc, "light cull root sig")
}
pub(in crate::directx) fn create_light_cull_pso(
device: &ID3D12Device,
root_sig: &ID3D12RootSignature,
cs: &[u8],
) -> Result<ID3D12PipelineState, String> {
let desc = D3D12_COMPUTE_PIPELINE_STATE_DESC {
pRootSignature: com::borrowed(root_sig),
CS: D3D12_SHADER_BYTECODE {
pShaderBytecode: cs.as_ptr() as _,
BytecodeLength: cs.len(),
},
..Default::default()
};
unsafe { crate::directx::pso_library::create_compute(device, &desc) }
.map_err(|e| format!("create light cull PSO: {e}"))
}
pub(in crate::directx) fn build_cluster_light_buffer(
device: &ID3D12Device,
) -> Result<ID3D12Resource, String> {
let len =
(CLUSTER_COUNT * CLUSTER_LIGHT_LIST_STRIDE) as u64 * std::mem::size_of::<u32>() as u64;
create_uav_buffer(device, len, D3D12_RESOURCE_STATE_COMMON)
}
pub(in crate::directx) fn build_cluster_params_buffers(
alloc: &DeviceAllocator,
frames: usize,
) -> Result<(Vec<PooledBuffer>, Vec<*mut u8>), String> {
let size = CLUSTER_PARAMS_SLOT_STRIDE * 2;
let mut resources = Vec::with_capacity(frames);
let mut ptrs = Vec::with_capacity(frames);
for _ in 0..frames {
let res = create_buffer(
alloc,
size,
D3D12_HEAP_TYPE_UPLOAD,
D3D12_RESOURCE_STATE_GENERIC_READ,
)?;
let mut ptr = std::ptr::null_mut();
unsafe { res.Map(0, None, Some(&mut ptr)) }
.map_err(|e| format!("map cluster params buffer: {e}"))?;
let ptr = ptr as *mut u8;
let unclustered = ClusterParams::ZERO;
unsafe {
std::ptr::copy_nonoverlapping(
&unclustered as *const ClusterParams as *const u8,
ptr.add((CLUSTER_SLOT_UNCLUSTERED * CLUSTER_PARAMS_SLOT_STRIDE) as usize),
std::mem::size_of::<ClusterParams>(),
);
}
resources.push(res);
ptrs.push(ptr);
}
Ok((resources, ptrs))
}
impl DxContext {
pub(in crate::directx) fn cluster_params_gva(&self, frame_idx: usize, clustered: bool) -> u64 {
let slot = if clustered {
CLUSTER_SLOT_CLUSTERED
} else {
CLUSTER_SLOT_UNCLUSTERED
};
let base = com::gpu_va(&self.light_cull.params_resources[frame_idx]);
base + slot * CLUSTER_PARAMS_SLOT_STRIDE
}
pub(in crate::directx) fn cluster_list_gva(&self) -> u64 {
com::gpu_va(&self.light_cull.cluster_buffer)
}
pub(in crate::directx) fn write_cluster_params(
&self,
frame_idx: usize,
params: &ClusterParams,
) {
let dst = self.light_cull.params_ptrs[frame_idx];
unsafe {
std::ptr::copy_nonoverlapping(
params as *const ClusterParams as *const u8,
dst.add((CLUSTER_SLOT_CLUSTERED * CLUSTER_PARAMS_SLOT_STRIDE) as usize),
std::mem::size_of::<ClusterParams>(),
);
}
}
pub(in crate::directx) fn encode_light_cull(
&self,
cmd: &ID3D12GraphicsCommandList,
frame_idx: usize,
) -> Result<(), String> {
let (pso, root_sig) = match (&self.light_cull.pso, &self.light_cull.root_sig) {
(Some(p), Some(r)) => (p, r),
_ => return Ok(()),
};
let cluster_buffer = &self.light_cull.cluster_buffer;
let params_gva = self.cluster_params_gva(frame_idx, true);
let lights_gva = com::gpu_va(&self.uniforms.local_light_buffer);
unsafe {
cmd.SetComputeRootSignature(root_sig);
cmd.SetPipelineState(pso);
cmd.SetComputeRootConstantBufferView(0, params_gva);
cmd.SetComputeRootShaderResourceView(1, lights_gva);
cmd.SetComputeRootUnorderedAccessView(2, com::gpu_va(cluster_buffer));
cmd.Dispatch(CLUSTER_COUNT.div_ceil(64), 1, 1);
}
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn kernel_cluster_constants_match_render_types() {
let kernel = concinnity_core::render::shaders::LIGHT_CULL;
assert!(kernel.contains(&format!(
"CLUSTER_LIGHT_LIST_STRIDE = {CLUSTER_LIGHT_LIST_STRIDE}u"
)));
let max_per_cluster = crate::gfx::render_types::MAX_LIGHTS_PER_CLUSTER;
assert!(kernel.contains(&format!("MAX_LIGHTS_PER_CLUSTER = {max_per_cluster}u")));
}
}