TensorFlow Models NLP 建模层库详解:official/nlp/modeling/layers 中 Transformer 与注意力机制的源码级剖析
本文围绕 Layers 模块文档 展开,系统讲解该仓库 NLP 建模层库中 20 余个核心 Keras 层——多头注意力、稀疏/线性注意力、带缓存的解码器、掩码与位置编码、任务头(MaskedLM、分类头)等——的设计意图与源码实现。读完后,你将能够直接选用、组合这些层构建新的 tf.keras NLP 模型,并理解 BertEncoder、AlbertEncoder 等上层网络是如何调用这些积木搭出完整 Transformer 的。
1. Layers:NLP 模型的基本构建块
模块入口文档的定义非常直接(见 README):
Layers are the fundamental building blocks for NLP models. They can be used to assemble new
tf.keraslayers 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(混合专家)、FactorizedEmbedding、PackBertEmbeddings、TNTransformerExpandCondense、TransformerScaffold、PerDimScaleAttention、MultiQueryAttention等进阶实现,均位于 layers 目录 下。
下文的组织方式与 README 的条目一一对应,并按"注意力 → 块结构 → 掩码/Softmax → 位置编码 → 任务头 → 文本前处理"展开。
2. 核心注意力层
2.1 MultiHeadAttention 与 CachedAttention
README 指出 MultiHeadAttention 实现了可选掩码的 query/key/value 注意力("Attention Is All You Need"),当 from_tensor 与 to_tensor 相同时即自注意力。从源码看,attention.py 中的实现是一行别名:
MultiHeadAttention = tf_keras.layers.MultiHeadAttention
即直接复用 Keras 内置多头注意力;本仓库真正新增的是同文件中的 CachedAttention——一个继承自 tf_keras.layers.MultiHeadAttention 的"带缓存注意力层",专门服务自回归解码:
- 普通路径(
decode_loop_step is None):_update_cache用tf.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.py 的
TransformerDecoderBlock可以看到它的典型用法——构造参数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,即该实现面向的序列长度上限。模块同时导出 BigBirdMasks(init.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.py 与 reuse_transformer.py。其思想是:相邻层注意力分数高度冗余,高层可复用小一层的分数,省去部分点积计算。源码要点:
reuse_attention参数:0不复用;-1表示全部头复用(源码中会归一化为num_heads);中间值表示部分头复用,且构造时会校验取值在[-1, num_heads]内;- 当"复用头数 < 总头数"时,
_build_from_signature把 query/key/value/输出投影拆成value_reuse与value_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 专用层
MobileBertEmbedding 与 MobileBertTransformer 分别实现 MobileBERT 论文提出的轻量嵌入层与 Transformer 层。init.py 同时还导出了同文件中的 MobileBertMaskedLM,说明该目录对 MobileBERT 提供了"嵌入 + 块 + MLM 头"的成套组件。
3. Transformer 块:Transformer、TransformerDecoderBlock 与 TransformerEncoderBlock
3.1 Transformer(已弃用的别名层)
Transformer 直接继承 TransformerEncoderBlock 并转发全部参数(intermediate_size→inner_dim、dropout_rate→output_dropout 等),但构造时会发出明确的弃用告警:
"The
Transformerlayer is deprecated. Please directly useTransformerEncoderBlock."(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.py、bert_encoder.py、albert_encoder.py(ALBERT 的"参数共享"层即复用同一个 TransformerEncoderBlock 实例)、seq2seq_transformer.py 都在实例化该块。
3.3 TransformerDecoderBlock
TransformerDecoderBlock 是解码器单层,由三个子层构成(源码 docstring 与 build 一致):
- 自注意力:默认类为
attention.CachedAttention(可经self_attention_cls替换),因此天然带cache机制; - encoder-decoder 交叉注意力:默认
MultiHeadAttention,multi_channel_cross_attention=True时换成MultiChannelAttention,也可用cross_attention_cls显式指定; - 位置前馈网络:
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_maskreshape 为[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)
PositionEmbedding 的 build 按 max_length 创建嵌入表、按输入末维确定宽度;同文件还导出 RelativePositionBias 与 RelativePositionEmbedding(见 init.py),分别服务"相对位置偏置"与"相对位置嵌入"两条路线。
6. 任务头与不确定性建模
6.1 MaskedLM
MaskedLM 是 BERT 的掩码语言模型头,README 强调"它假设外部传入嵌入表变量"。源码印证了这一约定:
- 构造参数:
embedding_table(必须来自 encoder 的get_embedding_table())、activation、initializer='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_classes、cls_token_idx=0(在序列第几位做池化,通常取 [CLS] 位)、activation="tanh"、dropout_rate=0.0;内部结构为 pooler_dense(可选)→ Dropout → logits 输出层。call(features, only_project=False) 支持只取池化向量不接分类投影。
同文件的 GaussianProcessClassificationHead 是 SNGP(光谱归一化神经网络高斯过程)分类头,依赖 gaussian_process.py 的 RandomFeatureGaussianProcess(随机特征 GP,见 "Random Features for Large-Scale Kernel Machines")与 spectral_normalization.py 的 SpectralNormalization(tf.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 均在其中:
- BertTokenizer:WordPiece 分词 + 特殊 token 处理;
- SentencepieceTokenizer:SentencePiece 分词(配合 train_sentencepiece.py 训练出的模型文件);
- BertPackInputs:把分词结果按 BERT 的
input_word_ids / input_mask / segment_ids约定打包、padding 成模型输入张量。
这使得数据管道可以不落地 TFRecord,而是在图中直接完成分词与打包。
8. 组合使用:从 README 条目到真实模型
把 README 的条目按"积木分类"汇总如下(全部文件路径均位于 official/nlp/modeling/layers/):
| 类别 | 层 | 实现文件 |
|---|---|---|
| 基础注意力 | MultiHeadAttention |
attention.py |
| 解码缓存注意力 | CachedAttention |
attention.py |
| 头间交互 | TalkingHeadsAttention、MultiChannelAttention |
talking_heads_attention.py、multi_channel_attention.py |
| 长序列稀疏注意力 | BigBirdAttention、BigBirdMasks |
bigbird_attention.py |
| 线性注意力 | KernelAttention、KernelMask |
kernel_attention.py |
| 注意力分数复用 | ReuseMultiHeadAttention、ReuseTransformer |
reuse_attention.py、reuse_transformer.py |
| 深层收敛 | ReZeroTransformer |
rezero_transformer.py |
| 相对位置/长上下文 | MultiHeadRelativeAttention、TwoStreamRelativeAttention、TransformerXL、TransformerXLBlock |
relative_attention.py、transformer_xl.py |
| 轻量 BERT | MobileBertEmbedding、MobileBertTransformer、MobileBertMaskedLM |
mobile_bert_layers.py |
| Transformer 块 | Transformer(弃用)、CompiledTransformer、TransformerDecoderBlock、TransformerEncoderBlock |
transformer.py、transformer_encoder_block.py |
| 掩码/Softmax | SelfAttentionMask、MaskedSoftmax |
self_attention_mask.py、masked_softmax.py |
| 位置编码 | PositionEmbedding、RelativePositionBias、RelativePositionEmbedding |
position_embedding.py |
| 任务头 | MaskedLM、ClassificationHead、GaussianProcessClassificationHead |
masked_lm.py、cls_head.py |
| 不确定性/正则 | RandomFeatureGaussianProcess、SpectralNormalization |
gaussian_process.py、spectral_normalization.py |
| 前馈/门控 | GatedFeedforward |
gated_feedforward.py |
| 排序 | MatMulWithMargin |
mat_mul_with_margin.py |
| TPU 嵌入 | OnDeviceEmbedding |
on_device_embedding.py |
| 文本前处理 | BertTokenizer、SentencepieceTokenizer、FastWordpieceBertTokenizer、BertPackInputs |
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.py、transformer_test.py),便于单独验证与替换;
- 面向性能演进提供了成体系的替代件:
BigBirdAttention(稀疏)、KernelAttention(线性)、ReuseMultiHeadAttention(跨层复用)、OnDeviceEmbedding(TPU); - 面向生产与序列化:层均注册为 Keras 可序列化对象,
TransformerEncoderBlock是现行推荐块,旧Transformer已弃用、CompiledTransformer提供自动编译加速。
理解本文内容后,你可以在 official/nlp/modeling/networks/ 与 official/nlp/modeling/models/ 中对照实际 encoder/模型实现,按需替换注意力或头结构,而不必重写整个训练管线。
atomcodeClaude Code 的开源替代方案。连接任意大模型,编辑代码,运行命令,自动验证 — 全自动执行。用 Rust 构建,极致性能。 | An open-source alternative to Claude Code. Connect any LLM, edit code, run commands, and verify changes — autonomously. Built in Rust for speed. Get StartedRust0624
Hy4-previewHy4 preview 是由腾讯混元团队研发的新一代混合专家(MoE)旗舰模型。模型总参数量 770B,每个 token 激活 49B,主干共包含78层,第一层采用标准 FFN,其余 77 层均为 MoE 结构,每层包含 256 个路由专家与 1 个共享专家,每个 token 激活 top-8 路由专家及共享专家。主干之外原生内置 1 层 MTP(总参数量 10B,激活 0.7B)以支持投机解码。Python00
GLM-5.3GLM-5.3 与 GLM-5.2 使用相同的基座模型——所有提升均来自后训练。与 GLM-5.2 相比,它在复杂编程和长程任务上的表现显著提升。Jinja00
GLM-5.3-FlashGLM-5.3-Flash (320B-A18B),是GLM-5系列的首个原生多模态模型。320B总参数,能力超过GLM-5.2Jinja00
Spark-X2.5-4BSpark-X2.5-4B 旨在让强大的 AI 更实用、更高效、更易获得。在广泛日常任务中表现强劲,涵盖对话、写作、翻译、推理、编码、工具调用以及智能体工作流,并在同等规模的开源模型中取得领先成绩。Spark-X2.5 将面向效率的架构与最高 1M tokens 的原生上下文窗口相结合,并支持 200 多种语言。Python00
Spark-X2.5-1.7BSpark-X2.5-1.7B 旨在让强大的 AI 更加实用、高效且易于获取。这些模型在广泛的日常任务中表现出色,涵盖对话、写作、翻译、推理、编程、工具调用和智能体工作流,并在同等规模的开源模型中取得领先结果。Spark-X2.5 将面向效率的架构与最高 1M tokens 的原生上下文窗口相结合,并支持 200 多种语言。Python00