diff --git a/megatron/core/datasets/data_schedule.py b/megatron/core/datasets/data_schedule.py index ab2d292ef60..d5bf07053cc 100644 --- a/megatron/core/datasets/data_schedule.py +++ b/megatron/core/datasets/data_schedule.py @@ -11,12 +11,9 @@ broadcast_tensor, build_packed_microbatches, create_data_iterator, - dcp_get_total_workload, - dcp_gpus_needed, - dcp_make_buckets_equal, get_batch_and_global_seqlens, get_cp_slice_for_thd, - next_hdp_group, + next_hdp_group_packing_aware, reroute_samples_to_dcp_ranks, ) from megatron.core.packed_seq_params import ( @@ -424,28 +421,17 @@ def get_groups_and_subsamples(self, sample_id_seqlens): """ mslpr = self.max_seq_len_per_rank min_cp = self.min_cp_size - workload_fn = lambda seq_len, cp_size=None: dcp_get_total_workload( - seq_len, mslpr, cp_size, min_cp - ) - gpus_fn = lambda seq_len: dcp_gpus_needed(seq_len, mslpr, min_cp) - buckets_fn = lambda sample_seqlens, compute_est: dcp_make_buckets_equal( - sample_seqlens, compute_est, mslpr, min_cp - ) - groups = [] sample_id_groups = [] sample_id_seqlens = sorted(sample_id_seqlens, key=lambda x: x[1], reverse=True) + while sample_id_seqlens: - mb, sample_id_seqlens, exec_times, sample_ids = next_hdp_group( + _, sample_id_seqlens, _, sample_ids = next_hdp_group_packing_aware( sample_id_seqlens, - workload_fn, self.total_hdp_gpus, - gpus_needed_fn=gpus_fn, - make_buckets_equal_fn=buckets_fn, max_seq_len_per_rank=mslpr, - get_total_workload_fn=workload_fn, + min_cp_size=min_cp, ) - groups.append(mb) sample_id_groups.append(sample_ids) if ( diff --git a/megatron/core/datasets/data_schedule_utils.py b/megatron/core/datasets/data_schedule_utils.py index 5f8bb6b3e01..190f898cfc3 100644 --- a/megatron/core/datasets/data_schedule_utils.py +++ b/megatron/core/datasets/data_schedule_utils.py @@ -1,15 +1,16 @@ # Copyright (c) 2025 NVIDIA CORPORATION. All rights reserved. -from collections import deque from functools import lru_cache from math import ceil, log2 -from typing import Callable, Dict, List, Optional, Sequence, Tuple +from typing import Dict, List, Optional, Sequence, Tuple import torch from megatron.core.extensions.transformer_engine import get_thd_partitioned_indices from megatron.core.rerun_state_machine import RerunDataIterator +_DYNAMIC_CP_WORKLOAD_CAP_DELTA = 0.05 + def get_cp_slice_for_thd(batch, cp_group, keys: Optional[Sequence[str]] = None): """Partition sequence data for context parallelism in THD format. @@ -530,28 +531,25 @@ def get_batch_and_global_seqlens(data_iterator, num_microbatches, dp_group): # ============================================================================= -def next_hdp_group( +def next_hdp_group_packing_aware( sample_seqlens: List[Tuple[int, int]], - compute_estimator: Callable[[int], float], total_gpus: int, - gpus_needed_fn: Callable[[int], int], - make_buckets_equal_fn: Callable, - max_seq_len_per_rank: float, - get_total_workload_fn: Callable, - delta: float = 0.05, - strategy: str = "dp", - eps_bucket: float = 0.10, + max_seq_len_per_rank: int, + min_cp_size: int = 1, ) -> Tuple[List[List[int]], List[Tuple[int, int]], List[float], List[List[int]]]: - """Form one balanced micro-batch group across DPxCP ranks. - - This is a standalone version of the scheduling algorithm extracted from - DefaultDynamicCPScheduler so it can live in a utils module. - - Extra args compared to the method version: - gpus_needed_fn: callable(seq_len) -> int - make_buckets_equal_fn: callable(sample_seqlens, compute_estimator) -> list[deque] - max_seq_len_per_rank: max tokens per rank for packing - get_total_workload_fn: callable(seq_len, cp_size) -> float + """Form one DCP microbatch with packing-aware CP group selection. + + This differs from the legacy DCP scheduler in two ways: + 1. Short sequences may use a larger CP group than their minimum required + CP size when that lowers the critical-path rank workload. + 2. Candidate placements are bounded by ``tall * max_seq_len_per_rank``, + the per-rank workload upper bound for packing sequences no longer than + the local tallest sequence in the microbatch. + + The scheduler keeps the legacy invariant that each returned microbatch has + no empty DPxCP rank after the fill step. For non-power-of-two DPxCP layouts, + it falls back to the full DPxCP group if power-of-two expansion cannot fill + every rank. """ if not sample_seqlens: return ( @@ -561,241 +559,216 @@ def next_hdp_group( [[] for _ in range(total_gpus)], ) - buckets = make_buckets_equal_fn(sample_seqlens, compute_estimator) - - micro_batches = [[] for _ in range(total_gpus)] - exec_times = [0.0 for _ in range(total_gpus)] - sample_ids_per_gpu = [[] for _ in range(total_gpus)] - packing_sequence_len = {} - - gpu_group_id = [None] * total_gpus - group_members = {} - group_size = {} - next_gid = 0 - - pp_cursor = 0 - prev_needed = None - check_balance = False - - while buckets: - sample_seq_tuple = bucket_idx = None - needed = None - - scan_order = ( - range(len(buckets)) - if strategy == "dp" - else [(pp_cursor + i) % len(buckets) for i in range(len(buckets))] - ) - - for idx in scan_order: - if not buckets[idx]: - continue - cand_tuple = buckets[idx][0] - cand_seq_len = cand_tuple[1] - needed = gpus_needed_fn(cand_seq_len) - - candidate_gids = [gid for gid, sz in group_size.items() if sz == needed] - free_ranks = [r for r, gid in enumerate(gpu_group_id) if gid is None] - if candidate_gids or len(free_ranks) >= needed: - sample_seq_tuple, bucket_idx = cand_tuple, idx - break + def cp_min_fn(seq_len: int) -> int: + return dcp_gpus_needed(seq_len, max_seq_len_per_rank, min_cp_size) - if sample_seq_tuple is None: - break + def workload(seq_len: int, cp_size: int) -> float: + return (seq_len * seq_len) / cp_size - if strategy == "pp": - pp_cursor = (bucket_idx + 1) % len(buckets) + sample_seqlens = sorted(sample_seqlens, key=lambda x: x[1], reverse=True) + local_tall = sample_seqlens[0][1] + cap = float(local_tall) * float(max_seq_len_per_rank) * (1.0 + _DYNAMIC_CP_WORKLOAD_CAP_DELTA) - sample_id, seq_len = sample_seq_tuple - needed = gpus_needed_fn(seq_len) - if prev_needed is None: - prev_needed = needed + micro_batches: List[List[int]] = [[] for _ in range(total_gpus)] + exec_times: List[float] = [0.0 for _ in range(total_gpus)] + sample_ids_per_gpu: List[List[int]] = [[] for _ in range(total_gpus)] + packing_sequence_len: Dict[int, float] = {} - candidate_gids = [ - gid - for gid, sz in group_size.items() - if sz == needed and packing_sequence_len[gid] + seq_len / needed <= max_seq_len_per_rank - ] - if candidate_gids: - best_gid, best_load = min( - ((gid, max(exec_times[r] for r in group_members[gid])) for gid in candidate_gids), - key=lambda t: t[1], - ) - else: - best_gid, best_load = None, float("inf") + gpu_group_id: List[Optional[int]] = [None] * total_gpus + group_members: Dict[int, List[int]] = {} + group_size: Dict[int, int] = {} + next_gid = 0 - free_ranks = [r for r, gid in enumerate(gpu_group_id) if gid is None] - if len(free_ranks) >= needed: - free_sorted = sorted(free_ranks, key=lambda r: exec_times[r]) - new_members = free_sorted[:needed] - new_load = exec_times[new_members[-1]] + sample_id, seq_len = sample_seqlens[0] + cp_size = cp_min_fn(seq_len) + assert cp_size <= total_gpus, ( + f"Sequence length {seq_len} requires CP size {cp_size}, " + f"but only {total_gpus} DPxCP ranks are available." + ) + group_id = next_gid + next_gid += 1 + members = list(range(cp_size)) + group_members[group_id] = members + group_size[group_id] = cp_size + packing_sequence_len[group_id] = seq_len / cp_size + per_gpu_cost = workload(seq_len, cp_size) + for rank in members: + gpu_group_id[rank] = group_id + micro_batches[rank].append(seq_len) + exec_times[rank] += per_gpu_cost + sample_ids_per_gpu[rank].append(sample_id) + + leftovers: List[Tuple[int, int]] = [] + for sample_id, seq_len in sample_seqlens[1:]: + min_needed = cp_min_fn(seq_len) + best = None + + cp_size = min_needed + while cp_size <= total_gpus: + per_gpu_cost = workload(seq_len, cp_size) + + for group_id, size in list(group_size.items()): + if size != cp_size: + continue + if packing_sequence_len.get(group_id, 0) + seq_len / cp_size > max_seq_len_per_rank: + continue + members = group_members[group_id] + member_set = set(members) + projected_max = max( + time + per_gpu_cost if rank in member_set else time + for rank, time in enumerate(exec_times) + ) + if projected_max <= cap and (best is None or projected_max < best[0]): + best = (projected_max, cp_size, "add", group_id, None) - if new_load < best_load: - best_gid = None - chosen_members = new_members - else: - chosen_members = group_members[best_gid] + free_ranks = [ + rank + for rank, assigned_group_id in enumerate(gpu_group_id) + if assigned_group_id is None + ] + if len(free_ranks) >= cp_size: + chosen_members = sorted(free_ranks, key=lambda rank: exec_times[rank])[:cp_size] + chosen_set = set(chosen_members) + projected_max = max( + time + per_gpu_cost if rank in chosen_set else time + for rank, time in enumerate(exec_times) + ) + if projected_max <= cap and (best is None or projected_max < best[0]): + best = (projected_max, cp_size, "new", None, chosen_members) + + cp_size *= 2 + + if best is None: + leftovers.append((sample_id, seq_len)) + continue + + _, selected_cp_size, action, group_id, chosen_members = best + per_gpu_cost = workload(seq_len, selected_cp_size) + if action == "add": + members = group_members[group_id] + packing_sequence_len[group_id] += seq_len / selected_cp_size + for rank in members: + micro_batches[rank].append(seq_len) + exec_times[rank] += per_gpu_cost + sample_ids_per_gpu[rank].append(sample_id) else: - if best_gid is None: - break - chosen_members = group_members[best_gid] - - if best_gid is None: - best_gid = next_gid + group_id = next_gid next_gid += 1 - group_members[best_gid] = chosen_members - group_size[best_gid] = needed - for r in chosen_members: - gpu_group_id[r] = best_gid - - per_gpu_cost = compute_estimator(seq_len) + group_members[group_id] = chosen_members + group_size[group_id] = selected_cp_size + packing_sequence_len[group_id] = seq_len / selected_cp_size + for rank in chosen_members: + gpu_group_id[rank] = group_id + micro_batches[rank].append(seq_len) + exec_times[rank] += per_gpu_cost + sample_ids_per_gpu[rank].append(sample_id) + + def fill_empty_gpus_once() -> bool: + nonlocal micro_batches, exec_times, sample_ids_per_gpu + + empty_ranks = [rank for rank, micro_batch in enumerate(micro_batches) if not micro_batch] + if not empty_ranks: + return False + assert all( + not micro_batches[rank] for rank in range(empty_ranks[0], total_gpus) + ), "fill_empty_gpus_once assumes empty ranks are contiguous at the tail" - packing_sequence_len[best_gid] = packing_sequence_len.get(best_gid, 0) + seq_len / needed - for r in chosen_members: - micro_batches[r].append(seq_len) - exec_times[r] += per_gpu_cost - sample_ids_per_gpu[r].append(sample_id) + existing_group_sizes = set(group_size.values()) + if not existing_group_sizes: + return False + min_group_size = min(existing_group_sizes) + next_power = min(min_group_size * 2, total_gpus) - buckets[bucket_idx].popleft() + for group_id, size in list(group_size.items()): + if size != min_group_size: + continue - while buckets and not buckets[0]: - buckets.pop(0) - pp_cursor %= max(1, len(buckets)) + members = group_members[group_id] + needed_count = next_power - min_group_size + group_start_rank = members[0] + group_end_rank = members[-1] + empty_rank = empty_ranks[0] + if group_end_rank + 1 > empty_rank or group_end_rank + needed_count >= total_gpus: + continue - if needed < prev_needed: - check_balance = True + work_to_push = micro_batches[group_end_rank + 1 : empty_rank] + exec_times_to_push = exec_times[group_end_rank + 1 : empty_rank] + sample_ids_to_push = sample_ids_per_gpu[group_end_rank + 1 : empty_rank] - if ( - check_balance - and buckets - and max(exec_times) - min(exec_times) <= delta * max(exec_times) - ): - break + new_micro_batches: List[List[int]] = [[] for _ in range(total_gpus)] + new_exec_times: List[float] = [0.0 for _ in range(total_gpus)] + new_sample_ids_per_gpu: List[List[int]] = [[] for _ in range(total_gpus)] - leftovers = [] - for b in buckets: - for sample_seq_tuple in b: - leftovers.append(sample_seq_tuple) - - def trim_overload(): - while True: - cur_max = max(exec_times) - cur_min = min(exec_times) - cur_slack = cur_max - cur_min - if cur_slack <= delta * cur_max: - break - if cur_min == 0: - break - - max_r = exec_times.index(cur_max) - gid = gpu_group_id[max_r] - members = group_members[gid] - - if not micro_batches[max_r] or len(micro_batches[max_r]) <= 1: - break - - seq = micro_batches[max_r][-1] - per_gpu_cost = compute_estimator(seq) - - proj_times = exec_times[:] - for r in members: - proj_times[r] -= per_gpu_cost - - proj_slack = max(proj_times) - min(proj_times) - - if proj_slack < cur_slack: - sample_id_to_remove = sample_ids_per_gpu[max_r][-1] - for r in members: - micro_batches[r].pop() - exec_times[r] -= per_gpu_cost - sample_ids_per_gpu[r].pop() - leftovers.append((sample_id_to_remove, seq)) - else: - break + for rank in range(group_start_rank): + new_micro_batches[rank] = micro_batches[rank] + new_exec_times[rank] = exec_times[rank] + new_sample_ids_per_gpu[rank] = sample_ids_per_gpu[rank] - # TODO(tailaim): uncomment this to support different ranks have different num_microbatches - # trim_overload() + for rank in range(group_start_rank, group_end_rank + needed_count + 1): + new_micro_batches[rank] = list(micro_batches[group_end_rank]) + new_exec_times[rank] = sum( + workload(length, next_power) for length in micro_batches[group_end_rank] + ) + new_sample_ids_per_gpu[rank] = list(sample_ids_per_gpu[group_end_rank]) - total_work_before = sum(len(mb) for mb in micro_batches) + for idx, work in enumerate(work_to_push): + target_rank = group_end_rank + needed_count + 1 + idx + new_micro_batches[target_rank] = work + new_exec_times[target_rank] = exec_times_to_push[idx] + new_sample_ids_per_gpu[target_rank] = sample_ids_to_push[idx] - def fill_empty_gpus(micro_batches, exec_times, sample_ids_per_gpu, group_members, group_size): - empty_gpus = [i for i in range(total_gpus) if not micro_batches[i]] - if not empty_gpus: - return (micro_batches, exec_times, sample_ids_per_gpu, group_members, group_size) + group_size[group_id] = next_power + group_members[group_id] = list( + range(group_start_rank, group_end_rank + needed_count + 1) + ) + for other_group_id in list(group_size.keys()): + if other_group_id == group_id: + continue + if min(group_members[other_group_id]) > group_end_rank: + group_members[other_group_id] = [ + rank + needed_count for rank in group_members[other_group_id] + ] - existing_group_sizes = set(group_size.values()) - assert ( - existing_group_sizes - ), "There should be at least one group existing, cannot redistribute, " - "try to increase 'max-seqlen-per-dp-cp-rank'." + micro_batches = new_micro_batches + exec_times = new_exec_times + sample_ids_per_gpu = new_sample_ids_per_gpu + return True - min_group_size = min(existing_group_sizes) - next_power = min(min_group_size * 2, total_gpus) + return False - for gid, size in group_size.items(): - if size == min_group_size: - members = group_members[gid] - needed_count = next_power - min_group_size - group_start_gpu = members[0] - group_end_gpu = members[-1] - empty_gpu = [idx for idx, work in enumerate(micro_batches) if not work][0] - assert not all( - work for work in micro_batches[empty_gpu : empty_gpu + needed_count] - ), "Empty GPUs were detected but not enough to expand." - work_to_push = micro_batches[group_end_gpu + 1 : empty_gpu] - exec_times_to_push = exec_times[group_end_gpu + 1 : empty_gpu] - sample_ids_to_push = sample_ids_per_gpu[group_end_gpu + 1 : empty_gpu] - - new_micro_batches = [[]] * len(micro_batches) - new_exec_times = [0.0] * len(exec_times) - new_sample_ids_per_gpu = [[]] * len(sample_ids_per_gpu) - - for i in range(group_start_gpu): - new_micro_batches[i] = micro_batches[i] - new_exec_times[i] = exec_times[i] - new_sample_ids_per_gpu[i] = sample_ids_per_gpu[i] - - for i in range(group_start_gpu, group_end_gpu + needed_count + 1): - new_micro_batches[i] = micro_batches[group_end_gpu] - new_exec_times[i] = get_total_workload_fn( - micro_batches[group_end_gpu][0], next_power - ) - new_sample_ids_per_gpu[i] = sample_ids_per_gpu[group_end_gpu] + def fill_with_full_dpxcp_group() -> None: + nonlocal micro_batches, exec_times, sample_ids_per_gpu, leftovers - for i, work in enumerate(work_to_push): - new_micro_batches[group_end_gpu + needed_count + 1 + i] = work - new_exec_times[group_end_gpu + needed_count + 1 + i] = exec_times_to_push[i] - new_sample_ids_per_gpu[group_end_gpu + needed_count + 1 + i] = ( - sample_ids_to_push[i] - ) + selected: List[Tuple[int, int]] = [] + next_leftovers: List[Tuple[int, int]] = [] + packed_sequence_len = 0.0 - group_size[gid] = next_power - group_members[gid] = list(range(members[0], members[-1] + needed_count + 1)) - for pushed_gid in group_size.keys(): - if pushed_gid > gid: - group_members[pushed_gid] = [ - x + needed_count for x in group_members[pushed_gid] - ] - - return ( - new_micro_batches, - new_exec_times, - new_sample_ids_per_gpu, - group_members, - group_size, - ) + for sample_id, seq_len in sample_seqlens: + per_rank_len = seq_len / total_gpus + if packed_sequence_len + per_rank_len <= max_seq_len_per_rank: + selected.append((sample_id, seq_len)) + packed_sequence_len += per_rank_len + else: + next_leftovers.append((sample_id, seq_len)) - empty_gpus = any([not micro_batches[i] for i in range(total_gpus)]) - while empty_gpus: - micro_batches, exec_times, sample_ids_per_gpu, group_members, group_size = fill_empty_gpus( - micro_batches, exec_times, sample_ids_per_gpu, group_members, group_size + assert selected, ( + "At least one sequence should fit in the full DPxCP group; " + "try to increase 'max-seqlen-per-dp-cp-rank'." ) - empty_gpus = any([not micro_batches[i] for i in range(total_gpus)]) - total_work_after = sum(len(mb) for mb in micro_batches) - assert ( - total_work_after >= total_work_before - ), f"Samples were removed: {total_work_before} -> {total_work_after}" + selected_ids = [sample_id for sample_id, _ in selected] + selected_lens = [seq_len for _, seq_len in selected] + per_rank_work = sum(workload(seq_len, total_gpus) for _, seq_len in selected) + + micro_batches = [list(selected_lens) for _ in range(total_gpus)] + exec_times = [per_rank_work for _ in range(total_gpus)] + sample_ids_per_gpu = [list(selected_ids) for _ in range(total_gpus)] + leftovers = next_leftovers + + while any(not micro_batch for micro_batch in micro_batches): + if not fill_empty_gpus_once(): + fill_with_full_dpxcp_group() + break return micro_batches, leftovers, exec_times, sample_ids_per_gpu @@ -903,49 +876,3 @@ def dcp_gpus_needed(seq_len: int, max_seq_len_per_rank: int, min_cp_size: int = """Number of GPUs needed, rounded up to the next power of 2, lower-bounded by min_cp_size.""" raw = max(1, 2 ** ceil(log2(seq_len / max_seq_len_per_rank))) return max(min_cp_size, raw) - - -@lru_cache(maxsize=128) -def dcp_get_total_workload( - seq_length: int, max_seq_len_per_rank: int, cp_size: Optional[int] = None, min_cp_size: int = 1 -) -> float: - """Estimate workload of a sub-sample for scheduling balance.""" - if cp_size is None: - cp_size = dcp_gpus_needed(seq_length, max_seq_len_per_rank, min_cp_size) - return (seq_length * seq_length) / cp_size - - -def dcp_make_buckets_equal( - sample_seqlens: List[Tuple[int, int]], - compute_estimator: Callable, - max_seq_len_per_rank: int, - min_cp_size: int = 1, -) -> List[deque]: - """Split samples into buckets of roughly equal work, one per unique CP size.""" - seqlens = [seq_len for _, seq_len in sample_seqlens] - k = len({dcp_gpus_needed(L, max_seq_len_per_rank, min_cp_size) for L in seqlens}) - - work = [] - for _, s in sample_seqlens: - cp_size = dcp_gpus_needed(s, max_seq_len_per_rank, min_cp_size) - work.append(compute_estimator(s, cp_size)) - total_work = sum(work) - target = total_work / k - buckets, cur, cur_work = [], [], 0.0 - remaining_k = k - - for i, (sample_id, seq_len) in enumerate(sample_seqlens): - w = compute_estimator(seq_len) - projected = cur_work + w - if cur and ( - projected > target * 1.1 or len(sample_seqlens) - i <= remaining_k - len(buckets) - ): - buckets.append(deque(cur)) - cur, cur_work = [], 0.0 - remaining_k -= 1 - cur.append((sample_id, seq_len)) - cur_work += w - - if cur: - buckets.append(deque(cur)) - return buckets diff --git a/tests/unit_tests/test_sequence_packing.py b/tests/unit_tests/test_sequence_packing.py index f1fbcb53dea..fe85b3ad8f9 100644 --- a/tests/unit_tests/test_sequence_packing.py +++ b/tests/unit_tests/test_sequence_packing.py @@ -9,12 +9,14 @@ from megatron.core import parallel_state from megatron.core.datasets.data_schedule import ( + DefaultDynamicCPScheduler, _build_thd_padding_mask, _get_scheduler_max_real_num_seqs, _sanitize_thd_padding_values, get_batch_on_this_rank_for_sequence_packing, wrap_data_iterator, ) +from megatron.core.datasets.data_schedule_utils import next_hdp_group_packing_aware from megatron.core.rerun_state_machine import RerunDataIterator from megatron.training.global_vars import unset_global_variables from tests.unit_tests.test_utilities import Utils @@ -169,6 +171,41 @@ def __next__(self): return batch +def test_next_hdp_group_packing_aware_can_use_larger_cp_group_for_short_sequences(): + micro_batches, leftovers, exec_times, sample_ids = next_hdp_group_packing_aware( + [(0, 6144), (1, 2048)], total_gpus=2, max_seq_len_per_rank=4096 + ) + + assert leftovers == [] + assert micro_batches == [[6144, 2048], [6144, 2048]] + assert sample_ids == [[0, 1], [0, 1]] + assert exec_times[0] == exec_times[1] + + +def test_next_hdp_group_packing_aware_fills_non_power_of_two_dpxcp_group(): + micro_batches, leftovers, exec_times, sample_ids = next_hdp_group_packing_aware( + [(0, 50), (1, 50)], total_gpus=14, max_seq_len_per_rank=100 + ) + + assert leftovers == [] + assert micro_batches == [[50, 50] for _ in range(14)] + assert sample_ids == [[0, 1] for _ in range(14)] + assert exec_times == [exec_times[0] for _ in range(14)] + + +def test_default_dynamic_cp_scheduler_uses_packing_aware_grouping_by_default(): + scheduler = DefaultDynamicCPScheduler( + max_seqlen_per_dp_cp_rank=4096, + cp_size=2, + dp_size=1, + microbatch_group_size_per_vp_stage=None, + ) + + sample_id_groups = scheduler.get_groups_and_subsamples([(0, 6144), (1, 2048)]) + + assert sample_id_groups == [[[0, 1], [0, 1]]] + + def _gather_tensor_from_tp_group(tensor): """Gather tensors from all TP ranks for comparison.""" assert tensor is not None, "Tensor should not be None" @@ -586,13 +623,16 @@ def _check_batch(batch_all, batch_keys): if is_pp_first and (vpp is None or vpp <= 1): dp_cp_group = parallel_state.get_data_parallel_group(with_context_parallel=True) cp_size = parallel_state.get_context_parallel_world_size() - cp_group = parallel_state.get_context_parallel_group() # Count each sequence exactly once using int64 for bitwise comparison. # THD (dp_balanced): CP siblings hold identical packed data, # so reduce across DP only (not CP) on both sides. - # DCP: data is redistributed uniquely across dp_cp ranks, - # so per-microbatch CP all_reduce + scale, then dp_cp all_reduce. + # DCP: wrap_data_iterator returns packed samples before + # get_batch_on_this_rank_for_sequence_packing applies CP + # slicing, so local CP siblings still hold identical packed + # tokens here. Scale each rank by max_cp / local_cp, then + # reduce across DPxCP. A local-CP all-reduce here would + # overcount the pre-slice tokens. # Both sides multiply by max_cp so DCP (with varying local_cp) # can be normalized to the same integer scale without division. max_cp = cp_size @@ -611,21 +651,11 @@ def _check_batch(batch_all, batch_keys): # After wrap. token_sum_after = torch.tensor(0, dtype=torch.int64, device='cuda') if is_dynamic_cp: - # DCP: per-microbatch CP all_reduce + scale to max_cp, - # then dp_cp all_reduce to aggregate unique contributions. for batch in batch_all: mb_sum = batch['tokens'].long().sum().clone() local_cp = batch['local_cp_size'] if isinstance(local_cp, torch.Tensor): local_cp = local_cp.item() - mb_cp_group = parallel_state.get_dynamic_data_context_parallel_groups( - group_size=local_cp - ) - torch.distributed.all_reduce( - mb_sum, op=torch.distributed.ReduceOp.SUM, group=mb_cp_group - ) - # all_reduce result = mb_sum * local_cp. - # Scale to mb_sum * max_cp. mb_sum *= max_cp // local_cp token_sum_after += mb_sum torch.distributed.all_reduce(