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를 그래프 안에 남길 여지가 생깁니다.
재현해 볼 실험 조건
- trip count가 입력 텐서에 의존하는
torch.while_loop를 준비합니다. - CUDA에서 동일 입력을 고정 max trip 설정과 함께 그래프 캡처 경로로 돌립니다.
- 캡처 성공 여부와, 캡처 실패 시 남는 graph break 위치를 프로파일로 확인합니다.
2.14는 CUDA while 조건부 노드를 통해 while_loop를 캡처 가능한 형태로 올립니다. 데이터 의존 trip count가 무조건 그래프를 깨던 이전 가정은, 고정된 최대 반복 한도와 맞물려 완화됩니다. 같은 입력을 두 버전에서 돌려 break 지점이 어디로 이동했는지 기록해 두면, 배포 전에 원인 분리가 빨라집니다.
실무에서 바로 볼 포인트
- 최대 반복 한도를 너무 크게 잡으면 캡처·메모리 비용이 커집니다. 워크로드 상한을 먼저 재십시오.
- 루프 안에서 장치 동기나 파이썬 부수효과가 있으면 그래프 밖으로 빠집니다. 본문을 순수 텐서 연산으로 좁히십시오.
- 배포·커스텀 루프 학습과 맞춰, compile 전후 스텝 시간을 같은 입력으로만 비교하십시오.
- 캡처가 실패해도 모델이 틀린 것은 아닙니다. break 위치를 좁힌 뒤 본문을 다시 순수 연산으로 다듬으면 됩니다.
관련 DAKER 학습
설명용 생성 이미지입니다. 캡처 성공 여부는 로컬 CUDA 환경에서 다시 확인하십시오.