diff --git a/AGENTS.md b/AGENTS.md
index 52ccaae..b8c3cab 100644
--- a/AGENTS.md
+++ b/AGENTS.md
@@ -9,7 +9,7 @@ For a new family, scale, or task, follow [docs/MODEL_BRINGUP.md](docs/MODEL_BRIN
## Repository map
-- `src/models/yolox/`: stable YOLOX nano/tiny/s/m/l/x implementation and official `.pth` import.
+- `src/models/yolox/`: stable YOLOX nano/tiny/s/m/l/x implementation and tensor-state import.
- `src/models/yolov3_tiny/`, `yolov8/`, `yolov10/`, `yolo11/`, `yolo12/`, `yolo26/`:
experimental Ultralytics-family graphs and native Burnpack loaders.
- YOLOv8, YOLO11, and YOLO26 also provide `-seg` variants; YOLOv8, YOLO11, and YOLO26 provide
@@ -32,30 +32,37 @@ user-facing text.
Run Python tools from the repository root with the tools project selected:
```console
-uv run --project tools tools/export_ultralytics_state.py yolo26n.pt target/yolo26n-state.pt
+uv run --project tools tools/export_checkpoint_state.py yolo26n.pt target/yolo26n-state.pt
```
Stable YOLOX and Ultralytics-family models both run from native Burnpacks:
```console
-montgomery pack-weights --model yolox-nano --input target/yolox_nano.pth
-montgomery predict --model yolox-nano --source docs/dog_bike_man.jpg
+uv run --project tools tools/export_checkpoint_state.py target/yolox_nano.pth target/yolox-nano-state.pt
+montgomery pack-weights --architecture yolox-nano --state target/yolox-nano-state.pt
+montgomery predict --model yolox-nano.bpk --source docs/dog_bike_man.jpg
-montgomery pack-weights --model yolo26n --input target/yolo26n-state.pt
-montgomery predict --model yolo26n --source docs/dog_bike_man.jpg
+montgomery pack-weights --architecture yolo26n --state target/yolo26n-state.pt
+montgomery predict --model yolo26n.bpk --source docs/dog_bike_man.jpg
```
Python/PyTorch is conversion- and development-time only; normal inference is Rust/Burn.
## Verification
-Run before handing off changes:
+CI installs the current stable Rust toolchain on every run. Before handing off changes, update the
+local stable toolchain (`rustup update stable`), confirm `rustc --version` matches current stable,
+and run the exact CI sequence below. Do not substitute `cargo check` for Clippy, filter the training
+tests, or omit `cargo build`:
```console
cargo fmt --check
+cargo build
cargo test
cargo clippy --all-targets -- -D warnings
-cargo check --no-default-features --lib
+cargo test --features training
+cargo clippy --features training --all-targets -- -D warnings
+cargo clippy --no-default-features --lib -- -D warnings
```
When external checkpoints and fixtures are available:
@@ -64,12 +71,8 @@ When external checkpoints and fixtures are available:
cargo test -- --ignored
```
-Training is opt-in. For training changes, run:
-
-```console
-cargo test --features training training
-cargo clippy --features training --all-targets -- -D warnings
-```
+Training is opt-in at runtime, but its full test and Clippy commands above are mandatory before
+handing off any change because Linux CI always runs them.
Real training, hardware smoke tests, and latency measurements must use `--release`. Single-image
latency tests must use `--test-threads 1` to avoid CPU contention. When touching a runtime/backend
diff --git a/README.md b/README.md
index 6930c96..fb6eacb 100644
--- a/README.md
+++ b/README.md
@@ -4,7 +4,7 @@
-Native object detection, instance segmentation, and image classification in Rust with [Burn](https://burn.dev).
+Native object detection, instance segmentation, and image classification in Rust with [Burn](https://burn.dev)
@@ -38,22 +38,24 @@ install [uv](https://docs.astral.sh/uv/) and run:
```console
uv run --project tools python -c "from ultralytics import YOLO; YOLO('yolo26n.pt')"
-uv run --project tools tools/export_ultralytics_state.py yolo26n.pt target/yolo26n-state.pt
-montgomery pack-weights --model yolo26n --input target/yolo26n-state.pt
+uv run --project tools tools/export_checkpoint_state.py yolo26n.pt target/yolo26n-state.pt
+montgomery pack-weights --architecture yolo26n --state target/yolo26n-state.pt
```
This creates `yolo26n.bpk`. Architecture, dataset, upstream version, format version, and licensing
-stay inside the artifact metadata. Run inference with the same short model name:
+stay inside the artifact metadata. Run inference by naming that artifact explicitly:
```console
-montgomery predict --model yolo26n --source image.jpg
+montgomery predict --model yolo26n.bpk --source image.jpg
```
Useful options: `--json`, `--confidence 0.30`, `--output result.png`, `--masks`, and `--device gpu`.
-Pass `--weights another-name.bpk` only when you deliberately use a nonstandard filename.
+There is no implicit filename and no separate architecture flag: `predict --model` only accepts a
+self-describing Montgomery `.bpk` Burnpack.
-YOLOX accepts its official `.pth` directly through `pack-weights`. Other families use the
-tensor-only conversion shown above. Prediction always uses a Montgomery `.bpk` Burnpack.
+The conversion workflow is identical for every family, including YOLOX: pass the trusted upstream
+`.pt` or `.pth` checkpoint through `export_checkpoint_state.py`, then pass its tensor-only output
+through `pack-weights`. Normal prediction and training never accept upstream checkpoint formats.
## Supported models
@@ -88,9 +90,20 @@ fn main() -> montgomery::Result<()> {
## Train
```console
-montgomery train --model yolo26n --data dataset.yaml --epochs 100
+# Fresh initialization
+montgomery train --architecture yolo26n --data dataset.yaml --epochs 100
+
+# Pretrained initialization
+montgomery train --model yolo26n.bpk --data dataset.yaml --epochs 100
+
+# Exact continuation (model and dataset come from the training checkpoint)
+montgomery train --resume runs/train/checkpoints/last
```
+Exactly one initialization mode is required: `--architecture` means scratch, `--model` requires a
+pretrained `.bpk`, and `--resume` requires a full native training checkpoint. A Burnpack initializes
+a new run; it is not a resumable optimizer checkpoint.
+
Every run contains:
- `results.csv`, `results.svg`, and `validation.jsonl`
@@ -109,11 +122,10 @@ and longer detection/segmentation runs still need work. Full methodology and lim
## Export ONNX
```console
-montgomery export-onnx --model yolo26n
+montgomery export-onnx --model yolo26n.bpk
```
-This reads `yolo26n.bpk` and writes `yolo26n.onnx`; use `--weights` or `--output` only to override
-those names.
+This reads the explicit Burnpack and writes `yolo26n.onnx`; use `--output` to select another path.
The offline exporter validates the graph with ONNX Runtime. Setup details are in
[tools/onnx/README.md](tools/onnx/README.md).
diff --git a/docs/MODEL_BRINGUP.md b/docs/MODEL_BRINGUP.md
index af4cf8e..185385a 100644
--- a/docs/MODEL_BRINGUP.md
+++ b/docs/MODEL_BRINGUP.md
@@ -17,9 +17,9 @@ Orientation first:
`src/models/yolo26/` (DFL-free end-to-end), or `src/models/yolov3_tiny/` (classic NMS path).
- Inference modes in the wild: one2one end-to-end heads are top-k selected + confidence filtered
(`end2end_topk_detections` in `src/lib.rs`); plain heads go through class-aware NMS.
-- Non-Ultralytics families need the same ground-truth discipline with different tooling. YOLOX
- imports its official `.pth` through `pack-weights` and consumes the resulting native `.bpk` at
- runtime. Its
+- Non-Ultralytics families need the same ground-truth discipline. Every family converts its
+ trusted upstream checkpoint to a tensor-only state with `tools/export_checkpoint_state.py`, then
+ packs that state into the native `.bpk` consumed at runtime. YOLOX's
golden fixtures come from the *official YOLOX repository sources* instead of the Ultralytics
package: `tools/export_yolox_fixtures.py` assembles a small import package from a plain YOLOX
checkout under `target/yolox-ref/`, loads the checkpoint with `strict=True`, and dumps per-stage
@@ -104,11 +104,11 @@ existing templates. Hard rules:
- the model preprocessing profile and appropriate shared runtime dispatch trait;
- `pack_weights` arm;
- the model-name, input-size, and packer-extension tests.
-- `src/main.rs`: `--model` help text lists the new name.
+- `src/main.rs`: keep `--architecture` for catalog identifiers and `--model` for `.bpk` paths.
## 5. Conversion tooling
-- `tools/export_ultralytics_state.py` is model-agnostic; run it against the official `.pt`.
+- `tools/export_checkpoint_state.py` is model-agnostic; run it against the official `.pt` or `.pth`.
- Copy `tools/export_yolo26_fixtures.py` to `tools/export__fixtures.py` and adjust: the hooked
body layer indices, the head index, the `preds["one2one"]` (or plain-tensor) access, and all
`` filenames. It writes the source/preprocessed reference PNGs and the golden JSON under
@@ -126,8 +126,8 @@ cargo check --no-default-features --lib
Then the parity loop (checkpoint and fixtures live under `target/`):
```console
-uv run --project tools tools\export_ultralytics_state.py target/.pt target/-state.pt
-montgomery pack-weights --model --input target/-state.pt
+uv run --project tools tools\export_checkpoint_state.py target/.pt target/-state.pt
+montgomery pack-weights --architecture --state target/-state.pt
uv run --project tools tools\export__fixtures.py target/.pt docs/dog_bike_man.jpg target
cargo test -- --ignored
```
@@ -180,7 +180,7 @@ YOLO26-cls bring-up (n/s/m/l/x) is the classification template.
one place — `Yolo11SegHead` wraps `Yolo11Head`) and add the task head modules with checkpoint
key names as field names. The mask tensors join the model output (`SegmentOutput`); the decode,
NMS, and mask assembly stay in the runtime.
-5. **Weights path** is unchanged: `tools/export_ultralytics_state.py` is model-agnostic (it dumps
+5. **Weights path** is unchanged: `tools/export_checkpoint_state.py` is model-agnostic (it dumps
whatever `state_dict()` holds), so only the Rust key remaps need the new head rules (one per
path-segment pattern), plus `ModelId` arms, packer arms, and verified artifact bytes/SHA-256.
6. **Fixtures and parity**: extend the family's fixture exporter for the task (the seg fixture adds
@@ -190,7 +190,7 @@ YOLO26-cls bring-up (n/s/m/l/x) is the classification template.
ignored Rust test can compare per-detection mask IoU (target >= 0.95).
7. **Public API**: a new result type (`SegmentationDetection` with `InstanceMask`), a new predictor
method that does not disturb `predict()`, letterbox geometry shared with the boxes, and CLI
- wiring (`--model -seg`, `--masks`) that leaves detect-model behavior untouched.
+ wiring (`--model -seg.bpk`, `--masks`) that leaves detect-model behavior untouched.
Classification (YOLO26-cls template) differs in these ways: the input is 224 px with Ultralytics'
classify transform (anti-aliased shortest-edge resize + centered crop, no letterbox), the class
diff --git a/src/data/letterbox.rs b/src/data/letterbox.rs
index 7e9c569..9f347b4 100644
--- a/src/data/letterbox.rs
+++ b/src/data/letterbox.rs
@@ -2,38 +2,80 @@ use fast_image_resize as fir;
use image::{DynamicImage, ImageBuffer, Rgb, RgbImage, imageops};
fn resize_opencv_linear(source: &DynamicImage, width: u32, height: u32) -> RgbImage {
- let source = source.to_rgb8();
- let source_width = source.width();
- let source_height = source.height();
+ // Borrow the RGB8 buffer when the source already is RGB8 to avoid a full-frame clone;
+ // otherwise convert once (same pixels `to_rgb8()` produced).
+ let owned;
+ let (source_raw, source_width, source_height) = match source.as_rgb8() {
+ Some(buffer) => (buffer.as_raw().as_slice(), buffer.width(), buffer.height()),
+ None => {
+ owned = source.to_rgb8();
+ (owned.as_raw().as_slice(), owned.width(), owned.height())
+ }
+ };
let scale_x = source_width as f64 / width as f64;
let scale_y = source_height as f64 / height as f64;
-
- ImageBuffer::from_fn(width, height, |x, y| {
+ let max_x = source_width as f64 - 1.0;
+ let max_y = source_height as f64 - 1.0;
+
+ // Hoist the per-column and per-row sampling math out of the inner loop. The formulas are
+ // unchanged from the per-pixel version; only the evaluation points move.
+ let width = width as usize;
+ let height = height as usize;
+ let source_width = source_width as usize;
+ let mut x0_table = Vec::with_capacity(width);
+ let mut x1_table = Vec::with_capacity(width);
+ let mut wx_table = Vec::with_capacity(width);
+ for x in 0..width {
let source_x = (x as f64 + 0.5) * scale_x - 0.5;
- let source_y = (y as f64 + 0.5) * scale_y - 0.5;
- let x0 = source_x.floor().clamp(0.0, source_width as f64 - 1.0) as u32;
- let y0 = source_y.floor().clamp(0.0, source_height as f64 - 1.0) as u32;
+ let x0 = source_x.floor().clamp(0.0, max_x) as usize;
let x1 = (x0 + 1).min(source_width - 1);
- let y1 = (y0 + 1).min(source_height - 1);
- let weight_x = source_x.clamp(0.0, source_width as f64 - 1.0) - x0 as f64;
- let weight_y = source_y.clamp(0.0, source_height as f64 - 1.0) - y0 as f64;
- let top_left = source.get_pixel(x0, y0).0;
- let top_right = source.get_pixel(x1, y0).0;
- let bottom_left = source.get_pixel(x0, y1).0;
- let bottom_right = source.get_pixel(x1, y1).0;
- let mut output = [0_u8; 3];
-
- for channel in 0..3 {
- let top =
- top_left[channel] as f64 * (1.0 - weight_x) + top_right[channel] as f64 * weight_x;
- let bottom = bottom_left[channel] as f64 * (1.0 - weight_x)
- + bottom_right[channel] as f64 * weight_x;
- output[channel] = (top * (1.0 - weight_y) + bottom * weight_y)
- .round()
- .clamp(0.0, 255.0) as u8;
+ x0_table.push(x0);
+ x1_table.push(x1);
+ wx_table.push(source_x.clamp(0.0, max_x) - x0 as f64);
+ }
+ let source_height_usize = source_height as usize;
+ let mut y0_table = Vec::with_capacity(height);
+ let mut y1_table = Vec::with_capacity(height);
+ let mut wy_table = Vec::with_capacity(height);
+ for y in 0..height {
+ let source_y = (y as f64 + 0.5) * scale_y - 0.5;
+ let y0 = source_y.floor().clamp(0.0, max_y) as usize;
+ let y1 = (y0 + 1).min(source_height_usize - 1);
+ y0_table.push(y0);
+ y1_table.push(y1);
+ wy_table.push(source_y.clamp(0.0, max_y) - y0 as f64);
+ }
+
+ let mut output = vec![0u8; width * height * 3];
+ for y in 0..height {
+ let y0 = y0_table[y] * source_width;
+ let y1 = y1_table[y] * source_width;
+ let weight_y = wy_table[y];
+ let inv_weight_y = 1.0 - weight_y;
+ let row = y * width * 3;
+ for x in 0..width {
+ let x0 = x0_table[x];
+ let x1 = x1_table[x];
+ let weight_x = wx_table[x];
+ let inv_weight_x = 1.0 - weight_x;
+ let tl = y0 + x0;
+ let tr = y0 + x1;
+ let bl = y1 + x0;
+ let br = y1 + x1;
+ let out = row + x * 3;
+ for channel in 0..3 {
+ let top = source_raw[tl * 3 + channel] as f64 * inv_weight_x
+ + source_raw[tr * 3 + channel] as f64 * weight_x;
+ let bottom = source_raw[bl * 3 + channel] as f64 * inv_weight_x
+ + source_raw[br * 3 + channel] as f64 * weight_x;
+ output[out + channel] = (top * inv_weight_y + bottom * weight_y)
+ .round()
+ .clamp(0.0, 255.0) as u8;
+ }
}
- Rgb(output)
- })
+ }
+ ImageBuffer::from_raw(width as u32, height as u32, output)
+ .expect("pre-sized resize buffer matches dimensions")
}
/// Resize an RGB8 image with the `fast_image_resize` crate (runtime-dispatched SIMD kernels).
diff --git a/src/lib.rs b/src/lib.rs
index b70d687..1846e52 100644
--- a/src/lib.rs
+++ b/src/lib.rs
@@ -293,6 +293,16 @@ impl ModelId {
format!("{}.bpk", self.as_str())
}
+ /// Read the architecture identifier embedded in a Montgomery Burnpack artifact.
+ #[cfg(feature = "pretrained")]
+ pub fn from_burnpack(path: impl AsRef) -> Result {
+ let path = path.as_ref();
+ require_burnpack_path(path)?;
+ artifact_metadata(path)?
+ .model_id
+ .ok_or_else(|| "Burnpack artifact is missing montgomery.model metadata".into())
+ }
+
/// Default square input side for this catalog model.
pub const fn default_input_size(self) -> usize {
match self {
@@ -919,7 +929,7 @@ fn require_burnpack_path(path: &Path) -> Result<()> {
if path.extension().and_then(|value| value.to_str()) != Some("bpk") {
return Err(
"Model requires a native .bpk artifact; convert upstream checkpoints with \
- pack_weights or the pack-weights CLI"
+ the pack-weights CLI"
.into(),
);
}
@@ -1295,30 +1305,45 @@ pub(crate) fn run_end_to_end_segmentations(
// Two-stage top-k exactly like `end2end_topk_detections`, with anchor indices kept so the
// mask coefficients of every survivor can be gathered.
let keep = max_detections.min(anchors);
- let best_scores = (0..anchors)
- .map(|anchor| {
- let row = &scores[anchor * num_classes..(anchor + 1) * num_classes];
- row.iter().copied().fold(f32::NEG_INFINITY, f32::max)
- })
- .collect::>();
- let mut anchor_order = (0..anchors).collect::>();
+ let mut best_scores = Vec::with_capacity(anchors);
+ for anchor in 0..anchors {
+ let row = &scores[anchor * num_classes..(anchor + 1) * num_classes];
+ let mut best = f32::NEG_INFINITY;
+ for &score in row {
+ if score > best {
+ best = score;
+ }
+ }
+ best_scores.push(best);
+ }
+ let mut anchor_order: Vec = (0..anchors)
+ .filter(|&anchor| best_scores[anchor] >= confidence_threshold)
+ .collect();
+ if anchor_order.len() > keep {
+ anchor_order
+ .select_nth_unstable_by(keep, |&a, &b| best_scores[b].total_cmp(&best_scores[a]));
+ anchor_order.truncate(keep);
+ }
anchor_order.sort_unstable_by(|&a, &b| best_scores[b].total_cmp(&best_scores[a]));
- anchor_order.truncate(keep);
- let mut candidates = Vec::with_capacity(keep * num_classes);
+ let mut candidates = Vec::with_capacity(anchor_order.len() * num_classes.min(8));
for (selected_index, &anchor) in anchor_order.iter().enumerate() {
+ let base = anchor * num_classes;
for class in 0..num_classes {
- candidates.push((scores[anchor * num_classes + class], selected_index, class));
+ let score = scores[base + class];
+ if score >= confidence_threshold {
+ candidates.push((score, selected_index, class));
+ }
}
}
+ if candidates.len() > keep {
+ candidates.select_nth_unstable_by(keep, |a, b| b.0.total_cmp(&a.0));
+ candidates.truncate(keep);
+ }
candidates.sort_unstable_by(|a, b| b.0.total_cmp(&a.0));
- candidates.truncate(keep);
- let mut survivors = Vec::new();
+ let mut survivors = Vec::with_capacity(candidates.len());
for (score, selected_index, class) in candidates {
- if score < confidence_threshold {
- continue;
- }
let anchor = anchor_order[selected_index];
let bbox = &boxes[anchor * 4..anchor * 4 + 4];
survivors.push(SegmentationCandidate {
@@ -1413,6 +1438,11 @@ pub(crate) fn run_classic_segmentations(
// Best class per anchor, thresholded (the `nms` helper's filter semantics).
let mut by_class: Vec> =
(0..num_classes).map(|_| Vec::new()).collect();
+ // Survivors are sparse after thresholding; avoid regrowth without over-allocating.
+ let per_class_cap = (anchors / num_classes.max(1)).clamp(4, 256);
+ for vec in by_class.iter_mut() {
+ vec.reserve(per_class_cap);
+ }
for anchor in 0..anchors {
let row = &scores[anchor * num_classes..(anchor + 1) * num_classes];
let (best_class, best_score) = row
@@ -1441,14 +1471,45 @@ pub(crate) fn run_classic_segmentations(
// with every previously kept box of the same class is at most the threshold).
let mut candidates = Vec::new();
for (class_id, mut class_candidates) in by_class.into_iter().enumerate() {
- class_candidates.sort_by(|a, b| b.0.confidence.partial_cmp(&a.0.confidence).unwrap());
- let mut kept: Vec<(BoundingBox, usize)> = Vec::new();
- for (bbox, anchor) in class_candidates {
- if kept
- .iter()
- .all(|(kept_box, _)| crate::postprocess::iou(kept_box, &bbox) <= iou_threshold)
- {
+ if class_candidates.is_empty() {
+ continue;
+ }
+ class_candidates.sort_unstable_by(|a, b| b.0.confidence.total_cmp(&a.0.confidence));
+ let len = class_candidates.len();
+ let mut areas = Vec::with_capacity(len);
+ for (bbox, _) in class_candidates.iter() {
+ areas.push((bbox.xmax - bbox.xmin).max(0.) * (bbox.ymax - bbox.ymin).max(0.));
+ }
+ let mut kept: Vec<(BoundingBox, usize)> = Vec::with_capacity(len);
+ let mut kept_areas: Vec = Vec::with_capacity(len);
+ for (index, (bbox, anchor)) in class_candidates.into_iter().enumerate() {
+ let area = areas[index];
+ let mut drop = false;
+ for (kept_idx, (kept_box, _)) in kept.iter().enumerate() {
+ if kept_box.xmax <= bbox.xmin
+ || kept_box.xmin >= bbox.xmax
+ || kept_box.ymax <= bbox.ymin
+ || kept_box.ymin >= bbox.ymax
+ {
+ continue;
+ }
+ let i_xmin = kept_box.xmin.max(bbox.xmin);
+ let i_xmax = kept_box.xmax.min(bbox.xmax);
+ let i_ymin = kept_box.ymin.max(bbox.ymin);
+ let i_ymax = kept_box.ymax.min(bbox.ymax);
+ let i_area = (i_xmax - i_xmin).max(0.) * (i_ymax - i_ymin).max(0.);
+ if i_area == 0.0 {
+ continue;
+ }
+ let union = kept_areas[kept_idx] + area - i_area;
+ if union > 0.0 && i_area / union > iou_threshold {
+ drop = true;
+ break;
+ }
+ }
+ if !drop {
kept.push((bbox, anchor));
+ kept_areas.push(area);
}
}
candidates.extend(
@@ -1510,34 +1571,53 @@ pub(crate) fn canvas_instance_mask(
}
// Bilinear upsample to the canvas, threshold at > 0, and crop to the box, per canvas pixel.
+ // Integerize the `crop_mask` window once (`x >= box[0] && x < box[2]` <=> `x in
+ // [ceil(box[0]), ceil(box[2]))`) so pixels outside the box are never visited instead of
+ // being branched away per pixel.
let scale_x = proto_width as f64 / canvas_width.max(1) as f64;
let scale_y = proto_height as f64 / canvas_height.max(1) as f64;
let mut mask = vec![false; canvas_width * canvas_height];
- for y in 0..canvas_height {
- // crop_mask keeps rows `y1 <= y < y2` (box edges in canvas pixels).
- if (y as f32) < box_canvas[1] || (y as f32) >= box_canvas[3] {
- continue;
- }
+ let x_start = (box_canvas[0].ceil().clamp(0.0, canvas_width as f32) as usize).min(canvas_width);
+ let x_end = (box_canvas[2].ceil().clamp(0.0, canvas_width as f32) as usize).min(canvas_width);
+ let y_start =
+ (box_canvas[1].ceil().clamp(0.0, canvas_height as f32) as usize).min(canvas_height);
+ let y_end = (box_canvas[3].ceil().clamp(0.0, canvas_height as f32) as usize).min(canvas_height);
+ if x_start >= x_end || y_start >= y_end {
+ return mask;
+ }
+ // Precompute the horizontal sampling tables once per mask; the same `source_x / x0 / x1 /
+ // lambda_x` was recomputed for every row before.
+ let width = x_end - x_start;
+ let mut x0_table = Vec::with_capacity(width);
+ let mut x1_table = Vec::with_capacity(width);
+ let mut lambda_x_table = Vec::with_capacity(width);
+ for x in x_start..x_end {
+ let source_x = ((x as f64 + 0.5) * scale_x - 0.5).max(0.0);
+ let x0 = source_x.floor();
+ let x1 = (x0 + 1.0).min((proto_width - 1) as f64);
+ x0_table.push(x0 as usize);
+ x1_table.push(x1 as usize);
+ lambda_x_table.push(source_x - x0);
+ }
+ for y in y_start..y_end {
let source_y = ((y as f64 + 0.5) * scale_y - 0.5).max(0.0);
let y0 = source_y.floor();
let y1 = (y0 + 1.0).min((proto_height - 1) as f64);
let y0 = y0 as usize;
+ let y1 = y1 as usize;
let lambda_y = source_y - y0 as f64;
- for x in 0..canvas_width {
- // crop_mask keeps columns `x1 <= x < x2`.
- if (x as f32) < box_canvas[0] || (x as f32) >= box_canvas[2] {
- continue;
- }
- let source_x = ((x as f64 + 0.5) * scale_x - 0.5).max(0.0);
- let x0 = source_x.floor();
- let x1 = (x0 + 1.0).min((proto_width - 1) as f64);
- let x0 = x0 as usize;
- let lambda_x = source_x - x0 as f64;
- let top = logits[y0 * proto_width + x0] * (1.0 - lambda_x)
- + logits[y0 * proto_width + x1 as usize] * lambda_x;
- let bottom = logits[y1 as usize * proto_width + x0] * (1.0 - lambda_x)
- + logits[y1 as usize * proto_width + x1 as usize] * lambda_x;
- mask[y * canvas_width + x] = top * (1.0 - lambda_y) + bottom * lambda_y > 0.0;
+ let inv_lambda_y = 1.0 - lambda_y;
+ let row = y * canvas_width;
+ let row0 = y0 * proto_width;
+ let row1 = y1 * proto_width;
+ for (i, x) in (x_start..x_end).enumerate() {
+ let x0 = x0_table[i];
+ let x1 = x1_table[i];
+ let lambda_x = lambda_x_table[i];
+ let inv_lambda_x = 1.0 - lambda_x;
+ let top = logits[row0 + x0] * inv_lambda_x + logits[row0 + x1] * lambda_x;
+ let bottom = logits[row1 + x0] * inv_lambda_x + logits[row1 + x1] * lambda_x;
+ mask[row + x] = top * inv_lambda_y + bottom * lambda_y > 0.0;
}
}
mask
@@ -2289,11 +2369,10 @@ impl Predictor {
}
}
-/// Convert imported upstream tensor state into Montgomery's versioned native Burnpack format.
-///
-/// YOLOX accepts its official `.pth` checkpoint; Ultralytics-family inputs are the tensor-only
-/// states generated by `tools/export_ultralytics_state.py`. The output is `.bpk` in the
-/// current directory and stores half-precision tensors. Existing files are never overwritten.
+/// Convert an imported tensor-only checkpoint state into Montgomery's versioned native Burnpack
+/// format. All families use states generated by `tools/export_checkpoint_state.py`. The output is
+/// `.bpk` in the current directory and stores half-precision tensors. Existing files are
+/// never overwritten.
#[cfg(feature = "pretrained")]
#[derive(Debug)]
pub struct PackedWeights {
@@ -2317,6 +2396,12 @@ pub fn pack_weights_to(
) -> Result {
let input = input.into();
let output = output.into();
+ if input.extension().and_then(|value| value.to_str()) != Some("pt") {
+ return Err(
+ "pack-weights input must be a tensor-only .pt state produced by tools/export_checkpoint_state.py"
+ .into(),
+ );
+ }
if output.extension().and_then(|value| value.to_str()) != Some("bpk") {
return Err("native weight artifact output must use the .bpk extension".into());
}
@@ -2452,12 +2537,25 @@ pub fn pack_weights_to(
fn image_to_tensor(image: DynamicImage, device: &Device) -> Tensor {
let rgb = image.into_rgb8();
- let shape = [rgb.height() as usize, rgb.width() as usize, 3];
+ let width = rgb.width() as usize;
+ let height = rgb.height() as usize;
+ let raw = rgb.into_raw();
+ // Fuse the HWC->CHW transpose with the u8->float cast on the host so the backend
+ // receives contiguous CHW directly (previously this uploaded HWC then ran a separate
+ // `permute` kernel over ~1.2M floats).
+ let mut chw = vec![0.0f32; width * height * 3];
+ let plane = width * height;
+ let (pixels, remainder) = raw.as_chunks::<3>();
+ debug_assert!(remainder.is_empty());
+ for (index, pixel) in pixels.iter().enumerate() {
+ chw[index] = pixel[0] as f32;
+ chw[plane + index] = pixel[1] as f32;
+ chw[2 * plane + index] = pixel[2] as f32;
+ }
Tensor::::from_data(
- TensorData::new(rgb.into_raw(), shape).convert::(),
+ TensorData::new(chw, [3, height, width]).convert::(),
device,
)
- .permute([2, 0, 1])
}
/// Select the strongest detections from decoded end-to-end one2one predictions.
@@ -2483,34 +2581,49 @@ pub(crate) fn end2end_topk_detections(
let image_scores = &scores[image * anchors * classes..(image + 1) * anchors * classes];
let image_boxes = &boxes[image * anchors * 4..(image + 1) * anchors * 4];
- let best_scores = (0..anchors)
- .map(|anchor| {
- let row = &image_scores[anchor * classes..(anchor + 1) * classes];
- row.iter().copied().fold(f32::NEG_INFINITY, f32::max)
- })
- .collect::>();
- let mut anchor_order = (0..anchors).collect::>();
+ let mut best_scores = Vec::with_capacity(anchors);
+ for anchor in 0..anchors {
+ let row = &image_scores[anchor * classes..(anchor + 1) * classes];
+ let mut best = f32::NEG_INFINITY;
+ for &score in row {
+ if score > best {
+ best = score;
+ }
+ }
+ best_scores.push(best);
+ }
+ // Anchors whose best score is below threshold cannot survive the confidence
+ // filter below, so drop them before the (anchor, class) expansion. When more
+ // than `keep` anchors remain above threshold this selects the same top-`keep`
+ // set as a full sort; otherwise the survivors are identical after filtering.
+ let mut anchor_order: Vec = (0..anchors)
+ .filter(|&anchor| best_scores[anchor] >= confidence_threshold)
+ .collect();
+ if anchor_order.len() > keep {
+ anchor_order
+ .select_nth_unstable_by(keep, |&a, &b| best_scores[b].total_cmp(&best_scores[a]));
+ anchor_order.truncate(keep);
+ }
anchor_order.sort_unstable_by(|&a, &b| best_scores[b].total_cmp(&best_scores[a]));
- anchor_order.truncate(keep);
- let mut candidates = Vec::with_capacity(keep * classes);
+ let mut candidates = Vec::with_capacity(anchor_order.len() * classes.min(8));
for (selected_index, &anchor) in anchor_order.iter().enumerate() {
+ let base = anchor * classes;
for class in 0..classes {
- candidates.push((
- image_scores[anchor * classes + class],
- selected_index,
- class,
- ));
+ let score = image_scores[base + class];
+ if score >= confidence_threshold {
+ candidates.push((score, selected_index, class));
+ }
}
}
+ if candidates.len() > keep {
+ candidates.select_nth_unstable_by(keep, |a, b| b.0.total_cmp(&a.0));
+ candidates.truncate(keep);
+ }
candidates.sort_unstable_by(|a, b| b.0.total_cmp(&a.0));
- candidates.truncate(keep);
let mut per_class = (0..classes).map(|_| Vec::new()).collect::>();
for (score, selected_index, class) in candidates {
- if score < confidence_threshold {
- continue;
- }
let anchor = anchor_order[selected_index];
let bbox = &image_boxes[anchor * 4..anchor * 4 + 4];
per_class[class].push(BoundingBox {
diff --git a/src/main.rs b/src/main.rs
index fa98dad..b821869 100644
--- a/src/main.rs
+++ b/src/main.rs
@@ -4,6 +4,8 @@ use std::path::PathBuf;
use burn::backend::Wgpu;
use burn::tensor::{Device, backend::Backend};
use burn_flex::Flex;
+#[cfg(feature = "training")]
+use clap::ArgGroup;
use clap::{Args as ClapArgs, Parser, Subcommand, ValueEnum};
#[cfg(feature = "onnx")]
use montgomery::export::{
@@ -13,7 +15,8 @@ use montgomery::export::{
use montgomery::training::automatic_worker_count;
#[cfg(feature = "training")]
use montgomery::training::runtime::{
- TrainingRequest, export as export_training, train as train_native, validate as validate_native,
+ TrainingInitialization, TrainingRequest, export as export_training, train as train_native,
+ validate as validate_native,
};
use montgomery::{
ModelId, ModelTask, PredictOptions, Predictor, annotate, annotate_segmentation, pack_weights,
@@ -44,7 +47,7 @@ struct Args {
enum Command {
/// Run detection, instance segmentation, or classification on an image.
Predict(PredictArgs),
- /// Pack an imported upstream checkpoint into a versioned native Burnpack artifact.
+ /// Pack an imported tensor-only state into a versioned native Burnpack artifact.
PackWeights(PackWeightsArgs),
/// Export the exact loaded Burn model weights to a validated portable ONNX artifact.
#[cfg(feature = "onnx")]
@@ -62,11 +65,25 @@ enum Command {
#[cfg(feature = "training")]
#[derive(Debug, ClapArgs)]
+#[command(group(
+ ArgGroup::new("initialization")
+ .required(true)
+ .multiple(false)
+ .args(["architecture", "model", "resume"])
+))]
struct TrainArgs {
- #[arg(long)]
- model: ModelId,
- #[arg(long)]
- data: PathBuf,
+ /// Initialize a new model architecture from scratch.
+ #[arg(long)]
+ architecture: Option,
+ /// Initialize from a pretrained Montgomery .bpk model; its architecture is read from metadata.
+ #[arg(long, value_name = "MODEL.bpk")]
+ model: Option,
+ /// Resume a full native training checkpoint; its model and dataset configuration are retained.
+ #[arg(long, value_name = "CHECKPOINT")]
+ resume: Option,
+ /// Dataset manifest for a scratch or pretrained run. Resume uses the checkpoint's dataset.
+ #[arg(long, required_unless_present = "resume", conflicts_with = "resume")]
+ data: Option,
#[arg(long, default_value_t = 100)]
epochs: usize,
#[arg(long, default_value_t = 8)]
@@ -90,12 +107,6 @@ struct TrainArgs {
/// Run data -> forward -> loss -> backward without mutating model or optimizer state.
#[arg(long)]
dry_run: bool,
- /// Resume a full native checkpoint. Model, task, classes, and dataset metadata are immutable.
- #[arg(long)]
- resume: Option,
- /// Initialize from an official tensor-only checkpoint. Mutually exclusive with --resume.
- #[arg(long)]
- weights: Option,
/// Confidence floor used for AP validation (low by default to preserve the PR curve).
#[arg(long)]
val_confidence: Option,
@@ -137,12 +148,9 @@ struct ExportTrainingArgs {
#[cfg(feature = "onnx")]
#[derive(Debug, ClapArgs)]
struct ExportOnnxArgs {
- /// Model architecture represented by the checkpoint.
- #[arg(long)]
- model: ModelId,
- /// Local Montgomery .bpk artifact (defaults to .bpk).
- #[arg(long)]
- weights: Option,
+ /// Montgomery .bpk model to export; architecture is read from artifact metadata.
+ #[arg(long, value_name = "MODEL.bpk")]
+ model: PathBuf,
/// Final ONNX path (defaults to .onnx). A missing suffix is added explicitly.
#[arg(long)]
output: Option,
@@ -193,11 +201,11 @@ struct ExportOnnxArgs {
struct PackWeightsArgs {
/// Model architecture represented by the checkpoint.
#[arg(long)]
- model: ModelId,
+ architecture: ModelId,
- /// Official YOLOX .pth or tensor-only state produced by the Ultralytics development bridge.
- #[arg(long)]
- input: PathBuf,
+ /// Tensor-only state produced by tools/export_checkpoint_state.py.
+ #[arg(long, value_name = "STATE.pt")]
+ state: PathBuf,
/// Native output artifact (defaults to .bpk); must not already exist.
#[arg(long)]
@@ -213,22 +221,16 @@ struct PredictArgs {
#[arg(long)]
source: PathBuf,
- /// Model architecture and scale to run: yolox-nano/tiny/s/m/l/x, yolov3-tinyu,
- /// yolov10n/s/m/b/l/x, yolo11n/s/m/l/x, yolo11n/s/m/l/x-seg, yolo11n/s/m/l/x-cls,
- /// yolov8n/s/m/l/x, yolov8n/s/m/l/x-seg, yolov8n/s/m/l/x-cls, yolo12n/s/m/l/x,
- /// yolo26n/s/m/l/x, yolo26n/s/m/l/x-seg, or yolo26n/s/m/l/x-cls.
- #[arg(long)]
- model: ModelId,
-
- /// Local Montgomery .bpk artifact (defaults to .bpk).
- #[arg(long)]
- weights: Option,
+ /// Montgomery .bpk model to run; architecture and task are read from artifact metadata.
+ #[arg(long, value_name = "MODEL.bpk")]
+ model: PathBuf,
/// Compute device for inference.
#[arg(long, value_enum, default_value = "cpu")]
device: DeviceSelection,
- /// Annotated output image (defaults to -detections.png).
+ /// Annotated output image (defaults to -detections.png, or
+ /// -segmentation.png with --masks).
#[arg(short, long)]
output: Option,
@@ -251,12 +253,15 @@ struct PredictArgs {
json: bool,
}
-fn default_output(input: &std::path::Path) -> PathBuf {
+fn default_output(input: &std::path::Path, masks: bool) -> PathBuf {
let stem = input
.file_stem()
.and_then(|s| s.to_str())
.unwrap_or("prediction");
- input.with_file_name(format!("{stem}-detections.png"))
+ // The suffix names the rendered content: mask outlines only appear with --masks, so a
+ // segmentation rendering must not default to a *-detections.png path.
+ let suffix = if masks { "segmentation" } else { "detections" };
+ input.with_file_name(format!("{stem}-{suffix}.png"))
}
fn main() -> montgomery::Result<()> {
@@ -273,12 +278,12 @@ fn main() -> montgomery::Result<()> {
Command::Predict(args) => predict(args),
Command::PackWeights(args) => {
let packed = match args.output {
- Some(output) => pack_weights_to(args.model, &args.input, output)?,
- None => pack_weights(args.model, &args.input)?,
+ Some(output) => pack_weights_to(args.architecture, &args.state, output)?,
+ None => pack_weights(args.architecture, &args.state)?,
};
eprintln!(
"Packed {} weights into {} ({} bytes, SHA-256 {})",
- args.model,
+ args.architecture,
packed.path.display(),
packed.bytes,
packed.sha256,
@@ -289,8 +294,14 @@ fn main() -> montgomery::Result<()> {
Command::ExportOnnx(args) => export_onnx_command(args),
#[cfg(feature = "training")]
Command::Train(args) => {
+ let initialization = match (args.architecture, args.model, args.resume) {
+ (Some(architecture), None, None) => TrainingInitialization::Scratch(architecture),
+ (None, Some(model), None) => TrainingInitialization::Pretrained(model),
+ (None, None, Some(checkpoint)) => TrainingInitialization::Resume(checkpoint),
+ _ => unreachable!("clap enforces exactly one training initialization mode"),
+ };
let run = train_native(TrainingRequest {
- model: args.model,
+ initialization,
data: args.data,
epochs: args.epochs,
batch_size: args.batch,
@@ -302,8 +313,6 @@ fn main() -> montgomery::Result<()> {
run_root: args.project,
name: args.name,
dry_run: args.dry_run,
- resume: args.resume,
- weights: args.weights,
val_confidence: args.val_confidence,
val_iou: args.val_iou,
max_detections: args.max_detections,
@@ -382,13 +391,12 @@ fn main() -> montgomery::Result<()> {
#[cfg(feature = "onnx")]
fn export_onnx_command(args: ExportOnnxArgs) -> montgomery::Result<()> {
- let weights = args
- .weights
- .unwrap_or_else(|| PathBuf::from(args.model.artifact_filename()));
+ let weights = args.model;
+ let model = ModelId::from_burnpack(&weights)?;
let output = args
.output
- .unwrap_or_else(|| PathBuf::from(format!("{}.onnx", args.model)));
- let mut options = OnnxExportOptions::for_model(args.model, output);
+ .unwrap_or_else(|| weights.with_extension("onnx"));
+ let mut options = OnnxExportOptions::for_model(model, output);
if let Some(imgsz) = &args.imgsz {
let (height, width) = parse_imgsz(imgsz)?;
options.input_shape = [args.batch, 3, height, width];
@@ -409,13 +417,13 @@ fn export_onnx_command(args: ExportOnnxArgs) -> montgomery::Result<()> {
options.force = args.force;
options.keep_intermediate = args.keep_intermediate;
options.reproducible = args.reproducible;
- let artifact = export_onnx(args.model, &weights, options)?;
+ let artifact = export_onnx(model, &weights, options)?;
if args.json {
println!("{}", serde_json::to_string_pretty(&artifact)?);
} else {
eprintln!(
"Exported {} to {} ({} bytes, SHA-256 {}); sidecar {}",
- args.model,
+ model,
artifact.path.display(),
artifact.bytes,
artifact.sha256,
@@ -467,26 +475,20 @@ fn run_predict(
options: PredictOptions,
device: Device,
) -> montgomery::Result<()> {
- let weights = args
- .weights
- .clone()
- .unwrap_or_else(|| PathBuf::from(args.model.artifact_filename()));
- if weights.extension().and_then(|value| value.to_str()) != Some("bpk") {
- return Err(
- "predict --weights requires a native .bpk artifact; convert upstream checkpoints with pack-weights"
- .into(),
- );
- }
+ let model = args.model.clone();
let output = args
.output
.clone()
- .unwrap_or_else(|| default_output(&args.source));
+ .unwrap_or_else(|| default_output(&args.source, args.masks));
- eprintln!("Loading {} weights with Burn...", args.model);
- let predictor: Predictor =
- Predictor::from_checkpoint_on_device(args.model, weights, device, options)?;
+ let predictor: Predictor = Predictor::with_options_on_device(model, device, options)?;
+ eprintln!(
+ "Loaded {} ({}) with Burn.",
+ args.model.display(),
+ predictor.model_id()
+ );
- match args.model.task() {
+ match predictor.task() {
ModelTask::Classification => {
let (image, classifications) = predictor.predict_classification_path(&args.source)?;
report_classifications(
@@ -736,8 +738,12 @@ mod tests {
#[test]
fn derives_default_output_path() {
assert_eq!(
- default_output(std::path::Path::new("photos/dog.jpg")),
+ default_output(std::path::Path::new("photos/dog.jpg"), false),
PathBuf::from("photos/dog-detections.png")
);
+ assert_eq!(
+ default_output(std::path::Path::new("photos/dog.jpg"), true),
+ PathBuf::from("photos/dog-segmentation.png")
+ );
}
}
diff --git a/src/models/yolo11/classification.rs b/src/models/yolo11/classification.rs
index 64fa3e6..82ac5ca 100644
--- a/src/models/yolo11/classification.rs
+++ b/src/models/yolo11/classification.rs
@@ -178,7 +178,7 @@ mod parity_tests {
let checkpoint = std::path::PathBuf::from(concat!("target/", $id, "-state.pt"));
assert!(
checkpoint.exists(),
- "convert {}.pt with tools/export_ultralytics_state.py first",
+ "convert {}.pt with tools/export_checkpoint_state.py first",
$id
);
let worker = std::thread::Builder::new()
diff --git a/src/models/yolo11/model.rs b/src/models/yolo11/model.rs
index 50730fe..ac7254b 100644
--- a/src/models/yolo11/model.rs
+++ b/src/models/yolo11/model.rs
@@ -1080,7 +1080,7 @@ mod tests {
let checkpoint = std::path::PathBuf::from(concat!("target/", $id, "-state.pt"));
assert!(
checkpoint.exists(),
- "convert {}.pt with tools/export_ultralytics_state.py first",
+ "convert {}.pt with tools/export_checkpoint_state.py first",
$id
);
let worker = std::thread::Builder::new()
@@ -1418,7 +1418,7 @@ mod tests {
let checkpoint = std::path::PathBuf::from(concat!("target/", $id, "-state.pt"));
assert!(
checkpoint.exists(),
- "convert {}.pt with tools/export_ultralytics_state.py first",
+ "convert {}.pt with tools/export_checkpoint_state.py first",
$id
);
let worker = std::thread::Builder::new()
diff --git a/src/models/yolo12/model.rs b/src/models/yolo12/model.rs
index 59a60fb..ce76588 100644
--- a/src/models/yolo12/model.rs
+++ b/src/models/yolo12/model.rs
@@ -581,7 +581,7 @@ mod tests {
let checkpoint = std::path::PathBuf::from(concat!("target/", $id, "-state.pt"));
assert!(
checkpoint.exists(),
- "convert {}.pt with tools/export_ultralytics_state.py first",
+ "convert {}.pt with tools/export_checkpoint_state.py first",
$id
);
let worker = std::thread::Builder::new()
diff --git a/src/models/yolo26/classification.rs b/src/models/yolo26/classification.rs
index 833955b..be86921 100644
--- a/src/models/yolo26/classification.rs
+++ b/src/models/yolo26/classification.rs
@@ -487,7 +487,7 @@ mod tests {
let checkpoint = std::path::PathBuf::from(concat!("target/", $id, "-state.pt"));
assert!(
checkpoint.exists(),
- "convert {}.pt with tools/export_ultralytics_state.py first",
+ "convert {}.pt with tools/export_checkpoint_state.py first",
$id
);
let worker = std::thread::Builder::new()
diff --git a/src/models/yolo26/model.rs b/src/models/yolo26/model.rs
index 3b3f986..907cdd1 100644
--- a/src/models/yolo26/model.rs
+++ b/src/models/yolo26/model.rs
@@ -698,7 +698,7 @@ mod tests {
let checkpoint = std::path::PathBuf::from(concat!("target/", $id, "-state.pt"));
assert!(
checkpoint.exists(),
- "convert {}.pt with tools/export_ultralytics_state.py first",
+ "convert {}.pt with tools/export_checkpoint_state.py first",
$id
);
let worker = std::thread::Builder::new()
diff --git a/src/models/yolo26/segmentation.rs b/src/models/yolo26/segmentation.rs
index 505e665..850d00d 100644
--- a/src/models/yolo26/segmentation.rs
+++ b/src/models/yolo26/segmentation.rs
@@ -854,7 +854,7 @@ mod parity_tests {
let checkpoint = std::path::PathBuf::from(concat!("target/", $id, "-state.pt"));
assert!(
checkpoint.exists(),
- "convert {}.pt with tools/export_ultralytics_state.py first",
+ "convert {}.pt with tools/export_checkpoint_state.py first",
$id
);
let worker = std::thread::Builder::new()
diff --git a/src/models/yolov10/model.rs b/src/models/yolov10/model.rs
index caf5bc6..532fca5 100644
--- a/src/models/yolov10/model.rs
+++ b/src/models/yolov10/model.rs
@@ -779,7 +779,7 @@ mod tests {
let checkpoint = std::path::PathBuf::from(concat!("target/", $id, "-state.pt"));
assert!(
checkpoint.exists(),
- "convert {}.pt with tools/export_ultralytics_state.py first",
+ "convert {}.pt with tools/export_checkpoint_state.py first",
$id
);
let worker = std::thread::Builder::new()
diff --git a/src/models/yolov3_tiny/model.rs b/src/models/yolov3_tiny/model.rs
index da45e88..a947fee 100644
--- a/src/models/yolov3_tiny/model.rs
+++ b/src/models/yolov3_tiny/model.rs
@@ -183,7 +183,7 @@ mod tests {
let checkpoint = std::path::PathBuf::from("target/yolov3-tinyu-state.pt");
assert!(
checkpoint.exists(),
- "convert yolov3-tinyu.pt with tools/export_ultralytics_state.py first"
+ "convert yolov3-tinyu.pt with tools/export_checkpoint_state.py first"
);
let worker = std::thread::Builder::new()
.stack_size(64 * 1024 * 1024)
diff --git a/src/models/yolov8/classification.rs b/src/models/yolov8/classification.rs
index 90cbef3..eb0d7a7 100644
--- a/src/models/yolov8/classification.rs
+++ b/src/models/yolov8/classification.rs
@@ -331,7 +331,7 @@ mod tests {
let checkpoint = std::path::PathBuf::from(concat!("target/", $id, "-state.pt"));
assert!(
checkpoint.exists(),
- "convert {}.pt with tools/export_ultralytics_state.py first",
+ "convert {}.pt with tools/export_checkpoint_state.py first",
$id
);
let worker = std::thread::Builder::new()
diff --git a/src/models/yolov8/model.rs b/src/models/yolov8/model.rs
index e43e805..712c0ee 100644
--- a/src/models/yolov8/model.rs
+++ b/src/models/yolov8/model.rs
@@ -1027,7 +1027,7 @@ mod tests {
let checkpoint = std::path::PathBuf::from(concat!("target/", $id, "-state.pt"));
assert!(
checkpoint.exists(),
- "convert {}.pt with tools/export_ultralytics_state.py first",
+ "convert {}.pt with tools/export_checkpoint_state.py first",
$id
);
let worker = std::thread::Builder::new()
@@ -1365,7 +1365,7 @@ mod tests {
let checkpoint = std::path::PathBuf::from(concat!("target/", $id, "-state.pt"));
assert!(
checkpoint.exists(),
- "convert {}.pt with tools/export_ultralytics_state.py first",
+ "convert {}.pt with tools/export_checkpoint_state.py first",
$id
);
let worker = std::thread::Builder::new()
diff --git a/src/models/yolox/weights.rs b/src/models/yolox/weights.rs
index 8dd1813..3129c74 100644
--- a/src/models/yolox/weights.rs
+++ b/src/models/yolox/weights.rs
@@ -21,7 +21,7 @@ pub struct OfficialCheckpoint {
pub sha256: &'static str,
}
-/// Official Apache-2.0 checkpoints used as inputs to `pack-weights` and parity tests.
+/// Official Apache-2.0 checkpoints used as inputs to the tensor-state converter and parity tests.
pub const OFFICIAL_CHECKPOINTS: &[OfficialCheckpoint] = &[
OfficialCheckpoint {
model: "yolox-nano",
diff --git a/src/postprocess.rs b/src/postprocess.rs
index f48bd7a..3bb6866 100644
--- a/src/postprocess.rs
+++ b/src/postprocess.rs
@@ -2,7 +2,6 @@
use alloc::vec::Vec;
use burn::tensor::{ElementConversion, Tensor, backend::Backend};
-use itertools::Itertools;
pub struct BoundingBox {
pub xmin: f32,
@@ -63,34 +62,36 @@ pub fn nms(
.map(|v| v.elem::())
.collect();
- // Per-class filtering based on score
- (0..num_classes)
- .map(|cls_id| {
- // [num_boxes, 1]
- (0..num_boxes)
- .filter_map(|box_idx| {
- let box_cls_idx = cls_idx[box_idx];
- if box_cls_idx != cls_id {
- return None;
- }
- let box_cls_score = cls_score[box_idx];
- if box_cls_score >= score_threshold {
- let bbox = &candidate_boxes[box_idx * 4..box_idx * 4 + 4];
- Some(BoundingBox {
- xmin: bbox[0] - bbox[2] / 2.,
- ymin: bbox[1] - bbox[3] / 2.,
- xmax: bbox[0] + bbox[2] / 2.,
- ymax: bbox[1] + bbox[3] / 2.,
- confidence: box_cls_score,
- })
- } else {
- None
- }
- })
- .sorted_unstable_by(|a, b| a.confidence.partial_cmp(&b.confidence).unwrap())
- .collect::>()
- })
- .collect::>()
+ // Per-class filtering based on score: single pass partitioned by argmax class.
+ // (Previously this scanned all boxes once per class plus a wasted ascending sort
+ // per class that `non_maximum_suppression` immediately re-sorted descending.)
+ let mut per_class: Vec> =
+ (0..num_classes).map(|_| Vec::new()).collect();
+ // Reserve roughly even distribution to avoid regrowth; over-reserve slightly when
+ // num_boxes is small.
+ let per_class_cap = (num_boxes / num_classes.max(1)).max(4);
+ for vec in per_class.iter_mut() {
+ vec.reserve(per_class_cap);
+ }
+ for box_idx in 0..num_boxes {
+ let box_cls_score = cls_score[box_idx];
+ if box_cls_score < score_threshold {
+ continue;
+ }
+ let box_cls_idx = cls_idx[box_idx];
+ if box_cls_idx >= num_classes {
+ continue;
+ }
+ let bbox = &candidate_boxes[box_idx * 4..box_idx * 4 + 4];
+ per_class[box_cls_idx].push(BoundingBox {
+ xmin: bbox[0] - bbox[2] / 2.,
+ ymin: bbox[1] - bbox[3] / 2.,
+ xmax: bbox[0] + bbox[2] / 2.,
+ ymax: bbox[1] + bbox[3] / 2.,
+ confidence: box_cls_score,
+ });
+ }
+ per_class
})
.collect::>();
@@ -102,6 +103,10 @@ pub fn nms(
}
/// Intersection over union of two bounding boxes.
+///
+/// Retained as the shared definition (unit tests, segmentation path reference); the NMS hot
+/// loop inlines the same math with precomputed areas and an AABB early-reject.
+#[allow(dead_code)]
pub fn iou(b1: &BoundingBox, b2: &BoundingBox) -> f32 {
let b1_area = (b1.xmax - b1.xmin).max(0.) * (b1.ymax - b1.ymin).max(0.);
let b2_area = (b2.xmax - b2.xmin).max(0.) * (b2.ymax - b2.ymin).max(0.);
@@ -116,19 +121,48 @@ pub fn iou(b1: &BoundingBox, b2: &BoundingBox) -> f32 {
/// Perform non-maximum suppression over boxes of the same class.
pub fn non_maximum_suppression(bboxes: &mut [Vec], threshold: f32) {
for bboxes_for_class in bboxes.iter_mut() {
- bboxes_for_class.sort_by(|b1, b2| b2.confidence.partial_cmp(&b1.confidence).unwrap());
+ if bboxes_for_class.len() < 2 {
+ continue;
+ }
+ bboxes_for_class.sort_unstable_by(|a, b| b.confidence.total_cmp(&a.confidence));
+ let len = bboxes_for_class.len();
+ let mut areas = Vec::with_capacity(len);
+ for bbox in bboxes_for_class.iter() {
+ areas.push((bbox.xmax - bbox.xmin).max(0.) * (bbox.ymax - bbox.ymin).max(0.));
+ }
let mut current_index = 0;
- for index in 0..bboxes_for_class.len() {
+ for index in 0..len {
+ let (xmin, ymin, xmax, ymax) = {
+ let current = &bboxes_for_class[index];
+ (current.xmin, current.ymin, current.xmax, current.ymax)
+ };
let mut drop = false;
for prev_index in 0..current_index {
- let iou = iou(&bboxes_for_class[prev_index], &bboxes_for_class[index]);
- if iou > threshold {
+ let kept = &bboxes_for_class[prev_index];
+ // AABB early-reject before the full IoU.
+ if kept.xmax <= xmin || kept.xmin >= xmax || kept.ymax <= ymin || kept.ymin >= ymax
+ {
+ continue;
+ }
+ let i_xmin = kept.xmin.max(xmin);
+ let i_xmax = kept.xmax.min(xmax);
+ let i_ymin = kept.ymin.max(ymin);
+ let i_ymax = kept.ymax.min(ymax);
+ let i_area = (i_xmax - i_xmin).max(0.) * (i_ymax - i_ymin).max(0.);
+ if i_area == 0.0 {
+ continue;
+ }
+ let union = areas[prev_index] + areas[index] - i_area;
+ if union > 0.0 && i_area / union > threshold {
drop = true;
break;
}
}
if !drop {
- bboxes_for_class.swap(current_index, index);
+ if current_index != index {
+ bboxes_for_class.swap(current_index, index);
+ areas.swap(current_index, index);
+ }
current_index += 1;
}
}
diff --git a/src/training/loss/yolox.rs b/src/training/loss/yolox.rs
index 13c63a0..d506fa9 100644
--- a/src/training/loss/yolox.rs
+++ b/src/training/loss/yolox.rs
@@ -119,10 +119,15 @@ pub fn tensor_loss(
if targets.len() != batch || classes == 0 {
return Err("YOLOX target batch or class count does not match predictions");
}
- let decoded_host = output.decoded_boxes.clone().detach().into_data();
- let regression_host = output.regression.clone().detach().into_data();
- let objectness_host = output.objectness_logits.clone().detach().into_data();
- let classes_host = output.class_logits.clone().detach().into_data();
+ let [decoded_host, regression_host, objectness_host, classes_host] =
+ burn::tensor::Transaction::default()
+ .register(output.decoded_boxes.clone().detach())
+ .register(output.regression.clone().detach())
+ .register(output.objectness_logits.clone().detach())
+ .register(output.class_logits.clone().detach())
+ .execute()
+ .try_into()
+ .expect("YOLOX assignment transaction must preserve four tensors");
let decoded = decoded_host
.as_slice::()
.map_err(|_| "YOLOX boxes are not f32")?;
diff --git a/src/training/runtime.rs b/src/training/runtime.rs
index 72bd934..3a7a0fe 100644
--- a/src/training/runtime.rs
+++ b/src/training/runtime.rs
@@ -770,10 +770,20 @@ where
Ok(target)
}
+#[derive(Debug, Clone)]
+pub enum TrainingInitialization {
+ /// Construct the named architecture with freshly initialized parameters.
+ Scratch(ModelId),
+ /// Initialize from a self-describing Montgomery inference Burnpack.
+ Pretrained(PathBuf),
+ /// Continue from a full native training checkpoint.
+ Resume(PathBuf),
+}
+
#[derive(Debug, Clone)]
pub struct TrainingRequest {
- pub model: ModelId,
- pub data: PathBuf,
+ pub initialization: TrainingInitialization,
+ pub data: Option,
pub epochs: usize,
pub batch_size: usize,
pub accumulation: usize,
@@ -784,8 +794,6 @@ pub struct TrainingRequest {
pub run_root: PathBuf,
pub name: String,
pub dry_run: bool,
- pub resume: Option,
- pub weights: Option,
pub val_confidence: Option,
pub val_iou: Option,
pub max_detections: Option,
@@ -805,23 +813,41 @@ pub fn train(request: TrainingRequest) -> Result Result> {
- if request.resume.is_some() && request.weights.is_some() {
- return Err("--resume and --weights are mutually exclusive".into());
- }
- let resume_manifest = request
- .resume
- .as_ref()
- .map(crate::training::checkpoint::load)
+ let resume = match &request.initialization {
+ TrainingInitialization::Resume(path) => Some(path),
+ TrainingInitialization::Scratch(_) | TrainingInitialization::Pretrained(_) => None,
+ };
+ let pretrained = match &request.initialization {
+ TrainingInitialization::Pretrained(path) => Some(path),
+ TrainingInitialization::Scratch(_) | TrainingInitialization::Resume(_) => None,
+ };
+ let resume_manifest = resume.map(crate::training::checkpoint::load).transpose()?;
+ let model_id = match &request.initialization {
+ TrainingInitialization::Scratch(model) => *model,
+ TrainingInitialization::Pretrained(path) => ModelId::from_burnpack(path)?,
+ TrainingInitialization::Resume(_) => {
+ resume_manifest
+ .as_ref()
+ .expect("resume initialization loaded a manifest")
+ .config
+ .model
+ .architecture
+ }
+ };
+ let pretrained_classes = pretrained
+ .map(|path| -> Result> {
+ let metadata = crate::artifact_metadata(path)?;
+ Ok(metadata.class_names.map_or_else(
+ || crate::catalog_class_names(model_id).len(),
+ |names| names.len(),
+ ))
+ })
.transpose()?;
- if let Some(manifest) = &resume_manifest
- && manifest.config.model.architecture != request.model
- {
- return Err("--model conflicts with the immutable resume checkpoint architecture".into());
- }
let dataset_path = resume_manifest
.as_ref()
.map(|manifest| manifest.config.data.clone())
- .unwrap_or_else(|| request.data.clone());
+ .or_else(|| request.data.clone())
+ .ok_or("scratch and pretrained training require --data")?;
let dataset = DatasetManifest::load(&dataset_path)?;
let spec = if let Some(manifest) = &resume_manifest {
if manifest.config.model.class_names != dataset.class_names {
@@ -830,7 +856,7 @@ fn train_inner(request: TrainingRequest) -> Result Result Result {{
let mut target = $target;
- if let Some(weights) = request.weights.as_ref() {
- if classes == $projection.official_classes() {
- target.load_pytorch_weights(weights)?;
+ if let Some(model) = pretrained {
+ let source_classes = pretrained_classes
+ .expect("pretrained initialization resolved its class count");
+ if classes == source_classes {
+ target.load_burnpack_weights(model)?;
} else {
+ if source_classes != $projection.official_classes() {
+ return Err(format!(
+ "cannot replace the class projection of a pretrained artifact with {source_classes} classes; this architecture's published pretraining has {} classes",
+ $projection.official_classes()
+ )
+ .into());
+ }
let mut official = $official;
- official.load_pytorch_weights(weights)?;
+ official.load_burnpack_weights(model)?;
target = transfer_pretrained(target, &official, $projection)?;
}
}
@@ -964,7 +999,7 @@ fn train_inner(request: TrainingRequest) -> Result Result Result run!(Yolo11ClsNConfig),
ModelId::Yolo11SCls => run!(Yolo11ClsSConfig),
ModelId::Yolo11MCls => run!(Yolo11ClsMConfig),
diff --git a/tests/integration.rs b/tests/integration.rs
index 04b0a74..e660f90 100644
--- a/tests/integration.rs
+++ b/tests/integration.rs
@@ -261,6 +261,15 @@ fn native_artifact_boundaries_reject_upstream_formats_before_io() {
.expect("upstream checkpoints must be rejected")
.to_string();
assert!(error.contains("native .bpk artifact"), "{error}");
+
+ let error = pack_weights_to(
+ ModelId::YoloxNano,
+ "upstream.pth",
+ "target/rejected-direct-yolox.bpk",
+ )
+ .unwrap_err()
+ .to_string();
+ assert!(error.contains("tensor-only .pt state"), "{error}");
}
#[test]
@@ -334,6 +343,44 @@ fn cli_help_exposes_the_supported_workflows_and_coordinate_contract() {
for contract in ["continuous XYXY pixel edges", "[0, width] x [0, height]"] {
assert!(help.contains(contract), "missing {contract} in:\n{help}");
}
+ assert!(help.contains("--model "), "{help}");
+
+ let output = montgomery(&["pack-weights", "--help"]);
+ assert!(output.status.success(), "{}", stderr(&output));
+ let help = String::from_utf8_lossy(&output.stdout);
+ assert!(help.contains("--architecture "), "{help}");
+ assert!(help.contains("--state "), "{help}");
+}
+
+#[cfg(feature = "training")]
+#[test]
+fn training_cli_requires_one_explicit_initialization_mode() {
+ let output = montgomery(&["train", "--help"]);
+ assert!(output.status.success(), "{}", stderr(&output));
+ let help = String::from_utf8_lossy(&output.stdout);
+ for selector in [
+ "--architecture ",
+ "--model ",
+ "--resume ",
+ ] {
+ assert!(help.contains(selector), "missing {selector} in:\n{help}");
+ }
+
+ let output = montgomery(&[
+ "train",
+ "--architecture",
+ "yolo26n",
+ "--model",
+ "yolo26n.bpk",
+ "--data",
+ "missing.yaml",
+ ]);
+ assert!(!output.status.success());
+ assert!(stderr(&output).contains("cannot be used with"));
+
+ let output = montgomery(&["train", "--model", "upstream.pt", "--data", "missing.yaml"]);
+ assert!(!output.status.success());
+ assert!(stderr(&output).contains("native .bpk artifact"));
}
#[test]
@@ -345,8 +392,6 @@ fn cli_rejects_bad_requests_before_loading_models() {
"--source",
"missing.png",
"--model",
- "yolox-nano",
- "--weights",
"upstream.pth",
],
"native .bpk artifact",
@@ -357,7 +402,7 @@ fn cli_rejects_bad_requests_before_loading_models() {
"--source",
"missing.png",
"--model",
- "yolox-nano",
+ "missing.bpk",
"--confidence",
"1.1",
],
@@ -366,9 +411,9 @@ fn cli_rejects_bad_requests_before_loading_models() {
(
&[
"pack-weights",
- "--model",
+ "--architecture",
"not-a-model",
- "--input",
+ "--state",
"missing.pt",
],
"unknown model 'not-a-model'",
@@ -394,8 +439,6 @@ fn cli_reports_the_gpu_feature_boundary() {
"--source",
"missing.png",
"--model",
- "yolox-nano",
- "--weights",
"missing.bpk",
"--device",
"gpu",
diff --git a/tools/bench_full_convergence.py b/tools/bench_full_convergence.py
index a3fd257..d3b16ec 100644
--- a/tools/bench_full_convergence.py
+++ b/tools/bench_full_convergence.py
@@ -140,18 +140,20 @@ def main() -> None:
model = "yolo26n-seg" if task == "segment" else "yolo26n"
pt = ROOT / "target" / f"{model}.pt"
state = ROOT / "target" / f"{model}-state.pt"
+ burnpack = ROOT / "target" / f"{model}.bpk"
if "ultralytics" in frameworks and not pt.exists():
subprocess.run([sys.executable, "-c", f"from ultralytics import YOLO; YOLO('{model}.pt')"], cwd=ROOT, check=True)
shutil.move(ROOT / f"{model}.pt", pt)
- if "native" in frameworks and not state.exists():
+ if "native" in frameworks and not burnpack.exists():
if not pt.exists():
subprocess.run([sys.executable, "-c", f"from ultralytics import YOLO; YOLO('{model}.pt')"], cwd=ROOT, check=True)
shutil.move(ROOT / f"{model}.pt", pt)
- subprocess.run([sys.executable, str(ROOT / "tools" / "export_ultralytics_state.py"), str(pt), str(state)], cwd=ROOT, check=True)
+ subprocess.run([sys.executable, str(ROOT / "tools" / "export_checkpoint_state.py"), str(pt), str(state)], cwd=ROOT, check=True)
+ subprocess.run([str(BINARY), "pack-weights", "--architecture", model, "--state", str(state), "--output", str(burnpack)], cwd=ROOT, check=True)
if "native" in frameworks:
command = [
- str(BINARY), "train", "--model", model, "--weights", str(state),
+ str(BINARY), "train", "--model", str(burnpack),
"--data", str(native_data), "--epochs", str(args.epochs), "--batch", str(args.batch),
"--imgsz", str(args.imgsz), "--workers", str(args.workers), "--prefetch", "2",
"--save-period", str(args.save_period), "--project", str(args.output / "native"),
diff --git a/tools/bench_training_matrix.py b/tools/bench_training_matrix.py
index e61560f..20b287b 100644
--- a/tools/bench_training_matrix.py
+++ b/tools/bench_training_matrix.py
@@ -71,7 +71,7 @@ class Scenario:
@property
def native_weights(self) -> Path:
- return ROOT / "target" / f"{self.model}-state.pt"
+ return ROOT / "target" / f"{self.model}.bpk"
@property
def ultralytics_weights(self) -> Path:
@@ -129,8 +129,6 @@ def command_for(framework: str, scenario: Scenario, project: Path, name: str) ->
str(NATIVE),
"train",
"--model",
- scenario.model,
- "--weights",
str(scenario.native_weights),
"--data",
str(scenario.native_data),
diff --git a/tools/export_checkpoint_state.py b/tools/export_checkpoint_state.py
new file mode 100644
index 0000000..0984e6f
--- /dev/null
+++ b/tools/export_checkpoint_state.py
@@ -0,0 +1,53 @@
+#!/usr/bin/env python3
+"""Convert a trusted upstream YOLO checkpoint into tensor-only state for Burn import.
+
+The same bridge handles official Ultralytics ``.pt`` and YOLOX ``.pth`` checkpoints. Python and
+PyTorch are development-time dependencies only; Montgomery runtime commands consume ``.bpk``.
+"""
+
+from __future__ import annotations
+
+import argparse
+from collections.abc import Mapping
+from pathlib import Path
+
+import torch
+
+
+def tensor_state(value: object) -> Mapping[str, torch.Tensor]:
+ if hasattr(value, "state_dict"):
+ value = value.state_dict()
+ if not isinstance(value, Mapping) or not value:
+ raise TypeError("checkpoint model payload is not a non-empty state dict or module")
+ invalid = [name for name, tensor in value.items() if not isinstance(name, str) or not torch.is_tensor(tensor)]
+ if invalid:
+ raise TypeError("checkpoint model payload contains non-tensor state entries")
+ return value
+
+
+def main() -> None:
+ parser = argparse.ArgumentParser()
+ parser.add_argument("input", type=Path, help="trusted upstream .pt or .pth checkpoint")
+ parser.add_argument("output", type=Path, help="tensor-only .pt state for pack-weights")
+ args = parser.parse_args()
+
+ # Official checkpoints contain pickled model classes. Only run this converter on trusted files.
+ checkpoint = torch.load(args.input, map_location="cpu", weights_only=False)
+ if not isinstance(checkpoint, Mapping):
+ raise TypeError("checkpoint root is not a mapping")
+ payload = checkpoint.get("ema") or checkpoint.get("model")
+ if payload is None:
+ raise KeyError("checkpoint has neither an 'ema' nor a 'model' payload")
+ state = {
+ name: tensor.detach().float().cpu()
+ for name, tensor in tensor_state(payload).items()
+ }
+
+ args.output.parent.mkdir(parents=True, exist_ok=True)
+ torch.save({"model": state}, args.output)
+ parameters = sum(tensor.numel() for tensor in state.values())
+ print(f"wrote {len(state)} tensors / {parameters:,} values to {args.output}")
+
+
+if __name__ == "__main__":
+ main()
diff --git a/tools/export_ultralytics_state.py b/tools/export_ultralytics_state.py
deleted file mode 100644
index 20b1507..0000000
--- a/tools/export_ultralytics_state.py
+++ /dev/null
@@ -1,36 +0,0 @@
-"""Convert a full Ultralytics checkpoint into a tensor-only state dict for Burn import.
-
-This is a development/build-time bridge. Python and PyTorch are not runtime dependencies of
-Montgomery. The output preserves the original parameter keys and contains no optimizer or model
-object metadata.
-"""
-
-from __future__ import annotations
-
-import argparse
-from pathlib import Path
-
-import torch
-from ultralytics.nn.tasks import torch_safe_load
-
-
-def main() -> None:
- parser = argparse.ArgumentParser()
- parser.add_argument("input", type=Path)
- parser.add_argument("output", type=Path)
- args = parser.parse_args()
-
- checkpoint, _ = torch_safe_load(args.input)
- model = checkpoint.get("ema") or checkpoint["model"]
- state = {name: tensor.detach().float().cpu() for name, tensor in model.state_dict().items()}
-
- args.output.parent.mkdir(parents=True, exist_ok=True)
- torch.save({"model": state}, args.output)
- parameters = sum(tensor.numel() for tensor in state.values())
- print(f"wrote {len(state)} tensors / {parameters:,} values to {args.output}")
- for name in state:
- print(name)
-
-
-if __name__ == "__main__":
- main()
diff --git a/tools/profile_wgpu_training.py b/tools/profile_wgpu_training.py
index 0a8da31..061fba6 100644
--- a/tools/profile_wgpu_training.py
+++ b/tools/profile_wgpu_training.py
@@ -36,10 +36,10 @@ class Workload:
WORKLOADS = (
- Workload("batch1-detect", "yolo26n", "coco8.yaml", "yolo26n-state.pt", 320, 1),
- Workload("batch1-segment", "yolo26n-seg", "coco8-seg.yaml", "yolo26n-seg-state.pt", 320, 1),
- Workload("highres-detect", "yolo26n", "coco8.yaml", "yolo26n-state.pt", 640, 2),
- Workload("medium-segment", "yolo26m-seg", "coco8-seg.yaml", "yolo26m-seg-state.pt", 320, 2),
+ Workload("batch1-detect", "yolo26n", "coco8.yaml", "yolo26n.bpk", 320, 1),
+ Workload("batch1-segment", "yolo26n-seg", "coco8-seg.yaml", "yolo26n-seg.bpk", 320, 1),
+ Workload("highres-detect", "yolo26n", "coco8.yaml", "yolo26n.bpk", 640, 2),
+ Workload("medium-segment", "yolo26m-seg", "coco8-seg.yaml", "yolo26m-seg.bpk", 320, 2),
)
@@ -69,8 +69,7 @@ def main() -> None:
for tasks_max in args.tasks:
for repeat in range(args.repeats):
command = [
- str(args.binary), "train", "--model", workload.model,
- "--weights", str(ROOT / "target" / workload.weights),
+ str(args.binary), "train", "--model", str(ROOT / "target" / workload.weights),
"--data", str(ROOT / "target" / "performance-comparison" / "data" / workload.data),
"--epochs", "1", "--batch", str(workload.batch), "--imgsz", str(workload.imgsz),
"--workers", "4", "--prefetch", "2", "--seed", "0",