AD-Styles/mini-llava-demo
0
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 