Team Ai
Datasetpublic

PerturbReason/PerturbReason_dataset_code

sourceHugging Facemitupdated 5mo agoView on Hugging Face
0likes50downloads
grpo_trainer_distributed_v2.py1312 linesDownload Raw Back to RL
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):

Showing the first 1,200 of 1312 lines. Download the file for the rest.