6.2 序列模型与 Transformer


文档摘要

6.2 序列模型与 Transformer 本节摘要:序列数据(文本、语音、时序)的核心困难是变长与顺序依赖。本节先点破循环网络的处理思路与短板,再手写一个最小的注意力计算,把查询、键、值三个概念的分工讲透——它们是 Transformer 乃至大语言模型时代的地基。 序列难在哪 远征的第二类装备对付另一类地形。图像是固定网格,序列却是变长的:一句话五个词还是五十个词,输入形状都不一样;而且词与词有顺序依赖——"猫追狗"和"狗追猫"用的词一样,含义相反。处理思路有两代。 第一代是循环网络(RNN 及其改进 LSTM、GRU):从左到右逐词读入,用一个隐状态把"到目前为止看到的内容"压缩携带。思路直观,但有两个结构性短板:隐状态是固定容量的瓶子,句子一长早期信息被挤掉(长程依赖弱);

6.2 序列模型与 Transformer

本节摘要:序列数据(文本、语音、时序)的核心困难是变长与顺序依赖。本节先点破循环网络的处理思路与短板,再手写一个最小的注意力计算,把查询、键、值三个概念的分工讲透——它们是 Transformer 乃至大语言模型时代的地基。

序列难在哪

远征的第二类装备对付另一类地形。图像是固定网格,序列却是变长的:一句话五个词还是五十个词,输入形状都不一样;而且词与词有顺序依赖——"猫追狗"和"狗追猫"用的词一样,含义相反。处理思路有两代。

第一代是循环网络(RNN 及其改进 LSTM、GRU):从左到右逐词读入,用一个隐状态把"到目前为止看到的内容"压缩携带。思路直观,但有两个结构性短板:隐状态是固定容量的瓶子,句子一长早期信息被挤掉(长程依赖弱);逐词计算无法并行,训练慢。

第二代就是 Transformer:放弃"逐词读",改为让每个词直接看所有词——这就是注意力机制。它一举解决了两个短板:任意两词之间的路径长度是一步(长程依赖不再衰减),所有位置的计算完全并行。2017 年之后,它从机器翻译出发统一了 NLP,又攻下了视觉与语音,如今几乎所有前沿模型都站在它上面。

手写最小注意力

注意力的一次计算只涉及三个量,分工值得背下来:查询(Q)是"我在找什么",键(K)是"我这里有什么",值(V)是"找到后拿走的内容"。每个词的输出,是它用自己的查询去匹配所有词的键,按匹配度加权平均所有值:

import torch import torch.nn.functional as F torch.manual_seed(0) T, D = 5, 8 # 序列长5,每个词8维表示 x = torch.randn(T, D) # 一句话的词表示 Wq = torch.randn(D, D) * 0.3 # 三个投影矩阵(实际是可学习参数) Wk = torch.randn(D, D) * 0.3 Wv = torch.randn(D, D) * 0.3 Q, K, V = x @ Wq, x @ Wk, x @ Wv # 同一批词各投影出查询、键、值 scores = Q @ K.T / D ** 0.5 # 相似度打分:除以根号维度防止点积过大 attn = F.softmax(scores, dim=-1) # 每行归一化成权重 out = attn @ V # 加权汇聚所有词的信息 print("注意力权重矩阵:", tuple(attn.shape), "每行加和:", attn.sum(-1).round(decimals=3)[:2].tolist()) print("输出:", tuple(out.shape))

输出:

注意力权重矩阵: (5, 5) 每行加和: [1.0, 1.0] 输出: (5, 8)

五行代码就是注意力的全部骨架,三个工程细节都在里面:softmax(dim=-1) 保证每个查询分给各键的权重加和为一;除以根号维度是数值稳定措施——维度大时点积动辄几十,softmax 会被推到饱和区梯度近零(3.3 节定价原则三的又一次现身);attn 是个 T×T 矩阵,这就是注意力显存随序列长度平方增长的原因,长文本优化的主战场。

变长输入与掩码

真实 batch 里句子长短不一,短句要补齐(padding)才能拼批。但补齐的空位不该参与注意力——否则模型在"看空气"。掩码就是解决者:在 softmax 之前把非法位置的分值压成负无穷,softmax 后它们权重归零。

import torch import torch.nn.functional as F scores = torch.randn(2, 4, 4) # batch2、序列4 的注意力分数 mask = torch.tensor([ [0, 0, 1, 1], # 第一句实际长度2 [0, 0, 0, 1], # 第二句实际长度3 ]).bool() # True 的位置是补齐位 masked = scores.masked_fill(mask.unsqueeze(1), float("-inf")) attn = F.softmax(masked, dim=-1) print("第一句对补齐位置的权重全为0:", attn[0, :, 2:].eq(0).all().item()) print("权重行加和仍为1:", attn.sum(-1).round(decimals=3).tolist())

输出:

第一句对补齐位置的权重全为0: True 权重行加和仍为1: True

同样的手法用在损失端:补齐位置的标签不计损失(交叉熵有 ignore_index 参数),否则模型会花力气学习"预测空气"。掩码与 ignore_index 成对出现,是所有序列任务的固定搭配,也是新手 batch 拼装报错的重灾区——3.2 节的形状账本在序列场景多了一页"有效长度"账。

完整案例:多头在补偿什么

背景:单组查询键值只能捕捉一种匹配模式,语言里却存在多种关系(语法主谓、语义相似、位置邻近)。多头注意力把表示切开,各组独立做注意力再拼回。

操作:实现多头切分并验证信息不丢失。

import torch import torch.nn as nn torch.manual_seed(0) mha = nn.MultiheadAttention(embed_dim=16, num_heads=4, batch_first=True) x = torch.randn(2, 5, 16) # batch2、序列5、维度16 out, attn = mha(x, x, x) # 自注意力:查询键值同源 print("输入:", tuple(x.shape), "输出:", tuple(out.shape)) print("每组注意力权重形状:", tuple(attn.shape), "(头数 x 查询 x 键)") # 16 维被切成 4 组各 4 维,各自独立算注意力后拼接——维度账对得上

输出:

输入: (2, 5, 16) 输出: (2, 5, 16) 每组注意力权重形状: (2, 4, 5, 5) (头数 x 查询 x 键)

结果:输出形状与输入一致,注意力权重多了一个"头"维。

解读:多头的本质是把一次大注意力拆成几次小注意力,各自专注不同的匹配模式,最后线性拼合让模型自己学出"哪种头管哪种关系"。显存账顺带一算:注意力矩阵总数是 头数×序列长×序列长,这也是为什么长序列任务要么减头数、要么用稀疏注意力的变体。至此 Transformer 的三件套——投影、缩放点积注意力、多头——你都摸过了,剩下的位置编码与层堆叠属于结构细节,阅读开源实现时按图索骥即可。

变式:把 num_heads 改成 1 和 8 各训一个小任务(第 5 章循环直接可用),对比收敛速度与显存占用——你会同时验证"多头提升表达"与"多头增加开销"两句看似矛盾的陈述。

本节要点回顾

  • 循环网络的短板:固定容量隐状态挤掉早期信息、无法并行;
  • 注意力三件套:查询找、键响应、值交付,权重是缩放点积加 softmax;
  • 掩码与 ignore_index 成对,补齐位在注意力和损失两端都要剔除;
  • 多头拆分匹配模式,代价是平方增长的注意力矩阵。

下一节转入后勤:用 TensorBoard 把训练过程记录成可回看、可对比的行军日志。


作者与出处
原作者: 灏天文库
来源:灏天文库
整理: 灏天文库整理
由灏天文库平台收录,内容或由平台用户上传,仅供学习交流
发布者: 作者: 灏天文库 转发
评论区 (0)
U