跳转至

sequence_parallel.py — 序列并行

文件路径: verl/utils/megatron/sequence_parallel.py

文件概述

Megatron 框架下的序列并行工具,处理去除 padding 后的 token 序列的 padding 对齐。

核心函数

def pad_to_sequence_parallel(unpad_tokens):
    """确保去 padding 后的 token 总数是 TP 并行度的倍数"""
    sp_world_size = mpu.get_tensor_model_parallel_world_size()
    pad_size = sp_world_size - total_nnz % sp_world_size
    if pad_size > 0:
        unpad_tokens = F.pad(unpad_tokens, (0, pad_size))
    return unpad_tokens

def mark_parameter_as_sequence_parallel(parameter):
    """标记参数为序列并行参数"""
    parameter.sequence_parallel = True

与其他模块的关系

  • 被 pipeline_parallel.py 用来计算输入形状
  • 被 Megatron 模型的前向传播使用

小结

确保变长序列在序列并行下正确对齐的工具。