Mở đầu: cùng bài toán, RNN sập về đoán bừa — attention không suy chuyển
Cho $1$ chuỗi $80$ ký hiệu, trong đó CHỈ ĐÚNG $1$ vị trí (ngẫu nhiên) mang tín hiệu thật, còn lại toàn nhiễu — nhiệm vụ: tìm đúng tín hiệu đó, bất kể nó nằm ở đâu. Verify bằng số thật: mạng RNN học HOÀN HẢO ($100\%$) khi chuỗi dài $20$ bước, nhưng sập thẳng về mức đoán bừa ($\approx50\%$) khi chuỗi dài $40$–$80$ bước. Cùng bài toán đó, attention giữ nguyên $100\%$ ở MỌI độ dài, không suy chuyển dù chuỗi dài gấp $4$ lần.
Đây không phải do attention "thông minh hơn" một cách mơ hồ — có lý do toán học cụ thể, và bài này verify từng mảnh: vì sao RNN "quên" (gradient bùng nổ/tiêu biến qua chuỗi dài — nợ hẹn từ Bài 7), và vì sao attention không mắc phải vấn đề đó.
embeddingLookup()). Tài nguyên ngoài:
PyTorch — nn.RNN,
Bahdanau et al. 2014 — Neural Machine Translation by Jointly Learning to Align and Translate.
1. Dữ liệu chuỗi khác gì: thứ tự mang nghĩa, độ dài biến thiên
MLP (Bài 6) nhận vector kích thước CỐ ĐỊNH — $784$ pixel MNIST luôn là $784$ số. CNN (Bài 11) xử lý được ảnh nhiều kích thước qua pooling, nhưng vẫn giả định cấu trúc LƯỚI 2D cố định vị trí tương đối. Câu văn thì khác hẳn ở $2$ điểm:
Thứ tự mang nghĩa. "Chó cắn người" và "người cắn chó" dùng ĐÚNG $3$ từ giống hệt nhau — chỉ đổi thứ tự mà nghĩa đảo ngược hoàn toàn. Một MLP nhận input dạng "túi từ" (bag of words, đếm tần suất không quan tâm thứ tự) sẽ thấy $2$ câu này Y HỆT NHAU.
Độ dài biến thiên. Câu có thể dài $3$ từ hay $30$ từ — MLP cần số input cố định ngay từ kiến trúc, không có cách tự nhiên xử lý input "co giãn".
Cần một kiến trúc có bộ nhớ — xử lý chuỗi từng phần tử một, mang theo trạng thái tích luỹ được từ các phần tử trước, và hoạt động với BẤT KỲ độ dài nào. Đó là ý tưởng của RNN (Recurrent Neural Network — mạng hồi quy).
2. RNN: unroll theo thời gian, hidden state là bộ nhớ nén
Công thức lõi — mỗi bước thời gian $t$ nhận input $x_t$ VÀ trạng thái ẩn bước trước $h_{t-1}$, tạo ra trạng thái ẩn mới:
$$h_t = \tanh(W_{xh} x_t + W_{hh} h_{t-1} + b_h)$$
$\tanh$ (Bài này mới cài, nén về $(-1,1)$, đối xứng gốc toạ độ — khác $\text{sigmoid}$ nén về $(0,1)$) là activation chuẩn của RNN cell. Điểm mấu chốt: CÙNG $1$ bộ trọng số $W_{xh}, W_{hh}, b_h$ được DÙNG LẠI ở MỌI bước thời gian — đây là chia sẻ trọng số theo THỜI GIAN, đối xứng đẹp với CNN (Bài 11) chia sẻ trọng số theo KHÔNG GIAN (cùng $1$ kernel quét mọi vị trí ảnh). $h_t$ đóng vai trò bộ nhớ nén: nó phải "gói" toàn bộ thông tin cần nhớ từ $x_1, \ldots, x_t$ vào ĐÚNG $1$ vector kích thước cố định.
"Unroll" (trải phẳng) nghĩa là vẽ ra đủ $T$ bản sao của cell RNN nối tiếp nhau — một khi đã unroll, graph
tính toán trông giống hệt các graph đã gặp từ Bài 7: mỗi phép toán là
$1$ nút, forward lưu giá trị, backward đi ngược topo. Backprop qua thời gian (BPTT —
backpropagation through time) không phải lý thuyết mới — nó CHÍNH LÀ backward()
đã xây ở Bài 7, áp thẳng lên graph đã unroll, không cần công thức nào khác.
function forwardRNN(params, window) {
let h = new Tensor(new Float32Array(H), [1, H]); // h_0 = 0
for (const idx of window) {
const x = embeddingLookup(params.emb, [idx]); // tra embedding tu Bai 12
const z = add(add(matmul(x, params.Wxh), matmul(h, params.Whh)), params.bh);
h = tanh(z); // CUNG 1 bo Wxh/Whh/bh dung lai MOI buoc
}
return add(matmul(h, params.Wo), params.bo); // du doan tu h CUOI CUNG
}
// goi .backward() tren loss binh thuong - BPTT tu dong xay ra qua vong lap tren
3. Vanishing/exploding trên chuỗi dài: món nợ Bài 7 tới hạn trả
Bài 7 Mục 5 đã nếm trước: tích nhiều đạo hàm liên tiếp $<1$ tiêu biến về $0$ theo cấp số nhân, tích nhiều đạo hàm $>1$ bùng nổ. RNN unroll qua $T$ bước chính là tích CHÍNH XÁC $T$ đạo hàm liên tiếp — gradient tại bước $0$ phải "xuyên qua" toàn bộ $T$ bước đó. Đây là lý do cụ thể RNN gặp khó với chuỗi dài mà MLP/CNN (không có phép nhân lặp theo chiều sâu THỜI GIAN) không gặp phải.
Verify bằng số thật (khởi tạo $W_{hh}$ CHỦ Ý lớn, $\sigma=2{,}0$, để rơi vào vùng bùng nổ): đo chuẩn gradient tại bước $0$ khi lan truyền ngược qua $T$ bước:
| $T$ (số bước) | Chuẩn gradient tại bước 0 |
|---|---|
| $1$ | $11{,}9$ |
| $5$ | $19{,}81$ |
| $10$ | $599{,}2$ |
| $20$ | $249{,}300$ |
| $40$ | $75{,}450{,}000$ |
Từ $11{,}9$ lên hơn $75$ triệu chỉ qua $40$ bước — một bước cập nhật với gradient cỡ đó sẽ phá huỷ toàn bộ trọng số ngay lập tức (NaN gần như chắc chắn). Gradient clipping chữa cháy: nếu chuẩn gradient vượt ngưỡng $\text{maxNorm}$, co lại đúng bằng ngưỡng đó, giữ nguyên HƯỚNG:
$$g \leftarrow g \cdot \min\left(1, \frac{\text{maxNorm}}{\|g\|}\right)$$
Verify: gradient chuẩn $75{,}450{,}000$ ở $T=40$, áp $\text{maxNorm}=5$ — chuẩn SAU clip đúng bằng $5{,}0$ (verify chính xác), hướng gradient giữ nguyên tuyệt đối (chỉ co độ dài, không đổi phương) — tỉ lệ mọi thành phần giữ y hệt nhau trước/sau clip.
LSTM (Long Short-Term Memory) và GRU (Gated Recurrent Unit) giải quyết vanishing gradient bằng cổng (gate) — các vector $0$–$1$ học được quyết định "quên bao nhiêu", "nhớ bao nhiêu", "xuất bao nhiêu" mỗi bước, cho phép gradient đi qua đường "highway" gần như không suy giảm ngay cả qua hàng trăm bước. Series này KHÔNG cài LSTM/GRU đầy đủ (nhiều tham số + cổng hơn hẳn RNN thuần) — bài học ở đây chỉ dừng ở mức khái niệm; thực tế công nghiệp ngày nay phần lớn đã bỏ qua cả RNN/LSTM để dùng thẳng attention (Mục 4), chính vì lý do sẽ verify ngay sau đây.
4. Attention — bước ngoặt: nhìn thẳng mọi vị trí, không nén nữa
Thay vì ép TOÀN BỘ quá khứ vào $1$ vector $h_t$ kích thước cố định (RNN), attention cho phép mô hình nhìn thẳng vào MỌI vị trí trong chuỗi mỗi khi cần, với trọng số HỌC ĐƯỢC quyết định vị trí nào đáng chú ý hơn. Công thức cơ bản ($3$ bước):
$$\text{score}_i = q \cdot k_i \qquad w = \text{softmax}(\text{score}) \qquad \text{context} = \sum_i w_i v_i$$
$q$ (query) là vector "đang tìm gì"; $k_i$ (key) là vector đại diện vị trí $i$ trong chuỗi;
$\text{score}_i$ đo độ "khớp" giữa query và từng vị trí; $\text{softmax}$ (Bài này mới cài, thuần — khác
softmaxCrossEntropy Bài 10 vì không gộp loss, trả về đúng phân bố xác suất) biến điểm số
thành trọng số cộng lại bằng $1$; $v_i$ (value) là nội dung thực sự lấy ra, $\text{context}$ là tổng có
trọng số của mọi vị trí. Bài này dùng bản ĐƠN GIẢN NHẤT ($k_i=v_i=$ embedding thô, $q$ là $1$ vector học
được duy nhất) — multi-head attention với $Q$/$K$/$V$ chiếu riêng để dành cho
Bài 14.
- Score (tích vô hướng): $\text{score}_1 = 0{,}25$, $\text{score}_2 = 1{,}3$, $\text{score}_3 = 0{,}55$ — vị trí $2$ khớp query nhất.
- Softmax: $w = (0{,}192,\ 0{,}549,\ 0{,}259)$ — tổng đúng bằng $1$, vị trí $2$ nhận trọng số chú ý cao nhất, đúng như điểm số dự đoán.
-
Context (tổng có trọng số): $\text{context} = 0{,}192 \times k_1 + 0{,}549 \times
k_2 + 0{,}259 \times k_3 \approx (0{,}558,\ 0{,}692)$ — nghiêng nhiều nhất về phía $k_2$ nhưng vẫn
pha trộn thông tin từ cả $3$ vị trí, không chọn cứng "tất cả hoặc không gì" như một phép
argmax.
2 hệ quả lịch sử của công thức trên:
1. Không còn quên. Gradient từ context về BẤT KỲ vị trí $i$ nào đi qua đúng $1$ bước (softmax + nhân trọng số) — không xuyên qua $T$ bước tanh liên tiếp như RNN. Khoảng cách trong chuỗi KHÔNG còn làm gradient suy biến.
2. Song song hoá được. RNN buộc phải tính $h_1$ xong mới tính được $h_2$ (phụ thuộc tuần tự) — không thể chia việc ra nhiều lõi CPU/GPU cùng lúc. Attention tính $\text{score}_i$ cho MỌI $i$ ĐỘC LẬP nhau — xử lý song song toàn bộ chuỗi cùng lúc, khai thác được phần cứng hiện đại. Đây là lý do trực tiếp mở đường cho Transformer (Bài 14).
5. Thực hành: verify "không còn quên" + attention thật trên Truyện Kiều
Demo $1$ — bài toán "nhớ xa": chuỗi $T$ ký hiệu, toàn nhiễu ngoại trừ ĐÚNG $1$ vị trí ngẫu nhiên mang tín hiệu ($2$ loại "marker", nhãn = loại marker đó). Train RNN và Attention CÙNG bài toán, CÙNG số epoch, chỉ khác kiến trúc:
| Kiến trúc | Train accuracy | Val accuracy |
|---|---|---|
| RNN | — | — |
| Attention | — | — |
Thử T=20 trước (cả 2 học tốt), rồi T=40/80 (RNN sập, Attention không đổi) — đúng kết quả đã verify: T=20 cả hai 100%; T=40/80 RNN về ~50% trong khi Attention vẫn 100%.
Demo $2$ — attention THẬT trên câu Truyện Kiều: train mô hình phát hiện câu thơ có chứa từ "hoa" hay không (cân bằng $130$ câu có/$130$ câu không, $208$ train/$52$ val). Verify: train accuracy $100\%$, val accuracy $100\%$. Gõ $1$ câu thơ (hoặc dùng câu mẫu) để xem trọng số attention rê chuột từng từ:
Mỗi từ tô đậm theo trọng số attention nó nhận được — với câu có "hoa", gần như toàn bộ trọng số dồn vào đúng từ đó bất kể nó đứng đầu hay cuối câu.
Đối chiếu công nghiệp: PyTorch nn.RNN — và vì sao thời nay hiếm dùng trực tiếp:
# Doi chieu 1-1 voi Muc 2 (RNN cell) va Muc 4 (attention co ban)
import torch.nn as nn
rnn = nn.RNN(input_size=D, hidden_size=H, batch_first=True) # dung Wxh/Whh/bh Muc 2
# nn.LSTM / nn.GRU ton tai (cong Muc 3) nhung it dung truc tiep trong kien truc
# moi hien nay - attention (duoi day) giai quyet triet de hon van de "quen"
query = nn.Parameter(torch.randn(D, 1))
scores = keys @ query # (T,1) - dung Muc 4
weights = torch.softmax(scores, dim=0)
context = (weights * values).sum(dim=0) # KHONG suy giam theo T, khong tuan tu
Tóm lược
- ✅ Chuỗi khác ảnh/vector cố định: thứ tự mang nghĩa, độ dài biến thiên — cần kiến trúc có bộ nhớ.
-
✅ RNN unroll theo thời gian, chia sẻ trọng số theo THỜI GIAN (đối xứng CNN chia sẻ
theo KHÔNG GIAN) — BPTT chỉ là
backward()Bài 7 áp trên graph đã unroll. - ✅ Vanishing/exploding là hệ quả trực tiếp của tích $T$ đạo hàm liên tiếp (nợ Bài 7); verified gradient bùng nổ $11{,}9\to75{,}450{,}000$ qua $40$ bước; gradient clipping co về đúng ngưỡng, giữ nguyên hướng; LSTM/GRU dùng cổng để giảm vấn đề này (khái niệm, không cài).
- ✅ Attention (score/softmax/weighted-sum) không nén quá khứ, nhìn thẳng mọi vị trí — verified KHÔNG suy giảm dù chuỗi dài gấp $4$ lần ($100\%$ ở $T=80$ so với RNN sập về $\approx50\%$), và song song hoá được (không như RNN tuần tự).
- ✅ Attention thật trên Truyện Kiều: phát hiện từ "hoa" đạt $100\%$ train/val accuracy, trọng số tự tập trung $>90\%$ vào đúng vị trí từ đó bất kể nằm đầu hay cuối câu.
Tải file code thực hành minh họa bài học
File JavaScript verify RNN vs Attention trên bài toán nhớ xa, gradient bùng nổ + clipping, và attention
thật tìm từ "hoa" trên Truyện Kiều — bằng số đo được (chạy node rnn_attention_demo.js, mất
khoảng $10$ giây):
Bình luận