%%writefile train_jirack_1b_lambda_warmup.py # train_jirack_1b_lambda_warmup.py # ============================================================================= # QAT training with lambda warmup -- адаптировано с train_..._10b_...py # под JiRack-1b (DeepSeek-R1-Distill-Qwen-1.5B, JiRackTernaryUltra_1b.py). # # ЧТО ИЗМЕНЕНО ОТНОСИТЕЛЬНО ВЕРСИИ ДЛЯ 10B (и почему): # # [A-1] Импорт из JiRackTernaryUltra_1b вместо JiRackTernaryPyTorch_10b_fixed. # Классы называются JiRackTransformer / JiRackConfig (без суффикса # размера) -- см. Ternarization_instructions_1b.md, раздел 0 и 3. # !! ПРОВЕРЬ реальные имена классов в файле перед запуском. # # [A-2] BASE_CHECKPOINT указывает на уже существующий # /mnt/nfs_share/JiRackUlrta_1/model.pt -- это архитектурно # сконвертированная из HF модель на lambda=0 (QAT ещё не запускали, # подтверждено в разговоре). Это ОБЫЧНЫЙ plain state_dict, не # {model, optimizer, global_step, lambda} -- значит он пойдёт по # "legacy" ветке load_any_checkpoint() и warmup начнётся с # global_step=0. Это ожидаемо и правильно для первого запуска. # # [A-3] Device определяется автоматически (cuda, если доступна, иначе # cpu). В версии для 10B было жёстко зашито .to("cuda") и # autocast("cuda", ...) в расчёте на 96GB Blackwell в Colab. # Для 1.5B модели это не обязательно тот же сервер/GPU -- поэтому # весь код ниже работает в обоих случаях без правки руками. # # [A-4] BATCH_SIZE/GRAD_ACCUM уменьшены как безопасный дефолт -- 1.5B # занимает на порядок меньше памяти, чем 10B, так что на GPU эти # числа наверняка можно поднять. Но если реально гоняешь на CPU # (jirack2) -- держи в уме, что QAT с backward-проходом на CPU # на порядки медленнее, чем только forward (для 27B на этом же # сервере forward был ~14-16 сек/токен -- backward + warmup на # тысячи шагов может занять очень долго). Стоит сначала прогнать # пробный запуск на малом числе шагов и посчитать время на шаг, # прежде чем оставлять это надолго. # # [A-5] Gradient checkpointing включается только если у model.blocks[i] # реально есть атрибут use_checkpoint -- в версии для 10B это # предполагалось безусловно. Для 1.5B память куда менее узкое # место, так что checkpointing может быть не нужен вообще (он # платит временем за экономию памяти, которая тут не так важна). # # [A-6] DATA_DIR/OUTPUT_DIR -- ЗАГЛУШКИ под реальные пути проекта. # Логика шардирования (sft_data_N.pt) скопирована как есть из # версии для 10B -- если для 1b данные готовятся иначе (другой # формат шардов, другое имя файлов), это нужно поправить отдельно. # # Всё остальное (lambda_schedule, invert_lambda_schedule, # load_any_checkpoint, save_checkpoint, цикл по шардам) перенесено # без изменений в логике -- она не зависит от размера модели. # ============================================================================= import os os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True") import torch import glob import re import gc import math import torch.nn as nn from torch.utils.data import Dataset, DataLoader from torch.nn.utils.rnn import pad_sequence from tqdm import tqdm from transformers.optimization import Adafactor # [A-1] импорт под 1b -- ПРОВЕРЬ реальные имена классов в файле from JiRackTernaryUltra_1b import JiRackTransformer, JiRackConfig # ========================= НАСТРОЙКИ ========================= # [A-6] заглушки -- поставь реальные пути DATA_DIR = "/mnt/nfs_share/JiRackUlrta_1/prepared_sft_data" OUTPUT_DIR = "/mnt/nfs_share/JiRackUlrta_1/qat_checkpoints" # [A-2] существующий архитектурно-сконвертированный чекпоинт, lambda=0 BASE_CHECKPOINT = "/mnt/nfs_share/JiRackUlrta_1/model.pt" # [A-4] уменьшенные дефолты под 1.5B -- подстроить по факту после # пробного прогона на нескольких шагах BATCH_SIZE = 4 GRAD_ACCUM = 4 LR = 4e-5 VAL_RATIO = 0.05 # === Lambda Warmup Settings === # Держим тот же порядок величины, что и для 10B (5000 микро-батчей). # Для 1.5B можно попробовать короче, но лучше сначала прогнать как есть # и посмотреть на кривую loss/lambda, прежде чем сокращать. LAMBDA_WARMUP_STEPS = 5000 # в МИКРО-батчах MAX_LAMBDA = 1.0 SAVE_OPTIMIZER = True AUTOSAVE_EVERY = 1000 # ============================================================= os.makedirs(OUTPUT_DIR, exist_ok=True) # [A-3] device auto-detect вместо жёсткого "cuda" DEVICE = "cuda" if torch.cuda.is_available() else "cpu" AUTOCAST_DTYPE = torch.bfloat16 # bf16 autocast работает и на cuda, и на cpu print(f"🖥️ Device: {DEVICE}") if DEVICE == "cpu": print("⚠️ CUDA недоступна -- обучение пойдёт на CPU. Backward-проход " "для полноценного warmup может быть очень медленным. Рекомендуется " "сначала прогнать 20-50 микро-батчей и замерить время на шаг, " "прежде чем оставлять на LAMBDA_WARMUP_STEPS шагов без присмотра.") torch.backends.cuda.matmul.allow_tf32 = True torch.backends.cudnn.allow_tf32 = True print("🚀 Loading JiRack 1b Ternary + Lambda Warmup...") config = JiRackConfig() model = JiRackTransformer(config) model.to(DEVICE) # [A-5] gradient checkpointing только если блоки его поддерживают if hasattr(model, "blocks"): n_checkpointed = 0 for block in model.blocks: if hasattr(block, "use_checkpoint"): block.use_checkpoint = True n_checkpointed += 1 if n_checkpointed: print(f"✅ Gradient checkpointing enabled on {n_checkpointed} blocks") else: print("ℹ️ Блоки не имеют атрибута use_checkpoint -- checkpointing " "пропущен (для 1.5B это обычно не критично).") optimizer = Adafactor( model.parameters(), lr=LR, eps=(1e-30, 1e-3), clip_threshold=1.0, decay_rate=-0.8, weight_decay=0.0001, scale_parameter=False, relative_step=False, warmup_init=False, ) criterion = nn.CrossEntropyLoss(ignore_index=-100) # ==================== LAMBDA SCHEDULE ==================== def lambda_schedule(step: int) -> float: """Sigmoid ramp over LAMBDA_WARMUP_STEPS micro-batches.""" if step >= LAMBDA_WARMUP_STEPS: return MAX_LAMBDA return MAX_LAMBDA / ( 1.0 + math.exp(-10 * (step - LAMBDA_WARMUP_STEPS / 2) / LAMBDA_WARMUP_STEPS) ) def invert_lambda_schedule(lam: float) -> int: """Reconstruct global_step from a saved lambda (legacy checkpoints that stored only the model state_dict). Inverse of the sigmoid above.""" if lam >= MAX_LAMBDA * 0.999: return LAMBDA_WARMUP_STEPS if lam <= 1e-6: return 0 p = lam / MAX_LAMBDA x = -math.log(1.0 / p - 1.0) # logit return int(round(x * LAMBDA_WARMUP_STEPS / 10 + LAMBDA_WARMUP_STEPS / 2)) # ==================== CHECKPOINT LOAD ==================== def load_any_checkpoint(path, model, optimizer): """Handles both new-format dicts and legacy plain state_dicts. Returns restored global_step.""" ckpt = torch.load(path, map_location="cpu", weights_only=True) if isinstance(ckpt, dict) and "model" in ckpt: # новый формат (уже был обучен этим же скриптом раньше) missing, unexpected = model.load_state_dict(ckpt["model"], strict=False) assert not unexpected, unexpected assert all(k.endswith("lambda_") for k in missing), missing if SAVE_OPTIMIZER and "optimizer" in ckpt and ckpt["optimizer"] is not None: try: optimizer.load_state_dict(ckpt["optimizer"]) print("✅ Optimizer state restored") except Exception as e: print(f"⚠️ Optimizer state not restored ({e}); continuing fresh") step = int(ckpt.get("global_step", 0)) print(f"✅ Resumed (new format): global_step={step}, " f"lambda={ckpt.get('lambda', 'n/a')}") return step # legacy: plain state_dict -- сюда попадёт исходный model.pt (lambda=0) missing, unexpected = model.load_state_dict(ckpt, strict=False) assert not unexpected, unexpected assert all(k.endswith("lambda_") for k in missing), missing # [A-2] get_lambda() может не существовать в JiRackTransformer (1b) -- # если так, считаем lambda=0.0, что для исходного model.pt и так верно. if hasattr(model, "get_lambda"): lam = model.get_lambda() else: lam = 0.0 print("ℹ️ У model нет get_lambda() -- считаю lambda=0.0 " "(корректно для непройденного через warmup model.pt).") step = invert_lambda_schedule(lam) print(f"✅ Resumed (legacy format): lambda={lam:.4f} -> " f"reconstructed global_step={step}") return step # Шард-чекпоинты (определяют, какие шарды уже пройдены)... checkpoints = sorted( glob.glob(os.path.join(OUTPUT_DIR, "jirack_1b_data_*.pt")), key=lambda x: int(re.search(r"data_(\d+)", x).group(1)), ) # ...но для СОСТОЯНИЯ МОДЕЛИ грузим самый свежий файл на диске -- это # может быть mid-shard autosave, записанный после последнего шарда. all_ckpts = glob.glob(os.path.join(OUTPUT_DIR, "*.pt")) global_step = 0 if all_ckpts: LATEST_CKPT = max(all_ckpts, key=os.path.getmtime) print(f"📦 Resuming from (newest on disk): {LATEST_CKPT}") global_step = load_any_checkpoint(LATEST_CKPT, model, optimizer) else: print(f"📦 Loading base checkpoint: {BASE_CHECKPOINT}") assert os.path.exists(BASE_CHECKPOINT), ( f"{BASE_CHECKPOINT} не найден -- проверь путь" ) global_step = load_any_checkpoint(BASE_CHECKPOINT, model, optimizer) model.set_lambda(lambda_schedule(global_step)) print(f"🔧 Warmup: {LAMBDA_WARMUP_STEPS} steps | starting at " f"step={global_step}, lambda={model.get_lambda() if hasattr(model, 'get_lambda') else lambda_schedule(global_step):.4f}") # ========================= DATASET ========================= class ShardDataset(Dataset): def __init__(self, data_list): self.data = data_list def __len__(self): return len(self.data) def __getitem__(self, idx): return self.data[idx] def collate_fn(batch): input_ids = pad_sequence( [item["input_ids"] for item in batch], batch_first=True, padding_value=0 ) attention_mask = pad_sequence( [item.get("attention_mask", torch.ones_like(item["input_ids"])) for item in batch], batch_first=True, padding_value=0, ) labels = input_ids.clone() labels[attention_mask == 0] = -100 return {"input_ids": input_ids, "labels": labels} # ========================= CHECKPOINT SAVE ========================= def save_checkpoint(path, model, optimizer, global_step, lam): model_sd = {k: v.detach().to(torch.bfloat16).cpu() for k, v in model.state_dict().items()} ckpt = { "model": model_sd, "optimizer": optimizer.state_dict() if SAVE_OPTIMIZER else None, "global_step": global_step, "lambda": lam, } torch.save(ckpt, path) # ========================= TRAINING ========================= # [A-6] шаблон имён шардов -- поправить под реальный формат данных для 1b all_shards = sorted( glob.glob(f"{DATA_DIR}/sft_data_*.pt"), key=lambda x: int(re.search(r"sft_data_(\d+)", x).group(1)), ) if not all_shards: print(f"⚠️ В {DATA_DIR} не найдено шардов sft_data_*.pt -- " f"проверь путь/формат перед запуском.") last_done = -1 if checkpoints: m = re.search(r"data_(\d+)", checkpoints[-1]) last_done = int(m.group(1)) if m else -1 for shard_path in all_shards: shard_name = os.path.basename(shard_path) shard_num = int(re.search(r"sft_data_(\d+)", shard_name).group(1)) if shard_num <= last_done: print(f"⏭ Skipping already processed: {shard_name}") continue print(f"\n🔥 Starting shard: {shard_name}") raw_shard_data = torch.load(shard_path, map_location="cpu", weights_only=False) val_size = int(len(raw_shard_data) * VAL_RATIO) train_size = len(raw_shard_data) - val_size train_data, val_data = torch.utils.data.random_split( raw_shard_data, [train_size, val_size], generator=torch.Generator().manual_seed(42), ) train_loader = DataLoader( ShardDataset(train_data), batch_size=BATCH_SIZE, shuffle=True, collate_fn=collate_fn, pin_memory=(DEVICE == "cuda"), ) model.train() pbar = tqdm(train_loader, desc=f"Shard {shard_num}", dynamic_ncols=True) optimizer.zero_grad() lambda_value = lambda_schedule(global_step) model.set_lambda(lambda_value) for step, batch in enumerate(pbar): if step % GRAD_ACCUM == 0: lambda_value = lambda_schedule(global_step) model.set_lambda(lambda_value) input_ids = batch["input_ids"].to(DEVICE, non_blocking=(DEVICE == "cuda")) labels = batch["labels"].to(DEVICE, non_blocking=(DEVICE == "cuda")) with torch.amp.autocast(DEVICE, dtype=AUTOCAST_DTYPE): logits = model(input_ids) if isinstance(logits, tuple): logits = logits[0] loss = criterion( logits[..., :-1, :].reshape(-1, config.vocab_size), labels[..., 1:].reshape(-1), ) loss = loss / GRAD_ACCUM if torch.isnan(loss) or torch.isinf(loss): print(f"\n⚠️ NaN/Inf loss at step {global_step} " f"(lambda={lambda_value:.4f}) -- window dropped") optimizer.zero_grad(set_to_none=True) if DEVICE == "cuda": torch.cuda.empty_cache() global_step += 1 continue loss.backward() if (step + 1) % GRAD_ACCUM == 0: torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) optimizer.step() optimizer.zero_grad() if step % 10 == 0: pbar.set_postfix({ "loss": f"{loss.item() * GRAD_ACCUM:.4f}", "lambda": f"{lambda_value:.4f}", "gstep": global_step, }) if (AUTOSAVE_EVERY > 0 and global_step > 0 and global_step % AUTOSAVE_EVERY == 0 and (step + 1) % GRAD_ACCUM == 0): autosave_path = os.path.join(OUTPUT_DIR, "autosave_latest.pt") tmp_path = autosave_path + ".tmp" save_checkpoint(tmp_path, model, optimizer, global_step, lambda_value) os.replace(tmp_path, autosave_path) pbar.write(f"💾 autosave @ gstep={global_step}, " f"lambda={lambda_value:.4f}") global_step += 1 # ==================== Validation ==================== print("🧪 Validating...") model.eval() total_val_loss = 0.0 val_steps = 0 val_loader = DataLoader( ShardDataset(val_data), batch_size=BATCH_SIZE, shuffle=False, collate_fn=collate_fn, pin_memory=(DEVICE == "cuda"), ) with torch.no_grad(): for batch in tqdm(val_loader, desc="Validating", leave=False): input_ids = batch["input_ids"].to(DEVICE, non_blocking=(DEVICE == "cuda")) labels = batch["labels"].to(DEVICE, non_blocking=(DEVICE == "cuda")) with torch.amp.autocast(DEVICE, dtype=AUTOCAST_DTYPE): logits = model(input_ids) if isinstance(logits, tuple): logits = logits[0] v_loss = criterion( logits[..., :-1, :].reshape(-1, config.vocab_size), labels[..., 1:].reshape(-1), ) if not (torch.isnan(v_loss) or torch.isinf(v_loss)): total_val_loss += v_loss.item() val_steps += 1 avg_val_loss = total_val_loss / val_steps if val_steps > 0 else float("inf") print(f"📊 Shard {shard_num} -- Val Loss: {avg_val_loss:.4f} " f"@ lambda={lambda_value:.4f} (gstep={global_step})") # ==================== Save ==================== save_path = os.path.join(OUTPUT_DIR, f"jirack_1b_data_{shard_num}.pt") save_checkpoint(save_path, model, optimizer, global_step, lambda_value) print(f"💾 Saved: {save_path}") del raw_shard_data if DEVICE == "cuda": torch.cuda.empty_cache() gc.collect() print("🏁 Training finished with Lambda Warmup!")