"""
Hierarchical YIELD Framework: Global Minimum via Coarse-to-Fine Replica Expansion

Key idea:
- Stage 1: Many replicas of a very small model (few parameters) explore globally.
- Collapse low-yield replicas based on entropy reduction per work.
- Expand survivors to higher-resolution models (more parameters) by transferring/upscaling weights.
- Repeat stages to focus compute on promising regions, approximating global minimum.

This approach leverages:
- Low cost of small models to sample many trajectories.
- Yield ratio to identify promising basins.
- Progressive refinement to increase capacity where needed.
"""

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 OrderedDict

# -------------------------------
# 1. Dynamic Model Builder (Coarse to Fine)
# -------------------------------

class FlexibleCIFAR10CNN(nn.Module):
    """
    A CNN whose width and depth can be scaled.
    stage=0: tiny (1 conv, 8 filters, no batchnorm)
    stage=1: small (2 convs, 16 filters, batchnorm)
    stage=2: medium (3 convs, 32 filters, batchnorm)
    stage=3: full (3 convs, 64 filters, batchnorm, dropout)
    """
    def __init__(self, stage=0, num_classes=10):
        super().__init__()
        self.stage = stage
        self.num_classes = num_classes
        
        if stage == 0:
            # Tiny: only one conv layer + simple classifier
            self.conv1 = nn.Conv2d(3, 8, 3, padding=1)
            self.pool = nn.MaxPool2d(2, 2)
            self.fc = nn.Linear(8 * 16 * 16, num_classes)  # after one pool: 32/2=16
            self._features = nn.Sequential(self.conv1, nn.ReLU(), self.pool)
        
        elif stage == 1:
            # Small: two convs, batchnorm, more filters
            self.conv1 = nn.Conv2d(3, 16, 3, padding=1)
            self.bn1 = nn.BatchNorm2d(16)
            self.conv2 = nn.Conv2d(16, 32, 3, padding=1)
            self.bn2 = nn.BatchNorm2d(32)
            self.pool = nn.MaxPool2d(2, 2)
            self.fc = nn.Linear(32 * 8 * 8, num_classes)  # after two pools: 32/4=8
            self._features = nn.Sequential(
                self.conv1, self.bn1, nn.ReLU(), self.pool,
                self.conv2, self.bn2, nn.ReLU(), self.pool
            )
        
        elif stage == 2:
            # Medium: three convs, dropout on fc
            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.fc = nn.Linear(128 * 4 * 4, 256)
            self.dropout = nn.Dropout(0.3)
            self.out = nn.Linear(256, num_classes)
            self._features = nn.Sequential(
                self.conv1, self.bn1, nn.ReLU(), self.pool,
                self.conv2, self.bn2, nn.ReLU(), self.pool,
                self.conv3, self.bn3, nn.ReLU(), self.pool
            )
        
        else:  # stage >= 3: full model (similar to original)
            self.conv1 = nn.Conv2d(3, 64, 3, padding=1)
            self.bn1 = nn.BatchNorm2d(64)
            self.conv2 = nn.Conv2d(64, 128, 3, padding=1)
            self.bn2 = nn.BatchNorm2d(128)
            self.conv3 = nn.Conv2d(128, 256, 3, padding=1)
            self.bn3 = nn.BatchNorm2d(256)
            self.pool = nn.MaxPool2d(2, 2)
            self.fc1 = nn.Linear(256 * 4 * 4, 512)
            self.dropout1 = nn.Dropout(0.3)
            self.fc2 = nn.Linear(512, 256)
            self.dropout2 = nn.Dropout(0.3)
            self.out = nn.Linear(256, num_classes)
            self._features = nn.Sequential(
                self.conv1, self.bn1, nn.ReLU(), self.pool,
                self.conv2, self.bn2, nn.ReLU(), self.pool,
                self.conv3, self.bn3, nn.ReLU(), self.pool
            )
    
    def forward(self, x):
        if self.stage <= 2:
            x = self._features(x)
            x = x.view(x.size(0), -1)
            if self.stage <= 1:
                x = self.fc(x)
            else:
                x = F.relu(self.fc(x))
                x = self.dropout(x)
                x = self.out(x)
        else:
            x = self._features(x)
            x = x.view(x.size(0), -1)
            x = F.relu(self.fc1(x))
            x = self.dropout1(x)
            x = F.relu(self.fc2(x))
            x = self.dropout2(x)
            x = self.out(x)
        return x
    
    def get_num_params(self):
        return sum(p.numel() for p in self.parameters() if p.requires_grad)

def expand_model(small_model, target_stage):
    """
    Transfer weights from a smaller stage model to a larger stage model.
    Uses heuristics: copy weights for matching layers, initialize new layers.
    """
    large_model = FlexibleCIFAR10CNN(stage=target_stage, num_classes=10)
    small_state = small_model.state_dict()
    large_state = large_model.state_dict()
    
    # Transfer common layer names (if dimensions match)
    for name in small_state:
        if name in large_state:
            s_tensor = small_state[name]
            l_tensor = large_state[name]
            # If shapes match, copy directly
            if s_tensor.shape == l_tensor.shape:
                large_state[name] = s_tensor.clone()
            else:
                # Handle dimension mismatches (e.g., conv1 weight: from 8->16 filters)
                # We can repeat or take subset; here we repeat for simplicity
                if 'weight' in name and len(s_tensor.shape) == 4:
                    # Repeat filters to match larger output channels
                    rep_factor = l_tensor.shape[0] // s_tensor.shape[0]
                    if rep_factor > 1:
                        repeated = s_tensor.repeat(rep_factor, 1, 1, 1)
                        # Also adjust input channels if needed (for later layers)
                        if repeated.shape != l_tensor.shape:
                            # Slice to exact size
                            repeated = repeated[:l_tensor.shape[0], :l_tensor.shape[1], :, :]
                        large_state[name] = repeated
                elif 'bias' in name:
                    rep_factor = l_tensor.shape[0] // s_tensor.shape[0]
                    if rep_factor > 1:
                        repeated = s_tensor.repeat(rep_factor)
                        large_state[name] = repeated[:l_tensor.shape[0]]
                # For linear layers, handle by repeating output neurons
                elif 'fc' in name or 'out' in name:
                    if len(s_tensor.shape) == 2:
                        # weight matrix: [out, in]
                        rep_factor = l_tensor.shape[0] // s_tensor.shape[0]
                        if rep_factor > 1:
                            repeated = s_tensor.repeat(rep_factor, 1)
                            large_state[name] = repeated[:l_tensor.shape[0], :l_tensor.shape[1]]
                    elif len(s_tensor.shape) == 1:
                        rep_factor = l_tensor.shape[0] // s_tensor.shape[0]
                        if rep_factor > 1:
                            repeated = s_tensor.repeat(rep_factor)
                            large_state[name] = repeated[:l_tensor.shape[0]]
    large_model.load_state_dict(large_state)
    return large_model

# -------------------------------
# 2. YIELD Replica with Stage Tracking
# -------------------------------

class HierarchicalYieldReplica:
    def __init__(self, model, optimizer_fn, hyperparams, device, stage, replica_id):
        self.device = device
        self.model = model.to(device)
        self.optimizer = optimizer_fn(self.model.parameters(), **hyperparams)
        self.hyperparams = hyperparams.copy()
        self.stage = stage
        self.id = replica_id
        self.work = 0.0
        self.init_entropy = None
        self.current_entropy = None
        self.yield_ratio = 0.0
        self.alive = True
        self._update_entropy()
        self.init_entropy = self.current_entropy
        
    def _update_entropy(self):
        total_var = 0.0
        n_params = 0
        for p in self.model.parameters():
            if p.requires_grad:
                var = p.data.var().item()
                total_var += var * p.numel()
                n_params += p.numel()
        if total_var > 0:
            entropy = 0.5 * math.log(2 * math.pi * math.e * (total_var / n_params))
        else:
            entropy = 0.0
        self.current_entropy = entropy
        
    def compute_yield(self, work_delta):
        self.work += work_delta
        self._update_entropy()
        entropy_red = self.init_entropy - self.current_entropy
        self.yield_ratio = entropy_red / self.work if self.work > 0 else 0.0
        return self.yield_ratio
    
    def add_work(self, amount):
        self.work += amount
    
    def train_step(self, inputs, targets):
        self.optimizer.zero_grad()
        outputs = self.model(inputs)
        loss = F.cross_entropy(outputs, targets)
        loss.backward()
        # optional curvature noise omitted for brevity
        self.optimizer.step()
        return loss.item()
    
    def get_model(self):
        return self.model

# -------------------------------
# 3. Hierarchical YIELD Trainer
# -------------------------------

class HierarchicalYIELDTrainer:
    def __init__(self, device, yield_threshold=0.001, num_stages=4, 
                 replicas_per_stage=[32, 16, 8, 4], epochs_per_stage=[5, 10, 20, 30]):
        self.device = device
        self.yield_threshold = yield_threshold
        self.num_stages = num_stages
        self.replicas_per_stage = replicas_per_stage
        self.epochs_per_stage = epochs_per_stage
        self.all_replicas = []  # keep for history
        
    def create_replicas(self, stage, num_replicas, base_model=None):
        replicas = []
        for i in range(num_replicas):
            if base_model is not None:
                # Expand from previous best model (or random variant)
                model = expand_model(base_model, stage)
            else:
                model = FlexibleCIFAR10CNN(stage=stage, num_classes=10)
            # Randomize hyperparams slightly per replica
            lr = np.random.uniform(0.001, 0.1) if stage < 2 else np.random.uniform(0.0001, 0.01)
            optimizer_fn = lambda params: optim.SGD(params, lr=lr, momentum=0.9, weight_decay=5e-4)
            hyperparams = {'lr': lr, 'momentum': 0.9, 'weight_decay': 5e-4}
            rep = HierarchicalYieldReplica(model, optimizer_fn, hyperparams, self.device, stage, i)
            replicas.append(rep)
        return replicas
    
    def train_stage(self, replicas, trainloader, epochs):
        for epoch in range(epochs):
            for replica in replicas:
                if not replica.alive:
                    continue
                replica.model.train()
                running_loss = 0.0
                work_epoch = 0
                for inputs, targets in trainloader:
                    inputs, targets = inputs.to(self.device), targets.to(self.device)
                    loss = replica.train_step(inputs, targets)
                    running_loss += loss
                    # Work = model parameters * batch_size
                    work_batch = replica.model.get_num_params() * inputs.size(0)
                    work_epoch += work_batch
                avg_loss = running_loss / len(trainloader)
                replica.compute_yield(work_epoch)
                print(f"Stage {replica.stage} | Rep {replica.id} | Epoch {epoch} | Loss {avg_loss:.4f} | Yield {replica.yield_ratio:.4f}")
    
    def collapse_replicas(self, replicas, keep_ratio=0.5):
        """Keep top-k replicas by yield ratio."""
        alive = [r for r in replicas if r.alive]
        if len(alive) == 0:
            return []
        alive.sort(key=lambda r: r.yield_ratio, reverse=True)
        keep_count = max(1, int(len(alive) * keep_ratio))
        survivors = alive[:keep_count]
        for r in alive[keep_count:]:
            r.alive = False
        return survivors
    
    def run(self, trainloader, testloader):
        current_replicas = []
        best_test_acc = 0.0
        best_model = None
        
        for stage in range(self.num_stages):
            print(f"\n=== STAGE {stage} (model capacity: {FlexibleCIFAR10CNN(stage=stage).get_num_params()} params) ===")
            if stage == 0:
                # Initialize fresh replicas
                current_replicas = self.create_replicas(stage, self.replicas_per_stage[stage])
            else:
                # Expand survivors from previous stage
                survivors = [r for r in current_replicas if r.alive]
                if len(survivors) == 0:
                    print("No survivors, restarting from random")
                    current_replicas = self.create_replicas(stage, self.replicas_per_stage[stage])
                else:
                    new_replicas = []
                    # Keep some expanded versions of each survivor
                    per_survivor = max(1, self.replicas_per_stage[stage] // len(survivors))
                    for surv in survivors:
                        for _ in range(per_survivor):
                            expanded_model = expand_model(surv.get_model(), stage)
                            # Slight mutation
                            lr = surv.hyperparams['lr'] * np.random.uniform(0.8, 1.2)
                            opt_fn = lambda params: optim.SGD(params, lr=lr, momentum=0.9, weight_decay=5e-4)
                            hyper = {'lr': lr, 'momentum': 0.9, 'weight_decay': 5e-4}
                            new_rep = HierarchicalYieldReplica(expanded_model, opt_fn, hyper, self.device, stage, len(new_replicas))
                            new_replicas.append(new_rep)
                    # If we need more, add random ones
                    while len(new_replicas) < self.replicas_per_stage[stage]:
                        new_replicas.extend(self.create_replicas(stage, 1))
                    current_replicas = new_replicas[:self.replicas_per_stage[stage]]
            
            # Train this stage
            self.train_stage(current_replicas, trainloader, self.epochs_per_stage[stage])
            
            # Collapse low-yield replicas
            current_replicas = self.collapse_replicas(current_replicas, keep_ratio=0.5)
            print(f"After collapse: {len(current_replicas)} replicas remain")
            
            # Evaluate best replica on test set
            best_rep = max(current_replicas, key=lambda r: r.yield_ratio)
            test_acc = self.evaluate(best_rep.model, testloader)
            print(f"Best replica test accuracy: {test_acc:.2f}%")
            if test_acc > best_test_acc:
                best_test_acc = test_acc
                best_model = copy.deepcopy(best_rep.model)
        
        return best_model, best_test_acc
    
    def evaluate(self, model, testloader):
        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()
        return 100.0 * correct / total

# -------------------------------
# 4. Main
# -------------------------------

def main():
    device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
    print(f"Using device: {device}")
    
    # Data
    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)
    
    trainer = HierarchicalYIELDTrainer(device, yield_threshold=0.001,
                                       num_stages=4,
                                       replicas_per_stage=[32, 16, 8, 4],
                                       epochs_per_stage=[5, 10, 20, 30])
    
    best_model, best_acc = trainer.run(trainloader, testloader)
    print(f"\nFinal best test accuracy: {best_acc:.2f}%")
    print("Hierarchical YIELD optimization completed. Global minimum approximated via coarse-to-fine search.")

if __name__ == "__main__":
    main()