"""
Rank Responses to align with Human Feedback (RRHF) for MK3

Implementation based on "RRHF: Rank Responses to align with Human Feedback"
(Yuan et al., 2023)

RRHF extends pairwise preferences to handle rankings of multiple responses.
Instead of just comparing two responses, it learns from full rankings.

Key insight: For K responses with ranking r_1 < r_2 < ... < r_K (lower is better),
optimize ranking loss that ensures better-ranked responses have higher probability.

Uses ListMLE (List Maximum Likelihood Estimation) or ranking-based contrastive loss.
"""

import torch
import torch.nn as nn
import torch.nn.functional as F
from typing import Optional, Dict, Tuple, List
import sys
import os

sys.path.append(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))

from .utils import compute_log_prob, compute_sequence_log_prob


class RRHFLoss(nn.Module):
    """
    Rank Responses Human Feedback loss.

    Supports multiple loss formulations:
    1. ListMLE: Maximum likelihood estimation for rankings
    2. Pairwise ranking: All pairwise comparisons in ranking
    3. Top-k ranking: Focus on top-k responses

    ListMLE: L = -Σ log(P(π_i | π_{i+1:K}))
    where π is the ranked permutation of responses.
    """

    def __init__(
        self,
        loss_type: str = 'listmle',
        temperature: float = 1.0,
        margin: float = 0.0,
        top_k: Optional[int] = None,
    ):
        """
        Initialize RRHF loss.

        Args:
            loss_type: 'listmle', 'pairwise', or 'topk'
            temperature: Temperature for softmax
            margin: Margin for pairwise ranking loss
            top_k: Number of top responses to focus on (for topk loss)
        """
        super().__init__()
        self.loss_type = loss_type.lower()
        self.temperature = temperature
        self.margin = margin
        self.top_k = top_k

        if self.loss_type not in ['listmle', 'pairwise', 'topk']:
            raise ValueError(f"Unknown loss_type: {loss_type}")

    def listmle_loss(
        self,
        logprobs: torch.Tensor,
        rankings: torch.Tensor,
    ) -> torch.Tensor:
        """
        ListMLE (List Maximum Likelihood Estimation) loss.

        For ranking r_1 < r_2 < ... < r_K:
        L = -Σ_i log(exp(s_i) / Σ_{j≥i} exp(s_j))

        where s_i are scores (log probabilities).

        Args:
            logprobs: [batch, num_responses] log probabilities
            rankings: [batch, num_responses] rankings (0 = best)

        Returns:
            loss: ListMLE loss
        """
        batch_size, num_responses = logprobs.shape

        # Sort by rankings to get order
        sorted_indices = torch.argsort(rankings, dim=1)

        # Gather logprobs in ranked order
        batch_indices = torch.arange(batch_size, device=logprobs.device).unsqueeze(1)
        sorted_logprobs = logprobs[batch_indices, sorted_indices]

        # Compute ListMLE loss
        # For each position i, compute: log(exp(s_i) / Σ_{j≥i} exp(s_j))
        loss = 0.0
        for i in range(num_responses - 1):
            # Current score
            current_score = sorted_logprobs[:, i] / self.temperature

            # Remaining scores (including current)
            remaining_scores = sorted_logprobs[:, i:] / self.temperature

            # Log-sum-exp for numerical stability
            log_sum_exp = torch.logsumexp(remaining_scores, dim=1)

            # ListMLE term: log(exp(s_i) / Σ_{j≥i} exp(s_j))
            #             = s_i - log(Σ_{j≥i} exp(s_j))
            loss = loss - (current_score - log_sum_exp).mean()

        return loss / (num_responses - 1)

    def pairwise_loss(
        self,
        logprobs: torch.Tensor,
        rankings: torch.Tensor,
    ) -> torch.Tensor:
        """
        Pairwise ranking loss.

        For all pairs (i, j) where rank_i < rank_j (i is better):
        L = max(0, margin - (logprob_i - logprob_j))

        Args:
            logprobs: [batch, num_responses] log probabilities
            rankings: [batch, num_responses] rankings (0 = best)

        Returns:
            loss: Pairwise ranking loss
        """
        batch_size, num_responses = logprobs.shape

        # Create all pairwise comparisons
        # i is better than j if ranking[i] < ranking[j]
        loss = 0.0
        num_pairs = 0

        for i in range(num_responses):
            for j in range(i + 1, num_responses):
                # Check which is better in each batch
                i_better = rankings[:, i] < rankings[:, j]
                j_better = rankings[:, j] < rankings[:, i]

                # Margin loss for i better than j
                if i_better.any():
                    diff_ij = logprobs[:, i] - logprobs[:, j]
                    loss_ij = torch.clamp(self.margin - diff_ij, min=0)
                    loss = loss + (loss_ij * i_better.float()).sum()
                    num_pairs += i_better.sum()

                # Margin loss for j better than i
                if j_better.any():
                    diff_ji = logprobs[:, j] - logprobs[:, i]
                    loss_ji = torch.clamp(self.margin - diff_ji, min=0)
                    loss = loss + (loss_ji * j_better.float()).sum()
                    num_pairs += j_better.sum()

        if num_pairs > 0:
            loss = loss / num_pairs
        else:
            loss = torch.tensor(0.0, device=logprobs.device)

        return loss

    def topk_loss(
        self,
        logprobs: torch.Tensor,
        rankings: torch.Tensor,
    ) -> torch.Tensor:
        """
        Top-k ranking loss.

        Focus on top-k responses in the ranking.
        Use ListMLE on top-k only.

        Args:
            logprobs: [batch, num_responses] log probabilities
            rankings: [batch, num_responses] rankings (0 = best)

        Returns:
            loss: Top-k ranking loss
        """
        batch_size, num_responses = logprobs.shape

        if self.top_k is None:
            k = num_responses
        else:
            k = min(self.top_k, num_responses)

        # Get top-k by ranking
        topk_indices = torch.argsort(rankings, dim=1)[:, :k]

        # Gather top-k logprobs and rankings
        batch_indices = torch.arange(batch_size, device=logprobs.device).unsqueeze(1)
        topk_logprobs = logprobs[batch_indices, topk_indices]
        topk_rankings = rankings[batch_indices, topk_indices]

        # Apply ListMLE on top-k
        return self.listmle_loss(topk_logprobs, topk_rankings)

    def forward(
        self,
        logprobs: torch.Tensor,
        rankings: torch.Tensor,
    ) -> Tuple[torch.Tensor, Dict[str, float]]:
        """
        Compute RRHF loss.

        Args:
            logprobs: [batch, num_responses] log probabilities
            rankings: [batch, num_responses] rankings (0 = best, lower is better)

        Returns:
            loss: RRHF loss
            metrics: Dictionary with metrics
        """
        if self.loss_type == 'listmle':
            loss = self.listmle_loss(logprobs, rankings)
        elif self.loss_type == 'pairwise':
            loss = self.pairwise_loss(logprobs, rankings)
        elif self.loss_type == 'topk':
            loss = self.topk_loss(logprobs, rankings)
        else:
            raise ValueError(f"Unknown loss_type: {self.loss_type}")

        # Metrics
        with torch.no_grad():
            # Rank correlation: Kendall's tau approximation
            # Check if predicted ranking matches true ranking
            pred_rankings = torch.argsort(torch.argsort(-logprobs, dim=1), dim=1)
            rank_accuracy = (pred_rankings == rankings).float().mean()

            # Top-1 accuracy: is highest prob response the best ranked?
            best_pred = torch.argmax(logprobs, dim=1)
            best_true = torch.argmin(rankings, dim=1)
            top1_accuracy = (best_pred == best_true).float().mean()

        metrics = {
            'loss': loss.item(),
            'rank_accuracy': rank_accuracy.item(),
            'top1_accuracy': top1_accuracy.item(),
            'mean_logprob': logprobs.mean().item(),
        }

        return loss, metrics


class RankResponsesHumanFeedback:
    """
    Complete RRHF training wrapper for MK3 continuous models.

    Handles ranking of multiple responses per prompt.
    """

    def __init__(
        self,
        policy_model,
        loss_type: str = 'listmle',
        temperature: float = 1.0,
        margin: float = 0.0,
        top_k: Optional[int] = None,
    ):
        """
        Initialize RRHF trainer.

        Args:
            policy_model: Model to train
            loss_type: 'listmle', 'pairwise', or 'topk'
            temperature: Temperature for softmax
            margin: Margin for pairwise loss
            top_k: Number of top responses (for topk loss)
        """
        self.policy_model = policy_model

        self.loss_fn = RRHFLoss(
            loss_type=loss_type,
            temperature=temperature,
            margin=margin,
            top_k=top_k
        )

    def compute_loss(
        self,
        prompt_vectors: torch.Tensor,
        response_vectors_list: List[torch.Tensor],
        rankings: torch.Tensor,
    ) -> Tuple[torch.Tensor, Dict[str, float]]:
        """
        Compute RRHF loss for ranked responses.

        Args:
            prompt_vectors: [batch, num_prompt_chunks, vector_dim]
            response_vectors_list: List of [batch, num_chunks, vector_dim], one per response
            rankings: [batch, num_responses] rankings (0 = best)

        Returns:
            loss: RRHF loss
            metrics: Dictionary with metrics
        """
        batch_size = prompt_vectors.shape[0]
        num_responses = len(response_vectors_list)

        # Compute log probability for each response
        logprobs_list = []
        for response_vectors in response_vectors_list:
            logprobs = compute_log_prob(
                self.policy_model,
                prompt_vectors,
                response_vectors,
                reduction='mean'
            )
            logprobs_list.append(logprobs)

        # Stack: [batch, num_responses]
        logprobs = torch.stack(logprobs_list, dim=1)

        # Compute RRHF loss
        loss, metrics = self.loss_fn(logprobs, rankings)

        return loss, metrics

    def train_step(
        self,
        prompt_tokens: torch.Tensor,
        response_tokens_list: List[torch.Tensor],
        rankings: torch.Tensor,
        optimizer: torch.optim.Optimizer,
    ) -> Dict[str, float]:
        """
        Single RRHF training step.

        Args:
            prompt_tokens: [batch, prompt_len] prompt token IDs
            response_tokens_list: List of [batch, response_len] response token IDs
            rankings: [batch, num_responses] rankings
            optimizer: Optimizer

        Returns:
            metrics: Training metrics
        """
        self.policy_model.train()

        # Convert tokens to continuous vectors
        prompt_vectors, _ = self.policy_model.tokenize_to_vectors(prompt_tokens)

        response_vectors_list = []
        for response_tokens in response_tokens_list:
            response_vectors, _ = self.policy_model.tokenize_to_vectors(response_tokens)
            response_vectors_list.append(response_vectors)

        # Compute loss
        loss, metrics = self.compute_loss(
            prompt_vectors,
            response_vectors_list,
            rankings
        )

        # Backward pass
        optimizer.zero_grad()
        loss.backward()
        optimizer.step()

        return metrics
