import torch
import torch.nn as nn
import torch.nn.functional as F
from torch.utils.data import DataLoader
from torchvision import datasets, transforms
from typing import List, Optional, Tuple, Union

# ==========================================
# 1. nRAN-T Model Definition (Adapted for 2D)
# ==========================================

class nRANTClassifier2D(nn.Module):
    """
    Multi-Stream Rational-Addition Network for MNIST (2D Adaptation).
    Instead of temporal scales, we use spatial scales (dilations).
    """
    def __init__(
        self,
        num_classes: int,
        in_channels: int = 1,
        hidden_dim: int = 32,
        scales: List[int] = [1, 2, 4, 8],
        max_n: Optional[int] = None,
        dropout: float = 0.1
    ):
        super().__init__()
        self.num_classes = num_classes
        self.scales = scales[:max_n] if max_n else scales
        self.max_n = len(self.scales)

        # Initial projection to hidden dimension
        self.input_proj = nn.Conv2d(in_channels, hidden_dim, kernel_size=1)

        # Parallel stream encoders (Rational Terms)
        self.streams = nn.ModuleList()
        for b in self.scales:
            stream = nn.Sequential(
                # Use dilation=b to capture different spatial scales
                nn.Conv2d(hidden_dim, hidden_dim, 3, padding=b, dilation=b),
                nn.BatchNorm2d(hidden_dim),
                nn.ReLU(inplace=True),
                nn.Dropout2d(dropout),
                nn.Conv2d(hidden_dim, hidden_dim, 3, padding=b, dilation=b),
                nn.BatchNorm2d(hidden_dim),
                nn.ReLU(inplace=True),
                nn.Conv2d(hidden_dim, hidden_dim, 1) # numerator projection
            )
            self.streams.append(stream)

        # Learnable denominators (the 'b' in a/b)
        self.denominators = nn.Parameter(
            torch.tensor(self.scales, dtype=torch.float32)
        )

        # Base offset (class prior 'c')
        self.base_offset = nn.Parameter(torch.zeros(num_classes))

        # Collapse: Global Average Pooling -> MLP
        self.pool = nn.AdaptiveAvgPool2d(1)
        self.head = nn.Sequential(
            nn.Linear(hidden_dim, hidden_dim),
            nn.ReLU(inplace=True),
            nn.Dropout(dropout),
            nn.Linear(hidden_dim, num_classes)
        )

    def forward(self, x: torch.Tensor, return_terms: bool = False) -> Union[torch.Tensor, Tuple[torch.Tensor, List[torch.Tensor]]]:
        # Project input: (B, 1, 28, 28) -> (B, H, 28, 28)
        h = self.input_proj(x)

        terms = []
        for i, stream in enumerate(self.streams):
            # a_i = stream(h)
            # term_i = a_i / b_i
            a_i = stream(h)
            b_i = self.denominators[i].abs() + 1e-6
            terms.append(a_i / b_i)

        # nRAN Addition: Sum of all scale-separated terms
        fused = torch.stack(terms, dim=0).sum(dim=0) # (B, H, 28, 28)

        # Collapse over spatial dimensions
        pooled = self.pool(fused).view(x.size(0), -1) # (B, H)

        # Final classification + base offset
        logits = self.head(pooled) + self.base_offset

        if return_terms:
            return logits, terms
        return logits

    def capacity_regularization(self, terms: List[torch.Tensor]) -> torch.Tensor:
        """Encourages small terms to vanish to simplify the network."""
        loss = 0.0
        for i, term in enumerate(terms):
            loss = loss + term.abs().mean() / (self.denominators[i].abs() + 1e-6)
        return loss / len(terms)

# ==========================================
# 2. Training Logic
# ==========================================

def train_nran(model, train_loader, optimizer, criterion, device, lambda_cap=1e-3):
    model.train()
    total_loss = 0
    correct = 0
    total = 0

    for data, target in train_loader:
        data, target = data.to(device), target.to(device)
        optimizer.zero_grad()

        logits, terms = model(data, return_terms=True)
        
        # Standard CE Loss + nRAN Capacity Regularization
        ce_loss = criterion(logits, target)
        cap_loss = model.capacity_regularization(terms)
        loss = ce_loss + lambda_cap * cap_loss
        
        loss.backward()
        optimizer.step()

        total_loss += loss.item()
        pred = logits.argmax(dim=1, keepdim=True)
        correct += pred.eq(target.view_as(pred)).sum().item()
        total += len(data)

    return total_loss / len(train_loader), correct / total

def test_nran(model, test_loader, criterion, device):
    model.eval()
    test_loss = 0
    correct = 0
    total = 0
    with torch.no_grad():
        for data, target in test_loader:
            data, target = data.to(device), target.to(device)
            logits = model(data)
            test_loss += criterion(logits, target).item()
            pred = logits.argmax(dim=1, keepdim=True)
            correct += pred.eq(target.view_as(pred)).sum().item()
            total += len(data)
    
    return test_loss / len(test_loader), correct / total

# ==========================================
# 3. Execution Execution
# ==========================================

if __name__ == "__main__":
    # Setup
    device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
    BATCH_SIZE = 64
    EPOCHS = 5
    LR = 0.001

    # Data
    transform = transforms.Compose([transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,))])
    train_set = datasets.MNIST('./data', train=True, download=True, transform=transform)
    test_set = datasets.MNIST('./data', train=False, transform=transform)
    train_loader = DataLoader(train_set, batch_size=BATCH_SIZE, shuffle=True)
    test_loader = DataLoader(test_set, batch_size=BATCH_SIZE, shuffle=False)

    # Model: 4 scales as per nRAN-T theory
    model = nRANTClassifier2D(num_classes=10, scales=[1, 2, 4, 8]).to(device)
    optimizer = torch.optim.AdamW(model.parameters(), lr=LR)
    criterion = nn.CrossEntropyLoss()

    print(f"Training nRAN-T on MNIST using device: {device}...")
    for epoch in range(1, EPOCHS + 1):
        train_loss, train_acc = train_nran(model, train_loader, optimizer, criterion, device)
        test_loss, test_acc = test_nran(model, test_loader, criterion, device)
        
        print(f"Epoch {epoch}/{EPOCHS} | Train Loss: {train_loss:.4f} Acc: {train_acc:.4f} | Test Loss: {test_loss:.4f} Acc: {test_acc:.4f}")

    print("\nFinal Test Accuracy:", f"{test_acc*100:.2f}%")