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

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

#err = yt - X@w1@w2
#err = yt@w2 - X@w1

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

yt = y_train[:100]
X = X_train[:100]

# np.eye(10)[yt]@w2 = X@w1

while True:
    a = np.eye(10)[yt]@w2
    b = X@w1
    err = a*a*a - a*a*b 
    w1 += 0.1 * X.T @ err 
    w2 += 0.1 * np.eye(10)[yt].T @ err
    forward = X@w1@np.linalg.pinv(w2)
    score = np.mean(forward.argmax(1)==yt)
    print(np.sum(err**2), score)
