인지야공

인지야공/딥러닝 기초 정리/6번째 글

LSTM 완전 풀이 — 초등학생도 따라오는 계산 과정

실행: python 02_lstm_손계산.py 이 문서의 모든 숫자는 실제로 계산해서 PyTorch nn.LSTM 과 대조 검증했습니다. (오차 3.8e-08 수준 = 완전히 같음)


1. LSTM이 뭐야? — 일기장 이야기

우리 반에 기억이가 있어. 기억이는 매일 새로운 이야기를 듣는데, 아주 낡은 일기장 하나를 들고 다녀.

매일 이야기를 들을 때마다 기억이는 딱 4가지를 결정해:

결정이름하는 일
① “오늘 이야기, 얼마나 적을까?”입력 게이트 (i)새 정보를 얼마나 받아들일지
② “일기장 옛날 내용, 얼마나 남길까?”망각 게이트 (f)옛 기억을 얼마나 지울지
③ “오늘 적을 내용은 이거야”후보 기억 (g)새로 적을 내용 자체
④ “친구한테 얼마나 말해줄까?”출력 게이트 (o)기억 중 얼마를 밖으로 내보낼지

그리고 기억이가 가진 건 두 종류의 기억이야.

  • c (셀 상태, cell state) = 일기장 그 자체. 오래 가는 기억(장기기억)
  • h (은닉 상태, hidden state) = 지금 친구한테 말해주는 내용. 바로 쓰는 기억(단기기억)

핵심: 일기장(c)은 조용히 계속 이어지고, 말(h)은 매번 새로 만들어진다.


2. 게이트는 “수도꼭지”다

게이트는 전부 0에서 1 사이의 숫자야. 수도꼭지를 얼마나 열지랑 똑같아.

  • 0 = 꽉 잠금 (하나도 안 통과)
  • 0.5 = 반만 열림 (반만 통과)
  • 1 = 활짝 열림 (다 통과)

이 0~1 숫자를 만들어주는 게 시그모이드(sigmoid) 함수야.

sigmoid(z) = 1 / (1 + e^(-z))

숫자 감각만 잡으면 돼:

zsigmoid(z)느낌
-30.05거의 잠김
-10.27조금 열림
00.5반쯤 열림
0.50.62좀 열림
1.50.82많이 열림
30.95거의 활짝

그리고 tanh는 -1에서 1 사이 값을 만들어. 내용에는 “좋다(+)/나쁘다(-)” 방향이 있어야 하니까, 후보 기억 g 와 출력 계산에는 sigmoid 대신 tanh 를 써.

ztanh(z)
0.80.664
1.70.935
2.60.990

3. 딱 4개의 공식

기호 정리:

  • x = 이번에 들어온 숫자 (오늘 들은 이야기)
  • h = 바로 전에 내가 말했던 것
  • c = 일기장 내용
  • W = 가중치(중요도), b = 편향(기본 성향)
i = sigmoid( W_ii·x + W_hi·h + b_i )      입력 게이트
f = sigmoid( W_if·x + W_hf·h + b_f )      망각 게이트
g = tanh   ( W_ig·x + W_hg·h + b_g )      후보 기억
o = sigmoid( W_io·x + W_ho·h + b_o )      출력 게이트

c_새것 = f × c_옛것  +  i × g             ← 일기장 갱신
h_새것 = o × tanh(c_새것)                  ← 밖으로 내보낼 말

말로 읽으면

새 일기장 = (옛 일기장 × 남길 비율) + (새 내용 × 받을 비율) 내보낼 말 = 일기장 내용을 적당히 눌러서(tanh) × 말할 비율


4. 실제 계산 — 숫자를 넣어보자

가장 작은 LSTM 을 만들자. 숫자 1개가 들어가고 숫자 1개가 나오는 LSTM.

입력: x = [1.0, 2.0, 3.0] (세 번에 나눠서 들어옴) 처음 상태: h = 0, c = 0 (아무 기억 없음)

가중치 (우리가 직접 정한 값):

입력 x 에 곱할 값이전 h 에 곱할 값편향 b
입력 게이트 i0.50.20.0
망각 게이트 f1.00.30.5
후보 기억 g0.80.40.0
출력 게이트 o0.60.10.0

⏱ 첫 번째 시각 : x₁ = 1.0 (h=0, c=0)

① 입력 게이트

z = 0.5 × 1.0 + 0.2 × 0 + 0.0 = 0.5
i = sigmoid(0.5) = 0.6225

→ 새 이야기를 62%만 받아들이겠다.

② 망각 게이트

z = 1.0 × 1.0 + 0.3 × 0 + 0.5 = 1.5
f = sigmoid(1.5) = 0.8176

→ 옛 기억은 82% 남기겠다. (지금은 일기장이 비어 있어서 의미는 없음)

③ 후보 기억

z = 0.8 × 1.0 + 0.4 × 0 + 0.0 = 0.8
g = tanh(0.8) = 0.6640

→ 오늘 적을 내용은 0.6640.

④ 출력 게이트

z = 0.6 × 1.0 + 0.1 × 0 + 0.0 = 0.6
o = sigmoid(0.6) = 0.6457

→ 기억의 65%만 밖으로 말하겠다.

⑤ 일기장 갱신

c₁ = f × c₀ + i × g
   = 0.8176 × 0      +  0.6225 × 0.6640
   = 0                +  0.4133
   = 0.4133

👉 앞부분이 0인 이유: 처음엔 일기장이 비어 있었으니까(c₀=0). 지울 옛 기억이 없어.

⑥ 내보낼 말

h₁ = o × tanh(c₁)
   = 0.6457 × tanh(0.4133)
   = 0.6457 × 0.3913
   = 0.2526

1번째 결과: c = 0.4133 (일기장), h = 0.2526 (말한 것)


⏱ 두 번째 시각 : x₂ = 2.0 (h=0.2526, c=0.4133)

이번엔 직전에 말했던 h=0.2526 도 같이 들어간다는 게 포인트!

① 입력 게이트

z = 0.5 × 2.0 + 0.2 × 0.2526 + 0.0 = 1.0 + 0.0505 = 1.0505
i = sigmoid(1.0505) = 0.7409

→ 아까(0.6225)보다 더 열렸다. 입력이 커졌으니까!

② 망각 게이트

z = 1.0 × 2.0 + 0.3 × 0.2526 + 0.5 = 2.0 + 0.0758 + 0.5 = 2.5758
f = sigmoid(2.5758) = 0.9293

→ 옛 기억을 93%나 남긴다. 거의 안 지우겠다는 뜻!

③ 후보 기억

z = 0.8 × 2.0 + 0.4 × 0.2526 + 0.0 = 1.6 + 0.1010 = 1.7011
g = tanh(1.7011) = 0.9355

④ 출력 게이트

z = 0.6 × 2.0 + 0.1 × 0.2526 + 0.0 = 1.2 + 0.0253 = 1.2253
o = sigmoid(1.2253) = 0.7730

⑤ 일기장 갱신 — 여기가 진짜 핵심

c₂ = f × c₁          +  i × g
   = 0.9293 × 0.4133 +  0.7409 × 0.9355
   = 0.3841          +  0.6931
   = 1.0772

👉 0.3841 은 “옛날에 적어둔 것 중 살아남은 부분”, 👉 0.6931 은 “오늘 새로 적은 부분” 둘을 더한 게 새 일기장!

⑥ 내보낼 말

h₂ = 0.7730 × tanh(1.0772) = 0.7730 × 0.7922 = 0.6123

2번째 결과: c = 1.0772, h = 0.6123


⏱ 세 번째 시각 : x₃ = 3.0 (h=0.6123, c=1.0772)

① 입력 게이트

z = 0.5 × 3.0 + 0.2 × 0.6123 = 1.5 + 0.1225 = 1.6225
i = sigmoid(1.6225) = 0.8351

② 망각 게이트

z = 1.0 × 3.0 + 0.3 × 0.6123 + 0.5 = 3.0 + 0.1837 + 0.5 = 3.6837
f = sigmoid(3.6837) = 0.9755

→ 97.5% 유지! 거의 통째로 기억한다.

③ 후보 기억

z = 0.8 × 3.0 + 0.4 × 0.6123 = 2.4 + 0.2449 = 2.6449
g = tanh(2.6449) = 0.9900

④ 출력 게이트

z = 0.6 × 3.0 + 0.1 × 0.6123 = 1.8 + 0.0612 = 1.8612
o = sigmoid(1.8612) = 0.8654

⑤ 일기장 갱신

c₃ = 0.9755 × 1.0772 + 0.8351 × 0.9900
   = 1.0508           + 0.8268
   = 1.8776

⑥ 내보낼 말

h₃ = 0.8654 × tanh(1.8776) = 0.8654 × 0.9543 = 0.8259

3번째 결과: c = 1.8776, h = 0.8259


5. PyTorch로 검증 — 진짜 맞았을까?

똑같은 가중치를 nn.LSTM 에 넣고 돌려봤다.

시각손계산 hPyTorch h차이
t=10.2526460.2526463.80e-08
t=20.6123380.6123381.07e-08
t=30.8258690.8258694.15e-08
마지막 c1.8775771.877577—

차이가 0.00000004 수준 = 컴퓨터의 소수점 반올림 오차뿐. 완전히 같다. ✔

⚠️ 주의할 점: PyTorch 는 게이트 순서가 반드시 [i, f, g, o] 다. weight_ih_l0 은 (4×hidden, input) 모양이고, 위에서부터 i / f / g / o 순서로 쌓여 있다. 편향도 bias_ih_l0 와 bias_hh_l0 두 개로 나뉘어 있어 실제 편향은 두 개의 합이다.


6. 망각 게이트가 왜 그렇게 중요한가

i × g = 0.5 로 고정해두고, f 값만 바꿔가며 일기장 c 가 어떻게 변하는지 봤다.

시각f=0.0 (다 잊음)f=0.5 (반만 기억)f=1.0 (다 기억)
10.5000.5000.500
20.5000.7501.000
30.5000.8751.500
40.5000.9382.000
50.5000.9692.500
60.5000.9843.000
  • f=0 → 매번 초기화. 바로 직전 것밖에 모른다. (= 기억상실)
  • f=0.5 → 금방 0.5+0.25+0.125… 로 수렴. 먼 과거는 절반의 절반의 절반… 으로 사라진다.
  • f=1 → 계속 쌓인다. 10칸, 20칸 전 정보도 그대로 살아 있다. ← 이게 장기기억!

왜 옛날 RNN 은 실패했나

기본 RNN 은 h_새것 = tanh(W·h_옛것 + U·x) 인데, 역전파할 때 W 가 계속 곱해진다.

W = 0.5 이면 →  0.5^20 = 0.00000095   (기울기 소실: 학습 신호가 사라짐)
W = 1.5 이면 →  1.5^20 = 3325         (기울기 폭발: 값이 터짐)

LSTM 의 c 는 다르다. c_새것 = f × c_옛것 + (새 것) 이라서 f 만 1 근처면 곱해지는 값이 1 → 아무리 많이 지나도 안 사라진다. 이걸 CEC(Constant Error Carousel, 상수 오차 회전목마) 라고 부른다.

덧셈으로 이어지는 c 의 길 = “기억 고속도로”. 이게 LSTM 의 발명 포인트다.


7. 실전 실험 — 20칸 전 정보를 진짜 기억할까?

과제: 길이 20짜리 수열을 주고, 맨 앞(0번째) 값이 뭐였는지 마지막에 맞히기. 19칸을 건너뛰고 기억해야 하는 어려운 문제다.

같은 LSTM(hidden=16, Adam lr=0.01, 40 epoch)으로 조건만 바꿔 두 번 돌렸다.

조건40 epoch 후 정확도
A. 입력 0/1, 망각 편향 기본값(0)52.2% ← 찍기(50%)와 같음. 학습 실패
B. 입력 -1/+1, 망각 편향 1.0 초기화100.0% ← 완벽

B는 10 epoch 만에 이미 100% 를 찍었다 (loss 0.0004).

여기서 배우는 것 2가지

① 망각 게이트 편향을 1.0 근처로 초기화하라. 처음에 f = sigmoid(1.0) = 0.73 이 되어 기억이 잘 안 지워진다. 편향이 0이면 f = sigmoid(0) = 0.5 라서 매 스텝마다 기억이 반토막 → 20스텝 뒤엔 0.5^20 ≈ 0.000001 로 완전히 소멸한다.

hidden = 16
with torch.no_grad():
    # bias 순서 [i, f, g, o] 중 f 구간만 1.0 으로
    lstm.bias_ih_l0[hidden:2*hidden].fill_(1.0)

② 입력을 0/1 로 주지 마라. 0은 아무리 큰 가중치를 곱해도 0이라 게이트를 전혀 못 움직인다. -1/+1 처럼 부호가 있는 값으로 바꾸면 신호가 훨씬 잘 전달된다.

결론: LSTM 은 구조만으로 자동 해결되지 않는다. 초기화와 입력 스케일이 함께 맞아야 한다. (교과서에 잘 안 나오지만 실전에서 제일 자주 발목 잡는 부분)


8. 한 장 요약 카드

입력   xtx_t 지금 값  ·  ht−1h_{t-1} 직전 출력  ·  ct−1c_{t-1} 일기장


i=σ(Wii x+Whi h+bi)새 정보 받을 비율  (0∼1)f=σ(Wif x+Whf h+bf)옛 기억 남길 비율  (0∼1)g=tanh⁡(Wig x+Whg h+bg)새로 적을 내용  (−1∼1)o=σ(Wio x+Who h+bo)내보낼 비율  (0∼1)\begin{aligned} i &= \sigma(W_{ii}\,x + W_{hi}\,h + b_i) && \text{새 정보 받을 비율}\ \ (0 \sim 1) \\ f &= \sigma(W_{if}\,x + W_{hf}\,h + b_f) && \text{옛 기억 남길 비율}\ \ (0 \sim 1) \\ g &= \tanh(W_{ig}\,x + W_{hg}\,h + b_g) && \text{새로 적을 내용}\ \ (-1 \sim 1) \\ o &= \sigma(W_{io}\,x + W_{ho}\,h + b_o) && \text{내보낼 비율}\ \ (0 \sim 1) \end{aligned}
ct=f⊙ct−1+i⊙g← 장기기억 (덧셈!)ht=o⊙tanh⁡(ct)← 단기기억 / 출력\begin{aligned} c_t &= f \odot c_{t-1} + i \odot g && \leftarrow\ \text{장기기억 (덧셈!)} \\ h_t &= o \odot \tanh(c_t) && \leftarrow\ \text{단기기억 / 출력} \end{aligned}

파라미터 개수 =4×(input_size×hidden+hidden×hidden+2×hidden)= 4 \times (\text{input\_size} \times \text{hidden} + \text{hidden} \times \text{hidden} + 2 \times \text{hidden})

꼭 기억할 3문장

  1. c 는 덧셈으로 이어지기 때문에 오래 살아남는다 (LSTM 의 전부).
  2. 게이트 3개는 전부 수도꼭지(01), 내용 g 만 tanh(-11).
  3. GRU 는 게이트를 2개(reset, update)로 줄이고 c 와 h 를 합친 경량 버전이다.

표시는 이 브라우저에만 남는다. 서버로 가는 것은 없다.