PaddleNLP attention_utils 模块深度解析:BigBird 稀疏注意力与 MultiHeadAttention 的实现与实战
2026/9/24 14:18:07 网站建设 项目流程
  • 人工智能
  • 大模型
  • 预训练
  • 微调
  • LoRA
  • RLHF
  • 强化学习
  • 分布式训练

【免费下载链接】PaddleNLP

Easy-to-use and powerful LLM and SLM library with awesome model zoo.

项目地址:https://gitcode.com/gh_mirrors/pa/PaddleNLP
点击查看免费下载

paddlenlp.transformers.attention_utils是 PaddleNLP 中与 Transformer 注意力机制直接相关的基础工具模块,承载了 BigBird 稀疏注意力的完整实现(全局注意力 + 窗口注意力 + 随机注意力)、默认的缩放点积注意力、可插拔的注意力实现注册机制,以及一个支持增量推理缓存的多头注意力封装。本文以 docs/zh/source/paddlenlp.transformers.attention_utils.rst 所列 API 为线索,深入 attention_utils.py 源码,讲解其中每个类与函数的内部逻辑、参数含义,并结合 BigBird 模型实现 与测试用例还原完整的调用链,帮助读者掌握如何在长序列场景下使用与扩展这套注意力基础设施。

一、模块定位:PaddleNLP 的注意力基础设施

docs/zh/source/paddlenlp.transformers.attention_utils.rst是一份 Sphinx 自动文档页,通过automodule指令把 paddlenlp/transformers/attention_utils.py 中所有带文档的公开成员自动渲染成 API 参考。因此,该模块的真实内容即其源码,核心成员包括:

成员类型作用
Registry/AttentionRegistry类 / 实例注意力实现的注册表,通过装饰器按名称注册实现
create_bigbird_rand_mask_idx函数生成单层 BigBird 随机注意力块的索引
create_bigbird_rand_mask_idx_list函数为每一层生成一份随机块索引列表
_convert_param_attr_to_list函数将单个ParamAttr统一展开为长度 n 的列表
Linear3D三维 Q/K/V 线性投影层,输出[B, H, T, D]张量
Attention抽象基类注意力实现基类,定义统一的前向接口
DefaultAttention注册名为default_attention的缩放点积注意力
BigBirdSparseAttention注册名为bigbird的块稀疏注意力
MultiHeadAttention多头注意力封装,含 Cache / StaticCache 推理缓存

从导入关系看,该模块是paddlenlp.transformers的公开基础设施之一:init.py 将create_bigbird_rand_mask_idx_list直接导出到顶层命名空间;BigBird 模型则同时导入MultiHeadAttention_convert_param_attr_to_list使用。可以推断,该模块是 PaddleNLP 中 BigBird 长序列模型专属的注意力工具箱。

二、注册表模式:AttentionRegistry 与可插拔注意力

模块用一个轻量注册表来解耦"注意力实现"与"使用方":

class Registry(object): def __init__(self): self.cls_dict = {} def register(self, name): def add_item(name, cls): self.cls_dict[name] = cls return cls return lambda cls: add_item(name, cls) AttentionRegistry = Registry()
  • register(name)返回一个装饰器,被装饰的类会以name为键存入AttentionRegistry.cls_dict,同时原类被原样返回,因此装饰不会改变类的行为(attention_utils.py#L26-L38)。
  • 模块内已有两个注册项:default_attentionDefaultAttentionbigbirdBigBirdSparseAttention
  • 消费方通过字符串名取实现:AttentionRegistry.cls_dictattention_type(见 attention_utils.py#L558-L560)。

这种设计让MultiHeadAttention无需感知具体算法:用户只需在配置中指定attention_type="bigbird""original_full",即可切换不同注意力实现;想要新增算法时,也只需实现Attention子类并注册即可,无需改动既有模型代码。

三、随机块索引生成:create_bigbird_rand_mask_idx 与 create_bigbird_rand_mask_idx_list

BigBird 稀疏注意力的关键难点在于:每个查询块除了关注全局块和窗口块之外,还要随机关注若干"随机块"。为了在多个 batch、多个 head 之间复用同一套随机采样结果,模块在 CPU 上用 NumPy 一次性生成好索引,再搬运到设备端使用。

create_bigbird_rand_mask_idx(num_layers, query_length, key_length, num_heads, block_size, window_size, num_global_blocks, num_rand_blocks, seed)(attention_utils.py#L41-L87)的核心流程:

  1. block_size把序列切块,得到num_key_blocks = key_length // block_sizenum_query_blocks = query_length // block_size,窗口半宽num_window_blocks = window_size // 2
  2. 对每个查询块,计算其"非法块集合":
    • 左右num_window_blocks范围内的窗口块(这些块走窗口注意力,不应再出现在随机块中);
    • num_global_blocks个全局块(这些块走全局注意力);
    • 序列边界处会回卷补齐,保证头部、尾部查询块的窗口块数量一致。
  3. 对每个 head 独立做np.random.permutation,从合法块中依次取出num_rand_blocks个作为该查询块的随机块。
  4. 最后把所有 head 的索引堆叠并做一次"减num_global_blocks // 2"的偏移变换,同时把 head 编号与块编号拼成[H*T, 2]形式的 gather 索引列表,供后续gather_nd使用。

create_bigbird_rand_mask_idx_list(num_layers, ...)(attention_utils.py#L90-L108)则是对上面函数按层数做列表推导,返回形状为[num_layers, H, L, 2](代码中经np.stack堆叠)的完整索引,保证每一层使用不同的随机采样结果。

调用示例(与 BigBirdModel 官方示例一致):

import paddle from paddlenlp.transformers import BigBirdModel, BigBirdTokenizer from paddlenlp.transformers import create_bigbird_rand_mask_idx_list tokenizer = BigBirdTokenizer.from_pretrained('bigbird-base-uncased') model = BigBirdModel.from_pretrained('bigbird-base-uncased') config = model.config max_seq_len = 512 text = "This is a docudrama story on the Lindy Chamberlain case ..." input_ids = tokenizer.convert_tokens_to_ids(tokenizer(text)) input_ids.extend([0] * (max_seq_len - len(input_ids))) seq_len = len(input_ids) input_ids = paddle.to_tensor([input_ids]) rand_mask_idx_list = create_bigbird_rand_mask_idx_list( config["num_layers"], seq_len, seq_len, config["nhead"], config["block_size"], config["window_size"], config["num_global_blocks"], config["num_rand_blocks"], config["seed"]) rand_mask_idx_list = [paddle.to_tensor(idx) for idx in rand_mask_idx_list] output = model(input_ids, rand_mask_idx_list=rand_mask_idx_list)

需要注意:从 BigBirdModel.forward 的实现看,模型内部会依据config里的num_layers / nhead / block_size / window_size / num_global_blocks / num_rand_blocks / seed重新生成rand_mask_idx_list,因此调用方即使不显式传入,模型也能自洽运行;生成逻辑本身仍是create_bigbird_rand_mask_idx_list。随机采样的可复现性由seed参数控制,配置中seed=None时则每次生成结果不同。

四、参数工具函数:_convert_param_attr_to_list

_convert_param_attr_to_list(param_attr, n)(attention_utils.py#L111-L137)用于把用户传入的ParamAttr统一规整为长度n的列表,是 PaddleNLP 模型中"一个配置同时驱动多个层"的常见做法:

  • 传入list/tuple:要求长度必须等于n,逐项规整;True转为默认ParamAttrFalse表示该层不创建参数(记为False),否则调用ParamAttr._to_attr归一化。
  • 传入单个bool:为True时生成n份默认ParamAttr,否则生成nFalse
  • 传入单个ParamAttr:深拷贝为n份,若属性带名字则在末尾追加_i后缀以避免参数名冲突。

在 BigBird 的 TransformerEncoderLayer 中,它被用来一次性为"自注意力层 + FFN 层"展开weight_attrbias_attr,再分别传入MultiHeadAttentionLinear

五、三维线性投影:Linear3D

常规nn.Linear输出[B, T, D],而多头注意力需要先把投影结果拆成多头。Linear3D(attention_utils.py#L140-L163)把这两步合并:

result = paddle.matmul(input, self.weight) # [B, T, D] x [D, D] result += paddle.reshape(self.bias, [1, 1, D]) # 加偏置 result = paddle.reshape(result, [B, T, H, -1]) # 拆出头维度 result = paddle.transpose(result, [0, 2, 1, 3]) # [B, H, T, D]

输入形状为[B, T, D](D 即hidden_size),权重形状为[hidden_size, hidden_size],输出直接就是多头注意力所需的[B, H, T, D]布局,省去了在MultiHeadAttention外层反复reshape + transpose的样板代码。MultiHeadAttentionq_proj / k_proj / v_proj全部使用该层实现(attention_utils.py#L553-L555)。

六、注意力基类与默认实现

6.1 Attention 基类

Attention(attention_utils.py#L166-L182)定义了所有注意力实现的统一前向协议:

def forward(self, query_matrix, key_matrix, value_matrix, d_head, attn_mask=None, rand_mask_idx=None, query_mask=None, key_mask=None, dropout=None): raise NotImplementedError

参数语义:

  • query_matrix / key_matrix / value_matrix:形状均为[B, H, T, D]
  • d_head:单头维度,用于缩放点积;
  • attn_mask:加法式注意力掩码(直接加到 logits 上);
  • rand_mask_idx:BigBird 随机块的 gather 索引;
  • query_mask / key_mask:形状分别为[B, 1, T, 1][B, 1, 1, T]的布尔掩码;
  • dropout:注意力权重上的 dropout 比例。

子类只需实现该协议即可接入MultiHeadAttention

6.2 DefaultAttention:缩放点积注意力

DefaultAttention(注册名default_attention,attention_utils.py#L185-L210)是标准的缩放点积注意力:

  1. 计算product = Q @ Kᵀ,再乘以d_head ** -0.5做缩放;
  2. (1 - Q_mask @ K_mask) * -1e6把 padding 位置压成很大的负值,使 softmax 后权重趋近于 0(掩码部分用矩阵乘法生成,等价于按行广播的 padding 掩码);
  3. 若传入attn_mask则继续累加(支持自定义加法掩码,例如单向因果掩码);
  4. softmax得到权重,若dropout非空则以upscale_in_train模式做训练期 dropout;
  5. out = weights @ V输出。

它对应的就是 BigBird 配置中attention_type="original_full"的 O(n²) 全注意力路径(见 BigBirdConfig 文档)。

七、BigBirdSparseAttention:全局 + 窗口 + 随机三路稀疏注意力

BigBirdSparseAttention(注册名bigbird,attention_utils.py#L213-L518)是整个模块最核心、最复杂的部分。它的目标是把 O(n²) 的全注意力降为近似线性复杂度:每个 token 只关注三类 key/value 块——序列首尾的全局块、自身附近的窗口块、以及按预生成索引采样的随机块。

7.1 超参与分块策略

__init__接收num_heads, block_size, window_size, num_global_blocks, num_rand_blocks, seed,并额外计算:

self.num_global_blocks_back = num_global_blocks // 2 self.num_global_blocks_front = (num_global_blocks // 2 if num_global_blocks % 2 == 0 else num_global_blocks // 2 + 1)

即把全局块均分到序列头(front)与尾(back)两侧;奇数个全局块时前端多分一块。前向开始时,输入[B, H, T, D]会被 reshape 成[B, H, L, bs, D]L = T // bs为块数),query/key/value 与对应的 mask 全部按块切分(attention_utils.py#L443-L447)。

7.2 全局注意力:_get_global_out

_get_global_out(query_matrix, key_matrix, value_matrix, key_mask, d_head, dropout, is_front)(attention_utils.py#L391-L404)让序列最前GF * bs个 token(或最后GB * bs个 token)作为 query,对完整序列做标准缩放点积注意力,产出全局块输出。从源码结构看,它内部同样用(1 - key_mask) * -1e6屏蔽 padding,并在 softmax 后做weights @ V

7.3 带内注意力:_get_band_mask 与 _get_band_matrix

对于中间的非全局查询块,需要同时聚合"前端全局块 + 窗口块 + 后端全局块",对应两个辅助函数:

  • _get_band_mask(blocked_query_mask, blocked_key_mask, batch_size, sequence_length)(attention_utils.py#L227-L289)用 mask 张量拼出形状[B, H, L-G, bs, (G+W)*bs]的合法注意力范围掩码。头部/尾部查询块通过zeros_like补零与concat实现窗口块的"回卷",保证每个查询块都恰好看到W个窗口块。
  • _get_band_matrix(blocked_matrix, B, T)(attention_utils.py#L291-L348)用同样的回卷逻辑,从已分块的 K/V 矩阵中取出对应的(G+W)个块,重排成[B, H, L-G, (G+W)*bs, D]的"带内 key/value 矩阵",其中全局块部分通过expand广播到每个查询位置。

7.4 随机注意力:_get_rand_mask 与 _gather_random_key_value

  • _get_rand_mask(blocked_query_mask, blocked_key_mask, rand_mask_idx, batch_size, sequence_length)(attention_utils.py#L350-L372)依据rand_mask_idxgather_nd把每个 head 对应的随机 key 掩码抓取出来,再与查询掩码做einsum得到形状[B, H, L-G, bs, R*bs]的随机注意力掩码。
  • _gather_random_key_value(blocked_matrix, rand_mask_idx, B, T)(attention_utils.py#L374-L389)对 K 和 V 分别执行相同的gather_nd,得到[B, H, L-G, R*bs, D]的随机 key/value 矩阵。

这两个函数都依赖第二节生成的rand_mask_idx(形状[H, T][head_id, block_id]对),因此随机块采样是"预计算 + 查表"而非前向中实时采样,这也是seed参数能保证可复现的原因。

7.5 forward:三路结果合并

forward(attention_utils.py#L410-L518)的完整流程:

  1. 计算前端全局块输出global_front_out与后端全局块输出global_back_out
  2. 拼接带内掩码与随机掩码得到second_mask,拼接带内 K/V 与随机 K/V 得到second_key_matrix / second_value_matrix
  3. 取中间查询块second_query_matrix = blocked_query_matrix[:, :, GF:-GB],通过einsum("bhlqd,bhlkd->bhlqk")计算分数、按d_head**-0.5缩放、加掩码、softmax;
  4. _get_splited_matrix把权重与 V 按"前端窗口 / 中间 / 后端窗口"切成三份分别加权求和——其中中间部分需要额外补上全局块与随机块对应的分值(见 attention_utils.py#L496-L508);
  5. 把三部分输出拼回[B, H, (L-G)*bs, D],再与全局前/后输出拼接成完整的[B, H, T, D],最后乘以query_mask屏蔽 padding。

整体数据流可概括为:

Q/K/V [B,H,T,D] --按块 reshape--> [B,H,L,bs,D] ├── 全局前/后块 → 标准全注意力(_get_global_out) ├── 中间查询块 → 带内块(_get_band_*)+ 随机块(_gather_random_key_value) └── 三路 concat → out [B,H,T,D] × query_mask

八、MultiHeadAttention:支持推理缓存的多头注意力封装

MultiHeadAttention(attention_utils.py#L521-L619)把投影、注意力实现、头合并与推理缓存整合成一个可直接使用的nn.Layer

8.1 构造参数

参数默认值说明
embed_dim必填模型隐藏维度,必须是num_heads的整数倍(构造时断言)
num_heads必填注意力头数,head_dim = embed_dim // num_heads
dropout0.0注意力权重 dropout 比例
kdim / vdimembed_dimK/V 投影输入维度,便于处理 cross-attention 中 K/V 与 Q 维度不同的场景
weight_attr / bias_attrNoneQ/K/V 与输出投影的参数属性
block_size1BigBird 块大小
window_size3BigBird 窗口块数量
num_global_blocks1BigBird 全局块数量
num_rand_blocks1BigBird 随机块数量
seedNone随机块采样种子
attention_type"bigbird"AttentionRegistry选择实现

内部结构:q_proj / k_proj / v_projLinear3Dout_projnn.Linearattn_impl由注册表按attention_type实例化(attention_utils.py#L553-L560)。

8.2 Cache 与 StaticCache

模块定义了两个命名元组(attention_utils.py#L523-L524):

Cache = collections.namedtuple("Cache", ["k", "v"]) StaticCache = collections.namedtuple("StaticCache", ["k", "v"])
  • Cache:增量式(incremental)缓存,用于自回归解码的 self-attention。_prepare_qkvisinstance(cache, self.Cache)时把新计算的 K/V 沿序列维concat到历史缓存上,实现"只算新 token"的逐 token 解码(attention_utils.py#L571-L575)。
  • StaticCache:静态 K/V 缓存,用于 encoder-decoder 场景(如 UniLM):_prepare_qkv检测到该类型时直接复用缓存中的 K/V,不再重新投影(attention_utils.py#L565-L568)。

gen_cache(key, value=None, type=Cache)(attention_utils.py#L584-L595)负责按类型构造缓存:StaticCache需要立即对key计算 K/V;Cache在未提供value时返回形状[-1, num_heads, 0, head_dim]的空缓存,传入value时则把初始 K/V 装入缓存(注释标明主要用于 UniLM 等场景)。

8.3 forward 数据流

q = self.q_proj(query) # [B, H, T, D] if isinstance(cache, self.StaticCache): k, v = cache.k, cache.v # 复用静态缓存 else: k, v = self.compute_kv(key, value) # k_proj / v_proj if isinstance(cache, self.Cache): k = paddle.concat([cache.k, k], axis=2) # 增量拼接 v = paddle.concat([cache.v, v], axis=2) out = self.attn_impl(q, k, v, self.head_dim, attn_mask, rand_mask_idx, query_mask, key_mask, self.dropout) out = paddle.transpose(out, [0, 2, 1, 3]) # [B, T, H, D] out = paddle.reshape(out, [0, 0, out.shape[2] * out.shape[3]]) out = self.out_proj(out) # 合并多头并投影

attention_type="bigbird"时,attn_impl就是上一节的BigBirdSparseAttention,因此MultiHeadAttention既保留了多头注意力的标准外接口(投影、多头合并、缓存),又完全继承了 BigBird 的稀疏计算。

九、在 BigBird 模型中的集成调用链

以 paddlenlp/transformers/bigbird/modeling.py 为参考,整条调用链为:

  1. 配置层BigBirdConfig定义attention_type(默认"bigbird")、block_sizewindow_sizenum_global_blocksnum_rand_blocks等超参。预训练配置bigbird-base-uncased中:block_size=16window_size=3num_global_blocks=2num_rand_blocks=3max_position_embeddings=4096(见 configuration.py#L23-L45)。
  2. 编码层TransformerEncoderLayer.__init___convert_param_attr_to_list展开参数后构造MultiHeadAttention(..., attention_type=config.attention_type, block_size=..., window_size=..., num_global_blocks=..., num_rand_blocks=..., seed=...)(modeling.py#L84-L96)。
  3. 模型层BigBirdModel._process_mask依据pad_token_id生成attention_mask / query_mask / key_mask(modeling.py#L380-L394);forward内部调用create_bigbird_rand_mask_idx_list为每一层生成随机块索引,再交给TransformerEncoder逐层前向(modeling.py#L501-L521)。
  4. 任务层BigBirdForSequenceClassification / BigBirdForQuestionAnswering / BigBirdForPretraining等在BigBirdModel之上加输出头,复用同一套稀疏注意力。

十、测试验证

仓库在 tests/transformers/bigbird/test_modeling.py 中对 BigBird 相关实现做了系统测试:

  • BigBirdModelTester覆盖batch_size=13、seq_length=7、hidden_size=32、num_attention_heads=4、num_hidden_layers=5等配置,通过BigBirdConfig(...)构造模型并验证各任务模型(BigBirdForMultipleChoiceBigBirdForQuestionAnsweringBigBirdForSequenceClassificationBigBirdForTokenClassificationBigBirdForPretraining)的输入输出(test_modeling.py#L40-L135)。
  • 测试通过parameterized_class参数化return_dict等选项,并复用通用ModelTesterMixin检查attention_utils相关实现与 PaddleNLP 模型基类约定的兼容性(test_modeling.py#L20-L37)。

这些测试一方面印证了attention_utils中注册表、随机索引生成、MultiHeadAttention等组件的正确性,另一方面也表明该模块是 PaddleNLP 模型库中可被独立测试与复用的公共组件。

十一、小结与实践建议

paddlenlp.transformers.attention_utils为 PaddleNLP 提供了一整套"即插即用"的注意力基础设施:

  • 需要标准全注意力时,使用attention_type="original_full"DefaultAttention);
  • 需要长序列稀疏注意力时,使用attention_type="bigbird"BigBirdSparseAttention),并配合create_bigbird_rand_mask_idx_list预生成随机块索引,注意block_size需能整除序列长度、embed_dim需能被num_heads整除;
  • 自回归解码时通过MultiHeadAttention.gen_cache初始化Cache实现增量 K/V 缓存;
  • 想接入自定义注意力算法时,继承Attention并实现统一forward协议,再用@AttentionRegistry.register("your_name")注册,即可通过配置字符串启用。

理解该模块,是深入 PaddleNLP 长序列模型(尤其是 BigBird 系列)源码、乃至在其基础上做注意力算法二次开发的第一站。

  • 人工智能
  • 大模型
  • 预训练
  • 微调
  • LoRA
  • RLHF
  • 强化学习
  • 分布式训练

【免费下载链接】PaddleNLP

Easy-to-use and powerful LLM and SLM library with awesome model zoo.

项目地址:https://gitcode.com/gh_mirrors/pa/PaddleNLP
点击查看免费下载

相关推荐

上一篇:告别9KB冗余!bignumber.js生产环境极致优化指南
下一篇:突破Windows文件系统限制:使用WinFsp实现压缩镜像实时挂载的终极指南

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询