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

# 1. Define Architecture
class MNISTNet(nn.Module):
    def __init__(self):
        super(MNISTNet, self).__init__()
        self.fc1 = nn.Linear(28 * 28, 128)
        self.fc2 = nn.Linear(128, 64)
        self.fc3 = nn.Linear(64, 10)

    def forward(self, x):
        x = x.view(-1, 28 * 28)
        x = F.relu(self.fc1(x))
        x = F.relu(self.fc2(x))
        return self.fc3(x)

# 2. Setup Data and Train Model
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
transform = transforms.Compose([transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,))])

train_loader = DataLoader(datasets.MNIST('../data', train=True, download=True, transform=transform), batch_size=64, shuffle=True)
test_loader = DataLoader(datasets.MNIST('../data', train=False, transform=transform), batch_size=1000, shuffle=False)

model = MNISTNet().to(device)
optimizer = optim.Adam(model.parameters(), lr=0.003)

print("Training baseline model...")
model.train()
for epoch in range(2):
    for data, target in 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()

# 3. Compute Baseline Probabilities AND Baseline Accuracy
model.eval()
baseline_logits = []
correct_baseline = 0
total_samples = 0

with torch.no_grad():
    for data, target in test_loader:
        data, target = data.to(device), target.to(device)
        logits = model(data)
        baseline_logits.append(logits)
        
        preds = logits.argmax(dim=-1)
        correct_baseline += preds.eq(target).sum().item()
        total_samples += target.size(0)

baseline_logits = torch.cat(baseline_logits, dim=0)
p_S = F.softmax(baseline_logits, dim=-1)
baseline_accuracy = (correct_baseline / total_samples) * 100

print(f"\nBaseline Test Accuracy: {baseline_accuracy:.2f}%")
print("---" * 15)

# 4. Compute Informational Gravity and Accuracy Drop per Neuron
print("Measuring Informational Gravity vs. Test Accuracy Drop per hidden neuron in fc2...")
results = []

def kl_divergence(p, q, eps=1e-10):
    return torch.sum(p * (torch.log(p + eps) - torch.log(q + eps)), dim=-1).mean().item()

with torch.no_grad():
    for neuron_idx in range(64):
        # Save weights
        orig_weight = model.fc2.weight.data[neuron_idx].clone()
        orig_bias = model.fc2.bias.data[neuron_idx].clone()
        
        # Remove component
        model.fc2.weight.data[neuron_idx] = 0.0
        model.fc2.bias.data[neuron_idx] = 0.0
        
        # Evaluate perturbed model for both KL and Accuracy
        perturbed_logits = []
        correct_perturbed = 0
        
        for data, target in test_loader:
            data, target = data.to(device), target.to(device)
            logits = model(data)
            perturbed_logits.append(logits)
            
            preds = logits.argmax(dim=-1)
            correct_perturbed += preds.eq(target).sum().item()
            
        perturbed_logits = torch.cat(perturbed_logits, dim=0)
        p_remove = F.softmax(perturbed_logits, dim=-1)
        
        # Calculations
        g_info = kl_divergence(p_S, p_remove)
        perturbed_accuracy = (correct_perturbed / total_samples) * 100
        acc_drop = baseline_accuracy - perturbed_accuracy
        
        results.append((neuron_idx, g_info, perturbed_accuracy, acc_drop))
        
        # Restore weights
        model.fc2.weight.data[neuron_idx] = orig_weight
        model.fc2.bias.data[neuron_idx] = orig_bias

# 5. Report Top 10 Most Critical "Load-Bearing" Neurons Found
print("\n--- Structural Analysis Results (Top 10 Most Significant Disruptions) ---")
for idx, g, new_acc, drop in sorted(results, key=lambda x: x[1], reverse=True)[:10]:
    if g > 0.1:
        role = "LOAD-BEARING"
    elif g > 0.01:
        role = "ENHANCING"
    else:
        role = "DECORATIVE"
        
    print(f"Neuron {idx:02d} | Pull ($\mathcal{G}$): {g:.5f} | New Acc: {new_acc:.2f}% | Drop: -{drop:.2f}% | Class: {role}")
