Instructions to use czyhust/finetune_moss-sortformer with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- NeMo
How to use czyhust/finetune_moss-sortformer with NeMo:
# tag did not correspond to a valid NeMo domain.
- Notebooks
- Google Colab
- Kaggle
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` |
|