import numpy as np

# Linear component
x = np.linspace(0, 1, 100)  # Linear input

# 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 with Linear Component
class MLPClassifier:
    def __init__(self, input_size, hidden_size, output_size, linear_dim=100, learning_rate=None):
        # Linear component: project input to linear_dim using x
        self.linear_weights = np.random.randn(input_size) * 0.01
        self.linear_bias = np.zeros((1, input_size))
        
        # MLP component
        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.linear_dim = linear_dim
        self.learning_rate = learning_rate if learning_rate is not None else np.random.rand(6)
        self.x = x  # Store linear component
    
    def linear_transform(self, X):
        """Apply linear transformation using x"""
        return np.dot(X, self.linear_weights) + self.linear_bias
    
    def relu(self, z):
        return np.maximum(0, z)
    
    def relu_derivative(self, z):
        return np.where(z > 0, 1, 0)
    
    def softmax(self, z):
        exp_z = np.exp(z - np.max(z, axis=1, keepdims=True))
        return exp_z / np.sum(exp_z, axis=1, keepdims=True)
    
    def forward(self, X):
        # Linear transformation
        #self.z0 = self.linear_transform(X)
        ids = np.argsort(X)
        self.z0 = self.linear_weights[ids]
        # MLP hidden layer
        self.z1 = np.dot(self.z0, self.W1) + self.b1
        self.a1 = self.relu(self.z1)
        
        # Output layer
        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]
        return -np.sum(y_true * np.log(y_pred + 1e-9)) / m
    
    def backward(self, X, y_true, y_pred):
        m = y_true.shape[0]
        
        # Output layer gradient
        dz2 = y_pred - y_true
        dW2 = np.dot(self.a1.T, dz2) / m
        db2 = np.sum(dz2, axis=0, keepdims=True) / m
        
        # Hidden layer gradient
        da1 = np.dot(dz2, self.W2.T)
        dz1 = da1 * self.relu_derivative(self.z1)
        dW1 = np.dot(self.z0.T, dz1) / m
        db1 = np.sum(dz1, axis=0, keepdims=True) / m
        
        # Linear component gradient
        d_linear_out = np.dot(dz1, X).sum(1) / m
        dW0 = np.dot(X.T, d_linear_out) / m
        db0 = np.sum(d_linear_out, axis=0, keepdims=True) / m
        
        return dW0, db0, dW1, db1, dW2, db2
    
    def update(self, X, y_true):
        y_pred = self.forward(X)
        dW0, db0, dW1, db1, dW2, db2 = self.backward(X, y_true, y_pred)
        
        self.linear_weights -= self.learning_rate[0] * dW0
        self.linear_bias -= self.learning_rate[1] * db0
        self.W1 -= self.learning_rate[2] * dW1
        self.b1 -= self.learning_rate[3] * db1
        self.W2 -= self.learning_rate[4] * dW2
        self.b2 -= self.learning_rate[5] * db2
    
    def predict(self, X):
        return np.argmax(self.forward(X), axis=1)
    
    def score(self, X, y_true):
        return np.mean(self.predict(X) == y_true)

# Initialize and train
learning_rate = np.random.rand(6)
f = MLPClassifier(input_size=784, hidden_size=100, output_size=10, 
                  linear_dim=100, learning_rate=learning_rate)

i = 0
while True:
    idx = np.random.randint(0, 60000, 100)
    X = X_train[idx]
    yt = y_train[idx]
    f.update(X, np.eye(10)[yt])
    print(i, f.score(X, yt))
    i += 1
