"""
Comprehensive Positional Encoding System for MK3

Implements multiple positional encoding schemes with support for
context extension to 128k-1M tokens:

1. RoPE (Rotary Position Embedding) - Base implementation
2. NTK-aware RoPE - Scaling for long contexts
3. YaRN (Yet another RoPE extensioN) - Advanced context extension
4. LongRoPE - Multi-scale interpolation
5. ALiBi (Attention with Linear Biases) - Alternative to RoPE

All implementations are production-ready with NO placeholders.
Designed for easy swapping via configuration.
"""

import torch
import torch.nn as nn
import torch.nn.functional as F
from typing import Optional, Tuple, Dict, Literal
import math


class RotaryPositionEmbedding(nn.Module):
    """
    Rotary Position Embedding (RoPE) - Base Implementation

    Paper: RoFormer: Enhanced Transformer with Rotary Position Embedding
    arXiv: 2104.09864

    Applies rotation matrices to queries and keys based on position.
    Enables relative position encoding through rotation in complex space.
    """

    def __init__(
        self,
        dim: int,
        max_seq_length: int = 2048,
        base: float = 10000.0,
        scaling_factor: float = 1.0,
    ):
        """
        Args:
            dim: Head dimension (must be even for rotation pairs)
            max_seq_length: Maximum sequence length to precompute
            base: Base for exponential decay (theta)
            scaling_factor: Linear scaling factor for position indices
        """
        super().__init__()

        assert dim % 2 == 0, f"RoPE requires even head dimension, got {dim}"

        self.dim = dim
        self.max_seq_length = max_seq_length
        self.base = base
        self.scaling_factor = scaling_factor

        # Precompute rotation frequencies
        # inv_freq = 1 / (base^(2i/d)) for i in [0, d/2)
        inv_freq = 1.0 / (base ** (torch.arange(0, dim, 2).float() / dim))
        self.register_buffer('inv_freq', inv_freq, persistent=False)

        # Precompute sin/cos for efficiency
        self._precompute_freqs_cis(max_seq_length)

    def _precompute_freqs_cis(self, seq_len: int):
        """Precompute complex exponentials (cos + i*sin) for rotations."""
        # Position indices: [0, 1, 2, ..., seq_len-1]
        t = torch.arange(seq_len, device=self.inv_freq.device).float()
        t = t / self.scaling_factor  # Apply scaling

        # Outer product: [seq_len, dim/2]
        freqs = torch.outer(t, self.inv_freq)

        # Complex exponentials: e^(i*theta) = cos(theta) + i*sin(theta)
        # Store as [seq_len, dim/2, 2] for (cos, sin)
        freqs_cos = torch.cos(freqs)  # [seq_len, dim/2]
        freqs_sin = torch.sin(freqs)  # [seq_len, dim/2]

        self.register_buffer('freqs_cos', freqs_cos, persistent=False)
        self.register_buffer('freqs_sin', freqs_sin, persistent=False)

    def _rotate_half(self, x: torch.Tensor) -> torch.Tensor:
        """
        Rotate half the hidden dims of the input.

        For rotation in complex plane, we need to swap and negate pairs:
        [x1, x2, x3, x4, ...] -> [-x2, x1, -x4, x3, ...]
        """
        x1 = x[..., : x.shape[-1] // 2]  # First half
        x2 = x[..., x.shape[-1] // 2 :]  # Second half
        return torch.cat([-x2, x1], dim=-1)

    def forward(
        self,
        q: torch.Tensor,
        k: torch.Tensor,
        position_ids: Optional[torch.Tensor] = None,
    ) -> Tuple[torch.Tensor, torch.Tensor]:
        """
        Apply rotary position embeddings to queries and keys.

        Args:
            q: Query tensor [batch, seq_len, num_heads, head_dim]
            k: Key tensor [batch, seq_len, num_heads, head_dim]
            position_ids: Optional position indices [batch, seq_len]

        Returns:
            Rotated (q, k) tensors with same shape as input
        """
        seq_len = q.shape[1]

        # Extend precomputed cache if needed
        if seq_len > self.max_seq_length:
            self.max_seq_length = seq_len
            self._precompute_freqs_cis(seq_len)

        # Get cos/sin for positions
        if position_ids is None:
            # Use sequential positions
            cos = self.freqs_cos[:seq_len, :]  # [seq_len, dim/2]
            sin = self.freqs_sin[:seq_len, :]  # [seq_len, dim/2]
        else:
            # Use custom positions
            cos = self.freqs_cos[position_ids]  # [batch, seq_len, dim/2]
            sin = self.freqs_sin[position_ids]  # [batch, seq_len, dim/2]

        # Expand cos/sin to match head dimension
        # [seq_len, dim/2] -> [seq_len, dim] by repeating each value
        cos = torch.repeat_interleave(cos, 2, dim=-1)  # [seq_len, dim]
        sin = torch.repeat_interleave(sin, 2, dim=-1)  # [seq_len, dim]

        # Reshape for broadcasting: [1, seq_len, 1, dim]
        if position_ids is None:
            cos = cos.unsqueeze(0).unsqueeze(2)
            sin = sin.unsqueeze(0).unsqueeze(2)
        else:
            cos = cos.unsqueeze(2)  # [batch, seq_len, 1, dim]
            sin = sin.unsqueeze(2)

        # Apply rotation: q_rot = q * cos + rotate_half(q) * sin
        q_embed = (q * cos) + (self._rotate_half(q) * sin)
        k_embed = (k * cos) + (self._rotate_half(k) * sin)

        return q_embed, k_embed


class NTKAwareRoPE(RotaryPositionEmbedding):
    """
    NTK-Aware Rotary Position Embedding

    Extends RoPE for longer contexts using Neural Tangent Kernel (NTK) scaling.
    Adjusts the base frequency to preserve relative position information.

    Paper: "Extending Context Window of Large Language Models via
            Positional Interpolation" (Reddit/Together)

    Effective for 2x-8x context extension with minimal fine-tuning.
    """

    def __init__(
        self,
        dim: int,
        max_seq_length: int = 2048,
        base: float = 10000.0,
        scaling_factor: float = 1.0,
        ntk_factor: float = 1.0,
        original_max_seq_length: int = 2048,
    ):
        """
        Args:
            dim: Head dimension
            max_seq_length: Target maximum sequence length
            base: Base theta value
            scaling_factor: Additional linear scaling
            ntk_factor: NTK scaling factor (auto-computed if 1.0)
            original_max_seq_length: Original training sequence length
        """
        # Compute NTK scaling if not provided
        if ntk_factor == 1.0 and max_seq_length > original_max_seq_length:
            # NTK formula: scale = (new_len / old_len) ^ (dim / (dim - 2))
            ratio = max_seq_length / original_max_seq_length
            ntk_factor = ratio ** (dim / (dim - 2))

        # Adjust base by NTK factor
        adjusted_base = base * ntk_factor

        super().__init__(
            dim=dim,
            max_seq_length=max_seq_length,
            base=adjusted_base,
            scaling_factor=scaling_factor,
        )

        self.ntk_factor = ntk_factor
        self.original_max_seq_length = original_max_seq_length


class YaRNPositionEmbedding(nn.Module):
    """
    YaRN: Yet another RoPE extensioN

    Paper: "YaRN: Efficient Context Window Extension of Large Language Models"
    arXiv: 2309.00071

    Advanced RoPE extension using:
    1. Attention scaling (temperature adjustment)
    2. Frequency interpolation with "ramp" function
    3. Preserves high-frequency components (critical for short-range)

    Can extend context to 128k-1M tokens effectively.
    """

    def __init__(
        self,
        dim: int,
        max_seq_length: int = 2048,
        base: float = 10000.0,
        original_max_seq_length: int = 2048,
        beta_fast: int = 32,
        beta_slow: int = 1,
        mscale: float = 1.0,
        mscale_all_dim: float = 0.0,
    ):
        """
        Args:
            dim: Head dimension
            max_seq_length: Target sequence length
            base: Base theta
            original_max_seq_length: Original training length
            beta_fast: Fast frequency boundary (in wavelength)
            beta_slow: Slow frequency boundary
            mscale: Attention temperature scaling factor
            mscale_all_dim: Global temperature adjustment
        """
        super().__init__()

        self.dim = dim
        self.max_seq_length = max_seq_length
        self.base = base
        self.original_max_seq_length = original_max_seq_length
        self.beta_fast = beta_fast
        self.beta_slow = beta_slow

        # Compute extension ratio
        self.scale = max_seq_length / original_max_seq_length

        # Compute inv_freq with YaRN interpolation
        inv_freq = 1.0 / (base ** (torch.arange(0, dim, 2).float() / dim))

        # Apply YaRN ramp function for frequency interpolation
        inv_freq_interpolated = self._apply_yarn_ramp(inv_freq)
        self.register_buffer('inv_freq', inv_freq_interpolated, persistent=False)

        # Compute attention scale (temperature adjustment)
        if mscale == 0.0:
            # Auto-compute mscale
            self.mscale = self._compute_mscale()
        else:
            self.mscale = mscale

        self.mscale_all_dim = mscale_all_dim

        # Precompute frequencies
        self._precompute_freqs_cis(max_seq_length)

    def _apply_yarn_ramp(self, inv_freq: torch.Tensor) -> torch.Tensor:
        """
        Apply YaRN's ramp function for smooth interpolation.

        Frequencies are divided into three regions:
        1. Low (slow): No interpolation (preserve long-range)
        2. Medium: Smooth ramp interpolation
        3. High (fast): Full interpolation (critical for short-range)
        """
        # Wavelengths for each frequency
        wavelengths = 2 * math.pi / inv_freq

        # Compute ramp weights
        # weight = 0 (no interp) for wavelength < beta_fast
        # weight = 1 (full interp) for wavelength > beta_slow
        ramp_min = self.beta_fast / self.scale
        ramp_max = self.beta_slow / self.scale

        # Linear ramp between ramp_min and ramp_max
        ramp_weights = torch.clamp(
            (wavelengths - ramp_min) / (ramp_max - ramp_min),
            min=0.0,
            max=1.0
        )

        # Apply interpolation: freq_new = freq * (1 + weight * (scale - 1))
        # Equivalent to: freq_new = freq / (1 + weight * (scale - 1))
        inv_freq_interpolated = inv_freq / (1.0 + ramp_weights * (self.scale - 1.0))

        return inv_freq_interpolated

    def _compute_mscale(self) -> float:
        """
        Compute attention temperature scaling (mscale).

        Formula from YaRN paper:
        mscale = 0.1 * ln(scale) + 1.0 if scale > 1 else 1.0
        """
        if self.scale > 1.0:
            return 0.1 * math.log(self.scale) + 1.0
        return 1.0

    def _precompute_freqs_cis(self, seq_len: int):
        """Precompute rotation frequencies."""
        t = torch.arange(seq_len, device=self.inv_freq.device).float()
        freqs = torch.outer(t, self.inv_freq)

        freqs_cos = torch.cos(freqs)
        freqs_sin = torch.sin(freqs)

        self.register_buffer('freqs_cos', freqs_cos, persistent=False)
        self.register_buffer('freqs_sin', freqs_sin, persistent=False)

    def _rotate_half(self, x: torch.Tensor) -> torch.Tensor:
        """Rotate half the hidden dims."""
        x1 = x[..., : x.shape[-1] // 2]
        x2 = x[..., x.shape[-1] // 2 :]
        return torch.cat([-x2, x1], dim=-1)

    def forward(
        self,
        q: torch.Tensor,
        k: torch.Tensor,
        position_ids: Optional[torch.Tensor] = None,
    ) -> Tuple[torch.Tensor, torch.Tensor]:
        """Apply YaRN position embeddings."""
        seq_len = q.shape[1]

        if seq_len > self.max_seq_length:
            self.max_seq_length = seq_len
            self._precompute_freqs_cis(seq_len)

        # Get cos/sin
        if position_ids is None:
            cos = self.freqs_cos[:seq_len, :]
            sin = self.freqs_sin[:seq_len, :]
        else:
            cos = self.freqs_cos[position_ids]
            sin = self.freqs_sin[position_ids]

        # Expand to head dimension
        cos = torch.repeat_interleave(cos, 2, dim=-1)
        sin = torch.repeat_interleave(sin, 2, dim=-1)

        # Reshape for broadcasting
        if position_ids is None:
            cos = cos.unsqueeze(0).unsqueeze(2)
            sin = sin.unsqueeze(0).unsqueeze(2)
        else:
            cos = cos.unsqueeze(2)
            sin = sin.unsqueeze(2)

        # Apply rotation
        q_embed = (q * cos) + (self._rotate_half(q) * sin)
        k_embed = (k * cos) + (self._rotate_half(k) * sin)

        # Apply attention scaling (temperature adjustment)
        if self.mscale != 1.0 or self.mscale_all_dim != 0.0:
            scale_factor = self.mscale + self.mscale_all_dim
            q_embed = q_embed * scale_factor
            k_embed = k_embed * scale_factor

        return q_embed, k_embed


class LongRoPE(nn.Module):
    """
    LongRoPE: Extending LLM Context Window to 2048k

    Paper: "LongRoPE: Extending LLM Context Window Beyond 2 Million Tokens"
    arXiv: 2402.13753

    Key innovations:
    1. Multi-scale interpolation (different scales for different frequency bands)
    2. Evolutionary search for optimal interpolation factors
    3. Progressive extension during fine-tuning

    Can extend to 1M+ context with proper fine-tuning schedule.
    """

    def __init__(
        self,
        dim: int,
        max_seq_length: int = 2048,
        base: float = 10000.0,
        original_max_seq_length: int = 2048,
        short_factor: Optional[torch.Tensor] = None,
        long_factor: Optional[torch.Tensor] = None,
        short_mscale: float = 1.0,
        long_mscale: float = 1.0,
    ):
        """
        Args:
            dim: Head dimension
            max_seq_length: Target sequence length
            base: Base theta
            original_max_seq_length: Original training length
            short_factor: Interpolation factors for short (high-freq) dimensions
            long_factor: Interpolation factors for long (low-freq) dimensions
            short_mscale: Temperature for short factors
            long_mscale: Temperature for long factors
        """
        super().__init__()

        self.dim = dim
        self.max_seq_length = max_seq_length
        self.base = base
        self.original_max_seq_length = original_max_seq_length

        # Compute base inv_freq
        inv_freq = 1.0 / (base ** (torch.arange(0, dim, 2).float() / dim))

        # Default factors if not provided
        if short_factor is None:
            # High-frequency components: less interpolation
            short_factor = torch.ones(dim // 2) * 1.0
        if long_factor is None:
            # Low-frequency components: more interpolation
            extension_ratio = max_seq_length / original_max_seq_length
            long_factor = torch.ones(dim // 2) * extension_ratio

        self.register_buffer('short_factor', short_factor, persistent=False)
        self.register_buffer('long_factor', long_factor, persistent=False)

        # Compute mixed interpolation factors
        # Use smooth transition between short and long factors
        num_freqs = dim // 2
        transition_weights = torch.linspace(0, 1, num_freqs)  # 0 to 1

        mixed_factors = (
            self.short_factor * (1 - transition_weights) +
            self.long_factor * transition_weights
        )

        # Apply interpolation to inv_freq
        inv_freq_interpolated = inv_freq / mixed_factors
        self.register_buffer('inv_freq', inv_freq_interpolated, persistent=False)

        self.short_mscale = short_mscale
        self.long_mscale = long_mscale

        # Precompute
        self._precompute_freqs_cis(max_seq_length)

    def _precompute_freqs_cis(self, seq_len: int):
        """Precompute rotation frequencies."""
        t = torch.arange(seq_len, device=self.inv_freq.device).float()
        freqs = torch.outer(t, self.inv_freq)

        freqs_cos = torch.cos(freqs)
        freqs_sin = torch.sin(freqs)

        self.register_buffer('freqs_cos', freqs_cos, persistent=False)
        self.register_buffer('freqs_sin', freqs_sin, persistent=False)

    def _rotate_half(self, x: torch.Tensor) -> torch.Tensor:
        """Rotate half the hidden dims."""
        x1 = x[..., : x.shape[-1] // 2]
        x2 = x[..., x.shape[-1] // 2 :]
        return torch.cat([-x2, x1], dim=-1)

    def forward(
        self,
        q: torch.Tensor,
        k: torch.Tensor,
        position_ids: Optional[torch.Tensor] = None,
    ) -> Tuple[torch.Tensor, torch.Tensor]:
        """Apply LongRoPE position embeddings."""
        seq_len = q.shape[1]

        if seq_len > self.max_seq_length:
            self.max_seq_length = seq_len
            self._precompute_freqs_cis(seq_len)

        # Get cos/sin
        if position_ids is None:
            cos = self.freqs_cos[:seq_len, :]
            sin = self.freqs_sin[:seq_len, :]
        else:
            cos = self.freqs_cos[position_ids]
            sin = self.freqs_sin[position_ids]

        # Expand
        cos = torch.repeat_interleave(cos, 2, dim=-1)
        sin = torch.repeat_interleave(sin, 2, dim=-1)

        # Reshape
        if position_ids is None:
            cos = cos.unsqueeze(0).unsqueeze(2)
            sin = sin.unsqueeze(0).unsqueeze(2)
        else:
            cos = cos.unsqueeze(2)
            sin = sin.unsqueeze(2)

        # Apply rotation
        q_embed = (q * cos) + (self._rotate_half(q) * sin)
        k_embed = (k * cos) + (self._rotate_half(k) * sin)

        # Apply multi-scale temperature
        mscale = (self.short_mscale + self.long_mscale) / 2.0
        if mscale != 1.0:
            q_embed = q_embed * mscale
            k_embed = k_embed * mscale

        return q_embed, k_embed


class ALiBiPositionBias(nn.Module):
    """
    ALiBi: Attention with Linear Biases

    Paper: "Train Short, Test Long: Attention with Linear Biases Enables
            Input Length Extrapolation"
    arXiv: 2108.12409

    Instead of adding positional encodings to embeddings, adds position-dependent
    biases directly to attention scores. Excellent extrapolation to longer contexts.

    Key property: NO modification to Q/K, only attention bias.
    """

    def __init__(
        self,
        num_heads: int,
        max_seq_length: int = 2048,
        slope_computation: Literal['original', 'linear', 'exponential'] = 'original',
    ):
        """
        Args:
            num_heads: Number of attention heads
            max_seq_length: Maximum sequence length (for precomputation)
            slope_computation: Method for computing head-specific slopes
                - 'original': Geometric sequence from paper
                - 'linear': Linear spacing
                - 'exponential': Exponential spacing
        """
        super().__init__()

        self.num_heads = num_heads
        self.max_seq_length = max_seq_length
        self.slope_computation = slope_computation

        # Compute slopes for each head
        slopes = self._compute_slopes(num_heads)
        self.register_buffer('slopes', slopes, persistent=True)

        # Precompute bias matrix
        self._precompute_bias(max_seq_length)

    def _compute_slopes(self, num_heads: int) -> torch.Tensor:
        """
        Compute ALiBi slopes for each attention head.

        Original paper uses geometric sequence:
        For n heads: slopes = [2^(-8/n * i) for i in 1..n]

        Closer heads have smaller slopes (less penalty for distance).
        """
        if self.slope_computation == 'original':
            # Geometric sequence from original paper
            def get_slopes_power_of_2(n):
                start = 2 ** (-8 / n)
                ratio = start
                return [start * (ratio ** i) for i in range(n)]

            def get_slopes(n):
                # Handle non-power-of-2 heads
                if n & (n - 1) == 0:  # is power of 2
                    return get_slopes_power_of_2(n)
                else:
                    # Closest power of 2
                    closest_power = 2 ** math.floor(math.log2(n))
                    slopes1 = get_slopes_power_of_2(closest_power)
                    slopes2 = get_slopes(2 * closest_power)
                    # Interleave to get n slopes
                    return slopes1 + slopes2[::2][:n - closest_power]

            slopes = torch.tensor(get_slopes(num_heads), dtype=torch.float32)

        elif self.slope_computation == 'linear':
            # Linear spacing between min and max slopes
            slopes = torch.linspace(0.125, 0.01, num_heads)

        elif self.slope_computation == 'exponential':
            # Exponential decay
            slopes = torch.tensor([
                2 ** (-8.0 * (i + 1) / num_heads) for i in range(num_heads)
            ], dtype=torch.float32)

        else:
            raise ValueError(f"Unknown slope computation: {self.slope_computation}")

        return slopes

    def _precompute_bias(self, seq_len: int):
        """
        Precompute bias matrix for efficiency.

        Bias[i,j] = -slope * |i - j| for each head
        Shape: [num_heads, seq_len, seq_len]
        """
        # Create position distance matrix
        positions = torch.arange(seq_len).unsqueeze(1)
        distances = torch.abs(positions - positions.t())  # [seq_len, seq_len]

        # Apply slopes: [num_heads, 1, 1] * [1, seq_len, seq_len]
        slopes = self.slopes.view(-1, 1, 1)  # [num_heads, 1, 1]
        bias = -slopes * distances.unsqueeze(0)  # [num_heads, seq_len, seq_len]

        self.register_buffer('bias', bias, persistent=False)

    def forward(
        self,
        attention_scores: torch.Tensor,
        seq_len: Optional[int] = None,
    ) -> torch.Tensor:
        """
        Add ALiBi bias to attention scores.

        Args:
            attention_scores: [batch, num_heads, seq_len, seq_len] attention logits
            seq_len: Optional sequence length (auto-detected if None)

        Returns:
            Biased attention scores with same shape
        """
        if seq_len is None:
            seq_len = attention_scores.shape[-1]

        # Extend cache if needed
        if seq_len > self.max_seq_length:
            self.max_seq_length = seq_len
            self._precompute_bias(seq_len)

        # Add bias: [batch, num_heads, seq_len, seq_len] + [num_heads, seq_len, seq_len]
        bias = self.bias[:, :seq_len, :seq_len]  # [num_heads, seq_len, seq_len]
        biased_scores = attention_scores + bias.unsqueeze(0)

        return biased_scores

    def get_bias(self, seq_len: int) -> torch.Tensor:
        """
        Get bias matrix for a specific sequence length.

        Useful for manual attention computation.

        Returns:
            bias: [num_heads, seq_len, seq_len]
        """
        if seq_len > self.max_seq_length:
            self._precompute_bias(seq_len)

        return self.bias[:, :seq_len, :seq_len]


class PositionalEncodingFactory:
    """
    Factory for creating positional encoding modules.

    Enables easy switching between PE schemes via configuration.
    """

    @staticmethod
    def create(
        pe_type: str,
        dim: int,
        num_heads: Optional[int] = None,
        max_seq_length: int = 2048,
        base: float = 10000.0,
        **kwargs
    ) -> nn.Module:
        """
        Create a positional encoding module.

        Args:
            pe_type: Type of positional encoding:
                - 'rope': Standard RoPE
                - 'ntk_rope': NTK-aware RoPE
                - 'yarn': YaRN extension
                - 'longrope': LongRoPE multi-scale
                - 'alibi': ALiBi linear biases
            dim: Embedding dimension (for RoPE variants)
            num_heads: Number of attention heads (for ALiBi)
            max_seq_length: Maximum sequence length
            base: Base frequency for RoPE
            **kwargs: Additional arguments for specific PE types

        Returns:
            Positional encoding module
        """
        pe_type = pe_type.lower()

        if pe_type == 'rope':
            return RotaryPositionEmbedding(
                dim=dim,
                max_seq_length=max_seq_length,
                base=base,
                scaling_factor=kwargs.get('scaling_factor', 1.0),
            )

        elif pe_type == 'ntk_rope':
            return NTKAwareRoPE(
                dim=dim,
                max_seq_length=max_seq_length,
                base=base,
                scaling_factor=kwargs.get('scaling_factor', 1.0),
                ntk_factor=kwargs.get('ntk_factor', 1.0),
                original_max_seq_length=kwargs.get('original_max_seq_length', 2048),
            )

        elif pe_type == 'yarn':
            return YaRNPositionEmbedding(
                dim=dim,
                max_seq_length=max_seq_length,
                base=base,
                original_max_seq_length=kwargs.get('original_max_seq_length', 2048),
                beta_fast=kwargs.get('beta_fast', 32),
                beta_slow=kwargs.get('beta_slow', 1),
                mscale=kwargs.get('mscale', 0.0),  # 0.0 = auto-compute
                mscale_all_dim=kwargs.get('mscale_all_dim', 0.0),
            )

        elif pe_type == 'longrope':
            return LongRoPE(
                dim=dim,
                max_seq_length=max_seq_length,
                base=base,
                original_max_seq_length=kwargs.get('original_max_seq_length', 2048),
                short_factor=kwargs.get('short_factor', None),
                long_factor=kwargs.get('long_factor', None),
                short_mscale=kwargs.get('short_mscale', 1.0),
                long_mscale=kwargs.get('long_mscale', 1.0),
            )

        elif pe_type == 'alibi':
            if num_heads is None:
                raise ValueError("ALiBi requires num_heads parameter")
            return ALiBiPositionBias(
                num_heads=num_heads,
                max_seq_length=max_seq_length,
                slope_computation=kwargs.get('slope_computation', 'original'),
            )

        else:
            raise ValueError(
                f"Unknown PE type: {pe_type}. "
                f"Choose from: rope, ntk_rope, yarn, longrope, alibi"
            )


def validate_context_extension(
    pe_module: nn.Module,
    original_length: int,
    target_length: int,
    test_batch_size: int = 2,
) -> Dict[str, float]:
    """
    Validate that a positional encoding can handle context extension.

    Tests:
    1. Forward pass completes without errors
    2. Output shape is correct
    3. Gradients flow properly
    4. Memory usage is reasonable

    Args:
        pe_module: Positional encoding module to test
        original_length: Original training sequence length
        target_length: Target extended sequence length
        test_batch_size: Batch size for testing

    Returns:
        Dictionary with validation metrics
    """
    device = next(pe_module.parameters()).device
    metrics = {}

    # Test original length
    try:
        if isinstance(pe_module, ALiBiPositionBias):
            # ALiBi: test attention scores
            attn_scores = torch.randn(
                test_batch_size, pe_module.num_heads,
                original_length, original_length,
                device=device, requires_grad=True
            )
            output = pe_module(attn_scores)
            metrics['original_length_pass'] = True
        else:
            # RoPE variants: test Q/K rotation
            head_dim = pe_module.dim
            num_heads = 8
            q = torch.randn(
                test_batch_size, original_length, num_heads, head_dim,
                device=device, requires_grad=True
            )
            k = torch.randn(
                test_batch_size, original_length, num_heads, head_dim,
                device=device, requires_grad=True
            )
            q_rot, k_rot = pe_module(q, k)
            output = q_rot
            metrics['original_length_pass'] = True
    except Exception as e:
        metrics['original_length_pass'] = False
        metrics['original_error'] = str(e)
        return metrics

    # Test target length
    try:
        if isinstance(pe_module, ALiBiPositionBias):
            attn_scores = torch.randn(
                test_batch_size, pe_module.num_heads,
                target_length, target_length,
                device=device, requires_grad=True
            )
            output = pe_module(attn_scores)
        else:
            q = torch.randn(
                test_batch_size, target_length, num_heads, head_dim,
                device=device, requires_grad=True
            )
            k = torch.randn(
                test_batch_size, target_length, num_heads, head_dim,
                device=device, requires_grad=True
            )
            q_rot, k_rot = pe_module(q, k)
            output = q_rot

        metrics['target_length_pass'] = True
        metrics['extension_ratio'] = target_length / original_length
    except Exception as e:
        metrics['target_length_pass'] = False
        metrics['target_error'] = str(e)
        return metrics

    # Test gradients
    try:
        loss = output.sum()
        loss.backward()
        metrics['gradients_pass'] = True
    except Exception as e:
        metrics['gradients_pass'] = False
        metrics['gradient_error'] = str(e)

    return metrics


# Example usage and configuration
def get_recommended_pe_config(
    target_context_length: int,
    original_context_length: int = 2048,
) -> Dict[str, any]:
    """
    Get recommended PE configuration for target context length.

    Recommendations based on extension ratio:
    - 1x-2x: Standard RoPE
    - 2x-8x: NTK-aware RoPE
    - 8x-64x: YaRN
    - 64x+: LongRoPE or ALiBi

    Args:
        target_context_length: Desired context window
        original_context_length: Original training context

    Returns:
        Configuration dictionary
    """
    ratio = target_context_length / original_context_length

    if ratio <= 2.0:
        return {
            'pe_type': 'rope',
            'max_seq_length': target_context_length,
            'scaling_factor': ratio,
        }
    elif ratio <= 8.0:
        return {
            'pe_type': 'ntk_rope',
            'max_seq_length': target_context_length,
            'original_max_seq_length': original_context_length,
        }
    elif ratio <= 64.0:
        return {
            'pe_type': 'yarn',
            'max_seq_length': target_context_length,
            'original_max_seq_length': original_context_length,
            'beta_fast': 32,
            'beta_slow': 1,
            'mscale': 0.0,  # auto-compute
        }
    else:
        # For extreme extension, recommend ALiBi
        return {
            'pe_type': 'alibi',
            'max_seq_length': target_context_length,
            'slope_computation': 'original',
        }
