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

try:
    torch.multiprocessing.set_sharing_strategy("file_system")
except RuntimeError:
    pass

# Scale factor for channel widths (60% of original)
SCALE = 0.6
NUM_WORKERS = 0

def scale(x):
    return int(x * SCALE) or 1  # Ensure at least 1 channel

# ============ ResNet Building Blocks ============

class BasicBlock(nn.Module):
    expansion = 1
    
    def __init__(self, in_planes, planes, stride=1, downsample=None):
        super().__init__()
        self.conv1 = nn.Conv2d(in_planes, planes, 3, stride=stride, padding=1, bias=False)
        self.bn1 = nn.BatchNorm2d(planes)
        self.conv2 = nn.Conv2d(planes, planes, 3, stride=1, padding=1, bias=False)
        self.bn2 = nn.BatchNorm2d(planes)
        self.downsample = downsample
    
    def forward(self, x):
        identity = x
        out = F.relu(self.bn1(self.conv1(x)))
        out = self.bn2(self.conv2(out))
        if self.downsample:
            identity = self.downsample(x)
        out += identity
        return F.relu(out)


class ResNet(nn.Module):
    def __init__(self, block, layers, num_classes=10, in_channels=1):
        super().__init__()
        self.in_planes = scale(64)
        
        # Initial conv
        self.conv1 = nn.Conv2d(in_channels, scale(64), 3, stride=1, padding=1, bias=False)
        self.bn1 = nn.BatchNorm2d(scale(64))
        
        # Residual layers
        self.layer1 = self._make_layer(block, scale(64), layers[0], stride=1)
        self.layer2 = self._make_layer(block, scale(128), layers[1], stride=2)
        self.layer3 = self._make_layer(block, scale(256), layers[2], stride=2)
        
        # Classifier
        self.avgpool = nn.AdaptiveAvgPool2d((1, 1))
        self.fc = nn.Linear(scale(256) * block.expansion, num_classes)
        
        # Weight init
        for m in self.modules():
            if isinstance(m, nn.Conv2d):
                nn.init.kaiming_normal_(m.weight, mode='fan_out', nonlinearity='relu')
            elif isinstance(m, nn.BatchNorm2d):
                nn.init.constant_(m.weight, 1)
                nn.init.constant_(m.bias, 0)
    
    def _make_layer(self, block, planes, blocks, stride=1):
        downsample = None
        if stride != 1 or self.in_planes != planes * block.expansion:
            downsample = nn.Sequential(
                nn.Conv2d(self.in_planes, planes * block.expansion, 1, stride=stride, bias=False),
                nn.BatchNorm2d(planes * block.expansion)
            )
        
        layers = [block(self.in_planes, planes, stride, downsample)]
        self.in_planes = planes * block.expansion
        for _ in range(1, blocks):
            layers.append(block(self.in_planes, planes))
        
        return nn.Sequential(*layers)
    
    def forward(self, x):
        x = F.relu(self.bn1(self.conv1(x)))
        x = self.layer1(x)
        x = self.layer2(x)
        x = self.layer3(x)
        x = self.avgpool(x)
        x = torch.flatten(x, 1)
        return self.fc(x)


def resnet18(in_channels=1, num_classes=10):
    return ResNet(BasicBlock, [2, 2, 2], num_classes, in_channels)


# ============ Training Function ============

def train_model(model, train_loader, epochs, device):
    criterion = nn.CrossEntropyLoss()
    optimizer = optim.Adam(model.parameters(), lr=0.001)
    scheduler = optim.lr_scheduler.StepLR(optimizer, step_size=5, gamma=0.5)
    
    model.train()
    for epoch in range(epochs):
        total_loss, correct, total = 0, 0, 0
        for batch_idx, (data, target) in enumerate(train_loader):
            data, target = data.to(device), target.to(device)
            
            optimizer.zero_grad()
            output = model(data)
            loss = criterion(output, target)
            loss.backward()
            optimizer.step()
            
            total_loss += loss.item()
            pred = output.argmax(dim=1)
            correct += pred.eq(target).sum().item()
            total += target.size(0)
            
            if batch_idx % 2 == 0:
                print(f"  Batch {batch_idx}/{len(train_loader)} - Loss: {loss.item():.4f}")
        
        scheduler.step()
        acc = 100. * correct / total
        print(f"  Epoch {epoch+1}/{epochs}: Loss={total_loss/len(train_loader):.4f}, Acc={acc:.2f}%")
    
    return model


def evaluate(model, loader, device):
    model.eval()
    correct, total = 0, 0
    with torch.no_grad():
        for data, target in loader:
            data, target = data.to(device), target.to(device)
            output = model(data)
            pred = output.argmax(dim=1)
            correct += pred.eq(target).sum().item()
            total += target.size(0)
    return 100. * correct / total


# ============ Main ============

if __name__ == "__main__":
    device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
    print(f"Using device: {device}")
    print(f"ResNet scale: {SCALE}")
    
    # Transforms
    mnist_transform = transforms.Compose([
        transforms.Resize(32),
        transforms.Grayscale(num_output_channels=1),
        transforms.ToTensor(),
        transforms.Normalize((0.1307,), (0.3081,))
    ])
    
    cifar_transform = transforms.Compose([
        transforms.Resize(32),
        transforms.Grayscale(num_output_channels=1),
        transforms.ToTensor(),
        transforms.Normalize((0.5,), (0.5,))
    ])
    
    # Datasets
    print("\n=== Loading MNIST ===")
    mnist_full = datasets.MNIST('../data', train=True, download=True, transform=mnist_transform)
    mnist_test = datasets.MNIST('../data', train=False, transform=mnist_transform)
    mnist_train, mnist_val = random_split(mnist_full, [50000, 10000])
    
    print("=== Loading CIFAR10 ===")
    cifar_full = datasets.CIFAR10('../data', train=True, download=True, transform=cifar_transform)
    cifar_test = datasets.CIFAR10('../data', train=False, transform=cifar_transform)
    
    # Create combined dataset (MNIST + CIFAR10) with distinct labels
    # MNIST labels: 0-9, CIFAR10 labels: 0-9 but we remap to 10-19
    # Actually, let's just train separate models and merge losses
    # For simplicity: train on MNIST first, then CIFAR10, then MNIST again
    
    # Phase 1: Train on combined dataset (alternating batches)
    print("\n========== Phase 1: Joint Training (MNIST + CIFAR10) ==========")
    model = resnet18(in_channels=1, num_classes=10).to(device)
    
    mnist_loader = DataLoader(mnist_train, batch_size=128, shuffle=True, num_workers=NUM_WORKERS)
    cifar_loader = DataLoader(cifar_full, batch_size=128, shuffle=True, num_workers=NUM_WORKERS)
    
    criterion = nn.CrossEntropyLoss()
    optimizer = optim.Adam(model.parameters(), lr=0.001)
    scheduler = optim.lr_scheduler.StepLR(optimizer, step_size=2, gamma=0.5)
    
    model.train()
    epochs_joint = 3
    
    for epoch in range(epochs_joint):
        print(f"\nEpoch {epoch+1}/{epochs_joint}")
        total_loss, correct, total = 0, 0, 0
        
        paired_batches = zip(mnist_loader, cifar_loader)

        for batch_idx, ((data_m, target_m), (data_c, target_c)) in enumerate(paired_batches):
            data_m, target_m = data_m.to(device), target_m.to(device)
            data_c, target_c = data_c.to(device), target_c.to(device)

            optimizer.zero_grad()

            out_m = model(data_m)
            loss_m = criterion(out_m, target_m)

            out_c = model(data_c)
            loss_c = criterion(out_c, target_c)

            loss = 0.5 * (loss_m + loss_c)
            loss.backward()
            optimizer.step()

            total_loss += loss.item()

            pred_m = out_m.argmax(dim=1)
            pred_c = out_c.argmax(dim=1)
            correct += pred_m.eq(target_m).sum().item()
            correct += pred_c.eq(target_c).sum().item()
            total += target_m.size(0) + target_c.size(0)

            if batch_idx % 100 == 0:
                print(
                    f"  Batch {batch_idx}/{min(len(mnist_loader), len(cifar_loader))} "
                    f"- Loss: {loss.item():.4f}"
                )

        scheduler.step()
        acc = 100.0 * correct / total
        avg_loss = total_loss / min(len(mnist_loader), len(cifar_loader))
        print(f"  Epoch {epoch+1}/{epochs_joint}: Loss={avg_loss:.4f}, Acc={acc:.2f}%")

    print("\n========== Phase 2: Evaluation ==========")
    mnist_test_loader = DataLoader(mnist_test, batch_size=256, shuffle=False, num_workers=NUM_WORKERS)
    cifar_test_loader = DataLoader(cifar_test, batch_size=256, shuffle=False, num_workers=NUM_WORKERS)

    mnist_acc = evaluate(model, mnist_test_loader, device)
    cifar_acc = evaluate(model, cifar_test_loader, device)

    print(f"MNIST test accuracy: {mnist_acc:.2f}%")
    print(f"CIFAR10 test accuracy: {cifar_acc:.2f}%")
