LLM WikiAccess-protected knowledge portal

WIKI

FP8 혼합 정밀도 LLM 학습: Transformer Engine과 Delayed Scaling으로 H100·B200 학습 비용을 절반으로 줄이는 방법

요약 대형 언어 모델 사전 학습 비용의 핵심 결정 요인은 연산 정밀도다. BF16으로 H100 하나를 4주 동안 돌리면 FP8로 2주면 끝난다는 단순 계산이 맞다. 그러나 FP8은 "정밀도가 낮은 BF16"이 아니다. 표현 범위·그래디언트 안정성·행렬 곱 누적 방식 이 근본적으로 다르기 때문에, 단순히 dtype만 바꾸면 학습이 발산하거나 품질이 하락한다. NVIDIA Transformer Engine 은 이 문제를 해결하기

경로human/study/content/ai-frontier/137-fp8-mixed-precision-llm-training-transformer-engine-delayed-scaling.md
카테고리Study
태그#ai-review #delayed #engine #scaling #study #training #transformer

요약

대형 언어 모델 사전 학습 비용의 핵심 결정 요인은 연산 정밀도다. BF16으로 H100 하나를 4주 동안 돌리면 FP8로 2주면 끝난다는 단순 계산이 맞다. 그러나 FP8은 "정밀도가 낮은 BF16"이 아니다. 표현 범위·그래디언트 안정성·행렬 곱 누적 방식이 근본적으로 다르기 때문에, 단순히 dtype만 바꾸면 학습이 발산하거나 품질이 하락한다.

NVIDIA Transformer Engine은 이 문제를 해결하기 위해 만들어진 라이브러리다. FP8 연산 단위(GEMM, LayerNorm, attention)를 블록별로 분리하고, 각 텐서에 스케일 팩터(scale factor)를 붙여 수치 범위를 동적으로 관리한다. DeepSeek-V3, Llama 3.3, Phi-4 등 2025~2026년의 주요 오픈 모델 대부분이 FP8 혼합 정밀도 학습을 채택했다.

핵심 요약:


배경: 왜 BF16으로는 충분하지 않은가

정밀도 포맷의 역사

LLM 학습 정밀도의 진화는 가용 GPU 하드웨어와 함께 진행됐다.

포맷지수 비트가수 비트동적 범위주요 용도
FP32823매우 넓음옵티마이저 마스터 가중치
FP16510좁음A100 이전 혼합 정밀도
BF1687FP32와 동등H100 기본 학습 포맷
FP8 E4M343좁음, 순전파용H100/B200 순전파
FP8 E5M252넓음, 역전파용H100/B200 그래디언트

BF16은 FP32와 같은 지수 비트를 가지므로 오버플로우 없이 대부분의 학습을 안정적으로 수행한다. 그러나 H100의 Tensor Core는 FP8 연산 시 2배의 FLOP/s를 제공한다. H100 SXM5 기준 BF16 989 TFLOP/s → FP8 1979 TFLOP/s. 이것이 FP8 학습의 동기다.

FP8이 어려운 이유

FP8 E4M3의 표현 범위는 [-448, 448]이다. LLM 학습에서 가중치·활성화·그래디언트의 절댓값은 이 범위를 훨씬 벗어난다. 단순히 dtype을 변경하면 오버플로우(NaN) 또는 언더플로우(0으로 소실)가 발생한다.

해결책은 텐서에 per-tensor 또는 per-block 스케일 팩터 s를 곱해서 FP8 범위 안으로 끌어들이는 것이다. x_fp8 = clip(x / s, -448, 448). 이 s 값을 어떻게 관리하느냐가 FP8 학습의 핵심이다.


Transformer Engine 아키텍처

설계 원칙

NVIDIA Transformer Engine은 세 가지 원칙으로 설계됐다.

  1. 레이어 단위 FP8 캡슐화: GEMM, LayerNorm, Attention 등 각 연산자를 FP8 입출력으로 감싸고, 스케일 팩터 관리를 내부에 숨긴다.
  2. Delayed Scaling: 이전 반복(iteration)의 통계로 다음 반복의 스케일을 결정한다. 실시간 abs-max 계산 오버헤드를 제거한다.
  3. 고정밀 마스터 가중치: 옵티마이저 상태(momentum, variance)와 마스터 가중치는 FP32 또는 BF16으로 유지한다.
FP8 혼합 정밀도 학습 데이터 흐름 (Transformer Engine) ▶ 순전파 (E4M3) 입력 X BF16 → FP8 E4M3 scale_fwd (지연 갱신) = amax(X[t-1])⁻¹ FP8 GEMM W×X (E4M3×E4M3→FP32 누적) LayerNorm BF16 출력 FlashAttention E4M3 (FA4 FP8) 손실 함수 FP32 출력 ◀ 역전파 (E5M2) dL/dY FP32 → FP8 E5M2 역전파 GEMM Wᵀ×dY (E5M2→FP32) dW 계산 dY×Xᵀ (FP32 누적) 옵티마이저 (AdamW) BF16/FP32 마스터 가중치 W 업데이트 FP32 → FP8 캐스트 Delayed Scaling — amax 이력 관리 (Transformer Engine 내부) 매 step 종료: amax_history[t] = max(|X|, |dX|) → scale[t+1] = margin / amax_history[t-window+1:t].max() window = 1024 (기본값) | 현재 step의 스케일은 이전 step 통계로 계산 → 추가 동기화 지점 불필요 핵심 구분 E4M3 (순전파): 정밀도 우선, 범위[-448,448] E5M2 (역전파): 범위 우선, 그래디언트 발산 방지 FP32/BF16: 마스터 가중치·옵티마이저 상태 ※ GEMM 누적(accumulate)은 항상 FP32: FP8 입력 × FP8 가중치 → FP32 부분합 → 최종 FP8 또는 BF16 출력
FP8 혼합 정밀도 학습 데이터 흐름 — Transformer Engine Delayed Scaling

Delayed Scaling 동작 원리

Transformer Engine의 Delayed Scaling은 다음 흐름으로 작동한다.

  1. amax 기록: 매 스텝 종료 시 해당 텐서의 절댓값 최대값 amaxamax_history 배열에 저장한다.
  2. 스케일 계산: 다음 스텝 시작 전, 지난 N 스텝(기본 1024)의 amax_history에서 최댓값을 꺼내 스케일 팩터를 계산한다. scale = fp8_max_val / (margin × amax_max).
  3. 적용: 계산된 스케일로 텐서를 FP8로 캐스팅한다.

margin은 미래 변동에 대비한 여유 계수다. NVIDIA 권고값은 1.0. 학습이 불안정하면 낮추고, FP8 활용도(포화 없음)가 낮으면 올린다.

실시간 스케일링(즉시 계산)과의 차이: 즉시 계산은 텐서를 먼저 훑어 amax를 구한 뒤 캐스팅하므로 두 번 읽어야 한다. Delayed Scaling은 이전 반복 통계를 쓰므로 텐서를 한 번만 읽는다. 1000~4000 토큰/스텝 규모에서 처리량 차이가 5~10% 발생한다.


두 FP8 서브포맷의 선택 이유

E4M3 vs E5M2

FP8에는 두 가지 표준 서브포맷이 있다. IEEE 754에 준하는 NaN과 무한대 처리 방식이 다르고, 지수·가수 비트 배분도 다르다.

특성E4M3E5M2
지수 비트45
가수 비트32
최대값44857344
특수값NaN만 (Inf 없음)NaN + Inf
사용 위치가중치, 활성화(순전파)그래디언트(역전파)

순전파에 E4M3를 쓰는 이유: 활성화·가중치값은 특별한 처리 없이 일정 범위 내에 있다. 상대 정밀도(가수 비트)가 더 중요하므로 E4M3가 적합하다.

역전파에 E5M2를 쓰는 이유: 그래디언트는 매우 작은 값이 많아 언더플로우가 발생하기 쉽다. 지수 비트가 많은 E5M2로 동적 범위를 넓혀 그래디언트 소실을 방지한다.


GEMM 누적 정밀도 문제

FP8 행렬 곱에서 가장 미묘한 부분은 부분합 누적(accumulation)이다.

H100/B200의 FP8 Tensor Core는 다음 방식으로 동작한다.

C_partial = A_fp8 @ B_fp8   # 16×16 타일 단위 계산
C_fp32 = sum(C_partial)     # FP32로 누적
C_out = cast_to_fp8(C_fp32) # 출력 포맷으로 변환

누적은 FP32다. 이것이 핵심이다. 타일 내 부분합만 FP8에서 이루어지고, 타일 간 합산은 FP32 정밀도로 이루어진다. 이 덕분에 수치 안정성이 BF16 학습과 큰 차이 없이 유지된다.

주의점: 타일 크기(16×16 = 256 원소)를 초과하는 누적은 모두 FP32다. 매우 큰 hidden size에서도 이 보호가 유지된다.


DeepSeek-V3 FP8 학습 사례

DeepSeek-V3(671B 파라미터, 37B 활성화 MoE)는 FP8 혼합 정밀도 학습을 체계적으로 적용한 최초의 프론티어 오픈 모델이다.

주요 설계 결정:

  1. 순전파 FP8: 모든 Linear 레이어의 활성화·가중치를 E4M3으로 캐스팅. 어텐션 레이어는 BF16 유지(수치 민감성 때문).
  2. 역전파 FP8: 그래디언트를 E5M2으로 캐스팅. 단, 가중치 업데이트(∂L/∂W)는 FP32.
  3. 전문가 병렬 통신 FP8: MoE 전문가 간 all-to-all 통신에서도 FP8 압축 적용. 네트워크 대역폭 절감.
  4. 마스터 가중치 BF16: 옵티마이저 상태와 마스터 가중치는 BF16으로 유지.

결과: 2048 H800 GPU에서 14.8T 토큰 학습 시 BF16 대비 약 40% 연산 비용 절감.


PyTorch 2.13으로 FP8 학습 도입하기

PyTorch 2.13은 torch.amp에 FP8 지원을 추가했다. Transformer Engine와 달리 PyTorch 네이티브 경로를 쓴다.

import torch
from transformer_engine.pytorch import fp8_autocast, DelayedScaling

# Transformer Engine 방식 (권장)
with fp8_autocast(enabled=True, fp8_recipe=DelayedScaling()):
    output = model(input)
    loss = criterion(output, target)
loss.backward()

# PyTorch 2.13 네이티브 방식 (실험적)
with torch.autocast(device_type="cuda", dtype=torch.float8_e4m3fn):
    output = model(input)

Transformer Engine 방식을 권장하는 이유: DelayedScaling 레시피가 amax 이력 관리, per-tensor 스케일, 그래디언트 스케일을 자동 처리한다. PyTorch 네이티브는 스케일 관리를 수동으로 해야 한다.


성능 수치와 운영 기준

처리량 향상

모델 크기GPUBF16 처리량FP8 처리량향상률
7BH100 SXM5~3,500 tok/s~5,200 tok/s1.49×
70B8×H100~450 tok/s~680 tok/s1.51×
671B MoE2048×H800(기준)+40%(DeepSeek-V3)
1T+B200 NVL72(기준)~1.8×(이론치, 아키텍처 의존)

Open question: B200에서 tcgen05 FP8 경로를 쓸 때 실측 처리량은 모델·시퀀스 길이에 크게 의존한다. 공개 벤치마크가 아직 제한적이다.

품질 영향

PPL(perplexity) 기준으로 BF16 학습 대비 FP8는 일반적으로 0.1~0.5% 이내의 차이를 보인다. 단, 다음 상황에서 품질 저하가 발생할 수 있다.

도입 체크리스트


References