NVIDIA, JAX용 Transformer Engine으로 드롭리스 MoE 학습 10배 가속
최근 30일 조회수 — 좋아요 —핵심 요약
NVIDIA가 JAX Transformer Engine 커널 최적화로 DeepSeek-V3 MoE 학습의 GPU당 처리량을 10.4배 높였다고 밝혔다.
MoE(Mixture of Experts) 모델을 대규모로 학습시킬 때 병목이 되는 것은 연산 자체가 아니라 통신과 불균등한 토큰 처리다. NVIDIA는 자사 기술 블로그를 통해 JAX용 Transformer Engine에 드롭리스(dropless) MoE 학습을 가속하는 커널 최적화를 적용해, DeepSeek-V3 학습 벤치마크에서 GPU당 처리량을 10.4배 끌어올렸다고 밝혔다.1
NVIDIA에 따르면 GB200 GPU에서 최적화 이전 DeepSeek-V3 학습 베이스라인은 GPU당 103 TFLOPS에 그쳤고, 누적 커널 시간의 84%를 GPU 간 통신이 차지했다. 즉 GPU 대부분이 실제 연산이 아니라 데이터를 기다리는 데 시간을 쓴 셈이다. JAX와 Transformer Engine의 타깃 커널 최적화를 적용한 뒤에는 이 수치가 GPU당 1,068 TFLOPS로 올라갔다.
MoE는 하나의 거대한 피드포워드 네트워크(FFN)를 모든 토큰이 공유하는 밀집(dense) 모델과 달리, 여러 개의 작은 전문가(expert) 네트워크와 학습된 라우터로 구성된다. 라우터가 각 토큰마다 상위 K개 전문가를 골라 활성화하는 방식이라, 학습이 진행될수록 라우터가 특정 전문가를 선호하게 되면서 전문가별 토큰 분포가 크게 치우친다. 배치마다 부하가 다르고, 같은 배치 안에서도 어떤 전문가는 다른 전문가보다 훨씬 많은 토큰을 받는다. 이 때문에 깔끔한 직사각형 행렬 연산(GEMM)으로 묶어 처리할 수 없는 불규칙한(ragged) 텐서가 발생하는데, 대부분의 라이브러리는 균일한 직사각형 데이터 구조를 전제로 최적화돼 있어 이 구조 자체가 성능 저하의 원인이 된다.
여기에 전문가 병렬화(EP, expert parallelism)를 쓰면 토큰을 여러 GPU에 분산시켰다가(dispatch) 다시 원래 순서로 모아야(combine) 하는데, 이 과정이 최적화돼 있지 않으면 통신이 전체 시간을 지배하고 GPU는 유휴 상태로 남는다. NVIDIA는 이런 all-to-all 통신 구간이 제대로 처리되지 않으면 GPU가 유용한 연산을 하기 전에 데이터를 기다리며 멈춰 선다고 설명한다.
이 문제를 다루는 방식은 크게 두 갈래로 나뉜다. 드롭리스 MoE는 부하가 아무리 불균등해도 모든 토큰을 선택된 전문가가 처리하도록 강제하는 방식으로, 모델 품질 면에서는 유리하지만 시스템에는 부담이 크다. NVIDIA가 인용한 MegaBlocks 연구는 전문가 연산을 블록 희소(block-sparse) 행렬곱으로 재구성해 각 전문가가 토큰을 버리거나 패딩하지 않고 서로 다른 개수를 처리할 수 있도록 했는데, 이를 위해서는 블록 희소 GPU 커널, 최적화된 그룹 GEMM, 가변 토큰 수에 맞춘 디스패치·컴바인 연산이 별도로 필요하다.
반대로 용량 기반(capacity-based) MoE는 각 전문가에게 고정된 토큰 예산을 부여하고 초과분은 잘라내거나 패딩해 맞추는 방식이다. 연산 구조는 규칙적이고 하드웨어 친화적이지만, 초과 토큰을 버리면 모델이 불완전한 데이터로 학습되고 패딩으로 버리지 않으면 연산과 메모리를 낭비하는 직접적인 트레이드오프가 생긴다.
드롭리스 방식을 택하면 학습 스택은 더 이상 고정된 전문가 형태에 의존할 수 없다. NVIDIA는 이 지점이 바로 Transformer Engine의 MoE 최적화가 겨냥한 문제라고 설명하며, 구체적인 커널 구현 방식은 이어지는 내용에서 다룬다고 밝혔다. 제공된 자료에는 최적화 기법의 세부 구현과 벤치마크 재현 조건까지는 담겨 있지 않아, 실제 적용 범위나 다른 하드웨어·프레임워크에서의 재현 여부는 추가 확인이 필요하다.
Footnotes
-
NVIDIA Technical Blog, “Accelerating Dropless MoE Training in JAX with NVIDIA Transformer Engine” ↩
읽기 목록은 이 브라우저에 저장됩니다.
출처
- Accelerating Dropless MoE Training in JAX with NVIDIA Transformer Engine | NVIDIA Technical Blog — NVIDIA Technical Blog
이 글은 위 출처를 근거로 자동 생성된 뒤 발행됐습니다. 원문을 함께 확인해 주세요. 교차 보도 없이 단독 출처로 작성됐습니다.