NeMo 说话人分离 API 参考:ClusteringDiarizer 与 SortformerEncLabelModel 全解析
2026/9/13 14:49:04 网站建设 项目流程

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混入类。读完本文,你可以掌握每个关键方法(diarizeprocess_signalforward_infertraining_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.ModuleModelDiarizationMixin(见 类定义);SortformerEncLabelModel则继承ModelPTExportableEncDecModelSpkDiarizationMixin(见 类定义),因此前者是"推理容器",后者是"可训练、可导出的神经网络"。

二、ClusteringDiarizer:级联式离线分离的完整流程

2.1 类定位与构造参数

ClusteringDiarizer是离线说话人分离的推理模型类,其文档字符串明确说明它负责分离流程的全部环节:Speech Activity Detection、Segmentation、Extract Embeddings、Clustering、Resegmentation 和 Scoring,所有参数均通过配置文件传入(见 类注释)。

构造函数签名:

def __init__(self, cfg: Union[DictConfig, Any], speaker_model=None):
  • cfg:OmegaConf 配置。若传入DictConfig,会先经model_utils.convert_model_config_to_dict_configmaybe_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-L158):

  • VAD:当oracle_vad=Falsevad.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表示沿用配置。

执行顺序与产出目录:

  1. 目录准备:out_dir/speaker_outputs(存在则清除旧结果)、out_dir/vad_outputs(内含vad_out.json)、out_dir/pred_rttms
  2. 语音活动检测_perform_speech_activity_detection():三种来源三选一,否则抛出ValueError——
    • NeMo VAD 模型:对长音频默认按split_duration=50(秒)切分以防显存溢出,再经prepare_manifest_setup_vad_test_data_run_vad得到帧级语音概率并生成分段表;
    • vad.external_vad_manifest:外部 VAD 结果 manifest;
    • oracle_vad:直接用 RTTM 真值生成 oracle 分段 manifest。
  3. 多尺度子分段 + 嵌入提取:对每个尺度(window, shift)依次调用_run_segmentation(把 VAD 分段切成语音子段)与_extract_embeddings(逐 batch 前向self._speaker_model.forward收集嵌入与时间戳)。若speaker_embeddings.parameters.save_embeddings=True,还会把中间嵌入保存到speaker_outputs/embeddings/以便调试复用。
  4. 聚类perform_clustering汇总多尺度嵌入与时间戳,对每条音频做聚类并写出pred_rttms
  5. 打分score_labels(..., collar=diarizer.collar, ignore_overlap=diarizer.ignore_overlap)计算 DER 等指标后返回。

2.3 保存与恢复

  • save_to(save_path)(源码):将model_config.yamlspeaker_model.nemo(及vad_model.nemo,若加载过 VAD)打包为一个.nemo归档。
  • restore_from(restore_path, override_config_path=None, map_location=None)(源码):解包归档后,把配置中 VAD 与说话人模型路径指向归档内的本地文件再重建实例;若归档不含 VAD 模型,会提示"需要提供 VAD 模型或含语音分段的 manifest"。
  • verbose属性直接读取cfg.verbose,控制 tqdm 进度条的显隐。

三、SortformerEncLabelModel:端到端分离模型 API

该类是 API 文档中列出成员最多的类。它要求配置包含preprocessorencoder(Transformer 或 FastConformer)、sortformer_modules,可选transformer_encoder(类注释)。

3.1 预训练模型:list_available_models

list_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_config+LhotseAudioToSpeechE2ESpkDiarDataset(Lhotse 数据管线);
  • 否则构建WaveformFeaturizerFilterbankFeatures,创建AudioToSpeechE2ESpkDiarDataset,并以eesd_train_collate_fn作为 collate 函数,shuffle=False

两个值得注意的细节:数据加载配置会注入config.subsampling_factor = self.output_subsampling_factor,保证标签与输出帧率对齐;setup_test_data之后可通过test_dataloader()属性取回self._test_dl

3.3 推理前向链:process_signal → frontend_encoder → forward_infer

process_signal(audio_signal, audio_signal_length)(源码):

  1. 将输入移到模型所在设备;非流式模式下做峰值归一化(1 / (max + eps)) * audio_signaleps默认1e-3,可在cfg.eps配置);
  2. 当批次总时长超过max_batch_dur(默认 20000 秒,通常意味着单条超长流式音频)时,切换为oom_safe_feature_extraction做显存安全的分块特征提取;
  3. 返回 mel 特征(B, num_features, num_frames)与帧长。

frontend_encoder(processed_signal, processed_signal_length, bypass_pre_encode=False)(源码):

  • 调用 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_resolution=True时掩码按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_encoder+forward_infer→ 按output_subsampling_factor计算输出长度、裁剪并做downsample_preds降采样。输出分辨率由_resolve_output_resolution校验:output_subsampling_factor必须是模型原生子采样因子的整数倍,否则回退并告警(源码)。

NeuralType 契约由input_types/output_types属性声明:输入('B','T')AudioSignalLengthsType,输出('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_weightats_weight(两者不能同时为 0),并校验ats_tolerance >= 0。ATS(Aligned Time-softmax Score)是无排列歧义的对齐损失;PIL(Permutation 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_idx=0)test_step:分别调用_get_aux_validation_evaluations/_get_aux_test_batch_evaluations统计验证/测试集指标。
  • multi_validation_epoch_end(outputs, dataloader_idx=0):汇总多验证集输出;on_validation_epoch_end覆写为super().on_validation_epoch_end(sync_metrics=True),确保多卡分布式下指标同步(源码)。
  • 指标体系由_init_eval_metrics()建立:训练/验证/测试各一套MultiBinaryAccuracy_accuracy_*_accuracy_*_ats),由_reset_train_metrics/_reset_valid_metrics复位。
  • 另有add_rttms_mask_mats(rttms_mask_mats, device):在需要对齐 GT 分离做评估时注入 RTTM 掩码矩阵,重复注入会抛错。

四、SpkDiarizationMixin:端到端分离的模板方法骨架

SpkDiarizationMixin(diarization.py)为"可分离模型"提供了统一的diarize()接口。它与具体模型解耦:模型只需实现三个抽象方法——_setup_diarize_dataloader_diarize_forward_diarize_output_processing

4.1 配置对象:DiarizeConfig 与 InternalDiarizeConfig

@dataclass 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] = None

InternalDiarizeConfig是推理过程内部的暂存区:devicedtypetraining_mode(推理前记住、推理后恢复)、target_sample_rate(默认 16000,会被 preprocessor 采样率覆盖)、dither_value/pad_to_value(推理时临时置 0 再还原)、temp_dirmanifest_filepathmax_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( audio="path/to/audio.wav", # 单文件/文件列表/manifest 路径/numpy 波形/DataLoader sample_rate=16000, # numpy 输入必需 batch_size=1, include_tensor_outputs=False, # True 时额外返回原始说话人概率张量 postprocessing_yaml=None, # VAD 式后处理阈值 yaml num_workers=0, verbose=True, ) # 返回 [[begin_sec, end_sec, spk_index], ...], # include_tensor_outputs=True 时返回 (上述列表, preds 张量列表)

mixin 内部的完整流程(diarize_generator,源码):

  1. _diarize_on_begin:把单字符串包装为列表;num_workers缺省为min(batch_size, cpu_count-1);记录并临时修改 preprocessor 的dither=0pad_to=0;切到eval();把日志压到 WARNING 级别。
  2. _diarize_input_processing:三种输入形态分别处理——
    • 单个.json/.jsonlmanifest 路径:直接audio_rttm_map解析,use_lhotse=True
    • 音频文件路径列表:_input_audio_to_rttm_processing为每个文件生成{uniq_id, audio_filepath, offset=0.0, text='-', label:'infer'}条目,再由_diarize_input_manifest_processing在临时目录写出manifest.json
    • np.ndarray波形:必须提供sample_rate;经_diarize_numpy_to_1d_float_tensor转单声道 float32 并用librosa.core.resample重采样到模型 preprocessor 采样率;CUDA + 多 worker 时自动强制num_workers=0以避免 worker 中创建 CUDA 张量;最终由NumpyAudioDataset+_diarize_collate_pad_to_device构建 DataLoader。
  3. 逐 batch 循环move_data_to_deviceself._diarize_forward(test_batch)self._diarize_output_processing(preds, uniq_ids, diarize_cfg)→ yield 结果,并torch.cuda.empty_cache()释放显存。
  4. _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_params=diarcfg.postprocessing_params, unit_10ms_frame_count=self.output_subsampling_factor)把帧概率转为说话人时间戳,再用generate_diarization_output_lines生成 RTTM 行;diarcfg.include_tensor_outputs=True时返回(RTTM 行列表, preds 张量列表)元组(源码)。

4.4 DiarizationMixin(级联侧的混入)

API 文档中的DiarizationMixin(mixins.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_marblenet
diarizer.vad.parameterswindow_length_in_secshift_length_in_secsmoothing(median/mean)、overlap等 VAD 后处理参数
diarizer.vad.external_vad_manifest外部 VAD 分段 manifest,三选一
diarizer.speaker_embeddings.model_path.nemo/.ckpt/预训练名(回退ecapa_tdnn
diarizer.speaker_embeddings.parameterswindow_length_in_secshift_length_in_secmultiscale_weights(多尺度嵌入)、save_embeddings
diarizer.clustering.parameters聚类算法参数,传入perform_clustering
diarizer.collar/diarizer.ignore_overlapDER 打分的容差与是否忽略重叠
diarizer.out_dir输出根目录,含vad_outputsspeaker_outputspred_rttms
verbose进度条开关

Sortformer 训练配置要点(依据 构造函数):preprocessorencoder(含subsampling_factor,默认按 8 处理)、sortformer_modules、可选transformer_encoderspec_augment;损失权重pil_weight/ats_weight/ats_toleranceeps(默认 1e-3)、high_resolution(bool)、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_outputs=True
  • 长音频/流式:diar_streaming_sortformer_4spk-v2.1配合streaming_mode=True的模型配置使用;输出降采样必须满足output_subsampling_factorchunk_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_endon_validation_epoch_end(sync_metrics=True)处聚合,多卡下无需额外处理。

六、小结

API 文档所列的四个对象构成两层能力:ClusteringDiarizer以"配置驱动"的方式把 VAD、嵌入、聚类、打分串成离线级联管线,其diarize/save_to/restore_from是完整的推理与持久化闭环;SortformerEncLabelModel则覆盖从setup_*_data数据加载、training_step/validation_step训练循环,到process_signalfrontend_encoderforward_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),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询