#!/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()
