인지야공/딥러닝 기초 정리/6번째 글
LSTM 완전 풀이 — 초등학생도 따라오는 계산 과정
실행:
python 02_lstm_손계산.py이 문서의 모든 숫자는 실제로 계산해서 PyTorchnn.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))
숫자 감각만 잡으면 돼:
| z | sigmoid(z) | 느낌 |
|---|---|---|
| -3 | 0.05 | 거의 잠김 |
| -1 | 0.27 | 조금 열림 |
| 0 | 0.5 | 반쯤 열림 |
| 0.5 | 0.62 | 좀 열림 |
| 1.5 | 0.82 | 많이 열림 |
| 3 | 0.95 | 거의 활짝 |
그리고 tanh는 -1에서 1 사이 값을 만들어. 내용에는 “좋다(+)/나쁘다(-)” 방향이 있어야 하니까, 후보 기억 g 와 출력 계산에는 sigmoid 대신 tanh 를 써.
| z | tanh(z) |
|---|---|
| 0.8 | 0.664 |
| 1.7 | 0.935 |
| 2.6 | 0.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 | |
|---|---|---|---|
| 입력 게이트 i | 0.5 | 0.2 | 0.0 |
| 망각 게이트 f | 1.0 | 0.3 | 0.5 |
| 후보 기억 g | 0.8 | 0.4 | 0.0 |
| 출력 게이트 o | 0.6 | 0.1 | 0.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 에 넣고 돌려봤다.
| 시각 | 손계산 h | PyTorch h | 차이 |
|---|---|---|---|
| t=1 | 0.252646 | 0.252646 | 3.80e-08 |
| t=2 | 0.612338 | 0.612338 | 1.07e-08 |
| t=3 | 0.825869 | 0.825869 | 4.15e-08 |
| 마지막 c | 1.877577 | 1.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 (다 기억) |
|---|---|---|---|
| 1 | 0.500 | 0.500 | 0.500 |
| 2 | 0.500 | 0.750 | 1.000 |
| 3 | 0.500 | 0.875 | 1.500 |
| 4 | 0.500 | 0.938 | 2.000 |
| 5 | 0.500 | 0.969 | 2.500 |
| 6 | 0.500 | 0.984 | 3.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. 한 장 요약 카드
입력 지금 값 · 직전 출력 · 일기장
파라미터 개수
꼭 기억할 3문장
- c 는 덧셈으로 이어지기 때문에 오래 살아남는다 (LSTM 의 전부).
- 게이트 3개는 전부 수도꼭지(0
1), 내용 g 만 tanh(-11). - GRU 는 게이트를 2개(reset, update)로 줄이고 c 와 h 를 합친 경량 버전이다.