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
from sklearn.neural_network import MLPClassifier

# 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(784,100)
w1 = 0.1 * np.random.randn(784,100)
#X_ = 0.001 * np.random.randn(10,784)

f = MLPClassifier()

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