PerturbReason/PerturbReason_dataset_code
050
1"""2rl_v1/grpo_trainer_distributed_v2.py3=====================================4Multi-node distributed GRPO training for PerturbReason v4.0.5 6Changes from v1 (grpo_trainer_distributed.py):7 - Loads pre-formatted JSONL from dataset_0326_noisy_hidden/{noisy,hidden}/3way_train/8 - Supports both chemical + genetic perturbation types9 - Supports both noisy (with retrieved knowledge) and hidden (no knowledge) tasks10 - Balanced oversampling across (task_type × pert_type × label) like SFT11 - Reward mode accuracy_triplet (model outputs <triplet> blocks)12 - Uses v4.0 SFT adapter as base13 14Architecture (per epoch):15 Phase 1 (prepare): Merge model, prepare data [single process]16 Phase 2 (generate): Parallel vLLM across nodes [1 proc/node, TP=2]17 Phase 3 (rewards): Compute rewards & advantages [single process]18 Phase 4 (train): DDP training with batched GRPO [1 proc/GPU]19 Phase 5 (merge): Merge LoRA for next epoch [single process]20 21Usage:22 See sh_grpo_distributed_v2.sh for the SLURM launch script.23"""24 25from __future__ import annotations26 27import gc28import json29import os30import random31import sys32import time33from collections import Counter, defaultdict34from pathlib import Path35from typing import Any, Dict, List, Optional36 37import numpy as np38import torch39import torch.distributed as dist40from torch.utils.data import DataLoader41 42import deepspeed43 44from peft import LoraConfig, PeftModel, get_peft_model45from transformers import AutoModelForCausalLM, AutoTokenizer, get_scheduler46 47# ── Ensure project root is importable ──48_PROJECT_ROOT = str(Path(__file__).resolve().parent.parent)49if _PROJECT_ROOT not in sys.path:50 sys.path.insert(0, _PROJECT_ROOT)51 52_RL_ROOT = str(Path(__file__).resolve().parent)53if _RL_ROOT not in sys.path:54 sys.path.insert(0, _RL_ROOT)55 56_SRC_ROOT = str(Path(__file__).resolve().parent / "src_0412")57if _SRC_ROOT not in sys.path:58 sys.path.insert(0, _SRC_ROOT)59 60_SRC_ROOT_OLD = str(Path(__file__).resolve().parent / "src_0302")61if _SRC_ROOT_OLD not in sys.path:62 sys.path.insert(0, _SRC_ROOT_OLD)63 64 65def make_deterministic(seed: int = 42):66 os.environ["PYTHONHASHSEED"] = str(seed)67 random.seed(seed)68 np.random.seed(seed)69 torch.manual_seed(seed)70 if torch.cuda.is_available():71 torch.cuda.manual_seed_all(seed)72 torch.backends.cudnn.deterministic = True73 torch.backends.cudnn.benchmark = False74 75 76def parse_args():77 import argparse78 p = argparse.ArgumentParser(description="Distributed GRPO training v2 (v4.0 model)")79 80 # Phase control81 p.add_argument("--phase", type=str, required=True,82 choices=["prepare", "generate", "rewards", "train", "merge"])83 p.add_argument("--epoch", type=int, default=1, help="Current epoch (1-based)")84 85 # Model86 p.add_argument("--model_path", type=str, required=True)87 p.add_argument("--base_adapter_paths", type=str, nargs="+", default=None)88 p.add_argument("--output_dir", type=str, required=True,89 help="Run output directory (created by shell script)")90 91 # Data — v2: noisy_dir and hidden_dir instead of reasoning_dir92 p.add_argument("--noisy_dir", type=str, default=None,93 help="Directory with noisy-context (retrieved knowledge) JSONL files")94 p.add_argument("--hidden_dir", type=str, default=None,95 help="Directory with hidden-context (no knowledge) JSONL files")96 p.add_argument("--data_jsonl", type=str, default=None,97 help="Direct path to a pre-merged JSONL file (overrides noisy_dir/hidden_dir)")98 p.add_argument("--task_origin", type=str, default="3way",99 choices=["dir", "de", "3way"])100 p.add_argument("--pert_type", type=str, default=None,101 choices=[None, "chemical", "genetic"],102 help="Filter by pert type (default: use both)")103 p.add_argument("--max_train_samples", type=int, default=None)104 105 # GRPO hyperparams106 p.add_argument("--num_generations", type=int, default=8)107 p.add_argument("--max_completion_length", type=int, default=768)108 p.add_argument("--max_prompt_length", type=int, default=2048)109 p.add_argument("--temperature", type=float, default=0.85)110 p.add_argument("--top_p", type=float, default=0.9)111 p.add_argument("--beta", type=float, default=0.02)112 113 # Training114 p.add_argument("--per_device_train_batch_size", type=int, default=4,115 help="Training batch size per GPU")116 p.add_argument("--gradient_accumulation_steps", type=int, default=2)117 p.add_argument("--learning_rate", type=float, default=5e-6)118 p.add_argument("--warmup_ratio", type=float, default=0.05)119 p.add_argument("--lr_scheduler_type", type=str, default="cosine")120 p.add_argument("--max_grad_norm", type=float, default=0.5)121 p.add_argument("--ref_batch_size", type=int, default=8,122 help="Batch size for reference log-prob computation")123 124 # LoRA125 p.add_argument("--lora_r", type=int, default=16)126 p.add_argument("--lora_alpha", type=int, default=32)127 p.add_argument("--lora_dropout", type=float, default=0.05)128 129 # Reward130 p.add_argument("--reward_mode", type=str, default="accuracy_triplet",131 choices=["composite", "multi", "accuracy_only", "accuracy_triplet",132 "accuracy_triplet_chain"])133 p.add_argument("--answer_weight", type=float, default=0.6,134 help="Weight for answer_correct in accuracy_triplet mode (default 0.6, triplet gets 1-w)")135 p.add_argument("--chain_weight", type=float, default=0.25,136 help="Weight for chain-answer consistency in accuracy_triplet_chain mode")137 138 # Exploration fixes (v4)139 p.add_argument("--gt_inject", action="store_true", default=False,140 help="Replace last rollout with GT completion to guarantee exploration")141 p.add_argument("--zero_std_baseline", type=float, default=0.0,142 help="When >0, assign ±this advantage to zero-std groups instead of skipping them")143 144 # vLLM145 p.add_argument("--vllm_gpu_utilization", type=float, default=0.85)146 p.add_argument("--vllm_tensor_parallel", type=int, default=2)147 p.add_argument("--vllm_batch_size", type=int, default=5000)148 149 # Generation sharding150 p.add_argument("--num_gen_nodes", type=int, default=1,151 help="Total number of nodes for generation")152 p.add_argument("--gen_node_rank", type=int, default=0,153 help="This node's rank for generation (set by SLURM)")154 155 # Misc156 p.add_argument("--bf16", action="store_true", default=True)157 p.add_argument("--gradient_checkpointing", action="store_true", default=True)158 p.add_argument("--seed", type=int, default=42)159 p.add_argument("--logging_steps", type=int, default=10)160 p.add_argument("--save_steps", type=int, default=200)161 p.add_argument("--save_total_limit", type=int, default=5)162 p.add_argument("--debug", action="store_true")163 p.add_argument("--sanity", action="store_true")164 p.add_argument("--no_think", action="store_true")165 p.add_argument("--num_train_epochs", type=int, default=3,166 help="Total epochs (used by merge phase)")167 168 args = p.parse_args()169 170 # Resolve full adapter chain (same as SFT training)171 if getattr(args, "base_adapter_paths", None):172 try:173 from model_utils_light import add_all_base_adapter_paths174 args = add_all_base_adapter_paths(args)175 print(f"Resolved adapter chain ({len(args.base_adapter_paths)} adapters):")176 for ap in args.base_adapter_paths:177 print(f" - {ap}")178 except ImportError:179 print("WARNING: model_utils_light not found, using adapter paths as-is")180 181 return args182 183 184# ══════════════════════════════════════════════════════════════185# Data loading for v2 (pre-formatted JSONL)186# ══════════════════════════════════════════════════════════════187 188def load_jsonl(path: str | Path) -> List[Dict[str, Any]]:189 """Load a JSONL file into a list of dicts."""190 records = []191 with open(path, "r", encoding="utf-8") as f:192 for line in f:193 line = line.strip()194 if line:195 records.append(json.loads(line))196 return records197 198 199def load_preformatted_data(200 noisy_dir: str | Path | None,201 hidden_dir: str | Path | None,202 pert_type: str | None = None,203 task_origin: str = "3way",204 max_samples: int | None = None,205 tokenizer=None,206) -> List[Dict[str, Any]]:207 """208 Load pre-formatted JSONL training data from noisy and hidden directories.209 210 Each JSONL file has records with: id, label, pert_type, prompt, response, output.211 For RL, we use prompt as input and response as ground truth for reward.212 213 Parameters214 ----------215 noisy_dir : Directory with noisy-context (retrieved knowledge) train/valid JSONLs216 hidden_dir : Directory with hidden-context (no knowledge) train/valid JSONLs217 pert_type : Filter by perturbation type ("chemical", "genetic", or None for both)218 task_origin : Task type for pattern matching ("3way", "dir", "de")219 max_samples : Limit total samples220 tokenizer : If provided, format prompts as chat messages221 """222 samples = []223 224 def _load_from_dir(data_dir: Path, data_type_label: str):225 """Load train JSONL files from a directory."""226 if data_dir is None or not Path(data_dir).exists():227 return228 229 data_dir = Path(data_dir)230 # Find train files matching pert_type filter231 for jsonl_path in sorted(data_dir.glob("*train*task_3way*.jsonl")):232 fname = jsonl_path.name233 # Skip cell-OOD files — these are evaluation splits and must not be in training234 if "cell_ood" in fname:235 print(f" Skipping cell-OOD file (reserved for evaluation): {fname}")236 continue237 # Apply pert_type filter238 if pert_type == "chemical" and "chemical" not in fname:239 continue240 if pert_type == "genetic" and "genetic" not in fname:241 continue242 243 # Determine pert_type from filename244 if "chemical" in fname:245 file_pert_type = "chemical"246 elif "genetic" in fname:247 file_pert_type = "genetic"248 else:249 file_pert_type = "unknown"250 251 records = load_jsonl(jsonl_path)252 print(f" Loaded {len(records)} records from {jsonl_path.name} "253 f"(data_type={data_type_label}, pert_type={file_pert_type})")254 255 for i, d in enumerate(records):256 prompt_text = d.get("prompt", "")257 response_text = d.get("response", "")258 label = d.get("label", "")259 260 if not prompt_text:261 continue262 263 sample = {264 "id": d.get("id", f"{data_type_label}_{file_pert_type}_{i}"),265 "prompt_text": prompt_text,266 "ground_truth": response_text,267 "label": label,268 "pert_type": d.get("pert_type", file_pert_type),269 "data_type": data_type_label, # "noisy" or "hidden"270 }271 272 # Format prompt as chat messages for GRPOTrainer273 if tokenizer is not None:274 sample["prompt"] = [{"role": "user", "content": prompt_text}]275 else:276 sample["prompt"] = prompt_text277 278 samples.append(sample)279 280 _load_from_dir(noisy_dir, "noisy")281 _load_from_dir(hidden_dir, "hidden")282 283 if max_samples and len(samples) > max_samples:284 random.shuffle(samples)285 samples = samples[:max_samples]286 287 return samples288 289 290def balance_by_pert_type_and_label(291 samples: List[Dict[str, Any]],292 pert_type_ratios: Dict[str, float],293 label_ratios: Dict[str, float],294 seed: int = 42,295) -> List[Dict[str, Any]]:296 """297 Two-stage balanced oversampling (matching SFT's pert_label_per_data_pert):298 1. Balance pert_type within each data_type group299 2. Balance labels within each (data_type, pert_type) group300 """301 rng = random.Random(seed)302 303 def _oversample(pool: List, target_n: int) -> List:304 if len(pool) >= target_n:305 return rng.sample(pool, target_n)306 else:307 extra = [rng.choice(pool) for _ in range(target_n - len(pool))]308 return pool + extra309 310 # Group by data_type311 dt_groups: Dict[str, List] = defaultdict(list)312 for s in samples:313 dt_groups[s.get("data_type", "unknown")].append(s)314 315 # Stage 1: Balance pert_type within each data_type316 pert_balanced = []317 for dt, dt_samples in dt_groups.items():318 dt_total = len(dt_samples)319 pt_groups: Dict[str, List] = defaultdict(list)320 for s in dt_samples:321 pt_groups[s.get("pert_type", "unknown")].append(s)322 323 for pt, ratio in pert_type_ratios.items():324 if pt not in pt_groups:325 continue326 target_n = int(dt_total * ratio)327 pert_balanced.extend(_oversample(pt_groups[pt], target_n))328 329 # Stage 2: Balance labels within each (data_type, pert_type) group330 final = []331 dtpt_groups: Dict[tuple, List] = defaultdict(list)332 for s in pert_balanced:333 key = (s.get("data_type", "unknown"), s.get("pert_type", "unknown"))334 dtpt_groups[key].append(s)335 336 for (dt, pt), group_samples in dtpt_groups.items():337 group_total = len(group_samples)338 label_groups: Dict[str, List] = defaultdict(list)339 for s in group_samples:340 label_groups[s.get("label", "unknown")].append(s)341 342 for label, ratio in label_ratios.items():343 if label not in label_groups:344 continue345 target_n = int(group_total * ratio)346 final.extend(_oversample(label_groups[label], target_n))347 348 rng.shuffle(final)349 return final350 351 352def build_grpo_dataset_v2(353 samples: List[Dict[str, Any]],354 shuffle: bool = True,355 seed: int = 42,356 task: str = "3way",357):358 """Build a HuggingFace Dataset for GRPO from pre-formatted samples."""359 from datasets import Dataset360 361 data = []362 for s in samples:363 data.append({364 "prompt": s["prompt"],365 "ground_truth": s.get("ground_truth", ""),366 "label": s.get("label", ""),367 "prompt_text": s.get("prompt_text", ""),368 "pert_type": s.get("pert_type", ""),369 "data_type": s.get("data_type", ""),370 "task": task,371 })372 373 ds = Dataset.from_list(data)374 if shuffle:375 ds = ds.shuffle(seed=seed)376 return ds377 378 379def print_dataset_stats(samples: List[Dict[str, Any]], name: str = "Dataset"):380 """Print dataset statistics."""381 print(f"\n{name} statistics:")382 print(f" Total samples: {len(samples)}")383 384 # Label distribution385 labels = Counter(s.get("label", "?") for s in samples)386 print(f" Labels: {dict(labels)}")387 388 # Pert type distribution389 pert_types = Counter(s.get("pert_type", "?") for s in samples)390 print(f" Pert types: {dict(pert_types)}")391 392 # Data type distribution (noisy vs hidden)393 data_types = Counter(s.get("data_type", "?") for s in samples)394 print(f" Data types: {dict(data_types)}")395 396 # Cross-tabulation: (data_type, pert_type, label)397 cross = Counter(398 (s.get("data_type", "?"), s.get("pert_type", "?"), s.get("label", "?"))399 for s in samples400 )401 print(f" Cross (data_type, pert_type, label): {dict(cross)}")402 403 404# ══════════════════════════════════════════════════════════════405# Batched GRPO utilities (reused from v1)406# ══════════════════════════════════════════════════════════════407 408class GRPOPairDataset(torch.utils.data.Dataset):409 """Pre-tokenized GRPO training pairs with optional reference log-probs."""410 411 def __init__(self, pairs: List[Dict], tokenizer, max_seq_len: int):412 self.items = []413 for pair in pairs:414 prompt_ids = tokenizer.encode(pair["prompt_text"])415 comp_ids = tokenizer.encode(pair["completion_text"], add_special_tokens=False)416 full_ids = (prompt_ids + comp_ids)[:max_seq_len]417 prompt_len = min(len(prompt_ids), max_seq_len)418 comp_len = len(full_ids) - prompt_len419 420 if comp_len <= 0:421 continue422 423 self.items.append({424 "input_ids": full_ids,425 "prompt_len": prompt_len,426 "comp_len": comp_len,427 "advantage": pair["advantage"],428 "ref_logprobs": None,429 })430 431 def __len__(self):432 return len(self.items)433 434 def __getitem__(self, idx):435 return self.items[idx]436 437 438def collate_grpo(batch, pad_token_id):439 """Right-pad batch to max sequence length."""440 max_len = max(len(b["input_ids"]) for b in batch)441 442 input_ids = []443 attention_mask = []444 prompt_lens = []445 comp_lens = []446 advantages = []447 ref_logprobs_list = []448 449 for b in batch:450 pad_len = max_len - len(b["input_ids"])451 input_ids.append(b["input_ids"] + [pad_token_id] * pad_len)452 attention_mask.append([1] * len(b["input_ids"]) + [0] * pad_len)453 prompt_lens.append(b["prompt_len"])454 comp_lens.append(b["comp_len"])455 advantages.append(b["advantage"])456 ref_logprobs_list.append(b.get("ref_logprobs"))457 458 return {459 "input_ids": torch.tensor(input_ids, dtype=torch.long),460 "attention_mask": torch.tensor(attention_mask, dtype=torch.long),461 "prompt_lens": prompt_lens,462 "comp_lens": comp_lens,463 "advantages": torch.tensor(advantages, dtype=torch.float32),464 "ref_logprobs": ref_logprobs_list,465 }466 467 468def extract_completion_logprobs(logits, input_ids, prompt_lens, comp_lens):469 """Extract per-token log-probs for completion tokens from batched logits."""470 log_probs = torch.log_softmax(logits, dim=-1)471 per_sample = []472 for i in range(logits.shape[0]):473 pl, cl = prompt_lens[i], comp_lens[i]474 sample_lp = log_probs[i, pl - 1 : pl + cl - 1, :]475 targets = input_ids[i, pl : pl + cl]476 token_lp = sample_lp.gather(-1, targets.unsqueeze(-1)).squeeze(-1)477 per_sample.append(token_lp)478 return per_sample479 480 481def compute_grpo_batch_loss(policy_lps, ref_lps, advantages, comp_lens, beta,482 clip_low=0.8, clip_high=1.2):483 """Compute GRPO clipped surrogate loss + KL penalty for a batch."""484 device = policy_lps[0].device485 total_loss = torch.tensor(0.0, device=device)486 total_adv_abs = 0.0487 n_valid = 0488 489 for i in range(len(policy_lps)):490 cl = comp_lens[i]491 pol_lp = policy_lps[i][:cl]492 ref_lp = ref_lps[i][:cl].to(device).detach()493 adv = advantages[i].to(device)494 495 log_ratio = pol_lp - ref_lp496 ratio = torch.exp(log_ratio)497 clipped = torch.clamp(ratio, clip_low, clip_high)498 surr = -torch.min(adv * ratio, adv * clipped).mean()499 kl = (ref_lp - pol_lp).mean()500 501 total_loss = total_loss + surr + beta * kl502 total_adv_abs += abs(adv.item())503 n_valid += 1504 505 if n_valid > 0:506 total_loss = total_loss / n_valid507 508 return total_loss, total_adv_abs / max(n_valid, 1)509 510 511# ══════════════════════════════════════════════════════════════512# Phase: PREPARE513# ══════════════════════════════════════════════════════════════514 515def phase_prepare(args):516 """Merge model + adapters, prepare and save training data."""517 output_dir = Path(args.output_dir)518 output_dir.mkdir(parents=True, exist_ok=True)519 520 with open(output_dir / "_args.json", "w") as f:521 json.dump(vars(args), f, indent=2, default=str)522 523 tokenizer = AutoTokenizer.from_pretrained(args.model_path, trust_remote_code=True)524 tokenizer.pad_token = tokenizer.eos_token525 tokenizer.padding_side = "left"526 527 # ── Load training data ──528 if args.data_jsonl:529 from rl_v1.data_loader import load_sft_jsonl_for_rl530 samples = load_sft_jsonl_for_rl(531 args.data_jsonl, tokenizer=tokenizer,532 max_samples=args.max_train_samples,533 )534 elif args.noisy_dir or args.hidden_dir:535 samples = load_preformatted_data(536 noisy_dir=args.noisy_dir,537 hidden_dir=args.hidden_dir,538 pert_type=args.pert_type,539 task_origin=args.task_origin,540 max_samples=args.max_train_samples,541 tokenizer=tokenizer,542 )543 else:544 raise ValueError("Must provide --data_jsonl or at least one of --noisy_dir / --hidden_dir")545 546 if args.debug:547 samples = samples[:min(50, len(samples))]548 print(f"Debug mode: trimmed to {len(samples)} samples")549 if args.sanity:550 samples = samples[:min(500, len(samples))]551 print(f"Sanity mode: trimmed to {len(samples)} samples")552 553 # Disable Qwen3 thinking mode554 if args.no_think:555 for s in samples:556 prompt = s.get("prompt")557 if isinstance(prompt, list):558 for msg in prompt:559 if msg.get("role") == "user":560 msg["content"] = msg["content"].rstrip() + " /no_think"561 break562 elif isinstance(prompt, str):563 s["prompt"] = prompt.rstrip() + " /no_think"564 print(f"no_think mode: appended /no_think to all {len(samples)} user prompts")565 566 # Balanced oversampling: pert_type → label (like SFT pert_label_per_data_pert)567 print_dataset_stats(samples, "Before oversampling")568 569 if args.task_origin == "3way":570 label_ratios = {"up": 0.33, "down": 0.33, "unchanged": 0.34}571 elif args.task_origin == "dir":572 label_ratios = {"up": 0.5, "down": 0.5}573 elif args.task_origin == "de":574 label_ratios = {"changed": 0.5, "unchanged": 0.5}575 else:576 label_ratios = {"up": 0.33, "down": 0.33, "unchanged": 0.34}577 578 if args.pert_type is None:579 pert_type_ratios = {"chemical": 0.5, "genetic": 0.5}580 else:581 pert_type_ratios = {args.pert_type: 1.0}582 583 samples = balance_by_pert_type_and_label(584 samples,585 pert_type_ratios=pert_type_ratios,586 label_ratios=label_ratios,587 seed=args.seed,588 )589 print_dataset_stats(samples, "After oversampling")590 591 dataset = build_grpo_dataset_v2(samples, shuffle=True, seed=args.seed, task=args.task_origin)592 593 # ── Prepare prompt texts (apply chat template for vLLM) ──594 prompt_texts = []595 dataset_records = []596 for i in range(len(dataset)):597 p = dataset[i]["prompt"]598 if isinstance(p, list):599 text = tokenizer.apply_chat_template(600 p, tokenize=False, add_generation_prompt=True601 )602 else:603 text = p604 prompt_texts.append(text)605 dataset_records.append({606 "prompt": dataset[i]["prompt"],607 "ground_truth": dataset[i].get("ground_truth", ""),608 "label": dataset[i].get("label", ""),609 "task": dataset[i].get("task", args.task_origin),610 })611 612 prepared_path = output_dir / "prepared_data.json"613 with open(prepared_path, "w") as f:614 json.dump({615 "prompt_texts": prompt_texts,616 "dataset_records": dataset_records,617 "num_samples": len(prompt_texts),618 }, f, ensure_ascii=False)619 print(f"Saved {len(prompt_texts)} prepared prompts to {prepared_path}")620 621 # ── Merge model ──622 epoch = args.epoch623 if epoch == 1:624 print("Merging base model + adapters on CPU ...")625 model = AutoModelForCausalLM.from_pretrained(626 args.model_path, torch_dtype=torch.bfloat16,627 trust_remote_code=True, device_map="cpu", low_cpu_mem_usage=True,628 )629 if args.base_adapter_paths:630 for ap in args.base_adapter_paths:631 print(f" Merging adapter: {ap}")632 model = PeftModel.from_pretrained(model, ap, device_map="cpu")633 model = model.merge_and_unload()634 635 merged_dir = str(output_dir / "_merged_model")636 print(f"Saving merged model to {merged_dir} ...")637 model.save_pretrained(merged_dir)638 tokenizer.save_pretrained(merged_dir)639 del model640 gc.collect()641 print(f"Merged model saved.")642 else:643 prev = output_dir / f"_merged_model_epoch{epoch - 1}"644 if not prev.exists():645 raise FileNotFoundError(f"Expected merged model not found: {prev}")646 print(f"Using existing merged model: {prev}")647 648 print("Phase PREPARE complete.")649 650 651# ══════════════════════════════════════════════════════════════652# Phase: GENERATE653# ══════════════════════════════════════════════════════════════654 655def phase_generate(args):656 """Generate completions for this node's shard of prompts using vLLM."""657 output_dir = Path(args.output_dir)658 epoch = args.epoch659 node_rank = args.gen_node_rank660 num_nodes = args.num_gen_nodes661 662 # Load prepared prompt texts663 with open(output_dir / "prepared_data.json") as f:664 prepared = json.load(f)665 all_prompts = prepared["prompt_texts"]666 667 # Contiguous sharding across nodes668 n = len(all_prompts)669 shard_size = (n + num_nodes - 1) // num_nodes670 start = node_rank * shard_size671 end = min(start + shard_size, n)672 my_prompts = all_prompts[start:end]673 674 print(f"Node {node_rank}/{num_nodes}: generating for prompts [{start}:{end}) "675 f"({len(my_prompts)} prompts × {args.num_generations} completions)")676 677 # Determine merged model directory678 if epoch == 1:679 merged_dir = str(output_dir / "_merged_model")680 else:681 merged_dir = str(output_dir / f"_merged_model_epoch{epoch - 1}")682 683 os.environ["VLLM_WORKER_MULTIPROC_METHOD"] = "spawn"684 from vllm import LLM, SamplingParams685 686 max_model_len = args.max_prompt_length + args.max_completion_length687 688 def _make_llm(enforce_eager: bool = False) -> LLM:689 eager_str = " (enforce_eager)" if enforce_eager else ""690 print(f"Initializing vLLM{eager_str} (TP={args.vllm_tensor_parallel}, "691 f"gpu_util={args.vllm_gpu_utilization}, max_len={max_model_len}) ...")692 return LLM(693 model=merged_dir,694 tensor_parallel_size=args.vllm_tensor_parallel,695 trust_remote_code=True,696 dtype="bfloat16",697 max_model_len=max_model_len,698 gpu_memory_utilization=args.vllm_gpu_utilization,699 enforce_eager=enforce_eager,700 )701 702 sampling_params = SamplingParams(703 temperature=args.temperature,704 top_p=args.top_p,705 max_tokens=args.max_completion_length,706 n=args.num_generations,707 )708 709 # Per-batch checkpoint so a mid-run crash doesn't lose all work.710 checkpoint_path = output_dir / f"gen_checkpoint_{node_rank}_epoch{epoch}.json"711 completions: List[List[str]] = []712 resume_batch = 0713 if checkpoint_path.exists():714 try:715 with open(checkpoint_path) as f:716 ckpt = json.load(f)717 completions = ckpt["completions"]718 resume_batch = ckpt["next_batch"]719 print(f"Node {node_rank}: resuming from checkpoint (batches done: {resume_batch}, "720 f"completions so far: {sum(len(c) for c in completions)})")721 except Exception:722 print(f"Node {node_rank}: checkpoint file corrupt, starting fresh.")723 completions, resume_batch = [], 0724 725 # Initialize vLLM726 llm = _make_llm(enforce_eager=False)727 728 batches = list(range(0, len(my_prompts), args.vllm_batch_size))729 for b_idx, bs in enumerate(batches):730 if b_idx < resume_batch:731 continue732 be = min(bs + args.vllm_batch_size, len(my_prompts))733 print(f" Node {node_rank}: batch [{bs}:{be}) ({be - bs} prompts) ...")734 try:735 outputs = llm.generate(my_prompts[bs:be], sampling_params)736 except Exception as exc:737 print(f" Node {node_rank}: vLLM error on batch {b_idx} ({exc!r}). "738 f"Reinitializing with enforce_eager=True and retrying ...")739 try:740 del llm741 except Exception:742 pass743 gc.collect()744 torch.cuda.empty_cache()745 time.sleep(3)746 llm = _make_llm(enforce_eager=True)747 outputs = llm.generate(my_prompts[bs:be], sampling_params)748 749 completions.extend([[o.text for o in out.outputs] for out in outputs])750 total_done = sum(len(c) for c in completions)751 print(f" Batch done. Total completions: {total_done}")752 753 # Save checkpoint after every batch754 with open(checkpoint_path, "w") as f:755 json.dump({"completions": completions, "next_batch": b_idx + 1}, f, ensure_ascii=False)756 757 del llm758 gc.collect()759 torch.cuda.empty_cache()760 761 # Save final shard to shared filesystem762 shard_path = output_dir / f"gen_shard_{node_rank}_epoch{epoch}.json"763 shard_data = {764 "node_rank": node_rank,765 "start_idx": start,766 "end_idx": end,767 "completions": completions,768 }769 with open(shard_path, "w") as f:770 json.dump(shard_data, f, ensure_ascii=False)771 print(f"Node {node_rank}: saved {len(completions)} prompt completions to {shard_path}")772 773 # Clean up checkpoint774 if checkpoint_path.exists():775 checkpoint_path.unlink()776 777 778# ══════════════════════════════════════════════════════════════779# Phase: REWARDS780# ══════════════════════════════════════════════════════════════781 782def phase_rewards(args):783 """Load all generation shards, compute rewards/advantages, save training pairs."""784 output_dir = Path(args.output_dir)785 epoch = args.epoch786 787 with open(output_dir / "prepared_data.json") as f:788 prepared = json.load(f)789 prompt_texts = prepared["prompt_texts"]790 records = prepared["dataset_records"]791 n = len(prompt_texts)792 793 # Reassemble completions from all node shards794 all_completions: List[Optional[List[str]]] = [None] * n795 shard_files = sorted(output_dir.glob(f"gen_shard_*_epoch{epoch}.json"))796 print(f"Loading {len(shard_files)} generation shards ...")797 798 for sf in shard_files:799 with open(sf) as f:800 shard = json.load(f)801 si = shard["start_idx"]802 for i, comps in enumerate(shard["completions"]):803 all_completions[si + i] = comps804 805 missing = [i for i, c in enumerate(all_completions) if c is None]806 if missing:807 print(f"WARNING: {len(missing)} / {n} prompts have no completions. "808 f"These prompts will be skipped.")809 if len(missing) == n:810 raise RuntimeError("No completions found — all shards missing.")811 812 # Compute rewards813 from rl_v1.rewards import get_reward_functions814 reward_funcs = get_reward_functions(815 mode=args.reward_mode,816 answer_weight=getattr(args, 'answer_weight', 0.6),817 chain_weight=getattr(args, 'chain_weight', 0.25),818 )819 820 gt_inject = getattr(args, 'gt_inject', False)821 zero_std_baseline = getattr(args, 'zero_std_baseline', 0.0)822 823 if gt_inject:824 print(f" GT injection: ON — replacing last rollout with GT completion")825 if zero_std_baseline > 0:826 print(f" Zero-std baseline: ±{zero_std_baseline} (instead of skipping)")827 828 pairs = []829 total_reward, n_samples = 0.0, 0830 n_zero_std, n_all_correct, n_all_wrong, n_mixed = 0, 0, 0, 0831 n_gt_injected = 0832 n_zero_std_rescued = 0833 comp_lengths = []834 sample_log = []835 836 for i in range(n):837 comps_i = all_completions[i]838 if comps_i is None:839 continue840 841 rec = records[i]842 gt = rec.get("ground_truth", "")843 task = rec.get("task", args.task_origin)844 prompt = rec.get("prompt", "")845 846 # ── GT Injection: replace last rollout with GT completion ──847 if gt_inject and gt:848 comps_i = list(comps_i) # copy to avoid mutating shard data849 comps_i[-1] = gt850 n_gt_injected += 1851 852 group_rewards = []853 for comp in comps_i:854 r = reward_funcs[0](855 [prompt], [comp], ground_truth=[gt], task=[task]856 )[0]857 group_rewards.append(r)858 comp_lengths.append(len(comp.split()))859 860 total_reward += sum(group_rewards)861 n_samples += len(group_rewards)862 863 mean_r = np.mean(group_rewards)864 std_r = np.std(group_rewards)865 if std_r < 1e-8:866 n_zero_std += 1867 if mean_r > 0.5:868 n_all_correct += 1869 else:870 n_all_wrong += 1871 872 # ── Zero-std rescue: assign ±baseline instead of skipping ──873 if zero_std_baseline > 0:874 n_zero_std_rescued += 1875 if mean_r > 0.5:876 # All correct → reinforce all877 adv_val = zero_std_baseline878 else:879 # All wrong → penalize all880 adv_val = -zero_std_baseline881 for comp in comps_i:882 pairs.append({883 "prompt_text": prompt_texts[i],884 "completion_text": comp,885 "advantage": adv_val,886 })887 continue888 889 n_mixed += 1890 advantages = [(r - mean_r) / std_r for r in group_rewards]891 892 if len(sample_log) < 5:893 sample_log.append({894 "ground_truth": gt,895 "group_rewards": group_rewards,896 "std": float(std_r),897 "completions": [c[:300] for c in comps_i],898 })899 900 for comp, adv in zip(comps_i, advantages):901 pairs.append({902 "prompt_text": prompt_texts[i],903 "completion_text": comp,904 "advantage": adv,905 })906 907 accuracy = total_reward / max(n_samples, 1)908 avg_len = np.mean(comp_lengths) if comp_lengths else 0909 max_len = max(comp_lengths) if comp_lengths else 0910 911 print(f" Mean reward (accuracy): {accuracy:.4f} (3-way random baseline: 0.333)")912 print(f" Zero-std groups: {n_zero_std}/{n} ({100*n_zero_std/max(n,1):.1f}%)")913 print(f" - All-correct (no signal): {n_all_correct}")914 print(f" - All-wrong (no signal): {n_all_wrong}")915 if gt_inject:916 print(f" GT injected into: {n_gt_injected}/{n} groups")917 if zero_std_baseline > 0:918 print(f" Zero-std rescued (±{zero_std_baseline}): {n_zero_std_rescued}")919 print(f" - Mixed (learnable): {n_mixed}")920 print(f" Training pairs (non-zero adv): {len(pairs)}")921 print(f" Completion length (words): avg={avg_len:.1f}, max={max_len}")922 923 # Save sample completions for inspection924 with open(output_dir / f"sample_completions_epoch{epoch}.json", "w") as f:925 json.dump(sample_log, f, indent=2, ensure_ascii=False)926 927 # Shuffle and save training pairs928 random.shuffle(pairs)929 pairs_path = output_dir / f"training_pairs_epoch{epoch}.jsonl"930 with open(pairs_path, "w") as f:931 for pair in pairs:932 f.write(json.dumps(pair, ensure_ascii=False) + "\n")933 print(f" Saved {len(pairs)} training pairs to {pairs_path}")934 935 print("Phase REWARDS complete.")936 937 938# ══════════════════════════════════════════════════════════════939# Phase: TRAIN (DDP across all GPUs)940# ══════════════════════════════════════════════════════════════941 942def phase_train(args):943 """DDP training across all GPUs with batched GRPO loss."""944 945 rank = int(os.environ.get("RANK", 0))946 local_rank = int(os.environ.get("LOCAL_RANK", 0))947 world_size = int(os.environ.get("WORLD_SIZE", 1))948 949 output_dir = Path(args.output_dir)950 epoch = args.epoch951 pairs_path = output_dir / f"training_pairs_epoch{epoch}.jsonl"952 953 if not pairs_path.exists():954 print(f"[Rank {rank}] ERROR: Training pairs file not found: {pairs_path}")955 print(f"[Rank {rank}] The REWARDS phase likely failed. Aborting training.")956 sys.exit(1)957 958 n_visible = torch.cuda.device_count()959 if local_rank >= n_visible:960 print(f"[Rank {rank}] LOCAL_RANK={local_rank} but only {n_visible} GPU(s) visible. "961 f"Clamping to device 0.")962 local_rank = 0963 964 # ── Init DDP ──965 dist.init_process_group(backend="nccl")966 torch.cuda.set_device(local_rank)967 device = torch.device(f"cuda:{local_rank}")968 make_deterministic(args.seed + rank)969 970 max_seq_len = args.max_prompt_length + args.max_completion_length971 972 if rank == 0:973 print(f"\n{'='*60}")974 print(f"DDP Training: epoch {epoch}, world_size={world_size}")975 print(f"{'='*60}")976 977 # ── Load training pairs ──978 all_pairs = []979 with open(pairs_path) as f:980 for line in f:981 line = line.strip()982 if line:983 all_pairs.append(json.loads(line))984 985 if not all_pairs:986 if rank == 0:987 print("WARNING: No training pairs. Skipping training.")988 dist.destroy_process_group()989 return990 991 # Pad to make divisible by world_size992 n_total = len(all_pairs)993 rem = n_total % world_size994 if rem != 0:995 all_pairs += all_pairs[:world_size - rem]996 997 # Shard: interleaved assignment across ranks998 my_pairs = all_pairs[rank::world_size]999 if rank == 0:1000 print(f"Total pairs: {n_total}, per-rank: {len(my_pairs)}")1001 1002 # ── Tokenizer ──1003 tokenizer = AutoTokenizer.from_pretrained(args.model_path, trust_remote_code=True)1004 tokenizer.pad_token = tokenizer.eos_token1005 tokenizer.padding_side = "right"1006 pad_id = tokenizer.pad_token_id1007 1008 # ── Pre-tokenize ──1009 dataset = GRPOPairDataset(my_pairs, tokenizer, max_seq_len)1010 if rank == 0:1011 print(f"Valid tokenized pairs per rank: {len(dataset)}")1012 1013 if len(dataset) == 0:1014 if rank == 0:1015 print("WARNING: No valid tokenized pairs. Skipping training.")1016 dist.destroy_process_group()1017 return1018 1019 # ── Load model with fresh LoRA ──1020 if epoch == 1:1021 merged_dir = str(output_dir / "_merged_model")1022 else:1023 merged_dir = str(output_dir / f"_merged_model_epoch{epoch - 1}")1024 1025 if rank == 0:1026 print(f"Loading merged model from {merged_dir} ...")1027 1028 model = AutoModelForCausalLM.from_pretrained(1029 merged_dir,1030 torch_dtype=torch.bfloat16,1031 trust_remote_code=True,1032 device_map={"": device},1033 )1034 1035 peft_config = LoraConfig(1036 r=args.lora_r,1037 lora_alpha=args.lora_alpha,1038 lora_dropout=args.lora_dropout,1039 bias="none",1040 task_type="CAUSAL_LM",1041 target_modules=[1042 "q_proj", "k_proj", "v_proj", "o_proj",1043 "gate_proj", "up_proj", "down_proj",1044 ],1045 )1046 model = get_peft_model(model, peft_config)1047 if rank == 0:1048 model.print_trainable_parameters()1049 1050 if args.gradient_checkpointing:1051 # use_reentrant=False: avoids illegal memory access with PEFT LoRA +1052 # DeepSpeed ZeRO-2 that occurs with the default reentrant checkpointing1053 # (caused v3 crash at rank15 ~step 3200; not triggered in v2 due to1054 # different dataset distribution with 4 vs 8 rollouts).1055 model.gradient_checkpointing_enable(1056 gradient_checkpointing_kwargs={"use_reentrant": False}1057 )1058 model.enable_input_require_grads()1059 1060 # ── Compute reference log-probs (batched, LoRA disabled) ──1061 if rank == 0:1062 print("Computing reference log-probs (batched) ...")1063 t0 = time.time()1064 1065 model.eval()1066 model.disable_adapter_layers()1067 1068 ref_logprobs = []1069 ref_loader = DataLoader(1070 dataset,1071 batch_size=args.ref_batch_size,1072 shuffle=False,1073 collate_fn=lambda b: collate_grpo(b, pad_id),1074 num_workers=0,1075 )1076 1077 with torch.no_grad():1078 for bi, batch in enumerate(ref_loader):1079 ids = batch["input_ids"].to(device)1080 mask = batch["attention_mask"].to(device)1081 logits = model(input_ids=ids, attention_mask=mask).logits1082 per_sample = extract_completion_logprobs(1083 logits, ids, batch["prompt_lens"], batch["comp_lens"]1084 )1085 ref_logprobs.extend([lp.cpu() for lp in per_sample])1086 if rank == 0 and (bi + 1) % 200 == 0:1087 print(f" Ref batch {bi + 1}/{len(ref_loader)}")1088 1089 model.enable_adapter_layers()1090 model.train()1091 1092 # Attach ref log-probs to dataset items1093 for i, rlp in enumerate(ref_logprobs):1094 dataset.items[i]["ref_logprobs"] = rlp1095 1096 if rank == 0:1097 print(f" Ref log-probs done in {time.time() - t0:.1f}s "1098 f"({len(ref_logprobs)} samples)")1099 1100 # ── DeepSpeed ZeRO-2 initialization ──1101 trainable_params = [p for p in model.parameters() if p.requires_grad]1102 optimizer = torch.optim.AdamW(trainable_params, lr=args.learning_rate)1103 1104 num_batches = len(dataset) // args.per_device_train_batch_size1105 num_update_steps = max(1, num_batches // args.gradient_accumulation_steps)1106 warmup_steps = int(args.warmup_ratio * num_update_steps)1107 scheduler = get_scheduler(1108 args.lr_scheduler_type,1109 optimizer,1110 num_warmup_steps=warmup_steps,1111 num_training_steps=num_update_steps,1112 )1113 1114 ds_config = {1115 "bf16": {"enabled": True},1116 "zero_optimization": {1117 "stage": 2,1118 "allgather_partitions": True,1119 "allgather_bucket_size": 200000000,1120 "overlap_comm": True,1121 "reduce_scatter": True,1122 "reduce_bucket_size": 200000000,1123 "contiguous_gradients": True,1124 },1125 "gradient_accumulation_steps": args.gradient_accumulation_steps,1126 "gradient_clipping": args.max_grad_norm,1127 "train_batch_size": (1128 args.per_device_train_batch_size * world_size1129 * args.gradient_accumulation_steps1130 ),1131 "train_micro_batch_size_per_gpu": args.per_device_train_batch_size,1132 "steps_per_print": 9999999,1133 }1134 1135 model_engine, optimizer, _, scheduler = deepspeed.initialize(1136 model=model,1137 optimizer=optimizer,1138 lr_scheduler=scheduler,1139 config=ds_config,1140 )1141 1142 if rank == 0:1143 eff_batch = world_size * args.per_device_train_batch_size * args.gradient_accumulation_steps1144 print(f"Training: {len(dataset)} pairs/rank, {num_update_steps} update steps")1145 print(f" Batch: {args.per_device_train_batch_size}/GPU × "1146 f"{args.gradient_accumulation_steps} accum × {world_size} GPUs = "1147 f"{eff_batch} effective")1148 print(f" DeepSpeed ZeRO-2 enabled")1149 1150 # ── Training loop ──1151 train_loader = DataLoader(1152 dataset,1153 batch_size=args.per_device_train_batch_size,1154 shuffle=True,1155 collate_fn=lambda b: collate_grpo(b, pad_id),1156 num_workers=0,1157 drop_last=True,1158 )1159 1160 global_step = 01161 log_loss = 0.01162 log_adv = 0.01163 log_n = 01164 1165 for step_idx, batch in enumerate(train_loader):1166 ids = batch["input_ids"].to(device)1167 mask = batch["attention_mask"].to(device)1168 1169 logits = model_engine(input_ids=ids, attention_mask=mask).logits1170 policy_lps = extract_completion_logprobs(1171 logits, ids, batch["prompt_lens"], batch["comp_lens"]1172 )1173 1174 loss, avg_adv = compute_grpo_batch_loss(1175 policy_lps, batch["ref_logprobs"], batch["advantages"],1176 batch["comp_lens"], args.beta,1177 )1178 1179 model_engine.backward(loss)1180 model_engine.step()1181 1182 log_loss += loss.detach().item()1183 log_adv += avg_adv1184 log_n += 11185 1186 if model_engine.is_gradient_accumulation_boundary():1187 global_step += 11188 1189 if rank == 0 and global_step % args.logging_steps == 0:1190 al = log_loss / max(log_n, 1)1191 aa = log_adv / max(log_n, 1)1192 lr = scheduler.get_last_lr()[0]1193 print(f" Step {global_step}/{num_update_steps} | "1194 f"Loss: {al:.4f} | |Adv|: {aa:.3f} | LR: {lr:.2e}")1195 log_loss = 0.01196 log_adv = 0.01197 log_n = 01198 1199 if (rank == 0 and args.save_steps1200 and global_step % args.save_steps == 0):