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

# Load training and test data
try:
    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)
except Exception:
    # Fallback to simulated data if files are not present
    X_train = np.random.randn(60000, 784)
    y_train = np.random.randint(0, 10, 60000)
    X_test = np.random.randn(1000, 784)
    y_test = np.random.randint(0, 10, 1000)

X20 = X_test[:1000]
yt20 = y_test[:1000]

class MLPClassifier:
    def __init__(self, input_size, hidden_size, output_size, learning_rate=None):
        if learning_rate is None:
            self.learning_rate = np.array([0.05, 0.05, 0.05, 0.05])
        else:
            self.learning_rate = learning_rate
            
        self.W1 = np.random.randn(input_size, hidden_size) * np.sqrt(2.0 / input_size)
        self.b1 = np.zeros((1, hidden_size))
        self.W2 = np.random.randn(hidden_size, output_size) * np.sqrt(2.0 / hidden_size)
        self.b2 = np.zeros((1, output_size))
        
        # --- NOVEL ADDITION: DEDICATED ELEMENT PROBABILITY MATRICES ---
        # These retain the exact dimensions of the target layers, maintaining 1:1 mapping
        self.P_W1 = np.ones_like(self.W1) * 0.5
        self.P_W2 = np.ones_like(self.W2) * 0.5
        
        self.w2_esd_history = []
        self.target_esd_w1 = 10.0
        self.target_esd_w2 = 6.0
        
    def relu(self, x):
        return np.maximum(0, x)
    
    def relu_derivative(self, x):
        return np.where(x > 0, 1, 0)
    
    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) + 1e-15)
    
    def forward(self, X):
        self.z1 = np.dot(X, self.W1) + self.b1
        self.a1 = self.relu(self.z1)
        self.z2 = np.dot(self.a1, self.W2) + self.b2
        output = self.softmax(self.z2)
        return output
    
    def forward_with_weights(self, X, W1, b1, W2, b2):
        z1 = np.dot(X, W1) + b1
        a1 = self.relu(z1)
        z2 = np.dot(a1, W2) + b2
        return a1, self.softmax(z2)

    def compute_w2_derivative(self, X, y_true, current_W2, current_b2):
        m = y_true.shape[0]
        a1, y_pred = self.forward_with_weights(X, self.W1, self.b1, current_W2, current_b2)
        dz2 = y_pred - y_true
        dW2 = np.dot(a1.T, dz2) / m
        db2 = np.sum(dz2, axis=0, keepdims=True) / m
        return dW2, db2

    def update(self, X, y_true):
        # --- LAYER 1: BACKPROPAGATION ---
        y_pred = self.forward(X)
        m = y_true.shape[0]
        dz2 = y_pred - y_true
        da1 = np.dot(dz2, self.W2.T)
        dz1 = da1 * self.relu_derivative(self.z1)
        dW1 = np.dot(X.T, dz1) / m
        db1 = np.sum(dz1, axis=0, keepdims=True) / m
        
        # Calculate global layer metrics
        current_w1_esd = self.calculate_esd(self.W1)
        error_w1 = self.target_esd_w1 - current_w1_esd
        dampening_factor_w1 = np.exp(-0.4 * error_w1)
        
        # Calculate dynamic updates for W1 element probabilities
        # Elements with large absolute updates relative to the mean get lower probability values (higher stability)
        mean_abs_dW1 = np.mean(np.abs(dW1)) + 1e-15
        self.P_W1 = 1.0 / (1.0 + np.abs(dW1) / mean_abs_dW1)
        
        # Apply element-wise updates using the probability matrix mapping
        step1 = (self.learning_rate[0] * dampening_factor_w1 * self.P_W1) / (1.0 + np.log1p(np.abs(dW1)))
        step2 = (self.learning_rate[1] * dampening_factor_w1) / (1.0 + np.log1p(np.linalg.norm(db1)))
        
        self.W1 -= step1 * dW1
        self.b1 -= step2 * db1

        # --- LAYER 2: RK4 ANTI-ALIASED INTEGRATION WITH ELEMENT PROBABILITY ---
        current_w2_esd = self.calculate_esd(self.W2)
        self.w2_esd_history.append(current_w2_esd)
        if len(self.w2_esd_history) > 10:
            self.w2_esd_history.pop(0)
            
        error_w2 = self.target_esd_w2 - current_w2_esd
        oscillation_amplitude = np.abs(self.w2_esd_history[-1] - self.w2_esd_history[-2]) if len(self.w2_esd_history) > 1 else 0.0
        dampening_factor_w2 = np.exp(-0.5 * error_w2 - 0.2 * oscillation_amplitude)
        
        # Calculate baseline RK4 updates
        dt_w2 = self.learning_rate[2] * dampening_factor_w2
        dt_b2 = self.learning_rate[3] * dampening_factor_w2
        
        dW2_k1, db2_k1 = self.compute_w2_derivative(X, y_true, self.W2, self.b2)
        dW2_k2, db2_k2 = self.compute_w2_derivative(X, y_true, self.W2 - 0.5 * dt_w2 * dW2_k1, self.b2 - 0.5 * dt_b2 * db2_k1)
        dW2_k3, db2_k3 = self.compute_w2_derivative(X, y_true, self.W2 - 0.5 * dt_w2 * dW2_k2, self.b2 - 0.5 * dt_b2 * db2_k2)
        dW2_k4, db2_k4 = self.compute_w2_derivative(X, y_true, self.W2 - dt_w2 * dW2_k3, self.b2 - dt_b2 * db2_k3)
        
        blended_dW2 = (dW2_k1 + 2.0 * dW2_k2 + 2.0 * dW2_k3 + dW2_k4) / 6.0
        blended_db2 = (db2_k1 + 2.0 * db2_k2 + 2.0 * db2_k3 + db2_k4) / 6.0
        
        # Update W2 element probabilities based on the integrated RK4 flow vectors
        mean_abs_dW2 = np.mean(np.abs(blended_dW2)) + 1e-15
        self.P_W2 = 1.0 / (1.0 + np.abs(blended_dW2) / mean_abs_dW2)
        
        # Element-wise step integration
        self.W2 -= (dt_w2 * self.P_W2) * blended_dW2
        self.b2 -= dt_b2 * blended_db2

    def predict(self, X):
        probabilities = self.forward(X)
        return np.argmax(probabilities, axis=1)

    def score(self, X, y_true):
        y = self.predict(X)
        return np.mean(y == y_true)

    def calculate_esd(self, matrix):
        abs_mat = np.abs(matrix)
        max_val = np.max(abs_mat)
        min_val = np.min(abs_mat[abs_mat > 0]) if np.any(abs_mat > 0) else 1e-15
        return np.log10(max_val / min_val)

if __name__ == "__main__":
    learning_rate = np.array([0.05, 0.05, 0.05, 0.05])
    f = MLPClassifier(input_size=784, hidden_size=128, output_size=10, learning_rate=learning_rate)

    i = 0
    print("Initiating Distributed Element Probability Matrix Engine...")
    while i <= 300:
        idx = np.random.randint(0, X_train.shape[0], 128)
        X = X_train[idx]
        yt = y_train[idx]
        
        if i % 50 == 0:
            train_acc = f.score(X, yt)
            test_acc = f.score(X20, yt20)
            esd_w1 = f.calculate_esd(f.W1)
            esd_w2 = f.calculate_esd(f.W2)
            print(f"Iteration: {i:03d} | Train Acc: {train_acc:.4f} | Test Acc: {test_acc:.4f} | W1 ESD: {esd_w1:.2f} | W2 ESD: {esd_w2:.2f}")
            
        f.update(X, np.eye(10)[yt])
        i += 1