#!/usr/bin/env python3 """Stage 2 (visual instruction tuning): LLaVA-Instruct-150K on COCO images. Starts from the stage-1 projector. Vision tower frozen; the LLM gets LoRA adapters; the projector keeps training at a lower LR. Same prompt layout as stage 1: [ user-header] [image tokens] [q1 assistant-header] a1 [user-header q2 assistant-header] a2 ... Loss only on assistant answers (+ their ). Gradient checkpointing on, length-bucketed batches, resume by samples seen. All settings via environment variables. """ import json import math import os import sys import time import zipfile import torch from PIL import Image from torch.utils.data import DataLoader, Dataset, Subset import train as T # stage-1 module: model layout, prompt strings, collate, zip check env = T.env HOME = os.path.expanduser("~") INSTRUCT_JSON = env("INSTRUCT_JSON", f"{HOME}/data/llava_instruct/llava_instruct_150k.json") COCO_ZIP = env("COCO_ZIP", f"{HOME}/data/coco/train2017.zip") STAGE1_PROJECTOR = env("STAGE1_PROJECTOR", f"{HOME}/checkpoints/stage1/projector_final.pt") CKPT_DIR = env("CKPT_DIR", f"{HOME}/checkpoints/stage2") NUM_SAMPLES = env("NUM_SAMPLES", 0, int) # 0 = all except the held-out set HELDOUT = env("HELDOUT", 1000, int) BATCH_SIZE = env("BATCH_SIZE", 8, int) GRAD_ACCUM = env("GRAD_ACCUM", 4, int) LORA_LR = env("LORA_LR", 2e-4, float) PROJ_LR = env("PROJ_LR", 2e-5, float) LORA_R = env("LORA_R", 64, int) LORA_ALPHA = env("LORA_ALPHA", 128, int) WARMUP_RATIO = env("WARMUP_RATIO", 0.03, float) MAX_TEXT_LEN = env("MAX_TEXT_LEN", 1024, int) MAX_STEPS = env("MAX_STEPS", 0, int) LOG_EVERY = env("LOG_EVERY", 25, int) SAVE_EVERY = env("SAVE_EVERY", 250, int) SEED = env("SEED", 42, int) NUM_WORKERS = T.NUM_WORKERS DEVICE, DTYPE, log = T.DEVICE, T.DTYPE, T.log LORA_TARGETS = ["q_proj", "k_proj", "v_proj", "o_proj", "gate_proj", "up_proj", "down_proj"] # ----------------------------- DATA ----------------------------- def split(ann): perm = torch.randperm(len(ann), generator=torch.Generator().manual_seed(SEED)).tolist() held, train = perm[-HELDOUT:], perm[:-HELDOUT] if NUM_SAMPLES: train = train[:NUM_SAMPLES] return train, held def build_ids(tok, conversations, max_len): ids, labels = [], [] first = True for turn in conversations: v = turn["value"].replace("", "").strip() if turn["from"] == "human": text = (v if first else T.USER_HEADER + v) + T.ASSIST_HEADER t = tok(text, add_special_tokens=False).input_ids ids += t labels += [-100] * len(t) first = False else: t = tok(v + T.EOT, add_special_tokens=False).input_ids ids += t labels += t return ids[:max_len], labels[:max_len] class InstructDataset(Dataset): def __init__(self, ann, zip_path, zip_prefix, processor, tok, max_len): self.ann, self.zip_path, self.prefix = ann, zip_path, zip_prefix self.processor, self.tok, self.max_len = processor, tok, max_len self.zip = None def __len__(self): return len(self.ann) def load_image(self, item): if self.zip is None: self.zip = zipfile.ZipFile(self.zip_path) try: with self.zip.open(self.prefix + item["image"]) as f: return Image.open(f).convert("RGB"), True except Exception: return Image.new("RGB", (224, 224), "white"), False def __getitem__(self, idx): item = self.ann[idx] image, ok = self.load_image(item) pv = self.processor(images=image, return_tensors="pt").pixel_values[0] ids, labels = build_ids(self.tok, item["conversations"], self.max_len) return {"pixel_values": pv, "input_ids": ids, "labels": labels, "ok": ok} def bucketed_order(indices, lengths, mb, seed): """Shuffle, sort by length inside chunks of 64 batches, keep only full batches, shuffle batches.""" g = torch.Generator().manual_seed(seed) perm = [indices[i] for i in torch.randperm(len(indices), generator=g).tolist()] chunk = mb * 64 batches = [] for s in range(0, len(perm), chunk): c = sorted(perm[s:s + chunk], key=lambda i: lengths[i]) batches += [c[j:j + mb] for j in range(0, len(c), mb) if len(c[j:j + mb]) == mb] order = torch.randperm(len(batches), generator=g).tolist() return [i for b in order for i in batches[b]] # ----------------------------- MODEL ----------------------------- def build_stage2_model(lora_state=None, projector_state=None, train_mode=True): from peft import LoraConfig, get_peft_model, set_peft_model_state_dict model, tok, pad_id, processor = T.build_model() # frozen vision + frozen LLM (bf16), fresh projector proj = projector_state if projector_state is not None else _load_proj(STAGE1_PROJECTOR) model.projector.load_state_dict(proj) model.projector.to(DEVICE, dtype=torch.float32) cfg = LoraConfig(r=LORA_R, lora_alpha=LORA_ALPHA, lora_dropout=0.05, target_modules=LORA_TARGETS, bias="none", task_type="CAUSAL_LM") model.llm = get_peft_model(model.llm, cfg) if lora_state is not None: set_peft_model_state_dict(model.llm, lora_state) for n, p in model.llm.named_parameters(): # LoRA weights in fp32 for stable AdamW updates if p.requires_grad: p.data = p.data.float() if train_mode: model.llm.base_model.model.gradient_checkpointing_enable(gradient_checkpointing_kwargs={"use_reentrant": False}) model.llm.train() else: model.llm.eval() model.vision.eval() return model, tok, pad_id, processor def _load_proj(path): sd = torch.load(path, map_location="cpu") return sd.get("projector_state_dict", sd) def lr_scale(step, total, warmup): if step < warmup: return (step + 1) / warmup progress = (step - warmup) / max(1, total - warmup) return 0.5 * (1 + math.cos(math.pi * min(1.0, progress))) def save_checkpoint(path, model, opt, samples_seen, step): from peft import get_peft_model_state_dict tmp = path + ".tmp" torch.save({ "lora_state_dict": get_peft_model_state_dict(model.llm), "projector_state_dict": model.projector.state_dict(), "optimizer_state_dict": opt.state_dict(), "samples_seen": samples_seen, "opt_step": step, "config": {"LORA_R": LORA_R, "LORA_ALPHA": LORA_ALPHA, "LORA_LR": LORA_LR, "PROJ_LR": PROJ_LR, "BATCH_SIZE": BATCH_SIZE, "GRAD_ACCUM": GRAD_ACCUM, "SEED": SEED}, }, tmp) os.replace(tmp, path) # ----------------------------- TRAIN ----------------------------- def main(): torch.manual_seed(SEED) os.makedirs(CKPT_DIR, exist_ok=True) ckpt_path = os.path.join(CKPT_DIR, "latest.pt") log(f"stage2: batch={BATCH_SIZE}x{GRAD_ACCUM} lora_r={LORA_R} lora_lr={LORA_LR} proj_lr={PROJ_LR} " f"max_text_len={MAX_TEXT_LEN} workers={NUM_WORKERS} max_steps={MAX_STEPS or 'full'}") ck = torch.load(ckpt_path, map_location="cpu") if os.path.exists(ckpt_path) else None init = None if ck is None and os.environ.get("INIT_FROM"): # continue from an earlier stage-2 run, fresh optimizer init = torch.load(os.environ["INIT_FROM"], map_location="cpu") log(f"initialising LoRA + projector from {os.environ['INIT_FROM']}") src = ck or init model, tok, pad_id, processor = build_stage2_model( lora_state=src["lora_state_dict"] if src else None, projector_state=src["projector_state_dict"] if src else None) del init lora_params = [p for n, p in model.llm.named_parameters() if p.requires_grad] proj_params = list(model.projector.parameters()) log(f"trainable: LoRA {sum(p.numel() for p in lora_params):,} + projector {sum(p.numel() for p in proj_params):,}") ann = json.load(open(INSTRUCT_JSON)) train_idx, _ = split(ann) prefix = T.find_zip_prefix([ann[i] for i in train_idx[:2000]], COCO_ZIP) dataset = InstructDataset(ann, COCO_ZIP, prefix, processor, tok, MAX_TEXT_LEN) lengths = [sum(len(t["value"]) for t in a["conversations"]) for a in ann] eff_bs = BATCH_SIZE * GRAD_ACCUM order = bucketed_order(train_idx, lengths, BATCH_SIZE, SEED + 1) total_steps = len(order) // eff_bs if MAX_STEPS: total_steps = min(total_steps, MAX_STEPS) warmup = max(1, int(total_steps * WARMUP_RATIO)) opt = torch.optim.AdamW([{"params": lora_params, "lr": LORA_LR, "base_lr": LORA_LR}, {"params": proj_params, "lr": PROJ_LR, "base_lr": PROJ_LR}], weight_decay=0.0) samples_seen, step = 0, 0 if ck: opt.load_state_dict(ck["optimizer_state_dict"]) samples_seen = ck["samples_seen"] step = samples_seen // eff_bs log(f"resumed: {samples_seen:,} samples -> step {step}/{total_steps}") del ck if step >= total_steps: log("already complete") return loader = DataLoader( Subset(dataset, order[samples_seen:]), batch_size=BATCH_SIZE, shuffle=False, num_workers=NUM_WORKERS, pin_memory=True, collate_fn=T.make_collate(pad_id), drop_last=True, persistent_workers=NUM_WORKERS > 0, prefetch_factor=4 if NUM_WORKERS > 0 else None) log(f"schedule: {total_steps} optimizer steps ({len(order):,} samples), {warmup} warmup, starting at {step}") params = lora_params + proj_params if DEVICE.type == "cuda": torch.cuda.reset_peak_memory_stats() opt.zero_grad(set_to_none=True) micro, loss_sum, loss_n, bad, skipped = 0, 0.0, 0, 0, 0 last_loss, t_window, steps_window, t_start = None, time.time(), 0, time.time() log_f = open(os.path.join(CKPT_DIR, "train_log.jsonl"), "a") for batch in loader: if step >= total_steps: break try: bad += batch["n_bad"] with torch.autocast(device_type=DEVICE.type, dtype=DTYPE): loss = model(batch["pixel_values"].to(DEVICE, non_blocking=True), batch["input_ids"].to(DEVICE, non_blocking=True), batch["attention_mask"].to(DEVICE, non_blocking=True), batch["labels"].to(DEVICE, non_blocking=True)).float() if not torch.isfinite(loss): skipped += 1 log(f"non-finite loss at step {step}; dropping accumulation window (skipped={skipped})") opt.zero_grad(set_to_none=True) micro = 0 continue (loss / GRAD_ACCUM).backward() except torch.cuda.OutOfMemoryError: log("CUDA OOM - rerun with a smaller BATCH_SIZE (and larger GRAD_ACCUM)") sys.exit(3) loss_sum += loss.item() loss_n += 1 micro += 1 if micro < GRAD_ACCUM: continue micro = 0 s = lr_scale(step, total_steps, warmup) for grp in opt.param_groups: grp["lr"] = grp["base_lr"] * s grad_norm = torch.nn.utils.clip_grad_norm_(params, max_norm=1.0).item() opt.step() opt.zero_grad(set_to_none=True) step += 1 steps_window += 1 samples_seen += eff_bs if step % LOG_EVERY == 0 or step == 1 or step == total_steps: dt = time.time() - t_window sps = dt / max(1, steps_window) avg = loss_sum / max(1, loss_n) last_loss = avg peak = torch.cuda.max_memory_allocated() / 1e9 if DEVICE.type == "cuda" else 0.0 log(f"step {step}/{total_steps} | loss {avg:.4f} | grad {grad_norm:.2f} | lr {LORA_LR * s:.2e} | " f"{sps:.3f} s/step | {eff_bs / sps:.1f} conv/s | ETA {(total_steps - step) * sps / 3600:.2f} h | " f"peak {peak:.1f} GB | bad_imgs {bad}") log_f.write(json.dumps({"step": step, "loss": avg, "grad_norm": grad_norm, "lr": LORA_LR * s, "s_per_step": sps, "peak_gb": peak, "samples_seen": samples_seen, "time": time.time()}) + "\n") log_f.flush() loss_sum, loss_n, t_window, steps_window = 0.0, 0, time.time(), 0 if step % SAVE_EVERY == 0: save_checkpoint(ckpt_path, model, opt, samples_seen, step) log(f"checkpoint saved at step {step} ({samples_seen:,} samples)") save_checkpoint(ckpt_path, model, opt, samples_seen, step) done = step >= total_steps if done and not MAX_STEPS: model.llm.save_pretrained(os.path.join(CKPT_DIR, "lora_adapter")) torch.save(model.projector.state_dict(), os.path.join(CKPT_DIR, "projector_stage2.pt")) log("saved lora_adapter/ and projector_stage2.pt") peak = torch.cuda.max_memory_allocated() / 1e9 if DEVICE.type == "cuda" else 0.0 log(f"finished: step {step}/{total_steps}, last avg loss {last_loss}, {(time.time() - t_start) / 60:.1f} min, " f"peak {peak:.1f} GB, skipped {skipped}, bad_imgs {bad}") T.write_result({"complete": done, "smoke": bool(MAX_STEPS), "step": step, "total_steps": total_steps, "last_loss": last_loss, "peak_gb": round(peak, 1), "skipped": skipped, "bad_imgs": bad, "batch_size": BATCH_SIZE}) if __name__ == "__main__": main()