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

        # REDUCED WEIGHTS: Scale down CCT terms to prevent explosion
        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.chaos_threshold = 1.0
        self.bias_bound = 1.0
        self.phase_bound = 10.0
        self.efficiency_threshold = 0.5
        self.curvature_threshold = 5.0
        self.compression_target = 0.5

    def forward(self, outputs, targets, step=None):
        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

        L_S = torch.tensor(0.0, device=self.device)
        L_P = torch.tensor(0.0, device=self.device)
        L_Q = torch.tensor(0.0, device=self.device)
        L_T = torch.tensor(0.0, device=self.device)

        if len(activations) > 0:
            h_norm = torch.norm(activations[-1], dim=1).mean()
            law_dev = torch.abs(h_norm - 1.0)
            L_S = L_S - self.alpha[0] * law_dev

        if step is not None and step > 0 and self.prev_loss is not None:
            grad_norm = 0.0
            for p in self.model.parameters():
                if p.grad is not None:
                    grad_norm += p.grad.norm().item()
            grad_change = abs(grad_norm - (self.prev_loss / 1000.0)) # scaled compare
            L_S = L_S + self.alpha[1] * grad_change

        sym_loss = 0.0
        for w in weights:
            if w.size(0) == w.size(1):
                sym_loss += torch.norm(w - w.T) / w.numel()
        L_S = L_S - self.alpha[2] * sym_loss

        for a in activations:
            L_S = L_S + self.alpha[3] * torch.mean(F.relu(torch.abs(a) - self.phase_bound))

        for b in biases:
            L_S = L_S + self.alpha[4] * torch.mean(F.relu(torch.abs(b) - self.bias_bound))

        target_onehot = F.one_hot(targets, num_classes=num_classes).float()
        dist_to_attractor = torch.norm(probs - target_onehot, dim=1).mean()
        L_S = L_S - self.alpha[5] * dist_to_attractor

        if len(activations) > 0:
            inv = torch.mean(activations[-1], dim=1).sum()
            if hasattr(self, 'prev_inv'):
                inv_change = torch.abs(inv - self.prev_inv)
                L_S = L_S - self.alpha[6] * inv_change
            self.prev_inv = inv.detach()

        if len(self.prev_weights) > 0:
            w_prev = self.prev_weights[-1]
            w_change = 0.0
            for w_curr, w_prev_layer in zip(weights, w_prev):
                w_change += torch.norm(w_curr - w_prev_layer)
            L_S = L_S - self.alpha[7] * w_change
        self.prev_weights.append([w.clone().detach() for w in weights])

        entropy = -torch.sum(probs * log_probs, dim=1).mean()
        if hasattr(self, 'prev_entropy'):
            L_P = L_P - self.beta[0] * torch.abs(entropy - self.prev_entropy)
        self.prev_entropy = entropy.detach()

        h_curr = activations[-1] if activations else outputs
        if len(self.prev_hidden) >= 1:
            h_prev = self.prev_hidden[-1]
            if h_curr.size(0) == h_prev.size(0):
                L_P = L_P - self.beta[1] * torch.norm(h_curr - h_prev, dim=1).mean()
        
        if len(self.prev_hidden) >= 2:
            h_prev_old = self.prev_hidden[-2]
            if h_curr.size(0) == h_prev_old.size(0):
                L_P = L_P + self.beta[4] * torch.norm(h_curr - (h_prev_old * 0.9), dim=1).mean()
        
        self.prev_hidden.append(h_curr.detach())

        L_P = L_P - self.beta[2] * entropy
        L_P = L_P + self.beta[5] * F.relu(torch.var(outputs, dim=0).mean() - 10.0)
        L_P = L_P + self.beta[6] * F.kl_div(log_probs, target_onehot, reduction='batchmean')
        cond_entropy = -torch.mean(torch.gather(log_probs, 1, targets.unsqueeze(1)))
        L_P = L_P - self.beta[7] * cond_entropy

        mi = entropy - cond_entropy
        L_Q = L_Q - self.gamma[0] * mi

        H_T = -torch.mean(torch.sum(target_onehot * torch.log(target_onehot+1e-8), dim=1))
        if len(activations) > 0:
            h_last = activations[-1]
            with torch.no_grad():
                _, idx = torch.sort(h_last[:, 0], dim=0)
                mid = batch_size // 2
                cond_ent_h = 0.0
                for split in [idx[:mid], idx[mid:]]:
                    sub_targets = targets[split]
                    if len(sub_targets) > 0:
                        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
            L_Q = L_Q - self.gamma[1] * F.relu(collapse_pot)

        if len(activations) > 0:
            param_per_neuron = weights[-1].numel() / weights[-1].size(0) if weights else 1.0
            eff = (collapse_pot if 'collapse_pot' in locals() else 0.0) / param_per_neuron
            L_Q = L_Q + self.gamma[2] * F.relu(torch.tensor(self.efficiency_threshold - eff, device=self.device))

        probs_sorted, _ = torch.sort(probs, dim=1, descending=True)
        L_Q = L_Q - self.gamma[3] * (probs_sorted[:, 0] - probs_sorted[:, 1]).mean()
        L_Q = L_Q - self.gamma[4] * entropy
        L_Q = L_Q + self.gamma[5] * F.relu(cond_entropy - self.efficiency_threshold)

        L_T = L_T + self.delta[0] * ce_loss
        L_T = L_T - self.delta[3] * entropy
        conv_gap = torch.abs(mi - H_T)
        L_T = L_T + self.delta[4] * conv_gap

        input_bits = 3072 * 8
        comp_bits = sum(p.numel() for p in self.model.parameters()) * 32
        ratio = comp_bits / input_bits
        L_T = L_T + self.delta[5] * F.relu(torch.tensor(self.compression_target - ratio, device=self.device))
        L_T = L_T - self.delta[6] * conv_gap
        L_T = L_T - self.delta[7] * float(abs(len(weights) - 3))

        total = self.lambda_standard * ce_loss + L_S + L_P + L_Q + L_T
        self.prev_loss = total.detach().item()
        return total

def train_model(model, device, trainloader, optimizer, criterion, epochs=20):
    model.train()
    for epoch in range(epochs):
        running_loss, correct, total = 0.0, 0, 0
        for i, (inputs, targets) in enumerate(trainloader):
            inputs, targets = inputs.to(device), targets.to(device)
            optimizer.zero_grad()
            outputs = model(inputs)
            loss = criterion(outputs, targets, step=epoch*len(trainloader)+i)
            loss.backward()
            optimizer.step()
            running_loss += loss.item()
            _, predicted = outputs.max(1)
            total += targets.size(0)
            correct += predicted.eq(targets).sum().item()
            if i % 100 == 99:
                print(f'Epoch {epoch+1}, Batch {i+1}: Loss {running_loss/100:.4f}, Acc {100.*correct/total:.2f}%')
                running_loss = 0.0
        print(f'Epoch {epoch+1} finished. Accuracy: {100.*correct/total:.2f}%')

def test_model(model, device, testloader):
    model.eval()
    correct, total = 0, 0
    with torch.no_grad():
        for inputs, targets in testloader:
            inputs, targets = inputs.to(device), targets.to(device)
            outputs = model(inputs)
            _, predicted = outputs.max(1)
            total += targets.size(0)
            correct += predicted.eq(targets).sum().item()
    print(f'Test Accuracy: {100.*correct/total:.2f}%')

def main():
    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
    print(f'Using device: {device}')
    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, num_workers=2)
    testset = torchvision.datasets.CIFAR10(root='./data', train=False, download=True, transform=transform)
    testloader = torch.utils.data.DataLoader(testset, batch_size=100, shuffle=False, num_workers=2)
    model = MLP(input_size=3072, hidden_sizes=[512, 256, 128], num_classes=10).to(device)
    criterion = CCT32Loss(model, lambda_standard=1.0, device=device)
    # LOWER LEARNING RATE: From 0.001 to 0.0001 for better stability
    optimizer = optim.Adam(model.parameters(), lr=0.0001)
    print("Starting training with stabilized CCT-32 loss...")
    train_model(model, device, trainloader, optimizer, criterion, epochs=20)
    test_model(model, device, testloader)

if __name__ == '__main__':
    main()
