"""
StreamingLLM Support - Attention Sinks for Infinite-Length Generation

Implements attention sink mechanism for efficient infinite-length text generation
without recomputing from scratch or performance degradation.

Key features:
- Attention sink tokens (initial tokens that stabilize attention)
- Rolling KV cache with fixed size
- Efficient long-sequence generation without full recomputation
- Perplexity stability for infinite-length generation

References:
- Efficient Streaming Language Models with Attention Sinks: https://arxiv.org/abs/2309.17453

Attention Sink Discovery:
Even when initial tokens are semantically irrelevant, they accumulate large attention
scores and act as "attention sinks" that stabilize the attention distribution.
Preserving these sink tokens in the KV cache allows models to maintain performance
with fixed-size caches.
"""

import torch
import torch.nn as nn
import torch.nn.functional as F
import math
from typing import Optional, Tuple, Dict
from collections import deque


class KVCache:
    """
    Key-Value cache with attention sink support.

    Maintains a fixed-size cache with special handling for sink tokens.
    """

    def __init__(
        self,
        max_cache_size: int,
        num_sink_tokens: int = 4,
        num_heads: int = 8,
        head_dim: int = 64,
        device: Optional[torch.device] = None,
        dtype: Optional[torch.dtype] = None,
    ):
        """
        Initialize KV cache.

        Args:
            max_cache_size: Maximum number of tokens to cache
            num_sink_tokens: Number of initial tokens to always keep (attention sinks)
            num_heads: Number of attention heads
            head_dim: Dimension per head
            device: Device to store cache
            dtype: Data type for cache
        """
        self.max_cache_size = max_cache_size
        self.num_sink_tokens = num_sink_tokens
        self.num_heads = num_heads
        self.head_dim = head_dim
        self.device = device or torch.device('cpu')
        self.dtype = dtype or torch.float32

        # Separate storage for sink tokens and rolling cache
        self.sink_keys = None
        self.sink_values = None
        self.rolling_keys = None
        self.rolling_values = None

        self.current_length = 0  # Total number of tokens seen
        self.cache_start_position = 0  # Position of first token in rolling cache

    def update(
        self,
        new_keys: torch.Tensor,
        new_values: torch.Tensor,
        position: Optional[int] = None,
    ) -> Tuple[torch.Tensor, torch.Tensor]:
        """
        Update cache with new keys and values.

        Args:
            new_keys: [batch, num_heads, seq_len, head_dim] new keys
            new_values: [batch, num_heads, seq_len, head_dim] new values
            position: Starting position of new tokens (None = append)

        Returns:
            cached_keys: [batch, num_heads, cache_len, head_dim] all cached keys
            cached_values: [batch, num_heads, cache_len, head_dim] all cached values
        """
        batch_size, num_heads, seq_len, head_dim = new_keys.shape

        # Move to cache device/dtype
        new_keys = new_keys.to(device=self.device, dtype=self.dtype)
        new_values = new_values.to(device=self.device, dtype=self.dtype)

        if self.sink_keys is None:
            # First update: initialize sink tokens
            num_sink = min(seq_len, self.num_sink_tokens)

            self.sink_keys = new_keys[:, :, :num_sink, :].clone()
            self.sink_values = new_values[:, :, :num_sink, :].clone()

            # Initialize rolling cache with remaining tokens
            if seq_len > num_sink:
                remaining_keys = new_keys[:, :, num_sink:, :]
                remaining_values = new_values[:, :, num_sink:, :]

                max_rolling = self.max_cache_size - self.num_sink_tokens
                if remaining_keys.shape[2] > max_rolling:
                    # Take only the most recent tokens
                    self.rolling_keys = remaining_keys[:, :, -max_rolling:, :].clone()
                    self.rolling_values = remaining_values[:, :, -max_rolling:, :].clone()
                else:
                    self.rolling_keys = remaining_keys.clone()
                    self.rolling_values = remaining_values.clone()

                self.cache_start_position = num_sink + max(0, remaining_keys.shape[2] - max_rolling)
            else:
                self.rolling_keys = torch.zeros(
                    batch_size, num_heads, 0, head_dim,
                    device=self.device, dtype=self.dtype
                )
                self.rolling_values = torch.zeros(
                    batch_size, num_heads, 0, head_dim,
                    device=self.device, dtype=self.dtype
                )
                self.cache_start_position = num_sink

            self.current_length = seq_len

        else:
            # Subsequent updates: append to rolling cache
            max_rolling = self.max_cache_size - self.num_sink_tokens

            # Concatenate new tokens to rolling cache
            if self.rolling_keys.shape[2] > 0:
                self.rolling_keys = torch.cat([self.rolling_keys, new_keys], dim=2)
                self.rolling_values = torch.cat([self.rolling_values, new_values], dim=2)
            else:
                self.rolling_keys = new_keys
                self.rolling_values = new_values

            # Trim if exceeds max size
            if self.rolling_keys.shape[2] > max_rolling:
                overflow = self.rolling_keys.shape[2] - max_rolling
                self.rolling_keys = self.rolling_keys[:, :, overflow:, :]
                self.rolling_values = self.rolling_values[:, :, overflow:, :]
                self.cache_start_position += overflow

            self.current_length += seq_len

        # Return combined cache (sinks + rolling)
        cached_keys = torch.cat([self.sink_keys, self.rolling_keys], dim=2)
        cached_values = torch.cat([self.sink_values, self.rolling_values], dim=2)

        return cached_keys, cached_values

    def get_cache(self) -> Tuple[Optional[torch.Tensor], Optional[torch.Tensor]]:
        """
        Get current cache contents.

        Returns:
            cached_keys: [batch, num_heads, cache_len, head_dim] or None
            cached_values: [batch, num_heads, cache_len, head_dim] or None
        """
        if self.sink_keys is None:
            return None, None

        cached_keys = torch.cat([self.sink_keys, self.rolling_keys], dim=2)
        cached_values = torch.cat([self.sink_values, self.rolling_values], dim=2)

        return cached_keys, cached_values

    def clear(self):
        """Clear the cache."""
        self.sink_keys = None
        self.sink_values = None
        self.rolling_keys = None
        self.rolling_values = None
        self.current_length = 0
        self.cache_start_position = 0

    def __len__(self) -> int:
        """Return current cache size."""
        if self.sink_keys is None:
            return 0
        return self.sink_keys.shape[2] + self.rolling_keys.shape[2]


class StreamingAttention(nn.Module):
    """
    Streaming attention with attention sinks.

    Enables efficient infinite-length generation with fixed memory.
    """

    def __init__(
        self,
        embed_dim: int,
        num_heads: int,
        max_cache_size: int = 2048,
        num_sink_tokens: int = 4,
        dropout: float = 0.0,
        causal: bool = True,
    ):
        super().__init__()

        self.embed_dim = embed_dim
        self.num_heads = num_heads
        self.head_dim = embed_dim // num_heads
        self.max_cache_size = max_cache_size
        self.num_sink_tokens = num_sink_tokens
        self.dropout = dropout
        self.causal = causal

        assert embed_dim % num_heads == 0
        assert max_cache_size >= num_sink_tokens, "Cache must be large enough for sink tokens"

        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

        # KV cache (initialized per forward pass)
        self._cache = None

    def forward(
        self,
        x: torch.Tensor,
        use_cache: bool = False,
        past_kv: Optional[KVCache] = None,
        attention_mask: Optional[torch.Tensor] = None,
    ) -> Tuple[torch.Tensor, Optional[KVCache]]:
        """
        Forward pass with streaming support.

        Args:
            x: [batch, seq_len, embed_dim] input tokens
            use_cache: Whether to use/update KV cache
            past_kv: Previous KV cache (if continuing generation)
            attention_mask: Optional attention mask

        Returns:
            output: [batch, seq_len, embed_dim] attended output
            kv_cache: Updated KV cache if use_cache=True
        """
        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]

        # Handle KV cache
        if use_cache:
            if past_kv is None:
                # Initialize cache
                past_kv = KVCache(
                    max_cache_size=self.max_cache_size,
                    num_sink_tokens=self.num_sink_tokens,
                    num_heads=self.num_heads,
                    head_dim=self.head_dim,
                    device=x.device,
                    dtype=x.dtype,
                )

            # Update cache with new K, V
            cached_k, cached_v = past_kv.update(k, v)

            # Use cached K, V for attention
            k_for_attn = cached_k
            v_for_attn = cached_v
        else:
            k_for_attn = k
            v_for_attn = v

        # Compute attention
        output = self._compute_attention(q, k_for_attn, v_for_attn, 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, past_kv if use_cache else None

    def _compute_attention(
        self,
        q: torch.Tensor,
        k: torch.Tensor,
        v: torch.Tensor,
        attention_mask: Optional[torch.Tensor] = None,
    ) -> torch.Tensor:
        """
        Compute attention.

        Args:
            q: [batch, num_heads, seq_len_q, head_dim]
            k: [batch, num_heads, seq_len_k, head_dim]
            v: [batch, num_heads, seq_len_k, head_dim]
            attention_mask: Optional mask

        Returns:
            output: [batch, num_heads, seq_len_q, head_dim]
        """
        # Attention scores
        scores = torch.matmul(q, k.transpose(-2, -1)) * self.scale

        # Apply causal mask
        if self.causal:
            seq_len_q = q.shape[2]
            seq_len_k = k.shape[2]

            # Causal mask: can attend to all cached tokens, but not future tokens
            if seq_len_q < seq_len_k:
                # During generation: query is just the new token
                # Can attend to all previous tokens (cache)
                pass  # No masking needed
            else:
                # During prefill: apply standard causal mask
                causal_mask = torch.triu(
                    torch.ones(seq_len_q, seq_len_k, device=q.device, dtype=torch.bool),
                    diagonal=1
                )
                scores = scores.masked_fill(causal_mask, float('-inf'))

        # Apply custom attention mask if provided
        if attention_mask is not None:
            scores = scores.masked_fill(~attention_mask, 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 to values
        output = torch.matmul(attn_weights, v)

        return output


class StreamingTransformerBlock(nn.Module):
    """
    Complete transformer block with streaming support.

    Combines streaming attention with feedforward network.
    """

    def __init__(
        self,
        embed_dim: int,
        num_heads: int,
        feedforward_dim: Optional[int] = None,
        max_cache_size: int = 2048,
        num_sink_tokens: int = 4,
        dropout: float = 0.1,
    ):
        super().__init__()

        self.embed_dim = embed_dim
        feedforward_dim = feedforward_dim or (embed_dim * 4)

        # Streaming attention
        self.attention = StreamingAttention(
            embed_dim=embed_dim,
            num_heads=num_heads,
            max_cache_size=max_cache_size,
            num_sink_tokens=num_sink_tokens,
            dropout=dropout,
            causal=True,
        )

        # Feedforward network
        self.feedforward = nn.Sequential(
            nn.Linear(embed_dim, feedforward_dim),
            nn.GELU(),
            nn.Dropout(dropout),
            nn.Linear(feedforward_dim, embed_dim),
            nn.Dropout(dropout),
        )

        # Layer norms
        self.norm1 = nn.LayerNorm(embed_dim)
        self.norm2 = nn.LayerNorm(embed_dim)

    def forward(
        self,
        x: torch.Tensor,
        use_cache: bool = False,
        past_kv: Optional[KVCache] = None,
        attention_mask: Optional[torch.Tensor] = None,
    ) -> Tuple[torch.Tensor, Optional[KVCache]]:
        """
        Forward pass with streaming support.

        Args:
            x: [batch, seq_len, embed_dim] input
            use_cache: Whether to use KV cache
            past_kv: Previous KV cache
            attention_mask: Optional attention mask

        Returns:
            output: [batch, seq_len, embed_dim]
            kv_cache: Updated KV cache if use_cache=True
        """
        # Self-attention with residual
        attn_out, new_kv = self.attention(
            self.norm1(x),
            use_cache=use_cache,
            past_kv=past_kv,
            attention_mask=attention_mask,
        )
        x = x + attn_out

        # Feedforward with residual
        ff_out = self.feedforward(self.norm2(x))
        x = x + ff_out

        return x, new_kv


class StreamingLLM(nn.Module):
    """
    Complete Streaming LLM with attention sinks.

    Supports infinite-length generation with fixed memory.
    """

    def __init__(
        self,
        vocab_size: int,
        embed_dim: int = 768,
        num_layers: int = 12,
        num_heads: int = 12,
        feedforward_dim: Optional[int] = None,
        max_cache_size: int = 2048,
        num_sink_tokens: int = 4,
        dropout: float = 0.1,
        max_seq_len: int = 8192,
    ):
        super().__init__()

        self.vocab_size = vocab_size
        self.embed_dim = embed_dim
        self.num_layers = num_layers
        self.max_cache_size = max_cache_size
        self.num_sink_tokens = num_sink_tokens

        # Token embedding
        self.token_embedding = nn.Embedding(vocab_size, embed_dim)

        # Positional embedding (learned or sinusoidal)
        self.position_embedding = nn.Embedding(max_seq_len, embed_dim)

        # Transformer blocks
        self.blocks = nn.ModuleList([
            StreamingTransformerBlock(
                embed_dim=embed_dim,
                num_heads=num_heads,
                feedforward_dim=feedforward_dim,
                max_cache_size=max_cache_size,
                num_sink_tokens=num_sink_tokens,
                dropout=dropout,
            )
            for _ in range(num_layers)
        ])

        # Output
        self.norm = nn.LayerNorm(embed_dim)
        self.lm_head = nn.Linear(embed_dim, vocab_size, bias=False)

        # Tie weights
        self.lm_head.weight = self.token_embedding.weight

    def forward(
        self,
        input_ids: torch.Tensor,
        use_cache: bool = False,
        past_kvs: Optional[List[KVCache]] = None,
        position_offset: int = 0,
    ) -> Tuple[torch.Tensor, Optional[List[KVCache]]]:
        """
        Forward pass with streaming support.

        Args:
            input_ids: [batch, seq_len] input token IDs
            use_cache: Whether to use KV cache
            past_kvs: List of past KV caches for each layer
            position_offset: Position offset for continuing generation

        Returns:
            logits: [batch, seq_len, vocab_size] output logits
            new_kvs: Updated KV caches if use_cache=True
        """
        batch_size, seq_len = input_ids.shape

        # Embeddings
        positions = torch.arange(
            position_offset,
            position_offset + seq_len,
            device=input_ids.device
        )
        x = self.token_embedding(input_ids) + self.position_embedding(positions)

        # Initialize cache list if needed
        if use_cache and past_kvs is None:
            past_kvs = [None] * self.num_layers

        # Process through transformer blocks
        new_kvs = []
        for i, block in enumerate(self.blocks):
            past_kv = past_kvs[i] if past_kvs else None
            x, new_kv = block(x, use_cache=use_cache, past_kv=past_kv)
            if use_cache:
                new_kvs.append(new_kv)

        # Output
        x = self.norm(x)
        logits = self.lm_head(x)

        return logits, new_kvs if use_cache else None

    @torch.no_grad()
    def generate(
        self,
        input_ids: torch.Tensor,
        max_new_tokens: int = 100,
        temperature: float = 1.0,
        top_k: Optional[int] = None,
        top_p: Optional[float] = None,
    ) -> torch.Tensor:
        """
        Generate tokens with streaming (infinite-length capable).

        Args:
            input_ids: [batch, seq_len] initial tokens
            max_new_tokens: Number of tokens to generate
            temperature: Sampling temperature
            top_k: Top-k sampling
            top_p: Nucleus sampling

        Returns:
            generated_ids: [batch, seq_len + max_new_tokens] generated tokens
        """
        batch_size = input_ids.shape[0]
        generated = input_ids.clone()

        # Initialize cache
        past_kvs = None
        position_offset = 0

        for _ in range(max_new_tokens):
            # Forward pass
            if past_kvs is None:
                # Prefill: process all tokens
                logits, past_kvs = self.forward(
                    generated,
                    use_cache=True,
                    past_kvs=None,
                    position_offset=0,
                )
                position_offset = generated.shape[1]
            else:
                # Generation: process only last token
                logits, past_kvs = self.forward(
                    generated[:, -1:],
                    use_cache=True,
                    past_kvs=past_kvs,
                    position_offset=position_offset,
                )
                position_offset += 1

            # Sample next token
            next_token_logits = logits[:, -1, :] / temperature

            # Top-k sampling
            if top_k is not None:
                indices_to_remove = next_token_logits < torch.topk(next_token_logits, top_k)[0][..., -1, None]
                next_token_logits[indices_to_remove] = float('-inf')

            # Top-p (nucleus) sampling
            if top_p is not None:
                sorted_logits, sorted_indices = torch.sort(next_token_logits, descending=True)
                cumulative_probs = torch.cumsum(F.softmax(sorted_logits, dim=-1), dim=-1)
                sorted_indices_to_remove = cumulative_probs > top_p
                sorted_indices_to_remove[..., 1:] = sorted_indices_to_remove[..., :-1].clone()
                sorted_indices_to_remove[..., 0] = 0
                indices_to_remove = sorted_indices_to_remove.scatter(1, sorted_indices, sorted_indices_to_remove)
                next_token_logits[indices_to_remove] = float('-inf')

            # Sample
            probs = F.softmax(next_token_logits, dim=-1)
            next_token = torch.multinomial(probs, num_samples=1)

            # Append to generated sequence
            generated = torch.cat([generated, next_token], dim=1)

        return generated
