在哪儿

已经连续学习了两个inference算法,本次学习第三个:

ctc prefix beam search,感觉这个貌似和attention rescoring非常类似。

在脚本中对应的是:

decode_mode="ctc_prefix_beam_search"算法

【attention rescoring中完整的包括了ctc-prefix-beam-search的所有逻辑!】

而且这个也要求batch-size=1:


本inference算法要求batch size=1

看细节。

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_prefix_beam_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_prefix_beam_search/text_bpe',
 reverse_weight=0.0, simulate_streaming=False, test_data='data/test1/data.list')

configs

这个和之前用的一样

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_prefix_beam_search

整体流程


整体上,还是这几步不变

这几步不变的:

数据准备

模型初始化

load已有的训练好的模型checkpoint

解码

one batch


对于一个batch的处理,包括解码,以及结果收集

解码部分:


还是熟悉的三个步骤!

还是三步:wav encoder(12层Conformer Encoder Layers),然后是单层的linear把512->5502词表;最后是基于ctc beam search来为每个time point (or, frame)确定若干候选。

beam size=10表示最后留下10个最好的候选序列,和它们的积累概率(对数)。

这个的细节可以参看:

http://www.speechhome.com/post/1504652982149582848

看几个中间变量的值:

hyps

ipdb>hyps
[((1396, 1396, 1396, 1396, 1396, 1396, 1885, 1670), -15.887155534713223), 
((1396, 1396, 1396, 1396, 1396, 1396, 1885), -15.996046160775117), 
((1396, 1396, 1396, 1396, 1396, 1396, 1396, 1670), -16.105350054413357), 
((1396, 1396, 1396, 1396, 1396, 1396, 1396, 1885, 1670), -16.110224930299356), 
((1396, 1396, 1396, 1396, 1396, 1885, 1670), -16.12109183842913), 
((1396, 1396, 1396, 1396, 1396, 1396, 1741, 1670), -16.169015181751966), 
((1396, 1396, 1396, 1396, 1396, 1396, 1670), -16.19584450996598), 
((1396, 1396, 1396, 1396, 1396, 1396, 1396), -16.20512037370705), 
((1396, 1396, 1396, 1396, 1396, 1396, 1396, 1885), -16.219115556361253), 
((1396, 1396, 1396, 1396, 1396, 1885), -16.229982464491023)]

hyps[0]:

((1396,1396,1396,1396,1396,1396,1885,1670),-15.887155534713223)

保存的结果:

>/workspace/asr/wenet/examples/csj/s0/wenet/bin/recognize.py(226)main()
225                 logging.info('{} {}'.format(key, content))
3-> 226                 fout.write('{} {}\n'.format(key, content))
    227

ipdb> key, content
('A03M0156_00000.612_00002.674', '××××××ます')

这样,ctc prefix beam search就算学习完毕了。