import numpy as np

# Linear component shared across all linear layers
x = np.linspace(0, 1, 100)

class DeepMLP:
    """
    10-layer MLP with linear components strategically placed:
    - Layer 1: Input projection (feature extraction)
    - Layer 4: Mid-level skip connection
    - Layer 7: High-level skip connection
    - Layer 10: Pre-output integration
    """
    
    def __init__(self, input_size=784, layer_sizes=[256, 128, 128, 256, 128, 128, 256, 128, 64, 10], 
                 linear_indices=[0, 3, 6, 9], linear_dim=100, learning_rate=None):
        
        self.layer_sizes = [input_size] + layer_sizes
        self.linear_indices = linear_indices
        self.linear_dim = linear_dim
        self.num_layers = len(layer_sizes)
        self.x = x
        
        # Initialize all weights and biases
        self.weights = []
        self.biases = []
        self.linear_weights = []
        self.linear_biases = []
        
        for i in range(self.num_layers):
            # Standard weights
            W = np.random.randn(self.layer_sizes[i], self.layer_sizes[i+1]) * np.sqrt(2.0 / self.layer_sizes[i])
            b = np.zeros((1, self.layer_sizes[i+1]))
            self.weights.append(W)
            self.biases.append(b)
            
            # Linear component weights (only for linear layers)
            if i in linear_indices:
                # Project to linear_dim, then back to layer size
                lw = np.random.randn(self.layer_sizes[i], linear_dim) * 0.01
                lb = np.zeros((1, linear_dim))
                self.linear_weights.append(lw)
                self.linear_biases.append(lb)
            else:
                self.linear_weights.append(None)
                self.linear_biases.append(None)
        
        # Learning rates for all parameters
        if learning_rate is None:
            self.learning_rate = 0.001
        else:
            self.learning_rate = learning_rate
        
        # For batch normalization (optional, very effective)
        self.gamma = [np.ones(s) for s in self.layer_sizes[1:]]
        self.beta = [np.zeros(s) for s in self.layer_sizes[1:]]
        self.running_mean = [np.zeros(s) for s in self.layer_sizes[1:]]
        self.running_var = [np.ones(s) for s in self.layer_sizes[1:]]
        self.bn_eps = 1e-5
        
    def linear_transform(self, X, layer_idx):
        """Apply linear component transformation"""
        if self.linear_weights[layer_idx] is None:
            return None
        return np.dot(X, self.linear_weights[layer_idx]) + self.linear_biases[layer_idx]
    
    def relu(self, z):
        return np.maximum(0, z)
    
    def relu_derivative(self, z):
        return np.where(z > 0, 1, 0)
    
    def leaky_relu(self, z, alpha=0.01):
        return np.where(z > 0, z, alpha * z)
    
    def leaky_relu_derivative(self, z, alpha=0.01):
        return np.where(z > 0, 1, alpha)
    
    def batch_norm(self, x, layer_idx, training=True):
        if training:
            mean = np.mean(x, axis=0)
            var = np.var(x, axis=0)
            self.running_mean[layer_idx] = 0.9 * self.running_mean[layer_idx] + 0.1 * mean
            self.running_var[layer_idx] = 0.9 * self.running_var[layer_idx] + 0.1 * var
        else:
            mean = self.running_mean[layer_idx]
            var = self.running_var[layer_idx]
        
        x_norm = (x - mean) / np.sqrt(var + self.bn_eps)
        return self.gamma[layer_idx] * x_norm + self.beta[layer_idx]
    
    def softmax(self, z):
        exp_z = np.exp(z - np.max(z, axis=1, keepdims=True))
        return exp_z / np.sum(exp_z, axis=1, keepdims=True)
    
    def forward(self, X, training=True):
        self.activations = [X]
        self.z_values = []
        self.linear_outputs = []
        self.bn_outputs = []
        
        current = X
        
        for i in range(self.num_layers):
            # Linear component (when applicable)
            linear_out = self.linear_transform(current, i)
            self.linear_outputs.append(linear_out)
            
            # Main linear transformation
            z = np.dot(current, self.weights[i]) + self.biases[i]
            self.z_values.append(z)
            
            # Batch normalization
            z = self.batch_norm(z, i, training)
            self.bn_outputs.append(z)
            
            # Activation (skip for last layer - use softmax)
            if i < self.num_layers - 1:
                # Leaky ReLU for better gradient flow
                current = self.leaky_relu(z)
            else:
                current = z
            
            self.activations.append(current)
        
        # Softmax output
        output = self.softmax(current)
        return output
    
    def compute_loss(self, y_true, y_pred):
        m = y_true.shape[0]
        return -np.sum(y_true * np.log(y_pred + 1e-9)) / m
    
    def backward(self, X, y_true, y_pred):
        m = y_true.shape[0]
        self.grad_weights = [None] * self.num_layers
        self.grad_biases = [None] * self.num_layers
        self.grad_linear_weights = [None] * self.num_layers
        self.grad_linear_biases = [None] * self.num_layers
        self.grad_gamma = [None] * self.num_layers
        self.grad_beta = [None] * self.num_layers
        
        # Output gradient
        dz = y_pred - y_true  # (m, 10)
        
        for i in reversed(range(self.num_layers)):
            a_prev = self.activations[i]  # (m, layer_size[i])
            
            # Gradient w.r.t. weights and biases
            self.grad_weights[i] = np.dot(a_prev.T, dz) / m
            self.grad_biases[i] = np.sum(dz, axis=0, keepdims=True) / m
            
            # Gradient w.r.t. input (for backprop)
            da_prev = np.dot(dz, self.weights[i].T)
            
            if i > 0:
                # Backprop through activation
                if i - 1 in self.linear_indices:
                    # Leaky ReLU derivative
                    dz = da_prev * self.leaky_relu_derivative(self.z_values[i-1])
                else:
                    dz = da_prev * self.leaky_relu_derivative(self.z_values[i-1])
                
                # Batch norm backward
                dz = self.batch_norm_backward(dz, i-1)
                
                # Linear component gradient (if applicable)
                if self.linear_weights[i-1] is not None:
                    linear_out = self.linear_outputs[i-1]
                    self.grad_linear_weights[i-1] = np.dot(a_prev[:, :-1].T if i-1 == 0 else a_prev.T, 
                                                          np.dot(dz, self.linear_weights[i-1].T) / m) if linear_out is not None else None
                    # Simplified: linear gradient
                    self.grad_linear_weights[i-1] = np.dot(a_prev.T, dz) / m
                    self.grad_linear_biases[i-1] = np.sum(dz, axis=0, keepdims=True) / m
    
    def batch_norm_backward(self, dz, layer_idx):
        """Simplified batch norm backward"""
        m = dz.shape[0]
        d_x_norm = dz * self.gamma[layer_idx]
        d_var = np.sum(d_x_norm * (self.activations[layer_idx + 1] - self.running_mean[layer_idx]), axis=0) * -0.5 * (self.running_var[layer_idx] + self.bn_eps) ** -1.5
        d_mean = np.sum(d_x_norm, axis=0) * -1 / np.sqrt(self.running_var[layer_idx] + self.bn_eps) + d_var * np.mean(-2 * (self.activations[layer_idx] - self.running_mean[layer_idx]), axis=0)
        dz = d_x_norm / np.sqrt(self.running_var[layer_idx] + self.bn_eps) + d_var * 2 * (self.activations[layer_idx] - self.running_mean[layer_idx]) / m + d_mean / m
        return dz
    
    def update(self, X, y_true, grad_clip=1.0):
        y_pred = self.forward(X, training=True)
        self.backward(X, y_true, y_pred)
        
        lr = self.learning_rate
        
        for i in range(self.num_layers):
            # Clip gradients
            if self.grad_weights[i] is not None:
                grad = np.clip(self.grad_weights[i], -grad_clip, grad_clip)
                self.weights[i] -= lr * grad
            
            if self.grad_biases[i] is not None:
                grad = np.clip(self.grad_biases[i], -grad_clip, grad_clip)
                self.biases[i] -= lr * grad
            
            if self.grad_linear_weights[i] is not None:
                grad = np.clip(self.grad_linear_weights[i], -grad_clip, grad_clip)
                self.linear_weights[i] -= lr * grad
            
            if self.grad_linear_biases[i] is not None:
                grad = np.clip(self.grad_linear_biases[i], -grad_clip, grad_clip)
                self.linear_biases[i] -= lr * grad
    
    def predict(self, X):
        probs = self.forward(X, training=False)
        return np.argmax(probs, axis=1)
    
    def score(self, X, y_true):
        return np.mean(self.predict(X) == y_true)


# Improved version with skip connections and better gradient flow
class DeepMLPWithSkipConnections:
    """
    10-layer MLP with skip connections through linear components:
    
    Architecture:
    Block 1: Linear(784->256) -> BN -> ReLU -> Linear(256->128) -> BN -> ReLU
    Block 2: Linear(128->256) + skip -> BN -> ReLU -> Linear(256->128) -> BN -> ReLU
    Block 3: Linear(128->256) + skip -> BN -> ReLU -> Linear(256->64) -> BN -> ReLU
    Output: Linear(64->10)
    """
    
    def __init__(self, input_size=784, output_size=10, linear_dim=128, learning_rate=0.001):
        self.lr = learning_rate
        self.linear_dim = linear_dim
        
        # Layer dimensions: [input, h1, h2, h3, h4, h5, h6, h7, h8, output]
        dims = [784, 256, 128, 256, 128, 256, 128, 64, 32, 10]
        
        self.params = {}
        idx = 0
        
        # Block 1: Input projection with linear component
        self.params['W1'] = np.random.randn(784, 256) * np.sqrt(2/784)
        self.params['b1'] = np.zeros(256)
        self.params['linear_W1'] = np.random.randn(784, linear_dim) * 0.01
        self.params['linear_b1'] = np.zeros(linear_dim)
        
        # Block 2-10
        for i in range(1, 9):
            self.params[f'W{i+1}'] = np.random.randn(dims[i], dims[i+1]) * np.sqrt(2/dims[i])
            self.params[f'b{i+1}'] = np.zeros(dims[i+1])
            
            # Linear components at strategic positions (layers 1, 4, 7)
            if i in [1, 4, 7]:
                self.params[f'linear_W{i+1}'] = np.random.randn(dims[i], linear_dim) * 0.01
                self.params[f'linear_b{i+1}'] = np.zeros(linear_dim)
        
        # Batch norm parameters
        for i in range(9):
            self.params[f'gamma{i+1}'] = np.ones(dims[i+1])
            self.params[f'beta{i+1}'] = np.zeros(dims[i+1])
            self.params[f'running_mean{i+1}'] = np.zeros(dims[i+1])
            self.params[f'running_var{i+1}'] = np.ones(dims[i+1])
    
    def linear_layer(self, x, name):
        if f'linear_W{name}' in self.params:
            return np.dot(x, self.params[f'linear_W{name}']) + self.params[f'linear_b{name}']
        return None
    
    def batch_norm(self, x, idx, training=True):
        gamma = self.params[f'gamma{idx}']
        beta = self.params[f'beta{idx}']
        
        if training:
            mean = np.mean(x, axis=0)
            var = np.var(x, axis=0)
            self.params[f'running_mean{idx}'] = 0.9 * self.params[f'running_mean{idx}'] + 0.1 * mean
            self.params[f'running_var{idx}'] = 0.9 * self.params[f'running_var{idx}'] + 0.1 * var
        else:
            mean = self.params[f'running_mean{idx}']
            var = self.params[f'running_var{idx}']
        
        return gamma * (x - mean) / np.sqrt(var + 1e-5) + beta
    
    def relu(self, x):
        return np.maximum(0, x)
    
    def softmax(self, x):
        exp_x = np.exp(x - np.max(x, axis=1, keepdims=True))
        return exp_x / np.sum(exp_x, axis=1, keepdims=True)

    def match_residual(self, residual, target_width):
        """Project a cached activation to the target width without broadcasting."""
        if residual is None:
            return None
        current_width = residual.shape[1]
        if current_width == target_width:
            return residual
        if current_width > target_width:
            return residual[:, :target_width]

        padded = np.zeros((residual.shape[0], target_width), dtype=residual.dtype)
        padded[:, :current_width] = residual
        return padded
    
    def forward(self, X, training=True):
        self.cache = {'X': X}
        h = X
        
        # Layer 1: Linear projection + skip
        z1 = np.dot(h, self.params['W1']) + self.params['b1']
        linear_out1 = self.linear_layer(h, 1)
        bn1 = self.batch_norm(z1, 1, training)
        h = self.relu(bn1)
        self.cache['z1'] = z1
        self.cache['bn1'] = bn1
        self.cache['h1'] = h
        self.cache['linear1'] = linear_out1
        
        # Layers 2-10: blocks with skip connections through linear components
        for i in range(2, 10):
            z = np.dot(h, self.params[f'W{i}']) + self.params[f'b{i}']
            linear_out = self.linear_layer(h, i)
            
            # Add shape-compatible residual connections from earlier blocks.
            residual_sources = {
                4: self.cache.get('h2'),
                7: self.cache.get('h6'),
            }
            residual = self.match_residual(residual_sources.get(i), z.shape[1])
            if residual is not None:
                z = z + residual
            
            bn = self.batch_norm(z, i, training)
            h = self.relu(bn)
            self.cache[f'z{i}'] = z
            self.cache[f'bn{i}'] = bn
            self.cache[f'h{i}'] = h
            self.cache[f'linear{i}'] = linear_out
        
        return self.softmax(h)
    
    def backward(self, y_true, y_pred):
        m = y_true.shape[0]
        self.grads = {}
        
        # Output layer gradient for softmax-cross-entropy.
        dz = y_pred - y_true
        
        for i in reversed(range(1, 10)):
            h_prev = self.cache.get(f'h{i-1}', self.cache['X'])
            self.grads[f'W{i}'] = np.dot(h_prev.T, dz) / m
            self.grads[f'b{i}'] = np.sum(dz, axis=0) / m
            
            if i > 1:
                da_prev = np.dot(dz, self.params[f'W{i}'].T)
                dz = da_prev * (self.cache[f'bn{i-1}'] > 0)
    
    def update(self, X, y_true):
        y_pred = self.forward(X, training=True)
        self.backward(y_true, y_pred)
        
        for key in self.grads:
            self.params[key] -= self.lr * self.grads[key]
    
    def predict(self, X):
        return np.argmax(self.forward(X, training=False), axis=1)
    
    def score(self, X, y):
        return np.mean(self.predict(X) == y)


# Usage example
if __name__ == "__main__":
    from scipy.io.wavfile import read
    
    # Load data
    X_train = read('../X_train.wav')[1].reshape(-1, 784)
    y_train = (read('../y_train.wav')[1] * 9).astype(int)
    X_test = read('../X_test.wav')[1].reshape(-1, 784)
    y_test = (read('../y_test.wav')[1] * 9).astype(int)
    
    # Initialize model
    model = DeepMLPWithSkipConnections(input_size=784, output_size=10, 
                                        linear_dim=128, learning_rate=0.001)
    
    # Training loop
    for epoch in range(100):
        idx = np.random.randint(0, 60000, 128)
        X_batch = X_train[idx]
        y_batch = y_train[idx]
        y_onehot = np.eye(10)[y_batch]
        
        model.update(X_batch, y_onehot)
        
        if epoch % 10 == 0:
            acc = model.score(X_test[:1000], y_test[:1000])
            print(f"Epoch {epoch}, Accuracy: {acc:.4f}")
