import torch
import torch.nn as nn
import torch.nn.functional as F
import torch.optim as optim
from torchvision import datasets, transforms
from torch.utils.data import DataLoader

## 1. Define the Spiral Worm Classifier
class SpiralWormClassifier(nn.Module):
    def __init__(self, grid_size=10, path_steps=50):
        super(SpiralWormClassifier, self).__init__()
        self.grid_size = grid_size
        self.path_steps = path_steps
        
        # Project MNIST input (28x28 = 784) into a 3D space (10x10x10 = 1000)
        self.fc1 = nn.Linear(784, grid_size ** 3)
        
        # Classifier head that takes the path integral (features collected by the worm)
        self.fc2 = nn.Linear(path_steps, 10)
        
        # Pre-compute the parametric spiral path coordinates (normalized to [-1, 1] for grid_sample)
        self.register_buffer('worm_path', self._generate_spiral_path())

    def _generate_spiral_path(self):
        """Generates a 3D spiral parametric path (the worm's trajectory)."""
        t = torch.linspace(0, 1, self.path_steps)
        
        # Spiral equations
        theta = 4 * torch.pi * t  # Multi-turn spiral
        r = t                     # Growing radius over time
        
        x = r * torch.cos(theta)
        y = r * torch.sin(theta)
        z = 2 * t - 1             # Linearly progress from -1 to 1 along Z-axis
        
        # Stack into [path_steps, 3] coordinates
        path = torch.stack([x, y, z], dim=-1) 
        return path

    def forward(self, x):
        batch_size = x.size(0)
        x = x.view(batch_size, -1) # Flatten image to 784
        
        # 1. Project into 3D representation volume [Batch, 1, 10, 10, 10]
        grid_features = self.fc1(x).view(batch_size, 1, self.grid_size, self.grid_size, self.grid_size)
        
        # 2. Prepare the worm path coordinates for grid sampling
        # grid_sample expects coordinates shaped as [Batch, Output_Deep, Output_Height, Output_Width, 3]
        # We model the worm path as a 1D sequence of points in 3D space: [Batch, 1, 1, path_steps, 3]
        sample_coords = self.worm_path.view(1, 1, 1, self.path_steps, 3).expand(batch_size, -1, -1, -1, -1)
        
        # 3. The "Worm" eats/samples features along its parametric spiral path
        # Output shape: [Batch, Channels(1), 1, 1, path_steps]
        path_features = F.grid_sample(grid_features, sample_coords, align_corners=True)
        path_features = path_features.view(batch_size, self.path_steps)
        
        # 4. Classify based on the sequence of values gathered along the trajectory
        logits = self.fc2(path_features)
        return logits

## 2. Training and Testing Functions
def train(model, device, train_loader, optimizer, epoch):
    model.train()
    for batch_idx, (data, target) in enumerate(train_loader):
        data, target = data.to(device), target.to(device)
        optimizer.zero_grad()
        output = model(data)
        loss = F.cross_entropy(output, target)
        loss.backward()
        optimizer.step()
        
        if batch_idx % 200 == 0:
            print(f"Train Epoch: {epoch} [{batch_idx * len(data)}/{len(train_loader.dataset)} "
                  f"({100. * batch_idx / len(train_loader):.0f}%)]\tLoss: {loss.item():.6f}")

def test(model, device, test_loader):
    model.eval()
    test_loss = 0
    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 += F.cross_entropy(output, target, reduction='sum').item()
            pred = output.argmax(dim=1, keepdim=True)
            correct += pred.eq(target.view_as(pred)).sum().item()

    test_loss /= len(test_loader.dataset)
    accuracy = 100. * correct / len(test_loader.dataset)
    print(f"\nTest set: Average loss: {test_loss:.4f}, Accuracy: {correct}/{len(test_loader.dataset)} ({accuracy:.2f}%)\n")

## 3. Main Execution Block
if __name__ == "__main__":
    # Device setup
    device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
    print(f"Using device: {device}")
    
    # Data transformations & loading
    transform = transforms.Compose([
        transforms.ToTensor(),
        transforms.Normalize((0.1307,), (0.3081,))
    ])
    
    train_dataset = datasets.MNIST('../data', train=True, download=True, transform=transform)
    test_dataset = datasets.MNIST('../data', train=False, transform=transform)
    
    train_loader = DataLoader(train_dataset, batch_size=64, shuffle=True)
    test_loader = DataLoader(test_dataset, batch_size=1000, shuffle=False)
    
    # Instantiate model, optimizer
    model = SpiralWormClassifier(grid_size=10, path_steps=50).to(device)
    optimizer = optim.Adam(model.parameters(), lr=0.001)
    
    # Run a quick 2-epoch training/testing routine
    for epoch in range(1, 3):
        train(model, device, train_loader, optimizer, epoch)
        test(model, device, test_loader)
