博客

Transformer 架构详解:一篇完整的梳理

2024年2月22日

English

当我刚进入 AI 领域、第一次接触 Transformer Architecture 时,我被铺天盖地的概念和教程压得有点喘不过气。我发现有些视频或文章默认你已经懂自然语言处理(NLP)的基础概念;另一些又太长、太绕,很难真正消化。为了理解 Transformer Architecture,我读了很多文章,也看了好几支视频。中间有不少时候,我会卡在一些非常基础的问题上,导致后面的内容怎么也接不上,比如:

  • 为什么 input embedding 可以表示一个词?
  • 矩阵乘法到底是什么,shape 又是怎么变化的?
  • softmax 在做什么?
  • 模型训练完成后,它到底被存在哪里? ...

这里面有太多复杂的细节,我自己也觉得要完全掌握这套架构并不容易。所以我希望你保持耐心,也建议你先看视频片段,因为很多概念一旦有人解释清楚,其实并不难。如果我们从基础开始,一步一步往前走,再慢慢进入更高级的部分,我相信你很快就能和我站在同一个地方看这件事。

高层概览

Transformer 模型最早由 2017 年的论文 《Attention is all you need》 提出。Transformer 架构最初是为了训练语言翻译模型而设计的。不过 OpenAI 团队后来发现,Transformer 架构正是字符预测的关键解法。一旦模型在整个互联网数据上完成训练,它就有可能理解任意文本的上下文,并像人一样连贯地补完整个句子。

模型由两部分组成:encoder 和 decoder。一般来说,encoder-only 架构擅长从文本中抽取信息,用于分类和回归等任务;而 decoder-only 模型更擅长生成文本。例如专注文本生成的 GPT,就属于 decoder-only 模型。

不过,GPT 模型只使用 Transformer 架构里的 decoder 部分。

我们先快速走一遍训练模型时,这套架构里的关键想法。

我画了一张图,用来说明 GPT 这类 decoder-only Transformer 架构的训练过程:

  1. 首先,我们需要一段输入字符序列作为训练数据。这些输入会被转换成 vector embedding 格式。

  2. 接着,我们会给这些 vector embeddings 加上 positional encoding,用来捕捉每个字符在序列中的位置。

  3. 然后,模型会通过一系列计算操作处理这些 input embeddings,最终为给定输入文本生成“下一个可能字符”的概率分布。

  4. 模型会把预测结果与训练数据里真实的下一个字符进行比较,并据此调整概率,也就是调整“weights”。

  5. 最后,模型会反复迭代这个过程,持续更新参数,从而提升未来预测的准确性。

下面我们拆开每一步的细节。

第 1 步:Tokenization(分词)

Tokenization 是 Transformer 模型的第一步,它做的事情是:

把输入句子翻译成一组数字表示。

Tokenization 是把文本切分成更小单位的过程,这些单位叫 tokens,可以是单词、子词、短语或字符。把句子拆成更小的片段,有助于模型识别文本背后的结构,并更高效地处理它。

例如:

Chapter 1: Building Rapport and Capturing

上面的句子可以被切成:

Chapter, ,1,:, ,Building, ,Rap,port, ,and, ,Capturing

它会被 tokenized 成 10 个数字:

[26072, 220, 16, 25, 17283, 23097, 403, 220, 323, 220, 17013, 220, 1711]

你可以看到,数字 220 被用来表示空格字符。 把字符 tokenize 成整数有很多种方法。对于我们的示例数据集,我们会使用 tiktoken library。

为了演示,我会使用一个很小的 textbook dataset (来自 Hugging Face,它包含 46 万个字符,用于我们的训练。

- File size:  450Kb
- Vocab size: 3,771 (means unique words/sub-words)

我们的训练数据包含 vocabulary size 为 3,771 的不同字符。 而用于 tokenize 这个 textbook dataset 的最大数字是 100069,它映射到字符 Clar

一旦有了 tokenized mapping,我们就能为数据集中的每个字符找到对应的整数索引。 后续与模型交互时,我们会使用这些分配好的整数索引作为 tokens,而不是直接使用完整单词。

第 2 步:Word Embeddings(词嵌入)

首先,让我们构建一张 look-up table,里面包含 vocabulary 中的所有字符。 本质上,这张表就是一个由随机初始化数字填满的矩阵。

考虑到我们最大的 token number 是 100069,并且这里选用 64(原始论文使用 512 维,记作 d_model),最终的查找表就是一个 100,069 × 64 的矩阵,这叫做 Token Embedding Look-up Table。 表示如下:

Token Embedding Look-Up Table:
               0         1         2         3         4         5         6         7         8         9   ...        54        55        56        57        58        59        60        61        62        63
0       0.625765  0.025510  0.954514  0.064349 -0.502401 -0.202555 -1.567081 -1.097956  0.235958 -0.239778  ...  0.420812  0.277596  0.778898  1.533269  1.609736 -0.403228 -0.274928  1.473840  0.068826  1.332708
1      -0.497006  0.465756 -0.257259 -1.067259  0.835319 -1.956048 -0.800265 -0.504499 -1.426664  0.905942  ...  0.008287 -0.252325 -0.657626  0.318449 -0.549586 -1.464924 -0.557690 -0.693927 -0.325247  1.243933
2       1.347121  1.690980 -0.124446 -1.682366  1.134614 -0.082384  0.289316  0.835773  0.306655 -0.747233  ...  0.543340 -0.843840 -0.687481  2.138219  0.511412  1.219090  0.097527 -0.978587 -0.432050 -1.493750
3       1.078523 -0.614952 -0.458853  0.567482  0.095883 -1.569957  0.373957 -0.142067 -1.242306 -0.961821  ... -0.882441  0.638720  1.119174 -1.907924 -0.527563  1.080655 -2.215207  0.203201 -1.115814 -1.258691
4       0.814849 -0.064297  1.423653  0.261726 -0.133177  0.211893  1.449790  3.055426 -1.783010 -0.832339  ...  0.665415  0.723436 -1.318454  0.785860 -1.150111  1.313207 -0.334949  0.149743  1.306531 -0.046524
...          ...       ...       ...       ...       ...       ...       ...       ...       ...       ...  ...       ...       ...       ...       ...       ...       ...       ...       ...       ...       ...
100064 -0.898191 -1.906910 -0.906910  1.838532  2.121814 -1.654444  0.082778  0.064536  0.345121  0.262247  ...  0.438956  0.163314  0.491996  1.721039 -0.124316  1.228242  0.368963  1.058280  0.406413 -0.326223
100065  1.354992 -1.203096 -2.184551 -1.745679 -0.005853 -0.860506  1.010784  0.355051 -1.489120 -1.936192  ...  1.354665 -1.338872 -0.263905  0.284906  0.202743 -0.487176 -0.421959  0.490739 -1.056457  2.636806
100066 -0.436116  0.450023 -1.381522  0.625508  0.415576  0.628877 -0.595811 -1.074244 -1.512645 -2.027422  ...  0.436522  0.068974  1.305852  0.005790 -0.583766 -0.797004  0.144952 -0.279772  1.522029 -0.629672
100067  0.147102  0.578953 -0.668165 -0.011443  0.236621  0.348374 -0.706088  1.368070 -1.428709 -0.620189  ...  1.130942 -0.739860 -1.546209 -1.475937 -0.145684 -1.744829  0.637790 -1.064455  1.290440 -1.110520
100068  0.415268 -0.345575  0.441546 -0.579085  1.110969 -1.303691  0.143943 -0.714082 -1.426512  1.646982  ... -2.502535  1.409418  0.159812 -0.911323  0.856282 -0.404213 -0.012741  1.333426  0.372255  0.722526

[100,069 rows x 64 columns]

其中,每一行表示一个字符(由它的 token number 索引),每一列表示一个 dimension。

现在,你可以先把“dimension”理解成一个字符的某种特征或侧面。在我们的例子里,我们指定了 64 个维度,也就是说模型会用 64 种不同方式去理解一个字符的文本含义,比如它是否更像名词、动词、形容词等等。

假设现在我们有一个 context_length16 的训练输入示例:

. By mastering the art of identifying underlying motivations and desires, we equip ourselves with

现在,我们使用每个 tokenized character 的整数索引去查 embedding table,从而取回对应的 embedding vector。 于是,我们得到它们各自的 input embeddings:

[ 627, 1383, 88861, 279, 1989, 315, 25607, 16940, 65931, 323, 32097, 11, 584, 26458, 13520, 449]

在 Transformer 架构中,多个输入序列通常会被并行处理,这经常被称为 multiple batches。 我们把 batch_size 设置为 4。也就是说,我们会一次处理四个随机选出的句子作为输入。

Input Sequence Batch:
       0     1      2      3     4      5      6      7      8      9      10     11     12     13     14     15
0    627  1383  88861    279  1989    315  25607  16940  65931    323  32097     11    584  26458  13520    449
1  15749   311   9615   3619   872   6444      6   3966     11  10742     11    323  32097     13   3296  22815
2  13189   315   1701   5557   304   6763    374  88861   7528  10758   7526     13   4314   7526   2997   2613
3    323  6376   2867  26470  1603  16661    264  49148    627     18     13  81745  48023  75311   7246  66044

[4 rows x 16 columns]

每一行表示一个句子;每一列表示这个句子中从第 0 到第 15 个位置的字符。

结果是,我们现在有了一个矩阵,用来表示 4 个 batch、每个 batch 16 个字符的输入。这个矩阵的 shape 是 (batch_size, context_length) = [4, 16]

回顾一下,我们前面定义了一个大小为 100,069 × 64 的 input embedding lookup table。 下一步,就是把 input sequences matrix 映射到这个 embedding matrix 上,得到我们的 Input Embedding

这里,我们先聚焦拆解 input sequence matrix 中的每一行,从第一行开始。 首先,把这一行从原始维度 (1, context_length) = [1, 16] reshape 成新的格式 (context_length, 1) = [16, 1]。随后,我们把这行重新组织后的数据叠到前面建立的 embedding matrix 上,embedding matrix 的大小是 (vocab_size, d_model) = [100069, 64],于是上下文窗口里的每个字符都会被替换成匹配的 embedding vector。 最终输出是一个 shape 为 (context_length, d_model) = [16, 64] 的矩阵。

input sequence batch 的第一行:

Input Embedding:
          0         1         2         3         4         5         6         7         8         9   ...        54        55        56        57        58        59        60        61        62        63
0   1.051807 -0.704369 -0.913199 -1.151564  0.582201 -0.898582  0.984299 -0.075260 -0.004821 -0.743642  ...  1.151378  0.119595  0.601200 -0.940352  0.289960  0.579749  0.428623  0.263096 -0.773865 -0.734220
1  -0.293959 -1.278850 -0.050731  0.862562  0.200148 -1.732625  0.374076 -1.128507  0.281203 -1.073113  ... -0.062417 -0.440599  0.800283  0.783043  1.602350 -0.676059 -0.246531  1.005652 -1.018667  0.604092
2  -0.292196  0.109248 -0.131576 -0.700536  0.326451 -1.885801 -0.150834  0.348330 -0.777281  0.986769  ...  0.382480  1.315575 -0.144037  1.280103  1.112829  0.438884 -0.275823 -2.226698  0.108984  0.701881
3   0.427942  0.878749 -0.176951  0.548772  0.226408 -0.070323 -1.865235  1.473364  1.032885  0.696173  ...  1.270187  1.028823 -0.872329 -0.147387 -0.083287  0.142618 -0.375903 -0.101887  0.989520 -0.062560
4  -1.064934 -0.131570  0.514266 -0.759037  0.294044  0.957125  0.976445 -1.477583 -1.376966 -1.171344  ...  0.231112  1.278687  0.254688  0.516287  0.621753  0.219179  1.345463 -0.927867  0.510172  0.656851
5   2.514588 -1.001251  0.391298 -0.845712  0.046932 -0.036732  1.396451  0.934358 -0.876228 -0.024440  ...  0.089804  0.646096 -0.206935  0.187104 -1.288239 -1.068143  0.696718 -0.373597 -0.334495 -0.462218
6   0.498423 -0.349237 -1.061968 -0.093099  1.374657 -0.512061 -1.238927 -1.342982 -1.611635  2.071445  ...  0.025505  0.638072  0.104059 -0.600942 -0.367796 -0.472189  0.843934  0.706170 -1.676522 -0.266379
7   1.684027 -0.651413 -0.768050  0.599159 -0.381595  0.928799  2.188572  1.579998 -0.122685 -1.026440  ... -0.313672  1.276962 -1.142109 -0.145139  1.207923 -0.058557 -0.352806  1.506868 -2.296642  1.378678
8  -0.041210 -0.834533 -1.243622 -0.675754 -1.776586  0.038765 -2.713090  2.423366 -1.711815  0.621387  ... -1.063758  1.525688 -1.762023  0.161098  0.026806  0.462347  0.732975  0.479750  0.942445 -1.050575
9   0.708754  1.058510  0.297560  0.210548  0.460551  1.016141  2.554897  0.254032  0.935956 -0.250423  ... -0.552835  0.084124  0.437348  0.596228  0.512168  0.289721 -0.028321 -0.932675 -0.411235  1.035754
10 -0.584553  1.395676  0.727354  0.641352  0.693481 -2.113973 -0.786199 -0.327758  1.278788 -0.156118  ...  1.204587 -0.131655 -0.595295 -0.433438 -0.863684  3.272247  0.101591  0.619058 -0.982174 -1.174125
11 -0.753828  0.098016 -0.945322  0.708373 -1.493744  0.394732  0.075629 -0.049392 -1.005564  0.356353  ...  2.452891 -0.233571  0.398788 -1.597272 -1.919085 -0.405561 -0.266644  1.237022  1.079494 -2.292414
12 -0.611864  0.006810  1.989711 -0.446170 -0.670108  0.045619 -0.092834  1.226774 -1.407549 -0.096695  ...  1.181310 -0.407162 -0.086341 -0.530628  0.042921  1.369478  0.823999 -0.312957  0.591755  0.516314
13 -0.584553  1.395676  0.727354  0.641352  0.693481 -2.113973 -0.786199 -0.327758  1.278788 -0.156118  ...  1.204587 -0.131655 -0.595295 -0.433438 -0.863684  3.272247  0.101591  0.619058 -0.982174 -1.174125
14 -1.174090  0.096075 -0.749195  0.395859 -0.622460 -1.291126  0.094431  0.680156 -0.480742  0.709318  ...  0.786663  0.237733  1.513797  0.296696  0.069533 -0.236719  1.098030 -0.442940 -0.583177  1.151497
15  0.401740 -0.529587  3.016675 -1.134723 -0.256546 -0.219896  0.637936  2.000511 -0.418684 -0.242720  ... -0.442287 -1.519394 -1.007496 -0.517480  0.307449 -0.316039 -0.880636 -1.424680 -1.901644  1.968463

[16 rows x 64 columns]

矩阵展示的是四行中的一行完成映射后的结果

剩下 3 行也做同样的事情,最后我们会得到 4 组 x [16 rows x 64 columns]。

这会产生 shape 为 (batch_size, context_length, d_model) = [4, 16, 64]Input Embedding 矩阵。

本质上,给每个词一个独特的 embedding,模型就能容纳语言中的变化,并处理那些有多个含义或形态的词。

即使我们还没有完全掌握背后的数学原理,也可以先把 input embedding matrix 理解成模型期望接收的输入格式,然后继续往前走。

第 3 步:Positional Encoding(位置编码)

在我看来,positional encoding 是 Transformer 架构里最难理解的概念。

概括一下 positional encoding 要解决的问题:

我们希望每个词都携带一些关于它在句子中位置的信息。

我们希望模型把彼此靠近的词看作“接近”,把相隔很远的词看作“遥远”。

我们希望 positional encoding 表达出一种模型可以学习的模式。

Position encoding 描述的是一个实体在序列中的位置,让每个位置都能被分配到一个独一无二的表示。

Position encoding 是另一组数字向量,会被加到每个 tokenized character 的 input embedding 上。 Position encoding 是正弦波和余弦波,它们的频率会根据 tokenized character 的位置而变化。

原始论文中,用来计算 position encoding 的方法是:


PE(pos, 2i) = sin(pos / 10000^(2i/d_model))
PE(pos, 2i+1) = cos(pos / 10000^(2i/d_model))

其中 pos 是位置,i 从 0 到 d_model/2。d_model 是训练模型时定义的模型维度 (在我们的例子中是 64,原始论文使用 512)

事实上,这个 positional encoding matrix 只需要创建一次,然后可以复用于每个 input sequence。

我们来看一下 positional encoding matrix:

Position Embedding Look-Up Table:
          0         1         2         3         4         5         6         7         8         9   ...        54        55        56        57        58        59        60        61        62        63
0   0.000000  1.000000  0.000000  1.000000  0.000000  1.000000  0.000000  1.000000  0.000000  1.000000  ...  0.000000  1.000000  0.000000  1.000000  0.000000  1.000000  0.000000  1.000000  0.000000  1.000000
1   0.841471  0.540302  0.681561  0.731761  0.533168  0.846009  0.409309  0.912396  0.310984  0.950415  ...  0.000422  1.000000  0.000316  1.000000  0.000237  1.000000  0.000178  1.000000  0.000133  1.000000
2   0.909297 -0.416147  0.997480  0.070948  0.902131  0.431463  0.746904  0.664932  0.591127  0.806578  ...  0.000843  1.000000  0.000632  1.000000  0.000474  1.000000  0.000356  1.000000  0.000267  1.000000
3   0.141120 -0.989992  0.778273 -0.627927  0.993253 -0.115966  0.953635  0.300967  0.812649  0.582754  ...  0.001265  0.999999  0.000949  1.000000  0.000711  1.000000  0.000533  1.000000  0.000400  1.000000
4  -0.756802 -0.653644  0.141539 -0.989933  0.778472 -0.627680  0.993281 -0.115730  0.953581  0.301137  ...  0.001687  0.999999  0.001265  0.999999  0.000949  1.000000  0.000711  1.000000  0.000533  1.000000
5  -0.958924  0.283662 -0.571127 -0.820862  0.323935 -0.946079  0.858896 -0.512150  0.999947 -0.010342  ...  0.002108  0.999998  0.001581  0.999999  0.001186  0.999999  0.000889  1.000000  0.000667  1.000000
6  -0.279415  0.960170 -0.977396 -0.211416 -0.230368 -0.973104  0.574026 -0.818837  0.947148 -0.320796  ...  0.002530  0.999997  0.001897  0.999998  0.001423  0.999999  0.001067  0.999999  0.000800  1.000000
7   0.656987  0.753902 -0.859313  0.511449 -0.713721 -0.700430  0.188581 -0.982058  0.800422 -0.599437  ...  0.002952  0.999996  0.002214  0.999998  0.001660  0.999999  0.001245  0.999999  0.000933  1.000000
8   0.989358 -0.145500 -0.280228  0.959933 -0.977262 -0.212036 -0.229904 -0.973213  0.574318 -0.818632  ...  0.003374  0.999994  0.002530  0.999997  0.001897  0.999998  0.001423  0.999999  0.001067  0.999999
9   0.412118 -0.911130  0.449194  0.893434 -0.939824  0.341660 -0.608108 -0.793854  0.291259 -0.956644  ...  0.003795  0.999993  0.002846  0.999996  0.002134  0.999998  0.001600  0.999999  0.001200  0.999999
10 -0.544021 -0.839072  0.937633  0.347628 -0.612937  0.790132 -0.879767 -0.475405 -0.020684 -0.999786  ...  0.004217  0.999991  0.003162  0.999995  0.002371  0.999997  0.001778  0.999998  0.001334  0.999999
11 -0.999990  0.004426  0.923052 -0.384674 -0.097276  0.995257 -0.997283 -0.073661 -0.330575 -0.943780  ...  0.004639  0.999989  0.003478  0.999994  0.002609  0.999997  0.001956  0.999998  0.001467  0.999999
12 -0.536573  0.843854  0.413275 -0.910606  0.448343  0.893862 -0.940067  0.340989 -0.607683 -0.794179  ...  0.005060  0.999987  0.003795  0.999993  0.002846  0.999996  0.002134  0.999998  0.001600  0.999999
13  0.420167  0.907447 -0.318216 -0.948018  0.855881  0.517173 -0.718144  0.695895 -0.824528 -0.565821  ...  0.005482  0.999985  0.004111  0.999992  0.003083  0.999995  0.002312  0.999997  0.001734  0.999998
14  0.990607  0.136737 -0.878990 -0.476839  0.999823 -0.018796 -0.370395  0.928874 -0.959605 -0.281349  ...  0.005904  0.999983  0.004427  0.999990  0.003320  0.999995  0.002490  0.999997  0.001867  0.999998
15  0.650288 -0.759688 -0.968206  0.250154  0.835838 -0.548975  0.042249  0.999107 -0.999519  0.031022  ...  0.006325  0.999980  0.004743  0.999989  0.003557  0.999994  0.002667  0.999996  0.002000  0.999998

[16 rows x 64 columns]

我们再稍微多聊一点 positional encoding 的技巧。

以我的理解,positional values 是围绕它们在序列中的相对位置建立的。而且,因为每个输入句子都有一致的 context length,我们就可以在不同输入之间复用同一套 positional encoding。因此,生成这些序列数字时必须小心:不能让数值过大,避免它压过 input embeddings;同时要让相邻位置之间差异较小,而距离远的位置之间差异更明显。

通过组合 sine 与 cosine vectors,模型看到的是一组独立于 word embedding 的 positional encoding vector,同时不会混淆 input embedding 中的语义信息。 很难想象这东西在 neural network 内部到底怎么工作,但它确实有效。

我们可以把 position embedding numbers 可视化,看一下其中的模式。

每条竖线是从 0 到 64 的一个 dimension;每一行表示一个字符。 这些值都在 -1 和 1 之间,因为它们来自 sine 和 cosine functions。颜色越深表示值越接近 -1,颜色越亮表示越接近 1。绿色表示中间值。

回到我们的 positional encoding matrix。你可以看到,这张 positional encoding table 与 input embedding table 中每个 batch 的 shape 相同:[4, 16, 64],它们都是 (context_length, d_model) = [16, 64]

由于两个 shape 相同的矩阵可以相加,我们就可以把 positional information 加到每一行 input embedding 上,得到 Final Input Embedding matrix。

batch 0:
          0         1         2         3         4         5         6         7         8         9   ...        54        55        56        57        58        59        60        61        62        63
0   1.051807  0.295631 -0.913199 -0.151564  0.582201  0.101418  0.984299  0.924740 -0.004821  0.256358  ...  1.151378  1.119595  0.601200  0.059648  0.289960  1.579749  0.428623  1.263096 -0.773865  0.265780
1   0.547512 -0.738548  0.630830  1.594323  0.733316 -0.886616  0.783385 -0.216111  0.592187 -0.122698  ... -0.061995  0.559401  0.800599  1.783043  1.602587  0.323941 -0.246353  2.005651 -1.018534  1.604092
2   0.617101 -0.306899  0.865904 -0.629588  1.228581 -1.454339  0.596070  1.013263 -0.186154  1.793348  ...  0.383324  2.315575 -0.143404  2.280102  1.113303  1.438884 -0.275467 -1.226698  0.109251  1.701881
3   0.569062 -0.111243  0.601322 -0.079154  1.219661 -0.186289 -0.911600  1.774332  1.845533  1.278927  ...  1.271452  2.028822 -0.871380  0.852612 -0.082575  1.142617 -0.375369  0.898113  0.989920  0.937440
4  -1.821736 -0.785214  0.655805 -1.748969  1.072516  0.329445  1.969725 -1.593312 -0.423386 -0.870206  ...  0.232799  2.278685  0.255953  1.516287  0.622701  1.219178  1.346175  0.072133  0.510705  1.656851
5   1.555663 -0.717588 -0.179829 -1.666574  0.370867 -0.982811  2.255347  0.422208  0.123719 -0.034782  ...  0.091912  1.646094 -0.205354  1.187103 -1.287054 -0.068144  0.697607  0.626403 -0.333828  0.537782
6   0.219007  0.610934 -2.039364 -0.304516  1.144289 -1.485164 -0.664902 -2.161820 -0.664487  1.750649  ...  0.028036  1.638068  0.105957  0.399056 -0.366373  0.527810  0.845001  1.706170 -1.675722  0.733621
7   2.341013  0.102489 -1.627363  1.110608 -1.095316  0.228369  2.377153  0.597940  0.677737 -1.625878  ... -0.310720  2.276958 -1.139895  0.854859  1.209583  0.941441 -0.351562  2.506867 -2.295708  2.378678
8   0.948148 -0.980033 -1.523850  0.284180 -2.753848 -0.173272 -2.942995  1.450153 -1.137498 -0.197246  ... -1.060385  2.525683 -1.759494  1.161095  0.028703  1.462346  0.734397  1.479749  0.943511 -0.050575
9   1.120872  0.147380  0.746753  1.103982 -0.479273  1.357801  1.946789 -0.539822  1.227215 -1.207067  ... -0.549040  1.084117  0.440194  1.596224  0.514303  1.289719 -0.026721  0.067324 -0.410035  2.035753
10 -1.128574  0.556604  1.664986  0.988980  0.080544 -1.323841 -1.665967 -0.803163  1.258105 -1.155904  ...  1.208804  0.868336 -0.592132  0.566557 -0.861313  4.272244  0.103369  1.619057 -0.980840 -0.174126
11 -1.753818  0.102441 -0.022270  0.323699 -1.591020  1.389990 -0.921654 -0.123053 -1.336139 -0.587427  ...  2.457530  0.766419  0.402266 -0.597278 -1.916476  0.594436 -0.264688  2.237020  1.080961 -1.292415
12 -1.148437  0.850664  2.402985 -1.356776 -0.221765  0.939481 -1.032902  1.567763 -2.015232 -0.890874  ...  1.186370  0.592825 -0.082546  0.469365  0.045767  2.369474  0.826133  0.687041  0.593355  1.516313
13 -0.164386  2.303123  0.409138 -0.306666  1.549362 -1.596800 -1.504343  0.368137  0.454260 -0.721938  ...  1.210069  0.868330 -0.591184  0.566554 -0.860601  4.272243  0.103903  1.619056 -0.980440 -0.174127
14 -0.183482  0.232812 -1.628186 -0.080981  0.377364 -1.309922 -0.275964  1.609030 -1.440347  0.427969  ...  0.792566  1.237715  1.518224  1.296686  0.072853  0.763276  1.100520  0.557057 -0.581310  2.151496
15  1.052028 -1.289275  2.048469 -0.884570  0.579293 -0.768871  0.680185  2.999618 -1.418203 -0.211697  ... -0.435962 -0.519414 -1.002752  0.482508  0.311006  0.683955 -0.877969 -0.424683 -1.899643  2.968462

[16 rows x 64 columns]

batch 1:
          0         1         2         3         4         5         6         7         8         9   ...        54        55        56        57        58        59        60        61        62        63
0  -0.264236  0.965681  1.909974 -0.338721 -0.554196  0.254583 -0.576111  1.766522 -0.652587  0.455450  ... -1.016426  0.458762 -0.513290  0.618411  0.877229  2.526591  0.614551  0.662366 -1.246907  1.128066
1   1.732205 -0.858178  0.324008  1.022650 -1.172865  0.513133 -0.121611  2.630085  0.072425  2.332296  ...  0.737660  1.988225  2.544661  1.995471  0.447863  3.174428  0.444989  0.860426  2.137797  1.537580
2  -1.348308 -1.080221  1.753394  0.156193  0.440652  1.015287 -0.790644  1.215537  2.037030  0.476560  ...  0.296941  1.100837 -0.153194  1.329375 -0.188958  1.229344 -1.301919  0.938138 -0.860689 -0.860137
3   0.601103 -0.156419  0.850114 -0.324190 -0.311584 -2.232454 -0.903112  0.242687  0.801908  2.502464  ... -0.397007  1.150545 -0.473907  0.318961 -1.970126  1.967961 -0.186831  0.131873  0.947445 -0.281573
4  -1.821736 -0.785214  0.655805 -1.748969  1.072516  0.329445  1.969725 -1.593312 -0.423386 -0.870206  ...  0.232799  2.278685  0.255953  1.516287  0.622701  1.219178  1.346175  0.072133  0.510705  1.656851
5   1.555663 -0.717588 -0.179829 -1.666574  0.370867 -0.982811  2.255347  0.422208  0.123719 -0.034782  ...  0.091912  1.646094 -0.205354  1.187103 -1.287054 -0.068144  0.697607  0.626403 -0.333828  0.537782
6   0.599841  0.943214 -1.397184 -0.607349 -0.333995 -1.222589 -0.731189 -0.997706  1.848611  0.254238  ...  0.340986  1.383113  1.674592  2.229903 -0.157415  0.362868 -0.493762  1.904136  0.027903  1.196017
7   0.072234  1.386670 -0.985962 -1.184486  0.958293 -0.295773 -1.529277 -0.727844  1.510503  1.268154  ... -0.356459  0.382331  0.138104 -0.360916 -0.638448  1.305404 -0.756442  0.299150  0.154600 -0.466154
8  -0.008645 -1.066763 -0.716555  2.148885 -0.709739 -0.137266  0.385401  0.699139  1.907906 -2.357567  ...  0.490190 -1.215412  1.216459  0.659227 -0.282908 -0.912266  0.595569  1.210701  0.737407  0.801672
9  -0.006332 -0.949928  0.192689  3.158421 -1.292153 -0.830248  0.966141 -2.056514  0.042364  1.485927  ...  0.480763 -0.318554  0.005837  3.031636 -0.448117  1.059403  0.598106  0.871427  0.327321  1.090921
10 -1.152681 -0.710162 -0.456591 -0.468090 -0.292566  0.747535 -0.149907 -0.395523  0.170872 -2.372754  ... -1.267461  0.043283 -0.114980  1.083042 -0.288776  1.442318  0.775591  0.728716 -0.576776 -0.727257
11 -0.955986 -0.277475  0.946888 -0.242687  1.257744  0.369994  0.460073  0.728078 -0.165204 -0.761762  ... -0.307983  2.078995 -1.067792  1.805637  0.608968  1.722982 -0.371174 -0.603182  0.285387  1.112932
12 -0.844347  0.883224  1.222388 -0.811387 -0.593557  0.157268 -0.650315  1.289236 -1.472027 -0.447092  ... -0.536433  2.465097 -0.822905  1.272786  0.703664  2.687270 -0.924388  0.596134 -0.367138  0.812242
13  0.776470  1.549248 -0.239693  0.133783  0.767255  1.996130 -0.436228 -0.327975 -0.650743  0.507769  ... -0.821793  1.387792 -1.052105  2.123603  1.421092  2.066746 -0.747766  0.627081 -1.749071 -0.679443
14  1.277579  0.653945  0.045632 -0.409790  0.829708  0.249433 -0.682051  0.601958 -1.932014 -2.077397  ...  0.160611  1.037856  0.656832  0.992817 -0.684056  1.031199 -0.180866  4.579140 -1.123555  0.181580
15  0.356328 -2.038538 -1.018938  1.112716  1.035987 -2.281600  0.416325 -0.129400 -0.718316 -1.042091  ... -0.056092  0.559381  0.805026  1.783032  1.605907  0.323934 -0.243863  2.005648 -1.016667  1.604090

[16 rows x 64 columns]

batch 2:
          0         1         2         3         4         5         6         7         8         9   ...        54        55        56        57        58        59        60        61        62        63
0   0.645854  1.291073 -1.588931  1.814376 -0.185270  0.846816 -1.686862  0.982995 -0.973108  1.297203  ...  0.852600  1.533231  0.692729  2.437029 -0.178137  0.493413  0.597484  1.909155  1.257821  2.644325
1   1.732205 -0.858178  0.324008  1.022650 -1.172865  0.513133 -0.121611  2.630085  0.072425  2.332296  ...  0.737660  1.988225  2.544661  1.995471  0.447863  3.174428  0.444989  0.860426  2.137797  1.537580
2   3.298391 -0.363908  0.376535 -0.276692  1.262433 -0.595659  1.694541  0.542514 -0.464756  0.368460  ... -0.169474  1.420809  0.304488  1.689731 -1.128037 -0.024476 -1.356808  2.160992 -2.110703 -0.472404
3   0.626955 -2.988524  0.915578  1.123503  0.635983  0.078006  0.466728 -0.930765  2.189286  1.505499  ...  2.496649  1.691578  0.642664  2.089205  1.926187  1.185045 -0.969952  0.666007 -0.030641  0.667574
4   0.396447 -2.116415  0.384262 -1.632779  0.859029 -0.726599  2.121946 -1.314046  0.744388 -0.227106  ... -1.937352  2.378620  0.029220  1.215336 -0.405487 -0.834419 -1.219825  0.000676 -0.821293  0.340797
5  -2.133014  0.379737 -1.320323 -0.425003 -0.298524 -2.237205  0.953327  0.168006  0.519205  0.698976  ...  0.788771  1.237731  1.515378  1.296695  0.070718  0.763281  1.098920  0.557059 -0.582510  2.151497
6  -0.390918  0.634039 -1.350461  0.032129  0.106428  0.370410  1.292387  0.986316 -0.095396  0.555067  ... -1.792372 -0.357599  0.912276  0.088746  0.866950  0.927208 -0.381643  2.532119  0.464615 -1.044299
7  -0.407947  0.622332 -0.345048 -0.247587 -0.419677  0.256695  1.165026 -2.459640 -0.576545 -1.770781  ...  0.234064  2.278682  0.256901  1.516285  0.623413  1.219177  1.346708  0.072133  0.511105  1.656851
8   3.503946 -1.146751  0.111070  0.114221 -0.930330 -0.248769  1.166547 -0.038856 -0.301910 -0.843072  ...  0.093177  1.646091 -0.204405  1.187101 -1.286342 -0.068145  0.698141  0.626402 -0.333428  0.537781
9  -1.946920 -0.443788  0.560103  3.584257 -0.134643 -1.538940 -1.059084 -0.128679  2.503847 -2.244587  ... -0.643552  1.608934 -0.488734 -0.291253  1.633294 -0.018763  0.696360 -0.657761  0.692395  1.741288
10  0.376520  0.583786 -0.705047  0.855548  0.471473  0.687240 -0.605646  0.463047  1.619052 -1.894214  ... -0.688652  1.974150 -1.399412  2.567682 -0.050040  1.782055 -0.297912  2.366196 -1.888527  0.635260
11 -0.109256 -1.394054  0.565499 -0.093785 -1.803309  0.662382 -1.528203  1.644028 -0.569133  0.438101  ...  0.741877  1.988214  2.547823  1.995465  0.450234  3.174424  0.446768  0.860424  2.139130  1.537579
12 -1.553993 -0.983421  0.392842 -1.473186  1.530387  1.894017 -0.732786 -1.601045 -0.740344  0.245303  ... -0.328828  3.013883  1.178296  1.263333  0.284824  0.791874  2.402131 -0.231270 -1.025411  0.178748
13 -0.757965  1.771306  0.805440 -0.509121  1.212250  0.388750 -0.606959  2.352489 -2.445346 -0.103223  ...  0.425556  1.783019  0.698336  1.871530  2.314023  0.424368 -1.002745  0.983784 -0.090133  0.905337
14 -0.183482  0.232812 -1.628186 -0.080981  0.377364 -1.309922 -0.275964  1.609030 -1.440347  0.427969  ...  0.792566  1.237715  1.518224  1.296686  0.072853  0.763276  1.100520  0.557057 -0.581310  2.151496
15 -0.151101 -0.257150 -0.478131 -1.170082  1.318685 -0.188166  0.146375  2.895475 -0.918949 -0.305261  ...  1.623350  1.656103 -0.600456  1.039260 -1.944202  0.894911  1.409396  1.722673 -0.172070  2.265543

[16 rows x 64 columns]

batch 3:
          0         1         2         3         4         5         6         7         8         9   ...        54        55        56        57        58        59        60        61        62        63
0   0.377847 -0.380613  1.958640  0.224087 -0.420293  0.915635 -1.077748  1.255988 -0.223147  0.977568  ... -1.290532  1.460963  1.365088 -2.037483 -2.213841  1.039091 -2.129649  0.108403 -0.356996  2.239356
1   0.527961  0.342787  0.096746  0.885016  0.706699  2.873656  0.139732  0.497379 -0.009022 -0.147825  ... -0.409913  0.785146 -0.138166  2.041000  0.277500  1.578947 -1.535113  0.912230 -0.312735  0.540365
2   1.054965 -0.134411  2.155045 -0.188724  0.651576 -0.265663 -0.777263  0.571080  1.508661  1.021718  ...  0.762458  2.297400 -0.624743 -0.979212  2.024008  1.295633  0.208825  0.953138 -2.962624  1.586901
3  -1.032970 -0.893918  0.029077 -0.232068  0.370793 -1.407092  1.048066  0.981123  0.331907  1.292072  ...  0.787928  1.237732  1.514746  1.296695  0.070244  0.763281  1.098564  0.557060 -0.582777  2.151497
4  -0.980037 -1.014605  1.875135 -2.459635  0.486067 -0.941092  1.205490  1.248531  1.801383  0.576983  ...  0.192097  1.784109 -0.201023  0.405095  0.982041  1.927637  0.008535  1.063376 -1.439787  2.967185
5  -0.369996 -1.151058 -0.126222  0.768431  0.107524 -0.481010  2.056029 -0.872815  1.522675 -0.440916  ...  0.246007 -1.032684  0.572565  0.944744  0.790383 -0.034063 -1.704374 -0.053319  1.739537  2.381506
6  -0.555136 -0.284736 -0.162689 -1.542923 -1.619371 -2.014224  0.957231 -0.338164  1.353500 -2.048436  ...  0.180549 -0.598603  0.427175  1.845072  0.924364 -0.013093 -0.054108 -0.082885 -0.719218  0.960552
7   0.548834  1.130444  1.207497  0.565839 -1.814344 -0.111523  0.480270 -1.741823  1.451116 -0.977640  ...  1.692325 -0.708754 -0.747591  1.373189 -0.224415 -0.074035 -0.323435  2.001849 -1.102584  1.644658
8   0.117209 -0.905490  0.272336  0.994848  0.648951  0.354459 -0.731171 -1.641071 -0.966286 -0.837498  ...  0.294006  1.008774  1.376944  2.969555  0.997452  2.076708  0.631358  1.080600  0.075384  1.819302
9   0.557786 -0.629395  1.606758  0.633762 -1.190379 -0.355466 -2.132275 -0.887707  1.208793 -0.741505  ...  0.765410  2.297393 -0.622529 -0.979216  2.025668  1.295631  0.210070  0.953136 -2.961691  1.586900
10  1.107697 -2.050459  1.399869  1.271179 -1.391529  1.103020 -0.910370 -0.398901 -0.803458 -2.081302  ...  1.462017 -0.115730  0.171052  0.594118  0.514388  1.593223  0.064085 -0.029184 -0.044621  1.206415
11 -1.771933  0.469475  0.961730  0.002798  1.386089  0.250342 -0.062900 -0.569053 -2.149857 -0.519952  ... -0.725692 -0.727693 -0.178683  1.675822 -0.401712  1.109331  0.980627 -0.357667 -0.484853  0.208340
12 -1.518213  1.899549 -0.320427 -0.929415 -0.701020  0.727833 -2.764498  0.612756  0.041370 -1.599998  ... -0.136314  1.068995  0.635501  0.765369  0.270007  0.319588 -0.652992  1.322658  1.724227  2.343042
13  0.094923  0.575470 -0.852224 -2.098593  0.998579  0.347285 -0.467688  0.773722 -1.664829 -0.412623  ... -1.274262  0.454381 -1.142107  1.853844 -1.912537  0.544311  0.667555 -1.187468  1.291108  2.275956
14 -0.183482  0.232812 -1.628186 -0.080981  0.377364 -1.309922 -0.275964  1.609030 -1.440347  0.427969  ...  0.792566  1.237715  1.518224  1.296686  0.072853  0.763276  1.100520  0.557057 -0.581310  2.151496
15  2.053710 -2.769740 -0.148796  0.983717 -0.038190 -0.655360  1.826909 -0.332533 -1.036128 -1.001430  ...  0.674310  0.695848 -0.181635  1.051397 -0.884897  1.590696 -1.375117  0.596254 -0.651398  0.797715

[16 rows x 64 columns]

这是将被送入 Transformer decoder blocks 进行训练的 final input embeddings。

这个最终结果矩阵叫做 Position-wise Input Embedding,它的 shape 是 (batch_size, context_length, d_model) = [4, 16, 64]

这就是 Positional Encoding 的全部。

但为什么要用 sine 和 cosine function 来编码位置?为什么不是随机数?为什么把两个数字相加之后,还能同时包含语义信息和位置信息? 一开始我也完全有同样的疑问。不过后来我发现,要训练一个模型,并不一定要完全吃透它背后的全部数学。因此,如果你想看更细的解释,可以参考单独章节,或者看我的视频片段。这里我们先讲到这里,然后进入下一步。

到目前为止,我们已经覆盖了模型中的 input encoding 和 positional encoding 部分。接下来进入 Transformer Block。

第 4 步:Transformer Block(Transformer 块)

一个 Transformer block 是三类层的堆叠:masked multi-head attention 机制、两个 normalization layers,以及一个 feed-forward network

Masked multi-head attention 是一组 self-attentions,其中每一个 self-attention 都叫做一个 head。所以我们先来看 self-attention 机制。

4.1 Multi-Head Attention 概览

Transformers 的力量来自一种叫 self-attention 的东西。有了 self-attention,模型就能把注意力集中在输入中最关键的部分。每个部分叫做一个 head。

一个 head 的工作方式是这样的: 它会让输入依次经过三个独特的层,分别叫 queries (Q)、keys (K) 和 values (V)。它首先比较 Q 和 K,调整结果,然后用这些比较结果生成一组分数,表示什么东西更重要。这些分数再被用来给 V 中的信息加权,让重要部分获得更多注意力。一个 head 的学习,就体现在训练过程中不断调整 Q、K、V 这些层里的参数。

Multi-head attention 简单说,就是把多个单独的 heads 堆在一起。 所有 heads 接收完全相同的输入,但在计算时使用各自独立的一组 weights。 处理完输入后,所有 heads 的输出会被 concatenated,然后再通过一个 linear layer。

下面这张图给出了一个 head 内部流程的可视化表示,也展示了 Multi-head Attention block 中的细节。

为了执行 attention 计算,我们先拿出原始论文 "Attention is all you need" 中的公式:

从公式来看,我们首先需要三个矩阵:Q (queries)、K (keys) 和 V (values)。为了计算 attention scores,需要执行下面几步:

  1. 用 Q 乘以 K 的转置(记作 K^T)
  2. 除以 K 的维度的平方根
  3. 应用 softmax function
  4. 乘以 V

我们一步步来看。

4.2 准备 Q,K,V

计算 attention 的第一步,是获得 Q、K、V 矩阵,它们分别表示 query、key 和 value。 这三个值会在 attention layer 中用于计算 attention probabilities(weights)。 它们是通过把前一步得到的 Position-wise Input Embedding 矩阵,也就是 X,分别输入到三个不同的 linear layers 中得到的,这三个层标记为 WqWkWv (所有值一开始都是随机分配,并且可学习)。 每个 linear layer 的输出会再被拆成若干个 heads,记作 num_heads,这里我们选择 4 个 heads。

Wq、Wk、Wv 是三个 shape 为 (d_model, d_model) = [64, 64] 的矩阵。所有值都是随机分配的。 在 neural network 中,这叫做一个 linear layertrainable parameters。Trainable parameters 就是模型会在训练过程中学习并自我更新的值。

为了得到 Q、K、V,我们会在 input embedding matrix X 与 Wq, Wk, Wv 三个矩阵之间分别做矩阵乘法 (再强调一次,它们的初始值是随机分配的)

  • Q = X*Wq
  • K = X*Wk
  • V = X*Wv

上面这个函数的计算(矩阵乘法)逻辑是:

X 的 shape 是 (batch_size, context_length, d_model) = [4, 16, 64],我们把它拆成 4 个 shape 为 [16, 64] 的子矩阵。 而 Wq, Wk, Wv 的 shape 是 (d_model, d_model) = [64, 64]。我们可以把 4 个 X 子矩阵中的每一个,分别与 Wq、Wk、Wv 做矩阵乘法。

如果你还记得线性代数,两个矩阵只有在第一个矩阵的列数等于第二个矩阵的行数时,才能相乘。在我们的例子里,X 的列数是 64,而 Wq, Wk, Wv 的行数也是 64。因此,这个乘法可以成立。

矩阵乘法的结果是 4 个 shape 为 [16, 64] 的子矩阵,组合起来可以表示为 (batch_size, context_length, d_model) = [4, 16, 64]

现在,我们有了 shape 为 (batch_size, context_length, d_model) = [4, 16, 64]Q, K, V 矩阵。接下来要把它们拆成多个 heads。这也就是 Transformer 架构把它命名为 Multi-Head Attention 的原因。

Splitting heads 的意思很简单:在 d_model 的 64 个维度里,我们把它们切成多个 heads,每个 head 包含一定数量的维度。每个 head 都可以学习输入中的某些模式或语义。

假设我们同样把 num_heads 设置为 4。这意味着我们会把 shape 为 [4, 16, 64]Q, K, V 矩阵拆成多个子矩阵。

实际切分是通过把最后一维 64 reshape 成 416 维的子维度完成的。

每个 Q, K, V 矩阵都会从 shape [4, 16, 64] 变成 [4, 16, 4, 16]。最后两个维度就是 heads。换句话说,它从:

[batch_size, context_length, d_model]

变成:

[batch_size, context_length, num_heads, head_size]

要理解 shape 相同的 Q、K、V 矩阵 [4, 16, 4, 16],可以这样看:

在这条 pipeline 中,有四个 batches。每个 batch 包含 16 个 tokens(words)。对于每个 token,有 4 个 heads,而每个 head 编码 16 维语义信息。

4.3 计算 Q,K Attention(注意力)

现在我们已经有了 Q、K、V 三个矩阵,可以开始一步步计算 single-head attention。

从 Transformer 图中可以看到,QK 矩阵会先相乘。

如果我们先从 Q 和 K 矩阵中去掉 batch_size,只保留最后三个维度,那么现在 Q = K = V = [context_length, num_heads, head_size] = [16, 4, 16]

我们还需要对前两个维度再做一次 transpose,让它们变成 Q = K = V = [num_heads, context_length, head_size] = [4 ,16, 16]。这是因为我们需要在最后两个维度上做矩阵乘法。

Q * K^T = [4, 16, 16] * [4, 16, 16] = [4, 16, 16]

为什么要这么做?这里的转置是为了方便不同 contexts 之间的矩阵乘法。 用图解释会更直接。最后两个维度,也就是 [16, 16],可以这样可视化:

这个矩阵中,每一行和每一列都表示我们示例句子上下文中的一个 token(word)。 矩阵乘法衡量的是每个词与上下文中每个其他词之间的相似度。数值越高,它们越相似。

我拿出其中一个 head 的 attention score:

[ 0.2712,  0.5608, -0.4975,  ..., -0.4172, -0.2944,  0.1899],
[-0.0456,  0.3352, -0.2611,  ...,  0.0419,  1.0149,  0.2020],
[-0.0627,  0.1498, -0.3736,  ..., -0.3537,  0.6299,  0.3374],
  ...,      ...,       ...,  ...,    ...,     ...,      ...,
  ...,      ...,       ...,  ...,    ...,     ...,      ...,
[-0.4166, -0.3364, -0.0458,  ..., -0.2498, -0.1401, -0.0726],
[ 0.4109,  1.3533, -0.9120,  ...,  0.7061, -0.0945,  0.2296],
[-0.0602,  0.2428, -0.3014,  ..., -0.0209, -0.6606, -0.3170]

[16 rows x 16 columns]

这个 16 × 16 矩阵里的数字,表示示例句子 ". By mastering the art of identifying underlying motivations and desires, we equip ourselves with" 的 attention scores。

画成图会更容易看:

横轴表示 Q 的某个 head,纵轴表示 K 的某个 head,彩色方块表示上下文中每个 token 与其他 token 之间的 similarity score。颜色越深,相似度越高。

当然,上面展示的相似度现在还没有太多意义,因为这些值只是随机分配出来的。但训练之后,similarity scores 就会变得有意义。

好,现在把 batch dimension,也就是 batch_size,重新带回 Q*K attention scores。最终结果会有 shape [batch_size, num_heads, context_length, head_size],也就是 [4, 4, 16, 16]

这就是当前步骤中的 Q*K Attention Score

4.4 Scale(缩放)

Scale 部分非常直接,我们只需要把 Q*K^T attention score 除以 K 的维度的平方根。

这里,K 的维度等于 Q 的维度,也就是 d_model 除以 num_heads:64/4 = 16。

然后我们取 16 的平方根,得到 4。再把 Q*K^T attention score 除以 4。

这样做的原因,是为了防止 Q*K^T attention score 过大。过大的值可能让 softmax function 饱和,进而导致 gradient vanish。

4.5 Mask(掩码)

在 decoder-only Transformer 模型中,masked self-attention 本质上起到了 sequence padding 的作用。

Decoder 只能看到前面的字符,不能看到未来的字符。所以未来字符会被 mask 掉,不参与 attention weights 的计算。

如果我们再把图画出来,这件事非常容易理解:

空白区域表示分数为 0,也就是被 masked out

Multi-headed attention layer 中 masking 的意义,是防止 decoder “看到未来”。 在我们的示例句子里,decoder 只能看到当前词,以及它之前出现过的所有词。

4.6 Softmax

Softmax 这一步会把一组数字变成一种特殊的列表,整个列表的总和为 1。 它会放大高数值、压低低数值,从而制造出更清晰的选择。

简而言之,softmax function 用于把 linear layer 的输出转换成 probability distribution。

在 PyTorch 这样的现代 deep learning frameworks 中,softmax function 是内置函数,使用起来非常直接:

torch.softmax(attention_score, dim=-1)

这行代码会对前一步计算出来的所有 attentions scores 应用 softmax,得到 0 到 1 之间的 probability distributions。

我们也拿出同一个 head 在应用 softmax 之后的 attention scores:

[1.0000, 0.0000, 0.0000,  ..., 0.0000, 0.0000, 0.0000],
[0.4059, 0.5941, 0.0000,  ..., 0.0000, 0.0000, 0.0000],
[0.3368, 0.4165, 0.2468,  ..., 0.0000, 0.0000, 0.0000],
...,
[0.0463, 0.0501, 0.0670,  ..., 0.0547, 0.0000, 0.0000],
[0.0769, 0.1974, 0.0205,  ..., 0.1034, 0.0464, 0.0000],
[0.0684, 0.0926, 0.0537,  ..., 0.0711, 0.0375, 0.0529]

所有 probability scores 现在都是正数,并且加起来等于 1。

4.7 计算 V Attention(注意力)

最后一步,是把 softmax output 与 V 矩阵相乘。

记住,我们的 V 矩阵也已经被拆成了多个 heads,shape 为 (batch_size, num_heads, context_length, head_size) = [4, 4, 16, 16]

而上一步 softmax 的输出 shape 也是 (batch_size, num_heads, context_length, head_size) = [4, 4, 16, 16]

这里,我们会在两个矩阵的最后两个维度上再做一次矩阵乘法。

softmax_output * V = [4, 4, 16, 16] * [4, 4, 16, 16] = [4, 4, 16, 16]

结果的 shape 是 [batch_size, num_heads, context_length, head_size] = [4, 4, 16, 16]

我们把这个结果称为 A

4.8 拼接与输出

Multi-head attention 的最后一步,是把所有 heads concatenated 在一起,并让它们通过一个 linear layer。

Concatenation 的想法是把所有 heads 里的信息合并到一起。因此,我们需要把 A 矩阵从 [batch_size, num_heads, context_length, head_size] = [4, 4, 16, 16] reshape 成 [batch_size, context_length, num_heads, head_size] = [4, 16, 4, 16]。 原因是我们需要把 num_headshead_size 放回最后两个维度,这样就可以轻松把它们组合 (通过矩阵乘法)d_model = 64 的大小。

这可以通过 PyTorch 的内置函数轻松完成:

A = A.transpose(1, 2) # [4, 16, 4, 16] [batch_size, context_length, num_heads, head_size]

接下来,我们需要把最后两个维度 [num_heads, head_size] = [4, 16] 合并成 [d_model] = [64]

A = A.reshape(batch_size, -1, d_model) # [4, 16, 64] [batch_size, context_length, d_model]

如你所见,经过一系列计算之后,我们的结果矩阵 A 又回到了与 input embedding matrix X 相同的 shape,也就是 [batch_size, context_length, d_model] = [4, 16, 64]。 因为这个输出结果接下来会作为 input 传给下一层,所以保持输入与输出 shape 一致是必要的。

但在把它传给下一层之前,我们还需要对它执行一次 linear transformation。这是通过让 concatenated matrix AWo 再做一次矩阵乘法完成的。

这个 Wo 是随机分配的,shape 为 [d_model, d_model],并且会在训练过程中被更新。

Output = A * Wo = [4, 16, 64] * [64, 64] = [4, 16, 64]

这个 linear layer 的输出就是 single-head attention 的输出,记作 output

恭喜!现在我们已经完成了 masked multi-head attention 部分! 接下来开始 Transformer block 的剩余部分。 这些部分相对直接,所以我会快速走一遍。

第 5 步:Residual Connection(残差连接)与 Layer Normalization(层归一化)

Residual connection,有时也叫 skip connection,是一种让原始输入 X 绕过一层或多层的连接。

它本质上就是把原始输入 X 与 multi-head attention layer 的 output 相加。因为它们 shape 相同,所以相加非常直接。

output = output + X

Residual connection 之后,流程进入 layer normalization。LayerNorm 是一种对网络中每一层输出做 normalization 的技术。 它通过减去该层输出的 mean,再除以 standard deviation 来完成。 这项技术用于防止一层的输出变得过大或过小,因为那会让网络变得不稳定。

在 PyTorch 中,这同样只是使用 nn.LayerNorm function 的一行代码。我们会在 Let's Code LLM 章节中看到它的实际用法。

在 "Attention is All You Need" 原始论文的图中,residual connection 和 layer normalization 被标记为 Add & Norm

第 6 步:Feed-Forward Network(前馈网络)

一旦我们得到了 normalized attention weights(probability scores),它们会被送入一个 position-wise feedforward network。

Feed-forward network (FFN) 由两个 linear layers 组成,中间夹着一个 ReLU activation function。我们来看 Python 代码如何实现:

# Define Feed Forward Network
output = nn.Linear(d_model, d_model * 4)(output)
output = nn.ReLU()(output)
output = nn.Linear(d_model * 4, d_model)(output)

让 ChatGPT 来解释上面的代码:

  1. output = nn.Linear(d_model, d_model * 4)(output): 这会对输入数据应用一个线性变换,也就是 y = xA^T + b。输入和输出大小分别是 d_model 与 d_model * 4。这个变换会增加输入数据的维度。
  2. output = nn.ReLU()(output): 这会逐元素应用 Rectified Linear Unit (ReLU) function。它作为 activation function 引入非线性,使模型能够学习更复杂的模式。
  3. output = nn.Linear(d_model * 4, d_model)(output): 这会应用另一个线性变换,把维度压回 d_model。这种“先扩张再收缩”的结构,是 neural networks 中很常见的模式。

作为刚开始学习 machine learning 或 LLM 的人,你和我一样,很可能会被这些解释绕晕。我第一次遇到这些术语时,也完全是同样的感觉。

但不用担心,我们可以这样理解:这个 feed-forward network 只是一个标准 neural network module,它的输入和输出都是 attention scores。它的目的,是把 attention scores 的维度从 64 扩展到 256,让信息更细粒度,也让模型能够学习更复杂的知识结构。然后,它再把维度压回 64,使其适合后续计算。

第 7 步:重复第 4 步到第 6 步

很好!我们已经完成了第一个 Transformer block。现在,需要对剩下想要拥有的 Transformer blocks 数量,重复同样的流程。

关于 heads,我引用 HuggingChat 的 AI 回答:

GPT-2 在最大配置(GPT-2-XL)中使用 48 个 transformer blocks;较小配置则使用更少的 transformer blocks(GPT-2-Large 为 36 个,GPT-2-Medium 为 24 个,GPT-2-Small 为 12 个)。每个 transformer block 都包含一个 multi-head self-attention mechanism,后面接 position-wise feed-forward networks。这些 transformer blocks 帮助模型捕捉长距离依赖,并生成连贯文本。

有了多个 blocks 之后,output 会被训练,并作为输入 X 传入下一个 block。经过多次迭代后,模型就能学习输入序列中词与词之间更复杂的模式和关系。

第 8 步:输出概率

在 inference 阶段,你想从模型中拿到下一个预测 token,但我们目前得到的其实是 vocabulary 中所有 tokens 的 probability distribution。 你还记得上面例子中的 vocabulary size 是 3,771 吗? 所以,为了选择概率最高的 token 之一,我们会构造一个矩阵,它的大小是 model dimensions d_model = 64 乘以 vocab_size = 3,771。 这一步在 training 和 inference 中没有区别。

# Apply the final linear layer to get the logits
logits = nn.Linear(d_model, vocab_size)(output)

我们把这个 linear layer 之后的输出称为 logits。Logits 是一个 shape 为 [batch_size, context_length, vocab_size] = [4, 16, 3771] 的矩阵。

然后,最后一个 softmax function 会被用来把 linear layer 的 logits 转换为 probability distribution。

logits = torch.softmax(logits, dim=-1)

注意:在 training 阶段,我们不需要在这里应用 softmax function,而是使用 nn.CrossEntropy function,因为它内置了 softmax 行为。

我们该如何理解 shape 为 [4, 16, 3771]logits?实际上,经过所有计算之后,它的想法非常简单:

我们有 4 条 batch pipelines,每条 pipeline 都包含该 input sequence 中的全部 16 个词,而每个词都会被映射成它相对于 vocabulary 中每个其他词的概率。

如果模型处于 training 阶段,我们会更新这些 probability parameters;如果模型处于 inference 阶段,我们就直接选择概率最高的那个。于是,一切就讲通了。

结语

通常来说,第一次接触 Transformer 架构时,很难一次性掌握它的全部复杂性。就我个人而言,我大约花了一个月,才彻底理解系统里的每一个组件。因此,我建议你继续阅读 reference page 中列出的其他资源,也从其他优秀的 Trailblazers 那里获得灵感。

个人来说,我觉得视频作为教程更有效。为了帮助学习,我已经准备了几支视频讲解,覆盖这套架构相关的概念和 hands-on coding experience。它们很快就会发布。

当你对 Transformer 架构的理解足够有信心之后,下一步就很适合通过一系列 guided steps 去实现代码。令人兴奋的进展就在前面。