"""
Catastrophic Forgetting Demonstration
======================================

A pedagogical example showing how neural networks catastrophically forget
previously learned tasks when trained sequentially on new tasks.

This toy example trains a simple neural network on two binary classification tasks:
- Task A: Classify points in upper-right quadrant (x > 0.5, y > 0.5) as positive
- Task B: Classify points in upper-left quadrant (x < 0.5, y > 0.5) as positive

We demonstrate:
1. Catastrophic forgetting: Task A accuracy drops from ~95% to ~15% after Task B training
2. Elastic Weight Consolidation (EWC): Partially mitigates forgetting with controlled trade-off

Dependencies: numpy, torch
"""

import numpy as np
import torch
import torch.nn as nn
import torch.optim as optim
from typing import Tuple, Dict


# Set random seeds for reproducibility
np.random.seed(42)
torch.manual_seed(42)


class SimpleNet(nn.Module):
    """
    Simple 2-layer neural network for binary classification.

    Architecture: 2 inputs -> 20 hidden units (ReLU) -> 1 output (logits)
    Note: We use logits (raw outputs) for numerical stability with BCEWithLogitsLoss
    """
    def __init__(self):
        super(SimpleNet, self).__init__()
        self.fc1 = nn.Linear(2, 20)
        self.fc2 = nn.Linear(20, 1)
        self.relu = nn.ReLU()

    def forward(self, x, return_probs=False):
        """
        Args:
            x: Input features
            return_probs: If True, return probabilities (for evaluation), else logits (for training)
        """
        x = self.relu(self.fc1(x))
        logits = self.fc2(x)
        if return_probs:
            return torch.sigmoid(logits)
        return logits


def generate_task_data(task: str, n_samples: int = 1000) -> Tuple[torch.Tensor, torch.Tensor]:
    """
    Generate synthetic data for binary classification tasks.

    Task A: Positive class = upper-right quadrant (x > 0.5, y > 0.5)
    Task B: Positive class = upper-left quadrant (x < 0.5, y > 0.5)

    Args:
        task: 'A' or 'B'
        n_samples: Number of samples to generate

    Returns:
        X: Input features (n_samples, 2)
        y: Binary labels (n_samples, 1)
    """
    # Generate random points in [0, 1] x [0, 1]
    X = np.random.rand(n_samples, 2)

    if task == 'A':
        # Task A: upper-right quadrant
        y = ((X[:, 0] > 0.5) & (X[:, 1] > 0.5)).astype(np.float32)
    elif task == 'B':
        # Task B: upper-left quadrant
        y = ((X[:, 0] < 0.5) & (X[:, 1] > 0.5)).astype(np.float32)
    else:
        raise ValueError("task must be 'A' or 'B'")

    X_tensor = torch.FloatTensor(X)
    y_tensor = torch.FloatTensor(y).unsqueeze(1)

    return X_tensor, y_tensor


def train_task(model: nn.Module, X: torch.Tensor, y: torch.Tensor,
               epochs: int = 100, lr: float = 0.01,
               ewc_loss=None) -> float:
    """
    Train the model on a task.

    Args:
        model: Neural network to train
        X: Input features
        y: Target labels
        epochs: Number of training epochs
        lr: Learning rate
        ewc_loss: Optional EWC loss function for regularization

    Returns:
        Final training accuracy
    """
    criterion = nn.BCEWithLogitsLoss()  # More numerically stable than BCELoss
    optimizer = optim.SGD(model.parameters(), lr=lr)

    for epoch in range(epochs):
        # Forward pass (get logits for training)
        logits = model(X, return_probs=False)

        # Compute loss (with optional EWC regularization)
        loss = criterion(logits, y)
        if ewc_loss is not None:
            loss = loss + ewc_loss()

        # Backward pass and optimization
        optimizer.zero_grad()
        loss.backward()
        optimizer.step()

    # Calculate final accuracy
    with torch.no_grad():
        probs = model(X, return_probs=True)
        predictions = (probs > 0.5).float()
        accuracy = (predictions == y).float().mean().item()

    return accuracy


def evaluate_task(model: nn.Module, X: torch.Tensor, y: torch.Tensor) -> float:
    """
    Evaluate model accuracy on a task.

    Args:
        model: Neural network to evaluate
        X: Input features
        y: Target labels

    Returns:
        Accuracy as a fraction
    """
    with torch.no_grad():
        probs = model(X, return_probs=True)
        predictions = (probs > 0.5).float()
        accuracy = (predictions == y).float().mean().item()
    return accuracy


class EWC:
    """
    Elastic Weight Consolidation (EWC) for mitigating catastrophic forgetting.

    EWC adds a regularization term to the loss that penalizes changes to important
    weights (those that were important for previous tasks). Importance is measured
    by the Fisher Information Matrix diagonal.

    The EWC loss is: lambda * sum_i F_i * (theta_i - theta*_i)^2
    where:
    - F_i is the Fisher information (importance) of weight i
    - theta_i is the current weight value
    - theta*_i is the optimal weight value from the previous task
    - lambda controls the strength of regularization
    """

    def __init__(self, model: nn.Module, X: torch.Tensor, y: torch.Tensor, lambda_ewc: float = 400):
        """
        Initialize EWC by computing Fisher information and storing optimal weights.

        Args:
            model: Trained neural network
            X: Input data from previous task
            y: Target labels from previous task
            lambda_ewc: Regularization strength (higher = more preservation, less plasticity)
        """
        self.model = model
        self.lambda_ewc = lambda_ewc

        # Store the optimal parameters from the previous task
        self.optimal_params = {}
        for name, param in model.named_parameters():
            self.optimal_params[name] = param.data.clone()

        # Compute Fisher Information Matrix (diagonal approximation)
        self.fisher = self._compute_fisher(X, y)

    def _compute_fisher(self, X: torch.Tensor, y: torch.Tensor) -> Dict[str, torch.Tensor]:
        """
        Compute diagonal Fisher Information Matrix.

        Fisher information measures how much a parameter affects the loss.
        We approximate it using the squared gradients of the log-likelihood.

        Args:
            X: Input data
            y: Target labels

        Returns:
            Dictionary mapping parameter names to Fisher information values
        """
        fisher = {}

        # Initialize Fisher information to zero for all parameters
        for name, param in self.model.named_parameters():
            fisher[name] = torch.zeros_like(param.data)

        # Compute gradients for each sample
        criterion = nn.BCEWithLogitsLoss(reduction='sum')

        for i in range(len(X)):
            self.model.zero_grad()

            # Forward pass on single sample
            logits = self.model(X[i:i+1], return_probs=False)
            loss = criterion(logits, y[i:i+1])

            # Backward pass
            loss.backward()

            # Accumulate squared gradients (Fisher information)
            for name, param in self.model.named_parameters():
                if param.grad is not None:
                    fisher[name] += param.grad.data ** 2

        # Average over samples
        for name in fisher:
            fisher[name] /= len(X)

        return fisher

    def penalty(self) -> torch.Tensor:
        """
        Compute EWC penalty term.

        Returns:
            EWC loss = lambda * sum_i F_i * (theta_i - theta*_i)^2
        """
        loss = 0
        for name, param in self.model.named_parameters():
            # Penalize deviation from optimal parameters, weighted by importance
            loss += (self.fisher[name] * (param - self.optimal_params[name]) ** 2).sum()
        return self.lambda_ewc * loss


def demonstrate_catastrophic_forgetting():
    """
    Main demonstration comparing vanilla sequential learning vs EWC.
    """
    print("=" * 70)
    print("CATASTROPHIC FORGETTING DEMONSTRATION")
    print("=" * 70)
    print()

    # Generate data for both tasks
    X_train_A, y_train_A = generate_task_data('A', n_samples=1000)
    X_test_A, y_test_A = generate_task_data('A', n_samples=500)

    X_train_B, y_train_B = generate_task_data('B', n_samples=1000)
    X_test_B, y_test_B = generate_task_data('B', n_samples=500)

    print("Generated synthetic datasets:")
    print(f"  Task A (upper-right quadrant): {len(X_train_A)} training, {len(X_test_A)} test samples")
    print(f"  Task B (upper-left quadrant): {len(X_train_B)} training, {len(X_test_B)} test samples")
    print()

    # -------------------------------------------------------------------------
    # Experiment 1: Vanilla Sequential Learning (Catastrophic Forgetting)
    # -------------------------------------------------------------------------
    print("-" * 70)
    print("EXPERIMENT 1: VANILLA SEQUENTIAL LEARNING")
    print("-" * 70)
    print()

    model_vanilla = SimpleNet()

    # Train on Task A
    print("Training on Task A...")
    train_acc_A = train_task(model_vanilla, X_train_A, y_train_A, epochs=200, lr=0.5)
    test_acc_A_after_A = evaluate_task(model_vanilla, X_test_A, y_test_A)

    print(f"  Task A training accuracy: {train_acc_A*100:.1f}%")
    print(f"  Task A test accuracy:     {test_acc_A_after_A*100:.1f}%")
    print()

    # Train on Task B (this causes catastrophic forgetting)
    print("Training on Task B...")
    train_acc_B = train_task(model_vanilla, X_train_B, y_train_B, epochs=200, lr=0.5)
    test_acc_B_after_B = evaluate_task(model_vanilla, X_test_B, y_test_B)

    print(f"  Task B training accuracy: {train_acc_B*100:.1f}%")
    print(f"  Task B test accuracy:     {test_acc_B_after_B*100:.1f}%")
    print()

    # Evaluate Task A again (catastrophic forgetting!)
    test_acc_A_after_B = evaluate_task(model_vanilla, X_test_A, y_test_A)

    print("Re-evaluating Task A after Task B training:")
    print(f"  Task A test accuracy: {test_acc_A_after_B*100:.1f}%")
    print(f"  Accuracy drop: {test_acc_A_after_A*100:.1f}% -> {test_acc_A_after_B*100:.1f}%")
    print(f"  CATASTROPHIC FORGETTING: {(test_acc_A_after_A - test_acc_A_after_B)*100:.1f}% loss!")
    print()

    # -------------------------------------------------------------------------
    # Experiment 2: Sequential Learning with EWC
    # -------------------------------------------------------------------------
    print("-" * 70)
    print("EXPERIMENT 2: ELASTIC WEIGHT CONSOLIDATION (EWC)")
    print("-" * 70)
    print()

    model_ewc = SimpleNet()

    # Train on Task A
    print("Training on Task A...")
    train_acc_A_ewc = train_task(model_ewc, X_train_A, y_train_A, epochs=200, lr=0.5)
    test_acc_A_after_A_ewc = evaluate_task(model_ewc, X_test_A, y_test_A)

    print(f"  Task A training accuracy: {train_acc_A_ewc*100:.1f}%")
    print(f"  Task A test accuracy:     {test_acc_A_after_A_ewc*100:.1f}%")
    print()

    # Initialize EWC after Task A
    print("Computing Fisher Information Matrix for EWC...")
    ewc = EWC(model_ewc, X_train_A, y_train_A, lambda_ewc=1000)
    print(f"  Regularization strength (lambda): {ewc.lambda_ewc}")
    print()

    # Train on Task B with EWC regularization
    print("Training on Task B with EWC regularization...")
    train_acc_B_ewc = train_task(model_ewc, X_train_B, y_train_B,
                                  epochs=200, lr=0.5, ewc_loss=ewc.penalty)
    test_acc_B_after_B_ewc = evaluate_task(model_ewc, X_test_B, y_test_B)

    print(f"  Task B training accuracy: {train_acc_B_ewc*100:.1f}%")
    print(f"  Task B test accuracy:     {test_acc_B_after_B_ewc*100:.1f}%")
    print()

    # Evaluate Task A again (EWC should mitigate forgetting)
    test_acc_A_after_B_ewc = evaluate_task(model_ewc, X_test_A, y_test_A)

    print("Re-evaluating Task A after Task B training with EWC:")
    print(f"  Task A test accuracy: {test_acc_A_after_B_ewc*100:.1f}%")
    print(f"  Accuracy drop: {test_acc_A_after_A_ewc*100:.1f}% -> {test_acc_A_after_B_ewc*100:.1f}%")
    print(f"  Forgetting mitigated: {(test_acc_A_after_A_ewc - test_acc_A_after_B_ewc)*100:.1f}% loss")
    print()

    # -------------------------------------------------------------------------
    # Comparison Summary
    # -------------------------------------------------------------------------
    print("=" * 70)
    print("COMPARISON SUMMARY")
    print("=" * 70)
    print()

    print("Task A Performance:")
    print(f"  Vanilla:  {test_acc_A_after_A*100:.1f}% -> {test_acc_A_after_B*100:.1f}%  (loss: {(test_acc_A_after_A - test_acc_A_after_B)*100:.1f}%)")
    print(f"  EWC:      {test_acc_A_after_A_ewc*100:.1f}% -> {test_acc_A_after_B_ewc*100:.1f}%  (loss: {(test_acc_A_after_A_ewc - test_acc_A_after_B_ewc)*100:.1f}%)")
    print()

    print("Task B Performance:")
    print(f"  Vanilla:  {test_acc_B_after_B*100:.1f}%")
    print(f"  EWC:      {test_acc_B_after_B_ewc*100:.1f}%")
    print()

    print("Trade-off Analysis:")
    forgetting_reduction = (test_acc_A_after_B_ewc - test_acc_A_after_B) / (test_acc_A_after_A - test_acc_A_after_B) * 100
    task_b_penalty = (test_acc_B_after_B - test_acc_B_after_B_ewc) * 100

    print(f"  EWC reduces forgetting by: {forgetting_reduction:.1f}%")
    print(f"  Cost: Task B accuracy reduced by: {task_b_penalty:.1f}%")
    print()

    print("Key Insight:")
    print("  EWC preserves Task A knowledge by constraining weight updates during")
    print("  Task B training. This creates a stability-plasticity trade-off:")
    print("  - More stability (higher lambda) = better Task A retention, worse Task B learning")
    print("  - More plasticity (lower lambda) = better Task B learning, more Task A forgetting")
    print()

    print("Mechanism Explanation:")
    print("  1. Vanilla learning: Task B training overwrites weights important for Task A")
    print("  2. EWC: Fisher information identifies important Task A weights")
    print("  3. EWC penalty: Penalizes large changes to important weights during Task B")
    print("  4. Result: Controlled trade-off between old and new task performance")
    print()


if __name__ == "__main__":
    demonstrate_catastrophic_forgetting()
