qat_ultra_1b_for_TQ_2.py
| 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!") |