vit-tiny-imagenet-demo / train_lora.py
turhancan97's picture
Upload folder using huggingface_hub
10af6f1 verified
Raw
History Blame Contribute Delete
9.76 kB
"""Fine-tune a LoRA adapter on Food-101 for a ViT image classifier.
Preserves the original pretrained weights (LoRA is additive) and saves the
adapter + the new classification head as a single PEFT-format artifact.
Example:
python train_lora.py \\
--rank 8 --alpha 16 --target-modules query value \\
--epochs 5 --batch-size 64 --lr 5e-4 \\
--push-to-hub turhancan97/vit-tiny-lora-food101
"""
from __future__ import annotations
import argparse
import json
import os
from dataclasses import asdict, dataclass
from pathlib import Path
import numpy as np
import torch
from datasets import load_dataset
from peft import LoraConfig, get_peft_model
from PIL import Image
from torchvision import transforms
from transformers import (
AutoImageProcessor,
AutoModelForImageClassification,
Trainer,
TrainingArguments,
)
@dataclass
class Args:
model_id: str
dataset_id: str
output_dir: str
rank: int
alpha: int
dropout: float
target_modules: list[str]
lr: float
batch_size: int
eval_batch_size: int
epochs: int
warmup_ratio: float
weight_decay: float
seed: int
push_to_hub: str | None
max_train_samples: int | None
max_eval_samples: int | None
eval_only: bool
num_workers: int
def parse_args() -> Args:
p = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
p.add_argument("--model-id", default="WinKawaks/vit-tiny-patch16-224")
p.add_argument("--dataset-id", default="food101")
p.add_argument("--output-dir", default="adapters/vit-tiny-lora-food101")
p.add_argument("--rank", type=int, default=8)
p.add_argument("--alpha", type=int, default=16)
p.add_argument("--dropout", type=float, default=0.1)
p.add_argument(
"--target-modules", nargs="+", default=["query", "value"],
help="Substring patterns matched against module names for LoRA injection.",
)
p.add_argument("--lr", type=float, default=5e-4)
p.add_argument("--batch-size", type=int, default=64)
p.add_argument("--eval-batch-size", type=int, default=128)
p.add_argument("--epochs", type=int, default=1)
p.add_argument("--warmup-ratio", type=float, default=0.03)
p.add_argument("--weight-decay", type=float, default=0.0)
p.add_argument("--seed", type=int, default=42)
p.add_argument("--push-to-hub", default=None, help="e.g. 'user/vit-tiny-lora-food101'")
p.add_argument("--max-train-samples", type=int, default=None, help="Smoke-test subset size.")
p.add_argument("--max-eval-samples", type=int, default=None)
p.add_argument("--eval-only", action="store_true")
p.add_argument("--num-workers", type=int, default=4)
ns = p.parse_args()
return Args(**{k.replace("-", "_"): v for k, v in vars(ns).items()})
def build_transforms(processor: AutoImageProcessor):
size = processor.size.get("height") or processor.size.get("shortest_edge") or 224
mean = processor.image_mean
std = processor.image_std
train_tf = transforms.Compose([
transforms.RandomResizedCrop(size, scale=(0.8, 1.0)),
transforms.RandomHorizontalFlip(),
transforms.ToTensor(),
transforms.Normalize(mean=mean, std=std),
])
eval_tf = transforms.Compose([
transforms.Resize(int(size * 256 / 224)),
transforms.CenterCrop(size),
transforms.ToTensor(),
transforms.Normalize(mean=mean, std=std),
])
return train_tf, eval_tf
def _ensure_rgb(img):
if isinstance(img, Image.Image):
return img.convert("RGB") if img.mode != "RGB" else img
return Image.fromarray(np.asarray(img)).convert("RGB")
def make_transform_fn(tf):
def _apply(batch):
batch["pixel_values"] = [tf(_ensure_rgb(img)) for img in batch["image"]]
return batch
return _apply
def collate_fn(examples):
pixel_values = torch.stack([ex["pixel_values"] for ex in examples])
labels = torch.tensor([ex["label"] for ex in examples], dtype=torch.long)
return {"pixel_values": pixel_values, "labels": labels}
def compute_metrics_topk(eval_pred):
logits, labels = eval_pred
logits = torch.as_tensor(logits)
labels = torch.as_tensor(labels)
top1 = (logits.argmax(dim=-1) == labels).float().mean().item()
k = min(5, logits.shape[-1])
topk = logits.topk(k=k, dim=-1).indices
top5 = (topk == labels.unsqueeze(-1)).any(dim=-1).float().mean().item()
return {"top1_accuracy": top1, "top5_accuracy": top5}
def main():
args = parse_args()
torch.manual_seed(args.seed)
np.random.seed(args.seed)
output_dir = Path(args.output_dir)
output_dir.mkdir(parents=True, exist_ok=True)
print(f"[1/5] Loading dataset: {args.dataset_id}")
ds = load_dataset(args.dataset_id)
train_split = "train" if "train" in ds else list(ds.keys())[0]
eval_split = "validation" if "validation" in ds else ("test" if "test" in ds else train_split)
train_ds = ds[train_split]
eval_ds = ds[eval_split]
label_feature = train_ds.features["label"]
num_labels = label_feature.num_classes
id2label = {i: label_feature.int2str(i) for i in range(num_labels)}
label2id = {v: k for k, v in id2label.items()}
print(f" train={len(train_ds)} eval={len(eval_ds)} num_labels={num_labels}")
if args.max_train_samples:
train_ds = train_ds.shuffle(seed=args.seed).select(range(min(args.max_train_samples, len(train_ds))))
if args.max_eval_samples:
eval_ds = eval_ds.shuffle(seed=args.seed).select(range(min(args.max_eval_samples, len(eval_ds))))
print(f"[2/5] Loading base model: {args.model_id}")
processor = AutoImageProcessor.from_pretrained(args.model_id, use_fast=True)
base_model = AutoModelForImageClassification.from_pretrained(
args.model_id,
num_labels=num_labels,
id2label=id2label,
label2id=label2id,
ignore_mismatched_sizes=True,
)
train_tf, eval_tf = build_transforms(processor)
train_ds.set_transform(make_transform_fn(train_tf))
eval_ds.set_transform(make_transform_fn(eval_tf))
print(f"[3/5] Wrapping with LoRA: rank={args.rank}, alpha={args.alpha}, "
f"target_modules={args.target_modules}")
lora_cfg = LoraConfig(
r=args.rank,
lora_alpha=args.alpha,
lora_dropout=args.dropout,
target_modules=list(args.target_modules),
bias="none",
)
model = get_peft_model(base_model, lora_cfg)
# PEFT freezes every non-LoRA parameter by default. Unfreeze the classifier
# so the new task head can be trained. We save it separately after training
# (rather than via `modules_to_save`) so the adapter artifact stays portable
# across base models with different original head sizes.
classifier = model.base_model.model.classifier
for p in classifier.parameters():
p.requires_grad_(True)
trainable, total = model.get_nb_trainable_parameters()
print(f" trainable params: {trainable:,} / {total:,} ({100 * trainable / total:.2f}%)")
training_args = TrainingArguments(
output_dir=str(output_dir / "trainer"),
per_device_train_batch_size=args.batch_size,
per_device_eval_batch_size=args.eval_batch_size,
learning_rate=args.lr,
num_train_epochs=args.epochs,
warmup_ratio=args.warmup_ratio,
weight_decay=args.weight_decay,
eval_strategy="epoch",
save_strategy="epoch",
save_total_limit=1,
load_best_model_at_end=True,
metric_for_best_model="top1_accuracy",
greater_is_better=True,
logging_strategy="steps",
logging_steps=25,
fp16=torch.cuda.is_available(),
dataloader_num_workers=args.num_workers,
remove_unused_columns=False,
report_to="none",
seed=args.seed,
)
trainer = Trainer(
model=model,
args=training_args,
train_dataset=train_ds,
eval_dataset=eval_ds,
data_collator=collate_fn,
compute_metrics=compute_metrics_topk,
)
if not args.eval_only:
print("[4/5] Training")
trainer.train()
else:
print("[4/5] Skipping training (--eval-only)")
print("[5/5] Evaluating on held-out split")
metrics = trainer.evaluate()
metrics["eval_samples"] = len(eval_ds)
print(json.dumps(metrics, indent=2))
(output_dir / "eval_metrics.json").write_text(json.dumps(metrics, indent=2))
print(f"Saving adapter to {output_dir}")
model.save_pretrained(str(output_dir))
processor.save_pretrained(str(output_dir))
(output_dir / "train_args.json").write_text(json.dumps(asdict(args), indent=2))
(output_dir / "labels.json").write_text(
json.dumps({str(i): id2label[i] for i in range(num_labels)}, indent=2)
)
torch.save(
{k: v.detach().cpu() for k, v in classifier.state_dict().items()},
output_dir / "classifier.pt",
)
if args.push_to_hub:
print(f"Pushing to Hugging Face Hub: {args.push_to_hub}")
model.push_to_hub(args.push_to_hub)
processor.push_to_hub(args.push_to_hub)
try:
from huggingface_hub import HfApi
api = HfApi()
for extra in ["labels.json", "classifier.pt"]:
api.upload_file(
path_or_fileobj=str(output_dir / extra),
path_in_repo=extra,
repo_id=args.push_to_hub,
repo_type="model",
commit_message=f"add {extra}",
)
except Exception as exc:
print(f"Warning: could not upload side-car files: {exc}")
print("Done.")
if __name__ == "__main__":
main()