import torch
import torch.nn as nn
import torch.nn.functional as F
import torch.optim as optim
import torchvision
import torchvision.transforms as transforms
import numpy as np
from collections import deque

# -------------------------------
# MLP Model for CIFAR-10
# -------------------------------
class MLP(nn.Module):
    def __init__(self, input_size=3072, hidden_sizes=[512, 256, 128], num_classes=10, dropout=0.2):
        super().__init__()
        layers = []
        prev = input_size
        for h in hidden_sizes:
            layers.append(nn.Linear(prev, h))
            layers.append(nn.ReLU())
            layers.append(nn.Dropout(dropout))
            prev = h
        layers.append(nn.Linear(prev, num_classes))
        self.net = nn.Sequential(*layers)
        self.activations = []
        self.weights = []
        self.biases = []
        self._register_hooks()

    def _register_hooks(self):
        def hook_fn(module, input, output):
            if isinstance(module, nn.Linear):
                self.activations.append(output.detach())
                self.weights.append(module.weight)
                self.biases.append(module.bias)
        for m in self.net:
            if isinstance(m, nn.Linear):
                m.register_forward_hook(hook_fn)

    def forward(self, x):
        self.activations.clear()
        self.weights.clear()
        self.biases.clear()
        x = x.view(x.size(0), -1)
        return self.net(x)

# -------------------------------
# CCT-32 Loss Implementation
# -------------------------------
class CCT32Loss(nn.Module):
    def __init__(self, model, lambda_standard=1.0, device='cuda'):
        super().__init__()
        self.model = model
        self.lambda_standard = lambda_standard
        self.device = device
        self.prev_weights = deque(maxlen=2)
        self.prev_hidden = deque(maxlen=5)
        self.prev_loss = None
        cct_scale = 0.001
        self.alpha = torch.ones(8, device=device) * cct_scale
        self.beta  = torch.ones(8, device=device) * cct_scale
        self.gamma = torch.ones(8, device=device) * cct_scale
        self.delta = torch.ones(8, device=device) * cct_scale
        self.phase_bound = 10.0
        self.bias_bound = 1.0
        self.efficiency_threshold = 0.5
        self.compression_target = 0.5

    def forward(self, outputs, targets, step=None, return_per_question=False):
        batch_size = outputs.size(0)
        num_classes = outputs.size(1)
        probs = F.softmax(outputs, dim=1)
        log_probs = F.log_softmax(outputs, dim=1)
        ce_loss = F.cross_entropy(outputs, targets)
        activations = self.model.activations
        weights = self.model.weights
        biases = self.model.biases
        target_onehot = F.one_hot(targets, num_classes=num_classes).float()
        entropy = -torch.sum(probs * log_probs, dim=1).mean()
        cond_entropy = -torch.mean(torch.gather(log_probs, 1, targets.unsqueeze(1)))
        mi = entropy - cond_entropy
        H_T = -torch.mean(torch.sum(target_onehot * torch.log(target_onehot+1e-8), dim=1))

        qs = []
        # --- Layer 1 (Stationary) ---
        qs.append(-self.alpha[0] * torch.abs(torch.norm(activations[-1], dim=1).mean() - 1.0) if activations else torch.tensor(0.0, device=self.device))
        qs.append(self.alpha[1] * torch.tensor(0.0, device=self.device))
        sym_val = sum(torch.norm(w - w.T)/w.numel() for w in weights if w.size(0)==w.size(1)) if weights else 0.0
        qs.append(-self.alpha[2] * torch.tensor(sym_val, device=self.device))
        qs.append(self.alpha[3] * sum(torch.mean(F.relu(torch.abs(a) - self.phase_bound)) for a in activations))
        qs.append(self.alpha[4] * sum(torch.mean(F.relu(torch.abs(b) - self.bias_bound)) for b in biases))
        qs.append(-self.alpha[5] * torch.norm(probs - target_onehot, dim=1).mean())
        qs.append(torch.tensor(0.0, device=self.device))
        qs.append(torch.tensor(0.0, device=self.device))

        # --- Layer 2 (Probability) ---
        qs.append(torch.tensor(0.0, device=self.device))
        h_curr = activations[-1] if activations else outputs
        q10 = -self.beta[1] * torch.norm(h_curr - self.prev_hidden[-1], dim=1).mean() if len(self.prev_hidden) > 0 and h_curr.size(0) == self.prev_hidden[-1].size(0) else torch.tensor(0.0, device=self.device)
        qs.append(q10)
        qs.append(-self.beta[2] * entropy)
        qs.append(torch.tensor(0.0, device=self.device))
        qs.append(torch.tensor(0.0, device=self.device))
        q14 = self.beta[4] * torch.norm(h_curr - (self.prev_hidden[-2] * 0.9), dim=1).mean() if len(self.prev_hidden) >= 2 and h_curr.size(0) == self.prev_hidden[-2].size(0) else torch.tensor(0.0, device=self.device)
        qs.append(q14)
        qs.append(self.beta[5] * F.relu(torch.var(outputs, dim=0).mean() - 10.0))
        qs.append(self.beta[6] * F.kl_div(log_probs, target_onehot, reduction='batchmean'))
        qs.append(-self.beta[7] * cond_entropy)

        # --- Layer 3 (Collapse) ---
        qs.append(-self.gamma[0] * mi)
        collapse_pot = torch.tensor(0.0, device=self.device)
        if activations:
             with torch.no_grad():
                mid = batch_size // 2
                sub_targets = targets[:mid]
                p = torch.bincount(sub_targets, minlength=num_classes).float() / (len(sub_targets) + 1e-8)
                cond_ent_h = -torch.sum(p * torch.log(p + 1e-8))
                collapse_pot = H_T - cond_ent_h
        qs.append(-self.gamma[1] * F.relu(collapse_pot))
        eff = collapse_pot / (weights[-1].numel() / weights[-1].size(0) if weights else 1.0)
        qs.append(self.gamma[2] * F.relu(torch.tensor(self.efficiency_threshold - eff.item(), device=self.device)))
        ps, _ = torch.sort(probs, dim=1, descending=True)
        qs.append(-self.gamma[3] * (ps[:, 0] - ps[:, 1]).mean())
        qs.append(-self.gamma[4] * entropy)
        qs.append(self.gamma[5] * F.relu(cond_entropy - self.efficiency_threshold))
        qs.append(torch.tensor(0.0, device=self.device))
        qs.append(torch.tensor(0.0, device=self.device))

        # --- Layer 4 (Taylor) ---
        qs.append(self.delta[0] * ce_loss)
        qs.append(torch.tensor(0.0, device=self.device))
        qs.append(torch.tensor(0.0, device=self.device))
        qs.append(-self.delta[3] * entropy)
        qs.append(self.delta[4] * torch.abs(mi - H_T))
        num_params = sum(p.numel() for p in self.model.parameters())
        ratio = (num_params * 32) / (3072 * 8)
        qs.append(self.delta[5] * F.relu(torch.tensor(self.compression_target - ratio, device=self.device)))
        qs.append(-self.delta[6] * torch.abs(mi - H_T))
        qs.append(-self.delta[7] * torch.tensor(float(abs(len(weights) - 3)), device=self.device))

        self.prev_hidden.append(h_curr.detach())
        if return_per_question:
            return qs
        total = self.lambda_standard * ce_loss + sum(qs)
        self.prev_loss = total.detach().item()
        return total

def train_model_sequential(model, device, trainloader, optimizer, criterion, epochs=10):
    model.train()
    for epoch in range(epochs):
        correct, total = 0, 0
        for batch_idx, (inputs, targets) in enumerate(trainloader):
            inputs, targets = inputs.to(device), targets.to(device)
            outputs = model(inputs)
            base_ce = F.cross_entropy(outputs, targets)
            losses_per_q = criterion(outputs, targets, return_per_question=True)
            for q_loss in losses_per_q:
                optimizer.zero_grad()
                combined = (0.01 * base_ce) + q_loss
                combined.backward(retain_graph=True)
                optimizer.step()
            with torch.no_grad():
                final_out = model(inputs)
                _, predicted = final_out.max(1)
                total += targets.size(0)
                correct += predicted.eq(targets).sum().item()
            if batch_idx % 100 == 99:
                print(f"Epoch {epoch+1}, Batch {batch_idx+1}: Acc {100.*correct/total:.2f}%")
        print(f"Epoch {epoch+1} Sequential finished. Accuracy: {100.*correct/total:.2f}%")

def run_sequential_experiment():
    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
    transform = transforms.Compose([transforms.ToTensor(), transforms.Normalize((0.5,0.5,0.5), (0.5,0.5,0.5))])
    trainset = torchvision.datasets.CIFAR10(root='./data', train=True, download=True, transform=transform)
    trainloader = torch.utils.data.DataLoader(trainset, batch_size=64, shuffle=True)
    model = MLP().to(device)
    criterion = CCT32Loss(model, device=device)
    optimizer = optim.Adam(model.parameters(), lr=0.0001)

    def train_model_fixed_sequential(model, device, trainloader, optimizer, criterion, epochs=5):
        model.train()
        for epoch in range(epochs):
            correct, total = 0, 0
            for batch_idx, (inputs, targets) in enumerate(trainloader):
                inputs, targets = inputs.to(device), targets.to(device)

                # 1. Base Loss Update
                optimizer.zero_grad()
                outputs = model(inputs)
                base_ce = F.cross_entropy(outputs, targets)
                base_ce.backward()
                optimizer.step()

                # 2. Sequential CCT Question Updates
                # We re-run the forward pass for each question (or group) 
                # because optimizer.step() invalidates the previous graph.
                for i in range(32):
                    optimizer.zero_grad()
                    current_outputs = model(inputs)
                    losses_per_q = criterion(current_outputs, targets, return_per_question=True)
                    
                    q_loss = losses_per_q[i]
                    if isinstance(q_loss, torch.Tensor) and q_loss.requires_grad:
                        q_loss.backward()
                        optimizer.step()

                with torch.no_grad():
                    final_out = model(inputs)
                    _, predicted = final_out.max(1)
                    total += targets.size(0)
                    correct += predicted.eq(targets).sum().item()

                if batch_idx % 100 == 99:
                    print(f"Epoch {epoch+1}, Batch {batch_idx+1}: Acc {100.*correct/total:.2f}%")
            print(f"Epoch {epoch+1} Sequential finished. Accuracy: {100.*correct/total:.2f}%")

    print("--- Starting Separate Sequential CCT-32 Training ---")
    train_model_fixed_sequential(model, device, trainloader, optimizer, criterion, epochs=5)

run_sequential_experiment()


