"""
Contrastive Decoding: Expert vs Amateur Models

Implements Contrastive Decoding to improve generation quality by contrasting
predictions from an expert model with an amateur (weaker) model.

The key insight: expert_logits - alpha * amateur_logits amplifies the expert's
advantages while suppressing common weaknesses.

Reference:
- Li et al. 2022: "Contrastive Decoding: Open-ended Text Generation as Optimization"
- O'Brien & Lewis 2023: "Contrastive Decoding Improves Reasoning in Large Language Models"
"""

import torch
import torch.nn.functional as F
from typing import Optional, Tuple, Dict, Callable, Union
from dataclasses import dataclass
import math


@dataclass
class ContrastiveConfig:
    """Configuration for contrastive decoding."""

    # Contrastive weight (alpha)
    alpha: float = 0.5  # Weight for amateur model logits

    # Plausibility constraint (beta)
    beta: float = 0.5  # Threshold for plausibility filtering

    # Adaptive masking
    use_adaptive_masking: bool = True
    adaptive_threshold: float = 0.1  # Relative threshold

    # Temperature for expert and amateur
    expert_temperature: float = 1.0
    amateur_temperature: float = 1.0

    # Whether to use log-space operations (more stable)
    use_log_space: bool = True


class ContrastiveDecoder:
    """
    Contrastive decoding using expert and amateur models.
    """

    def __init__(
        self,
        expert_model,
        amateur_model,
        config: Optional[ContrastiveConfig] = None
    ):
        """
        Initialize contrastive decoder.

        Args:
            expert_model: The expert (stronger) model
            amateur_model: The amateur (weaker) model
            config: Contrastive decoding configuration
        """
        self.expert_model = expert_model
        self.amateur_model = amateur_model
        self.config = config or ContrastiveConfig()

    def compute_contrastive_logits(
        self,
        expert_logits: torch.Tensor,
        amateur_logits: torch.Tensor,
        alpha: Optional[float] = None,
        beta: Optional[float] = None
    ) -> torch.Tensor:
        """
        Compute contrastive logits: expert - alpha * amateur.

        With plausibility constraint: only contrast when amateur is confident.

        Args:
            expert_logits: [batch_size, vocab_size] logits from expert model
            amateur_logits: [batch_size, vocab_size] logits from amateur model
            alpha: Contrastive weight (uses config default if None)
            beta: Plausibility threshold (uses config default if None)

        Returns:
            Contrastive logits [batch_size, vocab_size]
        """
        alpha = alpha if alpha is not None else self.config.alpha
        beta = beta if beta is not None else self.config.beta

        # Apply temperatures
        if self.config.expert_temperature != 1.0:
            expert_logits = expert_logits / self.config.expert_temperature
        if self.config.amateur_temperature != 1.0:
            amateur_logits = amateur_logits / self.config.amateur_temperature

        if self.config.use_log_space:
            # Work in log-probability space (more stable)
            expert_log_probs = F.log_softmax(expert_logits, dim=-1)
            amateur_log_probs = F.log_softmax(amateur_logits, dim=-1)

            # Plausibility constraint: only contrast where amateur is confident
            if self.config.use_adaptive_masking:
                # Adaptive masking based on amateur confidence
                amateur_probs = torch.exp(amateur_log_probs)
                plausibility_mask = amateur_probs > beta

                # Apply contrastive decoding only where plausible
                contrastive_log_probs = expert_log_probs.clone()
                contrastive_log_probs[plausibility_mask] -= alpha * amateur_log_probs[plausibility_mask]
            else:
                # Simple subtraction
                contrastive_log_probs = expert_log_probs - alpha * amateur_log_probs

            # Convert back to logits (for consistency with sampling interface)
            # We'll use the log probabilities directly, but scale them to logit range
            contrastive_logits = contrastive_log_probs

        else:
            # Work in probability space
            expert_probs = F.softmax(expert_logits, dim=-1)
            amateur_probs = F.softmax(amateur_logits, dim=-1)

            # Plausibility constraint
            if self.config.use_adaptive_masking:
                plausibility_mask = amateur_probs > beta
                contrastive_probs = expert_probs.clone()
                contrastive_probs[plausibility_mask] -= alpha * amateur_probs[plausibility_mask]
                # Ensure non-negative
                contrastive_probs = torch.clamp(contrastive_probs, min=1e-10)
                # Renormalize
                contrastive_probs = contrastive_probs / contrastive_probs.sum(dim=-1, keepdim=True)
            else:
                contrastive_probs = expert_probs - alpha * amateur_probs
                contrastive_probs = torch.clamp(contrastive_probs, min=1e-10)
                contrastive_probs = contrastive_probs / contrastive_probs.sum(dim=-1, keepdim=True)

            # Convert to logits
            contrastive_logits = torch.log(contrastive_probs)

        return contrastive_logits

    def generate_contrastive(
        self,
        initial_tokens: torch.Tensor,
        max_new_vectors: int = 10,
        temperature: float = 1.0,
        top_p: float = 0.9,
        alpha: Optional[float] = None,
        beta: Optional[float] = None,
        return_expert_only: bool = False,
        device: Optional[torch.device] = None
    ) -> Dict[str, torch.Tensor]:
        """
        Generate text using contrastive decoding.

        Args:
            initial_tokens: [batch_size, seq_len] initial tokens
            max_new_vectors: Number of vectors to generate
            temperature: Sampling temperature (applied after contrasting)
            top_p: Nucleus sampling threshold
            alpha: Contrastive weight
            beta: Plausibility threshold
            return_expert_only: If True, also return expert-only generation for comparison
            device: Device for computation

        Returns:
            Dictionary containing:
            - 'generated_tokens': Contrastively decoded tokens
            - 'expert_tokens': Expert-only tokens (if return_expert_only=True)
        """
        device = device or next(self.expert_model.parameters()).device
        alpha = alpha if alpha is not None else self.config.alpha
        beta = beta if beta is not None else self.config.beta

        self.expert_model.eval()
        self.amateur_model.eval()

        with torch.no_grad():
            # Convert initial tokens to continuous vectors for both models
            expert_vectors, _ = self.expert_model.tokenize_to_vectors(initial_tokens)
            amateur_vectors, _ = self.amateur_model.tokenize_to_vectors(initial_tokens)

            # Generate autoregressively
            for step in range(max_new_vectors):
                # Forward pass through both models
                expert_output = self.expert_model.forward(expert_vectors, use_cache=True)
                amateur_output = self.amateur_model.forward(amateur_vectors, use_cache=True)

                # Get predicted vectors
                expert_pred = expert_output['predicted_vectors'][:, -1:, :]
                amateur_pred = amateur_output['predicted_vectors'][:, -1:, :]

                # For continuous models, we need to get logits at the token level
                # Decode vectors to logits
                expert_logits = self.expert_model.autoencoder.decode(expert_pred.squeeze(1))
                amateur_logits = self.amateur_model.autoencoder.decode(amateur_pred.squeeze(1))

                # expert_logits, amateur_logits: [batch, chunk_size, vocab_size]
                # Process each position in chunk
                chunk_tokens = []
                for pos in range(self.expert_model.chunk_size):
                    expert_pos_logits = expert_logits[:, pos, :]
                    amateur_pos_logits = amateur_logits[:, pos, :]

                    # Compute contrastive logits
                    contrastive_logits = self.compute_contrastive_logits(
                        expert_pos_logits,
                        amateur_pos_logits,
                        alpha=alpha,
                        beta=beta
                    )

                    # Apply temperature
                    if temperature != 1.0:
                        contrastive_logits = contrastive_logits / temperature

                    # Sample
                    probs = F.softmax(contrastive_logits, dim=-1)
                    token = torch.multinomial(probs, num_samples=1)
                    chunk_tokens.append(token)

                # Stack chunk tokens: [batch, chunk_size]
                chunk = torch.cat(chunk_tokens, dim=-1)

                # Encode chunk back to vector
                next_vector = self.expert_model.autoencoder.encode(chunk)
                next_vector = next_vector.unsqueeze(1)  # [batch, 1, vector_dim]

                # Append to sequences
                expert_vectors = torch.cat([expert_vectors, next_vector], dim=1)
                amateur_vectors = torch.cat([amateur_vectors, next_vector], dim=1)

            # Convert final vectors to tokens
            generated_tokens = self.expert_model.vectors_to_tokens(expert_vectors, return_logits=False)

        result = {'generated_tokens': generated_tokens}

        if return_expert_only:
            # Generate with expert only for comparison
            expert_only_tokens = self.expert_model.generate(
                initial_tokens=initial_tokens,
                max_new_vectors=max_new_vectors,
                temperature=temperature,
                top_p=top_p
            )
            result['expert_tokens'] = expert_only_tokens

        return result

    def compute_adaptive_alpha(
        self,
        expert_logits: torch.Tensor,
        amateur_logits: torch.Tensor,
        base_alpha: Optional[float] = None
    ) -> torch.Tensor:
        """
        Compute adaptive contrastive weight based on model disagreement.

        When models disagree more, use higher alpha (more contrasting).
        When models agree, use lower alpha (less contrasting).

        Args:
            expert_logits: [batch_size, vocab_size] expert logits
            amateur_logits: [batch_size, vocab_size] amateur logits
            base_alpha: Base alpha value

        Returns:
            alpha: [batch_size] adaptive alpha values
        """
        base_alpha = base_alpha if base_alpha is not None else self.config.alpha

        # Compute distributions
        expert_probs = F.softmax(expert_logits, dim=-1)
        amateur_probs = F.softmax(amateur_logits, dim=-1)

        # Compute KL divergence (measures disagreement)
        kl_div = F.kl_div(
            amateur_probs.log(),
            expert_probs,
            reduction='none'
        ).sum(dim=-1)

        # Scale alpha based on disagreement
        # Higher disagreement -> higher alpha
        max_kl = math.log(expert_logits.shape[-1])  # Maximum possible KL
        normalized_kl = torch.clamp(kl_div / max_kl, 0, 1)

        # Alpha ranges from base_alpha/2 to base_alpha*2
        alpha = base_alpha * (0.5 + 1.5 * normalized_kl)

        return alpha


def create_amateur_from_expert(
    expert_model,
    method: str = "layer_subset",
    num_layers: Optional[int] = None
) -> 'ContinuousAutoregressiveModel':
    """
    Create an amateur model from an expert model.

    Methods:
    - "layer_subset": Use only first N layers of expert
    - "checkpoint": Load from earlier checkpoint
    - "separate": Use a separately trained smaller model

    Args:
        expert_model: The expert model
        method: Method for creating amateur
        num_layers: Number of layers for layer_subset method

    Returns:
        Amateur model
    """
    if method == "layer_subset":
        # Create a model with fewer layers
        num_layers = num_layers or max(1, expert_model.num_layers // 2)

        # Create amateur model config
        from MK3.calm.continuous_model import ContinuousAutoregressiveModel

        amateur_model = ContinuousAutoregressiveModel(
            vocab_size=expert_model.vocab_size,
            embedding_dim=expert_model.embedding_dim,
            vector_dim=expert_model.vector_dim,
            chunk_size=expert_model.chunk_size,
            num_layers=num_layers,  # Fewer layers
            num_heads=expert_model.num_heads,
            max_seq_length=expert_model.max_seq_length
        )

        # Copy autoencoder weights
        amateur_model.autoencoder.load_state_dict(
            expert_model.autoencoder.state_dict()
        )

        # Copy first N layers
        for i in range(num_layers):
            amateur_model.transformer_blocks[i].load_state_dict(
                expert_model.transformer_blocks[i].state_dict()
            )

        return amateur_model

    else:
        raise ValueError(f"Unknown method: {method}")


def compute_contrastive_score(
    expert_logits: torch.Tensor,
    amateur_logits: torch.Tensor,
    alpha: float = 0.5
) -> torch.Tensor:
    """
    Compute contrastive score for each token.

    Higher score = expert is more confident than amateur.

    Args:
        expert_logits: [batch_size, vocab_size] expert logits
        amateur_logits: [batch_size, vocab_size] amateur logits
        alpha: Contrastive weight

    Returns:
        scores: [batch_size, vocab_size] contrastive scores
    """
    expert_log_probs = F.log_softmax(expert_logits, dim=-1)
    amateur_log_probs = F.log_softmax(amateur_logits, dim=-1)

    contrastive_scores = expert_log_probs - alpha * amateur_log_probs

    return contrastive_scores


def analyze_expert_amateur_difference(
    expert_logits: torch.Tensor,
    amateur_logits: torch.Tensor,
    tokenizer,
    top_k: int = 10
) -> Dict[str, any]:
    """
    Analyze differences between expert and amateur predictions.

    Args:
        expert_logits: Expert model logits
        amateur_logits: Amateur model logits
        tokenizer: Tokenizer for decoding
        top_k: Number of top predictions to show

    Returns:
        Analysis dictionary
    """
    expert_probs = F.softmax(expert_logits, dim=-1)
    amateur_probs = F.softmax(amateur_logits, dim=-1)

    # Get top predictions from each
    expert_top_probs, expert_top_indices = torch.topk(expert_probs, top_k, dim=-1)
    amateur_top_probs, amateur_top_indices = torch.topk(amateur_probs, top_k, dim=-1)

    # Compute metrics
    kl_divergence = F.kl_div(amateur_probs.log(), expert_probs, reduction='batchmean')
    js_divergence = 0.5 * (
        F.kl_div(amateur_probs.log(), 0.5 * (expert_probs + amateur_probs), reduction='batchmean') +
        F.kl_div(expert_probs.log(), 0.5 * (expert_probs + amateur_probs), reduction='batchmean')
    )

    # Top predictions (first batch item)
    expert_top = [
        (tokenizer.decode([idx.item()]), prob.item())
        for idx, prob in zip(expert_top_indices[0], expert_top_probs[0])
    ]
    amateur_top = [
        (tokenizer.decode([idx.item()]), prob.item())
        for idx, prob in zip(amateur_top_indices[0], amateur_top_probs[0])
    ]

    return {
        'kl_divergence': kl_divergence.item(),
        'js_divergence': js_divergence.item(),
        'expert_top_predictions': expert_top,
        'amateur_top_predictions': amateur_top,
        'expert_entropy': -(expert_probs * expert_probs.log()).sum(dim=-1).mean().item(),
        'amateur_entropy': -(amateur_probs * amateur_probs.log()).sum(dim=-1).mean().item()
    }
