The weights landed last commit with nothing to read them. This is the decoder, and the shape of it follows from one property worth stating before the code: the categories must partition the image. ## Why a partition, and not a mask per category The scene tab applies one grade to every pixel of a category — lift the sky, desaturate foliage — and both grades meet at the horizon. If each category carried an independent mask, feathering them outward would make the boundary band belong to both, so both grades would land there and every horizon would acquire a visible seam. Feathering has to *blend* there, not accumulate. So `marginalise` takes one softmax over all 150 channels and sums within each category. Grouping cannot change a total of one, so the listed categories plus the unlisted remainder sum to one at every pixel, by construction rather than by normalising afterwards. `parse_categories` refuses a descriptor that claims a class twice, because that is the one input that would quietly make the property untrue. ## The descriptor is data, and hand-written `models/scene/categories.txt` groups ADE20K's 150 classes into the eight a photographer would recognise. It is a file rather than a table in Rust for the reason `models/LICENCE.md` predicted — a vocabulary is model metadata — and it is line-oriented with comments rather than JSON like the `.classes.json` beside it, because that file is generated and this one is argued. Why `swimming pool` is water and not architecture belongs next to the line that says so. Classes are named, not indexed. An index is silently wrong after a re-export; a name is loudly wrong, and the loader refuses one the model does not have. ## Resolution, kept visible `Scene` holds the native 80×80 logit grid and resamples on demand rather than upsampling once at load. The coarseness is real — it is what the graph produces — and a type that hides it behind an early resize invites callers to expect detail that was never there. `rasterise` is where the letterbox inverse lives, once. `Letterbox` and `Window` become `pub(crate)` and `to_proto` generalises to `to_grid`, because both dense outputs this crate reads are an even fraction of the same letterboxed square and differ only in the divisor. ## Verified by looking, which is the only way this gets verified `examples/scene.rs` writes the photograph dimmed outside each category. A transposed axis or an off-by-one in the inverse produces perfectly plausible weights over slightly the wrong pixels, and no unit test catches that. On an indoor frame the person mask lands on the person, including the outstretched arm, and sky reads ~5% against a bright ceiling. It doubles as the benchmark, because every timing quoted while this model was chosen came off a laptop compiling other things and none of them belong in a document. Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
526 lines
20 KiB
Rust
526 lines
20 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;
|
||
|
||
use crate::semantic::{install_backend, 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: ort::session::Session,
|
||
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> {
|
||
install_backend();
|
||
|
||
let session = ort::session::Session::builder()
|
||
.map_err(SegmentError::Inference)?
|
||
.commit_from_memory(bytes)
|
||
.map_err(SegmentError::Inference)?;
|
||
|
||
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(),
|
||
});
|
||
}
|
||
|
||
// Split the borrow: `run` needs the session mutably while
|
||
// `marginalise` needs the categories, and going through `self` for
|
||
// both at once is what the borrow checker objects to.
|
||
let Self {
|
||
session,
|
||
categories,
|
||
} = self;
|
||
|
||
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)
|
||
}
|
||
|
||
/// 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)
|
||
}
|
||
}
|
||
|
||
/// 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"));
|
||
}
|
||
|
||
/// 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}"
|
||
);
|
||
}
|
||
}
|
||
}
|