import torch
import torch.nn as nn
import torch.optim as optim
from torchvision import datasets, transforms
from torch.utils.data import DataLoader
from collections import deque
import math

# ---------------------------------------------------------
# 1. Custom Anti-Overfitting Geometric Optimizer
# ---------------------------------------------------------
class ScaledGeometricOptimizer(optim.Optimizer):
    def __init__(self, params, lr=1e-3, alpha=0.6, drift_coef=0.7, diff_coef=0.15, memory_len=8):
        defaults = dict(lr=lr, alpha=alpha, drift_coef=drift_coef, diff_coef=diff_coef, memory_len=memory_len)
        super(ScaledGeometricOptimizer, self).__init__(params, defaults)
        
        for group in self.param_groups:
            for p in group['params']:
                state = self.state[p]
                state['history'] = deque(maxlen=memory_len)
                state['curvature_buffer'] = torch.zeros_like(p.data)

    @torch.no_grad()
    def step(self, closure=None):
        loss = None
        for group in self.param_groups:
            lr = group['lr']
            alpha = group['alpha']
            drift_coef = group['drift_coef']
            diff_coef = group['diff_coef']
            
            for p in group['params']:
                if p.grad is None:
                    continue
                
                grad = p.grad.clone()
                state = self.state[p]
                
                # 1. Hessian Spectrum Diamond (Landscape Normalization)
                state['curvature_buffer'].mul_(0.95).add_(grad.pow(2), alpha=0.05)
                diamond_scale = 1.0 / (torch.sqrt(state['curvature_buffer']) + 1e-6)
                scaled_grad = grad * diamond_scale

                # 2. Wavelet Packet Tree (Multi-Scale Noise Filtering)
                coarse_trend = torch.mean(scaled_grad) * torch.ones_like(scaled_grad)
                fine_details = scaled_grad - coarse_trend
                filtered_grad = coarse_trend + (0.15 * fine_details)  # Suppress leaf-level structural noise

                # 3. Itô Strip (Drift vs. Continuous Exploration)
                drift = filtered_grad
                diffusion = torch.randn_like(grad) * diff_coef
                ito_step = (drift_coef * drift) + ((1.0 - drift_coef) * diffusion)

                # 4. Caputo Fractional Arc (Power-Law Path Integration)
                state['history'].append(ito_step)
                
                fractional_step = torch.zeros_like(grad)
                h_len = len(state['history'])
                for i, past_update in enumerate(state['history']):
                    k = h_len - i
                    weight = 1.0 / math.pow(k, 1.0 - alpha)
                    fractional_step.add_(past_update, alpha=weight)
                fractional_step.div_(h_len)

                # Execute parameter shift across the calculated trajectory
                p.add_(fractional_step, alpha=-lr)
        return loss

# ---------------------------------------------------------
# 2. Architecture & Data Processing Configuration
# ---------------------------------------------------------
class MNIST_MLP(nn.Module):
    def __init__(self):
        super().__init__()
        self.net = nn.Sequential(
            nn.Flatten(),
            nn.Linear(28*28, 256),
            nn.ReLU(),
            nn.Linear(256, 10)
        )
    def forward(self, x): 
        return self.net(x)

def main():
    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
    print(f"Executing environment pipeline on device context: {device}")
    
    transform = transforms.Compose([
        transforms.ToTensor(), 
        transforms.Normalize((0.1307,), (0.3081,))
    ])
    
    train_dataset = datasets.MNIST(root='../data', train=True, download=True, transform=transform)
    test_dataset = datasets.MNIST(root='../data', train=False, download=True, transform=transform)
    
    # Dynamic calculation to train on the FULL dataset using exactly 100 iterations per epoch
    # 60,000 samples / 100 iterations = Batch Size of 600
    batch_size = len(train_dataset) // 100 
    
    train_loader = DataLoader(train_dataset, batch_size=100, shuffle=True, drop_last=True)
    test_loader = DataLoader(test_dataset, batch_size=1000, shuffle=False)

    model = MNIST_MLP().to(device)
    criterion = nn.CrossEntropyLoss()
    
    # Custom Multi-Point Derivative Optimizer
    optimizer = ScaledGeometricOptimizer(
        model.parameters(), 
        lr=0.005, 
        diff_coef=0.0001, #0.18, 
        alpha=0.001 #0.55
    )

    epochs = 3
    print(f"Beginning full-scale training processing {len(train_dataset)} train samples over {epochs} Epochs...")
    print(f"Configured layout: {len(train_loader)} steps per epoch (Batch size: {batch_size}).\n")

    for epoch in range(1, epochs + 1):
        model.train()
        running_loss = 0.0
        
        for batch_idx, (data, target) in enumerate(train_loader):
            data, target = data.to(device), target.to(device)
            for i in range(20):
                optimizer.zero_grad()
                output = model(data)
                loss = criterion(output, target)
                loss.backward()
                optimizer.step()
                if i==0:
                    running_loss += loss.item()
                    
                    if batch_idx % 20 == 0:
                        print(f"Epoch [{epoch}/{epochs}] | Step [{batch_idx}] | Current Step Loss: {loss.item():.4f}")

        # Evaluation Loop across the FULL Test Set
        model.eval()
        test_loss = 0.0
        correct = 0
        total = 0
        
        with torch.no_grad():
            for test_images, test_labels in test_loader:
                test_images, test_labels = test_images.to(device), test_labels.to(device)
                preds = model(test_images)
                test_loss += criterion(preds, test_labels).item()
                correct += (preds.argmax(dim=1) == test_labels).sum().item()
                total += test_labels.size(0)

        avg_train_loss = running_loss / len(train_loader)
        avg_test_loss = test_loss / len(test_loader)
        accuracy = 100. * correct / total
        
        print(f"\n--> Epoch {epoch} Final Evaluation Summary:")
        print(f"    Average Training Loss   : {avg_train_loss:.4f}")
        print(f"    Full Test Set Loss      : {avg_test_loss:.4f}")
        print(f"    Full Test Set Accuracy  : {correct}/{total} ({accuracy:.2f}%)")
        print("-" * 65 + "\n")

if __name__ == '__main__':
    main()
