import numpy as np
from scipy.io.wavfile import read
from scipy.fft import dct
import time

# Load data (same as before)
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 LearnableSinCosClassifier:
    """
    Learnable classifier based on sin/cos matrix theory.
    
    Architecture:
    1. Linear projection: x -> z = W*x + b  (learnable)
    2. Extract phase from z: θ = arctan2(z[:,1], z[:,0])
    3. Build Fourier features: [cos(θ), sin(θ), cos(2θ), sin(2θ), ...]
    4. Linear classifier on these features.
    
    This corresponds to a parametric sin/cos matrix where the matrix
    entries are linear combinations of sin/cos of learned phase.
    """
    
    def __init__(self, input_dim=784, hidden_dim=64, n_frequencies=8, n_classes=10, lr=0.01):
        self.input_dim = input_dim
        self.hidden_dim = hidden_dim
        self.n_frequencies = n_frequencies
        self.n_classes = n_classes
        self.lr = lr
        
        # Projection layer: input -> hidden (to compute phase)
        self.W = np.random.randn(input_dim, hidden_dim) * 0.01
        self.b = np.zeros((1, hidden_dim))
        
        # Classification layer: Fourier features (2*n_frequencies) -> classes
        self.feat_dim = 2 * n_frequencies
        self.W_cls = np.random.randn(self.feat_dim, n_classes) * 0.01
        self.b_cls = np.zeros((1, n_classes))
        
        # Store for gradient computation
        self.z = None
        self.theta = None
        self.features = None
        
    def _fourier_features(self, theta):
        """Generate [cos(kθ), sin(kθ)] for k=1..n_frequencies."""
        features = []
        for k in range(1, self.n_frequencies + 1):
            features.append(np.cos(k * theta))
            features.append(np.sin(k * theta))
        return np.column_stack(features)
    
    def forward(self, X):
        """Forward pass: project, compute phase, generate features, classify."""
        # Projection
        self.z = np.dot(X, self.W) + self.b  # (n, hidden_dim)
        
        # Phase from first two components (or use PCA-like)
        # We'll use the first two dimensions of z to compute phase
        # but we could also use all dimensions (e.g., via a learned mapping)
        # For simplicity, use first two.
        self.theta = np.arctan2(self.z[:, 1] + 1e-9, self.z[:, 0] + 1e-9)  # (n,)
        
        # Fourier features
        self.features = self._fourier_features(self.theta)  # (n, feat_dim)
        
        # Classification
        logits = np.dot(self.features, self.W_cls) + self.b_cls  # (n, n_classes)
        return logits
    
    def softmax(self, logits):
        exp_logits = np.exp(logits - np.max(logits, axis=1, keepdims=True))
        return exp_logits / np.sum(exp_logits, axis=1, keepdims=True)
    
    def compute_loss(self, logits, y_true):
        """Cross-entropy loss."""
        probs = self.softmax(logits)
        m = y_true.shape[0]
        loss = -np.sum(np.log(probs[np.arange(m), y_true] + 1e-9)) / m
        return loss
    
    def backward(self, X, y_true, logits):
        """Gradient descent via backpropagation."""
        m = y_true.shape[0]
        # One-hot encoding
        y_onehot = np.eye(self.n_classes)[y_true]
        
        # Gradient w.r.t. logits
        probs = self.softmax(logits)
        dlogits = (probs - y_onehot) / m  # (n, n_classes)
        
        # Gradients for classification layer
        dW_cls = np.dot(self.features.T, dlogits)  # (feat_dim, n_classes)
        db_cls = np.sum(dlogits, axis=0, keepdims=True)  # (1, n_classes)
        
        # Gradient w.r.t. features
        dfeatures = np.dot(dlogits, self.W_cls.T)  # (n, feat_dim)
        
        # Gradient w.r.t. theta (chain rule through Fourier features)
        # d(features)/dθ: each feature is sin(kθ) or cos(kθ)
        # We need to compute dL/dθ = sum over features dL/df * df/dθ
        dtheta = np.zeros(m)
        for k in range(1, self.n_frequencies + 1):
            idx_cos = 2*(k-1)
            idx_sin = 2*(k-1) + 1
            # cos(kθ) derivative = -k*sin(kθ)
            dtheta += dfeatures[:, idx_cos] * (-k * np.sin(k * self.theta))
            # sin(kθ) derivative = k*cos(kθ)
            dtheta += dfeatures[:, idx_sin] * (k * np.cos(k * self.theta))
        
        # Gradient w.r.t. z (phase is atan2(z1, z0))
        # θ = atan2(z1, z0) => dθ/dz0 = -z1/(z0^2+z1^2), dθ/dz1 = z0/(z0^2+z1^2)
        z0 = self.z[:, 0] + 1e-9
        z1 = self.z[:, 1] + 1e-9
        norm2 = z0**2 + z1**2
        dz = np.zeros_like(self.z)
        dz[:, 0] = dtheta * (-z1 / norm2)
        dz[:, 1] = dtheta * (z0 / norm2)
        
        # Gradient w.r.t. W and b
        dW = np.dot(X.T, dz)  # (input_dim, hidden_dim)
        db = np.sum(dz, axis=0, keepdims=True)  # (1, hidden_dim)
        
        return dW, db, dW_cls, db_cls
    
    def update(self, X, y_true):
        """One training step."""
        logits = self.forward(X)
        loss = self.compute_loss(logits, y_true)
        dW, db, dW_cls, db_cls = self.backward(X, y_true, logits)
        
        self.W -= self.lr[0] * dW
        self.b -= self.lr[1] * db
        self.W_cls -= self.lr[2] * dW_cls
        self.b_cls -= self.lr[3] * db_cls
        
        return loss
    
    def predict(self, X):
        logits = self.forward(X)
        return np.argmax(logits, axis=1)
    
    def score(self, X, y_true):
        return np.mean(self.predict(X) == y_true)


# Main execution
if __name__ == "__main__":
    print("=" * 60)
    print("Learnable Sin/Cos Matrix Classifier for MNIST")
    print("=" * 60)
    
    # Instantiate
    model = LearnableSinCosClassifier(
        input_dim=784,
        hidden_dim=100,
        n_frequencies=2,
        n_classes=10,
        lr=np.random.rand(4)
    )
    
    # Training loop
    epoch = 0
    print("\nTraining Learnable Sin/Cos Classifier...")
    while True:
        idx = np.random.randint(0,60000,100)
        X = X_train[idx]
        yt = y_train[idx]
        loss = model.update(X, yt)
        # Evaluate on subset
        test_acc = model.score(X20, yt20)
        print(f"Step {epoch+1}: loss={loss:.4f}, test_acc={test_acc:.4f}")
        epoch += 1
    print(f"\nFinal test accuracy: {model.score(X20, yt20):.4f}")
    
