Setting the file. One moment.
Image Classification Training · Huggingface Vision Trainer · huggingface/skills · Skills Docs
ContentsBack to the top of the page Script
scripts/ image_classification_training.py
Python · 383 lines · 13 KB
17
import
logging
18 import os
19 import sys
20 from dataclasses import dataclass, field
21 from functools import partial
22 from typing import Any
23
24 import evaluate
25 import numpy as np
26 import torch
27 from datasets import load_dataset
28 from torchvision.transforms import (
29 CenterCrop,
30 Compose,
31 Normalize,
32 RandomHorizontalFlip,
33 RandomResizedCrop,
34 Resize,
35 ToTensor,
36 )
37
38 import trackio
39
40 import transformers
41 from transformers import (
42 AutoConfig,
43 AutoImageProcessor,
44 AutoModelForImageClassification,
45 DefaultDataCollator,
46 HfArgumentParser,
47 Trainer,
48 TrainingArguments,
49 )
50 from transformers.trainer import EvalPrediction
51 from transformers.utils import check_min_version
52 from transformers.utils.versions import require_version
53
54
55 logger = logging.getLogger( __name__ )
56
57 check_min_version( "4.57.0.dev0" )
58 require_version( "datasets>=2.0.0" )
59
60
61 @dataclass
62 class DataTrainingArguments :
63 dataset_name: str = field(
64 default = "ethz/food101" ,
65 metadata = { "help" : "Name of a dataset from the Hub." },
66 )
67 dataset_config_name: str | None = field(
68 default = None ,
69 metadata = { "help" : "The configuration name of the dataset to use (via the datasets library)." },
70 )
71 train_val_split: float | None = field(
72 default = 0.15 ,
73 metadata = { "help" : "Fraction to split off of train for validation (used only when no validation split exists)." },
74 )
75 max_train_samples: int | None = field(
76 default = None ,
77 metadata = { "help" : "Truncate training set to this many samples (for debugging / quick tests)." },
78 )
79 max_eval_samples: int | None = field(
80 default = None ,
81 metadata = { "help" : "Truncate evaluation set to this many samples." },
82 )
83 image_column_name: str = field(
84 default = "image" ,
85 metadata = { "help" : "The column name for images in the dataset." },
86 )
87 label_column_name: str = field(
88 default = "label" ,
89 metadata = { "help" : "The column name for labels in the dataset." },
90 )
91
92
93 @dataclass
94 class ModelArguments :
95 model_name_or_path: str = field(
96 default = "timm/mobilenetv3_small_100.lamb_in1k" ,
97 metadata = { "help" : "Path to pretrained model or model identifier from huggingface.co/models." },
98 )
99 config_name: str | None = field(
100 default = None ,
101 metadata = { "help" : "Pretrained config name or path if not the same as model_name." },
102 )
103 cache_dir: str | None = field(
104 default = None ,
105 metadata = { "help" : "Where to store pretrained models downloaded from the Hub." },
106 )
107 model_revision: str = field(
108 default = "main" ,
109 metadata = { "help" : "The specific model version to use (branch, tag, or commit id)." },
110 )
111 image_processor_name: str | None = field(
112 default = None ,
113 metadata = { "help" : "Name or path of image processor config." },
114 )
115 ignore_mismatched_sizes: bool = field(
116 default = True ,
117 metadata = { "help" : "Allow loading weights when num_labels differs from pretrained checkpoint." },
118 )
119 token: str | None = field(
120 default = None ,
121 metadata = { "help" : "Auth token for private models / datasets." },
122 )
123 trust_remote_code: bool = field(
124 default = False ,
125 metadata = { "help" : "Whether to trust remote code from Hub repos." },
126 )
127
128
129 def build_transforms (image_processor, is_training: bool ):
130 """Build torchvision transforms from the image processor's config."""
131 if hasattr (image_processor, "size" ):
132 size = image_processor.size
133 if "shortest_edge" in size:
134 img_size = size[ "shortest_edge" ]
135 elif "height" in size and "width" in size:
136 img_size = (size[ "height" ], size[ "width" ])
137 else :
138 img_size = 224
139 else :
140 img_size = 224
141
142 if hasattr (image_processor, "image_mean" ) and image_processor.image_mean:
143 normalize = Normalize( mean = image_processor.image_mean, std = image_processor.image_std)
144 else :
145 normalize = Normalize( mean = [ 0.485 , 0.456 , 0.406 ], std = [ 0.229 , 0.224 , 0.225 ])
146
147 if is_training:
148 return Compose([
149 RandomResizedCrop(img_size),
150 RandomHorizontalFlip(),
151 ToTensor(),
152 normalize,
153 ])
154 else :
155 if isinstance (img_size, int ):
156 resize_size = int (img_size / 0.875 ) # standard 87.5% center crop ratio
157 else :
158 resize_size = tuple ( int (s / 0.875 ) for s in img_size)
159 return Compose([
160 Resize(resize_size),
161 CenterCrop(img_size),
162 ToTensor(),
163 normalize,
164 ])
165
166
167 def main ():
168 parser = HfArgumentParser((ModelArguments, DataTrainingArguments, TrainingArguments))
169 if len (sys.argv) == 2 and sys.argv[ 1 ].endswith( ".json" ):
170 model_args, data_args, training_args = parser.parse_json_file( json_file = os.path.abspath(sys.argv[ 1 ]))
171 else :
172 model_args, data_args, training_args = parser.parse_args_into_dataclasses()
173
174 # --- Hub authentication ---
175 from huggingface_hub import login
176 hf_token = os.environ.get( "HF_TOKEN" ) or os.environ.get( "hfjob" )
177 if hf_token:
178 login( token = hf_token)
179 training_args.hub_token = hf_token
180 logger.info( "Logged in to Hugging Face Hub" )
181 elif training_args.push_to_hub:
182 logger.warning( "HF_TOKEN not found in environment. Hub push will likely fail." )
183
184 # --- Trackio ---
185 trackio.init( project = training_args.output_dir, name = training_args.run_name)
186
187 # --- Logging ---
188 logging.basicConfig(
189 format = " %(asctime)s - %(levelname)s - %(name)s - %(message)s " ,
190 datefmt = "%m/ %d /%Y %H:%M:%S" ,
191 handlers = [logging.StreamHandler(sys.stdout)],
192 )
193 if training_args.should_log:
194 transformers.utils.logging.set_verbosity_info()
195
196 log_level = training_args.get_process_log_level()
197 logger.setLevel(log_level)
198 transformers.utils.logging.set_verbosity(log_level)
199 transformers.utils.logging.enable_default_handler()
200 transformers.utils.logging.enable_explicit_format()
201
202 logger.warning(
203 f "Process rank: { training_args.local_process_index } , device: { training_args.device } , "
204 f "n_gpu: { training_args.n_gpu } , distributed training: "
205 f " { training_args.parallel_mode.value == 'distributed' } , 16-bits training: { training_args.fp16 } "
206 )
207 logger.info( f "Training/evaluation parameters { training_args } " )
208
209 # --- Load dataset ---
210 dataset = load_dataset(
211 data_args.dataset_name,
212 data_args.dataset_config_name,
213 cache_dir = model_args.cache_dir,
214 trust_remote_code = model_args.trust_remote_code,
215 )
216
217 # --- Resolve label column ---
218 label_col = data_args.label_column_name
219 if label_col not in dataset[ "train" ].column_names:
220 candidates = [c for c in dataset[ "train" ].column_names if c in ( "label" , "labels" , "class" , "fine_label" )]
221 if candidates:
222 label_col = candidates[ 0 ]
223 logger.info( f "Label column ' { data_args.label_column_name } ' not found, using ' { label_col } '" )
224 else :
225 raise ValueError (
226 f "Label column ' { data_args.label_column_name } ' not found. "
227 f "Available columns: { dataset[ 'train' ].column_names } "
228 )
229
230 # --- Discover labels ---
231 label_feature = dataset[ "train" ].features[label_col]
232 if hasattr (label_feature, "names" ):
233 label_names = label_feature.names
234 else :
235 unique_labels = sorted ( set (dataset[ "train" ][label_col]))
236 if all ( isinstance (l, str ) for l in unique_labels):
237 label_names = unique_labels
238 else :
239 label_names = [ str (l) for l in unique_labels]
240
241 num_labels = len (label_names)
242 id2label = dict ( enumerate (label_names))
243 label2id = {v: k for k, v in id2label.items()}
244 logger.info( f "Number of classes: { num_labels } " )
245
246 # --- Remap string labels to int if needed ---
247 sample_label = dataset[ "train" ][ 0 ][label_col]
248 if isinstance (sample_label, str ):
249 logger.info( "Remapping string labels to integer IDs" )
250 for split_name in list (dataset.keys()):
251 dataset[split_name] = dataset[split_name].map(
252 lambda ex: {label_col: label2id[ex[label_col]]},
253 )
254
255 # --- Shuffle + Train/val split ---
256 dataset[ "train" ] = dataset[ "train" ].shuffle( seed = training_args.seed)
257
258 data_args.train_val_split = None if "validation" in dataset else data_args.train_val_split
259 if isinstance (data_args.train_val_split, float ) and data_args.train_val_split > 0.0 :
260 split = dataset[ "train" ].train_test_split(data_args.train_val_split, seed = training_args.seed)
261 dataset[ "train" ] = split[ "train" ]
262 dataset[ "validation" ] = split[ "test" ]
263
264 # --- Truncate ---
265 if data_args.max_train_samples is not None :
266 max_train = min (data_args.max_train_samples, len (dataset[ "train" ]))
267 dataset[ "train" ] = dataset[ "train" ].select( range (max_train))
268 logger.info( f "Truncated training set to { max_train } samples" )
269 if data_args.max_eval_samples is not None and "validation" in dataset:
270 max_eval = min (data_args.max_eval_samples, len (dataset[ "validation" ]))
271 dataset[ "validation" ] = dataset[ "validation" ].select( range (max_eval))
272 logger.info( f "Truncated validation set to { max_eval } samples" )
273
274 # --- Load model & image processor ---
275 common_pretrained_args = {
276 "cache_dir" : model_args.cache_dir,
277 "revision" : model_args.model_revision,
278 "token" : model_args.token,
279 "trust_remote_code" : model_args.trust_remote_code,
280 }
281
282 config = AutoConfig.from_pretrained(
283 model_args.config_name or model_args.model_name_or_path,
284 num_labels = num_labels,
285 label2id = label2id,
286 id2label = id2label,
287 ** common_pretrained_args,
288 )
289
290 model = AutoModelForImageClassification.from_pretrained(
291 model_args.model_name_or_path,
292 config = config,
293 ignore_mismatched_sizes = model_args.ignore_mismatched_sizes,
294 ** common_pretrained_args,
295 )
296
297 image_processor = AutoImageProcessor.from_pretrained(
298 model_args.image_processor_name or model_args.model_name_or_path,
299 ** common_pretrained_args,
300 )
301
302 # --- Build transforms ---
303 train_transforms = build_transforms(image_processor, is_training = True )
304 val_transforms = build_transforms(image_processor, is_training = False )
305
306 image_col = data_args.image_column_name
307
308 def preprocess_train (examples):
309 return {
310 "pixel_values" : [train_transforms(img.convert( "RGB" )) for img in examples[image_col]],
311 "labels" : examples[label_col],
312 }
313
314 def preprocess_val (examples):
315 return {
316 "pixel_values" : [val_transforms(img.convert( "RGB" )) for img in examples[image_col]],
317 "labels" : examples[label_col],
318 }
319
320 dataset[ "train" ].set_transform(preprocess_train)
321 if "validation" in dataset:
322 dataset[ "validation" ].set_transform(preprocess_val)
323 if "test" in dataset:
324 dataset[ "test" ].set_transform(preprocess_val)
325
326 # --- Metrics ---
327 accuracy_metric = evaluate.load( "accuracy" )
328
329 def compute_metrics (eval_pred: EvalPrediction):
330 predictions = np.argmax(eval_pred.predictions, axis = 1 )
331 return accuracy_metric.compute( predictions = predictions, references = eval_pred.label_ids)
332
333 # --- Trainer ---
334 eval_dataset = None
335 if training_args.do_eval:
336 if "validation" in dataset:
337 eval_dataset = dataset[ "validation" ]
338 elif "test" in dataset:
339 eval_dataset = dataset[ "test" ]
340
341 trainer = Trainer(
342 model = model,
343 args = training_args,
344 train_dataset = dataset[ "train" ] if training_args.do_train else None ,
345 eval_dataset = eval_dataset,
346 processing_class = image_processor,
347 data_collator = DefaultDataCollator(),
348 compute_metrics = compute_metrics,
349 )
350
351 # --- Train ---
352 if training_args.do_train:
353 train_result = trainer.train( resume_from_checkpoint = training_args.resume_from_checkpoint)
354 trainer.save_model()
355 trainer.log_metrics( "train" , train_result.metrics)
356 trainer.save_metrics( "train" , train_result.metrics)
357 trainer.save_state()
358
359 # --- Evaluate ---
360 if training_args.do_eval:
361 test_dataset = dataset.get( "test" , dataset.get( "validation" ))
362 test_prefix = "test" if "test" in dataset else "eval"
363 if test_dataset is not None :
364 metrics = trainer.evaluate( eval_dataset = test_dataset, metric_key_prefix = test_prefix)
365 trainer.log_metrics(test_prefix, metrics)
366 trainer.save_metrics(test_prefix, metrics)
367
368 trackio.finish()
369
370 # --- Push to Hub ---
371 kwargs = {
372 "finetuned_from" : model_args.model_name_or_path,
373 "dataset" : data_args.dataset_name,
374 "tags" : [ "image-classification" , "vision" ],
375 }
376 if training_args.push_to_hub:
377 trainer.push_to_hub( ** kwargs)
378 else :
379 trainer.create_model_card( ** kwargs)
380
381
382 if __name__ == "__main__" :
383 main()