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

# ==========================================
# 1. The Continuous Frequency Architecture
# ==========================================

class ResonanceLayer(nn.Module):
    def __init__(self, in_channels, out_channels):
        super(ResonanceLayer, self).__init__()
        # Use 3x3 kernels as 'Tuning Probes'
        self.conv = nn.Conv2d(in_channels, out_channels, kernel_size=3, padding=1)
        self.bn = nn.BatchNorm2d(out_channels)

    def forward(self, x):
        res = self.conv(x)
        res = self.bn(res)
        # Tanh represents the oscillation between Tension (-1) and Resolution (1)
        return torch.tanh(res)

class CCTCollapseLayer(nn.Module):
    def __init__(self, in_features, num_classes, theta_coll=1.0):
        super(CCTCollapseLayer, self).__init__()
        self.linear = nn.Linear(in_features, num_classes)
        self.theta_coll = theta_coll

    def forward(self, x):
        logits = self.linear(x)
        # Semantic Pressure is the amplitude of the state vector
        pressure = torch.norm(logits, p=2, dim=1, keepdim=True)
        return logits, pressure

class ContinuousFrequencyNet(nn.Module):
    def __init__(self):
        super(ContinuousFrequencyNet, self).__init__()
        # CIFAR-10 images are 3x32x32
        self.res1 = ResonanceLayer(3, 32) # 3 channels for RGB
        self.res2 = ResonanceLayer(32, 64)
        self.res3 = ResonanceLayer(64, 128)
        
        self.pool = nn.AdaptiveAvgPool2d(1)
        self.collapse = CCTCollapseLayer(128, 10)

    def forward(self, x):
        x = self.res1(x)
        x = self.res2(x)
        x = self.res3(x)
        x = self.pool(x).view(x.size(0), -1)
        logits, pressure = self.collapse(x)
        return logits, pressure

# ==========================================
# 2. Training and Testing Pipeline
# ==========================================

def train():
    # Hyperparameters
    batch_size = 128 # Larger batch for CIFAR
    epochs = 10      # More epochs needed for complex images
    lr = 0.001

    # Data Loading (CIFAR-10)
    transform = transforms.Compose([
        transforms.ToTensor(), 
        transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5))
    ])
    
    train_dataset = datasets.CIFAR10('../.././data', train=True, download=True, transform=transform)
    test_dataset = datasets.CIFAR10('../.././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 = ContinuousFrequencyNet()
    optimizer = optim.Adam(model.parameters(), lr=lr)
    criterion = nn.CrossEntropyLoss()

    print("Initiating Resonance Training on CIFAR-10...\n")

    for epoch in range(epochs):
        model.train()
        total_loss = 0
        for batch_idx, (data, target) in enumerate(train_loader):
            optimizer.zero_grad()
            logits, pressure = model(data)
            loss = criterion(logits, target)
            loss.backward()
            optimizer.step()
            total_loss += loss.item()

        print(f"Epoch {epoch+1}/{epochs} | Loss: {total_loss/len(train_loader):.4f}")

    # ==========================================
    # 3. Evaluation using CCT Logic
    # ==========================================
    model.eval()
    num_correct_collapsed = 0
    num_unresolved = 0
    num_total_correct = 0 # Total matches regardless of pressure
    total_samples = 0
    
    # SET YOUR THRESHOLD HERE
    # If pressure < 5.0, the SI says "Insufficient Work"
    CCT_THRESHOLD = 5.0 

    with torch.no_grad():
        for data, target in test_loader:
            logits, pressure = model(data)
            
            pred = logits.argmax(dim=1)
            
            # Boolean mask for a successful collapse (Pressure > Threshold)
            collapsed_mask = (pressure.squeeze() >= CCT_THRESHOLD)
            # Boolean mask for a correct prediction
            correct_mask = (pred == target)
            
            # 1. Truly Correct: Predicted Correct AND Collapsed
            num_correct_collapsed += (correct_mask & collapsed_mask).sum().item()
            
            # 2. Unresolved: Did not cross the pressure threshold
            num_unresolved += (~collapsed_mask).sum().item()
            
            # 3. Total Correct: Just matches, ignore the CCT energy check
            num_total_correct += correct_mask.sum().item()
            
            total_samples += target.size(0)

    print(f"\n--- Final CCT Results ---")
    print(f"Total Samples: {total_samples}")
    print(f"Total Match Count: {num_total_correct} ({100 * num_total_correct / total_samples:.2f}%)")
    print(f"Correct & Collapsed: {num_correct_collapsed} ({100 * num_correct_collapsed / total_samples:.2f}%)")
    print(f"Unresolved (Insufficient Work): {num_unresolved} ({100 * num_unresolved / total_samples:.2f}%)")

if __name__ == "__main__":
    train()
