"""
Direct Preference Optimization (DPO) for MK3

Implementation of DPO from "Direct Preference Optimization: Your Language Model is
Secretly a Reward Model" (Rafailov et al., 2023)

DPO optimizes models directly on preference data without needing a separate reward model.

Key insight: The optimal policy under reward model r(x,y) can be written as:
    π*(y|x) ∝ π_ref(y|x) * exp(r(x,y) / β)

This allows direct optimization: maximize log π(y_w|x) - log π(y_l|x) where
y_w is preferred over y_l.
"""

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, estimate_kl_divergence


class DPOLoss(nn.Module):
    """
    Direct Preference Optimization loss.

    Loss = -log σ(β * (log π(y_w|x) - log π(y_l|x) - log π_ref(y_w|x) + log π_ref(y_l|x)))

    where:
    - y_w: chosen/preferred response
    - y_l: rejected/dispreferred response
    - π: policy model
    - π_ref: reference model
    - β: temperature parameter (controls deviation from reference)
    - σ: sigmoid function
    """

    def __init__(
        self,
        beta: float = 0.1,
        label_smoothing: float = 0.0,
        use_ipo: bool = False,
        reference_free: bool = False
    ):
        """
        Initialize DPO loss.

        Args:
            beta: Temperature parameter (higher = stay closer to reference)
            label_smoothing: Label smoothing for robustness
            use_ipo: Use IPO variant (Implicit Preference Optimization)
            reference_free: Skip reference model (use prior initialization)
        """
        super().__init__()
        self.beta = beta
        self.label_smoothing = label_smoothing
        self.use_ipo = use_ipo
        self.reference_free = reference_free

    def forward(
        self,
        policy_chosen_logprobs: torch.Tensor,
        policy_rejected_logprobs: torch.Tensor,
        reference_chosen_logprobs: Optional[torch.Tensor] = None,
        reference_rejected_logprobs: Optional[torch.Tensor] = None,
    ) -> Tuple[torch.Tensor, Dict[str, float]]:
        """
        Compute DPO loss.

        Args:
            policy_chosen_logprobs: Log P(y_w|x) under policy [batch]
            policy_rejected_logprobs: Log P(y_l|x) under policy [batch]
            reference_chosen_logprobs: Log P(y_w|x) under reference [batch]
            reference_rejected_logprobs: Log P(y_l|x) under reference [batch]

        Returns:
            loss: DPO loss
            metrics: Dictionary with metrics
        """
        if self.reference_free:
            # Reference-free variant: assume reference gives uniform probability
            reference_chosen_logprobs = torch.zeros_like(policy_chosen_logprobs)
            reference_rejected_logprobs = torch.zeros_like(policy_rejected_logprobs)
        else:
            if reference_chosen_logprobs is None or reference_rejected_logprobs is None:
                raise ValueError("Reference logprobs required when reference_free=False")

        # Compute preference logits
        # logit = β * (log π(y_w|x) - log π(y_l|x) - log π_ref(y_w|x) + log π_ref(y_l|x))
        policy_logratios = policy_chosen_logprobs - policy_rejected_logprobs
        reference_logratios = reference_chosen_logprobs - reference_rejected_logprobs
        logits = self.beta * (policy_logratios - reference_logratios)

        if self.use_ipo:
            # IPO loss: (logit - 1/(2*beta))^2
            # More stable than DPO, less sensitive to outliers
            loss = (logits - 1.0 / (2.0 * self.beta)) ** 2
        else:
            # Standard DPO loss: -log σ(logit)
            if self.label_smoothing > 0:
                # Label smoothing for robustness
                # y = (1 - ε) * 1 + ε * 0 = 1 - ε
                labels = 1.0 - self.label_smoothing
                loss = -F.logsigmoid(logits) * labels - F.logsigmoid(-logits) * (1 - labels)
            else:
                loss = -F.logsigmoid(logits)

        # Metrics
        with torch.no_grad():
            # Implicit reward: r(x,y) = β * log(π(y|x) / π_ref(y|x))
            chosen_rewards = self.beta * (policy_chosen_logprobs - reference_chosen_logprobs)
            rejected_rewards = self.beta * (policy_rejected_logprobs - reference_rejected_logprobs)
            reward_margin = chosen_rewards - rejected_rewards

            # Accuracy: how often is chosen preferred over rejected
            accuracy = (logits > 0).float()

        metrics = {
            'loss': loss.mean().item(),
            'accuracy': accuracy.mean().item(),
            'reward_margin': reward_margin.mean().item(),
            'chosen_reward': chosen_rewards.mean().item(),
            'rejected_reward': rejected_rewards.mean().item(),
            'policy_logratios': policy_logratios.mean().item(),
            'reference_logratios': reference_logratios.mean().item(),
        }

        return loss.mean(), metrics


class DirectPreferenceOptimization:
    """
    Complete DPO training wrapper for MK3 continuous models.

    Handles:
    - Reference model management
    - Log probability computation in continuous space
    - Preference data batching
    """

    def __init__(
        self,
        policy_model,
        reference_model: Optional[nn.Module] = None,
        beta: float = 0.1,
        label_smoothing: float = 0.0,
        use_ipo: bool = False,
        reference_free: bool = False,
    ):
        """
        Initialize DPO trainer.

        Args:
            policy_model: Model to train
            reference_model: Reference model (if None, use copy of policy)
            beta: Temperature parameter
            label_smoothing: Label smoothing
            use_ipo: Use IPO variant
            reference_free: Skip reference model
        """
        self.policy_model = policy_model
        self.reference_free = reference_free

        if reference_free:
            self.reference_model = None
        else:
            if reference_model is None:
                # Create frozen copy of policy as reference
                import copy
                self.reference_model = copy.deepcopy(policy_model)
                # Freeze reference
                for param in self.reference_model.parameters():
                    param.requires_grad = False
                self.reference_model.eval()
            else:
                self.reference_model = reference_model

        self.loss_fn = DPOLoss(
            beta=beta,
            label_smoothing=label_smoothing,
            use_ipo=use_ipo,
            reference_free=reference_free
        )

    def compute_loss(
        self,
        prompt_vectors: torch.Tensor,
        chosen_vectors: torch.Tensor,
        rejected_vectors: torch.Tensor,
    ) -> Tuple[torch.Tensor, Dict[str, float]]:
        """
        Compute DPO 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: DPO loss
            metrics: Dictionary with metrics
        """
        # Compute policy log probabilities
        policy_chosen_logprobs = compute_log_prob(
            self.policy_model,
            prompt_vectors,
            chosen_vectors,
            reduction='mean'
        )

        policy_rejected_logprobs = compute_log_prob(
            self.policy_model,
            prompt_vectors,
            rejected_vectors,
            reduction='mean'
        )

        # Compute reference log probabilities
        if self.reference_free:
            reference_chosen_logprobs = None
            reference_rejected_logprobs = None
        else:
            with torch.no_grad():
                reference_chosen_logprobs = compute_log_prob(
                    self.reference_model,
                    prompt_vectors,
                    chosen_vectors,
                    reduction='mean'
                )

                reference_rejected_logprobs = compute_log_prob(
                    self.reference_model,
                    prompt_vectors,
                    rejected_vectors,
                    reduction='mean'
                )

        # Compute DPO loss
        loss, metrics = self.loss_fn(
            policy_chosen_logprobs,
            policy_rejected_logprobs,
            reference_chosen_logprobs,
            reference_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 DPO 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 update_reference_model(self):
        """
        Update reference model to current policy (for online DPO).
        """
        if self.reference_model is not None:
            self.reference_model.load_state_dict(self.policy_model.state_dict())
            print("Reference model updated to current policy")
