GitHub Actions commited on
Commit
5624679
·
1 Parent(s): 703f4fd

Sync from GitHub Actions

Browse files
.gitattributes CHANGED
@@ -1,5 +1 @@
1
- notebooks/** linguist-documentation
2
- notebooks/*.ipynb linguist-documentation
3
- notebooks/**/*.ipynb linguist-documentation
4
- notebooks/** linguist-vendored
5
  *.ipynb linguist-documentation
 
 
 
 
 
1
  *.ipynb linguist-documentation
notebooks/autocatalog-model-comparsion.ipynb ADDED
@@ -0,0 +1,2557 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "cells": [
3
+ {
4
+ "cell_type": "markdown",
5
+ "id": "6c8e59c3",
6
+ "metadata": {},
7
+ "source": [
8
+ "### AutoCatalogAI — Reproducible Model Comparison\n",
9
+ "\n",
10
+ "This notebook creates a fair comparison between:\n",
11
+ "\n",
12
+ "1. **Majority Baseline**\n",
13
+ "2. **Frozen CLIP + Multi-task Heads (V1)**\n",
14
+ "3. **AutoCatalogAI V2 (fine-tuned production model)**"
15
+ ]
16
+ },
17
+ {
18
+ "cell_type": "code",
19
+ "execution_count": null,
20
+ "id": "4ce28167",
21
+ "metadata": {
22
+ "trusted": true
23
+ },
24
+ "outputs": [],
25
+ "source": [
26
+ "%pip uninstall -y torch torchvision torchaudio\n",
27
+ "\n",
28
+ "%pip install -q --no-cache-dir \\\n",
29
+ " torch==2.5.1 \\\n",
30
+ " torchvision==0.20.1 \\\n",
31
+ " torchaudio==2.5.1 \\\n",
32
+ " --index-url https://download.pytorch.org/whl/cu121\n",
33
+ "\n",
34
+ "%pip install -q \\\n",
35
+ " transformers==4.46.3 \\\n",
36
+ " datasets==3.1.0 \\\n",
37
+ " huggingface-hub==0.26.2 \\\n",
38
+ " scikit-learn==1.5.2 \\\n",
39
+ " pandas==2.2.3 \\\n",
40
+ " numpy==1.26.4 \\\n",
41
+ " Pillow==11.0.0 \\\n",
42
+ " tqdm==4.67.1"
43
+ ]
44
+ },
45
+ {
46
+ "cell_type": "markdown",
47
+ "id": "5aade65d",
48
+ "metadata": {},
49
+ "source": [
50
+ "## 2. Imports and Configuration"
51
+ ]
52
+ },
53
+ {
54
+ "cell_type": "code",
55
+ "execution_count": 1,
56
+ "id": "6f8acb93",
57
+ "metadata": {
58
+ "execution": {
59
+ "iopub.execute_input": "2026-07-06T14:50:58.717046Z",
60
+ "iopub.status.busy": "2026-07-06T14:50:58.716189Z",
61
+ "iopub.status.idle": "2026-07-06T14:51:09.100500Z",
62
+ "shell.execute_reply": "2026-07-06T14:51:09.099494Z",
63
+ "shell.execute_reply.started": "2026-07-06T14:50:58.717016Z"
64
+ },
65
+ "trusted": true
66
+ },
67
+ "outputs": [
68
+ {
69
+ "name": "stdout",
70
+ "output_type": "stream",
71
+ "text": [
72
+ "Python: 3.12.13\n",
73
+ "Torch: 2.5.1+cu121\n",
74
+ "Transformers: 4.46.3\n",
75
+ "CUDA available: True\n",
76
+ "GPU: Tesla P100-PCIE-16GB\n",
77
+ "CUDA runtime: 12.1\n"
78
+ ]
79
+ }
80
+ ],
81
+ "source": [
82
+ "import gc\n",
83
+ "import hashlib\n",
84
+ "import json\n",
85
+ "import os\n",
86
+ "import platform\n",
87
+ "import random\n",
88
+ "import time\n",
89
+ "from datetime import datetime, timezone\n",
90
+ "from pathlib import Path\n",
91
+ "\n",
92
+ "import numpy as np\n",
93
+ "import pandas as pd\n",
94
+ "import torch\n",
95
+ "import torch.nn as nn\n",
96
+ "import torch.nn.functional as F\n",
97
+ "import transformers\n",
98
+ "from datasets import load_dataset\n",
99
+ "from huggingface_hub import hf_hub_download\n",
100
+ "from PIL import Image\n",
101
+ "from sklearn.metrics import accuracy_score, f1_score\n",
102
+ "from sklearn.model_selection import train_test_split\n",
103
+ "from torch.utils.data import DataLoader, Dataset\n",
104
+ "from tqdm.auto import tqdm\n",
105
+ "from transformers import CLIPImageProcessor, CLIPModel\n",
106
+ "\n",
107
+ "print(\"Python:\", platform.python_version())\n",
108
+ "print(\"Torch:\", torch.__version__)\n",
109
+ "print(\"Transformers:\", transformers.__version__)\n",
110
+ "print(\"CUDA available:\", torch.cuda.is_available())\n",
111
+ "\n",
112
+ "if torch.cuda.is_available():\n",
113
+ " print(\"GPU:\", torch.cuda.get_device_name(0))\n",
114
+ " print(\"CUDA runtime:\", torch.version.cuda)\n"
115
+ ]
116
+ },
117
+ {
118
+ "cell_type": "code",
119
+ "execution_count": null,
120
+ "id": "bb00e756",
121
+ "metadata": {
122
+ "execution": {
123
+ "iopub.execute_input": "2026-07-06T14:51:12.581470Z",
124
+ "iopub.status.busy": "2026-07-06T14:51:12.580398Z",
125
+ "iopub.status.idle": "2026-07-06T14:51:12.589409Z",
126
+ "shell.execute_reply": "2026-07-06T14:51:12.588347Z",
127
+ "shell.execute_reply.started": "2026-07-06T14:51:12.581433Z"
128
+ },
129
+ "trusted": true
130
+ },
131
+ "outputs": [],
132
+ "source": [
133
+ "DATASET_NAME = \"ashraq/fashion-product-images-small\"\n",
134
+ "V1_REPO_ID = \"mohsin416/autocatalogai-clip-multitask\"\n",
135
+ "V2_REPO_ID = \"mohsin416/autocatalogai-clip-multitask-v2\"\n",
136
+ "\n",
137
+ "TASKS = [\n",
138
+ " \"gender\",\n",
139
+ " \"masterCategory\",\n",
140
+ " \"subCategory\",\n",
141
+ " \"articleType\",\n",
142
+ " \"baseColour\",\n",
143
+ " \"season\",\n",
144
+ " \"usage\",\n",
145
+ "]\n",
146
+ "\n",
147
+ "SEED = 42\n",
148
+ "TRAIN_RATIO = 0.70\n",
149
+ "VAL_RATIO = 0.15\n",
150
+ "TEST_RATIO = 0.15\n",
151
+ "\n",
152
+ "BATCH_SIZE = 64\n",
153
+ "NUM_WORKERS = 0\n",
154
+ "\n",
155
+ "LATENCY_WARMUP_RUNS = 20\n",
156
+ "LATENCY_MEASURED_RUNS = 100\n",
157
+ "\n",
158
+ "COLOR_IMAGE_SIZE = 128\n",
159
+ "COLOR_FEATURE_DIM = 37\n",
160
+ "\n",
161
+ "DEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n",
162
+ "ROOT_DIR = Path(\".\")\n",
163
+ "OUTPUT_DIR = ROOT_DIR / \"artifacts\" / \"evaluation\" / \"model_comparison\"\n",
164
+ "PROCESSED_DIR = ROOT_DIR / \"data\" / \"processed\"\n",
165
+ "\n",
166
+ "OUTPUT_DIR.mkdir(parents=True, exist_ok=True)\n",
167
+ "PROCESSED_DIR.mkdir(parents=True, exist_ok=True)"
168
+ ]
169
+ },
170
+ {
171
+ "cell_type": "markdown",
172
+ "id": "b03f3ae6",
173
+ "metadata": {},
174
+ "source": [
175
+ "## 3. Reproducibility"
176
+ ]
177
+ },
178
+ {
179
+ "cell_type": "code",
180
+ "execution_count": 3,
181
+ "id": "5027507b",
182
+ "metadata": {
183
+ "execution": {
184
+ "iopub.execute_input": "2026-07-06T14:51:14.492941Z",
185
+ "iopub.status.busy": "2026-07-06T14:51:14.492441Z",
186
+ "iopub.status.idle": "2026-07-06T14:51:14.500479Z",
187
+ "shell.execute_reply": "2026-07-06T14:51:14.499674Z",
188
+ "shell.execute_reply.started": "2026-07-06T14:51:14.492906Z"
189
+ },
190
+ "trusted": true
191
+ },
192
+ "outputs": [],
193
+ "source": [
194
+ "def set_seed(seed):\n",
195
+ " random.seed(seed)\n",
196
+ " np.random.seed(seed)\n",
197
+ " torch.manual_seed(seed)\n",
198
+ "\n",
199
+ " if torch.cuda.is_available():\n",
200
+ " torch.cuda.manual_seed_all(seed)\n",
201
+ "\n",
202
+ " torch.backends.cudnn.benchmark = False\n",
203
+ " torch.backends.cudnn.deterministic = True\n",
204
+ "\n",
205
+ "\n",
206
+ "set_seed(SEED)\n"
207
+ ]
208
+ },
209
+ {
210
+ "cell_type": "markdown",
211
+ "id": "08513a1e",
212
+ "metadata": {},
213
+ "source": [
214
+ "## 4. Download Published Model Metadata\n",
215
+ "\n",
216
+ "The comparison uses the exact published V1 and V2 checkpoints. \n",
217
+ "V1 represents the frozen-CLIP baseline. V2 represents the production model.\n"
218
+ ]
219
+ },
220
+ {
221
+ "cell_type": "code",
222
+ "execution_count": 4,
223
+ "id": "1ac434ae",
224
+ "metadata": {
225
+ "execution": {
226
+ "iopub.execute_input": "2026-07-06T14:51:15.077911Z",
227
+ "iopub.status.busy": "2026-07-06T14:51:15.076996Z",
228
+ "iopub.status.idle": "2026-07-06T14:51:16.662367Z",
229
+ "shell.execute_reply": "2026-07-06T14:51:16.661468Z",
230
+ "shell.execute_reply.started": "2026-07-06T14:51:15.077876Z"
231
+ },
232
+ "trusted": true
233
+ },
234
+ "outputs": [
235
+ {
236
+ "name": "stdout",
237
+ "output_type": "stream",
238
+ "text": [
239
+ "Base model: openai/clip-vit-base-patch32\n",
240
+ "Task classes: {'gender': 5, 'masterCategory': 7, 'subCategory': 45, 'articleType': 141, 'baseColour': 46, 'season': 4, 'usage': 8}\n"
241
+ ]
242
+ }
243
+ ],
244
+ "source": [
245
+ "def load_json(path):\n",
246
+ " with open(path, \"r\", encoding=\"utf-8\") as file:\n",
247
+ " return json.load(file)\n",
248
+ "\n",
249
+ "\n",
250
+ "def safe_torch_load(path, map_location=\"cpu\"):\n",
251
+ " try:\n",
252
+ " return torch.load(\n",
253
+ " path,\n",
254
+ " map_location=map_location,\n",
255
+ " weights_only=True,\n",
256
+ " )\n",
257
+ " except (TypeError, RuntimeError):\n",
258
+ " return torch.load(\n",
259
+ " path,\n",
260
+ " map_location=map_location,\n",
261
+ " )\n",
262
+ "\n",
263
+ "\n",
264
+ "def download_repo_artifacts(repo_id, include_rules=False):\n",
265
+ " filenames = [\n",
266
+ " \"model.pt\",\n",
267
+ " \"config.json\",\n",
268
+ " \"label_maps.json\",\n",
269
+ " ]\n",
270
+ "\n",
271
+ " if include_rules:\n",
272
+ " filenames.append(\"consistency_rules.json\")\n",
273
+ "\n",
274
+ " paths = {\n",
275
+ " filename: hf_hub_download(\n",
276
+ " repo_id=repo_id,\n",
277
+ " filename=filename,\n",
278
+ " repo_type=\"model\",\n",
279
+ " )\n",
280
+ " for filename in filenames\n",
281
+ " }\n",
282
+ "\n",
283
+ " artifacts = {\n",
284
+ " \"checkpoint\": safe_torch_load(paths[\"model.pt\"]),\n",
285
+ " \"config\": load_json(paths[\"config.json\"]),\n",
286
+ " \"label_maps\": load_json(paths[\"label_maps.json\"]),\n",
287
+ " }\n",
288
+ "\n",
289
+ " if include_rules:\n",
290
+ " artifacts[\"consistency_rules\"] = load_json(\n",
291
+ " paths[\"consistency_rules.json\"]\n",
292
+ " )\n",
293
+ "\n",
294
+ " return artifacts\n",
295
+ "\n",
296
+ "\n",
297
+ "v1_artifacts = download_repo_artifacts(V1_REPO_ID)\n",
298
+ "v2_artifacts = download_repo_artifacts(\n",
299
+ " V2_REPO_ID,\n",
300
+ " include_rules=True,\n",
301
+ ")\n",
302
+ "\n",
303
+ "if v1_artifacts[\"label_maps\"] != v2_artifacts[\"label_maps\"]:\n",
304
+ " raise ValueError(\n",
305
+ " \"V1 and V2 label maps are different. \"\n",
306
+ " \"A direct comparison would not be valid.\"\n",
307
+ " )\n",
308
+ "\n",
309
+ "label_maps = v2_artifacts[\"label_maps\"]\n",
310
+ "v1_config = v1_artifacts[\"config\"]\n",
311
+ "v2_config = v2_artifacts[\"config\"]\n",
312
+ "\n",
313
+ "MODEL_NAME = (\n",
314
+ " v2_config.get(\"base_model_name\")\n",
315
+ " or v2_config.get(\"model_name\")\n",
316
+ " or \"openai/clip-vit-base-patch32\"\n",
317
+ ")\n",
318
+ "\n",
319
+ "HIDDEN_DIM = int(v2_config.get(\"hidden_dim\", 512))\n",
320
+ "DROPOUT = float(v2_config.get(\"dropout\", 0.2))\n",
321
+ "\n",
322
+ "task_num_classes = {\n",
323
+ " task: len(label_maps[task][\"label2id\"])\n",
324
+ " for task in TASKS\n",
325
+ "}\n",
326
+ "\n",
327
+ "print(\"Base model:\", MODEL_NAME)\n",
328
+ "print(\"Task classes:\", task_num_classes)\n"
329
+ ]
330
+ },
331
+ {
332
+ "cell_type": "markdown",
333
+ "id": "279e5552",
334
+ "metadata": {},
335
+ "source": [
336
+ "## 5. Load and Clean the Dataset\n",
337
+ "\n",
338
+ "The same validation rules used in the training notebook are applied here.\n"
339
+ ]
340
+ },
341
+ {
342
+ "cell_type": "code",
343
+ "execution_count": 5,
344
+ "id": "bf1f57de",
345
+ "metadata": {
346
+ "execution": {
347
+ "iopub.execute_input": "2026-07-06T14:51:16.664194Z",
348
+ "iopub.status.busy": "2026-07-06T14:51:16.663758Z",
349
+ "iopub.status.idle": "2026-07-06T14:51:31.386130Z",
350
+ "shell.execute_reply": "2026-07-06T14:51:31.385294Z",
351
+ "shell.execute_reply.started": "2026-07-06T14:51:16.664164Z"
352
+ },
353
+ "trusted": true
354
+ },
355
+ "outputs": [
356
+ {
357
+ "data": {
358
+ "application/vnd.jupyter.widget-view+json": {
359
+ "model_id": "daa46350514b4fd99e16003161697efd",
360
+ "version_major": 2,
361
+ "version_minor": 0
362
+ },
363
+ "text/plain": [
364
+ "Filter: 0%| | 0/44072 [00:00<?, ? examples/s]"
365
+ ]
366
+ },
367
+ "metadata": {},
368
+ "output_type": "display_data"
369
+ },
370
+ {
371
+ "name": "stdout",
372
+ "output_type": "stream",
373
+ "text": [
374
+ "Raw samples: 44072\n",
375
+ "Clean samples: 44072\n",
376
+ "Dataset fingerprint: 5cacd020bfdb9ce5\n"
377
+ ]
378
+ }
379
+ ],
380
+ "source": [
381
+ "raw_dataset = load_dataset(\n",
382
+ " DATASET_NAME,\n",
383
+ " split=\"train\",\n",
384
+ ")\n",
385
+ "\n",
386
+ "missing_columns = [\n",
387
+ " task\n",
388
+ " for task in TASKS\n",
389
+ " if task not in raw_dataset.column_names\n",
390
+ "]\n",
391
+ "\n",
392
+ "if \"image\" not in raw_dataset.column_names:\n",
393
+ " raise ValueError(\"Dataset must contain an image column.\")\n",
394
+ "\n",
395
+ "if missing_columns:\n",
396
+ " raise ValueError(\n",
397
+ " f\"Dataset is missing task columns: {missing_columns}\"\n",
398
+ " )\n",
399
+ "\n",
400
+ "\n",
401
+ "def is_valid_row(row):\n",
402
+ " if row.get(\"image\") is None:\n",
403
+ " return False\n",
404
+ "\n",
405
+ " for task in TASKS:\n",
406
+ " value = row.get(task)\n",
407
+ "\n",
408
+ " if value is None:\n",
409
+ " return False\n",
410
+ "\n",
411
+ " value = str(value).strip()\n",
412
+ "\n",
413
+ " if not value:\n",
414
+ " return False\n",
415
+ "\n",
416
+ " if value not in label_maps[task][\"label2id\"]:\n",
417
+ " return False\n",
418
+ "\n",
419
+ " return True\n",
420
+ "\n",
421
+ "\n",
422
+ "clean_dataset = raw_dataset.filter(is_valid_row)\n",
423
+ "\n",
424
+ "print(\"Raw samples:\", len(raw_dataset))\n",
425
+ "print(\"Clean samples:\", len(clean_dataset))\n",
426
+ "print(\"Dataset fingerprint:\", clean_dataset._fingerprint)\n"
427
+ ]
428
+ },
429
+ {
430
+ "cell_type": "markdown",
431
+ "id": "c683a237",
432
+ "metadata": {},
433
+ "source": [
434
+ "## 6. Recreate the Exact 70/15/15 Split\n",
435
+ "\n",
436
+ "The split uses seed `42` and article-type stratification, matching the V2 training notebook.\n"
437
+ ]
438
+ },
439
+ {
440
+ "cell_type": "code",
441
+ "execution_count": 6,
442
+ "id": "2f6aa87a",
443
+ "metadata": {
444
+ "execution": {
445
+ "iopub.execute_input": "2026-07-06T14:51:31.388042Z",
446
+ "iopub.status.busy": "2026-07-06T14:51:31.387795Z",
447
+ "iopub.status.idle": "2026-07-06T14:51:33.455071Z",
448
+ "shell.execute_reply": "2026-07-06T14:51:33.454144Z",
449
+ "shell.execute_reply.started": "2026-07-06T14:51:31.388018Z"
450
+ },
451
+ "trusted": true
452
+ },
453
+ "outputs": [
454
+ {
455
+ "name": "stdout",
456
+ "output_type": "stream",
457
+ "text": [
458
+ "Train: 30850\n",
459
+ "Validation: 6611\n",
460
+ "Test: 6611\n",
461
+ "Test split SHA256: 106737acc60a436248d35e8a375a014d8b6e2f26f74b147b6444540b7781c0ec\n"
462
+ ]
463
+ }
464
+ ],
465
+ "source": [
466
+ "metadata = {\n",
467
+ " task: [\n",
468
+ " str(value).strip()\n",
469
+ " for value in clean_dataset[task]\n",
470
+ " ]\n",
471
+ " for task in TASKS\n",
472
+ "}\n",
473
+ "\n",
474
+ "df = pd.DataFrame(metadata)\n",
475
+ "df[\"dataset_idx\"] = np.arange(len(clean_dataset))\n",
476
+ "\n",
477
+ "\n",
478
+ "def make_safe_stratify_labels(series):\n",
479
+ " counts = series.value_counts()\n",
480
+ "\n",
481
+ " return series.apply(\n",
482
+ " lambda value: (\n",
483
+ " value\n",
484
+ " if counts[value] >= 2\n",
485
+ " else \"__rare__\"\n",
486
+ " )\n",
487
+ " )\n",
488
+ "\n",
489
+ "\n",
490
+ "all_indices = df.index.to_numpy()\n",
491
+ "\n",
492
+ "try:\n",
493
+ " train_idx, temporary_idx = train_test_split(\n",
494
+ " all_indices,\n",
495
+ " test_size=VAL_RATIO + TEST_RATIO,\n",
496
+ " random_state=SEED,\n",
497
+ " stratify=make_safe_stratify_labels(\n",
498
+ " df[\"articleType\"]\n",
499
+ " ),\n",
500
+ " )\n",
501
+ "except ValueError:\n",
502
+ " train_idx, temporary_idx = train_test_split(\n",
503
+ " all_indices,\n",
504
+ " test_size=VAL_RATIO + TEST_RATIO,\n",
505
+ " random_state=SEED,\n",
506
+ " )\n",
507
+ "\n",
508
+ "temporary_df = df.loc[temporary_idx]\n",
509
+ "\n",
510
+ "try:\n",
511
+ " val_idx, test_idx = train_test_split(\n",
512
+ " temporary_idx,\n",
513
+ " test_size=TEST_RATIO / (VAL_RATIO + TEST_RATIO),\n",
514
+ " random_state=SEED,\n",
515
+ " stratify=make_safe_stratify_labels(\n",
516
+ " temporary_df[\"articleType\"]\n",
517
+ " ),\n",
518
+ " )\n",
519
+ "except ValueError:\n",
520
+ " val_idx, test_idx = train_test_split(\n",
521
+ " temporary_idx,\n",
522
+ " test_size=TEST_RATIO / (VAL_RATIO + TEST_RATIO),\n",
523
+ " random_state=SEED,\n",
524
+ " )\n",
525
+ "\n",
526
+ "train_df = df.loc[train_idx].copy()\n",
527
+ "val_df = df.loc[val_idx].copy()\n",
528
+ "test_df = df.loc[test_idx].copy()\n",
529
+ "\n",
530
+ "assert set(train_idx).isdisjoint(val_idx)\n",
531
+ "assert set(train_idx).isdisjoint(test_idx)\n",
532
+ "assert set(val_idx).isdisjoint(test_idx)\n",
533
+ "\n",
534
+ "train_df.to_csv(\n",
535
+ " PROCESSED_DIR / \"train_v2.csv\",\n",
536
+ " index=False,\n",
537
+ ")\n",
538
+ "val_df.to_csv(\n",
539
+ " PROCESSED_DIR / \"val_v2.csv\",\n",
540
+ " index=False,\n",
541
+ ")\n",
542
+ "test_df.to_csv(\n",
543
+ " PROCESSED_DIR / \"test_v2.csv\",\n",
544
+ " index=False,\n",
545
+ ")\n",
546
+ "\n",
547
+ "test_index_bytes = np.asarray(\n",
548
+ " sorted(test_idx),\n",
549
+ " dtype=np.int64,\n",
550
+ ").tobytes()\n",
551
+ "\n",
552
+ "test_split_sha256 = hashlib.sha256(\n",
553
+ " test_index_bytes\n",
554
+ ").hexdigest()\n",
555
+ "\n",
556
+ "print(\"Train:\", len(train_df))\n",
557
+ "print(\"Validation:\", len(val_df))\n",
558
+ "print(\"Test:\", len(test_df))\n",
559
+ "print(\"Test split SHA256:\", test_split_sha256)\n"
560
+ ]
561
+ },
562
+ {
563
+ "cell_type": "markdown",
564
+ "id": "3e08e2d0",
565
+ "metadata": {},
566
+ "source": [
567
+ "## 7. Test Dataset and DataLoader"
568
+ ]
569
+ },
570
+ {
571
+ "cell_type": "code",
572
+ "execution_count": 7,
573
+ "id": "ee239b3e",
574
+ "metadata": {
575
+ "execution": {
576
+ "iopub.execute_input": "2026-07-06T14:51:33.456387Z",
577
+ "iopub.status.busy": "2026-07-06T14:51:33.456134Z",
578
+ "iopub.status.idle": "2026-07-06T14:51:33.553429Z",
579
+ "shell.execute_reply": "2026-07-06T14:51:33.552537Z",
580
+ "shell.execute_reply.started": "2026-07-06T14:51:33.456361Z"
581
+ },
582
+ "trusted": true
583
+ },
584
+ "outputs": [
585
+ {
586
+ "name": "stdout",
587
+ "output_type": "stream",
588
+ "text": [
589
+ "Test samples: 6611\n"
590
+ ]
591
+ }
592
+ ],
593
+ "source": [
594
+ "processor = CLIPImageProcessor.from_pretrained(\n",
595
+ " MODEL_NAME\n",
596
+ ")\n",
597
+ "\n",
598
+ "\n",
599
+ "def extract_color_features(\n",
600
+ " image,\n",
601
+ " image_size=COLOR_IMAGE_SIZE,\n",
602
+ "):\n",
603
+ " image = image.convert(\"RGB\").resize(\n",
604
+ " (image_size, image_size)\n",
605
+ " )\n",
606
+ "\n",
607
+ " margin = int(image_size * 0.10)\n",
608
+ "\n",
609
+ " image = image.crop(\n",
610
+ " (\n",
611
+ " margin,\n",
612
+ " margin,\n",
613
+ " image_size - margin,\n",
614
+ " image_size - margin,\n",
615
+ " )\n",
616
+ " )\n",
617
+ "\n",
618
+ " rgb = np.asarray(\n",
619
+ " image,\n",
620
+ " dtype=np.float32,\n",
621
+ " ) / 255.0\n",
622
+ "\n",
623
+ " hsv = np.asarray(\n",
624
+ " image.convert(\"HSV\"),\n",
625
+ " dtype=np.float32,\n",
626
+ " ) / 255.0\n",
627
+ "\n",
628
+ " rgb_flat = rgb.reshape(-1, 3)\n",
629
+ " hsv_flat = hsv.reshape(-1, 3)\n",
630
+ "\n",
631
+ " saturation = hsv_flat[:, 1]\n",
632
+ " value = hsv_flat[:, 2]\n",
633
+ "\n",
634
+ " foreground_mask = (\n",
635
+ " (saturation > 0.08)\n",
636
+ " | (value < 0.92)\n",
637
+ " )\n",
638
+ "\n",
639
+ " if foreground_mask.sum() < 256:\n",
640
+ " foreground_mask = np.ones(\n",
641
+ " len(hsv_flat),\n",
642
+ " dtype=bool,\n",
643
+ " )\n",
644
+ "\n",
645
+ " selected_rgb = rgb_flat[foreground_mask]\n",
646
+ " selected_hsv = hsv_flat[foreground_mask]\n",
647
+ "\n",
648
+ " hue_hist, _ = np.histogram(\n",
649
+ " selected_hsv[:, 0],\n",
650
+ " bins=12,\n",
651
+ " range=(0.0, 1.0),\n",
652
+ " )\n",
653
+ "\n",
654
+ " saturation_hist, _ = np.histogram(\n",
655
+ " selected_hsv[:, 1],\n",
656
+ " bins=8,\n",
657
+ " range=(0.0, 1.0),\n",
658
+ " )\n",
659
+ "\n",
660
+ " value_hist, _ = np.histogram(\n",
661
+ " selected_hsv[:, 2],\n",
662
+ " bins=8,\n",
663
+ " range=(0.0, 1.0),\n",
664
+ " )\n",
665
+ "\n",
666
+ " hue_hist = hue_hist.astype(np.float32)\n",
667
+ " saturation_hist = saturation_hist.astype(np.float32)\n",
668
+ " value_hist = value_hist.astype(np.float32)\n",
669
+ "\n",
670
+ " hue_hist /= max(hue_hist.sum(), 1.0)\n",
671
+ " saturation_hist /= max(\n",
672
+ " saturation_hist.sum(),\n",
673
+ " 1.0,\n",
674
+ " )\n",
675
+ " value_hist /= max(value_hist.sum(), 1.0)\n",
676
+ "\n",
677
+ " rgb_mean = selected_rgb.mean(\n",
678
+ " axis=0\n",
679
+ " ).astype(np.float32)\n",
680
+ "\n",
681
+ " rgb_std = selected_rgb.std(\n",
682
+ " axis=0\n",
683
+ " ).astype(np.float32)\n",
684
+ "\n",
685
+ " rgb_median = np.median(\n",
686
+ " selected_rgb,\n",
687
+ " axis=0,\n",
688
+ " ).astype(np.float32)\n",
689
+ "\n",
690
+ " features = np.concatenate(\n",
691
+ " [\n",
692
+ " hue_hist,\n",
693
+ " saturation_hist,\n",
694
+ " value_hist,\n",
695
+ " rgb_mean,\n",
696
+ " rgb_std,\n",
697
+ " rgb_median,\n",
698
+ " ]\n",
699
+ " ).astype(np.float32)\n",
700
+ "\n",
701
+ " if features.shape[0] != COLOR_FEATURE_DIM:\n",
702
+ " raise ValueError(\n",
703
+ " f\"Expected {COLOR_FEATURE_DIM} color features, \"\n",
704
+ " f\"got {features.shape[0]}\"\n",
705
+ " )\n",
706
+ "\n",
707
+ " return features\n",
708
+ "\n",
709
+ "\n",
710
+ "class ComparisonDataset(Dataset):\n",
711
+ " def __init__(\n",
712
+ " self,\n",
713
+ " source_dataset,\n",
714
+ " indices,\n",
715
+ " processor,\n",
716
+ " label_maps,\n",
717
+ " ):\n",
718
+ " self.source_dataset = source_dataset\n",
719
+ " self.indices = list(map(int, indices))\n",
720
+ " self.processor = processor\n",
721
+ " self.label_maps = label_maps\n",
722
+ "\n",
723
+ " def __len__(self):\n",
724
+ " return len(self.indices)\n",
725
+ "\n",
726
+ " def __getitem__(self, index):\n",
727
+ " global_index = self.indices[index]\n",
728
+ " item = self.source_dataset[global_index]\n",
729
+ "\n",
730
+ " image = item[\"image\"]\n",
731
+ "\n",
732
+ " if not isinstance(image, Image.Image):\n",
733
+ " image = Image.open(image)\n",
734
+ "\n",
735
+ " image = image.convert(\"RGB\")\n",
736
+ "\n",
737
+ " pixel_values = self.processor(\n",
738
+ " images=image,\n",
739
+ " return_tensors=\"pt\",\n",
740
+ " )[\"pixel_values\"].squeeze(0)\n",
741
+ "\n",
742
+ " color_features = torch.tensor(\n",
743
+ " extract_color_features(image),\n",
744
+ " dtype=torch.float32,\n",
745
+ " )\n",
746
+ "\n",
747
+ " labels = {\n",
748
+ " task: torch.tensor(\n",
749
+ " self.label_maps[task][\"label2id\"][\n",
750
+ " str(item[task]).strip()\n",
751
+ " ],\n",
752
+ " dtype=torch.long,\n",
753
+ " )\n",
754
+ " for task in TASKS\n",
755
+ " }\n",
756
+ "\n",
757
+ " return {\n",
758
+ " \"pixel_values\": pixel_values,\n",
759
+ " \"color_features\": color_features,\n",
760
+ " \"labels\": labels,\n",
761
+ " \"global_index\": global_index,\n",
762
+ " }\n",
763
+ "\n",
764
+ "\n",
765
+ "def collate_batch(batch):\n",
766
+ " return {\n",
767
+ " \"pixel_values\": torch.stack(\n",
768
+ " [item[\"pixel_values\"] for item in batch]\n",
769
+ " ),\n",
770
+ " \"color_features\": torch.stack(\n",
771
+ " [item[\"color_features\"] for item in batch]\n",
772
+ " ),\n",
773
+ " \"labels\": {\n",
774
+ " task: torch.stack(\n",
775
+ " [item[\"labels\"][task] for item in batch]\n",
776
+ " )\n",
777
+ " for task in TASKS\n",
778
+ " },\n",
779
+ " \"global_indices\": [\n",
780
+ " item[\"global_index\"]\n",
781
+ " for item in batch\n",
782
+ " ],\n",
783
+ " }\n",
784
+ "\n",
785
+ "\n",
786
+ "test_dataset = ComparisonDataset(\n",
787
+ " clean_dataset,\n",
788
+ " test_df[\"dataset_idx\"],\n",
789
+ " processor,\n",
790
+ " label_maps,\n",
791
+ ")\n",
792
+ "\n",
793
+ "test_loader = DataLoader(\n",
794
+ " test_dataset,\n",
795
+ " batch_size=BATCH_SIZE,\n",
796
+ " shuffle=False,\n",
797
+ " num_workers=NUM_WORKERS,\n",
798
+ " pin_memory=torch.cuda.is_available(),\n",
799
+ " collate_fn=collate_batch,\n",
800
+ ")\n",
801
+ "\n",
802
+ "print(\"Test samples:\", len(test_dataset))\n"
803
+ ]
804
+ },
805
+ {
806
+ "cell_type": "markdown",
807
+ "id": "c97a0b5b",
808
+ "metadata": {},
809
+ "source": [
810
+ "## 8. Model Architectures"
811
+ ]
812
+ },
813
+ {
814
+ "cell_type": "code",
815
+ "execution_count": 8,
816
+ "id": "65d75793",
817
+ "metadata": {
818
+ "execution": {
819
+ "iopub.execute_input": "2026-07-06T14:51:33.555302Z",
820
+ "iopub.status.busy": "2026-07-06T14:51:33.554960Z",
821
+ "iopub.status.idle": "2026-07-06T14:51:33.569958Z",
822
+ "shell.execute_reply": "2026-07-06T14:51:33.568888Z",
823
+ "shell.execute_reply.started": "2026-07-06T14:51:33.555271Z"
824
+ },
825
+ "trusted": true
826
+ },
827
+ "outputs": [],
828
+ "source": [
829
+ "class ClassificationHead(nn.Module):\n",
830
+ " def __init__(\n",
831
+ " self,\n",
832
+ " embedding_dim,\n",
833
+ " num_classes,\n",
834
+ " hidden_dim=512,\n",
835
+ " dropout=0.2,\n",
836
+ " ):\n",
837
+ " super().__init__()\n",
838
+ "\n",
839
+ " self.net = nn.Sequential(\n",
840
+ " nn.LayerNorm(embedding_dim),\n",
841
+ " nn.Linear(embedding_dim, hidden_dim),\n",
842
+ " nn.GELU(),\n",
843
+ " nn.Dropout(dropout),\n",
844
+ " nn.Linear(hidden_dim, num_classes),\n",
845
+ " )\n",
846
+ "\n",
847
+ " def forward(self, features):\n",
848
+ " return self.net(features)\n",
849
+ "\n",
850
+ "\n",
851
+ "class CLIPMultiTaskClassifier(nn.Module):\n",
852
+ " def __init__(\n",
853
+ " self,\n",
854
+ " model_name,\n",
855
+ " task_num_classes,\n",
856
+ " hidden_dim=512,\n",
857
+ " dropout=0.2,\n",
858
+ " ):\n",
859
+ " super().__init__()\n",
860
+ "\n",
861
+ " self.clip = CLIPModel.from_pretrained(\n",
862
+ " model_name\n",
863
+ " )\n",
864
+ "\n",
865
+ " embedding_dim = (\n",
866
+ " self.clip.config.projection_dim\n",
867
+ " )\n",
868
+ "\n",
869
+ " self.heads = nn.ModuleDict(\n",
870
+ " {\n",
871
+ " task: ClassificationHead(\n",
872
+ " embedding_dim,\n",
873
+ " num_classes,\n",
874
+ " hidden_dim,\n",
875
+ " dropout,\n",
876
+ " )\n",
877
+ " for task, num_classes\n",
878
+ " in task_num_classes.items()\n",
879
+ " }\n",
880
+ " )\n",
881
+ "\n",
882
+ " def forward(self, pixel_values):\n",
883
+ " image_features = (\n",
884
+ " self.clip.get_image_features(\n",
885
+ " pixel_values=pixel_values\n",
886
+ " )\n",
887
+ " )\n",
888
+ "\n",
889
+ " image_features = F.normalize(\n",
890
+ " image_features,\n",
891
+ " dim=-1,\n",
892
+ " )\n",
893
+ "\n",
894
+ " return {\n",
895
+ " task: head(image_features)\n",
896
+ " for task, head in self.heads.items()\n",
897
+ " }\n",
898
+ "\n",
899
+ "\n",
900
+ "class CLIPMultiTaskClassifierV2(nn.Module):\n",
901
+ " def __init__(\n",
902
+ " self,\n",
903
+ " model_name,\n",
904
+ " task_num_classes,\n",
905
+ " hidden_dim=512,\n",
906
+ " dropout=0.2,\n",
907
+ " color_feature_dim=37,\n",
908
+ " ):\n",
909
+ " super().__init__()\n",
910
+ "\n",
911
+ " self.clip = CLIPModel.from_pretrained(\n",
912
+ " model_name\n",
913
+ " )\n",
914
+ "\n",
915
+ " embedding_dim = (\n",
916
+ " self.clip.config.projection_dim\n",
917
+ " )\n",
918
+ "\n",
919
+ " self.heads = nn.ModuleDict(\n",
920
+ " {\n",
921
+ " task: ClassificationHead(\n",
922
+ " embedding_dim,\n",
923
+ " num_classes,\n",
924
+ " hidden_dim,\n",
925
+ " dropout,\n",
926
+ " )\n",
927
+ " for task, num_classes\n",
928
+ " in task_num_classes.items()\n",
929
+ " }\n",
930
+ " )\n",
931
+ "\n",
932
+ " self.master_to_sub = nn.Linear(\n",
933
+ " task_num_classes[\"masterCategory\"],\n",
934
+ " task_num_classes[\"subCategory\"],\n",
935
+ " bias=False,\n",
936
+ " )\n",
937
+ "\n",
938
+ " self.sub_to_article = nn.Linear(\n",
939
+ " task_num_classes[\"subCategory\"],\n",
940
+ " task_num_classes[\"articleType\"],\n",
941
+ " bias=False,\n",
942
+ " )\n",
943
+ "\n",
944
+ " self.article_to_season = nn.Linear(\n",
945
+ " task_num_classes[\"articleType\"],\n",
946
+ " task_num_classes[\"season\"],\n",
947
+ " bias=False,\n",
948
+ " )\n",
949
+ "\n",
950
+ " self.article_to_usage = nn.Linear(\n",
951
+ " task_num_classes[\"articleType\"],\n",
952
+ " task_num_classes[\"usage\"],\n",
953
+ " bias=False,\n",
954
+ " )\n",
955
+ "\n",
956
+ " self.color_branch = nn.Sequential(\n",
957
+ " nn.LayerNorm(color_feature_dim),\n",
958
+ " nn.Linear(color_feature_dim, 64),\n",
959
+ " nn.GELU(),\n",
960
+ " nn.Dropout(0.10),\n",
961
+ " nn.Linear(\n",
962
+ " 64,\n",
963
+ " task_num_classes[\"baseColour\"],\n",
964
+ " ),\n",
965
+ " )\n",
966
+ "\n",
967
+ " def forward(\n",
968
+ " self,\n",
969
+ " pixel_values,\n",
970
+ " color_features,\n",
971
+ " ):\n",
972
+ " image_features = (\n",
973
+ " self.clip.get_image_features(\n",
974
+ " pixel_values=pixel_values\n",
975
+ " )\n",
976
+ " )\n",
977
+ "\n",
978
+ " image_features = F.normalize(\n",
979
+ " image_features,\n",
980
+ " dim=-1,\n",
981
+ " )\n",
982
+ "\n",
983
+ " outputs = {\n",
984
+ " task: head(image_features)\n",
985
+ " for task, head in self.heads.items()\n",
986
+ " }\n",
987
+ "\n",
988
+ " master_probs = torch.softmax(\n",
989
+ " outputs[\"masterCategory\"].detach(),\n",
990
+ " dim=1,\n",
991
+ " )\n",
992
+ "\n",
993
+ " outputs[\"subCategory\"] = (\n",
994
+ " outputs[\"subCategory\"]\n",
995
+ " + self.master_to_sub(master_probs)\n",
996
+ " )\n",
997
+ "\n",
998
+ " sub_probs = torch.softmax(\n",
999
+ " outputs[\"subCategory\"].detach(),\n",
1000
+ " dim=1,\n",
1001
+ " )\n",
1002
+ "\n",
1003
+ " outputs[\"articleType\"] = (\n",
1004
+ " outputs[\"articleType\"]\n",
1005
+ " + self.sub_to_article(sub_probs)\n",
1006
+ " )\n",
1007
+ "\n",
1008
+ " article_probs = torch.softmax(\n",
1009
+ " outputs[\"articleType\"].detach(),\n",
1010
+ " dim=1,\n",
1011
+ " )\n",
1012
+ "\n",
1013
+ " outputs[\"season\"] = (\n",
1014
+ " outputs[\"season\"]\n",
1015
+ " + self.article_to_season(article_probs)\n",
1016
+ " )\n",
1017
+ "\n",
1018
+ " outputs[\"usage\"] = (\n",
1019
+ " outputs[\"usage\"]\n",
1020
+ " + self.article_to_usage(article_probs)\n",
1021
+ " )\n",
1022
+ "\n",
1023
+ " outputs[\"baseColour\"] = (\n",
1024
+ " outputs[\"baseColour\"]\n",
1025
+ " + self.color_branch(color_features)\n",
1026
+ " )\n",
1027
+ "\n",
1028
+ " return outputs\n"
1029
+ ]
1030
+ },
1031
+ {
1032
+ "cell_type": "markdown",
1033
+ "id": "6bcdc281",
1034
+ "metadata": {},
1035
+ "source": [
1036
+ "## 9. Shared Metric Functions"
1037
+ ]
1038
+ },
1039
+ {
1040
+ "cell_type": "code",
1041
+ "execution_count": 9,
1042
+ "id": "d76d6326",
1043
+ "metadata": {
1044
+ "execution": {
1045
+ "iopub.execute_input": "2026-07-06T14:51:33.571811Z",
1046
+ "iopub.status.busy": "2026-07-06T14:51:33.571394Z",
1047
+ "iopub.status.idle": "2026-07-06T14:51:33.589627Z",
1048
+ "shell.execute_reply": "2026-07-06T14:51:33.588833Z",
1049
+ "shell.execute_reply.started": "2026-07-06T14:51:33.571769Z"
1050
+ },
1051
+ "trusted": true
1052
+ },
1053
+ "outputs": [],
1054
+ "source": [
1055
+ "def evaluate_predictions(\n",
1056
+ " y_true,\n",
1057
+ " y_pred,\n",
1058
+ " y_top3,\n",
1059
+ "):\n",
1060
+ " task_metrics = {}\n",
1061
+ "\n",
1062
+ " for task in TASKS:\n",
1063
+ " top3_matches = [\n",
1064
+ " int(true_label)\n",
1065
+ " in set(map(int, top3_labels))\n",
1066
+ " for true_label, top3_labels\n",
1067
+ " in zip(\n",
1068
+ " y_true[task],\n",
1069
+ " y_top3[task],\n",
1070
+ " )\n",
1071
+ " ]\n",
1072
+ "\n",
1073
+ " task_metrics[task] = {\n",
1074
+ " \"accuracy\": float(\n",
1075
+ " accuracy_score(\n",
1076
+ " y_true[task],\n",
1077
+ " y_pred[task],\n",
1078
+ " )\n",
1079
+ " ),\n",
1080
+ " \"macro_f1\": float(\n",
1081
+ " f1_score(\n",
1082
+ " y_true[task],\n",
1083
+ " y_pred[task],\n",
1084
+ " average=\"macro\",\n",
1085
+ " zero_division=0,\n",
1086
+ " )\n",
1087
+ " ),\n",
1088
+ " \"top3_accuracy\": float(\n",
1089
+ " np.mean(top3_matches)\n",
1090
+ " ),\n",
1091
+ " }\n",
1092
+ "\n",
1093
+ " exact_matches = np.ones(\n",
1094
+ " len(y_true[TASKS[0]]),\n",
1095
+ " dtype=bool,\n",
1096
+ " )\n",
1097
+ "\n",
1098
+ " for task in TASKS:\n",
1099
+ " exact_matches &= (\n",
1100
+ " np.asarray(y_true[task])\n",
1101
+ " == np.asarray(y_pred[task])\n",
1102
+ " )\n",
1103
+ "\n",
1104
+ " overall = {\n",
1105
+ " \"average_accuracy\": float(\n",
1106
+ " np.mean(\n",
1107
+ " [\n",
1108
+ " task_metrics[task][\"accuracy\"]\n",
1109
+ " for task in TASKS\n",
1110
+ " ]\n",
1111
+ " )\n",
1112
+ " ),\n",
1113
+ " \"average_macro_f1\": float(\n",
1114
+ " np.mean(\n",
1115
+ " [\n",
1116
+ " task_metrics[task][\"macro_f1\"]\n",
1117
+ " for task in TASKS\n",
1118
+ " ]\n",
1119
+ " )\n",
1120
+ " ),\n",
1121
+ " \"average_top3_accuracy\": float(\n",
1122
+ " np.mean(\n",
1123
+ " [\n",
1124
+ " task_metrics[task][\"top3_accuracy\"]\n",
1125
+ " for task in TASKS\n",
1126
+ " ]\n",
1127
+ " )\n",
1128
+ " ),\n",
1129
+ " \"exact_match_accuracy\": float(\n",
1130
+ " exact_matches.mean()\n",
1131
+ " ),\n",
1132
+ " \"samples\": int(len(exact_matches)),\n",
1133
+ " }\n",
1134
+ "\n",
1135
+ " return {\n",
1136
+ " \"task_metrics\": task_metrics,\n",
1137
+ " \"overall_metrics\": overall,\n",
1138
+ " }\n",
1139
+ "\n",
1140
+ "\n",
1141
+ "@torch.inference_mode()\n",
1142
+ "def collect_model_predictions(\n",
1143
+ " model,\n",
1144
+ " loader,\n",
1145
+ " device,\n",
1146
+ " uses_color_features,\n",
1147
+ "):\n",
1148
+ " model.eval()\n",
1149
+ "\n",
1150
+ " y_true = {\n",
1151
+ " task: []\n",
1152
+ " for task in TASKS\n",
1153
+ " }\n",
1154
+ "\n",
1155
+ " y_pred = {\n",
1156
+ " task: []\n",
1157
+ " for task in TASKS\n",
1158
+ " }\n",
1159
+ "\n",
1160
+ " y_top3 = {\n",
1161
+ " task: []\n",
1162
+ " for task in TASKS\n",
1163
+ " }\n",
1164
+ "\n",
1165
+ " global_indices = []\n",
1166
+ "\n",
1167
+ " for batch in tqdm(\n",
1168
+ " loader,\n",
1169
+ " desc=\"Evaluating\",\n",
1170
+ " leave=False,\n",
1171
+ " ):\n",
1172
+ " pixel_values = batch[\n",
1173
+ " \"pixel_values\"\n",
1174
+ " ].to(device)\n",
1175
+ "\n",
1176
+ " if uses_color_features:\n",
1177
+ " outputs = model(\n",
1178
+ " pixel_values,\n",
1179
+ " batch[\"color_features\"].to(device),\n",
1180
+ " )\n",
1181
+ " else:\n",
1182
+ " outputs = model(pixel_values)\n",
1183
+ "\n",
1184
+ " for task in TASKS:\n",
1185
+ " probabilities = torch.softmax(\n",
1186
+ " outputs[task],\n",
1187
+ " dim=1,\n",
1188
+ " )\n",
1189
+ "\n",
1190
+ " top_k = min(\n",
1191
+ " 3,\n",
1192
+ " probabilities.shape[1],\n",
1193
+ " )\n",
1194
+ "\n",
1195
+ " top_values, top_indices = torch.topk(\n",
1196
+ " probabilities,\n",
1197
+ " k=top_k,\n",
1198
+ " dim=1,\n",
1199
+ " )\n",
1200
+ "\n",
1201
+ " predictions = top_indices[:, 0]\n",
1202
+ "\n",
1203
+ " y_true[task].extend(\n",
1204
+ " batch[\"labels\"][task]\n",
1205
+ " .numpy()\n",
1206
+ " .tolist()\n",
1207
+ " )\n",
1208
+ "\n",
1209
+ " y_pred[task].extend(\n",
1210
+ " predictions\n",
1211
+ " .cpu()\n",
1212
+ " .numpy()\n",
1213
+ " .tolist()\n",
1214
+ " )\n",
1215
+ "\n",
1216
+ " y_top3[task].extend(\n",
1217
+ " top_indices\n",
1218
+ " .cpu()\n",
1219
+ " .numpy()\n",
1220
+ " .tolist()\n",
1221
+ " )\n",
1222
+ "\n",
1223
+ " global_indices.extend(\n",
1224
+ " batch[\"global_indices\"]\n",
1225
+ " )\n",
1226
+ "\n",
1227
+ " for task in TASKS:\n",
1228
+ " y_true[task] = np.asarray(\n",
1229
+ " y_true[task],\n",
1230
+ " dtype=np.int64,\n",
1231
+ " )\n",
1232
+ "\n",
1233
+ " y_pred[task] = np.asarray(\n",
1234
+ " y_pred[task],\n",
1235
+ " dtype=np.int64,\n",
1236
+ " )\n",
1237
+ "\n",
1238
+ " y_top3[task] = np.asarray(\n",
1239
+ " y_top3[task],\n",
1240
+ " dtype=np.int64,\n",
1241
+ " )\n",
1242
+ "\n",
1243
+ " return (\n",
1244
+ " y_true,\n",
1245
+ " y_pred,\n",
1246
+ " y_top3,\n",
1247
+ " global_indices,\n",
1248
+ " )\n",
1249
+ "\n",
1250
+ "\n",
1251
+ "def get_state_dict(checkpoint):\n",
1252
+ " if \"model_state_dict\" in checkpoint:\n",
1253
+ " return checkpoint[\"model_state_dict\"]\n",
1254
+ "\n",
1255
+ " if \"state_dict\" in checkpoint:\n",
1256
+ " return checkpoint[\"state_dict\"]\n",
1257
+ "\n",
1258
+ " return checkpoint\n"
1259
+ ]
1260
+ },
1261
+ {
1262
+ "cell_type": "markdown",
1263
+ "id": "07ff619d",
1264
+ "metadata": {},
1265
+ "source": [
1266
+ "## 10. Majority Baseline\n",
1267
+ "\n",
1268
+ "Top-1 uses the most frequent training label for each task. \n",
1269
+ "Top-3 uses the three most frequent training labels for each task.\n"
1270
+ ]
1271
+ },
1272
+ {
1273
+ "cell_type": "code",
1274
+ "execution_count": 10,
1275
+ "id": "3f5ca12c",
1276
+ "metadata": {
1277
+ "execution": {
1278
+ "iopub.execute_input": "2026-07-06T14:51:39.992449Z",
1279
+ "iopub.status.busy": "2026-07-06T14:51:39.991906Z",
1280
+ "iopub.status.idle": "2026-07-06T14:51:40.154821Z",
1281
+ "shell.execute_reply": "2026-07-06T14:51:40.154030Z",
1282
+ "shell.execute_reply.started": "2026-07-06T14:51:39.992414Z"
1283
+ },
1284
+ "trusted": true
1285
+ },
1286
+ "outputs": [
1287
+ {
1288
+ "name": "stdout",
1289
+ "output_type": "stream",
1290
+ "text": [
1291
+ "{\n",
1292
+ " \"average_accuracy\": 0.42628087386822827,\n",
1293
+ " \"average_macro_f1\": 0.07927692465823664,\n",
1294
+ " \"average_top3_accuracy\": 0.7344469174752037,\n",
1295
+ " \"exact_match_accuracy\": 0.009529571925578581,\n",
1296
+ " \"samples\": 6611\n",
1297
+ "}\n"
1298
+ ]
1299
+ }
1300
+ ],
1301
+ "source": [
1302
+ "majority_top1 = {}\n",
1303
+ "majority_top3 = {}\n",
1304
+ "\n",
1305
+ "for task in TASKS:\n",
1306
+ " counts = train_df[task].value_counts()\n",
1307
+ "\n",
1308
+ " majority_top1[task] = (\n",
1309
+ " label_maps[task][\"label2id\"][\n",
1310
+ " counts.index[0]\n",
1311
+ " ]\n",
1312
+ " )\n",
1313
+ "\n",
1314
+ " majority_top3[task] = [\n",
1315
+ " label_maps[task][\"label2id\"][label]\n",
1316
+ " for label in counts.index[:3]\n",
1317
+ " ]\n",
1318
+ "\n",
1319
+ "\n",
1320
+ "majority_y_true = {\n",
1321
+ " task: test_df[task]\n",
1322
+ " .map(label_maps[task][\"label2id\"])\n",
1323
+ " .to_numpy(dtype=np.int64)\n",
1324
+ " for task in TASKS\n",
1325
+ "}\n",
1326
+ "\n",
1327
+ "majority_y_pred = {\n",
1328
+ " task: np.full(\n",
1329
+ " len(test_df),\n",
1330
+ " majority_top1[task],\n",
1331
+ " dtype=np.int64,\n",
1332
+ " )\n",
1333
+ " for task in TASKS\n",
1334
+ "}\n",
1335
+ "\n",
1336
+ "majority_y_top3 = {\n",
1337
+ " task: np.tile(\n",
1338
+ " np.asarray(\n",
1339
+ " majority_top3[task],\n",
1340
+ " dtype=np.int64,\n",
1341
+ " ),\n",
1342
+ " (len(test_df), 1),\n",
1343
+ " )\n",
1344
+ " for task in TASKS\n",
1345
+ "}\n",
1346
+ "\n",
1347
+ "majority_metrics = evaluate_predictions(\n",
1348
+ " majority_y_true,\n",
1349
+ " majority_y_pred,\n",
1350
+ " majority_y_top3,\n",
1351
+ ")\n",
1352
+ "\n",
1353
+ "print(\n",
1354
+ " json.dumps(\n",
1355
+ " majority_metrics[\"overall_metrics\"],\n",
1356
+ " indent=2,\n",
1357
+ " )\n",
1358
+ ")\n"
1359
+ ]
1360
+ },
1361
+ {
1362
+ "cell_type": "code",
1363
+ "execution_count": 11,
1364
+ "id": "74860669",
1365
+ "metadata": {
1366
+ "execution": {
1367
+ "iopub.execute_input": "2026-07-06T14:51:42.144938Z",
1368
+ "iopub.status.busy": "2026-07-06T14:51:42.144442Z",
1369
+ "iopub.status.idle": "2026-07-06T14:51:42.156211Z",
1370
+ "shell.execute_reply": "2026-07-06T14:51:42.155268Z",
1371
+ "shell.execute_reply.started": "2026-07-06T14:51:42.144902Z"
1372
+ },
1373
+ "trusted": true
1374
+ },
1375
+ "outputs": [
1376
+ {
1377
+ "name": "stdout",
1378
+ "output_type": "stream",
1379
+ "text": [
1380
+ "Majority lookup latency: 0.000565 ms\n"
1381
+ ]
1382
+ }
1383
+ ],
1384
+ "source": [
1385
+ "def benchmark_majority_lookup(\n",
1386
+ " measured_runs=10000,\n",
1387
+ "):\n",
1388
+ " start = time.perf_counter()\n",
1389
+ "\n",
1390
+ " for _ in range(measured_runs):\n",
1391
+ " _ = {\n",
1392
+ " task: majority_top1[task]\n",
1393
+ " for task in TASKS\n",
1394
+ " }\n",
1395
+ "\n",
1396
+ " elapsed_ms = (\n",
1397
+ " time.perf_counter() - start\n",
1398
+ " ) * 1000.0\n",
1399
+ "\n",
1400
+ " return elapsed_ms / measured_runs\n",
1401
+ "\n",
1402
+ "\n",
1403
+ "majority_latency_ms = benchmark_majority_lookup()\n",
1404
+ "\n",
1405
+ "print(\n",
1406
+ " \"Majority lookup latency:\",\n",
1407
+ " f\"{majority_latency_ms:.6f} ms\",\n",
1408
+ ")\n"
1409
+ ]
1410
+ },
1411
+ {
1412
+ "cell_type": "markdown",
1413
+ "id": "b4db8d68",
1414
+ "metadata": {},
1415
+ "source": [
1416
+ "## 11. Evaluate Frozen CLIP + Heads (V1)"
1417
+ ]
1418
+ },
1419
+ {
1420
+ "cell_type": "code",
1421
+ "execution_count": 12,
1422
+ "id": "f0ec76a3",
1423
+ "metadata": {
1424
+ "execution": {
1425
+ "iopub.execute_input": "2026-07-06T14:51:44.397516Z",
1426
+ "iopub.status.busy": "2026-07-06T14:51:44.396710Z",
1427
+ "iopub.status.idle": "2026-07-06T14:52:39.006069Z",
1428
+ "shell.execute_reply": "2026-07-06T14:52:39.005197Z",
1429
+ "shell.execute_reply.started": "2026-07-06T14:51:44.397482Z"
1430
+ },
1431
+ "trusted": true
1432
+ },
1433
+ "outputs": [
1434
+ {
1435
+ "data": {
1436
+ "application/vnd.jupyter.widget-view+json": {
1437
+ "model_id": "",
1438
+ "version_major": 2,
1439
+ "version_minor": 0
1440
+ },
1441
+ "text/plain": [
1442
+ "Evaluating: 0%| | 0/104 [00:00<?, ?it/s]"
1443
+ ]
1444
+ },
1445
+ "metadata": {},
1446
+ "output_type": "display_data"
1447
+ },
1448
+ {
1449
+ "name": "stdout",
1450
+ "output_type": "stream",
1451
+ "text": [
1452
+ "{\n",
1453
+ " \"average_accuracy\": 0.8335026038852995,\n",
1454
+ " \"average_macro_f1\": 0.6568447666058456,\n",
1455
+ " \"average_top3_accuracy\": 0.9711087581303888,\n",
1456
+ " \"exact_match_accuracy\": 0.27938284677053393,\n",
1457
+ " \"samples\": 6611\n",
1458
+ "}\n"
1459
+ ]
1460
+ }
1461
+ ],
1462
+ "source": [
1463
+ "v1_checkpoint = v1_artifacts[\"checkpoint\"]\n",
1464
+ "\n",
1465
+ "v1_model_name = (\n",
1466
+ " v1_config.get(\"base_model_name\")\n",
1467
+ " or v1_config.get(\"model_name\")\n",
1468
+ " or v1_checkpoint.get(\"model_name\")\n",
1469
+ " or MODEL_NAME\n",
1470
+ ")\n",
1471
+ "\n",
1472
+ "v1_hidden_dim = int(\n",
1473
+ " v1_config.get(\n",
1474
+ " \"hidden_dim\",\n",
1475
+ " v1_checkpoint.get(\n",
1476
+ " \"hidden_dim\",\n",
1477
+ " 512,\n",
1478
+ " ),\n",
1479
+ " )\n",
1480
+ ")\n",
1481
+ "\n",
1482
+ "v1_dropout = float(\n",
1483
+ " v1_config.get(\n",
1484
+ " \"dropout\",\n",
1485
+ " v1_checkpoint.get(\n",
1486
+ " \"dropout\",\n",
1487
+ " 0.2,\n",
1488
+ " ),\n",
1489
+ " )\n",
1490
+ ")\n",
1491
+ "\n",
1492
+ "v1_model = CLIPMultiTaskClassifier(\n",
1493
+ " model_name=v1_model_name,\n",
1494
+ " task_num_classes=task_num_classes,\n",
1495
+ " hidden_dim=v1_hidden_dim,\n",
1496
+ " dropout=v1_dropout,\n",
1497
+ ")\n",
1498
+ "\n",
1499
+ "v1_model.load_state_dict(\n",
1500
+ " get_state_dict(v1_checkpoint),\n",
1501
+ " strict=True,\n",
1502
+ ")\n",
1503
+ "\n",
1504
+ "v1_model.to(DEVICE)\n",
1505
+ "v1_model.eval()\n",
1506
+ "\n",
1507
+ "(\n",
1508
+ " v1_y_true,\n",
1509
+ " v1_y_pred,\n",
1510
+ " v1_y_top3,\n",
1511
+ " v1_global_indices,\n",
1512
+ ") = collect_model_predictions(\n",
1513
+ " v1_model,\n",
1514
+ " test_loader,\n",
1515
+ " DEVICE,\n",
1516
+ " uses_color_features=False,\n",
1517
+ ")\n",
1518
+ "\n",
1519
+ "v1_metrics = evaluate_predictions(\n",
1520
+ " v1_y_true,\n",
1521
+ " v1_y_pred,\n",
1522
+ " v1_y_top3,\n",
1523
+ ")\n",
1524
+ "\n",
1525
+ "print(\n",
1526
+ " json.dumps(\n",
1527
+ " v1_metrics[\"overall_metrics\"],\n",
1528
+ " indent=2,\n",
1529
+ " )\n",
1530
+ ")\n"
1531
+ ]
1532
+ },
1533
+ {
1534
+ "cell_type": "markdown",
1535
+ "id": "99fde9e8",
1536
+ "metadata": {},
1537
+ "source": [
1538
+ "## 12. Fair Latency Benchmark"
1539
+ ]
1540
+ },
1541
+ {
1542
+ "cell_type": "code",
1543
+ "execution_count": 13,
1544
+ "id": "bc279ecd",
1545
+ "metadata": {
1546
+ "execution": {
1547
+ "iopub.execute_input": "2026-07-06T14:52:43.888474Z",
1548
+ "iopub.status.busy": "2026-07-06T14:52:43.888075Z",
1549
+ "iopub.status.idle": "2026-07-06T14:52:44.676522Z",
1550
+ "shell.execute_reply": "2026-07-06T14:52:44.675520Z",
1551
+ "shell.execute_reply.started": "2026-07-06T14:52:43.888445Z"
1552
+ },
1553
+ "trusted": true
1554
+ },
1555
+ "outputs": [
1556
+ {
1557
+ "name": "stdout",
1558
+ "output_type": "stream",
1559
+ "text": [
1560
+ "{'average_ms': 6.226557110003341, 'p50_ms': 6.155085000045801, 'p95_ms': 6.70613664992743, 'warmup_runs': 20, 'measured_runs': 100}\n"
1561
+ ]
1562
+ }
1563
+ ],
1564
+ "source": [
1565
+ "@torch.inference_mode()\n",
1566
+ "def benchmark_model_latency(\n",
1567
+ " model,\n",
1568
+ " sample,\n",
1569
+ " device,\n",
1570
+ " uses_color_features,\n",
1571
+ " warmup_runs=20,\n",
1572
+ " measured_runs=100,\n",
1573
+ "):\n",
1574
+ " model.eval()\n",
1575
+ "\n",
1576
+ " pixel_values = sample[\n",
1577
+ " \"pixel_values\"\n",
1578
+ " ].unsqueeze(0).to(device)\n",
1579
+ "\n",
1580
+ " color_features = sample[\n",
1581
+ " \"color_features\"\n",
1582
+ " ].unsqueeze(0).to(device)\n",
1583
+ "\n",
1584
+ " for _ in range(warmup_runs):\n",
1585
+ " if uses_color_features:\n",
1586
+ " model(\n",
1587
+ " pixel_values,\n",
1588
+ " color_features,\n",
1589
+ " )\n",
1590
+ " else:\n",
1591
+ " model(pixel_values)\n",
1592
+ "\n",
1593
+ " if device.type == \"cuda\":\n",
1594
+ " torch.cuda.synchronize()\n",
1595
+ "\n",
1596
+ " times = []\n",
1597
+ "\n",
1598
+ " for _ in range(measured_runs):\n",
1599
+ " if device.type == \"cuda\":\n",
1600
+ " torch.cuda.synchronize()\n",
1601
+ "\n",
1602
+ " start = time.perf_counter()\n",
1603
+ "\n",
1604
+ " if uses_color_features:\n",
1605
+ " model(\n",
1606
+ " pixel_values,\n",
1607
+ " color_features,\n",
1608
+ " )\n",
1609
+ " else:\n",
1610
+ " model(pixel_values)\n",
1611
+ "\n",
1612
+ " if device.type == \"cuda\":\n",
1613
+ " torch.cuda.synchronize()\n",
1614
+ "\n",
1615
+ " times.append(\n",
1616
+ " (\n",
1617
+ " time.perf_counter() - start\n",
1618
+ " )\n",
1619
+ " * 1000.0\n",
1620
+ " )\n",
1621
+ "\n",
1622
+ " return {\n",
1623
+ " \"average_ms\": float(\n",
1624
+ " np.mean(times)\n",
1625
+ " ),\n",
1626
+ " \"p50_ms\": float(\n",
1627
+ " np.percentile(times, 50)\n",
1628
+ " ),\n",
1629
+ " \"p95_ms\": float(\n",
1630
+ " np.percentile(times, 95)\n",
1631
+ " ),\n",
1632
+ " \"warmup_runs\": warmup_runs,\n",
1633
+ " \"measured_runs\": measured_runs,\n",
1634
+ " }\n",
1635
+ "\n",
1636
+ "\n",
1637
+ "latency_sample = test_dataset[0]\n",
1638
+ "\n",
1639
+ "v1_latency = benchmark_model_latency(\n",
1640
+ " v1_model,\n",
1641
+ " latency_sample,\n",
1642
+ " DEVICE,\n",
1643
+ " uses_color_features=False,\n",
1644
+ " warmup_runs=LATENCY_WARMUP_RUNS,\n",
1645
+ " measured_runs=LATENCY_MEASURED_RUNS,\n",
1646
+ ")\n",
1647
+ "\n",
1648
+ "print(v1_latency)\n"
1649
+ ]
1650
+ },
1651
+ {
1652
+ "cell_type": "markdown",
1653
+ "id": "c5593fe6",
1654
+ "metadata": {},
1655
+ "source": [
1656
+ "## 13. Evaluate AutoCatalogAI V2"
1657
+ ]
1658
+ },
1659
+ {
1660
+ "cell_type": "code",
1661
+ "execution_count": 14,
1662
+ "id": "86188832",
1663
+ "metadata": {
1664
+ "execution": {
1665
+ "iopub.execute_input": "2026-07-06T14:52:47.240600Z",
1666
+ "iopub.status.busy": "2026-07-06T14:52:47.240043Z",
1667
+ "iopub.status.idle": "2026-07-06T14:53:39.405651Z",
1668
+ "shell.execute_reply": "2026-07-06T14:53:39.404397Z",
1669
+ "shell.execute_reply.started": "2026-07-06T14:52:47.240567Z"
1670
+ },
1671
+ "trusted": true
1672
+ },
1673
+ "outputs": [
1674
+ {
1675
+ "data": {
1676
+ "application/vnd.jupyter.widget-view+json": {
1677
+ "model_id": "",
1678
+ "version_major": 2,
1679
+ "version_minor": 0
1680
+ },
1681
+ "text/plain": [
1682
+ "Evaluating: 0%| | 0/104 [00:00<?, ?it/s]"
1683
+ ]
1684
+ },
1685
+ "metadata": {},
1686
+ "output_type": "display_data"
1687
+ },
1688
+ {
1689
+ "name": "stdout",
1690
+ "output_type": "stream",
1691
+ "text": [
1692
+ "{\n",
1693
+ " \"average_accuracy\": 0.8747758065561727,\n",
1694
+ " \"average_macro_f1\": 0.6740687827687263,\n",
1695
+ " \"average_top3_accuracy\": 0.981502690321326,\n",
1696
+ " \"exact_match_accuracy\": 0.4046286492209953,\n",
1697
+ " \"samples\": 6611\n",
1698
+ "}\n",
1699
+ "{'average_ms': 7.4737231199719645, 'p50_ms': 7.384413499949005, 'p95_ms': 8.047834450019309, 'warmup_runs': 20, 'measured_runs': 100}\n"
1700
+ ]
1701
+ }
1702
+ ],
1703
+ "source": [
1704
+ "del v1_model\n",
1705
+ "\n",
1706
+ "gc.collect()\n",
1707
+ "\n",
1708
+ "if torch.cuda.is_available():\n",
1709
+ " torch.cuda.empty_cache()\n",
1710
+ "\n",
1711
+ "\n",
1712
+ "v2_checkpoint = v2_artifacts[\"checkpoint\"]\n",
1713
+ "\n",
1714
+ "v2_model = CLIPMultiTaskClassifierV2(\n",
1715
+ " model_name=(\n",
1716
+ " v2_checkpoint.get(\n",
1717
+ " \"model_name\",\n",
1718
+ " MODEL_NAME,\n",
1719
+ " )\n",
1720
+ " ),\n",
1721
+ " task_num_classes=(\n",
1722
+ " v2_checkpoint.get(\n",
1723
+ " \"task_num_classes\",\n",
1724
+ " task_num_classes,\n",
1725
+ " )\n",
1726
+ " ),\n",
1727
+ " hidden_dim=int(\n",
1728
+ " v2_checkpoint.get(\n",
1729
+ " \"hidden_dim\",\n",
1730
+ " HIDDEN_DIM,\n",
1731
+ " )\n",
1732
+ " ),\n",
1733
+ " dropout=float(\n",
1734
+ " v2_checkpoint.get(\n",
1735
+ " \"dropout\",\n",
1736
+ " DROPOUT,\n",
1737
+ " )\n",
1738
+ " ),\n",
1739
+ " color_feature_dim=int(\n",
1740
+ " v2_checkpoint.get(\n",
1741
+ " \"color_feature_dim\",\n",
1742
+ " COLOR_FEATURE_DIM,\n",
1743
+ " )\n",
1744
+ " ),\n",
1745
+ ")\n",
1746
+ "\n",
1747
+ "v2_model.load_state_dict(\n",
1748
+ " get_state_dict(v2_checkpoint),\n",
1749
+ " strict=True,\n",
1750
+ ")\n",
1751
+ "\n",
1752
+ "v2_model.to(DEVICE)\n",
1753
+ "v2_model.eval()\n",
1754
+ "\n",
1755
+ "(\n",
1756
+ " v2_y_true,\n",
1757
+ " v2_y_pred,\n",
1758
+ " v2_y_top3,\n",
1759
+ " v2_global_indices,\n",
1760
+ ") = collect_model_predictions(\n",
1761
+ " v2_model,\n",
1762
+ " test_loader,\n",
1763
+ " DEVICE,\n",
1764
+ " uses_color_features=True,\n",
1765
+ ")\n",
1766
+ "\n",
1767
+ "v2_metrics = evaluate_predictions(\n",
1768
+ " v2_y_true,\n",
1769
+ " v2_y_pred,\n",
1770
+ " v2_y_top3,\n",
1771
+ ")\n",
1772
+ "\n",
1773
+ "v2_latency = benchmark_model_latency(\n",
1774
+ " v2_model,\n",
1775
+ " latency_sample,\n",
1776
+ " DEVICE,\n",
1777
+ " uses_color_features=True,\n",
1778
+ " warmup_runs=LATENCY_WARMUP_RUNS,\n",
1779
+ " measured_runs=LATENCY_MEASURED_RUNS,\n",
1780
+ ")\n",
1781
+ "\n",
1782
+ "print(\n",
1783
+ " json.dumps(\n",
1784
+ " v2_metrics[\"overall_metrics\"],\n",
1785
+ " indent=2,\n",
1786
+ " )\n",
1787
+ ")\n",
1788
+ "\n",
1789
+ "print(v2_latency)\n"
1790
+ ]
1791
+ },
1792
+ {
1793
+ "cell_type": "markdown",
1794
+ "id": "f8f35b7f",
1795
+ "metadata": {},
1796
+ "source": [
1797
+ "## 14. Build the Comparison Table\n",
1798
+ "\n",
1799
+ "The table uses **raw model predictions** for a fair comparison. \n",
1800
+ "V2 consistency correction is not applied here because the other baselines do not use post-processing.\n"
1801
+ ]
1802
+ },
1803
+ {
1804
+ "cell_type": "code",
1805
+ "execution_count": 16,
1806
+ "id": "303f42be",
1807
+ "metadata": {
1808
+ "execution": {
1809
+ "iopub.execute_input": "2026-07-06T14:53:47.885588Z",
1810
+ "iopub.status.busy": "2026-07-06T14:53:47.885021Z",
1811
+ "iopub.status.idle": "2026-07-06T14:53:47.904148Z",
1812
+ "shell.execute_reply": "2026-07-06T14:53:47.902890Z",
1813
+ "shell.execute_reply.started": "2026-07-06T14:53:47.885552Z"
1814
+ },
1815
+ "trusted": true
1816
+ },
1817
+ "outputs": [
1818
+ {
1819
+ "data": {
1820
+ "text/html": [
1821
+ "<div>\n",
1822
+ "<style scoped>\n",
1823
+ " .dataframe tbody tr th:only-of-type {\n",
1824
+ " vertical-align: middle;\n",
1825
+ " }\n",
1826
+ "\n",
1827
+ " .dataframe tbody tr th {\n",
1828
+ " vertical-align: top;\n",
1829
+ " }\n",
1830
+ "\n",
1831
+ " .dataframe thead th {\n",
1832
+ " text-align: right;\n",
1833
+ " }\n",
1834
+ "</style>\n",
1835
+ "<table border=\"1\" class=\"dataframe\">\n",
1836
+ " <thead>\n",
1837
+ " <tr style=\"text-align: right;\">\n",
1838
+ " <th></th>\n",
1839
+ " <th>Model</th>\n",
1840
+ " <th>Avg Accuracy</th>\n",
1841
+ " <th>Avg Macro F1</th>\n",
1842
+ " <th>Top-3 Accuracy</th>\n",
1843
+ " <th>Exact Match</th>\n",
1844
+ " <th>Latency (ms)</th>\n",
1845
+ " </tr>\n",
1846
+ " </thead>\n",
1847
+ " <tbody>\n",
1848
+ " <tr>\n",
1849
+ " <th>0</th>\n",
1850
+ " <td>Majority Baseline</td>\n",
1851
+ " <td>42.63%</td>\n",
1852
+ " <td>7.93%</td>\n",
1853
+ " <td>73.44%</td>\n",
1854
+ " <td>0.95%</td>\n",
1855
+ " <td>0.001</td>\n",
1856
+ " </tr>\n",
1857
+ " <tr>\n",
1858
+ " <th>1</th>\n",
1859
+ " <td>Frozen CLIP + Heads (V1)</td>\n",
1860
+ " <td>83.35%</td>\n",
1861
+ " <td>65.68%</td>\n",
1862
+ " <td>97.11%</td>\n",
1863
+ " <td>27.94%</td>\n",
1864
+ " <td>6.227</td>\n",
1865
+ " </tr>\n",
1866
+ " <tr>\n",
1867
+ " <th>2</th>\n",
1868
+ " <td>AutoCatalogAI V2</td>\n",
1869
+ " <td>87.48%</td>\n",
1870
+ " <td>67.41%</td>\n",
1871
+ " <td>98.15%</td>\n",
1872
+ " <td>40.46%</td>\n",
1873
+ " <td>7.474</td>\n",
1874
+ " </tr>\n",
1875
+ " </tbody>\n",
1876
+ "</table>\n",
1877
+ "</div>"
1878
+ ],
1879
+ "text/plain": [
1880
+ " Model Avg Accuracy Avg Macro F1 Top-3 Accuracy \\\n",
1881
+ "0 Majority Baseline 42.63% 7.93% 73.44% \n",
1882
+ "1 Frozen CLIP + Heads (V1) 83.35% 65.68% 97.11% \n",
1883
+ "2 AutoCatalogAI V2 87.48% 67.41% 98.15% \n",
1884
+ "\n",
1885
+ " Exact Match Latency (ms) \n",
1886
+ "0 0.95% 0.001 \n",
1887
+ "1 27.94% 6.227 \n",
1888
+ "2 40.46% 7.474 "
1889
+ ]
1890
+ },
1891
+ "metadata": {},
1892
+ "output_type": "display_data"
1893
+ }
1894
+ ],
1895
+ "source": [
1896
+ "comparison_rows = []\n",
1897
+ "\n",
1898
+ "\n",
1899
+ "def add_comparison_row(\n",
1900
+ " model_name,\n",
1901
+ " metrics,\n",
1902
+ " latency_ms,\n",
1903
+ "):\n",
1904
+ " overall = metrics[\"overall_metrics\"]\n",
1905
+ "\n",
1906
+ " comparison_rows.append(\n",
1907
+ " {\n",
1908
+ " \"Model\": model_name,\n",
1909
+ " \"Avg Accuracy\": overall[\n",
1910
+ " \"average_accuracy\"\n",
1911
+ " ],\n",
1912
+ " \"Avg Macro F1\": overall[\n",
1913
+ " \"average_macro_f1\"\n",
1914
+ " ],\n",
1915
+ " \"Top-3 Accuracy\": overall[\n",
1916
+ " \"average_top3_accuracy\"\n",
1917
+ " ],\n",
1918
+ " \"Exact Match\": overall[\n",
1919
+ " \"exact_match_accuracy\"\n",
1920
+ " ],\n",
1921
+ " \"Latency (ms)\": latency_ms,\n",
1922
+ " }\n",
1923
+ " )\n",
1924
+ "\n",
1925
+ "\n",
1926
+ "add_comparison_row(\n",
1927
+ " \"Majority Baseline\",\n",
1928
+ " majority_metrics,\n",
1929
+ " majority_latency_ms,\n",
1930
+ ")\n",
1931
+ "\n",
1932
+ "add_comparison_row(\n",
1933
+ " \"Frozen CLIP + Heads (V1)\",\n",
1934
+ " v1_metrics,\n",
1935
+ " v1_latency[\"average_ms\"],\n",
1936
+ ")\n",
1937
+ "\n",
1938
+ "add_comparison_row(\n",
1939
+ " \"AutoCatalogAI V2\",\n",
1940
+ " v2_metrics,\n",
1941
+ " v2_latency[\"average_ms\"],\n",
1942
+ ")\n",
1943
+ "\n",
1944
+ "\n",
1945
+ "comparison_df = pd.DataFrame(\n",
1946
+ " comparison_rows\n",
1947
+ ")\n",
1948
+ "\n",
1949
+ "display_df = comparison_df.copy()\n",
1950
+ "\n",
1951
+ "for column in [\n",
1952
+ " \"Avg Accuracy\",\n",
1953
+ " \"Avg Macro F1\",\n",
1954
+ " \"Top-3 Accuracy\",\n",
1955
+ " \"Exact Match\",\n",
1956
+ "]:\n",
1957
+ " display_df[column] = (\n",
1958
+ " display_df[column] * 100\n",
1959
+ " ).map(lambda value: f\"{value:.2f}%\")\n",
1960
+ "\n",
1961
+ "display_df[\"Latency (ms)\"] = (\n",
1962
+ " display_df[\"Latency (ms)\"]\n",
1963
+ " .map(lambda value: f\"{value:.3f}\")\n",
1964
+ ")\n",
1965
+ "\n",
1966
+ "display(display_df)\n"
1967
+ ]
1968
+ },
1969
+ {
1970
+ "cell_type": "markdown",
1971
+ "id": "3f4e3ac7",
1972
+ "metadata": {},
1973
+ "source": [
1974
+ "## 15. Per-Task Metrics"
1975
+ ]
1976
+ },
1977
+ {
1978
+ "cell_type": "code",
1979
+ "execution_count": 17,
1980
+ "id": "6c7e45bc",
1981
+ "metadata": {
1982
+ "execution": {
1983
+ "iopub.execute_input": "2026-07-06T14:54:06.403485Z",
1984
+ "iopub.status.busy": "2026-07-06T14:54:06.402755Z",
1985
+ "iopub.status.idle": "2026-07-06T14:54:06.419062Z",
1986
+ "shell.execute_reply": "2026-07-06T14:54:06.418316Z",
1987
+ "shell.execute_reply.started": "2026-07-06T14:54:06.403449Z"
1988
+ },
1989
+ "trusted": true
1990
+ },
1991
+ "outputs": [
1992
+ {
1993
+ "data": {
1994
+ "text/html": [
1995
+ "<div>\n",
1996
+ "<style scoped>\n",
1997
+ " .dataframe tbody tr th:only-of-type {\n",
1998
+ " vertical-align: middle;\n",
1999
+ " }\n",
2000
+ "\n",
2001
+ " .dataframe tbody tr th {\n",
2002
+ " vertical-align: top;\n",
2003
+ " }\n",
2004
+ "\n",
2005
+ " .dataframe thead th {\n",
2006
+ " text-align: right;\n",
2007
+ " }\n",
2008
+ "</style>\n",
2009
+ "<table border=\"1\" class=\"dataframe\">\n",
2010
+ " <thead>\n",
2011
+ " <tr style=\"text-align: right;\">\n",
2012
+ " <th></th>\n",
2013
+ " <th>Model</th>\n",
2014
+ " <th>Task</th>\n",
2015
+ " <th>Accuracy</th>\n",
2016
+ " <th>Macro F1</th>\n",
2017
+ " <th>Top-3 Accuracy</th>\n",
2018
+ " </tr>\n",
2019
+ " </thead>\n",
2020
+ " <tbody>\n",
2021
+ " <tr>\n",
2022
+ " <th>0</th>\n",
2023
+ " <td>Majority Baseline</td>\n",
2024
+ " <td>gender</td>\n",
2025
+ " <td>0.499168</td>\n",
2026
+ " <td>0.133185</td>\n",
2027
+ " <td>0.969899</td>\n",
2028
+ " </tr>\n",
2029
+ " <tr>\n",
2030
+ " <th>1</th>\n",
2031
+ " <td>Majority Baseline</td>\n",
2032
+ " <td>masterCategory</td>\n",
2033
+ " <td>0.484496</td>\n",
2034
+ " <td>0.108790</td>\n",
2035
+ " <td>0.948571</td>\n",
2036
+ " </tr>\n",
2037
+ " <tr>\n",
2038
+ " <th>2</th>\n",
2039
+ " <td>Majority Baseline</td>\n",
2040
+ " <td>subCategory</td>\n",
2041
+ " <td>0.349115</td>\n",
2042
+ " <td>0.012939</td>\n",
2043
+ " <td>0.584329</td>\n",
2044
+ " </tr>\n",
2045
+ " <tr>\n",
2046
+ " <th>3</th>\n",
2047
+ " <td>Majority Baseline</td>\n",
2048
+ " <td>articleType</td>\n",
2049
+ " <td>0.160339</td>\n",
2050
+ " <td>0.002176</td>\n",
2051
+ " <td>0.297837</td>\n",
2052
+ " </tr>\n",
2053
+ " <tr>\n",
2054
+ " <th>4</th>\n",
2055
+ " <td>Majority Baseline</td>\n",
2056
+ " <td>baseColour</td>\n",
2057
+ " <td>0.218726</td>\n",
2058
+ " <td>0.008158</td>\n",
2059
+ " <td>0.456209</td>\n",
2060
+ " </tr>\n",
2061
+ " <tr>\n",
2062
+ " <th>5</th>\n",
2063
+ " <td>Majority Baseline</td>\n",
2064
+ " <td>season</td>\n",
2065
+ " <td>0.489033</td>\n",
2066
+ " <td>0.164212</td>\n",
2067
+ " <td>0.939949</td>\n",
2068
+ " </tr>\n",
2069
+ " <tr>\n",
2070
+ " <th>6</th>\n",
2071
+ " <td>Majority Baseline</td>\n",
2072
+ " <td>usage</td>\n",
2073
+ " <td>0.783089</td>\n",
2074
+ " <td>0.125479</td>\n",
2075
+ " <td>0.944335</td>\n",
2076
+ " </tr>\n",
2077
+ " <tr>\n",
2078
+ " <th>7</th>\n",
2079
+ " <td>Frozen CLIP + Heads (V1)</td>\n",
2080
+ " <td>gender</td>\n",
2081
+ " <td>0.888670</td>\n",
2082
+ " <td>0.776600</td>\n",
2083
+ " <td>0.995613</td>\n",
2084
+ " </tr>\n",
2085
+ " <tr>\n",
2086
+ " <th>8</th>\n",
2087
+ " <td>Frozen CLIP + Heads (V1)</td>\n",
2088
+ " <td>masterCategory</td>\n",
2089
+ " <td>0.993344</td>\n",
2090
+ " <td>0.862994</td>\n",
2091
+ " <td>0.999395</td>\n",
2092
+ " </tr>\n",
2093
+ " <tr>\n",
2094
+ " <th>9</th>\n",
2095
+ " <td>Frozen CLIP + Heads (V1)</td>\n",
2096
+ " <td>subCategory</td>\n",
2097
+ " <td>0.942974</td>\n",
2098
+ " <td>0.710894</td>\n",
2099
+ " <td>0.994706</td>\n",
2100
+ " </tr>\n",
2101
+ " <tr>\n",
2102
+ " <th>10</th>\n",
2103
+ " <td>Frozen CLIP + Heads (V1)</td>\n",
2104
+ " <td>articleType</td>\n",
2105
+ " <td>0.841023</td>\n",
2106
+ " <td>0.663970</td>\n",
2107
+ " <td>0.971109</td>\n",
2108
+ " </tr>\n",
2109
+ " <tr>\n",
2110
+ " <th>11</th>\n",
2111
+ " <td>Frozen CLIP + Heads (V1)</td>\n",
2112
+ " <td>baseColour</td>\n",
2113
+ " <td>0.601119</td>\n",
2114
+ " <td>0.340958</td>\n",
2115
+ " <td>0.853880</td>\n",
2116
+ " </tr>\n",
2117
+ " <tr>\n",
2118
+ " <th>12</th>\n",
2119
+ " <td>Frozen CLIP + Heads (V1)</td>\n",
2120
+ " <td>season</td>\n",
2121
+ " <td>0.707003</td>\n",
2122
+ " <td>0.728013</td>\n",
2123
+ " <td>0.984874</td>\n",
2124
+ " </tr>\n",
2125
+ " <tr>\n",
2126
+ " <th>13</th>\n",
2127
+ " <td>Frozen CLIP + Heads (V1)</td>\n",
2128
+ " <td>usage</td>\n",
2129
+ " <td>0.860384</td>\n",
2130
+ " <td>0.514485</td>\n",
2131
+ " <td>0.998185</td>\n",
2132
+ " </tr>\n",
2133
+ " <tr>\n",
2134
+ " <th>14</th>\n",
2135
+ " <td>AutoCatalogAI V2</td>\n",
2136
+ " <td>gender</td>\n",
2137
+ " <td>0.919226</td>\n",
2138
+ " <td>0.812002</td>\n",
2139
+ " <td>0.997429</td>\n",
2140
+ " </tr>\n",
2141
+ " <tr>\n",
2142
+ " <th>15</th>\n",
2143
+ " <td>AutoCatalogAI V2</td>\n",
2144
+ " <td>masterCategory</td>\n",
2145
+ " <td>0.994403</td>\n",
2146
+ " <td>0.846441</td>\n",
2147
+ " <td>0.999395</td>\n",
2148
+ " </tr>\n",
2149
+ " <tr>\n",
2150
+ " <th>16</th>\n",
2151
+ " <td>AutoCatalogAI V2</td>\n",
2152
+ " <td>subCategory</td>\n",
2153
+ " <td>0.963546</td>\n",
2154
+ " <td>0.759314</td>\n",
2155
+ " <td>0.997277</td>\n",
2156
+ " </tr>\n",
2157
+ " <tr>\n",
2158
+ " <th>17</th>\n",
2159
+ " <td>AutoCatalogAI V2</td>\n",
2160
+ " <td>articleType</td>\n",
2161
+ " <td>0.876418</td>\n",
2162
+ " <td>0.663665</td>\n",
2163
+ " <td>0.981546</td>\n",
2164
+ " </tr>\n",
2165
+ " <tr>\n",
2166
+ " <th>18</th>\n",
2167
+ " <td>AutoCatalogAI V2</td>\n",
2168
+ " <td>baseColour</td>\n",
2169
+ " <td>0.697171</td>\n",
2170
+ " <td>0.364965</td>\n",
2171
+ " <td>0.906973</td>\n",
2172
+ " </tr>\n",
2173
+ " <tr>\n",
2174
+ " <th>19</th>\n",
2175
+ " <td>AutoCatalogAI V2</td>\n",
2176
+ " <td>season</td>\n",
2177
+ " <td>0.754803</td>\n",
2178
+ " <td>0.767844</td>\n",
2179
+ " <td>0.989714</td>\n",
2180
+ " </tr>\n",
2181
+ " <tr>\n",
2182
+ " <th>20</th>\n",
2183
+ " <td>AutoCatalogAI V2</td>\n",
2184
+ " <td>usage</td>\n",
2185
+ " <td>0.917864</td>\n",
2186
+ " <td>0.504250</td>\n",
2187
+ " <td>0.998185</td>\n",
2188
+ " </tr>\n",
2189
+ " </tbody>\n",
2190
+ "</table>\n",
2191
+ "</div>"
2192
+ ],
2193
+ "text/plain": [
2194
+ " Model Task Accuracy Macro F1 \\\n",
2195
+ "0 Majority Baseline gender 0.499168 0.133185 \n",
2196
+ "1 Majority Baseline masterCategory 0.484496 0.108790 \n",
2197
+ "2 Majority Baseline subCategory 0.349115 0.012939 \n",
2198
+ "3 Majority Baseline articleType 0.160339 0.002176 \n",
2199
+ "4 Majority Baseline baseColour 0.218726 0.008158 \n",
2200
+ "5 Majority Baseline season 0.489033 0.164212 \n",
2201
+ "6 Majority Baseline usage 0.783089 0.125479 \n",
2202
+ "7 Frozen CLIP + Heads (V1) gender 0.888670 0.776600 \n",
2203
+ "8 Frozen CLIP + Heads (V1) masterCategory 0.993344 0.862994 \n",
2204
+ "9 Frozen CLIP + Heads (V1) subCategory 0.942974 0.710894 \n",
2205
+ "10 Frozen CLIP + Heads (V1) articleType 0.841023 0.663970 \n",
2206
+ "11 Frozen CLIP + Heads (V1) baseColour 0.601119 0.340958 \n",
2207
+ "12 Frozen CLIP + Heads (V1) season 0.707003 0.728013 \n",
2208
+ "13 Frozen CLIP + Heads (V1) usage 0.860384 0.514485 \n",
2209
+ "14 AutoCatalogAI V2 gender 0.919226 0.812002 \n",
2210
+ "15 AutoCatalogAI V2 masterCategory 0.994403 0.846441 \n",
2211
+ "16 AutoCatalogAI V2 subCategory 0.963546 0.759314 \n",
2212
+ "17 AutoCatalogAI V2 articleType 0.876418 0.663665 \n",
2213
+ "18 AutoCatalogAI V2 baseColour 0.697171 0.364965 \n",
2214
+ "19 AutoCatalogAI V2 season 0.754803 0.767844 \n",
2215
+ "20 AutoCatalogAI V2 usage 0.917864 0.504250 \n",
2216
+ "\n",
2217
+ " Top-3 Accuracy \n",
2218
+ "0 0.969899 \n",
2219
+ "1 0.948571 \n",
2220
+ "2 0.584329 \n",
2221
+ "3 0.297837 \n",
2222
+ "4 0.456209 \n",
2223
+ "5 0.939949 \n",
2224
+ "6 0.944335 \n",
2225
+ "7 0.995613 \n",
2226
+ "8 0.999395 \n",
2227
+ "9 0.994706 \n",
2228
+ "10 0.971109 \n",
2229
+ "11 0.853880 \n",
2230
+ "12 0.984874 \n",
2231
+ "13 0.998185 \n",
2232
+ "14 0.997429 \n",
2233
+ "15 0.999395 \n",
2234
+ "16 0.997277 \n",
2235
+ "17 0.981546 \n",
2236
+ "18 0.906973 \n",
2237
+ "19 0.989714 \n",
2238
+ "20 0.998185 "
2239
+ ]
2240
+ },
2241
+ "metadata": {},
2242
+ "output_type": "display_data"
2243
+ }
2244
+ ],
2245
+ "source": [
2246
+ "per_task_rows = []\n",
2247
+ "\n",
2248
+ "models_and_metrics = {\n",
2249
+ " \"Majority Baseline\": majority_metrics,\n",
2250
+ " \"Frozen CLIP + Heads (V1)\": v1_metrics,\n",
2251
+ " \"AutoCatalogAI V2\": v2_metrics,\n",
2252
+ "}\n",
2253
+ "\n",
2254
+ "for model_name, metrics in models_and_metrics.items():\n",
2255
+ " for task in TASKS:\n",
2256
+ " task_result = metrics[\n",
2257
+ " \"task_metrics\"\n",
2258
+ " ][task]\n",
2259
+ "\n",
2260
+ " per_task_rows.append(\n",
2261
+ " {\n",
2262
+ " \"Model\": model_name,\n",
2263
+ " \"Task\": task,\n",
2264
+ " \"Accuracy\": task_result[\n",
2265
+ " \"accuracy\"\n",
2266
+ " ],\n",
2267
+ " \"Macro F1\": task_result[\n",
2268
+ " \"macro_f1\"\n",
2269
+ " ],\n",
2270
+ " \"Top-3 Accuracy\": task_result[\n",
2271
+ " \"top3_accuracy\"\n",
2272
+ " ],\n",
2273
+ " }\n",
2274
+ " )\n",
2275
+ "\n",
2276
+ "\n",
2277
+ "per_task_df = pd.DataFrame(\n",
2278
+ " per_task_rows\n",
2279
+ ")\n",
2280
+ "\n",
2281
+ "display(per_task_df)\n"
2282
+ ]
2283
+ },
2284
+ {
2285
+ "cell_type": "markdown",
2286
+ "id": "a877244a",
2287
+ "metadata": {},
2288
+ "source": [
2289
+ "## 16. Save Reproducibility Artifacts"
2290
+ ]
2291
+ },
2292
+ {
2293
+ "cell_type": "code",
2294
+ "execution_count": 18,
2295
+ "id": "2f752616",
2296
+ "metadata": {
2297
+ "execution": {
2298
+ "iopub.execute_input": "2026-07-06T14:54:13.885906Z",
2299
+ "iopub.status.busy": "2026-07-06T14:54:13.885429Z",
2300
+ "iopub.status.idle": "2026-07-06T14:54:13.905372Z",
2301
+ "shell.execute_reply": "2026-07-06T14:54:13.904349Z",
2302
+ "shell.execute_reply.started": "2026-07-06T14:54:13.885872Z"
2303
+ },
2304
+ "trusted": true
2305
+ },
2306
+ "outputs": [
2307
+ {
2308
+ "name": "stdout",
2309
+ "output_type": "stream",
2310
+ "text": [
2311
+ "| Model | Avg Accuracy | Avg Macro F1 | Top-3 Accuracy | Exact Match | Latency |\n",
2312
+ "|---|---:|---:|---:|---:|---:|\n",
2313
+ "| Majority Baseline | 42.63% | 7.93% | 73.44% | 0.95% | 0.001 ms |\n",
2314
+ "| Frozen CLIP + Heads (V1) | 83.35% | 65.68% | 97.11% | 27.94% | 6.227 ms |\n",
2315
+ "| AutoCatalogAI V2 | 87.48% | 67.41% | 98.15% | 40.46% | 7.474 ms |\n",
2316
+ "\n",
2317
+ "> All models were evaluated on the same 6,611-image held-out test split. Metrics use raw model predictions. Latency is batch-size-1 model-forward time with preprocessing excluded.\n",
2318
+ "\n",
2319
+ "Saved files:\n",
2320
+ "- artifacts/evaluation/model_comparison/README_model_comparison.md\n",
2321
+ "- artifacts/evaluation/model_comparison/benchmark_metadata.json\n",
2322
+ "- artifacts/evaluation/model_comparison/model_comparison.csv\n",
2323
+ "- artifacts/evaluation/model_comparison/model_comparison.json\n",
2324
+ "- artifacts/evaluation/model_comparison/model_comparison_per_task.csv\n"
2325
+ ]
2326
+ }
2327
+ ],
2328
+ "source": [
2329
+ "comparison_csv_path = (\n",
2330
+ " OUTPUT_DIR\n",
2331
+ " / \"model_comparison.csv\"\n",
2332
+ ")\n",
2333
+ "\n",
2334
+ "per_task_csv_path = (\n",
2335
+ " OUTPUT_DIR\n",
2336
+ " / \"model_comparison_per_task.csv\"\n",
2337
+ ")\n",
2338
+ "\n",
2339
+ "comparison_json_path = (\n",
2340
+ " OUTPUT_DIR\n",
2341
+ " / \"model_comparison.json\"\n",
2342
+ ")\n",
2343
+ "\n",
2344
+ "metadata_path = (\n",
2345
+ " OUTPUT_DIR\n",
2346
+ " / \"benchmark_metadata.json\"\n",
2347
+ ")\n",
2348
+ "\n",
2349
+ "readme_table_path = (\n",
2350
+ " OUTPUT_DIR\n",
2351
+ " / \"README_model_comparison.md\"\n",
2352
+ ")\n",
2353
+ "\n",
2354
+ "comparison_df.to_csv(\n",
2355
+ " comparison_csv_path,\n",
2356
+ " index=False,\n",
2357
+ ")\n",
2358
+ "\n",
2359
+ "per_task_df.to_csv(\n",
2360
+ " per_task_csv_path,\n",
2361
+ " index=False,\n",
2362
+ ")\n",
2363
+ "\n",
2364
+ "\n",
2365
+ "comparison_payload = {\n",
2366
+ " row[\"Model\"]: {\n",
2367
+ " \"average_accuracy\": float(\n",
2368
+ " row[\"Avg Accuracy\"]\n",
2369
+ " ),\n",
2370
+ " \"average_macro_f1\": float(\n",
2371
+ " row[\"Avg Macro F1\"]\n",
2372
+ " ),\n",
2373
+ " \"average_top3_accuracy\": float(\n",
2374
+ " row[\"Top-3 Accuracy\"]\n",
2375
+ " ),\n",
2376
+ " \"exact_match_accuracy\": float(\n",
2377
+ " row[\"Exact Match\"]\n",
2378
+ " ),\n",
2379
+ " \"latency_ms\": float(\n",
2380
+ " row[\"Latency (ms)\"]\n",
2381
+ " ),\n",
2382
+ " }\n",
2383
+ " for row in comparison_rows\n",
2384
+ "}\n",
2385
+ "\n",
2386
+ "with open(\n",
2387
+ " comparison_json_path,\n",
2388
+ " \"w\",\n",
2389
+ " encoding=\"utf-8\",\n",
2390
+ ") as file:\n",
2391
+ " json.dump(\n",
2392
+ " comparison_payload,\n",
2393
+ " file,\n",
2394
+ " indent=2,\n",
2395
+ " ensure_ascii=False,\n",
2396
+ " )\n",
2397
+ "\n",
2398
+ "\n",
2399
+ "benchmark_metadata = {\n",
2400
+ " \"created_at_utc\": datetime.now(\n",
2401
+ " timezone.utc\n",
2402
+ " ).isoformat(),\n",
2403
+ " \"dataset_name\": DATASET_NAME,\n",
2404
+ " \"dataset_fingerprint\": (\n",
2405
+ " clean_dataset._fingerprint\n",
2406
+ " ),\n",
2407
+ " \"v1_repo_id\": V1_REPO_ID,\n",
2408
+ " \"v2_repo_id\": V2_REPO_ID,\n",
2409
+ " \"seed\": SEED,\n",
2410
+ " \"train_samples\": int(\n",
2411
+ " len(train_df)\n",
2412
+ " ),\n",
2413
+ " \"validation_samples\": int(\n",
2414
+ " len(val_df)\n",
2415
+ " ),\n",
2416
+ " \"test_samples\": int(\n",
2417
+ " len(test_df)\n",
2418
+ " ),\n",
2419
+ " \"test_split_sha256\": (\n",
2420
+ " test_split_sha256\n",
2421
+ " ),\n",
2422
+ " \"tasks\": TASKS,\n",
2423
+ " \"metric_policy\": (\n",
2424
+ " \"Raw model predictions for all rows\"\n",
2425
+ " ),\n",
2426
+ " \"latency_policy\": {\n",
2427
+ " \"batch_size\": 1,\n",
2428
+ " \"model_forward_only\": True,\n",
2429
+ " \"preprocessing_excluded\": True,\n",
2430
+ " \"warmup_runs\": (\n",
2431
+ " LATENCY_WARMUP_RUNS\n",
2432
+ " ),\n",
2433
+ " \"measured_runs\": (\n",
2434
+ " LATENCY_MEASURED_RUNS\n",
2435
+ " ),\n",
2436
+ " },\n",
2437
+ " \"environment\": {\n",
2438
+ " \"python\": platform.python_version(),\n",
2439
+ " \"torch\": torch.__version__,\n",
2440
+ " \"transformers\": (\n",
2441
+ " transformers.__version__\n",
2442
+ " ),\n",
2443
+ " \"device\": str(DEVICE),\n",
2444
+ " \"gpu\": (\n",
2445
+ " torch.cuda.get_device_name(0)\n",
2446
+ " if torch.cuda.is_available()\n",
2447
+ " else None\n",
2448
+ " ),\n",
2449
+ " \"cuda_runtime\": torch.version.cuda,\n",
2450
+ " },\n",
2451
+ "}\n",
2452
+ "\n",
2453
+ "with open(\n",
2454
+ " metadata_path,\n",
2455
+ " \"w\",\n",
2456
+ " encoding=\"utf-8\",\n",
2457
+ ") as file:\n",
2458
+ " json.dump(\n",
2459
+ " benchmark_metadata,\n",
2460
+ " file,\n",
2461
+ " indent=2,\n",
2462
+ " ensure_ascii=False,\n",
2463
+ " )\n",
2464
+ "\n",
2465
+ "\n",
2466
+ "def markdown_percent(value):\n",
2467
+ " return f\"{value * 100:.2f}%\"\n",
2468
+ "\n",
2469
+ "\n",
2470
+ "markdown_lines = [\n",
2471
+ " \"| Model | Avg Accuracy | Avg Macro F1 | Top-3 Accuracy | Exact Match | Latency |\",\n",
2472
+ " \"|---|---:|---:|---:|---:|---:|\",\n",
2473
+ "]\n",
2474
+ "\n",
2475
+ "for row in comparison_rows:\n",
2476
+ " latency_text = (\n",
2477
+ " f\"{row['Latency (ms)']:.3f} ms\"\n",
2478
+ " )\n",
2479
+ "\n",
2480
+ " markdown_lines.append(\n",
2481
+ " \"| \"\n",
2482
+ " + \" | \".join(\n",
2483
+ " [\n",
2484
+ " row[\"Model\"],\n",
2485
+ " markdown_percent(\n",
2486
+ " row[\"Avg Accuracy\"]\n",
2487
+ " ),\n",
2488
+ " markdown_percent(\n",
2489
+ " row[\"Avg Macro F1\"]\n",
2490
+ " ),\n",
2491
+ " markdown_percent(\n",
2492
+ " row[\"Top-3 Accuracy\"]\n",
2493
+ " ),\n",
2494
+ " markdown_percent(\n",
2495
+ " row[\"Exact Match\"]\n",
2496
+ " ),\n",
2497
+ " latency_text,\n",
2498
+ " ]\n",
2499
+ " )\n",
2500
+ " + \" |\"\n",
2501
+ " )\n",
2502
+ "\n",
2503
+ "markdown_lines.extend(\n",
2504
+ " [\n",
2505
+ " \"\",\n",
2506
+ " (\n",
2507
+ " \"> All models were evaluated on the same \"\n",
2508
+ " f\"{len(test_df):,}-image held-out test split. \"\n",
2509
+ " \"Metrics use raw model predictions. \"\n",
2510
+ " \"Latency is batch-size-1 model-forward time \"\n",
2511
+ " \"with preprocessing excluded.\"\n",
2512
+ " ),\n",
2513
+ " ]\n",
2514
+ ")\n",
2515
+ "\n",
2516
+ "readme_markdown = \"\\n\".join(\n",
2517
+ " markdown_lines\n",
2518
+ ")\n",
2519
+ "\n",
2520
+ "readme_table_path.write_text(\n",
2521
+ " readme_markdown,\n",
2522
+ " encoding=\"utf-8\",\n",
2523
+ ")\n",
2524
+ "\n",
2525
+ "print(readme_markdown)\n",
2526
+ "\n",
2527
+ "print(\"\\nSaved files:\")\n",
2528
+ "\n",
2529
+ "for path in sorted(\n",
2530
+ " OUTPUT_DIR.iterdir()\n",
2531
+ "):\n",
2532
+ " print(\"-\", path)\n"
2533
+ ]
2534
+ }
2535
+ ],
2536
+ "metadata": {
2537
+ "kernelspec": {
2538
+ "display_name": "Python 3",
2539
+ "language": "python",
2540
+ "name": "python3"
2541
+ },
2542
+ "language_info": {
2543
+ "codemirror_mode": {
2544
+ "name": "ipython",
2545
+ "version": 3
2546
+ },
2547
+ "file_extension": ".py",
2548
+ "mimetype": "text/x-python",
2549
+ "name": "python",
2550
+ "nbconvert_exporter": "python",
2551
+ "pygments_lexer": "ipython3",
2552
+ "version": "3.12.13"
2553
+ }
2554
+ },
2555
+ "nbformat": 4,
2556
+ "nbformat_minor": 5
2557
+ }