跳到主要內容

004互動

在瀏覽器裡從零訓練一個 Transformer:看注意力自己長出來

不用 PyTorch,也不用任何函式庫。用 TypeScript 寫一個會自動微分的小引擎、一個一萬多個參數的 Transformer,然後按下按鈕,看它在幾秒內學會把一串數字反過來。

發布
閱讀時間
7 分鐘

CNN 那篇的權重是先在 PyTorch 裡訓練好、再搬進瀏覽器的。這一篇不一樣:訓練本身就發生在你的瀏覽器裡。下面這個模型現在的權重完全是隨機的,它什麼都不會。

我們要教它玩一個很簡單的遊戲:給它六個數字,請它倒過來寫一遍。看到 3 1 4 1 5 9,就要回答 9 5 1 4 1 3

對人來說這不用學,但模型一開始連「倒過來」是什麼意思都不知道。沒有人會告訴它規則,它只會一次又一次看到題目和正確答案,然後自己想辦法。

下面的儀器由上到下是三個步驟:先看清楚任務,按「開始訓練」,看最上面「它現在寫出來的」那一行從粉紅色(錯)一位一位變成青色(對),最後往下看注意力圖,那裡畫的是模型寫每一位數字時正在看輸入的哪裡。

fig 01/transformer / train
任務
13,728 個參數

讀六個數字,然後把它們倒過來寫。

模型讀到
314159
正確答案
951413
它現在寫出來的
速度
訓練步數
0
損失
整串答對率
%

還沒開始訓練。現在的權重是隨機的,所以上面它寫出來的答案是亂猜的。

注意力圖:它寫每一位的時候,在看輸入的哪裡

注意力頭 1

上半部是輸入,不計分

正在寫的位置

注意力頭 2

上半部是輸入,不計分

正在寫的位置

還沒訓練的時候,亮的位置是隨機的,沒有意義。按「開始訓練」,看它怎麼變。

選「慢」可以看清楚變化的過程,選「快」大約一兩秒就訓練完成。按「換一題」可以換一組數字:訓練用的題目每次都是隨機產生的,所以這一題模型幾乎不可能看過。

如果一切順利,你會在一百步之內看到三件事同時發生:損失掉到接近零、它寫出來的答案整串變成青色,以及其中一張注意力圖的下半部長出一條從右上到左下的斜線

那條斜線不是我畫上去的,也沒有任何一行程式碼叫它長成那樣。它是訓練的結果。這篇文章要講的就是:它是怎麼長出來的。

為什麼是六個數字,而不是一個會說話的模型

Transformer 的教學通常拿莎士比亞當訓練資料,訓練個幾十分鐘,得到一段看起來像英文的亂碼。這有兩個問題:在瀏覽器裡太慢,而且你很難判斷模型到底學會了什麼。

所以這裡換一個玩具任務:輸入六個數字,輸出它們的某種重新排列。它的好處是:

  • 對錯一目了然。答案不是「看起來像不像」,而是每一位數都對或不對。
  • 幾秒鐘就練得起來。模型只需要一萬多個參數。
  • 注意力圖可以直接讀。要寫出答案的第 1 位,模型必須去讀輸入的第 6 位;這件事會原封不動地出現在注意力矩陣上。

至於為什麼剛好是六個:長度固定,模型才能做得夠小,也才能在一張 12×12 的圖上把它的注意力整個畫出來。模型內部有一張「位置表」,這裡只做了 12 格(六個輸入、一個箭頭、五個已經寫出來的答案),所以它只認得這個長度。這是為了看得清楚而做的取捨,不是 Transformer 本身的限制。

模型本身沒有任何偷工減料:它是一個標準的 decoder-only Transformer,和 GPT 同一種結構,只是很小。

這篇的模型
詞彙表11 個 token:數字 0–9,加上一個「→」
序列長度12
寬度 d32
注意力頭2
層數1
參數總數13,728

把問題變成「預測下一個 token」

語言模型只會做一件事:看著前面的 token,猜下一個。所以我們把一題寫成一條序列:

3 1 4 1 5 9 → 9 5 1 4 1 3

然後要求模型在每個位置預測下一個 token。不過前半段是隨機的數字,沒有人猜得中,所以只有後半段(答案)會被計分:

content/posts/transformer-from-scratch/components/task.ts
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:如果你選了我,我給你這些資訊。

位置 ii 拿自己的 query 去和每個位置 jj 的 key 做內積,得到一個分數;分數經過 softmax 變成權重,再用這些權重把 value 加權平均:

Attention(Q,K,V)=softmax ⁣(QKdk+M)V\mathrm{Attention}(Q, K, V) = \mathrm{softmax}\!\left(\frac{QK^{\top}}{\sqrt{d_k}} + M\right) V

MM因果遮罩j>ij > i 的位置填上 -\infty,softmax 之後權重就是零。模型在寫答案的第 2 位時,不能偷看第 3 位。上面注意力圖的右上三角永遠是暗的,就是這個原因。

儀器裡畫的那兩張圖,就是 softmax()\mathrm{softmax}(\cdot) 的結果本身:第 ii 列第 jj 行有多亮,代表位置 ii 從位置 jj 讀了多少。

lib/ml/transformer.ts
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 位時,需要倒數第二位。也就是說:

6+i6 + i 個位置,應該去讀第 5i5 - i 個位置。

這是一條純粹由位置決定的規則,和數字是多少無關。模型一開始完全不知道這件事。但每一次答錯,梯度都會把 query 和 key 往「讓正確的那一格分數變高」的方向推一點。推個幾十次之後,那條斜線就出現了。

「排序」就沒有這麼乾淨的圖了。排序需要的不是「去讀第幾個位置」,而是「去找還沒用過的最小數字」,這取決於內容而不是位置。它大約要 600 到 800 步才會練到 100%,注意力也散得多。一層、兩個頭的模型可以學會它,但學到的解法不是人類一眼看得懂的那種。

模型的其餘部分

注意力負責在位置之間搬運資訊,其他零件都是逐位置運作的:

  1. Embedding:每個 token 查表得到一個 32 維向量,再加上一個代表「我在第幾個位置」的向量。沒有位置向量,模型根本分不出第 1 位和第 6 位,反轉任務就不可能學會。
  2. LayerNorm:把每個位置的向量調整成平均 0、變異數 1,讓訓練穩定。
  3. MLP:兩層全連接,中間放大成 4 倍寬再縮回來。注意力把資訊搬過來,MLP 負責處理它。
  4. 殘差連接:每個子層的輸出是「加回去」而不是「取代」。
  5. Unembedding:最後把 32 維向量投影回 11 個 token 的分數。
lib/ml/transformer.ts
// 自注意力:每個位置從前面的位置收集資訊
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 個數字。

手算當然不可能,所以我寫了一個很小的自動微分引擎。想法是:每做一個運算,就順手記下「如果有人告訴我輸出的梯度,我該怎麼把它傳回輸入」。以矩陣乘法 C=ABC = AB 為例,反向的規則是 A=CB\partial A = \partial C \, B^{\top}B=AC\partial B = A^{\top} \partial C

lib/ml/autograd.ts
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:它替每個參數各自記錄梯度的移動平均和平方的移動平均,藉此決定每個參數自己的步長。

content/posts/transformer-from-scratch/components/task.ts
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、算出損失、把梯度傳回去、把每個參數往對的方向推一點點。你剛才在瀏覽器裡看到的那幾秒鐘,就是這件事的全部。