"""
Speculative Decoding for CALM

Accelerates generation by K-2x using a smaller "draft" model to predict
multiple continuous vectors ahead, then verifying with the main model.

Key innovation: Speculative decoding in continuous space (not discrete tokens).

Algorithm:
1. Draft model predicts K vectors ahead
2. Main model verifies in parallel
3. Accept matching vectors, reject diverging ones
4. Expected speedup: 2-3x on top of CALM's K-token speedup
"""

import torch
import torch.nn as nn
import torch.nn.functional as F
from typing import Optional, Tuple, Dict, List
import math


class DraftCALMModel(nn.Module):
    """
    Lightweight draft model for speculative decoding.

    Much smaller than main model (e.g., 6 layers vs 12 layers)
    to enable fast parallel prediction.
    """

    def __init__(
        self,
        vector_dim: int = 1024,
        embedding_dim: int = 512,  # Smaller than main model
        num_layers: int = 6,       # Fewer layers than main model
        num_heads: int = 8,
        feedforward_dim: Optional[int] = None,
        dropout: float = 0.1,
    ):
        super().__init__()

        self.vector_dim = vector_dim
        self.embedding_dim = embedding_dim
        self.num_layers = num_layers
        feedforward_dim = feedforward_dim or (embedding_dim * 4)

        # Project continuous vectors to embedding space
        self.vector_to_embed = nn.Sequential(
            nn.Linear(vector_dim, embedding_dim),
            nn.LayerNorm(embedding_dim),
            nn.GELU(),
            nn.Dropout(dropout),
        )

        # Lightweight transformer layers
        encoder_layer = nn.TransformerEncoderLayer(
            d_model=embedding_dim,
            nhead=num_heads,
            dim_feedforward=feedforward_dim,
            dropout=dropout,
            activation='gelu',
            batch_first=True,
            norm_first=True
        )
        self.transformer = nn.TransformerEncoder(
            encoder_layer,
            num_layers=num_layers
        )

        # Project back to vector space
        self.embed_to_vector = nn.Sequential(
            nn.Linear(embedding_dim, vector_dim),
            nn.LayerNorm(vector_dim),
            nn.GELU(),
            nn.Linear(vector_dim, vector_dim)
        )

        # Position embeddings
        self.register_buffer(
            'position_encoding',
            self._create_position_encoding(2048, embedding_dim)
        )

        self.dropout = nn.Dropout(dropout)

    def _create_position_encoding(self, max_len: int, d_model: int) -> torch.Tensor:
        """Create sinusoidal position encodings."""
        position = torch.arange(max_len).unsqueeze(1)
        div_term = torch.exp(torch.arange(0, d_model, 2) * (-math.log(10000.0) / d_model))
        pe = torch.zeros(max_len, d_model)
        pe[:, 0::2] = torch.sin(position * div_term)
        pe[:, 1::2] = torch.cos(position * div_term)
        return pe

    def forward(
        self,
        continuous_vectors: torch.Tensor,
        causal_mask: Optional[torch.Tensor] = None
    ) -> torch.Tensor:
        """
        Fast forward pass for draft predictions.

        Args:
            continuous_vectors: [batch, seq_len, vector_dim] input vectors
            causal_mask: [seq_len, seq_len] causal attention mask

        Returns:
            predicted_vectors: [batch, seq_len, vector_dim] predicted next vectors
        """
        batch_size, seq_len, _ = continuous_vectors.shape

        # Project to embedding space
        embeddings = self.vector_to_embed(continuous_vectors)

        # Add position encodings
        position_enc = self.position_encoding[:seq_len, :].unsqueeze(0).expand(batch_size, -1, -1)
        embeddings = embeddings + position_enc
        embeddings = self.dropout(embeddings)

        # Create causal mask if not provided
        if causal_mask is None:
            causal_mask = torch.triu(
                torch.ones(seq_len, seq_len, device=embeddings.device, dtype=torch.bool),
                diagonal=1
            )

        # Transformer encoding
        hidden = self.transformer(embeddings, mask=causal_mask, is_causal=True)

        # Project to vector space
        predicted_vectors = self.embed_to_vector(hidden)

        return predicted_vectors


class SpeculativeCALMDecoder:
    """
    Speculative decoder for CALM models.

    Uses a small draft model to predict K vectors ahead, then verifies
    with the main model in parallel. Achieves 2-3x speedup on top of
    CALM's K-token speedup.

    Total speedup: K * 2-3 = 16-24x for K=8!
    """

    def __init__(
        self,
        main_model: nn.Module,
        draft_model: DraftCALMModel,
        autoencoder: nn.Module,
        lookahead_steps: int = 4,
        acceptance_threshold: float = 0.9,  # Cosine similarity threshold
        temperature: float = 1.0,
    ):
        """
        Args:
            main_model: Main CALM model
            draft_model: Lightweight draft model
            autoencoder: Autoencoder for token conversion
            lookahead_steps: Number of vectors to predict ahead
            acceptance_threshold: Cosine similarity threshold for accepting draft
            temperature: Sampling temperature
        """
        self.main_model = main_model
        self.draft_model = draft_model
        self.autoencoder = autoencoder
        self.lookahead_steps = lookahead_steps
        self.acceptance_threshold = acceptance_threshold
        self.temperature = temperature

        # Statistics
        self.total_draft_predictions = 0
        self.accepted_predictions = 0

    @torch.no_grad()
    def generate(
        self,
        initial_tokens: torch.Tensor,
        max_new_vectors: int = 10,
        temperature: Optional[float] = None,
        top_p: float = 0.9,
        top_k: Optional[int] = None,
    ) -> Tuple[torch.Tensor, Dict]:
        """
        Generate with speculative decoding.

        Args:
            initial_tokens: [batch, seq_len] initial tokens
            max_new_vectors: Number of vectors to generate
            temperature: Sampling temperature (overrides default)
            top_p: Nucleus sampling threshold
            top_k: Top-k sampling threshold

        Returns:
            generated_tokens: [batch, total_seq_len] generated tokens
            stats: Dict with generation statistics
        """
        self.main_model.eval()
        self.draft_model.eval()

        temperature = temperature or self.temperature
        batch_size = initial_tokens.size(0)

        # Convert initial tokens to vectors
        continuous_vectors, _ = self.main_model.tokenize_to_vectors(initial_tokens)
        num_vectors_generated = 0

        # Statistics
        total_draft = 0
        total_accepted = 0
        verification_steps = 0

        while num_vectors_generated < max_new_vectors:
            # Step 1: Draft model predicts lookahead_steps vectors
            draft_predictions = self._draft_predict(
                continuous_vectors,
                num_steps=min(self.lookahead_steps, max_new_vectors - num_vectors_generated)
            )  # [batch, lookahead_steps, vector_dim]

            # Step 2: Main model verifies all predictions in parallel
            accepted_vectors, num_accepted = self._verify_and_accept(
                continuous_vectors,
                draft_predictions,
                temperature=temperature
            )

            # Update statistics
            total_draft += draft_predictions.size(1)
            total_accepted += num_accepted
            verification_steps += 1

            # Append accepted vectors
            continuous_vectors = torch.cat([continuous_vectors, accepted_vectors], dim=1)
            num_vectors_generated += num_accepted

            # If no vectors accepted, generate one with main model
            if num_accepted == 0:
                # Fallback: single prediction with main model
                output = self.main_model(continuous_vectors, use_cache=True)
                next_vector = output['predicted_vectors'][:, -1:, :]

                # Apply sampling
                if temperature != 1.0:
                    noise = torch.randn_like(next_vector) * temperature
                    next_vector = next_vector + noise
                    next_vector = F.normalize(next_vector, p=2, dim=-1) * math.sqrt(self.main_model.vector_dim)

                continuous_vectors = torch.cat([continuous_vectors, next_vector], dim=1)
                num_vectors_generated += 1

        # Convert vectors to tokens
        generated_tokens = self.main_model.vectors_to_tokens(continuous_vectors, return_logits=False)

        # Statistics
        acceptance_rate = total_accepted / total_draft if total_draft > 0 else 0.0
        effective_speedup = (1 + acceptance_rate * self.lookahead_steps) if acceptance_rate > 0 else 1.0

        stats = {
            'total_draft_predictions': total_draft,
            'accepted_predictions': total_accepted,
            'acceptance_rate': acceptance_rate,
            'verification_steps': verification_steps,
            'effective_speedup': effective_speedup,
            'theoretical_max_speedup': 1 + self.lookahead_steps,
        }

        return generated_tokens, stats

    def _draft_predict(
        self,
        continuous_vectors: torch.Tensor,
        num_steps: int
    ) -> torch.Tensor:
        """
        Use draft model to predict multiple vectors ahead.

        Args:
            continuous_vectors: [batch, seq_len, vector_dim] current sequence
            num_steps: Number of vectors to predict ahead

        Returns:
            draft_predictions: [batch, num_steps, vector_dim] predicted vectors
        """
        predictions = []

        # Autoregressively predict with draft model
        current_sequence = continuous_vectors

        for _ in range(num_steps):
            # Draft model forward pass
            draft_output = self.draft_model(current_sequence)  # [batch, seq_len, vector_dim]

            # Get last prediction
            next_vector = draft_output[:, -1:, :]  # [batch, 1, vector_dim]
            predictions.append(next_vector)

            # Append to sequence for next prediction
            current_sequence = torch.cat([current_sequence, next_vector], dim=1)

        # Stack predictions: [batch, num_steps, vector_dim]
        draft_predictions = torch.cat(predictions, dim=1)

        return draft_predictions

    def _verify_and_accept(
        self,
        continuous_vectors: torch.Tensor,
        draft_predictions: torch.Tensor,
        temperature: float = 1.0
    ) -> Tuple[torch.Tensor, int]:
        """
        Verify draft predictions with main model and accept matching ones.

        Args:
            continuous_vectors: [batch, seq_len, vector_dim] current sequence
            draft_predictions: [batch, num_draft, vector_dim] draft predictions
            temperature: Sampling temperature

        Returns:
            accepted_vectors: [batch, num_accepted, vector_dim] accepted vectors
            num_accepted: Number of accepted vectors
        """
        batch_size, num_draft, vector_dim = draft_predictions.shape

        # Concatenate current sequence with draft predictions
        extended_sequence = torch.cat([continuous_vectors, draft_predictions], dim=1)

        # Main model forward pass (verifies all drafts in parallel)
        output = self.main_model(extended_sequence, use_cache=True)
        main_predictions = output['predicted_vectors']  # [batch, total_len, vector_dim]

        # Get main model's predictions for the positions where draft predicted
        main_pred_vectors = main_predictions[:, -num_draft-1:-1, :]  # [batch, num_draft, vector_dim]

        # Compare draft vs main predictions using cosine similarity
        draft_norm = F.normalize(draft_predictions, p=2, dim=-1)
        main_norm = F.normalize(main_pred_vectors, p=2, dim=-1)

        # Cosine similarity: [batch, num_draft]
        similarities = (draft_norm * main_norm).sum(dim=-1)

        # Accept predictions above threshold
        # We accept sequentially until first rejection
        num_accepted = 0
        for i in range(num_draft):
            if (similarities[:, i] >= self.acceptance_threshold).all():
                num_accepted = i + 1
            else:
                break  # Stop at first rejection

        if num_accepted > 0:
            # Use main model's predictions (more accurate than draft)
            accepted_vectors = main_pred_vectors[:, :num_accepted, :]
        else:
            # No vectors accepted
            accepted_vectors = torch.empty(batch_size, 0, vector_dim, device=draft_predictions.device)

        return accepted_vectors, num_accepted

    def get_acceptance_rate(self) -> float:
        """Get current acceptance rate."""
        if self.total_draft_predictions == 0:
            return 0.0
        return self.accepted_predictions / self.total_draft_predictions

    def reset_stats(self):
        """Reset statistics."""
        self.total_draft_predictions = 0
        self.accepted_predictions = 0


def train_draft_model(
    draft_model: DraftCALMModel,
    main_model: nn.Module,
    train_loader,
    num_epochs: int = 5,
    learning_rate: float = 1e-4,
    device: torch.device = None,
) -> Dict:
    """
    Train draft model to mimic main model's predictions.

    The draft model learns to approximate the main model using knowledge distillation.

    Args:
        draft_model: Draft model to train
        main_model: Main model (teacher)
        train_loader: DataLoader for training data
        num_epochs: Number of training epochs
        learning_rate: Learning rate
        device: Device for training

    Returns:
        metrics: Dict with training metrics
    """
    device = device or torch.device('cuda' if torch.cuda.is_available() else 'cpu')
    draft_model.to(device)
    main_model.to(device)
    main_model.eval()  # Main model is frozen (teacher)

    optimizer = torch.optim.AdamW(draft_model.parameters(), lr=learning_rate, weight_decay=0.01)

    for epoch in range(num_epochs):
        epoch_loss = 0.0
        epoch_cosine_sim = 0.0
        num_batches = 0

        for batch in train_loader:
            token_ids, mask = batch
            token_ids = token_ids.to(device)

            # Convert to continuous vectors
            with torch.no_grad():
                continuous_vectors, _ = main_model.tokenize_to_vectors(token_ids)

                # Get main model predictions (teacher)
                main_output = main_model(continuous_vectors, use_cache=False)
                main_predictions = main_output['predicted_vectors']

            # Draft model predictions (student)
            draft_predictions = draft_model(continuous_vectors)

            # Distillation loss: MSE + cosine similarity
            mse_loss = F.mse_loss(draft_predictions, main_predictions)
            cosine_sim = F.cosine_similarity(draft_predictions, main_predictions, dim=-1).mean()
            loss = mse_loss + 0.1 * (1 - cosine_sim)

            # Backward pass
            optimizer.zero_grad()
            loss.backward()
            torch.nn.utils.clip_grad_norm_(draft_model.parameters(), max_norm=1.0)
            optimizer.step()

            epoch_loss += loss.item()
            epoch_cosine_sim += cosine_sim.item()
            num_batches += 1

        # Epoch summary
        avg_loss = epoch_loss / num_batches
        avg_cosine_sim = epoch_cosine_sim / num_batches
        print(f"Epoch {epoch}: Loss={avg_loss:.4f}, CosineSim={avg_cosine_sim:.4f}")

    metrics = {
        'final_loss': avg_loss,
        'final_cosine_similarity': avg_cosine_sim,
    }

    return metrics
