
重访熵半环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 RecognitionOpenReview的配套开源代码其核心贡献是把 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引入 TensorFlowasr_loss.py 还使用了lingvo.core.py_utilssemiring.py 与 utils.py 依赖tensorflow_probabilitytfp测试文件额外依赖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)二元加法对数域即 LogSumExpadd_list(elems_list)列表加法往往比逐个add更高效multiplicative_identity(shape, dtype)乘法单位元 (1)通常返回0log 空间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 bconvert_logits恒等变换——网络 logits 本身就是 log 空间的值。实现见 LogSemiring其add/add_list调用utils.logsumexp_listmultiply调用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)把乘积的熵分量拆成两项 LogSumExpconvert_logitslog(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,hae, bf, LogSumExp(bg, cf), LogSumExp(bh, df)。其推导与熵半环同理-q1q2·log(q1q2) (-q1·log q1)·q2 q1·(-q2·log q2)-q1q2·log(p1p2) (-q1·log p1)·q2 q1·(-q2·log p2)convert_logitslog(p), log(q) - log(p), log(q), log(-q·log(q)), log(-q·log(p))。3.4 三个半环的对照表半环元素结构加法乘法典型用途LogSemiringlog(p)LogSumExpa bCTC / RNN-T 的 NLL 损失LogEntropySemiringlog p, log(-p·log q)分量级 LogSumExpac, LSE(ad, bc)熵正则化取 pqLogReverseKLSemiringlog p, log q, log(-q·log q), log(-q·log p)分量级 LogSumExpae, bf, LSE(bg, cf), LSE(bh, df)模型间反向 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只有 y1 时成立的位置做了防护避免log作用于非正数产生 NaN并把结果中的inf统一为-inflogzero(shape, dtype)返回log(0)即负无穷常量作为加法单位元safe_result(result)把结果中的inf替换为-inf保证0 与 -inf 相乘这类边界情形不会产生NaNtuple_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 是 blankoutput_labels形状[B, U]两个序列长度均为[B]。ctc_semiringasr_loss.py的关键步骤构建状态表用tf.one_hot把标签转成[B, U, V]经tf.einsum(buv, btv - but)抽取每个标签在每帧上的得分得到[B, U, T]再用interleave_with_blank把 blankstate[:, :, :1]即词表 0 号位插入相邻标签之间得到[B, 2U1, T]的 CTC 状态表重复标签掩码CTC 不允许直接跳过重复标签如AA必须拆成A, blank, A。代码用相邻标签比较生成is_label_distinct位掩码如AABBB - TFTFF插空后作用于跳过两步的转移起始掩码只允许从状态 0 与状态 1开头 blank 或第一个标签出发其余状态在 t0 时被掩成加法单位元-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_blankasr_loss.py把[..., U, ...]扩展为[..., 2U1, ...]例如AAA - bAbAbAbasr_loss_test.py 的testInterleaveWithBlank分别沿axis1与axis0手工枚举验证了该函数。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_semiringasr_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 包装函数返回 NLLctcasr_loss.py与rnntasr_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-Tlog_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.pysemiring_test.py 对三个半环统一做parameterized测试逐个检查半环定义的四项公理见 semiring_test.py加法是交换幺半群交换律、单位元、结合律乘法是幺半群单位元、结合律add_list与逐次add结果一致、multiply_list与逐次multiply结果一致乘法对加法满足左右分配律加法单位元是乘法零元湮灭律。这从代数层面保证了三个半环确实构成合法的半环。7.2 CTC / RNN-T 手工格asr_loss_test.pytestCTCByHandasr_loss_test.py构造了一个4×3的极小 logits 与标签[1, 2, 2]手工枚举出唯一合法路径(1, 2, blank, 2)因为标签 2 重复中间必须插入 blank把 4 个 logits 之和做负 LogSumExp 后与asr_loss.ctc的输出对拍随后又验证了两类边界输入序列过短无合法路径时损失被清零为 0以及输入序列未用尽的 logits 被正确掩码此时手工枚举出 5 条路径。testRNNTByHandasr_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.pyasr_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对拍容差atol1e-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),仅供参考