import torch
import torch.nn as nn
import torch.optim as optim
import torch.nn.functional as F
import torchvision
import torchvision.transforms as transforms
import numpy as np
import copy
import math
from collections import deque
from torch.autograd import grad

# -------------------------------
# 1. CNN Model for CIFAR10
# -------------------------------
class CIFAR10CNN(nn.Module):
    """Simple CNN for CIFAR10 (approx 1.2M parameters)"""
    def __init__(self):
        super().__init__()
        self.conv1 = nn.Conv2d(3, 32, 3, padding=1)
        self.bn1 = nn.BatchNorm2d(32)
        self.conv2 = nn.Conv2d(32, 64, 3, padding=1)
        self.bn2 = nn.BatchNorm2d(64)
        self.conv3 = nn.Conv2d(64, 128, 3, padding=1)
        self.bn3 = nn.BatchNorm2d(128)
        self.pool = nn.MaxPool2d(2, 2)
        self.fc1 = nn.Linear(128 * 4 * 4, 256)
        self.dropout = nn.Dropout(0.3)
        self.fc2 = nn.Linear(256, 10)

    def forward(self, x):
        x = self.pool(F.relu(self.bn1(self.conv1(x))))
        x = self.pool(F.relu(self.bn2(self.conv2(x))))
        x = self.pool(F.relu(self.bn3(self.conv3(x))))
        x = x.view(-1, 128 * 4 * 4)
        x = F.relu(self.fc1(x))
        x = self.dropout(x)
        x = self.fc2(x)
        return x

# -------------------------------
# 2. YIELD Replica Class
# -------------------------------
class YieldReplica:
    """
    A single solver replica with its own model, optimizer, and state.
    Tracks yield ratio (entropy reduction per work) and supports second-order corrections.
    """
    def __init__(self, model_fn, optimizer_fn, hyperparams, device):
        self.device = device
        self.model = model_fn().to(device)
        self.optimizer = optimizer_fn(self.model.parameters(), **hyperparams)
        self.hyperparams = hyperparams.copy()
        self.work = 0.0
        self.init_entropy = None
        self.current_entropy = None
        self.yield_ratio = 0.0
        self.alive = True
        self.saddle_escape_attempts = 0
        self._update_entropy()
        self.init_entropy = self.current_entropy

    def _update_entropy(self):
        total_variance = 0.0
        n_params = 0
        for p in self.model.parameters():
            if p.requires_grad:
                var = p.data.var().item()
                total_variance += var * p.numel()
                n_params += p.numel()
        if total_variance > 0:
            entropy = 0.5 * math.log(2 * math.pi * math.e * (total_variance / n_params))
        else:
            entropy = 0.0
        self.current_entropy = entropy

    def compute_yield(self, work_delta):
        self.work += work_delta
        self._update_entropy()
        entropy_reduction = self.init_entropy - self.current_entropy
        if self.work > 0:
            self.yield_ratio = entropy_reduction / self.work
        else:
            self.yield_ratio = 0.0
        return self.yield_ratio

    def second_order_correction(self, loss, max_iter=10, tolerance=1e-6):
        total_grad_norm = 0.0
        for p in self.model.parameters():
            if p.grad is not None:
                total_grad_norm += p.grad.norm().item()**2
        total_grad_norm = math.sqrt(total_grad_norm)
        if total_grad_norm > 1e-3:
            return False

        params = list(self.model.parameters())
        v = [torch.randn_like(p) for p in params]
        norm_v = math.sqrt(sum(p.norm().item()**2 for p in v))
        v = [p / norm_v for p in v]

        def hvp(v):
            grad_loss = torch.autograd.grad(loss, params, create_graph=True, retain_graph=True)
            grad_dot_v = sum(torch.sum(g * v_i) for g, v_i in zip(grad_loss, v))
            hvp_result = torch.autograd.grad(grad_dot_v, params, retain_graph=True)
            return hvp_result

        eigval_old = 0.0
        for _ in range(max_iter):
            hv = hvp(v)
            num = sum(torch.sum(h_i * v_i) for h_i, v_i in zip(hv, v))
            denom = sum(torch.sum(v_i * v_i) for v_i in v)
            eigval = num / (denom + 1e-9)
            norm_hv = math.sqrt(sum(h_i.norm().item()**2 for h_i in hv))
            if norm_hv > tolerance:
                v = [h_i / norm_hv for h_i in hv]
            if abs(eigval - eigval_old) < tolerance:
                break
            eigval_old = eigval

        if abs(eigval) < 1e-4:
            alpha = 0.1
            with torch.no_grad():
                for p, step in zip(params, v):
                    p.add_(step, alpha=alpha)
            self.saddle_escape_attempts += 1
            return True
        return False

# -------------------------------
# 3. YIELD Training Loop
# -------------------------------
class YIELDTrainer:
    def __init__(self, model_fn, replicas_config, device, yield_threshold=0.01,
                 curvature_noise_gamma=100.0, saddle_escape_freq=50):
        self.device = device
        self.model_fn = model_fn
        self.replicas = []
        for cfg in replicas_config:
            rep = YieldReplica(model_fn, cfg['optimizer_fn'], cfg['hyperparams'], device)
            self.replicas.append(rep)
        self.yield_threshold = yield_threshold
        self.curvature_noise_gamma = curvature_noise_gamma
        self.saddle_escape_freq = saddle_escape_freq
        self.iteration = 0

    def _compute_curvature_trace_estimate(self, model, loss, num_samples=5):
        params = list(model.parameters())
        trace_estimate = 0.0
        for _ in range(num_samples):
            v = [torch.randn_like(p) for p in params]
            grad_loss = torch.autograd.grad(loss, params, create_graph=True, retain_graph=True)
            grad_dot_v = sum(torch.sum(g * v_i) for g, v_i in zip(grad_loss, v))
            hv = torch.autograd.grad(grad_dot_v, params, retain_graph=True)
            trace_estimate += sum(torch.sum(h_i * v_i) for h_i, v_i in zip(hv, v))
        return trace_estimate / num_samples

    def _adaptive_noise(self, model, loss, base_lr, curvature_trace):
        T = max(0.01, math.exp(-curvature_trace / self.curvature_noise_gamma))
        noise_scale = math.sqrt(2 * base_lr * T)
        if noise_scale > 0:
            for p in model.parameters():
                if p.grad is not None:
                    noise = torch.randn_like(p.grad) * noise_scale
                    p.grad.add_(noise)

    def collapse_low_yield_replicas(self):
        alive_replicas = [r for r in self.replicas if r.alive and r.yield_ratio >= self.yield_threshold]
        if len(alive_replicas) == 0 and len(self.replicas) > 0:
            best = max(self.replicas, key=lambda r: r.yield_ratio)
            alive_replicas = [best]
        self.replicas = alive_replicas
        if len(self.replicas) == 1:
            best = self.replicas[0]
            new_hyperparams = best.hyperparams.copy()
            if 'lr' in new_hyperparams:
                new_hyperparams['lr'] *= np.random.uniform(0.5, 1.5)
            new_rep = YieldReplica(self.model_fn, optim.SGD, new_hyperparams, self.device)
            new_rep.model.load_state_dict(copy.deepcopy(best.model.state_dict()))
            self.replicas.append(new_rep)

    def train_epoch(self, trainloader, epoch):
        for replica in self.replicas:
            model = replica.model
            optimizer = replica.optimizer
            model.train()
            running_loss = 0.0
            work_this_epoch = 0
            for inputs, targets in trainloader:
                inputs, targets = inputs.to(self.device), targets.to(self.device)
                optimizer.zero_grad()
                outputs = model(inputs)
                loss = F.cross_entropy(outputs, targets)
                # Fix: Retain graph for second-order computations
                loss.backward(retain_graph=True)
                
                work_batch = sum(p.numel() for p in model.parameters() if p.requires_grad) * inputs.size(0)
                work_this_epoch += work_batch

                curvature_trace = self._compute_curvature_trace_estimate(model, loss, num_samples=2)
                base_lr = optimizer.param_groups[0]['lr']
                self._adaptive_noise(model, loss, base_lr, curvature_trace)

                if self.iteration % self.saddle_escape_freq == 0:
                    replica.second_order_correction(loss)

                optimizer.step()
                running_loss += loss.item()
                
            replica.compute_yield(work_this_epoch)
            print(f"Epoch {epoch} | Replica {self.replicas.index(replica)} | Loss {running_loss/len(trainloader):.4f} | Yield {replica.yield_ratio:.4f}")
        self.iteration += 1

    def evaluate(self, testloader):
        best_acc = 0.0
        for replica in self.replicas:
            model = replica.model
            model.eval()
            correct, total = 0, 0
            with torch.no_grad():
                for inputs, targets in testloader:
                    inputs, targets = inputs.to(self.device), targets.to(self.device)
                    outputs = model(inputs)
                    _, predicted = outputs.max(1)
                    total += targets.size(0)
                    correct += predicted.eq(targets).sum().item()
            acc = 100.0 * correct / total
            best_acc = max(best_acc, acc)
            print(f"Replica {self.replicas.index(replica)} Test Acc: {acc:.2f}%")
        return best_acc, None

def main():
    device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
    transform = transforms.Compose([transforms.ToTensor(), transforms.Normalize((0.4914, 0.4822, 0.4465), (0.2023, 0.1994, 0.2010))])
    trainset = torchvision.datasets.CIFAR10(root='./data', train=True, download=True, transform=transform)
    trainloader = torch.utils.data.DataLoader(trainset, batch_size=128, shuffle=True)
    testset = torchvision.datasets.CIFAR10(root='./data', train=False, download=True, transform=transform)
    testloader = torch.utils.data.DataLoader(testset, batch_size=100, shuffle=False)

    replicas_config = [
        {'optimizer_fn': optim.SGD, 'hyperparams': {'lr': 0.01, 'momentum': 0.9, 'weight_decay': 5e-4}},
        {'optimizer_fn': optim.Adam, 'hyperparams': {'lr': 0.001, 'weight_decay': 5e-4}}
    ]
    
    trainer = YIELDTrainer(CIFAR10CNN, replicas_config, device, yield_threshold=0.005)
    for epoch in range(10):
        trainer.train_epoch(trainloader, epoch)
        trainer.evaluate(testloader)

if __name__ == "__main__":
    main()