"""
description:
Launch MindSpeed RL Worker.
Expose the RayActorGroup class.
launch remote worker task:
def create_actor_handlers()
runtime_env = {
"env_vars": { "MASTER_ADDR": xxx }
}
return self.worker.options(
scheduling_strategy=PlacementGroupSchedulingStrategy(
placement_group=placement_group, --> control ray resource deployment
placement_group_bundle_index=rank_index --> control ray colocate workers
),
runtime_env=runtime_env --> pass environment variables to remote task
).remote(...) --> launch remote task
"""
from types import ModuleType
from typing import Type, Dict, Callable, Tuple, List, Optional, Union
from dataclasses import dataclass
import acl
import ray
from ray.util.placement_group import PlacementGroup
from ray.util.scheduling_strategies import PlacementGroupSchedulingStrategy
from mindspeed_rl.config_cls.megatron_config import MegatronConfig
from mindspeed_rl.config_cls.rl_config import RLConfig
from mindspeed_rl.config_cls.generate_config import GenerateConfig
from mindspeed_rl.config_cls.mindstudio_config import ProfilerConfig, MsprobeConfig
from mindspeed_rl.utils.tokenizer import BaseTokenizer
from mindspeed_rl.workers.base_worker import BaseWorker
from mindspeed_rl.workers.actor_hybrid_worker import ActorHybridWorker
from mindspeed_rl.workers.reference_worker import ReferenceWorker
from mindspeed_rl.workers.reward_worker import RewardWorker
from mindspeed_rl.workers.integrated_worker import IntegratedWorker
from mindspeed_rl.workers.critic_worker import CriticWorker
from mindspeed_rl.workers.vit_worker import VitWorker
resource_mapping: Dict[str, Callable[[RLConfig], int]] = {
ActorHybridWorker.__ray_actor_class__.__name__: lambda config: config.actor_resource,
IntegratedWorker.__ray_actor_class__.__name__: lambda config: config.actor_resource,
RewardWorker.__ray_actor_class__.__name__: lambda config: config.reward_resource,
ReferenceWorker.__ray_actor_class__.__name__: lambda config: config.reference_resource,
CriticWorker.__ray_actor_class__.__name__: lambda config: config.critic_resource,
VitWorker.__ray_actor_class__.__name__: lambda config: config.vit_resource
}
def get_rl_resource_by_worker_type(rl_config: RLConfig, worker: Type[BaseWorker]) -> Optional[int]:
actor_class = worker.__ray_actor_class__.__name__
return resource_mapping.get(actor_class, lambda _: None)(rl_config)
def get_npu_deployment(rl_config: RLConfig, worker: Type[BaseWorker]):
resource = get_rl_resource_by_worker_type(rl_config, worker)
return resource.num_npus if resource else 0
def get_device_information(num_npus: int) \
-> Tuple[int, int]:
try:
num_devices_per_node, ret = acl.rt.get_device_count()
if ret != 0:
num_devices_per_node = 8
except Exception:
num_devices_per_node = 8
if num_devices_per_node > num_npus:
num_devices_per_node = num_npus
return (num_devices_per_node,
(num_npus + num_devices_per_node - 1) // num_devices_per_node)
def construct_placement_groups(num_npus, num_cpus, num_devices_per_node, num_nodes) \
-> List[PlacementGroup]:
"""
构造ray placement group
构造原则:
1 基于节点连续分配资源
2 共置情形,相同rank索引使用相同placement group索引
3 STRICT_PACK部署策略,强制一个placement group部署在同一节点上
"""
placement_groups = []
for index in range(num_nodes):
if (num_npus % num_devices_per_node != 0 and
index == num_nodes - 1):
bundles = [{"NPU": 1, "CPU": num_cpus}
for _ in range(num_npus % num_devices_per_node)]
else:
bundles = [{"NPU": 1, "CPU": num_cpus}
for _ in range(num_devices_per_node)]
placement_group = ray.util.placement_group(bundles, strategy="STRICT_PACK")
ray.get(placement_group.ready())
placement_groups.append(placement_group)
return placement_groups
def construct_colocate_placement_groups(rl_config) \
-> List[PlacementGroup]:
num_npus = get_npu_deployment(rl_config, ActorHybridWorker)
num_deivces_per_node, num_nodes = get_device_information(num_npus)
return construct_placement_groups(num_npus, rl_config.num_cpus_for_placement_group,
num_deivces_per_node, num_nodes)
@dataclass
class ActorHandlerParams:
placement_group: PlacementGroup
world_size: int
rank_index: int
bundle_index: int
master_addr: str
master_port: int
class RayActorGroup:
def __init__(
self,
worker: Type[BaseWorker],
placement_group: Union[PlacementGroup, List[PlacementGroup]],
megatron_config: MegatronConfig,
rl_config: RLConfig,
model_provider: Callable,
initialize_func: Callable,
profiler_config: Optional[ProfilerConfig] = None,
msprobe_config: Optional[MsprobeConfig] = None,
tokenizer: BaseTokenizer = None,
generate_config: GenerateConfig = None,
resources: Dict[str, float] = None,
num_resources_per_node: int = None,
get_megatron_module: Callable = None,
**kwargs
):
"""
description:
ray actor group, all same work type deploy in one group
parameters:
worker : worker class, such as ReferenceWorker
placement_group : ray placement group
megatron_config : megatron config data
rl_config : reinforcement learning config data
model_provider : model provider function
initialize_func : model initialization function
tokenizer : tokenizer
generate_config : vllm config data
resources : user defined ray resource
num_resources_per_node : number of resources per node
kwargs : keyword arguments
"""
self.worker = worker
self.placement_group = placement_group
self.megatron_config = megatron_config
self.rl_config = rl_config
self.generate_config = generate_config
self.profiler_config = profiler_config
self.msprobe_config = msprobe_config
self.model_provider = model_provider
self.initialize_func = initialize_func
self.tokenizer = tokenizer
self.get_megatron_module = get_megatron_module
self.kwargs = kwargs
self.num_npus = get_npu_deployment(rl_config, worker)
self.resources = resources
self.num_resources_per_node = num_resources_per_node
self.actor_handlers = []
self.temp_actor_ref_objs = []
self.num_devices_per_node, self.num_nodes = (
get_device_information(self.num_npus))
self.initialize_actor_handlers(placement_group)
def initialize_actor_handlers(self, placement_group):
world_size = self.num_npus
placement_group = self.get_placement_group(placement_group=placement_group)
self.placement_group = placement_group
master_actor = self.build_master_actor(placement_group, world_size)
if world_size > 1:
self.build_worker_actor(master_actor, placement_group, world_size)
def get_placement_group(self, placement_group: PlacementGroup = None) \
-> Union[PlacementGroup, List[PlacementGroup]]:
if placement_group is not None:
return placement_group
return construct_placement_groups(self.num_npus, self.rl_config.num_cpus_for_placement_group,
self.num_devices_per_node, self.num_nodes)
def create_actor_handlers(self, param: ActorHandlerParams) \
-> ray.actor.ActorHandle:
runtime_env = {
"env_vars": {
"MASTER_ADDR": param.master_addr if param.master_addr else "localhost",
"MASTER_PORT": str(param.master_port) if param.master_port else "",
"WORLD_SIZE": str(param.world_size),
"RANK": str(param.rank_index),
}
}
return self.worker.options(
scheduling_strategy=PlacementGroupSchedulingStrategy(
placement_group=param.placement_group,
placement_group_bundle_index=param.bundle_index
),
runtime_env=runtime_env
).remote(
self.megatron_config,
self.rl_config,
self.generate_config,
model_provider=self.model_provider,
get_megatron_module=self.get_megatron_module,
initialize_func=self.initialize_func,
profiler_config=self.profiler_config,
msprobe_config=self.msprobe_config,
tokenizer=self.tokenizer,
**self.kwargs
)
def build_master_actor(self, placement_group, world_size) -> ray.actor.ActorHandle:
actor_handle = self.create_actor_handlers(
ActorHandlerParams(placement_group[0], world_size, 0, 0, None, None))
self.actor_handlers.append(actor_handle)
return actor_handle
def build_worker_actor(self, master_handler, placement_group, world_size) -> None:
master_addr, master_port = ray.get(master_handler.get_master_addr_port.remote())
for rank in range(1, self.num_devices_per_node):
self.actor_handlers.append(self.create_actor_handlers(
ActorHandlerParams(placement_group[0], world_size, rank,
rank, master_addr, master_port)))
rank_index = self.num_devices_per_node - 1
for node_index in range(1, self.num_nodes):
for bundle_index in range(0, self.num_devices_per_node):
rank_index += 1
self.actor_handlers.append(self.create_actor_handlers(
ActorHandlerParams(placement_group[node_index], world_size, rank_index,
bundle_index, master_addr, master_port)))
def execute_async_command(self, method_name: str, *args, **kwargs):
ray_objs = []
for handler in self.actor_handlers:
if hasattr(handler, method_name) and callable(getattr(handler, method_name)):
ray_objs.append(getattr(handler, method_name, None).remote(*args, **kwargs))
return ray_objs
def execute_sync_command(self, method_name: str, *args, **kwargs):
return ray.get(self.execute_async_command(method_name, *args, **kwargs))
def async_init_transfer_dock(self, transfer_dock, mm_transfer_dock=None):
for actor in self.actor_handlers:
self.temp_actor_ref_objs.append(actor.init_transfer_dock.remote(transfer_dock, mm_transfer_dock))
def sync_init_transfer_dock(self, transfer_dock, mm_transfer_dock=None, sampling_transfer_dock=None, mm_sampling_transfer_dock=None):
for actor in self.actor_handlers:
ray.get(actor.init_transfer_dock.remote(transfer_dock, mm_transfer_dock, sampling_transfer_dock, mm_sampling_transfer_dock))
def enter_infer_mode(self, blocking=False):
for actor in self.actor_handlers:
self.temp_actor_ref_objs.append(actor.enter_infer_mode.remote())
if blocking:
ray.get(self.temp_actor_ref_objs)
def exit_infer_mode(self, blocking=False):
for actor in self.actor_handlers:
self.temp_actor_ref_objs.append(actor.exit_infer_mode.remote())
if blocking:
ray.get(self.temp_actor_ref_objs)
def wait_all_ref_objs_run_over(self):
ray.get(self.temp_actor_ref_objs)
self.temp_actor_ref_objs.clear()
def get_iteration(self):
return ray.get(self.actor_handlers[0].get_iteration.remote())
def generate_sequences(self, blocking=False):
for actor in self.actor_handlers:
self.temp_actor_ref_objs.append(actor.generate_sequences.remote())
if blocking:
ray.get(self.temp_actor_ref_objs)
def compute_image_embeds(self, blocking=False):
for actor in self.actor_handlers:
self.temp_actor_ref_objs.append(actor.compute_image_embeds.remote())
if blocking:
ray.get(self.temp_actor_ref_objs)
def compute_log_prob(self, blocking=False):
for actor in self.actor_handlers:
self.temp_actor_ref_objs.append(actor.compute_log_prob.remote())
if blocking:
ray.get(self.temp_actor_ref_objs)
def compute_ref_log_prob(self, blocking=False):
for actor in self.actor_handlers:
self.temp_actor_ref_objs.append(actor.compute_ref_log_prob.remote())
if blocking:
ray.get(self.temp_actor_ref_objs)
def compute_rm_score(self, blocking=False):
for actor in self.actor_handlers:
self.temp_actor_ref_objs.append(actor.compute_rm_score.remote())
if blocking:
ray.get(self.temp_actor_ref_objs)
def update(self, kl_ctrl, skip_actor_log_prob):
actor_train_objs = []
for actor in self.actor_handlers:
actor_train_objs.append(actor.update.remote(kl_ctrl, skip_actor_log_prob))
return ray.get(actor_train_objs)
def update_actor(self, skip_actor_log_prob, kl_ctrl=None):
actor_train_objs = []
for actor in self.actor_handlers:
actor_train_objs.append(actor.update.remote(kl_ctrl, skip_actor_log_prob))
return ray.get(actor_train_objs)
def update_critic(self, blocking=False, kl_ctrl=None):
actor_train_objs = []
for actor in self.actor_handlers:
actor_train_objs.append(actor.update.remote(kl_ctrl))
if blocking:
ray.get(actor_train_objs)
def compute_values(self, blocking=False):
for actor in self.actor_handlers:
self.temp_actor_ref_objs.append(actor.compute_values.remote())
if blocking:
ray.get(self.temp_actor_ref_objs)
def save_checkpoint(self, iteration):
actor_train_objs = []
for actor in self.actor_handlers:
actor_train_objs.append(actor.save_ckpt.remote(iteration))
return ray.get(actor_train_objs)
def initialize(self):
for actor in self.actor_handlers:
self.temp_actor_ref_objs.append(actor.initialize.remote())
ray.get(self.temp_actor_ref_objs)
return self
def get_consumed_train_samples(self):
return ray.get(self.actor_handlers[0].get_consumed_train_samples.remote())