| 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() |