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

# ==========================================
# GPU Setup
# ==========================================
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
print(f"Using device: {device}")

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

class ResonanceLayer(nn.Module):
    def __init__(self, in_channels, out_channels):
        super(ResonanceLayer, self).__init__()
        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)
        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)
        pressure = torch.norm(logits, p=2, dim=1, keepdim=True)
        return logits, pressure

class ContinuousFrequencyNet(nn.Module):
    def __init__(self):
        super(ContinuousFrequencyNet, self).__init__()
        self.res1 = ResonanceLayer(3, 32)
        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():
    batch_size = 128
    epochs = 10
    lr = 0.001

    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)

    # MOVE MODEL TO GPU
    model = ContinuousFrequencyNet().to(device)
    optimizer = optim.Adam(model.parameters(), lr=lr)
    criterion = nn.CrossEntropyLoss()

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

    for epoch in range(epochs):
        model.train()
        total_loss = 0
        for batch_idx, (data, target) in enumerate(train_loader):
            # MOVE DATA TO GPU
            data, target = data.to(device), target.to(device)
            
            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_samples = 0
    
    CCT_THRESHOLD = 5.0 

    with torch.no_grad():
        for data, target in test_loader:
            # MOVE DATA TO GPU
            data, target = data.to(device), target.to(device)
            
            logits, pressure = model(data)
            pred = logits.argmax(dim=1)
            
            collapsed_mask = (pressure.squeeze() >= CCT_THRESHOLD)
            correct_mask = (pred == target)
            
            num_correct_collapsed += (correct_mask & collapsed_mask).sum().item()
            num_unresolved += (~collapsed_mask).sum().item()
            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()
