"""
Speculative Decoding for Faster Generation

Implements speculative decoding to accelerate autoregressive generation:
1. Use a fast "draft" model to generate K candidate tokens
2. Verify candidates in parallel with the main model
3. Accept tokens that match, reject and resample others

This can provide 2-3x speedup without changing output distribution.

References:
- Leviathan et al. 2023: "Fast Inference from Transformers via Speculative Decoding"
- Chen et al. 2023: "Accelerating Large Language Model Decoding with Speculative Sampling"
"""

import torch
import torch.nn.functional as F
from typing import Optional, List, Dict, Tuple, Callable
from dataclasses import dataclass
import math


@dataclass
class SpeculativeConfig:
    """Configuration for speculative decoding."""

    # Number of tokens to generate speculatively
    num_speculative_tokens: int = 4

    # Acceptance strategy
    acceptance_method: str = "probabilistic"  # "greedy", "probabilistic", "typical"

    # Temperature for draft and target models
    draft_temperature: float = 1.0
    target_temperature: float = 1.0

    # Probabilistic acceptance threshold
    acceptance_threshold: float = 0.8

    # Whether to use early exit on rejection
    early_exit: bool = True

    # Max retries for draft generation
    max_retries: int = 1


class SpeculativeDecoder:
    """
    Speculative decoding with draft and target models.
    """

    def __init__(
        self,
        target_model,
        draft_model,
        config: Optional[SpeculativeConfig] = None
    ):
        """
        Initialize speculative decoder.

        Args:
            target_model: Main (larger/slower) model
            draft_model: Draft (smaller/faster) model for speculation
            config: Speculative decoding configuration
        """
        self.target_model = target_model
        self.draft_model = draft_model
        self.config = config or SpeculativeConfig()

        # Statistics
        self.acceptance_stats = {
            'total_proposed': 0,
            'total_accepted': 0,
            'acceptance_rate': 0.0
        }

    def generate_speculative(
        self,
        initial_tokens: torch.Tensor,
        max_new_tokens: int = 100,
        temperature: float = 1.0,
        top_p: float = 0.9,
        num_speculative: Optional[int] = None,
        device: Optional[torch.device] = None,
        return_stats: bool = False
    ) -> Dict[str, any]:
        """
        Generate text using speculative decoding.

        Args:
            initial_tokens: [batch_size, seq_len] initial tokens
            max_new_tokens: Maximum number of new tokens to generate
            temperature: Sampling temperature
            top_p: Nucleus sampling threshold
            num_speculative: Number of speculative tokens per step (uses config if None)
            device: Device for computation
            return_stats: Whether to return acceptance statistics

        Returns:
            Dictionary containing:
            - 'generated_tokens': Generated token sequence
            - 'acceptance_stats': Acceptance statistics (if return_stats=True)
            - 'num_target_calls': Number of target model forward passes
        """
        device = device or next(self.target_model.parameters()).device
        num_speculative = num_speculative or self.config.num_speculative_tokens

        batch_size = initial_tokens.shape[0]
        current_tokens = initial_tokens.to(device)

        self.target_model.eval()
        self.draft_model.eval()

        tokens_generated = 0
        num_target_calls = 0

        # Local acceptance stats
        local_accepted = 0
        local_proposed = 0

        with torch.no_grad():
            while tokens_generated < max_new_tokens:
                # Determine how many tokens to speculate
                remaining = max_new_tokens - tokens_generated
                K = min(num_speculative, remaining)

                # Step 1: Generate K candidate tokens with draft model
                draft_candidates = self._generate_draft_candidates(
                    current_tokens,
                    K,
                    temperature=self.config.draft_temperature,
                    device=device
                )

                # draft_candidates: [batch_size, K]
                num_candidates = draft_candidates.shape[1]
                local_proposed += num_candidates

                # Step 2: Verify candidates with target model in parallel
                # Append candidates to current sequence
                extended_tokens = torch.cat([current_tokens, draft_candidates], dim=1)

                # Get target model predictions for the extended sequence
                target_logits = self._get_target_logits_batch(
                    extended_tokens,
                    device=device
                )
                num_target_calls += 1

                # target_logits: [batch_size, extended_seq_len, vocab_size]
                # We care about positions where we added candidates

                # Step 3: Accept/reject candidates
                num_accepted, accepted_tokens = self._verify_and_accept(
                    candidates=draft_candidates,
                    target_logits=target_logits,
                    current_seq_len=current_tokens.shape[1],
                    temperature=temperature,
                    device=device
                )

                local_accepted += num_accepted

                # Update current sequence
                if num_accepted > 0:
                    current_tokens = torch.cat([current_tokens, accepted_tokens], dim=1)
                    tokens_generated += num_accepted

                # If not all accepted, sample one more from target model at rejection point
                if num_accepted < num_candidates and tokens_generated < max_new_tokens:
                    # Sample next token from target model distribution
                    rejection_pos = current_tokens.shape[1]
                    if rejection_pos < target_logits.shape[1]:
                        rejection_logits = target_logits[:, rejection_pos - 1, :]
                        rejection_logits = rejection_logits / temperature
                        probs = F.softmax(rejection_logits, dim=-1)
                        next_token = torch.multinomial(probs, num_samples=1)
                        current_tokens = torch.cat([current_tokens, next_token], dim=1)
                        tokens_generated += 1
                        local_accepted += 1
                        local_proposed += 1

                # Safety check
                if tokens_generated >= max_new_tokens:
                    break

        # Update global stats
        self.acceptance_stats['total_proposed'] += local_proposed
        self.acceptance_stats['total_accepted'] += local_accepted
        if self.acceptance_stats['total_proposed'] > 0:
            self.acceptance_stats['acceptance_rate'] = (
                self.acceptance_stats['total_accepted'] / self.acceptance_stats['total_proposed']
            )

        result = {
            'generated_tokens': current_tokens,
            'num_target_calls': num_target_calls,
            'tokens_generated': tokens_generated
        }

        if return_stats:
            result['acceptance_stats'] = {
                'accepted': local_accepted,
                'proposed': local_proposed,
                'acceptance_rate': local_accepted / local_proposed if local_proposed > 0 else 0.0,
                'global_acceptance_rate': self.acceptance_stats['acceptance_rate'],
                'speedup_estimate': num_speculative * (local_accepted / local_proposed) if local_proposed > 0 else 1.0
            }

        return result

    def _generate_draft_candidates(
        self,
        current_tokens: torch.Tensor,
        K: int,
        temperature: float,
        device: torch.device
    ) -> torch.Tensor:
        """
        Generate K candidate tokens using draft model.

        Args:
            current_tokens: [batch_size, seq_len] current sequence
            K: Number of candidates to generate
            temperature: Sampling temperature
            device: Device for computation

        Returns:
            candidates: [batch_size, K] candidate tokens
        """
        # For continuous models, we need to work with vectors
        # Generate K tokens autoregressively with draft model

        draft_sequence = current_tokens.clone()
        candidates = []

        for _ in range(K):
            # Convert to vectors
            draft_vectors, _ = self.draft_model.tokenize_to_vectors(draft_sequence)

            # Forward pass
            draft_output = self.draft_model.forward(draft_vectors, use_cache=False)
            draft_pred_vector = draft_output['predicted_vectors'][:, -1:, :]

            # Decode to logits
            draft_logits = self.draft_model.autoencoder.decode(draft_pred_vector.squeeze(1))
            # draft_logits: [batch, chunk_size, vocab_size]

            # Take first position of chunk (or last, depending on model design)
            # For simplicity, take mean of chunk positions
            draft_logits_pos = draft_logits.mean(dim=1)  # [batch, vocab_size]

            # Apply temperature and sample
            draft_logits_pos = draft_logits_pos / temperature
            probs = F.softmax(draft_logits_pos, dim=-1)
            next_token = torch.multinomial(probs, num_samples=1)  # [batch, 1]

            candidates.append(next_token)
            draft_sequence = torch.cat([draft_sequence, next_token], dim=1)

        # Stack candidates: [batch, K]
        candidates = torch.cat(candidates, dim=1)

        return candidates

    def _get_target_logits_batch(
        self,
        tokens: torch.Tensor,
        device: torch.device
    ) -> torch.Tensor:
        """
        Get target model logits for a sequence in parallel.

        Args:
            tokens: [batch_size, seq_len] token sequence
            device: Device for computation

        Returns:
            logits: [batch_size, seq_len, vocab_size] logits for each position
        """
        # Convert to continuous vectors
        vectors, _ = self.target_model.tokenize_to_vectors(tokens)

        # Forward pass
        output = self.target_model.forward(vectors, use_cache=False)
        predicted_vectors = output['predicted_vectors']  # [batch, num_chunks, vector_dim]

        # Decode all vectors to logits
        batch_size, num_chunks, vector_dim = predicted_vectors.shape
        all_logits = []

        for i in range(num_chunks):
            chunk_logits = self.target_model.autoencoder.decode(predicted_vectors[:, i, :])
            # chunk_logits: [batch, chunk_size, vocab_size]
            all_logits.append(chunk_logits)

        # Concatenate: [batch, num_chunks * chunk_size, vocab_size]
        logits = torch.cat(all_logits, dim=1)

        return logits

    def _verify_and_accept(
        self,
        candidates: torch.Tensor,
        target_logits: torch.Tensor,
        current_seq_len: int,
        temperature: float,
        device: torch.device
    ) -> Tuple[int, torch.Tensor]:
        """
        Verify candidates against target model predictions and accept/reject.

        Args:
            candidates: [batch_size, K] candidate tokens
            target_logits: [batch_size, seq_len, vocab_size] target model logits
            current_seq_len: Length of sequence before candidates
            temperature: Sampling temperature
            device: Device for computation

        Returns:
            Tuple of (num_accepted, accepted_tokens)
        """
        batch_size, K = candidates.shape

        # Get target probabilities for positions where candidates were added
        accepted_tokens = []
        num_accepted = 0

        for k in range(K):
            # Position in target_logits for k-th candidate
            # The target_logits contains predictions for positions [0, ..., current_seq_len + K - 1]
            # Position current_seq_len predicts token at current_seq_len + 1
            pos = current_seq_len + k - 1

            if pos < 0 or pos >= target_logits.shape[1]:
                break

            # Get target distribution for this position
            target_pos_logits = target_logits[:, pos, :]  # [batch, vocab_size]
            target_pos_logits = target_pos_logits / temperature
            target_probs = F.softmax(target_pos_logits, dim=-1)

            # Get probability of candidate token under target model
            candidate_token = candidates[:, k:k+1]  # [batch, 1]
            candidate_prob = torch.gather(target_probs, 1, candidate_token).squeeze(-1)  # [batch]

            # Acceptance criterion (probabilistic)
            if self.config.acceptance_method == "probabilistic":
                # Accept with probability min(1, p_target / p_draft)
                # For simplicity, accept if prob > threshold
                accept = candidate_prob >= self.config.acceptance_threshold
            elif self.config.acceptance_method == "greedy":
                # Accept if candidate is the argmax of target
                target_argmax = target_probs.argmax(dim=-1, keepdim=True)  # [batch, 1]
                accept = (candidate_token == target_argmax).squeeze(-1)
            else:
                # Default: accept if prob above threshold
                accept = candidate_prob >= self.config.acceptance_threshold

            # For batch, we accept if ALL batch items accept (for simplicity)
            # In practice, handle per-batch-item
            if accept.all():
                accepted_tokens.append(candidate_token)
                num_accepted += 1
            else:
                # Rejection: stop accepting further candidates
                if self.config.early_exit:
                    break

        if accepted_tokens:
            accepted_tokens = torch.cat(accepted_tokens, dim=1)  # [batch, num_accepted]
        else:
            accepted_tokens = torch.empty(batch_size, 0, dtype=torch.long, device=device)

        return num_accepted, accepted_tokens

    def reset_stats(self):
        """Reset acceptance statistics."""
        self.acceptance_stats = {
            'total_proposed': 0,
            'total_accepted': 0,
            'acceptance_rate': 0.0
        }

    def get_speedup_estimate(self) -> float:
        """
        Estimate speedup factor from current acceptance rate.

        Returns:
            Estimated speedup multiplier
        """
        if self.acceptance_stats['acceptance_rate'] == 0:
            return 1.0

        # Expected speedup = K * acceptance_rate
        # where K is number of speculative tokens
        K = self.config.num_speculative_tokens
        acceptance_rate = self.acceptance_stats['acceptance_rate']

        # Average accepted tokens per target model call
        avg_accepted = K * acceptance_rate

        # Speedup compared to regular generation
        # Regular: 1 token per target call
        # Speculative: avg_accepted tokens per target call
        speedup = avg_accepted

        return speedup


def create_draft_model_from_checkpoint(
    target_model,
    checkpoint_path: Optional[str] = None,
    method: str = "layer_subset",
    num_layers: Optional[int] = None
) -> 'ContinuousAutoregressiveModel':
    """
    Create a draft model for speculative decoding.

    The draft model should be:
    - Significantly faster than target model (smaller)
    - Trained on similar data
    - Compatible architecture

    Args:
        target_model: Target model
        checkpoint_path: Path to draft model checkpoint (if using separate model)
        method: Creation method ("layer_subset", "checkpoint")
        num_layers: Number of layers (for layer_subset)

    Returns:
        Draft model
    """
    if method == "layer_subset":
        # Use first N layers of target model
        from MK3.calm.continuous_model import ContinuousAutoregressiveModel

        num_layers = num_layers or max(1, target_model.num_layers // 3)

        draft_model = ContinuousAutoregressiveModel(
            vocab_size=target_model.vocab_size,
            embedding_dim=target_model.embedding_dim // 2,  # Smaller embedding
            vector_dim=target_model.vector_dim,
            chunk_size=target_model.chunk_size,
            num_layers=num_layers,
            num_heads=max(1, target_model.num_heads // 2),  # Fewer heads
            max_seq_length=target_model.max_seq_length
        )

        # Copy autoencoder
        draft_model.autoencoder.load_state_dict(
            target_model.autoencoder.state_dict()
        )

        return draft_model

    elif method == "checkpoint":
        if checkpoint_path is None:
            raise ValueError("checkpoint_path required for checkpoint method")

        # Load from checkpoint
        checkpoint = torch.load(checkpoint_path, map_location='cpu')
        # Assume checkpoint contains model config
        # Build and load model
        # (Implementation depends on checkpoint format)
        raise NotImplementedError("Checkpoint loading not implemented")

    else:
        raise ValueError(f"Unknown method: {method}")


def benchmark_speculative_decoding(
    target_model,
    draft_model,
    test_prompts: List[str],
    tokenizer,
    max_new_tokens: int = 100,
    num_speculative_values: List[int] = [1, 2, 4, 8],
    device: Optional[torch.device] = None
) -> Dict[str, any]:
    """
    Benchmark speculative decoding with different settings.

    Args:
        target_model: Target model
        draft_model: Draft model
        test_prompts: List of test prompts
        tokenizer: Tokenizer
        max_new_tokens: Maximum tokens to generate
        num_speculative_values: List of K values to test
        device: Device for computation

    Returns:
        Benchmark results
    """
    import time

    device = device or next(target_model.parameters()).device
    results = {}

    for K in num_speculative_values:
        config = SpeculativeConfig(num_speculative_tokens=K)
        decoder = SpeculativeDecoder(target_model, draft_model, config)

        total_time = 0
        total_tokens = 0

        for prompt in test_prompts:
            # Encode prompt
            prompt_ids = tokenizer.encode(prompt, add_special_tokens=True)
            prompt_tensor = torch.tensor([prompt_ids], dtype=torch.long, device=device)

            # Generate with speculative decoding
            start_time = time.time()
            output = decoder.generate_speculative(
                initial_tokens=prompt_tensor,
                max_new_tokens=max_new_tokens,
                return_stats=True
            )
            end_time = time.time()

            total_time += (end_time - start_time)
            total_tokens += output['tokens_generated']

        avg_time = total_time / len(test_prompts)
        tokens_per_sec = total_tokens / total_time if total_time > 0 else 0

        results[f"K={K}"] = {
            'avg_time': avg_time,
            'tokens_per_sec': tokens_per_sec,
            'acceptance_rate': decoder.acceptance_stats['acceptance_rate'],
            'estimated_speedup': decoder.get_speedup_estimate()
        }

    return results
