Switch Transformer를 구현 관점에서 읽기: Top-1 라우팅과 Expert 병렬 처리
요약
Switch Transformer는 MoE 구조에서 라우팅을 Top-1 Expert로 제한하여 효율성을 높인 트랜스포머입니다. 이 글은 Switch layer가 텐서, 그래디언트, 통신 관점에서 어떻게 작동하는지 원 논문을 기반으로 심층 분석합니다. 특히 FFN 레이어를 대체하며 조건부 계산을 도입하고, Top-1 선택 과정에서 발생하는 학습 및 구현상의 기술적 세부 사항들을 다룹니다.
핵심 포인트
- Switch Transformer는 MoE 구조를 활용하여 효율성을 높인 트랜스포머입니다.
- Self-Attention 대신 FFN에 Switch layer를 적용하여 조건부 계산을 도입합니다.
- Top-1 라우팅은 이산적 선택이지만, Gradient 흐름과 보조 손실(auxiliary loss)로 학습이 가능합니다.
- 구현 시 토큰 순서 변경 및 재배열 과정에서 복잡한 텐서 처리가 필요합니다.
Switch Transformer는 MoE(Mixture of Experts: 여러 Expert 중 일부만 실행하는 구조)의 라우팅을 1 토큰당 1개의 Expert로 제한한 트랜스포머입니다.
Top-1이라고 들으면 '가장 높은 점수의 Expert를 하나만 선택'하는 것으로 보일 수 있습니다. 하지만 실제 구현에서는 다음 처리들이 통합되어야 비로소 Switch layer로서 작동합니다.
- 모든 토큰 및 모든 Expert의 Router 확률 계산
- 토큰별 Top-1 Expert 선택
- Expert Capacity 슬롯에 토큰 배치
- Expert별로 재배열하여 디스패치(dispatch) 수행
- Expert FFN을 배치 실행
- 원래 토큰 순서로 복원하고 게이트 값(gate value)을 곱해 잔차 경로에 합류
이 글에서는 Switch Transformers 원 논문을 기반으로 Top-1 라우팅을 텐서(tensor), 그래디언트(gradient), 통신(communication) 순서로 추적합니다.
Switch가 희소하게 만드는 부분은 주로 FFN
원 논문의 Switch Transformer는 T5를 기반으로 한 인코더-디코더 모델입니다. 트랜스포머 블록의 Self-Attention을 Expert화하는 것이 아니라, 일부 Dense FFN을 Router와 여러 Expert FFN으로 구성된 Switch layer로 대체합니다.

개념적으로는 다음 위치에 있습니다.
hidden states
-> Self-Attention
-> residual / normalization
...
Attention은 토큰 간의 정보를 섞습니다. 반면, FFN은 각 토큰 위치를 개별적으로 변환합니다. Switch layer는 후자에 조건부 계산을 도입합니다.
모든 FFN을 Switch layer로 만들 필요는 없습니다. 원 논문의 대표적인 실험에서도 Dense FFN과 Switch FFN을 교차 배치하는 구성이 사용되었습니다. 모델 이름뿐만 아니라, 몇 개 레이어마다 Expert layer가 들어가는지 확인할 필요가 있습니다.
Router의 tensor shape과 Top-1 수식
배치와 시퀀스를 합친 토큰 수를
Router weight를
하나의 토큰 표현으로
Switch는 최대 확률의 Expert ID를 선택합니다.
출력은 선택된 Expert의 출력에 Router 확률을 곱한 값입니다.
여기서 중요한 것은, Top-1로 하더라도
argmax가 있어도 Router를 학습할 수 있는 이유
argmax로 얻는 Expert ID는 이산값(discrete value)입니다. 선택 경계를 넘지 않는 범위에서는 ID 자체를 일반 미분으로 움직일 수 없습니다.
반면, 출력에는 연속값의

loss를
이 gradient는 softmax로부터
하지만 main task의 loss만으로는 일부 Expert에 할당이 집중될 수 있습니다. Switch Transformer에서는 후술할 보조 부하 분산 loss(auxiliary load balancing loss)를 병용합니다.
forward pass에서는 토큰 순서를 두 번 변경한다
Top-1 index를 얻었더라도, 원래 토큰 순서 그대로 1개씩 Expert를 호출하면 작은 행렬 곱셈이 대량으로 발생합니다. 구현에서는 같은 Expert를 선택한 토큰들을 모읍니다.
원래 순서: t1 t2 t3 t4 t5 t6
Expert ID: E2 E1 E2 E4 E1 E3
dispatch 후:
...
이 permute와 inverse permute 외에도, 각 토큰의 게이트 값과 capacity 슬롯을 유지해야 합니다. Expert를 다른 디바이스에 배치하는 경우, dispatch와 return은 All-to-All(각 디바이스가 서로 다른 데이터를 상호 교환하는 집합 통신)이 됩니다.
Expert Capacity를 구체적으로 계산한다
Switch의 Top-1에서는 1개의 Expert가 수용할 토큰 수의 기준으로 다음을 사용합니다.
[7, 4, 3, 2]
이었다고 가정해 봅시다.

| 용량 계수(Capacity Factor) | capacity / Expert | 전체 슬롯 (全slot) | 오버플로우 (overflow) | 여유 슬롯 (空きslot) |
|---|---|---|---|---|
| 1.0 | 4 | 16 | 3 | 3 |
| ... | ||||
| capacity를 늘리면 overflow는 줄어들지만, padding에 해당하는 여유 슬롯(empty slot), 메모리, 통신량은 증가합니다. 작게 설정하면 tensor shape을 제어하기 쉽지만, Expert FFN을 실행할 수 없는 token이 늘어납니다. |
원 논문의 Switch layer에서는 overflow token은 Expert 계산을 건너뛰고 잔차 경로(residual path)를 통해 흐릅니다. sequence에서 token 자체가 사라지는 것은 아니지만, 해당 레이어에서 기대했던 Expert 변환은 받지 못합니다.
보조 손실(Auxiliary loss)은 하드 할당과 소프트 확률을 결합한다
Expert
Switch Transformer의 보조 손실은 다음과 같은 형태입니다.
보조 손실과 capacity는 역할이 다릅니다.
| 메커니즘 (mechanism) | 단계 (stage) | 목적 (purpose) | 트레이드오프 (trade-off) |
|---|---|---|---|
| 보조 부하 분산 손실 (Auxiliary load balancing loss) | training objective | routing의 집중을 억제한다 | 계수가 강하면 task에 유용한 편향도 억제할 수 있다 |
| Expert Capacity | forward execution | 버퍼와 계산량에 상한을 설정한다 | 크면 여유 슬롯이, 작으면 overflow가 늘어난다 |
Router만 float32로 올리는 Selective Precision
Router의 softmax는 작은 수치 차이로 선택된 Expert가 달라지는 부분입니다. 원 논문은 모델 전체가 아닌 Router의 국소 계산만 float32로 하는 Selective Precision(불안정한 부분만 고정밀도로 계산하는 방법)을 제안했습니다.

bfloat16 hidden states
-> cast to float32
-> Router projection / softmax / Top-1
...
고정밀도를 Router 내부로 제한함으로써, float32 tensor를 장치 간 통신에 가져오지 않고 안정성을 개선하는 설계입니다.
원 논문은 작은 weight 초기화 스케일도 검증했습니다. Top-1만 추출하는 것이 아니라, precision, 초기화, 보조 손실을 학습 안정화의 조합으로 이해해야 합니다.
Expert Parallelism에서는 통신을 active compute와 분리한다
Expert Parallelism(Expert 단위의 병렬화)에서는 Expert마다 다른 장치에 배치합니다. 각 장치가 가진 token의 선택지가 다른 장치라면, hidden state를 교환하게 됩니다.

| 병렬화 (Parallelization) | 분할 대상 (split target) | 주요 통신 (main communication) |
|---|---|---|
| Data Parallelism | batch | gradient의 All-Reduce |
| ... | ||
| Switch의 Top-1은 1 token을 여러 Expert에 복제하지 않기 때문에, Top-2보다 dispatch가 단순화됩니다. 그럼에도 불구하고 All-to-All 통신은 남아 있습니다. |
Expert FFN의 이론 FLOPs가 작더라도, 다음 조건에서는 wall-clock time이 늘어납니다.
- 일부 Expert에 token이 집중되어 straggler(처리 완료가 느린 worker)가 발생하는 경우
- 1 Expert당 token 수가 적고 행렬 곱셈이 너무 작은 경우
- network topology에 비해 먼 장치 간 통신이 많은 경우
- communication overlap이 효과를 발휘하지 못하고, Expert 계산 전후에 대기하는 시간이 긴 경우
'active parameter가 적다'는 것이 '실측 지연 시간(latency)도 같은 비율로 짧다'는 것을 의미하지 않습니다.
production에서 분리하여 관찰할 지표
Switch layer를 운영할 때는, quality, routing, capacity, kernel, network을 개별적으로 기록하면 원인을 구분하기 쉽습니다.
| 지표 | 알 수 있는 것 | 단독으로는 알 수 없는 것 |
|---|---|---|
| main loss / auxiliary loss | 태스크 학습과 균등화의 힘 관계 | 디바이스 간 실행 시간 차 |
| ... | ||
| 평균값만으로는 hot Expert나 느린 device를 가릴 수 있습니다. 레이어별, 디바이스별, 분위수(percentile), 최대값도 저장하여 step time의 악화와 대조해야 합니다. |
Hugging Face의 SwitchTransformers 문서에서도 num_experts,
expert_capacity,
router_dtype,
Router logits, auxiliary loss는 각각 다른 설정 및 출력입니다. 하나의 'MoE 설정'으로 묶지 않고 관찰하는 것이 안전합니다.
원 논문의 속도 값을 일반화하지 말 것
arXiv 기록에서는 동일한 계산 자원을 사용하는 T5-Base / T5-Large 계열과의 비교에서 최대 약 7배의 pretraining speedup이 보고되었습니다. 가장 큰 규모의 Switch-C는 약 1.571조 파라미터, 2048 Expert입니다.
다만, 원 논문은 T5계 encoder-decoder와 C4의 masked span objective를 사용한 pretraining 실험입니다. decoder-only chat model의 serving 전반에서 7배 빨라진다는 결과는 아닙니다.
적어도 다음을 나누어 비교해야 합니다.
- 모델이 보유하는 총 파라미터
- 1 토큰으로 실행되는 active Expert
- FLOPs / sequence
- 목표 품질까지의 step 수와 실시간
- 토큰 처리량(token throughput)
- 메모리와 All-to-All 시간
원 논문 자체도 algorithm과 low-level 구현 양쪽 모두가 최종 속도에 영향을 준다고 언급하고 있습니다.
요약
Switch Transformer의 Top-1 Routing은 단순히 argmax를 추가하는 것만으로는 작동하지 않습니다.
- Router 확률에서 1 Expert를 선택하는 T×N - 연속값의 gate를 출력에 곱하고, Router로 gradient를 되돌립니다.
- capacity slot에 토큰을 채우고, overflow는 잔차 경로(residual path)로 흘려보냅니다.
- 보조 손실(auxiliary loss)로 hard 할당과 soft 확률의 집중을 억제합니다.
- Router 내부만 float32로 올리고, 통신 전에 bfloat16으로 되돌립니다.
- Expert Parallelism에서는 All-to-All을 포함한 실측 시간을 관찰합니다.
Top-1이 줄이는 것은 선택 후의 Expert 실행 경로입니다. 총 가중치(weight) 메모리, Router 점수(score), capacity 여유분, 디바이스 간 통신은 남아 있습니다. 이 경계를 분리하면 Switch 계열 모델의 규모와 효율을 같은 숫자로 말하지 않아도 됩니다.
원 논문의 실험 값, T5 전체 아키텍처, 16 토큰을 사용한 capacity 계산을 더 자세히 확인하고 싶다면 개인 블로그의 Switch Transformer 완전판도 참고해 주세요.
참고 자료
Discussion

AI 자동 생성 콘텐츠
본 콘텐츠는 Zenn ML의 원문을 AI가 자동으로 요약·번역·분석한 것입니다. 원 저작권은 원저작자에게 있으며, 정확한 내용은 반드시 원문을 확인해 주세요.
원문 바로가기