Download model/scripts/train.py from ViuAI/ViuMini-MoE-242M: direct link, hf CLI and curl.
- Browser
- Download file 98.7 kB
-
https://huggingface.co/ViuAI/ViuMini-MoE-242M/resolve/main/model/scripts/train.py
- Command line
-
hf download hf://ViuAI/ViuMini-MoE-242M/model/scripts/train.py
-
curl -L -o train.py https://huggingface.co/ViuAI/ViuMini-MoE-242M/resolve/main/model/scripts/train.py
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) | |
| 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) | |
| 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) | |
| 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)) | |
| 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() | |