From 381dcdeb9ac5b1b06aea88c18d8563af1ad2cf05 Mon Sep 17 00:00:00 2001 From: Matt Keeter Date: Sat, 12 Sep 2026 10:18:56 -0400 Subject: [PATCH] Lazy compute pipeline construction --- fidget-wgpu/src/lib.rs | 29 ++--- fidget-wgpu/src/pixel/effects/mod.rs | 9 +- fidget-wgpu/src/pixel/mod.rs | 27 +++-- fidget-wgpu/src/voxel/effects/mod.rs | 9 +- fidget-wgpu/src/voxel/mod.rs | 168 +++++++++++++++------------ 5 files changed, 136 insertions(+), 106 deletions(-) diff --git a/fidget-wgpu/src/lib.rs b/fidget-wgpu/src/lib.rs index 10ba0deb..3ef205b5 100644 --- a/fidget-wgpu/src/lib.rs +++ b/fidget-wgpu/src/lib.rs @@ -73,7 +73,6 @@ use fidget_core::{ use fidget_raster::RenderSize; use heck::ToShoutySnakeCase; -use std::collections::BTreeMap; use zerocopy::{FromBytes, Immutable, IntoBytes}; pub mod buf; @@ -358,15 +357,20 @@ impl Gpu { //////////////////////////////////////////////////////////////////////////////// /// Container of multiple pipelines, parameterized by register count -pub(crate) struct RegPipeline(BTreeMap); +pub(crate) struct RegPipeline { + thunk: Box wgpu::ComputePipeline>, + cache: + [std::cell::OnceCell; REG_PIPELINE_SIZES.len()], +} + +const REG_PIPELINE_SIZES: &[u8] = &[8, 16, 32, 64, 128, 192, 255]; impl RegPipeline { - pub fn build wgpu::ComputePipeline>(builder: F) -> Self { - let mut out = BTreeMap::new(); - for reg_count in [8, 16, 32, 64, 128, 192, 255] { - out.insert(reg_count, builder(reg_count)); + pub fn build(builder: Box wgpu::ComputePipeline>) -> Self { + Self { + thunk: builder, + cache: std::array::from_fn(|_| Default::default()), } - Self(out) } /// Returns the pipeline with sufficient registers to render `reg_count` @@ -374,13 +378,12 @@ impl RegPipeline { /// # Panics /// If `reg_count` is 256 (which is not allowed in bytecode tapes) pub fn get(&self, reg_count: u8) -> &wgpu::ComputePipeline { - let (r, v) = self - .0 - .range(reg_count..) - .next() + let (i, r) = REG_PIPELINE_SIZES + .iter() + .enumerate() + .find(|(_i, r)| reg_count <= **r) .expect("bytecode tape cannot use more than 255 registers"); - assert!(*r >= reg_count); - v + self.cache[i].get_or_init(|| (*self.thunk)(*r)) } } diff --git a/fidget-wgpu/src/pixel/effects/mod.rs b/fidget-wgpu/src/pixel/effects/mod.rs index de7b3a20..d647dcd1 100644 --- a/fidget-wgpu/src/pixel/effects/mod.rs +++ b/fidget-wgpu/src/pixel/effects/mod.rs @@ -475,14 +475,15 @@ impl ColorContext { ], immediate_size: 0u32, }); - let color_pipeline = RegPipeline::build(|reg_count| { + let device_ = device.clone(); + let color_pipeline = RegPipeline::build(Box::new(move |reg_count| { let shader_code = color_shader(reg_count); let shader_module = - device.create_shader_module(wgpu::ShaderModuleDescriptor { + device_.create_shader_module(wgpu::ShaderModuleDescriptor { label: None, source: wgpu::ShaderSource::Wgsl(shader_code.into()), }); - device.create_compute_pipeline(&wgpu::ComputePipelineDescriptor { + device_.create_compute_pipeline(&wgpu::ComputePipelineDescriptor { label: Some(&format!("color ({reg_count})")), layout: Some(&pipeline_layout), module: &shader_module, @@ -490,7 +491,7 @@ impl ColorContext { compilation_options: Default::default(), cache: None, }) - }); + })); Self { config_bind_group_layout, diff --git a/fidget-wgpu/src/pixel/mod.rs b/fidget-wgpu/src/pixel/mod.rs index be3e9eec..48a98650 100644 --- a/fidget-wgpu/src/pixel/mod.rs +++ b/fidget-wgpu/src/pixel/mod.rs @@ -162,14 +162,15 @@ impl RootContext { immediate_size: 0u32, }); - let root_pipeline = RegPipeline::build(|reg_count| { + let device_ = device.clone(); + let root_pipeline = RegPipeline::build(Box::new(move |reg_count| { let shader_code = interval_root_shader(reg_count); let shader_module = - device.create_shader_module(wgpu::ShaderModuleDescriptor { + device_.create_shader_module(wgpu::ShaderModuleDescriptor { label: None, source: wgpu::ShaderSource::Wgsl(shader_code.into()), }); - device.create_compute_pipeline(&wgpu::ComputePipelineDescriptor { + device_.create_compute_pipeline(&wgpu::ComputePipelineDescriptor { label: Some(&format!("interval root ({reg_count})")), layout: Some(&pipeline_layout), module: &shader_module, @@ -177,7 +178,7 @@ impl RootContext { compilation_options: Default::default(), cache: None, }) - }); + })); Self { bind_group_layout, @@ -239,14 +240,15 @@ impl IntervalTilesContext { immediate_size: 0u32, }); - let tiles_pipeline = RegPipeline::build(|reg_count| { + let device_ = device.clone(); + let tiles_pipeline = RegPipeline::build(Box::new(move |reg_count| { let shader_code = interval_tiles_shader(reg_count); let shader_module = - device.create_shader_module(wgpu::ShaderModuleDescriptor { + device_.create_shader_module(wgpu::ShaderModuleDescriptor { label: None, source: wgpu::ShaderSource::Wgsl(shader_code.into()), }); - device.create_compute_pipeline(&wgpu::ComputePipelineDescriptor { + device_.create_compute_pipeline(&wgpu::ComputePipelineDescriptor { label: Some(&format!("interval tiles ({reg_count})")), layout: Some(&pipeline_layout), module: &shader_module, @@ -254,7 +256,7 @@ impl IntervalTilesContext { compilation_options: Default::default(), cache: None, }) - }); + })); Self { bind_group_layout, @@ -314,14 +316,15 @@ impl PixelTilesContext { immediate_size: 0u32, }); - let tiles_pipeline = RegPipeline::build(|reg_count| { + let device_ = device.clone(); + let tiles_pipeline = RegPipeline::build(Box::new(move |reg_count| { let shader_code = pixel_tiles_shader(reg_count); let shader_module = - device.create_shader_module(wgpu::ShaderModuleDescriptor { + device_.create_shader_module(wgpu::ShaderModuleDescriptor { label: None, source: wgpu::ShaderSource::Wgsl(shader_code.into()), }); - device.create_compute_pipeline(&wgpu::ComputePipelineDescriptor { + device_.create_compute_pipeline(&wgpu::ComputePipelineDescriptor { label: Some(&format!("pixel tiles ({reg_count})")), layout: Some(&pipeline_layout), module: &shader_module, @@ -329,7 +332,7 @@ impl PixelTilesContext { compilation_options: Default::default(), cache: None, }) - }); + })); Self { bind_group_layout, diff --git a/fidget-wgpu/src/voxel/effects/mod.rs b/fidget-wgpu/src/voxel/effects/mod.rs index de04e69f..f61c3cb3 100644 --- a/fidget-wgpu/src/voxel/effects/mod.rs +++ b/fidget-wgpu/src/voxel/effects/mod.rs @@ -1251,14 +1251,15 @@ impl ColorContext { ], immediate_size: 0u32, }); - let color_pipeline = RegPipeline::build(|reg_count| { + let device_ = device.clone(); + let color_pipeline = RegPipeline::build(Box::new(move |reg_count| { let shader_code = color_shader(reg_count); let shader_module = - device.create_shader_module(wgpu::ShaderModuleDescriptor { + device_.create_shader_module(wgpu::ShaderModuleDescriptor { label: None, source: wgpu::ShaderSource::Wgsl(shader_code.into()), }); - device.create_compute_pipeline(&wgpu::ComputePipelineDescriptor { + device_.create_compute_pipeline(&wgpu::ComputePipelineDescriptor { label: Some(&format!("color ({reg_count})")), layout: Some(&pipeline_layout), module: &shader_module, @@ -1266,7 +1267,7 @@ impl ColorContext { compilation_options: Default::default(), cache: None, }) - }); + })); Self { config_bind_group_layout, diff --git a/fidget-wgpu/src/voxel/mod.rs b/fidget-wgpu/src/voxel/mod.rs index 6ea62110..11c4ca05 100644 --- a/fidget-wgpu/src/voxel/mod.rs +++ b/fidget-wgpu/src/voxel/mod.rs @@ -478,14 +478,15 @@ impl RootContext { ], immediate_size: 0u32, }); - let root_pipeline = RegPipeline::build(|reg_count| { + let device_ = device.clone(); + let root_pipeline = RegPipeline::build(Box::new(move |reg_count| { let shader_code = interval_root_shader(reg_count); let shader_module = - device.create_shader_module(wgpu::ShaderModuleDescriptor { + device_.create_shader_module(wgpu::ShaderModuleDescriptor { label: Some("interval root"), source: wgpu::ShaderSource::Wgsl(shader_code.into()), }); - device.create_compute_pipeline(&wgpu::ComputePipelineDescriptor { + device_.create_compute_pipeline(&wgpu::ComputePipelineDescriptor { label: Some(&format!("interval root ({reg_count})")), layout: Some(&pipeline_layout), module: &shader_module, @@ -493,7 +494,7 @@ impl RootContext { compilation_options: Default::default(), cache: None, }) - }); + })); Self { bind_group_layout, @@ -652,71 +653,90 @@ impl IntervalContext { immediate_size: 0u32, }); - let interval64_pipeline = RegPipeline::build(|reg_count| { - let shader_code = interval_tiles_shader(reg_count); - // SAFETY: the shader is carefully written - let shader_module = unsafe { - device.create_shader_module_trusted( - wgpu::ShaderModuleDescriptor { - label: Some(&format!( - "interval64 tiles shader ({reg_count})" - )), - source: wgpu::ShaderSource::Wgsl(shader_code.into()), - }, - wgpu::ShaderRuntimeChecks { - bounds_checks: false, - force_loop_bounding: false, - ray_query_initialization_tracking: false, - task_shader_dispatch_tracking: false, - mesh_shader_primitive_indices_clamp: false, + let device_ = device.clone(); + let interval_pipeline_layout_ = interval_pipeline_layout.clone(); + let interval64_pipeline = + RegPipeline::build(Box::new(move |reg_count| { + let shader_code = interval_tiles_shader(reg_count); + // SAFETY: the shader is carefully written + let shader_module = unsafe { + device_.create_shader_module_trusted( + wgpu::ShaderModuleDescriptor { + label: Some(&format!( + "interval64 tiles shader ({reg_count})" + )), + source: wgpu::ShaderSource::Wgsl( + shader_code.into(), + ), + }, + wgpu::ShaderRuntimeChecks { + bounds_checks: false, + force_loop_bounding: false, + ray_query_initialization_tracking: false, + task_shader_dispatch_tracking: false, + mesh_shader_primitive_indices_clamp: false, + }, + ) + }; + device_.create_compute_pipeline( + &wgpu::ComputePipelineDescriptor { + label: Some(&format!("interval64 ({reg_count})")), + layout: Some(&interval_pipeline_layout_), + module: &shader_module, + entry_point: Some("interval_tile_main"), + compilation_options: wgpu::PipelineCompilationOptions { + constants: &[ + ("TILE_SIZE", 64.0), + ("SUBTILE_SIZE", 16.0), + ], + ..Default::default() + }, + cache: None, }, ) - }; - device.create_compute_pipeline(&wgpu::ComputePipelineDescriptor { - label: Some(&format!("interval64 ({reg_count})")), - layout: Some(&interval_pipeline_layout), - module: &shader_module, - entry_point: Some("interval_tile_main"), - compilation_options: wgpu::PipelineCompilationOptions { - constants: &[("TILE_SIZE", 64.0), ("SUBTILE_SIZE", 16.0)], - ..Default::default() - }, - cache: None, - }) - }); - - let interval16_pipeline = RegPipeline::build(|reg_count| { - let shader_code = interval_tiles_shader(reg_count); - // SAFETY: the shader is carefully written - let shader_module = unsafe { - device.create_shader_module_trusted( - wgpu::ShaderModuleDescriptor { - label: Some(&format!( - "interval16 tiles shader ({reg_count})" - )), - source: wgpu::ShaderSource::Wgsl(shader_code.into()), - }, - wgpu::ShaderRuntimeChecks { - bounds_checks: false, - force_loop_bounding: false, - ray_query_initialization_tracking: false, - task_shader_dispatch_tracking: false, - mesh_shader_primitive_indices_clamp: false, + })); + + let device_ = device.clone(); + let interval16_pipeline = + RegPipeline::build(Box::new(move |reg_count| { + let shader_code = interval_tiles_shader(reg_count); + // SAFETY: the shader is carefully written + let shader_module = unsafe { + device_.create_shader_module_trusted( + wgpu::ShaderModuleDescriptor { + label: Some(&format!( + "interval16 tiles shader ({reg_count})" + )), + source: wgpu::ShaderSource::Wgsl( + shader_code.into(), + ), + }, + wgpu::ShaderRuntimeChecks { + bounds_checks: false, + force_loop_bounding: false, + ray_query_initialization_tracking: false, + task_shader_dispatch_tracking: false, + mesh_shader_primitive_indices_clamp: false, + }, + ) + }; + device_.create_compute_pipeline( + &wgpu::ComputePipelineDescriptor { + label: Some(&format!("interval16 ({reg_count})")), + layout: Some(&interval_pipeline_layout), + module: &shader_module, + entry_point: Some("interval_tile_main"), + compilation_options: wgpu::PipelineCompilationOptions { + constants: &[ + ("TILE_SIZE", 16.0), + ("SUBTILE_SIZE", 4.0), + ], + ..Default::default() + }, + cache: None, }, ) - }; - device.create_compute_pipeline(&wgpu::ComputePipelineDescriptor { - label: Some(&format!("interval16 ({reg_count})")), - layout: Some(&interval_pipeline_layout), - module: &shader_module, - entry_point: Some("interval_tile_main"), - compilation_options: wgpu::PipelineCompilationOptions { - constants: &[("TILE_SIZE", 16.0), ("SUBTILE_SIZE", 4.0)], - ..Default::default() - }, - cache: None, - }) - }); + })); let sort_bind_group_layout = device.create_bind_group_layout(&wgpu::BindGroupLayoutDescriptor { @@ -867,11 +887,12 @@ impl VoxelContext { ], immediate_size: 0u32, }); - let voxel_pipeline = RegPipeline::build(|reg_count| { + let device_ = device.clone(); + let voxel_pipeline = RegPipeline::build(Box::new(move |reg_count| { let shader_code = voxel_tiles_shader(reg_count); // SAFETY: The shader is careful, good luck let shader_module = unsafe { - device.create_shader_module_trusted( + device_.create_shader_module_trusted( wgpu::ShaderModuleDescriptor { label: Some("voxel shader module"), source: wgpu::ShaderSource::Wgsl(shader_code.into()), @@ -885,7 +906,7 @@ impl VoxelContext { }, ) }; - device.create_compute_pipeline(&wgpu::ComputePipelineDescriptor { + device_.create_compute_pipeline(&wgpu::ComputePipelineDescriptor { label: Some(&format!("voxels ({reg_count})")), layout: Some(&pipeline_layout), module: &shader_module, @@ -893,7 +914,7 @@ impl VoxelContext { compilation_options: Default::default(), cache: None, }) - }); + })); Self { bind_group_layout, @@ -951,14 +972,15 @@ impl NormalsContext { ], immediate_size: 0u32, }); - let normals_pipeline = RegPipeline::build(|reg_count| { + let device_ = device.clone(); + let normals_pipeline = RegPipeline::build(Box::new(move |reg_count| { let shader_code = normals_shader(reg_count); let shader_module = - device.create_shader_module(wgpu::ShaderModuleDescriptor { + device_.create_shader_module(wgpu::ShaderModuleDescriptor { label: Some("normals shader module"), source: wgpu::ShaderSource::Wgsl(shader_code.into()), }); - device.create_compute_pipeline(&wgpu::ComputePipelineDescriptor { + device_.create_compute_pipeline(&wgpu::ComputePipelineDescriptor { label: Some(&format!("normals ({reg_count})")), layout: Some(&pipeline_layout), module: &shader_module, @@ -966,7 +988,7 @@ impl NormalsContext { compilation_options: Default::default(), cache: None, }) - }); + })); Self { bind_group_layout,