"""
Advanced Sampling Strategies for Text Generation

Implements state-of-the-art sampling methods:
- Temperature sampling
- Top-k sampling
- Top-p (nucleus) sampling
- Min-p sampling (newer, often superior to top-p)
- Combined sampling strategies

References:
- Top-k: Fan et al. 2018 "Hierarchical Neural Story Generation"
- Top-p: Holtzman et al. 2019 "The Curious Case of Neural Text Degeneration"
- Min-p: https://github.com/ggerganov/llama.cpp/pull/3841
"""

import torch
import torch.nn.functional as F
from typing import Optional, Tuple, Callable, Dict, List
from dataclasses import dataclass
import math


@dataclass
class SamplingStrategy:
    """Base configuration for sampling strategies."""
    temperature: float = 1.0
    top_k: Optional[int] = None
    top_p: Optional[float] = None
    min_p: Optional[float] = None
    repetition_penalty: float = 1.0
    frequency_penalty: float = 0.0
    presence_penalty: float = 0.0


class TemperatureSampler:
    """
    Temperature sampling adjusts the randomness of predictions.

    - temperature < 1.0: More deterministic (sharper distribution)
    - temperature = 1.0: Unchanged distribution
    - temperature > 1.0: More random (flatter distribution)
    """

    @staticmethod
    def apply(logits: torch.Tensor, temperature: float = 1.0) -> torch.Tensor:
        """
        Apply temperature scaling to logits.

        Args:
            logits: [batch_size, vocab_size] or [batch_size, seq_len, vocab_size]
            temperature: Temperature value (> 0)

        Returns:
            Temperature-scaled logits
        """
        if temperature <= 0:
            raise ValueError(f"Temperature must be positive, got {temperature}")

        if temperature == 1.0:
            return logits

        return logits / temperature


class TopKSampler:
    """
    Top-k sampling: Sample from the k most likely tokens.

    Filters out all tokens except the top-k most probable ones.
    """

    @staticmethod
    def apply(logits: torch.Tensor, k: int) -> torch.Tensor:
        """
        Apply top-k filtering to logits.

        Args:
            logits: [batch_size, vocab_size] or [batch_size, seq_len, vocab_size]
            k: Number of top tokens to keep

        Returns:
            Filtered logits (non-top-k tokens set to -inf)
        """
        if k <= 0:
            return logits

        # Get vocab size
        vocab_size = logits.shape[-1]
        k = min(k, vocab_size)

        # Get top-k values and indices
        top_k_values, top_k_indices = torch.topk(logits, k, dim=-1)

        # Create mask: set all non-top-k to -inf
        filtered_logits = torch.full_like(logits, float('-inf'))
        filtered_logits.scatter_(-1, top_k_indices, top_k_values)

        return filtered_logits


class TopPSampler:
    """
    Top-p (nucleus) sampling: Sample from smallest set of tokens with cumulative prob >= p.

    More adaptive than top-k: adjusts the number of tokens based on their probabilities.
    """

    @staticmethod
    def apply(logits: torch.Tensor, p: float) -> torch.Tensor:
        """
        Apply top-p (nucleus) filtering to logits.

        Args:
            logits: [batch_size, vocab_size] or [batch_size, seq_len, vocab_size]
            p: Cumulative probability threshold (0 < p <= 1.0)

        Returns:
            Filtered logits (tokens outside nucleus set to -inf)
        """
        if not 0 < p <= 1.0:
            raise ValueError(f"top_p must be in (0, 1], got {p}")

        if p >= 1.0:
            return logits

        # Sort logits in descending order
        sorted_logits, sorted_indices = torch.sort(logits, descending=True, dim=-1)

        # Compute softmax probabilities
        sorted_probs = F.softmax(sorted_logits, dim=-1)

        # Compute cumulative probabilities
        cumulative_probs = torch.cumsum(sorted_probs, dim=-1)

        # Create mask: remove tokens with cumulative probability above threshold
        # Keep at least one token (the most probable one)
        sorted_indices_to_remove = cumulative_probs > p
        sorted_indices_to_remove[..., 0] = False  # Keep the most probable token

        # Shift right to keep the first token that exceeds threshold
        sorted_indices_to_remove[..., 1:] = sorted_indices_to_remove[..., :-1].clone()
        sorted_indices_to_remove[..., 0] = False

        # Create filtered logits
        filtered_logits = logits.clone()

        # Scatter -inf to removed positions
        # First, create indices in original order
        indices_to_remove = sorted_indices_to_remove.scatter(
            -1, sorted_indices, sorted_indices_to_remove
        )
        filtered_logits[indices_to_remove] = float('-inf')

        return filtered_logits


class MinPSampler:
    """
    Min-p sampling: Sample from tokens with probability >= p * max_prob.

    This is often superior to top-p because it adapts to the confidence
    of the model - when the model is confident (high max prob), fewer tokens
    are sampled; when uncertain (low max prob), more tokens are considered.

    Reference: https://github.com/ggerganov/llama.cpp/pull/3841
    """

    @staticmethod
    def apply(logits: torch.Tensor, min_p: float) -> torch.Tensor:
        """
        Apply min-p filtering to logits.

        Args:
            logits: [batch_size, vocab_size] or [batch_size, seq_len, vocab_size]
            min_p: Minimum probability threshold (0 < min_p <= 1.0)

        Returns:
            Filtered logits (tokens below threshold set to -inf)
        """
        if not 0 < min_p <= 1.0:
            raise ValueError(f"min_p must be in (0, 1], got {min_p}")

        if min_p >= 1.0:
            # Only keep the most probable token
            max_logit = logits.max(dim=-1, keepdim=True)[0]
            filtered_logits = torch.full_like(logits, float('-inf'))
            filtered_logits = torch.where(logits >= max_logit, logits, filtered_logits)
            return filtered_logits

        # Convert to probabilities
        probs = F.softmax(logits, dim=-1)

        # Get max probability
        max_prob = probs.max(dim=-1, keepdim=True)[0]

        # Calculate threshold
        threshold = min_p * max_prob

        # Create mask
        mask = probs < threshold

        # Apply mask
        filtered_logits = logits.clone()
        filtered_logits[mask] = float('-inf')

        return filtered_logits


class RepetitionPenaltySampler:
    """
    Apply repetition penalty to discourage repeated tokens.

    Implements the penalty from CTRL paper (Keskar et al. 2019).
    """

    @staticmethod
    def apply(
        logits: torch.Tensor,
        generated_tokens: torch.Tensor,
        penalty: float = 1.0,
        frequency_penalty: float = 0.0,
        presence_penalty: float = 0.0
    ) -> torch.Tensor:
        """
        Apply repetition penalties to logits.

        Args:
            logits: [batch_size, vocab_size] current logits
            generated_tokens: [batch_size, seq_len] previously generated tokens
            penalty: Repetition penalty multiplier (1.0 = no penalty, >1.0 = discourage)
            frequency_penalty: Penalty proportional to token frequency
            presence_penalty: Fixed penalty for any token that appeared

        Returns:
            Penalized logits
        """
        if penalty == 1.0 and frequency_penalty == 0.0 and presence_penalty == 0.0:
            return logits

        batch_size, vocab_size = logits.shape
        penalized_logits = logits.clone()

        for i in range(batch_size):
            # Get unique tokens and their counts
            unique_tokens, counts = torch.unique(generated_tokens[i], return_counts=True)

            for token, count in zip(unique_tokens, counts):
                if token >= vocab_size:
                    continue

                # Apply repetition penalty (CTRL-style)
                if penalty != 1.0:
                    if logits[i, token] < 0:
                        penalized_logits[i, token] *= penalty
                    else:
                        penalized_logits[i, token] /= penalty

                # Apply frequency penalty (proportional to count)
                if frequency_penalty != 0.0:
                    penalized_logits[i, token] -= frequency_penalty * count.float()

                # Apply presence penalty (binary)
                if presence_penalty != 0.0:
                    penalized_logits[i, token] -= presence_penalty

        return penalized_logits


class CombinedSampler:
    """
    Combined sampling strategy applying multiple techniques in sequence.

    Order of operations:
    1. Repetition penalties
    2. Temperature scaling
    3. Top-k filtering
    4. Top-p filtering
    5. Min-p filtering
    6. Sampling
    """

    def __init__(self, config: SamplingStrategy):
        """
        Initialize combined sampler.

        Args:
            config: Sampling configuration
        """
        self.config = config

    def sample(
        self,
        logits: torch.Tensor,
        generated_tokens: Optional[torch.Tensor] = None,
        num_samples: int = 1
    ) -> torch.Tensor:
        """
        Sample tokens using combined strategy.

        Args:
            logits: [batch_size, vocab_size] logits
            generated_tokens: [batch_size, seq_len] previously generated tokens (for penalties)
            num_samples: Number of samples to draw

        Returns:
            sampled_tokens: [batch_size, num_samples] sampled token indices
        """
        # Apply repetition penalties
        if generated_tokens is not None:
            logits = RepetitionPenaltySampler.apply(
                logits,
                generated_tokens,
                penalty=self.config.repetition_penalty,
                frequency_penalty=self.config.frequency_penalty,
                presence_penalty=self.config.presence_penalty
            )

        # Apply temperature
        if self.config.temperature != 1.0:
            logits = TemperatureSampler.apply(logits, self.config.temperature)

        # Apply top-k
        if self.config.top_k is not None:
            logits = TopKSampler.apply(logits, self.config.top_k)

        # Apply top-p
        if self.config.top_p is not None:
            logits = TopPSampler.apply(logits, self.config.top_p)

        # Apply min-p (usually use min-p OR top-p, not both)
        if self.config.min_p is not None:
            logits = MinPSampler.apply(logits, self.config.min_p)

        # Sample
        probs = F.softmax(logits, dim=-1)

        if num_samples == 1:
            # Single sample per batch item
            sampled_tokens = torch.multinomial(probs, num_samples=1)
        else:
            # Multiple samples per batch item
            sampled_tokens = torch.multinomial(probs, num_samples=num_samples)

        return sampled_tokens


def sample_from_logits(
    logits: torch.Tensor,
    temperature: float = 1.0,
    top_k: Optional[int] = None,
    top_p: Optional[float] = None,
    min_p: Optional[float] = None,
    num_samples: int = 1,
    generated_tokens: Optional[torch.Tensor] = None,
    repetition_penalty: float = 1.0,
    frequency_penalty: float = 0.0,
    presence_penalty: float = 0.0
) -> torch.Tensor:
    """
    Convenience function for sampling with various strategies.

    Args:
        logits: [batch_size, vocab_size] logits to sample from
        temperature: Temperature for sampling
        top_k: Top-k filtering (None = disabled)
        top_p: Top-p filtering (None = disabled)
        min_p: Min-p filtering (None = disabled)
        num_samples: Number of samples to draw
        generated_tokens: Previously generated tokens for repetition penalties
        repetition_penalty: Repetition penalty multiplier
        frequency_penalty: Frequency-based penalty
        presence_penalty: Presence-based penalty

    Returns:
        sampled_tokens: [batch_size, num_samples] or [batch_size] if num_samples=1
    """
    config = SamplingStrategy(
        temperature=temperature,
        top_k=top_k,
        top_p=top_p,
        min_p=min_p,
        repetition_penalty=repetition_penalty,
        frequency_penalty=frequency_penalty,
        presence_penalty=presence_penalty
    )

    sampler = CombinedSampler(config)
    samples = sampler.sample(logits, generated_tokens, num_samples)

    # Squeeze if single sample
    if num_samples == 1:
        samples = samples.squeeze(-1)

    return samples


def compute_entropy(logits: torch.Tensor) -> torch.Tensor:
    """
    Compute entropy of logits distribution.

    Higher entropy = more uncertainty = flatter distribution
    Lower entropy = more certainty = sharper distribution

    Args:
        logits: [batch_size, vocab_size] logits

    Returns:
        entropy: [batch_size] entropy values
    """
    probs = F.softmax(logits, dim=-1)
    log_probs = F.log_softmax(logits, dim=-1)
    entropy = -(probs * log_probs).sum(dim=-1)
    return entropy


def adaptive_temperature(
    logits: torch.Tensor,
    base_temperature: float = 1.0,
    entropy_threshold: float = 2.0,
    min_temperature: float = 0.5,
    max_temperature: float = 2.0
) -> Tuple[torch.Tensor, torch.Tensor]:
    """
    Adaptive temperature based on prediction entropy.

    When model is uncertain (high entropy), increase temperature.
    When model is confident (low entropy), decrease temperature.

    Args:
        logits: [batch_size, vocab_size] logits
        base_temperature: Base temperature value
        entropy_threshold: Entropy midpoint for scaling
        min_temperature: Minimum temperature
        max_temperature: Maximum temperature

    Returns:
        Tuple of:
        - adjusted_logits: Temperature-adjusted logits
        - temperatures: [batch_size] computed temperatures
    """
    # Compute entropy
    entropy = compute_entropy(logits)

    # Scale temperature based on entropy
    # High entropy -> higher temperature
    # Low entropy -> lower temperature
    temperature_scale = torch.sigmoid((entropy - entropy_threshold) / entropy_threshold)
    temperatures = min_temperature + (max_temperature - min_temperature) * temperature_scale
    temperatures = temperatures.unsqueeze(-1)  # [batch_size, 1]

    # Apply temperature
    adjusted_logits = logits / temperatures

    return adjusted_logits, temperatures.squeeze(-1)
