finetune_moss-sortformer / docs /train_gpu_usage.md
czyhust's picture
Full upload
74e4281 verified
|
Raw
History Blame Contribute Delete
11.1 kB

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