已开启
gbs_train大小写修改 #903
qsuai创建于 2月12日
gbs_train大小写修改 #903
已开启
共 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_age | 339 | self.max_age = max_age |
| 340 | - self.GBS_train = GBS_train | 340 | + self.gbs_train = gbs_train |
| 341 | self.rollout_completed = torch.zeros(self.max_len, dtype=torch.int32) # 标志当前样本是否完成rollout:eod || max_tokens | 341 | 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 None | 652 | 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_train | 700 | + self.stop_partial_rollout_signal = all_consumed_group_num >= self.gbs_train |
| 701 | self.global_ready_mask = global_ready_mask | 701 | self.global_ready_mask = global_ready_mask |
| 702 | - return all_consumed_group_num >= self.GBS_train | 702 | + 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_prompt | 704 | + 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_len | 706 | 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_train | 735 | + self.stop_partial_rollout_signal = all_consumed_group_num >= self.gbs_train |
| 736 | self.global_ready_mask = global_ready_mask | 736 | 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_train | 118 | + 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: float | 32 | 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] = meta | 215 | 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 None | 559 | 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 None | 563 | return None |
| 564 | # Sort usable indices by their corresponding age values in descending order | 564 | # Sort usable indices by their corresponding age values in descending order |
| 565 | step = meta.n_samples_per_prompt | 565 | 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_train | 673 | + 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_prompt | 691 | + 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): |