#!/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()