"""
Sparse Attention Patterns for MK3

Implements various sparse attention patterns for efficient long-sequence processing:
- Local (Sliding Window) Attention
- Global + Local (Longformer-style)
- Strided Attention
- Random Attention
- BigBird-style (Random + Window + Global)

These patterns reduce attention complexity from O(n²) to O(n) or O(n log n).

References:
- Longformer: https://arxiv.org/abs/2004.05150
- BigBird: https://arxiv.org/abs/2007.14062
- Sparse Transformers: https://arxiv.org/abs/1904.10509
"""

import torch
import torch.nn as nn
import torch.nn.functional as F
import math
from typing import Optional, Tuple, List, Literal


class LocalAttention(nn.Module):
    """
    Local (Sliding Window) Attention.

    Each token attends only to a local window of neighbors.
    Reduces complexity from O(n²) to O(n·w) where w is window size.
    """

    def __init__(
        self,
        embed_dim: int,
        num_heads: int,
        window_size: int = 256,
        dropout: float = 0.0,
        causal: bool = False,
    ):
        super().__init__()

        self.embed_dim = embed_dim
        self.num_heads = num_heads
        self.head_dim = embed_dim // num_heads
        self.window_size = window_size
        self.dropout = dropout
        self.causal = causal

        assert embed_dim % num_heads == 0

        self.scale = 1.0 / math.sqrt(self.head_dim)

        # QKV projection
        self.qkv_proj = nn.Linear(embed_dim, 3 * embed_dim, bias=False)
        self.out_proj = nn.Linear(embed_dim, embed_dim, bias=False)

        if dropout > 0:
            self.dropout_module = nn.Dropout(dropout)
        else:
            self.dropout_module = None

    def forward(
        self,
        x: torch.Tensor,
        attention_mask: Optional[torch.Tensor] = None,
    ) -> torch.Tensor:
        """
        Local attention forward pass.

        Args:
            x: [batch, seq_len, embed_dim]
            attention_mask: Optional [batch, seq_len, seq_len] mask

        Returns:
            output: [batch, seq_len, embed_dim]
        """
        batch_size, seq_len, _ = x.shape

        # Project to Q, K, V
        qkv = self.qkv_proj(x)
        qkv = qkv.reshape(batch_size, seq_len, 3, self.num_heads, self.head_dim)
        qkv = qkv.permute(2, 0, 3, 1, 4)
        q, k, v = qkv[0], qkv[1], qkv[2]
        # q, k, v: [batch, num_heads, seq_len, head_dim]

        # Compute local attention
        output = self._local_attention(q, k, v, attention_mask)

        # Reshape and project
        output = output.transpose(1, 2).contiguous().view(batch_size, seq_len, self.embed_dim)
        output = self.out_proj(output)

        return output

    def _local_attention(
        self,
        q: torch.Tensor,
        k: torch.Tensor,
        v: torch.Tensor,
        attention_mask: Optional[torch.Tensor] = None,
    ) -> torch.Tensor:
        """
        Compute local windowed attention.

        Args:
            q, k, v: [batch, num_heads, seq_len, head_dim]
            attention_mask: Optional mask

        Returns:
            output: [batch, num_heads, seq_len, head_dim]
        """
        batch_size, num_heads, seq_len, head_dim = q.shape

        # Pad sequence if needed to handle edge cases
        half_window = self.window_size // 2
        pad_size = half_window if not self.causal else 0

        if pad_size > 0:
            k_padded = F.pad(k, (0, 0, pad_size, pad_size), value=0)
            v_padded = F.pad(v, (0, 0, pad_size, pad_size), value=0)
        else:
            k_padded = k
            v_padded = v

        # Initialize output
        output = torch.zeros_like(q)

        # Process each position
        for i in range(seq_len):
            if self.causal:
                # Causal: attend to previous window_size tokens
                start = max(0, i - self.window_size + 1)
                end = i + 1
                k_window = k[:, :, start:end, :]
                v_window = v[:, :, start:end, :]
            else:
                # Bidirectional: attend to window_size/2 on each side
                start = i
                end = min(i + self.window_size + 1, k_padded.shape[2])
                k_window = k_padded[:, :, start:end, :]
                v_window = v_padded[:, :, start:end, :]

            # Compute attention for this position
            q_i = q[:, :, i:i+1, :]  # [batch, num_heads, 1, head_dim]
            scores = torch.matmul(q_i, k_window.transpose(-2, -1)) * self.scale

            # Apply mask if provided
            if attention_mask is not None:
                if self.causal:
                    mask_window = attention_mask[:, i:i+1, start:end]
                else:
                    mask_window = attention_mask[:, i:i+1, start:end]
                scores = scores.masked_fill(~mask_window.unsqueeze(1), float('-inf'))

            # Softmax
            attn_weights = F.softmax(scores, dim=-1)

            if self.dropout_module is not None:
                attn_weights = self.dropout_module(attn_weights)

            # Apply attention
            output[:, :, i:i+1, :] = torch.matmul(attn_weights, v_window)

        return output


class GlobalLocalAttention(nn.Module):
    """
    Global + Local Attention (Longformer-style).

    Combines local windowed attention with global attention on special tokens.
    - Most tokens use local attention (window)
    - Selected tokens (e.g., [CLS]) attend globally and are attended to globally
    """

    def __init__(
        self,
        embed_dim: int,
        num_heads: int,
        window_size: int = 256,
        num_global_tokens: int = 1,
        dropout: float = 0.0,
    ):
        super().__init__()

        self.embed_dim = embed_dim
        self.num_heads = num_heads
        self.head_dim = embed_dim // num_heads
        self.window_size = window_size
        self.num_global_tokens = num_global_tokens
        self.dropout = dropout

        assert embed_dim % num_heads == 0

        self.scale = 1.0 / math.sqrt(self.head_dim)

        # QKV projections
        self.qkv_proj = nn.Linear(embed_dim, 3 * embed_dim, bias=False)
        self.out_proj = nn.Linear(embed_dim, embed_dim, bias=False)

        if dropout > 0:
            self.dropout_module = nn.Dropout(dropout)
        else:
            self.dropout_module = None

    def forward(
        self,
        x: torch.Tensor,
        global_token_mask: Optional[torch.Tensor] = None,
    ) -> torch.Tensor:
        """
        Global + Local attention forward pass.

        Args:
            x: [batch, seq_len, embed_dim]
            global_token_mask: [batch, seq_len] boolean mask (True = global token)

        Returns:
            output: [batch, seq_len, embed_dim]
        """
        batch_size, seq_len, _ = x.shape

        # If no global mask provided, use first num_global_tokens
        if global_token_mask is None:
            global_token_mask = torch.zeros(batch_size, seq_len, dtype=torch.bool, device=x.device)
            global_token_mask[:, :self.num_global_tokens] = True

        # Project to Q, K, V
        qkv = self.qkv_proj(x)
        qkv = qkv.reshape(batch_size, seq_len, 3, self.num_heads, self.head_dim)
        qkv = qkv.permute(2, 0, 3, 1, 4)
        q, k, v = qkv[0], qkv[1], qkv[2]

        # Compute attention
        output = self._global_local_attention(q, k, v, global_token_mask)

        # Reshape and project
        output = output.transpose(1, 2).contiguous().view(batch_size, seq_len, self.embed_dim)
        output = self.out_proj(output)

        return output

    def _global_local_attention(
        self,
        q: torch.Tensor,
        k: torch.Tensor,
        v: torch.Tensor,
        global_token_mask: torch.Tensor,
    ) -> torch.Tensor:
        """Compute global + local attention."""
        batch_size, num_heads, seq_len, head_dim = q.shape

        # Separate global and local tokens
        global_mask_expanded = global_token_mask.unsqueeze(1).unsqueeze(-1)  # [batch, 1, seq_len, 1]

        # Extract global tokens
        num_global = global_token_mask.sum(dim=-1).max().item()
        global_indices = torch.where(global_token_mask[0])[0][:num_global]

        # Initialize output
        output = torch.zeros_like(q)

        # Process each position
        for i in range(seq_len):
            q_i = q[:, :, i:i+1, :]  # [batch, num_heads, 1, head_dim]

            if global_token_mask[0, i]:
                # Global token: attend to all tokens
                k_attended = k
                v_attended = v
            else:
                # Local token: attend to window + global tokens
                # Window indices
                start = max(0, i - self.window_size // 2)
                end = min(seq_len, i + self.window_size // 2 + 1)
                window_indices = list(range(start, end))

                # Combine window and global indices
                all_indices = sorted(set(window_indices + global_indices.tolist()))
                all_indices_tensor = torch.tensor(all_indices, device=q.device)

                k_attended = k[:, :, all_indices_tensor, :]
                v_attended = v[:, :, all_indices_tensor, :]

            # Compute attention
            scores = torch.matmul(q_i, k_attended.transpose(-2, -1)) * self.scale
            attn_weights = F.softmax(scores, dim=-1)

            if self.dropout_module is not None:
                attn_weights = self.dropout_module(attn_weights)

            output[:, :, i:i+1, :] = torch.matmul(attn_weights, v_attended)

        return output


class BigBirdAttention(nn.Module):
    """
    BigBird-style Sparse Attention.

    Combines three types of attention:
    1. Random attention: each token attends to r random tokens
    2. Window attention: each token attends to w local neighbors
    3. Global attention: g tokens attend to and are attended by all tokens

    This provides O(n) complexity while maintaining good performance.
    """

    def __init__(
        self,
        embed_dim: int,
        num_heads: int,
        window_size: int = 128,
        num_random_tokens: int = 64,
        num_global_tokens: int = 2,
        dropout: float = 0.0,
        block_size: int = 64,
    ):
        super().__init__()

        self.embed_dim = embed_dim
        self.num_heads = num_heads
        self.head_dim = embed_dim // num_heads
        self.window_size = window_size
        self.num_random_tokens = num_random_tokens
        self.num_global_tokens = num_global_tokens
        self.dropout = dropout
        self.block_size = block_size

        assert embed_dim % num_heads == 0

        self.scale = 1.0 / math.sqrt(self.head_dim)

        # QKV projection
        self.qkv_proj = nn.Linear(embed_dim, 3 * embed_dim, bias=False)
        self.out_proj = nn.Linear(embed_dim, embed_dim, bias=False)

        if dropout > 0:
            self.dropout_module = nn.Dropout(dropout)
        else:
            self.dropout_module = None

    def forward(
        self,
        x: torch.Tensor,
        global_token_mask: Optional[torch.Tensor] = None,
        random_seed: Optional[int] = None,
    ) -> torch.Tensor:
        """
        BigBird attention forward pass.

        Args:
            x: [batch, seq_len, embed_dim]
            global_token_mask: [batch, seq_len] boolean mask for global tokens
            random_seed: Optional seed for random attention (for reproducibility)

        Returns:
            output: [batch, seq_len, embed_dim]
        """
        batch_size, seq_len, _ = x.shape

        # Default global token mask (first num_global_tokens)
        if global_token_mask is None:
            global_token_mask = torch.zeros(batch_size, seq_len, dtype=torch.bool, device=x.device)
            global_token_mask[:, :self.num_global_tokens] = True

        # Project to Q, K, V
        qkv = self.qkv_proj(x)
        qkv = qkv.reshape(batch_size, seq_len, 3, self.num_heads, self.head_dim)
        qkv = qkv.permute(2, 0, 3, 1, 4)
        q, k, v = qkv[0], qkv[1], qkv[2]

        # Compute BigBird attention
        output = self._bigbird_attention(q, k, v, global_token_mask, random_seed)

        # Reshape and project
        output = output.transpose(1, 2).contiguous().view(batch_size, seq_len, self.embed_dim)
        output = self.out_proj(output)

        return output

    def _bigbird_attention(
        self,
        q: torch.Tensor,
        k: torch.Tensor,
        v: torch.Tensor,
        global_token_mask: torch.Tensor,
        random_seed: Optional[int] = None,
    ) -> torch.Tensor:
        """Compute BigBird sparse attention."""
        batch_size, num_heads, seq_len, head_dim = q.shape

        # Get global token indices
        global_indices = torch.where(global_token_mask[0])[0]

        # Initialize output
        output = torch.zeros_like(q)

        # Set random seed if provided for reproducibility
        if random_seed is not None:
            torch.manual_seed(random_seed)

        # Process each position
        for i in range(seq_len):
            q_i = q[:, :, i:i+1, :]

            # 1. Window attention (local neighbors)
            window_start = max(0, i - self.window_size // 2)
            window_end = min(seq_len, i + self.window_size // 2 + 1)
            window_indices = list(range(window_start, window_end))

            # 2. Random attention (sample random tokens)
            if not global_token_mask[0, i]:
                # Sample random tokens (excluding window and global)
                available_indices = set(range(seq_len)) - set(window_indices) - set(global_indices.tolist())
                if len(available_indices) > 0:
                    random_indices = list(available_indices)[:self.num_random_tokens]
                else:
                    random_indices = []
            else:
                random_indices = []

            # 3. Global tokens
            global_indices_list = global_indices.tolist()

            # Combine all attended indices
            if global_token_mask[0, i]:
                # Global token attends to everything
                attended_indices = list(range(seq_len))
            else:
                attended_indices = sorted(set(window_indices + random_indices + global_indices_list))

            attended_indices_tensor = torch.tensor(attended_indices, device=q.device)

            # Get K, V for attended positions
            k_attended = k[:, :, attended_indices_tensor, :]
            v_attended = v[:, :, attended_indices_tensor, :]

            # Compute attention
            scores = torch.matmul(q_i, k_attended.transpose(-2, -1)) * self.scale
            attn_weights = F.softmax(scores, dim=-1)

            if self.dropout_module is not None:
                attn_weights = self.dropout_module(attn_weights)

            output[:, :, i:i+1, :] = torch.matmul(attn_weights, v_attended)

        return output


class StridedAttention(nn.Module):
    """
    Strided Sparse Attention.

    Each token attends to every k-th token (strided pattern).
    Useful for hierarchical modeling.
    """

    def __init__(
        self,
        embed_dim: int,
        num_heads: int,
        stride: int = 8,
        window_size: int = 128,
        dropout: float = 0.0,
    ):
        super().__init__()

        self.embed_dim = embed_dim
        self.num_heads = num_heads
        self.head_dim = embed_dim // num_heads
        self.stride = stride
        self.window_size = window_size
        self.dropout = dropout

        assert embed_dim % num_heads == 0

        self.scale = 1.0 / math.sqrt(self.head_dim)

        # QKV projection
        self.qkv_proj = nn.Linear(embed_dim, 3 * embed_dim, bias=False)
        self.out_proj = nn.Linear(embed_dim, embed_dim, bias=False)

        if dropout > 0:
            self.dropout_module = nn.Dropout(dropout)
        else:
            self.dropout_module = None

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        """Strided attention forward pass."""
        batch_size, seq_len, _ = x.shape

        # Project to Q, K, V
        qkv = self.qkv_proj(x)
        qkv = qkv.reshape(batch_size, seq_len, 3, self.num_heads, self.head_dim)
        qkv = qkv.permute(2, 0, 3, 1, 4)
        q, k, v = qkv[0], qkv[1], qkv[2]

        # Compute strided attention
        output = self._strided_attention(q, k, 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

    def _strided_attention(
        self,
        q: torch.Tensor,
        k: torch.Tensor,
        v: torch.Tensor,
    ) -> torch.Tensor:
        """Compute strided attention."""
        batch_size, num_heads, seq_len, head_dim = q.shape

        # Initialize output
        output = torch.zeros_like(q)

        # Process each position
        for i in range(seq_len):
            q_i = q[:, :, i:i+1, :]

            # Local window
            window_start = max(0, i - self.window_size // 2)
            window_end = min(seq_len, i + self.window_size // 2 + 1)
            window_indices = list(range(window_start, window_end))

            # Strided indices
            strided_indices = list(range(0, seq_len, self.stride))

            # Combine
            attended_indices = sorted(set(window_indices + strided_indices))
            attended_indices_tensor = torch.tensor(attended_indices, device=q.device)

            k_attended = k[:, :, attended_indices_tensor, :]
            v_attended = v[:, :, attended_indices_tensor, :]

            # Compute attention
            scores = torch.matmul(q_i, k_attended.transpose(-2, -1)) * self.scale
            attn_weights = F.softmax(scores, dim=-1)

            if self.dropout_module is not None:
                attn_weights = self.dropout_module(attn_weights)

            output[:, :, i:i+1, :] = torch.matmul(attn_weights, v_attended)

        return output


def create_sparse_attention(
    embed_dim: int,
    num_heads: int,
    pattern: Literal["local", "global_local", "bigbird", "strided"] = "local",
    **kwargs
) -> nn.Module:
    """
    Factory function for creating sparse attention layers.

    Args:
        embed_dim: Embedding dimension
        num_heads: Number of attention heads
        pattern: Sparse attention pattern type
        **kwargs: Additional pattern-specific arguments

    Returns:
        Sparse attention module
    """
    if pattern == "local":
        return LocalAttention(embed_dim, num_heads, **kwargs)
    elif pattern == "global_local":
        return GlobalLocalAttention(embed_dim, num_heads, **kwargs)
    elif pattern == "bigbird":
        return BigBirdAttention(embed_dim, num_heads, **kwargs)
    elif pattern == "strided":
        return StridedAttention(embed_dim, num_heads, **kwargs)
    else:
        raise ValueError(f"Unknown sparse attention pattern: {pattern}")
