当我刚进入 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 架构的训练过程:

-
首先,我们需要一段输入字符序列作为训练数据。这些输入会被转换成 vector embedding 格式。
-
接着,我们会给这些 vector embeddings 加上 positional encoding,用来捕捉每个字符在序列中的位置。
-
然后,模型会通过一系列计算操作处理这些 input embeddings,最终为给定输入文本生成“下一个可能字符”的概率分布。
-
模型会把预测结果与训练数据里真实的下一个字符进行比较,并据此调整概率,也就是调整“weights”。
-
最后,模型会反复迭代这个过程,持续更新参数,从而提升未来预测的准确性。
下面我们拆开每一步的细节。
第 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_length 为 16 的训练输入示例:
. 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,需要执行下面几步:
- 用 Q 乘以 K 的转置(记作 K^T)
- 除以 K 的维度的平方根
- 应用 softmax function
- 乘以 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 中得到的,这三个层标记为 Wq、Wk 和 Wv (所有值一开始都是随机分配,并且可学习)。 每个 linear layer 的输出会再被拆成若干个 heads,记作 num_heads,这里我们选择 4 个 heads。
Wq、Wk、Wv 是三个 shape 为 (d_model, d_model) = [64, 64] 的矩阵。所有值都是随机分配的。 在 neural network 中,这叫做一个 linear layer 或 trainable 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 成 4 个 16 维的子维度完成的。
每个 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 图中可以看到,Q 和 K 矩阵会先相乘。
如果我们先从 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_heads 和 head_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 A 与 Wo 再做一次矩阵乘法完成的。
这个 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 来解释上面的代码:
- output = nn.Linear(d_model, d_model * 4)(output): 这会对输入数据应用一个线性变换,也就是 y = xA^T + b。输入和输出大小分别是 d_model 与 d_model * 4。这个变换会增加输入数据的维度。
- output = nn.ReLU()(output): 这会逐元素应用 Rectified Linear Unit (ReLU) function。它作为 activation function 引入非线性,使模型能够学习更复杂的模式。
- 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 去实现代码。令人兴奋的进展就在前面。