torch.switch로 다중 분기를 한 번에 올려 MoE·인덱스 디스패치를 중첩 cond 없이 추적합니다 | DAKER 커뮤니티

torch.switch 분기
설명용 생성 이미지입니다.

한 줄 요약
PyTorch 2.14의 torch.switch는 인덱스로 고르는 다중 분기를 한 번의 higher-order op로 올립니다. 예전에는 torch.cond를 겹쳐 n-way를 만들었고, 그래프가 커지고 의도가 흐려졌습니다. MoE처럼 전문가 인덱스로 갈라지는 경로를 짧게 추적할 때 맞습니다.

작은 MoE 또는 인덱스 디스패치 모듈을 하나 고른 뒤, 중첩 cond 버전과 switch 버전을 같은 입력으로 컴파일해 보십시오. 그래프 노드 수·재컴파일 횟수·스텝 시간만 비교하면 오늘 확인할 수 있습니다. 노트에는 PyTorch 버전, 분기 수, Dynamo 로그 한 줄을 남깁니다.

재현해 볼 실험 조건

  1. 분기 3~8개짜리 작은 함수를 torch.cond 중첩으로 작성해 torch.compile합니다.
  2. 같은 로직을 torch.switch로 바꿔 다시 컴파일합니다.
  3. 그래프 크기·가드·스텝 시간을 비교하고, 공유 인자가 분기마다 다시 lift되지 않는지 Dynamo 로그만 확인합니다.

이전에 다룬 @dynamic_spec·복소 텐서 컴파일과는 다른 축입니다. 오늘은 다중 분기 표현만 봅니다.

실무에서 바로 볼 포인트

관련 DAKER 학습

설명용 생성 이미지입니다. 동작은 로컬에서 중첩 cond와 switch를 같은 입력으로 비교해 확인하십시오.