"""
Kahneman-Tversky Optimization (KTO) for MK3

Implementation of KTO from "KTO: Model Alignment as Prospect Theoretic Optimization"
(Ethayarajh et al., 2024)

KTO uses prospect theory to model human preferences:
- Loss aversion: losses loom larger than gains
- Reference dependence: outcomes evaluated relative to reference point
- Diminishing sensitivity: marginal impact decreases with magnitude

Unlike DPO, KTO can use binary feedback (good/bad) without pairwise comparisons.

Key insight: Model utility using value function from prospect theory:
    v(x) = x^α if x ≥ 0 else -λ * (-x)^β
where λ > 1 captures loss aversion.
"""

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 KTOLoss(nn.Module):
    """
    Kahneman-Tversky Optimization loss.

    Uses prospect theory value function to weight preferences:
    - Gains: v(x) = x^α
    - Losses: v(x) = -λ * (-x)^β

    Loss for desirable output y_d:
        L_D = -log σ(β * (log π(y_d|x) - log π_ref(y_d|x)))

    Loss for undesirable output y_u:
        L_U = -log σ(β * (log π_ref(y_u|x) - log π(y_u|x)))
    """

    def __init__(
        self,
        beta: float = 0.1,
        lambda_loss_aversion: float = 2.0,
        alpha: float = 1.0,
        beta_sensitivity: float = 1.0,
        desirable_weight: float = 1.0,
        undesirable_weight: float = 1.0,
    ):
        """
        Initialize KTO loss.

        Args:
            beta: Temperature parameter
            lambda_loss_aversion: Loss aversion coefficient (λ > 1)
            alpha: Gain sensitivity exponent
            beta_sensitivity: Loss sensitivity exponent
            desirable_weight: Weight for desirable outputs
            undesirable_weight: Weight for undesirable outputs
        """
        super().__init__()
        self.beta = beta
        self.lambda_loss_aversion = lambda_loss_aversion
        self.alpha = alpha
        self.beta_sensitivity = beta_sensitivity
        self.desirable_weight = desirable_weight
        self.undesirable_weight = undesirable_weight

    def value_function(self, x: torch.Tensor, is_gain: bool) -> torch.Tensor:
        """
        Prospect theory value function.

        Args:
            x: Input values
            is_gain: Whether this represents a gain (True) or loss (False)

        Returns:
            v(x): Value function output
        """
        if is_gain:
            # Gains: v(x) = x^α
            return torch.pow(torch.clamp(x, min=0), self.alpha)
        else:
            # Losses: v(x) = -λ * (-x)^β
            return -self.lambda_loss_aversion * torch.pow(
                torch.clamp(-x, min=0),
                self.beta_sensitivity
            )

    def forward(
        self,
        policy_logprobs: torch.Tensor,
        reference_logprobs: torch.Tensor,
        is_desirable: torch.Tensor,
    ) -> Tuple[torch.Tensor, Dict[str, float]]:
        """
        Compute KTO loss.

        Args:
            policy_logprobs: Log P(y|x) under policy [batch]
            reference_logprobs: Log P(y|x) under reference [batch]
            is_desirable: Boolean mask [batch], True for desirable outputs

        Returns:
            loss: KTO loss
            metrics: Dictionary with metrics
        """
        # Compute KL divergence: KL(π || π_ref)
        kl = policy_logprobs - reference_logprobs

        # Apply prospect theory value function
        # For desirable outputs: want to increase KL (move away from reference)
        # For undesirable outputs: want to decrease KL (move toward reference)

        # Desirable: maximize log σ(β * KL)
        # Undesirable: maximize log σ(-β * KL)
        logits = torch.where(
            is_desirable,
            self.beta * kl,
            -self.beta * kl
        )

        # Base loss: -log σ(logit)
        base_loss = -F.logsigmoid(logits)

        # Apply prospect theory weighting
        # Desirable outputs are "gains", undesirable are "losses"
        weighted_loss = torch.where(
            is_desirable,
            self.desirable_weight * base_loss,
            self.undesirable_weight * base_loss
        )

        # Apply value function for prospect theory
        # This models diminishing sensitivity and loss aversion
        value_weighted_loss = torch.where(
            is_desirable,
            self.value_function(base_loss, is_gain=True),
            self.value_function(base_loss, is_gain=False)
        )

        # Final loss combines weighted and value-weighted components
        loss = 0.5 * weighted_loss + 0.5 * torch.abs(value_weighted_loss)

        # Metrics
        with torch.no_grad():
            # Implicit reward
            rewards = self.beta * kl

            # Separate metrics for desirable and undesirable
            desirable_mask = is_desirable.bool()
            undesirable_mask = ~desirable_mask

            if desirable_mask.any():
                desirable_reward = rewards[desirable_mask].mean()
                desirable_loss = loss[desirable_mask].mean()
            else:
                desirable_reward = torch.tensor(0.0)
                desirable_loss = torch.tensor(0.0)

            if undesirable_mask.any():
                undesirable_reward = rewards[undesirable_mask].mean()
                undesirable_loss = loss[undesirable_mask].mean()
            else:
                undesirable_reward = torch.tensor(0.0)
                undesirable_loss = torch.tensor(0.0)

            # Accuracy: desirable should have positive logits, undesirable negative
            accuracy = torch.where(
                is_desirable,
                (logits > 0).float(),
                (logits < 0).float()
            )

        metrics = {
            'loss': loss.mean().item(),
            'accuracy': accuracy.mean().item(),
            'desirable_reward': desirable_reward.item(),
            'undesirable_reward': undesirable_reward.item(),
            'desirable_loss': desirable_loss.item(),
            'undesirable_loss': undesirable_loss.item(),
            'kl_divergence': kl.mean().item(),
        }

        return loss.mean(), metrics


class KahnemanTverskyOptimization:
    """
    Complete KTO training wrapper for MK3 continuous models.

    Supports binary feedback (good/bad) without requiring pairwise preferences.
    """

    def __init__(
        self,
        policy_model,
        reference_model: Optional[nn.Module] = None,
        beta: float = 0.1,
        lambda_loss_aversion: float = 2.0,
        alpha: float = 1.0,
        beta_sensitivity: float = 1.0,
        desirable_weight: float = 1.0,
        undesirable_weight: float = 1.0,
    ):
        """
        Initialize KTO trainer.

        Args:
            policy_model: Model to train
            reference_model: Reference model (if None, use copy of policy)
            beta: Temperature parameter
            lambda_loss_aversion: Loss aversion coefficient
            alpha: Gain sensitivity exponent
            beta_sensitivity: Loss sensitivity exponent
            desirable_weight: Weight for desirable outputs
            undesirable_weight: Weight for undesirable outputs
        """
        self.policy_model = policy_model

        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 = KTOLoss(
            beta=beta,
            lambda_loss_aversion=lambda_loss_aversion,
            alpha=alpha,
            beta_sensitivity=beta_sensitivity,
            desirable_weight=desirable_weight,
            undesirable_weight=undesirable_weight
        )

    def compute_loss(
        self,
        prompt_vectors: torch.Tensor,
        response_vectors: torch.Tensor,
        is_desirable: torch.Tensor,
    ) -> Tuple[torch.Tensor, Dict[str, float]]:
        """
        Compute KTO loss for a batch.

        Args:
            prompt_vectors: [batch, num_prompt_chunks, vector_dim]
            response_vectors: [batch, num_response_chunks, vector_dim]
            is_desirable: [batch] Boolean mask for desirable outputs

        Returns:
            loss: KTO loss
            metrics: Dictionary with metrics
        """
        # Compute policy log probabilities
        policy_logprobs = compute_log_prob(
            self.policy_model,
            prompt_vectors,
            response_vectors,
            reduction='mean'
        )

        # Compute reference log probabilities
        with torch.no_grad():
            reference_logprobs = compute_log_prob(
                self.reference_model,
                prompt_vectors,
                response_vectors,
                reduction='mean'
            )

        # Compute KTO loss
        loss, metrics = self.loss_fn(
            policy_logprobs,
            reference_logprobs,
            is_desirable
        )

        return loss, metrics

    def train_step(
        self,
        prompt_tokens: torch.Tensor,
        response_tokens: torch.Tensor,
        is_desirable: torch.Tensor,
        optimizer: torch.optim.Optimizer,
    ) -> Dict[str, float]:
        """
        Single KTO training step.

        Args:
            prompt_tokens: [batch, prompt_len] prompt token IDs
            response_tokens: [batch, response_len] response token IDs
            is_desirable: [batch] Boolean mask for desirable outputs
            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, _ = self.policy_model.tokenize_to_vectors(response_tokens)

        # Compute loss
        loss, metrics = self.compute_loss(
            prompt_vectors,
            response_vectors,
            is_desirable
        )

        # Backward pass
        optimizer.zero_grad()
        loss.backward()
        optimizer.step()

        return metrics

    def update_reference_model(self):
        """
        Update reference model to current policy.
        """
        self.reference_model.load_state_dict(self.policy_model.state_dict())
        print("Reference model updated to current policy")
