|
17 | 17 | collect_all_to_all, |
18 | 18 | dispatch_one_to_all, |
19 | 19 | ) |
| 20 | +from roll.platforms import current_platform |
20 | 21 | from roll.utils.constants import RAY_NAMESPACE |
21 | 22 | from roll.distributed.scheduler.resource_manager import ResourceManager |
22 | 23 | from roll.utils.import_utils import safe_import_class |
@@ -122,14 +123,25 @@ def _create_workers(self): |
122 | 123 |
|
123 | 124 | runtime_env = RuntimeEnv(env_vars=env_vars) |
124 | 125 | self.worker_config.resource_placement_groups = pgs |
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) |
| 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) |
133 | 145 | self.workers.append(worker) |
134 | 146 | if rank == 0: |
135 | 147 | self.master_addr, self.master_port = ray.get(worker.get_master_addr_and_port.remote()) |
|
0 commit comments