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

import copy
import math
import numpy as np
import torch
from torch import nn
from sklearn.datasets import make_moons
from sklearn.model_selection import train_test_split
from sklearn.preprocessing import StandardScaler

torch.set_num_threads(1)
torch.manual_seed(51)
np.random.seed(51)
base = nn.Sequential(nn.Linear(3, 5), nn.Tanh(), nn.Linear(5, 1)).double()
micro = copy.deepcopy(base)
x = torch.randn(11, 3, dtype=torch.float64)
y = torch.randn(11, 1, dtype=torch.float64)
nn.functional.mse_loss(base(x), y).backward()
for indices in torch.arange(11).split(4):
    (nn.functional.mse_loss(micro(x[indices]), y[indices], reduction="sum") / 11).backward()
for whole_parameter, micro_parameter in zip(base.parameters(), micro.parameters()):
    torch.testing.assert_close(whole_parameter.grad, micro_parameter.grad, atol=1e-12, rtol=1e-12)

X, labels = make_moons(n_samples=602, noise=0.16, random_state=51)
X_train, X_hold, y_train, y_hold = train_test_split(
    X, labels, test_size=202, stratify=labels, random_state=52
)
X_val, X_test, y_val, y_test = train_test_split(
    X_hold, y_hold, test_size=101, stratify=y_hold, random_state=53
)
scaler = StandardScaler().fit(X_train)
def convert(a, b):
    return (torch.tensor(scaler.transform(a), dtype=torch.float32),
            torch.tensor(b, dtype=torch.long))
xt, yt = convert(X_train, y_train)
xv, yv = convert(X_val, y_val)
xs, ys = convert(X_test, y_test)
torch.manual_seed(54)
model = nn.Sequential(nn.Linear(2, 24), nn.Tanh(), nn.Linear(24, 2))
normalization_types = (nn.LayerNorm, nn.GroupNorm, nn.BatchNorm1d,
                       nn.BatchNorm2d, nn.BatchNorm3d)
no_decay_ids = set()
for module in model.modules():
    for name, parameter in module.named_parameters(recurse=False):
        if name == "bias" or isinstance(module, normalization_types):
            no_decay_ids.add(id(parameter))
decay, no_decay = [], []
for parameter in model.parameters():
    if parameter.requires_grad:
        (no_decay if id(parameter) in no_decay_ids else decay).append(parameter)
optimizer = torch.optim.AdamW(
    [{"params": decay, "weight_decay": 0.01},
     {"params": no_decay, "weight_decay": 0.0}], lr=0.02
)
epochs, micro_size, accumulation = 35, 37, 3
updates_per_epoch = math.ceil(math.ceil(len(xt) / micro_size) / accumulation)
total_updates = epochs * updates_per_epoch
warmup_updates = 8
generator = torch.Generator().manual_seed(55)
updates = 0
best = float("inf")
best_state = copy.deepcopy(model.state_dict())
with torch.inference_mode():
    initial_loss = nn.functional.cross_entropy(model(xt), yt).item()
for epoch in range(epochs):
    model.train()
    batches = list(torch.randperm(len(xt), generator=generator).split(micro_size))
    seen = 0
    for offset in range(0, len(batches), accumulation):
        group = batches[offset:offset + accumulation]
        denominator = sum(len(indices) for indices in group)
        optimizer.zero_grad(set_to_none=True)
        for indices in group:
            loss_sum = nn.functional.cross_entropy(model(xt[indices]), yt[indices], reduction="sum")
            assert torch.isfinite(loss_sum)
            (loss_sum / denominator).backward()
            seen += len(indices)
        gradient_norm = nn.utils.clip_grad_norm_(model.parameters(), 1.0)
        assert torch.isfinite(gradient_norm)
        if updates < warmup_updates:
            ratio = (updates + 1) / warmup_updates
        else:
            progress = (updates - warmup_updates) / max(1, total_updates - warmup_updates - 1)
            ratio = 0.1 + 0.9 * (1 + math.cos(math.pi * min(progress, 1))) / 2
        for group_parameters in optimizer.param_groups:
            group_parameters["lr"] = 0.02 * ratio
        optimizer.step()
        updates += 1
    assert seen == len(xt)
    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())
assert updates == total_updates
model.load_state_dict(best_state)
model.eval()
with torch.inference_mode():
    final_loss = nn.functional.cross_entropy(model(xt), yt).item()
    accuracy = (model(xs).argmax(1) == ys).float().mean().item()
assert final_loss < initial_loss * 0.6
assert accuracy > 0.88
print("Updates:", updates, "initial/final training loss:", initial_loss, final_loss)
print("Selected validation loss:", best, "test accuracy:", accuracy)
print("Unequal-microbatch gradient equivalence and complete-tail checks passed.")
