본문으로 건너뛰기
PyTorch조회 1

AMD GPU의 FP8 학습 최적화

TorchAO와 TorchTitan이 AMD GPU의 FP8 학습 속도와 MoE 양자화 비용을 함께 개선했습니다.

이 요약은 AI가 원문을 분석해 생성했습니다. 정확한 내용은 원문 기준으로 확인하세요.

TL;DR

AMD와 Meta/PyTorch 엔지니어들은 Primus-Turbo의 AMD 최적화를 TorchAO와 TorchTitan에 upstream해 AMD Instinct GPU에서 FP8 학습을 표준 PyTorch stack으로 실행할 수 있게 했습니다. Llama3-8B dense 모델에서는 FP8 matrix core를 활용해 BF16보다 13.4% 높은 처리량을 얻었고, peak memory는 약 39GB로 거의 변하지 않았습니다. DeepSeek-V3 671B 같은 MoE 모델에서는 e4m3fnuz 형식 자동 감지, grouped GEMM, Triton kernel fusion으로 양자화 overhead의 89%를 회복했으며 처리량을 5,996 tok/s에서 7,027 tok/s로 높였습니다. transpose copy 제거와 비연속 메모리 기록 개선도 각각 backward 4.2배, colwise scales 6.2배의 개선을 만들었지만, autotune 후보 확대는 성능 향상 없이 compile time만 늘려 되돌려졌습니다.

섹션별 상세

AMD Instinct GPU에서 TorchTitan의 FP8 학습을 기본 지원하려면 PyTorch의 수치 형식, MoE 연산, kernel 실행 경로를 함께 맞춰야 했습니다. Primus-Turbo에서 개발한 AMD 최적화를 TorchAO와 TorchTitan에 upstream하면서 별도 AMD 전용 설치 없이 표준 PyTorch 학습 스택에서 기능을 사용할 기반을 마련했습니다. dense Llama3-8B에서는 8×MI300X와 batch size 1, sequence length 8192 조건에서 rowwise FP8이 BF16보다 13.4% 높은 처리량을 기록했으며 peak memory는 약 39GB로 거의 같았습니다.
8×MI300X에서 Llama3-8B 학습 처리량을 BF16과 세 가지 FP8 설정으로 비교한 막대그래프입니다.
ChartBF16 baseline은 5,377 tokens/second이고 Rowwise FP8은 5,891, Rowwise FP8 + GW는 6,100, Tensorwise FP8은 6,166을 기록합니다. 그래프의 주석처럼 peak memory는 모든 설정에서 약 39GB로 비슷하므로 성능 차이는 메모리 절약보다 FP8 matrix core의 빠른 연산에서 발생합니다.
근거
  • 8×MI300X에서 Llama3-8B rowwise FP8 학습은 BF16보다 13.4% 높은 처리량을 기록했다. Figure 1의 막대그래프와 TorchAO PR #2736 조건 표기: BF16 5,377, Rowwise FP8 5,891, Rowwise FP8 + GW 6,100 tok/s 출처
AMD Instinct의 FP8 형식인 e4m3fnuz는 최대값이 240이고 NaN·Inf 인코딩이 없어, 다른 형식의 최대값을 사용하면 오류 없이 activation clipping과 gradient 손상이 발생할 수 있습니다. TorchAO는 하드웨어를 자동 감지해 올바른 FP8 dtype과 최대값을 선택하도록 바뀌었고, TorchTitan에는 MI300X의 peak FLOPS와 FNUZ용 loss baseline도 반영됐습니다. 따라서 형식 선택은 성능 조정이 아니라 모델 결과의 정확성을 좌우하는 필수 조건입니다.
ROCm에서 TorchTitan, TorchAO, PyTorch Core, AMD Instinct Hardware가 연결되는 FP8 학습 software stack입니다.
Diagram상위 TorchTitan은 FSDP2와 여러 parallelism 방식을 사용하고, TorchAO는 Float8Linear와 Triton FP8 Kernels, scaled_grouped_mm을 제공합니다. 아래 PyTorch Core와 ROCm·AMD Instinct Hardware가 HIP API, RCCL, Triton을 통해 이어지며 AMD upstream PR의 적용 계층을 구분합니다.
근거
  • AMD e4m3fnuz 형식의 최대값은 240이며 NaN과 Inf 인코딩이 없다. AMD FP8 format in TorchAO 절의 e4m3fnuz 속성 표
MoE 모델은 토큰을 일부 expert로 라우팅하므로 dense GEMM처럼 고정된 행렬 모양을 사용할 수 없고, expert별 weight column scale과 activation row scale, 라우팅 offset을 함께 처리해야 합니다. ROCm에서는 Composable Kernel backend를 이용해 토큰을 offset으로 expert에 배치하고, fused Triton kernel로 양자화한 뒤 grouped GEMM을 단일 launch 경로로 연결했습니다. 이 구조가 DeepSeek-V3와 Llama 4 같은 MoE 모델에서 FP8을 적용할 수 있는 실행 기반이 됐습니다.
MoE 토큰 라우팅과 FP8 grouped GEMM 실행 흐름을 dense GEMM과 비교한 다이어그램입니다.
Diagram입력 token은 offset에 따라 E0, E1, E2 expert로 나뉘고 Triton Scale-and-Cast가 activation row와 weight expert-column을 각각 양자화합니다. 이후 _scaled_grouped_mm이 expert별 GEMM을 단일 kernel launch로 처리하며, dense GEMM보다 scale과 routing offset 관리가 복잡한 구조를 나타냅니다.
초기 FP8 양자화 pipeline은 absmax 계산, scale 적용, clamp와 cast를 여러 kernel launch로 나누고 각 단계의 중간 tensor를 HBM에 저장했습니다. MoE layer에 expert weight tensor가 많아지자 8비트 산술보다 kernel launch와 HBM 왕복이 더 큰 비용이 됐으며, 최적화는 launch 수 감소, 메모리 접근 개선, 불필요한 synchronization 제거의 세 단계로 진행됐습니다. backward 경로에서는 .t().contiguous().t()에 따른 전체 tensor 복사를 없애고 scale-and-cast 단계를 fusion해 DeepSeek-MoE-16B에서 backward throughput을 4.2배 높였습니다.
Backward pass에서 FP8 양자화 kernel과 메모리 복사 흐름이 최적화 전후로 어떻게 달라지는지 나타낸 trace입니다.
Screenshottrace에는 FSDP::forward_prefetch와 여러 Triton·통신 구간이 시간축에 배치되어 있으며, 최적화된 경로는 quantization 단계의 반복 작업을 줄인 형태로 나타납니다. 본문에서 설명한 transpose copy 제거와 scale-and-cast fusion이 HBM 왕복 및 kernel launch 비용을 줄이는 근거로 연결됩니다.
근거
  • DeepSeek-MoE-16B의 backward fusion은 8×MI300X에서 backward pass throughput을 4.2배 높였다. Workload Optimization Result 표와 backward pass 최적화 절의 PR #3972, #4069
forward 경로에서는 expert weight 양자화에 필요했던 5개 generic kernel을 하나의 fused Triton kernel로 통합하고 expert와 output-dimension block을 함께 병렬화했습니다. 8×MI325X에서 DeepSeek-V3 671B의 forward 양자화 시간이 약 19ms에서 7ms로 줄었고, 전체 처리량은 5,996 tok/s에서 7,027 tok/s로 17% 증가했습니다. 이는 BF16 기준 7,156 tok/s와의 FP8 격차 중 89%를 회복한 결과입니다.
FP8 forward 경로 최적화 전의 Perfetto trace로, 여러 양자화 kernel이 반복 실행되는 모습을 보여줍니다.
ScreenshotFSDP::forward_prefetch 구간 안에 여러 개의 분리된 quantization 작업과 통신 작업이 이어져 있어 forward 양자화 chain의 반복 launch 구조가 드러납니다. 본문은 이 경로에서 5개 generic kernel이 expert별로 반복되어 약 90ms/step의 overhead를 만들었다고 수치화합니다.
FP8 forward 경로 최적화 후의 Perfetto trace로, 여러 작업이 fused Triton kernel으로 통합된 모습을 보여줍니다.
Screenshot최적화 후 trace에는 triton_fp8_colwise_3d_scale_and_cast로 대체된 단일 fused 구간이 나타나며, forward 양자화 시간이 약 19ms에서 7ms로 줄어든 결과와 대응합니다. kernel launch를 다섯 번에서 한 번으로 줄이면 주변 GEMM이 더 일찍 실행될 수 있어 전체 처리량 개선으로 이어집니다.
DeepSeek-V3 671B의 Attention, MoE, GEMM, Activation, Norm, Communication, Others별 GPU 시간을 세 설정으로 비교한 차트입니다.
ChartFP8 upstream 설정 V2에서는 Others가 127.1ms로 가장 큰 비중을 차지하고, fused FP8 설정 V4에서는 Others가 276.6ms로 표시된 비교 구조 속에서 전체 step time이 BF16 664ms, FP8 Upstream 791ms, FP8+Fused 693ms로 제시됩니다. 본문은 generic quantization kernel에 집중된 FP8 overhead의 92%가 Others에 있었고 fusion으로 BF16과 FP8 사이 격차의 89%를 회복했다고 해석합니다.
근거
  • DeepSeek-V3 671B forward fusion은 8×MI325X에서 처리량을 5,996 tok/s에서 7,027 tok/s로 높였다. Forward pass 절 및 Figure 5의 Perfetto trace 비교 출처
backward의 colwise scales kernel은 SIMD lane이 K 간격으로 떨어진 주소에 기록해 비연속 메모리 transaction을 일으키고 있었습니다. 출력 tile을 LDS에서 transpose한 뒤 저장하고 중복 HBM read를 없앤 single-pass variant를 추가하면서 MI300X의 DeepSeek-V3 671B MoE layer 시간이 7,290μs에서 1,170μs로 줄어 6.2배 빨라졌습니다. AMD GPU의 atomic 연산에는 불필요한 acquire-release fence가 들어가므로 torch.version.hip 조건에서 relaxed ordering을 적용해 synchronization 비용도 낮췄습니다.
근거
  • colwise scales 최적화는 MI300X의 DeepSeek-V3 671B MoE layer 시간을 7,290μs에서 1,170μs로 줄였다. Level 2 절의 메모리 접근 최적화 결과와 PR #4113
Triton autotune 후보를 1개에서 8~16개로 늘리는 접근은 Llama 4의 MI300X shape에서 측정 가능한 개선을 만들지 못했고 첫 iteration compile time만 늘렸습니다. 후보 구성을 무작정 확대하는 대신 wavefront size, LDS capacity, register pressure 같은 하드웨어 제약에 맞춰 search space를 설계해야 한다는 결과가 나왔습니다. 현재 upstream pipeline은 MI355X에서 MXFP8 grouped GEMM과 forward·backward 양자화 kernel로 확장되고 있습니다.
근거
  • Triton autotune 후보를 8~16개로 확대한 변경은 측정 가능한 성능 개선 없이 첫 iteration compile time을 늘려 되돌려졌다. What didn’t work 절의 Llama 4, MI300X benchmark와 PR #3952, #4024

용어 해설

FP8 양자화(FP8 Quantization)
FP8 양자화는 16비트 데이터를 8비트 부동소수점 형식으로 변환해 행렬 연산의 데이터 이동량과 계산 비용을 줄이는 기법입니다. 변환 전 스케일 계산과 clipping이 필요하며, AMD GPU에서는 e4m3fnuz 형식의 표현 범위를 정확히 맞춰야 모델 품질 저하를 피할 수 있습니다.
Grouped GEMM
Grouped GEMM은 Mixture-of-Experts 모델에서 토큰마다 선택된 expert의 가변 크기 배치를 한 번에 행렬 곱으로 처리하는 방식입니다. 일반 GEMM과 달리 expert별 weight scale, activation의 행별 scale, 토큰 라우팅 offset을 함께 관리해야 하므로 양자화와 kernel dispatch가 더 복잡합니다.
MXFP8
MXFP8은 데이터와 함께 group 단위 scale을 저장하는 FP8 scaling 전략입니다. tensorwise, rowwise, blockwise보다 세밀한 단위로 값을 조정할 수 있어 MoE와 같은 불규칙한 연산 형태에 적용되지만, 하드웨어와 kernel이 해당 scale 구조를 함께 처리해야 합니다.
Triton kernel fusion(Triton Kernel Fusion)
Triton kernel fusion은 absmax 계산, scale 산출, clipping, FP8 변환처럼 পৃথ개별 kernel로 실행되던 단계를 하나의 kernel에 결합하는 방식입니다. 중간 tensor를 HBM에 반복해서 기록하고 읽는 과정을 줄여 산술 연산보다 메모리 이동과 kernel launch가 병목인 MoE 양자화에서 효과를 냅니다.
FNUZ
FNUZ는 AMD Instinct GPU가 사용하는 FP8 형식으로, Finite, No NaN, Unsigned Zero를 뜻합니다. e4m3fnuz의 최대값은 240이며 NaN과 Inf 표현이 없어 범위를 넘은 값이 오류 대신 clipping으로 이어질 수 있으므로, TorchAO가 하드웨어에 맞는 dtype과 최대값을 선택해야 합니다.

기술

  • Primus-Turbo
  • TorchTitan
  • TorchAO
  • FP8
  • BF16
  • FNUZ
  • e4m3fnuz
  • e4m3fn
  • ROCm
  • Triton
  • Composable Kernel
  • torch.compile
  • FSDP2
  • LDS
  • HBM
  • MI300X
  • MI325X
  • MI350X
  • MI355X
  • DeepSeek-V3
  • Llama3-8B
  • Llama 4

활용 사례

  • AMD Instinct GPU에서 dense LLM의 FP8 학습
  • DeepSeek-V3와 Llama 4 같은 MoE 모델의 FP8 grouped GEMM
  • TorchAO와 TorchTitan 기반 분산 학습
  • ROCm 환경의 Triton 양자화 kernel 최적화
AI 분석 전체 내용 보기

AI 요약 · 북마크 · 개인 피드 설정 — 무료

출처 · 인용 안내

원문 발행 2026. 08. 14.수집 2026. 08. 14.출처 타입 RSS

인용 시 "요약 출처: AI Trends (aitrends.kr)"를 표기하고, 사실 확인은 원문 보기 기준으로 진행해 주세요. 자세한 기준은 운영 정책을 참고해 주세요.