
NeMo 说话人分离 API 参考ClusteringDiarizer 与 SortformerEncLabelModel 全解析【免费下载链接】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 仓库中的docs/source/asr/speaker_diarization/api.rst为骨架系统讲解该 API 页面覆盖的两类说话人分离Speaker Diarization, SD模型类——级联式ClusteringDiarizer与端到端SortformerEncLabelModel——以及支撑它们的DiarizationMixin/SpkDiarizationMixin混入类。读完本文你可以掌握每个关键方法diarize、process_signal、forward_infer、training_step等的职责、参数与调用链并能对照 实现源码 与 级联实现 完成推理、训练与微调的落地。一、API 总览文档结构对应的源码位置API 文档定义了四个 API 对象它们在源码中的位置如下表API 对象类型源码位置角色nemo.collections.asr.models.ClusteringDiarizer模型类clustering_diarizer.py级联式离线分离VAD 嵌入 聚类推理nemo.collections.asr.models.SortformerEncLabelModel模型类sortformer_diar_models.pySortformer 端到端分离训练、验证、推理nemo.collections.asr.parts.mixins.DiarizationMixin混入mixins.py级联分离器的公共工具manifest 构建等nemo.collections.asr.parts.mixins.diarization.SpkDiarizationMixin抽象混入diarization.py端到端分离器的diarize()模板方法骨架两者在继承关系上分工清晰ClusteringDiarizer继承torch.nn.Module、Model与DiarizationMixin见 类定义SortformerEncLabelModel则继承ModelPT、ExportableEncDecModel与SpkDiarizationMixin见 类定义因此前者是推理容器后者是可训练、可导出的神经网络。二、ClusteringDiarizer级联式离线分离的完整流程2.1 类定位与构造参数ClusteringDiarizer是离线说话人分离的推理模型类其文档字符串明确说明它负责分离流程的全部环节Speech Activity Detection、Segmentation、Extract Embeddings、Clustering、Resegmentation 和 Scoring所有参数均通过配置文件传入见 类注释。构造函数签名def __init__(self, cfg: Union[DictConfig, Any], speaker_modelNone):cfgOmegaConf 配置。若传入DictConfig会先经model_utils.convert_model_config_to_dict_config与maybe_update_config_version转换以兼容 Hydra 1.0 实例化源码。speaker_model可选的已加载说话人嵌入模型实例不传时从cfg.diarizer.speaker_embeddings.model_path加载。构造过程中从配置解析出三组参数self._diarizer_params self._cfg.diarizer总控参数、self._speaker_params嵌入提取窗口参数、self._cluster_params聚类参数。VAD 与说话人模型的加载规则源码 L103-L158VAD当oracle_vadFalse且vad.model_path非空时初始化。model_path以.nemo结尾时走EncDecClassificationModel.restore_from本地加载否则按预训练模型名从 NGC 拉取若请求的名称不可用则回退到vad_telephony_marblenet并打印警告。说话人嵌入模型支持三种来源——传入的speaker_model对象、.nemo文件restore_from、.ckpt文件load_from_checkpoint、或 NGC 预训练名不可用时回退到ecapa_tdnn。多尺度参数parse_scale_configs(window_length_in_sec, shift_length_in_sec, multiscale_weights)将嵌入提取的窗口/步长配置解析为多尺度字典供后续逐尺度做子分段与嵌入聚合。2.2 diarize()主推理入口diarize()是整个级联流程的驱动函数源码def diarize(self, paths2audio_files: List[str] None, batch_size: int 0)paths2audio_files音频文件路径列表。传入时会自动写入paths2audio_filepath.json作为 manifest否则使用配置中已有的diarizer.manifest_filepath。batch_size说话人嵌入提取与 VAD 推理的批大小0表示沿用配置。执行顺序与产出目录目录准备out_dir/speaker_outputs存在则清除旧结果、out_dir/vad_outputs内含vad_out.json、out_dir/pred_rttms。语音活动检测_perform_speech_activity_detection()三种来源三选一否则抛出ValueError——NeMo VAD 模型对长音频默认按split_duration50秒切分以防显存溢出再经prepare_manifest、_setup_vad_test_data、_run_vad得到帧级语音概率并生成分段表vad.external_vad_manifest外部 VAD 结果 manifestoracle_vad直接用 RTTM 真值生成 oracle 分段 manifest。多尺度子分段 嵌入提取对每个尺度(window, shift)依次调用_run_segmentation把 VAD 分段切成语音子段与_extract_embeddings逐 batch 前向self._speaker_model.forward收集嵌入与时间戳。若speaker_embeddings.parameters.save_embeddingsTrue还会把中间嵌入保存到speaker_outputs/embeddings/以便调试复用。聚类perform_clustering汇总多尺度嵌入与时间戳对每条音频做聚类并写出pred_rttms。打分score_labels(..., collardiarizer.collar, ignore_overlapdiarizer.ignore_overlap)计算 DER 等指标后返回。2.3 保存与恢复save_to(save_path)源码将model_config.yaml、speaker_model.nemo及vad_model.nemo若加载过 VAD打包为一个.nemo归档。restore_from(restore_path, override_config_pathNone, map_locationNone)源码解包归档后把配置中 VAD 与说话人模型路径指向归档内的本地文件再重建实例若归档不含 VAD 模型会提示需要提供 VAD 模型或含语音分段的 manifest。verbose属性直接读取cfg.verbose控制 tqdm 进度条的显隐。三、SortformerEncLabelModel端到端分离模型 API该类是 API 文档中列出成员最多的类。它要求配置包含preprocessor、encoderTransformer 或 FastConformer、sortformer_modules可选transformer_encoder类注释。3.1 预训练模型list_available_modelslist_available_models()返回三个可用的预训练模型源码pretrained_model_name说明diar_sortformer_4spk-v1离线 Sortformer最多 4 说话人diar_streaming_sortformer_4spk-v2流式 Sortformer最多 4 说话人diar_streaming_sortformer_4spk-v2.1流式 Sortformer v2.1这些模型也可通过SortformerEncLabelModel.from_pretrained(diar_sortformer_4spk-v1)直接实例化。3.2 数据加载setup_training_data / setup_validation_data / setup_test_data三个方法都委托给私有方法__setup_dataloader_from_config(config)源码其行为由配置决定配置含use_lhotse: true时走get_lhotse_dataloader_from_configLhotseAudioToSpeechE2ESpkDiarDatasetLhotse 数据管线否则构建WaveformFeaturizer与FilterbankFeatures创建AudioToSpeechE2ESpkDiarDataset并以eesd_train_collate_fn作为 collate 函数shuffleFalse。两个值得注意的细节数据加载配置会注入config.subsampling_factor self.output_subsampling_factor保证标签与输出帧率对齐setup_test_data之后可通过test_dataloader()属性取回self._test_dl。3.3 推理前向链process_signal → frontend_encoder → forward_inferprocess_signal(audio_signal, audio_signal_length)源码将输入移到模型所在设备非流式模式下做峰值归一化(1 / (max eps)) * audio_signaleps默认1e-3可在cfg.eps配置当批次总时长超过max_batch_dur默认 20000 秒通常意味着单条超长流式音频时切换为oom_safe_feature_extraction做显存安全的分块特征提取返回 mel 特征(B, num_features, num_frames)与帧长。frontend_encoder(processed_signal, processed_signal_length, bypass_pre_encodeFalse)源码调用 encoder 得到emb_seq转置为时间主序(B, T, D)若encoder.d_model ! model_defaults.tf_d_model则经sortformer_modules.encoder_proj线性投影对齐维度否则跳过。forward_infer(emb_seq, emb_seq_length)源码——离线推理的主前向由长度构造encoder_mask若配置了transformer_encoder先做一次 Transformer 编码sortformer_modules.upsample_hidden上采样到帧率high_resolutionTrue时掩码按upsample_factor重复展开forward_speaker_sigmoids输出排序后的说话人激活概率形状(batch_size, diar_frame_count, num_speakers)并按输出掩码置零无效帧。forward(audio_signal, audio_signal_length)源码是训练与推理的统一入口process_signal→ 训练态加spec_augmentation→ 流式模式走forward_streaming离线模式走frontend_encoderforward_infer→ 按output_subsampling_factor计算输出长度、裁剪并做downsample_preds降采样。输出分辨率由_resolve_output_resolution校验output_subsampling_factor必须是模型原生子采样因子的整数倍否则回退并告警源码。NeuralType 契约由input_types/output_types属性声明输入(B,T)的AudioSignal与LengthsType输出(B,T,C)的ProbsType。3.4 训练与验证training_step / validation_step / multi_validation_epoch_end损失权重由_init_loss_weights()初始化源码从cfg.pil_weight默认 0.0与cfg.ats_weight默认 1.0归一化出pil_weight与ats_weight两者不能同时为 0并校验ats_tolerance 0。ATSAligned Time-softmax Score是无排列歧义的对齐损失PILPermutation Invariant Loss覆盖排列搜索开销。training_step(batch, batch_idx)前向得到preds计算损失并通过_get_aux_train_evaluations(preds, targets, target_lens)记录 batch 级 F1/Precision/Recall 辅助指标最后_reset_train_metrics周期性复位精度指标。validation_step(batch, batch_idx, dataloader_idx0)与test_step分别调用_get_aux_validation_evaluations/_get_aux_test_batch_evaluations统计验证/测试集指标。multi_validation_epoch_end(outputs, dataloader_idx0)汇总多验证集输出on_validation_epoch_end覆写为super().on_validation_epoch_end(sync_metricsTrue)确保多卡分布式下指标同步源码。指标体系由_init_eval_metrics()建立训练/验证/测试各一套MultiBinaryAccuracy_accuracy_*与_accuracy_*_ats由_reset_train_metrics/_reset_valid_metrics复位。另有add_rttms_mask_mats(rttms_mask_mats, device)在需要对齐 GT 分离做评估时注入 RTTM 掩码矩阵重复注入会抛错。四、SpkDiarizationMixin端到端分离的模板方法骨架SpkDiarizationMixindiarization.py为可分离模型提供了统一的diarize()接口。它与具体模型解耦模型只需实现三个抽象方法——_setup_diarize_dataloader、_diarize_forward、_diarize_output_processing。4.1 配置对象DiarizeConfig 与 InternalDiarizeConfigdataclass class DiarizeConfig: session_len_sec: float -1 # 端到端分离会话长度上限秒 batch_size: int 1 num_workers: int 1 sample_rate: Optional[int] None # numpy 输入必须提供 postprocessing_yaml: Optional[str] None # VAD 式后处理参数 yaml verbose: bool True include_tensor_outputs: bool False postprocessing_params: PostProcessingParams None max_num_of_spks: Optional[int] None _internal: Optional[InternalDiarizeConfig] NoneInternalDiarizeConfig是推理过程内部的暂存区device、dtype、training_mode推理前记住、推理后恢复、target_sample_rate默认 16000会被 preprocessor 采样率覆盖、dither_value/pad_to_value推理时临时置 0 再还原、temp_dir、manifest_filepath、max_num_of_spks默认 4。辅助函数get_value_from_diarization_config(diarcfg, key, default)按属性名取值缺失时记录 debug 日志并返回默认值——这使得上层可以传入任意DiarizeConfig子类而不破坏兼容性。4.2 diarize() 与 diarize_generator() 的调用链SortformerEncLabelModel.diarize()源码是一键分离入口直接转发到 mixin 的diarize()output model.diarize( audiopath/to/audio.wav, # 单文件/文件列表/manifest 路径/numpy 波形/DataLoader sample_rate16000, # numpy 输入必需 batch_size1, include_tensor_outputsFalse, # True 时额外返回原始说话人概率张量 postprocessing_yamlNone, # VAD 式后处理阈值 yaml num_workers0, verboseTrue, ) # 返回 [[begin_sec, end_sec, spk_index], ...] # include_tensor_outputsTrue 时返回 (上述列表, preds 张量列表)mixin 内部的完整流程diarize_generator源码_diarize_on_begin把单字符串包装为列表num_workers缺省为min(batch_size, cpu_count-1)记录并临时修改 preprocessor 的dither0、pad_to0切到eval()把日志压到 WARNING 级别。_diarize_input_processing三种输入形态分别处理——单个.json/.jsonlmanifest 路径直接audio_rttm_map解析use_lhotseTrue音频文件路径列表_input_audio_to_rttm_processing为每个文件生成{uniq_id, audio_filepath, offset0.0, text-, label:infer}条目再由_diarize_input_manifest_processing在临时目录写出manifest.jsonnp.ndarray波形必须提供sample_rate经_diarize_numpy_to_1d_float_tensor转单声道 float32 并用librosa.core.resample重采样到模型 preprocessor 采样率CUDA 多 worker 时自动强制num_workers0以避免 worker 中创建 CUDA 张量最终由NumpyAudioDataset_diarize_collate_pad_to_device构建 DataLoader。逐 batch 循环move_data_to_device→self._diarize_forward(test_batch)→self._diarize_output_processing(preds, uniq_ids, diarize_cfg)→ yield 结果并torch.cuda.empty_cache()释放显存。_diarize_on_end恢复训练模式、preprocessor 参数与日志级别放在finally中保证异常时也能恢复。4.3 Sortformer 对抽象方法的实现_setup_diarize_dataloader(config)与训练数据加载同构按manifest_filepath/ 临时 manifest 构建推理 DataLoader源码。_diarize_forward(batch)在torch.no_grad()下调用self.forward并把 preds 移回 CPU 后清缓存源码。_diarize_output_processing(outputs, uniq_ids, diarcfg)按 batch 切分 preds调用predlist_to_timestamps(batch_preds_list, audio_rttm_map_dict, cfg_vad_paramsdiarcfg.postprocessing_params, unit_10ms_frame_countself.output_subsampling_factor)把帧概率转为说话人时间戳再用generate_diarization_output_lines生成 RTTM 行diarcfg.include_tensor_outputsTrue时返回(RTTM 行列表, preds 张量列表)元组源码。4.4 DiarizationMixin级联侧的混入API 文档中的DiarizationMixinmixins.py服务于ClusteringDiarizer提供path2audio_files_to_manifest等 manifest 工具使级联分离器能把任意文件路径列表转换为内部 manifest。ClusteringDiarizer.diarize()中self.path2audio_files_to_manifest(paths2audio_files, ...)即来自该混入。五、配置速查与实操建议ClusteringDiarizer 关键配置项依据 源码解析路径 归纳配置键作用diarizer.oracle_vad用 RTTM 真值做 oracle VAD评估上限用diarizer.vad.model_path.nemo路径或预训练名不可用时回退vad_telephony_marblenetdiarizer.vad.parameters含window_length_in_sec、shift_length_in_sec、smoothingmedian/mean、overlap等 VAD 后处理参数diarizer.vad.external_vad_manifest外部 VAD 分段 manifest三选一diarizer.speaker_embeddings.model_path.nemo/.ckpt/预训练名回退ecapa_tdnndiarizer.speaker_embeddings.parameterswindow_length_in_sec、shift_length_in_sec、multiscale_weights多尺度嵌入、save_embeddingsdiarizer.clustering.parameters聚类算法参数传入perform_clusteringdiarizer.collar/diarizer.ignore_overlapDER 打分的容差与是否忽略重叠diarizer.out_dir输出根目录含vad_outputs、speaker_outputs、pred_rttmsverbose进度条开关Sortformer 训练配置要点依据 构造函数preprocessor、encoder含subsampling_factor默认按 8 处理、sortformer_modules、可选transformer_encoder与spec_augment损失权重pil_weight/ats_weight/ats_toleranceeps默认 1e-3、high_resolutionbool、output_subsampling_factor须整除关系成立、streaming_mode/async_streaming流式模型初始化时会经_check_streaming_parameters校验 chunk 长度与降采样因子的整除关系、max_batch_dur默认 20000 秒。实操建议快速验证用SortformerEncLabelModel.from_pretrained(diar_sortformer_4spk-v1)加载离线模型后直接调diarize([a.wav, b.wav])即可获得每段[开始秒, 结束秒, 说话人序号]需要原始概率时传include_tensor_outputsTrue。长音频/流式diar_streaming_sortformer_4spk-v2.1配合streaming_modeTrue的模型配置使用输出降采样必须满足output_subsampling_factor与chunk_len * upsample_factor的整除约束校验逻辑。级联基线ClusteringDiarizer适合已有成熟 VAD 嵌入模型的场景或需要 oracle/外部 VAD 对照实验的评估流程注意其 VAD 阶段默认把长音频切成 50 秒段显存仍紧张时可调小split_duration见 _perform_speech_activity_detection 中的提示日志。训练/微调按第 3.2、3.4 节准备 train/val/test 三份数据配置可含use_lhotse: true经setup_training_data等方法挂接验证集指标在multi_validation_epoch_end与on_validation_epoch_end(sync_metricsTrue)处聚合多卡下无需额外处理。六、小结API 文档所列的四个对象构成两层能力ClusteringDiarizer以配置驱动的方式把 VAD、嵌入、聚类、打分串成离线级联管线其diarize/save_to/restore_from是完整的推理与持久化闭环SortformerEncLabelModel则覆盖从setup_*_data数据加载、training_step/validation_step训练循环到process_signal→frontend_encoder→forward_infer推理前向的全生命周期并可导出流式推理图。SpkDiarizationMixin以模板方法把输入归一化、DataLoader 构建、逐 batch 前向与 RTTM 输出固化为统一骨架DiarizationMixin补足级联侧的 manifest 工具。结合 说话人分离入门文档、模型文档 与 示例脚本即可在本仓库内完成从推理到训练的全部工作。【免费下载链接】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),仅供参考