import numpy as np
from scipy.io.wavfile import read
from scipy.signal import find_peaks
from numpy.fft import fft, ifft
from numpy.linalg import pinv

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

w2 = 0.1 * np.random.randn(10,100)
w1 = 0.1 * np.random.randn(784,100)

i=0
while True:
    idx = np.random.randint(0,60000,100)
    X = X_train[idx]
    yt = y_train[idx]
    a = np.eye(10)[yt]@w2
    b = X@w1
    err = a*a*a - a*a*b
    w1 += 0.01 * X.T @ err 
    w2 -= 0.01 * np.eye(10)[yt].T @ err
    if i%1000==0:
        forward = X@w1@pinv(w2)
        score = np.mean(forward.argmax(1)==yt)
        print(np.sum(err**2), score)
    i+=1
