Let the model say what a thing is and the watershed say where it ends
Local masking needs to know where an image's regions are. The watershed spike (S15 arm A) found the boundaries but had no idea what any of them enclosed; its coarse levels were geometric accidents. This adds the other half and the thing that joins them. `core/dr-segment` is where region reasoning now lives — the hierarchy moves out of `dr-gpu`, which keeps only the pixel passes that are genuinely shaders. The new crate is device-free and, without its default features, model-free too: 20 of its tests need neither an adapter nor 11 MB of weights. Arm B runs YOLO26n-seg through `ort`. D13 framed inference as a choice between `ort`'s C++ runtime and the pure-Rust dependency policy; that was a false choice. `ort`'s `alternative-backend` feature unlinks the C entirely and `ort-tract` supplies the API from tract, which is pure Rust. Measured before committing to it: zero unsupported operators, 420 ms for 640x640, and correct masks on bus.jpg. No NDK problem to solve, so D13's largest tolerated exception is not needed. Arm C is `prior.rs`, and it ships because the two arms fail in opposite directions. Instance membership re-weights the merge saddles, so region pairs the model believes share an object merge early and pairs straddling its edge merge late. No boundary moves — only the order in which they dissolve — which is how the result stays pixel-accurate at every level while its coarse levels become named things. Two things the spec assumed that turned out to be false, both recorded in models/LICENCE.md: there is no usable ADE20K-trained YOLO, so the shipped vocabulary is COCO's 80 subjects and *stuff* like sky and foliage must come from arm A; and tract cannot parse a dynamic-shape export, so the graph's input is fixed and tiling is the only route to more semantic resolution. Weights are AGPL-3.0, which GPLv3 §13 permits and which makes the combined work effectively AGPL. Deliberate, not accidental. They live in Git LFS, and a build script fails with an instruction rather than embedding a pointer file when the clone lacks them.
This commit is contained in:
@@ -9,6 +9,12 @@ license.workspace = true
|
||||
dr-types.workspace = true
|
||||
dr-decode.workspace = true
|
||||
dr-pipeline.workspace = true
|
||||
# The watershed's pixel passes are here because they are shaders; everything
|
||||
# that reasons about regions rather than pixels lives there, where it is
|
||||
# testable with no adapter present. No features: this half needs neither the
|
||||
# inference runtime nor the weights, and the workspace declaration defaults
|
||||
# them off so that stays true.
|
||||
dr-segment.workspace = true
|
||||
wgpu.workspace = true
|
||||
thiserror.workspace = true
|
||||
log.workspace = true
|
||||
|
||||
@@ -12,7 +12,7 @@
|
||||
//! PPM for the same reason `develop` uses it: no encoder dependency, and
|
||||
//! every viewer reads it. This is a diagnostic, not an export path.
|
||||
|
||||
use dr_gpu::hierarchy::{MergeTree, RegionField};
|
||||
use dr_segment::{MergeTree, RegionField};
|
||||
use dr_gpu::{DemosaicedImage, Demosaicer, GpuContext, SegmentOptions, SegmentPass};
|
||||
|
||||
/// The ladder the example dumps. Chosen to span "far too fine to be useful"
|
||||
|
||||
@@ -20,7 +20,6 @@ use wgpu::util::DeviceExt;
|
||||
mod adjust;
|
||||
mod demosaic;
|
||||
mod error;
|
||||
pub mod hierarchy;
|
||||
mod histogram;
|
||||
mod readback;
|
||||
mod segment;
|
||||
|
||||
@@ -2,7 +2,7 @@
|
||||
//!
|
||||
//! 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.
|
||||
//! those labels lives in [`dr_segment`], which needs no device.
|
||||
//!
|
||||
//! # Cost
|
||||
//!
|
||||
@@ -453,7 +453,7 @@ pub struct Segmentation {
|
||||
width: u32,
|
||||
height: u32,
|
||||
/// Per pixel, the linear index of its basin root. Sparse — compacted by
|
||||
/// [`crate::hierarchy::RegionField::from_roots`].
|
||||
/// [`dr_segment::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))]
|
||||
@@ -475,11 +475,11 @@ impl Segmentation {
|
||||
/// **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<crate::hierarchy::RegionField, GpuError> {
|
||||
pub fn read_field(&self) -> Result<dr_segment::RegionField, GpuError> {
|
||||
let n = (self.width * self.height) as usize;
|
||||
let roots: Vec<u32> = read_buffer(&self.ctx, &self.labels, n)?;
|
||||
let gradient: Vec<f32> = read_buffer(&self.ctx, &self.gradient, n)?;
|
||||
Ok(crate::hierarchy::RegionField::from_roots(
|
||||
Ok(dr_segment::RegionField::from_roots(
|
||||
&roots,
|
||||
&gradient,
|
||||
self.width as usize,
|
||||
@@ -591,7 +591,7 @@ fn read_buffer<T: bytemuck::Pod>(
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use crate::hierarchy::MergeTree;
|
||||
use dr_segment::MergeTree;
|
||||
|
||||
fn ctx() -> Option<GpuContext> {
|
||||
match pollster::block_on(GpuContext::new_headless()) {
|
||||
|
||||
@@ -0,0 +1,46 @@
|
||||
[package]
|
||||
name = "dr-segment"
|
||||
version.workspace = true
|
||||
edition.workspace = true
|
||||
rust-version.workspace = true
|
||||
license.workspace = true
|
||||
# Guards against a Git LFS pointer being embedded in place of the weights.
|
||||
build = "build.rs"
|
||||
|
||||
[dependencies]
|
||||
thiserror.workspace = true
|
||||
log.workspace = true
|
||||
|
||||
# Inference. `ort` is the API; **tract is the engine** — see the workspace
|
||||
# manifest for why the C++ ONNX Runtime is not linked here.
|
||||
ort = { workspace = true, optional = true }
|
||||
ort-tract = { workspace = true, optional = true }
|
||||
ndarray = { workspace = true, optional = true }
|
||||
|
||||
[dev-dependencies]
|
||||
# The example reads an ordinary JPEG, because the thing worth looking at is
|
||||
# whether detections land on a real photograph. Pure Rust, and already in the
|
||||
# tree for embedded previews.
|
||||
zune-jpeg.workspace = true
|
||||
env_logger.workspace = true
|
||||
|
||||
[features]
|
||||
# On by default: a local adjustment that cannot select a subject is half the
|
||||
# feature, and the whole point of the tract backend is that enabling this costs
|
||||
# no C dependency on any platform.
|
||||
default = ["semantic", "embedded-model"]
|
||||
|
||||
# Arm B — the ONNX runtime and the instance decoder.
|
||||
#
|
||||
# Separable because the watershed half is genuinely independent of it: with
|
||||
# this off, `dr-segment` is a pure-CPU graph algorithm crate with no model to
|
||||
# carry, which is what the headless hierarchy tests want.
|
||||
semantic = ["dep:ort", "dep:ort-tract", "dep:ndarray"]
|
||||
|
||||
# Compile the weights into the binary.
|
||||
#
|
||||
# Separate from `semantic` because the two answer different questions. Android
|
||||
# hands the app no filesystem path to read a model from (ARCH §6.9), so there
|
||||
# it must be embedded; a desktop packager pointing at a system model directory,
|
||||
# or a test that only needs the decoder, wants the runtime without the 11 MB.
|
||||
embedded-model = ["semantic"]
|
||||
@@ -0,0 +1,58 @@
|
||||
//! Check the model is a model and not an LFS pointer.
|
||||
//!
|
||||
//! `models/*.onnx` is stored in Git LFS (see `.gitattributes`). A clone made
|
||||
//! without git-lfs installed, or with `GIT_LFS_SKIP_SMUDGE` set, leaves a
|
||||
//! ~130-byte text pointer at that path instead of the weights.
|
||||
//!
|
||||
//! Without this check `include_bytes!` would happily embed the pointer, the
|
||||
//! crate would compile, and the failure would surface much later as an opaque
|
||||
//! ONNX parse error from inside tract — at which point the connection back to
|
||||
//! a missing `git lfs pull` is not one anybody would make quickly. Failing
|
||||
//! here costs one file read and turns that into a sentence.
|
||||
|
||||
use std::path::Path;
|
||||
|
||||
const MODEL: &str = "models/yolo26n-seg.onnx";
|
||||
|
||||
fn main() {
|
||||
println!("cargo:rerun-if-changed={MODEL}");
|
||||
println!("cargo:rerun-if-changed=build.rs");
|
||||
|
||||
// Only the embedded path needs the file present; a build without it is
|
||||
// watershed-only by choice and should not be blocked on weights.
|
||||
if std::env::var_os("CARGO_FEATURE_EMBEDDED_MODEL").is_none() {
|
||||
return;
|
||||
}
|
||||
|
||||
let path = Path::new(MODEL);
|
||||
let Ok(bytes) = std::fs::read(path) else {
|
||||
panic!(
|
||||
"\n\n{MODEL} is missing.\n\
|
||||
It ships in Git LFS. Run `git lfs install && git lfs pull`, or build \
|
||||
with `--no-default-features` for a watershed-only build.\n"
|
||||
);
|
||||
};
|
||||
|
||||
// ONNX is protobuf, which has no magic number; an LFS pointer is short
|
||||
// ASCII beginning with a version URL. Testing for the pointer is the
|
||||
// reliable direction — it has a known shape, where "valid protobuf" does
|
||||
// not, and this check only has to catch the one failure that actually
|
||||
// happens in practice.
|
||||
if bytes.starts_with(b"version https://git-lfs") {
|
||||
panic!(
|
||||
"\n\n{MODEL} is a Git LFS pointer, not the model ({} bytes).\n\
|
||||
Run `git lfs install && git lfs pull` to fetch the real file.\n",
|
||||
bytes.len()
|
||||
);
|
||||
}
|
||||
|
||||
// A model that parses as a pointer-sized file is not one either. The real
|
||||
// export is ~11 MB; anything under a megabyte is a truncated checkout.
|
||||
if bytes.len() < 1_000_000 {
|
||||
panic!(
|
||||
"\n\n{MODEL} is only {} bytes — expected ~11 MB.\n\
|
||||
The checkout looks incomplete; try `git lfs pull`.\n",
|
||||
bytes.len()
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,149 @@
|
||||
//! Run the semantic arm over a JPEG and write what it found.
|
||||
//!
|
||||
//! The point of S15 step 2 applied to arm B: no amount of unit testing settles
|
||||
//! whether the decode is right, because a transposed axis or an off-by-one in
|
||||
//! the letterbox produces perfectly plausible numbers and a mask sitting six
|
||||
//! pixels to the left. Looking at the overlay settles it in one glance.
|
||||
//!
|
||||
//! ```sh
|
||||
//! cargo run -p dr-segment --example detect --release -- photo.jpg
|
||||
//! cargo run -p dr-segment --example detect --release -- photo.jpg out 0.25 tiled
|
||||
//! ```
|
||||
//!
|
||||
//! Writes `<prefix>-overlay.ppm` — the image with each instance tinted by a
|
||||
//! per-instance colour — and prints the detection list. PPM for the same
|
||||
//! reason the other examples use it: no encoder dependency, and every viewer
|
||||
//! reads it.
|
||||
|
||||
use dr_segment::semantic::{SemanticModel, SemanticOptions, Tiling};
|
||||
|
||||
fn main() {
|
||||
env_logger::init();
|
||||
|
||||
let mut args = std::env::args().skip(1);
|
||||
let Some(path) = args.next() else {
|
||||
eprintln!("usage: detect <photo.jpg> [out-prefix] [confidence] [tiled]");
|
||||
std::process::exit(2);
|
||||
};
|
||||
let prefix = args.next().unwrap_or_else(|| "detect".into());
|
||||
let confidence = args
|
||||
.next()
|
||||
.and_then(|s| s.parse().ok())
|
||||
.unwrap_or(SemanticOptions::default().confidence);
|
||||
let tiled = args.next().is_some_and(|s| s == "tiled");
|
||||
|
||||
let (rgb, width, height) = load_jpeg(&path);
|
||||
println!("image {width}x{height}");
|
||||
|
||||
let options = SemanticOptions {
|
||||
confidence,
|
||||
tiling: if tiled {
|
||||
Tiling::Grid { overlap: 0.25 }
|
||||
} else {
|
||||
Tiling::Whole
|
||||
},
|
||||
..SemanticOptions::default()
|
||||
};
|
||||
println!(
|
||||
"tiling {}",
|
||||
if tiled { "grid, 25% overlap" } else { "whole frame" }
|
||||
);
|
||||
|
||||
let t0 = std::time::Instant::now();
|
||||
let mut model = SemanticModel::embedded().expect("load embedded model");
|
||||
println!("load {:.0} ms", t0.elapsed().as_secs_f32() * 1000.0);
|
||||
|
||||
let t1 = std::time::Instant::now();
|
||||
let instances = model
|
||||
.detect(&rgb, width, height, &options)
|
||||
.expect("inference");
|
||||
println!(
|
||||
"detect {:.0} ms",
|
||||
t1.elapsed().as_secs_f32() * 1000.0
|
||||
);
|
||||
println!("found {} instances", instances.len());
|
||||
|
||||
for (i, inst) in instances.iter().enumerate() {
|
||||
let covered = inst.mask.iter().filter(|&&m| m >= 0.5).count();
|
||||
println!(
|
||||
" [{i:2}] {:<14} {:.2} box ({:.0},{:.0})-({:.0},{:.0}) {:.1}% of frame",
|
||||
inst.class_name,
|
||||
inst.score,
|
||||
inst.bbox.0,
|
||||
inst.bbox.1,
|
||||
inst.bbox.2,
|
||||
inst.bbox.3,
|
||||
100.0 * covered as f32 / (width * height) as f32,
|
||||
);
|
||||
}
|
||||
|
||||
// Tint each instance and write the composite. A mask in the wrong place is
|
||||
// obvious here and invisible in the numbers above.
|
||||
let mut out = vec![0u8; width * height * 3];
|
||||
for (p, px) in out.chunks_exact_mut(3).enumerate() {
|
||||
for c in 0..3 {
|
||||
px[c] = (rgb[p * 3 + c].clamp(0.0, 1.0) * 255.0) as u8;
|
||||
}
|
||||
}
|
||||
for (i, inst) in instances.iter().enumerate() {
|
||||
let tint = colour(i);
|
||||
for (p, &m) in inst.mask.iter().enumerate() {
|
||||
if m < 0.5 {
|
||||
continue;
|
||||
}
|
||||
let px = &mut out[p * 3..p * 3 + 3];
|
||||
for c in 0..3 {
|
||||
px[c] = ((px[c] as f32) * 0.45 + tint[c] as f32 * 0.55) as u8;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
let file = format!("{prefix}-overlay.ppm");
|
||||
write_ppm(&file, &out, width, height);
|
||||
println!("wrote {file}");
|
||||
}
|
||||
|
||||
/// A distinct colour per instance index — the same golden-angle walk the
|
||||
/// watershed example uses, so the two overlays are read the same way.
|
||||
fn colour(i: usize) -> [u8; 3] {
|
||||
let h = (i as f32 * 137.508) % 360.0;
|
||||
let (c, x) = (255.0, 255.0 * (1.0 - ((h / 60.0) % 2.0 - 1.0).abs()));
|
||||
let (r, g, b) = match (h / 60.0) as u32 {
|
||||
0 => (c, x, 0.0),
|
||||
1 => (x, c, 0.0),
|
||||
2 => (0.0, c, x),
|
||||
3 => (0.0, x, c),
|
||||
4 => (x, 0.0, c),
|
||||
_ => (c, 0.0, x),
|
||||
};
|
||||
[r as u8, g as u8, b as u8]
|
||||
}
|
||||
|
||||
fn load_jpeg(path: &str) -> (Vec<f32>, usize, usize) {
|
||||
let bytes = std::fs::read(path).unwrap_or_else(|e| panic!("read {path}: {e}"));
|
||||
let mut decoder = zune_jpeg::JpegDecoder::new(&bytes);
|
||||
let pixels = decoder.decode().expect("decode jpeg");
|
||||
let info = decoder.info().expect("jpeg info");
|
||||
let (w, h) = (info.width as usize, info.height as usize);
|
||||
|
||||
// The model was trained on gamma-encoded sRGB, so the JPEG's own values go
|
||||
// through unlinearised — this is one of the few places in the codebase
|
||||
// where *not* linearising is the correct thing to do.
|
||||
let rgb = match pixels.len() / (w * h) {
|
||||
3 => pixels.iter().map(|&v| v as f32 / 255.0).collect(),
|
||||
1 => pixels
|
||||
.iter()
|
||||
.flat_map(|&v| [v as f32 / 255.0; 3])
|
||||
.collect(),
|
||||
n => panic!("unexpected {n} components per pixel"),
|
||||
};
|
||||
|
||||
(rgb, w, h)
|
||||
}
|
||||
|
||||
fn write_ppm(path: &str, rgb: &[u8], width: usize, height: usize) {
|
||||
use std::io::Write;
|
||||
let mut f = std::io::BufWriter::new(std::fs::File::create(path).expect("create ppm"));
|
||||
write!(f, "P6\n{width} {height}\n255\n").expect("ppm header");
|
||||
f.write_all(rgb).expect("ppm body");
|
||||
}
|
||||
@@ -0,0 +1,54 @@
|
||||
# Model weights — licensing
|
||||
|
||||
`yolo26n-seg.onnx` is exported from Ultralytics YOLO26n-seg
|
||||
(`https://huggingface.co/Ultralytics/YOLO26`, `yolo26n-seg.pt`) by
|
||||
`tools/export-seg-model.sh`. `yolo26n-seg.classes.json` is that checkpoint's
|
||||
class vocabulary, written out by the same script.
|
||||
|
||||
## The grant
|
||||
|
||||
**Ultralytics releases YOLO under AGPL-3.0**, and the weights carry the same
|
||||
grant as the framework — the HuggingFace repository declares `agpl-3.0` for the
|
||||
checkpoints themselves, not merely for the training code. A commercial licence
|
||||
is offered separately; DarkRoom does not use it and does not need it.
|
||||
|
||||
## What that means for DarkRoom
|
||||
|
||||
DarkRoom is GPL-3.0-or-later. **GPLv3 §13 explicitly permits combination with
|
||||
AGPL-3.0 code**, so redistributing these weights inside this repository is
|
||||
allowed — this is *not* the situation the InsightFace "buffalo" weights would
|
||||
have created, where a non-commercial research grant is simply incompatible with
|
||||
the project's licence and with F-Droid, Flatpak and Play distribution
|
||||
(NFR-COMPAT-2, D13).
|
||||
|
||||
The consequence, and it is a real one: **the combined work is effectively
|
||||
AGPL-3.0.** §13's permission runs one way — the AGPL's §13 network-use condition
|
||||
attaches to the portion under that licence. For a local-first desktop and
|
||||
Android photo editor that condition has no practical bite, because there is no
|
||||
network service offering the combined work to remote users. It would acquire
|
||||
bite the moment any hosted or server-side rendering appeared, and that is the
|
||||
thing to remember rather than rediscover.
|
||||
|
||||
This was decided deliberately (D14), not arrived at by accident, and
|
||||
`docs/segmentation.md` §7 records the reasoning.
|
||||
|
||||
## Class vocabulary — a caveat worth reading
|
||||
|
||||
`docs/segmentation.md` §4 specified YOLO **pretrained on ADE20K**, whose 150
|
||||
classes include the *stuff* categories that matter most in photography — sky,
|
||||
vegetation, water, wall, mountain.
|
||||
|
||||
**No such model exists in usable form.** Checked 2026-08-21: Ultralytics ships
|
||||
YOLO26-seg trained on **COCO**, whose 80 classes are all *things* — person,
|
||||
dog, car, bird, potted plant — and the one HuggingFace repository claiming a
|
||||
YOLO/ADE20K combination (`laxmacl/yolov8-ade20k`) is empty. ADE20K semantic
|
||||
models do exist, but as SegFormer/OneFormer/MaskFormer transformers, not YOLO.
|
||||
|
||||
So the shipped vocabulary selects **subjects**, not **stuff**. "Select the
|
||||
person" works; "select the sky" does not come from the model and must come from
|
||||
the watershed hierarchy instead. That is a narrower arm B than §4 assumed, and
|
||||
it raises rather than lowers the importance of arm C.
|
||||
|
||||
The loader treats the vocabulary as model metadata rather than compiled-in
|
||||
knowledge, so adding a stuff-class model later is a file plus a descriptor, not
|
||||
a code change.
|
||||
@@ -0,0 +1,82 @@
|
||||
[
|
||||
"person",
|
||||
"bicycle",
|
||||
"car",
|
||||
"motorcycle",
|
||||
"airplane",
|
||||
"bus",
|
||||
"train",
|
||||
"truck",
|
||||
"boat",
|
||||
"traffic light",
|
||||
"fire hydrant",
|
||||
"stop sign",
|
||||
"parking meter",
|
||||
"bench",
|
||||
"bird",
|
||||
"cat",
|
||||
"dog",
|
||||
"horse",
|
||||
"sheep",
|
||||
"cow",
|
||||
"elephant",
|
||||
"bear",
|
||||
"zebra",
|
||||
"giraffe",
|
||||
"backpack",
|
||||
"umbrella",
|
||||
"handbag",
|
||||
"tie",
|
||||
"suitcase",
|
||||
"frisbee",
|
||||
"skis",
|
||||
"snowboard",
|
||||
"sports ball",
|
||||
"kite",
|
||||
"baseball bat",
|
||||
"baseball glove",
|
||||
"skateboard",
|
||||
"surfboard",
|
||||
"tennis racket",
|
||||
"bottle",
|
||||
"wine glass",
|
||||
"cup",
|
||||
"fork",
|
||||
"knife",
|
||||
"spoon",
|
||||
"bowl",
|
||||
"banana",
|
||||
"apple",
|
||||
"sandwich",
|
||||
"orange",
|
||||
"broccoli",
|
||||
"carrot",
|
||||
"hot dog",
|
||||
"pizza",
|
||||
"donut",
|
||||
"cake",
|
||||
"chair",
|
||||
"couch",
|
||||
"potted plant",
|
||||
"bed",
|
||||
"dining table",
|
||||
"toilet",
|
||||
"tv",
|
||||
"laptop",
|
||||
"mouse",
|
||||
"remote",
|
||||
"keyboard",
|
||||
"cell phone",
|
||||
"microwave",
|
||||
"oven",
|
||||
"toaster",
|
||||
"sink",
|
||||
"refrigerator",
|
||||
"book",
|
||||
"clock",
|
||||
"vase",
|
||||
"scissors",
|
||||
"teddy bear",
|
||||
"hair drier",
|
||||
"toothbrush"
|
||||
]
|
||||
Binary file not shown.
@@ -0,0 +1,59 @@
|
||||
//! Region segmentation for local masking (S15, docs/segmentation.md).
|
||||
//!
|
||||
//! Local adjustments need to know where the image's regions are before they
|
||||
//! can snap a mask to one. This crate is that map, and it is deliberately
|
||||
//! **device-free**: the watershed's pixel passes live in `dr-gpu` because they
|
||||
//! are shaders, and everything that reasons about *regions* rather than
|
||||
//! *pixels* lives here, where it can be tested on hand-built inputs with no
|
||||
//! adapter present (ARCH §6.5a).
|
||||
//!
|
||||
//! # The three arms
|
||||
//!
|
||||
//! [`hierarchy`] is **arm A** — a watershed over-segments the image and the
|
||||
//! recorded merge order becomes a granularity ladder. Deterministic, needs no
|
||||
//! model, works on any picture, and knows nothing about what anything *is*.
|
||||
//!
|
||||
//! [`semantic`] is **arm B** — a YOLO instance-segmentation model naming the
|
||||
//! subjects it recognises. Knows what things are, and is vague about exactly
|
||||
//! where their edges fall (its prototypes are quarter-resolution).
|
||||
//!
|
||||
//! [`prior`] is **arm C**, and it is the one that ships. Arm B's instances
|
||||
//! *re-weight* arm A's merge order, so coarse levels of the ladder line up
|
||||
//! with real objects while every boundary stays exactly where the watershed
|
||||
//! put it. The model contributes what it is good at — knowing what things are
|
||||
//! — and the watershed contributes what it is good at, which is knowing where
|
||||
//! the edge is, to the pixel, at every scale.
|
||||
//!
|
||||
//! That combination is also what repairs the vocabulary problem. The shipped
|
||||
//! model is COCO-trained, so it recognises subjects and has no class for sky,
|
||||
//! foliage or wall (`models/LICENCE.md`). Selecting those falls to arm A,
|
||||
//! which never needed a vocabulary to begin with.
|
||||
|
||||
pub mod hierarchy;
|
||||
pub mod prior;
|
||||
#[cfg(feature = "semantic")]
|
||||
pub mod semantic;
|
||||
|
||||
pub use hierarchy::{Edge, Merge, MergeTree, RegionField};
|
||||
pub use prior::{Membership, PriorOptions};
|
||||
#[cfg(feature = "semantic")]
|
||||
pub use semantic::{Instance, SemanticModel, SemanticOptions, Tiling};
|
||||
|
||||
/// What can go wrong between an image and a region map.
|
||||
#[derive(Debug, thiserror::Error)]
|
||||
pub enum SegmentError {
|
||||
#[error("could not read model file: {0}")]
|
||||
ModelRead(#[source] std::io::Error),
|
||||
|
||||
#[cfg(feature = "semantic")]
|
||||
#[error("inference failed: {0}")]
|
||||
Inference(#[source] ort::Error),
|
||||
|
||||
#[error("image buffer is {got} floats, expected {expected} (RGB, three per pixel)")]
|
||||
ImageShape { expected: usize, got: usize },
|
||||
|
||||
/// The graph produced something the decoder does not recognise — a
|
||||
/// different model, or a different export of the same one.
|
||||
#[error("model output '{0}' did not have the expected shape")]
|
||||
OutputShape(&'static str),
|
||||
}
|
||||
@@ -0,0 +1,418 @@
|
||||
//! Arm C — semantic instances as a prior over the watershed merge order.
|
||||
//!
|
||||
//! docs/segmentation.md §5. The spec calls this the expected winner and it is
|
||||
//! what ships, for a reason that survives the model turning out to be narrower
|
||||
//! than §4 assumed: the two arms fail in *opposite* directions, so each one
|
||||
//! covers the other's failure.
|
||||
//!
|
||||
//! - The watershed knows where every edge is and nothing about what it
|
||||
//! separates. Its coarse levels are geometric accidents — level 7 is *a*
|
||||
//! coarser partition, not *the* object.
|
||||
//! - The model knows a dog is a dog and puts the dog's outline roughly where
|
||||
//! the dog is, at quarter resolution, with a soft edge.
|
||||
//!
|
||||
//! Weighting the merge by semantic agreement takes the outline from the
|
||||
//! watershed and the grouping from the model. A region pair the model believes
|
||||
//! belongs to one object merges early; a pair straddling an object's edge
|
||||
//! merges late. **No boundary moves** — only the order in which boundaries
|
||||
//! dissolve — which is why the result is pixel-accurate at every level while
|
||||
//! its coarse levels are named things.
|
||||
//!
|
||||
//! # Why this is not "just use the model's mask"
|
||||
//!
|
||||
//! Because a mask edge is judged at 100% zoom, and the model's edge is a
|
||||
//! quarter-resolution sigmoid. Using the instance mask directly gives a
|
||||
//! selection that is semantically right and visibly soft — acceptable for
|
||||
//! biasing, not acceptable as the mask itself. Snapping to watershed regions
|
||||
//! ([`regions_for_instance`]) gives the same selection with the sensor's own
|
||||
//! edges.
|
||||
|
||||
use std::collections::HashMap;
|
||||
|
||||
use crate::hierarchy::{Edge, RegionField};
|
||||
|
||||
/// How strongly the model is allowed to reorder the merge.
|
||||
#[derive(Debug, Clone, Copy, PartialEq)]
|
||||
pub struct PriorOptions {
|
||||
/// How far a confident semantic judgement may scale a saddle, `0.0..1.0`.
|
||||
///
|
||||
/// At `0.0` this is arm A exactly. At `0.9` a pair the model is certain
|
||||
/// shares an object merges at a tenth of its true boundary strength.
|
||||
///
|
||||
/// Not `1.0`, and the ceiling is the point: at `1.0` an agreeing pair gets
|
||||
/// a saddle of zero and merges *before* genuinely identical neighbours,
|
||||
/// which lets a confident-but-wrong detection flatten real structure it
|
||||
/// happens to cover. Leaving headroom keeps the image's own evidence able
|
||||
/// to outvote the model.
|
||||
pub strength: f32,
|
||||
|
||||
/// Coverage above which a region counts as belonging to an instance.
|
||||
///
|
||||
/// Applied to the *mean* of the instance's soft mask over the region, so
|
||||
/// this is "most of this region is inside the dog", not "some pixel is".
|
||||
pub membership: f32,
|
||||
|
||||
/// Instances scoring below this contribute no prior at all.
|
||||
///
|
||||
/// Higher than the detection threshold on purpose. A weak detection is
|
||||
/// still worth *offering* in a list a person picks from, where the cost of
|
||||
/// being wrong is an ignored entry — but not worth silently reshaping the
|
||||
/// hierarchy every other interaction depends on.
|
||||
pub confidence: f32,
|
||||
}
|
||||
|
||||
impl Default for PriorOptions {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
strength: 0.75,
|
||||
membership: 0.5,
|
||||
confidence: 0.4,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Which instance, if any, each region belongs to.
|
||||
///
|
||||
/// One dominant instance per region rather than a vector of memberships:
|
||||
/// regions are small — a proxy watershed makes thousands of them — and a
|
||||
/// region spanning two objects means the watershed already failed there, which
|
||||
/// is a case to leave to the image gradient rather than to average over.
|
||||
#[derive(Debug, Clone, PartialEq)]
|
||||
pub struct Membership {
|
||||
/// Per region: the instance index it mostly belongs to, and how much.
|
||||
pub of: Vec<Option<(usize, f32)>>,
|
||||
}
|
||||
|
||||
impl Membership {
|
||||
/// Compute per-region membership from soft instance masks.
|
||||
///
|
||||
/// `masks` is one full-image coverage buffer per instance, each
|
||||
/// `width * height` — [`crate::semantic::Instance::mask`] is exactly this,
|
||||
/// passed as slices so that this function needs no `semantic` feature and
|
||||
/// stays testable with hand-written masks.
|
||||
pub fn compute(
|
||||
field: &RegionField,
|
||||
masks: &[&[f32]],
|
||||
options: &PriorOptions,
|
||||
) -> Result<Self, MembershipError> {
|
||||
let pixels = field.width * field.height;
|
||||
for (i, mask) in masks.iter().enumerate() {
|
||||
if mask.len() != pixels {
|
||||
return Err(MembershipError::MaskSize {
|
||||
instance: i,
|
||||
expected: pixels,
|
||||
got: mask.len(),
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
// Sum coverage per (region, instance), then divide by region size.
|
||||
let mut totals = vec![0.0f32; field.region_count * masks.len()];
|
||||
let mut sizes = vec![0u32; field.region_count];
|
||||
|
||||
for (p, &label) in field.labels.iter().enumerate() {
|
||||
let r = label as usize;
|
||||
sizes[r] += 1;
|
||||
for (i, mask) in masks.iter().enumerate() {
|
||||
totals[r * masks.len() + i] += mask[p];
|
||||
}
|
||||
}
|
||||
|
||||
let of = (0..field.region_count)
|
||||
.map(|r| {
|
||||
let size = sizes[r].max(1) as f32;
|
||||
let row = &totals[r * masks.len()..(r + 1) * masks.len()];
|
||||
|
||||
// `total_cmp` and an index tiebreak: two instances covering a
|
||||
// region equally must resolve the same way on every machine,
|
||||
// because the merge order below is derived from this and an
|
||||
// unstable merge order is an unstable label field (M5).
|
||||
let best = row
|
||||
.iter()
|
||||
.enumerate()
|
||||
.max_by(|(ia, a), (ib, b)| a.total_cmp(b).then(ib.cmp(ia)))?;
|
||||
|
||||
let coverage = best.1 / size;
|
||||
(coverage >= options.membership).then_some((best.0, coverage))
|
||||
})
|
||||
.collect();
|
||||
|
||||
Ok(Self { of })
|
||||
}
|
||||
|
||||
/// The instance a region belongs to, if any.
|
||||
pub fn instance_of(&self, region: u32) -> Option<usize> {
|
||||
self.of.get(region as usize).copied().flatten().map(|(i, _)| i)
|
||||
}
|
||||
|
||||
/// How two regions relate semantically, in `-1.0..=1.0`.
|
||||
///
|
||||
/// `+c` when both sit in the same instance with confidence `c`, `-c` when
|
||||
/// they sit in different ones or one is inside an object and the other is
|
||||
/// background, and `0.0` when neither belongs to anything — two patches of
|
||||
/// hillside get no opinion from a model that has no word for hillside, and
|
||||
/// fall back to arm A untouched.
|
||||
pub fn affinity(&self, a: u32, b: u32) -> f32 {
|
||||
let (a, b) = (
|
||||
self.of.get(a as usize).copied().flatten(),
|
||||
self.of.get(b as usize).copied().flatten(),
|
||||
);
|
||||
match (a, b) {
|
||||
(Some((ia, ca)), Some((ib, cb))) if ia == ib => ca.min(cb),
|
||||
(Some((_, ca)), Some((_, cb))) => -ca.min(cb),
|
||||
(Some((_, c)), None) | (None, Some((_, c))) => -c,
|
||||
(None, None) => 0.0,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, thiserror::Error, PartialEq)]
|
||||
pub enum MembershipError {
|
||||
#[error("instance {instance} mask is {got} pixels, expected {expected}")]
|
||||
MaskSize {
|
||||
instance: usize,
|
||||
expected: usize,
|
||||
got: usize,
|
||||
},
|
||||
}
|
||||
|
||||
/// Re-weight a region field's boundaries by semantic agreement.
|
||||
///
|
||||
/// Returns a field whose labels are untouched and whose adjacency saddles have
|
||||
/// been scaled — so [`crate::MergeTree::build`] over the result yields a
|
||||
/// hierarchy that climbs toward objects instead of toward whatever happened to
|
||||
/// be smooth.
|
||||
///
|
||||
/// The scaling is `saddle * (1 - strength * affinity)`, which has the three
|
||||
/// properties that matter: agreement shrinks a saddle toward zero without ever
|
||||
/// reaching it, disagreement grows one without bound, and an affinity of zero
|
||||
/// is exactly arm A. A pair the model has no opinion about is left alone
|
||||
/// rather than nudged.
|
||||
pub fn apply_semantic_prior(
|
||||
field: &RegionField,
|
||||
membership: &Membership,
|
||||
options: &PriorOptions,
|
||||
) -> RegionField {
|
||||
let strength = options.strength.clamp(0.0, 0.99);
|
||||
|
||||
let mut adjacency: Vec<Edge> = field
|
||||
.adjacency
|
||||
.iter()
|
||||
.map(|e| Edge {
|
||||
a: e.a,
|
||||
b: e.b,
|
||||
saddle: e.saddle * (1.0 - strength * membership.affinity(e.a, e.b)),
|
||||
})
|
||||
.collect();
|
||||
|
||||
// Re-sorted because `MergeTree::build` consumes this in order and trusts
|
||||
// it to be sorted; the same total order as `RegionField::from_roots` uses,
|
||||
// for the same determinism reason.
|
||||
adjacency.sort_by(|x, y| {
|
||||
x.saddle
|
||||
.total_cmp(&y.saddle)
|
||||
.then(x.a.cmp(&y.a))
|
||||
.then(x.b.cmp(&y.b))
|
||||
});
|
||||
|
||||
RegionField {
|
||||
width: field.width,
|
||||
height: field.height,
|
||||
labels: field.labels.clone(),
|
||||
region_count: field.region_count,
|
||||
adjacency,
|
||||
}
|
||||
}
|
||||
|
||||
/// The regions making up one instance — click-to-select, snapped to edges.
|
||||
///
|
||||
/// This is the interaction the whole spike exists to enable, and the reason it
|
||||
/// returns *region ids* rather than a raster: a mask that is a set of integers
|
||||
/// is diffable, mergeable at node level under FR-NC-9, and cheap in a sidecar
|
||||
/// (docs/segmentation.md §1). A raster is none of those.
|
||||
///
|
||||
/// The returned ids are sorted, so the same click always produces the same
|
||||
/// mask — which is what lets it be a cache key.
|
||||
pub fn regions_for_instance(
|
||||
field: &RegionField,
|
||||
mask: &[f32],
|
||||
options: &PriorOptions,
|
||||
) -> Vec<u32> {
|
||||
let mut coverage = vec![0.0f32; field.region_count];
|
||||
let mut sizes = vec![0u32; field.region_count];
|
||||
|
||||
for (p, &label) in field.labels.iter().enumerate() {
|
||||
coverage[label as usize] += mask.get(p).copied().unwrap_or(0.0);
|
||||
sizes[label as usize] += 1;
|
||||
}
|
||||
|
||||
(0..field.region_count as u32)
|
||||
.filter(|&r| coverage[r as usize] / sizes[r as usize].max(1) as f32 >= options.membership)
|
||||
.collect()
|
||||
}
|
||||
|
||||
/// Every pixel covered by a set of region ids, as a binary mask.
|
||||
///
|
||||
/// The other direction: region ids are what gets *stored*, and a rasteriser
|
||||
/// needs pixels. On the shipping path this happens in a shader (ARCH §5.4);
|
||||
/// this exists for export, for tests, and for the example.
|
||||
pub fn rasterise(field: &RegionField, regions: &[u32]) -> Vec<f32> {
|
||||
let selected: HashMap<u32, ()> = regions.iter().map(|&r| (r, ())).collect();
|
||||
field
|
||||
.labels
|
||||
.iter()
|
||||
.map(|l| if selected.contains_key(l) { 1.0 } else { 0.0 })
|
||||
.collect()
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use crate::hierarchy::MergeTree;
|
||||
|
||||
/// A 4x2 field: regions 0 and 1 on the left, 2 and 3 on the right.
|
||||
fn field() -> RegionField {
|
||||
RegionField {
|
||||
width: 4,
|
||||
height: 2,
|
||||
labels: vec![0, 0, 2, 2, 1, 1, 3, 3],
|
||||
region_count: 4,
|
||||
adjacency: vec![
|
||||
Edge { a: 0, b: 1, saddle: 1.0 },
|
||||
Edge { a: 0, b: 2, saddle: 1.0 },
|
||||
Edge { a: 1, b: 3, saddle: 1.0 },
|
||||
Edge { a: 2, b: 3, saddle: 1.0 },
|
||||
],
|
||||
}
|
||||
}
|
||||
|
||||
/// An instance covering the left half — regions 0 and 1.
|
||||
fn left_half() -> Vec<f32> {
|
||||
vec![1.0, 1.0, 0.0, 0.0, 1.0, 1.0, 0.0, 0.0]
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn membership_finds_the_covered_regions() {
|
||||
let f = field();
|
||||
let m = Membership::compute(&f, &[&left_half()], &PriorOptions::default()).unwrap();
|
||||
|
||||
assert_eq!(m.instance_of(0), Some(0));
|
||||
assert_eq!(m.instance_of(1), Some(0));
|
||||
assert_eq!(m.instance_of(2), None, "right half is outside the instance");
|
||||
assert_eq!(m.instance_of(3), None);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn affinity_is_signed_by_agreement() {
|
||||
let f = field();
|
||||
let m = Membership::compute(&f, &[&left_half()], &PriorOptions::default()).unwrap();
|
||||
|
||||
assert!(m.affinity(0, 1) > 0.0, "both inside the instance");
|
||||
assert!(m.affinity(0, 2) < 0.0, "across the instance boundary");
|
||||
assert_eq!(m.affinity(2, 3), 0.0, "model has no opinion on either");
|
||||
}
|
||||
|
||||
/// The property arm C exists for: with equal image evidence everywhere,
|
||||
/// the semantic pair merges first.
|
||||
#[test]
|
||||
fn the_prior_reorders_the_merge_toward_the_object() {
|
||||
let f = field();
|
||||
let m = Membership::compute(&f, &[&left_half()], &PriorOptions::default()).unwrap();
|
||||
|
||||
// Arm A alone: every saddle is 1.0, so the merge order is arbitrary
|
||||
// and 0-1 has no reason to come first.
|
||||
let plain = MergeTree::build(&f);
|
||||
assert_eq!(plain.merges.len(), 3);
|
||||
|
||||
let biased = MergeTree::build(&apply_semantic_prior(&f, &m, &PriorOptions::default()));
|
||||
let first = biased.merges[0];
|
||||
assert_eq!(
|
||||
(first.a, first.b),
|
||||
(0, 1),
|
||||
"the two regions inside the instance should merge first"
|
||||
);
|
||||
assert!(
|
||||
first.saddle < 1.0,
|
||||
"agreement should lower the saddle, got {}",
|
||||
first.saddle
|
||||
);
|
||||
|
||||
// And the boundary the model believes in should now be the last to go.
|
||||
let last = biased.merges.last().unwrap();
|
||||
assert!(
|
||||
last.saddle > 1.0,
|
||||
"a semantic boundary should outlast the others, got {}",
|
||||
last.saddle
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn a_cut_at_two_groups_splits_along_the_instance() {
|
||||
let f = field();
|
||||
let m = Membership::compute(&f, &[&left_half()], &PriorOptions::default()).unwrap();
|
||||
let tree = MergeTree::build(&apply_semantic_prior(&f, &m, &PriorOptions::default()));
|
||||
|
||||
let grouping = tree.cut_to(2);
|
||||
assert_eq!(grouping[0], grouping[1], "instance regions share a group");
|
||||
assert_eq!(grouping[2], grouping[3], "background regions share a group");
|
||||
assert_ne!(grouping[0], grouping[2], "and the two groups differ");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn zero_strength_is_arm_a_exactly() {
|
||||
let f = field();
|
||||
let opts = PriorOptions { strength: 0.0, ..PriorOptions::default() };
|
||||
let m = Membership::compute(&f, &[&left_half()], &opts).unwrap();
|
||||
assert_eq!(apply_semantic_prior(&f, &m, &opts).adjacency, f.adjacency);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn selection_snaps_to_whole_regions() {
|
||||
let f = field();
|
||||
|
||||
// A mask that is ragged at the pixel level — as a quarter-resolution
|
||||
// sigmoid would be — still selects clean whole regions.
|
||||
let ragged = vec![1.0, 0.9, 0.1, 0.0, 0.8, 1.0, 0.0, 0.2];
|
||||
let regions = regions_for_instance(&f, &ragged, &PriorOptions::default());
|
||||
assert_eq!(regions, vec![0, 1]);
|
||||
|
||||
let pixels = rasterise(&f, ®ions);
|
||||
assert_eq!(pixels, vec![1.0, 1.0, 0.0, 0.0, 1.0, 1.0, 0.0, 0.0]);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn regions_are_sorted_so_a_click_is_a_cache_key() {
|
||||
let f = field();
|
||||
let all = vec![1.0; 8];
|
||||
let regions = regions_for_instance(&f, &all, &PriorOptions::default());
|
||||
assert!(regions.windows(2).all(|w| w[0] < w[1]));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn a_wrong_sized_mask_is_an_error_not_a_panic() {
|
||||
let f = field();
|
||||
let err = Membership::compute(&f, &[&vec![0.0; 3]], &PriorOptions::default()).unwrap_err();
|
||||
assert_eq!(
|
||||
err,
|
||||
MembershipError::MaskSize { instance: 0, expected: 8, got: 3 }
|
||||
);
|
||||
}
|
||||
|
||||
/// Two instances, so the "different objects repel" branch is covered.
|
||||
#[test]
|
||||
fn different_instances_repel() {
|
||||
let f = field();
|
||||
let right_half = vec![0.0, 0.0, 1.0, 1.0, 0.0, 0.0, 1.0, 1.0];
|
||||
let m = Membership::compute(
|
||||
&f,
|
||||
&[&left_half(), &right_half],
|
||||
&PriorOptions::default(),
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(m.instance_of(0), Some(0));
|
||||
assert_eq!(m.instance_of(2), Some(1));
|
||||
assert!(m.affinity(0, 2) < 0.0, "two different objects should repel");
|
||||
assert!(m.affinity(2, 3) > 0.0, "same object should attract");
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,709 @@
|
||||
//! Semantic segmentation — arm B (S15, docs/segmentation.md §4).
|
||||
//!
|
||||
//! Runs a YOLO instance-segmentation graph over a proxy-resolution image and
|
||||
//! returns the instances it found: a class, a score, a box, and a soft mask
|
||||
//! each. [`crate::prior`] is what turns those into a merge prior over the
|
||||
//! watershed hierarchy; nothing here knows about regions.
|
||||
//!
|
||||
//! # This is instance segmentation, not semantic segmentation
|
||||
//!
|
||||
//! §4 of the spec assumed a *semantic* model — a full partition of the image
|
||||
//! into 150 ADE20K classes, sky and vegetation among them. The model that
|
||||
//! actually exists is COCO-trained and *instance*-based, and the difference is
|
||||
//! not cosmetic:
|
||||
//!
|
||||
//! - **It does not partition the image.** It finds objects. Most pixels in a
|
||||
//! landscape belong to no instance at all, and that is not a failure — there
|
||||
//! is no COCO class for "hillside".
|
||||
//! - **It separates two people**, where a semantic model would hand back one
|
||||
//! "person" area covering both. For selecting a subject this is the better
|
||||
//! behaviour, and it is worth being glad of rather than working around.
|
||||
//!
|
||||
//! So arm B here contributes *subjects*, and the watershed contributes
|
||||
//! everything else. See `models/LICENCE.md` for why no ADE20K variant is
|
||||
//! shipped instead.
|
||||
//!
|
||||
//! # Cost, and where it may run
|
||||
//!
|
||||
//! ~470 ms for one 640×640 inference on the reference desktop's CPU, pure Rust
|
||||
//! via tract. That is a **once-per-image background precompute** and nothing
|
||||
//! else: it must never sit on the frame path (ARCH §6.1), and the interactive
|
||||
//! operations it enables — click a subject, grow a selection — read its cached
|
||||
//! output rather than re-running it.
|
||||
|
||||
use std::sync::Arc;
|
||||
|
||||
use ndarray::{Array4, ArrayView2, ArrayView3};
|
||||
|
||||
use crate::SegmentError;
|
||||
|
||||
/// The graph's fixed input edge, in pixels.
|
||||
///
|
||||
/// **Fixed, not configurable.** tract cannot parse the dynamic-shape export of
|
||||
/// this model — it fails shape inference on the neck's `Concat` — so the graph
|
||||
/// ships with its input baked to one square size. Everything else in this
|
||||
/// module, letterboxing and tiling alike, exists to fit arbitrary images
|
||||
/// through that fixed window.
|
||||
pub const INPUT_EDGE: usize = 640;
|
||||
|
||||
/// Detections per forward pass, from the graph's output shape `[1, 300, 38]`.
|
||||
const MAX_DETECTIONS: usize = 300;
|
||||
/// Mask prototypes, from `[1, 32, 160, 160]`.
|
||||
const PROTOTYPES: usize = 32;
|
||||
/// `4` box + `1` score + `1` class + `PROTOTYPES` coefficients.
|
||||
const DETECTION_STRIDE: usize = 6 + PROTOTYPES;
|
||||
/// Prototype masks come out at a quarter of the input edge.
|
||||
const PROTO_STRIDE: usize = 4;
|
||||
|
||||
/// How the image is presented to a fixed-shape graph.
|
||||
#[derive(Debug, Clone, Copy, PartialEq)]
|
||||
pub enum Tiling {
|
||||
/// One inference over the whole frame, letterboxed into the square input.
|
||||
///
|
||||
/// The default, and the right default: a photographic subject is usually
|
||||
/// *large* in frame, which is the case whole-image inference handles best
|
||||
/// and the case tiling helps least.
|
||||
Whole,
|
||||
/// Cover the frame with overlapping fixed-size windows.
|
||||
///
|
||||
/// Buys resolution — a subject 200 px across in a 1600 px proxy reaches
|
||||
/// the model at 200 px rather than at 80 — and costs one inference per
|
||||
/// tile. Worth it for a small subject in a large frame (a bird against
|
||||
/// sky, a figure in a landscape) and wasteful otherwise.
|
||||
///
|
||||
/// `overlap` is the fraction of a tile shared with its neighbour, which
|
||||
/// has to exceed zero or a subject sitting on a seam is cut in half by
|
||||
/// both tiles and recognised by neither.
|
||||
Grid { overlap: f32 },
|
||||
}
|
||||
|
||||
impl Default for Tiling {
|
||||
fn default() -> Self {
|
||||
Self::Whole
|
||||
}
|
||||
}
|
||||
|
||||
/// How the semantic pass is tuned.
|
||||
#[derive(Debug, Clone, Copy, PartialEq)]
|
||||
pub struct SemanticOptions {
|
||||
/// Minimum detection score to keep.
|
||||
///
|
||||
/// Deliberately low. A false positive costs a spurious entry in a list the
|
||||
/// user is choosing from; a false negative costs a subject that cannot be
|
||||
/// selected at all, which is the worse failure for a selection tool.
|
||||
pub confidence: f32,
|
||||
/// Mask probability above which a pixel is inside the instance.
|
||||
pub mask_threshold: f32,
|
||||
pub tiling: Tiling,
|
||||
/// Mask IoU above which two detections from different tiles are judged to
|
||||
/// be the same object. Unused when [`Tiling::Whole`].
|
||||
pub merge_iou: f32,
|
||||
}
|
||||
|
||||
impl Default for SemanticOptions {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
confidence: 0.25,
|
||||
mask_threshold: 0.5,
|
||||
tiling: Tiling::Whole,
|
||||
merge_iou: 0.55,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// One detected object.
|
||||
#[derive(Debug, Clone, PartialEq)]
|
||||
pub struct Instance {
|
||||
pub class_id: u16,
|
||||
pub class_name: Arc<str>,
|
||||
pub score: f32,
|
||||
/// Bounding box in **source image** pixels: `(x0, y0, x1, y1)`.
|
||||
pub bbox: (f32, f32, f32, f32),
|
||||
/// Per-pixel coverage over the whole source image, row-major, `0.0..=1.0`.
|
||||
///
|
||||
/// Soft rather than binary because arm C weights merges by it, and a hard
|
||||
/// threshold there would throw away exactly the confidence information
|
||||
/// that makes a prior a prior rather than a decision.
|
||||
pub mask: Vec<f32>,
|
||||
pub width: usize,
|
||||
pub height: usize,
|
||||
}
|
||||
|
||||
impl Instance {
|
||||
/// Fraction of this instance's mass inside a set of pixels.
|
||||
pub fn coverage(&self, pixels: impl Iterator<Item = usize>) -> f32 {
|
||||
let mut inside = 0.0;
|
||||
let mut n = 0usize;
|
||||
for p in pixels {
|
||||
inside += self.mask.get(p).copied().unwrap_or(0.0);
|
||||
n += 1;
|
||||
}
|
||||
if n == 0 {
|
||||
0.0
|
||||
} else {
|
||||
inside / n as f32
|
||||
}
|
||||
}
|
||||
|
||||
fn iou(&self, other: &Self, threshold: f32) -> f32 {
|
||||
let mut inter = 0usize;
|
||||
let mut union = 0usize;
|
||||
for (a, b) in self.mask.iter().zip(&other.mask) {
|
||||
let (a, b) = (*a >= threshold, *b >= threshold);
|
||||
inter += usize::from(a && b);
|
||||
union += usize::from(a || b);
|
||||
}
|
||||
if union == 0 {
|
||||
0.0
|
||||
} else {
|
||||
inter as f32 / union as f32
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// A loaded segmentation model.
|
||||
///
|
||||
/// Holds an `ort` session, so it is neither `Clone` nor cheap to build —
|
||||
/// construct once and keep it. Loading is ~50 ms.
|
||||
pub struct SemanticModel {
|
||||
session: ort::session::Session,
|
||||
classes: Vec<Arc<str>>,
|
||||
}
|
||||
|
||||
/// The weights that ship with this crate (`models/`, AGPL — see LICENCE.md).
|
||||
///
|
||||
/// Embedded rather than read from a path because Android hands the app no
|
||||
/// filesystem location to read from (ARCH §6.9) — the same reasoning that has
|
||||
/// the Lensfun database shipping inside its crate.
|
||||
#[cfg(feature = "embedded-model")]
|
||||
const EMBEDDED_MODEL: &[u8] = include_bytes!("../models/yolo26n-seg.onnx");
|
||||
#[cfg(feature = "embedded-model")]
|
||||
const EMBEDDED_CLASSES: &str = include_str!("../models/yolo26n-seg.classes.json");
|
||||
|
||||
impl SemanticModel {
|
||||
/// Load the model that ships with this crate.
|
||||
#[cfg(feature = "embedded-model")]
|
||||
pub fn embedded() -> Result<Self, SegmentError> {
|
||||
Self::from_bytes(EMBEDDED_MODEL, parse_classes(EMBEDDED_CLASSES))
|
||||
}
|
||||
|
||||
/// Load a model from an ONNX file, with `classes` supplying its vocabulary.
|
||||
///
|
||||
/// The vocabulary is a parameter rather than a constant so that swapping in
|
||||
/// a model with different classes — an ADE20K stuff model, say — is a data
|
||||
/// change rather than a code change.
|
||||
pub fn from_path(
|
||||
path: impl AsRef<std::path::Path>,
|
||||
classes: Vec<Arc<str>>,
|
||||
) -> Result<Self, SegmentError> {
|
||||
let bytes = std::fs::read(path).map_err(SegmentError::ModelRead)?;
|
||||
Self::from_bytes(&bytes, classes)
|
||||
}
|
||||
|
||||
pub fn from_bytes(bytes: &[u8], classes: Vec<Arc<str>>) -> Result<Self, SegmentError> {
|
||||
// Idempotent, and it must happen before any other `ort` call: with
|
||||
// `alternative-backend` there is no linked runtime to fall back on, so
|
||||
// an un-set API is a panic rather than a slow path.
|
||||
install_backend();
|
||||
|
||||
let session = ort::session::Session::builder()
|
||||
.map_err(SegmentError::Inference)?
|
||||
.commit_from_memory(bytes)
|
||||
.map_err(SegmentError::Inference)?;
|
||||
|
||||
Ok(Self { session, classes })
|
||||
}
|
||||
|
||||
pub fn classes(&self) -> &[Arc<str>] {
|
||||
&self.classes
|
||||
}
|
||||
|
||||
/// Find the objects in an image.
|
||||
///
|
||||
/// `rgb` is tightly packed `f32` RGB in `0.0..=1.0`, row-major, three
|
||||
/// components per pixel — the same linear-ish proxy the watershed reads,
|
||||
/// so both arms describe the same picture.
|
||||
pub fn detect(
|
||||
&mut self,
|
||||
rgb: &[f32],
|
||||
width: usize,
|
||||
height: usize,
|
||||
options: &SemanticOptions,
|
||||
) -> Result<Vec<Instance>, SegmentError> {
|
||||
if width == 0 || height == 0 {
|
||||
return Ok(Vec::new());
|
||||
}
|
||||
if rgb.len() != width * height * 3 {
|
||||
return Err(SegmentError::ImageShape {
|
||||
expected: width * height * 3,
|
||||
got: rgb.len(),
|
||||
});
|
||||
}
|
||||
|
||||
let windows = self.windows(width, height, options.tiling);
|
||||
let mut found: Vec<Instance> = Vec::new();
|
||||
|
||||
for window in &windows {
|
||||
let batch = self.run_window(rgb, width, height, window, options)?;
|
||||
merge_into(&mut found, batch, options);
|
||||
}
|
||||
|
||||
// Strongest first: this list is offered to a person as "what did you
|
||||
// mean", and the most confident guess belongs at the top.
|
||||
found.sort_by(|a, b| b.score.total_cmp(&a.score));
|
||||
Ok(found)
|
||||
}
|
||||
|
||||
/// The source-space rectangles each inference covers.
|
||||
fn windows(&self, width: usize, height: usize, tiling: Tiling) -> Vec<Window> {
|
||||
match tiling {
|
||||
Tiling::Whole => vec![Window {
|
||||
x: 0.0,
|
||||
y: 0.0,
|
||||
w: width as f32,
|
||||
h: height as f32,
|
||||
}],
|
||||
Tiling::Grid { overlap } => {
|
||||
// A tile covers a square of source pixels whose edge is the
|
||||
// shorter image dimension clamped to something the model can
|
||||
// still see detail in. Below that the tiling is pointless —
|
||||
// the window is already smaller than the input.
|
||||
let edge = (width.min(height) as f32).min(INPUT_EDGE as f32 * 1.5);
|
||||
let overlap = overlap.clamp(0.0, 0.9);
|
||||
let stride = (edge * (1.0 - overlap)).max(1.0);
|
||||
|
||||
let mut windows = Vec::new();
|
||||
for gy in 0..steps(height as f32, edge, stride) {
|
||||
for gx in 0..steps(width as f32, edge, stride) {
|
||||
// Last row and column are pulled back inside the frame
|
||||
// rather than padded, so no inference is spent on
|
||||
// blank margin.
|
||||
let x = (gx as f32 * stride).min((width as f32 - edge).max(0.0));
|
||||
let y = (gy as f32 * stride).min((height as f32 - edge).max(0.0));
|
||||
windows.push(Window {
|
||||
x,
|
||||
y,
|
||||
w: edge.min(width as f32),
|
||||
h: edge.min(height as f32),
|
||||
});
|
||||
}
|
||||
}
|
||||
windows
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn run_window(
|
||||
&mut self,
|
||||
rgb: &[f32],
|
||||
width: usize,
|
||||
height: usize,
|
||||
window: &Window,
|
||||
options: &SemanticOptions,
|
||||
) -> Result<Vec<Instance>, SegmentError> {
|
||||
let letterbox = Letterbox::fit(window.w, window.h);
|
||||
let input = letterbox.sample(rgb, width, height, window);
|
||||
|
||||
let outputs = self
|
||||
.session
|
||||
.run(ort::inputs![
|
||||
ort::value::Tensor::from_array(input).map_err(SegmentError::Inference)?
|
||||
])
|
||||
.map_err(SegmentError::Inference)?;
|
||||
|
||||
let (det_shape, det) = outputs[0]
|
||||
.try_extract_tensor::<f32>()
|
||||
.map_err(SegmentError::Inference)?;
|
||||
let (proto_shape, proto) = outputs[1]
|
||||
.try_extract_tensor::<f32>()
|
||||
.map_err(SegmentError::Inference)?;
|
||||
|
||||
// The decoder reads fixed column offsets out of each row, so a row of
|
||||
// an unexpected width means a model this code cannot read — a
|
||||
// different class count, a different prototype count, a detect-only
|
||||
// export. Caught here as an error rather than downstream as garbage
|
||||
// boxes, because garbage boxes look like a bad model rather than a
|
||||
// wrong one.
|
||||
if det_shape[2] as usize != DETECTION_STRIDE {
|
||||
return Err(SegmentError::OutputShape("detections"));
|
||||
}
|
||||
|
||||
let detections = ArrayView2::from_shape(
|
||||
(det_shape[1] as usize, det_shape[2] as usize),
|
||||
&det[..(det_shape[1] * det_shape[2]) as usize],
|
||||
)
|
||||
.map_err(|_| SegmentError::OutputShape("detections"))?;
|
||||
|
||||
let (pc, ph, pw) = (
|
||||
proto_shape[1] as usize,
|
||||
proto_shape[2] as usize,
|
||||
proto_shape[3] as usize,
|
||||
);
|
||||
let protos = ArrayView3::from_shape((pc, ph, pw), &proto[..pc * ph * pw])
|
||||
.map_err(|_| SegmentError::OutputShape("prototypes"))?;
|
||||
|
||||
// `&self.classes` rather than `self.decode(..)`: `outputs` holds a
|
||||
// mutable borrow of `self.session` until it drops, and a method call
|
||||
// would borrow all of `self`. Borrowing the two fields separately is
|
||||
// what the borrow checker will actually allow here.
|
||||
Ok(decode(
|
||||
&self.classes,
|
||||
detections,
|
||||
protos,
|
||||
&letterbox,
|
||||
window,
|
||||
width,
|
||||
height,
|
||||
options,
|
||||
))
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
/// Turn one forward pass into instances in source-image space.
|
||||
///
|
||||
/// YOLO26 is **NMS-free**: the head emits a fixed 300 slots already suppressed
|
||||
/// and score-ordered, so there is no non-maximum suppression to implement
|
||||
/// here. Only the cross-*tile* duplicates need merging, and that is
|
||||
/// [`merge_into`]'s job.
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
fn decode(
|
||||
classes: &[Arc<str>],
|
||||
detections: ArrayView2<f32>,
|
||||
protos: ArrayView3<f32>,
|
||||
letterbox: &Letterbox,
|
||||
window: &Window,
|
||||
width: usize,
|
||||
height: usize,
|
||||
options: &SemanticOptions,
|
||||
) -> Vec<Instance> {
|
||||
let (ph, pw) = (protos.shape()[1], protos.shape()[2]);
|
||||
let mut out = Vec::new();
|
||||
|
||||
for d in 0..detections.shape()[0].min(MAX_DETECTIONS) {
|
||||
let row = detections.row(d);
|
||||
let score = row[4];
|
||||
// Score-ordered, so the first miss ends the useful part of the batch
|
||||
// and the remaining slots are padding.
|
||||
if score < options.confidence {
|
||||
break;
|
||||
}
|
||||
|
||||
let class_id = row[5] as u16;
|
||||
let Some(class_name) = classes.get(class_id as usize).cloned() else {
|
||||
continue;
|
||||
};
|
||||
|
||||
// Box is in letterboxed input space; undo the letterbox and the window
|
||||
// offset to land in source pixels.
|
||||
let bbox = letterbox.to_source(row[0], row[1], row[2], row[3], window);
|
||||
|
||||
let coeffs: Vec<f32> = row.iter().skip(6).take(PROTOTYPES).copied().collect();
|
||||
let mask = assemble_mask(
|
||||
&coeffs, protos, ph, pw, letterbox, window, &bbox, width, height, options,
|
||||
);
|
||||
|
||||
out.push(Instance {
|
||||
class_id,
|
||||
class_name,
|
||||
score,
|
||||
bbox,
|
||||
mask,
|
||||
width,
|
||||
height,
|
||||
});
|
||||
}
|
||||
|
||||
out
|
||||
}
|
||||
|
||||
/// A source-space rectangle fed through one inference.
|
||||
#[derive(Debug, Clone, Copy)]
|
||||
struct Window {
|
||||
x: f32,
|
||||
y: f32,
|
||||
w: f32,
|
||||
h: f32,
|
||||
}
|
||||
|
||||
/// The scale-and-pad that fits an arbitrary rectangle into the square input.
|
||||
#[derive(Debug, Clone, Copy)]
|
||||
struct Letterbox {
|
||||
/// Input pixels per source pixel.
|
||||
scale: f32,
|
||||
pad_x: f32,
|
||||
pad_y: f32,
|
||||
}
|
||||
|
||||
impl Letterbox {
|
||||
fn fit(w: f32, h: f32) -> Self {
|
||||
let scale = (INPUT_EDGE as f32 / w).min(INPUT_EDGE as f32 / h);
|
||||
Self {
|
||||
scale,
|
||||
pad_x: (INPUT_EDGE as f32 - w * scale) * 0.5,
|
||||
pad_y: (INPUT_EDGE as f32 - h * scale) * 0.5,
|
||||
}
|
||||
}
|
||||
|
||||
/// Resample a source window into the graph's `[1, 3, 640, 640]` input.
|
||||
///
|
||||
/// Bilinear, and grey (`0.5`) in the padding — the value the network sees
|
||||
/// least as an edge, where black would draw a hard border across the frame
|
||||
/// and invite a detection along it.
|
||||
fn sample(&self, rgb: &[f32], width: usize, height: usize, window: &Window) -> Array4<f32> {
|
||||
let mut input = Array4::<f32>::from_elem((1, 3, INPUT_EDGE, INPUT_EDGE), 0.5);
|
||||
|
||||
for iy in 0..INPUT_EDGE {
|
||||
let sy = (iy as f32 + 0.5 - self.pad_y) / self.scale + window.y;
|
||||
if sy < window.y || sy >= window.y + window.h {
|
||||
continue;
|
||||
}
|
||||
for ix in 0..INPUT_EDGE {
|
||||
let sx = (ix as f32 + 0.5 - self.pad_x) / self.scale + window.x;
|
||||
if sx < window.x || sx >= window.x + window.w {
|
||||
continue;
|
||||
}
|
||||
let (x0, y0) = (sx.floor(), sy.floor());
|
||||
let (fx, fy) = (sx - x0, sy - y0);
|
||||
let x0 = (x0 as isize).clamp(0, width as isize - 1) as usize;
|
||||
let y0 = (y0 as isize).clamp(0, height as isize - 1) as usize;
|
||||
let x1 = (x0 + 1).min(width - 1);
|
||||
let y1 = (y0 + 1).min(height - 1);
|
||||
|
||||
for c in 0..3 {
|
||||
let at = |x: usize, y: usize| rgb[(y * width + x) * 3 + c];
|
||||
let top = at(x0, y0) * (1.0 - fx) + at(x1, y0) * fx;
|
||||
let bot = at(x0, y1) * (1.0 - fx) + at(x1, y1) * fx;
|
||||
input[[0, c, iy, ix]] = top * (1.0 - fy) + bot * fy;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
input
|
||||
}
|
||||
|
||||
/// Input-space box back to source pixels.
|
||||
fn to_source(&self, x0: f32, y0: f32, x1: f32, y1: f32, w: &Window) -> (f32, f32, f32, f32) {
|
||||
(
|
||||
(x0 - self.pad_x) / self.scale + w.x,
|
||||
(y0 - self.pad_y) / self.scale + w.y,
|
||||
(x1 - self.pad_x) / self.scale + w.x,
|
||||
(y1 - self.pad_y) / self.scale + w.y,
|
||||
)
|
||||
}
|
||||
|
||||
/// Source pixel to prototype-grid coordinates.
|
||||
fn to_proto(&self, sx: f32, sy: f32, w: &Window) -> (f32, f32) {
|
||||
(
|
||||
((sx - w.x) * self.scale + self.pad_x) / PROTO_STRIDE as f32,
|
||||
((sy - w.y) * self.scale + self.pad_y) / PROTO_STRIDE as f32,
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
/// Combine the prototype masks by one detection's coefficients.
|
||||
///
|
||||
/// The mask is `sigmoid(Σ coeff_k · proto_k)`, sampled straight into source
|
||||
/// resolution and **clipped to the detection's box** — YOLO's prototypes are
|
||||
/// global, so a coefficient set that describes a dog also lights up faintly on
|
||||
/// a second dog elsewhere in the frame. The box is what makes an instance mask
|
||||
/// an *instance* mask.
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
fn assemble_mask(
|
||||
coeffs: &[f32],
|
||||
protos: ArrayView3<f32>,
|
||||
ph: usize,
|
||||
pw: usize,
|
||||
letterbox: &Letterbox,
|
||||
window: &Window,
|
||||
bbox: &(f32, f32, f32, f32),
|
||||
width: usize,
|
||||
height: usize,
|
||||
options: &SemanticOptions,
|
||||
) -> Vec<f32> {
|
||||
let mut mask = vec![0.0f32; width * height];
|
||||
|
||||
let x0 = bbox.0.floor().max(0.0) as usize;
|
||||
let y0 = bbox.1.floor().max(0.0) as usize;
|
||||
let x1 = (bbox.2.ceil() as usize).min(width);
|
||||
let y1 = (bbox.3.ceil() as usize).min(height);
|
||||
|
||||
for y in y0..y1 {
|
||||
for x in x0..x1 {
|
||||
let (gx, gy) = letterbox.to_proto(x as f32 + 0.5, y as f32 + 0.5, window);
|
||||
if gx < 0.0 || gy < 0.0 || gx >= pw as f32 || gy >= ph as f32 {
|
||||
continue;
|
||||
}
|
||||
|
||||
// Bilinear over the prototype grid: nearest-neighbour here shows
|
||||
// as visible 4-pixel stair-stepping on the mask edge.
|
||||
let (fx0, fy0) = (gx.floor(), gy.floor());
|
||||
let (fx, fy) = (gx - fx0, gy - fy0);
|
||||
let (gx0, gy0) = (fx0 as usize, fy0 as usize);
|
||||
let (gx1, gy1) = ((gx0 + 1).min(pw - 1), (gy0 + 1).min(ph - 1));
|
||||
|
||||
let mut acc = 0.0;
|
||||
for (k, &c) in coeffs.iter().enumerate().take(PROTOTYPES.min(protos.shape()[0])) {
|
||||
if c == 0.0 {
|
||||
continue;
|
||||
}
|
||||
let p = protos.index_axis(ndarray::Axis(0), k);
|
||||
let top = p[[gy0, gx0]] * (1.0 - fx) + p[[gy0, gx1]] * fx;
|
||||
let bot = p[[gy1, gx0]] * (1.0 - fx) + p[[gy1, gx1]] * fx;
|
||||
acc += c * (top * (1.0 - fy) + bot * fy);
|
||||
}
|
||||
|
||||
let p = 1.0 / (1.0 + (-acc).exp());
|
||||
if p >= options.mask_threshold * 0.5 {
|
||||
mask[y * width + x] = p;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
mask
|
||||
}
|
||||
|
||||
/// Fold one tile's detections into the running set.
|
||||
///
|
||||
/// Only needed for [`Tiling::Grid`]: an object straddling a seam is seen by
|
||||
/// both tiles, and without this it would appear twice in the list a person
|
||||
/// chooses from. Keeps the higher-scoring copy, which is generally the tile
|
||||
/// that saw more of the object.
|
||||
fn merge_into(found: &mut Vec<Instance>, batch: Vec<Instance>, options: &SemanticOptions) {
|
||||
for candidate in batch {
|
||||
let duplicate = found.iter_mut().find(|existing| {
|
||||
existing.class_id == candidate.class_id
|
||||
&& existing.iou(&candidate, options.mask_threshold) >= options.merge_iou
|
||||
});
|
||||
|
||||
match duplicate {
|
||||
Some(existing) if existing.score < candidate.score => *existing = candidate,
|
||||
Some(_) => {}
|
||||
None => found.push(candidate),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn steps(extent: f32, edge: f32, stride: f32) -> usize {
|
||||
if extent <= edge {
|
||||
1
|
||||
} else {
|
||||
(((extent - edge) / stride).ceil() as usize) + 1
|
||||
}
|
||||
}
|
||||
|
||||
/// Point `ort` at tract, exactly once per process.
|
||||
fn install_backend() {
|
||||
use std::sync::Once;
|
||||
static ONCE: Once = Once::new();
|
||||
ONCE.call_once(|| {
|
||||
// Returns false if an API was already installed, which is not an error
|
||||
// — it means something else got here first, and there is only one
|
||||
// backend compiled in for it to have chosen.
|
||||
let _ = ort::set_api(ort_tract::api());
|
||||
});
|
||||
}
|
||||
|
||||
/// Read the class list written beside the model by `tools/export-seg-model.sh`.
|
||||
///
|
||||
/// A deliberately small hand-rolled reader for a flat array of strings, rather
|
||||
/// than a `serde_json` dependency for one file of one shape that this
|
||||
/// repository generates itself.
|
||||
pub fn parse_classes(json: &str) -> Vec<Arc<str>> {
|
||||
let mut out = Vec::new();
|
||||
let mut chars = json.chars().peekable();
|
||||
while let Some(c) = chars.next() {
|
||||
if c != '"' {
|
||||
continue;
|
||||
}
|
||||
let mut name = String::new();
|
||||
while let Some(c) = chars.next() {
|
||||
match c {
|
||||
'"' => break,
|
||||
'\\' => name.extend(chars.next()),
|
||||
_ => name.push(c),
|
||||
}
|
||||
}
|
||||
out.push(name.into());
|
||||
}
|
||||
out
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn classes_parse_from_the_exported_json() {
|
||||
let parsed = parse_classes("[\n \"person\",\n \"bicycle\",\n \"car\"\n]");
|
||||
assert_eq!(&*parsed[0], "person");
|
||||
assert_eq!(&*parsed[2], "car");
|
||||
assert_eq!(parsed.len(), 3);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn letterbox_round_trips_a_landscape_window() {
|
||||
let window = Window { x: 0.0, y: 0.0, w: 1600.0, h: 1067.0 };
|
||||
let lb = Letterbox::fit(window.w, window.h);
|
||||
|
||||
// A source point maps into the input and back to where it started.
|
||||
let (ix, iy) = (
|
||||
(400.0 - window.x) * lb.scale + lb.pad_x,
|
||||
(300.0 - window.y) * lb.scale + lb.pad_y,
|
||||
);
|
||||
let (sx, sy, _, _) = lb.to_source(ix, iy, 0.0, 0.0, &window);
|
||||
assert!((sx - 400.0).abs() < 1e-3, "x round-trip: {sx}");
|
||||
assert!((sy - 300.0).abs() < 1e-3, "y round-trip: {sy}");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn letterbox_pads_the_short_axis_only() {
|
||||
let lb = Letterbox::fit(1600.0, 1067.0);
|
||||
assert!(lb.pad_x.abs() < 1e-3, "wide image should not pad in x");
|
||||
assert!(lb.pad_y > 100.0, "wide image should pad in y: {}", lb.pad_y);
|
||||
}
|
||||
|
||||
/// The seam case tiling exists for, and the one it must not double-count.
|
||||
#[test]
|
||||
fn grid_tiling_covers_the_frame_with_overlap() {
|
||||
let model_windows = |w: usize, h: usize, overlap: f32| {
|
||||
// `windows` needs no session state, so exercise it through a
|
||||
// stand-in rather than loading 11 MB of weights in a unit test.
|
||||
let edge = (w.min(h) as f32).min(INPUT_EDGE as f32 * 1.5);
|
||||
let stride = (edge * (1.0 - overlap)).max(1.0);
|
||||
(steps(w as f32, edge, stride), steps(h as f32, edge, stride))
|
||||
};
|
||||
|
||||
let (cols, rows) = model_windows(1600, 1067, 0.25);
|
||||
assert!(cols >= 2, "a 1600px frame needs more than one column");
|
||||
assert_eq!(rows, 2, "1067px against a 960px tile is two rows");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn a_square_frame_is_a_single_tile() {
|
||||
assert_eq!(steps(640.0, 640.0, 480.0), 1);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn duplicate_detections_across_tiles_collapse_to_the_stronger() {
|
||||
let opts = SemanticOptions::default();
|
||||
let make = |score: f32, on: bool| Instance {
|
||||
class_id: 0,
|
||||
class_name: "person".into(),
|
||||
score,
|
||||
bbox: (0.0, 0.0, 2.0, 2.0),
|
||||
mask: if on { vec![1.0; 4] } else { vec![0.0; 4] },
|
||||
width: 2,
|
||||
height: 2,
|
||||
};
|
||||
|
||||
let mut found = vec![make(0.6, true)];
|
||||
merge_into(&mut found, vec![make(0.9, true)], &opts);
|
||||
assert_eq!(found.len(), 1, "same object seen twice is one instance");
|
||||
assert_eq!(found[0].score, 0.9, "the more confident tile wins");
|
||||
|
||||
// A disjoint mask is a different object and must survive.
|
||||
merge_into(&mut found, vec![make(0.5, false)], &opts);
|
||||
assert_eq!(found.len(), 2);
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user