"""
YIELD Framework for CIFAR10: Global Minimum Convergence via Paradox-Driven Optimization

Implements key concepts from the YIELD programming language:
- Superposition of multiple solver replicas with different hyperparameters
- Yield ratio (Δentropy / Δwork) to collapse low-performing replicas
- Second-order correction for flat saddles using Hessian eigenvector
- Curvature-aware adaptive noise (thermodynamic temperature)
- Collapse mechanism that keeps only high-yield trajectories

This model trains a CNN on CIFAR10, aiming to reach a near-global minimum.
"""

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           # total energy spent (e.g., FLOPs proxy)
        self.init_entropy = None  # initial weight entropy
        self.current_entropy = None
        self.yield_ratio = 0.0
        self.alive = True
        self.saddle_escape_attempts = 0
        
        # Compute initial weight entropy (variance-based)
        self._update_entropy()
        self.init_entropy = self.current_entropy
        
    def _update_entropy(self):
        """Compute entropy of weights as sum of per-layer log variance."""
        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()
        # Use log variance as entropy proxy (differential entropy for Gaussian)
        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):
        """Update yield ratio after performing 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 add_work(self, amount):
        """Accumulate work (e.g., number of forward/backward passes)."""
        self.work += amount
        
    def get_parameters(self):
        return self.model.parameters()
    
    def get_model(self):
        return self.model
    
    # -------------------------------
    # Flat saddle escape via second-order correction
    # -------------------------------
    def second_order_correction(self, loss, max_iter=10, tolerance=1e-6):
        """
        If gradient norm is very small and Hessian has near-zero eigenvalue,
        take a step along the smallest eigenvector to escape flat region.
        Uses power iteration on Hessian-vector product (approximate).
        Returns True if a correction was applied.
        """
        # Only apply if gradient is tiny
        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
        
        # Compute smallest eigenvalue direction using Hessian-vector product
        params = list(self.model.parameters())
        n_params = sum(p.numel() for p in params)
        # Power iteration to find eigenvector with smallest magnitude eigenvalue
        # We use the Lanczos-style approach: we want direction of minimal curvature.
        # Use a random vector v
        v = [torch.randn_like(p) for p in params]
        # Normalize
        norm_v = math.sqrt(sum(p.norm().item()**2 for p in v))
        v = [p / norm_v for p in v]
        
        # Function to compute Hv
        def hvp(v):
            # Compute gradient of (gradient dot 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
        
        # Power iteration to find eigenvector corresponding to smallest eigenvalue (in absolute)
        # We'll use few iterations to estimate curvature
        eigval_old = 0.0
        for _ in range(max_iter):
            hv = hvp(v)
            # Compute Rayleigh quotient
            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
            # Update v = hv / norm(hv)
            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 eigenvalue near zero (flat direction), step along it
        if abs(eigval) < 1e-4:
            # Step size alpha (boosted to escape flat plateau)
            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:
    """
    Manages superposition of replicas, collapse based on yield ratio,
    curvature-aware noise, and global convergence.
    """
    def __init__(self, model_fn, replicas_config, device, yield_threshold=0.01,
                 curvature_noise_gamma=100.0, saddle_escape_freq=50):
        """
        replicas_config: list of dicts, each with 'optimizer_fn' and 'hyperparams'
        """
        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):
        """
        Estimate trace of Hessian via Hutchinson's method.
        Returns average curvature (proxy for trace).
        """
        params = list(model.parameters())
        trace_estimate = 0.0
        for _ in range(num_samples):
            v = [torch.randn_like(p) for p in params]
            # Compute Hv
            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):
        """Add noise to gradients based on local curvature (inverse temperature)."""
        # Temperature high in flat regions (small 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):
        """Remove replicas with yield ratio below threshold."""
        alive_replicas = [r for r in self.replicas if r.alive and r.yield_ratio >= self.yield_threshold]
        # Also keep at least one replica even if below threshold
        if len(alive_replicas) == 0 and len(self.replicas) > 0:
            # Keep the best one
            best = max(self.replicas, key=lambda r: r.yield_ratio)
            alive_replicas = [best]
        # Optionally, if only one left, keep it always
        self.replicas = alive_replicas
        # Spawn a new replica from the best if only one remains? (Optional)
        if len(self.replicas) == 1:
            # Create a mutated copy with slightly different hyperparams
            best = self.replicas[0]
            new_hyperparams = best.hyperparams.copy()
            # Mutate learning rate slightly
            if 'lr' in new_hyperparams:
                new_hyperparams['lr'] *= np.random.uniform(0.5, 1.5)
            new_rep = YieldReplica(self.model_fn, 
                                   lambda params: optim.SGD(params, **new_hyperparams),
                                   new_hyperparams, self.device)
            # Copy weights from best
            new_rep.model.load_state_dict(copy.deepcopy(best.model.state_dict()))
            self.replicas.append(new_rep)
    
    def train_epoch(self, trainloader, epoch):
        """Train all replicas for one epoch."""
        for replica in self.replicas:
            if not replica.alive:
                continue
            model = replica.model
            optimizer = replica.optimizer
            model.train()
            running_loss = 0.0
            work_this_epoch = 0
            for batch_idx, (inputs, targets) in enumerate(trainloader):
                inputs, targets = inputs.to(self.device), targets.to(self.device)
                optimizer.zero_grad()
                outputs = model(inputs)
                loss = F.cross_entropy(outputs, targets)
                loss.backward()
                
                # Work: number of parameters processed (proxy)
                work_batch = sum(p.numel() for p in model.parameters() if p.requires_grad) * inputs.size(0)
                work_this_epoch += work_batch
                
                # Curvature-aware noise
                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)
                
                # Optional: second-order escape from flat saddles
                if self.iteration % self.saddle_escape_freq == 0:
                    replica.second_order_correction(loss)
                
                optimizer.step()
                running_loss += loss.item()
                
            avg_loss = running_loss / len(trainloader)
            # Update yield ratio for this replica
            replica.compute_yield(work_this_epoch)
            print(f"Epoch {epoch} | Replica {self.replicas.index(replica)} | Loss {avg_loss:.4f} | Yield {replica.yield_ratio:.4f}")
        self.iteration += 1
    
    def evaluate(self, testloader):
        """Evaluate all alive replicas on test set and return best accuracy."""
        best_acc = 0.0
        best_model = None
        for replica in self.replicas:
            if not replica.alive:
                continue
            model = replica.model
            model.eval()
            correct = 0
            total = 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
            if acc > best_acc:
                best_acc = acc
                best_model = copy.deepcopy(model)
            print(f"Replica {self.replicas.index(replica)} Test Accuracy: {acc:.2f}%")
        return best_acc, best_model

# -------------------------------
# 4. Main Experiment
# -------------------------------
def main():
    device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
    print(f"Using device: {device}")
    
    # Data preparation
    transform_train = transforms.Compose([
        transforms.RandomCrop(32, padding=4),
        transforms.RandomHorizontalFlip(),
        transforms.ToTensor(),
        transforms.Normalize((0.4914, 0.4822, 0.4465), (0.2023, 0.1994, 0.2010)),
    ])
    transform_test = 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_train)
    trainloader = torch.utils.data.DataLoader(trainset, batch_size=128, shuffle=True, num_workers=2)
    testset = torchvision.datasets.CIFAR10(root='./data', train=False, download=True, transform=transform_test)
    testloader = torch.utils.data.DataLoader(testset, batch_size=100, shuffle=False, num_workers=2)
    
    # Define hyperparameter configurations for replicas (superposition)
    replicas_config = [
        {'optimizer_fn': lambda params: optim.SGD(params, lr=0.01, momentum=0.9, weight_decay=5e-4),
         'hyperparams': {'lr': 0.01, 'momentum': 0.9, 'weight_decay': 5e-4}},
        {'optimizer_fn': lambda params: optim.SGD(params, lr=0.05, momentum=0.9, weight_decay=5e-4),
         'hyperparams': {'lr': 0.05, 'momentum': 0.9, 'weight_decay': 5e-4}},
        {'optimizer_fn': lambda params: optim.Adam(params, lr=0.001, weight_decay=5e-4),
         'hyperparams': {'lr': 0.001, 'weight_decay': 5e-4}},
        {'optimizer_fn': lambda params: optim.RMSprop(params, lr=0.001, weight_decay=5e-4),
         'hyperparams': {'lr': 0.001, 'weight_decay': 5e-4}},
    ]
    
    trainer = YIELDTrainer(CIFAR10CNN, replicas_config, device, yield_threshold=0.005,
                           curvature_noise_gamma=100.0, saddle_escape_freq=30)
    
    print("Starting YIELD training with superposition and collapse...")
    num_epochs = 50
    best_acc = 0.0
    for epoch in range(num_epochs):
        trainer.train_epoch(trainloader, epoch)
        # Collapse low-yield replicas every few epochs
        if epoch % 5 == 0 and epoch > 0:
            trainer.collapse_low_yield_replicas()
            print(f"After collapse, {len(trainer.replicas)} replicas remain.")
        test_acc, _ = trainer.evaluate(testloader)
        if test_acc > best_acc:
            best_acc = test_acc
            print(f"New best accuracy: {best_acc:.2f}%")
    
    print(f"Final best test accuracy: {best_acc:.2f}%")
    print("YIELD optimization complete. Global minimum approximated.")

if __name__ == "__main__":
    main()