ARTICLE · INTELLIGENCE

战地情报 · 详情页

来自尧图项目组的一线实战观察与深度解析

kws_streaming 关键词唤醒模型量化实战:基于 mfcc_op 的 12 标签 TFLite 量化训练与评估指南

kws_streaming 关键词唤醒模型量化实战:基于 mfcc_op 的 12 标签 TFLite 量化训练与评估指南 kws_streaming 关键词唤醒模型量化实战基于 mfcc_op 的 12 标签 TFLite 量化训练与评估指南【免费下载链接】google-researchGoogle Research项目地址: https://gitcode.com/gh_mirrors/go/google-research导读本文是 Google Research kws_streaming 项目中量化版 12 标签关键词唤醒KWS模型实验指南的完整解读对应仓库文档为 kws_experiments_quantized_12_labels.md。该文档基于论文Streaming keyword spotting on mobile devicesarXiv:2005.06720的模型体系展示了如何通过--feature_type mfcc_op让语音特征提取器以 TFLite 原生算子参与模型构建从而支持端到端的训练后量化PTQ并在 Speech Commands V2 数据集上复现 svdf、lstm_peep、crnn、crnn_state、dnn、att_mh_rnn 六类模型的 float/quant/stream 精度、模型大小与延迟指标。读完本文你将掌握从环境搭建、数据集准备、TFLite 基准工具编译到逐模型训练评估与量化结果解读的完整流程。一、背景为什么量化版实验要切换到 mfcc_op1.1 mfcc_tf 与 mfcc_op 的数值差异原始论文实验使用--feature_type mfcc_tf其特征提取基于 DFT 与 DCT二者通过矩阵乘法实现。这种方式虽然与任何推理引擎兼容但 DFT 权重会直接成为模型的一部分显著增加模型体积且对这类模型做训练后量化会带来明显的精度损失。量化版实验改用--feature_type mfcc_op语音 MFCC 特征提取器改为调用 TFLite 内部算子内部通过 FFT 执行 DFT/DCT数值上与 mfcc_tf 存在差异。文档明确提示由于没有针对 mfcc_op 重新做超参数优化精度可能出现一定下降。同时 mfcc_op 会调用audio_spectrogram()与mfcc()后者期望输入为平方后的 FFT 幅度因此所有命令中都必须显式设置--fft_magnitude_squared 1该参数在 model_flags.py 中被转换为布尔值。1.2 端到端模型preprocess raw本文档中的所有模型都同时指定--feature_type mfcc_op语音 MFCC 特征提取器使用 TFLite 内部算子--preprocess raw模型接收原始音频特征提取器作为模型的一部分。这样构建出的模型是端到端自包含的推理时只需向模型喂入原始音频即可直接得到分类结果无需在外部单独维护特征提取管线极大简化了移动端部署与测试。这一点在 svdf.py、att_mh_rnn.py 的源码中可以看到一致的实现模式当flags.preprocess raw时模型输入后直接接入speech_features.SpeechFeatures层。二、环境搭建与数据集准备2.1 获取 kws_streaming 代码# create main folder mkdir test # set path to a main folder KWS_PATH$PWD/test cd $KWS_PATH# copy content of kws_streaming to a folder /tmp/test/kws_streaming git clone https://github.com/google-research/google-research.git mv google-research/kws_streaming .2.2 安装 TensorFlow 与依赖# set up virtual env pip install virtualenv virtualenv --system-site-packages -p python3 ./venv3 source ./venv3/bin/activate # install TensorFlow, correct TensorFlow version is important pip install --upgrade pip pip install tf_nightly pip install tensorflow_addons pip install tensorflow_model_optimization # was tested on tf_nightly-2.3.0.dev20200515-cp36-cp36m-manylinux2010_x86_64.whl # install libs: pip install pydot pip install graphviz pip install numpy pip install absl-py文档特别强调TensorFlow 版本必须正确其验证过的版本为tf_nightly-2.3.0.dev20200515。tensorflow_model_optimization用于支持量化感知训练等高级量化方案见 README.md 中关于 functional API 与 subclass API 的 QAT 示例。2.3 数据集Speech Commands V22018kws_streaming 支持两版数据集V1 2017 与 V2 2018本文档统一使用 V2# download and set up path to data set V2 and set it up wget https://storage.googleapis.com/download.tensorflow.org/data/speech_commands_v0.02.tar.gz mkdir data2 mv ./speech_commands_v0.02.tar.gz ./data2 cd ./data2 tar -xf ./speech_commands_v0.02.tar.gz cd ../ # path to data sets V2 DATA_PATH$KWS_PATH/data22.4 模型输出目录# set up path for model training mkdir $KWS_PATH/models2_q # models trained on data V2 MODELS_PATH$KWS_PATH/models2_q完成上述步骤后KWS_PATH主目录下应包含如下结构kws_streaming/ colab/ data/ experiments/ ... data2 _background_noise_/ bed/ ... models2_q/ svdf/ ...2.5 编译 TFLite 基准测试工具Android若需在手机上评测延迟需使用 bazel 编译 TFLite 基准测试二进制# build benchmarking binary bazel build -c opt --configandroid_arm64 --cxxopt--stdc17 \ third_party/tensorflow/lite/tools/benchmark:benchmark_model # check that phone is connected adb devices # copy benchmarking binary to phone adb push bazel-bin/third_party/tensorflow/lite/tools/benchmark/benchmark_model /data/local/tmp # allow executing benchmarking file as a program adb shell chmod x /data/local/tmp/benchmark_model # build benchmarking binary - for a case if Flex is used by neural network model bazel build -c opt --configandroid_arm64 --cxxopt--stdc17 \ third_party/tensorflow/lite/tools/benchmark:benchmark_model_plus_flex # check that phone is connected adb devices # copy benchmarking binary to phone adb push bazel-bin/third_party/tensorflow/lite/tools/benchmark/benchmark_model_plus_flex /data/local/tmp # allow executing benchmarking file as a program adb shell chmod x /data/local/tmp/benchmark_model_plus_flex当模型用到 Flex 算子时需要使用benchmark_model_plus_flex变体。三、训练与评估的统一入口所有模型都通过 model_train_eval.py 统一驱动它注册了各模型的子解析器见 models.py 中MODELS字典与 model_train_eval.py 中按模型名注册model_parameters的代码。两种运行方式# CMD_TRAINbazel run -c opt --copt-mavx2 kws_streaming/train:model_train_eval -- CMD_TRAINpython -m kws_streaming.train.model_train_eval关键运行规则与 model_train_eval.py 的main()逻辑对应--train 1从头训练并评估若模型已训练好可设--train 0仅评估并生成 TFLite 精度报告重新训练时需删除$MODELS_PATH下对应的模型子文件夹否则会因目录已存在而报错训练结束后脚本会输出flags.json完整训练参数快照、TF/TFLite 非流式与流式模型精度报告、量化模型等产物。四、六大量化模型的完整训练命令与结果下文所有命令均省略统一前缀--data_url 、--data_dir $DATA_PATH/、--alsologtostderr等公共参数完整命令见 kws_experiments_quantized_12_labels.md这里给出每个模型的完整可执行命令。4.1 svdf参数量354Kfloat精度 96.0模型 1003KB延迟 2msquant精度 96.0模型 369KB延迟 1.6msstream float精度 96.0模型 1003KB延迟 0.4msstream quant精度 96.0模型 406KB延迟 0.4ms$CMD_TRAIN \ --data_url \ --data_dir $DATA_PATH/ \ --train_dir $MODELS_PATH/svdf/ \ --mel_upper_edge_hertz 7600 \ --how_many_training_steps 20000,20000,20000,20000 \ --learning_rate 0.001,0.0005,0.0001,0.00002 \ --window_size_ms 40.0 \ --window_stride_ms 20.0 \ --mel_num_bins 80 \ --dct_num_features 30 \ --resample 0.15 \ --alsologtostderr \ --time_shift_ms 100 \ --train 1 \ --feature_type mfcc_op \ --fft_magnitude_squared 1 \ svdf \ --svdf_memory_size 4,10,10,10,10,10 \ --svdf_units1 256,256,256,256,256,256 \ --svdf_act relu,relu,relu,relu,relu,relu \ --svdf_units2 128,128,128,128,128,-1 \ --svdf_dropout 0.0,0.0,0.0,0.0,0.0,0.0 \ --svdf_pad 0 \ --dropout1 0.0 \ --units2 \ --act2 svdf 各参数含义与默认值可在 svdf.py 的model_parameters中查到--svdf_memory_size表示各 SVDF 层在时间维保留的历史步数--svdf_units1/--svdf_units2分别表示 SVDF 模块前半部分与投影部分的单元数-1表示该层无投影--svdf_pad 0时使用 valid padding置 1 则用 causal padding源码中padding causal if flags.svdf_pad else valid流式模式更推荐 causal。4.2 lstm_peep参数量545Kfloat精度 97.3模型 2200KB延迟 10msquant精度 97.3模型 723KB延迟 3.8ms$CMD_TRAIN \ --data_url \ --data_dir $DATA_PATH/ \ --train_dir $MODELS_PATH/lstm_peep/ \ --mel_upper_edge_hertz 7600 \ --how_many_training_steps 20000,20000,20000,20000 \ --learning_rate 0.001,0.0005,0.0001,0.00002 \ --window_size_ms 40.0 \ --window_stride_ms 20.0 \ --mel_num_bins 40 \ --dct_num_features 20 \ --resample 0.15 \ --alsologtostderr \ --train 1 \ --lr_schedule exp \ --use_spec_augment 1 \ --time_masks_number 2 \ --time_mask_max_size 10 \ --frequency_masks_number 2 \ --frequency_mask_max_size 5 \ --feature_type mfcc_op \ --fft_magnitude_squared 1 \ lstm \ --lstm_units 500 \ --return_sequences 0 \ --use_peepholes 1 \ --num_proj 200 \ --dropout1 0.3 \ --units1 \ --act1 \ --stateful 0该模型启用 peephole 连接的 LSTM--use_peepholes 1并带投影--num_proj 200同时启用 SpecAugment--use_spec_augment 12 个时间掩码最大 10、2 个频率掩码最大 5提升鲁棒性。4.3 crnn参数量467Kfloat精度 97.4模型 1800KB延迟 7msquant精度 97.0模型 593KB延迟 2.6ms$CMD_TRAIN \ --data_url \ --data_dir $DATA_PATH/ \ --train_dir $MODELS_PATH/crnn/ \ --mel_upper_edge_hertz 7600 \ --how_many_training_steps 20000,20000,20000,20000 \ --learning_rate 0.001,0.0005,0.0001,0.00002 \ --window_size_ms 40.0 \ --window_stride_ms 20.0 \ --mel_num_bins 40 \ --dct_num_features 20 \ --resample 0.15 \ --alsologtostderr \ --train 1 \ --lr_schedule exp \ --use_spec_augment 1 \ --time_masks_number 2 \ --time_mask_max_size 10 \ --frequency_masks_number 2 \ --frequency_mask_max_size 5 \ --feature_type mfcc_op \ --fft_magnitude_squared 1 \ crnn \ --cnn_filters 16,16 \ --cnn_kernel_size (3,3),(5,3) \ --cnn_act relu,relu \ --cnn_dilation_rate (1,1),(1,1) \ --cnn_strides (1,1),(1,1) \ --gru_units 256 \ --return_sequences 0 \ --dropout1 0.1 \ --units1 128,256 \ --act1 linear,relu \ --stateful 0crnn 由两段卷积16/16 个滤波器核尺寸 (3,3) 与 (5,3)叠加 256 单元 GRU 构成最后接两段全连接层。4.4 crnn_state参数量467Kfloat精度 97.1模型 1800KB延迟 7.1msquant精度 96.9模型 593KB延迟 2.6msstream float精度 96.3模型 1700KB延迟 0.2msstream quant精度 95.8模型 472KB延迟 0.1ms$CMD_TRAIN \ --data_url \ --data_dir $DATA_PATH/ \ --train_dir $MODELS_PATH/crnn_state/ \ --mel_upper_edge_hertz 7600 \ --how_many_training_steps 20000,20000,20000,20000 \ --learning_rate 0.001,0.0005,0.0001,0.00002 \ --window_size_ms 40.0 \ --window_stride_ms 20.0 \ --mel_num_bins 40 \ --dct_num_features 20 \ --resample 0.15 \ --alsologtostderr \ --train 1 \ --lr_schedule exp \ --use_spec_augment 1 \ --time_masks_number 2 \ --time_mask_max_size 10 \ --frequency_masks_number 2 \ --frequency_mask_max_size 5 \ --feature_type mfcc_op \ --fft_magnitude_squared 1 \ crnn \ --cnn_filters 16,16 \ --cnn_kernel_size (3,3),(5,3) \ --cnn_act relu,relu \ --cnn_dilation_rate (1,1),(1,1) \ --cnn_strides (1,1),(1,1) \ --gru_units 256 \ --return_sequences 0 \ --dropout1 0.1 \ --units1 128,256 \ --act1 linear,relu \ --stateful 1crnn_state 与 crnn 拓扑相同唯一区别是--stateful 1有状态 GRU因此可以转换为流式推理模式20ms 音频包延迟低至 0.1–0.2ms是流式场景的代表模型。4.5 dnn参数量447Kfloat精度 90.4模型 1700KB延迟 1.2msquant精度 90.2模型 443KB延迟 1.1ms$CMD_TRAIN \ --data_url \ --data_dir $DATA_PATH/ \ --train_dir $MODELS_PATH/dnn/ \ --mel_upper_edge_hertz 7600 \ --how_many_training_steps 20000,20000,20000,20000 \ --learning_rate 0.001,0.0005,0.0001,0.00002 \ --window_size_ms 40.0 \ --window_stride_ms 20.0 \ --mel_num_bins 40 \ --dct_num_features 20 \ --resample 0.15 \ --alsologtostderr \ --train 1 \ --lr_schedule exp \ --use_spec_augment 1 \ --time_masks_number 2 \ --time_mask_max_size 10 \ --frequency_masks_number 2 \ --frequency_mask_max_size 5 \ --feature_type mfcc_op \ --fft_magnitude_squared 1 \ dnn \ --units1 64,128 \ --act1 relu,relu \ --pool_size 2 \ --strides 2 \ --dropout1 0.1 \ --units2 128,256 \ --act2 linear,reludnn 是纯全连接基线模型精度最低90.4/90.2但模型量化后体积从 1700KB 压缩到 443KB适合对精度要求不高的低资源场景。4.6 att_mh_rnn多注意力 BiRNN不可流式参数量700Kfloat精度 97.9模型 3400KB延迟 8msquant精度 97.8模型 1300KB延迟 4ms$CMD_TRAIN \ --data_url \ --data_dir $DATA_PATH/ \ --train_dir $MODELS_PATH/att_mh_rnn/ \ --mel_upper_edge_hertz 8000 \ --how_many_training_steps 20000,20000,20000,20000 \ --learning_rate 0.001,0.0005,0.0001,0.00002 \ --window_size_ms 40.0 \ --window_stride_ms 20.0 \ --mel_num_bins 40 \ --dct_num_features 20 \ --resample 0.15 \ --alsologtostderr \ --train 1 \ --lr_schedule exp \ --use_spec_augment 1 \ --time_masks_number 2 \ --time_mask_max_size 10 \ --frequency_masks_number 2 \ --frequency_mask_max_size 5 \ --feature_type mfcc_op \ --fft_magnitude_squared 1 \ att_mh_rnn \ --cnn_filters 10,1 \ --cnn_kernel_size (5,1),(5,1) \ --cnn_act relu,relu \ --cnn_dilation_rate (1,1),(1,1) \ --cnn_strides (1,1),(1,1) \ --rnn_layers 2 \ --rnn_type gru \ --rnn_units 128 \ --heads 4 \ --dropout1 0.2 \ --units2 64,32 \ --act2 relu,linearatt_mh_rnn 在六类模型中精度最高97.9/97.8结构为卷积 双向 RNN 多头注意力--heads 4即论文中的 multihead self-attention参数定义见 att_mh_rnn.py。由于使用双向 RNN该模型不可流式化model_train_eval.py 中将att_mh_rnn、att_rnn、tc_resnet显式列入non_streamable_models因此无 stream 指标。注意其--mel_upper_edge_hertz为 8000与其余模型不同。五、结果解读与量化收益分析5.1 指标总览模型参数量float 精度quant 精度float 模型/延迟quant 模型/延迟stream quant 精度/延迟svdf354K96.096.01003KB / 2ms369KB / 1.6ms96.0 / 0.4mslstm_peep545K97.397.32200KB / 10ms723KB / 3.8ms—crnn467K97.497.01800KB / 7ms593KB / 2.6ms—crnn_state467K97.196.91800KB / 7.1ms593KB / 2.6ms95.8 / 0.1msdnn447K90.490.21700KB / 1.2ms443KB / 1.1ms—att_mh_rnn700K97.997.83400KB / 1.3ms1300KB / 4ms—5.2 训练后量化的实现机制量化在评估脚本中由 TFLite 完成README.md 明确指出对应 model_train_eval.py 中的转换逻辑。从源码看name2opt字典为不量化与quantize_opt_for_size_tf.lite.Optimize.DEFAULT两条路径分别生成tflite_non_stream与quantize_opt_for_size_tflite_non_stream两套模型若未提供校准数据TFLite 会退化为混合量化int/float 混合运算提供校准数据后则为全整数量化转换入口为 test.py 的convert_model_tflite它加载best_weights后调用utils.model_to_tflite生成指定模式非流式 / 外部状态流式的 TFLite 模型流式模型会额外生成tflite_stream_state_external_model_accuracy_reset0.txt状态不重置模拟长时运行与reset1状态重置等价于非流式两份精度报告。5.3 关键结论量化精度损失极小得益于 mfcc_op六类模型量化后精度下降普遍不超过 0.5 个百分点svdf、lstm_peep 甚至无损验证了文档中mfcc_op 下量化精度损失可忽略的判断体积大幅压缩如 att_mh_rnn 从 3400KB 降至 1300KB约 -62%dnn 从 1700KB 降至 443KB约 -74%且模型体积仅由神经网络权重决定不含特征提取器权重流式量化进一步降低延迟crnn_state 从非流式 7.1ms 降至流式 0.1ms20ms 音频包粒度svdf 流式 float/quant 均为 0.4ms适用于实时唤醒词检测。六、注意事项与适用前提版本依赖以上命令基于文档验证的tf_nightly-2.3.0.dev20200515环境使用其他 TensorFlow 版本可能需要调整例如tf.compat.v1相关的会话/图模式 API见 model_train_eval.pymfcc_op 的数值差异与论文使用的 mfcc_tf 相比存在数值差异且未做超参优化如需更高的绝对精度可参考论文原版实验 kws_experiments_paper_12_labels.md 与最新实验 kws_experiments_12_labels.md后者在相同数据集上达到 96.4%–98.4% 精度不可流式模型att_mh_rnn双向 RNN 全局注意力与 att_rnn、tc_resnet 无法转换为流式推理使用前需确认模型的可流式属性完整模型清单见 README.md训练产物解读每个模型的$MODELS_PATH/model/目录会生成flags.json训练参数快照、labels.txt、model_summary.txt、TF/TFLite 各模式精度报告与模型文件详细目录结构说明见 README.md更多实验量化版文档聚焦 12 标签 Speech Commands V230K 参数量小模型见 kws_experiments_30k_12_labels.md35 标签自定义数据训练见 kws_experiments_35_labels.md。结语本文完整复现了 kws_streaming 量化版 12 标签关键词唤醒实验以mfcc_oppreprocess raw构建端到端可量化模型在 Speech Commands V2 上逐模型给出可复制的训练命令并结合源码说明了 TFLite 训练后量化、流式状态管理与不可流式模型判定的底层实现。无论是验证论文结论、为移动端部署选型还是作为自定义关键词数据集的起点这份指南都提供了完整的可操作路径。【免费下载链接】google-researchGoogle Research项目地址: https://gitcode.com/gh_mirrors/go/google-research创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
RELATED READING

延伸阅读

更多一线实战笔记与深度复盘,助您持续精进