torch.switch로 다중 분기를 한 번에 올려 MoE·인덱스 디스패치를 중첩 cond 없이 추적합니다 | DAKER 커뮤니티
한 줄 요약
PyTorch 2.14의 torch.switch는 인덱스로 고르는 다중 분기를 한 번의 higher-order op로 올립니다. 예전에는 torch.cond를 겹쳐 n-way를 만들었고, 그래프가 커지고 의도가 흐려졌습니다. MoE처럼 전문가 인덱스로 갈라지는 경로를 짧게 추적할 때 맞습니다.
작은 MoE 또는 인덱스 디스패치 모듈을 하나 고른 뒤, 중첩 cond 버전과 switch 버전을 같은 입력으로 컴파일해 보십시오. 그래프 노드 수·재컴파일 횟수·스텝 시간만 비교하면 오늘 확인할 수 있습니다. 노트에는 PyTorch 버전, 분기 수, Dynamo 로그 한 줄을 남깁니다.
재현해 볼 실험 조건
- 분기 3~8개짜리 작은 함수를
torch.cond중첩으로 작성해torch.compile합니다. - 같은 로직을
torch.switch로 바꿔 다시 컴파일합니다. - 그래프 크기·가드·스텝 시간을 비교하고, 공유 인자가 분기마다 다시 lift되지 않는지 Dynamo 로그만 확인합니다.
이전에 다룬 @dynamic_spec·복소 텐서 컴파일과는 다른 축입니다. 오늘은 다중 분기 표현만 봅니다.
실무에서 바로 볼 포인트
- MoE·라우터처럼 인덱스로 갈라지는 모델을 중첩 cond 없이 올립니다.
- API가 아직 불안정할 수 있으니, 재현 노트에 정확한 빌드 번호를 고정합니다.
- 분기 본문이 크게 다르면 여전히 재컴파일이 날 수 있으니, 공유 텐서만 먼저 맞춥니다.
관련 DAKER 학습
설명용 생성 이미지입니다. 동작은 로컬에서 중첩 cond와 switch를 같은 입력으로 비교해 확인하십시오.