文章来源于新一代Kaldi,作者NGK编辑部
本文介绍新一代 Kaldi 中模型平均:
相关代码:
https://github.com/k2-fsa/icefall/tree/master/egs/librispeech/ASR/pruned_transducer_stateless4
原来的模型平均策略
decode.py解码时,除了需要提供参数--epoch来指定要读取的模型外,还会搭配参数--avg来实现模型平均。epoch-*.pt。这么做一方面是为了方便用户在训练中断时继续训练,另一方面则是为了在解码时实现模型平均。decode.py解码时,指定参数--epoch 24 --avg 3, 我们会读取文件epoch-24.pt、epoch-23.pt、epoch-22.pt所保存的模型,利用它们求得一个平均模型来解码,如下图所示。
模型平均的目的
存在的局限性
每个 epoch 采样一个模型,中间经过了很多个 batch,这种采样方式可能过于稀疏。 如果每个 epoch 对训练数据的遍历顺序不是随机的,那么每个 epoch 结束时所保存的模型采样点,可能会“记住”了数据遍历的顺序。使用这些采样点来进行模型平均,显然不是我们所期望的。
潜在的解决方案
通过使用更密集的采样点,可以得到一个噪声(随机性)更低的平均模型。 用来作模型平均的采样点,覆盖了每个 epoch 中数据遍历的不同阶段,这样就不用担心前面提到的模型可能会“记住”了数据遍历顺序的问题。
训练时,我们需要保存大量的模型,例如每 100 个 batch 保存一个; 解码时,为了实现模型平均,我们需要读取大量的模型文件。
改进后的模型平均策略
训练中维护平均模型model_avg

为采样的 batch 下标集合。
,更新平均模型,得到
:
,就是 batch 下标集合
所对应的模型采样点的均值,即
:
train.py中的参数--average-period,其默认值为 100,即每 100 个 batch 作一次采样。可参考代码:
https://github.com/k2-fsa/icefall/blob/master/egs/librispeech/ASR/pruned_transducer_stateless4/train.py
epoch-*.pt中,除了保存当前的模型之外,也会保存当前所维护的平均模型epoch-*.pt个数并没有增多。解码时利用平均模型 model_avg
:epoch- 保存着平均模型
,其经过了p 次采样;epoch- 保存着平均模型
,其经过了q 次采样。
, 我们只需要根据
和
两个模型计算得到:
decode.py解码时,指定参数--epoch 24 --avg 3 --use-averaged-model 1, 我们会读取文件epoch-24.pt和epoch-21.pt中分别保存的平均模型
可参考代码:
https://github.com/k2-fsa/icefall/blob/master/egs/librispeech/ASR/pruned_transducer_stateless4/decode.py
实验结果
full librispeech数据集上使用Reworked Conformer训练 20 个 epoch,然后比较以下三种方式,在test-clean和test-other两个测试集上,对解码结果的影响:epoch-20:使用 epoch-20 的模型解码,即指定参数 --epoch 20;epoch-20-avg-5:使用 epoch-16 ~ epoch 20 这个五个模型的平均模型解码,即指定参数 --epoch 20 --avg 5;epoch-20-avg-5-use-averaged-model:使用 epoch-15 ~ epoch 20 这个区间内每隔 100 个 batch 作采样的平均模型解码,即指定参数 --epoch 20 --avg 5 --use-averaged-model 1。
| Decoding model | WER on test-clean (%) | WER on test-other (%) |
|---|---|---|
| epoch-20 | 3.34 | 8.18 |
| epoch-20-avg-5 | 2.93 | 7.1 |
| epoch-20-avg-5-use-averaged-model | 2.82 | 7.0 |
注意:本文没有讨论文件 checkpoint-*.pt,与文件epoch-*.pt的区别在于,它是默认每 8000 个 batch 保存一个文件,方便用户通过指定参数--iter使用 epoch 中间的模型。请注意区分 checkpoint-*.pt和上述的每 100 个 batch 更新一次平均模型的区别。上述的模型平均策略同样应用于文件checkpoint-*.pt。
