# Source: content/notes/deep-learning/cnns.md
# Independent CPU example; use the curriculum environment.
# See /notes/ml/#example-environment or /notes/deep-learning/#example-environment.

import copy
import numpy as np
import torch
from torch import nn
from sklearn.datasets import load_digits
from sklearn.model_selection import train_test_split

torch.set_num_threads(1)
torch.manual_seed(72)
np.random.seed(72)
data = load_digits()
images = data.images.astype(np.float32)[:, None, :, :] / 16.0
labels = data.target
train_x, hold_x, train_y, hold_y = train_test_split(
    images, labels, test_size=0.3, stratify=labels, random_state=73)
val_x, test_x, val_y, test_y = train_test_split(
    hold_x, hold_y, test_size=0.5, stratify=hold_y, random_state=74)
xt, yt = torch.tensor(train_x), torch.tensor(train_y, dtype=torch.long)
xv, yv = torch.tensor(val_x), torch.tensor(val_y, dtype=torch.long)
xs, ys = torch.tensor(test_x), torch.tensor(test_y, dtype=torch.long)
model = nn.Sequential(
    nn.Conv2d(1, 12, 3, padding=1), nn.ReLU(),
    nn.Conv2d(12, 24, 3, padding=1), nn.ReLU(), nn.MaxPool2d(2),
    nn.Flatten(), nn.Linear(24 * 4 * 4, 10)
)
assert model(xt[:7]).shape == (7, 10)
optimizer = torch.optim.AdamW(model.parameters(), lr=0.005, weight_decay=0.001)
generator = torch.Generator().manual_seed(75)
best = float("inf")
best_state = copy.deepcopy(model.state_dict())
for epoch in range(16):
    model.train()
    for indices in torch.randperm(len(xt), generator=generator).split(128):
        optimizer.zero_grad(set_to_none=True)
        loss = nn.functional.cross_entropy(model(xt[indices]), yt[indices])
        assert torch.isfinite(loss)
        loss.backward()
        optimizer.step()
    model.eval()
    with torch.inference_mode():
        validation = nn.functional.cross_entropy(model(xv), yv).item()
    if validation < best:
        best = validation
        best_state = copy.deepcopy(model.state_dict())
model.load_state_dict(best_state)
model.eval()
with torch.inference_mode():
    prediction = model(xs).argmax(1)
    accuracy = (prediction == ys).float().mean().item()
    test_loss = nn.functional.cross_entropy(model(xs), ys).item()
assert accuracy > 0.9
probe = xs[:1].detach().clone().requires_grad_()
score = model(probe)[0, int(ys[0])]
sensitivity, = torch.autograd.grad(score, probe)
assert sensitivity.shape == (1, 1, 8, 8)
assert torch.isfinite(sensitivity).all() and sensitivity.abs().sum() > 0
print("Parameters:", sum(p.numel() for p in model.parameters()))
print("Validation/test loss:", best, test_loss, "test accuracy:", accuracy)
print("Input sensitivity magnitude:", sensitivity.abs().sum().item())
