"""
Visualization Utilities

Plotting and monitoring for the AGI system.
"""

import numpy as np
import matplotlib.pyplot as plt
from typing import List, Dict, Any, Optional
import os


class AGIVisualizer:
    """
    Comprehensive visualization for AGI training.
    """
    
    def __init__(self, save_dir: str = "plots"):
        self.save_dir = save_dir
        os.makedirs(save_dir, exist_ok=True)
        
        # Color scheme
        self.colors = {
            's_prime': '#2196F3',      # Blue
            's_double_prime': '#4CAF50', # Green
            's_triple_prime': '#FF9800', # Orange
            'loss': '#F44336',          # Red
            'complexity': '#9C27B0',    # Purple
            'growth': '#E91E63'         # Pink
        }
    
    def plot_training_progress(self, 
                               losses: List[float],
                               saliences: List[float],
                               title: str = "Training Progress",
                               save_name: Optional[str] = None):
        """Plot loss and salience curves."""
        fig, (ax1, ax2) = plt.subplots(2, 1, figsize=(12, 8), sharex=True)
        
        ax1.plot(losses, color=self.colors['loss'], alpha=0.7, label='Loss')
        ax1.set_ylabel('Loss')
        ax1.set_title(title)
        ax1.legend()
        ax1.grid(True, alpha=0.3)
        ax1.set_yscale('log')
        
        ax2.plot(saliences, color=self.colors['s_prime'], alpha=0.7, label='Salience')
        ax2.set_xlabel('Epoch')
        ax2.set_ylabel('Salience')
        ax2.legend()
        ax2.grid(True, alpha=0.3)
        
        plt.tight_layout()
        
        if save_name:
            plt.savefig(os.path.join(self.save_dir, save_name), dpi=150)
        
        return fig
    
    def plot_three_layer_dynamics(self,
                                  s_prime: List[float],
                                  s_double_prime: List[float],
                                  s_triple_prime: List[float],
                                  save_name: Optional[str] = None):
        """Plot all three salience layers."""
        fig, axes = plt.subplots(3, 1, figsize=(14, 10), sharex=True)
        
        # S'
        axes[0].plot(s_prime, color=self.colors['s_prime'], linewidth=1.5)
        axes[0].set_ylabel("S' (Trajectory)")
        axes[0].set_title("Three-Layer Recursive Optimization Dynamics")
        axes[0].grid(True, alpha=0.3)
        axes[0].fill_between(range(len(s_prime)), s_prime, alpha=0.3, color=self.colors['s_prime'])
        
        # S''
        axes[1].plot(s_double_prime, color=self.colors['s_double_prime'], linewidth=1.5)
        axes[1].set_ylabel("S'' (Meta)")
        axes[1].grid(True, alpha=0.3)
        axes[1].fill_between(range(len(s_double_prime)), s_double_prime, alpha=0.3, color=self.colors['s_double_prime'])
        
        # S'''
        axes[2].plot(s_triple_prime, color=self.colors['s_triple_prime'], linewidth=1.5)
        axes[2].set_ylabel("S''' (Self)")
        axes[2].set_xlabel('Step')
        axes[2].grid(True, alpha=0.3)
        axes[2].fill_between(range(len(s_triple_prime)), s_triple_prime, alpha=0.3, color=self.colors['s_triple_prime'])
        
        plt.tight_layout()
        
        if save_name:
            plt.savefig(os.path.join(self.save_dir, save_name), dpi=150)
        
        return fig
    
    def plot_decision_boundary(self,
                               network,
                               X: np.ndarray,
                               y: np.ndarray,
                               title: str = "Decision Boundary",
                               save_name: Optional[str] = None):
        """Plot decision boundary for 2D classification."""
        fig, ax = plt.subplots(1, 1, figsize=(10, 8))
        
        # Create mesh
        x_min, x_max = X[:, 0].min() - 0.5, X[:, 0].max() + 0.5
        y_min, y_max = X[:, 1].min() - 0.5, X[:, 1].max() + 0.5
        xx, yy = np.meshgrid(np.arange(x_min, x_max, 0.02),
                            np.arange(y_min, y_max, 0.02))
        
        # Predict on mesh
        Z = network.forward(np.c_[xx.ravel(), yy.ravel()])
        Z = Z.reshape(xx.shape)
        
        # Plot decision boundary
        ax.contourf(xx, yy, Z, cmap=plt.cm.RdBu, alpha=0.8, levels=20)
        ax.contour(xx, yy, Z, levels=[0.5], colors='black', linewidths=2)
        
        # Plot data points
        scatter = ax.scatter(X[:, 0], X[:, 1], c=y.flatten(), 
                            edgecolors='black', cmap=plt.cm.RdBu, s=50)
        
        ax.set_title(title)
        ax.set_xlabel('Feature 1')
        ax.set_ylabel('Feature 2')
        
        plt.colorbar(scatter, ax=ax, label='Class')
        plt.tight_layout()
        
        if save_name:
            plt.savefig(os.path.join(self.save_dir, save_name), dpi=150)
        
        return fig
    
    def plot_network_growth(self,
                           neuron_counts: List[int],
                           growth_events: List[int],
                           save_name: Optional[str] = None):
        """Plot network architecture evolution."""
        fig, ax = plt.subplots(1, 1, figsize=(12, 6))
        
        ax.plot(neuron_counts, color=self.colors['complexity'], linewidth=2)
        
        # Mark growth events
        for event in growth_events:
            if event < len(neuron_counts):
                ax.axvline(x=event, color=self.colors['growth'], linestyle='--', alpha=0.7)
        
        ax.set_xlabel('Epoch')
        ax.set_ylabel('Hidden Neurons')
        ax.set_title('Network Architecture Evolution')
        ax.grid(True, alpha=0.3)
        
        if save_name:
            plt.savefig(os.path.join(self.save_dir, save_name), dpi=150)
        
        return fig
    
    def plot_sparse_topology(self,
                            connection_history: List[int],
                            sparsity_history: List[float],
                            neuron_salience: np.ndarray,
                            save_name: Optional[str] = None):
        """Plot sparse network topology evolution."""
        fig, axes = plt.subplots(2, 2, figsize=(14, 10))
        
        # Connections over time
        axes[0, 0].plot(connection_history, color='green')
        axes[0, 0].set_title('Active Connections (Homeostasis)')
        axes[0, 0].set_xlabel('Metabolic Cycle')
        axes[0, 0].set_ylabel('Connections')
        axes[0, 0].grid(True, alpha=0.3)
        
        # Sparsity over time
        axes[0, 1].plot(sparsity_history, color='orange')
        axes[0, 1].set_title('Network Sparsity')
        axes[0, 1].set_xlabel('Metabolic Cycle')
        axes[0, 1].set_ylabel('Sparsity')
        axes[0, 1].grid(True, alpha=0.3)
        
        # Neuron salience distribution
        axes[1, 0].bar(range(len(neuron_salience)), neuron_salience, color='purple', width=1.0)
        axes[1, 0].set_title('Neuron Salience (Specialization)')
        axes[1, 0].set_xlabel('Neuron Index')
        axes[1, 0].set_ylabel('Salience')
        
        # Salience histogram
        axes[1, 1].hist(neuron_salience, bins=50, color='purple', alpha=0.7)
        axes[1, 1].set_title('Salience Distribution')
        axes[1, 1].set_xlabel('Salience')
        axes[1, 1].set_ylabel('Count')
        axes[1, 1].axvline(x=np.mean(neuron_salience), color='red', linestyle='--', label='Mean')
        axes[1, 1].legend()
        
        plt.tight_layout()
        
        if save_name:
            plt.savefig(os.path.join(self.save_dir, save_name), dpi=150)
        
        return fig
    
    def plot_salience_components(self,
                                 components_history: Dict[str, List[float]],
                                 save_name: Optional[str] = None):
        """Plot individual salience components over time."""
        fig, axes = plt.subplots(2, 3, figsize=(15, 8))
        
        component_names = ['novelty', 'retention', 'meaning', 'continuity', 'fatigue', 'jerk']
        colors = ['#2196F3', '#4CAF50', '#FF9800', '#9C27B0', '#F44336', '#607D8B']
        
        for idx, (name, color) in enumerate(zip(component_names, colors)):
            ax = axes[idx // 3, idx % 3]
            if name in components_history:
                ax.plot(components_history[name], color=color, linewidth=1.5)
            ax.set_title(f'{name.capitalize()} ({"+" if name not in ["fatigue", "jerk"] else "-"})')
            ax.grid(True, alpha=0.3)
            ax.set_xlabel('Step')
        
        plt.suptitle('Salience Component Dynamics', fontsize=14)
        plt.tight_layout()
        
        if save_name:
            plt.savefig(os.path.join(self.save_dir, save_name), dpi=150)
        
        return fig
    
    def plot_comprehensive_dashboard(self,
                                     losses: List[float],
                                     network,
                                     X: np.ndarray,
                                     y: np.ndarray,
                                     neuron_counts: List[int],
                                     lr_history: List[float],
                                     save_name: Optional[str] = None):
        """Create comprehensive training dashboard."""
        fig = plt.figure(figsize=(16, 12))
        
        # Loss curve
        ax1 = fig.add_subplot(2, 2, 1)
        ax1.plot(losses, color='blue')
        ax1.set_title('Learning Curve')
        ax1.set_xlabel('Epoch')
        ax1.set_ylabel('Loss')
        ax1.set_yscale('log')
        ax1.grid(True, alpha=0.3)
        
        # Decision boundary
        ax2 = fig.add_subplot(2, 2, 2)
        x_min, x_max = X[:, 0].min() - 0.5, X[:, 0].max() + 0.5
        y_min, y_max = X[:, 1].min() - 0.5, X[:, 1].max() + 0.5
        xx, yy = np.meshgrid(np.arange(x_min, x_max, 0.05),
                            np.arange(y_min, y_max, 0.05))
        Z = network.forward(np.c_[xx.ravel(), yy.ravel()])
        Z = Z.reshape(xx.shape)
        ax2.contourf(xx, yy, Z, cmap=plt.cm.RdBu, alpha=0.8)
        ax2.scatter(X[:, 0], X[:, 1], c=y.flatten(), edgecolors='k', cmap=plt.cm.RdBu)
        ax2.set_title(f'Decision Boundary (Loss: {losses[-1]:.4f})')
        
        # Network growth
        ax3 = fig.add_subplot(2, 2, 3)
        ax3.plot(neuron_counts, color='purple')
        ax3.set_title('Brain Capacity')
        ax3.set_xlabel('Epoch')
        ax3.set_ylabel('Hidden Neurons')
        ax3.grid(True, alpha=0.3)
        
        # Learning rate
        ax4 = fig.add_subplot(2, 2, 4)
        ax4.plot(lr_history, color='green')
        ax4.set_title('Self-Regulated Learning Rate')
        ax4.set_xlabel('Epoch')
        ax4.set_ylabel('Learning Rate')
        ax4.grid(True, alpha=0.3)
        
        plt.tight_layout()
        
        if save_name:
            plt.savefig(os.path.join(self.save_dir, save_name), dpi=150)
        
        return fig
