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