NeMo Speech 说话人分离模型详解:Sortformer 端到端分离、Sort Loss 与流式 AOSC 机制
【免费下载链接】SpeechA scalable generative AI framework built for researchers and developers working on Large Language Models, Multimodal, and Speech AI (Automatic Speech Recognition and Text-to-Speech)项目地址: https://gitcode.com/GitHub_Trending/nem/Speech
本文围绕 NeMo Speech(nem/Speech 仓库)ASR 集合中的说话人分离(Speaker Diarization)模型体系展开,覆盖官方文档docs/source/asr/speaker_diarization/models.rst讲解的两大技术路线——端到端 Sortformer Diarizer 与级联(Cascaded)Clustering Diarizer,并结合nemo/collections/asr下的真实源码,深入剖析 Sortformer 的 Sort Loss/ATS 训练机制、流式版 Arrival-Order Speaker Cache(AOSC)实现,以及可复制的推理调用方式。读完后,你将理解两类分离系统各自的架构与适用边界,并能在源码层面定位关键配置参数与调用链。
两大分离系统总览
官方文档将 NeMo Speech AI 当前支持的说话人分离系统划分为两类(见 models.rst):
- 端到端说话人分离(End-to-end Speaker Diarization):Sortformer DiarizerSortformer 是基于 Transformer 编码器的端到端分离模型,直接从输入音频片段生成预测的说话人标签。官方同时提供离线(offline)与在线(online/streaming)两个版本,且流式版本也可以通过设置足够大的 chunk 尺寸来用于离线分离。
- 级联(Cascaded / Pipelined)说话人分离:Clustering Diarizer基于聚类的分离流水线由三段组成:先使用 MarbleNet 模型做语音活动检测(VAD),再用 TitaNet 模型抽取说话人嵌入,最后进行谱聚类(spectral clustering)。
从源码结构看,级联路线对应 ClusteringDiarizer 类,其 docstring 明确写出职责链为“Speech Activity Detection、Segmentation、Extract Embeddings、Clustering、Resegmentation and Scoring”。该类在初始化时依据cfg.diarizer.vad.model_path加载 VAD 模型(支持本地.nemo文件或预训练模型名,找不到时回退到vad_telephony_marblenet),并分别保存speaker_embeddings.parameters与clustering.parameters两组参数,聚类动作最终调用工具函数perform_clustering(见 clustering_diarizer.py 导入段)。对应的端到端路线则由 SortformerEncLabelModel 承担。两条路线的推理入口也做了统一抽象:Sortformer 通过SpkDiarizationMixin.diarize()模板方法驱动,聚类侧则通过DiarizationMixin。
Sortformer 的定位:从分离模型到多说话人 ASR 的监督组件
文档用较大篇幅解释了 Sortformer 诞生的背景:说话人分离要解决“谁在何时说话”的问题,而多说话人 ASR(又称 speaker-attributed / multitalker ASR)在此基础上要求转写文本同时带上说话人归属。文档指出两个现实难点:
- 数据稀缺:需要大量高质量、带精确时间戳的多说话人标注数据,采集难度远高于单人数据集,在低资源语言与受隐私法规约束的领域(如医疗)尤为突出;
- 长音频需求:真实场景常需处理长达数小时的音频,长时长数据的获取与标注进一步放大了供需差距。
为此,文档介绍了 Sortformer 的两项核心设计:
- Arrival Time Sort(ATS,到达时间排序):把 ASR 输出中的说话人 token 与分离输出中的说话人时间戳按“到达时间”排序来消解说话人排列(permutation)。这样多说话人 ASR 系统就能直接用 token 级交叉熵损失训练或微调,而不再依赖基于时间戳或帧级目标的 PIL 损失函数。
- Sort Loss:一种为 Transformer 模型生成梯度、使其学会按到达时间顺序(ATO)输出说话人标签时间戳的损失方法,用于训练神经分离器本身。
与传统端到端分离系统(如 EEND-SA、EEND-EDA 等,文档以外部文献链接引用)的关键区别在于类出现矩阵(class presence,文档记作 $\hat{\mathbf{Y}}$)的组织方式:PIL 通过搜索使预测与目标之间损失最小的排列来计算,而 Sort Loss 直接比较预测与目标“按到达时间排序后”的说话人活动序列。文档特别指出,相同的真实标签在 Sort Loss 与 PIL 下可能产生不同的目标矩阵。
在源码中,这两种目标矩阵的构造分别由get_ats_targets_hungarian与get_pil_targets_hungarian完成(见 sortformer_diar_models.py 导入)。训练/验证步_get_aux_train_evaluations会同时计算两套目标下的损失,并按混合权重组合(见 损失计算实现):
targets_ats, _ = get_ats_targets_hungarian(targets, preds, tolerance=self.ats_tolerance) targets_pil, _ = get_pil_targets_hungarian(targets, preds) ats_loss = self.loss(probs=preds, labels=targets_ats, target_lens=target_lens) pil_loss = self.loss(probs=preds, labels=targets_pil, target_lens=target_lens) loss = self.ats_weight * ats_loss + self.pil_weight * pil_loss对应的混合权重在 _init_loss_weights 中解析,均可在模型配置中指定:
| 配置键 | 默认值 | 含义 |
|---|---|---|
pil_weight | 0.0 | PIL 目标损失权重 |
ats_weight | 1.0 | ATS(Sort Loss)目标损失权重,两者之和不可为 0 |
ats_tolerance | 0 | 构造 ATS 目标时允许的到达时间容差(非负) |
high_resolution | False | 是否以预下采样帧率输出预测 |
output_subsampling_factor | 取encoder.subsampling_factor(默认 8) | 每个输出预测对应的 10ms 特征帧数,必须是模型原生下采样因子的整数倍(见 _resolve_output_resolution) |
streaming_mode/async_streaming | False/False | 流式推理开关与异步流式(ragged 行按最大容量 padding)开关 |
max_batch_dur | 20000 | 触发 OOM 安全分块特征提取的批内总时长(秒)上限 |
文档还强调 Sortformer 与 ASR 编码器直接集成:把说话人监督数据以 speaker kernels 的形式嵌入 ASR 编码器状态中,使说话人信息与转写信息统一处理。由此带来两个工程收益:多说话人 ASR 训练阶段无需专用损失计算函数,可直接复用单说话人 ASR 的标准训练框架;且 Sortformer 本身也可作为独立的端到端分离器使用——尤其在带精确时间戳的高质量模拟数据上训练后,只要把它作为Speaker Supervision模块挂入多说话人 ASR 的计算图,就能提升多说话人 ASR 性能(见 主数据流示意图)。
SortformerEncLabelModel 源码级解析
SortformerEncLabelModel 同时继承ModelPT、ExportableEncDecModel与SpkDiarizationMixin,其 docstring 说明模型配置需包含四个部分:preprocessor(特征前端)、encoder(FastConformer 编码器)、sortformer_modules(Sortformer 特有模块)、以及可选的transformer_encoder。
预训练模型清单。list_available_models()(L67-L98)注册了三个可用模型,与文档中“离线 + 流式”两条产品线一一对应:
| 预训练模型名 | 说明 |
|---|---|
diar_sortformer_4spk-v1 | 离线(4 说话人)Sortformer 分离器 |
diar_streaming_sortformer_4spk-v2 | 流式(4 说话人)Sortformer 分离器 |
diar_streaming_sortformer_4spk-v2.1 | 流式版本的 v2.1 更新 |
前向推理链路(离线)。forward()的流程为:process_signal(波形归一化 + preprocessor 提特征,必要时走oom_safe_feature_extraction分块防 OOM)→frontend_encoder(FastConformer 编码,必要时经encoder_proj投影到 Sortformer 模块维度)→forward_infer(可选再过transformer_encoder,然后upsample_hidden与forward_speaker_sigmoids输出 sigmoid 概率),最后按output_subsampling_factor对齐并可选下采样(forward 实现)。注意process_signal中一个细节:非流式模式下波形会按批内最大值归一化,而流式模式跳过了该步骤(L598-L601),这是流式与离线行为的一个可验证差异。
推理输出处理。diarize()一键推理接口(L1394-L1434)复用混入类 SpkDiarizationMixin.diarize() 的模板:输入可以是音频文件路径、路径列表、.json/.jsonlmanifest 或 numpy 波形;_diarize_output_processing再调用predlist_to_timestamps与generate_diarization_output_lines把逐帧 sigmoid 概率转成 RTTM 风格的[begin, end, speaker]段。推理参数由 DiarizeConfig 承载:
@dataclass class DiarizeConfig: session_len_sec: float = -1 # 端到端分离会话时长上限(秒),-1 表示不限制 batch_size: int = 1 num_workers: int = 1 sample_rate: Optional[int] = None postprocessing_yaml: Optional[str] = None # 后处理参数 yaml 路径 verbose: bool = True include_tensor_outputs: bool = False # 是否同时返回原始 speaker 概率张量 max_num_of_spks: Optional[int] = None据此可写出最小推理调用(audio支持单路径/路径列表/manifest,返回格式见diarizedocstring:[[begin_seconds, end_seconds, speaker_index], ...]):
model = SortformerEncLabelModel.from_pretrained("diar_sortformer_4spk-v1") segments = model.diarize(audio="path/to/wav_or_manifest", batch_size=2, include_tensor_outputs=False, verbose=True)训练侧数据支持 NeMo 原生 manifest 数据集AudioToSpeechE2ESpkDiarDataset与 Lhotse 数据集LhotseAudioToSpeechE2ESpkDiarDataset(use_lhotse开关,见 数据加载配置);评估指标由MultiBinaryAccuracy分别按 PIL 与 ATS 两种目标统计 F1/Precision/Recall(L231-L241)。
流式 Sortformer Diarizer:AOSC 与 FIFO 缓存机制
文档对Streaming Sortformer的描述是:为处理实时音频,流式版把声音切成小的、带重叠的 chunk 逐段处理,并引入Arrival-Order Speaker Cache(AOSC),存储音频流中此前检测到所有说话人的帧级声学嵌入,使当前 chunk 中的说话人可与历史说话人比对,从而保证同一个人贯穿整个流保持同一标签。流式版在 Fast-Conformer 中增加了一个pre-encoder 层来生成说话人缓存,每步都会对缓存做过滤、仅保留高质量的缓存向量;除缓存管理外,流式架构与离线版一致。
这些概念在源码中可以逐一对应:
- 流式状态结构。
SortformerStreamingState(定义于 sortformer_modules.py)持有spkcache / spkcache_lengths(到达序说话人缓存)、fifo / fifo_lengths(最近若干 chunk 的嵌入队列)及各自的预测spkcache_preds / fifo_preds、压缩状态spkcache_compressed、以及缓存排序信息spk_perm(见 forward_streaming_step 文档串)。 - 单步推理流程。forward_streaming_step 每步执行:pre-encode 当前 chunk → 将
spkcache + fifo + chunk拼接(同步模式用concat_embs,异步模式用concat_and_pad)→frontend_encoder(bypass_pre_encode=True)编码拼接序列 →forward_infer输出全段预测 →streaming_update/streaming_update_async更新缓存与 FIFO 状态 → 把本 chunk 有效区间的预测chunk_preds拼接到total_preds。left_offset/right_offset参数实现了文档所说的“重叠 chunk”(左右上下文)机制。 - 因果注意力的训练技巧。训练阶段,forward_streaming 会以
causal_attn_rate的概率把编码器的注意力窗口临时切换为因果形式(att_context_size = [-1, causal_attn_rc]),以此模拟流式时只能看到过去帧的约束,训练结束再恢复全上下文——这解释了离线架构如何复用于流式。 - 导出接口。流式模型导出(ONNX)的图输入为
chunk, chunk_lengths, spkcache, spkcache_lengths, fifo, fifo_lengths,输出为spkcache_fifo_chunk_preds, chunk_pre_encode_embs, chunk_pre_encode_lengths(forward_for_export 与 input_names/output_names)。streaming_input_examples依据chunk_left_context + chunk_len + chunk_right_context、spkcache_len、fifo_len等模块参数自动生成与模型尺寸匹配的示例张量(L679-L712),部署时可直接用model.streaming_export(output="...")触发。 - 参数一致性校验。启用
streaming_mode时,初始化会调用_check_streaming_parameters(),要求 chunk 预测长度与输出下采样因子整除对齐(L128-L142),这与文档“流式版可通过足够大的 chunk 尺寸做离线分离”的说法一致——chunk 参数正是决定延迟与离线复用能力的关键旋钮。
文档还配有一张三说话人实时分离动画热图,展示当前 chunk 中说话人活动如何被检测并更新进 AOSC 与 FIFO 队列(aosc_3spk_example.gif),以及逐步推理数据流图(streaming_steps.png)。
实操入口:训练、推理与测试脚本
围绕这两类模型,仓库提供了可直接运行的示例与测试入口:
- 分离示例脚本目录examples/speaker_tasks/diarization/:包含 5 个 Python 脚本与 9 个 YAML 配置,覆盖端到端分离器推理、流式分离推理、聚类分离器推理及多说话人数据仿真等场景;其中聚类路线的离线推理入口为 offline_diar_infer.py。
- 功能测试脚本tests/functional_tests/:
L2_Speaker_dev_run_EndtoEnd_Diarizer_Inference.sh、L2_Speaker_dev_run_EndtoEnd_Streaming_Diarizer_Inference.sh、L2_Speaker_dev_run_EndtoEnd_Speaker_Diarization_Sortformer.sh、L2_Speaker_dev_run_Clustering_Diarizer_Inference.sh、L2_Speaker_dev_run_Speaker_Diarization_with_ASR_Inference.sh分别验证了上述两类系统及其与 ASR 的组合链路。 - 模型级 e2e 测试tests/e2e_nightly/:
test_model_support_nvidia__diar_sortformer_4spk_v1.py与L2_Model_Support_nvidia__diar_streaming_sortformer_4spk_v2.sh(及 v2_1)对应上节列出的预训练模型名,可用于验证模型加载与推理链路。 - 单元测试tests/collections/speaker_tasks/:覆盖分离器模块与聚类工具的行为验证。
小结
NeMo Speech 的说话人分离体系以“端到端 Sortformer + 级联 Clustering”双轨布局:Sortformer 通过 Arrival Time Sort 与 Sort Loss 摆脱了 PIL 对排列搜索的依赖,使分离目标可用标准可微计算图训练,并天然与多说话人 ASR 的训练框架兼容;流式版则以 AOSC 说话人缓存加 FIFO 队列、pre-encoder 缓存过滤与因果注意力训练技巧,实现了跨 chunk 的说话人身份一致性。级联路线(MarbleNet VAD + TitaNet 嵌入 + 谱聚类)则提供了组件可替换、便于单独调优的经典方案。结合 models.rst 的概念讲解与 sortformer_diar_models.py、clustering_diarizer.py 的实现细节,可以完整地从原理、配置到推理/部署逐层理解当前仓库的分离能力边界与适用前提(如high_resolution、output_subsampling_factor与 chunk 参数需满足整除约束,流式与离线的归一化行为差异等)。
【免费下载链接】SpeechA scalable generative AI framework built for researchers and developers working on Large Language Models, Multimodal, and Speech AI (Automatic Speech Recognition and Text-to-Speech)项目地址: https://gitcode.com/GitHub_Trending/nem/Speech
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考