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.py from Modularcomputing/AtlasVision: direct link, hf CLI and curl.
- Browser
- Download file 16.3 kB
-
https://huggingface.co/Modularcomputing/AtlasVision/resolve/c8317c0100f694d3a46f0e00b32159e6829feecc/code/train.py
- Command line
-
hf download hf://Modularcomputing/AtlasVision@c8317c0100f694d3a46f0e00b32159e6829feecc/code/train.py
-
curl -L -o train.py https://huggingface.co/Modularcomputing/AtlasVision/resolve/c8317c0100f694d3a46f0e00b32159e6829feecc/code/train.py
16.3 kB
| #!/usr/bin/env python3 | |
| """Stage 1 (projector alignment): SigLIP2 vision tower + N-ATLaS LLM, only the MLP projector trains. | |
| Fixes vs. the Colab notebook: | |
| * LLM loaded in bf16 (notebook loaded fp16), bf16 autocast, fp32 projector, no GradScaler | |
| * dynamic padding + attention mask (notebook padded every sample to 512 tokens; captions are ~10-25) | |
| * image tokens placed after <BOS><user header>, not before BOS | |
| * fixed random subset, cosine LR schedule, resume by samples seen (instant, batch-size independent) | |
| * zip path check so missing images fail loudly instead of silently training on white images | |
| All settings come from environment variables (see CONFIG below). | |
| """ | |
| import json | |
| import math | |
| import os | |
| import sys | |
| import time | |
| import zipfile | |
| import torch | |
| import torch.nn as nn | |
| from PIL import Image | |
| from torch.utils.data import DataLoader, Dataset, Subset | |
| from transformers import AutoImageProcessor, AutoModel, AutoModelForCausalLM, AutoTokenizer | |
| def env(name, default, cast=str): | |
| v = os.environ.get(name) | |
| return cast(v) if v not in (None, "") else default | |
| # ----------------------------- CONFIG ----------------------------- | |
| HOME = os.path.expanduser("~") | |
| LLM_NAME = env("LLM_NAME", "NCAIR1/N-ATLaS") | |
| VISION_NAME = env("VISION_NAME", "google/siglip2-base-patch16-224") | |
| DATA_DIR = env("DATA_DIR", f"{HOME}/data/llava_pretrain") | |
| JSON_PATH = env("JSON_PATH", f"{DATA_DIR}/blip_laion_cc_sbu_558k.json") | |
| ZIP_PATH = env("ZIP_PATH", f"{DATA_DIR}/images.zip") | |
| CKPT_DIR = env("CKPT_DIR", f"{HOME}/checkpoints/stage1") | |
| NUM_SAMPLES = env("NUM_SAMPLES", 150_000, int) | |
| BATCH_SIZE = env("BATCH_SIZE", 16, int) | |
| GRAD_ACCUM = env("GRAD_ACCUM", 1, int) | |
| LR = env("LR", 2e-4, float) | |
| WARMUP_RATIO = env("WARMUP_RATIO", 0.03, float) | |
| MAX_TEXT_LEN = env("MAX_TEXT_LEN", 128, int) | |
| MAX_STEPS = env("MAX_STEPS", 0, int) # >0 caps optimizer steps (smoke tests) | |
| GRAD_CKPT = env("GRAD_CKPT", 0, int) | |
| LOG_EVERY = env("LOG_EVERY", 25, int) | |
| SAVE_EVERY = env("SAVE_EVERY", 500, int) | |
| SEED = env("SEED", 42, int) | |
| NUM_WORKERS = env("NUM_WORKERS", max(1, int(float(os.environ.get("GMN_CPU_LIMIT", "8"))) - 2), int) | |
| DEVICE = torch.device(env("DEVICE", "cuda" if torch.cuda.is_available() else "cpu")) | |
| DTYPE = torch.bfloat16 | |
| USER_HEADER = "<|start_header_id|>user<|end_header_id|>\n\n" | |
| ASSIST_HEADER = "<|eot_id|><|start_header_id|>assistant<|end_header_id|>\n\n" | |
| EOT = "<|eot_id|>" | |
| def log(msg): | |
| print(f"[{time.strftime('%H:%M:%S')}] {msg}", flush=True) | |
| # ----------------------------- DATA ----------------------------- | |
| def find_zip_prefix(annotations, zip_path, n_check=500): | |
| """Return the prefix to put before item['image'] inside the zip, verifying images exist.""" | |
| with zipfile.ZipFile(zip_path) as zf: | |
| names = set(zf.namelist()) | |
| sample = [a["image"] for a in annotations[:n_check]] | |
| prefix = "" | |
| if sum(s in names for s in sample) < 0.95 * len(sample): | |
| first = sample[0] | |
| hits = [n for n in names if n.endswith("/" + first)] | |
| if hits: | |
| prefix = hits[0][: -len(first)] | |
| found = sum((prefix + s) in names for s in sample) | |
| log(f"zip check: {found}/{len(sample)} images found (prefix={prefix!r}, {len(names):,} files in zip)") | |
| if found < 0.95 * len(sample): | |
| raise SystemExit(f"Too many images missing from {zip_path}; example: {sample[0]!r}") | |
| return prefix | |
| class LLaVAPretrainDataset(Dataset): | |
| def __init__(self, annotations, zip_path, zip_prefix, processor, tokenizer, max_text_len): | |
| self.ann = annotations | |
| self.zip_path = zip_path | |
| self.zip_prefix = zip_prefix | |
| self.processor = processor | |
| self.tok = tokenizer | |
| self.max_len = max_text_len | |
| self.zip = None # opened lazily, once per dataloader worker | |
| def __len__(self): | |
| return len(self.ann) | |
| def __getitem__(self, idx): | |
| if self.zip is None: | |
| self.zip = zipfile.ZipFile(self.zip_path) | |
| item = self.ann[idx] | |
| user, answer = "", "" | |
| for turn in item["conversations"]: | |
| if turn["from"] == "human": | |
| user = turn["value"].replace("<image>", "").strip() | |
| elif turn["from"] == "gpt": | |
| answer = turn["value"].strip() | |
| ok = True | |
| try: | |
| with self.zip.open(self.zip_prefix + item["image"]) as f: | |
| image = Image.open(f).convert("RGB") | |
| except Exception: | |
| image = Image.new("RGB", (224, 224), "white") | |
| ok = False | |
| pixel_values = self.processor(images=image, return_tensors="pt").pixel_values[0] | |
| q = self.tok(user + ASSIST_HEADER, add_special_tokens=False).input_ids | |
| a = self.tok(answer + EOT, add_special_tokens=False).input_ids | |
| ids = (q + a)[: self.max_len] | |
| labels = ([-100] * len(q) + a)[: self.max_len] | |
| return {"pixel_values": pixel_values, "input_ids": ids, "labels": labels, "ok": ok} | |
| def make_collate(pad_id): | |
| def collate(batch): | |
| L = max(len(b["input_ids"]) for b in batch) | |
| B = len(batch) | |
| ids = torch.full((B, L), pad_id, dtype=torch.long) | |
| labels = torch.full((B, L), -100, dtype=torch.long) | |
| mask = torch.zeros((B, L), dtype=torch.long) | |
| for i, b in enumerate(batch): | |
| n = len(b["input_ids"]) | |
| ids[i, :n] = torch.tensor(b["input_ids"]) | |
| labels[i, :n] = torch.tensor(b["labels"]) | |
| mask[i, :n] = 1 | |
| return { | |
| "pixel_values": torch.stack([b["pixel_values"] for b in batch]), | |
| "input_ids": ids, | |
| "labels": labels, | |
| "attention_mask": mask, | |
| "n_bad": sum(not b["ok"] for b in batch), | |
| } | |
| return collate | |
| # ----------------------------- MODEL ----------------------------- | |
| class ProjectionMLP(nn.Module): | |
| """Same layout as the notebook (net.0 / net.2), so Colab checkpoints load.""" | |
| def __init__(self, vision_dim, text_dim): | |
| super().__init__() | |
| self.net = nn.Sequential(nn.Linear(vision_dim, text_dim), nn.GELU(), nn.Linear(text_dim, text_dim)) | |
| def forward(self, x): | |
| return self.net(x) | |
| class AtlasVision(nn.Module): | |
| """[<BOS> user-header] [image tokens] [user text, assistant header, answer]""" | |
| def __init__(self, vision, llm, prefix_ids, vision_dim, text_dim): | |
| super().__init__() | |
| self.vision = vision | |
| self.llm = llm | |
| self.projector = ProjectionMLP(vision_dim, text_dim) | |
| self.register_buffer("prefix_ids", torch.tensor(prefix_ids, dtype=torch.long)[None], persistent=False) | |
| def build_inputs(self, pixel_values, input_ids, attention_mask, labels=None): | |
| B = input_ids.size(0) | |
| with torch.no_grad(): | |
| feats = self.vision(pixel_values=pixel_values.to(DTYPE)).last_hidden_state | |
| img = self.projector(feats.float()).to(DTYPE) | |
| emb = self.llm.get_input_embeddings() | |
| pre = emb(self.prefix_ids.expand(B, -1)).to(DTYPE) | |
| txt = emb(input_ids).to(DTYPE) | |
| n_fixed = pre.size(1) + img.size(1) | |
| inputs_embeds = torch.cat([pre, img, txt], dim=1) | |
| mask = torch.cat([attention_mask.new_ones(B, n_fixed), attention_mask], dim=1) | |
| full_labels = None | |
| if labels is not None: | |
| full_labels = torch.cat([labels.new_full((B, n_fixed), -100), labels], dim=1) | |
| return inputs_embeds, mask, full_labels | |
| def forward(self, pixel_values, input_ids, attention_mask, labels): | |
| inputs_embeds, mask, full_labels = self.build_inputs(pixel_values, input_ids, attention_mask, labels) | |
| return self.llm(inputs_embeds=inputs_embeds, attention_mask=mask, labels=full_labels, use_cache=False).loss | |
| def load_pretrained(cls, name, **kw): | |
| try: | |
| return cls.from_pretrained(name, dtype=DTYPE, **kw) | |
| except TypeError: | |
| return cls.from_pretrained(name, torch_dtype=DTYPE, **kw) | |
| def build_model(): | |
| tok = AutoTokenizer.from_pretrained(LLM_NAME) | |
| pad_id = tok.pad_token_id if tok.pad_token_id is not None else tok.eos_token_id | |
| processor = AutoImageProcessor.from_pretrained(VISION_NAME) | |
| full_vision = load_pretrained(AutoModel, VISION_NAME) | |
| vision = full_vision.vision_model | |
| vision_dim = full_vision.config.vision_config.hidden_size | |
| del full_vision # drops the unused SigLIP text tower | |
| llm = load_pretrained(AutoModelForCausalLM, LLM_NAME) | |
| text_dim = llm.config.hidden_size | |
| llm.config.use_cache = False | |
| vision.to(DEVICE).eval() | |
| llm.to(DEVICE) | |
| for p in vision.parameters(): | |
| p.requires_grad_(False) | |
| for p in llm.parameters(): | |
| p.requires_grad_(False) | |
| if GRAD_CKPT: | |
| llm.gradient_checkpointing_enable(gradient_checkpointing_kwargs={"use_reentrant": False}) | |
| llm.train() # HF only checkpoints in train mode; Llama has no dropout | |
| else: | |
| llm.eval() | |
| prefix_ids = tok(USER_HEADER, add_special_tokens=True).input_ids | |
| model = AtlasVision(vision, llm, prefix_ids, vision_dim, text_dim) | |
| model.projector.to(DEVICE, dtype=torch.float32) | |
| model.prefix_ids = model.prefix_ids.to(DEVICE) | |
| return model, tok, pad_id, processor | |
| # ----------------------------- TRAIN ----------------------------- | |
| def lr_at(step, total, warmup): | |
| if step < warmup: | |
| return LR * (step + 1) / warmup | |
| progress = (step - warmup) / max(1, total - warmup) | |
| return LR * 0.5 * (1 + math.cos(math.pi * min(1.0, progress))) | |
| def save_checkpoint(path, model, opt, samples_seen, step, extra=None): | |
| tmp = path + ".tmp" | |
| torch.save({ | |
| "projector_state_dict": model.projector.state_dict(), | |
| "optimizer_state_dict": opt.state_dict(), | |
| "samples_seen": samples_seen, | |
| "opt_step": step, | |
| "config": {"LR": LR, "BATCH_SIZE": BATCH_SIZE, "GRAD_ACCUM": GRAD_ACCUM, "NUM_SAMPLES": NUM_SAMPLES, | |
| "SEED": SEED, "LLM_NAME": LLM_NAME, "VISION_NAME": VISION_NAME, **(extra or {})}, | |
| }, tmp) | |
| os.replace(tmp, path) # atomic: a stop mid-save never corrupts the last good checkpoint | |
| def write_result(d): | |
| path = os.environ.get("GMN_RESULT_PATH") | |
| if path: | |
| with open(path, "w") as f: | |
| json.dump(d, f) | |
| def main(): | |
| torch.manual_seed(SEED) | |
| os.makedirs(CKPT_DIR, exist_ok=True) | |
| ckpt_path = os.path.join(CKPT_DIR, "latest.pt") | |
| log(f"device={DEVICE} batch={BATCH_SIZE}x{GRAD_ACCUM} lr={LR} samples={NUM_SAMPLES} " | |
| f"grad_ckpt={GRAD_CKPT} workers={NUM_WORKERS} max_steps={MAX_STEPS or 'full'}") | |
| model, tok, pad_id, processor = build_model() | |
| n_train = sum(p.numel() for p in model.parameters() if p.requires_grad) | |
| log(f"trainable params: {n_train:,} (projector only)") | |
| with open(JSON_PATH) as f: | |
| annotations = json.load(f) | |
| g = torch.Generator().manual_seed(SEED) | |
| keep = torch.randperm(len(annotations), generator=g)[:NUM_SAMPLES].tolist() | |
| annotations = [annotations[i] for i in keep] | |
| zip_prefix = find_zip_prefix(annotations, ZIP_PATH) | |
| dataset = LLaVAPretrainDataset(annotations, ZIP_PATH, zip_prefix, processor, tok, MAX_TEXT_LEN) | |
| eff_bs = BATCH_SIZE * GRAD_ACCUM | |
| total_steps = math.ceil(len(dataset) / eff_bs) | |
| if MAX_STEPS: | |
| total_steps = min(total_steps, MAX_STEPS) | |
| warmup = max(1, int(total_steps * WARMUP_RATIO)) | |
| params = [p for p in model.projector.parameters()] | |
| opt = torch.optim.AdamW(params, lr=LR, weight_decay=0.0) | |
| samples_seen, step = 0, 0 | |
| if os.path.exists(ckpt_path): | |
| ck = torch.load(ckpt_path, map_location="cpu") | |
| model.projector.load_state_dict(ck["projector_state_dict"]) | |
| opt.load_state_dict(ck["optimizer_state_dict"]) | |
| samples_seen = ck["samples_seen"] | |
| step = samples_seen // eff_bs # works even if batch size changed | |
| log(f"resumed: {samples_seen:,} samples seen -> optimizer step {step}/{total_steps}") | |
| del ck | |
| if step >= total_steps: | |
| log("already complete") | |
| return | |
| order = torch.randperm(len(dataset), generator=torch.Generator().manual_seed(SEED + 1))[samples_seen:].tolist() | |
| loader = DataLoader( | |
| Subset(dataset, order), batch_size=BATCH_SIZE, shuffle=False, num_workers=NUM_WORKERS, | |
| pin_memory=DEVICE.type == "cuda", collate_fn=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, {warmup} warmup, starting at {step}") | |
| model.projector.train() | |
| if DEVICE.type == "cuda": | |
| torch.cuda.reset_peak_memory_stats() | |
| opt.zero_grad(set_to_none=True) | |
| micro, loss_sum, loss_n, bad_imgs, skipped = 0, 0.0, 0, 0, 0 | |
| last_loss, t_window, steps_window = None, time.time(), 0 | |
| t_start = time.time() | |
| log_f = open(os.path.join(CKPT_DIR, "train_log.jsonl"), "a") | |
| for batch in loader: | |
| if step >= total_steps: | |
| break | |
| try: | |
| pv = batch["pixel_values"].to(DEVICE, non_blocking=True) | |
| ids = batch["input_ids"].to(DEVICE, non_blocking=True) | |
| mask = batch["attention_mask"].to(DEVICE, non_blocking=True) | |
| labels = batch["labels"].to(DEVICE, non_blocking=True) | |
| bad_imgs += batch["n_bad"] | |
| with torch.autocast(device_type=DEVICE.type, dtype=DTYPE): | |
| loss = model(pv, ids, mask, labels).float() | |
| if not torch.isfinite(loss): | |
| skipped += 1 | |
| log(f"non-finite loss at step {step}; dropping this 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 GRAD_CKPT=1 or a smaller BATCH_SIZE") | |
| sys.exit(3) | |
| loss_sum += loss.item() | |
| loss_n += 1 | |
| micro += 1 | |
| if micro < GRAD_ACCUM: | |
| continue | |
| micro = 0 | |
| lr = lr_at(step, total_steps, warmup) | |
| for grp in opt.param_groups: | |
| grp["lr"] = lr | |
| 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 | |
| eta_h = (total_steps - step) * sps / 3600 | |
| log(f"step {step}/{total_steps} | loss {avg:.4f} | grad {grad_norm:.2f} | lr {lr:.2e} | " | |
| f"{sps:.3f} s/step | {eff_bs / sps:.1f} img/s | ETA {eta_h:.2f} h | peak {peak:.1f} GB | " | |
| f"bad_imgs {bad_imgs}") | |
| log_f.write(json.dumps({"step": step, "loss": avg, "grad_norm": grad_norm, "lr": lr, | |
| "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: | |
| torch.save(model.projector.state_dict(), os.path.join(CKPT_DIR, "projector_final.pt")) | |
| log("saved projector_final.pt") | |
| peak = torch.cuda.max_memory_allocated() / 1e9 if DEVICE.type == "cuda" else 0.0 | |
| elapsed = time.time() - t_start | |
| log(f"finished: step {step}/{total_steps}, last avg loss {last_loss}, {elapsed / 60:.1f} min, " | |
| f"peak {peak:.1f} GB, skipped {skipped}, bad_imgs {bad_imgs}") | |
| 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_imgs, | |
| "grad_ckpt": GRAD_CKPT, "batch_size": BATCH_SIZE}) | |
| if __name__ == "__main__": | |
| main() | |