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

1.2 RNN-T

Encoder
Predictor
Joint Network
1.3 U2/U2++-T


Transducer Decoder流式/非流式解码Nbest Attention Decoder对Nbest做Rescore

2.1 Encoder
2.2 Precitor
RNN base (GRU LSTM RNN) Embedding[3] Conv1D
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]]:
... # 原始输入 [你 好 _we net]
# 正常输入 [ 你 好 _we net]
# 替换blank [ 你 好 _we net]
# Predictor 输入 [ 你 好 _we net]
代码详见 https://github.com/wenet-e2e/wenet/blob/main/wenet/transducer/transducer.py#L90 2.2 Joint
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

3.1 非流模型
数据集 aishell
模型:conformer
Predictor: LSTM

3.2 Predcitor type
数据集 aishell
模型:conformer

数据集 aishell
模型:U2++
Predictor: LSTM
3.3 流式模型

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]
