使用 自动混合精度(AMP, Automatic Mixed Precision)。
from torch.cuda.amp import autocast, GradScaler作用:
scaler = GradScaler()autocastfor x, y in dataloader:
optimizer.zero_grad()
with autocast(): # FP16 自动转换
y_pred = model(x)
loss = criterion(y_pred, y)✅ 在 autocast() 作用域内:
scaler.scale(loss) 进行反向传播scaler.scale(loss).backward()✅ 作用:
scaler.step(optimizer)
scaler.update()✅ scaler.step():
PyTorch AMP 自动维护 FP32 主权重:
autocast() 中 PyTorch 会自动处理⚠️ 注意:
验证阶段 不需要 AMP:
model.eval()
with torch.no_grad():
for x, y in val_loader:
y_pred = model(x)✅ 或显式关闭 autocast:
with autocast(enabled=False):
...from torch.cuda.amp import autocast, GradScaler
model = model.cuda()
optimizer = torch.optim.Adam(model.parameters())
criterion = torch.nn.CrossEntropyLoss()
scaler = GradScaler()
for epoch in range(epochs):
model.train()
for x, y in dataloader:
x, y = x.cuda(), y.cuda()
optimizer.zero_grad()
with autocast():
y_pred = model(x)
loss = criterion(y_pred, y)
scaler.scale(loss).backward()
scaler.step(optimizer)
scaler.update()tf.keras.mixed_precision.set_global_policy('mixed_float16')jax.numpy.float16如果你愿意,我也可以:
你更关心哪一部分?