Spaces:
Runtime error
Runtime error
File size: 4,986 Bytes
18870c4 580017f 18870c4 956b83d 18870c4 580017f 18870c4 580017f 18870c4 580017f 18870c4 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 | 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()
|