GitHub:https://github.com/Xianchao-Wu/wekws是2022年5月31号fork的,目标还是逐行基于脑图来学习分析代码。

按照epoch循环

四大步了,训练,cv,保存checkpoint,调整学习率


准备进入每个batch

train的内部


一个batch的处理:

七大步,搞定一个batch的训练


这七步分别是:
  1. batch拿到;

  2. 调用model的forward,得到logits;

  3. 计算loss和accuracy;

  4. loss.backward()

  5. 剪切gradient,

  6. optimizer.step()

  7. 输出log信息。(一次batch之后的)。

核心当然是model的forward,以及具体的loss的计算过程了,其他的没啥特别的。


如何构造一个batch

wekws里面是使用了七个相关的processor,

  1. parse_raw,读取wav文件,构造waveform;
  2. filter,按照长度过滤;
  3. resample,重新调整采样率(如果有必要,即,如果wave本身的采样率和指定的采样率是一样的,那就没有必要重新调整采样率了);
  4. compute_fbank,计算(这里是40维度)的梅尔谱;
  5. shuffle,重新洗牌;
  6. batch,按照batch size,封装成一个个的batch;
  7. padding,按照一个batch内部的不同sequence的长度,来做关于长度的padding。
七个processor来逐次嵌套调用,从而构造出来一个batch


model的forward

model的forward的内部核心,6个步骤
上面的model内部的forward,共有6步:
  1. global_cmvn;就是类似x=(x-mean)/std;
  2. preprocessing,是有个线性层,其负责的是,40 -> 256;
  3. backbone,这个是调用了四个blocks的tcn model;
  4. classifier,负责把256 -> 2,
  5. sigmoid,每个预测结果都给概率化一下;
  6. return logits,维度是[256, 1482, 2],其中256=batch size,1482=“time-sensive” 序列长度;2=判别结果。即,1482个“点”上,每个点,是属于两个“唤醒词”的其中一个的概率。

backbone的四大天王


backbone的四个blocks,i=0,1,2,3

一个DsCnnBlock

具体看一个天王:【这个cnn模块,其实和conformer中用的cnn module;或者citrinet中的cnn block,有很多相似的地方。。。】

这里面的核心,就是调用cnn模块,以及搞下quant和dequant,最后还有一个y=y+x的残差连接


loss和acc的计算

loss的计算


这里计算一下loss和accuracy

展开这个max_pooling_loss看一下:


1 max_pooling_loss:


max pooling loss的内部,分成两个循环,一个是计算loss的,一个是计算accuracy的。


继续看第一个for循环:


i是遍历所有的sequence(一个batch内部);而j是遍历所有的唤醒词(id)


这里,j就代表了0号唤醒词,1号唤醒词;

然后target[i=0] 和 j,相等或者不等,有这两种情况。

target[i] 代表的是第i个序列(语音),的取值,可以是0,1,-1,其中0,1代表两个唤醒词,-1代表没有唤醒词。

如果target[i] != j,则说明预测失败;反过来,如果target[i] = j,则说明这个序列预测成功了。



2 target[i] != j


target[i]和j不相等的时候

这段代码的操作有点飘。。。感觉还是-log(max(p))。


3 target[i]=j


target[i]和j相等的时候,说明第i个序列,找到了第j个候选词


各种操作猛如虎,仔细一看是头猪。。。

-torch.log(max_prob)是正解。


accuracy的计算:


这是遍历256个序列,例如:一个序列的1482个点上,最大的max_p大于0.5而且正好取值最大的这个点的索引,就是target[i]的话,那就是预测成功了。


上面的决策变量是0.5。

两种情况。

第一种是:max_p > 0.5,说明有个唤醒词的概率>0.5,即这个唤醒词被识别出来了,如果正好是参考答案target[i]里面的那个唤醒词,那就算成功预测,打赏+1,即num_correct += 1

第二种是:max_p < 0.5,即没有任何一个候选词被选拔出来,那如果参考答案是target[i]=-1,即该序列里面本来就不包括“唤醒词”,那也算成功预测,打赏+1,即num_correct += 1。


cv

脑图


进入cv valid部分的逻辑

这里面,进入with torch.no_grad()看看:

处理一个batch的逻辑


valid 里面的对于一个batch的处理,主要是为了收集loss和acc


一个batch的样子


一个batch的细节取值


最后保存ckpt

最后,就是一个个的保存checkpoints了。


最后了,构造一个软链接, fnal.pt


中间的log信息如下:

log信息


先到这里,待续。其他的几个网络结构。

-rw-rw-r-- 1 1003 1003 848 May 31 02:01 ds_tcn.yaml
-rw-rw-r-- 1 1003 1003 762 May 31 02:01 gru.yaml
-rw-rw-r-- 1 1003 1003 870 May 31 02:01 mdtc.yaml
-rw-rw-r-- 1 1003 1003 923 May 31 02:01 mdtc_small.yaml
-rw-rw-r-- 1 1003 1003 848 May 31 02:01 tcn.yaml


后续逐步介绍。不过大概的框架算是有了的。