前文介绍了端到端识别的基本概念,本文介绍Wenet中模型的设计与实现。
Wenet的代码借鉴了Espnet等开源实现,比较简洁,但是为了实现基于chunk的流式解码,以及处理batch内不等长序列,引入的一些实现技巧,比如cache和mask,使得多处的代码在初次阅读时不易理解,可在第一步学习代码时略过相关内容。
核心模型的代码位于wenet/transformer/目录
模型入口 ASRModel
wenet/transformer/asr_model.py
模型定义
使用pytorch Module构建神经网络时,在init中定义用到的子模块
class ASRModel(torch.nn.Module)
def __init__():
self.encoder = encoder
self.decoder = decoder
self.ctc = ctc
self.criterion_att = LabelSmoothingLoss(...)# AED的lossASRModel的init中定义了encoder, decoder, ctc, criterion_att几个基本模块。其整体网络拓扑如下图所示。
- encoder是Shared Encoder,其中也包括了Subsampling网络。
- decoder是Attention-based Decoder网络
- ctc是ctc Decoder网络(很简单,仅仅是前向网络和softmax)和ctc loss
- criterion_att是attention-based decoder的自回归似然loss,实际是一个LabelSmoothing的loss。

ASRModel中的模块又有自己的子模块,可以通过print打印出完整的模型结构。
model = ASRModel(...)
print(model)创建模型
def init_asr_model(config):该方法根据传入的config,创建一个ASRModel实例。 config内容由训练模型时使用的yaml文件提供。这个创建仅仅是构建了一个初始模型,其参数是随机的,可以通过model.load_state_dict(checkpoint)从训练好的模型中加载参数。
前向计算
pytorch框架下,只需定义模型的前向计算forword,。对于每个Module,可以通过阅读forward代码来学习其具体实现。
classASRModel(torch.nn.Module):
def forward()
...
# Encoder
encoder_out, encoder_mask = self.encoder(speech, speech_lengths)
encoder_out_lens = encoder_mask.squeeze(1).sum(1)
# Attention-decoder
loss_att, acc_att = self._calc_att_loss(encoder_out, encoder_mask,
text, text_lengths)
# CTC
loss_ctc = self.ctc(encoder_out, encoder_out_lens, text,text_lengths)
loss = self.ctc_weight * loss_ctc + (1 -self.ctc_weight) * loss_att
...其他接口
ASRModel除了定义模型结构和实现前向计算用于训练外,还有两个功能:
- 提供多种python的解码接口
- 提供runtime中需要使用的接口。
python解码接口
recognize() # attention decoder
attention_rescoring() # CTC + attention rescoring
ctc_prefix_beam_search() # CTC prefix beamsearch
ctc_greedy_search() # CTC greedy search用于Runtime的接口, 这些接口均有@torch.jit.export注解,可以在C++中调用
subsampling_rate()
right_context()
sos_symbol()
eos_symbol()
forward_encoder_chunk()
forward_attention_decoder()
ctc_activation()其中比较重要的是:
forward_attention_decoder()Attention Decoder的序列forward计算,非自回归模式。ctc_activation()CTC Decoder forward计算forward_encoder_chunk()基于chunk的Encoder forward计算
初学者可以先仅关注训练时使用的forward()函数。
Encoder网络
wenet/transformer/encoder.py
Wenet的encoder支持Transformer和Conformer两种网络结构,实现时使用了模版方法的设计模式进代码复用。BaseEncoder中定义了如下统一的前向过程,由TransformerEncoder,ConformerEncoder继承BaseEncoder后分别定义各自的self.encoders的结构。
class BaseEncoder(torch.nn.Module):
def forward(...):
xs, pos_emb, masks = self.embed(xs, masks)
chunk_masks = add_optional_chunk_mask(xs, ..)
for layer in self.encoders:
xs, chunk_masks, _ = layer(xs, chunk_masks, pos_emb, mask_pad)
if self.normalize_before:
xs = self.after_norm(xs)可以看到Encoder分为两大部分
- self.embed是Subsampling网络
- self.encoders是一组相同结构网络(Encoder Blocks)的堆叠
除了forward,Encoder还实现了两个方法,此处不展开介绍。
forward_chunk_by_chunk,python解码时,模拟流式解码模式基于chunk的前向计算。forward_chunk, 单次基于chunk的前向计算,通过ASRModel导出为forward_encoder_chunk()供runtime解码使用。
下面先介绍Subsampling部分,再介绍Encoder Block
Subsampling网络
wenet/transformer/subsampling.py
前文已经介绍了降采样或者降帧率的目的。这里不再重述。
语音任务里有两种使用CNN的方式,一种是2D-Conv,一种是1D-Conv:
- 2D-Conv: 输入数据看作是深度(通道数)为1,高度为F(Fbank特征维度,idim),宽度为T(帧数)的一张图.
- 1D-Conv: 输入数据看作是深度(通道数)为F(Fbank特征维度),高度为1,宽度为T(帧数)的一张图.
Kaldi中著名的TDNN就是是1D-Conv,在Wenet中采用2D-Conv来实现降采样。
Wenet中提供了多个降采样的网络,这里选择把帧率降低4倍的网络Conv2dSubsampling4来说明。
class Conv2dSubsampling4(BaseSubsampling):
def __init__(self, idim: int, odim: int, dropout_rate: float,
pos_enc_class: torch.nn.Module):
"""Construct an Conv2dSubsampling4 object."""
super().__init__()
self.conv = torch.nn.Sequential(
torch.nn.Conv2d(1, odim, 3, 2),
torch.nn.ReLU(),
torch.nn.Conv2d(odim, odim, 3, 2),
torch.nn.ReLU(),
)我们可以把一个语音帧序列[T,D]看作宽是T,高是D,深度为1的图像。Conv2dSubsampling4通过两个stride=2的2d-CNN,把图像的宽和高都降为1/4. 因为图像的宽即是帧数,所以帧数变为1/4.
torch.nn.Conv2d(1, odim, kernel_size=3, stride=2)
torch.nn.Conv2d(odim, odim, kernel_size=3, stride=2)具体的实现过程
defforward(...):
x = x.unsqueeze(1) # (b, c=1, t, f)
x = self.conv(x)
b, c, t, f = x.size()
x = self.out(x.transpose(1, 2).contiguous().view(b, t, c * f))
x, pos_emb = self.pos_enc(x, offset)
return x, pos_emb, x_mask[:, :, :-2:2][:, :, :-2:2]- x = x.unsqueeze(1) # (b, c=1, t, f) 增加channel维,以符合2dConv需要的数据格式。
- conv(x)中进行两次卷积,此时t维度约等于原来的1/4,因为没加padding,实际上是从长度T变为长度((T-1)/2-1)/2)。注意经过卷积后深度不再是1。
- view(b, t, c * f) 将深度和高度合并平铺到同一维,然后通过self.out()对每帧做Affine变换.
- pos_enc(x, offset) 经过subsampling之后,帧数变少了,此时再计算Positional Eembedding。
在纯self-attention层构建的网络里,为了保证序列的顺序不可变性而引入了PE,从而交换序列中的两帧,输出会不同。但是由于subsampling的存在,序列本身已经失去了交换不变性,所以其实PE可以省去。
x_mask是原始帧率下的记录batch各序列长度的mask,在计算attention以及ctc loss时均要使用,现在帧数降低了,x_mask也要跟着变化。
返回独立的pos_emb,是因为在relative position attention中,需要获取relative pos_emb的信息。在标准attention中该返回值不会被用到。

上下文依赖
注意Conv2dSubsampling4中的这两个变量。
self.subsampling_rate = 4
self.right_context = 6这两个变量都在asr_model中进行了导出,在runtime时被使用,他们的含义是什么?
在CTC或者WFST解码时,我们是一帧一帧解码,这里的帧指的是subsample之后的帧。我们称为解码帧,而模型输入的帧序列里的帧(subsample之前的)称为原始语音帧。
在图里可以看到
- 第1个解码帧,需要依赖第1到第7个原始语音帧。
- 第2个解码帧,需要依赖第5到第11个原始语音帧。
subsampling_rate: 对于相邻两个解码帧,在原始帧上的间隔。
right_context: 对于某个解码帧,其对应的第一个原始帧的右侧还需要额外依赖多少帧,才能获得这个解码帧的全部信息。
在runtime decoder中,每次会送入一组帧进行前向计算并解码,一组(chunk)帧是定义在解码帧级别的,在处理第一个chunk时,接受输入获得当前chunk需要的所有的context,之后每次根据chunk大小和subsampling_rate获取新需要的原始帧。比如,chunk_size=1,则第一个chunk需要1-7帧,第二个chunk只要新拿到8-11帧即可。
# runtime/core/decoder/torch_asr_decoder.cc
TorchAsrDecoder::AdvanceDecoding()
if (!start_) { // First chunk
int context = right_context + 1; // Add current frame
num_requried_frames = (opts_.chunk_size - 1) * subsampling_rate + context;
} else {
num_requried_frames = opts_.chunk_size * subsampling_rate;
}Encoder Block
wenet/transformer/encoder_layer.py
对于Encoder, Wenet提供了Transformer和Conformer两种结构,Conformer在Transformer里引入了卷积层,是目前语音识别任务效果最好的模型之一。 强烈建议阅读这篇文章The Annotated Transformer, 了解Transformer的结构和实现。
Transformer的self.encoders由一组TransformerEncoderLayer组成
self.encoders=torch.nn.ModuleList([
TransformerEncoderLayer(
output_size,
MultiHeadedAttention(attention_heads, output_size,
attention_dropout_rate),
PositionwiseFeedForward(output_size, linear_units,
dropout_rate), dropout_rate,
normalize_before, concat_after) for _ in range(num_blocks)
])Conformer的self.encoders由一组ConformerEncoderLayer组成
self.encoders = torch.nn.ModuleList([
ConformerEncoderLayer(
output_size,
RelPositionMultiHeadedAttention(*encoder_selfattn_layer_args),
PositionwiseFeedForward(*positionwise_layer_args),
PositionwiseFeedForward(*positionwise_layer_args)
if macaron_style else None,
ConvolutionModule(*convolution_layer_args)
if use_cnn_module else None,
dropout_rate,
normalize_before,
concat_after,
) for _ in range(num_blocks)
])仅介绍ConformerEncoderLayer,其涉及的主要模块有:
- RelPositionMultiHeadedAttention
- PositionwiseFeedForward
- ConvolutionModule
conformer论文中conformer block的结构如图

如果不考虑cache,使用normalize_before=True,feed_forward_macaron=True,则wenet中的ConformerEncoderLayer的forward可以简化为
class ConformerEncoderLayer(nn.Module):
def forward(...):
residual = x
x = self.norm_ff_macaron(x)
x = self.feed_forward_macaron(x)
x = residual + 0.5 * self.dropout(x)
residual = x
x = self.norm_mha(x)
x_att = self.self_attn(x, x, x, pos_emb, mask)
x = residual + self.dropout(x_att)
residual = x
x = self.norm_conv(x)
x, new_cnn_cache = self.conv_module(x, mask_pad, cnn_cache)
x = x + self.dropout(x)
residual = x
x = self.norm_ff(x)
x = self.feed_forward(x)
x = residual + 0.5 * self.dropout(x)
x = self.norm_final(x)可以看到,对于RelPositionMultiHeadedAttention,ConvolutionModule,PositionwiseFeedForward,都是前有Layernorm,后有Dropout,再搭配Residual。
Conformer Block - RelPositionMultiHeadedAttention
wenet/transformer/attention.py
attention.py中提供了两种attention的实现,MultiHeadedAttention和RelPositionMultiHeadedAttention。
MultiHeadedAttention用于Transformer。RelPositionMultiHeadedAttention用于Conformer。
原始的Conformer论文中提到的self-attention是Relative Position Multi Headed Attention,这是transformer-xl中提出的一种改进attention,和标准attention的区别在于,其中显示利用了相对位置信息,具体原理和实现可参考文章。Conformer ASR中的Relative Positional Embedding
注意,wenet中实现的Relative Position Multi Headed Attention是存在问题的, 但是由于采用正确的实现并没有什么提升,就没有更新成transformer-xl中实现。
Conformer Block - PositionwiseFeedForward
wenet/transformer/positionwise_feed_forward.py
PositionwiseFeedForward,对各个帧时刻输入均使用同一个矩阵权重去做前向Affine计算,即通过一个[H1, H2]的的前向矩阵,把[B, T, H1]变为[B,T,H2]。
Conformer Block - ConvolutionModule
wenet/transformer/convolution.py
ConvolutionModule结构如下
Wenet中使用了因果卷积(Causal Convolution),即不看右侧上下文,这样无论模型含有多少卷积层,对右侧的上下文都无依赖。
原始的对称卷积,如果不进行左右padding,则做完卷积后长度会减小。

因此标准的卷积,为了保证卷积后序列长度一致,需要在左右各pad长度为(kernel_size - 1) // 2的 0.
if causal: # 使用因果卷积
padding = 0 # Conv1D函数设置的padding长度
self.lorder = kernel_size - 1 # 因果卷积左侧手动padding的长度
else: # 使用标准卷积
# kernel_size should be an odd number for none causal convolution
assert (kernel_size - 1) % 2 == 0
padding = (kernel_size - 1) // 2 # Conv1D函数设置的padding长度
self.lorder = 0因果卷积的实现其实很简单,只在左侧pad长度为kernel_size - 1的0,即可实现。如图所示。
if self.lorder > 0:
if cache is None:
x = nn.functional.pad(x, (self.lorder, 0), 'constant', 0.0)Attention based Decoder网络
对于Attention based Decoder, Wenet提供了自回归Transformer和双向自回归Transformer结构。 所谓自回归,既上一时刻的网络输出要作为网络当前时刻的输入,产生当前时刻的输出。
在ASR整个任务中,Attention based Decoder的输入是当前已经产生的文本,输出接下来要产生的文本,因此这个模型建模了语言模型的信息。
这种网络在解码时,只能依次产生输出,而不能一次产生整个输出序列。
和Encoder中的attention层区别在于,Decoder网络里每层DecoderLayer,除了进行self attention操作(self.self_attn),也和encoder的输出进行cross attention操作(self.src_attn)
另外在实现上,由于自回归和cross attention,mask的使用也和encoder有所区别。
CTC Loss
wenet/transformer/ctc.py
CTC Loss包含了CTC decoder和CTC loss两部分,CTC decoder仅仅对Encoder做了一次前向线性计算,然后计算softmax.
# hs_pad: (B, L, NProj) -> ys_hat: (B, L, Nvocab)
ys_hat = self.ctc_lo(F.dropout(hs_pad, p=self.dropout_rate))
# ys_hat: (B, L, D) -> (L, B, D)
ys_hat = ys_hat.transpose(0, 1)
ys_hat = ys_hat.log_softmax(2)
loss = self.ctc_loss(ys_hat, ys_pad, hlens, ys_lens)
# Batch-size average
loss = loss / ys_hat.size(1)
return lossCTC loss的部分则直接使用的torch提供的函数torch.nn.CTCLoss.
self.ctc_loss = torch.nn.CTCLoss(reduction=reduction_type)Attention-based Decoder Loss
wenet/transformer/label_smoothing_loss.py
Attention-based Decoder的Loss是在最大化自回归的概率,在每个位置计算模型输出概率和样本标注概率的Cross Entropy。这个过程采用teacher forcing的方式,而不采用scheduled sampling。
每个位置上,样本标注概率是一个one-hot的表示,既真实的标注概率为1,其他概率为0. Smoothing Loss中,对于样本标注概率,将真实的标注概率设置为1-e,其他概率设为e/(V-1)。
网络的完整结构
通过print()打印出的ASRModel的网络结构。
ASRModel(
(encoder): ConformerEncoder(
(global_cmvn): GlobalCMVN()
(embed): Conv2dSubsampling4(
(conv): Sequential(
(0): Conv2d(1, 256, kernel_size=(3, 3), stride=(2, 2))
(1): ReLU()
(2): Conv2d(256, 256, kernel_size=(3, 3), stride=(2, 2))
(3): ReLU()
)
(out): Sequential(
(0): Linear(in_features=4864, out_features=256, bias=True)
)
(pos_enc): RelPositionalEncoding(
(dropout): Dropout(p=0.1, inplace=False)
)
)
(after_norm): LayerNorm((256,), eps=1e-12, elementwise_affine=True)
(encoders): ModuleList(
(0): ConformerEncoderLayer(
(self_attn): RelPositionMultiHeadedAttention(
(linear_q): Linear(in_features=256, out_features=256, bias=True)
(linear_k): Linear(in_features=256, out_features=256, bias=True)
(linear_v): Linear(in_features=256, out_features=256, bias=True)
(linear_out): Linear(in_features=256, out_features=256, bias=True)
(dropout): Dropout(p=0.0, inplace=False)
(linear_pos): Linear(in_features=256, out_features=256, bias=False)
)
(feed_forward): PositionwiseFeedForward(
(w_1): Linear(in_features=256, out_features=2048, bias=True)
(activation): Swish()
(dropout): Dropout(p=0.1, inplace=False)
(w_2): Linear(in_features=2048, out_features=256, bias=True)
)
(feed_forward_macaron): PositionwiseFeedForward(
(w_1): Linear(in_features=256, out_features=2048, bias=True)
(activation): Swish()
(dropout): Dropout(p=0.1, inplace=False)
(w_2): Linear(in_features=2048, out_features=256, bias=True)
)
(conv_module): ConvolutionModule(
(pointwise_conv1): Conv1d(256, 512, kernel_size=(1,), stride=(1,))
(depthwise_conv): Conv1d(256, 256, kernel_size=(15,), stride=(1,), groups=256)
(norm): LayerNorm((256,), eps=1e-05, elementwise_affine=True)
(pointwise_conv2): Conv1d(256, 256, kernel_size=(1,), stride=(1,))
(activation): Swish()
)
(norm_ff): LayerNorm((256,), eps=1e-12, elementwise_affine=True)
(norm_mha): LayerNorm((256,), eps=1e-12, elementwise_affine=True)
(norm_ff_macaron): LayerNorm((256,), eps=1e-12, elementwise_affine=True)
(norm_conv): LayerNorm((256,), eps=1e-12, elementwise_affine=True)
(norm_final): LayerNorm((256,), eps=1e-12, elementwise_affine=True)
(dropout): Dropout(p=0.1, inplace=False)
(concat_linear): Linear(in_features=512, out_features=256, bias=True)
)
...
)
)
(decoder): TransformerDecoder(
(embed): Sequential(
(0): Embedding(4233, 256)
(1): PositionalEncoding(
(dropout): Dropout(p=0.1, inplace=False)
)
)
(after_norm): LayerNorm((256,), eps=1e-12, elementwise_affine=True)
(output_layer): Linear(in_features=256, out_features=4233, bias=True)
(decoders): ModuleList(
(0): DecoderLayer(
(self_attn): MultiHeadedAttention(
(linear_q): Linear(in_features=256, out_features=256, bias=True)
(linear_k): Linear(in_features=256, out_features=256, bias=True)
(linear_v): Linear(in_features=256, out_features=256, bias=True)
(linear_out): Linear(in_features=256, out_features=256, bias=True)
(dropout): Dropout(p=0.0, inplace=False)
)
(src_attn): MultiHeadedAttention(
(linear_q): Linear(in_features=256, out_features=256, bias=True)
(linear_k): Linear(in_features=256, out_features=256, bias=True)
(linear_v): Linear(in_features=256, out_features=256, bias=True)
(linear_out): Linear(in_features=256, out_features=256, bias=True)
(dropout): Dropout(p=0.0, inplace=False)
)
(feed_forward): PositionwiseFeedForward(
(w_1): Linear(in_features=256, out_features=2048, bias=True)
(activation): ReLU()
(dropout): Dropout(p=0.1, inplace=False)
(w_2): Linear(in_features=2048, out_features=256, bias=True)
)
(norm1): LayerNorm((256,), eps=1e-12, elementwise_affine=True)
(norm2): LayerNorm((256,), eps=1e-12, elementwise_affine=True)
(norm3): LayerNorm((256,), eps=1e-12, elementwise_affine=True)
(dropout): Dropout(p=0.1, inplace=False)
(concat_linear1): Linear(in_features=512, out_features=256, bias=True)
(concat_linear2): Linear(in_features=512, out_features=256, bias=True)
)
...
)
)
(ctc): CTC(
(ctc_lo): Linear(in_features=256, out_features=4233, bias=True)
(ctc_loss): CTCLoss()
)
(criterion_att): LabelSmoothingLoss(
(criterion): KLDivLoss()
)
)