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 模型的前向传播使用
小结¶
确保变长序列在序列并行下正确对齐的工具。