重访熵半环:entropy_semiring 中基于半环框架的 CTC/RNN-T 实现、熵正则化与蒸馏解析
【免费下载链接】google-researchGoogle Research项目地址: https://gitcode.com/gh_mirrors/go/google-research
本篇技术指南围绕 Google Research 的entropy_semiring模块展开,该模块是论文Revisiting the Entropy Semiring for Neural Speech Recognition(OpenReview)的配套开源代码,其核心贡献是把 CTC 与 RNN-T 两种主流神经语音识别(ASR)损失统一到"半环(Semiring)"这一代数框架下,并给出熵半环的示例实现,为熵正则化(regularization)与知识蒸馏(distillation)类应用提供基础。读完本文,你将掌握半环的数学公理与代码接口、三个具体半环(对数半环、对数熵半环、对数反向 KL 半环)的运算规则、CTC/RNN-T 动态规划如何被半环化、以及如何借助测试文件中的"手工枚举格"验证实现正确性。
一、项目定位:一篇论文的配套开源实现
entropy_semiring目录规模很小,只包含 8 个 Python 文件,没有配置文件与运行入口,属于典型的研究型"库 + 测试"结构:
entropy_semiring/ ├── README.md ├── semiring.py # 半环抽象基类与三个具体半环实现 ├── utils.py # 数值稳定的对数域工具函数 ├── asr_loss.py # CTC / RNN-T 损失及各自的半环化版本 ├── asr_divergence.py # 熵与反向 KL 散度的计算接口 ├── semiring_test.py # 半环代数公理测试 ├── asr_loss_test.py # CTC / RNN-T 手工枚举格测试 └── asr_divergence_test.py # 熵 / 散度手工公式测试README.md 明确声明了代码展示的两项内容:
- 半环框架中的 CTC 与 RNN-T:把两种 ASR 序列损失放进同一个代数框架,动态规划图保持不变,只替换"加法"与"乘法"两种运算;
- 熵半环的示例实现:该实现可用于正则化与蒸馏等应用场景。
同时 README 特别提示读者去查看测试文件中的"手工计算的小型 CTC 与 RNN-T 格(lattice)"示例,并验证代码输出与手工结果一致——这正是 asr_loss_test.py 与 asr_divergence_test.py 所承担的角色。
从代码依赖看,该模块运行在 Lingvo 生态内:所有文件均通过from lingvo import compat as tf引入 TensorFlow,asr_loss.py 还使用了lingvo.core.py_utils,semiring.py 与 utils.py 依赖tensorflow_probability(tfp);测试文件额外依赖absl.testing.parameterized与numpy。这意味着要运行本模块,需要先具备 Lingvo 与 TensorFlow Probability 环境。
二、半环抽象:七方法接口与四项代数公理
Semiring 抽象基类(semiring.Semiring)是全部实现的数学起点。一个半环是装备了两种二元运算的集合:
- 加法 (+)是带单位元 (0) 的交换幺半群;
- 乘法 (*)是带单位元 (1) 的幺半群;
- 乘法对加法满足左右分配律;
- 加法单位元 (0) 是乘法的零元(annihilator),即任何元素与 (0) 相乘仍得 (0)。
代码中,Semiring是一个abc.ABC泛型抽象类,具体子类必须实现七个方法:
| 方法 | 语义 | 典型实现思路 |
|---|---|---|
additive_identity(shape, dtype) | 加法单位元 (0) | 通常返回-inf常量 |
add(elem_1, elem_2) | 二元加法 | 对数域即 LogSumExp |
add_list(elems_list) | 列表加法 | 往往比逐个add更高效 |
multiplicative_identity(shape, dtype) | 乘法单位元 (1) | 通常返回0(log 空间) |
multiply(elem_1, elem_2) | 二元乘法 | 对数域即对数域加法 |
multiply_list(elems_list) | 列表乘法 | 批量累乘,需数值稳定 |
convert_logits(logits) | 把网络输出的 logits 转换成语义单元输入 | 如<logp, logq> -> <logp, log(-plogq)> |
类注释特别解释了为什么除了二元add/multiply还要提供add_list/multiply_list:列表版本通常存在比迭代调用二元运算更高效、更稳定的实现方式。convert_logits则是半环与神经网络输出之间的"适配层",把网络的 logits 转换成半环元素所需的各个分量。
三、三个具体半环的实现与运算规则
semiring.py 提供了三个具体的半环类,全部工作在 log 域以保证数值稳定,元素以"若干个同形状 Tensor 组成的元组"表示(文件顶部定义了LogTensor、DualTensor、LogReverseKLTensor等类型别名,见 semiring.py)。
3.1 LogSemiring:标准对数半环
元素形式为log(p),p 是 [0,1] 区间实数,即一条对齐路径的对数概率:
- 加法单位元:
-inf; - 加法:
a (+) b = LogSumExp(a, b); - 乘法单位元:
0(即 log(1)); - 乘法:
a (*) b = a + b; convert_logits:恒等变换——网络 logits 本身就是 log 空间的值。
实现见 LogSemiring,其add/add_list调用utils.logsumexp_list,multiply调用utils.safe_result把可能产生的inf兜底为-inf。它等价于经典的 CTC/RNN-T forward 算法:对所有可行对齐路径做 LogSumExp 求和,取负后即为负对数似然损失。
3.2 LogEntropySemiring:对数熵半环(本模块核心示例)
这是 README 点名的"熵半环示例实现"。每个元素是一个二元组<log(p), log(-p·log(q))>,其中 p、q 都是 [0,1] 概率;类注释指出它遵循"双数系统(dual number)"并在两个分量上施加 log 同态,参见 LogEntropySemiring。
运算规则如下(记<a,b> (+) <c,d>、<a,b> (*) <c,d>):
- 加法单位元:
<-inf, -inf>; - 加法:
<LogSumExp(a,c), LogSumExp(b,d)>; - 乘法单位元:
<0, -inf>(因为 log(-1·log(1)) = log(0) = -inf); - 乘法:
<a + c, LogSumExp(a + d, b + c)>——这正是utils.logcrossmultiply(a, b, c, d)的实现:利用恒等式-p1p2·log(q1q2) = (-p1·log q1)·p2 + p1·(-p2·log q2),把"乘积的熵分量"拆成两项 LogSumExp; convert_logits:<log(p), log(q)> -> <log(p), log(-p·log(q))>,第二分量由utils.logminus(logp, logq)计算。
关键洞察在于:当 p = q 时,第二分量log(-p·log(p))恰好就是对数熵log H(p)。因此 asr_divergence.py 的log_entropy_ctc会把同一份 logits 同时传给半环输入的两个分量(sr_inputs=(input_logits, input_logits)),从而让动态规划在求和所有对齐路径概率的同时,同步累积路径熵——一次前向即得到log(-Σp·logp)。
3.3 LogReverseKLSemiring:对数反向 KL 半环
为了支持两个模型之间的散度计算(蒸馏场景),本模块还实现了四元组半环,元素为<log(p), log(q), log(-q·log(q)), log(-q·log(p))>,参见 LogReverseKLSemiring:
- 加法单位元:
<-inf, -inf, -inf, -inf>; - 乘法单位元:
<0, 0, -inf, -inf>; - 乘法(记元素一为
<a,b,c,d>、元素二为<e,f,g,h>):<a+e, b+f, LogSumExp(b+g, c+f), LogSumExp(b+h, d+f)>。其推导与熵半环同理:-q1q2·log(q1q2) = (-q1·log q1)·q2 + q1·(-q2·log q2),-q1q2·log(p1p2) = (-q1·log p1)·q2 + q1·(-q2·log p2); convert_logits:<log(p), log(q)> -> <log(p), log(q), log(-q·log(q)), log(-q·log(p))>。
3.4 三个半环的对照表
| 半环 | 元素结构 | 加法 | 乘法 | 典型用途 |
|---|---|---|---|---|
LogSemiring | log(p) | LogSumExp | a + b | CTC / RNN-T 的 NLL 损失 |
LogEntropySemiring | <log p, log(-p·log q)> | 分量级 LogSumExp | <a+c, LSE(a+d, b+c)> | 熵正则化(取 p=q) |
LogReverseKLSemiring | <log p, log q, log(-q·log q), log(-q·log p)> | 分量级 LogSumExp | <a+e, b+f, LSE(b+g, c+f), LSE(b+h, d+f)> | 模型间反向 KL 散度(蒸馏) |
其中LSE表示 LogSumExp。三个半环的乘法都刻意保持数值稳定,相关处理集中在utils.safe_result与utils.logcrossmultiply中。
四、数值稳定的工具层:utils.py 的六个核心函数
半环的高阶运算全部建立在 utils.py 这组小而关键的函数之上:
logsumexp_list(tensor_list):对一组形状相同的 Tensor 沿新堆叠轴做tf.reduce_logsumexp,是半环add/add_list的底层实现;weightedlogsumexp_list(tensor_list, weights_list):带权重的 LogSumExp(基于tfp.math.reduce_weighted_logsumexp),用于计算q·log q - q·log p这类"两项之差"的散度公式;logcrossmultiply(a, b, c, d):计算LogSumExp(a + d, b + c),是熵半环与反向 KL 半环乘法中"交叉分量"的通用实现;logminus(logx, logy):由log(x)、log(y)计算log(-x·log(y))。注意其内部对log(y) >= 0(即 y ≥ 1,只有 y=1 时成立)的位置做了防护,避免log作用于非正数产生 NaN,并把结果中的inf统一为-inf;logzero(shape, dtype):返回log(0)即负无穷常量,作为加法单位元;safe_result(result):把结果中的inf替换为-inf,保证"0 与 -inf 相乘"这类边界情形不会产生NaN;tuple_to_list(x):把"元组的列表"转成"列表的元组",用于ctc_semiring的转移生成阶段。
这组函数体现了全模块的设计基调:一切运算都在 log 域进行,并用-inf表示概率 0,从而避免概率下溢。
五、CTC 与 RNN-T 的半环化:同一张 DP 图,不同的"加法"与"乘法"
asr_loss.py 是 README 第一项内容的落地:CTC 与 RNN-T 的动态规划(DP)图保持不变,只是把加法与乘法替换成半环定义的操作。这也解释了为什么需要ctc_semiring/rnnt_semiring这两个"通用半环版本"。
5.1 CTC:半环化的 forward 算法
标准 CTC 的输入输出约定如下(B 批大小、T 输入帧数、U 输出标签数、V 词表大小):input_logits形状[B, T, V],且约定词表第 0 个 token 是 blank;output_labels形状[B, U];两个序列长度均为[B]。
ctc_semiring(asr_loss.py)的关键步骤:
- 构建状态表:用
tf.one_hot把标签转成[B, U, V],经tf.einsum('buv, btv -> but')抽取每个标签在每帧上的得分得到[B, U, T],再用interleave_with_blank把 blank(state[:, :, :1]即词表 0 号位)插入相邻标签之间,得到[B, 2U+1, T]的 CTC 状态表; - 重复标签掩码:CTC 不允许直接跳过重复标签(如
AA必须拆成A, blank, A)。代码用相邻标签比较生成is_label_distinct位掩码(如AABBB -> TFTFF),插空后作用于"跳过两步"的转移; - 起始掩码:只允许从状态 0 与状态 1(开头 blank 或第一个标签)出发,其余状态在 t=0 时被掩成加法单位元
-inf; - 逐帧前向:
_generate_transitions产生三条转移——plus_zero(留在当前状态,即 blank 自环)、plus_one(前进 1 步)、plus_two(前进 2 步,跳过 blank,受重复标签掩码约束);_step内先sr.add_list聚合三条路径再sr.multiply乘上当前帧得分,外层用tf.scan沿时间轴迭代; - 收尾:在
input_seq_len - 1时刻取末态2*output_seq_len与次末态2*output_seq_len - 1相加,即标准 CTC 的结束条件;最后把inf的无效损失清零。
interleave_with_blank(asr_loss.py)把[..., U, ...]扩展为[..., 2U+1, ...](例如AAA -> bAbAbAb),asr_loss_test.py 的testInterleaveWithBlank分别沿axis=1与axis=0手工枚举验证了该函数。
5.2 RNN-T:对角 loop skewing 的实现
RNN-T 处理两条序列(s1 习惯上为输入、s2 为输出),每一步"消费/产出"两条序列中的一条,约定最后一个 token 必来自 s1。其 DP 方程为:
alpha[s1, s2] = alpha[s1-1, s2] * s1_logits[s1-1, s2] + alpha[s1, s2-1] * s2_logits[s1, s2-1] 边界条件: loss = alpha[S1, S2] * s1_logits[S1, S2]rnnt_semiring(asr_loss.py)的实现要点:
- 掩码:用
tf.sequence_mask(s1_seq_len)把超出长度的 s2 得分掩成加法单位元,避免越界帧参与计算; - Loop skewing:为把二维嵌套循环转成一维迭代,代码按对角方向 D = S1 + S2 - 1 把
[B, S1, S2]表"斜切"重排成[D, B, S2],_skew函数完成该重排(代码注释引用 Bagby、Rao、Sim 2018 年在 IEEE SLT 提出的 RNN-T 高效实现思路); - 迭代:
_step中分别计算alpha * s1_d与alpha * s2_d,其中 s2 分支需经_shift_down_s2下移一个时间步以对齐对角线结构,再sr.add求和;累加器初值由乘法单位元([B,1])拼接加法单位元([B,S2-1])构成; - 边界条件:在
(s1_seq_len + s2_seq_len - 1, s2_seq_len)处取值,完成"最后一个 token 来自 s1"的约定;长度为 0 的序列没有合法对齐路径,损失直接置 0。
5.3 包装函数:返回 NLL
ctc(asr_loss.py)与rnnt(asr_loss.py)是面向日常使用的薄包装:它们以LogSemiring调用通用版本,取第一个分量log_sum并返回-log_sum,即常规的负对数似然损失。半环化版本返回的是"元组"——普通损失只需要第一个分量,而熵与散度场景需要其余分量(详见下一节)。
六、熵正则与蒸馏:asr_divergence.py 的四种实用接口
README 指出熵半环"对正则化与蒸馏类应用有帮助",asr_divergence.py 正是这一声明的代码落点。该文件实现了四个函数,统一约定:第一个返回值总是 NLL(负对数似然),第二个返回值才是熵或散度(见文件头注释)。目前仅支持同构模型对:CTC 对 CTC、RNN-T 对 RNN-T。
log_entropy_ctc(input_logits, output_labels, input_seq_len, output_seq_len):以LogEntropySemiring运行ctc_semiring,两份输入均为同一份 logits,返回(-logp, log(-Σp·logp))。第二项即对数熵log H(p),可直接作为熵正则化(最大化/最小化输出分布的确定性)的目标项;log_entropy_rnnt(...):同样逻辑作用于 RNN-T;log_reverse_kl_ctc_ctc(input_logits_pair, ...):input_logits_pair携带两个模型(如教师与学生的 logits),以LogReverseKLSemiring运行,再用utils.weightedlogsumexp_list([log(-q·log q), log(-q·log p)], [-1.0, 1.0])计算log(Σ q·log(q/p)),即对数反向 KL 散度log KL(q‖p);log_reverse_kl_rnnt_rnnt(...):RNN-T 版本。
从源码可以清楚看到散度的推导:KL(q‖p) = Σ q·log(q/p) = Σ q·log q - Σ q·log p,而半环的四元组恰好累积了log(-q·log q)与log(-q·log p)两类跨路径求和项,最后的加权 LogSumExp 正是把"两项之差"安全地合并成一个对数标量。因此,把其中一个模型的 logits 视为教师、另一个视为学生,log_reverse_kl_*就构成了标准的反向 KL 蒸馏目标;把同一模型的 logits 传入熵半环,则得到熵正则化目标。测试文件 asr_divergence_test.py 注释补充了一个重要细节:实践中计算的是非归一化的熵与散度(判别式 ASR 模型不会是 delta 函数),仅在测试中硬编码归一化分布以便对拍。
七、手工枚举的测试格:如何验证实现正确性
README 反复强调"请查看测试文件中的手工计算格并验证输出"。三个测试文件从不同粒度对实现做了对拍验证。
7.1 代数公理测试:semiring_test.py
semiring_test.py 对三个半环统一做parameterized测试,逐个检查半环定义的四项公理(见 semiring_test.py):
- 加法是交换幺半群(交换律、单位元、结合律);
- 乘法是幺半群(单位元、结合律);
add_list与逐次add结果一致、multiply_list与逐次multiply结果一致;- 乘法对加法满足左右分配律;
- 加法单位元是乘法零元(湮灭律)。
这从代数层面保证了三个半环确实构成合法的半环。
7.2 CTC / RNN-T 手工格:asr_loss_test.py
testCTCByHand(asr_loss_test.py)构造了一个4×3的极小 logits 与标签[1, 2, 2],手工枚举出唯一合法路径(1, 2, blank, 2)(因为标签 2 重复,中间必须插入 blank),把 4 个 logits 之和做负 LogSumExp 后与asr_loss.ctc的输出对拍;随后又验证了两类边界:输入序列过短(无合法路径)时损失被清零为 0,以及输入序列未用尽的 logits 被正确掩码(此时手工枚举出 5 条路径)。
testRNNTByHand(asr_loss_test.py)对3×3的 s1/s2 logits 手工枚举出全部 6 条满足"最后 token 来自 s1"的对齐路径,逐条求和取负 LogSumExp 与asr_loss.rnnt对拍;并验证了 s1 或 s2 长度为零时损失置 0、以及把未使用位置 logits 篡改成1.23后结果不变(证明掩码生效)。
7.3 熵与散度手工公式:asr_divergence_test.py
asr_divergence_test.py 先手工枚举 CTC 的 5 条路径与 RNN-T 的 6 条路径,再由DivergenceFormula辅助函数按定义式-Σ p·log p、Σ q·(log q - log p)计算 log-熵与 log-反向 KL(见 asr_divergence_test.py),最后与log_entropy_ctc、log_entropy_rnnt、log_reverse_kl_ctc_ctc、log_reverse_kl_rnnt_rnnt四个函数的输出逐一assertAllClose对拍,容差atol=1e-37。由于测试数据被硬编码为归一化分布,NLL 输出恰好为 0,从而把注意力集中在熵与散度本身。
八、运行与验证方式
本模块没有 CLI 入口,属于库式代码。在具备 Lingvo(提供lingvo.compat.tf与lingvo.core.py_utils)、TensorFlow Probability 依赖的环境中,可以从entropy_semiring目录直接运行三个测试文件验证 README 所述的手工格:
python -m semiring_test python -m asr_loss_test python -m asr_divergence_test三个测试均以tf.test.main()收尾(见各测试文件末尾),测试通过即表明:三个半环满足全部代数公理,CTC/RNN-T 半环化损失与手工枚举路径一致,熵与反向 KL 散度与手工公式一致。在自己的 Lingvo 工程中使用时,可按需import asr_loss(取ctc/rnnt损失)或import asr_divergence(取熵/散度用于正则化与蒸馏),并参照测试中[B, T, V]、[B, U]、[B, S1, S2]的张量约定组织输入。
九、总结
entropy_semiring以极小的代码量示范了一个高信息密度的研究思路:把 CTC 与 RNN-T 的序列求和从"专用实现"提升为"半环参数化实现"——DP 图只写一遍,通过替换add/multiply即可同时获得 NLL 损失、熵与模型间散度。对数熵半环与对数反向 KL 半环分别对应正则化与蒸馏两类应用,而三个测试文件中的手工枚举格则为实现正确性提供了最直接的证据。对想要在 Lingvo 系 ASR 模型中引入熵正则或 KL 蒸馏的开发者而言,本模块是一份可直接借鉴与扩展的参考实现。
【免费下载链接】google-researchGoogle Research项目地址: https://gitcode.com/gh_mirrors/go/google-research
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考