"""
Query-Key Normalization for Attention Mechanisms

Implements QK-Normalization which normalizes queries and keys before computing
attention scores, improving training stability and model performance.

Used in modern architectures like:
- ViT-22B (Dehghani et al., 2023)
- Gemma (Google, 2024)
- InternLM2 (2024)

Benefits:
- Prevents attention logits from growing too large
- Improves training stability, especially at scale
- Better gradient flow in deep networks
- Reduces need for careful initialization
"""

import torch
import torch.nn as nn
import torch.nn.functional as F
from typing import Optional, Tuple
from .normalization import RMSNorm


class QKNormalization(nn.Module):
    """
    Query-Key Normalization for attention.

    Normalizes query and key vectors before computing attention scores,
    preventing logits from becoming too large and improving stability.

    Three normalization strategies:
    1. 'l2': L2 normalization (cosine attention)
    2. 'rmsnorm': RMSNorm on queries and keys
    3. 'layernorm': LayerNorm on queries and keys (less common)

    Args:
        dim: Dimension of query/key vectors
        norm_type: Type of normalization ('l2', 'rmsnorm', 'layernorm')
        eps: Small value for numerical stability
    """

    def __init__(
        self,
        dim: int,
        norm_type: str = 'rmsnorm',
        eps: float = 1e-6,
    ):
        super().__init__()

        self.dim = dim
        self.norm_type = norm_type
        self.eps = eps

        if norm_type == 'rmsnorm':
            self.query_norm = RMSNorm(dim, eps=eps)
            self.key_norm = RMSNorm(dim, eps=eps)
        elif norm_type == 'layernorm':
            self.query_norm = nn.LayerNorm(dim, eps=eps)
            self.key_norm = nn.LayerNorm(dim, eps=eps)
        elif norm_type == 'l2':
            self.query_norm = None
            self.key_norm = None
        else:
            raise ValueError(f"Unknown norm_type: {norm_type}")

    def forward(
        self,
        query: torch.Tensor,
        key: torch.Tensor
    ) -> Tuple[torch.Tensor, torch.Tensor]:
        """
        Normalize queries and keys.

        Args:
            query: Query tensor of shape [..., seq_len_q, dim]
            key: Key tensor of shape [..., seq_len_k, dim]

        Returns:
            normalized_query: Normalized query tensor
            normalized_key: Normalized key tensor
        """
        if self.norm_type == 'l2':
            # L2 normalization (cosine similarity)
            query = F.normalize(query, p=2, dim=-1, eps=self.eps)
            key = F.normalize(key, p=2, dim=-1, eps=self.eps)
        else:
            # RMSNorm or LayerNorm
            query = self.query_norm(query)
            key = self.key_norm(key)

        return query, key

    def extra_repr(self) -> str:
        """String representation for debugging."""
        return f'dim={self.dim}, norm_type={self.norm_type}, eps={self.eps}'


class QKNormAttention(nn.Module):
    """
    Multi-head attention with QK-Normalization.

    Standard multi-head attention but with normalized queries and keys,
    improving training stability and performance.

    Args:
        embed_dim: Total dimension of the model
        num_heads: Number of attention heads
        dropout: Dropout probability
        bias: Whether to use bias in projections
        qk_norm_type: Type of QK normalization
        qk_scale: Optional scaling factor for attention scores
    """

    def __init__(
        self,
        embed_dim: int,
        num_heads: int = 8,
        dropout: float = 0.0,
        bias: bool = False,
        qk_norm_type: str = 'rmsnorm',
        qk_scale: Optional[float] = None,
    ):
        super().__init__()

        assert embed_dim % num_heads == 0, "embed_dim must be divisible by num_heads"

        self.embed_dim = embed_dim
        self.num_heads = num_heads
        self.head_dim = embed_dim // num_heads
        self.qk_scale = qk_scale or (self.head_dim ** -0.5)

        # Q, K, V projections
        self.q_proj = nn.Linear(embed_dim, embed_dim, bias=bias)
        self.k_proj = nn.Linear(embed_dim, embed_dim, bias=bias)
        self.v_proj = nn.Linear(embed_dim, embed_dim, bias=bias)

        # Output projection
        self.out_proj = nn.Linear(embed_dim, embed_dim, bias=bias)

        # QK normalization
        self.qk_norm = QKNormalization(
            dim=self.head_dim,
            norm_type=qk_norm_type
        )

        # Dropout
        self.attn_dropout = nn.Dropout(dropout)
        self.out_dropout = nn.Dropout(dropout)

    def forward(
        self,
        query: torch.Tensor,
        key: torch.Tensor,
        value: torch.Tensor,
        mask: Optional[torch.Tensor] = None,
        return_attention: bool = False,
    ) -> Tuple[torch.Tensor, Optional[torch.Tensor]]:
        """
        Apply multi-head attention with QK normalization.

        Args:
            query: Query tensor of shape [batch, seq_len_q, embed_dim]
            key: Key tensor of shape [batch, seq_len_k, embed_dim]
            value: Value tensor of shape [batch, seq_len_k, embed_dim]
            mask: Attention mask of shape [batch, seq_len_q, seq_len_k] or broadcastable
            return_attention: Whether to return attention weights

        Returns:
            output: Output tensor of shape [batch, seq_len_q, embed_dim]
            attention_weights: Optional attention weights [batch, num_heads, seq_len_q, seq_len_k]
        """
        batch_size, seq_len_q, _ = query.shape
        seq_len_k = key.shape[1]

        # Project Q, K, V
        Q = self.q_proj(query)  # [batch, seq_len_q, embed_dim]
        K = self.k_proj(key)    # [batch, seq_len_k, embed_dim]
        V = self.v_proj(value)  # [batch, seq_len_k, embed_dim]

        # Reshape for multi-head attention
        Q = Q.view(batch_size, seq_len_q, self.num_heads, self.head_dim)
        K = K.view(batch_size, seq_len_k, self.num_heads, self.head_dim)
        V = V.view(batch_size, seq_len_k, self.num_heads, self.head_dim)

        # Transpose to [batch, num_heads, seq_len, head_dim]
        Q = Q.transpose(1, 2)
        K = K.transpose(1, 2)
        V = V.transpose(1, 2)

        # Apply QK normalization per head
        # Need to reshape for normalization
        Q_normalized = []
        K_normalized = []

        for head_idx in range(self.num_heads):
            q_head = Q[:, head_idx, :, :]  # [batch, seq_len_q, head_dim]
            k_head = K[:, head_idx, :, :]  # [batch, seq_len_k, head_dim]

            q_norm, k_norm = self.qk_norm(q_head, k_head)

            Q_normalized.append(q_norm)
            K_normalized.append(k_norm)

        # Stack back
        Q = torch.stack(Q_normalized, dim=1)  # [batch, num_heads, seq_len_q, head_dim]
        K = torch.stack(K_normalized, dim=1)  # [batch, num_heads, seq_len_k, head_dim]

        # Compute attention scores
        attn_scores = torch.matmul(Q, K.transpose(-2, -1))  # [batch, num_heads, seq_len_q, seq_len_k]
        attn_scores = attn_scores * self.qk_scale

        # Apply mask if provided
        if mask is not None:
            if mask.dim() == 2:
                # [seq_len_q, seq_len_k] -> [batch, 1, seq_len_q, seq_len_k]
                mask = mask.unsqueeze(0).unsqueeze(0)
            elif mask.dim() == 3:
                # [batch, seq_len_q, seq_len_k] -> [batch, 1, seq_len_q, seq_len_k]
                mask = mask.unsqueeze(1)

            attn_scores = attn_scores.masked_fill(~mask.bool(), float('-inf'))

        # Softmax and dropout
        attn_weights = F.softmax(attn_scores, dim=-1)  # [batch, num_heads, seq_len_q, seq_len_k]
        attn_weights = self.attn_dropout(attn_weights)

        # Apply attention to values
        output = torch.matmul(attn_weights, V)  # [batch, num_heads, seq_len_q, head_dim]

        # Reshape back
        output = output.transpose(1, 2).contiguous()  # [batch, seq_len_q, num_heads, head_dim]
        output = output.view(batch_size, seq_len_q, self.embed_dim)

        # Output projection
        output = self.out_proj(output)
        output = self.out_dropout(output)

        if return_attention:
            return output, attn_weights
        else:
            return output, None


class QKNormCrossAttention(nn.Module):
    """
    Cross-attention with QK-Normalization.

    Used for encoder-decoder attention or any cross-attention scenario.
    Queries come from one sequence, keys and values from another.

    Args:
        embed_dim: Total dimension of the model
        num_heads: Number of attention heads
        dropout: Dropout probability
        bias: Whether to use bias in projections
        qk_norm_type: Type of QK normalization
    """

    def __init__(
        self,
        embed_dim: int,
        num_heads: int = 8,
        dropout: float = 0.0,
        bias: bool = False,
        qk_norm_type: str = 'rmsnorm',
    ):
        super().__init__()

        assert embed_dim % num_heads == 0, "embed_dim must be divisible by num_heads"

        self.embed_dim = embed_dim
        self.num_heads = num_heads
        self.head_dim = embed_dim // num_heads
        self.scale = self.head_dim ** -0.5

        # Projections
        self.q_proj = nn.Linear(embed_dim, embed_dim, bias=bias)
        self.k_proj = nn.Linear(embed_dim, embed_dim, bias=bias)
        self.v_proj = nn.Linear(embed_dim, embed_dim, bias=bias)
        self.out_proj = nn.Linear(embed_dim, embed_dim, bias=bias)

        # QK normalization
        self.qk_norm = QKNormalization(
            dim=self.head_dim,
            norm_type=qk_norm_type
        )

        # Dropout
        self.attn_dropout = nn.Dropout(dropout)
        self.out_dropout = nn.Dropout(dropout)

    def forward(
        self,
        query: torch.Tensor,
        context: torch.Tensor,
        mask: Optional[torch.Tensor] = None,
        return_attention: bool = False,
    ) -> Tuple[torch.Tensor, Optional[torch.Tensor]]:
        """
        Apply cross-attention with QK normalization.

        Args:
            query: Query tensor from target sequence [batch, seq_len_q, embed_dim]
            context: Context tensor from source sequence [batch, seq_len_ctx, embed_dim]
            mask: Attention mask [batch, seq_len_q, seq_len_ctx]
            return_attention: Whether to return attention weights

        Returns:
            output: Output tensor [batch, seq_len_q, embed_dim]
            attention_weights: Optional attention weights
        """
        batch_size, seq_len_q, _ = query.shape
        seq_len_ctx = context.shape[1]

        # Project Q from query, K and V from context
        Q = self.q_proj(query)     # [batch, seq_len_q, embed_dim]
        K = self.k_proj(context)   # [batch, seq_len_ctx, embed_dim]
        V = self.v_proj(context)   # [batch, seq_len_ctx, embed_dim]

        # Reshape for multi-head
        Q = Q.view(batch_size, seq_len_q, self.num_heads, self.head_dim).transpose(1, 2)
        K = K.view(batch_size, seq_len_ctx, self.num_heads, self.head_dim).transpose(1, 2)
        V = V.view(batch_size, seq_len_ctx, self.num_heads, self.head_dim).transpose(1, 2)

        # Apply QK normalization per head
        Q_normalized = []
        K_normalized = []

        for head_idx in range(self.num_heads):
            q_head = Q[:, head_idx, :, :]
            k_head = K[:, head_idx, :, :]

            q_norm, k_norm = self.qk_norm(q_head, k_head)

            Q_normalized.append(q_norm)
            K_normalized.append(k_norm)

        Q = torch.stack(Q_normalized, dim=1)
        K = torch.stack(K_normalized, dim=1)

        # Compute attention
        attn_scores = torch.matmul(Q, K.transpose(-2, -1)) * self.scale

        if mask is not None:
            if mask.dim() == 2:
                mask = mask.unsqueeze(0).unsqueeze(0)
            elif mask.dim() == 3:
                mask = mask.unsqueeze(1)
            attn_scores = attn_scores.masked_fill(~mask.bool(), float('-inf'))

        attn_weights = F.softmax(attn_scores, dim=-1)
        attn_weights = self.attn_dropout(attn_weights)

        output = torch.matmul(attn_weights, V)
        output = output.transpose(1, 2).contiguous().view(batch_size, seq_len_q, self.embed_dim)

        output = self.out_proj(output)
        output = self.out_dropout(output)

        if return_attention:
            return output, attn_weights
        else:
            return output, None


class AdaptiveQKNorm(nn.Module):
    """
    Adaptive QK normalization with learnable temperature.

    Allows the model to learn the optimal scale for attention scores
    after QK normalization.

    Args:
        dim: Dimension of query/key vectors
        norm_type: Type of normalization
        init_temperature: Initial temperature value
        learnable_temperature: Whether temperature is learnable
    """

    def __init__(
        self,
        dim: int,
        norm_type: str = 'rmsnorm',
        init_temperature: float = 1.0,
        learnable_temperature: bool = True,
    ):
        super().__init__()

        self.qk_norm = QKNormalization(dim, norm_type)

        # Learnable temperature parameter
        if learnable_temperature:
            self.temperature = nn.Parameter(
                torch.tensor(init_temperature).log()  # Log for positivity
            )
        else:
            self.register_buffer('temperature', torch.tensor(init_temperature).log())

    def forward(
        self,
        query: torch.Tensor,
        key: torch.Tensor
    ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
        """
        Apply adaptive QK normalization.

        Args:
            query: Query tensor
            key: Key tensor

        Returns:
            normalized_query: Normalized query
            normalized_key: Normalized key
            temperature: Current temperature value
        """
        query, key = self.qk_norm(query, key)

        # Apply learned temperature
        temp = self.temperature.exp()

        return query, key, temp

    def extra_repr(self) -> str:
        """String representation for debugging."""
        temp_value = self.temperature.exp().item() if isinstance(self.temperature, nn.Parameter) else self.temperature.exp().item()
        return f'temperature={temp_value:.4f}'
