学习第四个解码算法:解码算法
在脚本中对应的是:
decode_mode="attention"
怎么还是感觉和之前的有一定的重合的地方???[这个是特殊的!自回归+beam search]
这个允许batch_size > 1,所以我们设置为2.
整体流程

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

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搜索最佳文本序列。
这个算法涉及的代码组织的是非常漂亮的,值得学习!

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

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

上面的截屏是我们重点想学习的内容,思想很不错,代码还非常简洁。
如果简单用几句话概括就是,一个语音,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)
这个,就是根据准备好的,如下信息:
- encoder_out
- encoder_mask
- hyps
- hyps_mask
- cache=None
来调用decoder的forward_one_step函数,这个函数内部,就是调用left_decoder来解码。

forward_one_step的输入参数和细节过程
这里的forward_one_step里面,有:
- 目标文本序列的embed
- 遍历self.decoders的三层decoder layers,解码;
- 解码结束之后,调用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_sizeipdb>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)这个,就是根据准备好的,如下信息:
- encoder_out
- encoder_mask
- hyps
- hyps_mask
- cache=None
来调用decoder的forward_one_step函数,这个函数内部,就是调用left_decoder来解码。

forward_one_step的输入参数和细节过程。需要注意的是,x的形状应该是[20, 2, 512]。因为现在的hyps是长度为2了。
这里的forward_one_step里面,有:
- 目标文本序列(目前长度为2)的embed
- 遍历self.decoders的三层decoder layers,解码;
- 解码结束之后,调用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。

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。

因为经过这俩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_sizeipdb>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计算,以及语言模型的使用方面的。
也会对其他一些目前为止还没有涉及到的内容,查漏补缺。
