qat_ultra_14b_for_tq_2.py
| 1 | #%%writefile train_ultra_14b_bf16_ada_warmup.py |
| 2 | # ============================================================================= |
| 3 | # COPYRIGHT © 2026 Konstantin Vladimirovich Grabko. ALL RIGHTS RESERVED. |
| 4 | # JiRack Ultra Ternary Transformer |
| 5 | # |
| 6 | # CMS Manhattan JiRack Technology — PATENT PENDING |
| 7 | # train_ultra_14b_bf16_ada_warmup.py |
| 8 | # ============================================================================= |
| 9 | # QAT training with lambda warmup — JiRack Ultra 14B edition. |
| 10 | # Adapted from train_ultra_7b_bf16_ada_warmup.py (same structure). |
| 11 | # |
| 12 | # Changes vs the 7B script: |
| 13 | # [X-1] import from JiRackTernaryUltra_14b (14B constants: vocab 152064, |
| 14 | # hidden 5120, 48L, 40/8 heads, θ=1M, eps=1e-5 — same |
| 15 | # set_lambda/get_lambda API) |
| 16 | # [X-2] base weights: CMSManhattan/JiRackUltra_14b once published; until |
| 17 | # then set BASE_CHECKPOINT to a local .pt, or the HF load will 404. |
| 18 | # [X-3] memory: 14B is the ceiling for a single 96GB card — ~30 GB bf16 |
| 19 | # weights + ~30 GB grads + Adafactor factored state + activations. |
| 20 | # BATCH_SIZE=1, GRAD_ACCUM=10, gradient checkpointing mandatory, |
| 21 | # keep sequences <= 1024 to start. If OOM: shorten sequences first, |
| 22 | # then consider freezing embed/lm_head (~10% of params). |
| 23 | # [X-4] tokenizer gate vs vocab 152064: JiRackPrecisionTokenizer |
| 24 | # (151,779) fits — assert, NEVER resize. |
| 25 | # All 7B/1B/10B fixes preserved: [T-3] resume restores global_step |
| 26 | # (legacy checkpoints reconstruct it by inverting the sigmoid), |
| 27 | # [T-4] Adafactor state saved/restored, [T-5] lambda updated per |
| 28 | # accumulation window, [T-7] val loss logged with its lambda, atomic |
| 29 | # mid-shard autosave, pad via tokenizer's real pad_token_id. |
| 30 | # ============================================================================= |
| 31 | |
| 32 | import os |
| 33 | # Must be set BEFORE torch initializes CUDA. |
| 34 | os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True") |
| 35 | |
| 36 | import torch |
| 37 | import glob |
| 38 | import re |
| 39 | import gc |
| 40 | import math |
| 41 | import torch.nn as nn |
| 42 | from torch.utils.data import Dataset, DataLoader |
| 43 | from torch.nn.utils.rnn import pad_sequence |
| 44 | from tqdm import tqdm |
| 45 | from transformers.optimization import Adafactor |
| 46 | |
| 47 | # [X-1] user's 14B model class |
| 48 | from JiRackTernaryUltra_14b import JiRackTransformer, JiRackConfig |
| 49 | |
| 50 | # ========================= НАСТРОЙКИ ========================= |
| 51 | DATA_DIR = "/content/prepared_sft_data" |
| 52 | OUTPUT_DIR = "/content/JiRackUltra_14B_Checkpoints" |
| 53 | |
| 54 | # [X-2] Planned published repo — verify it exists before relying on it. |
| 55 | # Until it is published, point BASE_CHECKPOINT at a local .pt instead. |
| 56 | HF_BASE_MODEL = "CMSManhattan/JiRackUltra_14b" |
| 57 | BASE_CHECKPOINT = None # e.g. "/content/ultra14b_base.pt" |
| 58 | |
| 59 | # [X-3] 14B on a 96GB card — this is tight. Do not raise BATCH_SIZE |
| 60 | # before confirming headroom with nvidia-smi during a full fwd+bwd. |
| 61 | BATCH_SIZE = 1 |
| 62 | GRAD_ACCUM = 10 |
| 63 | LR = 4e-5 |
| 64 | VAL_RATIO = 0.05 |
| 65 | |
| 66 | # === Lambda Warmup Settings === |
| 67 | # Same gentle sigmoid; measured in MICRO-batches. |
| 68 | # 5000 micro-batches = 500 optimizer steps at GRAD_ACCUM=10. |
| 69 | LAMBDA_WARMUP_STEPS = 5000 |
| 70 | MAX_LAMBDA = 1.0 |
| 71 | SAVE_OPTIMIZER = True # [T-4]; ~30 GB checkpoints with optimizer — |
| 72 | # set False if disk is tight, warmup quality |
| 73 | # suffers only across restarts |
| 74 | AUTOSAVE_EVERY = 1000 # mid-shard autosave, 0 = off |
| 75 | |
| 76 | # [X-4] tokenizer — JiRackPrecisionTokenizer fits the 14B matrix too |
| 77 | TOKENIZER_REPO = "CMSManhattan/JiRackPrecisionTokenizer" |
| 78 | # ============================================================= |
| 79 | |
| 80 | os.makedirs(OUTPUT_DIR, exist_ok=True) |
| 81 | |
| 82 | torch.backends.cuda.matmul.allow_tf32 = True |
| 83 | torch.backends.cudnn.allow_tf32 = True |
| 84 | |
| 85 | print("🚀 Loading JiRack Ultra 14B Ternary + Lambda Warmup...") |
| 86 | |
| 87 | config = JiRackConfig() |
| 88 | |
| 89 | # [X-4] load once, keep it — assert fit against the padded matrix (152064), |
| 90 | # NEVER resize_token_embeddings (151,779 < 152,064: shrinking would corrupt |
| 91 | # the matrix; the "Must resize" note on the tokenizer's card only applies |
| 92 | # when tokenizer vocab > model vocab). |
| 93 | from transformers import AutoTokenizer |
| 94 | tokenizer = AutoTokenizer.from_pretrained(TOKENIZER_REPO) |
| 95 | assert len(tokenizer) <= config.vocab_size, ( |
| 96 | f"tokenizer ({len(tokenizer)}) > padded matrix ({config.vocab_size}) — " |
| 97 | f"do NOT resize_token_embeddings, fix the tokenizer/data instead" |
| 98 | ) |
| 99 | print(f"✅ Tokenizer fits padded matrix: {len(tokenizer)} <= {config.vocab_size}") |
| 100 | |
| 101 | # Trust the LOADED object over any card text; fail loudly if pad is unset. |
| 102 | PAD_ID = tokenizer.pad_token_id |
| 103 | if PAD_ID is None: |
| 104 | PAD_ID = tokenizer.eos_token_id |
| 105 | print(f"⚠️ tokenizer.pad_token_id is None, falling back to eos_token_id={PAD_ID}") |
| 106 | assert PAD_ID is not None, "tokenizer has neither pad_token_id nor eos_token_id set" |
| 107 | print(f"✅ Using pad_token_id={PAD_ID} (eos_token_id={tokenizer.eos_token_id})") |
| 108 | |
| 109 | # [X-3] gradient checkpointing is mandatory at 14B |
| 110 | model = JiRackTransformer(config, use_checkpoint=True) |
| 111 | model.to("cuda") |
| 112 | |
| 113 | for block in model.blocks: |
| 114 | block.use_checkpoint = True |
| 115 | print("✅ Gradient checkpointing enabled") |
| 116 | |
| 117 | optimizer = Adafactor( |
| 118 | model.parameters(), |
| 119 | lr=LR, |
| 120 | eps=(1e-30, 1e-3), |
| 121 | clip_threshold=1.0, |
| 122 | decay_rate=-0.8, |
| 123 | weight_decay=0.0001, |
| 124 | scale_parameter=False, |
| 125 | relative_step=False, |
| 126 | warmup_init=False, |
| 127 | ) |
| 128 | |
| 129 | criterion = nn.CrossEntropyLoss(ignore_index=-100) |
| 130 | |
| 131 | |
| 132 | # ==================== LAMBDA SCHEDULE ==================== |
| 133 | def lambda_schedule(step: int) -> float: |
| 134 | """Sigmoid ramp over LAMBDA_WARMUP_STEPS micro-batches.""" |
| 135 | if step >= LAMBDA_WARMUP_STEPS: |
| 136 | return MAX_LAMBDA |
| 137 | return MAX_LAMBDA / ( |
| 138 | 1.0 + math.exp(-10 * (step - LAMBDA_WARMUP_STEPS / 2) / LAMBDA_WARMUP_STEPS) |
| 139 | ) |
| 140 | |
| 141 | |
| 142 | def invert_lambda_schedule(lam: float) -> int: |
| 143 | """Reconstruct global_step from a saved lambda (legacy checkpoints). |
| 144 | Inverse of the sigmoid above.""" |
| 145 | if lam >= MAX_LAMBDA * 0.999: |
| 146 | return LAMBDA_WARMUP_STEPS |
| 147 | if lam <= 1e-6: |
| 148 | return 0 |
| 149 | p = lam / MAX_LAMBDA |
| 150 | x = -math.log(1.0 / p - 1.0) # logit |
| 151 | return int(round(x * LAMBDA_WARMUP_STEPS / 10 + LAMBDA_WARMUP_STEPS / 2)) |
| 152 | |
| 153 | |
| 154 | # ==================== CHECKPOINT LOAD ==================== |
| 155 | def load_any_checkpoint(path, model, optimizer): |
| 156 | """Handles both new-format dicts and legacy plain state_dicts. |
| 157 | Returns restored global_step.""" |
| 158 | ckpt = torch.load(path, map_location="cpu", weights_only=True) |
| 159 | |
| 160 | if isinstance(ckpt, dict) and "model" in ckpt: |
| 161 | # [T-3] new format |
| 162 | missing, unexpected = model.load_state_dict(ckpt["model"], strict=False) |
| 163 | assert not unexpected, unexpected |
| 164 | assert all(k.endswith("lambda_") for k in missing), missing |
| 165 | if SAVE_OPTIMIZER and "optimizer" in ckpt and ckpt["optimizer"] is not None: |
| 166 | try: |
| 167 | optimizer.load_state_dict(ckpt["optimizer"]) |
| 168 | print("✅ Optimizer state restored") |
| 169 | except Exception as e: |
| 170 | print(f"⚠️ Optimizer state not restored ({e}); continuing fresh") |
| 171 | step = int(ckpt.get("global_step", 0)) |
| 172 | print(f"✅ Resumed (new format): global_step={step}, " |
| 173 | f"lambda={ckpt.get('lambda', 'n/a')}") |
| 174 | return step |
| 175 | |
| 176 | # legacy: plain state_dict |
| 177 | missing, unexpected = model.load_state_dict(ckpt, strict=False) |
| 178 | assert not unexpected, unexpected |
| 179 | assert all(k.endswith("lambda_") for k in missing), missing |
| 180 | lam = model.get_lambda() |
| 181 | step = invert_lambda_schedule(lam) |
| 182 | print(f"✅ Resumed (legacy format): lambda={lam:.4f} -> " |
| 183 | f"reconstructed global_step={step}") |
| 184 | return step |
| 185 | |
| 186 | |
| 187 | def load_hf_base(model): |
| 188 | """[X-2] Pull the published JiRack Ultra 14B and map it in with the |
| 189 | model's own load_hf_state_dict. Use BASE_CHECKPOINT until the repo is |
| 190 | published.""" |
| 191 | from transformers import AutoModelForCausalLM |
| 192 | print(f"⬇️ Loading HF base: {HF_BASE_MODEL}") |
| 193 | # bf16, not fp32: a 14B fp32 state dict (~59 GB) would double peak host |
| 194 | # RAM during mapping; bf16 (~30 GB) is already the working format. |
| 195 | hf = AutoModelForCausalLM.from_pretrained(HF_BASE_MODEL, torch_dtype=torch.bfloat16) |
| 196 | real_missing, unexpected = model.load_hf_state_dict(hf.state_dict(), strict=True) |
| 197 | del hf |
| 198 | gc.collect() |
| 199 | torch.cuda.empty_cache() |
| 200 | return 0 # fresh QAT run starts at global_step 0 |
| 201 | |
| 202 | |
| 203 | # Shard-named checkpoints (define which shards are already done)... |
| 204 | checkpoints = sorted( |
| 205 | glob.glob(os.path.join(OUTPUT_DIR, "jirack_ultra14b_data_*.pt")), |
| 206 | key=lambda x: int(re.search(r"data_(\d+)", x).group(1)), |
| 207 | ) |
| 208 | # ...but for MODEL STATE, resume from whichever .pt is newest on disk. |
| 209 | all_ckpts = glob.glob(os.path.join(OUTPUT_DIR, "*.pt")) |
| 210 | |
| 211 | global_step = 0 |
| 212 | if all_ckpts: |
| 213 | LATEST_CKPT = max(all_ckpts, key=os.path.getmtime) |
| 214 | print(f"📦 Resuming from (newest on disk): {LATEST_CKPT}") |
| 215 | global_step = load_any_checkpoint(LATEST_CKPT, model, optimizer) |
| 216 | elif BASE_CHECKPOINT is not None: |
| 217 | print(f"📦 Loading base checkpoint: {BASE_CHECKPOINT}") |
| 218 | assert os.path.exists(BASE_CHECKPOINT), f"{BASE_CHECKPOINT} not found" |
| 219 | global_step = load_any_checkpoint(BASE_CHECKPOINT, model, optimizer) |
| 220 | else: |
| 221 | # [X-2] default path: published JiRack Ultra 14B (once live) |
| 222 | global_step = load_hf_base(model) |
| 223 | |
| 224 | model.to("cuda") |
| 225 | |
| 226 | # [T-3] lambda follows global_step from here on. |
| 227 | model.set_lambda(lambda_schedule(global_step)) |
| 228 | print(f"🔧 Warmup: {LAMBDA_WARMUP_STEPS} steps | starting at " |
| 229 | f"step={global_step}, lambda={model.get_lambda():.4f}") |
| 230 | |
| 231 | |
| 232 | # ========================= DATASET ========================= |
| 233 | class ShardDataset(Dataset): |
| 234 | def __init__(self, data_list): |
| 235 | self.data = data_list |
| 236 | |
| 237 | def __len__(self): |
| 238 | return len(self.data) |
| 239 | |
| 240 | def __getitem__(self, idx): |
| 241 | return self.data[idx] |
| 242 | |
| 243 | |
| 244 | def collate_fn(batch): |
| 245 | # pad with the tokenizer's real pad_token_id, not a hardcoded 0 |
| 246 | input_ids = pad_sequence( |
| 247 | [item["input_ids"] for item in batch], batch_first=True, padding_value=PAD_ID |
| 248 | ) |
| 249 | attention_mask = pad_sequence( |
| 250 | [item.get("attention_mask", torch.ones_like(item["input_ids"])) |
| 251 | for item in batch], |
| 252 | batch_first=True, padding_value=0, |
| 253 | ) |
| 254 | labels = input_ids.clone() |
| 255 | labels[attention_mask == 0] = -100 |
| 256 | return {"input_ids": input_ids, "labels": labels} |
| 257 | |
| 258 | |
| 259 | # ========================= CHECKPOINT SAVE ========================= |
| 260 | def save_checkpoint(path, model, optimizer, global_step, lam): |
| 261 | # bf16 on disk for economy (fp32 master precision lost across restarts only) |
| 262 | model_sd = {k: v.detach().to(torch.bfloat16).cpu() |
| 263 | for k, v in model.state_dict().items()} |
| 264 | ckpt = { |
| 265 | "model": model_sd, |
| 266 | "optimizer": optimizer.state_dict() if SAVE_OPTIMIZER else None, |
| 267 | "global_step": global_step, |
| 268 | "lambda": lam, |
| 269 | } |
| 270 | torch.save(ckpt, path) |
| 271 | |
| 272 | |
| 273 | # ========================= TRAINING ========================= |
| 274 | all_shards = sorted( |
| 275 | glob.glob(f"{DATA_DIR}/sft_data_*.pt"), |
| 276 | key=lambda x: int(re.search(r"sft_data_(\d+)", x).group(1)), |
| 277 | ) |
| 278 | |
| 279 | last_done = -1 # -1 = no shards processed yet (shard numbering starts at 0!) |
| 280 | if checkpoints: |
| 281 | m = re.search(r"data_(\d+)", checkpoints[-1]) |
| 282 | last_done = int(m.group(1)) if m else -1 |
| 283 | |
| 284 | for shard_path in all_shards: |
| 285 | shard_name = os.path.basename(shard_path) |
| 286 | shard_num = int(re.search(r"sft_data_(\d+)", shard_name).group(1)) |
| 287 | |
| 288 | if shard_num <= last_done: |
| 289 | print(f"⏭ Skipping already processed: {shard_name}") |
| 290 | continue |
| 291 | |
| 292 | print(f"\n🔥 Starting shard: {shard_name}") |
| 293 | |
| 294 | raw_shard_data = torch.load(shard_path, map_location="cpu", weights_only=False) |
| 295 | val_size = int(len(raw_shard_data) * VAL_RATIO) |
| 296 | train_size = len(raw_shard_data) - val_size |
| 297 | |
| 298 | train_data, val_data = torch.utils.data.random_split( |
| 299 | raw_shard_data, [train_size, val_size], |
| 300 | generator=torch.Generator().manual_seed(42), |
| 301 | ) |
| 302 | |
| 303 | train_loader = DataLoader( |
| 304 | ShardDataset(train_data), batch_size=BATCH_SIZE, shuffle=True, |
| 305 | collate_fn=collate_fn, pin_memory=True, |
| 306 | ) |
| 307 | |
| 308 | model.train() |
| 309 | pbar = tqdm(train_loader, desc=f"Shard {shard_num}", dynamic_ncols=True) |
| 310 | optimizer.zero_grad() |
| 311 | |
| 312 | lambda_value = lambda_schedule(global_step) |
| 313 | model.set_lambda(lambda_value) |
| 314 | |
| 315 | for step, batch in enumerate(pbar): |
| 316 | # [T-5] update lambda only at accumulation-window boundaries. |
| 317 | if step % GRAD_ACCUM == 0: |
| 318 | lambda_value = lambda_schedule(global_step) |
| 319 | model.set_lambda(lambda_value) |
| 320 | |
| 321 | input_ids = batch["input_ids"].to("cuda", non_blocking=True) |
| 322 | labels = batch["labels"].to("cuda", non_blocking=True) |
| 323 | |
| 324 | with torch.amp.autocast("cuda", dtype=torch.bfloat16): |
| 325 | logits = model(input_ids) |
| 326 | if isinstance(logits, tuple): |
| 327 | logits = logits[0] |
| 328 | loss = criterion( |
| 329 | logits[..., :-1, :].reshape(-1, config.vocab_size), |
| 330 | labels[..., 1:].reshape(-1), |
| 331 | ) |
| 332 | loss = loss / GRAD_ACCUM |
| 333 | |
| 334 | if torch.isnan(loss) or torch.isinf(loss): |
| 335 | print(f"\n⚠️ NaN/Inf loss at step {global_step} " |
| 336 | f"(lambda={lambda_value:.4f}) — window dropped") |
| 337 | optimizer.zero_grad(set_to_none=True) |
| 338 | torch.cuda.empty_cache() |
| 339 | global_step += 1 |
| 340 | continue |
| 341 | |
| 342 | loss.backward() |
| 343 | |
| 344 | if (step + 1) % GRAD_ACCUM == 0: |
| 345 | torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) |
| 346 | optimizer.step() |
| 347 | optimizer.zero_grad() |
| 348 | |
| 349 | if step % 10 == 0: |
| 350 | pbar.set_postfix({ |
| 351 | "loss": f"{loss.item() * GRAD_ACCUM:.4f}", |
| 352 | "lambda": f"{lambda_value:.4f}", |
| 353 | "gstep": global_step, |
| 354 | }) |
| 355 | |
| 356 | # Mid-shard autosave (atomic: tmp -> rename), at accumulation |
| 357 | # boundaries only (grads flushed, state clean). |
| 358 | if (AUTOSAVE_EVERY > 0 and global_step > 0 |
| 359 | and global_step % AUTOSAVE_EVERY == 0 |
| 360 | and (step + 1) % GRAD_ACCUM == 0): |
| 361 | autosave_path = os.path.join(OUTPUT_DIR, "autosave_latest.pt") |
| 362 | tmp_path = autosave_path + ".tmp" |
| 363 | save_checkpoint(tmp_path, model, optimizer, global_step, lambda_value) |
| 364 | os.replace(tmp_path, autosave_path) |
| 365 | pbar.write(f"💾 autosave @ gstep={global_step}, " |
| 366 | f"lambda={lambda_value:.4f}") |
| 367 | |
| 368 | global_step += 1 |
| 369 | |
| 370 | # ==================== Validation ==================== |
| 371 | print("🧪 Validating...") |
| 372 | model.eval() |
| 373 | total_val_loss = 0.0 |
| 374 | val_steps = 0 |
| 375 | |
| 376 | val_loader = DataLoader( |
| 377 | ShardDataset(val_data), batch_size=BATCH_SIZE, shuffle=False, |
| 378 | collate_fn=collate_fn, pin_memory=True, |
| 379 | ) |
| 380 | |
| 381 | with torch.no_grad(): |
| 382 | for batch in tqdm(val_loader, desc="Validating", leave=False): |
| 383 | input_ids = batch["input_ids"].to("cuda", non_blocking=True) |
| 384 | labels = batch["labels"].to("cuda", non_blocking=True) |
| 385 | with torch.amp.autocast("cuda", dtype=torch.bfloat16): |
| 386 | logits = model(input_ids) |
| 387 | if isinstance(logits, tuple): |
| 388 | logits = logits[0] |
| 389 | v_loss = criterion( |
| 390 | logits[..., :-1, :].reshape(-1, config.vocab_size), |
| 391 | labels[..., 1:].reshape(-1), |
| 392 | ) |
| 393 | if not (torch.isnan(v_loss) or torch.isinf(v_loss)): |
| 394 | total_val_loss += v_loss.item() |
| 395 | val_steps += 1 |
| 396 | |
| 397 | avg_val_loss = total_val_loss / val_steps if val_steps > 0 else float("inf") |
| 398 | # [T-7] val loss only comparable at the SAME lambda |
| 399 | print(f"📊 Shard {shard_num} — Val Loss: {avg_val_loss:.4f} " |
| 400 | f"@ lambda={lambda_value:.4f} (gstep={global_step})") |
| 401 | |
| 402 | # ==================== Save ==================== |
| 403 | save_path = os.path.join(OUTPUT_DIR, f"jirack_ultra14b_data_{shard_num}.pt") |
| 404 | save_checkpoint(save_path, model, optimizer, global_step, lambda_value) |
| 405 | print(f"💾 Saved: {save_path}") |
| 406 | |
| 407 | del raw_shard_data |
| 408 | torch.cuda.empty_cache() |
| 409 | gc.collect() |
| 410 | |
| 411 | print("🏁 Training finished with Lambda Warmup!") |