from __future__ import annotations import argparse from pathlib import Path import torch from torch import nn from .config import apply_overrides, load_config from .model import build_model class ExportModel(nn.Module): def __init__(self, model: nn.Module) -> None: super().__init__() self.model = model def forward(self, images: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: outputs = self.model(images) return outputs["pred_logits"], outputs["pred_boxes"] def main() -> None: parser = argparse.ArgumentParser(description="Export ObjectModel-v1 to ONNX") parser.add_argument("--config", default="configs/objectmodel_v1.yaml") parser.add_argument("--checkpoint", required=True) parser.add_argument("--output", default="objectmodel-v1.onnx") parser.add_argument("--opset", type=int, default=20) parser.add_argument("--set", action="append", default=[]) args = parser.parse_args() config = apply_overrides(load_config(args.config), args.set) model = build_model(config) checkpoint = torch.load(args.checkpoint, map_location="cpu", weights_only=False) model.load_state_dict(checkpoint.get("ema", checkpoint.get("model", checkpoint))) model.eval() wrapper = ExportModel(model) size = model.spec.input_size sample = torch.randn(1, 3, size, size) output = Path(args.output) output.parent.mkdir(parents=True, exist_ok=True) torch.onnx.export( wrapper, (sample,), output, input_names=["images"], output_names=["logits", "boxes"], dynamic_axes={ "images": {0: "batch"}, "logits": {0: "batch"}, "boxes": {0: "batch"}, }, opset_version=args.opset, dynamo=False, ) print(f"Exported {output} ({output.stat().st_size / 1024**2:.2f} MiB)") if __name__ == "__main__": main()