"""
Resonant Training Framework - Simplified & Fast
Core concept: Phase-aligned gradient updates with cross-agent resonance
"""

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 TRAINING FRAMEWORK (Simplified)")
print("=" * 60)


class ResonantLinear(nn.Module):
    """Linear layer with phase-coupled gradient updates"""
    
    def __init__(self, in_features, out_features):
        super().__init__()
        self.weight = nn.Parameter(torch.randn(out_features, in_features) * 0.1)
        self.bias = nn.Parameter(torch.zeros(out_features))
        
        # Phase state (simplified - one phase per layer)
        self.phase = torch.rand(()) * 2 * np.pi
        self.coherence = 0.0
    
    def forward(self, x):
        # Add small phase-based perturbation to weights
        perturbation = 0.01 * torch.sin(self.phase)
        effective_weight = self.weight * (1 + perturbation)
        return F.linear(x, effective_weight, self.bias)
    
    def align_phase(self, gradient_magnitude, lr=0.01):
        """Align phase toward gradient direction"""
        target = torch.atan2(torch.tensor(gradient_magnitude), torch.tensor(1.0)).item()
        diff = target - self.phase
        self.phase += lr * np.sin(diff)


class ResonantMLP(nn.Module):
    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([
            ResonantLinear(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 resonant_update(self, gradients, lr=0.001, phase_lr=0.1):
        """Update weights + align phases"""
        with torch.no_grad():
            for i, layer in enumerate(self.layers):
                layer.weight.data -= lr * gradients[f'layers.{i}.weight']
                layer.bias.data -= lr * gradients[f'layers.{i}.bias']
                
                grad_mag = gradients[f'layers.{i}.weight'].abs().mean().item()
                layer.align_phase(grad_mag, phase_lr)
    
    def compute_coherence(self):
        """Phase coherence across layers"""
        phases = [l.phase for l in self.layers]
        mean_x = np.mean([np.cos(p.item()) for p in phases])
        mean_y = np.mean([np.sin(p.item()) for p in phases])
        return np.sqrt(mean_x**2 + mean_y**2)


class ResonantAgent:
    """Single AI agent with resonant training"""
    
    def __init__(self, agent_id):
        self.id = agent_id
        self.model = ResonantMLP()#.to(device)
        self.phase = np.random.uniform(0, 2*np.pi)
        self.frequency = 1.0
        self.loss_history = []
        self.phase_history = []
    
    def train_step(self, x, y, criterion, grad_clip=5.0):
        self.model.train()
        self.model.zero_grad()
        output = self.model(x)
        loss = criterion(output, y)
        loss.backward()
        torch.nn.utils.clip_grad_norm_(self.model.parameters(), max_norm=grad_clip)
        
        grad_mag = sum(p.grad.abs().sum().item() for p in self.model.parameters())
        self.phase += 0.05 * np.sin(grad_mag * 0.01)
        self.phase_history.append(self.phase)
        self.loss_history.append(loss.item())
        
        correct = (output.argmax(1) == y).sum().item()
        return loss.item(), correct

    def collect_gradients(self):
        gradients = {}
        for name, param in self.model.named_parameters():
            if param.grad is not None:
                gradients[name] = param.grad.detach().clone()
        return gradients


def run_experiment(n_epochs=5, n_agents=3, batch_size=256):
    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
    print(f"\nDevice: {device}")
    # Load 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 TRAINING ==========
    print(f"\n🧠 Resonant Training ({n_agents} agents)")
    print("-" * 50)
    
    agents = [ResonantAgent(i) for i in range(n_agents)]
    for agent in agents:
        agent.model.to(device)
    resonant_losses = []
    coherence_history = []
    
    for epoch in range(n_epochs):
        epoch_losses = []
        epoch_coherences = []
        epoch_correct = 0
        epoch_total = 0
        
        for x, y in train_loader:
            x, y = x.to(device), y.to(device)
            
            # Local gradients
            local_gradients = []
            for agent in agents:
                loss, correct = agent.train_step(x, y, criterion)
                epoch_losses.append(loss)
                epoch_correct += correct
                epoch_total += y.size(0)
                local_gradients.append(agent.collect_gradients())
            
            # Cross-resonance: phase-aligned updates
            phases = np.array([a.phase for a in agents])
            phase_diff = phases[:, None] - phases[None, :]
            coupling = 0.3 * np.cos(phase_diff)
            epoch_coherences.append(np.mean(np.abs(np.cos(phase_diff))))
            
            # Apply resonant updates using bounded phase-weighted gradient blending.
            for i, agent in enumerate(agents):
                blended_gradients = {}
                for name, grad in local_gradients[i].items():
                    blended = grad.clone()
                    for j in range(n_agents):
                        if i == j:
                            continue
                        blended += 0.15 * coupling[i, j] * local_gradients[j][name]
                    blended_gradients[name] = torch.nan_to_num(blended, nan=0.0, posinf=1.0, neginf=-1.0)
                agent.model.resonant_update(blended_gradients, lr=0.001, phase_lr=0.05)
        
        # Track stats
        avg_loss = float(np.mean(epoch_losses))
        resonant_losses.append(avg_loss)
        coherence = float(np.mean(epoch_coherences))
        coherence_history.append(coherence)
        
        acc = epoch_correct / epoch_total
        
        print(f"Epoch {epoch+1}: Loss={avg_loss:.4f}, Coherence={coherence:.3f}, Acc={acc:.4f}")
    
    # Test
    resonant_acc = evaluate_agents(agents, test_loader)
    
    # ========== STANDARD TRAINING ==========
    print(f"\n⚡ Standard Training (baseline)")
    print("-" * 50)
    
    standard = ResonantMLP().to(device)
    optimizer = optim.Adam(standard.parameters(), lr=0.001)
    
    standard_losses = []
    for epoch in range(n_epochs):
        epoch_loss = 0
        epoch_correct = 0
        epoch_total = 0
        for x, y in train_loader:
            x, y = x.to(device), y.to(device)
            optimizer.zero_grad()
            output = standard(x)
            loss = criterion(output, y)
            loss.backward()
            torch.nn.utils.clip_grad_norm_(standard.parameters(), max_norm=5.0)
            optimizer.step()
            epoch_loss += loss.item()
            epoch_correct += (output.argmax(1) == y).sum().item()
            epoch_total += y.size(0)
        standard_losses.append(epoch_loss / len(train_loader))
        
        acc = epoch_correct / epoch_total
        print(f"Epoch {epoch+1}: Loss={standard_losses[-1]:.4f}, Acc={acc:.4f}")
    
    standard_acc = evaluate_single(standard, test_loader)
    
    # ========== VISUALIZATION ==========
    fig, axes = plt.subplots(1, 3, figsize=(12, 4))
    
    axes[0].plot(resonant_losses, 'b-o', label='Resonant')
    axes[0].plot(standard_losses, 'r--s', label='Standard')
    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)
    
    axes[1].plot(coherence_history, 'purple', linewidth=2)
    axes[1].set_xlabel('Epoch')
    axes[1].set_ylabel('Coherence')
    axes[1].set_title('Phase Coherence')
    axes[1].grid(True, alpha=0.3)
    
    # Phase trajectories
    for i, agent in enumerate(agents):
        axes[2].plot(agent.phase_history, label=f'Agent {i}', alpha=0.7)
    axes[2].set_xlabel('Step')
    axes[2].set_ylabel('Phase')
    axes[2].set_title('Phase Evolution')
    axes[2].legend()
    axes[2].grid(True, alpha=0.3)
    
    plt.tight_layout()
    plt.savefig('resonant_results.png', dpi=150)
    plt.show()
    
    # Summary
    print("\n" + "=" * 60)
    print("RESULTS")
    print("=" * 60)
    print(f"Resonant System: {resonant_acc:.4f}")
    print(f"Standard Network: {standard_acc:.4f}")
    print(f"Difference: {(resonant_acc - standard_acc)*100:+.2f}%")


def evaluate_agents(agents, loader):
    correct = 0
    total = 0
    for x, y in loader:
        #x, y = x.cuda(), y.cuda()
        for agent in agents:
            pred = agent.model(x).argmax(1)
            correct += (pred == y).sum().item()
            total += x.size(0)
    return correct / total

def evaluate_single(model, loader):
    correct = 0
    total = 0
    model.eval()
    with torch.no_grad():
        for x, y in loader:
            #x, y = x.cuda(), y.cuda()
            pred = model(x).argmax(1)
            correct += (pred == y).sum().item()
            total += x.size(0)
    return correct / total


if __name__ == "__main__":
    run_experiment(n_epochs=5, n_agents=3, batch_size=100)
