"""
Advanced Feed-Forward Network Variants for MK3

Implements modern FFN architectures including:
- SwiGLU (Swish + Gated Linear Unit) - used in PaLM, LLaMA
- GEGLU (GELU + Gated Linear Unit) - used in GLU Variants Improve Transformer
- GeGLU (Gated GELU) - variant of GEGLU
- Standard ReLU/GELU FFN for comparison

Gated variants have been shown to improve model capacity and training
dynamics compared to standard FFNs.
"""

import torch
import torch.nn as nn
import torch.nn.functional as F
from typing import Optional, Literal
from .normalization import RMSNorm


class SwiGLU(nn.Module):
    """
    SwiGLU Feed-Forward Network.

    SwiGLU(x) = (Swish(xW) ⊙ xV)U
    where Swish(x) = x * sigmoid(x), ⊙ is element-wise product

    Used in PaLM and LLaMA models. Provides better performance than
    standard FFN with similar parameter count.

    Args:
        dim: Input/output dimension
        hidden_dim: Hidden layer dimension (typically 4*dim, but can be adjusted)
        dropout: Dropout probability
        bias: Whether to use bias in linear layers
    """

    def __init__(
        self,
        dim: int,
        hidden_dim: Optional[int] = None,
        dropout: float = 0.0,
        bias: bool = False,
    ):
        super().__init__()

        hidden_dim = hidden_dim or int(dim * 8 / 3)  # Adjusted for gating
        # Make hidden_dim multiple of 256 for efficiency
        hidden_dim = ((hidden_dim + 255) // 256) * 256

        self.w1 = nn.Linear(dim, hidden_dim, bias=bias)  # Gate projection
        self.w2 = nn.Linear(hidden_dim, dim, bias=bias)  # Output projection
        self.w3 = nn.Linear(dim, hidden_dim, bias=bias)  # Up projection

        self.dropout = nn.Dropout(dropout) if dropout > 0 else nn.Identity()

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        """
        Apply SwiGLU transformation.

        Args:
            x: Input tensor of shape [..., dim]

        Returns:
            Output tensor of shape [..., dim]
        """
        # SwiGLU: (Swish(xW1) ⊙ xW3)W2
        gate = F.silu(self.w1(x))  # Swish/SiLU activation
        up = self.w3(x)
        hidden = gate * up  # Element-wise gating
        output = self.w2(hidden)
        output = self.dropout(output)

        return output


class GEGLU(nn.Module):
    """
    GEGLU (GELU + Gated Linear Unit) Feed-Forward Network.

    GEGLU(x) = (GELU(xW) ⊙ xV)U

    Shown to improve Transformer performance in "GLU Variants Improve Transformer".
    Used in various modern architectures.

    Args:
        dim: Input/output dimension
        hidden_dim: Hidden layer dimension
        dropout: Dropout probability
        bias: Whether to use bias in linear layers
        approximate_gelu: Use approximate GELU for faster computation
    """

    def __init__(
        self,
        dim: int,
        hidden_dim: Optional[int] = None,
        dropout: float = 0.0,
        bias: bool = False,
        approximate_gelu: bool = False,
    ):
        super().__init__()

        hidden_dim = hidden_dim or int(dim * 8 / 3)
        hidden_dim = ((hidden_dim + 255) // 256) * 256

        self.w1 = nn.Linear(dim, hidden_dim, bias=bias)
        self.w2 = nn.Linear(hidden_dim, dim, bias=bias)
        self.w3 = nn.Linear(dim, hidden_dim, bias=bias)

        self.dropout = nn.Dropout(dropout) if dropout > 0 else nn.Identity()
        self.approximate_gelu = approximate_gelu

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        """
        Apply GEGLU transformation.

        Args:
            x: Input tensor of shape [..., dim]

        Returns:
            Output tensor of shape [..., dim]
        """
        # GEGLU: (GELU(xW1) ⊙ xW3)W2
        if self.approximate_gelu:
            gate = F.gelu(self.w1(x), approximate='tanh')
        else:
            gate = F.gelu(self.w1(x))

        up = self.w3(x)
        hidden = gate * up
        output = self.w2(hidden)
        output = self.dropout(output)

        return output


class GeGLU(nn.Module):
    """
    GeGLU (Gated GELU) - Alternative GEGLU formulation.

    GeGLU(x) = GELU(xW) ⊙ xV

    Slightly different from GEGLU in that it doesn't have the output projection
    inside the gating operation. Can be more efficient in some cases.

    Args:
        dim: Input/output dimension
        hidden_dim: Hidden layer dimension
        dropout: Dropout probability
        bias: Whether to use bias in linear layers
    """

    def __init__(
        self,
        dim: int,
        hidden_dim: Optional[int] = None,
        dropout: float = 0.0,
        bias: bool = False,
    ):
        super().__init__()

        hidden_dim = hidden_dim or (dim * 4)

        self.gate_proj = nn.Linear(dim, hidden_dim, bias=bias)
        self.up_proj = nn.Linear(dim, hidden_dim, bias=bias)
        self.down_proj = nn.Linear(hidden_dim, dim, bias=bias)

        self.dropout = nn.Dropout(dropout) if dropout > 0 else nn.Identity()

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        """
        Apply GeGLU transformation.

        Args:
            x: Input tensor of shape [..., dim]

        Returns:
            Output tensor of shape [..., dim]
        """
        gate = F.gelu(self.gate_proj(x))
        up = self.up_proj(x)
        hidden = gate * up
        output = self.down_proj(hidden)
        output = self.dropout(output)

        return output


class StandardFFN(nn.Module):
    """
    Standard Feed-Forward Network with configurable activation.

    FFN(x) = activation(xW1)W2

    Supports ReLU, GELU, and other standard activations.

    Args:
        dim: Input/output dimension
        hidden_dim: Hidden layer dimension (typically 4*dim)
        dropout: Dropout probability
        activation: Activation function ('relu', 'gelu', 'swish', etc.)
        bias: Whether to use bias in linear layers
    """

    def __init__(
        self,
        dim: int,
        hidden_dim: Optional[int] = None,
        dropout: float = 0.0,
        activation: str = 'gelu',
        bias: bool = True,
    ):
        super().__init__()

        hidden_dim = hidden_dim or (dim * 4)

        self.w1 = nn.Linear(dim, hidden_dim, bias=bias)
        self.w2 = nn.Linear(hidden_dim, dim, bias=bias)
        self.dropout = nn.Dropout(dropout) if dropout > 0 else nn.Identity()

        # Select activation function
        if activation == 'relu':
            self.activation = nn.ReLU()
        elif activation == 'gelu':
            self.activation = nn.GELU()
        elif activation == 'swish' or activation == 'silu':
            self.activation = nn.SiLU()
        elif activation == 'tanh':
            self.activation = nn.Tanh()
        else:
            raise ValueError(f"Unknown activation: {activation}")

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        """
        Apply standard FFN transformation.

        Args:
            x: Input tensor of shape [..., dim]

        Returns:
            Output tensor of shape [..., dim]
        """
        hidden = self.activation(self.w1(x))
        output = self.w2(hidden)
        output = self.dropout(output)

        return output


class FeedForwardNetwork(nn.Module):
    """
    Configurable Feed-Forward Network with multiple variant support.

    Unified interface for different FFN types with optional pre-normalization.

    Args:
        dim: Input/output dimension
        hidden_dim: Hidden layer dimension
        dropout: Dropout probability
        ffn_type: Type of FFN ('swiglu', 'geglu', 'geglu_v2', 'gelu', 'relu')
        bias: Whether to use bias in linear layers
        pre_norm: Whether to apply RMSNorm before FFN
    """

    def __init__(
        self,
        dim: int,
        hidden_dim: Optional[int] = None,
        dropout: float = 0.0,
        ffn_type: Literal['swiglu', 'geglu', 'geglu_v2', 'gelu', 'relu', 'swish'] = 'swiglu',
        bias: bool = False,
        pre_norm: bool = True,
    ):
        super().__init__()

        self.ffn_type = ffn_type
        self.pre_norm = pre_norm

        # Pre-normalization
        if pre_norm:
            self.norm = RMSNorm(dim)
        else:
            self.norm = nn.Identity()

        # Select FFN variant
        if ffn_type == 'swiglu':
            self.ffn = SwiGLU(dim, hidden_dim, dropout, bias)
        elif ffn_type == 'geglu':
            self.ffn = GEGLU(dim, hidden_dim, dropout, bias)
        elif ffn_type == 'geglu_v2':
            self.ffn = GeGLU(dim, hidden_dim, dropout, bias)
        elif ffn_type in ['gelu', 'relu', 'swish']:
            self.ffn = StandardFFN(dim, hidden_dim, dropout, ffn_type, bias)
        else:
            raise ValueError(f"Unknown FFN type: {ffn_type}")

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        """
        Apply FFN with optional pre-normalization.

        Args:
            x: Input tensor of shape [..., dim]

        Returns:
            Output tensor of shape [..., dim]
        """
        normalized = self.norm(x)
        output = self.ffn(normalized)
        return output

    def extra_repr(self) -> str:
        """String representation for debugging."""
        return f'ffn_type={self.ffn_type}, pre_norm={self.pre_norm}'


class ParallelFFN(nn.Module):
    """
    Parallel Feed-Forward Network (as in some modern architectures).

    Processes input through FFN in parallel with attention, rather than sequentially.
    Can improve training efficiency and model capacity.

    Args:
        dim: Input/output dimension
        hidden_dim: Hidden layer dimension
        dropout: Dropout probability
        ffn_type: Type of FFN to use
    """

    def __init__(
        self,
        dim: int,
        hidden_dim: Optional[int] = None,
        dropout: float = 0.0,
        ffn_type: str = 'swiglu',
    ):
        super().__init__()

        self.norm = RMSNorm(dim)
        self.ffn = FeedForwardNetwork(
            dim=dim,
            hidden_dim=hidden_dim,
            dropout=dropout,
            ffn_type=ffn_type,
            pre_norm=False  # Already normalized
        )

    def forward(
        self,
        x: torch.Tensor,
        attn_output: torch.Tensor
    ) -> torch.Tensor:
        """
        Apply parallel FFN.

        Args:
            x: Original input (pre-attention) of shape [..., dim]
            attn_output: Attention output of shape [..., dim]

        Returns:
            Combined output of shape [..., dim]
        """
        # Normalize once for both branches
        normalized = self.norm(x)

        # FFN on normalized input (parallel to attention)
        ffn_out = self.ffn(normalized)

        # Combine: x + attn_out + ffn_out
        output = x + attn_output + ffn_out

        return output


class ExpertFFN(nn.Module):
    """
    Mixture-of-Experts style FFN with simple top-k routing.

    Uses multiple FFN experts and routes each token to top-k experts.
    Simplified version for demonstration - full MoE would need load balancing.

    Args:
        dim: Input/output dimension
        hidden_dim: Hidden layer dimension
        num_experts: Number of expert FFNs
        top_k: Number of experts to route each token to
        dropout: Dropout probability
        ffn_type: Type of FFN for experts
    """

    def __init__(
        self,
        dim: int,
        hidden_dim: Optional[int] = None,
        num_experts: int = 8,
        top_k: int = 2,
        dropout: float = 0.0,
        ffn_type: str = 'swiglu',
    ):
        super().__init__()

        self.num_experts = num_experts
        self.top_k = top_k

        # Router network
        self.router = nn.Linear(dim, num_experts)

        # Expert FFNs
        self.experts = nn.ModuleList([
            FeedForwardNetwork(
                dim=dim,
                hidden_dim=hidden_dim,
                dropout=dropout,
                ffn_type=ffn_type,
                pre_norm=False
            )
            for _ in range(num_experts)
        ])

        self.norm = RMSNorm(dim)

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        """
        Apply expert FFN with routing.

        Args:
            x: Input tensor of shape [batch, seq_len, dim]

        Returns:
            Output tensor of shape [batch, seq_len, dim]
        """
        original_shape = x.shape
        x_flat = x.view(-1, x.shape[-1])  # [batch*seq_len, dim]

        # Normalize
        x_norm = self.norm(x_flat)

        # Route to experts
        router_logits = self.router(x_norm)  # [batch*seq_len, num_experts]
        router_probs = F.softmax(router_logits, dim=-1)

        # Select top-k experts
        top_k_probs, top_k_indices = torch.topk(router_probs, self.top_k, dim=-1)
        top_k_probs = top_k_probs / top_k_probs.sum(dim=-1, keepdim=True)  # Normalize

        # Apply experts
        output = torch.zeros_like(x_flat)
        for i in range(self.top_k):
            expert_idx = top_k_indices[:, i]
            expert_prob = top_k_probs[:, i:i+1]

            # Process each token with its selected expert
            for expert_id in range(self.num_experts):
                mask = (expert_idx == expert_id)
                if mask.any():
                    expert_input = x_norm[mask]
                    expert_output = self.experts[expert_id](expert_input)
                    output[mask] += expert_prob[mask] * expert_output

        # Reshape back
        output = output.view(original_shape)

        return output
