#!/usr/bin/env python """Comprehensive reward model evaluation suite. Usage: python scripts/eval_reward_model.py \\ --checkpoint /path/to/model.pt \\ --data_dir /path/to/dpo/data \\ --output_dir ./reward_eval_results \\ --n_test_cases 10 \\ --n_viz_pairs 5 \\ --token_interval 30 \\ --max_val_samples 1000 This script: 1. Loads trained reward model 2. Tests on sample pairs 3. Visualizes reward progression every ~750 tokens 4. Runs full validation 5. Generates plots and statistics in timestamped output directory """ import argparse import json import os import sys from datetime import datetime from typing import List, Dict import matplotlib.pyplot as plt import numpy as np import torch from tqdm import tqdm # Add parent directory to path sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) from scripts.reward_eval_utils import ( load_reward_model, create_data_sampling_info, load_dpo_sample, extract_scalar_rewards, compute_loss_end_indices, ) def run_test_cases( model: torch.nn.Module, data_sampling_info: dict, output_dir: str, n_samples: int = 10, device: str = "cuda", ) -> List[Dict]: """Run model on sample test cases and save results. Args: model: Trained reward model data_sampling_info: Data sampling info dict output_dir: Directory to save results n_samples: Number of test pairs to evaluate device: Device for computation Returns: results: List of dicts with test case results """ results = [] errors = [] print(f"Evaluating {n_samples} test pairs...") for idx in range(n_samples): try: X_c, Y_c, meta_c, X_r, Y_r, meta_r, start_idx = load_dpo_sample( "val", idx, data_sampling_info ) # Get rewards (use autocast for bfloat16) with torch.no_grad(): with torch.amp.autocast(device_type="cuda", dtype=torch.bfloat16): output_c = model(X_c, return_logits=True) output_r = model(X_r, return_logits=True) # Extract reward_logits from dict rewards_c = output_c["reward_logits"] rewards_r = output_r["reward_logits"] # Compute end indices from Y end_idx_c = compute_loss_end_indices(Y_c, rewards_c.shape[1]) end_idx_r = compute_loss_end_indices(Y_r, rewards_r.shape[1]) # Extract scalars scalar_c = extract_scalar_rewards(rewards_c, [start_idx], end_idx_c) scalar_r = extract_scalar_rewards(rewards_r, [start_idx], end_idx_r) correct = scalar_c > scalar_r margin = (scalar_c - scalar_r).item() results.append( { "pair_idx": idx, "chosen_reward": scalar_c.item(), "rejected_reward": scalar_r.item(), "margin": margin, "correct": bool(correct), "chosen_text": meta_c.get("text", "")[:60], "rejected_text": meta_r.get("text", "")[:60], } ) except Exception as e: error_msg = f"Pair {idx}: {str(e)}" print(f" ⚠ Error: {error_msg}") import traceback errors.append({"pair_idx": idx, "error": str(e), "traceback": traceback.format_exc()}) continue # Save error log if any if errors: error_file = os.path.join(output_dir, "errors.txt") with open(error_file, "w") as f: f.write("ERRORS DURING EVALUATION\n") f.write("=" * 80 + "\n\n") for err in errors: f.write(f"Pair {err['pair_idx']}:\n") f.write(f" Error: {err['error']}\n") f.write(f" Traceback:\n{err['traceback']}\n") print(f"⚠ {len(errors)} errors logged to {error_file}") # Save results to file output_file = os.path.join(output_dir, "test_cases.txt") with open(output_file, "w") as f: f.write("REWARD MODEL TEST CASES\n") f.write("=" * 80 + "\n\n") for r in results: f.write(f"Pair {r['pair_idx']}:\n") f.write(f" Chosen: {r['chosen_reward']:.4f} - {r['chosen_text']}\n") f.write(f" Rejected: {r['rejected_reward']:.4f} - {r['rejected_text']}\n") f.write(f" Margin: {r['margin']:+.4f}\n") f.write(f" Correct: {'✓' if r['correct'] else '✗'}\n\n") accuracy = sum(r["correct"] for r in results) / len(results) if results else 0 f.write(f"\n") f.write(f"Accuracy: {accuracy:.1%} ({sum(r['correct'] for r in results)}/{len(results)})\n") f.write(f"Mean Margin: {np.mean([r['margin'] for r in results]):.4f}\n") print(f"✓ Test cases saved to {output_file}") return results def visualize_reward_progression( model: torch.nn.Module, data_sampling_info: dict, output_dir: str, n_pairs: int = 5, token_interval: int = 30, device: str = "cuda", ): """Plot how rewards change across sequence positions. Args: model: Trained reward model data_sampling_info: Data sampling info dict output_dir: Directory to save plots n_pairs: Number of pairs to visualize token_interval: Sample every N tokens (30 tokens ≈ 1.2 seconds at 25Hz) device: Device for computation """ print(f"Creating reward progression plots for {n_pairs} pairs...") fig, axes = plt.subplots(n_pairs, 1, figsize=(14, 3 * n_pairs)) if n_pairs == 1: axes = [axes] errors = [] for i in range(n_pairs): try: X_c, Y_c, meta_c, X_r, Y_r, meta_r, start_idx = load_dpo_sample("val", i, data_sampling_info) # Get token-level rewards (use autocast for bfloat16) with torch.no_grad(): with torch.amp.autocast(device_type="cuda", dtype=torch.bfloat16): output_c = model(X_c, return_logits=True) output_r = model(X_r, return_logits=True) # Extract reward_logits from dict rewards_c = output_c["reward_logits"][0] # (seq_len,) rewards_r = output_r["reward_logits"][0] # (seq_len,) # Compute end index from Y end_idx_list = compute_loss_end_indices(Y_c, len(rewards_c)) end_idx = end_idx_list[0] # Valid positions: from start_idx to end_idx valid_pos = np.arange(start_idx, end_idx) # Compute progressive/cumulative rewards # Point 1: Mean reward of tokens [start_idx : start_idx+750] # Point 2: Mean reward of tokens [start_idx : start_idx+1500] # etc. step_sizes = list(range(token_interval, len(valid_pos) + token_interval, token_interval)) progressive_rewards_c = [] progressive_rewards_r = [] progressive_times = [] for end_offset in step_sizes: # Get tokens from start to current end position end_pos = min(end_offset, len(valid_pos)) tokens_to_average = valid_pos[:end_pos] # Compute mean reward up to this point avg_reward_c = rewards_c[tokens_to_average].mean().item() avg_reward_r = rewards_r[tokens_to_average].mean().item() progressive_rewards_c.append(avg_reward_c) progressive_rewards_r.append(avg_reward_r) # Time = number of tokens / 25Hz progressive_times.append(end_pos / 25.0) if end_pos >= len(valid_pos): break ax = axes[i] ax.plot( progressive_times, progressive_rewards_c, "g-o", label="Chosen", linewidth=2, markersize=5, alpha=0.8, ) ax.plot( progressive_times, progressive_rewards_r, "r-s", label="Rejected", linewidth=2, markersize=5, alpha=0.8, ) ax.axhline(y=0, color="k", linestyle="--", alpha=0.3, linewidth=1) ax.set_ylabel("Cumulative Reward", fontsize=10) # Title with text sample title_text = meta_c.get("text", "")[:60] if len(meta_c.get("text", "")) > 60: title_text += "..." ax.set_title(f"Pair {i}: {title_text}", fontsize=11) ax.legend(loc="best", fontsize=9) ax.grid(True, alpha=0.3) except Exception as e: import traceback print(f" ⚠ Error on pair {i}: {e}") errors.append({"pair_idx": i, "error": str(e)}) # Leave this subplot empty or put error message ax = axes[i] ax.text(0.5, 0.5, f"Error loading pair {i}", ha="center", va="center") ax.set_title(f"Pair {i}: Error") continue if errors: print(f"⚠ {len(errors)} errors during visualization") axes[-1].set_xlabel("Sequence Length (seconds)", fontsize=11) fig.suptitle( f"Progressive Reward: Average from Start to T (steps every {token_interval} tokens ≈ {token_interval/25:.1f}s)", fontsize=14, y=0.998, ) plt.tight_layout() output_file = os.path.join(output_dir, "reward_progression.png") plt.savefig(output_file, dpi=150, bbox_inches="tight") plt.close() print(f"✓ Reward progression plot saved to {output_file}") def run_full_validation( model: torch.nn.Module, data_sampling_info: dict, output_dir: str, max_samples: int = None, device: str = "cuda", ) -> Dict: """Evaluate on full validation set. Args: model: Trained reward model data_sampling_info: Data sampling info dict output_dir: Directory to save results max_samples: Maximum number of pairs to evaluate (None = all) device: Device for computation Returns: stats: Dictionary with validation statistics """ # Get number of pairs from data metas = data_sampling_info["val"]["metas"] n_pairs = len(metas) // 2 if max_samples: n_pairs = min(n_pairs, max_samples) print(f"Running full validation on {n_pairs} pairs...") results = [] errors = [] for idx in tqdm(range(n_pairs), desc="Evaluating pairs"): try: X_c, Y_c, _, X_r, Y_r, _, start_idx = load_dpo_sample("val", idx, data_sampling_info) with torch.no_grad(): with torch.amp.autocast(device_type="cuda", dtype=torch.bfloat16): output_c = model(X_c, return_logits=True) output_r = model(X_r, return_logits=True) # Extract reward_logits from dict rewards_c = output_c["reward_logits"] rewards_r = output_r["reward_logits"] # Compute end indices from Y end_idx_c = compute_loss_end_indices(Y_c, rewards_c.shape[1]) end_idx_r = compute_loss_end_indices(Y_r, rewards_r.shape[1]) scalar_c = extract_scalar_rewards(rewards_c, [start_idx], end_idx_c) scalar_r = extract_scalar_rewards(rewards_r, [start_idx], end_idx_r) results.append( { "chosen": scalar_c.item(), "rejected": scalar_r.item(), "margin": (scalar_c - scalar_r).item(), "correct": bool(scalar_c > scalar_r), } ) except Exception as e: errors.append({"pair_idx": idx, "error": str(e)}) continue # Log errors if errors: error_file = os.path.join(output_dir, "validation_errors.txt") with open(error_file, "w") as f: f.write(f"ERRORS DURING FULL VALIDATION\n") f.write(f"=" * 80 + "\n\n") f.write(f"Total errors: {len(errors)} / {n_pairs} ({len(errors)/n_pairs*100:.1f}%)\n\n") for err in errors[:100]: # Log first 100 errors f.write(f"Pair {err['pair_idx']}: {err['error']}\n") if len(errors) > 100: f.write(f"\n... and {len(errors) - 100} more errors\n") print(f"\n⚠ {len(errors)} errors during validation (logged to {error_file})") # Compute statistics if not results: print("No valid results!") return {} accuracy = sum(r["correct"] for r in results) / len(results) margins = [r["margin"] for r in results] mean_margin = np.mean(margins) std_margin = np.std(margins) median_margin = np.median(margins) stats = { "n_pairs": len(results), "accuracy": float(accuracy), "mean_margin": float(mean_margin), "std_margin": float(std_margin), "median_margin": float(median_margin), "correct_count": sum(r["correct"] for r in results), "min_margin": float(np.min(margins)), "max_margin": float(np.max(margins)), } # Save statistics stats_file = os.path.join(output_dir, "validation_stats.json") with open(stats_file, "w") as f: json.dump(stats, f, indent=2) # Plot margin distribution plt.figure(figsize=(10, 6)) plt.hist(margins, bins=50, edgecolor="black", alpha=0.7, color="steelblue") plt.axvline(x=0, color="r", linestyle="--", linewidth=2, label="Zero margin", alpha=0.8) plt.axvline( x=mean_margin, color="g", linestyle="-", linewidth=2, label=f"Mean: {mean_margin:.3f}", alpha=0.8 ) plt.xlabel("Reward Margin (chosen - rejected)", fontsize=12) plt.ylabel("Count", fontsize=12) plt.title( f"Validation: Accuracy {accuracy:.1%}, Mean Margin {mean_margin:.3f} ± {std_margin:.3f}", fontsize=14, ) plt.legend(fontsize=10) plt.grid(True, alpha=0.3) margin_plot = os.path.join(output_dir, "margin_distribution.png") plt.savefig(margin_plot, dpi=150, bbox_inches="tight") plt.close() print(f"✓ Validation complete: {accuracy:.1%} accuracy, mean margin {mean_margin:.3f}") print(f"✓ Margin distribution saved to {margin_plot}") return stats def create_summary( output_dir: str, test_results: List[Dict], val_stats: Dict, checkpoint_path: str, ): """Create human-readable summary file. Args: output_dir: Directory to save summary test_results: Results from test cases val_stats: Statistics from full validation checkpoint_path: Path to model checkpoint used """ summary_path = os.path.join(output_dir, "summary.txt") with open(summary_path, "w") as f: f.write("REWARD MODEL EVALUATION SUMMARY\n") f.write("=" * 80 + "\n\n") f.write(f"Timestamp: {datetime.now().strftime('%Y-%m-%d %H:%M:%S')}\n") f.write(f"Checkpoint: {checkpoint_path}\n\n") f.write("TEST CASES (Sample Pairs)\n") f.write("-" * 80 + "\n") if test_results: test_acc = sum(r["correct"] for r in test_results) / len(test_results) test_margin = np.mean([r["margin"] for r in test_results]) f.write(f"Samples: {len(test_results)}\n") f.write( f"Accuracy: {test_acc:.1%} ({sum(r['correct'] for r in test_results)}/{len(test_results)})\n" ) f.write(f"Mean Margin: {test_margin:.4f}\n") f.write(f"See test_cases.txt for details\n\n") else: f.write("No test results\n\n") f.write("FULL VALIDATION\n") f.write("-" * 80 + "\n") if val_stats and "n_pairs" in val_stats: f.write(f"Total Pairs: {val_stats['n_pairs']}\n") f.write(f"Accuracy: {val_stats['accuracy']:.3f} ({val_stats['accuracy']:.1%})\n") f.write(f"Mean Margin: {val_stats['mean_margin']:.4f}\n") f.write(f"Std Margin: {val_stats['std_margin']:.4f}\n") f.write(f"Median Margin: {val_stats['median_margin']:.4f}\n") f.write(f"Min Margin: {val_stats['min_margin']:.4f}\n") f.write(f"Max Margin: {val_stats['max_margin']:.4f}\n") f.write(f"Correct: {val_stats['correct_count']} / {val_stats['n_pairs']}\n\n") else: f.write("Skipped (disabled in evaluation script)\n\n") f.write("GENERATED FILES\n") f.write("-" * 80 + "\n") f.write(f" - config.json: Evaluation configuration\n") f.write(f" - test_cases.txt: Detailed results on sample pairs\n") f.write(f" - reward_progression.png: Reward vs time for multiple pairs\n") f.write(f" - summary.txt: This file\n") # Note: errors.txt is generated by run_test_cases if errors occur f.write("\n") f.write("INTERPRETATION (Test Cases)\n") f.write("-" * 80 + "\n") if test_results: test_acc = sum(r["correct"] for r in test_results) / len(test_results) test_margin = np.mean([r["margin"] for r in test_results]) f.write(f"✓ Test Accuracy: {test_acc:.1%} - ") if test_acc > 0.70: f.write("Excellent! Model distinguishes preferences well.\n") elif test_acc > 0.60: f.write("Good. Model learns preferences (expected with ~70% noisy labels).\n") elif test_acc > 0.55: f.write("Fair. Model shows some learning.\n") else: f.write("Poor. Model may need more training or debugging.\n") f.write(f"✓ Test Mean Margin: {test_margin:.4f} - ") if test_margin > 1.0: f.write("Strong. Model is confident in preferences.\n") elif test_margin > 0.5: f.write("Good. Model shows clear preference signal.\n") elif test_margin > 0.2: f.write("Moderate. Model shows weak but positive signal.\n") else: f.write("Weak. Model barely distinguishes preferences.\n") f.write("\nNOTE: Full validation is disabled. Test cases provide quick sanity check.\n") print(f"✓ Summary saved to {summary_path}") def main(): """Main evaluation function.""" parser = argparse.ArgumentParser( description="Evaluate trained reward model", 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_val.bin, meta_val.jsonl, etc.)", ) parser.add_argument( "--output_dir", default="./reward_eval_results", help="Base directory for saving results (timestamped subdirectory will be created)", ) parser.add_argument( "--n_test_cases", type=int, default=10, help="Number of test case pairs to evaluate" ) parser.add_argument( "--n_viz_pairs", type=int, default=5, help="Number of pairs to visualize in progression plot" ) parser.add_argument( "--token_interval", type=int, default=30, help="Sample reward every N tokens (30 tokens ≈ 1.2s at 25Hz)", ) parser.add_argument( "--max_val_samples", type=int, default=None, help="Maximum validation samples to evaluate (None = all)", ) parser.add_argument("--device", default="cuda", help="Device for computation (cuda or cpu)") args = parser.parse_args() # Create timestamped output directory timestamp = datetime.now().strftime("%Y-%m-%d_%H-%M-%S") output_dir = os.path.join(args.output_dir, timestamp) os.makedirs(output_dir, exist_ok=True) print("=" * 80) print("REWARD MODEL EVALUATION") print("=" * 80) print(f"Output directory: {output_dir}") print(f"Checkpoint: {args.checkpoint}") print(f"Data directory: {args.data_dir}") print("") # Load model model, model_args = load_reward_model(args.checkpoint, args.device) cfg = model.config # Get tokenizer path from data_dir tokenizer_fp = os.path.join(args.data_dir, "tokenizer_60k.json") if not os.path.exists(tokenizer_fp): print(f"Warning: Tokenizer not found at {tokenizer_fp}") tokenizer_fp = None print("\n" + "=" * 80) print("Loading validation data...") print("=" * 80) # Create data_sampling_info EXACTLY like train_reward_model.py data_sampling_info = create_data_sampling_info(args.data_dir, "val", cfg, tokenizer_fp, args.device) print(f"✓ Loaded {len(data_sampling_info['val']['metas'])} validation samples") print(f"✓ Datasets: {data_sampling_info['val']['names']}") # Save configuration config_info = { "checkpoint": args.checkpoint, "data_dir": args.data_dir, "timestamp": timestamp, "n_test_cases": args.n_test_cases, "n_viz_pairs": args.n_viz_pairs, "token_interval": args.token_interval, "max_val_samples": args.max_val_samples, "device": args.device, "model_config": { k: str(v) if not isinstance(v, (int, float, bool, str, type(None))) else v for k, v in model_args.items() }, } with open(os.path.join(output_dir, "config.json"), "w") as f: json.dump(config_info, f, indent=2) # Run evaluations print("\n" + "=" * 80) print("1. Testing on sample cases...") print("=" * 80) test_results = run_test_cases(model, data_sampling_info, output_dir, args.n_test_cases, args.device) print("\n" + "=" * 80) print("2. Visualizing reward progression...") print("=" * 80) visualize_reward_progression( model, data_sampling_info, output_dir, args.n_viz_pairs, args.token_interval, args.device ) # Skip full validation by default (can be slow and error-prone) # Uncomment if you want full validation statistics print("\n" + "=" * 80) print("3. Running full validation...") print("=" * 80) val_stats = run_full_validation( model, data_sampling_info, output_dir, args.max_val_samples, args.device ) # val_stats = {} # Empty stats if not running full validation # Create summary print("\n" + "=" * 80) print("3. Creating summary...") print("=" * 80) create_summary(output_dir, test_results, val_stats, args.checkpoint) print("\n" + "=" * 80) print("✅ EVALUATION COMPLETE!") print("=" * 80) print(f"📁 Results saved to: {output_dir}") print("") print("Generated files:") for fname in sorted(os.listdir(output_dir)): fpath = os.path.join(output_dir, fname) if os.path.isfile(fpath): size = os.path.getsize(fpath) print(f" - {fname} ({size:,} bytes)") print("=" * 80) if __name__ == "__main__": main()