"""
Mixture-of-Experts (MoE) Attention for MK3

Implements sparse MoE attention with top-k routing for improved efficiency and capacity.

Key features:
- Top-k expert selection per token
- Load balancing loss
- Expert capacity handling
- Differentiable routing with Gumbel-Softmax
- Auxiliary loss for load balancing

References:
- Switch Transformers: https://arxiv.org/abs/2101.03961
- GShard: https://arxiv.org/abs/2006.16668
- Expert Choice Routing: https://arxiv.org/abs/2202.09368
"""

import torch
import torch.nn as nn
import torch.nn.functional as F
import math
from typing import Optional, Tuple, Dict


class TopKRouter(nn.Module):
    """
    Top-K router for selecting experts per token.

    Implements load-balanced routing with auxiliary loss to encourage
    uniform expert utilization.
    """

    def __init__(
        self,
        dim: int,
        num_experts: int,
        top_k: int = 2,
        capacity_factor: float = 1.25,
        noisy_gating: bool = True,
        noise_std: float = 0.1,
    ):
        super().__init__()

        self.dim = dim
        self.num_experts = num_experts
        self.top_k = top_k
        self.capacity_factor = capacity_factor
        self.noisy_gating = noisy_gating
        self.noise_std = noise_std

        # Router network
        self.router = nn.Linear(dim, num_experts, bias=False)

        # Initialize router with small weights for stability
        nn.init.normal_(self.router.weight, mean=0.0, std=0.01)

    def add_noise(self, logits: torch.Tensor) -> torch.Tensor:
        """Add noise to routing logits during training."""
        if self.training and self.noisy_gating:
            noise = torch.randn_like(logits) * self.noise_std
            return logits + noise
        return logits

    def forward(
        self,
        x: torch.Tensor,
        use_aux_loss: bool = True,
    ) -> Tuple[torch.Tensor, torch.Tensor, Optional[torch.Tensor]]:
        """
        Route tokens to top-k experts.

        Args:
            x: [batch, seq_len, dim] input tokens
            use_aux_loss: Whether to compute auxiliary load balancing loss

        Returns:
            expert_indices: [batch, seq_len, top_k] selected expert indices
            routing_weights: [batch, seq_len, top_k] routing weights (softmax over top-k)
            aux_loss: Optional auxiliary loss for load balancing
        """
        batch_size, seq_len, dim = x.shape

        # Compute routing logits: [batch, seq_len, num_experts]
        logits = self.router(x)

        # Add noise during training for exploration
        logits = self.add_noise(logits)

        # Top-k selection
        top_k_logits, top_k_indices = torch.topk(logits, self.top_k, dim=-1)
        # top_k_logits: [batch, seq_len, top_k]
        # top_k_indices: [batch, seq_len, top_k]

        # Compute routing weights (softmax over top-k)
        routing_weights = F.softmax(top_k_logits, dim=-1)

        # Compute auxiliary load balancing loss
        aux_loss = None
        if use_aux_loss and self.training:
            aux_loss = self._compute_load_balancing_loss(logits, top_k_indices)

        return top_k_indices, routing_weights, aux_loss

    def _compute_load_balancing_loss(
        self,
        logits: torch.Tensor,
        top_k_indices: torch.Tensor,
    ) -> torch.Tensor:
        """
        Compute load balancing auxiliary loss.

        Encourages uniform distribution of tokens across experts.

        Args:
            logits: [batch, seq_len, num_experts] routing logits
            top_k_indices: [batch, seq_len, top_k] selected expert indices

        Returns:
            loss: Scalar auxiliary loss
        """
        batch_size, seq_len, num_experts = logits.shape

        # Compute fraction of tokens routed to each expert
        # [batch, seq_len, top_k] -> [batch * seq_len * top_k]
        flat_indices = top_k_indices.reshape(-1)
        expert_counts = torch.bincount(
            flat_indices,
            minlength=num_experts
        ).float()
        expert_fraction = expert_counts / expert_counts.sum()

        # Compute routing probability for each expert (mean of softmax)
        routing_probs = F.softmax(logits, dim=-1)  # [batch, seq_len, num_experts]
        routing_probs_mean = routing_probs.mean(dim=[0, 1])  # [num_experts]

        # Load balancing loss: encourage expert_fraction to match routing_probs_mean
        # CV^2 (coefficient of variation squared) encourages uniform distribution
        loss = num_experts * torch.sum(expert_fraction * routing_probs_mean)

        return loss


class MoEAttentionExpert(nn.Module):
    """
    Single expert for MoE attention.

    Each expert is a complete multi-head attention module.
    """

    def __init__(
        self,
        embed_dim: int,
        num_heads: int,
        dropout: float = 0.0,
    ):
        super().__init__()

        self.embed_dim = embed_dim
        self.num_heads = num_heads
        self.head_dim = embed_dim // num_heads

        assert embed_dim % num_heads == 0

        # QKV projection
        self.qkv_proj = nn.Linear(embed_dim, 3 * embed_dim, bias=False)

        # Output projection
        self.out_proj = nn.Linear(embed_dim, embed_dim, bias=False)

        self.dropout = nn.Dropout(dropout) if dropout > 0 else None

        self.scale = 1.0 / math.sqrt(self.head_dim)

    def forward(
        self,
        x: torch.Tensor,
        attention_mask: Optional[torch.Tensor] = None,
    ) -> torch.Tensor:
        """
        Expert forward pass.

        Args:
            x: [batch, seq_len, embed_dim]
            attention_mask: [batch, seq_len, seq_len]

        Returns:
            output: [batch, seq_len, embed_dim]
        """
        batch_size, seq_len, _ = x.shape

        # Project to Q, K, V
        qkv = self.qkv_proj(x)  # [batch, seq_len, 3 * embed_dim]
        qkv = qkv.reshape(batch_size, seq_len, 3, self.num_heads, self.head_dim)
        qkv = qkv.permute(2, 0, 3, 1, 4)  # [3, batch, num_heads, seq_len, head_dim]
        q, k, v = qkv[0], qkv[1], qkv[2]

        # Attention scores
        attn_scores = torch.matmul(q, k.transpose(-2, -1)) * self.scale

        # Apply mask
        if attention_mask is not None:
            if attention_mask.dim() == 3:
                attention_mask = attention_mask.unsqueeze(1)
            attn_scores = attn_scores.masked_fill(attention_mask == 0, float('-inf'))

        # Softmax
        attn_weights = F.softmax(attn_scores, dim=-1)

        if self.dropout is not None:
            attn_weights = self.dropout(attn_weights)

        # Apply attention to values
        output = torch.matmul(attn_weights, v)

        # Reshape and project
        output = output.transpose(1, 2).contiguous().view(batch_size, seq_len, self.embed_dim)
        output = self.out_proj(output)

        return output


class MoEAttention(nn.Module):
    """
    Mixture-of-Experts Multi-Head Attention.

    Routes tokens to different attention experts based on learned routing.
    """

    def __init__(
        self,
        embed_dim: int,
        num_heads: int,
        num_experts: int = 8,
        top_k: int = 2,
        dropout: float = 0.0,
        capacity_factor: float = 1.25,
        noisy_gating: bool = True,
        aux_loss_weight: float = 0.01,
    ):
        super().__init__()

        self.embed_dim = embed_dim
        self.num_heads = num_heads
        self.num_experts = num_experts
        self.top_k = top_k
        self.dropout = dropout
        self.aux_loss_weight = aux_loss_weight

        # Router
        self.router = TopKRouter(
            dim=embed_dim,
            num_experts=num_experts,
            top_k=top_k,
            capacity_factor=capacity_factor,
            noisy_gating=noisy_gating,
        )

        # Expert modules
        self.experts = nn.ModuleList([
            MoEAttentionExpert(
                embed_dim=embed_dim,
                num_heads=num_heads,
                dropout=dropout,
            )
            for _ in range(num_experts)
        ])

        # Layer norm
        self.layer_norm = nn.LayerNorm(embed_dim)

    def forward(
        self,
        x: torch.Tensor,
        attention_mask: Optional[torch.Tensor] = None,
        return_aux_loss: bool = True,
    ) -> Tuple[torch.Tensor, Optional[torch.Tensor]]:
        """
        MoE attention forward pass.

        Args:
            x: [batch, seq_len, embed_dim] input
            attention_mask: [batch, seq_len, seq_len] attention mask
            return_aux_loss: Whether to return auxiliary load balancing loss

        Returns:
            output: [batch, seq_len, embed_dim] attended output
            aux_loss: Optional auxiliary loss for load balancing
        """
        batch_size, seq_len, embed_dim = x.shape
        residual = x

        # Normalize input
        x = self.layer_norm(x)

        # Route tokens to experts
        expert_indices, routing_weights, aux_loss = self.router(x, use_aux_loss=return_aux_loss)
        # expert_indices: [batch, seq_len, top_k]
        # routing_weights: [batch, seq_len, top_k]

        # Initialize output
        output = torch.zeros_like(x)

        # Process tokens through selected experts
        # For efficiency, we batch process tokens by expert
        for expert_idx in range(self.num_experts):
            # Find tokens routed to this expert
            # [batch, seq_len, top_k] -> [batch, seq_len, top_k]
            expert_mask = (expert_indices == expert_idx)  # [batch, seq_len, top_k]

            if not expert_mask.any():
                continue

            # Get routing weights for this expert
            # [batch, seq_len, top_k] -> [batch, seq_len]
            expert_weights = torch.where(
                expert_mask,
                routing_weights,
                torch.zeros_like(routing_weights)
            ).sum(dim=-1, keepdim=True)  # [batch, seq_len, 1]

            # Process through expert (process all tokens, will mask later)
            expert_output = self.experts[expert_idx](x, attention_mask)

            # Add weighted expert output
            output = output + expert_weights * expert_output

        # Residual connection
        output = output + residual

        # Scale auxiliary loss
        if aux_loss is not None:
            aux_loss = aux_loss * self.aux_loss_weight

        return output, aux_loss if return_aux_loss else None


class SparseMoEAttention(nn.Module):
    """
    Sparse MoE Attention with expert capacity limits.

    More efficient version that strictly enforces capacity constraints.
    """

    def __init__(
        self,
        embed_dim: int,
        num_heads: int,
        num_experts: int = 8,
        expert_capacity: Optional[int] = None,
        top_k: int = 2,
        dropout: float = 0.0,
        aux_loss_weight: float = 0.01,
    ):
        super().__init__()

        self.embed_dim = embed_dim
        self.num_heads = num_heads
        self.num_experts = num_experts
        self.expert_capacity = expert_capacity
        self.top_k = top_k
        self.aux_loss_weight = aux_loss_weight

        # Router
        self.router = TopKRouter(
            dim=embed_dim,
            num_experts=num_experts,
            top_k=top_k,
        )

        # Experts
        self.experts = nn.ModuleList([
            MoEAttentionExpert(
                embed_dim=embed_dim,
                num_heads=num_heads,
                dropout=dropout,
            )
            for _ in range(num_experts)
        ])

        self.layer_norm = nn.LayerNorm(embed_dim)

    def forward(
        self,
        x: torch.Tensor,
        attention_mask: Optional[torch.Tensor] = None,
        return_aux_loss: bool = True,
    ) -> Tuple[torch.Tensor, Optional[torch.Tensor], Dict]:
        """
        Sparse MoE forward with capacity limits.

        Args:
            x: [batch, seq_len, embed_dim]
            attention_mask: [batch, seq_len, seq_len]
            return_aux_loss: Whether to return auxiliary loss

        Returns:
            output: [batch, seq_len, embed_dim]
            aux_loss: Optional auxiliary loss
            stats: Dictionary with routing statistics
        """
        batch_size, seq_len, embed_dim = x.shape
        residual = x

        # Set expert capacity if not specified
        if self.expert_capacity is None:
            tokens_per_batch = batch_size * seq_len
            capacity = int((tokens_per_batch / self.num_experts) * 1.25)
        else:
            capacity = self.expert_capacity

        # Normalize
        x = self.layer_norm(x)

        # Route tokens
        expert_indices, routing_weights, aux_loss = self.router(x, use_aux_loss=return_aux_loss)

        # Initialize output and statistics
        output = torch.zeros_like(x)
        expert_usage = torch.zeros(self.num_experts, device=x.device)
        tokens_dropped = 0

        # Process each expert with capacity constraints
        for expert_idx in range(self.num_experts):
            # Find tokens routed to this expert
            expert_mask = (expert_indices == expert_idx)
            expert_tokens = expert_mask.any(dim=-1)  # [batch, seq_len]

            # Get number of tokens for this expert
            num_tokens = expert_tokens.sum().item()
            expert_usage[expert_idx] = num_tokens

            if num_tokens == 0:
                continue

            # Apply capacity limit
            if num_tokens > capacity:
                # Keep only top capacity tokens by routing weight
                expert_weights_for_cap = torch.where(
                    expert_mask,
                    routing_weights,
                    torch.zeros_like(routing_weights)
                ).max(dim=-1)[0]  # [batch, seq_len]

                # Get top-capacity tokens
                _, top_indices = torch.topk(expert_weights_for_cap, capacity)
                capacity_mask = torch.zeros_like(expert_tokens)
                capacity_mask.view(-1)[top_indices] = 1
                expert_tokens = expert_tokens & capacity_mask.bool()

                tokens_dropped += (num_tokens - capacity)

            # Get routing weights for selected tokens
            expert_weights = torch.where(
                expert_mask,
                routing_weights,
                torch.zeros_like(routing_weights)
            ).sum(dim=-1, keepdim=True)  # [batch, seq_len, 1]

            # Apply capacity mask to weights
            expert_weights = expert_weights * expert_tokens.unsqueeze(-1).float()

            # Process through expert
            expert_output = self.experts[expert_idx](x, attention_mask)

            # Accumulate weighted output
            output = output + expert_weights * expert_output

        # Residual connection
        output = output + residual

        # Statistics
        stats = {
            'expert_usage': expert_usage.cpu().numpy(),
            'tokens_dropped': tokens_dropped,
            'capacity': capacity,
            'load_balance': expert_usage.std().item() / (expert_usage.mean().item() + 1e-8),
        }

        # Scale auxiliary loss
        if aux_loss is not None:
            aux_loss = aux_loss * self.aux_loss_weight

        return output, aux_loss if return_aux_loss else None, stats


def create_moe_attention(
    embed_dim: int,
    num_heads: int,
    num_experts: int = 8,
    top_k: int = 2,
    dropout: float = 0.0,
    sparse: bool = False,
    **kwargs
) -> nn.Module:
    """
    Factory function for creating MoE attention layers.

    Args:
        embed_dim: Embedding dimension
        num_heads: Number of attention heads per expert
        num_experts: Number of expert modules
        top_k: Number of experts to activate per token
        dropout: Dropout probability
        sparse: Whether to use sparse MoE with capacity limits
        **kwargs: Additional arguments

    Returns:
        MoE attention module
    """
    if sparse:
        return SparseMoEAttention(
            embed_dim=embed_dim,
            num_heads=num_heads,
            num_experts=num_experts,
            top_k=top_k,
            dropout=dropout,
            **kwargs
        )
    else:
        return MoEAttention(
            embed_dim=embed_dim,
            num_heads=num_heads,
            num_experts=num_experts,
            top_k=top_k,
            dropout=dropout,
            **kwargs
        )
