"""
Advanced Normalization Layers for MK3

Implements RMSNorm (Root Mean Square Layer Normalization) which is more efficient
and stable than LayerNorm, as used in modern architectures like LLaMA and GPT-NeoX.

RMSNorm simplifies LayerNorm by removing mean centering and bias, using only
the root mean square for re-scaling. This provides:
- Faster computation (no mean calculation)
- Better numerical stability
- Comparable or better performance
- Reduced memory footprint
"""

import torch
import torch.nn as nn
import torch.nn.functional as F
from typing import Optional, Tuple


class RMSNorm(nn.Module):
    """
    Root Mean Square Layer Normalization.

    Normalizes inputs using RMS instead of mean and variance:
        y = (x / RMS(x)) * scale
    where RMS(x) = sqrt(mean(x^2) + eps)

    Used in modern LLMs like LLaMA, GPT-NeoX, and others for improved
    efficiency and stability over LayerNorm.

    Args:
        normalized_shape: Input shape from an expected input size
        eps: Small value to avoid division by zero
        elementwise_affine: Whether to learn scaling parameter
    """

    def __init__(
        self,
        normalized_shape: int,
        eps: float = 1e-6,
        elementwise_affine: bool = True,
    ):
        super().__init__()

        self.normalized_shape = (normalized_shape,)
        self.eps = eps
        self.elementwise_affine = elementwise_affine

        if self.elementwise_affine:
            self.weight = nn.Parameter(torch.ones(normalized_shape))
        else:
            self.register_parameter('weight', None)

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        """
        Apply RMSNorm to input tensor.

        Args:
            x: Input tensor of shape [..., normalized_shape]

        Returns:
            Normalized tensor of same shape as input
        """
        # Compute RMS: sqrt(mean(x^2) + eps)
        # Keep dims for broadcasting
        variance = x.pow(2).mean(-1, keepdim=True)
        x = x * torch.rsqrt(variance + self.eps)

        # Apply learned scaling if enabled
        if self.elementwise_affine:
            x = x * self.weight

        return x

    def extra_repr(self) -> str:
        """String representation for debugging."""
        return f'{self.normalized_shape[0]}, eps={self.eps}, elementwise_affine={self.elementwise_affine}'


class RMSNormGated(nn.Module):
    """
    Gated RMSNorm with learnable gating mechanism.

    Combines RMSNorm with a learned gate that can modulate the normalization:
        y = gate * RMSNorm(x) + (1 - gate) * x

    This allows the model to learn when to apply normalization strongly
    vs. when to preserve the original signal.

    Args:
        normalized_shape: Input shape from an expected input size
        eps: Small value to avoid division by zero
        gate_init: Initial value for gate (0.0 to 1.0)
    """

    def __init__(
        self,
        normalized_shape: int,
        eps: float = 1e-6,
        gate_init: float = 1.0,
    ):
        super().__init__()

        self.normalized_shape = (normalized_shape,)
        self.eps = eps

        # Learnable parameters
        self.weight = nn.Parameter(torch.ones(normalized_shape))
        self.gate = nn.Parameter(torch.full((normalized_shape,), gate_init))

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        """
        Apply gated RMSNorm.

        Args:
            x: Input tensor of shape [..., normalized_shape]

        Returns:
            Gated normalized tensor of same shape as input
        """
        # Compute RMS normalization
        variance = x.pow(2).mean(-1, keepdim=True)
        x_norm = x * torch.rsqrt(variance + self.eps)
        x_norm = x_norm * self.weight

        # Apply learned gate
        gate = torch.sigmoid(self.gate)
        output = gate * x_norm + (1 - gate) * x

        return output

    def extra_repr(self) -> str:
        """String representation for debugging."""
        return f'{self.normalized_shape[0]}, eps={self.eps}'


class AdaptiveRMSNorm(nn.Module):
    """
    Adaptive RMSNorm with conditional scaling.

    Allows external conditioning signal to modulate the normalization scale,
    useful for conditional generation or style transfer.

    Args:
        normalized_shape: Input shape from an expected input size
        condition_dim: Dimension of conditioning signal
        eps: Small value to avoid division by zero
    """

    def __init__(
        self,
        normalized_shape: int,
        condition_dim: Optional[int] = None,
        eps: float = 1e-6,
    ):
        super().__init__()

        self.normalized_shape = (normalized_shape,)
        self.eps = eps

        # Base scale parameter
        self.weight = nn.Parameter(torch.ones(normalized_shape))

        # Conditional modulation
        if condition_dim is not None:
            self.condition_proj = nn.Sequential(
                nn.Linear(condition_dim, normalized_shape * 2),
                nn.SiLU(),
                nn.Linear(normalized_shape * 2, normalized_shape)
            )
        else:
            self.condition_proj = None

    def forward(
        self,
        x: torch.Tensor,
        condition: Optional[torch.Tensor] = None
    ) -> torch.Tensor:
        """
        Apply adaptive RMSNorm.

        Args:
            x: Input tensor of shape [..., normalized_shape]
            condition: Optional conditioning tensor of shape [..., condition_dim]

        Returns:
            Conditionally normalized tensor of same shape as input
        """
        # Compute RMS normalization
        variance = x.pow(2).mean(-1, keepdim=True)
        x_norm = x * torch.rsqrt(variance + self.eps)

        # Compute scale
        scale = self.weight

        if condition is not None and self.condition_proj is not None:
            # Add conditional scale modulation
            cond_scale = self.condition_proj(condition)
            if cond_scale.dim() == 2 and x_norm.dim() == 3:
                cond_scale = cond_scale.unsqueeze(1)
            scale = scale * (1 + cond_scale)

        x_norm = x_norm * scale

        return x_norm

    def extra_repr(self) -> str:
        """String representation for debugging."""
        has_condition = self.condition_proj is not None
        return f'{self.normalized_shape[0]}, eps={self.eps}, conditional={has_condition}'


class LayerScale(nn.Module):
    """
    Layer Scale (Touvron et al., 2021).

    Applies learnable per-channel scaling after normalization, helping with
    training stability in very deep networks.

    Args:
        dim: Feature dimension
        init_values: Initial scale values (small values like 1e-5 or 1e-6)
    """

    def __init__(self, dim: int, init_values: float = 1e-5):
        super().__init__()
        self.gamma = nn.Parameter(torch.ones(dim) * init_values)

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        """Apply layer scaling."""
        return x * self.gamma


def replace_layernorm_with_rmsnorm(module: nn.Module, eps: float = 1e-6) -> nn.Module:
    """
    Recursively replace all LayerNorm modules with RMSNorm in a model.

    Args:
        module: PyTorch module to process
        eps: Epsilon value for RMSNorm

    Returns:
        Modified module with RMSNorm instead of LayerNorm
    """
    for name, child in module.named_children():
        if isinstance(child, nn.LayerNorm):
            # Replace with RMSNorm of same dimension
            old_shape = child.normalized_shape[0]
            new_norm = RMSNorm(old_shape, eps=eps)

            # Copy weight if it exists
            if child.weight is not None:
                new_norm.weight.data.copy_(child.weight.data)

            setattr(module, name, new_norm)
        else:
            # Recursively process child modules
            replace_layernorm_with_rmsnorm(child, eps=eps)

    return module
