ctx.set_output_grad_dtype으로 출력 저장 dtype과 다른 기울기 dtype을 선언하면 혼합정밀 Function이 깨지지 않습니다 | DAKER 커뮤니티

출력 기울기 dtype
설명용 생성 이미지입니다.

한 줄 요약
PyTorch 2.14의 ctx.set_output_grad_dtype은 사용자 정의 autograd.Function에서 출력 텐서의 저장 dtype과, 뒤로 들어오는 기울기 dtype을 따로 선언합니다. 혼합정밀에서 두 값이 어긋나 Function이 깨지던 경우를 줄입니다.

출력은 낮은 정밀로 두고, 기울기는 더 넓은 정밀로 받아야 할 때가 있습니다. 예전에는 출력 dtype과 기울기 dtype이 같다고 가정한 뒤, 맞지 않으면 오류나 묵시적 변환으로 디버깅이 어려웠습니다. 오늘은 작은 사용자 정의 Function 하나에서 출력은 반정밀, 기울기는 단정밀로 선언한 뒤, backward에 실제로 어떤 dtype이 들어오는지부터 확인하십시오. 참가자 노트에는 출력 dtype, 선언한 기울기 dtype, 실제 수신 dtype을 적습니다. “빨라졌다”보다 “선언한 dtype이 그대로 오는지”가 재현에 도움이 됩니다.

재현해 볼 실험 조건

  1. 간단한 autograd.Function을 만들고, forward에서 출력을 반정밀로 저장합니다.
  2. ctx.set_output_grad_dtype으로 기울기 dtype을 단정밀로 선언한 뒤, 스칼라 손실로 backward를 호출합니다.
  3. 같은 Function에서 선언을 뺀 경우와 비교해, 오류·묵시 변환·실제 기울기 dtype 차이를 한 줄로 남깁니다.

node_creation_hook은 노드가 생길 때 메타데이터를 붙이는 API이고, 이번 글의 set_output_grad_dtype은 기울기 dtype 계약입니다. 둘은 다른 확장점입니다. 실험이 끝나면 사용한 PyTorch 버전·출력 dtype·선언 dtype을 한 줄로 고정해 두면, 다음 참가자가 같은 조건을 바로 따라올 수 있습니다.

실무에서 바로 볼 포인트

관련 DAKER 학습

설명용 생성 이미지입니다. 동작은 로컬에서 ctx.set_output_grad_dtype 결과로 확인하십시오.