Spaces:
Sleeping
Sleeping
| """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, | |
| ) | |
| 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() | |