PaddleSpeech 中 Transformer Mask 模块解析:subsequent_mask 与 target_mask 的实现与调用
【免费下载链接】PaddleSpeechEasy-to-use Speech Toolkit including Self-Supervised Learning model, SOTA/Streaming ASR with punctuation, Streaming TTS with text frontend, Speaker Verification System, End-to-End Speech Translation and Keyword Spotting. Won NAACL2022 Best Demo Award.项目地址: https://gitcode.com/gh_mirrors/pa/PaddleSpeech
导读
本篇技术指南聚焦 PaddleSpeech 仓库中paddlespeech.t2s.modules.transformer.mask模块(对应的 API 文档入口为 paddlespeech.t2s.modules.transformer.mask.rst,由 Sphinxautomodule指令自动生成文档)。该模块承载着 TTS 与 ASR 中 Transformer 系列模型最基础也最关键的组件——注意力掩码(Attention Mask)的生成逻辑。读完本文,你将掌握subsequent_mask与target_mask两个核心函数的数学语义、PaddlePaddle 实现细节,以及它们如何在 TransformerTTS 训练与自回归推理、解码器 beam search 中被实际调用,并能据此在自己的模型中正确构造掩码。
一、为什么 Transformer 需要 Mask:模块定位
Transformer 的自注意力(Self-Attention)在计算任意两个位置间的注意力权重时,默认允许"每个位置看到序列中的所有其他位置"。但在以下两类场景中必须人为切断这种可见性:
- 解码器自回归训练(Masked Self-Attention):解码器第 t 步只能看到第 1..t 步的输出,否则训练时会"偷看"未来帧,导致推理与训练不一致;
- 变长 batch 的 Padding 抑制:同一 batch 内序列长度不一,短序列被 pad 到统一长度,padding 位置不能参与注意力计算与损失统计。
mask.py模块正是为上述两类需求提供统一、可复用的张量生成函数。整个模块非常精简,仅有 30 余行(见 mask.py),包含两个公开函数:
| 函数 | 功能 | 返回形状 |
|---|---|---|
subsequent_mask(size, dtype) | 生成下三角掩码,用于自回归掩码(causal mask) | (size, size) |
target_mask(ys_in_pad, ignore_id, dtype) | 生成解码器自注意力掩码,同时覆盖 padding 抑制与因果性 | (B, Lmax, Lmax) |
二、subsequent_mask:两行代码构建因果掩码
2.1 源码实现
def subsequent_mask(size, dtype=paddle.bool): """Create mask for subsequent steps (size, size).""" ret = paddle.ones([size, size], dtype=dtype) return paddle.tril(ret)实现只用了两个 Paddle 算子:
paddle.ones([size, size], dtype=dtype):构造一个全 1 的方阵;paddle.tril(ret):取该方阵的下三角部分,右上角全部置 0。
2.2 语义与示例
函数 docstring 中给出了直观示例:
subsequent_mask(3) [[1, 0, 0], [1, 1, 0], [1, 1, 1]]即矩阵元素M[i][j] = 1当且仅当j <= i。在自注意力中,第 i 行代表"第 i 个位置可以 attend 到的位置集合",因此该掩码保证每个位置只能关注自身及其左侧(过去)的位置,正是标准 Transformer 论文中 decoder 的因果掩码。默认dtype=paddle.bool,可直接作为 attention 的布尔掩码使用;也可通过dtype参数指定为float32等类型,用于需要加权(如加性掩码-inf)的场合。
2.3 与 s2t 模块的同名实现对比
值得注意的是,在 ASR 侧的 s2t/modules/mask.py 中存在一个同名subsequent_mask(size),实现完全一致(同样基于paddle.tril)。这印证了因果掩码在 PaddleSpeech 的 TTS(paddlespeech.t2s)与 ASR(paddlespeech.s2t)两大技术栈中是通用的基础组件,且实现被刻意保持为最小、最直观的形式。
三、target_mask:padding 抑制与因果掩码的组合
3.1 源码实现
def target_mask(ys_in_pad, ignore_id, dtype=paddle.bool): """Create mask for decoder self-attention.""" ys_mask = ys_in_pad != ignore_id # (B, Lmax):True 表示非 padding m = subsequent_mask(ys_mask.shape[-1]).unsqueeze(0) # (1, Lmax, Lmax):因果掩码 return ys_mask.unsqueeze(-2) & m # (B, 1, Lmax) & (1, Lmax, Lmax) -> (B, Lmax, Lmax)3.2 分步拆解
- 第 1 步
ys_in_pad != ignore_id:ys_in_pad是 padding 后的目标序列(B, Lmax),ignore_id是 padding 索引。比较得到(B, Lmax)的布尔张量,True 表示该位置是真实 token、False 表示 padding; - 第 2 步
subsequent_mask(Lmax).unsqueeze(0):先生成(Lmax, Lmax)因果掩码,再在 batch 维插入一维变成(1, Lmax, Lmax),以便广播; - 第 3 步广播与运算
ys_mask.unsqueeze(-2) & m:ys_mask变成(B, 1, Lmax)后,与(1, Lmax, Lmax)按位与,自动广播为(B, Lmax, Lmax)。最终元素为 True 当且仅当:该行位置是真实 token,且列位置 <= 行位置——同时满足"非 padding"与"不看向未来"两个条件。
3.3 返回值语义
函数返回(B, Lmax, Lmax)的三维掩码,可直接作为 decoder.py 中Decoder.forward(tgt, tgt_mask, memory, memory_mask)的tgt_mask参数传入,用于屏蔽解码器自注意力。
四、在 TransformerTTS 中的完整调用链
4.1 训练路径:_target_mask
在 transformer_tts.py 的TransformerTTS._forward中:
y_masks = self._target_mask(olens_in) zs, _ = self.decoder(ys_in, y_masks, hs, h_masks)其中_target_mask的实现(transformer_tts.py)与模块中的target_mask思路一致,但输入换成了各样本的真实长度olens:
def _target_mask(self, olens): y_masks = make_non_pad_mask(olens) # (B, Lmax) s_masks = subsequent_mask(y_masks.shape[-1]).unsqueeze(0) # (1, Lmax, Lmax) return paddle.logical_and(y_masks.unsqueeze(-2), s_masks)make_non_pad_mask来自 nets_utils.py,根据长度生成非 padding 掩码(1 表示有效位置),等价于target_mask中的ys_in_pad != ignore_id;- 最终同样以"非 padding 掩码 & 因果掩码"的方式组合出
(B, Lmax, Lmax)的 decoder 自注意力掩码。
这里有一个值得注意的预处理细节:_forward中在送入解码器前执行了ys_in = self._add_first_frame_and_remove_last_frame(ys_in)(transformer_tts.py),即头部补一个全零帧、去掉最后一帧,保证自回归目标对齐。掩码的Lmax维度也随之对齐。
4.2 推理路径:逐步生成与掩码增长
在自回归推理inference中(transformer_tts.py),每一步都重新构造当前长度的因果掩码:
y_masks = subsequent_mask(idx).unsqueeze(0) z, z_cache = self.decoder.forward_one_step(ys, y_masks, hs, cache=z_cache)由于推理时每次只生成一帧,idx从 1 递增,掩码尺寸随之增长为(1, idx, idx),保证已生成的帧只能看到更早的帧。该路径配合forward_one_step(见 decoder.py)实现逐步解码。
4.3 解码器内部:beam search 中的掩码复用
subsequent_mask还被Decoder的评分接口复用(decoder.py):
def score(self, ys, state, x): ys_mask = subsequent_mask(len(ys)).unsqueeze(0) logp, state = self.forward_one_step(ys.unsqueeze(0), ys_mask, x.unsqueeze(0), cache=state)以及在批量 beam search 的batch_score中(decoder.py):
ys_mask = subsequent_mask(ys.shape[-1]).unsqueeze(0)这说明因果掩码不仅在训练时使用,在解码器的搜索阶段同样需要每步/每候选序列实时生成,是自回归解码的通用基础设施。
五、ASR 侧的掩码扩展:同一设计思想的进阶
虽然本模块面向 TTS,但理解掩码设计能帮助你快速读懂 ASR 侧的进阶变体。在 s2t/modules/mask.py 中,除了与subsequent_mask等价的实现外,还提供了一系列衍生工具:
| 函数 | 作用 |
|---|---|
make_pad_mask(lengths) | 根据长度生成 padding 位置掩码(1 表示 padding),见 mask.py |
make_non_pad_mask(lengths) | make_pad_mask的逻辑取反,1 表示有效位置 |
subsequent_chunk_mask(size, chunk_size, num_left_chunks) | 流式解码所需的 chunk 掩码,支持只看左侧有限 chunk(见 mask.py) |
add_optional_chunk_mask(...) | 训练时在全局注意力、动态 chunk、静态 chunk 之间切换的可选掩码(见 mask.py 起) |
mask_finished_scores/mask_finished_preds | 搜索中对已结束序列的打分与预测屏蔽 |
例如在流式 ASR 中,add_optional_chunk_mask被 s2t/modules/encoder.py 引入,用于平衡"全局上下文质量"与"流式低延迟"——这正是 t2s 侧因果掩码思想在时序约束更强场景下的延伸。上述函数共同构成了 PaddleSpeech 中"按长度/按 chunk/按因果性"三种维度裁剪注意力可见范围的完整工具集。
六、实践要点与常见坑
dtype 选择:
subsequent_mask默认返回paddle.bool,适合直接作为布尔掩码;若你的注意力实现要求-inf加性掩码(如score + mask),需先mask.astype('float32')再乘以一个大的负数。t2s 的MultiHeadedAttention使用布尔掩码直接屏蔽(可参考 attention.py),不必显式构造-inf矩阵。维度广播:
target_mask返回(B, Lmax, Lmax),而_source_mask(encoder 掩码,见 transformer_tts.py)返回(B, 1, Tmax)。两者形状不同是刻意的:encoder 自注意力是"全可见 + 非 padding",只需要二维信息;decoder 自注意力额外多出因果维度,需要三维。组装Decoder.forward(tgt, tgt_mask, memory, memory_mask)时务必保持两者维度正确。ignore_id 与 padding_idx 的一致性:
target_mask的ignore_id必须与词嵌入层的padding_idx(TransformerTTS 中为 0,见 transformer_tts.py)保持一致,否则掩码会错误屏蔽真实 token。reduction_factor 对齐:当
reduction_factor > 1时,TransformerTTS 会对目标序列做时间维抽帧,olens_in = olens // reduction_factor(见 transformer_tts.py),此时掩码长度必须以抽帧后的长度计算,否则会产生维度不匹配。
七、小结
paddlespeech.t2s.modules.transformer.mask模块虽然只有两个函数、代码不足 40 行,却是 PaddleSpeech 中所有 Transformer 系 TTS/ASR 模型自回归解码正确性的基石:subsequent_mask用paddle.tril一行算子实现因果可见性,target_mask通过布尔广播将 padding 抑制与因果性组合成解码器掩码。理解这两个函数,也就掌握了从 TransformerTTS 训练(_target_mask)到推理(subsequent_mask(idx)逐步增长)、再到解码器 beam search(score/batch_score)整条调用链的掩码逻辑,并能为阅读 ASR 侧subsequent_chunk_mask、add_optional_chunk_mask等流式扩展打下基础。相关可继续深入阅读的文件包括:mask.py、decoder.py、transformer_tts.py 与 s2t/modules/mask.py。
【免费下载链接】PaddleSpeechEasy-to-use Speech Toolkit including Self-Supervised Learning model, SOTA/Streaming ASR with punctuation, Streaming TTS with text frontend, Speaker Verification System, End-to-End Speech Translation and Keyword Spotting. Won NAACL2022 Best Demo Award.项目地址: https://gitcode.com/gh_mirrors/pa/PaddleSpeech
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考