finetune_spk-sortformer / docs /ARCHITECTURE.md
czyhust's picture
Add files using upload-large-folder tool
bed2cee verified
|
Raw
History Blame Contribute Delete
9.72 kB

Sortformer V2:双流 Speaker-Aware 说话人日志

一、整体思路

在原始 Sortformer 的 18 层 Transformer Encoder 基础上,引入一条与主特征流(H1)并行的 speaker 特征流(H2)。两条流在每一层都做 self-attention 和互相之间的 cross-attention,让说话人表征与 diarization 表征持续交互。spk 分支有自己的输出头,用 chunk 级别的 PIL loss 单独监督。

spk encoder 的表征只在第一层之前作为输入,此后 H2 由每层的 H2 分支自己演化,不再重新注入。


二、网络结构

原始音频 (16kHz, 最长300s)
    │
    ├──────────────────────────────┐
    │                              │
    ▼                              ▼
Mel Spectrogram               CAM++ Speaker Encoder(冻结backbone)
(128维, 10ms)                Fbank → 卷积backbone → 帧级特征
    │                        → Conv下采样 → 192维@80ms
    ▼                              │
FastConformer (17层, 512维)        │
8x下采样 → 512维@80ms              │
    │                              │
    ▼                              ▼
Linear 512→192                 H2 初始表征 (B,T,192)
    │                              │
    ▼                              ▼
    ┌──────────────────────────────────────────┐
    │    DualStream Transformer (18层, 192维)  │
    │                                          │
    │   每层 Block:                            │
    │     H1: self-attn → H1/H2 互attn → FFN   │
    │     H2: self-attn → H1/H2 互attn → FFN   │
    │                                          │
    └──────────────────────────────────────────┘
    │                              │
    ▼                              ▼
主 Head (预训练权重)          spk Head (随机初始化)
    │                              │
    ▼                              ▼
ATS+PIL loss                  chunk级PIL loss

每层 DualStreamTransformerBlock 细节

输入: H1 (B,T,192), H2 (B,T,192)

1. H1 self-attention   (first_sub_layer, 预训练权重)
   H1 = LN1(H1 + MHA(H1,H1,H1))

2. H2 self-attention   (h2_self_attn, 随机初始化)
   H2 = LN1(H2 + MHA(H2,H2,H2))

3. 互 attention(无 gate, 随机初始化)
   H1 = crossLN1(H1 + MHA(H1→H2))   # H1 attend H2
   H2 = crossLN2(H2 + MHA(H2→H1))   # H2 attend H1

4. FFN
   H1 = LN2(H1 + FFN(H1))   (second_sub_layer, 预训练权重)
   H2 = LN2(H2 + FFN(H2))   (h2_ffn, 随机初始化)

所有 attention 均为 8 头(head_dim = 192/8 = 24)。

关键设计:H1 的参数名与原始 TransformerEncoderBlock 完全一致first_sub_layerlayer_norm_1second_sub_layerlayer_norm_2),因此原始 Sortformer checkpoint 可以无缝 warm-start 到 H1 流,新增的 H2 流和互 attention 保持随机初始化。

权重来源总结

模块 权重来源
FastConformer encoder 预训练
H1 self-attn + FFN + LN 预训练(同名加载)
H2 self-attn + FFN + LN 随机初始化
H1↔H2 互 attention ×2 + LN 随机初始化
主输出 head 预训练
spk 输出 head 随机初始化
CAM++ 卷积 backbone 预训练,永远冻结
CAM++ 下采样 conv (proj) 随机初始化,可训练

三、Speaker Encoder(CAM++)

提取流程

原始音频 (最长300s)
    │
    ▼
切分为无重叠窗口:每窗口 2 秒(chunk_dur_sec=2.0, chunk_stride_sec=2.0)
    │
    ▼  每窗口独立处理(encode_batch_size=16 批量并行):
Kaldi Fbank (80维, 10ms) → CAM++ 卷积backbone(冻结)
    → 帧级特征 (512维, 20ms),每帧编码周围大范围上下文
    │
    ▼
Conv1d(512→192, kernel=12, stride=4, padding=4)  ← 可训练下采样层
    → 每个输出帧融合 240ms(12帧×20ms)上下文
    → 192维 @80ms 帧率
    │
    ▼
窗口按时间顺序拼接 → (B, T_diar, 192),与 sortformer 帧率严格对齐

设计说明

  • 不做池化:CAM++ 原始的 utterance-level TSTP 全局池化不适用于逐帧场景。改为可学习卷积下采样,让 H2 流自己的 self-attention 做时间整合。
  • kernel=12(240ms):每个 H2 帧融合 240ms 的 CAM++ 帧级特征,比输出帧率(80ms)宽 3 倍,提供更丰富的局部上下文。
  • proj 可训练:下采样卷积是新参数(~98K),随训练学习如何把帧级特征映射到 H2 空间;CAM++ backbone 冻结。
  • 2s 窗口:与 CAM++ 训练时长一致,窗口间无重叠,显存可控。

四、训练算法

双损失

主分支:  loss1 = 0.5 × ATS_loss + 0.5 × PIL_loss   (原始 Sortformer 损失)
spk分支: loss2 = chunk级 PIL loss
总损失:  loss = loss1 + spk_pil_weight × loss2     (spk_pil_weight=1.0)

Chunk 级 PIL Loss(spk 分支)

主分支沿用原始做法:整段拼接后一次匈牙利排列匹配。spk 分支不同——按 chunk_len=188 帧(15.04秒)切块:

spk_preds / targets 按 188 帧切成 n 个 chunk
每个 chunk 内部独立做 PIL 排列匹配 → 各自算 BCE
最后所有 chunk 的 loss 取平均

每个 chunk 是独立的排列问题,模型被强制在短窗口内做说话人区分,与流式推理场景一致。

两种训练模式

模式 配置 可训练参数 说明
只训新增 freeze_base_model: true ~13.4M H2流 + 互attn + spk head + proj
全量训练 freeze_base_model: false ~131M 所有参数(除 CAM++ backbone)

两种模式下 CAM++ 卷积 backbone 均冻结(proj 下采样层可训练)。

学习率分组(全量模式)

参数组 学习率
原始 Sortformer 参数 lr=2e-5
新增 H2 流 + 互attn + spk head + proj speaker_lr=1e-4
CAM++ backbone 冻结

五、流式模式

  • 流式推理时每 chunk 处理 [spkcache | fifo | chunk] 拼接序列
  • spkcache_spk/fifo_spk 同步缓存对应帧的 speaker embedding,压缩时用相同的 topk_indices gather,保证 H2 与 H1 每帧时刻严格对齐
  • spk 分支(H2 → spk head)只用于训练监督,推理时只用主分支输出

六、超参数

新增超参数

参数 说明
use_dual_stream true 启用双流架构
spk_pil_weight 1.0 spk 分支 chunk PIL loss 权重
freeze_base_model false/true 只训新增 / 全量训练
speaker_lr 1e-4 新增参数学习率(全量模式)
speaker_encoder.chunk_dur_sec 2.0 spk encoder 窗口长度(无重叠)
speaker_encoder.chunk_stride_sec 2.0 窗口步长(=窗口长,无重叠)
speaker_encoder.encode_batch_size 24 多窗口批量并行
speaker_encoder.downsample_kernel 12 下采样卷积核(240ms 上下文)
speaker_encoder.checkpoint_path null 权重从 warm-start nemo 加载,不硬编码

沿用超参数

参数
采样率 16000
最大说话人数 8
最大音频时长 300s
batch size 1
优化器 AdamW (β=[0.9,0.98])
基础学习率 2e-5
权重衰减 2e-3
最大 epoch 10
chunk_len 188 (15.04s)
spkcache_len 376
因果注意力概率 0.5 (rc=7)
dropout (attn/ffn) 0.5

7. 文件结构

fintune_sortformer_v2/
├── docs/ARCHITECTURE.md                      ← 本文档
├── checkpoints/
│   └── dual_stream_init_8spk.nemo            ← 初始化nemo(含全部权重)
├── src/
│   ├── finetune_pipeline/
│   │   ├── conf/
│   │   │   └── streaming_sortformer_diarizer_8spk-v1-finetune.yaml
│   │   │       ← use_dual_stream, spk_pil_weight, freeze_base_model, speaker_lr
│   │   └── scripts/
│   │       ├── init_dual_stream_nemo.py       ← 生成含spk encoder权重的nemo
│   │       └── ckpt2nemo.py                  ← ckpt转nemo(含全部新权重)
│   └── third_party/nemo/collections/asr/
│       ├── models/sortformer_diar_models.py  ← 双流接入、spk head、chunk PIL loss、两种训练模式
│       └── modules/
│           ├── speaker_encoder.py            ← 冻结CAM++ backbone + 可训练Conv下采样
│           ├── sortformer_modules.py         ← spk cache 追踪(spkcache_spk/fifo_spk)
│           └── transformer/
│               ├── dual_stream_transformer.py ← 双流block + encoder
│               └── transformer_encoders.py    ← 原始单流(use_dual_stream=false时用)
└── wespeaker/                                ← CAM++ 模型

8. 初始化与训练流程

# 1. 生成初始化 nemo(H1=预训练, H2/互attn=随机, CAM++=冻结, proj=随机)
python src/finetune_pipeline/scripts/init_dual_stream_nemo.py \
    --pretrained /path/to/original_sortformer.nemo \
    --output /path/to/output_dual_stream.nemo

# 2. 训练时 warm-start 这个 nemo
bash run_train.sh --init_nemo_path /path/to/output_dual_stream.nemo

# 3. 训练后 ckpt 转 nemo(ckpt 已含全部权重)
python src/finetune_pipeline/scripts/ckpt2nemo.py \
    /path/to/checkpoint.ckpt /path/to/output.nemo