import numpy as np
from scipy.io.wavfile import read

# Linear component
x = np.linspace(0, 1, 100)

# Load training and test 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)
X20 = X_test[:1000]
yt20 = y_test[:1000]


class PolynomialLayer:
    """A layer using polynomial fitting instead of linear transformation"""
    def __init__(self, input_size, output_size, degree=3):
        self.input_size = input_size
        self.output_size = output_size
        self.degree = degree
        # Store polynomial coefficients for each output neuron
        self.poly_coeffs = [
            np.zeros((degree + 1, input_size)) for _ in range(output_size)
        ]
    
    def fit(self, X, Y):
        """Fit polynomial coefficients to map X -> Y"""
        # X: (m, input_size), Y: (m, output_size)
        for out_idx in range(self.output_size):
            for in_idx in range(self.input_size):
                # Fit polynomial: input_feature -> output_feature
                try:
                    coeffs = np.polyfit(X[:, in_idx], Y[:, out_idx], self.degree)
                    self.poly_coeffs[out_idx][:, in_idx] = coeffs
                except:
                    self.poly_coeffs[out_idx][:, in_idx] = np.zeros(self.degree + 1)
    
    def forward(self, X):
        """Evaluate polynomials"""
        # X: (m, input_size)
        output = np.zeros((X.shape[0], self.output_size))
        for out_idx in range(self.output_size):
            for in_idx in range(self.input_size):
                # Evaluate polynomial at X[:, in_idx]
                poly_val = np.polyval(self.poly_coeffs[out_idx][:, in_idx], X[:, in_idx])
                output[:, out_idx] += poly_val
        return output


class PolyMLPClassifier:
    def __init__(self, layer_sizes, poly_degrees=None, learning_rate=0.01):
        """
        layer_sizes: [input, hidden1, ..., output]
        poly_degrees: polynomial degree for each layer (default: 3)
        """
        self.layer_sizes = layer_sizes
        self.num_layers = len(layer_sizes) - 1
        
        # Default polynomial degrees
        if poly_degrees is None:
            poly_degrees = [3] * self.num_layers
        
        self.poly_degrees = poly_degrees
        self.learning_rate = learning_rate
        
        # Create polynomial layers (linear layers)
        self.poly_layers = []
        for i in range(self.num_layers):
            # Alternating: linear layer -> polynomial layer
            if i % 2 == 0:
                self.poly_layers.append(
                    PolynomialLayer(layer_sizes[i], layer_sizes[i+1], poly_degrees[i])
                )
            else:
                self.poly_layers.append(
                    PolynomialLayer(layer_sizes[i], layer_sizes[i+1], poly_degrees[i])
                )
        
        # Also keep standard linear layers for comparison
        self.linear_weights = []
        self.linear_biases = []
        for i in range(self.num_layers):
            W = np.random.randn(layer_sizes[i], layer_sizes[i+1]) * 0.01
            b = np.zeros((1, layer_sizes[i+1]))
            self.linear_weights.append(W)
            self.linear_biases.append(b)
    
    def relu(self, Z):
        return np.maximum(0, Z)
    
    def relu_derivative(self, Z):
        return np.where(Z > 0, 1, 0)
    
    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):
        """Forward pass using polynomial evaluation"""
        self.activations = [X]
        A = X
        
        for i in range(self.num_layers):
            # Polynomial transformation
            poly_out = self.poly_layers[i].forward(A)
            
            # Add bias (learned)
            Z = poly_out + self.linear_biases[i]
            
            # Activation
            if i == self.num_layers - 1:
                A = self.softmax(Z)
            else:
                A = self.relu(Z)
            
            self.activations.append(A)
        
        return A
    
    def backward(self, y_true, y_pred):
        """Backward pass with polynomial gradient updates"""
        m = y_true.shape[0]
        self.grad_weights = []
        self.grad_biases = []
        
        # Output layer
        dZ = y_pred - y_true
        dW = np.dot(self.activations[-2].T, dZ) / m
        db = np.sum(dZ, axis=0, keepdims=True) / m
        dA = np.dot(dZ, self.linear_weights[-1].T)
        
        self.grad_weights.append(dW)
        self.grad_biases.append(db)
        
        # Hidden layers
        for i in reversed(range(1, self.num_layers)):
            dZ = dA * self.relu_derivative(self.activations[i])
            dW = np.dot(self.activations[i-1].T, dZ) / m
            db = np.sum(dZ, axis=0, keepdims=True) / m
            dA = np.dot(dZ, self.linear_weights[i-1].T)
            
            self.grad_weights.insert(0, dW)
            self.grad_biases.insert(0, db)
        
        return self.grad_weights, self.grad_biases
    
    def update_poly_layers(self, X, y_true, y_pred):
        """Update polynomial coefficients using polyfit"""
        # For each layer, refit polynomials to better approximate the target
        for layer_idx in range(self.num_layers - 1):
            A = self.activations[layer_idx]
            target = self.grad_weights[layer_idx + 1].T  # Use gradients as target
            self.poly_layers[layer_idx].fit(A, target)
    
    def update(self, X, y_true):
        """Update both polynomial and linear layers"""
        y_pred = self.forward(X)
        dW, db = self.backward(y_true, y_pred)
        
        # Update linear layers
        for i in range(self.num_layers):
            self.linear_weights[i] -= self.learning_rate * dW[i]
            self.linear_biases[i] -= self.learning_rate * db[i]
        
        # Update polynomial layers periodically
        if np.random.rand() < 0.1:  # 10% chance to refit
            self.update_poly_layers(X, y_true, y_pred)
    
    def predict(self, X):
        return np.argmax(self.forward(X), axis=1)
    
    def score(self, X, y_true):
        return np.mean(self.predict(X) == y_true)


# Alternative: Pure polynomial regression replacing entire MLP
class PurePolyRegressor:
    """Replace MLP entirely with polynomial regression"""
    def __init__(self, input_size, output_size, degree=5):
        self.degree = degree
        self.output_size = output_size
        # Coefficients for each output class
        self.coeffs = [np.zeros(degree + 1) for _ in range(output_size)]
    
    def fit(self, X, y):
        """Fit polynomial for each class (one-vs-all)"""
        X_flat = X[:, :min(X.shape[1], 100)]  # Use subset of features
        
        for class_idx in range(self.output_size):
            y_binary = (y == class_idx).astype(float)
            try:
                self.coeffs[class_idx] = np.polyfit(
                    X_flat.mean(axis=1), 
                    y_binary, 
                    self.degree
                )
            except:
                self.coeffs[class_idx] = np.zeros(self.degree + 1)
    
    def predict_proba(self, X):
        """Predict probabilities using polynomial evaluation"""
        X_flat = X[:, :min(X.shape[1], 100)]
        x_mean = X_flat.mean(axis=1)
        
        probs = np.zeros((X.shape[0], self.output_size))
        for class_idx in range(self.output_size):
            probs[:, class_idx] = np.polyval(self.coeffs[class_idx], x_mean)
        
        # Softmax normalization
        probs = probs - probs.max(axis=1, keepdims=True)
        exp_probs = np.exp(probs)
        return exp_probs / exp_probs.sum(axis=1, keepdims=True)
    
    def predict(self, X):
        return np.argmax(self.predict_proba(X), axis=1)
    
    def score(self, X, y_true):
        return np.mean(self.predict(X) == y_true)


# Hybrid: Polynomial layers + Standard backprop
class HybridPolyMLP:
    """MLP where linear layers can be replaced with polynomial layers"""
    def __init__(self, layer_dims, use_poly=[False, True, False, True]):
        self.layer_dims = layer_dims
        self.use_poly = use_poly
        
        self.weights = []
        self.biases = []
        self.poly_layers = []
        
        for i in range(len(layer_dims) - 1):
            W = np.random.randn(layer_dims[i], layer_dims[i+1]) * 0.01
            b = np.zeros((1, layer_dims[i+1]))
            self.weights.append(W)
            self.biases.append(b)
            
            if use_poly[i]:
                self.poly_layers.append(PolynomialLayer(layer_dims[i], layer_dims[i+1], degree=3))
            else:
                self.poly_layers.append(None)
    
    def forward(self, X):
        self.activations = [X]
        self.z_values = []
        A = X
        for i in range(len(self.layer_dims) - 1):
            if self.use_poly[i] and self.poly_layers[i]:
                # Polynomial transformation
                poly_out = self.poly_layers[i].forward(A)
                Z = poly_out + self.biases[i]
            else:
                # Standard linear transformation
                Z = np.dot(A, self.weights[i]) + self.biases[i]

            self.z_values.append(Z)
            if i == len(self.layer_dims) - 2:
                A = self.softmax(Z)
            else:
                A = self.relu(Z)
            self.activations.append(A)
        
        return A
    
    def relu(self, Z):
        return np.maximum(0, Z)
    
    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 predict(self, X):
        return np.argmax(self.forward(X), axis=1)
    
    def score(self, X, y_true):
        return np.mean(self.predict(X) == y_true)

    def backward(self, y_true, y_pred):
        m = y_true.shape[0]
        grad_weights = [None] * len(self.weights)
        grad_biases = [None] * len(self.biases)

        dz = y_pred - y_true
        for layer_idx in reversed(range(len(self.weights))):
            a_prev = self.activations[layer_idx]
            grad_weights[layer_idx] = np.dot(a_prev.T, dz) / m
            grad_biases[layer_idx] = np.sum(dz, axis=0, keepdims=True) / m

            if layer_idx > 0:
                dz = np.dot(dz, self.weights[layer_idx].T)
                dz = dz * (self.z_values[layer_idx - 1] > 0)

        return grad_weights, grad_biases

    def update(self, X, y_true, learning_rate=0.01):
        y_pred = self.forward(X)
        grad_weights, grad_biases = self.backward(y_true, y_pred)

        for layer_idx in range(len(self.weights)):
            self.weights[layer_idx] -= learning_rate * grad_weights[layer_idx]
            self.biases[layer_idx] -= learning_rate * grad_biases[layer_idx]

        return y_pred


# Initialize and train
layer_dims = [784, 100, 100, 100, 100, 10]
use_poly = [True, True, True, True, True]  # Use polynomial layers

f = HybridPolyMLP(layer_dims=layer_dims, use_poly=use_poly)

# Training loop
i = 0
while True:
    idx = np.random.randint(0, 60000, 100)
    X = X_train[idx]
    yt = y_train[idx]
    y_onehot = np.eye(10)[yt]
    
    f.update(X, y_onehot, learning_rate=0.01)
    
    print(i, f.score(X20, yt20))
    i += 1
