PyTorch 2.14 torch.while_loop가 CUDA 그래프 캡처돼 trip count 데이터 의존에도 graph break가 줄습니다 | DAKER 커뮤니티

루프가 그래프에
설명용 생성 이미지입니다.

한 줄 요약
PyTorch 2.14부터 torch.while_loop는 CUDA while 조건부 노드로 CUDA 그래프 캡처가 가능합니다. 반복 횟수가 데이터에 따라 바뀌어도, 예전처럼 곧바로 graph break로 떨어지지 않는 경로가 열렸습니다.

동적 길이 루프를 compile에 넣다 보면, trip count가 입력에 묶여 그래프가 끊기는 장면을 자주 봅니다. 디코드처럼 길이가 샘플마다 다른 작업을 한 그래프에 남기고 싶을 때 특히 답답합니다. 오늘은 루프 본문이 매 스텝 같은 연산인지부터 확인하십시오. 본문이 안정적이면 while_loop를 그래프 안에 남길 여지가 생깁니다.

재현해 볼 실험 조건

  1. trip count가 입력 텐서에 의존하는 torch.while_loop를 준비합니다.
  2. CUDA에서 동일 입력을 고정 max trip 설정과 함께 그래프 캡처 경로로 돌립니다.
  3. 캡처 성공 여부와, 캡처 실패 시 남는 graph break 위치를 프로파일로 확인합니다.

2.14는 CUDA while 조건부 노드를 통해 while_loop를 캡처 가능한 형태로 올립니다. 데이터 의존 trip count가 무조건 그래프를 깨던 이전 가정은, 고정된 최대 반복 한도와 맞물려 완화됩니다. 같은 입력을 두 버전에서 돌려 break 지점이 어디로 이동했는지 기록해 두면, 배포 전에 원인 분리가 빨라집니다.

실무에서 바로 볼 포인트

관련 DAKER 학습

설명용 생성 이미지입니다. 캡처 성공 여부는 로컬 CUDA 환경에서 다시 확인하십시오.