import argparse
import time

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

from ex03 import AnurupyenaDepthwiseSeparableConv


torch.manual_seed(42)
torch.set_num_threads(min(4, torch.get_num_threads()))


class MNISTAnurupyenaNet(nn.Module):
    def __init__(self):
        super().__init__()
        self.features = nn.Sequential(
            AnurupyenaDepthwiseSeparableConv(1, 32),
            nn.BatchNorm2d(32),
            nn.ReLU(inplace=True),
            AnurupyenaDepthwiseSeparableConv(32, 64),
            nn.BatchNorm2d(64),
            nn.ReLU(inplace=True),
            nn.MaxPool2d(2),
            nn.Dropout2d(0.05),
            AnurupyenaDepthwiseSeparableConv(64, 96),
            nn.BatchNorm2d(96),
            nn.ReLU(inplace=True),
            AnurupyenaDepthwiseSeparableConv(96, 128),
            nn.BatchNorm2d(128),
            nn.ReLU(inplace=True),
            nn.MaxPool2d(2),
            nn.Dropout2d(0.10),
        )
        self.classifier = nn.Sequential(
            nn.Flatten(),
            nn.Linear(128 * 7 * 7, 128),
            nn.ReLU(inplace=True),
            nn.Dropout(0.20),
            nn.Linear(128, 10),
        )

    def forward(self, inputs):
        outputs = self.features(inputs)
        return self.classifier(outputs)


def benchmark_anurupyena_layer(device):
    layer = AnurupyenaDepthwiseSeparableConv(16, 32).to(device).eval()
    sample = torch.randn(1, 16, 28, 28, device=device)

    with torch.inference_mode():
        for _ in range(10):
            layer(sample)
        if device.type == "cuda":
            torch.cuda.synchronize()

        start = time.perf_counter()
        runs = 100
        for _ in range(runs):
            layer(sample)
        if device.type == "cuda":
            torch.cuda.synchronize()

    elapsed_ms = (time.perf_counter() - start) * 1000 / runs
    print(f"Anurupyena depthwise separable layer: {elapsed_ms:.2f} ms")


def make_loaders(batch_size, data_root):
    transform = transforms.Compose(
        [
            transforms.ToTensor(),
            transforms.Normalize((0.1307,), (0.3081,)),
        ]
    )

    train_dataset = datasets.MNIST(root=data_root, train=True, download=True, transform=transform)
    test_dataset = datasets.MNIST(root=data_root, train=False, download=True, transform=transform)

    train_loader = DataLoader(
        train_dataset,
        batch_size=batch_size,
        shuffle=True,
        num_workers=2,
        pin_memory=torch.cuda.is_available(),
    )
    test_loader = DataLoader(
        test_dataset,
        batch_size=batch_size,
        shuffle=False,
        num_workers=2,
        pin_memory=torch.cuda.is_available(),
    )
    return train_loader, test_loader


def train_one_epoch(model, loader, optimizer, criterion, device, scheduler=None):
    model.train()
    total_loss = 0.0
    total_correct = 0
    total_samples = 0

    for images, labels in loader:
        images = images.to(device, non_blocking=True)
        labels = labels.to(device, non_blocking=True)

        optimizer.zero_grad(set_to_none=True)
        logits = model(images)
        loss = criterion(logits, labels)
        loss.backward()
        optimizer.step()
        if scheduler is not None:
            scheduler.step()

        batch_size = labels.size(0)
        total_loss += loss.item() * batch_size
        total_correct += (logits.argmax(dim=1) == labels).sum().item()
        total_samples += batch_size

    return total_loss / total_samples, total_correct / total_samples


@torch.inference_mode()
def evaluate(model, loader, criterion, device):
    model.eval()
    total_loss = 0.0
    total_correct = 0
    total_samples = 0

    for images, labels in loader:
        images = images.to(device, non_blocking=True)
        labels = labels.to(device, non_blocking=True)

        logits = model(images)
        loss = criterion(logits, labels)

        batch_size = labels.size(0)
        total_loss += loss.item() * batch_size
        total_correct += (logits.argmax(dim=1) == labels).sum().item()
        total_samples += batch_size

    return total_loss / total_samples, total_correct / total_samples


def main():
    parser = argparse.ArgumentParser(description="MNIST model using Anurupyena depthwise separable convolutions.")
    parser.add_argument("--epochs", type=int, default=5)
    parser.add_argument("--batch-size", type=int, default=128)
    parser.add_argument("--lr", type=float, default=3e-3)
    parser.add_argument("--weight-decay", type=float, default=1e-4)
    parser.add_argument("--data-root", type=str, default="./data")
    parser.add_argument("--device", type=str, default="cuda" if torch.cuda.is_available() else "cpu")
    args = parser.parse_args()

    device = torch.device(args.device)
    train_loader, test_loader = make_loaders(args.batch_size, args.data_root)

    model = MNISTAnurupyenaNet().to(device)
    criterion = nn.CrossEntropyLoss()
    optimizer = torch.optim.AdamW(model.parameters(), lr=args.lr, weight_decay=args.weight_decay)
    scheduler = torch.optim.lr_scheduler.OneCycleLR(
        optimizer,
        max_lr=args.lr,
        epochs=args.epochs,
        steps_per_epoch=len(train_loader),
        pct_start=0.15,
        div_factor=10.0,
        final_div_factor=100.0,
    )

    benchmark_anurupyena_layer(device)

    for epoch in range(1, args.epochs + 1):
        train_loss, train_acc = train_one_epoch(model, train_loader, optimizer, criterion, device, scheduler)
        test_loss, test_acc = evaluate(model, test_loader, criterion, device)
        print(
            f"epoch {epoch:02d} | "
            f"train loss {train_loss:.4f} | train acc {train_acc:.4%} | "
            f"test loss {test_loss:.4f} | test acc {test_acc:.4%}"
        )


if __name__ == "__main__":
    main()
