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