File size: 1,906 Bytes
9b92c75 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 | 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() |