"""
Gibbs-Normalized Salience Scoring Formula

z_i = α·ΔA_i + β·R_i + γ·M_i + η·ln(C_i+ε) - λ·t_i - κ·φ_i
z̃_i = (z_i - μ[z]) / (RMS[z] + ε)
S_i = B · softmax(z̃_i / τ)

Key properties:
- Additive logit-space gates (no multiplicative gating)
- RMSNorm for stabilization (not √d scaling)
- Log-space continuity: ln(C + ε)
- Subtractive fatigue: -κφ
- Capacity budget B (learnable or fixed)
- Proper scope-based normalization
- Shift invariance via mean centering
- Scale invariance via RMS normalization
- Budget conservation via softmax

MODERNIZED:
- RMSNorm instead of LayerNorm for better efficiency
"""

import torch
import torch.nn as nn
import torch.nn.functional as F
from typing import Tuple, Optional, Dict
import math
from .normalization import RMSNorm


class GibbsSalienceFormula(nn.Module):
    """
    Gibbs-normalized salience formula with additive logit-space gates.

    Computes:
        z_i = α·ΔA_i + β·R_i + γ·M_i + η·ln(C_i+ε) - λ·t_i - κ·φ_i
        z̃_i = (z_i - μ[z]) / (RMS[z] + ε)
        S_i = B · softmax(z̃_i / τ)

    Components:
    - ΔA: Novelty (information gain) - logit space output
    - R: Retention (long-term value) - logit space output
    - M: Payoff (immediate utility) - logit space output
    - C: Continuity (coherence) - positive, for ln(C+ε)
    - φ: Fatigue (redundancy) - positive, for -κφ
    - τ: Temperature parameter (learnable)
    - B: Capacity budget (learnable or fixed)
    """

    def __init__(
        self,
        embedding_dim: int,
        alpha: float = 1.0,      # Novelty coefficient
        beta: float = 0.8,       # Retention coefficient
        gamma: float = 0.6,      # Payoff coefficient
        eta: float = 0.5,        # Continuity coefficient
        lambda_decay: float = 0.1,  # Time decay coefficient
        kappa: float = 0.3,      # Fatigue coefficient
        temperature: float = 1.0,
        capacity_budget: float = 1.0,
        learnable_coefficients: bool = True,
        learnable_temperature: bool = True,
        learnable_budget: bool = False,
        component_hidden_dim: Optional[int] = None,
        epsilon: float = 1e-8,
    ):
        super().__init__()

        self.embedding_dim = embedding_dim
        self.component_hidden_dim = component_hidden_dim or (embedding_dim * 2)
        self.epsilon = epsilon

        # Learnable coefficients
        if learnable_coefficients:
            self.alpha = nn.Parameter(torch.tensor(alpha))
            self.beta = nn.Parameter(torch.tensor(beta))
            self.gamma = nn.Parameter(torch.tensor(gamma))
            self.eta = nn.Parameter(torch.tensor(eta))
            self.lambda_decay = nn.Parameter(torch.tensor(lambda_decay))
            self.kappa = nn.Parameter(torch.tensor(kappa))
        else:
            self.register_buffer('alpha', torch.tensor(alpha))
            self.register_buffer('beta', torch.tensor(beta))
            self.register_buffer('gamma', torch.tensor(gamma))
            self.register_buffer('eta', torch.tensor(eta))
            self.register_buffer('lambda_decay', torch.tensor(lambda_decay))
            self.register_buffer('kappa', torch.tensor(kappa))

        # Temperature parameter τ
        if learnable_temperature:
            self.temperature = nn.Parameter(torch.tensor(temperature))
        else:
            self.register_buffer('temperature', torch.tensor(temperature))

        # Capacity budget B
        if learnable_budget:
            self.capacity_budget = nn.Parameter(torch.tensor(capacity_budget))
        else:
            self.register_buffer('capacity_budget', torch.tensor(capacity_budget))

        # Component neural networks - outputs for logit space
        # Novelty (ΔA): Unbounded output for logit space
        self.novelty_net = nn.Sequential(
            nn.Linear(embedding_dim * 2, self.component_hidden_dim),
            RMSNorm(self.component_hidden_dim),
            nn.GELU(),
            nn.Dropout(0.1),
            nn.Linear(self.component_hidden_dim, self.component_hidden_dim // 2),
            RMSNorm(self.component_hidden_dim // 2),
            nn.GELU(),
            nn.Linear(self.component_hidden_dim // 2, 1),
            # No activation - unbounded logit space
        )

        # Retention (R): Unbounded output for logit space
        self.retention_net = nn.Sequential(
            nn.Linear(embedding_dim, self.component_hidden_dim),
            RMSNorm(self.component_hidden_dim),
            nn.GELU(),
            nn.Dropout(0.1),
            nn.Linear(self.component_hidden_dim, self.component_hidden_dim // 2),
            RMSNorm(self.component_hidden_dim // 2),
            nn.GELU(),
            nn.Linear(self.component_hidden_dim // 2, 1),
            # No activation - unbounded logit space
        )

        # Payoff (M): Unbounded output for logit space
        self.payoff_net = nn.Sequential(
            nn.Linear(embedding_dim, self.component_hidden_dim),
            RMSNorm(self.component_hidden_dim),
            nn.GELU(),
            nn.Dropout(0.1),
            nn.Linear(self.component_hidden_dim, self.component_hidden_dim // 2),
            RMSNorm(self.component_hidden_dim // 2),
            nn.GELU(),
            nn.Linear(self.component_hidden_dim // 2, 1),
            # No activation - unbounded logit space
        )

        # Continuity (C): Positive output for ln(C+ε)
        self.continuity_net = nn.Sequential(
            nn.Linear(embedding_dim * 2, self.component_hidden_dim),
            RMSNorm(self.component_hidden_dim),
            nn.GELU(),
            nn.Dropout(0.1),
            nn.Linear(self.component_hidden_dim, self.component_hidden_dim // 2),
            RMSNorm(self.component_hidden_dim // 2),
            nn.GELU(),
            nn.Linear(self.component_hidden_dim // 2, 1),
            nn.Softplus()  # Ensures positive output for log
        )

        # Fatigue (φ): Positive output for -κφ
        self.fatigue_net = nn.Sequential(
            nn.Linear(embedding_dim + 1, self.component_hidden_dim),  # +1 for max similarity
            RMSNorm(self.component_hidden_dim),
            nn.GELU(),
            nn.Dropout(0.1),
            nn.Linear(self.component_hidden_dim, self.component_hidden_dim // 2),
            RMSNorm(self.component_hidden_dim // 2),
            nn.GELU(),
            nn.Linear(self.component_hidden_dim // 2, 1),
            nn.Softplus()  # Ensures positive output
        )

    def compute_novelty(
        self,
        current: torch.Tensor,
        context: torch.Tensor
    ) -> torch.Tensor:
        """
        Compute ΔA: Novelty (information gain) - unbounded logit space

        Args:
            current: [batch*seq_len, dim]
            context: [batch*seq_len, dim]

        Returns:
            novelty: [batch*seq_len] unbounded
        """
        combined = torch.cat([current, context], dim=-1)
        novelty = self.novelty_net(combined).squeeze(-1)
        return novelty

    def compute_retention(self, x: torch.Tensor) -> torch.Tensor:
        """
        Compute R: Retention (long-term value) - unbounded logit space

        Args:
            x: [batch*seq_len, dim]

        Returns:
            retention: [batch*seq_len] unbounded
        """
        retention = self.retention_net(x).squeeze(-1)
        return retention

    def compute_payoff(self, x: torch.Tensor) -> torch.Tensor:
        """
        Compute M: Payoff (immediate utility) - unbounded logit space

        Args:
            x: [batch*seq_len, dim]

        Returns:
            payoff: [batch*seq_len] unbounded
        """
        payoff = self.payoff_net(x).squeeze(-1)
        return payoff

    def compute_continuity(
        self,
        current: torch.Tensor,
        context: torch.Tensor
    ) -> torch.Tensor:
        """
        Compute C: Continuity (coherence) - positive for ln(C+ε)

        Args:
            current: [batch*seq_len, dim]
            context: [batch*seq_len, dim]

        Returns:
            continuity: [batch*seq_len] positive
        """
        combined = torch.cat([current, context], dim=-1)
        continuity = self.continuity_net(combined).squeeze(-1)
        return continuity

    def compute_fatigue(
        self,
        current: torch.Tensor,
        memory_buffer: Optional[torch.Tensor] = None
    ) -> torch.Tensor:
        """
        Compute φ: Fatigue (redundancy) - positive for -κφ

        Args:
            current: [batch*seq_len, dim]
            memory_buffer: [memory_size, dim] optional

        Returns:
            fatigue: [batch*seq_len] positive
        """
        if memory_buffer is None or memory_buffer.shape[0] == 0:
            # No memory, no fatigue
            batch_size = current.shape[0]
            return torch.zeros(batch_size, device=current.device, dtype=current.dtype)

        # Ensure same device
        memory_buffer = memory_buffer.to(current.device)

        # Filter out zero rows
        memory_norm = memory_buffer.norm(dim=-1)
        valid_mask = memory_norm > 1e-6
        if not valid_mask.any():
            return torch.zeros(current.shape[0], device=current.device, dtype=current.dtype)

        memory_buffer = memory_buffer[valid_mask]

        # Compute cosine similarity to memory
        current_norm = F.normalize(current, p=2, dim=-1)
        memory_norm = F.normalize(memory_buffer, p=2, dim=-1)

        # Dot product: [batch*seq_len, memory_size]
        similarities = torch.matmul(current_norm, memory_norm.T)
        similarities = torch.clamp(similarities, -1.0, 1.0)

        # Max similarity to any memory item
        max_similarity = similarities.max(dim=-1)[0]  # [batch*seq_len]

        # Compute fatigue from max similarity
        fatigue_input = torch.cat([current, max_similarity.unsqueeze(-1)], dim=-1)
        fatigue = self.fatigue_net(fatigue_input).squeeze(-1)

        return fatigue

    def rms_normalize(
        self,
        z: torch.Tensor,
        dim: int = -1
    ) -> torch.Tensor:
        """
        RMS normalization: z̃_i = (z_i - μ[z]) / (RMS[z] + ε)

        Args:
            z: Raw logits [batch, seq_len]
            dim: Dimension to normalize over

        Returns:
            z_tilde: Normalized logits [batch, seq_len]
        """
        # Mean centering for shift invariance
        mean = z.mean(dim=dim, keepdim=True)
        z_centered = z - mean

        # RMS normalization for scale invariance
        rms = torch.sqrt(torch.mean(z_centered ** 2, dim=dim, keepdim=True) + self.epsilon)
        z_normalized = z_centered / rms

        return z_normalized

    def forward(
        self,
        current: torch.Tensor,
        context: torch.Tensor,
        time_steps: Optional[torch.Tensor] = None,
        memory_buffer: Optional[torch.Tensor] = None,
        return_components: bool = False
    ) -> Tuple[torch.Tensor, Optional[Dict]]:
        """
        Compute Gibbs-normalized salience formula.

        z_i = α·ΔA_i + β·R_i + γ·M_i + η·ln(C_i+ε) - λ·t_i - κ·φ_i
        z̃_i = (z_i - μ[z]) / (RMS[z] + ε)
        S_i = B · softmax(z̃_i / τ)

        Args:
            current: [batch, seq_len, dim] current embeddings
            context: [batch, seq_len, dim] context embeddings
            time_steps: [batch, seq_len] time steps (optional)
            memory_buffer: [memory_size, dim] memory for fatigue (optional)
            return_components: Whether to return component breakdown

        Returns:
            scores: [batch, seq_len] salience scores (sum to B per scope)
            components: Optional dict with component values
        """
        batch_size, seq_len, embed_dim = current.shape

        # Flatten for efficient computation
        current_flat = current.view(-1, embed_dim)  # [batch*seq_len, dim]
        context_flat = context.view(-1, embed_dim)  # [batch*seq_len, dim]

        # Compute components
        novelty = self.compute_novelty(current_flat, context_flat)  # [batch*seq_len] unbounded
        retention = self.compute_retention(current_flat)  # [batch*seq_len] unbounded
        payoff = self.compute_payoff(current_flat)  # [batch*seq_len] unbounded
        continuity = self.compute_continuity(current_flat, context_flat)  # [batch*seq_len] positive
        fatigue = self.compute_fatigue(current_flat, memory_buffer)  # [batch*seq_len] positive

        # Additive logit-space combination
        # z_i = α·ΔA_i + β·R_i + γ·M_i + η·ln(C_i+ε) - λ·t_i - κ·φ_i

        # Positive terms
        z = self.alpha * novelty
        z = z + self.beta * retention
        z = z + self.gamma * payoff
        z = z + self.eta * torch.log(continuity + self.epsilon)

        # Negative terms
        if time_steps is not None:
            time_steps_flat = time_steps.view(-1).float()
            z = z - torch.abs(self.lambda_decay) * time_steps_flat

        z = z - torch.abs(self.kappa) * fatigue

        # Reshape to [batch, seq_len]
        z = z.view(batch_size, seq_len)

        # RMS normalization: z̃_i = (z_i - μ[z]) / (RMS[z] + ε)
        z_normalized = self.rms_normalize(z, dim=-1)

        # Temperature scaling
        temp = torch.abs(self.temperature) + 0.1  # Ensure positive
        z_scaled = z_normalized / temp

        # Softmax for probability distribution (ensures conservation)
        probabilities = F.softmax(z_scaled, dim=-1)

        # Scale by capacity budget: S_i = B · softmax(z̃_i / τ)
        budget = torch.abs(self.capacity_budget)
        scores = budget * probabilities

        # Prepare component dictionary if requested
        components = None
        if return_components:
            components = {
                'novelty': novelty.view(batch_size, seq_len),
                'retention': retention.view(batch_size, seq_len),
                'payoff': payoff.view(batch_size, seq_len),
                'continuity': continuity.view(batch_size, seq_len),
                'fatigue': fatigue.view(batch_size, seq_len),
                'log_continuity': torch.log(continuity + self.epsilon).view(batch_size, seq_len),
                'raw_logits': z,
                'normalized_logits': z_normalized,
                'probabilities': probabilities,
                'final_scores': scores,
                'coefficients': {
                    'alpha': self.alpha.item() if isinstance(self.alpha, torch.Tensor) else self.alpha,
                    'beta': self.beta.item() if isinstance(self.beta, torch.Tensor) else self.beta,
                    'gamma': self.gamma.item() if isinstance(self.gamma, torch.Tensor) else self.gamma,
                    'eta': self.eta.item() if isinstance(self.eta, torch.Tensor) else self.eta,
                    'lambda': self.lambda_decay.item() if isinstance(self.lambda_decay, torch.Tensor) else self.lambda_decay,
                    'kappa': self.kappa.item() if isinstance(self.kappa, torch.Tensor) else self.kappa,
                    'temperature': temp.item() if isinstance(temp, torch.Tensor) else temp,
                    'budget': budget.item() if isinstance(budget, torch.Tensor) else budget,
                }
            }

        return scores, components

    def compute_salience_loss(
        self,
        scores: torch.Tensor,
        targets: Optional[torch.Tensor] = None
    ) -> torch.Tensor:
        """
        Compute regularization loss to encourage proper salience distribution.

        Args:
            scores: [batch, seq_len] salience scores
            targets: [batch, seq_len] optional target salience distribution

        Returns:
            loss: Scalar loss value
        """
        loss = 0.0

        # Encourage diversity: scores shouldn't all be the same
        score_std = scores.std(dim=-1).mean()
        diversity_loss = torch.exp(-score_std)  # Penalize low diversity
        loss = loss + 0.1 * diversity_loss

        # Encourage sparsity: not everything should be highly salient
        sparsity_loss = (scores ** 2).mean()
        loss = loss + 0.01 * sparsity_loss

        # If targets provided, match them
        if targets is not None:
            target_loss = F.mse_loss(scores, targets)
            loss = loss + target_loss

        return loss

    # === Sanity Check Methods ===

    def scale_sweep(
        self,
        current: torch.Tensor,
        context: torch.Tensor,
        scales: Optional[torch.Tensor] = None
    ) -> Dict[str, torch.Tensor]:
        """
        Test scale invariance: S(c·x) ≈ S(x) for various c

        Args:
            current: [batch, seq_len, dim]
            context: [batch, seq_len, dim]
            scales: Optional tensor of scale factors to test

        Returns:
            results: Dict with scale factors and corresponding score variations
        """
        if scales is None:
            scales = torch.tensor([0.1, 0.5, 1.0, 2.0, 10.0], device=current.device)

        baseline_scores, _ = self.forward(current, context)

        variations = []
        for scale in scales:
            scaled_current = current * scale
            scaled_context = context * scale
            scaled_scores, _ = self.forward(scaled_current, scaled_context)

            # Measure relative difference
            rel_diff = torch.abs(scaled_scores - baseline_scores) / (baseline_scores + self.epsilon)
            variations.append(rel_diff.mean().item())

        return {
            'scales': scales,
            'mean_relative_variations': torch.tensor(variations),
            'max_variation': max(variations),
            'scale_invariant': max(variations) < 0.1,  # <10% variation
        }

    def shift_invariance(
        self,
        current: torch.Tensor,
        context: torch.Tensor,
        shifts: Optional[torch.Tensor] = None
    ) -> Dict[str, torch.Tensor]:
        """
        Test shift invariance: S(x + c) ≈ S(x) for various c

        Args:
            current: [batch, seq_len, dim]
            context: [batch, seq_len, dim]
            shifts: Optional tensor of shift values to test

        Returns:
            results: Dict with shift values and corresponding score variations
        """
        if shifts is None:
            shifts = torch.tensor([-10.0, -1.0, 0.0, 1.0, 10.0], device=current.device)

        baseline_scores, _ = self.forward(current, context)

        variations = []
        for shift in shifts:
            shifted_current = current + shift
            shifted_context = context + shift
            shifted_scores, _ = self.forward(shifted_current, shifted_context)

            # Measure relative difference
            rel_diff = torch.abs(shifted_scores - baseline_scores) / (baseline_scores + self.epsilon)
            variations.append(rel_diff.mean().item())

        return {
            'shifts': shifts,
            'mean_relative_variations': torch.tensor(variations),
            'max_variation': max(variations),
            'shift_invariant': max(variations) < 0.1,  # <10% variation
        }

    def budget_conservation(
        self,
        current: torch.Tensor,
        context: torch.Tensor,
        tolerance: float = 1e-5
    ) -> Dict[str, torch.Tensor]:
        """
        Test budget conservation: sum(S_i) = B for each scope

        Args:
            current: [batch, seq_len, dim]
            context: [batch, seq_len, dim]
            tolerance: Acceptable deviation from budget

        Returns:
            results: Dict with budget sums and conservation status
        """
        scores, _ = self.forward(current, context)

        # Sum over sequence dimension (within each scope)
        sums = scores.sum(dim=-1)  # [batch]

        # Expected budget
        expected_budget = torch.abs(self.capacity_budget)

        # Check conservation
        deviations = torch.abs(sums - expected_budget)
        max_deviation = deviations.max().item()

        return {
            'budget_sums': sums,
            'expected_budget': expected_budget.item(),
            'deviations': deviations,
            'max_deviation': max_deviation,
            'conserved': max_deviation < tolerance,
        }

    def friction_monotonicity(
        self,
        current: torch.Tensor,
        context: torch.Tensor,
        time_steps: Optional[torch.Tensor] = None,
        num_samples: int = 10
    ) -> Dict[str, torch.Tensor]:
        """
        Test friction monotonicity: S decreases with time and fatigue

        Args:
            current: [batch, seq_len, dim]
            context: [batch, seq_len, dim]
            time_steps: [batch, seq_len] optional
            num_samples: Number of time/fatigue samples to test

        Returns:
            results: Dict with monotonicity test results
        """
        batch_size, seq_len, _ = current.shape

        # Test time decay monotonicity
        if time_steps is None:
            time_values = torch.linspace(0, 10, num_samples, device=current.device)
        else:
            time_values = torch.linspace(time_steps.min(), time_steps.max(), num_samples, device=current.device)

        time_scores = []
        for t in time_values:
            t_steps = torch.full((batch_size, seq_len), t.item(), device=current.device)
            scores, _ = self.forward(current, context, time_steps=t_steps)
            time_scores.append(scores.mean().item())

        time_scores = torch.tensor(time_scores)
        time_monotonic = torch.all(time_scores[1:] <= time_scores[:-1] + 1e-6).item()

        # Test fatigue suppression (create synthetic memory with increasing similarity)
        fatigue_scores = []
        for i in range(num_samples):
            # Create memory buffer with increasing similarity to current
            similarity = i / (num_samples - 1)  # 0 to 1
            memory_buffer = current[0, 0:1, :] * similarity + torch.randn_like(current[0, 0:1, :]) * (1 - similarity)
            memory_buffer = memory_buffer.repeat(10, 1)  # Create buffer of 10 items

            scores, _ = self.forward(current, context, memory_buffer=memory_buffer)
            fatigue_scores.append(scores.mean().item())

        fatigue_scores = torch.tensor(fatigue_scores)
        fatigue_monotonic = torch.all(fatigue_scores[1:] <= fatigue_scores[:-1] + 1e-6).item()

        return {
            'time_values': time_values,
            'time_scores': time_scores,
            'time_monotonic': time_monotonic,
            'fatigue_scores': fatigue_scores,
            'fatigue_monotonic': fatigue_monotonic,
            'overall_monotonic': time_monotonic and fatigue_monotonic,
        }

    def run_all_sanity_checks(
        self,
        current: torch.Tensor,
        context: torch.Tensor,
        time_steps: Optional[torch.Tensor] = None,
        verbose: bool = True
    ) -> Dict[str, Dict]:
        """
        Run all sanity checks and return comprehensive results.

        Args:
            current: [batch, seq_len, dim]
            context: [batch, seq_len, dim]
            time_steps: [batch, seq_len] optional
            verbose: Whether to print results

        Returns:
            results: Dict with all sanity check results
        """
        results = {
            'scale_sweep': self.scale_sweep(current, context),
            'shift_invariance': self.shift_invariance(current, context),
            'budget_conservation': self.budget_conservation(current, context),
            'friction_monotonicity': self.friction_monotonicity(current, context, time_steps),
        }

        if verbose:
            print("\n=== Gibbs Salience Formula Sanity Checks ===\n")

            print("1. Scale Invariance:")
            print(f"   Max variation: {results['scale_sweep']['max_variation']:.6f}")
            print(f"   Status: {'✓ PASS' if results['scale_sweep']['scale_invariant'] else '✗ FAIL'}")

            print("\n2. Shift Invariance:")
            print(f"   Max variation: {results['shift_invariance']['max_variation']:.6f}")
            print(f"   Status: {'✓ PASS' if results['shift_invariance']['shift_invariant'] else '✗ FAIL'}")

            print("\n3. Budget Conservation:")
            print(f"   Max deviation: {results['budget_conservation']['max_deviation']:.8f}")
            print(f"   Status: {'✓ PASS' if results['budget_conservation']['conserved'] else '✗ FAIL'}")

            print("\n4. Friction Monotonicity:")
            print(f"   Time monotonic: {'✓ PASS' if results['friction_monotonicity']['time_monotonic'] else '✗ FAIL'}")
            print(f"   Fatigue monotonic: {'✓ PASS' if results['friction_monotonicity']['fatigue_monotonic'] else '✗ FAIL'}")
            print(f"   Status: {'✓ PASS' if results['friction_monotonicity']['overall_monotonic'] else '✗ FAIL'}")

            # Overall summary
            all_pass = (
                results['scale_sweep']['scale_invariant'] and
                results['shift_invariance']['shift_invariant'] and
                results['budget_conservation']['conserved'] and
                results['friction_monotonicity']['overall_monotonic']
            )
            print(f"\n{'='*45}")
            print(f"Overall: {'✓ ALL CHECKS PASSED' if all_pass else '✗ SOME CHECKS FAILED'}")
            print(f"{'='*45}\n")

        return results
