#!/usr/bin/env python """Cache reward model outputs for all DPO dataset samples. This script pre-computes reward scores for all samples in a DPO dataset and saves them to a JSON file for fast lookup during training. Similar to how train_dpo.py caches reference model losses. Usage: # Single node python scripts/cache_reward.py \\ --checkpoint /path/to/reward_model.pt \\ --data_dir /path/to/dpo/data \\ --output_name "reward_crow_r1" \\ --batch_size 16 # Multi-node (via SLURM) srun python scripts/cache_reward.py \\ --checkpoint /path/to/reward_model.pt \\ --data_dir /path/to/dpo/data \\ --output_name "reward_crow_r1" \\ --batch_size 16 """ import argparse import datetime import gc import json import logging import math import os import sys import time from collections import defaultdict from typing import Dict import numpy as np import torch from torch.distributed import destroy_process_group, init_process_group from tqdm import tqdm # Add parent directory to path sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) from data_utils_mmap import read_jsonl from modules.gpt import GPTTrainConfig from scripts.reward_eval_utils import ( load_reward_model, compute_loss_end_indices, extract_scalar_rewards, ) from utils.dpo_data_utils import get_batch from utils.helpers import dist_barrier, print_with_time, print_with_time_master # Turn down some annoying logging logging.getLogger("torch.distributed.fsdp._debug_utils").setLevel(logging.ERROR) logging.getLogger("torch.distributed.fsdp._optim_utils").setLevel(logging.ERROR) logging.getLogger("torch.distributed.checkpoint._dedup_tensors").setLevel(logging.ERROR) def setup_distributed(master_addr: str = "localhost", master_port: str = "12355"): """Setup distributed training environment. Copied from train_dpo.py lines 212-260. """ os.environ["MASTER_ADDR"] = str(master_addr) os.environ["MASTER_PORT"] = str(master_port) if "SLURM_PROCID" in os.environ: # Running on SLURM if int(os.environ["SLURM_NTASKS_PER_NODE"]) != torch.cuda.device_count(): raise ValueError( f"SLURM_NTASKS_PER_NODE ({os.environ['SLURM_NTASKS_PER_NODE']}) does not match" f" the number of CUDA devices ({torch.cuda.device_count()}) on node {os.environ['HOSTNAME']}" ) ddp_rank = int(os.environ["SLURM_PROCID"]) ddp_local_rank = int(os.environ["SLURM_LOCALID"]) world_size = int(os.environ["SLURM_JOB_NUM_NODES"]) * int(os.environ["SLURM_NTASKS_PER_NODE"]) else: # Running locally ddp_rank = 0 ddp_local_rank = 0 world_size = 1 os.environ["RANK"] = str(ddp_rank) os.environ["LOCAL_RANK"] = str(ddp_local_rank) # Init process group if distributed ddp = int(os.environ.get("RANK", -1)) != -1 if ddp: try: init_process_group( backend="nccl", timeout=datetime.timedelta(seconds=24 * 60 * 60), rank=ddp_rank, world_size=world_size, device_id=torch.device(f"cuda:{ddp_local_rank}"), ) except Exception as e: print(f"Distributed error on rank {ddp_rank} with host {os.environ['HOSTNAME']}") raise e device = f"cuda:{ddp_local_rank}" torch.cuda.set_device(device) master_process = ddp_rank == 0 print_with_time(f"ddp init, rank {ddp_rank}, local_rank {ddp_local_rank}") else: device = "cuda" torch.cuda.set_device(device) master_process = True ddp_rank = 0 ddp_local_rank = 0 world_size = 1 n_gpus_per_node = torch.cuda.device_count() dist_barrier() print_with_time_master(f"ddp init: world size {world_size} ddp_rank {ddp_rank}.") return ddp, ddp_rank, ddp_local_rank, world_size, n_gpus_per_node, device, master_process def load_dataset( data_dir: str, filename: str, info_filename: str, metas_filename: str, t_data_memmap: int, semantic_n_codebooks: int, data_coarse_n_codebooks: int, ) -> tuple: """Load dataset for reward caching. Simplified version of train_dpo.py load_dataset (lines 470-561). """ dataset_names = [] data_idx_lists = [] data_weights = [] data = np.memmap(os.path.join(data_dir, filename), dtype=np.uint16, mode="r") data = data.reshape(-1, t_data_memmap, semantic_n_codebooks + data_coarse_n_codebooks) with open(os.path.join(data_dir, info_filename)) as f: infos = json.load(f) metas = read_jsonl(os.path.join(data_dir, metas_filename)) assert len(data) == len(metas), (len(data), len(metas)) # Empty artist_to_songs (not needed for caching) artist_to_songs = defaultdict(list) # Process infos for dset_name in sorted(infos.keys()): if "idx_map" in infos[dset_name]: infos[dset_name]["idx_map"] = {int(k): v for k, v in infos[dset_name]["idx_map"].items()} for dset_name in sorted(infos.keys()): info = infos[dset_name] dataset_names.append(dset_name) if info.get("task", "default") == "default": assert "idx_list" in info idx_list = sorted(info["idx_list"]) else: raise ValueError(f"unsupported task for {dset_name} in info file") data_idx_lists.append(idx_list) data_weights.append(len(idx_list)) weights_norm = np.sum(data_weights) data_weights = [v / weights_norm for v in data_weights] print_with_time_master(f"{len(data):,} lines of {filename} loaded.") assert len(infos) == len(dataset_names) == len(data_weights) == len(data_idx_lists) return ( dataset_names, data_idx_lists, data_weights, data, metas, infos, artist_to_songs, ) def cache_rewards_for_split( split: str, model: torch.nn.Module, data_sampling_info: dict, eval_loss_batch_size: int, ddp_rank: int, ddp_local_rank: int, world_size: int, n_gpus_per_node: int, device: str, local_data_shard_dir: str, log_interval: int = 25, ) -> Dict[int, Dict]: """Cache rewards for a data split (train or val). Similar to update_loss_dict_lookup() in train_dpo.py (lines 1187-1266). Args: split: "train" or "val" model: Reward model in eval mode data_sampling_info: Data sampling info dict eval_loss_batch_size: Batch size for evaluation ddp_rank: Global rank ddp_local_rank: Local rank world_size: Total number of processes n_gpus_per_node: GPUs per node device: Device string local_data_shard_dir: Local data shard directory (None if not sharded) log_interval: Logging interval Returns: local_rewards: Dict mapping idx -> {"reward": float} """ print_with_time_master(f"Computing rewards for {split} split...") # Determine if using global or local rank if local_data_shard_dir is None: eval_ddp_rank = ddp_rank eval_world_size = world_size else: eval_ddp_rank = ddp_local_rank eval_world_size = n_gpus_per_node # Get all data indices data_idx_lists = data_sampling_info[split]["idx_lists"] data_idx_lists_flat = sorted([idx for sublist in data_idx_lists for idx in sublist]) # Estimate iterations est_eval_tot_iter_num = int( math.ceil(len(data_idx_lists_flat) / (eval_loss_batch_size * eval_world_size)) ) print_with_time_master( f"Estimated total iterations: {est_eval_tot_iter_num}, " f"data size: {len(data_idx_lists_flat)}, " f"batch size: {eval_loss_batch_size}, " f"ddp_rank: {eval_ddp_rank}" ) local_rewards = {} t0 = time.time() # Create eval data sampling info with larger batch size eval_data_sampling_info = data_sampling_info.copy() eval_data_sampling_info["batch_size"] = eval_loss_batch_size eval_data_sampling_info["batch_size_tokens"] = ( data_sampling_info["cfg"].block_size * eval_loss_batch_size ) model.eval() ctx = torch.amp.autocast(device_type="cuda", dtype=torch.bfloat16) for eval_iter_num in tqdm( range(est_eval_tot_iter_num), desc=f"Caching {split} rewards", disable=(ddp_rank != 0), ): # Get batch indices for all GPUs global_batch_row_idx_list = data_idx_lists_flat[ eval_iter_num * eval_loss_batch_size * eval_world_size : (eval_iter_num + 1) * eval_loss_batch_size * eval_world_size ] # Subslice for this GPU batch_row_idx_list = global_batch_row_idx_list[ eval_ddp_rank * eval_loss_batch_size : (eval_ddp_rank + 1) * eval_loss_batch_size ] real_batch_size = len(batch_row_idx_list) if real_batch_size == 0: continue # Pad to full batch if needed batch_row_idx_list += [0] * (eval_loss_batch_size - real_batch_size) # Load batch WITHOUT DPO pairing - each sample individually _, X, Y, loss_start_index_list = get_batch( eval_data_sampling_info, split, abs_row_idx=batch_row_idx_list, min_text_offs=0, suppress_text=False, dummy_data=False, inference=True, return_idx=True, load_dpo_pair=False, # Load individually, not paired ) X = X.to(device) Y = Y.to(device) # Compute rewards for batch with torch.no_grad(): with ctx: output = model(X, return_logits=True) reward_logits = output["reward_logits"] # (batch, seq_len) # Compute end indices from Y loss_end_index_list = compute_loss_end_indices(Y, reward_logits.shape[1]) # Extract scalar rewards scalar_rewards = extract_scalar_rewards( reward_logits, loss_start_index_list, loss_end_index_list ) # Store rewards for each sample for i in range(real_batch_size): idx = batch_row_idx_list[i] local_rewards[idx] = {"reward": scalar_rewards[i].item()} # Logging t1 = time.time() dt = t1 - t0 t0 = t1 if ddp_rank == 0 and (eval_iter_num % log_interval == 0): tokens_per_s = eval_loss_batch_size * data_sampling_info["cfg"].block_size / dt print_with_time_master( f"iter {eval_iter_num}/{est_eval_tot_iter_num}: " f"step_time {dt * 1000:.1f}ms, " f"throughput {tokens_per_s / 1e3:,.0f}k tok/s/node" ) print_with_time_master(f"Computed {len(local_rewards)} rewards for {split} split on rank {ddp_rank}") return local_rewards def main(): """Main entry point for reward caching script.""" parser = argparse.ArgumentParser( description="Cache reward model outputs for DPO dataset", formatter_class=argparse.ArgumentDefaultsHelpFormatter, ) parser.add_argument( "--checkpoint", required=True, help="Path to trained reward model checkpoint (.pt file)", ) parser.add_argument( "--data_dir", required=True, help="Path to DPO data directory (containing data_tr.bin, meta_tr.jsonl, etc.)", ) parser.add_argument( "--output_name", default="reward_model", help="Cache filename prefix (output: {output_name}_cached_reward.json)", ) parser.add_argument( "--batch_size", type=int, default=16, help="Batch size for evaluation (larger = faster)", ) parser.add_argument( "--master_addr", default="localhost", help="Master node address for distributed training", ) parser.add_argument( "--master_port", default="12355", help="Master node port for distributed training", ) parser.add_argument( "--local_data_shard_dir", default=None, help="Local data shard directory (for multi-node setups)", ) parser.add_argument( "--t_data_memmap", type=int, default=12_000, help="Time dimension of memmap data (must match training config)", ) args = parser.parse_args() # Setup distributed ( ddp, ddp_rank, ddp_local_rank, world_size, n_gpus_per_node, device, master_process, ) = setup_distributed(args.master_addr, args.master_port) # Disable GC for performance gc.disable() print_with_time_master("=" * 80) print_with_time_master("REWARD MODEL CACHING") print_with_time_master("=" * 80) print_with_time_master(f"Checkpoint: {args.checkpoint}") print_with_time_master(f"Data directory: {args.data_dir}") print_with_time_master(f"Output name: {args.output_name}") print_with_time_master(f"Batch size: {args.batch_size}") print_with_time_master(f"World size: {world_size}") print_with_time_master("") # Check if cache already exists cache_path = os.path.join(args.data_dir, f"{args.output_name}_cached_reward.json") if os.path.exists(cache_path) and master_process: print_with_time_master(f"Cache already exists at {cache_path}") print_with_time_master("Exiting...") if ddp: destroy_process_group() return # Load reward model print_with_time_master("=" * 80) print_with_time_master("Loading reward model...") print_with_time_master("=" * 80) # Suppress prints on non-master processes during model loading if not master_process: import io import contextlib with contextlib.redirect_stdout(io.StringIO()): model, model_args = load_reward_model(args.checkpoint, device) else: model, model_args = load_reward_model(args.checkpoint, device) cfg = model.config # Verify reward head if not cfg.use_reward_head: raise RuntimeError(f"Model use_reward_head is {cfg.use_reward_head}, expected True") if "reward_head" not in model.output_modules: raise RuntimeError("Model does not have reward_head in output_modules") print_with_time_master("✓ Reward model loaded with reward head") dist_barrier() # Constants from model config and args t_data_memmap = args.t_data_memmap semantic_n_codebooks = cfg.semantic_n_codebooks data_coarse_n_codebooks = cfg.coarse_n_codebooks print_with_time_master( f"Data shape config: t_data_memmap={t_data_memmap}, " f"semantic_n_codebooks={semantic_n_codebooks}, " f"data_coarse_n_codebooks={data_coarse_n_codebooks}" ) # Load datasets print_with_time_master("=" * 80) print_with_time_master("Loading datasets...") print_with_time_master("=" * 80) data_dir = args.local_data_shard_dir if args.local_data_shard_dir else args.data_dir ( val_dataset_names, val_data_idx_lists, val_data_weights, val_data, val_metas, val_info, val_artist_to_songs, ) = load_dataset( data_dir, "data_val.bin", "info_val.json", "meta_val.jsonl", t_data_memmap, semantic_n_codebooks, data_coarse_n_codebooks, ) ( train_dataset_names, train_data_idx_lists, train_data_weights, train_data, train_metas, train_info, train_artist_to_songs, ) = load_dataset( data_dir, "data_tr.bin", "info_tr.json", "meta_tr.jsonl", t_data_memmap, semantic_n_codebooks, data_coarse_n_codebooks, ) dist_barrier() # Create data_sampling_info tokenizer_fp = os.path.join(data_dir, "tokenizer_60k.json") if not os.path.exists(tokenizer_fp): print_with_time_master(f"Warning: Tokenizer not found at {tokenizer_fp}") tokenizer_fp = None data_sampling_info = { "cfg": cfg, "train_cfg": GPTTrainConfig(), "batch_size": args.batch_size, "batch_size_tokens": cfg.block_size * args.batch_size, "tokenizer_fp": tokenizer_fp, "device": device, "device_type": "cuda" if "cuda" in device else "cpu", "train": { "data": train_data, "metas": train_metas, "infos": train_info, "artist_to_songs": train_artist_to_songs, "names": train_dataset_names, "weights": train_data_weights, "idx_lists": train_data_idx_lists, "all_idx_lists": sorted([idx for sublist in train_data_idx_lists for idx in sublist]), }, "val": { "data": val_data, "metas": val_metas, "infos": val_info, "artist_to_songs": val_artist_to_songs, "names": val_dataset_names, "weights": val_data_weights, "idx_lists": val_data_idx_lists, "all_idx_lists": sorted([idx for sublist in val_data_idx_lists for idx in sublist]), }, } # Cache rewards for both splits print_with_time_master("=" * 80) print_with_time_master("Computing rewards...") print_with_time_master("=" * 80) local_rewards_train = cache_rewards_for_split( "train", model, data_sampling_info, args.batch_size, ddp_rank, ddp_local_rank, world_size, n_gpus_per_node, device, args.local_data_shard_dir, ) local_rewards_val = cache_rewards_for_split( "val", model, data_sampling_info, args.batch_size, ddp_rank, ddp_local_rank, world_size, n_gpus_per_node, device, args.local_data_shard_dir, ) # Gather all results across GPUs print_with_time_master("=" * 80) print_with_time_master("Gathering results...") print_with_time_master("=" * 80) all_rewards = [None for _ in range(world_size)] local_rewards = {"train": local_rewards_train, "val": local_rewards_val} if ddp: torch.distributed.all_gather_object(all_rewards, local_rewards) else: all_rewards = [local_rewards] # Merge and save on master process if master_process: print_with_time_master("Merging results from all ranks...") final_rewards = {"train": {}, "val": {}} for worker_rewards in all_rewards: final_rewards["train"].update(worker_rewards["train"]) final_rewards["val"].update(worker_rewards["val"]) print_with_time_master(f"Total train rewards: {len(final_rewards['train'])}") print_with_time_master(f"Total val rewards: {len(final_rewards['val'])}") # Verify we have all samples expected_train = len(data_sampling_info["train"]["all_idx_lists"]) expected_val = len(data_sampling_info["val"]["all_idx_lists"]) if len(final_rewards["train"]) != expected_train: print_with_time_master( f"WARNING: Expected {expected_train} train samples, got {len(final_rewards['train'])}" ) if len(final_rewards["val"]) != expected_val: print_with_time_master( f"WARNING: Expected {expected_val} val samples, got {len(final_rewards['val'])}" ) # Save to JSON print_with_time_master(f"Saving cached rewards to {cache_path}") with open(cache_path, "w") as f: json.dump(final_rewards, f) print_with_time_master("=" * 80) print_with_time_master("✅ CACHING COMPLETE!") print_with_time_master("=" * 80) print_with_time_master(f"Cache saved to: {cache_path}") print_with_time_master(f"File size: {os.path.getsize(cache_path) / 1024 / 1024:.2f} MB") print_with_time_master("") dist_barrier() if ddp: destroy_process_group() if __name__ == "__main__": main()