use concinnity_core::render::error::{RenderError, RenderResult};
use concinnity_core::render::reflection_probe::PrefilterPlan;
use concinnity_core::render::uniforms::ProbePrefilterParams;
use windows::Win32::Graphics::Direct3D12::*;
use windows::Win32::Graphics::Dxgi::Common::*;
use super::builtin_shaders::CompileProgram;
use super::com;
use super::context::DxContext;
use super::probe_set::{CubeResource, create_cube_resource, write_cube_mip_uav};
use super::pso::compute_pso;
use crate::directx::descriptor_slot::DescriptorTables;
use crate::directx::descriptor_slot::SrvSlot;
use crate::directx::root_constants::RootConstants;
use crate::directx::root_sig::{RootSig, SamplerState, Visibility};
pub(in crate::directx) const PROBE_CUBE_FORMAT: DXGI_FORMAT = DXGI_FORMAT_R16G16B16A16_FLOAT;
pub(in crate::directx) const PROBE_MAX_MIPS: usize = 12;
const PREFILTER_TILE: u32 = 8;
pub(in crate::directx) struct ProbePrefilterPipelines {
mip_root: ID3D12RootSignature,
ggx_root: ID3D12RootSignature,
mip0: ID3D12PipelineState,
downsample: ID3D12PipelineState,
ggx: ID3D12PipelineState,
}
pub(in crate::directx) fn typed_uav_load_supported(device: &ID3D12Device) -> bool {
let mut options = D3D12_FEATURE_DATA_D3D12_OPTIONS::default();
let ok = unsafe {
device.CheckFeatureSupport(
D3D12_FEATURE_D3D12_OPTIONS,
&mut options as *mut _ as *mut std::ffi::c_void,
size_of::<D3D12_FEATURE_DATA_D3D12_OPTIONS>() as u32,
)
};
ok.is_ok() && options.TypedUAVLoadAdditionalFormats.as_bool()
}
impl ProbePrefilterPipelines {
pub(in crate::directx) fn new(device: &ID3D12Device, hot_reload: bool) -> RenderResult<Self> {
use super::builtin_shaders;
let mip_root = create_mip_root_signature(device)?;
let ggx_root = create_ggx_root_signature(device)?;
let mip0 = compute_pso(
device,
&mip_root,
&builtin_shaders::PROBE_MIP0.compile(hot_reload)?,
"probe_mip0",
)?;
let downsample = compute_pso(
device,
&mip_root,
&builtin_shaders::PROBE_DOWNSAMPLE.compile(hot_reload)?,
"probe_downsample",
)?;
let ggx = compute_pso(
device,
&ggx_root,
&builtin_shaders::PROBE_GGX.compile(hot_reload)?,
"probe_ggx",
)?;
Ok(Self {
mip_root,
ggx_root,
mip0,
downsample,
ggx,
})
}
}
pub(in crate::directx) struct PrefilterGpu {
capture: ID3D12Resource,
cube: usize,
mips: u32,
}
impl PrefilterGpu {
pub(in crate::directx) fn new(
ctx: &DxContext,
plan: &PrefilterPlan,
cubes: &super::probe_set::ProbeCubeArray,
cube: usize,
) -> RenderResult<PrefilterGpu> {
let mips = plan.mips();
if mips as usize > PROBE_MAX_MIPS {
return Err(RenderError::Other(format!(
"probe: {mips} mips exceeds the {PROBE_MAX_MIPS} descriptors reserved for one bake"
)));
}
if cube >= cubes.capacity() || cubes.mips() != mips {
return Err(RenderError::Other(format!(
"probe: the cube array has no cube {cube} at {mips} mips"
)));
}
let capture = create_cube_resource(
&ctx.hw.device,
CubeResource {
face_size: plan.face_size(),
mips,
cubes: 1,
state: D3D12_RESOURCE_STATE_COPY_DEST,
label: "probe capture cube",
},
)?;
let device = &ctx.hw.device;
let d = &ctx.descriptors;
write_cube_srv(device, &capture, mips, ctx.probe_capture_srv_cpu());
for mip in 0..mips {
write_cube_mip_uav(
device,
&capture,
0,
mip,
cpu_slot(ctx, d.layout.probe_capture_uav_base_slot + mip as usize),
);
cubes.write_mip_uav(
device,
cube,
mip,
cpu_slot(ctx, d.layout.probe_cube_uav_base_slot + mip as usize),
);
}
write_cube_mip_uav(
device,
&capture,
0,
0,
cpu_slot(ctx, d.layout.probe_mip0_pair_slot),
);
cubes.write_mip_uav(
device,
cube,
0,
cpu_slot(ctx, d.layout.probe_mip0_pair_slot + 1),
);
Ok(PrefilterGpu {
capture,
cube,
mips,
})
}
pub(in crate::directx) fn capture(&self) -> &ID3D12Resource {
&self.capture
}
pub(in crate::directx) fn cube(&self) -> usize {
self.cube
}
pub(in crate::directx) fn mips(&self) -> u32 {
self.mips
}
}
impl DxContext {
pub(in crate::directx) fn probe_capture_srv_cpu(&self) -> D3D12_CPU_DESCRIPTOR_HANDLE {
cpu_slot(self, self.descriptors.layout.probe_capture_srv_slot)
}
pub(in crate::directx) fn encode_probe_pyramid(
&self,
cmd: &ID3D12GraphicsCommandList,
gpu: &PrefilterGpu,
plan: &PrefilterPlan,
) -> RenderResult<()> {
let pipelines =
self.probe.prefilter.as_ref().ok_or_else(|| {
RenderError::Other("probe: prefilter pipelines missing".to_string())
})?;
let mut barriers = vec![super::texture::transition_barrier(
gpu.capture(),
D3D12_RESOURCE_STATE_COPY_DEST,
D3D12_RESOURCE_STATE_UNORDERED_ACCESS,
)];
barriers.extend(self.probe.gpu.bound_cubes().cube_barriers(
gpu.cube(),
super::probe_set::PROBE_CUBES_STATE,
D3D12_RESOURCE_STATE_UNORDERED_ACCESS,
));
unsafe {
cmd.ResourceBarrier(&barriers);
cmd.SetComputeRootSignature(&pipelines.mip_root);
}
let d = &self.descriptors;
self.dispatch_prefilter(
cmd,
&pipelines.mip0,
gpu_slot(self, d.layout.probe_mip0_pair_slot),
None,
&plan.mip0_params(),
plan.face_size(),
);
for mip in 1..plan.mips() {
uav_barrier(cmd, gpu.capture());
self.dispatch_prefilter(
cmd,
&pipelines.downsample,
gpu_slot(
self,
d.layout.probe_capture_uav_base_slot + (mip - 1) as usize,
),
None,
&plan.downsample_params(mip),
plan.mip_face_size(mip),
);
}
unsafe {
cmd.ResourceBarrier(&[super::texture::transition_barrier(
gpu.capture(),
D3D12_RESOURCE_STATE_UNORDERED_ACCESS,
D3D12_RESOURCE_STATE_NON_PIXEL_SHADER_RESOURCE,
)]);
}
Ok(())
}
pub(in crate::directx) fn encode_probe_ggx_mip(
&self,
cmd: &ID3D12GraphicsCommandList,
plan: &PrefilterPlan,
dst_mip: u32,
) -> RenderResult<()> {
let pipelines =
self.probe.prefilter.as_ref().ok_or_else(|| {
RenderError::Other("probe: prefilter pipelines missing".to_string())
})?;
unsafe { cmd.SetComputeRootSignature(&pipelines.ggx_root) };
let d = &self.descriptors;
self.dispatch_prefilter(
cmd,
&pipelines.ggx,
gpu_slot(self, d.layout.probe_cube_uav_base_slot + dst_mip as usize),
Some(gpu_slot(self, d.layout.probe_capture_srv_slot)),
&plan.ggx_params(dst_mip),
plan.mip_face_size(dst_mip),
);
Ok(())
}
fn dispatch_prefilter(
&self,
cmd: &ID3D12GraphicsCommandList,
pso: &ID3D12PipelineState,
uav_table: SrvSlot,
srv_table: Option<SrvSlot>,
params: &ProbePrefilterParams,
size: u32,
) {
let groups = size.div_ceil(PREFILTER_TILE).max(1);
unsafe {
cmd.SetPipelineState(pso);
cmd.set_compute_root_constants(0, params);
match srv_table {
Some(srv) => {
cmd.set_compute_srv_table(1, srv);
cmd.set_compute_srv_table(2, uav_table);
}
None => cmd.set_compute_srv_table(1, uav_table),
}
cmd.Dispatch(groups, groups, 6);
}
}
}
fn create_mip_root_signature(device: &ID3D12Device) -> RenderResult<ID3D12RootSignature> {
RootSig::new()
.constants::<ProbePrefilterParams>(0, Visibility::All)
.uav_table(0, 2, Visibility::All) .build(device, "probe prefilter mip root sig")
}
fn create_ggx_root_signature(device: &ID3D12Device) -> RenderResult<ID3D12RootSignature> {
use Visibility::All;
RootSig::new()
.constants::<ProbePrefilterParams>(0, All)
.srv_table(0, 1, All) .uav_table(0, 1, All) .static_sampler(SamplerState::LinearClamp, 0, All)
.build(device, "probe prefilter ggx root sig")
}
fn write_cube_srv(
device: &ID3D12Device,
resource: &ID3D12Resource,
mips: u32,
srv_cpu: D3D12_CPU_DESCRIPTOR_HANDLE,
) {
let desc = D3D12_SHADER_RESOURCE_VIEW_DESC {
Format: PROBE_CUBE_FORMAT,
ViewDimension: D3D12_SRV_DIMENSION_TEXTURECUBE,
Shader4ComponentMapping: D3D12_DEFAULT_SHADER_4_COMPONENT_MAPPING,
Anonymous: D3D12_SHADER_RESOURCE_VIEW_DESC_0 {
TextureCube: D3D12_TEXCUBE_SRV {
MostDetailedMip: 0,
MipLevels: mips,
ResourceMinLODClamp: 0.0,
},
},
};
unsafe { device.CreateShaderResourceView(resource, Some(&desc), srv_cpu) };
}
fn uav_barrier(cmd: &ID3D12GraphicsCommandList, resource: &ID3D12Resource) {
let barrier = D3D12_RESOURCE_BARRIER {
Type: D3D12_RESOURCE_BARRIER_TYPE_UAV,
Flags: D3D12_RESOURCE_BARRIER_FLAG_NONE,
Anonymous: D3D12_RESOURCE_BARRIER_0 {
UAV: std::mem::ManuallyDrop::new(D3D12_RESOURCE_UAV_BARRIER {
pResource: com::borrowed(resource),
}),
},
};
unsafe { cmd.ResourceBarrier(&[barrier]) };
}
fn cpu_slot(ctx: &DxContext, slot: usize) -> D3D12_CPU_DESCRIPTOR_HANDLE {
let base = unsafe {
ctx.descriptors
.srv_heap
.GetCPUDescriptorHandleForHeapStart()
};
D3D12_CPU_DESCRIPTOR_HANDLE {
ptr: base.ptr + slot * ctx.descriptors.srv_descriptor_size,
}
}
fn gpu_slot(ctx: &DxContext, slot: usize) -> SrvSlot {
SrvSlot::at(
&ctx.descriptors.srv_heap,
ctx.descriptors.srv_descriptor_size,
slot,
)
}