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

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


#yt = real_ran(np.random.rand(100,10))
yh = real_ran(np.random.randn(100,100))
#x = real_ran(np.random.rand(100,784))
w1 = 0.001 * real_ran(np.random.randn(784,100))
w2 = 0.001 * real_ran(np.random.randn(100,10))


step = 0
while True:
    idx = np.random.randint(0,60000,100)
    x = X_train[idx]
    yt = np.eye(10)[y_train[idx]]
    y1 = x @ w1
    y2 = y1 @ w2
    err2 = yt - y2
    grad2 = y1.T @ err2
    err1 = yh - y1
    grad1 = x.T @ err1
    w2 += 0.0001 * grad2
    w1 += 0.0001 * grad1
#    break
    if step % 100 == 0:
        print(step, np.mean(err2.collapse()**2))
    step += 1
