Image Classification
vision
self-supervised
mathmanu commited on
Commit
5d6cefd
·
verified ·
1 Parent(s): b00511d

Add dino model files

Browse files
README.md ADDED
@@ -0,0 +1,187 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ ---
2
+ license: apache-2.0
3
+ tags:
4
+ - vision
5
+ - image-classification
6
+ - self-supervised
7
+ datasets:
8
+ - imagenet-1k
9
+ ---
10
+
11
+ <div align="center">
12
+
13
+ # DINO for TI EdgeAI
14
+
15
+ ### Self-Supervised Vision Transformer Backbone for Image Classification
16
+
17
+ [![License](https://img.shields.io/badge/License-Apache%202.0-blue?style=for-the-badge)](https://opensource.org/licenses/Apache-2.0)
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
+ **DINO** (Self-**Di**stillation with **No** labels) is a self-supervised Vision Transformer pre-training method from Meta AI. The backbone models produce rich feature embeddings that achieve strong performance on ImageNet classification without any labels during pre-training.
29
+
30
+ These ONNX models include the **full backbone + pretrained linear classification head**, outputting 1000-class ImageNet logits `[1, 1000]`. Feature extraction follows DINO's `eval_linear.py` conventions:
31
+ - **ViT-S models**: CLS tokens from last 4 blocks concatenated → `[B, 1536]`
32
+ - **ViT-B models**: CLS token + averaged patch tokens (interleaved) → `[B, 1536]`
33
+ - **ResNet-50**: avgpool output → `[B, 2048]`
34
+
35
+ > See [DINOv2](../DINOv2/) for the improved second-generation models.
36
+
37
+ ---
38
+
39
+ ## Model Variants
40
+
41
+ | Model | Architecture | Params | Linear Top-1 | k-NN Top-1 | Validated Devices | Config |
42
+ |-------|-------------|--------|-------------|-----------|----------|--------|
43
+ | `dino_vits16` | ViT-S/16 | 21M | 77.0% | 74.5% | TDA4VH | [dino_vits16_config.yaml](dino_vits16_config.yaml) |
44
+ | `dino_vits8` | ViT-S/8 | 21M | 79.7% | 78.3% | TDA4VH | [dino_vits8_config.yaml](dino_vits8_config.yaml) |
45
+ | `dino_vitb16` | ViT-B/16 | 85M | 78.2% | 76.1% | TDA4VH | [dino_vitb16_config.yaml](dino_vitb16_config.yaml) |
46
+ | `dino_vitb8` | ViT-B/8 | 85M | 80.1% | 77.4% | TDA4VH | [dino_vitb8_config.yaml](dino_vitb8_config.yaml) |
47
+ | `dino_resnet50` | ResNet-50 | 23M | 75.3% | 67.5% | TDA4VH | [dino_resnet50_config.yaml](dino_resnet50_config.yaml) |
48
+
49
+ **Recommended for edge deployment:** `dino_vits16` (best accuracy/compute trade-off)
50
+
51
+ ---
52
+
53
+ ## Quick Start
54
+
55
+ ### Prerequisites
56
+
57
+ ```bash
58
+ pip install onnx>=1.22.0 onnxruntime>=1.23.2
59
+ ```
60
+
61
+ ### Export the Model
62
+
63
+ ```bash
64
+ # Export the default model (ViT-S/16)
65
+ python prepare_model.py
66
+
67
+ # Export a specific model variant
68
+ python prepare_model.py --model dino_vitb16
69
+
70
+ # Export all supported models
71
+ python prepare_model.py --model all
72
+
73
+ # Re-run shape fixing on an already-exported ONNX
74
+ python prepare_model.py --model dino_vits16 --skip-export
75
+ ```
76
+
77
+ The script automatically:
78
+ - Loads pretrained backbone from PyTorch Hub (`facebookresearch/dino:main`)
79
+ - Downloads pretrained linear classification weights from Meta AI
80
+ - Combines backbone + linear head into a single classification model
81
+ - Exports to ONNX (opset 17) and fixes input shapes to [1, 3, 224, 224]
82
+ - Validates the model outputs `[1, 1000]` class logits
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/dino_vits16_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/dino_vits16_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
+ If you use these models, please cite:
117
+
118
+ ```bibtex
119
+ @inproceedings{caron2021emerging,
120
+ title={Emerging Properties in Self-Supervised Vision Transformers},
121
+ author={Caron, Mathilde and Touvron, Hugo and Misra, Ishan and
122
+ J{\'e}gou, Herv{\'e} and Mairal, Julien and Bojanowski, Piotr
123
+ and Joulin, Armand},
124
+ booktitle={Proceedings of the IEEE/CVF International Conference
125
+ on Computer Vision (ICCV)},
126
+ year={2021}
127
+ }
128
+ ```
129
+
130
+ ---
131
+
132
+ ## 🔗 Resources
133
+
134
+ | Resource | Link |
135
+ |----------|------|
136
+ | **Paper** | [arXiv:2104.14294](https://arxiv.org/abs/2104.14294) |
137
+ | **Source Code** | [facebookresearch/dino](https://github.com/facebookresearch/dino) |
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
+ | **DINOv2** | [Improved successor](../DINOv2/) |
142
+
143
+ ---
144
+
145
+ ## Related Models
146
+
147
+ <table>
148
+ <tr>
149
+ <td align="center">
150
+
151
+ **DINOv2**
152
+ Improved DINO
153
+ Higher accuracy
154
+
155
+ </td>
156
+ <td align="center">
157
+
158
+ **ViT-S/16**
159
+ Recommended
160
+ Best edge trade-off
161
+
162
+ </td>
163
+ <td align="center">
164
+
165
+ **ResNet-50**
166
+ CNN backbone
167
+ Lower compute
168
+
169
+ </td>
170
+ <td align="center">
171
+
172
+ **CLIP**
173
+ Vision-Language
174
+ Zero-shot capable
175
+
176
+ </td>
177
+ </tr>
178
+ </table>
179
+
180
+ ---
181
+
182
+ <div align="center">
183
+
184
+ **Maintained by:** Texas Instruments EdgeAI Team
185
+ **Last Updated:** August 2026
186
+
187
+ </div>
dino_resnet50_config.yaml ADDED
@@ -0,0 +1,34 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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: 256
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: dino_resnet50.onnx
26
+ model_id: cl-mh6020
27
+ input_details: null
28
+ output_details: null
29
+ num_inputs: 1
30
+ model_info:
31
+ metric_reference:
32
+ accuracy_top1%: 75.3
33
+ compact_name: DINO-ResNet-50
34
+ shortlisted: true
dino_vitb16_config.yaml ADDED
@@ -0,0 +1,34 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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: 256
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: dino_vitb16.onnx
26
+ model_id: cl-mh6018
27
+ input_details: null
28
+ output_details: null
29
+ num_inputs: 1
30
+ model_info:
31
+ metric_reference:
32
+ accuracy_top1%: 78.2
33
+ compact_name: DINO-ViT-B/16
34
+ shortlisted: true
dino_vitb8_config.yaml ADDED
@@ -0,0 +1,34 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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: 256
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: dino_vitb8.onnx
26
+ model_id: cl-mh6019
27
+ input_details: null
28
+ output_details: null
29
+ num_inputs: 1
30
+ model_info:
31
+ metric_reference:
32
+ accuracy_top1%: 80.1
33
+ compact_name: DINO-ViT-B/8
34
+ shortlisted: true
dino_vits16_config.yaml ADDED
@@ -0,0 +1,34 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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: 256
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: dino_vits16.onnx
26
+ model_id: cl-mh6016
27
+ input_details: null
28
+ output_details: null
29
+ num_inputs: 1
30
+ model_info:
31
+ metric_reference:
32
+ accuracy_top1%: 77.0
33
+ compact_name: DINO-ViT-S/16
34
+ shortlisted: true
dino_vits8_config.yaml ADDED
@@ -0,0 +1,34 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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: 256
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: dino_vits8.onnx
26
+ model_id: cl-mh6017
27
+ input_details: null
28
+ output_details: null
29
+ num_inputs: 1
30
+ model_info:
31
+ metric_reference:
32
+ accuracy_top1%: 79.7
33
+ compact_name: DINO-ViT-S/8
34
+ shortlisted: true
prepare_model.py ADDED
@@ -0,0 +1,511 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python3
2
+ """
3
+ Export DINO classification models (backbone + linear head) from PyTorch Hub
4
+ to ONNX with fixed input shapes for TI EdgeAI hardware deployment.
5
+
6
+ Each exported model includes the full DINO backbone and the pretrained linear
7
+ classification head, outputting 1000-class ImageNet logits [1, 1000].
8
+
9
+ Feature extraction follows DINO's eval_linear.py conventions:
10
+ ViT-S models : last 4 blocks' CLS tokens concatenated → [B, 384×4 = 1536]
11
+ ViT-B models : CLS token + averaged patch tokens (interleaved) → [B, 768×2 = 1536]
12
+ ResNet-50 : avgpool output → [B, 2048]
13
+
14
+ Supported models (backbone + linear head, 1000-class ImageNet):
15
+ dino_vits16 - ViT-S/16, 21M params, 77.0% linear top-1, 74.5% k-NN top-1
16
+ dino_vits8 - ViT-S/8, 21M params, 79.7% linear top-1, 78.3% k-NN top-1
17
+ dino_vitb16 - ViT-B/16, 85M params, 78.2% linear top-1, 76.1% k-NN top-1
18
+ dino_vitb8 - ViT-B/8, 85M params, 80.1% linear top-1, 77.4% k-NN top-1
19
+ dino_resnet50 - ResNet-50, 23M params, 75.3% linear top-1, 67.5% k-NN top-1
20
+
21
+ Usage:
22
+ python prepare_model.py --model dino_vits16
23
+ python prepare_model.py --model dino_vitb16 --no-simplifier
24
+ python prepare_model.py --model all
25
+ """
26
+
27
+ import sys
28
+ import subprocess
29
+ import tempfile
30
+ from pathlib import Path
31
+
32
+
33
+ # n_last_blocks / avgpool follow DINO's eval_linear.py default args per arch:
34
+ # ViT-S: n_last_blocks=4, avgpool=False → linear_in = 384 * 4 = 1536
35
+ # ViT-B: n_last_blocks=1, avgpool=True → linear_in = 768 * 2 = 1536
36
+ # ResNet: direct avgpool output → linear_in = 2048
37
+ SUPPORTED_MODELS = {
38
+ 'dino_vits16': {
39
+ 'arch': 'ViT-S/16', 'params': '21M', 'accuracy_top1': 77.0, 'knn_top1': 74.5,
40
+ 'n_last_blocks': 4, 'avgpool': False, 'linear_in': 384 * 4,
41
+ },
42
+ 'dino_vits8': {
43
+ 'arch': 'ViT-S/8', 'params': '21M', 'accuracy_top1': 79.7, 'knn_top1': 78.3,
44
+ 'n_last_blocks': 4, 'avgpool': False, 'linear_in': 384 * 4,
45
+ },
46
+ 'dino_vitb16': {
47
+ 'arch': 'ViT-B/16', 'params': '85M', 'accuracy_top1': 78.2, 'knn_top1': 76.1,
48
+ 'n_last_blocks': 1, 'avgpool': True, 'linear_in': 768 * 2,
49
+ },
50
+ 'dino_vitb8': {
51
+ 'arch': 'ViT-B/8', 'params': '85M', 'accuracy_top1': 80.1, 'knn_top1': 77.4,
52
+ 'n_last_blocks': 1, 'avgpool': True, 'linear_in': 768 * 2,
53
+ },
54
+ 'dino_resnet50': {
55
+ 'arch': 'ResNet-50', 'params': '23M', 'accuracy_top1': 75.3, 'knn_top1': 67.5,
56
+ 'n_last_blocks': None, 'avgpool': False, 'linear_in': 2048,
57
+ },
58
+ }
59
+
60
+ _BASE_URL = 'https://dl.fbaipublicfiles.com/dino/'
61
+ LINEAR_WEIGHTS_URLS = {
62
+ 'dino_vits16': _BASE_URL + 'dino_deitsmall16_pretrain/dino_deitsmall16_linearweights.pth',
63
+ 'dino_vits8': _BASE_URL + 'dino_deitsmall8_pretrain/dino_deitsmall8_linearweights.pth',
64
+ 'dino_vitb16': _BASE_URL + 'dino_vitbase16_pretrain/dino_vitbase16_linearweights.pth',
65
+ 'dino_vitb8': _BASE_URL + 'dino_vitbase8_pretrain/dino_vitbase8_linearweights.pth',
66
+ 'dino_resnet50': _BASE_URL + 'dino_resnet50_pretrain/dino_resnet50_linearweights.pth',
67
+ }
68
+
69
+
70
+ def _ensure_dependencies():
71
+ required = {
72
+ 'onnx': 'onnx',
73
+ 'onnxsim': 'onnx-simplifier',
74
+ 'torch': 'torch',
75
+ }
76
+ for module, package in required.items():
77
+ try:
78
+ __import__(module)
79
+ except ImportError:
80
+ print(f"Installing missing dependency: {package}")
81
+ subprocess.check_call([sys.executable, '-m', 'pip', 'install', package])
82
+
83
+
84
+ _ensure_dependencies()
85
+
86
+ import torch
87
+ import torch.nn as nn
88
+ import onnx
89
+ from onnx import shape_inference
90
+ import argparse
91
+
92
+
93
+ class _LinearClassifier(nn.Module):
94
+ """Linear head matching DINO's eval_linear.py LinearClassifier structure."""
95
+ def __init__(self, in_features, num_classes=1000):
96
+ super().__init__()
97
+ self.linear = nn.Linear(in_features, num_classes)
98
+
99
+ def forward(self, x):
100
+ return self.linear(x)
101
+
102
+
103
+ class _DinoViTClassifier(nn.Module):
104
+ """
105
+ DINO ViT backbone + linear head for classification.
106
+
107
+ Feature extraction matches eval_linear.py:
108
+ - Collects CLS tokens from the last n_last_blocks transformer blocks
109
+ - For ViT-B (avgpool=True): interleaves CLS with averaged patch tokens
110
+ using the same stack+flatten as the original code, preserving weight
111
+ compatibility: [CLS[0], patch[0], CLS[1], patch[1], ...]
112
+ """
113
+ def __init__(self, backbone, linear_head, n_last_blocks, avgpool):
114
+ super().__init__()
115
+ self.backbone = backbone
116
+ self.linear_head = linear_head
117
+ self.n = n_last_blocks
118
+ self.avgpool = avgpool
119
+
120
+ def forward(self, x):
121
+ intermediate = self.backbone.get_intermediate_layers(x, self.n)
122
+ feat = torch.cat([layer[:, 0] for layer in intermediate], dim=-1)
123
+ if self.avgpool:
124
+ # Interleave CLS and patch-average as in eval_linear.py:
125
+ # stack → [B, embed, 2] → flatten(1) → [B, embed*2]
126
+ patch_avg = torch.mean(intermediate[-1][:, 1:], dim=1)
127
+ feat = torch.stack([feat, patch_avg], dim=-1).flatten(1)
128
+ return self.linear_head(feat)
129
+
130
+
131
+ class _DinoResNetClassifier(nn.Module):
132
+ """DINO ResNet-50 backbone + linear head for classification."""
133
+ def __init__(self, backbone, linear_head):
134
+ super().__init__()
135
+ self.backbone = backbone
136
+ self.linear_head = linear_head
137
+
138
+ def forward(self, x):
139
+ return self.linear_head(self.backbone(x))
140
+
141
+
142
+ def _build_classifier(model_name, info):
143
+ """
144
+ Load DINO backbone from PyTorch Hub, load pretrained linear weights,
145
+ and return a combined classifier module ready for ONNX export.
146
+
147
+ Returns the combined nn.Module or None on failure.
148
+ """
149
+ print(f"\nLoading backbone from PyTorch Hub:")
150
+ print(f" torch.hub.load('facebookresearch/dino:main', '{model_name}')")
151
+ try:
152
+ backbone = torch.hub.load('facebookresearch/dino:main', model_name, pretrained=True)
153
+ except Exception as e:
154
+ print(f"✗ Failed to load backbone: {e}")
155
+ print(" Ensure you have an internet connection and PyTorch installed.")
156
+ return None
157
+ backbone.eval()
158
+
159
+ print(f"\nDownloading linear weights:")
160
+ print(f" URL: {LINEAR_WEIGHTS_URLS[model_name]}")
161
+ try:
162
+ ckpt = torch.hub.load_state_dict_from_url(
163
+ LINEAR_WEIGHTS_URLS[model_name], map_location='cpu', progress=True
164
+ )
165
+ state_dict = ckpt['state_dict']
166
+ # Saved under DDP → strip 'module.' prefix
167
+ state_dict = {k.replace('module.', ''): v for k, v in state_dict.items()}
168
+ except Exception as e:
169
+ print(f"✗ Failed to download linear weights: {e}")
170
+ return None
171
+
172
+ linear_head = _LinearClassifier(info['linear_in'])
173
+ try:
174
+ linear_head.load_state_dict(state_dict, strict=True)
175
+ print(f"✓ Linear weights loaded ({info['linear_in']} → 1000 classes)")
176
+ except Exception as e:
177
+ print(f"✗ Failed to load linear weights into head: {e}")
178
+ return None
179
+ linear_head.eval()
180
+
181
+ if info['n_last_blocks'] is None:
182
+ model = _DinoResNetClassifier(backbone, linear_head)
183
+ else:
184
+ model = _DinoViTClassifier(backbone, linear_head, info['n_last_blocks'], info['avgpool'])
185
+
186
+ model.eval()
187
+ return model
188
+
189
+
190
+ def export_to_onnx(model_name, output_path, height=224, width=224):
191
+ """
192
+ Build the DINO backbone + linear head and export to ONNX (opset 17).
193
+ """
194
+ info = SUPPORTED_MODELS[model_name]
195
+ print(f"\nDINO Model Export")
196
+ print("=" * 80)
197
+ print(f"Model: {model_name} ({info['arch']})")
198
+ print(f"Params: {info['params']}")
199
+ print(f"Top-1 (lin): {info['accuracy_top1']}%")
200
+ print(f"Top-1 (k-NN): {info['knn_top1']}%")
201
+ print(f"Input shape: [1, 3, {height}, {width}]")
202
+ print(f"Output shape: [1, 1000]")
203
+
204
+ model = _build_classifier(model_name, info)
205
+ if model is None:
206
+ return False
207
+
208
+ dummy_input = torch.randn(1, 3, height, width)
209
+
210
+ # Sanity-check output shape before export
211
+ with torch.no_grad():
212
+ out = model(dummy_input)
213
+ if list(out.shape) != [1, 1000]:
214
+ print(f"✗ Unexpected output shape: {list(out.shape)}, expected [1, 1000]")
215
+ return False
216
+ print(f"\n✓ Output shape verified: {list(out.shape)}")
217
+
218
+ print(f"\nExporting to ONNX (opset 17):")
219
+ print(f" Output: {output_path}")
220
+
221
+ try:
222
+ with tempfile.TemporaryDirectory() as tmpdir:
223
+ tmp_onnx = Path(tmpdir) / f"{model_name}.onnx"
224
+
225
+ torch.onnx.export(
226
+ model,
227
+ dummy_input,
228
+ str(tmp_onnx),
229
+ export_params=True,
230
+ opset_version=17,
231
+ do_constant_folding=True,
232
+ input_names=['input'],
233
+ output_names=['output'],
234
+ dynamic_axes={
235
+ 'input': {0: 'batch_size'},
236
+ 'output': {0: 'batch_size'},
237
+ },
238
+ )
239
+
240
+ probe = onnx.load(str(tmp_onnx), load_external_data=False)
241
+ has_external = any(
242
+ t.data_location == onnx.TensorProto.EXTERNAL
243
+ for t in probe.graph.initializer
244
+ )
245
+
246
+ if has_external:
247
+ # Should not happen for DINO models (<2 GB), but handle gracefully
248
+ print("\nMerging external tensor data...")
249
+ exported = onnx.load(str(tmp_onnx))
250
+ else:
251
+ exported = onnx.load(str(tmp_onnx))
252
+ onnx.save(exported, str(output_path))
253
+
254
+ except Exception as e:
255
+ print(f"✗ ONNX export failed: {e}")
256
+ return False
257
+
258
+ if not output_path.exists():
259
+ print("✗ Export failed: output file not created")
260
+ return False
261
+
262
+ file_size = output_path.stat().st_size
263
+ print(f"✓ Export completed: {file_size:,} bytes ({file_size / 1024 / 1024:.2f} MB)")
264
+ return True
265
+
266
+
267
+ def fix_model_shape(model_path, output_path, batch_size=1, channels=3, height=224, width=224, use_simplifier=True):
268
+ """
269
+ Convert dynamic ONNX model input shape to fixed shape in all layers.
270
+ """
271
+ print(f"\nFixing Model Shapes:")
272
+ print("=" * 80)
273
+ print(f"Input model: {model_path}")
274
+ print(f"Output model: {output_path}")
275
+
276
+ model = onnx.load(str(model_path))
277
+
278
+ graph = model.graph
279
+ input_tensor = None
280
+ for inp in graph.input:
281
+ if any(init.name == inp.name for init in graph.initializer):
282
+ continue
283
+ input_tensor = inp
284
+ break
285
+
286
+ if input_tensor is None:
287
+ print("✗ Error: No input tensor found!")
288
+ return False
289
+
290
+ original_shape = []
291
+ for dim in input_tensor.type.tensor_type.shape.dim:
292
+ if dim.dim_value:
293
+ original_shape.append(str(dim.dim_value))
294
+ elif dim.dim_param:
295
+ original_shape.append(f"'{dim.dim_param}'")
296
+ else:
297
+ original_shape.append("?")
298
+ print(f"Original shape: [{', '.join(original_shape)}]")
299
+
300
+ new_shape = [batch_size, channels, height, width]
301
+ print(f"Fixed shape: {new_shape}")
302
+
303
+ input_tensor.type.tensor_type.shape.ClearField('dim')
304
+ for dim_value in new_shape:
305
+ dim = input_tensor.type.tensor_type.shape.dim.add()
306
+ dim.dim_value = dim_value
307
+
308
+ try:
309
+ model = shape_inference.infer_shapes(model)
310
+ print(f"✓ Propagated shapes through {len(model.graph.value_info)} intermediate tensors")
311
+ except Exception as e:
312
+ print(f"âš  Warning: Shape inference issue: {e}")
313
+
314
+ try:
315
+ onnx.checker.check_model(model)
316
+ print("✓ Model validation passed")
317
+ except Exception as e:
318
+ print(f"✗ Model validation failed: {e}")
319
+ return False
320
+
321
+ if use_simplifier:
322
+ try:
323
+ import onnxsim
324
+ model_simplified, check = onnxsim.simplify(
325
+ model,
326
+ check_n=3,
327
+ perform_optimization=True,
328
+ skip_fuse_bn=False,
329
+ overwrite_input_shapes={input_tensor.name: new_shape},
330
+ )
331
+ if check:
332
+ orig_nodes = len(graph.node)
333
+ simp_nodes = len(model_simplified.graph.node)
334
+ model = model_simplified
335
+ print(f"✓ Model simplified ({orig_nodes} → {simp_nodes} nodes)")
336
+ else:
337
+ print("âš  Simplification validation failed, using non-simplified version")
338
+ except ImportError:
339
+ print("âš  onnx-simplifier not installed, skipping")
340
+ except Exception as e:
341
+ print(f"âš  Simplification failed: {e}, continuing without")
342
+
343
+ onnx.save(model, str(output_path))
344
+ output_size = output_path.stat().st_size
345
+ print(f"\n✓ Saved: {output_path} ({output_size / 1024 / 1024:.2f} MB)")
346
+
347
+ try:
348
+ verified = onnx.load(str(output_path))
349
+ onnx.checker.check_model(verified)
350
+ for inp in verified.graph.input:
351
+ if any(init.name == inp.name for init in verified.graph.initializer):
352
+ continue
353
+ shape = [dim.dim_value for dim in inp.type.tensor_type.shape.dim]
354
+ if all(isinstance(s, int) and s > 0 for s in shape):
355
+ print(f"✓ Input '{inp.name}': {shape}")
356
+ else:
357
+ print(f"âš  Input '{inp.name}' has dynamic dimensions: {shape}")
358
+ print("✨ Success! Fixed model ready for deployment")
359
+ return True
360
+ except Exception as e:
361
+ print(f"✗ Final verification failed: {e}")
362
+ return False
363
+
364
+
365
+ def _prepare_single_model(model_name, args, script_dir):
366
+ """Export backbone+head and fix shapes for one model. Returns True on success."""
367
+ info = SUPPORTED_MODELS[model_name]
368
+ final_output = script_dir / f"{model_name}.onnx"
369
+
370
+ print(f"\n{'=' * 80}")
371
+ print(f"DINO Model Preparation: {model_name}")
372
+ print(f"{'=' * 80}")
373
+ print(f"Architecture: {info['arch']} | Params: {info['params']}")
374
+ print(f"Top-1: {info['accuracy_top1']}% | k-NN: {info['knn_top1']}%")
375
+ print(f"Input shape: [{args.batch_size}, {args.channels}, {args.height}, {args.width}]")
376
+
377
+ if args.skip_export:
378
+ if not final_output.exists():
379
+ print(f"\n✗ Error: ONNX file not found: {final_output}")
380
+ print(" Run without --skip-export to export it first.")
381
+ return False
382
+ print(f"\nUsing existing ONNX file: {final_output.name}")
383
+ else:
384
+ if final_output.exists() and not args.force_export:
385
+ print(f"\nONNX file already exists: {final_output.name}")
386
+ print(f"File size: {final_output.stat().st_size / 1024 / 1024:.2f} MB")
387
+ print("Use --force-export to re-export.")
388
+ return True
389
+
390
+ success = export_to_onnx(model_name, final_output, args.height, args.width)
391
+ if not success:
392
+ return False
393
+
394
+ success = fix_model_shape(
395
+ final_output,
396
+ final_output,
397
+ batch_size=args.batch_size,
398
+ channels=args.channels,
399
+ height=args.height,
400
+ width=args.width,
401
+ use_simplifier=not args.no_simplifier,
402
+ )
403
+
404
+ if success:
405
+ print(f"\n{'=' * 80}")
406
+ print("COMPLETE!")
407
+ print(f"{'=' * 80}")
408
+ print(f"Model: {model_name}")
409
+ print(f"Output: {final_output.name} ({final_output.stat().st_size / 1024 / 1024:.2f} MB)")
410
+ print(f"Config: {model_name}_config.yaml")
411
+ else:
412
+ print("\n✗ Shape fixing failed")
413
+
414
+ return success
415
+
416
+
417
+ def main():
418
+ parser = argparse.ArgumentParser(
419
+ description=(
420
+ 'Export DINO classification models (backbone + linear head) from PyTorch Hub\n'
421
+ 'to ONNX format with fixed static shapes for TI EdgeAI hardware deployment.\n'
422
+ '\n'
423
+ 'Each model outputs 1000-class ImageNet logits [1, 1000].\n'
424
+ '\n'
425
+ 'Processing pipeline:\n'
426
+ ' 1. Load pretrained backbone from torch.hub (facebookresearch/dino:main)\n'
427
+ ' 2. Download pretrained linear classification weights from Meta AI\n'
428
+ ' 3. Combine backbone + linear head into a single module\n'
429
+ ' 4. Export to ONNX (opset 17) with dynamic batch axis\n'
430
+ ' 5. Fix dynamic input shapes to static [batch, channels, height, width]\n'
431
+ ' 6. Run ONNX shape inference and onnxsim simplification\n'
432
+ ' 7. Validate the final model'
433
+ ),
434
+ formatter_class=argparse.RawDescriptionHelpFormatter,
435
+ epilog="""
436
+ Examples:
437
+ # Export default model (ViT-S/16)
438
+ %(prog)s
439
+
440
+ # Export a specific variant
441
+ %(prog)s --model dino_vitb16
442
+
443
+ # Export all supported models in sequence
444
+ %(prog)s --model all
445
+
446
+ # Skip onnxsim (faster, larger output file)
447
+ %(prog)s --model dino_vits16 --no-simplifier
448
+
449
+ # Re-run shape inference + onnxsim on an already-exported ONNX file
450
+ %(prog)s --model dino_vits16 --skip-export
451
+
452
+ # Force re-export even if ONNX file exists
453
+ %(prog)s --model dino_vits16 --force-export
454
+
455
+ Available models:
456
+ dino_vits16 - ViT-S/16, 21M params, 77.0%% linear top-1 (recommended for edge)
457
+ dino_vits8 - ViT-S/8, 21M params, 79.7%% linear top-1
458
+ dino_vitb16 - ViT-B/16, 85M params, 78.2%% linear top-1
459
+ dino_vitb8 - ViT-B/8, 85M params, 80.1%% linear top-1
460
+ dino_resnet50 - ResNet-50, 23M params, 75.3%% linear top-1
461
+ """
462
+ )
463
+
464
+ parser.add_argument(
465
+ '--model', type=str, default='dino_vits16',
466
+ choices=list(SUPPORTED_MODELS.keys()) + ['all'],
467
+ help='Model variant to export, or "all" to export every model (default: dino_vits16)',
468
+ )
469
+ parser.add_argument('--batch-size', type=int, default=1,
470
+ help='Fixed batch size (default: 1)')
471
+ parser.add_argument('--channels', type=int, default=3,
472
+ help='Number of channels (default: 3)')
473
+ parser.add_argument('--height', type=int, default=224,
474
+ help='Image height (default: 224)')
475
+ parser.add_argument('--width', type=int, default=224,
476
+ help='Image width (default: 224)')
477
+ parser.add_argument('--force-export', action='store_true',
478
+ help='Force re-export even if ONNX file already exists')
479
+ parser.add_argument('--skip-export', action='store_true',
480
+ help='Skip export, only re-run shape inference + onnxsim on existing ONNX')
481
+ parser.add_argument('--no-simplifier', action='store_true',
482
+ help='Skip onnx-simplifier (onnxsim) step; shape inference still runs')
483
+
484
+ args = parser.parse_args()
485
+ script_dir = Path(__file__).parent
486
+
487
+ if args.model == 'all':
488
+ models = list(SUPPORTED_MODELS.keys())
489
+ print(f"Exporting {len(models)} DINO models...")
490
+ results = {}
491
+ for model_name in models:
492
+ results[model_name] = _prepare_single_model(model_name, args, script_dir)
493
+
494
+ print(f"\n{'=' * 80}")
495
+ print("ALL MODELS SUMMARY")
496
+ print(f"{'=' * 80}")
497
+ succeeded = [m for m, ok in results.items() if ok]
498
+ failed = [m for m, ok in results.items() if not ok]
499
+ for m in succeeded:
500
+ print(f" ✓ {m}")
501
+ for m in failed:
502
+ print(f" ✗ {m}")
503
+ print(f"\n{len(succeeded)}/{len(models)} models completed successfully.")
504
+ sys.exit(0 if not failed else 1)
505
+ else:
506
+ ok = _prepare_single_model(args.model, args, script_dir)
507
+ sys.exit(0 if ok else 1)
508
+
509
+
510
+ if __name__ == '__main__':
511
+ main()