"""
Advanced Generation Methods for Continuous Model

Extends the ContinuousAutoregressiveModel with advanced inference capabilities:
- Advanced sampling (min-p, top-k, top-p, temperature, penalties)
- Streaming generation
- Adaptive temperature
- Chain-of-Thought integration
- Self-Consistency integration
- Contrastive decoding support
"""

import torch
import torch.nn.functional as F
from typing import Optional, List, Dict, Tuple, Callable
import sys
import os

# Add parent to path
sys.path.append(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))

from inference.sampling import sample_from_logits, adaptive_temperature


def generate_advanced(
    model,
    initial_tokens: torch.Tensor,
    max_new_vectors: int = 10,
    temperature: float = 1.0,
    top_p: Optional[float] = None,
    top_k: Optional[int] = None,
    min_p: Optional[float] = None,
    repetition_penalty: float = 1.0,
    frequency_penalty: float = 0.0,
    presence_penalty: float = 0.0,
) -> torch.Tensor:
    """
    Generate text with advanced sampling strategies.

    Args:
        model: ContinuousAutoregressiveModel instance
        initial_tokens: [batch, seq_len] initial tokens
        max_new_vectors: Number of vectors to generate
        temperature: Sampling temperature
        top_p: Nucleus sampling threshold
        top_k: Top-k filtering
        min_p: Min-p filtering (often better than top-p)
        repetition_penalty: Penalty for repeated tokens
        frequency_penalty: Frequency-based penalty
        presence_penalty: Presence-based penalty

    Returns:
        generated_tokens: [batch, seq_len + max_new_vectors * chunk_size]
    """
    model.eval()

    with torch.no_grad():
        # Convert initial tokens to vectors
        continuous_vectors, _ = model.tokenize_to_vectors(initial_tokens)
        all_generated_tokens = [initial_tokens]

        # Generate new vectors
        for step in range(max_new_vectors):
            # Forward pass
            output = model.forward(continuous_vectors, use_cache=True)
            predicted_vectors = output['predicted_vectors']

            # Get last predicted vector
            next_vector = predicted_vectors[:, -1:, :]  # [batch, 1, vector_dim]

            # Decode to token-level logits
            chunk_logits = model.autoencoder.decode(next_vector.squeeze(1))  # [batch, chunk_size, vocab_size]

            # Sample each position in chunk
            chunk_tokens = []
            current_generated = torch.cat(all_generated_tokens, dim=1)

            for pos in range(model.chunk_size):
                pos_logits = chunk_logits[:, pos, :]  # [batch, vocab_size]

                # Use advanced sampling
                sampled_token = sample_from_logits(
                    logits=pos_logits,
                    temperature=temperature,
                    top_k=top_k,
                    top_p=top_p,
                    min_p=min_p,
                    generated_tokens=current_generated,
                    repetition_penalty=repetition_penalty,
                    frequency_penalty=frequency_penalty,
                    presence_penalty=presence_penalty
                )  # [batch]

                chunk_tokens.append(sampled_token.unsqueeze(1))
                current_generated = torch.cat([current_generated, sampled_token.unsqueeze(1)], dim=1)

            # Stack chunk: [batch, chunk_size]
            chunk = torch.cat(chunk_tokens, dim=1)
            all_generated_tokens.append(chunk)

            # Re-encode chunk to continuous vector
            next_vector = model.autoencoder.encode(chunk).unsqueeze(1)
            continuous_vectors = torch.cat([continuous_vectors, next_vector], dim=1)

        # Concatenate all tokens
        generated_tokens = torch.cat(all_generated_tokens, dim=1)

    return generated_tokens


def generate_streaming(
    model,
    initial_tokens: torch.Tensor,
    max_new_vectors: int = 10,
    temperature: float = 1.0,
    top_p: Optional[float] = None,
    top_k: Optional[int] = None,
    min_p: Optional[float] = None,
    callback: Optional[Callable[[torch.Tensor], None]] = None,
):
    """
    Generate text with streaming output (yields chunks as they're generated).

    This allows for real-time display of generation progress.

    Args:
        model: ContinuousAutoregressiveModel instance
        initial_tokens: [batch, seq_len] initial tokens
        max_new_vectors: Number of vectors to generate
        temperature: Sampling temperature
        top_p: Nucleus sampling threshold
        top_k: Top-k filtering
        min_p: Min-p filtering
        callback: Optional callback function called with each chunk

    Yields:
        Token chunks [batch, chunk_size] as they are generated
    """
    model.eval()

    with torch.no_grad():
        # Convert initial tokens to vectors
        continuous_vectors, _ = model.tokenize_to_vectors(initial_tokens)

        # Generate new vectors one at a time
        for step in range(max_new_vectors):
            # Forward pass
            output = model.forward(continuous_vectors, use_cache=True)
            predicted_vectors = output['predicted_vectors']

            # Get last predicted vector
            next_vector = predicted_vectors[:, -1:, :]

            # Decode to token logits
            chunk_logits = model.autoencoder.decode(next_vector.squeeze(1))  # [batch, chunk_size, vocab_size]

            # Sample each position
            chunk_tokens = []
            for pos in range(model.chunk_size):
                pos_logits = chunk_logits[:, pos, :]

                sampled_token = sample_from_logits(
                    logits=pos_logits,
                    temperature=temperature,
                    top_k=top_k,
                    top_p=top_p,
                    min_p=min_p
                )

                chunk_tokens.append(sampled_token.unsqueeze(1))

            # Stack chunk
            chunk = torch.cat(chunk_tokens, dim=1)  # [batch, chunk_size]

            # Re-encode for next iteration
            next_vector = model.autoencoder.encode(chunk).unsqueeze(1)
            continuous_vectors = torch.cat([continuous_vectors, next_vector], dim=1)

            # Yield chunk
            yield chunk

            # Call callback if provided
            if callback is not None:
                callback(chunk)


def generate_with_adaptive_temperature(
    model,
    initial_tokens: torch.Tensor,
    max_new_vectors: int = 10,
    base_temperature: float = 1.0,
    entropy_threshold: float = 2.0,
    min_temperature: float = 0.5,
    max_temperature: float = 2.0,
    **kwargs
) -> Tuple[torch.Tensor, List[float]]:
    """
    Generate with adaptive temperature based on prediction entropy.

    High entropy (uncertain) -> higher temperature (more random)
    Low entropy (confident) -> lower temperature (more deterministic)

    Args:
        model: ContinuousAutoregressiveModel instance
        initial_tokens: Initial token sequence
        max_new_vectors: Number of vectors to generate
        base_temperature: Base temperature
        entropy_threshold: Entropy midpoint for scaling
        min_temperature: Minimum temperature
        max_temperature: Maximum temperature
        **kwargs: Additional sampling arguments (top_p, top_k, etc.)

    Returns:
        Tuple of (generated_tokens, temperature_history)
    """
    model.eval()
    temperature_history = []

    with torch.no_grad():
        continuous_vectors, _ = model.tokenize_to_vectors(initial_tokens)
        generated_chunks = [initial_tokens]

        for step in range(max_new_vectors):
            output = model.forward(continuous_vectors, use_cache=True)
            predicted_vectors = output['predicted_vectors']
            next_vector = predicted_vectors[:, -1:, :]

            # Decode to logits
            chunk_logits = model.autoencoder.decode(next_vector.squeeze(1))

            # Sample each position with adaptive temperature
            chunk_tokens = []
            for pos in range(model.chunk_size):
                pos_logits = chunk_logits[:, pos, :]

                # Compute adaptive temperature
                adjusted_logits, temp = adaptive_temperature(
                    pos_logits,
                    base_temperature=base_temperature,
                    entropy_threshold=entropy_threshold,
                    min_temperature=min_temperature,
                    max_temperature=max_temperature
                )

                temperature_history.append(temp.mean().item())

                # Sample with adjusted logits (temperature already applied)
                sampled_token = sample_from_logits(
                    logits=adjusted_logits,
                    temperature=1.0,  # Already applied in adaptive_temperature
                    **kwargs
                )

                chunk_tokens.append(sampled_token.unsqueeze(1))

            chunk = torch.cat(chunk_tokens, dim=1)
            generated_chunks.append(chunk)

            # Re-encode
            next_vector = model.autoencoder.encode(chunk).unsqueeze(1)
            continuous_vectors = torch.cat([continuous_vectors, next_vector], dim=1)

        generated_tokens = torch.cat(generated_chunks, dim=1)

    return generated_tokens, temperature_history


def generate_with_chain_of_thought(
    model,
    tokenizer,
    prompt: str,
    max_new_vectors: int = 20,
    cot_config: Optional[dict] = None,
    **generation_kwargs
) -> Dict[str, any]:
    """
    Generate with Chain-of-Thought prompting.

    Args:
        model: ContinuousAutoregressiveModel instance
        tokenizer: Tokenizer for encoding/decoding
        prompt: Input prompt/question
        max_new_vectors: Maximum vectors to generate
        cot_config: Configuration for CoT (see inference.chain_of_thought.CoTConfig)
        **generation_kwargs: Additional generation arguments

    Returns:
        Dictionary with reasoning and final answer
    """
    from inference.chain_of_thought import ChainOfThoughtGenerator, CoTPrompt, CoTConfig

    # Create CoT config
    if cot_config is None:
        config = CoTConfig()
    else:
        config = CoTConfig(**cot_config)

    # Create CoT generator
    cot_generator = ChainOfThoughtGenerator(model, tokenizer, config)

    # Create prompt
    cot_prompt = CoTPrompt(question=prompt)

    # Generate with CoT
    result = cot_generator.generate_with_cot(
        cot_prompt,
        max_length=max_new_vectors * model.chunk_size,
        **generation_kwargs
    )

    return result


def generate_with_self_consistency(
    model,
    tokenizer,
    prompt: str,
    num_paths: int = 5,
    max_new_vectors: int = 20,
    temperature: float = 0.7,
    sc_config: Optional[dict] = None,
    **generation_kwargs
) -> Dict[str, any]:
    """
    Generate with Self-Consistency (multiple reasoning paths).

    Args:
        model: ContinuousAutoregressiveModel instance
        tokenizer: Tokenizer for encoding/decoding
        prompt: Input prompt/question
        num_paths: Number of reasoning paths to generate
        max_new_vectors: Maximum vectors per path
        temperature: Sampling temperature (higher = more diversity)
        sc_config: Configuration for self-consistency
        **generation_kwargs: Additional generation arguments

    Returns:
        Dictionary with consensus answer and all paths
    """
    from inference.self_consistency import SelfConsistencyGenerator, SelfConsistencyConfig

    # Create config
    if sc_config is None:
        config = SelfConsistencyConfig(num_paths=num_paths, temperature=temperature)
    else:
        config = SelfConsistencyConfig(**sc_config)

    # Create generator
    sc_generator = SelfConsistencyGenerator(model, tokenizer, config)

    # Generate with self-consistency
    result = sc_generator.generate_with_self_consistency(
        prompt=prompt,
        max_length=max_new_vectors * model.chunk_size,
        return_all_paths=True,
        **generation_kwargs
    )

    return result


def generate_with_contrastive_decoding(
    expert_model,
    amateur_model,
    initial_tokens: torch.Tensor,
    max_new_vectors: int = 10,
    alpha: float = 0.5,
    beta: float = 0.5,
    **generation_kwargs
) -> Dict[str, torch.Tensor]:
    """
    Generate using contrastive decoding (expert - amateur).

    Args:
        expert_model: Expert (stronger) model
        amateur_model: Amateur (weaker) model
        initial_tokens: Initial tokens
        max_new_vectors: Number of vectors to generate
        alpha: Contrastive weight
        beta: Plausibility threshold
        **generation_kwargs: Additional generation arguments

    Returns:
        Dictionary with generated tokens and statistics
    """
    from inference.contrastive_decoding import ContrastiveDecoder, ContrastiveConfig

    # Create config
    config = ContrastiveConfig(alpha=alpha, beta=beta)

    # Create decoder
    decoder = ContrastiveDecoder(expert_model, amateur_model, config)

    # Generate
    result = decoder.generate_contrastive(
        initial_tokens=initial_tokens,
        max_new_vectors=max_new_vectors,
        **generation_kwargs
    )

    return result


def generate_with_speculative_decoding(
    target_model,
    draft_model,
    initial_tokens: torch.Tensor,
    max_new_tokens: int = 100,
    num_speculative: int = 4,
    **generation_kwargs
) -> Dict[str, any]:
    """
    Generate using speculative decoding for faster generation.

    Args:
        target_model: Target (main) model
        draft_model: Draft (faster) model
        initial_tokens: Initial tokens
        max_new_tokens: Maximum new tokens to generate
        num_speculative: Number of speculative tokens per step
        **generation_kwargs: Additional generation arguments

    Returns:
        Dictionary with generated tokens and statistics
    """
    from inference.speculative import SpeculativeDecoder, SpeculativeConfig

    # Create config
    config = SpeculativeConfig(num_speculative_tokens=num_speculative)

    # Create decoder
    decoder = SpeculativeDecoder(target_model, draft_model, config)

    # Generate
    result = decoder.generate_speculative(
        initial_tokens=initial_tokens,
        max_new_tokens=max_new_tokens,
        return_stats=True,
        **generation_kwargs
    )

    return result


# Convenience class to attach these methods to model instances
class AdvancedGenerationMixin:
    """
    Mixin class to add advanced generation methods to ContinuousAutoregressiveModel.

    Usage:
        class EnhancedModel(ContinuousAutoregressiveModel, AdvancedGenerationMixin):
            pass

        model = EnhancedModel(...)
        tokens = model.generate_advanced(...)
    """

    def generate_advanced(self, *args, **kwargs):
        return generate_advanced(self, *args, **kwargs)

    def generate_streaming(self, *args, **kwargs):
        return generate_streaming(self, *args, **kwargs)

    def generate_with_adaptive_temperature(self, *args, **kwargs):
        return generate_with_adaptive_temperature(self, *args, **kwargs)

    def generate_with_chain_of_thought(self, tokenizer, *args, **kwargs):
        return generate_with_chain_of_thought(self, tokenizer, *args, **kwargs)

    def generate_with_self_consistency(self, tokenizer, *args, **kwargs):
        return generate_with_self_consistency(self, tokenizer, *args, **kwargs)
