DemosaicedImage::from_rgb_f32 takes the network's linear camera RGB and stands it beside the classical source of the same photograph: the matrix, profile tables and as-shot balance are that source's, the id is new, so nothing downstream can tell which demosaic ran and every cache keyed on the source sees a new one. GrainBlend is the denoise's live control. It returns only the brightness of the noise the network removed, taken after the as-shot balance and handed back divided by it, so the grain is neutral in the finished picture; colour speckle and demosaic false colour stay out. It writes a new source rather than adding a term to the adjust shader: the blend depends on two images and one number, a 20 MP pass is milliseconds, and a fresh source id is all the adjust pass's caches need. The test reads it back: at 0 the network's result, at 1 the same white-balanced step in every channel.
220 lines
8.2 KiB
Rust
220 lines
8.2 KiB
Rust
//! 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<f32>,
|
||
}
|
||
|
||
@group(0) @binding(0) var denoised: texture_2d<f32>;
|
||
@group(0) @binding(1) var classical: texture_2d<f32>;
|
||
@group(0) @binding(2) var<uniform> p: Params;
|
||
@group(0) @binding(3) var out: texture_storage_2d<rgba16float, write>;
|
||
|
||
@compute @workgroup_size(8, 8)
|
||
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
||
let dims = textureDimensions(denoised);
|
||
if (gid.x >= dims.x || gid.y >= dims.y) {
|
||
return;
|
||
}
|
||
let xy = vec2<i32>(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<f32>(0.2126, 0.7152, 0.0722), wb * (c - d));
|
||
textureStore(out, xy, vec4<f32>(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<Arc<DemosaicedImage>, 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)))
|
||
}
|
||
}
|