已合并
特性:use_dp_batch_balance #382
nuerxiati创建于 2025年6月17日
特性:use_dp_batch_balance #382
已合并
从refs/pull/382/head合入到master
共 8 个文件变更+166-16
| @@ -4,6 +4,7 @@ | |||
| 4 | 4 | ||
| 5 | 1. **填充移除(Remove padding)** | 5 | 1. **填充移除(Remove padding)** |
| 6 | 2. **动态批量大小(Dynamic Batch Size)** | 6 | 2. **动态批量大小(Dynamic Batch Size)** |
| 7 | +3. **数据并行负载均衡(DP Batch Balance)** | ||
| 7 | 8 | ||
| 8 | --- | 9 | --- |
| 9 | 10 | ||
| @@ -65,7 +66,7 @@ rl_config: | |||
| 65 | 66 | ||
| 66 | --- | 67 | --- |
| 67 | 68 | ||
| 68 | -## 参数说明:`max_packing_token_size | 69 | +## 参数说明:max_packing_token_size |
| 69 | 70 | ||
| 70 | `max_packing_token_size` 是动态批大小(Dynamic Batch Size)机制中的核心参数,用于限制每个拼接后的 micro batch 中 token 的总数,防止因拼接过多序列而导致显存溢出(OOM)。 | 71 | `max_packing_token_size` 是动态批大小(Dynamic Batch Size)机制中的核心参数,用于限制每个拼接后的 micro batch 中 token 的总数,防止因拼接过多序列而导致显存溢出(OOM)。 |
| 71 | 72 | ||
| @@ -91,3 +92,26 @@ rl_config: | |||
| 91 | use_dynamic_bsz: true | 92 | use_dynamic_bsz: true |
| 92 | max_packing_token_size: 8192 | 93 | max_packing_token_size: 8192 |
| 93 | ``` | 94 | ``` |
| 95 | + | ||
| 96 | +# 📦 数据并行负载均衡(DP Batch Balance)特性 | ||
| 97 | + | ||
| 98 | +## 背景介绍 | ||
| 99 | + | ||
| 100 | +在数据并行(Data Parallel, DP)训练中,若各 DP 节点的序列总长度不均衡,会导致计算量少的节点提前完成等待,形成「木桶效应」。该特性通过装箱算法均衡各 DP 节点的序列总长度,减少节点间等待时间,提升分布式训练效率。 | ||
| 101 | + | ||
| 102 | +## 实现原理 | ||
| 103 | + | ||
| 104 | +1. **序列长度收集**:获取当前批次所有样本的序列长度; | ||
| 105 | +2. **动态装箱分组**:使用堆排序装箱算法,按序列长度从大到小依次分配至当前总长度最小的分组,确保各 DP 节点的总序列长度均衡; | ||
| 106 | +3. **数据分配**:将分组后的样本分配至各 DP 节点,实现计算量均衡。 | ||
| 107 | + | ||
| 108 | +<p align="center"> | ||
| 109 | + <img src="../../sources/images/remove_padding/dp_balance.png" width="400"/> | ||
| 110 | +</p> | ||
| 111 | + | ||
| 112 | +## 配置方法 | ||
| 113 | + | ||
| 114 | +```yaml | ||
| 115 | +rl_config: | ||
| 116 | + use_dp_batch_balance: true | ||
| 117 | +``` | ||
| @@ -52,8 +52,8 @@ class RLConfig(BaseConfig): | |||
| 52 | num_cpus_for_placement_group: Number of CPUs for ray worker placement group | 52 | num_cpus_for_placement_group: Number of CPUs for ray worker placement group |
| 53 | 53 | ||
| 54 | is_multimodal: Whether base model is a multimodal model or not (default: False) | 54 | is_multimodal: Whether base model is a multimodal model or not (default: False) |
| 55 | + use_dp_batch_balance: Whether to use dynamic batch size balancing across data parallel ranks (default: False) | ||
| 55 | # Default values can still be defined if no config is provided | 56 | # Default values can still be defined if no config is provided |
| 56 | - | ||
| 57 | use_remove_padding: Whether to use packed sequences for forward (default: False) | 57 | use_remove_padding: Whether to use packed sequences for forward (default: False) |
| 58 | ''' | 58 | ''' |
| 59 | 59 | ||
| @@ -98,6 +98,7 @@ class RLConfig(BaseConfig): | |||
| 98 | self.num_cpus_for_local_task = 1 | 98 | self.num_cpus_for_local_task = 1 |
| 99 | self.num_cpus_for_placement_group = 8 | 99 | self.num_cpus_for_placement_group = 8 |
| 100 | self.use_integrated_worker = True | 100 | self.use_integrated_worker = True |
| 101 | + self.use_dp_batch_balance = False | ||
| 101 | self.ref_forward_micro_batch_size = None | 102 | self.ref_forward_micro_batch_size = None |
| 102 | self.actor_forward_micro_batch_size = None | 103 | self.actor_forward_micro_batch_size = None |
| 103 | 104 | ||
| @@ -38,6 +38,7 @@ class RayBaseTrainer(object): | |||
| 38 | dataset_additional_keys: List[str] = None, | 38 | dataset_additional_keys: List[str] = None, |
| 39 | blocking: bool = False, | 39 | blocking: bool = False, |
| 40 | guarantee_order: bool = False, | 40 | guarantee_order: bool = False, |
| 41 | + use_dp_batch_balance: bool = False, | ||
| 41 | num_cpus_for_local_task: float = 0.1, | 42 | num_cpus_for_local_task: float = 0.1, |
| 42 | **kwargs): | 43 | **kwargs): |
| 43 | 44 | ||
| @@ -62,6 +63,7 @@ class RayBaseTrainer(object): | |||
| 62 | self.dataset_additional_keys = dataset_additional_keys | 63 | self.dataset_additional_keys = dataset_additional_keys |
| 63 | self.blocking = blocking | 64 | self.blocking = blocking |
| 64 | self.guarantee_order = guarantee_order | 65 | self.guarantee_order = guarantee_order |
| 66 | + self.use_dp_batch_balance = use_dp_batch_balance | ||
| 65 | self.num_cpus_for_local_task = num_cpus_for_local_task | 67 | self.num_cpus_for_local_task = num_cpus_for_local_task |
| 66 | self.kwargs = kwargs | 68 | self.kwargs = kwargs |
| 67 | 69 | ||
| @@ -298,6 +298,12 @@ class GRPOTransferDock(TransferDock): | |||
| 298 | "reward_scores", | 298 | "reward_scores", |
| 299 | "grpo_metrics", | 299 | "grpo_metrics", |
| 300 | ] | 300 | ] |
| 301 | + self.batch_seqlen_balance_mapper = { | ||
| 302 | + "ref_log_prob": ["prompt_length", "response_length"], | ||
| 303 | + "actor_log_prob": ["prompt_length", "response_length"], | ||
| 304 | + "reward_scores": ["prompt_length", "response_length"], | ||
| 305 | + "actor_train": ["prompt_length", "response_length"] | ||
| 306 | + } | ||
| 301 | if addition_columns: | 307 | if addition_columns: |
| 302 | for column in addition_columns: | 308 | for column in addition_columns: |
| 303 | if column not in self.experience_columns: | 309 | if column not in self.experience_columns: |
| @@ -338,6 +344,7 @@ class GRPOTransferDock(TransferDock): | |||
| 338 | experience_count: int = None, | 344 | experience_count: int = None, |
| 339 | indexes: List[int] = None, | 345 | indexes: List[int] = None, |
| 340 | get_n_samples: bool = True, | 346 | get_n_samples: bool = True, |
| 347 | + use_batch_seqlen_balance: bool = False | ||
| 341 | ): | 348 | ): |
| 342 | """Get padded experience data from GRPOTransferDock. | 349 | """Get padded experience data from GRPOTransferDock. |
| 343 | 350 | ||
| @@ -350,6 +357,7 @@ class GRPOTransferDock(TransferDock): | |||
| 350 | multiple: The multiple of TP to pad. | 357 | multiple: The multiple of TP to pad. |
| 351 | get_n_samples: Whether to get n samples at the same time. | 358 | get_n_samples: Whether to get n samples at the same time. |
| 352 | target_seq_len: Target sequence length. | 359 | target_seq_len: Target sequence length. |
| 360 | + use_batch_seqlen_balance: Whether to enable batch balance with seq_len. | ||
| 353 | 361 | ||
| 354 | Returns: Data dict and row numbers. | 362 | Returns: Data dict and row numbers. |
| 355 | 363 | ||
| @@ -383,11 +391,13 @@ class GRPOTransferDock(TransferDock): | |||
| 383 | f"n_samples_per_prompt: {self.n_samples_per_prompt}" | 391 | f"n_samples_per_prompt: {self.n_samples_per_prompt}" |
| 384 | ) | 392 | ) |
| 385 | indexes = self._sample_ready_index_n_samples( | 393 | indexes = self._sample_ready_index_n_samples( |
| 386 | - consumer, experience_count, experience_columns | 394 | + consumer, experience_count, experience_columns, |
| 395 | + use_batch_seqlen_balance=use_batch_seqlen_balance | ||
| 387 | ) | 396 | ) |
| 388 | else: | 397 | else: |
| 389 | indexes = self._sample_ready_index( | 398 | indexes = self._sample_ready_index( |
| 390 | - consumer, experience_count, experience_columns | 399 | + consumer, experience_count, experience_columns, |
| 400 | + use_batch_seqlen_balance=use_batch_seqlen_balance | ||
| 391 | ) | 401 | ) |
| 392 | 402 | ||
| 393 | if not indexes: | 403 | if not indexes: |
| @@ -472,6 +482,7 @@ class GRPOTransferDock(TransferDock): | |||
| 472 | experience_count: int, | 482 | experience_count: int, |
| 473 | experience_columns: List[str], | 483 | experience_columns: List[str], |
| 474 | target_seq_len: int = None, | 484 | target_seq_len: int = None, |
| 485 | + use_batch_seqlen_balance: bool = False | ||
| 475 | ) -> Optional[List[int]]: | 486 | ) -> Optional[List[int]]: |
| 476 | """Randomly select a specified number of prepared experiences from TransferDock. | 487 | """Randomly select a specified number of prepared experiences from TransferDock. |
| 477 | 488 | ||
| @@ -496,13 +507,22 @@ class GRPOTransferDock(TransferDock): | |||
| 496 | if len(usable_indexes) < experience_count: | 507 | if len(usable_indexes) < experience_count: |
| 497 | return None | 508 | return None |
| 498 | 509 | ||
| 499 | - if experience_count > 0: | 510 | + if experience_count <= 0: |
| 511 | + return None | ||
| 512 | + | ||
| 513 | + if consumer in self.batch_seqlen_balance_mapper and use_batch_seqlen_balance and len( | ||
| 514 | + usable_indexes) % experience_count == 0: | ||
| 515 | + sampled_indexes = self.batch_seqlen_balance_sampler( | ||
| 516 | + consumer, usable_indexes, experience_count, get_n_samples=False | ||
| 517 | + ) | ||
| 518 | + if not sampled_indexes: | ||
| 519 | + return None | ||
| 520 | + else: | ||
| 500 | sampled_indexes = self.batch_balencing_sampler( | 521 | sampled_indexes = self.batch_balencing_sampler( |
| 501 | experience_columns, usable_indexes, experience_count, target_seq_len | 522 | experience_columns, usable_indexes, experience_count, target_seq_len |
| 502 | ) | 523 | ) |
| 503 | - self.experience_consumer_status[consumer][sampled_indexes] = 1 | 524 | + self.experience_consumer_status[consumer][sampled_indexes] = 1 |
| 504 | - else: | 525 | + |
| 505 | - sampled_indexes = None | ||
| 506 | 526 | ||
| 507 | return sampled_indexes | 527 | return sampled_indexes |
| 508 | 528 | ||
| @@ -512,6 +532,7 @@ class GRPOTransferDock(TransferDock): | |||
| 512 | experience_count: int, | 532 | experience_count: int, |
| 513 | experience_columns: List[str], | 533 | experience_columns: List[str], |
| 514 | target_seq_len: int = None, | 534 | target_seq_len: int = None, |
| 535 | + use_batch_seqlen_balance: bool = False | ||
| 515 | ) -> Optional[List[int]]: | 536 | ) -> Optional[List[int]]: |
| 516 | """Randomly select a specified number of prepared experiences from TransferDock at multiples of n_sample. | 537 | """Randomly select a specified number of prepared experiences from TransferDock at multiples of n_sample. |
| 517 | 538 | ||
| @@ -520,6 +541,7 @@ class GRPOTransferDock(TransferDock): | |||
| 520 | experience_count: Number for rows to sample. | 541 | experience_count: Number for rows to sample. |
| 521 | experience_columns: Columns from which to sample. | 542 | experience_columns: Columns from which to sample. |
| 522 | target_seq_len: Sample according with seq_len and target_seq_len. | 543 | target_seq_len: Sample according with seq_len and target_seq_len. |
| 544 | + use_batch_seqlen_balance: Balance bath with seq_len | ||
| 523 | 545 | ||
| 524 | Returns: Sampled row numbers. | 546 | Returns: Sampled row numbers. |
| 525 | 547 | ||
| @@ -557,12 +579,20 @@ class GRPOTransferDock(TransferDock): | |||
| 557 | if len(usable_indexes) < experience_count_n_samples: | 579 | if len(usable_indexes) < experience_count_n_samples: |
| 558 | return None | 580 | return None |
| 559 | 581 | ||
| 560 | - sampled_indexes_n_sample = self.batch_balencing_sampler( | 582 | + if consumer in self.batch_seqlen_balance_mapper and use_batch_seqlen_balance and len( |
| 561 | - experience_columns, | 583 | + usable_indexes) % experience_count_n_samples == 0: |
| 562 | - usable_indexes, | 584 | + sampled_indexes_n_sample = self.batch_seqlen_balance_sampler( |
| 563 | - experience_count_n_samples, | 585 | + consumer, usable_indexes, experience_count_n_samples, get_n_samples=True |
| 564 | - target_seq_len, | 586 | + ) |
| 565 | - ) | 587 | + if not sampled_indexes_n_sample: |
| 588 | + return None | ||
| 589 | + else: | ||
| 590 | + sampled_indexes_n_sample = self.batch_balencing_sampler( | ||
| 591 | + experience_columns, | ||
| 592 | + usable_indexes, | ||
| 593 | + experience_count_n_samples, | ||
| 594 | + target_seq_len, | ||
| 595 | + ) | ||
| 566 | 596 | ||
| 567 | sampled_indexes = [] | 597 | sampled_indexes = [] |
| 568 | for n_sample_index in sampled_indexes_n_sample: | 598 | for n_sample_index in sampled_indexes_n_sample: |
| @@ -611,6 +641,34 @@ class GRPOTransferDock(TransferDock): | |||
| 611 | """ | 641 | """ |
| 612 | return self.experience_consumer_status | 642 | return self.experience_consumer_status |
| 613 | 643 | ||
| 644 | + def batch_seqlen_balance_sampler( | ||
| 645 | + self, consumer, usable_indexes, experience_count, get_n_samples=False | ||
| 646 | + ): | ||
| 647 | + from mindspeed_rl.utils.seqlen_balancing import get_seqlen_balanced_partitions | ||
| 648 | + | ||
| 649 | + if len(usable_indexes) == experience_count: | ||
| 650 | + sampled_indexes = [int(usable_indexes[i]) for i in range(experience_count)] | ||
| 651 | + return sampled_indexes | ||
| 652 | + seq_len_columns = self.batch_seqlen_balance_mapper.get(consumer) | ||
| 653 | + if get_n_samples: | ||
| 654 | + seq_len_list = [ | ||
| 655 | + sum([self.experience_data[key][idx * self.n_samples_per_prompt + addition].item() | ||
| 656 | + for addition in range(self.n_samples_per_prompt) for key in seq_len_columns]) | ||
| 657 | + for idx in usable_indexes | ||
| 658 | + ] | ||
| 659 | + else: | ||
| 660 | + seq_len_list = [ | ||
| 661 | + sum([self.experience_data[key][idx].item() for key in seq_len_columns]) | ||
| 662 | + for idx in usable_indexes | ||
| 663 | + ] | ||
| 664 | + k_partitions = len(seq_len_list) // experience_count | ||
| 665 | + sampled_indexes_idx = get_seqlen_balanced_partitions(seq_len_list, k_partitions, equal_size=True) | ||
| 666 | + if len(sampled_indexes_idx) > 0: | ||
| 667 | + sampled_indexes = [int(usable_indexes[i]) for i in sampled_indexes_idx[0]] | ||
| 668 | + else: | ||
| 669 | + sampled_indexes = None | ||
| 670 | + return sampled_indexes | ||
| 671 | + | ||
| 614 | def batch_balencing_sampler( | 672 | def batch_balencing_sampler( |
| 615 | self, experience_columns, usable_indexes, experience_count, target_seq_len=None | 673 | self, experience_columns, usable_indexes, experience_count, target_seq_len=None |
| 616 | ): | 674 | ): |
| @@ -122,6 +122,70 @@ def karmarkar_karp(seqlen_list: List[int], k_partitions: int, equal_size: bool): | |||
| 122 | return partitions | 122 | return partitions |
| 123 | 123 | ||
| 124 | 124 | ||
| 125 | +def heapq_partition(seqlen_list: List[int], k_partitions: int, equal_size: bool): | ||
| 126 | + equal_part_num = len(seqlen_list) // k_partitions | ||
| 127 | + | ||
| 128 | + sorted_seqlen = sorted([(seqlen, i) for i, seqlen in enumerate(seqlen_list)], reverse=True) | ||
| 129 | + | ||
| 130 | + # Initialize the heap: each group maintains [current sum, number of elements, group index, elements in the group] | ||
| 131 | + groups = [[0, 0, i, []] for i in range(k_partitions)] | ||
| 132 | + heapq.heapify(groups) | ||
| 133 | + | ||
| 134 | + partitions = [] | ||
| 135 | + for seqlen, i in sorted_seqlen: | ||
| 136 | + current_group = heapq.heappop(groups) | ||
| 137 | + current_group[3].append(i) | ||
| 138 | + current_group[0] += seqlen | ||
| 139 | + current_group[1] += 1 | ||
| 140 | + if equal_size: | ||
| 141 | + if current_group[1] < equal_part_num: | ||
| 142 | + heapq.heappush(groups, current_group) | ||
| 143 | + else: | ||
| 144 | + partitions.append(current_group[3]) | ||
| 145 | + else: | ||
| 146 | + heapq.heappush(groups, current_group) | ||
| 147 | + | ||
| 148 | + partitions.extend([group[3] for group in groups]) | ||
| 149 | + | ||
| 150 | + if equal_size: | ||
| 151 | + for i, partition in enumerate(partitions): | ||
| 152 | + if len(partition) * k_partitions != len(seqlen_list): | ||
| 153 | + raise ValueError(f"Partition {i} has {len(partition)} items, expected {len(seqlen_list) // k_partitions}") | ||
| 154 | + return partitions | ||
| 155 | + | ||
| 156 | + | ||
| 157 | +def get_seqlen_balanced_partitions(seqlen_list: List[int], k_partitions: int, equal_size: bool): | ||
| 158 | + """get order of seq lengths to make partitions balanced, this is | ||
| 159 | + used in balancing sum of seq length across dp ranks and micro batches | ||
| 160 | + Parameters: | ||
| 161 | + seqlen_list (List[int]): | ||
| 162 | + seq lengths of each items | ||
| 163 | + k_partitions (int): | ||
| 164 | + resulting number of partitions | ||
| 165 | + equal_size (bool): | ||
| 166 | + if True, number of items in each partitions must be equal. | ||
| 167 | + if False, only consider balancing the sum, each partition can have | ||
| 168 | + variable number of items | ||
| 169 | + Returns: | ||
| 170 | + partitions (List[List[int]]): | ||
| 171 | + return k_partitions list containing the index of items. | ||
| 172 | + """ | ||
| 173 | + if k_partitions > len(seqlen_list): | ||
| 174 | + raise ValueError(f"number of items:[{len(seqlen_list)}] < k_partitions:[{k_partitions}]") | ||
| 175 | + | ||
| 176 | + def _check_and_sort_partitions(partitions): | ||
| 177 | + seen_idx = set() | ||
| 178 | + sorted_partitions = [None] * k_partitions | ||
| 179 | + for i, partition in enumerate(partitions): | ||
| 180 | + for idx in partition: | ||
| 181 | + seen_idx.add(idx) | ||
| 182 | + sorted_partitions[i] = sorted(partition) | ||
| 183 | + return sorted_partitions | ||
| 184 | + | ||
| 185 | + partitions = heapq_partition(seqlen_list=seqlen_list, k_partitions=k_partitions, equal_size=equal_size) | ||
| 186 | + return _check_and_sort_partitions(partitions) | ||
| 187 | + | ||
| 188 | + | ||
| 125 | def rearrange_micro_batches(seqlen_list: List[int], max_token_len: int, dp_group=None): | 189 | def rearrange_micro_batches(seqlen_list: List[int], max_token_len: int, dp_group=None): |
| 126 | """get order of seq lengths to make partitions balanced, this is | 190 | """get order of seq lengths to make partitions balanced, this is |
| 127 | used in balancing sum of seq length across dp ranks and micro batches | 191 | used in balancing sum of seq length across dp ranks and micro batches |
| @@ -136,7 +200,6 @@ def rearrange_micro_batches(seqlen_list: List[int], max_token_len: int, dp_group | |||
| 136 | if max(seqlen_list) > max_token_len: | 200 | if max(seqlen_list) > max_token_len: |
| 137 | raise ValueError(f"seqlen of items:[{max(seqlen_list)}] must <= max_token_len:[{max_token_len}]") | 201 | raise ValueError(f"seqlen of items:[{max(seqlen_list)}] must <= max_token_len:[{max_token_len}]") |
| 138 | 202 | ||
| 139 | - | ||
| 140 | # Calculate the minimum number of bins | 203 | # Calculate the minimum number of bins |
| 141 | total_sum_of_seqlen = sum(seqlen_list) | 204 | total_sum_of_seqlen = sum(seqlen_list) |
| 142 | if total_sum_of_seqlen % max_token_len == 0: | 205 | if total_sum_of_seqlen % max_token_len == 0: |
| @@ -261,7 +261,8 @@ class BaseWorker(BaseRayWorker, ABC): | |||
| 261 | if rank_flg: | 261 | if rank_flg: |
| 262 | batch_data, index = ray.get(self.td.get_experience.remote(experience_consumer_stage, experience_columns, | 262 | batch_data, index = ray.get(self.td.get_experience.remote(experience_consumer_stage, experience_columns, |
| 263 | experience_count, indexes=indexes, | 263 | experience_count, indexes=indexes, |
| 264 | - get_n_samples=get_n_samples)) # cpu数据 | 264 | + get_n_samples=get_n_samples, |
| 265 | + use_batch_seqlen_balance=self.rl_config.use_dp_batch_balance)) # cpu数据 | ||
| 265 | if not index: # 判断是否取出数据,未取出数据为-1 | 266 | if not index: # 判断是否取出数据,未取出数据为-1 |
| 266 | index = [-1] * experience_count | 267 | index = [-1] * experience_count |
| 267 | 268 | ||
| @@ -56,6 +56,7 @@ actor_config: | |||
| 56 | rl_config: | 56 | rl_config: |
| 57 | use_integrated_worker: true | 57 | use_integrated_worker: true |
| 58 | blocking: true | 58 | blocking: true |
| 59 | + use_dp_batch_balance: true | ||
| 59 | use_dynamic_bsz: true | 60 | use_dynamic_bsz: true |
| 60 | max_packing_token_size: 8192 | 61 | max_packing_token_size: 8192 |
| 61 | actor_forward_micro_batch_size: 8 | 62 | actor_forward_micro_batch_size: 8 |