#!/usr/bin/env python3 """ YOLO11-OBB batch export + trim Workflow: 1. Export raw ONNX from ultralytics (includes full post-processing) 2. Automatically trim to NPU-friendly format (9 NHWC head outputs) Usage: python3 export_and_trim.py # all variants, 1024 python3 export_and_trim.py --variant n s # only process n and s python3 export_and_trim.py --variant x --imgsz 640 python3 export_and_trim.py --skip-export # skip export, only trim existing raw onnx """ import argparse import os import re import shutil from typing import Dict, List, Tuple import onnx from onnx import TensorProto, helper, shape_inference ALL_VARIANTS = ["n", "s", "m", "l", "x"] HEAD_RE = re.compile(r".*/cv([234])\.(\d+)\.\d+/Conv_output_\d+$") # ============================================================ # Part 1: Export raw ONNX from ultralytics # ============================================================ def export_raw_onnx(variant: str, imgsz: int, out_dir: str) -> str: from ultralytics import YOLO model_name = f"yolo11{variant}-obb" onnx_final = os.path.join(out_dir, f"{model_name}_{imgsz}x{imgsz}_raw.onnx") os.makedirs(out_dir, exist_ok=True) os.chdir(out_dir) print(f"\n{'='*60}") print(f"[export] {model_name} @ {imgsz}x{imgsz}") print(f"{'='*60}") model = YOLO(f"{model_name}.pt") exported = model.export( format="onnx", imgsz=imgsz, dynamic=False, opset=11, simplify=True, nms=False, ) if exported and os.path.abspath(exported) != os.path.abspath(onnx_final): shutil.move(exported, onnx_final) print(f"[export] -> {onnx_final}") return onnx_final # ============================================================ # Part 2: Trim ONNX for NPU deployment # ============================================================ def trim_onnx(input_path: str, output_path: str, imgsz: int, reg_max: int = 16, ne: int = 1) -> None: model = onnx.load(input_path) print(f"[trim] Loaded: {input_path}") inferred = shape_inference.infer_shapes(model) shape_of: Dict[str, Tuple[int, ...]] = {} for vi in list(inferred.graph.value_info) + list(inferred.graph.output): d = [x.dim_value for x in vi.type.tensor_type.shape.dim] if len(d) == 4 and all(v > 0 for v in d): shape_of[vi.name] = tuple(d) head: Dict[Tuple[int, int], Tuple[str, int, int, int]] = {} for node in inferred.graph.node: if node.op_type != "Conv": continue out = node.output[0] m = HEAD_RE.match(out) if m and out in shape_of: _, c, h, w = shape_of[out] head[(int(m.group(1)), int(m.group(2)))] = (out, c, h, w) assert head, "no /cv{2,3,4}..2/Conv leaf found" scales = sorted({s for (_, s) in head}, key=lambda s: -head[(2, s)][2]) triples, strides = [], [] nc_set, rm_set, ne_set = set(), set(), set() for s in scales: bn, bc, bh, bw = head[(2, s)] cn, cc, _, _ = head[(3, s)] an, ac, _, _ = head[(4, s)] stride_h = imgsz // bh triples.append((bn, cn, an)) strides.append(stride_h) rm_set.add(bc // 4) nc_set.add(cc) ne_set.add(ac) assert len(nc_set) == 1, f"inconsistent nc: {nc_set}" nc = nc_set.pop() print(f"[trim] nc={nc} reg_max={rm_set.pop()} ne={ne_set.pop()} strides={strides}") new_outputs, new_nodes = [], [] for i, (stride, (box, cls, ang)) in enumerate(zip(strides, triples)): for kind, src in [("box", box), ("cls", cls), ("ang", ang)]: dst = f"{kind}_scale{i}_stride{stride}" new_nodes.append(helper.make_node( "Transpose", [src], [dst], perm=[0, 2, 3, 1], name=f"trim_t_{kind}_{i}", )) new_outputs.append(helper.make_tensor_value_info( dst, TensorProto.FLOAT, None)) graph = model.graph del graph.output[:] graph.output.extend(new_outputs) graph.node.extend(new_nodes) # Prune unreachable nodes producers = {o: n for n in graph.node for o in n.output} needed = {vi.name for vi in new_outputs} stack = list(needed) kept = set() while stack: t = stack.pop() n = producers.get(t) if n is None or id(n) in kept: continue kept.add(id(n)) for inp in n.input: if inp and inp not in needed: needed.add(inp) stack.append(inp) kept_nodes = [n for n in graph.node if id(n) in kept] del graph.node[:] graph.node.extend(kept_nodes) used = {inp for n in graph.node for inp in n.input if inp} del graph.initializer[:] graph.initializer.extend([i for i in onnx.load(input_path).graph.initializer if i.name in used]) del graph.value_info[:] model = shape_inference.infer_shapes(model) onnx.checker.check_model(model) onnx.save(model, output_path) print(f"[trim] -> {output_path} ({len(model.graph.node)} nodes)") # ============================================================ # Main: batch export + trim # ============================================================ def main(): ap = argparse.ArgumentParser( description="YOLO11-OBB batch export + trim for NPU (all variants)") ap.add_argument("--variant", type=str, nargs="+", default=ALL_VARIANTS, choices=ALL_VARIANTS, help="Model variants (default: all)") ap.add_argument("--imgsz", type=int, default=1024, help="Input resolution (default: 1024)") ap.add_argument("--out-dir", type=str, default=os.path.dirname(os.path.abspath(__file__))) ap.add_argument("--skip-export", action="store_true", help="Skip ultralytics export, only trim existing raw ONNX") ap.add_argument("--skip-trim", action="store_true", help="Skip trim, only export raw ONNX") args = ap.parse_args() print(f"Variants: {args.variant} imgsz: {args.imgsz}") print(f"Output dir: {args.out_dir}") for v in args.variant: model_name = f"yolo11{v}-obb" raw_path = os.path.join(args.out_dir, f"{model_name}_{args.imgsz}x{args.imgsz}_raw.onnx") trim_path = os.path.join(args.out_dir, f"{model_name}_{args.imgsz}x{args.imgsz}_trim.onnx") try: if not args.skip_export: export_raw_onnx(v, args.imgsz, args.out_dir) if not args.skip_trim: if not os.path.isfile(raw_path): print(f"[skip] raw ONNX not found: {raw_path}") continue trim_onnx(raw_path, trim_path, args.imgsz) except Exception as e: print(f"[FAILED] {model_name}: {e}") import traceback traceback.print_exc() print(f"\nDone. Processed {len(args.variant)} variants.") if __name__ == "__main__": main()