import argparse from pathlib import Path import evaluate import numpy as np from datasets import load_dataset from transformers import ( AutoImageProcessor, AutoModelForImageClassification, Trainer, TrainingArguments, ) def parse_args(): parser = argparse.ArgumentParser(description="Train a transfer-learning sports classifier.") parser.add_argument("--data_dir", type=str, default="data/sports", help="Folder with class subfolders.") parser.add_argument("--base_model", type=str, default="microsoft/resnet-18") parser.add_argument("--output_dir", type=str, default="sports-vit-transfer") parser.add_argument("--epochs", type=int, default=5) parser.add_argument("--batch_size", type=int, default=16) parser.add_argument("--learning_rate", type=float, default=5e-5) parser.add_argument("--validation_split", type=float, default=0.2) parser.add_argument("--push_to_hub", action="store_true") parser.add_argument("--hub_model_id", type=str, default="") return parser.parse_args() def build_transforms(processor): image_mean = processor.image_mean image_std = processor.image_std size_cfg = processor.size if isinstance(size_cfg, dict): size = size_cfg.get("shortest_edge") or size_cfg.get("height") or size_cfg.get("width") or 224 else: size = int(size_cfg) if size_cfg else 224 from torchvision.transforms import ( CenterCrop, Compose, Normalize, RandomHorizontalFlip, RandomResizedCrop, Resize, ToTensor, ) train_tfm = Compose( [ RandomResizedCrop(size), RandomHorizontalFlip(), ToTensor(), Normalize(mean=image_mean, std=image_std), ] ) val_tfm = Compose( [ Resize(size), CenterCrop(size), ToTensor(), Normalize(mean=image_mean, std=image_std), ] ) return train_tfm, val_tfm def main(): args = parse_args() data_dir = Path(args.data_dir) if not data_dir.exists(): raise FileNotFoundError(f"Dataset folder not found: {data_dir}") ds = load_dataset("imagefolder", data_dir=str(data_dir)) if "validation" not in ds: if "test" in ds: ds["validation"] = ds["test"] else: split = ds["train"].train_test_split(test_size=args.validation_split, seed=42) ds["train"] = split["train"] ds["validation"] = split["test"] labels = ds["train"].features["label"].names label2id = {label: i for i, label in enumerate(labels)} id2label = {i: label for i, label in enumerate(labels)} processor = AutoImageProcessor.from_pretrained(args.base_model) model = AutoModelForImageClassification.from_pretrained( args.base_model, num_labels=len(labels), id2label=id2label, label2id=label2id, ignore_mismatched_sizes=True, ) train_tfm, val_tfm = build_transforms(processor) def transform_train(batch): batch["pixel_values"] = [train_tfm(img.convert("RGB")) for img in batch["image"]] return batch def transform_val(batch): batch["pixel_values"] = [val_tfm(img.convert("RGB")) for img in batch["image"]] return batch ds["train"].set_transform(transform_train) ds["validation"].set_transform(transform_val) def collate_fn(batch): import torch return { "pixel_values": torch.stack([example["pixel_values"] for example in batch]), "labels": torch.tensor([example["label"] for example in batch]), } metric = evaluate.load("accuracy") def compute_metrics(eval_pred): logits, labels_ = eval_pred predictions = np.argmax(logits, axis=1) return metric.compute(predictions=predictions, references=labels_) train_args = TrainingArguments( output_dir=args.output_dir, remove_unused_columns=False, eval_strategy="epoch", save_strategy="no", logging_strategy="steps", logging_steps=20, learning_rate=args.learning_rate, per_device_train_batch_size=args.batch_size, per_device_eval_batch_size=args.batch_size, num_train_epochs=args.epochs, load_best_model_at_end=False, push_to_hub=args.push_to_hub, hub_model_id=args.hub_model_id if args.hub_model_id else None, report_to="none", ) trainer = Trainer( model=model, args=train_args, train_dataset=ds["train"], eval_dataset=ds["validation"], data_collator=collate_fn, compute_metrics=compute_metrics, ) trainer.train() metrics = trainer.evaluate() print("Validation metrics:", metrics) trainer.save_model(args.output_dir) processor.save_pretrained(args.output_dir) if args.push_to_hub: trainer.push_to_hub() if __name__ == "__main__": main()