import numpy as np
from sklearn.neural_network import MLPClassifier
from sklearn.linear_model import LinearRegression
from scipy.io.wavfile import read
from scipy.stats import entropy
import torch
from torch import nn



# 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]
X20 = torch.from_numpy(X20)
yt20 = torch.from_numpy(yt20)

w = nn.Parameter(torch.randn(100,1,784))
f = MLPClassifier()

g = torch.nn.Sequential(nn.Linear(784,100),nn.ReLU(), nn.Linear(100,10))
cross = nn.CrossEntropyLoss()
optim = torch.optim.Adam(g.parameters())
optim2 = torch.optim.Adam([w])

i=0
while True:
    idx = np.random.randint(0,60000,100)
    X = torch.from_numpy(X_train[idx])
    yt = torch.from_numpy(y_train[idx])
    loss = 0
    for k in range(100):
        loss = cross(g(X),yt)
        optim.zero_grad()
        loss.backward()
        optim.step()
        idx = np.random.randint(0,60000,100)
        X0 = torch.from_numpy(X_train[idx])
        y0 = torch.from_numpy(y_train[idx])
        optim.zero_grad()
        loss = cross(g(X0),y0)
        loss.backward()
        optim.step()    
    print(i, sum(g(X20).argmax(1)==yt20))
    i+=1


