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

def lowpass_filter(signal, cutoff_freq, sample_rate=784):
    """Apply low-pass filter using FFT, keeping only frequencies below cutoff."""
    freq_domain = fft(signal)
    frequencies = np.fft.fftfreq(len(signal), d=1.0/sample_rate)
    # Zero out frequencies above cutoff
    freq_domain[np.abs(frequencies) > cutoff_freq] = 0
    return np.real(ifft(freq_domain))

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

# Define the MLP Classifier
class MLPClassifier:
    def __init__(self, input_size, hidden_size, output_size, learning_rate=np.random.rand(4)):
        self.W1 = np.random.randn(input_size, hidden_size) * 0.01
        self.b1 = np.zeros((1, hidden_size))
        self.W2 = np.random.randn(hidden_size, output_size) * 0.01
        self.b2 = np.zeros((1, output_size))
        self.learning_rate = learning_rate
    
    def relu(self, x):
        return np.maximum(0, x)
    
    def relu_derivative(self, x):
        return np.where(x > 0, 1, 0)
    
    def softmax(self, x):
        exp_x = np.exp(x - np.max(x, axis=1, keepdims=True))
        return exp_x / np.sum(exp_x, axis=1, keepdims=True)
    
    def forward(self, X):
        self.z1 = np.dot(X, self.W1) + self.b1
        self.a1 = self.relu(self.z1)
        self.z2 = np.dot(self.a1, self.W2) + self.b2
        output = self.softmax(self.z2)
        return output
    
    def compute_loss(self, y_true, y_pred):
        m = y_true.shape[0]
        loss = -np.sum(y_true * np.log(y_pred + 1e-9)) / m
        return loss
    
    def backward(self, X, y_true, y_pred):
        m = y_true.shape[0]
        dz2 = y_pred - y_true
        dW2 = np.dot(self.a1.T, dz2) / m
        db2 = np.sum(dz2, axis=0, keepdims=True) / m
        da1 = np.dot(dz2, self.W2.T)
        dz1 = da1 * self.relu_derivative(self.z1)
        dW1 = np.dot(X.T, dz1) / m
        db1 = np.sum(dz1, axis=0, keepdims=True) / m
        return dW1, db1, dW2, db2
    
    def update(self, X, y_true):
        y_pred = self.forward(X)
        dW1, db1, dW2, db2 = self.backward(X, y_true, y_pred)
        self.W1 -= self.learning_rate[0] * (dW1 - lowpass_filter(dW1,100))
        self.b1 -= self.learning_rate[1] * db1
        self.W2 -= self.learning_rate[2] * (dW2 - lowpass_filter(dW2,100))
        self.b2 -= self.learning_rate[3] * db2
    
    def predict(self, X):
        probabilities = self.forward(X)
        return np.argmax(probabilities, axis=1)

    def score(self, X, y_true):
        y = self.predict(X)
        return np.mean(y == y_true)

# Initialize and train the MLP Classifier
learning_rate = np.random.rand(4)
#f = MLPClassifier(input_size=784, hidden_size=100, output_size=10, learning_rate=learning_rate)
t = np.linspace(0,2*np.pi,784)
f = MLPClassifier2([100]*3)
# Modes

i = 0
while True:
    idx = np.random.randint(0, 60000, 100)
    X = X_train[idx]
    yt = y_train[idx]
    for _ in range(2):
        idx = np.random.randint(0, 60000, 100)
        X0 = X_train[idx]
        y0 = y_train[idx]
        for _ in range(2):
            idx = np.random.randint(0, 60000, 100)
            X1 = X_train[idx]
            y1 = y_train[idx]
            for _ in range(1):
                idx = np.random.randint(0, 60000, 100)
                X2 = X_train[idx]
                y2 = y_train[idx]
                for _ in range(1):
                    idx = np.random.randint(0, 60000, 100)
                    X3 = X_train[idx]
                    y3 = y_train[idx]
                    for _ in range(1):
                        idx = np.random.randint(0, 60000, 100)
                        X4 = X_train[idx]
                        y4 = y_train[idx]
                        for _ in range(1):
                            idx = np.random.randint(0, 60000, 100)
                            X5 = X_train[idx]
                            y5 = y_train[idx]
                            for _ in range(1):
                                idx = np.random.randint(0, 60000, 100)
                                X6 = X_train[idx]
                                y6 = y_train[idx]
                                """f.update(X,np.eye(10)[yt])
                                f.update(X6,np.eye(10)[y6])
                                f.update(X5,np.eye(10)[y5])
                                f.update(X4,np.eye(10)[y4])
                                f.update(X3,np.eye(10)[y3])
                                f.update(X2,np.eye(10)[y2])
                                f.update(X1,np.eye(10)[y1])
                                f.update(X0,np.eye(10)[y0])"""
                                f.partial_fit(X6,y6,classes=range(10))
                                print(i, f.score(X20,yt20))
                                f.partial_fit(X5,f.predict(X5),classes=range(10))
                                f.partial_fit(X4,y4,classes=range(10))
                                f.partial_fit(X3,f.predict(X3),classes=range(10))
                                f.partial_fit(X2,y2,classes=range(10))
                                f.partial_fit(X1,y1,classes=range(10))
                                f.partial_fit(X0,y0,classes=range(10))
                                f.partial_fit(X,yt,classes=range(10))
                                

    i += 1

