已开启
gbs_train大小写修改 #903
gbs_train大小写修改 #903
已开启
qsuai创建于 2月12日
5 个文件变更+24-24
@@ -386,7 +386,7 @@ class DataStrategy:
386 experience_consumers=experience_consumers,386 experience_consumers=experience_consumers,
387 metrics=metrics,387 metrics=metrics,
388 max_age=max_age,388 max_age=max_age,
389- GBS_train=gbs_train,389+ gbs_train=gbs_train,
390 )390 )
391 additional_keys = dataset_additional_keys or []391 additional_keys = dataset_additional_keys or []
392 consumer_columns = {392 consumer_columns = {
@@ -466,7 +466,7 @@ class DataStrategy:
466 n_samples_per_prompt=n_samples_per_prompt,466 n_samples_per_prompt=n_samples_per_prompt,
467 metrics=metrics,467 metrics=metrics,
468 max_age=max_age,468 max_age=max_age,
469- GBS_train=gbs_train,469+ gbs_train=gbs_train,
470 addition_columns=addition_columns,470 addition_columns=addition_columns,
471 )471 )
472 if is_multimodal():472 if is_multimodal():
@@ -731,7 +731,7 @@ class DataStrategy:
731 max_num_prompt_in_batch,731 max_num_prompt_in_batch,
732 n_samples_per_prompt,732 n_samples_per_prompt,
733 max_age=1,733 max_age=1,
734- GBS_train=max_num_prompt_in_batch,734+ gbs_train=max_num_prompt_in_batch,
735 metrics=metrics,735 metrics=metrics,
736 addition_columns=addition_columns,736 addition_columns=addition_columns,
737 addition_consumers=addition_consumers,737 addition_consumers=addition_consumers,
@@ -740,7 +740,7 @@ class DataStrategy:
740 td_max_len,740 td_max_len,
741 n_samples_per_prompt,741 n_samples_per_prompt,
742 max_age=max_age,742 max_age=max_age,
743- GBS_train=gbs_train,743+ gbs_train=gbs_train,
744 metrics=metrics,744 metrics=metrics,
745 addition_columns=addition_columns,745 addition_columns=addition_columns,
746 addition_consumers=addition_consumers,746 addition_consumers=addition_consumers,
@@ -759,7 +759,7 @@ class DataStrategy:
759 td_max_len,759 td_max_len,
760 n_samples_per_prompt,760 n_samples_per_prompt,
761 max_age=max_age,761 max_age=max_age,
762- GBS_train=gbs_train,762+ gbs_train=gbs_train,
763 metrics=metrics,763 metrics=metrics,
764 addition_columns=addition_columns,764 addition_columns=addition_columns,
765 addition_consumers=addition_consumers,765 addition_consumers=addition_consumers,
@@ -258,7 +258,7 @@ class GRPOTransferDock(TransferDock):
258 n_samples_per_prompt: int,258 n_samples_per_prompt: int,
259 metrics=None,259 metrics=None,
260 max_age: int = 1,260 max_age: int = 1,
261- GBS_train: int = 0,261+ gbs_train: int = 0,
262 addition_columns: Union[List[str], None] = None,262 addition_columns: Union[List[str], None] = None,
263 addition_consumers: Union[List[str], None] = None,263 addition_consumers: Union[List[str], None] = None,
264 timeout: Union[int, None] = None,264 timeout: Union[int, None] = None,
@@ -337,7 +337,7 @@ class GRPOTransferDock(TransferDock):
337 }337 }
338 338 
339 self.max_age = max_age339 self.max_age = max_age
340- self.GBS_train = GBS_train340+ self.gbs_train = gbs_train
341 self.rollout_completed = torch.zeros(self.max_len, dtype=torch.int32) # 标志当前样本是否完成rollout:eod || max_tokens341 self.rollout_completed = torch.zeros(self.max_len, dtype=torch.int32) # 标志当前样本是否完成rollout:eod || max_tokens
342 self.age = torch.zeros(self.max_len, dtype=torch.int32) # 落后当前actor参数的训练步数,是否需要按age排序?age的更新需要在TD逐出和重排序的时候做342 self.age = torch.zeros(self.max_len, dtype=torch.int32) # 落后当前actor参数的训练步数,是否需要按age排序?age的更新需要在TD逐出和重排序的时候做
343 self.enable_partial_rollout = max_age > 1 # max_age = 1 是续推0次,因为rollout_completed的判断是在TD外面做的343 self.enable_partial_rollout = max_age > 1 # max_age = 1 是续推0次,因为rollout_completed的判断是在TD外面做的
@@ -499,7 +499,7 @@ class GRPOTransferDock(TransferDock):
499 )499 )
500 data_dict = remove_padding_tensor_dict_to_dict(data_dict)500 data_dict = remove_padding_tensor_dict_to_dict(data_dict)
501 501 
502- if self.enable_partial_rollout and self.GBS_train == 0:502+ if self.enable_partial_rollout and self.gbs_train == 0:
503 raise ValueError("GBS for update must be provided when enabling partial rollout")503 raise ValueError("GBS for update must be provided when enabling partial rollout")
504 504 
505 if self.enable_partial_rollout and 'responses' in data_dict.keys():505 if self.enable_partial_rollout and 'responses' in data_dict.keys():
@@ -648,7 +648,7 @@ class GRPOTransferDock(TransferDock):
648 usable_indexes = (not_consumed_indexes & data_ready_indexes & update_ready_group_indexes).nonzero(as_tuple=True)[0]648 usable_indexes = (not_consumed_indexes & data_ready_indexes & update_ready_group_indexes).nonzero(as_tuple=True)[0]
649 649 
650 if len(usable_indexes) < experience_count_n_samples or \650 if len(usable_indexes) < experience_count_n_samples or \
651- self.experience_consumer_status[consumer].sum() >= self.GBS_train * self.n_samples_per_prompt:651+ self.experience_consumer_status[consumer].sum() >= self.gbs_train * self.n_samples_per_prompt:
652 return None652 return None
653 653 
654 if self.enable_partial_rollout:654 if self.enable_partial_rollout:
@@ -693,15 +693,15 @@ class GRPOTransferDock(TransferDock):
693 693 
694 """694 """
695 if self.enable_partial_rollout:695 if self.enable_partial_rollout:
696- if self.GBS_train == 0:696+ if self.gbs_train == 0:
697 raise ValueError("GBS for update must be provided when enabling partial rollout")697 raise ValueError("GBS for update must be provided when enabling partial rollout")
698 if consumer == 'actor_rollout':698 if consumer == 'actor_rollout':
699 all_consumed_group_num, global_ready_mask, _ = self.find_all_consumed_n_samples_groups(consumer='actor_rollout')699 all_consumed_group_num, global_ready_mask, _ = self.find_all_consumed_n_samples_groups(consumer='actor_rollout')
700- self.stop_partial_rollout_signal = all_consumed_group_num >= self.GBS_train700+ self.stop_partial_rollout_signal = all_consumed_group_num >= self.gbs_train
701 self.global_ready_mask = global_ready_mask701 self.global_ready_mask = global_ready_mask
702- return all_consumed_group_num >= self.GBS_train702+ return all_consumed_group_num >= self.gbs_train
703 else:703 else:
704- return self.experience_consumer_status[consumer].sum() == self.GBS_train * self.n_samples_per_prompt704+ return self.experience_consumer_status[consumer].sum() == self.gbs_train * self.n_samples_per_prompt
705 else:705 else:
706 return self.experience_consumer_status[consumer].sum() == self.max_len706 return self.experience_consumer_status[consumer].sum() == self.max_len
707 707 
@@ -732,7 +732,7 @@ class GRPOTransferDock(TransferDock):
732 732 
733 def get_update_ready(self, require_max_age_all_finished=True):733 def get_update_ready(self, require_max_age_all_finished=True):
734 all_consumed_group_num, global_ready_mask, _ = self.find_all_consumed_n_samples_groups(consumer='actor_rollout')734 all_consumed_group_num, global_ready_mask, _ = self.find_all_consumed_n_samples_groups(consumer='actor_rollout')
735- self.stop_partial_rollout_signal = all_consumed_group_num >= self.GBS_train735+ self.stop_partial_rollout_signal = all_consumed_group_num >= self.gbs_train
736 self.global_ready_mask = global_ready_mask736 self.global_ready_mask = global_ready_mask
737 737 
738 if require_max_age_all_finished:738 if require_max_age_all_finished:
@@ -103,7 +103,7 @@ class TransferQueueClient:
103 metrics: Metric = None,103 metrics: Metric = None,
104 topic: str = DEFAULT_TOPIC,104 topic: str = DEFAULT_TOPIC,
105 max_age: int = 1,105 max_age: int = 1,
106- GBS_train: int = 0,106+ gbs_train: int = 0,
107 ) -> None:107 ) -> None:
108 """108 """
109 Register a topic on the manager and provision per-shard storage.109 Register a topic on the manager and provision per-shard storage.
@@ -115,7 +115,7 @@ class TransferQueueClient:
115 if not topic:115 if not topic:
116 raise ValueError("Topic must be non-empty")116 raise ValueError("Topic must be non-empty")
117 ray.get(self.manager.add_topic.remote(117 ray.get(self.manager.add_topic.remote(
118- topic, prompts_num, n_samples_per_prompt, experience_columns, experience_consumers, timeout=timeout, metrics=metrics, max_age=max_age, GBS_train=GBS_train118+ topic, prompts_num, n_samples_per_prompt, experience_columns, experience_consumers, timeout=timeout, metrics=metrics, max_age=max_age, gbs_train=gbs_train
119 ))119 ))
120 self.logger.info(120 self.logger.info(
121 f"TQ_Client: Created topic '{topic}' with prompts_num={prompts_num}, "121 f"TQ_Client: Created topic '{topic}' with prompts_num={prompts_num}, "
@@ -31,7 +31,7 @@ class TopicMeta:
31 experience_consumers: List[str]31 experience_consumers: List[str]
32 timeout: float32 timeout: float
33 max_age: Optional[int]33 max_age: Optional[int]
34- GBS_train: Optional[int]34+ gbs_train: Optional[int]
35 35 
36 max_len: int = field(init=False)36 max_len: int = field(init=False)
37 prompts_per_shard: int = field(init=False)37 prompts_per_shard: int = field(init=False)
@@ -192,7 +192,7 @@ class TransferQueueManager:
192 timeout: float,192 timeout: float,
193 metrics: Metric,193 metrics: Metric,
194 max_age: int,194 max_age: int,
195- GBS_train: int,195+ gbs_train: int,
196 ) -> None:196 ) -> None:
197 """197 """
198 Register a topic and create storage tables on all shards.198 Register a topic and create storage tables on all shards.
@@ -210,7 +210,7 @@ class TransferQueueManager:
210 experience_consumers=experience_consumers,210 experience_consumers=experience_consumers,
211 timeout=timeout,211 timeout=timeout,
212 max_age=max_age,212 max_age=max_age,
213- GBS_train=GBS_train,213+ gbs_train=gbs_train,
214 )214 )
215 self.topics[topic] = meta215 self.topics[topic] = meta
216 for i, actor in enumerate(self.data_actors):216 for i, actor in enumerate(self.data_actors):
@@ -559,7 +559,7 @@ class TransferQueueManager:
559 return None559 return None
560 560 
561 if meta.enable_partial_rollout:561 if meta.enable_partial_rollout:
562- if meta.experience_consumer_status[consumer].sum() >= meta.GBS_train * meta.n_samples_per_prompt:562+ if meta.experience_consumer_status[consumer].sum() >= meta.gbs_train * meta.n_samples_per_prompt:
563 return None563 return None
564 # Sort usable indices by their corresponding age values in descending order564 # Sort usable indices by their corresponding age values in descending order
565 step = meta.n_samples_per_prompt565 step = meta.n_samples_per_prompt
@@ -670,7 +670,7 @@ class TransferQueueManager:
670 ).all(dim=1)670 ).all(dim=1)
671 exhaustively_consumed_groups_count = exhaustively_consumed_groups_mask.sum().item()671 exhaustively_consumed_groups_count = exhaustively_consumed_groups_mask.sum().item()
672 672 
673- return exhaustively_consumed_groups_count >= meta.GBS_train673+ return exhaustively_consumed_groups_count >= meta.gbs_train
674 674 
675 def all_consumed(self, topic: str, consumer: str, get_n_samples: bool) -> bool:675 def all_consumed(self, topic: str, consumer: str, get_n_samples: bool) -> bool:
676 """676 """
@@ -685,10 +685,10 @@ class TransferQueueManager:
685 685 
686 # Verify if total consumed data volume is sufficient (equal to GBS * n_samples)686 # Verify if total consumed data volume is sufficient (equal to GBS * n_samples)
687 if meta.enable_partial_rollout:687 if meta.enable_partial_rollout:
688- if meta.GBS_train == 0:688+ if meta.gbs_train == 0:
689 raise ValueError("GBS for update must be provided when enabling partial rollout")689 raise ValueError("GBS for update must be provided when enabling partial rollout")
690 if get_n_samples:690 if get_n_samples:
691- return int(meta.experience_consumer_status[consumer].sum().item()) == meta.GBS_train * meta.n_samples_per_prompt691+ return int(meta.experience_consumer_status[consumer].sum().item()) == meta.gbs_train * meta.n_samples_per_prompt
692 else:692 else:
693 return self._check_exhaustively_consumed_groups(topic, consumer)693 return self._check_exhaustively_consumed_groups(topic, consumer)
694 else:694 else:
@@ -225,7 +225,7 @@ class TestTransferQueue(DistributedTest):
225 metrics=Metric(),225 metrics=Metric(),
226 topic=topic,226 topic=topic,
227 max_age=max_age,227 max_age=max_age,
228- GBS_train=gbs_train,228+ gbs_train=gbs_train,
229 )229 )
230 230 
231 def test_put_get(self):231 def test_put_get(self):