4.2 LSTM과 GRU

판다·6일 전

1. 바닐라 RNN의 한계

앞서 학습한 기본 형태의 RNN은 보통 바닐라 RNN(Vanilla RNN) 이라고 부른다. Keras에서는 SimpleRNN이 이에 해당한다.

바닐라 RNN은 이전 시점의 은닉 상태를 다음 시점으로 전달하기 때문에 순서가 있는 데이터를 처리할 수 있다. 하지만 시퀀스가 길어질수록 초반 정보가 뒤쪽까지 충분히 전달되지 못한다는 문제가 있다.

예를 들어 다음 문장을 보자.

모스크바에 여행을 왔는데 건물도 예쁘고 먹을 것도 맛있었어.
그런데 글쎄 직장 상사한테 전화가 왔어. 어디냐고 묻더라고. 그래서 나는 말했지.
저 여행 왔는데요. 여기 ___

마지막 빈칸에 들어갈 장소를 예측하려면 문장 초반의 모스크바라는 정보를 기억해야 한다. 그러나 기본 RNN은 시점이 길어질수록 앞쪽 정보를 잊어버릴 수 있어, 적절한 단어를 예측하지 못할 수 있다.

이처럼 긴 시퀀스에서 오래전 정보를 제대로 유지하지 못하는 문제를 장기 의존성 문제(Long-Term Dependencies) 라고 한다. wikidocs

2. 바닐라 RNN 복습

바닐라 RNN은 현재 입력 (xt)와 이전 은닉 상태 (h{t-1})를 활용하여 현재 은닉 상태 (h_t)를 계산한다.

[
ht = \tanh(W_xx_t + W_hh{t-1} + b)
]

기호의미
(x_t)현재 시점의 입력
(h_{t-1})이전 시점의 은닉 상태
(h_t)현재 시점의 은닉 상태
(W_x)현재 입력에 적용되는 가중치
(W_h)이전 은닉 상태에 적용되는 가중치
(b)편향
(\tanh)하이퍼볼릭 탄젠트 활성화 함수

기본 RNN은 구조가 단순하다는 장점이 있지만, 긴 문장·긴 시계열처럼 장기간의 맥락을 반영해야 하는 문제에서는 한계를 보인다. LSTM과 GRU는 이 한계를 보완하기 위해 제안된 RNN 계열 모델이다. wikidocs

3. LSTM이란?

LSTM(Long Short-Term Memory)은 긴 시퀀스에서 중요한 정보를 비교적 오래 유지하도록 설계된 RNN의 변형이다.

LSTM은 일반 RNN의 은닉 상태 외에 셀 상태(Cell State) 를 추가로 관리한다. 셀 상태는 긴 시간 동안 유지할 정보를 전달하는 통로이며, LSTM은 여러 게이트를 통해 어떤 정보를 저장하고, 삭제하고, 출력할지 결정한다. wikidocs

LSTM에는 다음 세 가지 핵심 게이트가 있다.

게이트역할직관적 의미
망각 게이트(Forget Gate)이전 기억을 얼마나 유지할지 결정무엇을 잊을까?
입력 게이트(Input Gate)현재 입력 중 무엇을 저장할지 결정무엇을 새로 기억할까?
출력 게이트(Output Gate)셀 상태 중 무엇을 은닉 상태로 출력할지 결정무엇을 현재 결과로 보여줄까?

각 게이트에는 시그모이드 함수가 사용된다.

[
\sigma(x) = \frac{1}{1 + e^{-x}}
]

시그모이드 함수의 출력은 0과 1 사이이다. 따라서 게이트 값이 0에 가까우면 해당 정보는 거의 통과하지 못하고, 1에 가까우면 정보가 많이 유지된다. wikidocs

4. LSTM의 동작 과정

입력 게이트

입력 게이트는 현재 입력 중에서 새롭게 기억할 정보를 선택한다.

[
it = \sigma(W{xi}xt + W{hi}h_{t-1} + b_i)
]

[
gt = \tanh(W{xg}xt + W{hg}h_{t-1} + b_g)
]

  • (i_t): 새 정보를 얼마나 저장할지 결정하는 값
  • (g_t): 현재 입력으로부터 생성한 후보 기억 정보

입력 게이트의 결과는 (i_t \circ g_t)로 계산된다. 즉, 후보 기억 (g_t) 중에서 실제로 셀 상태에 반영할 정보의 양을 (i_t)가 조절한다.

망각 게이트

망각 게이트는 이전 셀 상태의 정보를 얼마나 유지할지 결정한다.

[
ft = \sigma(W{xf}xt + W{hf}h_{t-1} + b_f)
]

망각 게이트의 값이 0에 가까우면 이전 기억을 많이 삭제하고, 1에 가까우면 이전 기억을 대부분 유지한다.

예를 들어 문장 주제가 바뀌었다면 이전 주제와 관련된 기억을 줄이고, 현재 문맥과 관련된 정보에 더 집중하도록 학습할 수 있다. wikidocs

셀 상태

LSTM의 핵심은 셀 상태 (C_t)이다. 셀 상태는 장기적으로 유지할 정보를 담기 때문에 장기 상태(long-term state) 라고도 부른다.

[
Ct = f_t \circ C{t-1} + i_t \circ g_t
]

이 수식은 두 부분으로 해석할 수 있다.

  • (ft \circ C{t-1}): 이전 기억 중 유지할 정보
  • (i_t \circ g_t): 현재 입력에서 새롭게 추가할 정보

즉, LSTM은 이전 기억을 무조건 전달하는 것이 아니라, 망각 게이트와 입력 게이트를 통해 선택적으로 기억을 유지·삭제·추가한다. wikidocs

출력 게이트와 은닉 상태

출력 게이트는 셀 상태 중 어떤 정보를 현재 은닉 상태로 내보낼지 결정한다.

[
ot = \sigma(W{xo}xt + W{ho}h_{t-1} + b_o)
]

[
h_t = o_t \circ \tanh(C_t)
]

  • (C_t): 장기적으로 유지되는 셀 상태
  • (h_t): 현재 시점에서 모델이 출력하는 은닉 상태
  • (o_t): 셀 상태를 얼마나 출력에 반영할지 결정하는 값

은닉 상태 (h_t)는 현재 시점의 출력에 직접 활용되므로 단기 상태(short-term state) 라고도 부른다. wikidocs

5. PyTorch LSTM 구현

PyTorch에서는 nn.LSTM()으로 LSTM 레이어를 구현할 수 있다.

import torch
import torch.nn as nn

input_size = 5
hidden_size = 8

lstm = nn.LSTM(
    input_size=input_size,
    hidden_size=hidden_size,
    batch_first=True
)

inputs = torch.randn(1, 10, 5)

outputs, (hidden_state, cell_state) = lstm(inputs)

print(outputs.shape)
print(hidden_state.shape)
print(cell_state.shape)

입력 텐서의 형태는 다음과 같다.

(batch_size, sequence_length, input_size)

예시 코드에서는 다음을 의미한다.

값의미
1배치 크기
10시퀀스 길이 또는 타임스텝 수
5각 시점 입력 벡터의 차원
8은닉 상태와 셀 상태의 차원

LSTM의 반환값은 다음과 같다.

outputs, (hidden_state, cell_state)
  • outputs: 모든 시점의 은닉 상태
  • hidden_state: 마지막 시점의 은닉 상태 (h_t)
  • cell_state: 마지막 시점의 셀 상태 (C_t)

기본 RNN은 마지막 은닉 상태만 관리하지만, LSTM은 은닉 상태와 셀 상태를 함께 관리한다는 점이 차이점이다. wikidocs

6. GRU란?

GRU(Gated Recurrent Unit)는 LSTM의 장기 의존성 문제 해결 방식은 유지하면서 구조를 더 단순하게 만든 RNN 모델이다.

GRU는 2014년 제안된 구조로, LSTM처럼 긴 시퀀스를 처리할 수 있도록 게이트를 사용한다. 다만 LSTM의 입력 게이트, 망각 게이트, 출력 게이트를 그대로 사용하지 않고, 다음 두 게이트로 단순화했다. wikidocs

게이트역할
리셋 게이트(Reset Gate)이전 은닉 상태를 얼마나 반영할지 조절
업데이트 게이트(Update Gate)이전 은닉 상태와 새 후보 상태를 얼마나 반영할지 조절

GRU는 별도의 셀 상태를 사용하지 않는다. 즉, LSTM의 장기 상태와 단기 상태를 분리하지 않고 하나의 은닉 상태로 관리한다.

7. GRU의 동작 원리

GRU의 리셋 게이트와 업데이트 게이트는 다음과 같이 계산한다.

[
rt = \sigma(W{xr}xt + W{hr}h_{t-1} + b_r)
]

[
zt = \sigma(W{xz}xt + W{hz}h_{t-1} + b_z)
]

새로운 후보 은닉 상태는 다음과 같다.

[
gt = \tanh(W{hg}(rt \circ h{t-1}) + W_{xg}x_t + b_g)
]

최종 은닉 상태는 이전 은닉 상태와 후보 은닉 상태를 업데이트 게이트 값으로 조합해 계산한다.

[
ht = (1-z_t) \circ g_t + z_t \circ h{t-1}
]

업데이트 게이트 (z_t)가 크면 기존 은닉 상태를 더 많이 유지하고, 작으면 새로 계산한 후보 은닉 상태를 더 많이 반영한다. 따라서 GRU도 상황에 따라 과거 정보를 유지하거나 새로운 정보로 갱신할 수 있다. wikidocs

8. PyTorch GRU 구현

PyTorch에서는 nn.GRU()를 사용한다.

import torch
import torch.nn as nn

input_size = 5
hidden_size = 8

gru = nn.GRU(
    input_size=input_size,
    hidden_size=hidden_size,
    batch_first=True
)

inputs = torch.randn(1, 10, 5)

outputs, hidden_state = gru(inputs)

print(outputs.shape)
print(hidden_state.shape)

GRU의 입력 형식은 RNN, LSTM과 동일하다.

(batch_size, sequence_length, input_size)

하지만 LSTM과 달리 GRU는 별도 셀 상태가 없으므로 다음 두 값만 반환한다.

outputs, hidden_state
  • outputs: 모든 시점의 은닉 상태
  • hidden_state: 마지막 시점의 은닉 상태

9. LSTM과 GRU 비교

구분LSTMGRU
게이트 수입력·망각·출력 게이트, 총 3개리셋·업데이트 게이트, 총 2개
셀 상태별도의 셀 상태 (C_t) 존재별도의 셀 상태 없음
구조 복잡도비교적 복잡함LSTM보다 단순함
파라미터 수상대적으로 많음상대적으로 적음
학습 속도상대적으로 느릴 수 있음일반적으로 더 빠를 수 있음
긴 시퀀스 처리우수함우수함
성능데이터와 문제에 따라 다름데이터와 문제에 따라 다름

GRU가 항상 LSTM보다 좋거나, LSTM이 항상 GRU보다 좋다고 단정할 수는 없다. 일반적으로 데이터가 적거나 모델 경량화와 빠른 학습이 중요하다면 GRU를 먼저 시도해볼 수 있다. 반대로 충분한 데이터가 있고 복잡한 장기 의존성을 정교하게 학습해야 한다면 LSTM을 실험 후보로 고려할 수 있다. 다만 실제 모델 선택은 데이터셋, 시퀀스 길이, 평가 지표, 학습 시간 등을 기준으로 실험을 통해 결정해야 한다. wikidocs

10. 정리

  • 바닐라 RNN은 긴 시퀀스에서 초반 정보를 잊어버리는 장기 의존성 문제가 있다.
  • LSTM은 입력·망각·출력 게이트와 셀 상태를 통해 중요한 정보를 장기간 유지한다.
  • LSTM의 셀 상태는 장기 상태, 은닉 상태는 단기 상태로 볼 수 있다.
  • GRU는 LSTM을 단순화한 구조로, 리셋 게이트와 업데이트 게이트를 사용한다.
  • GRU는 별도의 셀 상태 없이 하나의 은닉 상태를 관리한다.
  • LSTM과 GRU 중 어느 모델이 더 좋은지는 문제와 데이터에 따라 달라지므로 실험으로 비교하는 것이 중요하다.
profile
welcome to deeplearning

0개의 댓글