首页
/ TensorFlow Models NLP 建模层库详解:official/nlp/modeling/layers 中 Transformer 与注意力机制的源码级剖析

TensorFlow Models NLP 建模层库详解:official/nlp/modeling/layers 中 Transformer 与注意力机制的源码级剖析

2026-09-04 20:44:47作者:彭桢灵Jeremy

本文围绕 Layers 模块文档 展开,系统讲解该仓库 NLP 建模层库中 20 余个核心 Keras 层——多头注意力、稀疏/线性注意力、带缓存的解码器、掩码与位置编码、任务头(MaskedLM、分类头)等——的设计意图与源码实现。读完后,你将能够直接选用、组合这些层构建新的 tf.keras NLP 模型,并理解 BertEncoderAlbertEncoder 等上层网络是如何调用这些积木搭出完整 Transformer 的。

1. Layers:NLP 模型的基本构建块

模块入口文档的定义非常直接(见 README):

Layers are the fundamental building blocks for NLP models. They can be used to assemble new tf.keras layers or models.

也就是说,这一层级的库不做"端到端模型",而是把注意力机制、掩码、位置编码、任务头、分词打包输入等能力拆成可复用的 tf_keras.layers.Layer,供上层网络(encoder/network)自由组装。模块统一导出集中在 official/nlp/modeling/layers/init.py,其中可以看到:

  • 大量层使用 @tf_keras.utils.register_keras_serializable(package="Text") 注册,使其可参与 Keras 模型序列化/保存;
  • 导出列表比 README 覆盖的条目更宽,还包括 MoeLayer(混合专家)、FactorizedEmbeddingPackBertEmbeddingsTNTransformerExpandCondenseTransformerScaffoldPerDimScaleAttentionMultiQueryAttention 等进阶实现,均位于 layers 目录 下。

下文的组织方式与 README 的条目一一对应,并按"注意力 → 块结构 → 掩码/Softmax → 位置编码 → 任务头 → 文本前处理"展开。

2. 核心注意力层

2.1 MultiHeadAttention 与 CachedAttention

README 指出 MultiHeadAttention 实现了可选掩码的 query/key/value 注意力("Attention Is All You Need"),当 from_tensorto_tensor 相同时即自注意力。从源码看,attention.py 中的实现是一行别名:

MultiHeadAttention = tf_keras.layers.MultiHeadAttention

即直接复用 Keras 内置多头注意力;本仓库真正新增的是同文件中的 CachedAttention——一个继承自 tf_keras.layers.MultiHeadAttention 的"带缓存注意力层",专门服务自回归解码:

  • 普通路径decode_loop_step is None):_update_cachetf.concat 把历史 key/value 与新 key/value 沿序列维拼接,再写回 cache["key"]cache["value"],返回全长度张量参与后续 einsum;
  • TPU 特化路径(传入 decode_loop_step):缓存张量形状固定,无法动态拼接,于是用 tf.one_hot(decode_loop_step, key_seq_dim) 构造单步索引,把当前步的 key/value 以加权方式"叠加"进固定形状的缓存,避免 shape 变化导致 TPU 重编译;
  • call(query, value, key=None, attention_mask=None, cache=None, decode_loop_step=None, ...) 返回 (attention_output, cache)(请求 return_attention_scores=True 时再附带 attention_scores)。

这一"cache + decode_loop_step"接口正是后文 TransformerDecoderBlock 逐 token 生成的底层支撑。

2.2 TalkingHeadsAttention 与 MultiChannelAttention

  • TalkingHeadsAttention(Talking-Heads Attention 论文变体):在 softmax 前/后对各头之间做线性"交谈",打破多头之间完全独立的假设;
  • MultiChannelAttention:多头注意力的一种变体,可把多路流合并来做交叉注意力。从 transformer.pyTransformerDecoderBlock 可以看到它的典型用法——构造参数 multi_channel_cross_attention=True 时,encoder-decoder 交叉注意力默认替换为 MultiChannelAttention,且 call 的输入需要额外第 5 个张量(doc-attention 概率),见源码中对 inputs[-1] 的解包。

2.3 BigBirdAttention:把二次复杂度降为线性

BigBirdAttention 实现 Big Bird 论文的稀疏注意力。源码中与注意力本体配合的掩码构造函数最能体现其"局部窗口 + 全局 token + 随机块"的组合结构:

  • create_band_mask_from_inputs(from_blocked_mask, to_blocked_mask):由分块 2D 掩码生成局部窗口(band)的 3D 注意力掩码,形状为 [batch, 1, L/b - 4, b, 3*to_block_size],即每个 query 块只看邻近 3 个 key 块;
  • bigbird_block_rand_mask(...):按行生成"随机块"邻接表,last_idx 可限制随机块只能选在序列前缀内(保证因果性);
  • create_rand_mask_from_inputs(...):把随机块索引展开成逐头 3D 掩码,供注意力打分阶段屏蔽。

文件顶部还定义了 MAX_SEQ_LEN = 4096,即该实现面向的序列长度上限。模块同时导出 BigBirdMasksinit.py 第 24 行),用于封装上述三类掩码的生成。

2.4 KernelAttention:核特征图的线性注意力

KernelAttention 把自注意力表示为核特征图的线性点积,利用矩阵乘法结合律把复杂度从 O(n²) 降到 O(n);README 说明其涵盖了 Linear Attention、Performer、Random Feature Attention 三类方法。从 kernel_attention.py 源码看,实现依赖一组窗口/分块工具函数:

  • pad_to_chunk_length / split_tensor_into_chunks:把序列切成长度可整除的 chunk,配合 TPU 分块计算;
  • rectangular_window_sum:用前缀和差值在滑动矩形窗口上做求和(O(n) 实现"局部平滑");
  • weighted_window_sum:用 tf.nn.depthwise_conv2d 实现加权滑动窗,等价于对核特征做因果卷积式归一化;
  • 配套的 KernelMask 把普通 2D 输入掩码 [batch, seq] 转成 KernelAttention 需要的掩码格式。

2.5 ReuseMultiHeadAttention 与 ReuseTransformer

Reuse Transformer 一族的 TF 实现位于 reuse_attention.pyreuse_transformer.py。其思想是:相邻层注意力分数高度冗余,高层可复用小一层的分数,省去部分点积计算。源码要点:

  • reuse_attention 参数:0 不复用;-1 表示全部头复用(源码中会归一化为 num_heads);中间值表示部分头复用,且构造时会校验取值在 [-1, num_heads] 内;
  • 当"复用头数 < 总头数"时,_build_from_signature 把 query/key/value/输出投影拆成 value_reusevalue_new 两组 EinsumDense_compute_attention 中先对新头正常计算 new_scores(可选叠加相对位置偏置),再把传入的 reuse_scores[:, :reuse_heads, :, :] 与之沿头维 tf.concat;全部复用时(reuse_heads == num_heads)则直接使用 reuse_scores、完全跳过 Q/K 投影与点积;
  • 相对位置偏置:use_relative_pe=True 且非全复用时,额外创建形状 [1, num_heads - reuse_heads, 2*pe_max_seq_length - 1]relative_position_embeddings 变量(pe_max_seq_length 默认 512),_compute_relative_position 通过索引矩阵查表得到逐头偏置,dtype 跟随 Keras 混合精度全局策略(mixed_bfloat16/mixed_float16/float32)。

2.6 ReZeroTransformer

ReZeroTransformer 在标准 Transformer 块的残差上引入逐元素可学习标量门控(ReZero 论文),用于缓解深层 Transformer 的收敛问题。它与 TransformerEncoderBlock 结构同构,仅残差路径不同,可对照阅读 rezero_transformer.py 及其测试 rezero_transformer_test.py

2.7 相对位置注意力与 TransformerXL

  • MultiHeadRelativeAttention:Transformer-XL 的相对位置编码注意力变体,且源码层面扩展支持了 XLNet 提出的基于 segment 的注意力偏置;
  • TwoStreamRelativeAttention:XLNet 的双流(query 流 + content 流)相对自注意力;
  • TransformerXL:包含 TransformerXLBlock(一个或双流相对自注意力 + 前馈网络)与 TransformerXL(管理 attention bias 及堆叠多个 block)两个类,二者均从 init.py 第 76-77 行导出。

2.8 MobileBert 专用层

MobileBertEmbeddingMobileBertTransformer 分别实现 MobileBERT 论文提出的轻量嵌入层与 Transformer 层。init.py 同时还导出了同文件中的 MobileBertMaskedLM,说明该目录对 MobileBERT 提供了"嵌入 + 块 + MLM 头"的成套组件。

3. Transformer 块:Transformer、TransformerDecoderBlock 与 TransformerEncoderBlock

3.1 Transformer(已弃用的别名层)

Transformer 直接继承 TransformerEncoderBlock 并转发全部参数(intermediate_sizeinner_dimdropout_rateoutput_dropout 等),但构造时会发出明确的弃用告警:

"The Transformer layer is deprecated. Please directly use TransformerEncoderBlock."(transformer.py

同文件还定义了 CompiledTransformer:在 call 上叠加 @tf_function_if_eager(experimental_compile=True)(来自 util.py),用于 TF Function 自动编译加速的场景。阅读旧配置(gin 配置里常见 Transformer 键名)时需要知道它等价于 TransformerEncoderBlock

3.2 TransformerEncoderBlock:当前推荐的标准块

TransformerEncoderBlock 是"多头注意力 + 两层前馈"的标准编码器块实现,也是本仓库 encoder 网络的默认积木。构造参数(源码签名,transformer_encoder_block.py):

参数 默认值 说明
num_attention_heads 必填 注意力头数
inner_dim 必填 前馈中间层宽度
inner_activation 必填 前馈激活函数
output_range None 对输入序列取 [0, output_range) 切片,None 表示不切
norm_first False False 为 Post-LN(对块输出归一化),True 为 Pre-LN(对输入归一化)
norm_epsilon 1e-12 LayerNorm/RMSNorm 的 epsilon
use_rms_norm False 用同文件定义的 RMSNorm 替代 LayerNorm
output_dropout / attention_dropout / inner_dropout 0.0 三处独立 dropout
num_kv_heads None 指定 KV 头数(Multi-Query/GQA 风格)
linformer_dim None 低秩线性注意力投影维度
use_sigmoid_attn / sigmoid_attn_bias False / None sigmoid 注意力开关与偏置
return_attention_scores False 是否额外返回注意力分数

同文件的 RMSNorm 实现也很简洁:对输入先转 float32,计算平方均值后 inputs * rsqrt(var + epsilon) * scale,再转回原 dtype,且 scale 权重关闭 autocast 以保持精度。

上层网络的真实调用可佐证其地位:bert_encoder.pybert_encoder.pyalbert_encoder.py(ALBERT 的"参数共享"层即复用同一个 TransformerEncoderBlock 实例)、seq2seq_transformer.py 都在实例化该块。

3.3 TransformerDecoderBlock

TransformerDecoderBlock 是解码器单层,由三个子层构成(源码 docstring 与 build 一致):

  1. 自注意力:默认类为 attention.CachedAttention(可经 self_attention_cls 替换),因此天然带 cache 机制;
  2. encoder-decoder 交叉注意力:默认 MultiHeadAttentionmulti_channel_cross_attention=True 时换成 MultiChannelAttention,也可用 cross_attention_cls 显式指定;
  3. 位置前馈网络EinsumDense("abc,cd->abd") → 激活 → dropout → 输出投影。

call(inputs, cache=None, decode_loop_step=None)inputs 为四元组 (input_tensor, memory, attention_mask, self_attention_mask)(多通道交叉注意力时为五元组),源码按 norm_first 分支在 Pre-LN/Post-LN 两种残差排布间切换,最终返回 (layer_output, cache)——把缓存回传给下一层,从而支撑整段解码循环。build 中还有一个实用约束:输入必须为三维 [batch, sequence, width],且 width % num_attention_heads == 0,否则抛出 ValueError

4. 掩码与 Softmax

4.1 SelfAttentionMask

SelfAttentionMask 从 2D 掩码生成 3D 自注意力掩码。实现原理(get_mask 函数):

  • 输入 inputs 形状 [batch, from_seq_length, ...]to_mask 形状 [batch, to_seq_length](int32,1 表示有效、0 表示需屏蔽);
  • 先把 to_mask reshape 为 [batch, 1, to_seq_length],再用 tf.broadcast_to 广播到 [batch, from_seq_length, to_seq_length]

这避免了显式构造平方张量,内存开销仅为广播视图。该文件同时提供独立的 get_mask(inputs, to_mask, dtype=None) 函数供非 Layer 场景直接调用。

4.2 MaskedSoftmax

MaskedSoftmax 实现带可选掩码的 softmax,README 中"1 表示放行、0 表示屏蔽,被屏蔽位置输出近似为 0"的语义在源码中对应一个关键细节:

adder = (1.0 - tf.cast(mask, scores.dtype)) * _large_compatible_negative(scores.dtype)
scores += adder

其中 _large_compatible_negative 对 float32 返回 -1e9,而 float16 无法表示 -1e9,于是返回 tf.float16.min——这是半精度训练下防止"负得不够大"导致泄漏的实现要点。另外:

  • mask_expansion_axes:当掩码比分数张量维度少时,循环 tf.expand_dims 到指定轴,使 [B, T, S] 的掩码能适配 [B, H, T, S] 的分数;
  • normalization_axes 默认 (-1,);多轴归一化时改用 exp(scores - reduce_logsumexp(...)) 的数值稳定写法,而不是 tf.nn.softmax(后者只支持单轴)。

5. 位置编码

PositionEmbedding 按 BERT 论文方式创建可学习的位置嵌入。构造参数:max_length(必填,动态序列最大长度)、initializer="glorot_uniform"seq_axis=1(在哪个轴上加嵌入)。其文档示例即完整可运行片段:

position_embedding = PositionEmbedding(max_length=100)
inputs = tf_keras.Input((100, 32), dtype=tf.float32)
outputs = position_embedding(inputs)

PositionEmbeddingbuildmax_length 创建嵌入表、按输入末维确定宽度;同文件还导出 RelativePositionBiasRelativePositionEmbedding(见 init.py),分别服务"相对位置偏置"与"相对位置嵌入"两条路线。

6. 任务头与不确定性建模

6.1 MaskedLM

MaskedLM 是 BERT 的掩码语言模型头,README 强调"它假设外部传入嵌入表变量"。源码印证了这一约定:

  • 构造参数:embedding_table(必须来自 encoder 的 get_embedding_table())、activationinitializer='glorot_uniform'output='logits' | 'predictions'(非法取值直接抛 ValueError);
  • build 中从嵌入表形状 (vocab_size, hidden_size) 反推维度,依次创建 Dense(hidden_size)transform/dense)、LayerNormalization(epsilon=1e-12)transform/LayerNorm)与形状 [vocab_size]output_bias/bias
  • call(sequence_data, masked_positions)_gather_indexes 取出被掩码位置,再投影、归一化、与词表嵌入做点积输出 logits。文档给出的最小用法:
encoder = modeling.networks.BertEncoder(...)
lm_layer = MaskedLM(embedding_table=encoder.get_embedding_table())

6.2 ClassificationHead 与 GaussianProcessClassificationHead

ClassificationHead 是"在嵌入序列上做池化的分类头"。源码参数:inner_dim(0/None 时只建输出投影)、num_classescls_token_idx=0(在序列第几位做池化,通常取 [CLS] 位)、activation="tanh"dropout_rate=0.0;内部结构为 pooler_dense(可选)→ Dropout → logits 输出层。call(features, only_project=False) 支持只取池化向量不接分类投影。

同文件的 GaussianProcessClassificationHead 是 SNGP(光谱归一化神经网络高斯过程)分类头,依赖 gaussian_process.pyRandomFeatureGaussianProcess(随机特征 GP,见 "Random Features for Large-Scale Kernel Machines")与 spectral_normalization.pySpectralNormalizationtf.Wrapper,对内部层应用谱范数正则),三者组合实现了"距离感知"的不确定性估计分类头。

6.3 其他头/正则层

  • MatMulWithMargin:带 margin 的矩阵乘法层,用于检索/排序类任务(ADD 模型的双编码器加性 margin softmax);
  • GatedFeedforward:GLU 变体门控前馈层("GLU Variants Improve Transformer");
  • OnDeviceEmbedding:为 TPU 模型设计的高效嵌入查找层,适合词表极大的场景。

7. 文本前处理层:从原始文本到 BERT 输入

text_layers.py 提供把"原始文本 → 模型输入"整条流水线 Keras 层化的能力,README 列出的三个类及 init.py 补充导出的 FastWordpieceBertTokenizer 均在其中:

这使得数据管道可以不落地 TFRecord,而是在图中直接完成分词与打包。

8. 组合使用:从 README 条目到真实模型

把 README 的条目按"积木分类"汇总如下(全部文件路径均位于 official/nlp/modeling/layers/):

类别 实现文件
基础注意力 MultiHeadAttention attention.py
解码缓存注意力 CachedAttention attention.py
头间交互 TalkingHeadsAttentionMultiChannelAttention talking_heads_attention.pymulti_channel_attention.py
长序列稀疏注意力 BigBirdAttentionBigBirdMasks bigbird_attention.py
线性注意力 KernelAttentionKernelMask kernel_attention.py
注意力分数复用 ReuseMultiHeadAttentionReuseTransformer reuse_attention.pyreuse_transformer.py
深层收敛 ReZeroTransformer rezero_transformer.py
相对位置/长上下文 MultiHeadRelativeAttentionTwoStreamRelativeAttentionTransformerXLTransformerXLBlock relative_attention.pytransformer_xl.py
轻量 BERT MobileBertEmbeddingMobileBertTransformerMobileBertMaskedLM mobile_bert_layers.py
Transformer 块 Transformer(弃用)、CompiledTransformerTransformerDecoderBlockTransformerEncoderBlock transformer.pytransformer_encoder_block.py
掩码/Softmax SelfAttentionMaskMaskedSoftmax self_attention_mask.pymasked_softmax.py
位置编码 PositionEmbeddingRelativePositionBiasRelativePositionEmbedding position_embedding.py
任务头 MaskedLMClassificationHeadGaussianProcessClassificationHead masked_lm.pycls_head.py
不确定性/正则 RandomFeatureGaussianProcessSpectralNormalization gaussian_process.pyspectral_normalization.py
前馈/门控 GatedFeedforward gated_feedforward.py
排序 MatMulWithMargin mat_mul_with_margin.py
TPU 嵌入 OnDeviceEmbedding on_device_embedding.py
文本前处理 BertTokenizerSentencepieceTokenizerFastWordpieceBertTokenizerBertPackInputs text_layers.py

一个典型组合(对应仓库内 bert_encoder.py 的结构):BertTokenizer/BertPackInputs 产输入 → 词嵌入 + PositionEmbedding + segment 嵌入 → 堆叠 N 个 TransformerEncoderBlock(num_attention_heads=..., inner_dim=..., inner_activation="gelu")MaskedLM(预训练)或 ClassificationHead(微调)。每个层都带 get_config/from_config(如 reuse_attention.py 中保存 query/key/value 形状以便 from_config 触发重建),因此整条模型可以随 Keras 序列化保存与加载。

9. 小结

  • official/nlp/modeling/layers/ 是仓库 NLP 建模的最小抽象层:注意力、块、掩码、位置、任务头各自成文件、成层,测试一一对应(如 attention_test.pytransformer_test.py),便于单独验证与替换;
  • 面向性能演进提供了成体系的替代件:BigBirdAttention(稀疏)、KernelAttention(线性)、ReuseMultiHeadAttention(跨层复用)、OnDeviceEmbedding(TPU);
  • 面向生产与序列化:层均注册为 Keras 可序列化对象,TransformerEncoderBlock 是现行推荐块,旧 Transformer 已弃用、CompiledTransformer 提供自动编译加速。

理解本文内容后,你可以在 official/nlp/modeling/networks/official/nlp/modeling/models/ 中对照实际 encoder/模型实现,按需替换注意力或头结构,而不必重写整个训练管线。

登录后查看全文
热门项目推荐
相关项目推荐