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


def load_dataset():
    X_train = read("../X_train.wav")[1].reshape(-1, 784).astype(np.float32)
    y_train = np.rint(read("../y_train.wav")[1] * 9).astype(np.int64)
    X_test = read("../X_test.wav")[1].reshape(-1, 784).astype(np.float32)
    y_test = np.rint(read("../y_test.wav")[1] * 9).astype(np.int64)
    return X_train, y_train, X_test, y_test


class MLPWithLearnableW:
    def __init__(
        self,
        input_size,
        hidden_size,
        output_size,
        n_views=8,
        lr=0.05,
        projection_lr_scale=0.1,
        grad_clip=1.0,
        seed=0,
    ):
        rng = np.random.default_rng(seed)
        self.n_views = n_views
        self.lr = lr
        self.projection_lr_scale = projection_lr_scale
        self.grad_clip = grad_clip

        self.w = rng.normal(
            0.0,
            1.0 / np.sqrt(input_size),
            size=(n_views, input_size, input_size),
        ).astype(np.float32)

        self.W1 = rng.normal(
            0.0, np.sqrt(2.0 / input_size), size=(input_size, hidden_size)
        ).astype(np.float32)
        self.b1 = np.zeros((1, hidden_size), dtype=np.float32)
        self.W2 = rng.normal(
            0.0, np.sqrt(2.0 / hidden_size), size=(hidden_size, output_size)
        ).astype(np.float32)
        self.b2 = np.zeros((1, output_size), dtype=np.float32)

    def relu(self, x):
        return np.maximum(x, 0.0)

    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, view_idx):
        X_proj = X @ self.w[view_idx]
        z1 = X_proj @ self.W1 + self.b1
        a1 = self.relu(z1)
        z2 = a1 @ self.W2 + self.b2
        y_pred = self.softmax(z2)
        return y_pred, X_proj, z1, a1

    def train_batch(self, X, y):
        m = X.shape[0]
        y_true = np.eye(self.b2.shape[1], dtype=np.float32)[y]

        dW1 = np.zeros_like(self.W1)
        db1 = np.zeros_like(self.b1)
        dW2 = np.zeros_like(self.W2)
        db2 = np.zeros_like(self.b2)
        dw = np.zeros_like(self.w)

        for view_idx in range(self.n_views):
            y_pred, X_proj, z1, a1 = self.forward(X, view_idx)

            dz2 = (y_pred - y_true) / self.n_views
            dW2 += (a1.T @ dz2) / m
            db2 += np.sum(dz2, axis=0, keepdims=True) / m

            dz1 = (dz2 @ self.W2.T) * (z1 > 0.0)
            dW1 += (X_proj.T @ dz1) / m
            db1 += np.sum(dz1, axis=0, keepdims=True) / m

            dX_proj = dz1 @ self.W1.T
            dw[view_idx] = (X.T @ dX_proj) / m

        np.clip(dW1, -self.grad_clip, self.grad_clip, out=dW1)
        np.clip(db1, -self.grad_clip, self.grad_clip, out=db1)
        np.clip(dW2, -self.grad_clip, self.grad_clip, out=dW2)
        np.clip(db2, -self.grad_clip, self.grad_clip, out=db2)
        np.clip(dw, -self.grad_clip, self.grad_clip, out=dw)

        self.W1 -= self.lr * dW1
        self.b1 -= self.lr * db1
        self.W2 -= self.lr * dW2
        self.b2 -= self.lr * db2
        self.w -= (self.lr * self.projection_lr_scale) * dw

    def predict_proba(self, X):
        probs = np.zeros((X.shape[0], self.b2.shape[1]), dtype=np.float32)
        for view_idx in range(self.n_views):
            probs += self.forward(X, view_idx)[0]
        return probs / self.n_views

    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)


X_train, y_train, X_test, y_test = load_dataset()
X_eval = X_test[:1000]
y_eval = y_test[:1000]

model = MLPWithLearnableW(
    input_size=784,
    hidden_size=128,
    output_size=10,
    n_views=8,
    lr=0.05,
    projection_lr_scale=0.1,
    grad_clip=1.0,
    seed=0,
)

batch_size = 128
step = 0

while True:
    idx = np.random.randint(0, X_train.shape[0], batch_size)
    model.train_batch(X_train[idx], y_train[idx])

    if step % 10 == 0:
        score = model.score(X_eval, y_eval)
        print(f"{step}: {score:.3f}", flush=True)

    step += 1
