import numpy as np
from scipy.io.wavfile import read
from convmlp import ConvMLPClassifier

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]


# Define convolutional layers
conv_layers = [
    {'type': 'conv2d', 'filters': 32, 'kernel_size': 3, 'padding': 'same'},
    {'type': 'maxpool2d', 'pool_size': 2},
    {'type': 'flatten'},  # Important: flatten before MLP
]

# Create model
model = ConvMLPClassifier(
    conv_layers=conv_layers,
    mlp_layers=[128, 64],  # MLP hidden layers after conv
    input_shape=(1, 28, 28),  # For MNIST
    max_iter=100,
    batch_size=1
)

i=0
# Fit and predict
while True:
    idx = np.random.randint(0,60000,100)
    X = X_train[idx]
    yt = y_train[idx]
    if i>0:
        print(i, np.sum(model.predict(X)==yt))
    model.partial_fit(X, yt, classes=range(10))
    i+=1
