"""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()