Upload the learned demosaic's result, and blend grain back into it

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.
This commit is contained in:
2026-10-03 11:20:48 -04:00
parent d8304d7c82
commit 8ea3c3181a
5 changed files with 470 additions and 15 deletions
+89
View File
@@ -346,6 +346,95 @@ impl DemosaicedImage {
}
impl DemosaicedImage {
/// TRACES: FR-DEV-3g
/// The learned demosaic's output for the photograph `like` was
/// demosaiced from: `width × height` interleaved RGB, linear camera
/// space, normalised as the demosaic normalises — the same texture the
/// classical path made, with the noise gone (denoise.md §2).
///
/// Everything that describes the photograph rather than its pixels —
/// matrix, profile tables, as-shot balance — is `like`'s, so nothing
/// downstream can tell which demosaic ran. A new [`Self::id`], so every
/// cache keyed on the source sees a new source.
pub fn from_rgb_f32(
ctx: &GpuContext,
like: &DemosaicedImage,
width: u32,
height: u32,
rgb: &[f32],
) -> Result<Self, GpuError> {
let n = width as usize * height as usize;
if rgb.len() != n * 3 {
return Err(GpuError::TooLarge(format!(
"{} values for a {width}×{height} RGB image",
rgb.len()
)));
}
let limits = ctx.device.limits();
if width > limits.max_texture_dimension_2d || height > limits.max_texture_dimension_2d {
return Err(GpuError::TooLarge(format!(
"{width}×{height} exceeds the device limit of {}",
limits.max_texture_dimension_2d
)));
}
let mut half = vec![0u16; n * 4];
let one = f32_to_f16_bits(1.0);
let threads = std::thread::available_parallelism().map_or(1, |n| n.get());
let per = n.div_ceil(threads).max(1);
std::thread::scope(|scope| {
for (k, out) in half.chunks_mut(per * 4).enumerate() {
scope.spawn(move || {
for (i, texel) in out.chunks_mut(4).enumerate() {
let src = &rgb[(k * per + i) * 3..(k * per + i) * 3 + 3];
for c in 0..3 {
texel[c] = f32_to_f16_bits_unclamped(src[c]);
}
texel[3] = one;
}
});
}
});
let texture = ctx.device.create_texture_with_data(
&ctx.queue,
&wgpu::TextureDescriptor {
label: Some("learned-demosaic-source"),
size: wgpu::Extent3d {
width,
height,
depth_or_array_layers: 1,
},
mip_level_count: 1,
sample_count: 1,
dimension: wgpu::TextureDimension::D2,
format: Self::FORMAT,
usage: wgpu::TextureUsages::TEXTURE_BINDING | wgpu::TextureUsages::COPY_SRC,
view_formats: &[],
},
wgpu::util::TextureDataOrder::LayerMajor,
bytemuck::cast_slice(&half),
);
Ok(like.sibling(texture, width, height))
}
/// A new source standing for the same photograph as `self`: its
/// description kept, its pixels `texture`, a fresh id.
pub(crate) fn sibling(&self, texture: wgpu::Texture, width: u32, height: u32) -> Self {
let view = texture.create_view(&Default::default());
Self {
texture,
view,
width,
height,
color_matrix: self.color_matrix,
profile_tables: self.profile_tables.clone(),
as_shot_wb: self.as_shot_wb,
non_linear: self.non_linear,
id: next_image_id(),
frame: self.frame,
window: self.window,
}
}
/// TRACES: FR-MRG-3
/// A source that is already RGB in camera space: a linear DNG, which is
/// what a merge writes. No demosaic; the samples are normalised by the
+219
View File
@@ -0,0 +1,219 @@
//! 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(&params),
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)))
}
}
+2
View File
@@ -26,6 +26,7 @@ mod demosaic;
mod detail;
mod error;
mod focus;
mod grain;
mod histogram;
mod mask;
mod merge;
@@ -41,6 +42,7 @@ pub use demosaic::{DemosaicedImage, Demosaicer};
pub use detail::INTERMEDIATE_FORMAT as DETAIL_INTERMEDIATE_FORMAT;
pub use error::GpuError;
pub use focus::{FocusPeakPass, FocusPeaking, PeakColour, PeakSensitivity};
pub use grain::GrainBlend;
pub use merge::{Band, MergeFrame, MergeOutput, MergePass};
// Renamed on the way out: `BINS` says enough inside `histogram`, and nothing
// at all at a crate root shared with demosaic and segmentation.
+145
View File
@@ -0,0 +1,145 @@
//! TRACES: FR-DEV-3g
//! The grain blend, read back off the device.
use dr_decode::{CfaPattern, CropRect, RawImage};
use dr_gpu::{DemosaicedImage, Demosaicer, GpuContext, GrainBlend};
const W: u32 = 16;
const H: u32 = 8;
fn ctx() -> Option<GpuContext> {
pollster::block_on(GpuContext::new_headless()).ok()
}
/// A photograph to stand the uploads beside: its as-shot balance is what
/// the grain is made neutral under.
fn like(ctx: &GpuContext) -> DemosaicedImage {
let raw = RawImage {
width: W,
height: H,
data: vec![400; (W * H) as usize],
cfa_pattern: CfaPattern::Rggb,
black_level: [0; 4],
white_level: 4095,
wb_coeffs: [2.0, 1.0, 1.5, 1.0],
color_matrix: Some([1.0, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0, 0.0, 1.0]),
samples_per_pixel: 1,
profile: None,
profile_tables: None,
make: String::new(),
model: String::new(),
crop: CropRect {
x: 0,
y: 0,
width: W,
height: H,
},
};
Demosaicer::new(ctx).unwrap().run(&raw).unwrap()
}
fn read(ctx: &GpuContext, img: &DemosaicedImage) -> Vec<[f32; 4]> {
let (w, h) = (img.texture().width(), img.texture().height());
let padded =
(w * 8).div_ceil(wgpu::COPY_BYTES_PER_ROW_ALIGNMENT) * wgpu::COPY_BYTES_PER_ROW_ALIGNMENT;
let buf = ctx.device.create_buffer(&wgpu::BufferDescriptor {
label: None,
size: (padded * h) as u64,
usage: wgpu::BufferUsages::COPY_DST | wgpu::BufferUsages::MAP_READ,
mapped_at_creation: false,
});
let mut enc = ctx.device.create_command_encoder(&Default::default());
enc.copy_texture_to_buffer(
img.texture().as_image_copy(),
wgpu::TexelCopyBufferInfo {
buffer: &buf,
layout: wgpu::TexelCopyBufferLayout {
offset: 0,
bytes_per_row: Some(padded),
rows_per_image: Some(h),
},
},
wgpu::Extent3d {
width: w,
height: h,
depth_or_array_layers: 1,
},
);
ctx.queue.submit(Some(enc.finish()));
let slice = buf.slice(..);
slice.map_async(wgpu::MapMode::Read, |_| {});
ctx.device
.poll(wgpu::PollType::wait_indefinitely())
.unwrap();
let bytes = slice.get_mapped_range();
let mut out = Vec::new();
for y in 0..h as usize {
let row: &[u16] =
bytemuck::cast_slice(&bytes[y * padded as usize..y * padded as usize + w as usize * 8]);
for t in row.chunks(4) {
out.push([0, 1, 2, 3].map(|c| half_to_f32(t[c])));
}
}
out
}
fn half_to_f32(h: u16) -> f32 {
let s = if h & 0x8000 != 0 { -1.0 } else { 1.0 };
let e = ((h >> 10) & 0x1f) as i32;
let m = (h & 0x3ff) as f32;
if e == 0 {
s * m * 2f32.powi(-24)
} else {
s * (1.0 + m / 1024.0) * 2f32.powi(e - 15)
}
}
#[test]
fn grain_returns_only_neutral_brightness() {
let Some(ctx) = ctx() else {
eprintln!("no GPU adapter; skipping");
return;
};
let base = like(&ctx);
let n = (W * H) as usize;
let d: Vec<f32> = (0..n).flat_map(|_| [0.20, 0.30, 0.10]).collect();
// The classical result: the same colour plus noise, coloured noise too.
let c: Vec<f32> = (0..n)
.flat_map(|i| {
let a = ((i * 37) % 11) as f32 / 110.0 - 0.05;
let b = ((i * 53) % 7) as f32 / 140.0 - 0.025;
[0.20 + a, 0.30 + b, 0.10 - a]
})
.collect();
let denoised = DemosaicedImage::from_rgb_f32(&ctx, &base, W, H, &d).unwrap();
let classical = DemosaicedImage::from_rgb_f32(&ctx, &base, W, H, &c).unwrap();
let blend = GrainBlend::new(&ctx);
let wb = [2.0f32, 1.0, 1.5];
let none = read(&ctx, &blend.blend(&denoised, &classical, 0.0).unwrap());
for p in &none {
for ch in 0..3 {
assert!(
(p[ch] - d[ch]).abs() < 1e-3,
"grain 0 must be the network's result: {p:?}"
);
}
}
let all = read(&ctx, &blend.blend(&denoised, &classical, 1.0).unwrap());
for (i, p) in all.iter().enumerate() {
let want_dy: f32 = [0.2126f32, 0.7152, 0.0722]
.iter()
.enumerate()
.map(|(ch, k)| k * wb[ch] * (c[i * 3 + ch] - d[ch]))
.sum();
// After white balance every channel moved by the same amount.
for ch in 0..3 {
let moved = wb[ch] * (p[ch] - d[ch]);
assert!(
(moved - want_dy).abs() < 2e-3,
"pixel {i} channel {ch}: moved {moved}, want {want_dy}"
);
}
}
}
File diff suppressed because one or more lines are too long