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

import torch
from torch import nn

torch.manual_seed(72)
torch.set_num_threads(1)
model = nn.Linear(2, 1)
optimizer = torch.optim.SGD(model.parameters(), lr=.05)
scheduler = torch.optim.lr_scheduler.StepLR(optimizer, step_size=1, gamma=.9)
scaler = torch.amp.GradScaler("cpu", init_scale=128.)
X, y = torch.randn(11, 2), torch.randn(11, 1)
microbatches = [(X[i:i+3], y[i:i+3]) for i in range(0, 11, 3)]
successful, skipped = 0, 0
for attempt in range(3):
    optimizer.zero_grad(set_to_none=True)
    for xb, yb in microbatches:
        with torch.autocast(device_type="cpu", dtype=torch.float16):
            prediction = model(xb)
            loss = nn.functional.mse_loss(prediction, yb, reduction="sum")/len(X)
        scaler.scale(loss).backward()
    if attempt == 0:
        next(model.parameters()).grad.reshape(-1)[0] = float("inf")
    saved = [p.detach().clone() for p in model.parameters()]
    old_scale = scaler.get_scale()
    scaler.unscale_(optimizer)
    nn.utils.clip_grad_norm_(model.parameters(), 1.)
    scaler.step(optimizer)
    scaler.update()
    if scaler.get_scale() < old_scale:
        skipped += 1
        assert all(torch.equal(p, old) for p, old in zip(model.parameters(), saved))
    else:
        successful += 1
        scheduler.step()
assert skipped == 1 and successful == 2
assert scheduler.last_epoch == successful
assert all(torch.isfinite(p).all() for p in model.parameters())
print("successful / skipped updates:", successful, skipped)
