ViuAI commited on
Commit
374ce18
Β·
verified Β·
1 Parent(s): 57e0bf2

Upload code/train_sft.py with huggingface_hub

Browse files
Files changed (1) hide show
  1. code/train_sft.py +124 -72
code/train_sft.py CHANGED
@@ -1,12 +1,18 @@
1
  # ==============================================================================
2
- # πŸš€ ViuAI Sarus-500M β€” SFT v16 Master Training Engine (OOM-Resistant Edition)
3
  # ==============================================================================
4
- # KEY IMPROVEMENTS:
5
- # β€’ OOM-PROOF: Default micro-batch set to 2/4 with Grad Accum (Effective Batch = 128)
6
- # β€’ ZERO MEMORY SPIKES: torch.cuda.empty_cache() & efficient gradient accumulation
7
- # β€’ DYNAMIC VOCAB LOSS: Loss dynamically adapts to shift_logits.size(-1)
8
- # β€’ ACCURATE SPEED: Padded tokens excluded from throughput metrics
9
- # β€’ 100% RESUME-READY: Full state dict preservation & Hugging Face Auto-Sync
 
 
 
 
 
 
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
- # --- 1. Memory-Mapped High Performance Dataset ---
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
- # --- 2. Dynamic Padding & Strict 2048 Bound Collate Function ---
 
 
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
- # --- 3. Cosine Learning Rate Schedule with Warmup ---
 
 
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
- # --- 4. Auto-detect OOM-Safe Batch Size by Hardware VRAM ---
94
- def auto_detect_batch_params():
95
- if not torch.cuda.is_available():
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
- if not os.path.exists(path):
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
- if not os.path.exists(path):
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
- # --- 6. Main Training Runner ---
 
 
171
  def main():
172
- auto_b, auto_ga = auto_detect_batch_params()
173
 
174
- parser = argparse.ArgumentParser(description="ViuAI Sarus-500M β€” SFT v16 Master Training Engine")
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=auto_b, help=f"Micro batch size (Auto-detected: {auto_b})")
180
- parser.add_argument("--grad_accum", type=int, default=auto_ga, help=f"Gradient accumulation steps (Auto-detected: {auto_ga})")
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("=" * 80)
218
- print(f"πŸš€ ViuAI Sarus-500M β€” SFT {args.version.upper()} Unified Master Training Runner")
219
- print(f" β€’ Device: {torch.cuda.get_device_name(0) if torch.cuda.is_available() else 'CPU'}")
220
- print(f" β€’ Data Folder: {stage_data_dir}")
221
- print(f" β€’ Init Checkpoint: {init_ckpt}")
222
- print(f" β€’ Save Checkpoint: {save_path}")
223
- print(f" β€’ Micro Batch Size: {args.batch_size} (Grad Accum {args.grad_accum} -> Effective Batch {args.batch_size * args.grad_accum})")
224
- print(f" β€’ Epochs: {args.epochs}")
225
- print(f" β€’ Max / Min LR: {args.max_lr} / {args.min_lr} (Cosine Schedule)")
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=args.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=args.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) / args.grad_accum)
322
  total_steps = steps_per_epoch * args.epochs
323
  warmup_steps = max(10, int(total_steps * args.warmup_ratio))
324
- print(f"πŸ“Š Training Steps: {steps_per_epoch} steps/epoch | Total: {total_steps} steps | Warmup: {warmup_steps} steps")
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
- dtype = torch.bfloat16 if torch.cuda.is_bf16_supported() else torch.float16
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" + "=" * 80)
383
- print(f"🏁 STARTING SFT {args.version.upper()} MASTER TRAINING")
384
- print("=" * 80)
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 / args.grad_accum
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) % args.grad_accum == 0:
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 / args.grad_accum
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) % args.grad_accum != 0:
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" + "=" * 80)
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("=" * 80)
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 = {