PaddleSpeech 关键词识别(KWS)实战指南:基于 MDTC 模型的命令行与 Python API 使用详解
【免费下载链接】PaddleSpeechEasy-to-use Speech Toolkit including Self-Supervised Learning model, SOTA/Streaming ASR with punctuation, Streaming TTS with text frontend, Speaker Verification System, End-to-End Speech Translation and Keyword Spotting. Won NAACL2022 Best Demo Award.项目地址: https://gitcode.com/gh_mirrors/pa/PaddleSpeech
关键词识别(Keyword Spotting, KWS)是语音技术中的一项核心任务,旨在从一段连续的语音中判定是否包含指定的唤醒词或关键词。本文以 PaddleSpeech 仓库中 demos/keyword_spotting 为骨架,完整讲解如何通过单条命令行或几行 Python 代码,使用 PaddleSpeech 预训练模型mdtc_heysnips对 WAV 音频执行关键词识别,并结合仓库源码深入剖析其调用链、特征提取流程与 MDTC 模型结构,帮助读者既会"开箱即用",也能理解底层实现原理。
一、KWS 任务与 Demo 概述
KWS(Keyword Spotting)是一项从给定语音音频中识别是否包含特定关键词的技术,是智能音箱、语音助手、低功耗唤醒等场景的核心模块。PaddleSpeech 仓库在 demos/keyword_spotting 中提供了一个完整的演示实现:
- 输入:一段 WAV 格式的语音文件;
- 输出:该语音与目标关键词的匹配得分(Score)、判定阈值(Threshold)以及是否命中关键词(Is keyword)的布尔结论;
- 使用方式:既可执行单条
paddlespeech kws命令行命令,也可以通过paddlespeech.cli.kws.KWSExecutor以 Python API 的方式在几行代码内完成调用。
该 demo 默认使用的模型为mdtc_heysnips,针对英文 "Hey Snips" 唤醒词训练,采样率为 16k。除此之外,仓库还提供了配套的完整训练与评估示例(见 examples/hey_snips/kws0),支持从零训练 MDTC 模型并在 HeySnips 数据集上评估 DET 指标。
二、环境安装
在使用 KWS demo 之前,需要先完成 PaddleSpeech 的安装。PaddleSpeech 提供了 easy、medium、hard 三种安装方式:
- easy:通过 pip 直接安装发布包,适合大多数仅需调用预训练模型的场景;
- medium:安装包含完整依赖(如 kenlm、numpy 等)的标准环境;
- hard:从源码编译安装,适合需要二次开发或使用最新特性的用户。
详细的安装步骤请参阅仓库内的安装文档(中文版见 install_cn.md),根据自身环境选择合适的方式即可。安装完成后,可以通过paddlespeech命令或 Python 导入来验证环境是否就绪。
三、准备输入音频
KWS demo 的输入需要满足两个约束:
- 必须是 WAV 格式文件(
.wav); - 采样率必须与模型一致。默认模型
mdtc_heysnips的采样率为 16k,因此输入的 WAV 文件采样率也应为 16k。若采样率不匹配,识别结果将不可靠,建议使用sox等工具对音频进行重采样。
仓库提供了两个官方示例音频,可直接下载用于体验:
wget -c https://paddlespeech.cdn.bcebos.com/kws/hey_snips.wav https://paddlespeech.cdn.bcebos.com/kws/non-keyword.wavhey_snips.wav:包含 "Hey Snips" 关键词的阳性样本;non-keyword.wav:不包含关键词的阴性样本。
这两个文件同样出现在 demo 目录的 run.sh 脚本中,可以直接bash run.sh一键完成下载与推理演示。
四、命令行使用(推荐)
4.1 基本命令
安装完成后,在包含音频文件的目录下执行:
paddlespeech kws --input ./hey_snips.wav paddlespeech kws --input ./non-keyword.wav对应输出如下:
# Input file: ./hey_snips.wav Score: 1.000, Threshold: 0.8, Is keyword: True # Input file: ./non-keyword.wav Score: 0.000, Threshold: 0.8, Is keyword: False可以看到,模型对包含关键词的音频给出接近 1.0 的高分(1.000 > 0.8,判定为命中),对不含关键词的音频给出 0.0 分(0.000 <= 0.8,判定为未命中)。
4.2 全部参数说明
执行paddlespeech kws --help可查看完整的命令行参数。结合 paddlespeech/cli/kws/infer.py 中KWSExecutor的参数定义,各参数含义如下:
| 参数 | 类型 | 默认值 | 说明 |
|---|---|---|---|
--input | str | 必填 | 用于关键词识别的音频文件路径 |
--threshold | float | 0.8 | 判定是否命中关键词的得分阈值,得分高于该值判定为包含关键词 |
--model | str | mdtc_heysnips | KWS 任务的模型类型,可选值由预训练模型列表动态生成(tag[:tag.index('-')]),当前为mdtc_heysnips |
--config | str | None | KWS 任务的配置文件(YAML),不设置时使用预训练模型自带的默认配置 |
--ckpt_path | str | None | 模型参数文件(checkpoint),不设置时自动下载并使用预训练模型权重 |
--device | str | paddle.get_device() | 执行推理的设备,默认取当前环境中 PaddlePaddle 的默认设备 |
-d, --job_dump_result | flag | 关闭 | 将任务结果保存到文件 |
-v, --verbose | flag | 关闭 | 增加当前任务的 logger 输出信息 |
从源码实现看(infer.py),--model的 choices 并非硬编码,而是从task_resource.pretrained_models的键中动态截取'-'之前的模型名生成,因此预训练模型表更新后命令行选项也会自动扩展。
4.3 使用自定义配置与权重
当不满足于预训练模型默认配置时,可以通过--config与--ckpt_path指定自己训练或微调得到的模型。在 infer.py 的_init_from_path中可以看到两条路径:
- 使用预训练模型:
ckpt_path=None时,拼接资源标签model_type + '-16k',从云端下载模型,配置与权重均取自下载目录; - 使用本地模型:显式传入
config与ckpt_path时,两者被转换为绝对路径后直接加载,此时配置文件中的stack_num、stack_size、in_channels、res_channels、kernel_size、num_keywords、sample_rate、frame_shift、frame_length、n_mels等字段将决定模型结构与特征提取参数。
五、Python API 使用
除了命令行,PaddleSpeech 还提供面向开发者更友好的 Python API。核心入口是paddlespeech.cli.kws.KWSExecutor(其定义与导出见 paddlespeech/cli/kws/init.py 与 infer.py):
import paddle from paddlespeech.cli.kws import KWSExecutor kws_executor = KWSExecutor() result = kws_executor( audio_file='./hey_snips.wav', threshold=0.8, model='mdtc_heysnips', config=None, ckpt_path=None, device=paddle.get_device()) print('KWS Result: \n{}'.format(result))输出:
KWS Result: Score: 1.000, Threshold: 0.8, Is keyword: True__call__方法(infer.py)的执行流程可以归纳为四个阶段:
- 设备设置:
paddle.set_device(device); - 模型初始化:
_init_from_path(model, config, ckpt_path),构建 backbone 与分类头、加载权重并置为eval()模式; - 预处理与推理:
preprocess()完成音频加载与特征提取,infer()在paddle.no_grad()下执行前向计算得到 logits; - 后处理:
postprocess(threshold)将 logits 转换为可读的字符串结果并返回。
值得注意的是,该 executor 还支持批量输入:命令行模式下execute()会通过get_input_source()解析输入源,对多个输入逐一执行推理并将结果汇总(infer.py),任一输入出错时返回False并在结果中以异常类名: 错误信息的形式标注。
六、输出结果解读
KWS 推理的最终输出是一行形如Score: X.XXX, Threshold: Y, Is keyword: True/False的文本。其计算逻辑在 infer.py 的postprocess中:
kws_score = max(self._outputs['logits'][0, :, 0]).item() return 'Score: {:.3f}, Threshold: {}, Is keyword: {}'.format( kws_score, threshold, kws_score > threshold)关键点:
- 模型输出的
logits形状为(batch, time_steps, num_keywords),取[0, :, 0]表示对第一个样本、所有时间步、第一个关键词的输出取最大值,作为整段音频的得分; Score保留三位小数展示;- 判定规则为简单的大小比较:
Score > Threshold时Is keyword: True,否则为False; - 阈值
threshold是影响误报率(False Alarm)与漏报率(Miss)平衡的关键超参数:阈值越高,越不容易误触发,但可能漏掉真正包含关键词的语音;阈值越低则相反。实际落地时通常结合 DET 曲线(Detection Error Tradeoff)选择业务可接受的阈值。
七、预训练模型
PaddleSpeech 官方发布并内置到命令与 Python API 的 KWS 预训练模型如下:
| 模型 | 语言 | 采样率 |
|---|---|---|
| mdtc_heysnips | en | 16k |
目前官方预训练模型聚焦于英文 "Hey Snips" 唤醒词场景,采样率 16k。如需其他语言或自定义关键词,可参考仓库内的训练示例(见下文第九节)自行训练模型。
八、源码级原理剖析:从音频到得分
8.1 音频加载与特征提取
在preprocess阶段(infer.py),输入音频经过两条核心路径:
- 音频读取:通过
paddlespeech.audio.backends.soundfile_load加载 WAV 波形; - 特征提取:使用 Kaldi 风格的 FBank 滤波器组特征,即
paddlespeech.audio.compliance.kaldi.fbank,参数来自模型配置:
self.feature_extractor = lambda x: kaldi_fbank( x, sr=config['sample_rate'], frame_shift=config['frame_shift'], frame_length=config['frame_length'], n_mels=config['n_mels'])默认配置下(见 examples/hey_snips/kws0/conf/mdtc.yaml):采样率16000、帧移10ms、帧长25ms、80 维 Mel 滤波器组。提取后的特征经unsqueeze(0)增加 batch 维度后送入模型。
8.2 MDTC 模型结构
mdtc_heysnips对应的骨干网络是MDTC(Multi-scale Dilated Temporal Convolution),其实现位于 paddlespeech/kws/models/mdtc.py,整体由以下组件构成:
DSDilatedConv1d:深度可分离膨胀卷积(Depthwise Separable Dilated Conv1d),先用groups=in_channels的分组膨胀卷积捕获长时依赖,再以1x1逐点卷积混合通道,配合 BatchNorm,显著降低参数量;TCNBlock:单个时间卷积块,由两个卷积路径组成并带有残差连接(causal=True时对输入做因果裁剪,避免未来信息泄漏,适合流式/在线场景);TCNStack:按stack_size组、组内膨胀率2^l(l从 0 到stack_num-1)递增的方式堆叠多个TCNBlock,逐层扩大感受野;MDTC:整体骨干,包含一个预处理TCNBlock与stack_num个TCNStack,多个尺度(stack)的输出在时间维度对齐后求和融合;KWSModel:分类头,在 backbone 的隐藏表示上接nn.Linear(hidden_dim, num_keywords)线性层与 Sigmoid 激活,将输出归一化到(0, 1)区间,作为关键词存在的概率得分。
默认配置stack_num=3、stack_size=4、res_channels=32、kernel_size=5、num_keywords=1。训练过程中使用的损失与相关工具位于 paddlespeech/kws/models/loss.py,推理时为causal=True的因果模式。
九、从零训练与评估(进阶)
若想复现或训练自己的 KWS 模型,仓库提供了完整的 HeySnips 示例:examples/hey_snips/kws0/README.md。其使用步骤如下:
- 准备数据集:按照该 README 指向的 sonos/keyword-spotting-research-datasets 说明下载并解压 HeySnips 数据集,然后将
data_dir替换为实际路径; - 一键训练与评估:
CUDA_VISIBLE_DEVICES=0,1 ./run.sh conf/mdtc.yaml脚本通过stage/stop_stage控制执行阶段(脚本位于 examples/hey_snips/kws0/run.sh):
- stage 1:从零开始训练;
- stage 2:在测试集上评估模型,并计算所有触发阈值下的检测错误权衡(DET)指标;
- stage 3:绘制 DET 曲线用于可视化。
训练脚本与评分脚本分别位于 paddlespeech/kws/exps/mdtc/train.py、score.py 与 compute_det.py,路径环境由 examples/hey_snips/kws0/path.sh 提供。
配置文件关键项
conf/mdtc.yaml 是训练与推理共用的配置模板,按区块划分:
| 区块 | 关键参数(默认值) | 作用 |
|---|---|---|
| Data | dataset: 'paddleaudio.datasets:HeySnips'、data_dir | 指定数据集类与数据路径 |
| Network | num_keywords: 1、stack_num: 3、stack_size: 4、in_channels: 80、res_channels: 32、kernel_size: 5 | 定义 MDTC 网络结构与关键词类别数 |
| Feature | feat_type: 'kaldi_fbank'、sample_rate: 16000、frame_shift: 10、frame_length: 25、n_mels: 80 | 特征提取参数,推理阶段由 infer.py 读取使用 |
| Training | epochs: 100、batch_size: 100、learning_rate: 0.001、weight_decay: 0.00005、grad_clip: 5.0、checkpoint_dir等 | 训练超参数与日志/保存频率 |
| Scoring | checkpoint、score_file、stats_file、img_file | 评估阶段的权重路径与输出文件 |
十、常见注意事项
- 采样率对齐:输入 WAV 必须与模型采样率(16k)一致,否则特征与模型训练分布不匹配,导致得分失真;
- 阈值调节:
threshold=0.8为默认值,实际业务中应根据误报/漏报的代价,参考 DET 曲线调整; - 自定义模型:传入
--config与--ckpt_path时需保证配置中的网络参数与 checkpoint 匹配,且特征参数(frame_shift、frame_length、n_mels、sample_rate)应与训练时一致; - 设备指定:
--device支持cpu/gpu等 PaddlePaddle 设备标识,默认取环境中的paddle.get_device(); - 日志输出:默认关闭 verbose 日志,排查问题时可加
-v观察预处理与推理细节。
至此,读者应已掌握 PaddleSpeech KWS demo 从安装、数据准备、命令行/ Python API 推理,到输出解读、模型原理与训练评估的完整链路,可在自己的项目中直接落地关键词识别能力。
【免费下载链接】PaddleSpeechEasy-to-use Speech Toolkit including Self-Supervised Learning model, SOTA/Streaming ASR with punctuation, Streaming TTS with text frontend, Speaker Verification System, End-to-End Speech Translation and Keyword Spotting. Won NAACL2022 Best Demo Award.项目地址: https://gitcode.com/gh_mirrors/pa/PaddleSpeech
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考