Image Classification
vision
convnet
mathmanu commited on
Commit
6e45d35
Β·
verified Β·
1 Parent(s): 87f541c

Add convnext model files

Browse files
README.md ADDED
@@ -0,0 +1,186 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ ---
2
+ license: bsd-3-clause
3
+ tags:
4
+ - vision
5
+ - image-classification
6
+ - convnet
7
+ datasets:
8
+ - imagenet-1k
9
+ ---
10
+
11
+ <div align="center">
12
+
13
+ # ConvNeXt for TI EdgeAI
14
+
15
+ ### A ConvNet for the 2020s β€” Pure ConvNet Matching Transformer Accuracy
16
+
17
+ [![License](https://img.shields.io/badge/License-BSD--3--Clause-blue?style=for-the-badge)](https://opensource.org/licenses/BSD-3-Clause)
18
+ [![Framework](https://img.shields.io/badge/Framework-ONNX-orange?style=for-the-badge)](https://onnx.ai/)
19
+ [![Task](https://img.shields.io/badge/Task-Classification-green?style=for-the-badge)](https://github.com/TexasInstruments/edgeai)
20
+ [![Dataset](https://img.shields.io/badge/Dataset-ImageNet--1K-blueviolet?style=for-the-badge)](http://www.image-net.org/)
21
+
22
+ </div>
23
+
24
+ ---
25
+
26
+ ## Overview
27
+
28
+ **ConvNeXt** is a pure convolutional network modernized by incorporating design principles from Vision Transformers (Swin Transformer). Introduced in [*A ConvNet for the 2020s*](https://arxiv.org/abs/2201.03545) (Liu et al., CVPR 2022), ConvNeXt matches or surpasses Swin Transformers in accuracy while retaining the simplicity, efficiency, and hardware-friendliness of standard CNNs β€” no attention mechanisms, no positional encodings.
29
+
30
+ Pretrained weights are sourced from **torchvision** (BSD-3-Clause), trained on ImageNet-1K using a modernized training recipe.
31
+
32
+ All variants take a **224Γ—224** input with ImageNet normalization (mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]). The ideal pre-crop resize is variant-specific (236px for Tiny, 230px for Small, 232px for Base/Large) and is already set correctly in each variant's config YAML.
33
+
34
+ ---
35
+
36
+ ## Model Variants
37
+
38
+ | Model | Architecture | Params | GFLOPs | Top-1 Acc | Top-5 Acc | Validated Devices | Config |
39
+ |-------|-------------|--------|--------|-----------|-----------|--------------------|--------|
40
+ | `convnext_tiny` | ConvNeXt-Tiny | 28.6M | 4.46 | **82.52%** | 96.15% | TDA4VH | [convnext_tiny_config.yaml](convnext_tiny_config.yaml) |
41
+ | `convnext_small` | ConvNeXt-Small | 50.2M | 8.68 | **83.62%** | 96.65% | TDA4VH | [convnext_small_config.yaml](convnext_small_config.yaml) |
42
+ | `convnext_base` | ConvNeXt-Base | 88.6M | 15.36 | **84.06%** | 96.87% | TDA4VH | [convnext_base_config.yaml](convnext_base_config.yaml) |
43
+ | `convnext_large` | ConvNeXt-Large | 197.8M | 34.36 | **84.41%** | 96.98% | TDA4VH | [convnext_large_config.yaml](convnext_large_config.yaml) |
44
+
45
+ **Recommended for edge deployment:** `convnext_tiny` delivers competitive accuracy (82.5%) at the lowest compute (4.46 GFLOPs, 28.6M params), making it the most practical choice for edge deployment. Larger variants offer incremental accuracy gains at significantly higher compute cost.
46
+
47
+ ---
48
+
49
+ ## Quick Start
50
+
51
+ ### Prerequisites
52
+
53
+ ```bash
54
+ pip install torch torchvision onnx>=1.22.0 onnxruntime>=1.23.2
55
+ # Optional but recommended for model optimization:
56
+ pip install onnx-simplifier
57
+ ```
58
+
59
+ ### Export the Model
60
+
61
+ ```bash
62
+ # Export the default model (convnext_tiny)
63
+ python prepare_model.py
64
+
65
+ # Export a specific model variant
66
+ python prepare_model.py --model convnext_base
67
+
68
+ # Export all supported models
69
+ python prepare_model.py --model all
70
+
71
+ # List all available variants
72
+ python prepare_model.py --list-models
73
+
74
+ # Use a custom checkpoint
75
+ python prepare_model.py --model convnext_tiny --weights /path/to/checkpoint.pth
76
+ ```
77
+
78
+ The script automatically:
79
+ - Downloads pretrained ImageNet-1K weights from torchvision (first run only)
80
+ - Exports the model to ONNX (opset 17) with static input shape [1, 3, 224, 224]
81
+ - Runs ONNX shape inference
82
+ - Optionally simplifies the graph with onnx-simplifier
83
+
84
+ ### Compile and Infer uing edgeai-tidlrunner
85
+
86
+ > **Note:** Run the commands below from inside the `tidlrunner` directory (the cloned [edgeai-tidlrunner](https://github.com/TexasInstruments/edgeai-tidlrunner) repository), with `--config_path` pointing to this model's config file.
87
+
88
+ **Compile using edgeai-tidlrunner - on PC**
89
+
90
+ ```bash
91
+ cd /path/to/edgeai-tidlrunner
92
+ tidlrunner-cli compile --target_device J784S4 \
93
+ --config_path /path/to/convnext_tiny_config.yaml
94
+ ```
95
+
96
+ **Run Inference Benchmark - on device**
97
+
98
+ ```bash
99
+ cd /path/to/edgeai-tidlrunner
100
+ tidlrunner-cli infer --target_device J784S4 \
101
+ --config_path /path/to/convnext_tiny_config.yaml
102
+ ```
103
+
104
+ ### Compile and Infer using edgeai-tidl-tools (Advanced):
105
+
106
+ Follow the instructions at https://github.com/TexasInstruments/edgeai-tidl-tools
107
+
108
+ ### Deploy using edgeai-tidl-tools:
109
+
110
+ Deplyment can be done using **[edgeai-tidl-tools](https://github.com/TexasInstruments/edgeai-tidl-tools)**. For ONNX models, onnxruntime-tidl with TIDL acceleration can be used. Consult the documentation of edgeai-tidl-tools for more details.
111
+
112
+ ---
113
+
114
+ ## Citation
115
+
116
+ ```bibtex
117
+ @inproceedings{liu2022convnet,
118
+ title = {A ConvNet for the 2020s},
119
+ author = {Zhuang Liu and Hanzi Mao and Chao-Yuan Wu and
120
+ Christoph Feichtenhofer and Trevor Darrell and Saining Xie},
121
+ booktitle = {Proceedings of the IEEE/CVF Conference on Computer Vision
122
+ and Pattern Recognition (CVPR)},
123
+ year = {2022},
124
+ url = {https://arxiv.org/abs/2201.03545}
125
+ }
126
+ ```
127
+
128
+ ---
129
+
130
+ ## πŸ”— Resources
131
+
132
+ | Resource | Link |
133
+ |----------|------|
134
+ | **Paper** | [arXiv:2201.03545](https://arxiv.org/abs/2201.03545) |
135
+ | **Source Repo** | [facebookresearch/ConvNeXt](https://github.com/facebookresearch/ConvNeXt) |
136
+ | **Torchvision Docs** | [ConvNeXt](https://docs.pytorch.org/vision/main/models/convnext.html) |
137
+ | **HuggingFace** | [facebook/convnext-tiny-224](https://huggingface.co/facebook/convnext-tiny-224) |
138
+ | **edgeai-tidl-tools** | [GitHub](https://github.com/TexasInstruments/edgeai-tidl-tools) |
139
+ | **edgeai-tidlrunner** | [GitHub](https://github.com/TexasInstruments/edgeai-tidlrunner) |
140
+ | **EdgeAI SDK** | [Documentation](https://github.com/TexasInstruments/edgeai/blob/main/edgeai-mpu/readme_sdk.md) |
141
+
142
+ ---
143
+
144
+ ## Related Models
145
+
146
+ <table>
147
+ <tr>
148
+ <td align="center">
149
+
150
+ **ViT**
151
+ Pure Transformer
152
+ Attention-based
153
+
154
+ </td>
155
+ <td align="center">
156
+
157
+ **DINOv2**
158
+ Self-supervised ViT
159
+ Higher accuracy
160
+
161
+ </td>
162
+ <td align="center">
163
+
164
+ **DINO**
165
+ Self-supervised ViT
166
+ Linear head
167
+
168
+ </td>
169
+ <td align="center">
170
+
171
+ **ResNet**
172
+ CNN baseline
173
+ Lower compute
174
+
175
+ </td>
176
+ </tr>
177
+ </table>
178
+
179
+ ---
180
+
181
+ <div align="center">
182
+
183
+ **Maintained by:** Texas Instruments EdgeAI Team
184
+ **Last Updated:** August 2026
185
+
186
+ </div>
convnext_base_config.yaml ADDED
@@ -0,0 +1,36 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ task_type: classification
2
+ #dataset_category: imagenet
3
+ #calibration_dataset: imagenet
4
+ #input_dataset: imagenet
5
+ dataloader:
6
+ name: image_classification_dataloader
7
+ path: ./data/datasets/imagenetv2c/val
8
+ postprocess: {}
9
+ preprocess:
10
+ resize: 232
11
+ crop: 224
12
+ data_layout: NCHW
13
+ reverse_channels: false
14
+ backend: pil
15
+ interpolation: null
16
+ resize_with_pad: false
17
+ pad_color: 0
18
+ session:
19
+ session_name: onnxrt
20
+ target_device: null
21
+ input_optimization: false
22
+ input_data_layout: NCHW
23
+ input_mean: [123.675, 116.28, 103.53]
24
+ input_scale: [0.017125, 0.017507, 0.017429]
25
+ model_path: convnext_base.onnx
26
+ model_id: cl-mh6036
27
+ input_details: null
28
+ output_details: null
29
+ num_inputs: 1
30
+ model_info:
31
+ metric_reference:
32
+ accuracy_top1%: 84.062
33
+ model_shortlist: 10
34
+ compact_name: convnext-base-224x224
35
+ shortlisted: true
36
+ recommended: false
convnext_large_config.yaml ADDED
@@ -0,0 +1,36 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ task_type: classification
2
+ #dataset_category: imagenet
3
+ #calibration_dataset: imagenet
4
+ #input_dataset: imagenet
5
+ dataloader:
6
+ name: image_classification_dataloader
7
+ path: ./data/datasets/imagenetv2c/val
8
+ postprocess: {}
9
+ preprocess:
10
+ resize: 232
11
+ crop: 224
12
+ data_layout: NCHW
13
+ reverse_channels: false
14
+ backend: pil
15
+ interpolation: null
16
+ resize_with_pad: false
17
+ pad_color: 0
18
+ session:
19
+ session_name: onnxrt
20
+ target_device: null
21
+ input_optimization: false
22
+ input_data_layout: NCHW
23
+ input_mean: [123.675, 116.28, 103.53]
24
+ input_scale: [0.017125, 0.017507, 0.017429]
25
+ model_path: convnext_large.onnx
26
+ model_id: cl-mh6037
27
+ input_details: null
28
+ output_details: null
29
+ num_inputs: 1
30
+ model_info:
31
+ metric_reference:
32
+ accuracy_top1%: 84.414
33
+ model_shortlist: 10
34
+ compact_name: convnext-large-224x224
35
+ shortlisted: true
36
+ recommended: false
convnext_small_config.yaml ADDED
@@ -0,0 +1,36 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ task_type: classification
2
+ #dataset_category: imagenet
3
+ #calibration_dataset: imagenet
4
+ #input_dataset: imagenet
5
+ dataloader:
6
+ name: image_classification_dataloader
7
+ path: ./data/datasets/imagenetv2c/val
8
+ postprocess: {}
9
+ preprocess:
10
+ resize: 230
11
+ crop: 224
12
+ data_layout: NCHW
13
+ reverse_channels: false
14
+ backend: pil
15
+ interpolation: null
16
+ resize_with_pad: false
17
+ pad_color: 0
18
+ session:
19
+ session_name: onnxrt
20
+ target_device: null
21
+ input_optimization: false
22
+ input_data_layout: NCHW
23
+ input_mean: [123.675, 116.28, 103.53]
24
+ input_scale: [0.017125, 0.017507, 0.017429]
25
+ model_path: convnext_small.onnx
26
+ model_id: cl-mh6035
27
+ input_details: null
28
+ output_details: null
29
+ num_inputs: 1
30
+ model_info:
31
+ metric_reference:
32
+ accuracy_top1%: 83.616
33
+ model_shortlist: 10
34
+ compact_name: convnext-small-224x224
35
+ shortlisted: true
36
+ recommended: false
convnext_tiny_config.yaml ADDED
@@ -0,0 +1,36 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ task_type: classification
2
+ #dataset_category: imagenet
3
+ #calibration_dataset: imagenet
4
+ #input_dataset: imagenet
5
+ dataloader:
6
+ name: image_classification_dataloader
7
+ path: ./data/datasets/imagenetv2c/val
8
+ postprocess: {}
9
+ preprocess:
10
+ resize: 236
11
+ crop: 224
12
+ data_layout: NCHW
13
+ reverse_channels: false
14
+ backend: pil
15
+ interpolation: null
16
+ resize_with_pad: false
17
+ pad_color: 0
18
+ session:
19
+ session_name: onnxrt
20
+ target_device: null
21
+ input_optimization: false
22
+ input_data_layout: NCHW
23
+ input_mean: [123.675, 116.28, 103.53]
24
+ input_scale: [0.017125, 0.017507, 0.017429]
25
+ model_path: convnext_tiny.onnx
26
+ model_id: cl-mh6034
27
+ input_details: null
28
+ output_details: null
29
+ num_inputs: 1
30
+ model_info:
31
+ metric_reference:
32
+ accuracy_top1%: 82.520
33
+ model_shortlist: 10
34
+ compact_name: convnext-tiny-224x224
35
+ shortlisted: true
36
+ recommended: true
prepare_model.py ADDED
@@ -0,0 +1,532 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python3
2
+ """
3
+ Export ConvNeXt classification models from torchvision to ONNX
4
+ for TI EdgeAI hardware deployment.
5
+
6
+ Model variants (BSD-3-Clause, ImageNet-1K pretrained via torchvision):
7
+ convnext_tiny – 224Γ—224, 28.59M params, 4.46G FLOPs, top-1 82.520% [default, recommended]
8
+ convnext_small – 224Γ—224, 50.22M params, 8.68G FLOPs, top-1 83.616%
9
+ convnext_base – 224Γ—224, 88.59M params, 15.36G FLOPs, top-1 84.062%
10
+ convnext_large – 224Γ—224, 197.77M params, 34.36G FLOPs, top-1 84.414%
11
+
12
+ Reference paper: A ConvNet for the 2020s (Liu et al., CVPR 2022)
13
+ https://arxiv.org/abs/2201.03545
14
+
15
+ Usage:
16
+ python prepare_model.py
17
+ python prepare_model.py --model convnext_tiny
18
+ python prepare_model.py --model convnext_tiny convnext_small
19
+ python prepare_model.py --model convnext_tiny --shape 224 224
20
+ python prepare_model.py --model all
21
+ python prepare_model.py --model convnext_tiny --weights /path/to/custom.pth
22
+ python prepare_model.py --list-models
23
+ """
24
+
25
+ from __future__ import annotations
26
+
27
+ import argparse
28
+ import importlib
29
+ import os
30
+ import shutil
31
+ import subprocess
32
+ import sys
33
+ import tempfile
34
+
35
+
36
+ # ─────────────────────────────────────────────
37
+ # Model catalogue
38
+ # ─────────────────────────────────────────────
39
+
40
+ MODEL_CATALOG: dict[str, dict] = {
41
+ "convnext_tiny": {
42
+ "tv_weights_cls": "ConvNeXt_Tiny_Weights",
43
+ "backbone": "ConvNeXt-Tiny",
44
+ "shape": (224, 224),
45
+ "resize": 236,
46
+ "params_m": 28.59,
47
+ "flops_g": 4.46,
48
+ "top1_acc": 82.520,
49
+ "top5_acc": 96.146,
50
+ "license": "BSD-3-Clause",
51
+ },
52
+ "convnext_small": {
53
+ "tv_weights_cls": "ConvNeXt_Small_Weights",
54
+ "backbone": "ConvNeXt-Small",
55
+ "shape": (224, 224),
56
+ "resize": 230,
57
+ "params_m": 50.22,
58
+ "flops_g": 8.68,
59
+ "top1_acc": 83.616,
60
+ "top5_acc": 96.650,
61
+ "license": "BSD-3-Clause",
62
+ },
63
+ "convnext_base": {
64
+ "tv_weights_cls": "ConvNeXt_Base_Weights",
65
+ "backbone": "ConvNeXt-Base",
66
+ "shape": (224, 224),
67
+ "resize": 232,
68
+ "params_m": 88.59,
69
+ "flops_g": 15.36,
70
+ "top1_acc": 84.062,
71
+ "top5_acc": 96.870,
72
+ "license": "BSD-3-Clause",
73
+ },
74
+ "convnext_large": {
75
+ "tv_weights_cls": "ConvNeXt_Large_Weights",
76
+ "backbone": "ConvNeXt-Large",
77
+ "shape": (224, 224),
78
+ "resize": 232,
79
+ "params_m": 197.77,
80
+ "flops_g": 34.36,
81
+ "top1_acc": 84.414,
82
+ "top5_acc": 96.976,
83
+ "license": "BSD-3-Clause",
84
+ },
85
+ }
86
+
87
+ DEFAULT_MODEL = "convnext_tiny"
88
+
89
+
90
+ # ─────────────────────────────────────────────
91
+ # Dependency installer
92
+ # ─────────────────────────────────────────────
93
+
94
+ def _pip_install(*packages: str) -> None:
95
+ """Install *packages* via pip, suppressing verbose output."""
96
+ print(f"[DEP] Installing: {', '.join(packages)} …")
97
+ result = subprocess.run(
98
+ [sys.executable, "-m", "pip", "install", *packages],
99
+ stdout=subprocess.DEVNULL,
100
+ stderr=subprocess.PIPE,
101
+ text=True,
102
+ )
103
+ if result.returncode != 0:
104
+ print(f"[DEP] ERROR: pip install failed (exit code {result.returncode}).")
105
+ if result.stderr:
106
+ print(result.stderr.strip())
107
+ print("[DEP] Please install manually and re-run:")
108
+ print(f" pip install {' '.join(packages)}")
109
+ sys.exit(1)
110
+ print("[DEP] Installation complete.\n")
111
+
112
+
113
+ def ensure_dependencies() -> None:
114
+ """Ensure all runtime dependencies are available."""
115
+ needed: list[str] = []
116
+ checks = {
117
+ "torch": "torch",
118
+ "torchvision": "torchvision",
119
+ "onnx": "onnx",
120
+ "onnxsim": "onnx-simplifier",
121
+ "onnxscript": "onnxscript",
122
+ }
123
+ for mod, pkg in checks.items():
124
+ try:
125
+ importlib.import_module(mod)
126
+ print(f"[DEP] βœ” {mod} is already installed.")
127
+ except ImportError:
128
+ print(f"[DEP] ✘ {mod} not found – will install '{pkg}'.")
129
+ needed.append(pkg)
130
+ if needed:
131
+ _pip_install(*needed)
132
+ else:
133
+ print("[DEP] All dependencies satisfied.\n")
134
+
135
+
136
+ # ─────────────────────────────────────────────
137
+ # ONNX post-processing helpers
138
+ # ─────────────────────────────────────────────
139
+
140
+ def _run_shape_inference(onnx_path: str) -> None:
141
+ """Run ONNX shape inference in-place."""
142
+ try:
143
+ import onnx
144
+ import onnx.shape_inference
145
+ print("[POST] Running ONNX shape inference …")
146
+ model = onnx.load(onnx_path)
147
+ model = onnx.shape_inference.infer_shapes(model)
148
+ onnx.save(model, onnx_path)
149
+ print("[POST] Shape inference complete.\n")
150
+ except Exception as exc:
151
+ print(f"[POST] WARNING: shape inference failed ({exc}) – model unchanged.\n")
152
+
153
+
154
+ def _maybe_simplify(onnx_path: str) -> None:
155
+ """Simplify the ONNX model in-place using onnxsim."""
156
+ try:
157
+ import onnx
158
+ import onnxsim
159
+ except ImportError:
160
+ print("[POST] onnxsim not installed – skipping simplification.\n")
161
+ print("[POST] Install with: pip install onnx-simplifier\n")
162
+ return
163
+
164
+ print("[POST] Simplifying ONNX model with onnxsim …")
165
+ try:
166
+ model = onnx.load(onnx_path)
167
+ model_simp, ok = onnxsim.simplify(model)
168
+ if ok:
169
+ onnx.save(model_simp, onnx_path)
170
+ print("[POST] Simplification complete.\n")
171
+ else:
172
+ print("[POST] WARNING: onnxsim validation failed – using original.\n")
173
+ except Exception as exc:
174
+ print(f"[POST] WARNING: onnxsim failed ({exc}) – using original.\n")
175
+
176
+
177
+ # ─────────────────────────────────────────────
178
+ # Model catalogue helpers
179
+ # ─────────────────────────────────────────────
180
+
181
+ def print_model_table() -> None:
182
+ """Print a formatted table of all available models."""
183
+ header = (
184
+ f" {'Variant':<16} {'Backbone':<16} {'Shape':<10} "
185
+ f"{'Params(M)':<10} {'FLOPs(G)':<9} {'Top-1 %':<9} Top-5 %"
186
+ )
187
+ sep = " " + "-" * (len(header) - 2)
188
+ print("\n" + "=" * len(header))
189
+ print(" Available ConvNeXt model variants")
190
+ print("=" * len(header))
191
+ print(header)
192
+ print(sep)
193
+
194
+ for key, info in MODEL_CATALOG.items():
195
+ h, w = info["shape"]
196
+ print(
197
+ f" {key:<16} {info['backbone']:<16} {h}Γ—{w:<5} "
198
+ f"{info['params_m']:<10.2f} {info['flops_g']:<9.2f} "
199
+ f"{info['top1_acc']:<9.3f} {info['top5_acc']:.3f}"
200
+ )
201
+ print("=" * len(header) + "\n")
202
+ print(" Accuracy evaluated on ImageNet-1K val (torchvision pretrained weights).")
203
+ print(" License: BSD-3-Clause (torchvision / PyTorch).\n")
204
+
205
+
206
+ # ─────────────────────────────────────────────
207
+ # Core export
208
+ # ─────────────────────────────────────────────
209
+
210
+ def export_model(
211
+ model_key: str,
212
+ output_dir: str,
213
+ shape: tuple[int, int] | None,
214
+ opset: int,
215
+ batch_size: int,
216
+ verbose: bool,
217
+ custom_weights: str | None,
218
+ force: bool,
219
+ simplify: bool = True,
220
+ ) -> str:
221
+ """
222
+ Load a ConvNeXt model from torchvision and export to ONNX.
223
+
224
+ The exported graph has a single image input (NCHW) and one output:
225
+ output [batch_size, 1000] – raw class logits (ImageNet-1K)
226
+
227
+ Shape inference and optional onnxsim simplification are applied.
228
+
229
+ Args:
230
+ model_key : Key from MODEL_CATALOG (e.g. "convnext_tiny").
231
+ output_dir : Directory where the .onnx file will be saved.
232
+ shape : Custom (H, W) override, or None for model default.
233
+ opset : ONNX opset version (default 17).
234
+ batch_size : Batch size in the exported graph (default 1).
235
+ verbose : Print detailed loading messages.
236
+ custom_weights: Path to a local .pth checkpoint; None = torchvision pretrained.
237
+ force : Re-export even if the destination .onnx already exists.
238
+ simplify : Apply onnxsim after export (default: True).
239
+
240
+ Returns:
241
+ Absolute path of the saved .onnx file.
242
+ """
243
+ import torch
244
+ import torchvision.models as tvm
245
+
246
+ info = MODEL_CATALOG[model_key]
247
+ export_h, export_w = shape if shape is not None else info["shape"]
248
+
249
+ # ── Destination path ──────────────────────────────────────────────────────
250
+ os.makedirs(output_dir, exist_ok=True)
251
+ shape_tag = f"_{export_h}x{export_w}" if shape is not None else ""
252
+ dst_name = f"{model_key}{shape_tag}.onnx"
253
+ dst_path = os.path.join(output_dir, dst_name)
254
+
255
+ if not force and os.path.exists(dst_path):
256
+ print(f"[SKIP] {dst_name} already exists. Use --force to re-export.\n")
257
+ return dst_path
258
+
259
+ print(f"[INFO] Model variant : {model_key}")
260
+ print(f"[INFO] Backbone : {info['backbone']}")
261
+ print(f"[INFO] TV weights cls : {info['tv_weights_cls']}")
262
+ print(f"[INFO] Input shape : {export_h}Γ—{export_w}")
263
+ print(f"[INFO] Batch size : {batch_size}")
264
+ print(f"[INFO] ONNX opset : {opset}")
265
+ print()
266
+
267
+ # ── Load model ───────────────────────────────────────────────────────────
268
+ model_fn = getattr(tvm, model_key)
269
+
270
+ if custom_weights:
271
+ print(f"[INFO] Loading architecture from torchvision, weights from: {custom_weights}")
272
+ model = model_fn(weights=None)
273
+ checkpoint = torch.load(custom_weights, map_location="cpu", weights_only=True)
274
+ state = checkpoint.get("model", checkpoint)
275
+ if isinstance(state, dict) and "module" in state:
276
+ state = state["module"]
277
+ model.load_state_dict(state)
278
+ else:
279
+ print(f"[INFO] Loading pretrained weights from torchvision …")
280
+ print(f"[INFO] (First run may download weights ~110 MB – 755 MB)")
281
+ weights_cls = getattr(tvm, info["tv_weights_cls"])
282
+ model = model_fn(weights=weights_cls.IMAGENET1K_V1)
283
+
284
+ model.eval()
285
+ print(f"[INFO] Model loaded.\n")
286
+
287
+ # ── Dry-run to confirm output shape ──────────────────────────────────────
288
+ dummy = torch.zeros(batch_size, 3, export_h, export_w)
289
+ with torch.no_grad():
290
+ out = model(dummy)
291
+ print(f"[INFO] Output shape : {list(out.shape)}")
292
+ print()
293
+
294
+ # ── Export to ONNX ────────────────────────────────────────────────────────
295
+ print(f"[INFO] Exporting to ONNX (opset {opset}) …")
296
+ with tempfile.TemporaryDirectory(prefix="convnext_export_") as tmp_dir:
297
+ tmp_path = os.path.join(tmp_dir, dst_name)
298
+
299
+ torch.onnx.export(
300
+ model,
301
+ dummy,
302
+ tmp_path,
303
+ input_names=["input"],
304
+ output_names=["output"],
305
+ opset_version=opset,
306
+ do_constant_folding=True,
307
+ verbose=False,
308
+ dynamo=False,
309
+ )
310
+
311
+ # Large models may export with external-data tensor files alongside
312
+ # the .onnx file (e.g. "<name>.onnx.data") – move everything the
313
+ # exporter produced, not just the primary graph file.
314
+ for fname in os.listdir(tmp_dir):
315
+ shutil.move(os.path.join(tmp_dir, fname), os.path.join(output_dir, fname))
316
+
317
+ print(f"[INFO] Raw ONNX written to: {dst_path}")
318
+
319
+ # ── Post-processing ───────────────────────────────────────────────────────
320
+ _run_shape_inference(dst_path)
321
+ if simplify:
322
+ _maybe_simplify(dst_path)
323
+
324
+ size_mb = os.path.getsize(dst_path) / (1024 * 1024)
325
+ print(f"\n[SUCCESS] ONNX model saved to: {dst_path} ({size_mb:.1f} MB)\n")
326
+ return dst_path
327
+
328
+
329
+ # ─────────────────────────────────────────────
330
+ # CLI
331
+ # ─────────────────────────────────────────────
332
+
333
+ def build_parser() -> argparse.ArgumentParser:
334
+ default_output = os.path.dirname(os.path.abspath(__file__))
335
+
336
+ parser = argparse.ArgumentParser(
337
+ description=(
338
+ "Export ConvNeXt pretrained ONNX models.\n\n"
339
+ "Pretrained ImageNet-1K weights are downloaded automatically from\n"
340
+ "torchvision on first use. Run --list-models to see all variants."
341
+ ),
342
+ formatter_class=argparse.RawDescriptionHelpFormatter,
343
+ epilog=(
344
+ "Examples:\n"
345
+ " %(prog)s\n"
346
+ " %(prog)s --model convnext_tiny\n"
347
+ " %(prog)s --model convnext_tiny convnext_small\n"
348
+ " %(prog)s --model convnext_tiny --shape 224 224\n"
349
+ " %(prog)s --model all\n"
350
+ " %(prog)s --model convnext_tiny --weights /path/to/custom.pth\n"
351
+ " %(prog)s --list-models"
352
+ ),
353
+ )
354
+
355
+ # ── Model selection ───────────────────────────────────────────────────────
356
+ parser.add_argument(
357
+ "--model",
358
+ nargs="+",
359
+ default=[DEFAULT_MODEL],
360
+ choices=list(MODEL_CATALOG.keys()) + ["all"],
361
+ metavar="VARIANT",
362
+ help=(
363
+ f"Model variant(s) to export. Use 'all' for all variants. "
364
+ f"Default: {DEFAULT_MODEL}. Run --list-models to see all options."
365
+ ),
366
+ )
367
+
368
+ # ── Export parameters ───────────────────────────────────────────────────���─
369
+ parser.add_argument(
370
+ "--shape",
371
+ nargs=2,
372
+ type=int,
373
+ default=None,
374
+ metavar=("H", "W"),
375
+ help=(
376
+ "Custom input resolution (height width). "
377
+ "Default: each model's native resolution (224Γ—224)."
378
+ ),
379
+ )
380
+ parser.add_argument(
381
+ "--opset",
382
+ type=int,
383
+ default=17,
384
+ metavar="N",
385
+ help="ONNX opset version. Default: 17.",
386
+ )
387
+ parser.add_argument(
388
+ "--batch-size",
389
+ type=int,
390
+ default=1,
391
+ metavar="N",
392
+ help="Batch size embedded in the exported ONNX graph. Default: 1.",
393
+ )
394
+
395
+ # ── Weight source ─────────────────────────────────────────────────────────
396
+ parser.add_argument(
397
+ "--weights",
398
+ default=None,
399
+ metavar="PATH",
400
+ help=(
401
+ "Path to a local .pth checkpoint (optional). "
402
+ "When omitted the official ImageNet-1K pretrained weights are "
403
+ "downloaded automatically from torchvision."
404
+ ),
405
+ )
406
+
407
+ # ── Output ────────────────────────────────────────────────────────────────
408
+ parser.add_argument(
409
+ "--output-dir",
410
+ default=default_output,
411
+ metavar="DIR",
412
+ help=f"Directory where .onnx files will be saved. Default: {default_output}",
413
+ )
414
+ parser.add_argument(
415
+ "--force",
416
+ action="store_true",
417
+ default=False,
418
+ help="Re-export even if the destination .onnx file already exists.",
419
+ )
420
+
421
+ # ── Simplification ────────────────────────────────────────────────────────
422
+ parser.add_argument(
423
+ "--simplify",
424
+ action="store_true",
425
+ default=True,
426
+ help=(
427
+ "Apply onnx-simplifier after export (default: enabled). "
428
+ "Requires: pip install onnx-simplifier. Use --no-simplify to disable."
429
+ ),
430
+ )
431
+ parser.add_argument(
432
+ "--no-simplify",
433
+ dest="simplify",
434
+ action="store_false",
435
+ help="Disable onnx-simplifier after export.",
436
+ )
437
+
438
+ # ── Verbosity ─────────────────────────────────────────────────────────────
439
+ parser.add_argument(
440
+ "--quiet",
441
+ action="store_true",
442
+ default=False,
443
+ help="Suppress verbose output during model loading.",
444
+ )
445
+
446
+ # ── Utility ───────────────────────────────────────────────────────────────
447
+ parser.add_argument(
448
+ "--list-models",
449
+ action="store_true",
450
+ default=False,
451
+ help="Print the model catalogue table and exit.",
452
+ )
453
+
454
+ return parser
455
+
456
+
457
+ # ─────────────────────────────────────────────
458
+ # Entry point
459
+ # ─────────────────────────────────────────────
460
+
461
+ def main() -> None:
462
+ parser = build_parser()
463
+ args = parser.parse_args()
464
+
465
+ if args.list_models:
466
+ print_model_table()
467
+ return
468
+
469
+ # ── Expand "all" keyword ──────────────────────────────────────────────────
470
+ if "all" in args.model:
471
+ args.model = list(MODEL_CATALOG.keys())
472
+
473
+ # ── Warn when --weights is used with multiple models ──────────────────────
474
+ if args.weights and len(args.model) > 1:
475
+ print(
476
+ "[WARN] --weights applies the same checkpoint to every model in "
477
+ "--model.\n This is unusual; pass a single --model variant "
478
+ "when using custom weights.\n"
479
+ )
480
+
481
+ # ── Install dependencies ──────────────────────────────────────────────────
482
+ ensure_dependencies()
483
+
484
+ # ── Export each model ─────────────────────────────────────────────────────
485
+ shape = (args.shape[0], args.shape[1]) if args.shape else None
486
+ output_dir = os.path.abspath(args.output_dir)
487
+
488
+ exported: list[str] = []
489
+ failed: list[str] = []
490
+
491
+ for model_key in args.model:
492
+ print(f"\n{'='*60}")
493
+ print(f" Exporting: {model_key}")
494
+ print(f"{'='*60}\n")
495
+
496
+ try:
497
+ out_path = export_model(
498
+ model_key = model_key,
499
+ output_dir = output_dir,
500
+ shape = shape,
501
+ opset = args.opset,
502
+ batch_size = args.batch_size,
503
+ verbose = not args.quiet,
504
+ custom_weights = args.weights,
505
+ force = args.force,
506
+ simplify = args.simplify,
507
+ )
508
+ exported.append(out_path)
509
+ except SystemExit:
510
+ raise
511
+ except Exception as exc:
512
+ print(f"[ERROR] Export failed for '{model_key}': {exc}")
513
+ failed.append(model_key)
514
+
515
+ # ── Summary ───────────────────────────────────────────────────────────────
516
+ print("\n" + "=" * 60)
517
+ print(" Export Summary")
518
+ print("=" * 60)
519
+ for path in exported:
520
+ size_mb = os.path.getsize(path) / (1024 * 1024)
521
+ print(f" βœ” {os.path.basename(path)} ({size_mb:.1f} MB)")
522
+ print(f" {path}")
523
+ if failed:
524
+ for key in failed:
525
+ print(f" ✘ {key}")
526
+ print(f"\n {len(exported)}/{len(args.model)} model(s) exported successfully.")
527
+ if failed:
528
+ sys.exit(1)
529
+
530
+
531
+ if __name__ == "__main__":
532
+ main()