sherpa-onnx 中的 pyannote 说话人分割模型:PyanNet 架构、训练配置与 ONNX 推理实战
【免费下载链接】sherpa-onnxSpeech-to-text, text-to-speech, speaker diarization, speech enhancement, source separation, and VAD using next-gen Kaldi with onnxruntime without Internet connection. Support embedded systems, Android, iOS, HarmonyOS, Raspberry Pi, RISC-V, RK NPU, Axera NPU, Ascend NPU, x86_64 servers, websocket server/client, support 12 programming languages项目地址: https://gitcode.com/GitHub_Trending/sh/sherpa-onnx
导读
本文以 sherpa-onnx 仓库中 scripts/pyannote/segmentation/notes.md 为核心,系统讲解 pyannote 说话人分割(speaker segmentation)模型的完整技术画像:从config.yaml的训练配置、PyanNet 模型各层架构、specifications与hparams语义,到该模型在 sherpa-onnx 中被导出为 ONNX 后用于语音活动检测(VAD)与说话人日志(speaker diarization)推理的完整链路。读完本文,你将能理解 pyannote segmentation-3.0 模型的内部结构、输入输出张量语义与帧率换算关系,并能复现仓库中"PyTorch 模型 → ONNX 模型 → onnxruntime 推理"的完整流程。
一、背景:说话人分割模型在 sherpa-onnx 中的位置
sherpa-onnx 是一个无需联网、基于下一代 Kaldi 与 onnxruntime 的语音工具集,覆盖语音识别、语音合成、说话人日志、语音增强、声源分离与 VAD 等任务。说话人分割(segmentation)模型是整个说话人日志与语音活动检测能力的地基:它负责在帧级别判断"每个时刻有哪些说话人处于活跃状态"。
在 scripts/pyannote/segmentation 目录下,仓库提供了一整套围绕 pyannote segmentation-3.0 的工程化工具链:
- notes.md:记录模型训练配置、网络架构与关键参数,是理解模型内部结构的权威资料;
- export-onnx.py:将 PyTorch 权重导出为 ONNX(含 int8 动态量化版本)并写入推理所需的元数据;
- show-onnx.py:打印 ONNX 模型的输入输出张量签名;
- vad-onnx.py 与 vad-torch.py:分别用 ONNX Runtime 与 PyTorch 完成基于该分割模型的 VAD;
- speaker-diarization-onnx.py 与 speaker-diarization-torch.py:在此基础上完成完整的说话人日志。
此外,该模型在 sherpa-onnx 主库中也有对应消费端,例如 offline-speaker-diarization-c-api.c、offline-speaker-diarization-cxx-api.cc 与各语言示例,说明这一套导出链路最终服务于跨语言的生产级推理。
二、config.yaml:训练配置逐项解析
notes.md 中记录的训练配置完整如下(保持原样):
task: _target_: pyannote.audio.tasks.SpeakerDiarization duration: 10.0 max_speakers_per_chunk: 3 max_speakers_per_frame: 2 model: _target_: pyannote.audio.models.segmentation.PyanNet sample_rate: 16000 num_channels: 1 sincnet: stride: 10 lstm: hidden_size: 128 num_layers: 4 bidirectional: true monolithic: true linear: hidden_size: 128 num_layers: 2各配置项的含义与影响如下:
task.duration: 10.0:训练时每个样本的音频时长固定为 10 秒。这在导出脚本中有直接呼应——export-onnx.py 中window_size = int(model.specifications.duration) * 16000,即 10 秒 × 16 kHz = 160000 个采样点,对应 ONNX 模型的输入张量长度。task.max_speakers_per_chunk: 3:每个 10 秒片段中最多出现 3 个说话人,对应model.specifications.classes中的['speaker#1', 'speaker#2', 'speaker#3']。task.max_speakers_per_frame: 2:单帧内最多同时出现 2 个说话人,对应后文的powerset_max_classes=2,即分类问题采用"幂集标签"(powerset)编码,每帧最多允许两人重叠发声。model.sample_rate: 16000:模型工作于 16 kHz 采样率。hparams中sincnet的sample_rate: 16000与之对应,SincNet 滤波器组需要知道采样率以计算滤波器频率参数。model.num_channels: 1:单声道输入。sincnet.stride: 10:SincNet 卷积核的滑动步长,直接影响时间维度的下采样节奏。lstm.hidden_size: 128 / num_layers: 4 / bidirectional: true / monolithic: true:4 层双向 LSTM,隐藏单元 128;monolithic: true表示整个 LSTM 作为一个整体模块(对应print(model)中单一的LSTM(60, 128, ...),而非分层的 cell 列表)。linear.hidden_size: 128 / num_layers: 2:分类头之前的 2 层全连接投影。
三、PyanNet 模型架构深度解析
notes.md 记录了print(model)的完整输出,这是理解网络各层连接关系的权威快照:
PyanNet( (sincnet): SincNet( (wav_norm1d): InstanceNorm1d(1, eps=1e-05, momentum=0.1, affine=True, track_running_stats=False) (conv1d): ModuleList( (0): Encoder( (filterbank): ParamSincFB() ) (1): Conv1d(80, 60, kernel_size=(5,), stride=(1,)) (2): Conv1d(60, 60, kernel_size=(5,), stride=(1,)) ) (pool1d): ModuleList( (0-2): 3 x MaxPool1d(kernel_size=3, stride=3, padding=0, dilation=1, ceil_mode=False) ) (norm1d): ModuleList( (0): InstanceNorm1d(80, eps=1e-05, momentum=0.1, affine=True, track_running_stats=False) (1-2): 2 x InstanceNorm1d(60, eps=1e-05, momentum=0.1, affine=True, track_running_stats=False) ) ) (lstm): LSTM(60, 128, num_layers=4, batch_first=True, dropout=0.5, bidirectional=True) (linear): ModuleList( (0): Linear(in_features=256, out_features=128, bias=True) (1): Linear(in_features=128, out_features=128, bias=True) ) (classifier): Linear(in_features=128, out_features=7, bias=True) (activation): LogSoftmax(dim=-1) )结合该输出,网络可划分为四个阶段:
1. SincNet 前端(特征提取)
wav_norm1d:对原始波形做 InstanceNorm,稳定输入;Encoder(filterbank=ParamSincFB):可学习的 Sinc 滤波器组,将 1 通道波形映射为 80 个滤波器组特征;- 两级
Conv1d(80→60, 60→60),卷积核 5,步长 1,每级卷积后接InstanceNorm1d(norm1d中第 0 项对应 80 通道,第 1、2 项对应 60 通道); - 三级
MaxPool1d(kernel_size=3, stride=3)串联下采样,每次将时间维压缩 3 倍。
2. 双向 LSTM 序列建模LSTM(60, 128, num_layers=4, batch_first=True, dropout=0.5, bidirectional=True):输入 60 维(SincNet 输出通道数),4 层双向 LSTM,每层 128 个隐藏单元。由于是双向,后续线性层输入维度为 128 × 2 = 256。训练阶段使用 0.5 的 dropout。
3. 线性投影Linear(256→128)、Linear(128→128)两层全连接,将双向 LSTM 拼接后的 256 维压缩到 128 维。
4. 分类头与激活Linear(128→7)输出 7 类 logits,LogSoftmax(dim=-1)归一化为对数概率。7 类的构成与幂集标签(powerset)编码直接相关:1 类"无说话人(静音)" + 3 个单说话人类(speaker#1/2/3)+ 3 个双说话人组合类(1+2、1+3、2+3)= 1 + 3 + 3 = 7。这与max_speakers_per_chunk: 3、max_speakers_per_frame: 2的配置严丝合缝。
从 export-onnx.py 中的断言可以印证该结构:
model.dimension == 7(输出类别数);- 10 秒输入(
[1, 1, 16000*10])对应输出[1, 589, 7](589 帧 × 7 类); model.receptive_field.step == 0.016875(帧步长 270 采样,约 16.875 ms);model.receptive_field.duration == 0.0619375(感受野 991 采样,约 61.94 ms)。
四、specifications 与 hparams:模型契约
notes.md 记录的model.specifications是推理侧必须遵守的契约:
>>> list(model.specifications) [Specifications(problem=<Problem.MONO_LABEL_CLASSIFICATION: 1>, resolution=<Resolution.FRAME: 1>, duration=10.0, min_duration=None, warm_up=(0.0, 0.0), classes=['speaker#1', 'speaker#2', 'speaker#3'], powerset_max_classes=2, permutation_invariant=True)]关键字段语义:
problem=MONO_LABEL_CLASSIFICATION:单标签分类——每个时刻的输出在 7 类中取概率最大的一类。这一点在 vad-onnx.py 的to_multi_label中体现为np.argmax(y, axis=-1),先取 argmax 再通过幂集映射表还原为多说话人标签。resolution=FRAME:模型输出是帧级别的(而非 segment 级别的粗粒度)。duration=10.0:输入窗口时长 10 秒。classes=['speaker#1','speaker#2','speaker#3']:最多 3 个说话人。powerset_max_classes=2:单帧最多 2 人同时说话。permutation_invariant=True:标签具有置换不变性,这也是说话人日志中"说话人编号不跨段保持身份"这一设计的前提——跨片段一致性依赖下游 embedding 聚类完成。
model.hparams则是训练超参数的最终落盘结果:
"linear": {'hidden_size': 128, 'num_layers': 2} "lstm": {'hidden_size': 128, 'num_layers': 4, 'bidirectional': True, 'monolithic': True, 'dropout': 0.5, 'batch_first': True} "num_channels": 1 "sample_rate": 16000 "sincnet": {'stride': 10, 'sample_rate': 16000}注意:lstm中出现了dropout: 0.5与batch_first: True,与print(model)的 LSTM 行完全对应;monolithic: True在 hparams 中同样被记录,说明该超参同时影响模块组织方式与序列化格式。
五、从 PyTorch 到 ONNX:导出流程与元数据契约
理解了模型结构后,export-onnx.py 的导出逻辑就非常清晰了。其核心步骤为:
加载权重并断言结构契约:
Model.from_pretrained("./pytorch_model.bin")加载权重后,依次断言dimension==7、problem==MONO_LABEL_CLASSIFICATION、resolution==FRAME、duration==10.0、sample_rate==16000,并验证输入张量形状[1, 1, 160000]、输出张量形状[1, 589, 7]、感受野步长/时长换算(270 与 991 采样)。这些断言保证了导出产物与推理脚本假设完全一致。导出 ONNX:使用
opset_version=13,输入名为x(形状[1, 1, T]),输出名为y(形状[1, T, 7]),并声明动态轴——x的 batch 维(0)与时间维(2)、y的 batch 维(0)与时间维(1)均可变,因此导出模型可以处理任意时长的音频窗口。写入自定义元数据(export-onnx.py 的
add_meta_data):将推理阶段必需的参数以 key-value 形式写入 ONNX 模型:num_speakers=3、powerset_max_classes=2、num_classes=7;sample_rate=16000;window_size=160000(10 秒窗口);receptive_field_size=991、receptive_field_shift=270;model_type="pyannote-segmentation-3.0"、version="1"以及model_author、license等来源信息。
这些元数据正是 vad-onnx.py 与 speaker-diarization-onnx.py 运行时通过
model.get_modelmeta().custom_metadata_map读取的契约,从而无需在脚本中硬编码任何模型参数。生成 int8 动态量化版本:对导出的
model.onnx调用onnxruntime.quantization.quantize_dynamic(..., weight_type=QuantType.QUInt8)产出model.int8.onnx,用于在低算力设备上以更小体积、更快速度推理(以少量精度换取性能)。
六、ONNX 输入输出签名与帧率换算
preprocess.sh 先对导出的模型做onnxruntime.quantization.preprocess预处理,再用 show-onnx.py 打印签名。预处理后的输入输出为:
- 输入
x:tensor(float),形状[1, 1, T](batch、通道、采样点数); - 输出
y:tensor(float),形状[1, floor(...), 7]。
preprocess.sh 中记录了对输出时间维公式的完整推导(T为输入采样点数):
floor(floor(floor(floor(T/10 - 251/10)/3 - 2/3)/3)/3 - 8/3) + 1 = (T - 721)/270该推导揭示了三个对使用者至关重要的结论(脚本注释原文保留):
- 输入采样点数至少为 721,否则无法产生任何输出帧;
- 每输出一帧对应 270 个采样点(16 kHz 下约 16.875 ms,与
receptive_field_shift一致); - 若输入增加 270 个采样点(T + 270),则输出恰好多一帧。
据此可快速换算:10 秒输入 T = 160000 时,(160000 − 721)/270 ≈ 589.92,向下取整得 589 帧,与导出脚本中的断言example_output.shape == [1, 589, 7]完全吻合。这组"721 / 270 / 991"的数字,是整个推理端帧对齐、时间戳换算(receptive_field_shift / sample_rate作为每帧时长)的基础。
七、用 ONNX 模型做语音活动检测(VAD)
vad-onnx.py 是 ONNX 推理侧的完整示范,其流水线可作为"纯 Python + onnxruntime 消费该模型"的模板:
- 读取元数据:从 ONNX 的 custom metadata 读取
window_size、sample_rate、receptive_field_size、receptive_field_shift、num_speakers、powerset_max_classes、num_classes,其中window_shift = 0.1 * window_size,即分帧步长为 1 秒。 - 分帧:用
numpy.lib.stride_tricks.as_strided将整段音频切成(num_chunks, window_size)的视图(脚本注释提示也可用torch.Tensor.unfold实现),每帧 10 秒、帧移 1 秒;末尾不足一帧时补零并单独送入模型。 - 批量推理:以
batch_size=32循环调用 ONNX Runtime,得到(num_chunks, num_frames, num_classes)的输出。 - 幂集标签还原:
get_powerset_mapping按"1 个说话人 → 3 个单标签、2 个说话人 → 3 个两两组合标签"构建映射表(幂集最大为 3 时直接报错Unsupported),再经np.argmax将 7 类输出映射回(num_chunks, num_frames, num_speakers)的多标签形式。 - 帧级加权融合:对各 chunk 的预测按时间位置对齐,使用 Hamming 窗加权平均,得到整段音频统一的帧级活动得分。
- 活动段检测:以
onset=0.5、offset=0.5为阈值做状态机扫描,输出活动片段起止时间;时间戳通过receptive_field_shift / sample_rate(每帧时长)乘以帧序号、再叠加receptive_field_size / sample_rate * 0.5的中心偏移得到。
vad-torch.py 提供了对照实验:直接用pyannote.audio.pipelines.VoiceActivityDetection流水线处理同一音频,便于验证 ONNX 版本与原始 PyTorch 版本的输出是否一致。这正是 run.sh 中依次运行 torch 与 onnx 两个版本的目的。
八、进阶:基于分割模型的完整说话人日志
vad-onnx.py 只区分"有语音/无语音",而 speaker-diarization-onnx.py 在分割模型之上叠加了说话人 embedding 与聚类,输出"谁在什么时候说话":
- 分割与幂集还原:与 VAD 完全相同的分帧、推理、
argmax+ 映射表还原流程,得到帧级多说话人标签; - 每个 (chunk, speaker) 抽取音频段:将某说话人在该 chunk 内的活跃帧对应的原始波形拼接(少于 10 帧即约 0.2 秒的片段被跳过),通过 sherpa-onnx 的
SpeakerEmbeddingExtractor(配置见sherpa_onnx.SpeakerEmbeddingExtractorConfig,见 SpeakerEmbeddingExtractorConfig.kt 同构 API)计算说话人 embedding; - 聚类成身份:用 sherpa-onnx 的
FastClustering(FastClusteringConfig可指定num_clusters或threshold两种模式,脚本中默认num_clusters=2,并注释提示按需调整)对全部 embedding 聚类,得到说话人身份编号; - 重标注与后处理:按聚类结果重写标签矩阵,结合每帧说话人数(
speaker_count)排序输出;最终按onset=0.5、offset=0.5、min_duration_off=0.5、min_duration_on=0.3生成Segment列表,并用merge_segment_list(gap=0.5s)合并同说话人的邻近片段,输出形如00:00:01.200 --> 00:00:03.500 speaker_00的时间轴。
对照实现 speaker-diarization-torch.py 则基于 pyannote 官方SpeakerDiarization流水线(参数含clustering.method=centroid、min_cluster_size=12、threshold=0.7045654963945799与segmentation.min_duration_off=0.5),并演示了使用 ONNX 格式的 WeSpeaker embedding 模型替换在线 embedding 的写法,可作为 ONNX 版本正确性的参照基准。
九、端到端复现:run.sh 与测试资源
run.sh 给出了从零复现的完整命令序列:
# 1. 安装依赖 pip install pyannote.audio onnx onnxruntime # 2. 下载 PyTorch 权重与测试音频 # pytorch_model.bin:pyannote segmentation-3.0 权重 # lei-jun-test.wav:测试波形(run.sh 中实际下载) # 3. 导出 ONNX(含 int8 量化) ./export-onnx.py # 4. 预处理并打印模型签名 ./preprocess.sh # 5. 三路对照验证 ./vad-torch.py ./vad-onnx.py --model ./model.onnx --wav ./lei-jun-test.wav ./vad-onnx.py --model ./model.int8.onnx --wav ./lei-jun-test.wavREADME.md 则记录了仓库用于测试的波形文件来源与预处理方法,其中包含0-four-speakers-zh.wav(4 人中文录音)与多个英文双人测试音频,部分由原始 mp4/mp3 转换而来,转换命令示例(原文档记录):
# mp4 -> wav(单声道、16 kHz) ffmpeg -i ./fcf059e3-689f-47ec-a000-bdace87f0113.mp4 -ac 1 -ar 16000 ./2-two-speakers-en.wav # mp3 -> wav(重采样到 16k) sox ML16091-Audio.mp3 -r 16k 3-two-speakers-en.wav建议读者按"下载权重 → export-onnx.py → preprocess.sh → 三路推理对照"的顺序自行复现,用lei-jun-test.wav验证 ONNX 与 PyTorch 输出的一致性,再进一步用 speaker-diarization-onnx.py 体验完整说话人日志。
十、参考资料
notes.md 中列出了该模型与流水线的两篇原始论文(原文为外部链接,此处仅保留题名,完整链接见 scripts/pyannote/segmentation/notes.md):
- pyannote.audio 2.1 speaker diarization pipeline: principle, benchmark, and recipe
- pyannote.audio speaker diarization pipeline at VoxSRC 2023
总结
本文围绕 scripts/pyannote/segmentation/notes.md 的配置与架构记录,完整还原了 pyannote segmentation-3.0 模型的技术全貌:config.yaml中每个超参的语义、PyanNet"SincNet + 双向 LSTM + 线性投影 + 7 类分类头"的层级结构、specifications与hparams所定义的推理契约,并进一步结合 export-onnx.py、preprocess.sh、vad-onnx.py 与 speaker-diarization-onnx.py 讲解了从 PyTorch 权重到 ONNX 模型、再到 VAD 与说话人日志推理的完整工程链路。掌握"721 / 270 / 991"这三组关键数字与幂集标签还原逻辑,即可在任何支持 ONNX Runtime 的环境里独立部署并二次开发该模型。
【免费下载链接】sherpa-onnxSpeech-to-text, text-to-speech, speaker diarization, speech enhancement, source separation, and VAD using next-gen Kaldi with onnxruntime without Internet connection. Support embedded systems, Android, iOS, HarmonyOS, Raspberry Pi, RISC-V, RK NPU, Axera NPU, Ascend NPU, x86_64 servers, websocket server/client, support 12 programming languages项目地址: https://gitcode.com/GitHub_Trending/sh/sherpa-onnx
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考