import torch
import torch.nn as nn
import torch.optim as optim
import torchvision
import torchvision.transforms as transforms
import numpy as np
import time
import warnings
warnings.filterwarnings('ignore')

# ==========================================
# 1. MODEL: Standard MLP for CIFAR-10
# ==========================================
class CCTMLP(nn.Module):
    def __init__(self):
        super().__init__()
        self.flatten = nn.Flatten()
        self.network = nn.Sequential(
            nn.Linear(32 * 32 * 3, 512),
            nn.ReLU(),
            nn.Dropout(0.3),
            nn.Linear(512, 256),
            nn.ReLU(),
            nn.Dropout(0.2),
            nn.Linear(256, 10)
        )

    def forward(self, x):
        return self.network(self.flatten(x))

# ==========================================
# 2. CCT-ML TRAINER ENGINE
# ==========================================
class CCTMLTrainer:
    def __init__(self, model, lr=1e-3, device='cpu', 
                 collapse_threshold=0.05, ode_curvature_sensitivity=1e-5):
        self.model = model.to(device)
        self.optimizer = optim.Adam(model.parameters(), lr=lr)
        self.criterion = nn.CrossEntropyLoss()
        self.device = device
        
        # CCT State Variables
        self.H_T = 1.0                 # Theory Entropy Proxy (EMA of loss)
        self.H_T_var = 0.0             # Variance Proxy (Uncertainty Spread)
        self.collapse_threshold = collapse_threshold
        self.ode_sensitivity = ode_curvature_sensitivity
        self.loss_history = []         # For ODE curvature tracking
        
        # Energy/Work Tracking
        self.total_work = 0            # Compute steps spent
        self.collapse_path = []        # Trace of Δ_i per step
        
    def _compute_entropy_proxy(self, batch_loss):
        """CCT: H(T) = Stationary (mean loss) + Probability (variance/spread)"""
        alpha = 0.02
        self.H_T = (1 - alpha) * self.H_T + alpha * batch_loss
        # Simple variance proxy: squared deviation from moving average
        self.H_T_var = (1 - alpha) * self.H_T_var + alpha * (batch_loss - self.H_T)**2
        return self.H_T + 0.5 * self.H_T_var

    def _ode_lr_adjust(self, current_lr, min_lr=1e-6, max_lr=5e-3):
        """CCT-ODE: Detect harmonic oscillation in loss curvature.
        If d²L/dt² < 0 (oscillating), reduce LR. If monotonic decay, maintain/slightly increase."""
        if len(self.loss_history) < 4:
            return current_lr
            
        # Finite difference derivatives
        d1 = self.loss_history[-1] - self.loss_history[-2]
        d2 = (self.loss_history[-1] - self.loss_history[-2]) - \
             (self.loss_history[-2] - self.loss_history[-3])
             
        # Oscillation detected: collapse potential is bouncing
        if d2 < -self.ode_sensitivity and d1 * (self.loss_history[-2] - self.loss_history[-3]) < 0:
            return max(min_lr, current_lr * 0.85)  # Dampen oscillation
        else:
            return min(max_lr, current_lr * 1.01)  # Gentle acceleration during smooth collapse

    def _check_distributional_collapse(self, val_loss_window, tolerance=0.005):
        """CCT: Stop when empirical loss distribution stabilizes (KL-divergence proxy)."""
        if len(val_loss_window) < 5:
            return False
        m = np.mean(val_loss_window)
        v = np.var(val_loss_window)
        # Target: low mean loss + near-zero variance = distributional match to "solved" state
        return m < self.collapse_threshold and v < tolerance

    def train(self, train_loader, val_loader, max_epochs=15):
        self.model.train()
        best_val_acc = 0.0
        val_loss_window = []
        lr = self.optimizer.param_groups[0]['lr']
        
        print(f"{'Step':<6} | {'Batch Loss':<10} | {'H(T)':<8} | {'Δ_i':<8} | {'LR':<9} | {'CCT State'}")
        print("-" * 85)
        
        step = 0
        start_time = time.time()
        
        for epoch in range(max_epochs):
            for inputs, labels in train_loader:
                inputs, labels = inputs.to(self.device), labels.to(self.device)

                self.optimizer.zero_grad()

                loss = 0
                for _ in range(10):
                    inputs, labels = next(iter(train_loader))
                    inputs, labels = inputs.to(self.device), labels.to(self.device)
                
                    # Forward
                    outputs = self.model(inputs)
                    loss += 0.01 * self.criterion(outputs, labels)
                
                # CCT: Compute Entropy Proxy & Collapse Potential
                H_before = self.H_T + 0.5 * self.H_T_var
                loss_val = loss.item()
                self.loss_history.append(loss_val)
                H_after = self._compute_entropy_proxy(loss_val)
                
                delta_i = H_before - H_after  # Collapse Potential
                self.collapse_path.append(delta_i)
                
                # Backward & Step (Work Investment)
                loss.backward()
                self.optimizer.step()
                self.total_work += 1
                
                # ODE-CCT LR Adjustment
                lr = 1e-3 #+ self._ode_lr_adjust(lr)
                for param_group in self.optimizer.param_groups:
                    param_group['lr'] = lr
                    
                step += 1
                
                # Logging & Conditional Collapse Check
                if step % 50 == 0:
                    self.model.eval()
                    with torch.no_grad():
                        val_loss = 0.0
                        correct = 0
                        total = 0
                        for v_in, v_lab in val_loader:
                            v_in, v_lab = v_in.to(self.device), v_lab.to(self.device)
                            v_out = self.model(v_in)
                            val_loss += self.criterion(v_out, v_lab).item() * v_in.size(0)
                            correct += (v_out.argmax(1) == v_lab).sum().item()
                            total += v_in.size(0)
                        val_loss /= total
                        val_acc = correct / total
                        self.model.train()
                        
                    val_loss_window.append(val_loss)
                    collapsed = self._check_distributional_collapse(val_loss_window)
                    
                    state = "COLLAPSED" if collapsed else "NAVIGATING"
                    if val_acc > best_val_acc:
                        best_val_acc = val_acc
                        
                    print(f"{step:<6} | {loss_val:<10.4f} | {H_after:<8.4f} | {delta_i:<8.4f} | {lr:<9.6f} | {state}")
                    
                    if collapsed:
                        print(f"\n✅ Distributional Collapse Triggered at Step {step}. H(T) stabilized.")
                        print(f"⏱️ Total Work: {self.total_work} steps. Best Val Acc: {best_val_acc:.4f}")
                        return step, best_val_acc
                        
            # Epoch end
            print(f"--- Epoch {epoch+1} Complete | Val Acc: {val_acc:.4f} | H(T): {H_after:.4f} ---")
            
        elapsed = time.time() - start_time
        print(f"\n🏁 Max Epochs Reached. Final Val Acc: {best_val_acc:.4f} | Total Work: {self.total_work} | Time: {elapsed:.2f}s")
        return self.total_work, best_val_acc

    def test(self, test_loader):
        self.model.eval()
        correct, total = 0, 0
        with torch.no_grad():
            for inputs, labels in test_loader:
                inputs, labels = inputs.to(self.device), labels.to(self.device)
                outputs = self.model(inputs)
                correct += (outputs.argmax(1) == labels).sum().item()
                total += labels.size(0)
        return correct / total

# ==========================================
# 3. EXECUTION PIPELINE
# ==========================================
if __name__ == "__main__":
    DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")
    print(f"🚀 Initializing CCT-ML Framework on {DEVICE}...")

    # CIFAR-10 Loaders
    transform = transforms.Compose([
        transforms.ToTensor(),
        transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5))
    ])
    train_set = torchvision.datasets.CIFAR10(root='../data', train=True, download=True, transform=transform)
    test_set = torchvision.datasets.CIFAR10(root='../data', train=False, download=True, transform=transform)
    
    train_loader = torch.utils.data.DataLoader(train_set, batch_size=16, shuffle=True, num_workers=2)
    test_loader = torch.utils.data.DataLoader(test_set, batch_size=128, shuffle=False, num_workers=2)

    # Initialize CCT-ML
    model = CCTMLP()
    cct_trainer = CCTMLTrainer(model, lr=2e-3, device=DEVICE, collapse_threshold=0.15)

    # Train with Conditional Collapse Logic
    steps, best_val_acc = cct_trainer.train(train_loader, test_loader, max_epochs=10)

    # Final Evaluation
    test_acc = cct_trainer.test(test_loader)
    print(f"\n📊 Final Test Accuracy: {test_acc:.4f}")
    print(f"📈 Collapse Path Statistics: Δ_mean={np.mean(cct_trainer.collapse_path):.4f}, Δ_std={np.std(cct_trainer.collapse_path):.4f}")
