"""
Neural network layers using the novel scoring formula.
"""

import torch
import torch.nn as nn
import torch.nn.functional as F
from typing import Optional, Tuple
from .formula import ScoringFormula


class FormulaAttentionLayer(nn.Module):
    """
    Attention mechanism based on the novel scoring formula.
    
    Uses the formula to compute attention scores instead of standard dot-product attention.
    """
    
    def __init__(
        self,
        embedding_dim: int,
        formula_config: Optional[dict] = None,
        dropout: float = 0.1,
    ):
        super().__init__()
        self.embedding_dim = embedding_dim
        
        formula_config = formula_config or {}
        self.scoring_formula = ScoringFormula(embedding_dim=embedding_dim, **formula_config)
        
        # Projection layers
        self.query_proj = nn.Linear(embedding_dim, embedding_dim)
        self.key_proj = nn.Linear(embedding_dim, embedding_dim)
        self.value_proj = nn.Linear(embedding_dim, embedding_dim)
        self.output_proj = nn.Linear(embedding_dim, embedding_dim)
        
        self.dropout = nn.Dropout(dropout)
        self.layer_norm = nn.LayerNorm(embedding_dim)
    
    def forward(
        self,
        query: torch.Tensor,
        key: torch.Tensor,
        value: torch.Tensor,
        time_steps: Optional[torch.Tensor] = None,
        memory_buffer: Optional[torch.Tensor] = None,
        mask: Optional[torch.Tensor] = None
    ) -> Tuple[torch.Tensor, torch.Tensor]:
        """
        Apply formula-based attention.
        
        Args:
            query: [batch, seq_len_q, embed_dim]
            key: [batch, seq_len_k, embed_dim]
            value: [batch, seq_len_k, embed_dim]
            time_steps: [batch, seq_len_k] optional
            memory_buffer: [memory_size, embed_dim] optional
            mask: [batch, seq_len_q, seq_len_k] optional attention mask
            
        Returns:
            output: [batch, seq_len_q, embed_dim]
            attention_weights: [batch, seq_len_q, seq_len_k]
        """
        batch_size, seq_len_q, embed_dim = query.shape
        seq_len_k = key.shape[1]
        
        # Project queries, keys, values
        Q = self.query_proj(query)
        K = self.key_proj(key)
        V = self.value_proj(value)
        
        # Vectorized computation of attention scores using the formula
        # Reshape to compute all query-key pairs in parallel
        # Q: [batch, seq_len_q, embed_dim] -> [batch * seq_len_q, embed_dim]
        # K: [batch, seq_len_k, embed_dim] -> [batch * seq_len_k, embed_dim]
        
        # Expand Q and K to create all pairs: [batch, seq_len_q, seq_len_k, embed_dim]
        Q_expanded = Q.unsqueeze(2).expand(-1, -1, seq_len_k, -1)  # [batch, seq_len_q, seq_len_k, embed_dim]
        K_expanded = K.unsqueeze(1).expand(-1, seq_len_q, -1, -1)   # [batch, seq_len_q, seq_len_k, embed_dim]
        
        # Flatten for batch processing: [batch * seq_len_q * seq_len_k, embed_dim]
        Q_flat = Q_expanded.reshape(-1, embed_dim)
        K_flat = K_expanded.reshape(-1, embed_dim)
        
        # Prepare time steps if provided
        if time_steps is not None:
            # Expand time_steps to match all pairs: [batch, seq_len_q, seq_len_k]
            time_steps_expanded = time_steps.unsqueeze(1).expand(-1, seq_len_q, -1)
            time_steps_flat = time_steps_expanded.reshape(-1)  # [batch * seq_len_q * seq_len_k]
        else:
            time_steps_flat = None
        
        # Compute formula scores for all pairs at once
        scores_flat, _ = self.scoring_formula(
            current=K_flat,  # Each key as "current"
            context=Q_flat,  # Corresponding query as "context"
            time_steps=time_steps_flat,
            memory_buffer=memory_buffer
        )
        
        # Reshape back to attention score matrix: [batch, seq_len_q, seq_len_k]
        attention_scores = scores_flat.reshape(batch_size, seq_len_q, seq_len_k)
        
        # Apply mask if provided
        if mask is not None:
            # Ensure mask is correct shape: [batch, seq_len_q, seq_len_k]
            # Convert to bool if needed
            if mask.dtype != torch.bool:
                mask = mask.bool()
            
            # Handle different mask shapes
            if mask.dim() == 2:
                # Could be [batch, seq_len] or [seq_len, seq_len]
                if mask.shape[0] == batch_size:
                    # [batch, seq_len] - assume same for query and key
                    if mask.shape[1] == seq_len_q and seq_len_q == seq_len_k:
                        # Perfect match - expand to [batch, seq_len_q, seq_len_k]
                        mask = mask.unsqueeze(2).expand(-1, -1, seq_len_k)
                    else:
                        # Need to handle differently
                        if mask.shape[1] == seq_len_q:
                            mask = mask.unsqueeze(2).expand(-1, -1, seq_len_k)
                        elif mask.shape[1] == seq_len_k:
                            mask = mask.unsqueeze(1).expand(-1, seq_len_q, -1)
                        else:
                            # Resize - take first min(seq_len_q, seq_len_k) elements
                            min_len = min(seq_len_q, seq_len_k, mask.shape[1])
                            mask = mask[:, :min_len]
                            if mask.shape[1] == seq_len_q:
                                mask = mask.unsqueeze(2).expand(-1, -1, seq_len_k)
                            else:
                                mask = mask.unsqueeze(1).expand(-1, seq_len_q, -1)
                else:
                    # Assume [seq_len, seq_len] causal mask
                    if mask.shape[0] == mask.shape[1]:
                        # Expand to batch dimension: [batch, seq_len, seq_len]
                        min_len = min(seq_len_q, seq_len_k, mask.shape[0])
                        mask = mask[:min_len, :min_len]
                        mask = mask.unsqueeze(0).expand(batch_size, -1, -1)
                        # Expand to match seq_len_q and seq_len_k if needed
                        if mask.shape[1] < seq_len_q or mask.shape[2] < seq_len_k:
                            new_mask = torch.zeros(batch_size, seq_len_q, seq_len_k, dtype=torch.bool, device=mask.device)
                            new_mask[:, :mask.shape[1], :mask.shape[2]] = mask
                            mask = new_mask
            elif mask.dim() == 3:
                # [batch, seq_len, seq_len] - check and resize if needed
                if mask.shape[0] != batch_size:
                    mask = mask[0:1].expand(batch_size, -1, -1)
                
                # Resize to match seq_len_q and seq_len_k
                if mask.shape[1] != seq_len_q or mask.shape[2] != seq_len_k:
                    new_mask = torch.zeros(batch_size, seq_len_q, seq_len_k, dtype=torch.bool, device=mask.device)
                    min_q = min(mask.shape[1], seq_len_q)
                    min_k = min(mask.shape[2], seq_len_k)
                    new_mask[:, :min_q, :min_k] = mask[:, :min_q, :min_k]
                    mask = new_mask
            
            # Apply mask
            attention_scores = attention_scores.masked_fill(~mask, float('-inf'))
        
        # Normalize attention scores (softmax over key dimension)
        attention_weights = F.softmax(attention_scores, dim=-1)
        attention_weights = self.dropout(attention_weights)
        
        # Ensure attention_weights is 3D for bmm
        # attention_weights should be [batch, seq_len_q, seq_len_k]
        if attention_weights.dim() == 2:
            attention_weights = attention_weights.unsqueeze(0)
        
        # Apply attention to values
        # V is [batch, seq_len_k, embed_dim]
        # attention_weights is [batch, seq_len_q, seq_len_k]
        output = torch.bmm(attention_weights, V)  # [batch, seq_len_q, embed_dim]
        output = self.output_proj(output)
        output = self.dropout(output)
        
        # Residual connection and layer norm
        output = self.layer_norm(output + query)
        
        return output, attention_weights


class FormulaSelectiveLayer(nn.Module):
    """
    Layer that selectively processes information based on formula scores.
    
    Filters and prioritizes information flow using the scoring formula.
    """
    
    def __init__(
        self,
        embedding_dim: int,
        formula_config: Optional[dict] = None,
        selection_threshold: float = 0.1,
    ):
        super().__init__()
        self.embedding_dim = embedding_dim
        self.selection_threshold = selection_threshold
        
        formula_config = formula_config or {}
        self.scoring_formula = ScoringFormula(embedding_dim=embedding_dim, **formula_config)
        
        self.process_net = nn.Sequential(
            nn.Linear(embedding_dim, embedding_dim * 2),
            nn.GELU(),
            nn.Linear(embedding_dim * 2, embedding_dim),
            nn.LayerNorm(embedding_dim)
        )
    
    def forward(
        self,
        x: torch.Tensor,
        context: torch.Tensor,
        time_steps: Optional[torch.Tensor] = None,
        memory_buffer: Optional[torch.Tensor] = None
    ) -> Tuple[torch.Tensor, torch.Tensor]:
        """
        Selectively process input based on formula scores.
        
        Args:
            x: Input [batch, seq_len, embed_dim]
            context: Context [batch, context_len, embed_dim] or [batch, embed_dim]
            time_steps: [batch, seq_len] optional
            memory_buffer: [memory_size, embed_dim] optional
            
        Returns:
            output: Processed output [batch, seq_len, embed_dim]
            selection_mask: [batch, seq_len] binary mask of selected items
        """
        batch_size, seq_len, embed_dim = x.shape
        
        # Handle context dimension
        if context.dim() == 2:
            if context.shape[0] == batch_size:
                # [batch, embed_dim]
                context_agg = context
            else:
                # Aggregate context (e.g., mean pooling)
                context_agg = context.mean(dim=0, keepdim=True).expand(batch_size, -1)
        elif context.dim() == 3:
            context_agg = context.mean(dim=1)  # [batch, embed_dim]
        else:
            # Single context vector
            context_agg = context.expand(batch_size, -1)
        
        # Vectorized computation of scores for all positions
        # Expand context to match sequence length: [batch, seq_len, embed_dim]
        context_expanded = context_agg.unsqueeze(1).expand(-1, seq_len, -1)
        
        # Compute scores for all positions at once
        # x: [batch, seq_len, embed_dim]
        # context_expanded: [batch, seq_len, embed_dim]
        scores, _ = self.scoring_formula(
            current=x,
            context=context_expanded,
            time_steps=time_steps,
            memory_buffer=memory_buffer
        )
        
        # Ensure scores is [batch, seq_len]
        if scores.dim() > 2:
            scores = scores.squeeze(-1)
        if scores.dim() == 1:
            scores = scores.unsqueeze(0)
        
        # Create selection mask
        selection_mask = (scores > self.selection_threshold).float()
        
        # Process selected items
        processed = self.process_net(x)  # [batch, seq_len, embed_dim]
        
        # Apply selection mask
        output = processed * selection_mask.unsqueeze(-1) + x * (1 - selection_mask.unsqueeze(-1))
        
        return output, selection_mask


class FormulaTransformerBlock(nn.Module):
    """
    Transformer-style block using formula-based mechanisms.
    """
    
    def __init__(
        self,
        embedding_dim: int,
        formula_config: Optional[dict] = None,
        feedforward_dim: int = None,
        dropout: float = 0.1,
        use_formula_attention: bool = True,
    ):
        super().__init__()
        self.embedding_dim = embedding_dim
        feedforward_dim = feedforward_dim or (embedding_dim * 4)
        
        if use_formula_attention:
            self.attention = FormulaAttentionLayer(
                embedding_dim=embedding_dim,
                formula_config=formula_config,
                dropout=dropout
            )
        else:
            # Standard multi-head attention as fallback
            self.attention = nn.MultiheadAttention(
                embed_dim=embedding_dim,
                num_heads=8,
                dropout=dropout,
                batch_first=True
            )
        
        self.feedforward = nn.Sequential(
            nn.Linear(embedding_dim, feedforward_dim),
            nn.GELU(),
            nn.Dropout(dropout),
            nn.Linear(feedforward_dim, embedding_dim),
            nn.Dropout(dropout)
        )
        
        self.layer_norm1 = nn.LayerNorm(embedding_dim)
        self.layer_norm2 = nn.LayerNorm(embedding_dim)
    
    def forward(
        self,
        x: torch.Tensor,
        time_steps: Optional[torch.Tensor] = None,
        memory_buffer: Optional[torch.Tensor] = None,
        mask: Optional[torch.Tensor] = None
    ) -> Tuple[torch.Tensor, Optional[torch.Tensor]]:
        """
        Apply transformer block with formula-based attention.
        
        Args:
            x: Input [batch, seq_len, embed_dim]
            time_steps: [batch, seq_len] optional
            memory_buffer: [memory_size, embed_dim] optional
            mask: [batch, seq_len, seq_len] optional
            
        Returns:
            output: [batch, seq_len, embed_dim]
            attention_weights: Optional attention weights
        """
        # Self-attention
        if isinstance(self.attention, FormulaAttentionLayer):
            attn_out, attn_weights = self.attention(
                query=x,
                key=x,
                value=x,
                time_steps=time_steps,
                memory_buffer=memory_buffer,
                mask=mask
            )
        else:
            attn_out, attn_weights = self.attention(x, x, x, attn_mask=mask)
        
        x = self.layer_norm1(x + attn_out)
        
        # Feedforward
        ff_out = self.feedforward(x)
        x = self.layer_norm2(x + ff_out)
        
        return x, attn_weights

