Image-Text-to-Text
PEFT
Safetensors
vision-language
multimodal
llava
lora
siglip2
n-atlas
nigerian-languages
Instructions to use Modularcomputing/AtlasVision with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- PEFT
How to use Modularcomputing/AtlasVision with PEFT:
Task type is invalid.
- Notebooks
- Google Colab
- Kaggle
Download code/train_stage2.py from Modularcomputing/AtlasVision: direct link, hf CLI and curl.
- Browser
- Download file 13.3 kB
-
https://huggingface.co/Modularcomputing/AtlasVision/resolve/main/code/train_stage2.py
- Command line
-
hf download hf://Modularcomputing/AtlasVision/code/train_stage2.py
-
curl -L -o train_stage2.py https://huggingface.co/Modularcomputing/AtlasVision/resolve/main/code/train_stage2.py
13.3 kB
| #!/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: | |
| [<BOS> user-header] [image tokens] [q1 <eot> assistant-header] a1 <eot> [user-header q2 <eot> assistant-header] a2 <eot> ... | |
| Loss only on assistant answers (+ their <eot>). 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("<image>", "").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() | |