ViuTranslate / scripts /train.py
ViuAI's picture
Fix: disable progress bars and memory copy during base checkpoint download
51eb3b0 verified
Raw
History Blame Contribute Delete
19.6 kB
# ==============================================================================
# 🚀 ViuTranslate — Dedicated Neural Translation Master Training Engine
# ==============================================================================
# Model Architecture: ViuAI Sarus-500M
# Target Repository: ViuAI/ViuTranslate
# Dataset Repository: ViuAI/ViuTranslate-Data
# Hardware: Auto-Tuned (Kaggle T4 x 2, P100, A100, RTX 3090/4090)
# ==============================================================================
import os
import sys
import math
import time
import shutil
import argparse
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
from torch.utils.data import Dataset, DataLoader, Sampler
from huggingface_hub import HfApi, hf_hub_download
# Fix stdout encoding for Windows & Cloud
if hasattr(sys.stdout, "reconfigure"):
sys.stdout.reconfigure(encoding="utf-8", errors="replace")
if hasattr(sys.stderr, "reconfigure"):
sys.stderr.reconfigure(encoding="utf-8", errors="replace")
cur_dir = os.path.dirname(os.path.abspath(__file__)) if "__file__" in locals() else os.getcwd()
parent_dir = os.path.dirname(cur_dir)
for p in [cur_dir, parent_dir, os.getcwd()]:
if p not in sys.path:
sys.path.insert(0, p)
from model import ViuAI
from config import ViuAIConfig
PAD_TOKEN_ID = 64000
EOT_ID = 64002
DOMAIN_NAMES = {
0: "en_to_hi_direct",
1: "en_to_hi_command",
2: "hi_to_en_direct",
3: "hi_to_en_command"
}
# ------------------------------------------------------------------------------
# 1. Hardware Profiler
# ------------------------------------------------------------------------------
def auto_profile_hardware():
if not torch.cuda.is_available():
return {
"tier": "CPU", "device_name": "CPU", "vram_gb": 0.0,
"micro_batch": 1, "grad_accum": 64, "dtype": torch.float32,
"desc": "CPU fallback mode"
}
props = torch.cuda.get_device_properties(0)
device_name = props.name
vram_gb = props.total_memory / (1024 ** 3)
major, minor = props.major, props.minor
bf16_supported = torch.cuda.is_bf16_supported()
dtype = torch.bfloat16 if bf16_supported else torch.float16
if vram_gb >= 30:
micro_batch, grad_accum = 32, 2
desc = "NVIDIA RTX 5090 / High-Tier 32GB Blackwell Beast"
elif vram_gb >= 20:
micro_batch, grad_accum = 12, 6
desc = "Pro-Tier GPU (24GB)"
elif vram_gb >= 12:
micro_batch, grad_accum = 8, 8
desc = "Standard Cloud GPU / 16GB (Kaggle T4 / P100)"
else:
micro_batch, grad_accum = 4, 16
desc = "Budget GPU (< 12GB)"
return {
"tier": "GPU",
"device_name": device_name,
"vram_gb": vram_gb,
"compute_cap": f"{major}.{minor}",
"micro_batch": micro_batch,
"grad_accum": grad_accum,
"effective_batch": micro_batch * grad_accum,
"dtype": dtype,
"desc": desc
}
# ------------------------------------------------------------------------------
# 2. Dataset & Length Grouping
# ------------------------------------------------------------------------------
class TranslationDataset(Dataset):
def __init__(self, ids_path: str, labels_path: str, offsets_path: str, domains_path: str = None):
self.tokens_mmap = np.load(ids_path, mmap_mode="r")
self.labels_mmap = np.load(labels_path, mmap_mode="r")
self.offsets = np.load(offsets_path)
self.domains = np.load(domains_path) if (domains_path and os.path.exists(domains_path)) else None
self.num_samples = len(self.offsets) - 1
def __len__(self):
return self.num_samples
def __getitem__(self, idx):
start_idx = int(self.offsets[idx])
end_idx = int(self.offsets[idx + 1])
tokens = torch.from_numpy(self.tokens_mmap[start_idx:end_idx].astype(np.int64))
labels = torch.from_numpy(self.labels_mmap[start_idx:end_idx].astype(np.int64))
domain_id = int(self.domains[idx]) if self.domains is not None else 0
return tokens, labels, domain_id
class LengthGroupedBatchSampler(Sampler):
def __init__(self, dataset, batch_size: int, mega_batch_mult: int = 40, shuffle: bool = True):
self.dataset = dataset
self.batch_size = batch_size
self.mega_batch_mult = mega_batch_mult
self.shuffle = shuffle
self.lengths = dataset.offsets[1:] - dataset.offsets[:-1]
def __iter__(self):
indices = np.random.permutation(len(self.dataset)) if self.shuffle else np.arange(len(self.dataset))
mega_batch_size = self.batch_size * self.mega_batch_mult
for i in range(0, len(indices), mega_batch_size):
mega_batch = indices[i:i + mega_batch_size]
mega_batch = mega_batch[np.argsort(self.lengths[mega_batch])]
for j in range(0, len(mega_batch), self.batch_size):
yield mega_batch[j:j + self.batch_size].tolist()
def __len__(self):
return math.ceil(len(self.dataset) / self.batch_size)
def collate_fn(batch):
tokens_list, labels_list, domain_list = zip(*batch)
max_len = min(512, max(len(t) for t in tokens_list))
padded_tokens = torch.full((len(batch), max_len), PAD_TOKEN_ID, dtype=torch.long)
padded_labels = torch.full((len(batch), max_len), -100, dtype=torch.long)
for i, (tok, lab) in enumerate(zip(tokens_list, labels_list)):
l = min(len(tok), max_len)
padded_tokens[i, :l] = tok[:l]
padded_labels[i, :l] = lab[:l]
return padded_tokens, padded_labels, torch.tensor(domain_list, dtype=torch.long)
# ------------------------------------------------------------------------------
# 3. Learning Rate Scheduler (Cosine with Warmup)
# ------------------------------------------------------------------------------
def get_lr(step, warmup_steps, total_steps, max_lr, min_lr):
if step < warmup_steps:
return max_lr * (step + 1) / max(1, warmup_steps)
if step > total_steps:
return min_lr
decay_ratio = (step - warmup_steps) / max(1, (total_steps - warmup_steps))
coeff = 0.5 * (1.0 + math.cos(math.pi * decay_ratio))
return min_lr + coeff * (max_lr - min_lr)
# ------------------------------------------------------------------------------
# 4. Live Evaluation Previews
# ------------------------------------------------------------------------------
@torch.no_grad()
def run_live_eval_previews(model, tokenizer, device):
if tokenizer is None:
return
test_cases = [
("Direct EN -> HI", "<|user|>\nThe sun rises in the east and sets in the west.<|endofturn|>\n<|assistant|>\n"),
("Direct HI -> EN", "<|user|>\nसूरज पूर्व में उगता है और पश्चिम में डूबता है।<|endofturn|>\n<|assistant|>\n"),
("Command EN -> HI", "<|user|>\nTranslate to Hindi: 'Consistency and discipline are the keys to long term success.'<|endofturn|>\n<|assistant|>\n"),
("Command HI -> EN", "<|user|>\nTranslate to English: 'सफलता का कोई शॉर्टकट नहीं होता, निरंतर प्रयास ही कुंजी है।'<|endofturn|>\n<|assistant|>\n")
]
print("\n 💬 --- [LIVE TRANSLATION PREVIEWS] ---")
model.eval()
for label, prompt in test_cases:
ids = torch.tensor([tokenizer.encode(prompt).ids], dtype=torch.long, device=device)
out = model.generate(ids, max_new_tokens=45, temperature=0.2, eos_token_id=EOT_ID)
gen = tokenizer.decode(out[0][ids.shape[1]:].tolist()).replace("<|endofturn|>", "").strip()
print(f" • [{label:18s}]: \"{gen}\"")
model.train()
# ------------------------------------------------------------------------------
# 5. Cloud Auto-Download Helper
# ------------------------------------------------------------------------------
def ensure_dataset_and_base_ckpt(data_dir: str, base_ckpt_path: str, token: str = None):
os.makedirs(data_dir, exist_ok=True)
os.environ["HF_HUB_DISABLE_PROGRESS_BARS"] = "1"
try:
from huggingface_hub.utils import disable_progress_bars
disable_progress_bars()
except Exception:
pass
shards = [
"train_tokens.npy", "train_labels.npy", "train_offsets.npy", "train_domains.npy",
"val_tokens.npy", "val_labels.npy", "val_offsets.npy", "val_domains.npy",
"metadata.json"
]
missing = [s for s in shards if not os.path.exists(os.path.join(data_dir, s))]
if missing:
print(f"\n🌐 Downloading ViuTranslate-Data shards from Hugging Face Hub (ViuAI/ViuTranslate-Data)...")
for s in shards:
target = os.path.join(data_dir, s)
if not os.path.exists(target):
print(f" • Fetching {s}...")
dl = hf_hub_download(repo_id="ViuAI/ViuTranslate-Data", filename=s, repo_type="dataset", token=token)
if dl != target and not os.path.exists(target):
shutil.copy(dl, target)
print(" ✅ All dataset shards downloaded.")
if not os.path.exists(base_ckpt_path):
print(f"\n🌐 Base checkpoint not found. Downloading base weights (~5.9GB) from ViuAI/ViuAI-500M...")
os.makedirs(os.path.dirname(base_ckpt_path) if os.path.dirname(base_ckpt_path) else ".", exist_ok=True)
dl_b = hf_hub_download(repo_id="ViuAI/ViuAI-500M", filename="checkpoints/ckpt_latest.pt", token=token)
if dl_b != base_ckpt_path and not os.path.exists(base_ckpt_path):
shutil.copy(dl_b, base_ckpt_path)
print(" ✅ Base checkpoint ready.")
# ------------------------------------------------------------------------------
# 6. Main Training Function
# ------------------------------------------------------------------------------
def main():
hw = auto_profile_hardware()
parser = argparse.ArgumentParser(description="ViuTranslate Dedicated Training Engine")
parser.add_argument("--data_dir", type=str, default="data/tokenized", help="Tokenized dataset directory")
parser.add_argument("--base_ckpt", type=str, default="checkpoints/ckpt_latest.pt", help="Base pretrained weights")
parser.add_argument("--output_dir", type=str, default="checkpoints", help="Output directory for trained model")
parser.add_argument("--batch_size", type=int, default=None, help="Micro batch size")
parser.add_argument("--grad_accum", type=int, default=None, help="Gradient accumulation steps")
parser.add_argument("--epochs", type=int, default=3, help="Training epochs (Default: 3)")
parser.add_argument("--max_lr", type=float, default=3.2e-5, help="Peak learning rate")
parser.add_argument("--min_lr", type=float, default=2.0e-6, help="Min learning rate")
parser.add_argument("--neftune_alpha", type=float, default=5.0, help="NEFTune noise scale")
parser.add_argument("--eval_interval", type=int, default=250, help="Evaluation interval")
parser.add_argument("--push_to_hf", action="store_true", default=False, help="Upload directly to ViuAI/ViuTranslate")
parser.add_argument("--hf_token", type=str, default=None, help="Hugging Face API token")
args = parser.parse_args()
micro_b = args.batch_size or hw["micro_batch"]
grad_acc = args.grad_accum or hw["grad_accum"]
eff_batch = micro_b * grad_acc
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
print("=" * 80)
print("🚀 ViuTranslate-500M — Dedicated Neural Translation Training Engine")
print(f" • Device: {hw['device_name']} ({hw['vram_gb']:.2f} GB VRAM)")
print(f" • Batch Config: Micro-Batch {micro_b} × Accum {grad_acc} = Effective Batch {eff_batch}")
print(f" • Target Epochs: {args.epochs}")
print(f" • Precision: {hw['dtype']}")
print(f" • Target HF Repo: ViuAI/ViuTranslate")
print("=" * 80)
# Cloud Sync
ensure_dataset_and_base_ckpt(args.data_dir, args.base_ckpt, args.hf_token)
# Load Tokenizer
tokenizer = None
tok_candidates = ["tokenizer.json", "tokenizer/tokenizer.json", os.path.join(cur_dir, "tokenizer.json")]
for tc in tok_candidates:
if os.path.exists(tc):
try:
from tokenizers import Tokenizer
tokenizer = Tokenizer.from_file(tc)
print(f"✅ Tokenizer loaded successfully ({tokenizer.get_vocab_size():,} vocab)")
break
except Exception:
pass
# Datasets & Loaders
train_ds = TranslationDataset(
ids_path=os.path.join(args.data_dir, "train_tokens.npy"),
labels_path=os.path.join(args.data_dir, "train_labels.npy"),
offsets_path=os.path.join(args.data_dir, "train_offsets.npy"),
domains_path=os.path.join(args.data_dir, "train_domains.npy")
)
val_ds = TranslationDataset(
ids_path=os.path.join(args.data_dir, "val_tokens.npy"),
labels_path=os.path.join(args.data_dir, "val_labels.npy"),
offsets_path=os.path.join(args.data_dir, "val_offsets.npy"),
domains_path=os.path.join(args.data_dir, "val_domains.npy")
)
train_sampler = LengthGroupedBatchSampler(train_ds, batch_size=micro_b, shuffle=True)
train_loader = DataLoader(train_ds, batch_sampler=train_sampler, collate_fn=collate_fn, num_workers=2, pin_memory=True)
val_loader = DataLoader(val_ds, batch_size=micro_b * 2, shuffle=False, collate_fn=collate_fn, num_workers=2)
# Initialize Model
cfg = ViuAIConfig.sft(vocab_size=64003, context_length=2048, neftune_alpha=args.neftune_alpha)
model = ViuAI(cfg).to(device)
print(f"\n📦 Loading base pretrained weights from {args.base_ckpt}...")
base_state = torch.load(args.base_ckpt, map_location=device, weights_only=False)
weights = base_state.get("model_state_dict", base_state)
model.load_state_dict(weights, strict=False)
print("✅ Pretrained weights loaded.")
# Optimizer
fused_available = 'fused' in torch.optim.AdamW.__init__.__code__.co_varnames and torch.cuda.is_available()
optimizer = torch.optim.AdamW(model.parameters(), lr=args.max_lr, weight_decay=0.01, betas=(0.9, 0.95), fused=fused_available)
total_steps = (len(train_loader) // grad_acc) * args.epochs
warmup_steps = int(total_steps * 0.04)
print(f"📊 Total Optimization Steps: {total_steps:,} | Warmup Steps: {warmup_steps:,}")
autocast_ctx = torch.amp.autocast(device_type="cuda", dtype=hw["dtype"]) if torch.cuda.is_available() else contextlib.nullcontext()
# Validation Function
@torch.no_grad()
def evaluate():
model.eval()
total_loss, total_tokens = 0.0, 0
domain_losses = {k: 0.0 for k in DOMAIN_NAMES.keys()}
domain_counts = {k: 0 for k in DOMAIN_NAMES.keys()}
for inputs, labels, doms in val_loader:
inputs, labels = inputs.to(device), labels.to(device)
with autocast_ctx:
logits, loss = model(inputs, targets=labels, pad_id=PAD_TOKEN_ID, shift_labels=True)
tok_count = (labels != -100).sum().item()
total_loss += loss.item() * tok_count
total_tokens += tok_count
# Domain-level tracking
for d in doms.unique():
d_val = d.item()
mask = (doms == d)
if mask.sum() > 0:
with autocast_ctx:
_, d_l = model(inputs[mask], targets=labels[mask], pad_id=PAD_TOKEN_ID, shift_labels=True)
domain_losses[d_val] += d_l.item()
domain_counts[d_val] += 1
avg_loss = total_loss / max(1, total_tokens)
ppl = math.exp(min(20.0, avg_loss))
d_summary = {DOMAIN_NAMES[k]: (domain_losses[k] / max(1, domain_counts[k])) for k in DOMAIN_NAMES.keys()}
model.train()
return avg_loss, ppl, d_summary
# Training Loop
os.makedirs(args.output_dir, exist_ok=True)
save_path = os.path.join(args.output_dir, "viutranslate_final.pt")
best_val_loss = float("inf")
global_step = 0
start_time = time.time()
total_tokens_trained = 0
print("\n⚡ Starting Training...")
for epoch in range(1, args.epochs + 1):
accum_loss = 0.0
model.train()
for micro_idx, (inputs, labels, _) in enumerate(train_loader):
inputs, labels = inputs.to(device, non_blocking=True), labels.to(device, non_blocking=True)
active_tokens = (labels != -100).sum().item()
total_tokens_trained += active_tokens
with autocast_ctx:
logits, loss = model(inputs, targets=labels, pad_id=PAD_TOKEN_ID, shift_labels=True)
loss_scaled = loss / grad_acc
loss_scaled.backward()
accum_loss += loss.item()
if (micro_idx + 1) % grad_acc == 0:
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
lr = get_lr(global_step, warmup_steps, total_steps, args.max_lr, args.min_lr)
for param_group in optimizer.param_groups:
param_group["lr"] = lr
optimizer.step()
optimizer.zero_grad(set_to_none=True)
global_step += 1
step_loss = accum_loss / grad_acc
accum_loss = 0.0
if global_step % 10 == 0 or global_step == 1:
elapsed = time.time() - start_time
tok_s = total_tokens_trained / max(1.0, elapsed)
vram = torch.cuda.memory_allocated() / (1024**3) if torch.cuda.is_available() else 0.0
print(f"Step {global_step:4d}/{total_steps} | Epoch {epoch} | Loss: {step_loss:.4f} | LR: {lr:.2e} | Speed: {tok_s:,.0f} tok/s | VRAM: {vram:.1f}GB")
if global_step % args.eval_interval == 0:
v_loss, v_ppl, d_losses = evaluate()
print(f"\n🌟 [Eval @ Step {global_step}] Val Loss: {v_loss:.4f} | Perplexity: {v_ppl:.2f}")
print(" 📊 Direction Losses: " + " | ".join([f"{k}: {v:.3f}" for k, v in d_losses.items()]))
run_live_eval_previews(model, tokenizer, device)
if v_loss < best_val_loss:
best_val_loss = v_loss
torch.save({"model_state_dict": model.state_dict(), "global_step": global_step, "val_loss": best_val_loss}, save_path)
print(f" 🏆 Saved New Best Checkpoint -> {save_path}\n")
# Final Save
torch.save({"model_state_dict": model.state_dict(), "global_step": global_step, "best_val_loss": best_val_loss}, save_path)
print(f"\n🎉 ViuTranslate Training Complete! Final checkpoint: {save_path}")
# Direct Push to ViuAI/ViuTranslate
if args.push_to_hf:
print("\n🚀 Pushing model weights to Hugging Face Model Repository (ViuAI/ViuTranslate)...")
token = args.hf_token or os.environ.get("HF_TOKEN")
if token:
api = HfApi(token=token)
api.upload_file(
path_or_fileobj=save_path,
path_in_repo="viutranslate_final.pt",
repo_id="ViuAI/ViuTranslate",
repo_type="model"
)
print("✅ Successfully uploaded viutranslate_final.pt to ViuAI/ViuTranslate!")
else:
print("⚠️ Skipping upload: No HF_TOKEN provided.")
if __name__ == "__main__":
main()