学习第四个解码算法:解码算法

在脚本中对应的是:

decode_mode="attention"

怎么还是感觉和之前的有一定的重合的地方???[这个是特殊的!自回归+beam search]

这个允许batch_size > 1,所以我们设置为2.

整体流程


整体准备流程就不再详细讲了,和前面几个inference算法是一致的。我们就从上面脑图的最后一行展开。


准备好一个batch之后(batch-size=2),开启核心的算法之旅:model.recognize()

one batch

>/workspace/asr/wenet/examples/csj/s0/wenet/bin/recognize.py(174)main()
    173             keys, feats, target, feats_lengths, target_lengths = batch
--> 174             feats = feats.to(device)
    175             target = target.to(device)

ipdb> keys
['A03M0156_00000.612_00002.674', 'A03M0156_00002.989_00004.918']
ipdb> feats.shape
torch.Size([2, 204, 80])
ipdb> target.shape
torch.Size([2, 13])
ipdb> target
tensor([[3461, 3927, 2815, 2906, 1845, 1396, 1396, 1762, 4098, 1640, 1885, 1670,
           -1],
        [1554, 2161, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396,
         1396]])
ipdb> feats_lengths
tensor([204, 191], dtype=torch.int32)
ipdb> target_lengths
tensor([12, 13], dtype=torch.int32)

上面直接是one batch信息了。

核心思想:

就是拿decoder的left-decoder当作step-by-step(逐步解码的)自回归模型来用,加上beam search搜索最佳文本序列。

这个算法涉及的代码组织的是非常漂亮的,值得学习!


上面是四个比较重要的,开启自回归解码+beam search的算法

展开来瞧瞧,看看输入变量和中间变量等的shape信息:


展开来看看,hyps,scores等的各自的shape


逐步地,一个token一个token地,调用decoder,然后得到10-best,然后合并beam,即基于自回归+beam search的解码!!这也是传统机器翻译常用的套路。

上面的截屏是我们重点想学习的内容,思想很不错,代码还非常简洁。

如果简单用几句话概括就是,一个语音,10个文本候选,然后每个候选自己找10个后接词,这样一共就有100个序列,按照当前序列得分+新词得分的顺序,重新排列,然后再从这100个候选里面,按照得分选择10个最好的。

如此反复,即实现了,自回归+beam size=10的解码过程。

当然,这里是以batch size =2 为例来解说。这样就是每个语音输入,一次有100个候选,然后从中挑选10个。如此扩展,挑选,反复进行,直到所有候选序列都遇到了eos。

这里因为maxlen=50,所以i=1到50,i=0是已经给了sos = start of sequence。

i=1 降龙十八掌

一个i取值下,有十八行代码,我们称其为“降龙十八掌”!

看下目前的“输入”变量的取值:

1. if end_flag.sum() == running_size: break

这里running_size= batch_size * beam_size = 2 * 10 = 20

>/workspace/asr/wenet/wenet/transformer/asr_model.py(238)recognize()
    237             import ipdb; ipdb.set_trace()
--> 238             if end_flag.sum() == running_size:
    239                 break

2022-03-14 13:26:08,428 DEBUG Using selector: EpollSelector
ipdb> end_flag # [20, 1]
tensor([[False],
        [False],
        [False],
        [False],
        [False],
        [False],
        [False],
        [False],
        [False],
        [False],
        [False],
        [False],
        [False],
        [False],
        [False],
        [False],
        [False],
        [False],
        [False],
        [False]], device='cuda:0')
ipdb> end_flag.sum()
tensor(0, device='cuda:0')

还没有到结束的时候。

batch-size=2,每个wav是10个序列,所以一共是20个序列。都遇到eos的时候,对

i的循环结束。

【打完收工】

2. causal mask构造

>/workspace/asr/wenet/wenet/transformer/asr_model.py(241)recognize()
    240             # 2.1 Forward decoder step
--> 241             hyps_mask = subsequent_mask(i).unsqueeze(0).repeat(
    242                 running_size, 1, 1).to(device)  # (B*N, i, i)

根据i的取值,构造(B*N, i, i)样式的causal masking。例如,i=2的时候,

[True, False]

[True, True]

这样的。

这个mask是加在目标文本序列上的,为的是作为下一步“decoder一步”的输入,控制文本序列的可见范围。

目前是i=1,所以,hyps_mask的取值为:

ipdb>hyps_mask# [20, 1, 1]
tensor([[[True]],
        [[True]],
        [[True]],
        [[True]],
        [[True]],
        [[True]],
        [[True]],
        [[True]],
        [[True]],
        [[True]],
        [[True]],
        [[True]],
        [[True]],
        [[True]],
        [[True]],
        [[True]],
        [[True]],
        [[True]],
        [[True]],
        [[True]]], device='cuda:0')

3. decoder一步

>/workspace/asr/wenet/wenet/transformer/asr_model.py(244)recognize()
    243             # logp: (B*N, vocab)
--> 244             logp, cache = self.decoder.forward_one_step(
    245                 encoder_out, encoder_mask, hyps, hyps_mask, cache)

这个,就是根据准备好的,如下信息:

  1. encoder_out
  2. encoder_mask
  3. hyps
  4. hyps_mask
  5. cache=None

来调用decoder的forward_one_step函数,这个函数内部,就是调用left_decoder来解码。


forward_one_step的输入参数和细节过程

这里的forward_one_step里面,有:

  1. 目标文本序列的embed
  2. 遍历self.decoders的三层decoder layers,解码;
  3. 解码结束之后,调用linear layer, 从512映射到5502。

这相当于一次自回归(one step auto-regressive decoding)。

返回的是logp.shape=[20, 5502]的张量。

“logp" 变量的含义:记录的是20个序列,每个序列的下一个候选词的分别的概率(log, 因为经历了softmax -> log)。

4. 取前10

为每个序列的最后一个新增加的位置,从5502个候选中,挑选得分(log概率)最大的10个。

因为beam size=10。


取前10的结果

>/workspace/asr/wenet/wenet/transformer/asr_model.py(248)recognize()
    247             top_k_logp, top_k_index = logp.topk(beam_size)  # 从(B*N, 5502)到(B*N, N)
--> 248             top_k_logp = mask_finished_scores(top_k_logp, end_flag)
    249             top_k_index = mask_finished_preds(top_k_index, end_flag, self.eos)

ipdb> top_k_logp.shape
torch.Size([20, 10])
ipdb> top_k_logp
tensor([[-1.1374, -1.5913, -2.6713, -3.7542, -3.8323, -3.9647, -4.0388, -4.2002,
         -4.2253, -4.5013],
        [-1.1374, -1.5913, -2.6713, -3.7542, -3.8323, -3.9647, -4.0388, -4.2002,
         -4.2253, -4.5013],
        [-1.1374, -1.5913, -2.6713, -3.7542, -3.8323, -3.9647, -4.0388, -4.2002,
         -4.2253, -4.5013],
        [-1.1374, -1.5913, -2.6713, -3.7542, -3.8323, -3.9647, -4.0388, -4.2002,
         -4.2253, -4.5013],
        [-1.1374, -1.5913, -2.6713, -3.7542, -3.8323, -3.9647, -4.0388, -4.2002,
         -4.2253, -4.5013],
        [-1.1374, -1.5913, -2.6713, -3.7542, -3.8323, -3.9647, -4.0388, -4.2002,
         -4.2253, -4.5013],
        [-1.1374, -1.5913, -2.6713, -3.7542, -3.8323, -3.9647, -4.0388, -4.2002,
         -4.2253, -4.5013],
        [-1.1374, -1.5913, -2.6713, -3.7542, -3.8323, -3.9647, -4.0388, -4.2002,
         -4.2253, -4.5013],
        [-1.1374, -1.5913, -2.6713, -3.7542, -3.8323, -3.9647, -4.0388, -4.2002,
         -4.2253, -4.5013],
        [-1.1374, -1.5913, -2.6713, -3.7542, -3.8323, -3.9647, -4.0388, -4.2002,
         -4.2253, -4.5013],
        [-1.1464, -1.5922, -2.6569, -3.7605, -3.7612, -3.9860, -4.0502, -4.2169,
         -4.2186, -4.5011],
        [-1.1464, -1.5922, -2.6569, -3.7605, -3.7612, -3.9860, -4.0502, -4.2169,
         -4.2186, -4.5011],
        [-1.1464, -1.5922, -2.6569, -3.7605, -3.7612, -3.9860, -4.0502, -4.2169,
         -4.2186, -4.5011],
        [-1.1464, -1.5922, -2.6569, -3.7605, -3.7612, -3.9860, -4.0502, -4.2169,
         -4.2186, -4.5011],
        [-1.1464, -1.5922, -2.6569, -3.7605, -3.7612, -3.9860, -4.0502, -4.2169,
         -4.2186, -4.5011],
        [-1.1464, -1.5922, -2.6569, -3.7605, -3.7612, -3.9860, -4.0502, -4.2169,
         -4.2186, -4.5011],
        [-1.1464, -1.5922, -2.6569, -3.7605, -3.7612, -3.9860, -4.0502, -4.2169,
         -4.2186, -4.5011],
        [-1.1464, -1.5922, -2.6569, -3.7605, -3.7612, -3.9860, -4.0502, -4.2169,
         -4.2186, -4.5011],
        [-1.1464, -1.5922, -2.6569, -3.7605, -3.7612, -3.9860, -4.0502, -4.2169,
         -4.2186, -4.5011],
        [-1.1464, -1.5922, -2.6569, -3.7605, -3.7612, -3.9860, -4.0502, -4.2169,
         -4.2186, -4.5011]], device='cuda:0')
ipdb> top_k_index.shape
torch.Size([20, 10])
ipdb> top_k_index
tensor([[1554, 1396, 1516, 1609, 3585, 1885, 1741, 1677, 1762, 2392],
        [1554, 1396, 1516, 1609, 3585, 1885, 1741, 1677, 1762, 2392],
        [1554, 1396, 1516, 1609, 3585, 1885, 1741, 1677, 1762, 2392],
        [1554, 1396, 1516, 1609, 3585, 1885, 1741, 1677, 1762, 2392],
        [1554, 1396, 1516, 1609, 3585, 1885, 1741, 1677, 1762, 2392],
        [1554, 1396, 1516, 1609, 3585, 1885, 1741, 1677, 1762, 2392],
        [1554, 1396, 1516, 1609, 3585, 1885, 1741, 1677, 1762, 2392],
        [1554, 1396, 1516, 1609, 3585, 1885, 1741, 1677, 1762, 2392],
        [1554, 1396, 1516, 1609, 3585, 1885, 1741, 1677, 1762, 2392],
        [1554, 1396, 1516, 1609, 3585, 1885, 1741, 1677, 1762, 2392],
        [1554, 1396, 1516, 3585, 1609, 1885, 1741, 1677, 1762, 2392],
        [1554, 1396, 1516, 3585, 1609, 1885, 1741, 1677, 1762, 2392],
        [1554, 1396, 1516, 3585, 1609, 1885, 1741, 1677, 1762, 2392],
        [1554, 1396, 1516, 3585, 1609, 1885, 1741, 1677, 1762, 2392],
        [1554, 1396, 1516, 3585, 1609, 1885, 1741, 1677, 1762, 2392],
        [1554, 1396, 1516, 3585, 1609, 1885, 1741, 1677, 1762, 2392],
        [1554, 1396, 1516, 3585, 1609, 1885, 1741, 1677, 1762, 2392],
        [1554, 1396, 1516, 3585, 1609, 1885, 1741, 1677, 1762, 2392],
        [1554, 1396, 1516, 3585, 1609, 1885, 1741, 1677, 1762, 2392],
        [1554, 1396, 1516, 3585, 1609, 1885, 1741, 1677, 1762, 2392]],
       device='cuda:0')

top_k_logp是20行10列,表示20个文本序列【为每个wav构造10个文本序列】,每个序列的新预测出来的token的最好的10个得分;

top_k_index也是20行10列,表示20个文本序列,每个序列的新预测出来的10e个得分最高的token所在的位置, token.id。

5. mask_finished_scores

top_k_logp = mask_finished_scores(top_k_logp, end_flag)

根据end_flag来对top_k_logp进行mask。


6. mask_finished_preds

top_k_index = mask_finished_preds(top_k_index, end_flag, self.eos)

根据end_flag和self.eos来对top_k_index来进行mask。


因为经过这俩mask,还没啥变换,先这样了。

7. 分数叠加(logp)

scores = scores + top_k_logp  # (B*N, N), broadcast add

在两者相加之前:

【需要留心scores的初始赋值!直觉感觉都是0.也行?0=log1】

含义:scores:目前为止20个文本序列的叠加之后的得分(log p);【一个序列一个取值,叠加log概率】

top_k_logp: 每个序列,十个新候选的得分。【一个序列10个值,是该序列新加的十个(最有可能的)词的分别的得分】

相加:(一个序列一个取值-自我复制10次)分别和(一个序列10个值)相加,得到(一个序列10个值)。

输出:20个序列,每个序列10个”再叠加“得分。

ipdb>scores# [20, 1],20个文本序列(部分),目前为止的“累计”得分。
tensor([[0.],
        [-inf],
        [-inf],
        [-inf],
        [-inf],
        [-inf],
        [-inf],
        [-inf],
        [-inf],
        [-inf],
        [0.],
        [-inf],
        [-inf],
        [-inf],
        [-inf],
        [-inf],
        [-inf],
        [-inf],
        [-inf],
        [-inf]], device='cuda:0')
ipdb> top_k_logp #[20,10]
tensor([[-1.1374, -1.5913, -2.6713, -3.7542, -3.8323, -3.9647, -4.0388, -4.2002,
         -4.2253, -4.5013],
        [-1.1374, -1.5913, -2.6713, -3.7542, -3.8323, -3.9647, -4.0388, -4.2002,
         -4.2253, -4.5013],
        [-1.1374, -1.5913, -2.6713, -3.7542, -3.8323, -3.9647, -4.0388, -4.2002,
         -4.2253, -4.5013],
        [-1.1374, -1.5913, -2.6713, -3.7542, -3.8323, -3.9647, -4.0388, -4.2002,
         -4.2253, -4.5013],
        [-1.1374, -1.5913, -2.6713, -3.7542, -3.8323, -3.9647, -4.0388, -4.2002,
         -4.2253, -4.5013],
        [-1.1374, -1.5913, -2.6713, -3.7542, -3.8323, -3.9647, -4.0388, -4.2002,
         -4.2253, -4.5013],
        [-1.1374, -1.5913, -2.6713, -3.7542, -3.8323, -3.9647, -4.0388, -4.2002,
         -4.2253, -4.5013],
        [-1.1374, -1.5913, -2.6713, -3.7542, -3.8323, -3.9647, -4.0388, -4.2002,
         -4.2253, -4.5013],
        [-1.1374, -1.5913, -2.6713, -3.7542, -3.8323, -3.9647, -4.0388, -4.2002,
         -4.2253, -4.5013],
        [-1.1374, -1.5913, -2.6713, -3.7542, -3.8323, -3.9647, -4.0388, -4.2002,
         -4.2253, -4.5013],
        [-1.1464, -1.5922, -2.6569, -3.7605, -3.7612, -3.9860, -4.0502, -4.2169,
         -4.2186, -4.5011],
        [-1.1464, -1.5922, -2.6569, -3.7605, -3.7612, -3.9860, -4.0502, -4.2169,
         -4.2186, -4.5011],
        [-1.1464, -1.5922, -2.6569, -3.7605, -3.7612, -3.9860, -4.0502, -4.2169,
         -4.2186, -4.5011],
        [-1.1464, -1.5922, -2.6569, -3.7605, -3.7612, -3.9860, -4.0502, -4.2169,
         -4.2186, -4.5011],
        [-1.1464, -1.5922, -2.6569, -3.7605, -3.7612, -3.9860, -4.0502, -4.2169,
         -4.2186, -4.5011],
        [-1.1464, -1.5922, -2.6569, -3.7605, -3.7612, -3.9860, -4.0502, -4.2169,
         -4.2186, -4.5011],
        [-1.1464, -1.5922, -2.6569, -3.7605, -3.7612, -3.9860, -4.0502, -4.2169,
         -4.2186, -4.5011],
        [-1.1464, -1.5922, -2.6569, -3.7605, -3.7612, -3.9860, -4.0502, -4.2169,
         -4.2186, -4.5011],
        [-1.1464, -1.5922, -2.6569, -3.7605, -3.7612, -3.9860, -4.0502, -4.2169,
         -4.2186, -4.5011],
        [-1.1464, -1.5922, -2.6569, -3.7605, -3.7612, -3.9860, -4.0502, -4.2169,
         -4.2186, -4.5011]], device='cuda:0')

之后是:20个序列,每个序列10个”再叠加“得分。

-inf和任何值”相加“,结果还是-inf。

ipdb>scores# [20, 10]
tensor([[-1.1374, -1.5913, -2.6713, -3.7542, -3.8323, -3.9647, -4.0388, -4.2002,
         -4.2253, -4.5013],
        [   -inf,    -inf,    -inf,    -inf,    -inf,    -inf,    -inf,    -inf,
            -inf,    -inf],
        [   -inf,    -inf,    -inf,    -inf,    -inf,    -inf,    -inf,    -inf,
            -inf,    -inf],
        [   -inf,    -inf,    -inf,    -inf,    -inf,    -inf,    -inf,    -inf,
            -inf,    -inf],
        [   -inf,    -inf,    -inf,    -inf,    -inf,    -inf,    -inf,    -inf,
            -inf,    -inf],
        [   -inf,    -inf,    -inf,    -inf,    -inf,    -inf,    -inf,    -inf,
            -inf,    -inf],
        [   -inf,    -inf,    -inf,    -inf,    -inf,    -inf,    -inf,    -inf,
            -inf,    -inf],
        [   -inf,    -inf,    -inf,    -inf,    -inf,    -inf,    -inf,    -inf,
            -inf,    -inf],
        [   -inf,    -inf,    -inf,    -inf,    -inf,    -inf,    -inf,    -inf,
            -inf,    -inf],
        [   -inf,    -inf,    -inf,    -inf,    -inf,    -inf,    -inf,    -inf,
            -inf,    -inf],
        [-1.1464, -1.5922, -2.6569, -3.7605, -3.7612, -3.9860, -4.0502, -4.2169,
         -4.2186, -4.5011],
        [   -inf,    -inf,    -inf,    -inf,    -inf,    -inf,    -inf,    -inf,
            -inf,    -inf],
        [   -inf,    -inf,    -inf,    -inf,    -inf,    -inf,    -inf,    -inf,
            -inf,    -inf],
        [   -inf,    -inf,    -inf,    -inf,    -inf,    -inf,    -inf,    -inf,
            -inf,    -inf],
        [   -inf,    -inf,    -inf,    -inf,    -inf,    -inf,    -inf,    -inf,
            -inf,    -inf],
        [   -inf,    -inf,    -inf,    -inf,    -inf,    -inf,    -inf,    -inf,
            -inf,    -inf],
        [   -inf,    -inf,    -inf,    -inf,    -inf,    -inf,    -inf,    -inf,
            -inf,    -inf],
        [   -inf,    -inf,    -inf,    -inf,    -inf,    -inf,    -inf,    -inf,
            -inf,    -inf],
        [   -inf,    -inf,    -inf,    -inf,    -inf,    -inf,    -inf,    -inf,
            -inf,    -inf],
        [   -inf,    -inf,    -inf,    -inf,    -inf,    -inf,    -inf,    -inf,
            -inf,    -inf]], device='cuda:0')

8. 变换scores shape

从(B*N, N) -> (B, N*N)。

scores=scores.view(batch_size,beam_size*beam_size)# (B, N*N)

转变之后,scores为:

【含义为】2个wav,每个wav的10个候选,分别扩展了10次之后,就得到了100个候选的”得分“,如下所示。

下一步,就是从这100个里面,排序挑选10个最好的得分。【100选10】

ipdb>scores
tensor([[-1.1374, -1.5913, -2.6713, -3.7542, -3.8323, -3.9647, -4.0388, -4.2002,
         -4.2253, -4.5013,    -inf,    -inf,    -inf,    -inf,    -inf,    -inf,
            -inf,    -inf,    -inf,    -inf,    -inf,    -inf,    -inf,    -inf,
            -inf,    -inf,    -inf,    -inf,    -inf,    -inf,    -inf,    -inf,
            -inf,    -inf,    -inf,    -inf,    -inf,    -inf,    -inf,    -inf,
            -inf,    -inf,    -inf,    -inf,    -inf,    -inf,    -inf,    -inf,
            -inf,    -inf,    -inf,    -inf,    -inf,    -inf,    -inf,    -inf,
            -inf,    -inf,    -inf,    -inf,    -inf,    -inf,    -inf,    -inf,
            -inf,    -inf,    -inf,    -inf,    -inf,    -inf,    -inf,    -inf,
            -inf,    -inf,    -inf,    -inf,    -inf,    -inf,    -inf,    -inf,
            -inf,    -inf,    -inf,    -inf,    -inf,    -inf,    -inf,    -inf,
            -inf,    -inf,    -inf,    -inf,    -inf,    -inf,    -inf,    -inf,
            -inf,    -inf,    -inf,    -inf],
        [-1.1464, -1.5922, -2.6569, -3.7605, -3.7612, -3.9860, -4.0502, -4.2169,
         -4.2186, -4.5011,    -inf,    -inf,    -inf,    -inf,    -inf,    -inf,
            -inf,    -inf,    -inf,    -inf,    -inf,    -inf,    -inf,    -inf,
            -inf,    -inf,    -inf,    -inf,    -inf,    -inf,    -inf,    -inf,
            -inf,    -inf,    -inf,    -inf,    -inf,    -inf,    -inf,    -inf,
            -inf,    -inf,    -inf,    -inf,    -inf,    -inf,    -inf,    -inf,
            -inf,    -inf,    -inf,    -inf,    -inf,    -inf,    -inf,    -inf,
            -inf,    -inf,    -inf,    -inf,    -inf,    -inf,    -inf,    -inf,
            -inf,    -inf,    -inf,    -inf,    -inf,    -inf,    -inf,    -inf,
            -inf,    -inf,    -inf,    -inf,    -inf,    -inf,    -inf,    -inf,
            -inf,    -inf,    -inf,    -inf,    -inf,    -inf,    -inf,    -inf,
            -inf,    -inf,    -inf,    -inf,    -inf,    -inf,    -inf,    -inf,
            -inf,    -inf,    -inf,    -inf]], device='cuda:0')
ipdb> scores.shape
torch.Size([2, 100])

9. topk of scores

【百里挑十】

scores,offset_k_index=scores.topk(k=beam_size)# (B, N)

得到得分,以及对应的位置索引。

得到的结果为:

ipdb>scores
tensor([[-1.1374, -1.5913, -2.6713, -3.7542, -3.8323, -3.9647, -4.0388, -4.2002,
         -4.2253, -4.5013],
        [-1.1464, -1.5922, -2.6569, -3.7605, -3.7612, -3.9860, -4.0502, -4.2169,
         -4.2186, -4.5011]], device='cuda:0')
ipdb> offset_k_index
tensor([[0, 1, 2, 3, 4, 5, 6, 7, 8, 9],
        [0, 1, 2, 3, 4, 5, 6, 7, 8, 9]], device='cuda:0')

从结果看,第一个wav,是从100个候选里面,选择了编号为[0, 1, ..., 9]的;

同样,第二个wav,也是从100个候选文本序列里面,选择了编号为[0, 1, ..., 9]的。

10. scores reshape

scores=scores.view(-1,1)# (B*N, 1)

这个变形之后,含义就是,20个候选文本序列,扩展一次之后,新的累计对数概率的值。

经历过上面一行代码之后,scores的形状和取值分别为:

ipdb>scores.shape
torch.Size([20, 1])
ipdb> scores
tensor([[-1.1374],
        [-1.5913],
        [-2.6713],
        [-3.7542],
        [-3.8323],
        [-3.9647],
        [-4.0388],
        [-4.2002],
        [-4.2253],
        [-4.5013],
        [-1.1464],
        [-1.5922],
        [-2.6569],
        [-3.7605],
        [-3.7612],
        [-3.9860],
        [-4.0502],
        [-4.2169],
        [-4.2186],
        [-4.5011]], device='cuda:0')

11. 构造offset

-->258base_k_index=torch.arange(batch_size,device=device).view(
    259                 -1, 1).repeat([1, beam_size])  # (B, N)

得到的是:

ipdb>base_k_index
tensor([[0, 0, 0, 0, 0, 0, 0, 0, 0, 0],
        [1, 1, 1, 1, 1, 1, 1, 1, 1, 1]], device='cuda:0')

这个和下面的”12. offset *= 100“是个”连招“,需要结合来看。

12. offset *= 100

-->260base_k_index=base_k_index*beam_size*beam_size

这是因为,每个wav会有100个文本候选(10*10),这样第二个wav的100个候选的序号就是从100开始的。

>/workspace/asr/wenet/wenet/transformer/asr_model.py(261)recognize()
    260             base_k_index = base_k_index * beam_size * beam_size
--> 261             best_k_index = base_k_index.view(-1) + offset_k_index.view(
    262                 -1)  # (B*N)

ipdb> base_k_index
tensor([[  0,   0,   0,   0,   0,   0,   0,   0,   0,   0],
        [100, 100, 100, 100, 100, 100, 100, 100, 100, 100]], device='cuda:0')

13. 新的best_k_index

【新定的”座次“】

>/workspace/asr/wenet/wenet/transformer/asr_model.py(262)recognize()
    261             best_k_index = base_k_index.view(-1) + offset_k_index.view(
--> 262                 -1)  # (B*N)

回顾一下:

ipdb>base_k_index.view(-1)
tensor([  0,   0,   0,   0,   0,   0,   0,   0,   0,   0, 100, 100, 100, 100,
        100, 100, 100, 100, 100, 100], device='cuda:0')
ipdb> offset_k_index.view(-1)
tensor([0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 0, 1, 2, 3, 4, 5, 6, 7, 8, 9],
       device='cuda:0')
ipdb> n
> /workspace/asr/wenet/wenet/transformer/asr_model.py(261)recognize()
    260             base_k_index = base_k_index * beam_size * beam_size
--> 261             best_k_index = base_k_index.view(-1) + offset_k_index.view(
    262                 -1)  # (B*N)

ipdb> n
> /workspace/asr/wenet/wenet/transformer/asr_model.py(265)recognize()
    264             # 2.5 Update best hyps
--> 265             best_k_pred = torch.index_select(top_k_index.view(-1),
    266                                              dim=-1,

ipdb> best_k_index
tensor([  0,   1,   2,   3,   4,   5,   6,   7,   8,   9, 100, 101, 102, 103,
        104, 105, 106, 107, 108, 109], device='cuda:0')

可以看到第二个wav相关的index是100开始的了。

什么意思呢?

从结果看,第一个wav,是从100个候选里面,选择了编号为[0,1,...,9]的;
同样,第二个wav,也是从100个候选文本序列里面,选择了编号为[0, 1, ..., 9]的。

即:因为现在大家是一个锅里了,那么第二个wav,编号就要都+100才行,因为前100都是第一个wav的!

11. 12. 13. 感觉可以一招搞定:

-->258base_k_index=torch.arange(batch_size,device=device).view(
    259                 -1, 1).repeat([1, beam_size])  # (B, N)
--> 260             base_k_index = base_k_index * beam_size * beam_size
--> 261             best_k_index = base_k_index.view(-1) + offset_k_index.view(
    262                 -1)  # (B*N)

修改为:

best_k_index=offset_k_index.view(-1)+
    (torch.arange(batch_size, device=device).view(-1,1).repeat([1, beam_size]) 
    * beam_size * beam_size).view(-1).

拆招之后,容易理解一些。

14. 求best_k_pred

依据best_k_index从top_k_index中选择:

【200个里面,根据定好的”座次“best_k_index,来选择20个】

-->265best_k_pred=torch.index_select(top_k_index.view(-1),
    266                                              dim=-1, index=best_k_index) # (B*N)

相关的取值:

这个top_k_index是来自第四步(”4. 取前10“),一步解码之后,每个序列有10个最好的。

这里top_k_index的所谓"index",指的是word_index in vocabulary。【或者叫token.id】

ipdb>top_k_index.view(-1)
tensor([1554, 1396, 1516, 1609, 3585, 1885, 1741, 1677, 1762, 2392, 1554, 1396,
        1516, 1609, 3585, 1885, 1741, 1677, 1762, 2392, 1554, 1396, 1516, 1609,
        3585, 1885, 1741, 1677, 1762, 2392, 1554, 1396, 1516, 1609, 3585, 1885,
        1741, 1677, 1762, 2392, 1554, 1396, 1516, 1609, 3585, 1885, 1741, 1677,
        1762, 2392, 1554, 1396, 1516, 1609, 3585, 1885, 1741, 1677, 1762, 2392,
        1554, 1396, 1516, 1609, 3585, 1885, 1741, 1677, 1762, 2392, 1554, 1396,
        1516, 1609, 3585, 1885, 1741, 1677, 1762, 2392, 1554, 1396, 1516, 1609,
        3585, 1885, 1741, 1677, 1762, 2392, 1554, 1396, 1516, 1609, 3585, 1885,
        1741, 1677, 1762, 2392, 1554, 1396, 1516, 3585, 1609, 1885, 1741, 1677,
        1762, 2392, 1554, 1396, 1516, 3585, 1609, 1885, 1741, 1677, 1762, 2392,
        1554, 1396, 1516, 3585, 1609, 1885, 1741, 1677, 1762, 2392, 1554, 1396,
        1516, 3585, 1609, 1885, 1741, 1677, 1762, 2392, 1554, 1396, 1516, 3585,
        1609, 1885, 1741, 1677, 1762, 2392, 1554, 1396, 1516, 3585, 1609, 1885,
        1741, 1677, 1762, 2392, 1554, 1396, 1516, 3585, 1609, 1885, 1741, 1677,
        1762, 2392, 1554, 1396, 1516, 3585, 1609, 1885, 1741, 1677, 1762, 2392,
        1554, 1396, 1516, 3585, 1609, 1885, 1741, 1677, 1762, 2392, 1554, 1396,
        1516, 3585, 1609, 1885, 1741, 1677, 1762, 2392], device='cuda:0')
ipdb> best_k_index
tensor([  0,   1,   2,   3,   4,   5,   6,   7,   8,   9, 100, 101, 102, 103,
        104, 105, 106, 107, 108, 109], device='cuda:0')


ipdb> best_k_pred
tensor([1554, 1396, 1516, 1609, 3585, 1885, 1741, 1677, 1762, 2392, 1554, 1396,
        1516, 3585, 1609, 1885, 1741, 1677, 1762, 2392], device='cuda:0')

15. 重新规划best_hyps_index

-->268best_hyps_index=best_k_index//beam_size
ipdb>best_k_index
tensor([  0,   1,   2,   3,   4,   5,   6,   7,   8,   9, 100, 101, 102, 103,
        104, 105, 106, 107, 108, 109], device='cuda:0')

ipdb> best_hyps_index
tensor([ 0,  0,  0,  0,  0,  0,  0,  0,  0,  0, 10, 10, 10, 10, 10, 10, 10, 10,
        10, 10], device='cuda:0')

【注意】这个//beam_size的含义,其实是说,

[0, 1, ..., 9]这10个候选,都是从原来的0号候选【第0个wav的第0个】扩展出来的;

[100, 101, ..., 109]这10个候选,都是从原来的10号候选【第1个wav的第0个】扩展出来的。

如果,这里有标号”11“,则11//beam_size=1,表明这个11号候选是从原来的1号候选【第0个wav的第1个】扩展出来的;所谓”原来的“,指的是执行”一步解码“之前的那个”原来的“。

16. last_best_k_hyps

-->269last_best_k_hyps=torch.index_select(
    270                 hyps, dim=0, index=best_hyps_index)  # (B*N, i)

效果为:

ipdb>hyps
tensor([-->选它[5501],
        [5501],
        [5501],
        [5501],
        [5501],
        [5501],
        [5501],
        [5501],
        [5501],
        [5501],
        -->选它[5501],
        [5501],
        [5501],
        [5501],
        [5501],
        [5501],
        [5501],
        [5501],
        [5501],
        [5501]], device='cuda:0')
ipdb> best_hyps_index
tensor([ 0,  0,  0,  0,  0,  0,  0,  0,  0,  0, 10, 10, 10, 10, 10, 10, 10, 10,
        10, 10], device='cuda:0')

--->

ipdb> last_best_k_hyps
tensor([[5501],
        [5501],
        [5501],
        [5501],
        [5501],
        [5501],
        [5501],
        [5501],
        [5501],
        [5501],
        [5501],
        [5501],
        [5501],
        [5501],
        [5501],
        [5501],
        [5501],
        [5501],
        [5501],
        [5501]], device='cuda:0')

17. 新旧文本序列结合

-->271hyps=torch.cat((last_best_k_hyps,best_k_pred.view(-1,1)),
    272                              dim=1)  # (B*N, i+1)

得到的结果为:

ipdb>best_k_pred
tensor([1554, 1396, 1516, 1609, 3585, 1885, 1741, 1677, 1762, 2392, 1554, 1396,
        1516, 3585, 1609, 1885, 1741, 1677, 1762, 2392], device='cuda:0')

ipdb> hyps
tensor([[5501, 1554],
        [5501, 1396],
        [5501, 1516],
        [5501, 1609],
        [5501, 3585],
        [5501, 1885],
        [5501, 1741],
        [5501, 1677],
        [5501, 1762],
        [5501, 2392],
        [5501, 1554],
        [5501, 1396],
        [5501, 1516],
        [5501, 3585],
        [5501, 1609],
        [5501, 1885],
        [5501, 1741],
        [5501, 1677],
        [5501, 1762],
        [5501, 2392]], device='cuda:0')

18. 更新end_flag

-->275end_flag=torch.eq(hyps[:,-1],self.eos).view(-1,1)
ipdb>end_flag
tensor([[False],
        [False],
        [False],
        [False],
        [False],
        [False],
        [False],
        [False],
        [False],
        [False],
        [False],
        [False],
        [False],
        [False],
        [False],
        [False],
        [False],
        [False],
        [False],
        [False]], device='cuda:0')

至此,通过这降龙十八掌,就算把i=1的搞定了。

不过瘾。

i=2再打一遍。

i=2 降龙十八掌

一个i取值下,有十八行代码,我们称其为“降龙十八掌”!

看下目前的“输入”变量的取值:

1. if end_flag.sum() == running_size: break

这里running_size= batch_size * beam_size = 2 * 10 = 20

>/workspace/asr/wenet/wenet/transformer/asr_model.py(238)recognize()
    237             import ipdb; ipdb.set_trace()
--> 238             if end_flag.sum() == running_size:
    239                 break

2022-03-14 13:26:08,428 DEBUG Using selector: EpollSelector
ipdb> end_flag # [20, 1]
tensor([[False],
        [False],
        [False],
        [False],
        [False],
        [False],
        [False],
        [False],
        [False],
        [False],
        [False],
        [False],
        [False],
        [False],
        [False],
        [False],
        [False],
        [False],
        [False],
        [False]], device='cuda:0')
ipdb> end_flag.sum()
tensor(0, device='cuda:0')

还没有到结束的时候。

batch-size=2,每个wav是10个序列,所以一共是20个序列。都遇到eos的时候,对

i的循环结束。

【打完收工】

2. causal mask构造

>/workspace/asr/wenet/wenet/transformer/asr_model.py(241)recognize()
    240             # 2.1 Forward decoder step
--> 241             hyps_mask = subsequent_mask(i).unsqueeze(0).repeat(
    242                 running_size, 1, 1).to(device)  # (B*N, i, i)

根据i的取值,构造(B*N, i, i)样式的causal masking。例如,i=2的时候,

[True, False]

[True, True]

这样的。

这个mask是加在目标文本序列上的,为的是作为下一步“decoder一步”的输入,控制文本序列的可见范围。

目前是i=2,所以,hyps_mask的取值为:

ipdb>hyps_mask# [20, 2, 2]
ipdb> hyps_mask
tensor([[[ True, False],
         [ True,  True]],

        [[ True, False],
         [ True,  True]],...

3. decoder一步

>/workspace/asr/wenet/wenet/transformer/asr_model.py(244)recognize()
    243             # logp: (B*N, vocab)
--> 244             logp, cache = self.decoder.forward_one_step(
    245                 encoder_out, encoder_mask, hyps, hyps_mask, cache)

这个,就是根据准备好的,如下信息:

  1. encoder_out
  2. encoder_mask
  3. hyps
  4. hyps_mask
  5. cache=None

来调用decoder的forward_one_step函数,这个函数内部,就是调用left_decoder来解码。


forward_one_step的输入参数和细节过程。需要注意的是,x的形状应该是[20, 2, 512]。因为现在的hyps是长度为2了。

这里的forward_one_step里面,有:

  1. 目标文本序列(目前长度为2)的embed
  2. 遍历self.decoders的三层decoder layers,解码;
  3. 解码结束之后,调用linear layer, 从512映射到5502。

这相当于一次自回归(one step auto-regressive decoding)。

返回的是logp.shape=[20, 5502]的张量。

“logp" 变量的含义:记录的是20个序列,每个序列的下一个候选词的分别的概率(log, 因为经历了softmax -> log)。


4. 取前10

为每个序列的最后一个新增加的位置,从5502个候选中,挑选得分(log概率)最大的10个。

因为beam size=10。


取前10的结果

>/workspace/asr/wenet/wenet/transformer/asr_model.py(248)recognize()
    247             top_k_logp, top_k_index = logp.topk(beam_size)  # 从(B*N, 5502)到(B*N, N)
--> 248             top_k_logp = mask_finished_scores(top_k_logp, end_flag)
    249             top_k_index = mask_finished_preds(top_k_index, end_flag, self.eos)

ipdb> top_k_logp.shape
torch.Size([20, 10])
ipdb> top_k_logp
tensor([[-0.6715, -2.5915, -2.8176, -3.4658, -3.8633, -3.8935, -4.5001, -4.6793,
         -4.8774, -4.8832],
        [-0.1215, -5.1892, -5.8848, -6.0152, -6.3124, -6.3748, -6.4667, -6.9663,
         -7.0732, -7.1089],
        [-0.4446, -3.1427, -3.2962, -3.3160, -4.0385, -4.1486, -4.5409, -5.0743,
         -5.2987, -5.3061],
        [-0.8755, -0.8790, -3.1292, -3.8481, -4.1838, -4.7200, -5.9845, -6.2537,
         -6.7967, -6.7976],
        [-0.3759, -2.2512, -2.5317, -3.9502, -4.0729, -5.7477, -6.2002, -6.2813,
         -6.8066, -7.1728],
        [-1.5993, -1.8278, -2.6657, -2.6836, -3.1362, -3.5816, -3.6726, -3.9984,
         -4.0295, -4.0912],
        [-1.7494, -2.4282, -2.6882, -2.8315, -2.9888, -3.0518, -3.1217, -3.2439,
         -3.3015, -3.4649],
        [-0.8807, -1.1557, -2.5517, -3.1766, -3.2914, -4.4061, -4.7774, -5.3440,
         -6.8675, -7.4507],
        [-2.2175, -2.6802, -2.7597, -3.0684, -3.0971, -3.1796, -3.5121, -3.7264,
         -3.7264, -3.7774],
        [-0.0411, -5.8379, -6.4897, -6.8977, -7.2526, -7.7804, -8.5189, -8.5764,
         -8.6164, -8.6835],
        [-0.6919, -2.5796, -2.8339, -3.3455, -3.8600, -3.8766, -4.5340, -4.6946,
         -4.7719, -4.8475],
        [-0.1226, -5.1558, -5.8437, -5.9238, -6.2855, -6.3183, -6.4208, -6.9053,
         -7.0038, -7.0675],
        [-0.4595, -3.1104, -3.2596, -3.3166, -3.9737, -4.0894, -4.5400, -5.0650,
         -5.2631, -5.2880],
        [-0.3737, -2.2660, -2.5029, -3.9505, -4.0983, -5.7882, -6.2399, -6.2466,
         -6.7734, -7.1520],
        [-0.8648, -0.8968, -3.1275, -3.7561, -4.1587, -4.7638, -5.9310, -6.2513,
         -6.7717, -6.8080],
        [-1.6178, -1.7725, -2.7131, -2.7386, -3.2172, -3.6594, -3.6717, -3.9279,
         -3.9826, -4.0302],
        [-1.7378, -2.4239, -2.7700, -2.9142, -2.9606, -3.0824, -3.1202, -3.2501,
         -3.3159, -3.4040],
        [-0.9077, -1.1315, -2.5595, -3.1133, -3.2934, -4.4051, -4.6483, -5.3214,
         -6.7739, -7.5141],
        [-2.1232, -2.7240, -2.7966, -3.0531, -3.1118, -3.1962, -3.4481, -3.7452,
         -3.8523, -3.8608],
        [-0.0410, -5.8763, -6.5699, -6.8113, -7.2947, -7.7432, -8.5014, -8.5454,
         -8.5757, -8.6094]], device='cuda:0')

ipdb> top_k_index.shape
torch.Size([20, 10])
ipdb> top_k_index
tensor([[2161, 1715, 1396, 3585, 1609, 1762, 1677, 3559, 2011, 2238],
        [1396, 1631, 2392, 1762, 1845, 1575, 2620, 1741, 2171, 1715],
        [1845, 2161, 1396, 1950, 1715, 1516, 1554, 2815, 1677, 1667],
        [1845, 1952, 1609, 1980, 1548, 1708, 1762, 1715, 1516, 1677],
        [2238, 3033, 2533, 4638, 3823, 2815, 3846, 2553, 2182, 2616],
        [2161, 1694, 1516, 1609, 1673, 2248, 2392, 1701, 2375, 1677],
        [1609, 1516, 2248, 1554, 1677, 1867, 1396, 1885, 4267, 1908],
        [1952, 1845, 1640, 1548, 1609, 1527, 1980, 1961, 1715, 1516],
        [2616, 5366, 1527, 2180, 1516, 1396, 2011, 1609, 1554, 2248],
        [4283, 4088, 4783, 5047, 1816, 3474, 5361, 5207, 1592, 4324],
        [2161, 1715, 1396, 3585, 1762, 1609, 1677, 3559, 2238, 2011],
        [1396, 1631, 2392, 1762, 1845, 1575, 2620, 1741, 1715, 2171],
        [1845, 2161, 1396, 1950, 1715, 1516, 1554, 2815, 1677, 1562],
        [2238, 3033, 2533, 4638, 3823, 2815, 3846, 2553, 2182, 2616],
        [1845, 1952, 1609, 1980, 1548, 1708, 1762, 1715, 1677, 1516],
        [2161, 1694, 1609, 1516, 1673, 2248, 2392, 1701, 2375, 1396],
        [1609, 1516, 2248, 1554, 1677, 1396, 1867, 1885, 4267, 1908],
        [1952, 1845, 1640, 1548, 1609, 1527, 1980, 1961, 1715, 1516],
        [2616, 5366, 1527, 2180, 1516, 1396, 2011, 1609, 1554, 2248],
        [4283, 4088, 4783, 5047, 1816, 3474, 5361, 5207, 1592, 4324]],
       device='cuda:0')

top_k_logp是20行10列,表示20个文本序列【为每个wav构造10个文本序列】,每个序列的新预测出来的token的最好的10个得分;

top_k_index也是20行10列,表示20个文本序列,每个序列的新预测出来的10e个得分最高的token所在的位置, token.id。

5. mask_finished_scores

top_k_logp = mask_finished_scores(top_k_logp, end_flag)

根据end_flag来对top_k_logp进行mask。


先挖坑了,大概思想ok

top_k_logp no change!

6. mask_finished_preds

top_k_index = mask_finished_preds(top_k_index, end_flag, self.eos)

top_k_index no change!

根据end_flag和self.eos来对top_k_index来进行mask。


先挖坑了,大概思想ok

因为经过这俩mask,还没啥变换,先这样了。

7. 分数叠加(logp)

scores=scores+top_k_logp# (B*N, N), broadcast add

在两者相加之前:

【需要留心scores的初始赋值!直觉感觉都是0.也行?0=log1】

含义:scores:目前为止20个文本序列的叠加之后的得分(log p);【一个序列一个取值,叠加log概率】

top_k_logp: 每个序列,十个新候选的得分。【一个序列10个值,是该序列新加的十个(最有可能的)词的分别的得分】

相加:(一个序列一个取值-自我复制10次)分别和(一个序列10个值)相加,得到(一个序列10个值)。

输出:20个序列,每个序列10个”再叠加“得分。

ipdb>scores# [20, 1],20个文本序列(部分),目前为止的“累计”得分。
tensor([[-1.1374],
        [-1.5913],
        [-2.6713],
        [-3.7542],
        [-3.8323],
        [-3.9647],
        [-4.0388],
        [-4.2002],
        [-4.2253],
        [-4.5013],
        [-1.1464],
        [-1.5922],
        [-2.6569],
        [-3.7605],
        [-3.7612],
        [-3.9860],
        [-4.0502],
        [-4.2169],
        [-4.2186],
        [-4.5011]], device='cuda:0')

ipdb> top_k_logp #[20,10]
tensor([[-0.6715, -2.5915, -2.8176, -3.4658, -3.8633, -3.8935, -4.5001, -4.6793,
         -4.8774, -4.8832],
        [-0.1215, -5.1892, -5.8848, -6.0152, -6.3124, -6.3748, -6.4667, -6.9663,
         -7.0732, -7.1089],
        [-0.4446, -3.1427, -3.2962, -3.3160, -4.0385, -4.1486, -4.5409, -5.0743,
         -5.2987, -5.3061],
        [-0.8755, -0.8790, -3.1292, -3.8481, -4.1838, -4.7200, -5.9845, -6.2537,
         -6.7967, -6.7976],
        [-0.3759, -2.2512, -2.5317, -3.9502, -4.0729, -5.7477, -6.2002, -6.2813,
         -6.8066, -7.1728],
        [-1.5993, -1.8278, -2.6657, -2.6836, -3.1362, -3.5816, -3.6726, -3.9984,
         -4.0295, -4.0912],
        [-1.7494, -2.4282, -2.6882, -2.8315, -2.9888, -3.0518, -3.1217, -3.2439,
         -3.3015, -3.4649],
        [-0.8807, -1.1557, -2.5517, -3.1766, -3.2914, -4.4061, -4.7774, -5.3440,
         -6.8675, -7.4507],
        [-2.2175, -2.6802, -2.7597, -3.0684, -3.0971, -3.1796, -3.5121, -3.7264,
         -3.7264, -3.7774],
        [-0.0411, -5.8379, -6.4897, -6.8977, -7.2526, -7.7804, -8.5189, -8.5764,
         -8.6164, -8.6835],
        [-0.6919, -2.5796, -2.8339, -3.3455, -3.8600, -3.8766, -4.5340, -4.6946,
         -4.7719, -4.8475],
        [-0.1226, -5.1558, -5.8437, -5.9238, -6.2855, -6.3183, -6.4208, -6.9053,
         -7.0038, -7.0675],
        [-0.4595, -3.1104, -3.2596, -3.3166, -3.9737, -4.0894, -4.5400, -5.0650,
         -5.2631, -5.2880],
        [-0.3737, -2.2660, -2.5029, -3.9505, -4.0983, -5.7882, -6.2399, -6.2466,
         -6.7734, -7.1520],
        [-0.8648, -0.8968, -3.1275, -3.7561, -4.1587, -4.7638, -5.9310, -6.2513,
         -6.7717, -6.8080],
        [-1.6178, -1.7725, -2.7131, -2.7386, -3.2172, -3.6594, -3.6717, -3.9279,
         -3.9826, -4.0302],
        [-1.7378, -2.4239, -2.7700, -2.9142, -2.9606, -3.0824, -3.1202, -3.2501,
         -3.3159, -3.4040],
        [-0.9077, -1.1315, -2.5595, -3.1133, -3.2934, -4.4051, -4.6483, -5.3214,
         -6.7739, -7.5141],
        [-2.1232, -2.7240, -2.7966, -3.0531, -3.1118, -3.1962, -3.4481, -3.7452,
         -3.8523, -3.8608],
        [-0.0410, -5.8763, -6.5699, -6.8113, -7.2947, -7.7432, -8.5014, -8.5454,
         -8.5757, -8.6094]], device='cuda:0')

之后是:20个序列,每个序列10个”再叠加“得分。

[-inf和任何值”相加“,结果还是-inf]。no -inf anymore!

ipdb>scores# [20, 10]
tensor([[ -1.8089,  -3.7290,  -3.9550,  -4.6033,  -5.0007,  -5.0309,  -5.6375,
          -5.8167,  -6.0149,  -6.0206],
        [ -1.7128,  -6.7805,  -7.4761,  -7.6065,  -7.9037,  -7.9661,  -8.0580,
          -8.5576,  -8.6645,  -8.7002],
        [ -3.1158,  -5.8140,  -5.9674,  -5.9873,  -6.7098,  -6.8198,  -7.2122,
          -7.7455,  -7.9700,  -7.9774],
        [ -4.6297,  -4.6332,  -6.8834,  -7.6023,  -7.9379,  -8.4741,  -9.7386,
         -10.0079, -10.5509, -10.5518],
        [ -4.2082,  -6.0835,  -6.3640,  -7.7825,  -7.9051,  -9.5800, -10.0324,
         -10.1136, -10.6389, -11.0051],
        [ -5.5640,  -5.7925,  -6.6304,  -6.6483,  -7.1009,  -7.5463,  -7.6374,
          -7.9631,  -7.9942,  -8.0559],
        [ -5.7882,  -6.4670,  -6.7270,  -6.8702,  -7.0276,  -7.0906,  -7.1605,
          -7.2827,  -7.3403,  -7.5037],
        [ -5.0809,  -5.3559,  -6.7519,  -7.3767,  -7.4916,  -8.6063,  -8.9776,
          -9.5442, -11.0677, -11.6509],
        [ -6.4428,  -6.9055,  -6.9850,  -7.2938,  -7.3224,  -7.4049,  -7.7374,
          -7.9517,  -7.9517,  -8.0027],
        [ -4.5424, -10.3393, -10.9910, -11.3990, -11.7539, -12.2817, -13.0203,
         -13.0777, -13.1177, -13.1848],
        [ -1.8383,  -3.7260,  -3.9803,  -4.4919,  -5.0064,  -5.0230,  -5.6804,
          -5.8410,  -5.9183,  -5.9939],
        [ -1.7148,  -6.7480,  -7.4358,  -7.5160,  -7.8777,  -7.9105,  -8.0130,
          -8.4975,  -8.5960,  -8.6597],
        [ -3.1164,  -5.7673,  -5.9165,  -5.9735,  -6.6306,  -6.7462,  -7.1968,
          -7.7219,  -7.9200,  -7.9449],
        [ -4.1343,  -6.0265,  -6.2635,  -7.7111,  -7.8588,  -9.5487, -10.0004,
         -10.0071, -10.5339, -10.9125],
        [ -4.6260,  -4.6580,  -6.8887,  -7.5173,  -7.9199,  -8.5251,  -9.6923,
         -10.0125, -10.5329, -10.5692],
        [ -5.6037,  -5.7585,  -6.6990,  -6.7246,  -7.2032,  -7.6453,  -7.6577,
          -7.9139,  -7.9686,  -8.0162],
        [ -5.7880,  -6.4741,  -6.8202,  -6.9644,  -7.0108,  -7.1326,  -7.1704,
          -7.3003,  -7.3661,  -7.4542],
        [ -5.1245,  -5.3484,  -6.7763,  -7.3302,  -7.5102,  -8.6220,  -8.8652,
          -9.5383, -10.9908, -11.7309],
        [ -6.3417,  -6.9426,  -7.0152,  -7.2716,  -7.3303,  -7.4147,  -7.6666,
          -7.9638,  -8.0708,  -8.0794],
        [ -4.5421, -10.3775, -11.0710, -11.3124, -11.7958, -12.2444, -13.0026,
         -13.0466, -13.0769, -13.1105]], device='cuda:0')

8. 变换scores shape

从(B*N, N) -> (B, N*N)。

scores = scores.view(batch_size, beam_size * beam_size)  # (B, N*N)

转变之后,scores为:

【含义为】2个wav,每个wav的10个候选,分别扩展了10次之后,就得到了100个候选的”得分“,如下所示。

下一步,就是从这100个里面,排序挑选10个最好的得分。【100选10】

ipdb>scores
tensor([[ -1.8089,  -3.7290,  -3.9550,  -4.6033,  -5.0007,  -5.0309,  -5.6375,
          -5.8167,  -6.0149,  -6.0206,  -1.7128,  -6.7805,  -7.4761,  -7.6065,
          -7.9037,  -7.9661,  -8.0580,  -8.5576,  -8.6645,  -8.7002,  -3.1158,
          -5.8140,  -5.9674,  -5.9873,  -6.7098,  -6.8198,  -7.2122,  -7.7455,
          -7.9700,  -7.9774,  -4.6297,  -4.6332,  -6.8834,  -7.6023,  -7.9379,
          -8.4741,  -9.7386, -10.0079, -10.5509, -10.5518,  -4.2082,  -6.0835,
          -6.3640,  -7.7825,  -7.9051,  -9.5800, -10.0324, -10.1136, -10.6389,
         -11.0051,  -5.5640,  -5.7925,  -6.6304,  -6.6483,  -7.1009,  -7.5463,
          -7.6374,  -7.9631,  -7.9942,  -8.0559,  -5.7882,  -6.4670,  -6.7270,
          -6.8702,  -7.0276,  -7.0906,  -7.1605,  -7.2827,  -7.3403,  -7.5037,
          -5.0809,  -5.3559,  -6.7519,  -7.3767,  -7.4916,  -8.6063,  -8.9776,
          -9.5442, -11.0677, -11.6509,  -6.4428,  -6.9055,  -6.9850,  -7.2938,
          -7.3224,  -7.4049,  -7.7374,  -7.9517,  -7.9517,  -8.0027,  -4.5424,
         -10.3393, -10.9910, -11.3990, -11.7539, -12.2817, -13.0203, -13.0777,
         -13.1177, -13.1848],
        [ -1.8383,  -3.7260,  -3.9803,  -4.4919,  -5.0064,  -5.0230,  -5.6804,
          -5.8410,  -5.9183,  -5.9939,  -1.7148,  -6.7480,  -7.4358,  -7.5160,
          -7.8777,  -7.9105,  -8.0130,  -8.4975,  -8.5960,  -8.6597,  -3.1164,
          -5.7673,  -5.9165,  -5.9735,  -6.6306,  -6.7462,  -7.1968,  -7.7219,
          -7.9200,  -7.9449,  -4.1343,  -6.0265,  -6.2635,  -7.7111,  -7.8588,
          -9.5487, -10.0004, -10.0071, -10.5339, -10.9125,  -4.6260,  -4.6580,
          -6.8887,  -7.5173,  -7.9199,  -8.5251,  -9.6923, -10.0125, -10.5329,
         -10.5692,  -5.6037,  -5.7585,  -6.6990,  -6.7246,  -7.2032,  -7.6453,
          -7.6577,  -7.9139,  -7.9686,  -8.0162,  -5.7880,  -6.4741,  -6.8202,
          -6.9644,  -7.0108,  -7.1326,  -7.1704,  -7.3003,  -7.3661,  -7.4542,
          -5.1245,  -5.3484,  -6.7763,  -7.3302,  -7.5102,  -8.6220,  -8.8652,
          -9.5383, -10.9908, -11.7309,  -6.3417,  -6.9426,  -7.0152,  -7.2716,
          -7.3303,  -7.4147,  -7.6666,  -7.9638,  -8.0708,  -8.0794,  -4.5421,
         -10.3775, -11.0710, -11.3124, -11.7958, -12.2444, -13.0026, -13.0466,
         -13.0769, -13.1105]], device='cuda:0')

ipdb> scores.shape
torch.Size([2, 100])

9. topk of scores

【百里挑十】

scores,offset_k_index=scores.topk(k=beam_size)# (B, N)

得到得分,以及对应的位置索引。

得到的结果为:

ipdb>scores
tensor([[-1.7128, -1.8089, -3.1158, -3.7290, -3.9550, -4.2082, -4.5424, -4.6033,
         -4.6297, -4.6332],
        [-1.7148, -1.8383, -3.1164, -3.7260, -3.9803, -4.1343, -4.4919, -4.5421,
         -4.6260, -4.6580]], device='cuda:0')
ipdb> offset_k_index
tensor([[10,  0, 20,  1,  2, 40, 90,  3, 30, 31],
        [10,  0, 20,  1,  2, 30,  3, 90, 40, 41]], device='cuda:0')

从结果看,第一个wav,是从100个候选里面,选择了编号为[10, 0, ..., 31]的;

同样,第二个wav,也是从100个候选文本序列里面,选择了编号为[10, 0, ..., 41]的。

【我们不一样!我们不一样!我们不一样!】

10. scores reshape

scores=scores.view(-1,1)# (B*N, 1)

这个变形之后,含义就是,20个候选文本序列,扩展一次之后,新的累计对数概率的值。

经历过上面一行代码之后,scores的形状和取值分别为:

ipdb>scores.shape
torch.Size([20, 1])
ipdb> scores
tensor([[-1.7128],
        [-1.8089],
        [-3.1158],
        [-3.7290],
        [-3.9550],
        [-4.2082],
        [-4.5424],
        [-4.6033],
        [-4.6297],
        [-4.6332],
        [-1.7148],
        [-1.8383],
        [-3.1164],
        [-3.7260],
        [-3.9803],
        [-4.1343],
        [-4.4919],
        [-4.5421],
        [-4.6260],
        [-4.6580]], device='cuda:0')

11. 构造offset

-->258base_k_index=torch.arange(batch_size,device=device).view(
    259                 -1, 1).repeat([1, beam_size])  # (B, N)

得到的是:

ipdb>base_k_index
tensor([[0, 0, 0, 0, 0, 0, 0, 0, 0, 0],
        [1, 1, 1, 1, 1, 1, 1, 1, 1, 1]], device='cuda:0')

这个和下面的”12. offset *= 100“是个”连招“,需要结合来看。

12. offset *= 100

-->260base_k_index=base_k_index*beam_size*beam_size

这是因为,每个wav会有100个文本候选(10*10),这样第二个wav的100个候选的序号就是从100开始的。

>/workspace/asr/wenet/wenet/transformer/asr_model.py(261)recognize()
    260             base_k_index = base_k_index * beam_size * beam_size
--> 261             best_k_index = base_k_index.view(-1) + offset_k_index.view(
    262                 -1)  # (B*N)

ipdb> base_k_index
tensor([[  0,   0,   0,   0,   0,   0,   0,   0,   0,   0],
        [100, 100, 100, 100, 100, 100, 100, 100, 100, 100]], device='cuda:0')

13. 新的best_k_index

【新定的”座次“】

>/workspace/asr/wenet/wenet/transformer/asr_model.py(262)recognize()
    261             best_k_index = base_k_index.view(-1) + offset_k_index.view(
--> 262                 -1)  # (B*N)

回顾一下:

ipdb>best_k_index
tensor([ 10,   0,  20,   1,   2,  40,  90,   3,  30,  31, 110, 100, 120, 101,
        102, 130, 103, 190, 140, 141], device='cuda:0')
ipdb> base_k_index.view(-1)
tensor([  0,   0,   0,   0,   0,   0,   0,   0,   0,   0, 100, 100, 100, 100,
        100, 100, 100, 100, 100, 100], device='cuda:0')
ipdb> offset_k_index.view(-1)
tensor([10,  0, 20,  1,  2, 40, 90,  3, 30, 31, 10,  0, 20,  1,  2, 30,  3, 90,
        40, 41], device='cuda:0')

ipdb> best_k_index
tensor([ 10,   0,  20,   1,   2,  40,  90,   3,  30,  31, 

        110, 100, 120, 101, 102, 130, 103, 190, 140, 141], device='cuda:0')

可以看到第二个wav相关的index是100开始的了。

什么意思呢?

从结果看,第一个wav,是从100个候选(0...99)里面,
选择了编号为[10,   0,  20,   1,   2,  40,  90,   3,  30,  31]的;

同样,第二个wav,也是从100个候选(100 ... 199)文本序列里面,
选择了编号为[110, 100, 120, 101, 102, 130, 103, 190, 140, 141]的。

即:因为现在大家是一个锅里了,那么第二个wav,编号就要都+100才行,因为前100都是第一个wav的!

14. 求best_k_pred

依据best_k_index从top_k_index中选择:

【200个里面,根据定好的”座次“best_k_index,来选择20个】

-->265best_k_pred=torch.index_select(top_k_index.view(-1),
    266                                              dim=-1, index=best_k_index) # (B*N)

相关的取值:

这个top_k_index是来自第四步(”4. 取前10“),一步解码之后,每个序列有10个最好的。

这里top_k_index的所谓"index",指的是word_index in vocabulary。【或者叫token.id】

ipdb>top_k_index.view(-1)
tensor([2161, 1715, 1396, 3585, 1609, 1762, 1677, 3559, 2011, 2238, 1396, 1631,
        2392, 1762, 1845, 1575, 2620, 1741, 2171, 1715, 1845, 2161, 1396, 1950,
        1715, 1516, 1554, 2815, 1677, 1667, 1845, 1952, 1609, 1980, 1548, 1708,
        1762, 1715, 1516, 1677, 2238, 3033, 2533, 4638, 3823, 2815, 3846, 2553,
        2182, 2616, 2161, 1694, 1516, 1609, 1673, 2248, 2392, 1701, 2375, 1677,
        1609, 1516, 2248, 1554, 1677, 1867, 1396, 1885, 4267, 1908, 1952, 1845,
        1640, 1548, 1609, 1527, 1980, 1961, 1715, 1516, 2616, 5366, 1527, 2180,
        1516, 1396, 2011, 1609, 1554, 2248, 4283, 4088, 4783, 5047, 1816, 3474,
        5361, 5207, 1592, 4324, 2161, 1715, 1396, 3585, 1762, 1609, 1677, 3559,
        2238, 2011, 1396, 1631, 2392, 1762, 1845, 1575, 2620, 1741, 1715, 2171,
        1845, 2161, 1396, 1950, 1715, 1516, 1554, 2815, 1677, 1562, 2238, 3033,
        2533, 4638, 3823, 2815, 3846, 2553, 2182, 2616, 1845, 1952, 1609, 1980,
        1548, 1708, 1762, 1715, 1677, 1516, 2161, 1694, 1609, 1516, 1673, 2248,
        2392, 1701, 2375, 1396, 1609, 1516, 2248, 1554, 1677, 1396, 1867, 1885,
        4267, 1908, 1952, 1845, 1640, 1548, 1609, 1527, 1980, 1961, 1715, 1516,
        2616, 5366, 1527, 2180, 1516, 1396, 2011, 1609, 1554, 2248, 4283, 4088,
        4783, 5047, 1816, 3474, 5361, 5207, 1592, 4324], device='cuda:0')

ipdb> best_k_index
tensor([ 10,   0,  20,   1,   2,  40,  90,   3,  30,  31, 110, 100, 120, 101,
        102, 130, 103, 190, 140, 141], device='cuda:0')



ipdb> best_k_pred
tensor([1396, 2161, 1845, 1715, 1396, 2238, 4283, 3585, 1845, 1952, 1396, 2161,
        1845, 1715, 1396, 2238, 3585, 4283, 1845, 1952], device='cuda:0')

15. 重新规划best_hyps_index

-->268best_hyps_index=best_k_index//beam_size
ipdb>best_k_index
tensor([ 10,   0,  20,   1,   2,  40,  90,   3,  30,  31, 110, 100, 120, 101,
        102, 130, 103, 190, 140, 141], device='cuda:0')
ipdb> best_hyps_index
tensor([ 1,  0,  2,  0,  0,  4,  9,  0,  3,  3, 11, 10, 12, 10, 10, 13, 10, 19,
        14, 14], device='cuda:0')

【注意】这个//beam_size的含义,其实是说,

[10,0,20,1,2,40,90,3,30,31]这10个候选,
都是分别从原来的第0个wav的第[1,0,2,0,0,4,9,0,3,3]个扩展出来的;


[110, 100, 120, 101,102,130,103,190,140,141]这10个候选,
都是分别从原来的第1个wav的第[11,10,12,10,10,13,10,19,14,14]个扩展出来的;

这里有标号”10“,则10//beam_size=1,表明这个10号候选是从原来的0号候选【第0个wav的第0个】扩展出来的;所谓”原来的“,指的是执行”一步解码“之前的那个”原来的“。

16. last_best_k_hyps

-->269last_best_k_hyps=torch.index_select(
    270                 hyps, dim=0, index=best_hyps_index)  # (B*N, i)

效果为:

ipdb>hyps
tensor([[5501, 1554],  # 0
        [5501, 1396],  # 1
        [5501, 1516],  # 2
        [5501, 1609],
        [5501, 3585],
        [5501, 1885],
        [5501, 1741],
        [5501, 1677],
        [5501, 1762],
        [5501, 2392],
        [5501, 1554],
        [5501, 1396],
        [5501, 1516],
        [5501, 3585],
        [5501, 1609],
        [5501, 1885],
        [5501, 1741],
        [5501, 1677],
        [5501, 1762],
        [5501, 2392], # 19
], device='cuda:0')

ipdb> best_hyps_index
tensor([ 1,  0,  2,  0,  0,  4,  9,  0,  3,  3, 11, 10, 12, 10, 10, 13, 10, 19,
        14, 14], device='cuda:0')


--->

ipdb> last_best_k_hyps
tensor([[5501, 1396],   # 原来的1
        [5501, 1554],   # 原来的0
        [5501, 1516],   # 原来的2
        [5501, 1554],   # 原来的0
        [5501, 1554],   # 原来的0
        [5501, 3585],
        [5501, 2392],
        [5501, 1554],
        [5501, 1609],
        [5501, 1609],
        [5501, 1396],
        [5501, 1554],
        [5501, 1516],
        [5501, 1554],
        [5501, 1554],
        [5501, 3585],
        [5501, 1554],
        [5501, 2392],   # 原来的19
        [5501, 1609],   # 原来的14
        [5501, 1609]    # 原来的14
], device='cuda:0')

17. 新旧文本序列结合

-->271hyps=torch.cat((last_best_k_hyps,best_k_pred.view(-1,1)),
    272                              dim=1)  # (B*N, i+1)

得到的结果为:

ipdb>last_best_k_hyps
tensor([[5501, 1396],
        [5501, 1554],
        [5501, 1516],
        [5501, 1554],
        [5501, 1554],
        [5501, 3585],
        [5501, 2392],
        [5501, 1554],
        [5501, 1609],
        [5501, 1609],
        [5501, 1396],
        [5501, 1554],
        [5501, 1516],
        [5501, 1554],
        [5501, 1554],
        [5501, 3585],
        [5501, 1554],
        [5501, 2392],
        [5501, 1609],
        [5501, 1609]], device='cuda:0')
ipdb> best_k_pred.view(-1,1)
tensor([[1396],
        [2161],
        [1845],
        [1715],
        [1396],
        [2238],
        [4283],
        [3585],
        [1845],
        [1952],
        [1396],
        [2161],
        [1845],
        [1715],
        [1396],
        [2238],
        [3585],
        [4283],
        [1845],
        [1952]], device='cuda:0')
ipdb> hyps
tensor([[5501, 1396, 1396],
        [5501, 1554, 2161],
        [5501, 1516, 1845],
        [5501, 1554, 1715],
        [5501, 1554, 1396],
        [5501, 3585, 2238],
        [5501, 2392, 4283],
        [5501, 1554, 3585],
        [5501, 1609, 1845],
        [5501, 1609, 1952],
        [5501, 1396, 1396],
        [5501, 1554, 2161],
        [5501, 1516, 1845],
        [5501, 1554, 1715],
        [5501, 1554, 1396],
        [5501, 3585, 2238],
        [5501, 1554, 3585],
        [5501, 2392, 4283],
        [5501, 1609, 1845],
        [5501, 1609, 1952]], device='cuda:0')

即,i=3的时候,decoder的输入hyps形状为(20,3).

18. 更新end_flag

-->275end_flag=torch.eq(hyps[:,-1],self.eos).view(-1,1)
ipdb>end_flag
tensor([[False],
        [False],
        [False],
        [False],
        [False],
        [False],
        [False],
        [False],
        [False],
        [False],
        [False],
        [False],
        [False],
        [False],
        [False],
        [False],
        [False],
        [False],
        [False],
        [False]], device='cuda:0')

至此,通过这降龙十八掌,就算把i=2的搞定了。

基本打过瘾了。

收工

执行完毕i=1 to 50之后。

>/workspace/asr/wenet/wenet/transformer/asr_model.py(279)recognize()
    278         import ipdb; ipdb.set_trace()
--> 279         scores = scores.view(batch_size, beam_size)
    280         # TODO: length normalization

2022-03-14 23:10:51,163 DEBUG Using selector: EpollSelector
ipdb> scores.shape
torch.Size([20, 1])
ipdb> scores
tensor([[ -5.7977],
        [ -5.8297],
        [ -5.8950],
        [ -6.0271],
        [ -6.0991],
        [ -6.3089],
        [ -6.4063],
        [ -6.6441],
        [ -6.8040],
        [ -8.3476],
        [ -5.6698],
        [ -5.7080],
        [ -5.8226],
        [ -5.8654],
        [ -6.0950],
        [ -6.1412],
        [ -6.4247],
        [ -6.5081],
        [ -6.8087],
        [-10.0732]], device='cuda:0')

我们继续。

ipdb>scores
tensor([[ -5.7977,  -5.8297,  -5.8950,  -6.0271,  -6.0991,  -6.3089,  -6.4063,
          -6.6441,  -6.8040,  -8.3476],
        [ -5.6698,  -5.7080,  -5.8226,  -5.8654,  -6.0950,  -6.1412,  -6.4247,
          -6.5081,  -6.8087, -10.0732]], device='cuda:0')
ipdb> n
> /workspace/asr/wenet/wenet/transformer/asr_model.py(282)recognize()
    281         best_scores, best_index = scores.max(dim=-1)
--> 282         best_hyps_index = best_index + torch.arange(
    283             batch_size, dtype=torch.long, device=device) * beam_size

ipdb> best_scores
tensor([-5.7977, -5.6698], device='cuda:0')
ipdb> best_index
tensor([0, 0], device='cuda:0')

每个wav选择一个得分最高的文本序列。

进一步:

>/workspace/asr/wenet/wenet/transformer/asr_model.py(282)recognize()
    281         best_scores, best_index = scores.max(dim=-1)
--> 282         best_hyps_index = best_index + torch.arange(
    283             batch_size, dtype=torch.long, device=device) * beam_size

ipdb> n
> /workspace/asr/wenet/wenet/transformer/asr_model.py(284)recognize()
    283             batch_size, dtype=torch.long, device=device) * beam_size
--> 284         best_hyps = torch.index_select(hyps, dim=0, index=best_hyps_index)
    285         best_hyps = best_hyps[:, 1:]

ipdb> best_hyps_index
tensor([ 0, 10], device='cuda:0')

从而,可以根据best_hyps_index来选择best_hyps:

ipdb>hyps
tensor([[5501, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396,
         1396, 1396, 5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501,
         5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501],
        [5501, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396,
         1396, 1396, 1396, 5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501,
         5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501],
        [5501, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396,
         1396, 5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501,
         5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501],
        [5501, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396,
         1396, 1396, 1396, 1396, 5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501,
         5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501],
        [5501, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396,
         5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501,
         5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501],
        [5501, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396,
         1396, 1396, 1396, 1396, 1396, 5501, 5501, 5501, 5501, 5501, 5501, 5501,
         5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501],
        [5501, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 5501,
         5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501,
         5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501],
        [5501, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396,
         1396, 1396, 1396, 1396, 1396, 1396, 5501, 5501, 5501, 5501, 5501, 5501,
         5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501],
        [5501, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 5501, 5501,
         5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501,
         5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501],
        [5501, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396,
         1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1575, 5501,
         5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501],
        [5501, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396,
         1396, 1396, 5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501,
         5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501],
        [5501, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396,
         1396, 5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501,
         5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501],
        [5501, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396,
         1396, 1396, 1396, 5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501,
         5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501],
        [5501, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396,
         5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501,
         5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501],
        [5501, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396,
         1396, 1396, 1396, 1396, 5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501,
         5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501],
        [5501, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 5501,
         5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501,
         5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501],
        [5501, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396,
         1396, 1396, 1396, 1396, 1396, 5501, 5501, 5501, 5501, 5501, 5501, 5501,
         5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501],
        [5501, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 5501, 5501,
         5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501,
         5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501],
        [5501, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396,
         1396, 1396, 1396, 1396, 1396, 1396, 5501, 5501, 5501, 5501, 5501, 5501,
         5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501],
        [5501, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396,
         1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1575, 4148, 4791, 1631,
         1640, 1737, 1527, 1694, 1701, 1584, 1885, 1670, 5501]],
       device='cuda:0')
ipdb> n
> /workspace/asr/wenet/wenet/transformer/asr_model.py(285)recognize()
    284         best_hyps = torch.index_select(hyps, dim=0, index=best_hyps_index)
--> 285         best_hyps = best_hyps[:, 1:]
    286         return best_hyps, best_scores

ipdb> best_hyps
tensor([[5501, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396,
         1396, 1396, 5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501,
         5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501],
        [5501, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396, 1396,
         1396, 1396, 5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501,
         5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501, 5501]],
       device='cuda:0')

这就是选择了两个最佳候选文本序列了。

-->285best_hyps=best_hyps[:,1:]
    286         return best_hyps, best_scores


后续就没啥难度了,无非是收集一下每个wav的top-1的文本序列;以及把结果写入文件。

这样, 这个自回归+beam search的"attention"解码方法,就算学习完毕了。

待续。

还有关于wer计算,以及语言模型的使用方面的。

也会对其他一些目前为止还没有涉及到的内容,查漏补缺。