sport-classification / train_model.py
ochsncon's picture
Update train_model.py
580017f verified
Raw
History Blame Contribute Delete
4.99 kB
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()