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