"""
============================================================
  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
from torch.utils.data import DataLoader
import torchvision
import torchvision.transforms as transforms
import numpy as np
from dataclasses import dataclass
from typing import List, Dict

# 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
    
@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
    entropy_threshold: float = 0.1
    orbit_iterations: int = 24
    orbit_force_scale: float = 0.01
    
class NBodyOptimizer:
    def __init__(self, parameters, config: GNBOFConfig):
        self.params = list(parameters)
        self.config = config
        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):
        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.fc1 = nn.Linear(28 * 28, latent_dim)
        self.classifier = 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 = x.view(x.size(0), -1)
        return F.relu(self.fc1(x))
    
    def forward(self, x):
        z = self.get_embedding(x)
        logits = self.classifier(z)
        attractor_logits = -torch.cdist(z, self.attractor_positions)
        return logits + attractor_logits * 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, total_updates, correct, total = 0.0, 0, 0, 0
        for batch_idx, (inputs, targets) in enumerate(train_loader):
            inputs, targets = inputs.to(device), targets.to(device)
            orbit_indices = torch.unique(targets).tolist()
            last_outputs = None

            for orbit_idx in orbit_indices:
                orbit_idx = int(orbit_idx)
                for _ in range(self.config.orbit_iterations):
                    outputs = self.model(inputs)
                    loss = F.cross_entropy(outputs, targets)
                    self.model.zero_grad()
                    loss.backward()

                    with torch.no_grad():
                        batch_axis = self.model.get_embedding(inputs).mean(dim=0)
                        orbit_force = self.optimizer.attractors[orbit_idx].position - batch_axis

                        for i, p in enumerate(self.model.parameters()):
                            if p.grad is None:
                                continue

                            if p is self.model.attractor_positions:
                                p.grad[orbit_idx] += orbit_force * self.config.orbit_force_scale

                            if 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()
                    total_updates += 1
                    last_outputs = outputs

            if last_outputs is not None:
                _, predicted = last_outputs.max(1)
                total += targets.size(0)
                correct += predicted.eq(targets).sum().item()

        self.optimizer.update_attractor_mass(self.model, train_loader)
        return total_loss / max(1, total_updates), 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['train_loss'].append(loss)
            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,
        orbit_iterations=24,
        orbit_force_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()
