ViuAI-500M / code /train.py
ViuAI's picture
Update code/train.py
7dd033a verified
Raw
History Blame Contribute Delete
15.4 kB
import os
os.environ["PYTORCH_CUDA_ALLOC_CONF"] = "expandable_segments:True"
import json, math, time, random, sys, contextlib, threading
import numpy as np
import torch
torch.set_float32_matmul_precision('high')
torch.backends.cuda.matmul.allow_tf32 = True
torch.backends.cuda.matmul.allow_bf16_reduced_precision_reduction = True
import torch.distributed as dist
from torch.nn.parallel import DistributedDataParallel as DDP
from huggingface_hub import HfApi, hf_hub_download, create_repo, login
from datetime import timedelta
def setup_ddp():
dist.init_process_group(backend="nccl", timeout=timedelta(minutes=30))
rank = dist.get_rank()
local_rank = int(os.environ["LOCAL_RANK"])
world_size = dist.get_world_size()
torch.cuda.set_device(local_rank)
return rank, local_rank, world_size
rank, local_rank, world_size = setup_ddp()
DEVICE = f"cuda:{local_rank}"
IS_MAIN = rank == 0
def log(msg):
if IS_MAIN:
print(msg, flush=True)
HF_TOKEN = os.environ.get("HF_TOKEN")
if HF_TOKEN is None:
raise ValueError("HF_TOKEN environment variable is not set!")
login(token=HF_TOKEN)
api = HfApi(token=HF_TOKEN)
MODEL_REPO = "ViuAI/ViuAI-500M"
DATA_REPO = "ViuAI/viuai-500m-data"
CKPT_REPO = MODEL_REPO
TOKENIZER_SUBFOLDER = "tokenizer"
CODE_SUBFOLDER = "code"
CKPT_SUBFOLDER = "checkpoints"
if IS_MAIN:
create_repo(CKPT_REPO, repo_type="model", token=HF_TOKEN, exist_ok=True, private=True)
dist.barrier()
if os.path.exists("/kaggle/working"):
WORK = "/kaggle/working"
elif os.path.exists("/root/working"):
WORK = "/root/working"
else:
WORK = "./working"
os.makedirs(WORK, exist_ok=True)
sys.path.insert(0, f"{WORK}/{CODE_SUBFOLDER}")
hf_hub_download(MODEL_REPO, f"{CODE_SUBFOLDER}/config.py", local_dir=WORK, token=HF_TOKEN)
hf_hub_download(MODEL_REPO, f"{CODE_SUBFOLDER}/model.py", local_dir=WORK, token=HF_TOKEN)
from config import ViuAIConfig
from model import ViuAI
def with_retry(fn, *args, retries=4, delay=5, **kwargs):
for attempt in range(1, retries + 1):
try:
return fn(*args, **kwargs)
except Exception as e:
if "404" in str(e) or "EntryNotFoundError" in type(e).__name__:
raise
if attempt == retries:
raise
log(f" [retry {attempt}/{retries}] {type(e).__name__}: {e}")
time.sleep(delay)
delay *= 2
class ShardPool:
def __init__(self, shard_names, repo_id, token, local_dir, pool_size=3, seed=0):
self.all_shards = list(shard_names)
self.rng = random.Random(seed)
self.rng.shuffle(self.all_shards)
self.cursor = 0
self.repo_id, self.token, self.local_dir, self.pool_size = repo_id, token, local_dir, pool_size
os.makedirs(local_dir, exist_ok=True)
self.pool = []
self.fill()
def _next_name(self):
if self.cursor >= len(self.all_shards):
self.rng.shuffle(self.all_shards)
self.cursor = 0
name = self.all_shards[self.cursor]
self.cursor += 1
return name
def fill(self):
while len(self.pool) < self.pool_size:
name = self._next_name()
path = with_retry(hf_hub_download, self.repo_id, name, repo_type="dataset",
local_dir=self.local_dir, token=self.token)
self.pool.append((path, np.memmap(path, dtype=np.uint16, mode="r")))
def rotate_one(self):
old_path, old_mmap = self.pool.pop(0)
del old_mmap
try: os.remove(old_path)
except OSError: pass
self.fill()
def sample_batch(self, batch_size, ctx_len):
xs, ys = [], []
for _ in range(batch_size):
_, tokens = self.rng.choice(self.pool)
i = self.rng.randint(0, len(tokens) - ctx_len - 1)
# CLEAN SHIFT: xs and ys are same length. ys is just shifted +1.
xs.append(np.array(tokens[i:i+ctx_len], dtype=np.int64))
ys.append(np.array(tokens[i+1:i+1+ctx_len], dtype=np.int64))
return np.stack(xs), np.stack(ys)
def load_val_tokens(val_names, repo_id, token, local_dir):
os.makedirs(local_dir, exist_ok=True)
arrays = []
for name in val_names:
path = with_retry(hf_hub_download, repo_id, name, repo_type="dataset", local_dir=local_dir, token=token)
arrays.append(np.memmap(path, dtype=np.uint16, mode="r"))
return np.concatenate(arrays)
def get_batch(source, batch_size, ctx_len):
if isinstance(source, ShardPool):
x_np, y_np = source.sample_batch(batch_size, ctx_len)
else:
i = np.random.randint(0, len(source) - ctx_len - 1, size=batch_size)
x_np = np.stack([np.array(source[j:j+ctx_len], dtype=np.int64) for j in i])
y_np = np.stack([np.array(source[j+1:j+1+ctx_len], dtype=np.int64) for j in i])
x = torch.from_numpy(x_np).pin_memory().to(DEVICE, non_blocking=True)
y = torch.from_numpy(y_np).pin_memory().to(DEVICE, non_blocking=True)
return x, y
def extract_clean_state_dict(model):
"""Safely extracts state dict removing DDP and compile prefixes."""
raw_model = model.module if hasattr(model, "module") else model
state_dict = raw_model.state_dict()
# Remove _orig_mod. prefix if present (from torch.compile)
return {k.replace("_orig_mod.", ""): v for k, v in state_dict.items()}
_upload_threads = []
def wait_for_uploads():
if not _upload_threads:
return
log(f" ⏳ Waiting for {len(_upload_threads)} background upload(s) to finish...")
for t in _upload_threads:
t.join()
_upload_threads.clear()
def save_checkpoint(model, optimizer, step, val_loss, filename, blocking=False):
if not IS_MAIN:
return
path = f"{WORK}/{filename}"
state_dict = extract_clean_state_dict(model)
data = {"model": state_dict, "step": step, "val_loss": val_loss}
# Save optimizer only for latest checkpoint to keep best checkpoint small (1GB vs 5GB)
if filename == "ckpt_latest.pt" and optimizer is not None:
data["optimizer"] = optimizer.state_dict()
torch.save(data, path)
log(f" [{filename} saved @ step {step}, val_loss={val_loss:.4f}] locally")
if blocking:
log(f" ☁️ Uploading {filename} to HF (blocking, please wait)...")
with_retry(api.upload_file, path_or_fileobj=path, path_in_repo=f"{CKPT_SUBFOLDER}/{filename}",
repo_id=CKPT_REPO, repo_type="model", token=HF_TOKEN)
log(f" ✅ Upload complete for {filename} @ step {step}")
else:
def _bg_upload():
try:
with_retry(api.upload_file, path_or_fileobj=path,
path_in_repo=f"{CKPT_SUBFOLDER}/{filename}",
repo_id=CKPT_REPO, repo_type="model", token=HF_TOKEN)
log(f" ☁️ [Background upload complete for {filename} @ step {step}]")
except Exception as e:
log(f" ❌ [Background upload failed for {filename}: {e}]")
t = threading.Thread(target=_bg_upload, daemon=False)
t.start()
_upload_threads.append(t)
log(f" ☁️ Uploading {filename} in background...")
def load_checkpoint(model, optimizer):
try:
path = with_retry(hf_hub_download, CKPT_REPO, f"{CKPT_SUBFOLDER}/ckpt_latest.pt", repo_type="model",
local_dir=WORK, token=HF_TOKEN)
ckpt = torch.load(path, map_location="cpu", weights_only=False)
target = model.module._orig_mod if hasattr(model.module, '_orig_mod') else model.module
target.load_state_dict(ckpt["model"])
optimizer.load_state_dict(ckpt["optimizer"])
log(f"Resumed from step {ckpt['step']} (val_loss={ckpt['val_loss']:.4f})")
return ckpt["step"]
except Exception as e:
log(f"No checkpoint found — starting fresh. ({type(e).__name__})")
return 0
def get_lr(step, warmup, total, max_lr, min_lr):
if step < warmup:
return max_lr * (step + 1) / warmup
if step > total:
return min_lr
ratio = (step - warmup) / (total - warmup)
return min_lr + 0.5 * (1 + math.cos(math.pi * ratio)) * (max_lr - min_lr)
# === 500M GPU CONFIG ===
# Auto-detect GPU memory to set optimal MICRO_BATCH
mem_gb = torch.cuda.get_device_properties(DEVICE).total_memory / (1024**3)
if mem_gb > 100: # H200 (141GB)
MICRO_BATCH = 32
GRAD_ACCUM = 4
EMPTY_CACHE_EVERY = 500
USE_CHECKPOINT = False
elif mem_gb > 70: # A100 (80GB) or RTX 6000 Ada (96GB)
MICRO_BATCH = 16
GRAD_ACCUM = 8
EMPTY_CACHE_EVERY = 250
USE_CHECKPOINT = False
elif mem_gb > 28: # 32GB GPUs (e.g. V100 32GB)
MICRO_BATCH = 16
GRAD_ACCUM = 8
EMPTY_CACHE_EVERY = 500
USE_CHECKPOINT = True
else: # 24GB GPUs (e.g. L4, A10G, RTX 4090)
MICRO_BATCH = 8
GRAD_ACCUM = 16
EMPTY_CACHE_EVERY = 500
USE_CHECKPOINT = True
if IS_MAIN:
print(f"🖥️ Detected GPU with {mem_gb:.1f}GB VRAM -> Auto-Configured: MICRO_BATCH={MICRO_BATCH}, GRAD_ACCUM={GRAD_ACCUM}, USE_CHECKPOINT={USE_CHECKPOINT}")
# =======================
CONTEXT_LEN = 2048
MAX_LR = 3e-4
MIN_LR = 3e-5
WARMUP_STEPS = 1000
TOTAL_STEPS = 40_000
WEIGHT_DECAY = 0.1
GRAD_CLIP = 1.0
MAX_STEPS = int(os.environ.get("MAX_STEPS", TOTAL_STEPS))
EVAL_EVERY = 500
SAVE_EVERY = 4000
ROTATE_EVERY = 500
torch.backends.cudnn.benchmark = True
torch.backends.cudnn.allow_tf32 = True
cfg = ViuAIConfig(context_length=CONTEXT_LEN, use_checkpoint=USE_CHECKPOINT)
model = ViuAI(cfg).to(DEVICE)
if IS_MAIN:
log(f"Model params: {model.num_params()/1e6:.1f}M | world_size={world_size}")
log("Compiling model for Extreme Speed (this takes 1-2 mins on step 0)...")
compute_cap = torch.cuda.get_device_capability(local_rank)
USE_COMPILE = compute_cap[0] >= 8 # Only Ampere+ (A100, H100, H200, RTX 4090)
if USE_COMPILE:
model = torch.compile(model)
model = DDP(model, device_ids=[local_rank])
optimizer = torch.optim.AdamW(model.parameters(), lr=MAX_LR, weight_decay=WEIGHT_DECAY, betas=(0.9, 0.95), fused=True)
start_step = load_checkpoint(model, optimizer)
start_step_t = torch.tensor([start_step], device=DEVICE)
dist.broadcast(start_step_t, src=0)
start_step = start_step_t.item()
files = api.list_repo_files(DATA_REPO, repo_type="dataset")
train_shards = sorted(f for f in files if f.startswith("shards/") and f.endswith(".bin") and "_val" not in f)
val_shards = sorted(f for f in files if f.startswith("shards/") and f.endswith("_val.bin"))
shard_pool = ShardPool(train_shards, DATA_REPO, HF_TOKEN, f"{WORK}/train_shards_r{rank}", seed=rank + start_step)
val_tokens = load_val_tokens(val_shards, DATA_REPO, HF_TOKEN, f"{WORK}/val_shards") if IS_MAIN else None
@torch.no_grad()
def evaluate():
if not IS_MAIN:
return None
model.eval()
losses = []
for _ in range(50):
x, y = get_batch(val_tokens, 4, CONTEXT_LEN)
with torch.autocast(device_type="cuda", dtype=torch.bfloat16):
unwrapped = model.module._orig_mod if hasattr(model.module, '_orig_mod') else model.module
_, loss = unwrapped(x, y)
losses.append(loss.item())
model.train()
return sum(losses) / len(losses)
model.train()
val_loss = None
best_val_loss = None
ema_tok_s = None
global_t0 = time.time()
t0 = time.time()
end_step = min(start_step + MAX_STEPS, TOTAL_STEPS)
for step in range(start_step, end_step):
lr = get_lr(step, WARMUP_STEPS, TOTAL_STEPS, MAX_LR, MIN_LR)
for g in optimizer.param_groups:
g["lr"] = lr
optimizer.zero_grad(set_to_none=True)
accum_loss = 0.0
for micro_i in range(GRAD_ACCUM):
x, y = get_batch(shard_pool, MICRO_BATCH, CONTEXT_LEN)
sync_ctx = model.no_sync() if micro_i < GRAD_ACCUM - 1 else contextlib.nullcontext()
with sync_ctx:
with torch.autocast(device_type="cuda", dtype=torch.bfloat16):
_, loss = model(x, y)
loss = loss / GRAD_ACCUM
loss.backward()
accum_loss += loss.item()
del x, y, loss
grad_norm = torch.nn.utils.clip_grad_norm_(model.parameters(), GRAD_CLIP)
if not math.isfinite(accum_loss) or not torch.isfinite(grad_norm):
if IS_MAIN:
log(f"⚠️ NaN/Inf loss or grad norm at step {step}! Skipping optimizer step.")
optimizer.zero_grad(set_to_none=True)
continue
optimizer.step()
if IS_MAIN and step % 10 == 0:
if step == start_step:
t0 = time.time()
mem_alloc = torch.cuda.memory_allocated(DEVICE) / 1e9
mem_reserved = torch.cuda.memory_reserved(DEVICE) / 1e9
log(f"step {step:5d}/{end_step} | loss {accum_loss:.4f} | lr {lr:.2e} | Calibrating speed... | "
f"mem alloc={mem_alloc:.2f}GB reserved={mem_reserved:.2f}GB")
else:
dt = time.time() - t0
tok_s = (MICRO_BATCH * GRAD_ACCUM * CONTEXT_LEN * world_size * 10) / dt if dt > 0 else 0
ema_tok_s = tok_s if ema_tok_s is None else 0.7 * ema_tok_s + 0.3 * tok_s
t0 = time.time()
mem_alloc = torch.cuda.memory_allocated(DEVICE) / 1e9
mem_reserved = torch.cuda.memory_reserved(DEVICE) / 1e9
steps_done = step - start_step
if steps_done > 0:
elapsed = time.time() - global_t0
eta_s = elapsed / steps_done * (end_step - step)
eta_str = f"{eta_s/60:.0f}min" if eta_s < 3600 else f"{eta_s/3600:.1f}hr"
else:
eta_str = "..."
log(f"step {step:5d}/{end_step} | loss {accum_loss:.4f} | lr {lr:.2e} | {ema_tok_s:.0f} tok/s | "
f"ETA {eta_str} | mem alloc={mem_alloc:.2f}GB reserved={mem_reserved:.2f}GB")
if step > start_step and step % EMPTY_CACHE_EVERY == 0:
pass # torch.cuda.empty_cache() hata diya taaki speed slow na ho
if step > start_step and step % ROTATE_EVERY == 0:
shard_pool.rotate_one()
if step > start_step and step % EVAL_EVERY == 0:
val_loss = evaluate()
if IS_MAIN and val_loss is not None:
log(f"step {step} | val_loss {val_loss:.4f}")
if best_val_loss is None or val_loss < best_val_loss:
best_val_loss = val_loss
# User requested not to save/upload ckpt_best.pt to save time
# save_checkpoint(model.module._orig_mod if hasattr(model.module, '_orig_mod') else model.module,
# optimizer, step, val_loss, "ckpt_best.pt")
if step > start_step and step % SAVE_EVERY == 0:
dist.barrier()
if IS_MAIN:
# Save latest checkpoint (with optimizer, 5GB)
save_checkpoint(model.module._orig_mod if hasattr(model.module, '_orig_mod') else model.module,
optimizer, step, val_loss if val_loss is not None else -1, "ckpt_latest.pt")
dist.barrier()
final_val = evaluate()
if IS_MAIN:
wait_for_uploads()
save_checkpoint(model.module._orig_mod if hasattr(model.module, '_orig_mod') else model.module,
optimizer, end_step, final_val if final_val is not None else -1, "ckpt_latest.pt", blocking=True)
if IS_MAIN:
log(f"\n500M Training complete! Final val_loss: {final_val if final_val is not None else -1:.4f}")
dist.destroy_process_group()