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
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)
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) —— 当前训练使用
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 至显存饱和点
# 逐步测试:从小到大,观察 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 内存高效注意力(免改代码)
# 在 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
# config
batch_size: 4 # 每步的 micro batch
trainer.accumulate_grad_batches: 3 # 梯度累计 3 步 → 有效 batch=12
策略 4: 添加梯度检查点 (Activation Checkpointing)
在 ConformerEncoder 和 TransformerEncoder 的每层上应用 torch.utils.checkpoint.checkpoint,以时间换空间:
# 修改 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 参数
# 保守配置(低显存)
sortformer_modules:
chunk_len: 94 # 每 chunk 处理更少帧
spkcache_len: 94 # 相应缩小缓存
causal_attn_rate: 1.0 # 全部批次使用受限自注意力
causal_attn_rc: 7 # 仅 7 帧右侧上下文
- 优点:显存降低 ~50%,自注意力序列减半
- 缺点:模型可能损失一些长程建模能力
策略 6: 混合策略(推荐)
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 实时显存监控
# 终端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 |