GPU: add support for N-many input images and not just exactly one

This commit is contained in:
firestar99
2026-06-07 20:39:30 -07:00
committed by Keavon Chambers
parent 7a6e4f6f25
commit 953a9312f0
3 changed files with 94 additions and 74 deletions
@@ -33,7 +33,8 @@ impl PerPixelAdjustShaderRuntime {
} }
impl ShaderRuntime { impl ShaderRuntime {
pub async fn run_per_pixel_adjust<T: BufferStruct>(&self, shaders: &Shaders<'_>, textures: List<Raster<GPU>>, args: Option<&T>) -> List<Raster<GPU>> { pub async fn run_per_pixel_adjust<T: BufferStruct>(&self, shaders: &Shaders<'_>, textures: &[List<Raster<GPU>>], args: Option<&T>) -> List<Raster<GPU>> {
assert_eq!(shaders.input_images, textures.len());
let mut cache = self.per_pixel_adjust.pipeline_cache.lock().await; let mut cache = self.per_pixel_adjust.pipeline_cache.lock().await;
let pipeline = cache let pipeline = cache
.entry(shaders.fragment_shader_name.to_owned()) .entry(shaders.fragment_shader_name.to_owned())
@@ -54,11 +55,13 @@ impl ShaderRuntime {
pub struct Shaders<'a> { pub struct Shaders<'a> {
pub wgsl_shader: &'a str, pub wgsl_shader: &'a str,
pub fragment_shader_name: &'a str, pub fragment_shader_name: &'a str,
pub input_images: usize,
pub has_uniform: bool, pub has_uniform: bool,
} }
pub struct PerPixelAdjustGraphicsPipeline { pub struct PerPixelAdjustGraphicsPipeline {
name: String, name: String,
input_images: usize,
has_uniform: bool, has_uniform: bool,
pipeline: wgpu::RenderPipeline, pipeline: wgpu::RenderPipeline,
} }
@@ -76,32 +79,23 @@ impl PerPixelAdjustGraphicsPipeline {
source: ShaderSource::Wgsl(Cow::Borrowed(info.wgsl_shader)), source: ShaderSource::Wgsl(Cow::Borrowed(info.wgsl_shader)),
}); });
let entries: &[_] = if info.has_uniform { let mut binding_alloc = Counter::default();
&[ let mut entries = Vec::new();
BindGroupLayoutEntry { if info.has_uniform {
binding: 0, entries.push(BindGroupLayoutEntry {
visibility: ShaderStages::FRAGMENT, binding: binding_alloc.alloc(),
ty: BindingType::Buffer { visibility: ShaderStages::FRAGMENT,
ty: BufferBindingType::Storage { read_only: true }, ty: BindingType::Buffer {
has_dynamic_offset: false, ty: BufferBindingType::Storage { read_only: true },
min_binding_size: None, has_dynamic_offset: false,
}, min_binding_size: None,
count: None,
}, },
BindGroupLayoutEntry { count: None,
binding: 1, });
visibility: ShaderStages::FRAGMENT, }
ty: BindingType::Texture { for _ in 0..info.input_images {
sample_type: TextureSampleType::Float { filterable: false }, entries.push(BindGroupLayoutEntry {
view_dimension: TextureViewDimension::D2, binding: binding_alloc.alloc(),
multisampled: false,
},
count: None,
},
]
} else {
&[BindGroupLayoutEntry {
binding: 0,
visibility: ShaderStages::FRAGMENT, visibility: ShaderStages::FRAGMENT,
ty: BindingType::Texture { ty: BindingType::Texture {
sample_type: TextureSampleType::Float { filterable: false }, sample_type: TextureSampleType::Float { filterable: false },
@@ -109,13 +103,13 @@ impl PerPixelAdjustGraphicsPipeline {
multisampled: false, multisampled: false,
}, },
count: None, count: None,
}] });
}; }
let pipeline_layout = device.create_pipeline_layout(&PipelineLayoutDescriptor { let pipeline_layout = device.create_pipeline_layout(&PipelineLayoutDescriptor {
label: Some(&format!("PerPixelAdjust {name} PipelineLayout")), label: Some(&format!("PerPixelAdjust {name} PipelineLayout")),
bind_group_layouts: &[Some(&device.create_bind_group_layout(&BindGroupLayoutDescriptor { bind_group_layouts: &[Some(&device.create_bind_group_layout(&BindGroupLayoutDescriptor {
label: Some(&format!("PerPixelAdjust {name} BindGroupLayout 0")), label: Some(&format!("PerPixelAdjust {name} BindGroupLayout 0")),
entries, entries: &entries,
}))], }))],
..Default::default() ..Default::default()
}); });
@@ -157,61 +151,76 @@ impl PerPixelAdjustGraphicsPipeline {
pipeline, pipeline,
name, name,
has_uniform: info.has_uniform, has_uniform: info.has_uniform,
input_images: info.input_images,
} }
} }
pub fn dispatch(&self, context: &WgpuContext, textures: List<Raster<GPU>>, arg_buffer: Option<Buffer>) -> List<Raster<GPU>> { pub fn dispatch(&self, context: &WgpuContext, in_textures: &[List<Raster<GPU>>], arg_buffer: Option<Buffer>) -> List<Raster<GPU>> {
assert_eq!(self.has_uniform, arg_buffer.is_some()); assert_eq!(self.has_uniform, arg_buffer.is_some());
assert_eq!(self.input_images, in_textures.len());
let device = &context.device; let device = &context.device;
let name = self.name.as_str(); let name = self.name.as_str();
// Assumption: when we have multiple input images to our node, each input's List of images can have a different
// length. Only process the minimum between all input images, same as `impl Blend<Color> for List<Raster<CPU>>`.
let dispatch_cnt = match in_textures.iter().map(|t| t.len()).min() {
None => {
return List::new();
}
Some(e) => e,
};
let mut cmd = device.create_command_encoder(&wgpu::CommandEncoderDescriptor { let mut cmd = device.create_command_encoder(&wgpu::CommandEncoderDescriptor {
label: Some(&format!("{name} cmd encoder")), label: Some(&format!("{name} cmd encoder")),
}); });
let out = (0..textures.len()) let out = (0..dispatch_cnt)
.map(|index| { .map(|dispatch_id| {
let element = textures.element(index).unwrap(); let mut binding_alloc = Counter::default();
let tex_in = &element.texture; let mut entries = Vec::new();
let view_in = tex_in.create_view(&TextureViewDescriptor::default()); if let Some(arg_buffer) = arg_buffer.as_ref() {
let format = tex_in.format(); entries.push(BindGroupEntry {
binding: binding_alloc.alloc(),
let entries: &[_] = if let Some(arg_buffer) = arg_buffer.as_ref() { resource: BindingResource::Buffer(BufferBinding {
&[ buffer: arg_buffer,
BindGroupEntry { offset: 0,
binding: 0, size: None,
resource: BindingResource::Buffer(BufferBinding { }),
buffer: arg_buffer, });
offset: 0, }
size: None, let in_texture_views = in_textures
}), .iter()
}, .map(|texture| {
BindGroupEntry { let element = texture.element(dispatch_id).unwrap();
binding: 1, element.texture.create_view(&TextureViewDescriptor::default())
resource: BindingResource::TextureView(&view_in), })
}, .collect::<Vec<_>>();
] for view_in in &in_texture_views {
} else { entries.push(BindGroupEntry {
&[BindGroupEntry { binding: binding_alloc.alloc(),
binding: 0,
resource: BindingResource::TextureView(&view_in), resource: BindingResource::TextureView(&view_in),
}] });
}; }
let bind_group = device.create_bind_group(&BindGroupDescriptor { let bind_group = device.create_bind_group(&BindGroupDescriptor {
label: Some(&format!("{name} bind group")), label: Some(&format!("{name} bind group")),
// `get_bind_group_layout` allocates unnecessary memory, we could create it manually to not do that // `get_bind_group_layout` allocates unnecessary memory, we could create it manually to not do that
layout: &self.pipeline.get_bind_group_layout(0), layout: &self.pipeline.get_bind_group_layout(0),
entries, entries: &entries,
}); });
// Assumption: The output texture has the same size and format as the first input texture. Like the
// blend node, that writes the output directly back into the first texture.
let outref_list = &in_textures[0];
let outref_tex = &outref_list.element(dispatch_id).unwrap().texture;
let tex_out = device.create_texture(&TextureDescriptor { let tex_out = device.create_texture(&TextureDescriptor {
label: Some(&format!("{name} texture out")), label: Some(&format!("{name} texture out")),
size: tex_in.size(), size: outref_tex.size(),
mip_level_count: 1, mip_level_count: 1,
sample_count: 1, sample_count: 1,
dimension: TextureDimension::D2, dimension: TextureDimension::D2,
format, format: outref_tex.format(),
usage: wgpu::TextureUsages::TEXTURE_BINDING | wgpu::TextureUsages::COPY_DST | wgpu::TextureUsages::COPY_SRC | wgpu::TextureUsages::RENDER_ATTACHMENT, usage: wgpu::TextureUsages::TEXTURE_BINDING | wgpu::TextureUsages::COPY_DST | wgpu::TextureUsages::COPY_SRC | wgpu::TextureUsages::RENDER_ATTACHMENT,
view_formats: &[format], view_formats: &[outref_tex.format()],
}); });
let view_out = tex_out.create_view(&TextureViewDescriptor::default()); let view_out = tex_out.create_view(&TextureViewDescriptor::default());
@@ -233,7 +242,7 @@ impl PerPixelAdjustGraphicsPipeline {
rp.set_bind_group(0, Some(&bind_group), &[]); rp.set_bind_group(0, Some(&bind_group), &[]);
rp.draw(0..3, 0..1); rp.draw(0..3, 0..1);
let attributes = textures.clone_item_attributes(index); let attributes = outref_list.clone_item_attributes(dispatch_id);
Item::from_parts(Raster::new(GPU { texture: tex_out }), attributes) Item::from_parts(Raster::new(GPU { texture: tex_out }), attributes)
}) })
.collect::<List<_>>(); .collect::<List<_>>();
@@ -241,3 +250,14 @@ impl PerPixelAdjustGraphicsPipeline {
out out
} }
} }
#[derive(Clone, Debug, Default)]
pub struct Counter(pub u32);
impl Counter {
pub fn alloc(&mut self) -> u32 {
let out = self.0;
self.0 += 1;
out
}
}
@@ -248,16 +248,15 @@ impl PerPixelAdjustCodegen<'_> {
is_data_field: false, is_data_field: false,
}); });
// find exactly one gpu_image field, runtime doesn't support more than 1 atm // find gpu_image fields
let gpu_image_field = { let gpu_images = fields
let mut iter = fields.iter().filter(|f| matches!(f.ty, ParsedFieldType::Regular(RegularParsedField { gpu_image: true, .. }))); .iter()
match (iter.next(), iter.next()) { .filter_map(|f| match f.ty {
(Some(v), None) => Ok(v), ParsedFieldType::Regular(RegularParsedField { gpu_image: true, .. }) => Some(&f.pat_ident.ident),
(Some(_), Some(more)) => Err(syn::Error::new_spanned(&more.pat_ident, "No more than one parameter must be annotated with `#[gpu_image]`")), _ => None,
(None, _) => Err(syn::Error::new_spanned(&self.parsed.fn_name, "At least one parameter must be annotated with `#[gpu_image]`")), })
}? .collect::<Vec<_>>();
}; let input_images = gpu_images.len();
let gpu_image = &gpu_image_field.pat_ident.ident;
// uniform buffer struct construction // uniform buffer struct construction
let has_uniform = self.has_uniform; let has_uniform = self.has_uniform;
@@ -287,7 +286,8 @@ impl PerPixelAdjustCodegen<'_> {
wgsl_shader: crate::WGSL_SHADER, wgsl_shader: crate::WGSL_SHADER,
fragment_shader_name: super::#entry_point_name, fragment_shader_name: super::#entry_point_name,
has_uniform: #has_uniform, has_uniform: #has_uniform,
}, #gpu_image, #uniform_buffer).await input_images: #input_images,
}, &[#(#gpu_images),*], #uniform_buffer).await
} }
}; };
@@ -141,7 +141,7 @@ pub fn apply_blend_mode(foreground: Color, background: Color, blend_mode: BlendM
} }
} }
#[node_macro::node(category("Raster"), cfg(feature = "std"))] #[node_macro::node(category("Raster"), shader_node(PerPixelAdjust))]
fn mix<T: Blend<Color> + Send>( fn mix<T: Blend<Color> + Send>(
_: impl Ctx, _: impl Ctx,
#[implementations( #[implementations(