"""
Resonant Filter MLP - MNIST
Each layer applies a resonant filter with learnable resonance parameter
"""

import torch
import torch.nn as nn
import torch.nn.functional as F
import torch.optim as optim
from torch.utils.data import DataLoader
from torchvision import datasets, transforms
import numpy as np
import matplotlib.pyplot as plt

torch.manual_seed(42)
np.random.seed(42)

print("=" * 60)
print("RESONANT FILTER MLP - MNIST")
print("=" * 60)


class ResonantFilter(nn.Module):
    """
    A simple resonant filter layer.
    Output = activation(W @ x + r * y_prev)
    where r is a learnable resonance coefficient (feedback).
    
    This creates frequency-selective behavior - resonance amplifies
    certain signal components based on the feedback strength.
    """
    
    def __init__(self, in_features, out_features):
        super().__init__()
        self.weight = nn.Parameter(torch.randn(out_features, in_features) * 0.05)
        self.bias = nn.Parameter(torch.zeros(out_features))
        
        # Resonance: feedback coefficient (-1 to 1)
        self.resonance = nn.Parameter(torch.tensor(0.1))
        
        # Internal state for recurrent filtering
        self.register_buffer('state', torch.zeros(1, out_features), persistent=False)
    
    def forward(self, x):
        # Linear transformation
        out = F.linear(x, self.weight, self.bias)
        
        # Apply resonance using a batch-independent memory vector.
        out = out + self.resonance * self.state
        
        # Track the mean layer response so the state shape stays stable
        # even when the last batch is smaller.
        self.state = out.detach().mean(dim=0, keepdim=True)
        
        return out
    
    def reset(self):
        self.state.zero_()


class ResonantFilterMLP(nn.Module):
    """MLP with resonant filter layers"""
    
    def __init__(self, input_size=784, hidden_sizes=[256, 128], output_size=10):
        super().__init__()
        sizes = [input_size] + hidden_sizes + [output_size]
        
        self.filters = nn.ModuleList([
            ResonantFilter(sizes[i], sizes[i+1])
            for i in range(len(sizes) - 1)
        ])
    
    def forward(self, x):
        x = x.view(x.size(0), -1)
        
        for i, filter in enumerate(self.filters):
            x = filter(x)
            if i < len(self.filters) - 1:  # ReLU for hidden layers
                x = F.relu(x)
        
        return x
    
    def reset_states(self):
        for filter in self.filters:
            filter.reset()


class StandardMLP(nn.Module):
    """Baseline: standard MLP without resonance"""
    
    def __init__(self, input_size=784, hidden_sizes=[256, 128], output_size=10):
        super().__init__()
        sizes = [input_size] + hidden_sizes + [output_size]
        self.layers = nn.ModuleList([
            nn.Linear(sizes[i], sizes[i+1])
            for i in range(len(sizes) - 1)
        ])
    
    def forward(self, x):
        x = x.view(x.size(0), -1)
        for i, layer in enumerate(self.layers):
            x = layer(x)
            if i < len(self.layers) - 1:
                x = F.relu(x)
        return x


def train_epoch(model, loader, criterion, optimizer, device, reset_states=True):
    model.train()
    total_loss = 0
    correct = 0
    total = 0
    
    if reset_states:
        model.reset_states()
    
    for x, y in loader:
        x, y = x.to(device), y.to(device)
        
        optimizer.zero_grad()
        output = model(x)
        loss = criterion(output, y)
        loss.backward()
        optimizer.step()
        
        total_loss += loss.item() * x.size(0)
        pred = output.argmax(dim=1)
        correct += (pred == y).sum().item()
        total += x.size(0)
    
    return total_loss / total, correct / total


@torch.no_grad()
def evaluate(model, loader, device):
    model.eval()
    correct = 0
    total = 0
    
    # Reset states for evaluation
    if hasattr(model, 'reset_states'):
        model.reset_states()
    
    for x, y in loader:
        x, y = x.to(device), y.to(device)
        output = model(x)
        pred = output.argmax(dim=1)
        correct += (pred == y).sum().item()
        total += x.size(0)
    
    return correct / total


def run_experiment(n_epochs=10, batch_size=256, lr=0.01):
    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
    print(f"\nDevice: {device}")
    
    # MNIST data
    transform = transforms.Compose([
        transforms.ToTensor(),
        transforms.Normalize((0.1307,), (0.3081,))
    ])
    
    train_loader = DataLoader(
        datasets.MNIST('data', train=True, download=True, transform=transform),
        batch_size=batch_size, shuffle=True
    )
    test_loader = DataLoader(
        datasets.MNIST('data', train=False, transform=transform),
        batch_size=batch_size
    )
    
    criterion = nn.CrossEntropyLoss()
    
    # ========== RESONANT FILTER MODEL ==========
    print("\n🧠 Training Resonant Filter MLP")
    print("-" * 50)
    
    resonant = ResonantFilterMLP().to(device)
    opt_res = optim.Adam(resonant.parameters(), lr=lr)
    
    resonant_train_losses = []
    resonant_train_accs = []
    
    for epoch in range(n_epochs):
        loss, acc = train_epoch(resonant, train_loader, criterion, opt_res, device)
        resonant_train_losses.append(loss)
        resonant_train_accs.append(acc)
        
        print(f"Epoch {epoch+1:2d}: Loss={loss:.4f}, Train Acc={acc:.4f}")
    
    resonant_test_acc = evaluate(resonant, test_loader, device)
    print(f"📊 Resonant Filter Test Accuracy: {resonant_test_acc:.4f}")
    
    # Print resonance values
    print("\n🔮 Learned Resonance Coefficients:")
    for i, f in enumerate(resonant.filters):
        print(f"  Layer {i}: resonance = {f.resonance.item():.4f}")
    
    # ========== STANDARD MODEL ==========
    print("\n⚡ Training Standard MLP (baseline)")
    print("-" * 50)
    
    standard = StandardMLP().to(device)
    opt_std = optim.Adam(standard.parameters(), lr=lr)
    
    standard_losses = []
    standard_accs = []
    
    for epoch in range(n_epochs):
        loss, acc = train_epoch(standard, train_loader, criterion, opt_std, device, reset_states=False)
        standard_losses.append(loss)
        standard_accs.append(acc)
        
        print(f"Epoch {epoch+1:2d}: Loss={loss:.4f}, Train Acc={acc:.4f}")
    
    standard_test_acc = evaluate(standard, test_loader, device)
    print(f"📊 Standard Test Accuracy: {standard_test_acc:.4f}")
    
    # ========== VISUALIZATION ==========
    fig, axes = plt.subplots(1, 4, figsize=(16, 4))
    
    # Loss comparison
    axes[0].plot(resonant_train_losses, 'b-o', label='Resonant Filter', linewidth=2)
    axes[0].plot(standard_losses, 'r--s', label='Standard', linewidth=2)
    axes[0].set_xlabel('Epoch')
    axes[0].set_ylabel('Loss')
    axes[0].set_title('Training Loss')
    axes[0].legend()
    axes[0].grid(True, alpha=0.3)
    
    # Accuracy comparison
    axes[1].plot(resonant_train_accs, 'b-o', label='Resonant Filter', linewidth=2)
    axes[1].plot(standard_accs, 'r--s', label='Standard', linewidth=2)
    axes[1].set_xlabel('Epoch')
    axes[1].set_ylabel('Accuracy')
    axes[1].set_title('Training Accuracy')
    axes[1].legend()
    axes[1].grid(True, alpha=0.3)
    
    # Test accuracy bar
    axes[2].bar(['Resonant\nFilter', 'Standard'], 
                [resonant_test_acc, standard_test_acc],
                color=['blue', 'red'], alpha=0.7, edgecolor='black')
    axes[2].set_ylabel('Test Accuracy')
    axes[2].set_title('Test Accuracy Comparison')
    axes[2].set_ylim([0.9, 1.0])
    axes[2].grid(True, alpha=0.3, axis='y')
    
    # Resonance coefficients
    layer_names = [f'L{i}' for i in range(len(resonant.filters))]
    res_values = [f.resonance.item() for f in resonant.filters]
    axes[3].bar(layer_names, res_values, color='purple', alpha=0.7, edgecolor='black')
    axes[3].axhline(y=0, color='black', linestyle='-', linewidth=0.5)
    axes[3].set_ylabel('Resonance Value')
    axes[3].set_title('Learned Resonance Coefficients')
    axes[3].grid(True, alpha=0.3, axis='y')
    
    plt.tight_layout()
    plt.savefig('resonant_filter_mlp.png', dpi=150)
    plt.show()
    
    # Summary
    print("\n" + "=" * 60)
    print("SUMMARY")
    print("=" * 60)
    print(f"Resonant Filter MLP Test Accuracy: {resonant_test_acc:.4f}")
    print(f"Standard MLP Test Accuracy:        {standard_test_acc:.4f}")
    diff = (resonant_test_acc - standard_test_acc) * 100
    print(f"Difference: {diff:+.2f}%")
    
    return {
        'resonant': {'train_loss': resonant_train_losses, 'test_acc': resonant_test_acc},
        'standard': {'train_loss': standard_losses, 'test_acc': standard_test_acc}
    }


if __name__ == "__main__":
    results = run_experiment(n_epochs=10, batch_size=256, lr=1e-3)
