One crate names the runtime, the providers and the devices; dr-face and dr-segment ask it for a session by role. It hands ort an API table once per process — from a libonnxruntime it dlopens when the app names a directory holding one, otherwise from tract — so the Rust build stays free of C on every target and a package can install the runtime as a file (docs/inference.md §3). Sessions live in a registry behind a Model handle that holds the bytes, not the session: every use refreshes a timestamp and a reaper unloads whatever sat idle past the decay. A scan that runs the detector on each image never lets it go idle; a click in the develop view lets the segmenter go after thirty seconds; a handle used after that reloads, and reloads on a higher rung if a compiled engine has landed meanwhile. The probe walks the platform's ladder by building strict sessions and timing them against the CPU provider, caches the choice against a fingerprint of the runtime, driver, hardware and models, and compiles engines for the selected rung in the background, smallest model first. Nothing in this commit turns the native path on: the apps still run on tract until they call init with a runtime directory.
661 lines
26 KiB
Rust
661 lines
26 KiB
Rust
//! Per-category weights over the whole frame — what the scene tab grades.
|
||
//!
|
||
//! [`semantic`](crate::semantic) answers "what objects are in this picture, and
|
||
//! which pixels are each one". This module answers a different question: "how
|
||
//! much of each pixel is sky". They are not the same question and they do not
|
||
//! want the same model.
|
||
//!
|
||
//! # Why a second model rather than a second reading of the first
|
||
//!
|
||
//! The instance model is COCO-trained, and COCO is eighty classes of *things*.
|
||
//! There is no class for sky, none for foliage, none for water — the categories
|
||
//! a landscape is mostly made of. That gap is recorded in `models/LICENCE.md`
|
||
//! and it is why the scene model exists: ADE20K's 150 classes are a scene
|
||
//! parse, *stuff* included.
|
||
//!
|
||
//! Going the other way is just as impossible. A semantic model merges every
|
||
//! pixel of a class into one region, so it cannot tell three people apart, and
|
||
//! telling three people apart is exactly what clicking a subject needs. Neither
|
||
//! model substitutes for the other, which is why both ship.
|
||
//!
|
||
//! # The partition of unity, and why it is the point
|
||
//!
|
||
//! [`Scene::weight`] is not a mask per category that each independently says
|
||
//! yes or no. It is a *partition*: at every pixel the listed categories plus
|
||
//! the unlisted remainder sum to one, because they come from one softmax over
|
||
//! all 150 channels, summed within each category.
|
||
//!
|
||
//! That property is what makes feathering safe. Feather a hard label map
|
||
//! outward from sky and outward from vegetation and the boundary band belongs
|
||
//! to both, so a `+20` on sky and a `−10` on vegetation both land there and
|
||
//! every horizon acquires a visible seam. Feather a partition of unity and the
|
||
//! weights still sum to one — the band gets a blend of the two grades, which is
|
||
//! what a photographer drawing that boundary by hand would have painted.
|
||
//!
|
||
//! # Resolution, stated plainly
|
||
//!
|
||
//! The graph's logits are `[1, 150, 80, 80]`: an eighth of the input edge, and
|
||
//! that is the real spatial resolution of everything here. The stock export
|
||
//! ends with a `Resize` to 640×640 and an `ArgMax`, and
|
||
//! `tools/export-seg-model.sh` cuts both — the upsample adds no information and
|
||
//! the argmax destroys the per-class scores this module needs. [`Scene`] keeps
|
||
//! the native grid and resamples on demand ([`Scene::rasterise`]) so that the
|
||
//! coarseness is visible in the type rather than hidden behind an early
|
||
//! upsample.
|
||
//!
|
||
//! Practically: a graduated grade over sky or water is unbothered by 80×80. A
|
||
//! hard edge — a rooftop against sky at 100% zoom — will show it, and no
|
||
//! feather setting invents detail the model never had.
|
||
//!
|
||
//! # Cost
|
||
//!
|
||
//! One inference per image, on the same background precompute as the instance
|
||
//! pass and never on the frame path (ARCH §6.1). The scene tab's sliders read
|
||
//! [`Scene`] and re-run nothing.
|
||
|
||
use std::sync::Arc;
|
||
|
||
use ndarray::ArrayView3;
|
||
|
||
#[cfg(test)]
|
||
use crate::semantic::INPUT_EDGE;
|
||
use crate::semantic::{Letterbox, Window};
|
||
use crate::SegmentError;
|
||
|
||
/// Classes in the ADE20K vocabulary the scene model was trained on.
|
||
///
|
||
/// Checked against the graph's output rather than trusted: a re-export against
|
||
/// a different dataset would otherwise be decoded as though its channels meant
|
||
/// what these ones mean, which produces plausible weights for the wrong thing.
|
||
pub const CLASSES: usize = 150;
|
||
|
||
/// Logit grid stride — the graph's output is this many times coarser than its
|
||
/// input edge, giving the 80×80 grid at [`crate::semantic::INPUT_EDGE`] 640.
|
||
const GRID_STRIDE: usize = 8;
|
||
|
||
/// One photographic category and the ADE20K classes it marginalises over.
|
||
#[derive(Debug, Clone)]
|
||
pub struct Category {
|
||
pub name: Arc<str>,
|
||
/// Indices into the model's vocabulary. Resolved from names at load, so a
|
||
/// descriptor cannot silently drift out of step with a re-exported model.
|
||
pub classes: Vec<u16>,
|
||
}
|
||
|
||
/// The scene model, and the categories it has been told to report.
|
||
pub struct SceneModel {
|
||
session: dr_inference_engine::Model,
|
||
categories: Vec<Category>,
|
||
}
|
||
|
||
/// The weights that ship in `models/scene/` (AGPL — see `models/LICENCE.md`).
|
||
///
|
||
/// Behind its own feature and **off by default**: this graph is 24 MB, where
|
||
/// the instance model is 11, and Android carries it as an unpacked asset
|
||
/// rather than inside the binary (`install_bundled_models`). A desktop build
|
||
/// or a test that wants it compiled in opts in.
|
||
#[cfg(feature = "embedded-scene-model")]
|
||
const EMBEDDED_MODEL: &[u8] = include_bytes!("../../../models/scene/yolo26s-sem-ade20k.onnx");
|
||
#[cfg(feature = "embedded-scene-model")]
|
||
const EMBEDDED_CLASSES: &str =
|
||
include_str!("../../../models/scene/yolo26s-sem-ade20k.classes.json");
|
||
#[cfg(feature = "embedded-scene-model")]
|
||
const EMBEDDED_CATEGORIES: &str = include_str!("../../../models/scene/categories.txt");
|
||
|
||
impl SceneModel {
|
||
/// Load the scene model compiled into the binary.
|
||
#[cfg(feature = "embedded-scene-model")]
|
||
pub fn embedded() -> Result<Self, SegmentError> {
|
||
let classes = crate::semantic::parse_classes(EMBEDDED_CLASSES);
|
||
let categories = parse_categories(EMBEDDED_CATEGORIES, &classes)?;
|
||
Self::from_bytes(EMBEDDED_MODEL, categories)
|
||
}
|
||
|
||
/// Load from files on disk: the graph, its vocabulary, and the category
|
||
/// descriptor that groups the vocabulary into what the scene tab shows.
|
||
///
|
||
/// Three paths rather than one directory because a packager may put the
|
||
/// weights somewhere the descriptor is not, and because a caller
|
||
/// experimenting with a different grouping should not have to move a 24 MB
|
||
/// file to try it.
|
||
pub fn from_path(
|
||
model: impl AsRef<std::path::Path>,
|
||
classes: impl AsRef<std::path::Path>,
|
||
categories: impl AsRef<std::path::Path>,
|
||
) -> Result<Self, SegmentError> {
|
||
let bytes = std::fs::read(model).map_err(SegmentError::ModelRead)?;
|
||
let classes = std::fs::read_to_string(classes).map_err(SegmentError::ModelRead)?;
|
||
let categories = std::fs::read_to_string(categories).map_err(SegmentError::ModelRead)?;
|
||
let classes = crate::semantic::parse_classes(&classes);
|
||
let categories = parse_categories(&categories, &classes)?;
|
||
Self::from_bytes(&bytes, categories)
|
||
}
|
||
|
||
pub fn from_bytes(bytes: &[u8], categories: Vec<Category>) -> Result<Self, SegmentError> {
|
||
// f32, as for `SemanticModel`; see there.
|
||
let session = dr_inference_engine::open(
|
||
dr_inference_engine::Role::Scene,
|
||
dr_inference_engine::Form::F32,
|
||
bytes,
|
||
)?;
|
||
|
||
Ok(Self {
|
||
session,
|
||
categories,
|
||
})
|
||
}
|
||
|
||
pub fn categories(&self) -> &[Category] {
|
||
&self.categories
|
||
}
|
||
|
||
/// Weigh every category over one image.
|
||
///
|
||
/// `rgb` is tightly packed `f32` RGB in `0.0..=1.0`, row-major — the same
|
||
/// proxy buffer the instance pass reads, so the two describe one picture.
|
||
///
|
||
/// One inference over the whole frame. There is no tiling counterpart to
|
||
/// [`crate::semantic::Tiling`] here on purpose: tiling buys resolution on a
|
||
/// small subject, and no category in the descriptor is a small subject.
|
||
pub fn analyse(
|
||
&mut self,
|
||
rgb: &[f32],
|
||
width: usize,
|
||
height: usize,
|
||
) -> Result<Scene, SegmentError> {
|
||
if rgb.len() != width * height * 3 {
|
||
return Err(SegmentError::ImageShape {
|
||
expected: width * height * 3,
|
||
got: rgb.len(),
|
||
});
|
||
}
|
||
|
||
let categories = &self.categories;
|
||
let acquired = self.session.acquire()?;
|
||
let mut session = acquired.lock();
|
||
|
||
let window = Window {
|
||
x: 0.0,
|
||
y: 0.0,
|
||
w: width as f32,
|
||
h: height as f32,
|
||
};
|
||
let letterbox = Letterbox::fit(window.w, window.h);
|
||
let input = letterbox.sample(rgb, width, height, &window);
|
||
|
||
let outputs = session
|
||
.run(ort::inputs![
|
||
ort::value::Tensor::from_array(input).map_err(SegmentError::Inference)?
|
||
])
|
||
.map_err(SegmentError::Inference)?;
|
||
|
||
let (shape, logits) = outputs[0]
|
||
.try_extract_tensor::<f32>()
|
||
.map_err(|_| SegmentError::OutputShape("logits"))?;
|
||
|
||
// `[1, 150, gh, gw]`. Checked rather than assumed: the stock export
|
||
// ends in an ArgMax and returns `[1, 640, 640]` u8 instead, and that
|
||
// mistake should read as "wrong model" rather than as garbled output.
|
||
if shape.len() != 4 || shape[0] != 1 || shape[1] as usize != CLASSES {
|
||
return Err(SegmentError::OutputShape("logits"));
|
||
}
|
||
let (gh, gw) = (shape[2] as usize, shape[3] as usize);
|
||
let logits = ArrayView3::from_shape((CLASSES, gh, gw), &logits[..CLASSES * gh * gw])
|
||
.map_err(|_| SegmentError::OutputShape("logits"))?;
|
||
|
||
Ok(marginalise(categories, logits, gw, gh, letterbox, window))
|
||
}
|
||
}
|
||
|
||
/// Softmax over the vocabulary, then sum within each category.
|
||
///
|
||
/// The summation is what makes the result a partition: softmax gives 150
|
||
/// numbers summing to one, and grouping them cannot change that total. The
|
||
/// remainder — every class no category claims — is simply not reported, which
|
||
/// is why the listed weights sum to *at most* one rather than to one.
|
||
///
|
||
/// Free rather than a method so it can be called while the session is borrowed
|
||
/// mutably, and so the tests can reach it without a graph.
|
||
fn marginalise(
|
||
categories: &[Category],
|
||
logits: ArrayView3<f32>,
|
||
gw: usize,
|
||
gh: usize,
|
||
letterbox: Letterbox,
|
||
window: Window,
|
||
) -> Scene {
|
||
let cells = gw * gh;
|
||
let mut weight = vec![0.0f32; categories.len() * cells];
|
||
let mut probability = vec![0.0f32; CLASSES];
|
||
|
||
for cell in 0..cells {
|
||
let (y, x) = (cell / gw, cell % gw);
|
||
|
||
// Shift by the maximum before exponentiating. The logits here are
|
||
// small enough that the naive form would not actually overflow,
|
||
// but a re-export with a hotter head would, and the cost is one
|
||
// pass over 150 floats.
|
||
let mut peak = f32::NEG_INFINITY;
|
||
for c in 0..CLASSES {
|
||
peak = peak.max(logits[[c, y, x]]);
|
||
}
|
||
let mut total = 0.0f32;
|
||
for c in 0..CLASSES {
|
||
let p = (logits[[c, y, x]] - peak).exp();
|
||
probability[c] = p;
|
||
total += p;
|
||
}
|
||
let norm = if total > 0.0 { 1.0 / total } else { 0.0 };
|
||
|
||
for (k, category) in categories.iter().enumerate() {
|
||
let mut sum = 0.0f32;
|
||
for &class in &category.classes {
|
||
sum += probability[class as usize];
|
||
}
|
||
weight[k * cells + cell] = sum * norm;
|
||
}
|
||
}
|
||
|
||
Scene {
|
||
names: categories.iter().map(|c| c.name.clone()).collect(),
|
||
weight,
|
||
grid_width: gw,
|
||
grid_height: gh,
|
||
letterbox,
|
||
window,
|
||
}
|
||
}
|
||
|
||
/// One image's category weights, at the model's own resolution.
|
||
#[derive(Debug, Clone)]
|
||
pub struct Scene {
|
||
names: Vec<Arc<str>>,
|
||
/// `[category][y * grid_width + x]`, each in `0.0..=1.0`, and across
|
||
/// categories summing to at most one at every cell.
|
||
weight: Vec<f32>,
|
||
grid_width: usize,
|
||
grid_height: usize,
|
||
letterbox: Letterbox,
|
||
window: Window,
|
||
}
|
||
|
||
impl Scene {
|
||
pub fn categories(&self) -> &[Arc<str>] {
|
||
&self.names
|
||
}
|
||
|
||
pub fn grid_size(&self) -> (usize, usize) {
|
||
(self.grid_width, self.grid_height)
|
||
}
|
||
|
||
/// How many pixels of the analysed image one logit cell spans.
|
||
///
|
||
/// This is the module header's "resolution, stated plainly" as a number a
|
||
/// caller can act on: [`Self::rasterise`] will happily hand back a
|
||
/// full-resolution buffer, and this says how much of that resolution is
|
||
/// the model's and how much is the bilinear's. At the 1600px proxy the
|
||
/// application segments at, it is 20.
|
||
///
|
||
/// [`crate::refine`] is the one caller, and it needs this rather than a
|
||
/// pixel count of its own because the erosion that separates evidence from
|
||
/// smear is *defined* as a multiple of the model's own resolution. Derived
|
||
/// from the fitted letterbox rather than recomputed from the image size,
|
||
/// so the two cannot drift.
|
||
pub fn cell_pixels(&self) -> f32 {
|
||
GRID_STRIDE as f32 / self.letterbox.scale()
|
||
}
|
||
|
||
/// One category's weights over the logit grid.
|
||
pub fn weight(&self, category: usize) -> Option<&[f32]> {
|
||
let cells = self.grid_width * self.grid_height;
|
||
self.weight.get(category * cells..(category + 1) * cells)
|
||
}
|
||
|
||
pub fn index_of(&self, name: &str) -> Option<usize> {
|
||
self.names.iter().position(|n| &**n == name)
|
||
}
|
||
|
||
/// How much of the frame this category covers, `0.0..=1.0`.
|
||
///
|
||
/// Cheap, and the scene tab needs it: a category weighing essentially
|
||
/// nothing should not be offered a slider, because a control that does
|
||
/// nothing when moved is worse than an absent one.
|
||
pub fn coverage(&self, category: usize) -> f32 {
|
||
match self.weight(category) {
|
||
Some(w) if !w.is_empty() => w.iter().sum::<f32>() / w.len() as f32,
|
||
_ => 0.0,
|
||
}
|
||
}
|
||
|
||
/// Resample one category to source-image resolution.
|
||
///
|
||
/// Bilinear over the logit grid. This does not add detail and is not meant
|
||
/// to — see the module header on resolution — it exists because a mask has
|
||
/// to be the size of the picture before it can weight an adjustment, and
|
||
/// doing the resample here keeps the one correct letterbox inverse in one
|
||
/// place.
|
||
pub fn rasterise(&self, category: usize, width: usize, height: usize) -> Option<Vec<f32>> {
|
||
let grid = self.weight(category)?;
|
||
let mut out = vec![0.0f32; width * height];
|
||
|
||
for y in 0..height {
|
||
for x in 0..width {
|
||
let (gx, gy) = self.letterbox.to_grid(
|
||
x as f32 + 0.5,
|
||
y as f32 + 0.5,
|
||
&self.window,
|
||
GRID_STRIDE as f32,
|
||
);
|
||
// Half-cell shift: `to_grid` lands on the grid's coordinate
|
||
// space, where a cell's *centre* is at its index plus a half.
|
||
let (gx, gy) = (gx - 0.5, gy - 0.5);
|
||
let x0 = gx.floor();
|
||
let y0 = gy.floor();
|
||
let (fx, fy) = (gx - x0, gy - y0);
|
||
let x0 = (x0 as isize).clamp(0, self.grid_width as isize - 1) as usize;
|
||
let y0 = (y0 as isize).clamp(0, self.grid_height as isize - 1) as usize;
|
||
let x1 = (x0 + 1).min(self.grid_width - 1);
|
||
let y1 = (y0 + 1).min(self.grid_height - 1);
|
||
|
||
let at = |gx: usize, gy: usize| grid[gy * self.grid_width + gx];
|
||
let top = at(x0, y0) * (1.0 - fx) + at(x1, y0) * fx;
|
||
let bot = at(x0, y1) * (1.0 - fx) + at(x1, y1) * fx;
|
||
out[y * width + x] = top * (1.0 - fy) + bot * fy;
|
||
}
|
||
}
|
||
|
||
Some(out)
|
||
}
|
||
}
|
||
|
||
impl Scene {
|
||
/// Build a scene from known weights, for tests.
|
||
///
|
||
/// The geometry in [`Scene::rasterise`] is a letterbox inverse, and a
|
||
/// letterbox inverse is exactly the kind of code that produces confident,
|
||
/// plausible, wrong answers. Testing it needs weights whose correct
|
||
/// destination is known in advance, which a real inference can never give.
|
||
#[cfg(test)]
|
||
fn from_weights(names: Vec<Arc<str>>, weight: Vec<f32>, gw: usize, gh: usize) -> Self {
|
||
// A square source, so the letterbox is the identity and any offset
|
||
// this finds is the mapping's own rather than the padding's.
|
||
let window = Window {
|
||
x: 0.0,
|
||
y: 0.0,
|
||
w: INPUT_EDGE as f32,
|
||
h: INPUT_EDGE as f32,
|
||
};
|
||
Self {
|
||
names,
|
||
weight,
|
||
grid_width: gw,
|
||
grid_height: gh,
|
||
letterbox: Letterbox::fit(window.w, window.h),
|
||
window,
|
||
}
|
||
}
|
||
}
|
||
|
||
/// Read `models/scene/categories.txt`, resolving class names to indices.
|
||
///
|
||
/// Hand-written rather than generated, unlike the `.classes.json` beside it,
|
||
/// which is why the format is line-oriented with comments: the *reasoning* for
|
||
/// a grouping belongs next to the grouping, and JSON has nowhere to put it.
|
||
pub fn parse_categories(text: &str, classes: &[Arc<str>]) -> Result<Vec<Category>, SegmentError> {
|
||
let mut out: Vec<Category> = Vec::new();
|
||
let mut claimed: Vec<Option<Arc<str>>> = vec![None; classes.len()];
|
||
|
||
for line in text.lines() {
|
||
let line = line.split('#').next().unwrap_or("").trim();
|
||
if line.is_empty() {
|
||
continue;
|
||
}
|
||
let Some((name, members)) = line.split_once('=') else {
|
||
return Err(SegmentError::CategoryDescriptor(format!(
|
||
"line is not `name = class, class, ...`: {line}"
|
||
)));
|
||
};
|
||
let name: Arc<str> = name.trim().into();
|
||
|
||
let mut indices = Vec::new();
|
||
for member in members.split(',') {
|
||
let member = member.trim();
|
||
if member.is_empty() {
|
||
continue;
|
||
}
|
||
let Some(index) = classes.iter().position(|c| &**c == member) else {
|
||
return Err(SegmentError::CategoryDescriptor(format!(
|
||
"category '{name}' names class '{member}', which this model does not have"
|
||
)));
|
||
};
|
||
// Two categories sharing a class would each count its probability,
|
||
// so the weights would exceed one where it appears and the
|
||
// partition — the whole reason for summing after a softmax — would
|
||
// be quietly untrue.
|
||
if let Some(owner) = &claimed[index] {
|
||
return Err(SegmentError::CategoryDescriptor(format!(
|
||
"class '{member}' is claimed by both '{owner}' and '{name}'"
|
||
)));
|
||
}
|
||
claimed[index] = Some(name.clone());
|
||
indices.push(index as u16);
|
||
}
|
||
|
||
if indices.is_empty() {
|
||
return Err(SegmentError::CategoryDescriptor(format!(
|
||
"category '{name}' lists no classes"
|
||
)));
|
||
}
|
||
out.push(Category {
|
||
name,
|
||
classes: indices,
|
||
});
|
||
}
|
||
|
||
if out.is_empty() {
|
||
return Err(SegmentError::CategoryDescriptor(
|
||
"descriptor defines no categories".into(),
|
||
));
|
||
}
|
||
Ok(out)
|
||
}
|
||
|
||
#[cfg(test)]
|
||
mod tests {
|
||
use super::*;
|
||
|
||
fn vocabulary() -> Vec<Arc<str>> {
|
||
["sky", "tree", "grass", "person", "wall"]
|
||
.iter()
|
||
.map(|s| Arc::from(*s))
|
||
.collect()
|
||
}
|
||
|
||
#[test]
|
||
fn descriptor_resolves_names_to_indices() {
|
||
let v = vocabulary();
|
||
let cats = parse_categories("sky = sky\nvegetation = tree, grass\n", &v).unwrap();
|
||
assert_eq!(cats.len(), 2);
|
||
assert_eq!(&*cats[0].name, "sky");
|
||
assert_eq!(cats[0].classes, vec![0]);
|
||
assert_eq!(cats[1].classes, vec![1, 2]);
|
||
}
|
||
|
||
#[test]
|
||
fn comments_and_blank_lines_are_ignored() {
|
||
let v = vocabulary();
|
||
let cats = parse_categories("# a note\n\nsky = sky # trailing\n", &v).unwrap();
|
||
assert_eq!(cats.len(), 1);
|
||
assert_eq!(cats[0].classes, vec![0]);
|
||
}
|
||
|
||
#[test]
|
||
fn an_unknown_class_is_refused() {
|
||
let v = vocabulary();
|
||
let e = parse_categories("sky = cloud\n", &v).unwrap_err();
|
||
assert!(format!("{e}").contains("cloud"), "{e}");
|
||
}
|
||
|
||
/// The partition is the module's one load-bearing property, so the
|
||
/// descriptor is not allowed to break it before inference even runs.
|
||
#[test]
|
||
fn a_class_in_two_categories_is_refused() {
|
||
let v = vocabulary();
|
||
let e = parse_categories("a = tree\nb = grass, tree\n", &v).unwrap_err();
|
||
assert!(format!("{e}").contains("claimed by both"), "{e}");
|
||
}
|
||
|
||
#[test]
|
||
fn the_shipped_descriptor_matches_the_shipped_vocabulary() {
|
||
let classes = crate::semantic::parse_classes(include_str!(
|
||
"../../../models/scene/yolo26s-sem-ade20k.classes.json"
|
||
));
|
||
assert_eq!(classes.len(), CLASSES);
|
||
let cats = parse_categories(
|
||
include_str!("../../../models/scene/categories.txt"),
|
||
&classes,
|
||
)
|
||
.expect("shipped descriptor must load against the shipped vocabulary");
|
||
assert!(cats.iter().any(|c| &*c.name == "sky"));
|
||
assert!(cats.iter().any(|c| &*c.name == "vegetation"));
|
||
}
|
||
|
||
/// Weight put in the top half of the grid must come back in the top half
|
||
/// of the image.
|
||
///
|
||
/// The one property that makes a category mask worth anything: if the
|
||
/// model says "sky up here" and `rasterise` puts it down there, every
|
||
/// grade lands on the wrong half of the photograph and nothing about the
|
||
/// numbers looks wrong. A vertical split is the cheapest arrangement that
|
||
/// catches a flipped axis, and a flipped axis is the mistake this code is
|
||
/// actually prone to.
|
||
#[test]
|
||
fn rasterise_keeps_weight_on_the_side_it_came_from() {
|
||
let (gw, gh) = (8usize, 8usize);
|
||
let mut weight = vec![0.0f32; gw * gh];
|
||
for y in 0..gh / 2 {
|
||
for x in 0..gw {
|
||
weight[y * gw + x] = 1.0;
|
||
}
|
||
}
|
||
let scene = Scene::from_weights(vec![Arc::from("sky")], weight, gw, gh);
|
||
|
||
let (w, h) = (64usize, 64usize);
|
||
let mask = scene.rasterise(0, w, h).expect("category 0 exists");
|
||
|
||
let mean = |y0: usize, y1: usize| {
|
||
let band: f32 = (y0..y1)
|
||
.flat_map(|y| (0..w).map(move |x| (y, x)))
|
||
.map(|(y, x)| mask[y * w + x])
|
||
.sum();
|
||
band / ((y1 - y0) * w) as f32
|
||
};
|
||
let top = mean(0, h / 4);
|
||
let bottom = mean(3 * h / 4, h);
|
||
assert!(
|
||
top > 0.9,
|
||
"the half that had the weight should keep it: {top}"
|
||
);
|
||
assert!(
|
||
bottom < 0.1,
|
||
"the half that had none should stay empty: {bottom}"
|
||
);
|
||
}
|
||
|
||
/// A horizontal split too, because a transpose passes the vertical test.
|
||
///
|
||
/// Swapping x and y maps a top band onto a left band, and the check above
|
||
/// would still see the top band full. Two axes is what makes the pair
|
||
/// meaningful; either alone is not.
|
||
#[test]
|
||
fn rasterise_does_not_transpose() {
|
||
let (gw, gh) = (8usize, 8usize);
|
||
let mut weight = vec![0.0f32; gw * gh];
|
||
for y in 0..gh {
|
||
for x in 0..gw / 2 {
|
||
weight[y * gw + x] = 1.0;
|
||
}
|
||
}
|
||
let scene = Scene::from_weights(vec![Arc::from("left")], weight, gw, gh);
|
||
|
||
let (w, h) = (64usize, 64usize);
|
||
let mask = scene.rasterise(0, w, h).expect("category 0 exists");
|
||
let mean = |x0: usize, x1: usize| {
|
||
let band: f32 = (0..h)
|
||
.flat_map(|y| (x0..x1).map(move |x| (y, x)))
|
||
.map(|(y, x)| mask[y * w + x])
|
||
.sum();
|
||
band / (h * (x1 - x0)) as f32
|
||
};
|
||
assert!(mean(0, w / 4) > 0.9, "left stays left");
|
||
assert!(mean(3 * w / 4, w) < 0.1, "right stays empty");
|
||
}
|
||
|
||
/// Coverage is the fraction of the frame, not a count or a sum.
|
||
///
|
||
/// The scene tab hides a category below half a percent, so a coverage that
|
||
/// is off by a factor of the grid size would either hide everything or
|
||
/// hide nothing, and both look like the model failing rather than the
|
||
/// arithmetic.
|
||
#[test]
|
||
fn coverage_is_a_fraction_of_the_frame() {
|
||
let (gw, gh) = (10usize, 10usize);
|
||
let mut weight = vec![0.0f32; gw * gh];
|
||
for cell in weight.iter_mut().take(25) {
|
||
*cell = 1.0;
|
||
}
|
||
let scene = Scene::from_weights(vec![Arc::from("quarter")], weight, gw, gh);
|
||
let c = scene.coverage(0);
|
||
assert!(
|
||
(c - 0.25).abs() < 1e-6,
|
||
"a quarter of the cells is 0.25, got {c}"
|
||
);
|
||
}
|
||
|
||
/// Softmax then group: the reported weights must never exceed one, and
|
||
/// must equal one exactly when the categories name every class.
|
||
#[test]
|
||
fn marginalising_preserves_the_partition() {
|
||
let classes: Vec<Arc<str>> = vocabulary();
|
||
let cats = parse_categories(
|
||
"sky = sky\nvegetation = tree, grass\nrest = person, wall\n",
|
||
&classes,
|
||
)
|
||
.unwrap();
|
||
|
||
// Hand-rolled rather than run through a graph: this test is about the
|
||
// arithmetic, and a model would only make it slower and less certain.
|
||
let (gw, gh) = (2usize, 2usize);
|
||
let mut logits = vec![0.0f32; classes.len() * gw * gh];
|
||
for (i, v) in logits.iter_mut().enumerate() {
|
||
*v = (i % 7) as f32 * 0.3;
|
||
}
|
||
let view = ArrayView3::from_shape((classes.len(), gh, gw), &logits).unwrap();
|
||
|
||
// `marginalise` is a method for access to `self.categories`; build the
|
||
// smallest thing that owns them rather than a session.
|
||
let cells = gw * gh;
|
||
let mut weight = vec![0.0f32; cats.len() * cells];
|
||
for cell in 0..cells {
|
||
let (y, x) = (cell / gw, cell % gw);
|
||
let peak = (0..classes.len()).fold(f32::NEG_INFINITY, |m, c| m.max(view[[c, y, x]]));
|
||
let p: Vec<f32> = (0..classes.len())
|
||
.map(|c| (view[[c, y, x]] - peak).exp())
|
||
.collect();
|
||
let total: f32 = p.iter().sum();
|
||
for (k, category) in cats.iter().enumerate() {
|
||
let s: f32 = category.classes.iter().map(|&c| p[c as usize]).sum();
|
||
weight[k * cells + cell] = s / total;
|
||
}
|
||
}
|
||
|
||
for cell in 0..cells {
|
||
let sum: f32 = (0..cats.len()).map(|k| weight[k * cells + cell]).sum();
|
||
assert!(
|
||
(sum - 1.0).abs() < 1e-5,
|
||
"categories covering every class must sum to 1, got {sum}"
|
||
);
|
||
}
|
||
}
|
||
}
|