手撕 多头注意力机制
·
Happy-LLM:从零开始的大语言模型原理与实践教程.pdf P20

代码(含最详细注释)
import torch.nn as nn
import torch
'''多头自注意力计算模块'''
class MultiHeadAttention(nn.Module):
def __init__(self,args:ModelArgs,is_casual=False):
#def和 __init__ 之间要间隔一个空格
#__init__ 是 Python 类的构造方法(初始化方法)
#self 必选参数,表示类的实例本身(即当前对象)
#args: ModelArgs 一个自定义类,用于封装模型的配置参数
#is_causal 控制模型是否启用因果机制(Causal Mechanism)
super().__init__()
#用于调用父类的构造方法(__init__ 方法)
#在其下面添加子类特有的属性
assert args.n_embd % args.n_head==0
#断言语句,用于确保 args.n_embd 能被 args.n_head 整除
#如果不满足,程序会抛出 AssertionError ,强制终止运行
#隐藏层维度必须是头数的整数倍,因为后⾯我们会将输⼊拆成头数个矩阵
model_parallel_size = 1
# 模型并⾏处理⼤⼩,默认为1
self.n_local_heads = args.n_heads // model_parallel_size
#Python中的 // 是整除运算符,返回两个数相除的整数部分(向下取整)
#本地计算头数,等于总头数除以模型并⾏处理⼤⼩
self.head_dim = args.dim // args.n_heads
#计算每个注意力头的向量维度
#将总维度 args.dim 平均分到每个注意力头上,得到每个头的向量维度
self.wq=nn.Linear(args.dim,args.n_heads*self.head_dim,bias=False)
self.wk=nn.Linear(args.dim,args.n_heads*self.head_dim,bias=False)
self.wv=nn.Linear(args.dim,args.n_heads*self.head_dim,bias=False)
#输入维度:args.dim(模型隐藏层的总维度)
#输出维度:args.n_heads * self.head_dim
#通过三个组合矩阵来代替了n个参数矩阵的组合
#这三个线性层将输入向量映射到Query、Key和值Value的矩阵
#self.wq:将输入向量映射为 查询矩阵 Q
#self.wk:将输入向量映射为 键矩阵 K
#self.wv:将输入向量映射为值矩阵 V
#多头注意力通过矩阵拼接实现了并行计算,避免了逐个头的串行处理
self.wo=nn.Linear(args.n_heads*self.head_dim,args.dim,bias=False)
#输出权重矩阵,维度为 n_embd x n_embd(head_dim = n_embeds / n_heads)
#self.wo的权重矩阵是可学习的,允许模型通过训练决定如何加权各个头的信息
#self.wo将多头注意力的输出映射回原始的维度(args.dim)
self.attn_dropout=nn.Dropout(args.dropout)
self.resid_dropout=nn.Dropout(args.dropout)
#Dropout 层,其作用是在训练过程中随机丢弃(置零)注意力权重的某些部分
#1.防止过拟合 2.增强泛化能力
if is_casual:
mask=torch.full((1,1,args.max_seq_len,args.max_seq_len),float("-inf"))
mask=torch.triu(mask,diagonal=1)
#torch.triu用于提取矩阵的上三角部分,并将下三角部分的值设为0
#diagonal=1 表示只保留主对角线上方1行的元素,主对角线及其下方的元素被设为 0
self.register_buffer("mask",mask)
#将张量(如 mask)绑定到模型中,使其成为模型状态的一部分
#被自动保存/加载,设备迁移时自动同步,适用于固定张量
#创建一个因果掩码(Causal Mask),其目的是遮蔽未来信息
#通过将未来位置的注意力权重设置为负无穷,使得这些位置在softmax后变为0
def forward(self,q:torch.Tensor,k:torch.Tenser,v:torch.Tenser):
bsz,seqlen,_=q.shape
#解析输入张量的维度
xq, xk, xv = self.wq(q), self.wk(k), self.wv(v)
#通过线性层生成 Query、Key、Value 向量,为后续注意力计算做准备
xq=xq.view(bsz,seqlen,self.n_local_heads,self.head_dim)
xk=xk.view(bsz,seqlen,self.n_local_heads,self.head_dim)
xv=xv.view(bsz,seqlen,self.n_local_heads,self.head_dim)
#xq使用 view 后,形状变为 (bsz, seqlen, n_local_heads, head_dim)
#xq.view将张量的原始形状重新排列为指定的新形状,但不改变张量的数据
xq=xq.transpose(1,2)
xk=xk.transpose(1,2)
xv=xv.transpose(1,2)
#调整维度1和2的顺序,得到形状 (bsz, n_local_heads, seqlen, head_dim)
#1.每个头(Head)的 Key 矩阵:形状为 (seqlen, head_dim)
#2.每个样本(Batch)的 Key 矩阵数量:
# 每个样本有 self.n_local_heads 个头,对应 self.n_local_heads 个 Key 矩阵
#3.整个批次(Batch)的 Key 矩阵数量:
# 总共有 bsz × self.n_local_heads 个 Key 矩阵(形状为 seqlen × head_dim)
#注意力分数的计算
scores=torch.matmul(xq,xk.transpose(2,3))/math.sqrt(self.head_dim)
#xk.transpose(2,3)即为转置 sqrt(square root):平方根
if self.is_casual:
assert hasattr(self,'mask')
#assert 断言检查:确保模型已经定义了 self.mask 这个属性
#如果self.mask不存在,程序会抛出AssertionError,提示开发者必须提供因果掩码
#hasattr(object, name)检查对象是否有属性,如果有,返回 True;否则返回 False
scores=scores+self.mask[:,:,:seqlen,:seqlen]
#从self.mask中截取前 seqlen 行和列,得到形状为 (1,1,seqlen,seqlen) 的掩码
#因为self.mask 是为最大序列长度 max_seq_len 预定义的
#确保只对有效位置(0 到 seqlen-1)应用掩码,避免干扰
#将掩码与注意力分数相加,使得未来位置的注意力分数变为 −∞
scores=F.softmax(scores.float(),dim=-1).type_as(xq)
#F.softmax是一个函数,可以直接调用,而nn.Softmax是一个类实例
#scores.float()将 scores 转换为float32类型,低精度可能导致溢出或精度不足
#dim=-1 指定在最后一个维度(通常是序列长度或头数)上进行 softmax
#这里 dim=-1 对每个位置的注意力权重进行归一化
#type_as(xq)将 softmax 后的 scores 的数据类型(dtype)转换为与 xq 相同
scores=self.attn_dropout(scores)
output=matmul(scores,xv)
output=output.transpose(1,2).contiguous().view(bsz,seqlen,-1)
#transpose(1, 2)交换张量的第1和第2个维度
#把(bsz, n_local_heads, seqlen, head_dim)交换后方便合并多个头
#调用transpose后,张量的is_contiguous()返回 False,因为数据在内存中“非连续”
#view操作要求张量是连续的,因此需要显式调用 contiguous() 创建一个新的连续张量
#view(bsz, seqlen, -1)将张量重塑为新的形状
output=self.wo(output)
#self.wo对多头注意力的输出进行线性变换(全连接)
#将多头注意力的拼接输出恢复为原始维度(与输入特征维度一致)
output=self.resid_dropout(output)
#标准 Transformer 结构中,每个子层(如自注意力层)的输出
# 会先经过 Dropout,再与输入相加(残差连接)
return output
易错点
Tensor
matmul前面一定要有torch.
nn.Dropout D大写
args
args(也就是ModelArgs类的实例)包含以下属性:
n_embd:嵌入维度,要能被n_heads整除
n_heads:注意力头的数量
dim:模型的维度
dropout:dropout 概率
max_seq_len:最大序列长度,用于生成因果掩码
MHA为什么work?
- 单头注意力只能关注输入序列中的一种特征模式(例如,仅关注局部依赖或仅关注长距离依赖),容易陷入局部最优解。MHA通过 多个独立的注意力头,MHA 允许模型同时关注输入序列中的 不同子空间,使得模型能够从 不同角度 解析输入,减少对单一模式的依赖。
- MHA 的每个注意力头是 独立计算 的,因此可以充分利用 GPU/TPU 的并行计算能力。相比串行计算,这种设计显著加快了训练和推理速度。
- 传统 RNN/LSTM 模型在处理长序列时容易遗忘早期信息,而单头注意力虽然能捕捉全局依赖,但可能无法覆盖所有可能的关联模式。通过多个注意力头,MHA 能够同时关注 多个不同位置 的依赖关系。
参考文章
更多推荐
所有评论(0)