import torch
import torch.nn as nn
import torch.nn.functional as F
import torch.optim as optim
from torch.utils.data import DataLoader
from torchvision import datasets, transforms
import matplotlib.pyplot as plt
import numpy as np

# ------------------------------------------------------------
# 1. Gravitational Superposition Classifier (GST)
# ------------------------------------------------------------
class GravitationalSuperpositionModel(nn.Module):
    """
    Embodiment of the Gravitational Superposition Theory:
    - 'matter' = real MNIST digits
    - 'core'   = learnable prototypes (positions + masses)
    - 'gravity'= attractive potential between input embedding and cores
    - 'superposition' = softmax over potentials
    - 'collapse' = hardmax (inference decision)
    """
    def __init__(self, embedding_dim=64, num_classes=10):
        super().__init__()
        # Feature extractor: CNN to map image into "space"
        self.feature_extractor = nn.Sequential(
            nn.Conv2d(1, 16, kernel_size=5, stride=1, padding=2),
            nn.ReLU(),
            nn.MaxPool2d(2),                     # 14x14
            nn.Conv2d(16, 32, kernel_size=5, stride=1, padding=2),
            nn.ReLU(),
            nn.MaxPool2d(2),                     # 7x7
            nn.Flatten(),
            nn.Linear(32 * 7 * 7, embedding_dim),
        )

        # Core: learnable positions (gravitational centers) and log-masses
        self.core_positions = nn.Parameter(torch.randn(num_classes, embedding_dim) * 0.1)
        self.log_masses = nn.Parameter(torch.zeros(num_classes))   # mass = exp(log_mass) > 0

        self.num_classes = num_classes
        self.eps = 1e-6

    def forward(self, x, return_embedding=False):
        # Embed input (the "matter sample")
        z = self.feature_extractor(x)                # (batch, embedding_dim)

        # Input mass: set to 1 for pure field-effect, or let it scale with ||z||
        input_mass = 1.0  # or torch.norm(z, dim=1, keepdim=True) to make input dynamic

        # Core masses (always positive)
        core_mass = torch.exp(self.log_masses)       # (num_classes,)

        # Squared distances between input and all cores
        # (batch, 1, dim) - (1, num_classes, dim) -> (batch, num_classes, dim)
        diff = z.unsqueeze(1) - self.core_positions.unsqueeze(0)
        dist_sq = diff.pow(2).sum(dim=2) + self.eps  # (batch, num_classes)

        # Gravitational logits: proportional to (m_in * m_c) / r^2
        # The "potential" is U = - (m_in * m_c) / r^2, so logits = -U
        logits = (input_mass * core_mass) / dist_sq   # (batch, num_classes)

        if return_embedding:
            return logits, z
        return logits

# ------------------------------------------------------------
# 2. Training and Evaluation Utilities
# ------------------------------------------------------------
def train_one_epoch(model, loader, optimizer, device):
    model.train()
    running_loss = 0.0
    correct = 0
    total = 0

    for images, labels in loader:
        images, labels = images.to(device), labels.to(device)
        for i in range(10):
            ids = np.random.randint(0,1024,100)
            optimizer.zero_grad()
            logits = model(images[ids])
            loss = F.cross_entropy(logits, labels[ids])
            loss.backward()
            optimizer.step()
            if i==0:
                running_loss += loss.item() * images.size(0)
                _, predicted = logits.max(1)
                total += labels[ids].size(0)
                correct += predicted.eq(labels[ids]).sum().item()

    return running_loss / total, correct / total

def evaluate(model, loader, device):
    model.eval()
    correct = 0
    total = 0
    with torch.no_grad():
        for images, labels in loader:
            images, labels = images.to(device), labels.to(device)
            logits = model(images)
            _, predicted = logits.max(1)
            total += labels.size(0)
            correct += predicted.eq(labels).sum().item()
    return correct / total

# ------------------------------------------------------------
# 3. Main: MNIST Training + Test
# ------------------------------------------------------------
def main():
    # Hyperparameters
    batch_size = 1024
    epochs = 10
    learning_rate = 1e-3
    embedding_dim = 64

    device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
    print(f"Using device: {device}")

    # Data
    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)

    train_loader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True, num_workers=2,drop_last=True)
    test_loader  = DataLoader(test_dataset, batch_size=batch_size, shuffle=False, num_workers=2)

    # Model, optimizer
    model = GravitationalSuperpositionModel(embedding_dim=embedding_dim, num_classes=10).to(device)
    optimizer = optim.Adam(model.parameters(), lr=learning_rate)

    # Training loop
    for epoch in range(1, epochs + 1):
        train_loss, train_acc = train_one_epoch(model, train_loader, optimizer, device)
        test_acc = evaluate(model, test_loader, device)
        print(f"Epoch {epoch:2d} | Loss: {train_loss:.4f} | Train Acc: {train_acc:.4f} | Test Acc: {test_acc:.4f}")

    # Final test accuracy
    final_acc = evaluate(model, test_loader, device)
    print(f"\nFinal Test Accuracy: {final_acc*100:.2f}%")

    # --------------------------------------------------------
    # 4. Visualising the "gravitational field" (embedding space)
    # --------------------------------------------------------
    # Project the core positions and some test samples using PCA
    from sklearn.decomposition import PCA

    model.eval()
    all_embeddings = []
    all_labels = []
    with torch.no_grad():
        for images, labels in test_loader:
            images = images.to(device)
            logits, z = model(images, return_embedding=True)
            all_embeddings.append(z.cpu().numpy())
            all_labels.append(labels.numpy())
    all_embeddings = np.concatenate(all_embeddings, axis=0)
    all_labels = np.concatenate(all_labels, axis=0)

    # Add core positions
    core_pos_np = model.core_positions.detach().cpu().numpy()

    # PCA on combined set
    pca = PCA(n_components=2)
    combined = np.vstack([all_embeddings, core_pos_np])
    proj = pca.fit_transform(combined)
    samples_proj = proj[:len(all_embeddings)]
    cores_proj = proj[len(all_embeddings):]

    # Plot
    plt.figure(figsize=(10, 8))
    scatter = plt.scatter(samples_proj[:, 0], samples_proj[:, 1], c=all_labels, cmap='tab10', alpha=0.3, s=10)
    plt.colorbar(scatter, ticks=range(10))
    # Cores as large stars with mass
    core_masses = torch.exp(model.log_masses).detach().cpu().numpy()
    for c in range(10):
        plt.scatter(cores_proj[c, 0], cores_proj[c, 1],
                    marker='*', s=200+200*core_masses[c], edgecolors='black',
                    c=plt.cm.tab10(c), label=f'Core {c} (m={core_masses[c]:.2f})')
    plt.title("Gravitational Superposition Embedding Space\nStars = core attractors (size ∝ mass)")
    plt.xlabel("PCA 1")
    plt.ylabel("PCA 2")
    plt.legend(bbox_to_anchor=(1.05, 1), loc='upper left')
    plt.tight_layout()
    plt.savefig("gst_embedding.png")
    plt.show()

if __name__ == "__main__":
    main()
