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

노드 생성 훅
설명용 생성 이미지입니다.

한 줄 요약
PyTorch 2.14의 torch.autograd.graph.node_creation_hook은 역전파 그래프의 각 노드가 만들어질 때 바로 실행됩니다. 나중에 그래프를 다시 훑지 않고, 그 시점에 메타데이터를 붙이거나 훅을 등록할 수 있습니다. 동기는 역전파 메모리를 그 메모리를 만든 앞 구간에 연결하는 일입니다.

예전에는 역전파가 끝난 뒤에야 “이 메모리가 어느 앞 연산에서 왔는지”를 추측해야 했습니다. 오늘은 작은 모델에서 훅을 등록한 뒤, 노드가 생길 때마다 연산 이름·출력 shape를 리스트에 쌓아 보십시오. 참가자 노트에는 훅이 호출된 횟수, 남긴 필드, 훅을 끈 뒤 리스트가 비는지를 적습니다. “메모리가 줄었다”보다 “노드가 생길 때 기록이 붙었는지”가 재현에 도움이 됩니다.

재현해 볼 실험 조건

  1. 작은 선형층 한 번과 loss.backward()만 있는 경로에서 torch.autograd.graph.node_creation_hook을 등록합니다.
  2. 훅 안에서 노드 타입(또는 이름)과 출력 텐서 정보를 리스트에 남기고, backward가 끝난 뒤 리스트 길이가 0이 아닌지 확인합니다.
  3. 같은 코드를 훅 없이 한 번 더 돌려, 기록이 훅 등록에만 의존하는지 비교합니다. 필요하면 커스텀 autograd.Function 한 개에도 같은 훅이 도는지 봅니다.

같은 릴리스의 ctx.set_output_grad_dtypecdist/pdist 이중 역전파는 다른 주제입니다. 이번 글에서는 노드가 생길 때 훅이 도는지에만 초점을 둡니다. 실험이 끝나면 사용한 PyTorch 버전·모델 한 줄·훅 등록 코드를 고정해 두면, 다음 참가자가 같은 조건을 바로 따라올 수 있습니다.

실무에서 바로 볼 포인트

관련 DAKER 학습

설명용 생성 이미지입니다. 동작은 로컬 autograd 훅 로그로 확인하십시오.