ViuAI's picture
Fix CUDA OOM on T4: cast model to amp_dtype directly, batch_size=2, safe RMSNorm
f5cccb7 verified
Raw History Blame Contribute Delete
98.7 kB
"""
VIU-1 pretraining loop (single GPU/CPU, DDP-ready later)
- AdamW + cosine warmup + grad-clip + BF16/FP16 + grad-accum
- checkpoint save/resume, val perplexity, MoE-aux log, sample decode
- dataset: raw .txt pack (abhi) / .bin memmap (Phase-1 me)
Run smoke (CPU, 3 steps, tiny model):
python train.py --smoke
Real:
python train.py --config ../configs/train_config.yaml --model_config ../configs/viu1_moe_config.yaml
"""
import argparse
import concurrent.futures
import contextlib
import json
import math
import os
import random
import re
import sys
import time
from pathlib import Path
import torch
import torch.nn.functional as F
from torch.utils.data import Dataset
sys.path.insert(0, str(Path(__file__).parent))
try:
from viu_moe import ViuArgs, Viu1MoE, load_args as load_model_args
except ImportError:
from viu1_moe import ViuArgs, Viu1MoE, load_args as load_model_args
try:
import yaml
except ImportError:
yaml = None
try:
from tokenizers import Tokenizer
except ImportError:
Tokenizer = None
try:
from tqdm import tqdm
except ImportError:
tqdm = None
# ---------------- data ----------------
def get_lang_tag(name_or_text: str) -> str:
n = str(name_or_text).lower().strip()
# Explicit language codes and names
if n in ("en", "eng", "english") or "english" in n:
return "<|english|>"
if n in ("hi", "hin", "hindi") or "hindi" in n:
return "<|hindi|>"
if n in ("hinglish", "hi-en", "en-hi", "hi_en", "bilingual") or "hinglish" in n:
return "<|hinglish|>"
# Simple heuristic for raw text: Devanagari -> hindi, else ascii -> english, else hinglish
if any('\u0900' <= ch <= '\u097F' for ch in n):
return "<|hindi|>"
if n.isascii():
return "<|english|>"
return "<|hinglish|>"
class TxtPackDataset(Dataset):
"""Sab .txt lines tokenize karke concat + seq_len blocks. Small/mid data ke liye.
NOTE: poora corpus RAM me aata hai — 4-5B tokens ke liye memmap/streaming chahiye (Phase-1).
Bade data pe warning degi."""
def __init__(self, files, tok, seq_len, warn_tokens=20_000_000):
ids = []
eos = tok.token_to_id("<eos>")
if eos is None:
eos = 2
for fp in files:
tag = get_lang_tag(Path(fp).name)
tag_id = tok.token_to_id(tag) if (tag and tok is not None) else None
with open(fp, encoding="utf-8") as f:
for line in f:
line = line.strip()
if line:
if tag_id is not None:
ids.append(tag_id)
ids.extend(tok.encode(line).ids)
ids.append(eos)
if len(ids) > warn_tokens:
print(f"[warn] {len(ids)/1e6:.1f}M tokens RAM me — bade corpus ke liye memmap/HF-streaming use karo.")
self.raw_ids = torch.tensor(ids, dtype=torch.long)
self.seq_len = seq_len
self.rebuild_blocks()
def rebuild_blocks(self):
n = (len(self.raw_ids) // self.seq_len) * self.seq_len
self.blocks = self.raw_ids[:n].view(-1, self.seq_len) if n > 0 else torch.zeros((1, self.seq_len), dtype=torch.long)
def set_seq_len(self, new_seq_len):
if new_seq_len == self.seq_len:
return
self.seq_len = new_seq_len
self.rebuild_blocks()
def __len__(self):
return len(self.blocks)
def __getitem__(self, i):
x = self.blocks[i]
return x[:-1], x[1:] # input, target (shifted)
def collect_txt(data_dir):
files = sorted(Path(data_dir).rglob("*.txt"))
# eval file ko train se bahar rakho
files = [f for f in files if "eval" not in [p.lower() for p in f.parts]]
if not files:
raise FileNotFoundError(f"{data_dir} me koi .txt nahi mili")
return files
# ---------------- curriculum & sched ----------------
def resolve_curriculum_stage(step, stages):
"""Return (stage_idx, stage_dict) based on current step."""
if not stages:
return 0, {}
for idx, st in enumerate(stages):
m = st.get("max_step")
if m is None or step < m:
return idx, st
return len(stages) - 1, stages[-1]
def fmt_hms(sec):
sec = max(int(sec), 0)
h, sec = divmod(sec, 3600)
m, s = divmod(sec, 60)
return f"{h}:{m:02d}:{s:02d}" if h else f"{m}:{s:02d}"
def cosine_lr(step, max_steps, warmup, lr, min_lr):
if step < warmup:
return lr * step / max(warmup, 1)
p = (step - warmup) / max(max_steps - warmup, 1)
return min_lr + 0.5 * (lr - min_lr) * (1 + math.cos(math.pi * min(p, 1.0)))
# ---------------- Muon Optimizer (Keller Jordan / Kimi K2) ----------------
def zeropower_via_newtonschulz5(G, steps=5, eps=1e-7):
"""
Newton-Schulz quintic iteration to approximate matrix polar decomposition.
Stabilizes and orthogonalizes momentum updates for 2D weight matrices.
"""
assert len(G.shape) == 2
a, b, c = (3.4445, -4.7750, 2.0315)
X = G.bfloat16() if (G.is_cuda and torch.cuda.is_bf16_supported()) else G.float()
X = X / (X.norm() + eps)
transposed = False
if X.size(0) > X.size(1):
X = X.T
transposed = True
for _ in range(steps):
A = X @ X.T
B = b * A + c * A @ A
X = a * X + B @ X
if transposed:
X = X.T
return X.to(dtype=G.dtype)
class Muon(torch.optim.Optimizer):
"""
Muon (MomentUm Orthogonalized by Newton-schulz) Optimizer.
Applies orthogonalized momentum updates to 2D internal weight matrices.
"""
def __init__(self, params, lr=0.02, momentum=0.95, weight_decay=0.01, ns_steps=5):
defaults = dict(lr=lr, momentum=momentum, weight_decay=weight_decay, ns_steps=ns_steps)
super().__init__(params, defaults)
@torch.no_grad()
def step(self, closure=None):
loss = None
if closure is not None:
with torch.enable_grad():
loss = closure()
for group in self.param_groups:
lr = group["lr"]
momentum = group["momentum"]
wd = group["weight_decay"]
steps = group["ns_steps"]
for p in group["params"]:
if p.grad is None:
continue
g = p.grad
state = self.state[p]
# Initialize momentum buffer
if "momentum_buffer" not in state:
state["momentum_buffer"] = torch.zeros_like(g)
buf = state["momentum_buffer"]
buf.mul_(momentum).add_(g)
# Orthogonalize via Newton-Schulz
u = zeropower_via_newtonschulz5(buf, steps=steps)
# Scale update based on aspect ratio
scale = max(1.0, math.sqrt(p.size(0) / p.size(1)))
u.mul_(scale)
# Weight decay
if wd > 0:
p.mul_(1.0 - lr * wd)
# Apply update
p.add_(u, alpha=-lr)
return loss
class CombinedOptimizer:
"""Combines Muon (for 2D linear weights) and AdamW (for 1D norms, embeddings, routers)."""
def __init__(self, optimizers):
self.optimizers = optimizers
self.param_groups = []
for opt in optimizers:
self.param_groups.extend(opt.param_groups)
def step(self, closure=None):
loss = None
for opt in self.optimizers:
res = opt.step(closure=closure)
if res is not None:
loss = res
return loss
def zero_grad(self, set_to_none=True):
for opt in self.optimizers:
opt.zero_grad(set_to_none=set_to_none)
@property
def state(self):
# Merged param->state view so torch AMP GradScaler.unscale_() works
# (it reads optimizer.state[p]). FP16-only path; BF16 configs never touch it.
merged = {}
for opt in self.optimizers:
try:
merged.update(opt.state)
except Exception:
pass
return merged
def state_dict(self):
return {"opts": [opt.state_dict() for opt in self.optimizers]}
def load_state_dict(self, state_dict):
if isinstance(state_dict, dict) and "opts" in state_dict:
for opt, sd in zip(self.optimizers, state_dict["opts"]):
opt.load_state_dict(sd)
elif len(self.optimizers) > 1:
try:
self.optimizers[1].load_state_dict(state_dict)
except Exception:
pass
def split_muon_params(model):
"""
Separates parameters into:
- Muon: 2D internal linear weights in MLA attention and MoE experts
- AdamW: Embeddings, RMSNorms, routers, biases, output heads
"""
muon_params = []
adamw_params = []
for name, p in model.named_parameters():
if not p.requires_grad:
continue
# 1D params (norms, biases) always AdamW
if p.ndim < 2:
adamw_params.append(p)
# Embeddings, router heads, and unembedding heads use AdamW
elif "tok_emb" in name or "head" in name or "router" in name:
adamw_params.append(p)
else:
muon_params.append(p)
return muon_params, adamw_params
def pick_device(pref=None):
"""Auto GPU: sabse zyada free VRAM wala cuda, warna cpu. --device se override."""
if pref:
return torch.device(pref)
if not torch.cuda.is_available():
return torch.device("cpu")
n = torch.cuda.device_count()
if n == 1:
return torch.device("cuda:0")
best, best_free = 0, -1
for i in range(n):
try:
free, _ = torch.cuda.mem_get_info(i)
except Exception:
free = -1
if free > best_free:
best, best_free = i, free
print(f"[auto] {n} GPUs mile — cuda:{best} (max free VRAM), baaki idle (DDP future)")
return torch.device(f"cuda:{best}")
def gpu_caps(device):
"""(name, total_GB, bf16_ok) — T4 (cc 7.5) me BF16 nahi, Ampere+ me hai."""
try:
idx = device.index if device.index is not None else torch.cuda.current_device()
name = torch.cuda.get_device_name(idx)
total = torch.cuda.get_device_properties(idx).total_memory / (1024 ** 3)
major, _ = torch.cuda.get_device_capability(idx)
return name, total, major >= 8
except Exception:
return "unknown", 0.0, False
def resolve_amp(want, device, bf16_ok):
"""amp: auto/bf16/fp16/none → (use_bf16, use_fp16). Galat combo pe warn + safe fallback."""
if device.type != "cuda":
return False, False
w = str(want or "auto").lower()
if w == "auto":
return (True, False) if bf16_ok else (False, True)
if w == "bf16":
if bf16_ok:
return True, False
print("[auto] WARN: is GPU me BF16 nahi (T4-class) → fp16 fallback")
return False, True
if w == "fp16":
return False, True
return False, False
def probe_micro_batch(model, seq_len, vocab_size, device, amp_dtype):
"""Asli OOM-test: bada micro-batch try karo, jo fite wahi lo. Target eff-batch same rehta hai.
NOTE: probe DDP wrap SE PEHLE chalta hai — isliye AdamW (~2x params) + DDP grad buckets (~1x params)
+ fragmentation ka explicit reserve chahiye, warna wrap ke baad training step-1 pe OOM."""
model.train()
best = 1
# Reserve: Muon fp32 momentum (~1.7B 2D params ~= 6.8GB) + AdamW8bit states +
# DDP grad buckets + fragmentation. 1.856B model => try 8GB, warna 4GB, warna 2GB.
adamw_buf = None
if device.type == "cuda":
for n_elems in (2_000_000_000, 1_000_000_000, 500_000_000): # ~8.0 / 4.0 / 2.0 GB
try:
adamw_buf = torch.empty((n_elems,), dtype=torch.float32, device=device)
break
except RuntimeError:
adamw_buf = None
try:
for mb in (1, 2, 4, 8, 16, 32):
try:
torch.cuda.empty_cache()
x = torch.randint(0, vocab_size, (mb, seq_len), device=device)
y = torch.randint(0, vocab_size, (mb, seq_len), device=device)
model.zero_grad(set_to_none=True)
if amp_dtype is not None:
ctx = torch.amp.autocast("cuda", dtype=amp_dtype)
else:
ctx = torch.amp.autocast("cpu", enabled=False)
with ctx:
_, loss, _ = model(x, y)
loss.backward()
del x, y, loss
model.zero_grad(set_to_none=True)
best = mb
except RuntimeError as e:
torch.cuda.empty_cache()
if "out of memory" in str(e).lower():
break
raise
finally:
del adamw_buf
torch.cuda.empty_cache()
return best
def maybe_apply_fp8(model, want, device, R0=True):
"""EXPERIMENTAL 2025-tech (Blackwell/5090 only): torchao Float8 GEMM training.
Default OFF. Optimizer (Muon/AdamW) original-dtype params dekhta hai, isliye
compatible hai. Koi bhi condition miss -> loud BF16 fallback, silent nahi."""
if not want:
return model, "fp8=off"
if device.type != "cuda":
if R0:
print("[fp8] CPU box — BF16 fallback (FP8 needs Blackwell CUDA)")
return model, "fp8=off(cpu)"
try:
major, _ = torch.cuda.get_device_capability(device)
except Exception:
major = 0
if major < 9:
if R0:
print("[fp8] pre-Blackwell GPU (sm<90) — BF16 fallback (FP8 needs sm_90+)")
return model, "fp8=off(arch)"
try:
from torchao.float8 import convert_to_float8_training
except ImportError:
if R0:
print("[fp8] torchao missing (pip install torchao --index-url https://download.pytorch.org/whl/nightly/cu128) — BF16 fallback")
return model, "fp8=off(no-torchao)"
try:
convert_to_float8_training(model)
if R0:
print("[fp8] Float8 training ENABLED via torchao (EXPERIMENTAL — pehle 500 steps loss-watch karo)")
return model, "fp8=on"
except Exception as e:
if R0:
print(f"[fp8] convert failed ({str(e)[:120]}) — BF16 fallback")
return model, "fp8=off(err)"
def maybe_apply_compile(model, want, mode, is_smoke, R0=True, device=None, vocab_size=48000):
"""torch.compile (2024-25 mature): 10-30% free speedup. Smoke me skip (fast test).
DDP-wrap SE PEHLE lagta hai taaki ckpt keys prefix-free rahein.
Fail-FAST: chhota warmup forward backend verify karta hai (missing triton/compiler
pehle step pe crash ke bajaye turant eager fallback). Pehla compiled step slow
hota hai (one-time codegen, minutes) — ye normal hai."""
if not want:
return model, "compile=off"
if is_smoke:
if R0:
print("[compile] smoke me skip (fast CPU test chahiye)")
return model, "compile=off(smoke)"
try:
model = torch.compile(model, mode=mode, fullgraph=False)
if device is not None and str(device) != "cpu":
with torch.no_grad():
_ids = torch.randint(0, vocab_size, (1, 16), device=device)
_ = model(_ids)
if R0:
print(f"[compile] torch.compile ON (mode={mode}, fullgraph=False, backend-verified, DDP-se-pehle)")
return model, f"compile=on({mode})"
except Exception as e:
if R0:
print(f"[compile] backend unavailable ({str(e)[:100]}) — eager fallback (training safe hai)")
try:
# compiled wrapper hata ke original module wapas (eager, safe)
inner = model._orig_mod if hasattr(model, "_orig_mod") else model
if inner is not model:
return inner, "compile=off(err-fallback)"
except Exception:
pass
return model, "compile=off(err)"
def save_slim_ckpt(path, model, cfg=None):
"""Inference & HF deployment ke liye slim checkpoint (~3.7GB BF16 for 1.856B, optimizer states ke bina)."""
path = Path(path)
path.parent.mkdir(parents=True, exist_ok=True)
state = {k: v.to(dtype=torch.bfloat16).cpu() for k, v in model.state_dict().items()}
torch.save({"model": state, "cfg": cfg}, path)
def save_ckpt(path, model, opt, step, cfg, scaler=None):
path.parent.mkdir(parents=True, exist_ok=True)
ck = {"step": step, "model": model.state_dict(),
"opt": opt.state_dict(), "cfg": cfg}
if scaler is not None:
ck["scaler"] = scaler.state_dict()
torch.save(ck, path)
def load_ckpt(path, model, opt, scaler=None):
# torch>=2.6 me weights_only default True hai — hamare ckpt me cfg dict hai isliye False chahiye
try:
ck = torch.load(path, map_location="cpu", weights_only=False)
except TypeError:
ck = torch.load(path, map_location="cpu")
model.load_state_dict(ck["model"])
if opt is not None and "opt" in ck:
try:
opt.load_state_dict(ck["opt"])
except Exception as e:
print(f"[warn] opt state restore skip: {e}")
if scaler is not None and "scaler" in ck:
try:
scaler.load_state_dict(ck["scaler"])
except Exception as e:
print(f"[warn] scaler state restore skip: {e}")
return ck.get("step", 0)
def hf_push_ckpt(ckpt_dir, repo_id, token, path_in_repo="model/checkpoints", only_file=None, as_name=None, msg="ckpt"):
"""Har training ke baad ckpt backup mandatory — sirf FINAL + distinct naam. Fail = WARN, crash nahi."""
try:
from huggingface_hub import HfApi
except ImportError:
print("[hf-push] skip (huggingface_hub nahi hai)")
return False
if not token:
print("[hf-push] skip (HF_TOKEN env me nahi)")
return False
try:
api = HfApi(token=token)
if only_file is not None:
api.upload_file(path_or_fileobj=str(only_file), path_in_repo=f"{path_in_repo}/{as_name or Path(only_file).name}",
repo_id=repo_id, commit_message=msg)
else:
api.upload_folder(folder_path=str(ckpt_dir), repo_id=repo_id, path_in_repo=path_in_repo,
commit_message=msg, ignore_patterns=["*.tmp", "*.lock"])
print(f"[hf-push] ok → {repo_id}/{path_in_repo}")
return True
except Exception as e:
print(f"[hf-push] WARN fail (training safe hai): {str(e)[:200]}")
return False
class AsyncCheckpointManager:
"""Non-blocking background checkpoint manager for Hugging Face upload and disk maintenance.
Ensures GPU training loop never pauses or blocks during remote checkpoint sync.
"""
def __init__(self, repo_id, token, ckpt_dir=None, log_dir=None, enabled=True, keep_local_full=2, keep_remote_full=2):
self.repo_id = repo_id
self.token = token
self.ckpt_dir = Path(ckpt_dir) if ckpt_dir else None
self.log_dir = Path(log_dir) if log_dir else None
self.enabled = bool(enabled and repo_id and token)
self.keep_local_full = keep_local_full
self.keep_remote_full = keep_remote_full
self.executor = concurrent.futures.ThreadPoolExecutor(max_workers=1, thread_name_prefix="hf_ckpt_uploader") if self.enabled else None
self.active_future = None
def prune_local_full_ckpts(self, current_step=None):
"""Keep only newest N full checkpoints locally to prevent Kaggle 20GB disk overflow."""
if not self.ckpt_dir or not self.ckpt_dir.exists():
return
full_files = []
for f in self.ckpt_dir.glob("step_*.pt"):
m = re.match(r"^step_(\d+)\.pt$", f.name)
if m:
full_files.append((int(m.group(1)), f))
full_files.sort(key=lambda x: x[0])
if len(full_files) > self.keep_local_full:
for s, f in full_files[:-self.keep_local_full]:
try:
f.unlink()
print(f"[ckpt-local] pruned old local full ckpt: {f.name} (kept newest {self.keep_local_full})", flush=True)
except Exception as e:
print(f"[ckpt-local] WARN deleting {f.name}: {e}", flush=True)
def push_checkpoint_async(self, step, full_file=None, slim_file=None, log_file=None):
if not self.enabled or self.executor is None:
return None
def _worker():
t_start = time.time()
try:
from huggingface_hub import HfApi
api = HfApi(token=self.token)
# 1. Upload slim checkpoint if provided
if slim_file and Path(slim_file).exists():
p = Path(slim_file)
print(f"[hf-async] starting background upload: {p.name}...", flush=True)
api.upload_file(path_or_fileobj=str(p),
path_in_repo=f"model/checkpoints/{p.name}",
repo_id=self.repo_id,
commit_message=f"slim ckpt step_{step}")
print(f"[hf-async] {p.name} uploaded successfully", flush=True)
# 2. Upload full checkpoint if provided
if full_file and Path(full_file).exists():
p = Path(full_file)
print(f"[hf-async] starting background upload: {p.name} (~GBs)...", flush=True)
api.upload_file(path_or_fileobj=str(p),
path_in_repo=f"model/checkpoints/{p.name}",
repo_id=self.repo_id,
commit_message=f"full ckpt step_{step} (resume)")
print(f"[hf-async] {p.name} uploaded successfully", flush=True)
# 3. Clean up older full checkpoints on remote HF repo
try:
files = api.list_repo_files(self.repo_id, repo_type="model")
def _get_step(fn):
m = re.search(r"model/checkpoints/step_(\d+)\.pt$", fn)
return int(m.group(1)) if m else -1
full_ckpts = sorted([f for f in files if _get_step(f) >= 0], key=_get_step)
if len(full_ckpts) > self.keep_remote_full:
for old_ckpt in full_ckpts[:-self.keep_remote_full]:
api.delete_file(path_in_repo=old_ckpt, repo_id=self.repo_id,
repo_type="model", commit_message=f"prune old ckpt {old_ckpt}")
print(f"[hf-async] pruned remote old ckpt: {old_ckpt}", flush=True)
except Exception as e_prune:
print(f"[hf-async] WARN remote repo pruning: {str(e_prune)[:120]}", flush=True)
# 4. Upload live log file to HF for remote tracking
if log_file and Path(log_file).exists():
try:
api.upload_file(path_or_fileobj=str(log_file),
path_in_repo="logs/train.jsonl",
repo_id=self.repo_id,
commit_message=f"sync log step_{step}")
except Exception as e_log:
print(f"[hf-async] WARN log upload: {str(e_log)[:120]}", flush=True)
elapsed = time.time() - t_start
print(f"[hf-async] completed step {step} background sync in {elapsed:.1f}s (training was never paused)", flush=True)
except Exception as e:
print(f"[hf-async] WARN background upload step {step} failed: {str(e)[:200]} (training safe, local file intact)", flush=True)
self.active_future = self.executor.submit(_worker)
return self.active_future
def wait_pending(self, timeout=None):
"""Wait for any active background upload to complete (e.g. before exiting training)."""
if self.executor is not None:
self.executor.shutdown(wait=True)
@torch.no_grad()
def evaluate(model, ds, device, max_batches=20, use_amp=False, amp_dtype=torch.float16):
model.eval()
tot, n = 0.0, 0
if use_amp and device.type == "cuda":
ctx = torch.amp.autocast("cuda", dtype=amp_dtype)
else:
ctx = torch.amp.autocast("cpu", enabled=False)
for i in range(min(len(ds), max_batches)):
x, y = ds[i]
x, y = x.unsqueeze(0).to(device), y.unsqueeze(0).to(device)
with ctx:
_, loss, aux = model(x, y)
# val PPL ke liye pure LM loss chahiye — MoE-aux + mt-loss hatao
base = loss.item()
try:
mt_w = float(aux.get("mt_loss_weight", 0.1)) # match viu_moe.py forward()
base = base - float(aux.get("moe_aux", 0.0)) - mt_w * float(aux.get("mt_loss", 0.0))
except Exception:
pass
tot += base
n += 1
model.train()
avg = tot / max(n, 1)
return avg, math.exp(min(avg, 20))
@torch.no_grad()
def sample_text(model, tok, prompt, device, max_new=30, seq_len=1024, temperature=0.7, top_k=40):
model.eval()
ids = tok.encode(prompt).ids[-64:]
out = torch.tensor([ids], dtype=torch.long, device=device)
for _ in range(max_new):
ctx = out[:, -seq_len:]
logits, _, _ = model(ctx)
nxt_logits = logits[:, -1, :] / max(temperature, 1e-5)
if top_k > 0:
v, _ = torch.topk(nxt_logits, min(top_k, nxt_logits.size(-1)))
nxt_logits[nxt_logits < v[:, [-1]]] = float("-inf")
probs = F.softmax(nxt_logits, dim=-1)
nxt = torch.multinomial(probs, num_samples=1)
out = torch.cat([out, nxt], dim=1)
model.train()
return tok.decode(out[0].tolist())
def main():
# Fragmentation kam karo (T4 pe large-alloc OOM ke baad reserved-but-unused dikhta tha)
os.environ.setdefault("PYTORCH_ALLOC_CONF", "expandable_segments:True")
ap = argparse.ArgumentParser()
ap.add_argument("--config", default="../configs/train_rtx5090.yaml")
ap.add_argument("--model_config", default="../configs/model_config.yaml")
ap.add_argument("--smoke", action="store_true", help="CPU pe tiny model + 3 steps test")
ap.add_argument("--steps", "--max_steps", dest="steps", type=int, default=None, help="Max training steps override")
ap.add_argument("--resume", default=None, help="Checkpoint path to resume training from")
ap.add_argument("--hf_dataset", default=None, help="HF dataset repo for streaming, e.g. ViuAI/viu-mini-pretrain-40-30-30")
ap.add_argument("--device", default=None, help="cuda / cuda:1 / cpu (default: auto = sabse zyada free VRAM wala GPU)")
ap.add_argument("--dist", action="store_true", help="DDP multi-GPU (torchrun se chalao). Bina iske single-GPU.")
ap.add_argument("--hf_safety", default=None, choices=["strict", "research", "full"], help="Safety filter: strict, research, or full")
ap.add_argument("--hf_split", default="train", help="Dataset split (default: train)")
ap.add_argument("--batch_size", type=int, default=None, help="Micro-batch size override")
ap.add_argument("--grad_accum", type=int, default=None, help="Gradient accumulation steps override")
ap.add_argument("--seq_len", type=int, default=None, help="Sequence length override (bypasses curriculum)")
ap.add_argument("--curriculum", dest="curriculum", action="store_true", default=None, help="Enable 3-stage sequence length curriculum (512 -> 1024 -> 2048)")
ap.add_argument("--no_curriculum", "--no-curriculum", dest="curriculum", action="store_false", help="Disable sequence length curriculum")
ap.add_argument("--no_auto_tune", action="store_true", help="Disable auto-tune and enforce exact batch_size and grad_accum")
ap.add_argument("--optimizer", default=None, choices=["muon", "adamw", "paged_adamw_8bit", "8bit"], help="Optimizer choice (muon, adamw, paged_adamw_8bit)")
ap.add_argument("--muon_lr", type=float, default=None, help="Muon learning rate override (default: 0.02)")
ap.add_argument("--muon_weight_decay", type=float, default=None, help="Muon weight decay override (default: 0.01)")
ap.add_argument("--compile", action="store_true", help="torch.compile on (10-30%% free speedup, default off)")
ap.add_argument("--fp8", action="store_true", help="torchao Float8 training (Blackwell/5090 only, EXPERIMENTAL)")
args = ap.parse_args()
base = Path(__file__).parent
cfg = {}
if yaml and Path(args.config).exists():
cfg = yaml.safe_load(open(args.config, encoding="utf-8")) or {}
cfg_dir = Path(args.config).parent
elif (base / args.config).exists() and yaml:
cfg = yaml.safe_load(open(base / args.config, encoding="utf-8")) or {}
cfg_dir = (base / args.config).parent
else:
cfg_dir = base
def C(k, d):
return cfg.get(k, d)
def resolve(p, default_base=cfg_dir): # config me relative paths ko config-file ya script se resolve karo (CWD se nahi)
pp = Path(p)
if pp.is_absolute():
return pp
# 1) config dir ke relative, 2) script base ke relative, 3) CWD
for cand in [cfg_dir / pp, base / pp, Path.cwd() / pp]:
if cand.exists():
return cand.resolve()
return (cfg_dir / pp).resolve()
# Early rank (startup prints R0-only chahte hain; DDP init non-smoke block me hoga)
rank = int(os.environ.get("RANK", "0"))
world = int(os.environ.get("WORLD_SIZE", "1"))
use_dist, R0 = False, (rank == 0)
if args.smoke:
print("[mode] SMOKE — tiny model, CPU, 3 steps")
margs = ViuArgs(dim=64, n_layers=2, n_heads=4, n_kv_heads=2, vocab_size=512,
max_seq_len=64, sliding_window=32, moe_every=2,
moe_num_experts=4, moe_num_routed_experts=4, moe_shared_experts=1,
multi_token=False, switch_gate=False)
seq_len, batch, max_steps = 32, 2, 3
data_dir, tok_path = base / "../../data/raw", None
lr, min_lr, warmup, wd = 3e-4, 3e-5, 1, 0.01
accum, clip, eval_every, save_every, log_every = 1, 1.0, 10, 100, 1
sample_every = 1000
ckpt_dir, log_dir = base / "../checkpoints", base / "../../logs"
device = torch.device("cpu")
use_bf16 = False
use_fp16 = False
use_amp = False
amp_dtype = torch.float32
use_curriculum = False
curriculum_stages = None
curr_stage_idx = 0
else:
margs = load_model_args(args.model_config if Path(args.model_config).exists()
else str(base / args.model_config))
# Sequence Length Curriculum Setup (512 -> 1024 -> 2048)
use_curriculum = bool(C("curriculum", False))
if args.curriculum is not None:
use_curriculum = bool(args.curriculum)
if args.seq_len is not None:
use_curriculum = False # explicit --seq_len overrides curriculum
curriculum_stages = C("curriculum_stages", [
{"max_step": 50000, "seq_len": 512, "batch_size": 8, "grad_accum": 64},
{"max_step": 200000, "seq_len": 1024, "batch_size": 8, "grad_accum": 32},
{"max_step": None, "seq_len": 2048, "batch_size": 4, "grad_accum": 32},
]) if use_curriculum else None
curr_stage_idx = 0
if use_curriculum:
curr_stage_idx, curr_stage = resolve_curriculum_stage(0, curriculum_stages)
seq_len = int(curr_stage["seq_len"])
batch = int(args.batch_size or curr_stage["batch_size"])
accum = int(args.grad_accum or curr_stage["grad_accum"])
if R0:
print(f"[curriculum] Initialized Stage 1/{len(curriculum_stages)}: "
f"seq_len={seq_len} batch={batch} accum={accum} (target {batch*accum*seq_len:,} tok/step)")
else:
seq_len = int(args.seq_len or C("seq_len", 1024))
batch = int(args.batch_size or C("batch_size", 2))
accum = int(args.grad_accum or C("grad_accum", 8))
_cfg_steps = C("max_steps", None)
max_steps = args.steps or (int(_cfg_steps) if _cfg_steps not in ("auto", "none", "null", None) else None)
data_dir = resolve(C("data_dir", "../../data/raw"))
tok_path = resolve(C("tokenizer_path", "../../tokenizer/outputs/tokenizer.json"))
lr, min_lr, warmup, wd = C("lr", 3e-4), C("min_lr", 3e-5), C("warmup_steps", 2500), C("weight_decay", 0.1)
clip = float(C("grad_clip", 1.0))
eval_every, save_every, log_every = C("eval_every", 1000), C("save_every", 5000), C("log_every", 50)
sample_every = C("sample_every", 1000)
ckpt_dir = resolve(C("ckpt_dir", "../checkpoints"))
log_dir = resolve(C("log_dir", "../../logs"))
hf_repo = C("hf_repo", "ViuAI/Viu-1.5B-MoE")
hf_token = os.environ.get("HF_TOKEN")
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
if args.device:
device = torch.device(args.device)
else:
device = pick_device()
# DDP: torchrun --nproc_per_node=2 se chalao; auto-detect bhi (WORLD_SIZE>1).
# NOTE: smoke me DDP kabhi nahi (CPU tiny test).
# rank/world env se hi rakho (early R0 ke saath consistent) — init ke waqt overwrite mat karo.
local_rank = int(os.environ.get("LOCAL_RANK", "0"))
use_dist = bool(args.dist or world > 1)
if use_dist and (not torch.cuda.is_available() or world < 2):
if R0:
print("[dist] WARN: multi-GPU env nahi (torchrun bina / 1 GPU) — single-GPU fallback")
use_dist = False
rank, world = 0, 1
if use_dist:
import torch.distributed as dist
import warnings
warnings.filterwarnings("ignore", message=".*find_unused_parameters.*")
warnings.filterwarnings("ignore", category=UserWarning, module="torch.distributed")
torch.cuda.set_device(local_rank)
device = torch.device(f"cuda:{local_rank}")
try:
dist.init_process_group(backend="nccl", rank=rank, world_size=world, device_id=device)
except TypeError:
dist.init_process_group(backend="nccl", rank=rank, world_size=world)
if use_dist and not args.smoke:
print(f"[dist] rank {rank}/{world} on {device} (DDP, grad-sync every step)")
R0 = (rank == 0)
gname, ggb, bf16_ok = gpu_caps(device) if device.type == "cuda" else ("cpu", 0.0, False)
use_bf16, use_fp16 = resolve_amp(C("amp", "auto"), device, bf16_ok)
use_amp = use_bf16 or use_fp16
amp_dtype = torch.bfloat16 if use_bf16 else torch.float16
if device.type == "cuda":
torch.backends.cudnn.benchmark = True
torch.backends.cuda.matmul.allow_tf32 = True
torch.backends.cudnn.allow_tf32 = True
random.seed(C("seed", 42) if not args.smoke else 42)
torch.manual_seed(42)
# tokenizer
tok = None
if Tokenizer is not None and tok_path is not None:
p_tok = Path(str(tok_path))
needs_download = False
if not p_tok.exists() or p_tok.stat().st_size < 10000:
needs_download = True
else:
try:
tok = Tokenizer.from_file(str(p_tok))
except Exception:
needs_download = True
if needs_download:
if R0:
print(f"[tok] {tok_path} missing or LFS pointer — auto-downloading valid tokenizer.json from Hugging Face...", flush=True)
from huggingface_hub import hf_hub_download
import shutil
try:
real_tok = hf_hub_download(
repo_id=hf_repo or "ViuAI/Viu-1.5B-MoE",
filename="tokenizer/outputs/tokenizer.json",
token=hf_token
)
p_tok.parent.mkdir(parents=True, exist_ok=True)
shutil.copy(real_tok, str(p_tok))
tok = Tokenizer.from_file(str(p_tok))
except Exception as e:
if R0:
print(f"[tok] fallback to dataset repo for tokenizer: {e}")
real_tok = hf_hub_download(
repo_id="ViuAI/viu-mini-raw-pretrain",
filename="tokenizer/tokenizer.json",
repo_type="dataset",
token=hf_token
)
p_tok.parent.mkdir(parents=True, exist_ok=True)
shutil.copy(real_tok, str(p_tok))
tok = Tokenizer.from_file(str(p_tok))
if tok is not None and R0:
print(f"[tok] loaded {p_tok} vocab={tok.get_vocab_size()}")
elif not args.smoke:
if R0:
print(f"[warn] tokenizer nahi mila ({tok_path}) — FINAL tokenizer Phase-2 me banega.")
print("[error] real training ke liye FINAL tokenizer chahiye. --smoke se test karo ya tokenizer train karo.")
return
if args.smoke or tok is None:
# smoke/dummy: sample data se chhota BPE on-the-fly? simple: train.py smoke me char-level dummy
pass
# data
# HF streaming path (Phase-1 ke baad prefer karo)
hf_dataset = args.hf_dataset or (None if args.smoke else C("hf_dataset", None))
ds = None
if hf_dataset and not args.smoke:
if tok is None:
if R0:
print("[error] HF dataset ke liye tokenizer chahiye — pehle tokenizer train karo.")
return
try:
if R0:
print(f"[hf] streaming dataset {hf_dataset} (native PyArrow streaming)")
# Non-5-col schema folders (AUDIT A3) — inko stream me mat lo, warna
# shuffle/interleave par CastError (wikipedia: title/text/url vs text/lang/...).
# Note: distilled/, science/, and translation/ (Samanantar) handled in _stream_parquet_rows.
# Exclude only truly unformatted raw directories (already packaged into domains/)
# and empty/non-parquet files. All 1,874 valid Parquet files stream natively.
# ZERO EXCLUSIONS: Stream 100% of all files in repository
NON_UNIFIED_PREFIXES = ()
# Corpus-slice rotation (Option A: full-data coverage round-by-round).
# hf_round badhte hi: (1) file-list rotate, (2) row-offset skip, (3) per-round seed.
# Har round ~alag slice prefetch hota hai; cell round = start_step // ROUND_SPAN se deta hai.
hf_round = int(C("hf_round", 0))
hf_round_rows = int(C("hf_round_rows", 200000))
hf_token = os.environ.get("HF_TOKEN")
try:
from huggingface_hub import list_repo_files
_all = list_repo_files(hf_dataset, repo_type="dataset", token=hf_token)
_raw_parquet = sorted([
f for f in _all
if f.endswith(".parquet") and not any(f.startswith(p) for p in NON_UNIFIED_PREFIXES)
])
_n_excl = sum(1 for f in _all if f.endswith(".parquet")) - len(_raw_parquet)
# Round-robin category interleave to prevent single-domain clumping during warmup
# and prioritize newly added high-yield frontier datasets (History, Cricket, QMS, Chat)
from collections import defaultdict
import random as _r
_cat_buckets = defaultdict(list)
for f in _raw_parquet:
_cat_buckets[f.split("/")[0]].append(f)
PRIORITY_PATTERNS = ("real_indian_history_pib", "real_cricket", "erotica_adult_stories", "real_cinema", "real_mythology", "real_finance", "real_industrial_qms", "conversational")
_rng_files = _r.Random(int(C("seed", 42)) + hf_round)
for cat in _cat_buckets:
_rng_files.shuffle(_cat_buckets[cat])
pri_newest = [f for f in _cat_buckets[cat] if any(p in f for p in ("real_indian_history_pib", "real_cricket", "erotica_adult_stories", "real_cinema", "real_mythology", "real_finance"))]
pri_other = [f for f in _cat_buckets[cat] if any(p in f for p in PRIORITY_PATTERNS) and f not in pri_newest]
rest = [f for f in _cat_buckets[cat] if f not in pri_newest and f not in pri_other]
_cat_buckets[cat] = pri_newest + pri_other + rest
_interleaved = []
_max_len = max(len(v) for v in _cat_buckets.values()) if _cat_buckets else 0
for i in range(_max_len):
for cat in sorted(_cat_buckets.keys()):
if i < len(_cat_buckets[cat]):
_interleaved.append(_cat_buckets[cat][i])
_data_files = _interleaved
if hf_round and _data_files:
rot = hf_round % len(_data_files)
_data_files = _data_files[rot:] + _data_files[:rot]
if R0:
print(f"[hf] data_files: {len(_data_files)} unified parquet (interleaved across {len(_cat_buckets)} categories, prioritized real history/cricket/QMS, excluded {_n_excl} non-5col)"
+ (f" round={hf_round} start={_data_files[0] if _data_files else 'none'}" if hf_round else ""))
import pyarrow.parquet as _pq
from huggingface_hub import HfFileSystem as _HfFS
import itertools as _it
_fs = _HfFS(token=hf_token)
def _stream_parquet_rows(files, repo_id):
for fp in files:
p = f"datasets/{repo_id}/{fp}"
for _attempt in range(3):
try:
with _fs.open(p, "rb") as f:
pf = _pq.ParquetFile(f)
_names = pf.schema.names
if "toxicity" in _names and "text" in _names:
# Jigsaw civil_comments: score columns -> safety_tag (label preserved as tag, not text)
cols = ["text", "toxicity", "severe_toxicity", "obscene", "threat", "insult", "identity_attack", "sexual_explicit"]
cols = [c for c in cols if c in _names]
for batch in pf.iter_batches(batch_size=2048, columns=cols):
d = batch.to_pydict()
n = len(d.get("text", []))
for i in range(n):
tw = str(d["text"][i])
if len(tw.strip()) <= 5:
continue
mx = 0.0
for sc in cols[1:]:
try:
v = float(d[sc][i])
except Exception:
v = 0.0
if v > mx:
mx = v
tag = "toxic" if mx >= 0.5 else "safe"
yield {"text": tw, "lang": "en", "source": "jigsaw_civil_comments", "domain": "toxicity", "safety_tag": tag}
elif "text" in _names:
cols = [c for c in ["text", "lang", "source", "domain", "safety_tag"] if c in pf.schema.names]
for batch in pf.iter_batches(batch_size=2048, columns=cols):
d = batch.to_pydict()
n = len(d.get("text", []))
for i in range(n):
yield {k: d[k][i] for k in cols}
elif "generated_full_text" in pf.schema.names:
for batch in pf.iter_batches(batch_size=2048, columns=["generated_full_text"]):
d = batch.to_pydict()
n = len(d.get("generated_full_text", []))
for i in range(n):
yield {"text": str(d["generated_full_text"][i]), "lang": "en", "source": "science_textbooks", "domain": "science", "safety_tag": "safe"}
elif "src" in pf.schema.names and "tgt" in pf.schema.names:
for batch in pf.iter_batches(batch_size=2048, columns=["src", "tgt"]):
d = batch.to_pydict()
n = len(d.get("src", []))
for i in range(n):
s = str(d["src"][i]).strip()
t = str(d["tgt"][i]).strip()
if len(s) > 10 and len(t) > 10:
if i % 2 == 0:
txt = f"English: {s}\nHindi: {t}"
else:
txt = f"Hindi: {t}\nEnglish: {s}"
yield {"text": txt, "lang": "bilingual", "source": "samanantar_translation", "domain": "translation", "safety_tag": "safe"}
elif "context" in pf.schema.names and "question" in pf.schema.names and "response" in pf.schema.names:
for batch in pf.iter_batches(batch_size=2048, columns=["context", "question", "response"]):
d = batch.to_pydict()
n = len(d.get("question", []))
for i in range(n):
ctx = str(d["context"][i]).strip()
q = str(d["question"][i]).strip()
ans = str(d["response"][i]).strip()
parts = []
if ctx and len(ctx) > 30:
parts.append(f"### Case Background:\n{ctx}")
parts.append(f"### Legal Query:\n{q}")
parts.append(f"### Legal Determination:\n{ans}")
yield {"text": "\n\n".join(parts), "lang": "en", "source": "indian_courts_case_law", "domain": "governance_law", "safety_tag": "safe"}
elif "instruction" in pf.schema.names and "output" in pf.schema.names:
cols = [c for c in ["instruction", "input", "output"] if c in pf.schema.names]
for batch in pf.iter_batches(batch_size=2048, columns=cols):
d = batch.to_pydict()
n = len(d["instruction"])
for i in range(n):
ins = str(d["instruction"][i]).strip()
inp = str(d["input"][i]).strip() if "input" in d else ""
ans = str(d["output"][i]).strip()
if len(ins) > 5 and len(ans) > 5:
if inp and len(inp) > 5:
txt = f"### Instruction:\n{ins}\n\n### Input:\n{inp}\n\n### Response:\n{ans}"
else:
txt = f"### Instruction:\n{ins}\n\n### Response:\n{ans}"
lg = "hinglish" if any('\u0900' <= ch <= '\u097F' for ch in txt[:500]) else "en"
yield {"text": txt, "lang": lg, "source": "instruction_following", "domain": "instruction", "safety_tag": "safe"}
elif "tweet" in pf.schema.names:
for batch in pf.iter_batches(batch_size=2048, columns=["tweet"]):
d = batch.to_pydict()
n = len(d.get("tweet", []))
for i in range(n):
tw = str(d["tweet"][i]).strip()
if len(tw) > 5:
yield {"text": tw, "lang": "en", "source": "hate_speech_offensive", "domain": "toxicity", "safety_tag": "toxic"}
elif "question" in pf.schema.names and "answer" in pf.schema.names:
# Standalone Q/A pairs (math/problem sets without context column)
for batch in pf.iter_batches(batch_size=2048, columns=["question", "answer"]):
d = batch.to_pydict()
n = len(d.get("question", []))
for i in range(n):
q = str(d["question"][i]).strip()
a = str(d["answer"][i]).strip()
if len(q) > 5 and len(a) > 5:
yield {"text": f"### Problem:\n{q}\n\n### Solution:\n{a}", "lang": "en", "source": "math_qa_pairs", "domain": "reasoning_math", "safety_tag": "safe"}
elif "Patient" in pf.schema.names and "Doctor" in pf.schema.names:
cols = [c for c in ["Description", "Patient", "Doctor"] if c in pf.schema.names]
for batch in pf.iter_batches(batch_size=2048, columns=cols):
d = batch.to_pydict()
n = len(d["Patient"])
for i in range(n):
desc = str(d["Description"][i]).strip() if "Description" in d else ""
pat = str(d["Patient"][i]).strip()
doc = str(d["Doctor"][i]).strip()
if pat and doc:
txt = f"### Medical Case:\n{desc}\n\n### Patient:\n{pat}\n\n### Doctor:\n{doc}" if desc else f"### Patient:\n{pat}\n\n### Doctor:\n{doc}"
yield {"text": txt, "lang": "en", "source": "ai_medical_dialogues", "domain": "health", "safety_tag": "safe"}
break # success — move to next file
except GeneratorExit:
return # caller closed the generator
except Exception as e:
if _attempt < 2:
import time as _time
_time.sleep(2 ** _attempt) # backoff: 1s, 2s
else:
if R0:
print(f"[hf] WARNING: skipping {fp} after 3 retries: {e}")
def _stream_shuffled(gen, buffer_size=20000, seed=42):
if buffer_size <= 0:
yield from gen
return
import random as _r
_rng = _r.Random(seed)
buf = []
for row in gen:
buf.append(row)
if len(buf) >= buffer_size:
idx = _rng.randint(0, len(buf) - 1)
yield buf.pop(idx)
_rng.shuffle(buf)
for row in buf:
yield row
raw_stream = _stream_parquet_rows(_data_files, hf_dataset)
seed_r = int(C("seed", 42)) + hf_round
skip_n = max(hf_round, 0) * max(hf_round_rows, 0)
if skip_n > 0:
raw_stream = _it.islice(raw_stream, skip_n, None)
natural_shuffle = int(C("natural_shuffle_buffer", 20000))
enforce_mix = bool(C("enforce_mix", False))
hf_ds = _stream_shuffled(raw_stream, buffer_size=natural_shuffle if not enforce_mix else 0, seed=seed_r)
except Exception as e:
if R0:
print(f"[hf] parquet stream setup ({e}) — fallback to direct single shard")
hf_ds = _stream_parquet_rows(["domains/train-00000.parquet"], hf_dataset)
# Language-mix rule OPTIONAL (default OFF = natural order).
# enforce_mix=true par weighted block sampler chalta hai (lang_mix weights se).
max_blocks_cfg = int(C("hf_max_blocks", 50000))
safety_mode = str(args.hf_safety or C("hf_safety", "full")).lower()
max_rows_cfg = C("hf_max_rows", None)
max_rows_cfg = int(max_rows_cfg) if max_rows_cfg else None
enforce_mix = bool(C("enforce_mix", False))
lang_mix = C("lang_mix", {"hinglish": 0.4, "hindi": 0.3, "english": 0.3}) or {}
mix_seed = int(C("seed", 42))
mix_scan_mult = int(C("mix_scan_mult", 5))
mix_max_rows = int(C("mix_max_rows", 2000000))
allow_toxic = safety_mode in ("research", "full")
allow_uncensored = safety_mode in ("full",)
allow_erotica = safety_mode in ("full",)
if R0:
print(f"[hf] safety={safety_mode} (toxic={allow_toxic}, uncensored={allow_uncensored}, erotica={allow_erotica}) "
f"max_blocks={max_blocks_cfg} max_rows={max_rows_cfg or 'unlimited'} "
f"enforce_mix={enforce_mix} lang_mix={lang_mix if enforce_mix else 'OFF (shuffle-buffer order)'}"
+ (f" natural_shuffle={natural_shuffle}" if (not enforce_mix and natural_shuffle > 0) else "")
+ (f" round={hf_round} skip_rows={skip_n:,}" if hf_round else ""))
try:
from data.scripts.filter_by_safety import keep_row as _keep
def _safety_ok(r):
return _keep(r, allow_uncensored=allow_uncensored, allow_toxic=allow_toxic)
except Exception:
def _safety_ok(r): # fallback: legacy rows bina safety_tag ke
tag = str(r.get("safety_tag") or "safe").lower()
if tag == "safe":
return True
if tag == "toxic":
return allow_toxic
if tag == "uncensored":
return allow_uncensored
if tag in ("sexual_explicit", "erotica", "erotic"):
return allow_erotica
s = f"{r.get('source','')} {r.get('domain','')}".lower()
if "uncensor" in s or "unfilter" in s or "vortex" in s:
return allow_uncensored
if "erotic" in s:
return allow_erotica
if any(k in s for k in ("toxic", "hate", "offensive", "profan", "civil_comments", "hasoc", "prism")):
return allow_toxic
return True
class HFStreamDataset(Dataset):
def __init__(self, stream, tok, seq_len, max_blocks=50000, max_rows=None,
enforce_mix=False, lang_mix=None, mix_seed=42, mix_scan_mult=5,
mix_max_rows=2000000):
self.stream = stream
self.tok = tok
self.seq_len = seq_len
self.max_blocks = max_blocks
self.max_rows = max_rows
self.enforce_mix = enforce_mix
self.lang_mix = lang_mix
self.mix_seed = mix_seed
self.mix_scan_mult = mix_scan_mult
self.mix_max_rows = mix_max_rows
self.blocks = torch.zeros((0, seq_len), dtype=torch.long)
self.stream_exhausted = False
self.refill_count = 0
self._remainder = {"hinglish": [], "hindi": [], "english": []} # carry-over tokens across refill()
self.refill()
def set_seq_len(self, new_seq_len):
"""Dynamically adjust block size for curriculum training.
Re-chunks unconsumed tokens and remainder buffer without loss."""
if new_seq_len == self.seq_len:
return
self.seq_len = new_seq_len
# Collect unconsumed blocks back into remainder buffer
if hasattr(self, "blocks") and len(self.blocks) > 0:
unconsumed = self.blocks.flatten().tolist()
self._remainder["english"].extend(unconsumed)
self.blocks = torch.zeros((0, self.seq_len), dtype=torch.long)
self.refill()
def refill(self):
if self.stream_exhausted:
return 0
self.refill_count += 1
buckets = {"hinglish": list(self._remainder.get("hinglish", [])),
"hindi": list(self._remainder.get("hindi", [])),
"english": list(self._remainder.get("english", []))}
eos = self.tok.token_to_id("<eos>")
if eos is None:
eos = 2
cap_tok = self.max_blocks * self.seq_len
weights = {k: float((self.lang_mix or {}).get(k, 0)) for k in buckets}
if sum(weights.values()) <= 0:
weights = {"hinglish": 0.4, "hindi": 0.3, "english": 0.3}
tot_w = sum(weights.values())
quota = {k: cap_tok * w / tot_w for k, w in weights.items()}
scan_cap = cap_tok * (self.mix_scan_mult if self.enforce_mix else 1)
n_rows = n_kept = n_safe_skip = n_empty = n_noschema = 0
lang_rows = {"hinglish": 0, "hindi": 0, "english": 0}
rows_found = False
for row in self.stream:
rows_found = True
n_rows += 1
if self.max_rows and n_rows > self.max_rows:
break
txt = row.get("text", "") if isinstance(row, dict) else ""
if not txt or not str(txt).strip():
n_empty += 1
continue
txt = str(txt).strip()
# Auditor filter: skip degenerate colon strings (e.g. NCERT colon scrapes ': : : : :')
if len(txt) < 5 or (len(txt) < 40 and txt.count(':') >= 3 and txt.count(':') > len(txt) / 3):
n_empty += 1
continue
if not _safety_ok(row if isinstance(row, dict) else {}):
n_safe_skip += 1
continue
try:
tag = get_lang_tag(row.get("lang") or row.get("source") or txt)
except Exception:
n_noschema += 1
continue
bucket = "hinglish" if "hinglish" in tag else ("hindi" if "hindi" in tag else "english")
if self.enforce_mix and len(buckets[bucket]) >= quota[bucket]:
n_kept += 1
if all(len(buckets[k]) >= quota[k] for k in buckets):
break
if n_rows >= (self.max_rows or self.mix_max_rows):
break
continue
tag_id = self.tok.token_to_id(tag) if self.tok is not None else None
if tag_id is not None:
buckets[bucket].append(tag_id)
buckets[bucket].extend(self.tok.encode(txt).ids)
buckets[bucket].append(eos)
lang_rows[bucket] += 1
n_kept += 1
total_tok = sum(len(v) for v in buckets.values())
if self.enforce_mix:
if all(len(buckets[k]) >= quota[k] for k in buckets):
break
if n_rows >= (self.max_rows or self.mix_max_rows):
break
if total_tok >= scan_cap:
break
elif total_tok >= cap_tok:
break
total_read = sum(len(v) for v in buckets.values())
if not rows_found or total_read == 0:
self.stream_exhausted = True
return 0
if R0:
print(f"[hf] chunk {self.refill_count}: scanned={n_rows:,} kept={n_kept:,} mix_rows={lang_rows} "
f"skip_empty={n_empty:,} skip_safety({safety_mode})={n_safe_skip:,} skip_schema={n_noschema:,}")
import random as _rng
if self.enforce_mix:
per_lang_blocks = {}
for k, v in buckets.items():
arr = torch.tensor(v, dtype=torch.long)
n = (len(arr) // self.seq_len) * self.seq_len
self._remainder[k] = arr[n:].tolist() # carry over tail tokens
per_lang_blocks[k] = arr[:n].view(-1, self.seq_len) if n > 0 else torch.zeros((0, self.seq_len), dtype=torch.long)
avail = {k: len(v) for k, v in per_lang_blocks.items()}
quota_b = {k: int(quota[k] // self.seq_len) for k in avail}
take = {k: min(avail[k], quota_b[k]) for k in avail}
short = self.max_blocks - sum(take.values())
if short > 0:
left = {k: avail[k] - take[k] for k in avail}
tot_left = sum(left.values()) or 1
for k in avail:
extra = min(left[k], int(round(short * left[k] / tot_left)))
take[k] += extra
chosen = [per_lang_blocks[k][:take[k]] for k in avail if take[k] > 0]
if not chosen:
self.blocks = torch.zeros((1, self.seq_len), dtype=torch.long)
else:
allb = torch.cat(chosen, dim=0)
idx = list(range(len(allb)))
_rng.Random(self.mix_seed + self.refill_count).shuffle(idx)
self.blocks = allb[idx]
if R0:
print(f"[hf] mix ENFORCED target={weights}: blocks_taken={take} total={len(self.blocks)} (avail={avail})")
else:
# Save per-lang remainders before merging
for k in ["hinglish", "hindi", "english"]:
arr_k = torch.tensor(buckets[k], dtype=torch.long)
n_k = (len(arr_k) // self.seq_len) * self.seq_len
self._remainder[k] = arr_k[n_k:].tolist()
ids = buckets["hinglish"] + buckets["hindi"] + buckets["english"]
if R0:
print(f"[hf] mix NATURAL (chunk {self.refill_count}): prefetched {len(ids):,} tokens -> {len(ids)//max(self.seq_len,1):,} blocks")
ids = torch.tensor(ids, dtype=torch.long)
n = (len(ids) // self.seq_len) * self.seq_len
self.blocks = ids[:n].view(-1, self.seq_len) if n > 0 else torch.zeros((1, self.seq_len), dtype=torch.long)
return len(self.blocks)
def __len__(self):
return len(self.blocks)
def __getitem__(self, i):
x = self.blocks[i]
return x[:-1], x[1:]
ds = HFStreamDataset(hf_ds, tok, seq_len+1, max_blocks=max_blocks_cfg, max_rows=max_rows_cfg,
enforce_mix=enforce_mix, lang_mix=lang_mix, mix_seed=seed_r,
mix_scan_mult=mix_scan_mult, mix_max_rows=mix_max_rows)
except Exception as e:
if R0:
print(f"[hf] streaming failed ({e}), fallback to local data/raw")
hf_dataset = None
ds = None
if args.smoke:
try:
files = collect_txt(Path(data_dir) if Path(data_dir).exists() else base / "../../data/raw")
except FileNotFoundError:
files = []
print("[smoke] data/raw nahi mila (HF clone me excluded hai) — built-in dummy lines use kar raha hu.")
# smoke me dummy tok: quick train 512 vocab on sample lines
from tokenizers.models import BPE
from tokenizers.trainers import BpeTrainer
from tokenizers.pre_tokenizers import ByteLevel
t = Tokenizer(BPE(unk_token="<unk>"))
t.pre_tokenizer = ByteLevel(add_prefix_space=False, use_regex=True)
lines = []
for fp in files:
lines += [l.strip() for l in open(fp, encoding="utf-8") if l.strip()]
if not lines:
# HF Kaggle clone ke liye fallback — data/raw push nahi hota
lines = [
"bhai kal party me kya scene hai",
"mai tumse bahut pyaar karta hu",
"yaar ye phone ka network bahut slow hai",
"tumne khana khaya kya abhi tak",
"नमस्ते आप कैसे हैं",
"मुझे हिंदी में कहानी सुनाओ",
"The quick brown fox jumps over the lazy dog",
"Artificial intelligence is transforming the world",
] * 10
special_toks = ["<pad>", "<bos>", "<eos>", "<unk>", "<|hindi|>", "<|english|>", "<|hinglish|>"]
t.train_from_iterator(lines, trainer=BpeTrainer(vocab_size=512, special_tokens=special_toks))
tok = t
margs.vocab_size = tok.get_vocab_size()
if files:
ds = TxtPackDataset(files, tok, seq_len + 1)
else:
# files ke bina seedha lines se blocks banao
ids = []
eos = tok.token_to_id("<eos>")
if eos is None:
eos = 2
for ln in lines:
tag = get_lang_tag(ln)
tag_id = tok.token_to_id(tag) if tok is not None else None
if tag_id is not None:
ids.append(tag_id)
ids.extend(tok.encode(ln).ids)
ids.append(eos)
import torch as _torch
ids_t = _torch.tensor(ids, dtype=_torch.long)
class _DummyDS(_torch.utils.data.Dataset):
def __init__(self, ids_t, seq_len):
self.ids_t = ids_t
self.seq_len = seq_len
self.rebuild()
def rebuild(self):
n = (len(self.ids_t) // self.seq_len) * self.seq_len
self.blocks = self.ids_t[:n].view(-1, self.seq_len) if n > 0 else _torch.zeros((1, self.seq_len), dtype=_torch.long)
def set_seq_len(self, new_seq_len):
if new_seq_len == self.seq_len:
return
self.seq_len = new_seq_len
self.rebuild()
def __len__(self): return len(self.blocks)
def __getitem__(self, i):
x = self.blocks[i]; return x[:-1], x[1:]
ds = _DummyDS(ids_t, seq_len + 1)
elif ds is not None:
pass # ds already built via HF streaming above
else:
files = collect_txt(Path(data_dir))
print(f"[data] {len(files)} files")
ds = TxtPackDataset(files, tok, seq_len + 1)
# shuffle-split (language bias fix): pehle shuffle, phir train/val baanto.
# DDP me har rank apna stride-shard leta hai (train_ids[rank::world)] — data overlap zero.
seed = C("seed", 42) if not args.smoke else 42
rng = random.Random(seed)
all_idx = list(range(len(ds)))
rng.shuffle(all_idx)
train_ratio = C("train_split", 0.95) if not args.smoke else 0.8
n_train = max(int(len(ds) * train_ratio), 1)
# val ke liye kam se kam 1 block rakho (agar possible)
if len(ds) > n_train:
train_ids = all_idx[:n_train]
val_ids = all_idx[n_train:]
class _FrozenValDS(Dataset):
def __init__(self, blocks): self.blocks = blocks
def __len__(self): return len(self.blocks)
def __getitem__(self, i):
x = self.blocks[i]; return x[:-1], x[1:]
val_ds = _FrozenValDS(ds.blocks[val_ids].clone())
else:
train_ids = all_idx
val_ds = None
if use_dist and not args.smoke:
train_ids = train_ids[rank::world]
if not train_ids:
# blocks < world ya rank ko empty stride: ZeroDivision crash se pehle guard
print(f"[dist] WARN rank {rank}: 0 train blocks after stride — fallback 1 block (data overlap)")
train_ids = [all_idx[rank % len(all_idx)]] if all_idx else [0]
else:
if R0 or use_dist: # DDP me stride-count rank-wise useful hai
print(f"[dist] rank {rank}: {len(train_ids)} train blocks (stride-shard /{world})")
if R0:
print(f"[data] blocks={len(ds)} train={len(train_ids)} val={len(val_ids) if val_ds else 0} seq={seq_len} batch={batch} accum={accum} eff_batch={batch*accum} (x{world} ranks = {batch*accum*world} global)")
model = Viu1MoE(margs).to(device=device, dtype=amp_dtype if (use_amp and device.type == "cuda") else None)
# auto_tune PROBE DDP wrap SE PEHLE (raw model): wrap ke baad probe backward allreduce
# karta hai — dono ranks ka OOM point alag ho to probe-step count mismatch → NCCL hang.
do_auto_tune = bool(C("auto_tune", True)) and not args.no_auto_tune
if do_auto_tune and not args.smoke and device.type == "cuda":
# Target global effective batch = 128 sequences (131,072 tokens/step)
# In DDP multi-GPU, divide target among ranks; in single-GPU, take full global target.
global_target_eff = int(C("target_global_batch", 128))
target_eff = max(1, global_target_eff // world) if use_dist else global_target_eff
probe_dtype = (torch.bfloat16 if use_bf16 else torch.float16) if (use_bf16 or use_fp16) else None
micro = probe_micro_batch(model, seq_len, margs.vocab_size, device, probe_dtype)
micro = max(1, min(micro, target_eff))
accum = max(1, round(target_eff / micro))
batch = micro
if use_dist and not args.smoke:
# ranks ka VRAM alag ho to micro/accum alag → per-step backward count alag → hang.
# rank0 ka auto_tune choice sab pe broadcast karo.
import torch.distributed as dist
_cfg = torch.tensor([batch, accum], device=device, dtype=torch.long)
dist.broadcast(_cfg, src=0)
batch, accum = int(_cfg[0]), int(_cfg[1])
print(f"[auto] {gname} ({ggb:.0f}GB, bf16_ok={bf16_ok}) amp={'bf16' if use_bf16 else 'fp16'} "
f"micro={batch} accum={accum} eff={batch * accum} (target {target_eff}) tok/step={batch * accum * seq_len}"
+ (" [dist broadcast r0->all]" if (use_dist and not args.smoke) else "")) if R0 else None
elif not args.smoke and device.type == "cuda" and R0:
print(f"[batch] MANUAL LOCK: micro={batch} accum={accum} eff={batch * accum} tok/step={batch * accum * seq_len * (world if use_dist else 1)} amp={'bf16' if use_bf16 else 'fp16'}")
# 2026 speed stack (opt-in, DDP-wrap se pehle): FP8 phir compile
if not args.smoke:
model, _fp8info = maybe_apply_fp8(model, bool(args.fp8 or C("fp8", False)), device, R0)
model, _cmpinfo = maybe_apply_compile(model, bool(args.compile or C("compile", False)),
str(C("compile_mode", "reduce-overhead")), args.smoke, R0,
device=device, vocab_size=margs.vocab_size)
if use_dist and not args.smoke:
# MoE top-2 routing me kuch experts kabhi-kabhi empty rehte hain -> find_unused_parameters=True (safe)
model = torch.nn.parallel.DistributedDataParallel(model, device_ids=[device.index], find_unused_parameters=True)
if R0:
print(f"[dist] DDP wrapped (find_unused=True, micro={batch} accum={accum} synced)")
ckpt_model = model.module if (use_dist and not args.smoke) else model # ckpt bina module. prefix ke (resume-safe)
total, active = ckpt_model.param_counts()
if R0:
print(f"[model] total={total/1e6:.1f}M active={active/1e6:.1f}M device={device}" + (f" (DDP x{world})" if use_dist else ""))
tok_per_step = batch * accum * seq_len * (world if use_dist else 1)
if (max_steps is None or max_steps <= 0) and not args.smoke:
total_tokens = int(C("total_tokens", 126_500_000_000))
max_steps = max(1, total_tokens // tok_per_step)
if R0:
print(f"[auto] max_steps dynamically computed from dataset ({total_tokens/1e9:.1f}B tokens / {tok_per_step:,} tok/step): {max_steps:,} steps")
decay_params = [p for p in model.parameters() if p.requires_grad and p.ndim >= 2]
no_decay_params = [p for p in model.parameters() if p.requires_grad and p.ndim < 2]
optim_groups = [
{"params": decay_params, "weight_decay": wd},
{"params": no_decay_params, "weight_decay": 0.0},
]
if bool(C("grad_ckpt", False)):
target = model.module if hasattr(model, "module") else model
if hasattr(target, "gradient_checkpointing_enable"):
target.gradient_checkpointing_enable()
if R0:
print("[model] gradient checkpointing enabled (VRAM activation memory halved for 48L)")
elif R0:
print("[model] WARN: grad_ckpt requested but model has no gradient_checkpointing_enable()")
opt_name = str(args.optimizer or C("optimizer", "adamw")).lower()
if opt_name == "muon":
muon_p, adamw_p = split_muon_params(model)
muon_lr = float(args.muon_lr if getattr(args, "muon_lr", None) is not None else C("muon_lr", 0.02))
muon_wd = float(args.muon_weight_decay if getattr(args, "muon_weight_decay", None) is not None else C("muon_weight_decay", 0.01))
opt_muon = Muon(muon_p, lr=muon_lr, momentum=float(C("muon_momentum", 0.95)), weight_decay=muon_wd)
adamw_groups = [
{"params": [p for p in adamw_p if p.ndim >= 2], "weight_decay": wd},
{"params": [p for p in adamw_p if p.ndim < 2], "weight_decay": 0.0},
]
if device.type == "cuda":
try:
import bitsandbytes as bnb
opt_adamw = bnb.optim.PagedAdamW8bit(adamw_groups, lr=lr, betas=tuple(C("betas", [0.9, 0.95])), eps=1e-5)
if R0:
print(f"[opt] Muon ({len(muon_p)} 2D weights, lr={muon_lr}, wd={muon_wd}) + PagedAdamW8bit ({len(adamw_p)} 1D/emb/router, lr={lr}, wd={wd}) hybrid loaded!")
except ImportError:
opt_adamw = torch.optim.AdamW(adamw_groups, lr=lr, betas=tuple(C("betas", [0.9, 0.95])), eps=1e-5)
if R0:
print(f"[opt] Muon ({len(muon_p)} 2D weights, lr={muon_lr}, wd={muon_wd}) + AdamW ({len(adamw_p)} 1D/emb/router, lr={lr}, wd={wd}) hybrid loaded!")
else:
opt_adamw = torch.optim.AdamW(adamw_groups, lr=lr, betas=tuple(C("betas", [0.9, 0.95])), eps=1e-5)
if R0:
print(f"[opt] Muon ({len(muon_p)} 2D weights, lr={muon_lr}, wd={muon_wd}) + AdamW ({len(adamw_p)} 1D/emb/router, lr={lr}, wd={wd}) hybrid loaded (CPU)!")
opt = CombinedOptimizer([opt_muon, opt_adamw])
elif opt_name in ("8bit", "paged_adamw_8bit", "adamw8bit") and device.type == "cuda":
try:
import bitsandbytes as bnb
opt = bnb.optim.PagedAdamW8bit(
optim_groups, lr=lr, betas=tuple(C("betas", [0.9, 0.95])) if not args.smoke else (0.9, 0.95), eps=1e-5
)
if R0:
print("[opt] loaded bitsandbytes PagedAdamW8bit (VRAM memory optimized for 1.5B)")
except ImportError:
if R0:
print("[opt] bitsandbytes not found, falling back to standard AdamW")
opt = torch.optim.AdamW(optim_groups, lr=lr, betas=tuple(C("betas", [0.9, 0.95])) if not args.smoke else (0.9, 0.95), eps=1e-5)
else:
opt = torch.optim.AdamW(optim_groups, lr=lr, betas=tuple(C("betas", [0.9, 0.95])) if not args.smoke else (0.9, 0.95), eps=1e-5)
use_amp = use_bf16 or use_fp16
amp_dtype = torch.bfloat16 if use_bf16 else torch.float16
# PyTorch GradScaler is strictly designed for FP32 master weights. If model parameters are cast to
# Half/BF16 (to fit within 16GB VRAM on T4), scaler.unscale_ raises ValueError("Attempting to unscale FP16 gradients.").
first_p = next(model.parameters(), None)
is_half_model = first_p is not None and first_p.dtype in (torch.float16, torch.bfloat16)
scaler = torch.amp.GradScaler("cuda") if (use_fp16 and device.type == "cuda" and not is_half_model) else None
# BF16 me scaler nahi chahiye, FP16 me chahiye
# HF auto-push (har training ke baad ckpt backup mandatory) — token env HF_TOKEN se
do_hf_push = bool(C("hf_push", False)) and not args.smoke
hf_repo = C("hf_repo", None)
hf_token = os.environ.get("HF_TOKEN")
hf_every = int(C("hf_push_every", 0) or 0)
# FULL ckpt (optimizer sahit) HF push — har N steps (user design: 10k pe push,
# break ke baad HF se download + resume). Slim se resume nahi hota (opt missing).
hf_full_every = int(C("hf_push_full_every", 0) or 0)
if do_hf_push and not hf_repo:
if R0:
print("[hf-push] WARN: hf_push true par hf_repo missing — push off")
do_hf_push = False
if do_hf_push and not hf_token and R0:
print("[hf-push] HF_TOKEN env me nahi — push skip hoga (token set karo ya manual cell chalao)")
start = 0
resume = None if args.smoke else (args.resume or C("resume", None))
if resume:
resolved_ckpt = None
# 1. Local path check
p_res = Path(resume)
if p_res.exists():
resolved_ckpt = p_res
elif (Path(ckpt_dir) / resume).exists():
resolved_ckpt = Path(ckpt_dir) / resume
elif (Path(ckpt_dir) / f"{resume}.pt").exists():
resolved_ckpt = Path(ckpt_dir) / f"{resume}.pt"
# 2. Remote HF Hub check if not found locally
if resolved_ckpt is None and hf_repo:
try:
from huggingface_hub import HfApi, hf_hub_download
api = HfApi(token=hf_token)
repo_files = api.list_repo_files(repo_id=hf_repo, token=hf_token)
target_file = None
if str(resume).lower() in ("auto", "latest"):
# Find highest step full checkpoint
full_pts = []
for f in repo_files:
m = re.match(r"^model/checkpoints/step_(\d+)\.pt$", f)
if m:
full_pts.append((int(m.group(1)), f))
if full_pts:
full_pts.sort(key=lambda x: x[0], reverse=True)
target_file = full_pts[0][1]
else:
cand = f"model/checkpoints/{Path(resume).name}"
if cand in repo_files:
target_file = cand
elif str(resume) in repo_files:
target_file = str(resume)
if target_file:
if R0:
print(f"[resume] downloading remote checkpoint {target_file} from {hf_repo}...", flush=True)
# Download directly to repo root structure
downloaded_path = hf_hub_download(
repo_id=hf_repo,
filename=target_file,
token=hf_token,
local_dir=str(resolve(Path(ckpt_dir).parent.parent)),
)
resolved_ckpt = Path(downloaded_path)
except Exception as e:
if R0:
print(f"[resume] HF Hub check warning: {e}", flush=True)
if use_dist and not args.smoke:
import torch.distributed as dist
dist.barrier()
if resolved_ckpt and resolved_ckpt.exists():
start = load_ckpt(resolved_ckpt, ckpt_model, opt, scaler=scaler)
if R0:
print(f"[resume] successfully resumed at step {start} from {resolved_ckpt}")
if use_curriculum:
new_stage_idx, new_stage = resolve_curriculum_stage(start, curriculum_stages)
if new_stage_idx != curr_stage_idx:
curr_stage_idx = new_stage_idx
seq_len = int(new_stage["seq_len"])
batch = int(args.batch_size or new_stage["batch_size"])
accum = int(args.grad_accum or new_stage["grad_accum"])
if use_dist and not args.smoke:
accum = max(1, accum // world)
if hasattr(ds, "set_seq_len"):
ds.set_seq_len(seq_len + 1)
train_ids = list(range(len(ds)))
if use_dist and not args.smoke:
train_ids = train_ids[rank::world]
if not train_ids:
train_ids = [0]
rng.shuffle(train_ids)
idx = 0
tok_per_step = batch * accum * seq_len * (world if use_dist else 1)
if R0:
print(f"[curriculum] Resumed into Stage {curr_stage_idx + 1}/{len(curriculum_stages)}: "
f"seq_len={seq_len} batch={batch} accum={accum} (tok/step={tok_per_step:,})", flush=True)
else:
if R0:
print(f"[resume] checkpoint '{resume}' not found locally or on HF Hub — starting fresh from step 0")
log_dir = Path(log_dir)
if R0:
log_dir.mkdir(parents=True, exist_ok=True)
logf = open(log_dir / ("train_smoke.jsonl" if args.smoke else "train.jsonl"), "a", encoding="utf-8") if R0 else None
if R0:
Path(ckpt_dir).mkdir(parents=True, exist_ok=True)
async_uploader = None
if R0 and not args.smoke:
async_uploader = AsyncCheckpointManager(
repo_id=hf_repo,
token=hf_token,
ckpt_dir=ckpt_dir,
log_dir=log_dir,
enabled=do_hf_push,
keep_local_full=int(C("keep_local_full", 2)),
keep_remote_full=int(C("keep_remote_full", 2))
)
model.train()
t0 = time.time()
step = start
idx = 0
nan_streak = 0
max_nan_streak = int(C("max_nan_streak", 25))
usage_log_every = int(C("usage_log_every", 100))
tok_per_step = batch * accum * seq_len * (world if use_dist else 1) # DDP: global tokens/step
pbar = None # Industry standard: clean stdout lines (no carriage-return progress bar in web terminal)
while step < max_steps:
# Check Automatic Curriculum Upgrade
if use_curriculum:
new_stage_idx, new_stage = resolve_curriculum_stage(step, curriculum_stages)
if new_stage_idx != curr_stage_idx:
old_seq = seq_len
curr_stage_idx = new_stage_idx
seq_len = int(new_stage["seq_len"])
batch = int(args.batch_size or new_stage["batch_size"])
accum = int(args.grad_accum or new_stage["grad_accum"])
if use_dist and not args.smoke:
accum = max(1, accum // world)
tok_per_step = batch * accum * seq_len * (world if use_dist else 1)
if hasattr(ds, "set_seq_len"):
ds.set_seq_len(seq_len + 1)
train_ids = list(range(len(ds)))
if use_dist and not args.smoke:
train_ids = train_ids[rank::world]
if not train_ids:
train_ids = [0]
rng.shuffle(train_ids)
idx = 0
if R0:
print("\n" + "=" * 80)
print(f"[CURRICULUM UPGRADE] Step {step:,}: Upgraded to Stage {curr_stage_idx + 1}/{len(curriculum_stages)}!")
print(f" * Context Length: {old_seq} -> {seq_len} tokens")
print(f" * Micro Batch Size: {batch}")
print(f" * Grad Accumulation: {accum}")
print(f" * Global Tok/Step: {tok_per_step:,} (Constant Token Budget Maintained)")
print("=" * 80 + "\n", flush=True)
if idx >= len(train_ids):
if hasattr(ds, "refill") and not args.smoke:
refilled = ds.refill()
if refilled > 0:
train_ids = list(range(len(ds)))
if use_dist and not args.smoke:
train_ids = train_ids[rank::world]
if not train_ids:
train_ids = [0]
rng.shuffle(train_ids)
idx = 0
if R0:
print(f"[stream] Buffer dynamically refilled: loaded {len(ds):,} fresh blocks from 125B+ token stream (step {step:,})", flush=True)
else:
if getattr(ds, "stream_exhausted", False):
if R0:
print(f"[stream] All 2,109 dataset shards processed! Stream complete at step {step:,}. Exiting to save final checkpoint...", flush=True)
break
rng.shuffle(train_ids)
idx = 0
else:
rng.shuffle(train_ids)
idx = 0
# accumulate (batch_size fix: har micro-step me batch samples stack karo)
opt.zero_grad(set_to_none=True)
acc_loss = 0.0
acc_moe_aux = 0.0
acc_mt_loss = 0.0
acc_usages = []
for micro_idx in range(accum):
batch_idx = [train_ids[(idx + j) % len(train_ids)] for j in range(batch)]
idx += batch
xs, ys = [], []
for b in batch_idx:
xb, yb = ds[b]
xs.append(xb)
ys.append(yb)
x = torch.stack(xs).to(device)
y = torch.stack(ys).to(device)
if use_amp and device.type == "cuda":
ctx = torch.amp.autocast("cuda", dtype=amp_dtype)
else:
ctx = torch.amp.autocast("cpu", enabled=False)
is_last_micro = (micro_idx == accum - 1)
sync_ctx = contextlib.nullcontext() if (is_last_micro or not use_dist or not hasattr(model, "no_sync")) else model.no_sync()
with sync_ctx:
with ctx:
_, loss, aux = model(x, y)
loss_scaled = loss / accum
if scaler is not None:
scaler.scale(loss_scaled).backward()
else:
loss_scaled.backward()
acc_loss += loss.item() / accum
if isinstance(aux, dict):
acc_moe_aux += float(aux.get("moe_aux", 0.0)) / accum
acc_mt_loss += float(aux.get("mt_loss", 0.0)) / accum
if "usages" in aux and aux["usages"]:
acc_usages.append(aux["usages"])
avg_usages = []
if acc_usages:
n_layers = len(acc_usages[0])
for l_idx in range(n_layers):
n_exp = len(acc_usages[0][l_idx])
layer_avg = [sum(acc_usages[s][l_idx][e] for s in range(len(acc_usages))) / len(acc_usages) for e in range(n_exp)]
avg_usages.append(layer_avg)
aux = {"moe_aux": round(acc_moe_aux, 4), "mt_loss": round(acc_mt_loss, 4), "usages": avg_usages}
# NaN/Inf guard: poisoned step ko skip karo (optimizer state bachao).
# DDP me decision GLOBAL hoti hai (kisi ek rank pe NaN => sab skip),
# warna ranks ke params desync ho jayenge. Streak badhti jaye to
# divergence maan ke emergency ckpt + sab ranks ek saath ruko.
local_bad = not math.isfinite(acc_loss)
do_skip = local_bad
if use_dist and not args.smoke:
try:
import torch.distributed as _dist
_t = torch.tensor([1 if local_bad else 0], device=device)
_dist.all_reduce(_t, op=_dist.ReduceOp.MAX)
do_skip = bool(_t.item())
except Exception:
pass
if do_skip:
nan_streak += 1
opt.zero_grad(set_to_none=True)
if R0:
print(f"[nan-guard] step {step+1}: non-finite loss ({acc_loss}) — update skipped (streak {nan_streak}/{max_nan_streak})", flush=True)
if logf is not None:
logf.write(json.dumps({"step": step + 1, "skipped": True, "loss": str(acc_loss),
"nan_streak": nan_streak}) + "\n")
logf.flush()
if nan_streak >= max_nan_streak and not args.smoke:
if R0:
print(f"[nan-guard] DIVERGENCE: {nan_streak} consecutive NaN steps — emergency ckpt + abort", flush=True)
save_ckpt(Path(ckpt_dir) / f"emergency_nan_{step+1}.pt", ckpt_model, opt, step, cfg, scaler=scaler)
break
step += 1
continue
nan_streak = 0
if scaler is not None:
scaler.unscale_(opt)
total_norm = torch.nn.utils.clip_grad_norm_(model.parameters(), clip)
try:
total_norm = float(total_norm)
except Exception:
total_norm = -1.0
cur_lr = cosine_lr(step + 1, max_steps, warmup, lr, min_lr)
lr_ratio = cur_lr / max(lr, 1e-9)
for g in opt.param_groups:
if "initial_lr" not in g:
g["initial_lr"] = g["lr"]
g["lr"] = g["initial_lr"] * lr_ratio
if scaler is not None:
scaler.step(opt)
scaler.update()
else:
opt.step()
if device.type == "cuda":
torch.cuda.empty_cache()
step += 1
# Clean industry-standard pretraining logs: only essential metrics, no terminal garbling
if R0 and (step % log_every == 0 or step <= 5 or step == max_steps):
dt = time.time() - t0
done = max(step - start, 1)
rate = done / max(dt, 1e-6)
eta = (max_steps - step) / rate if rate > 0 else 0
tps = tok_per_step * rate
pct = 100.0 * step / max(max_steps, 1)
moe_usages = aux.get("usages") if isinstance(aux, dict) else None
max_expert_share = None
if moe_usages:
avg_u = [sum(layer_u[e] for layer_u in moe_usages) / len(moe_usages) for e in range(len(moe_usages[0]))]
max_expert_share = round(max(avg_u), 3)
if max_expert_share > 0.60:
print(f"[warn] ROUTER IMBALANCE: max expert share {max_expert_share} > 0.60!", flush=True)
# Full 48x28 usage matrix sirf kabhi-kabhi log karo (log-bloat bachao);
# har line me max_expert_share + grad_norm kaafi hai.
aux_log = dict(aux) if isinstance(aux, dict) else aux
if isinstance(aux_log, dict) and step % usage_log_every != 0 and step > 5 and step != max_steps:
aux_log = dict(aux_log)
aux_log["usages"] = []
msg = {"step": step, "of": max_steps, "pct": round(pct, 2), "seq_len": seq_len, "loss": round(acc_loss, 4),
"lr": cur_lr, "grad_norm": round(total_norm, 3), "tok_per_sec": round(tps, 1),
"elapsed_sec": round(dt, 1), "eta_sec": round(eta, 1), "aux": aux_log}
if max_expert_share is not None:
msg["max_expert_share"] = max_expert_share
if use_curriculum:
msg["curriculum_stage"] = curr_stage_idx + 1
# Essential clean log line
cur_tag = f"seq: {seq_len} (S{curr_stage_idx+1})" if use_curriculum else f"seq: {seq_len}"
print(f"[step {step:,}/{max_steps:,}] ({pct:.2f}%) | {cur_tag} | loss: {acc_loss:.4f} | gn: {total_norm:.2f} | lr: {cur_lr:.2e} | {tps:,.0f} tok/s | ETA: {fmt_hms(eta)}", flush=True)
if logf is not None:
logf.write(json.dumps(msg) + "\n")
logf.flush()
if R0 and (not args.smoke) and val_ds is not None and step % eval_every == 0:
vl, ppl = evaluate(ckpt_model, val_ds, device, use_amp=use_amp, amp_dtype=amp_dtype)
print(f"[eval] step {step} val_loss {vl:.4f} ppl {ppl:.1f}")
if logf is not None:
logf.write(json.dumps({"step": step, "val_loss": vl, "ppl": ppl}) + "\n")
logf.flush()
if R0 and (not args.smoke) and step % save_every == 0:
full_pt = Path(ckpt_dir) / f"step_{step}.pt"
slim_pt = Path(ckpt_dir) / f"slim_step_{step}.pt"
save_ckpt(full_pt, ckpt_model, opt, step, cfg, scaler=scaler)
save_slim_ckpt(slim_pt, ckpt_model, cfg)
print(f"[ckpt] saved step_{step}.pt + slim_step_{step}.pt to disk")
# Prune older local checkpoints to ensure Kaggle disk never fills up
if async_uploader is not None:
async_uploader.prune_local_full_ckpts(step)
# Asynchronous background upload to Hugging Face Hub (zero GPU blocking)
push_slim = (do_hf_push and hf_every > 0 and (step // max(save_every, 1)) % hf_every == 0)
push_full = (do_hf_push and hf_full_every > 0 and step % hf_full_every == 0)
if push_slim or push_full:
if logf is not None:
logf.flush()
if async_uploader is not None and async_uploader.enabled:
print(f"[ckpt] dispatching background HF upload for step {step} (GPU training continues immediately)...", flush=True)
async_uploader.push_checkpoint_async(
step=step,
full_file=full_pt if push_full else None,
slim_file=slim_pt if push_slim else None,
log_file=Path(log_dir) / "train.jsonl"
)
if R0 and (not args.smoke) and tok is not None and step % sample_every == 0:
try:
print("[sample]", sample_text(ckpt_model, tok, C("sample_prompt", "bhai kal party me"), device, seq_len=min(seq_len, margs.max_seq_len)))
except Exception as e:
print(f"[sample] skip: {e}")
if pbar is not None:
pbar.close()
total_dt = time.time() - t0
total_done = max(step - start, 1)
if R0:
print(f"[done] steps {total_done} time {fmt_hms(total_dt)} "
f"avg_tok/s {tok_per_step * total_done / max(total_dt, 1e-6):,.0f}", flush=True)
if R0 and async_uploader is not None and async_uploader.enabled:
print("[hf-async] waiting for any in-flight background checkpoint uploads to complete...", flush=True)
async_uploader.wait_pending()
if R0 and do_hf_push:
save_slim_ckpt(Path(ckpt_dir) / f"slim_final_{step}.pt", ckpt_model, cfg)
hf_push_ckpt(ckpt_dir, hf_repo, hf_token, only_file=Path(ckpt_dir) / f"slim_final_{step}.pt",
as_name="final_model.pt", msg=f"final slim ckpt step_{step}")
if R0 and not args.smoke and tok is not None:
try:
print("[sample]", sample_text(ckpt_model, tok, C("sample_prompt", "bhai kal party me"), device))
except Exception as e:
print(f"[sample] skip: {e}")
if args.smoke:
print("[smoke] OK — loop, ckpt-funcs, eval-funcs sab chal rahe hai.")
if logf is not None:
logf.close()
if use_dist:
import torch.distributed as _dist
_dist.barrier()
_dist.destroy_process_group()
if R0:
print("[dist] process group closed")
if __name__ == "__main__":
main()