Post

[PyTorch] 연산 그래프 제어

[PyTorch] 연산 그래프 제어

detach()

.detach()는 PyTorch의 연산 그래프를 끊어버리는 역할을 한다.

예를 들어, 신경망의 입력 $x$가 두 개의 layer $f$와 $g$를 지나 최종 예측값 $\hat{y}$로 만들어지는 과정이 있다고 해보자.

\[x \xrightarrow{f} a \xrightarrow{g} \hat{y}\]

여기서 예측값 $\hat{y}$과 실제 정답 $y$와의 손실 $\mathcal{L}(\hat{y},y)$을 계산해서 $x$를 최적화하고 싶다면, 역전파를 통해 아래와 같이 업데이트해야 한다.

\[x\leftarrow x-\eta\nabla_x\mathcal{L}\]

이때, 그라디언트 $\nabla_x\mathcal{L}$은 chain rule에 의해 연산의 역순으로 아래와 같이 계산된다.

\[\nabla_x\mathcal{L}=\frac{\partial\mathcal{L}}{\partial x}=\frac{\partial\mathcal{L}}{\partial\hat{y}}\cdot\frac{\partial\hat{y}}{\partial a}\cdot\frac{\partial a}{\partial x}\]

fig1

만약 y_detached=y_hat.detach()를 하게 되면, $\hat{y}$과 이전 연산 과정 사이의 연결 고리가 완전히 끊어지게 되고, y_detached를 이용해 손실 함수 $\mathcal{L}$을 계산하고 역전파 loss.backward()를 시도하면 아래와 같은 일이 벌어진다.

  • PyTorch는 y_detached를 기존 그래프에서 파생된 결과물이 아니라, 연산 기록이 없는 완전히 새로운 상수 또는 새로운 출발점으로 취급한다.
  • 따라서 chain rule의 가장 첫 단추인 $\frac{\partial\mathcal{L}}{\partial\hat{y}}$ 정보가 이전 변수인 $a$나 $x$로 넘어가지 못하고 차단된다.
  • 결과적으로 $\frac{\partial\mathcal{L}}{\partial x}$는 계산되지 않으며, $x$에 대한 그라디언트에는 값이 흐르지 않아 모델의 입력이나 가중치를 전혀 업데이트할 수 없게 된다.

fig2

1
2
3
4
5
6
7
8
9
10
11
12
13
14
import torch

x = torch.tensor([5.0], requires_grad=True)
a = x * 2
y_hat = a + 3

y_hat_detached = y_hat.detach()

y = torch.tensor([10.0])

loss = y - y_hat_detached

# dL/dx 계산
grad_x = torch.autograd.grad(loss, x)[0]    # 에러 발생!
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
import torch

x = torch.tensor([5.0], requires_grad=True)
a = x * 2
y_hat = a + 3

y = torch.tensor([10.0])

loss = y - y_hat

# dL/dx 계산
grad_x = torch.autograd.grad(loss, x)[0]
print(grad_x)
--------------------------------------------------------------
>>> tensor([-2.])

requires_grad

requires_grad는 PyTorch에게 해당 텐서에 대한 연산 기록을 추적할지 말지를 지시하는 속성이다.

연산 기록을 남기는 것은 메모리를 많이 잡아먹기 때문에, 학습이 필요한 텐서가 아니라면 requires_grad=False로 설정하여 미분 연산을 방지한다.

만약 어떤 텐서 xrequires_grad=True로 생성하였다면, PyTorch의 Autograd (자동 미분 엔진)는 그 순간부터 x가 거치는 모든 수학적 연산 과정을 백그라운드에서 추적하고 기록한다.

requires_grad=False 속성의 텐서와 requires_grad=True 속성의 텐서가 연산을 하게 된다면, 그 연산의 결과물도 자동으로 requires_grad=True가 되고 연산 과정이 기록되게 된다. 하지만, requires_grad=False 속성의 텐서들은 업데이트되지는 않는다.

PyTorch에서 requires_grad=True인 텐서를 가지고 어떤 수학적 연산을 수행하면, 결과로 나오는 텐서에는 grad_fn이라는 속성이 자동으로 부여된다.

예를 들어 텐서끼리 곱셈을 했다면 grad_fn=<MulBackward0>, 덧셈을 했다면 grad_fn=<AddBackward0> 같은 식별자가 출력창에 함께 나타난다.

1
2
3
4
5
6
7
8
import torch

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

--------------------------------------------------------------
>>> tensor([6.], grad_fn=<MulBackward0>)

torch.no_grad()

requires_grad가 개별 텐서 단위의 설정이라면, torch.no_grad()는 코드 블록 전체의 연산 기록을 멈추는 역할을 한다.

블록 내부에 있는 텐서들이 원래 requires_grad=True 속성이었더라도, 이 블록 안에서는 무조건 무시되고 기울기 추적이 중단된다.

1
2
3
4
5
import torch

with torch.no_grad()
    x = torch.tensor([2.0], requires_grad=True)
    y = torch.tensor([3.0], requires_grad=False)
This post is licensed under CC BY 4.0 by the author.