import os import torch import glob import re import sys import types import torch.nn as nn from torch.utils.data import Dataset, DataLoader from torch.optim import AdamW from tqdm import tqdm # Окружение ROCm (Безопасный менеджмент памяти для MI50 32GB) os.environ["HIP_BLAS_LT"] = "0" os.environ["PYTORCH_HIP_ALLOC_CONF"] = "max_split_size_mb:32,garbage_collection_threshold:0.6" # Импортируем архитектуру JiRack 1B from JiRackTernaryPyTorch_1b import TernaryTransformer1B, TernaryConfig from peft import LoraConfig, get_peft_model, TaskType # --- CONFIG --- DATA_DIR = "/mnt/nfs_clientshare/GammaCorpus-Fact-QA-JSON" OUTPUT_DIR = "/mnt/nfs_clientshare/checkpoints_GammaCopusQA_LoRA" BATCH_SIZE = 1 GRAD_ACCUM = 4 LR = 2e-4 DEVICE = "cuda" MAX_SEQ_LEN = 2048 # Стабильный предел для gfx906 без поддержки FlashAttention os.makedirs(OUTPUT_DIR, exist_ok=True) def natural_key(string_): return [int(s) if s.isdigit() else s for s in re.split(r'(\d+)', string_)] # --- CHECKPOINT RESUME --- checkpoints = sorted(glob.glob(f"{OUTPUT_DIR}/jirack_lora_*.pt"), key=natural_key) if checkpoints: LATEST_CHECKPOINT = checkpoints[-1] match = re.findall(r'data_(\d+)', os.path.basename(LATEST_CHECKPOINT)) last_shard_idx = int(match[-1]) if match else -1 print(f"🔄 Возобновляем обучение LoRA с шарда: {last_shard_idx}") else: LATEST_CHECKPOINT = "jiarck_pro_model.pt" last_shard_idx = -1 print(f"🚀 Старт. Замораживаем базовую модель: {LATEST_CHECKPOINT}") # --- DATASET --- class ShardDataset(Dataset): def __init__(self, shard_path): raw_data = torch.load(shard_path, map_location='cpu', weights_only=False) if isinstance(raw_data, dict): self.data = raw_data.get("input_ids", next(iter(raw_data.values()))) else: self.data = raw_data def __len__(self): return self.data.size(0) def __getitem__(self, idx): return self.data[idx] # --- FIND SHARDS --- all_shards = sorted(glob.glob(f"{DATA_DIR}/jirack_sft_gemma_facts_qa_data_*.pt"), key=natural_key) shards_to_train = [s for s in all_shards if int(re.findall(r'data_(\d+)', os.path.basename(s))[0]) > last_shard_idx] if not shards_to_train: print(f"✅ Обучать нечего. Последний индекс: {last_shard_idx}") sys.exit(0) # --- MODEL SETUP --- config = TernaryConfig() def config_get(self, key, default=None): return getattr(self, key, default) config.get = types.MethodType(config_get, config) config.tie_word_embeddings = False model = TernaryTransformer1B(config) model.config = config model.config.model_type = "custom_jirack" # Адаптация forward original_forward = model.forward def adaptive_forward(input_ids, *args, **kwargs): return original_forward(input_ids) model.forward = adaptive_forward # Включаем Gradient Checkpointing для базовых блоков for block in model.blocks: block.gradient_checkpointing = True print(f"📦 Загрузка базовой модели: {LATEST_CHECKPOINT}") sd = torch.load(LATEST_CHECKPOINT, map_location='cpu', weights_only=False) if isinstance(sd, dict) and "model_state_dict" in sd: model.load_state_dict(sd["model_state_dict"], strict=False) else: model.load_state_dict(sd, strict=False) # Возвращаем мощную LoRA конфигурацию (ранг 16 на все ключевые проекции) peft_config = LoraConfig( task_type=TaskType.CAUSAL_LM, r=16, lora_alpha=32, lora_dropout=0.05, target_modules=[ "q_proj", "k_proj", "v_proj", "out_proj", "ffn_w1", "ffn_w2", "ffn_w3" ], bias="none" ) model.prepare_inputs_for_generation = lambda *args, **kwargs: None print("🔗 Внедряем LoRA адаптеры...") model = get_peft_model(model, peft_config) if hasattr(model, "gradient_checkpointing_enable"): model.gradient_checkpointing_enable() model.print_trainable_parameters() if checkpoints: print(f"📥 Подгружаем веса LoRA из: {LATEST_CHECKPOINT}") if isinstance(sd, dict) and "lora_state_dict" in sd: model.load_state_dict(sd["lora_state_dict"], strict=False) model.to(DEVICE) model.train() trainable_params = [p for p in model.parameters() if p.requires_grad] optimizer = AdamW(trainable_params, lr=LR) criterion = nn.CrossEntropyLoss(ignore_index=-100) scaler = torch.amp.GradScaler('cuda') # --- TRAINING LOOP --- for shard_path in shards_to_train: shard_name = os.path.basename(shard_path) dataset = ShardDataset(shard_path) dataloader = DataLoader(dataset, batch_size=BATCH_SIZE, shuffle=True, num_workers=0, pin_memory=True) print(f"\n🔥 Начинаем шард: {shard_name} | Объем: {len(dataset):,}") pbar = tqdm(dataloader, desc=f"Shard {shard_name}") optimizer.zero_grad() current_loss = 0.0 torch.cuda.empty_cache() for step, batch in enumerate(pbar): if isinstance(batch, dict): input_ids = batch.get("input_ids", batch).to(DEVICE) else: input_ids = batch.to(DEVICE) if input_ids.dim() == 3: input_ids = input_ids.view(BATCH_SIZE, -1) if input_ids.size(1) > MAX_SEQ_LEN: input_ids = input_ids[:, :MAX_SEQ_LEN] labels = input_ids.clone() try: # Используем стандартный нативный фоллбэк без принудительных ядер with torch.amp.autocast('cuda', dtype=torch.float16): logits, _ = model(input_ids) loss = criterion( logits[..., :-1, :].contiguous().view(-1, config.vocab_size), labels[..., 1:].contiguous().view(-1) ) loss = loss / GRAD_ACCUM scaler.scale(loss).backward() if (step + 1) % GRAD_ACCUM == 0: scaler.unscale_(optimizer) torch.nn.utils.clip_grad_norm_(trainable_params, 1.0) scaler.step(optimizer) scaler.update() optimizer.zero_grad() current_loss = loss.item() * GRAD_ACCUM pbar.set_postfix({"loss": f"{current_loss:.4f}"}) if (step + 1) % 10 == 0 or (step + 1) == len(dataloader): print(f"\n[INFO] {shard_name} | Шаг {step+1:6d} | Loss: {current_loss:.4f}") except RuntimeError as e: if "out of memory" in str(e).lower(): print("⚠️ Обнаружен OOM - пропускаем батч") optimizer.zero_grad() torch.cuda.empty_cache() continue else: raise e save_path = os.path.join(OUTPUT_DIR, f"jirack_lora_{shard_name}") from peft import get_peft_model_state_dict lora_sd = get_peft_model_state_dict(model) torch.save({ "lora_state_dict": lora_sd, "shard_idx": int(re.findall(r'data_(\d+)', shard_name)[0]) }, save_path) print(f"💾 Успешно сохранены веса LoRA: {save_path}\n") print("🏁 Обучение восстановлено на стабильных 2К параметрах!")