본문 바로가기
C.W.K.
Stream
Lesson 02 of 10 · published

계산 그래프와 grad_fn

~12 min · graph, grad_fn, leaf

Level 0텐서 탐구자
0 XP0/62 lessons0/13 achievements
0/120 XP to next level120 XP to go0% complete

모든 연산은 노드를 만들어

requires_grad=True인 텐서에 연산을 적용하면 PyTorch가 그래프 노드를 만들어. 이 노드는 어떤 연산이 텐서를 만들었는지(grad_fn), 그리고 어떤 입력이 그 연산에 들어갔는지를 기억해. 그래야 역전파할 때 입력 쪽으로 되짚어 갈 수 있어.

텐서를 직접 살펴보면 차이가 보여:

  • 리프 텐서: nn.Parametertorch.tensor(..., requires_grad=True)처럼 직접 만든 텐서야. grad_fnNone이고 is_leaf는 True야. 기울기가 쌓이는 곳이지.
  • 중간 텐서: 연산이 만든 출력 텐서야. grad_fnAddBackward0, MulBackward0 같은 역전파 함수를 가리켜. 리프 텐서가 아니므로 기본 설정에서는 자신의 .grad를 저장하지 않아.

역전파로 그래프를 거슬러 가기

loss.backward()를 호출하면 autograd가 손실에서 시작해 grad_fn으로 연결된 모든 노드를 거슬러 가. 각 노드는 자신의 국소 야코비안, 더 정확히는 벡터-야코비안 곱을 계산하는 법을 알고 있어. 연쇄 법칙은 이 결과들을 조합해 각 리프 텐서의 최종 기울기를 구해.

한 가지 기억할 점이 있어. 기본 설정에서는 메모리를 아끼려고 역전파 직후 그래프를 해제해. 같은 그래프에서 다시 역전파해야 할 때만 retain_graph=True를 넘겨. 보통은 필요하지 않아.

Code

계산 그래프 살펴보기·python
import torch

x = torch.tensor(2.0, requires_grad=True)
w = torch.tensor(3.0, requires_grad=True)
b = torch.tensor(1.0, requires_grad=True)

z = w * x          # MulBackward0
y = z + b          # AddBackward0
loss = y ** 2      # PowBackward0

# Leaves
print(x.is_leaf, x.grad_fn)   # True None
print(w.is_leaf, w.grad_fn)   # True None

# Intermediates
print(z.is_leaf, z.grad_fn)   # False <MulBackward0>
print(y.is_leaf, y.grad_fn)   # False <AddBackward0>
print(loss.grad_fn)           # <PowBackward0>
선형 회귀 같은 작은 그래프 통한 역전파·python
import torch

x = torch.tensor(2.0, requires_grad=True)
w = torch.tensor(3.0, requires_grad=True)
b = torch.tensor(1.0, requires_grad=True)

z = w * x          # 6
y = z + b          # 7
loss = y ** 2      # 49
loss.backward()

# Gradients via chain rule
# dloss/dy = 2y       = 14
# dy/dz   = 1         → dloss/dz = 14
# dy/db   = 1         → dloss/db = 14
# dz/dw   = x = 2     → dloss/dw = 28
# dz/dx   = w = 3     → dloss/dx = 42
print(x.grad, w.grad, b.grad)
# tensor(42.) tensor(28.) tensor(14.)
retain_graph: 역전파 두 번 필요할 때·python
import torch

x = torch.tensor(2.0, requires_grad=True)
y = x ** 3

# First backward — frees graph by default
y.backward()
print(x.grad)         # 12.

# Second backward without retain_graph errors:
# y = x ** 3
# y.backward()  -- works because we built a fresh graph

# If you wanted to backward TWICE on the SAME graph:
x.grad = None
y = x ** 3
y.backward(retain_graph=True)
y.backward()          # second pass on the same retained graph
print(x.grad)         # 24. (gradients accumulate)

External links

Exercise

계산 그래프를 만들어 봐. x는 requires_grad=True인 리프 텐서고, y = x*2, z = y + 3, loss = z**2로 정의해. x=1일 때 손실을 x로 미분한 값을 손으로 예측한 뒤 .backward()를 실행해 확인해. 마지막으로 중간 텐서인 y와 z의 .grad가 모두 None인지 검증해.

Progress

Progress is local-only — sign in to sync across devices.
이 페이지에서 버그를 발견하셨거나 피드백이 있으세요?문제 신고

댓글 0

🔔 답글 알림 (로그인 필요)
로그인댓글을 남기려면 로그인해 주세요.

아직 댓글이 없어요. 첫 댓글을 남겨보세요.