经超哥同意,接下来几篇文章将转载超哥写的《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

    • 联合建模

    • 神经网络类型

    • 流式语音识别

      第二章可参考

    https://mp.weixin.qq.com/s/a3Du45RQVJtv5h0wPqMW3g

  • 第2节: Wenet中的神经网络设计与实现

    • Subsampling网络

    • Encoder Block

    • 模型定义

    • 创建模型

    • 前向计算

    • 其他接口

    • 模型入口 ASRModel

    • Encoder网络

    • Attention based Decoder网络

    • CTC Loss

    • Attention based Decoder Loss

    • 网络的完整结构

      第三章可参考

    https://mp.weixin.qq.com/s/3eNj1wczdNycS6vf8BE0OQ

  • 第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

    • 其他

      本文讲解第四章

    第4节: 进阶话题:Cache

    Runtime流式解码

    Python流式解码

    BaseEncoder.forward_chunk()分析

      • offset

      • subsampling内部

      • subsampling_cache

      • elayers_output_cache

      • conformer_cnn_cache



进阶话题: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流式解码

asr_model.py中的forward_encoder_chunk()通过jit导出,用于C++ runtime,其内部使用了encoder.py中的forward_chunk()函数。
# 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:]