train_jirack_1b_lambda_warmup.py
18.2 KB · 414 lines · python Raw
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