CoolFace
Apppublic

AD-Styles/mini-llava-demo

sourceHugging Facemitupdated 5mo agoView on Hugging Face
0likes
train.py175 linesDownload Raw Back to src
1"""Stage 1 학습 — projector만 학습하여 시각 특징을 LLM 임베딩 공간으로 정렬.2 3사용 예:4  python -m src.train \\5    --data-path data/coco_subset/manifest.json \\6    --output-dir checkpoints/stage1 \\7    --batch-size 8 --epochs 1 --lr 1e-38"""9from __future__ import annotations10 11import argparse12import math13import os14import random15 16import torch17from torch.optim import AdamW18from torch.optim.lr_scheduler import LambdaLR19from torch.utils.data import DataLoader20from tqdm import tqdm21 22from .config import TrainConfig23from .dataset import VQACollator, VQADataset24from .model import MiniLLaVA25 26 27def set_seed(seed: int):28    random.seed(seed)29    torch.manual_seed(seed)30    if torch.cuda.is_available():31        torch.cuda.manual_seed_all(seed)32 33 34def cosine_lr_lambda(total_steps: int, warmup_steps: int):35    def fn(step: int):36        if step < warmup_steps:37            return step / max(1, warmup_steps)38        progress = (step - warmup_steps) / max(1, total_steps - warmup_steps)39        return 0.5 * (1.0 + math.cos(math.pi * progress))40 41    return fn42 43 44def maybe_apply_lora(model: MiniLLaVA, cfg: TrainConfig):45    """Stage 2: 기존 projector는 그대로 학습 가능 + LLM에 LoRA 어댑터 추가."""46    if not cfg.use_lora:47        return model48    from peft import LoraConfig, get_peft_model49 50    lora_cfg = LoraConfig(51        r=cfg.lora_r,52        lora_alpha=cfg.lora_alpha,53        lora_dropout=cfg.lora_dropout,54        target_modules=["q_proj", "k_proj", "v_proj", "o_proj"],55        task_type="CAUSAL_LM",56    )57    model.llm = get_peft_model(model.llm, lora_cfg)58    # PEFT 가 base LLM을 자동 freeze. projector는 외부라 trainable 유지.59    return model60 61 62def parse_args() -> TrainConfig:63    p = argparse.ArgumentParser()64    p.add_argument("--data-path", type=str, required=True)65    p.add_argument("--output-dir", type=str, default="checkpoints/stage1")66    p.add_argument("--batch-size", type=int, default=8)67    p.add_argument("--grad-accum-steps", type=int, default=1)68    p.add_argument("--epochs", type=int, default=1)69    p.add_argument("--lr", type=float, default=1e-3)70    p.add_argument("--weight-decay", type=float, default=0.0)71    p.add_argument("--warmup-ratio", type=float, default=0.03)72    p.add_argument("--max-text-length", type=int, default=512)73    p.add_argument("--log-every", type=int, default=20)74    p.add_argument("--save-every", type=int, default=500)75    p.add_argument("--seed", type=int, default=42)76    p.add_argument("--use-lora", action="store_true",77                   help="Stage 2: LoRA adapter on LLM + projector 동시 학습")78    p.add_argument("--lora-r", type=int, default=16)79    p.add_argument("--lora-alpha", type=int, default=32)80    p.add_argument("--lora-dropout", type=float, default=0.05)81    p.add_argument("--init-projector", type=str, default=None,82                   help="기존 projector ckpt에서 시작 (Stage 1 → Stage 2 이어 학습)")83    args = p.parse_args()84    return TrainConfig(**vars(args))85 86 87def main():88    cfg = parse_args()89    set_seed(cfg.seed)90    os.makedirs(cfg.output_dir, exist_ok=True)91 92    device = torch.device("cuda" if torch.cuda.is_available() else "cpu")93    print(f"[device] {device}")94 95    print("[init] loading MiniLLaVA ...")96    model = MiniLLaVA(freeze_vision=True, freeze_llm=not cfg.use_lora)97    if cfg.init_projector and os.path.exists(cfg.init_projector):98        print(f"[init] loading existing projector → {cfg.init_projector}")99        model.load_projector(cfg.init_projector, map_location="cpu")100    model = maybe_apply_lora(model, cfg)101    model.to(device)102    print(f"[init] trainable params: {model.num_trainable():,}")103 104    print(f"[data] loading {cfg.data_path}")105    dataset = VQADataset(106        cfg.data_path, model.tokenizer, model.image_processor, cfg.max_text_length107    )108    collator = VQACollator(pad_token_id=model.tokenizer.pad_token_id)109    loader = DataLoader(110        dataset,111        batch_size=cfg.batch_size,112        shuffle=True,113        num_workers=2,114        pin_memory=True,115        collate_fn=collator,116    )117    print(f"[data] {len(dataset)} samples, {len(loader)} batches/epoch")118 119    optimizer = AdamW(120        model.trainable_parameters(), lr=cfg.lr, weight_decay=cfg.weight_decay121    )122 123    total_steps = (len(loader) // cfg.grad_accum_steps) * cfg.epochs124    warmup_steps = int(total_steps * cfg.warmup_ratio)125    scheduler = LambdaLR(optimizer, cosine_lr_lambda(total_steps, warmup_steps))126 127    global_step = 0128    model.train()129    if hasattr(model, "vision"):130        model.vision.eval()131 132    for epoch in range(cfg.epochs):133        pbar = tqdm(loader, desc=f"epoch {epoch + 1}/{cfg.epochs}")134        running_loss = 0.0135        for step, batch in enumerate(pbar):136            batch = {k: v.to(device, non_blocking=True) for k, v in batch.items()}137 138            outputs = model(**batch)139            loss = outputs.loss / cfg.grad_accum_steps140            loss.backward()141            running_loss += loss.item() * cfg.grad_accum_steps142 143            if (step + 1) % cfg.grad_accum_steps == 0:144                torch.nn.utils.clip_grad_norm_(model.trainable_parameters(), 1.0)145                optimizer.step()146                scheduler.step()147                optimizer.zero_grad(set_to_none=True)148                global_step += 1149 150                if global_step % cfg.log_every == 0:151                    avg = running_loss / (cfg.log_every * cfg.grad_accum_steps)152                    pbar.set_postfix(153                        loss=f"{avg:.4f}", lr=f"{scheduler.get_last_lr()[0]:.2e}"154                    )155                    running_loss = 0.0156 157                if global_step % cfg.save_every == 0:158                    ckpt = os.path.join(159                        cfg.output_dir, f"projector_step{global_step}.pt"160                    )161                    model.save_projector(ckpt)162 163    final_path = os.path.join(cfg.output_dir, "projector.pt")164    model.save_projector(final_path)165    print(f"[done] saved → {final_path}")166 167    if cfg.use_lora:168        lora_dir = os.path.join(cfg.output_dir, "lora_adapter")169        model.llm.save_pretrained(lora_dir)170        print(f"[done] saved LoRA → {lora_dir}")171 172 173if __name__ == "__main__":174    main()175