AI Briefing

NVIDIA, JAX용 TE로 MoE 학습 10배 가속

NVIDIA의 Transformer Engine으로 JAX의 Dropless MoE 학습 처리량 약 10배 달성하기 (105 → 1,025 TFLOPs)

·2026.09.16 15:30

핵심 내용

NVIDIA가 JAX용 Transformer Engine을 통해 Dropless MoE 학습에서 GB200 기준 GPU당 1,025 TFLOPs를 달성하며 처리량을 약 10배 개선했다.

1 / 6

자세히 보기

NVIDIA가 JAX 프레임워크용 **Transformer Engine(TE)**을 통해 Dropless MoE(Mixture of Experts) 학습 성능을 대폭 향상시켰다. 기존 라이브러리는 라우팅 불균형으로 인한 Ragged Tensor 처리와 GPU 간 통신 병목으로 인해 DeepSeek-V3 모델 학습 시 GPU당 약 105 TFLOPs에 그쳤으나, TE 적용 시 GB200 환경에서 GPU당 1,025 TFLOPs를 달성하며 약 10배의 성능 향상을 기록했다.

주요 기술적 개선 사항

  • Grouped GEMM 및 MXFP8 양자화: 전문가별 행렬곱을 단일 커널 호출로 처리하고, Blackwell 아키텍처에서 MXFP8 블록 스케일링을 지원하여 패딩 낭비를 제거했다.
  • EP(Expert Parallelism) 통신 최적화: 디스패치와 컴바인 연산을 NCCL EP 기반의 융합 커널로 통합하여 CPU 크리티컬 패스를 제거하고 대역폭을 절약했다.
  • 메모리 오프로딩: Grace Blackwell의 NVLink-C2C 대역폭을 활용해 활성값을 오프로딩함으로써 OOM(Out of Memory)을 방지하고 재계산 대비 57% 빠른 속도를 구현했다.

배포 및 적용 조건

해당 최적화는 NVIDIA NGC MaxText 컨테이너(2026-09-09 이미지 이후)와 TE JAX v2.19 이상 버전에서 사용 가능하다. DeepSeek-V3 671B 모델 재현 시 EP=8, FSDP=16 등의 설정과 함께 te_moe_block=true, te_gmm_quantization="te_mxfp8" 등의 옵션을 활성화해야 한다. 향후 NVFP4 지원 및 CuTe DSL 기반의 추가 커널 융합을 통해 DeepSeek-V3 사전학습 속도를 최대 8%, GPT-OSS는 최대 93%까지 향상시킬 계획이다.

이 한국어 요약은 AI가 자동으로 만들었습니다. 원문의 주장과 맥락은 원문에서 확인해 주세요. 저작권은 원저작자에게 있습니다.

AI 처리 방식을 확인하거나, 요약 오류와 출처 표기 문제, 삭제 요청을 문의 · 건의로 알려주세요.