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


# 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 MLPClassifier:
    def __init__(self, input_size, hidden_size, output_size, learning_rate):
        self.W1 = np.random.randn(input_size, hidden_size) * 1e-3
        self.b1 = np.zeros((1, hidden_size))
        self.W2 = np.random.randn(hidden_size, output_size) * 1e-2
        self.b2 = np.zeros((1, output_size))
        self.learning_rate = np.asarray(learning_rate, dtype=float)

    def leaky_relu(self, x, slope=1e-3):
        return np.where(x > 0, x, slope * x)

    def leaky_relu_derivative(self, x, slope=1e-3):
        return np.where(x > 0, 1.0, slope)

    def softmax(self, x):
        shifted = x - np.max(x, axis=1, keepdims=True)
        exp_x = np.exp(shifted)
        return exp_x / np.sum(exp_x, axis=1, keepdims=True)

    def forward(self, X):
        self.z1 = X @ self.W1 + self.b1
        self.a1 = self.leaky_relu(self.z1)
        self.z2 = self.a1 @ self.W2 + self.b2
        return self.softmax(self.z2)

    def backward(self, X, output_signal):
        m = X.shape[0]
        dz2 = output_signal
        dW2 = self.a1.T @ dz2 / m
        db2 = np.sum(dz2, axis=0, keepdims=True) / m
        da1 = dz2 @ self.W2.T
        dz1 = da1 * self.leaky_relu_derivative(self.z1)
        dW1 = X.T @ dz1 / m
        db1 = np.sum(dz1, axis=0, keepdims=True) / m
        return dW1, db1, dW2, db2

    def update(self, X, output_signal):
        self.forward(X)
        dW1, db1, dW2, db2 = self.backward(X, output_signal)
        self.W1 -= self.learning_rate[0] * dW1
        self.b1 -= self.learning_rate[1] * db1
        self.W2 -= self.learning_rate[2] * dW2
        self.b2 -= self.learning_rate[3] * db2

    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 one_hot(labels, num_classes=10):
    return np.eye(num_classes)[labels]


# Axiom:
# Parameters should adapt to the time dynamics of error, not just to one-step fixes.
# The update signal combines:
# 1. instantaneous error
# 2. accumulated error
# 3. derivative of accumulated error
eta_error = 1.0
eta_memory = 0.03
eta_dynamics = 0.07
eta_hidden = 1.0
amp_hidden = 0.97
decay = 0.09
batch_size = 100

learning_rate = np.random.rand(4)
f = MLPClassifier(input_size=784, hidden_size=100, output_size=10, learning_rate=learning_rate)

#X_mean = [X_train[y_train==i].mean(0) for i in range(10)]
#y = f.forward(X)
#Xinv = np.stack([y[:,i] * Xm[i] for i in range(10)]) # look at each Xinv and acculmulate err

accumulated_error = np.zeros((batch_size, 10))
i = 0

while True:
    idx = np.random.randint(1000, len(X_train), batch_size)
    ids = np.random.randint(0, 1000, batch_size)
    X = X_train[idx]
    yt = y_train[idx]
    y_target = one_hot(yt)

    hidden_signal_error = f.forward(amp_hidden * X_train[ids]) - one_hot(y_train[ids])
    instantaneous_error = f.forward(X) - y_target
    next_accumulated_error = decay * accumulated_error + instantaneous_error + hidden_signal_error
    error_derivative = next_accumulated_error - accumulated_error

    # Replace short-sighted fixes with a signal that mixes present error,
    # persistent bias, and the time derivative of that persistence.
    adjustment_signal = (
        eta_error * instantaneous_error
        + eta_memory * next_accumulated_error
        + eta_dynamics * error_derivative
        + eta_hidden * hidden_signal_error
    )

    f.update(X, adjustment_signal)
    accumulated_error = next_accumulated_error

    if i % 10 == 0:
        print(
            i,"batch_acc=",f.score(X, yt),"test_acc=",f.score(X20, yt20),
            "hidden_acc=",f.score(X_train[ids], y_train[ids])
        )
    i += 1
