본문 바로가기
Machine-Learning/Basic

혼합 정밀도와 수치 포맷 (fp16, bf16, fp8)

by AtoN 2022. 11. 30.

혼합 정밀도와 수치 포맷 (fp16, bf16, fp8)

학습 중 loss가 갑자기 NaN으로 터지는 사고가 있습니다. 저정밀도 학습에서는 표현 범위를 벗어난 값이 inf가 되거나, 아주 작은 gradient가 0으로 사라지는 수치 문제가 실제 원인일 수 있습니다.

다만 NaN의 원인을 dtype 하나로 단정하면 안 됩니다. 너무 큰 learning rate, 잘못된 loss, 데이터의 NaN/Inf, optimizer 설정, 분산 통신 오류도 같은 증상을 만듭니다. 이 문서에서는 그중 수치 포맷이 만드는 문제와 혼합 정밀도가 이를 어떻게 완화하는지에 집중합니다.


1장. 정밀도를 낮추는 이유

이득이 세 군데에서 납니다.

항목 이득
메모리 절반이면 두 배 큰 모델이 들어간다
대역폭 읽고 쓰는 양이 절반이다
연산 속도 전용 하드웨어 유닛이 훨씬 빠르다

Mixed Precision Training 초록의 표현입니다.

이 접근으로 딥러닝 모델의 메모리 소비를 거의 2배 줄일 수 있다.

대신 잃는 것. 표현할 수 있는 값의 범위가 좁아지고, 값 사이의 촘촘함도 떨어집니다.


2장. 부동소수점 구조

부동소수점 수는 세 부분으로 이루어집니다.

부호(sign)     양수인가 음수인가
지수(exponent) 크기의 자릿수. 범위를 결정
가수(mantissa) 세부 값. 정밀도를 결정

포맷마다 이 비트를 다르게 나눕니다.

fp32   부호 1  지수 8   가수 23
fp16   부호 1  지수 5   가수 10
bf16   부호 1  지수 8   가수 7

지수 비트가 만드는 차이. fp16은 fp32보다 지수가 3비트 적어 범위가 크게 좁습니다. bf16은 fp32와 지수가 같아 범위가 똑같은 대신 가수가 7비트뿐입니다. bf16의 b는 brain으로, Google Brain이 딥러닝에 맞춰 설계한 배분입니다.

결국 fp16은 정밀도를 사고 범위를 팔았고, bf16은 범위를 사고 정밀도를 팔았습니다. 딥러닝에서는 범위가 더 중요했습니다.


3장. fp16의 문제

그래디언트 언더플로. 그래디언트는 매우 작은 값입니다. fp16의 표현 가능한 최소값보다 작으면 그냥 0이 되고, 학습이 멈춘 것처럼 보입니다.

활성값 오버플로. 반대로 값이 크면 표현 범위를 넘습니다. inf 가 되고 곧 NaN으로 번집니다.

작은 갱신의 소실. 파라미터가 1.0인데 갱신량이 0.0001이면 fp16의 정밀도로는 더해도 1.0 그대로입니다. 학습이 진행되지 않습니다.


4장. 혼합 정밀도의 처방

fp32 마스터 사본. 초록의 첫 처방입니다.

첫째, 우리는 각 옵티마이저 스텝 이후 그래디언트를 누적하는 단정밀도 가중치 사본을 유지할 것을 권장한다. 이 단정밀도 사본은 학습 중 반정밀도 형식으로 반올림된다.

연산은 fp16으로 하고 가중치 갱신은 fp32 사본에 합니다. 그러면 작은 갱신이 사라지지 않고 누적됩니다.

손실 스케일링. 초록의 두 번째 처방입니다.

둘째, 우리는 반정밀도 그래디언트에서의 정보 손실을 다루기 위해 손실을 적절히 스케일링할 것을 제안한다.

역전파 전에 손실에 큰 수를 곱하면 그래디언트가 전부 그만큼 커져 언더플로를 벗어납니다. 갱신 전에 다시 나눠 원래 크기로 돌립니다.

동적 손실 스케일링. 스케일 값을 고정하지 않고 자동으로 조절합니다. inf 나 NaN이 나오면 스케일을 줄이고 그 스텝을 건너뛰며, 안정적이면 서서히 키웁니다. PyTorch AMP의 GradScaler 가 이 일을 합니다.

민감한 연산은 더 높은 정밀도로. 정규화, softmax, loss reduction처럼 수치적으로 민감한 연산은 프레임워크의 mixed-precision policy가 FP32 등 더 높은 정밀도로 실행하는 경우가 많습니다. 모든 normalization과 softmax가 항상 FP32라는 뜻은 아니고, kernel과 dtype 정책에 따라 달라집니다. PyTorch AMP의 autocast가 연산별로 적절한 dtype을 선택하는 이유가 여기에 있습니다.


5장. bf16

손실 스케일링이 빠지는 이유. 지수가 fp32와 같아 범위가 같습니다. 그래서 언더플로와 오버플로가 거의 나지 않고, 손실 스케일링이 대체로 필요 없습니다.

가수 7비트의 대가. 가수가 7비트라 fp16의 10비트보다 적고, 같은 크기의 값들을 구분하는 능력이 떨어집니다. 그런데 학습에서는 이것이 별로 문제가 되지 않았습니다.

그래디언트는 원래 잡음이 많아 소수점 아래 몇 자리가 더 정확해도 이득이 작습니다. 반면 값이 범위를 벗어나 NaN이 되는 것은 치명적입니다.

LLM 학습에서 널리 쓰이는 배경. 최근 대규모 Transformer 학습에서는 BF16이 매우 흔합니다. FP16보다 넓은 동적 범위 덕분에 loss scaling 의존도를 줄일 수 있기 때문입니다. 다만 모든 공개 LLM 가중치가 BF16인 것은 아니며, 저장 dtype과 실제 compute dtype도 다를 수 있습니다. Ling-3.0-tiny처럼 BF16을 명시한 모델은 그 한 사례입니다.

torch_dtype: bfloat16

하드웨어 조건. NVIDIA GPU에서는 Ampere 세대부터 BF16 Tensor Core 연산이 본격 지원됩니다. 이전 세대나 다른 가속기에서는 지원 범위가 다르므로 실제 하드웨어를 확인해야 합니다. BF16을 쓸 수 없다고 항상 FP16만 가능한 것도 아니고, FP16을 쓴다면 loss scaling 필요성이 커집니다.


6장. fp8

E4M3와 E5M2. FP8 Formats for Deep Learning(2022)은 E4M3(4비트 지수, 3비트 가수)와 E5M2(5비트 지수, 2비트 가수) 두 형식을 제안했습니다.

E4M3는 상대적으로 정밀도 쪽, E5M2는 동적 범위 쪽에 유리합니다. 논문과 초기 Hopper 계열 mixed-FP8 recipe에서는 이 차이를 이용해 다음처럼 역할을 나누는 구성이 대표적이었습니다.

forward의 weight / activation   → E4M3를 주로 사용
backward의 gradient             → 더 넓은 범위가 필요한 경우 E5M2 사용

중요한 것은 FP8은 dtype 이름만 바꿔서는 안정적으로 쓰기 어렵고 scaling이 핵심이라는 점입니다. 8비트 안에 값을 넣기 위해 tensor의 amax를 추적하고 scale을 조절합니다.

원 논문은 다양한 이미지·언어 과제와 최대 175B 파라미터 언어 모델 실험에서 FP8이 16비트 기준선과 비슷한 품질을 낼 수 있음을 보고했습니다. 이 결과는 FP8이 단순 추론 양자화가 아니라 학습에도 사용할 수 있다는 근거가 됐습니다.

2026년의 변화: block scaling. 초기 FP8 recipe는 tensor 전체에 scale 하나를 두는 per-tensor scaling을 많이 사용했습니다. tensor 안에 큰 값과 작은 값이 섞이면 하나의 scale로 둘을 모두 살리기 어렵습니다. 최근 하드웨어는 scale을 더 작은 block 단위로 나눕니다.

NVIDIA Blackwell의 MXFP8은 32개 연속 값마다 E8M0 scale 하나를 두고, 기본적으로 forward와 backward 모두 E4M3를 사용할 수 있습니다. 세밀한 scaling이 E5M2가 담당하던 동적 범위 부담을 줄인 것입니다.

전통적 FP8
   tensor 전체 → scale 1개

MXFP8
   32개 값 → scale 1개
   32개 값 → scale 1개
   ...

Blackwell 계열에서는 FP4까지 내려간 NVFP4 recipe도 등장했습니다. 흐름은 단순히 FP32 → FP16 → FP8 → FP4로 비트를 줄이는 것이 아니라, block scaling과 rounding 같은 보정 기법을 함께 발전시켜 저정밀도의 오차를 통제하는 방향입니다.

INT8과의 차이. INT8은 정수 양자화이고 FP8은 부동소수점입니다. 둘 다 8비트지만 값의 배치 방식과 scaling, 사용할 수 있는 kernel이 다릅니다. 특정 모델이 INT8에서 품질 손실이 컸다고 해서 FP8에서도 같은 결과가 난다고 볼 수 없습니다.


7장. 실무

NaN이 떴을 때는 이 순서로 봅니다.

① dtype 이 fp16 인가. bf16 으로 바꿀 수 있나
② 손실 스케일링이 켜져 있는가
③ 그래디언트 클리핑이 켜져 있는가
④ 수치적으로 민감한 연산이 어떤 dtype 으로 계산되는가
⑤ 학습률이 너무 큰가
⑥ 데이터에 이상값이 있는가

FP16의 overflow/underflow가 원인이었다면 BF16 전환으로 안정성이 좋아질 수 있습니다. 하지만 NaN의 원인이 learning rate, 데이터, loss 구현 등에 있다면 dtype을 바꿔도 해결되지 않으므로 최초의 inf/NaN 발생 위치를 함께 추적해야 합니다.

config에서 볼 것. torch_dtype 또는 최신 Transformers의 dtype 관련 설정은 가중치를 어떤 dtype으로 저장·로드할지 판단하는 단서입니다. 하지만 이 값만 보고 실제 GEMM의 compute dtype을 단정하면 안 됩니다. autocast, quantization 설정, kernel 구현에 따라 저장 dtype과 계산 dtype이 달라질 수 있습니다.

학습과 추론의 차이. 역사적인 FP16 mixed-precision recipe는 FP32 master weight를 유지했지만 현대 optimizer와 framework의 내부 상태 표현은 구현마다 다릅니다. 핵심은 학습에는 backward와 optimizer state까지 안정적으로 유지할 정밀도가 필요하다는 점입니다. 추론은 backward와 optimizer가 없어서 더 공격적인 FP8·INT8·INT4 같은 저정밀도 기법을 적용하기 쉽습니다.

dtype 변환 시 주의점. BF16으로 저장된 값을 FP16으로 단순 변환하면 FP16의 더 좁은 범위를 넘어 overflow할 수 있습니다. FP16에서 BF16으로 옮기면 범위 측면에서는 여유가 생기지만 BF16의 가수가 더 짧아 정밀도 손실은 생길 수 있습니다. 어느 방향도 '비트 수가 같으니 동일하다'고 보면 안 됩니다.

재현성. 저정밀도 연산은 누적 순서에 따라 결과가 달라집니다. 그래서 같은 시드로 돌려도 비트 단위로 같지 않을 수 있습니다.


8장. 정리

FP16은 좁은 동적 범위를 loss scaling과 고정밀도 update 경로로 보완해 왔고, BF16은 FP32와 같은 지수 폭을 유지해 대규모 학습에서 널리 쓰입니다. FP8은 E4M3/E5M2 형식뿐 아니라 scaling recipe가 성패를 좌우하며, 최근에는 MXFP8처럼 block 단위 scaling으로 더 세밀하게 값을 맞추는 방향으로 발전했습니다.

포맷을 한 표로 모으면 이렇습니다.

포맷 부호 지수 가수 특징
fp32 1 8 23 기준 정밀도. optimizer/update 경로 등에 사용 가능
fp16 1 5 10 BF16보다 정밀하지만 범위 좁음
bf16 1 8 7 FP32와 같은 지수 폭. LLM 학습에서 널리 사용
fp8 E4M3 1 4 3 상대적으로 정밀도 우선. 초기 recipe에서 forward에 흔함
fp8 E5M2 1 5 2 상대적으로 범위 우선. 초기 recipe에서 gradient에 흔함

FP16과 BF16은 같은 16비트라도 지수와 가수 배분이 달라 성격이 크게 다릅니다. BF16에서 loss scaling이 대체로 불필요하다는 것도 절대 규칙은 아닙니다. FP8과 INT8 역시 같은 8비트라는 이유만으로 같은 양자화 방식으로 취급하면 안 됩니다.


용어 정리

용어 한 줄 뜻
혼합 정밀도 연산은 저정밀도로, 가중치 갱신은 고정밀도로 하는 학습 방식
부호 / 지수 / 가수 부동소수점의 세 구성 요소. 지수가 범위, 가수가 정밀도를 결정
fp32 32비트 단정밀도. 지수 8, 가수 23
fp16 16비트 반정밀도. 지수 5, 가수 10
bf16 16비트. 지수 8로 fp32와 범위가 같고 가수는 7
언더플로 값이 표현 최소값보다 작아 0이 되는 현상
오버플로 값이 표현 최대값을 넘어 무한대가 되는 현상
NaN 정의되지 않은 수. 한 번 생기면 계산 전체로 번짐
마스터 가중치 갱신을 누적하기 위해 유지하는 fp32 가중치 사본
손실 스케일링 역전파 전 손실에 큰 수를 곱해 그래디언트 언더플로를 막는 기법
GradScaler PyTorch AMP에서 동적 손실 스케일링을 담당하는 객체
E4M3 / E5M2 fp8의 두 인코딩. 각각 정밀도 쪽과 범위 쪽
torch_dtype config에서 가중치 저장 형식을 지정하는 필드

참고자료

댓글