"""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,
    )
