启动脚本

第一个inference的算法 attention rescoring,就在上面part -6 学习了。

本次学习第二个inference算法:ctc greedy search。

这个只需要在运行脚本中增加:

decode_modes="ctc_greedy_search"即可。


具体的解码inference的脚本

recognize.py

[wenet/bin/recognize.py]

args

Namespace(batch_size=1,beam_size=10,bpe_model=None,
checkpoint='exp/sp_spec_aug_conformer_bidecoder_large/84.pt', 
config='exp/sp_spec_aug_conformer_bidecoder_large/train.yaml', 
ctc_weight=0.5, data_type='raw', decoding_chunk_size=-1, 
dict='data/lang_char/train_bpe4096_units.txt', gpu=0, 
mode='ctc_greedy_search', non_lang_syms=None, 
num_decoding_left_chunks=-1, override_config=[], penalty=0.0, 
result_file='exp/sp_spec_aug_conformer_bidecoder_large/test1_ctc_greedy_search/text_bpe', 
reverse_weight=0.0, simulate_streaming=False, test_data='data/test1/data.list')

这个inference算法,支持batch_size>1的情况!

test_conf

{'batch_conf':{'batch_size':12,'batch_type':'static'},

'fbank_conf': {'dither': 1.0, 'frame_length': 25, 'frame_shift': 10, 'num_mel_bins': 80}, 

'filter_conf': {'max_length': 2000, 'max_output_input_ratio': 10.0, 'min_length': 50,
 'min_output_input_ratio': 0.05, 'token_max_length': 400, 'token_min_length': 1}, 
'resample_conf': {'resample_rate': 16000}, 'shuffle': True, 
'shuffle_conf': {'shuffle_size': 1500}, 'sort': True, 
'sort_conf': {'sort_size': 500}, 'spec_aug': True, 
'spec_aug_conf': {'max_f': 10, 'max_t': 50, 'num_f_mask': 2, 'num_t_mask': 3}, 
'speed_perturb': True}

整体逻辑:


整体逻辑图,ctc-greedy-search和attention-rescoring相同

上图给出了几个步骤:

读取配置文件,构造测试数据dataset和dataloader,导入unigram词典,导入已有的训练好的checkpoint,以及最重要的解码工作。

每个batch下的操作


每个batch下的解码操作

为了方便学习,我这里分别设置了batch-size=1和batch-size=2两种情况,看代码是如何跑的。

核心是调用model.ctc_greedy_search()算法。

ctc_greedy_search()


ctc_greedy_search算法的脑图,重要的点是三个

上面给出了ctc_greedy_search算法的脑图,重要的点是三个:

其一,对输入wav frame的基于conformer encoder layers的编码,例如12层;

其二,使用一个线性层,512 -> 5502,为每个frame的原本的512维度向量,映射到词表,从而为每个frame获取候选词的概率,并对数化;

其三,对候选搞个topk,这个就是greedy的了,没有beam search啥事情了。每个frame都要最好的那个候选,即可。然后就是收集结果:去掉候选序列中的blank(token.id=0),以及如果是连续的重复的token,只要一个即可。

关于这个算法的截屏:


ctc_greedy_search的代码截屏,重要的是两点:编码器(12 conformerEncoderLayers,以及ctc.log_softmax把每个frame的512维度向量映射到5502词表,以及概率对数化。

上面给出的是:ctc_greedy_search的代码截屏,重要的是两点:编码器(12 conformerEncoderLayers,以及ctc.log_softmax把每个frame的512维度向量映射到5502词表,以及概率对数化。

_forward_encoder

[wenet/transformer/asr_model.py] 这个已经过了好几遍了,这里不冗述:


_forward_encoder 涉及的脑图

返回的是两个东西:

其一,encoder_out,[1=batch-size, 50=frame-num, 512=dimension of representation]

其二,encoder_mask,为frame num长度mask的东西,如果batch-size=1,则这个没啥用(都是true)。【即如果一个batch中,有序列的长度短,那么不足的部分,就填充一些padding id】

ctc.log_softmax


调用一个线性层,并搞下softmax 以及log

这个函数,在讲前一个attention-rescoring解码方法的时候,也用到了。

核心就是一个512 -> 5502的线性层,并搞一下softmax,概率化,然后log。

top-1 greedy

核心思想:每个frame只要概率最大的一个候选,然后每个frame的最大概率的候选,放一起,就是一个文本序列


每个frame只要概率最大的一个候选,然后每个frame的最大概率的候选,放一起,就是一个文本序列

topk示例

0=blank

ipdb>topk_index
tensor([[   0,    0,    0,    0, 1396,    0,    0,    0,    0,    0,    0,    0,
            0,    0,    0,    0,    0,    0,    0,    0,    0,    0,    0,    0,
            0,    0,    0,    0,    0, 1845,    0,    0,    0, 1396,    0,    0,
            0,    0,    0,    0,    0,    0,    0, 1885, 1670,    0,    0,    0,
            0,    0]], device='cuda:0')
ipdb> topk_prob
tensor([[[-0.9934],
         [-0.1779],
         [-0.2814],
         [-0.1027],
         [-1.0185],
         [-0.0290],
         [-1.1119],
         [-0.0324],
         [-0.8022],
         [-0.0493],
         [-0.9088],
         [-0.0329],
         [-0.7611],
         [-0.0296],
         [-0.7114],
         [-0.0318],
         [-0.7456],
         [-0.0733],
         [-0.5396],
         [-0.3997],
         [-0.4009],
         [-0.3177],
         [-0.2930],
         [-0.3159],
         [-0.2080],
         [-0.4096],
         [-0.1676],
         [-0.2456],
         [-0.1270],
         [-1.1715],
         [-0.0319],
         [-0.7106],
         [-0.0777],
         [-0.5603],
         [-0.0709],
         [-0.3034],
         [-0.3249],
         [-0.2274],
         [-0.6441],
         [-0.3882],
         [-0.3631],
         [-0.4921],
         [-0.7564],
         [-1.1378],
         [-0.7769],
         [-0.0975],
         [-0.0890],
         [-0.1064],
         [-0.1579],
         [-0.0545]]], device='cuda:0')

最后就是去除blank,以及对于连续的相同的候选词,只留一个的操作了。

batch-size=1的情况,有点简单,哈哈,不太过瘾。

继续看看batch-size=2的情况。

batch-size=2

一个batch的内容


speech (frame) length和text length长度都不一样,两者都需要mask

encoder的输出:

2022-03-1405:56:45,794DEBUGUsingselector:EpollSelector
> /workspace/asr/wenet/wenet/transformer/asr_model.py(178)_forward_encoder()
    177         import ipdb; ipdb.set_trace()
--> 178         return encoder_out, encoder_mask
    179

2022-03-14 05:56:45,890 DEBUG Using selector: EpollSelector
ipdb> encoder_out.shape
torch.Size([2, 50, 512])
ipdb> encoder_mask
tensor([[[ True,  True,  True,  True,  True,  True,  True,  True,  True,  True,
           True,  True,  True,  True,  True,  True,  True,  True,  True,  True,
           True,  True,  True,  True,  True,  True,  True,  True,  True,  True,
           True,  True,  True,  True,  True,  True,  True,  True,  True,  True,
           True,  True,  True,  True,  True,  True,  True,  True,  True,  True]],

        [[ True,  True,  True,  True,  True,  True,  True,  True,  True,  True,
           True,  True,  True,  True,  True,  True,  True,  True,  True,  True,
           True,  True,  True,  True,  True,  True,  True,  True,  True,  True,
           True,  True,  True,  True,  True,  True,  True,  True,  True,  True,
           True,  True,  True,  True,  True,  True,  True,  True, False, False]]],
       device='cuda:0')
ipdb>

上面可以看到,encoder对两个wav操作之后,分别得到长度为50和48(虽然,因为卷积等操作,它们已经不是真正的frame了,不过简单期间,还是以'frame'称呼它们。。。)的两个序列。

目前的位置:

ipdb>encoder_out_lens
tensor([50, 48], device='cuda:0')
ipdb> n
> /workspace/asr/wenet/wenet/transformer/asr_model.py(321)ctc_greedy_search()
    320         ctc_probs = self.ctc.log_softmax(
--> 321             encoder_out)  # (B, maxlen, vocab_size)
    322         topk_prob, topk_index = ctc_probs.topk(1, dim=2)  # (B, maxlen, 1)

ipdb>
> /workspace/asr/wenet/wenet/transformer/asr_model.py(320)ctc_greedy_search()
    319         encoder_out_lens = encoder_mask.squeeze(1).sum(1)
--> 320         ctc_probs = self.ctc.log_softmax(
    321             encoder_out)  # (B, maxlen, vocab_size)

ipdb>
> /workspace/asr/wenet/wenet/transformer/asr_model.py(322)ctc_greedy_search()
    321             encoder_out)  # (B, maxlen, vocab_size)
--> 322         topk_prob, topk_index = ctc_probs.topk(1, dim=2)  # (B, maxlen, 1)
    323         topk_index = topk_index.view(batch_size, maxlen)  # (B, maxlen)

ipdb> ctc_probs.shape
torch.Size([2, 50, 5502])

ctc_probs

的形状为[2, 50, 5502]。

然后还是322行,这个topk(1, dim=2)的贪心搜索!

看下它们的值:

ipdb>topk_prob
tensor([[[-0.9934],
         [-0.1779],
         ...
         [-0.0545]],

        [[-0.9958],
         [-0.1771],
         ...
         [-0.1645],
         [-0.1645]]], device='cuda:0')
ipdb> topk_prob.shape
torch.Size([2, 50, 1])
ipdb> topk_index.shape
torch.Size([2, 50, 1])
ipdb> topk_index
tensor([[[   0],
         [   0],
         [   0],
         [   0],
         [1396],
         [   0],
         ...
         [1885],
         [1670],
         [   0],
         [   0],
         [   0],
         [   0],
         [   0]],

        [[   0],
         [   0],
         [   0],
         [   0],
         [1396],
        ...
         [1885],
         [1670],
         [   0],
         [   0],
         [   0],
         [   0],
         [   0],
         [   0],
         [   0]]], device='cuda:0')

搞了维度精简之后:

ipdb>topk_index
tensor([[   0,    0,    0,    0, 1396,    0,    0,    0,    0,    0,    0,    0,
            0,    0,    0,    0,    0,    0,    0,    0,    0,    0,    0,    0,
            0,    0,    0,    0,    0, 1845,    0,    0,    0, 1396,    0,    0,
            0,    0,    0,    0,    0,    0,    0, 1885, 1670,    0,    0,    0,
            0,    0],
        [   0,    0,    0,    0, 1396,    0,    0,    0,    0,    0,    0,    0,
            0,    0,    0,    0,    0,    0,    0,    0,    0,    0,    0,    0,
            0,    0,    0, 1845,    0,    0,    0, 1396,    0,    0,    0,    0,
            0,    0,    0,    0,    0, 1885, 1670,    0,    0,    0,    0,    0,
            0,    0]], device='cuda:0')

然后是长度mask:

ipdb>mask
    324         mask = make_pad_mask(encoder_out_lens, maxlen)  # (B, maxlen)
tensor([[False, False, False, False, False, False, False, False, False, False,
         False, False, False, False, False, False, False, False, False, False,
         False, False, False, False, False, False, False, False, False, False,
         False, False, False, False, False, False, False, False, False, False,
         False, False, False, False, False, False, False, False, False, False],
        [False, False, False, False, False, False, False, False, False, False,
         False, False, False, False, False, False, False, False, False, False,
         False, False, False, False, False, False, False, False, False, False,
         False, False, False, False, False, False, False, False, False, False,
         False, False, False, False, False, False, False, False,  True,  True]],
       device='cuda:0')

经历过masked_fill_之后,topk_index为:

ipdb>topk_index
tensor([[   0,    0,    0,    0, 1396,    0,    0,    0,    0,    0,    0,    0,
            0,    0,    0,    0,    0,    0,    0,    0,    0,    0,    0,    0,
            0,    0,    0,    0,    0, 1845,    0,    0,    0, 1396,    0,    0,
            0,    0,    0,    0,    0,    0,    0, 1885, 1670,    0,    0,    0,
            0,    0],
        [   0,    0,    0,    0, 1396,    0,    0,    0,    0,    0,    0,    0,
            0,    0,    0,    0,    0,    0,    0,    0,    0,    0,    0,    0,
            0,    0,    0, 1845,    0,    0,    0, 1396,    0,    0,    0,    0,
            0,    0,    0,    0,    0, 1885, 1670,    0,    0,    0,    0,    0,
         5501, 5501]], device='cuda:0')

可以看到已经填充了5501= end of sentence标签了。

如此,就收集到了hyps:

[[0,0,0,0,1396,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,
0, 0, 0, 0, 0, 1845, 0, 0, 0, 1396, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1885, 
1670, 0, 0, 0, 0, 0], 

[0, 0, 0, 0, 1396, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 
0, 0, 0, 1845, 0, 0, 0, 1396, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1885, 1670, 
0, 0, 0, 0, 0, 5501, 5501]]

搞了remove_duplicates_and_blank之后:

[[1396,1845,1396,1885,1670],[1396,1845,1396,1885,1670,5501]]

就是两个语音的分别对应的文本id输出了。后续会简单把id映射回文字即可。

('A03M0156_00000.612_00002.674','×の×ます')
('A03M0156_00002.989_00004.918', '×の×ます')

这样这个解码算法,就算掰扯完了。