已合并
特性:use_dp_batch_balance #382
nuerxiati创建于 2025年6月17日
特性:use_dp_batch_balance #382
已合并
nuerxiati创建于 2025年6月17日
从refs/pull/382/head合入到master
共 8 个文件变更+166-16
@@ -4,6 +4,7 @@
4 4 
51. **填充移除(Remove padding)** 51. **填充移除(Remove padding)**
62. **动态批量大小(Dynamic Batch Size)**62. **动态批量大小(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_size69+## 参数说明: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: true92 use_dynamic_bsz: true
92 max_packing_token_size: 819293 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 group52 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 provided56 # 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 = 198 self.num_cpus_for_local_task = 1
99 self.num_cpus_for_placement_group = 899 self.num_cpus_for_placement_group = 8
100 self.use_integrated_worker = True100 self.use_integrated_worker = True
101+ self.use_dp_batch_balance = False
101 self.ref_forward_micro_batch_size = None102 self.ref_forward_micro_batch_size = None
102 self.actor_forward_micro_batch_size = None103 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_keys63 self.dataset_additional_keys = dataset_additional_keys
63 self.blocking = blocking64 self.blocking = blocking
64 self.guarantee_order = guarantee_order65 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_task67 self.num_cpus_for_local_task = num_cpus_for_local_task
66 self.kwargs = kwargs68 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_columns394+ 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_columns399+ 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 None508 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_len522 experience_columns, usable_indexes, experience_count, target_seq_len
502 )523 )
503- self.experience_consumer_status[consumer][sampled_indexes] = 1524+ self.experience_consumer_status[consumer][sampled_indexes] = 1
504- else:525+ 
505- sampled_indexes = None
506 526 
507 return sampled_indexes527 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 None580 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_status642 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=None673 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 partitions122 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+ 
125def rearrange_micro_batches(seqlen_list: List[int], max_token_len: int, dp_group=None):189def 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 is190 """get order of seq lengths to make partitions balanced, this is
127 used in balancing sum of seq length across dp ranks and micro batches191 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 bins203 # 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: # 判断是否取出数据,未取出数据为-1266 if not index: # 判断是否取出数据,未取出数据为-1
266 index = [-1] * experience_count267 index = [-1] * experience_count
267 268 
@@ -56,6 +56,7 @@ actor_config:
56rl_config:56rl_config:
57 use_integrated_worker: true57 use_integrated_worker: true
58 blocking: true58 blocking: true
59+ use_dp_batch_balance: true
59 use_dynamic_bsz: true60 use_dynamic_bsz: true
60 max_packing_token_size: 819261 max_packing_token_size: 8192
61 actor_forward_micro_batch_size: 862 actor_forward_micro_batch_size: 8