目前的位置:

left_decoder和right_decoder已经执行完毕,结果返回:


decoder执行结束,结果返回

开启loss计算之旅:


decoder执行完毕;现在需要根据model output和reference text来计算loss了

可以看到的是,目标序列是一次预测出来的。没有自回归。

举例子为:即一次性从y0 y1 y2 预测出来y1 y2 y3。

KL-loss

两个方向上的decoder完成之后,我们在这里:

self.criterion_att对象

[wenet/transformer/label_smoothing_loss.py] class LabelSmoothingLoss


基于label smoothing的KL loss的计算

下面这段代码还是非常经典的:

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。


ctc计算的基本逻辑,这个主要是调用外部的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的逻辑

上面的脑图,涵盖了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。

至此,整个训练流程就算学习完毕了。