PyTorch 2.14에서 torch.while_loop CUDA 그래프 캡처가 달라진 점과 점검할 부분 | DAKER 커뮤니티

반복 횟수가 입력에 따라 달라지는 루프는 PyTorch compile과 CUDA 그래프를 함께 쓸 때 늘 까다로운 지점이었습니다. 특히 디코드처럼 샘플마다 길이가 다른 작업은 graph break가 쉽게 생겨, 성능을 기대만큼 끌어올리기 어려운 경우가 많았습니다.
PyTorch 2.14에서는 torch.while_loop를 CUDA 그래프로 캡처할 수 있게 되면서 이 문제가 조금 다른 국면으로 들어왔습니다. 루프가 데이터에 따라 달라지더라도, 예전처럼 곧바로 그래프가 끊기는 상황이 줄어들 수 있기 때문입니다.
PyTorch 2.14부터 torch.while_loop를 CUDA 그래프로 캡처할 수 있습니다.
왜 달라졌는가
길이가 계속 바뀌는 루프를 compile에 넣으면, 반복 횟수인 trip count가 입력값에 따라 달라져 그래프가 자주 끊깁니다. trip count는 루프가 몇 번 도는지를 뜻합니다. 특히 디코드처럼 샘플마다 길이가 다른 작업은 한 그래프 안에 유지하기가 어려웠습니다.
이럴 때는 먼저 루프 본문이 매 단계 같은 연산인지 확인하는 것이 좋습니다. 본문이 안정적이면 while_loop를 그래프 안에 남길 가능성이 커집니다.
2.14는 CUDA while 조건부 노드를 사용해 while_loop를 캡처 가능한 형태로 올립니다. 그래서 데이터에 따라 trip count가 바뀐다고 해서 항상 그래프가 깨지지는 않습니다. 여기에 고정된 최대 반복 한도를 함께 두면, 이전보다 더 유연하게 처리할 수 있습니다.
직접 확인해 볼 실험 조건
실제로는 같은 입력을 기준으로 캡처 경로를 비교해 보는 것이 가장 분명합니다. 입력 텐서에 따라 trip count가 달라지는 torch.while_loop를 준비하고, CUDA에서 같은 입력을 사용해 고정 max trip 설정과 함께 그래프 캡처 경로로 실행해 보면 됩니다.
이후에는 캡처가 성공했는지 확인하고, 실패했다면 graph break가 어디서 생기는지 프로파일로 보면 됩니다. 프로파일은 실행 구간을 기록해 병목이나 끊김 지점을 찾는 도구입니다. 같은 입력을 두 버전에서 실행해 graph break 위치가 어떻게 달라졌는지 기록해 두면, 배포 전에 원인을 나누어 확인하기가 쉬워집니다.
같은 입력으로 비교해야 graph break 위치와 스텝 시간 차이를 제대로 볼 수 있습니다.
실무에서 먼저 볼 점
가장 먼저 볼 부분은 최대 반복 한도입니다. 최대 반복 한도를 너무 크게 잡으면 캡처 비용과 메모리 사용량이 커집니다. 먼저 워크로드의 실제 상한을 재는 것이 좋습니다.
루프 본문 안의 동작도 중요합니다. 루프 안에 장치 동기나 파이썬 부수효과가 있으면 그래프 밖으로 빠질 수 있습니다. 부수효과는 출력 외에 상태를 바꾸는 동작입니다. 가능하면 본문을 순수한 텐서 연산으로 좁히는 편이 안정적입니다.
배포나 커스텀 루프 학습을 비교할 때는 compile 전후의 스텝 시간을 반드시 같은 입력으로 비교해야 합니다. 캡처가 실패해도 모델이 틀렸다는 뜻은 아닙니다. 먼저 break 위치를 좁혀 보고, 루프 본문을 다시 순수 연산 중심으로 다듬으면 됩니다.
관련 DAKER 학습
참고 자료
https://daker.ai/public/learning/materials/pytorch-advanced-custom-loop-amp-debugging-1
https://daker.ai/public/learning/materials/pytorch-advanced-deployment-torchscript-onnx-perfo
https://daker.ai/public/learning/tracks/pytorch-practice-advanced-track
여러분은 torch.while_loop를 쓸 때 어떤 지점에서 graph break를 가장 자주 만나셨나요?