"""
Continuous embedding layer for CALM integration.

Handles both discrete token embeddings and continuous vector representations,
enabling seamless transitions between token-based and vector-based processing.

MODERNIZED:
- RMSNorm instead of LayerNorm for better efficiency
"""

import torch
import torch.nn as nn
import torch.nn.functional as F
from typing import Optional, Tuple, Union
from .normalization import RMSNorm


class ContinuousEmbedding(nn.Module):
    """
    Embedding layer that supports both discrete tokens and continuous vectors.

    Can:
    1. Embed discrete tokens (standard embedding)
    2. Process continuous vectors (learned projection)
    3. Interpolate between discrete and continuous representations
    """

    def __init__(
        self,
        vocab_size: int,
        embedding_dim: int,
        max_seq_length: int = 2048,
        dropout: float = 0.1,
        padding_idx: Optional[int] = None,
        continuous_mode: bool = False,
    ):
        super().__init__()

        self.vocab_size = vocab_size
        self.embedding_dim = embedding_dim
        self.max_seq_length = max_seq_length
        self.continuous_mode = continuous_mode

        # Token embeddings (for discrete tokens)
        self.token_embedding = nn.Embedding(
            vocab_size,
            embedding_dim,
            padding_idx=padding_idx
        )

        # Position embeddings
        self.position_embedding = nn.Embedding(max_seq_length, embedding_dim)

        # Continuous projection (for continuous vectors)
        # Projects continuous vectors to embedding space
        self.continuous_proj = nn.Sequential(
            nn.Linear(embedding_dim, embedding_dim * 2),
            RMSNorm(embedding_dim * 2),
            nn.GELU(),
            nn.Dropout(dropout),
            nn.Linear(embedding_dim * 2, embedding_dim),
            RMSNorm(embedding_dim)
        )

        # Mode indicator embedding (helps model distinguish discrete vs continuous)
        self.mode_embedding = nn.Parameter(torch.randn(2, embedding_dim) * 0.02)

        self.dropout = nn.Dropout(dropout)
        self.norm = RMSNorm(embedding_dim)

    def forward(
        self,
        input_data: torch.Tensor,
        is_continuous: bool = False,
        position_ids: Optional[torch.Tensor] = None,
        add_position_embeddings: bool = True,
    ) -> torch.Tensor:
        """
        Forward pass supporting both discrete tokens and continuous vectors.

        Args:
            input_data: Either:
                - [batch, seq_len] discrete token IDs (is_continuous=False)
                - [batch, seq_len, embed_dim] continuous vectors (is_continuous=True)
            is_continuous: Whether input is continuous vectors or discrete tokens
            position_ids: [batch, seq_len] position indices (optional)
            add_position_embeddings: Whether to add positional embeddings

        Returns:
            embeddings: [batch, seq_len, embed_dim]
        """
        if is_continuous:
            # Input is continuous vectors
            batch_size, seq_len, embed_dim = input_data.shape
            assert embed_dim == self.embedding_dim, \
                f"Continuous input dim {embed_dim} != embedding_dim {self.embedding_dim}"

            # Project continuous vectors
            embeddings = self.continuous_proj(input_data)

            # Add continuous mode indicator
            mode_embed = self.mode_embedding[1].unsqueeze(0).unsqueeze(0)  # [1, 1, embed_dim]
            embeddings = embeddings + mode_embed

        else:
            # Input is discrete tokens
            batch_size, seq_len = input_data.shape

            # Token embeddings
            embeddings = self.token_embedding(input_data)

            # Add discrete mode indicator
            mode_embed = self.mode_embedding[0].unsqueeze(0).unsqueeze(0)  # [1, 1, embed_dim]
            embeddings = embeddings + mode_embed

        # Add positional embeddings
        if add_position_embeddings:
            if position_ids is None:
                position_ids = torch.arange(
                    seq_len,
                    dtype=torch.long,
                    device=input_data.device
                ).unsqueeze(0).expand(batch_size, -1)

            position_embeds = self.position_embedding(position_ids)
            embeddings = embeddings + position_embeds

        # Layer norm and dropout
        embeddings = self.norm(embeddings)
        embeddings = self.dropout(embeddings)

        return embeddings

    def decode_to_logits(
        self,
        hidden_states: torch.Tensor,
        is_continuous: bool = False
    ) -> torch.Tensor:
        """
        Decode hidden states to vocabulary logits.

        Args:
            hidden_states: [batch, seq_len, embed_dim]
            is_continuous: Whether to decode for continuous or discrete output

        Returns:
            logits: [batch, seq_len, vocab_size]
        """
        if is_continuous:
            # For continuous output, just return the hidden states
            # (will be used by autoencoder decoder)
            return hidden_states
        else:
            # For discrete output, compute logits via embedding matrix
            # Use weight tying: logits = hidden @ embedding.T
            logits = F.linear(hidden_states, self.token_embedding.weight)
            return logits

    def interpolate_embeddings(
        self,
        discrete_tokens: torch.Tensor,
        continuous_vectors: torch.Tensor,
        alpha: float = 0.5
    ) -> torch.Tensor:
        """
        Interpolate between discrete token embeddings and continuous vectors.

        Useful for curriculum learning: start with discrete, gradually move to continuous.

        Args:
            discrete_tokens: [batch, seq_len] token IDs
            continuous_vectors: [batch, seq_len, embed_dim] continuous vectors
            alpha: Interpolation weight (0 = discrete, 1 = continuous)

        Returns:
            interpolated: [batch, seq_len, embed_dim]
        """
        # Get discrete embeddings
        discrete_embeds = self.forward(discrete_tokens, is_continuous=False)

        # Get continuous embeddings
        continuous_embeds = self.forward(continuous_vectors, is_continuous=True)

        # Interpolate
        interpolated = (1 - alpha) * discrete_embeds + alpha * continuous_embeds

        return interpolated


class HybridEmbedding(nn.Module):
    """
    Hybrid embedding that can process sequences with mixed discrete/continuous elements.

    This is useful for models that operate on both token sequences and compressed
    continuous representations simultaneously.
    """

    def __init__(
        self,
        vocab_size: int,
        embedding_dim: int,
        max_seq_length: int = 2048,
        dropout: float = 0.1,
    ):
        super().__init__()

        self.base_embedding = ContinuousEmbedding(
            vocab_size=vocab_size,
            embedding_dim=embedding_dim,
            max_seq_length=max_seq_length,
            dropout=dropout
        )

        # Gating mechanism to blend discrete and continuous
        self.gate = nn.Sequential(
            nn.Linear(embedding_dim * 2, embedding_dim),
            RMSNorm(embedding_dim),
            nn.GELU(),
            nn.Linear(embedding_dim, 1),
            nn.Sigmoid()
        )

    def forward(
        self,
        discrete_tokens: Optional[torch.Tensor] = None,
        continuous_vectors: Optional[torch.Tensor] = None,
        discrete_mask: Optional[torch.Tensor] = None,
        continuous_mask: Optional[torch.Tensor] = None,
    ) -> torch.Tensor:
        """
        Process hybrid input with both discrete and continuous elements.

        Args:
            discrete_tokens: [batch, seq_len] token IDs (optional)
            continuous_vectors: [batch, seq_len, embed_dim] vectors (optional)
            discrete_mask: [batch, seq_len] mask for discrete tokens (optional)
            continuous_mask: [batch, seq_len] mask for continuous vectors (optional)

        Returns:
            embeddings: [batch, seq_len, embed_dim]
        """
        assert discrete_tokens is not None or continuous_vectors is not None, \
            "Must provide either discrete_tokens or continuous_vectors"

        if discrete_tokens is not None and continuous_vectors is not None:
            # Both provided - blend them
            discrete_embeds = self.base_embedding(discrete_tokens, is_continuous=False)
            continuous_embeds = self.base_embedding(continuous_vectors, is_continuous=True)

            # Compute gating weights
            combined = torch.cat([discrete_embeds, continuous_embeds], dim=-1)
            gate_weights = self.gate(combined)  # [batch, seq_len, 1]

            # Blend
            embeddings = gate_weights * continuous_embeds + (1 - gate_weights) * discrete_embeds

            # Apply masks if provided
            if discrete_mask is not None:
                embeddings = embeddings * discrete_mask.unsqueeze(-1)
            if continuous_mask is not None:
                embeddings = embeddings * continuous_mask.unsqueeze(-1)

        elif discrete_tokens is not None:
            # Only discrete
            embeddings = self.base_embedding(discrete_tokens, is_continuous=False)
            if discrete_mask is not None:
                embeddings = embeddings * discrete_mask.unsqueeze(-1)

        else:
            # Only continuous
            embeddings = self.base_embedding(continuous_vectors, is_continuous=True)
            if continuous_mask is not None:
                embeddings = embeddings * continuous_mask.unsqueeze(-1)

        return embeddings
