"""Utility functions for reward model evaluation. This module provides helper functions for: - Loading trained reward models - Computing loss end indices (padding detection) - Extracting scalar rewards from token-level rewards - Creating data sampling info for evaluation """ import inspect import os from typing import Tuple import numpy as np import torch from collections import defaultdict from data_utils_mmap import read_jsonl from modules.gpt import GPTConfig, GPTTrainConfig, GPT def load_reward_model( checkpoint_path: str, device: str = "cuda", ) -> Tuple[torch.nn.Module, dict]: """Load trained reward model from checkpoint. Args: checkpoint_path: Path to .pt checkpoint file device: Device to load model on Returns: model: Loaded GPT model with reward head in eval mode model_args: Model configuration dict """ print(f"Loading checkpoint from {checkpoint_path}") ckpt = torch.load(checkpoint_path, map_location="cpu", weights_only=False) # Get model args and ensure reward head is enabled model_args = ckpt["model_args"].copy() model_args["use_reward_head"] = True # Filter unknown parameters (e.g., use_mt5 from old checkpoints) valid_params = set(inspect.signature(GPTConfig.__init__).parameters.keys()) - {"self"} filtered_args = {k: v for k, v in model_args.items() if k in valid_params} unknown_params = set(model_args.keys()) - valid_params if unknown_params: print(f"Filtering out unknown parameters: {unknown_params}") # Initialize model config = GPTConfig(**filtered_args) model = GPT(config, GPTTrainConfig()) # Load weights model.load_state_dict(ckpt["model"], strict=False) # Convert to bfloat16 and move to device in one call (order matters!) model = model.to(device=device, dtype=torch.bfloat16) model.eval() # Verify all parameters are bfloat16 for name, param in model.named_parameters(): if param.dtype != torch.bfloat16: print(f"Warning: Parameter {name} is {param.dtype}, converting to bfloat16") param.data = param.data.to(torch.bfloat16) print(f"✓ Model loaded successfully") print(f" use_reward_head: {model.config.use_reward_head}") print(f" Output modules: {list(model.output_modules.keys())}") print(f" Parameters: {sum(p.numel() for p in model.parameters()):,}") print(f" Dtype: bfloat16") return model, filtered_args def compute_loss_end_indices(Y: torch.Tensor, seq_len: int, semantic_pad_token: int = 4000) -> list[int]: """Compute where valid data ends for each sample (before padding). Y[j]=-1 OR Y[j]=pad_token means X[j+1] is padding, so end_index = last_valid_j + 2 Args: Y: Targets (batch, n_codebooks, seq_len-1) with -1 or pad_token for padding seq_len: Sequence length of X/rewards semantic_pad_token: Padding token value for semantic codebook (default 4000) Returns: loss_end_index_list: End index for each sample """ batch_size = Y.shape[0] loss_end_index_list = [] for i in range(batch_size): y_sample = Y[i, 0, :] # First codebook (semantic) # Check for BOTH -1 (marked by mask_middle_padding) AND actual pad token # Non-padding means: not -1 AND not semantic_pad_token non_pad_mask = (y_sample != -1) & (y_sample != semantic_pad_token) if non_pad_mask.any(): last_valid = torch.where(non_pad_mask)[0][-1].item() end_idx = last_valid + 2 # Y-shift: Y[j]=pad → X[j+1] padding else: end_idx = seq_len loss_end_index_list.append(end_idx) return loss_end_index_list def extract_scalar_rewards( reward_logits: torch.Tensor, loss_start_index_list: list[int], loss_end_index_list: list[int] = None, ) -> torch.Tensor: """Extract scalar rewards by averaging between start and end indices. Args: reward_logits: Token-level rewards (batch, seq_len) loss_start_index_list: Where generation starts for each sample loss_end_index_list: Where valid data ends (None = seq_len, i.e., no padding) Returns: scalar_rewards: (batch,) """ batch_size, seq_len = reward_logits.shape device = reward_logits.device # Position indices positions = torch.arange(seq_len, device=device).unsqueeze(0).expand(batch_size, -1) # Start indices start_indices = ( torch.tensor(loss_start_index_list, device=device).unsqueeze(1) if loss_start_index_list else torch.zeros(batch_size, 1, dtype=torch.long, device=device) ) # End indices (default to seq_len if not provided) end_indices = ( torch.tensor(loss_end_index_list, device=device).unsqueeze(1) if loss_end_index_list else torch.full((batch_size, 1), seq_len, dtype=torch.long, device=device) ) # Simple mask: start <= position < end mask = (positions >= start_indices) & (positions < end_indices) # Average masked_rewards = reward_logits * mask valid_counts = mask.sum(dim=1, keepdim=True).clamp(min=1) scalar_rewards = masked_rewards.sum(dim=1) / valid_counts.squeeze(1) return scalar_rewards def create_data_sampling_info( data_dir: str, split: str, cfg: GPTConfig, tokenizer_fp: str, device: str = "cuda", ) -> dict: """Create data_sampling_info dict exactly like train_reward_model.py. Args: data_dir: Path to DPO data directory split: "train" or "val" cfg: Model configuration tokenizer_fp: Path to tokenizer device: Device Returns: data_sampling_info: Dict with all data loading info """ # Load files - match train_reward_model.py exactly if split == "val": filename = "data_val.bin" metas_filename = "meta_val.jsonl" info_filename = "info_val.json" else: filename = "data_tr.bin" metas_filename = "meta_tr.jsonl" info_filename = "info_tr.json" # Load memory-mapped data t_data_memmap = 12000 # From training script # Use config values for flexibility (not hardcoded) semantic_n_codebooks = cfg.semantic_n_codebooks data_coarse_n_codebooks = cfg.coarse_n_codebooks 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) # Load metadata metas = read_jsonl(os.path.join(data_dir, metas_filename)) # Load info with open(os.path.join(data_dir, info_filename)) as f: import json infos = json.load(f) # Build dataset structure (from train_reward_model.py load_dataset function) dataset_names = [] data_idx_lists = [] data_weights = [] # 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": idx_list = sorted(info["idx_list"]) else: raise ValueError(f"Unsupported task: {info.get('task')}") data_idx_lists.append(idx_list) data_weights.append(len(idx_list)) # Normalize weights weights_norm = np.sum(data_weights) data_weights = [v / weights_norm for v in data_weights] # Create artist_to_songs (empty for simplicity) artist_to_songs = defaultdict(list) # Create data_sampling_info structure data_sampling_info = { "cfg": cfg, "train_cfg": GPTTrainConfig(), "batch_size": 2, # 1 pair = 2 samples "batch_size_tokens": cfg.block_size * 2, "tokenizer_fp": tokenizer_fp, "device": device, "device_type": "cuda" if "cuda" in device else "cpu", split: { "data": data, "metas": metas, "infos": infos, "artist_to_songs": artist_to_songs, "names": dataset_names, "weights": data_weights, "idx_lists": data_idx_lists, "all_idx_lists": sorted([idx for sublist in data_idx_lists for idx in sublist]), }, } return data_sampling_info def load_dpo_sample( split: str, pair_idx: int, data_sampling_info: dict, ) -> Tuple: """Load a DPO pair (chosen and rejected samples) using get_batch. Args: split: "train" or "val" pair_idx: Index of the pair to load data_sampling_info: Data sampling info dict Returns: X_chosen, Y_chosen, meta_chosen, X_rejected, Y_rejected, meta_rejected, loss_start_idx """ from utils.dpo_data_utils import get_batch # Load the pair using get_batch idxs, X, Y, loss_start_index_list = get_batch( data_sampling_info, split, dummy_data=False, row_idx=[pair_idx, pair_idx], suppress_text=False, load_dpo_pair=True, return_idx=True, ) # Get metadata metas = data_sampling_info[split]["metas"] # Split into chosen (odd) and rejected (even) X_rejected = X[0:1] Y_rejected = Y[0:1] meta_rejected = metas[idxs[0]] X_chosen = X[1:2] Y_chosen = Y[1:2] meta_chosen = metas[idxs[1]] loss_start_idx = loss_start_index_list[1] return ( X_chosen, Y_chosen, meta_chosen, X_rejected, Y_rejected, meta_rejected, loss_start_idx, )