PyTorch 2.14에서 ctx.set_output_grad_dtype으로 혼합정밀 autograd.Function 깨짐 줄이기 | DAKER 커뮤니티

출력 기울기 dtype

혼합정밀 학습에서는 출력 텐서를 어떤 dtype으로 저장할지와, backward로 들어오는 기울기를 어떤 dtype으로 받을지가 꼭 같지 않을 때가 있습니다. 문제는 이 차이가 사용자 정의 autograd.Function 안에서 분명하게 드러나지 않으면, 오류가 나거나 묵시적 변환이 일어나도 원인을 바로 찾기 어렵다는 점입니다.

PyTorch 2.14의 ctx.set_output_grad_dtype은 바로 이 지점을 다룹니다. 출력 저장 dtype과 기울기 dtype을 따로 선언할 수 있게 하면서, 혼합정밀 환경에서 Function이 깨지던 경우를 줄이는 데 도움이 됩니다.

PyTorch 2.14의 ctx.set_output_grad_dtype은 사용자 정의 autograd.Function에서 출력 텐서의 저장 dtype과, 뒤로 들어오는 기울기 dtype을 따로 선언합니다.

출력은 낮은 정밀로 두고, 기울기는 더 넓은 정밀로 받아야 할 때가 있습니다. 예전에는 출력 dtype과 기울기 dtype이 같다고 가정한 뒤, 맞지 않으면 오류나 묵시적 변환으로 이어져 디버깅이 까다로웠습니다. 이럴 때는 작은 사용자 정의 Function 하나로 먼저 재현해 보고, backward에 실제로 어떤 dtype이 들어오는지 확인해 두는 것이 좋습니다.

왜 이 차이를 먼저 확인해야 하는가

혼합정밀에서 중요한 것은 단순히 반정밀 출력을 쓰는지 여부가 아닙니다. 실제로는 출력 dtype, 선언한 기울기 dtype, 그리고 backward가 받은 실제 dtype이 서로 어떻게 맞물리는지가 핵심입니다. 성능 수치보다 먼저 이 계약이 지켜지는지 확인해 두면, 같은 조건에서 재현하고 비교하기가 훨씬 쉬워집니다.

빨라졌는지보다 선언한 dtype이 그대로 오는지 확인하는 편이 재현에 더 도움이 됩니다.

재현해 볼 실험 조건

실험은 아주 작게 시작하면 됩니다. 간단한 autograd.Function을 만들고, forward에서 출력을 반정밀로 저장합니다. 그다음 ctx.set_output_grad_dtype으로 기울기 dtype을 단정밀로 선언한 뒤, 스칼라 손실로 backward를 호출해 실제로 어떤 dtype이 들어오는지 확인합니다.

이때 같은 Function에서 선언을 뺀 경우도 함께 비교해 두면 좋습니다. 오류가 나는지, 묵시적 변환이 생기는지, 실제 기울기 dtype이 어떻게 달라지는지를 한 줄씩 남겨 두면 이후 비교가 쉬워집니다.

기록해 둘 항목

참가자 노트에는 출력 dtype, 선언한 기울기 dtype, 실제 수신 dtype을 적어 두면 됩니다. 실험이 끝난 뒤에는 사용한 PyTorch 버전, 출력 dtype, 선언 dtype을 한 줄로 정리해 두는 것이 좋습니다. 다음 참가자가 같은 조건을 바로 따라오기 쉬워집니다.

node_creation_hook과는 무엇이 다른가

node_creation_hook은 노드가 생길 때 메타데이터를 붙이는 API입니다. 반면 이번 글의 set_output_grad_dtype은 기울기 dtype 계약을 다루는 기능입니다. 둘은 이름이 비슷하게 보일 수 있어도 역할이 다릅니다.

node_creation_hook은 메타데이터를 붙이는 확장점이고, set_output_grad_dtype은 기울기 dtype 계약입니다.

따라서 이번 실험에서는 두 기능을 섞지 않고, 기울기 dtype 선언이 실제로 어떻게 동작하는지만 분리해서 보는 편이 좋습니다.

실무에서 바로 볼 포인트

혼합정밀 Function을 다룰 때는 출력 저장 dtype과 기울기 dtype을 먼저 적어 두는 것이 좋습니다. 그리고 학습 노트에는 선언이 있을 때와 없을 때의 오류 여부와 실제 dtype을 함께 남기면 됩니다. 이렇게 해 두면 문제가 생겼을 때 원인을 좁히기가 수월합니다.

관련 DAKER 학습

PyTorch 실전 고급 트랙
node_creation_hook 실습 글
DAKER 커뮤니티

참고 자료

https://daker.ai/public/learning/tracks/pytorch-practice-advanced-track
https://daker.ai/community/autograd-node-creation-hook-attach-metadata-hooks
https://daker.ai/community

같은 조건에서 실험해 보셨다면, 선언 유무에 따라 실제로 들어온 기울기 dtype이 어떻게 달랐는지 궁금합니다.