use std::collections::HashMap;
use std::sync::Arc;
use vulkano::device::Device;
use vulkano::pipeline::compute::ComputePipelineCreateInfo;
use vulkano::pipeline::layout::{
PipelineDescriptorSetLayoutCreateInfo, PipelineLayout,
};
use vulkano::pipeline::{ComputePipeline, PipelineShaderStageCreateInfo};
use vulkano::shader::ShaderModule;
use crate::shaders::compute::*;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub struct ShaderBindings {
pub needs_read_buffer: bool, pub needs_write_buffer: bool, pub needs_grid_counts: bool, pub needs_grid_objects: bool, pub needs_big_indices: bool, }
impl ShaderBindings {
pub fn basic() -> Self {
Self {
needs_read_buffer: true,
needs_write_buffer: true,
needs_grid_counts: false,
needs_grid_objects: false,
needs_big_indices: false,
}
}
pub fn grid_build() -> Self {
Self {
needs_read_buffer: true,
needs_write_buffer: true,
needs_grid_counts: true,
needs_grid_objects: true,
needs_big_indices: true,
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Default)]
pub enum ComputeShaderType {
#[default]
FullPhysics,
MidPhysic,
Static,
NoCollision,
GridBuild,
Empty,
Cull,
Test,
}
impl ComputeShaderType {
pub fn sort_key(&self) -> u32 {
match self {
ComputeShaderType::FullPhysics => 0,
ComputeShaderType::MidPhysic => 1,
ComputeShaderType::Static => 2,
ComputeShaderType::NoCollision => 3,
ComputeShaderType::GridBuild => 4,
ComputeShaderType::Empty => 5,
ComputeShaderType::Cull => 6,
ComputeShaderType::Test => 7,
}
}
pub fn needs_bindings(&self) -> ShaderBindings {
match self {
ComputeShaderType::GridBuild => ShaderBindings::grid_build(),
ComputeShaderType::Test => ShaderBindings::grid_build(),
_ => ShaderBindings::basic(),
}
}
}
pub struct ComputeShaderRegistry {
pipelines: HashMap<ComputeShaderType, Arc<ComputePipeline>>,
scene_shader: Option<ComputeShaderType>,
}
impl ComputeShaderRegistry {
pub fn new(device: &Arc<Device>) -> Self {
let mut pipelines = HashMap::new();
let cs_full = cs_full::load(device.clone())
.expect("Failed to load FullPhysics compute shader");
let cp_full = create_compute_pipeline(device, cs_full, "FullPhysics");
pipelines.insert(ComputeShaderType::FullPhysics, cp_full);
let cs_mid = cs_no_rot::load(device.clone())
.expect("Failed to load MidPhysics compute shader");
let cp_mid = create_compute_pipeline(device, cs_mid, "MidPhysics");
pipelines.insert(ComputeShaderType::MidPhysic, cp_mid);
let cs_static = cs_empty::load(device.clone())
.expect("Failed to load Static compute shader");
let cp_static = create_compute_pipeline(device, cs_static, "Static");
pipelines.insert(ComputeShaderType::Static, cp_static);
let cs_no_col = cs_no_coll::load(device.clone())
.expect("Failed to load NoCollision compute shader");
let cp_no_col =
create_compute_pipeline(device, cs_no_col, "NoCollision");
pipelines.insert(ComputeShaderType::NoCollision, cp_no_col);
let cs_grid = cs_grid_build::load(device.clone())
.expect("Failed to load GridBuild compute shader");
let cp_grid = create_compute_pipeline(device, cs_grid, "GridBuild");
pipelines.insert(ComputeShaderType::GridBuild, cp_grid);
let cs_empty = cs_empty::load(device.clone())
.expect("Failed to load Empty compute shader");
let cp_empty = create_compute_pipeline(device, cs_empty, "Empty");
pipelines.insert(ComputeShaderType::Empty, cp_empty);
let cs_cull = cs_cull::load(device.clone())
.expect("Failed to load Cull compute shader");
let cp_cull = create_compute_pipeline(device, cs_cull, "Cull");
pipelines.insert(ComputeShaderType::Cull, cp_cull);
let cs_test = cs_test::load(device.clone())
.expect("Failed to load Test compute shader");
let cp_test = create_compute_pipeline(device, cs_test, "Test");
pipelines.insert(ComputeShaderType::Test, cp_test);
Self {
pipelines,
scene_shader: None,
}
}
pub fn get_pipeline(
&self,
shader_type: ComputeShaderType,
) -> &Arc<ComputePipeline> {
self.pipelines
.get(&shader_type)
.expect("Compute pipeline not found")
}
pub fn set_scene_shader(&mut self, shader: ComputeShaderType) {
self.scene_shader = Some(shader);
}
pub fn get_default_shader(&self) -> ComputeShaderType {
ComputeShaderType::default()
}
pub fn clear_scene_shader(&mut self) {
self.scene_shader = None;
}
pub fn scene_shader(&self) -> ComputeShaderType {
self.scene_shader
.unwrap_or_else(|| self.get_default_shader())
}
pub fn scene_shader_optional(&self) -> Option<ComputeShaderType> {
self.scene_shader
}
}
fn create_compute_pipeline(
device: &Arc<Device>,
shader: Arc<ShaderModule>,
name: &str,
) -> Arc<ComputePipeline> {
let stage =
PipelineShaderStageCreateInfo::new(shader.entry_point("main").unwrap());
let layout = PipelineLayout::new(
device.clone(),
PipelineDescriptorSetLayoutCreateInfo::from_stages([&stage])
.into_pipeline_layout_create_info(device.clone())
.unwrap(),
)
.unwrap();
ComputePipeline::new(
device.clone(),
None,
ComputePipelineCreateInfo::stage_layout(stage, layout),
)
.unwrap_or_else(|error| {
panic!("Failed to create {name} compute pipeline: {error}")
})
}
#[repr(C)]
#[derive(Copy, Clone, Debug, bytemuck::Pod, bytemuck::Zeroable)]
pub struct CullPushConstants {
pub view_proj: [[f32; 4]; 4],
pub batch_offset: u32, pub batch_count: u32, pub visible_list_offset: u32, }