Skip to content

Commit bf06136

Browse files
FightingZhennoemotiovon
authored andcommitted
(refactor): extract gpu & ray utils.
1 parent ce82cc7 commit bf06136

1 file changed

Lines changed: 8 additions & 20 deletions

File tree

roll/distributed/executor/cluster.py

Lines changed: 8 additions & 20 deletions
Original file line numberDiff line numberDiff line change
@@ -17,7 +17,6 @@
1717
collect_all_to_all,
1818
dispatch_one_to_all,
1919
)
20-
from roll.platforms import current_platform
2120
from roll.utils.constants import RAY_NAMESPACE
2221
from roll.distributed.scheduler.resource_manager import ResourceManager
2322
from roll.utils.import_utils import safe_import_class
@@ -123,25 +122,14 @@ def _create_workers(self):
123122

124123
runtime_env = RuntimeEnv(env_vars=env_vars)
125124
self.worker_config.resource_placement_groups = pgs
126-
if current_platform.ray_device_key == "GPU":
127-
worker = self.worker_cls.options(
128-
scheduling_strategy=PlacementGroupSchedulingStrategy(placement_group=deploy_pg["placement_group"]),
129-
name=worker_name,
130-
namespace=RAY_NAMESPACE,
131-
runtime_env=runtime_env,
132-
num_cpus=0.01,
133-
num_gpus=0.01 if self.worker_config.device_mapping else 0,
134-
).remote(worker_config=self.worker_config)
135-
else:
136-
worker = self.worker_cls.options(
137-
scheduling_strategy=PlacementGroupSchedulingStrategy(placement_group=deploy_pg["placement_group"]),
138-
name=worker_name,
139-
namespace=RAY_NAMESPACE,
140-
runtime_env=runtime_env,
141-
num_cpus=0.01,
142-
num_gpus=0,
143-
resources={current_platform.ray_device_key: 0.01 if self.worker_config.device_mapping else 0},
144-
).remote(worker_config=self.worker_config)
125+
worker = self.worker_cls.options(
126+
scheduling_strategy=PlacementGroupSchedulingStrategy(placement_group=deploy_pg["placement_group"]),
127+
name=worker_name,
128+
namespace=RAY_NAMESPACE,
129+
runtime_env=runtime_env,
130+
num_cpus=0.01,
131+
num_gpus=0.01 if self.worker_config.device_mapping else 0,
132+
).remote(worker_config=self.worker_config)
145133
self.workers.append(worker)
146134
if rank == 0:
147135
self.master_addr, self.master_port = ray.get(worker.get_master_addr_and_port.remote())

0 commit comments

Comments
 (0)