执行摘要
回滚序列长度平衡分区的后备机制,恢复原始分区算法。
PR正文仅说明“Reverts THUDM/slime#1823”,未提供具体原因。结合上下文推测,可能是PR #1823引入的后备机制在实际运行中存在问题或不再需要,因此决定回滚到更简单的原始实现。
该PR值得关注,因为它回滚了一个重要的内存安全机制。建议精读以理解回滚动机,并关注后续是否会有更稳健的分区方案。
该PR没有review评论,直接由作者合并。
PR正文仅说明“Reverts THUDM/slime#1823”,未提供具体原因。结合上下文推测,可能是PR #1823引入的后备机制在实际运行中存在问题或不再需要,因此决定回滚到更简单的原始实现。
该PR值得关注,因为它回滚了一个重要的内存安全机制。建议精读以理解回滚动机,并关注后续是否会有更稳健的分区方案。
该PR没有review评论,直接由作者合并。
_get_capped_partitions 函数,该函数原本用于在平衡分区超出token预算时提供cap-aware分区作为后备。max_tokens 的检查,不再调用后备函数,直接使用 get_seqlen_balanced_partitions 的结果。| 文件 | 模块 | 状态 | 重要度 |
|---|---|---|---|
slime/backends/megatron_utils/data.py |
数据迭代器 | modified | 7.01 |
slime/backends/megatron_utils/data.py
core-logic
这是唯一变更的文件,包含了动态批处理的核心逻辑,移除后备机制直接影响内存安全性。
def get_data_iterator(
args: Namespace,
model: torch.nn.Module | Sequence[torch.nn.Module],
rollout_data: RolloutBatch,
) -> tuple[list[DataIterator], list[int]]:
# ... 省略前部分代码 ...
if not args.use_dynamic_batch_size:
# 固定批处理路径
num_microbatches = [num_local_gbs // args.micro_batch_size for _ in range(num_steps_per_rollout)]
data_iterator = _generate_data_iterator(rollout_data, args.micro_batch_size)
else:
# 动态批处理路径
assert args.max_tokens_per_gpu is not None
# 计算每个 step 的 microbatch 数量
samples = rollout_data["total_lengths"]
num_microbatches = []
for i in range(num_steps_per_rollout):
start, end = i * num_local_gbs, (i + 1) * num_local_gbs
num_microbatches.append(
get_minimum_num_micro_batch_size(samples[start:end], args.max_tokens_per_gpu * cp_size)
)
# 全局同步最大 microbatch 数量
num_microbatches = torch.tensor(num_microbatches, dtype=torch.int, device=torch.cuda.current_device())
dist.all_reduce(num_microbatches, op=dist.ReduceOp.MAX, group=dp_group)
if vpp_size > 1:
# VPP 对齐逻辑:从向上取整改为向下取整,仅需被 vpp_size 整除
num_microbatches = torch.clamp(
num_microbatches // microbatch_group_size_per_vp_stage * microbatch_group_size_per_vp_stage,
min=1,
)
num_microbatches = num_microbatches.tolist()
# 平衡每个 microbatch 的序列长度
micro_batch_indices = []
for i, num_mbs in enumerate(num_microbatches):
start, end = i * num_local_gbs, (i + 1) * num_local_gbs
samples = rollout_data["total_lengths"][start:end]
# 直接使用平衡分区,不再检查是否超出 max_tokens
partitions = get_seqlen_balanced_partitions(samples, num_mbs, equal_size=False)
# 调整索引偏移
for j in range(num_mbs):
for k in range(len(partitions[j])):
partitions[j][k] += start
micro_batch_indices.extend(partitions)
assert len(set(sum(micro_batch_indices, []))) == num_local_samples
data_iterator = _generate_data_iterator(rollout_data, None, micro_batch_indices)
return data_iterator, num_microbatches
当前评论区没有形成足够清晰的争议点或结论,后续有更多讨论时会体现在这里。
get_seqlen_balanced_partitions 在某些边缘情况下产生超出GPU内存限制的分区,可能导致OOM错误。get_seqlen_balanced_partitions 的可靠性。
参与讨论