import torch from torch import nn torch.manual_seed(7) device_type = "cuda" if torch.cuda.is_available() else "cpu" device = torch.device(device_type) amp_dtype = torch.float16 if device_type == "cuda" else torch.bfloat16 if device_type == "cpu": torch.backends.mkldnn.enabled = False model = nn.Sequential( nn.Linear(4, 8), nn.ReLU(), nn.Linear(8, 1), ).to(device) optimizer = torch.optim.SGD(model.parameters(), lr=0.05) loss_fn = nn.MSELoss() scaler = torch.amp.GradScaler(device_type, enabled=(amp_dtype == torch.float16)) inputs = torch.randn(16, 4, device=device) target = torch.randn(16, 1, device=device) before = model[0].weight.detach().clone() optimizer.zero_grad(set_to_none=True) with torch.amp.autocast(device_type=device_type, dtype=amp_dtype): prediction = model(inputs) loss = loss_fn(prediction, target) scaler.scale(loss).backward() gradients_finite = bool(torch.isfinite(model[0].weight.grad).all()) scaler.step(optimizer) scaler.update() weight_changed = not torch.equal(before, model[0].weight.detach()) print(f"device_type={device_type}") print(f"amp_dtype={amp_dtype}") print(f"scaler_enabled={scaler.is_enabled()}") print(f"prediction_dtype={prediction.dtype}") print(f"loss_dtype={loss.dtype}") print(f"gradients_finite={gradients_finite}") print(f"weight_changed={weight_changed}") print(f"loss_value={loss.detach().float().item():.6f}")