import numpy as np
from scipy.io.wavfile import read
from ran_array_complete import RANArray as RAN, real_ran

# 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)

# Normalize to [0, 1] (critical for stability)
X_train = X_train.astype(np.float64) / 255.0
X_test = X_test.astype(np.float64) / 255.0

# Network: 784 -> 100 -> 10, linear
w1 = 0.01 * RAN(np.random.randn(784, 10),np.ones((784,10)),np.random.randn(784, 10))
w2 = 0.01 * RAN(np.zeros((10,10)),np.ones((10,10)),np.random.randn(10, 10))

lr = np.random.rand(2)
step = 0

while True:
    idx = np.random.randint(0, 60000, 100)
    x = real_ran(X_train[idx])
    yt = real_ran(np.eye(10)[y_train[idx]])

    # Forward
    y1 = x @ w1
    y2 = y1 @ w2

    # Backprop
    err2 = yt - y2.collapse()
    grad2 = y1.T @ err2

    # Correct hidden layer error: backprop through w2
    err1 = err2 @ w2.T
    grad1 = x.T @ err1

    # Update
    w2 = w2 + lr[0] * grad2
    w1 = w1 + lr[1] * grad1

    if step % 100 == 0:
        print(step, np.mean(err2.collapse()**2))
    step += 1
