import numpy as np
from sklearn.neural_network import MLPClassifier
from scipy.io.wavfile import read

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

ref = X_train[-200:]

v0 = np.var(np.concatenate([X[0,None],ref]),0)
v1 = np.var(np.concatenate([X[1,None],ref]),0)

f = MLPClassifier()
k=0
while True:
    idx = np.random.randint(0,60000,100)
    X = X_train[idx]
    yt = y_train[idx]
    Xvar = [X[i] * np.var(np.concatenate([X[i,None],ref]),0) for i in range(100)]
    f.partial_fit(Xvar, yt, classes=range(10))
    print(k, f.score(Xvar,yt))
    
    k+=1
