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

class ODECCTClassifier(nn.Module):
    def __init__(self, input_dim=784, hidden_dim=64, num_classes=10, steps=5, dt=0.1):
        super(ODECCTClassifier, self).__init__()
        self.hidden_dim = hidden_dim
        self.steps = steps
        self.dt = dt
        
        # Physics Parameters: alpha (field decay), gamma (friction), omega (natural frequency)
        self.alpha = 0.5
        self.gamma = 0.3
        self.omega = 1.0
        
        # Initial Core Stationary Matrix Backbone (Base topology skeleton)
        self.W_0 = nn.Parameter(torch.randn(hidden_dim, hidden_dim) * 0.01)
        
        # Input Projection Map (Translates pixel potentials to latent nodes)
        self.input_projection = nn.Linear(input_dim, hidden_dim)
        # Readout head mapping the collapsed hidden state to digit labels
        self.classifier_head = nn.Linear(hidden_dim, num_classes)
        
        # Structural Spatial Metric Node Coordinates (For potential field function)
        # Initializing pseudo-coordinates for hidden nodes in a 2D space
        grid_size = int(np.ceil(np.sqrt(hidden_dim)))
        x = torch.linspace(-1, 1, grid_size)
        y = torch.linspace(-1, 1, grid_size)
        grid_x, grid_y = torch.meshgrid(x, y, indexing='ij')
        coords = torch.stack([grid_x.flatten(), grid_y.flatten()], dim=1)[:hidden_dim]
        self.register_buffer('coords', coords)
        
        # Precompute Spatial Influence Field Metric phi(s, k) between all node pairs
        # phi_metric[i, s, k] represents influence of node i on edge (s->k)
        dists = torch.cdist(self.coords, self.coords, p=2) # [hidden_dim, hidden_dim]
        phi_metric = torch.zeros(hidden_dim, hidden_dim, hidden_dim)
        for i in range(hidden_dim):
            # Distance from node i to s and node i to k
            d_is = dists[i, :].unsqueeze(1) # [hidden_dim, 1]
            d_ik = dists[i, :].unsqueeze(0) # [1, hidden_dim]
            spatial_decay = torch.exp(-self.alpha * (d_is + d_ik))
            # Stabilized denominator with larger epsilon to prevent explosion
            denom = (d_is + d_ik) + 0.5 
            phi_metric[i] = spatial_decay / denom
            
        # Normalize phi_metric to keep the force field magnitude reasonable
        phi_metric = phi_metric / (phi_metric.norm() + 1e-5)
        self.register_buffer('phi_metric', phi_metric)

    def forward(self, x):
        batch_size = x.size(0)
        
        # 1. Image as initial potential energy surface
        h = torch.tanh(self.input_projection(x)) # Shape: [batch_size, hidden_dim]
        
        # Initialize Dynamic Manifold Weights W[s, k] to original backbone
        W_sk = self.W_0.clone().unsqueeze(0).repeat(batch_size, 1, 1) # [B, H, H]
        dW_sk_dt = torch.zeros_like(W_sk)
        
        # Initialize Self-Influence Diagonals W[i, i]
        W_ii = torch.zeros(batch_size, self.hidden_dim, device=x.device)
        
        # History tracker for State Hashing / Periodicity Collapse Checks
        state_history = []
        
        # 2. Step through the internal Update ODE System over Time t
        for t in range(self.steps):
            # Node resolution shifts based on hidden layer state energy
            W_ii_new = torch.sigmoid(h) 
            dW_ii_dt = (W_ii_new - W_ii) / self.dt
            delta_W_ii = W_ii_new - W_ii
            W_ii = W_ii_new
            
            # Compute Physics Force Field Propagation F from W[i,i] to W[s,k]
            # Optimized: Using matrix multiplication to avoid large intermediate tensors
            impulse_vec = dW_ii_dt * delta_W_ii # [B, H]
            F = torch.matmul(impulse_vec, self.phi_metric.view(self.hidden_dim, -1)).view(batch_size, self.hidden_dim, self.hidden_dim)
            
            # Second order ODE tracking physical momentum changes
            d2W_sk_dt2 = -self.gamma * dW_sk_dt - (self.omega**2) * (W_sk - self.W_0) + F
            
            # Integrate positions and velocities
            dW_sk_dt = dW_sk_dt + d2W_sk_dt2 * self.dt
            W_sk = W_sk + dW_sk_dt * self.dt
            
            # Apply evolved manifold weights back to system states
            h = torch.tanh(torch.bmm(W_sk, h.unsqueeze(-1)).squeeze(-1))
            
            # 3. State Hashing for Periodicity Recognition
            state_hash_representation = torch.round(h * 10) / 10
            
            # Look for collisions in the temporal timeline
            periodicity_detected = False
            for prev_h in state_history:
                if torch.allclose(state_hash_representation, prev_h, atol=1e-2):
                    periodicity_detected = True
                    break
                    
            if periodicity_detected:
                break
                
            state_history.append(state_hash_representation)
            
        # Output through readout mapping
        out = self.classifier_head(h)
        return out

# --- Train and Test Pipeline Execution Setup ---

def main():
    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
    print(f"Executing ODE-CCT Simulation Engine on Device: {device}")
    
    # Load and Normalize MNIST Image Assets
    transform = transforms.Compose([
        transforms.ToTensor(),
        transforms.Normalize((0.1307,), (0.3081,)),
        transforms.Lambda(lambda x: x.view(-1)) # Flatten 28x28 images into 784 arrays
    ])
    
    train_dataset = datasets.MNIST(root='../data', train=True, download=True, transform=transform)
    test_dataset = datasets.MNIST(root='../data', train=False, download=True, transform=transform)
    
    train_loader = DataLoader(train_dataset, batch_size=128, shuffle=True)
    test_loader = DataLoader(test_dataset, batch_size=1000, shuffle=False)
    
    # Initialize Model Instance
    model = ODECCTClassifier(input_dim=784, hidden_dim=64, num_classes=10, steps=6, dt=0.1).to(device)
    criterion = nn.CrossEntropyLoss()
    optimizer = optim.AdamW(model.parameters(), lr=0.003, weight_decay=1e-4)
    
    # Training Phase
    epochs = 3
    print("\n--- Initializing Model Optimization Loops ---")
    for epoch in range(1, epochs + 1):
        model.train()
        total_loss = 0
        correct = 0
        for batch_idx, (data, target) in enumerate(train_loader):
            data, target = data.to(device), target.to(device)
            
            optimizer.zero_grad()
            output = model(data)
            loss = criterion(output, target)
            loss.backward()
            optimizer.step()
            
            total_loss += loss.item()
            pred = output.argmax(dim=1, keepdim=True)
            correct += pred.eq(target.view_as(pred)).sum().item()
            
            if batch_idx % 150 == 0:
                print(f"Epoch {epoch} [{batch_idx * len(data)}/{len(train_loader.dataset)}] "
                      f"Current Loss: {loss.item():.4f}")
                
        avg_loss = total_loss / len(train_loader)
        accuracy = 100. * correct / len(train_loader.dataset)
        print(f"📊 Epoch {epoch} Complete -> Average Loss: {avg_loss:.4f}, Accuracy: {accuracy:.2f}%")

    # Evaluation / Testing Phase
    print("\n--- Running Inference Testing Pass ---")
    model.eval()
    test_loss = 0
    test_correct = 0
    with torch.no_grad():
        for data, target in test_loader:
            data, target = data.to(device), target.to(device)
            output = model(data)
            test_loss += criterion(output, target).item()
            pred = output.argmax(dim=1, keepdim=True)
            test_correct += pred.eq(target.view_as(pred)).sum().item()

    test_loss /= len(test_loader)
    test_acc = 100. * test_correct / len(test_loader.dataset)
    print(f"🏁 Final Test Results -> Mean Evaluation Loss: {test_loss:.4f}, Accuracy: {test_acc:.2f}%\n")

if __name__ == '__main__':
    main()
