已合并
worker初始化优化&&DAPO去除冗余代码&&其他一些小优化 #442
LookAround0301创建于 2025年6月30日
worker初始化优化&&DAPO去除冗余代码&&其他一些小优化 #442
已合并
从refs/pull/442/head合入到master
共 17 个文件变更+53-96
| @@ -17,7 +17,7 @@ from ray.util import placement_group | |||
| 17 | 17 | ||
| 18 | from mindspeed_rl.utils import seed_all | 18 | from mindspeed_rl.utils import seed_all |
| 19 | from mindspeed_rl.utils import get_tokenizer | 19 | from mindspeed_rl.utils import get_tokenizer |
| 20 | -from mindspeed_rl.utils.utils import MsProbe | 20 | +from mindspeed_rl.utils.utils import MsProbe, get_node_nums |
| 21 | from mindspeed_rl.utils.loggers import Loggers | 21 | from mindspeed_rl.utils.loggers import Loggers |
| 22 | from mindspeed_rl.utils.utils import parse_args_from_config | 22 | from mindspeed_rl.utils.utils import parse_args_from_config |
| 23 | from mindspeed_rl.config_cls.validate_config import validate_rl_args | 23 | from mindspeed_rl.config_cls.validate_config import validate_rl_args |
| @@ -109,10 +109,6 @@ def train(config): | |||
| 109 | 109 | ||
| 110 | reward_list.append(reward_worker) | 110 | reward_list.append(reward_worker) |
| 111 | 111 | ||
| 112 | - def get_node_nums(): | ||
| 113 | - nodes = ray.nodes() | ||
| 114 | - return len([node for node in nodes if node.get("Alive", False)]) | ||
| 115 | - | ||
| 116 | rule_reward_num_process = get_node_nums() | 112 | rule_reward_num_process = get_node_nums() |
| 117 | if rl_config.rule_reward: | 113 | if rl_config.rule_reward: |
| 118 | pg = placement_group( | 114 | pg = placement_group( |
| @@ -19,7 +19,7 @@ from mindspeed_rl.config_cls.validate_config import validate_rl_args | |||
| 19 | from mindspeed_rl.utils import get_tokenizer | 19 | from mindspeed_rl.utils import get_tokenizer |
| 20 | from mindspeed_rl.datasets.build_dataset import build_train_valid_test_datasets | 20 | from mindspeed_rl.datasets.build_dataset import build_train_valid_test_datasets |
| 21 | from mindspeed_rl.utils import seed_all | 21 | from mindspeed_rl.utils import seed_all |
| 22 | -from mindspeed_rl.utils.utils import MsProbe | 22 | +from mindspeed_rl.utils.utils import MsProbe, get_node_nums |
| 23 | from mindspeed_rl.utils.loggers import Loggers | 23 | from mindspeed_rl.utils.loggers import Loggers |
| 24 | from mindspeed_rl.utils.utils import parse_args_from_config | 24 | from mindspeed_rl.utils.utils import parse_args_from_config |
| 25 | from mindspeed_rl.config_cls.megatron_config import MegatronConfig | 25 | from mindspeed_rl.config_cls.megatron_config import MegatronConfig |
| @@ -125,10 +125,6 @@ def train(config): | |||
| 125 | 125 | ||
| 126 | reward_list.append(reward_worker) | 126 | reward_list.append(reward_worker) |
| 127 | 127 | ||
| 128 | - def get_node_nums(): | ||
| 129 | - nodes = ray.nodes() | ||
| 130 | - return len([node for node in nodes if node.get("Alive", False)]) | ||
| 131 | - | ||
| 132 | rule_reward_num_process = get_node_nums() | 128 | rule_reward_num_process = get_node_nums() |
| 133 | if rl_config.rule_reward: | 129 | if rl_config.rule_reward: |
| 134 | pg = placement_group( | 130 | pg = placement_group( |
| @@ -19,7 +19,7 @@ from mindspeed_rl.config_cls.validate_config import validate_rl_args | |||
| 19 | from mindspeed_rl.utils import get_tokenizer | 19 | from mindspeed_rl.utils import get_tokenizer |
| 20 | from mindspeed_rl.datasets.build_dataset import build_train_valid_test_datasets | 20 | from mindspeed_rl.datasets.build_dataset import build_train_valid_test_datasets |
| 21 | from mindspeed_rl.utils import seed_all | 21 | from mindspeed_rl.utils import seed_all |
| 22 | -from mindspeed_rl.utils.utils import MsProbe | 22 | +from mindspeed_rl.utils.utils import MsProbe, get_node_nums |
| 23 | from mindspeed_rl.utils.loggers import Loggers | 23 | from mindspeed_rl.utils.loggers import Loggers |
| 24 | from mindspeed_rl.utils.utils import parse_args_from_config | 24 | from mindspeed_rl.utils.utils import parse_args_from_config |
| 25 | from mindspeed_rl.config_cls.megatron_config import MegatronConfig | 25 | from mindspeed_rl.config_cls.megatron_config import MegatronConfig |
| @@ -29,7 +29,7 @@ from mindspeed_rl.config_cls.mindstudio_config import ProfilerConfig, MsprobeCon | |||
| 29 | from mindspeed_rl.datasets.prompt_dataset import PromptDataset | 29 | from mindspeed_rl.datasets.prompt_dataset import PromptDataset |
| 30 | from mindspeed_rl.datasets.dataloader import PromptDataLoader | 30 | from mindspeed_rl.datasets.dataloader import PromptDataLoader |
| 31 | from mindspeed_rl.workers.rule_reward import RuleReward | 31 | from mindspeed_rl.workers.rule_reward import RuleReward |
| 32 | -from mindspeed_rl.trainer.ppo_trainer import RayPPOTrainer | 32 | +from mindspeed_rl.trainer.ppo_trainer_hybrid import RayPPOTrainer |
| 33 | from mindspeed_rl.workers.scheduler.launcher import RayActorGroup | 33 | from mindspeed_rl.workers.scheduler.launcher import RayActorGroup |
| 34 | from mindspeed_rl.workers.actor_hybrid_worker import ActorHybridWorker | 34 | from mindspeed_rl.workers.actor_hybrid_worker import ActorHybridWorker |
| 35 | from mindspeed_rl.workers.reference_woker import ReferenceWorker | 35 | from mindspeed_rl.workers.reference_woker import ReferenceWorker |
| @@ -148,10 +148,6 @@ def train(config): | |||
| 148 | global_batch_size=actor_config.global_batch_size * rl_config.n_samples_per_prompt | 148 | global_batch_size=actor_config.global_batch_size * rl_config.n_samples_per_prompt |
| 149 | ).initialize() | 149 | ).initialize() |
| 150 | 150 | ||
| 151 | - def get_node_nums(): | ||
| 152 | - nodes = ray.nodes() | ||
| 153 | - return len([node for node in nodes if node.get("Alive", False)]) | ||
| 154 | - | ||
| 155 | rule_reward_num_process = get_node_nums() | 151 | rule_reward_num_process = get_node_nums() |
| 156 | if rl_config.rule_reward: | 152 | if rl_config.rule_reward: |
| 157 | pg = placement_group( | 153 | pg = placement_group( |
| @@ -34,7 +34,6 @@ megatron_training: | |||
| 34 | no_shuffle: true | 34 | no_shuffle: true |
| 35 | full_shuffle_instruction_dataset: false | 35 | full_shuffle_instruction_dataset: false |
| 36 | 36 | ||
| 37 | - | ||
| 38 | actor_config: | 37 | actor_config: |
| 39 | model: qwen25_7b | 38 | model: qwen25_7b |
| 40 | micro_batch_size: 1 | 39 | micro_batch_size: 1 |
| @@ -3,6 +3,7 @@ import os | |||
| 3 | from mindspeed_rl.config_cls.rl_config import RLConfig | 3 | from mindspeed_rl.config_cls.rl_config import RLConfig |
| 4 | from mindspeed_rl.config_cls.megatron_config import MegatronConfig | 4 | from mindspeed_rl.config_cls.megatron_config import MegatronConfig |
| 5 | from mindspeed_rl.config_cls.generate_config import GenerateConfig | 5 | from mindspeed_rl.config_cls.generate_config import GenerateConfig |
| 6 | +from mindspeed_rl.utils.utils import get_node_nums | ||
| 6 | 7 | ||
| 7 | 8 | ||
| 8 | def validate_rl_args( | 9 | def validate_rl_args( |
| @@ -226,9 +227,14 @@ def validate_rl_args( | |||
| 226 | (reward_config.global_batch_size * rl_config.n_samples_per_prompt // reward_data_parallel_size) | 227 | (reward_config.global_batch_size * rl_config.n_samples_per_prompt // reward_data_parallel_size) |
| 227 | ) | 228 | ) |
| 228 | else: | 229 | else: |
| 230 | + rule_reward_num_process = get_node_nums() | ||
| 229 | rl_config.reward_dispatch_size = ( | 231 | rl_config.reward_dispatch_size = ( |
| 230 | - rl_config.reward_dispatch_size or (reward_config.global_batch_size * rl_config.n_samples_per_prompt) | 232 | + rl_config.reward_dispatch_size or (reward_config.global_batch_size * rl_config.n_samples_per_prompt // rule_reward_num_process) |
| 231 | ) | 233 | ) |
| 234 | + if reward_config.global_batch_size % rule_reward_num_process != 0: | ||
| 235 | + raise ValueError( | ||
| 236 | + f"Reward dispatch size configuration error!" | ||
| 237 | + f"global_batch_size {reward_config.global_batch_size} must be divisible by the number of nodes in the ray cluster") | ||
LP 这段校验可以删掉,后面会有校验经验计数与全局批次关系 ![]() ![]() | |||
| 232 | 238 | ||
| 233 | # 若开启dapo动态采样,需要更新用于计算adv和logp的dispatch_size | 239 | # 若开启dapo动态采样,需要更新用于计算adv和logp的dispatch_size |
| 234 | if rl_config.filter_groups_enable: | 240 | if rl_config.filter_groups_enable: |
| @@ -287,12 +293,12 @@ def validate_rl_args( | |||
| 287 | rl_config.critic_update_dispatch_size, | 293 | rl_config.critic_update_dispatch_size, |
| 288 | "Critic Update") | 294 | "Critic Update") |
| 289 | 295 | ||
| 290 | - if not rl_config.use_integrated_worker: | 296 | + rl_config.actor_update_dispatch_size = ( |
| 291 | - rl_config.actor_update_dispatch_size = ( | 297 | + rl_config.actor_update_dispatch_size or |
| 292 | - rl_config.actor_update_dispatch_size or | 298 | + (actor_config.global_batch_size * rl_config.n_samples_per_prompt // actor_data_parallel_size) |
| 293 | - (actor_config.global_batch_size * rl_config.n_samples_per_prompt // actor_data_parallel_size) | 299 | + ) |
| 294 | - ) | ||
| 295 | 300 | ||
| 301 | + if not rl_config.use_integrated_worker: | ||
| 296 | # 若开启dapo动态采样,需要更新用于actor更新的dispatch_size | 302 | # 若开启dapo动态采样,需要更新用于actor更新的dispatch_size |
| 297 | if rl_config.filter_groups_enable: | 303 | if rl_config.filter_groups_enable: |
| 298 | rl_config.actor_update_dispatch_size = \ | 304 | rl_config.actor_update_dispatch_size = \ |
| @@ -116,32 +116,20 @@ def compute_verifier_score(batch, megatron_config, rl_config, tokenizer, ignore_ | |||
| 116 | for score, reward in zip(scores, overlong_reward): | 116 | for score, reward in zip(scores, overlong_reward): |
| 117 | score += reward | 117 | score += reward |
| 118 | 118 | ||
| 119 | - scores = torch.tensor( | ||
| 120 | - scores, | ||
| 121 | - dtype=torch.float64, | ||
| 122 | - device=reward_index.device | ||
| 123 | - ) | ||
| 124 | - | ||
| 125 | - original_scores = torch.tensor( | ||
| 126 | - scores, | ||
| 127 | - dtype=torch.float32, | ||
| 128 | - device=reward_index.device | ||
| 129 | - ).reshape(reward_index.shape) | ||
| 130 | - | ||
| 131 | scores = torch.tensor( | 119 | scores = torch.tensor( |
| 132 | scores, | 120 | scores, |
| 133 | dtype=torch.float32, | 121 | dtype=torch.float32, |
| 134 | device=reward_index.device | 122 | device=reward_index.device |
| 135 | ) | 123 | ) |
| 136 | - | ||
| 137 | scores = scores.reshape(reward_index.shape) | 124 | scores = scores.reshape(reward_index.shape) |
| 125 | + | ||
| 138 | end_time = time.time() | 126 | end_time = time.time() |
| 139 | metrics["timing/rule_reward"] = [round(end_time, 4), round(start_time, 4)] | 127 | metrics["timing/rule_reward"] = [round(end_time, 4), round(start_time, 4)] |
| 140 | metrics["start_time/rule_reward"] = [round(start_time, 4)] | 128 | metrics["start_time/rule_reward"] = [round(start_time, 4)] |
| 141 | metrics["end_time/rule_reward"] = [round(end_time, 4)] | 129 | metrics["end_time/rule_reward"] = [round(end_time, 4)] |
| 142 | 130 | ||
| 143 | 131 | ||
| 144 | - return scores, metrics, original_scores | 132 | + return scores, metrics |
| 145 | 133 | ||
| 146 | 134 | ||
| 147 | def verifier(responses, data, config, **kwargs): | 135 | def verifier(responses, data, config, **kwargs): |
| @@ -1,5 +1,6 @@ | |||
| 1 | # Copyright (c) 2024, HUAWEI CORPORATION. All rights reserved. | 1 | # Copyright (c) 2024, HUAWEI CORPORATION. All rights reserved. |
| 2 | +__all__ = ['RayGRPOTrainer', 'RayDAPOTrainer', 'RayPPOTrainer'] | ||
| 3 | + | ||
| 2 | from .grpo_trainer_hybrid import RayGRPOTrainer | 4 | from .grpo_trainer_hybrid import RayGRPOTrainer |
| 3 | from .dapo_trainer_hybrid import RayDAPOTrainer | 5 | from .dapo_trainer_hybrid import RayDAPOTrainer |
| 4 | - | 6 | +from .ppo_trainer_hybrid import RayPPOTrainer |
| 5 | -__all__ = ['RayGRPOTrainer', 'RayDAPOTrainer'] | ||
| @@ -95,7 +95,7 @@ class RayPPOTrainer(RayBaseTrainer): | |||
| 95 | self.kwargs = kwargs | 95 | self.kwargs = kwargs |
| 96 | self.set_actor_log_prob_skip_flag() | 96 | self.set_actor_log_prob_skip_flag() |
| 97 | self.addition_columns = ['values'] | 97 | self.addition_columns = ['values'] |
| 98 | - self.addition_consumers = ["compute_kl", "critic_train", "critic_compute_values"] | 98 | + self.addition_consumers = ["compute_kl", "critic_train", "critic_compute_values", "ppo_metrics"] |
| 99 | if self.dataset_additional_keys: | 99 | if self.dataset_additional_keys: |
| 100 | self.addition_columns.extend(self.dataset_additional_keys) | 100 | self.addition_columns.extend(self.dataset_additional_keys) |
| 101 | self.transfer_dock_init() | 101 | self.transfer_dock_init() |
| @@ -129,7 +129,7 @@ class RayPPOTrainer(RayBaseTrainer): | |||
| 129 | """ | 129 | """ |
| 130 | The utils loop of PPO | 130 | The utils loop of PPO |
| 131 | """ | 131 | """ |
| 132 | - logger = Loggers('ppo_trainer') | 132 | + logger = Loggers('ppo_trainer_hybrid') |
| 133 | metrics = Metric() | 133 | metrics = Metric() |
| 134 | iteration = self.actor_worker.get_iteration() | 134 | iteration = self.actor_worker.get_iteration() |
| 135 | 135 | ||
| @@ -165,7 +165,7 @@ def dynamic_sampling(num_prompt_in_batch, data_num, n_samples_per_prompt, sampli | |||
| 165 | logger = Loggers('dynamic_sampling') | 165 | logger = Loggers('dynamic_sampling') |
| 166 | experience_consumer_stage = 'dynamic_sampling' | 166 | experience_consumer_stage = 'dynamic_sampling' |
| 167 | experience_columns = ['prompts', 'prompt_length', 'responses', 'labels', 'response_length', | 167 | experience_columns = ['prompts', 'prompt_length', 'responses', 'labels', 'response_length', |
| 168 | - 'input_ids', 'rm_scores', 'metric_for_dapo', 'reward_for_dapo'] | 168 | + 'input_ids', 'rm_scores', 'metric_for_dapo'] |
| 169 | sorted_indexes = get_current_dp_range_indexes(experience_count=data_num, | 169 | sorted_indexes = get_current_dp_range_indexes(experience_count=data_num, |
| 170 | assign_batch_size=data_num) if guarantee_order else None | 170 | assign_batch_size=data_num) if guarantee_order else None |
| 171 | while not ray.get(sampling_transfer_dock.all_consumed.remote(experience_consumer_stage)): | 171 | while not ray.get(sampling_transfer_dock.all_consumed.remote(experience_consumer_stage)): |
| @@ -401,7 +401,6 @@ def compute_dapo_data_metrics( | |||
| 401 | "returns", | 401 | "returns", |
| 402 | "prompt_length", | 402 | "prompt_length", |
| 403 | "response_length", | 403 | "response_length", |
| 404 | - "reward_for_dapo" | ||
| 405 | ] | 404 | ] |
| 406 | pad_token_id = tokenizer.pad if tokenizer.pad is not None else tokenizer.eod | 405 | pad_token_id = tokenizer.pad if tokenizer.pad is not None else tokenizer.eod |
| 407 | sorted_indexes = get_current_dp_range_indexes(experience_count=experience_count, | 406 | sorted_indexes = get_current_dp_range_indexes(experience_count=experience_count, |
| @@ -416,7 +415,6 @@ def compute_dapo_data_metrics( | |||
| 416 | sequence_score = batch["rm_scores"].sum(-1) | 415 | sequence_score = batch["rm_scores"].sum(-1) |
| 417 | prompt_length = batch["prompt_length"] | 416 | prompt_length = batch["prompt_length"] |
| 418 | response_length = batch["response_length"] | 417 | response_length = batch["response_length"] |
| 419 | - reward_for_dapo = batch["reward_for_dapo"] | ||
| 420 | 418 | ||
| 421 | metrics = { | 419 | metrics = { |
| 422 | # score | 420 | # score |
| @@ -431,8 +429,6 @@ def compute_dapo_data_metrics( | |||
| 431 | "prompt_length/mean": torch.mean(prompt_length, dtype=torch.float32).detach().item(), | 429 | "prompt_length/mean": torch.mean(prompt_length, dtype=torch.float32).detach().item(), |
| 432 | "prompt_length/max": torch.max(prompt_length).detach().item(), | 430 | "prompt_length/max": torch.max(prompt_length).detach().item(), |
| 433 | "prompt_length/min": torch.min(prompt_length).detach().item(), | 431 | "prompt_length/min": torch.min(prompt_length).detach().item(), |
| 434 | - | ||
| 435 | - "dapo/reward_for_dapo/mean": torch.mean(reward_for_dapo).detach().item(), | ||
| 436 | } | 432 | } |
| 437 | return metrics | 433 | return metrics |
| 438 | 434 | ||
PL 后面的compute_ppo_data_metrics函数中的experience_consumer_stage应该是ppo_metrics,并且在transfer dock里面增加ppo_metrics这一experience_consumer ![]() ![]() 已添加 ![]() ![]() | |||
| @@ -453,7 +449,7 @@ def compute_ppo_data_metrics( | |||
| 453 | Returns: | 449 | Returns: |
| 454 | Dictionary containing various metric values | 450 | Dictionary containing various metric values |
| 455 | """ | 451 | """ |
| 456 | - experience_consumer_stage = "grpo_metrics" | 452 | + experience_consumer_stage = "ppo_metrics" |
| 457 | experience_columns = [ | 453 | experience_columns = [ |
| 458 | "rm_scores", | 454 | "rm_scores", |
| 459 | "token_level_rewards", | 455 | "token_level_rewards", |
| @@ -288,7 +288,6 @@ class GRPOTransferDock(TransferDock): | |||
| 288 | "advantages", | 288 | "advantages", |
| 289 | "returns", | 289 | "returns", |
| 290 | "metric_for_dapo", | 290 | "metric_for_dapo", |
| 291 | - "reward_for_dapo" | ||
| 292 | ] | 291 | ] |
| 293 | self.experience_consumers = [ | 292 | self.experience_consumers = [ |
| 294 | "trainer", | 293 | "trainer", |
P self.experience_consumers 要增加ppo_metrics ![]() ![]() | |||
| @@ -4,13 +4,13 @@ | |||
| 4 | import os | 4 | import os |
| 5 | import sys | 5 | import sys |
| 6 | import json | 6 | import json |
| 7 | - | ||
| 8 | import time | 7 | import time |
| 9 | import math | 8 | import math |
| 10 | import random | 9 | import random |
| 11 | from functools import wraps | 10 | from functools import wraps |
| 12 | from typing import Dict, List | 11 | from typing import Dict, List |
| 13 | 12 | ||
| 13 | +import ray | ||
| 14 | import omegaconf | 14 | import omegaconf |
| 15 | import numpy as np | 15 | import numpy as np |
| 16 | import torch | 16 | import torch |
| @@ -18,6 +18,11 @@ import torch_npu | |||
| 18 | from torch import Tensor | 18 | from torch import Tensor |
| 19 | 19 | ||
| 20 | 20 | ||
| 21 | +def get_node_nums(): | ||
| 22 | + nodes = ray.nodes() | ||
| 23 | + return len([node for node in nodes if node.get("Alive", False)]) | ||
| 24 | + | ||
| 25 | + | ||
| 21 | def get_current_dp_range_indexes(experience_count, assign_batch_size, current_dp_rank=0): | 26 | def get_current_dp_range_indexes(experience_count, assign_batch_size, current_dp_rank=0): |
| 22 | all_indexes = list(range(assign_batch_size * current_dp_rank, assign_batch_size * (current_dp_rank + 1))) | 27 | all_indexes = list(range(assign_batch_size * current_dp_rank, assign_batch_size * (current_dp_rank + 1))) |
| 23 | return [all_indexes[i:i + experience_count] for i in range(0, len(all_indexes), experience_count)] | 28 | return [all_indexes[i:i + experience_count] for i in range(0, len(all_indexes), experience_count)] |
| @@ -198,18 +198,13 @@ class ActorHybridWorkerBase(BaseWorker): | |||
| 198 | if is_multimodal(): | 198 | if is_multimodal(): |
| 199 | experience_columns.extend(['attention_mask', 'position_ids']) | 199 | experience_columns.extend(['attention_mask', 'position_ids']) |
| 200 | 200 | ||
| 201 | - if self.rl_config.use_integrated_worker: | 201 | + experience_count = self.rl_config.actor_update_dispatch_size |
LP 在参数校验处已经修改 ![]() ![]() experience_count 默认值应该是( self.megatron_config.global_batch_size // self.parallel_state.get_data_parallel_world_size() ) ![]() ![]() | |||
| 202 | + | ||
| 203 | + if self.rl_config.filter_groups_enable: | ||
| 202 | experience_count = ( | 204 | experience_count = ( |
| 203 | - self.megatron_config.global_batch_size // | 205 | + self.rl_config.filter_groups_train_batch_size * self.rl_config.n_samples_per_prompt // |
| 204 | self.parallel_state.get_data_parallel_world_size() | 206 | self.parallel_state.get_data_parallel_world_size() |
| 205 | ) | 207 | ) |
| 206 | - if self.rl_config.filter_groups_enable: | ||
| 207 | - experience_count = ( | ||
| 208 | - self.rl_config.filter_groups_train_batch_size * self.rl_config.n_samples_per_prompt // | ||
| 209 | - self.parallel_state.get_data_parallel_world_size() | ||
| 210 | - ) | ||
| 211 | - else: | ||
| 212 | - experience_count = self.rl_config.actor_update_dispatch_size | ||
| 213 | 208 | ||
| 214 | if skip_actor_log_prob: | 209 | if skip_actor_log_prob: |
| 215 | experience_columns.remove('old_log_prob') | 210 | experience_columns.remove('old_log_prob') |
| @@ -163,14 +163,8 @@ class CriticWorkerBase(BaseWorker): | |||
| 163 | experience_columns = ['responses', 'advantages', 'old_log_prob', 'values', 'returns', | 163 | experience_columns = ['responses', 'advantages', 'old_log_prob', 'values', 'returns', |
| 164 | 'input_ids', 'response_length', 'prompt_length'] | 164 | 'input_ids', 'response_length', 'prompt_length'] |
| 165 | 165 | ||
| 166 | - if self.rl_config.use_integrated_worker: | 166 | + experience_count = self.rl_config.critic_update_dispatch_size |
| 167 | - experience_count = ( | 167 | + |
| 168 | - self.megatron_config.global_batch_size // | ||
| 169 | - self.parallel_state.get_data_parallel_world_size() | ||
| 170 | - ) | ||
| 171 | - else: | ||
| 172 | - experience_count = self.rl_config.critic_update_dispatch_size | ||
| 173 | - | ||
| 174 | learning_rate = None | 168 | learning_rate = None |
| 175 | for param_group in self.optimizer.param_groups: | 169 | for param_group in self.optimizer.param_groups: |
| 176 | learning_rate = param_group['lr'] | 170 | learning_rate = param_group['lr'] |
| @@ -92,13 +92,7 @@ class RewardWorkerBase(BaseWorker): | |||
| 92 | experience_columns = ['input_ids', 'prompt_length', "responses", "response_length", | 92 | experience_columns = ['input_ids', 'prompt_length', "responses", "response_length", |
| 93 | *self.megatron_config.dataset_additional_keys] | 93 | *self.megatron_config.dataset_additional_keys] |
| 94 | 94 | ||
| 95 | - if self.rl_config.use_integrated_worker: | 95 | + experience_count = self.rl_config.reward_dispatch_size |
| 96 | - experience_count = ( | ||
| 97 | - self.megatron_config.global_batch_size // | ||
| 98 | - self.parallel_state.get_data_parallel_world_size() | ||
| 99 | - ) | ||
| 100 | - else: | ||
| 101 | - experience_count = self.rl_config.reward_dispatch_size | ||
| 102 | 96 | ||
| 103 | sorted_indexes = self.get_dp_range_indexes(experience_count, | 97 | sorted_indexes = self.get_dp_range_indexes(experience_count, |
| 104 | use_vllm=False) if self.rl_config.guarantee_order else None | 98 | use_vllm=False) if self.rl_config.guarantee_order else None |
| @@ -63,7 +63,7 @@ class RuleReward(object): | |||
| 63 | use_verifier_mask[::self.n_samples_per_prompt]] for key, value in batch_data.items()} | 63 | use_verifier_mask[::self.n_samples_per_prompt]] for key, value in batch_data.items()} |
| 64 | ignore_token = self.tokenizer.pad if self.tokenizer.pad else self.tokenizer.eod | 64 | ignore_token = self.tokenizer.pad if self.tokenizer.pad else self.tokenizer.eod |
| 65 | 65 | ||
| 66 | - rm_scores, metrics, original_scores = compute_verifier_score( | 66 | + rm_scores, metrics = compute_verifier_score( |
| 67 | batch_data, | 67 | batch_data, |
| 68 | self.megatron_config, | 68 | self.megatron_config, |
| 69 | self.rl_config, | 69 | self.rl_config, |
| @@ -74,7 +74,7 @@ class RuleReward(object): | |||
| 74 | for key, value in metrics.items(): | 74 | for key, value in metrics.items(): |
| 75 | ray.get(self.td.update_metrics.remote(key, value=value, cumulate=True)) | 75 | ray.get(self.td.update_metrics.remote(key, value=value, cumulate=True)) |
| 76 | 76 | ||
| 77 | - output = {"rm_scores": rm_scores, "reward_for_dapo": original_scores} | 77 | + output = {"rm_scores": rm_scores} |
| 78 | if self.rl_config.filter_groups_enable: | 78 | if self.rl_config.filter_groups_enable: |
| 79 | metric = torch.tensor(metrics[self.rl_config.filter_groups_metric], dtype=torch.float32, | 79 | metric = torch.tensor(metrics[self.rl_config.filter_groups_metric], dtype=torch.float32, |
| 80 | device=rm_scores.device) | 80 | device=rm_scores.device) |
| @@ -40,31 +40,23 @@ from mindspeed_rl.workers.integrated_worker import IntegratedWorker | |||
| 40 | from mindspeed_rl.workers.critic_worker import CriticWorker | 40 | from mindspeed_rl.workers.critic_worker import CriticWorker |
| 41 | 41 | ||
| 42 | 42 | ||
| 43 | -def get_rl_resource_by_worker_type(rl_config: RLConfig, worker: Type[BaseWorker]): | 43 | +resource_mapping: Dict[str, Callable[[RLConfig], int]] = { |
| 44 | - if (worker.__ray_actor_class__.__name__ == | 44 | + ActorHybridWorker.__ray_actor_class__.__name__: lambda config: config.actor_resource, |
| 45 | - ActorHybridWorker.__ray_actor_class__.__name__): | 45 | + IntegratedWorker.__ray_actor_class__.__name__: lambda config: config.actor_resource, |
| 46 | - return rl_config.actor_resource | 46 | + RewardWorker.__ray_actor_class__.__name__: lambda config: config.reward_resource, |
| 47 | - elif (worker.__ray_actor_class__.__name__ == | 47 | + ReferenceWorker.__ray_actor_class__.__name__: lambda config: config.reference_resource, |
| 48 | - IntegratedWorker.__ray_actor_class__.__name__): | 48 | + CriticWorker.__ray_actor_class__.__name__: lambda config: config.critic_resource |
| 49 | - return rl_config.actor_resource | 49 | +} |
| 50 | - elif (worker.__ray_actor_class__.__name__ == | 50 | + |
| 51 | - RewardWorker.__ray_actor_class__.__name__): | 51 | + |
| 52 | - return rl_config.reward_resource | 52 | +def get_rl_resource_by_worker_type(rl_config: RLConfig, worker: Type[BaseWorker]) -> Optional[int]: |
| 53 | - elif (worker.__ray_actor_class__.__name__ == | 53 | + actor_class = worker.__ray_actor_class__.__name__ |
| 54 | - ReferenceWorker.__ray_actor_class__.__name__): | 54 | + return resource_mapping.get(actor_class, lambda _: None)(rl_config) |
| 55 | - return rl_config.reference_resource | ||
| 56 | - elif (worker.__ray_actor_class__.__name__ == | ||
| 57 | - CriticWorker.__ray_actor_class__.__name__): | ||
| 58 | - return rl_config.critic_resource | ||
| 59 | - else: | ||
| 60 | - return None | ||
| 61 | 55 | ||
| 62 | 56 | ||
| 63 | def get_npu_deployment(rl_config: RLConfig, worker: Type[BaseWorker]): | 57 | def get_npu_deployment(rl_config: RLConfig, worker: Type[BaseWorker]): |
| 64 | resource = get_rl_resource_by_worker_type(rl_config, worker) | 58 | resource = get_rl_resource_by_worker_type(rl_config, worker) |
| 65 | - if resource is None: | 59 | + return resource.num_npus if resource else 0 |
| 66 | - return 0 | ||
| 67 | - return resource.num_npus | ||
| 68 | 60 | ||
| 69 | 61 | ||
| 70 | def get_device_information(num_npus: int) \ | 62 | def get_device_information(num_npus: int) \ |


没有直接删除,优化了一下校验逻辑