Download src/objectmodel_v1/export.py from bench-labs/objectmodel-v1: direct link, hf CLI and curl.
- Browser
- Download file 1.91 kB
-
https://huggingface.co/bench-labs/objectmodel-v1/resolve/main/src/objectmodel_v1/export.py
- Command line
-
hf download hf://bench-labs/objectmodel-v1/src/objectmodel_v1/export.py
-
curl -L -o export.py https://huggingface.co/bench-labs/objectmodel-v1/resolve/main/src/objectmodel_v1/export.py
1.91 kB
| 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() |