Upload code/train_sft.py with huggingface_hub
Browse files- code/train_sft.py +124 -72
code/train_sft.py
CHANGED
|
@@ -1,12 +1,18 @@
|
|
| 1 |
# ==============================================================================
|
| 2 |
-
# π ViuAI Sarus-500M β SFT v16
|
| 3 |
# ==============================================================================
|
| 4 |
-
#
|
| 5 |
-
# β’
|
| 6 |
-
# β’
|
| 7 |
-
# β’
|
| 8 |
-
# β’
|
| 9 |
-
# β’
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 10 |
# ==============================================================================
|
| 11 |
|
| 12 |
import os
|
|
@@ -23,6 +29,9 @@ import torch.nn.functional as F
|
|
| 23 |
from torch.utils.data import Dataset, DataLoader
|
| 24 |
from huggingface_hub import HfApi, hf_hub_download
|
| 25 |
|
|
|
|
|
|
|
|
|
|
| 26 |
# Fix output encoding for Windows & Cloud terminals
|
| 27 |
if hasattr(sys.stdout, "reconfigure"):
|
| 28 |
sys.stdout.reconfigure(encoding="utf-8", errors="replace")
|
|
@@ -36,7 +45,68 @@ if cur_dir not in sys.path:
|
|
| 36 |
from model import Transformer, ModelArgs
|
| 37 |
|
| 38 |
|
| 39 |
-
# ---
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 40 |
class SFTDataset(Dataset):
|
| 41 |
def __init__(self, ids_path: str, labels_path: str, offsets_path: str, domains_path: str = None):
|
| 42 |
self.tokens_mmap = np.load(ids_path, mmap_mode="r")
|
|
@@ -52,7 +122,6 @@ class SFTDataset(Dataset):
|
|
| 52 |
start_idx = int(self.offsets[idx])
|
| 53 |
end_idx = int(self.offsets[idx + 1])
|
| 54 |
|
| 55 |
-
# Load slices from memory-mapped files
|
| 56 |
tokens = torch.from_numpy(self.tokens_mmap[start_idx:end_idx].astype(np.int64))
|
| 57 |
labels = torch.from_numpy(self.labels_mmap[start_idx:end_idx].astype(np.int64))
|
| 58 |
domain_id = int(self.domains[idx]) if self.domains is not None else 0
|
|
@@ -60,7 +129,9 @@ class SFTDataset(Dataset):
|
|
| 60 |
return tokens, labels, domain_id
|
| 61 |
|
| 62 |
|
| 63 |
-
# ---
|
|
|
|
|
|
|
| 64 |
def sft_collate_fn(batch, pad_token_id=64000, ignore_index=-100):
|
| 65 |
inputs, labels, domain_ids = zip(*batch)
|
| 66 |
|
|
@@ -79,7 +150,9 @@ def sft_collate_fn(batch, pad_token_id=64000, ignore_index=-100):
|
|
| 79 |
return batch_inputs, batch_labels, torch.tensor(domain_ids, dtype=torch.long)
|
| 80 |
|
| 81 |
|
| 82 |
-
# ---
|
|
|
|
|
|
|
| 83 |
def get_lr(it, warmup_steps, total_steps, max_lr, min_lr):
|
| 84 |
if it < warmup_steps:
|
| 85 |
return max_lr * (it + 1) / max(1, warmup_steps)
|
|
@@ -90,35 +163,14 @@ def get_lr(it, warmup_steps, total_steps, max_lr, min_lr):
|
|
| 90 |
return min_lr + coeff * (max_lr - min_lr)
|
| 91 |
|
| 92 |
|
| 93 |
-
# ---
|
| 94 |
-
|
| 95 |
-
|
| 96 |
-
return 2, 64
|
| 97 |
-
vram_gb = torch.cuda.get_device_properties(0).total_memory / (1024 ** 3)
|
| 98 |
-
if vram_gb >= 100: # H100, H200 (80GB-140GB)
|
| 99 |
-
return 16, 8 # 16 micro batch * 8 accum = 128 effective
|
| 100 |
-
elif vram_gb >= 40: # A100-40GB/80GB, A6000
|
| 101 |
-
return 8, 16 # 8 micro batch * 16 accum = 128 effective
|
| 102 |
-
elif vram_gb >= 20: # RTX 3090, RTX 4090, L4, A10G (24GB)
|
| 103 |
-
return 4, 32 # 4 micro batch * 32 accum = 128 effective (OOM-Proof!)
|
| 104 |
-
else: # T4, V100, RTX 4080 (16GB)
|
| 105 |
-
return 2, 64 # 2 micro batch * 64 accum = 128 effective (OOM-Proof!)
|
| 106 |
-
|
| 107 |
-
|
| 108 |
-
# --- 5. Self-Contained Cloud Download Helper ---
|
| 109 |
def is_valid_checkpoint(path: str) -> bool:
|
| 110 |
-
|
| 111 |
-
return False
|
| 112 |
-
if os.path.getsize(path) < 10 * 1024 * 1024:
|
| 113 |
-
return False
|
| 114 |
-
return True
|
| 115 |
|
| 116 |
def is_valid_data_file(path: str) -> bool:
|
| 117 |
-
|
| 118 |
-
return False
|
| 119 |
-
if os.path.getsize(path) < 1000:
|
| 120 |
-
return False
|
| 121 |
-
return True
|
| 122 |
|
| 123 |
def ensure_cloud_data_and_checkpoint(version: str, data_dir: str, ckpt_path: str, token: str = None):
|
| 124 |
stage_subfolder = f"sft_{version}"
|
|
@@ -167,17 +219,19 @@ def ensure_cloud_data_and_checkpoint(version: str, data_dir: str, ckpt_path: str
|
|
| 167 |
print(f"β οΈ Error downloading base checkpoint: {e}")
|
| 168 |
|
| 169 |
|
| 170 |
-
# ---
|
|
|
|
|
|
|
| 171 |
def main():
|
| 172 |
-
|
| 173 |
|
| 174 |
-
parser = argparse.ArgumentParser(description="ViuAI Sarus-500M β SFT v16
|
| 175 |
parser.add_argument("--version", type=str, default="v16", help="Dataset/Checkpoint version: v16 (default)")
|
| 176 |
parser.add_argument("--data_dir", type=str, default=None, help="Root directory containing tokenized_data")
|
| 177 |
parser.add_argument("--init_ckpt", type=str, default=None, help="Initial checkpoint path")
|
| 178 |
parser.add_argument("--output_dir", type=str, default=None, help="Directory to save checkpoints")
|
| 179 |
-
parser.add_argument("--batch_size", type=int, default=
|
| 180 |
-
parser.add_argument("--grad_accum", type=int, default=
|
| 181 |
parser.add_argument("--epochs", type=int, default=3, help="Number of training epochs (Default: 3 for SFT v16)")
|
| 182 |
parser.add_argument("--max_lr", type=float, default=3.2e-5, help="Peak learning rate for Cosine Schedule")
|
| 183 |
parser.add_argument("--min_lr", type=float, default=2.0e-6, help="Minimum learning rate")
|
|
@@ -189,6 +243,11 @@ def main():
|
|
| 189 |
parser.add_argument("--hf_token", type=str, default=None, help="Hugging Face API token")
|
| 190 |
args = parser.parse_args()
|
| 191 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 192 |
root_dir = os.path.abspath(os.path.join(cur_dir, ".."))
|
| 193 |
data_dir = args.data_dir or os.path.join(root_dir, "tokenized_data")
|
| 194 |
output_dir = args.output_dir or os.path.join(root_dir, "sft_checkpoints", f"sft_{args.version}")
|
|
@@ -214,18 +273,15 @@ def main():
|
|
| 214 |
if torch.cuda.is_available():
|
| 215 |
torch.cuda.empty_cache()
|
| 216 |
|
| 217 |
-
print("=" *
|
| 218 |
-
print(f"π ViuAI Sarus-500M β SFT {args.version.upper()}
|
| 219 |
-
print(f" β’
|
| 220 |
-
print(f" β’
|
| 221 |
-
print(f" β’
|
| 222 |
-
print(f" β’
|
| 223 |
-
print(f" β’
|
| 224 |
-
print(f" β’ Epochs:
|
| 225 |
-
print(
|
| 226 |
-
print(f" β’ Push to HF: {args.push_to_hf}")
|
| 227 |
-
print(f" β’ Resume Mode: {args.resume}")
|
| 228 |
-
print("=" * 80)
|
| 229 |
|
| 230 |
# 1. Load Data
|
| 231 |
train_dataset = SFTDataset(
|
|
@@ -243,7 +299,7 @@ def main():
|
|
| 243 |
|
| 244 |
train_loader = DataLoader(
|
| 245 |
train_dataset,
|
| 246 |
-
batch_size=
|
| 247 |
shuffle=True,
|
| 248 |
collate_fn=sft_collate_fn,
|
| 249 |
num_workers=2,
|
|
@@ -251,7 +307,7 @@ def main():
|
|
| 251 |
)
|
| 252 |
val_loader = DataLoader(
|
| 253 |
val_dataset,
|
| 254 |
-
batch_size=
|
| 255 |
shuffle=False,
|
| 256 |
collate_fn=sft_collate_fn,
|
| 257 |
num_workers=2,
|
|
@@ -318,10 +374,10 @@ def main():
|
|
| 318 |
]
|
| 319 |
optimizer = torch.optim.AdamW(optimizer_grouped_parameters, lr=args.max_lr, betas=(0.9, 0.95), eps=1e-8)
|
| 320 |
|
| 321 |
-
steps_per_epoch = math.ceil(len(train_loader) /
|
| 322 |
total_steps = steps_per_epoch * args.epochs
|
| 323 |
warmup_steps = max(10, int(total_steps * args.warmup_ratio))
|
| 324 |
-
print(f"π Training
|
| 325 |
|
| 326 |
start_epoch = 1
|
| 327 |
global_step = 0
|
|
@@ -339,12 +395,9 @@ def main():
|
|
| 339 |
|
| 340 |
# Mixed Precision Setup
|
| 341 |
if torch.cuda.is_available():
|
| 342 |
-
|
| 343 |
-
autocast_ctx = torch.amp.autocast(device_type="cuda", dtype=dtype)
|
| 344 |
else:
|
| 345 |
-
dtype = torch.float32
|
| 346 |
autocast_ctx = contextlib.nullcontext()
|
| 347 |
-
print(f"β‘ Mixed Precision: {dtype}")
|
| 348 |
|
| 349 |
# Evaluation Helper
|
| 350 |
@torch.no_grad()
|
|
@@ -379,9 +432,9 @@ def main():
|
|
| 379 |
total_tokens_trained = 0
|
| 380 |
model.train()
|
| 381 |
|
| 382 |
-
print("\n" + "=" *
|
| 383 |
-
print(f"π STARTING SFT {args.version.upper()} MASTER TRAINING")
|
| 384 |
-
print("=" *
|
| 385 |
|
| 386 |
for epoch in range(start_epoch, args.epochs + 1):
|
| 387 |
print(f"\n--- Epoch {epoch}/{args.epochs} ---")
|
|
@@ -405,14 +458,14 @@ def main():
|
|
| 405 |
shift_logits = logits[..., :-1, :].contiguous()
|
| 406 |
shift_labels = labels[..., 1:].contiguous()
|
| 407 |
loss = F.cross_entropy(shift_logits.view(-1, shift_logits.size(-1)), shift_labels.view(-1), ignore_index=-100)
|
| 408 |
-
loss_scaled = loss /
|
| 409 |
|
| 410 |
loss_scaled.backward()
|
| 411 |
accum_loss += loss.item()
|
| 412 |
epoch_loss += loss.item()
|
| 413 |
epoch_batches += 1
|
| 414 |
|
| 415 |
-
if (micro_idx + 1) %
|
| 416 |
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
|
| 417 |
|
| 418 |
lr = get_lr(global_step, warmup_steps, total_steps, args.max_lr, args.min_lr)
|
|
@@ -423,7 +476,7 @@ def main():
|
|
| 423 |
optimizer.zero_grad(set_to_none=True)
|
| 424 |
global_step += 1
|
| 425 |
|
| 426 |
-
step_avg_loss = accum_loss /
|
| 427 |
accum_loss = 0.0
|
| 428 |
|
| 429 |
if global_step % 10 == 0 or global_step == 1:
|
|
@@ -448,7 +501,6 @@ def main():
|
|
| 448 |
best_val_loss = v_loss
|
| 449 |
print(f" π New Best Validation Loss: {best_val_loss:.4f}! Saving checkpoint...")
|
| 450 |
|
| 451 |
-
# Save checkpoint
|
| 452 |
save_payload = {
|
| 453 |
"model_state_dict": model.state_dict(),
|
| 454 |
"optimizer_state_dict": optimizer.state_dict(),
|
|
@@ -463,7 +515,7 @@ def main():
|
|
| 463 |
print(f" πΎ Saved checkpoint -> {save_path}\n")
|
| 464 |
|
| 465 |
# End of Epoch Handling
|
| 466 |
-
if (micro_idx + 1) %
|
| 467 |
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
|
| 468 |
lr = get_lr(global_step, warmup_steps, total_steps, args.max_lr, args.min_lr)
|
| 469 |
for param_group in optimizer.param_groups:
|
|
@@ -477,13 +529,13 @@ def main():
|
|
| 477 |
|
| 478 |
# Final Evaluation & Save
|
| 479 |
final_val_loss, final_val_ppl = evaluate()
|
| 480 |
-
print("\n" + "=" *
|
| 481 |
print(f"π SFT {args.version.upper()} MASTER TRAINING COMPLETED!")
|
| 482 |
print(f" β’ Final Validation Loss: {final_val_loss:.4f}")
|
| 483 |
print(f" β’ Final Perplexity: {final_val_ppl:.2f}")
|
| 484 |
print(f" β’ Total Active Tokens: {total_tokens_trained:,}")
|
| 485 |
print(f" β’ Total Time Taken: {(time.time() - start_time)/60:.2f} minutes")
|
| 486 |
-
print("=" *
|
| 487 |
|
| 488 |
# Final Save
|
| 489 |
final_payload = {
|
|
|
|
| 1 |
# ==============================================================================
|
| 2 |
+
# π ViuAI Sarus-500M β SFT v16 Universal Auto-Adaptive Training Engine
|
| 3 |
# ==============================================================================
|
| 4 |
+
# UNIVERSAL HARDWARE SUPPORT (Zero Manual Tuning Needed!):
|
| 5 |
+
# β’ π Ultra-Tier (H200 141GB, H100 80GB, GH200) -> Batch 16 x Accum 8 = 128
|
| 6 |
+
# β’ π High-Tier (RTX 5090 32GB, A100 40GB, A6000) -> Batch 8 x Accum 16 = 128
|
| 7 |
+
# β’ π Pro-Tier (RTX 4090 24GB, RTX 3090 24GB, L4) -> Batch 4 x Accum 32 = 128
|
| 8 |
+
# β’ β‘ Entry-Tier (T4 16GB, V100 16GB, RTX 4080 16GB) -> Batch 2 x Accum 64 = 128
|
| 9 |
+
# β’ π» Budget (<12GB VRAM / CPU) -> Batch 1 x Accum 128 = 128
|
| 10 |
+
#
|
| 11 |
+
# KEY FEATURES:
|
| 12 |
+
# β’ 100% OOM-PROOF: Auto-probes VRAM, sets expandable segments & clears cache.
|
| 13 |
+
# β’ TARGET EFFECTIVE BATCH = 128 (Mathematical invariance across all hardware).
|
| 14 |
+
# β’ NATIVE BF16 / FP16 AUTO-SWITCH: Uses bfloat16 on Ampere/Ada/Hopper/Blackwell.
|
| 15 |
+
# β’ RESUME & HF AUTO-PUSH: Seamless checkpointing and Hub synchronization.
|
| 16 |
# ==============================================================================
|
| 17 |
|
| 18 |
import os
|
|
|
|
| 29 |
from torch.utils.data import Dataset, DataLoader
|
| 30 |
from huggingface_hub import HfApi, hf_hub_download
|
| 31 |
|
| 32 |
+
# Optimize CUDA allocator to prevent fragmentation
|
| 33 |
+
os.environ["PYTORCH_CUDA_ALLOC_CONF"] = "expandable_segments:True"
|
| 34 |
+
|
| 35 |
# Fix output encoding for Windows & Cloud terminals
|
| 36 |
if hasattr(sys.stdout, "reconfigure"):
|
| 37 |
sys.stdout.reconfigure(encoding="utf-8", errors="replace")
|
|
|
|
| 45 |
from model import Transformer, ModelArgs
|
| 46 |
|
| 47 |
|
| 48 |
+
# ------------------------------------------------------------------------------
|
| 49 |
+
# 1. Universal Hardware Prober & Auto-Tuner
|
| 50 |
+
# ------------------------------------------------------------------------------
|
| 51 |
+
def auto_profile_hardware():
|
| 52 |
+
"""Auto-detects GPU model, VRAM capacity, compute capability and selects golden parameters."""
|
| 53 |
+
if not torch.cuda.is_available():
|
| 54 |
+
return {
|
| 55 |
+
"tier": "CPU", "device_name": "CPU", "vram_gb": 0.0,
|
| 56 |
+
"micro_batch": 1, "grad_accum": 128, "dtype": torch.float32,
|
| 57 |
+
"desc": "CPU fallback mode"
|
| 58 |
+
}
|
| 59 |
+
|
| 60 |
+
props = torch.cuda.get_device_properties(0)
|
| 61 |
+
device_name = props.name
|
| 62 |
+
vram_gb = props.total_memory / (1024 ** 3)
|
| 63 |
+
major, minor = props.major, props.minor
|
| 64 |
+
bf16_supported = torch.cuda.is_bf16_supported()
|
| 65 |
+
dtype = torch.bfloat16 if bf16_supported else torch.float16
|
| 66 |
+
|
| 67 |
+
# Tier Classification for 2048 Context Length
|
| 68 |
+
if vram_gb >= 75: # H200 (141GB), H100 (80GB), GH200, A100-80GB
|
| 69 |
+
tier = "Ultra-Tier"
|
| 70 |
+
micro_batch = 16
|
| 71 |
+
grad_accum = 8
|
| 72 |
+
desc = "NVIDIA Hopper / Datacenter Beast"
|
| 73 |
+
elif vram_gb >= 30: # RTX 5090 (32GB), A100 (40GB), A6000 (48GB), RTX 6000 Ada
|
| 74 |
+
tier = "High-Tier"
|
| 75 |
+
micro_batch = 8
|
| 76 |
+
grad_accum = 16
|
| 77 |
+
desc = "NVIDIA Blackwell / High-End Workstation"
|
| 78 |
+
elif vram_gb >= 20: # RTX 4090 (24GB), RTX 3090 (24GB), L4 (24GB), A10G (24GB)
|
| 79 |
+
tier = "Pro-Tier"
|
| 80 |
+
micro_batch = 4
|
| 81 |
+
grad_accum = 32
|
| 82 |
+
desc = "NVIDIA Ada Lovelace / Ampere Pro (24GB)"
|
| 83 |
+
elif vram_gb >= 12: # T4 (16GB), V100 (16GB), RTX 4080 (16GB), RTX 4070Ti (12GB)
|
| 84 |
+
tier = "Entry-Tier"
|
| 85 |
+
micro_batch = 2
|
| 86 |
+
grad_accum = 64
|
| 87 |
+
desc = "Standard Cloud GPU / 16GB"
|
| 88 |
+
else: # < 12GB VRAM
|
| 89 |
+
tier = "Budget-Tier"
|
| 90 |
+
micro_batch = 1
|
| 91 |
+
grad_accum = 128
|
| 92 |
+
desc = "Budget GPU (< 12GB)"
|
| 93 |
+
|
| 94 |
+
return {
|
| 95 |
+
"tier": tier,
|
| 96 |
+
"device_name": device_name,
|
| 97 |
+
"vram_gb": vram_gb,
|
| 98 |
+
"compute_cap": f"{major}.{minor}",
|
| 99 |
+
"micro_batch": micro_batch,
|
| 100 |
+
"grad_accum": grad_accum,
|
| 101 |
+
"effective_batch": micro_batch * grad_accum,
|
| 102 |
+
"dtype": dtype,
|
| 103 |
+
"desc": desc
|
| 104 |
+
}
|
| 105 |
+
|
| 106 |
+
|
| 107 |
+
# ------------------------------------------------------------------------------
|
| 108 |
+
# 2. Memory-Mapped High Performance Dataset
|
| 109 |
+
# ------------------------------------------------------------------------------
|
| 110 |
class SFTDataset(Dataset):
|
| 111 |
def __init__(self, ids_path: str, labels_path: str, offsets_path: str, domains_path: str = None):
|
| 112 |
self.tokens_mmap = np.load(ids_path, mmap_mode="r")
|
|
|
|
| 122 |
start_idx = int(self.offsets[idx])
|
| 123 |
end_idx = int(self.offsets[idx + 1])
|
| 124 |
|
|
|
|
| 125 |
tokens = torch.from_numpy(self.tokens_mmap[start_idx:end_idx].astype(np.int64))
|
| 126 |
labels = torch.from_numpy(self.labels_mmap[start_idx:end_idx].astype(np.int64))
|
| 127 |
domain_id = int(self.domains[idx]) if self.domains is not None else 0
|
|
|
|
| 129 |
return tokens, labels, domain_id
|
| 130 |
|
| 131 |
|
| 132 |
+
# ------------------------------------------------------------------------------
|
| 133 |
+
# 3. Dynamic Padding & Strict 2048 Bound Collate Function
|
| 134 |
+
# ------------------------------------------------------------------------------
|
| 135 |
def sft_collate_fn(batch, pad_token_id=64000, ignore_index=-100):
|
| 136 |
inputs, labels, domain_ids = zip(*batch)
|
| 137 |
|
|
|
|
| 150 |
return batch_inputs, batch_labels, torch.tensor(domain_ids, dtype=torch.long)
|
| 151 |
|
| 152 |
|
| 153 |
+
# ------------------------------------------------------------------------------
|
| 154 |
+
# 4. Cosine Learning Rate Schedule with Warmup
|
| 155 |
+
# ------------------------------------------------------------------------------
|
| 156 |
def get_lr(it, warmup_steps, total_steps, max_lr, min_lr):
|
| 157 |
if it < warmup_steps:
|
| 158 |
return max_lr * (it + 1) / max(1, warmup_steps)
|
|
|
|
| 163 |
return min_lr + coeff * (max_lr - min_lr)
|
| 164 |
|
| 165 |
|
| 166 |
+
# ------------------------------------------------------------------------------
|
| 167 |
+
# 5. Cloud Auto-Sync (Dataset & Checkpoint Auto-Download)
|
| 168 |
+
# ------------------------------------------------------------------------------
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 169 |
def is_valid_checkpoint(path: str) -> bool:
|
| 170 |
+
return os.path.exists(path) and (os.path.getsize(path) >= 10 * 1024 * 1024)
|
|
|
|
|
|
|
|
|
|
|
|
|
| 171 |
|
| 172 |
def is_valid_data_file(path: str) -> bool:
|
| 173 |
+
return os.path.exists(path) and (os.path.getsize(path) >= 1000)
|
|
|
|
|
|
|
|
|
|
|
|
|
| 174 |
|
| 175 |
def ensure_cloud_data_and_checkpoint(version: str, data_dir: str, ckpt_path: str, token: str = None):
|
| 176 |
stage_subfolder = f"sft_{version}"
|
|
|
|
| 219 |
print(f"β οΈ Error downloading base checkpoint: {e}")
|
| 220 |
|
| 221 |
|
| 222 |
+
# ------------------------------------------------------------------------------
|
| 223 |
+
# 6. Main Universal Training Runner
|
| 224 |
+
# ------------------------------------------------------------------------------
|
| 225 |
def main():
|
| 226 |
+
hw = auto_profile_hardware()
|
| 227 |
|
| 228 |
+
parser = argparse.ArgumentParser(description="ViuAI Sarus-500M β SFT v16 Universal Auto-Adaptive Training Engine")
|
| 229 |
parser.add_argument("--version", type=str, default="v16", help="Dataset/Checkpoint version: v16 (default)")
|
| 230 |
parser.add_argument("--data_dir", type=str, default=None, help="Root directory containing tokenized_data")
|
| 231 |
parser.add_argument("--init_ckpt", type=str, default=None, help="Initial checkpoint path")
|
| 232 |
parser.add_argument("--output_dir", type=str, default=None, help="Directory to save checkpoints")
|
| 233 |
+
parser.add_argument("--batch_size", type=int, default=None, help="Micro batch size (Auto-configured if omitted)")
|
| 234 |
+
parser.add_argument("--grad_accum", type=int, default=None, help="Gradient accumulation steps (Auto-configured if omitted)")
|
| 235 |
parser.add_argument("--epochs", type=int, default=3, help="Number of training epochs (Default: 3 for SFT v16)")
|
| 236 |
parser.add_argument("--max_lr", type=float, default=3.2e-5, help="Peak learning rate for Cosine Schedule")
|
| 237 |
parser.add_argument("--min_lr", type=float, default=2.0e-6, help="Minimum learning rate")
|
|
|
|
| 243 |
parser.add_argument("--hf_token", type=str, default=None, help="Hugging Face API token")
|
| 244 |
args = parser.parse_args()
|
| 245 |
|
| 246 |
+
# Determine batch parameters: explicit user override OR auto-detected golden values
|
| 247 |
+
micro_b = args.batch_size if args.batch_size is not None else hw["micro_batch"]
|
| 248 |
+
grad_acc = args.grad_accum if args.grad_accum is not None else hw["grad_accum"]
|
| 249 |
+
eff_batch = micro_b * grad_acc
|
| 250 |
+
|
| 251 |
root_dir = os.path.abspath(os.path.join(cur_dir, ".."))
|
| 252 |
data_dir = args.data_dir or os.path.join(root_dir, "tokenized_data")
|
| 253 |
output_dir = args.output_dir or os.path.join(root_dir, "sft_checkpoints", f"sft_{args.version}")
|
|
|
|
| 273 |
if torch.cuda.is_available():
|
| 274 |
torch.cuda.empty_cache()
|
| 275 |
|
| 276 |
+
print("=" * 85)
|
| 277 |
+
print(f"π ViuAI Sarus-500M β SFT {args.version.upper()} Universal Auto-Adaptive Training Runner")
|
| 278 |
+
print(f" β’ Hardware Tier: {hw['tier']} ({hw['desc']})")
|
| 279 |
+
print(f" β’ Device Name: {hw['device_name']} | VRAM: {hw['vram_gb']:.2f} GB")
|
| 280 |
+
print(f" β’ Auto-Tuned Batch: Micro-Batch {micro_b} Γ Grad-Accum {grad_acc} = Effective Batch {eff_batch}")
|
| 281 |
+
print(f" β’ Precision Mode: {hw['dtype']}")
|
| 282 |
+
print(f" β’ Checkpoint Path: {save_path}")
|
| 283 |
+
print(f" β’ Target Epochs: {args.epochs}")
|
| 284 |
+
print("=" * 85)
|
|
|
|
|
|
|
|
|
|
| 285 |
|
| 286 |
# 1. Load Data
|
| 287 |
train_dataset = SFTDataset(
|
|
|
|
| 299 |
|
| 300 |
train_loader = DataLoader(
|
| 301 |
train_dataset,
|
| 302 |
+
batch_size=micro_b,
|
| 303 |
shuffle=True,
|
| 304 |
collate_fn=sft_collate_fn,
|
| 305 |
num_workers=2,
|
|
|
|
| 307 |
)
|
| 308 |
val_loader = DataLoader(
|
| 309 |
val_dataset,
|
| 310 |
+
batch_size=micro_b,
|
| 311 |
shuffle=False,
|
| 312 |
collate_fn=sft_collate_fn,
|
| 313 |
num_workers=2,
|
|
|
|
| 374 |
]
|
| 375 |
optimizer = torch.optim.AdamW(optimizer_grouped_parameters, lr=args.max_lr, betas=(0.9, 0.95), eps=1e-8)
|
| 376 |
|
| 377 |
+
steps_per_epoch = math.ceil(len(train_loader) / grad_acc)
|
| 378 |
total_steps = steps_per_epoch * args.epochs
|
| 379 |
warmup_steps = max(10, int(total_steps * args.warmup_ratio))
|
| 380 |
+
print(f"π Training Plan: {steps_per_epoch} steps/epoch | Total: {total_steps} steps | Warmup: {warmup_steps} steps")
|
| 381 |
|
| 382 |
start_epoch = 1
|
| 383 |
global_step = 0
|
|
|
|
| 395 |
|
| 396 |
# Mixed Precision Setup
|
| 397 |
if torch.cuda.is_available():
|
| 398 |
+
autocast_ctx = torch.amp.autocast(device_type="cuda", dtype=hw["dtype"])
|
|
|
|
| 399 |
else:
|
|
|
|
| 400 |
autocast_ctx = contextlib.nullcontext()
|
|
|
|
| 401 |
|
| 402 |
# Evaluation Helper
|
| 403 |
@torch.no_grad()
|
|
|
|
| 432 |
total_tokens_trained = 0
|
| 433 |
model.train()
|
| 434 |
|
| 435 |
+
print("\n" + "=" * 85)
|
| 436 |
+
print(f"π STARTING SFT {args.version.upper()} MASTER TRAINING (AUTO-ADAPTIVE GPU ENGINE)")
|
| 437 |
+
print("=" * 85)
|
| 438 |
|
| 439 |
for epoch in range(start_epoch, args.epochs + 1):
|
| 440 |
print(f"\n--- Epoch {epoch}/{args.epochs} ---")
|
|
|
|
| 458 |
shift_logits = logits[..., :-1, :].contiguous()
|
| 459 |
shift_labels = labels[..., 1:].contiguous()
|
| 460 |
loss = F.cross_entropy(shift_logits.view(-1, shift_logits.size(-1)), shift_labels.view(-1), ignore_index=-100)
|
| 461 |
+
loss_scaled = loss / grad_acc
|
| 462 |
|
| 463 |
loss_scaled.backward()
|
| 464 |
accum_loss += loss.item()
|
| 465 |
epoch_loss += loss.item()
|
| 466 |
epoch_batches += 1
|
| 467 |
|
| 468 |
+
if (micro_idx + 1) % grad_acc == 0:
|
| 469 |
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
|
| 470 |
|
| 471 |
lr = get_lr(global_step, warmup_steps, total_steps, args.max_lr, args.min_lr)
|
|
|
|
| 476 |
optimizer.zero_grad(set_to_none=True)
|
| 477 |
global_step += 1
|
| 478 |
|
| 479 |
+
step_avg_loss = accum_loss / grad_acc
|
| 480 |
accum_loss = 0.0
|
| 481 |
|
| 482 |
if global_step % 10 == 0 or global_step == 1:
|
|
|
|
| 501 |
best_val_loss = v_loss
|
| 502 |
print(f" π New Best Validation Loss: {best_val_loss:.4f}! Saving checkpoint...")
|
| 503 |
|
|
|
|
| 504 |
save_payload = {
|
| 505 |
"model_state_dict": model.state_dict(),
|
| 506 |
"optimizer_state_dict": optimizer.state_dict(),
|
|
|
|
| 515 |
print(f" πΎ Saved checkpoint -> {save_path}\n")
|
| 516 |
|
| 517 |
# End of Epoch Handling
|
| 518 |
+
if (micro_idx + 1) % grad_acc != 0:
|
| 519 |
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
|
| 520 |
lr = get_lr(global_step, warmup_steps, total_steps, args.max_lr, args.min_lr)
|
| 521 |
for param_group in optimizer.param_groups:
|
|
|
|
| 529 |
|
| 530 |
# Final Evaluation & Save
|
| 531 |
final_val_loss, final_val_ppl = evaluate()
|
| 532 |
+
print("\n" + "=" * 85)
|
| 533 |
print(f"π SFT {args.version.upper()} MASTER TRAINING COMPLETED!")
|
| 534 |
print(f" β’ Final Validation Loss: {final_val_loss:.4f}")
|
| 535 |
print(f" β’ Final Perplexity: {final_val_ppl:.2f}")
|
| 536 |
print(f" β’ Total Active Tokens: {total_tokens_trained:,}")
|
| 537 |
print(f" β’ Total Time Taken: {(time.time() - start_time)/60:.2f} minutes")
|
| 538 |
+
print("=" * 85)
|
| 539 |
|
| 540 |
# Final Save
|
| 541 |
final_payload = {
|