use ash::vk;
use crate::vulkan::owned::{
OwnedDescriptorPool, OwnedPipeline, OwnedPipelineLayout, OwnedSampler, OwnedSetLayout, VkDevice,
};
use super::allocator::{DeviceAllocator, PooledBuffer, PooledImage};
use super::pipeline::spv_module;
use super::resources::alloc_descriptor_sets;
use super::texture::{
LayoutTransition, SubresourceRange, one_shot_submit, transition_image_layout_range,
};
const MAX_HIZ_MIPS: usize = 16;
const HIZ_TILE: u32 = 8;
pub(in crate::vulkan) use crate::vulkan::uniforms::CullHizParams;
use concinnity_render::uniforms::HizParams;
pub(super) fn hiz_mip_count(width: u32, height: u32) -> u32 {
let m = width.max(height).max(1);
32 - m.leading_zeros()
}
pub(super) struct HiZResources {
init_pipeline: OwnedPipeline,
downsample_pipeline: OwnedPipeline,
init_pipeline_layout: OwnedPipelineLayout,
downsample_pipeline_layout: OwnedPipelineLayout,
init_set_layout: OwnedSetLayout,
downsample_set_layout: OwnedSetLayout,
pub(super) read_set_layout: OwnedSetLayout,
descriptor_pool: OwnedDescriptorPool,
pub(super) pyramid: PooledImage,
sampled_view: vk::ImageView,
mip_views: Vec<vk::ImageView>,
sampler: OwnedSampler,
init_sets: Vec<vk::DescriptorSet>,
downsample_sets: Vec<vk::DescriptorSet>,
pub(super) read_sets: Vec<vk::DescriptorSet>,
pub(super) cull_ubos: Vec<PooledBuffer>,
pub(super) read_sets2: Vec<vk::DescriptorSet>,
pub(super) cull_ubos2: Vec<PooledBuffer>,
pub(super) width: u32,
pub(super) height: u32,
pub(super) mip_count: u32,
sample_count: u32,
}
fn create_hiz_image(
alloc: &DeviceAllocator,
width: u32,
height: u32,
mip_count: u32,
) -> Result<PooledImage, String> {
let img_info = vk::ImageCreateInfo::default()
.image_type(vk::ImageType::TYPE_2D)
.extent(vk::Extent3D {
width: width.max(1),
height: height.max(1),
depth: 1,
})
.mip_levels(mip_count.max(1))
.array_layers(1)
.format(vk::Format::R32_SFLOAT)
.tiling(vk::ImageTiling::OPTIMAL)
.initial_layout(vk::ImageLayout::UNDEFINED)
.usage(vk::ImageUsageFlags::STORAGE | vk::ImageUsageFlags::SAMPLED)
.sharing_mode(vk::SharingMode::EXCLUSIVE)
.samples(vk::SampleCountFlags::TYPE_1);
alloc
.create_image(&img_info, vk::MemoryPropertyFlags::DEVICE_LOCAL)
.map_err(|e| format!("hiz image: {e}"))
}
fn create_hiz_view(
device: &VkDevice,
image: vk::Image,
base_mip: u32,
level_count: u32,
) -> Result<vk::ImageView, String> {
let info = vk::ImageViewCreateInfo::default()
.image(image)
.view_type(vk::ImageViewType::TYPE_2D)
.format(vk::Format::R32_SFLOAT)
.subresource_range(vk::ImageSubresourceRange {
aspect_mask: vk::ImageAspectFlags::COLOR,
base_mip_level: base_mip,
level_count,
base_array_layer: 0,
layer_count: 1,
});
unsafe { device.create_image_view(&info, None) }.map_err(|e| format!("hiz view: {e}"))
}
fn build_hiz_pipelines(
device: &VkDevice,
init_layout: vk::PipelineLayout,
downsample_layout: vk::PipelineLayout,
sample_count: u32,
hot_reload: bool,
) -> Result<(OwnedPipeline, OwnedPipeline), String> {
let ctx = super::builtins::Ctx::plain(hot_reload);
let init_spv = if sample_count > 1 {
super::slang_builtins::HIZ_INIT_MSAA.compile(&ctx)?
} else {
super::slang_builtins::HIZ_INIT_SINGLE.compile(&ctx)?
};
let downsample_spv = super::slang_builtins::HIZ_DOWNSAMPLE.compile(&ctx)?;
let init = create_compute_pipeline(device, init_layout, &init_spv)?;
let downsample = create_compute_pipeline(device, downsample_layout, &downsample_spv)?;
Ok((init, downsample))
}
fn create_compute_pipeline(
device: &VkDevice,
layout: vk::PipelineLayout,
spv: &[u8],
) -> Result<OwnedPipeline, String> {
let module = spv_module(device, spv)?;
let entry = std::ffi::CString::new("main").unwrap();
let stage = vk::PipelineShaderStageCreateInfo::default()
.stage(vk::ShaderStageFlags::COMPUTE)
.module(module.handle())
.name(&entry);
let info = vk::ComputePipelineCreateInfo::default()
.stage(stage)
.layout(layout);
let pipeline = crate::vulkan::pipeline_cache::create_compute_pipeline(device, &info)
.map_err(|e| format!("create hiz pipeline: {e}"))?;
Ok(pipeline)
}
#[derive(Clone, Copy)]
pub(super) struct HiZDeviceCtx<'a> {
pub(super) alloc: &'a DeviceAllocator,
pub(super) device: &'a VkDevice,
pub(super) command_pool: vk::CommandPool,
pub(super) queue: vk::Queue,
}
#[derive(Clone, Copy)]
pub(super) struct HiZTarget<'a> {
pub(super) width: u32,
pub(super) height: u32,
pub(super) depth_views: &'a [vk::ImageView],
}
impl HiZResources {
pub(super) fn read_set_sources(&self) -> (vk::ImageView, vk::Sampler) {
(self.sampled_view, self.sampler.handle())
}
pub(super) fn new(
ctx: HiZDeviceCtx,
target: HiZTarget,
sample_count: u32,
frames: usize,
two_pass: bool,
hot_reload: bool,
) -> Result<Self, String> {
let HiZDeviceCtx { alloc, device, .. } = ctx;
let HiZTarget { width, height, .. } = target;
let init_set_layout = create_set_layout(
device,
&[
(0, vk::DescriptorType::SAMPLED_IMAGE),
(1, vk::DescriptorType::STORAGE_IMAGE),
],
)?;
let downsample_set_layout = create_set_layout(
device,
&[
(0, vk::DescriptorType::STORAGE_IMAGE),
(1, vk::DescriptorType::STORAGE_IMAGE),
],
)?;
let read_set_layout = create_set_layout(
device,
&[
(0, vk::DescriptorType::COMBINED_IMAGE_SAMPLER),
(1, vk::DescriptorType::UNIFORM_BUFFER),
],
)?;
let push_range = vk::PushConstantRange::default()
.stage_flags(vk::ShaderStageFlags::COMPUTE)
.offset(0)
.size(std::mem::size_of::<HizParams>() as u32);
let init_pipeline_layout =
create_pipeline_layout(device, init_set_layout.handle(), push_range)?;
let downsample_pipeline_layout =
create_pipeline_layout(device, downsample_set_layout.handle(), push_range)?;
let (init_pipeline, downsample_pipeline) = build_hiz_pipelines(
device,
init_pipeline_layout.handle(),
downsample_pipeline_layout.handle(),
sample_count,
hot_reload,
)?;
let descriptor_pool = create_pool(device, frames, two_pass)?;
let ubo_size = std::mem::size_of::<CullHizParams>() as u64;
let alloc_ubo_ring = |count: usize| -> Result<Vec<PooledBuffer>, String> {
(0..count)
.map(|_| {
alloc
.create_buffer(
ubo_size,
vk::BufferUsageFlags::UNIFORM_BUFFER,
vk::MemoryPropertyFlags::HOST_VISIBLE
| vk::MemoryPropertyFlags::HOST_COHERENT,
)
.map_err(String::from)
})
.collect()
};
let cull_ubos = alloc_ubo_ring(frames)?;
let cull_ubos2 = alloc_ubo_ring(if two_pass { frames } else { 0 })?;
let mut res = Self {
init_pipeline,
downsample_pipeline,
init_pipeline_layout,
downsample_pipeline_layout,
init_set_layout,
downsample_set_layout,
read_set_layout,
descriptor_pool,
pyramid: PooledImage::null(),
sampled_view: vk::ImageView::null(),
mip_views: Vec::new(),
sampler: create_sampler(device)?,
init_sets: Vec::new(),
downsample_sets: Vec::new(),
read_sets: Vec::new(),
cull_ubos,
read_sets2: Vec::new(),
cull_ubos2,
width,
height,
mip_count: 0,
sample_count,
};
res.create_image_and_sets(ctx, target)?;
Ok(res)
}
fn create_image_and_sets(
&mut self,
ctx: HiZDeviceCtx,
target: HiZTarget,
) -> Result<(), String> {
let HiZDeviceCtx {
alloc,
device,
command_pool,
queue,
} = ctx;
let HiZTarget {
width,
height,
depth_views,
} = target;
let mip_count = hiz_mip_count(width, height).min(MAX_HIZ_MIPS as u32).max(1);
let pyramid = create_hiz_image(alloc, width, height, mip_count)?;
one_shot_submit(device, command_pool, queue, |cmd| {
transition_image_layout_range(
device,
cmd,
pyramid.image(),
LayoutTransition {
old_layout: vk::ImageLayout::UNDEFINED,
new_layout: vk::ImageLayout::SHADER_READ_ONLY_OPTIMAL,
aspect: vk::ImageAspectFlags::COLOR,
},
SubresourceRange {
base_layer: 0,
layer_count: 1,
base_mip: 0,
mip_count,
},
);
})?;
let sampled_view = create_hiz_view(device, pyramid.image(), 0, mip_count)?;
pyramid.attach_view(sampled_view);
let mut mip_views = Vec::with_capacity(mip_count as usize);
for mip in 0..mip_count {
let view = create_hiz_view(device, pyramid.image(), mip, 1)?;
pyramid.attach_view(view);
mip_views.push(view);
}
unsafe {
device
.reset_descriptor_pool(
self.descriptor_pool.handle(),
vk::DescriptorPoolResetFlags::empty(),
)
.map_err(|e| format!("reset hiz pool: {e}"))?;
}
let frames = self.cull_ubos.len();
let init_layouts: Vec<_> = (0..frames).map(|_| self.init_set_layout.handle()).collect();
let init_sets =
alloc_descriptor_sets(device, self.descriptor_pool.handle(), &init_layouts)?;
let downsample_layouts: Vec<_> = (1..mip_count)
.map(|_| self.downsample_set_layout.handle())
.collect();
let downsample_sets =
alloc_descriptor_sets(device, self.descriptor_pool.handle(), &downsample_layouts)?;
let read_layouts: Vec<_> = (0..frames).map(|_| self.read_set_layout.handle()).collect();
let read_sets =
alloc_descriptor_sets(device, self.descriptor_pool.handle(), &read_layouts)?;
let two_pass = !self.cull_ubos2.is_empty();
let read_layouts2: Vec<_> = (0..if two_pass { frames } else { 0 })
.map(|_| self.read_set_layout.handle())
.collect();
let read_sets2 =
alloc_descriptor_sets(device, self.descriptor_pool.handle(), &read_layouts2)?;
for (i, &set) in init_sets.iter().enumerate() {
let depth = depth_views[i.min(depth_views.len().saturating_sub(1))];
write_sampled_image(device, set, 0, depth);
write_storage_image(device, set, 1, mip_views[0]);
}
for (step, &set) in downsample_sets.iter().enumerate() {
let m = step + 1;
write_storage_image(device, set, 0, mip_views[m - 1]);
write_storage_image(device, set, 1, mip_views[m]);
}
for (i, &set) in read_sets.iter().enumerate() {
write_sampler(device, set, 0, sampled_view, self.sampler.handle());
write_uniform_buffer(
device,
set,
1,
self.cull_ubos[i].buffer(),
std::mem::size_of::<CullHizParams>() as u64,
);
}
for (i, &set) in read_sets2.iter().enumerate() {
write_sampler(device, set, 0, sampled_view, self.sampler.handle());
write_uniform_buffer(
device,
set,
1,
self.cull_ubos2[i].buffer(),
std::mem::size_of::<CullHizParams>() as u64,
);
}
self.pyramid = pyramid;
self.sampled_view = sampled_view;
self.mip_views = mip_views;
self.init_sets = init_sets;
self.downsample_sets = downsample_sets;
self.read_sets = read_sets;
self.read_sets2 = read_sets2;
self.width = width;
self.height = height;
self.mip_count = mip_count;
Ok(())
}
pub(super) fn resize_to(&mut self, ctx: HiZDeviceCtx, target: HiZTarget) -> Result<(), String> {
self.create_image_and_sets(ctx, target)
}
pub(super) fn swap_pipelines(&mut self, init: OwnedPipeline, downsample: OwnedPipeline) {
self.init_pipeline = init;
self.downsample_pipeline = downsample;
}
pub(super) fn recompile_pipelines(
&self,
device: &VkDevice,
hot_reload: bool,
) -> Result<(OwnedPipeline, OwnedPipeline), String> {
build_hiz_pipelines(
device,
self.init_pipeline_layout.handle(),
self.downsample_pipeline_layout.handle(),
self.sample_count,
hot_reload,
)
}
pub(super) fn destroy(&mut self, _device: &VkDevice) {
self.pyramid = PooledImage::null();
self.sampled_view = vk::ImageView::null();
self.mip_views.clear();
self.cull_ubos.clear();
self.cull_ubos2.clear();
}
}
impl crate::vulkan::context::VkContext {
pub(in crate::vulkan) fn encode_hiz_build(&self, cmd: vk::CommandBuffer, frame_idx: usize) {
let Some(hiz) = self.cull.hiz.as_ref() else {
return;
};
if hiz.mip_count == 0 || hiz.mip_views.is_empty() {
return;
}
let device = &self.device;
let init_params = HizParams {
dst_width: hiz.width,
dst_height: hiz.height,
src_mip: 0,
sample_count: hiz.sample_count.max(1),
};
unsafe {
device.cmd_bind_pipeline(
cmd,
vk::PipelineBindPoint::COMPUTE,
hiz.init_pipeline.handle(),
);
device.cmd_bind_descriptor_sets(
cmd,
vk::PipelineBindPoint::COMPUTE,
hiz.init_pipeline_layout.handle(),
0,
std::slice::from_ref(&hiz.init_sets[frame_idx]),
&[],
);
device.cmd_push_constants(
cmd,
hiz.init_pipeline_layout.handle(),
vk::ShaderStageFlags::COMPUTE,
0,
as_bytes(&init_params),
);
device.cmd_dispatch(
cmd,
hiz.width.div_ceil(HIZ_TILE),
hiz.height.div_ceil(HIZ_TILE),
1,
);
}
let mut cur_w = hiz.width;
let mut cur_h = hiz.height;
for mip in 1..hiz.mip_count {
unsafe {
device.cmd_pipeline_barrier(
cmd,
vk::PipelineStageFlags::COMPUTE_SHADER,
vk::PipelineStageFlags::COMPUTE_SHADER,
vk::DependencyFlags::empty(),
&[],
&[],
&[hiz_image_barrier(
hiz.pyramid.image(),
hiz.mip_count,
vk::ImageLayout::GENERAL,
vk::ImageLayout::GENERAL,
vk::AccessFlags::SHADER_WRITE,
vk::AccessFlags::SHADER_READ,
)],
);
}
let next_w = (cur_w / 2).max(1);
let next_h = (cur_h / 2).max(1);
let params = HizParams {
dst_width: next_w,
dst_height: next_h,
src_mip: mip - 1,
sample_count: 0,
};
unsafe {
device.cmd_bind_pipeline(
cmd,
vk::PipelineBindPoint::COMPUTE,
hiz.downsample_pipeline.handle(),
);
device.cmd_bind_descriptor_sets(
cmd,
vk::PipelineBindPoint::COMPUTE,
hiz.downsample_pipeline_layout.handle(),
0,
std::slice::from_ref(&hiz.downsample_sets[(mip - 1) as usize]),
&[],
);
device.cmd_push_constants(
cmd,
hiz.downsample_pipeline_layout.handle(),
vk::ShaderStageFlags::COMPUTE,
0,
as_bytes(¶ms),
);
device.cmd_dispatch(cmd, next_w.div_ceil(HIZ_TILE), next_h.div_ceil(HIZ_TILE), 1);
}
cur_w = next_w;
cur_h = next_h;
}
}
}
fn as_bytes<T: bytemuck::NoUninit>(v: &T) -> &[u8] {
bytemuck::bytes_of(v)
}
fn hiz_image_barrier(
image: vk::Image,
mip_count: u32,
old: vk::ImageLayout,
new: vk::ImageLayout,
src: vk::AccessFlags,
dst: vk::AccessFlags,
) -> vk::ImageMemoryBarrier<'static> {
vk::ImageMemoryBarrier::default()
.src_access_mask(src)
.dst_access_mask(dst)
.old_layout(old)
.new_layout(new)
.src_queue_family_index(vk::QUEUE_FAMILY_IGNORED)
.dst_queue_family_index(vk::QUEUE_FAMILY_IGNORED)
.image(image)
.subresource_range(vk::ImageSubresourceRange {
aspect_mask: vk::ImageAspectFlags::COLOR,
base_mip_level: 0,
level_count: mip_count,
base_array_layer: 0,
layer_count: 1,
})
}
fn create_set_layout(
device: &VkDevice,
bindings: &[(u32, vk::DescriptorType)],
) -> Result<OwnedSetLayout, String> {
let binds: Vec<_> = bindings
.iter()
.map(|&(b, ty)| {
vk::DescriptorSetLayoutBinding::default()
.binding(b)
.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!("hiz set layout: {e}"))
}
fn create_pipeline_layout(
device: &VkDevice,
set_layout: vk::DescriptorSetLayout,
push_range: 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_range)),
)
.map_err(|e| format!("hiz pipeline layout: {e}"))
}
fn create_pool(
device: &VkDevice,
frames: usize,
two_pass: bool,
) -> Result<OwnedDescriptorPool, String> {
let f = frames as u32;
let read_rings = if two_pass { 2 } else { 1 };
let sizes = [
vk::DescriptorPoolSize::default()
.ty(vk::DescriptorType::COMBINED_IMAGE_SAMPLER)
.descriptor_count(read_rings * f),
vk::DescriptorPoolSize::default()
.ty(vk::DescriptorType::SAMPLED_IMAGE)
.descriptor_count(f),
vk::DescriptorPoolSize::default()
.ty(vk::DescriptorType::STORAGE_IMAGE)
.descriptor_count(f + 2 * MAX_HIZ_MIPS as u32),
vk::DescriptorPoolSize::default()
.ty(vk::DescriptorType::UNIFORM_BUFFER)
.descriptor_count(read_rings * f),
];
let max_sets = (1 + read_rings) * f + MAX_HIZ_MIPS as u32;
device
.create_descriptor_pool(
&vk::DescriptorPoolCreateInfo::default()
.pool_sizes(&sizes)
.max_sets(max_sets),
)
.map_err(|e| format!("hiz descriptor pool: {e}"))
}
fn create_sampler(device: &VkDevice) -> Result<OwnedSampler, String> {
let info = vk::SamplerCreateInfo::default()
.mag_filter(vk::Filter::NEAREST)
.min_filter(vk::Filter::NEAREST)
.mipmap_mode(vk::SamplerMipmapMode::NEAREST)
.address_mode_u(vk::SamplerAddressMode::CLAMP_TO_EDGE)
.address_mode_v(vk::SamplerAddressMode::CLAMP_TO_EDGE)
.address_mode_w(vk::SamplerAddressMode::CLAMP_TO_EDGE)
.min_lod(0.0)
.max_lod(MAX_HIZ_MIPS as f32);
device
.create_sampler(&info)
.map_err(|e| format!("hiz sampler: {e}"))
}
fn write_sampler(
device: &VkDevice,
set: vk::DescriptorSet,
binding: u32,
view: vk::ImageView,
sampler: vk::Sampler,
) {
let info = vk::DescriptorImageInfo::default()
.image_layout(vk::ImageLayout::SHADER_READ_ONLY_OPTIMAL)
.image_view(view)
.sampler(sampler);
let write = vk::WriteDescriptorSet::default()
.dst_set(set)
.dst_binding(binding)
.descriptor_type(vk::DescriptorType::COMBINED_IMAGE_SAMPLER)
.image_info(std::slice::from_ref(&info));
unsafe { device.update_descriptor_sets(std::slice::from_ref(&write), &[]) };
}
fn write_sampled_image(
device: &VkDevice,
set: vk::DescriptorSet,
binding: u32,
view: vk::ImageView,
) {
let info = vk::DescriptorImageInfo::default()
.image_layout(vk::ImageLayout::SHADER_READ_ONLY_OPTIMAL)
.image_view(view);
let write = vk::WriteDescriptorSet::default()
.dst_set(set)
.dst_binding(binding)
.descriptor_type(vk::DescriptorType::SAMPLED_IMAGE)
.image_info(std::slice::from_ref(&info));
unsafe { device.update_descriptor_sets(std::slice::from_ref(&write), &[]) };
}
fn write_storage_image(
device: &VkDevice,
set: vk::DescriptorSet,
binding: u32,
view: vk::ImageView,
) {
let info = vk::DescriptorImageInfo::default()
.image_layout(vk::ImageLayout::GENERAL)
.image_view(view);
let write = vk::WriteDescriptorSet::default()
.dst_set(set)
.dst_binding(binding)
.descriptor_type(vk::DescriptorType::STORAGE_IMAGE)
.image_info(std::slice::from_ref(&info));
unsafe { device.update_descriptor_sets(std::slice::from_ref(&write), &[]) };
}
fn write_uniform_buffer(
device: &VkDevice,
set: vk::DescriptorSet,
binding: u32,
buffer: vk::Buffer,
range: u64,
) {
let info = vk::DescriptorBufferInfo::default()
.buffer(buffer)
.offset(0)
.range(range);
let write = vk::WriteDescriptorSet::default()
.dst_set(set)
.dst_binding(binding)
.descriptor_type(vk::DescriptorType::UNIFORM_BUFFER)
.buffer_info(std::slice::from_ref(&info));
unsafe { device.update_descriptor_sets(std::slice::from_ref(&write), &[]) };
}
#[cfg(test)]
mod tests {
use super::hiz_mip_count;
#[test]
fn mip_count_power_of_two() {
assert_eq!(hiz_mip_count(1, 1), 1);
assert_eq!(hiz_mip_count(2, 2), 2);
assert_eq!(hiz_mip_count(256, 256), 9);
assert_eq!(hiz_mip_count(1024, 1024), 11);
}
#[test]
fn mip_count_uses_larger_dimension() {
assert_eq!(hiz_mip_count(1920, 1080), hiz_mip_count(1920, 1920));
assert_eq!(hiz_mip_count(1920, 1080), 11);
assert_eq!(hiz_mip_count(1280, 720), 11);
}
#[test]
fn mip_count_clamps_zero() {
assert_eq!(hiz_mip_count(0, 0), 1);
assert_eq!(hiz_mip_count(0, 8), 4);
}
}