from __future__ import annotations

from pathlib import Path

import torch
import torch.nn as nn
import torch.utils.data as Data


STEP3_DIR = Path("./step3")
MNIST_DIR = STEP3_DIR / "mnist"
MODEL_PATH = STEP3_DIR / "cnn.pkl"


class CNN(nn.Module):
    """2xConv+Pool -> FC (MNIST)."""

    def __init__(self):
        super().__init__()
        self.conv1 = nn.Sequential(
            nn.Conv2d(1, 16, 5, 1, 2), nn.ReLU(), nn.MaxPool2d(2)
        )
        self.conv2 = nn.Sequential(
            nn.Conv2d(16, 32, 5, 1, 2), nn.ReLU(), nn.MaxPool2d(2)
        )
        self.out = nn.Linear(32 * 7 * 7, 10)

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        x = self.conv1(x)
        x = self.conv2(x)
        x = x.view(x.size(0), -1)
        return self.out(x)


def _try_load_mnist_offline(limit_train: int = 6000, limit_test: int = 600):
    try:
        import torchvision
        from torchvision import transforms

        transform = transforms.ToTensor()
        train_set = torchvision.datasets.MNIST(root=str(MNIST_DIR), train=True, transform=transform, download=False)
        test_set = torchvision.datasets.MNIST(root=str(MNIST_DIR), train=False, transform=transform, download=False)
        train_subset = [(train_set[i][0], train_set[i][1]) for i in range(min(limit_train, len(train_set)))]
        test_subset = [(test_set[i][0], test_set[i][1]) for i in range(min(limit_test, len(test_set)))]
        return train_subset, test_subset
    except Exception as e:
        print(f"[MNIST] 离线读取失败：{e}")
        return None


def _make_synthetic(limit_train: int = 6000, limit_test: int = 600):
    x_train = torch.randn(limit_train, 1, 28, 28)
    y_train = torch.randint(0, 10, (limit_train,))
    x_test = torch.randn(limit_test, 1, 28, 28)
    y_test = torch.randint(0, 10, (limit_test,))
    return Data.TensorDataset(x_train, y_train), Data.TensorDataset(x_test, y_test)


def train_cnn(epochs: int = 2, batch_size: int = 64, lr: float = 1e-3, device: str | None = None) -> None:
    device = device or ("cuda" if torch.cuda.is_available() else "cpu")
    local = _try_load_mnist_offline()
    if local is None:
        train_ds, test_ds = _make_synthetic()
    else:
        xtr = torch.stack([img for img, _ in local[0]], 0)
        ytr = torch.tensor([int(lbl) for _, lbl in local[0]], dtype=torch.long)
        xte = torch.stack([img for img, _ in local[1]], 0)
        yte = torch.tensor([int(lbl) for _, lbl in local[1]], dtype=torch.long)
        train_ds = Data.TensorDataset(xtr, ytr)
        test_ds = Data.TensorDataset(xte, yte)

    train_loader = Data.DataLoader(train_ds, batch_size=batch_size, shuffle=True)
    test_loader = Data.DataLoader(test_ds, batch_size=batch_size, shuffle=False)

    model = CNN().to(device)
    criterion = nn.CrossEntropyLoss()
    opt = torch.optim.Adam(model.parameters(), lr=lr)

    for ep in range(1, epochs + 1):
        model.train()
        loss_sum = 0.0
        correct = 0
        total = 0
        for xb, yb in train_loader:
            xb, yb = xb.to(device), yb.to(device)
            logits = model(xb)
            loss = criterion(logits, yb)
            opt.zero_grad(); loss.backward(); opt.step()
            loss_sum += float(loss.item()) * xb.size(0)
            correct += int((logits.argmax(1) == yb).sum().item())
            total += int(xb.size(0))
        print(f"[CNN] Epoch {ep:02d} | Loss {loss_sum/total:.4f} | Acc {correct/total:.4f}")

        model.eval(); t_correct = 0; t_total = 0
        with torch.no_grad():
            for xb, yb in test_loader:
                xb, yb = xb.to(device), yb.to(device)
                pred = model(xb).argmax(1)
                t_correct += int((pred == yb).sum().item()); t_total += int(xb.size(0))
        print(f"[CNN] Test Acc {t_correct/max(1,t_total):.4f}")

    STEP3_DIR.mkdir(parents=True, exist_ok=True)
    torch.save(model.state_dict(), MODEL_PATH)
    print(f"[CNN] saved: {MODEL_PATH}")


if __name__ == "__main__":
    train_cnn()

