№ 004互動
在瀏覽器裡從零訓練一個 Transformer:看注意力自己長出來
不用 PyTorch,也不用任何函式庫。用 TypeScript 寫一個會自動微分的小引擎、一個一萬多個參數的 Transformer,然後按下按鈕,看它在幾秒內學會把一串數字反過來。
- 發布
- 閱讀時間
- 7 分鐘
CNN 那篇的權重是先在 PyTorch 裡訓練好、再搬進瀏覽器的。這一篇不一樣:訓練本身就發生在你的瀏覽器裡。下面這個模型現在的權重完全是隨機的,它什麼都不會。
我們要教它玩一個很簡單的遊戲:給它六個數字,請它倒過來寫一遍。看到 3 1 4 1 5 9,就要回答 9 5 1 4 1 3。
對人來說這不用學,但模型一開始連「倒過來」是什麼意思都不知道。沒有人會告訴它規則,它只會一次又一次看到題目和正確答案,然後自己想辦法。
下面的儀器由上到下是三個步驟:先看清楚任務,按「開始訓練」,看最上面「它現在寫出來的」那一行從粉紅色(錯)一位一位變成青色(對),最後往下看注意力圖,那裡畫的是模型寫每一位數字時正在看輸入的哪裡。
讀六個數字,然後把它們倒過來寫。
- 模型讀到
- 314159→
- 正確答案
- 951413
- 它現在寫出來的
還沒開始訓練。現在的權重是隨機的,所以上面它寫出來的答案是亂猜的。
注意力圖:它寫每一位的時候,在看輸入的哪裡
注意力頭 1
↓ 正在寫的位置
注意力頭 2
↓ 正在寫的位置
還沒訓練的時候,亮的位置是隨機的,沒有意義。按「開始訓練」,看它怎麼變。
如果一切順利,你會在一百步之內看到三件事同時發生:損失掉到接近零、它寫出來的答案整串變成青色,以及其中一張注意力圖的下半部長出一條從右上到左下的斜線。
那條斜線不是我畫上去的,也沒有任何一行程式碼叫它長成那樣。它是訓練的結果。這篇文章要講的就是:它是怎麼長出來的。
為什麼是六個數字,而不是一個會說話的模型
Transformer 的教學通常拿莎士比亞當訓練資料,訓練個幾十分鐘,得到一段看起來像英文的亂碼。這有兩個問題:在瀏覽器裡太慢,而且你很難判斷模型到底學會了什麼。
所以這裡換一個玩具任務:輸入六個數字,輸出它們的某種重新排列。它的好處是:
- 對錯一目了然。答案不是「看起來像不像」,而是每一位數都對或不對。
- 幾秒鐘就練得起來。模型只需要一萬多個參數。
- 注意力圖可以直接讀。要寫出答案的第 1 位,模型必須去讀輸入的第 6 位;這件事會原封不動地出現在注意力矩陣上。
至於為什麼剛好是六個:長度固定,模型才能做得夠小,也才能在一張 12×12 的圖上把它的注意力整個畫出來。模型內部有一張「位置表」,這裡只做了 12 格(六個輸入、一個箭頭、五個已經寫出來的答案),所以它只認得這個長度。這是為了看得清楚而做的取捨,不是 Transformer 本身的限制。
模型本身沒有任何偷工減料:它是一個標準的 decoder-only Transformer,和 GPT 同一種結構,只是很小。
| 這篇的模型 | |
|---|---|
| 詞彙表 | 11 個 token:數字 0–9,加上一個「→」 |
| 序列長度 | 12 |
寬度 d | 32 |
| 注意力頭 | 2 |
| 層數 | 1 |
| 參數總數 | 13,728 |
把問題變成「預測下一個 token」
語言模型只會做一件事:看著前面的 token,猜下一個。所以我們把一題寫成一條序列:
3 1 4 1 5 9 → 9 5 1 4 1 3然後要求模型在每個位置預測下一個 token。不過前半段是隨機的數字,沒有人猜得中,所以只有後半段(答案)會被計分:
export function example(digits: number[], task: TaskName) {
const full = [...digits, SEP, ...TASKS[task](digits)];
const ids = full.slice(0, -1);
// −1 表示「這個位置不計分」
const targets = full.slice(1).map((t, i) => (i < digits.length ? -1 : t));
return { ids, targets };
}這也是為什麼注意力圖的上半部被蓋上一層灰:那幾列不影響損失,它們長什麼樣子都無所謂。
注意力:每個位置自己決定要讀誰
Transformer 的核心是一個很簡單的想法。每個位置會產生三個向量:
- query:我在找什麼?
- key:我這裡有什麼?
- value:如果你選了我,我給你這些資訊。
位置 拿自己的 query 去和每個位置 的 key 做內積,得到一個分數;分數經過 softmax 變成權重,再用這些權重把 value 加權平均:
是因果遮罩: 的位置填上 ,softmax 之後權重就是零。模型在寫答案的第 2 位時,不能偷看第 3 位。上面注意力圖的右上三角永遠是暗的,就是這個原因。
儀器裡畫的那兩張圖,就是 的結果本身:第 列第 行有多亮,代表位置 從位置 讀了多少。
const scores = tape.scale(tape.matmul(qh, tape.transpose(kh)), 1 / Math.sqrt(dh));
const weights = tape.causalSoftmax(scores); // 這就是畫出來的那張圖
mixed.push(tape.matmul(weights, vh));為什麼會長出斜線
想一下反轉任務需要什麼。寫答案的第 1 位時,模型站在「→」的位置,需要的是輸入的最後一位;寫第 2 位時,需要倒數第二位。也就是說:
第 個位置,應該去讀第 個位置。
這是一條純粹由位置決定的規則,和數字是多少無關。模型一開始完全不知道這件事。但每一次答錯,梯度都會把 query 和 key 往「讓正確的那一格分數變高」的方向推一點。推個幾十次之後,那條斜線就出現了。
「排序」就沒有這麼乾淨的圖了。排序需要的不是「去讀第幾個位置」,而是「去找還沒用過的最小數字」,這取決於內容而不是位置。它大約要 600 到 800 步才會練到 100%,注意力也散得多。一層、兩個頭的模型可以學會它,但學到的解法不是人類一眼看得懂的那種。
模型的其餘部分
注意力負責在位置之間搬運資訊,其他零件都是逐位置運作的:
- Embedding:每個 token 查表得到一個 32 維向量,再加上一個代表「我在第幾個位置」的向量。沒有位置向量,模型根本分不出第 1 位和第 6 位,反轉任務就不可能學會。
- LayerNorm:把每個位置的向量調整成平均 0、變異數 1,讓訓練穩定。
- MLP:兩層全連接,中間放大成 4 倍寬再縮回來。注意力把資訊搬過來,MLP 負責處理它。
- 殘差連接:每個子層的輸出是「加回去」而不是「取代」。
- Unembedding:最後把 32 維向量投影回 11 個 token 的分數。
// 自注意力:每個位置從前面的位置收集資訊
x = tape.add(x, tape.matmul(tape.concatCols(mixed), P[p + "wo"]));
// MLP:每個位置自己消化收集到的東西
const m = tape.layerNorm(x, P[p + "ln2.g"], P[p + "ln2.b"]);
const hidden = tape.relu(tape.addRow(tape.matmul(m, P[p + "w1"]), P[p + "b1"]));
x = tape.add(x, tape.addRow(tape.matmul(hidden, P[p + "w2"]), P[p + "b2"]));訓練:梯度從哪裡來
前向傳播只是一連串的矩陣運算。訓練需要的是另一個方向:損失對每一個參數的偏導數。13,728 個參數,就要 13,728 個數字。
手算當然不可能,所以我寫了一個很小的自動微分引擎。想法是:每做一個運算,就順手記下「如果有人告訴我輸出的梯度,我該怎麼把它傳回輸入」。以矩陣乘法 為例,反向的規則是 和 :
matmul(a: Mat, b: Mat): Mat {
const out = /* … 前向:照常計算 a·b … */;
this.record(() => {
for (let i = 0; i < n; i++)
for (let j = 0; j < m; j++)
for (let p = 0; p < k; p++) {
a.grad[i * k + p] += out.grad[i * m + j] * b.data[p * m + j]; // dA = dC · Bᵀ
b.grad[p * m + j] += out.grad[i * m + j] * a.data[i * k + p]; // dB = Aᵀ · dC
}
});
return out;
}整個引擎只有十二種運算:矩陣乘法、加法、加偏置、縮放、轉置、ReLU、LayerNorm、因果 softmax、切開和接回注意力頭、embedding 查表,以及最後的交叉熵損失。前向跑完之後,把記下來的函式倒著執行一遍,每個參數的梯度就都算好了。這就是反向傳播。
有了梯度,剩下的就是更新參數。這裡用的是 Adam:它替每個參數各自記錄梯度的移動平均和平方的移動平均,藉此決定每個參數自己的步長。
step(batch = 16, lr = 3e-3): number {
this.model.zeroGrad();
for (let b = 0; b < batch; b++) {
const { ids, targets } = example(randomDigits(this.rng), this.task);
const tape = new Tape();
tape.crossEntropy(this.model.forward(tape, ids).logits, targets, 1 / batch);
tape.backward(); // 16 題的梯度加在一起
}
this.adam.step(lr);
}每一步都是 16 題全新的隨機題目。六位數總共有一百萬種組合,模型在一百步裡只看過其中 1,600 種,卻能答對沒看過的題目。它不是把答案背起來,而是學到了規則。
這和真正的大型語言模型差在哪裡
結構上幾乎沒有差別。差別在規模,以及規模帶來的一切:
- 參數:這裡是 1.4 萬個,現在的大模型是幾千億個。
- 層數:這裡是 1 層。上面的排序任務已經暗示了深度的用處:一層只能做「一次查找」,要組合多個步驟就需要更多層。
- 資料:這裡的資料是無限的、沒有雜訊的,而且任務有唯一正確答案。真實的文字三者皆非。
- 位置編碼:這裡用的是學出來的絕對位置向量,而且位置表只有 12 格,所以它只會處理剛好六位數,七位數連放都放不進去。
不過核心是一樣的:預測下一個 token、算出損失、把梯度傳回去、把每個參數往對的方向推一點點。你剛才在瀏覽器裡看到的那幾秒鐘,就是這件事的全部。