import torch
import torch.nn as nn
import torch.optim as optim
from torchvision import datasets, transforms
from torch.utils.data import DataLoader

# 1. Define the ODE-CCT Iterative Classifier Architecture
class ODECCTClassifier(nn.Module):
    def __init__(self, input_dim=784, num_classes=10, hidden_dim=128):
        super(ODECCTClassifier, self).__init__()
        self.hidden_dim = hidden_dim
        
        # State update network: F(x, prob_t)
        self.fc_x = nn.Linear(input_dim, hidden_dim)
        self.fc_p = nn.Linear(num_classes, hidden_dim)
        self.fc_out = nn.Linear(hidden_dim, num_classes)
        self.relu = nn.ReLU()

    def forward(self, x, init_prob, clock_freq=10.0, max_cycles=20, entropy_threshold=None):
        """
        x: Input image tensor [batch_size, 784]
        init_prob: Initial state probability distribution [batch_size, 10]
        clock_freq: Simulated computer clock frequency (determines step size dt = 1/f)
        max_cycles: Maximum processing clock cycles allowed (budget constraint)
        entropy_threshold: Dynamic collapse criteria. If None, runs full budget.
        """
        batch_size = x.size(0)
        dt = 1.0 / clock_freq
        
        # Convert initial probability to logit state space
        # Add small epsilon to prevent log(0)
        logits = torch.log(init_prob + 1e-8) 
        
        # Track clock steps spent per sample
        cycles_spent = torch.zeros(batch_size, device=x.device)
        active_mask = torch.ones(batch_size, dtype=torch.bool, device=x.device)
        
        x_emb = self.relu(self.fc_x(x))
        
        for step in range(max_cycles):
            if not active_mask.any():
                break
                
            # Current probabilities
            p = torch.softmax(logits, dim=-1)
            
            # Compute ODE derivative state modification: F(x, p)
            p_emb = self.relu(self.fc_p(p))
            derivative = self.fc_out(self.relu(x_emb + p_emb))
            
            # Update step (Euler Integration): u_{t+1} = u_t + dt * F(x, p)
            # Only update samples that haven't collapsed yet
            new_logits = logits + dt * derivative
            logits = torch.where(active_mask.unsqueeze(-1), new_logits, logits)
            
            # Update cycle statistics
            cycles_spent += active_mask.float()
            
            # Evaluate Entropy Collapse Criteria: H(p) = -sum(p * log(p))
            if entropy_threshold is not None:
                current_p = torch.softmax(logits, dim=-1)
                entropy = -torch.sum(current_p * torch.log(current_p + 1e-8), dim=-1)
                # Collapse occurred if entropy drops below threshold
                collapsed = entropy < entropy_threshold
                active_mask = active_mask & ~collapsed

        return logits, cycles_spent

# 2. Main execution wrapper for training and testing
def run_simulation():
    device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
    print(f"Running on processing unit: {device}")

    # Hyperparameters
    BATCH_SIZE = 64
    EPOCHS = 2
    CLOCK_FREQ = 5.0        # System clock rate
    MAX_CYCLES = 15         # Computational upper bound
    ENTROPY_THETA = 0.3     # Collapse threshold target

    # MNIST Data pipelines
    transform = transforms.Compose([transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,)), transforms.Lambda(lambda x: x.view(-1))])
    train_dataset = datasets.MNIST('../data', train=True, download=True, transform=transform)
    test_dataset = datasets.MNIST('../data', train=False, transform=transform)
    train_loader = DataLoader(train_dataset, batch_size=BATCH_SIZE, shuffle=True)
    test_loader = DataLoader(test_dataset, batch_size=BATCH_SIZE, shuffle=False)

    model = ODECCTClassifier().to(device)
    optimizer = optim.Adam(model.parameters(), lr=0.002)
    criterion = nn.CrossEntropyLoss()

    for _ in range(10):
        # Training phase
        model.train()
        print("\n--- Training Phase (Building Stationary Rules) ---")
        for epoch in range(EPOCHS):
            total_loss = 0
            for batch_idx, (data, target) in enumerate(train_loader):
                data, target = data.to(device), target.to(device)
                
                # Initial Condition: Uniform probability distribution (maximum uncertainty)
                init_prob = torch.ones(data.size(0), 10, device=device) / 10.0
                
                optimizer.zero_grad()
                # Train over the standard clock processing parameters
                outputs, _ = model(data, init_prob, clock_freq=CLOCK_FREQ, max_cycles=MAX_CYCLES)
                loss = criterion(outputs, target)
                loss.backward()
                optimizer.step()
                
                total_loss += loss.item()
            print(f"Epoch {epoch+1}/{EPOCHS} | Optimization Loss: {total_loss/len(train_loader):.4f}")

        # Evaluation phase using conditional early-exit collapse
        model.eval()
        correct = 0
        total_cycles = 0
        total_samples = 0
        
        print("\n--- Evaluation Phase (Conditional Collapse Testing) ---")
        with torch.no_grad():
            for data, target in test_loader:
                data, target = data.to(device), target.to(device)
                init_prob = torch.ones(data.size(0), 10, device=device) / 10.0
                
                # Execute with active early-exit tracking using the entropy threshold
                outputs, cycles = model(data, init_prob, clock_freq=CLOCK_FREQ, max_cycles=MAX_CYCLES, entropy_threshold=ENTROPY_THETA)
                
                preds = outputs.argmax(dim=-1)
                correct += (preds == target).sum().item()
                total_cycles += cycles.sum().item()
                total_samples += data.size(0)

        accuracy = (correct / total_samples) * 100
        avg_cycles = total_cycles / total_samples
        print(f"\nFinal Classification Accuracy: {accuracy:.2f}%")
        print(f"Average Clock Cycles Expended per Sample: {avg_cycles:.2f} / {MAX_CYCLES} max cycles")
        print(f"Energy Efficiency Ratio (Compute Saved): {((1.0 - (avg_cycles / MAX_CYCLES)) * 100):.2f}%")

if __name__ == '__main__':
    run_simulation()
