//! Watershed segmentation — arm A's GPU half (S15, docs/segmentation.md). //! //! Runs the five passes in `shaders/watershed.wgsl` over a demosaiced image //! and leaves a basin label per pixel on the GPU. The hierarchy built from //! those labels lives in [`crate::hierarchy`], which needs no device. //! //! # Cost //! //! Every pass is a trivial kernel and the whole chain is a handful of //! milliseconds at proxy resolution. It runs **once per image**, off the //! interactive path — the point of precomputing a region map is that //! selection afterwards is a label comparison rather than a flood fill. //! //! # The open question this leaves //! //! [`Segmentation::read_field`] copies the label and gradient buffers back to //! the CPU to build the region adjacency graph, and is gated behind the //! `readback` feature for the same reason `read_pixels` is. That gate is not //! ceremony: a shipping build cannot take this path (ARCH §6.1, AC-8), so the //! RAG would have to be accumulated GPU-side with atomics instead. //! //! For a spike that trade is the right way round — the readback is once per //! image and off the frame path, and building the GPU-side RAG before knowing //! whether the granularity ladder is any good would be work spent on a //! question not yet asked. But it is a real gap between this and something //! shippable, and it should be read as one. use wgpu::util::DeviceExt; use crate::{DemosaicedImage, GpuContext, GpuError}; /// How the watershed is tuned for one image. #[derive(Debug, Clone, Copy, PartialEq)] pub struct SegmentOptions { /// Longest proxy edge. The segmentation runs here, not at sensor /// resolution: a 24 MP watershed costs 12× the memory to place boundaries /// a person cannot see, and the boundary refinement that matters at 1:1 /// is a separate stage (docs/segmentation.md §4). pub max_edge: u32, /// Pre-smoothing radius in proxy pixels. The caller's to raise with ISO — /// this is the single knob that decides whether a noisy file segments /// into regions or into grain. pub blur_radius: i32, pub w_luma: f32, pub w_chroma: f32, } impl Default for SegmentOptions { fn default() -> Self { Self { // ~1.3 MP at 3:2. Large enough that a boundary is within a pixel // or two of where it belongs, small enough that the whole chain // fits comfortably in memory on a phone. max_edge: 1600, blur_radius: 2, w_luma: 1.0, // Chroma carries most of the sensor noise and few of the // boundaries anyone would draw, so it counts for less — but not // zero, or a red flower on green leaves has no edge at all. w_chroma: 0.5, } } } #[repr(C)] #[derive(Copy, Clone, bytemuck::Pod, bytemuck::Zeroable)] struct SegParams { width: u32, height: u32, src_width: u32, src_height: u32, blur_radius: i32, non_linear: u32, w_luma: f32, w_chroma: f32, } /// One compute stage: its layout and its compiled pipeline. struct Stage { layout: wgpu::BindGroupLayout, pipeline: wgpu::ComputePipeline, } /// Runs the watershed chain. pub struct SegmentPass { ctx: GpuContext, features: Stage, blur: Stage, gradient: Stage, flow: Stage, jump: Stage, } impl SegmentPass { pub fn new(ctx: &GpuContext) -> Result { // A validation failure here is a bug in the shader, not a user error. // Surfaced as a Result rather than wgpu's default panic, matching how // `AdjustPass` handles its generated source. let scope = ctx.device.push_error_scope(wgpu::ErrorFilter::Validation); let module = ctx .device .create_shader_module(wgpu::ShaderModuleDescriptor { label: Some("watershed"), source: wgpu::ShaderSource::Wgsl(include_str!("shaders/watershed.wgsl").into()), }); let features = { let layout = ctx .device .create_bind_group_layout(&wgpu::BindGroupLayoutDescriptor { label: Some("watershed-features-bgl"), entries: &[ uniform_entry(0), wgpu::BindGroupLayoutEntry { binding: 1, visibility: wgpu::ShaderStages::COMPUTE, ty: wgpu::BindingType::Texture { sample_type: wgpu::TextureSampleType::Float { filterable: true }, view_dimension: wgpu::TextureViewDimension::D2, multisampled: false, }, count: None, }, storage_entry(2, false), ], }); let pipeline = compute(ctx, &module, &layout, "features"); Stage { layout, pipeline } }; let buffer_stage = |in_binding: u32, out_binding: u32, entry: &str, label: &str| { let layout = ctx .device .create_bind_group_layout(&wgpu::BindGroupLayoutDescriptor { label: Some(label), entries: &[ uniform_entry(0), storage_entry(in_binding, true), storage_entry(out_binding, false), ], }); let pipeline = compute(ctx, &module, &layout, entry); Stage { layout, pipeline } }; let blur = buffer_stage(3, 4, "blur", "watershed-blur-bgl"); let gradient = buffer_stage(5, 6, "gradient", "watershed-gradient-bgl"); let flow = buffer_stage(7, 8, "flow", "watershed-flow-bgl"); let jump = buffer_stage(9, 10, "jump", "watershed-jump-bgl"); if let Some(err) = pollster::block_on(scope.pop()) { return Err(GpuError::ShaderCompilation(err.to_string())); } Ok(Self { ctx: ctx.clone(), features, blur, gradient, flow, jump, }) } /// Segment an image into basins. pub fn run( &self, source: &DemosaicedImage, opts: SegmentOptions, ) -> Result { let (src_w, src_h) = source.size(); let (width, height) = proxy_size(src_w, src_h, opts.max_edge); let n = (width * height) as u64; let params = self .ctx .device .create_buffer_init(&wgpu::util::BufferInitDescriptor { label: Some("watershed-params"), contents: bytemuck::bytes_of(&SegParams { width, height, src_width: src_w, src_height: src_h, blur_radius: opts.blur_radius, non_linear: u32::from(source.is_non_linear()), w_luma: opts.w_luma, w_chroma: opts.w_chroma, }), usage: wgpu::BufferUsages::UNIFORM, }); // `vec4` rather than `vec3` for the feature buffers: a WGSL storage // array of vec3 still strides by 16 bytes, so packing to three floats // would save nothing and cost an index calculation. let feat_a = self.buffer("watershed-feat-a", n * 16, false); let feat_b = self.buffer("watershed-feat-b", n * 16, false); let gradient = self.buffer("watershed-gradient", n * 4, true); let parent_a = self.buffer("watershed-parent-a", n * 4, true); let parent_b = self.buffer("watershed-parent-b", n * 4, true); let mut enc = self .ctx .device .create_command_encoder(&wgpu::CommandEncoderDescriptor { label: Some("watershed-encoder"), }); let groups = (width.div_ceil(8), height.div_ceil(8)); let features_bg = self .ctx .device .create_bind_group(&wgpu::BindGroupDescriptor { label: Some("watershed-features-bg"), layout: &self.features.layout, entries: &[ wgpu::BindGroupEntry { binding: 0, resource: params.as_entire_binding(), }, wgpu::BindGroupEntry { binding: 1, resource: wgpu::BindingResource::TextureView(source.view()), }, wgpu::BindGroupEntry { binding: 2, resource: feat_a.as_entire_binding(), }, ], }); let blur_bg = self.bind(&self.blur.layout, ¶ms, 3, &feat_a, 4, &feat_b); let gradient_bg = self.bind(&self.gradient.layout, ¶ms, 5, &feat_b, 6, &gradient); let flow_bg = self.bind(&self.flow.layout, ¶ms, 7, &gradient, 8, &parent_a); let jump_ab = self.bind(&self.jump.layout, ¶ms, 9, &parent_a, 10, &parent_b); let jump_ba = self.bind(&self.jump.layout, ¶ms, 9, &parent_b, 10, &parent_a); // Pointer jumping halves every path per pass, so log2 of the pixel // count bounds it — that is the longest possible descent chain. A // convergence test would cost a readback per iteration to save a // handful of dispatches of a two-line kernel. let jumps = (n as f64).log2().ceil() as u32 + 1; { let mut pass = enc.begin_compute_pass(&wgpu::ComputePassDescriptor { label: Some("watershed-pass"), timestamp_writes: None, }); for (pipeline, bg) in [ (&self.features.pipeline, &features_bg), (&self.blur.pipeline, &blur_bg), (&self.gradient.pipeline, &gradient_bg), (&self.flow.pipeline, &flow_bg), ] { pass.set_pipeline(pipeline); pass.set_bind_group(0, bg, &[]); pass.dispatch_workgroups(groups.0, groups.1, 1); } pass.set_pipeline(&self.jump.pipeline); for i in 0..jumps { let bg = if i % 2 == 0 { &jump_ab } else { &jump_ba }; pass.set_bind_group(0, bg, &[]); pass.dispatch_workgroups(groups.0, groups.1, 1); } } self.ctx.queue.submit(Some(enc.finish())); // An odd number of jumps leaves the result in B. let labels = if jumps % 2 == 1 { parent_b } else { parent_a }; Ok(Segmentation { ctx: self.ctx.clone(), width, height, labels, gradient, }) } fn buffer(&self, label: &str, size: u64, copyable: bool) -> wgpu::Buffer { let mut usage = wgpu::BufferUsages::STORAGE; if copyable { usage |= wgpu::BufferUsages::COPY_SRC; } self.ctx.device.create_buffer(&wgpu::BufferDescriptor { label: Some(label), size, usage, mapped_at_creation: false, }) } fn bind( &self, layout: &wgpu::BindGroupLayout, params: &wgpu::Buffer, in_binding: u32, input: &wgpu::Buffer, out_binding: u32, output: &wgpu::Buffer, ) -> wgpu::BindGroup { self.ctx .device .create_bind_group(&wgpu::BindGroupDescriptor { label: Some("watershed-bg"), layout, entries: &[ wgpu::BindGroupEntry { binding: 0, resource: params.as_entire_binding(), }, wgpu::BindGroupEntry { binding: in_binding, resource: input.as_entire_binding(), }, wgpu::BindGroupEntry { binding: out_binding, resource: output.as_entire_binding(), }, ], }) } } /// The result of one segmentation: a basin label per pixel, on the GPU. pub struct Segmentation { /// Only [`Self::read_field`] reads this, so a build without `readback` /// carries it unread. That is now the ordinary build: dr-ui used to turn /// the feature on for the whole workspace and stopped when S1 removed the /// display readback, which is what made the field look dead. #[cfg_attr(not(any(test, feature = "readback")), allow(dead_code))] ctx: GpuContext, width: u32, height: u32, /// Per pixel, the linear index of its basin root. Sparse — compacted by /// [`crate::hierarchy::RegionField::from_roots`]. labels: wgpu::Buffer, /// As with `ctx` above: read only by [`Self::read_field`]. #[cfg_attr(not(any(test, feature = "readback")), allow(dead_code))] gradient: wgpu::Buffer, } impl Segmentation { pub fn size(&self) -> (u32, u32) { (self.width, self.height) } /// The label buffer, for a shader that masks by region id. pub fn labels(&self) -> &wgpu::Buffer { &self.labels } /// Build the region adjacency graph, reading the labels back to the CPU. /// /// **Not a shipping path** — see this module's header. Gated so it cannot /// be reached from a production build by accident. #[cfg(any(test, feature = "readback"))] pub fn read_field(&self) -> Result { let n = (self.width * self.height) as usize; let roots: Vec = read_buffer(&self.ctx, &self.labels, n)?; let gradient: Vec = read_buffer(&self.ctx, &self.gradient, n)?; Ok(crate::hierarchy::RegionField::from_roots( &roots, &gradient, self.width as usize, self.height as usize, )) } } fn uniform_entry(binding: u32) -> wgpu::BindGroupLayoutEntry { wgpu::BindGroupLayoutEntry { binding, visibility: wgpu::ShaderStages::COMPUTE, ty: wgpu::BindingType::Buffer { ty: wgpu::BufferBindingType::Uniform, has_dynamic_offset: false, min_binding_size: None, }, count: None, } } fn storage_entry(binding: u32, read_only: bool) -> wgpu::BindGroupLayoutEntry { wgpu::BindGroupLayoutEntry { binding, visibility: wgpu::ShaderStages::COMPUTE, ty: wgpu::BindingType::Buffer { ty: wgpu::BufferBindingType::Storage { read_only }, has_dynamic_offset: false, min_binding_size: None, }, count: None, } } fn compute( ctx: &GpuContext, module: &wgpu::ShaderModule, layout: &wgpu::BindGroupLayout, entry: &str, ) -> wgpu::ComputePipeline { let pipeline_layout = ctx .device .create_pipeline_layout(&wgpu::PipelineLayoutDescriptor { label: Some("watershed-layout"), bind_group_layouts: &[Some(layout)], immediate_size: 0, }); ctx.device .create_compute_pipeline(&wgpu::ComputePipelineDescriptor { label: Some(entry), layout: Some(&pipeline_layout), module, entry_point: Some(entry), compilation_options: Default::default(), cache: None, }) } /// The proxy size for a source, preserving aspect and never upscaling. fn proxy_size(src_w: u32, src_h: u32, max_edge: u32) -> (u32, u32) { let longest = src_w.max(src_h); if longest <= max_edge || longest == 0 { return (src_w.max(1), src_h.max(1)); } let scale = f64::from(max_edge) / f64::from(longest); ( ((f64::from(src_w) * scale).round() as u32).max(1), ((f64::from(src_h) * scale).round() as u32).max(1), ) } #[cfg(any(test, feature = "readback"))] fn read_buffer( ctx: &GpuContext, buffer: &wgpu::Buffer, len: usize, ) -> Result, GpuError> { let size = (len * std::mem::size_of::()) as u64; let staging = ctx.device.create_buffer(&wgpu::BufferDescriptor { label: Some("watershed-readback"), size, 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_buffer_to_buffer(buffer, 0, &staging, 0, size); ctx.queue.submit(Some(enc.finish())); let slice = staging.slice(..); let (tx, rx) = std::sync::mpsc::channel(); slice.map_async(wgpu::MapMode::Read, move |r| { let _ = tx.send(r); }); ctx.device .poll(wgpu::PollType::wait_indefinitely()) .map_err(|e| GpuError::Readback(e.to_string()))?; rx.recv() .map_err(|e| GpuError::Readback(e.to_string()))? .map_err(|e| GpuError::Readback(e.to_string()))?; let data = slice.get_mapped_range(); let out = bytemuck::cast_slice::(&data).to_vec(); drop(data); staging.unmap(); Ok(out) } #[cfg(test)] mod tests { use super::*; use crate::hierarchy::MergeTree; fn ctx() -> Option { match pollster::block_on(GpuContext::new_headless()) { Ok(c) => Some(c), Err(e) => { eprintln!("skipping: no GPU adapter ({e})"); None } } } #[test] fn a_proxy_preserves_aspect_and_never_upscales() { assert_eq!(proxy_size(6000, 4000, 1600), (1600, 1067)); assert_eq!(proxy_size(4000, 6000, 1600), (1067, 1600)); // A thumbnail must not be blown up to the proxy size — there is no // detail there to find basins in. assert_eq!(proxy_size(800, 600, 1600), (800, 600)); assert_eq!(proxy_size(0, 0, 1600), (1, 1)); } /// Two flat halves split by a hard vertical edge. fn two_tone(w: u32, h: u32) -> Vec { let mut px = Vec::with_capacity((w * h * 4) as usize); for _ in 0..h { for x in 0..w { let v = if x < w / 2 { 30u8 } else { 220u8 }; px.extend_from_slice(&[v, v, v, 255]); } } px } #[test] fn a_hard_edge_produces_two_regions_at_the_top_of_the_ladder() { // The end-to-end property, on an image whose answer is not in doubt: // whatever the watershed does with texture, it must not lose an edge // this obvious, and the coarsest non-trivial cut must be exactly the // two halves. let Some(ctx) = ctx() else { return }; let (w, h) = (64u32, 64u32); let src = DemosaicedImage::from_rgba8(&ctx, &two_tone(w, h), w, h).expect("source"); let pass = SegmentPass::new(&ctx).expect("segment pass"); let seg = pass.run(&src, SegmentOptions::default()).expect("run"); assert_eq!(seg.size(), (w, h)); let field = seg.read_field().expect("read field"); let tree = MergeTree::build(&field); let px = field.apply(&tree.cut_to(2)); for y in 0..h as usize { let left = px[y * w as usize]; let right = px[y * w as usize + w as usize - 1]; assert_ne!(left, right, "the two halves must not share a region"); } } #[test] fn a_flat_image_does_not_fragment() { // The noise case in miniature. A gradient of zero everywhere is one // enormous plateau, which is exactly where a watershed without a // strict tie-break either hangs or shatters into per-pixel basins. let Some(ctx) = ctx() else { return }; let (w, h) = (32u32, 32u32); let flat = vec![128u8; (w * h * 4) as usize]; let src = DemosaicedImage::from_rgba8(&ctx, &flat, w, h).expect("source"); let pass = SegmentPass::new(&ctx).expect("segment pass"); let seg = pass.run(&src, SegmentOptions::default()).expect("run"); let field = seg.read_field().expect("read field"); assert_eq!( field.region_count, 1, "a plateau should resolve to one basin, not {}", field.region_count ); } #[test] fn the_same_image_segments_identically_twice() { // M5 on one device — the weaker half of the determinism question, but // the half that catches a race in the pointer jumping. Cross-vendor // is the part that needs hardware this test cannot assume. let Some(ctx) = ctx() else { return }; let (w, h) = (48u32, 48u32); let src = DemosaicedImage::from_rgba8(&ctx, &two_tone(w, h), w, h).expect("source"); let pass = SegmentPass::new(&ctx).expect("segment pass"); let a = pass .run(&src, SegmentOptions::default()) .expect("run") .read_field() .expect("field"); let b = pass .run(&src, SegmentOptions::default()) .expect("run") .read_field() .expect("field"); assert_eq!(a, b, "segmentation must be reproducible run to run"); } }