torch.while_loop를 CUDA 그래프에 담아 가변 반복도 한 번의 캡처로 재생합니다 | DAKER 커뮤니티

while_loop 그래프
설명용 생성 이미지입니다.

한 줄 요약
PyTorch 2.14부터 torch.while_loop를 CUDA 그래프에 담을 수 있습니다. 조건은 부모 스트림에서 평가되고, 본문 끝에서 다시 검사되는 CUDA while 조건부 노드로 이어집니다. 가변 길이 인덱스·패킹된 시퀀스처럼 반복 횟수가 런타임에 정해지는 구간을 캡처 밖으로 밀어내지 않습니다.

가변 trip count가 있는 작은 루프를 하나 고른 뒤, 그래프 캡처 전후 스텝 시간과 동기화 횟수만 비교해 보십시오. 노트에는 PyTorch·CUDA 버전, 최대 trip count, 캡처 성공 여부를 남깁니다.

재현해 볼 실험 조건

  1. 고정 상한 trip count를 둔 torch.while_loop로 가변 길이 축소를 만듭니다.
  2. 같은 루프를 CUDA 그래프 캡처·재생으로 돌립니다.
  3. 캡처 실패 지점·디바이스-호스트 동기화·스텝 시간을 비교하고, while_loop 없이 풀린 버전과 섞지 마십시오.

이전에 다룬 Inductor simple_overlap·compile-on-one-rank와는 다른 축입니다. 오늘은 루프가 캡처를 깨지 않는지만 봅니다.

실무에서 바로 볼 포인트

관련 DAKER 학습

설명용 생성 이미지입니다. 동작은 로컬 CUDA에서 while_loop 캡처 전후를 비교해 확인하십시오.