一字一诗
LESSON 13 / 26

多头:同时做几组汇总

这一课完成什么

把一次注意力汇总分成几组并行计算,再合并结果。每一组叫一个“头”。头数不是诗句数,也不是模型里几个独立人格;它只是把表示宽度分成若干部分,分别计算匹配和汇总。

多头,是把特征分组后分别汇总

一个位置的向量被分成几组,每组各自计算注意力,最后拼回原宽度。不同组可以学出不同的处理方式,但我们没有事先给它们分派语法或押韵任务。

多头,是把特征分组后分别汇总
用文字逐步读这张图
  1. 原来的宽度:每位置 8 个数字;C = 8
  2. 分为 2 头:每头 4 个数字;8 ÷ 2 = 4
  3. 第一组计算:自己的 Q、K、V;得到一份汇总
  4. 第二组计算:另一组 Q、K、V;可有不同权重
  5. 拼回 8 格:先拼接,再做变换;不是求平均
  6. 尺寸要兼容:8 可以分 1、2、4、8 头;不能平均分成 3 头
增加头数一定增加参数吗?

在总宽度固定、当前实现不变时,主要是重新分组,QKV 的总投影尺寸不因此增加。

用尺寸追踪一次分头

假设 B=1、T=3、C=8,头数 H=2。每个头宽度为 8÷2=4。Q 原本是 [1,3,8],先改成 [1,3,2,4],再交换位置维和头维,得到 [1,2,3,4]

环节尺寸含义
Q、K、V[1,2,3,4]两个头各处理三个位置
匹配分数[1,2,3,3]每个头有一张位置匹配表
汇总后的内容[1,2,3,4]每个头各输出四项信息
合并[1,3,8]每个位置重新拼回八项

每个头都执行第 11、12 课的比较、遮挡、归一化与汇总。合并后再经过输出投影,允许不同头的信息相互组合。

运行并查看

bash
./poet lesson 13
python
q = q.view(batch, length, heads, width // heads).transpose(1, 2)
merged = values.transpose(1, 2).contiguous().view(batch, length, width)

先不改代码,用纸列出两次 transpose 前后的每个维度。随后把头数从 2 改为 4,再运行。总宽度不变,每头宽度由 4 变成 2,权重表头维由 2 变成 4。

头数增加会发生什么

在本项目固定总宽度的实现中,只改变头数,QKV 变换矩阵的尺寸不会增加,每头分到的宽度反而变小。因此不能把“头数更多”直接等同于“模型参数更多”或“效果一定更好”。不同头可能学到不同的关联,也可能重复;需要用固定的验证方式比较。

注意力头的含义不是事先写死的“一个看押韵、一个看思乡”。模型没有收到这样的分工指令。为了说明结构可以用多组汇总来理解,但不能把比喻当作已验证的解释。

验收与排错

小练习与答案

正式配置宽度 256、头数 4、前文长度 64,单首诗每个头的分数表多大?所有头共多少格?

查看答案

每个头是 64×64,共四张表,因此有 4×64×64=16,384 个分数格;每头的 Q、K、V 宽度是 64。这里是中间计算数量,并非新增了这么多可学习参数。批数量增加还会复制这些中间计算。

动手补全一小段

先运行上面的完整实验,再复制本课起始文件为自己的练习。starter 有意留空 solve 函数;补全后运行文件,底部检查会告诉你是否符合本课要求。卡住时打开参考答案,比较每一步。

bash
cp lessons/13/starter.py lessons/13/my_exercise.py
# 编辑 my_exercise.py 中的 solve 函数,然后运行:
.venv/bin/python lessons/13/my_exercise.py
# 对照完整答案:
.venv/bin/python lessons/13/solution.py

下载起始代码 · 下载参考答案

展开本课补全练习的完整参考答案
python
"""第 13 课补全练习:把表示分成多个头。修改 solve,保持下方检查不变。"""
import math, io, argparse
import torch
from torch import nn
from torch.nn import functional as F
torch.set_num_threads(2)
torch.manual_seed(26)

def solve(x, heads):
    b,t,c=x.shape
    if c % heads: raise ValueError('宽度不能整除头数')
    return x.view(b,t,heads,c//heads).transpose(1,2)

x=torch.arange(24).view(1,3,8)
y=solve(x,2)
assert y.shape==(1,2,3,4)
assert torch.equal(y[0,1,0],x[0,0,4:])
print("本课补全练习通过。")
展开本课完整、可独立运行的实验代码
python
"""本课独立实验;在仓库根目录执行 .venv/bin/python lessons/13/experiment.py。"""
import sys, json, math
from pathlib import Path
sys.path.insert(0, str(Path(__file__).resolve().parents[2]))
import torch
from torch import nn
from torch.nn import functional as F
from poetry_gpt.common import ROOT, DEFAULT_DATA, read_json, read_jsonl
from poetry_gpt.model import ModelConfig, PoetryGPT, CausalAttention, Block
from poetry_gpt.data import Tokenizer, SPECIAL, clean_record, keywords_for
from poetry_gpt.labs import show
torch.set_num_threads(2)
torch.manual_seed(26)

config = ModelConfig(20, context=8, width=8, heads=2, layers=1, dropout=0)
attention = CausalAttention(config).eval()
out, weights = attention(torch.randn(1, 3, 8), inspect=True)
show('权重 [批,头,查询位置,被看位置]', list(weights.shape))
show('合并输出', list(out.shape))
下载本课实验