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


一个batch的处理:

batch拿到;
调用model的forward,得到logits;
计算loss和accuracy;
loss.backward()
剪切gradient,
optimizer.step()
输出log信息。(一次batch之后的)。
如何构造一个batch
wekws里面是使用了七个相关的processor,
parse_raw,读取wav文件,构造waveform; filter,按照长度过滤; resample,重新调整采样率(如果有必要,即,如果wave本身的采样率和指定的采样率是一样的,那就没有必要重新调整采样率了); compute_fbank,计算(这里是40维度)的梅尔谱; shuffle,重新洗牌; batch,按照batch size,封装成一个个的batch; padding,按照一个batch内部的不同sequence的长度,来做关于长度的padding。


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

一个DsCnnBlock

loss的计算

1 max_pooling_loss:

继续看第一个for循环:

这里,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

这段代码的操作有点飘。。。感觉还是-log(max(p))。
3 target[i]=j

各种操作猛如虎,仔细一看是头猪。。。
-torch.log(max_prob)是正解。
accuracy的计算:

上面的决策变量是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。
脑图

这里面,进入with torch.no_grad()看看:
处理一个batch的逻辑

一个batch的样子

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

中间的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后续逐步介绍。不过大概的框架算是有了的。
