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:
2026-08-22 08:39:16 +02:00
parent ecd6df686c
commit 0da8271836
18 changed files with 2287 additions and 19 deletions
+429
View File
@@ -0,0 +1,429 @@
//! Region adjacency and the merge hierarchy — arm A's CPU half (S15).
//!
//! **Deliberately device-free.** Nothing here touches wgpu, and that is the
//! point: the hierarchy is where arm A's behaviour actually lives, so it has
//! to be assertable on hand-built inputs rather than only on whatever a GPU
//! happened to produce (ARCH §6.5a). Every test below runs on CI machines
//! with no adapter.
//!
//! It is also why this is a legitimate CPU stage rather than a violation of
//! ARCH §6.1. The watershed's basins are found on the GPU, over pixels; what
//! follows runs on the **region adjacency graph**, which for a 2 MP proxy is
//! a few thousand nodes. Pixels stay on the GPU, the graph is CPU-side —
//! the same split ARCH §3.4 already draws for the edit graph.
//!
//! # Why a hierarchy rather than one partition
//!
//! A single segmentation has exactly one granularity, and no threshold is
//! right for both "her eye" and "her face". So the watershed deliberately
//! *over*-segments, and the merge order recorded here is what lets one
//! interaction — a scroll, a drag — walk from a fragment up to the object
//! that contains it.
use std::collections::HashMap;
/// One boundary between two adjacent regions.
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct Edge {
pub a: u32,
pub b: u32,
/// The **saddle**: the lowest point on the ridge separating the two
/// basins, which is the height at which water joining them would first
/// spill over.
///
/// The minimum over the shared boundary, not the mean or the maximum. A
/// long boundary that is strong everywhere except one weak gap describes
/// two regions that should merge early — the gap is exactly where a
/// person would say the edge fails.
pub saddle: f32,
}
/// A partition of the image into labelled regions, plus how they adjoin.
///
/// The shared interface from docs/segmentation.md §2: arm A produces this
/// from a watershed, arm B would produce it from a class map, and the
/// consumers above cannot tell which.
#[derive(Debug, Clone, PartialEq)]
pub struct RegionField {
pub width: usize,
pub height: usize,
/// Region id per pixel, compacted to `0..region_count`.
pub labels: Vec<u32>,
pub region_count: usize,
/// Deduplicated, `a < b`, sorted for reproducibility.
pub adjacency: Vec<Edge>,
}
impl RegionField {
/// Build from the watershed's raw output.
///
/// `roots` holds, per pixel, the linear index of its basin root — what
/// the pointer-jumping pass converges to. Those indices are sparse and
/// arbitrary, so they are compacted here to `0..region_count` in order of
/// first appearance, which makes them stable to serialise and cheap to
/// index.
pub fn from_roots(roots: &[u32], gradient: &[f32], width: usize, height: usize) -> Self {
assert_eq!(roots.len(), width * height, "roots must cover every pixel");
assert_eq!(
gradient.len(),
width * height,
"gradient must cover every pixel"
);
// Compact in raster order rather than by hashing the root value, so
// the same image always yields the same numbering (M5).
let mut compact: HashMap<u32, u32> = HashMap::new();
let mut labels = vec![0u32; roots.len()];
for (i, &root) in roots.iter().enumerate() {
let next = compact.len() as u32;
labels[i] = *compact.entry(root).or_insert(next);
}
let region_count = compact.len();
// Saddles, over 4-neighbour adjacency. Crossing a boundary means
// climbing to the higher of the two pixels, so a pair's pass height
// is the max; the region pair's saddle is the min over all its pairs.
let mut saddles: HashMap<(u32, u32), f32> = HashMap::new();
let note = |p: usize, q: usize, saddles: &mut HashMap<(u32, u32), f32>| {
let (la, lb) = (labels[p], labels[q]);
if la == lb {
return;
}
let key = (la.min(lb), la.max(lb));
let pass = gradient[p].max(gradient[q]);
saddles
.entry(key)
.and_modify(|s| {
if pass < *s {
*s = pass;
}
})
.or_insert(pass);
};
for y in 0..height {
for x in 0..width {
let i = y * width + x;
if x + 1 < width {
note(i, i + 1, &mut saddles);
}
if y + 1 < height {
note(i, i + width, &mut saddles);
}
}
}
let mut adjacency: Vec<Edge> = saddles
.into_iter()
.map(|((a, b), saddle)| Edge { a, b, saddle })
.collect();
// Sorted by (saddle, a, b) with `total_cmp` rather than `partial_cmp`:
// a total order over floats, so the sequence cannot depend on how the
// hash map happened to iterate. The merge order below is derived from
// this, and an unstable merge order is an unstable label field.
adjacency.sort_by(|x, y| {
x.saddle
.total_cmp(&y.saddle)
.then(x.a.cmp(&y.a))
.then(x.b.cmp(&y.b))
});
Self {
width,
height,
labels,
region_count,
adjacency,
}
}
/// Apply a cut's grouping to every pixel, for display or masking.
pub fn apply(&self, grouping: &[u32]) -> Vec<u32> {
self.labels.iter().map(|&l| grouping[l as usize]).collect()
}
}
/// One merge in the hierarchy.
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct Merge {
pub a: u32,
pub b: u32,
pub saddle: f32,
}
/// The merge order over a [`RegionField`] — the whole hierarchy.
///
/// Kruskal over the adjacency edges: sort by saddle, union-find, and record
/// each union that actually joined two distinct sets. The recording *is* the
/// dendrogram, which is why the multiscale part costs one `Vec` rather than a
/// second algorithm.
#[derive(Debug, Clone, PartialEq)]
pub struct MergeTree {
/// In merge order, so saddles are non-decreasing.
pub merges: Vec<Merge>,
pub region_count: usize,
}
impl MergeTree {
pub fn build(field: &RegionField) -> Self {
let mut uf = UnionFind::new(field.region_count);
let mut merges = Vec::new();
// `field.adjacency` is already sorted by (saddle, a, b).
for e in &field.adjacency {
if uf.union(e.a as usize, e.b as usize) {
merges.push(Merge {
a: e.a,
b: e.b,
saddle: e.saddle,
});
}
}
Self {
merges,
region_count: field.region_count,
}
}
/// The grouping after applying every merge weaker than `threshold`.
///
/// Returns one group id per original region, compacted to `0..groups`.
pub fn cut_at(&self, threshold: f32) -> Vec<u32> {
self.cut(self.merges.iter().take_while(|m| m.saddle <= threshold))
}
/// The grouping with at most `target` groups.
///
/// The interface a scroll wheel wants: thresholds are in gradient units
/// and mean nothing to anyone, where "about two hundred regions" is a
/// granularity a person can ask for.
pub fn cut_to(&self, target: usize) -> Vec<u32> {
let target = target.max(1);
let take = self.region_count.saturating_sub(target);
self.cut(self.merges.iter().take(take))
}
fn cut<'a>(&self, merges: impl Iterator<Item = &'a Merge>) -> Vec<u32> {
let mut uf = UnionFind::new(self.region_count);
for m in merges {
uf.union(m.a as usize, m.b as usize);
}
// Compact roots in region order, so a cut's numbering is stable and
// independent of the merge sequence that produced it.
let mut seen: HashMap<usize, u32> = HashMap::new();
(0..self.region_count)
.map(|r| {
let root = uf.find(r);
let next = seen.len() as u32;
*seen.entry(root).or_insert(next)
})
.collect()
}
/// How many groups `cut_to` / `cut_at` would leave, without building one.
pub fn groups_at(&self, threshold: f32) -> usize {
let merged = self.merges.iter().filter(|m| m.saddle <= threshold).count();
self.region_count - merged
}
}
struct UnionFind {
parent: Vec<usize>,
rank: Vec<u8>,
}
impl UnionFind {
fn new(n: usize) -> Self {
Self {
parent: (0..n).collect(),
rank: vec![0; n],
}
}
fn find(&mut self, mut x: usize) -> usize {
while self.parent[x] != x {
// Path halving — the compression that keeps this near-linear
// without the second pass full compression needs.
self.parent[x] = self.parent[self.parent[x]];
x = self.parent[x];
}
x
}
/// Returns whether this union joined two distinct sets.
fn union(&mut self, a: usize, b: usize) -> bool {
let (ra, rb) = (self.find(a), self.find(b));
if ra == rb {
return false;
}
// Union by rank, but with a deterministic tie-break: equal ranks
// attach the higher index under the lower. Without it the tree shape
// depends on argument order, and `find` would return a different
// representative for the same input on a different day.
let (lo, hi) = if self.rank[ra] > self.rank[rb] {
(ra, rb)
} else if self.rank[rb] > self.rank[ra] {
(rb, ra)
} else {
let (lo, hi) = (ra.min(rb), ra.max(rb));
self.rank[lo] += 1;
(lo, hi)
};
self.parent[hi] = lo;
true
}
}
#[cfg(test)]
mod tests {
use super::*;
/// A 4×2 field split down the middle, with a weak gap in the boundary.
///
/// labels 0 0 1 1 gradient 0 5 5 0
/// 0 0 1 1 0 2 2 0
///
/// The saddle is 2, not 5: the bottom row is where the ridge is lowest.
fn split_field() -> RegionField {
let roots = vec![0, 0, 2, 2, 0, 0, 2, 2];
let gradient = vec![0.0, 5.0, 5.0, 0.0, 0.0, 2.0, 2.0, 0.0];
RegionField::from_roots(&roots, &gradient, 4, 2)
}
#[test]
fn roots_compact_to_dense_ids() {
// Basin roots are pixel indices and so are sparse and arbitrary.
// Anything that indexes by region needs them dense.
let f = split_field();
assert_eq!(f.region_count, 2);
assert_eq!(f.labels, vec![0, 0, 1, 1, 0, 0, 1, 1]);
}
#[test]
fn the_saddle_is_the_lowest_pass_not_the_typical_one() {
// The property that makes the hierarchy match human judgement: two
// regions joined by one weak gap belong together, however strong the
// rest of the boundary is.
let f = split_field();
assert_eq!(f.adjacency.len(), 1);
assert_eq!(f.adjacency[0].saddle, 2.0);
}
#[test]
fn a_uniform_image_is_one_region() {
// No minima to separate, so nothing to merge — and the tree must not
// invent a merge it cannot justify.
let f = RegionField::from_roots(&[0; 16], &[0.0; 16], 4, 4);
assert_eq!(f.region_count, 1);
assert!(f.adjacency.is_empty());
assert!(MergeTree::build(&f).merges.is_empty());
}
/// Three regions in a row: 0 | 1 | 2, with the 1–2 boundary weaker.
fn chain_field() -> RegionField {
let roots = vec![0, 0, 2, 2, 4, 4];
let gradient = vec![0.0, 9.0, 0.0, 3.0, 0.0, 0.0];
RegionField::from_roots(&roots, &gradient, 6, 1)
}
#[test]
fn the_weakest_boundary_merges_first() {
// The whole basis of the granularity ladder: walking up the tree must
// dissolve the least convincing edge before a strong one.
let tree = MergeTree::build(&chain_field());
assert_eq!(tree.merges.len(), 2);
assert_eq!(tree.merges[0].saddle, 3.0);
assert_eq!(tree.merges[1].saddle, 9.0);
assert_eq!((tree.merges[0].a, tree.merges[0].b), (1, 2));
}
#[test]
fn merges_are_ordered_by_saddle() {
let tree = MergeTree::build(&chain_field());
assert!(
tree.merges.windows(2).all(|w| w[0].saddle <= w[1].saddle),
"a cut at a threshold is only meaningful if merges are ordered"
);
}
#[test]
fn a_cut_walks_from_every_region_to_one() {
let tree = MergeTree::build(&chain_field());
// The bottom of the ladder: nothing merged.
let fine = tree.cut_to(3);
assert_eq!(fine, vec![0, 1, 2]);
// The middle: the weak boundary is gone, the strong one survives.
let mid = tree.cut_to(2);
assert_eq!(mid[1], mid[2], "the weak boundary should have dissolved");
assert_ne!(mid[0], mid[1], "the strong boundary should survive");
// The top: one region.
let coarse = tree.cut_to(1);
assert_eq!(coarse, vec![0, 0, 0]);
}
#[test]
fn cut_to_is_clamped_rather_than_panicking() {
// A scroll wheel runs past both ends of the ladder, and neither end
// is an error.
let tree = MergeTree::build(&chain_field());
assert_eq!(tree.cut_to(0), tree.cut_to(1));
assert_eq!(tree.cut_to(99), vec![0, 1, 2]);
}
#[test]
fn cut_at_and_cut_to_agree() {
let tree = MergeTree::build(&chain_field());
// Above the weak saddle, below the strong one.
assert_eq!(tree.cut_at(5.0), tree.cut_to(2));
assert_eq!(tree.groups_at(5.0), 2);
}
#[test]
fn a_grouping_maps_back_onto_pixels() {
let f = chain_field();
let tree = MergeTree::build(&f);
let px = f.apply(&tree.cut_to(2));
assert_eq!(px, vec![0, 0, 1, 1, 1, 1]);
}
#[test]
fn the_same_input_gives_a_bit_identical_result() {
// M5 in miniature. The CPU half must contribute no nondeterminism of
// its own, or there is no point asking whether the GPU half does:
// hash map iteration order is the obvious way to fail this, which is
// why both the adjacency list and the cut numbering are sorted rather
// than taken as the map yields them.
let (a, b) = (chain_field(), chain_field());
assert_eq!(a, b);
let (ta, tb) = (MergeTree::build(&a), MergeTree::build(&b));
assert_eq!(ta, tb);
assert_eq!(ta.cut_to(2), tb.cut_to(2));
}
#[test]
fn a_ladder_over_many_regions_is_monotone() {
// Region counts must fall as you climb, with no level skipped —
// otherwise a scroll step could jump past the granularity someone
// wanted.
let n = 32;
let roots: Vec<u32> = (0..n).map(|i| i as u32).collect();
let gradient: Vec<f32> = (0..n).map(|i| (i % 7) as f32).collect();
let f = RegionField::from_roots(&roots, &gradient, n, 1);
let tree = MergeTree::build(&f);
let mut last = usize::MAX;
for target in (1..=n).rev() {
let groups = tree.cut_to(target).iter().max().map(|m| *m as usize + 1);
let groups = groups.unwrap_or(0);
assert_eq!(groups, target, "cut_to({target}) should leave {target}");
assert!(groups < last);
last = groups;
}
}
}
+59
View File
@@ -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),
}
+418
View File
@@ -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, &regions);
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");
}
}
+709
View File
@@ -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);
}
}