目前的位置:
left_decoder和right_decoder已经执行完毕,结果返回:

开启loss计算之旅:

可以看到的是,目标序列是一次预测出来的。没有自回归。
举例子为:即一次性从y0 y1 y2 预测出来y1 y2 y3。
KL-loss
两个方向上的decoder完成之后,我们在这里:
self.criterion_att对象
[wenet/transformer/label_smoothing_loss.py] class LabelSmoothingLoss

下面这段代码还是非常经典的:
wenet/transformer/label_smoothing_loss.py里面的。
class LabelSmoothingLoss的forward函数:
defforward(self,x:torch.Tensor,target:torch.Tensor)->torch.Tensor:
"""Compute loss between x and target.
The model outputs and data labels tensors are flatten to
(batch*seqlen, class) shape and a mask is applied to the
padding part which should not be calculated for loss.
Args:
x (torch.Tensor): prediction (batch, seqlen, class)
target (torch.Tensor):
target signal masked with self.padding_id (batch, seqlen)
Returns:
loss (torch.Tensor) : The KL loss, scalar float value
"""
assert x.size(2) == self.size
batch_size = x.size(0)
x = x.view(-1, self.size)
target = target.view(-1)
# use zeros_like instead of torch.no_grad() for true_dist,
# since no_grad() can not be exported by JIT
true_dist = torch.zeros_like(x)
# self.smoothing = 1 - self.confidence
true_dist.fill_(self.smoothing / (self.size - 1))
ignore = target == self.padding_idx # (B,)
total = len(target) - ignore.sum().item()
target = target.masked_fill(ignore, 0) # avoid -1 index
true_dist.scatter_(1, target.unsqueeze(1), self.confidence)
# self.confidence + self.smoothing = 1.0
# x -> model output distribution, [84, 5502]
# true_dist = true reference distribution, [84, 5502] 经过了label smoothing
kl = self.criterion(torch.log_softmax(x, dim=1), true_dist) # KLDivLoss()
denom = total if self.normalize_length else batch_size # denom=12
return kl.masked_fill(ignore.unsqueeze(1), 0).sum() / denom在调用self.criterion (KLDivLoss())之前,两个参数的取值分别为:
ipdb>target[0]
tensor(1762, device='cuda:0')
ipdb> true_dist[0, 1760:1770]
tensor([1.8179e-05, 1.8179e-05, 9.0000e-01, 1.8179e-05, 1.8179e-05, 1.8179e-05,
1.8179e-05, 1.8179e-05, 1.8179e-05, 1.8179e-05], device='cuda:0')
ipdb> torch.log_softmax(x, dim=1)[0, 1760:1770]
tensor([-8.3035, -8.8391, -9.1146, -8.9233, -9.7420, -8.3117, -9.3750, -7.9063,
-8.9433, -9.4808], device='cuda:0', grad_fn=) 注意上面的ref的9.0000e-01=0.9的取值。
tensor(38.2568,device='cuda:0',grad_fn=) 
两个方向上的KL-div-loss的计算
tensor(38.6594,device='cuda:0',grad_fn=) 是right-decoder的结果对应的KL-div-loss。
-->145loss_att=loss_att*(
146 1-self.reverse_weight)+r_loss_att*self.reverse_weight两者加权相加之后,得到:
tensor(38.3776,device='cuda:0',grad_fn=) th_accuracy
[wenet/utils/common.py]
defth_accuracy(pad_outputs:torch.Tensor,pad_targets:torch.Tensor,
ignore_label: int) -> float:
"""Calculate accuracy.
Args:
pad_outputs (Tensor): Prediction tensors (B * Lmax, D).
pad_targets (LongTensor): Target label tensors (B, Lmax, D).
ignore_label (int): Ignore label id.
Returns:
float: Accuracy value (0.0 - 1.0).
"""
pad_pred = pad_outputs.view(pad_targets.size(0), pad_targets.size(1),
pad_outputs.size(1)).argmax(2)
mask = pad_targets != ignore_label
numerator = torch.sum(
pad_pred.masked_select(mask) == pad_targets.masked_select(mask))
denominator = torch.sum(mask)
return float(numerator) / float(denominator)
这个的逻辑还是比较容易理解的。
首先是从[12,7,512]中选择dim(2)最大的那个index,然后那个词就作为模型预测的输出。
然后得到[12,7]和reference相同,直接比较token.id就可以了。
注意是有关于长度的mask,那些长度不够本batch max-len的,pad=-1的位置,就不计算在内了。
ctc
[wenet/transformer/ctc.py] class CTC
ipdb>self.ctc
CTC(
(ctc_lo): Linear(in_features=512, out_features=5502, bias=True)
(ctc_loss): CTCLoss()
)上面是CTC的定义,里面有个linear,是从512映射到5502,然后是torch自己的CTCLoss。

需要注意的是,CTCLoss接受的是(sequence.length, batch.size, vocab.size)这样的输入。所以上面有
【13=输入frame相关的长度, 12=批大小, 5502=词表大小】
CTCLoss - PyTorch 1.11.0 documentation

最后是loss_att和loss_ctc的加权求和:
ipdb>n
> /workspace/asr/wenet/wenet/transformer/asr_model.py(116)forward()
115 else:
--> 116 loss = self.ctc_weight * loss_ctc + (1 -
117 self.ctc_weight) * loss_att得到的是:
ipdb>loss,loss_att,loss_ctc
(tensor(57.6924, device='cuda:0', grad_fn=),
tensor(38.3776, device='cuda:0', grad_fn=),
tensor(102.7603, device='cuda:0', grad_fn=)) 优化前进一步:

前进一步,优化
上面是囊括了,model.forward,以及loss的获取,optimizer前进一步,lr scheduler也是前进一步。
输出Log。包括三种loss等。
目前的位置:
【wenet/bin/train.py】 刚才执行完毕了一个epoch的train之后,
结下来就是对val set进行cv操作了。(交叉检验)。

一个epoch的训练结束,现在开始executor.cv操作。

上面的脑图,涵盖了cv之后,以及每个epoch之后的保存checkpoint的逻辑。
然后是所有的epoch结束之后,搞出来final.pt。
至此,除了executor.cv,其他的逻辑都算是比较直接的。
os.symlink() 方法用于创建一个软链接,即把最后的例如199.pt加个别名final.pt。
cv
[wenet/utils/executor.py] 里面的cv方法。
下面的整体逻辑,和train的基本类似。
defcv(self,model,data_loader,device,args):
''' Cross validation on
'''
import ipdb; ipdb.set_trace()
model.eval()
rank = args.get('rank', 0)
epoch = args.get('epoch', 0)
log_interval = args.get('log_interval', 10)
# in order to avoid division by 0
num_seen_utts = 1
total_loss = 0.0
with torch.no_grad(): # 不要梯度!
for batch_idx, batch in enumerate(data_loader):
key, feats, target, feats_lengths, target_lengths = batch
# 得到一个evaluation set的batch
feats = feats.to(device)
target = target.to(device)
feats_lengths = feats_lengths.to(device)
target_lengths = target_lengths.to(device)
num_utts = target_lengths.size(0) # 样本数量
if num_utts == 0:
continue
loss, loss_att, loss_ctc = model(feats, feats_lengths, target,
target_lengths) # 三个loss
if torch.isfinite(loss):
num_seen_utts += num_utts
total_loss += loss.item() * num_utts
if batch_idx % log_interval == 0:
log_str = 'CV Batch {}/{} loss {:.6f} '.format(
epoch, batch_idx, loss.item())
if loss_att is not None:
log_str += 'loss_att {:.6f} '.format(loss_att.item())
if loss_ctc is not None:
log_str += 'loss_ctc {:.6f} '.format(loss_ctc.item())
log_str += 'history loss {:.6f}'.format(total_loss /
num_seen_utts)
log_str += ' rank {}'.format(rank)
logging.debug(log_str)
return total_loss, num_seen_utts整个逻辑,没啥说的了。
本来以为是cv = cross validation,其实就是标准的validation。
至此,整个训练流程就算学习完毕了。
