Spaces:
Runtime error
Runtime error
| 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() | |