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

class MNIST_AttractorNet(nn.Module):
    """
    CCT-inspired architecture. 
    Treated as a vector field mapping image space to class attractors.
    """
    def __init__(self):
        super().__init__()
        self.flatten = nn.Flatten()
        self.field = nn.Sequential(
            nn.Linear(28*28, 256),
            nn.ReLU(),
            nn.Linear(256, 128),
            nn.ReLU(),
            nn.Linear(128, 10) # 10 Class Attractors
        )

    def forward(self, x):
        x = self.flatten(x)
        return self.field(x)

class TectonicController:
    """
    Version 2.0: Not just pressure, but Tectonic Shifting.
    Handles the 'Stalled Waterfall' problem identified in ex05.py.
    """
    def __init__(self, target_slope=-0.01, kp=0.1, ki=0.005, kd=0.02):
        self.target_slope = target_slope
        self.kp, self.ki, self.kd = kp, ki, kd
        
        self.prev_loss = None
        self.prev_slope = 0
        self.integral = 0
        self.stagnation_counter = 0

    def compute_control(self, current_loss):
        if self.prev_loss is None:
            self.prev_loss = current_loss
            return 1.0, False

        slope = current_loss - self.prev_loss
        error = slope - self.target_slope
        
        self.integral += error
        derivative = slope - self.prev_slope
        
        pressure = 1.0 + (self.kp * error) + (self.ki * self.integral) + (self.kd * derivative)
        pressure = max(0.1, min(pressure, 15.0))
        
        # --- STAGNATION DETECTION (The Tectonic Trigger) ---
        # If loss is flat (slope ~ 0) but pressure is high, we are stuck.
        if abs(slope) < 1e-4 and pressure > 2.0:
            self.stagnation_counter += 1
        else:
            self.stagnation_counter = 0
        
        # Trigger a 'Tectonic Shift' (Morphic Perturbation) if stuck for too long
        shift_needed = self.stagnation_counter > 10
        
        self.prev_loss = current_loss
        self.prev_slope = slope
        return pressure, shift_needed

def train_tectonic():
    # Data
    transform = transforms.Compose([transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,))])
    train_set = datasets.MNIST('../data', train=True, download=True, transform=transform)
    test_set = datasets.MNIST('../data', train=False, transform=transform)
    train_loader = DataLoader(train_set, batch_size=64, shuffle=True)
    test_loader = DataLoader(test_set, batch_size=1000, shuffle=False)

    model = MNIST_AttractorNet()
    criterion = nn.CrossEntropyLoss()
    optimizer = optim.Adam(model.parameters(), lr=1e-3)
    controller = TectonicController()
    
    history = []
    
    for epoch in range(10): # 10 epochs for demonstration
        model.train()
        epoch_losses = []
        
        for batch_idx, (data, target) in enumerate(train_loader):
            optimizer.zero_grad()
            output = model(data)
            loss = criterion(output, target)
            
            # Get Control Signals
            pressure, shift = controller.compute_control(loss.item())
            
            loss.backward()
            
            # Apply Pressure
            for param in model.parameters():
                if param.grad is not None:
                    param.grad.data.mul_(pressure)
            
            # Execute Tectonic Shift if stuck
            if shift:
                with torch.no_grad():
                    for param in model.parameters():
                        # Inject a small morphic rotation to jump out of the basin
                        param.add_(torch.randn_like(param) * 0.01)
            
            optimizer.step()
            epoch_losses.append(loss.item())
            
        avg_loss = np.mean(epoch_losses)
        history.append(avg_loss)
        
        # Evaluation
        model.eval()
        correct = 0
        with torch.no_grad():
            for data, target in test_loader:
                output = model(data)
                pred = output.argmax(dim=1, keepdim=True)
                correct += pred.eq(target.view_as(pred)).sum().item()
        
        accuracy = 100. * correct / len(test_loader.dataset)
        print(f"Epoch {epoch} | Loss: {avg_loss:.4f} | Acc: {accuracy:.2f}% | Last Pressure: {pressure:.2f}")

    return history

if __name__ == "__main__":
    history = train_tectonic()
    plt.plot(history)
    plt.title("Tectonic Waterfall Loss")
    plt.ylabel("Loss")
    plt.xlabel("Epoch")
    plt.show()
