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

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

class CCT_MLP:
    def __init__(self, input_size=784, hidden_size=100, output_size=10, num_classes=10, learning_rate=np.random.rand(4)):
        # X_ becomes the CLASS DRIVER matrix (learned)
        self.X_ = np.random.randn(num_classes, input_size) * 0.01  # Shape: (10, 784)
        
        # Standard weights
        self.W1 = np.random.randn(input_size, hidden_size) * 0.01
        self.W1_k = 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.W2_k = np.random.randn(hidden_size, output_size) * 0.01
        self.b2 = np.zeros((1, output_size))
        
        self.num_classes = num_classes
        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 drive_X(self, X, yt):
        """
        Use X_[yt] to DRIVE X toward the target class.
        X_: (num_classes, 784) - lookup table per class
        yt: (batch_size,) - target class indices
        Returns: Driven input shaped (batch_size, 784)
        """
        # Select the driver for each sample in batch
        drivers = self.X_[yt]  # Shape: (batch_size, 784)
        
        # Option 1: Gating - modulate input by class driver
        driven_X = X * (1 + drivers)  # Soft gating
        
        # Option 2: Affine transformation - X transformed by class driver
        # driven_X = X + drivers  # Additive shift
        
        # Option 3: Linear combination - blend input with class prototype
        # driven_X = 0.7 * X + 0.3 * drivers
        
        return driven_X
    
    def forward(self, X, yt=None, training=True):
        # Training: use class driver to modulate input
        if training and yt is not None:
            X_driven = self.drive_X(X, yt)
        else:
            X_driven = X
            
        if i % 2 == 0:
            self.z1 = np.dot(X_driven, self.W1) + self.b1
            self.a1 = self.relu(self.z1)
            self.z2 = np.dot(self.a1, self.W2) + self.b2
        else:
            self.z1 = np.dot(X_driven, self.W1_k) + self.b1
            self.a1 = self.relu(self.z1)
            self.z2 = np.dot(self.a1, self.W2_k) + self.b2
            
        return self.softmax(self.z2)
    
    def backward(self, X, y_true, y_pred, yt):
        m = y_true.shape[0]
        
        # Get driven input for gradient computation
        X_driven = self.drive_X(X, yt)
        
        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 if i % 2 == 0 else self.W2_k.T)
        dz1 = da1 * self.relu_derivative(self.z1)
        dW1 = np.dot(X_driven.T, dz1) / m
        db1 = np.sum(dz1, axis=0, keepdims=True) / m
        
        # Gradient for X_ (class drivers) - how to update drivers
        # dL/dX_[yt] = dL/dX_driven * dX_driven/dX_
        dX_driven = np.dot(dz1, (self.W1 if i % 2 == 0 else self.W1_k).T)
        dX_ = dX_driven  # Since X_driven = X * (1 + X_[yt])
        
        return dW1, db1, dW2, db2, dX_
    
    def update(self, X, y_true, yt):
        y_pred = self.forward(X, yt, training=True)
        dW1, db1, dW2, db2, dX_ = self.backward(X, y_true, y_pred, yt)
        
        if i % 2 == 0:
            self.W1 -= self.learning_rate[0] * dW1
        else:
            self.W1_k -= self.learning_rate[0] * dW1
            self.W1_k += 0.1 * self.learning_rate[0] * self.W1
            
        self.b1 -= self.learning_rate[1] * db1
        
        if i % 2 == 0:
            self.W2 -= self.learning_rate[2] * dW2
        else:
            self.W2_k -= self.learning_rate[2] * dW2
            self.W2_k += 0.1 * self.learning_rate[2] * self.W2
            
        self.b2 -= self.learning_rate[3] * db2
        
        # Update class drivers - this is the CCT component
        self.X_[yt] -= 0.1 * self.learning_rate[3] * dX_  # Only update drivers for target classes
    
    def predict(self, X):
        probabilities = self.forward(X, training=False)
        return np.argmax(probabilities, axis=1)
    
    def score(self, X, y):
        return np.mean(self.predict(X) == y)


# Training loop
learning_rate = np.random.rand(4)/2
f = CCT_MLP(input_size=784, hidden_size=100, output_size=10, learning_rate=learning_rate)

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

"""# 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.W1_k = 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.W2_k = np.random.randn(hidden_size, output_size) * 0.01
        self.b2 = np.zeros((1, output_size))
        self.learning_rate = learning_rate
        self.k = np.random.randn(1,784)
    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):
        if i%2==0:
            self.z1 = np.dot(X_[yt], 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)    
        else:
            self.z1 = np.dot(X, self.W1_k) + self.b1
            self.a1 = self.relu(self.z1)
            self.z2 = np.dot(self.a1, self.W2_k) + 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)
        if i%2 == 0:self.W1 -= self.learning_rate[0] * dW1
        if i%2 != 0:
            self.W1_k -= self.learning_rate[0] * dW1
            self.W1_k += 0.01 * self.W1
        self.b1 -= self.learning_rate[1] * db1
        if i%2 == 0:self.W2 -= self.learning_rate[2] * dW2
        if i%2 != 0:
            self.W2_k -= self.learning_rate[2] * dW2
            self.W2_k += 0.01 * self.W2
        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)
X_ = np.random.randn(10,784)
k = np.random.randn(1,784)

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

    f.update(X, np.eye(10)[yt])

    i += 1
"""
