arcisvlm / scripts /train_stage1_ddp.py
Hardik Sanghvi
feat: integrate Gemma 4 E2B backbone for production-quality VLM inference
7a564e3
Raw
History Blame Contribute Delete
19 kB
#!/usr/bin/env python3
"""
Stage 1: JEPA Alignment — DDP Training on 8x A100.
Trains X-Encoder + Predictor + Y-Encoder with InfoNCE loss across
multiple GPUs using PyTorch DistributedDataParallel.
Usage:
torchrun --nproc_per_node=8 scripts/train_stage1_ddp.py --config configs/scale_1.3b.yaml
torchrun --nproc_per_node=8 scripts/train_stage1_ddp.py --config configs/scale_1.3b.yaml --resume checkpoints/stage1_epoch2.pt
"""
import argparse
import math
import os
import sys
import time
import torch
import torch.distributed as dist
import torch.nn as nn
from torch.nn.parallel import DistributedDataParallel as DDP
from torch.utils.data import DataLoader, Dataset
from torch.utils.data.distributed import DistributedSampler
import yaml
# Add project root to path
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
from model.vlm import VLJEPAModel
from model.tokenizer import BPETokenizer
# ---------------------------------------------------------------------------
# Dataset helpers
# ---------------------------------------------------------------------------
def build_stage1_dataset(config: dict, tokenizer) -> Dataset:
"""
Build Stage 1 caption dataset from CC3M/SBU/LAION-COCO.
Raises RuntimeError if no real data is found.
"""
img_size = config["vision"]["img_size"]
vocab_size = config["decoder"]["vocab_size"]
# FIRST: Check for pre-downloaded JSONL data (from download_all_data.py)
for jsonl_dir in ["data/downloads/stage1", "data/downloads/stage1_fullscale"]:
if os.path.exists(jsonl_dir):
try:
import json as _json
from PIL import Image as _Image
from torchvision import transforms as _transforms
class Stage1JSONLDataset(Dataset):
"""Load Stage 1 image-caption pairs from JSONL with real images."""
def __init__(self, jsonl_dir, tokenizer, img_size, max_cap=128):
self.samples = []
self.tokenizer = tokenizer
self.max_cap = max_cap
self.transform = _transforms.Compose([
_transforms.Resize((img_size, img_size)),
_transforms.ToTensor(),
_transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]),
])
for fname in sorted(os.listdir(jsonl_dir)):
if fname.endswith('.jsonl'):
with open(os.path.join(jsonl_dir, fname)) as f:
for line in f:
try:
self.samples.append(_json.loads(line.strip()))
except _json.JSONDecodeError:
continue
def __len__(self):
return len(self.samples)
def __getitem__(self, idx):
item = self.samples[idx]
image_path = item.get("image_path")
if not image_path or not os.path.exists(image_path):
raise FileNotFoundError(f"Image not found: {image_path}")
image = self.transform(_Image.open(image_path).convert("RGB"))
caption = item.get("answer", item.get("caption", ""))
cap_ids = self.tokenizer.encode(str(caption))[:self.max_cap]
cap_ids += [self.tokenizer.pad_id] * (self.max_cap - len(cap_ids))
cap_t = torch.tensor(cap_ids, dtype=torch.long)
return {
"image": image,
"caption_ids": cap_t,
"caption_mask": (cap_t != self.tokenizer.pad_id).long(),
}
dataset = Stage1JSONLDataset(jsonl_dir, tokenizer, img_size)
if len(dataset) > 100:
print(f" [REAL DATA] Stage 1 JSONL: {len(dataset)} samples from {jsonl_dir}")
return dataset
except Exception as e:
print(f" [WARN] Stage 1 JSONL loading failed: {e}")
# SECOND: Try loading CC3M from local paired JSON
cc3m_json = "data/cc3m_paired.json"
if os.path.exists(cc3m_json):
try:
import json
from PIL import Image
from torchvision import transforms
class CC3MLocalDataset(Dataset):
def __init__(self, pairs, tokenizer, img_size, vocab_size, max_cap=128):
self.pairs = pairs
self.tokenizer = tokenizer
self.max_cap = max_cap
self.vocab_size = vocab_size
self.transform = transforms.Compose([
transforms.Resize((img_size, img_size)),
transforms.ToTensor(),
transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]),
])
def __len__(self):
return len(self.pairs)
def __getitem__(self, idx):
p = self.pairs[idx]
try:
img = self.transform(Image.open(p["image"]).convert("RGB"))
except Exception as e:
raise FileNotFoundError(f"Image not found: {p['image']}. Data may be corrupted. Original error: {e}")
cap_ids = self.tokenizer.encode(p["caption"])[:self.max_cap]
cap_ids += [self.tokenizer.pad_id] * (self.max_cap - len(cap_ids))
cap_t = torch.tensor(cap_ids, dtype=torch.long)
return {
"image": img,
"caption_ids": cap_t,
"caption_mask": (cap_t != self.tokenizer.pad_id).long(),
}
with open(cc3m_json) as f:
pairs = json.load(f)
if len(pairs) > 100:
dataset = CC3MLocalDataset(pairs, tokenizer, img_size, vocab_size)
print(f"[REAL DATA] CC3M local: {len(dataset)} image-caption pairs")
return dataset
except Exception as e:
print(f"[WARN] CC3M local loading failed: {e}")
# Try loading real datasets via data/multi_dataset.py
try:
from data.multi_dataset import build_stage1_dataset as _build
dataset = _build(config, tokenizer)
if len(dataset) > 0:
return dataset
except (ImportError, Exception) as e:
pass
# Try the local CaptionDataset with Flickr8k
try:
from data.dataset import CaptionDataset
dataset = CaptionDataset(
image_dir="data/flickr8k/Images",
captions_file="data/flickr8k/captions.txt",
tokenizer=tokenizer,
img_size=img_size,
)
if len(dataset) > 0:
return dataset
except Exception:
pass
raise RuntimeError(
"FATAL: No Stage 1 training data found.\n"
"Download real data first: python3 scripts/download_all_data.py --stage 1\n"
"Required: data/downloads/stage1/ with JSONL files containing image_path"
)
# ---------------------------------------------------------------------------
# LR scheduler with linear warmup + cosine decay
# ---------------------------------------------------------------------------
class CosineWarmupScheduler(torch.optim.lr_scheduler._LRScheduler):
"""Linear warmup for `warmup_steps`, then cosine decay to `min_lr`."""
def __init__(self, optimizer, warmup_steps: int, total_steps: int,
min_lr: float = 1e-6, last_epoch: int = -1):
self.warmup_steps = warmup_steps
self.total_steps = total_steps
self.min_lr = min_lr
super().__init__(optimizer, last_epoch)
def get_lr(self):
step = self.last_epoch
if step < self.warmup_steps:
# Linear warmup
scale = step / max(1, self.warmup_steps)
return [base_lr * scale for base_lr in self.base_lrs]
else:
# Cosine decay
progress = (step - self.warmup_steps) / max(1, self.total_steps - self.warmup_steps)
cosine = 0.5 * (1.0 + math.cos(math.pi * progress))
return [
self.min_lr + (base_lr - self.min_lr) * cosine
for base_lr in self.base_lrs
]
# ---------------------------------------------------------------------------
# Training
# ---------------------------------------------------------------------------
def setup_distributed():
"""Initialize distributed process group."""
dist.init_process_group(backend="nccl")
local_rank = int(os.environ["LOCAL_RANK"])
torch.cuda.set_device(local_rank)
return local_rank
def cleanup():
"""Destroy process group."""
if dist.is_initialized():
dist.destroy_process_group()
def is_rank0():
return not dist.is_initialized() or dist.get_rank() == 0
def log(msg: str):
"""Print only on rank 0."""
if is_rank0():
print(msg, flush=True)
def save_checkpoint(model: nn.Module, optimizer, scheduler, epoch: int,
global_step: int, loss: float, path: str):
"""Save checkpoint from rank 0 only."""
if not is_rank0():
return
os.makedirs(os.path.dirname(path), exist_ok=True)
# Unwrap DDP module
state_dict = model.module.state_dict() if hasattr(model, "module") else model.state_dict()
torch.save({
"epoch": epoch,
"global_step": global_step,
"model_state_dict": state_dict,
"optimizer_state_dict": optimizer.state_dict(),
"scheduler_state_dict": scheduler.state_dict(),
"loss": loss,
}, path)
log(f" Checkpoint saved: {path}")
def push_checkpoints():
"""Push checkpoints to GitHub LFS. Disabled during training to avoid git lock issues.
Call scripts/push_checkpoints.py manually after training completes."""
pass # Disabled — run push_checkpoints.py separately after training
def main():
parser = argparse.ArgumentParser(description="Stage 1 DDP: JEPA Alignment")
parser.add_argument("--config", type=str, required=True, help="Path to YAML config")
parser.add_argument("--resume", type=str, default=None, help="Path to checkpoint to resume from")
args = parser.parse_args()
# ---- Distributed setup ----
local_rank = setup_distributed()
world_size = dist.get_world_size()
global_rank = dist.get_rank()
device = torch.device(f"cuda:{local_rank}")
# ---- Config ----
with open(args.config) as f:
config = yaml.safe_load(f)
stage_cfg = config["train_stage1"]
per_gpu_batch = stage_cfg.get("batch_size", 4)
grad_accum = stage_cfg.get("gradient_accumulation", 1)
effective_batch = per_gpu_batch * world_size * grad_accum
max_epochs = stage_cfg["max_epochs"]
lr = stage_cfg["learning_rate"]
warmup_steps = stage_cfg["warmup_steps"]
grad_clip = stage_cfg["gradient_clip"]
log("=" * 70)
log("ArcisVLM — Stage 1: JEPA Alignment (DDP)")
log("=" * 70)
log(f" World size: {world_size}")
log(f" Per-GPU batch: {per_gpu_batch}")
log(f" Grad accum: {grad_accum}")
log(f" Effective batch: {effective_batch}")
log(f" Max epochs: {max_epochs}")
log(f" Learning rate: {lr}")
log(f" Warmup steps: {warmup_steps}")
log(f" Precision: {stage_cfg.get('precision', 'bf16')}")
log(f" Grad checkpoint: enabled (ViT encoder)")
# ---- Tokenizer ----
tokenizer = BPETokenizer(vocab_size=config["decoder"]["vocab_size"])
tok_path = "checkpoints/tokenizer_32k.json"
if os.path.exists(tok_path):
tokenizer.load(tok_path)
log(f" Tokenizer: {len(tokenizer)} tokens (from {tok_path})")
else:
tok_fallback = "checkpoints/tokenizer.json"
if os.path.exists(tok_fallback):
tokenizer.load(tok_fallback)
log(f" Tokenizer: {len(tokenizer)} tokens (from {tok_fallback})")
else:
log(" [WARN] No tokenizer found — using untrained tokenizer")
# ---- Dataset ----
dataset = build_stage1_dataset(config, tokenizer)
sampler = DistributedSampler(dataset, num_replicas=world_size, rank=global_rank, shuffle=True)
loader = DataLoader(
dataset,
batch_size=per_gpu_batch,
sampler=sampler,
num_workers=4,
pin_memory=True,
drop_last=True,
)
log(f" Dataset: {len(dataset)} samples, {len(loader)} batches/GPU")
# ---- Model ----
model = VLJEPAModel(config).to(device)
# Enable gradient checkpointing on ViT encoder to save ~40% VRAM
if hasattr(model, 'x_encoder'):
model.x_encoder._gradient_checkpointing = True
log(" Gradient checkpointing: enabled on x_encoder (ViT)")
if hasattr(model, 'y_encoder') and hasattr(model.y_encoder, 'blocks'):
# Y-encoder is smaller, but checkpoint it too for safety
pass
if is_rank0():
params = model.count_parameters()
for k, v in params.items():
log(f" {k}: {v:,}")
# ---- Optimizer (Y-Encoder gets slower LR) ----
y_params = list(model.y_encoder.parameters())
y_param_ids = {id(p) for p in y_params}
other_params = [p for p in model.parameters() if id(p) not in y_param_ids and p.requires_grad]
y_lr = lr * config["y_encoder"]["lr_multiplier"]
optimizer = torch.optim.AdamW([
{"params": other_params, "lr": lr},
{"params": y_params, "lr": y_lr},
], weight_decay=0.01)
# ---- Scheduler ----
total_steps = max_epochs * len(loader)
scheduler = CosineWarmupScheduler(optimizer, warmup_steps=warmup_steps, total_steps=total_steps)
# ---- Mixed precision ----
use_bf16 = stage_cfg.get("precision", "bf16") == "bf16"
scaler = torch.amp.GradScaler("cuda", enabled=(not use_bf16)) # GradScaler not needed for bf16
autocast_dtype = torch.bfloat16 if use_bf16 else torch.float16
# ---- Resume ----
start_epoch = 0
global_step = 0
if args.resume and os.path.exists(args.resume):
ckpt = torch.load(args.resume, map_location=device, weights_only=False)
model.load_state_dict(ckpt["model_state_dict"])
optimizer.load_state_dict(ckpt["optimizer_state_dict"])
if "scheduler_state_dict" in ckpt:
scheduler.load_state_dict(ckpt["scheduler_state_dict"])
start_epoch = ckpt["epoch"]
global_step = ckpt.get("global_step", start_epoch * len(loader))
log(f" Resumed from {args.resume} (epoch {start_epoch}, loss {ckpt['loss']:.4f})")
# ---- DDP wrap ----
model = DDP(model, device_ids=[local_rank], output_device=local_rank,
find_unused_parameters=False)
# ---- Training loop with gradient accumulation ----
model.train()
os.makedirs("checkpoints", exist_ok=True)
for epoch in range(start_epoch, max_epochs):
sampler.set_epoch(epoch)
epoch_loss = 0.0
epoch_steps = 0
epoch_start = time.time()
optimizer.zero_grad(set_to_none=True)
for batch_idx, batch in enumerate(loader):
images = batch["image"].to(device, non_blocking=True)
cap_ids = batch["caption_ids"].to(device, non_blocking=True)
cap_mask = batch["caption_mask"].to(device, non_blocking=True)
# Gradient accumulation: scale loss by accum steps
with torch.amp.autocast("cuda", dtype=autocast_dtype):
output = model.module.forward_stage1(images, None, None, cap_ids, cap_mask)
loss = output["loss"] / grad_accum
if use_bf16:
loss.backward()
else:
scaler.scale(loss).backward()
# Only step optimizer every grad_accum batches
if (batch_idx + 1) % grad_accum == 0 or (batch_idx + 1) == len(loader):
# Gradient clipping
if use_bf16:
torch.nn.utils.clip_grad_norm_(model.parameters(), grad_clip)
optimizer.step()
else:
scaler.unscale_(optimizer)
torch.nn.utils.clip_grad_norm_(model.parameters(), grad_clip)
scaler.step(optimizer)
scaler.update()
scheduler.step()
optimizer.zero_grad(set_to_none=True)
global_step += 1
# Track unscaled loss
epoch_loss += loss.item() * grad_accum
epoch_steps += 1
# Log every 50 optimizer steps
if global_step > 0 and global_step % 50 == 0 and (batch_idx + 1) % grad_accum == 0:
current_lr = scheduler.get_last_lr()[0]
gpu_mem = torch.cuda.max_memory_allocated(device) / 1e9
log(f" [Step {global_step}] loss={loss.item() * grad_accum:.4f} lr={current_lr:.2e} GPU mem={gpu_mem:.1f}GB")
# ---- Epoch summary ----
# All-reduce loss across ranks for accurate average
avg_loss_tensor = torch.tensor([epoch_loss, epoch_steps], device=device, dtype=torch.float64)
dist.all_reduce(avg_loss_tensor, op=dist.ReduceOp.SUM)
avg_loss = (avg_loss_tensor[0] / avg_loss_tensor[1]).item()
epoch_time = time.time() - epoch_start
current_lr = scheduler.get_last_lr()[0]
gpu_mem = torch.cuda.max_memory_allocated(device) / 1e9
log(f"\nEpoch {epoch + 1}/{max_epochs}: loss={avg_loss:.4f} lr={current_lr:.2e} "
f"time={epoch_time:.0f}s GPU mem={gpu_mem:.1f}GB")
# ---- Go/No-Go Gate 1: loss < 3.0 after epoch 1 ----
gate1_threshold = 3.0
if epoch == 0 and avg_loss >= gate1_threshold:
log(f"\n*** GO/NO-GO GATE 1 WARNING: loss={avg_loss:.4f} >= {gate1_threshold} ***")
log("*** Check data pipeline and hyperparameters. Continuing training. ***")
elif epoch == 0:
log(f" Go/No-Go Gate 1 PASSED: loss={avg_loss:.4f} < {gate1_threshold}")
# ---- Checkpoint every epoch ----
ckpt_path = f"checkpoints/stage1_epoch{epoch + 1}.pt"
save_checkpoint(model, optimizer, scheduler, epoch + 1, global_step, avg_loss, ckpt_path)
push_checkpoints()
dist.barrier()
# ---- Final checkpoint ----
save_checkpoint(model, optimizer, scheduler, max_epochs, global_step, avg_loss, "checkpoints/stage1_final.pt")
push_checkpoints()
log("\n" + "=" * 70)
log(f"Stage 1 complete. Final loss: {avg_loss:.4f}")
log("=" * 70)
cleanup()
if __name__ == "__main__":
main()