메모리는 절반, 속도는 두 배 가까이
혼합 정밀도 학습은 대부분의 연산을 float16이나 bfloat16으로 실행하고, 소프트맥스·정규화·손실 축소처럼 수치 오차에 민감한 일부 연산만 float32로 유지해. Ampere 이후 GPU와 Apple Silicon에서는 보통 1.5~2배 빨라지고 활성화 메모리는 약 절반으로 줄어.
현대적인 PyTorch API는 torch.amp야. 예전의 torch.cuda.amp는 사용 중단됐어. 구성 요소는 두 가지야:
autocast(device_type='cuda', dtype=...): 연산마다 알맞은 정밀도를 선택하는 문맥 관리자야.GradScaler(device_type): fp16에서만 필요해. 역전파 전에 손실 크기를 키워서 좁은 fp16 범위에서 기울기가 언더플로되지 않게 하고, 옵티마이저가 갱신하기 전에 원래 크기로 되돌려.
fp16과 bf16
- bfloat16: float32와 같은 지수 범위를 갖지만 가수부 정밀도는 낮아. Ampere 이후 NVIDIA GPU(A100, RTX 30/40, H100)와 Apple Silicon에서 쓸 수 있어. GradScaler는 필요하지 않아. 현대적인 기본 선택으로 권장해.
- float16: 지수 범위가 좁아서 GradScaler가 필요해. V100, T4, RTX 20 시리즈 같은 이전 세대 GPU에서는 여전히 중요해.
어디를 감싸야 할까
자동 형변환 문맥으로 순전파와 손실 계산을 감싸. 역전파와 옵티마이저 갱신은 그 문맥 밖에서 실행해.