PyTorch 2.14 torch.switch로 다중 분기를 더 짧게 추적하는 방법 | DAKER 커뮤니티
한 줄 답: PyTorch 2.14 torch.switch는 인덱스로 고르는 다중 분기를 한 번의 higher-order op로 올립니다. MoE·인덱스 디스패치에서 중첩 torch.cond보다 트레이스 그래프를 짧게 읽기 쉬운 경우가 많습니다. 같은 입력으로 그래프 크기·가드·스텝 시간만 비교하고, 성능 배율은 단정하지 마세요.
DAKER PyTorch 학습 독자를 위한 점검 노트입니다. API가 아직 바뀔 수 있으니 빌드 번호를 노트에 고정합니다.
왜 중첩 cond 대신 switch를 보게 되는가
n-way 분기를 torch.cond로 겹치면 그래프가 커지고 의도가 흐려지기 쉽습니다. torch.switch는 인덱스로 한 경로를 고르는 상황을 더 직접 표현합니다.
torch.switch는 인덱스로 고르는 다중 분기를 한 번의 higher-order op로 올립니다.
재현해 볼 실험 조건
- 분기 3~8개짜리 작은 함수를 중첩
torch.cond로 작성해torch.compile합니다. - 같은 로직을
torch.switch로 바꿔 다시 컴파일합니다. - 그래프 크기·가드·스텝 시간을 비교합니다.
- 공유 인자가 분기마다 다시 lift되는지 Dynamo 로그를 확인합니다.
실무에서 바로 볼 포인트
MoE·라우터처럼 인덱스로 갈라지는 모델에 자연스럽습니다. 분기 본문이 크게 다르면 재컴파일이 날 수 있으므로, 먼저 공유 텐서 맞물림을 보세요.
측정표 템플릿
- PyTorch 버전 / 분기 수
- cond 중첩 vs switch 그래프 크기
- 재컴파일 횟수·스텝 시간
- Dynamo 로그 한 줄
오늘 바로 할 측정
작은 인덱스 디스패치 모듈 하나로 cond vs switch만 비교합니다. @dynamic_spec·복소 텐서 컴파일과 축을 섞지 마세요.
마치며
torch.switch는 “무조건 더 빠름”이 아니라 “다중 분기 표현을 짧게 추적하는 도구”에 가깝습니다. 공식 릴리스 노트·PR을 기준으로 재현하세요.
관련 DAKER 글과 주제를 나누는 법
while_loop CUDA 그래프·TunableOp는 인접 주제지만 측정 축이 다릅니다. switch 노트에는 분기 수·그래프 크기·재컴파일만 남깁니다.
실패를 문서로 남기는 이유
switch로 바꿨더니 가드가 늘었다면 분기 본문 차이와 공유 텐서를 한 줄로 남기세요.
중첩 cond와의 비교 축을 고정하기
“더 빠르다”만 남기면 다음 실험이 흔들립니다. 같은 입력·같은 분기 수에서 그래프 노드 수, 가드 수, 재컴파일 횟수, 스텝 시간 네 칸만 채웁니다. 분기 본문 길이가 크게 다르면 switch로 바꿔도 재컴파일이 날 수 있으니, 본문 차이를 표에 한 줄로 적습니다.
MoE 라우터에 옮길 때
인덱스 디스패치·전문가 선택처럼 “하나의 정수 인덱스로 경로를 고르는” 구조부터 옮깁니다. 서로 다른 제어 흐름이 얽힌 중첩 조건을 한 번에 switch로 바꾸려 하지 마세요. 작은 라우터 모듈에서 그래프가 짧아지는지 확인한 뒤 확장합니다.
출처
- PyTorch 2.14 Release Blog: https://pytorch.org/blog/pytorch-2-14-release/
- PR #182902 (torch.switch): https://github.com/pytorch/pytorch/pull/182902
- torch.cond docs: https://docs.pytorch.org/docs/stable/generated/torch.cond.html
- v2.14.0 release tag: https://github.com/pytorch/pytorch/releases/tag/v2.14.0
- PyTorch 실전 고급 트랙: https://daker.ai/public/learning/tracks/pytorch-practice-advanced-track