"""
Dynamic Neural Network

A self-growing neural network that can modify its architecture
based on salience signals. Implements:
    - Neuron addition (brain growth)
    - Weight pruning (synaptic elimination)  
    - Knowledge preservation during growth
"""

import numpy as np
from typing import Optional, Tuple, List, Dict, Any
from dataclasses import dataclass


@dataclass
class DynamicNetConfig:
    """Configuration for dynamic network."""
    input_dim: int = 2
    initial_hidden: int = 4
    output_dim: int = 1
    max_hidden: int = 100
    growth_rate: int = 2
    prune_threshold: float = 0.01


class DynamicNeuralNetwork:
    """
    Self-modifying neural network with growth and pruning.
    
    Key features:
        - Starts small (underpowered)
        - Grows when optimization stagnates
        - Prunes dead neurons
        - Preserves knowledge during modifications
    """
    
    def __init__(self, config: Optional[DynamicNetConfig] = None):
        self.config = config or DynamicNetConfig()
        
        self.input_dim = self.config.input_dim
        self.hidden_dim = self.config.initial_hidden
        self.output_dim = self.config.output_dim
        
        # Initialize weights
        self._init_weights()
        
        # Tracking
        self.neuron_salience = np.zeros(self.hidden_dim)
        self.growth_history: List[int] = []
        self.prune_history: List[int] = []
        
    def _init_weights(self):
        """Xavier initialization."""
        self.W1 = np.random.randn(self.input_dim, self.hidden_dim) * np.sqrt(2/self.input_dim)
        self.b1 = np.zeros((1, self.hidden_dim))
        self.W2 = np.random.randn(self.hidden_dim, self.output_dim) * np.sqrt(2/self.hidden_dim)
        self.b2 = np.zeros((1, self.output_dim))
    
    def forward(self, X: np.ndarray) -> np.ndarray:
        """Forward pass."""
        X = X.reshape(-1, self.input_dim) if X.ndim == 1 else X
        
        self.z1 = np.dot(X, self.W1) + self.b1
        self.a1 = np.tanh(self.z1)
        self.z2 = np.dot(self.a1, self.W2) + self.b2
        self.a2 = 1 / (1 + np.exp(-np.clip(self.z2, -500, 500)))  # Sigmoid with clipping
        
        return self.a2
    
    def train_step(self, X: np.ndarray, y: np.ndarray, lr: float = 0.01) -> float:
        """One training step with backprop."""
        X = X.reshape(-1, self.input_dim) if X.ndim == 1 else X
        y = y.reshape(-1, self.output_dim) if y.ndim == 1 else y
        m = X.shape[0]
        
        # Forward
        output = self.forward(X)
        
        # Loss (binary cross-entropy)
        eps = 1e-8
        loss = -np.mean(y * np.log(output + eps) + (1 - y) * np.log(1 - output + eps))
        
        # Backward
        d_z2 = output - y
        d_W2 = np.dot(self.a1.T, d_z2) / m
        d_b2 = np.mean(d_z2, axis=0, keepdims=True)
        
        d_a1 = np.dot(d_z2, self.W2.T)
        d_z1 = d_a1 * (1 - self.a1 ** 2)
        d_W1 = np.dot(X.T, d_z1) / m
        d_b1 = np.mean(d_z1, axis=0, keepdims=True)
        
        # Update weights
        self.W2 -= lr * d_W2
        self.b2 -= lr * d_b2
        self.W1 -= lr * d_W1
        self.b1 -= lr * d_b1
        
        # Track neuron salience (gradient flow × activation)
        batch_salience = np.mean(np.abs(self.a1 * d_a1), axis=0)
        self.neuron_salience = 0.9 * self.neuron_salience + 0.1 * batch_salience
        
        return float(loss)
    
    def grow_brain(self, n_neurons: Optional[int] = None) -> int:
        """
        Add neurons to hidden layer.
        
        New neurons are initialized with small random weights
        to allow gradual integration without disrupting existing knowledge.
        """
        if n_neurons is None:
            n_neurons = self.config.growth_rate
            
        if self.hidden_dim + n_neurons > self.config.max_hidden:
            n_neurons = self.config.max_hidden - self.hidden_dim
            if n_neurons <= 0:
                return self.hidden_dim
        
        # New weights for W1 (input -> new neurons)
        new_W1 = np.random.randn(self.input_dim, n_neurons) * 0.1
        new_b1 = np.zeros((1, n_neurons))
        
        # New weights for W2 (new neurons -> output)
        new_W2 = np.random.randn(n_neurons, self.output_dim) * 0.1
        
        # Concatenate
        self.W1 = np.hstack([self.W1, new_W1])
        self.b1 = np.hstack([self.b1, new_b1])
        self.W2 = np.vstack([self.W2, new_W2])
        
        # Expand salience tracking
        self.neuron_salience = np.concatenate([self.neuron_salience, np.zeros(n_neurons)])
        
        self.hidden_dim += n_neurons
        self.growth_history.append(self.hidden_dim)
        
        return self.hidden_dim
    
    def prune_neurons(self, threshold: Optional[float] = None) -> int:
        """
        Remove neurons with low salience.
        
        This implements synaptic elimination - removing
        connections that don't contribute to the solution.
        """
        if threshold is None:
            threshold = self.config.prune_threshold
            
        # Find neurons to keep
        keep_mask = self.neuron_salience > threshold
        
        # Always keep at least some neurons
        if np.sum(keep_mask) < self.config.initial_hidden:
            # Keep the top initial_hidden neurons
            indices = np.argsort(self.neuron_salience)[-self.config.initial_hidden:]
            keep_mask = np.zeros_like(keep_mask, dtype=bool)
            keep_mask[indices] = True
        
        n_pruned = self.hidden_dim - np.sum(keep_mask)
        
        if n_pruned > 0:
            # Prune weights
            self.W1 = self.W1[:, keep_mask]
            self.b1 = self.b1[:, keep_mask]
            self.W2 = self.W2[keep_mask, :]
            self.neuron_salience = self.neuron_salience[keep_mask]
            
            self.hidden_dim = np.sum(keep_mask)
            self.prune_history.append(n_pruned)
        
        return n_pruned
    
    def get_parameters(self) -> np.ndarray:
        """Flatten all parameters."""
        return np.concatenate([
            self.W1.flatten(), self.b1.flatten(),
            self.W2.flatten(), self.b2.flatten()
        ])
    
    def get_complexity(self) -> int:
        """Return current network complexity (parameter count)."""
        return self.W1.size + self.b1.size + self.W2.size + self.b2.size
    
    def get_effective_complexity(self) -> float:
        """Return effective complexity based on active neurons."""
        active_neurons = np.sum(self.neuron_salience > self.config.prune_threshold)
        return float(active_neurons * (self.input_dim + self.output_dim))


class SalienceGuidedNetwork(DynamicNeuralNetwork):
    """
    Neural network that uses salience signals to guide growth.
    
    Integrates with the AGI framework by:
        - Growing when novelty stagnates
        - Pruning when complexity exceeds budget
        - Preserving high-retention neurons
    """
    
    def __init__(self, config: Optional[DynamicNetConfig] = None):
        super().__init__(config)
        
        self.stagnation_counter = 0
        self.best_loss = float('inf')
        self.loss_history: List[float] = []
        
        # Salience component weights for growth decision
        self.novelty_threshold = 0.1
        self.stagnation_threshold = 20
        
    def train_step_with_salience(self, 
                                 X: np.ndarray, 
                                 y: np.ndarray, 
                                 lr: float = 0.01) -> Tuple[float, bool]:
        """
        Train step that returns whether growth should be triggered.
        """
        loss = self.train_step(X, y, lr)
        self.loss_history.append(loss)
        
        # Check for stagnation
        if loss < self.best_loss - 0.001:
            self.best_loss = loss
            self.stagnation_counter = 0
            should_grow = False
        else:
            self.stagnation_counter += 1
            should_grow = self.stagnation_counter >= self.stagnation_threshold
            
            if should_grow:
                self.grow_brain()
                self.stagnation_counter = 0
                self.best_loss = loss + 0.05  # Reset threshold
        
        return loss, should_grow
    
    def metabolic_cycle(self, prune: bool = True, grow: bool = False) -> Dict[str, int]:
        """
        Perform metabolic maintenance cycle.
        
        Like biological sleep - consolidate learning, prune waste, grow if needed.
        """
        results = {'pruned': 0, 'grown': 0, 'hidden_dim': self.hidden_dim}
        
        if prune:
            results['pruned'] = self.prune_neurons()
        
        if grow:
            old_dim = self.hidden_dim
            self.grow_brain()
            results['grown'] = self.hidden_dim - old_dim
        
        results['hidden_dim'] = self.hidden_dim
        return results
