From a33093a46d0dadfd18ea62ae6424ebf27ea1a750 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Jos=C3=A9=20D=C3=ADaz?= Date: Thu, 3 Sep 2026 21:54:05 +0200 Subject: [PATCH 1/4] performance improvements --- src/data/letterbox.rs | 96 +++++++++++----- src/lib.rs | 222 ++++++++++++++++++++++++++----------- src/postprocess.rs | 102 +++++++++++------ src/training/loss/yolox.rs | 13 ++- 4 files changed, 305 insertions(+), 128 deletions(-) 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..8b4bed9 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -1295,30 +1295,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 +1428,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 +1461,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 +1561,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 @@ -2452,12 +2522,23 @@ 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; + for (index, pixel) in raw.chunks_exact(3).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 +2564,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/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")?; From 4f00865039b037a7a51910675b239fa0e26ef17e Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Jos=C3=A9=20D=C3=ADaz?= Date: Thu, 3 Sep 2026 22:00:39 +0200 Subject: [PATCH 2/4] tiny error in cli --- src/main.rs | 18 +++++++++++++----- 1 file changed, 13 insertions(+), 5 deletions(-) diff --git a/src/main.rs b/src/main.rs index fa98dad..57d8c48 100644 --- a/src/main.rs +++ b/src/main.rs @@ -228,7 +228,8 @@ struct PredictArgs { #[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 +252,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<()> { @@ -480,7 +484,7 @@ fn run_predict( 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 = @@ -736,8 +740,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") + ); } } From 4c3deb8ac1f9df6812bd42730406a6079befda5e Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Jos=C3=A9=20D=C3=ADaz?= Date: Fri, 4 Sep 2026 02:37:56 +0200 Subject: [PATCH 3/4] more strict CLI regarding model passing --- AGENTS.md | 13 ++-- README.md | 36 ++++++--- docs/MODEL_BRINGUP.md | 18 ++--- src/lib.rs | 27 +++++-- src/main.rs | 116 ++++++++++++++-------------- src/models/yolo11/classification.rs | 2 +- src/models/yolo11/model.rs | 4 +- src/models/yolo12/model.rs | 2 +- src/models/yolo26/classification.rs | 2 +- src/models/yolo26/model.rs | 2 +- src/models/yolo26/segmentation.rs | 2 +- src/models/yolov10/model.rs | 2 +- src/models/yolov3_tiny/model.rs | 2 +- src/models/yolov8/classification.rs | 2 +- src/models/yolov8/model.rs | 4 +- src/models/yolox/weights.rs | 2 +- src/training/runtime.rs | 89 ++++++++++++++------- tests/integration.rs | 57 ++++++++++++-- tools/bench_full_convergence.py | 8 +- tools/bench_training_matrix.py | 4 +- tools/export_checkpoint_state.py | 53 +++++++++++++ tools/export_ultralytics_state.py | 36 --------- tools/profile_wgpu_training.py | 11 ++- 23 files changed, 307 insertions(+), 187 deletions(-) create mode 100644 tools/export_checkpoint_state.py delete mode 100644 tools/export_ultralytics_state.py diff --git a/AGENTS.md b/AGENTS.md index 52ccaae..ad45fe3 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,17 +32,18 @@ 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. diff --git a/README.md b/README.md index 6930c96..fb6eacb 100644 --- a/README.md +++ b/README.md @@ -4,7 +4,7 @@ Montgomery -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/lib.rs b/src/lib.rs index 8b4bed9..3cf11ff 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(), ); } @@ -2359,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 { @@ -2387,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()); } diff --git a/src/main.rs b/src/main.rs index 57d8c48..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,16 +221,9 @@ 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")] @@ -277,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, @@ -293,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, @@ -306,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, @@ -386,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]; @@ -413,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, @@ -471,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, 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( 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/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", From 71859480b18a117c527a8c491b46637fc2537103 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Jos=C3=A9=20D=C3=ADaz?= Date: Fri, 4 Sep 2026 02:53:45 +0200 Subject: [PATCH 4/4] as chunks fix --- AGENTS.md | 18 ++++++++++-------- src/lib.rs | 4 +++- 2 files changed, 13 insertions(+), 9 deletions(-) diff --git a/AGENTS.md b/AGENTS.md index ad45fe3..b8c3cab 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -50,13 +50,19 @@ Python/PyTorch is conversion- and development-time only; normal inference is Rus ## 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: @@ -65,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/src/lib.rs b/src/lib.rs index 3cf11ff..1846e52 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -2545,7 +2545,9 @@ fn image_to_tensor(image: DynamicImage, device: &Device) -> Tenso // `permute` kernel over ~1.2M floats). let mut chw = vec![0.0f32; width * height * 3]; let plane = width * height; - for (index, pixel) in raw.chunks_exact(3).enumerate() { + 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;