经超哥同意,接下来几篇文章将转载超哥写的《Wenet网络设计与实现》。感兴趣的读者可以关注知乎杨超,同时wenet1.0正式发布( wenet1.0正式发布:更快更高更强更有生产力https://mp.weixin.qq.com/s/CG6g1TBSvdVRYp_9EqanNA),欢迎大家关注,具体的代码
https://github.com/wenet-e2e/wenet
第一章节可参考
https://mp.weixin.qq.com/s/5JbPrWYkbft2OQ-_vzNVSw
第1节: 端到端语音识别基础
CTC目标函数
Attention-based Encoder Decoder
联合建模
神经网络类型
流式语音识别
第二章可参考
第2节: Wenet中的神经网络设计与实现
Subsampling网络
Encoder Block
模型定义
创建模型
前向计算
其他接口
模型入口 ASRModel
Encoder网络
Attention based Decoder网络
CTC Loss
Attention based Decoder Loss
网络的完整结构
第三章可参考
第3节: 进阶话题:Mask
Subsampling中的mask
Conformer Block中的Conv的mask
MultiHeadedAttention Module的Mask实现
Chunk-based mask
处理Padding对Loss的影响
处理模型输入Padding
问题1:Batch Padding
问题2: 自回归
问题3: Chunk-Based Model
Encoder中的mask
Decoder中的mask
其他
本文讲解第四章
offset
subsampling内部
subsampling_cache
elayers_output_cache
conformer_cnn_cache
https://mp.weixin.qq.com/s/a3Du45RQVJtv5h0wPqMW3g
https://mp.weixin.qq.com/s/3eNj1wczdNycS6vf8BE0OQ
第4节: 进阶话题:Cache
Runtime流式解码
Python流式解码
BaseEncoder.forward_chunk()分析
进阶话题:Cache
标准的forward是整个序列进行计算,但是在流式推断时,需要chunk级别的forward,因此需要引入cache的概念,即当前chunk的进行前向计算时,需要拿到上次前向的一些结果作为输入。
什么是cache?
对于流式推断,输入是一个个chunk的到来,对第i个chunk,当计算第k层网络的输出时,由于网络结构存在对左侧上下文的依赖,需要依赖第k-1层网络里在i之前的一些chunks的输出。如果对于当前到来chunk,将其和依赖的chunk序列(比如10层self-attention层,每层依赖左侧4个chunk,则累积起来需要依赖左侧40个chunk)拼起来作为网络输入进行前向,其计算量会比较大。对于那些已经计算过的chunk,可以将那些在计算下一个chunk的输出时需要的中间量保存下来,从而减少重复计算。这种方式就叫cache。
另外,wenet的网络在设计时,对于因果卷积和self-attention的左侧上下文都使用有限长度,因此无论序列多长,每次cache的大小是不变的(不增长)。
仅仅encoder部分涉及chunk计算时的cache。
对于CTC decoder,由于是线性层,不需要cache。
对于AED decoder,是在计算完整个序列的encoder输出后进行rescoring,不涉及chunk。
Runtime流式解码
# wenet/transformer/asr_model.py
@torch.jit.export
def forward_encoder_chunk() Python流式解码
如果设置simulate_streaming为True,则会模拟runtime流时解码的过程,将数据分成chunk,依次进行前向计算。该方法的结果,和送入整个序列通过mask进行流式模拟的结果应该是一致的。
recognize() -> _forward_encoder() -> BaseEncoder.forward_chunk_by_chunk() forward_chunk_by_chunk()的内部也是使用的forward_chunk()函数。
BaseEncoder.forward_chunk()分析
forward_chunk()是对单个chunk进行前向计算的核心函数。下面从该函数的内容来了解cache的实现。
# wenet/transformer/encoder.py
def forward_chunk(
self,
xs: torch.Tensor,
offset: int,
required_cache_size: int,
subsampling_cache: Optional[torch.Tensor] = None,
elayers_output_cache: Optional[List[torch.Tensor]] = None,
conformer_cnn_cache: Optional[List[torch.Tensor]] = None,
) -> Tuple[torch.Tensor, torch.Tensor, List[torch.Tensor],
List[torch.Tensor]]: xs是当前的chunk输入,由于对于单个chunk的前向计算,需要之前的chunk的计算得到的信息,因此这里需要传入相关的三个cache信息。
subsampling_cache:torch.Tensorsubsampling的输出的cache。即第一个conformer block的输入。
elayers_output_cache:List[torch.Tensor]第1个到最后1个conformer block的输出的cache。也就是第2个conformer block的输入和CTC层的输入。
conformer_cnn_cache:List[torch.Tensor]conformer block里的conv层的左侧依赖的输入cache。
cache的大小
subsampling_cache和elayers_output_cache的大小 由self-attention是对左侧的依赖长度required_cache_size决定。decoding_chunk_size是解码帧级别的chunk大小, num_decoding_left_chunks是self-attention依赖的左侧chunk数。
required_cache_size = decoding_chunk_size * num_decoding_left_chunks conformer_cnn_cache的大小和required_cache_size无关,由casual网络的左侧上下文lorder决定。
函数返回了四个值,包括当前chunk输入对应的输出,更新后的三个cache。
该函数的整个计算过程请参考下图

offset
当按chunk进行输入时,不能直接得到chunk在序列中的位置,需要传入offset给出该chunk在整个序列里的偏移,用于计算positional encoding。
xs, pos_emb, _ = self.embed(xs, tmp_masks, offset) subsampling内部
subsampling内部的计算虽然存在冗余,但是不进行cache。一个是其实现比较复杂,另一个原因是subsampling的计算量占比不大。
subsampling_cache
subsampling的输出的cache。即第一个conformer block的输入。
if subsampling_cache is not None:
cache_size = subsampling_cache.size(1)
# xs是第一个conformer block的输入
xs = torch.cat((subsampling_cache, xs), dim=1)
else:
cache_size = 0
pos_emb = self.embed.position_encoding(offset - cache_size, xs.size(1))
if required_cache_size < 0:
next_cache_start = 0
elif required_cache_size == 0:
next_cache_start = xs.size(1)
else:
next_cache_start = max(xs.size(1) - required_cache_size, 0)
# 更新subsampling_cache
r_subsampling_cache = xs[:, next_cache_start:, :] elayers_output_cache
第1个到最后1个conformer block的输出的cache。也就是第2个conformer block的输入和CTC层的输入。
for i, layer in enumerate(self.encoders):
attn_cache = elayers_output_cache[i]
cnn_cache = conformer_cnn_cache[i]
xs, _, new_cnn_cache = layer(xs,
masks,
pos_emb,
output_cache=attn_cache,
cnn_cache=cnn_cache)
# 更新elayers_output_cache
r_elayers_output_cache.append(xs[:, next_cache_start:, :]) 注意,此处的xs不是当前的chunk,而是当前chunk+cache输入,所以其长度不是chunk_size, 而是chunk_size + required_cache_size。
# wenet/transformer/encoder.py BaseEncoder.forward_chunk()
# 第一个conformer block输入的xs
xs = torch.cat((subsampling_cache, xs), dim=1)
# wenet/transformer/encoder_layer.py ConformerEncoderLayer.forward()
# 之后的conformer block输入的xs
if output_cache is not None:
x = torch.cat([output_cache, x], dim=1) layer()对应着wenet/transformer/encoder_layer.py中的ConformerEncoderLayer.forward()。下面是其具体过程。
# 计算feedforwad/res/norm(包含当前chunk和左侧num_decoding_left_chunks个chunk)
# 使用cache时,只要计算当前chunk x_q的self-attentionattention和residual
chunk = x.size(1) - output_cache.size(1)
x_q = x[:, -chunk:, :]
# 只选择当前chunk对应的部分做residual计算
residual = residual[:, -chunk:, :]
# 选取当前chunk对应的mask,
mask = mask[:, -chunk:, :]
# 使用当前chunk的x_q去和其依赖的x做attention
x = residual + self.dropout(self.self_attn(x_q, x, x, mask))
# 仅计算计算当前chunk的conv
x, new_cnn_cache = self.conv_module(x, mask_pad, cnn_cache)
# 仅计算当前chunk的feedforwad/res/norm
x = self.norm2(x)
x = residual + self.dropout(self.feed_forward(x))
# 可以看到通过cache节省了x[:, :-chunk, :]部分的attention/conv以及之后的feedforwad/res/norm计算
# chunk的输出和cache拼在一起,作为网络的最终输出。
x = torch.cat([output_cache, x], dim=1) 注意,self-attention之前的一些前向计算其实仍然存在冗余,如果对attention层的输入进行cache,而不是对conformer block层的输入cache,可以进一步降低计算量。
conformer_cnn_cache
conformer block里的conv层的左侧依赖的输入cache。
conformer_cnn_cache大小为lorder,即因果卷积左侧依赖,。
# wenet/transformer/encoder_layer.py ConformerEncoderLayer.forward()
# conformer_cnn_cache通过ConvolutionModule.forward()返回的新cache来更新
x, new_cnn_cache = self.conv_module(x, mask_pad, cnn_cache) # wenet/transformer/convolution.py ConvolutionModule.forward()
if self.lorder > 0:
if cache is None:
x = nn.functional.pad(x, (self.lorder, 0), 'constant', 0.0)
else:
x = torch.cat((cache, x), dim=2)
# 更新 conformer_cnn_cache
new_cache = x[:, :, -self.lorder:] 