"""
Core scoring formula implementation: S' = (w₁·ΔA + w₂·R + w₃·M) × C × e^(-λt) × (1 - kφ)
"""

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


class ScoringFormula(nn.Module):
    """
    Implements the novel scoring formula for information selection and processing.
    
    S' = (w₁·ΔA + w₂·R + w₃·M) × C × e^(-λt) × (1 - kφ)
    
    Where:
    - ΔA = Novelty (information gain)
    - R = Retention (long-term value)
    - M = Payoff (immediate utility)
    - C = Continuity (coherence)
    - φ = Fatigue (redundancy)
    """
    
    def __init__(
        self,
        embedding_dim: int,
        novelty_dim: Optional[int] = None,
        w1: float = 0.4,  # Weight for novelty
        w2: float = 0.3,  # Weight for retention
        w3: float = 0.3,  # Weight for payoff
        lambda_decay: float = 0.1,  # Time decay parameter
        k_fatigue: float = 0.2,  # Fatigue coefficient
        learnable_weights: bool = True,
        learnable_decay: bool = True,
    ):
        super().__init__()
        self.embedding_dim = embedding_dim
        self.novelty_dim = novelty_dim or embedding_dim
        
        # Learnable parameters
        if learnable_weights:
            self.w1 = nn.Parameter(torch.tensor(w1))
            self.w2 = nn.Parameter(torch.tensor(w2))
            self.w3 = nn.Parameter(torch.tensor(w3))
        else:
            self.register_buffer('w1', torch.tensor(w1))
            self.register_buffer('w2', torch.tensor(w2))
            self.register_buffer('w3', torch.tensor(w3))
            
        if learnable_decay:
            self.lambda_decay = nn.Parameter(torch.tensor(lambda_decay))
            self.k_fatigue = nn.Parameter(torch.tensor(k_fatigue))
        else:
            self.register_buffer('lambda_decay', torch.tensor(lambda_decay))
            self.register_buffer('k_fatigue', torch.tensor(k_fatigue))
        
        # Neural networks for computing components
        # Novelty (ΔA): Information gain - how much new information does this provide?
        self.novelty_net = nn.Sequential(
            nn.Linear(embedding_dim * 2, self.novelty_dim),
            nn.LayerNorm(self.novelty_dim),
            nn.GELU(),
            nn.Linear(self.novelty_dim, embedding_dim),
            nn.Sigmoid()
        )
        
        # Retention (R): Long-term value - how valuable will this be in the future?
        self.retention_net = nn.Sequential(
            nn.Linear(embedding_dim, embedding_dim),
            nn.LayerNorm(embedding_dim),
            nn.GELU(),
            nn.Linear(embedding_dim, embedding_dim),
            nn.Tanh()
        )
        
        # Payoff (M): Immediate utility - how useful is this right now?
        self.payoff_net = nn.Sequential(
            nn.Linear(embedding_dim, embedding_dim),
            nn.LayerNorm(embedding_dim),
            nn.GELU(),
            nn.Linear(embedding_dim, embedding_dim),
            nn.Tanh()
        )
        
        # Continuity (C): Coherence - how well does this fit with context?
        self.continuity_net = nn.Sequential(
            nn.Linear(embedding_dim * 2, embedding_dim),
            nn.LayerNorm(embedding_dim),
            nn.GELU(),
            nn.Linear(embedding_dim, embedding_dim),
            nn.Sigmoid()
        )
        
        # Fatigue (φ): Redundancy - how similar is this to recent items?
        self.fatigue_net = nn.Sequential(
            nn.Linear(embedding_dim * 2, embedding_dim),
            nn.LayerNorm(embedding_dim),
            nn.GELU(),
            nn.Linear(embedding_dim, 1),
            nn.Sigmoid()
        )
        
        # Memory state for tracking recent items (for fatigue calculation)
        self.memory_buffer_size = 32
        
    def compute_novelty(self, current: torch.Tensor, context: torch.Tensor) -> torch.Tensor:
        """
        Compute ΔA: Novelty (information gain)
        
        Measures how much new information the current input provides
        relative to the context.
        """
        # Concatenate current and context
        combined = torch.cat([current, context], dim=-1)
        # Compute information gain
        novelty = self.novelty_net(combined)
        return novelty
    
    def compute_retention(self, x: torch.Tensor) -> torch.Tensor:
        """
        Compute R: Retention (long-term value)
        
        Estimates the long-term importance and memorability of the input.
        """
        return self.retention_net(x)
    
    def compute_payoff(self, x: torch.Tensor) -> torch.Tensor:
        """
        Compute M: Payoff (immediate utility)
        
        Measures the immediate usefulness and relevance of the input.
        """
        return self.payoff_net(x)
    
    def compute_continuity(self, current: torch.Tensor, context: torch.Tensor) -> torch.Tensor:
        """
        Compute C: Continuity (coherence)
        
        Measures how well the current input fits with the existing context.
        """
        combined = torch.cat([current, context], dim=-1)
        continuity = self.continuity_net(combined)
        # Aggregate to a scalar per item (mean across embedding dim)
        continuity = continuity.mean(dim=-1, keepdim=True)
        return continuity
    
    def compute_fatigue(
        self,
        current: torch.Tensor,
        memory_buffer: Optional[torch.Tensor] = None
    ) -> torch.Tensor:
        """
        Compute φ: Fatigue (redundancy)
        
        Measures how redundant the current input is relative to recent items.
        If memory_buffer is None, returns low fatigue.
        """
        if memory_buffer is None or memory_buffer.shape[0] == 0:
            # No memory, so no fatigue
            return torch.zeros(current.shape[0], 1, device=current.device, dtype=current.dtype)
        
        # Ensure memory_buffer is on same device and has valid content
        memory_buffer = memory_buffer.to(current.device)
        
        # Filter out zero rows in memory buffer
        memory_norm = memory_buffer.norm(dim=-1)
        valid_indices = memory_norm > 1e-6
        if valid_indices.sum() == 0:
            return torch.zeros(current.shape[0], 1, device=current.device, dtype=current.dtype)
        
        memory_buffer = memory_buffer[valid_indices]
        
        # Compute similarity to recent items
        # More stable computation: compute dot product then normalize
        current_norm = F.normalize(current, p=2, dim=-1)  # [batch, dim]
        memory_norm = F.normalize(memory_buffer, p=2, dim=-1)  # [memory_size, dim]
        
        # Compute cosine similarity: dot product of normalized vectors
        similarities = torch.matmul(current_norm, memory_norm.T)  # [batch, memory_size]
        
        # Clamp to valid range
        similarities = torch.clamp(similarities, -1.0, 1.0)
        
        # Take max similarity (most similar item in memory)
        max_similarity = similarities.max(dim=-1)[0]  # [batch]
        
        # Compute fatigue from similarity
        max_sim_expanded = max_similarity.unsqueeze(-1).expand(-1, current.shape[-1])
        fatigue_input = torch.cat([current, max_sim_expanded], dim=-1)
        
        fatigue = self.fatigue_net(fatigue_input)  # [batch, 1]
        return fatigue
    
    def forward(
        self,
        current: torch.Tensor,
        context: torch.Tensor,
        time_steps: Optional[torch.Tensor] = None,
        memory_buffer: Optional[torch.Tensor] = None
    ) -> Tuple[torch.Tensor, dict]:
        """
        Compute the full scoring formula S'
        
        Args:
            current: Current input embeddings [batch, seq_len, embedding_dim] or [batch, embedding_dim]
            context: Context embeddings [batch, seq_len, embedding_dim] or [batch, embedding_dim]
            time_steps: Time step for each item [batch, seq_len] or [batch] (optional)
            memory_buffer: Recent items for fatigue computation [memory_size, embedding_dim] (optional)
            
        Returns:
            scores: Computed scores [batch, seq_len] or [batch]
            components: Dictionary with all component values for analysis
        """
        # Ensure 3D tensors: [batch, seq_len, dim]
        if current.dim() == 2:
            current = current.unsqueeze(1)
            context = context.unsqueeze(1)
            was_2d = True
        else:
            was_2d = False
        
        batch_size, seq_len, embed_dim = current.shape
        
        # Compute components
        # Novelty: [batch, seq_len, dim]
        novelty = self.compute_novelty(current, context)
        # Aggregate novelty to scalar (mean across embedding)
        novelty_scalar = novelty.mean(dim=-1)  # [batch, seq_len]
        
        # Retention: [batch, seq_len, dim]
        retention = self.compute_retention(current)
        retention_scalar = retention.mean(dim=-1)  # [batch, seq_len]
        
        # Payoff: [batch, seq_len, dim]
        payoff = self.compute_payoff(current)
        payoff_scalar = payoff.mean(dim=-1)  # [batch, seq_len]
        
        # Combined weighted sum: [batch, seq_len]
        weighted_sum = (
            self.w1 * novelty_scalar +
            self.w2 * retention_scalar +
            self.w3 * payoff_scalar
        )
        
        # Continuity: [batch, seq_len, 1] -> [batch, seq_len]
        continuity = self.compute_continuity(current, context).squeeze(-1)
        
        # Fatigue: [batch, seq_len, 1] -> [batch, seq_len]
        if memory_buffer is not None:
            # Reshape for fatigue computation
            current_flat = current.view(-1, embed_dim)  # [batch * seq_len, dim]
            fatigue_flat = self.compute_fatigue(current_flat, memory_buffer)
            fatigue = fatigue_flat.view(batch_size, seq_len)  # [batch, seq_len]
        else:
            fatigue = torch.zeros(batch_size, seq_len, device=current.device)
        
        # Time decay: e^(-λt)
        if time_steps is not None:
            if time_steps.dim() == 1:
                time_steps = time_steps.unsqueeze(1)  # [batch, 1]
            time_decay = torch.exp(-self.lambda_decay * time_steps.float())
            if time_decay.shape[1] == 1:
                time_decay = time_decay.expand(-1, seq_len)  # [batch, seq_len]
        else:
            # Default: no time decay (t=0)
            time_decay = torch.ones(batch_size, seq_len, device=current.device)
        
        # Apply softmax to weights for stability
        weights_sum = torch.abs(self.w1) + torch.abs(self.w2) + torch.abs(self.w3)
        w1_norm = torch.abs(self.w1) / (weights_sum + 1e-8)
        w2_norm = torch.abs(self.w2) / (weights_sum + 1e-8)
        w3_norm = torch.abs(self.w3) / (weights_sum + 1e-8)
        
        weighted_sum = (
            w1_norm * novelty_scalar +
            w2_norm * retention_scalar +
            w3_norm * payoff_scalar
        )
        
        # Final score: S' = (w₁·ΔA + w₂·R + w₃·M) × C × e^(-λt) × (1 - kφ)
        scores = (
            weighted_sum *
            continuity *
            time_decay *
            (1 - self.k_fatigue * fatigue)
        )
        
        # Squeeze if input was 2D
        if was_2d:
            scores = scores.squeeze(1)
            novelty_scalar = novelty_scalar.squeeze(1)
            retention_scalar = retention_scalar.squeeze(1)
            payoff_scalar = payoff_scalar.squeeze(1)
            continuity = continuity.squeeze(1)
            fatigue = fatigue.squeeze(1)
            if time_steps is not None:
                time_decay = time_decay.squeeze(1)
        
        components = {
            'novelty': novelty_scalar,
            'retention': retention_scalar,
            'payoff': payoff_scalar,
            'continuity': continuity,
            'fatigue': fatigue,
            'time_decay': time_decay if time_steps is not None else None,
            'weighted_sum': weighted_sum,
            'weights': {
                'w1': self.w1.item() if not isinstance(self.w1, torch.Tensor) else self.w1.data,
                'w2': self.w2.item() if not isinstance(self.w2, torch.Tensor) else self.w2.data,
                'w3': self.w3.item() if not isinstance(self.w3, torch.Tensor) else self.w3.data,
            }
        }
        
        return scores, components

