qat_ultra_14b_for_tq_2.py
15.7 KB · 411 lines · python Raw
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!")