已合并
worker初始化优化&&DAPO去除冗余代码&&其他一些小优化 #442
LookAround0301创建于 2025年6月30日
worker初始化优化&&DAPO去除冗余代码&&其他一些小优化 #442
已合并
LookAround0301创建于 2025年6月30日
refs/pull/442/head合入到master
17 个文件变更+53-96
@@ -17,7 +17,7 @@ from ray.util import placement_group
17 17 
18from mindspeed_rl.utils import seed_all18from mindspeed_rl.utils import seed_all
19from mindspeed_rl.utils import get_tokenizer19from mindspeed_rl.utils import get_tokenizer
20-from mindspeed_rl.utils.utils import MsProbe20+from mindspeed_rl.utils.utils import MsProbe, get_node_nums
21from mindspeed_rl.utils.loggers import Loggers21from mindspeed_rl.utils.loggers import Loggers
22from mindspeed_rl.utils.utils import parse_args_from_config22from mindspeed_rl.utils.utils import parse_args_from_config
23from mindspeed_rl.config_cls.validate_config import validate_rl_args23from 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
19from mindspeed_rl.utils import get_tokenizer19from mindspeed_rl.utils import get_tokenizer
20from mindspeed_rl.datasets.build_dataset import build_train_valid_test_datasets20from mindspeed_rl.datasets.build_dataset import build_train_valid_test_datasets
21from mindspeed_rl.utils import seed_all21from mindspeed_rl.utils import seed_all
22-from mindspeed_rl.utils.utils import MsProbe22+from mindspeed_rl.utils.utils import MsProbe, get_node_nums
23from mindspeed_rl.utils.loggers import Loggers23from mindspeed_rl.utils.loggers import Loggers
24from mindspeed_rl.utils.utils import parse_args_from_config24from mindspeed_rl.utils.utils import parse_args_from_config
25from mindspeed_rl.config_cls.megatron_config import MegatronConfig25from 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
19from mindspeed_rl.utils import get_tokenizer19from mindspeed_rl.utils import get_tokenizer
20from mindspeed_rl.datasets.build_dataset import build_train_valid_test_datasets20from mindspeed_rl.datasets.build_dataset import build_train_valid_test_datasets
21from mindspeed_rl.utils import seed_all21from mindspeed_rl.utils import seed_all
22-from mindspeed_rl.utils.utils import MsProbe22+from mindspeed_rl.utils.utils import MsProbe, get_node_nums
23from mindspeed_rl.utils.loggers import Loggers23from mindspeed_rl.utils.loggers import Loggers
24from mindspeed_rl.utils.utils import parse_args_from_config24from mindspeed_rl.utils.utils import parse_args_from_config
25from mindspeed_rl.config_cls.megatron_config import MegatronConfig25from mindspeed_rl.config_cls.megatron_config import MegatronConfig
@@ -29,7 +29,7 @@ from mindspeed_rl.config_cls.mindstudio_config import ProfilerConfig, MsprobeCon
29from mindspeed_rl.datasets.prompt_dataset import PromptDataset29from mindspeed_rl.datasets.prompt_dataset import PromptDataset
30from mindspeed_rl.datasets.dataloader import PromptDataLoader30from mindspeed_rl.datasets.dataloader import PromptDataLoader
31from mindspeed_rl.workers.rule_reward import RuleReward31from mindspeed_rl.workers.rule_reward import RuleReward
32-from mindspeed_rl.trainer.ppo_trainer import RayPPOTrainer32+from mindspeed_rl.trainer.ppo_trainer_hybrid import RayPPOTrainer
33from mindspeed_rl.workers.scheduler.launcher import RayActorGroup33from mindspeed_rl.workers.scheduler.launcher import RayActorGroup
34from mindspeed_rl.workers.actor_hybrid_worker import ActorHybridWorker34from mindspeed_rl.workers.actor_hybrid_worker import ActorHybridWorker
35from mindspeed_rl.workers.reference_woker import ReferenceWorker35from 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_prompt148 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(
Rconfigs/ppo_qwen25_32b_integrated.yamlconfigs/ppo_qwen25_32b_A3.yaml+0-0
文件重命名但无更改。
Rconfigs/ppo_qwen25_7b_integrated.yamlconfigs/ppo_qwen25_7b_A2.yaml+0-1
@@ -34,7 +34,6 @@ megatron_training:
34 no_shuffle: true34 no_shuffle: true
35 full_shuffle_instruction_dataset: false35 full_shuffle_instruction_dataset: false
36 36 
37- 
38actor_config:37actor_config:
39 model: qwen25_7b38 model: qwen25_7b
40 micro_batch_size: 139 micro_batch_size: 1
@@ -3,6 +3,7 @@ import os
3from mindspeed_rl.config_cls.rl_config import RLConfig3from mindspeed_rl.config_cls.rl_config import RLConfig
4from mindspeed_rl.config_cls.megatron_config import MegatronConfig4from mindspeed_rl.config_cls.megatron_config import MegatronConfig
5from mindspeed_rl.config_cls.generate_config import GenerateConfig5from mindspeed_rl.config_cls.generate_config import GenerateConfig
6+from mindspeed_rl.utils.utils import get_node_nums
6 7 
7 8 
8def validate_rl_args(9def 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
LLookAround03012025年7月1日

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

likedislike
Ppengnuoheng2025年6月30日

这段校验可以删掉,后面会有校验经验计数与全局批次关系

likedislike
232 238 
233 # 若开启dapo动态采样,需要更新用于计算adv和logp的dispatch_size239 # 若开启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 or298+ (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_size302 # 若开启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 += reward117 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.device122 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_scores132+ return scores, metrics
145 133 
146 134 
147def verifier(responses, data, config, **kwargs):135def 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+ 
2from .grpo_trainer_hybrid import RayGRPOTrainer4from .grpo_trainer_hybrid import RayGRPOTrainer
3from .dapo_trainer_hybrid import RayDAPOTrainer5from .dapo_trainer_hybrid import RayDAPOTrainer
4- 6+from .ppo_trainer_hybrid import RayPPOTrainer
5-__all__ = ['RayGRPOTrainer', 'RayDAPOTrainer']
Rmindspeed_rl/trainer/ppo_trainer.pymindspeed_rl/trainer/ppo_trainer_hybrid.py+2-2
@@ -95,7 +95,7 @@ class RayPPOTrainer(RayBaseTrainer):
95 self.kwargs = kwargs95 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 PPO130 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 None170 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.eod405 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 # score420 # 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 metrics433 return metrics
438 434 
PL
Ppengnuoheng2025年7月1日

后面的compute_ppo_data_metrics函数中的experience_consumer_stage应该是ppo_metrics,并且在transfer dock里面增加ppo_metrics这一experience_consumer

likedislike
LLookAround03012025年7月1日

已添加

likedislike
@@ -453,7 +449,7 @@ def compute_ppo_data_metrics(
453 Returns:449 Returns:
454 Dictionary containing various metric values450 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
Ppengnuoheng2025年7月1日

self.experience_consumers 要增加ppo_metrics

likedislike
@@ -4,13 +4,13 @@
4import os4import os
5import sys5import sys
6import json6import json
7- 
8import time7import time
9import math8import math
10import random9import random
11from functools import wraps10from functools import wraps
12from typing import Dict, List11from typing import Dict, List
13 12 
13+import ray
14import omegaconf14import omegaconf
15import numpy as np15import numpy as np
16import torch16import torch
@@ -18,6 +18,11 @@ import torch_npu
18from torch import Tensor18from 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+ 
21def get_current_dp_range_indexes(experience_count, assign_batch_size, current_dp_rank=0):26def 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
LLookAround03012025年7月1日

在参数校验处已经修改

likedislike
Ppengnuoheng2025年7月1日

experience_count 默认值应该是( self.megatron_config.global_batch_size // self.parallel_state.get_data_parallel_world_size() )

likedislike
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 = None168 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 None98 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.eod64 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
40from mindspeed_rl.workers.critic_worker import CriticWorker40from 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_resource46+ 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_resource49+}
50- elif (worker.__ray_actor_class__.__name__ ==50+ 
51- RewardWorker.__ray_actor_class__.__name__):51+ 
52- return rl_config.reward_resource52+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 
63def get_npu_deployment(rl_config: RLConfig, worker: Type[BaseWorker]):57def 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 
70def get_device_information(num_npus: int) \62def get_device_information(num_npus: int) \