"""
============================================================
  10-ATTACTOR N-BODY MNIST CLASSIFIER (CUDA ENHANCED)
  Based on Conditional Collapse Theory (CCT) + 
  Gravitational N-Body Optimization Framework (GNBOF)
============================================================
"""

import torch
import torch.nn as nn
import torch.nn.functional as F
import torch.optim as optim
from torch.utils.data import DataLoader, TensorDataset
import torchvision
import torchvision.transforms as transforms
import numpy as np
import matplotlib.pyplot as plt
from matplotlib.patches import Circle
import time
import math
from dataclasses import dataclass, field
from typing import List, Dict, Tuple, Optional

# Setup Device
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
print(f"[SYSTEM] Using device: {device}")

@dataclass
class AttractorConfig:
    digit: int
    mass: float = 1.0
    position: torch.Tensor = None
    is_black_hole: bool = False
    event_horizon_radius: float = 0.0
    orbital_energy: float = 0.0
    
@dataclass
class GNBOFConfig:
    G: float = 0.1
    beta: float = 0.9
    singularity_threshold: float = 5.0
    hawking_rate: float = 0.01
    escape_velocity_scale: float = 1.0
    lagrange_threshold: float = 0.5
    periodicity_tolerance: float = 1e-4
    compute_budget: float = 1e6
    entropy_threshold: float = 0.1
    
class NBodyOptimizer:
    def __init__(self, parameters, config: GNBOFConfig, num_attractors=10):
        self.params = list(parameters)
        self.config = config
        self.num_attractors = num_attractors
        self.velocity = [torch.zeros_like(p).to(device) for p in self.params]
        self.position_history = []
        self.attractors: List[AttractorConfig] = []
        self.current_entropy = float('inf')
        self.total_energy_spent = 0.0
        self.collapse_achieved = False
        self.black_holes: List[int] = []
        
    def initialize_attractors(self, model, dataloader):
        print("\n[GNBOF] Initializing 10 digit attractors on device...")
        model.eval()
        all_embeddings = []
        all_labels = []
        max_samples = 500
        collected = 0
        
        with torch.no_grad():
            for inputs, targets in dataloader:
                if collected >= max_samples:
                    break
                inputs = inputs.to(device)
                if hasattr(model, 'get_embedding'):
                    emb = model.get_embedding(inputs)
                else:
                    emb = inputs.view(inputs.size(0), -1)

                take = min(emb.size(0), max_samples - collected)
                all_embeddings.append(emb[:take].cpu())
                all_labels.append(targets[:take].cpu())
                collected += take
        
        all_embeddings = torch.cat(all_embeddings, dim=0)
        all_labels = torch.cat(all_labels, dim=0)
        
        self.attractors = []
        for digit in range(10):
            mask = all_labels == digit
            if mask.sum() > 0:
                center = all_embeddings[mask].mean(dim=0).to(device)
            else:
                center = (torch.randn(all_embeddings.shape[1]) * 0.1).to(device)
            
            self.attractors.append(AttractorConfig(
                digit=digit, mass=1.0,
                position=center.clone().detach().requires_grad_(True)
            ))
        
        return self.attractors
    
    def compute_gravitational_force(self, position: torch.Tensor) -> torch.Tensor:
        total_force = torch.zeros_like(position)
        for i, attractor in enumerate(self.attractors):
            if i in self.black_holes:
                direction = position - attractor.position
                distance = torch.norm(direction) + 1e-6
                repulsive = -self.config.G * attractor.mass / (distance ** 2 + 1.0)
                total_force += repulsive * (direction / distance)
            else:
                direction = attractor.position - position
                distance = torch.norm(direction) + 1e-6
                attraction = self.config.G * attractor.mass / (distance ** 2 + 1e-6)
                total_force += attraction * (direction / distance)
        return total_force

    def update_attractor_mass(self, model, dataloader, epoch):
        model.eval()
        digit_losses = {i: [] for i in range(10)}
        with torch.no_grad():
            for inputs, targets in dataloader:
                inputs, targets = inputs.to(device), targets.to(device)
                outputs = model(inputs)
                for d in range(10):
                    mask = targets == d
                    if mask.sum() > 0:
                        digit_losses[d].append(F.cross_entropy(outputs[mask], targets[mask]).item())
        
        for i, attractor in enumerate(self.attractors):
            if len(digit_losses[i]) > 0:
                avg_loss = np.mean(digit_losses[i])
                new_mass = 1.0 / (1.0 + avg_loss)
                if new_mass > self.config.singularity_threshold and not attractor.is_black_hole:
                    attractor.is_black_hole = True
                    self.black_holes.append(i)
                attractor.mass = new_mass

    def get_cct_entropy(self, model, dataloader) -> float:
        model.eval()
        all_probs = []
        with torch.no_grad():
            for inputs, _ in dataloader:
                inputs = inputs.to(device)
                outputs = model(inputs)
                all_probs.append(F.softmax(outputs, dim=1))
        
        all_probs = torch.cat(all_probs, dim=0)
        entropy = -torch.sum(all_probs * torch.log(all_probs + 1e-10)) / all_probs.size(0)
        self.current_entropy = entropy.item()
        return self.current_entropy

    def check_collapse(self, model, dataloader) -> bool:
        entropy = self.get_cct_entropy(model, dataloader)
        if entropy < self.config.entropy_threshold:
            self.collapse_achieved = True
            return True
        return False

class AttractorMNIST(nn.Module):
    def __init__(self, latent_dim=64, num_attractors=10):
        super().__init__()
        self.latent_dim = latent_dim
        self.num_attractors = num_attractors
        self.conv1 = nn.Conv2d(1, 32, 3, padding=1)
        self.conv2 = nn.Conv2d(32, 64, 3, padding=1)
        self.pool = nn.MaxPool2d(2, 2)
        self.fc1 = nn.Linear(64 * 7 * 7, latent_dim)
        self.attractor_projection = nn.Linear(latent_dim, num_attractors)
        self.attractor_positions = nn.Parameter(torch.randn(num_attractors, latent_dim) * 0.1)
        
    def get_embedding(self, x):
        x = self.pool(F.relu(self.conv1(x)))
        x = self.pool(F.relu(self.conv2(x)))
        x = x.view(x.size(0), -1)
        return F.relu(self.fc1(x))
    
    def forward(self, x):
        z = self.get_embedding(x)
        logits = self.attractor_projection(z)
        attractor_logits = []
        for i in range(self.num_attractors):
            dist = torch.norm(z - self.attractor_positions[i], dim=1, keepdim=True)
            attractor_logits.append(-dist)
        return logits + torch.cat(attractor_logits, dim=1) * 0.1

class CCTNBodyTrainer:
    def __init__(self, model, optimizer, config: GNBOFConfig):
        self.model = model.to(device)
        self.optimizer = optimizer
        self.config = config
        self.history = {'epoch': [], 'train_loss': [], 'test_acc': [], 'entropy': [], 'attractor_masses': []}

    def train_epoch(self, train_loader, epoch):
        self.model.train()
        total_loss, correct, total = 0, 0, 0
        for batch_idx, (inputs, targets) in enumerate(train_loader):
            inputs, targets = inputs.to(device), targets.to(device)
            outputs = self.model(inputs)
            loss = F.cross_entropy(outputs, targets)
            self.model.zero_grad()
            loss.backward()
            with torch.no_grad():
                for i, p in enumerate(self.model.parameters()):
                    if p.grad is not None and i < len(self.optimizer.velocity):
                        v = self.optimizer.velocity[i]
                        v.mul_(self.config.beta).add_(p.grad, alpha=1 - self.config.beta)
                        p.sub_(v, alpha=self.config.escape_velocity_scale)
            total_loss += loss.item()
            _, predicted = outputs.max(1)
            total += targets.size(0)
            correct += predicted.eq(targets).sum().item()
        self.optimizer.update_attractor_mass(self.model, train_loader, epoch)
        return total_loss / len(train_loader), 100. * correct / total, self.optimizer.check_collapse(self.model, train_loader)

    def test(self, test_loader):
        self.model.eval()
        correct, total = 0, 0
        with torch.no_grad():
            for inputs, targets in test_loader:
                inputs, targets = inputs.to(device), targets.to(device)
                outputs = self.model(inputs)
                _, predicted = outputs.max(1)
                total += targets.size(0)
                correct += predicted.eq(targets).sum().item()
        return 100. * correct / total

    def train(self, train_loader, test_loader, epochs=10):
        self.optimizer.initialize_attractors(self.model, train_loader)
        for epoch in range(epochs):
            loss, acc, collapsed = self.train_epoch(train_loader, epoch)
            test_acc = self.test(test_loader)
            entropy = self.optimizer.get_cct_entropy(self.model, test_loader)
            self.history['epoch'].append(epoch)
            self.history['test_acc'].append(test_acc)
            self.history['entropy'].append(entropy)
            self.history['attractor_masses'].append([a.mass for a in self.optimizer.attractors])
            print(f"[Epoch {epoch+1}] Acc: {test_acc:.2f}%, Entropy: {entropy:.4f}")
            if collapsed: break
        return self.history

def main():
    transform = transforms.Compose([transforms.ToTensor(), transforms.Normalize((0.5,), (0.5,))])
    train_dataset = torchvision.datasets.MNIST(root='./data', train=True, download=True, transform=transform)
    test_dataset = torchvision.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=128, shuffle=False)
    model = AttractorMNIST().to(device)
    config = GNBOFConfig(G=0.05, beta=0.9, singularity_threshold=3.0, entropy_threshold=0.3, escape_velocity_scale=0.01)
    optimizer = NBodyOptimizer(model.parameters(), config)
    trainer = CCTNBodyTrainer(model, optimizer, config)
    history = trainer.train(train_loader, test_loader)
    return model, history

if __name__ == '__main__':
    model, history = main()