目前,WeNet实验性加入了对RNN-T的支持。本文从模型架构、代码实现、部分实验结果等三方面来介绍。该工作由 WeNet 社区姚卓远、丁涵宇、张悦铠、周鼎皓等同学共同完成,目前还在进一步完善。


1. U2/U2++-T

1.1 回顾U2/U2++

WeNet采用U2/U2++模型,其结构为:


模型主要分为三部分:Shared Encoder、CTC Decoder、Attention Decoder。
在训练阶段,loss包含CTC Loss和Attention Decoder loss
在推理阶段,CTC Decoder可以流式/非流式解码出Nbest,Attention Decoder对Nbest做Rescore。

1.2 RNN-T

RNN-T结构为:
RNN-T
  • Encoder

  • Predictor

  • Joint Network

训练阶段:(不考虑降采样)
1 Encoder输出shape为[B,T,V]
2 Predictor输出shape为[B,U,V]
3 joint联合[B,T,V] [B,U,V],输出为[B,T,U,V]
RNNTLoss在该四维张量上计算损失
推理阶段:Predictor使用之前非Blank输出和Encoder的当前帧,Joint来进行解码,
有关训练和解码,详细可参考李宏毅rnnt[1]

1.3 U2/U2++-T

不难想象,在现有WeNet代码上支持RNN-T就落到Predictor和Joint的实现上
首先我们可以简化RNN-T结构图,其中Transducer Decoder内部负责Predictor、Encoder 和Joint的数据流。

那么WeNet版的Transducer呼之欲出:

U2/U2++-T
训练阶段,loss包含RNNTLoss和Attention Decoder loss
推理阶段:
  • Transducer Decoder流式/非流式解码Nbest
  • Attention Decoder对Nbest做Rescore
甚至,LAS、CTC、RNN-T可以理解为不同的对齐拓扑,在训练阶段,不考虑显存占用的情况下,我们可以把CTC Decoder也放进去,做更多的multi task learning。


2. 模型代码实现

2.1 Encoder

来自WeNet TransformerEncoder/ConformerEncoder

2.2 Precitor

虽然目前共识是Predictor的作用有限[2],但是我们也实现了三种类型的Predictor:
  • RNN base (GRU LSTM RNN)
  • Embedding[3]
  • Conv1D
Predictor 在推理时”自回归“特性:每一个step,需要之前的state和当前输入。但是在训练阶段不要中间state,所以我们实现了forward_step用于推理和jit导出,forward用于训练
class PredictorBase(torch.nn.Module):

    def __init__(self) -> None:
        super().__init__()

    def forward(
        self,
        input: torch.Tensor,
        cache: Optional[List[torch.Tensor]] = None,
    ):
      ...

    def forward_step(
            self, input: torch.Tensor, padding: torch.Tensor,
            cache: List[torch.Tensor]
    ) -> Tuple[torch.Tensor, List[torch.Tensor]]:
       ...
RNNTLoss要求文本前置个blank,我们把Predictor的输入做了些修改:
# 原始输入 [你 好 _we net]
# 正常输入 [<s> 你 好 _we net]
# 替换blank [<blank> 你 好 _we net]
# Predictor 输入 [<blank> 你 好 _we net]
代码详见 https://github.com/wenet-e2e/wenet/blob/main/wenet/transducer/transducer.py#L90
如果你有兴趣预训练Predictor,那么请将 替换成
Embedding 是个轻量级的Predictor,有兴趣可以参考下文献,目前我们没有实现tie embedding。

2.2 Joint

训练阶段,Encoder输出为[B,T,E], Predictor输出为[B,U,P] 这两个的维度未必是V(词汇表大小),所以这里我们做了简单的映射
if (self.prejoin_linear and self.enc_ffn is not None
                and self.pred_ffn is not None):
        enc_out = self.enc_ffn(enc_out)  # [B,T,E] -> [B,T,V]
        pred_out = self.pred_ffn(pred_out) # [B,U,P]-> [B,U,V]

2.3 Loss

torchaudio目前原生支持RNNTLoss



3. 部分实验结果

3.1 非流模型

  • 数据集 aishell

  • 模型:conformer

  • Predictor: LSTM


非流RNN-T,可以取得和旧模型一致的效果,并且在多任务的加持下,U2/U2++-T的ctc和rescore分支4.51也取得小幅提升(4.51<4.61)

3.2 Predcitor type

  • 数据集 aishell

  • 模型:conformer


  • 数据集 aishell

  • 模型:U2++

  • Predictor: LSTM

Embedding的Predictor可以取得和LSTM大致相同的结果,关于multi-head 和history的影响可以参考文献[2]

3.3 流式模型


流式模型并非当前最优,还需要社区小伙伴共同努力💪
更多实验结果请参考aishell[4]。


4. 补充
  • RNN-T是个很“费卡”的模型,但是得益于代码的一致性,我们可以很轻松的从一个预训练的WeNet模型初始化RNN-T:设置参数 --checkpoint
  • torchaudio 目前支持的RNNTLoss不太完备,比如FastEmit、[B,T,U,V]->[B,T,U,2]等,但是随着torchaudio的完善这些功能也会逐步加上去[5]。或者你可以集成第三方的RNNTLoss:HawkAaron / warp-transducer[6]、1ytic / warp-rnnt[7]、optimized_transducer[8]


5. 参考资料
[1]rnnt训练: https://www.bilibili.com/video/BV1me411x7DZ/,[2]ECHO STATE SPEECH RECOGNITION: https://arxiv.org/pdf/2102.09114.pdf,
[3]Tied & Reduced RNN-T Decoder: https://arxiv.org/pdf/2109.07513.pdf,
[4]aishell: https://github.com/wenet-e2e/wenet/tree/main/examples/aishell/rnnt,
[5]Fastemit: https://github.com/pytorch/audio/issues/2593,
[6]warp-transducer: https://github.com/HawkAaron/warp-transducer,
[7]1ytic / warp-rnnt: https://github.com/1ytic/warp-rnnt,
[8]optimized_transducer: https://github.com/csukuangfj/optimized_transducer,