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