File size: 11,072 Bytes
74e4281
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
# Streaming Sortformer 长音频训练原理与 GPU 显存优化

## 1. 模型架构概览

```
Audio (16kHz, T秒)

  ├─ AudioToMelSpectrogramPreprocessor  (25ms窗, 10ms步长, 128 mel bins)
  │   └─ 输出: (B, 128, N_feat_frames)   where N_feat_frames ≈ T × 100

  ├─ [SpecAugment] (仅训练时)

  ├─ ConformerEncoder (17层, d_model=512, rel_pos attention, 8头, 8x subsampling via dw_striding)
  │   │  pre_encode: ConvSubsampling → (B, N_feat_frames/8, 512)
  │   │  encode: 17 × ConformerLayer (MHA → Conv → FFN)
  │   └─ 输出: (B, N_feat_frames/8, 512)  每帧分辨率 = 80ms

  ├─ encoder_proj: Linear(512 → 192)  (仅当 fc_d_model ≠ tf_d_model 时)

  ├─ TransformerEncoder (18层, d_model=192, 8头, inner_size=768)
  │   └─ 输出: (B, N_feat_frames/8, 192)

  └─ SortformerModules
      ├─ 10-speaker split-head (4 base + 6 new, 差分学习率)
      ├─ Sigmoid activation
      └─ 输出: (B, N_feat_frames/8, 10)
```

**参数量**: ~70-80M (~140-160 MB in bf16)

## 2. 流式训练的两种正向传播路径

### 离线模式 (streaming_mode=False)

```python
forward():
    mel = preprocessor(audio)                          # (B, 128, F)
    emb = ConformerEncoder(mel)                        # (B, F/8, 512) — 全序列自注意力
    emb = encoder_proj(emb)                            # (B, F/8, 192)
    preds = TransformerEncoder(emb) + Sigmoid head     # (B, F/8, N_spk)
```

- 全序列同时编码,自注意力 O(F²)
- GPU 显存瓶颈:**注意力矩阵** `(B, n_heads, F, F)`

### 流式模式 (streaming_mode=True) —— 当前训练使用

```python
forward():
    mel = preprocessor(audio)                          # (B, 128, F)  全量保持在显存中

    streaming_state = init_streaming_state()            # spkcache=(B, 0, 512), fifo=(B, 0, 512)

    for chunk in streaming_feat_loader(mel):            # 每次 yield 一个 chunk
        # chunk = mel[:, left_offset : chunk_len]

        forward_streaming_step(chunk, streaming_state):
            chunk_emb  = Conformer.pre_encode(chunk)                   # 上下文无关的预编码

            combined   = concat([spkcache, fifo, chunk_emb])           # ~376 frames max
            encoded    = Conformer.encode(combined) + Transformer(combined)  # 自注意力仅限 376 帧
            preds      = Sigmoid_head(encoded)

            streaming_state = streaming_update(combined, preds)        # 更新 spkcache / fifo
            total_preds     = cat([total_preds, chunk_preds])

    return total_preds                                                 # (B, F/8, N_spk)
```

**关键优势**: 每次自注意力范围 = `spkcache_len(188) + fifo_len(0) + chunk_len(188)`**376 帧**,与总序列长度无关。

## 3. 长音频(1000s)的 GPU 显存消耗分析

### 参数设定(当前 10spk config)

| 参数 | 值 | 含义 |
|------|-----|------|
| `session_len_sec` | 1200s | 单样本最长加载音频 |
| `chunk_len` | 188 | 每次处理的帧数 (80ms/帧) |
| `spkcache_len` | 188 | 说话人缓存帧数 |
| `fifo_len` | 0 | FIFO 队列(禁用) |
| `batch_size` | 12 | 每批样本数 |
| `causal_attn_rate` | 0.5 | 50% 批次使用受限右上下文 |
| `causal_attn_rc` | 7 | 受限右上下文帧数 |

### 1000s 音频的帧数

```
Feat frames: 1000s × 100Hz  = 100,000 帧
Diar frames: 1000s / 0.08s  =  12,500 帧  (8x subsampling)
Chunks:      12,500 / 188   ≈  67 个 chunk
```

### 逐阶段显存分析(以 batch_size=1, bf16 为例)

| 阶段 | 形状 | 估算显存 |
|------|------|---------|
| 原始波形 | (1, 16,000,000) | ~32 MB |
| Mel spectrogram | (1, 128, 100,000) | ~25 MB |
| 单个 chunk Mel | (1, 128, ~1504) | ~0.4 MB |
| chunk pre_encode | (1, 188, 256) | ~0.1 MB |
| combined (spkcache+fifo+chunk) | (1, 376, 512) | ~0.4 MB |
| Conformer 17层 (376帧) | 注意力: (1, 8, 376, 376) × 17 | ~324 MB (梯度前向) |
| Transformer 18层 (376帧) | 注意力: (1, 8, 376, 376) × 18 | ~144 MB |
| total_preds (累积) | (1, 12500, 10) | ~0.25 MB |
| spkcache | (1, 188, 512) | ~0.2 MB |

> **注**: 以上是单次前向传播的估算。由于 spkcache 在 chunk 间传递且保持梯度图连接,训练时每个 chunk 的 35 层中间激活值都需要保留用于反向传播。

### 离线模式 vs 流式训练模式(1000s 音频, 单样本, bf16)

| 指标 | 离线模式 | 流式训练模式 |
|------|---------|-------------|
| 每次自注意力序列长度 | **12,500 帧** | **~376 帧** |
| 单一注意力矩阵 | ~600 MB | ~0.5 MB |
| 35 层注意力合计 | ~21 GB ❌ OOM | ~17 MB ✓ |
| 峰值显存估算 | >40 GB (严重 OOM) | ~15-25 GB |
| 训练可行性 | 不可行 | 可行 |

**结论**: 流式训练模式是训练长音频的核心机制——通过在 spkcache 尺度(~376 帧)上进行自注意力,避免了 O(F²) 的显存爆炸。

## 4. 影响显存的关键参数

### 4.1 `batch_size`

显存占用与 `batch_size` 近似呈**线性关系**。每个样本独立维护一份 spkcache 和 fifo 状态。

```
峰值显存 ≈ batch_size × (spkcache_per_sample + chunk_activations + total_preds)
```

### 4.2 `session_len_sec`

- 不直接影响自注意力显存(流式模式下注意力固定为 ~376 帧)
- 影响 **total_preds** 累积张量大小:`(B, N_frames, N_spk)`
- 影响 **processed_signal** Mel 频谱的显存占用

### 4.3 `chunk_len` / `spkcache_len`

直接影响每次自注意力的序列长度 `(B, heads, chunk+spkcache+fifo, chunk+spkcache+fifo)`。

| chunk_len | 自注意力矩阵大小 | 相对显存 |
|-----------|----------------|---------|
| 94 | (B, 8, 282, 282) | ~0.5x |
| 188 (当前) | (B, 8, 376, 376) | 1x |
| 376 | (B, 8, 564, 564) | ~2.3x |

### 4.4 `causal_attn_rate` / `causal_attn_rc`

训练时,`causal_attn_rate` 比例的批次将 Conformer 的 `att_context_size` 临时设置为 `[-1, 7]`(仅保留 7 帧右侧上下文),同时将 Transformer 的 `diag` 参数设为 7(三角掩码)。

这会将自注意力的有效范围从 376×376 降至约 376×7 → 显存从 ~0.5MB 降至 ~0.01MB,约为原来的 **2%**

### 4.5 模型维度参数

| 参数 | 值 | 对显存的影响 |
|------|-----|------------|
| `fc_d_model` (Conformer) | 512 | 注意力的 K/V 维度 |
| `tf_d_model` (Transformer) | 192 | Transformer 的每层激活值 |
| `n_heads` | 8 | 注意力头的数量 |
| `n_layers` (Conformer) | 17 | 层数直接影响中间激活值累积 |

## 5. 最大化 1000s 音频训练的显存利用率

### 5.1 当前配置分析

当前 config (`streaming_sortformer_diarizer_10spk-v1.yaml`) 已将 `session_len_sec` 从 90s 提升至 1200s,`batch_size` 从 4 提升至 12。

### 5.2 优化策略(由易到难)

#### 策略 1: 调整 batch_size 至显存饱和点

```bash
# 逐步测试:从小到大,观察 nvidia-smi 显存使用率
for bs in 1 2 4 6 8 10 12 14 16; do
    echo "Testing batch_size=$bs"
    # 在 config 中设置 batch_size=$bs,训练几个 step 后观察显存
done
```

#### 策略 2: 启用 PyTorch 内存高效注意力(免改代码)

```python
# 在 train 脚本开头添加:
torch.backends.cuda.enable_mem_efficient_sdp(True)
```

这会将 Attention 模块的 `torch.matmul(Q, K.transpose())` + `softmax` + `matmul(A, V)` 合并为 `F.scaled_dot_product_attention(Q, K, V)`,使用 PyTorch 内置的 memory-efficient 实现,可节省 20-40% 注意力显存。

#### 策略 3: 使用梯度累计 + 小 batch

```yaml
# config
batch_size: 4                    # 每步的 micro batch
trainer.accumulate_grad_batches: 3  # 梯度累计 3 步 → 有效 batch=12
```

#### 策略 4: 添加梯度检查点 (Activation Checkpointing)

在 ConformerEncoder 和 TransformerEncoder 的每层上应用 `torch.utils.checkpoint.checkpoint`,以时间换空间:

```python
# 修改 conformer_encoder.py 和 transformer_encoders.py,在每层 forward 外包 checkpoint
for layer in self.layers:
    output = torch.utils.checkpoint.checkpoint(layer, output, mask, ...)
```

每个 checkpointed 层仅保留输入,反向传播时重新计算中间激活值。**约可节省 50-70% 编码器显存**,代价是训练速度降低 20-30%。

#### 策略 5: 调整 chunk 参数

```yaml
# 保守配置(低显存)
sortformer_modules:
    chunk_len: 94             # 每 chunk 处理更少帧
    spkcache_len: 94          # 相应缩小缓存
    causal_attn_rate: 1.0     # 全部批次使用受限自注意力
    causal_attn_rc: 7          # 仅 7 帧右侧上下文
```

- 优点:显存降低 ~50%,自注意力序列减半
- 缺点:模型可能损失一些长程建模能力

#### 策略 6: 混合策略(推荐)

```yaml
trainer:
    precision: bf16                    # bf16 混合精度(已启用)
    accumulate_grad_batches: 2          # 梯度累计

batch_size: 6                          # micro batch

model:
    train_ds:
        session_len_sec: 600            # 截断至 10 分钟(训练时合理上界)
        use_lhotse: False               # 当前不可用

    sortformer_modules:
        chunk_len: 188                  # 保持默认
        spkcache_len: 188               # 保持默认
        causal_attn_rate: 0.7           # 提高受限注意力比例
        causal_attn_rc: 7
```

预期效果:V100-32G 上显存利用率 ~80-90%,训练速度良好。

### 5.3 实时显存监控

```bash
# 终端1: 启动监控
bash dev/gpu_usage/gpu_monitor.sh gpu_usage.csv

# 终端2: 启动训练
bash run_train_10spk.sh

# 训练结束后,查看显存峰值
python -c "
import pandas as pd
df = pd.read_csv('gpu_usage.csv')
for col in df.columns:
    if 'util' in col or 'mem_used' in col:
        print(f'{col}: max={df[col].max()}, mean={df[col].mean():.1f}')
"
```

### 5.4 常见 OOM 原因排查

| 症状 | 可能原因 | 解决方案 |
|------|---------|---------|
| mel spectrogram 阶段 OOM | `session_len_sec × batch_size` 过大导致原始波形占满显存 | 减小 `batch_size` |
| "GPU out of memory with X bytes" | 某层自注意力超过显存 | 启用 mem_efficient_sdp 或减小 `batch_size` |
| 第一次 step 就 OOM | batch 内 padding 后最大长度远大于 `session_len_sec` | 检查 `session_len_sec` 对齐 |
| 训练几个 step 后 OOM | 梯度累积或显存碎片 | `torch.cuda.empty_cache()` |
| 验证阶段 OOM | `validation_ds.session_len_sec` 过大且 `batch_size` 过大 | 验证时减小 `batch_size` |

## 6. 总结

| 要点 | 说明 |
|------|------|
| 核心机制 | 流式训练通过 chunk-by-chunk 处理 + spkcache 状态传递,将自注意力从 O(T²) 降为 O(chunk²) |
| 瓶颈 | 35 层编码器的中间激活值 × batch_size |
| 关键参数 | `batch_size`, `chunk_len`, `causal_attn_rate`, `session_len_sec` |
| 推荐起始配置 | `batch_size=6`, `session_len_sec=600`, `accumulate_grad_batches=2`, `causal_attn_rate=0.7` |
| 进阶优化 | `mem_efficient_sdp`, `gradient_checkpointing`, `accumulate_grad_batches` |