torch.autograd.graph.node_creation_hook으로 역전파 노드가 생길 때 메타데이터와 훅을 바로 붙입니다 | DAKER 커뮤니티

한 줄 요약
PyTorch 2.14의 torch.autograd.graph.node_creation_hook은 역전파 그래프의 각 노드가 만들어질 때 바로 실행됩니다. 나중에 그래프를 다시 훑지 않고, 그 시점에 메타데이터를 붙이거나 훅을 등록할 수 있습니다. 동기는 역전파 메모리를 그 메모리를 만든 앞 구간에 연결하는 일입니다.
예전에는 역전파가 끝난 뒤에야 “이 메모리가 어느 앞 연산에서 왔는지”를 추측해야 했습니다. 오늘은 작은 모델에서 훅을 등록한 뒤, 노드가 생길 때마다 연산 이름·출력 shape를 리스트에 쌓아 보십시오. 참가자 노트에는 훅이 호출된 횟수, 남긴 필드, 훅을 끈 뒤 리스트가 비는지를 적습니다. “메모리가 줄었다”보다 “노드가 생길 때 기록이 붙었는지”가 재현에 도움이 됩니다.
재현해 볼 실험 조건
- 작은 선형층 한 번과
loss.backward()만 있는 경로에서torch.autograd.graph.node_creation_hook을 등록합니다. - 훅 안에서 노드 타입(또는 이름)과 출력 텐서 정보를 리스트에 남기고, backward가 끝난 뒤 리스트 길이가 0이 아닌지 확인합니다.
- 같은 코드를 훅 없이 한 번 더 돌려, 기록이 훅 등록에만 의존하는지 비교합니다. 필요하면 커스텀
autograd.Function한 개에도 같은 훅이 도는지 봅니다.
같은 릴리스의 ctx.set_output_grad_dtype나 cdist/pdist 이중 역전파는 다른 주제입니다. 이번 글에서는 노드가 생길 때 훅이 도는지에만 초점을 둡니다. 실험이 끝나면 사용한 PyTorch 버전·모델 한 줄·훅 등록 코드를 고정해 두면, 다음 참가자가 같은 조건을 바로 따라올 수 있습니다.
실무에서 바로 볼 포인트
- 훅은 그래프가 만들어진 뒤가 아니라, 노드가 생기는 순간에 붙습니다. 등록 시점을 forward 전으로 두십시오.
- 학습 노트에 훅 호출 횟수와 남긴 필드 이름을 적습니다.
- 메모리 귀속(어느 앞 구간이 역전파 메모리를 만들었는지)을 보려면, 훅에서 그 구간 식별자를 같이 저장하십시오.
관련 DAKER 학습
설명용 생성 이미지입니다. 동작은 로컬 autograd 훅 로그로 확인하십시오.