"""
Odds-Ratio Preference Optimization (ORPO) for MK3

Implementation of ORPO from "ORPO: Monolithic Preference Optimization without
Reference Model" (Hong et al., 2024)

ORPO combines standard language modeling with preference learning in a single objective,
without requiring a reference model.

Key insight: Use odds ratio between chosen and rejected to create preference signal:
    OR(y_w, y_l | x) = P(y_w|x) / P(y_l|x)

Loss = L_SFT + λ * L_OR
where:
- L_SFT: Standard supervised fine-tuning loss on chosen responses
- L_OR: Odds ratio preference loss
"""

import torch
import torch.nn as nn
import torch.nn.functional as F
from typing import Optional, Dict, Tuple
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 ORPOLoss(nn.Module):
    """
    Odds-Ratio Preference Optimization loss.

    L_ORPO = L_SFT + λ * L_OR

    L_SFT = -log P(y_w|x)  [standard language modeling on chosen]
    L_OR = -log σ(log(P(y_w|x) / P(y_l|x)))
         = -log σ(log P(y_w|x) - log P(y_l|x))

    This combines supervised learning with preference optimization without
    needing a reference model.
    """

    def __init__(
        self,
        lambda_or: float = 0.1,
        use_log_odds: bool = True,
        sft_weight: float = 1.0,
    ):
        """
        Initialize ORPO loss.

        Args:
            lambda_or: Weight for odds ratio loss
            use_log_odds: Use log odds ratio (more stable)
            sft_weight: Weight for SFT loss
        """
        super().__init__()
        self.lambda_or = lambda_or
        self.use_log_odds = use_log_odds
        self.sft_weight = sft_weight

    def forward(
        self,
        chosen_logprobs: torch.Tensor,
        rejected_logprobs: torch.Tensor,
    ) -> Tuple[torch.Tensor, Dict[str, float]]:
        """
        Compute ORPO loss.

        Args:
            chosen_logprobs: Log P(y_w|x) [batch]
            rejected_logprobs: Log P(y_l|x) [batch]

        Returns:
            loss: ORPO loss
            metrics: Dictionary with metrics
        """
        # SFT loss: maximize log probability of chosen
        # We want to maximize log P(y_w|x), so minimize -log P(y_w|x)
        sft_loss = -chosen_logprobs

        # Odds ratio loss
        if self.use_log_odds:
            # Log odds ratio: log(P(y_w|x) / P(y_l|x)) = log P(y_w|x) - log P(y_l|x)
            log_odds = chosen_logprobs - rejected_logprobs

            # OR loss: -log σ(log odds)
            # We want log odds > 0 (chosen more likely than rejected)
            or_loss = -F.logsigmoid(log_odds)
        else:
            # Direct odds ratio (less stable with log probabilities)
            odds = torch.exp(chosen_logprobs - rejected_logprobs)
            or_loss = -torch.log(torch.sigmoid(odds))

        # Combined loss
        total_loss = self.sft_weight * sft_loss + self.lambda_or * or_loss

        # Metrics
        with torch.no_grad():
            log_odds = chosen_logprobs - rejected_logprobs

            # Accuracy: how often is chosen preferred
            accuracy = (log_odds > 0).float()

            # Implicit reward difference
            reward_margin = log_odds

        metrics = {
            'loss': total_loss.mean().item(),
            'sft_loss': sft_loss.mean().item(),
            'or_loss': or_loss.mean().item(),
            'accuracy': accuracy.mean().item(),
            'log_odds': log_odds.mean().item(),
            'reward_margin': reward_margin.mean().item(),
            'chosen_logprob': chosen_logprobs.mean().item(),
            'rejected_logprob': rejected_logprobs.mean().item(),
        }

        return total_loss.mean(), metrics


class OddsRatioPreferenceOptimization:
    """
    Complete ORPO training wrapper for MK3 continuous models.

    Advantage: No reference model needed, combines SFT with preference learning.
    """

    def __init__(
        self,
        policy_model,
        lambda_or: float = 0.1,
        use_log_odds: bool = True,
        sft_weight: float = 1.0,
    ):
        """
        Initialize ORPO trainer.

        Args:
            policy_model: Model to train
            lambda_or: Weight for odds ratio loss
            use_log_odds: Use log odds ratio
            sft_weight: Weight for SFT loss
        """
        self.policy_model = policy_model

        self.loss_fn = ORPOLoss(
            lambda_or=lambda_or,
            use_log_odds=use_log_odds,
            sft_weight=sft_weight
        )

    def compute_loss(
        self,
        prompt_vectors: torch.Tensor,
        chosen_vectors: torch.Tensor,
        rejected_vectors: torch.Tensor,
    ) -> Tuple[torch.Tensor, Dict[str, float]]:
        """
        Compute ORPO loss for a batch of preferences.

        Args:
            prompt_vectors: [batch, num_prompt_chunks, vector_dim]
            chosen_vectors: [batch, num_chosen_chunks, vector_dim]
            rejected_vectors: [batch, num_rejected_chunks, vector_dim]

        Returns:
            loss: ORPO loss
            metrics: Dictionary with metrics
        """
        # Compute log probabilities
        chosen_logprobs = compute_log_prob(
            self.policy_model,
            prompt_vectors,
            chosen_vectors,
            reduction='mean'
        )

        rejected_logprobs = compute_log_prob(
            self.policy_model,
            prompt_vectors,
            rejected_vectors,
            reduction='mean'
        )

        # Compute ORPO loss
        loss, metrics = self.loss_fn(chosen_logprobs, rejected_logprobs)

        return loss, metrics

    def train_step(
        self,
        prompt_tokens: torch.Tensor,
        chosen_tokens: torch.Tensor,
        rejected_tokens: torch.Tensor,
        optimizer: torch.optim.Optimizer,
    ) -> Dict[str, float]:
        """
        Single ORPO training step.

        Args:
            prompt_tokens: [batch, prompt_len] prompt token IDs
            chosen_tokens: [batch, chosen_len] chosen response token IDs
            rejected_tokens: [batch, rejected_len] rejected response token IDs
            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)
        chosen_vectors, _ = self.policy_model.tokenize_to_vectors(chosen_tokens)
        rejected_vectors, _ = self.policy_model.tokenize_to_vectors(rejected_tokens)

        # Compute loss
        loss, metrics = self.compute_loss(
            prompt_vectors,
            chosen_vectors,
            rejected_vectors
        )

        # Backward pass
        optimizer.zero_grad()
        loss.backward()
        optimizer.step()

        return metrics

    def compute_sft_loss(
        self,
        prompt_tokens: torch.Tensor,
        target_tokens: torch.Tensor,
    ) -> Tuple[torch.Tensor, Dict[str, float]]:
        """
        Compute standard SFT loss on supervised data.

        Can be used to pre-train before preference optimization.

        Args:
            prompt_tokens: [batch, prompt_len] prompt token IDs
            target_tokens: [batch, target_len] target token IDs

        Returns:
            loss: SFT loss
            metrics: Dictionary with metrics
        """
        # Convert to vectors
        prompt_vectors, _ = self.policy_model.tokenize_to_vectors(prompt_tokens)
        target_vectors, _ = self.policy_model.tokenize_to_vectors(target_tokens)

        # Compute log probability
        logprobs = compute_log_prob(
            self.policy_model,
            prompt_vectors,
            target_vectors,
            reduction='mean'
        )

        # SFT loss: -log P(y|x)
        loss = -logprobs.mean()

        metrics = {
            'sft_loss': loss.item(),
            'logprob': logprobs.mean().item(),
        }

        return loss, metrics
