train_jirack_1b_lambda_warmup.py
| 1 | %%writefile train_jirack_1b_lambda_warmup.py |
| 2 | |
| 3 | # train_jirack_1b_lambda_warmup.py |
| 4 | # ============================================================================= |
| 5 | # QAT training with lambda warmup -- адаптировано с train_..._10b_...py |
| 6 | # под JiRack-1b (DeepSeek-R1-Distill-Qwen-1.5B, JiRackTernaryUltra_1b.py). |
| 7 | # |
| 8 | # ЧТО ИЗМЕНЕНО ОТНОСИТЕЛЬНО ВЕРСИИ ДЛЯ 10B (и почему): |
| 9 | # |
| 10 | # [A-1] Импорт из JiRackTernaryUltra_1b вместо JiRackTernaryPyTorch_10b_fixed. |
| 11 | # Классы называются JiRackTransformer / JiRackConfig (без суффикса |
| 12 | # размера) -- см. Ternarization_instructions_1b.md, раздел 0 и 3. |
| 13 | # !! ПРОВЕРЬ реальные имена классов в файле перед запуском. |
| 14 | # |
| 15 | # [A-2] BASE_CHECKPOINT указывает на уже существующий |
| 16 | # /mnt/nfs_share/JiRackUlrta_1/model.pt -- это архитектурно |
| 17 | # сконвертированная из HF модель на lambda=0 (QAT ещё не запускали, |
| 18 | # подтверждено в разговоре). Это ОБЫЧНЫЙ plain state_dict, не |
| 19 | # {model, optimizer, global_step, lambda} -- значит он пойдёт по |
| 20 | # "legacy" ветке load_any_checkpoint() и warmup начнётся с |
| 21 | # global_step=0. Это ожидаемо и правильно для первого запуска. |
| 22 | # |
| 23 | # [A-3] Device определяется автоматически (cuda, если доступна, иначе |
| 24 | # cpu). В версии для 10B было жёстко зашито .to("cuda") и |
| 25 | # autocast("cuda", ...) в расчёте на 96GB Blackwell в Colab. |
| 26 | # Для 1.5B модели это не обязательно тот же сервер/GPU -- поэтому |
| 27 | # весь код ниже работает в обоих случаях без правки руками. |
| 28 | # |
| 29 | # [A-4] BATCH_SIZE/GRAD_ACCUM уменьшены как безопасный дефолт -- 1.5B |
| 30 | # занимает на порядок меньше памяти, чем 10B, так что на GPU эти |
| 31 | # числа наверняка можно поднять. Но если реально гоняешь на CPU |
| 32 | # (jirack2) -- держи в уме, что QAT с backward-проходом на CPU |
| 33 | # на порядки медленнее, чем только forward (для 27B на этом же |
| 34 | # сервере forward был ~14-16 сек/токен -- backward + warmup на |
| 35 | # тысячи шагов может занять очень долго). Стоит сначала прогнать |
| 36 | # пробный запуск на малом числе шагов и посчитать время на шаг, |
| 37 | # прежде чем оставлять это надолго. |
| 38 | # |
| 39 | # [A-5] Gradient checkpointing включается только если у model.blocks[i] |
| 40 | # реально есть атрибут use_checkpoint -- в версии для 10B это |
| 41 | # предполагалось безусловно. Для 1.5B память куда менее узкое |
| 42 | # место, так что checkpointing может быть не нужен вообще (он |
| 43 | # платит временем за экономию памяти, которая тут не так важна). |
| 44 | # |
| 45 | # [A-6] DATA_DIR/OUTPUT_DIR -- ЗАГЛУШКИ под реальные пути проекта. |
| 46 | # Логика шардирования (sft_data_N.pt) скопирована как есть из |
| 47 | # версии для 10B -- если для 1b данные готовятся иначе (другой |
| 48 | # формат шардов, другое имя файлов), это нужно поправить отдельно. |
| 49 | # |
| 50 | # Всё остальное (lambda_schedule, invert_lambda_schedule, |
| 51 | # load_any_checkpoint, save_checkpoint, цикл по шардам) перенесено |
| 52 | # без изменений в логике -- она не зависит от размера модели. |
| 53 | # ============================================================================= |
| 54 | |
| 55 | import os |
| 56 | os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True") |
| 57 | |
| 58 | import torch |
| 59 | import glob |
| 60 | import re |
| 61 | import gc |
| 62 | import math |
| 63 | import torch.nn as nn |
| 64 | from torch.utils.data import Dataset, DataLoader |
| 65 | from torch.nn.utils.rnn import pad_sequence |
| 66 | from tqdm import tqdm |
| 67 | from transformers.optimization import Adafactor |
| 68 | |
| 69 | # [A-1] импорт под 1b -- ПРОВЕРЬ реальные имена классов в файле |
| 70 | from JiRackTernaryUltra_1b import JiRackTransformer, JiRackConfig |
| 71 | |
| 72 | # ========================= НАСТРОЙКИ ========================= |
| 73 | # [A-6] заглушки -- поставь реальные пути |
| 74 | DATA_DIR = "/mnt/nfs_share/JiRackUlrta_1/prepared_sft_data" |
| 75 | OUTPUT_DIR = "/mnt/nfs_share/JiRackUlrta_1/qat_checkpoints" |
| 76 | # [A-2] существующий архитектурно-сконвертированный чекпоинт, lambda=0 |
| 77 | BASE_CHECKPOINT = "/mnt/nfs_share/JiRackUlrta_1/model.pt" |
| 78 | |
| 79 | # [A-4] уменьшенные дефолты под 1.5B -- подстроить по факту после |
| 80 | # пробного прогона на нескольких шагах |
| 81 | BATCH_SIZE = 4 |
| 82 | GRAD_ACCUM = 4 |
| 83 | LR = 4e-5 |
| 84 | VAL_RATIO = 0.05 |
| 85 | |
| 86 | # === Lambda Warmup Settings === |
| 87 | # Держим тот же порядок величины, что и для 10B (5000 микро-батчей). |
| 88 | # Для 1.5B можно попробовать короче, но лучше сначала прогнать как есть |
| 89 | # и посмотреть на кривую loss/lambda, прежде чем сокращать. |
| 90 | LAMBDA_WARMUP_STEPS = 5000 # в МИКРО-батчах |
| 91 | MAX_LAMBDA = 1.0 |
| 92 | SAVE_OPTIMIZER = True |
| 93 | AUTOSAVE_EVERY = 1000 |
| 94 | # ============================================================= |
| 95 | |
| 96 | os.makedirs(OUTPUT_DIR, exist_ok=True) |
| 97 | |
| 98 | # [A-3] device auto-detect вместо жёсткого "cuda" |
| 99 | DEVICE = "cuda" if torch.cuda.is_available() else "cpu" |
| 100 | AUTOCAST_DTYPE = torch.bfloat16 # bf16 autocast работает и на cuda, и на cpu |
| 101 | print(f"🖥️ Device: {DEVICE}") |
| 102 | if DEVICE == "cpu": |
| 103 | print("⚠️ CUDA недоступна -- обучение пойдёт на CPU. Backward-проход " |
| 104 | "для полноценного warmup может быть очень медленным. Рекомендуется " |
| 105 | "сначала прогнать 20-50 микро-батчей и замерить время на шаг, " |
| 106 | "прежде чем оставлять на LAMBDA_WARMUP_STEPS шагов без присмотра.") |
| 107 | |
| 108 | torch.backends.cuda.matmul.allow_tf32 = True |
| 109 | torch.backends.cudnn.allow_tf32 = True |
| 110 | |
| 111 | print("🚀 Loading JiRack 1b Ternary + Lambda Warmup...") |
| 112 | |
| 113 | config = JiRackConfig() |
| 114 | model = JiRackTransformer(config) |
| 115 | model.to(DEVICE) |
| 116 | |
| 117 | # [A-5] gradient checkpointing только если блоки его поддерживают |
| 118 | if hasattr(model, "blocks"): |
| 119 | n_checkpointed = 0 |
| 120 | for block in model.blocks: |
| 121 | if hasattr(block, "use_checkpoint"): |
| 122 | block.use_checkpoint = True |
| 123 | n_checkpointed += 1 |
| 124 | if n_checkpointed: |
| 125 | print(f"✅ Gradient checkpointing enabled on {n_checkpointed} blocks") |
| 126 | else: |
| 127 | print("ℹ️ Блоки не имеют атрибута use_checkpoint -- checkpointing " |
| 128 | "пропущен (для 1.5B это обычно не критично).") |
| 129 | |
| 130 | optimizer = Adafactor( |
| 131 | model.parameters(), |
| 132 | lr=LR, |
| 133 | eps=(1e-30, 1e-3), |
| 134 | clip_threshold=1.0, |
| 135 | decay_rate=-0.8, |
| 136 | weight_decay=0.0001, |
| 137 | scale_parameter=False, |
| 138 | relative_step=False, |
| 139 | warmup_init=False, |
| 140 | ) |
| 141 | |
| 142 | criterion = nn.CrossEntropyLoss(ignore_index=-100) |
| 143 | |
| 144 | |
| 145 | # ==================== LAMBDA SCHEDULE ==================== |
| 146 | def lambda_schedule(step: int) -> float: |
| 147 | """Sigmoid ramp over LAMBDA_WARMUP_STEPS micro-batches.""" |
| 148 | if step >= LAMBDA_WARMUP_STEPS: |
| 149 | return MAX_LAMBDA |
| 150 | return MAX_LAMBDA / ( |
| 151 | 1.0 + math.exp(-10 * (step - LAMBDA_WARMUP_STEPS / 2) / LAMBDA_WARMUP_STEPS) |
| 152 | ) |
| 153 | |
| 154 | |
| 155 | def invert_lambda_schedule(lam: float) -> int: |
| 156 | """Reconstruct global_step from a saved lambda (legacy checkpoints |
| 157 | that stored only the model state_dict). Inverse of the sigmoid above.""" |
| 158 | if lam >= MAX_LAMBDA * 0.999: |
| 159 | return LAMBDA_WARMUP_STEPS |
| 160 | if lam <= 1e-6: |
| 161 | return 0 |
| 162 | p = lam / MAX_LAMBDA |
| 163 | x = -math.log(1.0 / p - 1.0) # logit |
| 164 | return int(round(x * LAMBDA_WARMUP_STEPS / 10 + LAMBDA_WARMUP_STEPS / 2)) |
| 165 | |
| 166 | |
| 167 | # ==================== CHECKPOINT LOAD ==================== |
| 168 | def load_any_checkpoint(path, model, optimizer): |
| 169 | """Handles both new-format dicts and legacy plain state_dicts. |
| 170 | Returns restored global_step.""" |
| 171 | ckpt = torch.load(path, map_location="cpu", weights_only=True) |
| 172 | |
| 173 | if isinstance(ckpt, dict) and "model" in ckpt: |
| 174 | # новый формат (уже был обучен этим же скриптом раньше) |
| 175 | missing, unexpected = model.load_state_dict(ckpt["model"], strict=False) |
| 176 | assert not unexpected, unexpected |
| 177 | assert all(k.endswith("lambda_") for k in missing), missing |
| 178 | if SAVE_OPTIMIZER and "optimizer" in ckpt and ckpt["optimizer"] is not None: |
| 179 | try: |
| 180 | optimizer.load_state_dict(ckpt["optimizer"]) |
| 181 | print("✅ Optimizer state restored") |
| 182 | except Exception as e: |
| 183 | print(f"⚠️ Optimizer state not restored ({e}); continuing fresh") |
| 184 | step = int(ckpt.get("global_step", 0)) |
| 185 | print(f"✅ Resumed (new format): global_step={step}, " |
| 186 | f"lambda={ckpt.get('lambda', 'n/a')}") |
| 187 | return step |
| 188 | |
| 189 | # legacy: plain state_dict -- сюда попадёт исходный model.pt (lambda=0) |
| 190 | missing, unexpected = model.load_state_dict(ckpt, strict=False) |
| 191 | assert not unexpected, unexpected |
| 192 | assert all(k.endswith("lambda_") for k in missing), missing |
| 193 | |
| 194 | # [A-2] get_lambda() может не существовать в JiRackTransformer (1b) -- |
| 195 | # если так, считаем lambda=0.0, что для исходного model.pt и так верно. |
| 196 | if hasattr(model, "get_lambda"): |
| 197 | lam = model.get_lambda() |
| 198 | else: |
| 199 | lam = 0.0 |
| 200 | print("ℹ️ У model нет get_lambda() -- считаю lambda=0.0 " |
| 201 | "(корректно для непройденного через warmup model.pt).") |
| 202 | step = invert_lambda_schedule(lam) |
| 203 | print(f"✅ Resumed (legacy format): lambda={lam:.4f} -> " |
| 204 | f"reconstructed global_step={step}") |
| 205 | return step |
| 206 | |
| 207 | |
| 208 | # Шард-чекпоинты (определяют, какие шарды уже пройдены)... |
| 209 | checkpoints = sorted( |
| 210 | glob.glob(os.path.join(OUTPUT_DIR, "jirack_1b_data_*.pt")), |
| 211 | key=lambda x: int(re.search(r"data_(\d+)", x).group(1)), |
| 212 | ) |
| 213 | # ...но для СОСТОЯНИЯ МОДЕЛИ грузим самый свежий файл на диске -- это |
| 214 | # может быть mid-shard autosave, записанный после последнего шарда. |
| 215 | all_ckpts = glob.glob(os.path.join(OUTPUT_DIR, "*.pt")) |
| 216 | |
| 217 | global_step = 0 |
| 218 | if all_ckpts: |
| 219 | LATEST_CKPT = max(all_ckpts, key=os.path.getmtime) |
| 220 | print(f"📦 Resuming from (newest on disk): {LATEST_CKPT}") |
| 221 | global_step = load_any_checkpoint(LATEST_CKPT, model, optimizer) |
| 222 | else: |
| 223 | print(f"📦 Loading base checkpoint: {BASE_CHECKPOINT}") |
| 224 | assert os.path.exists(BASE_CHECKPOINT), ( |
| 225 | f"{BASE_CHECKPOINT} не найден -- проверь путь" |
| 226 | ) |
| 227 | global_step = load_any_checkpoint(BASE_CHECKPOINT, model, optimizer) |
| 228 | |
| 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() if hasattr(model, 'get_lambda') else lambda_schedule(global_step):.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 | input_ids = pad_sequence( |
| 248 | [item["input_ids"] for item in batch], batch_first=True, padding_value=0 |
| 249 | ) |
| 250 | attention_mask = pad_sequence( |
| 251 | [item.get("attention_mask", torch.ones_like(item["input_ids"])) |
| 252 | for item in batch], |
| 253 | batch_first=True, padding_value=0, |
| 254 | ) |
| 255 | labels = input_ids.clone() |
| 256 | labels[attention_mask == 0] = -100 |
| 257 | return {"input_ids": input_ids, "labels": labels} |
| 258 | |
| 259 | |
| 260 | # ========================= CHECKPOINT SAVE ========================= |
| 261 | def save_checkpoint(path, model, optimizer, global_step, lam): |
| 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 | # [A-6] шаблон имён шардов -- поправить под реальный формат данных для 1b |
| 275 | all_shards = sorted( |
| 276 | glob.glob(f"{DATA_DIR}/sft_data_*.pt"), |
| 277 | key=lambda x: int(re.search(r"sft_data_(\d+)", x).group(1)), |
| 278 | ) |
| 279 | if not all_shards: |
| 280 | print(f"⚠️ В {DATA_DIR} не найдено шардов sft_data_*.pt -- " |
| 281 | f"проверь путь/формат перед запуском.") |
| 282 | |
| 283 | last_done = -1 |
| 284 | if checkpoints: |
| 285 | m = re.search(r"data_(\d+)", checkpoints[-1]) |
| 286 | last_done = int(m.group(1)) if m else -1 |
| 287 | |
| 288 | for shard_path in all_shards: |
| 289 | shard_name = os.path.basename(shard_path) |
| 290 | shard_num = int(re.search(r"sft_data_(\d+)", shard_name).group(1)) |
| 291 | |
| 292 | if shard_num <= last_done: |
| 293 | print(f"⏭ Skipping already processed: {shard_name}") |
| 294 | continue |
| 295 | |
| 296 | print(f"\n🔥 Starting shard: {shard_name}") |
| 297 | |
| 298 | raw_shard_data = torch.load(shard_path, map_location="cpu", weights_only=False) |
| 299 | val_size = int(len(raw_shard_data) * VAL_RATIO) |
| 300 | train_size = len(raw_shard_data) - val_size |
| 301 | |
| 302 | train_data, val_data = torch.utils.data.random_split( |
| 303 | raw_shard_data, [train_size, val_size], |
| 304 | generator=torch.Generator().manual_seed(42), |
| 305 | ) |
| 306 | |
| 307 | train_loader = DataLoader( |
| 308 | ShardDataset(train_data), batch_size=BATCH_SIZE, shuffle=True, |
| 309 | collate_fn=collate_fn, pin_memory=(DEVICE == "cuda"), |
| 310 | ) |
| 311 | |
| 312 | model.train() |
| 313 | pbar = tqdm(train_loader, desc=f"Shard {shard_num}", dynamic_ncols=True) |
| 314 | optimizer.zero_grad() |
| 315 | |
| 316 | lambda_value = lambda_schedule(global_step) |
| 317 | model.set_lambda(lambda_value) |
| 318 | |
| 319 | for step, batch in enumerate(pbar): |
| 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(DEVICE, non_blocking=(DEVICE == "cuda")) |
| 325 | labels = batch["labels"].to(DEVICE, non_blocking=(DEVICE == "cuda")) |
| 326 | |
| 327 | with torch.amp.autocast(DEVICE, dtype=AUTOCAST_DTYPE): |
| 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 | if DEVICE == "cuda": |
| 342 | torch.cuda.empty_cache() |
| 343 | global_step += 1 |
| 344 | continue |
| 345 | |
| 346 | loss.backward() |
| 347 | |
| 348 | if (step + 1) % GRAD_ACCUM == 0: |
| 349 | torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) |
| 350 | optimizer.step() |
| 351 | optimizer.zero_grad() |
| 352 | |
| 353 | if step % 10 == 0: |
| 354 | pbar.set_postfix({ |
| 355 | "loss": f"{loss.item() * GRAD_ACCUM:.4f}", |
| 356 | "lambda": f"{lambda_value:.4f}", |
| 357 | "gstep": global_step, |
| 358 | }) |
| 359 | |
| 360 | if (AUTOSAVE_EVERY > 0 and global_step > 0 |
| 361 | and global_step % AUTOSAVE_EVERY == 0 |
| 362 | and (step + 1) % GRAD_ACCUM == 0): |
| 363 | autosave_path = os.path.join(OUTPUT_DIR, "autosave_latest.pt") |
| 364 | tmp_path = autosave_path + ".tmp" |
| 365 | save_checkpoint(tmp_path, model, optimizer, global_step, lambda_value) |
| 366 | os.replace(tmp_path, autosave_path) |
| 367 | pbar.write(f"💾 autosave @ gstep={global_step}, " |
| 368 | f"lambda={lambda_value:.4f}") |
| 369 | |
| 370 | global_step += 1 |
| 371 | |
| 372 | # ==================== Validation ==================== |
| 373 | print("🧪 Validating...") |
| 374 | model.eval() |
| 375 | total_val_loss = 0.0 |
| 376 | val_steps = 0 |
| 377 | |
| 378 | val_loader = DataLoader( |
| 379 | ShardDataset(val_data), batch_size=BATCH_SIZE, shuffle=False, |
| 380 | collate_fn=collate_fn, pin_memory=(DEVICE == "cuda"), |
| 381 | ) |
| 382 | |
| 383 | with torch.no_grad(): |
| 384 | for batch in tqdm(val_loader, desc="Validating", leave=False): |
| 385 | input_ids = batch["input_ids"].to(DEVICE, non_blocking=(DEVICE == "cuda")) |
| 386 | labels = batch["labels"].to(DEVICE, non_blocking=(DEVICE == "cuda")) |
| 387 | with torch.amp.autocast(DEVICE, dtype=AUTOCAST_DTYPE): |
| 388 | logits = model(input_ids) |
| 389 | if isinstance(logits, tuple): |
| 390 | logits = logits[0] |
| 391 | v_loss = criterion( |
| 392 | logits[..., :-1, :].reshape(-1, config.vocab_size), |
| 393 | labels[..., 1:].reshape(-1), |
| 394 | ) |
| 395 | if not (torch.isnan(v_loss) or torch.isinf(v_loss)): |
| 396 | total_val_loss += v_loss.item() |
| 397 | val_steps += 1 |
| 398 | |
| 399 | avg_val_loss = total_val_loss / val_steps if val_steps > 0 else float("inf") |
| 400 | print(f"📊 Shard {shard_num} -- Val Loss: {avg_val_loss:.4f} " |
| 401 | f"@ lambda={lambda_value:.4f} (gstep={global_step})") |
| 402 | |
| 403 | # ==================== Save ==================== |
| 404 | save_path = os.path.join(OUTPUT_DIR, f"jirack_1b_data_{shard_num}.pt") |
| 405 | save_checkpoint(save_path, model, optimizer, global_step, lambda_value) |
| 406 | print(f"💾 Saved: {save_path}") |
| 407 | |
| 408 | del raw_shard_data |
| 409 | if DEVICE == "cuda": |
| 410 | torch.cuda.empty_cache() |
| 411 | gc.collect() |
| 412 | |
| 413 | print("🏁 Training finished with Lambda Warmup!") |
| 414 | |