mirror of
https://github.com/wassname/ray.git
synced 2026-08-14 12:40:23 +08:00
[docker] Detect CPUs in container correctly (#10507)
Co-authored-by: simon-mo <simon.mo@hey.com> Co-authored-by: Richard Liaw <rliaw@berkeley.edu> Co-authored-by: Alex Wu <itswu.alex@gmail.com>
This commit is contained in:
co-authored by
simon-mo
Richard Liaw
Alex Wu
parent
660aee6311
commit
5bc2ba38fd
@@ -109,7 +109,10 @@ class Training(TrainingOperator):
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
ray.init(address=None if args.local else "auto")
|
||||
if args.local:
|
||||
ray.init(num_cpus=2)
|
||||
else:
|
||||
ray.init(address="auto")
|
||||
num_workers = 2 if args.local else int(ray.cluster_resources().get(device))
|
||||
from ray.util.sgd.torch.examples.train_example import LinearDataset
|
||||
|
||||
|
||||
@@ -289,7 +289,10 @@ if __name__ == "__main__":
|
||||
default=False,
|
||||
help="Enables GPU training")
|
||||
args = parser.parse_args()
|
||||
ray.init(address=args.address)
|
||||
if args.smoke_test:
|
||||
ray.init(num_cpus=2)
|
||||
else:
|
||||
ray.init(address=args.address)
|
||||
|
||||
trainer = train_example(
|
||||
num_workers=args.num_workers,
|
||||
|
||||
@@ -128,8 +128,10 @@ def main():
|
||||
setup_default_logging()
|
||||
|
||||
args, args_text = parse_args()
|
||||
|
||||
ray.init(address=args.ray_address)
|
||||
if args.smoke_test:
|
||||
ray.init(num_cpus=int(args.ray_num_workers))
|
||||
else:
|
||||
ray.init(address=args.ray_address)
|
||||
|
||||
CustomTrainingOperator = TrainingOperator.from_creators(
|
||||
model_creator=model_creator,
|
||||
|
||||
@@ -115,10 +115,17 @@ if __name__ == "__main__":
|
||||
help="Enables GPU training")
|
||||
parser.add_argument(
|
||||
"--tune", action="store_true", default=False, help="Tune training")
|
||||
parser.add_argument(
|
||||
"--smoke-test",
|
||||
action="store_true",
|
||||
default=False,
|
||||
help="Finish quickly for testing.")
|
||||
|
||||
args, _ = parser.parse_known_args()
|
||||
|
||||
import ray
|
||||
|
||||
ray.init(address=args.address)
|
||||
if args.smoke_test:
|
||||
ray.init(num_cpus=2)
|
||||
else:
|
||||
ray.init(address=args.address)
|
||||
train_example(num_workers=args.num_workers, use_gpu=args.use_gpu)
|
||||
|
||||
@@ -59,6 +59,8 @@ def tune_example(operator_cls, num_workers=1, use_gpu=False):
|
||||
if __name__ == "__main__":
|
||||
import argparse
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument(
|
||||
"--smoke-test", action="store_true", help="Finish quickly for testing")
|
||||
parser.add_argument(
|
||||
"--address",
|
||||
type=str,
|
||||
@@ -77,7 +79,10 @@ if __name__ == "__main__":
|
||||
|
||||
args, _ = parser.parse_known_args()
|
||||
|
||||
ray.init(address=args.address)
|
||||
if args.smoke_test:
|
||||
ray.init(num_cpus=2)
|
||||
else:
|
||||
ray.init(address=args.address)
|
||||
CustomTrainingOperator = TrainingOperator.from_creators(
|
||||
model_creator=model_creator, optimizer_creator=optimizer_creator,
|
||||
data_creator=data_creator, loss_creator=nn.MSELoss)
|
||||
|
||||
Reference in New Issue
Block a user