YAML Metadata Warning:empty or missing yaml metadata in repo card

Check out the documentation for more information.

G-STAR Sortformer 微调流水线

多说话人日志化-AST(自动语音转录)微调流水线,适用于 G-STAR 流式 Sortformer 模型。通过 FastMSS 生成合成多说话人对话数据,并使用 split-head 差异化学习率训练支持任意说话人数量扩展(4→N)的日志化 ASR 模型。

核心功能

  • 多说话人数据模拟 — 基于 FastMSS 的会议对话生成,包含 HMM 话轮切换、RIR 混响、噪声增强和三角数加权说话人采样
  • Sortformer 微调 — 流式 Sortformer 训练,支持 Hungarian PIT 损失、split-head 架构和 bf16 精度
  • 任意 N 说话人扩展 — 基于 SVD 正交初始化,将预训练的 4 说话人检查点扩展到 N 说话人,无需从头训练
  • 流式推理 — 基于 FIFO 说话人缓存的逐块日志化,支持 ATS(到达时间排序)标签分配
  • 数据分析 — 生成数据的统计分析(说话人分布、时长分布及 CPD 检测)
  • 新数据源集成 — 支持通过 G-STAR 对齐流水线扩展到无预对齐的音视频+文本数据源

目录结构

finetune_pipeline/
├── train.sh                              # Sortformer 训练入口
├── infer.sh                              # 流式推理入口
├── run_train.sh                          # 训练示例脚本
├── run_infer.sh                          # 推理示例脚本
├── path.sh                               # Python 环境配置
│
├── src/
│   ├── finetune_pipeline/
│   │   ├── bin/
│   │   │   ├── streaming_sortformer_diar_train.py   # 训练入口脚本
│   │   │   └── e2e_diarize_speech.py                # 推理入口脚本
│   │   ├── scripts/
│   │   │   ├── ckpt2nemo.py               # Lightning ckpt → NeMo .nemo 格式
│   │   │   ├── extend_output_layer.py     # SVD 正交初始化 N 说话人扩展
│   │   │   ├── merge_split_head.py        # 微调后合并 split head → 统一层
│   │   │   ├── convert_infer_log_to_csv.py # 推理日志 → CSV 格式
│   │   │   ├── split_manifest_by_spk.py   # 按说话人数量分割 manifest
│   │   │   └── parse_options.sh           # Bash 命令行参数解析器
│   │   └── conf/
│   │       ├── streaming_sortformer_diarizer_4spk-v2.yaml
│   │       ├── streaming_sortformer_diarizer_8spk-v1-finetune.yaml
│   │       └── streaming_sortformer_diarizer_10spk-v1.yaml
│   └── third_party/
│       ├── FastMSS/          # FastMSS 分支(子模块)— 数据模拟引擎
│       └── nemo/             # 裁剪版 NeMo — Sortformer 模型及工具
│
├── data/
│   ├── fastmss/              # FastMSS 数据模拟配置及脚本
│   │   ├── simulate.sh       # FastMSS 单次运行入口
│   │   └── scripts/          # CutSet 生成、分析等工具
│   ├── shared/               # 共享工具脚本
│   └── conf/                 # 模拟 Hydra 配置
│
├── docs/                     # 项目文档(见下方索引)
├── checkpoints/              # 预训练模型检查点
├── exp/                      # 训练实验输出
└── dev/                      # 开发/调试脚本

快速开始

前置条件

# 激活环境
. ./path.sh

1. 生成训练数据 (FastMSS)

# 完整的 train/dev/test 模拟(6 说话人,总计 750 小时)
bash data/fastmss/run_6spk_v2.sh

# 单次运行(自定义参数)
bash data/fastmss/simulate.sh \
    --output-dir data/dump_fastmss/my_run \
    --config-path extend_sortformer/6spk \
    --config-name v2 \
    --profile train \
    --spk-weights "[1,3,6,10,15,21]" \
    --n-meetings 30000 \
    --seed 42

2. 训练模型

# 使用预配置的 4 说话人训练
bash run_train.sh \
    --train_manifest_path data/train_manifest.json \
    --dev_manifest_path data/dev_manifest.json \
    --exp_name 4spk_v1

# 或使用原始 train.sh
bash train.sh \
    --train_manifest_path data/train_manifest.json \
    --dev_manifest_path data/dev_manifest.json \
    --exp_name 4spk_v1

3. 推理

# 单数据集推理
bash run_infer.sh \
    --model_path checkpoints/model.nemo \
    --dataset_manifest data/test_manifest.jsonl \
    --exp_dir exp/infer/test

# 批量推理多个数据集
bash run_infer.sh \
    --model_path checkpoints/model.nemo \
    --dataset_base_dir data/global_testset \
    --datasets "ami fisher mlc" \
    --exp_dir exp/infer/batch_test

支持任意说话人数量的微调方法

概述

Sortformer 原始设计支持固定数量的说话人(如 4 说话人)。为了扩展到更多说话人(N>4),本项目提供了基于 split-head 架构的扩展方法:

  1. 扩展输出层:将原始的 Linear(192, 4) 分解为 Linear(192, 4) + Linear(192, N-4)
  2. 差异化学习率:基础说话人使用较低学习率(1e-5),新说话人使用较高学习率(1e-4)
  3. 训练后合并:将 split head 合并为统一的 Linear(192, N) 便于后续推理或继续微调

完整流程

步骤 1:扩展输出层(SVD 正交初始化)

# 将 4 说话人模型扩展到 8 说话人
python src/finetune_pipeline/scripts/extend_output_layer.py \
    --src checkpoints/pretrained/sortformer_4spk.nemo \
    --dst-spk 8 \
    --out checkpoints/8spk_extended.nemo

原理说明

  • 对原始权重矩阵 W 进行 SVD 分解:U, S, Vh = torch.linalg.svd(W, full_matrices=True)
  • 新说话人的权重行取自 Vh[4:],与原有行空间正交
  • 按原矩阵平均范数进行归一化,保持数值稳定性

步骤 2:使用差异化学习率微调

# 使用扩展后的模型进行微调
bash run_train.sh \
    --train_manifest_path data/train_8spk.json \
    --dev_manifest_path data/dev_8spk.json \
    --config_name streaming_sortformer_diarizer_8spk-v1-finetune.yaml \
    --init_nemo_path checkpoints/8spk_extended.nemo \
    --exp_name 8spk_finetune

配置要求: 在配置文件中设置:

model:
  sortformer_modules:
    num_spks: 8
    n_base_spks: 4  # 启用 split-head 架构
  lr: 1e-5          # 基础说话人学习率
  optim_new_lr: 1e-4  # 新说话人学习率

步骤 3:合并 Split Head(可选但推荐)

# 微调完成后合并 split head
python src/finetune_pipeline/scripts/merge_split_head.py \
    --src exp/streaming_sortformer_diar_train/8spk_finetune/checkpoints/model--val_loss=xxx.nemo \
    --out checkpoints/8spk_merged.nemo

何时需要合并

  • ✅ 需要继续微调到更多说话人(如 8→10)
  • ✅ 需要简化推理模型结构
  • ✅ 需要发布最终模型
  • ❌ 仅用于单次推理(可直接使用 split-head 模型)

示例:4 → 8 → 10 说话人扩展

# 第一阶段:4 → 8 说话人
python src/finetune_pipeline/scripts/extend_output_layer.py \
    --src checkpoints/4spk.nemo --dst-spk 8 --out checkpoints/8spk_extended.nemo

bash run_train.sh \
    --config_name streaming_sortformer_diarizer_8spk-v1-finetune.yaml \
    --init_nemo_path checkpoints/8spk_extended.nemo \
    --exp_name 4to8_finetune

# 合并 8 说话人模型
python src/finetune_pipeline/scripts/merge_split_head.py \
    --src exp/.../checkpoints/model.nemo --out checkpoints/8spk_merged.nemo

# 第二阶段:8 → 10 说话人
python src/finetune_pipeline/scripts/extend_output_layer.py \
    --src checkpoints/8spk_merged.nemo --dst-spk 10 --out checkpoints/10spk_extended.nemo

bash run_train.sh \
    --config_name streaming_sortformer_diarizer_10spk-v1.yaml \
    --init_nemo_path checkpoints/10spk_extended.nemo \
    --exp_name 8to10_finetune

注意事项

  1. 数据要求:微调数据必须包含目标说话人数量的样本(如 8 说话人训练需要 6-8 说话人的数据)
  2. 学习率设置:新说话人学习率应比基础说话人高约 10 倍
  3. 训练稳定性:建议先用小学习率预热,再逐步增大
  4. 内存需求:split-head 模型比统一模型占用稍多内存

文档索引

文档 主题
Ultra-Sortformer.md 项目概述:SVD 扩展、差异化学习率训练、已发布检查点
fastmss_6spk_data_generation_plan.md 6 说话人 FastMSS 数据生成计划
fastmss_data_source_extension_pipeline.md 通过 G-STAR 对齐扩展新数据源的流水线
necessary_revision.md NeMo 修改的规范清单
simplify_nemo_files.md NeMo 文件的可删除与必需清单
reduce_permutation_complexity.md 用 Hungarian 算法替代 O(N!) PIT 枚举
pit_and_sort_loss.md PIT vs ATS 损失函数详解
precision_loss_bf16.md bf16 尾数精度对 PIT 和 Hungarian 匹配的影响
gpu_memory_scaling_law.md GPU 内存 ≈ k × session_len × batch_size 缩放规律
convert_extended_sortformer.md 4spk → N-spk 检查点扩展工作流
layer_extension_utils.md 扩展脚本架构分析

依赖项

  • Python: PyTorch, NeMo (裁剪版), Lhotse, Pyroomacoustics, Omegaconf, SoundFile
  • 外部工具: Qwen3-ForcedAligner(用于基于对齐的数据源扩展)
  • 数据: LibriSpeech, WHAM 噪声数据
  • 硬件: A800-80G / A100 GPU(已测试 bf16 训练)

许可证

本项目包含来自 NVIDIA NeMo (Apache 2.0) 和 FastMSS (MIT) 的裁剪代码。详细信息请参阅各源文件头部。

Downloads last month
69
Inference Providers NEW
This model isn't deployed by any Inference Provider. 🙋 Ask for provider support