File size: 9,719 Bytes
bed2cee
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
# 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_layer`、`layer_norm_1`、`second_sub_layer`、`layer_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. 初始化与训练流程

```bash
# 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
```