use ash::vk;
use concinnity_core::render::reflection_probe::PrefilterPlan;
use concinnity_core::render::uniforms::ProbePrefilterParams;
use super::allocator::{DeviceAllocator, PooledImage};
use super::owned::{
OwnedDescriptorPool, OwnedPipeline, OwnedPipelineLayout, OwnedSampler, OwnedSetLayout, VkDevice,
};
use super::resources::alloc_descriptor_sets;
pub(super) const PROBE_CUBE_FORMAT: vk::Format = vk::Format::R16G16B16A16_SFLOAT;
const PREFILTER_TILE: u32 = 8;
fn cube_range(mips: u32) -> vk::ImageSubresourceRange {
vk::ImageSubresourceRange {
aspect_mask: vk::ImageAspectFlags::COLOR,
base_mip_level: 0,
level_count: mips,
base_array_layer: 0,
layer_count: 6,
}
}
pub(super) struct ProbePrefilterPipelines {
mip_set_layout: OwnedSetLayout,
ggx_set_layout: OwnedSetLayout,
mip_pipeline_layout: OwnedPipelineLayout,
ggx_pipeline_layout: OwnedPipelineLayout,
mip0: OwnedPipeline,
downsample: OwnedPipeline,
ggx: OwnedPipeline,
sampler: OwnedSampler,
}
impl ProbePrefilterPipelines {
pub(super) fn new(device: &VkDevice, hot_reload: bool) -> Result<Self, String> {
use super::slang_builtins::SlangCompile;
let mip_set_layout = create_set_layout(
device,
&[
vk::DescriptorType::STORAGE_IMAGE,
vk::DescriptorType::STORAGE_IMAGE,
],
)?;
let ggx_set_layout = create_set_layout(
device,
&[
vk::DescriptorType::SAMPLED_IMAGE,
vk::DescriptorType::SAMPLER,
vk::DescriptorType::STORAGE_IMAGE,
],
)?;
let push = vk::PushConstantRange::default()
.stage_flags(vk::ShaderStageFlags::COMPUTE)
.offset(0)
.size(size_of::<ProbePrefilterParams>() as u32);
let mip_pipeline_layout = create_pipeline_layout(device, mip_set_layout.handle(), push)?;
let ggx_pipeline_layout = create_pipeline_layout(device, ggx_set_layout.handle(), push)?;
let ctx = super::builtins::Ctx::plain(hot_reload);
let mip0 = create_compute_pipeline(
device,
mip_pipeline_layout.handle(),
&super::slang_builtins::PROBE_MIP0.compile(&ctx)?,
"probe_mip0",
)?;
let downsample = create_compute_pipeline(
device,
mip_pipeline_layout.handle(),
&super::slang_builtins::PROBE_DOWNSAMPLE.compile(&ctx)?,
"probe_downsample",
)?;
let ggx = create_compute_pipeline(
device,
ggx_pipeline_layout.handle(),
&super::slang_builtins::PROBE_GGX.compile(&ctx)?,
"probe_ggx",
)?;
let sampler = super::texture::create_sampler_cube_linear(device)?;
Ok(Self {
mip_set_layout,
ggx_set_layout,
mip_pipeline_layout,
ggx_pipeline_layout,
mip0,
downsample,
ggx,
sampler,
})
}
}
pub(super) struct PrefilterGpu {
capture: PooledImage,
probe: PooledImage,
probe_cube_view: vk::ImageView,
probe_mip_views: Vec<vk::ImageView>,
mip0_set: vk::DescriptorSet,
downsample_sets: Vec<vk::DescriptorSet>,
ggx_sets: Vec<vk::DescriptorSet>,
#[expect(dead_code, reason = "owns the sets its handles name")]
pool: OwnedDescriptorPool,
mips: u32,
}
impl PrefilterGpu {
pub(super) fn new(
device: &VkDevice,
alloc: &DeviceAllocator,
pipelines: &ProbePrefilterPipelines,
plan: &PrefilterPlan,
) -> Result<PrefilterGpu, String> {
let mips = plan.mips();
let capture = create_cube_image(
alloc,
plan.face_size(),
mips,
vk::ImageUsageFlags::TRANSFER_DST
| vk::ImageUsageFlags::STORAGE
| vk::ImageUsageFlags::SAMPLED,
)?;
let probe = create_cube_image(
alloc,
plan.face_size(),
mips,
vk::ImageUsageFlags::STORAGE | vk::ImageUsageFlags::SAMPLED,
)?;
let capture_cube_view = create_cube_view(device, capture.image(), mips)?;
let probe_cube_view = create_cube_view(device, probe.image(), mips)?;
let capture_mip_views = mip_storage_views(device, capture.image(), mips)?;
let probe_mip_views = mip_storage_views(device, probe.image(), mips)?;
capture.attach_view(capture_cube_view);
probe.attach_view(probe_cube_view);
for &view in &capture_mip_views {
capture.attach_view(view);
}
for &view in &probe_mip_views {
probe.attach_view(view);
}
let steps = mips.saturating_sub(1) as usize;
let pool = create_pool(device, steps)?;
let mip_layouts = vec![pipelines.mip_set_layout.handle(); steps + 1];
let ggx_layouts = vec![pipelines.ggx_set_layout.handle(); steps];
let mut mip_sets = alloc_descriptor_sets(device, pool.handle(), &mip_layouts)?;
let ggx_sets = alloc_descriptor_sets(device, pool.handle(), &ggx_layouts)?;
let mip0_set = mip_sets.remove(0);
let downsample_sets = mip_sets;
write_storage_pair(device, mip0_set, capture_mip_views[0], probe_mip_views[0]);
for (step, &set) in downsample_sets.iter().enumerate() {
let dst = step + 1;
write_storage_pair(
device,
set,
capture_mip_views[dst - 1],
capture_mip_views[dst],
);
}
for (step, &set) in ggx_sets.iter().enumerate() {
write_ggx_set(
device,
set,
capture_cube_view,
pipelines.sampler.handle(),
probe_mip_views[step + 1],
);
}
Ok(PrefilterGpu {
capture,
probe,
probe_cube_view,
probe_mip_views,
mip0_set,
downsample_sets,
ggx_sets,
pool,
mips,
})
}
pub(super) fn capture_image(&self) -> vk::Image {
self.capture.image()
}
pub(super) fn probe_image(&self) -> vk::Image {
self.probe.image()
}
pub(super) fn into_probe_cube(self) -> super::texture::GpuImage {
super::texture::GpuImage::from_pooled_with_aux(
self.probe,
self.probe_cube_view,
self.probe_mip_views,
)
}
}
impl super::context::VkContext {
pub(in crate::vulkan) fn encode_probe_pyramid(
&self,
cmd: vk::CommandBuffer,
gpu: &PrefilterGpu,
plan: &PrefilterPlan,
) -> Result<(), String> {
let pipelines = self
.probe
.prefilter
.as_ref()
.ok_or("probe: prefilter pipelines missing")?;
let device = &self.device;
transition(
device,
cmd,
gpu.capture.image(),
gpu.mips,
LayoutSide {
layout: vk::ImageLayout::TRANSFER_DST_OPTIMAL,
access: vk::AccessFlags::TRANSFER_WRITE,
stage: vk::PipelineStageFlags::TRANSFER,
},
LayoutSide {
layout: vk::ImageLayout::GENERAL,
access: vk::AccessFlags::SHADER_READ | vk::AccessFlags::SHADER_WRITE,
stage: vk::PipelineStageFlags::COMPUTE_SHADER,
},
);
transition(
device,
cmd,
gpu.probe.image(),
gpu.mips,
LayoutSide {
layout: vk::ImageLayout::UNDEFINED,
access: vk::AccessFlags::empty(),
stage: vk::PipelineStageFlags::TOP_OF_PIPE,
},
LayoutSide {
layout: vk::ImageLayout::GENERAL,
access: vk::AccessFlags::SHADER_WRITE,
stage: vk::PipelineStageFlags::COMPUTE_SHADER,
},
);
self.dispatch_prefilter(
cmd,
pipelines.mip_pipeline_layout.handle(),
pipelines.mip0.handle(),
gpu.mip0_set,
&plan.mip0_params(),
plan.face_size(),
);
for (step, &set) in gpu.downsample_sets.iter().enumerate() {
let dst = step as u32 + 1;
storage_barrier(device, cmd);
self.dispatch_prefilter(
cmd,
pipelines.mip_pipeline_layout.handle(),
pipelines.downsample.handle(),
set,
&plan.downsample_params(dst),
plan.mip_face_size(dst),
);
}
transition(
device,
cmd,
gpu.capture.image(),
gpu.mips,
LayoutSide {
layout: vk::ImageLayout::GENERAL,
access: vk::AccessFlags::SHADER_WRITE,
stage: vk::PipelineStageFlags::COMPUTE_SHADER,
},
LayoutSide {
layout: vk::ImageLayout::SHADER_READ_ONLY_OPTIMAL,
access: vk::AccessFlags::SHADER_READ,
stage: vk::PipelineStageFlags::COMPUTE_SHADER,
},
);
Ok(())
}
pub(in crate::vulkan) fn encode_probe_ggx_mip(
&self,
cmd: vk::CommandBuffer,
gpu: &PrefilterGpu,
plan: &PrefilterPlan,
dst_mip: u32,
) -> Result<(), String> {
let pipelines = self
.probe
.prefilter
.as_ref()
.ok_or("probe: prefilter pipelines missing")?;
let set = *gpu
.ggx_sets
.get(dst_mip as usize - 1)
.ok_or("probe: convolution mip out of range")?;
self.dispatch_prefilter(
cmd,
pipelines.ggx_pipeline_layout.handle(),
pipelines.ggx.handle(),
set,
&plan.ggx_params(dst_mip),
plan.mip_face_size(dst_mip),
);
Ok(())
}
pub(in crate::vulkan) fn encode_probe_cube_readable(
&self,
cmd: vk::CommandBuffer,
image: vk::Image,
mips: u32,
) {
transition(
&self.device,
cmd,
image,
mips,
LayoutSide {
layout: vk::ImageLayout::GENERAL,
access: vk::AccessFlags::SHADER_WRITE,
stage: vk::PipelineStageFlags::COMPUTE_SHADER,
},
LayoutSide {
layout: vk::ImageLayout::SHADER_READ_ONLY_OPTIMAL,
access: vk::AccessFlags::SHADER_READ,
stage: vk::PipelineStageFlags::FRAGMENT_SHADER,
},
);
}
fn dispatch_prefilter(
&self,
cmd: vk::CommandBuffer,
layout: vk::PipelineLayout,
pipeline: vk::Pipeline,
set: vk::DescriptorSet,
params: &ProbePrefilterParams,
size: u32,
) {
let groups = size.div_ceil(PREFILTER_TILE).max(1);
unsafe {
self.device
.cmd_bind_pipeline(cmd, vk::PipelineBindPoint::COMPUTE, pipeline);
self.device.cmd_bind_descriptor_sets(
cmd,
vk::PipelineBindPoint::COMPUTE,
layout,
0,
std::slice::from_ref(&set),
&[],
);
self.device.cmd_push_constants(
cmd,
layout,
vk::ShaderStageFlags::COMPUTE,
0,
bytemuck::bytes_of(params),
);
self.device.cmd_dispatch(cmd, groups, groups, 6);
}
}
}
fn create_cube_image(
alloc: &DeviceAllocator,
face_size: u32,
mips: u32,
usage: vk::ImageUsageFlags,
) -> Result<PooledImage, String> {
let info = vk::ImageCreateInfo::default()
.flags(vk::ImageCreateFlags::CUBE_COMPATIBLE)
.image_type(vk::ImageType::TYPE_2D)
.extent(vk::Extent3D {
width: face_size,
height: face_size,
depth: 1,
})
.mip_levels(mips)
.array_layers(6)
.format(PROBE_CUBE_FORMAT)
.tiling(vk::ImageTiling::OPTIMAL)
.initial_layout(vk::ImageLayout::UNDEFINED)
.usage(usage)
.sharing_mode(vk::SharingMode::EXCLUSIVE)
.samples(vk::SampleCountFlags::TYPE_1);
alloc
.create_image(&info, vk::MemoryPropertyFlags::DEVICE_LOCAL)
.map_err(|e| format!("probe cube image: {e}"))
}
fn create_cube_view(
device: &VkDevice,
image: vk::Image,
mips: u32,
) -> Result<vk::ImageView, String> {
let info = vk::ImageViewCreateInfo::default()
.image(image)
.view_type(vk::ImageViewType::CUBE)
.format(PROBE_CUBE_FORMAT)
.subresource_range(cube_range(mips));
unsafe { device.create_image_view(&info, None) }.map_err(|e| format!("probe cube view: {e}"))
}
fn mip_storage_views(
device: &VkDevice,
image: vk::Image,
mips: u32,
) -> Result<Vec<vk::ImageView>, String> {
(0..mips)
.map(|mip| {
let info = vk::ImageViewCreateInfo::default()
.image(image)
.view_type(vk::ImageViewType::TYPE_2D_ARRAY)
.format(PROBE_CUBE_FORMAT)
.subresource_range(vk::ImageSubresourceRange {
aspect_mask: vk::ImageAspectFlags::COLOR,
base_mip_level: mip,
level_count: 1,
base_array_layer: 0,
layer_count: 6,
});
unsafe { device.create_image_view(&info, None) }
.map_err(|e| format!("probe mip {mip} view: {e}"))
})
.collect()
}
fn create_set_layout(
device: &VkDevice,
types: &[vk::DescriptorType],
) -> Result<OwnedSetLayout, String> {
let binds: Vec<_> = types
.iter()
.enumerate()
.map(|(i, &ty)| {
vk::DescriptorSetLayoutBinding::default()
.binding(i as u32)
.descriptor_type(ty)
.descriptor_count(1)
.stage_flags(vk::ShaderStageFlags::COMPUTE)
})
.collect();
device
.create_descriptor_set_layout(
&vk::DescriptorSetLayoutCreateInfo::default().bindings(&binds),
)
.map_err(|e| format!("probe prefilter set layout: {e}"))
}
fn create_pipeline_layout(
device: &VkDevice,
set_layout: vk::DescriptorSetLayout,
push: vk::PushConstantRange,
) -> Result<OwnedPipelineLayout, String> {
let layouts = [set_layout];
device
.create_pipeline_layout(
&vk::PipelineLayoutCreateInfo::default()
.set_layouts(&layouts)
.push_constant_ranges(std::slice::from_ref(&push)),
)
.map_err(|e| format!("probe prefilter pipeline layout: {e}"))
}
fn create_pool(device: &VkDevice, steps: usize) -> Result<OwnedDescriptorPool, String> {
let steps = steps as u32;
let sizes = [
vk::DescriptorPoolSize::default()
.ty(vk::DescriptorType::STORAGE_IMAGE)
.descriptor_count(2 + 3 * steps),
vk::DescriptorPoolSize::default()
.ty(vk::DescriptorType::SAMPLED_IMAGE)
.descriptor_count(steps),
vk::DescriptorPoolSize::default()
.ty(vk::DescriptorType::SAMPLER)
.descriptor_count(steps),
];
device
.create_descriptor_pool(
&vk::DescriptorPoolCreateInfo::default()
.pool_sizes(&sizes)
.max_sets(1 + 2 * steps),
)
.map_err(|e| format!("probe prefilter descriptor pool: {e}"))
}
fn write_storage_pair(
device: &VkDevice,
set: vk::DescriptorSet,
src: vk::ImageView,
dst: vk::ImageView,
) {
let src_info = storage_info(src);
let dst_info = storage_info(dst);
let writes = [
storage_write(set, 0, std::slice::from_ref(&src_info)),
storage_write(set, 1, std::slice::from_ref(&dst_info)),
];
unsafe { device.update_descriptor_sets(&writes, &[]) };
}
fn write_ggx_set(
device: &VkDevice,
set: vk::DescriptorSet,
cube: vk::ImageView,
sampler: vk::Sampler,
dst: vk::ImageView,
) {
let cube_info = vk::DescriptorImageInfo::default()
.image_layout(vk::ImageLayout::SHADER_READ_ONLY_OPTIMAL)
.image_view(cube);
let sampler_info = vk::DescriptorImageInfo::default().sampler(sampler);
let dst_info = storage_info(dst);
let writes = [
vk::WriteDescriptorSet::default()
.dst_set(set)
.dst_binding(0)
.descriptor_type(vk::DescriptorType::SAMPLED_IMAGE)
.image_info(std::slice::from_ref(&cube_info)),
vk::WriteDescriptorSet::default()
.dst_set(set)
.dst_binding(1)
.descriptor_type(vk::DescriptorType::SAMPLER)
.image_info(std::slice::from_ref(&sampler_info)),
storage_write(set, 2, std::slice::from_ref(&dst_info)),
];
unsafe { device.update_descriptor_sets(&writes, &[]) };
}
fn storage_info(view: vk::ImageView) -> vk::DescriptorImageInfo {
vk::DescriptorImageInfo::default()
.image_layout(vk::ImageLayout::GENERAL)
.image_view(view)
}
fn storage_write<'a>(
set: vk::DescriptorSet,
binding: u32,
info: &'a [vk::DescriptorImageInfo],
) -> vk::WriteDescriptorSet<'a> {
vk::WriteDescriptorSet::default()
.dst_set(set)
.dst_binding(binding)
.descriptor_type(vk::DescriptorType::STORAGE_IMAGE)
.image_info(info)
}
fn storage_barrier(device: &VkDevice, cmd: vk::CommandBuffer) {
let barrier = vk::MemoryBarrier::default()
.src_access_mask(vk::AccessFlags::SHADER_WRITE)
.dst_access_mask(vk::AccessFlags::SHADER_READ | vk::AccessFlags::SHADER_WRITE);
unsafe {
device.cmd_pipeline_barrier(
cmd,
vk::PipelineStageFlags::COMPUTE_SHADER,
vk::PipelineStageFlags::COMPUTE_SHADER,
vk::DependencyFlags::empty(),
std::slice::from_ref(&barrier),
&[],
&[],
);
}
}
#[derive(Clone, Copy)]
struct LayoutSide {
layout: vk::ImageLayout,
access: vk::AccessFlags,
stage: vk::PipelineStageFlags,
}
fn transition(
device: &VkDevice,
cmd: vk::CommandBuffer,
image: vk::Image,
mips: u32,
from: LayoutSide,
to: LayoutSide,
) {
let barrier = vk::ImageMemoryBarrier::default()
.src_access_mask(from.access)
.dst_access_mask(to.access)
.old_layout(from.layout)
.new_layout(to.layout)
.src_queue_family_index(vk::QUEUE_FAMILY_IGNORED)
.dst_queue_family_index(vk::QUEUE_FAMILY_IGNORED)
.image(image)
.subresource_range(cube_range(mips));
unsafe {
device.cmd_pipeline_barrier(
cmd,
from.stage,
to.stage,
vk::DependencyFlags::empty(),
&[],
&[],
std::slice::from_ref(&barrier),
);
}
}
fn create_compute_pipeline(
device: &VkDevice,
layout: vk::PipelineLayout,
spv: &[u8],
label: &str,
) -> Result<OwnedPipeline, String> {
let module = super::pipeline::spv_module(device, spv)?;
let entry = std::ffi::CString::new("main").expect("static entry name has no interior nul");
let stage = vk::PipelineShaderStageCreateInfo::default()
.stage(vk::ShaderStageFlags::COMPUTE)
.module(module.handle())
.name(&entry);
let info = vk::ComputePipelineCreateInfo::default()
.stage(stage)
.layout(layout);
crate::vulkan::pipeline_cache::create_compute_pipeline(device, &info)
.map_err(|e| format!("create {label} pipeline: {e}"))
}