//! TRACES: FR-DEV-3g //! Grain back into a denoised photograph, as brightness only. //! //! The learned denoise's one live control. The network's result and the //! classical demosaic of the same mosaic differ by the noise the network //! removed — plus the classical path's colour speckle and demosaic false //! colour, which nobody wants back. So only the brightness of the difference //! is returned, in proportion to `grain`: //! //! `out = denoised + grain · ΔY / wb`, with `ΔY = Y(wb · (classical − denoised))` //! //! `Y` is taken after the as-shot balance and handed back divided by it, so //! the grain is neutral in the finished picture rather than tinted the //! colour of the sensor's raw response. At 0 the result is the network's //! exactly; at 1 the brightness noise is all back, the colour noise none. //! //! A pass of its own producing a new source rather than a term in the //! adjust shader: the blend depends only on the two images and one number, //! a 20 MP pass is a few milliseconds, and a new source id is all the //! adjust pass's caches need to know it changed. use std::sync::Arc; use crate::demosaic::DemosaicedImage; use crate::{GpuContext, GpuError}; const SHADER: &str = r#" struct Params { grain: f32, _pad0: f32, _pad1: f32, _pad2: f32, wb: vec4, } @group(0) @binding(0) var denoised: texture_2d; @group(0) @binding(1) var classical: texture_2d; @group(0) @binding(2) var p: Params; @group(0) @binding(3) var out: texture_storage_2d; @compute @workgroup_size(8, 8) fn main(@builtin(global_invocation_id) gid: vec3) { let dims = textureDimensions(denoised); if (gid.x >= dims.x || gid.y >= dims.y) { return; } let xy = vec2(gid.xy); let d = textureLoad(denoised, xy, 0).rgb; let c = textureLoad(classical, xy, 0).rgb; let wb = p.wb.rgb; let dy = p.grain * dot(vec3(0.2126, 0.7152, 0.0722), wb * (c - d)); textureStore(out, xy, vec4(d + dy / wb, 1.0)); } "#; #[repr(C)] #[derive(Copy, Clone, bytemuck::Pod, bytemuck::Zeroable)] struct Params { grain: f32, _pad: [f32; 3], wb: [f32; 4], } pub struct GrainBlend { ctx: GpuContext, pipeline: wgpu::ComputePipeline, layout: wgpu::BindGroupLayout, } impl GrainBlend { pub fn new(ctx: &GpuContext) -> Self { let device = &ctx.device; let module = device.create_shader_module(wgpu::ShaderModuleDescriptor { label: Some("grain-blend"), source: wgpu::ShaderSource::Wgsl(SHADER.into()), }); let texture = |binding| wgpu::BindGroupLayoutEntry { binding, visibility: wgpu::ShaderStages::COMPUTE, ty: wgpu::BindingType::Texture { sample_type: wgpu::TextureSampleType::Float { filterable: false }, view_dimension: wgpu::TextureViewDimension::D2, multisampled: false, }, count: None, }; let layout = device.create_bind_group_layout(&wgpu::BindGroupLayoutDescriptor { label: Some("grain-blend-layout"), entries: &[ texture(0), texture(1), wgpu::BindGroupLayoutEntry { binding: 2, visibility: wgpu::ShaderStages::COMPUTE, ty: wgpu::BindingType::Buffer { ty: wgpu::BufferBindingType::Uniform, has_dynamic_offset: false, min_binding_size: None, }, count: None, }, wgpu::BindGroupLayoutEntry { binding: 3, visibility: wgpu::ShaderStages::COMPUTE, ty: wgpu::BindingType::StorageTexture { access: wgpu::StorageTextureAccess::WriteOnly, format: DemosaicedImage::FORMAT, view_dimension: wgpu::TextureViewDimension::D2, }, count: None, }, ], }); let pipeline_layout = device.create_pipeline_layout(&wgpu::PipelineLayoutDescriptor { label: Some("grain-blend-pipeline-layout"), bind_group_layouts: &[Some(&layout)], immediate_size: 0, }); let pipeline = device.create_compute_pipeline(&wgpu::ComputePipelineDescriptor { label: Some("grain-blend"), layout: Some(&pipeline_layout), module: &module, entry_point: Some("main"), compilation_options: Default::default(), cache: None, }); Self { ctx: ctx.clone(), pipeline, layout, } } /// `denoised` with `grain` (0–1) of `classical`'s brightness noise back. /// Both must be the same photograph at the same size. pub fn blend( &self, denoised: &DemosaicedImage, classical: &DemosaicedImage, grain: f32, ) -> Result, GpuError> { let (w, h) = (denoised.texture().width(), denoised.texture().height()); if (classical.texture().width(), classical.texture().height()) != (w, h) { return Err(GpuError::TooLarge(format!( "grain from a {}×{} source into a {w}×{h} one", classical.texture().width(), classical.texture().height() ))); } let device = &self.ctx.device; let texture = device.create_texture(&wgpu::TextureDescriptor { label: Some("grain-blended-source"), size: wgpu::Extent3d { width: w, height: h, depth_or_array_layers: 1, }, mip_level_count: 1, sample_count: 1, dimension: wgpu::TextureDimension::D2, format: DemosaicedImage::FORMAT, usage: wgpu::TextureUsages::STORAGE_BINDING | wgpu::TextureUsages::TEXTURE_BINDING | wgpu::TextureUsages::COPY_SRC, view_formats: &[], }); let out_view = texture.create_view(&Default::default()); let wb = denoised.as_shot_wb(); let g = wb[1].max(1e-6); let params = Params { grain: grain.clamp(0.0, 1.0), _pad: [0.0; 3], // Green-normalised, and never zero: the shader divides by it. wb: [(wb[0] / g).max(1e-3), 1.0, (wb[2] / g).max(1e-3), 1.0], }; use wgpu::util::DeviceExt; let buffer = device.create_buffer_init(&wgpu::util::BufferInitDescriptor { label: Some("grain-blend-params"), contents: bytemuck::bytes_of(¶ms), usage: wgpu::BufferUsages::UNIFORM, }); let bind = device.create_bind_group(&wgpu::BindGroupDescriptor { label: Some("grain-blend-bg"), layout: &self.layout, entries: &[ wgpu::BindGroupEntry { binding: 0, resource: wgpu::BindingResource::TextureView(denoised.view()), }, wgpu::BindGroupEntry { binding: 1, resource: wgpu::BindingResource::TextureView(classical.view()), }, wgpu::BindGroupEntry { binding: 2, resource: buffer.as_entire_binding(), }, wgpu::BindGroupEntry { binding: 3, resource: wgpu::BindingResource::TextureView(&out_view), }, ], }); let mut enc = device.create_command_encoder(&wgpu::CommandEncoderDescriptor { label: Some("grain-blend"), }); { let mut pass = enc.begin_compute_pass(&wgpu::ComputePassDescriptor { label: Some("grain-blend"), timestamp_writes: None, }); pass.set_pipeline(&self.pipeline); pass.set_bind_group(0, &bind, &[]); pass.dispatch_workgroups(w.div_ceil(8), h.div_ceil(8), 1); } self.ctx.queue.submit(Some(enc.finish())); Ok(Arc::new(denoised.sibling(texture, w, h))) } }